diff --git a/tests/pytorch/test_torch_compile.py b/tests/pytorch/test_torch_compile.py index eae6f0a8a2..19eb4e245e 100644 --- a/tests/pytorch/test_torch_compile.py +++ b/tests/pytorch/test_torch_compile.py @@ -2720,3 +2720,320 @@ def test_te_ops_forward_kwargs_compile(): for gain in (3.0, 5.0): _check_ops(compiled, model, x, dy, {"gain": gain}) _assert_custom_ops(graphs[-1:], "_affineop", present=False) + + +@pytest.mark.skipif(not _opaque_available, reason="torch opaque object API not available") +@pytest.mark.parametrize("compile_mode", ["default", "reduce-overhead"]) +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16, torch.float32]) +@pytest.mark.parametrize("equal_splits", [False, True]) +def test_te_split_parameters_compile(compile_mode, dtype, equal_splits): + """Every split remains an autograd input across storage and parameter changes.""" + import copy + + torch._dynamo.reset() + counters.clear() + splits = ("q", "k", "v") if equal_splits else dict(q=64, k=32, v=32) + model = te.Linear(64, 192 if equal_splits else 128, parameters_split=splits, params_dtype=dtype) + model.k_weight.requires_grad_(False) + model.v_bias.requires_grad_(False) + reference = copy.deepcopy(model) + compiled = torch.compile(model, fullgraph=True, mode=compile_mode) + optimizer = torch.optim.SGD(model.parameters(), lr=0.01) + ref_optimizer = torch.optim.SGD(reference.parameters(), lr=0.01) + for mutation in ("none", "optimizer", "data", "parameter", "assign", "dtype"): + for current in (model, reference): + if mutation == "data": + current.k_weight.data = current.k_weight.detach().clone() + 0.01 + current.k_bias.data = current.k_bias.detach().clone() + 0.01 + elif mutation == "parameter": + current.q_weight = torch.nn.Parameter(current.q_weight.detach().clone() + 0.01) + elif mutation == "assign": + state = {name: value.clone() for name, value in current.state_dict().items()} + current.load_state_dict(state, assign=True) + elif mutation == "dtype": + current.to(torch.float64).to(dtype) + if mutation == "optimizer": + optimizer.step() + ref_optimizer.step() + for _ in range(3): + torch.compiler.cudagraph_mark_step_begin() + inp = torch.randn(16, 64, device="cuda", dtype=dtype, requires_grad=True) + ref_inp = inp.detach().clone().requires_grad_() + model.zero_grad(set_to_none=True) + reference.zero_grad(set_to_none=True) + out = compiled(inp) + grad = torch.randn_like(out) + out.backward(grad) + actual = [out.detach().clone(), inp.grad.clone()] + actual.extend(None if p.grad is None else p.grad.clone() for p in model.parameters()) + ref_out = reference(ref_inp) + ref_out.backward(grad) + expected = [ref_out, ref_inp.grad, *(p.grad for p in reference.parameters())] + for result, target in zip(actual, expected): + if target is None: + assert result is None + else: + torch.testing.assert_close(result, target, **dtype_tols(dtype)) + if compile_mode == "reduce-overhead": + assert not counters["inductor"]["cudagraph_skips"] + + +@pytest.mark.skipif(not _opaque_available, reason="torch opaque object API not available") +@pytest.mark.parametrize("fp8_recipe", [None, *_all_recipes], ids=recipe_id) +@pytest.mark.parametrize("bias_mode", ["fused", "none", "returned"]) +def test_te_split_parameters_recipes(fp8_recipe, bias_mode): + autocast_recipe = fp8_recipe or recipe.Float8CurrentScaling() + options = dict(bias=bias_mode != "none", return_bias=bias_mode == "returned") + model = te.Linear( + 128, + 256, + parameters_split=dict(q=128, k=64, v=64), + params_dtype=torch.bfloat16, + **options, + ) + + def fn(inp): + with te.autocast(enabled=fp8_recipe is not None, recipe=autocast_recipe): + out = model(inp) + if bias_mode == "returned": + out, bias = out + out = out + bias + return out + + torch._dynamo.reset() + compiled = torch.compile(fn, fullgraph=True) + inp = torch.randn(128, 128, device="cuda", dtype=torch.bfloat16, requires_grad=True) + grad = torch.randn(128, 256, device="cuda", dtype=torch.bfloat16) + expected = fn(inp) + expected.backward(grad) + expected_grads = [t.grad.clone() for t in (inp, *model.parameters()) if t.requires_grad] + inp.grad = None + model.zero_grad(set_to_none=True) + actual = compiled(inp) + actual.backward(grad) + actual_grads = [t.grad for t in (inp, *model.parameters()) if t.requires_grad] + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + torch.testing.assert_close(actual_grads, expected_grads, rtol=0, atol=0) + + +@pytest.mark.skipif(not _opaque_available, reason="torch opaque object API not available") +@pytest.mark.parametrize("training", [False, True]) +def test_te_split_parameters_no_cat(training): + """Adjacent parameter storage does not launch a concatenation in either pass.""" + model = te.Linear(64, 128, parameters_split=dict(q=64, k=32, v=32), params_dtype=torch.bfloat16) + compiled = torch.compile(model, fullgraph=True) + inp = torch.randn(16, 64, device="cuda", dtype=torch.bfloat16, requires_grad=training) + + def run(): + with torch.set_grad_enabled(training): + out = compiled(inp) + if training: + out.sum().backward() + model.zero_grad(set_to_none=True) + inp.grad = None + + for _ in range(3): + run() + with torch.profiler.profile(activities=[torch.profiler.ProfilerActivity.CPU]) as prof: + run() + assert "aten::cat" not in {event.key for event in prof.key_averages()} + model.k_weight.data = model.k_weight.detach().clone() + with torch.profiler.profile(activities=[torch.profiler.ProfilerActivity.CPU]) as prof: + run() + assert "aten::cat" in {event.key for event in prof.key_averages()} + + +@pytest.mark.skipif(not _opaque_available, reason="torch opaque object API not available") +def test_te_split_parameters_saved_versions(): + model = te.Linear(64, 128, parameters_split=dict(q=64, k=32, v=32), params_dtype=torch.bfloat16) + compiled = torch.compile(model, fullgraph=True, backend="aot_eager") + inp = torch.randn(16, 64, device="cuda", dtype=torch.bfloat16, requires_grad=True) + out = compiled(inp) + with torch.no_grad(): + model.v_weight.add_(0.01) + with pytest.raises(RuntimeError, match="modified by an inplace operation"): + out.sum().backward() + + +@pytest.mark.skipif(not _opaque_available, reason="torch opaque object API not available") +def test_te_split_parameters_saved_hooks(): + import copy + + model = te.Linear(64, 128, parameters_split=dict(q=64, k=32, v=32), params_dtype=torch.bfloat16) + reference = copy.deepcopy(model) + compiled = torch.compile(model, fullgraph=True, backend="aot_eager") + inp = torch.randn(16, 64, device="cuda", dtype=torch.bfloat16, requires_grad=True) + ref_inp = inp.detach().clone().requires_grad_() + with torch.autograd.graph.saved_tensors_hooks(lambda t: t.clone(), lambda t: t): + out = compiled(inp) + out.sum().backward() + expected = reference(ref_inp) + expected.sum().backward() + torch.testing.assert_close(out, expected, rtol=0, atol=0) + torch.testing.assert_close(inp.grad, ref_inp.grad, rtol=0, atol=0) + for actual, target in zip(model.parameters(), reference.parameters()): + torch.testing.assert_close(actual.grad, target.grad, rtol=0, atol=0) + + +def test_te_split_parameters_fallback_check(monkeypatch): + model = te.Linear(64, 128, parameters_split=dict(q=64, k=32, v=32)) + inp = torch.empty(16, 64, device="cuda") + + def unexpected_materialization(*args, **kwargs): + raise AssertionError("Fallback eligibility must not concatenate parameters") + + monkeypatch.setattr(model, "_get_weight_and_bias_tensors", unexpected_materialization) + assert model._compile_eager_fallback_reason(inp, None, False, False, True, False) is None + assert ( + model._compile_eager_fallback_reason(inp, None, True, False, True, False) + == "differentiable fp8_output=True" + ) + + +@pytest.mark.parametrize("defer", [False, True]) +def test_concat_input_single_tensor(defer): + from transformer_engine.pytorch.module._common import concat_input + + tensor = torch.nn.Parameter(torch.randn(4, 8, device="cuda")) + assert concat_input([tensor], defer=defer) is tensor + + +@pytest.mark.parametrize("disjoint", [False, True]) +def test_concat_input_eager_gradients(disjoint): + from transformer_engine.pytorch.module._common import concat_input + + parts = [torch.nn.Parameter(t) for t in torch.randn(8, 8, device="cuda").split((3, 5))] + if disjoint: + parts[1] = torch.nn.Parameter(parts[1].detach().clone()) + result = concat_input(parts) + expected = torch.cat(parts) + grad = torch.randn_like(result) + torch.testing.assert_close(result, expected, rtol=0, atol=0) + torch.testing.assert_close( + torch.autograd.grad(result, parts, grad), + torch.autograd.grad(expected, parts, grad), + rtol=0, + atol=0, + ) + + +@pytest.mark.parametrize( + "layout", ["adjacent", "offset", "disjoint", "strided", "singleton", "negative", "conjugate"] +) +def test_parameter_parts_storage(layout): + from transformer_engine.pytorch.dynamo.parameter_parts import ParameterParts + + storage = torch.arange(128, device="cuda").view(16, 8) + if layout == "offset": + storage = storage[2:10] + elif layout == "strided": + storage = storage[:, ::2] + elif layout == "singleton": + storage = storage[:, :1] + elif layout == "negative": + storage = torch._neg_view(storage) + elif layout == "conjugate": + storage = torch.complex(storage.float(), storage.float()).conj() + parts = list(storage.split([1, storage.shape[0] - 3, 2])) + if layout == "disjoint": + parts[1] = parts[1].clone() + result = ParameterParts(parts).materialize() + torch.testing.assert_close(result, torch.cat(parts), rtol=0, atol=0) + if layout in ("adjacent", "offset"): + assert result.data_ptr() == parts[0].data_ptr() + else: + assert result.untyped_storage().data_ptr() != parts[0].untyped_storage().data_ptr() + + +@pytest.mark.parametrize("fake", [False, True], ids=["eager", "fake"]) +@pytest.mark.parametrize("mixed_dtype", [False, True]) +def test_parameter_parts_spec(fake, mixed_dtype): + from transformer_engine.pytorch.dynamo.parameter_parts import ParameterParts + + with FakeTensorMode() if fake else contextlib.nullcontext(): + parts = [ + torch.empty(3, 8, device="cuda", dtype=torch.bfloat16), + torch.empty( + 5, + 8, + device="cuda", + dtype=torch.float32 if mixed_dtype else torch.bfloat16, + requires_grad=True, + ), + ] + expected = torch.cat(parts) + spec = ParameterParts(parts).to_spec() + assert spec.shape == tuple(expected.shape) + assert spec.dtype == expected.dtype + assert spec.device == expected.device + assert spec.requires_grad == expected.requires_grad + assert not spec.is_quantized + + +@pytest.mark.parametrize("internal", [False, True]) +def test_parameter_parts_rejects_quantized_parts(internal): + from transformer_engine.pytorch.dynamo.parameter_parts import ParameterParts + from transformer_engine.pytorch.module._common import concat_input + + quantizer = _current_scaling() + quantizer.internal = internal + tensor = TensorSpec( + shape=(4, 8), dtype=torch.bfloat16, quantizer=quantizer, device=torch.device("cpu") + ).create_tensor() + with pytest.raises(TypeError, match="non-quantized tensors"): + ParameterParts([tensor, tensor]) + assert concat_input([tensor], defer=True) is tensor + + +@pytest.mark.skipif(not _opaque_available, reason="torch opaque object API not available") +@pytest.mark.parametrize("grad_target", ["input", "weight", "bias"]) +def test_te_split_parameters_grad_targets(grad_target): + model = te.Linear(64, 128, parameters_split=dict(q=64, k=32, v=32), params_dtype=torch.bfloat16) + for parameter in model.parameters(): + parameter.requires_grad_(False) + if grad_target != "input": + getattr(model, f"v_{grad_target}").requires_grad_(True) + inp = torch.randn(16, 64, device="cuda", dtype=torch.bfloat16) + inp.requires_grad_(grad_target == "input") + compiled = torch.compile(model, fullgraph=True) + out = compiled(inp) + grad = torch.randn_like(out) + out.backward(grad) + if grad_target == "bias": + torch.testing.assert_close(model.v_bias.grad, grad[:, -32:].sum(0), rtol=0, atol=0) + actual = [out.detach().clone()] + tensors = [inp, *model.parameters()] + actual.extend(None if t.grad is None else t.grad.clone() for t in tensors) + for tensor in tensors: + tensor.grad = None + expected = model(inp) + expected.backward(grad) + for result, target in zip(actual, [expected, *(t.grad for t in tensors)]): + if target is None: + assert result is None + else: + torch.testing.assert_close(result, target, rtol=0, atol=0) + + +@pytest.mark.skipif(not _opaque_available, reason="torch opaque object API not available") +@pytest.mark.parametrize("bf16_part", ["q", "k"]) +def test_te_split_parameters_mixed_dtype_autocast(bf16_part): + model = te.Linear(64, 192, parameters_split=("q", "k", "v"), params_dtype=torch.float32) + name = f"{bf16_part}_weight" + setattr(model, name, torch.nn.Parameter(getattr(model, name).to(torch.bfloat16))) + inp = torch.randn(16, 64, device="cuda", dtype=torch.float32, requires_grad=True) + grad = torch.randn(16, 192, device="cuda", dtype=torch.bfloat16) + torch._dynamo.reset() + compiled = torch.compile(model, fullgraph=True) + with torch.autocast("cuda", dtype=torch.bfloat16): + expected = model(inp) + expected.backward(grad) + expected_grads = [tensor.grad.clone() for tensor in (inp, *model.parameters())] + inp.grad = None + model.zero_grad(set_to_none=True) + with torch.autocast("cuda", dtype=torch.bfloat16): + actual = compiled(inp) + actual.backward(grad) + actual_grads = [tensor.grad for tensor in (inp, *model.parameters())] + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + torch.testing.assert_close(actual_grads, expected_grads, rtol=0, atol=0) diff --git a/transformer_engine/pytorch/dynamo/custom_op.py b/transformer_engine/pytorch/dynamo/custom_op.py index a87374095a..a81ebccae8 100644 --- a/transformer_engine/pytorch/dynamo/custom_op.py +++ b/transformer_engine/pytorch/dynamo/custom_op.py @@ -37,6 +37,10 @@ * ``TENSOR_OR_QUANTIZED`` -- a field that may be a plain tensor, a bare quantized storage, or ``None``: three slots (the tensor, its flat inner buffers, and a ``__kind__`` tag) so a quantized tensor crosses as its buffers. + Unions including ``ParameterParts`` also accept deferred row concatenations: + all parts cross in the buffer list, with full gradients split outside the op. + Saved parts use ``parameter_parts.restore_from_func_ctx`` in the backward + container; reconstruction of the full view happens only in the real kernel. * ``SIMPLE`` -- every remaining simple value (scalars, enums, sizes, quantizers -- value-opaque constants baked into the graph -- and nested collections of them), gathered into one shared ``OpaqueValueBundle`` slot. @@ -106,6 +110,7 @@ from torch.utils._pytree import tree_flatten, tree_unflatten from .tensor_spec import TensorSpec, to_tensor_spec +from .parameter_parts import ParameterParts from ..quantized_tensor import ( QuantizedTensor, QuantizedTensorStorage, @@ -408,13 +413,13 @@ class _TensorOrQuantizedKind(Enum): NONE = "none" TENSOR = "tensor" STORAGE = "storage" + CONCATENATED = "concatenated" _TQ_KIND_KEY = "__kind__" _SIMPLE_META_SLOT = "_simple_meta" -# Matched by exact member set, so a bare quantized annotation or an accidental -# extra union member is rejected rather than silently taken as tensor-or-quantized. +# Tensor unions are matched by exact member sets. _TQ_MEMBERS = frozenset(get_args(TensorOrQuantized)) @@ -437,14 +442,19 @@ class _FieldPlan: name: str kind: _FieldKind slots: Tuple[_SlotSpec, ...] + allows_concatenation: bool = False def _is_tensor_storage_union(annot: Any) -> bool: - """Whether ``annot`` is exactly the tensor-or-quantized union.""" + """Whether a tensor union can use the tensor / buffers / metadata slots.""" if not _is_union(annot): return False members = frozenset(a for a in get_args(annot) if a is not type(None)) - return members == _TQ_MEMBERS + return members in ( + _TQ_MEMBERS, + _TQ_MEMBERS | {ParameterParts}, + frozenset((torch.Tensor, ParameterParts)), + ) def _is_process_group_annot(annot: Any) -> bool: @@ -486,7 +496,9 @@ def _parse_field(name: str, annot: Any) -> _FieldPlan: _SlotSpec(name + "__tensors", "Tensor[]"), _SlotSpec(name + "__meta", _OPAQUE_VALUE_BUNDLE_TYPE_NAME), ) - return _FieldPlan(name, _FieldKind.TENSOR_OR_QUANTIZED, slots) + return _FieldPlan( + name, _FieldKind.TENSOR_OR_QUANTIZED, slots, ParameterParts in get_args(annot) + ) stripped, is_optional = _strip_optional(annot) if stripped is torch.Tensor: slot = _SlotSpec(name, "Tensor?" if is_optional else "Tensor") @@ -527,6 +539,12 @@ def _pack_tensor_or_quantized(field: _FieldPlan, value: Any, slots: Dict[str, An slots[tensor_slot] = None slots[inner_slot] = [] slots[meta_slot] = OpaqueValueBundle({_TQ_KIND_KEY: _TensorOrQuantizedKind.NONE}) + elif isinstance(value, ParameterParts): + if not field.allows_concatenation: + raise TypeError(f"field {field.name!r} does not accept ParameterParts") + slots[tensor_slot] = None + slots[inner_slot] = list(value.parts) + slots[meta_slot] = OpaqueValueBundle({_TQ_KIND_KEY: _TensorOrQuantizedKind.CONCATENATED}) elif isinstance(value, torch.Tensor): # Plain tensor *and* subclass (e.g. Float8Tensor) pass through the # ``Tensor?`` slot; subclass flattening (if any) is done by the @@ -555,6 +573,8 @@ def _unpack_tensor_or_quantized(field: _FieldPlan, slots: Dict[str, Any]) -> Any return None if kind == _TensorOrQuantizedKind.TENSOR: return slots[tensor_slot] + if kind == _TensorOrQuantizedKind.CONCATENATED: + return ParameterParts(tuple(slots[inner_slot])) return _storage_unflatten(meta, slots[inner_slot]) @@ -741,7 +761,9 @@ def _spec_view(obj: Any, tensor_field_names: Sequence[str]) -> Any: overrides: Dict[str, Any] = {} for name in tensor_field_names: value = getattr(obj, name, None) - if value is not None and not isinstance(value, TensorSpec): + if isinstance(value, ParameterParts): + overrides[name] = value.to_spec() + elif value is not None and not isinstance(value, TensorSpec): overrides[name] = to_tensor_spec(value) if not overrides: return obj @@ -984,8 +1006,14 @@ def _register_base_op( ``pack_result``. """ + concatenated_fields = tuple(f.name for f in plan.fields if f.allows_concatenation) + def _impl(*flat: Any) -> List[torch.Tensor]: obj = plan.unpack(dict(zip(plan.slot_names, flat))) + for name in concatenated_fields: + value = getattr(obj, name) + if isinstance(value, ParameterParts): + setattr(obj, name, value.materialize()) return pack_result(impl(obj)) def _fake(*flat: Any) -> List[torch.Tensor]: @@ -1039,7 +1067,16 @@ def _setup_context(ctx, inputs, output): out_plan.ctx_attrs, tuple(saved_list), ) - tensors_to_save, tensor_objects = prepare_for_saving(*(tensors_to_save_from_setup or ())) + saved = [] + ctx.concatenated_saved_lengths = [] + for value in tensors_to_save_from_setup or (): + if isinstance(value, ParameterParts): + ctx.concatenated_saved_lengths.append(len(value.parts)) + saved.extend(value.parts) + else: + ctx.concatenated_saved_lengths.append(None) + saved.append(value) + tensors_to_save, tensor_objects = prepare_for_saving(*saved) ctx.tensor_objects = tensor_objects ctx.save_for_backward(*tensors_to_save) ctx.backward_objects = bwd_obj @@ -1047,6 +1084,13 @@ def _setup_context(ctx, inputs, output): # Input shapes for the grad slots (SymInt-safe on ctx): a bwd impl may # rederive shapes lossily (e.g. rank-1 inputs come back rank-2), so the # returned grads are viewed back to the true input shapes below. + ctx.concatenated_grad_shapes = { + pos: [tensor.shape for tensor in inputs[pos + 1]] + for pos in grad_targets + if pos in fwd_plan.tensor_or_quantized_offsets() + and inputs[pos] is None + and inputs[pos + 2][_TQ_KIND_KEY] == _TensorOrQuantizedKind.CONCATENATED + } ctx.grad_input_shapes = { pos: inputs[pos].shape for pos in grad_targets if isinstance(inputs[pos], torch.Tensor) } @@ -1073,12 +1117,18 @@ def _autograd_backward(ctx, *grad_outputs): for pos, length in ctx.fwd_tensor_list_lengths.items(): out[pos] = [None] * length for pos, g in zip(grad_targets, grads): + if pos in ctx.concatenated_grad_shapes: + shapes = ctx.concatenated_grad_shapes[pos] + if g is not None: + out[pos + 1] = list(torch.split(g, [shape[0] for shape in shapes], dim=0)) + continue if g is not None: shape = ctx.grad_input_shapes.get(pos) if shape is not None and g.shape != shape: g = g.view(shape) out[pos] = g ctx.grad_input_shapes = None + ctx.concatenated_grad_shapes = None return tuple(out) fwd_op.register_autograd(_autograd_backward, setup_context=_setup_context) diff --git a/transformer_engine/pytorch/dynamo/parameter_parts.py b/transformer_engine/pytorch/dynamo/parameter_parts.py new file mode 100644 index 0000000000..c0ad545d78 --- /dev/null +++ b/transformer_engine/pytorch/dynamo/parameter_parts.py @@ -0,0 +1,98 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Deferred row concatenation for tensors consumed inside a custom op.""" + +from dataclasses import dataclass +from typing import Tuple, TypeVar, Union + +import torch +from torch._prims_common import make_contiguous_strides_for + +from .tensor_spec import TensorSpec +from ..quantized_tensor import ( + QuantizedTensor, + QuantizedTensorStorage, + restore_from_func_ctx as restore_tensor_ctx, +) + + +@dataclass(frozen=True, slots=True) +class ParameterParts: + """Keep every parameter visible to autograd until the opaque consumer runs.""" + + parts: Tuple[torch.Tensor, ...] + + def __post_init__(self): + object.__setattr__(self, "parts", tuple(self.parts)) + if not self.parts: + raise ValueError("ParameterParts requires at least one tensor") + if any( + not isinstance(part, torch.Tensor) + or isinstance(part, (QuantizedTensor, QuantizedTensorStorage)) + for part in self.parts + ): + raise TypeError("ParameterParts only supports non-quantized tensors") + + def to_spec(self) -> TensorSpec: + """Describe the concatenation without accessing its storage.""" + first = self.parts[0] + dtype = first.dtype + for tensor in self.parts[1:]: + dtype = torch.promote_types(dtype, tensor.dtype) + return TensorSpec( + shape=(sum(t.shape[0] for t in self.parts), *first.shape[1:]), + dtype=dtype, + device=first.device, + requires_grad=any(t.requires_grad for t in self.parts), + ) + + def materialize(self): + """Build a checked view, or copy parts whose storage is not adjacent.""" + first = self.parts[0] + storage = first.untyped_storage() + offset = first.storage_offset() + for tensor in self.parts: + if ( + tensor.dtype != first.dtype + or tensor.device != first.device + or tensor.is_neg() + or tensor.is_conj() + ): + return torch.cat(self.parts) + if ( + tensor.shape[1:] != first.shape[1:] + or not tensor.is_contiguous() + or tensor.untyped_storage().data_ptr() != storage.data_ptr() + or tensor.storage_offset() != offset + ): + return torch.cat(self.parts) + offset += tensor.numel() + if offset * first.element_size() > storage.nbytes(): + return torch.cat(self.parts) + shape = (sum(t.shape[0] for t in self.parts), *first.shape[1:]) + return first.as_strided(shape, make_contiguous_strides_for(shape)) + + +_TensorT = TypeVar("_TensorT") +ConcatInput = Union[_TensorT, ParameterParts] + + +def restore_from_func_ctx(ctx): + """Restore original split parameters saved by the custom-op framework.""" + tensors = restore_tensor_ctx(ctx) + lengths = getattr(ctx, "concatenated_saved_lengths", None) + if lengths is None: + return tensors + restored = [] + offset = 0 + for length in lengths: + if length is None: + restored.append(tensors[offset]) + offset += 1 + else: + restored.append(ParameterParts(tuple(tensors[offset : offset + length]))) + offset += length + ctx.concatenated_saved_lengths = None + return restored diff --git a/transformer_engine/pytorch/module/_common.py b/transformer_engine/pytorch/module/_common.py index 3d04882698..7785eda7eb 100644 --- a/transformer_engine/pytorch/module/_common.py +++ b/transformer_engine/pytorch/module/_common.py @@ -6,14 +6,17 @@ import dataclasses import queue -from typing import Any, Callable, List, Optional, Tuple, Union +from typing import Any, Callable, List, Optional, Sequence, Tuple, Union import torch from .. import cpp_extensions as tex from ..constants import TE_DType from ..distributed import in_fp8_activation_recompute_phase +from ..dynamo import TensorOrQuantized +from ..dynamo.parameter_parts import ConcatInput, ParameterParts from ..export import is_in_onnx_export_mode +from ..quantized_tensor import QuantizedTensor, QuantizedTensorStorage from ..quantization import FP8GlobalStateManager from ..tensor.hybrid_tensor import HybridQuantizer from ..utils import get_default_init_method @@ -239,6 +242,27 @@ def noop_cat( return _NoopCatFunc.apply(dim, *tensors) +def concat_input( + tensors: Sequence[TensorOrQuantized], *, defer: bool = False +) -> ConcatInput[TensorOrQuantized]: + """Prepare one parameter operand for eager or a compiled consumer.""" + if len(tensors) == 1: + return tensors[0] + if ( + defer + and tensors + and all( + isinstance(t, torch.Tensor) + and not isinstance(t, (QuantizedTensor, QuantizedTensorStorage)) + for t in tensors + ) + ): + return ParameterParts(tuple(tensors)) + if torch.compiler.is_compiling(): + return torch.cat(tensors) + return noop_cat(tensors) + + @dataclasses.dataclass class _ParameterInitMeta: """ diff --git a/transformer_engine/pytorch/module/base.py b/transformer_engine/pytorch/module/base.py index 9ead1e7365..8909720d02 100644 --- a/transformer_engine/pytorch/module/base.py +++ b/transformer_engine/pytorch/module/base.py @@ -46,6 +46,7 @@ _fsdp_gather_tensors, ) from ..constants import dist_group_type +from ..dynamo.tensor_spec import TensorSpec from ..cpp_extensions.gemm import _NUM_MAX_UB_STREAMS from ..quantized_tensor import QuantizedTensor, QuantizedTensorStorage, Quantizer from ..tensor.float8_tensor import Float8Quantizer, Float8CurrentScalingQuantizer @@ -1322,7 +1323,7 @@ def _get_weight_quantizers(self) -> List[Quantizer]: def _enable_weight_preswizzle( self, quantizer: Quantizer, - weight: torch.Tensor, + weight: Union[torch.Tensor, TensorSpec], ) -> bool: """Whether to fuse scale-factor swizzling into weight quantization. @@ -1339,7 +1340,7 @@ def _enable_weight_preswizzle( if isinstance(quantizer, MXFP8Quantizer): return True if isinstance(quantizer, NVFP4Quantizer): - rows, cols = weight.numel() // weight.shape[-1], weight.shape[-1] + rows, cols = math.prod(weight.shape[:-1]), weight.shape[-1] arch_supported = get_device_compute_capability() >= (10, 0) if quantizer.with_rht: return arch_supported and rows % 64 == 0 and cols % 128 == 0 diff --git a/transformer_engine/pytorch/module/linear.py b/transformer_engine/pytorch/module/linear.py index 7f00e19644..8f66102c40 100644 --- a/transformer_engine/pytorch/module/linear.py +++ b/transformer_engine/pytorch/module/linear.py @@ -32,7 +32,8 @@ from ._common import ( can_reconstruct_wgrad_input_from_original, check_fp8_reduce_and_update, - noop_cat, + concat_input, + sum_bias_grad, set_quantizer_amax_reduction_group, set_quantizer_usage_for_wgrad_all_gather, WeightGradStore, @@ -82,8 +83,8 @@ QuantizedTensorStorage, Quantizer, prepare_for_saving, - restore_from_func_ctx, ) +from ..dynamo.parameter_parts import ConcatInput, ParameterParts, restore_from_func_ctx from ..dynamo import ( TensorSpec, TensorOrQuantized, @@ -110,9 +111,9 @@ class LinearFwdArgs: """Single-argument bag for the forward path of :class:`_Linear`.""" # --- Differentiable tensors (also passed positionally to autograd) --- - weight: TensorOrQuantized + weight: ConcatInput[TensorOrQuantized] inp: torch.Tensor - bias: Optional[torch.Tensor] + bias: Optional[ConcatInput[torch.Tensor]] # --- Non-differentiable cached tensors --- # TensorOrQuantized so a cached quantized workspace can cross the op boundary. @@ -238,9 +239,9 @@ class LinearBwdArgs: # --- Saved / restored tensors (populated at backward entry) --- grad_output: Optional[torch.Tensor] = None inputmat: Optional[TensorOrQuantized] = None - weight_fp8: Optional[TensorOrQuantized] = None - saved_weight: Optional[TensorOrQuantized] = None - bias: Optional[torch.Tensor] = None + weight_fp8: Optional[ConcatInput[TensorOrQuantized]] = None + saved_weight: Optional[ConcatInput[TensorOrQuantized]] = None + bias: Optional[ConcatInput[torch.Tensor]] = None # --- Quantizers --- input_quantizer: Optional[Quantizer] = None @@ -1263,6 +1264,8 @@ def _linear_backward_impl(args: LinearBwdArgs) -> Tuple[Union[torch.Tensor, None grad_output_quantizer, ) nvtx_range_pop(f"{nvtx_label}.grad_output_preprocess") + if bwd_args.use_bias and grad_bias is None and not bwd_args.requires_wgrad: + grad_bias = sum_bias_grad(grad_output) # -------------------------------------------------- # Grad output tensor is ready for computing grad input... @@ -1775,11 +1778,7 @@ def _linear_backward_fake( ) grad_bias = None - # FP8 backward computes bgrad in grad_output_preprocess whenever bias is - # used; in high precision it is fused into the wgrad GEMM, so it only - # exists when wgrad runs. - fp8_bwd = args.fp8 and args.backward_override is None - if args.use_bias and (args.requires_wgrad or fp8_bwd): + if args.use_bias: grad_bias = TensorSpec( shape=(out_features,), dtype=out_dtype, device=args.grad_output.device ) @@ -1931,7 +1930,9 @@ class Linear(TransformerEngineBaseModule): (preferably an OrderedDict) is provided, the keys are used as names and values as split sizes along dim 0. The resulting parameters will have names that end in ``_weight`` or ``_bias``, so trailing underscores are - stripped from any provided names. + stripped from any provided names. Under ``torch.compile``, adjacent + parts sharing storage are consumed without a concatenation copy. + Disjoint parts and returned split biases still require concatenation. device : Union[torch.device, str], default = "cuda" The device on which the parameters of the model will be allocated. It is the user's responsibility to ensure all parameters are moved to the GPU before running the @@ -2399,7 +2400,17 @@ def forward( inp = self.prepare_forward(inp, allow_non_contiguous=isinstance(inp, QuantizedTensor)) try: - weight_tensor, bias_tensor = self._get_weight_and_bias_tensors() + weight_tensor, bias_tensor = self._get_weight_and_bias_tensors( + defer_concatenation=torch.compiler.is_compiling() and _linear_op is not None + ) + weight_meta = ( + weight_tensor.to_spec() + if isinstance(weight_tensor, ParameterParts) + else weight_tensor + ) + bias_meta = ( + bias_tensor.to_spec() if isinstance(bias_tensor, ParameterParts) else bias_tensor + ) quantizers = ( self._get_quantizers(fp8_output, fp8_grad, is_grad_enabled) @@ -2421,7 +2432,7 @@ def forward( ) = quantizers if weight_quantizer is not None and not debug: weight_quantizer.optimize_for_gemm = self._enable_weight_preswizzle( - weight_quantizer, weight_tensor + weight_quantizer, weight_meta ) use_compiled_op = torch.compiler.is_compiling() and _linear_op is not None @@ -2451,7 +2462,7 @@ def forward( backward_override = None custom = is_custom(input_quantizer) or is_custom(weight_quantizer) backward_input_needs_gather = ( - weight_tensor.requires_grad + weight_meta.requires_grad and self.parallel_mode == "column" and self.sequence_parallel ) @@ -2484,9 +2495,9 @@ def forward( weight_workspace=weight_workspace, # requires_grad flags input_requires_grad=inp.requires_grad, - weight_requires_grad=weight_tensor.requires_grad, + weight_requires_grad=weight_meta.requires_grad, bias_requires_grad=( - linear_bias_tensor.requires_grad if linear_bias_tensor is not None else False + bias_meta.requires_grad if linear_bias_tensor is not None else False ), # quantizers input_quantizer=input_quantizer, @@ -2546,9 +2557,15 @@ def forward( msg=f"te.Linear falling back to eager: {fallback_reason}" ) use_compiled_op = False + weight_tensor, bias_tensor = self._get_weight_and_bias_tensors() + linear_bias_tensor = ( + bias_tensor if self.apply_bias and not self.gemm_bias_unfused_add else None + ) + fwd_args.weight = weight_tensor + fwd_args.bias = linear_bias_tensor if use_compiled_op: - check_gemm_dims(inp, weight_tensor, self.fp8) + check_gemm_dims(inp.shape, weight_meta.shape, self.fp8) out, new_weight_workspace = _linear_op(fwd_args) else: out, new_weight_workspace = _linear_eager( @@ -2643,8 +2660,7 @@ def _compile_eager_fallback_reason( prepare_forward. Quantizer checks stay in compile_unsupported_reason.""" if debug: return "debug instrumentation (nvidia-dlfw-inspect)" - weight_tensor, bias_tensor = self._get_weight_and_bias_tensors() - if is_distributed_weight(weight_tensor): + if any(is_distributed_weight(getattr(self, name)) for name in self.weight_names): return "a DistributedWeight (custom weight parallelism, e.g. GTP)" if isinstance(inp, (QuantizedTensor, QuantizedTensorStorage)): return "a quantized input tensor" @@ -2652,8 +2668,10 @@ def _compile_eager_fallback_reason( return "manual TE FSDP (fsdp_group); use FSDP2 or MCore FSDP" any_requires_grad = ( inp.requires_grad - or weight_tensor.requires_grad - or (bias_tensor is not None and bias_tensor.requires_grad) + or any(getattr(self, name).requires_grad for name in self.weight_names) + or ( + self.use_bias and any(getattr(self, name).requires_grad for name in self.bias_names) + ) ) if fp8_output and is_grad_enabled and any_requires_grad: return "differentiable fp8_output=True" @@ -2693,14 +2711,16 @@ def _forward_eager_fallback( fp8_grad=fp8_grad, ) - def _get_weight_and_bias_tensors(self) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: - # Get concatenated weight and bias tensors - unfused_weights = self._get_weight_tensors() - weight_tensor = noop_cat(unfused_weights) + def _get_weight_and_bias_tensors(self, defer_concatenation=False): + weights = self._get_weight_tensors() + weight_tensor = concat_input(weights, defer=defer_concatenation) + bias_tensor = None if self.use_bias: - bias_tensor = noop_cat([getattr(self, name) for name in self.bias_names]) - else: - bias_tensor = None + biases = [getattr(self, name) for name in self.bias_names] + bias_tensor = concat_input( + biases, + defer=defer_concatenation and self.apply_bias and not self.gemm_bias_unfused_add, + ) return weight_tensor, bias_tensor def onnx_forward( diff --git a/transformer_engine/pytorch/utils.py b/transformer_engine/pytorch/utils.py index ca77f7b8bb..53e27ae6f6 100644 --- a/transformer_engine/pytorch/utils.py +++ b/transformer_engine/pytorch/utils.py @@ -716,21 +716,21 @@ def assert_dim_for_fp8_exec(*tensors: List[torch.Tensor]) -> None: ) -def check_gemm_dims(inp: torch.Tensor, weight: torch.Tensor, fp8: bool) -> None: +def check_gemm_dims(inp_shape: Sequence[int], weight_shape: Sequence[int], fp8: bool) -> None: """Emit the TN GEMM (``y = x @ w^T``) dim constraints as ``torch._check`` guards at trace time. torch.compile path only; eager validation lives in the op impl. Messages are constant: Dynamo forbids tensor closures here. """ # pylint: disable=protected-access torch._check( - inp.shape[-1] == weight.shape[-1], + inp_shape[-1] == weight_shape[-1], lambda: "GEMM not possible: input last dim must equal in_features", ) if not fp8: return - for tensor, name in ((inp, "input"), (weight, "weight")): + for shape, name in ((inp_shape, "input"), (weight_shape, "weight")): torch._check( - math.prod(tensor.shape[:-1]) % 8 == 0 and tensor.shape[-1] % 16 == 0, + math.prod(shape[:-1]) % 8 == 0 and shape[-1] % 16 == 0, lambda n=name: ( f"FP8 execution requires the {n}'s product of all dimensions except the" " last to be divisible by 8 and its last dimension to be divisible by 16"