From fdaea9e1f31542453f9d99e53587cf6376e6a803 Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Mon, 5 Oct 2026 11:42:18 +0200 Subject: [PATCH 1/4] [PyTorch] Avoid concatenation copies for compiled split Linear parameters Pass split Linear parameters through the custom-op boundary, reconstruct checked views at execution time and split full gradients outside backward. Preserve concatenation for disjoint storage and returned biases, handle promoted dtypes, and compute bias gradients when weights are frozen. Signed-off-by: Pawel Gadzinski --- tests/pytorch/test_torch_compile.py | 214 ++++++++++++++++++ .../pytorch/dynamo/concatenated_tensor.py | 89 ++++++++ .../pytorch/dynamo/custom_op.py | 60 ++++- transformer_engine/pytorch/module/linear.py | 59 +++-- 4 files changed, 396 insertions(+), 26 deletions(-) create mode 100644 transformer_engine/pytorch/dynamo/concatenated_tensor.py diff --git a/tests/pytorch/test_torch_compile.py b/tests/pytorch/test_torch_compile.py index eae6f0a8a2..195b9bc339 100644 --- a/tests/pytorch/test_torch_compile.py +++ b/tests/pytorch/test_torch_compile.py @@ -2720,3 +2720,217 @@ 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): + 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=fp8_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.parametrize( + "layout", ["adjacent", "offset", "disjoint", "strided", "singleton", "negative", "conjugate"] +) +def test_concatenated_tensor_storage(layout): + from transformer_engine.pytorch.dynamo.concatenated_tensor import ConcatenatedTensor + + 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 = ConcatenatedTensor(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.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/concatenated_tensor.py b/transformer_engine/pytorch/dynamo/concatenated_tensor.py new file mode 100644 index 0000000000..6f6c39ebc6 --- /dev/null +++ b/transformer_engine/pytorch/dynamo/concatenated_tensor.py @@ -0,0 +1,89 @@ +# 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 +import math +from typing import List + +import torch +from torch._prims_common import make_contiguous_strides_for + +from ..quantized_tensor import restore_from_func_ctx as restore_tensor_ctx + + +@dataclass +class ConcatenatedTensor: + """Keep every parameter visible to autograd until the opaque consumer runs.""" + + tensors: List[torch.Tensor] + + @property + def shape(self): + """Combined shape along dimension zero.""" + return (sum(t.shape[0] for t in self.tensors), *self.tensors[0].shape[1:]) + + @property + def dtype(self): + """Parameter dtype.""" + dtype = self.tensors[0].dtype + for tensor in self.tensors[1:]: + dtype = torch.promote_types(dtype, tensor.dtype) + return dtype + + @property + def device(self): + """Parameter device.""" + return self.tensors[0].device + + @property + def requires_grad(self): + """Whether any part requires a gradient.""" + return any(t.requires_grad for t in self.tensors) + + def numel(self): + """Combined element count.""" + return math.prod(self.shape) + + def materialize(self): + """Build a checked view, or copy parts whose storage is not adjacent.""" + first = self.tensors[0] + storage = first.untyped_storage() + offset = first.storage_offset() + for tensor in self.tensors: + if tensor.dtype != first.dtype or tensor.device != first.device: + return torch.cat(self.tensors) + if ( + tensor.shape[1:] != first.shape[1:] + or not tensor.is_contiguous() + or tensor.is_neg() + or tensor.is_conj() + or tensor.untyped_storage().data_ptr() != storage.data_ptr() + or tensor.storage_offset() != offset + ): + return torch.cat(self.tensors) + offset += tensor.numel() + if offset * first.element_size() > storage.nbytes(): + return torch.cat(self.tensors) + return first.as_strided(self.shape, make_contiguous_strides_for(self.shape)) + + +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(ConcatenatedTensor(list(tensors[offset : offset + length]))) + offset += length + ctx.concatenated_saved_lengths = None + return restored diff --git a/transformer_engine/pytorch/dynamo/custom_op.py b/transformer_engine/pytorch/dynamo/custom_op.py index a87374095a..5225e29345 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 ``ConcatenatedTensor`` also accept deferred row concatenations: + all parts cross in the buffer list, with full gradients split outside the op. + Saved parts use ``concatenated_tensor.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 .concatenated_tensor import ConcatenatedTensor 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 | {ConcatenatedTensor}, + frozenset((torch.Tensor, ConcatenatedTensor)), + ) 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, ConcatenatedTensor 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, ConcatenatedTensor): + if not field.allows_concatenation: + raise TypeError(f"field {field.name!r} does not accept ConcatenatedTensor") + slots[tensor_slot] = None + slots[inner_slot] = value.tensors + 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 ConcatenatedTensor(slots[inner_slot]) return _storage_unflatten(meta, slots[inner_slot]) @@ -984,8 +1004,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, ConcatenatedTensor): + setattr(obj, name, value.materialize()) return pack_result(impl(obj)) def _fake(*flat: Any) -> List[torch.Tensor]: @@ -1039,7 +1065,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, ConcatenatedTensor): + ctx.concatenated_saved_lengths.append(len(value.tensors)) + saved.extend(value.tensors) + 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 +1082,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 +1115,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/module/linear.py b/transformer_engine/pytorch/module/linear.py index 7f00e19644..acde633499 100644 --- a/transformer_engine/pytorch/module/linear.py +++ b/transformer_engine/pytorch/module/linear.py @@ -33,6 +33,7 @@ can_reconstruct_wgrad_input_from_original, check_fp8_reduce_and_update, noop_cat, + 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.concatenated_tensor import ConcatenatedTensor, 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: Union[TensorOrQuantized, ConcatenatedTensor] inp: torch.Tensor - bias: Optional[torch.Tensor] + bias: Optional[Union[torch.Tensor, ConcatenatedTensor]] # --- 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[Union[TensorOrQuantized, ConcatenatedTensor]] = None + saved_weight: Optional[Union[TensorOrQuantized, ConcatenatedTensor]] = None + bias: Optional[Union[torch.Tensor, ConcatenatedTensor]] = 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,9 @@ 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 + ) quantizers = ( self._get_quantizers(fp8_output, fp8_grad, is_grad_enabled) @@ -2546,6 +2549,12 @@ 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) @@ -2693,14 +2702,24 @@ 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 = ( + ConcatenatedTensor(weights) + if defer_concatenation and len(weights) > 1 + else noop_cat(weights) + ) + 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 = ( + ConcatenatedTensor(biases) + if defer_concatenation + and len(biases) > 1 + and self.apply_bias + and not self.gemm_bias_unfused_add + else noop_cat(biases) + ) return weight_tensor, bias_tensor def onnx_forward( From 289cff12aedad7d8ce09442ca2366b16b42795b3 Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Mon, 5 Oct 2026 12:40:05 +0200 Subject: [PATCH 2/4] [PyTorch] Make deferred parameter concatenation explicit Replace ConcatenatedTensor with DeferredCat(parts), expose metadata through to_spec and keep materialization inside the custom op. Preserve deferred backward operands for non-FP8 training, reject quantized parts, and make dimension checks consume explicit metadata. Signed-off-by: Pawel Gadzinski --- tests/pytorch/test_torch_compile.py | 44 +++++++++- .../pytorch/dynamo/custom_op.py | 32 ++++---- ...concatenated_tensor.py => deferred_cat.py} | 80 ++++++++++--------- transformer_engine/pytorch/module/base.py | 5 +- transformer_engine/pytorch/module/linear.py | 39 +++++---- transformer_engine/pytorch/utils.py | 8 +- 6 files changed, 132 insertions(+), 76 deletions(-) rename transformer_engine/pytorch/dynamo/{concatenated_tensor.py => deferred_cat.py} (50%) diff --git a/tests/pytorch/test_torch_compile.py b/tests/pytorch/test_torch_compile.py index 195b9bc339..1f8081d9f6 100644 --- a/tests/pytorch/test_torch_compile.py +++ b/tests/pytorch/test_torch_compile.py @@ -2857,8 +2857,8 @@ def test_te_split_parameters_saved_versions(): @pytest.mark.parametrize( "layout", ["adjacent", "offset", "disjoint", "strided", "singleton", "negative", "conjugate"] ) -def test_concatenated_tensor_storage(layout): - from transformer_engine.pytorch.dynamo.concatenated_tensor import ConcatenatedTensor +def test_deferred_cat_storage(layout): + from transformer_engine.pytorch.dynamo.deferred_cat import DeferredCat storage = torch.arange(128, device="cuda").view(16, 8) if layout == "offset": @@ -2874,7 +2874,7 @@ def test_concatenated_tensor_storage(layout): parts = list(storage.split([1, storage.shape[0] - 3, 2])) if layout == "disjoint": parts[1] = parts[1].clone() - result = ConcatenatedTensor(parts).materialize() + result = DeferredCat(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() @@ -2882,6 +2882,44 @@ def test_concatenated_tensor_storage(layout): 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_deferred_cat_spec(fake, mixed_dtype): + from transformer_engine.pytorch.dynamo.deferred_cat import DeferredCat + + 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 = DeferredCat(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_deferred_cat_rejects_quantized_parts(internal): + from transformer_engine.pytorch.dynamo.deferred_cat import DeferredCat + + 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"): + DeferredCat([tensor, 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): diff --git a/transformer_engine/pytorch/dynamo/custom_op.py b/transformer_engine/pytorch/dynamo/custom_op.py index 5225e29345..e3cddbd6d6 100644 --- a/transformer_engine/pytorch/dynamo/custom_op.py +++ b/transformer_engine/pytorch/dynamo/custom_op.py @@ -37,9 +37,9 @@ * ``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 ``ConcatenatedTensor`` also accept deferred row concatenations: + Unions including ``DeferredCat`` also accept deferred row concatenations: all parts cross in the buffer list, with full gradients split outside the op. - Saved parts use ``concatenated_tensor.restore_from_func_ctx`` in the backward + Saved parts use ``deferred_cat.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 @@ -110,7 +110,7 @@ from torch.utils._pytree import tree_flatten, tree_unflatten from .tensor_spec import TensorSpec, to_tensor_spec -from .concatenated_tensor import ConcatenatedTensor +from .deferred_cat import DeferredCat from ..quantized_tensor import ( QuantizedTensor, QuantizedTensorStorage, @@ -452,8 +452,8 @@ def _is_tensor_storage_union(annot: Any) -> bool: members = frozenset(a for a in get_args(annot) if a is not type(None)) return members in ( _TQ_MEMBERS, - _TQ_MEMBERS | {ConcatenatedTensor}, - frozenset((torch.Tensor, ConcatenatedTensor)), + _TQ_MEMBERS | {DeferredCat}, + frozenset((torch.Tensor, DeferredCat)), ) @@ -497,7 +497,7 @@ def _parse_field(name: str, annot: Any) -> _FieldPlan: _SlotSpec(name + "__meta", _OPAQUE_VALUE_BUNDLE_TYPE_NAME), ) return _FieldPlan( - name, _FieldKind.TENSOR_OR_QUANTIZED, slots, ConcatenatedTensor in get_args(annot) + name, _FieldKind.TENSOR_OR_QUANTIZED, slots, DeferredCat in get_args(annot) ) stripped, is_optional = _strip_optional(annot) if stripped is torch.Tensor: @@ -539,11 +539,11 @@ 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, ConcatenatedTensor): + elif isinstance(value, DeferredCat): if not field.allows_concatenation: - raise TypeError(f"field {field.name!r} does not accept ConcatenatedTensor") + raise TypeError(f"field {field.name!r} does not accept DeferredCat") slots[tensor_slot] = None - slots[inner_slot] = value.tensors + slots[inner_slot] = 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 @@ -574,7 +574,7 @@ def _unpack_tensor_or_quantized(field: _FieldPlan, slots: Dict[str, Any]) -> Any if kind == _TensorOrQuantizedKind.TENSOR: return slots[tensor_slot] if kind == _TensorOrQuantizedKind.CONCATENATED: - return ConcatenatedTensor(slots[inner_slot]) + return DeferredCat(slots[inner_slot]) return _storage_unflatten(meta, slots[inner_slot]) @@ -761,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, DeferredCat): + 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 @@ -1010,7 +1012,7 @@ 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, ConcatenatedTensor): + if isinstance(value, DeferredCat): setattr(obj, name, value.materialize()) return pack_result(impl(obj)) @@ -1068,9 +1070,9 @@ def _setup_context(ctx, inputs, output): saved = [] ctx.concatenated_saved_lengths = [] for value in tensors_to_save_from_setup or (): - if isinstance(value, ConcatenatedTensor): - ctx.concatenated_saved_lengths.append(len(value.tensors)) - saved.extend(value.tensors) + if isinstance(value, DeferredCat): + ctx.concatenated_saved_lengths.append(len(value.parts)) + saved.extend(value.parts) else: ctx.concatenated_saved_lengths.append(None) saved.append(value) diff --git a/transformer_engine/pytorch/dynamo/concatenated_tensor.py b/transformer_engine/pytorch/dynamo/deferred_cat.py similarity index 50% rename from transformer_engine/pytorch/dynamo/concatenated_tensor.py rename to transformer_engine/pytorch/dynamo/deferred_cat.py index 6f6c39ebc6..cae9527807 100644 --- a/transformer_engine/pytorch/dynamo/concatenated_tensor.py +++ b/transformer_engine/pytorch/dynamo/deferred_cat.py @@ -5,69 +5,73 @@ """Deferred row concatenation for tensors consumed inside a custom op.""" from dataclasses import dataclass -import math from typing import List import torch from torch._prims_common import make_contiguous_strides_for -from ..quantized_tensor import restore_from_func_ctx as restore_tensor_ctx +from .tensor_spec import TensorSpec +from ..quantized_tensor import ( + QuantizedTensor, + QuantizedTensorStorage, + restore_from_func_ctx as restore_tensor_ctx, +) @dataclass -class ConcatenatedTensor: +class DeferredCat: """Keep every parameter visible to autograd until the opaque consumer runs.""" - tensors: List[torch.Tensor] + parts: List[torch.Tensor] - @property - def shape(self): - """Combined shape along dimension zero.""" - return (sum(t.shape[0] for t in self.tensors), *self.tensors[0].shape[1:]) + def __post_init__(self): + if not self.parts: + raise ValueError("DeferredCat requires at least one tensor") + if any( + not isinstance(part, torch.Tensor) + or isinstance(part, (QuantizedTensor, QuantizedTensorStorage)) + for part in self.parts + ): + raise TypeError("DeferredCat only supports non-quantized tensors") - @property - def dtype(self): - """Parameter dtype.""" - dtype = self.tensors[0].dtype - for tensor in self.tensors[1:]: + 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 dtype - - @property - def device(self): - """Parameter device.""" - return self.tensors[0].device - - @property - def requires_grad(self): - """Whether any part requires a gradient.""" - return any(t.requires_grad for t in self.tensors) - - def numel(self): - """Combined element count.""" - return math.prod(self.shape) + 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.tensors[0] + first = self.parts[0] storage = first.untyped_storage() offset = first.storage_offset() - for tensor in self.tensors: - if tensor.dtype != first.dtype or tensor.device != first.device: - return torch.cat(self.tensors) + for tensor in self.parts: if ( - tensor.shape[1:] != first.shape[1:] - or not tensor.is_contiguous() + 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.tensors) + return torch.cat(self.parts) offset += tensor.numel() if offset * first.element_size() > storage.nbytes(): - return torch.cat(self.tensors) - return first.as_strided(self.shape, make_contiguous_strides_for(self.shape)) + 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)) def restore_from_func_ctx(ctx): @@ -83,7 +87,7 @@ def restore_from_func_ctx(ctx): restored.append(tensors[offset]) offset += 1 else: - restored.append(ConcatenatedTensor(list(tensors[offset : offset + length]))) + restored.append(DeferredCat(list(tensors[offset : offset + length]))) offset += length ctx.concatenated_saved_lengths = None return restored 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 acde633499..e9b3f1d7cd 100644 --- a/transformer_engine/pytorch/module/linear.py +++ b/transformer_engine/pytorch/module/linear.py @@ -84,7 +84,7 @@ Quantizer, prepare_for_saving, ) -from ..dynamo.concatenated_tensor import ConcatenatedTensor, restore_from_func_ctx +from ..dynamo.deferred_cat import DeferredCat, restore_from_func_ctx from ..dynamo import ( TensorSpec, TensorOrQuantized, @@ -111,9 +111,9 @@ class LinearFwdArgs: """Single-argument bag for the forward path of :class:`_Linear`.""" # --- Differentiable tensors (also passed positionally to autograd) --- - weight: Union[TensorOrQuantized, ConcatenatedTensor] + weight: Union[TensorOrQuantized, DeferredCat] inp: torch.Tensor - bias: Optional[Union[torch.Tensor, ConcatenatedTensor]] + bias: Optional[Union[torch.Tensor, DeferredCat]] # --- Non-differentiable cached tensors --- # TensorOrQuantized so a cached quantized workspace can cross the op boundary. @@ -239,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[Union[TensorOrQuantized, ConcatenatedTensor]] = None - saved_weight: Optional[Union[TensorOrQuantized, ConcatenatedTensor]] = None - bias: Optional[Union[torch.Tensor, ConcatenatedTensor]] = None + weight_fp8: Optional[Union[TensorOrQuantized, DeferredCat]] = None + saved_weight: Optional[Union[TensorOrQuantized, DeferredCat]] = None + bias: Optional[Union[torch.Tensor, DeferredCat]] = None # --- Quantizers --- input_quantizer: Optional[Quantizer] = None @@ -2403,6 +2403,12 @@ def forward( 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, DeferredCat) else weight_tensor + ) + bias_meta = ( + bias_tensor.to_spec() if isinstance(bias_tensor, DeferredCat) else bias_tensor + ) quantizers = ( self._get_quantizers(fp8_output, fp8_grad, is_grad_enabled) @@ -2424,7 +2430,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 @@ -2454,7 +2460,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 ) @@ -2487,9 +2493,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, @@ -2557,7 +2563,7 @@ def forward( 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( @@ -2705,19 +2711,24 @@ def _forward_eager_fallback( def _get_weight_and_bias_tensors(self, defer_concatenation=False): weights = self._get_weight_tensors() weight_tensor = ( - ConcatenatedTensor(weights) - if defer_concatenation and len(weights) > 1 + DeferredCat(weights) + if defer_concatenation + and len(weights) > 1 + and all(not isinstance(w, (QuantizedTensor, QuantizedTensorStorage)) for w in weights) else noop_cat(weights) ) bias_tensor = None if self.use_bias: biases = [getattr(self, name) for name in self.bias_names] bias_tensor = ( - ConcatenatedTensor(biases) + DeferredCat(biases) if defer_concatenation and len(biases) > 1 and self.apply_bias and not self.gemm_bias_unfused_add + and all( + not isinstance(b, (QuantizedTensor, QuantizedTensorStorage)) for b in biases + ) else noop_cat(biases) ) return weight_tensor, bias_tensor 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" From 43a33bb348549658408a145ecd2c85c8836651ee Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Mon, 5 Oct 2026 16:23:23 +0200 Subject: [PATCH 3/4] [PyTorch] Simplify split Linear parameter handling Keep singleton operands unwrapped and represent split parameters with immutable ParameterParts through the ConcatInput alias. Centralize eager and compiled input preparation, let the custom-op adapter restore saved parts, and select the execution path before concatenating parameters. Name backward weights by their role and update the existing MLA caller. Signed-off-by: Pawel Gadzinski --- tests/pytorch/test_torch_compile.py | 87 +++++- .../pytorch/attention/fused_mla_q_uproj.py | 4 +- .../pytorch/dynamo/custom_op.py | 116 ++++---- .../{deferred_cat.py => parameter_parts.py} | 33 +-- transformer_engine/pytorch/module/_common.py | 26 +- transformer_engine/pytorch/module/linear.py | 252 ++++++------------ 6 files changed, 270 insertions(+), 248 deletions(-) rename transformer_engine/pytorch/dynamo/{deferred_cat.py => parameter_parts.py} (72%) diff --git a/tests/pytorch/test_torch_compile.py b/tests/pytorch/test_torch_compile.py index 1f8081d9f6..ba0432b68c 100644 --- a/tests/pytorch/test_torch_compile.py +++ b/tests/pytorch/test_torch_compile.py @@ -2048,7 +2048,7 @@ def fn(inp): ) -# Configs rejected by LinearFwdArgs.compile_unsupported_reason() that a +# Configs rejected by Linear._compile_eager_fallback_reason() that a # single-GPU unit test can construct. Distributed-only reasons (fsdp_group, # DistributedWeight) and CPU offloading need machinery this file doesn't have; # delayed scaling is a hard error (check_recipe_support), tested separately. @@ -2782,6 +2782,7 @@ def test_te_split_parameters_compile(compile_mode, dtype, equal_splits): @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, @@ -2792,7 +2793,7 @@ def test_te_split_parameters_recipes(fp8_recipe, bias_mode): ) def fn(inp): - with te.autocast(enabled=fp8_recipe is not None, recipe=fp8_recipe): + with te.autocast(enabled=fp8_recipe is not None, recipe=autocast_recipe): out = model(inp) if bias_mode == "returned": out, bias = out @@ -2854,11 +2855,73 @@ def test_te_split_parameters_saved_versions(): 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_deferred_cat_storage(layout): - from transformer_engine.pytorch.dynamo.deferred_cat import DeferredCat +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": @@ -2874,7 +2937,7 @@ def test_deferred_cat_storage(layout): parts = list(storage.split([1, storage.shape[0] - 3, 2])) if layout == "disjoint": parts[1] = parts[1].clone() - result = DeferredCat(parts).materialize() + 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() @@ -2884,8 +2947,8 @@ def test_deferred_cat_storage(layout): @pytest.mark.parametrize("fake", [False, True], ids=["eager", "fake"]) @pytest.mark.parametrize("mixed_dtype", [False, True]) -def test_deferred_cat_spec(fake, mixed_dtype): - from transformer_engine.pytorch.dynamo.deferred_cat import DeferredCat +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 = [ @@ -2899,7 +2962,7 @@ def test_deferred_cat_spec(fake, mixed_dtype): ), ] expected = torch.cat(parts) - spec = DeferredCat(parts).to_spec() + spec = ParameterParts(parts).to_spec() assert spec.shape == tuple(expected.shape) assert spec.dtype == expected.dtype assert spec.device == expected.device @@ -2908,8 +2971,9 @@ def test_deferred_cat_spec(fake, mixed_dtype): @pytest.mark.parametrize("internal", [False, True]) -def test_deferred_cat_rejects_quantized_parts(internal): - from transformer_engine.pytorch.dynamo.deferred_cat import DeferredCat +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 @@ -2917,7 +2981,8 @@ def test_deferred_cat_rejects_quantized_parts(internal): shape=(4, 8), dtype=torch.bfloat16, quantizer=quantizer, device=torch.device("cpu") ).create_tensor() with pytest.raises(TypeError, match="non-quantized tensors"): - DeferredCat([tensor, tensor]) + ParameterParts([tensor, tensor]) + assert concat_input([tensor], defer=True) is tensor @pytest.mark.skipif(not _opaque_available, reason="torch opaque object API not available") diff --git a/transformer_engine/pytorch/attention/fused_mla_q_uproj.py b/transformer_engine/pytorch/attention/fused_mla_q_uproj.py index b15e53e6ee..913aea2cca 100644 --- a/transformer_engine/pytorch/attention/fused_mla_q_uproj.py +++ b/transformer_engine/pytorch/attention/fused_mla_q_uproj.py @@ -299,8 +299,8 @@ def backward_linear( bwd_args = LinearBwdArgs( grad_output=grad_output, inputmat=x_saved, - weight_fp8=w_q, - saved_weight=w_q, + weight_for_dgrad=w_q, + original_weight=w_q, grad_output_quantizer=grad_output_quantizer, inp_shape=x_saved.shape, activation_dtype=act_dtype, diff --git a/transformer_engine/pytorch/dynamo/custom_op.py b/transformer_engine/pytorch/dynamo/custom_op.py index e3cddbd6d6..503e3b8d4c 100644 --- a/transformer_engine/pytorch/dynamo/custom_op.py +++ b/transformer_engine/pytorch/dynamo/custom_op.py @@ -37,10 +37,9 @@ * ``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 ``DeferredCat`` also accept deferred row concatenations: + 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 ``deferred_cat.restore_from_func_ctx`` in the backward - container; reconstruction of the full view happens only in the real kernel. + The adapter saves/restores the parts; only the real kernel materializes them. * ``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. @@ -72,7 +71,7 @@ * on ``backward()`` the incoming flat grads are sliced per user output from the stashed plan (a ``grad_outputs`` field on the backward args receives the whole tuple; otherwise ``grad_output`` receives the first output's grad), - the container's optional ``setup_saved_tensors`` hook restores the saved + the container's optional ``set_saved_tensors`` hook receives the restored tensors, then the *backward op* runs the real ``bwd_impl`` and returns the flat grads (``bwd_fake_impl`` is its data-free fake). @@ -110,13 +109,14 @@ from torch.utils._pytree import tree_flatten, tree_unflatten from .tensor_spec import TensorSpec, to_tensor_spec -from .deferred_cat import DeferredCat +from .parameter_parts import ParameterParts from ..quantized_tensor import ( QuantizedTensor, QuantizedTensorStorage, Quantizer, _quantized_tensor_passthrough_ops, prepare_for_saving, + restore_from_func_ctx, ) from ..utils import record_compile_disabled @@ -413,7 +413,7 @@ class _TensorOrQuantizedKind(Enum): NONE = "none" TENSOR = "tensor" STORAGE = "storage" - CONCATENATED = "concatenated" + PARTS = "parts" _TQ_KIND_KEY = "__kind__" @@ -452,8 +452,8 @@ def _is_tensor_storage_union(annot: Any) -> bool: members = frozenset(a for a in get_args(annot) if a is not type(None)) return members in ( _TQ_MEMBERS, - _TQ_MEMBERS | {DeferredCat}, - frozenset((torch.Tensor, DeferredCat)), + _TQ_MEMBERS | {ParameterParts}, + frozenset((torch.Tensor, ParameterParts)), ) @@ -497,7 +497,7 @@ def _parse_field(name: str, annot: Any) -> _FieldPlan: _SlotSpec(name + "__meta", _OPAQUE_VALUE_BUNDLE_TYPE_NAME), ) return _FieldPlan( - name, _FieldKind.TENSOR_OR_QUANTIZED, slots, DeferredCat in get_args(annot) + name, _FieldKind.TENSOR_OR_QUANTIZED, slots, ParameterParts in get_args(annot) ) stripped, is_optional = _strip_optional(annot) if stripped is torch.Tensor: @@ -539,12 +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, DeferredCat): + elif isinstance(value, ParameterParts): if not field.allows_concatenation: - raise TypeError(f"field {field.name!r} does not accept DeferredCat") + raise TypeError(f"field {field.name!r} does not accept ParameterParts") slots[tensor_slot] = None - slots[inner_slot] = value.parts - slots[meta_slot] = OpaqueValueBundle({_TQ_KIND_KEY: _TensorOrQuantizedKind.CONCATENATED}) + slots[inner_slot] = list(value.parts) + slots[meta_slot] = OpaqueValueBundle({_TQ_KIND_KEY: _TensorOrQuantizedKind.PARTS}) 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 @@ -564,7 +564,9 @@ def _pack_tensor_or_quantized(field: _FieldPlan, value: Any, slots: Dict[str, An ) -def _unpack_tensor_or_quantized(field: _FieldPlan, slots: Dict[str, Any]) -> Any: +def _unpack_tensor_or_quantized( + field: _FieldPlan, slots: Dict[str, Any], *, materialize_parts: bool = False +) -> Any: """Inverse of :func:`_pack_tensor_or_quantized`.""" tensor_slot, inner_slot, meta_slot = (s.name for s in field.slots) meta = slots[meta_slot] @@ -573,8 +575,9 @@ 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 DeferredCat(slots[inner_slot]) + if kind == _TensorOrQuantizedKind.PARTS: + parts = ParameterParts(tuple(slots[inner_slot])) + return parts.materialize() if materialize_parts else parts return _storage_unflatten(meta, slots[inner_slot]) @@ -676,7 +679,7 @@ def pack(self, obj: Any) -> Dict[str, Any]: slots[_SIMPLE_META_SLOT] = OpaqueValueBundle(simple) return slots - def unpack(self, slots: Dict[str, Any]) -> Any: + def unpack(self, slots: Dict[str, Any], *, materialize_parts: bool = False) -> Any: """Rebuild a fresh ``arg_type`` instance from the op's flat slot dict. Inverse of :meth:`pack`. @@ -688,7 +691,9 @@ def unpack(self, slots: Dict[str, Any]) -> Any: case _FieldKind.TENSOR: kwargs[field.name] = slots[field.slots[0].name] case _FieldKind.TENSOR_OR_QUANTIZED: - kwargs[field.name] = _unpack_tensor_or_quantized(field, slots) + kwargs[field.name] = _unpack_tensor_or_quantized( + field, slots, materialize_parts=materialize_parts + ) case _FieldKind.PROCESS_GROUP: if bundle is not None: name = bundle[field.name] @@ -761,7 +766,7 @@ 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 isinstance(value, DeferredCat): + 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) @@ -1006,14 +1011,8 @@ 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, DeferredCat): - setattr(obj, name, value.materialize()) + obj = plan.unpack(dict(zip(plan.slot_names, flat)), materialize_parts=True) return pack_result(impl(obj)) def _fake(*flat: Any) -> List[torch.Tensor]: @@ -1029,6 +1028,32 @@ def _fake(*flat: Any) -> List[torch.Tensor]: return op +def _flatten_saved_parameter_parts(ctx, tensors): + lengths = [len(t.parts) if isinstance(t, ParameterParts) else None for t in tensors] + ctx.parameter_parts_lengths = lengths if any(n is not None for n in lengths) else None + if ctx.parameter_parts_lengths is None: + return tensors + return [part for t in tensors for part in (t.parts if isinstance(t, ParameterParts) else (t,))] + + +def _restore_saved_tensors(ctx): + tensors = restore_from_func_ctx(ctx) + lengths = getattr(ctx, "parameter_parts_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.parameter_parts_lengths = None + return restored + + def _register_autograd_for_op( *, fwd_op: Any, @@ -1047,6 +1072,9 @@ def _register_autograd_for_op( the plan on ``ctx`` so backward can slice its grads per user output. """ bwd_takes_grad_tuple = any(f.name == "grad_outputs" for f in bwd_plan.fields) + supports_parts = any(f.allows_concatenation for f in fwd_plan.fields) + tensor_storage_offsets = frozenset(fwd_plan.tensor_or_quantized_offsets()) + set_saved_tensors = getattr(bwd_plan.arg_type, "set_saved_tensors", None) def _setup_context(ctx, inputs, output): ctx.fwd_tensor_list_lengths = { @@ -1067,15 +1095,9 @@ def _setup_context(ctx, inputs, output): out_plan.ctx_attrs, tuple(saved_list), ) - saved = [] - ctx.concatenated_saved_lengths = [] - for value in tensors_to_save_from_setup or (): - if isinstance(value, DeferredCat): - ctx.concatenated_saved_lengths.append(len(value.parts)) - saved.extend(value.parts) - else: - ctx.concatenated_saved_lengths.append(None) - saved.append(value) + saved = tensors_to_save_from_setup or () + if supports_parts: + saved = _flatten_saved_parameter_parts(ctx, saved) tensors_to_save, tensor_objects = prepare_for_saving(*saved) ctx.tensor_objects = tensor_objects ctx.save_for_backward(*tensors_to_save) @@ -1084,12 +1106,12 @@ 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 = { + ctx.parameter_parts_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() + if pos in tensor_storage_offsets and inputs[pos] is None - and inputs[pos + 2][_TQ_KIND_KEY] == _TensorOrQuantizedKind.CONCATENATED + and inputs[pos + 2][_TQ_KIND_KEY] == _TensorOrQuantizedKind.PARTS } ctx.grad_input_shapes = { pos: inputs[pos].shape for pos in grad_targets if isinstance(inputs[pos], torch.Tensor) @@ -1097,7 +1119,9 @@ def _setup_context(ctx, inputs, output): def _autograd_backward(ctx, *grad_outputs): bwd_obj = ctx.backward_objects - if hasattr(bwd_obj, "setup_saved_tensors"): + if set_saved_tensors is not None: + set_saved_tensors(bwd_obj, _restore_saved_tensors(ctx)) + elif hasattr(bwd_obj, "setup_saved_tensors"): bwd_obj.setup_saved_tensors(ctx) ctx.tensor_objects = None user_grads = _slice_user_grads(ctx.output_ranges, grad_outputs[0]) @@ -1117,8 +1141,8 @@ 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 pos in ctx.parameter_parts_grad_shapes: + shapes = ctx.parameter_parts_grad_shapes[pos] if g is not None: out[pos + 1] = list(torch.split(g, [shape[0] for shape in shapes], dim=0)) continue @@ -1128,7 +1152,7 @@ def _autograd_backward(ctx, *grad_outputs): g = g.view(shape) out[pos] = g ctx.grad_input_shapes = None - ctx.concatenated_grad_shapes = None + ctx.parameter_parts_grad_shapes = None return tuple(out) fwd_op.register_autograd(_autograd_backward, setup_context=_setup_context) @@ -1387,15 +1411,17 @@ def register_custom_op_with_autograd( non-differentiable input). * ``bwd_fake_impl(bwd_args)`` -- data-free twin of ``bwd_impl`` returning :class:`TensorSpec` grads. - * ``bwd_arg_type.setup_saved_tensors(ctx)`` -- optional hook; skipped if - absent. + * ``bwd_arg_type.set_saved_tensors(tensors)`` -- optional hook receiving + restored tensors, storage objects or parameter parts in forward save order. + Otherwise, the legacy ``setup_saved_tensors(ctx)`` hook is used if present. How the backward container is populated: ``setup_context`` fills the ``bwd_arg_type`` instance's non-tensor fields (quantizers, config) from forward state and returns the tensors to persist; the framework saves them via ``ctx.save_for_backward``. Before ``bwd_impl`` runs, the framework restores them into the container's tensor fields through the - ``setup_saved_tensors`` hook and sets the incoming gradient directly -- + ``set_saved_tensors`` (or legacy ``setup_saved_tensors``) hook and sets the + incoming gradient directly -- into a ``grad_outputs`` field (tuple, one grad per user output) if ``bwd_arg_type`` declares one, else into ``grad_output`` (the first user output's grad) -- so ``bwd_impl`` receives a fully-populated diff --git a/transformer_engine/pytorch/dynamo/deferred_cat.py b/transformer_engine/pytorch/dynamo/parameter_parts.py similarity index 72% rename from transformer_engine/pytorch/dynamo/deferred_cat.py rename to transformer_engine/pytorch/dynamo/parameter_parts.py index cae9527807..b314b2efb3 100644 --- a/transformer_engine/pytorch/dynamo/deferred_cat.py +++ b/transformer_engine/pytorch/dynamo/parameter_parts.py @@ -5,7 +5,7 @@ """Deferred row concatenation for tensors consumed inside a custom op.""" from dataclasses import dataclass -from typing import List +from typing import Tuple, TypeVar, Union import torch from torch._prims_common import make_contiguous_strides_for @@ -14,25 +14,25 @@ from ..quantized_tensor import ( QuantizedTensor, QuantizedTensorStorage, - restore_from_func_ctx as restore_tensor_ctx, ) -@dataclass -class DeferredCat: +@dataclass(frozen=True, slots=True) +class ParameterParts: """Keep every parameter visible to autograd until the opaque consumer runs.""" - parts: List[torch.Tensor] + parts: Tuple[torch.Tensor, ...] def __post_init__(self): + object.__setattr__(self, "parts", tuple(self.parts)) if not self.parts: - raise ValueError("DeferredCat requires at least one tensor") + 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("DeferredCat only supports non-quantized tensors") + raise TypeError("ParameterParts only supports non-quantized tensors") def to_spec(self) -> TensorSpec: """Describe the concatenation without accessing its storage.""" @@ -74,20 +74,5 @@ def materialize(self): return first.as_strided(shape, make_contiguous_strides_for(shape)) -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(DeferredCat(list(tensors[offset : offset + length]))) - offset += length - ctx.concatenated_saved_lengths = None - return restored +_TensorT = TypeVar("_TensorT") +ConcatInput = Union[_TensorT, ParameterParts] 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/linear.py b/transformer_engine/pytorch/module/linear.py index e9b3f1d7cd..a9691acaee 100644 --- a/transformer_engine/pytorch/module/linear.py +++ b/transformer_engine/pytorch/module/linear.py @@ -32,7 +32,7 @@ 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, @@ -83,8 +83,9 @@ QuantizedTensorStorage, Quantizer, prepare_for_saving, + restore_from_func_ctx, ) -from ..dynamo.deferred_cat import DeferredCat, restore_from_func_ctx +from ..dynamo.parameter_parts import ConcatInput, ParameterParts from ..dynamo import ( TensorSpec, TensorOrQuantized, @@ -111,9 +112,9 @@ class LinearFwdArgs: """Single-argument bag for the forward path of :class:`_Linear`.""" # --- Differentiable tensors (also passed positionally to autograd) --- - weight: Union[TensorOrQuantized, DeferredCat] + weight: ConcatInput[TensorOrQuantized] inp: torch.Tensor - bias: Optional[Union[torch.Tensor, DeferredCat]] + bias: Optional[ConcatInput[torch.Tensor]] # --- Non-differentiable cached tensors --- # TensorOrQuantized so a cached quantized workspace can cross the op boundary. @@ -180,57 +181,6 @@ class LinearFwdArgs: cpu_offloading: bool is_grad_enabled: bool - def compile_unsupported_reason(self) -> Optional[str]: - """Reason this config can't use the torch.compile custom-op path (else None).""" - if self.debug: - return "debug instrumentation (nvidia-dlfw-inspect)" - if is_distributed_weight(self.weight): - return "a DistributedWeight (custom weight parallelism, e.g. GTP)" - if isinstance(self.inp, (QuantizedTensor, QuantizedTensorStorage)): - return "a quantized input tensor" - if self.fsdp_group is not None: - return "manual TE FSDP (fsdp_group); use FSDP2 or MCore FSDP" - if ( - self.fp8_output - and self.is_grad_enabled - and (self.input_requires_grad or self.weight_requires_grad or self.bias_requires_grad) - ): - return "differentiable fp8_output=True" - if self.cpu_offloading: - return "CPU activation offloading" - if self.wgrad_store is not None: - # Non-None only when delayed wgrad compute is on (see Linear.forward). - return "delayed wgrad compute (wgrad_store)" - if ( - self.grad_input_quantizer is not None - and self.is_grad_enabled - and self.input_requires_grad - and not (self.ub_overlap_rs_dgrad or self.ub_bulk_wgrad) - ): - # A quantized dgrad can't cross the op boundary: grads are packed - # one plain Tensor[] slot each (_pack_bwd_result). - return "a quantized input grad (fp8_grad=True)" - if self.cache_weight and self.fp8: - # The cached workspace is updated in place on the first microbatch, - # which the functional op (mutates_args=()) can't express. Without - # FP8 no workspace exists, so is_first_microbatch is inert. - return "FP8 weight caching (is_first_microbatch)" - if self.fuse_wgrad_accumulation: - return "fuse_wgrad_accumulation (main_grad)" - for quantizer in ( - self.input_quantizer, - self.weight_quantizer, - self.output_quantizer, - self.grad_input_quantizer, - self.grad_weight_quantizer, - self.grad_output_quantizer, - ): - # e.g. delayed-scaling Float8Quantizer and unregistered custom-recipe - # quantizers are not value-opaque and can't cross the custom-op boundary. - if quantizer is not None and not is_value_opaque_quantizer(quantizer): - return "a quantizer not registered as a torch.compile value-opaque type" - return None - @dataclass(slots=True) class LinearBwdArgs: @@ -239,9 +189,9 @@ class LinearBwdArgs: # --- Saved / restored tensors (populated at backward entry) --- grad_output: Optional[torch.Tensor] = None inputmat: Optional[TensorOrQuantized] = None - weight_fp8: Optional[Union[TensorOrQuantized, DeferredCat]] = None - saved_weight: Optional[Union[TensorOrQuantized, DeferredCat]] = None - bias: Optional[Union[torch.Tensor, DeferredCat]] = None + weight_for_dgrad: Optional[ConcatInput[TensorOrQuantized]] = None + original_weight: Optional[ConcatInput[TensorOrQuantized]] = None + bias: Optional[ConcatInput[torch.Tensor]] = None # --- Quantizers --- input_quantizer: Optional[Quantizer] = None @@ -305,16 +255,9 @@ class LinearBwdArgs: # --- Per-backward scratch state (populated inside _linear_backward_impl) --- ub_obj_gradout: Optional[Any] = None - def setup_saved_tensors(self, ctx: torch.autograd.function.FunctionCtx) -> None: - """Pull saved tensors from ``ctx`` into the fields backward consumes.""" - ( - self.inputmat, - self.weight_fp8, - self.saved_weight, - self.bias, - ) = restore_from_func_ctx( - ctx - ) # pylint: disable=unbalanced-tuple-unpacking + def set_saved_tensors(self, tensors) -> None: + """Assign restored operands in their forward save order.""" + self.inputmat, self.weight_for_dgrad, self.original_weight, self.bias = tensors def _out_leading_from_inp(leading: int, args: Union[LinearFwdArgs, LinearBwdArgs]) -> int: @@ -745,7 +688,7 @@ def _linear_forward_impl( saved_tensor_aliases = ( "inp" if saved_inputmat is inp else None, wt_alias, - "weight", # ``saved_weight`` slot is always the weight parameter + "weight", # ``original_weight`` slot is always the weight parameter "bias" if bias is not None else None, ) tensors_to_save_from_forward = ( @@ -911,7 +854,7 @@ def _linear_forward_fake( # ------------------------------------------------------ # Backward state -- saved-tensor layout - # (saved_inputmat, wt_save, saved_weight, bias) with name-based aliasing. + # (saved_inputmat, wt_save, original_weight, bias) with name-based aliasing. # ------------------------------------------------------ tensors_to_save_from_forward = None ctx_attrs = None @@ -965,7 +908,7 @@ def _linear_forward_fake( shape=tuple(weight.shape), dtype=activation_dtype, device=weight.device ) - # Slot 2 -- ``saved_weight`` (always aliased to ``weight``). + # Slot 2 -- ``original_weight`` (always aliased to ``weight``). # Slot 3 -- ``bias`` (aliased to ``bias`` when present, else absent). saved_tensor_aliases = ( inputmat_alias, @@ -1085,8 +1028,8 @@ def _linear_setup_ctx( bwd_args.grad_weight_quantizer = None bwd_args.grad_output_quantizer = None - saved_inputmat, wt_save, saved_weight, saved_bias = tensors_to_save_from_forward - inputmat_alias, wt_save_alias, saved_weight_alias, bias_alias = ctx_attrs[ + saved_inputmat, wt_save, original_weight, saved_bias = tensors_to_save_from_forward + inputmat_alias, wt_save_alias, original_weight_alias, bias_alias = ctx_attrs[ "saved_tensor_aliases" ] bwd_args.owns_input = inputmat_alias != "inp" @@ -1098,26 +1041,26 @@ def _linear_setup_ctx( wt_save = fwd_outputs[1] elif wt_save_alias == "weight_workspace": wt_save = fwd_args.weight_workspace - if saved_weight_alias == "weight": - saved_weight = weight + if original_weight_alias == "weight": + original_weight = weight if bias_alias == "bias": saved_bias = bias - return (saved_inputmat, wt_save, saved_weight, saved_bias) + return (saved_inputmat, wt_save, original_weight, saved_bias) def _linear_backward_impl(args: LinearBwdArgs) -> Tuple[Union[torch.Tensor, None], ...]: """Backward implementation for the linear layer. Caller must have populated ``args.grad_output`` and run - ``args.setup_saved_tensors(ctx)`` before invocation. + ``args.set_saved_tensors(tensors)`` before invocation. """ bwd_args = args grad_output = args.grad_output assert grad_output is not None inputmat = args.inputmat - weight_fp8 = args.weight_fp8 - saved_weight = args.saved_weight - is_dist_weight = is_distributed_weight(saved_weight) + weight_for_dgrad = args.weight_for_dgrad + original_weight = args.original_weight + is_dist_weight = is_distributed_weight(original_weight) bias = args.bias input_quantizer = args.input_quantizer weight_quantizer = args.weight_quantizer @@ -1166,20 +1109,20 @@ def _linear_backward_impl(args: LinearBwdArgs) -> Tuple[Union[torch.Tensor, None origin_weight_python_object.main_grad = main_grad # Gather intermediate/activation tensors if needed - # NOTE: weight_fp8 = weight when bwd_args.fp8 == False and torch.disttributed.FSDP already + # NOTE: weight_for_dgrad = weight when bwd_args.fp8 == False and torch.disttributed.FSDP already # shards/unshards the base weights so we don't do it ourselves nvtx_range_push(f"{nvtx_label}.fsdp_gather") _fsdp_gather_tensors( bwd_args.fsdp_group, bwd_args.fsdp_shapes, inputmat, - weight_fp8, + weight_for_dgrad, ) nvtx_range_pop(f"{nvtx_label}.fsdp_gather") # Reconstruct inp_shape when not stored (compiled mode with dynamic shapes). if bwd_args.inp_shape is None: - in_features = saved_weight.shape[-1] + in_features = original_weight.shape[-1] inp_leading = _inp_leading_from_out(grad_output.shape[0], bwd_args) bwd_args.inp_shape = ( torch.Size([in_features]) @@ -1341,32 +1284,32 @@ def _linear_backward_impl(args: LinearBwdArgs) -> Tuple[Union[torch.Tensor, None # Distributed weight (e.g. GTP): re-gather the sharded weight; runs even when # requires_dgrad=False so the prev_w prefetch is issued for the next layer's bwd. if is_dist_weight: - weight_fp8 = materialize_weight_for_backward(saved_weight)[0] + weight_for_dgrad = materialize_weight_for_backward(original_weight)[0] if bwd_args.requires_dgrad: # FSDP2: Re-create workspace from all-gathered weight when # workspace was not saved. (Issue #2681) - # Use saved_weight (the original weight parameter) since - # weight_fp8 is only set when workspace was saved. - if weight_fp8 is None: - if isinstance(saved_weight, QuantizedTensorStorage): + # Use original_weight (the original weight parameter) since + # weight_for_dgrad is only set when workspace was saved. + if weight_for_dgrad is None: + if isinstance(original_weight, QuantizedTensorStorage): # saved weight is already set to right usages by # fsdp2 quantized-tensor hooks when workspace was not saved. - weight_fp8 = saved_weight + weight_for_dgrad = original_weight elif bwd_args.weight_quantizer is not None: bwd_args.weight_quantizer.set_usage(rowwise=True, columnwise=True) - weight_fp8 = bwd_args.weight_quantizer(saved_weight) + weight_for_dgrad = bwd_args.weight_quantizer(original_weight) elif ( is_dist_weight and bwd_args.fp8 and bwd_args.weight_quantizer is not None - and not isinstance(weight_fp8, QuantizedTensorStorage) + and not isinstance(weight_for_dgrad, QuantizedTensorStorage) ): # Distributed weight re-gathered a BF16 weight: quantize with the layer quantizer # so the dgrad operand isn't cast by the delayed recipe. bwd_args.weight_quantizer.set_usage(rowwise=True, columnwise=True) - weight_fp8 = bwd_args.weight_quantizer(weight_fp8) + weight_for_dgrad = bwd_args.weight_quantizer(weight_for_dgrad) # Make sure required data is available if isinstance(grad_output, QuantizedTensorStorage): @@ -1374,9 +1317,9 @@ def _linear_backward_impl(args: LinearBwdArgs) -> Tuple[Union[torch.Tensor, None if ( bwd_args.fp8 and weight_quantizer is not None - and isinstance(weight_fp8, QuantizedTensorStorage) + and isinstance(weight_for_dgrad, QuantizedTensorStorage) ): - weight_fp8.update_usage(columnwise_usage=True) + weight_for_dgrad.update_usage(columnwise_usage=True) # Choose whether to use GEMM kernel with split accumulator use_split_accumulator = bwd_args.dgrad_use_split_accumulator @@ -1403,18 +1346,18 @@ def _linear_backward_impl(args: LinearBwdArgs) -> Tuple[Union[torch.Tensor, None # Note: dx = dy * w nvtx_range_push(f"{nvtx_label}.dgrad_gemm") - weight_for_dgrad = weight_fp8 + gemm_weight = weight_for_dgrad if bwd_args.backward_override == "dequantized": - if isinstance(weight_for_dgrad, QuantizedTensorStorage): - weight_for_dgrad = weight_for_dgrad.dequantize(dtype=bwd_args.activation_dtype) + if isinstance(gemm_weight, QuantizedTensorStorage): + gemm_weight = gemm_weight.dequantize(dtype=bwd_args.activation_dtype) else: - weight_for_dgrad = cast_if_needed(weight_for_dgrad, bwd_args.activation_dtype) + gemm_weight = cast_if_needed(gemm_weight, bwd_args.activation_dtype) elif bwd_args.backward_override == "high_precision": - weight_for_dgrad = saved_weight - if isinstance(weight_for_dgrad, QuantizedTensorStorage): - weight_for_dgrad = weight_for_dgrad.dequantize(dtype=bwd_args.activation_dtype) + gemm_weight = original_weight + if isinstance(gemm_weight, QuantizedTensorStorage): + gemm_weight = gemm_weight.dequantize(dtype=bwd_args.activation_dtype) gemm_out, *_, reduce_scatter_out = general_gemm( - weight_for_dgrad, + gemm_weight, grad_output, layout="NN", grad=True, @@ -1434,8 +1377,8 @@ def _linear_backward_impl(args: LinearBwdArgs) -> Tuple[Union[torch.Tensor, None # and 2d block-scaled weights in TE managed memory. So we need to clear # it here. # (Issues #2681, #2717) - if bwd_args.is_fsdp2 and isinstance(weight_fp8, QuantizedTensorStorage): - clear_columnwise_cache(weight_fp8) + if bwd_args.is_fsdp2 and isinstance(weight_for_dgrad, QuantizedTensorStorage): + clear_columnwise_cache(weight_for_dgrad) # Prepare grad input tensor # Note: Perform tensor-parallel communication @@ -1644,7 +1587,7 @@ def wgrad_gemm( # Distributed weight (e.g. GTP): reduce-scatter the freshly computed wgrad # (async; overlap with the next layer's bwd via the cascade). if is_dist_weight: - wgrad = finalize_weight_grads(saved_weight, [wgrad])[0] + wgrad = finalize_weight_grads(original_weight, [wgrad])[0] # Update grad bias if needed if grad_bias is None: @@ -1716,7 +1659,7 @@ def wgrad_gemm( # Scatter fp8 weight buffers if bwd_args.fp8 and not bwd_args.is_weight_param_quantized: - _fsdp_scatter_tensors(bwd_args.fsdp_group, weight_fp8) + _fsdp_scatter_tensors(bwd_args.fsdp_group, weight_for_dgrad) return ( wgrad, dgrad if bwd_args.requires_dgrad else None, @@ -1738,7 +1681,7 @@ def _linear_backward_fake( "(fsdp_group is not None); use FSDP2 or MCore FSDP." ) - weight = args.saved_weight + weight = args.original_weight out_dtype = args.activation_dtype out_features, in_features = weight.shape @@ -1866,7 +1809,7 @@ def backward( """Backward pass: compute gradients and reduce FP8 scaling factors.""" bwd_args: LinearBwdArgs = ctx.backward_objects bwd_args.grad_output = grad_output - bwd_args.setup_saved_tensors(ctx) + bwd_args.set_saved_tensors(restore_from_func_ctx(ctx)) nvtx_label = "transformer_engine._Linear.backward" if bwd_args.ub_name is not None: nvtx_label = f"{nvtx_label}.{bwd_args.ub_name}" @@ -2388,7 +2331,10 @@ def forward( if get_ub_is_fp8(self.ub_name + "_dgrad", FP8GlobalStateManager.is_fp8_enabled()): fp8_grad = True - if torch.compiler.is_compiling() and _linear_op is not None: + use_compiled_op = torch.compiler.is_compiling() and _linear_op is not None + if _linear_op is None and torch.compiler.is_compiling(): + warn_if_compile_disabled() + if use_compiled_op: reason = self._compile_eager_fallback_reason( inp, is_first_microbatch, fp8_output, fp8_grad, is_grad_enabled, debug ) @@ -2400,16 +2346,6 @@ def forward( inp = self.prepare_forward(inp, allow_non_contiguous=isinstance(inp, QuantizedTensor)) try: - 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, DeferredCat) else weight_tensor - ) - bias_meta = ( - bias_tensor.to_spec() if isinstance(bias_tensor, DeferredCat) else bias_tensor - ) - quantizers = ( self._get_quantizers(fp8_output, fp8_grad, is_grad_enabled) if not debug @@ -2420,6 +2356,26 @@ def forward( debug = False quantizers = self._get_quantizers(fp8_output, fp8_grad, is_grad_enabled) + if use_compiled_op and any( + q is not None and not is_value_opaque_quantizer(q) for q in quantizers + ): + reason = "a quantizer not registered as a torch.compile value-opaque type" + warn_compile_eager_fallback(reason) + torch._dynamo.graph_break(msg=f"te.Linear falling back to eager: {reason}") + use_compiled_op = False + + weight_tensor, bias_tensor = self._get_weight_and_bias_tensors( + defer_concatenation=use_compiled_op + ) + 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 + ) + ( input_quantizer, weight_quantizer, @@ -2433,9 +2389,6 @@ def forward( weight_quantizer, weight_meta ) - use_compiled_op = torch.compiler.is_compiling() and _linear_op is not None - if _linear_op is None and torch.compiler.is_compiling(): - warn_if_compile_disabled() if use_compiled_op: # Process groups cross the op boundary separately from quantizers. for quantizer in (input_quantizer, grad_output_quantizer): @@ -2546,22 +2499,6 @@ def forward( is_grad_enabled=is_grad_enabled, ) - if use_compiled_op: - # Safety net for quantizer-dependent conditions only. - fallback_reason = fwd_args.compile_unsupported_reason() - if fallback_reason is not None: - warn_compile_eager_fallback(fallback_reason) - torch._dynamo.graph_break( - 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.shape, weight_meta.shape, self.fp8) out, new_weight_workspace = _linear_op(fwd_args) @@ -2654,12 +2591,10 @@ def _compile_eager_fallback_reason( is_grad_enabled: bool, debug: bool, ) -> Optional[str]: - """Why this call can't use the compiled op (else None), decided before - prepare_forward. Quantizer checks stay in compile_unsupported_reason.""" + """Check configuration before prepare_forward, without concatenating parameters.""" 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" @@ -2667,8 +2602,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" @@ -2709,29 +2646,14 @@ def _forward_eager_fallback( ) def _get_weight_and_bias_tensors(self, defer_concatenation=False): - weights = self._get_weight_tensors() - weight_tensor = ( - DeferredCat(weights) - if defer_concatenation - and len(weights) > 1 - and all(not isinstance(w, (QuantizedTensor, QuantizedTensorStorage)) for w in weights) - else noop_cat(weights) - ) - bias_tensor = None + weight = concat_input(self._get_weight_tensors(), defer=defer_concatenation) + bias = None if self.use_bias: - biases = [getattr(self, name) for name in self.bias_names] - bias_tensor = ( - DeferredCat(biases) - if defer_concatenation - and len(biases) > 1 - and self.apply_bias - and not self.gemm_bias_unfused_add - and all( - not isinstance(b, (QuantizedTensor, QuantizedTensorStorage)) for b in biases - ) - else noop_cat(biases) + bias = concat_input( + [getattr(self, name) for name in self.bias_names], + defer=defer_concatenation and self.apply_bias and not self.gemm_bias_unfused_add, ) - return weight_tensor, bias_tensor + return weight, bias def onnx_forward( self, From 5d5c1802c10e281a2b64b1be42a65e4631f7cb77 Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Mon, 5 Oct 2026 16:56:29 +0200 Subject: [PATCH 4/4] [PyTorch] Limit the split-parameter API refactor Restore compile_unsupported_reason, setup_saved_tensors(ctx), the original backward operand names, and the existing adapter contract. Remove the unrelated MLA changes. Retain the parameter-parts alias and helper, split-parameter correctness fixes, and focused regression coverage. Signed-off-by: Pawel Gadzinski --- tests/pytorch/test_torch_compile.py | 2 +- .../pytorch/attention/fused_mla_q_uproj.py | 4 +- .../pytorch/dynamo/custom_op.py | 98 +++----- .../pytorch/dynamo/parameter_parts.py | 20 ++ transformer_engine/pytorch/module/linear.py | 220 ++++++++++++------ 5 files changed, 203 insertions(+), 141 deletions(-) diff --git a/tests/pytorch/test_torch_compile.py b/tests/pytorch/test_torch_compile.py index ba0432b68c..19eb4e245e 100644 --- a/tests/pytorch/test_torch_compile.py +++ b/tests/pytorch/test_torch_compile.py @@ -2048,7 +2048,7 @@ def fn(inp): ) -# Configs rejected by Linear._compile_eager_fallback_reason() that a +# Configs rejected by LinearFwdArgs.compile_unsupported_reason() that a # single-GPU unit test can construct. Distributed-only reasons (fsdp_group, # DistributedWeight) and CPU offloading need machinery this file doesn't have; # delayed scaling is a hard error (check_recipe_support), tested separately. diff --git a/transformer_engine/pytorch/attention/fused_mla_q_uproj.py b/transformer_engine/pytorch/attention/fused_mla_q_uproj.py index 913aea2cca..b15e53e6ee 100644 --- a/transformer_engine/pytorch/attention/fused_mla_q_uproj.py +++ b/transformer_engine/pytorch/attention/fused_mla_q_uproj.py @@ -299,8 +299,8 @@ def backward_linear( bwd_args = LinearBwdArgs( grad_output=grad_output, inputmat=x_saved, - weight_for_dgrad=w_q, - original_weight=w_q, + weight_fp8=w_q, + saved_weight=w_q, grad_output_quantizer=grad_output_quantizer, inp_shape=x_saved.shape, activation_dtype=act_dtype, diff --git a/transformer_engine/pytorch/dynamo/custom_op.py b/transformer_engine/pytorch/dynamo/custom_op.py index 503e3b8d4c..a81ebccae8 100644 --- a/transformer_engine/pytorch/dynamo/custom_op.py +++ b/transformer_engine/pytorch/dynamo/custom_op.py @@ -39,7 +39,8 @@ 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. - The adapter saves/restores the parts; only the real kernel materializes them. + 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. @@ -71,7 +72,7 @@ * on ``backward()`` the incoming flat grads are sliced per user output from the stashed plan (a ``grad_outputs`` field on the backward args receives the whole tuple; otherwise ``grad_output`` receives the first output's grad), - the container's optional ``set_saved_tensors`` hook receives the restored + the container's optional ``setup_saved_tensors`` hook restores the saved tensors, then the *backward op* runs the real ``bwd_impl`` and returns the flat grads (``bwd_fake_impl`` is its data-free fake). @@ -116,7 +117,6 @@ Quantizer, _quantized_tensor_passthrough_ops, prepare_for_saving, - restore_from_func_ctx, ) from ..utils import record_compile_disabled @@ -413,7 +413,7 @@ class _TensorOrQuantizedKind(Enum): NONE = "none" TENSOR = "tensor" STORAGE = "storage" - PARTS = "parts" + CONCATENATED = "concatenated" _TQ_KIND_KEY = "__kind__" @@ -544,7 +544,7 @@ def _pack_tensor_or_quantized(field: _FieldPlan, value: Any, slots: Dict[str, An 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.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 @@ -564,9 +564,7 @@ def _pack_tensor_or_quantized(field: _FieldPlan, value: Any, slots: Dict[str, An ) -def _unpack_tensor_or_quantized( - field: _FieldPlan, slots: Dict[str, Any], *, materialize_parts: bool = False -) -> Any: +def _unpack_tensor_or_quantized(field: _FieldPlan, slots: Dict[str, Any]) -> Any: """Inverse of :func:`_pack_tensor_or_quantized`.""" tensor_slot, inner_slot, meta_slot = (s.name for s in field.slots) meta = slots[meta_slot] @@ -575,9 +573,8 @@ def _unpack_tensor_or_quantized( return None if kind == _TensorOrQuantizedKind.TENSOR: return slots[tensor_slot] - if kind == _TensorOrQuantizedKind.PARTS: - parts = ParameterParts(tuple(slots[inner_slot])) - return parts.materialize() if materialize_parts else parts + if kind == _TensorOrQuantizedKind.CONCATENATED: + return ParameterParts(tuple(slots[inner_slot])) return _storage_unflatten(meta, slots[inner_slot]) @@ -679,7 +676,7 @@ def pack(self, obj: Any) -> Dict[str, Any]: slots[_SIMPLE_META_SLOT] = OpaqueValueBundle(simple) return slots - def unpack(self, slots: Dict[str, Any], *, materialize_parts: bool = False) -> Any: + def unpack(self, slots: Dict[str, Any]) -> Any: """Rebuild a fresh ``arg_type`` instance from the op's flat slot dict. Inverse of :meth:`pack`. @@ -691,9 +688,7 @@ def unpack(self, slots: Dict[str, Any], *, materialize_parts: bool = False) -> A case _FieldKind.TENSOR: kwargs[field.name] = slots[field.slots[0].name] case _FieldKind.TENSOR_OR_QUANTIZED: - kwargs[field.name] = _unpack_tensor_or_quantized( - field, slots, materialize_parts=materialize_parts - ) + kwargs[field.name] = _unpack_tensor_or_quantized(field, slots) case _FieldKind.PROCESS_GROUP: if bundle is not None: name = bundle[field.name] @@ -1011,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)), materialize_parts=True) + 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]: @@ -1028,32 +1029,6 @@ def _fake(*flat: Any) -> List[torch.Tensor]: return op -def _flatten_saved_parameter_parts(ctx, tensors): - lengths = [len(t.parts) if isinstance(t, ParameterParts) else None for t in tensors] - ctx.parameter_parts_lengths = lengths if any(n is not None for n in lengths) else None - if ctx.parameter_parts_lengths is None: - return tensors - return [part for t in tensors for part in (t.parts if isinstance(t, ParameterParts) else (t,))] - - -def _restore_saved_tensors(ctx): - tensors = restore_from_func_ctx(ctx) - lengths = getattr(ctx, "parameter_parts_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.parameter_parts_lengths = None - return restored - - def _register_autograd_for_op( *, fwd_op: Any, @@ -1072,9 +1047,6 @@ def _register_autograd_for_op( the plan on ``ctx`` so backward can slice its grads per user output. """ bwd_takes_grad_tuple = any(f.name == "grad_outputs" for f in bwd_plan.fields) - supports_parts = any(f.allows_concatenation for f in fwd_plan.fields) - tensor_storage_offsets = frozenset(fwd_plan.tensor_or_quantized_offsets()) - set_saved_tensors = getattr(bwd_plan.arg_type, "set_saved_tensors", None) def _setup_context(ctx, inputs, output): ctx.fwd_tensor_list_lengths = { @@ -1095,9 +1067,15 @@ def _setup_context(ctx, inputs, output): out_plan.ctx_attrs, tuple(saved_list), ) - saved = tensors_to_save_from_setup or () - if supports_parts: - saved = _flatten_saved_parameter_parts(ctx, saved) + 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) @@ -1106,12 +1084,12 @@ 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.parameter_parts_grad_shapes = { + ctx.concatenated_grad_shapes = { pos: [tensor.shape for tensor in inputs[pos + 1]] for pos in grad_targets - if pos in tensor_storage_offsets + if pos in fwd_plan.tensor_or_quantized_offsets() and inputs[pos] is None - and inputs[pos + 2][_TQ_KIND_KEY] == _TensorOrQuantizedKind.PARTS + 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) @@ -1119,9 +1097,7 @@ def _setup_context(ctx, inputs, output): def _autograd_backward(ctx, *grad_outputs): bwd_obj = ctx.backward_objects - if set_saved_tensors is not None: - set_saved_tensors(bwd_obj, _restore_saved_tensors(ctx)) - elif hasattr(bwd_obj, "setup_saved_tensors"): + if hasattr(bwd_obj, "setup_saved_tensors"): bwd_obj.setup_saved_tensors(ctx) ctx.tensor_objects = None user_grads = _slice_user_grads(ctx.output_ranges, grad_outputs[0]) @@ -1141,8 +1117,8 @@ 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.parameter_parts_grad_shapes: - shapes = ctx.parameter_parts_grad_shapes[pos] + 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 @@ -1152,7 +1128,7 @@ def _autograd_backward(ctx, *grad_outputs): g = g.view(shape) out[pos] = g ctx.grad_input_shapes = None - ctx.parameter_parts_grad_shapes = None + ctx.concatenated_grad_shapes = None return tuple(out) fwd_op.register_autograd(_autograd_backward, setup_context=_setup_context) @@ -1411,17 +1387,15 @@ def register_custom_op_with_autograd( non-differentiable input). * ``bwd_fake_impl(bwd_args)`` -- data-free twin of ``bwd_impl`` returning :class:`TensorSpec` grads. - * ``bwd_arg_type.set_saved_tensors(tensors)`` -- optional hook receiving - restored tensors, storage objects or parameter parts in forward save order. - Otherwise, the legacy ``setup_saved_tensors(ctx)`` hook is used if present. + * ``bwd_arg_type.setup_saved_tensors(ctx)`` -- optional hook; skipped if + absent. How the backward container is populated: ``setup_context`` fills the ``bwd_arg_type`` instance's non-tensor fields (quantizers, config) from forward state and returns the tensors to persist; the framework saves them via ``ctx.save_for_backward``. Before ``bwd_impl`` runs, the framework restores them into the container's tensor fields through the - ``set_saved_tensors`` (or legacy ``setup_saved_tensors``) hook and sets the - incoming gradient directly -- + ``setup_saved_tensors`` hook and sets the incoming gradient directly -- into a ``grad_outputs`` field (tuple, one grad per user output) if ``bwd_arg_type`` declares one, else into ``grad_output`` (the first user output's grad) -- so ``bwd_impl`` receives a fully-populated diff --git a/transformer_engine/pytorch/dynamo/parameter_parts.py b/transformer_engine/pytorch/dynamo/parameter_parts.py index b314b2efb3..c0ad545d78 100644 --- a/transformer_engine/pytorch/dynamo/parameter_parts.py +++ b/transformer_engine/pytorch/dynamo/parameter_parts.py @@ -14,6 +14,7 @@ from ..quantized_tensor import ( QuantizedTensor, QuantizedTensorStorage, + restore_from_func_ctx as restore_tensor_ctx, ) @@ -76,3 +77,22 @@ def materialize(self): _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/linear.py b/transformer_engine/pytorch/module/linear.py index a9691acaee..8f66102c40 100644 --- a/transformer_engine/pytorch/module/linear.py +++ b/transformer_engine/pytorch/module/linear.py @@ -83,9 +83,8 @@ QuantizedTensorStorage, Quantizer, prepare_for_saving, - restore_from_func_ctx, ) -from ..dynamo.parameter_parts import ConcatInput, ParameterParts +from ..dynamo.parameter_parts import ConcatInput, ParameterParts, restore_from_func_ctx from ..dynamo import ( TensorSpec, TensorOrQuantized, @@ -181,6 +180,57 @@ class LinearFwdArgs: cpu_offloading: bool is_grad_enabled: bool + def compile_unsupported_reason(self) -> Optional[str]: + """Reason this config can't use the torch.compile custom-op path (else None).""" + if self.debug: + return "debug instrumentation (nvidia-dlfw-inspect)" + if is_distributed_weight(self.weight): + return "a DistributedWeight (custom weight parallelism, e.g. GTP)" + if isinstance(self.inp, (QuantizedTensor, QuantizedTensorStorage)): + return "a quantized input tensor" + if self.fsdp_group is not None: + return "manual TE FSDP (fsdp_group); use FSDP2 or MCore FSDP" + if ( + self.fp8_output + and self.is_grad_enabled + and (self.input_requires_grad or self.weight_requires_grad or self.bias_requires_grad) + ): + return "differentiable fp8_output=True" + if self.cpu_offloading: + return "CPU activation offloading" + if self.wgrad_store is not None: + # Non-None only when delayed wgrad compute is on (see Linear.forward). + return "delayed wgrad compute (wgrad_store)" + if ( + self.grad_input_quantizer is not None + and self.is_grad_enabled + and self.input_requires_grad + and not (self.ub_overlap_rs_dgrad or self.ub_bulk_wgrad) + ): + # A quantized dgrad can't cross the op boundary: grads are packed + # one plain Tensor[] slot each (_pack_bwd_result). + return "a quantized input grad (fp8_grad=True)" + if self.cache_weight and self.fp8: + # The cached workspace is updated in place on the first microbatch, + # which the functional op (mutates_args=()) can't express. Without + # FP8 no workspace exists, so is_first_microbatch is inert. + return "FP8 weight caching (is_first_microbatch)" + if self.fuse_wgrad_accumulation: + return "fuse_wgrad_accumulation (main_grad)" + for quantizer in ( + self.input_quantizer, + self.weight_quantizer, + self.output_quantizer, + self.grad_input_quantizer, + self.grad_weight_quantizer, + self.grad_output_quantizer, + ): + # e.g. delayed-scaling Float8Quantizer and unregistered custom-recipe + # quantizers are not value-opaque and can't cross the custom-op boundary. + if quantizer is not None and not is_value_opaque_quantizer(quantizer): + return "a quantizer not registered as a torch.compile value-opaque type" + return None + @dataclass(slots=True) class LinearBwdArgs: @@ -189,8 +239,8 @@ class LinearBwdArgs: # --- Saved / restored tensors (populated at backward entry) --- grad_output: Optional[torch.Tensor] = None inputmat: Optional[TensorOrQuantized] = None - weight_for_dgrad: Optional[ConcatInput[TensorOrQuantized]] = None - original_weight: Optional[ConcatInput[TensorOrQuantized]] = None + weight_fp8: Optional[ConcatInput[TensorOrQuantized]] = None + saved_weight: Optional[ConcatInput[TensorOrQuantized]] = None bias: Optional[ConcatInput[torch.Tensor]] = None # --- Quantizers --- @@ -255,9 +305,16 @@ class LinearBwdArgs: # --- Per-backward scratch state (populated inside _linear_backward_impl) --- ub_obj_gradout: Optional[Any] = None - def set_saved_tensors(self, tensors) -> None: - """Assign restored operands in their forward save order.""" - self.inputmat, self.weight_for_dgrad, self.original_weight, self.bias = tensors + def setup_saved_tensors(self, ctx: torch.autograd.function.FunctionCtx) -> None: + """Pull saved tensors from ``ctx`` into the fields backward consumes.""" + ( + self.inputmat, + self.weight_fp8, + self.saved_weight, + self.bias, + ) = restore_from_func_ctx( + ctx + ) # pylint: disable=unbalanced-tuple-unpacking def _out_leading_from_inp(leading: int, args: Union[LinearFwdArgs, LinearBwdArgs]) -> int: @@ -688,7 +745,7 @@ def _linear_forward_impl( saved_tensor_aliases = ( "inp" if saved_inputmat is inp else None, wt_alias, - "weight", # ``original_weight`` slot is always the weight parameter + "weight", # ``saved_weight`` slot is always the weight parameter "bias" if bias is not None else None, ) tensors_to_save_from_forward = ( @@ -854,7 +911,7 @@ def _linear_forward_fake( # ------------------------------------------------------ # Backward state -- saved-tensor layout - # (saved_inputmat, wt_save, original_weight, bias) with name-based aliasing. + # (saved_inputmat, wt_save, saved_weight, bias) with name-based aliasing. # ------------------------------------------------------ tensors_to_save_from_forward = None ctx_attrs = None @@ -908,7 +965,7 @@ def _linear_forward_fake( shape=tuple(weight.shape), dtype=activation_dtype, device=weight.device ) - # Slot 2 -- ``original_weight`` (always aliased to ``weight``). + # Slot 2 -- ``saved_weight`` (always aliased to ``weight``). # Slot 3 -- ``bias`` (aliased to ``bias`` when present, else absent). saved_tensor_aliases = ( inputmat_alias, @@ -1028,8 +1085,8 @@ def _linear_setup_ctx( bwd_args.grad_weight_quantizer = None bwd_args.grad_output_quantizer = None - saved_inputmat, wt_save, original_weight, saved_bias = tensors_to_save_from_forward - inputmat_alias, wt_save_alias, original_weight_alias, bias_alias = ctx_attrs[ + saved_inputmat, wt_save, saved_weight, saved_bias = tensors_to_save_from_forward + inputmat_alias, wt_save_alias, saved_weight_alias, bias_alias = ctx_attrs[ "saved_tensor_aliases" ] bwd_args.owns_input = inputmat_alias != "inp" @@ -1041,26 +1098,26 @@ def _linear_setup_ctx( wt_save = fwd_outputs[1] elif wt_save_alias == "weight_workspace": wt_save = fwd_args.weight_workspace - if original_weight_alias == "weight": - original_weight = weight + if saved_weight_alias == "weight": + saved_weight = weight if bias_alias == "bias": saved_bias = bias - return (saved_inputmat, wt_save, original_weight, saved_bias) + return (saved_inputmat, wt_save, saved_weight, saved_bias) def _linear_backward_impl(args: LinearBwdArgs) -> Tuple[Union[torch.Tensor, None], ...]: """Backward implementation for the linear layer. Caller must have populated ``args.grad_output`` and run - ``args.set_saved_tensors(tensors)`` before invocation. + ``args.setup_saved_tensors(ctx)`` before invocation. """ bwd_args = args grad_output = args.grad_output assert grad_output is not None inputmat = args.inputmat - weight_for_dgrad = args.weight_for_dgrad - original_weight = args.original_weight - is_dist_weight = is_distributed_weight(original_weight) + weight_fp8 = args.weight_fp8 + saved_weight = args.saved_weight + is_dist_weight = is_distributed_weight(saved_weight) bias = args.bias input_quantizer = args.input_quantizer weight_quantizer = args.weight_quantizer @@ -1109,20 +1166,20 @@ def _linear_backward_impl(args: LinearBwdArgs) -> Tuple[Union[torch.Tensor, None origin_weight_python_object.main_grad = main_grad # Gather intermediate/activation tensors if needed - # NOTE: weight_for_dgrad = weight when bwd_args.fp8 == False and torch.disttributed.FSDP already + # NOTE: weight_fp8 = weight when bwd_args.fp8 == False and torch.disttributed.FSDP already # shards/unshards the base weights so we don't do it ourselves nvtx_range_push(f"{nvtx_label}.fsdp_gather") _fsdp_gather_tensors( bwd_args.fsdp_group, bwd_args.fsdp_shapes, inputmat, - weight_for_dgrad, + weight_fp8, ) nvtx_range_pop(f"{nvtx_label}.fsdp_gather") # Reconstruct inp_shape when not stored (compiled mode with dynamic shapes). if bwd_args.inp_shape is None: - in_features = original_weight.shape[-1] + in_features = saved_weight.shape[-1] inp_leading = _inp_leading_from_out(grad_output.shape[0], bwd_args) bwd_args.inp_shape = ( torch.Size([in_features]) @@ -1284,32 +1341,32 @@ def _linear_backward_impl(args: LinearBwdArgs) -> Tuple[Union[torch.Tensor, None # Distributed weight (e.g. GTP): re-gather the sharded weight; runs even when # requires_dgrad=False so the prev_w prefetch is issued for the next layer's bwd. if is_dist_weight: - weight_for_dgrad = materialize_weight_for_backward(original_weight)[0] + weight_fp8 = materialize_weight_for_backward(saved_weight)[0] if bwd_args.requires_dgrad: # FSDP2: Re-create workspace from all-gathered weight when # workspace was not saved. (Issue #2681) - # Use original_weight (the original weight parameter) since - # weight_for_dgrad is only set when workspace was saved. - if weight_for_dgrad is None: - if isinstance(original_weight, QuantizedTensorStorage): + # Use saved_weight (the original weight parameter) since + # weight_fp8 is only set when workspace was saved. + if weight_fp8 is None: + if isinstance(saved_weight, QuantizedTensorStorage): # saved weight is already set to right usages by # fsdp2 quantized-tensor hooks when workspace was not saved. - weight_for_dgrad = original_weight + weight_fp8 = saved_weight elif bwd_args.weight_quantizer is not None: bwd_args.weight_quantizer.set_usage(rowwise=True, columnwise=True) - weight_for_dgrad = bwd_args.weight_quantizer(original_weight) + weight_fp8 = bwd_args.weight_quantizer(saved_weight) elif ( is_dist_weight and bwd_args.fp8 and bwd_args.weight_quantizer is not None - and not isinstance(weight_for_dgrad, QuantizedTensorStorage) + and not isinstance(weight_fp8, QuantizedTensorStorage) ): # Distributed weight re-gathered a BF16 weight: quantize with the layer quantizer # so the dgrad operand isn't cast by the delayed recipe. bwd_args.weight_quantizer.set_usage(rowwise=True, columnwise=True) - weight_for_dgrad = bwd_args.weight_quantizer(weight_for_dgrad) + weight_fp8 = bwd_args.weight_quantizer(weight_fp8) # Make sure required data is available if isinstance(grad_output, QuantizedTensorStorage): @@ -1317,9 +1374,9 @@ def _linear_backward_impl(args: LinearBwdArgs) -> Tuple[Union[torch.Tensor, None if ( bwd_args.fp8 and weight_quantizer is not None - and isinstance(weight_for_dgrad, QuantizedTensorStorage) + and isinstance(weight_fp8, QuantizedTensorStorage) ): - weight_for_dgrad.update_usage(columnwise_usage=True) + weight_fp8.update_usage(columnwise_usage=True) # Choose whether to use GEMM kernel with split accumulator use_split_accumulator = bwd_args.dgrad_use_split_accumulator @@ -1346,18 +1403,18 @@ def _linear_backward_impl(args: LinearBwdArgs) -> Tuple[Union[torch.Tensor, None # Note: dx = dy * w nvtx_range_push(f"{nvtx_label}.dgrad_gemm") - gemm_weight = weight_for_dgrad + weight_for_dgrad = weight_fp8 if bwd_args.backward_override == "dequantized": - if isinstance(gemm_weight, QuantizedTensorStorage): - gemm_weight = gemm_weight.dequantize(dtype=bwd_args.activation_dtype) + if isinstance(weight_for_dgrad, QuantizedTensorStorage): + weight_for_dgrad = weight_for_dgrad.dequantize(dtype=bwd_args.activation_dtype) else: - gemm_weight = cast_if_needed(gemm_weight, bwd_args.activation_dtype) + weight_for_dgrad = cast_if_needed(weight_for_dgrad, bwd_args.activation_dtype) elif bwd_args.backward_override == "high_precision": - gemm_weight = original_weight - if isinstance(gemm_weight, QuantizedTensorStorage): - gemm_weight = gemm_weight.dequantize(dtype=bwd_args.activation_dtype) + weight_for_dgrad = saved_weight + if isinstance(weight_for_dgrad, QuantizedTensorStorage): + weight_for_dgrad = weight_for_dgrad.dequantize(dtype=bwd_args.activation_dtype) gemm_out, *_, reduce_scatter_out = general_gemm( - gemm_weight, + weight_for_dgrad, grad_output, layout="NN", grad=True, @@ -1377,8 +1434,8 @@ def _linear_backward_impl(args: LinearBwdArgs) -> Tuple[Union[torch.Tensor, None # and 2d block-scaled weights in TE managed memory. So we need to clear # it here. # (Issues #2681, #2717) - if bwd_args.is_fsdp2 and isinstance(weight_for_dgrad, QuantizedTensorStorage): - clear_columnwise_cache(weight_for_dgrad) + if bwd_args.is_fsdp2 and isinstance(weight_fp8, QuantizedTensorStorage): + clear_columnwise_cache(weight_fp8) # Prepare grad input tensor # Note: Perform tensor-parallel communication @@ -1587,7 +1644,7 @@ def wgrad_gemm( # Distributed weight (e.g. GTP): reduce-scatter the freshly computed wgrad # (async; overlap with the next layer's bwd via the cascade). if is_dist_weight: - wgrad = finalize_weight_grads(original_weight, [wgrad])[0] + wgrad = finalize_weight_grads(saved_weight, [wgrad])[0] # Update grad bias if needed if grad_bias is None: @@ -1659,7 +1716,7 @@ def wgrad_gemm( # Scatter fp8 weight buffers if bwd_args.fp8 and not bwd_args.is_weight_param_quantized: - _fsdp_scatter_tensors(bwd_args.fsdp_group, weight_for_dgrad) + _fsdp_scatter_tensors(bwd_args.fsdp_group, weight_fp8) return ( wgrad, dgrad if bwd_args.requires_dgrad else None, @@ -1681,7 +1738,7 @@ def _linear_backward_fake( "(fsdp_group is not None); use FSDP2 or MCore FSDP." ) - weight = args.original_weight + weight = args.saved_weight out_dtype = args.activation_dtype out_features, in_features = weight.shape @@ -1809,7 +1866,7 @@ def backward( """Backward pass: compute gradients and reduce FP8 scaling factors.""" bwd_args: LinearBwdArgs = ctx.backward_objects bwd_args.grad_output = grad_output - bwd_args.set_saved_tensors(restore_from_func_ctx(ctx)) + bwd_args.setup_saved_tensors(ctx) nvtx_label = "transformer_engine._Linear.backward" if bwd_args.ub_name is not None: nvtx_label = f"{nvtx_label}.{bwd_args.ub_name}" @@ -2331,10 +2388,7 @@ def forward( if get_ub_is_fp8(self.ub_name + "_dgrad", FP8GlobalStateManager.is_fp8_enabled()): fp8_grad = True - use_compiled_op = torch.compiler.is_compiling() and _linear_op is not None - if _linear_op is None and torch.compiler.is_compiling(): - warn_if_compile_disabled() - if use_compiled_op: + if torch.compiler.is_compiling() and _linear_op is not None: reason = self._compile_eager_fallback_reason( inp, is_first_microbatch, fp8_output, fp8_grad, is_grad_enabled, debug ) @@ -2346,26 +2400,8 @@ def forward( inp = self.prepare_forward(inp, allow_non_contiguous=isinstance(inp, QuantizedTensor)) try: - quantizers = ( - self._get_quantizers(fp8_output, fp8_grad, is_grad_enabled) - if not debug - else self._get_debug_quantizers(fp8_output, fp8_grad, is_grad_enabled) - ) - if debug: - if self.no_debug_features_active(quantizers): - debug = False - quantizers = self._get_quantizers(fp8_output, fp8_grad, is_grad_enabled) - - if use_compiled_op and any( - q is not None and not is_value_opaque_quantizer(q) for q in quantizers - ): - reason = "a quantizer not registered as a torch.compile value-opaque type" - warn_compile_eager_fallback(reason) - torch._dynamo.graph_break(msg=f"te.Linear falling back to eager: {reason}") - use_compiled_op = False - weight_tensor, bias_tensor = self._get_weight_and_bias_tensors( - defer_concatenation=use_compiled_op + defer_concatenation=torch.compiler.is_compiling() and _linear_op is not None ) weight_meta = ( weight_tensor.to_spec() @@ -2376,6 +2412,16 @@ def forward( bias_tensor.to_spec() if isinstance(bias_tensor, ParameterParts) else bias_tensor ) + quantizers = ( + self._get_quantizers(fp8_output, fp8_grad, is_grad_enabled) + if not debug + else self._get_debug_quantizers(fp8_output, fp8_grad, is_grad_enabled) + ) + if debug: + if self.no_debug_features_active(quantizers): + debug = False + quantizers = self._get_quantizers(fp8_output, fp8_grad, is_grad_enabled) + ( input_quantizer, weight_quantizer, @@ -2389,6 +2435,9 @@ def forward( weight_quantizer, weight_meta ) + use_compiled_op = torch.compiler.is_compiling() and _linear_op is not None + if _linear_op is None and torch.compiler.is_compiling(): + warn_if_compile_disabled() if use_compiled_op: # Process groups cross the op boundary separately from quantizers. for quantizer in (input_quantizer, grad_output_quantizer): @@ -2499,6 +2548,22 @@ def forward( is_grad_enabled=is_grad_enabled, ) + if use_compiled_op: + # Safety net for quantizer-dependent conditions only. + fallback_reason = fwd_args.compile_unsupported_reason() + if fallback_reason is not None: + warn_compile_eager_fallback(fallback_reason) + torch._dynamo.graph_break( + 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.shape, weight_meta.shape, self.fp8) out, new_weight_workspace = _linear_op(fwd_args) @@ -2591,7 +2656,8 @@ def _compile_eager_fallback_reason( is_grad_enabled: bool, debug: bool, ) -> Optional[str]: - """Check configuration before prepare_forward, without concatenating parameters.""" + """Why this call can't use the compiled op (else None), decided before + prepare_forward. Quantizer checks stay in compile_unsupported_reason.""" if debug: return "debug instrumentation (nvidia-dlfw-inspect)" if any(is_distributed_weight(getattr(self, name)) for name in self.weight_names): @@ -2646,14 +2712,16 @@ def _forward_eager_fallback( ) def _get_weight_and_bias_tensors(self, defer_concatenation=False): - weight = concat_input(self._get_weight_tensors(), defer=defer_concatenation) - bias = None + weights = self._get_weight_tensors() + weight_tensor = concat_input(weights, defer=defer_concatenation) + bias_tensor = None if self.use_bias: - bias = concat_input( - [getattr(self, name) for name in self.bias_names], + 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, bias + return weight_tensor, bias_tensor def onnx_forward( self,