diff --git a/docs/_static/css/diagram-colors.css b/docs/_static/css/diagram-colors.css index f5dc7da4dd..9ee5827bd1 100644 --- a/docs/_static/css/diagram-colors.css +++ b/docs/_static/css/diagram-colors.css @@ -279,3 +279,18 @@ html[data-theme="dark"] .subtitle, html[data-theme="dark"] .memory-label { fill: #e0e0e0; } html[data-theme="dark"] .connector { stroke: #bdbdbd; } + +/* fine_grained_quantization diagrams */ +html[data-theme="dark"] .fmt-mxfp8 { fill: #10375c; stroke: #64b5f6; } +html[data-theme="dark"] .fmt-nvfp4 { fill: #5c3a10; stroke: #ffb74d; } +html[data-theme="dark"] .fmt-bf16 { fill: #1e4620; stroke: #81c784; } +html[data-theme="dark"] .fmt-mxfp8-text { fill: #90caf9; } +html[data-theme="dark"] .fmt-nvfp4-text { fill: #ffcc80; } +html[data-theme="dark"] .fmt-bf16-text { fill: #a5d6a7; } +html[data-theme="dark"] .source { fill: #3a2f5c; stroke: #b39ddb; } +html[data-theme="dark"] .quantizer { fill: #1e4620; stroke: #81c784; } +html[data-theme="dark"] .representation { fill: #10375c; stroke: #64b5f6; } +html[data-theme="dark"] .dequantize { fill: #5c3a10; stroke: #ffb74d; } +html[data-theme="dark"] .rowlabel, +html[data-theme="dark"] .legend, +html[data-theme="dark"] .op { fill: #e0e0e0; } diff --git a/docs/api/pytorch.rst b/docs/api/pytorch.rst index 54981c9086..8ab9dfbc9d 100644 --- a/docs/api/pytorch.rst +++ b/docs/api/pytorch.rst @@ -6,7 +6,12 @@ PyTorch ======= -.. autoapiclass:: transformer_engine.pytorch.Linear(in_features, out_features, bias=True, **kwargs) +.. autoapiclass:: transformer_engine.pytorch.autocast(enabled=True, calibrating=False, recipe=None, amax_reduction_group=None) + +Standard layers +--------------- + +.. autoapiclass:: transformer_engine.pytorch.Linear(in_features, out_features, **kwargs) :members: forward, set_tensor_parallel_group .. autoapiclass:: transformer_engine.pytorch.GroupedLinear(in_features, out_features, bias=True, **kwargs) @@ -34,20 +39,34 @@ PyTorch .. autoapiclass:: transformer_engine.pytorch.TransformerLayer(hidden_size, ffn_hidden_size, num_attention_heads, **kwargs) :members: forward, set_context_parallel_group, set_tensor_parallel_group +Model-specific layers +--------------------- + +DeepSeek-V3 +^^^^^^^^^^^ + +.. autoapiclass:: transformer_engine.pytorch.models.DeepSeekV3Layer(hidden_size, num_attention_heads, **kwargs) + :members: forward + +.. autoapiclass:: transformer_engine.pytorch.models.DeepSeekV3MoE(hidden_size, moe_ffn_hidden_size, num_experts, **kwargs) + :members: forward, update_expert_bias + +.. autoapiclass:: transformer_engine.pytorch.models.MultiLatentAttention(hidden_size, num_attention_heads, **kwargs) + :members: forward + +Other +----- + .. autoapiclass:: transformer_engine.pytorch.dot_product_attention.inference.InferenceParams(max_batch_size, max_sequence_length) :members: reset, allocate_memory, pre_step, get_seqlens_pre_step, convert_paged_to_nonpaged, step .. autoapiclass:: transformer_engine.pytorch.CudaRNGStatesTracker() :members: reset, get_states, set_states, add, fork - -.. autoapiclass:: transformer_engine.pytorch.autocast(enabled=True, calibrating=False, recipe=None, amax_reduction_group=None) - .. autoapifunction:: transformer_engine.pytorch.quantized_model_init .. autoapifunction:: transformer_engine.pytorch.checkpoint - .. autoapifunction:: transformer_engine.pytorch.make_graphed_callables .. autoapifunction:: transformer_engine.pytorch.get_cpu_offload_context @@ -112,6 +131,12 @@ Communication-computation overlap :members: FP8, NONE +Fine-grained quantization recipes +--------------------------------- + +.. autoapiclass:: transformer_engine.pytorch.QuantizerRole(module_type="", tensor_type="", name="") + + Quantized tensors ----------------- @@ -129,6 +154,10 @@ Quantized tensors .. autoapiclass:: transformer_engine.pytorch.NVFP4TensorStorage(rowwise_data, rowwise_scale_inv, columnwise_data, columnwise_scale_inv, amax_rowwise, amax_columnwise, fp4_dtype, quantizer) +.. autoapiclass:: transformer_engine.pytorch.HybridQuantizedTensorStorage(*, rowwise_storage, columnwise_storage, quantizer, fake_dtype=None) + +.. autoapiclass:: transformer_engine.pytorch.IdentityTensorStorage(*, hp_data, fake_dtype=None, quantizer=None) + .. autoapiclass:: transformer_engine.pytorch.Float8Tensor(shape, dtype, data, fp8_scale_inv, fp8_dtype, requires_grad=False, data_transpose=None, quantizer=None) .. autoapiclass:: transformer_engine.pytorch.MXFP8Tensor(rowwise_data, rowwise_scale_inv, columnwise_data, columnwise_scale_inv, fp8_dtype, quantizer) @@ -137,6 +166,10 @@ Quantized tensors .. autoapiclass:: transformer_engine.pytorch.NVFP4Tensor(rowwise_data, rowwise_scale_inv, columnwise_data, columnwise_scale_inv, amax_rowwise, amax_columnwise, fp4_dtype, quantizer) +.. autoapiclass:: transformer_engine.pytorch.HybridQuantizedTensor(shape, dtype, *, rowwise_storage, columnwise_storage, quantizer, requires_grad=False, device=None) + +.. autoapiclass:: transformer_engine.pytorch.IdentityTensor(shape, dtype, *, hp_data, quantizer=None, requires_grad=False, device=None) + Quantizers ---------- @@ -153,6 +186,10 @@ Quantizers .. autoapiclass:: transformer_engine.pytorch.NVFP4Quantizer(fp4_dtype, *, rowwise=True, columnwise=True, **kwargs) +.. autoapiclass:: transformer_engine.pytorch.HybridQuantizer(*, rowwise_quantizer, columnwise_quantizer, columnwise_source="original") + +.. autoapiclass:: transformer_engine.pytorch.IdentityQuantizer(*, dtype=None, rowwise=True, columnwise=True) + Tensor saving and restoring functions ------------------------------------- diff --git a/docs/examples/fine_grained_quantization/pytorch_fine_grained_quantization_example.py b/docs/examples/fine_grained_quantization/pytorch_fine_grained_quantization_example.py new file mode 100644 index 0000000000..b511fc41dc --- /dev/null +++ b/docs/examples/fine_grained_quantization/pytorch_fine_grained_quantization_example.py @@ -0,0 +1,130 @@ +#!/usr/bin/env python3 +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Runnable fine-grained quantization recipe example. + +The factory assigns one precision to each ``demo.fc1`` Linear GEMM: + +* fprop: ``weight.row(MXFP8) x input.row(MXFP8)`` +* dgrad: ``weight.col(NVFP4) x grad_output.row(NVFP4)`` +* wgrad: ``input.col(original BF16) x grad_output.col(original BF16)`` + +``demo.fc2`` runs every GEMM in high precision. ``demo.output`` is not +special-cased and therefore exercises the MXFP8 base-factory fallback. + +Run from the Transformer Engine repository root:: + + python docs/examples/fine_grained_quantization/\ + pytorch_fine_grained_quantization_example.py +""" + +from __future__ import annotations + +import torch +import transformer_engine.pytorch as te + + +def require_supported_hardware() -> None: + """Fail early with TE's reason when either required format is unavailable.""" + + if not torch.cuda.is_available(): + raise SystemExit("This example requires a CUDA-capable NVIDIA GPU.") + + failures = [] + for name, check in ( + ("MXFP8", te.is_mxfp8_available), + ("NVFP4", te.is_nvfp4_available), + ): + available, reason = check(return_reason=True) + if not available: + failures.append(f"{name}: {reason}") + if failures: + raise SystemExit("Required formats are unavailable: " + "; ".join(failures)) + + +require_supported_hardware() + +# START_FINE_GRAINED_QUANTIZATION_EXAMPLE + +from typing import Optional + +import torch + +import transformer_engine.pytorch as te +from transformer_engine.common.recipe import CustomRecipe +from transformer_engine.pytorch.custom_recipes.quantizer_factories import ( + mxfp8_factory, + nvfp4_factory, +) + + +THREE_FORMAT_MODULE = "demo.fc1" +HIGH_PRECISION_MODULE = "demo.fc2" +BASE_FACTORY = mxfp8_factory + + +def quantizer_factory(role: Optional[te.QuantizerRole]): + """Return a fresh quantizer for every role, including ``None``. + + ``BASE_FACTORY`` makes the factory total: unknown roles, future role values, + and untargeted modules all retain valid MXFP8 behavior. + """ + + if role is not None and role.name == THREE_FORMAT_MODULE: + # Constructing fresh child quantizers for every call is recommended. + if role.tensor_type == "input": + # Wgrad retains the original BF16 input. + return te.HybridQuantizer( + rowwise_quantizer=mxfp8_factory(role), + columnwise_quantizer=te.IdentityQuantizer(), + columnwise_source="original", + ) + if role.tensor_type == "weight": + # Dgrad uses NVFP4 quantized from the dequantized MXFP8 fprop weight. + return te.HybridQuantizer( + rowwise_quantizer=mxfp8_factory(role), + columnwise_quantizer=nvfp4_factory(role), + columnwise_source="rowwise_dequantized", + ) + if role.tensor_type == "grad_output": + # Dgrad uses NVFP4 while wgrad retains the original BF16 gradient. + return te.HybridQuantizer( + rowwise_quantizer=nvfp4_factory(role), + columnwise_quantizer=te.IdentityQuantizer(), + columnwise_source="original", + ) + + if role is not None and role.name == HIGH_PRECISION_MODULE: + return te.IdentityQuantizer() + + return BASE_FACTORY(role) + + +linear_options = {"bias": False, "params_dtype": torch.bfloat16, "device": "cuda"} +model = torch.nn.Sequential( + te.Linear(128, 256, name=THREE_FORMAT_MODULE, **linear_options), + torch.nn.GELU(), + te.Linear(256, 256, name=HIGH_PRECISION_MODULE, **linear_options), + torch.nn.GELU(), + te.Linear(256, 128, name="demo.output", **linear_options), +) +inputs = torch.randn(64, 128, device="cuda", dtype=torch.bfloat16, requires_grad=True) +recipe = CustomRecipe(qfactory=quantizer_factory) + +with te.autocast(enabled=True, recipe=recipe): + outputs = model(inputs) + +loss = outputs.float().square().mean() +loss.backward() + +# END_FINE_GRAINED_QUANTIZATION_EXAMPLE + +gradients = [inputs.grad, *(parameter.grad for parameter in model.parameters())] +assert all(gradient is not None for gradient in gradients) +assert all(torch.isfinite(gradient).all() for gradient in gradients) + +print(f"GPU: {torch.cuda.get_device_name()}") +print(f"TE Linear names: {[model[index].name for index in (0, 2, 4)]}") +print(f"loss: {loss.item():.6f}; forward and backward completed") diff --git a/docs/features/low_precision_training/fine_grained_quantization/fine_grained_quantization.rst b/docs/features/low_precision_training/fine_grained_quantization/fine_grained_quantization.rst new file mode 100644 index 0000000000..301b655b3f --- /dev/null +++ b/docs/features/low_precision_training/fine_grained_quantization/fine_grained_quantization.rst @@ -0,0 +1,364 @@ +.. + Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + + See LICENSE for license information. + +.. _fine-grained-quantization-recipes: +.. _heterogeneous-quantization-recipes: + +Fine-grained quantization recipes +================================= + +Standard TE recipes quantize the whole model the same way. That is often too +coarse: one sensitive layer may need BF16 while the rest runs in MXFP8, or a +gradient GEMM may tolerate a cheaper format than the forward pass. Fine-grained +recipes lift this restriction: you write a small factory function that picks a +quantizer for each slot TE asks about, and pass it via +:class:`~transformer_engine.common.recipe.CustomRecipe` to the usual +:class:`~transformer_engine.pytorch.autocast`. +"Fine-grained" refers to the granularity of that choice (per module, tensor +role, and GEMM direction), not to the block size of the scaling factors. + +.. warning:: + + Fine-grained recipes are currently available only in the PyTorch API of + TE. + +.. warning:: + + Fine-grained recipes and their construction APIs are experimental: API, + validation, and kernel coverage may change without notice. This guide does + not define a supported recipe or an expected accuracy/performance ordering. + + +Example: mixing MXFP8, NVFP4, and BF16 +-------------------------------------- + +The `runnable example `__ +makes the following assignments: + +.. raw:: html + :file: img/fine_grained_assignments.svg + +*Figure 1. Precision assignments per module and GEMM used throughout this +guide.* + +A minimal factory implementing these assignments, plugged into the standard +TE autocast path: + +.. tabs:: + + .. tab:: PyTorch + + .. code-block:: python + + import transformer_engine.pytorch as te + from transformer_engine.common.recipe import CustomRecipe + from transformer_engine.pytorch.custom_recipes.quantizer_factories import ( + mxfp8_factory, + nvfp4_factory, + ) + + + def quantizer_factory(role): + if role is not None and role.name == "demo.fc1": + if role.tensor_type == "input": + # wgrad keeps the original BF16 input + return te.HybridQuantizer( + rowwise_quantizer=mxfp8_factory(role), + columnwise_quantizer=te.IdentityQuantizer(), + columnwise_source="original", + ) + if role.tensor_type == "weight": + # fprop in MXFP8, dgrad in NVFP4 + return te.HybridQuantizer( + rowwise_quantizer=mxfp8_factory(role), + columnwise_quantizer=nvfp4_factory(role), + columnwise_source="rowwise_dequantized", + ) + if role.tensor_type == "grad_output": + # dgrad in NVFP4, wgrad keeps the original BF16 gradient + return te.HybridQuantizer( + rowwise_quantizer=nvfp4_factory(role), + columnwise_quantizer=te.IdentityQuantizer(), + columnwise_source="original", + ) + if role is not None and role.name == "demo.fc2": + return te.IdentityQuantizer() # whole module stays in BF16 + return mxfp8_factory(role) # every other TE module in MXFP8 + + + recipe = CustomRecipe(qfactory=quantizer_factory) + + with te.autocast(enabled=True, recipe=recipe): + output = model(inputs) + +The complete, runnable version is available +`on GitHub `__ +(requires Blackwell or later); run it from the repository root after +installing TE: + +.. code-block:: bash + + python docs/examples/fine_grained_quantization/pytorch_fine_grained_quantization_example.py + +CustomRecipe and quantizer factory +---------------------------------- + +:class:`~transformer_engine.common.recipe.CustomRecipe` is used like any +other TE recipe (``DelayedScaling``, ``MXFP8BlockScaling``, ...), but carries +no quantization logic of its own: TE asks your ``qfactory`` for a quantizer +whenever a module needs one. + +Each TE module defines an ordered role list for the forward and backward +quantizer slots it needs. When module recipe state is initialized or rebuilt, +a ``CustomRecipe`` calls ``qfactory(role)`` once for every slot in that list. +It does not call the factory on every unchanged forward. + +.. tabs:: + + .. tab:: PyTorch + + .. code-block:: python + + # QuantizerRole describes the slot being configured (fields below): + # + # @dataclasses.dataclass(frozen=True) + # class QuantizerRole: + # module_type: str = "" + # tensor_type: str = "" + # name: str = "" + + + def quantizer_factory(role: Optional[te.QuantizerRole]): + # construct a fresh quantizer on every call + ... + # Boundary slots may pass role=None or a role with empty fields, so + # always end with a default that covers every remaining role. + return mxfp8_factory(role) + + + # The factory plugs into the standard TE autocast path: + recipe = CustomRecipe(qfactory=quantizer_factory) + + with te.autocast(enabled=True, recipe=recipe): + output = model(inputs) + + **Module type** + + The kind of TE module that owns the slot, filled in by TE itself: + + * ``"linear"`` — ``Linear``, ``LayerNormLinear``, ``fc1``/``fc2`` in + ``LayerNormMLP``, ``qkv``/``proj`` in ``MultiheadAttention``; + * ``"grouped_linear"`` — ``GroupedLinear``; + * ``"dpa"`` — ``DotProductAttention``. + + **Tensor type** + + Which tensor of that module the quantizer will process, also filled in + by TE. For ``"linear"`` and ``"grouped_linear"``: + + * ``"input"`` — the activation (fprop, wgrad); + * ``"weight"`` — (fprop, dgrad); + * ``"grad_output"`` — the incoming gradient (dgrad, wgrad). + + For ``"dpa"``: + + * ``"qkv"`` — the query/key/value tensor; + * ``"s"`` — the softmax output; + * ``"do"`` — the output gradient; + * ``"dp"`` — the gradient of ``"s"``. + + **Name** + + The identity of one concrete module instance, supplied by the caller: + ``te.Linear(..., name="decoder.39.fc2")``. Composite TE modules may + append suffixes such as ``.fc1``, ``.fc2``, and ``.proj``. + + The role vocabulary is experimental and may grow between releases — + one more reason to end the factory with a total default. Treat the role + strings as selectors, not a fixed enumeration. Prefer a module-level + function for the factory itself, so that launchers and checkpointing + setups can import or pickle it. + + TE provides factories for its native quantizers in + ``transformer_engine.pytorch.custom_recipes.quantizer_factories`` + (``mxfp8_factory``, ``nvfp4_factory``, ...). They can be used as + defaults or to construct ``HybridQuantizer`` children. Additional + specialized recipes are available in + ``transformer_engine.pytorch.custom_recipes.quantizer_factory_zoo``. + + The factory is not limited to TE-native quantizers: it may return your + own :class:`~transformer_engine.pytorch.Quantizer` subclass, and custom + quantizers can also serve as ``HybridQuantizer`` children. The GEMMs + still need to receive representations in formats they support. + +HybridQuantizer +--------------- + +During training, each tensor of a ``Linear`` or ``GroupedLinear`` layer feeds +two different GEMMs: its rowwise representation feeds one, its columnwise +representation the other (the exact operand layout is described in the +:doc:`Introduction <../introduction/introduction>`). Since those two GEMMs may +want different formats, the tensor needs a quantizer per direction: +:class:`~transformer_engine.pytorch.HybridQuantizer` composes a rowwise and a +columnwise quantizer, and its output, +:class:`~transformer_engine.pytorch.HybridQuantizedTensor`, composes the +corresponding representations. + +.. tabs:: + + .. tab:: PyTorch + + The following is pseudocode illustrating the composition: + + .. code-block:: text + + quantizer = te.HybridQuantizer( + rowwise_quantizer=MXFP8Quantizer(fp8_dtype=DType.kFloat8E4M3), + columnwise_quantizer=NVFP4Quantizer(), + columnwise_source="original", # or "rowwise_dequantized" + ) + + # Quantization yields a HybridQuantizedTensor whose rowwise + # representation is MXFP8 and columnwise representation is NVFP4; + # each GEMM consumes the representation it needs. + qtensor = quantizer(tensor) + +.. raw:: html + :file: img/hybrid_quantizer.svg + +*Figure 2. HybridQuantizer composes a rowwise and a columnwise quantizer; each +representation of the result feeds a different GEMM.* + +Choosing the columnwise source +^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +``columnwise_source`` is a separate numerical recipe choice that controls the +source for the columnwise representation: + +.. list-table:: + :header-rows: 1 + :widths: 30 70 + + * - Value + - Columnwise source + * - ``"original"`` + - The original high-precision tensor. + * - ``"rowwise_dequantized"`` + - Dequantized rowwise representation. + +.. raw:: html + :file: img/hybrid_columnwise_source.svg + +*Figure 3. The columnwise representation can be derived from the original +high-precision tensor or from the dequantized rowwise representation.* + +For forward inputs and weights, ``"rowwise_dequantized"`` derives the backward +representation from the value consumed in the forward direction. This +can improve forward/backward numerical consistency and may affect convergence. +It does not recover information discarded by rowwise quantization. +``"original"`` instead derives both representations from the original tensor. +Choose the provenance as part of the numerical recipe. + +IdentityQuantizer +----------------- + +:class:`~transformer_engine.pytorch.IdentityQuantizer` stores its input in the +held compute dtype, typically BF16, FP16, or FP32. It can keep a complete slot +in high precision or act as one child of a ``HybridQuantizer``: + +.. tabs:: + + .. tab:: PyTorch + + .. code-block:: python + + # whole slot in high precision (e.g. a module kept in BF16) + quantizer = te.IdentityQuantizer() + + # one direction in high precision, the other quantized + quantizer = te.HybridQuantizer( + rowwise_quantizer=mxfp8_factory(role), + columnwise_quantizer=te.IdentityQuantizer(), + columnwise_source="rowwise_dequantized", + ) + +Note that in the second example the columnwise direction is high precision but +holds the value reconstructed from MXFP8, not the original input — see +`Choosing the columnwise source`_ above. + +Example: one format per GEMM +---------------------------- + +A natural way to design a recipe is to pick one format for each GEMM. To +translate that into quantizers, look at what each GEMM consumes — both of its +operands must be in that GEMM's format: + +* **fprop** consumes ``input.rowwise`` and ``weight.rowwise``; +* **dgrad** consumes ``grad_output.rowwise`` and ``weight.columnwise``; +* **wgrad** consumes ``input.columnwise`` and ``grad_output.columnwise``. + +Reading the same table per tensor gives the ``HybridQuantizer`` for each role. +For the example assignments (fprop in MXFP8, dgrad in NVFP4, wgrad in BF16): + +.. code-block:: text + + input = HybridQuantizer(rowwise=MXFP8, columnwise=BF16) # fprop | wgrad + weight = HybridQuantizer(rowwise=MXFP8, columnwise=NVFP4) # fprop | dgrad + grad_output = HybridQuantizer(rowwise=NVFP4, columnwise=BF16) # dgrad | wgrad + +.. raw:: html + :file: img/fine_grained_linear_mapping.svg + +*Figure 4. Each GEMM consumes one representation of each of its two operand +tensors; giving both operands the same format sets that GEMM's precision.* + +If two directions use the same quantizer configuration, a plain quantizer may +replace the corresponding hybrid; one factory may return both plain and hybrid +quantizers. +The two operands of each GEMM still need a combination supported by that GEMM +backend. TE may reject incompatible quantizer pairs or unsupported layouts. + +.. note:: + + On supported hardware these recipes run TE's regular quantized kernels: the + tensors are quantized on the GPU and the GEMMs execute in the selected + low-precision formats. TE does not fall back to fake quantization + (quantize-dequantize followed by a high-precision GEMM). + + +Validating and optimizing a recipe +---------------------------------- + +The factory API can express more recipes than TE has kernels for, so any +assignment lands in one of three buckets: + +* **Fast** — quantization hits TE's fused kernels and every GEMM runs a + native low-precision implementation. +* **Correct but potentially unoptimized** — the recipe executes, but some + selected paths may not have fused or optimized implementations in the + current TE release. For example, ``HybridQuantizer`` may produce its rowwise + and columnwise representations in separate kernel launches; future releases + may fuse this work. +* **Rejected** — the two operands of some GEMM end up in a combination of + formats or layouts that no GEMM backend supports, and TE raises an error. + This can happen with plain and hybrid quantizers alike. + +Before adopting a recipe for a real workload, check that: + +* it executes at all on the target GPU, software version, and modules; +* it runs on optimized kernels rather than fallback paths; +* accuracy and convergence hold on the target model and distributed setup; +* throughput and memory actually improve on the target workload. + +The unoptimized paths are still useful: accuracy and convergence experiments can run +on them before dedicated kernels exist, so the precision of each GEMM can be +treated as an accuracy/performance trade-off to explore. + +API reference +------------- + +See the :doc:`PyTorch API <../../../api/pytorch>` for ``QuantizerRole``, +``HybridQuantizer``, ``IdentityQuantizer``, and their returned tensor types. +See the :doc:`Common API <../../../api/common>` for ``CustomRecipe``. diff --git a/docs/features/low_precision_training/fine_grained_quantization/img/fine_grained_assignments.svg b/docs/features/low_precision_training/fine_grained_quantization/img/fine_grained_assignments.svg new file mode 100644 index 0000000000..b9aa48333d --- /dev/null +++ b/docs/features/low_precision_training/fine_grained_quantization/img/fine_grained_assignments.svg @@ -0,0 +1,95 @@ + + + Precision assignment by tensor role and module + Each tensor role provides a rowwise and a columnwise representation, consumed by fprop, dgrad, and wgrad GEMMs. demo.fc1: input is MXFP8 rowwise and original BF16 columnwise; weight is MXFP8 rowwise and NVFP4 columnwise; grad_output is NVFP4 rowwise and original BF16 columnwise. demo.fc2 keeps every tensor in BF16. Other TE modules use MXFP8 everywhere. + + + + + Precision assignment by tensor role and module + + demo.fc1 + demo.fc2 + Other TE modules + + + + input + rowwise (fprop) + + MXFP8 + + BF16 + + MXFP8 + + columnwise (wgrad) + + BF16 (original) + + BF16 + + MXFP8 + + + + + weight + rowwise (fprop) + + MXFP8 + + BF16 + + MXFP8 + + columnwise (dgrad) + + NVFP4 + + BF16 + + MXFP8 + + + + + grad_output + rowwise (dgrad) + + NVFP4 + + BF16 + + MXFP8 + + columnwise (wgrad) + + BF16 (original) + + BF16 + + MXFP8 + + + + + MXFP8 + + NVFP4 + + BF16 (high precision) + + diff --git a/docs/features/low_precision_training/fine_grained_quantization/img/fine_grained_linear_mapping.svg b/docs/features/low_precision_training/fine_grained_quantization/img/fine_grained_linear_mapping.svg new file mode 100644 index 0000000000..0e17ffdfa4 --- /dev/null +++ b/docs/features/low_precision_training/fine_grained_quantization/img/fine_grained_linear_mapping.svg @@ -0,0 +1,59 @@ + + + Per-GEMM formats and the operands each GEMM consumes + Three GEMM cards. Fprop consumes input.rowwise and weight.rowwise, both MXFP8. Dgrad consumes weight.columnwise and grad_output.rowwise, both NVFP4. Wgrad consumes input.columnwise and grad_output.columnwise, both BF16. + + + + + + + fprop + format: MXFP8 + + input.rowwise + MXFP8 + × + + weight.rowwise + MXFP8 + + + + + dgrad + format: NVFP4 + + grad_output.rowwise + NVFP4 + × + + weight.columnwise + NVFP4 + + + + + wgrad + format: BF16 + + input.columnwise + BF16 + × + + grad_output.columnwise + BF16 + + diff --git a/docs/features/low_precision_training/fine_grained_quantization/img/hybrid_columnwise_source.svg b/docs/features/low_precision_training/fine_grained_quantization/img/hybrid_columnwise_source.svg new file mode 100644 index 0000000000..ebc077eb04 --- /dev/null +++ b/docs/features/low_precision_training/fine_grained_quantization/img/hybrid_columnwise_source.svg @@ -0,0 +1,74 @@ + + + Hybrid quantizer columnwise source choices + With original provenance, both quantizers consume the original high-precision tensor. With rowwise-dequantized provenance, the columnwise quantizer consumes the dequantized rowwise representation. + + + + + + + + Choosing the columnwise source + + + + columnwise_source="original" + + + High-precision tensor + + + + same original source + + + Rowwise quantizer + + Columnwise quantizer + + + + + Rowwise + representation + + Columnwise + representation + + + + + columnwise_source="rowwise_dequantized" + + + High-precision tensor + + + + Rowwise quantizer + + + + Rowwise + representation + + + + Dequantize + + + + Columnwise quantizer + + + Columnwise + representation + + diff --git a/docs/features/low_precision_training/fine_grained_quantization/img/hybrid_quantizer.svg b/docs/features/low_precision_training/fine_grained_quantization/img/hybrid_quantizer.svg new file mode 100644 index 0000000000..6e542306c7 --- /dev/null +++ b/docs/features/low_precision_training/fine_grained_quantization/img/hybrid_quantizer.svg @@ -0,0 +1,58 @@ + + + HybridQuantizer data flow + A high-precision tensor enters a HybridQuantizer whose rowwise child is an MXFP8 quantizer and columnwise child is an NVFP4 quantizer. The result is a HybridQuantizedTensor with an MXFP8 rowwise representation and an NVFP4 columnwise representation, each consumed by a different GEMM. + + + + + + + + + + tensor + high precision (BF16) + + + + + + + HybridQuantizer + + rowwise_quantizer + MXFP8Quantizer + + columnwise_quantizer + NVFP4Quantizer + + + + + + + HybridQuantizedTensor + + rowwise + MXFP8 + + columnwise + NVFP4 + + + + + GEMM 1 + GEMM 2 + diff --git a/docs/features/low_precision_training/fp8_blockwise_scaling/fp8_blockwise_scaling.rst b/docs/features/low_precision_training/fp8_blockwise_scaling/fp8_blockwise_scaling.rst index 48d17db8d5..557a4b09e8 100644 --- a/docs/features/low_precision_training/fp8_blockwise_scaling/fp8_blockwise_scaling.rst +++ b/docs/features/low_precision_training/fp8_blockwise_scaling/fp8_blockwise_scaling.rst @@ -180,6 +180,37 @@ Blackwell and later (SM >= 10.0) – the recipe is emulated with MXFP8. Note tha ---- + +Quantizer +--------- + +.. tabs:: + + .. tab:: PyTorch + + Blockwise scaling uses + :class:`~transformer_engine.pytorch.Float8BlockQuantizer`. Each block of + the tensor gets its own power-of-two scale: ``block_scaling_dim=1`` + scales 1x128 blocks, ``block_scaling_dim=2`` (the default) scales + 128x128 blocks. This recipe is not available in TE/JAX. + + .. code-block:: python + + import torch + import transformer_engine.pytorch as te + + tensor = torch.randn(256, 256, device="cuda", dtype=torch.bfloat16) + + quantizer = te.Float8BlockQuantizer( + fp8_dtype=te.DType.kFloat8E4M3, + rowwise=True, + columnwise=True, + block_scaling_dim=1, + ) + + qtensor = quantizer(tensor) + roundtrip = qtensor.dequantize() + Developer Notes --------------- diff --git a/docs/features/low_precision_training/fp8_current_scaling/fp8_current_scaling.rst b/docs/features/low_precision_training/fp8_current_scaling/fp8_current_scaling.rst index cac3792194..2436a07566 100644 --- a/docs/features/low_precision_training/fp8_current_scaling/fp8_current_scaling.rst +++ b/docs/features/low_precision_training/fp8_current_scaling/fp8_current_scaling.rst @@ -164,6 +164,57 @@ Here's how to use FP8 Current Scaling recipe in PyTorch and JAX: ---- + +Quantizer +--------- + +.. tabs:: + + .. tab:: PyTorch + + Current scaling uses + :class:`~transformer_engine.pytorch.Float8CurrentScalingQuantizer`. It + needs no external state: at each call it computes the amax of the input + tensor, derives the scale from it, and then quantizes. + + .. code-block:: python + + import torch + import transformer_engine.pytorch as te + + tensor = torch.randn(256, 256, device="cuda", dtype=torch.bfloat16) + + quantizer = te.Float8CurrentScalingQuantizer( + fp8_dtype=te.DType.kFloat8E4M3, + device="cuda", + ) + + qtensor = quantizer(tensor) + roundtrip = qtensor.dequantize() + + .. tab:: JAX + + Current scaling uses ``CurrentScaleQuantizer``. At each call it computes + the amax of the input tensor, derives the scale from it, and then + quantizes. + + .. code-block:: python + + import jax.numpy as jnp + from transformer_engine.jax.quantize import ( + QuantizerFactory, ScalingMode, QuantizeLayout, + ) + + x = jnp.ones((256, 256), dtype=jnp.bfloat16) + + quantizer = QuantizerFactory.create( + scaling_mode=ScalingMode.CURRENT_TENSOR_SCALING, + q_dtype=jnp.float8_e4m3fn, + q_layout=QuantizeLayout.ROWWISE, + ) + qtensor = quantizer.quantize(x) + roundtrip = qtensor.dequantize() + Developer Notes --------------- @@ -177,4 +228,4 @@ On Blackwell and later, rowwise and columnwise tensors share the same memory lay so all-gather of columnwise tensors is directly supported. For Hopper and Ada, all-gather of transposed FP8 tensors is not supported. -The rowwise tensor is gathered first, then transposed to columnwise format. \ No newline at end of file +The rowwise tensor is gathered first, then transposed to columnwise format. diff --git a/docs/features/low_precision_training/fp8_delayed_scaling/fp8_delayed_scaling.rst b/docs/features/low_precision_training/fp8_delayed_scaling/fp8_delayed_scaling.rst index d39787f6f5..99a379eed1 100644 --- a/docs/features/low_precision_training/fp8_delayed_scaling/fp8_delayed_scaling.rst +++ b/docs/features/low_precision_training/fp8_delayed_scaling/fp8_delayed_scaling.rst @@ -160,4 +160,69 @@ However, amax reduction works slightly differently in different frameworks. Supported devices ----------------- -Ada and later (SM 8.9+) \ No newline at end of file +Ada and later (SM 8.9+) + +Quantizer +--------- + +.. tabs:: + + .. tab:: PyTorch + + Delayed scaling uses + :class:`~transformer_engine.pytorch.Float8Quantizer`. It does not + compute the scaling factor from the current tensor: one-element + ``scale`` and ``amax`` buffers are supplied at construction. + Quantization applies the given scale and records the tensor's amax into + the ``amax`` buffer. + + During training both buffers are views into the recipe state: ``scale`` + into its per-quantizer scale vector, ``amax`` into the current row of + its ``(amax_history_len, num_quantizers)`` amax history. At the end of + each step the recipe state computes a new scale from the history (its + max or most recent entry, per ``amax_compute_algo``), rolls the history + by one slot, and zeroes the current row — all in place, so the views + held by the quantizer stay valid for the whole training run. + + .. code-block:: python + + import torch + import transformer_engine.pytorch as te + + tensor = torch.randn(256, 256, device="cuda", dtype=torch.bfloat16) + + quantizer = te.Float8Quantizer( + scale=torch.ones(1, device="cuda"), + amax=torch.zeros(1, device="cuda"), + fp8_dtype=te.DType.kFloat8E4M3, + ) + + qtensor = quantizer(tensor) + roundtrip = qtensor.dequantize() + + .. tab:: JAX + + Delayed scaling uses ``DelayedScaleQuantizer``. The ``scale`` and the + ``amax_history`` (1024 entries by default) are fields of the quantizer + itself, carried through JAX transformations as its pytree state. Each + ``quantize()`` call applies the current ``scale``, then updates the + state: the tensor's amax is written into the history, a new scale is + computed from the history (max or most-recent entry, per + ``amax_compute_algo``), and the history is rolled by one slot. + + .. code-block:: python + + import jax.numpy as jnp + from transformer_engine.jax.quantize import ( + QuantizerFactory, ScalingMode, QuantizeLayout, + ) + + x = jnp.ones((256, 256), dtype=jnp.bfloat16) + + quantizer = QuantizerFactory.create( + scaling_mode=ScalingMode.DELAYED_TENSOR_SCALING, + q_dtype=jnp.float8_e4m3fn, + q_layout=QuantizeLayout.ROWWISE, + ) + qtensor = quantizer.quantize(x) + roundtrip = qtensor.dequantize() diff --git a/docs/features/low_precision_training/index.rst b/docs/features/low_precision_training/index.rst index 0a798f1364..b9649c00a4 100644 --- a/docs/features/low_precision_training/index.rst +++ b/docs/features/low_precision_training/index.rst @@ -15,4 +15,5 @@ Low precision training fp8_blockwise_scaling/fp8_blockwise_scaling.rst mxfp8/mxfp8.rst nvfp4/nvfp4.rst + fine_grained_quantization/fine_grained_quantization.rst speedups.rst diff --git a/docs/features/low_precision_training/introduction/introduction.rst b/docs/features/low_precision_training/introduction/introduction.rst index fba7796ece..2255308b04 100644 --- a/docs/features/low_precision_training/introduction/introduction.rst +++ b/docs/features/low_precision_training/introduction/introduction.rst @@ -283,3 +283,71 @@ so GEMM with tensors ``A`` and ``B`` returns ``B * A^T``. :file: img/fp8_linear_flow.svg *Figure 4: Forward pass of a Linear layer with low precision data flow.* + +Quantizers +---------- + +Every recipe implements its quantization logic in a **quantizer** — an object +that converts a high-precision tensor into a quantized one. TE modules create +and use quantizers internally according to the active recipe, but a quantizer +can also be used directly: + +.. tabs:: + + .. tab:: PyTorch + + .. code-block:: python + + import torch + import transformer_engine.pytorch as te + + tensor = torch.randn(256, 256, device="cuda", dtype=torch.bfloat16) + + quantizer = te.MXFP8Quantizer(fp8_dtype=te.DType.kFloat8E4M3) + qtensor = quantizer(tensor) # quantize + roundtrip = qtensor.dequantize() # back to high precision + + The main parts of the interface are: + + * ``quantize(tensor)`` — quantizes a high-precision tensor and returns a + ``QuantizedTensor``; calling the quantizer (``quantizer(tensor)``) is + a shorthand; + * ``update_quantized(src, dst)`` — quantizes ``src`` in place into an + already-allocated quantized tensor ``dst``; + * ``make_empty(shape)`` — allocates an uninitialized quantized tensor to + be filled later; + * ``rowwise_usage`` / ``columnwise_usage`` — flags selecting which of + the two GEMM-oriented representations the produced tensor holds; + * the returned ``QuantizedTensor`` supports ``dequantize()`` back to + high precision. + + .. tab:: JAX + + .. code-block:: python + + import jax.numpy as jnp + from transformer_engine.jax.quantize import ( + QuantizerFactory, ScalingMode, QuantizeLayout, + ) + + x = jnp.ones((256, 256), dtype=jnp.bfloat16) + + quantizer = QuantizerFactory.create( + scaling_mode=ScalingMode.MXFP8_1D_SCALING, + q_dtype=jnp.float8_e4m3fn, + q_layout=QuantizeLayout.ROWWISE, + ) + qtensor = quantizer.quantize(x) + roundtrip = qtensor.dequantize() + + The main parts of the interface are: + + * ``quantize(x, is_rowwise=..., is_colwise=...)`` — quantizes a tensor + and returns a ``ScaledTensor`` holding the requested representations + (the default comes from the quantizer's ``q_layout``); + * the returned ``ScaledTensor`` supports ``dequantize()`` back to high + precision; + * quantizers are registered pytrees, so they can be passed through JAX + transformations. + +Each recipe section ends with a short description of that recipe's quantizer. diff --git a/docs/features/low_precision_training/mxfp8/mxfp8.rst b/docs/features/low_precision_training/mxfp8/mxfp8.rst index 1fbcc43af9..1827d42cb6 100644 --- a/docs/features/low_precision_training/mxfp8/mxfp8.rst +++ b/docs/features/low_precision_training/mxfp8/mxfp8.rst @@ -152,6 +152,57 @@ SM 10.0, SM 10.3 ---- + +Quantizer +--------- + +.. tabs:: + + .. tab:: PyTorch + + MXFP8 uses :class:`~transformer_engine.pytorch.MXFP8Quantizer`. Every + 32-element block shares one power-of-two (E8M0) scale, computed from the + block's amax at quantization time; no external state is needed. + + .. code-block:: python + + import torch + import transformer_engine.pytorch as te + + tensor = torch.randn(256, 256, device="cuda", dtype=torch.bfloat16) + + quantizer = te.MXFP8Quantizer(fp8_dtype=te.DType.kFloat8E4M3) + + qtensor = quantizer(tensor) + roundtrip = qtensor.dequantize() + + .. tab:: JAX + + MXFP8 uses ``BlockScaleQuantizer`` — the JAX quantizer for block-based + scaling, selected by ``ScalingMode.MXFP8_1D_SCALING``. Instead of one + scale per tensor, the tensor is split along the quantization axis into + 32-element blocks and each block gets its own power-of-two (E8M0) scale, + computed from that block's amax at quantization time. Because the scale + is derived from the current data, no external state (scale buffers or + amax history) is needed. + + .. code-block:: python + + import jax.numpy as jnp + from transformer_engine.jax.quantize import ( + QuantizerFactory, ScalingMode, QuantizeLayout, + ) + + x = jnp.ones((256, 256), dtype=jnp.bfloat16) + + quantizer = QuantizerFactory.create( + scaling_mode=ScalingMode.MXFP8_1D_SCALING, + q_dtype=jnp.float8_e4m3fn, + q_layout=QuantizeLayout.ROWWISE, + ) + qtensor = quantizer.quantize(x) + roundtrip = qtensor.dequantize() + Developer Notes --------------- @@ -210,4 +261,4 @@ All-gather of columnwise tensors All-gather of columnwise tensors is supported and necessary because: - columnwise quantized tensors cannot be computed from rowwise quantized ones, -- gathering high-precision tensors is avoided in most cases for performance reasons. \ No newline at end of file +- gathering high-precision tensors is avoided in most cases for performance reasons. diff --git a/docs/features/low_precision_training/nvfp4/nvfp4.rst b/docs/features/low_precision_training/nvfp4/nvfp4.rst index 900789b0d3..26d798651d 100644 --- a/docs/features/low_precision_training/nvfp4/nvfp4.rst +++ b/docs/features/low_precision_training/nvfp4/nvfp4.rst @@ -250,6 +250,60 @@ Supported devices ---- + +Quantizer +--------- + +.. tabs:: + + .. tab:: PyTorch + + NVFP4 uses :class:`~transformer_engine.pytorch.NVFP4Quantizer`. It + implements the two-level scaling described above: an FP8 (E4M3) scale + per 16-element block plus one FP32 scale per tensor. Further keyword + options select the recipe variations from this page (random Hadamard + transforms, stochastic rounding, 2D weight scaling); they are internal + knobs and may change without notice. + + .. code-block:: python + + import torch + import transformer_engine.pytorch as te + + tensor = torch.randn(256, 256, device="cuda", dtype=torch.bfloat16) + + quantizer = te.NVFP4Quantizer( + fp4_dtype=te.DType.kFloat4E2M1, + rowwise=True, + columnwise=True, + ) + + qtensor = quantizer(tensor) + roundtrip = qtensor.dequantize() + + .. tab:: JAX + + NVFP4 uses its own ``NVFP4Quantizer``, with the same two-level scaling. + ``ScalingMode.NVFP4_1D_SCALING`` selects per-block scaling only, + ``ScalingMode.NVFP4_2D_SCALING`` adds 2D weight scaling. + + .. code-block:: python + + import jax.numpy as jnp + from transformer_engine.jax.quantize import ( + QuantizerFactory, ScalingMode, QuantizeLayout, + ) + + x = jnp.ones((256, 256), dtype=jnp.bfloat16) + + quantizer = QuantizerFactory.create( + scaling_mode=ScalingMode.NVFP4_1D_SCALING, + q_dtype=jnp.float4_e2m1fn, + q_layout=QuantizeLayout.ROWWISE, + ) + qtensor = quantizer.quantize(x) + roundtrip = qtensor.dequantize() + Developer Notes --------------- diff --git a/docs/features/low_precision_training/performance_considerations/performance_considerations.rst b/docs/features/low_precision_training/performance_considerations/performance_considerations.rst index 2c21799dd6..afa5d16c86 100644 --- a/docs/features/low_precision_training/performance_considerations/performance_considerations.rst +++ b/docs/features/low_precision_training/performance_considerations/performance_considerations.rst @@ -143,6 +143,61 @@ Transformer Engine chooses the best possible fusion internally taking the recipe *Figure 3: Three scenarios of producing quantized tensors in rowwise and columnwise usages.* +**Usages in the quantizer API** + +The usages are visible directly in the quantizer API: + +.. tabs:: + + .. tab:: PyTorch + + At quantization time, the quantizer's ``rowwise_usage`` and + ``columnwise_usage`` flags select which representations ``quantize()`` + produces; when both are set, the representations are computed together + in one fused kernel (scenario 1 above). + + After quantization, ``update_usage()`` on the quantized tensor removes a + representation or, when supported by the format, generates a missing one. + Passing ``rowwise_usage=False`` after the forward pass frees the rowwise + data while keeping the columnwise data for backward. Some formats also + support ``columnwise_usage=True`` to create the columnwise representation + from the data already present (e.g. by a transpose on Hopper — scenario 3 + above); unsupported requests raise an error. Arguments left as ``None`` + preserve the current state. + + .. code-block:: python + + quantizer = te.MXFP8Quantizer( + fp8_dtype=te.DType.kFloat8E4M3, + rowwise=True, + columnwise=True, + ) + + qtensor = quantizer(tensor) # both representations, one fused kernel + + qtensor.update_usage(rowwise_usage=False) # drop rowwise, keep columnwise + + .. tab:: JAX + + The usages are selected when the tensor is quantized: the quantizer's + ``q_layout`` (``QuantizeLayout.ROWWISE``, ``COLWISE``, or + ``ROWWISE_COLWISE``) sets the default, and ``quantize()`` accepts + ``is_rowwise``/``is_colwise`` overrides. Requesting both usages returns + a ``ScaledTensor2x`` holding the two representations. There is no + in-place ``update_usage()``: JAX arrays are immutable, so a + representation is not added or dropped later — unneeded ones are simply + not requested and get dropped by XLA's dead-code elimination. + + .. code-block:: python + + quantizer = QuantizerFactory.create( + scaling_mode=ScalingMode.MXFP8_1D_SCALING, + q_dtype=jnp.float8_e4m3fn, + q_layout=QuantizeLayout.ROWWISE_COLWISE, + ) + + qtensor = quantizer.quantize(x) # ScaledTensor2x, both representations + rowwise_only = quantizer.quantize(x, is_rowwise=True, is_colwise=False) Memory usage @@ -470,4 +525,3 @@ Actual behavior depends on the recipe and module configuration. *Figure 5: All-gather of quantized tensors for input and gradient tensors. This is one possible scenario — actual behavior varies depending on the recipe and module configuration.* - diff --git a/qa/L0_pytorch_unittest/test.sh b/qa/L0_pytorch_unittest/test.sh index a78a99d7f9..56aef3b840 100644 --- a/qa/L0_pytorch_unittest/test.sh +++ b/qa/L0_pytorch_unittest/test.sh @@ -58,6 +58,7 @@ python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_backward_overrid python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_permutation.xml $TE_PATH/tests/pytorch/test_permutation.py || test_fail "test_permutation.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_cross_entropy.xml $TE_PATH/tests/pytorch/test_cross_entropy.py || test_fail "test_cross_entropy.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_cpu_offloading.xml $TE_PATH/tests/pytorch/test_cpu_offloading.py || test_fail "test_cpu_offloading.py" +python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_models.xml $TE_PATH/tests/pytorch/test_models.py || test_fail "test_models.py" NVTE_FLASH_ATTN=0 NVTE_CPU_OFFLOAD_V1=1 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_cpu_offloading_v1.xml $TE_PATH/tests/pytorch/test_cpu_offloading_v1.py || test_fail "test_cpu_offloading_v1.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_hybrid_quantization.xml $TE_PATH/tests/pytorch/test_hybrid_quantization.py || test_fail "test_hybrid_quantization.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_identity_quantizer.xml $TE_PATH/tests/pytorch/test_identity_quantizer.py || test_fail "test_identity_quantizer.py" diff --git a/qa/L1_pytorch_distributed_unittest/test.sh b/qa/L1_pytorch_distributed_unittest/test.sh index f1de313fdc..68f242870c 100644 --- a/qa/L1_pytorch_distributed_unittest/test.sh +++ b/qa/L1_pytorch_distributed_unittest/test.sh @@ -56,6 +56,7 @@ python3 -m pytest -v -s --junitxml=$XML_LOG_DIR/pytest_test_cu_seqlens_cache.xml python3 -m pytest -v -s --junitxml=$XML_LOG_DIR/pytest_test_cast_master_weights_to_fp8.xml $TE_PATH/tests/pytorch/distributed/test_cast_master_weights_to_fp8.py || test_fail "test_cast_master_weights_to_fp8.py" python3 -m pytest -v -s --junitxml=$XML_LOG_DIR/pytest_test_newton_schulz.xml $TE_PATH/tests/pytorch/distributed/test_newton_schulz.py || test_fail "test_newton_schulz.py" python3 -m pytest -v -s --junitxml=$XML_LOG_DIR/pytest_test_ep.xml $TE_PATH/tests/pytorch/distributed/test_ep.py || test_fail "test_ep.py" +python3 -m pytest -v -s --junitxml=$XML_LOG_DIR/pytest_test_models.xml $TE_PATH/tests/pytorch/distributed/test_models.py || test_fail "distributed/test_models.py" # debug tests diff --git a/tests/jax/test_fused_attn.py b/tests/jax/test_fused_attn.py index b6fc8f7794..1ca4121a4f 100644 --- a/tests/jax/test_fused_attn.py +++ b/tests/jax/test_fused_attn.py @@ -612,11 +612,6 @@ def _check_configs(self): pytest.skip( "B1SS, BHSS and 11SS bias shapes are only supported for non-padding mask" ) - elif self.backend != NVTE_Fused_Attn_Backend.NVTE_F16_arbitrary_seqlen: - pytest.skip( - "B1SS, BHSS and 11SS bias shapes are only supported for " - "the F16_arbitrary_seqlen backend." - ) def _setup_inputs(self): self._check_configs() diff --git a/tests/pytorch/attention/test_linear_mxfp8_attention.py b/tests/pytorch/attention/test_linear_mxfp8_attention.py index f1bba7bc9a..95770a6bb8 100644 --- a/tests/pytorch/attention/test_linear_mxfp8_attention.py +++ b/tests/pytorch/attention/test_linear_mxfp8_attention.py @@ -36,7 +36,11 @@ _current_file = pathlib.Path(__file__).resolve() sys.path = [str(_current_file.parent.parent)] + sys.path from utils import ModelConfig, compare_and_assert, get_available_attention_backends -from mla_rope_utils import apply_mla_rope, build_rope_tables +from transformer_engine.pytorch.models.deepseek_v3.mla_rope import ( + apply_mla_rope_kv, + apply_mla_rope_q, + build_rope_tables, +) try: @@ -183,6 +187,13 @@ def _run_projections( return q_flat, kv_flat, q, kv, k_pos_emb +def _apply_rope(q, kv, k_pos_emb, rope_tables): + cos, sin = rope_tables + q = apply_mla_rope_q(q, cos, sin, HEAD_DIM_NOPE, HEAD_DIM_ROPE) + k, v = apply_mla_rope_kv(kv, k_pos_emb, cos, sin, HEAD_DIM_NOPE, HEAD_DIM_ROPE, HEAD_DIM_V) + return q, k, v + + def _run_forward_bf16( modules: tuple, x: torch.Tensor, @@ -190,7 +201,7 @@ def _run_forward_bf16( ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: q_proj, kv_proj, dpa, out_linear = modules _, _, q, kv, k_pos_emb = _run_projections(q_proj, kv_proj, x) - q, k, v = apply_mla_rope(q, kv, k_pos_emb, cos_table=rope_tables[0], sin_table=rope_tables[1]) + q, k, v = _apply_rope(q, kv, k_pos_emb, rope_tables) attn_out = dpa(q, k, v, qkv_format="sbhd") return q, k, v, out_linear(attn_out.view(x.shape[0], x.shape[1], HIDDEN_SIZE)) @@ -212,13 +223,7 @@ def _run_forward_mxfp8( x, is_first_microbatch, ) - q, k, v = apply_mla_rope( - q, - kv, - k_pos_emb, - cos_table=rope_tables[0], - sin_table=rope_tables[1], - ) + q, k, v = _apply_rope(q, kv, k_pos_emb, rope_tables) attn_out = dpa(q, k, v, qkv_format="sbhd") out = out_linear( attn_out.view(x.shape[0], x.shape[1], HIDDEN_SIZE), @@ -292,7 +297,7 @@ def test_accuracy(self, batch_size: int, seq_len: int) -> None: _set_seed() baseline_modules, mxfp8_modules = _build_modules() x = torch.randn(seq_len, batch_size, HIDDEN_SIZE, dtype=torch.bfloat16, device="cuda") - rope_tables = build_rope_tables(seq_len, device=x.device) + rope_tables = build_rope_tables(seq_len, HEAD_DIM_ROPE, device=x.device) q_bf16, k_bf16, v_bf16, out_bf16 = _run_forward_bf16(baseline_modules, x, rope_tables) q_mxfp8, k_mxfp8, v_mxfp8, out_mxfp8 = _run_forward_mxfp8( @@ -378,7 +383,7 @@ def test_backward(self, batch_size: int, seq_len: int) -> None: device="cuda", requires_grad=True, ) - rope_tables = build_rope_tables(seq_len, device=x.device) + rope_tables = build_rope_tables(seq_len, HEAD_DIM_ROPE, device=x.device) *_, out_mxfp8 = _run_forward_mxfp8(mxfp8_modules, x, fp8_recipe, rope_tables) out_mxfp8.sum().backward() @@ -412,7 +417,7 @@ def test_performance(self, batch_size: int, seq_len: int) -> None: device="cuda", requires_grad=True, ) - rope_tables = build_rope_tables(seq_len, device=x.device) + rope_tables = build_rope_tables(seq_len, HEAD_DIM_ROPE, device=x.device) mxfp8_fprop_ms, mxfp8_bprop_ms = _benchmark_training_step( _run_forward_mxfp8, mxfp8_modules, x, fp8_recipe, rope_tables diff --git a/tests/pytorch/distributed/fsdp2_tests/run_fsdp2_fused_adam.py b/tests/pytorch/distributed/fsdp2_tests/run_fsdp2_fused_adam.py index 08e762045a..fc482ce5b2 100644 --- a/tests/pytorch/distributed/fsdp2_tests/run_fsdp2_fused_adam.py +++ b/tests/pytorch/distributed/fsdp2_tests/run_fsdp2_fused_adam.py @@ -1733,6 +1733,10 @@ def test_fused_adam_hybrid_scale_uniform_across_shards(hybrid_recipe_name): ), f"missing hybrid current-scaling directions: {checked}" +@pytest.mark.skipif( + not te.is_fp8_available(), + reason=te.is_fp8_available(return_reason=True)[1], +) def test_fused_adam_hybrid_identity_fp8_master_weights(): """FSDP2 + FusedAdam with Hybrid(FP8 current rowwise, Identity columnwise). diff --git a/tests/pytorch/distributed/fsdp2_tests/run_fsdp2_model.py b/tests/pytorch/distributed/fsdp2_tests/run_fsdp2_model.py index 39d825e701..c2a9df765e 100644 --- a/tests/pytorch/distributed/fsdp2_tests/run_fsdp2_model.py +++ b/tests/pytorch/distributed/fsdp2_tests/run_fsdp2_model.py @@ -510,6 +510,10 @@ def _hybrid_param_count(): _check_fp8_fsdp2_allgather(model, tols=dict(atol=5e-4, rtol=5e-3)) +@pytest.mark.skipif( + not te.is_fp8_available(), + reason=te.is_fp8_available(return_reason=True)[1], +) def test_distributed_hybrid_identity_all(): """FSDP2 training/all-gather with an all-Identity CustomRecipe. diff --git a/tests/pytorch/distributed/run_models.py b/tests/pytorch/distributed/run_models.py new file mode 100644 index 0000000000..9561d20117 --- /dev/null +++ b/tests/pytorch/distributed/run_models.py @@ -0,0 +1,169 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. +"""Multi-process tests for model-specific layers (te.models), launched via torchrun.""" + +import os +import sys + +import torch +import torch.distributed as dist + +from transformer_engine.pytorch.ep import ep_bootstrap, ep_finalize, release_symm_mem_pool +from transformer_engine.pytorch.models import DeepSeekV3Layer + +HIDDEN = 256 +MOE_FFN = 128 +SHARED_FFN = 128 +NUM_LOCAL_EXPERTS = 2 +TOP_K = 2 +TOKENS_PER_RANK = 64 +HEADS = 4 +DTYPE = torch.bfloat16 + +MLA_KWARGS = dict( + q_lora_rank=96, + kv_lora_rank=64, + qk_nope_head_dim=64, + qk_rope_head_dim=32, + v_head_dim=64, +) + + +def _device_sm() -> int: + major, minor = torch.cuda.get_device_capability() + return major * 10 + minor + + +def _recv_capacity(ep_size: int) -> int: + cap = ep_size * TOKENS_PER_RANK * TOP_K + NUM_LOCAL_EXPERTS * 128 + return -(-cap // 128) * 128 + + +def _broadcast_params(module: torch.nn.Module) -> None: + for t in list(module.parameters()) + list(module.buffers()): + dist.broadcast(t.detach(), src=0) + + +def _make_layer(ep_group, ep_size: int, num_experts: int) -> DeepSeekV3Layer: + ep = ep_group is not None + return DeepSeekV3Layer( + HIDDEN, + HEADS, + num_experts=num_experts, + moe_ffn_hidden_size=MOE_FFN, + topk=TOP_K, + shared_expert_ffn_hidden_size=SHARED_FFN, + params_dtype=DTYPE, + ep_group=ep_group, + ep_max_tokens_per_rank=TOKENS_PER_RANK if ep else None, + **MLA_KWARGS, + ) + + +def _copy_weights(ep_layer: DeepSeekV3Layer, ref: DeepSeekV3Layer, rank: int) -> None: + ref_params = dict(ref.named_parameters()) + ref_bufs = dict(ref.named_buffers()) + with torch.no_grad(): + for name, p in ep_layer.named_parameters(): + if not name.startswith("mlp.experts."): + p.copy_(ref_params[name]) + for name, b in ep_layer.named_buffers(): + if name in ref_bufs and b.shape == ref_bufs[name].shape: + b.copy_(ref_bufs[name]) + ep_fc1, _, ep_fc2 = ep_layer.mlp.experts + ref_fc1, _, ref_fc2 = ref.mlp.experts + for local_e in range(NUM_LOCAL_EXPERTS): + global_e = rank * NUM_LOCAL_EXPERTS + local_e + getattr(ep_fc1, f"weight{local_e}").copy_(getattr(ref_fc1, f"weight{global_e}")) + getattr(ep_fc2, f"weight{local_e}").copy_(getattr(ref_fc2, f"weight{global_e}")) + + +def test_layer_ep_matches_local(rank: int, ep_size: int, ep_group) -> None: + """Full DeepSeekV3Layer with EP must match the all-experts-local layer numerically.""" + num_experts = NUM_LOCAL_EXPERTS * ep_size + torch.manual_seed(0) + ref = _make_layer(None, ep_size, num_experts) + _broadcast_params(ref) + ep_layer = _make_layer(ep_group, ep_size, num_experts) + _copy_weights(ep_layer, ref, rank) + + torch.manual_seed(1234 + rank) + x = torch.randn(TOKENS_PER_RANK // 2, 2, HIDDEN, dtype=DTYPE, device="cuda") + x_ep = x.clone().requires_grad_(True) + x_ref = x.clone().requires_grad_(True) + + out_ep = ep_layer(x_ep) + out_ref = ref(x_ref) + assert out_ep.shape == x.shape + torch.testing.assert_close(out_ep, out_ref, rtol=0.05, atol=0.05) + + grad_out = torch.randn_like(out_ep) + out_ep.backward(grad_out) + out_ref.backward(grad_out) + torch.testing.assert_close(x_ep.grad, x_ref.grad, rtol=0.05, atol=0.05) + + ref_params = dict(ref.named_parameters()) + for name, p in ep_layer.named_parameters(): + if name.startswith("mlp.experts.") or p.grad is None: + continue + torch.testing.assert_close(p.grad, ref_params[name].grad, rtol=0.1, atol=0.1, msg=name) + + # A local expert's wgrad on its owner rank equals the sum of the + # reference wgrads over all ranks. all_reduce is collective, so every + # rank must reduce every expert's grad (in the same order). + ep_fc1, _, ep_fc2 = ep_layer.mlp.experts + ref_fc1, _, ref_fc2 = ref.mlp.experts + for ep_fc, ref_fc in ((ep_fc1, ref_fc1), (ep_fc2, ref_fc2)): + ref_grads = [getattr(ref_fc, f"weight{e}").grad.float().clone() for e in range(num_experts)] + for g in ref_grads: + dist.all_reduce(g) + for local_e in range(NUM_LOCAL_EXPERTS): + global_e = rank * NUM_LOCAL_EXPERTS + local_e + ep_grad = getattr(ep_fc, f"weight{local_e}").grad.float() + torch.testing.assert_close(ep_grad, ref_grads[global_e], rtol=0.1, atol=0.1) + + counts = ep_layer.mlp._last_tokens_per_expert.clone() + dist.all_reduce(counts) + assert counts.sum().item() == ep_size * TOKENS_PER_RANK * TOP_K + + ep_layer.mlp.update_expert_bias() + assert torch.isfinite(ep_layer.mlp.expert_bias).all() + + +def main() -> int: + dist.init_process_group(backend="nccl") + torch.cuda.set_device(int(os.environ["LOCAL_RANK"])) + from torch.distributed import _symmetric_memory as _symm_mem + + _symm_mem.set_backend("NCCL") + + rank = dist.get_rank() + ep_size = dist.get_world_size() + if _device_sm() < 90: + if rank == 0: + print(f"NCCL EP requires SM>=90 (got SM{_device_sm()}); skipping.") + dist.destroy_process_group() + return 0 + + ep_group = dist.new_group(ranks=list(range(ep_size)), backend="nccl") + ep_bootstrap( + ep_group, + num_experts=NUM_LOCAL_EXPERTS * ep_size, + max_tokens_per_rank=TOKENS_PER_RANK, + hidden_dim=HIDDEN, + num_topk=TOP_K, + recv_capacity_per_rank=_recv_capacity(ep_size), + ) + test_layer_ep_matches_local(rank, ep_size, ep_group) + print(f"[rank {rank}] PASSED") + + dist.barrier() + ep_finalize() + release_symm_mem_pool() + dist.destroy_process_group() + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tests/pytorch/distributed/test_fusible_ops.py b/tests/pytorch/distributed/test_fusible_ops.py index d733286093..c314bf5bb6 100644 --- a/tests/pytorch/distributed/test_fusible_ops.py +++ b/tests/pytorch/distributed/test_fusible_ops.py @@ -12,6 +12,7 @@ import pathlib import subprocess import sys +import tempfile from typing import Optional import pytest @@ -58,7 +59,8 @@ def world_group() -> torch.distributed.ProcessGroup: torch.cuda.set_device(rank) group = torch.distributed.init_process_group( "nccl", - init_method="file:///tmp/rdzv", + # Each parallel job must use a fresh FileStore shared by only its ranks. + init_method=f"file://{os.environ['NVTE_TEST_RDZV_PATH']}", world_size=world_size, rank=rank, ) @@ -1053,10 +1055,10 @@ def test_distributed_fuser_ops(world_size: int) -> None: current_file, "--parallel", ] - result = subprocess.run( - command, - check=True, - ) + with tempfile.TemporaryDirectory(prefix="te-test-fusible-ops-") as temp_dir: + env = dict(os.environ) + env["NVTE_TEST_RDZV_PATH"] = str(pathlib.Path(temp_dir) / "rdzv") + subprocess.run(command, check=True, env=env) def main() -> None: diff --git a/tests/pytorch/distributed/test_fusible_ops_with_userbuffers.py b/tests/pytorch/distributed/test_fusible_ops_with_userbuffers.py index 07dffebf5f..38f49a96cb 100644 --- a/tests/pytorch/distributed/test_fusible_ops_with_userbuffers.py +++ b/tests/pytorch/distributed/test_fusible_ops_with_userbuffers.py @@ -12,6 +12,7 @@ import pathlib import subprocess import sys +import tempfile import pytest import torch @@ -107,7 +108,8 @@ def world_group() -> torch.distributed.ProcessGroup: torch.cuda.set_device(local_rank) group = torch.distributed.init_process_group( "nccl", - init_method="file:///tmp/rdzv", + # Each parallel job must use a fresh FileStore shared by only its ranks. + init_method=f"file://{os.environ['NVTE_TEST_RDZV_PATH']}", world_size=world_size, rank=rank, device_id=torch.device(f"cuda:{local_rank}"), @@ -471,7 +473,9 @@ def test_fuser_ops_with_userbuffers( env["NVTE_ALLOW_NONDETERMINISTIC_ALGO"] = "0" # Launch parallel job - run_distributed(command, env=env) + with tempfile.TemporaryDirectory(prefix="te-test-fusible-ops-userbuffers-") as temp_dir: + env["NVTE_TEST_RDZV_PATH"] = str(pathlib.Path(temp_dir) / "rdzv") + run_distributed(command, env=env) def main() -> None: diff --git a/tests/pytorch/distributed/test_models.py b/tests/pytorch/distributed/test_models.py new file mode 100644 index 0000000000..1b96eae2aa --- /dev/null +++ b/tests/pytorch/distributed/test_models.py @@ -0,0 +1,31 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +import os +import subprocess +from pathlib import Path + +import pytest +import torch + +TEST_ROOT = Path(__file__).parent.resolve() +NUM_PROCS = min(8, torch.cuda.device_count()) +LAUNCH_CMD = ["torchrun", f"--nproc_per_node={NUM_PROCS}"] + + +def _has_nvlink() -> bool: + # NCCL EP falls back to the network transport and deadlocks on PCIe-only nodes. + out = subprocess.run( + ["nvidia-smi", "nvlink", "--status"], capture_output=True, text=True, check=False + ).stdout + return "GB/s" in out + + +@pytest.mark.skipif(NUM_PROCS < 2, reason="EP requires >= 2 GPUs") +@pytest.mark.skipif(not _has_nvlink(), reason="NCCL EP requires NVLink") +def test_deepseek_layer_ep(): + result = subprocess.run( + LAUNCH_CMD + [str(TEST_ROOT / "run_models.py")], env=os.environ, check=False, timeout=300 + ) + assert result.returncode == 0 diff --git a/tests/pytorch/test_cuda_graphs.py b/tests/pytorch/test_cuda_graphs.py index 5a848dc0e8..8f85f57f32 100644 --- a/tests/pytorch/test_cuda_graphs.py +++ b/tests/pytorch/test_cuda_graphs.py @@ -5,6 +5,8 @@ from typing import Callable, Dict, Iterable, List, Tuple, Union import pytest import copy +import gc +import weakref import torch from transformer_engine.pytorch import ( @@ -994,6 +996,206 @@ def hook(module: torch.nn.Module) -> None: ] +def test_ordered_warmup_releases_consumed_outputs() -> None: + """Ordered warmup should only retain outputs until their corresponding backward.""" + + class OutputLifetimeModule(torch.nn.Module): + def __init__(self) -> None: + super().__init__() + self.previous_output = None + + def forward(self, input_: torch.Tensor) -> torch.Tensor: + is_warmup = not torch.cuda.is_current_stream_capturing() + if is_warmup and self.previous_output is not None: + assert self.previous_output() is None + output = input_ * 2 + if is_warmup: + self.previous_output = weakref.ref(output) + return output + + module = OutputLifetimeModule() + sample_args = tuple((torch.ones(4, 8, device="cuda", requires_grad=True),) for _ in range(2)) + graphed_callables = make_graphed_callables( + (module,), + sample_args, + num_warmup_iters=2, + _order=[1, -1, 1, -1], + _num_layers_per_chunk=[1], + ) + assert module.previous_output is not None + assert module.previous_output() is None + reset_graphs(graphed_callables) + + +def test_unordered_warmup_releases_consumed_outputs() -> None: + """Unordered warmup should release each output after its corresponding backward.""" + + class OutputLifetimeModule(torch.nn.Module): + def __init__(self, output_refs: list, module_idx: int) -> None: + super().__init__() + self.output_refs = output_refs + self.module_idx = module_idx + self.capture_started = False + + def forward(self, input_: torch.Tensor) -> torch.Tensor: + output = input_ * 2 + if torch.cuda.is_current_stream_capturing(): + self.capture_started = True + else: + self.output_refs[self.module_idx] = weakref.ref(output) + return output + + output_refs = [None, None] + modules = tuple(OutputLifetimeModule(output_refs, module_idx) for module_idx in range(2)) + + def first_module_backward_pre_hook(_module: torch.nn.Module) -> None: + if not modules[0].capture_started: + assert output_refs[1] is not None + assert output_refs[1]() is None + + graphed_callables = make_graphed_callables( + modules, + tuple((torch.ones(4, 8, device="cuda", requires_grad=True),) for _ in modules), + num_warmup_iters=2, + capture_time_hooks=[ + {"backward_pre_hooks": {0: first_module_backward_pre_hook}}, + None, + ], + ) + assert all(output_ref is not None and output_ref() is None for output_ref in output_refs) + reset_graphs(graphed_callables) + + +def test_inference_warmup_does_not_retain_outputs() -> None: + """Inference warmup should release outputs as soon as each forward returns.""" + + class OutputLifetimeModule(torch.nn.Module): + def __init__(self, previous_output: list) -> None: + super().__init__() + self.previous_output = previous_output + + def forward(self, input_: torch.Tensor) -> torch.Tensor: + is_warmup = not torch.cuda.is_current_stream_capturing() + if is_warmup and self.previous_output[0] is not None: + assert self.previous_output[0]() is None + output = input_ * 2 + if is_warmup: + self.previous_output[0] = weakref.ref(output) + return output + + previous_output = [None] + modules = tuple(OutputLifetimeModule(previous_output).eval() for _ in range(2)) + graphed_callables = make_graphed_callables( + modules, + tuple((torch.ones(4, 8, device="cuda"),) for _ in modules), + num_warmup_iters=2, + ) + assert previous_output[0] is not None + assert previous_output[0]() is None + reset_graphs(graphed_callables) + + +def test_reused_capture_buffers_release_outputs_after_backward() -> None: + """Capture locals must not keep weak-refed output buffers alive.""" + + class OutputLifetimeModule(torch.nn.Module): + def __init__(self) -> None: + super().__init__() + self.previous_capture_output = None + + def forward(self, input_: torch.Tensor) -> torch.Tensor: + if ( + torch.cuda.is_current_stream_capturing() + and self.previous_capture_output is not None + ): + assert self.previous_capture_output() is None + output = input_ * 2 + if torch.cuda.is_current_stream_capturing(): + self.previous_capture_output = weakref.ref(output) + return output + + module = OutputLifetimeModule() + sample_args = tuple((torch.ones(4, 8, device="cuda", requires_grad=True),) for _ in range(2)) + graphed_callables = make_graphed_callables( + (module,), + sample_args, + _order=[1, -1, 1, -1], + _num_layers_per_chunk=[1], + _reuse_graph_input_output_buffers=True, + ) + assert module.previous_capture_output is not None + assert module.previous_capture_output() is None + reset_graphs(graphed_callables) + + +def test_reset_releases_only_the_selected_callable() -> None: + """Reset releases one callable's graph state without retaining its peers.""" + + class CaptureOutputModule(torch.nn.Module): + def __init__(self) -> None: + super().__init__() + self.capture_output = None + + def forward(self, input_: torch.Tensor) -> torch.Tensor: + output = input_ * 2 + if torch.cuda.is_current_stream_capturing(): + self.capture_output = weakref.ref(output) + return output + + modules = tuple(CaptureOutputModule().cuda() for _ in range(2)) + graphed_callables = make_graphed_callables( + modules, + tuple((torch.ones(4, device="cuda", requires_grad=True),) for _ in modules), + ) + capture_outputs = tuple(module.capture_output for module in modules) + assert all(output is not None and output() is not None for output in capture_outputs) + + graphed_callables[0].reset() + graphed_callables[0].reset() + gc.collect() + assert capture_outputs[0]() is None + assert capture_outputs[1]() is not None + + output = graphed_callables[1](torch.randn(4, device="cuda", requires_grad=True)) + output.sum().backward() + del output + graphed_callables[1].reset() + gc.collect() + assert capture_outputs[1]() is None + + +@pytest.mark.parametrize("with_order", (False, True)) +def test_reset_rejects_all_replay_entry_points(with_order: bool) -> None: + """Reset is idempotent and terminal for forward and backward replay.""" + + class TestModule(torch.nn.Module): + def forward(self, input_: torch.Tensor) -> torch.Tensor: + return input_ * 2 + + module = TestModule().cuda() + sample_input = torch.ones(4, device="cuda", requires_grad=True) + graph_options = {} + if with_order: + graph_options = {"_order": [1, -1], "_num_layers_per_chunk": [1]} + graphed_callable = make_graphed_callables(module, (sample_input,), **graph_options) + output = graphed_callable(torch.randn_like(sample_input, requires_grad=True)) + torch.cuda.synchronize() + + graphed_callable.reset() + graphed_callable.reset() + if not with_order: + # The eager fallback for a different training state is invalid after reset too. + graphed_callable.eval() + + error = "has been reset and can no longer be used" + with pytest.raises(RuntimeError, match=error): + graphed_callable(torch.randn_like(sample_input, requires_grad=True)) + with pytest.raises(RuntimeError, match=error): + graphed_callable.backward_dw() + with pytest.raises(RuntimeError, match=error): + output.sum().backward() + + @pytest.mark.parametrize("with_order", (False, True)) def test_make_graphed_callables_with_capture_time_hooks(with_order: bool) -> None: """Test capture-time hooks around warmup and graph capture.""" diff --git a/tests/pytorch/test_fusible_ops.py b/tests/pytorch/test_fusible_ops.py index 6cd1fc3065..2adce717c9 100644 --- a/tests/pytorch/test_fusible_ops.py +++ b/tests/pytorch/test_fusible_ops.py @@ -146,6 +146,173 @@ def maybe_skip_quantization( pytest.skip("NVFP4 quantization is only supported with BF16 data") +def test_operation_fuser_caches_plans_by_grad_requirement(monkeypatch) -> None: + """Cache and restore fusion plans for checkpoint forward and recompute.""" + + # Count fusion-plan construction without depending on any particular real + # fusion implementation. Each distinct fusion configuration invokes this + # hook once, while a cache hit must bypass it entirely. + fusion_calls = 0 + + def track_fusion(ops, *, recipe): # pylint: disable=unused-argument + nonlocal fusion_calls + fusion_calls += 1 + # Preserve the operation list so this hook observes plan construction + # without changing the topology under test. + return ops + + # The fusion registries are class attributes shared by every OperationFuser. + # pytest's monkeypatch fixture restores all three after the test, preventing + # this synthetic fusion function from leaking into other tests. Keep only a + # joint forward-backward fusion hook so each plan build has one countable + # callback and no registered TE fusion can affect the result. + monkeypatch.setattr(OperationFuser, "forward_backward_fusion_functions", [track_fusion]) + monkeypatch.setattr(OperationFuser, "forward_fusion_functions", []) + monkeypatch.setattr(OperationFuser, "backward_fusion_functions", []) + + # One Identity op is enough to exercise the cache. With one basic op, + # first_op_requiring_backward has an intentionally simple interpretation: + # 0: backward starts at the Identity op; + # 1: the boundary is past the only op, so no backward work is required. + fuser = OperationFuser([te_ops.Identity()]) + x = torch.ones(1, requires_grad=True) + # maybe_fuse_ops expects one extra-input collection per basic op. Identity + # has no extra inputs, so its collection is an empty tuple. + extra_inputs = [()] + + # Phase 1: the original checkpointed forward runs with grad disabled. This + # is the first invocation, so the fuser must construct and cache the no-grad + # configuration. The runtime backward boundary is past the only op. + fuser.maybe_fuse_ops(False, None, x, extra_inputs) + assert fusion_calls == 1 + assert fuser.first_op_requiring_backward == 1 + no_grad_forward_ops = fuser._forward_ops + no_grad_backward_ops = fuser._backward_ops + + # Phase 2: backward replays the checkpointed region with grad enabled. The + # backward boundary is part of the fusion key, allowing future fusion rules + # to choose a training-specific topology. The first grad-enabled invocation + # therefore constructs and caches a second configuration. + fuser.maybe_fuse_ops(True, None, x, extra_inputs) + assert fusion_calls == 2 + assert fuser.first_op_requiring_backward == 0 + grad_forward_ops = fuser._forward_ops + grad_backward_ops = fuser._backward_ops + assert grad_forward_ops is not no_grad_forward_ops + assert grad_backward_ops is not no_grad_backward_ops + + # Phase 3: the next checkpointed forward must select the exact no-grad lists + # cached in phase 1. Before the cache was added, every boundary transition + # rebuilt the fused operations and called track_fusion again. + fuser.maybe_fuse_ops(False, None, x, extra_inputs) + assert fusion_calls == 2 + assert fuser.first_op_requiring_backward == 1 + assert fuser._forward_ops is no_grad_forward_ops + assert fuser._backward_ops is no_grad_backward_ops + + # Phase 4: another recomputation must likewise restore the grad-enabled + # lists from phase 2. The full alternating sequence has built only the two + # configurations represented by its two fusion keys. + fuser.maybe_fuse_ops(True, None, x, extra_inputs) + assert fusion_calls == 2 + assert fuser.first_op_requiring_backward == 0 + assert fuser._forward_ops is grad_forward_ops + assert fuser._backward_ops is grad_backward_ops + + +def test_operation_fuser_resets_recipe_state_independently_from_plan_cache(monkeypatch) -> None: + """Track recipe-state resets independently from fusion-plan construction.""" + + fusion_calls = 0 + + def track_fusion(ops, *, recipe): # pylint: disable=unused-argument + nonlocal fusion_calls + fusion_calls += 1 + return ops + + # Replace the process-wide fusion registries so one callback corresponds to + # one plan construction. monkeypatch restores the registries after the test. + monkeypatch.setattr(OperationFuser, "forward_backward_fusion_functions", [track_fusion]) + monkeypatch.setattr(OperationFuser, "forward_fusion_functions", []) + monkeypatch.setattr(OperationFuser, "backward_fusion_functions", []) + + op = te_ops.Identity() + reset_recipes = [] + first_forward_calls = 0 + + def track_recipe_reset(*, recipe): + reset_recipes.append(recipe) + + def track_first_forward(): + nonlocal first_forward_calls + first_forward_calls += 1 + + # Identity has no quantizers, so replace its state hooks with counters. This + # keeps the test CPU-only and isolates OperationFuser's reset decisions. + monkeypatch.setattr(op, "reset_recipe_state", track_recipe_reset) + monkeypatch.setattr(op, "pre_first_fuser_forward", track_first_forward) + + fuser = OperationFuser([op]) + x = torch.ones(1) + extra_inputs = [()] + + current_scaling = transformer_engine.common.recipe.Float8CurrentScaling(backward_override=None) + fuser.maybe_fuse_ops(False, current_scaling, x, extra_inputs) + assert reset_recipes == [current_scaling] + assert first_forward_calls == 1 + assert fusion_calls == 1 + + # A fresh but equivalent recipe does not invalidate state or the plan. + equivalent_current_scaling = transformer_engine.common.recipe.Float8CurrentScaling( + backward_override=None + ) + fuser.maybe_fuse_ops(False, equivalent_current_scaling, x, extra_inputs) + assert reset_recipes == [current_scaling] + assert first_forward_calls == 1 + assert fusion_calls == 1 + + # Backward override affects both recipe state and fusion topology, so it + # triggers one reset and constructs a distinct cached plan. + overridden_current_scaling = transformer_engine.common.recipe.Float8CurrentScaling( + backward_override="high_precision" + ) + fuser.maybe_fuse_ops(False, overridden_current_scaling, x, extra_inputs) + assert reset_recipes == [current_scaling, overridden_current_scaling] + assert first_forward_calls == 1 + assert fusion_calls == 2 + + delayed_scaling = transformer_engine.common.recipe.DelayedScaling( + amax_history_len=8, + backward_override=None, + ) + fuser.maybe_fuse_ops(False, delayed_scaling, x, extra_inputs) + assert reset_recipes == [current_scaling, overridden_current_scaling, delayed_scaling] + assert first_forward_calls == 1 + assert fusion_calls == 3 + + # Amax history length only affects delayed-scaling recipe state. Reset that + # state, but restore the existing DelayedScaling fusion plan from the cache. + resized_delayed_scaling = transformer_engine.common.recipe.DelayedScaling( + amax_history_len=16, + backward_override=None, + ) + fuser.maybe_fuse_ops(False, resized_delayed_scaling, x, extra_inputs) + assert reset_recipes == [ + current_scaling, + overridden_current_scaling, + delayed_scaling, + resized_delayed_scaling, + ] + assert first_forward_calls == 1 + assert fusion_calls == 3 + + # Repeating the exact recipe parameters performs neither operation again. + fuser.maybe_fuse_ops(False, resized_delayed_scaling, x, extra_inputs) + assert len(reset_recipes) == 4 + assert first_forward_calls == 1 + assert fusion_calls == 3 + + @torch.no_grad() def make_reference_and_test_tensors( shape: int | Iterable[int], diff --git a/tests/pytorch/test_grouped_mlp.py b/tests/pytorch/test_grouped_mlp.py index d48f7afae6..15c1ff6c51 100644 --- a/tests/pytorch/test_grouped_mlp.py +++ b/tests/pytorch/test_grouped_mlp.py @@ -5,6 +5,7 @@ from __future__ import annotations from collections.abc import Iterable +import contextlib import functools import os import math @@ -1067,6 +1068,46 @@ def train_step( class TestGroupedMLPFusedOp: """Tests for grouped MLP fused op""" + def test_fusion_requires_supported_grad_output_format(self, monkeypatch) -> None: + """Fuse E4M3 MXFP8 and NVFP4, but decline MXFP8 with an E5M2 backward.""" + from transformer_engine.common.recipe import Format, MXFP8BlockScaling, NVFP4BlockScaling + + fused_op_cls = grouped_mlp_module.GroupedMLP_CuTeGEMMGLU + monkeypatch.setattr(fused_op_cls, "is_supported", classmethod(lambda cls: True)) + + fc1 = te.ops.GroupedLinear(1, 64, 128, bias=False, device="cuda") + activation = te.ops.ScaledSwiGLU(glu_interleave_size=32) + fc2 = te.ops.GroupedLinear(1, 64, 64, bias=False, device="cuda") + ops = [fc1, activation, fc2] + + def fuse(recipe): + return grouped_mlp_module.fuse_grouped_mlp_ops( + ops, + recipe=recipe, + fused_op_cls=fused_op_cls, + ) + + def assert_fused(recipe): + fused_ops = fuse(recipe) + assert len(fused_ops) == 1 + fused_op = fused_ops[0] + assert isinstance(fused_op, fused_op_cls) + assert list(fused_op.basic_ops) == ops + + hybrid = MXFP8BlockScaling(fp8_format=Format.HYBRID) + assert fuse(hybrid) is ops + + e4m3 = MXFP8BlockScaling(fp8_format=Format.E4M3) + assert_fused(e4m3) + + # NVFP4 quantizes gradients to FP4, so the FP8 format must not gate it. Forcing the + # lookup to E5M2 is what an NVFP4 recipe would hit if the check were not MXFP8-only. + monkeypatch.setattr( + grouped_mlp_module, "get_fp8_torch_dtype", lambda *_, **__: torch.float8_e5m2 + ) + nvfp4 = NVFP4BlockScaling(disable_rht=False) + assert_fused(nvfp4) + @pytest.mark.parametrize("bias", (False, True)) @pytest.mark.parametrize("quantization", _grouped_mlp_quantization_list) @pytest.mark.parametrize("single_grouped_weight", (False, True)) @@ -2805,6 +2846,194 @@ def train_step( assert_close(graph_grad, param.grad, **tols) +class TestGroupedMLPDeterminism: + """Determinism coverage for the CuTe DSL fused grouped MLP. + + Only the dSReLU wrapper can make ``dprob`` bit-exact, and only from cuDNN FE 1.28.0 on. + Anything else must refuse a determinism request rather than run non-deterministically. + """ + + @pytest.fixture + def _restore_torch_determinism(self): + """``use_deterministic_algorithms`` is process-global, so put it back.""" + previous = torch.are_deterministic_algorithms_enabled() + yield + torch.use_deterministic_algorithms(previous) + + @pytest.mark.parametrize( + "allow_nondeterministic,torch_flag,expected", + ( + (None, False, False), # default: non-deterministic algorithms are allowed + ("1", False, False), + ("0", False, True), # the TE variable alone + (None, True, True), # the torch flag alone, which TE must not ignore + ("1", True, True), # ... including when the TE variable says otherwise + ("0", True, True), + ), + ) + def test_either_knob_requests_determinism( + self, + monkeypatch, + _restore_torch_determinism, + *, + allow_nondeterministic: Optional[str], + torch_flag: bool, + expected: bool, + ) -> None: + """``=1`` is the absence of a request, not a request for non-determinism.""" + if allow_nondeterministic is None: + monkeypatch.delenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", raising=False) + else: + monkeypatch.setenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", allow_nondeterministic) + torch.use_deterministic_algorithms(torch_flag) + assert grouped_mlp_module._deterministic_algorithms_required() is expected + + def test_only_the_srelu_path_can_be_deterministic(self) -> None: + """The capability belongs to the wrapper, not the environment. Needs no GPU.""" + glu = grouped_mlp_module.GroupedMLP_CuTeGEMMGLU + unary = grouped_mlp_module.GroupedMLP_CuTeGEMMUnary + assert glu.grouped_gemm_dactivation_is_deterministic() is False + assert isinstance(unary.grouped_gemm_dactivation_is_deterministic(), bool) + + @pytest.mark.skipif(not mxfp8_available, reason=reason_for_no_mxfp8) + @pytest.mark.parametrize( + "activation,fused_cls", + ( + ("scaled_srelu", grouped_mlp_module.GroupedMLP_CuTeGEMMUnary), + ("scaled_swiglu", grouped_mlp_module.GroupedMLP_CuTeGEMMGLU), + ), + ) + def test_determinism_either_runs_or_refuses( + self, monkeypatch, *, activation, fused_cls + ) -> None: + """A request TE cannot honor must fail loudly; one it can must still be correct.""" + if not fused_cls.is_supported(): + pytest.skip("MXFP8 fused grouped MLP is not supported on this system") + + monkeypatch.setenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "0") + expectation = ( + contextlib.nullcontext() + if fused_cls.grouped_gemm_dactivation_is_deterministic() + else pytest.raises(RuntimeError, match="dprob") + ) + with expectation: + TestGroupedMLPFusedOp().test_grouped_mlp( + bias=False, + hidden_size=128, + quantization="mxfp8", + single_grouped_weight=False, + activation=activation, + ) + + @pytest.mark.skipif(not mxfp8_available, reason=reason_for_no_mxfp8) + def test_scale_bias_refuses_under_the_torch_flag( + self, monkeypatch, _restore_torch_determinism + ) -> None: + """``scale_bias`` finishes ``dprob`` in a Triton kernel that reads only the env var. + + So the torch flag alone is the combination that used to pass this op's own check and + then reduce nondeterministically anyway, on a front-end new enough to say yes. + """ + fused_cls = grouped_mlp_module.GroupedMLP_CuTeGEMMUnary + if not fused_cls.is_supported(): + pytest.skip("MXFP8 fused grouped MLP is not supported on this system") + + monkeypatch.delenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", raising=False) + # warn_only so torch's own enforcement cannot raise first and mask what TE does. + torch.use_deterministic_algorithms(True, warn_only=True) + with pytest.raises(RuntimeError, match="dprob"): + TestGroupedMLPFusedOp().test_grouped_mlp( + bias=True, + hidden_size=128, + quantization="mxfp8", + single_grouped_weight=False, + activation="scaled_srelu", + ) + + @pytest.mark.skipif(not mxfp8_available, reason=reason_for_no_mxfp8) + def test_dprob_is_bit_exact_across_runs(self, monkeypatch) -> None: + """Repeated identical runs must give a bit-identical ``dprob``. + + An ulp of reordering passes every tolerance in this file, so only an exact + comparison across runs can see it. + """ + fused_cls = grouped_mlp_module.GroupedMLP_CuTeGEMMUnary + if not fused_cls.is_supported(): + pytest.skip("MXFP8 fused grouped MLP is not supported on this system") + if not fused_cls.grouped_gemm_dactivation_is_deterministic(): + pytest.skip("dSReLU determinism needs cuDNN frontend 1.28.0 or later") + + monkeypatch.setenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "0") + + device = torch.device("cuda") + dtype = torch.bfloat16 + # Measured on GB300, determinism off, 8 launches per shape (job 538058): this shape + # gives 7/7 runs differing from run 0, so the assertion below can actually fail. + # Shapes matter more than they look -- l=8 with the same n and tokens/group varies + # only 2/7, which an 8-run sample reports as stable often enough to be useless, and + # cudnn-frontend#521 measured its own l=4 / [256]*4 / n=512 as never varying. + group_size = 16 + hidden_size = 2048 + tokens_per_group = 1024 + split_sizes = torch.tensor([tokens_per_group] * group_size, dtype=torch.int, device=device) + num_tokens = tokens_per_group * group_size + + recipe = make_recipe("mxfp8") + + # Plain random tensors, not make_reference_and_test_tensors: this test compares two + # runs against each other, never against a reference, so the fp64 companion and the + # MXFP8 representability round-trip would both be allocated and thrown away. + def _rand(*shape, requires_grad=True) -> torch.Tensor: + out = torch.empty(shape, dtype=dtype, device=device).uniform_(-0.25, 0.25) + return out.requires_grad_() if requires_grad else out + + x = _rand(num_tokens, hidden_size) + dy = _rand(num_tokens, hidden_size, requires_grad=False) + probs = _rand(num_tokens) + + # No bias, or probs.grad comes from the Triton dbias kernel instead of cuDNN. + with te.quantized_model_init(enabled=True, recipe=recipe): + module = te.ops.Sequential( + te.ops.GroupedLinear( + group_size, hidden_size, hidden_size, bias=False, device=device, dtype=dtype + ), + te.ops.ScaledSReLU(), + te.ops.GroupedLinear( + group_size, hidden_size, hidden_size, bias=False, device=device, dtype=dtype + ), + ) + + def _run() -> torch.Tensor: + x.grad = None + probs.grad = None + with te.autocast(enabled=True, recipe=recipe): + y = module(x, split_sizes, probs, split_sizes) + y.backward(dy) + return probs.grad.detach().clone() + + runs = [_run()] + # Without the fusion there is no cuDNN dprob and the comparison proves nothing. + forward_ops = module._module_groups[0]._forward_ops + assert len(forward_ops) == 1 + assert isinstance(forward_ops[0][0], fused_cls) + # More than two, as cudnn-frontend#521 does: the cross-CTA order that determinism + # removes is set by the scheduler, so two runs can agree by luck. + runs += [_run() for _ in range(int(os.getenv("NVTE_TEST_DETERMINISM_REPEATS", "4")) - 1)] + torch.cuda.synchronize() + + assert torch.isfinite(runs[0]).all(), "dprob is not finite; the comparison would be moot" + # Bytes, not values: torch.equal calls +0.0 and -0.0 equal, and a change in reduction + # order can produce exactly that. Weight grads are excluded from the comparison -- + # the CuTe DSL wgrad kernel has its own K-split atomics, which this change leaves. + for index, later in enumerate(runs[1:], start=1): + assert torch.equal( + runs[0].contiguous().view(torch.uint8), later.contiguous().view(torch.uint8) + ), ( + f"dprob differs between run 0 and run {index} under determinism; max |delta| =" + f" {(runs[0].float() - later.float()).abs().max().item()}" + ) + + def test_grouped_gemm_quant_cute_matches_mxfp8_quantized() -> None: if not mxfp8_available: pytest.skip(reason_for_no_mxfp8) diff --git a/tests/pytorch/test_models.py b/tests/pytorch/test_models.py new file mode 100644 index 0000000000..cd1903e79f --- /dev/null +++ b/tests/pytorch/test_models.py @@ -0,0 +1,176 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +import math + +import pytest +import torch + +from transformer_engine.pytorch.utils import deinterleave_glu_tensor +from transformer_engine.pytorch.models import DeepSeekV3MoE, MultiLatentAttention + +SEQ_LEN = 128 +BATCH = 2 +HIDDEN = 256 +HEADS = 4 +DTYPE = torch.bfloat16 + +MLA_KWARGS = dict( + q_lora_rank=96, + kv_lora_rank=64, + qk_nope_head_dim=64, + qk_rope_head_dim=32, + v_head_dim=64, +) + + +def _input(requires_grad=True): + torch.manual_seed(1234) + return torch.randn( + SEQ_LEN, BATCH, HIDDEN, dtype=DTYPE, device="cuda", requires_grad=requires_grad + ) + + +def test_mla_rope_triton_matches_pytorch(): + from transformer_engine.pytorch.models.deepseek_v3 import mla_rope + + if not mla_rope.HAVE_TRITON: + pytest.skip("Triton unavailable") + s, b, h = 64, 2, 4 + nope, rope, vdim = 64, 32, 64 + cos, sin = mla_rope.build_rope_tables(s, rope, device="cuda") + + torch.manual_seed(0) + q_leaf = torch.randn(s, b, h, nope + rope, device="cuda", requires_grad=True) + kv_leaf = torch.randn(s, b, h, nope + vdim, device="cuda", requires_grad=True) + pos_leaf = torch.randn(s, b, 1, rope, device="cuda", requires_grad=True) + grad_q = torch.randn(s, b, h, nope + rope, device="cuda") + grad_k = torch.randn(s, b, h, nope + rope, device="cuda") + grad_v = torch.randn(s, b, h, vdim, device="cuda") + + def run(fmt): + # non-leaf copies: the Triton q kernel rotates in place + q, kv, pos = q_leaf * 1.0, kv_leaf * 1.0, pos_leaf * 1.0 + q_out = mla_rope.apply_mla_rope_q(q, cos, sin, nope, rope, fmt) + k_out, v_out = mla_rope.apply_mla_rope_kv(kv, pos, cos, sin, nope, rope, vdim, fmt) + # fresh grad clones: the Triton q backward modifies its input grad in place + torch.autograd.backward( + [q_out, k_out, v_out], [grad_q.clone(), grad_k.clone(), grad_v.clone()] + ) + grads = (q_leaf.grad.clone(), kv_leaf.grad.clone(), pos_leaf.grad.clone()) + q_leaf.grad = kv_leaf.grad = pos_leaf.grad = None + return (q_out.clone(), k_out, v_out), grads + + (q_t, k_t, v_t), grads_t = run("sbhd") + + seq_dim = 0 + q_ref = torch.cat( + ( + (q_leaf * 1.0)[..., :nope], + mla_rope._rotate_interleaved_to_neox((q_leaf * 1.0)[..., nope:], cos, sin, seq_dim), + ), + dim=-1, + ) + k_ref = torch.cat( + ( + (kv_leaf * 1.0)[..., :nope], + mla_rope._rotate_interleaved_to_neox(pos_leaf * 1.0, cos, sin, seq_dim).expand( + s, b, h, rope + ), + ), + dim=-1, + ) + v_ref = (kv_leaf * 1.0)[..., nope:] + torch.autograd.backward([q_ref, k_ref, v_ref], [grad_q.clone(), grad_k.clone(), grad_v.clone()]) + + torch.testing.assert_close(q_t, q_ref, rtol=1e-5, atol=1e-5) + torch.testing.assert_close(k_t, k_ref, rtol=1e-5, atol=1e-5) + torch.testing.assert_close(v_t, v_ref, rtol=1e-5, atol=1e-5) + torch.testing.assert_close(grads_t[0], q_leaf.grad, rtol=1e-5, atol=1e-5) + torch.testing.assert_close(grads_t[1], kv_leaf.grad, rtol=1e-5, atol=1e-5) + torch.testing.assert_close(grads_t[2], pos_leaf.grad, rtol=1e-5, atol=1e-5) + + +def test_rope_tables_yarn(): + from transformer_engine.pytorch.models.deepseek_v3 import mla_rope + + s, rope = 8192, 64 + cos, sin = mla_rope.build_rope_tables(s, rope, device="cuda") + cos_none, sin_none = mla_rope.build_rope_tables(s, rope, device="cuda", scaling_factor=None) + assert torch.equal(cos, cos_none) and torch.equal(sin, sin_none) + + yarn = dict(scaling_factor=40.0, original_max_position_embeddings=4096) + cos_y, sin_y = mla_rope.build_rope_tables(s, rope, device="cuda", **yarn) + factor = mla_rope.yarn_concentration_factor(40.0, 1.0, 0.0) + assert factor == pytest.approx(0.1 * math.log(40.0) + 1.0) + # amplitude scaled by the concentration factor + torch.testing.assert_close(cos_y**2 + sin_y**2, torch.full_like(cos_y, factor**2)) + # high-frequency dims untouched, low-frequency dims interpolated by 1/scaling_factor + torch.testing.assert_close(cos_y[:, 0] / factor, cos[:, 0]) + angle_y = torch.atan2(sin_y[:, rope // 2 - 1], cos_y[:, rope // 2 - 1]) + angle = torch.atan2(sin[:, rope // 2 - 1], cos[:, rope // 2 - 1]) + torch.testing.assert_close(angle_y[:64], angle[:64] / 40.0, atol=1e-4, rtol=0) + + +@pytest.mark.parametrize("mscale_all_dim", [0.0, 1.0]) +def test_mla_yarn_softmax_scale(mscale_all_dim): + mla = MultiLatentAttention( + HIDDEN, + HEADS, + params_dtype=DTYPE, + rope_scaling_factor=40.0, + original_max_position_embeddings=64, + mscale_all_dim=mscale_all_dim, + **MLA_KWARGS, + ) + m = 0.1 * mscale_all_dim * math.log(40.0) + 1.0 + qk_head_dim = MLA_KWARGS["qk_nope_head_dim"] + MLA_KWARGS["qk_rope_head_dim"] + assert mla.softmax_scale == pytest.approx(m * m / math.sqrt(qk_head_dim)) + + +@pytest.mark.parametrize("shared", [False, True], ids=["no_shared", "shared"]) +@pytest.mark.parametrize("grouped", [False, True], ids=["ungrouped", "grouped"]) +@pytest.mark.parametrize("topk", [2, 4]) +def test_moe_matches_dense_reference(shared, grouped, topk): + """Routed output must equal the prob-weighted sum of the selected expert MLPs.""" + torch.manual_seed(0) + num_experts = 4 + moe = DeepSeekV3MoE( + HIDDEN, + moe_ffn_hidden_size=128, + num_experts=num_experts, + topk=topk, + num_groups=2 if grouped else None, + group_topk=topk // 2 if grouped else None, + shared_expert_ffn_hidden_size=128 if shared else None, + params_dtype=DTYPE, + ) + x = _input() + out = moe(x) + assert out.shape == x.shape + out.sum().backward() + assert torch.isfinite(x.grad).all() + + tokens = x.detach().reshape(-1, HIDDEN) + probs, _ = moe._route(moe.gate(tokens).float()) + assert (probs > 0).sum(dim=1).eq(topk).all() + assert moe._last_tokens_per_expert.sum().item() == tokens.shape[0] * topk + + fc1, _, fc2 = moe.experts + ref = torch.zeros_like(tokens) + for e in range(num_experts): + w1 = deinterleave_glu_tensor(getattr(fc1, f"weight{e}"), 32) + w2 = getattr(fc2, f"weight{e}") + gate_part, lin_part = (tokens @ w1.t()).chunk(2, dim=-1) + act = torch.nn.functional.silu(gate_part.float()) * lin_part.float() + ref += (act.to(DTYPE) * probs[:, e : e + 1].to(DTYPE)) @ w2.t() + if shared: + ref += moe.shared_expert(tokens) + torch.testing.assert_close(out.reshape(-1, HIDDEN), ref, rtol=0.05, atol=0.05) + + bias_before = moe.expert_bias.clone() + moe.update_expert_bias() + assert torch.isfinite(moe.expert_bias).all() + if topk < num_experts: + assert not torch.equal(bias_before, moe.expert_bias) diff --git a/tests/pytorch/test_sanity.py b/tests/pytorch/test_sanity.py index c9b620fa1e..60d17e39af 100644 --- a/tests/pytorch/test_sanity.py +++ b/tests/pytorch/test_sanity.py @@ -35,6 +35,7 @@ is_bf16_available, ) from transformer_engine.common import recipe +from transformer_engine.pytorch.models import DeepSeekV3Layer from transformer_engine.pytorch.cpp_extensions import general_gemm from transformer_engine.pytorch.tensor.utils import replace_raw_data from transformer_engine.pytorch.module import is_module_grouped_tensor_path_supported @@ -736,6 +737,39 @@ def test_sanity_layernorm_mlp( _test_sanity_common(block, dtype, config, fp8_recipe, skip_wgrad, skip_dgrad, microbatching) +@pytest.mark.parametrize("dtype", param_types) +@pytest.mark.parametrize("fp8_recipe", fp8_recipes, ids=recipe_id) +@pytest.mark.parametrize("moe", all_boolean, ids=["dense", "moe"]) +def test_sanity_deepseek_v3_layer(dtype, fp8_recipe, moe): + config = model_configs["small"] + + if fp8_recipe is not None: + if not is_fp8_supported(config): + pytest.skip("Model config does not support FP8") + if fp8_recipe.nvfp4() and dtype == torch.float16: + pytest.skip("FP16 output for NVFP4 not supported") + + mlp_kwargs = ( + dict(num_experts=4, topk=2, moe_ffn_hidden_size=32, shared_expert_ffn_hidden_size=32) + if moe + else dict(ffn_hidden_size=4 * config.hidden_size) + ) + block = DeepSeekV3Layer( + config.hidden_size, + config.num_heads, + q_lora_rank=16, + kv_lora_rank=16, + qk_nope_head_dim=16, + qk_rope_head_dim=16, + v_head_dim=16, + params_dtype=dtype, + device="cuda", + **mlp_kwargs, + ) + + _test_sanity_e2e(block, dtype, config, fp8_recipe, skip_wgrad=False) + + @pytest.mark.parametrize("dtype", param_types) @pytest.mark.parametrize("fp8_recipe", fp8_recipes, ids=recipe_id) @pytest.mark.parametrize("model", ["small"]) diff --git a/transformer_engine/common/cast/core/grouped_tma.cuh b/transformer_engine/common/cast/core/grouped_tma.cuh index 61218d654a..8603fd1fd2 100644 --- a/transformer_engine/common/cast/core/grouped_tma.cuh +++ b/transformer_engine/common/cast/core/grouped_tma.cuh @@ -53,12 +53,6 @@ inline bool dimensions_supported_by_TMA(const Tensor *const t) { return cols % alignment_requirement == 0; } -__device__ __forceinline__ unsigned char *align_smem_ptr_per_TMA_requirements(unsigned char *p) { - size_t addr = reinterpret_cast(p); - addr = (addr + TMA_SHMEM_ALIGNMENT - 1) & ~(TMA_SHMEM_ALIGNMENT - 1); - return reinterpret_cast(addr); -} - // Copies the base tensor map to shmem, modifies the copy, stores the modified tensor map at index __device__ __forceinline__ void modify_base_tensor_map(const CUtensorMap base_tensor_map, CUtensorMap *global_tensor_map, diff --git a/transformer_engine/common/cast/fp8_blockwise/group_quantize_fp8_blockwise.cuh b/transformer_engine/common/cast/fp8_blockwise/group_quantize_fp8_blockwise.cuh index 203f569471..31feaf833d 100644 --- a/transformer_engine/common/cast/fp8_blockwise/group_quantize_fp8_blockwise.cuh +++ b/transformer_engine/common/cast/fp8_blockwise/group_quantize_fp8_blockwise.cuh @@ -314,9 +314,9 @@ __global__ void __launch_bounds__(kThreadsPerBlock, 4) group_block_scaled_2d_tma // Dynamic smem holds the IType input tile (TMA dest, must be 128 B aligned). // warp_amaxes and tma_mbar are static smem. - extern __shared__ unsigned char smem_raw_2d_tma[]; - IType(*smem_in)[kTileDim] = reinterpret_cast( - common::align_smem_ptr_per_TMA_requirements(smem_raw_2d_tma)); + extern __shared__ char smem_raw_2d_tma[]; + IType(*smem_in)[kTileDim] = + reinterpret_cast(align_up(smem_raw_2d_tma, TMA_SHMEM_ALIGNMENT)); __shared__ CType warp_amaxes[kNumWarps]; __shared__ size_t warp_offset_partials[kNumWarps]; @@ -603,8 +603,8 @@ __global__ void __launch_bounds__(kThreadsPerBlock) group_block_scaled_1d_tma_ke // Dynamic smem: IType[kTileDim][kTileDim], 128 B aligned for TMA. Static smem // (smem_T when CW, tma_mbar) lives outside the dynamic region. - extern __shared__ unsigned char smem_raw_1d_tma[]; - unsigned char* smem_base = common::align_smem_ptr_per_TMA_requirements(smem_raw_1d_tma); + extern __shared__ char smem_raw_1d_tma[]; + char* smem_base = align_up(smem_raw_1d_tma, TMA_SHMEM_ALIGNMENT); IType(*smem)[kTileDim] = reinterpret_cast(smem_base); __shared__ uint64_t tma_mbar; diff --git a/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8.cuh b/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8.cuh index 980a77db0a..b0383d95f3 100644 --- a/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8.cuh +++ b/transformer_engine/common/cast/mxfp8/group_quantize_mxfp8.cuh @@ -704,8 +704,8 @@ __global__ void __launch_bounds__(CastTraits::THREADS_PER_CHUNK) group_quantize_ constexpr size_t out_mem_rowwise = (ROWWISE_SCALING ? buff_size_aligned_out : 0); // The destination shared memory buffer of a bulk tensor operation should be 16-byte aligned - extern __shared__ unsigned char dynamic_shmem[]; - unsigned char *dshmem = align_smem_ptr_per_TMA_requirements(dynamic_shmem); + extern __shared__ char dynamic_shmem[]; + char *dshmem = align_up(dynamic_shmem, TMA_SHMEM_ALIGNMENT); // The destination shared memory buffer of a bulk tensor operation should be 16-byte aligned IType *sIn_ptr = reinterpret_cast(dshmem); diff --git a/transformer_engine/common/cast/mxfp8/group_scaled_swiglu_mxfp8.cuh b/transformer_engine/common/cast/mxfp8/group_scaled_swiglu_mxfp8.cuh index 24f84fa359..878fc93107 100644 --- a/transformer_engine/common/cast/mxfp8/group_scaled_swiglu_mxfp8.cuh +++ b/transformer_engine/common/cast/mxfp8/group_scaled_swiglu_mxfp8.cuh @@ -209,8 +209,8 @@ __global__ void __launch_bounds__(THREADS_PER_CHUNK) group_scaled_swiglu_mxfp8_k DIVUP_TO_MULTIPLE(CHUNK_DIM_Y * sizeof(float), TMA_SHMEM_ALIGNMENT); // shmem layout: [act input][gate input][colwise output][prob] - extern __shared__ unsigned char dynamic_shmem[]; - unsigned char *dshmem = align_smem_ptr_per_TMA_requirements(dynamic_shmem); + extern __shared__ char dynamic_shmem[]; + char *dshmem = align_up(dynamic_shmem, TMA_SHMEM_ALIGNMENT); IType *sInAct_ptr = reinterpret_cast(dshmem); IType *sInGate_ptr = reinterpret_cast(dshmem + buff_size_aligned_in); diff --git a/transformer_engine/common/cast/nvfp4/specialized/quantize_transpose_nvfp4_tuned_1D.cuh b/transformer_engine/common/cast/nvfp4/specialized/quantize_transpose_nvfp4_tuned_1D.cuh index cdd0d4916a..58af1f7938 100644 --- a/transformer_engine/common/cast/nvfp4/specialized/quantize_transpose_nvfp4_tuned_1D.cuh +++ b/transformer_engine/common/cast/nvfp4/specialized/quantize_transpose_nvfp4_tuned_1D.cuh @@ -410,8 +410,8 @@ __global__ void __launch_bounds__(THREADS_NUM) quantize_transpose_nvfp4_tuned_1D TunableConfig::CHUNK_DIM_Y * SCALES_PER_CHUNK_X * sizeof(nvfp4_scale_t), TMA_SHMEM_ALIGNMENT); // The destination shared memory buffer of a bulk tensor operation should be 16-byte aligned - extern __shared__ unsigned char dynamic_shmem[]; - unsigned char *dshmem = common::align_smem_ptr_per_TMA_requirements(dynamic_shmem); + extern __shared__ char dynamic_shmem[]; + char *dshmem = align_up(dynamic_shmem, TMA_SHMEM_ALIGNMENT); IType *sIn_ptr = reinterpret_cast(dshmem); fp4e2m1x2 *sOut_ptr = reinterpret_cast(dshmem + in_mem); diff --git a/transformer_engine/common/recipe/__init__.py b/transformer_engine/common/recipe/__init__.py index 5d5ce1f6cf..128e8280bb 100644 --- a/transformer_engine/common/recipe/__init__.py +++ b/transformer_engine/common/recipe/__init__.py @@ -639,13 +639,18 @@ class CustomRecipe(Recipe): ---------- qfactory : Callable Factory callable that returns a quantizer instance *or* a - ``QuantizerRequest`` subclass for a given ``QuantizerRole``. + ``QuantizerRequest`` subclass for a given optional ``QuantizerRole``. The callable is invoked as:: qfactory( - role: QuantizerRole, + role: Optional[QuantizerRole], ) -> Union[Quantizer, QuantizerRequest] + Boundary slots may provide ``None`` or a role with empty fields. The + factory must return a valid object for every call. Return an + ``IdentityQuantizer`` for an intentional high-precision slot instead + of returning ``None``. + ``QuantizerRole`` is a frozen dataclass with the following fields: - ``module_type`` (str): module type (empty string when not set), e.g. @@ -663,7 +668,8 @@ class CustomRecipe(Recipe): See ``transformer_engine.pytorch.quantization.QuantizerRole`` and ``transformer_engine.pytorch.quantization.DelayedScalingRequest`` - for full documentation. + for API details. See :ref:`heterogeneous-quantization-recipes` for + construction rules and direction mapping. backward_override : {None, 'high_precision', 'dequantized'}, default = None Backward precision mode. None does not modify backward behavior, diff --git a/transformer_engine/pytorch/__init__.py b/transformer_engine/pytorch/__init__.py index 576fb57c5e..5940ce7b3d 100644 --- a/transformer_engine/pytorch/__init__.py +++ b/transformer_engine/pytorch/__init__.py @@ -35,6 +35,7 @@ from transformer_engine.pytorch.attention import InferenceParams from transformer_engine.pytorch.attention import RotaryPositionEmbedding from transformer_engine.pytorch.transformer import TransformerLayer +from transformer_engine.pytorch import models from transformer_engine.pytorch.permutation import ( moe_permute, moe_permute_with_probs, diff --git a/transformer_engine/pytorch/csrc/extensions/ep.cpp b/transformer_engine/pytorch/csrc/extensions/ep.cpp index c74d4ddb6d..97bc70bddd 100644 --- a/transformer_engine/pytorch/csrc/extensions/ep.cpp +++ b/transformer_engine/pytorch/csrc/extensions/ep.cpp @@ -59,6 +59,21 @@ std::atomic g_zero_copy_enabled{false}; // is not symm-mem-backed; the backend treats it as "no window, use staged copy". constexpr NVTECommWindow kNoWindow = {nullptr, 0}; +#ifdef NCCL_HAS_SYMMEM_SUPPORT +// Offset of a symm-mem allocation relative to the start of its NCCL window. +// Newer torch places the signal pad at the front of the allocation and exposes +// get_window_offset() for it; on older torch the window starts at the buffer +// base, where get_offset() is already window-relative. +template +auto symm_mem_window_offset(T* sm, int) -> decltype(sm->get_window_offset()) { + return sm->get_window_offset(); +} +template +size_t symm_mem_window_offset(T* sm, ...) { + return sm->get_offset(); +} +#endif + // Resolve ``t`` to an NCCL symm-mem window for the zero-copy one-sided path. // Returns ``kNoWindow`` when symm-mem support isn't compiled in, zero-copy is // disabled, no group is set, or ``t`` isn't symm-mem-backed; callers pass the @@ -78,11 +93,11 @@ NVTECommWindow maybe_make_window(const at::Tensor& t) { NVTE_CHECK(nccl_sm != nullptr, "Symm-mem backend mismatch: expected NCCLSymmetricMemory. Set the backend to " "\"NCCL\" before allocating EP payload buffers."); - // NCCL EP consumes window-relative offsets (the NCCL window starts at the signal pad, - // not at the buffer base). get_window_offset() = buffer_offset + get_offset(); add - // ``t``'s own storage offset for slice/view positioning. + // NCCL EP consumes window-relative offsets. Add ``t``'s own storage offset so a + // slice/view of a symm-mem allocation (e.g. the scale region carved from a shared + // recv buffer) resolves to its true position in the window. const uint64_t offset = - static_cast(nccl_sm->get_window_offset()) + + static_cast(symm_mem_window_offset(nccl_sm, 0)) + static_cast(t.storage_offset()) * static_cast(t.element_size()); return NVTECommWindow{static_cast(nccl_sm->get_window()), offset}; #else diff --git a/transformer_engine/pytorch/graph.py b/transformer_engine/pytorch/graph.py index b298b3d8ff..04fa56721d 100644 --- a/transformer_engine/pytorch/graph.py +++ b/transformer_engine/pytorch/graph.py @@ -8,7 +8,7 @@ import gc import warnings from math import ceil -from typing import Any, Callable, Dict, List, Optional, Tuple, TypeVar, Union +from typing import Any, Callable, Dict, List, NamedTuple, Optional, Tuple, TypeVar, Union import torch from torch.utils._pytree import tree_flatten as _tree_flatten @@ -44,6 +44,13 @@ ) +class _GraphedCallableHelpers(NamedTuple): + """Lifecycle helpers owned by one graphed callable invocation.""" + + ensure_not_reset: Callable[[], None] + release_static_state: Callable[[], None] + + def set_capture_start() -> None: """Record beginning of `make_graphed_callables`.""" global _IS_GRAPH_CAPTURING @@ -633,13 +640,17 @@ def _run_warmup_backward(func_idx, func, outputs, warmup_iter, callable_idx): warmup_outputs = [] for func_idx, func in zip(warmup_func_idx, warmup_func): outputs = _run_warmup_forward(func_idx, func, func_idx) - warmup_outputs.append((func_idx, func, outputs)) - if is_training: - for func_idx, func, outputs in reversed(warmup_outputs): - _run_warmup_backward(func_idx, func, outputs, warmup_iter, func_idx) + if is_training: + warmup_outputs.append((func_idx, func, outputs)) + else: + del outputs + while warmup_outputs: + func_idx, func, outputs = warmup_outputs.pop() + _run_warmup_backward(func_idx, func, outputs, warmup_iter, func_idx) + del outputs else: # Follow _order exactly, mirroring the capture phase. - per_fwd_outputs = {} # per_callable_fwd_idx -> flattened outputs + per_fwd_outputs = {} # per_callable_fwd_idx -> outstanding flattened outputs fwd_idx = [0] * num_model_chunks bwd_idx = [0] * num_model_chunks for c_id in _order: @@ -653,7 +664,10 @@ def _run_warmup_backward(func_idx, func, outputs, warmup_iter, callable_idx): ) + (fwd_idx[m_chunk] * _num_layers_per_chunk[m_chunk] + l_no) func = callables[callable_idx] outputs = _run_warmup_forward(per_callable_fwd_idx, func, callable_idx) - per_fwd_outputs[per_callable_fwd_idx] = outputs + if is_training: + per_fwd_outputs[per_callable_fwd_idx] = outputs + else: + del outputs fwd_idx[m_chunk] += 1 elif ceil(c_id) == c_id: # Backward pass for chunk -c_id. @@ -665,10 +679,11 @@ def _run_warmup_backward(func_idx, func, outputs, warmup_iter, callable_idx): _prefix_num_layers[m_chunk] * num_microbatches ) + (bwd_idx[m_chunk] * _num_layers_per_chunk[m_chunk] + l_no) func = callables[callable_idx] - outputs = per_fwd_outputs[per_callable_bwd_idx] + outputs = per_fwd_outputs.pop(per_callable_bwd_idx) _run_warmup_backward( per_callable_bwd_idx, func, outputs, warmup_iter, callable_idx ) + del outputs bwd_idx[m_chunk] += 1 if post_warmup_hook is not None: @@ -729,6 +744,7 @@ def _run_warmup_backward(func_idx, func, outputs, warmup_iter, callable_idx): per_callable_static_outputs[per_callable_fwd_idx] = tuple(flatten_outputs) per_callable_output_unflatten_spec[per_callable_fwd_idx] = spec graph_callables[per_callable_fwd_idx] = func + del outputs, flatten_outputs fwd_idx[m_chunk] += 1 else: # Capture backward graph for model chunk c_id, microbatch bwd_idx[-c_id-1] @@ -917,6 +933,11 @@ def _run_warmup_backward(func_idx, func, outputs, warmup_iter, callable_idx): per_callable_static_grad_inputs[idx] ) previous_chunk_last_callable_bwd_idx = per_callable_bwd_idx + + # The per-callable containers now own all tensors that must survive + # capture. Drop local strong references so weak-refed graph buffers can + # be returned to the shared CUDA graph pool before the next capture. + del static_outputs, static_grad_inputs, grad_inputs if ceil(c_id) == c_id: bwd_idx[m_chunk] += 1 else: @@ -1028,12 +1049,22 @@ def make_graphed_autograd_function( static_grad_inputs, returned_param_grad_clone_slots, ): + is_reset = False + + def ensure_not_reset(): + """Reject replay after this callable's graph state has been released.""" + if is_reset: + raise RuntimeError( + "This graphed callable has been reset and can no longer be used." + ) + class Graphed(torch.autograd.Function): """Autograd function for graph replay.""" @staticmethod def forward(ctx, skip_fp8_weight_update, cuda_graph_stream, cuda_graph_event, *inputs): # pylint: disable=missing-function-docstring + ensure_not_reset() # Set flag for whether to update FP8 weight updates ctx.is_first_module = FP8GlobalStateManager.is_first_fp8_module() @@ -1071,6 +1102,7 @@ def forward(ctx, skip_fp8_weight_update, cuda_graph_stream, cuda_graph_event, *i @torch.autograd.function.once_differentiable def backward(ctx, *grads): # pylint: disable=missing-function-docstring + ensure_not_reset() # Replay backward graph if len(grads) != len(static_grad_outputs): @@ -1119,6 +1151,7 @@ def backward(ctx, *grads): return (None, None, None) + tuple(grad_inputs) def functionalized(*user_args, **user_kwargs): + ensure_not_reset() # Decide whether to update FP8 weights skip_fp8_weight_update = None @@ -1170,16 +1203,43 @@ def functionalized(*user_args, **user_kwargs): ) return _tree_unflatten(out, output_unflatten_spec) - return functionalized - - def make_graphed_attribute_functions(graph_idx): - # Get te modules for current graph + def release_static_state(): + """Release per-callable state captured by replay closures.""" + nonlocal fwd_graph, bwd_graph, is_reset + nonlocal module_params + nonlocal static_input_surface, static_outputs + nonlocal static_grad_outputs, static_grad_inputs + + is_reset = True + + # Drop the per-callable references that can own graph-pool storage. + fwd_graph = None + bwd_graph = None + module_params = () + static_input_surface = () + static_outputs = () + static_grad_outputs = () + static_grad_inputs = () + + helpers = _GraphedCallableHelpers( + ensure_not_reset=ensure_not_reset, + release_static_state=release_static_state, + ) + return functionalized, helpers + + def make_graphed_attribute_functions(graph_idx, helpers): + # Snapshot per-callable state so returned closures do not retain the outer lists. + fwd_graph = fwd_graphs[graph_idx] + bwd_graph = bwd_graphs[graph_idx] + bwd_dw_graph = bwd_dw_graphs[graph_idx] + need_bwd_dw = need_bwd_dw_graph.get(graph_idx, False) te_modules = visited_te_modules.get(graph_idx, set()) # Attach backward_dw as an attribute to the graphed callable. def backward_dw(): - if need_bwd_dw_graph.get(graph_idx, False): - bwd_dw_graphs[graph_idx].replay() + helpers.ensure_not_reset() + if need_bwd_dw: + bwd_dw_graph.replay() # Trigger the grad accumulation hook for wgrad graphs. for module in te_modules: @@ -1191,16 +1251,24 @@ def backward_dw(): # Attach reset as an attribute to the graphed callable. def reset(): - fwd_graphs[graph_idx].reset() - bwd_graphs[graph_idx].reset() - bwd_dw_graphs[graph_idx].reset() + nonlocal fwd_graph, bwd_graph, bwd_dw_graph, te_modules + + for graph in (fwd_graph, bwd_graph, bwd_dw_graph): + if graph is not None: + graph.reset() + + fwd_graph = None + bwd_graph = None + bwd_dw_graph = None + te_modules = () + helpers.release_static_state() return backward_dw, reset # Put together the final graphed callables ret = [] for i in range(len(sample_args)): - graphed = make_graphed_autograd_function( + graphed, helpers = make_graphed_autograd_function( fwd_graphs[i], bwd_graphs[i], per_callable_module_params[i], @@ -1218,8 +1286,17 @@ def reset(): te_modules = visited_te_modules.get(i, set()) if isinstance(func, torch.nn.Module): - def make_graphed_forward(func, graph_training_state, graphed, orig_fwd, te_modules): + def make_graphed_forward( + func, + graph_training_state, + graphed, + orig_fwd, + te_modules, + helpers, + ): def new_fwd(*user_args, **user_kwargs): + helpers.ensure_not_reset() + # If the module's training-or-eval state matches what we graphed, # run the graph, otherwise run the original forward method if func.training == graph_training_state: @@ -1264,7 +1341,14 @@ def new_fwd(*user_args, **user_kwargs): return new_fwd - forward = make_graphed_forward(func, func.training, graphed, func.forward, te_modules) + forward = make_graphed_forward( + func, + func.training, + graphed, + func.forward, + te_modules, + helpers, + ) if _order is None: func.forward = forward ret.append(func) @@ -1273,7 +1357,10 @@ def new_fwd(*user_args, **user_kwargs): else: ret.append(graphed) - backward_dw_func, reset_func = make_graphed_attribute_functions(i) + backward_dw_func, reset_func = make_graphed_attribute_functions( + i, + helpers, + ) setattr(ret[-1], "backward_dw", backward_dw_func) setattr(ret[-1], "reset", reset_func) diff --git a/transformer_engine/pytorch/models/__init__.py b/transformer_engine/pytorch/models/__init__.py new file mode 100644 index 0000000000..bee5474c81 --- /dev/null +++ b/transformer_engine/pytorch/models/__init__.py @@ -0,0 +1,13 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Model-specific transformer layers composed from Transformer Engine modules.""" + +from transformer_engine.pytorch.models.deepseek_v3 import ( + DeepSeekV3Layer, + DeepSeekV3MoE, + MultiLatentAttention, +) + +__all__ = ["DeepSeekV3Layer", "DeepSeekV3MoE", "MultiLatentAttention"] diff --git a/transformer_engine/pytorch/models/deepseek_v3/__init__.py b/transformer_engine/pytorch/models/deepseek_v3/__init__.py new file mode 100644 index 0000000000..a7cbb50ae2 --- /dev/null +++ b/transformer_engine/pytorch/models/deepseek_v3/__init__.py @@ -0,0 +1,13 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""DeepSeekV3 transformer layer built from Transformer Engine MoE building blocks.""" + +from transformer_engine.pytorch.models.deepseek_v3.multi_latent_attention import ( + MultiLatentAttention, +) +from transformer_engine.pytorch.models.deepseek_v3.moe import DeepSeekV3MoE +from transformer_engine.pytorch.models.deepseek_v3.transformer_layer import DeepSeekV3Layer + +__all__ = ["DeepSeekV3Layer", "DeepSeekV3MoE", "MultiLatentAttention"] diff --git a/tests/pytorch/attention/mla_rope_utils.py b/transformer_engine/pytorch/models/deepseek_v3/mla_rope.py similarity index 63% rename from tests/pytorch/attention/mla_rope_utils.py rename to transformer_engine/pytorch/models/deepseek_v3/mla_rope.py index 90eebfc66a..0aea284afd 100644 --- a/tests/pytorch/attention/mla_rope_utils.py +++ b/transformer_engine/pytorch/models/deepseek_v3/mla_rope.py @@ -2,16 +2,18 @@ # # See LICENSE for license information. -"""MLA RoPE for DSv3 671B - Triton forward and backward kernels. +"""Fused MLA RoPE kernels (DeepSeekV3-style decoupled RoPE/NoPE). -Source: Megatron-LM megatron/core/fusions/fused_mla_yarn_rope_apply.py -Falls back to pure PyTorch when Triton is unavailable. +The query kernel rotates the trailing ``head_dim_rope`` slice in place; the KV +kernel builds the key (nope | broadcast-rotated shared rope head) and value +tensors in a single pass. Falls back to pure PyTorch when Triton is unavailable +or for the ``bshd`` layout. -Note: DSv3 uses YaRN-scaled RoPE for long-context extrapolation. This test -intentionally uses plain RoPE (base=10000) because it only validates MXFP8 -attention path wiring, tensor shapes, forward/backward flow, and relative BF16 -vs MXFP8 behavior. Both reference and MXFP8 paths use the same RoPE tables. -""" +The rope slice is read interleaved (checkpoint layout) and written in NeoX +half-split layout.""" + +import math +from typing import Optional, Tuple import torch @@ -23,25 +25,78 @@ except ImportError: HAVE_TRITON = False -HEAD_DIM_ROPE = 64 -HEAD_DIM_NOPE = 128 -HEAD_DIM_V = 128 -ROTARY_BASE = 10000 +__all__ = [ + "build_rope_tables", + "apply_mla_rope_q", + "apply_mla_rope_kv", + "yarn_mscale", + "yarn_concentration_factor", +] + + +def _yarn_correction_dim(num_rotations, dim, base, max_pos): + return (dim * math.log(max_pos / (num_rotations * 2 * math.pi))) / (2 * math.log(base)) + + +def _yarn_correction_range(beta_fast, beta_slow, dim, base, max_pos, round_to_int=True): + low = _yarn_correction_dim(beta_fast, dim, base, max_pos) + high = _yarn_correction_dim(beta_slow, dim, base, max_pos) + if round_to_int: + low, high = math.floor(low), math.ceil(high) + return max(low, 0), min(high, dim - 1) + + +def _yarn_linear_ramp(low, high, dim, device): + if low == high: + high += 0.001 + ramp = (torch.arange(dim, dtype=torch.float32, device=device) - low) / (high - low) + return torch.clamp(ramp, 0, 1) + + +def yarn_mscale(scale: float, mscale: float = 1.0) -> float: + """YaRN attention temperature factor ``0.1 * mscale * ln(scale) + 1`` (1 for scale <= 1).""" + if scale <= 1: + return 1.0 + return 0.1 * mscale * math.log(scale) + 1.0 + + +def yarn_concentration_factor(scaling_factor: float, mscale: float, mscale_all_dim: float) -> float: + """Factor multiplied into cos/sin tables.""" + return yarn_mscale(scaling_factor, mscale) / yarn_mscale(scaling_factor, mscale_all_dim) def build_rope_tables( seq_len: int, - emb_dim: int = HEAD_DIM_ROPE, - base: int = ROTARY_BASE, - device: torch.device = None, -) -> tuple[torch.Tensor, torch.Tensor]: - inv_freq = 1.0 / ( - base ** (torch.arange(0, emb_dim, 2, dtype=torch.float32, device=device) / emb_dim) - ) + emb_dim: int, + base: float = 10000.0, + device: Optional[torch.device] = None, + scaling_factor: Optional[float] = None, + original_max_position_embeddings: int = 4096, + beta_fast: float = 32.0, + beta_slow: float = 1.0, + mscale: float = 1.0, + mscale_all_dim: float = 0.0, +) -> Tuple[torch.Tensor, torch.Tensor]: + """cos/sin tables of shape ``[seq_len, emb_dim]`` (fp32, NeoX duplicated halves). + + With ``scaling_factor`` set, frequencies follow YaRN (NTK-by-parts ramp between + ``beta_fast``/``beta_slow`` rotations over ``original_max_position_embeddings``) and the + tables are scaled by the YaRN concentration factor. + """ + exponent = torch.arange(0, emb_dim, 2, dtype=torch.float32, device=device) / emb_dim + inv_freq = 1.0 / (base**exponent) + factor = 1.0 + if scaling_factor is not None: + low, high = _yarn_correction_range( + beta_fast, beta_slow, emb_dim, base, original_max_position_embeddings + ) + extra_mask = 1.0 - _yarn_linear_ramp(low, high, emb_dim // 2, device) + inv_freq = (inv_freq / scaling_factor) * (1 - extra_mask) + inv_freq * extra_mask + factor = yarn_concentration_factor(scaling_factor, mscale, mscale_all_dim) t = torch.arange(seq_len, device=device, dtype=torch.float32) freqs = torch.outer(t, inv_freq) freqs = torch.cat([freqs, freqs], dim=-1) - return torch.cos(freqs).contiguous(), torch.sin(freqs).contiguous() + return (torch.cos(freqs) * factor).contiguous(), (torch.sin(freqs) * factor).contiguous() if HAVE_TRITON: @@ -69,20 +124,9 @@ def _get_thd_token_idx(cu_seqlens, pid_m, seq_num, cp_rank, cp_size): ) * this_seq_len // 2 return token_idx - @triton.autotune( - configs=[ - triton.Config({"BLOCK_H": 1}), - triton.Config({"BLOCK_H": 2}), - triton.Config({"BLOCK_H": 4}), - triton.Config({"BLOCK_H": 8}), - triton.Config({"BLOCK_H": 16}), - triton.Config({"BLOCK_H": 32}), - triton.Config({"BLOCK_H": 64}), - triton.Config({"BLOCK_H": 128}), - ], - key=["emb_dim", "head_num"], - restore_value=["Q"], - ) + _AUTOTUNE_CONFIGS = [triton.Config({"BLOCK_H": h}) for h in (1, 2, 4, 8, 16, 32, 64, 128)] + + @triton.autotune(configs=_AUTOTUNE_CONFIGS, key=["emb_dim", "head_num"], restore_value=["Q"]) @triton.jit def rotary_fwd_q_kernel( Q, @@ -100,6 +144,7 @@ def rotary_fwd_q_kernel( cp_size, BLOCK_H: tl.constexpr, ): + """In-place RoPE fwd on the trailing rope slice of q.""" pid_m = tl.program_id(axis=0) pid_head = tl.program_id(axis=1) if cu_seqlens_q is None: @@ -129,20 +174,7 @@ def rotary_fwd_q_kernel( tl.store(Q + x_left_off, x_left, mask=mask) tl.store(Q + x_right_off, x_right, mask=mask) - @triton.autotune( - configs=[ - triton.Config({"BLOCK_H": 1}), - triton.Config({"BLOCK_H": 2}), - triton.Config({"BLOCK_H": 4}), - triton.Config({"BLOCK_H": 8}), - triton.Config({"BLOCK_H": 16}), - triton.Config({"BLOCK_H": 32}), - triton.Config({"BLOCK_H": 64}), - triton.Config({"BLOCK_H": 128}), - ], - key=["emb_dim", "head_num"], - restore_value=["DO"], - ) + @triton.autotune(configs=_AUTOTUNE_CONFIGS, key=["emb_dim", "head_num"], restore_value=["DO"]) @triton.jit def rotary_bwd_q_kernel( DO, @@ -160,6 +192,7 @@ def rotary_bwd_q_kernel( cp_size, BLOCK_H: tl.constexpr, ): + """In-place RoPE bwd on the trailing rope slice of dq.""" pid_m = tl.program_id(axis=0) pid_head = tl.program_id(axis=1) if cu_seqlens_q is None: @@ -189,19 +222,7 @@ def rotary_bwd_q_kernel( tl.store(DO + x_1_off, x_1, mask=mask) tl.store(DO + x_2_off, x_2, mask=mask) - @triton.autotune( - configs=[ - triton.Config({"BLOCK_H": 1}), - triton.Config({"BLOCK_H": 2}), - triton.Config({"BLOCK_H": 4}), - triton.Config({"BLOCK_H": 8}), - triton.Config({"BLOCK_H": 16}), - triton.Config({"BLOCK_H": 32}), - triton.Config({"BLOCK_H": 64}), - triton.Config({"BLOCK_H": 128}), - ], - key=["emb_dim", "k_dim", "v_dim", "head_num"], - ) + @triton.autotune(configs=_AUTOTUNE_CONFIGS, key=["emb_dim", "k_dim", "v_dim", "head_num"]) @triton.jit def rotary_fwd_kv_kernel( KV, @@ -228,6 +249,7 @@ def rotary_fwd_kv_kernel( cp_size, BLOCK_H: tl.constexpr, ): + """Fwd: build (key, value) from kv and the shared rotated rope head.""" pid_m = tl.program_id(axis=0) pid_head = tl.program_id(axis=1) if cu_seqlens_kv is None: @@ -268,19 +290,7 @@ def rotary_fwd_kv_kernel( tl.store(K_ptr + x_left_off, x_left, mask=mask) tl.store(K_ptr + x_right_off, x_right, mask=mask) - @triton.autotune( - configs=[ - triton.Config({"BLOCK_H": 1}), - triton.Config({"BLOCK_H": 2}), - triton.Config({"BLOCK_H": 4}), - triton.Config({"BLOCK_H": 8}), - triton.Config({"BLOCK_H": 16}), - triton.Config({"BLOCK_H": 32}), - triton.Config({"BLOCK_H": 64}), - triton.Config({"BLOCK_H": 128}), - ], - key=["emb_dim", "k_dim", "v_dim", "head_num"], - ) + @triton.autotune(configs=_AUTOTUNE_CONFIGS, key=["emb_dim", "k_dim", "v_dim", "head_num"]) @triton.jit def rotary_bwd_kv_kernel( dK, @@ -307,6 +317,7 @@ def rotary_bwd_kv_kernel( cp_size, BLOCK_H: tl.constexpr, ): + """Bwd: scatter (dk, dv) into dkv and reduce rope-slice grads into demb.""" pid_m = tl.program_id(axis=0) pid_head = tl.program_id(axis=1) if cu_seqlens_kv is None: @@ -357,19 +368,23 @@ def rotary_bwd_kv_kernel( tl.store(dEMB_ptr + tl.arange(0, emb_dim // 2) * 2, x_1) tl.store(dEMB_ptr + tl.arange(0, emb_dim // 2) * 2 + 1, x_2) - def _flattened_token_stride(tensor: torch.Tensor) -> int: - if tensor.dim() == 4: - return tensor.stride(1) - return tensor.stride(0) + def _token_stride(tensor: torch.Tensor) -> int: + return tensor.stride(1) if tensor.dim() == 4 else tensor.stride(0) class _MLARoPEQTriton(torch.autograd.Function): + """In-place RoPE on the trailing rope slice of q [s, b, h, nope+rope].""" + @staticmethod def forward(ctx, q, cos, sin, head_dim_nope, head_dim_rope): + """Rotate the rope slice of q in place.""" + if not q.is_contiguous(): + q = q.contiguous() s, b, nheads, _ = q.shape - total = s * b - grid_q = lambda META: (total, triton.cdiv(nheads, META["BLOCK_H"])) - rotary_fwd_q_kernel[grid_q]( + def grid(meta): + return (s * b, triton.cdiv(nheads, meta["BLOCK_H"])) + + rotary_fwd_q_kernel[grid]( q, cos, sin, @@ -379,38 +394,38 @@ def forward(ctx, q, cos, sin, head_dim_nope, head_dim_rope): b, None, None, - _flattened_token_stride(q), + _token_stride(q), q.stride(2), 0, 1, ) - ctx.save_for_backward(cos, sin) - ctx.head_dim_nope = head_dim_nope - ctx.head_dim_rope = head_dim_rope - ctx.nheads = nheads - ctx.s = s - ctx.b = b + ctx.dims = (s, b, nheads, head_dim_nope, head_dim_rope) return q @staticmethod def backward(ctx, dq): + """Counter-rotate the rope slice of dq (in place on the copy).""" cos, sin = ctx.saved_tensors - s, b, nheads = ctx.s, ctx.b, ctx.nheads - total = s * b + # attention backward may hand over a strided grad; the kernel + # assumes a contiguous [s, b, h, d] layout + dq = dq.contiguous() + s, b, nheads, head_dim_nope, head_dim_rope = ctx.dims + + def grid(meta): + return (s * b, triton.cdiv(nheads, meta["BLOCK_H"])) - grid_q = lambda META: (total, triton.cdiv(nheads, META["BLOCK_H"])) - rotary_bwd_q_kernel[grid_q]( + rotary_bwd_q_kernel[grid]( dq, cos, sin, - ctx.head_dim_nope, - ctx.head_dim_rope, + head_dim_nope, + head_dim_rope, nheads, b, None, None, - _flattened_token_stride(dq), + _token_stride(dq), dq.stride(2), 0, 1, @@ -418,15 +433,21 @@ def backward(ctx, dq): return dq, None, None, None, None class _MLARoPEKVTriton(torch.autograd.Function): + """kv [s, b, h, nope+v] + shared rope head [s, b, 1, rope] -> (k, v).""" + @staticmethod def forward(ctx, kv, k_pos_emb, cos, sin, head_dim_nope, head_dim_rope, head_dim_v): + """Build (k, v) from kv and the shared rope head.""" + if not kv.is_contiguous(): + kv = kv.contiguous() s, b, nheads, _ = kv.shape - total = s * b - o_key = kv.new_empty(s, b, nheads, head_dim_nope + head_dim_rope) o_value = kv.new_empty(s, b, nheads, head_dim_v) - grid_kv = lambda META: (total, triton.cdiv(nheads, META["BLOCK_H"])) - rotary_fwd_kv_kernel[grid_kv]( + + def grid(meta): + return (s * b, triton.cdiv(nheads, meta["BLOCK_H"])) + + rotary_fwd_kv_kernel[grid]( kv, k_pos_emb, o_key, @@ -440,37 +461,34 @@ def forward(ctx, kv, k_pos_emb, cos, sin, head_dim_nope, head_dim_rope, head_dim b, None, None, - _flattened_token_stride(kv), + _token_stride(kv), kv.stride(2), - _flattened_token_stride(k_pos_emb), - _flattened_token_stride(o_key), + _token_stride(k_pos_emb), + _token_stride(o_key), o_key.stride(2), - _flattened_token_stride(o_value), + _token_stride(o_value), o_value.stride(2), 0, 1, ) - ctx.save_for_backward(cos, sin) - ctx.head_dim_nope = head_dim_nope - ctx.head_dim_rope = head_dim_rope - ctx.head_dim_v = head_dim_v - ctx.nheads = nheads - ctx.s = s - ctx.b = b + ctx.dims = (s, b, nheads, head_dim_nope, head_dim_rope, head_dim_v) return o_key, o_value @staticmethod def backward(ctx, dk_out, dv_out): + """Gradients for (kv, k_pos_emb) from (dk, dv).""" cos, sin = ctx.saved_tensors - s, b, nheads = ctx.s, ctx.b, ctx.nheads - ndp, ndr, ndv = ctx.head_dim_nope, ctx.head_dim_rope, ctx.head_dim_v - total = s * b - + s, b, nheads, ndp, ndr, ndv = ctx.dims + dk_out = dk_out.contiguous() + dv_out = dv_out.contiguous() d_kv = dk_out.new_empty(s, b, nheads, ndp + ndv) d_emb = dk_out.new_empty(s, b, 1, ndr) - grid_kv = lambda META: (total, triton.cdiv(nheads, META["BLOCK_H"])) - rotary_bwd_kv_kernel[grid_kv]( + + def grid(meta): + return (s * b, triton.cdiv(nheads, meta["BLOCK_H"])) + + rotary_bwd_kv_kernel[grid]( dk_out, dv_out, d_kv, @@ -484,185 +502,66 @@ def backward(ctx, dk_out, dv_out): b, None, None, - _flattened_token_stride(dk_out), + _token_stride(dk_out), dk_out.stride(2), - _flattened_token_stride(dv_out), + _token_stride(dv_out), dv_out.stride(2), - _flattened_token_stride(d_kv), + _token_stride(d_kv), d_kv.stride(2), - _flattened_token_stride(d_emb), + _token_stride(d_emb), 0, 1, ) return d_kv, d_emb, None, None, None, None, None -def _apply_mla_rope_q_with_tables( +def _rotate_interleaved_to_neox(x, cos_table, sin_table, seq_dim): + shape = [1, 1, 1, cos_table.shape[-1]] + shape[seq_dim] = cos_table.shape[0] + cos_ = cos_table.view(shape).to(x.dtype) + sin_ = sin_table.view(shape).to(x.dtype) + half = x.shape[-1] // 2 + x_1 = x[..., 0::2] + x_2 = x[..., 1::2] + x_left = x_1 * cos_[..., :half] - x_2 * sin_[..., :half] + x_right = x_2 * cos_[..., half:] + x_1 * sin_[..., half:] + return torch.cat((x_left, x_right), dim=-1) + + +def apply_mla_rope_q( q: torch.Tensor, cos_table: torch.Tensor, sin_table: torch.Tensor, - head_dim_nope: int = HEAD_DIM_NOPE, - head_dim_rope: int = HEAD_DIM_ROPE, + head_dim_nope: int, + head_dim_rope: int, + tensor_format: str = "sbhd", ) -> torch.Tensor: - if HAVE_TRITON: - return _MLARoPEQTriton.apply( - q, - cos_table, - sin_table, - head_dim_nope, - head_dim_rope, - ) - return _apply_pytorch_q(q, cos_table, sin_table, head_dim_nope, head_dim_rope) + """RoPE on the trailing ``head_dim_rope`` slice of q; in place on the Triton path.""" + if HAVE_TRITON and tensor_format == "sbhd": + return _MLARoPEQTriton.apply(q, cos_table, sin_table, head_dim_nope, head_dim_rope) + seq_dim = 0 if tensor_format == "sbhd" else 1 + q_rope = _rotate_interleaved_to_neox(q[..., head_dim_nope:], cos_table, sin_table, seq_dim) + return torch.cat((q[..., :head_dim_nope], q_rope), dim=-1) -def _apply_mla_rope_kv_with_tables( +def apply_mla_rope_kv( kv: torch.Tensor, k_pos_emb: torch.Tensor, cos_table: torch.Tensor, sin_table: torch.Tensor, - head_dim_nope: int = HEAD_DIM_NOPE, - head_dim_rope: int = HEAD_DIM_ROPE, - head_dim_v: int = HEAD_DIM_V, -) -> tuple[torch.Tensor, torch.Tensor]: - if HAVE_TRITON: + head_dim_nope: int, + head_dim_rope: int, + head_dim_v: int, + tensor_format: str = "sbhd", +) -> Tuple[torch.Tensor, torch.Tensor]: + """Build (k, v) from kv ``[.., h, nope+v]`` and the shared rope head ``[.., 1, rope]``.""" + if HAVE_TRITON and tensor_format == "sbhd": return _MLARoPEKVTriton.apply( - kv, - k_pos_emb, - cos_table, - sin_table, - head_dim_nope, - head_dim_rope, - head_dim_v, + kv, k_pos_emb, cos_table, sin_table, head_dim_nope, head_dim_rope, head_dim_v ) - return _apply_pytorch_kv( - kv, - k_pos_emb, - cos_table, - sin_table, - head_dim_nope, - head_dim_rope, - head_dim_v, - ) - - -def apply_mla_rope_q( - q: torch.Tensor, - head_dim_nope: int = HEAD_DIM_NOPE, - head_dim_rope: int = HEAD_DIM_ROPE, - base: int = ROTARY_BASE, - cos_table: torch.Tensor | None = None, - sin_table: torch.Tensor | None = None, -) -> torch.Tensor: - if cos_table is None or sin_table is None: - s = q.shape[0] - cos_table, sin_table = build_rope_tables( - s, - emb_dim=head_dim_rope, - base=base, - device=q.device, - ) - return _apply_mla_rope_q_with_tables( - q, - cos_table, - sin_table, - head_dim_nope, - head_dim_rope, - ) - - -def apply_mla_rope_kv( - kv: torch.Tensor, - k_pos_emb: torch.Tensor, - head_dim_nope: int = HEAD_DIM_NOPE, - head_dim_rope: int = HEAD_DIM_ROPE, - head_dim_v: int = HEAD_DIM_V, - base: int = ROTARY_BASE, - cos_table: torch.Tensor | None = None, - sin_table: torch.Tensor | None = None, -) -> tuple[torch.Tensor, torch.Tensor]: - if cos_table is None or sin_table is None: - s = kv.shape[0] - cos_table, sin_table = build_rope_tables( - s, - emb_dim=head_dim_rope, - base=base, - device=kv.device, - ) - return _apply_mla_rope_kv_with_tables( - kv, - k_pos_emb, - cos_table, - sin_table, - head_dim_nope, - head_dim_rope, - head_dim_v, - ) - - -def apply_mla_rope( - q: torch.Tensor, - kv: torch.Tensor, - k_pos_emb: torch.Tensor, - head_dim_nope: int = HEAD_DIM_NOPE, - head_dim_rope: int = HEAD_DIM_ROPE, - head_dim_v: int = HEAD_DIM_V, - base: int = ROTARY_BASE, - cos_table: torch.Tensor | None = None, - sin_table: torch.Tensor | None = None, -) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - if cos_table is None or sin_table is None: - s = q.shape[0] - cos_table, sin_table = build_rope_tables( - s, - emb_dim=head_dim_rope, - base=base, - device=q.device, - ) - q = _apply_mla_rope_q_with_tables(q, cos_table, sin_table, head_dim_nope, head_dim_rope) - k, v = _apply_mla_rope_kv_with_tables( - kv, - k_pos_emb, - cos_table, - sin_table, - head_dim_nope, - head_dim_rope, - head_dim_v, - ) - return q, k, v - - -def _rotate_interleaved_to_neox( - x: torch.Tensor, cos_table: torch.Tensor, sin_table: torch.Tensor -) -> torch.Tensor: - cos_ = cos_table[:, None, None, :].to(x.dtype) - sin_ = sin_table[:, None, None, :].to(x.dtype) - half_dim = x.shape[-1] // 2 - x_1 = x[..., 0::2] - x_2 = x[..., 1::2] - x_left = x_1 * cos_[..., :half_dim] - x_2 * sin_[..., :half_dim] - x_right = x_2 * cos_[..., half_dim:] + x_1 * sin_[..., half_dim:] - return torch.cat((x_left, x_right), dim=-1) - - -def _apply_pytorch_q(q, cos_table, sin_table, head_dim_nope, head_dim_rope): - q_nope = q[..., :head_dim_nope] - q_rope = q[..., head_dim_nope : head_dim_nope + head_dim_rope] - q_rope = _rotate_interleaved_to_neox(q_rope, cos_table, sin_table) - return torch.cat((q_nope, q_rope), dim=-1) - - -def _apply_pytorch_kv( - kv, - k_pos_emb, - cos_table, - sin_table, - head_dim_nope, - head_dim_rope, - head_dim_v, -): + seq_dim = 0 if tensor_format == "sbhd" else 1 k_nope = kv[..., :head_dim_nope] v = kv[..., head_dim_nope : head_dim_nope + head_dim_v] - k_rope = _rotate_interleaved_to_neox(k_pos_emb, cos_table, sin_table).expand( - -1, -1, kv.shape[2], -1 - ) - return torch.cat((k_nope, k_rope), dim=-1), v + k_rope = _rotate_interleaved_to_neox(k_pos_emb, cos_table, sin_table, seq_dim) + k_rope = k_rope.expand(*k_nope.shape[:-1], -1) + return torch.cat((k_nope, k_rope), dim=-1), v.contiguous() diff --git a/transformer_engine/pytorch/models/deepseek_v3/moe.py b/transformer_engine/pytorch/models/deepseek_v3/moe.py new file mode 100644 index 0000000000..42f5048e6e --- /dev/null +++ b/transformer_engine/pytorch/models/deepseek_v3/moe.py @@ -0,0 +1,281 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""DeepSeekV3 MoE block: sigmoid router with aux-loss-free bias, shared + +routed experts.""" + +from typing import Optional, Union + +import torch + +import transformer_engine.pytorch.ops as te_ops +from transformer_engine.pytorch.router import fused_topk_with_score_function +from transformer_engine.pytorch.permutation import ( + moe_permute_and_pad_with_probs, + moe_permute_with_probs, + moe_unpermute, +) +from transformer_engine.pytorch.quantization import ( + FP8GlobalStateManager, + get_align_size_for_quantization, +) + +__all__ = ["DeepSeekV3MoE"] + + +_EP_ALIGNMENT = 128 + + +def _make_swiglu_mlp(hidden_size, ffn_hidden_size, dtype, device, num_experts=None): + """Dense SwiGLU MLP, or a grouped one (probs applied inside the activation) per expert. + + The grouped variant fuses into a single CuTe grouped MLP on supported hardware. + """ + common = {"bias": False, "dtype": dtype, "device": device} + if num_experts is None: + return te_ops.Sequential( + te_ops.Linear(hidden_size, 2 * ffn_hidden_size, **common), + te_ops.SwiGLU(), + te_ops.Linear(ffn_hidden_size, hidden_size, **common), + ) + return te_ops.Sequential( + te_ops.GroupedLinear(num_experts, hidden_size, 2 * ffn_hidden_size, **common), + te_ops.ScaledSwiGLU(glu_interleave_size=32), + te_ops.GroupedLinear(num_experts, ffn_hidden_size, hidden_size, **common), + ) + + +class DeepSeekV3MoE(torch.nn.Module): + """ + DeepSeekV3 Mixture-of-Experts block. + + Each token is scored by a sigmoid router with a non-trainable expert bias + updated by ``update_expert_bias()`` (aux-loss-free load balancing) and, + optionally, group-limited routing: experts are split into ``num_groups`` + groups, the top ``group_topk`` groups are selected by their summed scores, + and the final ``topk`` experts are chosen only from those groups. Selected + tokens run through the routed experts, a SwiGLU MLP shared across experts + as a grouped GEMM, with the routing probability applied inside the MLP. An + optional shared expert (dense SwiGLU MLP) is added to every token. On + hardware that supports it the expert MLP runs as a single fused + grouped-GEMM kernel. + + Without ``ep_group`` all experts live on the local device. With + ``ep_group`` the experts are split across the group and tokens are + exchanged over NCCL; this requires ``ep_bootstrap`` to be called once per + process before constructing the module, and bfloat16 inputs. + + Parameters + ---------- + hidden_size : int + size of each input sample. + moe_ffn_hidden_size : int + ffn size of each routed expert. + num_experts : int + total number of routed experts. + topk : int, default = 8 + number of experts per token. + num_groups : int, optional + number of expert groups for node-limited routing. + group_topk : int, optional + number of groups each token is limited to. + routed_scaling_factor : float, default = 2.5 + scaling applied to the routing probabilities. + shared_expert_ffn_hidden_size : int, optional + ffn size of the shared expert; ``None`` + disables the shared expert. + expert_bias_update_rate : float, default = 1e-3 + step size of the aux-loss-free bias update + (see :meth:`update_expert_bias`). + params_dtype : torch.dtype, optional + dtype of module parameters. + ep_group : ProcessGroup, optional + expert-parallel process group; enables the NCCL EP path. + ep_max_tokens_per_rank : int, optional + max local tokens per forward (required with EP). + """ + + def __init__( + self, + hidden_size: int, + moe_ffn_hidden_size: int, + num_experts: int, + topk: int = 8, + num_groups: Optional[int] = None, + group_topk: Optional[int] = None, + routed_scaling_factor: float = 2.5, + shared_expert_ffn_hidden_size: Optional[int] = None, + expert_bias_update_rate: float = 1e-3, + params_dtype: Optional[torch.dtype] = None, + device: Union[torch.device, str] = "cuda", + ep_group: Optional[torch.distributed.ProcessGroup] = None, + ep_max_tokens_per_rank: Optional[int] = None, + ) -> None: + super().__init__() + + dtype = params_dtype if params_dtype is not None else torch.get_default_dtype() + self.hidden_size = hidden_size + self.num_experts = num_experts + self.topk = topk + self.num_groups = num_groups + self.group_topk = group_topk + self.routed_scaling_factor = routed_scaling_factor + self.expert_bias_update_rate = expert_bias_update_rate + + self.gate = torch.nn.Linear( + hidden_size, num_experts, bias=False, dtype=dtype, device=device + ) + self.register_buffer( + "expert_bias", torch.zeros(num_experts, dtype=torch.float32, device=device) + ) + self._last_tokens_per_expert: Optional[torch.Tensor] = None + + self.ep_group = ep_group + self.ep_size = 1 if ep_group is None else torch.distributed.get_world_size(ep_group) + assert num_experts % self.ep_size == 0 + num_local_experts = num_experts // self.ep_size + + self.experts = _make_swiglu_mlp( + hidden_size, moe_ffn_hidden_size, dtype, device, num_experts=num_local_experts + ) + + self.shared_expert = None + if shared_expert_ffn_hidden_size is not None: + self.shared_expert = _make_swiglu_mlp( + hidden_size, shared_expert_ffn_hidden_size, dtype, device + ) + + self.ep_buffer = None + if ep_group is not None: + from transformer_engine.pytorch.ep import EpBuffer + + assert ep_max_tokens_per_rank is not None, "EP requires ep_max_tokens_per_rank." + # Worst case plus per-expert alignment padding, rounded up to + # the multiple of 128 required by the fused grouped MLP. + cap = self.ep_size * ep_max_tokens_per_rank * topk + cap += num_local_experts * _EP_ALIGNMENT + cap = -(-cap // _EP_ALIGNMENT) * _EP_ALIGNMENT + self.ep_buffer = EpBuffer( + top_k=topk, + max_tokens_per_rank=ep_max_tokens_per_rank, + hidden_dim=hidden_size, + num_local_experts=num_local_experts, + recv_capacity_per_rank=cap, + alignment=_EP_ALIGNMENT, + device=device, + ) + + def _route(self, logits: torch.Tensor, topk_indices: Optional[torch.Tensor] = None): + return fused_topk_with_score_function( + logits=logits, + topk=self.topk, + use_pre_softmax=False, + num_groups=self.num_groups, + group_topk=self.group_topk, + scaling_factor=self.routed_scaling_factor, + score_function="sigmoid", + expert_bias=self.expert_bias, + topk_indices=topk_indices, + ) + + def _forward_local(self, tokens: torch.Tensor) -> torch.Tensor: + probs, routing_map = self._route(self.gate(tokens).float()) + tokens_per_expert = routing_map.sum(dim=0) + self._last_tokens_per_expert = tokens_per_expert.detach() + + # Quantized grouped GEMMs need every expert's row count aligned. + align = 1 + if FP8GlobalStateManager.is_fp8_enabled(): + align = get_align_size_for_quantization(FP8GlobalStateManager.get_fp8_recipe()) + if align > 1: + permuted, permuted_probs, row_id_map, pad_offsets, tokens_per_expert = ( + moe_permute_and_pad_with_probs(tokens, probs, routing_map, tokens_per_expert, align) + ) + else: + permuted, permuted_probs, row_id_map = moe_permute_with_probs( + tokens, probs, routing_map, num_out_tokens=tokens.shape[0] * self.topk + ) + pad_offsets = None + + # The fused grouped MLP requires the total row count to be a multiple + # of 128; rows beyond sum(tokens_per_expert) fall outside every group. + num_rows = permuted.shape[0] + pad = (-num_rows) % 128 + if pad: + permuted = torch.nn.functional.pad(permuted, (0, 0, 0, pad)) + permuted_probs = torch.nn.functional.pad(permuted_probs, (0, pad)) + + out = self.experts( + permuted, tokens_per_expert, permuted_probs.to(tokens.dtype), tokens_per_expert + ) + return moe_unpermute( + out[:num_rows], row_id_map, restore_shape=tokens.shape, pad_offsets=pad_offsets + ) + + def _forward_ep(self, tokens: torch.Tensor) -> torch.Tensor: + from transformer_engine.pytorch.ep import ep_dispatch, ep_combine + + assert tokens.dtype == torch.bfloat16, "The EP path requires bfloat16 inputs." + topk_idx = torch.empty( + (tokens.shape[0], self.topk), dtype=torch.int64, device=tokens.device + ) + probs, topk_idx = self._route(self.gate(tokens).float(), topk_indices=topk_idx) + flat_idx = topk_idx.flatten() + self._last_tokens_per_expert = torch.zeros( + self.num_experts, dtype=torch.long, device=tokens.device + ).scatter_add_(0, flat_idx, torch.ones_like(flat_idx)) + topk_weights = probs.gather(1, topk_idx) + + # Zero-filled recv/grad buffers: per-expert alignment padding lands + # inside the grouped-GEMM m_splits, so uninitialized rows would poison + # the expert wgrads. + cap = self.ep_buffer.recv_capacity_per_rank + recv_tokens, recv_weights, tokens_per_expert = ep_dispatch( + self.ep_buffer, + tokens, + topk_idx, + topk_weights, + recv_tokens=torch.zeros( + (cap, self.hidden_size), dtype=tokens.dtype, device=tokens.device + ), + recv_topk_weights=torch.zeros((cap,), dtype=torch.float32, device=tokens.device), + ) + expert_out = self.experts( + recv_tokens, tokens_per_expert, recv_weights.to(tokens.dtype), tokens_per_expert + ) + return ep_combine( + self.ep_buffer, + expert_out, + num_local_tokens=tokens.shape[0], + grad_out=torch.zeros_like(expert_out), + ) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + """ + Parameters + ---------- + hidden_states : torch.Tensor + input of shape ``[..., hidden_size]``. + """ + tokens = hidden_states.reshape(-1, self.hidden_size) + if self.ep_group is not None: + out = self._forward_ep(tokens) + else: + out = self._forward_local(tokens) + if self.shared_expert is not None: + out = out + self.shared_expert(tokens) + return out.view_as(hidden_states) + + @torch.no_grad() + def update_expert_bias(self) -> None: + """Aux-loss-free bias update from the last forward's routing counts. + + With data/expert parallelism, all-reduce ``_last_tokens_per_expert`` + across ranks before calling (or call on identically-routed ranks). + """ + counts = self._last_tokens_per_expert + if counts is None: + return + err = counts.float().mean() - counts.float() + self.expert_bias += self.expert_bias_update_rate * torch.sign(err) diff --git a/transformer_engine/pytorch/models/deepseek_v3/multi_latent_attention.py b/transformer_engine/pytorch/models/deepseek_v3/multi_latent_attention.py new file mode 100644 index 0000000000..56b4d3d0d1 --- /dev/null +++ b/transformer_engine/pytorch/models/deepseek_v3/multi_latent_attention.py @@ -0,0 +1,257 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Multi-Latent Attention (MLA) block as used in DeepSeekV3.""" + +import math +from typing import Optional, Union + +import torch + +from transformer_engine.pytorch.module import Linear, LayerNormLinear +from transformer_engine.pytorch.attention import DotProductAttention +from transformer_engine.pytorch.models.deepseek_v3.mla_rope import ( + apply_mla_rope_kv, + apply_mla_rope_q, + build_rope_tables, + yarn_mscale, +) + +__all__ = ["MultiLatentAttention"] + + +class MultiLatentAttention(torch.nn.Module): + """ + Multi-Latent Attention as used in DeepSeekV3. + + Queries and key-values are projected through low-rank latents + (``q_lora_rank``, ``kv_lora_rank``); RMSNorm on each latent is fused into + the up-projection (:class:`LayerNormLinear` with RMSNorm). Each query/key + head is split into a ``qk_nope_head_dim`` part and a ``qk_rope_head_dim`` + part; RoPE is applied only to the rope part, and the key rope part comes + from a single shared head broadcast to all heads. Attention runs through + :class:`DotProductAttention` with asymmetric head dims + ``kv_channels=(qk_nope_head_dim + qk_rope_head_dim, v_head_dim)``, which + supports the cuDNN fused attention backend. + + RoPE uses the fused MLA kernels from :mod:`.mla_rope` (in-place on the + query rope slice, single-pass key/value assembly); the rope slice follows + the DeepSeekV3 checkpoint convention (interleaved weights, NeoX output). + + Parameters + ---------- + hidden_size : int + size of each input sample. + num_attention_heads : int + number of attention heads. + q_lora_rank : int, default = 1536 + rank of the query latent. + kv_lora_rank : int, default = 512 + rank of the key-value latent. + qk_nope_head_dim : int, default = 128 + per-head dim of the non-rotary query/key part. + qk_rope_head_dim : int, default = 64 + per-head dim of the rotary query/key part. + v_head_dim : int, default = 128 + per-head dim of the values. + attention_dropout : float, default = 0.0 + dropout probability on attention scores. + attn_mask_type : str, default = "causal" + attention mask type passed to :class:`DotProductAttention`. + layernorm_epsilon : float, default = 1e-6 + epsilon of the latent RMSNorms (matches DeepSeekV3). + rotary_base : float, default = 10000.0 + RoPE base. + rope_scaling_factor : float, optional + YaRN context-extension factor; ``None`` disables YaRN. + original_max_position_embeddings : int, default = 4096 + pre-extension context length (YaRN). + beta_fast : float, default = 32.0 + YaRN high-frequency rotation bound. + beta_slow : float, default = 1.0 + YaRN low-frequency rotation bound. + mscale : float, default = 1.0 + YaRN mscale of the rope part. + mscale_all_dim : float, default = 0.0 + YaRN mscale of all dims; sets the default softmax scale to + ``m**2 / sqrt(qk head dim)`` with ``m = 0.1 * mscale_all_dim * ln(factor) + 1``. + softmax_scale : float, optional + softmax scale; defaults to ``1/sqrt(qk head dim)`` (times the YaRN + ``m**2`` when YaRN is enabled). + qkv_format : str, default = "sbhd" + layout of the input/output tensors. + params_dtype : torch.dtype, optional + dtype of module parameters. + tp_group : ProcessGroup, optional + tensor-parallel process group for the up/output projections. + tp_size : int, default = 1 + tensor-parallel world size. + """ + + def __init__( + self, + hidden_size: int, + num_attention_heads: int, + q_lora_rank: int = 1536, + kv_lora_rank: int = 512, + qk_nope_head_dim: int = 128, + qk_rope_head_dim: int = 64, + v_head_dim: int = 128, + attention_dropout: float = 0.0, + attn_mask_type: str = "causal", + layernorm_epsilon: float = 1e-6, + rotary_base: float = 10000.0, + rope_scaling_factor: Optional[float] = None, + original_max_position_embeddings: int = 4096, + beta_fast: float = 32.0, + beta_slow: float = 1.0, + mscale: float = 1.0, + mscale_all_dim: float = 0.0, + softmax_scale: Optional[float] = None, + qkv_format: str = "sbhd", + params_dtype: Optional[torch.dtype] = None, + tp_group: Optional[torch.distributed.ProcessGroup] = None, + tp_size: int = 1, + device: Union[torch.device, str] = "cuda", + ) -> None: + super().__init__() + + assert qkv_format in ("sbhd", "bshd"), "MultiLatentAttention supports sbhd/bshd formats." + assert num_attention_heads % tp_size == 0 + + self.qkv_format = qkv_format + self.num_attention_heads = num_attention_heads + self.num_attention_heads_per_partition = num_attention_heads // tp_size + self.qk_nope_head_dim = qk_nope_head_dim + self.qk_rope_head_dim = qk_rope_head_dim + self.qk_head_dim = qk_nope_head_dim + qk_rope_head_dim + self.v_head_dim = v_head_dim + self.kv_lora_rank = kv_lora_rank + + common = {"bias": False, "params_dtype": params_dtype, "device": device} + tp = {"tp_group": tp_group, "tp_size": tp_size} + + self.q_down_proj = Linear(hidden_size, q_lora_rank, **common) + self.q_up_proj = LayerNormLinear( + q_lora_rank, + num_attention_heads * self.qk_head_dim, + normalization="RMSNorm", + eps=layernorm_epsilon, + parallel_mode="column" if tp_size > 1 else None, + **tp, + **common, + ) + self.kv_down_proj = Linear(hidden_size, kv_lora_rank + qk_rope_head_dim, **common) + self.kv_up_proj = LayerNormLinear( + kv_lora_rank, + num_attention_heads * (qk_nope_head_dim + v_head_dim), + normalization="RMSNorm", + eps=layernorm_epsilon, + parallel_mode="column" if tp_size > 1 else None, + **tp, + **common, + ) + self.out_proj = Linear( + num_attention_heads * v_head_dim, + hidden_size, + parallel_mode="row" if tp_size > 1 else None, + **tp, + **common, + ) + + self.rotary_base = rotary_base + self._yarn_kwargs = { + "scaling_factor": rope_scaling_factor, + "original_max_position_embeddings": original_max_position_embeddings, + "beta_fast": beta_fast, + "beta_slow": beta_slow, + "mscale": mscale, + "mscale_all_dim": mscale_all_dim, + } + self._rope_tables: Optional[tuple] = None + + if softmax_scale is None and rope_scaling_factor is not None: + m = yarn_mscale(rope_scaling_factor, mscale_all_dim) + softmax_scale = m * m / math.sqrt(self.qk_head_dim) + self.softmax_scale = softmax_scale + + self.core_attention = DotProductAttention( + num_attention_heads, + kv_channels=(self.qk_head_dim, v_head_dim), + attention_dropout=attention_dropout, + qkv_format=qkv_format, + attn_mask_type=attn_mask_type, + softmax_scale=softmax_scale, + tp_group=tp_group, + tp_size=tp_size, + ) + + def _rope_tables_for(self, seq_len: int, device: torch.device): + if self._rope_tables is None or self._rope_tables[0].shape[0] < seq_len: + self._rope_tables = build_rope_tables( + seq_len, + self.qk_rope_head_dim, + base=self.rotary_base, + device=device, + **self._yarn_kwargs, + ) + cos, sin = self._rope_tables + return cos[:seq_len], sin[:seq_len] + + def forward( + self, + hidden_states: torch.Tensor, + attention_mask: Optional[torch.Tensor] = None, + attn_mask_type: Optional[str] = None, + checkpoint_core_attention: bool = False, + ) -> torch.Tensor: + """ + Parameters + ---------- + hidden_states : torch.Tensor + input of shape ``[sq, b, h]`` (sbhd) or ``[b, sq, h]`` (bshd). + attention_mask : torch.Tensor, optional + boolean mask passed to :class:`DotProductAttention`. + attn_mask_type : str, optional + override of the constructor's mask type. + checkpoint_core_attention : bool, default = False + checkpoint the core attention computation. + """ + seq_dim = 0 if self.qkv_format == "sbhd" else 1 + seq_len = hidden_states.shape[seq_dim] + heads = self.num_attention_heads_per_partition + + q = self.q_up_proj(self.q_down_proj(hidden_states)) + q = q.view(*q.shape[:-1], heads, self.qk_head_dim) + + kv_down = self.kv_down_proj(hidden_states) + kv_latent, k_pos = torch.split(kv_down, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1) + kv = self.kv_up_proj(kv_latent) + kv = kv.view(*kv.shape[:-1], heads, self.qk_nope_head_dim + self.v_head_dim) + + cos, sin = self._rope_tables_for(seq_len, hidden_states.device) + q = apply_mla_rope_q( + q, cos, sin, self.qk_nope_head_dim, self.qk_rope_head_dim, self.qkv_format + ) + k, v = apply_mla_rope_kv( + kv, + k_pos.unsqueeze(-2), + cos, + sin, + self.qk_nope_head_dim, + self.qk_rope_head_dim, + self.v_head_dim, + self.qkv_format, + ) + + context = self.core_attention( + q, + k, + v, + attention_mask=attention_mask, + qkv_format=self.qkv_format, + attn_mask_type=attn_mask_type, + checkpoint_core_attention=checkpoint_core_attention, + ) + return self.out_proj(context) diff --git a/transformer_engine/pytorch/models/deepseek_v3/transformer_layer.py b/transformer_engine/pytorch/models/deepseek_v3/transformer_layer.py new file mode 100644 index 0000000000..ab11fe6394 --- /dev/null +++ b/transformer_engine/pytorch/models/deepseek_v3/transformer_layer.py @@ -0,0 +1,179 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""DeepSeekV3 transformer layer.""" + +from typing import Optional, Union + +import torch + +from transformer_engine.pytorch.module import LayerNormMLP, RMSNorm +from transformer_engine.pytorch.models.deepseek_v3.multi_latent_attention import ( + MultiLatentAttention, +) +from transformer_engine.pytorch.models.deepseek_v3.moe import DeepSeekV3MoE + +__all__ = ["DeepSeekV3Layer"] + + +class DeepSeekV3Layer(torch.nn.Module): + """ + A full DeepSeekV3 transformer layer, analogous to + :class:`TransformerLayer`: pre-RMSNorm + :class:`MultiLatentAttention`, + then either a dense SwiGLU MLP (:class:`LayerNormMLP` with RMSNorm, used + for the first dense layers of DeepSeekV3) or :class:`DeepSeekV3MoE`, each + with a residual connection. + + Parameters + ---------- + hidden_size : int + size of each input sample. + num_attention_heads : int + number of attention heads. + ffn_hidden_size : int + ffn size of the dense MLP (used when ``num_experts`` is + ``None``). + num_experts : int, optional + number of routed experts; ``None`` makes this a dense layer. + moe_ffn_hidden_size : int, optional + ffn size of each routed expert (required with MoE). + hidden_dropout : float, default = 0.0 + dropout probability on the residual branches. + **kwargs + kwargs common to the submodules (``q_lora_rank``, ``kv_lora_rank``, + ``qk_nope_head_dim``, ``qk_rope_head_dim``, ``v_head_dim``, + ``attention_dropout``, ``attn_mask_type``, ``qkv_format``, ``topk``, + ``num_groups``, ``group_topk``, ``routed_scaling_factor``, + ``shared_expert_ffn_hidden_size``, EP options, ...), forwarded to + :class:`MultiLatentAttention` and :class:`DeepSeekV3MoE`. + """ + + _MLA_KWARGS = frozenset( + { + "q_lora_rank", + "kv_lora_rank", + "qk_nope_head_dim", + "qk_rope_head_dim", + "v_head_dim", + "attention_dropout", + "attn_mask_type", + "rotary_base", + "rope_scaling_factor", + "original_max_position_embeddings", + "beta_fast", + "beta_slow", + "mscale", + "mscale_all_dim", + "softmax_scale", + "qkv_format", + "tp_group", + "tp_size", + } + ) + _MOE_KWARGS = frozenset( + { + "topk", + "num_groups", + "group_topk", + "routed_scaling_factor", + "shared_expert_ffn_hidden_size", + "expert_bias_update_rate", + "ep_group", + "ep_max_tokens_per_rank", + } + ) + + def __init__( + self, + hidden_size: int, + num_attention_heads: int, + ffn_hidden_size: Optional[int] = None, + num_experts: Optional[int] = None, + moe_ffn_hidden_size: Optional[int] = None, + hidden_dropout: float = 0.0, + layernorm_epsilon: float = 1e-5, + params_dtype: Optional[torch.dtype] = None, + device: Union[torch.device, str] = "cuda", + **kwargs, + ) -> None: + super().__init__() + + unknown = set(kwargs) - self._MLA_KWARGS - self._MOE_KWARGS + if unknown: + raise TypeError(f"Unexpected keyword arguments: {sorted(unknown)}") + mla_kwargs = {k: v for k, v in kwargs.items() if k in self._MLA_KWARGS} + moe_kwargs = {k: v for k, v in kwargs.items() if k in self._MOE_KWARGS} + + self.hidden_dropout = hidden_dropout + + self.input_layernorm = RMSNorm( + hidden_size, eps=layernorm_epsilon, device=device, dtype=params_dtype + ) + self.self_attention = MultiLatentAttention( + hidden_size, + num_attention_heads, + params_dtype=params_dtype, + device=device, + **mla_kwargs, + ) + + if num_experts is None: + assert ffn_hidden_size is not None, "Dense layers require ffn_hidden_size." + self.pre_mlp_layernorm = None + self.mlp = LayerNormMLP( + hidden_size, + ffn_hidden_size, + eps=layernorm_epsilon, + normalization="RMSNorm", + activation="swiglu", + bias=False, + params_dtype=params_dtype, + device=device, + ) + else: + assert moe_ffn_hidden_size is not None, "MoE layers require moe_ffn_hidden_size." + self.pre_mlp_layernorm = RMSNorm( + hidden_size, eps=layernorm_epsilon, device=device, dtype=params_dtype + ) + self.mlp = DeepSeekV3MoE( + hidden_size, + moe_ffn_hidden_size, + num_experts, + params_dtype=params_dtype, + device=device, + **moe_kwargs, + ) + + def _residual_add(self, out: torch.Tensor, residual: torch.Tensor) -> torch.Tensor: + out = torch.nn.functional.dropout(out, p=self.hidden_dropout, training=self.training) + return residual + out + + def forward( + self, + hidden_states: torch.Tensor, + attention_mask: Optional[torch.Tensor] = None, + checkpoint_core_attention: bool = False, + ) -> torch.Tensor: + """ + Parameters + ---------- + hidden_states : torch.Tensor + input of shape ``[sq, b, h]`` (sbhd) or ``[b, sq, h]`` (bshd). + attention_mask : torch.Tensor, optional + boolean attention mask. + checkpoint_core_attention : bool, default = False + checkpoint the core attention computation. + """ + attention_out = self.self_attention( + self.input_layernorm(hidden_states), + attention_mask=attention_mask, + checkpoint_core_attention=checkpoint_core_attention, + ) + hidden_states = self._residual_add(attention_out, hidden_states) + + if self.pre_mlp_layernorm is not None: + mlp_out = self.mlp(self.pre_mlp_layernorm(hidden_states)) + else: + mlp_out = self.mlp(hidden_states) + return self._residual_add(mlp_out, hidden_states) diff --git a/transformer_engine/pytorch/module/linear.py b/transformer_engine/pytorch/module/linear.py index 55fc69ef7f..94de69e975 100644 --- a/transformer_engine/pytorch/module/linear.py +++ b/transformer_engine/pytorch/module/linear.py @@ -1901,7 +1901,13 @@ def _linear_eager( class Linear(TransformerEngineBaseModule): """Applies a linear transformation to the incoming data :math:`y = xA^T + b` - On NVIDIA GPUs it is a drop-in replacement for ``torch.nn.Linear``. + On NVIDIA GPUs, this module implements the same linear transformation as + ``torch.nn.Linear``. + + .. note:: + + Its constructor signature differs from ``torch.nn.Linear``. Pass optional + arguments, including ``bias``, by keyword. Parameters ---------- diff --git a/transformer_engine/pytorch/ops/basic/activation.py b/transformer_engine/pytorch/ops/basic/activation.py index 5c33c08b44..26d5261b13 100644 --- a/transformer_engine/pytorch/ops/basic/activation.py +++ b/transformer_engine/pytorch/ops/basic/activation.py @@ -16,7 +16,7 @@ from ...constants import DType from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload from ...tensor.float8_tensor import Float8CurrentScalingQuantizer, Quantizer -from ...utils import clear_tensor_data +from ...utils import _compile_safe_warn, clear_tensor_data from ..op import BasicOperation, OperationContext from .._common import maybe_dequantize @@ -404,14 +404,14 @@ def fuser_forward( next_op_input_quantizer: Optional[Quantizer], # pylint: disable=unused-argument basic_op_kwargs: list[dict[str, Any]], # pylint: disable=unused-argument ) -> tuple[torch.Tensor, Sequence[Sequence[torch.Tensor]]]: + extra_input = basic_op_extra_inputs[0][0] + if self.activation_recompute_in_mlp: - raise RuntimeError( - f"{self.__class__.__name__}(activation_recompute_in_mlp=True) requires the " - "fused grouped MLP path." + _compile_safe_warn( + f"{self.__class__.__name__}(activation_recompute_in_mlp=True) is only supported " + "in the fused grouped MLP path." ) - extra_input = basic_op_extra_inputs[0][0] - if torch.is_autocast_enabled(): dtype = torch.get_autocast_dtype("cuda") elif isinstance(input_, torch.Tensor): @@ -447,18 +447,18 @@ def fuser_backward( ]: del basic_op_grad_extra_outputs - if self.activation_recompute_in_mlp: - raise RuntimeError( - f"{self.__class__.__name__}(activation_recompute_in_mlp=True) requires the " - "fused grouped MLP path." - ) - ctx = basic_op_ctxs[0] x, scales = ctx.saved_tensors x = maybe_dequantize(x.contiguous(), ctx.dtype) scales = maybe_dequantize(scales, ctx.dtype) grad_output = maybe_dequantize(grad_output.contiguous(), ctx.dtype) + if self.activation_recompute_in_mlp: + _compile_safe_warn( + f"{self.__class__.__name__}(activation_recompute_in_mlp=True) is only supported " + "in the fused grouped MLP path." + ) + grad_input, grad_extra_input = self._scaled_unary_backward( grad_output, x, @@ -483,7 +483,9 @@ class ScaledSReLU(_ScaledUnary): ---------- activation_recompute_in_mlp : bool, default = ``False`` Enable fused grouped MLP kernels to recompute activation outputs - during backward when supported instead of saving them. + during backward when supported instead of saving them. Outside the + fused grouped MLP path this option has no effect and a warning is + emitted. """ def _scaled_unary_forward( diff --git a/transformer_engine/pytorch/ops/fused/grouped_mlp.py b/transformer_engine/pytorch/ops/fused/grouped_mlp.py index 61f80b9d9f..58f22c874c 100644 --- a/transformer_engine/pytorch/ops/fused/grouped_mlp.py +++ b/transformer_engine/pytorch/ops/fused/grouped_mlp.py @@ -17,7 +17,7 @@ from packaging.version import Version as PkgVersion import transformer_engine_torch as tex -from ...constants import MXFP8_BLOCK_SCALING_SIZE, NVFP4_BLOCK_SCALING_SIZE, TE_DType +from ...constants import DType, MXFP8_BLOCK_SCALING_SIZE, NVFP4_BLOCK_SCALING_SIZE, TE_DType from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload, start_offload from ...cpp_extensions import general_gemm, general_grouped_gemm_for_grouped_tensor from ...distributed_weight import ( @@ -27,7 +27,7 @@ finalize_weight_grads, ) from ...module.base import _2X_ACC_WGRAD -from ...quantization import Recipe +from ...quantization import Recipe, get_fp8_torch_dtype from ...tensor import NVFP4Quantizer, NVFP4Tensor, NVFP4TensorStorage, Quantizer from ...tensor.grouped_tensor import GroupedTensor from ...tensor.mxfp8_tensor import MXFP8Quantizer, MXFP8Tensor @@ -170,6 +170,17 @@ def _cudnn_frontend_supports_single_group_runtime_offsets( ) and _cudnn_frontend_version_at_least("1.27.0") +def _deterministic_algorithms_required() -> bool: + """Whether bit-exact reproducibility was asked for. Same union as ``DotProductAttention``. + + Uncached: both knobs can change during the process. + """ + return ( + not bool(int(os.getenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "1"))) + or torch.are_deterministic_algorithms_enabled() + ) + + def _wrap_single_quantized_as_grouped( tensor: torch.Tensor, quantized: MXFP8Tensor | NVFP4Tensor | NVFP4TensorStorage, @@ -877,6 +888,12 @@ def fuse_grouped_mlp_ops( # NVFP4 fused grouped MLP uses graph-safe grouped quantize, which currently requires RHT. if recipe.nvfp4() and recipe.disable_rht: return ops + # The fused MXFP8 backward reinterprets the grad output's storage as E4M3, so an E5M2 + # backward format would have its gradients misread rather than converted. This declines + # MXFP8 with Format.HYBRID. fp8_format does not describe NVFP4 gradients, so NVFP4 is + # excluded from the check rather than relying on its value. + if recipe.mxfp8() and get_fp8_torch_dtype(recipe, fprop_tensor=False) != torch.float8_e4m3fn: + return ops if activation_op_types is None: activation_op_types = [ScaledSwiGLU, ScaledClampedQGeGLU] if _cudnn_frontend_supports_grouped_gemm_situglu(): @@ -950,6 +967,11 @@ def grouped_gemm_dactivation_kernel(cls) -> Callable: """Fused kernel for grouped GEMM, activation backward, and scale grad.""" raise NotImplementedError + @classmethod + def grouped_gemm_dactivation_is_deterministic(cls) -> bool: + """Whether this op's dactivation kernel can produce a bit-exact ``dprob``.""" + return False + @classmethod @functools.lru_cache(maxsize=None) def grouped_gemm_quant_kernel(cls) -> Callable: @@ -1398,7 +1420,7 @@ def fuser_forward( alpha_tensor = get_cached_ones_tensor(num_groups, dtype, device) norm_const_tensor = get_cached_ones_tensor(1, torch.float32, device) - current_stream = torch.cuda.current_stream().cuda_stream + current_stream = torch.cuda.current_stream(device.index).cuda_stream fc1_bias_packed = _pack_grouped_linear_bias_for_cudnn(fc1_op) fc2_bias_packed = _pack_grouped_linear_bias_for_cudnn(fc2_op) @@ -2026,6 +2048,13 @@ def fuser_backward( or isinstance(fc1_weight_param, NVFP4Tensor) or isinstance(fc2_weight_param, NVFP4Tensor) ) + if not use_nvfp4 and fc2_grad_output_quantizer.dtype != DType.kFloat8E4M3: + # The pack below reinterprets the grad output's storage as E4M3 rather than + # converting it, so anything else would be read as the wrong format. + raise RuntimeError( + "Fused grouped MLP backward requires an E4M3 grad output, but the recipe " + f"produced {fc2_grad_output_quantizer.dtype}." + ) data_dtype = torch.float4_e2m1fn_x2 if use_nvfp4 else torch.float8_e4m3fn scale_view_dtype = torch.float8_e4m3fn if use_nvfp4 else torch.float8_e8m0fnu sf_vec_size = NVFP4_BLOCK_SCALING_SIZE if use_nvfp4 else MXFP8_BLOCK_SCALING_SIZE @@ -2082,9 +2111,31 @@ def fuser_backward( # Kernel scaling factors alpha_tensor = get_cached_ones_tensor(num_groups, dtype, device) norm_const_tensor = get_cached_ones_tensor(1, torch.float32, device) - current_stream = torch.cuda.current_stream().cuda_stream + current_stream = torch.cuda.current_stream(device.index).cuda_stream unit_activation_scale = bool(getattr(fc1_ctx, "unit_activation_scale", False)) + # A unit activation scale produces no dprob, so there is nothing to make deterministic. + deterministic_dactivation = ( + not unit_activation_scale and _deterministic_algorithms_required() + ) + if deterministic_dactivation: + # Two kernels write dprob and both have to be exact. The cuDNN dactivation + # epilogue produces it below; then, when scale_bias is set, it is passed to + # compute_grouped_dbias_dscales as the ``dscales`` accumulator and atomically + # added into (see triton/grouped_dbias_dscales.py). That Triton kernel is never + # deterministic, so scale_bias rules out a bit-exact dprob on its own. + dprob_is_deterministic = ( + self.grouped_gemm_dactivation_is_deterministic() and not scale_bias + ) + if not dprob_is_deterministic: + raise RuntimeError( + "Deterministic execution was requested" + " (NVTE_ALLOW_NONDETERMINISTIC_ALGO=0 or" + " torch.use_deterministic_algorithms), but the scale gradient (dprob) is" + " accumulated with nondeterministic atomics on this configuration." + " A bit-exact dprob requires the scaled-SReLU activation," + " nvidia-cudnn-frontend 1.28.0 or later, and an FC2 without scale_bias." + ) scales_f32 = None scales_tensor = None dscales_tensor = None @@ -2139,6 +2190,9 @@ def fuser_backward( "use_dynamic_sched": True, } dactivation_kernel = self.grouped_gemm_dactivation_kernel() + if deterministic_dactivation: + # Never passed to a wrapper that would reject it -- the check above raises first. + fc2_dactivation_kwargs["deterministic"] = True if _cudnn_frontend_supports_single_group_runtime_offsets(type(activation_op)): fc2_dactivation_kwargs["use_single_group_runtime_offsets"] = num_groups == 1 if self._cudnn_dact_func is not None: @@ -2669,6 +2723,19 @@ def grouped_gemm_dactivation_kernel(cls) -> Callable: return grouped_gemm_dsrelu_wrapper_sm100 + @classmethod + @functools.lru_cache(maxsize=None) + def grouped_gemm_dactivation_is_deterministic(cls) -> bool: + """Feature-detect the dSReLU wrapper's ``deterministic`` argument (cuDNN FE 1.28.0+).""" + try: + kernel = cls.grouped_gemm_dactivation_kernel() + except ImportError: + return False + try: + return "deterministic" in inspect.signature(kernel).parameters + except (TypeError, ValueError): + return False + def fuse_ops( ops: list[FusibleOperation], diff --git a/transformer_engine/pytorch/ops/fuser.py b/transformer_engine/pytorch/ops/fuser.py index fd66529ba8..2500002700 100644 --- a/transformer_engine/pytorch/ops/fuser.py +++ b/transformer_engine/pytorch/ops/fuser.py @@ -11,7 +11,7 @@ import torch -from ..quantization import FP8GlobalStateManager, Recipe, DelayedScaling +from ..quantization import FP8GlobalStateManager, Recipe from ..quantized_tensor import prepare_for_saving, restore_from_func_ctx from .op import ( BasicOperation, @@ -48,6 +48,8 @@ def _is_graph_capturing() -> bool: OperationFusionFunction: TypeAlias = ( "Callable[tuple[list[FusibleOperation], ...], list[FusibleOperation]]" ) +_FusedOpList: TypeAlias = list[tuple[FusibleOperation, list[int]]] +_FusionParams: TypeAlias = tuple[type, int, Optional[str]] class _OperationFuserAutogradFunction(torch.autograd.Function): @@ -535,15 +537,22 @@ def __init__( op._lock_extra_tensor_channels() # Ops for forward and backward pass, will be populated in maybe_fuse_ops - self._forward_ops: list[tuple[FusibleOperation, list[int]]] - self._backward_ops: list[tuple[FusibleOperation, list[int]]] + self._forward_ops: _FusedOpList + self._backward_ops: _FusedOpList + + # Fused operation configurations are reusable wrappers around the basic + # ops, so cache each configuration by the state that selected it. + self._fused_ops_cache: dict[_FusionParams, tuple[_FusedOpList, _FusedOpList]] = {} # Cache and detect change of state relevant for fusing operations self.recipe_type = None - self.first_op_requiring_backward = 0 self.backward_override = None self._last_amax_history_len = 0 + # Runtime backward boundary. Full activation recompute alternates this + # between the checkpointed forward and the grad-enabled recomputation. + self.first_op_requiring_backward = 0 + # Flatten list of parameters self._basic_op_params = [list(op.parameters()) for op in self._basic_ops] self._basic_op_num_params = list(map(len, self._basic_op_params)) @@ -626,36 +635,48 @@ def maybe_fuse_ops( first_op_requiring_backward = op_idx break - # Early exit if fusion parameters haven't changed - need_reset = False + # Update the runtime backward boundary on every invocation, including + # paths that reuse a cached fused operation configuration. + self.first_op_requiring_backward = first_op_requiring_backward + + # Check if recipe parameters don't match cached values. In this case, + # the recipe state in the basic ops might be invalid, so reset it. recipe_type = type(recipe) + need_to_reset_recipe_state = self.recipe_type != recipe_type + backward_override = recipe.backward_override if recipe is not None else None - fusion_params = (recipe_type, first_op_requiring_backward, backward_override) - if fusion_params != ( - self.recipe_type, - self.first_op_requiring_backward, - self.backward_override, - ): - # Recipe type, backward override, or grad requirements have changed - need_reset = True - elif ( + if backward_override != self.backward_override: + self.backward_override = backward_override + need_to_reset_recipe_state = True + + if ( recipe is not None and recipe.delayed() and self._last_amax_history_len != recipe.amax_history_len ): - # FP8 delayed scaling has changed amax history length - need_reset = True - if not need_reset: - return - - # Reset recipe state - for op in self._basic_ops: - op.reset_recipe_state(recipe=recipe) + self._last_amax_history_len = recipe.amax_history_len + need_to_reset_recipe_state = True - # Check if this is the first iteration - if self.recipe_type is None: + if need_to_reset_recipe_state: for op in self._basic_ops: - op.pre_first_fuser_forward() + op.reset_recipe_state(recipe=recipe) + + # Check if this is the first iteration + if self.recipe_type is None: + for op in self._basic_ops: + op.pre_first_fuser_forward() + + self.recipe_type = recipe_type + + # Training and inference may support different fusions. Keep the + # backward boundary in the key, but pay construction cost only once for + # each configuration. Full recompute therefore builds at most one + # no-grad plan and one grad-enabled plan for a stable recipe. + fusion_params = (recipe_type, first_op_requiring_backward, backward_override) + cached_ops = self._fused_ops_cache.get(fusion_params) + if cached_ops is not None: + self._forward_ops, self._backward_ops = cached_ops + return # Apply joint forward-backward fusions first joint_ops = OperationFuser._apply_fusions( @@ -682,14 +703,9 @@ def maybe_fuse_ops( self._basic_ops, ) - # Save current fusion params - self.recipe_type, self.first_op_requiring_backward, self.backward_override = fusion_params - - # Save amax history length - if isinstance(recipe, DelayedScaling): - self._last_amax_history_len = recipe.amax_history_len - else: - self._last_amax_history_len = 0 + # The FusedOperation contract excludes parameters and per-invocation + # state, so the mapped lists can be selected directly on cache hits. + self._fused_ops_cache[fusion_params] = (self._forward_ops, self._backward_ops) def __call__( self, diff --git a/transformer_engine/pytorch/tensor/hybrid_tensor.py b/transformer_engine/pytorch/tensor/hybrid_tensor.py index 8df2ec8b4b..26d0798b92 100644 --- a/transformer_engine/pytorch/tensor/hybrid_tensor.py +++ b/transformer_engine/pytorch/tensor/hybrid_tensor.py @@ -19,6 +19,10 @@ class HybridQuantizer(Quantizer): """Quantizer that composes rowwise and columnwise representations. + .. warning:: + **EXPERIMENTAL**: ``HybridQuantizer`` is under active development and + its API is subject to change without notice. + When both representations are requested, applies ``rowwise_quantizer`` to produce the rowwise representation and ``columnwise_quantizer`` to produce the columnwise representation. The results are wrapped in a diff --git a/transformer_engine/pytorch/tensor/identity_tensor.py b/transformer_engine/pytorch/tensor/identity_tensor.py index 8310afc653..ec171564fe 100644 --- a/transformer_engine/pytorch/tensor/identity_tensor.py +++ b/transformer_engine/pytorch/tensor/identity_tensor.py @@ -26,6 +26,10 @@ class IdentityQuantizer(Quantizer): """Quantizer that produces a high-precision passthrough representation. + .. warning:: + **EXPERIMENTAL**: ``IdentityQuantizer`` is under active development and + its API is subject to change without notice. + Returns an :class:`IdentityTensorStorage` (or :class:`IdentityTensor`) holding the tensor directly, without a low-precision encoding. ``general_gemm`` materializes it as a plain tensor, so a GEMM consumes it @@ -174,6 +178,21 @@ class IdentityTensor(IdentityTensorStorage, QuantizedTensor): Presents as a standard tensor of its nominal dtype; internally it just holds data directly in that dtype, without a low-precision encoding. + + Parameters + ---------- + shape : iterable of int + Tensor dimensions. + dtype : torch.dtype + Logical tensor datatype. + hp_data : torch.Tensor + Held high-precision data. + quantizer : IdentityQuantizer, optional + Quantizer that produced the tensor. + requires_grad : bool, default = False + Whether to compute gradients for this tensor. + device : torch.device, optional + Device containing the tensor. """ def __repr__(self, *, tensor_contents=None):