From 18d47b7fece818a78c783d64865d1635262a251f Mon Sep 17 00:00:00 2001 From: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> Date: Thu, 10 Sep 2026 06:38:58 -0700 Subject: [PATCH 01/19] [TRTLLM-16304][feat] In-tree implementation of staircase Staircase is one self-contained forward per (checkpoint, GPU architecture, parallel topology) triple, assembled only from a catalog of certified ops and trusted through accuracy gates rather than shared abstractions. Where _torch/models/modeling_deepseekv3.py is one class serving every checkpoint and topology, _torch/staircase/models/deepseek_v3/ is one flat codebase per target. The two live side by side and the contrast is the point. Entry is an environment variable: TRTLLM_STAIRCASE=require "off" -- unset, the default -- is byte-for-byte today's behaviour: the resolver returns immediately and nothing in the package is imported. "auto" uses a target when one matches. "require" raises instead of falling back, because a silent fallback would attribute the built-in implementation's numbers to staircase. An unrecognised value raises rather than reading as "off", for the same reason. A variable rather than an LLM-API field, so that nothing outside this package carries the concept and the only upstream change staircase needs is the resolver hook itself. It has to be exported before the ranks start, not merely before LLM(...): worker ranks receive the environment as it stood when MPI initialized, so a value set later reaches the driver and not them -- and a driver resolving a staircase target while its workers resolve the built-in is exactly the silent split "require" exists to prevent. Routing follows an existing upstream pattern. Targets are keyed by a synthetic architecture name no checkpoint declares, reached through a rewrite in AutoModelForCausalLM._resolve_class -- the same shape as MTPDraftModelForCausalLM. A small table maps architectures[0] to a routing module; that module owns one forward-reading decision tree, so reading a single file tells you where any configuration lands, and explain.py replays the same tree to say why. The package sits beside _torch/models/ rather than inside it so its registrations count as external and win their slot, and so the zoo's non-recursive staleness scan is not broken. The upstream footprint is seventeen lines: the hook in modeling_auto.py, and two package_data patterns in setup.py for the TARGET.md and configs/ that a gate command reads from the installed tree. This batch brings two targets and the 30 catalog entries they call (19 trtllm ops with contract, wrapper and GPU test; 11 thin torch mirrors). Validated on GB300 (sm_103), trtllm 1.3.0rc26: catalog 19/19 entries, 312 cells, plus both 4-rank collective matrices gpt-oss-120b / tp1 boot 10/10; gsm8k 90.6748 vs threshold 85.5989 r1-0528-nvfp4 / dep4 boot 10/10; gsm8k 95.0720 vs threshold 89.9962 r1 configs/mtp3.yaml boot 10/10; paired gsm8k delta -0.3791 against a |delta| < 1.2 criterion; acceptance_length 3.3514 against stock's 3.2752 Every checkpoint digest was verified against the target's own record, so the new gate records and the pre-move sm_100 ones were measured on byte-identical weights. Four differences surfaced while re-certifying the catalog on sm_103, none resolved by widening a tolerance: - Three ops' schemas changed between rc21 and rc26 (renamed and added parameters). The wrappers now mirror their schemas argument for argument. - The MoE FC1 epilogue's MXFP8 block scale is floor(log2(amax))-8 on sm_100 and ceil(log2(amax/448)) on sm_103, bit-exactly. The reference is now architecture-keyed and each arch refutes the other's recipe. - torch 2.12 defaults fp32 matmul to TF32, so cublas_mm's *reference* was the imprecise side; the op is bit-identical to a TF32-disabled product. - The MLA append op now accepts NVFP4 latent pools as well as fp8. Measured statements in the migrated contracts are left exactly as written. They are true records of what was observed on sm_100; rewriting them would manufacture GB300 evidence that does not exist. Receipts are per architecture. Signed-off-by: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> --- setup.py | 8 + tensorrt_llm/_torch/models/modeling_auto.py | 9 + tensorrt_llm/_torch/staircase/README.md | 274 + tensorrt_llm/_torch/staircase/__init__.py | 45 + tensorrt_llm/_torch/staircase/_claim_test.py | 171 + .../_torch/staircase/_router_index.py | 256 + .../_torch/staircase/_routing_test.py | 230 + .../_torch/staircase/catalog/__init__.py | 6 + .../staircase/catalog/activation/__init__.py | 3 + .../activation/flashinfer_silu_and_mul.md | 85 + .../activation/flashinfer_silu_and_mul.py | 12 + .../flashinfer_silu_and_mul_test.py | 56 + .../staircase/catalog/attention/__init__.py | 3 + .../catalog/attention/fused_qk_norm_rope.md | 160 + .../catalog/attention/fused_qk_norm_rope.py | 56 + .../attention/fused_qk_norm_rope_test.py | 252 + .../attention/load_paged_kv_cache_for_mla.md | 317 + .../attention/load_paged_kv_cache_for_mla.py | 50 + .../load_paged_kv_cache_for_mla_test.py | 1003 +++ .../mla_rope_append_paged_kv_assign_q.md | 375 + .../mla_rope_append_paged_kv_assign_q.py | 63 + .../mla_rope_append_paged_kv_assign_q_test.py | 955 +++ .../catalog/attention/mla_rope_generation.md | 427 ++ .../catalog/attention/mla_rope_generation.py | 122 + .../attention/mla_rope_generation_test.py | 1459 ++++ .../catalog/attention/thop_attention.md | 1937 +++++ .../catalog/attention/thop_attention.py | 281 + .../catalog/attention/thop_attention_test.py | 6669 +++++++++++++++++ .../_torch/staircase/catalog/comm/__init__.py | 3 + .../staircase/catalog/comm/_rank_job.py | 88 + .../staircase/catalog/comm/allgather.md | 379 + .../staircase/catalog/comm/allgather.py | 39 + .../staircase/catalog/comm/allgather_test.py | 990 +++ .../_torch/staircase/catalog/comm/conftest.py | 19 + .../staircase/catalog/comm/reducescatter.md | 424 ++ .../staircase/catalog/comm/reducescatter.py | 57 + .../catalog/comm/reducescatter_test.py | 1619 ++++ .../catalog/comm/test_allgather_op_matrix.py | 24 + .../comm/test_reducescatter_op_matrix.py | 24 + .../_torch/staircase/catalog/gemm/__init__.py | 3 + .../_torch/staircase/catalog/gemm/bmm_out.md | 101 + .../_torch/staircase/catalog/gemm/bmm_out.py | 21 + .../staircase/catalog/gemm/bmm_out_test.py | 75 + .../staircase/catalog/gemm/cublas_mm.md | 130 + .../staircase/catalog/gemm/cublas_mm.py | 35 + .../staircase/catalog/gemm/cublas_mm_test.py | 139 + .../staircase/catalog/gemm/nvfp4_gemm.md | 306 + .../staircase/catalog/gemm/nvfp4_gemm.py | 51 + .../staircase/catalog/gemm/nvfp4_gemm_test.py | 650 ++ .../_torch/staircase/catalog/index.yaml | 200 + .../_torch/staircase/catalog/moe/__init__.py | 3 + .../catalog/moe/fp4_block_scale_moe_runner.md | 660 ++ .../catalog/moe/fp4_block_scale_moe_runner.py | 129 + .../moe/fp4_block_scale_moe_runner_test.py | 1926 +++++ .../_torch/staircase/catalog/moe/fused_moe.md | 427 ++ .../_torch/staircase/catalog/moe/fused_moe.py | 133 + .../staircase/catalog/moe/fused_moe_test.py | 990 +++ .../mxe4m3_mxe2m1_block_scale_moe_runner.md | 513 ++ .../mxe4m3_mxe2m1_block_scale_moe_runner.py | 129 + ...e4m3_mxe2m1_block_scale_moe_runner_test.py | 1604 ++++ .../staircase/catalog/moe/noaux_tc_op.md | 244 + .../staircase/catalog/moe/noaux_tc_op.py | 49 + .../staircase/catalog/moe/noaux_tc_op_test.py | 499 ++ .../_torch/staircase/catalog/norm/__init__.py | 3 + .../norm/flashinfer_fused_add_rmsnorm.md | 100 + .../norm/flashinfer_fused_add_rmsnorm.py | 14 + .../norm/flashinfer_fused_add_rmsnorm_test.py | 85 + .../catalog/norm/flashinfer_rmsnorm.md | 74 + .../catalog/norm/flashinfer_rmsnorm.py | 12 + .../catalog/norm/flashinfer_rmsnorm_test.py | 74 + .../catalog/quantization/__init__.py | 3 + .../catalog/quantization/fp4_quantize.md | 341 + .../catalog/quantization/fp4_quantize.py | 37 + .../catalog/quantization/fp4_quantize_test.py | 730 ++ .../catalog/quantization/mxfp8_quantize.md | 170 + .../catalog/quantization/mxfp8_quantize.py | 21 + .../quantization/mxfp8_quantize_test.py | 289 + .../staircase/catalog/torch/__init__.py | 7 + .../_torch/staircase/catalog/torch/add.py | 14 + .../_torch/staircase/catalog/torch/concat.py | 10 + .../_torch/staircase/catalog/torch/copy_.py | 16 + .../staircase/catalog/torch/embedding.py | 11 + .../_torch/staircase/catalog/torch/empty.py | 15 + .../_torch/staircase/catalog/torch/expand.py | 15 + .../_torch/staircase/catalog/torch/pad.py | 16 + .../_torch/staircase/catalog/torch/reshape.py | 14 + .../_torch/staircase/catalog/torch/split.py | 12 + .../staircase/catalog/torch/transpose.py | 15 + .../staircase/catalog/torch/view_dtype.py | 16 + .../docs/models/expert-weight-packing.md | 553 ++ .../docs/models/multi-token-prediction.md | 299 + .../references/trtllm-runtime-integration.md | 1053 +++ tensorrt_llm/_torch/staircase/explain.py | 111 + .../_torch/staircase/models/__init__.py | 8 + .../staircase/models/deepseek_v3/__init__.py | 3 + .../staircase/models/deepseek_v3/routing.py | 76 + .../models/deepseek_v3/targets/__init__.py | 3 + .../targets/r1_0528_nvfp4/__init__.py | 3 + .../targets/r1_0528_nvfp4/sm_103/__init__.py | 3 + .../r1_0528_nvfp4/sm_103/dep4/TARGET.md | 1631 ++++ .../r1_0528_nvfp4/sm_103/dep4/__init__.py | 3 + .../sm_103/dep4/configs/identity.yaml | 21 + .../sm_103/dep4/configs/mtp1.yaml | 69 + .../sm_103/dep4/configs/mtp2.yaml | 69 + .../sm_103/dep4/configs/mtp3.yaml | 79 + .../sm_103/dep4/configs/trtllm-ref-boot.yaml | 30 + .../sm_103/dep4/configs/trtllm-ref-mtp3.yaml | 75 + .../r1_0528_nvfp4/sm_103/dep4/modeling.py | 2229 ++++++ .../r1_0528_nvfp4/sm_103/dep4/weights.py | 448 ++ .../staircase/models/gpt_oss/__init__.py | 3 + .../staircase/models/gpt_oss/routing.py | 59 + .../models/gpt_oss/targets/__init__.py | 3 + .../gpt_oss/targets/gpt_oss_120b/__init__.py | 3 + .../targets/gpt_oss_120b/sm_103/__init__.py | 3 + .../targets/gpt_oss_120b/sm_103/tp1/TARGET.md | 480 ++ .../gpt_oss_120b/sm_103/tp1/__init__.py | 3 + .../gpt_oss_120b/sm_103/tp1/modeling.py | 677 ++ .../gpt_oss_120b/sm_103/tp1/weights.py | 246 + .../_torch/staircase/references/accuracy.yaml | 190 + 119 files changed, 38514 insertions(+) create mode 100644 tensorrt_llm/_torch/staircase/README.md create mode 100644 tensorrt_llm/_torch/staircase/__init__.py create mode 100644 tensorrt_llm/_torch/staircase/_claim_test.py create mode 100644 tensorrt_llm/_torch/staircase/_router_index.py create mode 100644 tensorrt_llm/_torch/staircase/_routing_test.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/__init__.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/activation/__init__.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul.md create mode 100644 tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul_test.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/attention/__init__.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope.md create mode 100644 tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope_test.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla.md create mode 100644 tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla_test.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q.md create mode 100644 tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q_test.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation.md create mode 100644 tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation_test.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/attention/thop_attention.md create mode 100644 tensorrt_llm/_torch/staircase/catalog/attention/thop_attention.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/attention/thop_attention_test.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/comm/__init__.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/comm/_rank_job.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/comm/allgather.md create mode 100644 tensorrt_llm/_torch/staircase/catalog/comm/allgather.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/comm/allgather_test.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/comm/conftest.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/comm/reducescatter.md create mode 100644 tensorrt_llm/_torch/staircase/catalog/comm/reducescatter.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/comm/reducescatter_test.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/comm/test_allgather_op_matrix.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/comm/test_reducescatter_op_matrix.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/gemm/__init__.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out.md create mode 100644 tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out_test.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm.md create mode 100644 tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm_test.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm.md create mode 100644 tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm_test.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/index.yaml create mode 100644 tensorrt_llm/_torch/staircase/catalog/moe/__init__.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner.md create mode 100644 tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner_test.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/moe/fused_moe.md create mode 100644 tensorrt_llm/_torch/staircase/catalog/moe/fused_moe.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/moe/fused_moe_test.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md create mode 100644 tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner_test.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op.md create mode 100644 tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op_test.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/norm/__init__.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm.md create mode 100644 tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm_test.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm.md create mode 100644 tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm_test.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/quantization/__init__.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize.md create mode 100644 tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize_test.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize.md create mode 100644 tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize_test.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/torch/__init__.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/torch/add.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/torch/concat.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/torch/copy_.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/torch/embedding.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/torch/empty.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/torch/expand.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/torch/pad.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/torch/reshape.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/torch/split.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/torch/transpose.py create mode 100644 tensorrt_llm/_torch/staircase/catalog/torch/view_dtype.py create mode 100644 tensorrt_llm/_torch/staircase/docs/models/expert-weight-packing.md create mode 100644 tensorrt_llm/_torch/staircase/docs/models/multi-token-prediction.md create mode 100644 tensorrt_llm/_torch/staircase/docs/references/trtllm-runtime-integration.md create mode 100644 tensorrt_llm/_torch/staircase/explain.py create mode 100644 tensorrt_llm/_torch/staircase/models/__init__.py create mode 100644 tensorrt_llm/_torch/staircase/models/deepseek_v3/__init__.py create mode 100644 tensorrt_llm/_torch/staircase/models/deepseek_v3/routing.py create mode 100644 tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/__init__.py create mode 100644 tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/__init__.py create mode 100644 tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/__init__.py create mode 100644 tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/TARGET.md create mode 100644 tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/__init__.py create mode 100644 tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/identity.yaml create mode 100644 tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp1.yaml create mode 100644 tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp2.yaml create mode 100644 tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp3.yaml create mode 100644 tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/trtllm-ref-boot.yaml create mode 100644 tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/trtllm-ref-mtp3.yaml create mode 100644 tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py create mode 100644 tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/weights.py create mode 100644 tensorrt_llm/_torch/staircase/models/gpt_oss/__init__.py create mode 100644 tensorrt_llm/_torch/staircase/models/gpt_oss/routing.py create mode 100644 tensorrt_llm/_torch/staircase/models/gpt_oss/targets/__init__.py create mode 100644 tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/__init__.py create mode 100644 tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/__init__.py create mode 100644 tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/TARGET.md create mode 100644 tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/__init__.py create mode 100644 tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py create mode 100644 tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/weights.py create mode 100644 tensorrt_llm/_torch/staircase/references/accuracy.yaml diff --git a/setup.py b/setup.py index a4150a691261..2176f4f8166d 100644 --- a/setup.py +++ b/setup.py @@ -209,6 +209,14 @@ def has_ext_modules(self): 'bindings/**/*.pyi', 'evaluate/lm_eval_tasks/**/*', 'usage/schemas/*.json', + # Two staircase patterns, both load-bearing. A target's configs/ are + # `--extra_llm_api_options` files that TARGET.md's verification commands + # pass to trtllm-eval by path, and the claim test asserts TARGET.md sits + # beside the modeling.py it vouches for; both read the installed tree. The + # contracts, catalog index and docs are read by people in a checkout and + # by no code, so they stay out of the wheel. + '_torch/staircase/models/**/*.md', + '_torch/staircase/models/**/configs/*.yaml', ] diff --git a/tensorrt_llm/_torch/models/modeling_auto.py b/tensorrt_llm/_torch/models/modeling_auto.py index babb46378737..cfe12520247f 100644 --- a/tensorrt_llm/_torch/models/modeling_auto.py +++ b/tensorrt_llm/_torch/models/modeling_auto.py @@ -1,6 +1,7 @@ from typing import Generic, Optional, Type from ..model_config import ModelConfig +from ..staircase import staircase_resolve from ..utils import model_extra_attrs from .modeling_utils import (DecoderModelForCausalLM, TConfig, TModel, get_registered_model_class, @@ -33,6 +34,14 @@ def _resolve_class(config: ModelConfig) -> Optional[Type]: "") # Strip the appended EAGLE3 model_arch = "EAGLE3" + model_arch + # Staircase targets are keyed by a synthetic architecture name that no + # checkpoint declares -- the same shape as the Eagle3 rewrite above. + # Returns None unless `staircase` is on and a target claims this exact + # (checkpoint, GPU arch, parallel topology), so the default path is + # byte-for-byte unchanged. + if (staircase_arch := staircase_resolve(config)) is not None: + model_arch = staircase_arch + return get_registered_model_class(model_arch) @staticmethod diff --git a/tensorrt_llm/_torch/staircase/README.md b/tensorrt_llm/_torch/staircase/README.md new file mode 100644 index 000000000000..df4396ca5476 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/README.md @@ -0,0 +1,274 @@ +# Staircase + +One self-contained modeling codebase per deployment target, beside the +built-in model zoo rather than inside it. + +Where `_torch/models/modeling_deepseekv3.py` is one class serving V3, V3-Lite, +R1 and V3.2 across every GPU generation and parallel topology, +`_torch/staircase/models/deepseek_v3/` is one flat forward per (checkpoint, +GPU architecture, parallel topology) triple — assembled only from `catalog/` +entries, sharing nothing with its siblings, and trusted through accuracy gates +instead of shared abstractions. The one-to-one correspondence between +`models//` here and `modeling_.py` there is the point of the exercise. + +## Using it + +```bash +export TRTLLM_STAIRCASE=require # off | auto | require +``` + +| `TRTLLM_STAIRCASE` | Behaviour | +|---|---| +| unset or `off` (default) | The resolver returns immediately. Nothing in this package is imported and behaviour is byte-for-byte what it is today. | +| `auto` | Uses a target when one matches this exact configuration; falls back to the built-in implementation when none does. | +| `require` | Raises instead of falling back, naming the criterion that did not match. | + +**Use `require` for anything you will attribute to staircase.** Under `auto`, +a configuration that misses a target's criteria silently gets the built-in +implementation — and a performance curve measured that way reads as +staircase's. That is the single most expensive mistake available here. + +**Export it before the ranks start**, not merely before `LLM(...)`. Worker +ranks receive the environment as it stood when MPI initialized, and long-lived +ranks under `trtllm-llmapi-launch` receive it once at launch, so a value set +later reaches the driver and not them -- and a driver resolving a target while +its workers resolve the built-in is the silent split this package exists to +prevent. + +An environment variable rather than an LLM-API field is a deliberate trade: it +keeps the entire concept inside this package, so the only change staircase +needs anywhere else is the `_resolve_class` hook. The cost is that the switch +does not appear in a run's recorded `llm.args` and cannot be set through +`--extra_llm_api_options`. + +The predecessor `STAIRCASE_TARGET` is gone. It named a *target*; this names +only a mode, and routing picks the target from the configuration. + +The checkpoint is read exactly as published, with no target-owned +`config.json`. + +## How a config finds its target + +``` +LLM(model=...) -> ModelLoader -> AutoModelForCausalLM._resolve_class + | + +-- staircase_resolve(config) + | + _router_index.py +-- architectures[0] -> routing module + models//routing.py +-- one forward-reading decision tree + +-- returns a synthetic class name + | + get_registered_model_class(name) +``` + +The synthetic name (`StaircaseGptOss120bSm103Tp1`) is a registry key that no +checkpoint declares. Upstream already does exactly this for +`MTPDraftModelForCausalLM`, which also exists only as a `_resolve_class` +rewrite. + +To ask why a configuration landed where it did: + +``` +python -m tensorrt_llm._torch.staircase.explain \ + --model /path/to/DeepSeek-R1-0528-NVFP4 --tp 4 --ep 4 --attention-dp +``` + +### What may decide a target + +A quantity may be a routing criterion only if it is known when +`_resolve_class` runs **and** constant for the engine's life — and only if it +changes the *structure of the forward*. Config shape, SM version, parallel +topology qualify. `max_num_tokens` and `cuda_graph_config.max_batch_size` pass +the first test and fail the second: they are tuning knobs and belong in a +target's `configs/` variant. Batch composition and `num_contexts` fail the +first outright — those move every step, and a target chosen from them would +be chosen once and then be wrong. Per-step specialization is a separate +mechanism (dispatch inside a target's `forward`), not a routing dimension. + +### Checkpoint identity is a shape fingerprint + +Routing recognizes a checkpoint by `(num_hidden_layers, hidden_size, ...)`, +the same sniffing idiom as `is_mla` / `is_nemotron_hybrid` upstream. It cannot +tell a fine-tune from the original. That is a deliberate trade, and it changes +what a gate record means: not "this target passed" but "this modeling code +passed **on the checkpoint whose digests TARGET.md records**". Running it on +any other checkpoint of the same shape is ungated. Each `TARGET.md` says so. + +## Layout + +``` +_router_index.py architectures[0] -> routing module. Small, stable, test-guarded. +explain.py why a configuration routed where it did +models// + routing.py one forward-reading decision tree per architecture family + targets//// + modeling.py weights.py TARGET.md [configs/] +catalog/ the kernel vocabulary: contract .md + wrapper .py +references/ accuracy anchors, keyed by HF repo name +docs/ cross-target mechanism and runtime notes +``` + +The third piece of a catalog entry, its GPU test, lives in the tests tree: + +``` +tests/unittest/_torch/staircase/ + test_staircase_claims.py routing tables vs the targets they name (no GPU) + test_staircase_routing.py what staircase_resolve does (no GPU) + /test_staircase_.py +``` + +That split is not a preference; it is where this repo's CI collects from, and +an in-package test is on no list. It costs one thing worth stating: a receipt +is valid only if it post-dates the last write to *every* file of its entry, so +that check now has to look in both trees. + +The two collective entries keep their rank bodies in `catalog/comm/` +(`allgather_test.py`, `reducescatter_test.py`) because the launcher re-execs +them as `python -m` and the ranks need the package context; only the collected +shells moved. + +Identity is the path. `targets/` keeps all three segments rather than +flattening them, and the class name carries the same triple; +`test_staircase_claims.py` asserts they agree. + +### Why beside `_torch/models/`, not inside it + +Two upstream mechanisms, and together they are the worst combination — +inheriting the built-in constraints without inheriting the built-in guards: + +* `is_builtin_zoo_module` matches on the zoo's package prefix. Inside it, + staircase registrations would count as built-in and only fill *empty* + registry slots. Outside it they are external and always win their slot. +* `test_lazy_model_zoo.py`'s scan of the zoo directory is **non-recursive**, + so decorators in a subpackage are invisible to it — putting synthetic names + in the built-in static index would fail its staleness assertion. + +Conceptually `models/` means "one architecture, one class, shared across +checkpoints", which is the opposite of what this package is for. + +## Gates + +1. **Boot** — minutes, binary. `examples/llm-api/quickstart_advanced.py` with + the target's topology flags and `TRTLLM_STAIRCASE=require` exported. Engine + cold start, weight-manifest coverage, a handful of greedy continuations. + Catches catastrophes, not accuracy. A target whose `configs/` holds a + variant that changes the forward has its own file there, so that path gets + a minutes-scale gate too. +2. **Accuracy** — the release criterion. `trtllm-eval` with the protocol from + `references/accuracy.yaml` and `TRTLLM_STAIRCASE=require` exported. One-sided: measured >= reference − tol. + This is the gate CI runs, as `accuracy/test_staircase.py`. +3. **Acceptance** — required whenever a variant is distribution-preserving by + construction, speculative decoding above all. Rejection sampling holds the + emitted distribution to the target model's, so a *miscomputed* draft path + produces correct text more slowly: boot passes, accuracy passes, only + speed moves. The detector is `acceptance_length` against the same + checkpoint under stock in-tree modeling at the same workload and the same + `speculative_config` — i.e. `TRTLLM_STAIRCASE=off` versus `=require`, + which is now one variable rather than two harnesses. + +Compare a variant against an identity run **in the same session**: one target +measured 94.7688 and 95.0720 on a bit-identical forward a day apart, so a +cross-session delta carries session variance into the judgement. + +Perf is measured, never gated. + +## Status of every record in this tree + +**The catalog is fully certified on sm_103. The targets construct but have +never executed.** + +Two things voided every receipt in the move: each catalog test file was +rewritten, and the targets moved from sm_100 (B200) to sm_103 (GB300), where +certification is per architecture. The whole catalog was therefore re-run on +GB300 under 1.3.0rc26 -- **19 of 19 entries pass, 312 certified cells**, plus +both 4-rank collective matrices. + +Getting there surfaced four real differences. None was resolved by widening a +tolerance, and each is written up in its own contract: + +* **Op schema drift rc21 -> rc26** (`thop_attention`, `mla_rope_generation`, + `mla_rope_append_paged_kv_assign_q`). Parameters were renamed and added. The + wrappers now mirror their schemas argument for argument, so the next drift + fails loudly rather than shifting a positional list silently. +* **The MoE FC1 epilogue changed block-scale recipe**, bit-exactly: + `floor(log2(amax))-8` on sm_100, `ceil(log2(amax/448))` on sm_103. The + reference is architecture-keyed and each arch refutes the other's recipe. +* **torch 2.12 made fp32 matmul default to TF32**, so `cublas_mm`'s *reference* + was the imprecise side; the op is bit-identical to a TF32-disabled product. +* **The MLA append op now accepts NVFP4 latent pools** as well as fp8 (accepted + by the op, not certified here). + +**gpt-oss-120b / sm_103 / tp1 is gated on GB300.** Its checkpoint was +downloaded and every digest re-verified against `TARGET.md`, so the new records +and the old sm_100 ones were measured on byte-identical weights. Both gates +pass with `TRTLLM_STAIRCASE=require` in force, which is what rules out the built-in +implementation having been measured instead: + +| Gate | Result | +|---|---| +| boot, 10 greedy continuations | **10/10** keyword asserts | +| gsm8k, full 1319 | **90.6748** (`exact_match,flexible-extract`) vs threshold 85.5989 -- pass by 5.08 | + +**deepseek-r1-0528-nvfp4 / sm_103 / dep4 is gated too**, identity and the +`mtp3` variant, on 4 GB300s through `trtllm-llmapi-launch`: + +| Gate | Result | +|---|---| +| boot, 10 greedy continuations | **10/10** | +| gsm8k, full 1319 | **95.0720** vs threshold 89.9962 -- pass by 5.08 | +| boot, `configs/mtp3.yaml` | **10/10**, with the layer-61 MTP module loaded | +| gsm8k, identity vs mtp3 **paired in one session** | delta **-0.3791** against a `\|delta\| < 1.2` criterion -- 0.63 sigma | +| acceptance vs stock | `acceptance_length` **3.3514** vs **3.2752**, ratio **1.023** | + +That last row is the one that matters for a speculative variant: rejection +sampling makes a miscomputed draft layer *slower*, not wrong, so boot and +accuracy are blind to it and only the acceptance rate against a reference can +see it. `configs/mtp{1,2}.yaml` remain ungated -- they are the dominated end of +the measured draft-length axis. + +Measured statements throughout the contracts are left exactly as written. +They are true records of what was observed on sm_100, and rewriting them would +manufacture GB300 evidence that does not exist. + +### What replaced the version pin + +Out of tree each target hard-asserted `tensorrt_llm == 1.3.0rc21` at import, +because a target reads private engine surface and a drifting engine silently +voids its records. In tree that assert is meaningless — the target moves with +the trunk — so it is gone, replaced by an SM assert, which is the part of the +identity that does *not* move. + +The pin was earning its keep, though, and the migration paid the bill +immediately: between rc21 and rc26 the attention backends moved from +`_torch/attention_backend/{interface,trtllm}.py` to +`_torch/attention/backends/`. The compatibility shim left behind re-exports +the names but not the submodule paths, so all five files importing them +failed at import. They now use the canonical path. + +The lesson generalizes: with no pin, a target's contact with private engine +surface is checked only by running it. The import-time +`_check_static_contract` op-existence loop and the first-forward metadata +field check are what turn that from a wrong answer into a loud failure, which +is why both survived the move. + +## Two facts this migration surfaced about upstream + +* **`pip` is a runtime dependency of TensorRT-LLM itself.** `libtensorrt_llm.so` + locates its bundled kernel headers on the NVRTC JIT path by running + `pip show tensorrt_llm`. Staircase was merely the first thing to write that + down (a `uv`-managed venv does not install `pip` by default), and it is not + a staircase dependency. +* **The catalog contracts are now documentation of upstream ops.** For + example `_torch/modules/linear.py` calls `cublas_mm(input, module.weight.t(), ...)` + correctly, but nothing there says why the `.t()` is mandatory or that a + non-contiguous weight computes a silently wrong answer — `catalog/gemm/cublas_mm.md` + does. The reverse duty comes with it: if one of these ops changes and its + contract does not follow, the misleading is no longer confined to staircase. + +## Not migrated in this batch + +The 13 catalog entries no migrated target calls (listed in `catalog/index.yaml`), +the five other targets, and the agent definitions. A later target that needs +one of those entries must bring its contract **and** its receipt, not just the +wrapper — otherwise that target consumes an op with no certification record at +all. diff --git a/tensorrt_llm/_torch/staircase/__init__.py b/tensorrt_llm/_torch/staircase/__init__.py new file mode 100644 index 000000000000..155310cbcb67 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/__init__.py @@ -0,0 +1,45 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Staircase: one self-contained modeling codebase per deployment target. + +Where the built-in zoo has one class per architecture serving every +checkpoint, parallel topology and GPU generation, staircase has one flat, +self-contained forward per (checkpoint, GPU arch, parallel) triple, assembled +only from ``catalog/`` entries and trusted through accuracy gates rather than +shared abstractions. The two live side by side: ``models//`` here +corresponds one-to-one with the zoo's ``modeling_.py``, and the contrast is +the point. + +Entry is a single environment variable:: + + TRTLLM_STAIRCASE = require + +``off`` (unset, the default) is byte-for-byte today's behaviour: the resolver returns +immediately and nothing in this package is imported. ``auto`` uses a target +when one matches and falls back to the built-in implementation when none +does. ``require`` raises instead of falling back -- see ``_router_index`` for +why that mode is not optional. + +This package sits *beside* ``_torch/models/`` rather than inside it, which is +load-bearing twice over: ``is_builtin_zoo_module`` matches on the zoo's +package prefix, so registrations from here count as external and always win +their architecture slot; and ``test_lazy_model_zoo``'s non-recursive scan of +the zoo directory would not see these modules, so putting the synthetic names +in the built-in static index would fail its staleness assertion. +""" + +from ._router_index import ( + STAIRCASE_ENV, + STAIRCASE_ROUTERS, + StaircaseContext, + StaircaseMode, + staircase_resolve, +) + +__all__ = [ + "STAIRCASE_ENV", + "STAIRCASE_ROUTERS", + "StaircaseContext", + "StaircaseMode", + "staircase_resolve", +] diff --git a/tensorrt_llm/_torch/staircase/_claim_test.py b/tensorrt_llm/_torch/staircase/_claim_test.py new file mode 100644 index 000000000000..fba88c277ce3 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/_claim_test.py @@ -0,0 +1,171 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""The routing tables and the targets they name must not drift apart. + +Everything here is read off disk as text -- no target module is imported, so +this runs anywhere, including a CI machine with no GPU and no built +extensions. That is deliberate: the failure mode being guarded against is a +rename or a move, and those are visible without executing anything. + +What is *not* guarded here is whether a target is correct; that is what its +gate records are for. +""" + +from __future__ import annotations + +import re +from pathlib import Path + +import pytest + +from ._router_index import STAIRCASE_ROUTERS, routing_module + +_ROOT = Path(__file__).resolve().parent +_ARCHS = sorted(STAIRCASE_ROUTERS) + + +def _module_path(dotted: str) -> Path: + """Resolve a package-relative dotted module name to its file.""" + return _ROOT / (dotted.replace(".", "/") + ".py") + + +def test_every_routed_architecture_has_an_importable_routing_module(): + for arch in _ARCHS: + assert routing_module(arch) is not None, arch + + +def test_no_routing_module_is_orphaned(): + """A routing.py the index does not name would never be consulted.""" + indexed = {_module_path(m) for m in STAIRCASE_ROUTERS.values()} + on_disk = set((_ROOT / "models").glob("*/routing.py")) + assert on_disk == indexed, ( + f"routing modules on disk but not in STAIRCASE_ROUTERS: " + f"{sorted(p.relative_to(_ROOT) for p in on_disk - indexed)}" + ) + + +@pytest.mark.parametrize("arch", _ARCHS) +def test_targets_table_and_target_modules_agree(arch): + """Every name the tree can return must have a module that registers it.""" + routing = routing_module(arch) + produced = set(routing._TARGETS.values()) + declared = set(routing.TARGET_MODULES) + assert produced == declared, ( + f"{arch}: _TARGETS yields {sorted(produced - declared)} with no " + f"TARGET_MODULES entry; TARGET_MODULES declares " + f"{sorted(declared - produced)} the tree cannot return" + ) + + +@pytest.mark.parametrize("arch", _ARCHS) +def test_each_target_module_registers_its_own_name(arch): + """The synthetic name is only a registry key -- the module must fill it. + + Nothing else can: the built-in static index does not carry staircase + names, so a target whose decorator says something different resolves to + None and the engine reports an unknown architecture. + """ + routing = routing_module(arch) + for name, dotted in routing.TARGET_MODULES.items(): + path = _module_path(dotted) + assert path.is_file(), f"{arch}: {name} -> missing {path}" + source = path.read_text() + assert f'@register_auto_model("{name}")' in source, ( + f"{arch}: {path.relative_to(_ROOT)} does not register {name!r}" + ) + + +@pytest.mark.parametrize("arch", _ARCHS) +def test_target_identity_matches_its_path(arch): + """Identity is the path: //. + + The class name encodes the same triple, and the routing module's ``_SM`` + has to be the arch segment those targets actually live under -- a target + moved to a new SM directory without its routing constant following is the + one drift that would still route, and route wrong. + """ + routing = routing_module(arch) + major, minor = routing._SM + expected_segment = f"sm_{major}{minor}" + + for name, dotted in routing.TARGET_MODULES.items(): + parts = dotted.split(".") + assert parts[-1] == "modeling", dotted + parallel, sm_segment, checkpoint = parts[-2], parts[-3], parts[-4] + + assert sm_segment == expected_segment, ( + f"{arch}: {name} lives under {sm_segment} but its routing module " + f"only ever matches {expected_segment}" + ) + + def camel(segment: str) -> str: + return "".join(w.capitalize() for w in segment.split("_")) + + for segment in (checkpoint, sm_segment, parallel): + assert camel(segment).lower() in name.lower(), ( + f"{arch}: {name} does not carry path segment {segment!r}" + ) + + assert (checkpoint, parallel) in routing._TARGETS, ( + f"{arch}: no _TARGETS entry keyed ({checkpoint!r}, {parallel!r})" + ) + assert routing._TARGETS[(checkpoint, parallel)] == name + + +@pytest.mark.parametrize("arch", _ARCHS) +def test_checkpoint_fingerprints_are_distinct(arch): + """Two checkpoints sharing a fingerprint would route to one target.""" + routing = routing_module(arch) + names = list(routing._CHECKPOINTS.values()) + assert len(names) == len(set(names)), f"{arch}: duplicate checkpoint names in _CHECKPOINTS" + assert set(names) == {c for c, _ in routing._TARGETS}, ( + f"{arch}: _CHECKPOINTS and _TARGETS name different checkpoints" + ) + + +@pytest.mark.parametrize("arch", _ARCHS) +def test_no_routing_module_reads_an_unplumbed_dimension(arch): + """``ctx.is_disagg`` is declared but nothing sets it, so it is always + False. A branch on it would take the wrong side in a disaggregated + deployment and say nothing -- exactly the silent-wrong-system failure + ``require`` exists to prevent. Delete this test when it is plumbed. + """ + source = _module_path(STAIRCASE_ROUTERS[arch]).read_text() + assert "is_disagg" not in source, ( + f"{arch}: routing reads ctx.is_disagg, which no caller sets yet; " + f"plumb it onto ModelConfig first" + ) + + +def test_every_target_ships_the_four_products(): + """modeling.py, weights.py, smoke.py and TARGET.md travel together.""" + for arch in _ARCHS: + routing = routing_module(arch) + for name, dotted in routing.TARGET_MODULES.items(): + target_dir = _module_path(dotted).parent + for product in ("modeling.py", "weights.py", "smoke.py", "TARGET.md"): + assert (target_dir / product).is_file(), f"{name}: missing {product}" + + +def test_targets_do_not_share_files(): + """Isolation is the property being demonstrated; measure it. + + Targets are allowed to import the catalog and nothing else of each + other's. A shared helper between two targets is the first step back to + the abstraction this package exists to avoid. + """ + for arch in _ARCHS: + routing = routing_module(arch) + for name, dotted in routing.TARGET_MODULES.items(): + target_dir = _module_path(dotted).parent + for source_file in target_dir.glob("*.py"): + text = source_file.read_text() + for match in re.finditer(r"^from (\.+)([\w.]*) import", text, re.MULTILINE): + dots, tail = match.group(1), match.group(2) + if len(dots) == 1: + continue # sibling within the target + assert tail.startswith("catalog"), ( + f"{name}: {source_file.name} reaches outside its own " + f"directory for {tail!r}; targets may import the " + f"catalog and nothing else" + ) diff --git a/tensorrt_llm/_torch/staircase/_router_index.py b/tensorrt_llm/_torch/staircase/_router_index.py new file mode 100644 index 000000000000..3c62838bb6d0 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/_router_index.py @@ -0,0 +1,256 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Architecture name -> routing module, and the resolver that reads the table. + +Staircase targets are keyed by a *synthetic* architecture name that no +checkpoint declares. The checkpoint's own ``architectures[0]`` stays what it +always was (``GptOssForCausalLM``, ``DeepseekV3ForCausalLM``), so it cannot +also be the target selector; ``staircase_resolve`` rewrites it into the +synthetic name, and ``AutoModelForCausalLM._resolve_class`` then looks that up +in the ordinary registry. + +The pattern is upstream's own: ``MTPDraftModelForCausalLM`` is likewise a name +no config.json declares, reached only through a rewrite in ``_resolve_class``. + +Two levels, on purpose: + +* This table maps ``architectures[0]`` -> routing module. Its keys are things + that already exist in every checkpoint's config.json, so adding a new + checkpoint, GPU arch or parallel topology never touches it -- only a new + architecture family does. +* The routing module owns the rest of the decision as one forward-reading + tree. Reading that single file tells you where any configuration lands, and + ``explain.py`` replays the same tree to say *why*. + +Routing modules are imported lazily: resolving a GptOss config never imports +the DeepSeek tree, and a target's modeling code is imported only once its +routing module has claimed the config. +""" + +from __future__ import annotations + +import enum +import os +from dataclasses import dataclass +from importlib import import_module +from typing import TYPE_CHECKING, Any, List, Optional, Tuple + +if TYPE_CHECKING: + from tensorrt_llm._torch.model_config import ModelConfig + +_PACKAGE = "tensorrt_llm._torch.staircase" + +#: The switch. An environment variable rather than an LLM-API field, so that +#: nothing outside this package has to carry the concept: the only upstream +#: change staircase needs is the ``_resolve_class`` hook itself. +#: +#: It has to be exported **before the ranks start**, not merely before +#: ``LLM(...)``. Worker ranks receive the environment as it stood when MPI +#: initialized, and long-lived ranks under ``trtllm-llmapi-launch`` receive it +#: once at launch, so a value set later reaches the driver and not them -- and +#: a driver that resolves a staircase target while its workers resolve the +#: built-in is exactly the silent split this package exists to prevent. Export +#: it in the shell, or before ``import tensorrt_llm``. +STAIRCASE_ENV = "TRTLLM_STAIRCASE" + +# architectures[0] -> routing module, relative to this package. +STAIRCASE_ROUTERS = { + "GptOssForCausalLM": "models.gpt_oss.routing", + "DeepseekV3ForCausalLM": "models.deepseek_v3.routing", +} + + +class StaircaseMode(str, enum.Enum): + """What to do when a config reaches the staircase resolver.""" + + OFF = "off" + AUTO = "auto" + REQUIRE = "require" + + @classmethod + def from_env(cls) -> "StaircaseMode": + """Read ``TRTLLM_STAIRCASE``; unset means off. + + An unknown value raises rather than falling back. Silently reading a + typo as "off" would hand back the built-in implementation while the + caller believed they had asked for a staircase target, and a number + measured that way is attributed to the wrong system -- the failure the + ``require`` mode below exists to prevent, arriving through the door + instead of the window. + """ + raw = os.environ.get(STAIRCASE_ENV) + if raw is None or raw == "": + return cls.OFF + try: + return cls(raw.strip().lower()) + except ValueError: + raise ValueError( + f"{STAIRCASE_ENV}={raw!r} is not a staircase mode; expected " + f"one of {', '.join(m.value for m in cls)}" + ) from None + + +@dataclass(frozen=True) +class StaircaseContext: + """Everything a routing decision is allowed to depend on. + + The admission rule is one line: a quantity may live here only if it is + already known when ``_resolve_class`` runs *and* does not change for the + rest of the engine's life. Config shape, SM version, the parallel mapping, + the quant and speculative configs and the disagg flag all qualify. Batch + composition, ``num_contexts``, "is this step pure decode" do not -- they + move every forward, and a target selected from them would be selected + once and then be wrong. + + That is a necessary condition, not a sufficient one. ``max_num_tokens`` + and ``cuda_graph_config.max_batch_size`` are per-instance constants too, + but they are tuning knobs: they belong in a ``configs/`` variant. Only an + instance constant that changes the *structure of the forward* earns a + target of its own. + + ``is_disagg`` is declared but **not yet plumbed**: nothing sets it on + ``ModelConfig``, so it reads False in every deployment, disaggregated or + not. It is here so the field exists the day it is, and a claim test + forbids any routing module from reading it until then -- a branch on a + value that is always False would take the wrong side silently, which is + the failure this package's ``require`` mode exists to prevent. The useful + disagg dimension is the instance's ``ServerRole``, and that stops at the + server layer today; wiring it into ``LlmArgs`` is a prerequisite, not + part of this work. + """ + + pretrained_config: Any + mapping: Any + sm: Tuple[int, int] + quant_config: Any + spec_config: Any + is_disagg: bool + + @classmethod + def from_model_config(cls, config: "ModelConfig") -> "StaircaseContext": + import torch + + assert torch.cuda.is_available(), ( + "staircase routes on the SM version of the device it will run on; " + "no CUDA device is visible" + ) + return cls( + pretrained_config=config.pretrained_config, + mapping=config.mapping, + sm=torch.cuda.get_device_capability(), + quant_config=config.quant_config, + spec_config=config.spec_config, + is_disagg=getattr(config, "is_disagg", False), + ) + + +class Trace: + """Records the criteria a routing tree evaluated, for ``explain``. + + Routing modules call ``check``/``resolve`` instead of a bare ``if`` so the + same tree that decides can also narrate. ``NULL_TRACE`` makes both a no-op + and is the default, so the resolve path pays nothing. + """ + + __slots__ = ("steps",) + + def __init__(self) -> None: + self.steps: List[Tuple[str, Any, Any]] = [] + + def check(self, label: str, value: Any, ok: bool) -> bool: + """Record a pass/fail criterion and return it unchanged.""" + self.steps.append((label, value, ok or None)) + return ok + + def resolve(self, label: str, value: Any, outcome: Any) -> Any: + """Record a criterion that names something, and return that name.""" + self.steps.append((label, value, outcome)) + return outcome + + +class _NullTrace(Trace): + __slots__ = () + + def __init__(self) -> None: # no list to append to + pass + + def check(self, label: str, value: Any, ok: bool) -> bool: + return ok + + def resolve(self, label: str, value: Any, outcome: Any) -> Any: + return outcome + + +NULL_TRACE = _NullTrace() + + +def routing_module(arch: str): + """Import and return the routing module for ``arch``, or None.""" + name = STAIRCASE_ROUTERS.get(arch) + if name is None: + return None + return import_module(f"{_PACKAGE}.{name}") + + +def staircase_resolve(config: "ModelConfig") -> Optional[str]: + """Rewrite ``architectures[0]`` into a staircase target's class name. + + Driven by ``TRTLLM_STAIRCASE``; see ``STAIRCASE_ENV`` for why it is an + environment variable and when it has to be set. + + Returns None when staircase is off, when no routing module claims the + architecture, or when the routing tree finds no matching target -- in + ``auto`` the caller then falls back to the built-in implementation. In + ``require`` a non-match raises instead, because the failure this mode + exists to prevent is silent: asking for a target that does not exist, + getting the in-tree implementation, and reading the resulting curve as + staircase's. + """ + mode = StaircaseMode.from_env() + if mode is StaircaseMode.OFF: + return None + + pretrained_config = config.pretrained_config + if not getattr(pretrained_config, "architectures", None): + return None + + ctx = StaircaseContext.from_model_config(config) + arch = pretrained_config.architectures[0] + trace = Trace() if mode is StaircaseMode.REQUIRE else NULL_TRACE + + routing = routing_module(arch) + target = None if routing is None else routing.route(ctx, trace) + + if target is None: + if mode is StaircaseMode.REQUIRE: + raise ValueError(explain_no_match(arch, routing, trace)) + return None + + # The synthetic name is only a registry key: importing the target module + # is what puts the class behind it. Nothing else would -- the built-in + # static index does not, and must not, carry staircase names. + import_module(f"{_PACKAGE}.{routing.TARGET_MODULES[target]}") + return target + + +def explain_no_match(arch: str, routing, trace: Trace) -> str: + """Say which criterion the configuration failed, not just that it did.""" + if routing is None: + known = ", ".join(sorted(STAIRCASE_ROUTERS)) or "(none)" + return ( + f"staircase is set to 'require' but no staircase target " + f"exists for architecture {arch!r}; routed architectures: " + f"{known}" + ) + lines = [ + f"staircase is set to 'require' but no target matched " + f"architecture {arch!r}. The decision tree in " + f"{routing.__name__.rpartition('.')[0].replace('.', '/')}/routing.py " + f"got as far as:" + ] + for label, value, outcome in trace.steps: + mark = "no match" if outcome is None else f"-> {outcome}" + lines.append(f" {label:<10} {value!r:<48} {mark}") + if not trace.steps: + lines.append(" (the tree rejected the configuration before its first recorded criterion)") + return "\n".join(lines) diff --git a/tensorrt_llm/_torch/staircase/_routing_test.py b/tensorrt_llm/_torch/staircase/_routing_test.py new file mode 100644 index 000000000000..41ee7e59bd16 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/_routing_test.py @@ -0,0 +1,230 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""What ``staircase_resolve`` actually does, driven by synthetic configs. + +No checkpoint and no weights: routing reads config *shape*, the mapping and +the SM version, all of which can be stated directly. The SM version is +monkeypatched so these run on any device -- the point here is the decision +logic, not the kernels. + +The three things worth proving: + +* ``off`` changes nothing. This is the whole safety argument for putting the + hook in ``_resolve_class`` at all. +* a matching configuration reaches a target class, and that class is + *external* -- so it wins the registry slot rather than losing it to the + built-in provider, which the lazy zoo may import afterwards. +* ``require`` raises on a near-miss and says which criterion missed. +""" + +from __future__ import annotations + +import pytest +import torch +from transformers import PretrainedConfig + +from tensorrt_llm._torch.model_config import ModelConfig +from tensorrt_llm._torch.models.modeling_auto import AutoModelForCausalLM +from tensorrt_llm._torch.models.modeling_utils import ( + _is_builtin_model_class, + get_registered_model_class, +) +from ._router_index import ( + STAIRCASE_ENV, + StaircaseMode, + staircase_resolve, +) +from tensorrt_llm.mapping import Mapping + +_SM103 = (10, 3) + + +@pytest.fixture(autouse=True) +def _on_sm103(monkeypatch): + """Route as if this were a GB300, wherever the test actually runs.""" + monkeypatch.setattr(torch.cuda, "is_available", lambda: True) + monkeypatch.setattr(torch.cuda, "get_device_capability", lambda *a, **k: _SM103) + + +def _gpt_oss_config(**overrides): + """The shape gpt-oss-120b's own config.json declares.""" + fields = dict( + architectures=["GptOssForCausalLM"], + model_type="gpt_oss", + num_hidden_layers=36, + hidden_size=2880, + num_local_experts=128, + ) + fields.update(overrides) + return PretrainedConfig(**fields) + + +def _r1_config(**overrides): + """The shape DeepSeek-R1-0528-NVFP4's own config.json declares.""" + fields = dict( + architectures=["DeepseekV3ForCausalLM"], + model_type="deepseek_v3", + num_hidden_layers=61, + hidden_size=7168, + n_routed_experts=256, + q_lora_rank=1536, + ) + fields.update(overrides) + return PretrainedConfig(**fields) + + +@pytest.fixture(autouse=True) +def _mode(monkeypatch, request): + """Default every case to 'auto'; a case that wants another mode sets it. + + The switch is an environment variable, so the tests set one too -- that is + the surface under test. + """ + monkeypatch.setenv(STAIRCASE_ENV, "auto") + + +def _set_mode(monkeypatch, mode): + if mode is None: + monkeypatch.delenv(STAIRCASE_ENV, raising=False) + else: + monkeypatch.setenv(STAIRCASE_ENV, mode) + + +def _model_config(pretrained_config, **mapping_kwargs): + mapping = Mapping(**mapping_kwargs) if mapping_kwargs else Mapping() + return ModelConfig(pretrained_config=pretrained_config, mapping=mapping) + + +_DEP4 = dict(world_size=4, tp_size=4, moe_ep_size=4, moe_tp_size=1, enable_attention_dp=True) + + +@pytest.mark.parametrize("unset", [True, False], ids=["env-unset", "env-off"]) +def test_off_resolves_nothing(monkeypatch, unset): + """The default must be indistinguishable from staircase not existing. + + Unset and an explicit "off" have to behave identically: the common case is + that nobody has heard of this package. + """ + _set_mode(monkeypatch, None if unset else "off") + config = _model_config(_gpt_oss_config()) + assert staircase_resolve(config) is None + + +def test_off_still_reaches_the_builtin_implementation(monkeypatch): + _set_mode(monkeypatch, "off") + config = _model_config(_gpt_oss_config()) + resolved = AutoModelForCausalLM._resolve_class(config) + assert resolved is not None + assert resolved.__module__ == "tensorrt_llm._torch.models.modeling_gpt_oss" + + +@pytest.mark.parametrize("mode", ["auto", "require"]) +def test_gpt_oss_tp1_matches(monkeypatch, mode): + _set_mode(monkeypatch, mode) + config = _model_config(_gpt_oss_config()) + assert staircase_resolve(config) == "StaircaseGptOss120bSm103Tp1" + + +@pytest.mark.parametrize("mode", ["auto", "require"]) +def test_r1_dep4_matches(monkeypatch, mode): + _set_mode(monkeypatch, mode) + config = _model_config(_r1_config(), **_DEP4) + assert staircase_resolve(config) == "StaircaseDeepseekR10528Nvfp4Sm103Dep4" + + +def test_resolving_registers_the_target_class(): + """The synthetic name is a key; the import behind it is what fills it.""" + config = _model_config(_gpt_oss_config()) + name = staircase_resolve(config) + cls = get_registered_model_class(name) + assert cls is not None, f"{name} resolved to no class" + assert cls.__name__ == name + assert cls.__module__.endswith( + "staircase.models.gpt_oss.targets.gpt_oss_120b.sm_103.tp1.modeling" + ) + + +def test_the_target_registration_counts_as_external(): + """External registrations always win their slot; built-ins only fill + empty ones. Living beside the zoo rather than inside it is what buys + this, and a move into _torch/models/ would silently reverse it.""" + config = _model_config(_gpt_oss_config()) + cls = get_registered_model_class(staircase_resolve(config)) + assert not _is_builtin_model_class(cls) + + +def test_resolve_class_rewrites_the_architecture_end_to_end(): + config = _model_config(_gpt_oss_config()) + resolved = AutoModelForCausalLM._resolve_class(config) + assert resolved.__name__ == "StaircaseGptOss120bSm103Tp1" + + +@pytest.mark.parametrize( + "config_kwargs, mapping_kwargs, missed", + [ + # a GptOss checkpoint of another size + (dict(num_hidden_layers=24, num_local_experts=32), {}, "shape"), + # the right checkpoint, a topology no target implements + (dict(), dict(world_size=2, tp_size=2), "parallel"), + ], +) +def test_gpt_oss_near_misses_do_not_match(monkeypatch, config_kwargs, mapping_kwargs, missed): + config = _model_config(_gpt_oss_config(**config_kwargs), **mapping_kwargs) + assert staircase_resolve(config) is None + + _set_mode(monkeypatch, "require") + with pytest.raises(ValueError, match=missed): + staircase_resolve(config) + + +def test_r1_without_attention_dp_does_not_match(): + """dep4 and tep4 differ in where the attention weights are split, so + attention DP is an identity criterion rather than a knob.""" + mapping_kwargs = dict(_DEP4, enable_attention_dp=False) + config = _model_config(_r1_config(), **mapping_kwargs) + assert staircase_resolve(config) is None + + +def test_require_names_the_criterion_that_missed(monkeypatch): + monkeypatch.setattr(torch.cuda, "get_device_capability", lambda *a, **k: (10, 0)) + _set_mode(monkeypatch, "require") + config = _model_config(_gpt_oss_config()) + with pytest.raises(ValueError) as excinfo: + staircase_resolve(config) + message = str(excinfo.value) + assert "sm" in message and "(10, 0)" in message + assert "no match" in message + + +def test_an_unrouted_architecture_is_not_an_error_under_auto(): + config = _model_config(PretrainedConfig(architectures=["LlamaForCausalLM"])) + assert staircase_resolve(config) is None + + +def test_an_unrouted_architecture_raises_under_require(monkeypatch): + _set_mode(monkeypatch, "require") + config = _model_config(PretrainedConfig(architectures=["LlamaForCausalLM"])) + with pytest.raises(ValueError, match="LlamaForCausalLM"): + staircase_resolve(config) + + +@pytest.mark.parametrize( + "raw,expected", [("off", "off"), ("AUTO", "auto"), (" require ", "require"), ("", "off")] +) +def test_the_env_var_is_read_leniently(monkeypatch, raw, expected): + monkeypatch.setenv(STAIRCASE_ENV, raw) + assert StaircaseMode.from_env().value == expected + + +def test_an_unknown_mode_raises_rather_than_falling_back(monkeypatch): + """A typo must not read as "off". + + That would hand back the built-in implementation while the caller believed + they had asked for a target -- the exact mis-attribution the require mode + exists to prevent. + """ + # "yes" rather than a misspelling: it is what someone reaching for a + # boolean would write, and it is the reading that must not be invented. + monkeypatch.setenv(STAIRCASE_ENV, "yes") + with pytest.raises(ValueError, match="not a staircase mode"): + StaircaseMode.from_env() diff --git a/tensorrt_llm/_torch/staircase/catalog/__init__.py b/tensorrt_llm/_torch/staircase/catalog/__init__.py new file mode 100644 index 000000000000..3854e1b320f8 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/__init__.py @@ -0,0 +1,6 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""The kernel vocabulary: one entry = contract .md + wrapper .py + GPU test. + +Nothing is imported here -- an entry is pulled in by the target that uses +it, so a process loads only the ops its target calls.""" diff --git a/tensorrt_llm/_torch/staircase/catalog/activation/__init__.py b/tensorrt_llm/_torch/staircase/catalog/activation/__init__.py new file mode 100644 index 000000000000..cea4f64ab601 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/activation/__init__.py @@ -0,0 +1,3 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Activation entries.""" diff --git a/tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul.md b/tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul.md new file mode 100644 index 000000000000..4121d98883d4 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul.md @@ -0,0 +1,85 @@ +--- +receipts: + sm_100: {status: passed, trtllm: 1.3.0rc21} + sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 4} +--- + +# flashinfer_silu_and_mul + +**Wraps** `torch.ops.trtllm.flashinfer_silu_and_mul` (one call). + +## Semantics + +SiLU-gated elementwise multiply (the SwiGLU MLP activation), fused in one +kernel. The last dim of the input holds the gate half followed by the up +half: + +``` +d = x.shape[-1] // 2 +out[..., i] = silu(x[..., i]) * x[..., d + i] for i in [0, d) +silu(v) = v / (1 + exp(-v)) +``` + +Both halves are loaded and the silu/multiply are computed in fp32 inside +the kernel; the result is cast back to the input dtype. + +Fusion boundary: the single call computes activation and gating multiply +only. The caller owns the gate/up projection that produces `x` (whether as +one fused matmul or two) and the down projection that consumes the output. +There is no quantization of the output (a separate +`torch.ops.trtllm.silu_and_mul` op exists with an optional quant scale) and +no gelu variant (a separate `torch.ops.trtllm.flashinfer_gelu_tanh_and_mul` +op exists). + +## Signature + +```python +def flashinfer_silu_and_mul(x: torch.Tensor) -> torch.Tensor +``` + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `x` | `[..., 2 * d]`, e.g. `[num_tokens, 2 * intermediate]` | fp16 / bf16 | contiguous | CUDA | +| returns | `[..., d]` | same as `x` | newly allocated, contiguous | CUDA (same device as `x`) | + +The output is a new tensor (`x` is not mutated). + +## Metadata consumed + +None. Stateless. + +## Preconditions + +- `x` is on a CUDA device and fully contiguous: the kernel indexes rows as + `token * 2d` on the raw pointer, so any non-contiguity silently corrupts + results. +- `x.shape[-1]` is even and a multiple of 16 elements for fp16/bf16 + (equivalently `d * itemsize % 16 == 0`). This is stricter than the + op's own runtime check — see Notes. +- `x.shape[-1] >= 16` for fp16/bf16 (`d` must hold at least one 16-byte + vector, or the computed block size is 0 and the launch is invalid). +- Dtype is fp16 or bf16. Other dtypes (including fp32) are rejected by the + kernel's dtype dispatch. +- No upper bound on `d` beyond memory: `d` larger than + `1024 * (16 / itemsize)` (e.g. > 8192 for bf16) takes a scalar remainder + loop, verified correct on this machine. + +## Notes + +- The op is registered only when flashinfer is importable + (`IS_FLASHINFER_AVAILABLE`); this pinned install ships + flashinfer-python 0.6.14. +- Alignment trap: the op's Python-side check only validates + `x.shape[-1] * itemsize % 16 == 0` (raising `ValueError`), but the + vectorized load of the up half starts at element offset `d`, so `d` + itself must be 16-byte aligned. Inputs that pass the check with + `d * itemsize % 16 != 0` (e.g. bf16 with last dim 24 or 2728) crash with + CUDA `misaligned address` — observed on sm_100. Keep the last dim a + multiple of 16 elements. +- The misaligned-address failure is sticky: it poisons the CUDA context + for the rest of the process. +- Programmatic dependent launch (PDL) is controlled by the env var + `TRTLLM_ENABLE_PDL` (default enabled) inside the trtllm custom op; it + affects scheduling only, not results. +- silu uses the fast-math `__expf`; observed error vs an fp32 torch + reference stays within default `assert_close` tolerances for fp16/bf16. diff --git a/tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul.py b/tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul.py new file mode 100644 index 000000000000..611c999e0338 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul.py @@ -0,0 +1,12 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""SiLU-gated multiply (SwiGLU activation) via the flashinfer silu_and_mul kernel.""" + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + + +def flashinfer_silu_and_mul(x: torch.Tensor) -> torch.Tensor: + """Return `silu(x[..., :d]) * x[..., d:]` with `d = x.shape[-1] // 2` as a new tensor.""" + return torch.ops.trtllm.flashinfer_silu_and_mul(x) diff --git a/tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul_test.py b/tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul_test.py new file mode 100644 index 000000000000..af19e2ef7d1e --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul_test.py @@ -0,0 +1,56 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the flashinfer_silu_and_mul catalog entry.""" + +import torch +import torch.nn.functional as F + +from .flashinfer_silu_and_mul import flashinfer_silu_and_mul + +assert torch.cuda.is_available(), "flashinfer_silu_and_mul requires a CUDA device" + + +def _ref_silu_and_mul(x: torch.Tensor) -> torch.Tensor: + """fp32-accumulated reference: silu(x[..., :d]) * x[..., d:].""" + gate, up = x.float().chunk(2, dim=-1) + return (F.silu(gate) * up).to(x.dtype) + + +def _check(x: torch.Tensor) -> None: + out = flashinfer_silu_and_mul(x) + ref = _ref_silu_and_mul(x) + assert out.shape == x.shape[:-1] + (x.shape[-1] // 2,) + assert out.dtype == x.dtype + torch.testing.assert_close(out, ref) + + +def test_bf16_2d() -> None: + torch.manual_seed(0) + # decode-like (few tokens) and prefill-like (many tokens) shapes; + # last dim is 2 * intermediate_size (gate half then up half) + for num_tokens, two_d in [(1, 8192), (4, 28672), (2048, 8192)]: + x = torch.randn(num_tokens, two_d, dtype=torch.bfloat16, device="cuda") + _check(x) + + +def test_bf16_3d() -> None: + torch.manual_seed(1) + x = torch.randn(4, 16, 2048, dtype=torch.bfloat16, device="cuda") + _check(x) + + +def test_bf16_edge_sizes() -> None: + # 16 is the smallest legal last dim (d = 8 = one vector); + # 16400 gives d = 8200 > 8192 = blockDim * vec_size, exercising the + # scalar remainder loop after the vectorized loop + torch.manual_seed(2) + for two_d in [16, 16400]: + x = torch.randn(16, two_d, dtype=torch.bfloat16, device="cuda") + _check(x) + + +def test_fp16_2d() -> None: + torch.manual_seed(3) + for num_tokens, two_d in [(2, 8192), (1024, 4096)]: + x = torch.randn(num_tokens, two_d, dtype=torch.float16, device="cuda") + _check(x) diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/__init__.py b/tensorrt_llm/_torch/staircase/catalog/attention/__init__.py new file mode 100644 index 000000000000..e45ba989dd6a --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/attention/__init__.py @@ -0,0 +1,3 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Attention entries.""" diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope.md b/tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope.md new file mode 100644 index 000000000000..156153e4bf79 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope.md @@ -0,0 +1,160 @@ +--- +receipts: + sm_100: {status: passed, trtllm: 1.3.0rc21} + sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 9} +--- + +# fused_qk_norm_rope + +**Wraps** `torch.ops.trtllm.fused_qk_norm_rope` (one call). + +## Semantics + +In-place, on a packed QKV activation `qkv` of shape +`[num_tokens, (num_heads_q + num_heads_k + num_heads_v) * head_dim]` whose +rows lay out `num_heads_q` query heads, then `num_heads_k` key heads, then +`num_heads_v` value heads, each `head_dim` wide. One call computes, per +token, for every **q and k head** `x` (fp32 internally, one final cast back +to bf16): + +``` +# 1. per-head RMS norm over the full head_dim (skipped if is_qk_norm=False) +h = x / sqrt(mean(x^2) + eps) * w # w = q_weight or k_weight + # use_gemma: * (1 + w) instead +# 2. RoPE on the first rotary_dim dims of h; h[rotary_dim:] passes through +inv_freq[j] = base^(-2j/rotary_dim) for j in [0, rotary_dim/2) +angle[j] = position_id * inv_freq[j] +cos, sin = cos(angle) * attention_factor, sin(angle) * attention_factor +# is_neox=True — half-split pairs (x1, x2) = (h[j], h[j + rotary_dim/2]) +# is_neox=False — interleaved pairs (x1, x2) = (h[2j], h[2j+1]) +(x1, x2) -> (x1*cos - x2*sin, x1*sin + x2*cos) +``` + +Value heads are passed through bit-exactly untouched. + +YaRN scaling (`factor`, `low`, `high`, `attention_factor`) replaces +`inv_freq` with the blend + +``` +ramp[j] = clamp((j - low) / (high - low), 0, 1) # high==low: high += 0.001 +inv_freq[j] = (1 - ramp[j]) * base^(-2j/rotary_dim) + + ramp[j] * base^(-2j/rotary_dim) / factor +``` + +which reduces to plain RoPE at `factor=1.0` (the `low`/`high` values are +then irrelevant). `attention_factor` scales cos/sin — but **not at +`factor == 1.0` exactly**, where any `attention_factor != 1.0` raises +`Assertion failed: attention_factor == 1.0f` +(`fusedQKNormRopeKernel.cu:322`). `factor = 1.0000001` accepts it. +Measured 2026-07-28; the gate is on `factor` alone and is not otherwise +documented. + +Interleaved mRoPE (`use_mrope=True`) takes 3 position rows +(temporal/height/width) per token. Frequency index `j` reads its position +from row 1 when `j % 3 == 1 and j < 3*mrope_section1`, from row 2 when +`j % 3 == 2 and j < 3*mrope_section2`, and from row 0 otherwise. Only this +interleaved layout is implemented — not the contiguous-section mRoPE +variant. + +Fusion boundary: q-norm + k-norm + RoPE, nothing else. The caller owns the +QKV projection before the call and everything after it (KV-cache append, +attention, output projection). A sibling op `fused_dit_qk_norm_rope` +exists for the DiT-style variant of this fusion. + +## Signature + +```python +def fused_qk_norm_rope( + qkv: torch.Tensor, + num_heads_q: int, + num_heads_k: int, + num_heads_v: int, + head_dim: int, + rotary_dim: int, + eps: float, + q_weight: torch.Tensor, + k_weight: torch.Tensor, + base: float, + is_neox: bool, + position_ids: torch.Tensor, + factor: float = 1.0, + low: float = 0.0, + high: float = 0.0, + attention_factor: float = 1.0, + is_qk_norm: bool = True, + use_gemma: bool = False, + use_mrope: bool = False, + mrope_section1: int = 0, + mrope_section2: int = 0, +) -> None +``` + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `qkv` | `[num_tokens, (num_heads_q+num_heads_k+num_heads_v)*head_dim]` (2D only) | bf16 | contiguous | CUDA | +| `num_heads_q/k/v` | scalar | Python int | — | — | +| `head_dim` | scalar, one of 64 / 128 / 256 | Python int | — | — | +| `rotary_dim` | scalar, even, `<= head_dim` | Python int | — | — | +| `eps` | scalar | Python float | — | — | +| `q_weight`, `k_weight` | `[head_dim]` | bf16 | contiguous | CUDA (same device) | +| `base` | scalar (RoPE theta) | Python float | — | — | +| `is_neox` | scalar | Python bool | — | — | +| `position_ids` | `[num_tokens]`, or `[3, num_tokens]` row-major when `use_mrope` | int32 | contiguous | CUDA (same device) | +| `factor`, `low`, `high`, `attention_factor` | scalar (YaRN; defaults = plain RoPE) | Python float | — | — | +| `is_qk_norm` | scalar (False: RoPE only, weights ignored) | Python bool | — | — | +| `use_gemma` | scalar (True: scale by `1 + w`) | Python bool | — | — | +| `use_mrope`, `mrope_section1`, `mrope_section2` | scalar (interleaved mRoPE) | Python bool / int | — | — | +| returns | — | — | — | — | + +Returns `None`. `qkv` is mutated in place: q and k head slices are +overwritten with the normed + rotated values; v head slices are untouched. + +## Metadata consumed + +None. Stateless — all RoPE frequencies are derived on the fly from the +scalar arguments; no cos/sin cache and no attention metadata are read. +`position_ids` are the absolute positions of the `num_tokens` tokens, in +row order of `qkv` (batch layout is the caller's concern; the kernel is +purely per-token). + +## Preconditions + +- `qkv` is 2D, contiguous, bf16, on CUDA, with + `qkv.shape[1] == (num_heads_q + num_heads_k + num_heads_v) * head_dim`. + bf16 is the only accepted dtype (fp16/fp32 raise). +- `head_dim` is one of 64, 128, 256. Others (32, 80, 96, 192 verified) + raise `Unsupported head dimension for fusedQKNormRope`. +- `q_weight` and `k_weight` are contiguous bf16 `[head_dim]` tensors on + the same device; they must be passed (and valid) even when + `is_qk_norm=False`, though their values are then unused. +- `position_ids` is contiguous int32 on the same device: flat + `[num_tokens]` normally, `[3, num_tokens]` row-major when + `use_mrope=True`. int64 raises. A 3D upstream view must be reshaped and + made contiguous by the caller. +- `rotary_dim` is even and at most `head_dim`; dims beyond `rotary_dim` + are normed (if `is_qk_norm`) but not rotated. **Only the evenness half + is enforced** (`Assertion failed: rotary_dim must be even`); + `rotary_dim > head_dim` is accepted silently and rewrites the q and k + head columns — measured at 160 and 256 against `head_dim = 128`, + 2026-07-28. This is the caller's obligation, not the op's. +- Everything else above is validated by the op with a clear error, and in + every raising case `qkv` comes back bit-identical to its pre-call + contents. The two exceptions are the `rotary_dim` bound just named and + the `attention_factor` gate under *Semantics*. + +## Notes + +- The kernel checks dtype/shape/contiguity eagerly (thop-level asserts), + so contract violations fail loudly rather than corrupting `qkv`. +- Accuracy: internal math is fp32 with a single bf16 round at the end. + Because RoPE rotates a bf16-rounded pair, the per-element error is + absolute in the pair magnitude — with unit-scale inputs, max abs error + observed on sm_100 was 0.031, at most 78% of the test's + `rtol=1.6e-2, atol=2e-2` allowance (compare with `atol` sized to the q/k + magnitude, not the default bf16 `atol=1e-5`). The ceiling sits on the + neox-prefill and Gemma-norm cases; most of the file sits at 0.008-0.016. +- `low`/`high` are float half-dim indices in `[0, rotary_dim/2)` + (fractional values legal, matching YaRN's `truncate=False`). +- TRT-LLM's own caller derives `factor/low/high/attention_factor` from a + YaRN config and uses `rotary_dim = head_dim * partial_rotary_factor`; + this entry exposes them raw. diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope.py b/tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope.py new file mode 100644 index 000000000000..3dfc7ae6b29b --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope.py @@ -0,0 +1,56 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""In-place fused per-head QK RMS norm + RoPE on a packed QKV tensor.""" + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + + +def fused_qk_norm_rope( + qkv: torch.Tensor, + num_heads_q: int, + num_heads_k: int, + num_heads_v: int, + head_dim: int, + rotary_dim: int, + eps: float, + q_weight: torch.Tensor, + k_weight: torch.Tensor, + base: float, + is_neox: bool, + position_ids: torch.Tensor, + factor: float = 1.0, + low: float = 0.0, + high: float = 0.0, + attention_factor: float = 1.0, + is_qk_norm: bool = True, + use_gemma: bool = False, + use_mrope: bool = False, + mrope_section1: int = 0, + mrope_section2: int = 0, +) -> None: + """In-place on qkv: per-head RMS norm of q/k heads, then RoPE. Returns None.""" + torch.ops.trtllm.fused_qk_norm_rope( + qkv, + num_heads_q, + num_heads_k, + num_heads_v, + head_dim, + rotary_dim, + eps, + q_weight, + k_weight, + base, + is_neox, + position_ids, + factor, + low, + high, + attention_factor, + is_qk_norm, + use_gemma, + use_mrope, + mrope_section1, + mrope_section2, + ) diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope_test.py b/tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope_test.py new file mode 100644 index 000000000000..66126930e796 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope_test.py @@ -0,0 +1,252 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the fused_qk_norm_rope catalog entry.""" + +import torch + +from .fused_qk_norm_rope import fused_qk_norm_rope + +assert torch.cuda.is_available(), "fused_qk_norm_rope requires a CUDA device" + +# The kernel rotates the bf16-rounded (x1, x2) pair, so the per-element +# error is absolute in the pair magnitude (up to ~5 after RMS norm with +# weights near 1), not relative to each output element: 0.02 ~= 5 ulp of +# bf16 (2^-8) at magnitude 5. Observed max abs err across all cases is +# 0.016; the default bf16 atol of 1e-5 is unreachable for near-zero +# outputs produced by rotating large pairs. +ATOL = 0.02 +RTOL = 1.6e-2 # torch default for bf16 + + +def _inv_freq(rotary_dim: int, base: float, factor: float, low: float, high: float) -> torch.Tensor: + """YaRN-blended inverse frequencies; plain RoPE when factor == 1.""" + half = rotary_dim // 2 + j = torch.arange(half, dtype=torch.float32, device="cuda") + pos_freqs = base ** (2.0 * j / rotary_dim) + inv_extrapolation = 1.0 / pos_freqs + inv_interpolation = 1.0 / (factor * pos_freqs) + if high == low: + high = high + 0.001 # kernel guards the ramp singularity the same way + ramp = ((j - low) / (high - low)).clamp(0.0, 1.0) + extrapolation_factor = 1.0 - ramp + return ( + inv_interpolation * (1.0 - extrapolation_factor) + inv_extrapolation * extrapolation_factor + ) + + +def _ref( + qkv: torch.Tensor, + num_heads_q: int, + num_heads_k: int, + num_heads_v: int, + head_dim: int, + rotary_dim: int, + eps: float, + q_weight: torch.Tensor, + k_weight: torch.Tensor, + base: float, + is_neox: bool, + position_ids: torch.Tensor, + factor: float = 1.0, + low: float = 0.0, + high: float = 0.0, + attention_factor: float = 1.0, + is_qk_norm: bool = True, + use_gemma: bool = False, + use_mrope: bool = False, + mrope_section1: int = 0, + mrope_section2: int = 0, +) -> torch.Tensor: + """fp32 reference: per-head RMS norm on q/k heads, then RoPE; v untouched.""" + num_heads = num_heads_q + num_heads_k + num_heads_v + x = qkv.float().view(-1, num_heads, head_dim).clone() + half = rotary_dim // 2 + inv_freq = _inv_freq(rotary_dim, base, factor, low, high) + if use_mrope: + angle = position_ids.float()[:, :, None] * inv_freq # [3, T, half] + cos, sin = angle.cos(), angle.sin() + + def pick(c: torch.Tensor) -> torch.Tensor: + out = c[0].clone() + out[:, 1 : mrope_section1 * 3 : 3] = c[1][:, 1 : mrope_section1 * 3 : 3] + out[:, 2 : mrope_section2 * 3 : 3] = c[2][:, 2 : mrope_section2 * 3 : 3] + return out + + cos, sin = pick(cos), pick(sin) + else: + angle = position_ids.float()[:, None] * inv_freq # [T, half] + cos, sin = angle.cos(), angle.sin() + cos = cos[:, None, :] * attention_factor + sin = sin[:, None, :] * attention_factor + for start, count, weight in ( + (0, num_heads_q, q_weight.float()), + (num_heads_q, num_heads_k, k_weight.float()), + ): + h = x[:, start : start + count, :] + if is_qk_norm: + h = h * torch.rsqrt(h.pow(2).mean(-1, keepdim=True) + eps) + h = h * ((1.0 + weight) if use_gemma else weight) + r = h[..., :rotary_dim] + if is_neox: + x1, x2 = r[..., :half], r[..., half:] + rotated = torch.cat([x1 * cos - x2 * sin, x1 * sin + x2 * cos], dim=-1) + else: + x1, x2 = r[..., ::2], r[..., 1::2] + rotated = torch.stack([x1 * cos - x2 * sin, x1 * sin + x2 * cos], dim=-1).flatten(-2) + x[:, start : start + count, :] = torch.cat([rotated, h[..., rotary_dim:]], -1) + return x.view(qkv.shape[0], -1) + + +def _check( + num_tokens: int, + num_heads_q: int, + num_heads_kv: int, + head_dim: int, + rotary_dim: int, + eps: float, + base: float, + is_neox: bool, + factor: float = 1.0, + low: float = 0.0, + high: float = 0.0, + attention_factor: float = 1.0, + is_qk_norm: bool = True, + use_gemma: bool = False, + use_mrope: bool = False, + mrope_section1: int = 0, + mrope_section2: int = 0, +) -> None: + width = (num_heads_q + 2 * num_heads_kv) * head_dim + qkv = torch.randn(num_tokens, width, dtype=torch.bfloat16, device="cuda") + q_w = torch.randn(head_dim, dtype=torch.bfloat16, device="cuda") * 0.1 + 1.0 + k_w = torch.randn(head_dim, dtype=torch.bfloat16, device="cuda") * 0.1 + 1.0 + shape = (3, num_tokens) if use_mrope else (num_tokens,) + pos = torch.randint(0, 32768, shape, dtype=torch.int32, device="cuda") + expected = _ref( + qkv, num_heads_q, num_heads_kv, num_heads_kv, head_dim, rotary_dim, + eps, q_w, k_w, base, is_neox, pos, factor, low, high, attention_factor, + is_qk_norm, use_gemma, use_mrope, mrope_section1, mrope_section2, + ).to(torch.bfloat16) # fmt: skip + fused_qk_norm_rope( + qkv, num_heads_q, num_heads_kv, num_heads_kv, head_dim, rotary_dim, + eps, q_w, k_w, base, is_neox, pos, factor, low, high, attention_factor, + is_qk_norm, use_gemma, use_mrope, mrope_section1, mrope_section2, + ) # fmt: skip + torch.testing.assert_close(qkv, expected, rtol=RTOL, atol=ATOL) + # v heads must pass through bit-exactly + v = qkv.view(num_tokens, -1, head_dim)[:, num_heads_q + num_heads_kv :, :] + ev = expected.view(num_tokens, -1, head_dim)[:, num_heads_q + num_heads_kv :, :] + assert torch.equal(v, ev) + + +def test_bf16_neox_head_dims() -> None: + torch.manual_seed(0) + # decode-like shapes over every supported head_dim + for head_dim in (64, 128, 256): + _check( + 2, 8, 2, head_dim=head_dim, rotary_dim=head_dim, eps=1e-6, + base=10000.0, is_neox=True, + ) # fmt: skip + + +def test_bf16_neox_prefill() -> None: + torch.manual_seed(1) + # prefill-like: many tokens, Qwen3-32B-like head layout + _check( + 8192, 32, 8, head_dim=128, rotary_dim=128, eps=1e-6, + base=1000000.0, is_neox=True, + ) # fmt: skip + + +def test_bf16_interleaved() -> None: + torch.manual_seed(2) + for num_tokens in (1, 65): + _check( + num_tokens, 8, 2, head_dim=128, rotary_dim=128, eps=1e-6, + base=10000.0, is_neox=False, + ) # fmt: skip + + +def test_bf16_gemma_norm() -> None: + torch.manual_seed(3) + _check( + 17, 4, 4, head_dim=64, rotary_dim=64, eps=1e-6, + base=10000.0, is_neox=True, use_gemma=True, + ) # fmt: skip + + +def test_bf16_rope_only() -> None: + torch.manual_seed(4) + # is_qk_norm=False skips the norm entirely; weights are ignored + _check( + 17, 8, 2, head_dim=128, rotary_dim=128, eps=1e-6, + base=10000.0, is_neox=True, is_qk_norm=False, + ) # fmt: skip + + +def test_bf16_partial_rotary() -> None: + torch.manual_seed(5) + # norm covers the full head_dim; rope covers only the first rotary_dim + for is_neox in (True, False): + _check( + 17, 8, 2, head_dim=128, rotary_dim=64, eps=1e-6, + base=10000.0, is_neox=is_neox, + ) # fmt: skip + + +def test_bf16_yarn() -> None: + torch.manual_seed(6) + _check( + 33, 8, 2, head_dim=128, rotary_dim=128, eps=1e-6, + base=1000000.0, is_neox=True, + factor=4.0, low=4.0, high=20.0, attention_factor=1.2, + ) # fmt: skip + + +def test_bf16_mrope_interleaved() -> None: + torch.manual_seed(7) + _check( + 19, 8, 2, head_dim=128, rotary_dim=128, eps=1e-6, + base=10000.0, is_neox=True, + use_mrope=True, mrope_section1=16, mrope_section2=24, + ) # fmt: skip + + +def test_rejects_out_of_contract() -> None: + torch.manual_seed(8) + pos = torch.arange(4, dtype=torch.int32, device="cuda") + w = torch.ones(128, dtype=torch.bfloat16, device="cuda") + + def call(qkv, head_dim=128, q_w=w, position_ids=pos): + fused_qk_norm_rope( + qkv, 8, 2, 2, head_dim, head_dim, 1e-6, q_w, w, + 10000.0, True, position_ids, + ) # fmt: skip + + # non-bf16 qkv + for dtype in (torch.float16, torch.float32): + try: + call(torch.randn(4, 12 * 128, dtype=dtype, device="cuda")) + raise AssertionError(f"{dtype} qkv unexpectedly accepted") + except RuntimeError: + pass + # unsupported head_dim + try: + w96 = torch.ones(96, dtype=torch.bfloat16, device="cuda") + call( + torch.randn(4, 12 * 96, dtype=torch.bfloat16, device="cuda"), + head_dim=96, + q_w=w96, + ) + raise AssertionError("head_dim=96 unexpectedly accepted") + except RuntimeError: + pass + # non-int32 position_ids + try: + call( + torch.randn(4, 12 * 128, dtype=torch.bfloat16, device="cuda"), + position_ids=pos.to(torch.int64), + ) + raise AssertionError("int64 position_ids unexpectedly accepted") + except RuntimeError: + pass diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla.md b/tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla.md new file mode 100644 index 000000000000..df9ca8b988ca --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla.md @@ -0,0 +1,317 @@ +--- +receipts: + sm_100: {status: passed, trtllm: 1.3.0rc21} + sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 18} +--- + +# load_paged_kv_cache_for_mla + +**Wraps** `torch.ops.trtllm.load_paged_kv_cache_for_mla` (one call). + +## Semantics + +The cache-read step of MLA context-phase attention with KV reuse: after a +context batch's full latent KV (cached prefix + this step's new tokens) is +resident in the paged MLA cache, one call gathers it back into dense +tensors so the caller can up-project and run context FMHA over the full +`[past + new]` range. + +For each context sequence `s` in `[0, num_contexts)`, with total latent-KV +length `L_s = cu_ctx_kv_lens[s+1] - cu_ctx_kv_lens[s]`, the call reads the +sequence's paged latent rows at positions `[0, L_s)` — each row is the +per-token `[compressed_kv | k_pe]` of width `C + R` +(`C = kv_lora_rank`, `R = qk_rope_head_dim`) — and writes them, +sequences concatenated in batch order, into two freshly allocated +contiguous tensors: + +``` +T = num_ctx_kv_tokens = sum(L_s) +compressed_kv[t, :] = cache_row(t)[:C] # [T, C] +k_pe[t, :] = cache_row(t)[C:] # [T, R] +``` + +The paged pool is never modified, and sequences at batch index +`>= num_contexts` (generation sequences) contribute nothing to the outputs. +Two pool configurations are certified, selected by `quant_mode`: + +**High-precision pool (`quant_mode = 0`).** The cache element type is +`out_dtype` and the gather is a bitwise copy. + +**fp8-e4m3 pool (`quant_mode` with the fp8-KV bit, `128`).** The cache +holds one `__nv_fp8_e4m3` byte per element and the gather dequantizes. +With `r = kv_scale_quant_orig[0]` (`1.0` when that tensor is omitted): + +``` +compressed_kv[t, :] = out_dtype( float(cache_row(t)[:C]) * r ) +k_pe[t, :] = out_dtype( float(cache_row(t)[C:]) * r ) +``` + +One fp32 multiply per element, by the exact fp32 number in +`kv_scale_quant_orig[0]`, rounded once to `out_dtype` — not a multiply in +`out_dtype`, which differs on 24.7% of the elements when `r` is not exactly +representable there (measured, `r = 1/1.5`, bf16). The byte→float step is +exact for all 256 e4m3 byte values, so at `r = 1` the gather is a lossless +widening. There is no per-block scale anywhere: the fp8 latent pool stores +nothing but e4m3 elements, and `r` is a single static number the caller +supplies. + +Fusion boundary: only this cache read happens inside the call. The caller +still owns writing the latent rows in the first place (the cached prefix +from earlier steps plus this step's new tokens — a sibling +context-preprocessing op appends the latter, quantizing them on the way in +under its own **write-side** scale), the `kv_b_proj` up-projection of +`compressed_kv`, and the subsequent context attention. **This call applies +`kv_scale_quant_orig` unconditionally and checks nothing about how the pool +came to hold what it holds** — see Preconditions. + +## Signature + +```python +def load_paged_kv_cache_for_mla( + out_dtype: torch.dtype, + num_contexts: int, + num_ctx_kv_tokens: int, + max_ctx_kv_len: int, + cu_ctx_kv_lens: torch.Tensor, + kv_cache_block_offsets: torch.Tensor, + host_kv_cache_pool_pointers: torch.Tensor, + host_kv_cache_pool_mapping: torch.Tensor, + kv_scale_quant_orig: Optional[torch.Tensor], + layer_idx: int, + kv_lora_rank: int, + qk_rope_head_dim: int, + tokens_per_block: int, + attention_window_size: int, + beam_width: int, + quant_mode: int, +) -> Tuple[torch.Tensor, torch.Tensor] +``` + +With `S` = total sequences in the batch (contexts first), `C` = +`kv_lora_rank`, `R` = `qk_rope_head_dim`, `T` = `num_ctx_kv_tokens`: + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `out_dtype` | — | the **output** element type, not the cache's: `torch.bfloat16` / `torch.float16` / `torch.float32`. On the fp8 path all three are certified over the same e4m3 pool; on the high-precision path it also fixes the cache's element width (see Preconditions) and bf16 / fp16 are certified there | — | — | +| `num_contexts` | leading context-phase sequence count | Python int | — | — | +| `num_ctx_kv_tokens` | `T = sum(L_s)` over context seqs, cached **plus** new tokens (sizes the outputs) | Python int | — | — | +| `max_ctx_kv_len` | `max(L_s)` over context seqs | Python int | — | — | +| `cu_ctx_kv_lens` | `[>= num_contexts+1]`; `[0, cumsum(L_s)]`; only the first `num_contexts+1` entries are read | int64 | contiguous | CUDA | +| `kv_cache_block_offsets` | `[1, >= S, 2, max_blocks_per_seq]`; raw block ids per seq (K and V rows identical, kv_factor=1 pool); contexts first | int32 | contiguous | CUDA | +| `host_kv_cache_pool_pointers` | `[num_pools, 2]`: (primary ptr, secondary ptr=0) | int64 | contiguous | CPU | +| `host_kv_cache_pool_mapping` | `[num_layers, 2]`: (pool index, layer-within-pool) row per layer | int32 | contiguous | CPU | +| `kv_scale_quant_orig` | the **read-side** dequantization scale; element `[0]` is the only one read. `None` behaves exactly as `1.0` and is what the engine's own call site passes. Read only when `quant_mode` has the fp8-KV bit — silently ignored otherwise | fp32 (others raise) | any 1-element-or-larger shape; a 0-dim scalar also works | CUDA (a CPU tensor is also accepted — see Notes) | +| `layer_idx` | row into `host_kv_cache_pool_mapping` | Python int | — | — | +| `kv_lora_rank` | `C` (512 certified) | Python int | — | — | +| `qk_rope_head_dim` | `R` (64 certified) | Python int | — | — | +| `tokens_per_block` | pool page size; 32 and 64 certified on both paths (32 is what a default `KvCacheConfig` produces) | Python int | — | — | +| `attention_window_size` | `>= max(L_s)`; production passes the manager's `max_seq_len` (smaller values imply cyclic-cache addressing, not certified) | Python int | — | — | +| `beam_width` | `1` | Python int | — | — | +| `quant_mode` | `0` = high-precision (bf16/fp16) KV cache; `128` (`QuantMode.FP8_KV_CACHE`) = e4m3 KV cache. Both certified. Other quantization bits set **alongside** `128` are not read — `1152` (`FP8_KV_CACHE \| FP8_1x128_128x128`, the combination trtllm's auto-deploy MLA path builds for an fp8 latent cache) and `384` (`\| FP8_QDQ`) return bit-identical results. `64` and `8192` are rejected (see Notes) | Python int | — | — | +| returns `compressed_kv` | `[T, C]`, freshly allocated | `out_dtype` | contiguous | CUDA | +| returns `k_pe` | `[T, R]`, freshly allocated | `out_dtype` | contiguous | CUDA | + +## Metadata consumed + +No thread-local or registered-layer state: all inputs are explicit +arguments. The addressing tensors are exactly what a +`TrtllmAttentionMetadata` prepared with +`enable_context_mla_with_cached_kv=True` and its `KVCacheManager` expose: +`ctx_kv_indptr` → `cu_ctx_kv_lens`, +`num_ctx_cached_tokens + num_ctx_tokens` → `num_ctx_kv_tokens`, +`max_ctx_kv_len` → `max_ctx_kv_len`, `metadata.kv_cache_block_offsets`, +`manager.kv_cache_pool_pointers`, `manager.kv_cache_pool_mapping`, +`manager.tokens_per_block`, `manager.max_seq_len` → +`attention_window_size`. Nothing about `kv_scale_quant_orig` comes from +metadata; it is the caller's number. + +## Preconditions + +- The paged pool addressed by `host_kv_cache_pool_pointers` + + `host_kv_cache_pool_mapping[layer_idx]` is an MLA latent cache: + kv_factor 1, one kv head, row width `C + R`, page size + `tokens_per_block`. Per page the layout is `[tokens_per_block, C + R]` + (page base = pool base + `block_id * tokens_per_block * (C + R)` + elements). +- **`quant_mode` alone fixes the pool's element width**, because the pool + arrives as a raw pointer and nothing else describes it: one byte per + element when the fp8-KV bit (`128`) is set, `sizeof(out_dtype)` when it + is not. So the caller owes a matching pool — e4m3 under `quant_mode=128`, + `out_dtype` under `quant_mode=0`. A mismatch in either direction is + **silent** (see Notes). +- `out_dtype` is one of fp16 / fp32 / bf16; anything else raises + `RuntimeError: out_dtype only support float16, float32, bfloat16` (tested + with e4m3, int8 and fp64). The output tensors are allocated before that + check, but no kernel runs. Note that under `quant_mode=128` `out_dtype` + is necessarily **not** the cache dtype: the cache is e4m3 and e4m3 is not + an accepted `out_dtype`. +- Every position `[0, L_s)` of every context sequence has already been + written to the pool (cached prefix and this step's new tokens alike). + The op reads raw pool memory, so slots never written simply propagate + whatever they currently hold. +- `kv_cache_block_offsets` covers `ceil(L_s / tokens_per_block)` pages for + every context sequence. +- `cu_ctx_kv_lens` is int64 (the kernel rejects other index dtypes), on + device, starts at 0, and is consistent with `num_ctx_kv_tokens` + (`cu_ctx_kv_lens[num_contexts] == num_ctx_kv_tokens`) and + `max_ctx_kv_len` (`>=` every `L_s`). +- `num_contexts > 0`, `num_ctx_kv_tokens > 0` and `max_ctx_kv_len > 0`: all + three are checked host-side and raise `RuntimeError` before any device + work. Production cannot reach a violation — the MLA context path runs only + under `num_contexts > 0`, and the cached-KV branch additionally requires + `num_ctx_cached_tokens > 0` — but a caller driving the op directly owns + them. +- `L_s <= attention_window_size` for every context sequence. +- `beam_width == 1`. + +**fp8 pool only (`quant_mode` fp8-KV bit set):** + +- `kv_scale_quant_orig` is an fp32 tensor whose element `[0]` is the + dequantization scale, or `None` for `1.0`. A `float16` or `float64` + tensor raises `RuntimeError: expected scalar type Float but found + Half`/`Double`; nothing else about it is validated (see Notes). +- **`kv_scale_quant_orig` must be the reciprocal of the scale the pool was + written at, and nothing anywhere checks that.** The write-side scale is + an argument of a different op (`kv_scale_orig_quant` on the sibling that + appends the new context rows) and this call never sees it. A pool written + at `w` and read at `r` yields exactly `float(e4m3(row * w)) * r`, so any + `r != 1/w` is a silently mis-scaled prefill with no error at any layer — + measured at `(w, r)` = (2, 2), (4, 1), (1, 3), each landing exactly + `w * r` away from the original rows. Production for + DeepSeek-R1-0528-FP4 is `w = r = 1.0`: all 124 per-layer `k_scale` / + `v_scale` tensors in that checkpoint are 0-dim fp32 scalars equal to 1.0 + (read from its safetensors shards), and the engine's own call sites + hard-code `None` on both the write and the read side. +- `out_dtype` must have the range to hold `max|e4m3 element| * r`. The + product is **not** clamped to `out_dtype`: with an e4m3 `±448` in the + pool and `r = 256`, fp16 output is `±inf` while bf16 and fp32 hold + `±114688` (measured). A NaN byte in the pool (`0x7f` / `0xff`) becomes a + NaN output. +- The pool must actually hold e4m3 bytes at `C + R` bytes per row. The op + reads them as raw bytes; no header, no per-block scale, no metadata. + +A caller violating none of these gets a correct result. + +## Notes + +- Schema argument names disagree with certified usage: the schema calls the + third and fourth arguments `num_ctx_cached_tokens` and + `max_ctx_cached_kv_len`, but the production call path passes cached + **plus** new token totals (`num_ctx_cached_tokens + num_ctx_tokens`, + `max_ctx_kv_len`). The totals are what size the outputs and bound the + read range; the wrapper parameter names reflect that. +- Precision, high-precision pool (sm_100): bf16→bf16 and fp16→fp16 gathers + are bitwise (`rtol=0, atol=0` holds), including rows crossing page + boundaries and pool-mapping rows other than 0 (`layer_idx=1` of a + two-layer pool certified). +- Precision, fp8 pool (sm_100): the gather is **bit-exact** against + `(float(e4m3 byte) * r)` rounded once to `out_dtype`, over + `r ∈ {omitted, 1.0, 1/1.5, 0.5, 0.25, 2.0}` × `out_dtype ∈ {bf16, fp16, + fp32}`, and separately over pool contents written at `w ∈ {0.25, 0.5, + 1.0, 1.5, 2.0, 4.0}`. Both output halves are covered by the same gate, so + the `C`/`R` split of the 1-byte row is pinned at byte `C`. +- fp8 domain (sm_100): all 256 e4m3 byte values — `±448`, `±0`, the + subnormals down to `2^-9`, and both NaN bytes — come back **exactly** at + `r = 1.0` in fp16, bf16 and fp32. The byte→float step neither saturates + nor flushes to zero; it is a plain conversion, and only the multiply that + follows can leave `out_dtype`'s range (see Preconditions). The *write* + side of the round trip behaves differently — the sibling append op's own + contract records that it clamps to `±448` where + `torch.Tensor.to(float8_e4m3fn)` yields the NaN byte — so a torch mirror + of the whole round trip is only valid while every `|value * w| <= 448`. + Every case certified here stays inside that. +- Page size (sm_100): certified at `tokens_per_block` 32 and 64 on both + paths. 32 is what a `KvCacheConfig` at its defaults hands the op; 64 is + what a tuned MLA target opts into. Both were run through the same bitwise + gate on real `KVCacheManager` state, over gathered ranges that end on a + page boundary and mid-page alike — high-precision pool, at 32: 37 / 96 + (three exact pages) / 400 (twelve pages plus sixteen rows) / 512 (sixteen + exact pages) in bf16 through pool-mapping row 1 of a two-layer pool, a + mixed batch with trailing generation sequences (128 = four exact pages, + and 63), and 64 / 608 / 7 in fp16; the deepest sequence walked nineteen + offsets-row entries. fp8 pool, at 32: gathered ranges of 37 / 96 / 400 / + 512 (1045 rows in one call) plus a mixed batch (128 and 63); at 64: + 1024 / 1024 / 135 (2183 rows in one call). The gather stayed bitwise in + every case, so + `page = t // tokens_per_block`, `slot = t % tokens_per_block` holds at + both sizes and both element widths. Other page sizes are untested rather + than known-bad. +- The outputs are new allocations sized `[T, C]`/`[T, R]`; a `T == 0` batch + is **rejected** by the op rather than merely unexercised (see + Preconditions), and production never reaches it anyway — the path is + guarded behind `num_ctx_cached_tokens > 0`. +- The pool is not written on either path: a whole-pool byte snapshot around + a 1045-row fp8 gather shows zero changed bytes, and a sentinel row placed + one slot past each context sequence's range never appears in the output. + Two identical calls return bitwise-identical tensors. +- A cache element-type or element-width mismatch (sm_100) is **silent in + every form measured**. The two mismatches an fp8 pool makes reachable — + reading it two bytes wide, and reading a two-byte pool one byte wide — + are gated in this entry's test; the other observations below are from + scratch probes. + - *Same width, wrong type* (bf16 pool read as fp16, or fp16 pool read as + bf16, `quant_mode=0` both times): addressing is unchanged and the cache + bits are copied out verbatim, so the output is bit-identical to the + correct gather and numerically wrong — max abs diff 2.0 one way and 512 + the other, nothing raised, nothing out of bounds. + - *Read wider than the pool* (fp32 `out_dtype` over a 2-byte pool, or + `quant_mode=0` over an e4m3 pool): the stride doubles, so token `t` of + block `b` is read from pool row + `2 * (b * tokens_per_block + t % tokens_per_block)` and consumes two + rows. Verified bit for bit both ways — 565 gathered rows in the + fp32-over-bf16 case, of which 310 came back holding a sentinel written + nowhere near that sequence, and the exact doubled row indices in the + e4m3 case (block 0 token 1 → row 2, block 1 token 0 → row 64). Past + half the pool it leaves the allocation entirely: at block 1584 of a + 2048-block pool compute-sanitizer reported invalid 16-byte `__global__` + reads 77.9 MB beyond the nearest allocation, while the call still + returned without raising. + - *Read narrower than the pool* (`quant_mode=128` over a bf16 pool): the + 1-byte slab for block `b` starts at byte + `b * tokens_per_block * (C + R)`, i.e. inside bf16 block `b // 2`, and + each token consumes `C + R` bytes. Verified bit for bit at blocks 0 and + 1. This direction stays inside the allocation — the read walks half the + bytes the pool holds — but every value is garbage, including NaNs + wherever a bf16 byte happens to be `0x7f`/`0xff` (107 of 36864 elements + in one scratch probe). +- `kv_scale_quant_orig` is read **only** when `quant_mode` has the fp8-KV + bit. With `quant_mode = 0` a scale of 4.0 produces byte-identical output + to `None`: a silently ignored argument, and the shape a caller who set up + the scale but forgot the `quant_mode` bit lands in. +- `kv_scale_quant_orig` validation is otherwise thin. Only the dtype is + checked, and loudly. A `[2]` fp32 CUDA tensor is accepted and only `[0]` + is used, and a 0-dim fp32 scalar tensor is accepted — both tested. A + **CPU** fp32 tensor is accepted too: the pointer is handed to the kernel + unchecked, and in a scratch probe on this machine it produced results + bit-identical to the equivalent CUDA tensor at two distinct values. That + works because this device reports + `CU_DEVICE_ATTRIBUTE_PAGEABLE_MEMORY_ACCESS = 1` (driver 595.58.03), so + it is a property of the machine rather than of the op — not something to + rely on, and left out of the test for that reason. Device placement is + not validated either way, so a wrong-device scale will not announce + itself. +- Rejected `quant_mode` values, both tested: + `64` (`QuantMode.INT8_KV_CACHE`) raises + `RuntimeError: [TensorRT-LLM][ERROR] Assertion failed: Only FP8 KV cache + is supported for now (../tensorrt_llm/thop/mlaPreprocessOp.cpp:134)`; + `8192` (`NVFP4_KV_CACHE`) fails earlier and for a different reason — + `Expected hostKvCachePoolPointers.dim() == 3`, the nvfp4 pool-pointer + layout, checked before the quant_mode branch — so it is not the same + guard. What the guard rejects is a *KV-cache* quantization bit other than + fp8; bits outside that group ride along unread, which is why `1152` + (`FP8_KV_CACHE | FP8_1x128_128x128`, the combination trtllm's + auto-deploy MLA path builds for an fp8 latent cache) and `384` + (`| FP8_QDQ`) are bit-identical to plain `128` (tested). +- The engine's own MLA backend passes `None` for `kv_scale_quant_orig` at + this call site while still passing its real `quant_mode`, so an fp8 KV + cache is dequantized at 1.0 there regardless of any calibrated per-layer + scale. Harmless for a checkpoint whose scales are 1.0; the `r != 1` + behaviour certified here is reachable only from a direct op caller. +- Sibling ops exist for adjacent MLA context-phase roles: + `trtllm::mla_rope_append_paged_kv_assign_q` (RoPE + append of the new + context tokens that this op then reads back, and the owner of the + write-side `kv_scale_orig_quant`), + `trtllm::load_chunked_kv_cache_for_mla` (slice-wise variant for chunked + prefill, with its own `kv_scale_quant_orig`), + `trtllm::merge_chunked_attention_for_mla`, and the context FMHA that + consumes the up-projected K/V. diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla.py b/tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla.py new file mode 100644 index 000000000000..2855bf2d1667 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla.py @@ -0,0 +1,50 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Gather each context sequence's full MLA latent KV from the paged cache.""" + +from typing import Optional, Tuple + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + + +def load_paged_kv_cache_for_mla( + out_dtype: torch.dtype, + num_contexts: int, + num_ctx_kv_tokens: int, + max_ctx_kv_len: int, + cu_ctx_kv_lens: torch.Tensor, + kv_cache_block_offsets: torch.Tensor, + host_kv_cache_pool_pointers: torch.Tensor, + host_kv_cache_pool_mapping: torch.Tensor, + kv_scale_quant_orig: Optional[torch.Tensor], + layer_idx: int, + kv_lora_rank: int, + qk_rope_head_dim: int, + tokens_per_block: int, + attention_window_size: int, + beam_width: int, + quant_mode: int, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Copy positions [0, L_s) of every context sequence out of the paged MLA + latent cache into two new contiguous tensors (compressed_kv, k_pe).""" + compressed_kv, k_pe = torch.ops.trtllm.load_paged_kv_cache_for_mla( + out_dtype, + num_contexts, + num_ctx_kv_tokens, + max_ctx_kv_len, + cu_ctx_kv_lens, + kv_cache_block_offsets, + host_kv_cache_pool_pointers, + host_kv_cache_pool_mapping, + kv_scale_quant_orig, + layer_idx, + kv_lora_rank, + qk_rope_head_dim, + tokens_per_block, + attention_window_size, + beam_width, + quant_mode, + ) + return compressed_kv, k_pe diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla_test.py b/tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla_test.py new file mode 100644 index 000000000000..26405494c7bf --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla_test.py @@ -0,0 +1,1003 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the load_paged_kv_cache_for_mla catalog entry. + +The op reads paged-KV-cache addressing tensors that the runtime normally +derives from a KVCacheManager and a prepared TrtllmAttentionMetadata. The +test builds that state for real — an actual MLA (SELFKONLY, kv_factor=1) +KVCacheManager and a TrtllmAttentionMetadata prepared with +enable_context_mla_with_cached_kv=True — writes known random latent rows +directly into the paged pool at every position of every context sequence, +then checks that the op gathers them back into two contiguous tensors +(compressed_kv, k_pe), in batch order, ignoring trailing generation +sequences. + +Two configurations are covered: + +1. Matching-dtype pool (`quant_mode=0`, `kv_scale_quant_orig=None`), where + the gather is a bitwise copy. Covered at both pool page sizes: 64, and + 32 (the engine default), where every sequence spans twice as many pages + and the exact-fill boundaries fall elsewhere. +2. fp8-e4m3 latent pool (`quant_mode=128`) at the DeepSeek-R1-0528 cell — + `kv_lora_rank=512`, `qk_rope_head_dim=64`, `tokens_per_block=32`, + `beam_width=1` — where the pool holds one byte per element and the + gather dequantizes to bf16 / fp16 / fp32 by multiplying with + `kv_scale_quant_orig`. The pool is filled with `e4m3(row * w)` for an + explicit write-side scale `w` (the write-side op's certified formula) so + that a deliberately inconsistent `(w, r)` pair can be driven through the + read side. +""" + +from typing import List, Optional, Tuple + +import torch + +from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata +from tensorrt_llm._torch.metadata import KVCacheParams +from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager +from tensorrt_llm.bindings import DataType +from tensorrt_llm.bindings.internal.batch_manager import CacheType +from tensorrt_llm.llmapi.llm_args import KvCacheConfig +from tensorrt_llm.mapping import Mapping + +from .load_paged_kv_cache_for_mla import load_paged_kv_cache_for_mla + +assert torch.cuda.is_available(), "load_paged_kv_cache_for_mla requires a CUDA device" + +# DeepSeek-V3 MLA latent geometry. +KV_LORA_RANK = 512 +QK_ROPE_HEAD_DIM = 64 +HEAD_SIZE = KV_LORA_RANK + QK_ROPE_HEAD_DIM +# Pool page size. 32 is the engine default (KvCacheConfig.tokens_per_block), +# 64 the value a tuned MLA target opts into; both are covered below. +TOKENS_PER_BLOCK = 64 +PAGE32 = 32 +MAX_SEQ_LEN = 1024 + +# quant_mode bits (tensorrt_llm.quantization.mode.QuantMode). +QM_INT8_KV_CACHE = 64 +QM_FP8_KV_CACHE = 128 +QM_FP8_QDQ = 256 +QM_FP8_1X128_128X128 = 1024 +QM_NVFP4_KV_CACHE = 8192 + +# e4m3 constants used to derive tolerances. The format carries 3 explicit +# mantissa bits, so a round-to-nearest quantization is within half an ulp = +# 2**-4 relative; the smallest subnormal is 2**-9, so the absolute error near +# zero is at most 2**-10. +E4M3_HALF_ULP_REL = 2.0**-4 +E4M3_MIN_SUBNORMAL = 2.0**-9 + +_TRTLLM_TO_TORCH_DTYPE = { + DataType.BF16: torch.bfloat16, + DataType.HALF: torch.float16, +} + + +class _MlaCacheEnv: + """Real op state: MLA (kv_factor=1) paged KV cache manager. + + With fp8_pool the manager allocates an e4m3 latent pool (one byte per + element) while the gather's out_dtype stays a high-precision type. + """ + + def __init__( + self, + cache_dtype: DataType, + num_layers: int = 1, + max_batch_size: int = 8, + tokens_per_block: int = TOKENS_PER_BLOCK, + fp8_pool: bool = False, + max_tokens: int = 131072, + ) -> None: + self.torch_dtype = _TRTLLM_TO_TORCH_DTYPE[cache_dtype] + self.max_batch_size = max_batch_size + self.tokens_per_block = tokens_per_block + self.fp8_pool = fp8_pool + self.pool_dtype = torch.float8_e4m3fn if fp8_pool else self.torch_dtype + self.quant_mode = QM_FP8_KV_CACHE if fp8_pool else 0 + self.kv_cache_manager = KVCacheManager( + KvCacheConfig(max_tokens=max_tokens, enable_block_reuse=False), + CacheType.SELFKONLY, # MLA latent cache: kv_factor=1, one kv head + num_layers=num_layers, + num_kv_heads=1, + head_dim=HEAD_SIZE, + tokens_per_block=tokens_per_block, + max_seq_len=MAX_SEQ_LEN, + max_batch_size=max_batch_size, + mapping=Mapping(world_size=1, tp_size=1, rank=0), + dtype=DataType.FP8 if fp8_pool else cache_dtype, + ) + # The pool must really be paged at the size and element type the case + # claims: the op derives both from tokens_per_block and quant_mode, + # not from anything the manager tells it. + assert self.kv_cache_manager.tokens_per_block == tokens_per_block + pool = self.kv_cache_manager.get_buffers(0) + assert pool is not None and pool.dtype == self.pool_dtype + + def prepare_metadata( + self, + request_ids: List[int], + seq_lens: List[int], + num_contexts: int, + cached_lens: List[int], + ) -> TrtllmAttentionMetadata: + metadata = TrtllmAttentionMetadata( + max_num_requests=self.max_batch_size, + max_num_tokens=8192, + kv_cache_manager=self.kv_cache_manager, + enable_context_mla_with_cached_kv=True, + ) + metadata.seq_lens = torch.tensor(seq_lens, dtype=torch.int) + metadata.num_contexts = num_contexts + metadata.request_ids = request_ids + metadata.prompt_lens = [c + s for c, s in zip(cached_lens, seq_lens)] + metadata.kv_cache_params = KVCacheParams( + use_cache=True, + num_cached_tokens_per_seq=cached_lens, + ) + metadata.prepare() + return metadata + + def pool_tensor(self, layer_idx: int) -> torch.Tensor: + pool = self.kv_cache_manager.get_buffers(layer_idx) + assert pool is not None + return pool + + def blocks(self, request_id: int) -> List[int]: + return list(self.kv_cache_manager.get_batch_cache_indices([request_id])[0]) + + def fill_layer( + self, + layer_idx: int, + request_ids: List[int], + kv_lens: List[int], + write_scale: float = 1.0, + ) -> Tuple[List[torch.Tensor], List[torch.Tensor]]: + """Write known random latent rows at positions [0, L_s) of each + sequence's pages. + + Returns (stored, source): `stored` is what physically sits in the + pool, `source` the unquantized fp32 rows it came from. On the fp8 + pool `stored = e4m3(source * write_scale)` — the write-side op's + certified formula, mirrored in torch. Every drawn value stays well + inside e4m3's +-448 range, which is where the torch cast and the + kernel's saturating cast agree. + """ + pool = self.pool_tensor(layer_idx) # [pages,1,tpb,1,head] + tpb = self.tokens_per_block + block_ids = self.kv_cache_manager.get_batch_cache_indices(request_ids) + stored_per_seq, source_per_seq = [], [] + for blocks, length in zip(block_ids, kv_lens): + if self.fp8_pool: + source = torch.randn(length, HEAD_SIZE, dtype=torch.float32, device="cuda") + stored = (source * write_scale).to(torch.float8_e4m3fn) + else: + stored = torch.randn(length, HEAD_SIZE, dtype=self.torch_dtype, device="cuda") + source = stored.float() + for start in range(0, length, tpb): + page = blocks[start // tpb] + n = min(tpb, length - start) + pool[page, 0, :n, 0].copy_(stored[start : start + n]) + stored_per_seq.append(stored) + source_per_seq.append(source) + return stored_per_seq, source_per_seq + + def shutdown(self) -> None: + self.kv_cache_manager.shutdown() + + +def _scale_tensor(value: Optional[float]) -> Optional[torch.Tensor]: + if value is None: + return None + return torch.full((1,), value, dtype=torch.float32, device="cuda") + + +def _call( + env: _MlaCacheEnv, + metadata: TrtllmAttentionMetadata, + num_contexts: int, + out_dtype: torch.dtype, + kv_scale_quant_orig: Optional[torch.Tensor], + layer_idx: int = 0, + quant_mode: Optional[int] = None, +) -> Tuple[torch.Tensor, torch.Tensor]: + """One wrapper call driven straight off prepared metadata.""" + total_ctx_kv = int(metadata.num_ctx_cached_tokens + metadata.num_ctx_tokens) + block_offsets = metadata.kv_cache_block_offsets + pool_pointers = env.kv_cache_manager.kv_cache_pool_pointers + pool_mapping = env.kv_cache_manager.kv_cache_pool_mapping + assert block_offsets is not None + assert pool_pointers is not None and pool_mapping is not None + compressed_kv, k_pe = load_paged_kv_cache_for_mla( + out_dtype, + num_contexts, + total_ctx_kv, + int(metadata.max_ctx_kv_len), + metadata.ctx_kv_indptr, + block_offsets, + pool_pointers, + pool_mapping, + kv_scale_quant_orig, + layer_idx, + KV_LORA_RANK, + QK_ROPE_HEAD_DIM, + env.tokens_per_block, + MAX_SEQ_LEN, # attention_window_size + 1, # beam_width + env.quant_mode if quant_mode is None else quant_mode, + ) + torch.cuda.synchronize() + return compressed_kv, k_pe + + +def _check_outputs( + compressed_kv: torch.Tensor, + k_pe: torch.Tensor, + total_ctx_kv: int, + out_dtype: torch.dtype, +) -> None: + assert compressed_kv.shape == (total_ctx_kv, KV_LORA_RANK), compressed_kv.shape + assert k_pe.shape == (total_ctx_kv, QK_ROPE_HEAD_DIM), k_pe.shape + assert compressed_kv.dtype == out_dtype and k_pe.dtype == out_dtype + assert compressed_kv.is_cuda and k_pe.is_cuda + assert compressed_kv.is_contiguous() and k_pe.is_contiguous() + + +def _run_and_check( + env: _MlaCacheEnv, + request_ids: List[int], + seq_lens: List[int], + num_contexts: int, + cached_lens: List[int], + layer_idx: int = 0, +) -> None: + """Fill the context sequences' cache rows, run the op, and verify the + gathered compressed_kv / k_pe against the written rows.""" + kv_lens = [c + s for c, s in zip(cached_lens, seq_lens)] + metadata = env.prepare_metadata(request_ids, seq_lens, num_contexts, cached_lens) + ctx_kv_lens = kv_lens[:num_contexts] + stored, _ = env.fill_layer(layer_idx, request_ids[:num_contexts], ctx_kv_lens) + + total_ctx_kv = int(metadata.num_ctx_cached_tokens + metadata.num_ctx_tokens) + assert total_ctx_kv == sum(ctx_kv_lens) + assert int(metadata.max_ctx_kv_len) == max(ctx_kv_lens) + + compressed_kv, k_pe = _call( + env, + metadata, + num_contexts, + env.torch_dtype, + None, # kv_scale_quant_orig: fp8-KV-cache path only + layer_idx=layer_idx, + ) + + expected = torch.cat(stored, dim=0) # [total_ctx_kv, HEAD_SIZE] + _check_outputs(compressed_kv, k_pe, total_ctx_kv, env.torch_dtype) + # Dtype-preserving gather: outputs must be bitwise equal to the cache rows. + torch.testing.assert_close(compressed_kv, expected[:, :KV_LORA_RANK], rtol=0.0, atol=0.0) + torch.testing.assert_close(k_pe, expected[:, KV_LORA_RANK:], rtol=0.0, atol=0.0) + + +def _dequant_mirror( + stored: torch.Tensor, + read_scale: Optional[torch.Tensor], + out_dtype: torch.dtype, +) -> torch.Tensor: + """float(e4m3 byte) * kv_scale_quant_orig[0] in fp32, rounded once to + out_dtype — the fp8 gather's reference, built from torch alone.""" + wide = stored.float() + if read_scale is not None: + wide = wide * read_scale[0] + return wide.to(out_dtype) + + +def _run_fp8( + env: _MlaCacheEnv, + request_ids: List[int], + seq_lens: List[int], + num_contexts: int, + cached_lens: List[int], + out_dtype: torch.dtype, + read_scale: Optional[float], + write_scale: float = 1.0, + layer_idx: int = 0, + extra_slots: int = 0, + add_requests: bool = True, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: + """Fill an fp8 pool with e4m3(row * write_scale) and gather it back. + + Returns (compressed_kv, k_pe, stored, source, read_scale_tensor) with + `stored` / `source` concatenated over the context sequences. + """ + assert env.fp8_pool + kv_lens = [c + s for c, s in zip(cached_lens, seq_lens)] + if add_requests: + env.kv_cache_manager.add_dummy_requests( + request_ids, token_nums=[n + extra_slots for n in kv_lens] + ) + metadata = env.prepare_metadata(request_ids, seq_lens, num_contexts, cached_lens) + ctx_kv_lens = kv_lens[:num_contexts] + stored, source = env.fill_layer(layer_idx, request_ids[:num_contexts], ctx_kv_lens, write_scale) + total_ctx_kv = int(metadata.num_ctx_cached_tokens + metadata.num_ctx_tokens) + assert total_ctx_kv == sum(ctx_kv_lens) + scale_t = _scale_tensor(read_scale) + compressed_kv, k_pe = _call( + env, metadata, num_contexts, out_dtype, scale_t, layer_idx=layer_idx + ) + _check_outputs(compressed_kv, k_pe, total_ctx_kv, out_dtype) + return ( + compressed_kv, + k_pe, + torch.cat(stored, dim=0), + torch.cat(source, dim=0), + scale_t, + ) + + +def _assert_fp8_exact( + compressed_kv: torch.Tensor, + k_pe: torch.Tensor, + stored: torch.Tensor, + read_scale: Optional[torch.Tensor], + out_dtype: torch.dtype, +) -> None: + """Both halves bitwise equal to the fp32-multiply mirror.""" + mirror = _dequant_mirror(stored, read_scale, out_dtype) + torch.testing.assert_close(compressed_kv, mirror[:, :KV_LORA_RANK], rtol=0.0, atol=0.0) + torch.testing.assert_close(k_pe, mirror[:, KV_LORA_RANK:], rtol=0.0, atol=0.0) + + +def _raises(fn) -> str: + try: + fn() + except RuntimeError as exc: + return str(exc) + raise AssertionError("expected the call to raise RuntimeError") + + +# ───────────────────────── matching-dtype pool ────────────────────────── + + +def test_bf16_mixed_lengths_layer1() -> None: + """Prefill-like batch (~1k gathered tokens): cached lengths of 0, exactly + one block, and multi-block; addressed through pool-mapping row 1 of a + two-layer pool.""" + torch.manual_seed(0) + env = _MlaCacheEnv(DataType.BF16, num_layers=2) + try: + cached = [0, 64, 100, 511] + new = [37, 64, 300, 1] + rids = [0, 1, 2, 3] + env.kv_cache_manager.add_dummy_requests( + rids, token_nums=[c + n for c, n in zip(cached, new)] + ) + _run_and_check( + env, + request_ids=rids, + seq_lens=new, + num_contexts=4, + cached_lens=cached, + layer_idx=1, + ) + finally: + env.shutdown() + + +def test_bf16_single_tiny_sequence() -> None: + """Smallest cached-context case: one sequence, one cached token plus one + new token (two gathered rows).""" + torch.manual_seed(1) + env = _MlaCacheEnv(DataType.BF16) + try: + env.kv_cache_manager.add_dummy_requests([0], token_nums=[2]) + _run_and_check( + env, + request_ids=[0], + seq_lens=[1], + num_contexts=1, + cached_lens=[1], + ) + finally: + env.shutdown() + + +def test_bf16_trailing_generation_seqs_ignored() -> None: + """Mixed batch: two context sequences followed by two generation + sequences. The op must gather exactly the context sequences' rows and + index per-seq tensors over [0, num_contexts) only.""" + torch.manual_seed(2) + env = _MlaCacheEnv(DataType.BF16) + try: + cached = [128, 3, 200, 77] + new = [40, 60, 1, 1] + rids = [0, 1, 2, 3] + env.kv_cache_manager.add_dummy_requests( + rids, token_nums=[c + n for c, n in zip(cached, new)] + ) + _run_and_check( + env, + request_ids=rids, + seq_lens=new, + num_contexts=2, + cached_lens=cached, + ) + finally: + env.shutdown() + + +def test_fp16_mixed_lengths() -> None: + """fp16 latent cache with fp16 out_dtype over block-crossing lengths.""" + torch.manual_seed(3) + env = _MlaCacheEnv(DataType.HALF) + try: + cached = [65, 640, 1] + new = [63, 128, 6] + rids = [0, 1, 2] + env.kv_cache_manager.add_dummy_requests( + rids, token_nums=[c + n for c, n in zip(cached, new)] + ) + _run_and_check( + env, + request_ids=rids, + seq_lens=new, + num_contexts=3, + cached_lens=cached, + ) + finally: + env.shutdown() + + +def test_bf16_page32_mixed_lengths_layer1() -> None: + """Page size 32 (the engine default). Gathered ranges of 37, 96, 400 and + 512 rows: 96 and 512 fill 3 and 16 pages exactly at 32 (neither is a + page multiple at 64), 400 spans 12 full pages plus 16 rows, 37 spans two + — every sequence crosses at least one boundary and the deepest one walks + 16 offsets-row entries. Addressed through pool-mapping row 1 of a + two-layer pool.""" + torch.manual_seed(4) + env = _MlaCacheEnv(DataType.BF16, num_layers=2, tokens_per_block=PAGE32) + try: + cached = [0, 32, 100, 511] + new = [37, 64, 300, 1] + rids = [0, 1, 2, 3] + env.kv_cache_manager.add_dummy_requests( + rids, token_nums=[c + n for c, n in zip(cached, new)] + ) + _run_and_check( + env, + request_ids=rids, + seq_lens=new, + num_contexts=4, + cached_lens=cached, + layer_idx=1, + ) + finally: + env.shutdown() + + +def test_bf16_page32_trailing_generation_seqs_ignored() -> None: + """Page size 32, mixed batch: two context sequences followed by two + generation sequences. The gathered lengths (128 = 4 exact pages, 63 = + two pages minus one row) sit either side of a page boundary, and the + ignored generation sequences own pages of their own.""" + torch.manual_seed(5) + env = _MlaCacheEnv(DataType.BF16, tokens_per_block=PAGE32) + try: + cached = [96, 32, 200, 77] + new = [32, 31, 1, 1] + rids = [0, 1, 2, 3] + env.kv_cache_manager.add_dummy_requests( + rids, token_nums=[c + n for c, n in zip(cached, new)] + ) + _run_and_check( + env, + request_ids=rids, + seq_lens=new, + num_contexts=2, + cached_lens=cached, + ) + finally: + env.shutdown() + + +def test_fp16_page32_mixed_lengths() -> None: + """fp16 latent cache at page size 32: gathered ranges of 64 (2 exact + pages), 608 (19 exact pages) and 7 (a partial first page).""" + torch.manual_seed(6) + env = _MlaCacheEnv(DataType.HALF, tokens_per_block=PAGE32) + try: + cached = [33, 512, 0] + new = [31, 96, 7] + rids = [0, 1, 2] + env.kv_cache_manager.add_dummy_requests( + rids, token_nums=[c + n for c, n in zip(cached, new)] + ) + _run_and_check( + env, + request_ids=rids, + seq_lens=new, + num_contexts=3, + cached_lens=cached, + ) + finally: + env.shutdown() + + +# ───────────────────────── fp8-e4m3 latent pool ───────────────────────── + + +def _fp8_env( + num_layers: int = 1, + tokens_per_block: int = PAGE32, + max_tokens: int = 131072, +) -> _MlaCacheEnv: + """DeepSeek-R1-0528 cell: e4m3 latent pool, page 32, C/R = 512/64.""" + return _MlaCacheEnv( + DataType.BF16, + num_layers=num_layers, + tokens_per_block=tokens_per_block, + fp8_pool=True, + max_tokens=max_tokens, + ) + + +def test_fp8_production_config() -> None: + """The production fp8 context call: quant_mode=128 with + kv_scale_quant_orig omitted (what the engine's own call site passes), a + unit-scale pool and bf16 output. Gathered ranges of 37 / 96 / 400 / 512 + rows over page 32. Also locks the three things a pure gather owes: the + pool is not modified, the slot one past each sequence's range is not + read, and a second identical call reproduces the first bitwise.""" + torch.manual_seed(40) + env = _fp8_env() + try: + cached = [0, 32, 100, 511] + new = [37, 64, 300, 1] + rids = [0, 1, 2, 3] + # One spare slot per sequence so a sentinel can sit just past L_s. + ckv, kpe, stored, _, scale_t = _run_fp8( + env, + request_ids=rids, + seq_lens=new, + num_contexts=4, + cached_lens=cached, + out_dtype=torch.bfloat16, + read_scale=None, + extra_slots=1, + ) + assert ckv.shape[0] == 1045, ckv.shape + _assert_fp8_exact(ckv, kpe, stored, scale_t, torch.bfloat16) + + # Sentinel just past each context sequence's gathered range, plus a + # whole-pool byte snapshot: the call must read neither and write + # nothing. + pool = env.pool_tensor(0) + sentinel = torch.full((HEAD_SIZE,), 240.0, dtype=torch.float32, device="cuda").to( + torch.float8_e4m3fn + ) + for rid, length in zip(rids, [c + n for c, n in zip(cached, new)]): + page = env.blocks(rid)[length // env.tokens_per_block] + pool[page, 0, length % env.tokens_per_block, 0].copy_(sentinel) + before = pool.view(torch.uint8).clone() + metadata = env.prepare_metadata(rids, new, 4, cached) + ckv2, kpe2 = _call(env, metadata, 4, torch.bfloat16, None) + assert int((before != pool.view(torch.uint8)).sum().item()) == 0 + assert not bool((ckv2.float() == 240.0).any().item()) + assert not bool((kpe2.float() == 240.0).any().item()) + # Re-running the same prepared step reproduces the gather bitwise. + torch.testing.assert_close(ckv2, ckv, rtol=0.0, atol=0.0) + torch.testing.assert_close(kpe2, kpe, rtol=0.0, atol=0.0) + finally: + env.shutdown() + + +def test_fp8_scale_and_out_dtype_matrix() -> None: + """kv_scale_quant_orig sweep x out_dtype. The gather is + `float(e4m3 byte) * r` evaluated in fp32 and rounded once to out_dtype, + bit for bit, for every combination — including r omitted, which behaves + exactly as 1.0.""" + for out_dtype in (torch.bfloat16, torch.float16, torch.float32): + torch.manual_seed(41) + env = _fp8_env() + try: + for i, read_scale in enumerate([None, 1.0, 1.0 / 1.5, 0.5, 0.25, 2.0]): + ckv, kpe, stored, _, scale_t = _run_fp8( + env, + request_ids=[i], + seq_lens=[40], + num_contexts=1, + cached_lens=[24], + out_dtype=out_dtype, + read_scale=read_scale, + ) + _assert_fp8_exact(ckv, kpe, stored, scale_t, out_dtype) + finally: + env.shutdown() + + +def test_fp8_wrong_mirrors_are_discriminated() -> None: + """The fp8 gate is bitwise (rtol=atol=0) and is never approached, so the + same comparison is run in the same process against five wrong mirrors of + the same 36864-element block. Every one of them must fail it. + + The interesting control is `multiply in out_dtype`: it differs from the + truth only by where the scale is rounded, so it is invisible to any + tolerance-based gate and only a bitwise one catches it — and only when + the scale is not exactly representable in out_dtype (at r = 1.5 or 0.25 + the two mirrors coincide exactly, measured).""" + torch.manual_seed(42) + env = _fp8_env() + try: + out_dtype = torch.bfloat16 + ckv, kpe, stored, source, scale_t = _run_fp8( + env, + request_ids=[0], + seq_lens=[40], + num_contexts=1, + cached_lens=[24], + out_dtype=out_dtype, + read_scale=1.0 / 1.5, + write_scale=1.5, + ) + assert scale_t is not None + got = torch.cat([ckv, kpe], dim=1) + correct = _dequant_mirror(stored, scale_t, out_dtype) + torch.testing.assert_close(got, correct, rtol=0.0, atol=0.0) + + n = got.numel() + wrong = { + # measured: 36843 / 36864 elements differ (99.94%) + "scale ignored": stored.float().to(out_dtype), + # measured: 36843 / 36864 (99.94%) + "reciprocal scale": (stored.float() / scale_t[0]).to(out_dtype), + # measured: 9113 / 36864 (24.72%) — the weakest control + "multiply in out_dtype": (stored.to(out_dtype) * scale_t[0].to(out_dtype)).to( + out_dtype + ), + # measured: 36358 / 36864 (98.63%) + "rows shifted by one token": torch.roll(correct, 1, 0), + # measured: 34362 / 36864 (93.21%) + "quantization skipped": (source * scale_t[0] * 1.5).to(out_dtype), + } + for name, mirror in wrong.items(): + differing = int((got != mirror).sum().item()) + # An allowance of 1% of the block: the correct mirror scores 0 and + # the weakest wrong one 24.8x this, the rest 93-100x. + assert differing > n // 100, (name, differing, n) + finally: + env.shutdown() + + +def test_fp8_inconsistent_write_read_scale_pair() -> None: + """No layer in the chain relates the write scale to the read scale. A + pool written at `w` and read at `r` comes back as + `float(e4m3(row * w)) * r` — exactly, with no error anywhere — so only + `r == 1/w` recovers the original rows and every other pair is a silently + mis-scaled prefill.""" + torch.manual_seed(43) + env = _fp8_env() + try: + # (write scale, read scale, product): three inconsistent pairs and + # one consistent one. + cases = [(2.0, 2.0, 4.0), (4.0, 1.0, 4.0), (1.0, 3.0, 3.0), (0.25, 4.0, 1.0)] + for i, (w, r, product) in enumerate(cases): + ckv, kpe, stored, source, scale_t = _run_fp8( + env, + request_ids=[i], + seq_lens=[40], + num_contexts=1, + cached_lens=[24], + out_dtype=torch.bfloat16, + read_scale=r, + write_scale=w, + ) + # Accepted, and exactly the mis-scaled value. + _assert_fp8_exact(ckv, kpe, stored, scale_t, torch.bfloat16) + got = torch.cat([ckv, kpe], dim=1).float() + # Tolerance derived from the e4m3 format: one round-to-nearest + # quantization at scale w, undone by r. + rtol = E4M3_HALF_ULP_REL + atol = 0.5 * E4M3_MIN_SUBNORMAL * r + torch.testing.assert_close(got, source * product, rtol=rtol, atol=atol) + if product != 1.0: + # ...and therefore not the original rows. + try: + torch.testing.assert_close(got, source, rtol=rtol, atol=atol) + except AssertionError: + pass + else: + raise AssertionError( + f"w={w} r={r} recovered the source rows; the op cannot be " + "applying both scales independently" + ) + finally: + env.shutdown() + + +def test_fp8_trailing_generation_seqs_ignored() -> None: + """fp8 mixed batch: two context sequences followed by two generation + sequences whose pages are filled too. Gathered lengths 128 (four exact + pages) and 63 straddle a page boundary.""" + torch.manual_seed(44) + env = _fp8_env() + try: + cached = [96, 32, 200, 77] + new = [32, 31, 1, 1] + rids = [0, 1, 2, 3] + ckv, kpe, stored, _, scale_t = _run_fp8( + env, + request_ids=rids, + seq_lens=new, + num_contexts=2, + cached_lens=cached, + out_dtype=torch.bfloat16, + read_scale=2.0, + write_scale=0.5, + ) + assert ckv.shape[0] == 191, ckv.shape + _assert_fp8_exact(ckv, kpe, stored, scale_t, torch.bfloat16) + finally: + env.shutdown() + + +def test_fp8_two_layer_pool_layer1() -> None: + """fp8 pool-mapping row 1 of a two-layer pool, with decoy rows written + at the same positions of layer 0.""" + torch.manual_seed(45) + env = _fp8_env(num_layers=2) + try: + cached = [0, 96] + new = [64, 33] + rids = [0, 1] + env.kv_cache_manager.add_dummy_requests( + rids, token_nums=[c + n for c, n in zip(cached, new)] + ) + env.fill_layer(0, rids, [c + n for c, n in zip(cached, new)], 1.0) + ckv, kpe, stored, _, scale_t = _run_fp8( + env, + request_ids=rids, + seq_lens=new, + num_contexts=2, + cached_lens=cached, + out_dtype=torch.bfloat16, + read_scale=1.0, + layer_idx=1, + add_requests=False, + ) + _assert_fp8_exact(ckv, kpe, stored, scale_t, torch.bfloat16) + # Layer 0 holds different rows at the same positions; reading it must + # return them, not layer 1's. + metadata = env.prepare_metadata(rids, new, 2, cached) + ckv0, _ = _call(env, metadata, 2, torch.bfloat16, _scale_tensor(1.0)) + assert not torch.equal(ckv0, ckv) + finally: + env.shutdown() + + +def test_fp8_page64_prefill_scale() -> None: + """fp8 at page size 64 and prefill scale: 2183 gathered rows across + three sequences (16 exact pages, 16 exact pages, three partial), fp32 + output, write scale 0.5 undone by read scale 2.0.""" + torch.manual_seed(46) + env = _fp8_env(tokens_per_block=TOKENS_PER_BLOCK) + try: + cached = [512, 0, 128] + new = [512, 1024, 7] + rids = [0, 1, 2] + ckv, kpe, stored, source, scale_t = _run_fp8( + env, + request_ids=rids, + seq_lens=new, + num_contexts=3, + cached_lens=cached, + out_dtype=torch.float32, + read_scale=2.0, + write_scale=0.5, + ) + assert ckv.shape[0] == 2183, ckv.shape + _assert_fp8_exact(ckv, kpe, stored, scale_t, torch.float32) + # The consistent pair recovers the pre-quantization rows to within + # one e4m3 rounding. + got = torch.cat([ckv, kpe], dim=1) + torch.testing.assert_close( + got, + source, + rtol=E4M3_HALF_ULP_REL, + atol=0.5 * E4M3_MIN_SUBNORMAL * 2.0, + ) + finally: + env.shutdown() + + +def test_fp8_byte_domain_and_out_dtype_range() -> None: + """The dequantization is a plain conversion over the whole e4m3 domain: + all 256 byte values (including +-448, +-0, subnormals and both NaN + bytes) come back exactly at r = 1.0. The product is *not* clamped to + out_dtype's range either — 448 * 256 overflows fp16 to +-inf while bf16 + and fp32 hold 114688.""" + torch.manual_seed(47) + env = _fp8_env() + try: + rids = [0] + env.kv_cache_manager.add_dummy_requests(rids, token_nums=[32]) + metadata = env.prepare_metadata(rids, [16], 1, [16]) + pool = env.pool_tensor(0) + page = env.blocks(0)[0] + all_bytes = torch.arange(256, dtype=torch.uint8, device="cuda").view(torch.float8_e4m3fn) + pool[page, 0, 0, 0, :256].copy_(all_bytes) + # Row 1 carries the extremes the range check needs. + extremes = torch.zeros(HEAD_SIZE, dtype=torch.float32, device="cuda") + extremes[:4] = torch.tensor([448.0, -448.0, E4M3_MIN_SUBNORMAL, 0.0]) + pool[page, 0, 1, 0].copy_(extremes.to(torch.float8_e4m3fn)) + + for out_dtype in (torch.float32, torch.bfloat16, torch.float16): + ckv, _ = _call(env, metadata, 1, out_dtype, None) + torch.testing.assert_close( + ckv[0, :256], + all_bytes.float().to(out_dtype), + rtol=0.0, + atol=0.0, + equal_nan=True, + ) + assert ckv[0, :256].isnan().sum().item() == 2 # 0x7f and 0xff + torch.testing.assert_close( + ckv[1, :4], + extremes[:4].to(out_dtype), + rtol=0.0, + atol=0.0, + ) + + big = _scale_tensor(256.0) + ckv, _ = _call(env, metadata, 1, torch.float16, big) + assert ckv[1, 0].isinf().item() and ckv[1, 0] > 0 + assert ckv[1, 1].isinf().item() and ckv[1, 1] < 0 + for out_dtype in (torch.bfloat16, torch.float32): + ckv, _ = _call(env, metadata, 1, out_dtype, big) + torch.testing.assert_close( + ckv[1, :2], + torch.tensor([114688.0, -114688.0], device="cuda").to(out_dtype), + rtol=0.0, + atol=0.0, + ) + finally: + env.shutdown() + + +def test_fp8_scale_tensor_forms() -> None: + """Only element [0] of kv_scale_quant_orig is read, and only its dtype is + validated: a `[2]` tensor, a `[1]` tensor and a 0-dim scalar carrying the + same leading value are interchangeable, while a non-fp32 tensor raises.""" + torch.manual_seed(51) + env = _fp8_env() + try: + rids = [0] + env.kv_cache_manager.add_dummy_requests(rids, token_nums=[32]) + metadata = env.prepare_metadata(rids, [16], 1, [16]) + stored, _ = env.fill_layer(0, rids, [32], 1.0) + reference = _dequant_mirror(stored[0], _scale_tensor(2.0), torch.bfloat16) + forms = [ + torch.full((1,), 2.0, dtype=torch.float32, device="cuda"), + torch.tensor([2.0, 99.0], dtype=torch.float32, device="cuda"), + torch.tensor(2.0, dtype=torch.float32, device="cuda"), + ] + for scale in forms: + ckv, kpe = _call(env, metadata, 1, torch.bfloat16, scale) + torch.testing.assert_close(torch.cat([ckv, kpe], dim=1), reference, rtol=0.0, atol=0.0) + for dtype in (torch.float16, torch.float64): + bad = torch.full((1,), 2.0, dtype=dtype, device="cuda") + message = _raises(lambda s=bad: _call(env, metadata, 1, torch.bfloat16, s)) + assert "expected scalar type Float but found" in message, message + finally: + env.shutdown() + + +def test_fp8_quant_mode_and_out_dtype_rejections() -> None: + """quant_mode: any KV-cache quantization other than fp8 is rejected; + extra quantization bits alongside the fp8 one are not. out_dtype: only + fp16 / fp32 / bf16 are accepted.""" + torch.manual_seed(48) + env = _fp8_env() + try: + rids = [0] + env.kv_cache_manager.add_dummy_requests(rids, token_nums=[32]) + metadata = env.prepare_metadata(rids, [16], 1, [16]) + env.fill_layer(0, rids, [32], 1.0) + scale_t = _scale_tensor(1.0) + + message = _raises( + lambda: _call(env, metadata, 1, torch.bfloat16, scale_t, quant_mode=QM_INT8_KV_CACHE) + ) + assert "Only FP8 KV cache is supported for now" in message, message + # NVFP4_KV_CACHE fails earlier and for a different reason: the nvfp4 + # pool-pointer layout, checked before the quant_mode branch. + message = _raises( + lambda: _call(env, metadata, 1, torch.bfloat16, scale_t, quant_mode=QM_NVFP4_KV_CACHE) + ) + assert "hostKvCachePoolPointers.dim() == 3" in message, message + + for out_dtype in (torch.float8_e4m3fn, torch.int8, torch.float64): + message = _raises(lambda od=out_dtype: _call(env, metadata, 1, od, scale_t)) + assert "out_dtype only support float16, float32, bfloat16" in message, message + + # The fp8 bit is what selects the path; other quantization bits ride + # along unread. 1152 is what an fp8-KV checkpoint's quant config + # produces (FP8_KV_CACHE | FP8_1x128_128x128). + base = torch.cat(_call(env, metadata, 1, torch.bfloat16, scale_t), dim=1) + for quant_mode in ( + QM_FP8_KV_CACHE | QM_FP8_1X128_128X128, + QM_FP8_KV_CACHE | QM_FP8_QDQ, + ): + other = torch.cat( + _call(env, metadata, 1, torch.bfloat16, scale_t, quant_mode=quant_mode), + dim=1, + ) + torch.testing.assert_close(other, base, rtol=0.0, atol=0.0) + finally: + env.shutdown() + + +def test_quant_mode_selects_cache_element_width() -> None: + """The pool arrives as a raw pointer, so quant_mode alone decides how + wide a cache element is — one byte with the fp8 bit set, sizeof(out_dtype) + without it. Both mismatches are silent and both are gated here, which is + also this file's control that a wrong read is visible to its comparison: + the same bitwise check that passes on every certified case fails on + these.""" + torch.manual_seed(49) + # An fp8 pool read with quant_mode=0: 2-byte elements over a 1-byte pool, + # so block b is read from byte 2 * b * tokens_per_block * (C+R) and each + # token consumes 2 * (C+R) bytes. + env = _fp8_env() + try: + rids = [0, 1] + env.kv_cache_manager.add_dummy_requests(rids, token_nums=[32, 32]) + metadata = env.prepare_metadata(rids, [16, 16], 2, [16, 16]) + stored, _ = env.fill_layer(0, rids, [32, 32], 1.0) + truth = torch.cat(stored, dim=0).float().to(torch.bfloat16) + ckv, kpe = _call(env, metadata, 2, torch.bfloat16, None, quant_mode=0) + got = torch.cat([ckv, kpe], dim=1) + assert not torch.equal(got, truth) + pool_bytes = env.pool_tensor(0).view(torch.uint8).reshape(-1) + stride = 2 * env.tokens_per_block * HEAD_SIZE + for out_row, rid in ((0, 0), (32, 1)): + base = env.blocks(rid)[0] * stride + expected = pool_bytes[base : base + 2 * HEAD_SIZE].view(torch.bfloat16) + torch.testing.assert_close(got[out_row], expected, rtol=0.0, atol=0.0, equal_nan=True) + finally: + env.shutdown() + + torch.manual_seed(50) + # A bf16 pool read with quant_mode=128: 1-byte elements over a 2-byte + # pool, so block b is read from byte b * tokens_per_block * (C+R) — half + # the intended stride, always in bounds and always wrong. + env = _MlaCacheEnv(DataType.BF16, tokens_per_block=PAGE32, max_tokens=16384) + try: + rids = [0, 1] + env.kv_cache_manager.add_dummy_requests(rids, token_nums=[32, 32]) + metadata = env.prepare_metadata(rids, [16, 16], 2, [16, 16]) + stored, _ = env.fill_layer(0, rids, [32, 32]) + truth = torch.cat(stored, dim=0) + ckv, kpe = _call(env, metadata, 2, torch.bfloat16, None, quant_mode=QM_FP8_KV_CACHE) + got = torch.cat([ckv, kpe], dim=1) + assert not torch.equal(got, truth) + pool_bytes = env.pool_tensor(0).view(torch.uint8).reshape(-1) + stride = env.tokens_per_block * HEAD_SIZE + for out_row, rid in ((0, 0), (32, 1)): + base = env.blocks(rid)[0] * stride + expected = ( + pool_bytes[base : base + HEAD_SIZE] + .view(torch.float8_e4m3fn) + .float() + .to(torch.bfloat16) + ) + torch.testing.assert_close(got[out_row], expected, rtol=0.0, atol=0.0, equal_nan=True) + + # And off the fp8 path kv_scale_quant_orig is read by nothing: the + # same call with a scale of 4.0 returns byte-identical output. + unscaled = torch.cat(_call(env, metadata, 2, torch.bfloat16, None), dim=1) + scaled = torch.cat(_call(env, metadata, 2, torch.bfloat16, _scale_tensor(4.0)), dim=1) + torch.testing.assert_close(scaled, unscaled, rtol=0.0, atol=0.0) + torch.testing.assert_close(unscaled, truth, rtol=0.0, atol=0.0) + finally: + env.shutdown() diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q.md b/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q.md new file mode 100644 index 000000000000..8b31d6ad40bb --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q.md @@ -0,0 +1,375 @@ +--- +receipts: + sm_100: {status: passed, trtllm: 1.3.0rc21} + sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 13} +--- + +# mla_rope_append_paged_kv_assign_q + +**Wraps** `torch.ops.trtllm.mla_rope_append_paged_kv_assign_q` (one call). + +## Semantics + +The context-phase (prefill) preprocessing step of MLA attention when a +cached KV prefix exists (KV reuse or chunked prefill), over one prepared +batch. Two configurations are certified, `beam_width=1` and DeepSeek MLA +geometry (`lora_size=512`, `rope_size=64`, `nope_size=128`) in both: + +- **matching-dtype latent pool** (`quant_mode=0`, + `kv_scale_orig_quant=None`): bf16 or fp16 activations with a paged latent + KV cache of the same dtype. The appended row is a plain dtype-preserving + write. +- **fp8-e4m3 latent pool** (`quant_mode=128`): bf16 activations with a + one-byte-per-element e4m3 pool. Only the **appended pool row** is + quantized — `q` and `latent_cache` are still rotated in place as bf16. + +The rest of this section describes the matching-dtype pool; the fp8 +subsection below states exactly which of its three effects changes (only +the third) and which do not. + +The batch has `num_contexts` context-phase sequences first; trailing +generation sequences are ignored entirely. Rows of `q` and `latent_cache` +are the context sequences' **new (uncached) tokens**, concatenated in batch +order. For context sequence `s`, with + +``` +cached_s = cu_ctx_cached_kv_lens[s+1] - cu_ctx_cached_kv_lens[s] # prefix already in cache +kv_s = cu_seq_lens[s+1] - cu_seq_lens[s] # cached + new +new_s = kv_s - cached_s # rows this op consumes +``` + +one call computes, for each new token `i in [0, new_s)` (tensor row `t`, +absolute position `p = cached_s + i`), with `R = rope_size`, +`C = lora_size`, `N = nope_size`, `H = head_num`: + +``` +# GPT-J (interleaved-pair) RoPE, fp32 math, one rounding to the I/O dtype. +# (cos_d, sin_d) = cos_sin_cache row p, pair d, d in [0, R/2) +rope(x)[2d] = x[2d] * cos_d - x[2d+1] * sin_d +rope(x)[2d+1] = x[2d] * sin_d + x[2d+1] * cos_d + +# 1. q RoPE in place (per head h): the rope tail of each head is rotated, +# reading its own pre-rotation contents ("assign q"). +q[t, h*(N+R)+N : (h+1)*(N+R)] = rope(q[t, h*(N+R)+N : (h+1)*(N+R)]) +# 2. k_pe RoPE in place on the latent row: +latent_cache[t, C:] = rope(latent_cache[t, C:]) +# 3. latent-cache append, at page/slot of position p in seq s's blocks: +cache_row(s, p) = concat(latent_cache[t, :C], # bitwise copy + rope(k_pe)) # same rotated values as 2. +``` + +Fusion boundary: q RoPE, k_pe RoPE, and the paged-cache append happen +inside the call. The caller still owns the projections that produced `q` +and `latent_cache`, reading the full `[cached + new]` latent KV back out of +the pool, the `kv_b_proj` up-projection, and the context FMHA itself. +`q`'s nope regions, `latent_cache`'s compressed region, and every +already-cached pool row (positions `< cached_s`) are bitwise untouched. + +### fp8-e4m3 latent pool (`quant_mode=128`) + +`quant_mode=128` is `QuantMode`'s fp8-KV-cache bit. It switches the +**latent pool's** element type to fp8-e4m3, one byte per element, while +`q` and `latent_cache` stay bf16. Certified at the DeepSeek-R1-0528 cell — +`head_num = 128`, `tokens_per_block = 32`, `lora/nope/rope = 512/128/64`, +`beam_width = 1`, single-layer pool — with the KV scaling factor omitted +(the production call) and explicit 1.0 / 1/1.5 / 0.5 / 0.25 / 2.0 swept +beside it, over fresh-prefill (`cached_s = 0`) and cached-prefix +(`cached_s > 0`) context sequences in the same call. + +One fp32 scalar steers it: + +``` +w = kv_scale_orig_quant[0] # 1.0 when that tensor is None +e4m3(x) = round-to-nearest cast of an fp32 x to float8_e4m3fn, + saturating at +-448 (see Notes) +``` + +Effects 1 and 2 above are **unchanged** — `q`'s rope tails and +`latent_cache`'s k_pe are rotated in place and written back as bf16, +un-quantized. Effect 3 becomes, with `bf16(.)` the rounding to the +activation dtype that the in-place write already performs: + +``` +# 3. latent-cache append, quantized, at page/slot of position p in seq s: +cache_row(s, p) = e4m3(concat(latent_cache[t, :C], + bf16(rope(k_pe))) * w) +``` + +`w` is a **per-tensor static scale**, not a block or dynamic one: every +appended row is the same multiply-and-round, bit-exact against that mirror +at each scale tested. The multiply happens in fp32 *after* the rounding to +bf16, so `e4m3(bf16(x) * w)` and `e4m3(bf16(x * w))` are different tensors +for a non-power-of-two `w` and the kernel computes the first (measured, see +Notes). + +**`q` is not quantized here, and neither is `latent_cache`.** Both come +back as ordinary bf16 rotations, matching the same reference the +`quant_mode=0` path matches, and neither is an e4m3 round trip (6.2% of +their elements survive an e4m3 round trip unchanged — the same fraction as +the untouched nope slice of the same tensor, and far from the 100% a +quantized tensor would show). The FMHA that consumes `q` afterwards is a +separate call — `trtllm::attention` with +`attention_input_type = context_only` and `latent_cache=None` — and on the +fp8 MLA context path that op was measured to quantize its own `q`/`k`/`v` +to e4m3 internally. That is that op's fact, established separately and on +its fresh-prefill flavor rather than the cached-KV flavor this op feeds; +what matters here is that nothing this op writes into `q` is pre-quantized, +so the pair is not a double quantization. + +The `w` on the pool row is likewise **not** applied to `q`: two runs over +identical inputs at `w = 1.0` and `w = 0.25` returned bitwise-identical `q` +and `latent_cache` while the appended rows moved (tested). + +`w` is a write-side factor only. Whatever later reads the pool back has to +supply its own read-side factor `1/w`; this op takes no such argument and +checks no relation. The generation-phase sibling +`trtllm::mla_rope_generation`, which appends into the same pool for decode +steps, takes both (`kv_scale_orig_quant` and `kv_scale_quant_orig`) and +uses them independently; its own append was separately measured to be the +same `e4m3(row * w)` static rule, so a pool written by both phases is +uniformly scaled as long as the caller passes the same `w` to each. + +## Signature + +```python +def mla_rope_append_paged_kv_assign_q( + q: torch.Tensor, + latent_cache: torch.Tensor, + num_contexts: int, + cu_ctx_cached_kv_lens: torch.Tensor, + cu_seq_lens: torch.Tensor, + max_input_uncached_seq_len: int, + cos_sin_cache: torch.Tensor, + head_num: int, + nope_size: int, + rope_size: int, + lora_size: int, + kv_cache_block_offsets: torch.Tensor, + host_kv_cache_pool_pointers: torch.Tensor, + host_kv_cache_pool_mapping: torch.Tensor, + kv_scale_orig_quant: Optional[torch.Tensor], + layer_idx: int, + tokens_per_block: int, + attention_window_size: int, + beam_width: int, + quant_mode: int, +) -> None +``` + +With `nc` = `num_contexts`, `S` = total sequences in the batch (contexts +first), `T` = `sum(new_s)` over context sequences, and `C`/`R`/`N`/`H` as +above: + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `q` | `[T, H*(N+R)]`; per head only the trailing `R` written, in the activation dtype on both paths (never quantized) | bf16 / fp16 (bf16 only on the fp8 path) | 2-D enforced (other ranks rejected), contiguous | CUDA | +| `latent_cache` | `[T, C+R]` = per-token `[compressed_kv \| k_pe]`; only `[..., C:]` written, in the activation dtype on both paths | same as `q` | 2-D enforced, contiguous | CUDA | +| `num_contexts` | leading context-phase sequence count | Python int | — | — | +| `cu_ctx_cached_kv_lens` | `[>= nc+1]`; `[0, cumsum(cached_s)]`; only the first `nc+1` entries read | int64 | contiguous | CUDA | +| `cu_seq_lens` | `[>= nc+1]`; `[0, cumsum(kv_s)]`; only the first `nc+1` entries read | int64 | contiguous | CUDA | +| `max_input_uncached_seq_len` | `max(new_s)` over context seqs (grid bound; production passes the exact max) | Python int | — | — | +| `cos_sin_cache` | `[1, max_pos * R * 2]`: per position `R` fp32 `(cos, sin)` pairs, second `R/2` a duplicate of the first (duplicated/`duplicate_data=True` layout; only pairs `[0, R/2)` are read per position) | fp32 | contiguous | CUDA | +| `head_num` | query heads `H`; 16 (matching-dtype pool) and 128 (fp8 pool) certified | Python int | — | — | +| `nope_size` | `N` (128 certified) | Python int | — | — | +| `rope_size` | `R` (64 certified) | Python int | — | — | +| `lora_size` | `C` (512 certified) | Python int | — | — | +| `kv_cache_block_offsets` | `[1, >= S, 2, max_blocks_per_seq]`; raw block ids per seq (K and V rows identical, kv_factor=1 pool); contexts first | int32 | contiguous | CUDA | +| `host_kv_cache_pool_pointers` | `[num_pools, 2]`: (primary ptr, secondary ptr=0) | int64 | contiguous | CPU | +| `host_kv_cache_pool_mapping` | `[num_layers, 2]`: (pool index, layer-within-pool) row per layer | int32 | contiguous | CPU | +| `kv_scale_orig_quant` | `quant_mode=0`: `None`. `quant_mode=128`: `[1]` holding the **write-side** factor `w` — everything this call quantizes is multiplied by it. `None` = 1.0, which is what the engine's own MLA call site passes. Only element `[0]` is read (a longer tensor is accepted) | fp32 (fp16 rejected, see *Notes*) | contiguous | CUDA | +| `residual_dim` | — | `int` | **rc26 addition.** Must be `0` or `rope_size`; the op rejects non-zero unless the KV pool is FP4. `0` on every path this entry certifies (bf16 and fp8-e4m3 pools), which is also what the in-tree caller passes | — | +| `layer_idx` | row into `host_kv_cache_pool_mapping` | Python int | — | — | +| `tokens_per_block` | pool page size; 32 and 64 certified on the matching-dtype pool, 32 on the fp8 pool (32 is what a default `KvCacheConfig` produces) | Python int | — | — | +| `attention_window_size` | `>= max(kv_s)`; production passes the manager's `max_seq_len` (smaller values imply cyclic-cache addressing, not certified) | Python int | — | — | +| `beam_width` | `1` | Python int | — | — | +| `quant_mode` | `0` (matching-dtype pool) and `128` = `QuantMode.FP8_KV_CACHE` (fp8-e4m3 pool) certified. A **non-fp8 quantized** KV cache is rejected — see *Notes*. `128 \| 256` (an fp8-weights bit alongside) was accepted rather than rejected; its output was not compared | Python int | — | — | +| returns | — (mutates `q`, `latent_cache`, and the paged pool) | — | — | — | + +## Metadata consumed + +No thread-local or registered-layer state: all inputs are explicit +arguments. The length/addressing tensors are exactly what a +`TrtllmAttentionMetadata` prepared with +`enable_context_mla_with_cached_kv=True` and its `KVCacheManager` expose: +`ctx_cached_token_indptr` → `cu_ctx_cached_kv_lens`, `ctx_kv_indptr` → +`cu_seq_lens`, `max_ctx_seq_len` → `max_input_uncached_seq_len`, +`metadata.kv_cache_block_offsets`, `manager.kv_cache_pool_pointers`, +`manager.kv_cache_pool_mapping`, `manager.tokens_per_block`, +`manager.max_seq_len` → `attention_window_size`. + +## Preconditions + +- The paged pool addressed by `host_kv_cache_pool_pointers` + + `host_kv_cache_pool_mapping[layer_idx]` is an MLA latent cache: + kv_factor 1, one kv head, row width `C + R`, page size + `tokens_per_block`. Per page the layout is `[tokens_per_block, C + R]` + and the page slab is `tokens_per_block * (C + R)` **elements**, whose + byte width the op takes from `quant_mode` and the activation dtype, not + from anything the cache manager tells it: the activation dtype's width at + `quant_mode=0` (2 bytes for the certified bf16/fp16), 1 byte (e4m3) at + `quant_mode=128`. The pool's real element type must match — a + `KVCacheManager` built with `DataType.BF16` / `DataType.HALF` and with + `DataType.FP8` respectively is what was certified. +- Every context sequence has block capacity for `kv_s` tokens **before** + the call (block offsets in `kv_cache_block_offsets` cover + `ceil(kv_s / tokens_per_block)` pages). +- `q` and `latent_cache` are 2-D, contiguous, of the same dtype, with + exactly `T` rows: the context sequences' new tokens in batch order + (production slices `q[:num_ctx_tokens]`; generation-token rows must not + be included). +- `q`'s rope tails and `latent_cache`'s k_pe hold **pre-rotation** values; + the call rotates them in place. A table with cos=1/sin=0 makes the RoPE + identity (pre-rotated inputs pass through). +- `cu_ctx_cached_kv_lens` / `cu_seq_lens` are int64 (the kernel rejects + other index dtypes), on device, start at 0, and satisfy + `cached_s <= kv_s` per sequence. +- `max_input_uncached_seq_len >= new_s` for every context sequence + (certified with the exact max), and `kv_s <= attention_window_size`, + `kv_s - 1 < max_pos` of the table. +- `cos_sin_cache` uses the duplicated `(cos, sin)`-pair layout above + (`RopeParams(..., duplicate_data=True).create_rope_const_params()` + produces it; MLA models set `duplicate_data=True`). +- **Matching-dtype pool only:** `quant_mode=0` and + `kv_scale_orig_quant=None`. +- **fp8 pool only** (`quant_mode=128`): + - `q` and `latent_cache` are bf16 (fp16 activations over an fp8 pool were + not run). + - `kv_scale_orig_quant` is a `[1]` fp32 CUDA tensor or `None` (= 1.0). + `None` is the production configuration and is what the engine's own MLA + call site passes. There is no read-side factor here: a consumer that + dequantizes the pool needs `1/w` from somewhere else, and nothing in + this call checks it. + - Input magnitudes must survive e4m3: `|value * w|` below `2**-10` + (~9.8e-4) rounds to zero, and anything above 448 saturates to `±448` + (measured — see *Notes*). The certified runs are unit-normal + activations at `w <= 2`. + +## Notes + +- Schema under-annotation: the registered schema marks **nothing** mutable — + it declares `(Tensor q, Tensor latent_cache, ...) -> ()`, with no + `Tensor(a!)` anywhere — yet the call rotates `q` and `latent_cache` in + place and appends to the paged pool. Do not rely on the schema's alias + info (e.g. under torch.compile functionalization). +- Precision of the in-place rotations (bf16 and fp16, sm_100, **both** + paths — the fp8 path rotates `q` and `latent_cache` exactly as the + matching-dtype path does): the `compressed_kv` copy is bitwise. Both RoPE + outputs match an fp32 reference rounded once to the I/O dtype and are + bitwise on all but 1.4e-5 to 1.7e-5 of elements: the kernel and a torch + fp32 reference evaluate `x*cos ∓ y*sin` in different orders, so where the + correctly-rounded fp32 result lands near a rounding midpoint of the I/O + dtype the two round to adjacent values. Measured in bf16 at `H = 16` on a + ~400-token context call: 7 such elements in 411 648 rotated `q` values at + page size 32 and 8 at page size 64, all within one bf16 ulp (0 in 25 728 + rotated `k_pe` values either way) — it is arithmetic, not addressing. At + `H = 128` on the same batch shape the tail reaches further: 46 elements in + 3 293 184 differ, one of them by 4 bf16 ulps, but the **absolute** + deviation stays ≤ 0.0078 because the elements that drift are the ones + where `x*cos - y*sin` nearly cancels. Default `assert_close` tolerances + absorb all of it with margin, which is what the test gates on; a caller + needing bit-exactness against a torch reference will not get it. +- Precision of the fp8 append (sm_100): **everything came back bit-exact** + against `e4m3(row * w)`, the `compressed_kv` half and the roped `k_pe` + half alike, at `w` ∈ {1.0 (both as `None` and as an explicit tensor), + 1/1.5, 0.5, 0.25, 2.0}. e4m3's 3-bit mantissa absorbs the evaluation-order + difference the bf16 rotation pays an ulp for. The test still gates the + roped half at one e4m3 ulp (`2**-3` relative) plus a bit-exact-fraction + floor rather than bitwise, because the tie it covers is arithmetic; + neither half has been approached. +- **The double rounding is real and a mirror must reproduce it.** The + quantized `k_pe` is `e4m3(bf16(rope(k_pe)) * w)` — and that bf16 value is + the one the call also writes back into `latent_cache` (both came back + bit-exact against the same torch reference), so a caller already holding + the post-call `latent_cache` can mirror the pool from it directly. + Quantizing the unrounded fp32 RoPE result instead disagrees with the + kernel on 2.4-3.3% of the appended `k_pe` bytes (measured at every scale), + which is enough to fail a bit-exact gate and small enough to pass a sloppy + one. Applying the scale before the bf16 rounding — + `e4m3(bf16(rope * w))` — disagrees on 1.05% of bytes at `w = 1/1.5` and is + identical at the power-of-two scales, so only a non-power-of-two scale + separates the two orders. +- A `quant_mode` asking for a **non-fp8 quantized** KV cache is rejected: + `quant_mode = 64` (`QuantMode.INT8_KV_CACHE`) raises + `RuntimeError: [TensorRT-LLM][ERROR] Assertion failed: Only FP8 KV cache + is supported for now (../tensorrt_llm/thop/mlaPreprocessOp.cpp:316)`. + Tested. `quant_mode = 8192` (`NVFP4_KV_CACHE`) also fails, but earlier and + for a different reason — `Expected hostKvCachePoolPointers.dim() == 3`, + the nvfp4 pool-pointer layout — so it is not the same guard. +- e4m3 range behaviour at the boundaries (scratch probe, `w = 1.0`): the + kernel **saturates**, writing `+448` for an input of 1000 or `+inf` and + `-448` for -1000, and flushes 1e-5 to `+0`. `torch.Tensor.to(float8_e4m3fn)` + does *not* saturate — it yields the NaN byte `0x7F` for 1000 and for + `inf` — so a torch mirror and the kernel agree only while every + `|value * w|` stays ≤ 448. Every certified run does. +- `kv_scale_orig_quant` validation is thin: a `[2]` fp32 CUDA tensor is + accepted and only `[0]` is used, and a **CPU** fp32 tensor is accepted + too — in one scratch probe on this machine it even produced the same bytes + as the equivalent CUDA tensor, so device placement is not checked and a + wrong-device scale will not announce itself. A `float16` tensor is + rejected loudly (`RuntimeError: expected scalar type Float but found + Half`). +- **Not idempotent**: the q and k_pe rotations read their own pre-call + contents in place, so calling twice on the same buffers composes two + rotations (unlike ops that read a separate source). Re-preparing inputs + is required before any retry. +- `cached_s = 0` (fresh prefill, nothing cached) is valid and certified on + both paths; positions then start at 0. `cached_s > 0` — a context call + over a prefix already in the pool, which is what block reuse makes the + engine do — is certified on both paths too, and the prefix rows come back + bitwise unchanged. +- Pool writes are exactly bounded on both paths: one `C+R`-element row per + new context token, at the `(page, slot)` its absolute position addresses, + and nothing else in the pool moves — checked by a whole-pool byte snapshot + on every case, including a 402-row single call spanning four sequences and + a mixed batch whose trailing generation sequences own their own pages. + That is also what pins the page-slab geometry described under + *Preconditions*. +- Page size (sm_100): certified at `tokens_per_block` 32 and 64 on the + matching-dtype pool, 32 on the fp8 pool. 32 is what + a `KvCacheConfig` at its defaults hands the op; 64 is what a tuned MLA + target opts into. At 32 the cached prefixes were chosen to put the first + new token at every alignment a 32-slot page has — 0 (fresh prefill), + 32 and 96 (slot 0 of a fresh page, after one and three full pages), + 1 / 33 / 100 (mid-page), 31 (a page's last slot, so the run crosses on its + very first token), 512 (sixteen full pages) and 511 (the last slot of page + 15, which the single new token closes) — with new-token runs of 1 to 300 + tokens, the longest walking ten pages in one call, in bf16 (including + through pool-mapping row 1 of a two-layer pool) and in fp16, plus a mixed + batch whose trailing generation sequences own their own pages. Every effect + held at the same gates as at page 64: RoPE within the one-ulp band above, + `compressed_kv` copy bitwise, cached prefix and the untouched slices + bitwise. The fp8 cases run page 32 only, over the same alignment set — + first new token at 0, 17, 31, 32, 96, 100 and 511, runs of 1 to 300 tokens, + and a 402-row four-sequence call. Other page sizes are untested rather than + known-bad. +- Both `q.dim() == 2` and `latent_cache.dim() == 2` are enforced with + runtime errors. A comment in the runtime describes `latent_cache` as + `[tokens, 1, C+R]`, but the op itself rejects 3-D — callers flatten to + rows first. +- Only the DeepSeek-V3 geometry (`C=512`, `R=64`, `N=128`) was exercised; + other sizes are unverified. +- Sibling ops exist for adjacent MLA roles: + `trtllm::load_paged_kv_cache_for_mla` (gathers the full `[cached + new]` + latent KV this op just completed, and takes its own read-side + `kv_scale_quant_orig`), `trtllm::load_chunked_kv_cache_for_mla` + and `trtllm::merge_chunked_attention_for_mla` (chunked-prefill variants), + `trtllm::mla_rope_generation` (the generation-phase counterpart), and the + context FMHA that consumes the up-projected K/V. +- Not certified on either path: `beam_width > 1`, + `attention_window_size` below the KV length (cyclic-cache addressing), and + MTP / speculative-decoding batches. Not certified on the fp8 path + specifically: + `tokens_per_block = 64`, head counts other than 128, fp16 activations, + multi-layer latent pools, `|value * w| > 448`, and a YaRN-scaled rope + table (the fp8 cases run the unscaled theta-10000 table; the table's + *content* is not an axis this op's gate exercises). + + +## rc26 change to the accepted KV-cache formats + +The op accepted only an fp8-e4m3 latent pool in 1.3.0rc21 and now also +accepts NVFP4; its rejection message changed accordingly from +`Only FP8 KV cache is supported for now` to `Only FP8 and NVFP4 KV +caches are supported for now`. An int8 pool (`quant_mode=64`) is still +rejected. **NVFP4 is accepted by the op but not certified here** — no +cell in this entry's test drives it, so it stays outside the envelope. diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q.py b/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q.py new file mode 100644 index 000000000000..d88e22599ebc --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q.py @@ -0,0 +1,63 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""MLA context-phase RoPE of q_pe/k_pe in place + latent paged-KV-cache append.""" + +from typing import Optional + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + + +def mla_rope_append_paged_kv_assign_q( + q: torch.Tensor, + latent_cache: torch.Tensor, + num_contexts: int, + cu_ctx_cached_kv_lens: torch.Tensor, + cu_seq_lens: torch.Tensor, + max_input_uncached_seq_len: int, + cos_sin_cache: torch.Tensor, + head_num: int, + nope_size: int, + rope_size: int, + lora_size: int, + kv_cache_block_offsets: torch.Tensor, + host_kv_cache_pool_pointers: torch.Tensor, + host_kv_cache_pool_mapping: torch.Tensor, + kv_scale_orig_quant: Optional[torch.Tensor], + # ``residual_dim`` (rc26; absent in rc21) must be 0 or ``rope_size``, + # and the op rejects non-zero unless the KV pool is FP4. Every caller + # here runs a bf16 or fp8-e4m3 pool, so 0 is the only legal value. + residual_dim: int, + layer_idx: int, + tokens_per_block: int, + attention_window_size: int, + beam_width: int, + quant_mode: int, +) -> None: + """RoPE each new context token's q_pe (in q) and k_pe (in latent_cache) + in place at its absolute position, and append the token's latent row + [compressed_kv | rope(k_pe)] to the paged MLA KV cache. Returns None.""" + torch.ops.trtllm.mla_rope_append_paged_kv_assign_q( + q, + latent_cache, + num_contexts, + cu_ctx_cached_kv_lens, + cu_seq_lens, + max_input_uncached_seq_len, + cos_sin_cache, + head_num, + nope_size, + rope_size, + lora_size, + kv_cache_block_offsets, + host_kv_cache_pool_pointers, + host_kv_cache_pool_mapping, + kv_scale_orig_quant, + residual_dim, + layer_idx, + tokens_per_block, + attention_window_size, + beam_width, + quant_mode, + ) diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q_test.py b/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q_test.py new file mode 100644 index 000000000000..6077938c440c --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q_test.py @@ -0,0 +1,955 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the mla_rope_append_paged_kv_assign_q catalog entry. + +The op reads paged-KV-cache addressing tensors and cumulative-length +tensors that the runtime normally derives from a KVCacheManager and a +TrtllmAttentionMetadata prepared with enable_context_mla_with_cached_kv=True. +The test builds that state for real, pre-writes known cached-prefix rows +into the paged pool, then checks the three kernel effects against a torch +fp32 reference: GPT-J RoPE of each new context token's q_pe written back +into q at the token's absolute position, the same RoPE of k_pe written back +into latent_cache, and [compressed_kv | rope(k_pe)] appended into the paged +latent cache at that position. Cached-prefix rows, q's nope slices, and +latent_cache's compressed_kv slice must be bitwise untouched. + +Two surfaces: + +1. Matching-dtype latent pool (`quant_mode=0`, `kv_scale_orig_quant=None`), + H = 16, in bf16 (production MLA dtype) and fp16, at both pool page + sizes: 64, and 32 (the engine default), where the append walks twice as + many pages and the cached prefixes place the first new token at every + 32-slot alignment. + +2. fp8-e4m3 latent pool (`quant_mode=128`) at the DeepSeek-R1-0528 cell — + H = 128, tokens_per_block = 32, C/R/nope = 512/64/128, beam_width = 1, + single-layer pool, bf16 activations. Here q and latent_cache are still + roped in place in bf16 (neither is quantized), and only the appended + pool row lands as e4m3, scaled by kv_scale_orig_quant. Swept over that + scaling factor (omitted = the production call, and explicit + 1.0 / 1/1.5 / 0.5 / 0.25 / 2.0), over fresh-prefill and cached-prefix + context calls in the same batch, and over a mixed batch whose trailing + generation sequences must be ignored. +""" + +from typing import List, NamedTuple, Optional, Sequence, Union + +import torch + +from tensorrt_llm._torch.attention.backends.interface import RopeParams +from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata +from tensorrt_llm._torch.metadata import KVCacheParams +from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager +from tensorrt_llm.bindings import DataType +from tensorrt_llm.bindings.internal.batch_manager import CacheType +from tensorrt_llm.llmapi.llm_args import KvCacheConfig +from tensorrt_llm.mapping import Mapping + +from .mla_rope_append_paged_kv_assign_q import mla_rope_append_paged_kv_assign_q + +assert torch.cuda.is_available(), "mla_rope_append_paged_kv_assign_q requires a CUDA device" + +# DeepSeek-V3 MLA head geometry (num_heads reduced to a TP-slice-like 16). +NUM_HEADS = 16 +KV_LORA_RANK = 512 +QK_ROPE_HEAD_DIM = 64 +QK_NOPE_HEAD_DIM = 128 +QK_HEAD_DIM = QK_NOPE_HEAD_DIM + QK_ROPE_HEAD_DIM +LATENT_SIZE = KV_LORA_RANK + QK_ROPE_HEAD_DIM +# Pool page size. 32 is the engine default (KvCacheConfig.tokens_per_block), +# 64 the value a tuned MLA target opts into; both are covered below. +TOKENS_PER_BLOCK = 64 +PAGE32 = 32 +MAX_SEQ_LEN = 1024 + +# ─── fp8-e4m3 latent pool surface ───────────────────────────────────── +# DeepSeek-R1-0528 cell: full 128-head MLA over an fp8 latent pool. +NUM_HEADS_R1 = 128 +QUANT_MODE_FP8_KV_CACHE = 128 # QuantMode.FP8_KV_CACHE +QUANT_MODE_INT8_KV_CACHE = 64 # QuantMode.INT8_KV_CACHE — rejected, see below + +# One e4m3 ulp. e4m3 keeps 3 mantissa bits, so 2**-3 relative is one ulp — the +# gate for the appended row's k_pe half, the only quantized output that passes +# through the in-kernel RoPE, where the kernel and a torch fp32 reference can +# evaluate x*cos -+ y*sin in different orders. Paired with a bit-exact-fraction +# floor, because in practice the coarse e4m3 rounding absorbs that difference +# completely: every appended k_pe byte measured on sm_100 under quant_mode 128 +# (all scales, all cases below) was bit-exact, so neither half of the gate has +# ever been approached. That both halves still discriminate is measured, not +# assumed — test_fp8_kv_explicit_unit_scale runs four wrong mirrors of the same +# bytes and asserts each blows the gate it is able to blow. +E4M3_ULP_RTOL = 2**-3 +E4M3_MAX_INEXACT_FRACTION = 1e-3 + +# The roped q must come back as an ordinary bf16 rotation, NOT quantized: this +# op leaves q to the context FMHA, which does its own e4m3 quantization. A bf16 +# value drawn from a continuous distribution survives an e4m3 round trip only +# when its mantissa already fits in 3 bits, which happens for ~6% of elements; +# a quantized q would sit at 100%. The gate is 25% — 4x the measured 6.2%, and +# 4x below a quantized tensor — and every case also compares its own untouched +# nope slice as an in-run baseline. +MAX_E4M3_ROUNDTRIP_FRACTION = 0.25 + +_TRTLLM_TO_TORCH_DTYPE = { + DataType.BF16: torch.bfloat16, + DataType.HALF: torch.float16, +} + + +class _MlaCtxEnv: + """Real op state: MLA (kv_factor=1) paged KV cache manager + duplicated- + layout RoPE table + context-MLA-with-cached-KV metadata. + + With fp8_pool the manager allocates an e4m3 latent pool (one byte per + element) while the activations (`q`, `latent_cache`) stay `cache_dtype`, + and `orig_quant` becomes the op's write-side KV scaling factor. + """ + + def __init__( + self, + cache_dtype: DataType = DataType.BF16, + num_layers: int = 1, + max_batch_size: int = 8, + tokens_per_block: int = TOKENS_PER_BLOCK, + num_heads: int = NUM_HEADS, + fp8_pool: bool = False, + orig_quant: Optional[float] = None, + ) -> None: + self.torch_dtype = _TRTLLM_TO_TORCH_DTYPE[cache_dtype] + self.max_batch_size = max_batch_size + self.tokens_per_block = tokens_per_block + self.num_heads = num_heads + self.fp8_pool = fp8_pool + self.pool_dtype = torch.float8_e4m3fn if fp8_pool else self.torch_dtype + self.quant_mode = QUANT_MODE_FP8_KV_CACHE if fp8_pool else 0 + self.kv_cache_manager = KVCacheManager( + KvCacheConfig(max_tokens=131072, enable_block_reuse=False), + CacheType.SELFKONLY, # MLA latent cache: kv_factor=1, one kv head + num_layers=num_layers, + num_kv_heads=1, + head_dim=LATENT_SIZE, + tokens_per_block=tokens_per_block, + max_seq_len=MAX_SEQ_LEN, + max_batch_size=max_batch_size, + mapping=Mapping(world_size=1, tp_size=1, rank=0), + dtype=DataType.FP8 if fp8_pool else cache_dtype, + ) + # The pool must really be paged at the size and element type the case + # claims: the op sizes its page slabs from tokens_per_block and + # quant_mode, not from anything the manager tells it. + assert self.kv_cache_manager.tokens_per_block == tokens_per_block + pool = self.kv_cache_manager.get_buffers(0) + assert pool is not None and pool.dtype == self.pool_dtype + # KV scaling factor. None is what the engine's own call site passes + # (TrtllmAttention.mla_rope_append_paged_kv_assign_q hard-codes it), + # and the op then behaves as if it were 1.0. + self.kv_scale_orig_quant = ( + None + if orig_quant is None + else torch.full((1,), orig_quant, dtype=torch.float32, device="cuda") + ) + # The exact fp32 number the op sees, as a 0-dim tensor, so the mirror + # cannot disagree with the kernel in the last bit. + self.write_scale = torch.tensor( + 1.0 if orig_quant is None else orig_quant, dtype=torch.float32 + ).cuda() + # Duplicated-layout fp32 (cos, sin) table, as the MLA backend builds + # it (RopeParams.from_config sets duplicate_data=True for MLA models). + rope = RopeParams( + dim=QK_ROPE_HEAD_DIM, + theta=10000.0, + max_positions=MAX_SEQ_LEN, + duplicate_data=True, + ) + _, self.rotary_cos_sin = rope.create_rope_const_params() + + def prepare_metadata( + self, + request_ids: List[int], + seq_lens: List[int], + num_contexts: int, + cached_lens: List[int], + ) -> TrtllmAttentionMetadata: + metadata = TrtllmAttentionMetadata( + max_num_requests=self.max_batch_size, + max_num_tokens=8192, + kv_cache_manager=self.kv_cache_manager, + enable_context_mla_with_cached_kv=True, + ) + metadata.seq_lens = torch.tensor(seq_lens, dtype=torch.int) + metadata.num_contexts = num_contexts + metadata.request_ids = request_ids + metadata.prompt_lens = [c + s for c, s in zip(cached_lens, seq_lens)] + metadata.kv_cache_params = KVCacheParams( + use_cache=True, + num_cached_tokens_per_seq=cached_lens, + ) + metadata.prepare() + return metadata + + def rope_fp32(self, x: torch.Tensor, positions: Union[int, Sequence[int]]) -> torch.Tensor: + """GPT-J interleaved rotation of the last dim in fp32, unrounded. + x is [n, ..., R] with positions of length n (or a single position for + x of shape [..., R]).""" + half = QK_ROPE_HEAD_DIM // 2 + table = self.rotary_cos_sin.view(-1, QK_ROPE_HEAD_DIM, 2) + if isinstance(positions, int): + cos = table[positions, :half, 0] + sin = table[positions, :half, 1] + else: + pos = torch.tensor(list(positions), dtype=torch.long, device="cuda") + # [n, half] broadcast over x's middle dims. + shape = [len(positions)] + [1] * (x.dim() - 2) + [half] + cos = table[pos, :half, 0].view(shape) + sin = table[pos, :half, 1].view(shape) + pairs = x.float().reshape(*x.shape[:-1], half, 2) + out = torch.empty_like(pairs) + out[..., 0] = pairs[..., 0] * cos - pairs[..., 1] * sin + out[..., 1] = pairs[..., 0] * sin + pairs[..., 1] * cos + return out.reshape(x.shape) + + def rope_ref(self, x: torch.Tensor, positions: Union[int, Sequence[int]]) -> torch.Tensor: + """rope_fp32 with one rounding back to x's dtype — what the kernel + writes into q and latent_cache.""" + return self.rope_fp32(x, positions).to(x.dtype) + + def quantize(self, x: torch.Tensor) -> torch.Tensor: + """e4m3(x * kv_scale_orig_quant) — the op's write-side quantization of + the appended latent row.""" + return (x.float() * self.write_scale).to(torch.float8_e4m3fn) + + def pool_tensor(self, layer_idx: int) -> torch.Tensor: + pool = self.kv_cache_manager.get_buffers(layer_idx) + assert pool is not None + return pool + + def blocks(self, request_id: int) -> List[int]: + return list(self.kv_cache_manager.get_batch_cache_indices([request_id])[0]) + + def slot(self, request_id: int, position: int) -> tuple[int, int]: + tpb = self.tokens_per_block + return self.blocks(request_id)[position // tpb], position % tpb + + def fill_rows(self, layer_idx: int, request_id: int, length: int) -> torch.Tensor: + """Write known random latent rows at positions [0, length) of one + sequence's pages; return the written rows [length, LATENT_SIZE].""" + draw_dtype = torch.float32 if self.fp8_pool else self.torch_dtype + rows = torch.randn(length, LATENT_SIZE, dtype=draw_dtype, device="cuda").to(self.pool_dtype) + pool = self.pool_tensor(layer_idx) # [pages,1,tpb,1,head] + tpb = self.tokens_per_block + blocks = self.blocks(request_id) + for start in range(0, length, tpb): + page = blocks[start // tpb] + n = min(tpb, length - start) + pool[page, 0, :n, 0].copy_(rows[start : start + n]) + return rows + + def read_rows(self, layer_idx: int, request_id: int, start: int, end: int) -> torch.Tensor: + """Read latent rows at positions [start, end) of one sequence back + from the paged pool as a dense [end-start, LATENT_SIZE] tensor.""" + pool = self.pool_tensor(layer_idx) + tpb = self.tokens_per_block + blocks = self.blocks(request_id) + out = torch.empty(end - start, LATENT_SIZE, dtype=self.pool_dtype, device="cuda") + for pos in range(start, end): + page = blocks[pos // tpb] + out[pos - start] = pool[page, 0, pos % tpb, 0] + return out + + def shutdown(self) -> None: + self.kv_cache_manager.shutdown() + + +def _assert_bytes_equal(got: torch.Tensor, want: torch.Tensor, what: str) -> None: + """Bit-exact comparison of two tensors, signed zeros included.""" + torch.testing.assert_close( + got.reshape(-1).view(torch.uint8), + want.reshape(-1).view(torch.uint8), + rtol=0.0, + atol=0.0, + msg=lambda m: f"{what} not bit-exact\n{m}", + ) + + +def _inexact_bytes(got: torch.Tensor, want: torch.Tensor) -> int: + return int((got.reshape(-1).view(torch.uint8) != want.reshape(-1).view(torch.uint8)).sum()) + + +def _assert_e4m3_roped(got: torch.Tensor, want: torch.Tensor, what: str) -> None: + """One-e4m3-ulp gate plus a bit-exact-fraction floor, for the part of the + fp8 output that passes through the in-kernel RoPE.""" + torch.testing.assert_close(got.float(), want.float(), rtol=E4M3_ULP_RTOL, atol=0.0) + inexact = _inexact_bytes(got, want) + allowed = max(1, int(E4M3_MAX_INEXACT_FRACTION * got.numel())) + assert inexact <= allowed, ( + f"{what} not bit-exact enough: {inexact}/{got.numel()} bytes differ (allowed {allowed})" + ) + + +def _fails_e4m3_tolerance(got: torch.Tensor, mirror: torch.Tensor) -> bool: + """Whether the tolerance half of _assert_e4m3_roped rejects this mirror.""" + try: + torch.testing.assert_close(got.float(), mirror.float(), rtol=E4M3_ULP_RTOL, atol=0.0) + except AssertionError: + return True + return False + + +def _e4m3_roundtrip_fraction(x: torch.Tensor) -> float: + """Fraction of elements of a bf16/fp16 tensor that survive an e4m3 round + trip unchanged — ~6% for continuous data, 100% for a quantized tensor.""" + rt = x.float().to(torch.float8_e4m3fn).to(x.dtype) + return float((x.reshape(-1) == rt.reshape(-1)).float().mean()) + + +class _Run(NamedTuple): + """What one op call was driven with and produced.""" + + q: torch.Tensor # post-call q, [T, H*(N+R)] + latent_cache: torch.Tensor # post-call latent_cache, [T, C+R] + latent_orig: torch.Tensor # pre-call latent_cache + ref_k_pe: torch.Tensor # torch reference for rope(k_pe), activation dtype + ref_k_pe_fp32: torch.Tensor # the same reference before rounding + rows: torch.Tensor # appended pool rows in row order, [T, C+R] + positions: List[int] # absolute position of every row + + +def _run_and_check( + env: _MlaCtxEnv, + request_ids: List[int], + seq_lens: List[int], + num_contexts: int, + cached_lens: List[int], + layer_idx: int = 0, +) -> _Run: + """Pre-fill every sequence's cached prefix, run the op over the context + sequences' new tokens, and verify every kernel effect.""" + kv_lens = [c + s for c, s in zip(cached_lens, seq_lens)] + metadata = env.prepare_metadata(request_ids, seq_lens, num_contexts, cached_lens) + # Known cached-prefix rows (all sequences) to check the op leaves them be. + prefix_rows = [env.fill_rows(layer_idx, rid, c) for rid, c in zip(request_ids, cached_lens)] + + num_heads = env.num_heads + ctx_new = seq_lens[:num_contexts] + num_tokens = sum(ctx_new) + assert int(metadata.num_ctx_tokens) == num_tokens + assert int(metadata.max_ctx_seq_len) == max(ctx_new) + # Absolute position of every new context token, in q/latent row order. + positions = [cached_lens[s] + i for s in range(num_contexts) for i in range(ctx_new[s])] + + q = torch.randn(num_tokens, num_heads * QK_HEAD_DIM, dtype=env.torch_dtype, device="cuda") + # The op rejects other ranks (thop checks latent_cache.dim() == 2). + latent_cache = torch.randn(num_tokens, LATENT_SIZE, dtype=env.torch_dtype, device="cuda") + q_orig = q.clone() + latent_orig = latent_cache.clone() + + block_offsets = metadata.kv_cache_block_offsets + pool_pointers = env.kv_cache_manager.kv_cache_pool_pointers + pool_mapping = env.kv_cache_manager.kv_cache_pool_mapping + assert block_offsets is not None + assert pool_pointers is not None and pool_mapping is not None + pool = env.pool_tensor(layer_idx) + pool_before = pool.clone() + + mla_rope_append_paged_kv_assign_q( + q, + latent_cache, + num_contexts, + metadata.ctx_cached_token_indptr, + metadata.ctx_kv_indptr, + int(metadata.max_ctx_seq_len), + env.rotary_cos_sin, + num_heads, + QK_NOPE_HEAD_DIM, + QK_ROPE_HEAD_DIM, + KV_LORA_RANK, + block_offsets, + pool_pointers, + pool_mapping, + env.kv_scale_orig_quant, # None outside the fp8-KV-cache path + 0, # residual_dim: 0 or rope_size; non-zero needs an FP4 KV pool + layer_idx, + env.tokens_per_block, + MAX_SEQ_LEN, # attention_window_size + 1, # beam_width + env.quant_mode, + ) + torch.cuda.synchronize() + + # 1. q: per head, the rope slice is rotated at the token's absolute + # position; the nope slice is bitwise untouched. This holds on the fp8 + # path too — q stays an unquantized activation tensor there. + q_heads = q.view(num_tokens, num_heads, QK_HEAD_DIM) + q_orig_heads = q_orig.view(num_tokens, num_heads, QK_HEAD_DIM) + ref_q_pe = env.rope_ref(q_orig_heads[..., QK_NOPE_HEAD_DIM:], positions) + torch.testing.assert_close(q_heads[..., QK_NOPE_HEAD_DIM:], ref_q_pe) + torch.testing.assert_close( + q_heads[..., :QK_NOPE_HEAD_DIM], + q_orig_heads[..., :QK_NOPE_HEAD_DIM], + rtol=0.0, + atol=0.0, # caller-owned region: must be bitwise untouched + ) + + # 2. latent_cache: k_pe rotated in place; compressed_kv bitwise untouched. + ref_k_pe_fp32 = env.rope_fp32(latent_orig[:, KV_LORA_RANK:], positions) + ref_k_pe = ref_k_pe_fp32.to(env.torch_dtype) + torch.testing.assert_close(latent_cache[:, KV_LORA_RANK:], ref_k_pe) + torch.testing.assert_close( + latent_cache[:, :KV_LORA_RANK], + latent_orig[:, :KV_LORA_RANK], + rtol=0.0, + atol=0.0, # read-only region: must be bitwise untouched + ) + + if env.fp8_pool: + # Neither in-place rotation is quantized, so the pair of this op and + # the context FMHA (which quantizes q/k/v itself) is not a double + # quantization. Compared against the untouched nope slice of the same + # tensor, which is the in-run baseline for un-quantized bf16 data. + for name, values in ( + ("roped q", q_heads[..., QK_NOPE_HEAD_DIM:]), + ("roped k_pe", latent_cache[:, KV_LORA_RANK:]), + ): + frac = _e4m3_roundtrip_fraction(values) + assert frac < MAX_E4M3_ROUNDTRIP_FRACTION, ( + f"{name} looks e4m3-quantized: {frac:.4f} of elements survive " + "an e4m3 round trip unchanged" + ) + baseline = _e4m3_roundtrip_fraction(q_orig_heads[..., :QK_NOPE_HEAD_DIM]) + assert baseline < MAX_E4M3_ROUNDTRIP_FRACTION, ( + f"un-quantized baseline is already {baseline:.4f} — the " + "round-trip check cannot discriminate at this input distribution" + ) + + # 3. Paged cache: rows [cached_s, kv_s) hold [compressed_kv | rope(k_pe)], + # quantized by kv_scale_orig_quant on the fp8 path; the pre-filled cached + # prefix [0, cached_s) is bitwise untouched. + appended = [] + row_start = 0 + for s in range(num_contexts): + rid = request_ids[s] + rows = env.read_rows(layer_idx, rid, cached_lens[s], kv_lens[s]) + appended.append(rows) + row_end = row_start + ctx_new[s] + want_ckv = latent_orig[row_start:row_end, :KV_LORA_RANK] + want_k_pe = ref_k_pe[row_start:row_end] + if env.fp8_pool: + _assert_bytes_equal( + rows[:, :KV_LORA_RANK], + env.quantize(want_ckv), + f"appended compressed_kv (request {rid})", + ) + _assert_e4m3_roped( + rows[:, KV_LORA_RANK:], + env.quantize(want_k_pe), + f"appended k_pe (request {rid})", + ) + else: + torch.testing.assert_close( + rows[:, :KV_LORA_RANK], + want_ckv, + rtol=0.0, + atol=0.0, # dtype-preserving copy: must be bitwise equal + ) + torch.testing.assert_close(rows[:, KV_LORA_RANK:], want_k_pe) + row_start = row_end + for s, rid in enumerate(request_ids): + if cached_lens[s] == 0: + continue + kept = env.read_rows(layer_idx, rid, 0, cached_lens[s]) + # Cached prefix (and every generation sequence's rows): bitwise kept. + _assert_bytes_equal(kept, prefix_rows[s], f"cached prefix (request {rid})") + + # 4. Nothing else in the pool moved: exactly one (C+R)-element row per new + # context token, at the slot its position addresses. This is also what + # pins the page-slab geometry, whose element width the op derives from + # quant_mode alone (1 byte under quant_mode=128). + changed = pool.reshape(-1).view(torch.uint8) != pool_before.reshape(-1).view(torch.uint8) + row_bytes = pool.element_size() * LATENT_SIZE + idx = changed.nonzero().reshape(-1) // row_bytes + got = set( + zip( + (idx // env.tokens_per_block).tolist(), + (idx % env.tokens_per_block).tolist(), + ) + ) + expected = { + env.slot(request_ids[s], cached_lens[s] + i) + for s in range(num_contexts) + for i in range(ctx_new[s]) + } + assert got <= expected, ( + f"op touched pool slots {sorted(got - expected)} outside the " + f"{len(expected)} (page, slot) pairs its inputs address" + ) + # Upper bound, not equality: a written byte that happens to match the byte + # already there is invisible to a snapshot diff. The rows' contents are + # pinned above. + assert idx.numel() <= num_tokens * row_bytes, ( + f"op wrote {idx.numel()} pool bytes, at most {num_tokens * row_bytes} rows' worth expected" + ) + + return _Run( + q=q, + latent_cache=latent_cache, + latent_orig=latent_orig, + ref_k_pe=ref_k_pe, + ref_k_pe_fp32=ref_k_pe_fp32, + rows=torch.cat(appended, dim=0), + positions=positions, + ) + + +def test_bf16_mixed_cached_lengths_layer1() -> None: + """Prefill-like batch (~400 new tokens): cached prefixes of 0 (fresh + prefill), exactly one block, mid-block, and near max_seq_len; addressed + through pool-mapping row 1 of a two-layer pool.""" + torch.manual_seed(0) + env = _MlaCtxEnv(num_layers=2) + try: + cached = [0, 64, 100, 511] + new = [37, 64, 300, 1] + rids = [0, 1, 2, 3] + env.kv_cache_manager.add_dummy_requests( + rids, token_nums=[c + n for c, n in zip(cached, new)] + ) + _run_and_check( + env, + request_ids=rids, + seq_lens=new, + num_contexts=4, + cached_lens=cached, + layer_idx=1, + ) + finally: + env.shutdown() + + +def test_bf16_trailing_generation_seqs_ignored() -> None: + """Mixed batch: two context sequences followed by two generation + sequences. The op must touch only the context sequences' new rows and + index per-seq tensors over [0, num_contexts).""" + torch.manual_seed(1) + env = _MlaCtxEnv() + try: + cached = [128, 3, 200, 77] + new = [40, 60, 1, 1] + rids = [0, 1, 2, 3] + env.kv_cache_manager.add_dummy_requests( + rids, token_nums=[c + n for c, n in zip(cached, new)] + ) + _run_and_check( + env, + request_ids=rids, + seq_lens=new, + num_contexts=2, + cached_lens=cached, + ) + finally: + env.shutdown() + + +def test_bf16_single_token_append() -> None: + """Smallest cached-context case: one sequence, one cached token plus one + new token (decode-like single-row call, appended at position 1).""" + torch.manual_seed(2) + env = _MlaCtxEnv() + try: + env.kv_cache_manager.add_dummy_requests([0], token_nums=[2]) + _run_and_check( + env, + request_ids=[0], + seq_lens=[1], + num_contexts=1, + cached_lens=[1], + ) + finally: + env.shutdown() + + +def test_fp16_mixed_cached_lengths() -> None: + """fp16 activations with an fp16 latent cache over block-crossing + cached/new lengths.""" + torch.manual_seed(3) + env = _MlaCtxEnv(cache_dtype=DataType.HALF) + try: + cached = [65, 640, 1] + new = [63, 128, 6] + rids = [0, 1, 2] + env.kv_cache_manager.add_dummy_requests( + rids, token_nums=[c + n for c, n in zip(cached, new)] + ) + _run_and_check( + env, + request_ids=rids, + seq_lens=new, + num_contexts=3, + cached_lens=cached, + ) + finally: + env.shutdown() + + +def test_bf16_page32_mixed_cached_lengths_layer1() -> None: + """Page size 32 (the engine default). The four cached prefixes place the + first new token at every alignment a 32-token page has: 0 (fresh + prefill), 32 (page 1 slot 0, an exact boundary), 100 (page 3 slot 4, + mid-page) and 511 (page 15 slot 31, a page's last slot). The 300-token + sequence then walks ten pages in one call. Addressed through + pool-mapping row 1 of a two-layer pool.""" + torch.manual_seed(4) + env = _MlaCtxEnv(num_layers=2, tokens_per_block=PAGE32) + try: + cached = [0, 32, 100, 511] + new = [37, 64, 300, 1] + rids = [0, 1, 2, 3] + env.kv_cache_manager.add_dummy_requests( + rids, token_nums=[c + n for c, n in zip(cached, new)] + ) + _run_and_check( + env, + request_ids=rids, + seq_lens=new, + num_contexts=4, + cached_lens=cached, + layer_idx=1, + ) + finally: + env.shutdown() + + +def test_bf16_page32_trailing_generation_seqs_ignored() -> None: + """Page size 32, mixed batch: two context sequences followed by two + generation sequences. Sequence 0's new tokens fill page 3 exactly + (96..127); sequence 1 starts on page 0's last slot and crosses on its + very first token (31..63). The op must still touch only the context + rows and index per-seq tensors over [0, num_contexts).""" + torch.manual_seed(5) + env = _MlaCtxEnv(tokens_per_block=PAGE32) + try: + cached = [96, 31, 200, 77] + new = [32, 33, 1, 1] + rids = [0, 1, 2, 3] + env.kv_cache_manager.add_dummy_requests( + rids, token_nums=[c + n for c, n in zip(cached, new)] + ) + _run_and_check( + env, + request_ids=rids, + seq_lens=new, + num_contexts=2, + cached_lens=cached, + ) + finally: + env.shutdown() + + +def test_fp16_page32_mixed_cached_lengths() -> None: + """fp16 activations and fp16 latent cache at page size 32: a prefix + ending mid-page (33 + 31 closes page 1 exactly), a 16-page prefix + followed by three full pages of new tokens (512 + 96), and a + single-page tail.""" + torch.manual_seed(6) + env = _MlaCtxEnv(cache_dtype=DataType.HALF, tokens_per_block=PAGE32) + try: + cached = [33, 512, 1] + new = [31, 96, 6] + rids = [0, 1, 2] + env.kv_cache_manager.add_dummy_requests( + rids, token_nums=[c + n for c, n in zip(cached, new)] + ) + _run_and_check( + env, + request_ids=rids, + seq_lens=new, + num_contexts=3, + cached_lens=cached, + ) + finally: + env.shutdown() + + +def _fp8_env(max_batch_size: int = 8, orig_quant: Optional[float] = None) -> _MlaCtxEnv: + """DeepSeek-R1-0528 cell over an fp8-e4m3 latent pool: H = 128, page 32, + bf16 activations, single-layer pool.""" + return _MlaCtxEnv( + max_batch_size=max_batch_size, + tokens_per_block=PAGE32, + num_heads=NUM_HEADS_R1, + fp8_pool=True, + orig_quant=orig_quant, + ) + + +def test_fp8_kv_production_config() -> None: + """The production fp8 context call: the KV scaling factor omitted (what + TrtllmAttention.mla_rope_append_paged_kv_assign_q passes) at H = 128 and + page size 32. Four context sequences append T = 402 rows in one call, with + cached prefixes placing the first new token at every alignment a 32-slot + page has — 0 (fresh prefill), 32 (page 1 slot 0), 100 (page 3 slot 4) and + 511 (page 15's last slot) — and the 300-token sequence walking ten + pages.""" + torch.manual_seed(20) + env = _fp8_env() + try: + cached = [0, 32, 100, 511] + new = [37, 64, 300, 1] + rids = [0, 1, 2, 3] + env.kv_cache_manager.add_dummy_requests( + rids, token_nums=[c + n for c, n in zip(cached, new)] + ) + _run_and_check( + env, + request_ids=rids, + seq_lens=new, + num_contexts=4, + cached_lens=cached, + ) + finally: + env.shutdown() + + +def test_fp8_kv_explicit_unit_scale() -> None: + """Explicit scaling factor 1.0 (a checkpoint whose k_scale/v_scale are the + production 1.0, passed as a real tensor), then the wrong-variant controls + that measure what the e4m3 gate on the appended k_pe half discriminates. + Every variant below is a plausible mis-derivation of the same bytes.""" + torch.manual_seed(21) + env = _fp8_env(orig_quant=1.0) + try: + cached = [0, 31] + new = [33, 40] + rids = [0, 1] + env.kv_cache_manager.add_dummy_requests( + rids, token_nums=[c + n for c, n in zip(cached, new)] + ) + run = _run_and_check( + env, + request_ids=rids, + seq_lens=new, + num_contexts=2, + cached_lens=cached, + ) + # Controls. Every appended byte above came back bit-exact, which is + # only meaningful if the same comparison moves for a wrong mirror. + # Measured on sm_100 at this case — the whole appended k_pe block + # (73 rows x 64 elements = 4672 bytes, bit-exact allowance 4) against + # four plausible mis-derivations, counting bytes differing and the + # elements the one-ulp tolerance itself rejects: + # + # mirror inexact bytes outside tolerance + # un-roped k_pe 3278 (819x) 2609 + # rope at position + 1 1799 (450x) 1133 + # the previous token's row 4607 (1152x) 4469 + # no bf16 rounding before e4m3 152 (38x) 1 + # (the correct mirror) 0 0 + # + # The first three blow both halves of the gate. The fourth — + # quantizing the unrounded fp32 RoPE result instead of the bf16 one + # the kernel rounds to first — differs from the truth by a single + # rounding step, so it is invisible to the tolerance by construction + # (1 element out of 4672 crosses it, which is not something to rely + # on) and only the bit-exact-fraction floor catches it, at 38x the + # allowance. Both halves are asserted below to actually fire. + got = run.rows[:, KV_LORA_RANK:] + allowed = max(1, int(E4M3_MAX_INEXACT_FRACTION * got.numel())) + shifted = [p + 1 for p in run.positions] + blunt = { + "un-roped k_pe": env.quantize(run.latent_orig[:, KV_LORA_RANK:]), + "rope at position + 1": env.quantize( + env.rope_ref(run.latent_orig[:, KV_LORA_RANK:], shifted) + ), + "the previous token's row": env.quantize(run.ref_k_pe.roll(1, dims=0)), + } + for name, mirror in blunt.items(): + inexact = _inexact_bytes(got, mirror) + assert inexact > 100 * allowed, ( + f"control '{name}' moved only {inexact}/{got.numel()} bytes — " + "the bit-exact-fraction floor would not have caught it" + ) + assert _fails_e4m3_tolerance(got, mirror), ( + f"control '{name}' stayed inside the one-ulp tolerance — that " + "half of the gate would not have caught it" + ) + # The rounding-level control: the fraction floor has to carry it alone. + no_bf16 = (run.ref_k_pe_fp32 * env.write_scale).to(torch.float8_e4m3fn) + inexact = _inexact_bytes(got, no_bf16) + assert inexact > 8 * allowed, ( + f"control 'no bf16 rounding before e4m3' moved only {inexact}/" + f"{got.numel()} bytes — the bit-exact-fraction floor would not " + "have caught it" + ) + finally: + env.shutdown() + + +def test_fp8_kv_scale_factor_sweep() -> None: + """KV scaling factors other than the production 1.0. 1/1.5 is not a power + of two, so the write-side multiply is a real rescale rather than an + exponent shift; 2.0 covers a factor above 1. At 1/1.5 the sweep also runs + the control that pins the order of the two operations — the kernel rounds + the fp32 RoPE result to bf16 and *then* multiplies by the scale in fp32, + so e4m3(bf16(rope * w)) is a different tensor.""" + for scale in (1.0 / 1.5, 0.5, 0.25, 2.0): + torch.manual_seed(22) + env = _fp8_env(orig_quant=scale) + try: + cached = [0, 31] + new = [33, 40] + rids = [0, 1] + env.kv_cache_manager.add_dummy_requests( + rids, token_nums=[c + n for c, n in zip(cached, new)] + ) + run = _run_and_check( + env, + request_ids=rids, + seq_lens=new, + num_contexts=2, + cached_lens=cached, + ) + got = run.rows[:, KV_LORA_RANK:] + allowed = max(1, int(E4M3_MAX_INEXACT_FRACTION * got.numel())) + if scale == 1.0 / 1.5: + # Scaling before the bf16 rounding instead of after is exact + # for a power of two, so only the non-power-of-two scale + # separates the two orders. Measured: 49/4672 bytes differ + # (1.05%), against a bit-exact allowance of 4. Exact at + # 0.5 / 0.25 / 2.0, which is why the control runs only here. + scale_first = ( + (run.ref_k_pe_fp32 * env.write_scale).to(torch.bfloat16).to(torch.float8_e4m3fn) + ) + inexact = _inexact_bytes(got, scale_first) + assert inexact > 4 * allowed, ( + f"control 'scale before the bf16 rounding' moved only " + f"{inexact}/{got.numel()} bytes — the bit-exact-fraction " + "floor would not have caught it" + ) + # And the scale is really applied: a mirror that ignores it (the + # w = 1.0 quantization) blows both halves of the gate at every + # scale swept here. + unscaled = run.ref_k_pe.float().to(torch.float8_e4m3fn) + inexact = _inexact_bytes(got, unscaled) + assert inexact > 100 * allowed and _fails_e4m3_tolerance(got, unscaled), ( + f"a mirror ignoring kv_scale_orig_quant={scale} moved only " + f"{inexact}/{got.numel()} bytes or stayed inside the one-ulp " + "tolerance — the scale is not discriminated" + ) + finally: + env.shutdown() + + +def test_fp8_kv_trailing_generation_seqs_ignored() -> None: + """fp8 mixed batch: two context sequences followed by two generation + sequences, all with cached prefixes. Sequence 0's new tokens fill page 3 + exactly (96..127); sequence 1 starts on page 0's last slot and crosses on + its very first token (31..63). The op must touch only the context rows, + index per-seq tensors over [0, num_contexts), and leave both generation + sequences' pages bitwise untouched.""" + torch.manual_seed(23) + env = _fp8_env(orig_quant=1.0 / 1.5) + try: + cached = [96, 31, 200, 77] + new = [32, 33, 1, 1] + rids = [0, 1, 2, 3] + env.kv_cache_manager.add_dummy_requests( + rids, token_nums=[c + n for c, n in zip(cached, new)] + ) + _run_and_check( + env, + request_ids=rids, + seq_lens=new, + num_contexts=2, + cached_lens=cached, + ) + finally: + env.shutdown() + + +def test_fp8_kv_scale_does_not_reach_q() -> None: + """The KV scaling factor is a write-side quantization scale for the pool + alone. Two runs over identical inputs (same seed, same batch) at + kv_scale_orig_quant 1.0 and 0.25 must produce bitwise-identical q and + latent_cache, and appended rows that differ.""" + runs = [] + for scale in (1.0, 0.25): + torch.manual_seed(24) + env = _fp8_env(orig_quant=scale) + try: + cached = [17, 0] + new = [20, 45] + rids = [0, 1] + env.kv_cache_manager.add_dummy_requests( + rids, token_nums=[c + n for c, n in zip(cached, new)] + ) + run = _run_and_check( + env, + request_ids=rids, + seq_lens=new, + num_contexts=2, + cached_lens=cached, + ) + runs.append((run.q.clone(), run.latent_cache.clone(), run.rows.clone())) + finally: + env.shutdown() + (q_a, latent_a, rows_a), (q_b, latent_b, rows_b) = runs + _assert_bytes_equal(q_a, q_b, "roped q across kv_scale_orig_quant") + _assert_bytes_equal(latent_a, latent_b, "roped latent_cache across kv_scale_orig_quant") + # And the appended rows really did move, so the comparison above is not + # vacuously true of a call that ignored the scale everywhere. + assert _inexact_bytes(rows_a, rows_b) > rows_a.numel() // 2, ( + "appended rows barely changed between scale 1.0 and 0.25 — the scale " + "does not reach the pool either" + ) + + +def test_fp8_kv_rejects_int8_kv_cache_quant_mode() -> None: + """A quant_mode asking for a non-fp8 quantized KV cache is rejected, not + silently treated as fp8 or as high precision.""" + torch.manual_seed(25) + env = _fp8_env(orig_quant=1.0) + try: + env.kv_cache_manager.add_dummy_requests([0], token_nums=[8]) + metadata = env.prepare_metadata([0], [8], 1, [0]) + q = torch.randn(8, NUM_HEADS_R1 * QK_HEAD_DIM, dtype=torch.bfloat16, device="cuda") + latent_cache = torch.randn(8, LATENT_SIZE, dtype=torch.bfloat16, device="cuda") + block_offsets = metadata.kv_cache_block_offsets + pool_pointers = env.kv_cache_manager.kv_cache_pool_pointers + pool_mapping = env.kv_cache_manager.kv_cache_pool_mapping + assert block_offsets is not None + assert pool_pointers is not None and pool_mapping is not None + raised = "" + try: + mla_rope_append_paged_kv_assign_q( + q, + latent_cache, + 1, + metadata.ctx_cached_token_indptr, + metadata.ctx_kv_indptr, + int(metadata.max_ctx_seq_len), + env.rotary_cos_sin, + NUM_HEADS_R1, + QK_NOPE_HEAD_DIM, + QK_ROPE_HEAD_DIM, + KV_LORA_RANK, + block_offsets, + pool_pointers, + pool_mapping, + env.kv_scale_orig_quant, + 0, # residual_dim + 0, # layer_idx + env.tokens_per_block, + MAX_SEQ_LEN, + 1, # beam_width + QUANT_MODE_INT8_KV_CACHE, + ) + torch.cuda.synchronize() + except RuntimeError as exc: + raised = str(exc) + # rc26 added NVFP4 latent pools, so the rejection message now + # enumerates two accepted formats. int8 is still rejected -- + # what this test certifies -- only the wording widened. + assert "Only FP8 and NVFP4 KV caches are supported for now" in raised, ( + f"quant_mode={QUANT_MODE_INT8_KV_CACHE} was not rejected as expected; got: {raised!r}" + ) + finally: + env.shutdown() diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation.md b/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation.md new file mode 100644 index 000000000000..e234b5a642ef --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation.md @@ -0,0 +1,427 @@ +--- +receipts: + sm_100: {status: passed, trtllm: 1.3.0rc21} + sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 22} +--- + +# mla_rope_generation + +**Wraps** `torch.ops.trtllm.mla_rope_generation` (one call). + +## Semantics + +The generation-phase (decode) preprocessing step of MLA absorbed attention, +over one prepared batch. Two configurations are certified, and they differ +in what the call writes — read the one you are on: + +- **bf16 latent pool** (`quant_mode=0`, every fp8/scale tensor `None`): the + roped q lands in `fused_q`'s tail and the appended cache row is bf16. +- **fp8-e4m3 latent pool** (`quant_mode=128`): `fused_q` is **not written at + all**; the roped q lands quantized in `quant_q_buffer`, the appended cache + row lands quantized, and the two decode-FMHA scale buffers are filled. + +Both are certified with bf16 activations (`fused_q`, `q_pe`, +`latent_cache`), `rope_append=True`, `beam_width=1`, +`block_ids_per_seq=None`, `out_scale=None` and no helix parallelism, at +`predicted_tokens_per_seq` (`P` below) 1, 2, 3 and 4. + +The batch has `num_contexts` context-phase sequences first, then `G` +generation-phase sequences. The op processes **only the generation +sequences**. Each of them contributes `P` new query tokens — `P = 1` is an +ordinary decode step, `P > 1` the MTP shape, where `P = max_draft_len + 1`. +`fused_q`/`q_pe`/`latent_cache`/`quant_q_buffer` therefore have `G * P` +rows, **token-major within a sequence**: row `n` is generation sequence +`g = n // P`'s `t = n % P`-th new token. Per-sequence tensors +(`sequence_length`, `host_past_key_value_lengths`, `host_context_lengths`, +`kv_cache_block_offsets`) cover the whole batch in order and stay one entry +per sequence; the op indexes them at `num_contexts + g`. + +**Every row gets its own rope angle.** With +`L_g = sequence_length[num_contexts + g]` the total KV length of generation +sequence `g` *including all `P` of its new tokens*, row `(g, t)` sits at +0-based absolute position + +``` +pos(g, t) = L_g - P + t # t = 0 .. P-1, consecutive, ending at L_g - 1 +``` + +so the `P` tokens of one sequence occupy `P` consecutive positions and each +is rotated at its own. At `P = 1` this is the familiar `L_g - 1`. The +position comes from the **device** `sequence_length` tensor: running the +same prepared step with `host_past_key_value_lengths` perturbed to +`sequence_length - 1` left every position, every pool slot and +`cu_kv_seqlens` unchanged (observed on the fp8 path at `P = 3`). + +**On the bf16 path**, one call then computes, with `R = qk_rope_head_dim` +and `C = kv_lora_rank` (the fp8 path replaces effects 1 and 2 — see below): + +``` +# GPT-J (interleaved-pair) RoPE, fp32 math, one rounding to bf16. +# (cos_d, sin_d) = rotary_cos_sin row pos(g,t), pair d, d in [0, R/2) +rope(x, p)[2d] = x[2d] * cos_d - x[2d+1] * sin_d +rope(x, p)[2d+1] = x[2d] * sin_d + x[2d+1] * cos_d + +# for every row n = g*P + t: +# 1. q RoPE into the fused-q tail (per head h): +fused_q[n, h, C:] = rope(q_pe[n, h, :], pos(g,t)) +# 2. latent-cache append, at page/slot of position pos(g,t) in seq g's blocks: +cache_row(g, pos(g,t)) = concat(latent_cache[n, :C], # bitwise copy + rope(latent_cache[n, C:], pos(g,t))) # k_pe rotated +# 3. decode-FMHA scheduler buffers (trtllm-gen MQA layout): +cu_q_seqlens[0:G+1] = arange(G+1) * num_heads * P +cu_kv_seqlens[0:G+1] = [0, cumsum(L_g over generation seqs)] +fmha_scheduler_counter[0] = 0 +``` + +Fusion boundary: q RoPE, k_pe RoPE, paged-cache append and scheduler-buffer +fill happen inside the call. The caller still owns the up/down projections, +the absorbed-q BMM that fills `fused_q[..., :C]`, and the subsequent +latent-space attention. `q_pe` and `latent_cache` are inputs only: despite +the mutable schema annotation on `q_pe` neither is modified (observed). + +On the bf16 path `fused_q[..., :C]` is bitwise untouched and never read, so +the absorbed-q BMM and this call may run concurrently. **On the fp8 path +they may not** — see below. + +### fp8-e4m3 latent pool (`quant_mode=128`) + +`quant_mode=128` is `QuantMode`'s fp8-KV-cache bit. It switches the +**latent pool's** element type to fp8-e4m3, one byte per element, while +`fused_q`, `q_pe` and `latent_cache` all stay bf16. Certified at the +DeepSeek-R1-0528 cell — `H = 128`, `tokens_per_block = 32`, +`C/R/nope/v = 512/64/128/128`, `q_lora_rank = 1536`, single-layer pool — at +`P` 1, 2, 3 and 4, with the KV scaling factor omitted (the production call) +and explicit 1.0 / 1.5 / 2.0 swept beside it, at `q_scaling` 1.0 and +DeepSeek-R1's YaRN value `1/mscale^2 = 0.5336594`. + +Two independent fp32 scalars steer it, each with its own role and neither +derived from the other: + +``` +w = kv_scale_orig_quant[0] # 1.0 when that tensor is None +r = kv_scale_quant_orig[0] # 1.0 when that tensor is None +e4m3(x) = round-to-nearest cast of an fp32 x to float8_e4m3fn +``` + +One call then computes, with `rope` as above and `bf16(.)` the rounding to +the activation dtype that happens before quantization, for every row +`n = g*P + t` at `p = pos(g,t)`: + +``` +# 1. q RoPE, quantized. fused_q is READ, never written: +quant_q_buffer[n, h, :C] = e4m3(fused_q[n, h, :C] * w) +quant_q_buffer[n, h, C:] = e4m3(bf16(rope(q_pe[n, h, :], p)) * w) +# 2. latent-cache append, quantized, at page/slot of position p: +cache_row(g, p) = e4m3(concat(latent_cache[n, :C], + bf16(rope(latent_cache[n, C:], p))) * w) +# 3. decode-FMHA scales (a per-batch scalar pair — not per sequence, and +# P does not enter them): +x = r*r / (q_scaling * sqrt(qk_nope_head_dim + qk_rope_head_dim)) +mla_bmm1_scale = [x, x * log2(e)] +mla_bmm2_scale = [r] +# 4. scheduler buffers: exactly as on the bf16 path +``` + +`w` is a **per-tensor static scale**, not a block or dynamic one: every +appended row and every quantized q element is the same multiply-and-round, +bit-exact against that mirror at each scale tested. + +Fusion boundary on this path: the same RoPE / append / scheduler fill, plus +the q quantization and the two scale derivations. The caller still owns the +absorbed-q BMM — and now must have **finished** it before this call, since +`fused_q[..., :C]` is read here to build `quant_q_buffer`'s head. (The +engine's own MLA module enforces exactly that by dropping the auxiliary +stream when the KV cache is fp8.) + +**Why `r` appears in both scales.** The fp8 MLA generation FMHA — the +sibling op `trtllm::attention` called with +`attention_input_type = generation_only`, `is_mla_enable=True` — takes its +query from `quant_q_buffer` rather than from the bf16 fused q, takes its +BMM1 scale from `mla_bmm1_scale[1]` (the log2-domain copy; `[0]` is inert +there) and its BMM2 scale from `mla_bmm2_scale[0]`, and reads neither KV +scale tensor. Dequantization is therefore the producer's job: `r*r` undoes +the `w` on the query and on `K`, and `r` undoes the `w` on `V = K[:, :C]`, +leaving the plain softmax scale `1 / (q_scaling * sqrt(nope + R))`. That +cancellation is only correct when `r = 1/w`; this op does not check it, and +a non-reciprocal pair produces a silently mis-scaled decode (pinned here by +running `w = 0.25` against `r = 3.0` and observing each drive its own half). +The consuming call's behaviour above is that op's fact, established +separately — this entry's test gates the values written, not their use. + +## Signature + +```python +def mla_rope_generation( + fused_q: torch.Tensor, + q_pe: torch.Tensor, + latent_cache: torch.Tensor, + rotary_cos_sin: Optional[torch.Tensor], + cu_q_seqlens: torch.Tensor, + cu_kv_seqlens: torch.Tensor, + fmha_scheduler_counter: torch.Tensor, + mla_bmm1_scale: Optional[torch.Tensor], + mla_bmm2_scale: Optional[torch.Tensor], + quant_q_buffer: Optional[torch.Tensor], + sequence_length: torch.Tensor, + host_past_key_value_lengths: torch.Tensor, + host_context_lengths: torch.Tensor, + num_contexts: int, + kv_cache_block_offsets: Optional[torch.Tensor], + host_kv_cache_pool_pointers: Optional[torch.Tensor], + host_kv_cache_pool_mapping: Optional[torch.Tensor], + kv_scale_orig_quant: Optional[torch.Tensor], + kv_scale_quant_orig: Optional[torch.Tensor], + out_scale: Optional[torch.Tensor], + block_ids_per_seq: Optional[torch.Tensor], + helix_tensor_params: List[Optional[torch.Tensor]], + predicted_tokens_per_seq: int, + layer_idx: int, + num_heads: int, + num_kv_heads: int, + head_size: int, + tokens_per_block: int, + attention_window_size: int, + beam_width: int, + quant_mode: int, + q_scaling: float, + q_lora_rank: int, + kv_lora_rank: int, + qk_nope_head_dim: int, + qk_rope_head_dim: int, + v_head_dim: int, + rope_append: bool, +) -> None +``` + +With `G` = generation sequences, `P` = `predicted_tokens_per_seq`, +`S` = total sequences (`num_contexts + G`), `H` = `num_heads`, +`C` = `kv_lora_rank`, `R` = `qk_rope_head_dim`, `D = C + R` +(= `head_size`). The four per-token tensors have `G * P` rows, ordered +token-major within a sequence (row `n = g*P + t`); everything else is one +entry per sequence and does **not** scale with `P`: + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `fused_q` | `[G*P, H, D]`; bf16 pool: only `[..., C:]` written, `[..., :C]` untouched and unread. fp8 pool: **entirely read-only**, and `[..., :C]` must already hold the absorbed q | bf16 | contiguous | CUDA | +| `q_pe` | `[G*P, H, R]`; read only | bf16 | last dim contiguous; strided head-dim views (e.g. a slice of packed `[G*P, H, nope+R]` q) work | CUDA | +| `latent_cache` | `[G*P, D]` = per-token `[compressed_kv \| k_pe]`; read only | bf16 | contiguous | CUDA | +| `rotary_cos_sin` | `[1, max_pos * R * 2]`: per position `R` fp32 `(cos, sin)` pairs, second `R/2` a duplicate of the first (duplicated/`duplicate_data=True` layout; only pairs `[0, R/2)` are read per position) | fp32 | contiguous | CUDA | +| `cu_q_seqlens` | `[>= G+1]`; exactly `G+1` entries written whatever `P` is (entries past `G+1` are left alone — observed) | int32 | contiguous | CUDA | +| `cu_kv_seqlens` | `[>= G+1]`; same, `G+1` entries | int32 | contiguous | CUDA | +| `fmha_scheduler_counter` | `[1]`; written (zeroed) | uint32 | — | CUDA | +| `mla_bmm1_scale` | bf16 pool: `None`. fp8 pool: `[2]`, written with `[x, x*log2(e)]`; may be `None`, and is then simply not written | fp32 | contiguous | CUDA | +| `mla_bmm2_scale` | bf16 pool: `None`. fp8 pool: `[1]`, written with `[kv_scale_quant_orig]`; may be `None`, and is then simply not written | fp32 | contiguous | CUDA | +| `quant_q_buffer` | bf16 pool: `None`. fp8 pool: `[G*P, H, D]`, **required** — one e4m3 byte per element, fully overwritten | uint8 (a `float8_e4m3fn` view of the same bytes behaves identically) | contiguous | CUDA | +| `sequence_length` | `[S]`; per-seq total KV length incl. **all `P`** of this step's tokens. This is the tensor the `P` positions and `cu_kv_seqlens` are derived from | int32 | contiguous | CUDA | +| `host_past_key_value_lengths` | `[S]`; pass the same values as `sequence_length` (total KV length, **not** past-only, despite the name). Perturbing it to `sequence_length - 1` changed no observable effect at the certified fp8 cell, so it is `sequence_length` that steers this op; larger divergences were not measured | int32 | contiguous | CPU | +| `host_context_lengths` | `[S]`; per-seq prompt length | int32 | contiguous | CPU | +| `num_contexts` | leading context-phase sequence count | Python int | — | — | +| `kv_cache_block_offsets` | `[1, >= S, 2, max_blocks_per_seq]`; raw block ids per seq (K and V rows identical, kv_factor=1 pool) | int32 | contiguous | CUDA | +| `host_kv_cache_pool_pointers` | `[num_pools, 2]`: (primary ptr, secondary ptr=0) | int64 | contiguous | CPU | +| `host_kv_cache_pool_mapping` | `[num_layers, 2]`: (pool index, layer-within-pool) row per layer | int32 | contiguous | CPU | +| `kv_scale_orig_quant` | `[1]`, the **write-side** factor `w`: everything quantized by this call is multiplied by it. `None` = 1.0 (what the engine's own MLA call site passes). Ignored on the bf16 path | fp32 | contiguous | CUDA | +| `kv_scale_quant_orig` | `[1]`, the **read-side** factor `r`: folded into `mla_bmm1_scale` (as `r*r`) and copied into `mla_bmm2_scale`. `None` = 1.0. Not used for anything else, and not derived from `kv_scale_orig_quant`. Ignored on the bf16 path | fp32 | contiguous | CUDA | +| `out_scale` | `None`; fp8-output scale otherwise (not certified) | — | — | — | +| `block_ids_per_seq` | `None`; flash-MLA layout otherwise (not certified) | — | — | — | +| `helix_tensor_params` | `[None, None]`; helix position offsets / inactive-rank mask otherwise (not certified) | Python list | — | — | +| `predicted_tokens_per_seq` | `P`, the query tokens each generation sequence carries. `1` (ordinary decode) and `2`, `3`, `4` (the MTP path, `max_draft_len + 1`) certified. It is a single scalar for the whole batch — a ragged per-sequence draft length cannot be expressed here. Values above 4 untested | Python int | — | — | +| `layer_idx` | row into `host_kv_cache_pool_mapping` | Python int | — | — | +| `num_heads` | query heads `H`; 16 and 128 certified | Python int | — | — | +| `num_kv_heads` | `1` (MLA latent cache) | Python int | — | — | +| `head_size` | `D = kv_lora_rank + qk_rope_head_dim` | Python int | — | — | +| `tokens_per_block` | pool page size; 32 and 64 certified on the bf16 path, 32 on the fp8 path (32 is what a default `KvCacheConfig` produces) | Python int | — | — | +| `attention_window_size` | `>= max total KV length`; smaller values imply cyclic-cache addressing (not certified) | Python int | — | — | +| `beam_width` | `1` | Python int | — | — | +| `quant_mode` | `0` (bf16 latent pool) and `128` (fp8-e4m3 latent pool) certified | Python int | — | — | +| `q_scaling` | softmax-scale divisor baked into `mla_bmm1_scale`; live on the fp8 path (1.0 and 0.5336594 certified), inert on the bf16 path, which writes no scale at all (bitwise-identical outputs at 1.0 and 7.5, observed) | Python float | — | — | +| `q_lora_rank` | unused by both certified paths (`0` and `1536` both fine) | Python int | — | — | +| `kv_lora_rank` | `C` | Python int | — | — | +| `qk_nope_head_dim`, `qk_rope_head_dim` | model MLA dims (`R` must match `q_pe`/table; their sum is the `sqrt` in `mla_bmm1_scale`) | Python int | — | — | +| `v_head_dim` | model v head dim; not read on either certified path | Python int | — | — | +| `rope_append` | `True` certified (rotate k_pe + append to cache) | Python bool | — | — | +| returns | — (mutates `cu_q_seqlens`, `cu_kv_seqlens`, `fmha_scheduler_counter`, the paged pool, plus `fused_q`'s tail on the bf16 path or `quant_q_buffer`/`mla_bmm1_scale`/`mla_bmm2_scale` on the fp8 path) | — | — | — | + +## Metadata consumed + +No thread-local or registered-layer state: all inputs are explicit +arguments. The length/addressing tensors are exactly what a prepared +`TrtllmAttentionMetadata` and its `KVCacheManager` expose: +`kv_lens_cuda_runtime` → `sequence_length`, `kv_lens_runtime` → +`host_past_key_value_lengths`, `prompt_lens_cpu_runtime` → +`host_context_lengths`, `metadata.kv_cache_block_offsets`, +`manager.kv_cache_pool_pointers`, `manager.kv_cache_pool_mapping`. + +## Preconditions + +- The paged pool addressed by `host_kv_cache_pool_pointers` + + `host_kv_cache_pool_mapping[layer_idx]` is an MLA latent cache: + kv_factor 1, one kv head, row width `D`, page size `tokens_per_block`. + Per page the layout is `[tokens_per_block, D]` and the page slab is + `tokens_per_block * D` **elements**, whose byte width the op derives from + `quant_mode` alone: 2 bytes (bf16) at `quant_mode=0`, 1 byte (e4m3) at + `quant_mode=128`. The pool's real element type must match — a + `KVCacheManager` built with `DataType.BF16` and `DataType.FP8` + respectively is what was certified. +- Every generation sequence has block capacity for `L_g` tokens **before** + the call (block offsets in `kv_cache_block_offsets` cover + `ceil(L_g / tokens_per_block)` pages); with a `KVCacheManager`, call + `impl.add_token(request_id)` **`P` times** per sequence per step, then + prepare the metadata. The `P` rows of a sequence may straddle a page + boundary; the op follows the block-offset row across it. +- `sequence_length` / `host_past_key_value_lengths` hold `L_g` (past + all + `P` new tokens), and `L_g <= attention_window_size`, + `L_g - 1 < max_pos` of the table, `L_g >= P`. +- Batch order: context sequences first; `fused_q`/`q_pe`/`latent_cache` + contain generation tokens only, `G * P` rows in generation-sequence order + with a sequence's `P` tokens contiguous and in position order. +- The `G * P` cache rows written by one call must address **distinct + physical slots**. This holds automatically for a `KVCacheManager`-allocated + batch (each row goes to its own position, distinct sequences hold distinct + pages), and it is the reason `P > 1` needs no extra care from the caller. If + a caller-built block-offset table aliases them onto one slot, the writes + race and leave torn cache rows, nondeterministically and with no error — + reproduced here on purpose as a control, at `P = 2` with 8 sequences + aliased onto one page: 6 of 6 identical armed calls left different pool + images, while the un-aliased geometry repeated bitwise in the same + process. +- `rotary_cos_sin` uses the duplicated `(cos, sin)`-pair layout above + (`RopeParams(..., duplicate_data=True).create_rope_const_params()` + produces it; MLA models set `duplicate_data=True`). A table with + cos=1/sin=0 makes both RoPEs identity (pre-rotated inputs pass through). +- `cu_q_seqlens`/`cu_kv_seqlens`/`fmha_scheduler_counter` are + caller-allocated with the dtypes above; contents need not be initialized. +- **bf16 path only:** `quant_mode=0` and every optional fp8/scale tensor + `None`. +- **fp8 path only** (`quant_mode=128`): + - `quant_q_buffer` must be allocated `[G*P, H, D]` at one byte per element. + It is **not** presence-checked: passing `None` is not rejected, the + kernel launches anyway and the run dies with `CUDA error: an illegal + memory access was encountered` (observed in a scratch probe at this + configuration; the CUDA context is lost, the process aborts, and the KV + manager's own `release_pools` fails on the way out). Its previous + contents are irrelevant — the call overwrites every byte. + - `fused_q[..., :C]` must already hold the absorbed q **when the call is + issued**: this path reads it. Overlapping the absorbed-q BMM with this + call on a second stream is a race here, unlike on the bf16 path. + - `mla_bmm1_scale` `[2]` fp32 and `mla_bmm2_scale` `[1]` fp32 may be + `None` — accepted, and the buffer is then simply not written (observed; + the rest of the outputs are unaffected). A decode consuming those scales + still needs them, so in practice both are allocated. + - `kv_scale_orig_quant` and `kv_scale_quant_orig` are `[1]` fp32 CUDA + tensors or `None` (= 1.0). They are used independently; pass reciprocals + (`kv_scale_orig_quant = 1/s`, `kv_scale_quant_orig = s`) unless a + deliberately asymmetric round trip is what you want. `None` for both is + the production configuration and is what the engine's own MLA call site + passes. + - Input magnitudes must survive e4m3: `|value * w|` below `2**-10` + (~9.8e-4) rounds to zero, and 448 is the largest finite e4m3 value (what + the kernel does above it was not measured). The certified runs are + unit-normal activations at `w <= 1`. + +## Notes + +- Schema annotation is wrong in both directions: it marks `fused_q` and + `q_pe` mutable (`Tensor(a!)`), but `q_pe` is never modified — and on the + fp8 path neither is `fused_q` — while `cu_q_seqlens`, `cu_kv_seqlens`, + `fmha_scheduler_counter`, `quant_q_buffer`, `mla_bmm1_scale` and + `mla_bmm2_scale` — all filled by the call — carry no annotation at all. Do + not rely on the schema's alias info (e.g. under torch.compile + functionalization). +- Precision, bf16 pool (sm_100): the `compressed_kv` copy is bitwise. Both + RoPE outputs match an fp32 reference rounded once to bf16 to within **one + bf16 ulp**, and are bitwise on all but ~1.5e-5 of elements: the kernel and + a torch fp32 reference evaluate `x*cos ∓ y*sin` in different orders, so + where the correctly-rounded fp32 result lands exactly on a bf16 rounding + midpoint the two round to adjacent values. Measured on a 64-sequence + decode step: 1 such element in 65 536 rotated `q_pe` values and 0 in 4096 + rotated `k_pe` values, with identical counts at page sizes 32 and 64 — it + is arithmetic, not addressing. Default `assert_close` tolerances absorb it + with margin, which is what the test gates on. +- Precision, fp8 pool (sm_100): **everything came back bit-exact**, both the + pure copies (`quant_q_buffer`'s absorbed-q head, the appended + `compressed_kv` half) and the roped halves. e4m3's 3-bit mantissa absorbs + the evaluation-order difference the bf16 pool pays one ulp for, but the + double rounding is real and a mirror must reproduce it: the kernel rounds + the fp32 RoPE result to **bf16 first** and quantizes that. Quantizing the + fp32 RoPE result directly disagrees with the kernel on ~3% of the roped q + elements (and up to 9% of the much smaller per-row `k_pe` samples). The + test still gates the roped halves at one e4m3 ulp + (`2**-3` relative) plus a bit-exact-fraction floor rather than bitwise, + because the tie it covers is arithmetic; neither half has been approached. +- The fp8 scales are exact enough to gate tightly: `mla_bmm2_scale` is a + bitwise copy of `kv_scale_quant_orig`, not a computation — a scratch probe + passing 0.7 (whose fp32 value is 0.699999988079071) got that exact fp32 + number back, and the test gates the copy bitwise at 1.0 / 1.5 / 2.0 / 3.0. + `mla_bmm1_scale` matches a double-precision reference rounded once to fp32 + to within 7.2e-8 relative (0.6 fp32 ulp) across the swept grid, exactly on + several cases. +- Pool writes are exactly bounded: `P` `D`-byte rows per generation + sequence, at the slots its `P` positions address, and nothing else in the + pool moves — checked by a whole-pool byte snapshot on every fp8 case, at + `P` 1 through 4, 64-sequence batches included. +- Re-running the same prepared step overwrites the same `P` cache slots + (positions are derived from `sequence_length`), so a repeated call is + idempotent, not double-appending — and on the fp8 path the second run + reproduces `quant_q_buffer` and the pool bitwise. This holds at `P > 1` + too: the rows do not advance by `P` on the second call. +- The scheduler buffers feed the trtllm-gen decode MQA: `cu_q_seqlens` is + in units of q rows (`H` rows per generation *token*, so `H * P` per + generation sequence), `cu_kv_seqlens` in tokens (the cumulative + `sequence_length` over the generation sequences, which already includes + all `P` new tokens). They are filled identically on both paths, and both + fills were observed at `P` 1, 2, 3 and 4 — `cu_q_seqlens[i] = i * H * P` + exactly, and no entry past index `G` is written. +- Page size (sm_100): certified at `tokens_per_block` 32 and 64 on the bf16 + path, 32 on the fp8 path. 32 is what a `KvCacheConfig` at its defaults + hands the op; 64 is what a tuned MLA target opts into. At 32 the appended + rows were placed at every alignment a 32-slot page has — position 31 (a + page's last slot), 32 and 96 (slot 0 of a fresh page after one and three + full pages), 64 (after two), and mid-page — including 64-sequence batches + run for two consecutive steps in which sequences cross into a new page + between the steps and others open a fresh page on the first, and mixed + batches whose leading context sequence is skipped. At `P > 1` the same + sweep covers a sequence whose `P` rows **straddle** the boundary, at every + split from 1 row in the old page to `P-1` (`P` = 2, 3 and 4, fp8 and bf16 + pools) — a shape unreachable at `P = 1`, where one call writes one slot + per sequence. Every effect held at the same gates as at page 64. Other + page sizes are untested rather than known-bad. +- `q_pe` in production is a strided slice of the packed q tensor + (`[G*P, H, nope+R]` split); the kernel handles that stride — certified for + both contiguous and packed-slice layouts, on both paths, at `P` 1-4. +- The `rotary_cos_sin` *content* is not a certified axis: what the op is + gated on is that it applies the table's `(cos, sin)` pairs, and both + certified configurations ran the unscaled theta-10000 table. A YaRN-scaled + table of the same layout has not been run through this op here. +- Sibling ops exist for adjacent MLA roles: + `trtllm::mla_rope_append_paged_kv_assign_q` (context-phase RoPE + cache + append), `trtllm::attention` / `trtllm::mla_custom_op_inplace` (the + attention core that consumes `fused_q` or `quant_q_buffer` and these + scheduler buffers), and `trtllm::merge_chunked_attention_for_mla`. +- `predicted_tokens_per_seq` is the *only* way a multi-token generation + query reaches this op: its signature carries no spec-decoding mask, tree + or draft-token argument of any kind, so there is nothing else for an MTP + caller to set here. (What the resulting query block is then allowed to + attend to is the attention op's fact, not this one's.) +- Not certified on either path: `predicted_tokens_per_seq` above 4, + `block_ids_per_seq` (the flash-MLA layout), helix, `out_scale` (fp8 + attention output), `attention_window_size` below the KV length, and + multi-layer pools. Not certified on the fp8 path specifically: + `tokens_per_block = 64`, head counts other than 128, and activation dtypes + other than bf16. + + +## rc26 additions + +Parameters that did not exist in 1.3.0rc21. Every value this entry certifies +reproduces the op's pre-rc26 behaviour, and matches what the in-tree caller +passes on the same path. + +| Parameter | Certified value | Why | +|---|---|---| +| `kv_cache_scale_orig_quant` | `None` | With no value the op uses `kv_scale_orig_quant` for the cache scale (`dsv3RopeOp.cpp`), which is exactly what it did before this parameter existed. | +| `residual_dim` | `0` | Must be `0` or `rope_size`; non-zero requires an FP4 KV pool, and this entry certifies bf16 and fp8-e4m3 pools. | +| `kv_norm_weight` / `kv_norm_eps` | `None` / `1e-6` | A non-`None` weight folds the `kv_a_layernorm` into this kernel, which then reads `latent_cache` **raw**. A caller that already normalized would normalize twice. Only DeepSeek-V4's sparse module folds it upstream. | +| `precomputed_cu_seqlens` | `False` | The kernel fills the cu-seqlens buffers itself; `True` only when the Q half is skipped. | +| `precomputed_fmha_scheduler` | `False` | The scheduler counter and the two bmm scales come from this call, not from a sparse index kernel. | +| `kv_only` / `kv_done_elsewhere` | `False` / `False` | Both halves run in this one call. | +| `quant_scale_qkv` | `None` | `q_nope` in `quant_q_buffer` is not pre-quantized. | diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation.py b/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation.py new file mode 100644 index 000000000000..20865dc2ba53 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation.py @@ -0,0 +1,122 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""MLA generation-phase RoPE + latent KV-cache append + FMHA scheduler-buffer fill.""" + +from typing import List, Optional + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + + +def mla_rope_generation( + fused_q: torch.Tensor, + q_pe: torch.Tensor, + latent_cache: torch.Tensor, + rotary_cos_sin: Optional[torch.Tensor], + cu_q_seqlens: torch.Tensor, + cu_kv_seqlens: torch.Tensor, + fmha_scheduler_counter: torch.Tensor, + mla_bmm1_scale: Optional[torch.Tensor], + mla_bmm2_scale: Optional[torch.Tensor], + quant_q_buffer: Optional[torch.Tensor], + sequence_length: torch.Tensor, + host_past_key_value_lengths: torch.Tensor, + host_context_lengths: torch.Tensor, + num_contexts: int, + kv_cache_block_offsets: Optional[torch.Tensor], + host_kv_cache_pool_pointers: Optional[torch.Tensor], + host_kv_cache_pool_mapping: Optional[torch.Tensor], + kv_scale_orig_quant: Optional[torch.Tensor], + kv_scale_quant_orig: Optional[torch.Tensor], + # rc26: when None the op falls back to kv_scale_orig_quant, which is what + # it did before this parameter existed (dsv3RopeOp.cpp:280). + kv_cache_scale_orig_quant: Optional[torch.Tensor], + out_scale: Optional[torch.Tensor], + block_ids_per_seq: Optional[torch.Tensor], + helix_tensor_params: List[Optional[torch.Tensor]], + predicted_tokens_per_seq: int, + layer_idx: int, + num_heads: int, + num_kv_heads: int, + head_size: int, + # rc26: 0 or rope_size, and non-zero requires an FP4 KV pool. + residual_dim: int, + tokens_per_block: int, + attention_window_size: int, + beam_width: int, + quant_mode: int, + q_scaling: float, + q_lora_rank: int, + kv_lora_rank: int, + qk_nope_head_dim: int, + qk_rope_head_dim: int, + v_head_dim: int, + rope_append: bool, + # Added in rc26; every default below reproduces the op's pre-rc26 + # behaviour. kv_norm_weight non-None would fold the kv_a_layernorm into + # this kernel, which then reads latent_cache RAW -- a caller that already + # normalized would be normalizing twice. + kv_norm_weight: Optional[torch.Tensor] = None, + kv_norm_eps: float = 1e-6, + precomputed_cu_seqlens: bool = False, + precomputed_fmha_scheduler: bool = False, + kv_only: bool = False, + kv_done_elsewhere: bool = False, + quant_scale_qkv: Optional[torch.Tensor] = None, +) -> None: + """RoPE q_pe, append latent_cache rows to the paged MLA KV cache, and fill + the decode-FMHA scheduler buffers. Returns None. + + Where the roped q_pe lands depends on the pool: into fused_q's tail over a + bf16 pool, into quant_q_buffer over an fp8 one (where fused_q is read, not + written). See the contract.""" + torch.ops.trtllm.mla_rope_generation( + fused_q, + q_pe, + latent_cache, + rotary_cos_sin, + cu_q_seqlens, + cu_kv_seqlens, + fmha_scheduler_counter, + mla_bmm1_scale, + mla_bmm2_scale, + quant_q_buffer, + sequence_length, + host_past_key_value_lengths, + host_context_lengths, + num_contexts, + kv_cache_block_offsets, + host_kv_cache_pool_pointers, + host_kv_cache_pool_mapping, + kv_scale_orig_quant, + kv_scale_quant_orig, + kv_cache_scale_orig_quant, + out_scale, + block_ids_per_seq, + helix_tensor_params, + predicted_tokens_per_seq, + layer_idx, + num_heads, + num_kv_heads, + head_size, + residual_dim, + tokens_per_block, + attention_window_size, + beam_width, + quant_mode, + q_scaling, + q_lora_rank, + kv_lora_rank, + qk_nope_head_dim, + qk_rope_head_dim, + v_head_dim, + rope_append, + kv_norm_weight, + kv_norm_eps, + precomputed_cu_seqlens, + precomputed_fmha_scheduler, + kv_only, + kv_done_elsewhere, + quant_scale_qkv, + ) diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation_test.py b/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation_test.py new file mode 100644 index 000000000000..cead752aad1d --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation_test.py @@ -0,0 +1,1459 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the mla_rope_generation catalog entry. + +The op reads paged-KV-cache addressing tensors and per-sequence length +tensors that the runtime normally derives from a KVCacheManager and a +prepared TrtllmAttentionMetadata. The test builds that state for real — +an actual MLA (SELFKONLY, kv_factor=1) KVCacheManager and a prepared +TrtllmAttentionMetadata — then checks the kernel effects against a torch +fp32 reference. + +Two surfaces: + +1. bf16 latent pool (`quant_mode=0`, all fp8 buffers None), H = 16: GPT-J + RoPE of q_pe written into fused_q's tail, [compressed_kv | rope(k_pe)] + appended into the paged latent cache, and the decode-FMHA scheduler + buffers (cu_q/cu_kv/counter) filled. Covered at both pool page sizes: + 64, and 32 (the engine default), where the appended slot is placed at + every 32-slot alignment and multi-step decodes cross a page boundary + between steps. + +2. fp8-e4m3 latent pool (`quant_mode=128`) at the DeepSeek-R1-0528 cell — + H = 128, tokens_per_block = 32, C/R/nope/v = 512/64/128/128, + beam_width = 1, rope_append = True. Here the op stops writing fused_q + entirely and instead writes the three buffers the fp8 MLA decode FMHA + consumes: an e4m3 copy of the fused query in `quant_q_buffer`, and the + two folded FMHA scales in `mla_bmm1_scale` / `mla_bmm2_scale`; the cache + append lands as e4m3. Swept over the KV scaling factor (omitted = the + production call, and explicit 1.0 / 1.5 / 2.0), over q_scaling (1.0 and + DeepSeek-R1's YaRN attention temperature), and over a deliberately + inconsistent scale pair that separates which tensor drives the write side + from which drives the read-side scales. + +3. predicted_tokens_per_seq (`P`) > 1 — the MTP path — over both pools, at + P in {1, 2, 3, 4}. One generation sequence then arrives with P query + tokens at P consecutive absolute positions, and both helpers below are + parameterised on P so P = 1 is the same code path it always was. What the + P > 1 cases pin: each row's own rope position, the P cache rows per + sequence (page-boundary straddles at every 32-slot alignment included), + the scheduler-buffer fills, and which length tensor the position comes + from. Two controls sit beside them — wrong-position mirrors, and the + documented paged-append race armed at P > 1 so a clean pool comparison + is distinguishable from a blind one. +""" + +import math +from typing import List, NamedTuple, Optional + +import torch + +from tensorrt_llm._torch.attention.backends.interface import RopeParams +from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata +from tensorrt_llm._torch.metadata import KVCacheParams +from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager +from tensorrt_llm.bindings import DataType +from tensorrt_llm.bindings.internal.batch_manager import CacheType +from tensorrt_llm.llmapi.llm_args import KvCacheConfig +from tensorrt_llm.mapping import Mapping + +from .mla_rope_generation import mla_rope_generation + +assert torch.cuda.is_available(), "mla_rope_generation requires a CUDA device" + +# DeepSeek-V3 MLA head geometry (num_heads reduced to a TP-slice-like 16). +NUM_HEADS = 16 +KV_LORA_RANK = 512 +QK_ROPE_HEAD_DIM = 64 +QK_NOPE_HEAD_DIM = 128 +V_HEAD_DIM = 128 +GEN_HEAD_SIZE = KV_LORA_RANK + QK_ROPE_HEAD_DIM +# Pool page size. 32 is the engine default (KvCacheConfig.tokens_per_block), +# 64 the value a tuned MLA target opts into; both are covered below. +TOKENS_PER_BLOCK = 64 +PAGE32 = 32 +MAX_SEQ_LEN = 1024 + +# ─── fp8-e4m3 latent pool surface ───────────────────────────────────── +# DeepSeek-R1-0528 cell: full 128-head MLA over an fp8 latent pool. +NUM_HEADS_R1 = 128 +Q_LORA_RANK_R1 = 1536 +QUANT_MODE_FP8_KV_CACHE = 128 # QuantMode.FP8_KV_CACHE +# DeepSeek-R1's YaRN attention temperature: mscale = 0.1 * mscale_all_dim * +# ln(factor) + 1 with mscale_all_dim = 1.0 and factor = 40. trtllm carries it +# as q_scaling = 1 / mscale**2, so the softmax scale mscale**2 / sqrt(nope + R) +# comes out of 1 / (q_scaling * sqrt(nope + R)). ~0.53366 — the value the +# target passes, and the reason q_scaling != 1 has to be certified here. +R1_MSCALE = 0.1 * math.log(40.0) + 1.0 +Q_SCALING_R1 = 1.0 / (R1_MSCALE * R1_MSCALE) + +# The op derives mla_bmm1_scale in fp32 on the device; the reference computes +# the same expression in double and rounds once. Max observed deviation over +# the swept (scale, q_scaling) grid: 7.2e-8 relative — 0.6 fp32 ulp, and zero +# on several of the cases. The gate is torch's default fp32 rtol with atol +# tightened to 0 (18x the observed deviation). It is not a loosening: every +# way of getting this scale wrong misses by a factor, not by ulps. In +# |got - want| / |want|, the metric assert_close gates on: a reference that +# drops the s**2 fold at s = 1.5 is off by 1.25 (9.6e5x the gate), one that +# drops q_scaling = 0.53366 by 0.874 (6.7e5x), one using sqrt(C + R) instead +# of sqrt(nope + R) by 0.732 (5.6e5x), and one leaving log2(e) off element +# [1] by 0.443 (3.4e5x). +BMM_SCALE_RTOL = 1.3e-6 + +# One e4m3 ulp. e4m3 keeps 3 mantissa bits, so 2**-3 relative is one ulp — the +# gate for the two parts of the fp8 output that go through the in-kernel RoPE +# (quant_q_buffer's tail and the appended row's k_pe half), where the kernel +# and a torch fp32 reference can evaluate x*cos -+ y*sin in different orders. +# Paired with a bit-exact-fraction floor, because in practice the coarse e4m3 +# rounding absorbs that difference completely: every roped element measured on +# sm_100 under quant_mode 128 (all scales, all cases below) was bit-exact, so +# neither half of the gate has ever been approached. That both halves still +# discriminate is measured, not assumed — test_fp8_kv_explicit_unit_scale runs +# three wrong mirrors of the same bytes and asserts each blows both. +E4M3_ULP_RTOL = 2**-3 +E4M3_MAX_INEXACT_FRACTION = 1e-3 + + +class _MlaEnv: + """Real op state: MLA KV cache manager + duplicated-layout RoPE table. + + With fp8_pool the manager allocates an e4m3 latent pool (one byte per + element) and the two KV scaling-factor tensors are built; orig_quant and + quant_orig are independent on purpose so a test can pass a deliberately + inconsistent pair. + """ + + def __init__( + self, + max_batch_size: int = 8, + tokens_per_block: int = TOKENS_PER_BLOCK, + num_heads: int = NUM_HEADS, + fp8_pool: bool = False, + orig_quant: Optional[float] = None, + quant_orig: Optional[float] = None, + ) -> None: + self.max_batch_size = max_batch_size + self.tokens_per_block = tokens_per_block + self.num_heads = num_heads + self.fp8_pool = fp8_pool + self.kv_cache_manager = KVCacheManager( + KvCacheConfig(max_tokens=131072, enable_block_reuse=False), + CacheType.SELFKONLY, # MLA latent cache: kv_factor=1, one kv head + num_layers=1, + num_kv_heads=1, + head_dim=GEN_HEAD_SIZE, + tokens_per_block=tokens_per_block, + max_seq_len=MAX_SEQ_LEN, + max_batch_size=max_batch_size, + mapping=Mapping(world_size=1, tp_size=1, rank=0), + dtype=DataType.FP8 if fp8_pool else DataType.BF16, + ) + # The pool must really be paged at the size and element type the case + # claims: the op sizes its page slabs from tokens_per_block and + # quant_mode, not from anything the manager tells it. + assert self.kv_cache_manager.tokens_per_block == tokens_per_block + pool = self.kv_cache_manager.get_buffers(0) + assert pool is not None + assert pool.dtype == (torch.float8_e4m3fn if fp8_pool else torch.bfloat16) + # KV scaling-factor tensors. None for both is what the production call + # site passes (TrtllmAttention.mla_rope_generation hard-codes None), and + # the op then behaves as if both were 1.0. + self.kv_scale_orig_quant = _scale_tensor(orig_quant) + self.kv_scale_quant_orig = _scale_tensor(quant_orig) + # The exact fp32 numbers the op sees, as 0-dim tensors, so the mirror + # cannot disagree with the kernel in the last bit. + self.write_scale = _scale_value(orig_quant) # multiplies on quantize + self.read_scale = _scale_value(quant_orig) # folded into the bmm scales + # Duplicated-layout fp32 (cos, sin) table, as the MLA backend builds + # it (RopeParams.from_config sets duplicate_data=True for MLA models). + rope = RopeParams( + dim=QK_ROPE_HEAD_DIM, + theta=10000.0, + max_positions=MAX_SEQ_LEN, + duplicate_data=True, + ) + _, self.rotary_cos_sin = rope.create_rope_const_params() + + def prepare_metadata( + self, + request_ids: List[int], + seq_lens: List[int], + num_contexts: int, + cached_lens: List[int], + ) -> TrtllmAttentionMetadata: + metadata = TrtllmAttentionMetadata( + max_num_requests=self.max_batch_size, + max_num_tokens=8192, + kv_cache_manager=self.kv_cache_manager, + ) + metadata.seq_lens = torch.tensor(seq_lens, dtype=torch.int) + metadata.num_contexts = num_contexts + metadata.request_ids = request_ids + metadata.prompt_lens = cached_lens + metadata.kv_cache_params = KVCacheParams( + use_cache=True, + num_cached_tokens_per_seq=cached_lens, + ) + metadata.prepare() + return metadata + + def rope_ref(self, x: torch.Tensor, position: int) -> torch.Tensor: + """GPT-J interleaved rotation of the last dim, fp32 math, bf16 result.""" + half = QK_ROPE_HEAD_DIM // 2 + table = self.rotary_cos_sin.view(-1, QK_ROPE_HEAD_DIM, 2) + cos = table[position, :half, 0] + sin = table[position, :half, 1] + pairs = x.float().reshape(*x.shape[:-1], half, 2) + out = torch.empty_like(pairs) + out[..., 0] = pairs[..., 0] * cos - pairs[..., 1] * sin + out[..., 1] = pairs[..., 0] * sin + pairs[..., 1] * cos + return out.reshape(x.shape).to(x.dtype) + + def quantize(self, x: torch.Tensor) -> torch.Tensor: + """e4m3(x * kv_scale_orig_quant) — the op's write-side quantization, + applied to both the appended cache row and the fused query.""" + return (x.float() * self.write_scale).to(torch.float8_e4m3fn) + + def pool_tensor(self) -> torch.Tensor: + pool = self.kv_cache_manager.get_buffers(0) + assert pool is not None + return pool + + def blocks(self, request_id: int) -> List[int]: + return list(self.kv_cache_manager.get_batch_cache_indices([request_id])[0]) + + def slot(self, request_id: int, position: int) -> tuple[int, int]: + tpb = self.tokens_per_block + return self.blocks(request_id)[position // tpb], position % tpb + + def cache_row(self, request_id: int, position: int) -> torch.Tensor: + """Read one token's latent row back from the paged pool.""" + page, offset = self.slot(request_id, position) + return self.pool_tensor()[page, 0, offset, 0] + + def shutdown(self) -> None: + self.kv_cache_manager.shutdown() + + +def _scale_tensor(value: Optional[float]) -> Optional[torch.Tensor]: + if value is None: + return None + return torch.full((1,), value, dtype=torch.float32, device="cuda") + + +def _scale_value(value: Optional[float]) -> torch.Tensor: + """The fp32 number the op uses for this role — 1.0 when the tensor is + omitted. 0-dim so it broadcasts into an fp32 multiply.""" + return torch.tensor(1.0 if value is None else value, dtype=torch.float32).cuda() + + +def _assert_bytes_equal(got: torch.Tensor, want: torch.Tensor, what: str) -> None: + """Bit-exact comparison of two e4m3 tensors, signed zeros included.""" + torch.testing.assert_close( + got.reshape(-1).view(torch.uint8), + want.reshape(-1).view(torch.uint8), + rtol=0.0, + atol=0.0, # pure scale-and-round: must be bit-exact + msg=lambda m: f"{what} not bit-exact\n{m}", + ) + + +def _assert_e4m3_roped(got: torch.Tensor, want: torch.Tensor, what: str) -> None: + """One-e4m3-ulp gate plus a bit-exact-fraction floor, for the parts of the + fp8 output that pass through the in-kernel RoPE.""" + torch.testing.assert_close(got.float(), want.float(), rtol=E4M3_ULP_RTOL, atol=0.0) + inexact = int((got.reshape(-1).view(torch.uint8) != want.reshape(-1).view(torch.uint8)).sum()) + allowed = max(1, int(E4M3_MAX_INEXACT_FRACTION * got.numel())) + assert inexact <= allowed, ( + f"{what} not bit-exact enough: {inexact}/{got.numel()} bytes differ (allowed {allowed})" + ) + + +class _Fp8Run(NamedTuple): + """The tensors one fp8 generation step was driven with and produced. + + Rows of the tensors are generation tokens: row `r` belongs to generation + sequence `r // predicted_tokens_per_seq` and sits at `positions[r]`.""" + + fused_q: torch.Tensor + q_pe: torch.Tensor + latent_cache: torch.Tensor + quant_q_buffer: torch.Tensor + gen_ids: List[int] + positions: List[int] + predicted_tokens_per_seq: int + + +def _positions( + kv_lens: List[int], num_contexts: int, num_gen: int, tokens_per_seq: int +) -> List[int]: + """0-based absolute position of every generation row, in row order. + + Row `r` is generation sequence `r // P`'s `r % P`-th new token; the P new + tokens of one sequence occupy consecutive positions ending at its total + KV length minus one.""" + return [ + kv_lens[num_contexts + g] - tokens_per_seq + t + for g in range(num_gen) + for t in range(tokens_per_seq) + ] + + +def _run_and_check( + env: _MlaEnv, + request_ids: List[int], + seq_lens: List[int], + num_contexts: int, + cached_lens: List[int], + q_pe_contiguous: bool, + predicted_tokens_per_seq: int = 1, +) -> None: + """One generation step over a bf16 latent pool: allocate P slots per + generation sequence, run the op over the generation tokens, and verify + every kernel effect.""" + gen_ids = request_ids[num_contexts:] + num_gen = len(gen_ids) + p = predicted_tokens_per_seq + assert seq_lens[num_contexts:] == [p] * num_gen + num_heads = env.num_heads + for rid in gen_ids: + for _ in range(p): + env.kv_cache_manager.impl.add_token(rid) + metadata = env.prepare_metadata(request_ids, seq_lens, num_contexts, cached_lens) + kv_lens = [c + s for c, s in zip(cached_lens, seq_lens)] + positions = _positions(kv_lens, num_contexts, num_gen, p) + rows = num_gen * p + + fused_q = torch.randn(rows, num_heads, GEN_HEAD_SIZE, dtype=torch.bfloat16, device="cuda") + q_pe = _make_q_pe(rows, num_heads, q_pe_contiguous) + latent_cache = torch.randn(rows, GEN_HEAD_SIZE, dtype=torch.bfloat16, device="cuda") + fused_q_orig = fused_q.clone() + q_pe_orig = q_pe.clone() + cu_q_seqlens = torch.full((num_gen + 1,), -1, dtype=torch.int32, device="cuda") + cu_kv_seqlens = torch.full((num_gen + 1,), -1, dtype=torch.int32, device="cuda") + fmha_scheduler_counter = torch.full((1,), 7, dtype=torch.uint32, device="cuda") + + mla_rope_generation( + fused_q, + q_pe, + latent_cache, + env.rotary_cos_sin, + cu_q_seqlens, + cu_kv_seqlens, + fmha_scheduler_counter, + None, # mla_bmm1_scale: fp8-KV-cache path only + None, # mla_bmm2_scale + None, # quant_q_buffer + metadata.kv_lens_cuda_runtime, + metadata.kv_lens_runtime, + metadata.prompt_lens_cpu_runtime, + num_contexts, + metadata.kv_cache_block_offsets, + env.kv_cache_manager.kv_cache_pool_pointers, + env.kv_cache_manager.kv_cache_pool_mapping, + None, # kv_scale_orig_quant + None, # kv_scale_quant_orig + None, # kv_cache_scale_orig_quant + None, # out_scale + None, # block_ids_per_seq + [None, None], # helix_tensor_params + p, # predicted_tokens_per_seq + 0, # layer_idx + num_heads, + 1, # num_kv_heads + GEN_HEAD_SIZE, + 0, # residual_dim + env.tokens_per_block, + MAX_SEQ_LEN, # attention_window_size + 1, # beam_width + 0, # quant_mode: bf16 KV cache + 1.0, # q_scaling + 0, # q_lora_rank + KV_LORA_RANK, + QK_NOPE_HEAD_DIM, + QK_ROPE_HEAD_DIM, + V_HEAD_DIM, + True, # rope_append + ) + torch.cuda.synchronize() + + # 1. fused_q tail = rope(q_pe) at that row's own position; the absorbed-q + # head slice is untouched. + for r in range(rows): + ref = env.rope_ref(q_pe_orig[r], positions[r]) + torch.testing.assert_close(fused_q[r, :, KV_LORA_RANK:], ref) + torch.testing.assert_close( + fused_q[..., :KV_LORA_RANK], + fused_q_orig[..., :KV_LORA_RANK], + rtol=0.0, + atol=0.0, # caller-owned region: must be bitwise untouched + ) + # q_pe is an input only (mutable in the schema, not mutated in practice). + torch.testing.assert_close(q_pe, q_pe_orig, rtol=0.0, atol=0.0) + + # 2. Cache append: [compressed_kv | rope(k_pe)] at each row's own slot. + for r in range(rows): + rid = gen_ids[r // p] + row = env.cache_row(rid, positions[r]) + torch.testing.assert_close( + row[:KV_LORA_RANK], + latent_cache[r, :KV_LORA_RANK], + rtol=0.0, + atol=0.0, # dtype-preserving copy: must be bitwise equal + ) + ref_k = env.rope_ref(latent_cache[r, KV_LORA_RANK:], positions[r]) + torch.testing.assert_close(row[KV_LORA_RANK:], ref_k) + + # 3. Scheduler buffers over generation sequences only. + _assert_scheduler_buffers( + cu_q_seqlens, + cu_kv_seqlens, + fmha_scheduler_counter, + num_gen, + num_heads, + kv_lens[num_contexts:], + p, + ) + + +def _run_and_check_fp8( + env: _MlaEnv, + request_ids: List[int], + seq_lens: List[int], + num_contexts: int, + cached_lens: List[int], + q_pe_contiguous: bool, + q_scaling: float, + q_lora_rank: int = 0, + pass_bmm1: bool = True, + pass_bmm2: bool = True, + repeats: int = 1, + predicted_tokens_per_seq: int = 1, +) -> _Fp8Run: + """One generation step over an fp8-e4m3 latent pool. Verifies the three + fp8 buffers, the quantized cache append, the scheduler buffers, that the + op writes nothing else in the pool, and that fused_q/q_pe/latent_cache all + come back bitwise untouched.""" + gen_ids = request_ids[num_contexts:] + num_gen = len(gen_ids) + p = predicted_tokens_per_seq + assert seq_lens[num_contexts:] == [p] * num_gen + num_heads = env.num_heads + for rid in gen_ids: + for _ in range(p): + env.kv_cache_manager.impl.add_token(rid) + metadata = env.prepare_metadata(request_ids, seq_lens, num_contexts, cached_lens) + kv_lens = [c + s for c, s in zip(cached_lens, seq_lens)] + positions = _positions(kv_lens, num_contexts, num_gen, p) + rows = num_gen * p + + fused_q = torch.randn(rows, num_heads, GEN_HEAD_SIZE, dtype=torch.bfloat16, device="cuda") + q_pe = _make_q_pe(rows, num_heads, q_pe_contiguous) + latent_cache = torch.randn(rows, GEN_HEAD_SIZE, dtype=torch.bfloat16, device="cuda") + fused_q_orig = fused_q.clone() + q_pe_orig = q_pe.clone() + latent_orig = latent_cache.clone() + cu_q_seqlens = torch.full((num_gen + 1,), -1, dtype=torch.int32, device="cuda") + cu_kv_seqlens = torch.full((num_gen + 1,), -1, dtype=torch.int32, device="cuda") + fmha_scheduler_counter = torch.full((1,), 7, dtype=torch.uint32, device="cuda") + # Production allocates quant_q_buffer as uint8 [tokens, heads, C + R]. + quant_q_buffer = torch.full( + (rows, num_heads, GEN_HEAD_SIZE), 255, dtype=torch.uint8, device="cuda" + ) + bmm1 = torch.full((2,), -99.0, dtype=torch.float32, device="cuda") if pass_bmm1 else None + bmm2 = torch.full((1,), -99.0, dtype=torch.float32, device="cuda") if pass_bmm2 else None + pool = env.pool_tensor() + pool_before = pool.clone() + + def call() -> None: + mla_rope_generation( + fused_q, + q_pe, + latent_cache, + env.rotary_cos_sin, + cu_q_seqlens, + cu_kv_seqlens, + fmha_scheduler_counter, + bmm1, + bmm2, + quant_q_buffer, + metadata.kv_lens_cuda_runtime, + metadata.kv_lens_runtime, + metadata.prompt_lens_cpu_runtime, + num_contexts, + metadata.kv_cache_block_offsets, + env.kv_cache_manager.kv_cache_pool_pointers, + env.kv_cache_manager.kv_cache_pool_mapping, + env.kv_scale_orig_quant, + env.kv_scale_quant_orig, + None, # kv_cache_scale_orig_quant + None, # out_scale + None, # block_ids_per_seq + [None, None], # helix_tensor_params + p, # predicted_tokens_per_seq + 0, # layer_idx + num_heads, + 1, # num_kv_heads + GEN_HEAD_SIZE, + 0, # residual_dim + env.tokens_per_block, + MAX_SEQ_LEN, # attention_window_size + 1, # beam_width + QUANT_MODE_FP8_KV_CACHE, + q_scaling, + q_lora_rank, + KV_LORA_RANK, + QK_NOPE_HEAD_DIM, + QK_ROPE_HEAD_DIM, + V_HEAD_DIM, + True, # rope_append + ) + torch.cuda.synchronize() + + call() + + # 1. fused_q is read-only on this path: the roped q goes to quant_q_buffer + # instead, so neither half of fused_q is written. + torch.testing.assert_close(fused_q, fused_q_orig, rtol=0.0, atol=0.0) + torch.testing.assert_close(q_pe, q_pe_orig, rtol=0.0, atol=0.0) + torch.testing.assert_close(latent_cache, latent_orig, rtol=0.0, atol=0.0) + + # 2. quant_q_buffer = e4m3(fused_q * orig_quant) with the tail replaced by + # the roped q_pe (rounded to bf16 first, as the kernel's own dtype does). + quant_q = quant_q_buffer.view(torch.float8_e4m3fn) + roped_q = torch.stack([env.rope_ref(q_pe_orig[r], positions[r]) for r in range(rows)]) + _assert_bytes_equal( + quant_q[..., :KV_LORA_RANK], + env.quantize(fused_q_orig[..., :KV_LORA_RANK]), + "quant_q_buffer absorbed-q head", + ) + _assert_e4m3_roped( + quant_q[..., KV_LORA_RANK:], + env.quantize(roped_q), + "quant_q_buffer roped tail", + ) + + # 3. The two decode-FMHA scales, folded with the read-side scaling factor. + if bmm1 is not None: + s = float(env.read_scale) + x = s * s / (q_scaling * math.sqrt(QK_NOPE_HEAD_DIM + QK_ROPE_HEAD_DIM)) + want1 = torch.tensor([x, x * math.log2(math.e)], dtype=torch.float32, device="cuda") + torch.testing.assert_close(bmm1, want1, rtol=BMM_SCALE_RTOL, atol=0.0) + if bmm2 is not None: + torch.testing.assert_close( + bmm2, + env.read_scale.reshape(1), + rtol=0.0, + atol=0.0, # a copy of kv_scale_quant_orig, not a computation + ) + + # 4. Cache append, quantized: e4m3([compressed_kv | rope(k_pe)] * orig_quant), + # one row per generation token at that token's own position. + for r in range(rows): + rid = gen_ids[r // p] + row = env.cache_row(rid, positions[r]) + _assert_bytes_equal( + row[:KV_LORA_RANK], + env.quantize(latent_orig[r, :KV_LORA_RANK]), + f"appended compressed_kv (request {rid}, position {positions[r]})", + ) + ref_k = env.rope_ref(latent_orig[r, KV_LORA_RANK:], positions[r]) + _assert_e4m3_roped( + row[KV_LORA_RANK:], + env.quantize(ref_k), + f"appended k_pe (request {rid}, position {positions[r]})", + ) + + # 5. Nothing else in the pool moved: exactly P C+R-byte rows per + # generation sequence, at the slots their positions address. + changed = pool.view(torch.uint8) != pool_before.view(torch.uint8) + expected_slots = {env.slot(gen_ids[r // p], positions[r]) for r in range(rows)} + idx = changed.nonzero().cpu() + got_slots = set(zip(idx[:, 0].tolist(), idx[:, 2].tolist())) + assert got_slots == expected_slots, ( + f"op touched pool slots {sorted(got_slots)}, expected {sorted(expected_slots)}" + ) + # Upper bound, not equality: a written byte that happens to match the byte + # already there (an element quantizing to +0 over a zeroed pool) is + # invisible to a snapshot diff. The rows' contents are pinned above. + assert idx.shape[0] <= rows * GEN_HEAD_SIZE, ( + f"op wrote {idx.shape[0]} pool bytes, at most {rows * GEN_HEAD_SIZE} rows' worth expected" + ) + + # 6. Scheduler buffers over generation sequences only. + _assert_scheduler_buffers( + cu_q_seqlens, + cu_kv_seqlens, + fmha_scheduler_counter, + num_gen, + num_heads, + kv_lens[num_contexts:], + p, + ) + + # 7. Repeating the prepared step rewrites the same bytes: the position is + # derived from sequence_length, so the call is idempotent rather than + # double-appending, and the fp8 outputs are run-to-run stable. + for _ in range(repeats - 1): + quant_q_before = quant_q_buffer.clone() + pool_after_first = pool.clone() + call() + torch.testing.assert_close(quant_q_buffer, quant_q_before, rtol=0.0, atol=0.0) + torch.testing.assert_close( + pool.view(torch.uint8), + pool_after_first.view(torch.uint8), + rtol=0.0, + atol=0.0, + ) + + return _Fp8Run( + fused_q=fused_q_orig, + q_pe=q_pe_orig, + latent_cache=latent_orig, + quant_q_buffer=quant_q_buffer, + gen_ids=gen_ids, + positions=positions, + predicted_tokens_per_seq=p, + ) + + +def _make_q_pe(rows: int, num_heads: int, contiguous: bool) -> torch.Tensor: + if contiguous: + return torch.randn(rows, num_heads, QK_ROPE_HEAD_DIM, dtype=torch.bfloat16, device="cuda") + # mla.py-style strided view: q_pe sliced out of a packed q tensor. + q = torch.randn( + rows, + num_heads, + QK_NOPE_HEAD_DIM + QK_ROPE_HEAD_DIM, + dtype=torch.bfloat16, + device="cuda", + ) + return q.split([QK_NOPE_HEAD_DIM, QK_ROPE_HEAD_DIM], dim=-1)[1] + + +def _assert_scheduler_buffers( + cu_q_seqlens: torch.Tensor, + cu_kv_seqlens: torch.Tensor, + fmha_scheduler_counter: torch.Tensor, + num_gen: int, + num_heads: int, + gen_kv_lens: List[int], + predicted_tokens_per_seq: int, +) -> None: + # cu_q counts q rows: num_heads rows per generation token, and a + # generation sequence contributes predicted_tokens_per_seq of them. + expected_cu_q = ( + torch.arange(num_gen + 1, dtype=torch.int32) * num_heads * predicted_tokens_per_seq + ) + expected_cu_kv = torch.zeros(num_gen + 1, dtype=torch.int32) + expected_cu_kv[1:] = torch.tensor(gen_kv_lens, dtype=torch.int32).cumsum(0) + torch.testing.assert_close(cu_q_seqlens.cpu(), expected_cu_q) + torch.testing.assert_close(cu_kv_seqlens.cpu(), expected_cu_kv) + assert fmha_scheduler_counter.item() == 0 + + +def test_bf16_decode_batch_strided_q_pe() -> None: + """Pure-decode batch; one sequence's new slot crosses a block boundary + (cached 64 = one full 64-token block); q_pe is a packed-q strided view.""" + torch.manual_seed(0) + env = _MlaEnv() + try: + env.kv_cache_manager.add_dummy_requests([0, 1], token_nums=[64, 32]) + _run_and_check( + env, + request_ids=[0, 1], + seq_lens=[1, 1], + num_contexts=0, + cached_lens=[64, 32], + q_pe_contiguous=False, + ) + finally: + env.shutdown() + + +def test_bf16_mixed_batch_skips_context() -> None: + """Context sequence leads the batch; the op consumes only the generation + tokens and indexes length/block tensors starting at num_contexts.""" + torch.manual_seed(1) + env = _MlaEnv() + try: + env.kv_cache_manager.add_dummy_requests([0, 1, 2], token_nums=[40, 100, 7]) + _run_and_check( + env, + request_ids=[0, 1, 2], + seq_lens=[40, 1, 1], + num_contexts=1, + cached_lens=[0, 100, 7], + q_pe_contiguous=True, + ) + finally: + env.shutdown() + + +def test_bf16_large_decode_batch_multi_step() -> None: + """64-sequence decode batch (many tokens per call), two consecutive + steps so the second call appends after the first call's tokens.""" + torch.manual_seed(2) + env = _MlaEnv(max_batch_size=64) + try: + rids = list(range(64)) + cached = [(37 * (i + 1)) % 800 + 1 for i in range(64)] + env.kv_cache_manager.add_dummy_requests(rids, token_nums=cached) + for step in range(2): + _run_and_check( + env, + request_ids=rids, + seq_lens=[1] * 64, + num_contexts=0, + cached_lens=[c + step for c in cached], + q_pe_contiguous=True, + ) + finally: + env.shutdown() + + +def test_bf16_page32_decode_batch_alignments() -> None: + """Page size 32 (the engine default): a pure-decode batch whose three + new slots land at every alignment a 32-token page has — position 32 + (page 1 slot 0, a fresh page after one full page), 31 (page 0's last + slot) and 64 (page 2 slot 0, after two full pages). q_pe is a packed-q + strided view.""" + torch.manual_seed(3) + env = _MlaEnv(tokens_per_block=PAGE32) + try: + env.kv_cache_manager.add_dummy_requests([0, 1, 2], token_nums=[32, 31, 64]) + _run_and_check( + env, + request_ids=[0, 1, 2], + seq_lens=[1, 1, 1], + num_contexts=0, + cached_lens=[32, 31, 64], + q_pe_contiguous=False, + ) + finally: + env.shutdown() + + +def test_bf16_page32_mixed_batch_skips_context() -> None: + """Page size 32, context sequence leading the batch: the op consumes + only the generation tokens and indexes length/block tensors from + num_contexts. The generation slots sit at position 96 (page 3 slot 0) + and 7 (mid first page).""" + torch.manual_seed(4) + env = _MlaEnv(tokens_per_block=PAGE32) + try: + env.kv_cache_manager.add_dummy_requests([0, 1, 2], token_nums=[40, 96, 7]) + _run_and_check( + env, + request_ids=[0, 1, 2], + seq_lens=[40, 1, 1], + num_contexts=1, + cached_lens=[0, 96, 7], + q_pe_contiguous=True, + ) + finally: + env.shutdown() + + +def test_bf16_page32_large_decode_batch_multi_step() -> None: + """Page size 32, 64-sequence decode batch over two consecutive steps. + The cached lengths spread the new slots over 25 pages; two sequences + (i = 5, 37) sit on a page's last slot at step 0 and cross into the next + page at step 1, and two (i = 18, 50) start a fresh page at step 0.""" + torch.manual_seed(5) + env = _MlaEnv(max_batch_size=64, tokens_per_block=PAGE32) + try: + rids = list(range(64)) + cached = [(37 * (i + 1)) % 800 + 1 for i in range(64)] + assert [i for i in range(64) if cached[i] % PAGE32 == PAGE32 - 1] == [5, 37] + assert [i for i in range(64) if cached[i] % PAGE32 == 0] == [18, 50] + env.kv_cache_manager.add_dummy_requests(rids, token_nums=cached) + for step in range(2): + _run_and_check( + env, + request_ids=rids, + seq_lens=[1] * 64, + num_contexts=0, + cached_lens=[c + step for c in cached], + q_pe_contiguous=True, + ) + finally: + env.shutdown() + + +def _fp8_env( + max_batch_size: int = 8, + orig_quant: Optional[float] = None, + quant_orig: Optional[float] = None, +) -> _MlaEnv: + """DeepSeek-R1-0528 cell over an fp8-e4m3 latent pool: H = 128, page 32.""" + return _MlaEnv( + max_batch_size=max_batch_size, + tokens_per_block=PAGE32, + num_heads=NUM_HEADS_R1, + fp8_pool=True, + orig_quant=orig_quant, + quant_orig=quant_orig, + ) + + +def test_fp8_kv_decode_production_config() -> None: + """The production fp8 generation call: both KV scale tensors omitted (what + TrtllmAttention.mla_rope_generation passes), DeepSeek-R1's YaRN q_scaling, + q_lora_rank 1536, q_pe a packed-q strided view. Three decode slots at + positions 32 / 31 / 64 — a fresh page, a page's last slot, and a fresh + page after two full ones. Run twice to pin run-to-run stability.""" + torch.manual_seed(10) + env = _fp8_env() + try: + env.kv_cache_manager.add_dummy_requests([0, 1, 2], token_nums=[32, 31, 64]) + _run_and_check_fp8( + env, + request_ids=[0, 1, 2], + seq_lens=[1, 1, 1], + num_contexts=0, + cached_lens=[32, 31, 64], + q_pe_contiguous=False, + q_scaling=Q_SCALING_R1, + q_lora_rank=Q_LORA_RANK_R1, + repeats=2, + ) + finally: + env.shutdown() + + +def test_fp8_kv_explicit_unit_scale() -> None: + """Explicit scaling factor 1.0 (a checkpoint whose k_scale/v_scale are the + production 1.0, passed as real tensors) at q_scaling 1.0, then the + wrong-variant controls that measure what the e4m3-ulp gate discriminates. + Every variant below is a plausible mis-derivation of the same buffer.""" + torch.manual_seed(11) + env = _fp8_env(orig_quant=1.0, quant_orig=1.0) + try: + env.kv_cache_manager.add_dummy_requests([0, 1], token_nums=[32, 31]) + run = _run_and_check_fp8( + env, + request_ids=[0, 1], + seq_lens=[1, 1], + num_contexts=0, + cached_lens=[32, 31], + q_pe_contiguous=True, + q_scaling=1.0, + ) + # Controls. Every fp8 comparison above came back bit-exact, which is + # only meaningful if the same comparison moves for a wrong mirror. + # Measured on sm_100 at this case — the k_pe half of request 0's + # appended row (position 32) against three plausible mis-derivations: + # + # mirror bytes differing max rel dev + # un-roped k_pe 49/64 (76.6%) 92.4 (739x) + # rope at position + 1 28/64 (43.8%) 3.2 (25.6x) + # the other request's roped row 63/64 (98.4%) 16.6 (133x) + # + # with the multiplier against E4M3_ULP_RTOL. The bit-exact-fraction + # floor allows 1 byte in a 64-element row, so even the weakest control + # — the off-by-one position, where at position 32 the low-frequency + # rope pairs barely move and over half the e4m3 bytes stay put — is 28x + # past the fraction allowance and 25.6x past the tolerance. + # quant_q_buffer's tail against fused_q's own never-roped q_pe differs + # in 74.5% of 16 384 bytes, against an allowance of 16. + gen_ids, positions = run.gen_ids, run.positions + row0 = env.cache_row(gen_ids[0], positions[0])[KV_LORA_RANK:] + row0_bytes = row0.view(torch.uint8) + wrong = { + "un-roped k_pe": env.quantize(run.latent_cache[0, KV_LORA_RANK:]), + "rope at position + 1": env.quantize( + env.rope_ref(run.latent_cache[0, KV_LORA_RANK:], positions[0] + 1) + ), + "other request's row": env.quantize( + env.rope_ref(run.latent_cache[1, KV_LORA_RANK:], positions[1]) + ), + } + for name, mirror in wrong.items(): + frac = float((row0_bytes != mirror.view(torch.uint8)).float().mean()) + rel = float( + ((row0.float() - mirror.float()).abs() / mirror.float().abs().clamp(min=1e-9)).max() + ) + assert frac > 100 * E4M3_MAX_INEXACT_FRACTION, ( + f"control '{name}' moved only {frac:.3f} of bytes — the " + "bit-exact-fraction floor would not have caught it" + ) + assert rel > 4 * E4M3_ULP_RTOL, ( + f"control '{name}' stayed within {rel:.3g} relative — the " + "one-ulp tolerance would not have caught it" + ) + quant_tail = run.quant_q_buffer[..., KV_LORA_RANK:].view(torch.float8_e4m3fn) + unroped = env.quantize(run.q_pe) + frac = float((quant_tail.view(torch.uint8) != unroped.view(torch.uint8)).float().mean()) + assert frac > 0.5, f"quant_q tail control only moved {frac:.3f} of bytes" + finally: + env.shutdown() + + +def test_fp8_kv_scale_factor_sweep() -> None: + """KV scaling factors other than the production 1.0, passed as the + reciprocal pair a calibrated fp8 checkpoint would produce. 1.5 is not a + power of two, so 1 / 1.5 is inexact in fp32 and the write-side multiply is + a real rescale rather than an exponent shift. Both q_scaling values are + exercised across the sweep.""" + for scale, q_scaling in ((1.5, Q_SCALING_R1), (2.0, 1.0)): + torch.manual_seed(12) + env = _fp8_env(orig_quant=1.0 / scale, quant_orig=scale) + try: + env.kv_cache_manager.add_dummy_requests([0, 1], token_nums=[32, 95]) + _run_and_check_fp8( + env, + request_ids=[0, 1], + seq_lens=[1, 1], + num_contexts=0, + cached_lens=[32, 95], + q_pe_contiguous=True, + q_scaling=q_scaling, + ) + finally: + env.shutdown() + + +def test_fp8_kv_scale_tensors_are_independent() -> None: + """A deliberately inconsistent scale pair (orig_quant 0.25, quant_orig 3.0) + separates the two roles: orig_quant alone drives the write side (the + appended row and quant_q_buffer), quant_orig alone drives both decode-FMHA + scales. The reference is built the same way, so a passing run pins the + split — the op does not derive either tensor from the other, and a caller + passing a non-reciprocal pair gets a silently inconsistent round trip.""" + torch.manual_seed(13) + env = _fp8_env(orig_quant=0.25, quant_orig=3.0) + try: + env.kv_cache_manager.add_dummy_requests([0, 1], token_nums=[32, 31]) + _run_and_check_fp8( + env, + request_ids=[0, 1], + seq_lens=[1, 1], + num_contexts=0, + cached_lens=[32, 31], + q_pe_contiguous=True, + q_scaling=Q_SCALING_R1, + ) + finally: + env.shutdown() + + +def test_fp8_kv_mixed_batch_skips_context() -> None: + """Context sequence leading the batch over an fp8 pool: the op consumes + only the generation tokens and indexes length/block tensors from + num_contexts. The generation slots sit at position 96 (page 3 slot 0) and + 7 (mid first page).""" + torch.manual_seed(14) + env = _fp8_env(orig_quant=1.0, quant_orig=1.0) + try: + env.kv_cache_manager.add_dummy_requests([0, 1, 2], token_nums=[40, 96, 7]) + _run_and_check_fp8( + env, + request_ids=[0, 1, 2], + seq_lens=[40, 1, 1], + num_contexts=1, + cached_lens=[0, 96, 7], + q_pe_contiguous=True, + q_scaling=Q_SCALING_R1, + ) + finally: + env.shutdown() + + +def test_fp8_kv_large_decode_batch_multi_step() -> None: + """64-sequence fp8 decode batch (8192 quantized q rows per call) over two + consecutive steps, at the production scale (both tensors omitted). The + cached lengths spread the new slots over 25 pages; two sequences sit on a + page's last slot at step 0 and cross into the next page at step 1, and two + start a fresh page at step 0.""" + torch.manual_seed(15) + env = _fp8_env(max_batch_size=64) + try: + rids = list(range(64)) + cached = [(37 * (i + 1)) % 800 + 1 for i in range(64)] + assert [i for i in range(64) if cached[i] % PAGE32 == PAGE32 - 1] == [5, 37] + assert [i for i in range(64) if cached[i] % PAGE32 == 0] == [18, 50] + env.kv_cache_manager.add_dummy_requests(rids, token_nums=cached) + for step in range(2): + _run_and_check_fp8( + env, + request_ids=rids, + seq_lens=[1] * 64, + num_contexts=0, + cached_lens=[c + step for c in cached], + q_pe_contiguous=True, + q_scaling=Q_SCALING_R1, + ) + finally: + env.shutdown() + + +def test_fp8_kv_bmm_scale_buffers_are_optional() -> None: + """Omitting mla_bmm1_scale or mla_bmm2_scale is accepted, not rejected: + the op simply does not write that buffer and produces everything else + unchanged. (quant_q_buffer is different — see the contract; omitting it is + an illegal memory access, so it cannot be exercised in-process.)""" + for pass_bmm1, pass_bmm2 in ((False, True), (True, False)): + torch.manual_seed(16) + env = _fp8_env(orig_quant=1.0, quant_orig=1.0) + try: + env.kv_cache_manager.add_dummy_requests([0], token_nums=[31]) + _run_and_check_fp8( + env, + request_ids=[0], + seq_lens=[1], + num_contexts=0, + cached_lens=[31], + q_pe_contiguous=True, + q_scaling=Q_SCALING_R1, + pass_bmm1=pass_bmm1, + pass_bmm2=pass_bmm2, + ) + finally: + env.shutdown() + + +class _Fp8Step: + """One prepared fp8 generation step, callable with argument overrides. + + `_run_and_check_fp8` allocates its own buffers and verifies every effect; + this is its low-level twin, for the two cases that have to vary an + argument it does not expose — the host length tensor and the block-offset + table. Pure decode (`num_contexts = 0`), production KV scale tensors.""" + + def __init__( + self, + env: _MlaEnv, + request_ids: List[int], + cached_lens: List[int], + predicted_tokens_per_seq: int, + q_scaling: float = Q_SCALING_R1, + ) -> None: + self.env = env + self.p = predicted_tokens_per_seq + self.q_scaling = q_scaling + self.gen_ids = request_ids + num_gen = len(request_ids) + for rid in request_ids: + for _ in range(self.p): + env.kv_cache_manager.impl.add_token(rid) + self.metadata = env.prepare_metadata(request_ids, [self.p] * num_gen, 0, cached_lens) + kv_lens = [c + self.p for c in cached_lens] + self.positions = _positions(kv_lens, 0, num_gen, self.p) + self.rows = num_gen * self.p + num_heads = env.num_heads + self.fused_q = torch.randn( + self.rows, num_heads, GEN_HEAD_SIZE, dtype=torch.bfloat16, device="cuda" + ) + self.q_pe = torch.randn( + self.rows, num_heads, QK_ROPE_HEAD_DIM, dtype=torch.bfloat16, device="cuda" + ) + self.latent_cache = torch.randn( + self.rows, GEN_HEAD_SIZE, dtype=torch.bfloat16, device="cuda" + ) + self.quant_q_buffer = torch.zeros( + self.rows, num_heads, GEN_HEAD_SIZE, dtype=torch.uint8, device="cuda" + ) + self.cu_q_seqlens = torch.zeros(num_gen + 1, dtype=torch.int32, device="cuda") + self.cu_kv_seqlens = torch.zeros(num_gen + 1, dtype=torch.int32, device="cuda") + self.fmha_scheduler_counter = torch.zeros(1, dtype=torch.uint32, device="cuda") + self.mla_bmm1_scale = torch.zeros(2, dtype=torch.float32, device="cuda") + self.mla_bmm2_scale = torch.zeros(1, dtype=torch.float32, device="cuda") + + def call( + self, + host_past: Optional[torch.Tensor] = None, + block_offsets: Optional[torch.Tensor] = None, + ) -> None: + md = self.metadata + mla_rope_generation( + self.fused_q, + self.q_pe, + self.latent_cache, + self.env.rotary_cos_sin, + self.cu_q_seqlens, + self.cu_kv_seqlens, + self.fmha_scheduler_counter, + self.mla_bmm1_scale, + self.mla_bmm2_scale, + self.quant_q_buffer, + md.kv_lens_cuda_runtime, + md.kv_lens_runtime if host_past is None else host_past, + md.prompt_lens_cpu_runtime, + 0, # num_contexts + md.kv_cache_block_offsets if block_offsets is None else block_offsets, + self.env.kv_cache_manager.kv_cache_pool_pointers, + self.env.kv_cache_manager.kv_cache_pool_mapping, + self.env.kv_scale_orig_quant, + self.env.kv_scale_quant_orig, + None, # kv_cache_scale_orig_quant + None, # out_scale + None, # block_ids_per_seq + [None, None], # helix_tensor_params + self.p, + 0, # layer_idx + self.env.num_heads, + 1, # num_kv_heads + GEN_HEAD_SIZE, + 0, # residual_dim + self.env.tokens_per_block, + MAX_SEQ_LEN, # attention_window_size + 1, # beam_width + QUANT_MODE_FP8_KV_CACHE, + self.q_scaling, + Q_LORA_RANK_R1, + KV_LORA_RANK, + QK_NOPE_HEAD_DIM, + QK_ROPE_HEAD_DIM, + V_HEAD_DIM, + True, # rope_append + ) + torch.cuda.synchronize() + + def assert_rows_at_sequence_length_slots(self) -> None: + """Every appended row is e4m3([compressed_kv | rope(k_pe)]) at the slot + `sequence_length - P + t` addresses.""" + for r in range(self.rows): + rid = self.gen_ids[r // self.p] + row = self.env.cache_row(rid, self.positions[r]) + _assert_bytes_equal( + row[:KV_LORA_RANK], + self.env.quantize(self.latent_cache[r, :KV_LORA_RANK]), + f"appended compressed_kv (request {rid}, row {r})", + ) + _assert_e4m3_roped( + row[KV_LORA_RANK:], + self.env.quantize( + self.env.rope_ref(self.latent_cache[r, KV_LORA_RANK:], self.positions[r]) + ), + f"appended k_pe (request {rid}, row {r})", + ) + + +def test_fp8_kv_mtp_production_sweep() -> None: + """The production MTP generation call at predicted_tokens_per_seq 1, 2, 3 + and 4 — 1 as the regression check that the single-token path did not move, + 2-4 because a target sweeping max_draft_len over 1/2/3 passes + max_draft_len + 1. Both KV scale tensors omitted (what the engine's own + MLA call site passes), DeepSeek-R1's YaRN q_scaling, q_lora_rank 1536, + q_pe a packed-q strided view. The three cached lengths put the P rows of + each sequence at a different 32-slot alignment: ending on a page's last + slot, starting a fresh page, and straddling the boundary. Every case runs + twice to pin that a repeated call is idempotent at P > 1 too.""" + for p in (1, 2, 3, 4): + torch.manual_seed(20 + p) + env = _fp8_env() + try: + cached = [PAGE32 - p, PAGE32, 2 * PAGE32 - 1] + env.kv_cache_manager.add_dummy_requests([0, 1, 2], token_nums=cached) + _run_and_check_fp8( + env, + request_ids=[0, 1, 2], + seq_lens=[p] * 3, + num_contexts=0, + cached_lens=cached, + q_pe_contiguous=False, + q_scaling=Q_SCALING_R1, + q_lora_rank=Q_LORA_RANK_R1, + repeats=2, + predicted_tokens_per_seq=p, + ) + finally: + env.shutdown() + + +def test_fp8_kv_mtp_page32_alignments() -> None: + """P = 2, 3 and 4 with the P rows of a sequence placed at every alignment + a 32-slot page allows: entirely inside a page, ending on its last slot, + straddling the boundary after each of the 1..P-1 possible splits, and + starting a fresh page at slot 0. A straddling generation sequence is new + at P > 1 — at P = 1 one call could only ever write one slot per + sequence.""" + for p in (2, 3, 4): + torch.manual_seed(30 + p) + env = _fp8_env() + try: + # First-token position of sequence k, walking the boundary past + # the whole P-row block. + cached = [PAGE32 - p - 1 + k for k in range(p + 2)] + rids = list(range(len(cached))) + env.kv_cache_manager.add_dummy_requests(rids, token_nums=cached) + _run_and_check_fp8( + env, + request_ids=rids, + seq_lens=[p] * len(rids), + num_contexts=0, + cached_lens=cached, + q_pe_contiguous=True, + q_scaling=Q_SCALING_R1, + predicted_tokens_per_seq=p, + ) + finally: + env.shutdown() + + +def test_fp8_kv_mtp_mixed_batch_skips_context() -> None: + """Context sequence leading a batch whose generation sequences each carry + P = 3 tokens: the op consumes only the generation rows and indexes the + length and block tensors from num_contexts. Both generation sequences + straddle a page boundary.""" + torch.manual_seed(34) + env = _fp8_env(orig_quant=1.0, quant_orig=1.0) + try: + env.kv_cache_manager.add_dummy_requests([0, 1, 2], token_nums=[40, 30, 63]) + _run_and_check_fp8( + env, + request_ids=[0, 1, 2], + seq_lens=[40, 3, 3], + num_contexts=1, + cached_lens=[0, 30, 63], + q_pe_contiguous=True, + q_scaling=Q_SCALING_R1, + predicted_tokens_per_seq=3, + ) + finally: + env.shutdown() + + +def test_fp8_kv_mtp_scale_factor() -> None: + """A calibrated fp8 checkpoint's reciprocal scale pair at P = 3, with + q_scaling 1.0: the write-side factor applies to all P appended rows and + all P quantized query rows, and the two decode-FMHA scales stay per-batch + scalars that P does not enter.""" + torch.manual_seed(35) + env = _fp8_env(orig_quant=1.0 / 1.5, quant_orig=1.5) + try: + env.kv_cache_manager.add_dummy_requests([0, 1], token_nums=[30, 95]) + _run_and_check_fp8( + env, + request_ids=[0, 1], + seq_lens=[3, 3], + num_contexts=0, + cached_lens=[30, 95], + q_pe_contiguous=True, + q_scaling=1.0, + predicted_tokens_per_seq=3, + ) + finally: + env.shutdown() + + +def test_fp8_kv_mtp_large_decode_batch_multi_step() -> None: + """64-sequence fp8 MTP batch at P = 2 (16 384 quantized q rows and 128 + appended cache rows per call) over two consecutive steps, at the + production scale. The cached lengths spread the rows over 25 pages; two + sequences straddle a page boundary inside a single call at step 0, and two + start a fresh page.""" + torch.manual_seed(36) + env = _fp8_env(max_batch_size=64) + try: + rids = list(range(64)) + cached = [(37 * (i + 1)) % 800 + 1 for i in range(64)] + assert [i for i in range(64) if cached[i] % PAGE32 == PAGE32 - 1] == [5, 37] + assert [i for i in range(64) if cached[i] % PAGE32 == 0] == [18, 50] + env.kv_cache_manager.add_dummy_requests(rids, token_nums=cached) + for step in range(2): + _run_and_check_fp8( + env, + request_ids=rids, + seq_lens=[2] * 64, + num_contexts=0, + cached_lens=[c + 2 * step for c in cached], + q_pe_contiguous=True, + q_scaling=Q_SCALING_R1, + predicted_tokens_per_seq=2, + ) + finally: + env.shutdown() + + +def test_fp8_kv_mtp_position_controls() -> None: + """P = 4 at the R1 cell, plus the wrong-position mirrors that measure what + the per-row rope gate discriminates. The failure mode they guard is the + one no shape check catches: every one of the P rows roped at a single + position. Measured on sm_100 at this case (positions 29-32 and 40-43), + against the appended k_pe half of the named row — 64 e4m3 bytes, gated at + one e4m3 ulp with an allowance of 1 byte: + + row mirror bytes differ max rel dev + 0 roped at the sequence's last position 35/64 (54.7%) 13.0 (104x) + 0 roped at the next row's position 21/64 (32.8%) 11.3 (90x) + 0 the next row's latent at row 0's pos 62/64 (96.9%) 1.4e9 + 3 roped at the sequence's first position 30/64 (46.9%) 40.1 (321x) + + with the multiplier against E4M3_ULP_RTOL. The weakest control is the + off-by-one position — at these positions the low-frequency rope pairs + barely move, so two thirds of the e4m3 bytes are genuinely unchanged — + and it is still 21x past the fraction allowance and 90x past the + tolerance. (The row-mixup mirror's relative deviation is degenerate + rather than informative: one of its elements quantizes to zero, so the + ratio is bounded only by the reference's 1e-9 clamp. Its byte fraction is + the meaningful half.) The single-position collapse is checked from both + ends, because a bug pinning all P positions to L - 1 is invisible on the + last row and one pinning them to L - P is invisible on the first.""" + torch.manual_seed(37) + env = _fp8_env(orig_quant=1.0, quant_orig=1.0) + try: + env.kv_cache_manager.add_dummy_requests([0, 1], token_nums=[29, 40]) + run = _run_and_check_fp8( + env, + request_ids=[0, 1], + seq_lens=[4, 4], + num_contexts=0, + cached_lens=[29, 40], + q_pe_contiguous=True, + q_scaling=Q_SCALING_R1, + predicted_tokens_per_seq=4, + ) + p = run.predicted_tokens_per_seq + latent, positions = run.latent_cache, run.positions + controls = [ + (0, "roped at the sequence's last position", positions[p - 1], 0), + (0, "roped at the next row's position", positions[1], 0), + (0, "the next row's latent at row 0's position", positions[0], 1), + (p - 1, "roped at the sequence's first position", positions[0], p - 1), + ] + for row, name, mirror_pos, src_row in controls: + got = env.cache_row(run.gen_ids[0], positions[row])[KV_LORA_RANK:] + mirror = env.quantize(env.rope_ref(latent[src_row, KV_LORA_RANK:], mirror_pos)) + frac = float((got.view(torch.uint8) != mirror.view(torch.uint8)).float().mean()) + rel = float( + ((got.float() - mirror.float()).abs() / mirror.float().abs().clamp(min=1e-9)).max() + ) + assert frac > 100 * E4M3_MAX_INEXACT_FRACTION, ( + f"control '{name}' (row {row}) moved only {frac:.3f} of bytes — " + "the bit-exact-fraction floor would not have caught it" + ) + assert rel > 4 * E4M3_ULP_RTOL, ( + f"control '{name}' (row {row}) stayed within {rel:.3g} relative — " + "the one-ulp tolerance would not have caught it" + ) + # The same collapse in quant_q_buffer: row 0's roped tail against the + # sequence's last position. 50.3% of its 8192 bytes moved (measured); + # the floor is set well below that because the fraction is a property + # of how far apart two rope angles are, not of the gate. + quant_tail = run.quant_q_buffer[0, :, KV_LORA_RANK:].view(torch.float8_e4m3fn) + mirror_q = env.quantize(env.rope_ref(run.q_pe[0], positions[p - 1])) + frac = float((quant_tail.view(torch.uint8) != mirror_q.view(torch.uint8)).float().mean()) + assert frac > 0.2, f"quant_q single-position control moved {frac:.3f} of bytes" + finally: + env.shutdown() + + +def test_fp8_kv_mtp_position_basis_is_sequence_length() -> None: + """Which length tensor supplies the P consecutive positions: the device + `sequence_length`, not its host twin. The same prepared P = 3 step is + issued twice — once as prepared, once with host_past_key_value_lengths + perturbed to sequence_length - 1 — and both land the rows at the + sequence_length-derived slots with bitwise identical pool, quant_q_buffer + and scheduler buffers.""" + torch.manual_seed(38) + env = _fp8_env() + try: + env.kv_cache_manager.add_dummy_requests([0, 1], token_nums=[40, 31]) + step = _Fp8Step(env, [0, 1], [40, 31], predicted_tokens_per_seq=3) + pool = env.pool_tensor() + + pool.zero_() + step.call() + step.assert_rows_at_sequence_length_slots() + pool_ref = pool.view(torch.uint8).clone() + quant_ref = step.quant_q_buffer.clone() + cu_kv_ref = step.cu_kv_seqlens.clone() + + perturbed = step.metadata.kv_lens_runtime.clone() - 1 + pool.zero_() + step.call(host_past=perturbed) + step.assert_rows_at_sequence_length_slots() + torch.testing.assert_close(pool.view(torch.uint8), pool_ref, rtol=0.0, atol=0.0) + torch.testing.assert_close(step.quant_q_buffer, quant_ref, rtol=0.0, atol=0.0) + torch.testing.assert_close(step.cu_kv_seqlens, cu_kv_ref, rtol=0.0, atol=0.0) + finally: + env.shutdown() + + +def test_fp8_kv_mtp_append_race_control() -> None: + """Harness-blindness control for every pool comparison above. + + The documented defect this surface neighbours: when one call's new tokens + do not map to distinct physical slots, the paged append's writes race and + leave torn cache rows, nondeterministically and with no error. At P > 1 a + single call writes P rows per sequence, so that arming sequence is + reachable here — this replays it by aliasing every block-offset entry onto + one physical page, so all 8 sequences' rows contend for the same two + slots. It fires (6 of 6 identical armed calls left different pool images + on sm_100, 380-1140 bytes apart), and the certified geometry repeated in + the same process reproduces bitwise. Without the armed half, a clean pool + comparison would be indistinguishable from a comparison that cannot see + an append at all.""" + torch.manual_seed(39) + env = _fp8_env(max_batch_size=8) + try: + rids = list(range(8)) + env.kv_cache_manager.add_dummy_requests(rids, token_nums=[40] * 8) + step = _Fp8Step(env, rids, [40] * 8, predicted_tokens_per_seq=2) + pool = env.pool_tensor() + honest = step.metadata.kv_cache_block_offsets + assert honest is not None + aliased = torch.full_like(honest, int(honest[0, 0, 0, 0])) + + armed = [] + for _ in range(4): + pool.zero_() + step.call(block_offsets=aliased) + armed.append(pool.view(torch.uint8).clone()) + assert any(not torch.equal(armed[0], a) for a in armed[1:]), ( + "the armed aliased-slot call reproduced bitwise across 4 runs — " + "this harness cannot see the append race it is meant to detect" + ) + + certified = [] + for _ in range(4): + pool.zero_() + step.call() + certified.append(pool.view(torch.uint8).clone()) + for run in certified[1:]: + torch.testing.assert_close(run, certified[0], rtol=0.0, atol=0.0) + step.assert_rows_at_sequence_length_slots() + finally: + env.shutdown() + + +def test_bf16_mtp_decode_batch() -> None: + """bf16 latent pool at P = 2 and 3 over the engine-default page size: the + roped q lands in fused_q's tail at each row's own position, and each + sequence's P cache rows land at consecutive slots — one sequence per case + inside a page, one straddling the boundary. The last case leads with a + context sequence.""" + for p, q_pe_contiguous in ((2, True), (3, False)): + torch.manual_seed(40 + p) + env = _MlaEnv(tokens_per_block=PAGE32) + try: + cached = [PAGE32 - p, PAGE32 - 1, 2 * PAGE32 - 1] + env.kv_cache_manager.add_dummy_requests([0, 1, 2], token_nums=cached) + _run_and_check( + env, + request_ids=[0, 1, 2], + seq_lens=[p] * 3, + num_contexts=0, + cached_lens=cached, + q_pe_contiguous=q_pe_contiguous, + predicted_tokens_per_seq=p, + ) + finally: + env.shutdown() + torch.manual_seed(44) + env = _MlaEnv(tokens_per_block=PAGE32) + try: + env.kv_cache_manager.add_dummy_requests([0, 1, 2], token_nums=[40, 31, 63]) + _run_and_check( + env, + request_ids=[0, 1, 2], + seq_lens=[40, 2, 2], + num_contexts=1, + cached_lens=[0, 31, 63], + q_pe_contiguous=True, + predicted_tokens_per_seq=2, + ) + finally: + env.shutdown() diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/thop_attention.md b/tensorrt_llm/_torch/staircase/catalog/attention/thop_attention.md new file mode 100644 index 000000000000..534da795ac24 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/attention/thop_attention.md @@ -0,0 +1,1937 @@ +--- +receipts: + sm_100: {status: passed, trtllm: 1.3.0rc21} + sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 96} +--- + +# thop_attention + +**Wraps** `tensorrt_llm.bindings.internal.thop.attention` (one call). + +This is a pybind binding, not a `torch.ops` op; its inclusion in the catalog +is an approved policy exception to the `torch.ops.trtllm.*` entry shape. + +## Semantics + +The fused attention core of one decoder layer over one prepared batch, with +**fully explicit state**: unlike the registered-layer attention entry +points, this binding reads no Python-side registry, no thread-local +metadata, and no layer objects. Every piece of batch state, cache +addressing, and layer config is an argument. + +The certified surface is bf16 activations over a caller-owned paged pool — +bf16 (`quant_mode=0`) or fp8-e4m3 (`quant_mode=128`, see *Standard +configuration over an fp8-e4m3 KV pool* and *MLA over an fp8-e4m3 latent +pool*; the two configurations quantize differently and are certified +separately) — with no other quantization and no spec-dec / sparse / cross / +helix / beam features (their arguments held at the inert values listed in +*Signature*), in two configurations: + +- **Standard** (`is_mla_enable=False`): packed-QKV GQA/MHA self-attention, + paged KV-cache append + masked FMHA, mixed batches in one call, over a + pool holding one layer or several layers side by side (see *Paged KV + cache addressing*), on either context execution path + (`use_paged_context_fmha` `False` or `True` — the latter is what lets a + context call carry a cached prefix, see *Paged-context FMHA*), optionally + with per-query-head **attention sinks** (bf16 pool, causal mask — see + *Attention sinks*) and/or a **sliding window** (`attention_window_size` + smaller than the sequence — see *Sliding window*). +- **MLA** (`is_mla_enable=True`): DeepSeek-style multi-head latent + attention over a paged **latent** cache (kv_factor 1: one latent row of + width `C + R` per token, `C = kv_lora_rank`, `R = qk_rope_head_dim`), + bf16 or fp8-e4m3. + A batch is served by **one call per phase** — `attention_input_type=1` + (context_only) over the context rows, then `=2` (generation_only) over + the generation rows — sharing the same full-batch state tensors. Mixed + `attention_input_type=0` is rejected for MLA. The context phase has two + certified flavors selected by `latent_cache`: fresh prefill + (`latent_cache` given: in-kernel RoPE + latent append) and no-append + (`latent_cache=None`: pure FMHA over explicit K/V — the cached-KV and + chunked-prefill context flows, including `softmax_stats_tensor` output). + Both flavors, the generation call and the mixed-batch pairing of two of + them are certified over **both** pool element types; the chunked + partial-pass pattern of the no-append flavor is bf16-only. The generation + call is additionally certified at `predicted_tokens_per_seq` 1, 2, 3 and 4 + — one speculative-decoding step, where a generation sequence contributes a + `P`-token draft chain instead of a single token. + Mixed `attention_input_type=0` is **not** rejected — nothing in the + binding checks it (the "MLA cannot be mixed" `ValueError` lives one layer + up, in the Python attention backend, which this op bypasses). On a + single-phase batch it is bitwise identical to that phase's own value; on a + genuinely mixed batch it runs **both** phase paths off one set of scalars, + and `head_size`, the `q`/`output` row widths and + `cu_q_seqlens`/`cu_kv_seqlens` can serve only one of them — so at most one + row group comes back right, the other is silently wrong or never written, + and the generation write overruns a context-sized `output`. Pass the + phase's own value. + +### Standard configuration + +One call, for a batch of `ns` sequences — context-phase sequences first +(`host_request_types[i] == 0`, `i < num_contexts`), generation-phase after +(`== 1`). Rows of `q` are: all context tokens sequence by sequence +(`num_ctx_tokens` rows total), then one row per generation sequence +(`predicted_tokens_per_seq == 1`). For each sequence `s` with new-token +count `l_s` and total KV length `kv_s = sequence_length[s]`: + +``` +# 1. slice K and V out of the packed q rows and append them to the paged +# cache at token positions kv_s - l_s .. kv_s - 1 +# (a generation sequence appends one token at position kv_s - 1; a context +# sequence appends l_s of them, l_s < kv_s only on the paged-context path +# — see *Paged-context FMHA*) +# 2. for each new query token i (absolute position p = kv_s - l_s + i), +# each query head h, over kv head h // (Hq/Hkv): +keys = [0, p] # mask_type == 1 (causal) +keys = [p - W + 1, p] # mask_type == 1, W = attention_window_size + # < p + 1 (see *Sliding window*) +keys = [0, kv_s - 1] # mask_type == 0 (padding: no mask) +logits = q[i, h] @ K[keys].T / (q_scaling * sqrt(D)) +out[i, h] = softmax(logits) @ V[keys] # fp32 accum +# with attention_sinks given (see *Attention sinks*): +out[i, h] = softmax(concat(logits, sink[h]))[:-1] @ V[keys] +``` + +Results are written into `output[:num_tokens]` in place; the op returns +nothing. Generation sequences always attend over all `kv_s` cached+current +tokens (a decode step is `l_s = 1`), or over the newest `W` of them when a +sliding window is active. + +Fusion boundary: KV-cache append, masking, softmax, and GQA attention +happen inside the call (plus RoPE when `position_embedding_type` selects +it — not certified for this configuration). The caller owns the QKV +projection, any q/k norm, external RoPE, and the output projection. + +#### Paged-context FMHA + +`use_paged_context_fmha` selects the **context** execution path; it does +not change the generation path. Both values are certified. + +- `False` — the context FMHA reads K and V out of the **packed `q` rows** + and never touches the pool, so a context sequence can only attend to its + own new tokens: `kv_s` must equal `l_s` (the sequence starts empty). +- `True` — the context FMHA reads K and V out of the **paged pool**, after + this call's own append has written the new tokens into it. `kv_s > l_s` + is then legal: the sequence's first `kv_s - l_s` tokens are a cached + prefix that earlier calls wrote, query row `i` sits at absolute position + `p = kv_s - l_s + i`, and the causal mask is bottom-right aligned exactly + as *Standard configuration* states. This is what KV-cache reuse and + chunked prefill need, and an engine turns it on for every call of the + batch as soon as either is enabled (trtllm's defaults enable reuse). + +Two consequences pinned by test: + +- **With nothing cached the flag is observationally inert.** A fresh + prefill and a decode step run from identical state at both values come + back **bitwise identical** in `output` and in pool content (certified on + 32/8/128, 32/4/128 and 64/8/64). The K/V source does move even then — + aliasing every page of a fresh-prefill sequence onto one garbage page + leaves the `False` output bitwise unchanged but moves the `True` output + 31x outside the tolerance band — it is the append running first that + leaves the pool holding exactly the packed rows. +- **With a cached prefix the flag is load-bearing and its absence is + silent.** The same cached-prefix context call at `False` returns without + raising, 51x outside the band around the correct answer. Nothing in the + op checks the combination. + +The append is unchanged by the flag: the new tokens land bit-exactly at +their absolute positions and nothing outside the token ranges is written, +so a repeated identical call is idempotent on both paths. + +Certified surface for `True`: the standard bf16-pool configuration, +`mask_type=1` (causal), `is_fused_qkv=True`, `q_scaling=1.0`, +`tokens_per_block` 32, geometries 32/8/128, 32/4/128 and 64/8/64, on a +single-layer pool and on a real 4-layer `KVCacheManager` pool, with and +without `attention_sinks` and a sliding window. `True` with `mask_type=0` +(padding), with the fp8-e4m3 pool, or with MLA is not certified (MLA +context takes separate q/k/v and its own cached-KV flavor instead). + +#### Head geometry + +Head **counts are a free axis**, not an enumerated list: `num_heads` (`Hq`) +and `num_kv_heads` (`Hkv`) are plain runtime parameters of the same +kernels, and every pair satisfying `Hq % Hkv == 0` is served — MHA +(`Hq == Hkv`), MQA (`Hkv == 1`), and any GQA ratio, power of two or not. +Certified on the bf16 pool: `Hq`/`Hkv` of 32/8 and 32/4 (the two shipped +targets), 64/8, 16/1, 28/4, 12/3, 8/2, 4/4 — ratios 1, 4, 7, 8 and 16, with +non-power-of-2 head counts (28, 12, 3) among them. All of them run the same +kernels at the same accuracy (see *Notes*); nothing distinguishes the +"round" geometries. + +`head_size` (`D`) is the **enumerated** axis: it selects the FMHA kernel, +and only the shipped head dims exist. Certified: 64, 128, 256 (all on the +bf16 pool; 128 also on the fp8-e4m3 pool). A scratch probe additionally saw +80 work (its generation kernel is JIT-generated), while 32, 96 and 192 are +rejected — hard, see *Preconditions*. Treat any `D` outside {64, 128, 256} +as untested rather than merely unsupported. Certified `Hq`/`Hkv`/`D` +triples on the bf16 pool: 32/8/128, 32/4/128, 64/8/64, 16/1/128, 28/4/128, +12/3/128, 8/2/128, 8/2/256, 4/4/64. + +`Hq % Hkv != 0` is the one geometry error the op does not report. A +context-only call at 6/4/128 returned normally, having computed only the +first `(Hq // Hkv) * Hkv = 4` head columns of `output` and left the other +two all-zero; the same geometry does raise once a generation-phase call +selects a decode kernel (`numQHeads should be multiple of numKVHeads`, from +the XQA dispatcher), and `Hq = 7`, `Hkv = 2` aborts the process outright +inside kernel selection. Because the context path fails silently, the +wrapper asserts `num_heads % num_kv_heads == 0` before the call — one of the +entry's two guard asserts (the other guards `attention_sinks`), holding for +every certified configuration including MLA (8/8, 16/16, 32/32 and 128/128 +context, 8/1, 16/1, 32/1 and 128/1 generation). + +#### Attention sinks + +`attention_sinks` is an optional fp32 `[Hq]` CUDA tensor of per-query-head +**sink logits**. When given, each head's softmax gains one extra logit +column that is dropped from the weighted sum: the sink lands in the +denominator only, so the attention weights of every row sum to less than 1 +and the output shrinks accordingly. For query row `i`, head `h`, over the +masked logit range of *Standard configuration*: + +``` +out[i, h] = softmax(concat(logits, sink[h]))[:-1] @ V[:n_keys] + = mass[i, h] * (softmax(logits) @ V[:n_keys]) +mass[i, h] = sigmoid(logsumexp(logits) - sink[h]) # < 1, per (row, head) +``` + +Facts pinned by test, not by inspection: + +- The sink is a logit in the **scaled**-score domain — it is compared + against `q @ K.T / (q_scaling * sqrt(D))`, *after* the softmax scale, not + against the raw dot products. (Certified at `head_size = 64`, + `q_scaling = 1.0`, so a scale of 1/8; the rival "sink joins the unscaled + row" hypothesis sits 12.8x to 41x outside the tolerance band on every + certified case.) +- It is honoured in **both** phases, which take different FMHA kernels: + the context kernel, the short-history decode kernel, and the + multi-CTA-KV decode kernel that folds partial softmax states across CTAs + (`...MultiCtasKvCga...`, selected at 600 and 2000 cached tokens). A + single `attention_input_type=0` call mixing a context and a generation + sequence applies it to both row groups. +- Indexing is by **query** head: sink `h` belongs to `output`'s head column + `h`, not to the kv head `h // (Hq/Hkv)`. +- The KV-cache append is unaffected — the appended pages stay bit-exact and + nothing outside the token ranges is written. +- `attention_sinks=None` and a sink that cannot contribute (`-100.0` or + `-inf` per head) produce **bitwise identical** output in both phases: the + argument's only effect is the one added denominator term. + +Certified surface for sinks: the standard configuration over a bf16 pool +(`quant_mode=0`), `mask_type=1` (causal), `is_fused_qkv=True`, +`q_scaling=1.0`, geometry 64/8/64 — the gpt-oss-120b tp1 layer shape — +both without a window (`attention_window_size = max_seq_len`) and with the +sliding window of *Sliding window* below, and on both context execution +paths including a context call over a cached prefix (see *Paged-context +FMHA*). Sinks combined with `mask_type=0`, with the fp8-e4m3 pool, with +`q_scaling != 1.0`, or with MLA are not certified. + +#### Sliding window + +`attention_window_size` (`W`) is a **per-call token count**. When a +sequence's KV length exceeds it, query row at absolute position `p` attends + +``` +keys [max(0, p - W + 1), p] # exactly min(p + 1, W) keys +``` + +— `W` keys, never `W + 1`; a sequence still shorter than `W` is plain +causal. `W >= every sequence's total KV` is the no-window behavior. The +window applies identically to context rows and generation rows, and with +`attention_sinks` the sink logit joins the denominator of that **windowed** +logit set (each key of a uniform-logit row of `n = min(p + 1, W)` keys +weighs `1 / (n + exp(sink[h]))`). + +Facts pinned by test, not by inspection: + +- The boundary is exact. A one-hot-V probe (all-zero K, so every unmasked + logit is 0 and the softmax is uniform) reads the per-key weights straight + out of `output`: keys below `p - W + 1` come back **bitwise zero**, keys + from `p - W + 1` up to `p` all carry the same `1 / (n + exp(sink[h]))` + with `n = min(p + 1, W)`. Certified at `W = 128` in both phases and at + `W = 33` and `100` — the window is a token count, not a page count, and + may be smaller than `tokens_per_block`. +- **The window is a mask only: cache addressing is unchanged.** The op + appends every new token at its **absolute** position — page + `t // tokens_per_block` of the sequence's `kv_cache_block_offsets` row, + slot `t % tokens_per_block` — and wraps nothing, neither at `W` nor at + `max_seq_len`. Whatever cyclic reuse of pool memory happens is the + caller's page mapping, not the op's (see *Paged KV cache addressing under + a sliding window*). +- A single context prefill longer than `W` is legal in one call: each row + gets its own window. At `use_paged_context_fmha=False` the context FMHA + reads the packed `q` rows, never the pool — aliasing every page of the + sequence onto one garbage page left the context output bitwise identical + (the append still runs, so the pool itself is then wrong). At `True` the + same aliasing moves the output far off: that path reads the pool (see + *Paged-context FMHA*). +- `sequence_length` must stay the **global** cached+new count: it is both + the origin the window is measured back from and the append position. + Capping it to `W` moves the write to slot `W - 1` and truncates the + attended range — silently, no error. +- `attention_window_size` is per call, so layers with different windows + share one pool, one layer→pool mapping and one block-offset table inside + the same batch (certified on a real 4-layer `KVCacheManager` pool with + windows alternating 128 / full — the gpt-oss layer alternation). The + shared pool must then still be sized for the full-attention layers; a + smaller pool per window group needs multi-pool addressing, which is not + certified. + +Certified surface for the window: the standard bf16-pool configuration, +`mask_type=1`, `is_fused_qkv=True`, `q_scaling=1.0`, geometry 64/8/64, +`tokens_per_block` 32, `W` of 128 / 100 / 33, with and without +`attention_sinks`, over prefill (including prefills 1.5-2.3x the window), +decode, mixed batches, histories up to 2000 cached tokens, and 250 decode +steps over a wrapping page ring; at `W = 128` also over a context call +whose cached prefix (200 tokens) runs well past the window, on the +paged-context path. `W` with `mask_type=0`, with the fp8-e4m3 pool, or with +MLA is not certified. + +### Standard configuration over an fp8-e4m3 KV pool + +`quant_mode=128` (the FP8-KV-cache bit alone — a bf16 checkpoint serving +with a quantized cache) switches the pool element type to fp8-e4m3 while +`q` and `output` stay bf16. Both fp32 `[1]` CUDA scale tensors are then +live: `kv_scale_orig_quant` holding `1/s` (applied on write) and +`kv_scale_quant_orig` holding `s` (applied on read), `s` = the KV-cache +scaling factor — `s = 1.0` for uncalibrated bf16-checkpoint serving, the +production default. Passing `None` for either is silently equivalent to +`1.0` rather than an error (see *Preconditions*), so at any other `s` the +tensors are the caller's only defence. The call's three roles were pinned +against per-role references: + +- **Append** (both phases): each new K/V token row lands in the pool as + `e4m3(k * kv_scale_orig_quant)` — fp32 multiply, round-to-nearest cast. + Bit-exact against that mirror for power-of-2 scales (1.0 and 2.0 + certified); at `s = 1.5` ~0.6% of elements differ from the mirror by + exactly one e4m3 ulp (the kernel's intermediate precision differs + slightly from a pure fp32 multiply), never by more. Pages outside the + written ranges stay bitwise untouched. +- **Context FMHA** reads the **bf16 packed `q` rows, not the pool**: + context output matches the bf16-KV fp32 reference at the same + tolerances as the bf16-pool surface (an fp8-round-trip reference is + ~1.4e-1 off). Prefill accuracy is unaffected by cache quantization; + only what later reads the cache sees e4m3. +- **Generation FMHA** attends over dequantized pool rows + `k_hat = v_hat = pool_row * kv_scale_quant_orig`, with `q` additionally + put through the same e4m3 round-trip in-kernel (quantize by + `kv_scale_orig_quant`, dequantize by `kv_scale_quant_orig` — pinned by + a non-power-of-2 scale sweep: only that scale collapses the error). + Decode output matched an fp32 reference over `(q_hat, K_hat, V_hat)` + within `2^-4 * (1 + |ref|)` — observed max abs err 3.1e-2 on + magnitude-~1 outputs, at most 44% of that allowance; the residual is + the kernel's internal e4m3 handling of the softmax probabilities (the + decode kernel computes every MMA operand in e4m3). + +Certified fp8-pool surface: GQA 32/8/128, `tokens_per_block` 32, causal +mask; on a single-layer pool: context-only / generation-only / mixed +batches, scales 1.0, 1.5, 2.0, decode histories up to 200 tokens with +page-boundary crossings; on a 4-layer shared pool (real `DataType.FP8` +manager state, `s = 1.0`): per-layer context-only prefill and +generation-only page-crossing decode with bit-exact per-layer appends and +bitwise sibling-layer isolation (see the multi-layer note in *Notes*; the +non-1.0 scale sweep and mixed batches were not re-run multi-layer). The +bf16-pool-only extras (padding mask, the paged-context path, and every head +geometry other than 32/8/128) were not re-run under fp8. MLA over an fp8 +pool is certified separately, and behaves differently in every role but the +append — see *MLA over an fp8-e4m3 latent pool*. Nothing in this section +transfers to it: in particular "context FMHA reads the bf16 `q` rows, not +the pool" is a fact about **this** configuration only. + +### MLA configuration + +Both phases scale `QK^T` by +`1 / (q_scaling * sqrt(qk_nope_head_dim + qk_rope_head_dim))` — the +generation phase too, despite its `head_size` being `C + R`. `q_scaling` is +a live argument, not a fixed 1.0: it was swept over four values in both +phases and moves only that scale (see the `q_scaling` row of *Geometry and +masking scalars*, and *Notes*). + +#### MLA head count + +`H` = `num_heads` is a plain runtime parameter of the context call, which +runs as MHA (`num_kv_heads = H`, `head_size = nope + R`); the generation +call runs as latent MQA (`num_kv_heads = 1`, `head_size = C + R`) and there +`H` does reach kernel selection (see the JIT paragraph below). Certified: +**`H` = 8, 16, 32 and 128** — 32 is the deepseek-v3-lite tp1 layer +shape, 16 a TP-slice-like count, 8 the tep4 slice of that same checkpoint +(32 attention heads split over 4 tensor-parallel ranks), and 128 the +DeepSeek-R1-0528 layer shape, which an attention-DP rank runs whole because +DP partitions requests rather than heads — all at the same +`C`/`R`/`nope`/`v_head_dim` (512/64/128/128). The first three were run +through the +identical four-case set at `tokens_per_block = 64` — fresh prefill including +a page-crossing sequence, cached-KV no-append context, generation decode over +a page boundary, and a mixed batch served by two phase calls — against the +same fp32 references and the same tolerances, and they land in the same +accuracy envelope: every `H = 8` case sits between 36.8% and 45.9% of the +tolerance allowance, at or below the `H = 32` run of the same case in all +eight comparisons and below the `H = 16` run in all but one (the page-64 +mixed batch, 38.5% against 37.9%) — per-case fractions in *Notes*. `H = 8` +runs that whole four-case set at +`tokens_per_block = 32` as well, on the page-32 geometries of *MLA page +size*, so it is certified at both page sizes across all four cases — the +coverage `H = 32` has. Those four cases all pass `q_lora_rank = 1536`, but +`q_lora_rank = 0` at `H = 8` no longer rests on the `H = 32` measurement: +the inertness sweep of *MLA q-LoRA rank* is driven at `H = 8` as well, which +puts it on the `...VarSeqQ8...` decode kernel this head count compiles +instead of on a kernel a tep4 target never runs. + +**`H = 128`** runs that same four-case set at `tokens_per_block = 32` only, +on the page-32 geometries of *MLA page size*, and runs it **twice**: once on +the unscaled rope table with `q_scaling = 1.0`, so the head count is the only +thing moving against the `H = 8` and `H = 32` runs of the same cases, and +once at the complete DeepSeek-R1-0528 cell — the checkpoint's YaRN-scaled +rope table plus `q_scaling = 1/mscale² ≈ 0.53366` (see *In-kernel RoPE +arguments* and the `q_scaling` row of *Geometry and masking scalars*), which +is the combination a DeepSeek-R1 layer actually passes. It lands in the same +accuracy envelope as the smaller counts: 37.2-63.3% of the tolerance +allowance on the baseline set, 37.2-60.7% at the R1 cell — per-case fractions +in *Notes*. `H = 128` is **not** certified at `tokens_per_block = 64`, and +none of the four cases was re-run at `q_lora_rank = 0` (they all pass 1536, +the R1 value); the rank's inertness rests on the `H = 32` and `H = 8` sweeps. + +`H = 128` is also the **only** head count certified over an fp8-e4m3 latent +pool, where the same four cases run again — baseline cell and R1 cell alike — +against fp8 references at the wider fp8 tolerance (see *MLA over an fp8-e4m3 +latent pool*). + +Everything scales with `H` exactly as the shape tables below state: `q`/`k` +row width `H * (nope+R)`, context `output` width `H * v_head_dim`, the `v` +split view's row stride `H * (nope + v_head_dim)` elements, generation `q` +width `H * (C+R)`, generation `output` width `H * C`, `q_pe` `[G*P, H, R]`, +and `cu_q_seqlens = arange(G+1) * H * P` (`P` = `predicted_tokens_per_seq`, +1 unless a speculative-decoding step). The paged latent pool is untouched by +the head count (its row width is `C + R`). + +The **context** phase JIT-compiles nothing at any certified head count. The +**generation** phase compiles a decode kernel whose name carries a q-tile +that the head count moves at the low end only: `H = 8` takes +`...PagedKvDenseP{32,64}VarSeqQ8Kv128StaticSwapsAbForGen` where `H` = 16, 32 +**and 128** all take `...VarSeqQ16...`. So `H = 8` decode is a genuinely +different compiled kernel, while 16, 32 and 128 are separate compile-cache +entries behind one kernel name — each still pays its own ~5.5 s compile (see +*Notes*). Head counts other than 8, 16, 32 and 128 are untested rather than +known-bad, and a new one may pay a decode JIT compile of its own. + +Two parts of the MLA surface stay certified at `H = 16` only, because no +other head count was run through them: the +`softmax_stats_tensor` output, and the chunked partial-pass pattern of the +no-append flavor (padding mask, per-pass stats, zero-KV sequences, +downstream merge). Both are independent of the head count in the shape +tables — the stats tensor is `[>= Tc, >= H, 2]` — but neither has a run at +any other head count behind it. + +#### MLA page size + +`tokens_per_block` is certified at **32 and 64** for MLA. 32 is what a +`KvCacheConfig` at its defaults produces, so it is the page size an MLA +target hits before anyone tunes it; 64 is the tuned value. Like the head +count this is a kernel-selection axis for the generation phase only: it +JIT-compiles a decode kernel whose name carries the page size — +`...PagedKvDenseP32VarSeqQ16Kv128StaticSwapsAbForGen` at 32 versus +`...P64...` at 64 (see *Notes*) — while the context phase triggers no JIT +compile at either size. + +Certified at page 32: all three MLA call flavors at `H = 32` (the +deepseek-v3-lite tp1 layer shape), at `H = 8` (its tep4 slice) and at +`H = 128` (the DeepSeek-R1-0528 layer shape, on the baseline rope/scale and +again at the full R1 cell), including +the mixed-batch pairing of two of them at all three counts, plus the +generation flavor also at `H = 16` because that is the one flavor where the +page size changes the compiled kernel and the decode compile cache is keyed +by head count. The page-32 geometry is chosen +for 32-slot pages rather than reused from the page-64 cases — a 64-token +prefill fills a 64-slot page exactly but spans two 32-slot pages, so the +crossing arithmetic differs: + +- **fresh prefill**: sequences of 96 (three exact pages — a 64-slot page + never ends there) and 33 (a one-token spill into a second page), append + checked page by page; +- **generation decode**: a 64-token history filling two pages exactly, so + the first decode token opens page 2, alongside a 31-token history whose + first decode token takes page 0's last slot and whose second opens page + 1 — a decode read range crossing a boundary mid-case, two steps, all + four head counts; +- **mixed batch**: a 64-token history decoded next to a fresh 33-token + context sequence, the two phase calls sharing one page-32 offsets table; +- **cached-KV no append**: prefixes of 96 / 31 / 0 reaching KV lengths of + 128 (four exact pages), 40 and 25, with the pages reserved as production + reserves them and the pool verified bitwise untouched. + +Page 32 is also the only size certified over an fp8-e4m3 latent pool, where +all four geometries above run again at `H = 128` (see *MLA over an fp8-e4m3 +latent pool*). + +Accuracy at page 32 lands on the page-64 magnitudes (per-case fractions in +*Notes*, including the one page-32 case that sets the file's headroom +ceiling — a flavor that never touches the pool, so its number tracks its KV +geometry rather than the page size), the paged latent append holds to the +same gate, and the generation call still writes nothing. Page sizes other +than 32 and 64 are untested for MLA rather than known-bad; the standard +configuration is certified at 32 only. + +#### MLA q-LoRA rank + +`q_lora_rank` names the rank of the checkpoint's `q_a_proj`: DeepSeek-V3 +down-projects the query through a rank-1536 LoRA pair, while a checkpoint +whose config carries `"q_lora_rank": null` (deepseek-v3-lite) has no q-LoRA +at all and projects the query directly. Certified: **1536 and 0**. + +The argument is **inert**: no observable of any certified MLA call depends +on it. That matches where the query projection sits — entirely outside this +call, which receives finished `q` rows (absorbed or not) whether or not a +q-LoRA produced them. Measured on sm_100 by sweeping +`1536 → 0 → 1536 → 4096` (the last larger than `C`, so arithmetic that +touched the value would move something) over identical inputs, in all three +MLA call flavors at `tokens_per_block = 32`, on the page-32 geometries of +*MLA page size*, at **two** of the certified head counts — `H = 32` (the +deepseek-v3-lite tp1 layer shape) and `H = 8` (its tep4 slice), two +independent sweeps with their own seeds; `H = 128` and `H = 16` were not +swept. Both swept counts are measured rather than one inferred from the other +because the generation phase compiles a different decode kernel at each +(`...P32VarSeqQ16...` at `H = 32`, `...P32VarSeqQ8...` at `H = 8` — see +*MLA head count*), so an `H = 32` sweep alone is evidence about a kernel a +tep4 target never runs. Not swept: the mixed batch (a pairing of two of +those flavors), `H = 16`, `tokens_per_block = 64` (both sweeps are +page-32), and the no-append flavor's chunked partial-pass pattern with +`softmax_stats_tensor` — the sweep takes its one-shot cached-KV pattern. +Every observable came back **bitwise identical** at both head counts: the +`output` rows, the whole paged latent pool, the in-place-roped `q`/`k` of +the fresh-prefill flavor, and the byte count the op resized `workspace_` +to. The repeated `1536` is the run-to-run determinism control that makes +those comparisons mean something. Kernel selection does not see it either: +a process running only the two sweeps compiled exactly one decode kernel +per head count — one `...VarSeqQ16...` serving the four `H = 32` rank +values and one `...VarSeqQ8...` serving the four `H = 8` ones — where a +real selection axis (the head count) compiles once per value; see *Notes*. + +`None` is **not** an accepted value on the MLA path even though the +signature admits it: the C++ unwraps the optional unconditionally and the +call raises `RuntimeError: bad optional access`. Pass the checkpoint's rank, +or `0` when it has none. + +**Context phase — fresh prefill** (`attention_input_type=1`, separate +`q`/`k`/`v`, `is_fused_qkv=False`, `latent_cache` given): certified context +sequences of this flavor have no cached prefix, so token `i` of a sequence +sits at absolute position `i`. Per context sequence of length `L`, with +per-token latent row `latent_cache[t] = [ckv_t | k_pe_t]`: + +``` +# GPT-J (interleaved-pair) RoPE at position i, fp32 math, from the +# duplicated-layout rotary_cos_sin table: +# rope(x)[2d] = x[2d]*cos_d - x[2d+1]*sin_d +# rope(x)[2d+1] = x[2d]*sin_d + x[2d+1]*cos_d +q'[i, h] = [ q[i, h, :nope] | rope_i(q[i, h, nope:]) ] # written back into q +k'[i, h] = [ k[i, h, :nope] | rope_i(k_pe_i) ] # k_pe broadcast to all + # heads, written back into k +cache_row(seq, i) = [ ckv_i | rope_i(k_pe_i) ] # paged latent append +out[i, h] = softmax_causal(q'[i, h] @ K'.T * scale) @ V # headDim nope+R, + # headDimV v_head_dim +``` + +`output[:num_ctx_tokens]` is written. The rope slice of the incoming `k` +(`k[..., nope:]`) is ignored (the kernel fills it from `latent_cache`); +the nope slices of `q`/`k` are left bitwise intact, and `latent_cache` is +read-only. Fusion boundary (context): q_pe/k_pe RoPE, latent-cache append, +and causal FMHA happen inside; the caller owns the projections that build +`q`, `k_nope`, `v`, and `latent_cache`, and the output projection. + +**Context phase — no append** (`attention_input_type=1`, separate +`q`/`k`/`v`, `is_fused_qkv=False`, `latent_cache=None`): the in-kernel +RoPE and cache append are skipped entirely — the call is a pure masked +FMHA over caller-supplied K/V, used for context over a cached KV prefix +and for chunked-prefill partial passes. Per context sequence `s` with +`n = host_context_lengths[s]` q rows and `m = sequence_length[s]` K/V rows +(this call's KV range; `m` is independent of `n`): + +``` +out[i, h] = softmax(q[i, h] @ K[:j_i].T * scale) @ V[:j_i] # fp32 accum + j_i = (m - n) + i + 1 # mask_type == 1: causal, bottom-right aligned + # (query i sits at absolute position m - n + i) + j_i = m # mask_type == 0: padding (whole KV range) +``` + +and, when `softmax_stats_tensor` is given, per (q row `t`, head `h`) over +the same masked logit range (natural-log domain, fp32): + +``` +softmax_stats[t, h] = (max_j logit_j, sum_j exp(logit_j - max)) +``` + +Nothing is rotated, appended, or mutated except `output` (and the stats +tensor): `q` is consumed as-is (production pre-rotates its rope tails +upstream — a sibling op, `trtllm::mla_rope_append_paged_kv_assign_q`, +exists for that step and the latent append), `k` carries rope(k_pe) in its +per-head tails, and `q`, `k`, `v`, and the paged pool were verified +bitwise untouched by the call. This flavor is certified on pure-context +batches (`ns == num_contexts`). Two usage patterns are certified: + +- **One-shot cached-KV context** (causal): K/V cover each sequence's full + `[cached + new]` range (`sequence_length = cached + new >` q length), + one call computes the final context output. +- **Chunked partial passes**: per chunk loop, K/V cover one chunk slice of + the cached range per sequence — `sequence_length[s]` = that sequence's + chunk length, `0` allowed (the sequence sits the loop out; its output + and stats rows are then **undefined** and must be skipped downstream) — + with `mask_type=0` and stats emitted; then one final causal pass over + the new tokens only (`sequence_length == context_lengths`), stats + emitted. Each pass's output/stats pair is folded into the running + full-range result by a sibling merge op + (`trtllm::merge_chunked_attention_for_mla`) under the production + copy/merge/skip plan — verified end to end: the merged output and stats + equal single-pass attention over the full `[cached + new]` range. + +**Generation phase** (`attention_input_type=2`, `q` = absorbed "fused q", +`is_fused_qkv=True`, `k`/`v` `None`): a pure latent-MQA **read** of the +paged cache. With `P = predicted_tokens_per_seq`, each generation sequence +contributes `P` query rows — one for an ordinary decode step, `P > 1` for a +speculative-decoding step whose `P` rows are that sequence's draft chain. +The rows are **token-major within a sequence**: row `n` is sequence +`g = n // P`'s `t = n % P`-th token, and it sits at absolute position +`L_g - P + t`, where `L_g = sequence_length[num_contexts + g]` is the total +KV length **including all `P` of this step's tokens**. Per row: + +``` +K = cache rows [0, L_g - P + t] # [L_g-P+t+1, C+R] +V = K[:, :C] +out[n, h] = softmax(q[n, h] @ K.T * scale) @ V # [H, C] per token +``` + +`output[:G*P]` is written (`G` = generation-sequence count). At `P = 1` the +key range is the whole cache, which is the ordinary decode rule. At `P > 1` +the **within-block mask is causal and bottom-right aligned against the +sequence's own KV length**: draft row `t` sees the cache up to and including +its own position and **not** rows `t+1 .. P-1` of its own block, whose tokens +do not exist yet. + +That is measured rather than inferred, and by a readout rather than a +tolerance: with each cached latent row's `ckv` half set to a one-hot vector, +the decode's `V` is that one-hot and the output row becomes the attention +weight vector itself, so the attended key set can be read off column by +column. At `P` 1-4, `H = 128`, page 32, over both pool element types, every +row weighted exactly `[0, L_g - P + t]` and every other column came back +**bitwise zero** — including the 1 to 3 future-sibling columns a full +within-block mask would have weighted at 0.023-0.039 each. `mask_type` does +not change this (see its row in *Geometry and masking scalars*), and on +sm_100 no mask tensor is involved at all: for a linear-tree draft the Python +backend forces `is_spec_decoding_enabled` off there, so +`spec_decoding_packed_mask` and its siblings are `None` and +`predicted_tokens_per_seq` is the only thing producing the mask. + +No RoPE is applied and nothing is written to the cache: the generation-phase +preprocessing — the RoPE of `q_pe` at each row's own position, the latent +append of all `P` rows at positions `L_g - P .. L_g - 1`, and the +decode-scheduler-buffer fill — happens **before** this call (a sibling op, +`trtllm::mla_rope_generation`, exists for it). + +Where the roped `q_pe` has to **land** is pool-dependent, and the line above +describes only the bf16 case. Over a **bf16** latent pool this call reads its +query from `q`, so the preprocessing must leave the roped `q_pe` in fused q's +tail (`q[..., C:]`) — that is the layout certified here. Over an +**fp8-e4m3** pool this call does not read `q` at all: the query arrives +already quantized in `quant_q_buffer` (see *MLA over an fp8-e4m3 latent +pool*), so whether a producer ever materializes a roped tail in `q` is +outside what this op observes. What it requires on that path is only that +`quant_q_buffer` hold the right bytes. + +The `latent_cache` and `q_pe` arguments must be non-None (presence-checked) +but their data is **not consumed** — verified by passing garbage in both +while output, pool, and `q` stayed correct/untouched, at `P` 1 through 4; they +are not mutated either. The caller owns the absorbed-q BMM that fills +`q[..., :C]` and the `v_b_proj` BMM applied to the `[G*P, H*C]` output +afterwards. + +#### MLA over an fp8-e4m3 latent pool + +`quant_mode=128` switches the **latent pool's** element type to fp8-e4m3 +(one byte per element; the pool is `[num_blocks, tokens_per_block, C+R]` +e4m3) while `q`, `k`, `v`, `latent_cache`, `q_pe` and `output` all stay +bf16. Certified at the DeepSeek-R1-0528 layer geometry — `H = 128`, +`tokens_per_block = 32`, `C/R/nope/v = 512/64/128/128`, +`q_lora_rank = 1536`, causal mask, single-layer pool, one call per phase — +for **all three MLA call flavors** (fresh-prefill context, no-append context, +generation) and for the **mixed batch** that pairs two of them, at KV-cache +scaling factor `s = 1.0` (the production value: a +`kv_cache_quant_algo: FP8` checkpoint whose per-layer `k_scale`/`v_scale` +are 1.0). `s` enters as the same +two fp32 `[1]` CUDA tensors the standard configuration uses — +`kv_scale_orig_quant` = `1/s`, `kv_scale_quant_orig` = `s` — and passing +`None` for both is again silently the `s = 1.0` path (bitwise identical +output and pool, verified). `s` = 1.5 and 2.0 are swept beside 1.0 on the +fresh-prefill, generation and no-append flavors; the mixed batch runs +`s = 1.0` only. + +Two rope/scale cells are covered, and not uniformly. The **fresh-prefill** +and **generation** flavors run both: the unscaled theta-10000 table at +`q_scaling = 1.0`, and the complete DeepSeek-R1-0528 cell — that checkpoint's +YaRN-scaled `rotary_cos_sin` table together with +`q_scaling = 1/mscale² ≈ 0.53366` (see *In-kernel RoPE arguments*). The +**no-append** flavor, the **mixed batch** and the `quant_mode`-bit checks run +the R1 cell only. Both axes are certified over the bf16 pool as axes of their +own; the fp8 runs pin them **in combination with the pool**, which is what a +DeepSeek-R1 layer actually passes and what separate certifications do not +cover. Both reach the +fp8 path and both are gated: an fp8 R1-cell prefill sits 12.7x outside a +reference built at `q_scaling = 1.0` and 3.7x outside one built from the +unscaled table (the matching reference uses 0.40 of the same band); its +decode sits 4.2x outside the `q_scaling = 1.0` reference and the no-append +flavor 10.8x. The rope table is +pinned far harder by the append than by either output gate: the appended +`k_pe` half is bit-exact against a table-aware mirror while a mirror built +from the unscaled table differs in 2336 of the 6144 e4m3 bytes of a 96-token +request (38.0%, against a gate that allows 6). + +**Only the append resembles the standard configuration's fp8 surface.** The +two read paths are different code and different arguments; read them here, +not there. + +- **Append** (context call, both halves of the row): each new latent row + lands as `e4m3(row * kv_scale_orig_quant)` — fp32 multiply, round-to- + nearest cast. Bit-exact against that mirror at every scale run (1.0, 1.5, + 2.0), for the bitwise-copied `ckv` half **and** the in-kernel-roped `k_pe` + half alike: e4m3's 3-bit mantissa absorbs the fp32 evaluation-order + difference that costs the bf16 pool its ~8-in-500 000 one-ulp elements. + Pages outside the written rows stay bitwise zero, including when a + request's pages are scattered and out of order. +- **Context FMHA runs in fp8, not bf16 — in *both* context flavors.** The op + quantizes `q`, `k` and `v` to e4m3 itself and computes the masked FMHA on + those operands, so **prefill accuracy is affected by cache quantization** — + not because the context reads the pool back (the `s != 1.0` runs match a + model built from the caller's own bf16 operands quantized at 1.0, which is + not what the pool rows hold), but because the op quantizes its inputs. + This is not inherited from one flavor to the other: it was measured + separately on the fresh-prefill flavor (`latent_cache` given) and on the + no-append flavor (`latent_cache=None`), which appends nothing and whose K/V + arrive **already dequantized** from an fp8 pool — it quantizes them again + anyway. + Pinned bitwise on each: with the softmax + collapsed onto one key, the output row comes back as `e4m3(V_row)` + exactly (max abs deviation 0, 99.98% of the raw bytes equal — the rest + signed zeros), where a bf16 path would return `V_row`, 60-63 bf16 ulps + away. On + realistic inputs each flavor matches an fp32 reference over e4m3-rounded + q/k/v within `2^-3 + 2^-4 * |ref|` (fresh prefill 40-49% of it, no-append + 41%); the fresh-prefill run sits 26x (baseline cell) to 49.8x (R1 cell) + outside the bf16 band around the un-quantized reference — the mixed batch's + context call 58.1x — while the no-append run sits 2.2x + outside the *fp8* band around it. + The quantization uses **scale 1.0** — `kv_scale_orig_quant` is *not* + applied to `q`/`k`/`v` — while the FMHA nevertheless applies the + read-side dequantization factors that scale implies: `s^2` on the softmax + scale and `s` on the output. The two only cancel at `s = 1.0`, so + **both context flavors are correct at `s = 1.0` alone**. At `s` = 1.5 and + 2.0 the output is what the `s = 1.0` math would give with those two factors + applied (measured within 0.58x / 0.84x of the same band on the fresh flavor + and 0.86x on the no-append one), i.e. 21x-48x + away from the correct result, silently; the no-append peaked readout returns + exactly `s * e4m3(V_row)`, bit for bit, which is where the output factor is + read off directly. There is no caller-side knob that + fixes it: `quant_scale_qkv`, the argument that looks like it would set the + quantization scale, is inert here — a scratch probe passing `[1/s]` in it + at `s = 2.0` returned bitwise-identical output. +- **The no-append flavor still touches no pool.** Under `quant_mode=128` as + under `0`, that call reads and writes nothing but `output`: `q`, `k`, `v` + and the whole e4m3 pool came back bitwise intact from a call whose pages + held real cached content. Its operands in the certified run were built the + way a target builds them — the cached prefix written into the pool as + `e4m3(row * kv_scale_orig_quant)`, read back as + `bf16(float(cache_byte) * kv_scale_quant_orig)` and pushed through a + `kv_b_proj`-shaped matmul, with only the new tokens fresh — so the flavor + is certified on operands that have been through the pool, not just on + operands of the right shape. +- **Generation FMHA takes its query from `quant_q_buffer`, and its two + scales from `mla_bmm1_scale` / `mla_bmm2_scale`.** `q` is not read at all + on this path (replacing it with garbage left the output bitwise + identical), and neither kv scale tensor is read either (dropping both is + bitwise identical at `s = 2.0`, where a read-side dequant would be + glaring). Mechanically the kernel computes, over the **raw** e4m3 numbers + in `quant_q_buffer` and in the pool: + + ``` + out[g, h] = softmax2(qq[g, h] @ Kraw.T * mla_bmm1_scale[1]) + @ Kraw[:, :C] * mla_bmm2_scale[0] + # softmax2 = exp2-domain softmax, so mla_bmm1_scale[1] is the natural-log + # scale times log2(e); mla_bmm1_scale[0] is never read (zeroing it changes + # nothing, zeroing [1] flattens the softmax) + ``` + + The caller therefore owns the dequantization. Filling the three buffers + like this — + + ``` + quant_q_buffer = e4m3(fused_q * kv_scale_orig_quant) # [G*P, H, C+R] + x = 1 / (q_scaling * sqrt(nope + R)) * s**2 + mla_bmm1_scale = [x, x * log2(e)] # fp32 [2] + mla_bmm2_scale = [s] # fp32 [1] + ``` + + — makes the call compute exactly latent MQA over the e4m3 round trip of + both operands (`s^2` undoing the `1/s` on the query and on `K`, `s` the + one on `V = K[:, :C]`), at the plain softmax scale: + + ``` + q_hat = e4m3(q * 1/s) * s K_hat = pool_row * s V_hat = K_hat[:, :C] + out[n, h] = softmax(q_hat[n, h] @ K_hat.T / (q_scaling*sqrt(nope+R))) @ V_hat + ``` + + with `K_hat` cut to row `n`'s own key range when `predicted_tokens_per_seq` + is above 1 (see *Semantics*), and the two bmm scales unchanged by it — they + are per-batch scalars, so an fp8 MTP decode needs no per-draft-row + rescaling. This whole recipe was re-run at `P` = 1, 2, 3 and 4. + + and that holds at **every** scale certified — 1.0, 1.5 and 2.0 all matched + this reference within 18% of the fp8 band, the non-power-of-two scale + included. Leaving `s` out of the two bmm + scales — the mistake the standard configuration's "dequantized on read" + wording invites — lands 2.9x (`s = 1.5`) and 3.6x (`s = 2.0`) outside it. + All three buffers are presence-checked: `None` raises + `Assertion failed: quant_q_buf is nullptr. (attentionOp.cpp:1050)`, + `bmm1_scale is nullptr. (:1051)`, `bmm2_scale is nullptr. (:1052)`. + Their *contents* are unchecked, and `quant_q_buffer` is read through a raw + pointer (an identical `uint8` view of the same bytes behaves identically). + The generation call still writes nothing: the pool comes back bitwise + unchanged. + + That block is a recipe for **this** op's caller — the certified runs passed + exactly those bytes, built in plain torch — and not a description of how a + producer arrives at them. In particular `quant_q_buffer` is written in terms + of the finished bf16 fused q this call would have read on the bf16 path; + a producer may instead assemble it from the absorbed `q[..., :C]` plus a + freshly roped `q_pe`, and never materialize a roped tail in `q` at all + (which is what the sibling preprocessing op does here). Both routes have to + land on the same bytes; only the bytes are this op's business. + +**Quantization bits outside the KV-cache group ride along unread.** A target +derives `quant_mode` from its checkpoint's quant config and will not get a +bare `128`: an nvfp4 checkpoint served with an fp8 KV cache produces `1152` +(`FP8_KV_CACHE | FP8_1x128_128x128`), and fp8-QDQ weights add `FP8_QDQ` for +`384`. Both were run against a bare `128` on all three MLA flavors at the R1 +cell and are **bit-identical** — output and pool alike. Nothing validates the +extra bits in either direction; this is a measurement, not a guarantee about +bits that were not tried (`64` = INT8 KV and `8192` = NVFP4 KV were not tried +on this op). + +The generation flavor is additionally certified over this pool at +`predicted_tokens_per_seq` = 1, 2, 3 and 4 — one speculative-decoding step per +call — including the mixed batch that pairs it with a context call. See the +`predicted_tokens_per_seq` and `mask_type` rows of *Geometry and masking +scalars* and the generation block of *Semantics*; the context flavors run at +1 only. + +Not certified over the fp8 latent pool: the no-append flavor's **chunked +partial-pass pattern** (`mask_type=0`, `softmax_stats_tensor`, zero-KV +sequences — only its one-shot cached-KV pattern is certified), +`tokens_per_block = 64`, head counts other than 128, multi-layer latent +pools, `predicted_tokens_per_seq` above 4, and `s != 1.0` in either context +flavor (measured-wrong, see above, rather than untested). + +## Signature + +```python +def thop_attention( + q, k, v, output, output_sf, workspace_, + sequence_length, host_past_key_value_lengths, host_total_kv_lens, + context_lengths, host_context_lengths, host_request_types, + max_context_q_len_override, + kv_cache_block_offsets, host_kv_cache_pool_pointers, + host_kv_cache_pool_mapping, cache_indirection, + kv_scale_orig_quant, kv_scale_quant_orig, out_scale, + rotary_inv_freq, rotary_cos_sin, latent_cache, q_pe, + block_ids_per_seq, attention_sinks, + is_fused_qkv, update_kv_cache, predicted_tokens_per_seq, + local_layer_idx, num_heads, num_kv_heads, head_size, tokens_per_block, + max_num_requests, max_context_length, max_seq_len, + attention_window_size, beam_width, mask_type, quant_mode, q_scaling, + position_embedding_type, rope_dim, rope_base, rope_scale_type, + rope_scale, rope_short_m_scale, rope_long_m_scale, rope_max_positions, + rope_original_max_positions, use_paged_context_fmha, + attention_input_type, is_mla_enable, chunked_prefill_buffer_batch_size, + q_lora_rank, kv_lora_rank, qk_nope_head_dim, qk_rope_head_dim, + v_head_dim, rope_append, mrope_rotary_cos_sin, mrope_position_deltas, + helix_position_offsets, helix_is_inactive_rank, attention_chunk_size, + softmax_stats_tensor, is_spec_decoding_enabled, use_spec_decoding, + is_spec_dec_tree, spec_decoding_generation_lengths, + spec_decoding_position_offsets_for_cpp, spec_decoding_packed_mask, + spec_decoding_bl_tree_mask_offset, spec_decoding_bl_tree_mask, + spec_bl_tree_first_sparse_mask_offset_kv, + sparse_kv_indices, sparse_kv_offsets, sparse_attn_indices, + sparse_attn_offsets, sparse_attn_indices_block_size, + # keyword tail, defaults mirror the binding: + num_sparse_topk=None, sparse_attn_kv_lens=None, + skip_softmax_threshold_scale_factor_prefill=None, + skip_softmax_threshold_scale_factor_decode=None, skip_softmax_stat=None, + cu_q_seqlens=None, cu_kv_seqlens=None, fmha_scheduler_counter=None, + mla_bmm1_scale=None, mla_bmm2_scale=None, quant_q_buffer=None, + flash_mla_tile_scheduler_metadata=None, flash_mla_num_splits=None, + sage_attn_num_elts_per_blk_q=0, sage_attn_num_elts_per_blk_k=0, + sage_attn_num_elts_per_blk_v=0, sage_attn_qk_int8=False, + num_contexts=0, num_ctx_tokens=0, trtllm_gen_jit_warmup=False, + aux_kv_cache_pool_ptr=None, is_cross=False, cross_kv=None, + relative_attention_bias=None, relative_attention_max_distance=0, + spec_decoding_target_max_draft_tokens=None, quant_scale_qkv=None, + dsv4_inv_rope_cos_sin_cache=None, enable_dsv4_epilogue_fusion=False, +) -> None +``` + +Exact per-argument types are in the wrapper; the binding is keyword-callable +and the wrapper forwards every argument by keyword. + +With `Hq` = `num_heads`, `Hkv` = `num_kv_heads`, `D` = `head_size`, +`ns` = batch size, `T` = total new tokens in the batch +(`num_ctx_tokens + (ns - num_contexts) * predicted_tokens_per_seq`). `T` is +the standard configuration's `q`/`output` row count; an **MLA** batch is +served by one call per phase, so its context call has `num_ctx_tokens` rows +and its generation call `(ns - num_contexts) * predicted_tokens_per_seq`. +Certified MLA dims (DeepSeek-V3 / R1 geometry): `C = kv_lora_rank = 512`, +`R = qk_rope_head_dim = 64`, `nope = qk_nope_head_dim = 128`, +`v_head_dim = 128`, at `H` = 8, 16, 32 **or** 128 query heads — a +plain runtime parameter for the context call as in the standard +configuration, and a decode-kernel selection axis for the generation call +(see *MLA head count*); every MLA shape below is written in terms of `H`. + +### Data tensors — standard configuration + +| Argument | Shape / value | Dtype | Device | +|---|---|---|---| +| `q` | `[T, (Hq + 2*Hkv) * D]`, packed per-token `q \| k \| v`, contiguous (both pool dtypes: activations stay bf16) | bf16 | CUDA | +| `k`, `v` | `None` — K/V ride inside packed `q` (`is_fused_qkv=True`; separate-KV mode not certified for this configuration) | — | — | +| `output` | `[T, Hq * D]`, contiguous; rows `[:T]` overwritten in place (both pool dtypes) | bf16 (same as `q`) | CUDA | +| `kv_scale_orig_quant` | `quant_mode=0`: `None`; `quant_mode=128`: `[1]` holding `1/s` — **required for any `s != 1.0`** (`None` is silently read as `1.0`, never rejected) | fp32 | CUDA | +| `kv_scale_quant_orig` | `quant_mode=0`: `None`; `quant_mode=128`: `[1]` holding `s` — **required for any `s != 1.0`** (same silent `1.0` fallback) | fp32 | CUDA | +| `attention_sinks` | `None` (no sink), or exactly `Hq` values, **contiguous** — one sink logit per query head (see *Attention sinks*; certified on the bf16 pool with `mask_type=1`). Shape beyond the element count is ignored (`[Hq]`, `[1, Hq]`, `[Hq, 1]` all behave identically) | fp32 | CUDA | +| `latent_cache`, `q_pe` | `None` | — | — | +| `cu_q_seqlens`, `cu_kv_seqlens`, `fmha_scheduler_counter` | `None` (the op fills internal ones when absent) | — | — | + +### Data tensors — MLA context call, fresh prefill (`latent_cache` given) + +`Tc` = `num_ctx_tokens` (this call's `q` rows), `Hq == Hkv == H`, +`head_size = nope + R` (192 certified). + +| Argument | Shape / value | Dtype | Device | +|---|---|---|---| +| `q` | `[Tc, H * (nope+R)]`, contiguous; rope slice of each head **overwritten in place** with rope(q_pe) | bf16 | CUDA | +| `k` | `[Tc, H * (nope+R)]`, contiguous; incoming rope slice ignored (may be uninitialized) and **overwritten in place** with rope(k_pe); nope slice consumed as K | bf16 | CUDA | +| `v` | `[Tc, H * v_head_dim]` row view whose **row stride is `H * (nope + v_head_dim)` elements** — the `[.., H*nope:]` split of a packed `[Tc, H*(nope+v)]` kv-projection buffer (certified); the context FMHA hard-codes this stride, so a fully contiguous `[Tc, H*v_head_dim]` tensor is misread (see *Notes*) | bf16 | CUDA | +| `output` | `[Tc, H * v_head_dim]`, contiguous; rows `[:Tc]` overwritten | bf16 | CUDA | +| `latent_cache` | `[Tc, C + R]` per-token `[ckv \| k_pe]`, contiguous; read-only | bf16 | CUDA | +| `q_pe` | `None` | — | — | +| `kv_scale_orig_quant`, `kv_scale_quant_orig` | `quant_mode=0`: `None`; `quant_mode=128`: `[1]` holding `1/s` and `s`. Only the append consumes `1/s`; the context FMHA applies `s^2`/`s` to a computation it quantized at 1.0, so **`s = 1.0` (or `None`) is the only correct value here** — see *MLA over an fp8-e4m3 latent pool* | fp32 | CUDA | +| `mla_bmm1_scale`, `mla_bmm2_scale`, `quant_q_buffer` | `None` (the context call neither checks nor reads them at either `quant_mode`) | — | — | +| `softmax_stats_tensor` | `None` (stats certified only for the no-append flavor) | — | — | +| `cu_q_seqlens`, `cu_kv_seqlens`, `fmha_scheduler_counter` | `None` (the op fills internal ones when absent) | — | — | + +### Data tensors — MLA context call, no append (`latent_cache=None`) + +`Tc` = `num_ctx_tokens` (this call's `q` rows); `Tkv` = sum of +`sequence_length[s]` over the context sequences = total K/V rows, the +per-sequence KV ranges concatenated in batch order. + +| Argument | Shape / value | Dtype | Device | +|---|---|---|---| +| `q` | `[Tc, H * (nope+R)]`, contiguous; consumed as-is (rope tails prepared upstream) and **not mutated** | bf16 | CUDA | +| `k` | `[Tkv, H * (nope+R)]`, contiguous, read-only; per-head tails hold rope(k_pe) broadcast to all heads (caller-built) | bf16 | CUDA | +| `v` | `[Tkv, H * v_head_dim]` row view, read-only; **row stride `H * (nope + v_head_dim)` elements required** (same hard-coded stride as the fresh flavor — certified as the split of packed `[Tkv, H*(nope+v)]` buffers; contiguous rows are misread) | bf16 | CUDA | +| `output` | `[Tc, H * v_head_dim]`, contiguous; rows `[:Tc]` overwritten (rows of zero-KV sequences: undefined) | bf16 | CUDA | +| `latent_cache` | `None` — selects this flavor (skips in-kernel RoPE + append; the paged pool is neither read nor written, at either `quant_mode`) | — | — | +| `q_pe` | `None` | — | — | +| `kv_scale_orig_quant`, `kv_scale_quant_orig` | `quant_mode=0`: `None`; `quant_mode=128`: `[1]` holding `1/s` and `s`. This flavor appends nothing, so `1/s` reaches nothing at all, while the FMHA still applies `s^2`/`s` to a computation it quantized at 1.0 — **`s = 1.0` (or `None`) is the only correct value here**, exactly as for the fresh-prefill flavor. See *MLA over an fp8-e4m3 latent pool* | fp32 | CUDA | +| `softmax_stats_tensor` | `None`, or `[>= Tc, >= H, 2]` contiguous, rows `[:Tc]` overwritten per (token, head) with `(max, sum)` as in *Semantics* (rows of zero-KV sequences: undefined). Presence-checked by the op: fp32, dim 0 `>=` the call's q rows, dim 1 `>= num_heads`, dim 2 `== 2`. Certified on the bf16 pool only | fp32 | CUDA | +| `mla_bmm1_scale`, `mla_bmm2_scale`, `quant_q_buffer` | `None` (this flavor neither checks nor reads them at either `quant_mode`) | — | — | +| `cu_q_seqlens`, `cu_kv_seqlens`, `fmha_scheduler_counter` | `None` (the op fills internal ones when absent) | — | — | + +### Data tensors — MLA generation call + +`G` = generation-sequence count, `P` = `predicted_tokens_per_seq` (query rows +per generation sequence: 1 for an ordinary decode step, `P > 1` for a +speculative-decoding step — 1, 2, 3 and 4 certified), `head_size = C + R` +(576 certified), `Hkv = 1`. Every per-row tensor is `G*P` rows tall and +**token-major within a sequence**: row `n` is sequence `n // P`'s +`n % P`-th token. + +| Argument | Shape / value | Dtype | Device | +|---|---|---|---| +| `q` | `[G*P, H * (C+R)]` fused q (absorbed q_nope in `[..., :C]`, rope(q_pe) in `[..., C:]`, both prepared before the call — at `P > 1` each row must have been roped at **its own** absolute position `L_g - P + n % P`; this call applies no rope and cannot tell), contiguous; not mutated. Under `quant_mode=128` it is **not read either** — the query comes from `quant_q_buffer` below — but it is still the tensor whose shape and dtype the call is configured around | bf16 | CUDA | +| `k`, `v` | `None` (`is_fused_qkv=True`) | — | — | +| `output` | `[G*P, H * C]`, contiguous; rows `[:G*P]` overwritten (per-head latent output — `v_b_proj` still to apply). A taller buffer's rows past `G*P` are left bitwise untouched | bf16 | CUDA | +| `latent_cache` | non-None required; **data not consumed and not mutated** (this step's rows must already be in the pool). Certified passing `[G*P, C + R]`, the production shape; the row count is not checked (`[G, …]` and `[1, …]` were also accepted, bitwise identically) | bf16 | CUDA | +| `q_pe` | non-None required; **data not consumed and not mutated**. Certified passing `[G*P, H, R]`, same unchecked-row-count caveat | bf16 | CUDA | +| `cu_q_seqlens` | `[G+1]` **required** (presence only): `arange(G+1) * H * P` (q-row prefix sums — this is what the sibling preprocessing op writes at every `P`) | int32 | CUDA | +| `cu_kv_seqlens` | `[G+1]`, `[0, cumsum(L_g)]` over generation sequences — **neither checked nor read on this path** | int32 | CUDA | +| `fmha_scheduler_counter` | `[1]` **required** (presence only), zeroed before the call | uint32 | CUDA | +| `quant_q_buffer` | `quant_mode=0`: `None`; `quant_mode=128`: `[G*P, H, C+R]` **required and read as the query** — `e4m3(fused_q * kv_scale_orig_quant)`, contiguous; `q` itself is then unread. Read through a raw pointer, so an equal-bytes `uint8` view behaves identically — and so a buffer shorter than `G*P` rows is read past its end rather than rejected | e4m3 | CUDA | +| `mla_bmm1_scale` | `quant_mode=0`: `None`; `quant_mode=128`: `[2]` **required**, `[x, x*log2(e)]` with `x = s^2 / (q_scaling * sqrt(nope+R))`. Element `[1]` is the one read. Per batch, **not** per `P` — the same pair serves every draft row | fp32 | CUDA | +| `mla_bmm2_scale` | `quant_mode=0`: `None`; `quant_mode=128`: `[1]` **required**, `[s]` — multiplies the output | fp32 | CUDA | +| `kv_scale_orig_quant`, `kv_scale_quant_orig` | **not read on this path at either `quant_mode`** (bitwise-identical output with both `None` at `s = 2.0`); pass them anyway if the same tensors serve the context call | fp32 | CUDA | + +`None` fails for two of the three — `cu_q_seqlens` (`seqQOffset is nullptr`, +`attentionOp.cpp:1045`) and `fmha_scheduler_counter` (`fmha_tile_counter is +nullptr`, `:1047`) — and both are presence-checked only, their contents +bitwise inert here. `cu_kv_seqlens` is accepted as `None`, and `None`, +all-zeros, non-monotonic and 1000x-inflated values all returned bitwise +identical output at `G = 4` and `G = 16` with pairwise-distinct `L_g`: the +paged decode kernel takes each sequence's KV length from `sequence_length` +(a one-token change there does move the output, from 0.016 to 0.499 max abs +error). Fill all three as the sibling generation-preprocessing op does — +that is the production shape and the certified one. + +**The scheduler buffers stay inert at `P > 1`** — the taller query block is +exactly the change that could have made them start mattering, so all of it +was re-measured at `P = 3`, `G = 2`, over the fp8 R1 cell. `cu_q_seqlens` +zeroed, left in the `P = 1` form `i*H`, given `i*P` without the head factor, +inflated 1000x, and made non-monotonic all returned **bitwise-identical** +output; so did `cu_kv_seqlens` zeroed, inflated and non-monotonic. Pass the +production values anyway: they are what the sibling op writes, and nothing +here promises the next version reads them the same way. + +### Shared data tensors + +| Argument | Shape / value | Dtype | Device | +|---|---|---|---| +| `output_sf` | `None` (NVFP4-output path not certified) | — | — | +| `workspace_` | persistent scratch tensor, any length (0 ok); the op grows it **in place** via `resize_()` when too small — 33 MB to 109 MB across the certified shapes. The MLA context requirement scales with tokens x heads and still sets the peak: 113 924 096 B for an `H = 128` context call over 545 tokens, against 52 579 072 B for the same call over 129 tokens and 100 024 320 B for the standard-configuration head-geometry sweep. The MLA **generation** call is well below that and, measured, does not grow with `predicted_tokens_per_seq`: 39 100 416 B at `H = 128`, page 32 for every `(G, P, L)` tried — `G` 2 and 4, `P` 1 and 4, `L` 50 and 500 — so a `P`-times-taller query block moves nothing here. Pass a plain resizable tensor, not a view; reuse it across calls to avoid re-allocation | int8 | CUDA | + +### Batch state (the "prepared metadata" of this op — all int32) + +| Argument | Shape | Content | Device | +|---|---|---|---| +| `sequence_length` | `[ns]` | per-sequence KV length **of this call**: total `cached + new` everywhere except an MLA chunked partial pass, where a context row holds that sequence's chunk length instead (0 = sits this loop out). For a generation row at `predicted_tokens_per_seq = P` the `new` part is all `P` of this step's tokens, and this tensor — not its host twin — is what the MLA generation call takes its per-sequence key range from (shifting it by one token moved every draft row's attended set by one, read out directly) | CUDA | +| `host_past_key_value_lengths` | `[ns]` | same values as `sequence_length`. On the MLA generation path the per-sequence values are otherwise inert, but the vector must not be **all** zero — see *Preconditions* | CPU | +| `host_total_kv_lens` | `[2]` | `[0]` = sum of `sequence_length` over context sequences (= the K/V row count for a no-append MLA context call), `[1]` = sum over generation sequences | CPU | +| `context_lengths` | `[ns]` | context sequences: this call's new-token (q-row) count — equals the prompt length only when nothing is cached; generation sequences: the original prompt length. On a context row it is load-bearing and unchecked (see *Preconditions*); on a generation row it is inert under a sliding window (0, 1 and the prompt length gave **bitwise identical** output) and inert on the MLA generation path at `predicted_tokens_per_seq` above 1 as well (0 and the full KV length both bitwise identical at `P = 3`) | CUDA | +| `host_context_lengths` | `[ns]` | same values as `context_lengths` | CPU | +| `host_request_types` | `[ns]` | 0 = context, 1 = generation; all 0s must precede all 1s | CPU | +| `num_contexts` (int) | — | number of context sequences. Binding default 0 is only correct for generation-only batches — always pass the true count | — | +| `num_ctx_tokens` (int) | — | sum of new-token counts over context sequences; same caveat | — | +| `max_context_q_len_override` | — | `None` (encoder CUDA-graph override; not certified) | — | + +CPU tensors may be pageable or pinned (pinning is a perf optimization). +For a mixed MLA batch, the two per-phase calls share these full-batch +(`ns`-sized) tensors and the same `num_contexts`/`num_ctx_tokens`; only +`q`/`output` and the phase-specific arguments differ between the calls. +The generation call indexes sequences starting at `num_contexts`. + +### Paged KV cache addressing + +The cache is a caller-owned pool addressed in **slab** units: one slab = +one layer's K (or V) block, `tokens_per_block * Hkv * D` elements +(standard) or `tokens_per_block * (C + R)` (MLA). A pool page holds the +slabs of every layer sharing the pool, layers side by side. Certified: +one pool, holding 1 layer (both configurations) or 4 layers (standard +configuration, bf16 and fp8-e4m3 pools both; multi-layer MLA pools not +certified). The certified multi-layer state — pool, mapping, block +offsets — was produced by a real 4-layer `KVCacheManager` (one per pool +dtype) and consumed unmodified, so the shapes below are exactly what a +multi-layer manager emits. + +Standard configuration (kv_factor 2, HND), `L` layers in the pool: pool +tensor `[num_blocks, L, 2, Hkv, tokens_per_block, D]` CUDA contiguous — +element type bf16 (`quant_mode=0`) or fp8-e4m3 (`quant_mode=128`) — +page `b` holds layer `l`'s K slab at +`[b, l, 0]` and its V slab at `[b, l, 1]` (single layer is the `L = 1` +case, `[num_blocks, 2, Hkv, tokens_per_block, D]`). The op sizes slabs +from `quant_mode`, never from the tensor (it sees only a raw pointer): a +pool whose dtype disagrees with `quant_mode` is silently corrupted. + +MLA (kv_factor 1, single layer): pool tensor +`[num_blocks, tokens_per_block, C + R]` CUDA contiguous — element type bf16 +(`quant_mode=0`) or fp8-e4m3 (`quant_mode=128`) — one latent row per token, +no separate K/V slabs. The slab width is `tokens_per_block * (C + R)` +**elements** regardless of the per-call `num_kv_heads`/`head_size` (context +passes `H`/192, generation `1`/576; both address the same pool), so the slab +is that many bytes under `quant_mode=128` and twice that under `quant_mode=0` +— the op derives the element width from `quant_mode` alone. Certified at +both element types with the requests' pages scattered out of order, every +page outside the reserved sets left bitwise zero. How many *bytes per token* +a KV-cache manager reserves for such a pool is the manager's business, not +this op's: the op is handed a base pointer plus slab offsets and addresses +exactly the geometry above. + +| Argument | Shape / value | Dtype | Device | +|---|---|---|---| +| `kv_cache_block_offsets` | `[1, max_num_requests, 2, max_blocks_per_seq]`; row `[0, s, 0, j]` = K-slab offset of sequence `s`'s `j`-th page, `[0, s, 1, j]` = V-slab offset. Offsets count slabs and are **layer-agnostic** — every layer of the pool shares one offsets row per sequence; layer selection comes from the pool mapping, not the offsets. Standard `L`-layer pool → page `p` has K `p * 2L`, V `p * 2L + 1` (single layer: K `2p`, V `2p + 1`); MLA pool → raw block id `p` in **both** rows. Only the first `ceil(kv_s / tokens_per_block)` entries per row are read | int32 | CUDA | +| `host_kv_cache_pool_pointers` | `[1, 2]`: `[0, 0]` = pool base address (`pool.data_ptr()`), `[0, 1]` = secondary-pool address (0 = none) | int64 | CPU | +| `host_kv_cache_pool_mapping` | `[num_layers, 2]`, one row per layer: row `local_layer_idx` = (pool index, layer index within pool). The row's **layer column drives the pool-base shift** — the call addresses slabs starting `layer_in_pool * kv_factor` slabs from the pool base (see *Notes*). Certified: the identity rows a real manager produces — `[[0, 0] .. [0, 3]]` for the 4-layer standard pool; `[1, 2]` zeros for single-layer pools | int32 | CPU | +| `local_layer_idx` (int) | row of `host_kv_cache_pool_mapping` for this layer (production: the layer's index among the rank's local layers); 0-3 certified (standard), 0 (MLA) | — | — | +| `tokens_per_block` (int) | pool page size; must be a power of two; 32 certified (standard), 32 and 64 certified (MLA — see *MLA page size*) | — | — | +| `update_kv_cache` (bool) | `True` (also for the MLA generation and no-append context calls, which nevertheless write nothing) | — | — | +| `cache_indirection` | `None` (beam search only) | — | — | +| `block_ids_per_seq` | `None` (FlashMLA path only) | — | — | + +A sequence's page set is shared by every layer of the pool: the per-layer +calls of one batch pass identical offsets and pool pointers and differ +only in `local_layer_idx`. Multiple pools (`num_pools > 1`: extra +offsets/pointer rows, mapping rows with pool index > 0) are not certified. + +### Paged KV cache addressing under a sliding window + +`attention_window_size` changes **nothing** about addressing: token `t` of a +sequence always lives at entry `t // tokens_per_block` of that sequence's +offsets row, slot `t % tokens_per_block`, and the append writes there. Every +memory saving is the caller's page mapping. What the op requires: + +- The offsets row is indexed by **absolute** page index, so it must have at + least `ceil(total_kv / tokens_per_block)` entries — that count keeps + growing with the sequence, window or no window, and bounds + `kv_cache_block_offsets.shape[3]`. +- Only entries holding at least one **in-window** token need to point at + valid live pages. An entry whose tokens are all older than + `total_kv - W` is never read: pointing it at a page filled with garbage + left a decode's output **bitwise identical** (checked entry by entry over + a 7-page sequence with `W = 128`; the first entry holding an in-window + token does change the output). Aged-out pages may therefore be recycled, + and stale ids may stay in the row — as long as they remain in-pool slab + offsets (out-of-pool ids were not tried and must not be). +- The same holds for a **context** call on the paged-context path, with the + in-window range measured from that call's rows: query row `i` of a + context sequence reads keys `[max(0, p - W + 1), p]` at + `p = kv_s - l_s + i`, so the call's read span is + `[max(0, kv_s - l_s + 1 - W), kv_s - 1]`. Measured entry by entry over an + 8-page sequence (`W = 128`, 200 cached tokens, 40 new): every entry + holding an in-window **cached** token changed the output (318x outside + the band with a decoy there), while entries below the window and the one + entry holding only this call's own new tokens left it **bitwise + identical** — that last one because the append lands wherever the offsets + point and the FMHA reads it straight back. The ring bound that follows + from that span (not separately certified) is + `P * tokens_per_block >= W + l_s - 1` for a cached-prefix context call, + on top of the decode bound below. +- A **bounded ring** is the resulting production shape: give the sequence + `P` physical pages and map absolute page index `j` to `ring[j % P]`. The + op's absolute-position append then lands token `t` at physical slot + `t mod (P * tokens_per_block)`, so the pool holds a moving window of the + sequence and genuinely wraps, overwriting exactly the slot of the token + that has just aged out. Requires + `P * tokens_per_block >= attention_window_size`; certified at + `P * tokens_per_block` of 128 (= `W`) and 192 over 250 wrapping decode + steps. One page short (96 < 128) is **silently wrong** — no exception, + output far off the reference. +- The same call's new tokens must map to **distinct** physical slots: + `P * tokens_per_block >= that call's new-token count`. A prefill longer + than the ring makes several tokens of one append target one slot; those + writes race (observed: two identical runs left different pool bytes, + torn rows). So a sliding-window sequence still needs + `ceil(prompt_len / tokens_per_block)` distinct pages for its prefill call + — the window shrinks the steady-state decode footprint, not the prefill's. +- The append remains surgical: one decode step changes exactly one + `(K, V)` row pair in the whole pool, the evicted slot's, and touches + nothing else. No compaction, zeroing, or relocation ever happens. + +### Geometry and masking scalars + +| Argument | Standard (certified) | MLA context | MLA generation | +|---|---|---|---| +| `is_fused_qkv` | `True` (packed QKV in `q`) | `False` | `True` | +| `attention_input_type` | 0 = mixed (either subset may be empty; 1/2 not certified here) | 1 = context_only | 2 = generation_only | +| `num_heads`, `num_kv_heads`, `head_size` | must match packed `q` and pool. Head counts are a free axis subject to `num_heads % num_kv_heads == 0`; `head_size` is enumerated by the available FMHA kernels (see *Head geometry*). Certified on the bf16 pool: 32/8/128, 32/4/128, 64/8/64, 16/1/128, 28/4/128, 12/3/128, 8/2/128, 8/2/256, 4/4/64; on the fp8-e4m3 pool: 32/8/128 | `H`/`H`/192 (`H`/`H`/`nope+R`); 8/8/192, 16/16/192, 32/32/192 and 128/128/192 certified on the bf16 latent pool, 128/128/192 only on the fp8-e4m3 one | `H`/1/576 (`H`/1/`C+R`); 8/1/576, 16/1/576, 32/1/576 and 128/1/576 certified on the bf16 latent pool, 128/1/576 only on the fp8-e4m3 one | +| `mask_type` | 1 = causal, 0 = padding — both certified on the bf16 pool (padding on context-only batches); fp8 pool certified with 1 only | 1 = causal (bottom-right aligned when KV > q); 0 = padding certified for the no-append chunked partial passes, on the bf16 latent pool only | **Inert — this path takes its mask from `predicted_tokens_per_seq` alone.** Pass 1. At `P = 1` a decode row attends to all `L_g` cached tokens; at `P > 1` draft row `t` attends to `[0, L_g - P + t]`, bottom-right-aligned causal within the block. `mask_type = 0` returns **bitwise-identical** output at `P` = 2 and 4 (fp8, `H = 128`, page 32) — the plausible guess that padding drops the within-block mask and lets a draft row see its later siblings is wrong | +| `q_lora_rank` | `None` | the checkpoint's `q_a_proj` rank — `1536` and `0` (no q-LoRA) certified and **inert**: no MLA path reads it (see *MLA q-LoRA rank*). `None` raises here, unlike in the standard configuration | same value as context | +| `kv_lora_rank`, `qk_nope_head_dim`, `qk_rope_head_dim` | all `None` | 512 / 128 / 64 certified | same values as context | +| `v_head_dim` | `None` | 128 (context V/output head width) | 512 (= `C`, latent output width) | +| `rope_append` | `None` | `True` | `True` (`False` widens the output to `C+R` per source — not certified) | +| `predicted_tokens_per_seq` | 1 | 1 | **1, 2, 3 and 4 certified** (`P`): the number of query rows each generation sequence contributes. 1 is an ordinary decode step; `P > 1` is one speculative-decoding step verifying a `P`-token draft chain, and `P` is then the *only* thing that produces the within-block mask (see `mask_type`). One scalar for the whole batch — a ragged per-sequence draft length cannot be expressed here. Certified over the fp8-e4m3 latent pool at the complete DeepSeek-R1-0528 cell and over a bf16 latent pool, both at `H = 128`, page 32; above 4, and at other head counts or page sizes, untested. `P` also reaches decode-kernel selection as the kernel's `maxSeqLenQ` — see *Notes* | +| `q_scaling` | softmax scale `1 / (q_scaling * sqrt(head_size))`; 1.0 certified | free positive scalar; scale uses `sqrt(nope+R)`. Certified 1.0, 0.53366 (`= 1/mscale²`, DeepSeek-R1's YaRN temperature — see *In-kernel RoPE arguments*), 2.0 and 0.25 on the bf16 latent pool at `H = 128`; 1.0 and 0.53366 on the fp8-e4m3 one, in both context flavors | same values; the scale still uses `sqrt(nope+R)`, **not** `sqrt(head_size)`. Over an fp8 pool it reaches the kernel only through the caller's `mla_bmm1_scale` (see *MLA over an fp8-e4m3 latent pool*) | +| `quant_mode` | 0 = bf16 pool; 128 = fp8-e4m3 pool (see *Standard configuration over an fp8-e4m3 KV pool*; other bits not certified) | 0, or 128 for an fp8-e4m3 latent pool. `1152` (`\| FP8_1x128_128x128`) and `384` (`\| FP8_QDQ`) are bit-identical to `128` here — only the KV-cache bit is read (see *MLA over an fp8-e4m3 latent pool*) | same value as context | +| `beam_width` | 1 | 1 | 1 | +| `max_num_requests` | `>= ns`; equals `kv_cache_block_offsets.shape[1]` | same | same | +| `max_context_length` | host launch/workspace bound `>=` max per-sequence q-row count of the call (certified `= max_seq_len`) | same | same | +| `max_seq_len` | `>=` max `sequence_length`; per-sequence page capacity bound. Does not select what a query attends to: with a window active, decodes at `max_seq_len` of `W`, `total + 1` and `4 * (total + 1)` returned **bitwise identical** output (it does steer kernel selection — see *Notes*) | same | same | +| `attention_window_size` | sliding-window token count `W`; `>= max_seq_len` = no window. Smaller = each query attends only the newest `W` keys (see *Sliding window*); 128 / 100 / 33 certified on the bf16 pool with `mask_type=1`. Per-call scalar: layers sharing a pool may differ in it | same (no window certified) | same (no window certified) | +| `use_paged_context_fmha` | context execution path; both values certified. `False` = context FMHA over the packed `q` rows, so every context sequence must start empty (`sequence_length == host_context_lengths`). `True` = context FMHA over the paged pool, which is what a context call over a cached prefix requires (`sequence_length > host_context_lengths`) — see *Paged-context FMHA*. With nothing cached the two are bitwise identical | `False` (MLA context always takes separate q/k/v) | `False` | +| `chunked_prefill_buffer_batch_size` | 1 | 1 | 1 | +| `attention_chunk_size` | `None` | `None` | `None` | +| `is_cross` / `cross_kv` | `False` / `None` | same | same | +| `trtllm_gen_jit_warmup` | `False` | `False` | `False` | + +### In-kernel RoPE arguments + +Standard configuration: RoPE disabled — `position_embedding_type=0` +(learned-absolute; no rotation applied), `rotary_inv_freq=None`, +`rotary_cos_sin=None`, `rope_dim=0`, `rope_base=10000.0`, +`rope_scale_type=0`, `rope_scale=1.0`, `rope_short_m_scale=1.0`, +`rope_long_m_scale=1.0`, `rope_max_positions=1024`, +`rope_original_max_positions=1024`, `mrope_rotary_cos_sin=None`, +`mrope_position_deltas=None`. Non-zero `rope_dim` with a rope-type +`position_embedding_type` applies RoPE to q/k inside the call — not +certified for this configuration. + +MLA (all calls receive the same values; only the fresh-prefill context +call rotates anything): `position_embedding_type=8` (yarn, as MLA models +set), `rope_dim=R`, and `rotary_cos_sin` = fp32 +`[1, max_positions * R * 2]` duplicated-layout table — per position, `R` +`(cos, sin)` pairs at flat offsets `p*2R + 2d` and `p*2R + 2d + 1`, whose +second `R/2` pairs duplicate the first; only pairs `[0, R/2)` are read per +position. + +**The table's content is the only rope input the op reads.** Two contents +are certified, both at `H = 128`, page 32, on the fresh-prefill context call +(the only MLA flavor that rotates anything): + +- the **unscaled** table, `angle(p, d) = p / theta^(2d/R)` at + `theta = 10000` — what a config without rope scaling produces, and what + every MLA case at `H` = 8, 16 and 32 above uses; +- the **YaRN-scaled** table a DeepSeek-R1-0528 config produces: `theta` + 10000, `factor` 40, `original_max_position_embeddings` 4096, + `beta_fast` 32, `beta_slow` 1, `mscale` 1.0, `mscale_all_dim` 1.0. + +The YaRN content, written out so a caller can build it without reaching into +TensorRT-LLM (fp32 throughout; `d` indexes the half-dimension `[0, R/2)` and +`j` the table's `R` pairs per position, `[0, R)`): + +``` +freq(d) = theta ** (2d / R) +low = max(0, floor(R * ln(orig_max_pos / (beta_fast * 2*pi)) + / (2 * ln(theta)))) # 10 for R1 +high = min(R - 1, ceil (R * ln(orig_max_pos / (beta_slow * 2*pi)) + / (2 * ln(theta)))) # 23 for R1 +ramp(d) = clamp((d - low) / max(high - low, 0.001), 0, 1) +inv_freq(d) = ramp(d) / (factor * freq(d)) + (1 - ramp(d)) / freq(d) +m(x) = 1.0 if factor <= 1 else 0.1 * x * ln(factor) + 1.0 +amplitude = m(mscale) / m(mscale_all_dim) # 1.0 for R1 +angle(p, j) = p * inv_freq(j mod (R/2)) # the duplication lives here +table[p*2R + 2j] = cos(angle(p, j)) * amplitude +table[p*2R + 2j + 1] = sin(angle(p, j)) * amplitude +``` + +Setting `factor = 1` collapses this to the unscaled table (verified: the two +constructions agree to 3.8e-6 max abs, fp32 evaluation-order noise). Two +consequences +worth naming: YaRN rescales only the **low-frequency** half of the spectrum +(`ramp` is 0 below `low`), so the two tables agree closely at small positions +and diverge with distance; and the table's `amplitude` is **1.0 whenever the +config's two mscales are equal**, which is R1's case — the model's YaRN +attention temperature `mscale = 0.1*ln(40) + 1 = 1.36885` is folded into the +softmax scale through `q_scaling = 1/mscale² ≈ 0.53366` instead, never into +the table. + +**Everything else about the rope configuration is inert on the MLA path** — +measured, not inferred. With the table held fixed, the seven scalars +`rope_base`, `rope_scale_type`, `rope_scale`, `rope_short_m_scale`, +`rope_long_m_scale`, `rope_max_positions`, `rope_original_max_positions` and +the `rotary_inv_freq` tensor moved **no** observable of a fresh-prefill +context call: output rows, the whole paged latent pool, the in-place-roped +`q`/`k`, all bitwise identical across R1's production scalar set, the +unscaled set above, a deliberately out-of-range set (`rope_base` 500000, +`rope_scale_type` 3, `rope_scale` 7.5, m-scales 3.0/9.0, position windows 77 +and 13 — both shorter than the sequences in flight), and `rotary_inv_freq` +zeroed or passed as `None`. Swapping the table in the same comparison does +move all of them, which is what makes the bitwise result mean something. +So: pass whatever scalars your config carries, and get the table right. + +Two obligations the table brings, neither of them checked: + +- It must hold **one row per absolute position the call ropes**. Nothing + bounds-checks it and `rope_max_positions` is not consulted: a 64-row table + under a 96-token prefill returned without raising and was wrong from output + row 64 on (max abs diff 3.9 against magnitude-4 outputs) — the kernel reads + past the end of the tensor. Certified with 1024 rows over sequences + reaching position 511. +- Row content does not depend on the row count, so a table truncated to + `max_seq_len` rows is **bit-identical** to the leading rows of the full + 163840-row table an engine builds for R1 (verified element for element). + Truncating is therefore safe; short-changing the positions in flight is + not. + +`rotary_inv_freq` = the `[R/2]` fp32 tensor from the same construction. It +is accepted, unread on this path (above), and still worth passing as +production does. + +### Feature groups not certified — pass these inert values + +| Group | Arguments → inert value | +|---|---| +| Quantized output / quantized QKV input | `out_scale=None`, `quant_scale_qkv=None` (fp8/fp4-quantized attention *output* and the fused DSv4 QKV-quantization path are not certified — on the MLA fp8 path `quant_scale_qkv` was additionally observed to be inert, a `[1]` fp32 `1/s` there leaving the context output bitwise identical. The fp8-e4m3 **KV pool** is certified in both configurations — see those sections, whose two `kv_scale_*` tensors then carry the scale, `None` being read as `1.0` rather than rejected. On the bf16 pool both stay `None`, and so does `quant_q_buffer`.) | +| Folded KV RMSNorm (MLA) | `kv_norm_weight=None`, `kv_norm_eps=1e-6`. Non-`None` weight folds the `kv_a_layernorm` into the KV kernel, which then reads `latent_cache` **raw** — a caller that already normalized would normalize twice. Only DeepSeek-V4's sparse module uses it upstream | +| Skip-correction (MLA) | `skip_correction_threshold=0.0`. A lossy trtllm-gen MLA option, SM100/SM103 only, off by default upstream (`enable_mla_skip_correction`). The engine forces it to 0.0 on a non-MLA layer regardless | +| Spec-dec tree mask | `force_prepare_spec_dec_tree_mask` — engine-set, `True` only for a dynamic tree; a linear-tree draft leaves it `False` | +| Sequence-count sizing | `max_num_sequences` — engine-set, defaults to `max_num_requests` | +| MLA DSv4 / FlashMLA / sparse-MLA sub-features | `flash_mla_tile_scheduler_metadata=None`, `flash_mla_num_splits=None`, `dsv4_inv_rope_cos_sin_cache=None`, `enable_dsv4_epilogue_fusion=False`, `aux_kv_cache_pool_ptr=None`, `sparse_attn_kv_lens=None` (`mla_bmm1_scale` / `mla_bmm2_scale` / `quant_q_buffer` are **certified** for the MLA generation call over an fp8-e4m3 latent pool, where all three are required — see that section; `None` everywhere else) | +| Speculative decoding — the **mask/offset tensor** group | `is_spec_decoding_enabled=False`, `use_spec_decoding=False`, `is_spec_dec_tree=False`, all `spec_decoding_*=None`, `spec_bl_tree_first_sparse_mask_offset_kv=None`, `spec_decoding_target_max_draft_tokens=None`. This does **not** mean speculative decoding is unavailable: `predicted_tokens_per_seq > 1` is certified for the MLA generation call with exactly these inert values, and that is also what the engine passes on sm_100 for a linear-tree draft, where `is_spec_decoding_enabled` is forced off (the Python backend computes it as `is_spec_decoding_enabled and (not trtllm_gen_arch or is_spec_dec_dynamic_tree)`), leaving the mask, position-offset and generation-length tensors all `None`. The within-block mask then comes from `predicted_tokens_per_seq` alone — see its row above. A **tree** draft, which is what these tensors describe, is not certified | +| Sparse / skip-softmax | `sparse_kv_indices=None`, `sparse_kv_offsets=None`, `sparse_attn_indices=None`, `sparse_attn_offsets=None`, `sparse_attn_indices_block_size=0`, `num_sparse_topk=None`, `skip_softmax_threshold_scale_factor_prefill=None`, `skip_softmax_threshold_scale_factor_decode=None`, `skip_softmax_stat=None` (`softmax_stats_tensor` is certified for the no-append MLA context call — see that table; `None` everywhere else) | +| SageAttention | `sage_attn_num_elts_per_blk_q=0`, `sage_attn_num_elts_per_blk_k=0`, `sage_attn_num_elts_per_blk_v=0`, `sage_attn_qk_int8=False` | +| Helix CP | `helix_position_offsets=None`, `helix_is_inactive_rank=None` | +| Relative bias (T5) | `relative_attention_bias=None`, `relative_attention_max_distance=0` | + +`attention_sinks` is **certified** for the standard bf16-pool causal +configuration — see *Attention sinks*; pass `None` everywhere else +(MLA calls, the fp8-e4m3 pool, `mask_type=0`). + +## Metadata consumed + +None — the op reads no thread-local model attrs, no registered layers, and +no global Python state. Every input above is an explicit argument. (The +batch-state tensor block plays the role that prepared attention metadata +plays for the registered-layer entry points.) + +## Preconditions + +- Batch order: context sequences before generation sequences; `q` rows in + that order. `num_contexts` / `num_ctx_tokens` passed explicitly and + consistent with `host_request_types`. +- `num_heads % num_kv_heads == 0` (the wrapper asserts it: a context-only + call at a non-multiple geometry returns a silently wrong `output`). +- `head_size` is one of the kernel's head dims. On sm_100 an unsupported + `head_size` **aborts the process**: the kernel-selection failure is + thrown where it cannot unwind (`std::terminate`), so it is neither + catchable nor loggable by the caller — observed at 32 (`Invalid TilePvM + as MMA only supports 64 or 128`), 96 and 192 (`Unsupported HeadDim for + BMM2-N`). Same for a `num_heads` that kernel selection cannot group into + CTAs (observed at `Hq = 7`, `Hkv = 2`: `Internal error numHeadsQ=7, + numHeadsPerCta=3, numCtasForAllHeads=2`). Verify a new `head_size` in a + throwaway process before wiring it into a target. +- Standard-configuration context sequences must start empty + (`sequence_length == host_context_lengths ==` new-token count) **unless** + `use_paged_context_fmha=True`; a cached prefix at `False` is silently + wrong (see *Paged-context FMHA*). The MLA fresh-prefill flavor requires + empty starts unconditionally; MLA context over cached tokens is certified + only via the no-append flavor (`latent_cache=None`, explicit full-range + or chunked K/V), and `use_paged_context_fmha` stays `False` for MLA. +- On the standard paged-context path (`use_paged_context_fmha=True`), a + context row's `context_lengths` / `host_context_lengths` is **this + call's** new-token (q-row) count, not the sequence's prompt or KV length. + Nothing checks it: passing the full KV length landed 101x outside the + tolerance band and `0` landed 53x outside, both without raising, and both + also corrupt the pool. `sequence_length` stays the global cached+new + count, as everywhere else. +- The paged-context path additionally requires every page holding a key the + context call reads to be valid at call time, and those pages to be + distinct from one another — the whole `[0, sequence_length)` range + without a window, the narrower range given in *Paged KV cache addressing + under a sliding window* with one. The packed path (`False`) reads no page + at all during a context call, so this obligation is new and nothing + checks it: aliasing a sequence's pages onto a single page returns a + result 31x outside the tolerance band without raising. +- No-append MLA context calls additionally require: `k`/`v` rows are + exactly the per-sequence KV ranges described by `sequence_length`, + concatenated in batch order, with `host_total_kv_lens[0]` equal to that + row count; `v`'s row stride is `H * (nope + v_head_dim)` elements; under + `mask_type=1` every context sequence satisfies `sequence_length >= + context_lengths` (the causal mask offset is their difference); under + `mask_type=0` any KV length including 0 is allowed, but the output and + stats rows of zero-KV sequences are undefined and must never be read + (production merge plans skip them). `q` arrives fully prepared — the + call applies no RoPE even though the rope-table arguments are passed. + Under `quant_mode=128` all of the above is unchanged, and the pool stays + untouched (it is not read to build `k`/`v` — assembling those from the + cached latent rows is entirely the caller's step); what changes is that the + op quantizes the `k`/`v` it is handed to e4m3 before the FMHA, so operands + the caller has already dequantized off an fp8 pool go through e4m3 a second + time. +- A caller-provided `softmax_stats_tensor` (no-append MLA context only) + is fp32 `[>= Tc, >= H, 2]`, contiguous, on the same device (the op + rejects other dtypes/sizes); certified at the exact size `[Tc, H, 2]`. +- For every generation sequence, the cached tokens at positions + `[0, kv_s - 1)` were actually written by prior calls addressing the same + pool pages, and `sequence_length` equals the true cached count + 1. +- Per-sequence pages cover `ceil(total_kv / tokens_per_block)` entries in + its `kv_cache_block_offsets` row, where `total_kv` is the sequence's + full `[cached + new]` KV length (also for chunked partial passes, whose + `sequence_length` entries are smaller); distinct sequences use distinct + pages; all offsets lie inside the pool. `total_kv <= max_seq_len` for + every sequence. +- With a sliding window the entry count above is unchanged (the row is + indexed by absolute page index), but only the entries holding in-window + tokens must be valid, and the pages behind them must be distinct over + each call's new-token range plus the whole window — see *Paged KV cache + addressing under a sliding window* for the ring rule + (`P * tokens_per_block >= max(attention_window_size, this call's + new tokens)`), which is not checked anywhere and fails silently. +- `attention_sinks`, when given, is fp32 (the op raises `Expected + attention_sinks to have float dtype` for anything else) and holds exactly + `num_heads` values in contiguous memory on the current CUDA device. Only + the dtype is checked by the op: the buffer is read as `num_heads` raw + fp32 values from `data_ptr()`, so every other violation is **silent** — + a shorter tensor is read past its end, a longer one has its tail ignored, + a strided view is read as its underlying memory (a stride-2 view holding + the right values returned results ~2.1 off), and an empty tensor + disables the sink entirely. The wrapper therefore asserts contiguity and + `numel() == num_heads`; nothing else guards this argument. A pageable + **CPU** fp32 tensor was also accepted and gave the correct result (the + values do reach the kernel), but CUDA is the certified device. +- All tensors contiguous unless noted (MLA context `v` is the strided + split view described in *Signature*), on the devices listed above; every + CUDA tensor on the current device; dtypes exactly as listed (host length + tensors are int32 — int64 is not certified). +- `quant_mode` and the pool element type must agree (0 ↔ bf16, + 128 ↔ fp8-e4m3): the op sizes and interprets slabs from `quant_mode` + alone and cannot detect a mismatched pool allocation. +- With `quant_mode=128`, `kv_scale_orig_quant` and `kv_scale_quant_orig` + are fp32 `[1]` CUDA tensors holding `1/s` and `s` for one scaling factor + `s` (mutual consistency is the caller's job). **`None` in either slot or + in both is not rejected — it is silently read as `1.0`**, bitwise + identical to passing 1.0 tensors in both the pool bytes and the output, + in both configurations, in every phase, and on both the single-layer and + the 4-layer manager pool. + A forgotten scale is therefore a wrong-number bug, not a crash: inert at + the production `s = 1.0`, and 1.19 off on magnitude-1.5 decode outputs + over a standard-configuration pool written at `s = 2.0`. + Certified `s`: 1.0 (the bf16-checkpoint production value), 1.5, 2.0. +- On the **MLA** path under `quant_mode=128` the same two tensors reach far + less of the computation, and `s != 1.0` is not survivable in the context + phase (full measurements in *MLA over an fp8-e4m3 latent pool*): + - the fresh-prefill append consumes `kv_scale_orig_quant` and is bit-exact + at any `s`; + - **both** context flavors quantize `q`/`k`/`v` at 1.0 but apply `s^2` and + `s` as if they had not, so a context call at `s != 1.0` returns a silently + wrong result — measured on the fresh-prefill flavor (21x / 48x outside the + band at 1.5 / 2.0) and on the no-append flavor (48x at 2.0), which appends + nothing and so has no use for `kv_scale_orig_quant` at all. Pass + `s = 1.0`, or `None`; + - the generation call reads **neither** tensor. Its query comes from + `quant_q_buffer` and its scales from `mla_bmm1_scale`/`mla_bmm2_scale`, + all three required (`None` raises `quant_q_buf is nullptr.` / + `bmm1_scale is nullptr.` / `bmm2_scale is nullptr.`, from + `attentionOp.cpp:1050-1052`), and the caller must fold `s` into them as + shown in that section. Nothing checks their contents. +- Bits of `quant_mode` outside the KV-cache group are **not** a reason to + strip the value a checkpoint's quant config produces: `1152` and `384` were + measured bit-identical to a bare `128` on every MLA flavor (see *MLA over an + fp8-e4m3 latent pool*). Bits that were not tried are unknown, not inert. +- MLA batches take one call per phase; both calls of a mixed batch use the + same full-batch state tensors. `attention_input_type=0` is **accepted, not + rejected**: inert (bitwise identical) on a single-phase batch, silently + unsatisfiable on a mixed one — see *Semantics*. +- Every MLA call needs an int `q_lora_rank` — `0` for a checkpoint without + a q-LoRA — even though the value is never read: `None` raises + `RuntimeError: bad optional access` (see *MLA q-LoRA rank*). +- MLA generation additionally requires, before the call: every generation + sequence's full latent history `[0, L_g)` resident in the pool — the + context call's append covers the prefill rows; this step's `P` rows at + positions `L_g - P .. L_g - 1` come from the caller's generation + preprocessing — and the caller-allocated `cu_q_seqlens` / `cu_kv_seqlens` / + `fmha_scheduler_counter` filled as listed in *Signature*. Over an + fp8-e4m3 pool the same preprocessing must also have filled + `quant_q_buffer` with the e4m3 query and both `mla_bmm*_scale` buffers + with the scales above; a decode row appended by the caller must be + quantized the same way the op's own append quantizes + (`e4m3(row * kv_scale_orig_quant)`), since the kernel applies no + per-row scale of its own. +- With `predicted_tokens_per_seq = P` above 1, MLA generation additionally + requires `L_g >= P` for every generation sequence — draft row 0 attends to + `[0, L_g - P]`, so a shorter `L_g` leaves it no keys. `L_g == P` is legal + and certified (row 0 then attends to exactly one key, its own). `q` / + `output` / `quant_q_buffer` must be `G*P` rows tall and token-major within + a sequence; `sequence_length[num_contexts + g]` must count **all `P`** of + this step's tokens. Nothing checks the row count — `q` and `quant_q_buffer` + are addressed from `P` and the sequence index, so a `G`-row buffer is read + past its end rather than rejected. +- `host_past_key_value_lengths` must be non-zero for at least one sequence. + Its per-sequence values are otherwise inert on the MLA generation path + (`sequence_length - 1`, `- P` and all-ones are each bitwise identical to + the true values at `P = 3`), but an **all-zero** vector makes the call + return without writing a single element of `output` — a sentinel-filled + buffer comes back untouched, at `P` = 1 and above alike, with no error. + Passing the same values as `sequence_length`, as everywhere else in this + contract, satisfies it. +- The MLA fresh-prefill context call clobbers the rope slice of each + `q`/`k` head in place: treat both tensors as consumed (their nope slices + stay intact, but do not rely on pre-call rope-slice contents + afterwards). The no-append context call mutates neither. +- `rotary_cos_sin` must hold a row for every absolute position the + fresh-prefill context call ropes, in the duplicated layout and with the + content its rope configuration implies (see *In-kernel RoPE arguments*). + Nothing checks either: a short table is read past its end and comes back + silently wrong from the first over-long position on, and the scalar rope + arguments beside it — including `rope_max_positions` — are not consulted at + all, so they cannot substitute for getting the table right. +- First generation-phase call per process JIT-compiles its FMHA kernel via + NVRTC (~5-10 s; one kernel per configuration — standard GQA and MLA + decode compile separately): the C++ locates the shipped kernel headers + (`tensorrt_llm/include/trtllm_gen_kernels/fmha/`) by running + `pip show tensorrt_llm` in a subprocess, so a `pip` that resolves the + installed package must be on PATH (this repo pins `pip` as a project + dependency for exactly this reason); otherwise the call raises + `NVRTC_ERROR_COMPILATION` ("could not open source file cuda.h"). + Context-phase kernels are precompiled into `libtensorrt_llm.so` and need + no JIT. +- The op is synchronous-schema but asynchronous: it launches on the current + CUDA stream and returns; synchronize before reading `output` on the host. + +## Notes + +- One Python-level invocation launches several device kernels (QKV/KV-cache + preprocessing, FMHA, postprocessing) inside the single C++ op — the same + attention core the TRTLLM backend entry points route into, certified here + as a single catalog call. +- On sm_100 the op selects trtllm-gen FMHA kernels. Generation kernels are + JIT-compiled per process (in-memory kernel cache; the log warns "Possible + JIT Cache Missing ... generateAndCompileKernel took ~5-10 s" on first + use of each kernel variant — the MLA decode variant is + `...HQk576HV512...PagedKv...`; the fp8-KV GQA decode variant is + `fmhaSm100aKernel_QkvE4m3OBfloat16H128PagedKvDense...` — its name pins + that q, K, and V are all e4m3 in the decode MMAs). Re-runs within the + process are fast. Context kernels — including the padding-mask and + softmax-stats variants of the no-append flavor, and the bf16 context + kernel that keeps serving the context phase under `quant_mode=128` — are + precompiled (no JIT observed). +- The generation-kernel variant is keyed by `head_size` and the q-tile size + the head group maps to, rather than by the head counts directly — several + distinct geometries share one name: the bf16 + variants observed are `...H128PagedKvDenseP32VarSeqQ8Kv128Static...` for + `Hq/Hkv` of 32/8, 32/4 and 28/4, `...Q16...` for 16/1, and `...H256...` / + `...H80...` for those head sizes. So a first call at a new head geometry + costs at most one extra ~6 s JIT compile, and often none. The mapping is + not constant in the head count either, though: MLA generation at + `H` = 16, 32 **and 128** all report + `fmhaSm100aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvDenseP64VarSeqQ16Kv128StaticSwapsAbForGen` + (the `P32` sibling at page 32), while `H = 8` reports + `...VarSeqQ8Kv128StaticSwapsAbForGen` — same kernel + family, different q-tile, hence a different compiled kernel. The q-tile + therefore stops growing well below the head count: 16x the heads of the + `H = 8` case still map to `Q16`. Where head + counts do share a name the **compile cache is still keyed more finely**: a + process that runs 16, 32 and 128 pays the ~5.5 s `...Q16...` compile three + times, once per head count, while any number of calls at one head count — + across batch sizes and histories — compile once. The **page size is part of + the name** too: the same MLA generation call at `tokens_per_block = 32` + compiles + `fmhaSm100aKernel_QkvBfloat16OBfloat16HQk576HV512HVPerCta128PagedKvDenseP32VarSeqQ16Kv128StaticSwapsAbForGen` + (~5.4-5.9 s, observed once per head count that maps to it), a genuinely + different kernel from the `...P64...` one — so page 32 and page 64 MLA + decode are separate compiled variants, both certified. Full-test-file counts on sm_100 measured + from the run behind the current receipt, all at + `predicted_tokens_per_seq = 1`: `...P64VarSeqQ16...` twice (the + page-64 16/32 pair), `...P32VarSeqQ16...` **three times** (16/32/128), + `...P64VarSeqQ8...` once and + `...P32VarSeqQ8...` once (`H = 8`, each serving both that page size's decode + and mixed-batch cases, the page-32 one additionally serving the `H = 8` + q-LoRA-rank sweep's decode calls) — all seven ~5.4-5.7 s, a wall-clock + figure that moves ~0.15 s run to run. One `H = 128` compile serves all ten + of the file's `H = 128` decode calls at `P = 1` — 3 from the baseline decode + and mixed-batch cases, 3 from their R1-cell twins, 4 from the `q_scaling` + sweep. The MTP section adds two more bf16 compiles on top, one per `P` it + runs there (`...HVPerCta128...P32VarSeqQ16...` at `P = 2` and + `...HVPerCta256...P32VarSeqQ16...` at `P = 4`), for nine bf16 MLA decode + compiles in the file. The counts are the stable part: identical across + full-file runs. + Every JIT compile logged across the file is a `...ForGen` decode variant — + no MLA context call at any certified head count, page size or pool dtype + triggered one. `quant_mode` is part of the decode kernel name too: MLA + generation over an fp8-e4m3 latent pool compiles + `fmhaSm100aKernel_QkvE4m3OBfloat16HQk576HV512HVPerCta128PagedKvDenseP32VarSeqQ16Kv128StaticSwapsAbForGen` + — the `QkvE4m3` sibling of the bf16 `...P32VarSeqQ16...` name, so q, K and V + are all e4m3 in the decode MMAs — once for every fp8 MLA decode call at + `predicted_tokens_per_seq = 1`: the baseline-cell decode, the R1-cell + decode, the mixed-batch decode and the `quant_mode`-bits decodes all share + it. So on the fp8 path neither `q_scaling`, nor the rope table, nor the extra + `quant_mode` bits reach decode kernel selection, and no fp8 MLA **context** + call — either flavor — triggers a JIT compile at all. +- **`predicted_tokens_per_seq` does reach decode kernel selection**: it becomes + the kernel's `maxSeqLenQ`, which appears in the JIT log line and moves the + `HVPerCta` split in the name. Observed at `H = 128`, page 32, both pool + dtypes: `P` = 1 and 2 take `...HV512HVPerCta128PagedKvDenseP32VarSeqQ16...` + and `P` = 3 and 4 take the `...HVPerCta256...` sibling, with `P` = 1 and + `P` = 2 paying **separate** ~5.4-5.7 s compiles despite sharing a name — the + same finer-than-the-name cache keying the head count shows. The file's fp8 + MLA decode compiles are therefore three, at `maxSeqLenQ` 1, 2 and 3, with + `P = 4` reusing the `maxSeqLenQ = 3` entry. Batch size can move it as well: + a probe at `G = 4`, `P = 4` compiled a further variant whose name carries no + `HVPerCta` segment at all. So a target sweeping `max_draft_len` should + expect one decode JIT compile per `P` it runs, not one for the model. +- First call logs "Attention workspace size is not enough" and resizes + `workspace_` in place — expected when starting from an empty tensor. The + size tracks the call's tokens x heads: ~33 MB standard, ~39 MB MLA and + ~35 MB no-append MLA context at the 8/16/32-head shapes, but ~50 MB for an + `H = 128` MLA context call over 129 tokens and ~109 MB over 545, which is + the file's peak (the standard-configuration head-geometry sweep reaches + ~95 MB). `quant_mode=128` raises the MLA context requirement — the fp8 + path stages quantized q/k/v in the workspace: 61 033 728 B for the same + 129-token `H = 128` call that needs 52 579 072 B in bf16, and 142 612 480 B + over 512 tokens (measured outside the shipped cases, which stay below the + file's bf16 peak). The **no-append** flavor pays the same surcharge — it + quantizes the K/V it is handed, so it stages them too: the shipped R1-cell + case (66 q rows over 193 K/V rows at `H = 128`) asks for 52 816 640 B under + `quant_mode=128` against 43 288 832 B at the same shape in bf16. An engine + at production `max_num_tokens` will be well above all of + these — budget from the shape, not from these figures. +- The paged-cache append is bit-exact wherever it is a copy: after a + standard call, page `p` of a + sequence holds exactly the bf16 K/V slices packed into `q`; after an MLA + fresh-prefill context call it holds `[ckv | rope(k_pe)]` with `ckv` + bitwise-copied (verified for both configurations, including page-boundary + crossings, and for MLA at page sizes 32 and 64). The MLA RoPE half is + bit-identical to an fp32 reference rounded once to bf16 **except at + fp32→bf16 rounding ties**: kernel and torch reference evaluate + `x*cos ∓ y*sin` in different fp32 orders, and where the correctly-rounded + fp32 result lands exactly on a bf16 midpoint the two round to adjacent + bf16 values. Measured on sm_100 over 60 prefill runs spanning both page + sizes (495 360 roped elements): 8 such elements — 1.6e-5 of the total, + never more than one bf16 ulp, and the same count at page 32 as at page 64, + so it is arithmetic rather than addressing. The entry's test gates the + RoPE half at one bf16 ulp (`rtol = 2^-7`, `atol = 0`) plus a hard "at most + one element in a thousand differs at all" floor, and keeps the `ckv` half + strictly bitwise; that gate leaves the real failure modes far outside it — + measured on one 96-token page-32 request, an append shifted one slot lands + at 6.1e5 ulp with 99.9% of elements differing, the pool read with page + size 64 instead of 32 at 2.6e5 ulp / 33%, RoPE applied at position `i+1` + at 5.1e5 ulp / 69%, `k_pe` left un-roped at 4.4e5 ulp / 93%. The + standard-configuration fp8-pool + append is bit-exact against + `e4m3(k * kv_scale_orig_quant)` (fp32 multiply, RN cast) for power-of-2 + scales; at `s=1.5` ~0.6% of elements sit one e4m3 ulp off that mirror + (never more) — see the fp8 section. The **MLA** fp8 append is bit-exact + against the same mirror at every scale run (1.0, 1.5, 2.0) and on both + halves of the latent row, roped `k_pe` included: the e4m3 grid is coarse + enough to swallow the rounding tie that costs the bf16 latent append its + one-ulp elements. Its gate is one e4m3 ulp plus the same + one-in-a-thousand inexact floor, and no run has approached it — across the + appended-row checks the fp8 MLA cases make (three scales, both rope tables, + all four call shapes), not one byte differs from the mirror, in either half. + Two things make that bit-exactness load-bearing rather than merely tidy. + First, it is the sharpest place the **rope table** shows up: a mirror roped + by the unscaled theta-10000 table instead of R1's YaRN one differs in 2336 + of the 6144 e4m3 bytes of a 96-token request (38.0%) against a gate that + allows 6 — ~389x — so the same comparison that passes on the certified + content demonstrably fails on the wrong one. Second, the comparison can see + a **racing** append: replaying this entry's own documented arming sequence + (a call whose new tokens do not map to distinct physical slots — here a + 96-token prefill with all three of its absolute pages aliased onto one + physical page) made an identical seeded call leave different pool bytes on + every repeat — five of five in a scratch probe, 11 177 to 12 513 bytes + apart — while the certified geometry reproduced bitwise. The shipped + control asserts exactly that (four armed repeats must not all agree, three + certified ones must), and it fired in five of five fresh processes. Both + controls run in the test file, in the same processes as the certified + cases. +- Multi-layer pool addressing (standard configuration, sm_100): certified + against state produced by a real 4-layer `KVCacheManager` (bf16, + GQA 8/2/128, `tokens_per_block` 32) — identity mapping rows, + layer-agnostic offsets (K `8p`, V `8p + 1`), prefill plus page-crossing + decode at `local_layer_idx` 0-3. Each layer's append landed bit-exactly + in its own slab group with sibling layers bitwise untouched, and decode + read back each layer's own history. A doctored-mapping call + (`local_layer_idx=1`, row 1 rewritten to layer-in-pool 3) landed its + append in layer 3's slabs while layer 1 stayed bitwise intact: the + pool-base shift is `mapping[local_layer_idx][1] * kv_factor` slabs — + the row's layer column, `local_layer_idx` only selecting the row. The + same 4-layer prefill + page-crossing-decode certification was re-run + over an fp8-e4m3 pool from a `DataType.FP8` manager (`quant_mode=128`, + GQA 32/8/128, `tokens_per_block` 32, `s = 1.0` — identical mapping and + offset layout): every layer's context and generation append landed + bit-exactly against the `e4m3(k * kv_scale_orig_quant)` mirror in its + own slab group, sibling layers stayed bitwise untouched by every call, + and each layer's decode read back its own history — pinning that the + layer-base shift is computed in e4m3-sized slabs under `quant_mode=128`. + The doctored-mapping probe was not repeated under fp8. +- Accuracy (bf16, sm_100) vs fp32 references, compared with `atol=5e-3, + rtol=1.6e-2` (not the default bf16 `atol=1e-5`): standard max abs err + 1.6e-2 on magnitude-~1 outputs (4-layer shared pool the same: 1.6e-2 + prefill, 7.8e-3 decode); MLA fresh-prefill context max abs err + 1.6e-2; MLA generation max abs err 7.8e-3; no-append MLA context max abs + err 1.6e-2 (one-shot cached-KV), 7.8e-3 (chunked partial passes) rising to + 1.6e-2 on the final causal pass, 1.6e-2 (chunk passes merged vs a + single-pass full-range reference). The largest fraction of the combined + `atol + rtol*|ref|` allowance any case in the test file uses is 66% — the + maximum sits on the `q_scaling` sweep's `q_scaling = 1.0` context run at + `H = 128` (66.1%), just ahead of the paged-context 4-layer shared-pool case + (65.8%), the page-32 cached-KV no-append MLA context case at `H = 32` + (63.9%), the page-32 `H = 128` fresh prefill (63.3%), the fp8-pool + multi-layer context (61.5%) and the YaRN-table case (61.3%); + this bullet's standard GQA/MHA cases sit at 36-52% and its MLA + single-pass cases at 36-64% (the sink, sliding-window and paged-context + bullets below carry their own numbers; the latent append's one-ulp gate + is a separate, tighter one, and one case does approach it — the page-32 + `H = 32` mixed batch reaches 84% of the one-ulp gate on its roped + elements). Every figure in this bullet and this ceiling + were re-measured across the whole file in the state the current receipt + certifies, by an instrumented sibling run of the same seeded cases. + Head geometry does not move these numbers: the two shipped-target + geometries (32/8/128, 32/4/128) and the MQA / odd-ratio / odd-count / + `head_size` 256 sweep all land at max abs err 1.6e-2 and at most 53% of + the allowance, the same envelope the 8/2/128 and 4/4/64 cases sit in. + A scratch probe of wider ratios (64/8/128, 64/1/128, 128/1/128) stayed at + the same 1.6e-2 max abs err but pushed decode to ~75% of the allowance — + headroom narrows as the head group grows, so re-measure before adopting a + ratio above 16. The 64/8/64 geometry without sinks sits in the same + envelope: max abs err 1.6e-2, at most 50% of the allowance. +- MLA head count does not move the numbers either (bf16, sm_100, same fp32 + references and `atol=5e-3, rtol=1.6e-2`). At `H = 32`, page 64: + fresh-prefill + context max abs err 1.6e-2 (53% of the allowance), cached-KV no-append + context 7.8e-3 (50%), generation decode 7.8e-3 (37%), mixed batch 1.6e-2 + context (51%) / 7.8e-3 decode (36%) — the same magnitudes the `H = 16` + runs of those four cases produce in the same process (1.6e-2 at 49%, + 1.6e-2 at 40%, 7.8e-3 at 39%, 1.6e-2 at 38% / 7.8e-3 at 36%), so doubling + the head count costs no headroom. + Halving it below 16 costs none either, and the `H = 8` cell has the + mildest worst case of the three (45.9% against 48.5% at `H = 16` and 63.9% + at `H = 32`): at page 64, fresh-prefill context 1.6e-2 (46%), + cached-KV no-append context 1.6e-2 (40%), generation decode 7.8e-3 (37% + then 36%) over two steps, mixed batch 7.8e-3 context (38%) / 7.8e-3 decode + (36%); at page 32, fresh-prefill 7.8e-3 (39%), cached-KV no-append 7.8e-3 + (40%), decode 7.8e-3 (37% then 34%), mixed batch 7.8e-3 context (37%) / + 7.8e-3 decode (36%). Case by case, every `H = 8` figure is at or below the + `H = 32` one for the same case; against `H = 16` the only one above is the + page-64 mixed batch's context row, 38.5% where `H = 16` reads 37.9%. + Multiplying it by 4 to `H = 128` costs no headroom either, at page 32: + on the baseline rope/scale, fresh-prefill context 1.6e-2 (63%), generation + decode 7.8e-3 (40%) then 1.6e-2 (40%) over two steps, mixed batch 1.6e-2 + context (52%) / 7.8e-3 decode (37%), cached-KV no-append context 1.6e-2 + (48%); at the full R1 cell (YaRN table, `q_scaling ≈ 0.53366`) the same + four cases read 1.6e-2 (61%), 1.6e-2 (49% then 43%), 1.6e-2 context (58%) / + 7.8e-3 decode (37%), and 1.6e-2 (54%). Same max abs errors as every smaller + count; the fractions sit inside the 34-64% band the others span, and the + head count is not what orders them (the `H = 128` fresh prefill's 63% is + the highest of its set, where at `H = 8` and `H = 32` the no-append case + is). + The paged latent append held to its gate (see the append bullet above) at + all four head counts, including page-boundary crossings, and every count + leaves `latent_cache`, `q_pe`, `q`/`k`/`v` and the pool untouched exactly + where the per-flavor tables say they do. At `H = 8` the appended rows came + back bitwise exact on both halves in all six append-checking cases (no + fp32→bf16 rounding tie among their roped elements), so the one-ulp gate was + not approached there; the same held at `H = 128`, whose 7 append-checking + cases make 16 appended-row checks between them — every one bitwise exact on + both halves, on both table contents. +- MLA page size does not move the numbers either (bf16, sm_100, same fp32 + references and `atol=5e-3, rtol=1.6e-2`). At `tokens_per_block = 32`, + `H = 32`: fresh-prefill context max abs err 1.6e-2 (50% of the allowance), + generation decode 1.6e-2 (42%) then 7.8e-3 (37%) over two steps, mixed + batch 1.6e-2 context (42%) / 7.8e-3 decode (37%), cached-KV no-append + context 1.6e-2 (64%); the page-32 decode case at `H = 16` sits at 7.8e-3 / + 37% on both steps, and the four page-32 cases at `H = 8` (numbers in the + head-count bullet above) span 34-40%. Those are the page-64 + magnitudes and, the no-append case aside, the page-64 band. The no-append flavor never reads or writes + the pool, so its 64% tracks that case's KV geometry (prefixes 96/31/0 + reaching 128/40/25 keys) rather than the page size — it was the file's + ceiling until the `H = 128` `q_scaling` sweep edged past it at 66%. + The paged latent append held to the same gate at page 32, the generation + call still wrote nothing, and the no-append call still left `q`/`k`/`v` + and the pool bitwise intact. +- MLA q-LoRA rank moves nothing whatsoever, at either shipped head count + (bf16, sm_100, `tokens_per_block = 32`; `H = 32` is the deepseek-v3-lite + tp1 cell, `H = 8` its tep4 slice; same fp32 references and `atol=5e-3, + rtol=1.6e-2`). At `q_lora_rank = 0`, `H = 32`: fresh-prefill context max + abs err 1.6e-2 (47% of the allowance), generation decode 7.8e-3 (36% then + 37%) over two steps, cached-KV no-append context 1.6e-2 (44%). At + `q_lora_rank = 0`, `H = 8`: fresh-prefill context 1.6e-2 (44%), + generation decode 7.8e-3 (37% then 36%), cached-KV no-append context + 1.6e-2 (37%). All six sit between 35.9% and 46.7% of the allowance — the + band the page-32 cases above sit in — and the two head counts interleave + there rather than separating (the `H = 8` decode's first step, 36.5%, is + the one figure above its `H = 32` counterpart, 35.9%). Each fresh-prefill + case's appended latent rows came back bitwise exact on both halves (no + fp32→bf16 rounding tie among either run's 8256 roped elements, so the + one-ulp gate was not even approached), and the no-append calls again left + `q`/`k` and the pool bitwise intact. Since the four-value rank sweep is + bitwise identical throughout (see *MLA q-LoRA rank*), these are element + for element the numbers the `q_lora_rank = 1536` runs of the same six + cases produce. The rank reaches neither kernel selection nor the compile + cache at either count: a process running only the two four-value sweeps + generated + `...HQk576HV512HVPerCta128PagedKvDenseP32VarSeqQ16Kv128StaticSwapsAbForGen` + exactly **once** (the `H = 32` sweep, 8 decode calls) and + `...HQk576HV512HVPerCta128PagedKvDenseP32VarSeqQ8Kv128StaticSwapsAbForGen` + exactly **once** (the `H = 8` sweep, 8 decode calls); and in the entry's + full test run, measured with and without the `H = 8` sweep present, the + counts are identical either way — the first variant twice, once each for + the two head counts that share it (`H = 16` and `H = 32`), the second once + (`H = 8`) — so those 8 extra decode calls add no compile at all, against + one compile per value on a real selection axis. +- MLA `q_scaling` is a live axis of both phases, and only of the softmax + scale (bf16, sm_100, `H = 128`, `tokens_per_block = 32`, YaRN table). Swept + over 1.0, 0.53366, 2.0 and 0.25 on identical inputs, each run compared + against fp32 references built at all four scales: the matching reference is + the only one within tolerance, at 30-66% of the allowance, while the full + 4x4 cross matrix spans 20.1x-178x the allowance in both phases — the + tightest cell is 1.0 against 2.0 in the context phase (20.1x), and a + reference at 1.0 for a run at any other value (the "argument silently + ignored" hypothesis) spans 20.1x-70.5x. The axis is finer than the sweep + resolves: 0.5 against 0.53366, a 6.3% change of scale, separates by only + 2.95x-3.61x, so the shipped values are kept apart rather than the gate + loosened. Across the sweep the paged latent pool comes back **bitwise + identical** — `q_scaling` reaches the softmax scale and nothing else, not + the RoPE and not the append. It does not reach kernel selection either: the + sweep's four decode calls compile nothing beyond the one `H = 128` decode + kernel the file already pays for. +- The MLA rope table is read for its content only (bf16, sm_100, `H = 128`, + `tokens_per_block = 32`). A fresh-prefill context call driven by the + DeepSeek-R1-0528 YaRN table reproduces a table-aware fp32 reference at 61% + of the allowance and sits 27x outside a reference built from the unscaled + theta-10000 table; the converse run sits 27x outside the YaRN reference. + That separation is position-dependent — 8.6x over a 96-token sequence, + 16.9x at 256, 27.0x at 512, 31.7x at 960 — because YaRN rescales only the + low-frequency half of the spectrum, so a short probe would not have + resolved the two tables at all. Meanwhile the seven scalar rope arguments + and `rotary_inv_freq` are bitwise inert beside the table (see *In-kernel + RoPE arguments*), and a plain-torch rebuild of the YaRN formula matches + TensorRT-LLM's own table to 1.9e-6 max abs on cos/sin values in [-1, 1] — + a few fp32 ulp from a different evaluation order of the same blend, where + formula slips land 4-5 orders of magnitude away (dropping the + interpolation, or moving `beta_fast` 32 → 64, both differ by 2.0; taking + the amplitude as `mscale` rather than the `mscale`/`mscale_all_dim` ratio + differs by 3.7e-1). +- Accuracy with attention sinks (bf16, sm_100, 64/8/64 causal) vs fp32 + sink-aware references at the same `atol=5e-3, rtol=1.6e-2`: max abs err + 1.6e-2, at most 50% of the allowance, across context prefill, two decode + steps, a mixed batch, decode over 600- and 2000-token histories, and the + constant-V probe that reads the softmax row mass directly. Every case + also had to sit **outside** the tolerance band of two rival references — + "sink silently ignored" and "sink pre-scaled" — and the tightest of those + separations was 12.8x the allowance (decode over 2000 cached tokens, + where a small sink is genuinely negligible against 2000 keys, so the test + draws the sink near `log(kv_len)` to keep the mechanism observable). +- Accuracy with the sliding window (bf16, sm_100, 64/8/64 causal, `W` 128 / + 100 / 33) vs fp32 windowed references, sinks on and off, at the same + `atol=5e-3, rtol=1.6e-2`: max abs err 3.9e-3 (at most 30% of the + allowance) across prefill past the window, decode, a mixed batch, a + 2000-token history, a 300-token prefill, and 250 wrapping decode steps; + the per-layer alternating-window case sits at 1.6e-2 / 49%, its + full-attention layers being the wider ones. Each case is additionally + gated **outside** the tolerance band of a "window ignored" (full causal) + reference on top of the two sink rivals. The exact per-key weights read + out by the one-hot probe match `1 / (n + exp(sink[h]))` to 3.6e-3 + relative — 59% of the 6e-3 (~1.5 bf16 ulp) gate — while out-of-window + keys are bitwise zero, which is what pins the boundary at `W` rather than + `W + 1` keys. +- Accuracy on the paged-context path (bf16, sm_100, `use_paged_context_fmha + = True`) vs the same fp32 references at the same `atol=5e-3, rtol=1.6e-2`: + max abs err 1.6e-2 across every case — fresh prefill and decode at the + three shipped geometries (53% of the allowance), cached prefixes at + 32/8/128 and 32/4/128 including page-boundary crossings and a mixed + cached/fresh/generation batch (54%), the six-prefix page-grid batch + (59%), two chunked-prefill loops (42%), and the 4-layer shared-pool + cached-context cycle (66%, the file's ceiling). The gpt-oss cell sits + lower: 7.8e-3 / 40% over sinks with and without the 128-token window, and + 1.6e-2 / 49% on the 200-token windowed prefill that sets up the read-set + probe; each sink case is gated outside the sink-ignored (11.9x-56x), + sink-pre-scaled (11.3x-53x) and — with the window — window-ignored + (32x-43x) rivals. The one-hot probe on a cached-prefix context row reads + the per-key weights back to 2.2e-3 relative, 36% of the same 6e-3 gate, + with out-of-window keys bitwise zero. Wrong batch states land 31x-318x + outside the band + (aliased pages 31x, `use_paged_context_fmha=False` over a cached prefix + 51x, `context_lengths` = full KV length 101x or 0 53x, a decoy in place + of an in-window cached page 248x-318x); the test gates them at 20x. +- Sliding-window kernel variants (sm_100): the FMHA family is chosen from + the batch's KV length against the window, **not** from `max_seq_len`. A + call moves from `...PagedKvDense...` / `...PackedQkvCausal...` to + `...SlidingOrChunkedCausal...` exactly when the batch's longest KV run + *exceeds* `attention_window_size` — at window 128, KV 127 and 128 take the + dense kernel and KV 129 the sliding one. `max_seq_len` does not enter the + choice: at window 128 and KV 41 the values 128, 256 and 1024 all give + `...PagedKvDense...`, and at KV 200 both 128 and 256 give + `...SlidingOrChunkedCausal...`; the dispatcher's `maxSeqLenKv` is the + batch's real KV length in every case (41, 128, 129, 200), never clamped to + the window. One windowed sequence therefore pays two decode JIT compiles + over its lifetime: `...PagedKvDense...` while its history still fits the + window, `...SlidingOrChunkedCausal...` from the first step past it. Decode + JIT-compiles + `fmhaSm100aKernel_QkvBfloat16OBfloat16H64PagedKvSlidingOrChunkedCausalP32VarSeqQ8Kv128StaticSwapsAbForGen` + once (~6 s) and `...SlidingOrChunkedCausalP32MultiCtasKvCga...` for long + histories — selected on the *total* cached length, so a 128-token window + over 2001 cached tokens still takes it; the context variants are + precompiled (no JIT observed). Where both families can compute the same + row (a short sequence alone, then batched behind one past the window) they + agree bit for bit. +- Paged-context kernels are precompiled too: across a full run of this + entry's test file every JIT compile logged is a `...ForGen` decode + variant, none a context one, so switching `use_paged_context_fmha` on + costs no compile. It does not steer the decode kernel either — a decode + step at both flag values is bitwise identical. +- Attention-sink quirks (sm_100): the sink pointer is read as `num_heads` + fp32 values with no size, stride, or device check (see *Preconditions*). + Enabling sinks costs no extra JIT compile: at 64/8/64 the decode kernel + `fmhaSm100aKernel_QkvBfloat16OBfloat16H64PagedKvDenseP32VarSeqQ8Kv128StaticSwapsAbForGen` + was generated once and then served both the sink and the no-sink calls in + the same process — adding or removing the argument triggered no further + generation. Long histories select the + `...H64PagedKvDenseP32MultiCtasKvCga...` variant, which folds partial + softmax states across CTAs and applies the sink in that reduction. A + `-inf` or `-100.0` sink reproduces the no-sink output bit for bit — the + sink adds exactly one term to the softmax denominator and changes nothing + else in the code path. +- Accuracy (fp8-e4m3 pool, standard configuration, sm_100): context matches + the bf16 numbers above + (same kernel, same inputs). Decode vs the fp32 quantization-aware + reference of the fp8 section: max abs err 3.1e-2 across a 10-seed sweep, + at most 44% of the `2^-4 * (1 + |ref|)` allowance; a 200-token-history + decode sat at 9.3e-3. A reference that skips the read-side dequant + diverges to ~4.4e-1 — the gate separates the failure mode by ~7x. The + 4-layer shared pool sits at the same levels per layer: context max abs + err 1.6e-2 (61% of the bf16-tolerance allowance — the highest fraction + observed on any certified case), decode max abs err 3.5e-2 (46% of the + fp8 allowance). +- Accuracy (fp8-e4m3 **latent** pool, MLA, sm_100, `H = 128`, page 32, + `s = 1.0`) vs fp32 references over e4m3-rounded operands, compared with + `atol = 2^-3, rtol = 2^-4`: context prefill (96 + 33 tokens) max abs err + 7.0e-2, 49% of the allowance; generation decode 2.1e-2 / 2.3e-2 over two + steps, 17-18%; the `s` = 1.5 and 2.0 decodes 2.1e-2, 16%. The floor is 2^-3 + rather than the standard configuration's 2^-4 because the residual is the + kernel's e4m3 handling of the softmax probabilities, whose error scales + with the `|V| ~ 1` rows rather than with the output element: at a 2^-4 + floor the context case runs 0.72-1.08 of the allowance over six seeds — + one seed *over* — because a `129 x 128 x 128` context output samples that + noise 16x more often than a two-row decode does. Rival separations from + the same runs: the bf16-KV reference (cache quantization ignored) 26x + outside the **bf16** band and 1.3-1.6x outside this one; the `s = 1.0` + math against `s` = 1.5 / 2.0 context runs 21x / 48x; a zeroed + `mla_bmm1_scale[1]` 4.2x; a decode whose bmm scales omit the kv scale + 2.9x at `s = 1.5` and 3.6x at `s = 2.0`. What the loosened gate cannot resolve is the query's + own e4m3 rounding in decode (0.23x, against the correct model's 0.17x), + so no claim rests on it; the "is the context math fp8" question is settled + bitwise instead by the peaked-softmax V readout described in that + section. +- Accuracy over the fp8 latent pool at the **complete DeepSeek-R1-0528 cell** + (YaRN table + `q_scaling = 1/mscale²`, same `atol = 2^-3, rtol = 2^-4`): + fresh-prefill context (96 + 33) 40.1% of the allowance, generation decode + 19.8% / 21.4% over two steps, mixed batch 36.4% (context) and 14.7% + (decode), no-append context (prefixes 96/31/0 reaching 128/40/25) 41.4%, + and the no-append realistic case of the operand/scale test 34.6% — the same + 14-49% envelope the baseline-cell fp8 cases sit in, so neither the rope + table nor the softmax scale costs headroom. Rival separations measured on + these runs: a reference at `q_scaling = 1.0` sits 12.7x outside the band in + fresh prefill, 10.8x in the no-append flavor and 4.2x in decode; an + unscaled-rope-table reference 3.7x in fresh prefill at 96 tokens (10.5x at + 512, since YaRN rescales only the low-frequency half of the spectrum and the + two tables diverge with distance); the bf16-KV reference 49.8x outside the + **bf16** band in fresh prefill, 58.1x in the mixed batch, and 2.2x outside + the fp8 band in the no-append flavor. The rope table is pinned far harder + by the append than by any of these: see the append bullet above. +- Softmax stats (no-append MLA context, sm_100): the emitted `(max, sum)` + match an fp32 reference over the same bf16 inputs to 2e-6 abs (max stat) + / 1.2e-6 rel (sum stat) — natural-log domain, scaled logits. Feeding each + pass's output/stats into `trtllm::merge_chunked_attention_for_mla` under + the production copy/merge/skip plan reproduced single-pass full-range + attention (output within the tolerances above, merged stats at the same + 1e-6-level agreement), certifying the stats as directly consumable by + that sibling op. +- V-stride hard-coding: for bf16 separate-QKV context MLA, the trtllm-gen + kernel computes V's row stride as `num_kv_heads * (head_size - 64 + + v_head_dim)` elements — the packed kv-projection width — rather than + reading it from the tensor (per source; the binding's `v_stride_in_bytes` + plumbing is unused in this version). A scratch probe passing a fully + contiguous `[Tkv, H*v_head_dim]` V returned wrong results (max abs err + ~4) and such a call reads out of bounds past the buffer. Always pass the + strided split view, in both context flavors. +- `chunked_prefill_buffer_batch_size` is consumed only for fp8-context-MLA + workspace sizing per source; inert in the certified bf16 configurations + (certified at 1, including the chunked partial passes). Every fp8 MLA + context case — both flavors — passes 1 as well and none was swept over it, + so the argument is certified at that one value on the path where source says + it is live. The chunked partial-pass *pattern* of the no-append flavor + (`mask_type=0` + `softmax_stats_tensor`) is not certified over an fp8 pool; + its one-shot cached-KV pattern is. +- Re-running the same prepared batch overwrites the same cache slots + (append position derives from `sequence_length` minus new-token count), + so a repeated call is idempotent, not double-appending. The MLA + generation and no-append MLA context calls write nothing at all (pool + bitwise unchanged). +- MLA generation head geometry: trtllm-gen decode kernels exist for + `(C+R, C)` of (576, 512) and (320, 256) per source; only (576, 512) is + certified, at 8, 16, 32 and 128 query heads (the compiled kernel reports + `HVPerCta128` at all four — the head count moves the q-tile, not the + per-CTA head-value width, and it stops moving it above 16). +- Sibling ops exist for adjacent roles: + `torch.ops.trtllm.attn_custom_op_inplace` and + `torch.ops.trtllm.mla_custom_op_inplace` (registered-layer wrappers over + the same attention core, state via model extra attrs), + `torch.ops.trtllm.create_attn_outputs` / `create_mla_outputs` + (output-buffer allocation), `torch.ops.trtllm.mla_rope_generation` + (MLA generation preprocessing: q_pe RoPE + latent append + the scheduler + buffers this op consumes), + `torch.ops.trtllm.mla_rope_append_paged_kv_assign_q`, + `load_paged_kv_cache_for_mla`, `load_chunked_kv_cache_for_mla` and + `merge_chunked_attention_for_mla` (the MLA cached-KV / chunked-prefill + context flow), and the split-phase trtllm-gen bindings + `thop.trtllm_gen_context_preprocess` / `thop.trtllm_gen_generation_preprocess` + / `thop.trtllm_gen_context_postprocess`. diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/thop_attention.py b/tensorrt_llm/_torch/staircase/catalog/attention/thop_attention.py new file mode 100644 index 000000000000..9ee461204103 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/attention/thop_attention.py @@ -0,0 +1,281 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Fused attention core with fully explicit state: paged KV-cache append + masked FMHA. + +Wraps the pybind binding ``tensorrt_llm.bindings.internal.thop.attention`` +(approved policy exception to the ``torch.ops.trtllm.*`` entry shape): the +same C++ attention op behind the TRTLLM backend, but every piece of batch +state and layer config arrives as an explicit argument — no registered +layers, no thread-local metadata. +""" + +from typing import Optional + +import torch + +from tensorrt_llm.bindings.internal import thop + + +def thop_attention( + q: torch.Tensor, + k: Optional[torch.Tensor], + v: Optional[torch.Tensor], + output: torch.Tensor, + output_sf: Optional[torch.Tensor], + workspace_: Optional[torch.Tensor], + sequence_length: torch.Tensor, + host_past_key_value_lengths: torch.Tensor, + host_total_kv_lens: torch.Tensor, + context_lengths: torch.Tensor, + host_context_lengths: torch.Tensor, + host_request_types: torch.Tensor, + max_context_q_len_override: Optional[int], + kv_cache_block_offsets: Optional[torch.Tensor], + host_kv_cache_pool_pointers: Optional[torch.Tensor], + host_kv_cache_pool_mapping: Optional[torch.Tensor], + cache_indirection: Optional[torch.Tensor], + kv_scale_orig_quant: Optional[torch.Tensor], + kv_scale_quant_orig: Optional[torch.Tensor], + out_scale: Optional[torch.Tensor], + rotary_inv_freq: Optional[torch.Tensor], + rotary_cos_sin: Optional[torch.Tensor], + latent_cache: Optional[torch.Tensor], + q_pe: Optional[torch.Tensor], + block_ids_per_seq: Optional[torch.Tensor], + attention_sinks: Optional[torch.Tensor], + is_fused_qkv: bool, + update_kv_cache: bool, + predicted_tokens_per_seq: int, + local_layer_idx: int, + num_heads: int, + num_kv_heads: int, + head_size: int, + tokens_per_block: Optional[int], + max_num_requests: int, + max_context_length: int, + max_seq_len: int, + attention_window_size: int, + beam_width: int, + mask_type: int, + quant_mode: int, + q_scaling: float, + position_embedding_type: int, + rope_dim: int, + rope_base: float, + rope_scale_type: int, + rope_scale: float, + rope_short_m_scale: float, + rope_long_m_scale: float, + rope_max_positions: int, + rope_original_max_positions: int, + use_paged_context_fmha: bool, + attention_input_type: Optional[int], + is_mla_enable: bool, + chunked_prefill_buffer_batch_size: Optional[int], + q_lora_rank: Optional[int], + kv_lora_rank: Optional[int], + qk_nope_head_dim: Optional[int], + qk_rope_head_dim: Optional[int], + v_head_dim: Optional[int], + rope_append: Optional[bool], + mrope_rotary_cos_sin: Optional[torch.Tensor], + mrope_position_deltas: Optional[torch.Tensor], + helix_position_offsets: Optional[torch.Tensor], + helix_is_inactive_rank: Optional[torch.Tensor], + attention_chunk_size: Optional[int], + softmax_stats_tensor: Optional[torch.Tensor], + is_spec_decoding_enabled: bool, + use_spec_decoding: bool, + is_spec_dec_tree: bool, + spec_decoding_generation_lengths: Optional[torch.Tensor], + spec_decoding_position_offsets_for_cpp: Optional[torch.Tensor], + spec_decoding_packed_mask: Optional[torch.Tensor], + spec_decoding_bl_tree_mask_offset: Optional[torch.Tensor], + spec_decoding_bl_tree_mask: Optional[torch.Tensor], + spec_bl_tree_first_sparse_mask_offset_kv: Optional[torch.Tensor], + sparse_kv_indices: Optional[torch.Tensor], + sparse_kv_offsets: Optional[torch.Tensor], + sparse_attn_indices: Optional[torch.Tensor], + sparse_attn_offsets: Optional[torch.Tensor], + sparse_attn_indices_block_size: int, + num_sparse_topk: Optional[int] = None, + sparse_attn_kv_lens: Optional[torch.Tensor] = None, + skip_softmax_threshold_scale_factor_prefill: Optional[float] = None, + skip_softmax_threshold_scale_factor_decode: Optional[float] = None, + skip_softmax_stat: Optional[torch.Tensor] = None, + cu_q_seqlens: Optional[torch.Tensor] = None, + cu_kv_seqlens: Optional[torch.Tensor] = None, + fmha_scheduler_counter: Optional[torch.Tensor] = None, + mla_bmm1_scale: Optional[torch.Tensor] = None, + mla_bmm2_scale: Optional[torch.Tensor] = None, + quant_q_buffer: Optional[torch.Tensor] = None, + flash_mla_tile_scheduler_metadata: Optional[torch.Tensor] = None, + flash_mla_num_splits: Optional[torch.Tensor] = None, + sage_attn_num_elts_per_blk_q: int = 0, + sage_attn_num_elts_per_blk_k: int = 0, + sage_attn_num_elts_per_blk_v: int = 0, + sage_attn_qk_int8: bool = False, + num_contexts: int = 0, + num_ctx_tokens: int = 0, + trtllm_gen_jit_warmup: bool = False, + aux_kv_cache_pool_ptr: Optional[int] = None, + is_cross: bool = False, + cross_kv: Optional[torch.Tensor] = None, + relative_attention_bias: Optional[torch.Tensor] = None, + relative_attention_max_distance: int = 0, + spec_decoding_target_max_draft_tokens: Optional[int] = None, + quant_scale_qkv: Optional[torch.Tensor] = None, + dsv4_inv_rope_cos_sin_cache: Optional[torch.Tensor] = None, + enable_dsv4_epilogue_fusion: bool = False, + # Added between 1.3.0rc21 and 1.3.0rc26. Defaults reproduce the behaviour + # the op had before they existed, and match what the in-tree caller + # (attention/backends/fmha/fallback.py) passes on a dense, non-sparse, + # non-folded path -- which is the path both migrated targets are on. + max_num_sequences: Optional[int] = None, + force_prepare_spec_dec_tree_mask: bool = False, + kv_norm_weight: Optional[torch.Tensor] = None, + kv_norm_eps: float = 1e-6, + skip_correction_threshold: float = 0.0, +) -> None: + """Run the fused attention core over one prepared batch — paged KV-cache + append + masked FMHA (standard), or the MLA context/generation phase + selected by attention_input_type; writes rows [:num_tokens] of output. + Returns None.""" + # num_heads must be an integer multiple of num_kv_heads. A context-only + # call with a non-multiple returns without raising on sm_100 (observed at + # 6q/4kv d128: only the first (num_heads // num_kv_heads) * num_kv_heads + # head columns of output are computed, the rest come back all-zero), so + # the geometry cannot be left for the op to reject. + assert num_heads % num_kv_heads == 0, ( + f"num_heads ({num_heads}) must be a multiple of num_kv_heads ({num_kv_heads})" + ) + # attention_sinks is consumed as a raw buffer of exactly num_heads fp32 + # values read from data_ptr(); only its dtype is validated by the op. On + # sm_100 a stride-2 view whose values were correct produced silently wrong + # output (the underlying memory is what gets read), a 32-element tensor was + # read 64 elements deep past its end, a 128-element one had its tail + # silently ignored, and an empty tensor silently disabled the sink. + if attention_sinks is not None: + assert attention_sinks.is_contiguous() and attention_sinks.numel() == num_heads, ( + f"attention_sinks must be contiguous with exactly num_heads " + f"({num_heads}) elements, got shape " + f"{tuple(attention_sinks.shape)} contiguous=" + f"{attention_sinks.is_contiguous()}" + ) + thop.attention( + q=q, + k=k, + v=v, + output=output, + output_sf=output_sf, + workspace_=workspace_, + sequence_length=sequence_length, + host_past_key_value_lengths=host_past_key_value_lengths, + host_total_kv_lens=host_total_kv_lens, + context_lengths=context_lengths, + host_context_lengths=host_context_lengths, + host_request_types=host_request_types, + max_context_q_len_override=max_context_q_len_override, + kv_cache_block_offsets=kv_cache_block_offsets, + host_kv_cache_pool_pointers=host_kv_cache_pool_pointers, + host_kv_cache_pool_mapping=host_kv_cache_pool_mapping, + cache_indirection=cache_indirection, + kv_scale_orig_quant=kv_scale_orig_quant, + kv_scale_quant_orig=kv_scale_quant_orig, + out_scale=out_scale, + rotary_inv_freq=rotary_inv_freq, + rotary_cos_sin=rotary_cos_sin, + latent_cache=latent_cache, + q_pe=q_pe, + block_ids_per_seq=block_ids_per_seq, + attention_sinks=attention_sinks, + is_fused_qkv=is_fused_qkv, + update_kv_cache=update_kv_cache, + predicted_tokens_per_seq=predicted_tokens_per_seq, + local_layer_idx=local_layer_idx, + num_heads=num_heads, + num_kv_heads=num_kv_heads, + head_size=head_size, + tokens_per_block=tokens_per_block, + max_num_requests=max_num_requests, + max_context_length=max_context_length, + max_seq_len=max_seq_len, + attention_window_size=attention_window_size, + beam_width=beam_width, + mask_type=mask_type, + quant_mode=quant_mode, + q_scaling=q_scaling, + position_embedding_type=position_embedding_type, + rope_dim=rope_dim, + rope_base=rope_base, + rope_scale_type=rope_scale_type, + rope_scale=rope_scale, + rope_short_m_scale=rope_short_m_scale, + rope_long_m_scale=rope_long_m_scale, + rope_max_positions=rope_max_positions, + rope_original_max_positions=rope_original_max_positions, + use_paged_context_fmha=use_paged_context_fmha, + attention_input_type=attention_input_type, + is_mla_enable=is_mla_enable, + chunked_prefill_buffer_batch_size=chunked_prefill_buffer_batch_size, + q_lora_rank=q_lora_rank, + kv_lora_rank=kv_lora_rank, + qk_nope_head_dim=qk_nope_head_dim, + qk_rope_head_dim=qk_rope_head_dim, + v_head_dim=v_head_dim, + rope_append=rope_append, + mrope_rotary_cos_sin=mrope_rotary_cos_sin, + mrope_position_deltas=mrope_position_deltas, + helix_position_offsets=helix_position_offsets, + helix_is_inactive_rank=helix_is_inactive_rank, + attention_chunk_size=attention_chunk_size, + softmax_stats_tensor=softmax_stats_tensor, + is_spec_decoding_enabled=is_spec_decoding_enabled, + use_spec_decoding=use_spec_decoding, + is_spec_dec_tree=is_spec_dec_tree, + spec_decoding_generation_lengths=spec_decoding_generation_lengths, + spec_decoding_position_offsets_for_cpp=spec_decoding_position_offsets_for_cpp, + spec_decoding_packed_mask=spec_decoding_packed_mask, + spec_decoding_bl_tree_mask_offset=spec_decoding_bl_tree_mask_offset, + spec_decoding_bl_tree_mask=spec_decoding_bl_tree_mask, + spec_bl_tree_first_sparse_mask_offset_kv=spec_bl_tree_first_sparse_mask_offset_kv, + sparse_kv_indices=sparse_kv_indices, + sparse_kv_offsets=sparse_kv_offsets, + sparse_attn_indices=sparse_attn_indices, + sparse_attn_offsets=sparse_attn_offsets, + sparse_attn_indices_block_size=sparse_attn_indices_block_size, + num_sparse_topk=num_sparse_topk, + sparse_attn_kv_lens=sparse_attn_kv_lens, + skip_softmax_threshold_scale_factor_prefill=skip_softmax_threshold_scale_factor_prefill, + skip_softmax_threshold_scale_factor_decode=skip_softmax_threshold_scale_factor_decode, + skip_softmax_stat=skip_softmax_stat, + cu_q_seqlens=cu_q_seqlens, + cu_kv_seqlens=cu_kv_seqlens, + fmha_scheduler_counter=fmha_scheduler_counter, + mla_bmm1_scale=mla_bmm1_scale, + mla_bmm2_scale=mla_bmm2_scale, + quant_q_buffer=quant_q_buffer, + flash_mla_tile_scheduler_metadata=flash_mla_tile_scheduler_metadata, + flash_mla_num_splits=flash_mla_num_splits, + sage_attn_num_elts_per_blk_q=sage_attn_num_elts_per_blk_q, + sage_attn_num_elts_per_blk_k=sage_attn_num_elts_per_blk_k, + sage_attn_num_elts_per_blk_v=sage_attn_num_elts_per_blk_v, + sage_attn_qk_int8=sage_attn_qk_int8, + num_contexts=num_contexts, + num_ctx_tokens=num_ctx_tokens, + trtllm_gen_jit_warmup=trtllm_gen_jit_warmup, + aux_kv_cache_pool_ptr=aux_kv_cache_pool_ptr, + is_cross=is_cross, + cross_kv=cross_kv, + relative_attention_bias=relative_attention_bias, + relative_attention_max_distance=relative_attention_max_distance, + spec_decoding_target_max_draft_tokens=spec_decoding_target_max_draft_tokens, + quant_scale_qkv=quant_scale_qkv, + dsv4_inv_rope_cos_sin_cache=dsv4_inv_rope_cos_sin_cache, + enable_dsv4_epilogue_fusion=enable_dsv4_epilogue_fusion, + max_num_sequences=max_num_sequences, + force_prepare_spec_dec_tree_mask=force_prepare_spec_dec_tree_mask, + kv_norm_weight=kv_norm_weight, + kv_norm_eps=kv_norm_eps, + skip_correction_threshold=skip_correction_threshold, + ) diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/thop_attention_test.py b/tensorrt_llm/_torch/staircase/catalog/attention/thop_attention_test.py new file mode 100644 index 000000000000..7fc5375d5d29 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/attention/thop_attention_test.py @@ -0,0 +1,6669 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the thop_attention catalog entry. + +The binding is stateless on the Python side: every piece of state it reads +is an explicit argument. The test builds that state for real — a caller- +owned paged KV pool, pool pointer/mapping tensors, block-offset tables +(hand-filled for single-layer pools, produced by a real KVCacheManager for +the multi-layer case), and the int32 host/device length tensors — and +compares against fp32 torch references. The certified configurations +exercised: + +1. Standard GQA/MHA: packed-QKV self-attention over a + [num_blocks, 2, Hkv, tokens_per_block, D] bf16 pool (kv_factor 2), with + the paged-cache append additionally checked bit-exactly. Head geometry is + swept across the two shipped-target geometries (32q/8kv and 32q/4kv, + d128) plus MQA, a non-power-of-2 GQA ratio, non-power-of-2 head counts, + and head_size 64/128/256; a non-multiple (Hq, Hkv) pair is a negative + test, since the op computes it silently wrong. +2. MLA (is_mla_enable=True): context prefill (separate q/k/v, in-kernel + GPT-J RoPE, latent-cache append) and generation decode (latent MQA over + the paged cache) over a [num_blocks, tokens_per_block, C+R] bf16 latent + pool (kv_factor 1), split into context_only / generation_only calls. + Every MLA case except the chunked one (3b) runs at three query-head + counts: 32 (the deepseek-v3-lite tp1 layer shape), 16 and 8 (its tep4 + slice), same C/R/nope/v. All three MLA call flavors additionally run at + latent-pool page size 32 (the engine default; 64 is the tuned value) on + page-32-specific geometry, at 32 and 8 query heads — and at 16 too for + the decode flavor, the one where the page size and the head count both + reach the compiled kernel. Each flavor is finally + swept over q_lora_rank (1536 = DeepSeek-V3, 0 = a checkpoint with no + q-LoRA, 1536 again as a determinism control, 4096 as an out-of-range + probe) on identical inputs, asserting every observable bitwise equal — + the argument turns out never to be read on the MLA paths; passing None + for it is a negative test. That sweep runs at both shipped head counts, + 32 (tp1) and 8 (tep4), because the two compile different decode kernels. +2b. The DeepSeek-R1-0528 cell, at page 32: all four MLA call flavors at + 128 query heads (128/128/192 context, 128/1/576 decode — attention DP + replicates the head count whole, so a dep4 rank runs all of them), first + on the baseline rope/scale to isolate the head count and then at the two + further values R1 moves: a YaRN-scaled rotary_cos_sin table and + q_scaling = 1/mscale^2 != 1. Both of those are also driven as axes of + their own — q_scaling swept over four values in both phases with every + run gated outside the other three's references, and the rope table pinned + as the *only* rope input the op reads (a plain-torch rebuild of the YaRN + formula drives the op; the seven scalar rope arguments and + rotary_inv_freq are bitwise inert beside it, with a table swap as the + control that the comparison can see a rope change at all). +2c. The DeepSeek-R1-0528 layer geometry over an **fp8-e4m3 latent pool** + (quant_mode 128, H = 128, page 32, one C+R-byte row per token): fresh + prefill and generation decode on the baseline rope/scale, at the + production kv scale s = 1.0 and swept over 1.5 and 2.0. Three things the + standard configuration's fp8 + section does not carry over, all measured here: the context FMHA runs on + e4m3 operands rather than staying bf16 (pinned bitwise by a peaked-softmax + V readout, and gated against the bf16-KV reference); the decode kernel + takes its query from quant_q_buffer and its two scales from + mla_bmm1_scale/mla_bmm2_scale, reading neither `q` nor the kv scale + tensors; and only the append side honours kv_scale_orig_quant, so the + context phase is correct at s = 1.0 alone while decode is correct at any s + once the caller folds it into the bmm scales. The requests' pages are + scattered out of order, which is what pins the e4m3-sized slab geometry. +2d. The **complete R1 cell over that fp8 pool** — YaRN table and + q_scaling = 1/mscale^2 together with the fp8 pool, which is what every + call of a DeepSeek-R1-0528 rank passes and what (2b) and (2c) cover only + separately. All four call flavors: fresh prefill, generation decode, the + mixed batch pairing them off one offsets table, and the no-append + (latent_cache=None) context an engine with block reuse runs for every + cached prefix — three of which had never been run over an fp8 pool at all. + Each run is gated against a q_scaling = 1.0 reference, prefill also + against an unscaled-rope-table one, and the appended rows carry the table + bitwise (with a mirror-swap control proving that comparison separates the + two tables at the positions in flight). The no-append flavor's K/V are + built the way a target builds them — cached prefix written into the pool + as e4m3(row/s), read back as bf16(float(byte) * s), pushed through a + kv_b_proj-shaped matmul, only the new tokens fresh — and it turns out to + quantize them to e4m3 all over again, settled bitwise by its own peaked-V + readout and swept over s. quant_mode 1152 (| FP8_1x128_128x128) and 384 + (| FP8_QDQ), which a checkpoint's quant config produces where a bare 128 + does not, are bit-identical to 128 on every flavor. One test is a positive + control rather than a certification: it replays this entry's own + documented torn-append race (a prefill whose pages are aliased onto one + physical page) and asserts the pool byte comparison every append check + rests on can see it. +2e. The MLA **generation** call at predicted_tokens_per_seq > 1 — one + speculative-decoding step, where a generation sequence contributes P query + rows (its draft chain) instead of one. Swept over P = 1, 2, 3, 4 (the + target's max_draft_len 0..3 plus one), at the R1 cell over the fp8 pool and + again over a bf16 one. The question the section exists to answer is the + within-block mask, and on sm_100 no mask tensor is involved: a linear-tree + draft has is_spec_decoding_enabled forced off there, so the answer has to + come from predicted_tokens_per_seq alone. It is read out directly — cache + rows carrying one-hot compressed_kv turn the output into the attention + weight vector — and the mask is bottom-right-aligned causal: row t attends + to keys [0, L - P + t] and its later siblings' columns come back bitwise + zero. mask_type 0 and 1 are bitwise identical here. Around that: the + token-major row order, the batch-state tensors this call reads at P > 1 + (cu_q_seqlens / cu_kv_seqlens / context_lengths contents still inert, + sequence_length still the KV extent, an all-zero + host_past_key_value_lengths skipping the call outright), mixed batches, + realistic-input accuracy against the causal fp32 reference with the + full-mask model as the rival, and a control that the decode-side comparison + can see a pool torn by the entry's own documented append race. +3. MLA context with latent_cache=None (no in-kernel RoPE or append): + a) one-shot cached-KV context — q pre-rotated upstream, K/V covering the + full [cached + new] range, causal mask bottom-right aligned; + b) chunked partial passes — K/V covering one chunk slice per call under + a padding mask with softmax_stats_tensor emitted, then a causal + new-token pass, folded together by the downstream trtllm merge op and + checked against a single-pass full-range fp32 reference. +4. Standard GQA over one paged pool shared by 4 layers: pool, layer->pool + mapping, and block offsets produced by a real multi-layer KVCacheManager + and consumed by the op as-is, with local_layer_idx 0-3 selecting mapping + rows. Per-layer outputs, bit-exact per-layer appends, sibling-layer + isolation, and a doctored-mapping call pinning the pool-base shift to + the mapping row's layer-in-pool column. +5. Standard GQA (32q/8kv d128, tpb 32) over an fp8-e4m3 paged pool + (quant_mode 128, fp32 [1] CUDA kv scale tensors): context prefill + (bf16 FMHA + bitwise-checked e4m3 append), decode over the quantized + cache (page-crossing, long-history, and mixed-batch), and non-1.0 + scale semantics (bit-exact at the power-of-2 scale, one-ulp-bounded at + a non-power-of-2 scale, dequant-on-read pinned by scale-aware + references). +6. The intersection of (4) and (5): one fp8-e4m3 paged pool shared by + 4 layers (quant_mode 128, s=1.0, GQA 32q/8kv d128, tpb 32), manager + state consumed as-is. Per-layer context and page-crossing decode + outputs, bit-exact per-layer e4m3 appends, and bitwise sibling-layer + isolation across both the context and generation append paths. +7. attention_sinks over the gpt-oss-120b tp1 geometry (64q/8kv d64, + causal, bf16 pool): the per-head sink as one extra softmax-denominator + column dropped from the output, gated in the context phase, the + generation phase, a mixed batch, and the multi-CTA-KV decode reduction; + a constant-V probe reading the softmax row mass directly; per-q-head + indexing; bitwise inertness of attention_sinks=None and of a sink that + cannot contribute; and rejection of non-fp32 dtypes (by the op) and of + wrong-size / strided buffers (by the wrapper). +8. The cyclic sliding window (attention_window_size < max_seq_len) with + sinks active, on the same 64q/8kv d64 geometry with window 128 — the + gpt-oss-120b sliding layer. A one-hot-V probe reads the attention + weights out directly and pins the attended key set to + [p - window + 1, p] with exact zeros outside it, in both phases; the + realistic-input cases are gated against sink-ignored, sink-pre-scaled + and window-ignored rivals in prefill, decode, a mixed batch, and a + 2000-token history. The caller-side contract is pinned too: the append + still lands at the absolute token position, so a bounded ring of pages + (absolute page j -> ring[j % P], P * tokens_per_block >= window) makes + the pool genuinely wrap while every decode keeps its full window; + fully-aged-out pages are bitwise unread and one page short is silently + wrong; a prefill longer than the window is legal in one call and its + output never reads the pool; sequence_length must stay global while + context_lengths and max_seq_len are inert; and layers with different + windows share one pool and one block-offset table in the same batch. +9. The paged-context execution path (use_paged_context_fmha=True), which an + engine prepares as soon as KV-cache reuse or chunked prefill is on: the + context FMHA sources K/V from the paged pool, so a context call may carry + a cached prefix (KV range beyond the q range, bottom-right-aligned causal + mask). Covered on the three shipped bf16 geometries: bitwise equivalence + with the packed path when nothing is cached; cached prefixes across the + page grid, chunked prefill loops, and mixed cached/fresh/generation + batches; the same over a real 4-layer KVCacheManager pool, where the read + now also goes through the layer base shift; and the gpt-oss cell — sinks + plus a 128-token window over a cached prefix past the window, with the + attended key set read out one-hot and the read page set measured page by + page. Wrong batch states (the flag left off, context_lengths carrying the + full KV length or 0, a replaced in-window cached page) are gated far + outside the tolerance band. +""" + +import math +from typing import Dict, List, NamedTuple, Optional, Tuple + +import torch +import torch.nn.functional as F + +from tensorrt_llm._torch.attention.backends.interface import RopeParams +from tensorrt_llm._torch.pyexecutor.resource_manager import CacheTypeCpp, DataType, KVCacheManager +from tensorrt_llm.functional import RotaryScalingType +from tensorrt_llm.llmapi.llm_args import KvCacheConfig +from tensorrt_llm.mapping import Mapping + +from .thop_attention import thop_attention + +assert torch.cuda.is_available(), "thop_attention requires a CUDA device" + +# The output is a softmax-weighted combination of unit-scale bf16 v rows; +# one bf16 ulp at magnitude 1 is 2^-8 ~= 3.9e-3. The kernel's online-softmax +# fp32 accumulation order differs from the reference and K/V round-trip +# through the bf16 KV cache. Observed on sm_100 across all cases: max abs +# err 1.6e-2, on magnitude-~1 elements (within rtol); at most 64% of the +# combined atol + rtol*|ref| allowance is used (the maximum is the page-32 +# cached-KV no-append MLA context case at 63.9%, then the fp8-pool +# multi-layer context at 62% and the MLA chunked-prefill final causal pass +# at 61%; the other MLA cases, at 8, 16 and 32 query heads and page sizes 32 +# and 64, sit at 36-53%). +ATOL = 5e-3 +RTOL = 1.6e-2 # torch default for bf16 + +MASK_CAUSAL = 1 # AttentionMaskType.causal +MASK_PADDING = 0 # AttentionMaskType.padding: no mask (bidirectional) + +QUANT_MODE_FP8_KV_CACHE = 128 # QuantMode.FP8_KV_CACHE (sole bit set) + +# fp8-e4m3 KV-pool decode tolerance. The reference dequantizes the mirrored +# e4m3 pool content and puts q through the same e4m3 round-trip the kernel +# applies (quantized with kv_scale_orig_quant — pinned by a non-power-of-2 +# scale sweep: only that scale collapses the error). The residual is the +# kernel's internal e4m3 handling of the softmax probabilities before BMM2 +# (decode kernel variant QkvE4m3...: every MMA operand is e4m3) plus fp8-MMA +# accumulation order. e4m3's max relative rounding error is 2^-4; observed +# max abs err vs this reference on sm_100: 3.6e-2 across 10 seeds x 2 decode +# steps (magnitude-~1 outputs), at most 53% of the 2^-4 * (1 + |ref|) +# allowance, while a reference that skips the kv dequant (the wrong-scale +# failure mode) sits at ~4.4e-1 — several times past the gate. +FP8_DECODE_ATOL = 2**-4 +FP8_DECODE_RTOL = 2**-4 + +# fp8-e4m3 latent-pool MLA tolerance, both phases. Under quant_mode 128 both +# MLA phases run their FMHA on e4m3 operands: context quantizes q/k/v itself, +# generation reads an e4m3 query out of quant_q_buffer and e4m3 rows out of +# the pool. The references round the same operands through e4m3 and then +# accumulate in fp32, so what is left is the kernel's own fp8 handling — +# chiefly the softmax probabilities, which are an MMA operand too and carry +# e4m3's 2**-4 relative error into a weighted sum of |V| ~ 1 rows. That error +# does not shrink with the output element's own magnitude, which is why the +# floor is 2**-3 rather than the 2**-4 of the standard-configuration decode +# gate above. Measured on sm_100 at H = 128, page 32 over six seeds: at a +# 2**-4 floor the 129-token context case runs 0.72-1.08 of the allowance +# (one seed *over* the gate — a 129x128x128 output samples the same noise 16x +# more often than a two-row decode does), and at this 2**-3 floor it runs +# 0.40-0.60 while decode runs 0.14-0.20 (0.28-0.39 at a 2**-4 floor). +# +# Discrimination. This gate is deliberately not where the "is the math really +# fp8" question is settled: the difference between quantized and unquantized +# operands is the same size as the gate, so the bf16-KV context reference — +# cache quantization ignored, what the bf16 surface is certified against — +# only misses it by 1.3-1.6x (its mean abs error is 2.5x the e4m3 model's, +# and it sits 23.5-32.3x outside the *bf16* band, which is what the context +# test gates on). That question is settled bitwise instead, by the peaked +# attention case: with the softmax collapsed onto one key the output *is* the +# V row, and it comes back bit-exactly e4m3(V) while the unquantized V row is +# 63 bf16 ulps away. What this gate does separate: the s = 1.0 math against +# s = 1.5 / 2.0 context runs (21x / 48x), a zeroed mla_bmm1_scale[1] (4.2x), +# and a decode whose mla_bmm scales leave the kv scale out (3.6x at s = 2.0). +# The rival it does not resolve at all is q left unquantized in the decode +# reference (0.23-0.24x against the correct model's 0.17x — e4m3 rounding of +# the query is invisible in an output that averages hundreds of cached rows), +# so no claim here rests on it. +FP8_MLA_ATOL = 2**-3 +FP8_MLA_RTOL = 2**-4 + + +class _PagedAttnEnv: + """Real op state: caller-owned paged KV pool + explicit metadata tensors. + + Also mirrors every K/V fed to the op per request, so references can be + built from the true sequence history, and the paged cache can be checked + against it. + """ + + def __init__( + self, + num_heads: int, + num_kv_heads: int, + head_dim: int, + tokens_per_block: int = 32, + num_blocks: int = 64, + max_batch: int = 4, + max_blocks_per_seq: int = 8, + pool_dtype: torch.dtype = torch.bfloat16, + quant_mode: int = 0, + kv_scaling_factor: Optional[float] = None, + attention_window_size: Optional[int] = None, + page_ring: Optional[int] = None, + ) -> None: + self.num_heads = num_heads + self.num_kv_heads = num_kv_heads + self.head_dim = head_dim + self.tokens_per_block = tokens_per_block + self.max_batch = max_batch + self.max_seq_len = max_blocks_per_seq * tokens_per_block + self.quant_mode = quant_mode + # None = no sliding window (window == max_seq_len). page_ring is the + # caller-side page budget: with it set, absolute page index j of a + # sequence is served by the j % page_ring-th physical page, so the pool + # holds a bounded window of the sequence and genuinely wraps. + self.attention_window_size = ( + self.max_seq_len if attention_window_size is None else attention_window_size + ) + self.page_ring = page_ring + self.rings: dict[int, List[int]] = {} + + # Single-layer, single-pool paged KV cache (HND layout). + self.kv_cache = torch.zeros( + num_blocks, + 2, + num_kv_heads, + tokens_per_block, + head_dim, + dtype=pool_dtype, + device="cuda", + ) + self.pool_pointers = torch.zeros(1, 2, dtype=torch.int64, device="cpu") + self.pool_pointers[0, 0] = self.kv_cache.data_ptr() + # KV-cache scaling factor s: dequant = quantized * s. Mirrors the + # production construction (quant_orig = s, orig_quant = 1/s, both + # fp32 [1] CUDA); None = unquantized pool, scale args passed as None. + if kv_scaling_factor is None: + self.kv_scale_quant_orig: Optional[torch.Tensor] = None + self.kv_scale_orig_quant: Optional[torch.Tensor] = None + else: + factor = torch.full((1,), kv_scaling_factor, dtype=torch.float32, device="cuda") + self.kv_scale_quant_orig = factor + self.kv_scale_orig_quant = 1.0 / factor + self.pool_mapping = torch.zeros(1, 2, dtype=torch.int32, device="cpu") + self.block_offsets = torch.zeros( + 1, max_batch, 2, max_blocks_per_seq, dtype=torch.int32, device="cuda" + ) + # Auto-resized in place by the op on first call. + self.workspace = torch.empty(0, dtype=torch.int8, device="cuda") + + self._next_free_page = 0 + self.pages: dict[int, List[int]] = {} + self.prompt_lens: dict[int, int] = {} + self.k_history: dict[int, List[torch.Tensor]] = {} + self.v_history: dict[int, List[torch.Tensor]] = {} + + def add_request(self, request_id: int, prompt_len: int) -> None: + self.pages[request_id] = [] + self.prompt_lens[request_id] = prompt_len + self.k_history[request_id] = [] + self.v_history[request_id] = [] + + def cached_len(self, request_id: int) -> int: + return sum(t.shape[0] for t in self.k_history[request_id]) + + def _ensure_pages(self, request_id: int, total_tokens: int) -> None: + tpb = self.tokens_per_block + needed = (total_tokens + tpb - 1) // tpb + pages = self.pages[request_id] + if self.page_ring is not None: + ring = self.rings.setdefault(request_id, []) + while len(ring) < self.page_ring: + ring.append(self._next_free_page) + self._next_free_page += 1 + while len(pages) < needed: + pages.append(ring[len(pages) % self.page_ring]) + return + while len(pages) < needed: + pages.append(self._next_free_page) + self._next_free_page += 1 + + def call_op( + self, + qkv: torch.Tensor, + seq_lens: List[int], + num_contexts: int, + request_ids: List[int], + mask_type: int, + attention_sinks: Optional[torch.Tensor] = None, + record: bool = True, + seq_lens_override: Optional[List[int]] = None, + ctx_lens_override: Optional[List[int]] = None, + max_seq_len_override: Optional[int] = None, + use_paged_context_fmha: bool = False, + ) -> torch.Tensor: + """One thop_attention call over explicitly constructed batch state. + + record=False skips the K/V history mirror, so the identical call can be + repeated (the append is idempotent) to compare code paths bit for bit. + The *_override arguments feed the op varied batch state + (sequence_length, context_lengths, max_seq_len) to pin what each one + actually has to carry — deliberately wrong values in the batch-state + probes, and the correct per-call context length on the paged-context + path, where a context row's context_lengths is this call's new-token + count rather than the registered prompt length. + use_paged_context_fmha=True selects the paged-context execution path + (see the paged-context section). + """ + ns = len(request_ids) + kv_lens = [] # cached + new, per sequence + ctx_lens = [] # prompt length, per sequence + for rid, new in zip(request_ids, seq_lens): + total = self.cached_len(rid) + new + assert total <= self.max_seq_len + self._ensure_pages(rid, total) + kv_lens.append(total) + ctx_lens.append(self.prompt_lens[rid]) + # K page offset of page p is 2*p, V is 2*p + 1 (single-layer pool). + row = self.block_offsets[0, len(kv_lens) - 1] + for j, p in enumerate(self.pages[rid]): + row[0, j] = 2 * p + row[1, j] = 2 * p + 1 + if seq_lens_override is not None: + kv_lens = list(seq_lens_override) + if ctx_lens_override is not None: + ctx_lens = list(ctx_lens_override) + + req_types = [0 if i < num_contexts else 1 for i in range(ns)] + total_ctx_kv = sum(kv_lens[:num_contexts]) + total_gen_kv = sum(kv_lens[num_contexts:]) + num_ctx_tokens = sum(seq_lens[:num_contexts]) + + num_tokens = sum(seq_lens) + output = torch.empty( + num_tokens, + self.num_heads * self.head_dim, + dtype=qkv.dtype, + device=qkv.device, + ) + thop_attention( + q=qkv, + k=None, # packed QKV rides inside q + v=None, + output=output, + output_sf=None, + workspace_=self.workspace, + sequence_length=torch.tensor(kv_lens, dtype=torch.int32, device="cuda"), + host_past_key_value_lengths=torch.tensor(kv_lens, dtype=torch.int32), + host_total_kv_lens=torch.tensor([total_ctx_kv, total_gen_kv], dtype=torch.int32), + context_lengths=torch.tensor(ctx_lens, dtype=torch.int32, device="cuda"), + host_context_lengths=torch.tensor(ctx_lens, dtype=torch.int32), + host_request_types=torch.tensor(req_types, dtype=torch.int32), + max_context_q_len_override=None, + kv_cache_block_offsets=self.block_offsets, + host_kv_cache_pool_pointers=self.pool_pointers, + host_kv_cache_pool_mapping=self.pool_mapping, + cache_indirection=None, + kv_scale_orig_quant=self.kv_scale_orig_quant, + kv_scale_quant_orig=self.kv_scale_quant_orig, + out_scale=None, + rotary_inv_freq=None, + rotary_cos_sin=None, + latent_cache=None, + q_pe=None, + block_ids_per_seq=None, + attention_sinks=attention_sinks, + is_fused_qkv=True, + update_kv_cache=True, + predicted_tokens_per_seq=1, + local_layer_idx=0, + num_heads=self.num_heads, + num_kv_heads=self.num_kv_heads, + head_size=self.head_dim, + tokens_per_block=self.tokens_per_block, + max_num_requests=self.max_batch, + max_context_length=self.max_seq_len, + max_seq_len=( + self.max_seq_len if max_seq_len_override is None else max_seq_len_override + ), + attention_window_size=self.attention_window_size, + beam_width=1, + mask_type=mask_type, + quant_mode=self.quant_mode, + q_scaling=1.0, + position_embedding_type=0, # no in-kernel RoPE + rope_dim=0, + rope_base=10000.0, + rope_scale_type=0, + rope_scale=1.0, + rope_short_m_scale=1.0, + rope_long_m_scale=1.0, + rope_max_positions=1024, + rope_original_max_positions=1024, + use_paged_context_fmha=use_paged_context_fmha, + attention_input_type=0, # mixed + is_mla_enable=False, + chunked_prefill_buffer_batch_size=1, + q_lora_rank=None, + kv_lora_rank=None, + qk_nope_head_dim=None, + qk_rope_head_dim=None, + v_head_dim=None, + rope_append=None, + mrope_rotary_cos_sin=None, + mrope_position_deltas=None, + helix_position_offsets=None, + helix_is_inactive_rank=None, + attention_chunk_size=None, + softmax_stats_tensor=None, + is_spec_decoding_enabled=False, + use_spec_decoding=False, + is_spec_dec_tree=False, + spec_decoding_generation_lengths=None, + spec_decoding_position_offsets_for_cpp=None, + spec_decoding_packed_mask=None, + spec_decoding_bl_tree_mask_offset=None, + spec_decoding_bl_tree_mask=None, + spec_bl_tree_first_sparse_mask_offset_kv=None, + sparse_kv_indices=None, + sparse_kv_offsets=None, + sparse_attn_indices=None, + sparse_attn_offsets=None, + sparse_attn_indices_block_size=0, + num_contexts=num_contexts, + num_ctx_tokens=num_ctx_tokens, + ) + torch.cuda.synchronize() + + # Record the K/V slices just appended to the cache, per request. + if record: + q_w = self.num_heads * self.head_dim + kv_w = self.num_kv_heads * self.head_dim + start = 0 + for rid, sl in zip(request_ids, seq_lens): + rows = qkv[start : start + sl] + self.k_history[rid].append( + rows[:, q_w : q_w + kv_w].view(sl, self.num_kv_heads, self.head_dim) + ) + self.v_history[rid].append( + rows[:, q_w + kv_w :].view(sl, self.num_kv_heads, self.head_dim) + ) + start += sl + return output + + def reference( + self, + qkv: torch.Tensor, + seq_lens: List[int], + request_ids: List[int], + cached_lens: List[int], + mask_type: int, + q_transform=None, + kv_transform=None, + window: Optional[int] = None, + ) -> torch.Tensor: + """fp32 SDPA over each sequence's full K/V history (must be called + after call_op so the history includes this call's K/V). q_transform / + kv_transform, when given, map the bf16 q / K/V history first (e.g. + the e4m3 round-trip a quantized cache imposes). window, when given, + additionally drops keys older than the sliding window.""" + q_w = self.num_heads * self.head_dim + outs = [] + start = 0 + for rid, sl, cached in zip(request_ids, seq_lens, cached_lens): + q_seq = qkv[start : start + sl, :q_w].view(sl, self.num_heads, self.head_dim) + if q_transform is not None: + q_seq = q_transform(q_seq) + k_seq = torch.cat(self.k_history[rid]) + v_seq = torch.cat(self.v_history[rid]) + if kv_transform is not None: + k_seq = kv_transform(k_seq) + v_seq = kv_transform(v_seq) + skv = k_seq.shape[0] + if mask_type == MASK_CAUSAL: + mask: Optional[torch.Tensor] = torch.zeros( + sl, skv, dtype=torch.bool, device=qkv.device + ) + for i in range(sl): + lo = 0 if window is None else max(0, cached + i + 1 - window) + mask[i, lo : cached + i + 1] = True + else: # padding: every query attends to every key + assert window is None, "sliding window certified for causal only" + mask = None + o = F.scaled_dot_product_attention( + q_seq.transpose(0, 1).float(), + k_seq.transpose(0, 1).float(), + v_seq.transpose(0, 1).float(), + attn_mask=mask, + enable_gqa=True, + ) + outs.append(o.transpose(0, 1).reshape(sl, -1)) + start += sl + return torch.cat(outs).to(qkv.dtype) + + def expected_cache(self, kv: torch.Tensor) -> torch.Tensor: + """Map bf16 K/V fed to the op to the expected pool content: identity + for a bf16 pool; fp32-scale-then-RN-cast for a quantized pool.""" + if self.kv_cache.dtype == torch.bfloat16: + return kv + assert self.kv_scale_orig_quant is not None + return (kv.float() * self.kv_scale_orig_quant).to(self.kv_cache.dtype) + + def cache_pages_content(self, request_id: int) -> Tuple[torch.Tensor, ...]: + """(pool K rows, pool V rows, expected K, expected V), token-major.""" + k_seq = torch.cat(self.k_history[request_id]) # [total, Hkv, D] + v_seq = torch.cat(self.v_history[request_id]) + total = k_seq.shape[0] + tpb = self.tokens_per_block + got_k, got_v = [], [] + for j, p in enumerate(self.pages[request_id]): + n = min(tpb, total - j * tpb) + if n <= 0: + break + got_k.append(self.kv_cache[p, 0, :, :n, :].permute(1, 0, 2)) + got_v.append(self.kv_cache[p, 1, :, :n, :].permute(1, 0, 2)) + return ( + torch.cat(got_k).contiguous(), + torch.cat(got_v).contiguous(), + self.expected_cache(k_seq), + self.expected_cache(v_seq), + ) + + def check_cache(self, request_id: int) -> None: + """The paged cache must hold exactly the expected rows, bitwise (the + bf16 K/V fed so far, or their scaled e4m3 casts for an fp8 pool).""" + got_k, got_v, exp_k, exp_v = self.cache_pages_content(request_id) + assert torch.equal(got_k.view(torch.uint8), exp_k.view(torch.uint8)), ( + f"K cache mismatch: request {request_id}" + ) + assert torch.equal(got_v.view(torch.uint8), exp_v.view(torch.uint8)), ( + f"V cache mismatch: request {request_id}" + ) + + def check_unwritten_pool_zero(self) -> None: + """Pool bytes no append should have written must still hold the + initial zeros: whole pages never allocated to a request, and the + unwritten tail rows of each request's partially filled last page.""" + tpb = self.tokens_per_block + written: dict[int, int] = {} + for rid, pages in self.pages.items(): + total = self.cached_len(rid) + for j, p in enumerate(pages): + written[p] = min(tpb, total - j * tpb) + for p in range(self.kv_cache.shape[0]): + n = written.get(p, 0) + if n >= tpb: + continue + tail = self.kv_cache[p, :, :, n:, :].contiguous() + assert (tail.view(torch.uint8) == 0).all(), f"page {p} written beyond token row {n}" + + def random_qkv(self, num_tokens: int) -> torch.Tensor: + width = (self.num_heads + 2 * self.num_kv_heads) * self.head_dim + return torch.randn(num_tokens, width, dtype=torch.bfloat16, device="cuda") + + +def _run_and_check( + env: _PagedAttnEnv, + seq_lens: List[int], + num_contexts: int, + request_ids: List[int], + mask_type: int = MASK_CAUSAL, +) -> None: + cached_lens = [env.cached_len(rid) for rid in request_ids] + qkv = env.random_qkv(sum(seq_lens)) + out = env.call_op(qkv, seq_lens, num_contexts, request_ids, mask_type) + ref = env.reference(qkv, seq_lens, request_ids, cached_lens, mask_type) + torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) + + +def test_bf16_context_prefill_gqa_d128() -> None: + """Prefill-like: pure-context batch, mixed lengths crossing a block + boundary, GQA 8q/2kv, head_dim 128; cache append checked bit-exactly.""" + torch.manual_seed(0) + env = _PagedAttnEnv(num_heads=8, num_kv_heads=2, head_dim=128) + env.add_request(0, 48) + env.add_request(1, 17) + _run_and_check(env, [48, 17], 2, [0, 1]) + env.check_cache(0) + env.check_cache(1) + + +def test_bf16_decode_reads_cached_kv_gqa_d128() -> None: + """Decode-like: prefill writes the cache, then two gen-only steps read it. + + ctx len 64 exactly fills two 32-token blocks, so the first decode token + lands in a freshly allocated block (block-boundary crossing). + """ + torch.manual_seed(1) + env = _PagedAttnEnv(num_heads=8, num_kv_heads=2, head_dim=128) + env.add_request(0, 64) + env.add_request(1, 17) + _run_and_check(env, [64, 17], 2, [0, 1]) + for _ in range(2): # two decode steps: one token per sequence each + _run_and_check(env, [1, 1], 0, [0, 1]) + env.check_cache(0) + env.check_cache(1) + + +def test_bf16_mixed_batch_gqa_d128() -> None: + """One batch mixing a context-phase sequence and a generation-phase one. + + Context sequences must precede generation sequences in the batch. + """ + torch.manual_seed(2) + env = _PagedAttnEnv(num_heads=8, num_kv_heads=2, head_dim=128) + env.add_request(0, 40) + _run_and_check(env, [40], 1, [0]) + env.add_request(1, 23) + _run_and_check(env, [23, 1], 1, [1, 0]) + env.check_cache(0) + env.check_cache(1) + + +def test_bf16_padding_mask_context() -> None: + """mask_type=padding: bidirectional attention over a context batch.""" + torch.manual_seed(3) + env = _PagedAttnEnv(num_heads=8, num_kv_heads=2, head_dim=128) + env.add_request(0, 31) + env.add_request(1, 9) + _run_and_check(env, [31, 9], 2, [0, 1], mask_type=MASK_PADDING) + + +def test_bf16_mha_head_dim64() -> None: + """MHA (4q/4kv), head_dim 64: prefill then one decode step.""" + torch.manual_seed(4) + env = _PagedAttnEnv(num_heads=4, num_kv_heads=4, head_dim=64) + env.add_request(0, 50) + env.add_request(1, 3) + _run_and_check(env, [50, 3], 2, [0, 1]) + _run_and_check(env, [1, 1], 0, [0, 1]) + env.check_cache(0) + env.check_cache(1) + + +def test_bf16_shipped_target_geometries() -> None: + """The two head geometries the shipped bf16 targets run: 32q/8kv d128 + (GQA 4:1, qwen3-8b) and 32q/4kv d128 (GQA 8:1, qwen3-30b-a3b). + + Each gets the full standard-configuration cycle: a page-crossing prefill + batch, two decode steps, and a mixed context+generation batch, with the + paged append checked bit-exactly and no writes outside the token ranges. + """ + for i, (num_heads, num_kv_heads) in enumerate([(32, 8), (32, 4)]): + torch.manual_seed(30 + i) + env = _PagedAttnEnv(num_heads=num_heads, num_kv_heads=num_kv_heads, head_dim=128) + env.add_request(0, 48) # crosses a 32-token page boundary + env.add_request(1, 17) + _run_and_check(env, [48, 17], 2, [0, 1]) + for _ in range(2): # two decode steps: one token per sequence each + _run_and_check(env, [1, 1], 0, [0, 1]) + env.add_request(2, 23) # context sequence joining a generation one + _run_and_check(env, [23, 1], 1, [2, 0]) + for rid in (0, 1, 2): + env.check_cache(rid) + env.check_unwritten_pool_zero() + + +def test_bf16_head_count_axis() -> None: + """Head counts are a free axis, not an enumerated list: any (Hq, Hkv) + with Hq % Hkv == 0 works, and head_size selects the FMHA kernel. + + Covered here: MQA (16q/1kv), a non-power-of-2 GQA ratio (28q/4kv, + ratio 7), non-power-of-2 head counts (12q/3kv), and head_size 256 — + none of them a power-of-2 grouping the earlier cases already pinned. + """ + geometries = [(16, 1, 128), (28, 4, 128), (12, 3, 128), (8, 2, 256)] + for i, (num_heads, num_kv_heads, head_dim) in enumerate(geometries): + torch.manual_seed(40 + i) + env = _PagedAttnEnv(num_heads=num_heads, num_kv_heads=num_kv_heads, head_dim=head_dim) + env.add_request(0, 40) + env.add_request(1, 7) + _run_and_check(env, [40, 7], 2, [0, 1]) + _run_and_check(env, [1, 1], 0, [0, 1]) + env.check_cache(0) + env.check_cache(1) + + +def test_rejects_non_divisible_head_counts() -> None: + """Hq must be an integer multiple of Hkv, and the wrapper must be the + one to say so: a context-only call at 6q/4kv d128 returns without + raising, having computed only the first (6 // 4) * 4 = 4 head columns + and left the other two all-zero (the decode path does raise, from + xqaDispatcher: 'numQHeads should be multiple of numKVHeads').""" + torch.manual_seed(50) + env = _PagedAttnEnv(num_heads=6, num_kv_heads=4, head_dim=128) + env.add_request(0, 16) + try: + env.call_op(env.random_qkv(16), [16], 1, [0], MASK_CAUSAL) + except AssertionError: + return + raise AssertionError("6q/4kv geometry was not rejected by the wrapper") + + +# ─── Standard configuration, multi-layer shared pool ─────────────────── + + +def _bitwise_equal(a: torch.Tensor, b: torch.Tensor) -> bool: + """Bitwise tensor equality (byte view), dtype-agnostic: float equality + would miss NaN-payload or signed-zero byte changes in an fp8 pool.""" + return torch.equal(a.contiguous().view(torch.uint8), b.contiguous().view(torch.uint8)) + + +class _MultiLayerPagedAttnEnv: + """Real multi-layer op state: a trtllm KVCacheManager hosting one paged + pool shared by several layers, whose pool pointers, layer->pool mapping, + and block offsets the op consumes exactly as produced. The pool element + type follows the manager dtype (bf16, or fp8-e4m3 with the same quant + args as the single-layer fp8 env). + + The manager owns page allocation; the test mirrors every K/V fed per + (layer, request) so references and bitwise cache checks come from the + true per-layer history. Page ids are re-derived from the manager's + offsets, pinning the multi-layer stride: page p of an L-layer pool has + K-slab offset p * L * 2 and V-slab offset p * L * 2 + 1. + """ + + def __init__( + self, + num_layers: int, + num_heads: int, + num_kv_heads: int, + head_dim: int, + tokens_per_block: int = 32, + num_blocks: int = 64, + max_batch: int = 4, + max_seq_len: int = 256, + dtype: DataType = DataType.BF16, + quant_mode: int = 0, + kv_scaling_factor: Optional[float] = None, + ) -> None: + self.num_layers = num_layers + self.num_heads = num_heads + self.num_kv_heads = num_kv_heads + self.head_dim = head_dim + self.tokens_per_block = tokens_per_block + self.max_batch = max_batch + self.max_seq_len = max_seq_len + self.quant_mode = quant_mode + + self.mgr = KVCacheManager( + KvCacheConfig( + max_tokens=num_blocks * tokens_per_block, + enable_block_reuse=False, + ), + CacheTypeCpp.SELF, + num_layers=num_layers, + num_kv_heads=num_kv_heads, + head_dim=head_dim, + tokens_per_block=tokens_per_block, + max_seq_len=max_seq_len, + max_batch_size=max_batch, + mapping=Mapping(world_size=1, tp_size=1, rank=0), + dtype=dtype, + ) + # KV-cache scaling factor s, same construction as the single-layer + # env: quant_orig = s, orig_quant = 1/s, both fp32 [1] CUDA; None = + # unquantized pool, scale args passed as None. + if kv_scaling_factor is None: + self.kv_scale_quant_orig: Optional[torch.Tensor] = None + self.kv_scale_orig_quant: Optional[torch.Tensor] = None + else: + factor = torch.full((1,), kv_scaling_factor, dtype=torch.float32, device="cuda") + self.kv_scale_quant_orig = factor + self.kv_scale_orig_quant = 1.0 / factor + assert self.mgr.num_pools == 1, "test expects one pool holding all layers" + pool_mapping = self.mgr.kv_cache_pool_mapping + assert pool_mapping is not None + self.pool_mapping: torch.Tensor = pool_mapping + self.block_offsets = torch.zeros( + 1, + max_batch, + 2, + self.mgr.max_blocks_per_seq, + dtype=torch.int32, + device="cuda", + ) + # Per-layer HND views into the shared pool, each + # [num_blocks, 2, Hkv, tokens_per_block, D] strided over whole pages. + self.layer_views: List[torch.Tensor] = [] + for idx in range(num_layers): + view = self.mgr.get_buffers(idx, kv_layout="HND") + assert view is not None + view.zero_() + self.layer_views.append(view) + # Auto-resized in place by the op on first call. + self.workspace = torch.empty(0, dtype=torch.int8, device="cuda") + + self.prompt_lens: Dict[int, int] = {} + self.k_history: Dict[Tuple[int, int], List[torch.Tensor]] = {} + self.v_history: Dict[Tuple[int, int], List[torch.Tensor]] = {} + + def add_request(self, request_id: int, prompt_len: int) -> None: + """Register the request with the manager, allocating its prompt pages.""" + self.prompt_lens[request_id] = prompt_len + added = self.mgr.add_dummy_requests([request_id], token_nums=[prompt_len]) + assert added is not None, "KV cache manager out of blocks" + for layer_idx in range(self.num_layers): + self.k_history[layer_idx, request_id] = [] + self.v_history[layer_idx, request_id] = [] + + def add_decode_token(self, request_id: int) -> None: + """Extend the manager's allocation by one generated token.""" + self.mgr.impl.add_token(request_id) + + def refresh_offsets(self, request_ids: List[int], num_contexts: int) -> None: + """Copy the manager-produced block offsets for this batch to device.""" + self.mgr.copy_batch_block_offsets( + self.block_offsets, request_ids, 1, num_contexts, len(request_ids) + ) + torch.cuda.synchronize() # the copy is staged non-blocking + + def cached_len(self, layer_idx: int, request_id: int) -> int: + return sum(t.shape[0] for t in self.k_history[layer_idx, request_id]) + + def call_op( + self, + layer_idx: int, + qkv: torch.Tensor, + seq_lens: List[int], + num_contexts: int, + request_ids: List[int], + pool_mapping: Optional[torch.Tensor] = None, + record: bool = True, + attention_window_size: Optional[int] = None, + use_paged_context_fmha: bool = False, + ctx_lens: Optional[List[int]] = None, + ) -> torch.Tensor: + """One thop_attention call for one layer of the shared pool. + + Offsets must have been refreshed for exactly this request order. + pool_mapping overrides the manager's mapping (doctored-mapping + probe); record=False skips the K/V history mirror for such calls. + attention_window_size defaults to max_seq_len (no sliding window) and + is a per-call scalar, so layers sharing the pool may differ in it. + use_paged_context_fmha=True selects the paged-context execution path; + ctx_lens then carries this call's per-context-row new-token count, + which a cached prefix makes differ from the registered prompt length. + """ + ns = len(request_ids) + kv_lens = [self.cached_len(layer_idx, rid) + new for rid, new in zip(request_ids, seq_lens)] + if ctx_lens is None: + ctx_lens = [self.prompt_lens[rid] for rid in request_ids] + req_types = [0 if i < num_contexts else 1 for i in range(ns)] + num_ctx_tokens = sum(seq_lens[:num_contexts]) + if pool_mapping is None: + pool_mapping = self.pool_mapping + + output = torch.empty( + sum(seq_lens), + self.num_heads * self.head_dim, + dtype=qkv.dtype, + device=qkv.device, + ) + thop_attention( + q=qkv, + k=None, # packed QKV rides inside q + v=None, + output=output, + output_sf=None, + workspace_=self.workspace, + sequence_length=torch.tensor(kv_lens, dtype=torch.int32, device="cuda"), + host_past_key_value_lengths=torch.tensor(kv_lens, dtype=torch.int32), + host_total_kv_lens=torch.tensor( + [sum(kv_lens[:num_contexts]), sum(kv_lens[num_contexts:])], + dtype=torch.int32, + ), + context_lengths=torch.tensor(ctx_lens, dtype=torch.int32, device="cuda"), + host_context_lengths=torch.tensor(ctx_lens, dtype=torch.int32), + host_request_types=torch.tensor(req_types, dtype=torch.int32), + max_context_q_len_override=None, + kv_cache_block_offsets=self.block_offsets, + host_kv_cache_pool_pointers=self.mgr.kv_cache_pool_pointers, + host_kv_cache_pool_mapping=pool_mapping, + cache_indirection=None, + kv_scale_orig_quant=self.kv_scale_orig_quant, + kv_scale_quant_orig=self.kv_scale_quant_orig, + out_scale=None, + rotary_inv_freq=None, + rotary_cos_sin=None, + latent_cache=None, + q_pe=None, + block_ids_per_seq=None, + attention_sinks=None, + is_fused_qkv=True, + update_kv_cache=True, + predicted_tokens_per_seq=1, + local_layer_idx=layer_idx, + num_heads=self.num_heads, + num_kv_heads=self.num_kv_heads, + head_size=self.head_dim, + tokens_per_block=self.tokens_per_block, + max_num_requests=self.max_batch, + max_context_length=self.max_seq_len, + max_seq_len=self.max_seq_len, + attention_window_size=( + self.max_seq_len if attention_window_size is None else attention_window_size + ), + beam_width=1, + mask_type=MASK_CAUSAL, + quant_mode=self.quant_mode, + q_scaling=1.0, + position_embedding_type=0, # no in-kernel RoPE + rope_dim=0, + rope_base=10000.0, + rope_scale_type=0, + rope_scale=1.0, + rope_short_m_scale=1.0, + rope_long_m_scale=1.0, + rope_max_positions=1024, + rope_original_max_positions=1024, + use_paged_context_fmha=use_paged_context_fmha, + attention_input_type=0, # mixed + is_mla_enable=False, + chunked_prefill_buffer_batch_size=1, + q_lora_rank=None, + kv_lora_rank=None, + qk_nope_head_dim=None, + qk_rope_head_dim=None, + v_head_dim=None, + rope_append=None, + mrope_rotary_cos_sin=None, + mrope_position_deltas=None, + helix_position_offsets=None, + helix_is_inactive_rank=None, + attention_chunk_size=None, + softmax_stats_tensor=None, + is_spec_decoding_enabled=False, + use_spec_decoding=False, + is_spec_dec_tree=False, + spec_decoding_generation_lengths=None, + spec_decoding_position_offsets_for_cpp=None, + spec_decoding_packed_mask=None, + spec_decoding_bl_tree_mask_offset=None, + spec_decoding_bl_tree_mask=None, + spec_bl_tree_first_sparse_mask_offset_kv=None, + sparse_kv_indices=None, + sparse_kv_offsets=None, + sparse_attn_indices=None, + sparse_attn_offsets=None, + sparse_attn_indices_block_size=0, + num_contexts=num_contexts, + num_ctx_tokens=num_ctx_tokens, + ) + torch.cuda.synchronize() + + if record: + q_w = self.num_heads * self.head_dim + kv_w = self.num_kv_heads * self.head_dim + start = 0 + for rid, sl in zip(request_ids, seq_lens): + rows = qkv[start : start + sl] + self.k_history[layer_idx, rid].append( + rows[:, q_w : q_w + kv_w].view(sl, self.num_kv_heads, self.head_dim) + ) + self.v_history[layer_idx, rid].append( + rows[:, q_w + kv_w :].view(sl, self.num_kv_heads, self.head_dim) + ) + start += sl + return output + + def reference( + self, + layer_idx: int, + qkv: torch.Tensor, + seq_lens: List[int], + request_ids: List[int], + cached_lens: List[int], + q_transform=None, + kv_transform=None, + window: Optional[int] = None, + ) -> torch.Tensor: + """fp32 causal SDPA over this layer's full K/V history (must be + called after call_op so the history includes this call's K/V). + q_transform / kv_transform, when given, map the bf16 q / K/V history + first (e.g. the e4m3 round-trip a quantized cache imposes). window, + when given, additionally drops keys older than the sliding window.""" + q_w = self.num_heads * self.head_dim + outs = [] + start = 0 + for rid, sl, cached in zip(request_ids, seq_lens, cached_lens): + q_seq = qkv[start : start + sl, :q_w].view(sl, self.num_heads, self.head_dim) + if q_transform is not None: + q_seq = q_transform(q_seq) + k_seq = torch.cat(self.k_history[layer_idx, rid]) + v_seq = torch.cat(self.v_history[layer_idx, rid]) + if kv_transform is not None: + k_seq = kv_transform(k_seq) + v_seq = kv_transform(v_seq) + mask = torch.zeros(sl, k_seq.shape[0], dtype=torch.bool, device=qkv.device) + for i in range(sl): + lo = 0 if window is None else max(0, cached + i + 1 - window) + mask[i, lo : cached + i + 1] = True + o = F.scaled_dot_product_attention( + q_seq.transpose(0, 1).float(), + k_seq.transpose(0, 1).float(), + v_seq.transpose(0, 1).float(), + attn_mask=mask, + enable_gqa=True, + ) + outs.append(o.transpose(0, 1).reshape(sl, -1)) + start += sl + return torch.cat(outs).to(qkv.dtype) + + def expected_cache(self, kv: torch.Tensor) -> torch.Tensor: + """Map bf16 K/V fed to the op to the expected pool content: identity + for a bf16 pool; fp32-scale-then-RN-cast for a quantized pool.""" + pool_dtype = self.layer_views[0].dtype + if pool_dtype == torch.bfloat16: + return kv + assert self.kv_scale_orig_quant is not None + return (kv.float() * self.kv_scale_orig_quant).to(pool_dtype) + + def check_caches(self, request_ids: List[int]) -> None: + """Every layer's slabs must hold exactly the expected rows, bitwise + (the bf16 K/V fed to that layer, or their scaled e4m3 casts for an + fp8 pool), at the pages the manager assigned (offsets are + layer-agnostic and count slabs: one page = num_layers * 2 slabs).""" + tpb = self.tokens_per_block + slabs_per_page = self.num_layers * 2 + for layer_idx in range(self.num_layers): + view = self.layer_views[layer_idx] + for s, rid in enumerate(request_ids): + k_seq = self.expected_cache(torch.cat(self.k_history[layer_idx, rid])) + v_seq = self.expected_cache(torch.cat(self.v_history[layer_idx, rid])) + total = k_seq.shape[0] + for j in range((total + tpb - 1) // tpb): + k_off = int(self.block_offsets[0, s, 0, j].item()) + v_off = int(self.block_offsets[0, s, 1, j].item()) + assert k_off % slabs_per_page == 0, "multi-layer K offset stride" + assert v_off == k_off + 1, "V slab follows K within the page" + page = k_off // slabs_per_page + n = min(tpb, total - j * tpb) + k_blk = view[page, 0, :, :n, :].permute(1, 0, 2) + v_blk = view[page, 1, :, :n, :].permute(1, 0, 2) + assert _bitwise_equal(k_blk, k_seq[j * tpb : j * tpb + n]), ( + f"K cache mismatch: layer {layer_idx}, request {rid}, page {j}" + ) + assert _bitwise_equal(v_blk, v_seq[j * tpb : j * tpb + n]), ( + f"V cache mismatch: layer {layer_idx}, request {rid}, page {j}" + ) + + def random_qkv(self, num_tokens: int) -> torch.Tensor: + width = (self.num_heads + 2 * self.num_kv_heads) * self.head_dim + return torch.randn(num_tokens, width, dtype=torch.bfloat16, device="cuda") + + +def test_bf16_multilayer_shared_pool_gqa_d128() -> None: + """One paged pool shared by 4 layers, addressed as a real multi-layer + KVCacheManager lays it out: pages interleave every layer's K/V slabs, + the block-offset table is layer-agnostic, and each call selects its + layer via local_layer_idx -> pool-mapping row. Prefill plus two decode + steps per layer (decode crosses a page boundary), GQA 8q/2kv d128: + per-layer outputs vs fp32 references, per-layer appends bit-exact, + sibling layers bitwise untouched by each call. A final doctored-mapping + call pins the pool-base shift to the mapping row's layer-in-pool column + (local_layer_idx only selects the row).""" + torch.manual_seed(10) + num_layers = 4 + env = _MultiLayerPagedAttnEnv(num_layers=num_layers, num_heads=8, num_kv_heads=2, head_dim=128) + # The real manager maps layer l to (pool 0, layer-in-pool l): one + # mapping row per layer, rows beyond 0 shifting the pool base. + assert env.pool_mapping.tolist() == [[0, layer] for layer in range(num_layers)] + + env.add_request(0, 64) # exactly two pages: decode crosses into a third + env.add_request(1, 17) + env.refresh_offsets([0, 1], num_contexts=2) + for layer in range(num_layers): + qkv = env.random_qkv(81) + siblings_before = [env.layer_views[m].clone() for m in range(num_layers) if m != layer] + out = env.call_op(layer, qkv, [64, 17], 2, [0, 1]) + ref = env.reference(layer, qkv, [64, 17], [0, 1], [0, 0]) + torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) + siblings_after = [env.layer_views[m] for m in range(num_layers) if m != layer] + for before, after in zip(siblings_before, siblings_after): + assert torch.equal(before, after) # no cross-layer write + env.check_caches([0, 1]) + + for _ in range(2): # decode: one token per sequence per layer per step + env.add_decode_token(0) + env.add_decode_token(1) + env.refresh_offsets([0, 1], num_contexts=0) + for layer in range(num_layers): + cached = [env.cached_len(layer, 0), env.cached_len(layer, 1)] + qkv = env.random_qkv(2) + out = env.call_op(layer, qkv, [1, 1], 0, [0, 1]) + ref = env.reference(layer, qkv, [1, 1], [0, 1], cached) + torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) + env.check_caches([0, 1]) + + # Doctored mapping: local_layer_idx=1 whose row claims layer-in-pool 3. + # The append must land in layer 3's slabs — the shift is driven by the + # mapping row's layer column, not by local_layer_idx itself. (Appends + # one token past each sequence's history, inside already-allocated + # pages; run after all correctness checks since it plants garbage.) + doctored = env.pool_mapping.clone() + doctored[1, 1] = 3 + views_before = [env.layer_views[m].clone() for m in range(num_layers)] + env.call_op(1, env.random_qkv(2), [1, 1], 0, [0, 1], pool_mapping=doctored, record=False) + assert torch.equal(env.layer_views[1], views_before[1]) # row idx not the shift + assert not torch.equal(env.layer_views[3], views_before[3]) + assert torch.equal(env.layer_views[0], views_before[0]) + assert torch.equal(env.layer_views[2], views_before[2]) + + +# ─── Standard configuration, fp8-e4m3 paged KV pool (quant_mode 128) ─── + + +def _fp8_env(kv_scaling_factor: float = 1.0) -> _PagedAttnEnv: + """GQA 32q/8kv d128 tpb 32 (qwen3-8b-like geometry), fp8-e4m3 pool.""" + return _PagedAttnEnv( + num_heads=32, + num_kv_heads=8, + head_dim=128, + pool_dtype=torch.float8_e4m3fn, + quant_mode=QUANT_MODE_FP8_KV_CACHE, + kv_scaling_factor=kv_scaling_factor, + ) + + +def _e4m3_roundtrip(env: "_PagedAttnEnv | _MultiLayerPagedAttnEnv"): + """bf16 -> e4m3 (scaled by orig_quant) -> bf16 (scaled by quant_orig): + what a value fed into the fp8 pool looks like when read back out.""" + orig_quant, quant_orig = env.kv_scale_orig_quant, env.kv_scale_quant_orig + assert orig_quant is not None and quant_orig is not None + + def transform(t: torch.Tensor) -> torch.Tensor: + q = (t.float() * orig_quant).to(torch.float8_e4m3fn) + return (q.float() * quant_orig).to(t.dtype) + + return transform + + +def test_fp8_kv_context_prefill_gqa32_8_d128() -> None: + """fp8 pool, context prefill: the context FMHA computes over the bf16 + packed QKV — accuracy identical to the bf16-pool surface (an fp8-KV + reference is ~1.4e-1 off; the pool plays no part in context math) — + while the in-op append writes e4m3(K * orig_quant) into the pool, + bit-exactly, touching nothing else.""" + torch.manual_seed(20) + env = _fp8_env() + env.add_request(0, 48) # crosses a 32-token page boundary + env.add_request(1, 17) + qkv = env.random_qkv(65) + out = env.call_op(qkv, [48, 17], 2, [0, 1], MASK_CAUSAL) + ref = env.reference(qkv, [48, 17], [0, 1], [0, 0], MASK_CAUSAL) + torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) + env.check_cache(0) + env.check_cache(1) + env.check_unwritten_pool_zero() + + +def test_fp8_kv_decode_reads_fp8_cache() -> None: + """Decode over the fp8 pool: prefill [64, 17] (64 fills exactly two + pages, so the first decode token opens a fresh page), then two decode + steps checked against the quantized-KV reference; appends stay bit-exact + through decode. A 200-token-history decode covers prefill-like KV + extents (observed err there is ~4x smaller than the short-history max).""" + torch.manual_seed(21) + env = _fp8_env() + env.add_request(0, 64) + env.add_request(1, 17) + env.call_op(env.random_qkv(81), [64, 17], 2, [0, 1], MASK_CAUSAL) + rt = _e4m3_roundtrip(env) + for _ in range(2): + cached = [env.cached_len(0), env.cached_len(1)] + qkv_d = env.random_qkv(2) + out = env.call_op(qkv_d, [1, 1], 0, [0, 1], MASK_CAUSAL) + ref = env.reference( + qkv_d, [1, 1], [0, 1], cached, MASK_CAUSAL, q_transform=rt, kv_transform=rt + ) + torch.testing.assert_close(out, ref, rtol=FP8_DECODE_RTOL, atol=FP8_DECODE_ATOL) + env.check_cache(0) + env.check_cache(1) + env.check_unwritten_pool_zero() + + torch.manual_seed(22) + env = _fp8_env() + env.add_request(0, 200) + env.call_op(env.random_qkv(200), [200], 1, [0], MASK_CAUSAL) + rt = _e4m3_roundtrip(env) + qkv_d = env.random_qkv(1) + out = env.call_op(qkv_d, [1], 0, [0], MASK_CAUSAL) + ref = env.reference(qkv_d, [1], [0], [200], MASK_CAUSAL, q_transform=rt, kv_transform=rt) + torch.testing.assert_close(out, ref, rtol=FP8_DECODE_RTOL, atol=FP8_DECODE_ATOL) + + +def test_fp8_kv_mixed_batch() -> None: + """One call mixing a context and a generation sequence over the fp8 + pool: context rows match the bf16-KV reference (bf16 context FMHA), the + generation row matches the quantized-KV reference (fp8 decode kernel).""" + torch.manual_seed(23) + env = _fp8_env() + env.add_request(0, 40) + env.call_op(env.random_qkv(40), [40], 1, [0], MASK_CAUSAL) + env.add_request(1, 23) + cached_gen = env.cached_len(0) + qkv = env.random_qkv(24) + out = env.call_op(qkv, [23, 1], 1, [1, 0], MASK_CAUSAL) + ref_ctx = env.reference(qkv[:23], [23], [1], [0], MASK_CAUSAL) + torch.testing.assert_close(out[:23], ref_ctx, rtol=RTOL, atol=ATOL) + rt = _e4m3_roundtrip(env) + ref_gen = env.reference( + qkv[23:], [1], [0], [cached_gen], MASK_CAUSAL, q_transform=rt, kv_transform=rt + ) + torch.testing.assert_close(out[23:], ref_gen, rtol=FP8_DECODE_RTOL, atol=FP8_DECODE_ATOL) + env.check_cache(0) + env.check_cache(1) + + +def test_fp8_kv_scale_semantics() -> None: + """Non-1.0 kv scales. s=2.0 (a power of two: scaling is an exact + exponent shift): the append stays bit-exact vs the e4m3(K * orig_quant) + mirror — orig_quant is consumed on write — and decode matches the + scale-aware reference — quant_orig is consumed on read (ignoring it + shows as ~4.4e-1). s=1.5 (not a power of two, so the e4m3 round-trip + genuinely depends on the scale value): the kernel's quantization + arithmetic differs from fp32-multiply-then-round-to-nearest on a small + fraction of elements (observed 0.6%), each within one e4m3 ulp; decode + matches within the same fp8 tolerance.""" + torch.manual_seed(24) + env = _fp8_env(kv_scaling_factor=2.0) + env.add_request(0, 40) + qkv = env.random_qkv(40) + out = env.call_op(qkv, [40], 1, [0], MASK_CAUSAL) + ref = env.reference(qkv, [40], [0], [0], MASK_CAUSAL) + torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) # bf16 context + env.check_cache(0) + rt = _e4m3_roundtrip(env) + qkv_d = env.random_qkv(1) + out = env.call_op(qkv_d, [1], 0, [0], MASK_CAUSAL) + ref = env.reference(qkv_d, [1], [0], [40], MASK_CAUSAL, q_transform=rt, kv_transform=rt) + torch.testing.assert_close(out, ref, rtol=FP8_DECODE_RTOL, atol=FP8_DECODE_ATOL) + + torch.manual_seed(25) + env = _fp8_env(kv_scaling_factor=1.5) + env.add_request(0, 64) + env.call_op(env.random_qkv(64), [64], 1, [0], MASK_CAUSAL) + got_k, got_v, exp_k, exp_v = env.cache_pages_content(0) + for got, exp in ((got_k, exp_k), (got_v, exp_v)): + exact = (got.view(torch.uint8) == exp.view(torch.uint8)).float().mean().item() + assert exact >= 0.99, f"append bit-exact fraction only {exact:.4f}" + diff = (got.float() - exp.float()).abs() + # One e4m3 ulp at the element's magnitude: 2^(floor(log2 |x|) - 3) + # for normals, 2^-9 below the min normal 2^-6. + mag = torch.maximum(got.float().abs(), exp.float().abs()).clamp(min=2**-6) + ulp = 2.0 ** (torch.floor(torch.log2(mag)) - 3) + assert bool((diff <= ulp).all()), "append off by more than one e4m3 ulp" + rt = _e4m3_roundtrip(env) + qkv_d = env.random_qkv(1) + out = env.call_op(qkv_d, [1], 0, [0], MASK_CAUSAL) + ref = env.reference(qkv_d, [1], [0], [64], MASK_CAUSAL, q_transform=rt, kv_transform=rt) + torch.testing.assert_close(out, ref, rtol=FP8_DECODE_RTOL, atol=FP8_DECODE_ATOL) + + +def test_fp8_kv_multilayer_shared_pool_gqa32_8_d128() -> None: + """One fp8-e4m3 paged pool shared by 4 layers (quant_mode 128, s=1.0), + manager state consumed as-is — the production serving shape of the fp8 + cache, at the production GQA 32q/8kv d128 tpb 32 geometry. The op sizes + slabs from quant_mode alone, so the layer-base shift must be computed in + e4m3 slab units for the append to land where the fp8 manager laid the + layer out. Prefill plus two decode steps per layer (decode crosses a + page boundary): per-layer context outputs vs bf16 references (context + FMHA reads the packed bf16 q rows, not the pool), per-layer decode + outputs vs quantized-KV references over that layer's own history, + per-layer e4m3 appends bit-exact, and sibling layers bitwise untouched + by every call — context and generation append paths both.""" + torch.manual_seed(26) + num_layers = 4 + env = _MultiLayerPagedAttnEnv( + num_layers=num_layers, + num_heads=32, + num_kv_heads=8, + head_dim=128, + dtype=DataType.FP8, + quant_mode=QUANT_MODE_FP8_KV_CACHE, + kv_scaling_factor=1.0, + ) + # A DataType.FP8 manager allocates a real e4m3 pool with identity + # mapping rows, exactly like the bf16 one. + assert env.layer_views[0].dtype == torch.float8_e4m3fn + assert env.pool_mapping.tolist() == [[0, layer] for layer in range(num_layers)] + + def check_sibling_isolation(layer: int, fn) -> None: + before = [env.layer_views[m].clone() for m in range(num_layers) if m != layer] + fn() + after = [env.layer_views[m] for m in range(num_layers) if m != layer] + for b, a in zip(before, after): + assert _bitwise_equal(b, a), f"call on layer {layer} wrote a sibling" + + env.add_request(0, 64) # exactly two pages: decode crosses into a third + env.add_request(1, 17) + env.refresh_offsets([0, 1], num_contexts=2) + for layer in range(num_layers): + qkv = env.random_qkv(81) + + def ctx_call(layer=layer, qkv=qkv): + out = env.call_op(layer, qkv, [64, 17], 2, [0, 1]) + ref = env.reference(layer, qkv, [64, 17], [0, 1], [0, 0]) + torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) + + check_sibling_isolation(layer, ctx_call) + env.check_caches([0, 1]) + + rt = _e4m3_roundtrip(env) + for _ in range(2): # decode: one token per sequence per layer per step + env.add_decode_token(0) + env.add_decode_token(1) + env.refresh_offsets([0, 1], num_contexts=0) + for layer in range(num_layers): + cached = [env.cached_len(layer, 0), env.cached_len(layer, 1)] + qkv = env.random_qkv(2) + + def gen_call(layer=layer, qkv=qkv, cached=cached): + out = env.call_op(layer, qkv, [1, 1], 0, [0, 1]) + ref = env.reference( + layer, + qkv, + [1, 1], + [0, 1], + cached, + q_transform=rt, + kv_transform=rt, + ) + torch.testing.assert_close(out, ref, rtol=FP8_DECODE_RTOL, atol=FP8_DECODE_ATOL) + + check_sibling_isolation(layer, gen_call) + env.check_caches([0, 1]) + + +# ─── Attention sinks (standard configuration, bf16 pool) ─────────────── +# gpt-oss-120b tp1 attention geometry: 64 q heads, 8 kv heads, head_size 64. +SINK_HQ, SINK_HKV, SINK_D = 64, 8, 64 + + +def _sink_env(**kwargs) -> _PagedAttnEnv: + return _PagedAttnEnv(num_heads=SINK_HQ, num_kv_heads=SINK_HKV, head_dim=SINK_D, **kwargs) + + +def _sinks(low: float, high: float, seed: int) -> torch.Tensor: + """One fp32 sink logit per q head, spread over [low, high] so every head + carries a different value (a kernel indexing sinks by kv head, or by a + shifted head index, cannot match the reference).""" + gen = torch.Generator(device="cuda").manual_seed(seed) + return torch.empty(SINK_HQ, dtype=torch.float32, device="cuda").uniform_( + low, high, generator=gen + ) + + +def _sink_reference( + env: _PagedAttnEnv, + qkv: torch.Tensor, + seq_lens: List[int], + request_ids: List[int], + cached_lens: List[int], + sink: Optional[torch.Tensor], + prescale: bool = False, + window: Optional[int] = None, +) -> torch.Tensor: + """fp32 causal attention with one extra per-q-head logit in the softmax + denominator that is dropped from the numerator: + + out[i, h] = softmax([scores[i, h, :], sink[h]])[:-1] @ V + + scores are the scaled logits (QK^T / sqrt(D)); the rows of the resulting + weight matrix therefore sum to less than 1. sink=None gives plain softmax. + prescale=True is the rival hypothesis — the sink joins the *unscaled* + score row and so gets multiplied by the softmax scale too. window, when + given, keeps only the newest `window` keys per query row. + + Built from each sequence's full K/V history, so it must be called after + call_op (same convention as _PagedAttnEnv.reference). + """ + hq, hkv, dim = env.num_heads, env.num_kv_heads, env.head_dim + scale = 1.0 / math.sqrt(dim) + rep = hq // hkv + q_w = hq * dim + outs = [] + start = 0 + for rid, sl, cached in zip(request_ids, seq_lens, cached_lens): + q = qkv[start : start + sl, :q_w].view(sl, hq, dim).float() + k = torch.cat(env.k_history[rid]).float().repeat_interleave(rep, dim=1) + v = torch.cat(env.v_history[rid]).float().repeat_interleave(rep, dim=1) + s = torch.einsum("ihd,jhd->hij", q, k) * scale # [hq, sl, kv] + keep = torch.zeros(sl, k.shape[0], dtype=torch.bool, device=q.device) + for i in range(sl): + lo = 0 if window is None else max(0, cached + i + 1 - window) + keep[i, lo : cached + i + 1] = True + s = s.masked_fill(~keep.unsqueeze(0), float("-inf")) + if sink is None: + p = torch.softmax(s, dim=-1) + else: + sk = sink.float().view(hq, 1) * (scale if prescale else 1.0) + m = torch.maximum(s.max(dim=-1).values, sk) # [hq, sl] + e = torch.exp(s - m.unsqueeze(-1)) + p = e / (e.sum(-1) + torch.exp(sk - m)).unsqueeze(-1) + outs.append(torch.einsum("hij,jhd->ihd", p, v).reshape(sl, -1)) + start += sl + return torch.cat(outs).to(qkv.dtype) + + +def _assert_sink_effect(out: torch.Tensor, rival: torch.Tensor, label: str) -> None: + """The observed output must sit far outside the tolerance band around a + rival hypothesis. Without this, matching the sink reference would prove + nothing: a kernel that silently ignored the sink argument would pass the + positive comparison too whenever the sink's contribution is small.""" + gap = (out.float() - rival.float()).abs().max().item() + allowance = ATOL + RTOL * rival.float().abs().max().item() + assert gap > 5.0 * allowance, ( + f"{label}: max|out - rival| = {gap:.4g} is not 5x outside the " + f"{allowance:.4g} tolerance band — the sink's effect is unresolvable here" + ) + + +def _check_sink_case( + env: _PagedAttnEnv, + qkv: torch.Tensor, + out: torch.Tensor, + seq_lens: List[int], + request_ids: List[int], + cached_lens: List[int], + sink: torch.Tensor, + label: str, + window: Optional[int] = None, +) -> None: + """Positive gate against the sink reference plus the separations that make + it meaningful: the sink is honoured at all, it is a logit in the + *scaled*-score domain rather than a pre-scaling one, and — when a sliding + window is active — the window is honoured too.""" + args = (env, qkv, seq_lens, request_ids, cached_lens) + torch.testing.assert_close( + out, _sink_reference(*args, sink, window=window), rtol=RTOL, atol=ATOL + ) + _assert_sink_effect(out, _sink_reference(*args, None, window=window), f"{label}: sink ignored") + _assert_sink_effect( + out, + _sink_reference(*args, sink, prescale=True, window=window), + f"{label}: sink pre-scaled", + ) + if window is not None: + _assert_sink_effect(out, _sink_reference(*args, sink), f"{label}: window ignored") + + +def test_bf16_sinks_gpt_oss_geometry() -> None: + """attention_sinks over the gpt-oss-120b tp1 geometry (64q/8kv d64, causal, + packed QKV, bf16 pool): prefill, two decode steps, and a mixed + context+generation batch — the three shapes attention_input_type=0 serves. + + Context and generation take different FMHA kernels, so each phase is gated + separately; a sink honoured in one and dropped in the other would be + silent. The cache append must be unaffected (bit-exact, no writes outside + the token ranges). + """ + torch.manual_seed(60) + env = _sink_env() + sink = _sinks(-2.0, 4.0, seed=60) + + env.add_request(0, 48) # crosses a 32-token page boundary + env.add_request(1, 17) + qkv = env.random_qkv(65) + out = env.call_op(qkv, [48, 17], 2, [0, 1], MASK_CAUSAL, attention_sinks=sink) + # The hand-written softmax must agree with the file's SDPA-based reference + # when the sink is absent, so the sink cases test the sink and not the math. + torch.testing.assert_close( + _sink_reference(env, qkv, [48, 17], [0, 1], [0, 0], None), + env.reference(qkv, [48, 17], [0, 1], [0, 0], MASK_CAUSAL), + rtol=RTOL, + atol=ATOL, + ) + _check_sink_case(env, qkv, out, [48, 17], [0, 1], [0, 0], sink, "context") + + for step in range(2): # generation phase: one token per sequence per step + cached = [env.cached_len(0), env.cached_len(1)] + qkv_d = env.random_qkv(2) + out_d = env.call_op(qkv_d, [1, 1], 0, [0, 1], MASK_CAUSAL, attention_sinks=sink) + _check_sink_case(env, qkv_d, out_d, [1, 1], [0, 1], cached, sink, f"decode step {step}") + + env.add_request(2, 23) # context sequence sharing a call with a decode one + cached_gen = env.cached_len(0) + qkv_m = env.random_qkv(24) + out_m = env.call_op(qkv_m, [23, 1], 1, [2, 0], MASK_CAUSAL, attention_sinks=sink) + _check_sink_case(env, qkv_m, out_m, [23, 1], [2, 0], [0, cached_gen], sink, "mixed batch") + + for rid in (0, 1, 2): + env.check_cache(rid) + env.check_unwritten_pool_zero() + + +def test_bf16_sinks_denominator_only() -> None: + """The sink enters the denominator and never the numerator. + + With every V row set to 1, the output of head h at token i is exactly the + softmax row mass, so the sink's effect is directly readable: + + mass[i, h] = sum_j exp(s_ij - m) / (sum_j exp(s_ij - m) + exp(sink_h - m)) + = sigmoid(logsumexp_j(s_ij) - sink_h) + + with s the *scaled* logits. Without a sink the mass is exactly 1; with one + it drops below 1 by an amount that depends on the sink alone. The rival + pre-scaling hypothesis predicts sigmoid(lse - sink/sqrt(D)) and is gated out. + """ + torch.manual_seed(61) + env = _sink_env() + env.add_request(0, 37) + n = 37 + qkv = env.random_qkv(n) + qkv[:, (SINK_HQ + SINK_HKV) * SINK_D :] = 1.0 # V slice: all ones + sink = _sinks(-1.0, 3.0, seed=61) + + # Both runs are the same (idempotent) call over an empty-cache sequence, so + # neither records: the second must see the first's exact batch state. + out = env.call_op(qkv, [n], 1, [0], MASK_CAUSAL, attention_sinks=sink, record=False) + out_ns = env.call_op(qkv, [n], 1, [0], MASK_CAUSAL, record=False) + + scale = 1.0 / math.sqrt(SINK_D) + q_w, kv_w = SINK_HQ * SINK_D, SINK_HKV * SINK_D + q = qkv[:, :q_w].view(n, SINK_HQ, SINK_D).float() + k = ( + qkv[:, q_w : q_w + kv_w] + .view(n, SINK_HKV, SINK_D) + .float() + .repeat_interleave(SINK_HQ // SINK_HKV, dim=1) + ) + s = torch.einsum("ihd,jhd->hij", q, k) * scale + s = s.masked_fill( + ~torch.ones(n, n, dtype=torch.bool, device="cuda").tril().unsqueeze(0), + float("-inf"), + ) + lse = torch.logsumexp(s, dim=-1) # [HQ, n] + + def mass_to_out(mass: torch.Tensor) -> torch.Tensor: + return ( + mass.transpose(0, 1) + .unsqueeze(-1) + .expand(n, SINK_HQ, SINK_D) + .reshape(n, -1) + .to(torch.bfloat16) + ) + + torch.testing.assert_close( + out, mass_to_out(torch.sigmoid(lse - sink.view(-1, 1))), rtol=RTOL, atol=ATOL + ) + # Control: no sink => every row mass is exactly 1. + torch.testing.assert_close(out_ns, torch.ones_like(out_ns), rtol=RTOL, atol=ATOL) + _assert_sink_effect(out, out_ns, "denominator: sink ignored") + _assert_sink_effect( + out, + mass_to_out(torch.sigmoid(lse - sink.view(-1, 1) * scale)), + "denominator: sink pre-scaled", + ) + # Every row mass strictly below 1: the dropped sink column is real. + assert bool((out.float() < 1.0).all()), "some softmax row still sums to 1" + + +def test_bf16_sinks_per_head_indexing() -> None: + """Sinks are indexed by *q* head. A checkerboard sink (+30 on even heads, + -100 on odd ones) makes the indexing structurally visible: heads with the + saturating sink must come back at ~0 (their whole softmax mass moves to the + dropped column), heads with the inert sink must be bit-identical to the + same call with attention_sinks=None. Any other indexing (kv head, an + offset, a transpose) scrambles which columns are which.""" + torch.manual_seed(62) + env = _sink_env() + env.add_request(0, 33) + env.add_request(1, 9) + qkv = env.random_qkv(42) + heads = torch.arange(SINK_HQ, device="cuda") + sink = torch.where( + heads % 2 == 0, + torch.full_like(heads, 30, dtype=torch.float32), + torch.full_like(heads, -100, dtype=torch.float32), + ) + out = env.call_op(qkv, [33, 9], 2, [0, 1], MASK_CAUSAL, attention_sinks=sink, record=False) + out_ns = env.call_op(qkv, [33, 9], 2, [0, 1], MASK_CAUSAL, record=False) + per_head = out.view(-1, SINK_HQ, SINK_D) + per_head_ns = out_ns.view(-1, SINK_HQ, SINK_D) + # exp(max_logit - 30) <= exp(-20) at these score magnitudes: the surviving + # mass is ~1e-9, far below one bf16 ulp of the ~1 no-sink output. + assert per_head[:, 0::2].abs().max().item() < 1e-6, "saturated head not drained" + assert _bitwise_equal(per_head[:, 1::2], per_head_ns[:, 1::2]), ( + "inert-sink head differs from the no-sink run" + ) + assert per_head_ns[:, 0::2].abs().max().item() > 0.1, "no-sink control is degenerate" + + +def test_bf16_sinks_long_history_decode() -> None: + """The generation kernel splits long KV ranges across CTAs and folds the + partial softmax states in a separate reduction; the sink is applied in that + reduction, so a short-history decode does not cover it. Histories of 600 + and 2000 tokens select the MultiCtasKvCga decode variant. + + Sinks are drawn near logsumexp(scaled scores) ~ log(kv_len) for such a + range, so the sink column carries a resolvable share of the mass: a + "realistic" small sink is genuinely negligible against 2000 keys and would + make the separation gate vacuous rather than the semantics different. + """ + for i, hist in enumerate((600, 2000)): + torch.manual_seed(63 + i) + env = _sink_env(num_blocks=160, max_blocks_per_seq=72) + env.add_request(0, hist) + env.call_op(env.random_qkv(hist), [hist], 1, [0], MASK_CAUSAL) + sink = _sinks(math.log(hist) - 1.0, math.log(hist) + 2.0, seed=63 + i) + qkv_d = env.random_qkv(1) + out = env.call_op(qkv_d, [1], 0, [0], MASK_CAUSAL, attention_sinks=sink) + _check_sink_case(env, qkv_d, out, [1], [0], [hist], sink, f"decode over {hist} cached") + + +def test_bf16_sinks_inert_path() -> None: + """attention_sinks=None must keep reproducing the certified no-sink + semantics, and a sink that cannot contribute must be indistinguishable + from absent — bit for bit, in both phases. + + Repeating an identical call is safe: the append derives its position from + sequence_length minus the new-token count, so it rewrites the same slots. + """ + torch.manual_seed(64) + env = _sink_env() + env.add_request(0, 40) + env.add_request(1, 41) + qkv = env.random_qkv(81) + ctx = env.call_op(qkv, [40, 41], 2, [0, 1], MASK_CAUSAL, record=False) + assert _bitwise_equal(ctx, env.call_op(qkv, [40, 41], 2, [0, 1], MASK_CAUSAL, record=False)), ( + "context call is not bitwise reproducible" + ) + for value in (-100.0, float("-inf")): + inert = torch.full((SINK_HQ,), value, dtype=torch.float32, device="cuda") + assert _bitwise_equal( + ctx, + env.call_op( + qkv, + [40, 41], + 2, + [0, 1], + MASK_CAUSAL, + attention_sinks=inert, + record=False, + ), + ), f"context: sink={value} is not bitwise inert" + # Record the prefill once, then the same checks on the generation kernel. + env.call_op(qkv, [40, 41], 2, [0, 1], MASK_CAUSAL) + torch.testing.assert_close( + ctx, + env.reference(qkv, [40, 41], [0, 1], [0, 0], MASK_CAUSAL), + rtol=RTOL, + atol=ATOL, + ) + qkv_d = env.random_qkv(2) + gen = env.call_op(qkv_d, [1, 1], 0, [0, 1], MASK_CAUSAL, record=False) + assert _bitwise_equal(gen, env.call_op(qkv_d, [1, 1], 0, [0, 1], MASK_CAUSAL, record=False)), ( + "generation call is not bitwise reproducible" + ) + for value in (-100.0, float("-inf")): + inert = torch.full((SINK_HQ,), value, dtype=torch.float32, device="cuda") + assert _bitwise_equal( + gen, + env.call_op( + qkv_d, + [1, 1], + 0, + [0, 1], + MASK_CAUSAL, + attention_sinks=inert, + record=False, + ), + ), f"generation: sink={value} is not bitwise inert" + env.call_op(qkv_d, [1, 1], 0, [0, 1], MASK_CAUSAL) # record the decode step + torch.testing.assert_close( + gen, + env.reference(qkv_d, [1, 1], [0, 1], [40, 41], MASK_CAUSAL), + rtol=RTOL, + atol=ATOL, + ) + + +def test_sinks_reject_bad_dtype_and_layout() -> None: + """Only the dtype is checked by the op ('Expected attention_sinks to have + float dtype'). Size and contiguity are not: the buffer is read as + num_heads raw fp32 values from data_ptr(), so a short tensor is read past + its end and a strided view is read as its underlying memory — both + silently wrong. The wrapper is what must reject those.""" + torch.manual_seed(65) + env = _sink_env() + env.add_request(0, 16) + qkv = env.random_qkv(16) + good = _sinks(-2.0, 4.0, seed=65) + other = _sinks(-2.0, 4.0, seed=66) + + for dtype in (torch.bfloat16, torch.float16, torch.float64): + try: + env.call_op( + qkv, + [16], + 1, + [0], + MASK_CAUSAL, + attention_sinks=good.to(dtype), + record=False, + ) + except RuntimeError as exc: + assert "float dtype" in str(exc), f"unexpected rejection for {dtype}: {exc}" + else: + raise AssertionError(f"{dtype} attention_sinks was not rejected") + + # Interleaving good/other and taking every second element yields a tensor + # whose *values* are `good` but whose memory is not — the op would read the + # interleaved bytes. + interleaved = torch.stack([good, other], dim=1).reshape(-1) + bad_layouts = { + "stride-2 view": interleaved[::2], + "too few elements": good[: SINK_HQ // 2].contiguous(), + "too many elements": torch.cat([good, other]), + "empty": torch.empty(0, dtype=torch.float32, device="cuda"), + } + for label, sinks in bad_layouts.items(): + try: + env.call_op(qkv, [16], 1, [0], MASK_CAUSAL, attention_sinks=sinks, record=False) + except AssertionError: + continue + raise AssertionError(f"{label} attention_sinks was not rejected by the wrapper") + + +# ─── Cyclic sliding window (attention_window_size < max_seq_len) + sinks ─ +# gpt-oss-120b tp1 sliding layers: 64q/8kv d64, window 128, per-head sink. +SWA_W = 128 +# Relative tolerance for the exact per-key attention weights read out by the +# one-hot probe below. Those weights are pure fp32 softmax constants rounded +# once to bf16: half an ulp is 2^-9 = 2.0e-3 relative, and the kernel's own +# fp32 softmax rounding lands on top. 6e-3 is ~1.5 bf16 ulp; the worst +# observed on sm_100 across the probe cases is 3.6e-3, 59% of it. +SWA_WEIGHT_RTOL = 6e-3 + + +def _swa_env(**kwargs) -> _PagedAttnEnv: + """Sink-geometry env with the cyclic sliding window switched on.""" + kwargs.setdefault("attention_window_size", SWA_W) + return _sink_env(**kwargs) + + +def _swa_indicator_qkv(n_tokens: int, positions: List[int], t0: int) -> torch.Tensor: + """Packed QKV that turns the op into an attention-weight read-out. + + Every K row is zero, so every unmasked logit is exactly 0 and the softmax + is uniform over precisely the attended set. V is a one-hot indicator — + V[t] = e_{t - t0} for t in [t0, t0 + head_dim) and 0 otherwise — so output + column d of a query row IS the attention weight that row gives key t0 + d. + A key outside the window must therefore come back as an exact zero. + """ + q_w, kv_w = SINK_HQ * SINK_D, SINK_HKV * SINK_D + qkv = torch.zeros(n_tokens, q_w + 2 * kv_w, dtype=torch.bfloat16, device="cuda") + qkv[:, :q_w] = torch.randn(n_tokens, q_w, dtype=torch.bfloat16, device="cuda") + v = torch.zeros(n_tokens, SINK_HKV, SINK_D, dtype=torch.bfloat16, device="cuda") + for i, t in enumerate(positions): + d = t - t0 + if 0 <= d < SINK_D: + v[i, :, d] = 1.0 + qkv[:, q_w + kv_w :] = v.reshape(n_tokens, kv_w) + return qkv + + +def _assert_window_weights( + out: torch.Tensor, + row: int, + pos: int, + t0: int, + sink: Optional[torch.Tensor], + label: str, + window: int = SWA_W, +) -> None: + """Output row `row` belongs to the query at absolute position `pos`. Under + the one-hot probe its column d holds the weight of key t0 + d, which must + be exactly 0 outside [pos - window + 1, pos] and 1 / (n + exp(sink_h)) + inside it, with n = min(pos + 1, window) keys.""" + w = out.view(-1, SINK_HQ, SINK_D)[row].float() # [Hq, D] weights + lo, n = max(0, pos - window + 1), min(pos + 1, window) + denom = float(n) + ( + torch.zeros(SINK_HQ, dtype=torch.float64, device="cuda") + if sink is None + else torch.exp(sink.double()) + ) + expect = (1.0 / denom).float() + inside = [d for d in range(SINK_D) if lo <= t0 + d <= pos] + outside = [d for d in range(SINK_D) if not (lo <= t0 + d <= pos)] + if outside: + assert bool((w[:, outside] == 0).all()), ( + f"{label}: keys outside [{lo}, {pos}] carry weight " + f"{w[:, outside].abs().max().item():.4g}" + ) + for d in inside: + torch.testing.assert_close( + w[:, d], + expect, + rtol=SWA_WEIGHT_RTOL, + atol=0.0, + msg=lambda m, d=d: f"{label}: key {t0 + d} weight off\n{m}", + ) + + +def test_bf16_swa_window_boundary_exact() -> None: + """Which keys does a query at absolute position p actually attend to? + + The one-hot probe answers it without any tolerance on the boundary: an + out-of-window key is an exact zero, an in-window key is the uniform weight + 1 / (n + exp(sink_h)). Two probe placements pin both ends — + t0 = 0 catches the moment the window starts biting (row 128 must drop key + 0) and t0 = 66 catches the moving lower edge (row 199 must drop keys + 66..71, keeping 72). The result: the attended set is [p - W + 1, p], + exactly W keys, never W + 1 — in the context phase and in the generation + phase, with the sink active in both. + """ + torch.manual_seed(70) + sink = _sinks(-1.0, 3.0, seed=70) + prefill = 200 # runs well past the 128 window + for t0 in (0, 66): + for sk in (sink, None): + env = _swa_env() + env.add_request(0, prefill) + qkv = _swa_indicator_qkv(prefill, list(range(prefill)), t0) + out = env.call_op(qkv, [prefill], 1, [0], MASK_CAUSAL, attention_sinks=sk) + for row in (66, 100, 127, 128, 150, 199): + _assert_window_weights(out, row, row, t0, sk, f"context t0={t0} row={row}") + qd = _swa_indicator_qkv(1, [prefill], t0) + out_d = env.call_op(qd, [1], 0, [0], MASK_CAUSAL, attention_sinks=sk) + _assert_window_weights(out_d, 0, prefill, t0, sk, f"decode t0={t0}") + # Sanity: the probe is not vacuous — with the window off, the same prefill + # gives every key of row 199 a non-zero weight (1 / 200), so the zeros + # above come from the window and not from the construction. + env = _sink_env() + env.add_request(0, prefill) + qkv = _swa_indicator_qkv(prefill, list(range(prefill)), 66) + out = env.call_op(qkv, [prefill], 1, [0], MASK_CAUSAL) + w199 = out.view(-1, SINK_HQ, SINK_D)[199].float() + torch.testing.assert_close( + w199, torch.full_like(w199, 1.0 / prefill), rtol=SWA_WEIGHT_RTOL, atol=0.0 + ) + + +def test_bf16_swa_window_values_not_page_aligned() -> None: + """The window is a token count, not a page count. Windows that are not + multiples of tokens_per_block (33, 100 against pages of 32) — including one + shorter than a single page — put the boundary at exactly the same + [p - window + 1, p], read out one-hot with the sink active.""" + torch.manual_seed(79) + sink = _sinks(-1.0, 3.0, seed=79) + prefill = 150 + for window in (33, 100): + env = _sink_env(attention_window_size=window) + env.add_request(0, prefill) + t0 = prefill - 1 - window - 3 # straddles the lower edge of the last rows + qkv = _swa_indicator_qkv(prefill, list(range(prefill)), t0) + out = env.call_op(qkv, [prefill], 1, [0], MASK_CAUSAL, attention_sinks=sink) + _assert_window_weights( + out, prefill - 1, prefill - 1, t0, sink, f"context w={window}", window + ) + qd = _swa_indicator_qkv(1, [prefill], t0) + out_d = env.call_op(qd, [1], 0, [0], MASK_CAUSAL, attention_sinks=sink) + _assert_window_weights(out_d, 0, prefill, t0, sink, f"decode w={window}", window) + + +def test_bf16_swa_sinks_context_and_decode() -> None: + """The gpt-oss-120b sliding layer on realistic inputs: window 128 with a + per-head sink, over a prefill that runs past the window plus three decode + steps. Every case is gated against three rivals — sink ignored, sink + pre-scaled, window ignored (full causal) — so neither mechanism can hide + behind the other. The paged append is unaffected by the window: it writes + each new token at its absolute position, so the pool still holds the whole + history bit-exactly and nothing outside the token ranges is touched. + + Sinks are drawn near log(window) — the logsumexp scale of a 128-key + softmax — so the sink column keeps a resolvable share of the mass once the + window is full; a smaller sink is genuinely negligible there and would + make the sink-ignored separation gate vacuous rather than the semantics + different.""" + torch.manual_seed(71) + env = _swa_env(max_blocks_per_seq=16, num_blocks=96) + sink = _sinks(math.log(SWA_W) - 1.0, math.log(SWA_W) + 2.0, seed=71) + env.add_request(0, 200) # 200 > 128: the window bites inside the prefill + env.add_request(1, 96) # shorter than the window: plain causal + qkv = env.random_qkv(296) + out = env.call_op(qkv, [200, 96], 2, [0, 1], MASK_CAUSAL, attention_sinks=sink) + # The hand-written windowed softmax must agree with the file's SDPA-based + # reference when the sink is absent, so the sink cases test the sink. + torch.testing.assert_close( + _sink_reference(env, qkv, [200, 96], [0, 1], [0, 0], None, window=SWA_W), + env.reference(qkv, [200, 96], [0, 1], [0, 0], MASK_CAUSAL, window=SWA_W), + rtol=RTOL, + atol=ATOL, + ) + _check_sink_case(env, qkv, out, [200, 96], [0, 1], [0, 0], sink, "swa context", window=SWA_W) + + for step in range(3): + cached = [env.cached_len(0), env.cached_len(1)] + qkv_d = env.random_qkv(2) + out_d = env.call_op(qkv_d, [1, 1], 0, [0, 1], MASK_CAUSAL, attention_sinks=sink) + _check_sink_case( + env, + qkv_d, + out_d, + [1, 1], + [0, 1], + cached, + sink, + f"swa decode step {step}", + window=SWA_W, + ) + for rid in (0, 1): + env.check_cache(rid) # whole history, bit-exact, at absolute positions + env.check_unwritten_pool_zero() + + +def test_bf16_swa_sinks_mixed_batch_and_long_history() -> None: + """Two shapes the window has to survive besides the plain ones: a mixed + context+generation call, and a decode over a history far longer than the + window, which selects the multi-CTA-KV decode variant that folds partial + softmax states across CTAs. Sinks are drawn near log(window) so the sink + column carries a resolvable share of a 128-key softmax — a small sink is + negligible there and would make the separation gate vacuous.""" + torch.manual_seed(72) + env = _swa_env(max_blocks_per_seq=16, num_blocks=96) + sink = _sinks(math.log(SWA_W) - 1.0, math.log(SWA_W) + 2.0, seed=72) + env.add_request(0, 300) + env.call_op(env.random_qkv(300), [300], 1, [0], MASK_CAUSAL, attention_sinks=sink) + env.add_request(1, 210) # context sequence sharing the call with a decode one + cached_gen = env.cached_len(0) + qkv_m = env.random_qkv(211) + out_m = env.call_op(qkv_m, [210, 1], 1, [1, 0], MASK_CAUSAL, attention_sinks=sink) + _check_sink_case( + env, + qkv_m, + out_m, + [210, 1], + [1, 0], + [0, cached_gen], + sink, + "swa mixed batch", + window=SWA_W, + ) + + long_env = _swa_env(num_blocks=200, max_blocks_per_seq=72) + long_env.add_request(0, 2000) + long_env.call_op(long_env.random_qkv(2000), [2000], 1, [0], MASK_CAUSAL) + qd = long_env.random_qkv(1) + out_l = long_env.call_op(qd, [1], 0, [0], MASK_CAUSAL, attention_sinks=sink) + _check_sink_case( + long_env, + qd, + out_l, + [1], + [0], + [2000], + sink, + "swa decode over 2000 cached", + window=SWA_W, + ) + + +def test_bf16_swa_cyclic_pool_wraps() -> None: + """The cache genuinely wraps, not merely gets masked. + + The caller gives each sequence a bounded ring of physical pages and maps + absolute page index j to ring[j % P]. Because the op appends every token + at its *absolute* position (page j = t // tokens_per_block, slot + t % tokens_per_block), that mapping makes the pool hold exactly the last + P * tokens_per_block tokens, overwriting aged-out slots in place. With + P * tokens_per_block >= window every decode still sees its full window. + Here the ring holds 192 tokens and the sequence runs to 400 — two full + wraps — with the sink on and every step checked against fp32.""" + torch.manual_seed(73) + tpb, ring = 32, 6 + env = _swa_env( + tokens_per_block=tpb, + num_blocks=16, + max_blocks_per_seq=16, + page_ring=ring, + ) + cap = ring * tpb # 192 token slots >= window 128 + sink = _sinks(math.log(SWA_W) - 1.0, math.log(SWA_W) + 2.0, seed=73) + prefill = 150 # <= cap: one call's new tokens must map to distinct slots + env.add_request(0, prefill) + qkv = env.random_qkv(prefill) + out = env.call_op(qkv, [prefill], 1, [0], MASK_CAUSAL, attention_sinks=sink) + _check_sink_case(env, qkv, out, [prefill], [0], [0], sink, "ring context", window=SWA_W) + + wraps = 0 + for step in range(250): + pos = env.cached_len(0) + qkv_d = env.random_qkv(1) + page, slot = (pos // tpb) % ring, pos % tpb + before = env.kv_cache.clone() + out_d = env.call_op(qkv_d, [1], 0, [0], MASK_CAUSAL, attention_sinks=sink) + if pos >= cap: # this slot already held token pos - cap + wraps += 1 + evicted = torch.cat(env.k_history[0])[pos - cap] + assert _bitwise_equal(before[env.rings[0][page], 0, :, slot, :], evicted), ( + f"step {pos}: ring slot did not hold the token about to age out" + ) + # Exactly one (K, V) row pair moves per decode: the append overwrites + # the evicted slot in place and touches nothing else in the pool. + changed = { + (int(b), int(s)) + for b, _kv, _h, s in (before != env.kv_cache).any(-1).nonzero().tolist() + } + assert changed == {(env.rings[0][page], slot)}, ( + f"step {pos}: decode wrote pool slots {sorted(changed)}" + ) + if step % 25 == 0 or step == 249: # full rival gating, periodically + _check_sink_case( + env, + qkv_d, + out_d, + [1], + [0], + [pos], + sink, + f"ring decode p={pos}", + window=SWA_W, + ) + else: + torch.testing.assert_close( + out_d, + _sink_reference(env, qkv_d, [1], [0], [pos], sink, window=SWA_W), + rtol=RTOL, + atol=ATOL, + ) + assert wraps > 200, f"the ring never wrapped enough ({wraps} wrapping steps)" + # The ring physically holds the last `cap` tokens, each at slot t % cap. + history = torch.cat(env.k_history[0]) + total = history.shape[0] + rows = torch.cat([env.kv_cache[p, 0].permute(1, 0, 2) for p in env.rings[0]]).contiguous() + newest = {t % cap: t for t in range(total - cap, total)} + assert _bitwise_equal(rows, torch.stack([history[newest[s]] for s in range(cap)])), ( + "the ring does not hold the last cap tokens at slot t % cap" + ) + + +def test_bf16_swa_aged_out_pages_and_ring_bound() -> None: + """What the caller still owes the op once tokens have aged out. + + A page whose tokens all sit below the window is never read — pointing its + block-offset entry at a loud decoy leaves the output bitwise unchanged, so + those pages may be recycled. Every page holding at least one in-window + token must be valid: the same substitution there changes the result. The + ring bound follows: P * tokens_per_block >= window, and one page short is + silently wrong rather than rejected.""" + torch.manual_seed(74) + tpb = 32 + env = _swa_env(tokens_per_block=tpb, num_blocks=64, max_blocks_per_seq=16) + env.add_request(0, 200) + env.call_op(env.random_qkv(200), [200], 1, [0], MASK_CAUSAL) + qd = env.random_qkv(1) + # record=False throughout: every call below is the same decode (kv length + # 201, window [73, 200]) with only the page table changing, and the append + # is idempotent, so the outputs are directly comparable bit for bit. + base = env.call_op(qd, [1], 0, [0], MASK_CAUSAL, record=False) + decoy = env.kv_cache.shape[0] - 1 + env.kv_cache[decoy] = 7.0 + live_from = (201 - SWA_W) // tpb # first page holding an in-window token + pages = env.pages[0] + for j in range(len(pages)): + keep, pages[j] = pages[j], decoy + out = env.call_op(qd, [1], 0, [0], MASK_CAUSAL, record=False) + pages[j] = keep + if j < live_from: + assert _bitwise_equal(out, base), ( + f"page {j} is fully aged out but its content still reached the output" + ) + else: + assert not _bitwise_equal(out, base), ( + f"page {j} holds in-window tokens but the decoy changed nothing" + ) + + for ring, ok in ((SWA_W // tpb, True), (SWA_W // tpb - 1, False)): + renv = _swa_env(tokens_per_block=tpb, num_blocks=16, max_blocks_per_seq=16, page_ring=ring) + renv.add_request(0, 100) + renv.call_op(renv.random_qkv(100), [100], 1, [0], MASK_CAUSAL) + worst, allowance = 0.0, 0.0 + for _ in range(120): + pos = renv.cached_len(0) + qkv_d = renv.random_qkv(1) + out_d = renv.call_op(qkv_d, [1], 0, [0], MASK_CAUSAL) + ref = _sink_reference(renv, qkv_d, [1], [0], [pos], None, window=SWA_W) + worst = max(worst, (out_d.float() - ref.float()).abs().max().item()) + allowance = max(allowance, ATOL + RTOL * ref.float().abs().max().item()) + if ok: + assert worst <= allowance, ( + f"a ring of {ring} pages holds the whole {SWA_W}-token window but " + f"the output is off by {worst:.4g}" + ) + else: + assert worst > 10 * allowance, ( + f"a ring of {ring} pages cannot hold the window, yet the output " + f"is within {worst:.4g} of the reference" + ) + + +def test_bf16_swa_context_prefill_longer_than_window() -> None: + """A single context prefill longer than the window is legal in one call: + the output matches the windowed reference row by row. The context FMHA + reads the packed QKV, never the pool — aliasing every page of the sequence + onto one decoy page leaves the context output bitwise identical (the + append still runs, so the pool itself becomes garbage; that call is last). + """ + torch.manual_seed(75) + env = _swa_env(num_blocks=64, max_blocks_per_seq=16) + sink = _sinks(math.log(SWA_W) - 1.0, math.log(SWA_W) + 2.0, seed=75) + prefill = 300 # 2.3x the window in one call + env.add_request(0, prefill) + qkv = env.random_qkv(prefill) + out = env.call_op(qkv, [prefill], 1, [0], MASK_CAUSAL, attention_sinks=sink, record=False) + # Same call with every page of the sequence aliased onto one loud decoy: + # the pool the call reads back is nonsense, the output must not move. + decoy = env.kv_cache.shape[0] - 1 + env.kv_cache[decoy] = 7.0 + keep, env.pages[0] = env.pages[0], [decoy] * len(env.pages[0]) + aliased = env.call_op(qkv, [prefill], 1, [0], MASK_CAUSAL, attention_sinks=sink, record=False) + env.pages[0] = keep + assert _bitwise_equal(aliased, out), ( + "context output depends on the pool content; it must read packed QKV only" + ) + # Re-run over the real pages (the append is idempotent) and record it, so + # the reference and the cache check see the true history. + again = env.call_op(qkv, [prefill], 1, [0], MASK_CAUSAL, attention_sinks=sink) + assert _bitwise_equal(again, out), "context call is not bitwise reproducible" + _check_sink_case(env, qkv, out, [prefill], [0], [0], sink, "long prefill", window=SWA_W) + env.check_cache(0) # the prefill wrote every token at its absolute position + + +def test_bf16_swa_batch_state_requirements() -> None: + """What the length tensors must carry once tokens have aged out. + + sequence_length stays the *global* cached+new count: it is what the window + is measured back from and where the append lands, so capping it to the + window both relocates the write and truncates the attended range. + context_lengths on a generation row and max_seq_len are inert — the same + decode comes back bitwise identical for every value tried.""" + torch.manual_seed(76) + prefill = 200 + + def prepared() -> Tuple[_PagedAttnEnv, torch.Tensor]: + env = _swa_env(num_blocks=64, max_blocks_per_seq=16) + env.add_request(0, prefill) + torch.manual_seed(76) + env.call_op(env.random_qkv(prefill), [prefill], 1, [0], MASK_CAUSAL) + torch.manual_seed(77) + return env, env.random_qkv(1) + + env, qd = prepared() + base = env.call_op(qd, [1], 0, [0], MASK_CAUSAL) + ref = _sink_reference(env, qd, [1], [0], [prefill], None, window=SWA_W) + torch.testing.assert_close(base, ref, rtol=RTOL, atol=ATOL) + + capped_env, capped_q = prepared() + capped = capped_env.call_op( + capped_q, [1], 0, [0], MASK_CAUSAL, record=False, seq_lens_override=[SWA_W] + ) + allowance = ATOL + RTOL * ref.float().abs().max().item() + assert (capped.float() - base.float()).abs().max().item() > 10 * allowance, ( + "a window-capped sequence_length silently produced the same answer" + ) + kv_w = SINK_HKV * SINK_D + q_w = SINK_HQ * SINK_D + relocated = capped_q[0, q_w : q_w + kv_w].view(SINK_HKV, SINK_D) + got_k, _, _, _ = capped_env.cache_pages_content(0) + assert _bitwise_equal(got_k[SWA_W - 1], relocated), ( + "sequence_length also drives the append position: capping it must have " + "written the new token at window - 1" + ) + + for ctx_len in (0, 1, prefill): + env_c, q_c = prepared() + out_c = env_c.call_op( + q_c, [1], 0, [0], MASK_CAUSAL, record=False, ctx_lens_override=[ctx_len] + ) + assert _bitwise_equal(out_c, base), f"context_lengths={ctx_len} changed a decode" + for msl in (SWA_W, prefill + 1, 4 * (prefill + 1)): + env_m, q_m = prepared() + out_m = env_m.call_op(q_m, [1], 0, [0], MASK_CAUSAL, record=False, max_seq_len_override=msl) + assert _bitwise_equal(out_m, base), f"max_seq_len={msl} changed a decode" + + +def test_bf16_swa_per_layer_windows_shared_pool() -> None: + """attention_window_size is a per-call scalar and cache addressing is by + absolute position, so layers with different windows share one pool, one + layer->pool mapping and one block-offset table inside the same batch — + the alternation gpt-oss runs (sliding, full, sliding, full). Each layer is + checked against its own window's reference and against the other layer's + window, which must be far away.""" + torch.manual_seed(78) + num_layers, prefill = 4, 200 + env = _MultiLayerPagedAttnEnv( + num_layers=num_layers, + num_heads=SINK_HQ, + num_kv_heads=SINK_HKV, + head_dim=SINK_D, + max_seq_len=512, + ) + windows = [SWA_W, None, SWA_W, None] # None = full attention + env.add_request(0, prefill) + env.refresh_offsets([0], num_contexts=1) + for layer, window in enumerate(windows): + qkv = env.random_qkv(prefill) + out = env.call_op(layer, qkv, [prefill], 1, [0], attention_window_size=window) + ref = env.reference(layer, qkv, [prefill], [0], [0], window=window) + torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) + other = env.reference(layer, qkv, [prefill], [0], [0], window=None if window else SWA_W) + gap = (out.float() - other.float()).abs().max().item() + assert gap > 5.0 * (ATOL + RTOL * other.float().abs().max().item()), ( + f"layer {layer}: the two windows are indistinguishable here" + ) + for _ in range(2): + env.add_decode_token(0) + env.refresh_offsets([0], num_contexts=0) + for layer, window in enumerate(windows): + cached = env.cached_len(layer, 0) + qkv_d = env.random_qkv(1) + out_d = env.call_op(layer, qkv_d, [1], 0, [0], attention_window_size=window) + ref_d = env.reference(layer, qkv_d, [1], [0], [cached], window=window) + torch.testing.assert_close(out_d, ref_d, rtol=RTOL, atol=ATOL) + env.check_caches([0]) # every layer's full history, bit-exact + + +# ─── Paged-context FMHA (use_paged_context_fmha=True) ────────────────── +# The execution path an engine prepares as soon as KV-cache reuse or chunked +# prefill is on — trtllm's default, so the path every bf16 target here runs. +# With the flag set the context FMHA sources K/V from the paged pool instead +# of the packed q rows, which is what lets a context call carry a cached +# prefix: its KV range exceeds its q range and the causal mask is +# bottom-right aligned. + +# A rival batch state must land this far outside the tolerance band around +# the correct output. The wrong states measured below sit at 31x-318x on +# sm_100 (aliased pages 31x, the flag left off over a cached prefix 51x, +# context_lengths carrying the full KV length 101x or 0 53x, a decoy in +# place of an in-window cached page 248x-318x), so the gate is nowhere near +# any of them, while a correct call uses at most 66% of the same band. +PAGED_CTX_MIN_SEPARATION = 20.0 + + +def _paged_ctx_call( + env: _PagedAttnEnv, + qkv: torch.Tensor, + seq_lens: List[int], + num_contexts: int, + request_ids: List[int], + **kwargs, +) -> torch.Tensor: + """One use_paged_context_fmha=True call. A context row's context_lengths + is this call's new-token (q-row) count — what the engine puts in + prompt_lens for a context request served over a cached prefix — while a + generation row keeps the registered prompt length.""" + ctx_lens = [ + seq_lens[i] if i < num_contexts else env.prompt_lens[rid] + for i, rid in enumerate(request_ids) + ] + return env.call_op( + qkv, + seq_lens, + num_contexts, + request_ids, + MASK_CAUSAL, + ctx_lens_override=ctx_lens, + use_paged_context_fmha=True, + **kwargs, + ) + + +def _run_paged_ctx_and_check( + env: _PagedAttnEnv, + seq_lens: List[int], + num_contexts: int, + request_ids: List[int], + window: Optional[int] = None, +) -> None: + """Positive gate for one paged-context call: fp32 attention over each + sequence's full [cached + new] history, causal mask bottom-right aligned + (query i of a context row sits at absolute position cached + i).""" + cached = [env.cached_len(rid) for rid in request_ids] + qkv = env.random_qkv(sum(seq_lens)) + out = _paged_ctx_call(env, qkv, seq_lens, num_contexts, request_ids) + ref = env.reference(qkv, seq_lens, request_ids, cached, MASK_CAUSAL, window=window) + torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) + + +def _assert_far_outside(out: torch.Tensor, correct: torch.Tensor, label: str) -> float: + """A wrong batch state must land far outside the tolerance band around the + correct output. A non-finite result counts as far outside — some wrong + states also corrupt the pool, and a later read of it can produce NaNs. + Returns the observed separation in units of the allowance.""" + diff = (out.float() - correct.float()).abs() + if not bool(torch.isfinite(diff).all()): + return float("inf") + allowance = ATOL + RTOL * correct.float().abs().max().item() + ratio = diff.max().item() / allowance + assert ratio > PAGED_CTX_MIN_SEPARATION, ( + f"{label}: only {ratio:.4g}x outside the tolerance band" + ) + return ratio + + +def test_bf16_paged_context_no_cached_tokens_matches_packed_path() -> None: + """With nothing cached, the flag changes no observable. + + Every fresh prefill of a target running with cache reuse takes this path, + so it has to reproduce the packed-QKV one exactly. A page-crossing prefill + batch and a decode step, run from identical state at both flag values, + come back bitwise identical in output and in pool content — on all three + shipped bf16 geometries. The K/V source does move (see the aliasing case + below); it is the append running first that leaves the pool holding + exactly the packed rows the other path reads. + """ + for i, (num_heads, num_kv_heads, head_dim) in enumerate( + [(32, 8, 128), (32, 4, 128), (64, 8, 64)] + ): + runs = [] + for use_paged in (False, True): + torch.manual_seed(80 + i) + env = _PagedAttnEnv(num_heads=num_heads, num_kv_heads=num_kv_heads, head_dim=head_dim) + env.add_request(0, 48) # crosses a 32-token page boundary + env.add_request(1, 17) + qkv = env.random_qkv(65) + out = env.call_op( + qkv, + [48, 17], + 2, + [0, 1], + MASK_CAUSAL, + use_paged_context_fmha=use_paged, + ) + torch.testing.assert_close( + out, + env.reference(qkv, [48, 17], [0, 1], [0, 0], MASK_CAUSAL), + rtol=RTOL, + atol=ATOL, + ) + cached = [env.cached_len(0), env.cached_len(1)] + qkv_d = env.random_qkv(2) + out_d = env.call_op( + qkv_d, [1, 1], 0, [0, 1], MASK_CAUSAL, use_paged_context_fmha=use_paged + ) + torch.testing.assert_close( + out_d, + env.reference(qkv_d, [1, 1], [0, 1], cached, MASK_CAUSAL), + rtol=RTOL, + atol=ATOL, + ) + env.check_cache(0) + env.check_cache(1) + env.check_unwritten_pool_zero() + runs.append((out, out_d, env.kv_cache)) + packed_run, paged_run = runs + for what, packed, paged in zip( + ("prefill output", "decode output", "pool"), packed_run, paged_run + ): + assert _bitwise_equal(packed, paged), ( + f"{num_heads}/{num_kv_heads}/{head_dim}: {what} differs between " + f"the packed and paged context paths with nothing cached" + ) + + +def test_bf16_paged_context_cached_prefix_shipped_geometries() -> None: + """The default path of the two plain-causal shipped targets: 32q/8kv d128 + (qwen3-8b) and 32q/4kv d128 (qwen3-30b-a3b). + + One full cycle per geometry: a fresh prefill, a context call over the + cached prefix that crosses a page boundary, a second one landing exactly + on a page end, a decode step, a second sequence, and finally a mixed batch + pairing a cached-prefix context row with a fresh context row and a + generation row — the batch shape a server with prefix reuse produces. The + appends stay bit-exact and nothing outside the token ranges is written. + """ + for i, (num_heads, num_kv_heads) in enumerate([(32, 8), (32, 4)]): + torch.manual_seed(83 + i) + env = _PagedAttnEnv(num_heads=num_heads, num_kv_heads=num_kv_heads, head_dim=128) + env.add_request(0, 96) # served by three context calls: 48 + 33 + 15 + _run_paged_ctx_and_check(env, [48], 1, [0]) # fresh, crosses page 1 + _run_paged_ctx_and_check(env, [33], 1, [0]) # 48 -> 81, crosses page 2 + _run_paged_ctx_and_check(env, [15], 1, [0]) # 81 -> 96, exact page end + _run_paged_ctx_and_check(env, [1], 0, [0]) # decode over the same pool + env.add_request(1, 23) + env.add_request(2, 57) + _run_paged_ctx_and_check(env, [40], 1, [2]) # a second cached prefix + _run_paged_ctx_and_check(env, [23, 17, 1], 2, [1, 2, 0]) + for rid in (0, 1, 2): + env.check_cache(rid) + env.check_unwritten_pool_zero() + + +def test_bf16_paged_context_prefix_page_geometry() -> None: + """Cached prefixes against the 32-token page grid, six of them advanced in + one call so the batch mixes them: a one-token prefix, one ending mid-page + (31), ones ending exactly on a page boundary (32 and 64), one that has + just opened a page (33), and a multi-page one (200). Each sequence's new + tokens land its KV range on a different side of a boundary.""" + torch.manual_seed(85) + prefixes = [1, 31, 32, 33, 64, 200] + new_tokens = [1, 2, 32, 31, 40, 40] + env = _PagedAttnEnv( + num_heads=32, + num_kv_heads=8, + head_dim=128, + num_blocks=96, + max_batch=len(prefixes), + max_blocks_per_seq=16, + ) + for rid, prefix in enumerate(prefixes): + env.add_request(rid, prefix) + _run_paged_ctx_and_check(env, [prefix], 1, [rid]) # fresh prefill + _run_paged_ctx_and_check(env, new_tokens, len(new_tokens), list(range(len(prefixes)))) + for rid in range(len(prefixes)): + env.check_cache(rid) + env.check_unwritten_pool_zero() + + +def test_bf16_paged_context_chunked_prefill_loop() -> None: + """Chunked prefill: two sequences advanced together over three calls, one + on page-aligned chunks (32/32/17) and one on ragged ones (20/45/16), so + every call is a context call over whatever each sequence has cached so + far. Every chunk is gated against the full-history reference and the pool + ends up holding both prompts bit-exactly.""" + torch.manual_seed(86) + env = _PagedAttnEnv( + num_heads=32, + num_kv_heads=8, + head_dim=128, + num_blocks=96, + max_blocks_per_seq=16, + ) + env.add_request(0, 81) + env.add_request(1, 81) + for chunk_a, chunk_b in ((32, 20), (32, 45), (17, 16)): + _run_paged_ctx_and_check(env, [chunk_a, chunk_b], 2, [0, 1]) + env.check_cache(0) + env.check_cache(1) + env.check_unwritten_pool_zero() + + +def test_bf16_paged_context_reads_the_pool_and_batch_state() -> None: + """What the flag actually moves, and what the caller then owes it. + + 1. The K/V source moves to the paged pool. Aliasing every page of a + *fresh* prefill onto one loud decoy leaves the packed path's output + bitwise identical — it never reads the pool — and moves the paged + path's far outside the band. + 2. A cached prefix is read from the pool: swapping a page holding it for + the decoy changes the output, while the page holding only this call's + own new tokens is invisible (the append lands wherever the offsets + point and is read straight back from there). + 3. use_paged_context_fmha=False on a context call with a cached prefix is + silently wrong — no exception, an answer far off the reference. + 4. context_lengths on a context row must be this call's new-token count. + The sequence's full KV length and 0 are both far off. + + The wrong batch states of (3) and (4) also corrupt the pool, so every + rival runs against its own freshly prepared copy of one fixed state: + request 0 with 64 cached tokens, advanced by a 20-token context call. + """ + # (1) which memory the context FMHA reads + for use_paged in (False, True): + torch.manual_seed(87) + env = _PagedAttnEnv(num_heads=32, num_kv_heads=8, head_dim=128) + env.add_request(0, 40) + qkv = env.random_qkv(40) + base = env.call_op( + qkv, + [40], + 1, + [0], + MASK_CAUSAL, + record=False, + use_paged_context_fmha=use_paged, + ) + decoy = env.kv_cache.shape[0] - 1 + env.kv_cache[decoy] = 7.0 + keep, env.pages[0] = env.pages[0], [decoy] * len(env.pages[0]) + aliased = env.call_op( + qkv, + [40], + 1, + [0], + MASK_CAUSAL, + record=False, + use_paged_context_fmha=use_paged, + ) + env.pages[0] = keep + if use_paged: + _assert_far_outside(aliased, base, "paged context over aliased pages") + else: + assert _bitwise_equal(aliased, base), "the packed path must not read the pool at all" + + cached, new = 64, 20 + + def prepared() -> Tuple[_PagedAttnEnv, torch.Tensor]: + torch.manual_seed(88) + env = _PagedAttnEnv(num_heads=32, num_kv_heads=8, head_dim=128) + env.add_request(0, cached + new) + env.call_op( + env.random_qkv(cached), + [cached], + 1, + [0], + MASK_CAUSAL, + ctx_lens_override=[cached], + use_paged_context_fmha=True, + ) + return env, env.random_qkv(new) + + # The baseline is the correct answer, gated against fp32 here so every + # separation below is measured against a result that is known right. + env, qkv = prepared() + base = _paged_ctx_call(env, qkv, [new], 1, [0]) + torch.testing.assert_close( + base, + env.reference(qkv, [new], [0], [cached], MASK_CAUSAL), + rtol=RTOL, + atol=ATOL, + ) + env.check_cache(0) + env.check_unwritten_pool_zero() + + # (3) the flag itself, and (4) the context length + penv, pq = prepared() + assert _bitwise_equal(pq, qkv), "the prepared state is not reproducible" + _assert_far_outside( + penv.call_op( + pq, + [new], + 1, + [0], + MASK_CAUSAL, + record=False, + ctx_lens_override=[new], + use_paged_context_fmha=False, + ), + base, + "cached-prefix context at use_paged_context_fmha=False", + ) + for ctx_len in (cached + new, 0): + cenv, cq = prepared() + _assert_far_outside( + cenv.call_op( + cq, + [new], + 1, + [0], + MASK_CAUSAL, + record=False, + ctx_lens_override=[ctx_len], + use_paged_context_fmha=True, + ), + base, + f"context_lengths={ctx_len}", + ) + + # (2) page by page. Pages 0-1 hold the cached prefix and must be read; + # page 2 holds only tokens 64..83, which this call appends itself. + last_cached_page = (cached - 1) // env.tokens_per_block + denv, dq = prepared() + decoy = denv.kv_cache.shape[0] - 1 + for page in range(len(denv.pages[0])): + denv.kv_cache[decoy] = 7.0 + keep, denv.pages[0][page] = denv.pages[0][page], decoy + swapped = _paged_ctx_call(denv, dq, [new], 1, [0], record=False) + denv.pages[0][page] = keep + if page <= last_cached_page: + _assert_far_outside(swapped, base, f"cached prefix page {page} replaced") + else: + assert _bitwise_equal(swapped, base), ( + f"page {page} holds only this call's own new tokens, yet " + f"redirecting it changed the output" + ) + + +def test_bf16_paged_context_multilayer_shared_pool() -> None: + """The paged-context read goes through the same layer base shift the + append does. + + One pool shared by 4 layers — real KVCacheManager state, GQA 32q/8kv d128 + — with every layer served a fresh context chunk, then a context chunk over + its own cached prefix, then a decode step. Each layer's output is gated + against that layer's own history, so a read landing in a sibling's slabs + cannot pass; the appends stay bit-exact and no call touches another + layer's slabs.""" + torch.manual_seed(89) + num_layers = 4 + env = _MultiLayerPagedAttnEnv( + num_layers=num_layers, + num_heads=32, + num_kv_heads=8, + head_dim=128, + num_blocks=96, + max_seq_len=512, + ) + env.add_request(0, 96) # the manager allocates the whole prompt up front + env.add_request(1, 40) + env.refresh_offsets([0, 1], num_contexts=2) + for chunk in ([64, 25], [32, 15]): + for layer in range(num_layers): + cached = [env.cached_len(layer, 0), env.cached_len(layer, 1)] + qkv = env.random_qkv(sum(chunk)) + siblings_before = [env.layer_views[m].clone() for m in range(num_layers) if m != layer] + out = env.call_op( + layer, + qkv, + chunk, + 2, + [0, 1], + use_paged_context_fmha=True, + ctx_lens=chunk, + ) + ref = env.reference(layer, qkv, chunk, [0, 1], cached) + torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) + siblings_after = [env.layer_views[m] for m in range(num_layers) if m != layer] + for before, after in zip(siblings_before, siblings_after): + assert torch.equal(before, after) # no cross-layer write + env.check_caches([0, 1]) + + env.add_decode_token(0) + env.add_decode_token(1) + env.refresh_offsets([0, 1], num_contexts=0) + for layer in range(num_layers): + cached = [env.cached_len(layer, 0), env.cached_len(layer, 1)] + qkv_d = env.random_qkv(2) + out_d = env.call_op(layer, qkv_d, [1, 1], 0, [0, 1], use_paged_context_fmha=True) + ref_d = env.reference(layer, qkv_d, [1, 1], [0, 1], cached) + torch.testing.assert_close(out_d, ref_d, rtol=RTOL, atol=ATOL) + env.check_caches([0, 1]) + + +def test_bf16_paged_context_gpt_oss_sinks_and_window() -> None: + """The gpt-oss-120b tp1 cell on this path: 64q/8kv d64 with per-head + sinks, first as a full-attention layer, then as a sliding one (window 128) + whose cached prefix runs well past the window. + + Both mechanisms are already certified on the packed path; what is new is + that the keys they weigh now come out of the pool and the mask is + bottom-right aligned. Every case is gated against the sink-ignored and + sink-pre-scaled rivals, the windowed ones additionally against a + window-ignored (full causal) rival.""" + torch.manual_seed(90) + sink = _sinks(-2.0, 4.0, seed=90) + env = _sink_env() + env.add_request(0, 81) + qkv = env.random_qkv(48) + out = _paged_ctx_call(env, qkv, [48], 1, [0], attention_sinks=sink) + _check_sink_case(env, qkv, out, [48], [0], [0], sink, "paged fresh context") + qkv2 = env.random_qkv(33) + out2 = _paged_ctx_call(env, qkv2, [33], 1, [0], attention_sinks=sink) + _check_sink_case(env, qkv2, out2, [33], [0], [48], sink, "paged cached context") + env.check_cache(0) + env.check_unwritten_pool_zero() + + torch.manual_seed(91) + # Sinks near log(window), as on the packed sliding-window cases: against a + # full 128-key softmax a smaller sink is negligible and the sink-ignored + # separation gate would be vacuous rather than the semantics different. + swa_sink = _sinks(math.log(SWA_W) - 1.0, math.log(SWA_W) + 2.0, seed=91) + swa = _swa_env(num_blocks=96, max_blocks_per_seq=16) + swa.add_request(0, 240) + q_pre = swa.random_qkv(200) # 200 > 128: the window bites inside the prefill + out_pre = _paged_ctx_call(swa, q_pre, [200], 1, [0], attention_sinks=swa_sink) + _check_sink_case( + swa, + q_pre, + out_pre, + [200], + [0], + [0], + swa_sink, + "paged swa prefill", + window=SWA_W, + ) + q_new = swa.random_qkv(40) + out_new = _paged_ctx_call(swa, q_new, [40], 1, [0], attention_sinks=swa_sink) + _check_sink_case( + swa, + q_new, + out_new, + [40], + [0], + [200], + swa_sink, + "paged swa cached context", + window=SWA_W, + ) + swa.check_cache(0) + + +def test_bf16_paged_context_window_boundary_and_read_set() -> None: + """Which keys a cached-prefix context row attends to, read out exactly, + and which pages the paged read touches. + + The one-hot probe (all-zero K, one-hot V, both cached and new) turns the + output into the attention weights: for a query row at absolute position p + every key outside [p - W + 1, p] comes back bitwise zero and the ones + inside all carry 1 / (n + exp(sink_h)) with n = min(p + 1, W). That pins + the mask as bottom-right aligned over the cached prefix, and as a token + count rather than a page count. + + The read set follows, measured page by page against a decoy: a page + changes the output exactly when it holds an in-window key this call does + not write itself, i.e. an in-window *cached* token. Pages fully below the + window are unread (recyclable), and so are pages holding only this call's + new tokens, whose append lands wherever the offsets point. Under the + packed path a context call read no page at all.""" + torch.manual_seed(92) + sink = _sinks(-1.0, 3.0, seed=92) + prefill, new, t0 = 150, 40, 66 + env = _swa_env(num_blocks=96, max_blocks_per_seq=16) + env.add_request(0, prefill + new) + _paged_ctx_call( + env, + _swa_indicator_qkv(prefill, list(range(prefill)), t0), + [prefill], + 1, + [0], + attention_sinks=sink, + ) + probe = _swa_indicator_qkv(new, list(range(prefill, prefill + new)), t0) + out = _paged_ctx_call(env, probe, [new], 1, [0], attention_sinks=sink) + for row in (0, 10, new - 1): + _assert_window_weights(out, row, prefill + row, t0, sink, f"paged cached context row={row}") + + torch.manual_seed(93) + tpb, cached = 32, 200 + renv = _swa_env(tokens_per_block=tpb, num_blocks=96, max_blocks_per_seq=16) + renv.add_request(0, cached + 40) + _run_paged_ctx_and_check(renv, [cached], 1, [0], window=SWA_W) + qkv = renv.random_qkv(40) + base = _paged_ctx_call(renv, qkv, [40], 1, [0], record=False) + decoy = renv.kv_cache.shape[0] - 1 + # This call's oldest in-window key belongs to its first query row, at + # absolute position `cached`; its last cached key is `cached - 1`. + first_read = max(0, cached + 1 - SWA_W) // tpb + last_cached = (cached - 1) // tpb + for page in range(len(renv.pages[0])): + renv.kv_cache[decoy] = 7.0 + keep, renv.pages[0][page] = renv.pages[0][page], decoy + got = _paged_ctx_call(renv, qkv, [40], 1, [0], record=False) + renv.pages[0][page] = keep + if first_read <= page <= last_cached: + _assert_far_outside(got, base, f"in-window cached page {page} replaced") + else: + assert _bitwise_equal(got, base), ( + f"page {page} holds no in-window cached token, yet the decoy reached the output" + ) + assert _bitwise_equal(_paged_ctx_call(renv, qkv, [40], 1, [0], record=False), base), ( + "the page-swap probe left the pool in a different state" + ) + + +# ─── MLA configuration (is_mla_enable=True) ──────────────────────────── +# DeepSeek MLA head geometry at four query-head counts: 32, the +# deepseek-v3-lite tp1 layer shape; 16, its tp2 slice; 8, its tep4 slice +# (32 checkpoint heads split over 4 tensor-parallel ranks); and 128, the +# DeepSeek-R1-0528 layer shape, which attention DP replicates whole onto +# every rank. C/R/nope/v are identical for all four — only the head count +# moves, so every case below runs at each count. +MLA_NUM_HEADS = 16 +MLA_NUM_HEADS_H32 = 32 +MLA_NUM_HEADS_H8 = 8 +MLA_NUM_HEADS_H128 = 128 +KV_LORA_RANK = 512 # C +QK_ROPE_HEAD_DIM = 64 # R +QK_NOPE_HEAD_DIM = 128 +QK_HEAD_DIM = QK_NOPE_HEAD_DIM + QK_ROPE_HEAD_DIM # 192, context head size +V_HEAD_DIM = 128 # context v head dim +LATENT_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM # 576, generation head size +# Latent-pool page size. 32 is the engine default (KvCacheConfig +# .tokens_per_block), 64 the value a tuned MLA target opts into. Both are +# covered; the page count per sequence is scaled so max_seq_len stays 1024 +# either way, leaving the page size as the only moving axis. +MLA_TOKENS_PER_BLOCK = 64 +MLA_PAGE32 = 32 +MLA_MAX_SEQ_LEN = 1024 +POSITION_EMBEDDING_TYPE_YARN = 8 # PositionEmbeddingType.yarn (MLA models) +# q-LoRA rank. DeepSeek-V3 down-projects q through a rank-1536 q_a_proj; +# checkpoints with "q_lora_rank": null (deepseek-v3-lite) have no q-LoRA at +# all and pass 0. Both are covered — see the inertness cases at the end. +Q_LORA_RANK_DSV3 = 1536 +Q_LORA_RANK_ZERO = 0 + +# Both MLA phases scale QK^T by 1 / (q_scaling * sqrt(QK_HEAD_DIM)) — the +# generation phase too, despite its head_size being C+R. +MLA_SOFTMAX_SCALE = 1.0 / math.sqrt(QK_HEAD_DIM) + +ROTARY_SCALING_TYPE_NONE = 0 # RotaryScalingType.none +ROTARY_SCALING_TYPE_YARN = 5 # RotaryScalingType.yarn + +# DeepSeek-R1-0528's rope configuration, from its config.json: theta 10000, +# YaRN over an original 4096-position window with factor 40, beta_fast 32 / +# beta_slow 1, both mscales 1.0, 163840 trained positions. +R1_ROPE_THETA = 10000.0 +R1_ROPE_FACTOR = 40.0 +R1_ROPE_ORIGINAL_MAX_POSITIONS = 4096 +R1_ROPE_BETA_FAST = 32 +R1_ROPE_BETA_SLOW = 1 +R1_ROPE_MSCALE = 1.0 +R1_ROPE_MSCALE_ALL_DIM = 1.0 +R1_MAX_POSITION_EMBEDDINGS = 163840 + +# YaRN's attention-temperature term: mscale = 0.1 * mscale_cfg * ln(factor) + 1 +# at factor > 1. R1 folds mscale^2 into the softmax scale rather than into the +# rope table (the table's own amplitude factor is +# yarn_get_mscale(factor, mscale) / yarn_get_mscale(factor, mscale_all_dim) = 1 +# when the two config mscales are equal), so the model wants +# 1 / (q_scaling * sqrt(192)) = mscale^2 / sqrt(192). +R1_MSCALE = 0.1 * R1_ROPE_MSCALE * math.log(R1_ROPE_FACTOR) + 1.0 +R1_Q_SCALING = 1.0 / (R1_MSCALE * R1_MSCALE) + + +class _RopeScalars(NamedTuple): + """The op's seven scalar rope arguments, carried as one group. + + They arrive alongside the rotary_cos_sin table, which already encodes + theta and any scaling in its content. Whether the op reads them at all on + the MLA path is what the scalar sweep at the end of the MLA section + settles; every MLA case here passes a set explicitly so the answer is + never assumed. + """ + + rope_base: float = 10000.0 + rope_scale_type: int = ROTARY_SCALING_TYPE_NONE + rope_scale: float = 1.0 + rope_short_m_scale: float = 1.0 + rope_long_m_scale: float = 1.0 + rope_max_positions: int = MLA_MAX_SEQ_LEN + rope_original_max_positions: int = MLA_MAX_SEQ_LEN + + +# What a TrtllmAttention module forwards for a DeepSeek-R1-0528 layer: +# rope_params.theta / scale_type / scale / max_positions / +# original_max_positions straight off the checkpoint's rope_scaling block. +R1_ROPE_SCALARS = _RopeScalars( + rope_base=R1_ROPE_THETA, + rope_scale_type=ROTARY_SCALING_TYPE_YARN, + rope_scale=R1_ROPE_FACTOR, + rope_max_positions=R1_MAX_POSITION_EMBEDDINGS, + rope_original_max_positions=R1_ROPE_ORIGINAL_MAX_POSITIONS, +) + + +def _r1_rope_params(max_positions: int) -> RopeParams: + """The rope construction a DeepSeek-R1-0528 config produces, truncated to + max_positions rows. Row content is independent of the row count (position + p contributes t = p against a position-independent inv_freq), so this is + the leading slice of the 163840-row table an engine builds.""" + return RopeParams( + dim=QK_ROPE_HEAD_DIM, + theta=R1_ROPE_THETA, + scale_type=RotaryScalingType.yarn, + scale=R1_ROPE_FACTOR, + max_positions=max_positions, + original_max_positions=R1_ROPE_ORIGINAL_MAX_POSITIONS, + beta_fast=R1_ROPE_BETA_FAST, + beta_slow=R1_ROPE_BETA_SLOW, + mscale=R1_ROPE_MSCALE, + mscale_all_dim=R1_ROPE_MSCALE_ALL_DIM, + duplicate_data=True, + ) + + +# Gate for the appended latent rows' k_pe half. The kernel rotates in fp32 and +# stores bf16; the reference rotates in fp32 with torch's evaluation order. The +# two fp32 results differ by one fp32 ulp on values whose correctly-rounded +# fp32 form lands exactly on a bf16 rounding midpoint, which flips the stored +# bf16 by one ulp. Measured on sm_100 over 60 prefill runs at page sizes 32 and +# 64 (495360 roped elements): 8 such elements, 1.6e-5 of the total, never more +# than one ulp, and the same count at both page sizes — so it is an arithmetic +# tie, not an addressing effect. rtol 2**-7 covers exactly +/-1 bf16 ulp at any +# magnitude; the bit-exact-fraction floor keeps the claim "essentially bitwise" +# rather than merely "within an ulp". Wrong variants measured against this gate +# on one 96-token page-32 request (6144 roped elements): an append shifted one +# slot 6.1e5 ulp / 99.9% of elements differing, the pool read with +# tokens_per_block=64 instead of 32 2.6e5 ulp / 33%, rope applied at position +# i+1 5.1e5 ulp / 69%, k_pe left un-roped 4.4e5 ulp / 93%, another request's +# rows 4.1e5 ulp / 99.9% — every one of them >=2.6e5x past the ulp gate and +# >=330x past the fraction gate, and all but the rope-position variants also +# break the strictly bitwise compressed_kv half. +LATENT_ROPE_ULP_RTOL = 2**-7 +LATENT_ROPE_MAX_INEXACT_FRACTION = 1e-3 + + +def _assert_latent_rope(got: torch.Tensor, expected: torch.Tensor, request_id: int) -> None: + """Dual gate on the roped k_pe half of a request's appended latent rows.""" + a = got[:, KV_LORA_RANK:].float() + b = expected[:, KV_LORA_RANK:].float() + torch.testing.assert_close(a, b, rtol=LATENT_ROPE_ULP_RTOL, atol=0.0) + inexact = int((a != b).sum()) + allowed = max(1, int(LATENT_ROPE_MAX_INEXACT_FRACTION * a.numel())) + assert inexact <= allowed, ( + f"latent cache rope not bit-exact enough: request {request_id}, " + f"{inexact}/{a.numel()} elements differ (allowed {allowed})" + ) + + +# The same gate one e4m3 ulp wide, for an fp8 pool. e4m3 keeps 3 mantissa +# bits, so one ulp is 2**-3 relative — coarse enough to swallow the fp32 +# evaluation-order difference that costs the bf16 pool its 8-in-495360 +# elements: every appended row measured on sm_100 under quant_mode 128 (both +# halves, scales 1.0 / 1.5 / 2.0, H = 16 and 128, page-crossing sequences) was +# bit-exact against e4m3(row * kv_scale_orig_quant), so this gate has never +# been approached. It stays a gate rather than a bitwise assert because the +# rounding tie it covers is arithmetic, not addressing: a mis-addressed append +# lands orders of magnitude outside it (see LATENT_ROPE_ULP_RTOL's wrong-variant +# figures, which the fp8 append shares — the addressing is the same code). +LATENT_ROPE_E4M3_ULP_RTOL = 2**-3 + + +def _assert_latent_rope_e4m3(got: torch.Tensor, expected: torch.Tensor, request_id: int) -> None: + """One-e4m3-ulp gate plus a bit-exact-fraction floor on the roped k_pe half + of a request's appended latent rows in an fp8 pool.""" + a = got[:, KV_LORA_RANK:].float() + b = expected[:, KV_LORA_RANK:].float() + torch.testing.assert_close(a, b, rtol=LATENT_ROPE_E4M3_ULP_RTOL, atol=0.0) + inexact = int( + ( + got[:, KV_LORA_RANK:].view(torch.uint8) != expected[:, KV_LORA_RANK:].view(torch.uint8) + ).sum() + ) + allowed = max(1, int(LATENT_ROPE_MAX_INEXACT_FRACTION * a.numel())) + assert inexact <= allowed, ( + f"fp8 latent cache rope not bit-exact enough: request {request_id}, " + f"{inexact}/{a.numel()} elements differ (allowed {allowed})" + ) + + +class _MlaPagedEnv: + """Real MLA op state: caller-owned paged latent pool (kv_factor 1) + + explicit metadata tensors + the duplicated-layout MLA RoPE table. + + Mirrors the expected latent-cache rows per request (computed with test- + side torch math, never read back from the op), so decode references can + be built from the true cache content and the pool can be checked + bitwise against it. + """ + + def __init__( + self, + num_heads: int = MLA_NUM_HEADS, + num_blocks: int = 64, + max_batch: int = 4, + max_blocks_per_seq: int = 16, + tokens_per_block: int = MLA_TOKENS_PER_BLOCK, + q_lora_rank: Optional[int] = Q_LORA_RANK_DSV3, + q_scaling: float = 1.0, + rope: Optional[RopeParams] = None, + rope_scalars: Optional[_RopeScalars] = None, + pool_dtype: torch.dtype = torch.bfloat16, + quant_mode: int = 0, + kv_scaling_factor: Optional[float] = None, + ) -> None: + self.num_heads = num_heads + self.max_batch = max_batch + self.tokens_per_block = tokens_per_block + self.q_lora_rank = q_lora_rank + self.max_seq_len = max_blocks_per_seq * tokens_per_block + self.q_scaling = q_scaling + self.quant_mode = quant_mode + # KV-cache scaling factor s (fp8 pool only): a latent row lands in the + # pool as e4m3(row * kv_scale_orig_quant) with orig_quant = 1/s, and + # dequantizes as value * kv_scale_quant_orig with quant_orig = s. None + # leaves both tensors unpassed, which the op reads as s = 1.0. + self.kv_scale = 1.0 if kv_scaling_factor is None else kv_scaling_factor + if kv_scaling_factor is None: + self.kv_scale_quant_orig: Optional[torch.Tensor] = None + self.kv_scale_orig_quant: Optional[torch.Tensor] = None + else: + factor = torch.full((1,), kv_scaling_factor, dtype=torch.float32, device="cuda") + self.kv_scale_quant_orig = factor + self.kv_scale_orig_quant = 1.0 / factor + # Both MLA phases scale QK^T by 1 / (q_scaling * sqrt(nope + R)). + self.softmax_scale = 1.0 / (q_scaling * math.sqrt(QK_HEAD_DIM)) + + # Single-layer, single-pool paged MLA latent cache: kv_factor 1, one + # kv head, row width C+R. Page p occupies pool[p] — one slab of + # tokens_per_block * (C+R) elements, one byte each under quant_mode + # 128 (e4m3) and two under quant_mode 0 (bf16). + self.pool = torch.zeros( + num_blocks, + tokens_per_block, + LATENT_DIM, + dtype=pool_dtype, + device="cuda", + ) + self.pool_pointers = torch.zeros(1, 2, dtype=torch.int64, device="cpu") + self.pool_pointers[0, 0] = self.pool.data_ptr() + self.pool_mapping = torch.zeros(1, 2, dtype=torch.int32, device="cpu") + self.block_offsets = torch.zeros( + 1, max_batch, 2, max_blocks_per_seq, dtype=torch.int32, device="cuda" + ) + # Auto-resized in place by the op on first call. + self.workspace = torch.empty(0, dtype=torch.int8, device="cuda") + + # Duplicated-layout fp32 (cos, sin) table, as the MLA backend builds + # it (RopeParams.from_config sets duplicate_data=True for MLA models). + # Default: the unscaled theta-10000 table. rope, when given, supplies + # another construction — the YaRN one a DeepSeek-R1 config produces. + rope = rope or RopeParams( + dim=QK_ROPE_HEAD_DIM, + theta=10000.0, + max_positions=self.max_seq_len, + duplicate_data=True, + ) + self.rotary_inv_freq, self.rotary_cos_sin = rope.create_rope_const_params() + self.rope_scalars = rope_scalars or _RopeScalars( + rope_max_positions=self.max_seq_len, + rope_original_max_positions=self.max_seq_len, + ) + + self._next_free_page = 0 + self.pages: Dict[int, List[int]] = {} + self.prompt_lens: Dict[int, int] = {} + # Expected cache rows per request, [n, C+R] each, test-computed. + self.latent_rows: Dict[int, List[torch.Tensor]] = {} + + def add_request(self, request_id: int, prompt_len: int) -> None: + self.pages[request_id] = [] + self.prompt_lens[request_id] = prompt_len + self.latent_rows[request_id] = [] + + def cached_len(self, request_id: int) -> int: + return sum(t.shape[0] for t in self.latent_rows[request_id]) + + @property + def fp8_pool(self) -> bool: + return self.pool.dtype == torch.float8_e4m3fn + + def quantize_e4m3(self, x: torch.Tensor) -> torch.Tensor: + """e4m3(x * kv_scale_orig_quant) — the op's write-side quantization, + and the same one the caller applies to the decode query. One + expression so the mirror, the reference and the buffer the op reads + cannot disagree in the last bit.""" + return (x.float() * (1.0 / self.kv_scale)).to(torch.float8_e4m3fn) + + def to_pool(self, rows: torch.Tensor) -> torch.Tensor: + """The bytes a latent row lands as in the pool: e4m3(row * orig_quant) + for an fp8 pool, the bf16 row itself otherwise.""" + return self.quantize_e4m3(rows) if self.fp8_pool else rows + + def e4m3_roundtrip(self, x: torch.Tensor) -> torch.Tensor: + """fp32 value -> e4m3 at the pool's write scale -> fp32 again: what a + value looks like once it has been through the fp8 pool (or through the + caller-side q quantization, which uses the same scale).""" + return self.quantize_e4m3(x).float() * self.kv_scale + + def fp8_decode_buffers( + self, fused_q: torch.Tensor + ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """The three buffers an fp8-pool MLA generation call requires, filled + the way the production generation-preprocessing step fills them (values + read back from torch.ops.trtllm.mla_rope_generation in a probe, rebuilt + here in plain torch): + + - quant_q_buffer: e4m3(fused_q * kv_scale_orig_quant) — the query the + decode kernel actually reads (`q` itself is not read); + - mla_bmm1_scale: [x, x * log2(e)] with x = softmax_scale * s**2, the + s**2 undoing the 1/s the query and the pool rows were quantized by. + Only element [1] (the log2-domain copy) is read by the kernel; + - mla_bmm2_scale: [s], undoing the 1/s of the V (= pool row) side. + """ + num_gen = fused_q.shape[0] + quant_q = self.quantize_e4m3(fused_q.view(num_gen, self.num_heads, LATENT_DIM)).contiguous() + bmm1 = self.softmax_scale * self.kv_scale * self.kv_scale + return ( + quant_q, + torch.tensor([bmm1, bmm1 * math.log2(math.e)], dtype=torch.float32, device="cuda"), + torch.tensor([self.kv_scale], dtype=torch.float32, device="cuda"), + ) + + def _ensure_pages(self, request_id: int, total_tokens: int) -> None: + tpb = self.tokens_per_block + needed = (total_tokens + tpb - 1) // tpb + pages = self.pages[request_id] + while len(pages) < needed: + pages.append(self._next_free_page) + self._next_free_page += 1 + + def rope_ref(self, x: torch.Tensor, position: int) -> torch.Tensor: + """GPT-J interleaved rotation of the last dim, fp32 math, bf16 result.""" + half = QK_ROPE_HEAD_DIM // 2 + table = self.rotary_cos_sin.view(-1, QK_ROPE_HEAD_DIM, 2) + cos = table[position, :half, 0] + sin = table[position, :half, 1] + pairs = x.float().reshape(*x.shape[:-1], half, 2) + out = torch.empty_like(pairs) + out[..., 0] = pairs[..., 0] * cos - pairs[..., 1] * sin + out[..., 1] = pairs[..., 0] * sin + pairs[..., 1] * cos + return out.reshape(x.shape).to(x.dtype) + + def append_decode_latent(self, request_id: int, row: torch.Tensor) -> None: + """Write one decode token's latent row into the pool, standing in for + the generation-preprocessing append the op itself does not perform.""" + pos = self.cached_len(request_id) + self._ensure_pages(request_id, pos + 1) + tpb = self.tokens_per_block + page = self.pages[request_id][pos // tpb] + self.pool[page, pos % tpb] = self.to_pool(row.unsqueeze(0))[0] + self.latent_rows[request_id].append(row.unsqueeze(0)) + + def _fill_offsets(self, request_ids: List[int]) -> None: + # kv_factor-1 pool: raw block id in both the K and the V row. + for s, rid in enumerate(request_ids): + row = self.block_offsets[0, s] + for j, p in enumerate(self.pages[rid]): + row[0, j] = p + row[1, j] = p + + def _call_op( + self, + q: torch.Tensor, + k: Optional[torch.Tensor], + v: Optional[torch.Tensor], + output: torch.Tensor, + kv_lens: List[int], + ctx_lens: List[int], + num_contexts: int, + num_ctx_tokens: int, + attention_input_type: int, + head_size: int, + num_kv_heads: int, + v_head_dim: int, + latent_cache: Optional[torch.Tensor], + q_pe: Optional[torch.Tensor], + cu_q_seqlens: Optional[torch.Tensor], + cu_kv_seqlens: Optional[torch.Tensor], + fmha_scheduler_counter: Optional[torch.Tensor], + mask_type: int = MASK_CAUSAL, + softmax_stats_tensor: Optional[torch.Tensor] = None, + mla_bmm1_scale: Optional[torch.Tensor] = None, + mla_bmm2_scale: Optional[torch.Tensor] = None, + quant_q_buffer: Optional[torch.Tensor] = None, + predicted_tokens_per_seq: int = 1, + host_past_lens: Optional[List[int]] = None, + ) -> None: + ns = len(kv_lens) + req_types = [0 if i < num_contexts else 1 for i in range(ns)] + thop_attention( + q=q, + k=k, + v=v, + output=output, + output_sf=None, + workspace_=self.workspace, + sequence_length=torch.tensor(kv_lens, dtype=torch.int32, device="cuda"), + host_past_key_value_lengths=torch.tensor( + kv_lens if host_past_lens is None else host_past_lens, + dtype=torch.int32, + ), + host_total_kv_lens=torch.tensor( + [sum(kv_lens[:num_contexts]), sum(kv_lens[num_contexts:])], + dtype=torch.int32, + ), + context_lengths=torch.tensor(ctx_lens, dtype=torch.int32, device="cuda"), + host_context_lengths=torch.tensor(ctx_lens, dtype=torch.int32), + host_request_types=torch.tensor(req_types, dtype=torch.int32), + max_context_q_len_override=None, + kv_cache_block_offsets=self.block_offsets, + host_kv_cache_pool_pointers=self.pool_pointers, + host_kv_cache_pool_mapping=self.pool_mapping, + cache_indirection=None, + kv_scale_orig_quant=self.kv_scale_orig_quant, + kv_scale_quant_orig=self.kv_scale_quant_orig, + out_scale=None, + rotary_inv_freq=self.rotary_inv_freq, + rotary_cos_sin=self.rotary_cos_sin, + latent_cache=latent_cache, + q_pe=q_pe, + block_ids_per_seq=None, + attention_sinks=None, + is_fused_qkv=k is None, + update_kv_cache=True, + predicted_tokens_per_seq=predicted_tokens_per_seq, + local_layer_idx=0, + num_heads=self.num_heads, + num_kv_heads=num_kv_heads, + head_size=head_size, + tokens_per_block=self.tokens_per_block, + max_num_requests=self.max_batch, + max_context_length=self.max_seq_len, + max_seq_len=self.max_seq_len, + attention_window_size=self.max_seq_len, + beam_width=1, + mask_type=mask_type, + quant_mode=self.quant_mode, + q_scaling=self.q_scaling, + position_embedding_type=POSITION_EMBEDDING_TYPE_YARN, + rope_dim=QK_ROPE_HEAD_DIM, + rope_base=self.rope_scalars.rope_base, + rope_scale_type=self.rope_scalars.rope_scale_type, + rope_scale=self.rope_scalars.rope_scale, + rope_short_m_scale=self.rope_scalars.rope_short_m_scale, + rope_long_m_scale=self.rope_scalars.rope_long_m_scale, + rope_max_positions=self.rope_scalars.rope_max_positions, + rope_original_max_positions=self.rope_scalars.rope_original_max_positions, + use_paged_context_fmha=False, + attention_input_type=attention_input_type, + is_mla_enable=True, + chunked_prefill_buffer_batch_size=1, + q_lora_rank=self.q_lora_rank, + kv_lora_rank=KV_LORA_RANK, + qk_nope_head_dim=QK_NOPE_HEAD_DIM, + qk_rope_head_dim=QK_ROPE_HEAD_DIM, + v_head_dim=v_head_dim, + rope_append=True, + mrope_rotary_cos_sin=None, + mrope_position_deltas=None, + helix_position_offsets=None, + helix_is_inactive_rank=None, + attention_chunk_size=None, + softmax_stats_tensor=softmax_stats_tensor, + is_spec_decoding_enabled=False, + use_spec_decoding=False, + is_spec_dec_tree=False, + spec_decoding_generation_lengths=None, + spec_decoding_position_offsets_for_cpp=None, + spec_decoding_packed_mask=None, + spec_decoding_bl_tree_mask_offset=None, + spec_decoding_bl_tree_mask=None, + spec_bl_tree_first_sparse_mask_offset_kv=None, + sparse_kv_indices=None, + sparse_kv_offsets=None, + sparse_attn_indices=None, + sparse_attn_offsets=None, + sparse_attn_indices_block_size=0, + cu_q_seqlens=cu_q_seqlens, + cu_kv_seqlens=cu_kv_seqlens, + fmha_scheduler_counter=fmha_scheduler_counter, + mla_bmm1_scale=mla_bmm1_scale, + mla_bmm2_scale=mla_bmm2_scale, + quant_q_buffer=quant_q_buffer, + num_contexts=num_contexts, + num_ctx_tokens=num_ctx_tokens, + ) + torch.cuda.synchronize() + + def call_context( + self, + request_ids: List[int], + seq_lens: List[int], + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + latent_cache: torch.Tensor, + gen_request_ids: Optional[List[int]] = None, + ) -> torch.Tensor: + """context_only call over fresh context sequences (positions 0..len-1). + + gen_request_ids, if given, are generation-phase sequences sharing the + batch metadata (their q rows belong to a separate generation_only + call). Records the expected latent rows [ckv | rope(k_pe)] per + context request. + """ + gen_request_ids = gen_request_ids or [] + num_contexts = len(request_ids) + num_ctx_tokens = sum(seq_lens) + for rid, ln in zip(request_ids, seq_lens): + assert self.cached_len(rid) == 0, "certified MLA context starts empty" + self._ensure_pages(rid, ln) + all_ids = request_ids + gen_request_ids + self._fill_offsets(all_ids) + kv_lens = list(seq_lens) + [self.cached_len(r) for r in gen_request_ids] + ctx_lens = list(seq_lens) + [self.prompt_lens[r] for r in gen_request_ids] + + output = torch.empty( + num_ctx_tokens, + self.num_heads * V_HEAD_DIM, + dtype=q.dtype, + device=q.device, + ) + self._call_op( + q=q, + k=k, + v=v, + output=output, + kv_lens=kv_lens, + ctx_lens=ctx_lens, + num_contexts=num_contexts, + num_ctx_tokens=num_ctx_tokens, + attention_input_type=1, # context_only + head_size=QK_HEAD_DIM, + num_kv_heads=self.num_heads, # context runs as MHA + v_head_dim=V_HEAD_DIM, + latent_cache=latent_cache, + q_pe=None, + cu_q_seqlens=None, + cu_kv_seqlens=None, + fmha_scheduler_counter=None, + ) + + # Record the expected appended rows: [ckv | rope_pos(k_pe)]. + start = 0 + for rid, ln in zip(request_ids, seq_lens): + rows = latent_cache[start : start + ln].clone() + for i in range(ln): + rows[i, KV_LORA_RANK:] = self.rope_ref(rows[i, KV_LORA_RANK:], i) + self.latent_rows[rid].append(rows) + start += ln + return output + + def reserve_cache_pages(self, request_id: int, total_tokens: int) -> None: + """Pre-allocate page capacity for a request's full [cached + new] KV + range. The latent_cache=None context calls never touch the pool, but + production always runs them with the pages already allocated.""" + self._ensure_pages(request_id, total_tokens) + + def call_context_no_append( + self, + request_ids: List[int], + new_lens: List[int], + pass_kv_lens: List[int], + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + mask_type: int = MASK_CAUSAL, + softmax_stats_tensor: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + """context_only call with latent_cache=None (cached-KV / chunked + context): no in-kernel RoPE and no cache append — q arrives + pre-rotated, K/V arrive explicitly with pass_kv_lens[s] rows per + sequence. sequence_length carries pass_kv_lens while context_lengths + stays at the new-token (q) counts, so KV and q lengths differ.""" + num_contexts = len(request_ids) + num_ctx_tokens = sum(new_lens) + self._fill_offsets(request_ids) + output = torch.empty( + num_ctx_tokens, + self.num_heads * V_HEAD_DIM, + dtype=q.dtype, + device=q.device, + ) + self._call_op( + q=q, + k=k, + v=v, + output=output, + kv_lens=list(pass_kv_lens), + ctx_lens=list(new_lens), + num_contexts=num_contexts, + num_ctx_tokens=num_ctx_tokens, + attention_input_type=1, # context_only + head_size=QK_HEAD_DIM, + num_kv_heads=self.num_heads, # context runs as MHA + v_head_dim=V_HEAD_DIM, + latent_cache=None, # selects the no-RoPE / no-append context path + q_pe=None, + cu_q_seqlens=None, + cu_kv_seqlens=None, + fmha_scheduler_counter=None, + mask_type=mask_type, + softmax_stats_tensor=softmax_stats_tensor, + ) + return output + + def call_generation( + self, + request_ids: List[int], + fused_q: torch.Tensor, + ctx_request_ids: Optional[List[int]] = None, + ctx_seq_lens: Optional[List[int]] = None, + fp8_buffers: Optional[ + Tuple[Optional[torch.Tensor], Optional[torch.Tensor], Optional[torch.Tensor]] + ] = None, + predicted_tokens_per_seq: int = 1, + mask_type: int = MASK_CAUSAL, + cu_q_override: Optional[torch.Tensor] = None, + cu_kv_override: Optional[torch.Tensor] = None, + host_past_override: Optional[List[int]] = None, + ctx_lens_override: Optional[List[int]] = None, + output_buffer: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + """generation_only call: latent MQA over the paged cache. + + The new tokens' latent rows must already be in the pool + (append_decode_latent). latent_cache/q_pe are presence-validated by + the op but their data is not consumed: filled with garbage here on + purpose. ctx_request_ids/ctx_seq_lens describe leading context-phase + sequences sharing the batch metadata (mixed batch, separate call). + + predicted_tokens_per_seq = P gives each generation sequence P query + rows instead of one (speculative decoding: a draft chain verified in + one step). The rows are token-major within a sequence, so sequence g + owns rows [g*P, (g+1)*P), and every per-row tensor is P times taller. + + Over an fp8 pool the call additionally needs the + (quant_q_buffer, mla_bmm1_scale, mla_bmm2_scale) triple, built by + fp8_decode_buffers() from fused_q unless fp8_buffers overrides it — + which is how the probes that pin each buffer's role feed deliberately + wrong or missing ones. The *_override arguments feed deliberately + wrong batch state the same way. + """ + ctx_request_ids = ctx_request_ids or [] + ctx_seq_lens = ctx_seq_lens or [] + num_contexts = len(ctx_request_ids) + num_gen = len(request_ids) + p = predicted_tokens_per_seq + rows = num_gen * p + all_ids = ctx_request_ids + request_ids + self._fill_offsets(all_ids) + gen_kv_lens = [self.cached_len(r) for r in request_ids] + kv_lens = list(ctx_seq_lens) + gen_kv_lens + ctx_lens = list(ctx_seq_lens) + [self.prompt_lens[r] for r in request_ids] + if ctx_lens_override is not None: + ctx_lens = list(ctx_lens_override) + + # Decode-FMHA scheduler buffers, filled as the generation-phase + # preprocessing op fills them: q rows / kv tokens cumulated over + # generation sequences only, counter zeroed. At P > 1 a generation + # sequence contributes num_heads * P q rows. + cu_q = torch.arange(num_gen + 1, dtype=torch.int32) * (self.num_heads * p) + cu_kv = torch.zeros(num_gen + 1, dtype=torch.int32) + cu_kv[1:] = torch.tensor(gen_kv_lens, dtype=torch.int32).cumsum(0) + if cu_q_override is not None: + cu_q = cu_q_override + if cu_kv_override is not None: + cu_kv = cu_kv_override + counter = torch.zeros(1, dtype=torch.uint32, device="cuda") + + garbage_latent = torch.randn(rows, LATENT_DIM, dtype=fused_q.dtype, device="cuda") + garbage_q_pe = torch.randn( + rows, + self.num_heads, + QK_ROPE_HEAD_DIM, + dtype=fused_q.dtype, + device="cuda", + ) + output = ( + torch.empty( + rows, + self.num_heads * KV_LORA_RANK, + dtype=fused_q.dtype, + device="cuda", + ) + if output_buffer is None + else output_buffer + ) + if fp8_buffers is None and self.fp8_pool: + fp8_buffers = self.fp8_decode_buffers(fused_q) + quant_q, bmm1, bmm2 = fp8_buffers or (None, None, None) + self._call_op( + q=fused_q, + k=None, + v=None, + output=output, + kv_lens=kv_lens, + ctx_lens=ctx_lens, + num_contexts=num_contexts, + num_ctx_tokens=sum(ctx_seq_lens), + attention_input_type=2, # generation_only + head_size=LATENT_DIM, + num_kv_heads=1, # latent MQA + v_head_dim=KV_LORA_RANK, + latent_cache=garbage_latent, + q_pe=garbage_q_pe, + cu_q_seqlens=cu_q.cuda(), + cu_kv_seqlens=cu_kv.cuda(), + fmha_scheduler_counter=counter, + mask_type=mask_type, + mla_bmm1_scale=bmm1, + mla_bmm2_scale=bmm2, + quant_q_buffer=quant_q, + predicted_tokens_per_seq=p, + host_past_lens=host_past_override, + ) + return output + + def context_reference( + self, + seq_lens: List[int], + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + latent_cache: torch.Tensor, + softmax_scale: Optional[float] = None, + e4m3_inputs: bool = False, + kv_scale_factors: bool = False, + ) -> torch.Tensor: + """fp32 causal MHA over per-position-roped q/k, headDimV=V_HEAD_DIM. + + Must be given the pre-call q/k/latent contents (the op clobbers the + q/k rope slices in place). softmax_scale overrides the env's own + scale, which is what lets a q_scaling sweep build rival references. + + e4m3_inputs rounds q, k and v through e4m3 at scale 1.0, which is what + the context FMHA quantizes them to under quant_mode 128 (the pool plays + no part in the context math; it is written, not read). kv_scale_factors + adds the two factors that path applies on top — softmax scale * s**2 + and output * s — which cancel only at s = 1.0. + """ + scale = self.softmax_scale if softmax_scale is None else softmax_scale + out_scale = 1.0 + if kv_scale_factors: + scale *= self.kv_scale * self.kv_scale + out_scale = self.kv_scale + outs = [] + start = 0 + h = self.num_heads + for ln in seq_lens: + q_bf = q[start : start + ln].view(ln, h, QK_HEAD_DIM) + q_seq = q_bf.float() + k_seq = k[start : start + ln].view(ln, h, QK_HEAD_DIM).float() + v_seq = v[start : start + ln].view(ln, h, V_HEAD_DIM).float() + latent = latent_cache[start : start + ln] + for i in range(ln): + q_seq[i, :, QK_NOPE_HEAD_DIM:] = self.rope_ref( + q_bf[i, :, QK_NOPE_HEAD_DIM:], i + ).float() + # k_pe comes from latent_cache, roped and broadcast per head. + k_seq[i, :, QK_NOPE_HEAD_DIM:] = self.rope_ref(latent[i, KV_LORA_RANK:], i).float() + if e4m3_inputs: + q_seq = q_seq.to(torch.float8_e4m3fn).float() + k_seq = k_seq.to(torch.float8_e4m3fn).float() + v_seq = v_seq.to(torch.float8_e4m3fn).float() + scores = torch.einsum("ihd,jhd->hij", q_seq, k_seq) * scale + mask = torch.tril(torch.ones(ln, ln, dtype=torch.bool, device="cuda")) + probs = torch.softmax(scores.masked_fill(~mask, float("-inf")), dim=-1) + outs.append( + (torch.einsum("hij,jhd->ihd", probs, v_seq) * out_scale) + .reshape(ln, -1) + .to(torch.bfloat16) + ) + start += ln + return torch.cat(outs) + + def generation_reference( + self, + request_ids: List[int], + fused_q: torch.Tensor, + softmax_scale: Optional[float] = None, + predicted_tokens_per_seq: int = 1, + full_mask: bool = False, + ) -> torch.Tensor: + """fp32 latent MQA over each sequence's mirrored cache rows: + K = rows [L, C+R], V = K[:, :C], scale 1/(q_scaling*sqrt(nope+rope)). + + At P = predicted_tokens_per_seq > 1 a sequence contributes P query + rows, token-major, and row t is the draft token at absolute position + L - P + t, so it attends to keys [0, L - P + t] — the bottom-right + aligned causal mask measured in the MTP section below. full_mask=True + builds the rival model instead (every row sees all L cached rows, + including the sibling draft tokens that do not exist yet at row t). + + Over an fp8 pool both operands go through the e4m3 round trip the op + sees: K rows because that is how they were written, and q because the + caller quantizes it into quant_q_buffer by the same scale (the decode + kernel reads no bf16 query at all). The two kv-scale factors the caller + folds into mla_bmm1_scale / mla_bmm2_scale cancel the 1/s exactly, so + the scale here stays the plain softmax scale. + """ + scale = self.softmax_scale if softmax_scale is None else softmax_scale + p = predicted_tokens_per_seq + outs = [] + for g, rid in enumerate(request_ids): + k_all = torch.cat(self.latent_rows[rid]).float() # [L, C+R] + total = k_all.shape[0] + if self.fp8_pool: + k_all = self.e4m3_roundtrip(k_all) + for t in range(p): + k_seq = k_all if full_mask else k_all[: total - p + t + 1] + q_g = fused_q[g * p + t].view(self.num_heads, LATENT_DIM).float() + if self.fp8_pool: + q_g = self.e4m3_roundtrip(q_g) + probs = torch.softmax(q_g @ k_seq.T * scale, dim=-1) + outs.append((probs @ k_seq[:, :KV_LORA_RANK]).to(torch.bfloat16)) + return torch.stack(outs).view(len(request_ids) * p, -1) + + def check_cache(self, request_id: int) -> None: + """The pool must hold the mirrored latent rows, page by page. + + compressed_kv is a dtype-preserving copy and is gated bitwise per + page; the roped k_pe half is gated by _assert_latent_rope over the + request's whole row range. Over an fp8 pool the mirror is the same + rows put through e4m3(row * kv_scale_orig_quant) and both halves are + compared on their raw bytes. + """ + expected = self.to_pool(torch.cat(self.latent_rows[request_id])) + total = expected.shape[0] + tpb = self.tokens_per_block + pages_read = [] + for j, p in enumerate(self.pages[request_id]): + n = min(tpb, total - j * tpb) + if n <= 0: + break + page = self.pool[p, :n] + ref = expected[j * tpb : j * tpb + n] + if self.fp8_pool: + same = torch.equal( + page[:, :KV_LORA_RANK].view(torch.uint8), + ref[:, :KV_LORA_RANK].view(torch.uint8), + ) + else: + same = torch.equal(page[:, :KV_LORA_RANK], ref[:, :KV_LORA_RANK]) + assert same, f"latent cache compressed_kv mismatch: request {request_id}, page {j}" + pages_read.append(page) + gathered = torch.cat(pages_read) + assert gathered.shape[0] == total + if self.fp8_pool: + _assert_latent_rope_e4m3(gathered, expected, request_id) + else: + _assert_latent_rope(gathered, expected, request_id) + + def check_unwritten_pool_zero(self, request_ids: List[int]) -> None: + """Every page outside the given requests' page sets is still all-zero. + + This is what pins the slab geometry: the op sizes a page from + tokens_per_block, C+R and quant_mode alone, so an element-width or + page-stride mistake writes into pages nobody reserved. + """ + used = set() + for rid in request_ids: + used.update(self.pages[rid]) + rest = [p for p in range(self.pool.shape[0]) if p not in used] + assert bool((self.pool[rest].float() == 0).all()), ( + "the append wrote outside the requests' pages" + ) + + +def _random_context_inputs( + num_tokens: int, + num_heads: int, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Production-shaped MLA context inputs: contiguous q and k (the k rope + slice deliberately garbage — the op fills it from latent_cache), v a + strided split view of a packed kv_b_proj-style buffer, latent + [ckv | k_pe].""" + q = torch.randn(num_tokens, num_heads * QK_HEAD_DIM, dtype=torch.bfloat16, device="cuda") + kv = torch.randn( + num_tokens, + num_heads * (QK_NOPE_HEAD_DIM + V_HEAD_DIM), + dtype=torch.bfloat16, + device="cuda", + ) + k_nope, v = kv.split([num_heads * QK_NOPE_HEAD_DIM, num_heads * V_HEAD_DIM], dim=-1) + k = torch.empty_like(q).view(num_tokens, num_heads, QK_HEAD_DIM) + k[..., :QK_NOPE_HEAD_DIM] = k_nope.view(num_tokens, num_heads, QK_NOPE_HEAD_DIM) + k[..., QK_NOPE_HEAD_DIM:] = torch.randn( + num_tokens, num_heads, QK_ROPE_HEAD_DIM, dtype=torch.bfloat16, device="cuda" + ) + k = k.view(num_tokens, num_heads * QK_HEAD_DIM) + latent = torch.randn(num_tokens, LATENT_DIM, dtype=torch.bfloat16, device="cuda") + return q, k, v, latent + + +def _mla_context_prefill_case( + num_heads: int, + seed: int, + tokens_per_block: int, + seq_lens: List[int], + q_scaling: float = 1.0, + rope: Optional[RopeParams] = None, + rope_scalars: Optional[_RopeScalars] = None, +) -> None: + """MLA context_only prefill: two fresh sequences, at least one crossing + a page boundary. Verifies the FMHA output against a rope-aware fp32 + reference, the latent-cache append page by page (check_cache), that + latent_cache is read-only, and that only the q/k rope slices are + clobbered.""" + torch.manual_seed(seed) + env = _MlaPagedEnv( + num_heads=num_heads, + max_blocks_per_seq=MLA_MAX_SEQ_LEN // tokens_per_block, + tokens_per_block=tokens_per_block, + q_scaling=q_scaling, + rope=rope, + rope_scalars=rope_scalars, + ) + num_tokens = sum(seq_lens) + for rid, ln in enumerate(seq_lens): + env.add_request(rid, ln) + q, k, v, latent = _random_context_inputs(num_tokens, num_heads) + q_orig, k_orig, latent_orig = q.clone(), k.clone(), latent.clone() + + rids = list(range(len(seq_lens))) + out = env.call_context(rids, seq_lens, q, k, v, latent) + ref = env.context_reference(seq_lens, q_orig, k_orig, v, latent_orig) + torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) + + # The append this case checks must really walk more than one page. + assert max(len(env.pages[rid]) for rid in rids) > 1 + for rid in rids: + env.check_cache(rid) + # latent_cache is an input only. + assert torch.equal(latent, latent_orig) + # The op ropes q_pe/k_pe in place: nope slices intact, rope slices clobbered. + q3 = q.view(num_tokens, num_heads, QK_HEAD_DIM) + k3 = k.view(num_tokens, num_heads, QK_HEAD_DIM) + q3_orig = q_orig.view(num_tokens, num_heads, QK_HEAD_DIM) + k3_orig = k_orig.view(num_tokens, num_heads, QK_HEAD_DIM) + assert torch.equal(q3[..., :QK_NOPE_HEAD_DIM], q3_orig[..., :QK_NOPE_HEAD_DIM]) + assert torch.equal(k3[..., :QK_NOPE_HEAD_DIM], k3_orig[..., :QK_NOPE_HEAD_DIM]) + assert not torch.equal(q3[..., QK_NOPE_HEAD_DIM:], q3_orig[..., QK_NOPE_HEAD_DIM:]) + assert not torch.equal(k3[..., QK_NOPE_HEAD_DIM:], k3_orig[..., QK_NOPE_HEAD_DIM:]) + + +def test_bf16_mla_context_prefill() -> None: + """Fresh-prefill MLA context at 16 query heads, page size 64 (100 > 64 + crosses a boundary).""" + _mla_context_prefill_case(MLA_NUM_HEADS, 5, MLA_TOKENS_PER_BLOCK, [100, 17]) + + +def test_bf16_mla_context_prefill_h32() -> None: + """The same fresh-prefill case at 32 query heads (32/32/192).""" + _mla_context_prefill_case(MLA_NUM_HEADS_H32, 105, MLA_TOKENS_PER_BLOCK, [100, 17]) + + +def test_bf16_mla_context_prefill_h8() -> None: + """The same fresh-prefill case at 8 query heads (8/8/192) — the tep4 + slice of a 32-head checkpoint.""" + _mla_context_prefill_case(MLA_NUM_HEADS_H8, 805, MLA_TOKENS_PER_BLOCK, [100, 17]) + + +def test_bf16_mla_page32_context_prefill_h32() -> None: + """Fresh-prefill MLA context at page size 32 (the engine default), 32 + query heads. 96 fills three 32-token pages exactly — a 64-token page + never ends there — and 33 crosses into a second page by one token, so + the append's page/slot arithmetic is exercised at both a page-aligned + end and a one-token spill.""" + _mla_context_prefill_case(MLA_NUM_HEADS_H32, 205, MLA_PAGE32, [96, 33]) + + +def test_bf16_mla_page32_context_prefill_h8() -> None: + """The same page-32 fresh-prefill case at 8 query heads.""" + _mla_context_prefill_case(MLA_NUM_HEADS_H8, 815, MLA_PAGE32, [96, 33]) + + +def test_bf16_mla_page32_context_prefill_h128() -> None: + """The same page-32 fresh-prefill case at 128 query heads (128/128/192) — + the DeepSeek-R1-0528 layer shape, which attention DP replicates whole onto + every rank rather than slicing. Baseline rope/scale so this case isolates + the head count; the R1 cell adds the other two axes further down.""" + _mla_context_prefill_case(MLA_NUM_HEADS_H128, 905, MLA_PAGE32, [96, 33]) + + +def _mla_generation_decode_case( + num_heads: int, + seed: int, + tokens_per_block: int, + prefill_lens: List[int], + q_scaling: float = 1.0, + rope: Optional[RopeParams] = None, + rope_scalars: Optional[_RopeScalars] = None, +) -> None: + """MLA generation_only decode over cache written by the context call. + + The first prefill length is an exact multiple of the page size, so the + first decode token lands in a fresh page. Two decode steps; each step + the test appends the new latent row (the op does not append in + generation) and checks the FMHA output against an fp32 latent-MQA + reference. The garbage latent_cache/q_pe arguments plus cache/fused_q + invariance pin the fusion boundary: the generation call only reads the + paged pool. + """ + torch.manual_seed(seed) + env = _MlaPagedEnv( + num_heads=num_heads, + max_blocks_per_seq=MLA_MAX_SEQ_LEN // tokens_per_block, + tokens_per_block=tokens_per_block, + q_scaling=q_scaling, + rope=rope, + rope_scalars=rope_scalars, + ) + assert prefill_lens[0] % tokens_per_block == 0 + rids = list(range(len(prefill_lens))) + for rid, ln in zip(rids, prefill_lens): + env.add_request(rid, ln) + q, k, v, latent = _random_context_inputs(sum(prefill_lens), num_heads) + env.call_context(rids, prefill_lens, q, k, v, latent) + for rid in rids: + env.check_cache(rid) + pages_after_prefill = {rid: len(env.pages[rid]) for rid in rids} + + for _ in range(2): + for rid in rids: + env.append_decode_latent( + rid, torch.randn(LATENT_DIM, dtype=torch.bfloat16, device="cuda") + ) + fused_q = torch.randn( + len(rids), num_heads * LATENT_DIM, dtype=torch.bfloat16, device="cuda" + ) + fused_q_orig = fused_q.clone() + pool_before = env.pool.clone() + out = env.call_generation(rids, fused_q) + ref = env.generation_reference(rids, fused_q_orig) + torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) + assert torch.equal(fused_q, fused_q_orig) + assert torch.equal(env.pool, pool_before) # generation never writes + # The decode steps must really have opened a new page on some sequence, + # i.e. the read range crossed a page boundary the prefill had not. + assert any(len(env.pages[rid]) > pages_after_prefill[rid] for rid in rids) + + +def test_bf16_mla_generation_decode() -> None: + """Latent-MQA MLA decode at 16 query heads, page size 64.""" + _mla_generation_decode_case(MLA_NUM_HEADS, 6, MLA_TOKENS_PER_BLOCK, [64, 30]) + + +def test_bf16_mla_generation_decode_h32() -> None: + """The same decode case at 32 query heads (32/1/576).""" + _mla_generation_decode_case(MLA_NUM_HEADS_H32, 106, MLA_TOKENS_PER_BLOCK, [64, 30]) + + +def test_bf16_mla_generation_decode_h8() -> None: + """The same decode case at 8 query heads (8/1/576). Unlike 16 and 32, + which share the ...VarSeqQ16... decode kernel, 8 heads map to a + ...VarSeqQ8... one — a differently named compiled variant, so this is the + flavor where the head count actually changes the kernel.""" + _mla_generation_decode_case(MLA_NUM_HEADS_H8, 806, MLA_TOKENS_PER_BLOCK, [64, 30]) + + +def test_bf16_mla_page32_generation_decode_h32() -> None: + """Latent-MQA MLA decode at page size 32 (the engine default), 32 query + heads — the decode path is where the page size selects a different + compiled trtllm-gen kernel (...PagedKvDenseP32... rather than ...P64...). + Sequence 0's 64-token prefill fills two pages exactly, so its first + decode token opens page 2; sequence 1's 31-token prefill leaves its + first decode token on page 0's last slot and its second one opens page + 1, so a decode read range crosses a boundary mid-case.""" + _mla_generation_decode_case(MLA_NUM_HEADS_H32, 206, MLA_PAGE32, [64, 31]) + + +def test_bf16_mla_page32_generation_decode_h16() -> None: + """The same page-32 decode case at 16 query heads. The decode kernel is + JIT-compiled once per head count even though its name does not carry the + count, so the page-32 variant is exercised at both certified counts.""" + _mla_generation_decode_case(MLA_NUM_HEADS, 216, MLA_PAGE32, [64, 31]) + + +def test_bf16_mla_page32_generation_decode_h8() -> None: + """The same page-32 decode case at 8 query heads. Both axes that reach + the compiled decode kernel move here at once: the page size is in the + kernel name (...P32...) and 8 heads take the ...VarSeqQ8... q-tile.""" + _mla_generation_decode_case(MLA_NUM_HEADS_H8, 816, MLA_PAGE32, [64, 31]) + + +def test_bf16_mla_page32_generation_decode_h128() -> None: + """The same page-32 decode case at 128 query heads (128/1/576). Decode is + the phase where the head count reaches kernel selection: 128 reports the + same ...P32VarSeqQ16Kv128... name 16 and 32 do, yet pays its own compile + (the cache is keyed more finely than the name).""" + _mla_generation_decode_case(MLA_NUM_HEADS_H128, 906, MLA_PAGE32, [64, 31]) + + +def _mla_mixed_batch_case( + num_heads: int, + seed: int, + tokens_per_block: int, + first_len: int, + second_len: int, + q_scaling: float = 1.0, + rope: Optional[RopeParams] = None, + rope_scalars: Optional[_RopeScalars] = None, +) -> None: + """A mixed batch is two calls sharing full-batch metadata: context_only + over the leading context rows, generation_only over the trailing + generation rows (indexed from num_contexts).""" + torch.manual_seed(seed) + env = _MlaPagedEnv( + num_heads=num_heads, + max_blocks_per_seq=MLA_MAX_SEQ_LEN // tokens_per_block, + tokens_per_block=tokens_per_block, + q_scaling=q_scaling, + rope=rope, + rope_scalars=rope_scalars, + ) + # Prefill request 0 alone, then decode it alongside a new context request. + env.add_request(0, first_len) + q, k, v, latent = _random_context_inputs(first_len, num_heads) + env.call_context([0], [first_len], q, k, v, latent) + + env.add_request(1, second_len) + env.append_decode_latent(0, torch.randn(LATENT_DIM, dtype=torch.bfloat16, device="cuda")) + q, k, v, latent = _random_context_inputs(second_len, num_heads) + q_orig, k_orig, latent_orig = q.clone(), k.clone(), latent.clone() + out_ctx = env.call_context([1], [second_len], q, k, v, latent, gen_request_ids=[0]) + ref_ctx = env.context_reference([second_len], q_orig, k_orig, v, latent_orig) + torch.testing.assert_close(out_ctx, ref_ctx, rtol=RTOL, atol=ATOL) + + fused_q = torch.randn(1, num_heads * LATENT_DIM, dtype=torch.bfloat16, device="cuda") + out_gen = env.call_generation([0], fused_q, ctx_request_ids=[1], ctx_seq_lens=[second_len]) + ref_gen = env.generation_reference([0], fused_q) + torch.testing.assert_close(out_gen, ref_gen, rtol=RTOL, atol=ATOL) + env.check_cache(0) + env.check_cache(1) + + +def test_bf16_mla_mixed_batch_split_calls() -> None: + """Mixed-batch MLA (two phase calls) at 16 query heads, page size 64.""" + _mla_mixed_batch_case(MLA_NUM_HEADS, 7, MLA_TOKENS_PER_BLOCK, 40, 23) + + +def test_bf16_mla_mixed_batch_split_calls_h32() -> None: + """The same mixed-batch case at 32 query heads.""" + _mla_mixed_batch_case(MLA_NUM_HEADS_H32, 107, MLA_TOKENS_PER_BLOCK, 40, 23) + + +def test_bf16_mla_mixed_batch_split_calls_h8() -> None: + """The same mixed-batch case at 8 query heads: the two phase calls of one + batch run at 8/8/192 and 8/1/576 off the same full-batch state tensors.""" + _mla_mixed_batch_case(MLA_NUM_HEADS_H8, 807, MLA_TOKENS_PER_BLOCK, 40, 23) + + +def test_bf16_mla_page32_mixed_batch_split_calls_h32() -> None: + """Mixed-batch MLA at page size 32, 32 query heads: the decoding + sequence's 64-token history fills two pages exactly, so its appended + token opens page 2 and the decode call reads across three pages, while + the context sequence sharing the batch spans two.""" + _mla_mixed_batch_case(MLA_NUM_HEADS_H32, 207, MLA_PAGE32, 64, 33) + + +def test_bf16_mla_page32_mixed_batch_split_calls_h8() -> None: + """The same page-32 mixed-batch case at 8 query heads.""" + _mla_mixed_batch_case(MLA_NUM_HEADS_H8, 817, MLA_PAGE32, 64, 33) + + +def test_bf16_mla_page32_mixed_batch_split_calls_h128() -> None: + """The same page-32 mixed-batch case at 128 query heads: the two phase + calls run at 128/128/192 and 128/1/576 off one set of batch tensors.""" + _mla_mixed_batch_case(MLA_NUM_HEADS_H128, 907, MLA_PAGE32, 64, 33) + + +# ─── MLA context with latent_cache=None (cached KV / chunked prefill) ─── + + +def _random_explicit_kv(num_tokens: int, num_heads: int) -> Tuple[torch.Tensor, torch.Tensor]: + """Production-shaped explicit K/V sources for the latent_cache=None + context calls: a contiguous K and a packed [T, H*(nope+v)] + kv_b_proj-style buffer whose _v_split_view is the V argument.""" + k = torch.randn(num_tokens, num_heads * QK_HEAD_DIM, dtype=torch.bfloat16, device="cuda") + packed_kv = torch.randn( + num_tokens, + num_heads * (QK_NOPE_HEAD_DIM + V_HEAD_DIM), + dtype=torch.bfloat16, + device="cuda", + ) + return k, packed_kv + + +def _v_split_view(packed_kv: torch.Tensor, num_heads: int) -> torch.Tensor: + """The [.., H*nope:] split view of a packed [T, H*(nope+v)] buffer. The + context FMHA hard-codes V's row stride to H*(nope+v_head_dim) elements, + so V must keep this stride — a contiguous [T, H*v] V is misread.""" + return packed_kv.split([num_heads * QK_NOPE_HEAD_DIM, num_heads * V_HEAD_DIM], dim=-1)[1] + + +def _explicit_kv_reference( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + q_lens: List[int], + kv_lens: List[int], + mask_type: int, + num_heads: int, + softmax_scale: float = MLA_SOFTMAX_SCALE, + e4m3_inputs: bool = False, + out_scale: float = 1.0, +) -> Tuple[torch.Tensor, torch.Tensor]: + """fp32 attention over explicit per-sequence K/V ranges plus softmax + stats: per (token, head), m = row max of the scaled logits and + sigma = sum(exp(logit - m)) over this pass's KV range. Causal is + bottom-right aligned (query i of a sequence sits at absolute position + kv_s - q_s + i). Rows of zero-KV sequences stay NaN — the op leaves them + undefined and downstream merge plans skip them. Returns + (bf16 output [Tq, H*v], fp32 stats [Tq, H, 2]). + + e4m3_inputs rounds q, k and v through e4m3 at scale 1.0, which is what this + flavor quantizes them to under quant_mode 128; out_scale multiplies the + output, the second of the two dequantization factors that path applies (the + first, s**2, rides in softmax_scale). Both cancel only at s = 1.0. The + stats stay in the un-scaled domain the op writes them in and are not + certified on the fp8 path.""" + total_q = sum(q_lens) + out = torch.full((total_q, num_heads * V_HEAD_DIM), float("nan"), device="cuda") + stats = torch.full((total_q, num_heads, 2), float("nan"), device="cuda") + q_start = kv_start = 0 + for q_len, kv_len in zip(q_lens, kv_lens): + if kv_len == 0: + q_start += q_len + continue + q_seq = q[q_start : q_start + q_len].view(q_len, num_heads, QK_HEAD_DIM).float() + k_seq = k[kv_start : kv_start + kv_len].view(kv_len, num_heads, QK_HEAD_DIM).float() + v_seq = v[kv_start : kv_start + kv_len].view(kv_len, num_heads, V_HEAD_DIM).float() + if e4m3_inputs: + q_seq = q_seq.to(torch.float8_e4m3fn).float() + k_seq = k_seq.to(torch.float8_e4m3fn).float() + v_seq = v_seq.to(torch.float8_e4m3fn).float() + scores = torch.einsum("ihd,jhd->hij", q_seq, k_seq) * softmax_scale + if mask_type == MASK_CAUSAL: + offset = kv_len - q_len + mask = torch.zeros(q_len, kv_len, dtype=torch.bool, device="cuda") + for i in range(q_len): + mask[i, : offset + i + 1] = True + scores = scores.masked_fill(~mask, float("-inf")) + row_max = scores.max(dim=-1).values # [H, q_len] + row_sum = torch.exp(scores - row_max.unsqueeze(-1)).sum(dim=-1) + probs = torch.softmax(scores, dim=-1) + out[q_start : q_start + q_len] = ( + torch.einsum("hij,jhd->ihd", probs, v_seq).reshape(q_len, -1) * out_scale + ) + stats[q_start : q_start + q_len, :, 0] = row_max.transpose(0, 1) + stats[q_start : q_start + q_len, :, 1] = row_sum.transpose(0, 1) + q_start += q_len + kv_start += kv_len + return out.to(torch.bfloat16), stats + + +def _assert_valid_rows_close( + actual: torch.Tensor, ref: torch.Tensor, ref_stats: torch.Tensor +) -> None: + """Compare only rows of sequences that had KV in this pass (non-NaN in + the reference); zero-KV rows are undefined op output.""" + valid = ~torch.isnan(ref_stats[:, 0, 0]) + torch.testing.assert_close(actual[valid], ref[valid], rtol=RTOL, atol=ATOL) + + +def _assert_stats_close(actual: torch.Tensor, ref_stats: torch.Tensor) -> None: + valid = ~torch.isnan(ref_stats[:, 0, 0]) + # Both sides are fp32 reductions over the same bf16-rounded q/k, so they + # differ only by accumulation order: observed max abs err 2e-6 (max stat) + # and rel err 1.2e-6 (sum stat) on sm_100, on stat magnitudes ~1-40. + # 1e-4 gives ~50x margin while still catching wrong-domain (log2) or + # unscaled-logit stats outright. + torch.testing.assert_close(actual[valid], ref_stats[valid], rtol=1e-4, atol=1e-4) + + +def _mla_context_cached_kv_case( + num_heads: int, + seed: int, + tokens_per_block: int, + cached_lens: List[int], + new_lens: List[int], + q_scaling: float = 1.0, + rope: Optional[RopeParams] = None, + rope_scalars: Optional[_RopeScalars] = None, +) -> None: + """MLA context over a cached KV prefix: latent_cache=None, q pre-rotated + upstream, K/V supplied for the full [cached + new] range so KV length + exceeds q length. Causal masking is bottom-right aligned. Verifies the + output against an fp32 reference and that the call mutates nothing but + output: q, k, v, and the paged pool stay bitwise intact (no in-kernel + RoPE, no append).""" + torch.manual_seed(seed) + env = _MlaPagedEnv( + num_heads=num_heads, + max_blocks_per_seq=MLA_MAX_SEQ_LEN // tokens_per_block, + tokens_per_block=tokens_per_block, + q_scaling=q_scaling, + rope=rope, + rope_scalars=rope_scalars, + ) + kv_lens = [c + n for c, n in zip(cached_lens, new_lens)] + for rid, total in enumerate(kv_lens): + env.add_request(rid, new_lens[rid]) + env.reserve_cache_pages(rid, total) + env.pool.normal_() # op must not read or write the pool in this mode + pool_before = env.pool.clone() + + q = torch.randn(sum(new_lens), num_heads * QK_HEAD_DIM, dtype=torch.bfloat16, device="cuda") + k, packed_kv = _random_explicit_kv(sum(kv_lens), num_heads) + v = _v_split_view(packed_kv, num_heads) + q_orig, k_orig, v_orig = q.clone(), k.clone(), v.clone() + + out = env.call_context_no_append([0, 1, 2], new_lens, kv_lens, q, k, v) + ref, _ = _explicit_kv_reference( + q, k, v, new_lens, kv_lens, MASK_CAUSAL, num_heads, env.softmax_scale + ) + torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) + + # Fusion boundary: everything is an input; only output rows are written. + assert torch.equal(q, q_orig) + assert torch.equal(k, k_orig) + assert torch.equal(v, v_orig) + assert torch.equal(env.pool, pool_before) + + +def test_bf16_mla_context_cached_kv_no_append() -> None: + """Cached-KV (no-append) MLA context at 16 query heads, page size 64. + Prefixes: page-crossing, mid-page, and empty.""" + _mla_context_cached_kv_case(MLA_NUM_HEADS, 8, MLA_TOKENS_PER_BLOCK, [80, 33, 0], [40, 7, 25]) + + +def test_bf16_mla_context_cached_kv_no_append_h32() -> None: + """The same cached-KV context case at 32 query heads.""" + _mla_context_cached_kv_case( + MLA_NUM_HEADS_H32, 108, MLA_TOKENS_PER_BLOCK, [80, 33, 0], [40, 7, 25] + ) + + +def test_bf16_mla_context_cached_kv_no_append_h8() -> None: + """The same cached-KV context case at 8 query heads.""" + _mla_context_cached_kv_case( + MLA_NUM_HEADS_H8, 808, MLA_TOKENS_PER_BLOCK, [80, 33, 0], [40, 7, 25] + ) + + +def test_bf16_mla_page32_context_cached_kv_no_append_h32() -> None: + """Cached-KV (no-append) MLA context at page size 32, 32 query heads. + The pages are reserved as production does even though this flavor never + touches the pool, so the batch carries a page-32 offsets table: prefixes + of 96 (three exact pages), 31 (one short of a page) and 0, reaching KV + lengths of 128 (four exact pages), 40 and 25.""" + _mla_context_cached_kv_case(MLA_NUM_HEADS_H32, 208, MLA_PAGE32, [96, 31, 0], [32, 9, 25]) + + +def test_bf16_mla_page32_context_cached_kv_no_append_h8() -> None: + """The same page-32 cached-KV context case at 8 query heads.""" + _mla_context_cached_kv_case(MLA_NUM_HEADS_H8, 818, MLA_PAGE32, [96, 31, 0], [32, 9, 25]) + + +def test_bf16_mla_page32_context_cached_kv_no_append_h128() -> None: + """The same page-32 cached-KV context case at 128 query heads. This is the + flavor an engine with block reuse on runs for every context request that + hits a cached prefix, so it carries the head count as much as the fresh + one does.""" + _mla_context_cached_kv_case(MLA_NUM_HEADS_H128, 908, MLA_PAGE32, [96, 31, 0], [32, 9, 25]) + + +def test_bf16_mla_context_chunked_prefill_with_merge() -> None: + """MLA chunked context: one padding-masked partial pass per cached-KV + chunk with softmax_stats_tensor emitted, a final causal pass over the + new tokens, each folded into the running output by the downstream + trtllm merge op (the production consumer of the emitted stats — the + reference is still built from torch alone). Every pass's output and + stats are checked against fp32 partial-attention references, and the + fully merged output/stats against a single-pass full-range reference.""" + torch.manual_seed(9) + env = _MlaPagedEnv() + cached_lens = [96, 48, 0] + new_lens = [32, 17, 23] + kv_lens = [c + n for c, n in zip(cached_lens, new_lens)] + num_ctx_tokens = sum(new_lens) + for rid, total in enumerate(kv_lens): + env.add_request(rid, new_lens[rid]) + env.reserve_cache_pages(rid, total) + + # Per-sequence full KV timeline; chunk passes slice the cached region. + # V buffers are sliced and concatenated as packed rows so every pass's V + # keeps the required H*(nope+v) row stride. + k_full = [] + packed_full = [] + for total in kv_lens: + k_seq, packed_seq = _random_explicit_kv(total, MLA_NUM_HEADS) + k_full.append(k_seq) + packed_full.append(packed_seq) + q = torch.randn( + num_ctx_tokens, MLA_NUM_HEADS * QK_HEAD_DIM, dtype=torch.bfloat16, device="cuda" + ) + + # Production chunk plan (greedy, 64-token buffer over the cached + # regions 96/48/0): loop 0 takes 64 of seq 0; loop 1 the remaining 32 of + # seq 0 plus 32 of seq 1; loop 2 the remaining 16 of seq 1. Merge ops: + # 2 = copy on a sequence's first pass, 1 = merge, 0 = skip; the final + # new-token pass merges (or copies for the cache-less seq 2). + chunk_lens = [[64, 0, 0], [32, 32, 0], [0, 16, 0]] + chunk_offsets = [[0, 0, 0], [64, 0, 0], [96, 32, 0]] + merge_ops = [[2, 0, 0], [1, 2, 0], [0, 1, 0], [1, 1, 2]] + + merged = torch.empty( + num_ctx_tokens, MLA_NUM_HEADS * V_HEAD_DIM, dtype=torch.bfloat16, device="cuda" + ) + merged_stats = torch.empty(num_ctx_tokens, MLA_NUM_HEADS, 2, dtype=torch.float32, device="cuda") + temp_stats = torch.empty_like(merged_stats) + cu_q = torch.tensor( + [0, *torch.tensor(new_lens).cumsum(0).tolist()], + dtype=torch.int64, + device="cuda", + ) + + def run_pass( + pass_kv_lens: List[int], + offsets: List[int], + mask_type: int, + ops: List[int], + ) -> None: + slices = list(enumerate(zip(offsets, pass_kv_lens))) + k_buf = torch.cat([k_full[s][o : o + n] for s, (o, n) in slices]) + v_buf = _v_split_view( + torch.cat([packed_full[s][o : o + n] for s, (o, n) in slices]), + MLA_NUM_HEADS, + ) + temp_out = env.call_context_no_append( + [0, 1, 2], + new_lens, + pass_kv_lens, + q, + k_buf, + v_buf, + mask_type=mask_type, + softmax_stats_tensor=temp_stats, + ) + ref_out, ref_stats = _explicit_kv_reference( + q, k_buf, v_buf, new_lens, pass_kv_lens, mask_type, MLA_NUM_HEADS + ) + _assert_valid_rows_close(temp_out, ref_out, ref_stats) + _assert_stats_close(temp_stats, ref_stats) + # The downstream merge op consumes the emitted stats directly. + torch.ops.trtllm.merge_chunked_attention_for_mla( + merged, + temp_out, + merged_stats, + temp_stats, + len(new_lens), + cu_q, + max(new_lens), + torch.tensor(ops, dtype=torch.int64, device="cuda"), + MLA_NUM_HEADS, + V_HEAD_DIM, + ) + torch.cuda.synchronize() + + for loop_idx, (lens, offs) in enumerate(zip(chunk_lens, chunk_offsets)): + run_pass(lens, offs, MASK_PADDING, merge_ops[loop_idx]) + # Final pass: causal attention of the new tokens over themselves only. + run_pass(new_lens, cached_lens, MASK_CAUSAL, merge_ops[-1]) + + # The merged result must equal single-pass attention over the full + # [cached + new] range (bottom-right-aligned causal). + k_all = torch.cat(k_full) + v_all = _v_split_view(torch.cat(packed_full), MLA_NUM_HEADS) + ref_full, ref_full_stats = _explicit_kv_reference( + q, k_all, v_all, new_lens, kv_lens, MASK_CAUSAL, MLA_NUM_HEADS + ) + torch.testing.assert_close(merged, ref_full, rtol=RTOL, atol=ATOL) + _assert_stats_close(merged_stats, ref_full_stats) + + +# ─── MLA at q_lora_rank = 0 (checkpoints with no q-LoRA) ─── + +# The values swept per MLA call flavor. Index 0 is the reference run and +# index 2 repeats it — the run-to-run determinism control that makes the +# bitwise comparisons meaningful. 4096 is larger than any real q_a_proj rank +# and larger than C, so any arithmetic actually reading the argument would +# move something. +Q_LORA_RANK_SWEEP = [Q_LORA_RANK_DSV3, Q_LORA_RANK_ZERO, Q_LORA_RANK_DSV3, 4096] + +# KV-cache scaling factors swept over the fp8 latent pool: 1.0 is the +# production value (DeepSeek-R1-0528-FP4 declares kv_cache_quant_algo FP8 +# with per-layer k_scale/v_scale both 1.0), 1.5 is not a power of two, so the +# e4m3 grid genuinely depends on it, and 2.0 is. +KV_SCALE_SWEEP = [1.0, 1.5, 2.0] + + +def _mla_page32_env(num_heads: int, q_lora_rank: int) -> _MlaPagedEnv: + """A page-32 MLA env at one head count and one q_lora_rank. Head count and + page size are certified axes of their own; fixing both within a sweep + leaves q_lora_rank the only variable. The two counts swept are the shipped + ones: 32 (deepseek-v3-lite tp1) and 8 (its tep4 slice), which is the count + that takes its own compiled decode kernel (...VarSeqQ8...).""" + return _MlaPagedEnv( + num_heads=num_heads, + max_blocks_per_seq=MLA_MAX_SEQ_LEN // MLA_PAGE32, + tokens_per_block=MLA_PAGE32, + q_lora_rank=q_lora_rank, + ) + + +def _assert_q_lora_rank_inert( + flavor: str, + runs: List[Tuple[int, Dict[str, torch.Tensor], int]], +) -> None: + """Every observable of the reference run — outputs, the whole paged pool, + the in-place-written inputs, and the size the op grew the workspace to — + must come back bitwise identical at every other q_lora_rank.""" + base_rank, base, base_ws = runs[0] + for rank, cur, ws in runs[1:]: + for key, expected in base.items(): + got = cur[key] + assert torch.equal(expected, got), ( + f"{flavor}: {key} differs at q_lora_rank={rank} vs " + f"{base_rank}: {int((expected != got).sum())}/{expected.numel()} " + f"elements, max abs " + f"{(expected.float() - got.float()).abs().max().item():.3e}" + ) + assert ws == base_ws, ( + f"{flavor}: workspace grew to {ws} bytes at q_lora_rank={rank}, " + f"{base_ws} at {base_rank}" + ) + + +def _mla_qlora_context_prefill_case(num_heads: int, seed: int) -> None: + """Fresh-prefill MLA context swept over q_lora_rank (page 32). + + This is the flavor with the most machinery behind the MLA meta params — + in-kernel GPT-J RoPE of q_pe/k_pe plus the paged latent append — so it is + where a q_lora_rank-dependent layout would surface. The zero run is + checked against the fp32 reference and its append against the mirrored + latent rows; the whole sweep is then checked bitwise against the + q_lora_rank=1536 run on identical inputs. + """ + seq_lens = [96, 33] + num_tokens = sum(seq_lens) + h = num_heads + torch.manual_seed(seed) + q_src, k_src, v, latent_src = _random_context_inputs(num_tokens, h) + v_src = v.clone() # v is shared across the sweep: it must stay an input + + runs: List[Tuple[int, Dict[str, torch.Tensor], int]] = [] + for rank in Q_LORA_RANK_SWEEP: + env = _mla_page32_env(h, rank) + for rid, ln in enumerate(seq_lens): + env.add_request(rid, ln) + q, k, latent = q_src.clone(), k_src.clone(), latent_src.clone() + out = env.call_context([0, 1], seq_lens, q, k, v, latent) + + if rank == Q_LORA_RANK_ZERO: + # The surface under test has to be right, not merely reproducible. + ref = env.context_reference(seq_lens, q_src, k_src, v, latent_src) + torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) + assert max(len(env.pages[rid]) for rid in (0, 1)) > 1 + for rid in (0, 1): + env.check_cache(rid) + assert torch.equal(latent, latent_src) # latent_cache is read-only + q3 = q.view(num_tokens, h, QK_HEAD_DIM) + k3 = k.view(num_tokens, h, QK_HEAD_DIM) + q3_src = q_src.view(num_tokens, h, QK_HEAD_DIM) + k3_src = k_src.view(num_tokens, h, QK_HEAD_DIM) + nope = QK_NOPE_HEAD_DIM + assert torch.equal(q3[..., :nope], q3_src[..., :nope]) + assert torch.equal(k3[..., :nope], k3_src[..., :nope]) + assert not torch.equal(q3[..., nope:], q3_src[..., nope:]) + assert not torch.equal(k3[..., nope:], k3_src[..., nope:]) + + runs.append( + ( + rank, + { + "output": out, + "pool": env.pool.clone(), + "q_after_call": q, + "k_after_call": k, + }, + env.workspace.numel(), + ) + ) + assert torch.equal(v, v_src), "the sweep must have run on identical inputs" + _assert_q_lora_rank_inert(f"fresh-prefill context H={h}", runs) + + +def test_bf16_mla_qlora0_context_prefill_h32() -> None: + _mla_qlora_context_prefill_case(MLA_NUM_HEADS_H32, 305) + + +def test_bf16_mla_qlora0_context_prefill_h8() -> None: + _mla_qlora_context_prefill_case(MLA_NUM_HEADS_H8, 315) + + +def _mla_qlora_generation_decode_case(num_heads: int, seed: int) -> None: + """Latent-MQA MLA decode swept over q_lora_rank (page 32). + + The generation phase is the one that JIT-compiles its FMHA kernel, so + this is where q_lora_rank would show up as a kernel-selection axis: the + whole sweep runs in one process against a single compiled decode kernel + (the compile cache is keyed by head count and page size, both fixed + here). Each variant builds its own history through its own context call, + so a rank that changed the append would already separate the pools. + """ + prefill_lens = [64, 31] + h = num_heads + torch.manual_seed(seed) + q_src, k_src, v, latent_src = _random_context_inputs(sum(prefill_lens), h) + v_src = v.clone() + # Decode rows and fused q are drawn once and replayed per variant. + decode_rows = [ + [torch.randn(LATENT_DIM, dtype=torch.bfloat16, device="cuda") for _ in prefill_lens] + for _ in range(2) + ] + fused_qs = [ + torch.randn(len(prefill_lens), h * LATENT_DIM, dtype=torch.bfloat16, device="cuda") + for _ in range(2) + ] + + runs: List[Tuple[int, Dict[str, torch.Tensor], int]] = [] + for rank in Q_LORA_RANK_SWEEP: + env = _mla_page32_env(h, rank) + for rid, ln in enumerate(prefill_lens): + env.add_request(rid, ln) + env.call_context([0, 1], prefill_lens, q_src.clone(), k_src.clone(), v, latent_src.clone()) + observed = {"pool_after_prefill": env.pool.clone()} + for step in range(2): + for rid in range(len(prefill_lens)): + env.append_decode_latent(rid, decode_rows[step][rid]) + # call_generation draws the (unconsumed) latent_cache/q_pe + # garbage from the global RNG: reseed so every variant is fed the + # same garbage and q_lora_rank stays the only difference. + torch.manual_seed(seed * 10 + step) + fused_q = fused_qs[step].clone() + out = env.call_generation([0, 1], fused_q) + if rank == Q_LORA_RANK_ZERO: + ref = env.generation_reference([0, 1], fused_qs[step]) + torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) + assert torch.equal(fused_q, fused_qs[step]) + observed[f"decode_out_{step}"] = out + observed[f"pool_after_decode_{step}"] = env.pool.clone() + runs.append((rank, observed, env.workspace.numel())) + assert torch.equal(v, v_src), "the sweep must have run on identical inputs" + _assert_q_lora_rank_inert(f"generation decode H={h}", runs) + + +def test_bf16_mla_qlora0_generation_decode_h32() -> None: + _mla_qlora_generation_decode_case(MLA_NUM_HEADS_H32, 306) + + +def test_bf16_mla_qlora0_generation_decode_h8() -> None: + """H=8 decode compiles ...VarSeqQ8Kv128... where 16 and 32 both take + ...VarSeqQ16...: a genuinely different kernel, so the H=32 sweep says + nothing about it. This is the tep4 deepseek-v3-lite cell, whose config + carries "q_lora_rank": null.""" + _mla_qlora_generation_decode_case(MLA_NUM_HEADS_H8, 316) + + +def _mla_qlora_context_cached_kv_no_append_case(num_heads: int, seed: int) -> None: + """Cached-KV (no-append) MLA context swept over q_lora_rank (page 32). + + latent_cache=None skips the in-kernel RoPE and the append, so this flavor + reaches a different context kernel from the fresh-prefill one; the paged + pool is pre-filled with noise and must come back bitwise untouched at + every rank. + """ + cached_lens, new_lens = [96, 31, 0], [32, 9, 25] + kv_lens = [c + n for c, n in zip(cached_lens, new_lens)] + h = num_heads + torch.manual_seed(seed) + q_src = torch.randn(sum(new_lens), h * QK_HEAD_DIM, dtype=torch.bfloat16, device="cuda") + k_src, packed_kv = _random_explicit_kv(sum(kv_lens), h) + v = _v_split_view(packed_kv, h) + v_src = v.clone() + pool_fill: Optional[torch.Tensor] = None + + runs: List[Tuple[int, Dict[str, torch.Tensor], int]] = [] + for rank in Q_LORA_RANK_SWEEP: + env = _mla_page32_env(h, rank) + for rid, total in enumerate(kv_lens): + env.add_request(rid, new_lens[rid]) + env.reserve_cache_pages(rid, total) + if pool_fill is None: + pool_fill = torch.randn_like(env.pool) + env.pool.copy_(pool_fill) # the op must neither read nor write it + q, k = q_src.clone(), k_src.clone() + out = env.call_context_no_append([0, 1, 2], new_lens, kv_lens, q, k, v) + + if rank == Q_LORA_RANK_ZERO: + ref, _ = _explicit_kv_reference(q, k, v, new_lens, kv_lens, MASK_CAUSAL, h) + torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) + assert torch.equal(q, q_src) + assert torch.equal(k, k_src) + assert torch.equal(env.pool, pool_fill) + + runs.append( + ( + rank, + {"output": out, "pool": env.pool.clone()}, + env.workspace.numel(), + ) + ) + assert torch.equal(v, v_src), "the sweep must have run on identical inputs" + _assert_q_lora_rank_inert(f"no-append context H={h}", runs) + + +def test_bf16_mla_qlora0_context_cached_kv_no_append_h32() -> None: + _mla_qlora_context_cached_kv_no_append_case(MLA_NUM_HEADS_H32, 307) + + +def test_bf16_mla_qlora0_context_cached_kv_no_append_h8() -> None: + _mla_qlora_context_cached_kv_no_append_case(MLA_NUM_HEADS_H8, 317) + + +# ─── The DeepSeek-R1-0528 cell: H=128 + YaRN rope table + q_scaling != 1 ─── +# +# Three axes move at once relative to the deepseek-v3-lite cells above: 128 +# query heads (attention DP replicates them whole onto every rank), a +# YaRN-scaled rotary_cos_sin table, and a softmax scale that folds YaRN's +# mscale^2 in through q_scaling. They are certified together — the four call +# flavors below at exactly the values a DeepSeek-R1-0528 layer passes — and +# separately, by the two axis sweeps after them. + + +def _assert_outside_band( + out: torch.Tensor, rival: torch.Tensor, threshold: float, label: str +) -> float: + """A rival hypothesis must land this far outside the tolerance band around + the observed output. Without it, matching the positive reference proves + little: an op that silently ignored the argument under test would pass the + positive comparison too wherever the argument's effect happens to be + small. Returns the separation in units of the allowance.""" + diff = (out.float() - rival.float()).abs().max().item() + allowance = ATOL + RTOL * rival.float().abs().max().item() + ratio = diff / allowance + assert ratio > threshold, f"{label}: only {ratio:.4g}x outside the tolerance band" + return ratio + + +def test_bf16_mla_r1_cell_context_prefill_h128() -> None: + """Fresh-prefill MLA context at the full R1 cell: 128/128/192, page 32, + the checkpoint's YaRN table, q_scaling = 1/mscale^2.""" + _mla_context_prefill_case( + MLA_NUM_HEADS_H128, + 925, + MLA_PAGE32, + [96, 33], + q_scaling=R1_Q_SCALING, + rope=_r1_rope_params(MLA_MAX_SEQ_LEN), + rope_scalars=R1_ROPE_SCALARS, + ) + + +def test_bf16_mla_r1_cell_generation_decode_h128() -> None: + """Latent-MQA decode at the full R1 cell: 128/1/576, page 32. The rope + table is passed but nothing is rotated here — the scale is what carries.""" + _mla_generation_decode_case( + MLA_NUM_HEADS_H128, + 926, + MLA_PAGE32, + [64, 31], + q_scaling=R1_Q_SCALING, + rope=_r1_rope_params(MLA_MAX_SEQ_LEN), + rope_scalars=R1_ROPE_SCALARS, + ) + + +def test_bf16_mla_r1_cell_mixed_batch_h128() -> None: + """Mixed batch at the full R1 cell: the two phase calls share one set of + batch tensors, one YaRN table and one q_scaling.""" + _mla_mixed_batch_case( + MLA_NUM_HEADS_H128, + 927, + MLA_PAGE32, + 64, + 33, + q_scaling=R1_Q_SCALING, + rope=_r1_rope_params(MLA_MAX_SEQ_LEN), + rope_scalars=R1_ROPE_SCALARS, + ) + + +def test_bf16_mla_r1_cell_context_cached_kv_no_append_h128() -> None: + """Cached-KV (no-append) context at the full R1 cell. This flavor applies + no rope at all — q arrives pre-rotated — so it is where q_scaling is the + only one of the three op-side axes with an effect.""" + _mla_context_cached_kv_case( + MLA_NUM_HEADS_H128, + 928, + MLA_PAGE32, + [96, 31, 0], + [32, 9, 25], + q_scaling=R1_Q_SCALING, + rope=_r1_rope_params(MLA_MAX_SEQ_LEN), + rope_scalars=R1_ROPE_SCALARS, + ) + + +# q_scaling values swept in both phases. 1.0 is the value certified before +# this run, R1_Q_SCALING (~0.53366) is what the checkpoint's YaRN config +# wants, and 2.0 / 0.25 bracket it from both sides so the argument is +# exercised as an axis rather than pinned at one number. +Q_SCALING_SWEEP = [1.0, R1_Q_SCALING, 2.0, 0.25] + +# A run at one q_scaling must land this far outside the tolerance band of a +# reference built at any other swept value. Measured over the full 4x4 matrix +# on sm_100 at H=128, page 32: every cross pair sits between 20.1x and 178x +# (the tightest is 1.0 against 2.0 in the context phase; the "q_scaling +# silently ignored" hypothesis — a reference at 1.0 for a run at any other +# value — spans 20.1x-70.5x), while a matching reference uses at most 0.23 of +# the same band. The 5x gate sits between those two populations by ~22x on +# one side and ~4x on the other. It is deliberately not tight enough to +# resolve two adjacent values: 0.5 against 0.53366, a 6.3% change of scale, +# separates by only 2.95x-3.61x, so the swept values are kept well apart. +Q_SCALING_MIN_SEPARATION = 5.0 + + +def test_bf16_mla_q_scaling_axis_h128() -> None: + """q_scaling is the softmax-scale axis of both MLA phases: QK^T is scaled + by 1 / (q_scaling * sqrt(nope + R)) in the context call and in the + generation call alike, the latter despite its head_size being C + R. + + Four values on identical inputs. Each run is checked against a reference + built at its own value and gated outside the references of the other + three, and the paged latent append is asserted bitwise identical across + the sweep — q_scaling moves the softmax scale and nothing else. + """ + h = MLA_NUM_HEADS_H128 + prefill_lens = [96, 33] + torch.manual_seed(935) + q_src, k_src, v, latent_src = _random_context_inputs(sum(prefill_lens), h) + decode_rows = [ + torch.randn(LATENT_DIM, dtype=torch.bfloat16, device="cuda") for _ in prefill_lens + ] + fused_q = torch.randn(len(prefill_lens), h * LATENT_DIM, dtype=torch.bfloat16, device="cuda") + + runs: List[Tuple[float, _MlaPagedEnv, torch.Tensor, torch.Tensor]] = [] + for q_scaling in Q_SCALING_SWEEP: + env = _MlaPagedEnv( + num_heads=h, + max_blocks_per_seq=MLA_MAX_SEQ_LEN // MLA_PAGE32, + tokens_per_block=MLA_PAGE32, + q_scaling=q_scaling, + rope=_r1_rope_params(MLA_MAX_SEQ_LEN), + rope_scalars=R1_ROPE_SCALARS, + ) + for rid, ln in enumerate(prefill_lens): + env.add_request(rid, ln) + out_ctx = env.call_context( + [0, 1], prefill_lens, q_src.clone(), k_src.clone(), v, latent_src.clone() + ) + for rid, row in enumerate(decode_rows): + env.append_decode_latent(rid, row) + out_gen = env.call_generation([0, 1], fused_q) + runs.append((q_scaling, env, out_ctx, out_gen)) + + for q_scaling, env, out_ctx, out_gen in runs: + for rival in Q_SCALING_SWEEP: + scale = 1.0 / (rival * math.sqrt(QK_HEAD_DIM)) + ref_ctx = env.context_reference( + prefill_lens, q_src, k_src, v, latent_src, softmax_scale=scale + ) + ref_gen = env.generation_reference([0, 1], fused_q, softmax_scale=scale) + if rival == q_scaling: + torch.testing.assert_close(out_ctx, ref_ctx, rtol=RTOL, atol=ATOL) + torch.testing.assert_close(out_gen, ref_gen, rtol=RTOL, atol=ATOL) + else: + _assert_outside_band( + out_ctx, + ref_ctx, + Q_SCALING_MIN_SEPARATION, + f"context at q_scaling={q_scaling} vs a reference at {rival}", + ) + _assert_outside_band( + out_gen, + ref_gen, + Q_SCALING_MIN_SEPARATION, + f"decode at q_scaling={q_scaling} vs a reference at {rival}", + ) + + base_pool = runs[0][1].pool + for q_scaling, env, _, _ in runs[1:]: + assert torch.equal(env.pool, base_pool), ( + f"the paged latent pool moved at q_scaling={q_scaling}" + ) + + +def _yarn_cos_sin_table( + num_positions: int, + dim: int, + theta: float, + factor: float, + original_max_positions: int, + beta_fast: int, + beta_slow: int, + mscale: float, + mscale_all_dim: float, +) -> torch.Tensor: + """The duplicated-layout fp32 (cos, sin) table a YaRN rope config + produces, built from the published YaRN formula with plain torch. + + A caller that cannot reach TensorRT-LLM's own table builder has to + reproduce this; asserting it against the builder's output is what makes + the formula usable as a contract statement rather than a description. + Layout: dim (cos, sin) pairs per position, the second dim/2 duplicating + the first, flattened to [1, num_positions * dim * 2]. + """ + half = dim // 2 + + def correction_dim(rotations: float) -> float: + return ( + dim + * math.log(original_max_positions / (rotations * 2 * math.pi)) + / (2 * math.log(theta)) + ) + + def attention_mscale(cfg_mscale: float) -> float: + return 1.0 if factor <= 1 else 0.1 * cfg_mscale * math.log(factor) + 1.0 + + low = max(0, math.floor(correction_dim(beta_fast))) + high = min(dim - 1, math.ceil(correction_dim(beta_slow))) + # The table's own amplitude factor — 1.0 whenever the two config mscales + # agree, which is where a model folds mscale into the softmax scale + # instead (see R1_Q_SCALING). + amplitude = attention_mscale(mscale) / attention_mscale(mscale_all_dim) + + pos_freqs = theta ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim) + ramp = torch.clamp( + (torch.arange(half, dtype=torch.float32) - low) / max(high - low, 0.001), 0, 1 + ) + # Interpolated (position-stretched) frequencies below the ramp, the + # original ones above it, blended across it. + inv_freq = ramp / (factor * pos_freqs) + (1 - ramp) / pos_freqs + angles = torch.outer(torch.arange(num_positions, dtype=torch.float32), inv_freq) + angles = torch.cat([angles, angles], dim=-1) # duplicate_data layout + table = torch.stack([torch.cos(angles) * amplitude, torch.sin(angles) * amplitude], dim=-1) + return table.reshape(1, -1).cuda() + + +def _unscaled_cos_sin_table() -> torch.Tensor: + """The plain theta-10000 table every MLA case above this section uses.""" + return RopeParams( + dim=QK_ROPE_HEAD_DIM, + theta=10000.0, + max_positions=MLA_MAX_SEQ_LEN, + duplicate_data=True, + ).create_rope_const_params()[1] + + +# A YaRN-table run must land this far outside the tolerance band of a +# reference built from the unscaled theta-10000 table. The separation grows +# with position, because YaRN only rescales the low-frequency half of the +# spectrum and those angles barely differ near position 0: measured on sm_100 +# at H=128, page 32, it is 8.6x over a 96-token sequence, 16.9x at 256, 27.0x +# at 512 and 31.7x at 960, against a matching reference's 0.23. The case +# below therefore prefills 512 tokens — 96 would not clear this gate, which +# is the point of choosing a length rather than loosening the gate. +YARN_TABLE_MIN_SEPARATION = 10.0 + + +def test_bf16_mla_yarn_rope_table_h128() -> None: + """The in-kernel RoPE of the fresh-prefill context call reads the + rotary_cos_sin table's *content*. + + Pinned here: the YaRN table a DeepSeek-R1-0528 config produces matches a + plain-torch build of the published formula; truncating it to max_seq_len + rows is the leading slice of the 163840-row table an engine allocates; + and the op driven from the torch-built table reproduces a table-aware + fp32 reference while landing far outside a reference built from the + unscaled theta-10000 table (and vice versa), so the table content is + established as load-bearing rather than assumed. + """ + trtllm_inv_freq, trtllm_table = _r1_rope_params(MLA_MAX_SEQ_LEN).create_rope_const_params() + torch_table = _yarn_cos_sin_table( + MLA_MAX_SEQ_LEN, + QK_ROPE_HEAD_DIM, + R1_ROPE_THETA, + R1_ROPE_FACTOR, + R1_ROPE_ORIGINAL_MAX_POSITIONS, + R1_ROPE_BETA_FAST, + R1_ROPE_BETA_SLOW, + R1_ROPE_MSCALE, + R1_ROPE_MSCALE_ALL_DIM, + ) + # Not bitwise: the two evaluate the same blend in different fp32 orders + # ((1 - (1 - ramp)) against ramp). Observed max abs diff 1.9e-6 on cos/sin + # values in [-1, 1] — a few fp32 ulp, 19% of this gate — while formula + # slips land 4-5 orders of magnitude past it: dropping the YaRN + # interpolation entirely (the unscaled table) or moving beta_fast 32 -> 64 + # both differ by 2.0, and taking the amplitude as mscale rather than the + # mscale/mscale_all_dim ratio differs by 3.7e-1. + torch.testing.assert_close(torch_table, trtllm_table, rtol=0.0, atol=1e-5) + assert trtllm_inv_freq is not None and trtllm_inv_freq.numel() == (QK_ROPE_HEAD_DIM // 2) + + # The unscaled table every MLA case above this section uses is the same + # construction at factor = 1: the ramp drops out and inv_freq(d) becomes + # 1/theta^(2d/R). Gated here so the formula is certified for both + # contents rather than only for the YaRN one (observed 3.8e-6, 38% of + # the gate). + unscaled_table = _unscaled_cos_sin_table() + torch.testing.assert_close( + _yarn_cos_sin_table( + MLA_MAX_SEQ_LEN, + QK_ROPE_HEAD_DIM, + R1_ROPE_THETA, + 1.0, + R1_ROPE_ORIGINAL_MAX_POSITIONS, + R1_ROPE_BETA_FAST, + R1_ROPE_BETA_SLOW, + R1_ROPE_MSCALE, + R1_ROPE_MSCALE_ALL_DIM, + ), + unscaled_table, + rtol=0.0, + atol=1e-5, + ) + + # Row content does not depend on the row count, so a table truncated to + # max_seq_len is the production table's leading slice. + _, full_table = _r1_rope_params(R1_MAX_POSITION_EMBEDDINGS).create_rope_const_params() + row = QK_ROPE_HEAD_DIM * 2 + assert torch.equal(full_table.view(-1, row)[:MLA_MAX_SEQ_LEN], trtllm_table.view(-1, row)) + del full_table + torch.cuda.empty_cache() + + h = MLA_NUM_HEADS_H128 + seq_lens = [512, 33] # far enough past the ramp for YaRN to bite + + def prefill(table: torch.Tensor) -> Tuple[_MlaPagedEnv, torch.Tensor, tuple]: + torch.manual_seed(936) + env = _MlaPagedEnv( + num_heads=h, + max_blocks_per_seq=MLA_MAX_SEQ_LEN // MLA_PAGE32, + tokens_per_block=MLA_PAGE32, + q_scaling=R1_Q_SCALING, + rope=_r1_rope_params(MLA_MAX_SEQ_LEN), + rope_scalars=R1_ROPE_SCALARS, + ) + env.rotary_cos_sin = table # op and reference both read it here + for rid, ln in enumerate(seq_lens): + env.add_request(rid, ln) + q, k, v, latent = _random_context_inputs(sum(seq_lens), h) + pre = (list(seq_lens), q.clone(), k.clone(), v, latent.clone()) + out = env.call_context([0, 1], seq_lens, q, k, v, latent) + return env, out, pre + + for driven, rival, label in ( + (torch_table, unscaled_table, "YaRN table"), + (unscaled_table, torch_table, "unscaled table"), + ): + env, out, pre = prefill(driven) + torch.testing.assert_close(out, env.context_reference(*pre), rtol=RTOL, atol=ATOL) + env.rotary_cos_sin = rival + _assert_outside_band( + out, + env.context_reference(*pre), + YARN_TABLE_MIN_SEPARATION, + f"{label} run against a reference built from the other table", + ) + env.rotary_cos_sin = driven + for rid in (0, 1): + env.check_cache(rid) # rope(k_pe) in the append follows the table + + +def _r1_prefill_observables( + rope_scalars: _RopeScalars, + cos_sin: Optional[torch.Tensor] = None, + inv_freq_mode: str = "keep", +) -> Dict[str, torch.Tensor]: + """One R1-cell fresh-prefill context call on fixed inputs, returning every + observable a rope argument could move: the output rows, the whole paged + latent pool (which receives rope(k_pe)), and the in-place-roped q/k.""" + seq_lens = [96, 33] + torch.manual_seed(936) + env = _MlaPagedEnv( + num_heads=MLA_NUM_HEADS_H128, + max_blocks_per_seq=MLA_MAX_SEQ_LEN // MLA_PAGE32, + tokens_per_block=MLA_PAGE32, + q_scaling=R1_Q_SCALING, + rope=_r1_rope_params(MLA_MAX_SEQ_LEN), + rope_scalars=rope_scalars, + ) + if cos_sin is not None: + env.rotary_cos_sin = cos_sin + if inv_freq_mode == "zeros": + env.rotary_inv_freq = torch.zeros_like(env.rotary_inv_freq) + elif inv_freq_mode == "none": + env.rotary_inv_freq = None + for rid, ln in enumerate(seq_lens): + env.add_request(rid, ln) + q, k, v, latent = _random_context_inputs(sum(seq_lens), MLA_NUM_HEADS_H128) + out = env.call_context([0, 1], seq_lens, q, k, v, latent) + return { + "output": out, + "pool": env.pool.clone(), + "q_after_call": q, + "k_after_call": k, + } + + +def test_bf16_mla_rope_scalars_inert_h128() -> None: + """With the table held fixed, the seven scalar rope arguments and + rotary_inv_freq move no observable of an MLA call. + + That is what makes the YaRN certification portable: a caller supplies the + scaled table and may pass whatever scalars its config carries. The last + check is the control that gives the bitwise comparisons meaning — the one + rope input the op does read (the table) is swapped, and the same + comparison has to see it. + """ + base = _r1_prefill_observables(R1_ROPE_SCALARS) + variants = { + # The set every MLA case above this section passes. + "unscaled-config scalars": _r1_prefill_observables( + _RopeScalars( + rope_max_positions=MLA_MAX_SEQ_LEN, + rope_original_max_positions=MLA_MAX_SEQ_LEN, + ) + ), + # Run-to-run determinism control for the comparisons around it. + "R1 scalars again": _r1_prefill_observables(R1_ROPE_SCALARS), + # Nothing a config would produce: a different theta, a different + # scaling family, m-scales far from 1, and position windows shorter + # than the sequences in flight. + "out-of-range scalars": _r1_prefill_observables( + _RopeScalars( + rope_base=500000.0, + rope_scale_type=3, + rope_scale=7.5, + rope_short_m_scale=3.0, + rope_long_m_scale=9.0, + rope_max_positions=77, + rope_original_max_positions=13, + ) + ), + "rotary_inv_freq zeroed": _r1_prefill_observables(R1_ROPE_SCALARS, inv_freq_mode="zeros"), + "rotary_inv_freq=None": _r1_prefill_observables(R1_ROPE_SCALARS, inv_freq_mode="none"), + } + for label, observed in variants.items(): + for key, expected in base.items(): + got = observed[key] + assert torch.equal(expected, got), ( + f"{label}: {key} is not bitwise identical — " + f"{int((expected != got).sum())}/{expected.numel()} elements, " + f"max abs {(expected.float() - got.float()).abs().max().item():.3e}" + ) + + swapped = _r1_prefill_observables(R1_ROPE_SCALARS, cos_sin=_unscaled_cos_sin_table()) + for key in ("output", "pool", "q_after_call", "k_after_call"): + assert not torch.equal(base[key], swapped[key]), ( + f"swapping the rope table left {key} bitwise unchanged — the " + f"inertness comparisons above cannot see a rope change at all" + ) + + +# ─── MLA over an fp8-e4m3 latent pool (quant_mode 128) ───────────────── +# The DeepSeek-R1-0528-FP4 cell: H = 128, page 32, C/R/nope/v = +# 512/64/128/128, q_lora_rank 1536, one latent row of C+R e4m3 bytes per +# token. The checkpoint's per-layer k_scale/v_scale are both 1.0, so s = 1.0 +# is the production scale; 1.5 and 2.0 are swept beside it because the +# scale tensors are the caller's only defence (None is read as 1.0). + + +def _fp8_mla_env( + num_heads: int = MLA_NUM_HEADS_H128, + kv_scaling_factor: Optional[float] = 1.0, + q_scaling: float = 1.0, + rope: Optional[RopeParams] = None, + rope_scalars: Optional[_RopeScalars] = None, + quant_mode: int = QUANT_MODE_FP8_KV_CACHE, +) -> _MlaPagedEnv: + """Page-32 MLA env over an fp8-e4m3 latent pool at the R1 head geometry. + + kv_scaling_factor=None leaves both kv scale tensors unpassed, which is the + silent-1.0 case a caller hits by forgetting them. q_scaling/rope/ + rope_scalars default to the baseline cell (unscaled theta-10000 table, + scale 1.0); _fp8_r1_env below fills the DeepSeek-R1-0528 values. + """ + return _MlaPagedEnv( + num_heads=num_heads, + max_blocks_per_seq=MLA_MAX_SEQ_LEN // MLA_PAGE32, + tokens_per_block=MLA_PAGE32, + q_lora_rank=Q_LORA_RANK_DSV3, + q_scaling=q_scaling, + rope=rope, + rope_scalars=rope_scalars, + pool_dtype=torch.float8_e4m3fn, + quant_mode=quant_mode, + kv_scaling_factor=kv_scaling_factor, + ) + + +def _fp8_mla_prefill( + env: _MlaPagedEnv, seq_lens: List[int], seed: int +) -> Tuple[torch.Tensor, Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]]: + """One fresh-prefill context call over the fp8 pool, returning the output + and the pre-call (q, k, v, latent) a reference has to be built from (the + call clobbers the q/k rope slices in place).""" + torch.manual_seed(seed) + q, k, v, latent = _random_context_inputs(sum(seq_lens), env.num_heads) + pre = (q.clone(), k.clone(), v, latent.clone()) + for rid, ln in enumerate(seq_lens): + env.add_request(rid, ln) + out = env.call_context(list(range(len(seq_lens))), seq_lens, q, k, v, latent) + return out, pre + + +def _band_fraction(out: torch.Tensor, ref: torch.Tensor, atol: float, rtol: float) -> float: + """How much of the atol + rtol*|ref| allowance the largest error uses — + elementwise, the statistic torch.testing.assert_close itself gates on, so + a value above 1.0 is exactly an assert_close failure, and by how much.""" + diff = (out.float() - ref.float()).abs() + return float((diff / (atol + rtol * ref.float().abs())).max()) + + +def _assert_far_outside_band( + out: torch.Tensor, + rival: torch.Tensor, + threshold: float, + label: str, + atol: float, + rtol: float, +) -> None: + """A rival hypothesis has to fail assert_close by at least this factor.""" + ratio = _band_fraction(out, rival, atol, rtol) + assert ratio > threshold, f"{label}: only {ratio:.4g}x outside the band" + + +def test_fp8_mla_context_prefill_h128() -> None: + """MLA fresh-prefill context over an fp8-e4m3 latent pool, s = 1.0. + + Two things the standard configuration's fp8 section does NOT carry over, + both measured here rather than inherited: + + 1. The context FMHA does *not* stay in bf16. Under quant_mode 128 the MLA + context path quantizes q, k and v to e4m3 (at scale 1.0 — the pool's + write scale is not applied to them) and runs the FMHA on those, so + prefill accuracy *is* affected by cache quantization. The bf16-KV + reference — the one the bf16 pool is certified against, and what a + caller would assume from the standard-configuration section — sits far + outside the bf16 band (observed 26x), while the e4m3-input reference + matches within the fp8 band (observed 49% of it). The sharp form of + this claim is the peaked-attention readout in the next case. + 2. The appended latent row lands as e4m3(row * kv_scale_orig_quant) on + *both* halves — compressed_kv bitwise, and the in-kernel-roped k_pe + half within the e4m3-ulp gate (observed bitwise too). + + The requests' pages are deliberately scattered and out of order, so a page + stride computed with the wrong element width (2-byte bf16 rather than + 1-byte e4m3) lands in a page this test then finds non-zero. + """ + env = _fp8_mla_env() + lens = [96, 33] # three exact 32-slot pages, and a one-token spill + torch.manual_seed(400) + env.add_request(0, lens[0]) + env.add_request(1, lens[1]) + # Hand-assigned, complete page sets (3 pages for 96 tokens, 2 for 33), so + # the env allocates none of its own: scattered and non-monotonic, sharing + # no page between the two requests. + env.pages[0] = [5, 1, 9] + env.pages[1] = [4, 7] + q, k, v, latent = _random_context_inputs(sum(lens), env.num_heads) + pre = (q.clone(), k.clone(), v, latent.clone()) + out = env.call_context([0, 1], lens, q, k, v, latent) + + ref = env.context_reference(lens, *pre, e4m3_inputs=True) + torch.testing.assert_close(out, ref, rtol=FP8_MLA_RTOL, atol=FP8_MLA_ATOL) + # Rival: the bf16-KV reference, i.e. "cache quantization does not reach + # the context math". Gated at 10x the bf16 band; observed 32x. + _assert_far_outside_band( + out, + env.context_reference(lens, *pre), + 10.0, + "fp8 MLA context against the bf16-KV reference", + ATOL, + RTOL, + ) + for rid in (0, 1): + env.check_cache(rid) + env.check_unwritten_pool_zero([0, 1]) + + +def test_fp8_mla_context_v_operand_is_e4m3_h128() -> None: + """The context FMHA's V operand really is e4m3, read out bitwise. + + The realistic-input case above compares two models that differ by the size + of the tolerance itself, so it cannot settle "is the context math fp8?" on + its own. This one can: give key 0 a score every other key cannot approach + (matching nope halves at magnitude 4, k_pe zeroed so the roped tails + contribute nothing) and every query row's softmax collapses onto it, so the + output row *is* V row 0. It comes back as e4m3(V[0]) — bit-exact, max abs + deviation 0, with only signed zeros differing in raw bytes — while the + unquantized V[0] a bf16-pool call would return is up to 63 bf16 ulps away + and matches in only 52% of its bytes. + """ + env = _fp8_mla_env() + seq_len = 40 + torch.manual_seed(500) + h = env.num_heads + q = torch.zeros(seq_len, h, QK_HEAD_DIM, dtype=torch.bfloat16, device="cuda") + q[:, :, :QK_NOPE_HEAD_DIM] = 4.0 + k = torch.zeros(seq_len, h, QK_HEAD_DIM, dtype=torch.bfloat16, device="cuda") + k[0, :, :QK_NOPE_HEAD_DIM] = 4.0 + packed = torch.zeros( + seq_len, + h * (QK_NOPE_HEAD_DIM + V_HEAD_DIM), + dtype=torch.bfloat16, + device="cuda", + ) + packed[:, : h * QK_NOPE_HEAD_DIM] = k[:, :, :QK_NOPE_HEAD_DIM].reshape( + seq_len, h * QK_NOPE_HEAD_DIM + ) + v = _v_split_view(packed, h) + v.view(seq_len, h, V_HEAD_DIM).copy_( + torch.randn(seq_len, h, V_HEAD_DIM, dtype=torch.bfloat16, device="cuda") + ) + latent = torch.randn(seq_len, LATENT_DIM, dtype=torch.bfloat16, device="cuda") + latent[:, KV_LORA_RANK:] = 0.0 + v0 = v.view(seq_len, h, V_HEAD_DIM)[0].clone() + + env.add_request(0, seq_len) + rows = env.call_context( + [0], + [seq_len], + q.view(seq_len, h * QK_HEAD_DIM), + k.view(seq_len, h * QK_HEAD_DIM), + v, + latent, + ).view(seq_len, h, V_HEAD_DIM) + + quantized = v0.float().to(torch.float8_e4m3fn).float().to(torch.bfloat16) + expected = quantized.unsqueeze(0).expand_as(rows) + torch.testing.assert_close(rows, expected, rtol=2**-8, atol=0.0) # one bf16 ulp + exact = float( + (rows.reshape(-1).view(torch.uint8) == expected.reshape(-1).view(torch.uint8)) + .float() + .mean() + ) + assert exact > 0.99, f"only {exact:.4f} of the readout is bit-exact e4m3(V[0])" + rival = v0.float().unsqueeze(0) + ulps = float(((rows.float() - rival).abs() / (2**-8 * rival.abs().clamp(min=2**-8))).max()) + assert ulps > 8.0, ( + f"the unquantized V row is only {ulps:.3g} bf16 ulps away — this probe " + f"cannot tell an e4m3 V operand from a bf16 one" + ) + + +def test_fp8_mla_generation_decode_h128() -> None: + """MLA generation decode over the fp8-e4m3 latent pool, s = 1.0. + + Histories of 64 (two exact pages, so the first decode token opens page 2) + and 31 (first decode token takes page 0's last slot, second opens page 1), + two steps. The decode kernel reads its query from quant_q_buffer, never + from `q`, and applies no kv scale of its own — the caller's + mla_bmm1_scale/mla_bmm2_scale carry the dequantization (see + fp8_decode_buffers). What the call then computes is latent MQA over the + e4m3 round trip of both operands, which is what generation_reference + builds. The pool must come back bitwise unchanged: the generation phase + reads it and nothing else. + """ + env = _fp8_mla_env() + prefill = [64, 31] + _fp8_mla_prefill(env, prefill, 401) + for rid in (0, 1): + env.check_cache(rid) + for step in range(2): + for rid in (0, 1): + env.append_decode_latent( + rid, + torch.randn(LATENT_DIM, dtype=torch.bfloat16, device="cuda") * 0.5, + ) + pool_before = env.pool.clone() + fused_q = ( + torch.randn(2, env.num_heads * LATENT_DIM, dtype=torch.bfloat16, device="cuda") * 0.3 + ) + out = env.call_generation([0, 1], fused_q) + ref = env.generation_reference([0, 1], fused_q) + torch.testing.assert_close(out, ref, rtol=FP8_MLA_RTOL, atol=FP8_MLA_ATOL) + assert _bitwise_equal(pool_before, env.pool), ( + "the MLA generation call wrote to the fp8 pool" + ) + + +def test_fp8_mla_decode_buffer_roles_h128() -> None: + """The three buffers an fp8-pool MLA generation call requires, one role at + a time. All three are presence-checked by the op; what each one carries is + unchecked, so a caller gets these wrong silently. + + - quant_q_buffer holds the query: replacing `q` with garbage leaves the + output bitwise identical, so `q` is not read on this path at all. + - mla_bmm1_scale[1] (the log2-domain copy) is the softmax scale the kernel + applies; element [0] is inert. + - mla_bmm2_scale[0] multiplies the output. + """ + env = _fp8_mla_env() + _fp8_mla_prefill(env, [40], 402) + torch.manual_seed(403) + env.append_decode_latent(0, torch.randn(LATENT_DIM, dtype=torch.bfloat16, device="cuda") * 0.5) + fused_q = torch.randn(1, env.num_heads * LATENT_DIM, dtype=torch.bfloat16, device="cuda") * 0.3 + quant_q, bmm1, bmm2 = env.fp8_decode_buffers(fused_q) + ref = env.generation_reference([0], fused_q) + + base = env.call_generation([0], fused_q, fp8_buffers=(quant_q, bmm1, bmm2)) + torch.testing.assert_close(base, ref, rtol=FP8_MLA_RTOL, atol=FP8_MLA_ATOL) + + for label, buffers, message in ( + ("quant_q_buffer", (None, bmm1, bmm2), "quant_q_buf is nullptr"), + ("mla_bmm1_scale", (quant_q, None, bmm2), "bmm1_scale is nullptr"), + ("mla_bmm2_scale", (quant_q, bmm1, None), "bmm2_scale is nullptr"), + ): + try: + env.call_generation([0], fused_q, fp8_buffers=buffers) + except RuntimeError as exc: + assert message in str(exc), f"{label}: unexpected message: {exc}" + else: + raise AssertionError(f"{label}=None was accepted under quant_mode 128") + + garbage_q = torch.randn_like(fused_q) * 4 + with_garbage = env.call_generation([0], garbage_q, fp8_buffers=(quant_q, bmm1, bmm2)) + assert _bitwise_equal(base, with_garbage), ( + "replacing q moved the output — the decode does read the bf16 query" + ) + # The buffer is consumed as raw bytes through a pointer, so the same bytes + # handed over as uint8 are the same query: nothing checks its dtype. + assert _bitwise_equal( + base, + env.call_generation([0], fused_q, fp8_buffers=(quant_q.view(torch.uint8), bmm1, bmm2)), + ), "quant_q_buffer's dtype changed the result" + + inert_head = torch.tensor([0.0, float(bmm1[1])], dtype=torch.float32, device="cuda") + assert _bitwise_equal( + base, env.call_generation([0], fused_q, fp8_buffers=(quant_q, inert_head, bmm2)) + ), "mla_bmm1_scale[0] moved the output" + dead_log2 = torch.tensor([float(bmm1[0]), 0.0], dtype=torch.float32, device="cuda") + # scale 0 flattens the softmax to a uniform average of the cached rows. + _assert_far_outside_band( + env.call_generation([0], fused_q, fp8_buffers=(quant_q, dead_log2, bmm2)), + ref, + 2.0, + "mla_bmm1_scale[1] zeroed", + FP8_MLA_ATOL, + FP8_MLA_RTOL, + ) + + doubled = torch.tensor([2.0 * env.kv_scale], dtype=torch.float32, device="cuda") + out2 = env.call_generation([0], fused_q, fp8_buffers=(quant_q, bmm1, doubled)) + torch.testing.assert_close( + out2, (ref.float() * 2.0).to(ref.dtype), rtol=FP8_MLA_RTOL, atol=FP8_MLA_ATOL + ) + + +def test_fp8_mla_kv_scale_semantics_h128() -> None: + """What the two kv scale tensors do on the MLA path, per phase. + + Write side (both phases' appends): kv_scale_orig_quant is applied, and the + row lands bit-exactly as e4m3(row * orig_quant) at every scale. + + Context read side: the FMHA quantizes q/k/v at scale 1.0 but still applies + the dequantization factors s**2 (softmax scale) and s (output), so at + s != 1.0 the context result is silently wrong — the s = 1.0 math sits + 21x-48x outside the fp8 band, while the distorted model reproduces the run. + + Generation read side: neither scale tensor is read at all. The dequant has + to arrive folded into mla_bmm1_scale/mla_bmm2_scale (as the production + generation-preprocessing step writes them), and once it does, decode is + correct at any scale — verified at s = 2.0, where dropping the fold lands + 3.6x outside the band. + + Passing None for both is bitwise identical to passing 1.0 tensors, in the + output and in the pool: a forgotten scale is a wrong-number bug at any + other s, not a crash. + """ + lens = [48, 17] + for scale in KV_SCALE_SWEEP: + env = _fp8_mla_env(kv_scaling_factor=scale) + out, pre = _fp8_mla_prefill(env, lens, 404) + for rid in (0, 1): + env.check_cache(rid) # e4m3(row * orig_quant), both halves + env.check_unwritten_pool_zero([0, 1]) + true_ref = env.context_reference(lens, *pre, e4m3_inputs=True) + distorted = env.context_reference(lens, *pre, e4m3_inputs=True, kv_scale_factors=True) + if scale == 1.0: + torch.testing.assert_close(out, true_ref, rtol=FP8_MLA_RTOL, atol=FP8_MLA_ATOL) + else: + _assert_far_outside_band( + out, + true_ref, + 10.0, + f"fp8 MLA context at s={scale} against the s=1.0 math", + FP8_MLA_ATOL, + FP8_MLA_RTOL, + ) + # The distortion is exactly the two kv-scale factors: the model + # lands inside the fp8 band itself (observed 0.58x at s=1.5 and + # 0.84x at s=2.0, against 21x / 48x for the true-value hypothesis), + # so it is gated at 1.5x rather than 1.0x only to leave the same + # headroom the certified s=1.0 case has. + assert _band_fraction(out, distorted, FP8_MLA_ATOL, FP8_MLA_RTOL) < 1.5, ( + f"the s**2 / s model does not explain the s={scale} run" + ) + + # None in both slots == 1.0 tensors, bitwise, in the output and the pool. + env_none = _fp8_mla_env(kv_scaling_factor=None) + out_none, _ = _fp8_mla_prefill(env_none, lens, 404) + env_one = _fp8_mla_env(kv_scaling_factor=1.0) + out_one, _ = _fp8_mla_prefill(env_one, lens, 404) + assert _bitwise_equal(out_none, out_one) and _bitwise_equal(env_none.pool, env_one.pool), ( + "kv_scale_* = None is not the s = 1.0 path" + ) + + # Generation at s != 1.0: correct with the folded scales — at a + # non-power-of-two scale too, where the e4m3 grid genuinely depends on the + # value — and blind to the scale tensors themselves. + for scale in (1.5, 2.0): + env = _fp8_mla_env(kv_scaling_factor=scale) + _fp8_mla_prefill(env, [40], 405) + torch.manual_seed(406) + env.append_decode_latent( + 0, torch.randn(LATENT_DIM, dtype=torch.bfloat16, device="cuda") * 0.5 + ) + fused_q = ( + torch.randn(1, env.num_heads * LATENT_DIM, dtype=torch.bfloat16, device="cuda") * 0.3 + ) + quant_q, bmm1, bmm2 = env.fp8_decode_buffers(fused_q) + out = env.call_generation([0], fused_q, fp8_buffers=(quant_q, bmm1, bmm2)) + ref = env.generation_reference([0], fused_q) + torch.testing.assert_close(out, ref, rtol=FP8_MLA_RTOL, atol=FP8_MLA_ATOL) + # Unfolded scales — what a caller gets by reading the standard + # configuration's fp8 section and assuming kv_scale_quant_orig is + # applied on read. + unfolded_1 = torch.tensor( + [env.softmax_scale, env.softmax_scale * math.log2(math.e)], + dtype=torch.float32, + device="cuda", + ) + unfolded_2 = torch.tensor([1.0], dtype=torch.float32, device="cuda") + _assert_far_outside_band( + env.call_generation([0], fused_q, fp8_buffers=(quant_q, unfolded_1, unfolded_2)), + ref, + 2.0, + f"fp8 MLA decode at s={scale} with the kv scale left out of the bmm scales", + FP8_MLA_ATOL, + FP8_MLA_RTOL, + ) + # The scale tensors themselves reach nothing: dropping both is bitwise + # identical at a scale where a read-side dequant would be visible. + env.kv_scale_orig_quant = None + env.kv_scale_quant_orig = None + assert _bitwise_equal( + out, env.call_generation([0], fused_q, fp8_buffers=(quant_q, bmm1, bmm2)) + ), "the MLA decode path does read the kv scale tensors" + + +# ─── The complete DeepSeek-R1-0528 cell over the fp8-e4m3 latent pool ── +# +# The fp8 cases above run the unscaled theta-10000 table at q_scaling = 1.0; +# the R1-cell cases further up run the YaRN table and q_scaling = 1/mscale^2 +# over a *bf16* pool. A DeepSeek-R1-0528 rank passes all three in every call, +# and certifying axes separately does not certify their combination, so the +# whole cell runs here: fp8 pool + YaRN table + q_scaling, H = 128, page 32, +# C/R/nope/v = 512/64/128/128, q_lora_rank 1536, s = 1.0. + + +def _fp8_r1_env(**kwargs) -> _MlaPagedEnv: + """_fp8_mla_env at the complete R1 cell: the checkpoint's YaRN rope table, + its scalar rope set, and q_scaling = 1/mscale^2.""" + kwargs.setdefault("q_scaling", R1_Q_SCALING) + kwargs.setdefault("rope", _r1_rope_params(MLA_MAX_SEQ_LEN)) + kwargs.setdefault("rope_scalars", R1_ROPE_SCALARS) + return _fp8_mla_env(**kwargs) + + +def _unscaled_table_env(**kwargs) -> _MlaPagedEnv: + """The same env on the *unscaled* theta-10000 table — the rival whose + references and append mirrors show that the table content is what the + kernel reads. Only its rope table and the mirrors built from it are used; + its pool is never called into.""" + kwargs.setdefault("q_scaling", R1_Q_SCALING) + return _fp8_mla_env(**kwargs) + + +def _rope_rows(env: _MlaPagedEnv, latent: torch.Tensor, seq_len: int) -> torch.Tensor: + """The latent rows a fresh-prefill append is expected to leave in the pool + for one sequence, roped position by position with env's table.""" + rows = latent[:seq_len].clone() + for i in range(seq_len): + rows[i, KV_LORA_RANK:] = env.rope_ref(latent[i, KV_LORA_RANK:], i) + return rows + + +def _assert_table_pins_the_append( + env: _MlaPagedEnv, + rival_env: _MlaPagedEnv, + request_id: int, + latent: torch.Tensor, + seq_len: int, +) -> None: + """check_cache's bit-exact comparison must *fail* against a mirror roped by + the unscaled table. + + The appended k_pe half is the sharpest place the rope table shows up: it is + gated bitwise rather than by tolerance. But a bit-exact match proves nothing + about which table was read unless the same comparison can tell the two + tables apart at the positions in flight, and YaRN rescales only the + low-frequency half of the spectrum, so the two agree closely at small + positions — exactly where a 96-token prefill lives. Measured here at + L = 96, page 32: 2336 of the 6144 e4m3 bytes of the roped half differ + between the two mirrors (38.0%), against an allowance of 6, so the + comparison separates them by ~389x. + """ + correct = env.latent_rows[request_id] + env.latent_rows[request_id] = [_rope_rows(rival_env, latent, seq_len)] + try: + env.check_cache(request_id) + except AssertionError: + pass + else: + raise AssertionError( + "the appended rows match a mirror roped by the unscaled table too — " + "this comparison cannot tell the two rope tables apart here" + ) + finally: + env.latent_rows[request_id] = correct + + +def test_fp8_mla_r1_cell_context_prefill_h128() -> None: + """Fresh-prefill MLA context at the complete R1 cell over an fp8-e4m3 + latent pool: 128/128/192, page 32, YaRN table, q_scaling = 1/mscale^2. + + All three axes are gated, not just run: the output matches an fp32 + reference over e4m3-rounded q/k/v built from this table and this scale + (observed 40% of the fp8 band), and sits far outside three rivals — the + bf16-KV reference the bf16 pool is certified against (49.8x the bf16 band), + a reference at q_scaling = 1.0 (12.7x the fp8 band), and one built from the + unscaled theta-10000 table (3.7x). The bit-exact append carries the table + a second time and much harder, with its own rival control. + """ + env = _fp8_r1_env() + rival = _unscaled_table_env() + lens = [96, 33] # three exact 32-slot pages, and a one-token spill + torch.manual_seed(950) + for rid, ln in enumerate(lens): + env.add_request(rid, ln) + q, k, v, latent = _random_context_inputs(sum(lens), env.num_heads) + pre = (q.clone(), k.clone(), v, latent.clone()) + out = env.call_context([0, 1], lens, q, k, v, latent) + + torch.testing.assert_close( + out, + env.context_reference(lens, *pre, e4m3_inputs=True), + rtol=FP8_MLA_RTOL, + atol=FP8_MLA_ATOL, + ) + _assert_far_outside_band( + out, + env.context_reference(lens, *pre), + 10.0, + "fp8 R1-cell MLA context against the bf16-KV reference", + ATOL, + RTOL, + ) + _assert_far_outside_band( + out, + env.context_reference( + lens, + *pre, + e4m3_inputs=True, + softmax_scale=1.0 / math.sqrt(QK_HEAD_DIM), + ), + 4.0, + "fp8 R1-cell MLA context against a q_scaling = 1.0 reference", + FP8_MLA_ATOL, + FP8_MLA_RTOL, + ) + _assert_far_outside_band( + out, + rival.context_reference(lens, *pre, e4m3_inputs=True), + 2.0, + "fp8 R1-cell MLA context against an unscaled-rope-table reference", + FP8_MLA_ATOL, + FP8_MLA_RTOL, + ) + + for rid in (0, 1): + env.check_cache(rid) + env.check_unwritten_pool_zero([0, 1]) + _assert_table_pins_the_append(env, rival, 0, pre[3], lens[0]) + + +def test_fp8_mla_r1_cell_generation_decode_h128() -> None: + """Latent-MQA decode at the complete R1 cell over the fp8-e4m3 pool. + + Nothing is rotated in this phase, so of the cell's three axes only + q_scaling has an effect here — it reaches the kernel through the caller's + mla_bmm1_scale, which fp8_decode_buffers builds from the env's softmax + scale. Histories of 64 (two exact pages, first decode token opens page 2) + and 31 (first decode token takes page 0's last slot, second opens page 1), + two steps. Gated against a reference at q_scaling = 1.0 (observed 4.2x the + fp8 band, where the matching reference uses 0.20). + """ + env = _fp8_r1_env() + prefill = [64, 31] + _fp8_mla_prefill(env, prefill, 951) + for rid in (0, 1): + env.check_cache(rid) + for _ in range(2): + for rid in (0, 1): + env.append_decode_latent( + rid, + torch.randn(LATENT_DIM, dtype=torch.bfloat16, device="cuda") * 0.5, + ) + pool_before = env.pool.clone() + fused_q = ( + torch.randn(2, env.num_heads * LATENT_DIM, dtype=torch.bfloat16, device="cuda") * 0.3 + ) + out = env.call_generation([0, 1], fused_q) + torch.testing.assert_close( + out, + env.generation_reference([0, 1], fused_q), + rtol=FP8_MLA_RTOL, + atol=FP8_MLA_ATOL, + ) + _assert_far_outside_band( + out, + env.generation_reference([0, 1], fused_q, softmax_scale=1.0 / math.sqrt(QK_HEAD_DIM)), + 2.0, + "fp8 R1-cell MLA decode against a q_scaling = 1.0 reference", + FP8_MLA_ATOL, + FP8_MLA_RTOL, + ) + assert _bitwise_equal(pool_before, env.pool), ( + "the MLA generation call wrote to the fp8 pool" + ) + + +def test_fp8_mla_r1_cell_mixed_batch_h128() -> None: + """A mixed batch over the fp8-e4m3 pool at the complete R1 cell: a 64-token + history decoded next to a fresh 33-token context sequence, the two phase + calls sharing one set of batch tensors and one page-32 offsets table. + + This is the pairing dispatch 02 did not run under fp8: the context call + carries a trailing generation sequence in its metadata (and must ignore its + rows), and the generation call carries a leading context sequence (and must + index its own from num_contexts). Both appends are checked bit-exactly + afterwards, so neither call may have written into the other's pages. + """ + env = _fp8_r1_env() + first, second = 64, 33 + torch.manual_seed(952) + env.add_request(0, first) + q, k, v, latent = _random_context_inputs(first, env.num_heads) + env.call_context([0], [first], q, k, v, latent) + + env.add_request(1, second) + env.append_decode_latent(0, torch.randn(LATENT_DIM, dtype=torch.bfloat16, device="cuda") * 0.5) + q, k, v, latent = _random_context_inputs(second, env.num_heads) + pre = (q.clone(), k.clone(), v, latent.clone()) + out_ctx = env.call_context([1], [second], q, k, v, latent, gen_request_ids=[0]) + torch.testing.assert_close( + out_ctx, + env.context_reference([second], *pre, e4m3_inputs=True), + rtol=FP8_MLA_RTOL, + atol=FP8_MLA_ATOL, + ) + _assert_far_outside_band( + out_ctx, + env.context_reference([second], *pre), + 10.0, + "fp8 R1-cell mixed-batch context against the bf16-KV reference", + ATOL, + RTOL, + ) + + fused_q = torch.randn(1, env.num_heads * LATENT_DIM, dtype=torch.bfloat16, device="cuda") * 0.3 + out_gen = env.call_generation([0], fused_q, ctx_request_ids=[1], ctx_seq_lens=[second]) + torch.testing.assert_close( + out_gen, + env.generation_reference([0], fused_q), + rtol=FP8_MLA_RTOL, + atol=FP8_MLA_ATOL, + ) + env.check_cache(0) + env.check_cache(1) + env.check_unwritten_pool_zero([0, 1]) + + +def _fp8_kv_b_proj(num_heads: int) -> torch.Tensor: + """A fixed absorbed-KV projection weight `[C, H*(nope+v)]`, scaled by + 1/sqrt(C) so its output keeps the unit magnitude every other operand in + this file has. Stands in for the checkpoint's kv_b_proj, which is what + turns a gathered compressed_kv row into K's nope half and V.""" + g = torch.Generator(device="cuda").manual_seed(9001) + w = torch.randn( + KV_LORA_RANK, + num_heads * (QK_NOPE_HEAD_DIM + V_HEAD_DIM), + generator=g, + dtype=torch.float32, + device="cuda", + ) + return (w / math.sqrt(KV_LORA_RANK)).to(torch.bfloat16) + + +def _fp8_pool_backed_kv( + env: _MlaPagedEnv, cached_lens: List[int], new_lens: List[int] +) -> Tuple[torch.Tensor, torch.Tensor]: + """Build a no-append context call's K/V the way a target on an fp8 latent + pool builds them, and write the cached prefixes into the pool. + + Cached rows are stored exactly as the op's own append stores them, + `e4m3(row * kv_scale_orig_quant)`, and read back with the fp8 latent + gather's formula — `bf16(float(cache_byte) * kv_scale_quant_orig)`, + mirrored in plain torch here so the operands depend on no other entry. The + gathered compressed_kv goes through _fp8_kv_b_proj into the packed + `[nope | v]` buffer the context FMHA wants; the gathered k_pe is already + roped (that is how the pool holds it) and becomes k's per-head tail. The + new tokens never touch the pool: their latent is fresh and their k_pe is + roped here at its absolute position, which is what the caller does. + + Returns `(k, v)` with every sequence's `[cached | new]` range concatenated + in batch order, `v` carrying the required `H * (nope + v)` row stride. + """ + h = env.num_heads + w = _fp8_kv_b_proj(h).float() + k_parts: List[torch.Tensor] = [] + packed_parts: List[torch.Tensor] = [] + for rid, (c_len, n_len) in enumerate(zip(cached_lens, new_lens)): + cached = torch.randn(c_len, LATENT_DIM, dtype=torch.bfloat16, device="cuda") + for i in range(c_len): + cached[i, KV_LORA_RANK:] = env.rope_ref(cached[i, KV_LORA_RANK:], i) + stored = env.to_pool(cached) + tpb = env.tokens_per_block + for i in range(c_len): + env.pool[env.pages[rid][i // tpb], i % tpb] = stored[i] + gathered = (stored.float() * env.kv_scale).to(torch.bfloat16) + + fresh = torch.randn(n_len, LATENT_DIM, dtype=torch.bfloat16, device="cuda") + for i in range(n_len): + fresh[i, KV_LORA_RANK:] = env.rope_ref(fresh[i, KV_LORA_RANK:], c_len + i) + + rows = torch.cat([gathered, fresh]) + total = c_len + n_len + packed = (rows[:, :KV_LORA_RANK].float() @ w).to(torch.bfloat16) + k = torch.empty(total, h, QK_HEAD_DIM, dtype=torch.bfloat16, device="cuda") + k[..., :QK_NOPE_HEAD_DIM] = packed[:, : h * QK_NOPE_HEAD_DIM].view( + total, h, QK_NOPE_HEAD_DIM + ) + k[..., QK_NOPE_HEAD_DIM:] = rows[:, KV_LORA_RANK:].unsqueeze(1).expand(-1, h, -1) + k_parts.append(k.reshape(total, h * QK_HEAD_DIM)) + packed_parts.append(packed) + return torch.cat(k_parts), _v_split_view(torch.cat(packed_parts), h) + + +def _fp8_no_append_reference( + env: _MlaPagedEnv, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + new_lens: List[int], + kv_lens: List[int], + e4m3_inputs: bool = True, + kv_scale_factors: bool = False, + softmax_scale: Optional[float] = None, +) -> torch.Tensor: + """The no-append context flavor's fp8 reference — the counterpart of + _MlaPagedEnv.context_reference for this flavor's explicit K/V. + + e4m3_inputs (the default here) rounds q, k and v through e4m3 at scale 1.0, + which is what the call quantizes them to under quant_mode 128; setting it + False builds the bf16-KV rival. kv_scale_factors adds the two factors that + path applies on top — s**2 on the softmax scale and s on the output — + which cancel only at s = 1.0. softmax_scale overrides the env's own scale, + which is what builds a q_scaling rival. + """ + scale = env.softmax_scale if softmax_scale is None else softmax_scale + s = env.kv_scale if kv_scale_factors else 1.0 + return _explicit_kv_reference( + q, + k, + v, + new_lens, + kv_lens, + MASK_CAUSAL, + env.num_heads, + scale * s * s, + e4m3_inputs=e4m3_inputs, + out_scale=s, + )[0] + + +def test_fp8_mla_r1_cell_context_cached_kv_no_append_h128() -> None: + """Cached-KV (no-append) MLA context over an fp8-e4m3 latent pool, at the + complete R1 cell. This is the flavor an engine with block reuse on runs for + every context request that hits a cached prefix, and it had never been run + over an fp8 pool. + + Its K/V come from where production's come from: the cached prefix is real + e4m3 pool content read back through the fp8 latent gather's formula and put + through a kv_b_proj-shaped matmul; only the new tokens are fresh. Prefixes + of 96 / 31 / 0 reach KV lengths of 128 / 40 / 25. + + The call quantizes those operands to e4m3 itself, exactly as the + fresh-prefill flavor does (settled bitwise in the next test) — so the + reference rounds q/k/v through e4m3, observed at 41% of the fp8 band, with + the bf16-KV model 2.2x outside it and a q_scaling = 1.0 reference 10.8x. + Nothing is mutated: q, k, v and the whole pool come back bitwise intact. + """ + env = _fp8_r1_env() + cached_lens = [96, 31, 0] + new_lens = [32, 9, 25] + kv_lens = [c + n for c, n in zip(cached_lens, new_lens)] + torch.manual_seed(953) + for rid, total in enumerate(kv_lens): + env.add_request(rid, new_lens[rid]) + env.reserve_cache_pages(rid, total) + k, v = _fp8_pool_backed_kv(env, cached_lens, new_lens) + q = torch.randn(sum(new_lens), env.num_heads * QK_HEAD_DIM, dtype=torch.bfloat16, device="cuda") + q_orig, k_orig, v_orig = q.clone(), k.clone(), v.clone() + pool_before = env.pool.clone() + + out = env.call_context_no_append([0, 1, 2], new_lens, kv_lens, q, k, v) + operands = (q, k, v, new_lens, kv_lens) + + torch.testing.assert_close( + out, + _fp8_no_append_reference(env, *operands), + rtol=FP8_MLA_RTOL, + atol=FP8_MLA_ATOL, + ) + _assert_far_outside_band( + out, + _fp8_no_append_reference(env, *operands, e4m3_inputs=False), + 1.5, + "fp8 no-append MLA context against the bf16-KV reference", + FP8_MLA_ATOL, + FP8_MLA_RTOL, + ) + _assert_far_outside_band( + out, + _fp8_no_append_reference(env, *operands, softmax_scale=1.0 / math.sqrt(QK_HEAD_DIM)), + 4.0, + "fp8 no-append MLA context against a q_scaling = 1.0 reference", + FP8_MLA_ATOL, + FP8_MLA_RTOL, + ) + + assert torch.equal(q, q_orig) + assert torch.equal(k, k_orig) + assert torch.equal(v, v_orig) + assert _bitwise_equal(env.pool, pool_before), ( + "the no-append MLA context call touched the fp8 pool" + ) + + +def test_fp8_mla_no_append_operands_and_scale_h128() -> None: + """What quant_mode 128 does to the no-append context flavor, settled the + two ways the realistic case cannot settle on its own. + + 1. The operands really are e4m3 — the same fact the fresh-prefill flavor + has, and it could not be assumed here, because this flavor's K/V arrive + already dequantized from an fp8 pool and it appends nothing. Peaked + readout: give key 0 a score no other key approaches, and every output + row *is* V row 0. It comes back as e4m3(V[0]) with max abs deviation 0 + (99.98% of the raw bytes equal, the rest signed zeros), while the + unquantized V[0] a bf16-pool call would return is 60.5 bf16 ulps away. + 2. The two kv-scale factors land on it exactly as they land on the fresh + flavor: quantization at 1.0, then s**2 on the softmax scale and s on the + output. At s = 2.0 the peaked readout returns 2 * e4m3(V[0]) bit for + bit, and on realistic inputs the true-value model sits 48x outside the + band while the distorted one explains the run at 0.86 of it. So this + flavor is correct at s = 1.0 alone, like the other context flavor. + """ + h = MLA_NUM_HEADS_H128 + seq_len = 40 + for scale in KV_SCALE_SWEEP: + env = _fp8_r1_env(kv_scaling_factor=scale) + torch.manual_seed(700) + env.add_request(0, seq_len) + env.reserve_cache_pages(0, seq_len) + # Matching nope halves at magnitude 4 on key 0 only; k_pe zeroed so the + # roped tails contribute nothing to any score. + q = torch.zeros(seq_len, h, QK_HEAD_DIM, dtype=torch.bfloat16, device="cuda") + q[:, :, :QK_NOPE_HEAD_DIM] = 4.0 + k = torch.zeros(seq_len, h, QK_HEAD_DIM, dtype=torch.bfloat16, device="cuda") + k[0, :, :QK_NOPE_HEAD_DIM] = 4.0 + packed = torch.zeros( + seq_len, + h * (QK_NOPE_HEAD_DIM + V_HEAD_DIM), + dtype=torch.bfloat16, + device="cuda", + ) + packed[:, : h * QK_NOPE_HEAD_DIM] = k[:, :, :QK_NOPE_HEAD_DIM].reshape( + seq_len, h * QK_NOPE_HEAD_DIM + ) + v = _v_split_view(packed, h) + v.view(seq_len, h, V_HEAD_DIM).copy_( + torch.randn(seq_len, h, V_HEAD_DIM, dtype=torch.bfloat16, device="cuda") + ) + v0 = v.view(seq_len, h, V_HEAD_DIM)[0].clone() + + rows = env.call_context_no_append( + [0], + [seq_len], + [seq_len], + q.view(seq_len, h * QK_HEAD_DIM), + k.view(seq_len, h * QK_HEAD_DIM), + v, + ).view(seq_len, h, V_HEAD_DIM) + + expected = ( + (v0.float().to(torch.float8_e4m3fn).float() * scale) + .to(torch.bfloat16) + .unsqueeze(0) + .expand_as(rows) + ) + torch.testing.assert_close(rows, expected, rtol=2**-8, atol=0.0) # one bf16 ulp + exact = float( + (rows.reshape(-1).view(torch.uint8) == expected.reshape(-1).view(torch.uint8)) + .float() + .mean() + ) + assert exact > 0.99, ( + f"only {exact:.4f} of the readout is bit-exact {scale} * e4m3(V[0]) at s={scale}" + ) + rival = (v0.float() * scale).unsqueeze(0) + ulps = float(((rows.float() - rival).abs() / (2**-8 * rival.abs().clamp(min=2**-8))).max()) + assert ulps > 8.0, ( + f"the unquantized V row is only {ulps:.3g} bf16 ulps away at " + f"s={scale} — this probe cannot tell an e4m3 V operand from a bf16 one" + ) + + # Realistic inputs: which of the two models explains a run at s != 1.0. + new_lens = [32, 9] + kv_lens = [64, 24] + for scale in KV_SCALE_SWEEP: + env = _fp8_r1_env(kv_scaling_factor=scale) + torch.manual_seed(954) + for rid, total in enumerate(kv_lens): + env.add_request(rid, new_lens[rid]) + env.reserve_cache_pages(rid, total) + q = torch.randn(sum(new_lens), h * QK_HEAD_DIM, dtype=torch.bfloat16, device="cuda") + k, packed = _random_explicit_kv(sum(kv_lens), h) + v = _v_split_view(packed, h) + out = env.call_context_no_append([0, 1], new_lens, kv_lens, q, k, v) + operands = (q, k, v, new_lens, kv_lens) + + true_ref = _fp8_no_append_reference(env, *operands) + if scale == 1.0: + torch.testing.assert_close(out, true_ref, rtol=FP8_MLA_RTOL, atol=FP8_MLA_ATOL) + else: + _assert_far_outside_band( + out, + true_ref, + 10.0, + f"fp8 no-append MLA context at s={scale} against the s=1.0 math", + FP8_MLA_ATOL, + FP8_MLA_RTOL, + ) + # Same 1.5x allowance the fresh-prefill scale case uses: the + # distorted model lands inside the band (observed 0.86x at s=2.0). + distorted = _fp8_no_append_reference(env, *operands, kv_scale_factors=True) + assert _band_fraction(out, distorted, FP8_MLA_ATOL, FP8_MLA_RTOL) < 1.5, ( + f"the s**2 / s model does not explain the s={scale} run" + ) + + +def test_fp8_mla_quant_mode_extra_bits_h128() -> None: + """Quantization bits outside the KV-cache group ride along unread. + + A target derives quant_mode from its checkpoint's quant config and will not + get a bare 128: DeepSeek-R1-0528-FP4 with an fp8 KV cache produces + 1152 (FP8_KV_CACHE | FP8_1x128_128x128), and a checkpoint with fp8 QDQ + weights produces 384 (| FP8_QDQ). Both are bit-identical to 128 here — in + the output *and* in the pool — across all three MLA call flavors, so the + op reads only the KV-cache bit. Nothing checks the other bits either way; + this test is what says they are inert rather than assumed to be. + """ + lens = [96, 33] + prefill = [64] + new_lens, kv_lens = [32, 9], [64, 24] + baseline: Dict[str, torch.Tensor] = {} + for quant_mode in (QUANT_MODE_FP8_KV_CACHE, 1152, 384): + env = _fp8_r1_env(quant_mode=quant_mode) + torch.manual_seed(960) + for rid, ln in enumerate(lens): + env.add_request(rid, ln) + q, k, v, latent = _random_context_inputs(sum(lens), env.num_heads) + results = { + "context": env.call_context([0, 1], lens, q, k, v, latent).clone(), + "pool": env.pool.clone(), + } + + gen_env = _fp8_r1_env(quant_mode=quant_mode) + _fp8_mla_prefill(gen_env, prefill, 961) + torch.manual_seed(962) + gen_env.append_decode_latent( + 0, torch.randn(LATENT_DIM, dtype=torch.bfloat16, device="cuda") * 0.5 + ) + fused_q = ( + torch.randn(1, gen_env.num_heads * LATENT_DIM, dtype=torch.bfloat16, device="cuda") + * 0.3 + ) + results["generation"] = gen_env.call_generation([0], fused_q).clone() + + na_env = _fp8_r1_env(quant_mode=quant_mode) + torch.manual_seed(963) + for rid, total in enumerate(kv_lens): + na_env.add_request(rid, new_lens[rid]) + na_env.reserve_cache_pages(rid, total) + q_na = torch.randn( + sum(new_lens), + na_env.num_heads * QK_HEAD_DIM, + dtype=torch.bfloat16, + device="cuda", + ) + k_na, packed_na = _random_explicit_kv(sum(kv_lens), na_env.num_heads) + results["no_append"] = na_env.call_context_no_append( + [0, 1], + new_lens, + kv_lens, + q_na, + k_na, + _v_split_view(packed_na, na_env.num_heads), + ).clone() + + if quant_mode == QUANT_MODE_FP8_KV_CACHE: + baseline = results + continue + for key, value in results.items(): + assert _bitwise_equal(value, baseline[key]), ( + f"quant_mode={quant_mode} changed the {key} result against a " + f"bare {QUANT_MODE_FP8_KV_CACHE}" + ) + + +def test_fp8_mla_append_race_control_h128() -> None: + """Positive control on the pool byte comparison every append check above + rests on: it must be able to see a *racing* append. + + This entry documents one intermittent defect on its own write path — when + a call's new tokens do not map to distinct physical slots, the writes race + and leave torn cache rows, nondeterministically and with no error. That is + the arming sequence replayed here on the fp8 latent pool: a 96-token + prefill whose three absolute pages are all mapped onto one physical page, + so three tokens contend for every slot. Four repeats of the identical + seeded call must not all agree — measured on sm_100, every one of five + repeats differed from the first, by 11 177 to 12 513 of the pool's bytes. + + Then the same comparison over the certified geometry (distinct pages, the + R1-cell prefill), repeated in the same process, must be bitwise stable — + which is what makes "the append is bit-exact" a result rather than the + absence of a look. + """ + lens = [96] + armed: List[torch.Tensor] = [] + for _ in range(4): + env = _fp8_r1_env() + torch.manual_seed(964) + env.add_request(0, lens[0]) + env.pages[0] = [2, 2, 2] # every absolute page aliased onto one page + q, k, v, latent = _random_context_inputs(sum(lens), env.num_heads) + env.call_context([0], lens, q, k, v, latent) + armed.append(env.pool.view(torch.uint8).clone()) + assert not all(_bitwise_equal(p, armed[0]) for p in armed[1:]), ( + "the aliased-page append came back identical in four repeats — this " + "pool comparison cannot see a racing append at all" + ) + + stable: List[torch.Tensor] = [] + for _ in range(3): + env = _fp8_r1_env() + torch.manual_seed(965) + env.add_request(0, lens[0]) + q, k, v, latent = _random_context_inputs(sum(lens), env.num_heads) + env.call_context([0], lens, q, k, v, latent) + env.check_cache(0) + stable.append(env.pool.clone()) + for i, pool in enumerate(stable[1:], start=1): + assert _bitwise_equal(pool, stable[0]), ( + f"the certified fp8 R1-cell append is not reproducible: repeat {i} " + f"differs from repeat 0" + ) + + +# --------------------------------------------------------------------------- +# MLA generation at predicted_tokens_per_seq > 1 +# --------------------------------------------------------------------------- + +# The draft-chain lengths a DeepSeek-R1-0528 MTP target sweeps: max_draft_len +# 0/1/2/3, passed as predicted_tokens_per_seq = max_draft_len + 1. P = 1 is the +# regression check — it drives the code path the entry was already certified on. +MTP_SWEEP = [1, 2, 3, 4] + +# Per-key score values for the mask readout below, and per-draft-row query +# values. All of them are exact in bf16 and in e4m3 (3 mantissa bits), so the +# logits the kernel forms are the ones the fp32 reference forms and the only +# inexactness left in the readout is the kernel's own e4m3 handling of the +# softmax probabilities. +MTP_KEY_SCORES = [1.0, 1.5, 2.0, 2.5, 3.0, 3.5, 4.0, 4.5] +MTP_Q_SCORES = [2.0, 4.0, 6.0, 8.0] + +# Gate on the readout's in-set weights. Its output elements *are* the softmax +# probabilities, so against an exact fp32 softmax what is left is the kernel's +# own e4m3 handling of those probabilities before BMM2 — a relative error, +# hence a pure rtol with atol 0. Measured on sm_100 at H = 128, page 32 over +# P = 1..4, weights spanning 0.0061 to 0.0690: max relative deviation 5.5% +# against e4m3's 2**-4 = 6.25% single-rounding bound, so a gate at 2**-4 would +# sit at 88% of its allowance and be fragile, while 2**-3 (two e4m3 ulps) puts +# it at 44%. Over a bf16 pool the same readout deviates by at most 0.53%, so +# 2**-6 puts that at 34%. +# +# The entry's FP8_MLA_ATOL of 2**-3 must NOT be used here: every weight in this +# readout is smaller than that floor, so an assert_close carrying it is +# vacuous — measured, by replacing the expected weights with the rival +# full-mask model and watching the fp8 cases still pass. +# +# This gate is also deliberately not what separates the two candidate +# within-block masks: a full mask moves the in-set weights by at most 0.0042 +# here against a 0.0024 residual (1.8x), which no tolerance resolves. The mask +# is separated by the *excluded* columns, gated bitwise zero. +MTP_READOUT_FP8_RTOL = 2**-3 +MTP_READOUT_BF16_RTOL = 2**-6 + + +def _mtp_readout_env(lens: List[int], offsets: List[int], fp8: bool = True) -> _MlaPagedEnv: + """An MLA env whose cached latent rows turn the generation call's output + into a direct readout of the attention weights it used. + + Cache row j of sequence g carries compressed_kv = one-hot at column + offsets[g] + j and exactly one non-zero k_pe element. A query whose + compressed-kv half is zero therefore scores only through k_pe, while the + decode's V is K[:, :C] — the one-hot — so + + output[row, h, offsets[g] + j] = the weight `row` put on key j + + with an exact zero on every key the mask excluded. The sequences' key + blocks are disjoint column ranges, so a row's non-zero columns also say + which sequence the kernel assigned it to. + """ + for ln, off in zip(lens, offsets): + assert off + ln <= KV_LORA_RANK, "readout columns must fit in C" + env = _MlaPagedEnv( + num_heads=MLA_NUM_HEADS_H128, + max_batch=max(4, len(lens)), + max_blocks_per_seq=MLA_MAX_SEQ_LEN // MLA_PAGE32, + tokens_per_block=MLA_PAGE32, + q_lora_rank=Q_LORA_RANK_DSV3, + pool_dtype=torch.float8_e4m3fn if fp8 else torch.bfloat16, + quant_mode=QUANT_MODE_FP8_KV_CACHE if fp8 else 0, + kv_scaling_factor=1.0 if fp8 else None, + ) + for g, ln in enumerate(lens): + env.add_request(g, ln) + for j in range(ln): + row = torch.zeros(LATENT_DIM, dtype=torch.bfloat16, device="cuda") + row[offsets[g] + j] = 1.0 + row[KV_LORA_RANK] = MTP_KEY_SCORES[j % len(MTP_KEY_SCORES)] + env.append_decode_latent(g, row) + return env + + +def _mtp_readout_q(rows: int, num_heads: int) -> torch.Tensor: + """Query rows for the readout env: the compressed-kv half is zero, so the + one-hot keys never enter a score, and one q_pe element carries a per-row + value, so the P rows of a draft block do not share one softmax.""" + q = torch.zeros(rows, num_heads, LATENT_DIM, dtype=torch.bfloat16, device="cuda") + for n in range(rows): + q[n, :, KV_LORA_RANK] = MTP_Q_SCORES[n % len(MTP_Q_SCORES)] + return q.view(rows, num_heads * LATENT_DIM) + + +def _mtp_expected_weights(env: _MlaPagedEnv, row: int, num_keys: int) -> torch.Tensor: + """The exact fp32 softmax the readout env's query row puts on num_keys.""" + z = torch.tensor( + [MTP_KEY_SCORES[j % len(MTP_KEY_SCORES)] for j in range(num_keys)], + dtype=torch.float32, + device="cuda", + ) + return torch.softmax(MTP_Q_SCORES[row % len(MTP_Q_SCORES)] * z * env.softmax_scale, dim=0) + + +def _assert_mtp_causal_readout( + env: _MlaPagedEnv, + out: torch.Tensor, + lens: List[int], + offsets: List[int], + p: int, + label: str, + rtol: float = MTP_READOUT_FP8_RTOL, +) -> None: + """Read every draft row's attended key set out of `out`, all heads. + + Three assertions per row, and the rival within-block mask — row t sees its + later siblings too — breaks all three: + + - every column outside [0, L - P + t] of the sequence's own key block is + *bitwise zero*, in particular the columns of draft tokens t+1..P-1, which + do not exist yet when row t is verified; + - so is every column of the other sequences' key blocks, which is what pins + the token-major row order; + - the weights inside the set match the exact fp32 softmax, at the pure-rtol + gate above. + + The gate on the excluded columns is exact zero rather than a tolerance, so + the separation from the rival is absolute; the assert at the end of the + loop keeps that meaningful by checking the rival would in fact have put + non-trivial mass there. + """ + heads = env.num_heads + for n in range(len(lens) * p): + g, t = n // p, n % p + ln = lens[g] + last = ln - p + t # inclusive index of the newest key row n may see + weights = out[n].view(heads, KV_LORA_RANK).float() + expected = _mtp_expected_weights(env, n, last + 1) + torch.testing.assert_close( + weights[:, offsets[g] : offsets[g] + last + 1], + expected.unsqueeze(0).expand(heads, -1), + rtol=rtol, + atol=0.0, + ) + outside = torch.ones(KV_LORA_RANK, dtype=torch.bool, device="cuda") + outside[offsets[g] : offsets[g] + last + 1] = False + stray = int((weights[:, outside] != 0).sum()) + assert stray == 0, ( + f"{label}: row {n} (sequence {g}, draft token {t} of {p}) put " + f"non-zero weight on {stray} key column(s) it must not see" + ) + if t + 1 < p: + rival = _mtp_expected_weights(env, n, ln) + diverted = float(rival[last + 1 :].sum()) + assert diverted > 0.01, ( + f"{label}: row {n}'s future siblings would carry only " + f"{diverted:.4g} of the mass under a full within-block mask, " + f"so this readout could not separate the two masks here" + ) + + +def test_fp8_mla_mtp_mask_is_bottom_right_causal_h128() -> None: + """What a draft row of an MLA generation call attends to at P > 1. + + This is the question the whole MTP surface turns on. At P = 1 there is no + within-block mask to get wrong, and mask_type = 1 simply means "attend to + everything cached". At P > 1 the P query rows of one sequence are its draft + chain: row t sits at absolute position L - P + t and must attend to the + cache up to and including its own position, and *not* to rows t+1..P-1 of + its own block, which are tokens that do not exist yet. If the kernel let + row t see row t+1, every verification past the first would be computed + against a future token; nothing raises, and rejection sampling would still + emit correct text with the drafts always rejected, so no downstream + accuracy gate could see it either. + + On sm_100 no mask tensor does this job — a linear-tree draft has + is_spec_decoding_enabled forced off there, so spec_decoding_packed_mask and + its siblings are all None. Whatever masking happens comes from + predicted_tokens_per_seq alone. + + Measured, at P = 1, 2, 3 and 4 over the fp8-e4m3 latent pool at the R1 + layer geometry, by reading the attention weights straight out of the + output: the mask **is** causal and bottom-right aligned against the + sequence's own KV length. Row t weights keys [0, L - P + t] and every other + column comes back bitwise zero, including the 1 to 3 future-sibling columns + that a full within-block mask would have weighted at 0.023-0.039 each. + + Two sequences of different lengths (33 and 50 cached rows, neither a + multiple of the 32-slot page, so the causal cut crosses a page boundary) + with disjoint readout column blocks, which also pins the row order: rows + [g*P, (g+1)*P) belong to sequence g, token-major. + """ + lens = [33, 50] + offsets = [0, KV_LORA_RANK // 2] + for p in MTP_SWEEP: + torch.manual_seed(1100 + p) + env = _mtp_readout_env(lens, offsets) + fused_q = _mtp_readout_q(len(lens) * p, env.num_heads) + pool_before = env.pool.clone() + out = env.call_generation([0, 1], fused_q, predicted_tokens_per_seq=p) + _assert_mtp_causal_readout(env, out, lens, offsets, p, f"fp8 P={p}") + assert _bitwise_equal(pool_before, env.pool), ( + "the MLA generation call wrote to the fp8 pool" + ) + + +def test_fp8_mla_mtp_mask_type_and_page_alignment_h128() -> None: + """Two follow-ups to the mask readout, on the same fp8 R1 geometry. + + 1. mask_type = 0 (padding) does **not** differ from mask_type = 1 here at + P > 1: both return the same bottom-right-causal attended sets, bitwise. + The entry certifies mask_type = 0 for context flavors only, and the + obvious guess — that padding would drop the within-block mask and let a + draft row see its later siblings — is wrong. mask_type does not reach + this path. + 2. The causal cut lands at L - P + t, so it walks across a 32-slot page + boundary as L moves. Swept at P = 4 over L = 30..36, which puts the cut + one slot before a boundary, exactly on it, and one slot after, and also + covers L = P + small (a sequence whose whole cached history is barely + longer than this step's own draft chain). + """ + lens = [33, 50] + offsets = [0, KV_LORA_RANK // 2] + for p in (2, 4): + torch.manual_seed(1200 + p) + env = _mtp_readout_env(lens, offsets) + fused_q = _mtp_readout_q(len(lens) * p, env.num_heads) + causal = env.call_generation([0, 1], fused_q, predicted_tokens_per_seq=p) + padding = env.call_generation( + [0, 1], fused_q, predicted_tokens_per_seq=p, mask_type=MASK_PADDING + ) + _assert_mtp_causal_readout(env, padding, lens, offsets, p, f"fp8 mask_type=0 P={p}") + assert _bitwise_equal(causal, padding), ( + f"mask_type changed the MLA generation result at P={p} — the two " + f"attended sets agree but not bitwise" + ) + + p = 4 + for ln in range(30, 37): + torch.manual_seed(1300 + ln) + env = _mtp_readout_env([ln], [0]) + fused_q = _mtp_readout_q(p, env.num_heads) + out = env.call_generation([0], fused_q, predicted_tokens_per_seq=p) + _assert_mtp_causal_readout(env, out, [ln], [0], p, f"fp8 page alignment L={ln} P={p}") + + +def test_fp8_mla_mtp_r1_cell_decode_h128() -> None: + """Realistic-input MLA decode at the complete DeepSeek-R1-0528 cell over + the fp8-e4m3 latent pool, swept over P = 1, 2, 3, 4. + + Random latent history and random fused q, two consecutive steps per P, two + sequences whose KV lengths sit at different page-32 alignments. The gate is + the fp32 latent-MQA reference under the measured bottom-right-causal mask, + at the fp8 band. + + The rival full-mask model is gated *relatively* rather than absolutely: at + these shapes it adds only 1-3 keys to a set of 30-70, and the fp8 band is + wide enough that the rival sometimes stays inside it (measured 0.92-1.57 of + the allowance over the seeds here, against 0.21-0.24 for the matching + model), so "the rival fails assert_close" would not be a stable statement. + What is stable is that the rival uses several times more of the allowance — + 4.4-6.9x measured, gated at 3x. The absolute mask evidence is the + bitwise-zero readout above and the bf16-pool case below (17-23x outside its + own band); what this case gates is the arithmetic of a production-shaped + decode. + """ + for p in MTP_SWEEP: + torch.manual_seed(1400 + p) + env = _fp8_r1_env() + prefill = [64, 31] + _fp8_mla_prefill(env, prefill, 1400 + p) + for rid in (0, 1): + env.check_cache(rid) + for _ in range(2): + for rid in (0, 1): + for _ in range(p): + env.append_decode_latent( + rid, + torch.randn(LATENT_DIM, dtype=torch.bfloat16, device="cuda") * 0.5, + ) + pool_before = env.pool.clone() + fused_q = ( + torch.randn( + 2 * p, + env.num_heads * LATENT_DIM, + dtype=torch.bfloat16, + device="cuda", + ) + * 0.3 + ) + out = env.call_generation([0, 1], fused_q, predicted_tokens_per_seq=p) + ref = env.generation_reference([0, 1], fused_q, predicted_tokens_per_seq=p) + torch.testing.assert_close(out, ref, rtol=FP8_MLA_RTOL, atol=FP8_MLA_ATOL) + assert _bitwise_equal(pool_before, env.pool), ( + "the MLA generation call wrote to the fp8 pool" + ) + if p > 1: + matched = _band_fraction(out, ref, FP8_MLA_ATOL, FP8_MLA_RTOL) + rival = _band_fraction( + out, + env.generation_reference( + [0, 1], fused_q, predicted_tokens_per_seq=p, full_mask=True + ), + FP8_MLA_ATOL, + FP8_MLA_RTOL, + ) + assert rival > 3.0 * matched, ( + f"fp8 R1-cell MTP decode at P={p}: the full-mask rival uses " + f"{rival:.3g} of the allowance against the matching model's " + f"{matched:.3g} — only {rival / matched:.3g}x apart" + ) + + +def test_fp8_mla_mtp_batch_state_h128() -> None: + """Which batch-state tensors the MLA generation call reads at P > 1. + + The entry records, at P = 1, that cu_q_seqlens is required by presence but + bitwise inert in its contents, that cu_kv_seqlens is inert outright, and + that the per-sequence KV length comes from sequence_length. A taller query + block is exactly the change that could have made the scheduler buffers + start mattering, so all of it is re-measured at P = 3. + + Measured: cu_q_seqlens contents stay inert — zeroed, left in the P = 1 + i*H form, given i*P without the head factor, inflated 1000x and made + non-monotonic all return bitwise-identical output. So do cu_kv_seqlens and + both context_lengths copies. sequence_length remains the KV extent: a + one-token change moves every draft row's attended set by one, read out + one-hot. host_past_key_value_lengths is inert as long as it is non-zero + (sequence_length - 1, - P and all-ones are all bitwise identical), but an + all-zero vector makes the call skip its work entirely and return without + writing a single element of `output` — not a P > 1 effect, it reproduces + at P = 1. + + Also gated: the call writes exactly the first G*P rows of `output` and + leaves a taller buffer's tail untouched. + """ + lens = [33, 50] + p = 3 + offsets = [0, KV_LORA_RANK // 2] + torch.manual_seed(1500) + env = _fp8_r1_env() + for rid, ln in enumerate(lens): + env.add_request(rid, ln - p) + for _ in range(ln): + env.append_decode_latent( + rid, torch.randn(LATENT_DIM, dtype=torch.bfloat16, device="cuda") * 0.5 + ) + fused_q = ( + torch.randn( + len(lens) * p, + env.num_heads * LATENT_DIM, + dtype=torch.bfloat16, + device="cuda", + ) + * 0.3 + ) + base = env.call_generation([0, 1], fused_q, predicted_tokens_per_seq=p) + assert _bitwise_equal(base, env.call_generation([0, 1], fused_q, predicted_tokens_per_seq=p)), ( + "the MLA generation call is not run-to-run deterministic here" + ) + + h, g = env.num_heads, len(lens) + for label, cu_q in ( + ("zeroed", torch.zeros(g + 1, dtype=torch.int32)), + ("the P=1 form i*H", torch.arange(g + 1, dtype=torch.int32) * h), + ("i*P without the head factor", torch.arange(g + 1, dtype=torch.int32) * p), + ("1000x", torch.arange(g + 1, dtype=torch.int32) * (h * p) * 1000), + ("non-monotonic", torch.tensor([0, 7 * h * p, 2 * h * p], dtype=torch.int32)), + ): + assert _bitwise_equal( + base, + env.call_generation([0, 1], fused_q, predicted_tokens_per_seq=p, cu_q_override=cu_q), + ), f"cu_q_seqlens {label} moved the output at P={p}" + + for label, cu_kv in ( + ("zeroed", torch.zeros(g + 1, dtype=torch.int32)), + ("1000x", torch.tensor([0, 33000, 83000], dtype=torch.int32)), + ("non-monotonic", torch.tensor([0, 90, 5], dtype=torch.int32)), + ): + assert _bitwise_equal( + base, + env.call_generation([0, 1], fused_q, predicted_tokens_per_seq=p, cu_kv_override=cu_kv), + ), f"cu_kv_seqlens {label} moved the output at P={p}" + + for label, host_past in ( + ("sequence_length - 1", [lens[0] - 1, lens[1] - 1]), + ("sequence_length - P", [lens[0] - p, lens[1] - p]), + ("all ones", [1, 1]), + ): + assert _bitwise_equal( + base, + env.call_generation( + [0, 1], + fused_q, + predicted_tokens_per_seq=p, + host_past_override=host_past, + ), + ), f"host_past_key_value_lengths = {label} moved the output at P={p}" + + for label, ctx_lens in (("zeroed", [0, 0]), ("the full KV length", list(lens))): + assert _bitwise_equal( + base, + env.call_generation( + [0, 1], fused_q, predicted_tokens_per_seq=p, ctx_lens_override=ctx_lens + ), + ), f"context_lengths {label} moved the output at P={p}" + + # An all-zero host_past_key_value_lengths is the one non-inert perturbation + # found, and what it does is skip the work entirely: the call returns + # without writing a single element of `output`, whatever that buffer held. + # Stated against a sentinel fill rather than against zero, so it is a claim + # about "not written" and not about a value that happens to be zero. + sentinel_fill = torch.full( + (len(lens) * p, env.num_heads * KV_LORA_RANK), + 7.0, + dtype=torch.bfloat16, + device="cuda", + ) + untouched = sentinel_fill.clone() + env.call_generation( + [0, 1], + fused_q, + predicted_tokens_per_seq=p, + host_past_override=[0, 0], + output_buffer=untouched, + ) + assert _bitwise_equal(untouched, sentinel_fill), ( + "an all-zero host_past_key_value_lengths no longer skips the call — " + "this control has stopped measuring what it claims" + ) + written = sentinel_fill.clone() + env.call_generation([0, 1], fused_q, predicted_tokens_per_seq=p, output_buffer=written) + assert _bitwise_equal(written, base), ( + "the same call with a true host_past_key_value_lengths did not write the certified result" + ) + + # sequence_length is the KV extent: shift it down and every row's attended + # set shifts with it, keeping the L - P + t shape. + for delta in (0, -1, -2): + shifted = _mtp_readout_env([lens[0] + delta, lens[1] + delta], offsets) + out = shifted.call_generation( + [0, 1], + _mtp_readout_q(len(lens) * p, shifted.num_heads), + predicted_tokens_per_seq=p, + ) + _assert_mtp_causal_readout( + shifted, + out, + [lens[0] + delta, lens[1] + delta], + offsets, + p, + f"sequence_length shifted by {delta}", + ) + + tall = torch.full( + (len(lens) * p + 4, env.num_heads * KV_LORA_RANK), + 7.0, + dtype=torch.bfloat16, + device="cuda", + ) + sentinel = tall[len(lens) * p :].clone() + env.call_generation([0, 1], fused_q, predicted_tokens_per_seq=p, output_buffer=tall) + assert _bitwise_equal(tall[len(lens) * p :], sentinel), ( + "the MLA generation call wrote past row G*P of `output`" + ) + assert _bitwise_equal(tall[: len(lens) * p], base), ( + "writing into a taller output buffer changed the result" + ) + + +def test_fp8_mla_mtp_mixed_batch_h128() -> None: + """A mixed batch whose generation sequences each carry P draft tokens. + + One leading context sequence plus one generation sequence, two phase calls + sharing the full-batch state tensors and one page-32 offsets table, at the + complete R1 cell over the fp8 pool. The context call's token accounting is + unchanged by P — it owns num_ctx_tokens rows — while the generation call + owns G*P rows indexed from num_contexts, and the two must still agree. + Swept over P = 2, 3, 4. + """ + for p in (2, 3, 4): + torch.manual_seed(1600 + p) + env = _fp8_r1_env() + first_len, second_len = 40, 33 + env.add_request(0, first_len) + q, k, v, latent = _random_context_inputs(first_len, env.num_heads) + env.call_context([0], [first_len], q, k, v, latent) + + env.add_request(1, second_len) + for _ in range(p): + env.append_decode_latent( + 0, torch.randn(LATENT_DIM, dtype=torch.bfloat16, device="cuda") * 0.5 + ) + q, k, v, latent = _random_context_inputs(second_len, env.num_heads) + pre = (q.clone(), k.clone(), v, latent.clone()) + out_ctx = env.call_context([1], [second_len], q, k, v, latent, gen_request_ids=[0]) + torch.testing.assert_close( + out_ctx, + env.context_reference([second_len], *pre, e4m3_inputs=True), + rtol=FP8_MLA_RTOL, + atol=FP8_MLA_ATOL, + ) + + fused_q = ( + torch.randn(p, env.num_heads * LATENT_DIM, dtype=torch.bfloat16, device="cuda") * 0.3 + ) + out_gen = env.call_generation( + [0], + fused_q, + ctx_request_ids=[1], + ctx_seq_lens=[second_len], + predicted_tokens_per_seq=p, + ) + torch.testing.assert_close( + out_gen, + env.generation_reference([0], fused_q, predicted_tokens_per_seq=p), + rtol=FP8_MLA_RTOL, + atol=FP8_MLA_ATOL, + ) + env.check_cache(0) + env.check_cache(1) + env.check_unwritten_pool_zero([0, 1]) + + +def test_bf16_mla_mtp_decode_h128() -> None: + """The same MTP generation surface over a bf16 latent pool. + + Two things this adds over the fp8 cases. The bf16 band is 8x tighter, so + the full-mask rival separates properly here — observed 17-23x outside, + against 0.22 for the matching model, which is what makes "the within-block + mask is causal" a gated result on realistic inputs and not only a readout. + And it pins that P > 1 is a property of the MLA decode kernel rather than + of the fp8 decode path, which is the only one the R1 target runs. + """ + lens = [33, 50] + offsets = [0, KV_LORA_RANK // 2] + for p in (2, 4): + torch.manual_seed(1700 + p) + env = _mtp_readout_env(lens, offsets, fp8=False) + fused_q = _mtp_readout_q(len(lens) * p, env.num_heads) + out = env.call_generation([0, 1], fused_q, predicted_tokens_per_seq=p) + _assert_mtp_causal_readout( + env, out, lens, offsets, p, f"bf16 P={p}", rtol=MTP_READOUT_BF16_RTOL + ) + + torch.manual_seed(1750 + p) + real = _MlaPagedEnv( + num_heads=MLA_NUM_HEADS_H128, + max_blocks_per_seq=MLA_MAX_SEQ_LEN // MLA_PAGE32, + tokens_per_block=MLA_PAGE32, + q_lora_rank=Q_LORA_RANK_DSV3, + ) + for rid, ln in enumerate(lens): + real.add_request(rid, ln - p) + for _ in range(ln): + real.append_decode_latent( + rid, + torch.randn(LATENT_DIM, dtype=torch.bfloat16, device="cuda") * 0.5, + ) + fq = ( + torch.randn( + len(lens) * p, + real.num_heads * LATENT_DIM, + dtype=torch.bfloat16, + device="cuda", + ) + * 0.3 + ) + pool_before = real.pool.clone() + got = real.call_generation([0, 1], fq, predicted_tokens_per_seq=p) + torch.testing.assert_close( + got, + real.generation_reference([0, 1], fq, predicted_tokens_per_seq=p), + rtol=RTOL, + atol=ATOL, + ) + _assert_far_outside_band( + got, + real.generation_reference([0, 1], fq, predicted_tokens_per_seq=p, full_mask=True), + 5.0, + f"bf16 MTP decode at P={p} against a full-mask rival", + ATOL, + RTOL, + ) + assert _bitwise_equal(pool_before, real.pool), ( + "the MLA generation call wrote to the bf16 pool" + ) + + +def test_fp8_mla_mtp_decode_sees_a_torn_pool_h128() -> None: + """Positive control on the decode-side comparison every claim above rests + on: it must be able to see a pool whose content moved under it. + + Every MTP result here reads a cache the caller wrote, and this entry + documents one intermittent defect on that write path — a call whose new + tokens do not map to distinct physical slots races and leaves torn rows, + nondeterministically and with no error. At P > 1 the caller's + generation-phase preprocessing writes P rows per sequence per step instead + of one, so that arming sequence is newly reachable from a decode step. + + Armed here the way the append control arms it: a 96-token prefill whose + three absolute pages are mapped onto one physical page. Four repeats must + not all leave the same pool — and, the part this control adds, the P = 3 + decode run over those differing pools must not all return the same output, + which is what makes a clean decode result a measurement rather than the + absence of a look. Then the certified geometry, in the same process: + distinct pages, identical seeds, pool and decode output both bitwise + stable across repeats. + """ + p = 3 + lens = [96] + torch.manual_seed(1800) + probe_q = ( + torch.randn(p, MLA_NUM_HEADS_H128 * LATENT_DIM, dtype=torch.bfloat16, device="cuda") * 0.3 + ) + + armed_pools: List[torch.Tensor] = [] + armed_outs: List[torch.Tensor] = [] + for _ in range(4): + env = _fp8_r1_env() + torch.manual_seed(1801) + env.add_request(0, lens[0]) + env.pages[0] = [2, 2, 2] # every absolute page aliased onto one page + q, k, v, latent = _random_context_inputs(sum(lens), env.num_heads) + env.call_context([0], lens, q, k, v, latent) + armed_pools.append(env.pool.view(torch.uint8).clone()) + armed_outs.append(env.call_generation([0], probe_q, predicted_tokens_per_seq=p).clone()) + assert not all(_bitwise_equal(x, armed_pools[0]) for x in armed_pools[1:]), ( + "the aliased-page append came back identical in four repeats — this " + "control cannot arm the defect it is meant to arm" + ) + assert not all(_bitwise_equal(x, armed_outs[0]) for x in armed_outs[1:]), ( + "the P>1 decode returned the same output over four demonstrably " + "different pools — the decode-side comparison is blind to pool content" + ) + + stable_pools: List[torch.Tensor] = [] + stable_outs: List[torch.Tensor] = [] + for _ in range(3): + env = _fp8_r1_env() + torch.manual_seed(1802) + env.add_request(0, lens[0]) + q, k, v, latent = _random_context_inputs(sum(lens), env.num_heads) + env.call_context([0], lens, q, k, v, latent) + env.check_cache(0) + stable_pools.append(env.pool.clone()) + stable_outs.append(env.call_generation([0], probe_q, predicted_tokens_per_seq=p).clone()) + for i in range(1, 3): + assert _bitwise_equal(stable_pools[i], stable_pools[0]), ( + f"the certified fp8 R1-cell append is not reproducible: repeat {i}" + ) + assert _bitwise_equal(stable_outs[i], stable_outs[0]), ( + f"the certified P>1 decode is not reproducible: repeat {i}" + ) + + +def test_mla_rejects_null_q_lora_rank() -> None: + """q_lora_rank must be an int on the MLA path: the C++ unwraps the + optional unconditionally, so None raises rather than defaulting.""" + torch.manual_seed(309) + h = MLA_NUM_HEADS_H32 + q, k, v, latent = _random_context_inputs(32, h) + env = _MlaPagedEnv( + num_heads=h, + max_blocks_per_seq=MLA_MAX_SEQ_LEN // MLA_PAGE32, + tokens_per_block=MLA_PAGE32, + q_lora_rank=None, + ) + env.add_request(0, 32) + try: + env.call_context([0], [32], q, k, v, latent) + except RuntimeError as exc: + assert "bad optional access" in str(exc), f"unexpected message: {exc}" + else: + raise AssertionError("q_lora_rank=None was accepted on the MLA path") diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/__init__.py b/tensorrt_llm/_torch/staircase/catalog/comm/__init__.py new file mode 100644 index 000000000000..f62792f5e815 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/comm/__init__.py @@ -0,0 +1,3 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Collective-communication entries.""" diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/_rank_job.py b/tensorrt_llm/_torch/staircase/catalog/comm/_rank_job.py new file mode 100644 index 000000000000..c4a9510746e0 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/comm/_rank_job.py @@ -0,0 +1,88 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Run a collective entry's rank job from pytest, and report what it did. + +The entry's own launcher already owns the hard parts -- one rank per visible +device, its own process group, and a deadline it enforces by killing that +group, which is what keeps a wedged collective from taking the calling run +down with it. A collective that breaks *hangs* rather than raising, so the +deadline is load-bearing and the ``mpi_pool_executor`` fixture has none. + +So this does not reimplement any of that. It selects the devices, starts the +launcher in a fresh interpreter, and turns its exit code into an assertion. + +Fresh interpreter is required, not tidiness: the launcher must not have +initialized MPI, and a pytest process that has imported ``tensorrt_llm`` +already has. + +For the same reason the launcher is started **by file path, not by +``-m``**. ``-m`` on a module inside this package imports every parent +package on the way to it, ``tensorrt_llm`` included, and that import calls +``MPI_Init``; the launcher would then fail to start ``mpirun`` at all +(measured: it exits 1 with no output from any rank). Run by path the file +has no package context and imports nothing but torch, which is all its +launcher half needs -- the entry keeps its relative imports inside +``_run_one_rank``, and the ranks it spawns *are* started with ``-m`` so +they get one. +""" + +from __future__ import annotations + +import os +import subprocess +import sys +from importlib import import_module +from pathlib import Path + +WORLD_SIZE = 4 +"""dep4's world size -- the topology these entries are certified for.""" + +_LAUNCHER_GRACE_S = 300 +"""Headroom over the entry's own deadline, so its message wins the race.""" + + +def _devices() -> str: + visible = os.environ.get("CUDA_VISIBLE_DEVICES") + if visible: + devices = [d for d in visible.split(",") if d.strip()] + else: + import torch + + devices = [str(i) for i in range(torch.cuda.device_count())] + assert len(devices) >= WORLD_SIZE, ( + f"this entry is certified at world size {WORLD_SIZE}; only " + f"{len(devices)} device(s) are visible" + ) + return ",".join(devices[:WORLD_SIZE]) + + +def run(entry: str) -> None: + """Run ``_test``'s launcher over ``WORLD_SIZE`` devices.""" + module = f"{__package__}.{entry}_test" + launcher = Path(__file__).with_name(f"{entry}_test.py") + env = dict(os.environ, CUDA_VISIBLE_DEVICES=_devices()) + # The launcher re-execs itself per rank and needs this package importable + # from the ranks; by path it has no package context of its own to inherit. + repo_root = Path(__file__).resolve().parents[5] + env["PYTHONPATH"] = os.pathsep.join(p for p in (str(repo_root), env.get("PYTHONPATH", "")) if p) + + # Read the budget off the entry rather than restating it: an entry that + # raises its own deadline would otherwise be killed by this one first, + # and the message a reader needs ("wedged") would be lost. reducescatter + # runs a second, separately capped job after the main one. + entry_module = import_module(module) + timeout = entry_module.DEADLINE_S + getattr(entry_module, "WEDGE_CAP_S", 0) + _LAUNCHER_GRACE_S + + completed = subprocess.run( + [sys.executable, str(launcher)], + env=env, + capture_output=True, + text=True, + timeout=timeout, + ) + if completed.returncode != 0: + raise AssertionError( + f"{launcher.name} exited {completed.returncode}\n" + f"--- stdout ---\n{completed.stdout}\n" + f"--- stderr ---\n{completed.stderr}" + ) diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/allgather.md b/tensorrt_llm/_torch/staircase/catalog/comm/allgather.md new file mode 100644 index 000000000000..dd289ecc9652 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/comm/allgather.md @@ -0,0 +1,379 @@ +--- +receipts: + sm_100: {status: passed, trtllm: 1.3.0rc21, world_size: 4} + sm_103: {status: passed, trtllm: 1.3.0rc26, world_size: 4} +--- + +# allgather + +**Wraps** `torch.ops.trtllm.allgather` (one call). + +## Semantics + +Every rank in `group` calls with its own slice of a tensor; every rank gets +back the whole thing, concatenated along **dim 0 in ascending rank order**: + +``` +out = cat([input_r0, input_r1, ..., input_r(G-1)], dim=0) +``` + +It is **not** in place — the result is a freshly allocated tensor, and +`input` is left untouched, so clobbering `input` right after the call +cannot reach the result. + +Only dim 0 is a gather axis. Every other dimension is carried through +untouched: `[n, 2560]` gathers to `[G*n, 2560]`, `[n, 2, 1280]` to +`[G*n, 2, 1280]`, and a 1-D `[n]` to `[G*n]`. Nothing is computed — the +call moves bytes, so the result is bit-for-bit the inputs' concatenation +in every dtype below. + +The `sizes` argument picks between two forms: + +**`sizes = None`** — every rank must hold the same number of rows `n`, and +`out.shape[0] = n * len(group)`. + +**`sizes = [n_0, ..., n_(G-1)]`** — rank `i` contributes `n_i` rows and +`out.shape[0] = sum(sizes)`. The list is in ascending rank order, must be +identical on every rank, and `sizes[my_rank]` must equal this rank's +`input.shape[0]`. Entries may be `0`: a rank with no rows contributes +nothing and the gather is still correct. This is the form attention data +parallelism needs, where per-rank token counts differ by construction. + +Both forms are certified, and they are certified separately because they are +not the same call underneath — an all-gather primitive takes one element +count. Read off NCCL's own trace of this entry's test: `sizes = None` +issues **one `ncclAllGather`** whose count is this rank's element count, +and the ragged form issues **one grouped `ncclBroadcast` per rank**, each +rooted at that rank with that rank's element count (for +`sizes = [1, 5, 9, 13]` at hidden 2560: counts 2560, 12800, 23040, 33280 at +roots 0, 1, 2, 3). So one ragged call enqueues `len(group)` NCCL operations +where a uniform call enqueues one — worth knowing when reasoning about the +call-order precondition, though the two forms were never mixed across ranks +in a probe, and the `sizes` list already has to be identical everywhere. +The two cost the same (see *Notes*), so the choice between them is about +what the caller can guarantee, not about speed. + +**Fusion boundary.** Inside the call: the collective and the output +allocation, nothing else. Outside: everything that produced `input` +(under attention DP, this rank's own attention output), any padding of +row counts to a uniform value, any quantization, and any reshaping — +including the reshape a caller needs to gather along an axis other than +dim 0, which this op cannot do. + +**Group membership.** `group` names **ranks in trtllm's MPI session +communicator** (`MPI_COMM_WORLD` under `mpirun`), not device ordinals or +`torch.distributed` ranks. A subset is legal: with `group = [0, 1]` at +world size 4, ranks 0 and 1 gather only between themselves and ranks 2 +and 3 must not call at all. The list is treated as a **set** — passing +`[3, 2, 1, 0]` produces the identical result to `[0, 1, 2, 3]`, still in +ascending rank order — so `group` cannot be used to permute the output. + +## Signature + +```python +def allgather( + input: torch.Tensor, + sizes: Optional[List[int]], + group: List[int], +) -> torch.Tensor +``` + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `input` | `[rows, ...]`, at least 1-D; `[T, 2560]`, `[T, 2, 1280]` and `[T]` tested | bf16; also fp16, fp32, int32, uint8, float8_e4m3fn | **contiguous** | CUDA | +| `sizes` | `None`, or a list of `len(group)` non-negative ints in ascending rank order | Python `int` | — | host | +| `group` | list of MPI session ranks | Python `int` | — | host | +| returns | `[n * len(group), *input.shape[1:]]` for `sizes = None`, `[sum(sizes), *input.shape[1:]]` otherwise | = `input.dtype` | contiguous, newly allocated | = `input.device` | + +There are no other arguments — no strategy, no workspace, no output +buffer, no autotuner. Nothing selects a transport, so unlike the sibling +all-reduce there is no strategy axis to certify. + +Every dtype in the table was certified in both forms (uniform and ragged) +as a bitwise move: bf16 for hidden states, uint8 for packed NVFP4 bytes +and their scale factors, float8_e4m3fn for fp8 activations, int32 for +expert ids, fp32/fp16 for routing weights and non-bf16 activations. int64 +and bool were also observed to move correctly in probing and are not +certified. There is no arithmetic here for a dtype to be wrong about — +this op does not reinterpret bytes the way a summing collective does. + +## Metadata consumed + +No attention metadata, no KV cache, no registered layer, no workspace. +Two pieces of process state stand behind the call: + +1. **trtllm's MPI session communicator.** The op resolves `group` against + it. Under `mpirun` the session is `MPI_COMM_WORLD`; a trtllm engine's + worker ranks are already inside one, so a target's forward needs no + setup. World size 1 was not exercised; this entry's receipt is world + size 4. The sibling op `torch.ops.trtllm.allgather_pg` takes a + `c10d.ProcessGroup` instead, for the `TLLM_DISABLE_MPI=1` path; + `torch.ops.trtllm.allgather_list` gathers a list of tensors in one + call. Both are different ops and neither is this entry. +2. **The NCCL communicator cache**, keyed by the rank set. Built on first + use for a given `group`, which makes the first call for a group much + slower than the rest — and means that first call **cannot be inside a + CUDA-graph capture** (see *Preconditions*). Inside a trtllm serving + engine this communicator is the **caller's alone**: the runtime issues + nothing on it, so the only calls whose order has to line up are the + ones the caller makes. Measured — see *Notes*, "Inside a live serving + engine under attention data parallelism". + +## Preconditions + +- **Every rank in `group` calls, and their arguments agree**: same dtype, + same trailing shape, same `sizes` list, and — when `sizes is None` — + the same row count. Disagreement is **not survivable**: it hangs rather + than raising. Measured at world size 4, three ways, each in its own + process: ranks passing different row counts with `sizes = None`; a rank + holding more rows than its `sizes` entry; a rank holding fewer. In each + case some ranks returned a plausible-looking tensor and the rest never + came back, and the job had to be killed. This is why this entry's test + carries an external deadline. Do not read the hang as a guarantee, + though: every one of those three cases changes how many bytes a rank + moves, and a disagreement that leaves the byte counts equal on every rank + is silently wrong instead — see the call-order bullet below, where that + is measured. +- `input` is **contiguous**. The op reads `input.numel()` elements from + `input.data_ptr()` and ignores strides, so a strided view is gathered + from the wrong bytes and returns a plausible-looking wrong answer with + no error. The wrapper asserts this. +- `input.dim() >= 1`. A 0-d tensor **segfaults** inside + `AllgatherOp::run_list` and kills every rank in the job with no + diagnostic beyond the fault handler's stack. The wrapper asserts this. +- `len(sizes) == len(group)` when `sizes` is given. The op sizes its + output from `sizes` alone and does not cross-check it against `group`: + a list one entry short silently drops the highest rank's rows (measured + — the result is a correct gather of the remaining ranks, at the shorter + length), and a list one entry too long returns + `sum(sizes)` rows whose tail is uninitialized memory (measured once, in + probing; indexing past the group is undefined behaviour and not + something to rely on). The wrapper asserts this. +- `sizes[my_rank] == input.shape[0]`, in ascending rank order. The op + does not check it and cannot be made to — it takes no rank argument — + and a mismatch hangs, as above. The wrapper cannot check it either, for + the same reason: this one is the caller's. +- `input.device` is the device this rank set with + `torch.cuda.set_device`, and one rank owns one device. +- **A `group`'s first call must not be made inside a CUDA-graph capture.** + It is where the group's NCCL communicator gets built, and that build + raises under capture — see the *CUDA graphs* note for the error text and + the one-call fix. +- **Every rank issues the same sequence of calls on `group`, in the same + order**, graph replays included. Position on the communicator is what + pairs one rank's call with another's, so order is the caller's whole + responsibility — and disagreeing on it is *worse* than disagreeing on the + arguments, because it usually does not hang. Measured at world size 4 in + the certified configuration: + + | how the ranks disagreed | what happened | + |---|---| + | two same-shaped gathers issued in swapped order on one rank | **silently wrong on every rank.** No error, no hang. The swapping rank had 3/4 of its elements wrong in both results, every other rank exactly 1/4 — the swapping rank's block. A swapped *pair* puts the positions back, so a plain gather straight afterwards is bitwise correct. In this entry's test. | + | one rank issuing one gather more than the others, each iteration | **silently wrong on every rank, permanently**: the positions never realign, so every later call stays paired one off. Probing, outside the test, for that reason. | + | a gather swapped against a `torch.ops.trtllm.reducescatter` of the **same** byte count on one rank | **silently wrong on every rank** — the gather came back holding reduce-scatter payload and vice versa. Probing, outside the test. | + | the same swap where the two calls carry **different** byte counts | **wedged.** No rank returned; the job had to be killed. Probing, outside the test, for that reason. | + + So an ordering divergence surfaces as an accuracy loss whenever the + mispaired calls happen to move the same number of bytes, and as a hang + only when they do not. A caller must not wait for a hang to tell it that + its ranks have diverged. +- **The stream is the engine's to choose and the ranks need not agree on + it.** The call runs on `torch.cuda.current_stream()`, and a serving + engine moves that stream under the model — the same forward runs on + torch's graph-capture stream while a decode graph is being captured and + on the serving stream otherwise, so a target cannot pin it. Stream + identity plays no part in pairing the calls: certified with every rank on + a side stream, with one rank on a side stream while the others stayed on + the default, and with ranks alternating in opposite patterns so they + disagreed at every call index — all bitwise correct. Certified with the + side stream joined to the current one on both ends, which is what the + engine and a target's forward both do; two calls on one `group` running + *concurrently* on two streams of the same rank was not exercised. +- `group` entries must be ranks of the MPI session. With + `group = [0, 1, 2, 3, 4]` at world size 4, **rank 0 raised**: + + ``` + RuntimeError: [TensorRT-LLM][ERROR] Assertion failed: Failed: MPI error + ../tensorrt_llm/runtime/utils/mpiUtils.cpp:277 '6' + (../tensorrt_llm/runtime/utils/mpiUtils.cpp:277) + ``` + + and ranks 1, 2 and 3 never returned from the call — the job had to be + killed, twice out of twice. So this is a raise on one rank and a wedge + everywhere else, not a survivable error: measured in probing, outside + the test, for exactly that reason. +- `rows = 0` on **every** rank (an empty `[0, H]` input with + `sizes = None`) is accepted and returns an empty `[0, H]` tensor + (observed in probing, outside the test). A single rank at zero rows is + certified through the `sizes` form. + +## Notes + +The certified path: `mpirun`-launched ranks whose session communicator is +`MPI_COMM_WORLD`, one rank per B200 (sm_100), world size 4, group +`[0, 1, 2, 3]` and `[0, 1]`, bf16 hidden 2560 unless the dtype table says +otherwise. The op has no strategy, workspace or autotuner state, so there +is no second execution path here for a precondition to be quietly true +of — the only axis with two paths is `sizes`, and both are certified. + +**Inside a live serving engine under attention data parallelism.** The +question this answers is whether the runtime is a second party on the +communicator — whether its own cross-rank work can interleave with a +model's calls and pair against them. It is not. Measured by running a +4-rank trtllm serving engine (attention DP on, MoE expert-parallel over the +same four ranks) whose model issues 29 of these gathers and 29 +reduce-scatters per forward, with NCCL's own collective trace on for the +whole process — engine warm-up, 68 decode-graph captures and the served +requests: + +- **The runtime issues nothing on this communicator.** NCCL built exactly + **one** communicator per rank in the whole process, and all 14,036 + collectives logged on it were the model's own — 7,018 `AllGather` and + 7,018 `ReduceScatter` per rank, over 242 forwards. The engine contributed + none. +- **What the engine does instead is host-side MPI on a different + communicator.** Under attention DP it agrees the per-rank token counts + once per step — that is where `attn_metadata.all_rank_num_tokens` comes + from — with `MPIDist.tp_allgather`, an `MPI_Allgather` / `MPI_Allgatherv` + on a sub-communicator it builds with `MPI_Comm_create_group`: congruent + to the session communicator, never identical to it, and never NCCL. + Graph-capture consensus, batch-size consensus, request broadcast and + response gathering all go the same way. So a model's gather cannot be + paired against an engine collective — there is none on this communicator + to pair it against. Certified in this entry's test by driving the engine's + own `MPIDist` object between batches of gathers. +- **All four ranks' call sequences agreed exactly**: the 14,036 trace + entries matched across ranks in op kind, element count and stream index, + position by position. +- **The engine issues a model's collective on two different streams.** The + trace splits into 137 runs: one of 580 calls on the serving stream, then + 68 pairs of [58 calls on torch's graph-capture stream][116 on the serving + stream] — one forward captured, two run eagerly as the runner's warm-up. + Which stream a call lands on is therefore the engine's choice, and the + *Preconditions* entry above records that this is harmless: the ranks + agreed on it here, and they do not have to. + +What those runs do **not** establish. Ten runs of that engine were made; +eight finished and two stopped with every rank spinning in +`cuLaunchKernelEx` (`sched_yield`) — the launch queue full because the +device was not retiring work. One stopped in engine warm-up, all four ranks +in `cudaDeviceSynchronize` after the same warm-up forward; one stopped in a +served forward, all four ranks past the gather and inside the +non-collective work behind it. Nothing measured attributes either to this +call: the ranks' enqueued sequences agreed everywhere they were traced, and +the same collective pattern driven without an engine — the engine's real +host-side step sync plus 29 gathers and 29 reduce-scatters per step at +ragged attention-DP row counts up to 2048, nothing synchronising — ran +2000 steps (58,000 gathers per rank, more collectives than a whole +benchmark run of that engine) clean on those same shared GPUs. Both +stalls landed on the four GPUs that were sharing SMs roughly 50/50 with an +unrelated process; the four exclusive GPUs took four of the ten runs and +stalled on none. That is a correlation on a small sample, not a mechanism. +Note also that stock trtllm does **not** take this path at that +configuration: its MoE communication factory selected `DeepEPLowLatency` +there (measured), so a stock run is not a control for this op. + +**CUDA graphs — the surface a decode step actually needs.** Certified at +world size 4, group `[0, 1, 2, 3]`, bf16, under torch's default +`capture_error_mode="global"`, with every rank capturing and replaying the +same graph: + +| form | under capture | +|---|---| +| `sizes = None` (uniform) | certified at **all 35 of `1..32, 64, 128, 256` rows** | +| `sizes = [...]` (ragged) | certified — but the split is frozen at capture, see below | + +A trtllm engine instantiates one decode graph per configured batch size — +with `cuda_graph_config.max_batch_size = 256` and no explicit +`batch_sizes` list, 35 of them at `1..32, 64, 128, 256` — keeps all of +them alive in one memory pool, and replays them interleaved as the served +batch size moves. Under attention DP a decode call's per-rank row count +*is* that batch size, so the whole set is certified: all 35 captured into +one shared pool, none released, then replayed in four orders — ascending, +descending, shuffled, and largest-jump-first (`256, 1, 128, 2, 64, 3, ...`) +— each order twice, once one replay at a time with a synchronize and a +full check between graphs, and once with all 35 issued back to back and +nothing synchronizing between them. Every replay's payload is one no +earlier call used, and every result is bitwise equal to the concatenation +reference. The same is certified with **29 independent gathers inside each +of the 35 graphs** — 1015 captured collectives in one pool, the shape this +checkpoint's decode graph has (30 layers, layer 0 dense, so 29 +expert-parallel MoE calls each preceded by one gather). + +A replay re-runs the collective over whatever the input buffer holds at +replay time and writes into the same tensor the capture returned — keep it +and read it after each `replay()`. Replays also stay correct with eager +traffic in between: an eager uniform gather at 1500, 2048 or 8192 rows and +an eager ragged gather between every pair of replays, all bitwise correct, +which is the shape a server has when a prefill runs between two decode +steps. + +**`sizes` is frozen into a graph; row counts must be padded to capture +one.** `sizes` is a host-side argument, so a capture bakes in the split it +was given and a replay re-runs *that* split. Certified by capturing two +graphs with different sizes vectors into one pool and replaying them in +both orders across two payload rounds: each keeps its own row split and +its own output length. The consequence for an attention-DP target is +concrete — a graph-captured decode step must pad every rank to the same +row count and pass `sizes = None`, because the ragged split cannot vary +per replay; the ragged form belongs to the eager (non-captured) steps. + +**A `group`'s first-ever call cannot be captured**, because it is where the +NCCL communicator is built. Four ranks capturing their first call each +raised, out of `tensorrt_llm::_v1::getComm`: + +``` +RuntimeError: Failed, NCCL error ../tensorrt_llm/common/opUtils.cpp:175 +'unhandled cuda error (run with NCCL_DEBUG=INFO for details)' +``` + +which invalidates the capture, so the `with torch.cuda.graph(...)` block +itself then raises + +``` +AcceleratorError: CUDA error: operation failed due to a previous error +during capture +``` + +(`cudaErrorStreamCaptureInvalidated`). It is survivable, and the fix is one +line: make one eager call per `group` before any capture. After the failure +the same group gathers correctly eagerly, and a capture taken after that +replays correctly — both certified. The same failure was seen in probing +for a subgroup whose communicator did not yet exist in a process that +already had the full group's. One caveat for anyone catching it: torch's +graph context manager ends the capture before it restores the stream, so a +failed capture leaves its own stream current — put the resting one back +with `torch.cuda.set_stream`. + +**Numerics.** There are none. The op copies bytes, so every assertion in +this entry's test is bitwise (`rtol=0, atol=0`) rather than +tolerance-based, in every dtype, in both forms, eager and captured. A +caller does not have to reason about accumulation order the way the +sibling all-reduce forces it to. + +**Gathering along another axis is the caller's problem.** This op only +concatenates dim 0. trtllm's module-level helper reaches other axes by +reshaping around this call — `view` to 2-D before, `chunk`/`split` and +`cat` after — which is Python-level composition and would have to be +written from catalog entries, not hidden inside a wrapper. For the +attention-DP use (gather token rows before an expert-parallel MoE call) +dim 0 is already the token axis and no reshape is needed. Note also that +"dim 0" is the buffer's first dimension, not necessarily the token axis: +a scale-factor buffer in a swizzled layout has no token rows to gather. + +**Cost: the two forms are indistinguishable.** Measured on the certified +path (world size 4, bf16, hidden 2560, 200 timed iterations after 20 warm-up +calls, CUDA events, eager): `sizes = None` against an explicit uniform +`sizes` vector of the same total size came out at 12.5 vs 12.4 us at 1 row, +13.0 vs 13.2 at 8, 20.2 vs 20.1 at 128 and 80.1 vs 80.2 at 2048 — ratio +1.00-1.01 throughout. So a caller that already has the per-rank counts loses +nothing by passing them, and the reason a captured decode step pads to +uniform rows is the frozen split above, not the cost of the ragged form. +Below ~128 rows the call is launch-bound (12-13 us for a message 20x apart +in size), which is the regime a decode step runs in. + +**The inverse op exists.** `torch.ops.trtllm.reducescatter` has the same +`(input, sizes, group)` signature and undoes this one — it is a different +op and is not this entry. diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/allgather.py b/tensorrt_llm/_torch/staircase/catalog/comm/allgather.py new file mode 100644 index 000000000000..aa3ef337825e --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/comm/allgather.py @@ -0,0 +1,39 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Concatenate every rank's rows of a tensor into the group-wide full set.""" + +from typing import List, Optional + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + + +def allgather( + input: torch.Tensor, + sizes: Optional[List[int]], + group: List[int], +) -> torch.Tensor: + """Gather `input` from every rank in `group`, concatenated along dim 0. + + `sizes` gives each rank's `input.shape[0]` in ascending rank order, or is + `None` when every rank holds the same number of rows. The result is a + freshly allocated tensor on every rank; `input` is left untouched. + """ + # Pure-metadata guards, each for a violation measured on this machine to + # end in silent corruption or a process kill rather than an error. + assert input.dim() >= 1, ( + "input must have at least one dimension; a 0-d tensor segfaults inside " + "AllgatherOp::run_list" + ) + assert input.is_contiguous(), ( + "input must be contiguous; the op reads it as packed memory and a " + "strided view is silently gathered from the wrong elements" + ) + assert sizes is None or len(sizes) == len(group), ( + f"len(sizes)={len(sizes) if sizes is not None else None} must equal " + f"len(group)={len(group)}; the op sizes its output from `sizes` alone, " + "so a short list silently drops the trailing ranks and a long one " + "appends uninitialized rows" + ) + return torch.ops.trtllm.allgather(input, sizes, group) diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/allgather_test.py b/tensorrt_llm/_torch/staircase/catalog/comm/allgather_test.py new file mode 100644 index 000000000000..ad375cbd2702 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/comm/allgather_test.py @@ -0,0 +1,990 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the allgather catalog entry. + +A collective cannot be exercised in one process, so this script is its own +launcher: run it plainly and it re-executes itself under `mpirun` with one +rank per device named in CUDA_VISIBLE_DEVICES, under a deadline the parent +enforces by killing the whole process group. A wedged collective hangs +rather than raising — and this op wedges rather than raising for every +rank-argument-disagreement case probed — so the deadline is what keeps a +broken kernel from taking the calling run down with it. + +Beyond the lockstep surface (uniform / ragged / dtypes / CUDA graphs) the +file covers what a serving engine adds: the engine's own attention-DP +synchronisation interleaved with the op, the engine's stream switching, and +what disagreeing on call order actually does. + + CUDA_VISIBLE_DEVICES=0,1,2,3 uv run python catalog/comm/allgather_test.py +""" + +import os +import random +import signal +import subprocess +import sys +from typing import Any, Dict, List, Optional, Sequence, Tuple + +import torch + +assert torch.cuda.is_available(), "allgather requires CUDA devices" + +_WORKER_FLAG = "--rank-worker" +DEADLINE_S = 1200 + +HIDDEN = 2560 # DeepSeek-V3-Lite hidden size, the caller that motivates this + +# The row counts one engine's decode graphs are captured at. trtllm's graph +# runner instantiates one graph per configured batch size, and under attention +# data parallelism a decode call's per-rank row count *is* that batch size: +# with `cuda_graph_config.max_batch_size = 256` and the default list that is 35 +# graphs at [1..32, 64, 128, 256], all alive at once in one memory pool and +# replayed interleaved as the served batch size moves. So the surface a +# captured gather has to be correct over is the whole set plus the transitions +# between its members, not one representative row count. +GRAPH_BATCH_SIZES: Tuple[int, ...] = tuple(range(1, 33)) + (64, 128, 256) + +# Gathers inside one decode graph for this checkpoint: 30 layers with +# `first_k_dense_replace = 1`, so 29 expert-parallel MoE calls, each preceded +# by one gather of the rank's own rows into the full token set. +SITES_PER_DECODE_GRAPH = 29 + +# Bound inside the rank body, never in the launcher: importing the entry +# pulls in tensorrt_llm, which calls MPI_Init at import, and an +# MPI-initialized process cannot launch `mpirun` — measured on this host, +# mpirun then exits 1 with no output from any rank. +allgather: Any = None +COMM: Any = None +RANK = 0 +WORLD = 1 +GROUP: List[int] = [0] +# The engine's own cross-rank synchronisation object, built from the real +# Mapping a serving engine builds under attention data parallelism. Bound in +# the rank body for the same reason as the rest. +DIST: Any = None +MPI: Any = None + + +def _payload( + rank: int, + rows: int, + seed: int, + dtype: torch.dtype = torch.bfloat16, + trailing: Tuple[int, ...] = (HIDDEN,), +) -> torch.Tensor: + """One rank's contribution: values distinct per (rank, seed), exact in dtype. + + Every rank can regenerate every other rank's payload — cuRAND is + deterministic for a given (seed, shape, dtype) — which is what lets the + reference below be arithmetic rather than a second collective. The float + value sets are multiples of 1/8 (bounded at 2.0 for e4m3, which only + represents that spacing below it), so a payload survives the dtype cast + exactly and every assertion in this file can be bitwise. + """ + gen = torch.Generator(device="cuda").manual_seed(seed * 977 + rank + 1) + shape = (rows, *trailing) + if dtype is torch.uint8: + return torch.randint(0, 256, shape, generator=gen, device="cuda", dtype=torch.int32).to( + torch.uint8 + ) + if dtype in (torch.int32, torch.int64): + return torch.randint( + -(2**20), 2**20, shape, generator=gen, device="cuda", dtype=torch.int64 + ).to(dtype) + amp = 16 if dtype is torch.float8_e4m3fn else 120 + raw = torch.randint(-amp, amp + 1, shape, generator=gen, device="cuda", dtype=torch.int32) + return (raw.float() / 8.0).to(dtype) + + +def _gather_ref( + sizes: Sequence[int], + seed: int, + dtype: torch.dtype = torch.bfloat16, + trailing: Tuple[int, ...] = (HIDDEN,), + ranks: Optional[Sequence[int]] = None, +) -> torch.Tensor: + """Arithmetic reference: the concatenation the collective is supposed to make. + + Built on this rank alone from the same seeds every rank uses, in ascending + rank order. Never from a second collective. + """ + ranks = list(range(WORLD)) if ranks is None else list(ranks) + return torch.cat( + [_payload(r, sizes[i], seed, dtype, trailing) for i, r in enumerate(ranks)], + dim=0, + ) + + +def _assert_bitwise(out: torch.Tensor, ref: torch.Tensor, where: str) -> None: + """Gate of exactly zero: the op moves bytes, so nothing may differ. + + Tightened from the default dtype-aware tolerances rather than loosened — + a gather computes nothing, so any difference at all is a wrong gather. + """ + if out.dtype is torch.float8_e4m3fn: + # torch.testing cannot compare float8 tensors; for a pure data move the + # byte pattern is the honest gate anyway. + out, ref = out.view(torch.uint8), ref.view(torch.uint8) + torch.testing.assert_close(out, ref, rtol=0, atol=0, msg=lambda built: f"{where}: {built}") + + +def _sizes_vectors() -> List[List[int]]: + """Per-rank row counts an attention-DP step produces, adapted to WORLD.""" + return [ + [1] * WORLD, # uniform, but stated as an explicit sizes vector + [1 + 4 * r for r in range(WORLD)], # steadily uneven decode batches + [7 * (WORLD - r) for r in range(WORLD)], # uneven the other way + [0] + [3 + 2 * r for r in range(WORLD - 1)], # one rank with no rows + [2048] + [1 + 388 * r for r in range(WORLD - 1)], # prefill-sized, lopsided + ] + + +def _adp_step_counts(steps: int = 8) -> List[List[int]]: + """Per-rank token counts a serving engine hands its layers, step by step. + + One seeded host RNG, so the vector is identical on every rank without a + collective — which is also what lets the engine's own `tp_allgather` of it + be checked against something rather than trusted. The mix is what a live + attention-DP engine produces: steps where it has equalized the per-rank + batch (it does that whenever a decode step is graph-eligible), ragged + decode steps, and ragged prefill-sized steps. + """ + rng = random.Random(90210) + out: List[List[int]] = [] + for step in range(steps): + if step % 3 == 0: + out.append([rng.choice([1, 2, 8, 17, 32])] * WORLD) + elif step % 3 == 1: + out.append([rng.randint(1, 32) for _ in range(WORLD)]) + else: + out.append([rng.randint(1, 1024) for _ in range(WORLD)]) + return out + + +def _gather_on_a_side_stream(x: torch.Tensor, side: torch.cuda.Stream) -> torch.Tensor: + """Issue one gather with `side` current, joined on both ends.""" + side.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(side): + out = allgather(x, None, GROUP) + torch.cuda.current_stream().wait_stream(side) + return out + + +def _warm_up_off_capture_stream(body, reps: int = 2) -> None: + """Run `body` on a side stream, then rejoin, so a capture can follow. + + Two things have to be done before a capture and cannot be done inside one: + the group's NCCL communicator has to exist (see + test_cuda_graph_capture_of_a_first_call_raises), and torch wants the work + warmed on a non-default stream. + """ + side = torch.cuda.Stream() + side.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(side): + for _ in range(reps): + body() + torch.cuda.current_stream().wait_stream(side) + torch.cuda.synchronize() + COMM.Barrier() + + +def _capture_per_batch_size( + sites: int = 1, seed: int = 11000 +) -> Tuple[Dict[int, List[torch.Tensor]], Dict[int, Any], Dict[int, List[torch.Tensor]]]: + """One graph per GRAPH_BATCH_SIZES entry, one shared pool, all left alive. + + That is the state a graph runner ends up in: every capture after the first + goes into the first one's pool, and none of them is ever destroyed. + `sites` gathers go into each graph, each reading its own persistent input + buffer, so the payload of every site can be moved independently. + Returns (inputs, graphs, outputs), all keyed by row count. + """ + inputs: Dict[int, List[torch.Tensor]] = {} + graphs: Dict[int, Any] = {} + outputs: Dict[int, List[torch.Tensor]] = {} + pool = None + for rows in GRAPH_BATCH_SIZES: + xs = [_payload(RANK, rows, seed + 13 * s) for s in range(sites)] + + def body(xs: List[torch.Tensor] = xs) -> List[torch.Tensor]: + return [allgather(x, None, GROUP) for x in xs] + + _warm_up_off_capture_stream(body) + graph = torch.cuda.CUDAGraph() + ctx = torch.cuda.graph(graph) if pool is None else torch.cuda.graph(graph, pool=pool) + with ctx: + outputs[rows] = body() + # Without this a capture that returned nothing would make every + # `_verify_replay` below an empty loop, i.e. a vacuous pass. + assert len(outputs[rows]) == sites, ( + f"rows={rows}: captured {len(outputs[rows])} outputs, expected {sites}" + ) + if pool is None: + pool = graph.pool() + inputs[rows], graphs[rows] = xs, graph + COMM.Barrier() + return inputs, graphs, outputs + + +def _replay_orders() -> List[Tuple[str, List[int]]]: + """Four ways a served batch size moves through the captured set.""" + asc = list(GRAPH_BATCH_SIZES) + shuffled = list(asc) + random.Random(1234).shuffle(shuffled) + # The largest jump still available at every step: 256, 1, 128, 2, 64, 3... + extremes: List[int] = [] + lo, hi = 0, len(asc) - 1 + while lo <= hi: + extremes.append(asc[hi]) + if lo != hi: + extremes.append(asc[lo]) + lo, hi = lo + 1, hi - 1 + return [ + ("ascending", asc), + ("descending", list(reversed(asc))), + ("shuffled", shuffled), + ("extremes", extremes), + ] + + +def _fresh_payload(inputs: Dict[int, List[torch.Tensor]], rows: int, seed: int) -> None: + """Overwrite every site's persistent input buffer for this row count.""" + for site, x in enumerate(inputs[rows]): + x.copy_(_payload(RANK, rows, seed + 13 * site)) + + +def _verify_replay( + outputs: Dict[int, List[torch.Tensor]], rows: int, seed: int, where: str +) -> None: + """Every site's captured output tensor holds this payload's gather, bitwise.""" + for site, out in enumerate(outputs[rows]): + _assert_bitwise(out, _gather_ref([rows] * WORLD, seed + 13 * site), f"{where} site={site}") + + +def test_cuda_graph_capture_of_a_first_call_raises() -> None: + """A group's first-ever call cannot be captured; the build inside fails. + + Must run before anything else touches GROUP — the failure is specifically + the NCCL communicator being built inside the capture, and it only happens + once per rank set per process. The failure is survivable, and the two + assertions after it are the workaround a caller needs: one eager call + first, then capture. + """ + rows = 16 + x = _payload(RANK, rows, 100) + ref = _gather_ref([rows] * WORLD, 100) + COMM.Barrier() + + graph = torch.cuda.CUDAGraph() + resting_stream = torch.cuda.current_stream() + raised: Optional[BaseException] = None + try: + with torch.cuda.graph(graph): + allgather(x, None, GROUP) + except RuntimeError as exc: + raised = exc + finally: + # torch's context manager ends the capture before restoring the + # stream, so a capture that fails at capture_end leaves its own + # stream current. Put the resting one back by hand. + torch.cuda.set_stream(resting_stream) + del graph + + assert raised is not None, "capturing a group's first-ever call was accepted" + chain, exc = [], raised + while exc is not None and len(chain) < 10: + chain.append(str(exc)) + exc = exc.__context__ + text = "\n".join(chain) + assert "operation failed due to a previous error during capture" in text, text + assert "NCCL error" in text and "opUtils.cpp" in text, text + + # Survivable: the group works eagerly straight afterwards... + torch.cuda.synchronize() + COMM.Barrier() + _assert_bitwise(allgather(x, None, GROUP), ref, "eager after failed capture") + COMM.Barrier() + + # ...and a capture taken after that warm-up replays correctly. + _warm_up_off_capture_stream(lambda: allgather(x, None, GROUP)) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + captured = allgather(x, None, GROUP) + x.copy_(_payload(RANK, rows, 700)) + graph.replay() + torch.cuda.synchronize() + _assert_bitwise(captured, _gather_ref([rows] * WORLD, 700), "replay after warm-up") + del graph + COMM.Barrier() + + +def test_uniform_gather() -> None: + """sizes=None: every rank holds `rows` rows, the result is WORLD * rows. + + Decode-like (1, 2, 8, 32) through prefill-like (2048) row counts, which is + the whole span an attention-DP rank feeds its expert-parallel MoE call. + """ + for rows in (1, 2, 8, 32, 2048): + seed = 1000 + rows + x = _payload(RANK, rows, seed) + out = allgather(x, None, GROUP) + assert out.shape == (rows * WORLD, HIDDEN), out.shape + assert out.dtype is x.dtype and out.device == x.device + assert out.is_contiguous() + _assert_bitwise(out, _gather_ref([rows] * WORLD, seed), f"uniform rows={rows}") + COMM.Barrier() + + +def test_ragged_gather() -> None: + """sizes=[...]: per-rank row counts differ, the result is sum(sizes). + + Attention data parallelism produces exactly this — each rank owns its own + requests — so the ragged form is the one a non-padded step takes, including + the case where one rank has no rows at all. + """ + for i, sizes in enumerate(_sizes_vectors()): + seed = 3000 + 7 * i + x = _payload(RANK, sizes[RANK], seed) + out = allgather(x, sizes, GROUP) + assert out.shape == (sum(sizes), HIDDEN), (sizes, out.shape) + assert out.dtype is x.dtype and out.is_contiguous() + _assert_bitwise(out, _gather_ref(sizes, seed), f"ragged sizes={sizes}") + COMM.Barrier() + + +def test_output_is_fresh_and_input_is_untouched() -> None: + """The op allocates its result; the caller keeps ownership of `input`.""" + rows, seed = 8, 4100 + x = _payload(RANK, rows, seed) + before = x.clone() + out = allgather(x, None, GROUP) + torch.cuda.synchronize() + assert out.data_ptr() != x.data_ptr(), "output aliased the input buffer" + _assert_bitwise(x, before, "input after the call") + + # Clobbering the input afterwards cannot reach the result. + x.fill_(-7.0) + torch.cuda.synchronize() + _assert_bitwise(out, _gather_ref([rows] * WORLD, seed), "output after input clobber") + COMM.Barrier() + + +def test_trailing_dims_are_preserved() -> None: + """dim 0 is the gather axis; every other dim is carried through untouched.""" + cases: List[Tuple[Tuple[int, ...], Optional[List[int]]]] = [ + ((2, HIDDEN // 2), None), + ((2, HIDDEN // 2), [1 + 4 * r for r in range(WORLD)]), + ((), None), # 1-D: a flat vector gathers to WORLD * rows elements + ((), [1 + 4 * r for r in range(WORLD)]), + ] + for i, (trailing, sizes) in enumerate(cases): + seed = 5000 + 11 * i + rows = 6 if sizes is None else sizes[RANK] + x = _payload(RANK, rows, seed, trailing=trailing) + out = allgather(x, sizes, GROUP) + total = rows * WORLD if sizes is None else sum(sizes) + assert out.shape == (total, *trailing), (trailing, sizes, out.shape) + _assert_bitwise( + out, + _gather_ref([rows] * WORLD if sizes is None else sizes, seed, trailing=trailing), + f"trailing={trailing} sizes={sizes}", + ) + COMM.Barrier() + + +def test_dtypes_move_bitwise() -> None: + """The payload dtypes an attention-DP dispatch moves, in both forms. + + bf16 hidden states, plus what a post-quantization dispatch carries next to + them: packed NVFP4 bytes and their scale factors (uint8), fp8 activations, + int32 expert ids and fp32 routing scales. Nothing is computed, so each is + a byte-for-byte move — the multi-tensor sibling op `allgather_list` exists + for gathering them in one call and is not this entry. + """ + sizes = [1 + 4 * r for r in range(WORLD)] + for i, dtype in ( + (0, torch.float16), + (1, torch.float32), + (2, torch.int32), + (3, torch.uint8), + (4, torch.float8_e4m3fn), + ): + seed = 6000 + 13 * i + x = _payload(RANK, 16, seed, dtype) + out = allgather(x, None, GROUP) + assert out.dtype is dtype and out.shape == (16 * WORLD, HIDDEN) + _assert_bitwise(out, _gather_ref([16] * WORLD, seed, dtype), f"uniform {dtype}") + COMM.Barrier() + + xr = _payload(RANK, sizes[RANK], seed, dtype) + outr = allgather(xr, sizes, GROUP) + assert outr.dtype is dtype and outr.shape == (sum(sizes), HIDDEN) + _assert_bitwise(outr, _gather_ref(sizes, seed, dtype), f"ragged {dtype}") + COMM.Barrier() + + +def test_group_selects_a_rank_subset() -> None: + """`group` names MPI session ranks; ranks outside it must not call.""" + subset = [0, 1] + rows, seed = 4, 7100 + if RANK in subset: + out = allgather(_payload(RANK, rows, seed), None, subset) + assert out.shape == (rows * len(subset), HIDDEN), out.shape + _assert_bitwise(out, _gather_ref([rows] * len(subset), seed, ranks=subset), "subset gather") + COMM.Barrier() + + +def test_group_order_does_not_change_the_output_order() -> None: + """The result is ordered by ascending rank, whatever order `group` lists.""" + rows, seed = 4, 7200 + x = _payload(RANK, rows, seed) + out = allgather(x, None, list(reversed(GROUP))) + _assert_bitwise(out, _gather_ref([rows] * WORLD, seed), "reversed group list") + COMM.Barrier() + + +def test_cuda_graph_at_every_engine_batch_size() -> None: + """Captured at all 35 engine batch sizes, then replayed interleaved. + + Capturing one row count proves nothing about an engine, which holds every + graph in GRAPH_BATCH_SIZES alive at once and replays them in whatever order + the served batch size takes. Each replay here gets a payload no earlier + call used, so a graph that silently did not re-run keeps the previous + answer and fails the bitwise gate. + + Two passes per order: one replay at a time with a synchronize and a full + check before the next graph runs, then the whole order issued back to back + with no synchronization in between — 35 differently sized collectives + queued as one run of work, which is the only pass that can catch a hazard + that needs the next call to already be in flight. + """ + inputs, graphs, outputs = _capture_per_batch_size() + assert len(graphs) == len(GRAPH_BATCH_SIZES) == 35 + seed = 110000 + for name, order in _replay_orders(): + for pos, rows in enumerate(order): + seed += 1000 + _fresh_payload(inputs, rows, seed) + graphs[rows].replay() + torch.cuda.synchronize() + _verify_replay(outputs, rows, seed, f"{name}/checked@{pos} rows={rows}") + COMM.Barrier() + staged: Dict[int, int] = {} + for rows in order: + seed += 1000 + staged[rows] = seed + _fresh_payload(inputs, rows, seed) + for rows in order: + graphs[rows].replay() + torch.cuda.synchronize() + for pos, rows in enumerate(order): + _verify_replay(outputs, rows, staged[rows], f"{name}/burst@{pos} rows={rows}") + COMM.Barrier() + # Discrimination for a gate of exactly 0: what a replay that did not happen + # leaves behind is the previous payload's gather. Measured by mutating this + # file to skip one row count's refresh — the checked pass then failed on all + # four ranks at that row count, 163158 of 163840 elements wrong (99.6%), + # greatest absolute difference 30.0. The assertion below keeps that + # separation from silently shrinking. + prev = _gather_ref([16] * WORLD, seed - 1000) + last = _gather_ref([16] * WORLD, seed) + margin = (last.float() - prev.float()).abs().max().item() + assert margin > 8.0, f"stale-replay margin {margin}" + del graphs, outputs, inputs + COMM.Barrier() + + +def test_cuda_graph_one_site_per_moe_layer_in_every_batch_size_graph() -> None: + """The decode graph this checkpoint captures, at all 35 batch sizes. + + 1015 captured gathers, one memory pool. The sites are independent — each + reads its own buffer and every one gets a distinct payload, so a site that + returned a neighbour's data would fail the bitwise gate. Two replay orders + rather than four: the per-order cost is 29x the single-site test's and the + axis this one adds is site count, not order. + """ + inputs, graphs, outputs = _capture_per_batch_size(sites=SITES_PER_DECODE_GRAPH, seed=21000) + seed = 310000 + for name, order in _replay_orders()[2:]: + for pos, rows in enumerate(order): + seed += 1000 + _fresh_payload(inputs, rows, seed) + graphs[rows].replay() + torch.cuda.synchronize() + _verify_replay(outputs, rows, seed, f"{name}/checked@{pos} rows={rows}") + COMM.Barrier() + staged: Dict[int, int] = {} + for rows in order: + seed += 1000 + staged[rows] = seed + _fresh_payload(inputs, rows, seed) + for rows in order: + graphs[rows].replay() + torch.cuda.synchronize() + for pos, rows in enumerate(order): + _verify_replay(outputs, rows, staged[rows], f"{name}/burst@{pos} rows={rows}") + COMM.Barrier() + del graphs, outputs, inputs + COMM.Barrier() + + +def test_cuda_graph_replays_survive_eager_calls_of_other_shapes() -> None: + """Between two decode replays, a server runs work no graph holds. + + A prefill runs eagerly at a row count far outside the captured set, and a + ragged gather — the other transport this op has — runs eagerly on the same + communicator the graphs baked in. + """ + inputs, graphs, outputs = _capture_per_batch_size(seed=41000) + seed = 510000 + for i, rows in enumerate(reversed(GRAPH_BATCH_SIZES)): + wide = (2048, 8192, 1500)[i % 3] + _assert_bitwise( + allgather(_payload(RANK, wide, 610000 + i), None, GROUP), + _gather_ref([wide] * WORLD, 610000 + i), + f"eager uniform {wide} rows", + ) + ragged = [(11 * i + 3 * r) % 97 for r in range(WORLD)] + _assert_bitwise( + allgather(_payload(RANK, ragged[RANK], 710000 + i), ragged, GROUP), + _gather_ref(ragged, 710000 + i), + f"eager ragged {ragged}", + ) + seed += 1000 + _fresh_payload(inputs, rows, seed) + graphs[rows].replay() + torch.cuda.synchronize() + _verify_replay(outputs, rows, seed, f"after-eager@{i} rows={rows}") + COMM.Barrier() + del graphs, outputs, inputs + COMM.Barrier() + + +def test_cuda_graph_holds_the_sizes_vector_it_captured() -> None: + """`sizes` is a host argument: a replay re-runs the split it was captured with. + + Two graphs with different sizes vectors are captured into one pool and + replayed in both orders; each keeps its own row split and its own output + length. This is why a graph-captured decode step needs padded (uniform) + row counts — the ragged split cannot be varied per replay. + """ + first = [1 + 4 * r for r in range(WORLD)] + second = [7 * (WORLD - r) for r in range(WORLD)] + seed = 81000 + buffers, graphs, captured = {}, {}, {} + pool = None + for tag, sizes in (("first", first), ("second", second)): + x = _payload(RANK, sizes[RANK], seed) + _warm_up_off_capture_stream(lambda x=x, sizes=sizes: allgather(x, sizes, GROUP)) + graph = torch.cuda.CUDAGraph() + ctx = torch.cuda.graph(graph) if pool is None else torch.cuda.graph(graph, pool=pool) + with ctx: + captured[tag] = allgather(x, sizes, GROUP) + pool = graph.pool() if pool is None else pool + buffers[tag], graphs[tag] = x, graph + assert captured[tag].shape == (sum(sizes), HIDDEN), captured[tag].shape + COMM.Barrier() + + for round_seed in (82000, 83000): + for tag, sizes in (("second", second), ("first", first)): + buffers[tag].copy_(_payload(RANK, sizes[RANK], round_seed)) + graphs[tag].replay() + torch.cuda.synchronize() + _assert_bitwise( + captured[tag], + _gather_ref(sizes, round_seed), + f"ragged replay {tag} sizes={sizes}", + ) + COMM.Barrier() + del graphs, captured, buffers + COMM.Barrier() + + +def test_the_gate_discriminates_a_wrong_gather() -> None: + """The bitwise gate rejects every plausible wrong gather, by a wide margin. + + A gate of exactly 0 cannot be too loose, but it can be blind: if every + rank's payload looked alike, a misordered or short gather would pass it. + These are the four ways this op could plausibly be wrong — rank blocks in + the wrong order, one rank's block standing in for another's, a rank that + never contributed, and the right data landing at the wrong offsets under a + rotated ragged split. Every one of them has the right shape, so shape + checking alone would let all four through. + + Measured on the certified path (world size 4, bf16, hidden 2560), as + (fraction of elements differing, greatest absolute difference): reversed + (0.9958, 30.0), one block substituted (0.2489, 30.0), one rank missing + (0.2489, 15.0), rotated split (0.9602, 30.0) — against a gate of 0. The + two 0.2489 figures sit just under the arithmetic ceiling of 1/WORLD: a + variant that corrupts one rank's block cannot move more than a quarter of + a 4-rank gather, which is why the floors below are per variant. + """ + rows, seed = 8, 91000 + uniform = allgather(_payload(RANK, rows, seed), None, GROUP) + blocks = [_payload(r, rows, seed) for r in range(WORLD)] + _assert_bitwise(uniform, torch.cat(blocks, dim=0), "uniform baseline") + + sizes = [1 + 4 * r for r in range(WORLD)] + ragged = allgather(_payload(RANK, sizes[RANK], seed), sizes, GROUP) + _assert_bitwise(ragged, _gather_ref(sizes, seed), "ragged baseline") + + rotated_split = sizes[1:] + sizes[:1] + one_block = 0.9 / WORLD # a variant that corrupts a single rank's block + wrong = [ + ("rank blocks reversed", uniform, torch.cat(blocks[::-1], dim=0), 0.9), + ( + "rank 0's block twice", + uniform, + torch.cat([blocks[0]] + blocks[1:-1] + [blocks[0]], dim=0), + one_block, + ), + ( + "one rank never contributed", + uniform, + torch.cat(blocks[:-1] + [torch.zeros_like(blocks[-1])], dim=0), + one_block, + ), + ( + "ragged split rotated", + ragged, + torch.cat([_payload(r, rotated_split[r], seed) for r in range(WORLD)], 0), + 0.9, + ), + ] + for name, got, variant, floor in wrong: + assert variant.shape == got.shape, (name, variant.shape, got.shape) + diff = (variant.float() - got.float()).abs() + fraction = (diff != 0).float().mean().item() + assert fraction > floor, f"{name}: only {fraction:.4f} of elements differ" + assert diff.max().item() > 8.0, f"{name}: max difference {diff.max().item()}" + COMM.Barrier() + + +def test_the_engines_own_cross_rank_step_is_not_on_this_communicator() -> None: + """What a serving engine synchronises with, next to what this op uses. + + Under attention data parallelism the engine agrees the per-rank token + counts once per step, and that is the only cross-rank traffic it issues in + this configuration. It does it with `MPIDist.tp_allgather` — a **host-side + MPI** collective on a sub-communicator the engine builds itself — not with + a device collective, and not on the MPI session communicator this op + resolves `group` against. The two never share an ordering. + """ + assert type(DIST).__name__ == "MPIDist", type(DIST).__name__ + tp_comm = DIST.tp_comm + assert isinstance(tp_comm, MPI.Comm), type(tp_comm) + # Same ranks in the same order, but a communicator of its own: MPI's own + # comparison says congruent, never identical. + assert MPI.Comm.Compare(tp_comm, MPI.COMM_WORLD) != MPI.IDENT + assert tp_comm.Get_size() == WORLD and tp_comm.Get_rank() == RANK + counts = [3 + 2 * r for r in range(WORLD)] + assert list(DIST.tp_allgather(counts[RANK])) == counts + COMM.Barrier() + + +def test_interleaved_with_the_engines_attention_dp_synchronisation() -> None: + """The op inside a forward, between the engine's own cross-rank steps. + + The shape a served forward has: once per step the engine agrees the + per-rank token counts on the host (the previous test's collective, which is + where `attn_metadata.all_rank_num_tokens` comes from), then the model + issues one gather per expert-parallel layer on rows padded to + `max(all_rank_num_tokens)`. + + Nothing synchronises here. Every step's payloads and gathers, and the host + collective between steps, are issued back to back for all 8 steps before a + single result is read — 232 gathers per rank in flight, which is what makes + the launch queue saturate the way it does in a served forward rather than + in a lockstep test loop. All of them are checked afterwards, bitwise. + """ + steps = _adp_step_counts() + pending: List[Tuple[int, int, int, int, torch.Tensor]] = [] + for step, counts in enumerate(steps): + # The engine's step-level sync, on the object the engine builds. It is + # a host collective, so the device work queued above it stays in + # flight across it — that interleaving is the point of this test. + observed = list(DIST.tp_allgather(counts[RANK])) + assert observed == counts, (step, observed, counts) + rows = max(counts) + for layer in range(SITES_PER_DECODE_GRAPH): + seed = 130000 + 1000 * step + 13 * layer + out = allgather(_payload(RANK, rows, seed), None, GROUP) + pending.append((step, layer, rows, seed, out)) + assert len(pending) == len(steps) * SITES_PER_DECODE_GRAPH + for step, layer, rows, seed, out in pending: + _assert_bitwise( + out, + _gather_ref([rows] * WORLD, seed), + f"adp step={step} layer={layer} rows={rows}", + ) + COMM.Barrier() + + +def test_the_stream_the_call_lands_on_is_not_part_of_the_match() -> None: + """The op runs on whatever stream is current, and ranks need not agree. + + A serving engine moves the current stream under the model: the same + forward runs on torch's graph-capture stream while a decode graph is being + captured and on the serving stream otherwise, and a target cannot pin + either. So "must the ranks agree on the stream?" is a precondition + question. They do not have to — what pairs the calls is their order on the + communicator. Three shapes: every rank on a side stream, one rank on a side + stream while the others stay on the default, and ranks alternating in + opposite patterns so that at every call index they disagree. + """ + side = torch.cuda.Stream() + rows = 24 + cases = [ + ("every rank on a side stream", lambda i: True), + ("only rank 0 on a side stream", lambda i: RANK == 0), + ("ranks alternating opposite", lambda i: (i + RANK) % 2 == 0), + ] + for c, (name, use_side) in enumerate(cases): + for i in range(4): + seed = 150000 + 1000 * c + 13 * i + x = _payload(RANK, rows, seed) + out = _gather_on_a_side_stream(x, side) if use_side(i) else allgather(x, None, GROUP) + _assert_bitwise(out, _gather_ref([rows] * WORLD, seed), f"{name} i={i}") + COMM.Barrier() + + +def test_wrapper_guards_a_non_contiguous_input() -> None: + """The wrapper's assert stands where the op itself is silently wrong.""" + rows, seed = 8, 92000 + values = _payload(RANK, rows, seed) + padded = torch.zeros(rows, HIDDEN * 2, dtype=torch.bfloat16, device="cuda") + padded[:, :HIDDEN] = values + view = padded[:, :HIDDEN] + assert torch.equal(view, values) and not view.is_contiguous() + + raw = torch.ops.trtllm.allgather(view, None, GROUP) + torch.cuda.synchronize() + # It read `rows * HIDDEN` packed elements from the start of `padded`, which + # is this rank's first rows/2 values interleaved with the zero half — so + # every second row of every rank's block comes back zero. + assert raw.shape == (rows * WORLD, HIDDEN), raw.shape + assert raw[1::2].abs().max().item() == 0.0, "expected the zero half to show" + assert (raw != _gather_ref([rows] * WORLD, seed)).float().mean().item() > 0.4 + COMM.Barrier() + + try: + allgather(view, None, GROUP) + except AssertionError: + pass + else: + raise AssertionError("wrapper accepted a non-contiguous input") + COMM.Barrier() + + +def test_wrapper_guards_a_sizes_list_of_the_wrong_length() -> None: + """A short `sizes` list silently drops the trailing ranks.""" + rows, seed = 4, 93000 + short = [rows] * (WORLD - 1) + raw = torch.ops.trtllm.allgather(_payload(RANK, rows, seed), short, GROUP) + torch.cuda.synchronize() + assert raw.shape == (rows * (WORLD - 1), HIDDEN), raw.shape + _assert_bitwise(raw, _gather_ref(short, seed, ranks=range(WORLD - 1)), "short sizes list") + COMM.Barrier() + + for sizes in (short, [rows] * (WORLD + 1)): + try: + allgather(_payload(RANK, rows, seed), sizes, GROUP) + except AssertionError: + pass + else: + raise AssertionError(f"wrapper accepted sizes of length {len(sizes)}") + COMM.Barrier() + + +def test_wrapper_guards_a_zero_dim_input() -> None: + """A 0-d input segfaults inside the op, so the wrapper stops it first. + + The op is deliberately not called here: the crash is in + `AllgatherOp::run_list` and kills every rank in the job. + """ + try: + allgather(torch.tensor(1.0, device="cuda"), None, GROUP) + except AssertionError: + pass + else: + raise AssertionError("wrapper accepted a 0-d input") + COMM.Barrier() + + +def test_call_order_disagreement_corrupts_silently() -> None: + """Ranks that disagree on the *order* of two equal-sized gathers get wrong + data back, with no error and no hang. + + Argument disagreement wedges (the contract's first precondition). Order + disagreement does not, as long as the two calls carry the same number of + bytes: what pairs the calls is their position on the communicator, so a + swapped pair pairs each rank's first call with the others' second, every + call returns a correctly *shaped* tensor and nothing reports anything. That + is the failure mode a target has to design against — it surfaces as an + accuracy loss, not as a crash. + + Measured on the certified path (world size 4, bf16, hidden 2560): the rank + that swapped gets exactly 3/4 of its elements wrong in both results, every + other rank exactly 1/4 — the disagreeing rank's block — against floors of + 0.7 and 0.2 here. Nothing raised and nothing hung, in 20 probe iterations + out of 20 and in every run of this test. + + Kept last, and it puts the communicator back: a swapped *pair* leaves the + per-communicator positions aligned once both calls are made, so the final + assertion is a plain gather coming back bitwise correct. + """ + rows, seed = 12, 160000 + first, second = _payload(RANK, rows, seed), _payload(RANK, rows, seed + 500) + ref_first = _gather_ref([rows] * WORLD, seed) + ref_second = _gather_ref([rows] * WORLD, seed + 500) + COMM.Barrier() + + if RANK == 0: + got_second = allgather(second, None, GROUP) + got_first = allgather(first, None, GROUP) + else: + got_first = allgather(first, None, GROUP) + got_second = allgather(second, None, GROUP) + torch.cuda.synchronize() + + floor = 0.7 if RANK == 0 else 0.2 + for name, got, ref in ( + ("first", got_first, ref_first), + ("second", got_second, ref_second), + ): + assert got.shape == ref.shape, (name, got.shape, ref.shape) + wrong = (got != ref).float().mean().item() + assert wrong > floor, f"{name}: only {wrong:.4f} of elements differ" + COMM.Barrier() + + x = _payload(RANK, rows, seed + 900) + _assert_bitwise( + allgather(x, None, GROUP), + _gather_ref([rows] * WORLD, seed + 900), + "gather after an order disagreement", + ) + COMM.Barrier() + + +TESTS = ( + # Stays first: it is the only test that can observe GROUP's first-ever + # call, and every later test needs the communicator it builds. + test_cuda_graph_capture_of_a_first_call_raises, + test_uniform_gather, + test_ragged_gather, + test_output_is_fresh_and_input_is_untouched, + test_trailing_dims_are_preserved, + test_dtypes_move_bitwise, + test_group_selects_a_rank_subset, + test_group_order_does_not_change_the_output_order, + test_cuda_graph_at_every_engine_batch_size, + test_cuda_graph_one_site_per_moe_layer_in_every_batch_size_graph, + test_cuda_graph_replays_survive_eager_calls_of_other_shapes, + test_cuda_graph_holds_the_sizes_vector_it_captured, + test_the_gate_discriminates_a_wrong_gather, + test_the_engines_own_cross_rank_step_is_not_on_this_communicator, + test_interleaved_with_the_engines_attention_dp_synchronisation, + test_the_stream_the_call_lands_on_is_not_part_of_the_match, + test_wrapper_guards_a_non_contiguous_input, + test_wrapper_guards_a_sizes_list_of_the_wrong_length, + test_wrapper_guards_a_zero_dim_input, + # Stays last: it deliberately disagrees on call order, and although a + # swapped pair realigns the communicator (its final assertion proves it), + # nothing after it should depend on that. + test_call_order_disagreement_corrupts_silently, +) + + +def _run_one_rank() -> int: + """Body of one MPI rank: run every test, abort the job if any fails.""" + global COMM, RANK, WORLD, GROUP, allgather, DIST, MPI + from mpi4py import MPI as _MPI + + from tensorrt_llm._torch.distributed import Distributed + from tensorrt_llm.mapping import Mapping + + from . import allgather as entry + + allgather = entry.allgather + MPI = _MPI + COMM = MPI.COMM_WORLD + RANK = COMM.Get_rank() + WORLD = COMM.Get_size() + GROUP = list(range(WORLD)) + assert WORLD >= 2, f"a collective needs at least 2 ranks, got {WORLD}" + torch.cuda.set_device(RANK) + # The engine's own cross-rank object, from the Mapping a serving engine + # builds for this topology: attention DP over `WORLD` ranks with the MoE + # expert-parallel over the same set. Real state, not a stand-in — it is the + # class the executor calls every step. + # `Distributed.get` is declared to return the abstract base; the concrete + # class under mpirun is MPIDist and only it carries `tp_comm`, which the + # tests below assert on — hence the widening. + engine_dist: Any = Distributed.get( + Mapping( + world_size=WORLD, + rank=RANK, + tp_size=WORLD, + moe_ep_size=WORLD, + enable_attention_dp=True, + ) + ) + DIST = engine_dist + + for test in TESTS: + try: + test() + except BaseException: + import traceback + + print(f"[rank {RANK}] FAILED {test.__name__}", flush=True) + traceback.print_exc() + sys.stdout.flush() + sys.stderr.flush() + # Abort rather than return: a rank that leaves a collective early + # wedges every other rank in it. + COMM.Abort(1) + COMM.Barrier() + print(f"[rank {RANK}] {len(TESTS)} tests passed", flush=True) + return 0 + + +def _spawn_ranks() -> None: + """Re-exec this file under mpirun, one rank per claimed device.""" + visible = os.environ.get("CUDA_VISIBLE_DEVICES") + assert visible, ( + "set CUDA_VISIBLE_DEVICES to the devices this run owns, " + "e.g. export CUDA_VISIBLE_DEVICES=0,1,2,3" + ) + world_size = len([d for d in visible.split(",") if d.strip()]) + assert world_size >= 2, ( + f"CUDA_VISIBLE_DEVICES names {world_size} device(s); a collective test needs at least 2" + ) + command = [ + "mpirun", + "-n", + str(world_size), + sys.executable, + "-m", + "tensorrt_llm._torch.staircase.catalog.comm.allgather_test", + _WORKER_FLAG, + ] + print(f"[launcher] {' '.join(command)}", flush=True) + # Own process group so the deadline can kill wedged grandchildren too. + process = subprocess.Popen(command, start_new_session=True) + try: + code = process.wait(timeout=DEADLINE_S) + except subprocess.TimeoutExpired: + os.killpg(process.pid, signal.SIGKILL) + process.wait() + raise AssertionError( + f"the {world_size}-rank run did not finish in {DEADLINE_S}s (wedged)" + ) from None + assert code == 0, f"the {world_size}-rank run exited {code}" + + +if __name__ == "__main__": + if _WORKER_FLAG in sys.argv: + sys.exit(_run_one_rank()) + _spawn_ranks() + print("OK") diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/conftest.py b/tensorrt_llm/_torch/staircase/catalog/comm/conftest.py new file mode 100644 index 000000000000..88bd8941aa50 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/comm/conftest.py @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Keep pytest from collecting the collective entries' rank bodies. + +``allgather_test.py`` and ``reducescatter_test.py`` are their own launchers: +their ``test_*`` functions read module-global rank state that only +``_run_one_rank`` binds, and several of them assert on communicator state the +previous test left behind, so they are a fixed sequence inside one 4-rank job +rather than independent cases. Collected directly they would run at world size +1 against unbound globals. + +``tests/unittest/_torch/staircase/comm/test_staircase_*_op_matrix.py`` are the +collected entry points; each starts the 4-rank job and reports its result. +These two files stay here rather than moving with them because the launcher +re-execs them as ``python -m`` and the ranks need this package context for +their relative imports. +""" + +collect_ignore = ["allgather_test.py", "reducescatter_test.py"] diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/reducescatter.md b/tensorrt_llm/_torch/staircase/catalog/comm/reducescatter.md new file mode 100644 index 000000000000..3f5ca642d0b5 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/comm/reducescatter.md @@ -0,0 +1,424 @@ +--- +receipts: + sm_100: {status: passed, trtllm: 1.3.0rc21, world_size: 4} + sm_103: {status: passed, trtllm: 1.3.0rc26, world_size: 4} +--- + +# reducescatter + +**Wraps** `torch.ops.trtllm.reducescatter` (one call). + +## Semantics + +Every rank in `group` calls with a tensor of the same full shape. The op sums +them elementwise, splits the sum along **dim 0**, and gives each rank the +slice belonging to **its position in `group`, in ascending rank order**: + +``` +total = input_r0 + input_r1 + ... + input_r(G-1) # elementwise +out_i = total[offset_i : offset_i + n_i] # rank at group position i +``` + +It is **not** in place — the result is a freshly allocated tensor, and +`input` is left untouched, so clobbering `input` right after the call cannot +reach the result. + +Only dim 0 is split. Every other dimension is carried through untouched: +`[G*n, 2560]` scatters to `[n, 2560]`, `[G*n, 2, 1280]` to `[n, 2, 1280]`, +and a 1-D `[G*n]` to `[n]`. + +The `sizes` argument picks between two splits: + +**`sizes = None`** — the split is even. `input.shape[0]` must be a multiple of +`len(group)`; rank at position `i` gets rows `[i*n, (i+1)*n)` where +`n = input.shape[0] / len(group)`. + +**`sizes = [n_0, ..., n_(G-1)]`** — rank at position `i` gets `n_i` rows +starting at `sum(sizes[:i])`. The list is in ascending rank order, must be +identical on every rank, and `sum(sizes)` must equal `input.shape[0]`. +Entries may be `0`: a rank that keeps no rows still has to call, and the +reduction is still correct. This is the form attention data parallelism needs, +where per-rank token counts differ by construction. + +Both forms are certified. They are **not the same code path**, and NCCL's own +trace shows why: `sizes = None` issues **one `ncclReduceScatter`** whose count +is the elements one rank keeps (rows kept times the trailing dims), while a +`sizes` vector issues **one grouped `ncclReduce` per rank**, rooted at that +rank (for `sizes = [1, 5, 9, 13]` at hidden 2560: counts 2560 / 12800 / 23040 / +33280 rooted at 0 / 1 / 2 / 3, in that one call). An even explicit `sizes` +vector returns bits identical to `sizes = None` and costs the same, while a +genuinely uneven split costs about 1.5x more at 128 rows and above (see +*Notes*). Either way the whole call is bracketed by one +`ncclGroupStart`/`ncclGroupEnd` pair, so one call is **one** position on the +communicator whichever form it takes, and the ragged form's `G` reduces pair +across ranks atomically rather than one root at a time — measured, and matching +the disassembly of `ReducescatterOp::run_list` in `libth_common.so`. That is +what makes the call-order rule in *Preconditions* form-independent. + +**It is the inverse of the sibling all-gather, in layout.** The rows this op +hands back to rank `i` are exactly the rows rank `i` would have contributed to +a gather with the same `sizes` and `group` — certified end to end, including +`reducescatter(allgather(x)) == G * x`. What it is *not* is a byte-exact +inverse: the gather moves bytes, this one computes a sum, and the sum is taken +in the input dtype rather than in fp32. Read *Numerics* before relying on the +result of a chain. + +**Fusion boundary.** Inside the call: the reduction, the split, and the output +allocation, nothing else. Outside: everything that produced `input` (under +attention DP, this rank's expert-window output over the whole gathered token +set), any residual add, any scaling by routing weights, any quantization, any +padding of row counts to a uniform value, and any reshaping — including the +reshape a caller needs to split along an axis other than dim 0, which this op +cannot do. + +**Group membership.** `group` names **ranks in trtllm's MPI session +communicator** (`MPI_COMM_WORLD` under `mpirun`), not device ordinals or +`torch.distributed` ranks. A subset is legal: with `group = [2, 3]` at world +size 4, ranks 2 and 3 reduce only between themselves — rank 2 gets slice 0 and +rank 3 gets slice 1, i.e. the **position in the group**, not the MPI rank +(certified; a subset excluding rank 0 is what pins this). Ranks 0 and 1 must +not call. The list is treated as a **set** — passing `[3, 2, 1, 0]` produces +the identical result to `[0, 1, 2, 3]` — so `group` cannot be used to permute +the split. + +## Signature + +```python +def reducescatter( + input: torch.Tensor, + sizes: Optional[List[int]], + group: List[int], +) -> torch.Tensor +``` + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `input` | `[rows, ...]`, at least 1-D; `[T, 2560]`, `[T, 2, 1280]` and `[T]` tested | bf16; also fp16, fp32, int32, int64, uint8, int8 | **contiguous** | CUDA | +| `sizes` | `None`, or a list of `len(group)` non-negative ints in ascending rank order summing to `input.shape[0]` | Python `int` | — | host | +| `group` | list of MPI session ranks | Python `int` | — | host | +| returns | `[input.shape[0] // len(group), *input.shape[1:]]` for `sizes = None`, `[sizes[my position], *input.shape[1:]]` otherwise | = `input.dtype` | contiguous, newly allocated | = `input.device` | + +There are no other arguments — no reduction operator (the reduction is always +a sum), no strategy, no workspace, no output buffer, no autotuner. Unlike the +sibling all-reduce there is no transport to select. + +Every float dtype in the table reduces as an arithmetic sum in **that dtype** +(see *Numerics*). The integer dtypes reduce as integer sums and wrap on +overflow in the narrow ones: four ranks of `int8` 100 came back as `-112` +(measured, outside the test, which keeps its integer payloads inside range). + +Dtypes that are **not** in the table divide into two kinds, and both are +traps: + +- `torch.float8_e4m3fn` is **accepted and silently wrong**. The op sums the + e4m3 *bit patterns* as unsigned bytes, wrapping at 256: four ranks each + holding `1.0` (byte `0x38`) return byte `0xE0`, i.e. `-32.0`, not `4.0`. + The wrapper rejects the dtype for this reason. Note that the sibling + all-gather moves fp8 correctly — it only copies — so a post-quantization + dispatch that works on the way out does **not** work on the way back. + `torch.bool` is the same shape of trap: the underlying bytes are summed + (four ranks of `True` leave byte value 4), which torch then reads as `True`, + so it behaves like an OR. Neither is certified. +- `torch.float64`, `torch.float8_e5m2`, `torch.float8_e4m3fnuz`, + `torch.float8_e5m2fnuz` and `torch.float8_e8m0fnu` **raise** + `RuntimeError: unsupported data type (../tensorrt_llm/runtime/torchUtils.h:113)` + — and that raise is terminal for the whole process. See *Preconditions*. + +## Metadata consumed + +No attention metadata, no KV cache, no registered layer, no workspace. +Two pieces of process state stand behind the call: + +1. **trtllm's MPI session communicator.** The op resolves `group` against it + and finds this rank's position in it. Under `mpirun` the session is + `MPI_COMM_WORLD`; a trtllm engine's worker ranks are already inside one, so + a target's forward needs no setup. World size 1 was not exercised; this + entry's receipt is world size 4. The sibling op + `torch.ops.trtllm.reducescatter_pg` takes a `c10d.ProcessGroup` instead, + for the `TLLM_DISABLE_MPI=1` path; `torch.ops.trtllm.reducescatter_list` + reduce-scatters a list of tensors in one call. Both are different ops and + neither is this entry. +2. **The NCCL communicator cache**, keyed by the rank set and shared with the + other collectives. Built on first use for a given `group`, which makes the + first call for a group much slower than the rest — and means that first + call **cannot be inside a CUDA-graph capture** (see *Preconditions*). + +## Preconditions + +- **Every rank in `group` calls, and their arguments agree**: same dtype, same + trailing shape, same `sizes` list, and the same `input.shape[0]`. + Disagreement is **not survivable**: it hangs rather than raising. Measured at + world size 4, each in its own process — ranks passing different row counts + with `sizes = None` (three ranks returned a plausible-looking tensor, one + never came back), and ranks passing different `sizes` vectors (no rank came + back). Both jobs had to be killed. This is why this entry's test carries an + external deadline. Do not read the hang as a guarantee, though: both of those + disagreements change how many bytes a rank moves, and a disagreement that + leaves the byte counts equal on every rank is silently wrong instead — see + the call-order bullet below, where that is measured. +- **A rank outside `group` must not call.** Measured: with `group = [0, 1]` at + world size 4, ranks 0 and 1 returned correctly and ranks 2 and 3 **never + returned** — they block inside the communicator bootstrap without ever + reaching the op's own membership check. The job had to be killed. +- `input` is **contiguous**. The op reads `input.numel()` elements from + `input.data_ptr()` and ignores strides, so a strided view is reduced from the + wrong bytes and returns a plausible-looking wrong answer with no error. The + wrapper asserts this. +- `input.dim() >= 1`. A 0-d tensor **segfaults** inside + `ReducescatterOp::run_list` and kills every rank in the job with no + diagnostic beyond the fault handler's stack. The wrapper asserts this. +- **The split must cover the input exactly**, and nothing checks it but the + wrapper: + - `sizes is None` requires `input.shape[0] % len(group) == 0`. Otherwise the + op takes `shape[0] // len(group)` rows per rank and the remainder is + **silently dropped** — 9 rows over 4 ranks returns 2 rows per rank, and the + 9th row is never reduced (measured; the 8 rows that are returned are + correct). + - `sizes` given requires `sum(sizes) == input.shape[0]`. A sum smaller than + the input silently ignores the tail rows (measured: the rows the split does + cover come back correct). A sum larger than the input makes the op read + **past the end of the buffer** — observed once, during this entry's test + bring-up, to return a correctly shaped tensor full of `NaN`. + - `len(sizes) == len(group)`. One entry short raises + `IndexError: vector::_M_range_check: __n (which is 3) >= this->size() + (which is 3)` on the highest rank while the others sit in the collective — + a raise on one rank and a wedge everywhere else, measured, not recoverable + inside the job. One entry too long is the out-of-bounds read above. +- `input.dtype` must be one the op supports, and this is a check the caller + has to make **before** calling rather than by catching. The unsupported + dtypes raise `RuntimeError: unsupported data type`, but the raise happens + after the op has opened an NCCL group and before it closes it, so the + group is left unbalanced and **every collective the process issues + afterwards silently returns garbage** — measured for this op (99.6% of + elements wrong on the very next call), and measured for + `torch.ops.trtllm.allgather` and `torch.ops.trtllm.allreduce` on the same + group too. It is the process that is finished, not just this entry's + communicator. Every other failure in this section that returns at all leaves + the group usable — the test checks that explicitly after the float8, + non-contiguous and uncovered-split calls. +- `input.device` is the device this rank set with `torch.cuda.set_device`, and + one rank owns one device. +- **A `group`'s first call must not be made inside a CUDA-graph capture.** It + is where the group's NCCL communicator gets built, and that build raises + under capture — see the *CUDA graphs* note for the error text and the + one-call fix. The cache is shared across collectives, so a group whose + communicator another op already built is warm for this one. +- **Every rank issues the same sequence of calls on `group`, in the same + order**, graph replays included. What pairs one rank's call with another's is + its **position** on the communicator — not the arguments, not the op, not the + stream — so call order is the caller's whole responsibility, and disagreeing + about it is *worse* than disagreeing about the arguments, because it does not + hang. Certified at world size 4 in the configuration of *Notes*: + + | how the ranks disagreed | what happened | + |---|---| + | two same-shaped calls issued in swapped order on one rank | **silently wrong on every rank.** No error, no hang, right shape, finite values. Each position returned the sum of whatever payloads met there, matched **bitwise** against a locally computed mix. Both `sizes` forms behave identically. A swapped *pair* puts the positions back, so a plain call straight afterwards is bitwise correct. In this entry's test. | + | one rank issuing one call more than the others | **silently wrong from the extra call onward, and it does not heal.** Every position after it stays paired one off; there is no resynchronization point inside a stream of collectives. Certified over five consecutive positions, every one of them bitwise equal to its own mix, and alignment came back only once the call counts were equalized. In this entry's test — where the other ranks issue one catch-up call at the end, which is what lets it terminate. What a rank left *permanently* ahead does was not probed: its last call has no partner, so the test would not have returned. | + | this op swapped against `torch.ops.trtllm.allgather` of the **same** element count on one rank | **silently wrong on every rank**, both calls returning at the shapes their own arguments imply. An even reduce-scatter of `[G*n, H]` and an all-gather of `[n, H]` both move `n*H` elements as NCCL counts them. The pair realigns. In this entry's test. | + | the same swap where the two calls carry **different** element counts | **wedged.** All four ranks entered the pair, none came out of it, and the job had to be ended from outside. Certified by the test's own second job, which is the only place it can live: a job that wedges never reports. | + | ranks disagreeing about the **stream** | **correct.** See the next bullet. | + + **How wrong** follows from the fact that this op computes rather than moves. + A collective that only copies spoils just the block the disagreeing rank + contributed; here that rank's payload is an addend of **every** element of + **every** rank's slice, so a single rank out of step makes every rank wrong + nearly everywhere — measured 0.982-0.993 of elements differing from the + intended result, on every rank, for every divergence above that returned. + The residue is elements where the two payloads happened to agree. + + **A mispaired result is deterministic, not noise**, which is the trap. It + carries no `NaN`, no `Inf` and no shape anomaly, it is bitwise identical + across repeats, and it is bitwise the **ring-order sum of the addends that + met** — the accumulation order of *Numerics*, applied to the wrong operands. + Certified on payloads with no exactness property, where a changed + accumulation order would show, and against a reference computed locally from + seeds, which is what makes it reproducible in any process that mispairs the + same way rather than only in this one (the cross-op mispairing was also + checked directly, returning identical bits in two separate processes). So a + divergence cannot be found by re-running the step and looking for + instability, and under attention data parallelism — where every rank pads to + `max(all_rank_num_tokens)` and every call therefore moves the same number of + bytes — it surfaces as a stable accuracy loss and nothing else. **Do not wait + for a hang to tell you the ranks have diverged.** +- **The stream is the engine's to choose and the ranks need not agree on it.** + The call runs on `torch.cuda.current_stream()`, and a serving engine moves + that stream under the model — the same forward runs on torch's graph-capture + stream while a decode graph is being captured and on the serving stream + otherwise, so a target cannot pin it. Stream identity plays no part in + pairing the calls, and none in the reduction order either: certified with + every rank on a side stream, with one rank on a side stream while the others + stayed on the default, and with ranks alternating in opposite patterns so + they disagreed at every call index — every result bitwise correct, including + against the ring-order chain of *Numerics* on payloads where a changed + accumulation order would show. Certified with the side stream joined to the + current one on both ends, which is what the engine and a target's forward + both do; two calls on one `group` running *concurrently* on two streams of + the same rank was not exercised. + +## Notes + +The certified path: `mpirun`-launched ranks whose session communicator is +`MPI_COMM_WORLD`, one rank per B200 (sm_100), world size 4, group +`[0, 1, 2, 3]` (subsets `[0, 1]` and `[2, 3]` for the membership claims), +bf16 hidden 2560 unless the dtype table says otherwise, NCCL as the only +transport. There is no strategy, workspace or autotuner state here, so the +only axis with two execution paths is `sizes`, and both are certified. + +**Who else is on the communicator, inside a serving engine.** Nobody — so the +only calls whose order has to line up are the ones the caller makes. Measured +on this machine by tracing a live 4-rank trtllm serving engine under attention +data parallelism with `NCCL_DEBUG=INFO NCCL_DEBUG_SUBSYS=INIT,COLL`: one NCCL +communicator per rank, 14,036 collectives on it over 242 forwards (7,018 +all-gathers and 7,018 reduce-scatters, all issued by the model, none by the +runtime), and per-rank sequences bit-identical on (op kind, element count, +stream index) position by position. The runtime's own cross-rank step under +attention DP — agreeing `all_rank_num_tokens` once per forward — is a +**host-side MPI** collective on a communicator it builds itself, congruent to +the session communicator but never identical to it, and never NCCL. That trace +also showed the engine issuing a model's collectives on **two** different +streams (torch's graph-capture stream during capture, the serving stream +otherwise), which is why the stream bullet in *Preconditions* matters. This +paragraph is measured evidence from this repo's own runs, not something this +entry's test carries — the test certifies the pairing rule itself, in isolation. + +**Numerics — this op computes, so the result is not the exact sum.** All of +the following is measured at world size 4, bf16, on payloads with no exactness +property (standard normal cast to bf16): + +- It **is deterministic**. Eight identical eager calls return bitwise + identical tensors; eight replays of a captured call return bitwise identical + tensors and agree bitwise with the eager result; and — in probing, outside + the test — three separate 4-rank processes returned byte-identical results + for the same payloads at 1, 16, 256 and 2048 rows and for a ragged split. A + target's accuracy gate can therefore expect run-to-run reproducibility. +- It is **not** the correctly rounded exact sum. About a third of the elements + differ from an fp32 reduction rounded once to bf16 (measured 0.33 at 1, 8, + 256 and 2048 rows). +- It **is a sequential summation in the input dtype**, in ring order starting + at the destination's successor: for the rank at group position `i` the order + is `x_{i+1} + x_{i+2} + ... + x_{i+G-1} + x_i` (indices mod `G`), each + partial sum rounded to the input dtype. Matched bitwise on every element at + 1, 8 and 256 rows with hidden 2560, and at hidden 64. A consequence worth + stating plainly: **two ranks holding identical rows can get different + answers**, because their orders differ. With rank 0 holding `1.0` and the + other three holding `2^-9`, rank 0 came back with `1.0078125` (its order + accumulates the three small terms first) and rank 2 with `1.0` (its order + starts from the large one and the small terms round away). +- Its distance from the exact sum stays inside the classical bound for a + sequential sum of `G` terms, `(G-1) * u * sum_r |x_r|` elementwise, with + `u = 2^-8` for bf16, `2^-11` for fp16 and `2^-24` for fp32. Largest observed + ratio to that bound: 0.89 for bf16 (over 1, 8, 64, 256 and 2048 rows), 0.82 + for fp16, 0.80 for fp32. +- The ring order is NCCL's algorithm choice for this topology, not a promise + of the op's interface. The entry's test asserts it, so a change shows up as + a test failure rather than as drifting accuracy. + +Because of the exactness of small dyadic values, everything else in this +entry's test is asserted **bitwise** (`rtol=0, atol=0`): its payloads are +multiples of 1/8 bounded so that every partial sum, in any order, is exactly +representable. That is a tightening of the default tolerances, not a +loosening — but it is only available to a test that controls its inputs. A +caller's real activations get the bullets above. + +**CUDA graphs — the surface a decode step actually needs.** Certified at world +size 4, group `[0, 1, 2, 3]`, bf16, under torch's default +`capture_error_mode="global"`, with every rank capturing and replaying the +same graph: + +| form | under capture | +|---|---| +| `sizes = None` (even) | certified at **all 35 per-rank row counts `1..32, 64, 128, 256`** | +| `sizes = [...]` (uneven) | certified — but the split is frozen at capture, see below | + +A trtllm engine instantiates one decode graph per configured batch size — +with `cuda_graph_config.max_batch_size = 256` and no explicit `batch_sizes` +list, 35 of them at `1..32, 64, 128, 256` — keeps all of them alive in one +memory pool, and replays them interleaved as the served batch size moves. +Under attention DP a decode call's per-rank row count *is* that batch size, so +this op's captured input is `G` times it. The whole set is certified: all 35 +captured into one shared pool, none released, then replayed in four orders — +ascending, descending, shuffled, and largest-jump-first (`256, 1, 128, 2, 64, +3, ...`) — each order twice, once one replay at a time with a synchronize and +a full check between graphs, and once with all 35 issued back to back and +nothing synchronizing between them. Every replay's payload is one no earlier +call used, and every result is bitwise equal to the arithmetic reference. The +same is certified with **29 independent reduce-scatters inside each of the 35 +graphs** — 1015 captured collectives in one pool, the shape this checkpoint's +decode graph has (30 layers, layer 0 dense, so 29 expert-parallel MoE calls +each followed by one reduce-scatter). + +A replay re-runs the collective over whatever the input buffer holds at replay +time and writes into the same tensor the capture returned — keep it and read +it after each `replay()`. Replays also stay correct with eager traffic in +between: an eager even reduce-scatter at 1500, 2048 or 8192 rows per rank and +an eager uneven one between every pair of replays, all bitwise correct, which +is the shape a server has when a prefill runs between two decode steps. + +**`sizes` is frozen into a graph; row counts must be padded to capture one.** +`sizes` is a host-side argument, so a capture bakes in the split it was given +and a replay re-runs *that* split. Certified by capturing two graphs with +different sizes vectors into one pool and replaying them in both orders across +two payload rounds: each keeps its own row split and its own output length. +The consequence for an attention-DP target is concrete — a graph-captured +decode step must pad every rank to the same row count and pass `sizes = None`, +because the uneven split cannot vary per replay; the uneven form belongs to +the eager (non-captured) steps. + +**A `group`'s first-ever call cannot be captured**, because it is where the +NCCL communicator is built. Four ranks capturing their first call each raised, +out of the op: + +``` +RuntimeError: Failed, NCCL error ../tensorrt_llm/common/opUtils.cpp:175 +'unhandled cuda error (run with NCCL_DEBUG=INFO for details)' +``` + +which invalidates the capture, so the `with torch.cuda.graph(...)` block then +raises + +``` +AcceleratorError: CUDA error: operation failed due to a previous error +during capture +``` + +(`cudaErrorStreamCaptureInvalidated`). It is survivable, and the fix is one +line: make one eager call per `group` before any capture. After the failure +the same group reduce-scatters correctly eagerly, and a capture taken after +that replays correctly — both certified. One caveat for anyone catching it: +torch's graph context manager ends the capture before it restores the stream, +so a failed capture leaves its own stream current — put the resting one back +with `torch.cuda.set_stream`. + +**Cost: the even split is the cheap one.** Measured on the certified path +(world size 4, bf16, hidden 2560, 200 timed iterations after 20 warm-up calls, +CUDA events, eager), per-rank row count `n`, all times per call: + +| `n` | `sizes = None` | even `sizes` vector | uneven `sizes` vector | +|---|---|---|---| +| 1 | 23.9 us | 24.1 us | 10.2 us (degenerate: `[4, 0, 0, 0]`) | +| 8 | 24.5 us | 24.3 us | 28.1 us | +| 128 | 31.0 us | 30.7 us | 47.2 us | +| 2048 | 163.4 us | 163.6 us | 248.7 us | + +`sizes = None` and an even explicit vector are indistinguishable (ratio +1.00-1.01) and return identical bits, so a caller with even counts loses +nothing by passing them. A genuinely uneven split costs about 1.5x from 128 +rows up — so unlike the sibling all-gather, where the two forms measured +identical, here padding to an even split is a small win on cost as well as +being the only form a graph can replay. Below ~128 rows the call is +launch-bound (24-25 us across a 128x message-size range), which is the regime +a decode step runs in. + +**Splitting along another axis is the caller's problem.** This op only splits +dim 0. trtllm's module-level helper reaches other axes by reshaping around +this call — `chunk`/`split` along the target dim, `reshape`, `cat`, then a +`view` of the result — which is Python-level composition and would have to be +written from catalog entries, not hidden inside a wrapper. For the attention-DP +use (return each rank its own token rows after an expert-parallel MoE call) +dim 0 is already the token axis and no reshape is needed. + +**The forward op exists.** `torch.ops.trtllm.allgather` has the same +`(input, sizes, group)` signature and is the other half of the round trip — +it is a different op and is not this entry. diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/reducescatter.py b/tensorrt_llm/_torch/staircase/catalog/comm/reducescatter.py new file mode 100644 index 000000000000..101b0735bb5a --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/comm/reducescatter.py @@ -0,0 +1,57 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Sum every rank's copy of a tensor and hand each rank back its own rows.""" + +from typing import List, Optional + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + + +def reducescatter( + input: torch.Tensor, + sizes: Optional[List[int]], + group: List[int], +) -> torch.Tensor: + """Sum `input` across every rank in `group`, then keep this rank's slice. + + Every rank passes a tensor of the same full shape; the elementwise sum is + split along dim 0 in ascending rank order and rank `i` of the group gets + slice `i`. `sizes` gives each rank's slice height in that order, or is + `None` when the split is even. The result is a freshly allocated tensor on + every rank; `input` is left untouched. + """ + # Pure-metadata guards, each for a violation measured on this machine to + # end in silent corruption, a wedge, or a process kill rather than an error. + assert input.dim() >= 1, ( + "input must have at least one dimension; a 0-d tensor segfaults inside " + "ReducescatterOp::run_list" + ) + assert input.is_contiguous(), ( + "input must be contiguous; the op reads it as packed memory and a " + "strided view is silently reduced over the wrong elements" + ) + assert input.dtype is not torch.float8_e4m3fn, ( + "float8_e4m3fn is accepted by the op but summed as raw unsigned bytes, " + "not as floats, so the result is meaningless; reduce in bf16/fp16/fp32 " + "and quantize afterwards" + ) + if sizes is None: + assert input.shape[0] % len(group) == 0, ( + f"input.shape[0]={input.shape[0]} must be divisible by " + f"len(group)={len(group)} when sizes is None; the op takes " + "shape[0] // len(group) rows per rank and silently drops the remainder" + ) + else: + assert len(sizes) == len(group), ( + f"len(sizes)={len(sizes)} must equal len(group)={len(group)}; a short " + "list raises on the highest rank and wedges the others, a long one " + "reduces more rows than the input holds and returns garbage" + ) + assert sum(sizes) == input.shape[0], ( + f"sum(sizes)={sum(sizes)} must equal input.shape[0]={input.shape[0]}; " + "the op sizes the reduction from `sizes` alone, so a short split " + "silently drops the trailing rows and a long one reads past the input" + ) + return torch.ops.trtllm.reducescatter(input, sizes, group) diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/reducescatter_test.py b/tensorrt_llm/_torch/staircase/catalog/comm/reducescatter_test.py new file mode 100644 index 000000000000..0b74117f01db --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/comm/reducescatter_test.py @@ -0,0 +1,1619 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the reducescatter catalog entry. + +A collective cannot be exercised in one process, so this script is its own +launcher: run it plainly and it re-executes itself under `mpirun` with one +rank per device named in CUDA_VISIBLE_DEVICES, under a deadline the parent +enforces by killing the whole process group. A wedged collective hangs +rather than raising — and this op wedges rather than raising when the ranks +disagree about the split — so the deadline is what keeps a broken kernel +from taking the calling run down with it. + + CUDA_VISIBLE_DEVICES=0,1,2,3 uv run python catalog/comm/reducescatter_test.py + +That runs two jobs, in this order. The first is the test proper: every +`TESTS` entry on `world_size` ranks, and it must exit 0. The second is a +four-rank sub-job of its own that deliberately mispairs this op against an +all-gather of a *different* byte count, which is the one call-order +divergence that hangs instead of returning wrong data — it cannot live in +the first job, because a job that wedges never reports. The launcher +certifies it out of band: every rank marks the file system before issuing +the pair, none marks it afterwards, and the job is ended by its own +watchdog. That sub-job doubles as the harness's positive control, since it +is a real collective deadlock this launcher has to survive. +""" + +import os +import random +import signal +import subprocess +import sys +import tempfile +import threading +import time +from typing import Any, Dict, List, Optional, Sequence, Tuple + +import torch + +assert torch.cuda.is_available(), "reducescatter requires CUDA devices" + +_WORKER_FLAG = "--rank-worker" +_WEDGE_FLAG = "--wedge-worker" +DEADLINE_S = 1800 + +# The wedge sub-job's budget. A pair that is going to return does so in +# microseconds, so a rank still inside it GRACE seconds later is wedged; the +# watchdog then ends its own process, because a rank blocked in a NCCL kernel +# that will never complete has no other way out. CAP covers the sub-job's +# import and communicator bring-up before the pair is issued (~50 s measured) +# plus the driver's teardown of four contexts holding a spinning kernel. +WEDGE_GRACE_S = 60 +WEDGE_CAP_S = 420 + +HIDDEN = 2560 # DeepSeek-V3-Lite hidden size, the caller that motivates this + +# The row counts one engine's decode graphs are captured at. trtllm's graph +# runner instantiates one graph per configured batch size, and under attention +# data parallelism a decode call's per-rank row count *is* that batch size: +# with `cuda_graph_config.max_batch_size = 256` and the default list that is 35 +# graphs at [1..32, 64, 128, 256], all alive at once in one memory pool and +# replayed interleaved as the served batch size moves. This op's input carries +# every rank's rows, so its captured input is WORLD times that. +GRAPH_BATCH_SIZES: Tuple[int, ...] = tuple(range(1, 33)) + (64, 128, 256) + +# Reduce-scatters inside one decode graph for this checkpoint: 30 layers with +# `first_k_dense_replace = 1`, so 29 expert-parallel MoE calls, each followed +# by one reduce-scatter of the four expert windows' partial outputs. +SITES_PER_DECODE_GRAPH = 29 + +# bf16 unit roundoff: the format carries 8 significand bits, so eps = 2^-8 and +# a round-to-nearest step is off by at most 2^-9 of the running value. Used by +# the accumulation-order test's error budget. +BF16_U = 2.0**-8 + +# Bound inside the rank body, never in the launcher: importing the entry +# pulls in tensorrt_llm, which calls MPI_Init at import, and an +# MPI-initialized process cannot launch `mpirun` — measured on this host, +# mpirun then exits 1 with no output from any rank. +reducescatter: Any = None +COMM: Any = None +RANK = 0 +WORLD = 1 +GROUP: List[int] = [0] + + +def _block( + rank: int, + dest: int, + rows: int, + seed: int, + dtype: torch.dtype = torch.bfloat16, + trailing: Tuple[int, ...] = (HIDDEN,), + amp: int = 31, +) -> torch.Tensor: + """The part of `rank`'s input that ends up on group position `dest`. + + Every rank's input is the concatenation of one block per destination, so a + rank can regenerate exactly the WORLD blocks its own result is the sum of + without materializing anybody's whole input — cuRAND is deterministic for a + given (seed, shape, dtype). That is what lets the reference below be + arithmetic rather than a second collective. + + The float value set is multiples of 1/8 with |x| <= amp/8. bf16, fp16 and + fp32 all represent k/8 exactly for |k| <= 255, and a sum of WORLD of them + reaches |k| <= 4 * 31 = 124, so **every partial sum is exact in every + summation order**. That is what lets the assertions be bitwise despite the + op summing in the input dtype (see test_reduction_is_deterministic_and_ + accumulates_in_the_input_dtype for what happens when they are not). + """ + gen = torch.Generator(device="cuda").manual_seed((seed * 977 + rank) * 131 + dest + 1) + shape = (rows, *trailing) + if dtype in (torch.uint8, torch.int8, torch.int32, torch.int64): + lo, hi = (0, 8) if dtype is torch.uint8 else (-15, 16) + return torch.randint(lo, hi, shape, generator=gen, device="cuda", dtype=torch.int32).to( + dtype + ) + raw = torch.randint(-amp, amp + 1, shape, generator=gen, device="cuda", dtype=torch.int32) + return (raw.float() / 8.0).to(dtype) + + +def _input( + rank: int, + sizes: Sequence[int], + seed: int, + dtype: torch.dtype = torch.bfloat16, + trailing: Tuple[int, ...] = (HIDDEN,), + amp: int = 31, +) -> torch.Tensor: + """One rank's whole contribution: every destination's block, concatenated.""" + return torch.cat( + [_block(rank, d, n, seed, dtype, trailing, amp) for d, n in enumerate(sizes)], + dim=0, + ) + + +def _ref( + sizes: Sequence[int], + pos: int, + seed: int, + dtype: torch.dtype = torch.bfloat16, + trailing: Tuple[int, ...] = (HIDDEN,), + ranks: Optional[Sequence[int]] = None, + amp: int = 31, +) -> torch.Tensor: + """Arithmetic reference: the sum this group position is supposed to get. + + Accumulated in fp32 (int64 for integer dtypes) on this rank alone, from the + same seeds every rank uses. Never from a second collective. + """ + ranks = list(range(WORLD)) if ranks is None else list(ranks) + acc_dtype = torch.float32 if dtype.is_floating_point else torch.int64 + acc = torch.zeros((sizes[pos], *trailing), dtype=acc_dtype, device="cuda") + for r in ranks: + acc += _block(r, pos, sizes[pos], seed, dtype, trailing, amp).to(acc_dtype) + return acc.to(dtype) + + +def _assert_bitwise(out: torch.Tensor, ref: torch.Tensor, where: str) -> None: + """Gate of exactly zero, on payloads whose every partial sum is exact. + + Tightened from the default dtype-aware tolerances rather than loosened: the + reduction is a sum in the input dtype, but `_block`'s value set makes every + intermediate representable, so any difference at all is a wrong reduction + or a wrong slice, not rounding. + """ + torch.testing.assert_close(out, ref, rtol=0, atol=0, msg=lambda built: f"{where}: {built}") + + +def _sizes_vectors() -> List[List[int]]: + """Per-rank row counts an attention-DP step produces, adapted to WORLD.""" + return [ + [1] * WORLD, # uniform, but stated as an explicit sizes vector + [1 + 4 * r for r in range(WORLD)], # steadily uneven decode batches + [7 * (WORLD - r) for r in range(WORLD)], # uneven the other way + [0] + [3 + 2 * r for r in range(WORLD - 1)], # one rank with no rows + [2048] + [1 + 388 * r for r in range(WORLD - 1)], # prefill-sized, lopsided + ] + + +def _mispaired_ref(sizes: Sequence[int], pos: int, seeds: Sequence[int]) -> torch.Tensor: + """The sum a call actually computes when the ranks disagree on call order. + + `seeds[r]` is the payload rank `r` happened to be holding when this + position on the communicator came round. With `_block`'s value set every + partial sum is exact, so the prediction is bitwise regardless of the order + the addends are accumulated in. + """ + acc = torch.zeros((sizes[pos], HIDDEN), dtype=torch.float32, device="cuda") + for rank, seed in enumerate(seeds): + acc += _block(rank, pos, sizes[pos], seed).float() + return acc.to(torch.bfloat16) + + +def _rand(rank: int, dest: int, rows: int, seed: int) -> torch.Tensor: + """Payload with no exactness property, for the accumulation-order tests.""" + gen = torch.Generator(device="cuda").manual_seed((seed * 977 + rank) * 131 + dest + 1) + return torch.randn((rows, HIDDEN), generator=gen, device="cuda", dtype=torch.float32).to( + torch.bfloat16 + ) + + +def _rand_input(rank: int, rows: int, seed: int) -> torch.Tensor: + """One rank's whole contribution, drawn from the inexact value set.""" + return torch.cat([_rand(rank, d, rows, seed) for d in range(WORLD)], dim=0) + + +def _ring_chain(rows: int, pos: int, seeds: Sequence[int]) -> torch.Tensor: + """Ring-order sequential sum for group position `pos`, rounded every step. + + The order this op reduces in (certified by + test_reduction_is_deterministic_and_accumulates_in_the_input_dtype): + `x_{pos+1} + x_{pos+2} + ... + x_{pos+G-1} + x_pos`, indices mod WORLD, + where `x_r` is rank `r`'s block for `pos` drawn from `seeds[r]`. Per-rank + seeds, so a mispaired call's exact bits can be predicted too — on payloads + where a different accumulation order would give different bits. + """ + first = (pos + 1) % WORLD + chain = _rand(first, pos, rows, seeds[first]).clone() + for k in range(2, WORLD + 1): + rank = (pos + k) % WORLD + chain = chain + _rand(rank, pos, rows, seeds[rank]) + return chain + + +def _swapped_pair( + first: torch.Tensor, second: torch.Tensor, sizes: Optional[List[int]] +) -> Tuple[torch.Tensor, torch.Tensor]: + """Issue two calls, with rank 0 issuing them in the opposite order. + + Returns the results of communicator **positions** 0 and 1 — which is not + the same as the results of `first` and `second`, and that is the point. + """ + if RANK == 0: + return ( + reducescatter(second, sizes, GROUP), + reducescatter(first, sizes, GROUP), + ) + return reducescatter(first, sizes, GROUP), reducescatter(second, sizes, GROUP) + + +def _reduce_scatter_on_a_side_stream(x: torch.Tensor, side: torch.cuda.Stream) -> torch.Tensor: + """Issue one call with `side` current, joined to the caller on both ends.""" + side.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(side): + out = reducescatter(x, None, GROUP) + torch.cuda.current_stream().wait_stream(side) + return out + + +def _fraction_wrong(got: torch.Tensor, ref: torch.Tensor) -> float: + """Share of elements that differ, over the rows the two tensors share.""" + rows = min(got.shape[0], ref.shape[0]) + assert rows > 0 + return (got[:rows] != ref[:rows]).float().mean().item() + + +def _warm_up_off_capture_stream(body, reps: int = 2) -> None: + """Run `body` on a side stream, then rejoin, so a capture can follow. + + Two things have to be done before a capture and cannot be done inside one: + the group's NCCL communicator has to exist (see + test_cuda_graph_capture_of_a_first_call_raises), and torch wants the work + warmed on a non-default stream. + """ + side = torch.cuda.Stream() + side.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(side): + for _ in range(reps): + body() + torch.cuda.current_stream().wait_stream(side) + torch.cuda.synchronize() + COMM.Barrier() + + +def _capture_per_batch_size( + sites: int = 1, seed: int = 11000 +) -> Tuple[Dict[int, List[torch.Tensor]], Dict[int, Any], Dict[int, List[torch.Tensor]]]: + """One graph per GRAPH_BATCH_SIZES entry, one shared pool, all left alive. + + That is the state a graph runner ends up in: every capture after the first + goes into the first one's pool, and none of them is ever destroyed. + `sites` reduce-scatters go into each graph, each reading its own persistent + input buffer, so the payload of every site can be moved independently. + Returns (inputs, graphs, outputs), all keyed by the per-rank row count. + """ + inputs: Dict[int, List[torch.Tensor]] = {} + graphs: Dict[int, Any] = {} + outputs: Dict[int, List[torch.Tensor]] = {} + pool = None + for rows in GRAPH_BATCH_SIZES: + xs = [_input(RANK, [rows] * WORLD, seed + 13 * s) for s in range(sites)] + + def body(xs: List[torch.Tensor] = xs) -> List[torch.Tensor]: + return [reducescatter(x, None, GROUP) for x in xs] + + _warm_up_off_capture_stream(body) + graph = torch.cuda.CUDAGraph() + ctx = torch.cuda.graph(graph) if pool is None else torch.cuda.graph(graph, pool=pool) + with ctx: + outputs[rows] = body() + # Without this a capture that returned nothing would make every + # `_verify_replay` below an empty loop, i.e. a vacuous pass. + assert len(outputs[rows]) == sites, ( + f"rows={rows}: captured {len(outputs[rows])} outputs, expected {sites}" + ) + if pool is None: + pool = graph.pool() + inputs[rows], graphs[rows] = xs, graph + COMM.Barrier() + return inputs, graphs, outputs + + +def _replay_orders() -> List[Tuple[str, List[int]]]: + """Four ways a served batch size moves through the captured set.""" + asc = list(GRAPH_BATCH_SIZES) + shuffled = list(asc) + random.Random(1234).shuffle(shuffled) + # The largest jump still available at every step: 256, 1, 128, 2, 64, 3... + extremes: List[int] = [] + lo, hi = 0, len(asc) - 1 + while lo <= hi: + extremes.append(asc[hi]) + if lo != hi: + extremes.append(asc[lo]) + lo, hi = lo + 1, hi - 1 + return [ + ("ascending", asc), + ("descending", list(reversed(asc))), + ("shuffled", shuffled), + ("extremes", extremes), + ] + + +def _fresh_payload(inputs: Dict[int, List[torch.Tensor]], rows: int, seed: int) -> None: + """Overwrite every site's persistent input buffer for this row count.""" + for site, x in enumerate(inputs[rows]): + x.copy_(_input(RANK, [rows] * WORLD, seed + 13 * site)) + + +def _verify_replay( + outputs: Dict[int, List[torch.Tensor]], rows: int, seed: int, where: str +) -> None: + """Every site's captured output tensor holds this payload's result, bitwise.""" + for site, out in enumerate(outputs[rows]): + _assert_bitwise(out, _ref([rows] * WORLD, RANK, seed + 13 * site), f"{where} site={site}") + + +def test_cuda_graph_capture_of_a_first_call_raises() -> None: + """A group's first-ever call cannot be captured; the build inside fails. + + Must run before anything else touches GROUP — the failure is specifically + the NCCL communicator being built inside the capture, and it only happens + once per rank set per process. The failure is survivable, and the two + assertions after it are the workaround a caller needs: one eager call + first, then capture. + """ + rows = 16 + sizes = [rows] * WORLD + x = _input(RANK, sizes, 100) + ref = _ref(sizes, RANK, 100) + COMM.Barrier() + + graph = torch.cuda.CUDAGraph() + resting_stream = torch.cuda.current_stream() + inner: Optional[BaseException] = None + outer: Optional[BaseException] = None + try: + with torch.cuda.graph(graph): + try: + reducescatter(x, None, GROUP) + except RuntimeError as exc: + inner = exc + raise + except RuntimeError as exc: + outer = exc + finally: + # torch's context manager ends the capture before restoring the + # stream, so a capture that fails at capture_end leaves its own + # stream current. Put the resting one back by hand. + torch.cuda.set_stream(resting_stream) + del graph + + assert inner is not None, "the op accepted a first call inside a capture" + assert "NCCL error" in str(inner) and "opUtils.cpp" in str(inner), str(inner) + assert outer is not None, "capturing a group's first-ever call was accepted" + assert "operation failed due to a previous error during capture" in str(outer), str(outer) + + # Survivable: the group works eagerly straight afterwards... + torch.cuda.synchronize() + COMM.Barrier() + _assert_bitwise(reducescatter(x, None, GROUP), ref, "eager after failed capture") + COMM.Barrier() + + # ...and a capture taken after that replays correctly. + _warm_up_off_capture_stream(lambda: reducescatter(x, None, GROUP)) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + captured = reducescatter(x, None, GROUP) + x.copy_(_input(RANK, sizes, 700)) + graph.replay() + torch.cuda.synchronize() + _assert_bitwise(captured, _ref(sizes, RANK, 700), "replay after warm-up") + del graph + COMM.Barrier() + + +def test_uniform_reduce_scatter() -> None: + """sizes=None: every rank sends WORLD * rows rows and keeps `rows` of them. + + Decode-like (1, 2, 8, 32) through prefill-like (2048) per-rank row counts, + which is the whole span an attention-DP rank gets back from its + expert-parallel MoE call. + """ + for rows in (1, 2, 8, 32, 2048): + seed = 1000 + rows + sizes = [rows] * WORLD + x = _input(RANK, sizes, seed) + assert x.shape == (rows * WORLD, HIDDEN), x.shape + out = reducescatter(x, None, GROUP) + assert out.shape == (rows, HIDDEN), out.shape + assert out.dtype is x.dtype and out.device == x.device + assert out.is_contiguous() + _assert_bitwise(out, _ref(sizes, RANK, seed), f"uniform rows={rows}") + COMM.Barrier() + + +def test_ragged_reduce_scatter() -> None: + """sizes=[...]: the split is uneven, this rank keeps sizes[my position]. + + Attention data parallelism produces exactly this — each rank owns its own + requests — so the ragged form is the one a non-padded step takes, including + the case where one rank has no rows at all. + """ + for i, sizes in enumerate(_sizes_vectors()): + seed = 3000 + 7 * i + x = _input(RANK, sizes, seed) + assert x.shape == (sum(sizes), HIDDEN), x.shape + out = reducescatter(x, sizes, GROUP) + assert out.shape == (sizes[RANK], HIDDEN), (sizes, out.shape) + assert out.dtype is x.dtype and out.is_contiguous() + _assert_bitwise(out, _ref(sizes, RANK, seed), f"ragged sizes={sizes}") + COMM.Barrier() + + +def test_output_is_fresh_and_input_is_untouched() -> None: + """The op allocates its result; the caller keeps ownership of `input`.""" + rows, seed = 8, 4100 + sizes = [rows] * WORLD + x = _input(RANK, sizes, seed) + before = x.clone() + out = reducescatter(x, None, GROUP) + torch.cuda.synchronize() + assert out.data_ptr() != x.data_ptr(), "output aliased the input buffer" + _assert_bitwise(x, before, "input after the call") + + # Clobbering the input afterwards cannot reach the result. + x.fill_(-7.0) + torch.cuda.synchronize() + _assert_bitwise(out, _ref(sizes, RANK, seed), "output after input clobber") + COMM.Barrier() + + +def test_trailing_dims_are_preserved() -> None: + """dim 0 is the scatter axis; every other dim is carried through untouched.""" + cases: List[Tuple[Tuple[int, ...], bool]] = [ + ((2, HIDDEN // 2), False), + ((2, HIDDEN // 2), True), + ((), False), # 1-D: a flat vector scatters into WORLD pieces + ((), True), + ] + for i, (trailing, ragged) in enumerate(cases): + seed = 5000 + 11 * i + sizes = [1 + 4 * r for r in range(WORLD)] if ragged else [6] * WORLD + x = _input(RANK, sizes, seed, trailing=trailing) + out = reducescatter(x, sizes if ragged else None, GROUP) + assert out.shape == (sizes[RANK], *trailing), (trailing, sizes, out.shape) + _assert_bitwise( + out, + _ref(sizes, RANK, seed, trailing=trailing), + f"trailing={trailing} ragged={ragged}", + ) + COMM.Barrier() + + +def test_dtypes_reduce_arithmetically() -> None: + """The dtypes whose sum this op actually computes, in both forms. + + fp16 and fp32 alongside bf16, and the integer widths, which reduce as + integer sums (kept small enough here not to overflow the narrow ones). + float8 is a separate test: it is accepted and summed as raw bytes. + """ + ragged = [1 + 4 * r for r in range(WORLD)] + uniform = [16] * WORLD + for i, dtype in ( + (0, torch.float16), + (1, torch.float32), + (2, torch.int32), + (3, torch.int64), + (4, torch.uint8), + (5, torch.int8), + ): + seed = 6000 + 13 * i + x = _input(RANK, uniform, seed, dtype) + out = reducescatter(x, None, GROUP) + assert out.dtype is dtype and out.shape == (16, HIDDEN) + _assert_bitwise(out, _ref(uniform, RANK, seed, dtype), f"uniform {dtype}") + COMM.Barrier() + + xr = _input(RANK, ragged, seed, dtype) + outr = reducescatter(xr, ragged, GROUP) + assert outr.dtype is dtype and outr.shape == (ragged[RANK], HIDDEN) + _assert_bitwise(outr, _ref(ragged, RANK, seed, dtype), f"ragged {dtype}") + COMM.Barrier() + + +def test_group_selects_a_rank_subset() -> None: + """`group` names MPI session ranks; the slice index is the position in it. + + The second subset is the discriminating one: it excludes rank 0, so a rank + that indexed the split by its MPI rank instead of by its position in the + group would read a different slice (and rank WORLD-1 would read past the + end). Ranks outside the subset must not call. + """ + subsets = [[0, 1]] + if WORLD >= 4: + subsets.append([WORLD - 2, WORLD - 1]) + for i, subset in enumerate(subsets): + rows, seed = 4, 7100 + 31 * i + sizes = [rows] * len(subset) + if RANK in subset: + pos = subset.index(RANK) + out = reducescatter(_input(RANK, sizes, seed), None, subset) + assert out.shape == (rows, HIDDEN), out.shape + _assert_bitwise(out, _ref(sizes, pos, seed, ranks=subset), f"subset {subset} uniform") + ragged = [2, 6] if len(subset) == 2 else [2] * len(subset) + outr = reducescatter(_input(RANK, ragged, seed), ragged, subset) + assert outr.shape == (ragged[pos], HIDDEN), outr.shape + _assert_bitwise(outr, _ref(ragged, pos, seed, ranks=subset), f"subset {subset} ragged") + COMM.Barrier() + + +def test_group_order_does_not_change_the_output() -> None: + """The split is ordered by ascending rank, whatever order `group` lists.""" + rows, seed = 4, 7200 + sizes = [rows] * WORLD + x = _input(RANK, sizes, seed) + out = reducescatter(x, None, list(reversed(GROUP))) + _assert_bitwise(out, _ref(sizes, RANK, seed), "reversed group list") + COMM.Barrier() + + +def test_reduction_is_deterministic_and_accumulates_in_the_input_dtype() -> None: + """What the sum is, exactly — the fact a target's accuracy gate rests on. + + On payloads with no exactness property (standard normal, cast to bf16): + + 1. It is **deterministic**: eight identical eager calls return bitwise + identical tensors, and so do eight replays of a captured call. + 2. It is **not** the correctly rounded exact sum. The op accumulates in the + input dtype, so about a third of the elements differ from an fp32 + reduction rounded once. + 3. It **is** a sequential summation in the input dtype, in ring order + starting at the destination's successor — matched bitwise on every + element here. That order is NCCL's algorithm choice for this topology; + if this assertion ever fails, the reduction order changed and the + contract's Numerics note has to be re-measured rather than this gate + loosened. + 4. Its distance from the exact sum stays inside the classical bound for a + sequential sum of WORLD terms, `(WORLD-1) * u * sum_r |x_r|` with + `u = 2^-8` for bf16 — the bound holds elementwise with the largest + observed ratio 0.89 (measured at 1, 8, 64, 256 and 2048 rows). + """ + for rows in (1, 8, 256): + seed = 8000 + rows + x = torch.cat([_rand(RANK, d, rows, seed) for d in range(WORLD)], dim=0) + outs = [reducescatter(x, None, GROUP) for _ in range(8)] + torch.cuda.synchronize() + for k, o in enumerate(outs[1:], start=1): + _assert_bitwise(o, outs[0], f"repeat {k} at rows={rows}") + + mine = [_rand(r, RANK, rows, seed) for r in range(WORLD)] + exact = torch.zeros(rows, HIDDEN, dtype=torch.float32, device="cuda") + absum = torch.zeros(rows, HIDDEN, dtype=torch.float32, device="cuda") + for chunk in mine: + exact += chunk.float() + absum += chunk.float().abs() + + # (2) not the correctly rounded exact sum + differs = (outs[0] != exact.to(torch.bfloat16)).float().mean().item() + assert differs > 0.05, ( + f"rows={rows}: only {differs:.4f} of elements differ from an fp32 " + "reduction — the op may have started accumulating in fp32" + ) + + # (3) a sequential chain in ring order, starting at my successor + chain = mine[(RANK + 1) % WORLD].clone() + for k in range(2, WORLD + 1): + chain = chain + mine[(RANK + k) % WORLD] + _assert_bitwise(outs[0], chain, f"ring-order chain at rows={rows}") + + # (4) inside the classical sequential-summation error bound + err = (outs[0].float() - exact).abs() + budget = (WORLD - 1) * BF16_U * absum + over = (err > budget).sum().item() + assert over == 0, ( + f"rows={rows}: {over} elements outside the sequential-summation " + f"bound, worst ratio {(err / budget.clamp_min(1e-30)).max().item():.4f}" + ) + COMM.Barrier() + + # Determinism under capture, and the captured result equals the eager one. + rows, seed = 16, 8500 + x = torch.cat([_rand(RANK, d, rows, seed) for d in range(WORLD)], dim=0) + eager = reducescatter(x, None, GROUP) + torch.cuda.synchronize() + _warm_up_off_capture_stream(lambda: reducescatter(x, None, GROUP)) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + captured = reducescatter(x, None, GROUP) + for k in range(8): + graph.replay() + torch.cuda.synchronize() + _assert_bitwise(captured, eager, f"replay {k} against the eager result") + del graph + COMM.Barrier() + + +def test_uniform_form_and_an_explicit_even_split_agree_bitwise() -> None: + """`sizes=None` and an explicit even `sizes` vector return the same bits. + + Worth pinning because the two are not obliged to take the same path — and + because an attention-DP target switches between them (padded, captured + decode steps pass None; eager steps pass the counts), so a caller needs to + know that the switch does not perturb the arithmetic. Driven on payloads + with no exactness property, where a different reduction order would show. + """ + for rows in (1, 8, 256, 2048): + seed = 9000 + rows + x = torch.cat([_rand(RANK, d, rows, seed) for d in range(WORLD)], dim=0) + a = reducescatter(x, None, GROUP) + b = reducescatter(x, [rows] * WORLD, GROUP) + torch.cuda.synchronize() + _assert_bitwise(b, a, f"explicit even split at rows={rows}") + COMM.Barrier() + + +def test_round_trip_with_a_gather_returns_each_rank_its_own_rows() -> None: + """The attention-DP MoE round trip, end to end, against a local reference. + + Gather every rank's tokens, let each rank apply its own expert window to + the whole token set, reduce-scatter the four partial outputs. The claim + under test is that the scatter undoes the gather: rank i gets back exactly + the rows rank i contributed, holding the sum of the four windows. + + `torch.ops.trtllm.allgather` appears here as a **fixture**, to build the + input the way the target will — never as the reference. The reference is + this rank's own rows times the sum of the four window scales, computed + locally. The gathered tensor is also checked against a local `torch.cat`, + so a wrong gather cannot be mistaken for a right reduce-scatter. + """ + sizes = [1 + 4 * r for r in range(WORLD)] + seed = 10500 + # Scales are k/4 and payloads k/8 with |k| <= 7, so every product and every + # partial sum stays exactly representable and the gate can be bitwise. + own = [_block(r, 0, sizes[r], seed, amp=7) for r in range(WORLD)] + mine = own[RANK] + gathered = torch.ops.trtllm.allgather(mine, sizes, GROUP) + torch.cuda.synchronize() + _assert_bitwise(gathered, torch.cat(own, dim=0), "gather fixture") + + scale = float(RANK + 1) / 4.0 + partial = (gathered.float() * scale).to(torch.bfloat16) + out = reducescatter(partial, sizes, GROUP) + torch.cuda.synchronize() + total_scale = sum((r + 1) / 4.0 for r in range(WORLD)) + assert out.shape == mine.shape, (out.shape, mine.shape) + _assert_bitwise(out, (mine.float() * total_scale).to(torch.bfloat16), "round trip") + COMM.Barrier() + + # The bare identity: reduce-scattering a gather returns WORLD copies summed. + out2 = reducescatter(torch.ops.trtllm.allgather(mine, sizes, GROUP), sizes, GROUP) + torch.cuda.synchronize() + _assert_bitwise(out2, (mine.float() * WORLD).to(torch.bfloat16), "rs(ag(x))") + COMM.Barrier() + + +def test_cuda_graph_at_every_engine_batch_size() -> None: + """Captured at all 35 engine batch sizes, then replayed interleaved. + + Capturing one row count proves nothing about an engine, which holds every + graph in GRAPH_BATCH_SIZES alive at once and replays them in whatever order + the served batch size takes. Each replay here gets a payload no earlier + call used, so a graph that silently did not re-run keeps the previous + answer and fails the bitwise gate. + + Two passes per order: one replay at a time with a synchronize and a full + check before the next graph runs, then the whole order issued back to back + with no synchronization in between — 35 differently sized collectives + queued as one run of work, which is the only pass that can catch a hazard + that needs the next call to already be in flight. + """ + inputs, graphs, outputs = _capture_per_batch_size() + assert len(graphs) == len(GRAPH_BATCH_SIZES) == 35 + seed = 110000 + for name, order in _replay_orders(): + for pos, rows in enumerate(order): + seed += 1000 + _fresh_payload(inputs, rows, seed) + graphs[rows].replay() + torch.cuda.synchronize() + _verify_replay(outputs, rows, seed, f"{name}/checked@{pos} rows={rows}") + COMM.Barrier() + staged: Dict[int, int] = {} + for rows in order: + seed += 1000 + staged[rows] = seed + _fresh_payload(inputs, rows, seed) + for rows in order: + graphs[rows].replay() + torch.cuda.synchronize() + for pos, rows in enumerate(order): + _verify_replay(outputs, rows, staged[rows], f"{name}/burst@{pos} rows={rows}") + COMM.Barrier() + # Discrimination for a gate of exactly 0: what a replay that did not happen + # leaves behind is the previous payload's result. The two consecutive + # references have to be far apart for that to be caught, which this checks. + prev = _ref([16] * WORLD, RANK, seed - 1000) + last = _ref([16] * WORLD, RANK, seed) + margin = (last.float() - prev.float()).abs().max().item() + assert margin > 8.0, f"stale-replay margin {margin}" + del graphs, outputs, inputs + COMM.Barrier() + + +def test_cuda_graph_one_site_per_moe_layer_in_every_batch_size_graph() -> None: + """The decode graph this checkpoint captures, at all 35 batch sizes. + + 1015 captured reduce-scatters, one memory pool. The sites are independent — + each reads its own buffer and every one gets a distinct payload, so a site + that returned a neighbour's data would fail the bitwise gate. Two replay + orders rather than four: the per-order cost is 29x the single-site test's + and the axis this one adds is site count, not order. + """ + inputs, graphs, outputs = _capture_per_batch_size(sites=SITES_PER_DECODE_GRAPH, seed=21000) + seed = 310000 + for name, order in _replay_orders()[2:]: + for pos, rows in enumerate(order): + seed += 1000 + _fresh_payload(inputs, rows, seed) + graphs[rows].replay() + torch.cuda.synchronize() + _verify_replay(outputs, rows, seed, f"{name}/checked@{pos} rows={rows}") + COMM.Barrier() + staged: Dict[int, int] = {} + for rows in order: + seed += 1000 + staged[rows] = seed + _fresh_payload(inputs, rows, seed) + for rows in order: + graphs[rows].replay() + torch.cuda.synchronize() + for pos, rows in enumerate(order): + _verify_replay(outputs, rows, staged[rows], f"{name}/burst@{pos} rows={rows}") + COMM.Barrier() + del graphs, outputs, inputs + COMM.Barrier() + + +def test_cuda_graph_replays_survive_eager_calls_of_other_shapes() -> None: + """Between two decode replays, a server runs work no graph holds. + + A prefill runs eagerly at a row count far outside the captured set, and a + ragged reduce-scatter — the other split this op has — runs eagerly on the + same communicator the graphs baked in. + """ + inputs, graphs, outputs = _capture_per_batch_size(seed=41000) + seed = 510000 + for i, rows in enumerate(reversed(GRAPH_BATCH_SIZES)): + wide = (2048, 8192, 1500)[i % 3] + wide_sizes = [wide] * WORLD + _assert_bitwise( + reducescatter(_input(RANK, wide_sizes, 610000 + i), None, GROUP), + _ref(wide_sizes, RANK, 610000 + i), + f"eager uniform {wide} rows", + ) + ragged = [(11 * i + 3 * r) % 97 for r in range(WORLD)] + _assert_bitwise( + reducescatter(_input(RANK, ragged, 710000 + i), ragged, GROUP), + _ref(ragged, RANK, 710000 + i), + f"eager ragged {ragged}", + ) + seed += 1000 + _fresh_payload(inputs, rows, seed) + graphs[rows].replay() + torch.cuda.synchronize() + _verify_replay(outputs, rows, seed, f"after-eager@{i} rows={rows}") + COMM.Barrier() + del graphs, outputs, inputs + COMM.Barrier() + + +def test_cuda_graph_holds_the_sizes_vector_it_captured() -> None: + """`sizes` is a host argument: a replay re-runs the split it was captured with. + + Two graphs with different sizes vectors are captured into one pool and + replayed in both orders; each keeps its own row split and its own output + length. This is why a graph-captured decode step needs padded (uniform) + row counts — the ragged split cannot be varied per replay. + """ + first = [1 + 4 * r for r in range(WORLD)] + second = [7 * (WORLD - r) for r in range(WORLD)] + seed = 81000 + buffers, graphs, captured = {}, {}, {} + pool = None + for tag, sizes in (("first", first), ("second", second)): + x = _input(RANK, sizes, seed) + _warm_up_off_capture_stream(lambda x=x, sizes=sizes: reducescatter(x, sizes, GROUP)) + graph = torch.cuda.CUDAGraph() + ctx = torch.cuda.graph(graph) if pool is None else torch.cuda.graph(graph, pool=pool) + with ctx: + captured[tag] = reducescatter(x, sizes, GROUP) + pool = graph.pool() if pool is None else pool + buffers[tag], graphs[tag] = x, graph + assert captured[tag].shape == (sizes[RANK], HIDDEN), captured[tag].shape + COMM.Barrier() + + for round_seed in (82000, 83000): + for tag, sizes in (("second", second), ("first", first)): + buffers[tag].copy_(_input(RANK, sizes, round_seed)) + graphs[tag].replay() + torch.cuda.synchronize() + _assert_bitwise( + captured[tag], + _ref(sizes, RANK, round_seed), + f"ragged replay {tag} sizes={sizes}", + ) + COMM.Barrier() + del graphs, captured, buffers + COMM.Barrier() + + +def test_the_gate_discriminates_a_wrong_reduce_scatter() -> None: + """The bitwise gate rejects every plausible wrong result, by a wide margin. + + A gate of exactly 0 cannot be too loose, but it can be blind: if every + rank's contribution looked alike, a missing addend or a wrong slice would + pass it. These are the four ways this op could plausibly be wrong — the + slice of a neighbouring group position, one rank's contribution missing, + no reduction at all (this rank's own slice returned untouched), and the + right data cut at the wrong offsets under a rotated ragged split. Every one + of them has the right shape, so shape checking alone would let all four + through. + """ + rows, seed = 8, 91000 + sizes = [rows] * WORLD + uniform = reducescatter(_input(RANK, sizes, seed), None, GROUP) + _assert_bitwise(uniform, _ref(sizes, RANK, seed), "uniform baseline") + ragged = [1 + 4 * r for r in range(WORLD)] + rag_out = reducescatter(_input(RANK, ragged, seed), ragged, GROUP) + _assert_bitwise(rag_out, _ref(ragged, RANK, seed), "ragged baseline") + + mine = [_block(r, RANK, rows, seed) for r in range(WORLD)] + missing_one = mine[0].float() + for c in mine[1:-1]: + missing_one = missing_one + c.float() + wrong = [ + ("neighbouring slice", uniform, _ref(sizes, (RANK + 1) % WORLD, seed)), + ("one rank never contributed", uniform, missing_one.to(torch.bfloat16)), + ("not reduced at all", uniform, mine[RANK]), + ( + "neighbouring slice, ragged split", + rag_out, + _ref(ragged, (RANK + 1) % WORLD, seed), + ), + ] + for name, got, variant in wrong: + # A wrong ragged slice has a different row count, so the comparison is + # over the rows the two share — the point is the values, not the shape. + n = min(variant.shape[0], got.shape[0]) + assert n > 0, name + diff = (variant[:n].float() - got[:n].float()).abs() + fraction = (diff != 0).float().mean().item() + assert fraction > 0.9, f"{name}: only {fraction:.4f} of elements differ" + assert diff.max().item() > 1.0, f"{name}: max difference {diff.max().item()}" + COMM.Barrier() + + +def test_calls_pair_by_position_and_a_swapped_pair_realigns() -> None: + """Calls pair by their **position** on the communicator, not by intent. + + One rank issuing two same-shaped calls in the opposite order to everybody + else is not detected: both calls return, at the right shape, with finite + values and no diagnostic of any kind. What each rank gets back is the sum + of whatever payloads happened to meet at that position — asserted here + bitwise against a locally computed mix, which is the strongest form of + "silently wrong" there is. + + This op differs from the sibling all-gather in *how much* it corrupts. A + gather hands each rank one block per rank, so a single disagreeing rank + spoils one block: 1/4 of the elements at world size 4. Here the mispaired + rank's payload is an addend of **every** element of every rank's slice, so + every rank is wrong nearly everywhere (measured 0.982-0.993 against the + floor of 0.9 below; the residue is elements where the two payloads happened + to agree). + + Both `sizes` forms are driven, because they are not the same NCCL traffic: + the even form issues one `ncclReduceScatter` while the ragged form issues a + grouped `ncclReduce` per rank. The group pairs atomically — the ragged form + mispairs exactly like the even one, rather than interleaving per-root. + + Finally, a swapped *pair* leaves the positions aligned once both calls are + made, so the plain call at the end of each round is bitwise correct. + """ + rows = 12 + ragged = [1 + 4 * r for r in range(WORLD)] + for tag, sizes in (("even", None), ("ragged", ragged)): + split = [rows] * WORLD if sizes is None else sizes + base = 160000 if sizes is None else 165000 + first_seed, second_seed = base, base + 500 + first = _input(RANK, split, first_seed) + second = _input(RANK, split, second_seed) + COMM.Barrier() + + pos0, pos1 = _swapped_pair(first, second, sizes) + torch.cuda.synchronize() + for name, got in (("pos0", pos0), ("pos1", pos1)): + assert got.shape == (split[RANK], HIDDEN), (tag, name, got.shape) + assert bool(torch.isfinite(got.float()).all()), f"{tag} {name}: not finite" + + # What each position actually computed: rank 0 was one call out of step. + _assert_bitwise( + pos0, + _mispaired_ref(split, RANK, [second_seed] + [first_seed] * (WORLD - 1)), + f"{tag}: position 0 is the mix of second(rank 0) and first(others)", + ) + _assert_bitwise( + pos1, + _mispaired_ref(split, RANK, [first_seed] + [second_seed] * (WORLD - 1)), + f"{tag}: position 1 is the mix of first(rank 0) and second(others)", + ) + + # ...and that is not what this rank asked for, on every rank. + intended = _ref(split, RANK, second_seed if RANK == 0 else first_seed) + wrong = _fraction_wrong(pos0, intended) + assert wrong > 0.9, f"{tag}: only {wrong:.4f} of elements differ from intent" + COMM.Barrier() + + _assert_bitwise( + reducescatter(_input(RANK, split, base + 900), sizes, GROUP), + _ref(split, RANK, base + 900), + f"{tag}: plain call after a swapped pair", + ) + COMM.Barrier() + + +def test_a_mispaired_result_is_deterministic_rather_than_noise() -> None: + """Because this op computes, "wrong" could have meant "unreproducible". + + It does not. On payloads with no exactness property, a mispaired call is + bitwise identical across repeats *and* bitwise equal to the ring-order + chain of the addends that met — the same accumulation order the aligned + call uses, applied to the wrong operands. So the mispairing perturbs which + tensors are summed and nothing else. + + That is the worse outcome for a caller, not the better one: re-running the + step reproduces the wrong answer exactly, so a divergence cannot be found + by looking for run-to-run instability. + """ + rows = 12 + first_seed, second_seed = 162000, 162500 + first = _rand_input(RANK, rows, first_seed) + second = _rand_input(RANK, rows, second_seed) + COMM.Barrier() + + seen: List[Tuple[torch.Tensor, torch.Tensor]] = [] + for _ in range(4): + pos0, pos1 = _swapped_pair(first, second, None) + torch.cuda.synchronize() + seen.append((pos0, pos1)) + COMM.Barrier() + for k, (pos0, pos1) in enumerate(seen[1:], start=1): + _assert_bitwise(pos0, seen[0][0], f"repeat {k} at position 0") + _assert_bitwise(pos1, seen[0][1], f"repeat {k} at position 1") + + _assert_bitwise( + seen[0][0], + _ring_chain(rows, RANK, [second_seed] + [first_seed] * (WORLD - 1)), + "position 0 is the ring-order sum of the mispaired addends", + ) + _assert_bitwise( + seen[0][1], + _ring_chain(rows, RANK, [first_seed] + [second_seed] * (WORLD - 1)), + "position 1 is the ring-order sum of the mispaired addends", + ) + # Discrimination for those two bitwise gates: the aligned result of the + # same call is a different tensor nearly everywhere. + aligned = _ring_chain(rows, RANK, [second_seed if RANK == 0 else first_seed] * WORLD) + wrong = _fraction_wrong(seen[0][0], aligned) + assert wrong > 0.9, f"only {wrong:.4f} of elements differ from the aligned sum" + COMM.Barrier() + + _assert_bitwise( + reducescatter(_input(RANK, [rows] * WORLD, 163000), None, GROUP), + _ref([rows] * WORLD, RANK, 163000), + "plain call after four swapped pairs", + ) + COMM.Barrier() + + +def test_an_extra_call_on_one_rank_misaligns_until_the_counts_match() -> None: + """An odd number of extra calls does not realign; a swapped pair does. + + Rank 0 issues one call the others never issue, then all ranks issue four + calls they intend to agree on. Every one of those five positions is + mispaired by exactly one place, and it stays that way — there is no + resynchronization point inside a stream of collectives, so the divergence + is permanent rather than transient. + + The other ranks issue one catch-up call at the end, which is what makes + this test terminate: the counts have to match for the last position to + complete at all. Once they do, alignment is restored and the plain call at + the end is bitwise correct. A caller with one rank permanently ahead gets + the wedge instead, at whatever point the process next waits on the device. + """ + rows, calls = 10, 4 + sizes = [rows] * WORLD + shared = [170000 + 100 * k for k in range(calls)] + extra_seed, catchup_seed = 179000, 179500 + COMM.Barrier() + + outs: List[torch.Tensor] = [] + if RANK == 0: + outs.append(reducescatter(_input(RANK, sizes, extra_seed), None, GROUP)) + for seed in shared: + outs.append(reducescatter(_input(RANK, sizes, seed), None, GROUP)) + else: + for seed in shared: + outs.append(reducescatter(_input(RANK, sizes, seed), None, GROUP)) + outs.append(reducescatter(_input(RANK, sizes, catchup_seed), None, GROUP)) + torch.cuda.synchronize() + assert len(outs) == calls + 1 + + # Position p carried rank 0's p-th issued payload and everybody else's. + rank0_order = [extra_seed] + shared + others_order = list(shared) + [catchup_seed] + for pos in range(calls + 1): + mix = [rank0_order[pos]] + [others_order[pos]] * (WORLD - 1) + _assert_bitwise( + outs[pos], _mispaired_ref(sizes, RANK, mix), f"off-by-one at position {pos}" + ) + intended = _ref(sizes, RANK, rank0_order[pos] if RANK == 0 else others_order[pos]) + wrong = _fraction_wrong(outs[pos], intended) + assert wrong > 0.9, f"position {pos}: only {wrong:.4f} of elements differ" + COMM.Barrier() + + _assert_bitwise( + reducescatter(_input(RANK, sizes, 178000), None, GROUP), + _ref(sizes, RANK, 178000), + "plain call after the call counts were equalized", + ) + COMM.Barrier() + + +def test_mispaired_against_an_all_gather_of_equal_byte_count_corrupts_silently() -> None: + """A different collective at the same byte count is not detected either. + + What pairs is position, not the identity of the op: with rank 0 issuing + `allgather` where the others issue `reducescatter`, both calls return, at + the shapes their own arguments imply, with finite values and no error. Two + calls of this pair move the same number of elements as NCCL counts them — + an even reduce-scatter of `[WORLD*n, H]` and an all-gather of `[n, H]` both + carry `n*H` — which is what keeps it from hanging. The unequal-byte-count + version of exactly this mispairing wedges instead, and is certified by the + launcher's sub-job rather than here. + + Measured on the certified path: the reduce-scatter comes back 0.984-0.989 + wrong and the gather 0.246-0.741 wrong (rank-dependent, since a gather's + blocks are spoiled individually), both bitwise stable across repeats and + across processes. The floors below are 0.9 and 0.2. `allgather` is a + fixture here, never a reference — every reference in this file is + arithmetic. + + The pair realigns: a plain reduce-scatter afterwards is bitwise correct, + checked after each repeat. + """ + rows = 16 + sizes = [rows] * WORLD + rs_seed, ag_seed = 190000, 190500 + rs_in = _input(RANK, sizes, rs_seed) + ag_in = _block(RANK, 0, rows, ag_seed) + assert rs_in.shape[0] // WORLD * HIDDEN == ag_in.numel(), "byte counts differ" + rs_ref = _ref(sizes, RANK, rs_seed) + ag_ref = torch.cat([_block(r, 0, rows, ag_seed) for r in range(WORLD)], dim=0) + COMM.Barrier() + + seen: List[Tuple[torch.Tensor, torch.Tensor]] = [] + for repeat in range(2): + if RANK == 0: + got_ag = torch.ops.trtllm.allgather(ag_in, None, GROUP) + got_rs = reducescatter(rs_in, None, GROUP) + else: + got_rs = reducescatter(rs_in, None, GROUP) + got_ag = torch.ops.trtllm.allgather(ag_in, None, GROUP) + torch.cuda.synchronize() + assert got_rs.shape == (rows, HIDDEN), got_rs.shape + assert got_ag.shape == (rows * WORLD, HIDDEN), got_ag.shape + assert bool(torch.isfinite(got_rs.float()).all()), "reduce-scatter not finite" + assert bool(torch.isfinite(got_ag.float()).all()), "gather not finite" + rs_wrong = _fraction_wrong(got_rs, rs_ref) + ag_wrong = _fraction_wrong(got_ag, ag_ref) + assert rs_wrong > 0.9, f"repeat {repeat}: reduce-scatter {rs_wrong:.4f} wrong" + assert ag_wrong > 0.2, f"repeat {repeat}: gather {ag_wrong:.4f} wrong" + seen.append((got_rs, got_ag)) + COMM.Barrier() + + _assert_bitwise( + reducescatter(_input(RANK, sizes, 191000 + repeat), None, GROUP), + _ref(sizes, RANK, 191000 + repeat), + f"plain call after cross-op repeat {repeat}", + ) + COMM.Barrier() + + _assert_bitwise(seen[1][0], seen[0][0], "cross-op reduce-scatter across repeats") + _assert_bitwise(seen[1][1], seen[0][1], "cross-op gather across repeats") + COMM.Barrier() + + +def test_the_stream_the_call_lands_on_is_not_part_of_the_match() -> None: + """The op runs on whatever stream is current, and ranks need not agree. + + A serving engine moves the current stream under the model: the same forward + runs on torch's graph-capture stream while a decode graph is being captured + and on the serving stream otherwise, and a target cannot pin either. So + "must the ranks agree on the stream?" is a precondition question. They do + not have to — what pairs the calls is their order on the communicator. + Three shapes: every rank on a side stream, one rank on a side stream while + the others stay on the default, and ranks alternating in opposite patterns + so that at every call index they disagree. + + Driven twice: once on the exact value set, where a wrong pairing shows, and + once on payloads with no exactness property against the ring-order chain, + where a changed *accumulation* order would show as well. Neither moves. + """ + side = torch.cuda.Stream() + rows = 24 + sizes = [rows] * WORLD + cases = ( + ("every rank on a side stream", lambda i: True), + ("only rank 0 on a side stream", lambda i: RANK == 0), + ("ranks alternating opposite", lambda i: (i + RANK) % 2 == 0), + ) + for c, (name, use_side) in enumerate(cases): + for i in range(4): + seed = 150000 + 1000 * c + 13 * i + x = _input(RANK, sizes, seed) + out = ( + _reduce_scatter_on_a_side_stream(x, side) + if use_side(i) + else reducescatter(x, None, GROUP) + ) + _assert_bitwise(out, _ref(sizes, RANK, seed), f"{name} i={i}") + COMM.Barrier() + + for c, (name, use_side) in enumerate(cases): + for i in range(4): + seed = 155000 + 1000 * c + 13 * i + x = _rand_input(RANK, rows, seed) + out = ( + _reduce_scatter_on_a_side_stream(x, side) + if use_side(i) + else reducescatter(x, None, GROUP) + ) + torch.cuda.synchronize() + _assert_bitwise(out, _ring_chain(rows, RANK, [seed] * WORLD), f"{name} ring i={i}") + COMM.Barrier() + + +def test_float8_is_summed_as_raw_bytes() -> None: + """float8_e4m3fn is accepted and reduced as unsigned bytes, not as floats. + + Measured, and the reason the wrapper rejects the dtype: an all-gather moves + fp8 activations correctly, so a caller doing post-quantization dispatch + would reasonably expect the return leg to work too. It does not — the sum + is of the e4m3 *bit patterns*, wrapping at 256. + """ + rows = 4 + x = torch.full((rows * WORLD, 8), float(RANK + 1), dtype=torch.float32, device="cuda").to( + torch.float8_e4m3fn + ) + raw = torch.ops.trtllm.reducescatter(x, None, GROUP) + torch.cuda.synchronize() + byte_sum = torch.zeros(rows, 8, dtype=torch.int64, device="cuda") + for r in range(WORLD): + one = torch.full((rows, 8), float(r + 1), dtype=torch.float32, device="cuda").to( + torch.float8_e4m3fn + ) + byte_sum += one.view(torch.uint8).to(torch.int64) + _assert_bitwise( + raw.view(torch.uint8), + (byte_sum % 256).to(torch.uint8), + "float8 reduced as bytes", + ) + # ...and that is nowhere near the float sum it looks like it should be. + assert abs(raw.float()[0, 0].item() - sum(range(1, WORLD + 1))) > 1.0 + COMM.Barrier() + + try: + reducescatter(x, None, GROUP) + except AssertionError: + pass + else: + raise AssertionError("wrapper accepted a float8_e4m3fn input") + COMM.Barrier() + + +def test_unsupported_dtypes_raise_and_poison_every_later_collective() -> None: + """fp64 and the other float8 formats raise — and the raise is terminal. + + Runs last, and has to: the raise leaves NCCL's group state unbalanced (the + op opens a group before it converts the dtype and never closes it), so + every collective the process issues afterwards returns garbage. Measured + here for this op; measured in probing for the sibling all-gather and + all-reduce on the same group too, so it is the process that is finished, + not this entry's communicator. + + The last assertion therefore asserts a defect. If it ever fails, upstream + has balanced the group and the contract's *Preconditions* wording — "not + survivable" — has to be re-measured, not this gate loosened. + """ + rows, seed = 8, 96000 + sizes = [rows] * WORLD + _assert_bitwise( + reducescatter(_input(RANK, sizes, seed), None, GROUP), + _ref(sizes, RANK, seed), + "baseline before the unsupported dtypes", + ) + COMM.Barrier() + + for dtype in (torch.float64, torch.float8_e5m2): + x = torch.ones(4 * WORLD, 8, dtype=torch.float32, device="cuda").to(dtype) + try: + reducescatter(x, None, GROUP) + except RuntimeError as exc: + assert "unsupported data type" in str(exc), str(exc) + else: + raise AssertionError(f"{dtype} was accepted") + COMM.Barrier() + + after = reducescatter(_input(RANK, sizes, seed), None, GROUP) + torch.cuda.synchronize() + wrong = (after != _ref(sizes, RANK, seed)).float().mean().item() + assert wrong > 0.5, ( + f"only {wrong:.4f} of elements are wrong after the unsupported-dtype " + "raise — the raise used to destroy every later collective in the " + "process; re-measure the contract's dtype precondition" + ) + COMM.Barrier() + + +def test_wrapper_guards_a_non_contiguous_input() -> None: + """The wrapper's assert stands where the op itself is silently wrong.""" + rows, seed = 8, 92000 + sizes = [rows] * WORLD + values = _input(RANK, sizes, seed) + padded = torch.zeros(rows * WORLD, HIDDEN * 2, dtype=torch.bfloat16, device="cuda") + padded[:, :HIDDEN] = values + view = padded[:, :HIDDEN] + assert torch.equal(view, values) and not view.is_contiguous() + + raw = torch.ops.trtllm.reducescatter(view, None, GROUP) + torch.cuda.synchronize() + # It reduced `rows * WORLD * HIDDEN` packed elements from the start of + # `padded`, which is the first half of the rows interleaved with the zero + # half — so the result is neither the right sum nor the right rows. + assert raw.shape == (rows, HIDDEN), raw.shape + ref = _ref(sizes, RANK, seed) + assert (raw != ref).float().mean().item() > 0.4 + COMM.Barrier() + + try: + reducescatter(view, None, GROUP) + except AssertionError: + pass + else: + raise AssertionError("wrapper accepted a non-contiguous input") + COMM.Barrier() + + +def test_wrapper_guards_a_zero_dim_input() -> None: + """A 0-d input segfaults inside the op, so the wrapper stops it first. + + The op is deliberately not called here: the crash is in + `ReducescatterOp::run_list` and kills every rank in the job. + """ + try: + reducescatter(torch.tensor(1.0, device="cuda"), None, GROUP) + except AssertionError: + pass + else: + raise AssertionError("wrapper accepted a 0-d input") + COMM.Barrier() + + +def test_wrapper_guards_a_sizes_list_of_the_wrong_length() -> None: + """Neither wrong length is survivable, so the wrapper stops both. + + The raw op is deliberately not called with either. A list one entry short + raises `IndexError: vector::_M_range_check: __n (which is WORLD) >= + this->size()` on the highest rank while the others sit in the collective, + which cannot be recovered from inside the job. A list one entry long makes + the op reduce `sum(sizes)` rows out of an input that holds fewer — an + out-of-bounds read; it returned a correctly shaped tensor of garbage (NaN) + when this test was brought up, and reading past a live allocation is not + something to repeat on every run. + """ + rows, seed = 4, 93000 + sizes = [rows] * WORLD + long_list = [rows] * (WORLD + 1) + for bad in ([rows] * (WORLD - 1), long_list): + try: + reducescatter(_input(RANK, sizes, seed), bad, GROUP) + except AssertionError: + pass + else: + raise AssertionError(f"wrapper accepted sizes of length {len(bad)}") + COMM.Barrier() + + +def test_wrapper_guards_a_split_that_does_not_cover_the_input() -> None: + """Rows the split does not reach are silently dropped, not flagged. + + Two ways to get there, both exercised on the raw op because both are + survivable: a row count that `len(group)` does not divide, and a `sizes` + vector summing to less than the input holds. A vector summing to *more* + than the input holds is not exercised — the op reads past the end of the + buffer — and the wrapper rejects it on the same assert. + """ + seed = 94000 + odd = WORLD * 2 + 1 + x = _input(RANK, [1] * odd, seed) # `odd` blocks of one row each + assert x.shape == (odd, HIDDEN) + raw = torch.ops.trtllm.reducescatter(x, None, GROUP) + torch.cuda.synchronize() + assert raw.shape == (odd // WORLD, HIDDEN), raw.shape # the last row vanished + COMM.Barrier() + + short_split = [2] * WORLD + y = _input(RANK, [4] * WORLD, seed) + raw2 = torch.ops.trtllm.reducescatter(y, short_split, GROUP) + torch.cuda.synchronize() + assert raw2.shape == (2, HIDDEN), raw2.shape + COMM.Barrier() + + for x_, sizes_ in ((x, None), (y, short_split), (y, [6] * WORLD)): + try: + reducescatter(x_, sizes_, GROUP) + except AssertionError: + pass + else: + raise AssertionError(f"wrapper accepted an uncovered split {sizes_}") + COMM.Barrier() + + +def test_the_group_still_works_after_the_negative_tests() -> None: + """The raises above leave the communicator usable — checked, not assumed.""" + rows, seed = 8, 95000 + sizes = [rows] * WORLD + out = reducescatter(_input(RANK, sizes, seed), None, GROUP) + _assert_bitwise(out, _ref(sizes, RANK, seed), "after the negative tests") + COMM.Barrier() + + +TESTS = ( + # Stays first: it is the only test that can observe GROUP's first-ever + # call, and every later test needs the communicator it builds. + test_cuda_graph_capture_of_a_first_call_raises, + test_uniform_reduce_scatter, + test_ragged_reduce_scatter, + test_output_is_fresh_and_input_is_untouched, + test_trailing_dims_are_preserved, + test_dtypes_reduce_arithmetically, + test_group_selects_a_rank_subset, + test_group_order_does_not_change_the_output, + test_reduction_is_deterministic_and_accumulates_in_the_input_dtype, + test_uniform_form_and_an_explicit_even_split_agree_bitwise, + test_round_trip_with_a_gather_returns_each_rank_its_own_rows, + test_cuda_graph_at_every_engine_batch_size, + test_cuda_graph_one_site_per_moe_layer_in_every_batch_size_graph, + test_cuda_graph_replays_survive_eager_calls_of_other_shapes, + test_cuda_graph_holds_the_sizes_vector_it_captured, + test_the_gate_discriminates_a_wrong_reduce_scatter, + # The call-order block. Each of these deliberately disagrees about call + # order and each restores alignment before it returns — the plain call + # every one of them ends on is what proves it. + test_calls_pair_by_position_and_a_swapped_pair_realigns, + test_a_mispaired_result_is_deterministic_rather_than_noise, + test_an_extra_call_on_one_rank_misaligns_until_the_counts_match, + test_mispaired_against_an_all_gather_of_equal_byte_count_corrupts_silently, + test_the_stream_the_call_lands_on_is_not_part_of_the_match, + test_float8_is_summed_as_raw_bytes, + test_wrapper_guards_a_non_contiguous_input, + test_wrapper_guards_a_zero_dim_input, + test_wrapper_guards_a_sizes_list_of_the_wrong_length, + test_wrapper_guards_a_split_that_does_not_cover_the_input, + test_the_group_still_works_after_the_negative_tests, + # Stays last: the raise it asserts leaves every later collective in the + # process returning garbage, so nothing can run after it. + test_unsupported_dtypes_raise_and_poison_every_later_collective, +) + + +def _run_one_rank() -> int: + """Body of one MPI rank: run every test, abort the job if any fails.""" + global COMM, RANK, WORLD, GROUP, reducescatter + from mpi4py import MPI + + from . import reducescatter as entry + + reducescatter = entry.reducescatter + COMM = MPI.COMM_WORLD + RANK = COMM.Get_rank() + WORLD = COMM.Get_size() + GROUP = list(range(WORLD)) + assert WORLD >= 2, f"a collective needs at least 2 ranks, got {WORLD}" + torch.cuda.set_device(RANK) + + for test in TESTS: + try: + test() + except BaseException: + import traceback + + print(f"[rank {RANK}] FAILED {test.__name__}", flush=True) + traceback.print_exc() + sys.stdout.flush() + sys.stderr.flush() + # Abort rather than return: a rank that leaves a collective early + # wedges every other rank in it. + COMM.Abort(1) + COMM.Barrier() + print(f"[rank {RANK}] {len(TESTS)} tests passed", flush=True) + return 0 + + +def _mark(kind: str) -> None: + """Record that this rank reached `kind`, durably, for the launcher to read.""" + path = os.path.join(os.environ["RS_WEDGE_MARKS"], f"{kind}.{RANK}") + with open(path, "w") as handle: + handle.write(f"{time.time():.3f}\n") + handle.flush() + os.fsync(handle.fileno()) + + +def _wedge_watchdog() -> None: + """End this rank once the mispaired pair has had its chance to return. + + Runs in a daemon thread, which can only make progress because the rank is + blocked in `torch.cuda.synchronize()` with the GIL released. There is no + gentler exit: the rank is waiting on a NCCL kernel that will never + complete, so no exception can be raised into it and no collective can be + cancelled. + """ + time.sleep(WEDGE_GRACE_S) + _mark("wedged") + os._exit(7) + + +def _run_wedge_rank() -> int: + """Body of one rank of the sub-job that certifies the wedge. + + Same mispairing as + test_mispaired_against_an_all_gather_of_equal_byte_count_corrupts_silently, + with one difference: the two calls carry **different** element counts + (`rows*HIDDEN` against `5*HIDDEN`). That is the case NCCL cannot serve out + of the buffers it was given, and it hangs rather than returning wrong data. + + The rank marks the file system either side of the pair so the launcher can + tell a wedge from a crash: every rank reaching `ready` and none reaching + `returned` is the wedge, and it means nothing without the first half. + """ + global COMM, RANK, WORLD, GROUP, reducescatter + from mpi4py import MPI + + from . import reducescatter as entry + + reducescatter = entry.reducescatter + COMM = MPI.COMM_WORLD + RANK = COMM.Get_rank() + WORLD = COMM.Get_size() + GROUP = list(range(WORLD)) + torch.cuda.set_device(RANK) + + # One aligned call first: the group's communicator has to already exist, so + # that what wedges below is the mispaired pair and not the bootstrap. + rows, seed = 16, 970000 + sizes = [rows] * WORLD + _assert_bitwise( + reducescatter(_input(RANK, sizes, seed), None, GROUP), + _ref(sizes, RANK, seed), + "wedge sub-job warm-up", + ) + COMM.Barrier() + + rs_in = _input(RANK, sizes, seed + 1000) + gather_in = _block(RANK, 0, 5, seed + 2000) + assert gather_in.numel() != rs_in.shape[0] // WORLD * HIDDEN, "counts must differ" + _mark("ready") + COMM.Barrier() + # Started after the barrier, so the grace covers the pair and nothing else. + threading.Thread(target=_wedge_watchdog, daemon=True).start() + + if RANK == 0: + torch.ops.trtllm.allgather(gather_in, None, GROUP) + reducescatter(rs_in, None, GROUP) + else: + reducescatter(rs_in, None, GROUP) + torch.ops.trtllm.allgather(gather_in, None, GROUP) + torch.cuda.synchronize() + _mark("returned") + print(f"[rank {RANK}] the mispaired pair RETURNED", flush=True) + return 0 + + +def _mpirun(world_size: int, flag: str, env: Optional[Dict[str, str]] = None) -> Any: + """Start one mpirun job in its own process group, so it can be killed.""" + command = [ + "mpirun", + "-n", + str(world_size), + sys.executable, + "-m", + "tensorrt_llm._torch.staircase.catalog.comm.reducescatter_test", + flag, + ] + print(f"[launcher] {' '.join(command)}", flush=True) + return subprocess.Popen(command, start_new_session=True, env=env) + + +def _certify_the_unequal_byte_count_wedge(world_size: int) -> None: + """Second job: the one call-order divergence that hangs instead of lying. + + It cannot be a `TESTS` entry, because the job that runs it never reports. + So it runs on its own, and the evidence is the marks its ranks leave: all + of them entered the mispaired pair, none came out of it within + WEDGE_GRACE_S, and the job ended by its own watchdog rather than by + completing. A crash before the pair fails this too — the `ready` count is + what separates the two. + """ + marks = tempfile.mkdtemp(prefix="reducescatter_wedge_") + started = time.time() + process = _mpirun(world_size, _WEDGE_FLAG, env=dict(os.environ, RS_WEDGE_MARKS=marks)) + try: + code: Optional[int] = process.wait(timeout=WEDGE_CAP_S) + ended = "its own watchdog" + except subprocess.TimeoutExpired: + os.killpg(process.pid, signal.SIGKILL) + process.wait() + code, ended = None, f"the launcher, after {WEDGE_CAP_S}s" + + seen = os.listdir(marks) + ready = [m for m in seen if m.startswith("ready.")] + returned = [m for m in seen if m.startswith("returned.")] + wedged = [m for m in seen if m.startswith("wedged.")] + for name in seen: + os.remove(os.path.join(marks, name)) + os.rmdir(marks) + print( + f"[launcher] wedge sub-job: exit={code} ready={len(ready)} " + f"returned={len(returned)} wedged={len(wedged)} " + f"ended by {ended} in {time.time() - started:.0f}s", + flush=True, + ) + + assert len(ready) == world_size, ( + f"only {len(ready)}/{world_size} ranks reached the mispaired pair — the " + "sub-job failed before it could be mispaired, so nothing was certified" + ) + assert not returned, ( + "the mispaired pair returned at unequal byte counts; it used to wedge, " + "so the contract's call-order table has to be re-measured" + ) + assert len(wedged) == world_size, ( + f"{len(wedged)}/{world_size} ranks were still inside the pair after " + f"{WEDGE_GRACE_S}s; expected every one of them" + ) + + +def _spawn_ranks() -> None: + """Re-exec this file under mpirun, one rank per claimed device.""" + visible = os.environ.get("CUDA_VISIBLE_DEVICES") + assert visible, ( + "set CUDA_VISIBLE_DEVICES to the devices this run owns, " + "e.g. export CUDA_VISIBLE_DEVICES=0,1,2,3" + ) + world_size = len([d for d in visible.split(",") if d.strip()]) + assert world_size >= 2, ( + f"CUDA_VISIBLE_DEVICES names {world_size} device(s); a collective test needs at least 2" + ) + # Own process group so the deadline can kill wedged grandchildren too. + process = _mpirun(world_size, _WORKER_FLAG) + try: + code = process.wait(timeout=DEADLINE_S) + except subprocess.TimeoutExpired: + os.killpg(process.pid, signal.SIGKILL) + process.wait() + raise AssertionError( + f"the {world_size}-rank run did not finish in {DEADLINE_S}s (wedged)" + ) from None + assert code == 0, f"the {world_size}-rank run exited {code}" + _certify_the_unequal_byte_count_wedge(world_size) + + +if __name__ == "__main__": + if _WORKER_FLAG in sys.argv: + sys.exit(_run_one_rank()) + if _WEDGE_FLAG in sys.argv: + sys.exit(_run_wedge_rank()) + _spawn_ranks() + print("OK") diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/test_allgather_op_matrix.py b/tensorrt_llm/_torch/staircase/catalog/comm/test_allgather_op_matrix.py new file mode 100644 index 000000000000..a1af97045f6d --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/comm/test_allgather_op_matrix.py @@ -0,0 +1,24 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Collected entry point for the allgather op's certification matrix. + +Distinct from ``tests/unittest/_torch/multi_gpu/test_allreduce.py`` and not a +duplicate of it: that file covers the *fusion patterns* through the +``AllReduce`` module, while this covers the op itself cell by cell -- strategy +x operation, which strategies are bitwise identical, and which combinations +are silently wrong rather than loud. The names are kept apart so review does +not read one as a copy of the other. + +The matrix itself lives in ``allgather_test.py``, which is its own 4-rank +launcher; see ``_rank_job`` for why that is left intact. +""" + +import torch + +from . import _rank_job + +assert torch.cuda.is_available(), "allgather requires CUDA devices" + + +def test_allgather_op_matrix() -> None: + _rank_job.run("allgather") diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/test_reducescatter_op_matrix.py b/tensorrt_llm/_torch/staircase/catalog/comm/test_reducescatter_op_matrix.py new file mode 100644 index 000000000000..b4db607ac5b7 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/comm/test_reducescatter_op_matrix.py @@ -0,0 +1,24 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Collected entry point for the reducescatter op's certification matrix. + +Distinct from ``tests/unittest/_torch/multi_gpu/test_allreduce.py`` and not a +duplicate of it: that file covers the *fusion patterns* through the +``AllReduce`` module, while this covers the op itself cell by cell. The names +are kept apart so review does not read one as a copy of the other. + +The matrix lives in ``reducescatter_test.py``, which is its own launcher and +runs two jobs: the ordered test sequence, then a separately capped job that +certifies the one call-order divergence that wedges instead of lying (it +cannot be a normal test, because the job that runs it never reports). +""" + +import torch + +from . import _rank_job + +assert torch.cuda.is_available(), "reducescatter requires CUDA devices" + + +def test_reducescatter_op_matrix() -> None: + _rank_job.run("reducescatter") diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/__init__.py b/tensorrt_llm/_torch/staircase/catalog/gemm/__init__.py new file mode 100644 index 000000000000..7effad5b0a94 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/gemm/__init__.py @@ -0,0 +1,3 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GEMM entries.""" diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out.md b/tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out.md new file mode 100644 index 000000000000..e7a5d35ec814 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out.md @@ -0,0 +1,101 @@ +--- +receipts: + sm_100: {status: passed, trtllm: 1.3.0rc21} + sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 5} +--- + +# bmm_out + +**Wraps** `torch.ops.trtllm.bmm_out` (one call). + +## Semantics + +Batched matrix multiply written into a caller-provided output buffer: + +``` +out[i, m, n] = sum_k a[i, m, k] * b[i, k, n] for every batch index i +``` + +Accumulation is fp32 for bf16/fp16 inputs: kernel and reference sum the same +exactly-representable products in a different order, so bf16 results agree to +within one rounding step but are **not** bit-identical. The per-element +mismatch rate against `torch.bmm(a.float(), b.float())` tracks K (~7e-6 at +K=16, ~5e-5 at K=64, ~3e-4 at K=576), so whether any one call happens to come +back bit-exact is a matter of how many elements it produces — and some shapes +sit in cuBLAS kernel-selection windows that are bit-identical over millions of +elements. Compare with this test's `atol=1e-3`, which is load-bearing: at +torch's default bf16 `atol` of 1e-5 a prefill-shaped +`(B=8, M=2048, K=512, N=128)` output fails 4 of 10 seeds. There is no +broadcasting: all three tensors are strictly 3D with equal batch sizes. + +Fusion boundary: the single call computes the batched gemm and nothing else — +no bias, no scaling, no activation, no quantization. The caller owns the +allocation of `out`, any transposition of `b` (pass a transpose *view*), and +any packing of head/group dims into the batch dim. Sibling ops exist for +quantized or arch-specialized batched gemms (`fp8_block_scaling_bmm_out`, +`fp4_bmm`, `cute_dsl_bf16_bmm_blackwell`). + +The op's reason to exist over plain `torch.bmm(..., out=)`: it is registered +as an opaque custom op so a torch.compile graph does not break when `out` is +a non-contiguous view. Strided views are first-class for all three arguments +(verified: transposed `out`, transposed `b`, row-strided `a`). + +## Signature + +```python +def bmm_out(a: torch.Tensor, b: torch.Tensor, out: torch.Tensor) -> None +``` + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `a` | `[B, M, K]` (3D only) | bf16 / fp16 / fp32 | any strides (strided views verified) | CUDA | +| `b` | `[B, K, N]` (3D only) | same as `a` | any strides (e.g. `weight.transpose(1, 2)` view) | CUDA | +| `out` | `[B, M, N]` (3D only) | same as `a` | any strides; written in place, never reallocated when the shape is right | CUDA | +| returns | — | — | `None`; the result is the mutation of `out` | — | + +## Metadata consumed + +None. Stateless. + +## Preconditions + +- All three tensors are 3D, on the same CUDA device. 2D inputs raise + (`batch1 must be a 3D tensor`); there is no implicit batch broadcasting. +- Batch sizes and the contraction dim `K` must match between `a` and `b` + (mismatch raises). +- `out.shape == (B, M, N)` exactly. **The op does not validate this**: a + wrong-shaped `out` is silently resized (with only a deprecation + warning). The hazard is real but is **shape/stride corruption, not a + lost write**: measured 2026-07-28, an aliased view with room to grow + keeps its `data_ptr` and the result *does* land in the caller's buffer, + now under a silently rewritten shape; a genuine reallocation moves the + shared storage, so sibling views follow it rather than being detached. + The wrapper asserts the shape. +- One dtype across `a`, `b`, `out`. **Mixed input dtypes always raise** — + there is no promotion path and no silent hazard here. The meta check + demands `out` in `b.dtype` while the kernel demands `a.dtype`, so when + `a.dtype != b.dtype` the two can never both be satisfied: all six + combinations over {bf16, fp16, fp32} raise, including the bf16-`a` / + fp32-`b` / fp32-`out` case an earlier revision of this bullet described + as working (measured 2026-07-28). The wrapper's single-dtype assert is + therefore redundant rather than load-bearing. +- `out.dtype` must equal the input dtype; a mismatch raises + (`Expected out tensor to have dtype ...`). +- Dtypes verified: bf16, fp16, fp32. float8_e4m3fn raises + (`"baddbmm_cuda" not implemented for 'Float8_e4m3fn'`). fp64 and integer + dtypes are untested (unknown). +- Arbitrary strides are supported for all three tensors, including + zero-copy transpose views of `b` and `out` (verified correct). + +## Notes + +- The op body is exactly one `torch.bmm(a, b, out=out)`; the launched kernel + is torch's cuBLAS strided-batched gemm, so numerical behavior follows the + process-wide torch matmul settings (tf32 flags for fp32, bf16 + reduced-precision-reduction flag). Receipts here were taken under torch + defaults. +- Not arch-gated; TRT-LLM uses it as the bf16 batched-gemm path in MLA + weight-absorption and output projections, with the batch dim carrying + head groups. +- Behavior above was established empirically under trtllm 1.3.0rc21 / + torch 2.11.0 on sm_100. diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out.py b/tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out.py new file mode 100644 index 000000000000..cf0979296def --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out.py @@ -0,0 +1,21 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Batched matmul into a caller-provided output buffer via the trtllm bmm_out op.""" + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + + +def bmm_out(a: torch.Tensor, b: torch.Tensor, out: torch.Tensor) -> None: + """Compute `out[i] = a[i] @ b[i]` for every batch index, in one bmm_out call.""" + # A wrong-shaped `out` is silently resized and re-allocated by the op, + # detaching it from any buffer the caller aliased it with. + assert a.dim() == 3 and b.dim() == 3 and out.dim() == 3, "a, b, out must be 3D" + assert out.shape == (a.shape[0], a.shape[1], b.shape[2]), ( + "out must be [B, M, N] matching a [B, M, K] and b [B, K, N]" + ) + # Mixed input dtypes are type-promoted instead of rejected, changing the + # required out dtype; the catalog exposes only the single-dtype form. + assert a.dtype == b.dtype == out.dtype, "a, b, out must share one dtype" + torch.ops.trtllm.bmm_out(a, b, out) diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out_test.py b/tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out_test.py new file mode 100644 index 000000000000..af91668f9622 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out_test.py @@ -0,0 +1,75 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the bmm_out catalog entry.""" + +import torch + +from .bmm_out import bmm_out + +assert torch.cuda.is_available(), "bmm_out requires a CUDA device" + + +def _check(a: torch.Tensor, b: torch.Tensor, out: torch.Tensor) -> None: + ref = torch.bmm(a.float(), b.float()).to(out.dtype) + bmm_out(a, b, out) + # rtol: torch.testing defaults per output dtype. atol: kernel and reference + # both accumulate in fp32 but in different summation orders; for K <= 1024 + # unit-variance inputs the order-dependent absolute noise is up to + # ~K * 2^-24 ~= 6e-5, which dominates on near-zero outputs produced by + # cancellation, so atol=1e-3 instead of the ~1e-5 defaults. + rtol = {torch.bfloat16: 1.6e-2, torch.float16: 1e-3, torch.float32: 1.3e-6} + torch.testing.assert_close(out, ref, rtol=rtol[out.dtype], atol=1e-3) + + +def _make( + batch: int, m: int, k: int, n: int, dtype: torch.dtype +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + a = torch.randn(batch, m, k, device="cuda").to(dtype) + b = torch.randn(batch, k, n, device="cuda").to(dtype) + out = torch.empty(batch, m, n, device="cuda", dtype=dtype) + return a, b, out + + +def test_bf16() -> None: + torch.manual_seed(0) + # decode-like (few tokens per batch entry) and prefill-like (many) shapes + for batch, m, k, n in [(64, 1, 512, 128), (16, 8, 576, 512), (8, 2048, 512, 128)]: + _check(*_make(batch, m, k, n, torch.bfloat16)) + + +def test_bf16_noncontiguous_views() -> None: + # The op exists so `out` may be a strided view (e.g. a transposed slice of + # a [tokens, groups, rank] buffer, as MLA uses it); `a` and `b` may be + # strided views too. Exercise all three. + torch.manual_seed(1) + batch, m, k, n = 16, 32, 256, 64 + a_wide = torch.randn(batch, m, 2 * k, device="cuda").to(torch.bfloat16) + a = a_wide[:, :, :k] # row-strided view + b = ( + torch.randn(batch, n, k, device="cuda").to(torch.bfloat16).transpose(1, 2) + ) # transposed view + out_buf = torch.empty(m, batch, n, device="cuda", dtype=torch.bfloat16) + out = out_buf.transpose(0, 1) # non-contiguous out + _check(a, b, out) + # writes landed in the aliased buffer, not a reallocation + ref = torch.bmm(a.float(), b.float()).to(torch.bfloat16) + torch.testing.assert_close(out_buf.transpose(0, 1), ref, rtol=1.6e-2, atol=1e-3) + + +def test_bf16_unaligned_shapes() -> None: + # dims not multiples of typical tile/vector widths + torch.manual_seed(2) + for batch, m, k, n in [(3, 5, 100, 60), (7, 13, 333, 129)]: + _check(*_make(batch, m, k, n, torch.bfloat16)) + + +def test_fp16() -> None: + torch.manual_seed(3) + for batch, m, k, n in [(64, 1, 512, 128), (8, 1024, 512, 256)]: + _check(*_make(batch, m, k, n, torch.float16)) + + +def test_fp32() -> None: + torch.manual_seed(4) + for batch, m, k, n in [(32, 2, 256, 128), (4, 1024, 512, 256)]: + _check(*_make(batch, m, k, n, torch.float32)) diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm.md b/tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm.md new file mode 100644 index 000000000000..de74e6a9b5a1 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm.md @@ -0,0 +1,130 @@ +--- +receipts: + sm_100: {status: passed, trtllm: 1.3.0rc21} + sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 7} +--- + +# cublas_mm + +**Wraps** `torch.ops.trtllm.cublas_mm` (one call). + +## Semantics + +General matrix multiply with an optional fused bias epilogue, executed as a +single cuBLASLt matmul: + +``` +out[m, n] = sum_k mat_a[m, k] * mat_b[k, n] (+ bias[n]) +``` + +Accumulation is fp32 (verified: bf16/fp16 results match an fp32-accumulated +reference exactly at K=256). The result is cast to the output dtype. + +Fusion boundary: the single call computes the gemm and, when `bias` is given, +adds it per output column inside the epilogue. Nothing else happens inside: +no activation, no quantization, no input/output scaling (fp8 inputs are +multiplied with alpha=1; a sibling op `cublas_scaled_mm` exists for the +scaled variant). The caller owns weight transposition metadata (pass a +transpose *view*, see Preconditions) and any flattening of leading dims. + +## Signature + +```python +def cublas_mm( + mat_a: torch.Tensor, + mat_b: torch.Tensor, + bias: torch.Tensor | None = None, + out_dtype: torch.dtype | None = None, + output_buffer_kind: int = 0, + group: list[int] | None = None, +) -> torch.Tensor +``` + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `mat_a` | `[M, K]` (2D only) | bf16 / fp16 / fp32 / float8_e4m3fn | dense row-major (contiguous) | CUDA | +| `mat_b` | `[K, N]` (2D only) | same as `mat_a` | dense column-major: strides `(1, K)`, e.g. `weight.t()` of a contiguous `[N, K]` weight | CUDA | +| `bias` | `[N]` or `None` | must equal the output dtype | contiguous | CUDA | +| `out_dtype` | — | `None` → same as `mat_a`; conversions verified: bf16→fp32, fp8e4m3→bf16 | — | — | +| `output_buffer_kind` | scalar | Python int: 0=DEFAULT, 1=USERBUFFERS, 2=NCCL_WINDOW | — | — | +| `group` | list of ranks or `None` | Python ints | only meaningful with NCCL_WINDOW | — | +| returns | `[M, N]` | `out_dtype` or `mat_a.dtype` | newly allocated, contiguous | CUDA | + +## Metadata consumed + +None. Stateless for the default path. `output_buffer_kind` selects where the +output tensor is allocated: 0 is a plain CUDA allocation; 1 (USERBUFFERS) and +2 (NCCL_WINDOW, together with `group` = the TP rank list) draw from +communication buffer pools so a downstream allreduce can consume the output +in place — those pools must have been initialized by the runtime beforehand. +Single-GPU callers pass `0, None` (2 with `group=None` silently falls back to +a plain allocation). The enum is importable as +`tensorrt_llm.bindings.internal.thop.BufferKind`. + +## Preconditions + +- `mat_a` and `mat_b` are 2D, on the same CUDA device, same dtype. 3D+ + inputs raise; flatten leading dims to `[M, K]` and unflatten after. +- `mat_a` is dense row-major. **The kernel ignores `mat_a` row strides**: a + row-strided view (e.g. a column slice of a wider buffer) is accepted and + produces silently wrong results (verified). The wrapper asserts + contiguity. +- `mat_b` has strides exactly `(1, K)` — a dense column-major buffer, i.e. + the `.t()` view of a contiguous `[N, K]` weight. The op checks + `stride(0) == 1` and raises otherwise, but does **not** check + `stride(1) == K`: the transpose view of a row-strided weight is accepted + and produces silently wrong results (verified). The wrapper asserts both. +- Dtypes: bf16, fp16, fp32 work with `out_dtype=None`; float8_e4m3fn + requires an explicit `out_dtype` (`None` raises a cuBLAS error). + `mat_a.dtype != mat_b.dtype` raises `CUBLAS_STATUS_NOT_SUPPORTED`. +- `bias`, when given: + - shape exactly `[N]`, contiguous, on the same device. **The op does not + validate the bias shape** (a wrong-length bias was accepted); the + wrapper asserts it. + - dtype must equal the *output* dtype (bf16 bias for bf16 out, fp32 bias + for bf16→fp32, bf16 bias for fp8→bf16). A mismatched dtype (e.g. fp32 + bias with bf16 output) is accepted and produces silently wrong results + (verified). + - not supported with fp32 inputs: the bias is accepted and **silently + ignored** (verified: output equals the unbiased product). The wrapper + asserts against this combination. +- `out_dtype` conversions other than the verified ones are not guaranteed: + bf16→fp16 raises a cuBLASLt runtime error on this machine; combinations + not listed above are untested (unknown). + +## Notes + +- The fake (meta) registration of this op accepts N-D `mat_a`, but the real + kernel requires 2D — shape inference under compile can diverge from + eager behavior for N-D inputs. +- Error traces name `cpp/tensorrt_llm/thop/cublasScaledMM.cpp` as the + implementation (shared with the scaled variant); the C++ source is not + shipped in the wheel, so all behavior above was established empirically + on sm_100. +- TRT-LLM routes bf16 linears to this op on sm>=100 because cuBLASLt picks + single-pass cluster-mode kernels for small-M (decode) gemms there; the + op itself is not gated on arch. +- All silently-wrong-result behaviors listed under Preconditions were + observed under trtllm 1.3.0rc21 on sm_100. + + +## The fp32 cell needs its reference pinned, not its tolerance widened + +torch 2.12 defaults `torch.backends.cuda.matmul.fp32_precision` to `tf32` (and +`allow_tf32` to True) on this hardware, so the natural reference +`mat_a.float() @ mat_b.float()` is a **TF32** product. Against a correct +kernel that reads as a ~0.035 absolute error at K=1024 -- about 30x the op's +own -- and it is the *reference* that is wrong. + +Measured on sm_103 against a float64 product on the same inputs: + +| | max abs error vs float64 | +|---|---| +| `cublas_mm` | 1.9e-5 | +| torch, `allow_tf32=False` | 1.9e-5 (**bit-identical to the op**) | +| torch, `allow_tf32=True` | 3.5e-2 | + +So the op is exact to fp32 and this entry's `_ref_mm` now forces TF32 off for +the duration of the reference product. The same trap applies to any entry +whose reference multiplies in fp32; the bf16 and fp16 cells were merely +tolerant enough to hide it. diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm.py b/tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm.py new file mode 100644 index 000000000000..9a0af503e9c8 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm.py @@ -0,0 +1,35 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""General matmul with optional fused bias via the trtllm cuBLASLt gemm op.""" + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + + +def cublas_mm( + mat_a: torch.Tensor, + mat_b: torch.Tensor, + bias: torch.Tensor | None = None, + out_dtype: torch.dtype | None = None, + output_buffer_kind: int = 0, + group: list[int] | None = None, +) -> torch.Tensor: + """Return `mat_a @ mat_b (+ bias)`, fp32-accumulated, in one cublas_mm call.""" + # The kernel reads mat_a as dense row-major and mat_b as dense + # column-major; other layouts produce silently wrong results. + assert mat_a.is_contiguous(), "mat_a must be dense row-major [M, K]" + assert mat_b.stride(0) == 1 and mat_b.stride(1) == mat_b.shape[0], ( + "mat_b must be dense column-major [K, N] (e.g. weight.t())" + ) + if bias is not None: + out_dt = out_dtype if out_dtype is not None else mat_a.dtype + # A bias in the wrong dtype or shape is accepted by the op and + # produces silently wrong results; with fp32 inputs the bias is + # accepted but silently ignored. + assert mat_a.dtype != torch.float32, "bias is silently ignored for fp32 inputs" + assert bias.dtype == out_dt, "bias dtype must equal the output dtype" + assert bias.shape == (mat_b.shape[1],) and bias.is_contiguous(), ( + "bias must be a contiguous [N] tensor" + ) + return torch.ops.trtllm.cublas_mm(mat_a, mat_b, bias, out_dtype, output_buffer_kind, group) diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm_test.py b/tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm_test.py new file mode 100644 index 000000000000..eb1c13784950 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm_test.py @@ -0,0 +1,139 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the cublas_mm catalog entry.""" + +import contextlib + +import torch + +from .cublas_mm import cublas_mm + +assert torch.cuda.is_available(), "cublas_mm requires a CUDA device" + + +@contextlib.contextmanager +def _true_fp32_matmul(): + """Make torch's fp32 matmul actually fp32 for the duration. + + torch 2.12 defaults ``matmul.fp32_precision`` to ``tf32`` and + ``allow_tf32`` to True on this hardware, so a plain ``a.float() @ + b.float()`` is a *TF32* product -- ~1e-3 relative, which is 30x the error + of the op being tested. Left alone the reference is the inaccurate side of + the comparison and the entry fails against a correct kernel. Measured + here: with TF32 off the op is bit-identical to torch and both sit 1.9e-5 + from a float64 product; with TF32 on the reference alone moves by 0.035. + """ + prev_allow = torch.backends.cuda.matmul.allow_tf32 + prev_prec = getattr(torch.backends.cuda.matmul, "fp32_precision", None) + torch.backends.cuda.matmul.allow_tf32 = False + if prev_prec is not None: + torch.backends.cuda.matmul.fp32_precision = "ieee" + try: + yield + finally: + torch.backends.cuda.matmul.allow_tf32 = prev_allow + if prev_prec is not None: + torch.backends.cuda.matmul.fp32_precision = prev_prec + + +def _ref_mm( + mat_a: torch.Tensor, + mat_b: torch.Tensor, + bias: torch.Tensor | None, + out_dtype: torch.dtype | None, +) -> torch.Tensor: + """fp32-accumulated reference: mat_a @ mat_b (+ bias), cast to out dtype.""" + with _true_fp32_matmul(): + ref = mat_a.float() @ mat_b.float() + if bias is not None: + ref = ref + bias.float() + return ref.to(out_dtype if out_dtype is not None else mat_a.dtype) + + +def _check( + mat_a: torch.Tensor, + mat_b: torch.Tensor, + bias: torch.Tensor | None = None, + out_dtype: torch.dtype | None = None, +) -> None: + out = cublas_mm(mat_a, mat_b, bias, out_dtype) + ref = _ref_mm(mat_a, mat_b, bias, out_dtype) + expected_dtype = out_dtype if out_dtype is not None else mat_a.dtype + assert out.shape == (mat_a.shape[0], mat_b.shape[1]) + assert out.dtype == expected_dtype + # rtol: torch.testing defaults per output dtype. atol: kernel and reference + # both accumulate in fp32 but in different summation orders; for K <= 4096 + # unit-variance inputs the order-dependent absolute noise is up to + # ~K * 2^-24 ~= 2.4e-4, which dominates on near-zero outputs produced by + # cancellation, so atol=1e-3 instead of the ~1e-5 defaults. + rtol = {torch.bfloat16: 1.6e-2, torch.float16: 1e-3, torch.float32: 1.3e-6} + torch.testing.assert_close(out, ref, rtol=rtol[expected_dtype], atol=1e-3) + + +def _make(m: int, k: int, n: int, dtype: torch.dtype) -> tuple[torch.Tensor, ...]: + """Build mat_a [M,K] row-major, mat_b [K,N] column-major, bias [N].""" + mat_a = torch.randn(m, k, device="cuda").to(dtype) + weight = torch.randn(n, k, device="cuda").to(dtype) # linear weight [N, K] + bias = torch.randn(n, device="cuda").to(dtype) + return mat_a, weight.t(), bias + + +def test_bf16_no_bias() -> None: + torch.manual_seed(0) + # decode-like (few tokens) and prefill-like (many tokens) shapes + for m, k, n in [(1, 4096, 4096), (8, 4096, 11008), (2048, 4096, 4096)]: + mat_a, mat_b, _ = _make(m, k, n, torch.bfloat16) + _check(mat_a, mat_b) + + +def test_bf16_bias() -> None: + torch.manual_seed(1) + for m, k, n in [(1, 4096, 4096), (512, 2048, 6144)]: + mat_a, mat_b, bias = _make(m, k, n, torch.bfloat16) + _check(mat_a, mat_b, bias) + + +def test_bf16_out_fp32() -> None: + # bias must match the output dtype (fp32 here), not the input dtype + torch.manual_seed(2) + mat_a, mat_b, _ = _make(16, 1024, 2048, torch.bfloat16) + bias_fp32 = torch.randn(2048, device="cuda", dtype=torch.float32) + _check(mat_a, mat_b, out_dtype=torch.float32) + _check(mat_a, mat_b, bias_fp32, out_dtype=torch.float32) + + +def test_bf16_unaligned_shapes() -> None: + # dims not multiples of typical tile/vector widths + torch.manual_seed(3) + for m, k, n in [(5, 100, 60), (7, 333, 129)]: + mat_a, mat_b, bias = _make(m, k, n, torch.bfloat16) + _check(mat_a, mat_b) + _check(mat_a, mat_b, bias) + + +def test_fp16() -> None: + torch.manual_seed(4) + for m, k, n in [(1, 4096, 4096), (1024, 2048, 2048)]: + mat_a, mat_b, bias = _make(m, k, n, torch.float16) + _check(mat_a, mat_b) + _check(mat_a, mat_b, bias) + + +def test_fp32() -> None: + # no bias: the op silently ignores bias when inputs are fp32 + # (contract precondition; guarded by an assert in the wrapper) + torch.manual_seed(5) + for m, k, n in [(2, 1024, 1024), (256, 2048, 1024)]: + mat_a, mat_b, _ = _make(m, k, n, torch.float32) + _check(mat_a, mat_b) + + +def test_fp8_e4m3_to_bf16() -> None: + # fp8 inputs require an explicit out_dtype; no scales are applied (alpha=1) + torch.manual_seed(6) + for m, k, n in [(1, 1024, 1024), (128, 1024, 512)]: + mat_a = (torch.randn(m, k, device="cuda") * 0.1).to(torch.float8_e4m3fn) + weight = (torch.randn(n, k, device="cuda") * 0.1).to(torch.float8_e4m3fn) + bias = torch.randn(n, device="cuda", dtype=torch.bfloat16) + _check(mat_a, weight.t(), out_dtype=torch.bfloat16) + _check(mat_a, weight.t(), bias, out_dtype=torch.bfloat16) diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm.md b/tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm.md new file mode 100644 index 000000000000..b326f8d640f4 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm.md @@ -0,0 +1,306 @@ +--- +receipts: + sm_100: {status: passed, trtllm: 1.3.0rc21} + sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 19} +--- + +# nvfp4_gemm + +**Wraps** `torch.ops.trtllm.nvfp4_gemm` (one call). + +## Semantics + +Dense GEMM in nn.Linear orientation over two **block-scaled NVFP4** operands. +Let `A` be the activation and `B` the weight, both stored as packed e2m1 +nibbles plus one e4m3 scale per 16 contiguous elements along `K`: + +``` +a[m, k] = e2m1(act_fp4 nibble (m, k)) * e4m3(act_sf scale (m, k // 16)) +b[n, k] = e2m1(weight nibble (n, k)) * e4m3(weight_scale scale (n, k // 16)) + +out[m, n] = output_dtype( alpha * sum_k a[m, k] * b[n, k] + bias[n] ) +``` + +i.e. `out = alpha * (a @ b.T) + bias`. The sum is accumulated in **fp32**; +`alpha` (a device fp32 scalar) multiplies the whole accumulator; `bias` (an +optional per-column vector) is added after `alpha`; the result is cast to +`output_dtype` once, at the end. + +Fusion boundary — inside the call: both operands' **block-scale +dequantization**, the fp32 accumulation, the per-tensor `alpha`, the optional +per-column bias, and the output cast. Outside the call, still the caller's: +producing `act_fp4`/`act_sf` (an activation quantizer such as +`torch.ops.trtllm.fp4_quantize`), putting **both** scale buffers into the +swizzled layout below, folding the two per-tensor global scales into `alpha`, +any activation function, any reduction across ranks, and any narrowing of a +padded `N` back to `out_features`. + +### The two global scales live in `alpha` + +An NVFP4 tensor quantized with a per-tensor global scale `g` reconstructs as +`data * sf / g` — this op applies `data * sf` only. With `g_act` for the +activation and `g_w` for the weight, the caller owes + +``` +alpha = 1 / (g_act * g_w) +``` + +For a modelopt/HF checkpoint the two conventions meet like this: the stored +`input_scale` is `amax_act / (448*6)` = `1 / g_act` (the quantizer's +`global_scale` is its reciprocal) and the stored `weight_scale_2` is +`amax_w / (448*6)` = `1 / g_w`, so + +``` +alpha = input_scale * weight_scale_2 # both as stored on disk + = weight_scale_2 / g_act # g_act = the quantizer's global_scale +``` + +`alpha` is a load-time scalar for weight-only-static checkpoints; nothing in +this op derives it. + +### Scale-factor layout (both operands) + +Each scale buffer is **1-D** and holds one e4m3 byte per (row, 16-element +block) in the **128x4 swizzled** order: the order the trtllm NVFP4 activation +quantizer emits with `is_sf_swizzled_layout=True` +(`torch.ops.trtllm.fp4_quantize`), and the order +`torch.ops.trtllm.block_scale_interleave` produces from a row-major +`[rows, cols]` tensor — the latter verified byte-identical to the formula below +on this machine. With `cols = K / 16`, the byte for `(row r, block c)` sits at +flat offset + +``` +(c % 4) ++ (c // 4) * 512 # 4 cols x 128 rows per column group ++ (r % 32) * 16 ++ ((r % 128) // 32) * 4 ++ (r // 128) * 128 * pad_up(cols, 4) +``` + +and the buffer's length is `pad_up(rows, 128) * pad_up(cols, 4)` bytes — +`rows = M` for the activation, `rows = N` for the weight. Bytes at offsets no +real `(r, c)` addresses (row padding up to a multiple of 128, column padding up +to a multiple of 4) are **ignored**: filling them with `0x00` or with `0x7E` +(448) gives bitwise-identical output. + +**A checkpoint that stores `weight_scale` row-major `[N, K/16]` therefore owes a +load-time relayout.** Feeding the linear buffer instead is accepted and silently +wrong — the observed error was of the same order as the output itself. Both +scale buffers are passed as `uint8`; an fp8 `float8_e4m3fn` view of the same +bytes is rejected. + +## Signature + +```python +def nvfp4_gemm( + act_fp4: torch.Tensor, + weight: torch.Tensor, + act_sf: torch.Tensor, + weight_scale: torch.Tensor, + alpha: torch.Tensor, + output_dtype: torch.dtype, + output_buffer_kind: int = 0, + allowed_backends: str = "cutlass,cublaslt,cuda_core", + group: list[int] | None = None, + bias: torch.Tensor | None = None, +) -> torch.Tensor +``` + +| Argument | Shape | Dtype | Layout / value | Device | +|---|---|---|---|---| +| `act_fp4` | `[M, K/2]` (2-D only) | uint8 | contiguous; element `2i` in the low nibble, `2i+1` in the high nibble | CUDA | +| `weight` | `[N, K/2]` (2-D only) | uint8 | contiguous **row-major nn.Linear weight**, not its `.t()` view | CUDA | +| `act_sf` | 1-D, `pad_up(M,128) * pad_up(K/16,4)` | uint8 | contiguous, 128x4 swizzled (above) | CUDA | +| `weight_scale` | 1-D, `pad_up(N,128) * pad_up(K/16,4)` | uint8 | contiguous, 128x4 swizzled (above) | CUDA | +| `alpha` | exactly 1 element (`[1]` or 0-D) | float32 | `1/(g_act * g_w)` | CUDA | +| `output_dtype` | scalar | `torch.dtype` | bfloat16, float16 or float32; the `cutedsl` backend supports bfloat16 only | — | +| `output_buffer_kind` | scalar | Python int | `0` = plain allocation (certified). `1` = userbuffers, `2` = NCCL window — for multi-rank output buffers, not certified here | — | +| `allowed_backends` | scalar | Python str | comma-separated subset of `cutlass`, `cublaslt`, `cuda_core`, `cutedsl`, `marlin`; all compute the same result | — | +| `group` | list of ranks or `None` | Python list | only meaningful with a window output buffer; `None` otherwise | — | +| `bias` | `[N]`, 1-D | **equal to `output_dtype`** | contiguous; added per output column | CUDA | +| returns | `[M, N]` | `output_dtype` | freshly allocated, contiguous | same as `act_fp4` | + +`K` is `2 * act_fp4.shape[1]`. Inputs are never written (verified bitwise), and +the returned buffer aliases nothing. + +`allowed_backends` is a *selection* set, not a composition: exactly one backend +runs per call and every one of them computes the formula above (all verified +equal to the reference on this machine). The op's own schema spells the +parameters `act_fp4, weight, act_sf, weight_scale, alpha, output_dtype, +output_buffer_kind, allowed_backends, group, bias`. + +## Metadata consumed + +The process-global AutoTuner profiling cache +(`tensorrt_llm._torch.autotuner.AutoTuner`). Nothing must be prepared: + +- Outside a tuning context, a cache miss logs a warning and runs the **fallback + tactic**: the `cutlass` backend with its default config if `cutlass` is in + `allowed_backends`, otherwise the first backend listed. +- Inside `tensorrt_llm._torch.autotuner.autotune()` the call profiles every + valid (backend, tactic) pair and caches the winner per M-bucket; one tuning + call at rows `M` fills every power-of-2 bucket `1..last_pow2(M)` for that + `(K, N)` (14 buckets at `M = 8192`, verified at four shapes). Later calls map + their `M` to its bucket. Observed winners on this machine span `cublaslt` and + `cutlass` depending on shape; across the 56 cache entries the four R1 shapes + produce, every winner carried a profiled tactic id rather than the fallback + marker. + +Correctness is certified for both cache states. At the four DeepSeek-R1-0528 +shapes below the tuned tactic's output was **bitwise identical** to the cold +fallback's at `M` in {1, 3, 8, 4097, 8192} — a measured property of this op at +those shapes, not a general guarantee that tactic choice is bit-neutral. No +workspace, runtime object or attention metadata is involved; the op is +otherwise stateless. + +## Preconditions + +Violations marked *silent* were observed on this machine to produce wrong +results without raising; the wrapper rejects each with a metadata assert. +Everything else raises. + +- `K % 32 == 0` and `N % 32 == 0` (16-byte operand lines). Both raise otherwise + — `Expected k to be divisible by 32` / `Expected n to be divisible by 32`. + `M` is unconstrained (1, 3, 7, 9, 127, 128, 129, 1000, 4097, 8192, ... all + verified); `M == 0` raises. This rule is the whole shape domain, and it was + re-measured as such: `K` up to 18432, `N` up to 36864 and `M` up to 8192 + behave exactly like the smaller shapes (see the shape-dependence note below). +- `act_fp4` and `weight` are 2-D, `uint8`, on CUDA, and **contiguous**. A 3-D + activation raises; a non-uint8 view raises; a CPU tensor raises. Contiguity is + *silent*: the `cutlass` and `cuda_core` backends raise, but `cublaslt` and + `cutedsl` ignore the stride and return wrong results, and `cublaslt` is both + in the default backend list and a frequent tuner pick. +- `weight.shape[1] == act_fp4.shape[1]` (same `K`); mismatch raises. +- `act_sf` and `weight_scale` are `uint8`, **contiguous**, and hold the + swizzled layout at the full padded size above. Non-contiguity is *silent* + under `cublaslt`/`cutedsl` (and under `cuda_core` for `weight_scale`); an + unswizzled (row-major linear) buffer of the right size is *silent* on the + default path, and no backend has any way to detect it; a buffer shorter than + the padded size reads out of bounds. Padding bytes may hold anything. +- `alpha` is a float32 CUDA tensor with **exactly one element**. A half or CPU + tensor raises. More than one element is *silent*: element 0 is used and the + rest ignored — this build has no per-token alpha. +- `bias`, when given, is 1-D with exactly `N` elements and dtype equal to + `output_dtype`. Every violation raises. +- `output_dtype` is bfloat16, float16 or float32. The `cutedsl` backend accepts + bfloat16 only and raises `ValueError` otherwise — forcing + `allowed_backends="cutedsl"` with fp16/fp32 raises, and so does an + `autotune()` pass over any list that merely contains `cutedsl` (its runner is + constructed while enumerating tactics); outside tuning such a list is fine, + because the fallback never constructs it. +- `allowed_backends` must be a non-empty comma-separated subset of the five + names; an empty or misspelled string raises `ValueError`. `marlin` is + Hopper-only and raises on sm_100. `cuda_core` requires SM >= 100 and is only + ever *selected* by the tuner for `M <= 8`; forcing it computes correctly up to + `M = 16` and raises `Failed to dispatch cudaCoreGemmLauncher` above that. +- `output_buffer_kind=1` (userbuffers) raises unless a userbuffers workspace of + sufficient size was allocated by the runtime. +- The op is deterministic: repeating a call with the same inputs and the same + tuner state gives bitwise-identical output. + +## Notes + +- Certified on sm_100 (B200) only. The kernels behind every backend are + Blackwell block-scaled MMA paths (the CUTLASS one is instantiated for + `cutlass::arch::Sm100`/`Sm103` and traps on other archs); no receipt is + claimed elsewhere. +- Receipt coverage, shapes: the four DeepSeek-V3-Lite NVFP4 dense linears in + `[out, in]` orientation — `(K, N)` = (2560, 24576), (12288, 2560), + (2560, 6144), (3072, 2560) — at `M` in {1, 2, 8, 64, 1024, 4096} (first + shape) and {1, 8, 64, 1024} (rest); the four DeepSeek-R1-0528 dense linears — + `(K, N)` = (7168, 36864), (18432, 7168), (7168, 4096), (2048, 7168) — each at + `M` in {1, 2, 8, 64, 1024, 4096, 8192}; `M` in {1, 3, 7, 9, 127, 128, 129, + 1000} at (2560, 512); `N = 160` (not a multiple of 128, so the weight scale + buffer pads rows); and a wide block-scale *range* case (scale bytes spanning + 2^-5..2^5, which is what makes the accumulation round) at `K = 12288` and at + `K = 18432`. +- Receipt coverage, everything else: bf16 / fp16 / fp32 outputs — bf16 at every + shape, fp16 at (2560, 512), fp32 at (2560, 512) and at both wide-range shapes + (`K` = 12288 and `K` = 18432); `alpha` in {1, 0.5, 0.03125, 1e-4}; bf16 and + fp32 bias at (2560, 512) only, the R1 shapes having been driven without bias + as the dense path uses none; the default `allowed_backends` + string at every shape above, plus `cutlass`, `cublaslt`, `cuda_core` and + `cutedsl` each forced alone at (2560, 512) and, at every R1 shape, `cutlass` / + `cublaslt` / `cutedsl` forced alone at `M` = 8 and `M` = 8192 with + `cuda_core` alongside them at `M` = 8; the autotuned path at (2560, 6144) and + at each R1 shape (tuned at `M` = 8192, then replayed warm at `M` in + {1, 3, 8, 4097, 8192} and compared bitwise against the same calls on a cold + cache); scale-padding insensitivity; input non-mutation and call-to-call + determinism; the swizzled-vs-linear weight-scale contrast at (2560, 512) and + at both extreme R1 shapes; 20 rejected-domain calls; and the five silent + domains the wrapper guards. +- **Nothing shape-dependent changes past the smaller shapes' maxima.** What is + shape-dependent is the tactic space, not the arithmetic: the number of + cuBLASLt heuristic algorithms varies with `(M, K, N)` (1..8 observed), the + CuTe DSL tile/cluster candidate count varies (120..240 observed), and + `cuda_core` is offered only for `M <= 8`. Measured at the four R1 shapes + against the smaller certified ones, every one of those counts landed inside + the range the smaller shapes already produce, no shape produced an empty + tactic list, the CUTLASS tactic list is 32 configs everywhere (its + enumeration takes no shape at all), and the scale-buffer size formula holds + unchanged at its widest here (`K/16 = 1152` columns, `K = 18432`). At those + shapes `cutlass`, `cublaslt`, `cutedsl` and — at `M <= 8` — `cuda_core` + returned **bitwise identical** outputs, each matching the native-torch + reference, and the warm tuner cache reproduced the cold fallback's bits + exactly. So the domain rule in Preconditions was sufficient on its own at + this scale-up: the shapes' only effect was on which tactic wins. +- Numerics. Against a native-torch reference built from the operand bytes, the + bf16 output matches `torch.testing.assert_close` at **default bf16 + tolerances** (rtol 1.6e-2, atol 1e-5) at every shape above; the observed max + relative deviation was 3.9e-3, i.e. one bf16 ulp. **That figure is against + the unrounded fp32 reference** (measured 3.891e-3 vs a 2^-8 = 3.906e-3 ulp). + The test itself compares against `ref.to(torch.bfloat16)`, so its own + assertions come out **bit-exact** — an instrumented run that records + `assert_close` deviations will therefore report zero here, and that is not a + contradiction of this number. When the block scales sit + in a narrow band the products and partial sums are exactly representable and + the **fp32 output is bit-exact** against an fp64 reference. With block scales + spanning 2^-5..2^5 the fp32 sum does round, and the deviation stays inside + the recursive-summation bound `K * 2^-24 * sum|a_i b_i|` with ~600x margin — + which is what an fp32 accumulator looks like and a bf16/fp16 one would not. + Both regimes were re-measured at `K = 18432`, the largest `K` certified: + narrow band still bit-exact in fp32, wide band inside the same bound with + 715x margin. +- Discrimination at the largest shapes was measured, not assumed: moving a + single e4m3 scale byte, or a single e2m1 nibble, by one code — the smallest + perturbation either operand admits — makes the default-tolerance comparison + fail at `(K, N, M)` = (18432, 7168, 8192) and (7168, 36864, 8192). Feeding + the weight scale row-major (unswizzled) instead of swizzled lands ~0.5x the + output's own magnitude off at those shapes. +- One call is one kernel launch on the certified default path: profiling a + single call showed exactly one device kernel for `cutlass` (bias fused into + its `LinCombPerColBias` epilogue) and for `cublaslt` (bias fused, the + `..._bias_TNT` nvjet variant). Two internal implementation details do not + change the result but are worth knowing: the `cuda_core` backend first + un-swizzles `act_sf` with its own kernel (2 launches), and `cutedsl` adds + `bias` as a post-GEMM elementwise add (2 launches with bias, 1 without). +- `cutedsl` JIT-compiles on first use (seconds) and additionally validates + `alpha.numel() == 1` itself. +- The op's Python body is registered straight through + `torch.library.Library.define/impl` (trtllm's `fast_custom_op`), so + `torch.ops.trtllm.nvfp4_gemm` is an ordinary dispatcher op with a registered + fake kernel returning `[M, N]` in `output_dtype`. There is no autograd + kernel — inference only. +- Upstream (`_torch/modules/linear.py`, `NVFP4LinearMethod`) calls this op with + exactly these operands: `fp4_quantize`'d activations in the **swizzled** + layout, the raw `[N, K/2]` weight, `module.weight_scale` — which its loader + fills with `block_scale_interleave(weight_scale)` — and + `alpha = input_scale * weight_scale_2`. It flattens 3-D activations to + `[M, K]` before the call and reshapes after, and slices `N` back to + `out_features` when the weight was padded. +- Sibling ops with the same operand vocabulary exist: + `torch.ops.trtllm.nvfp4_gemm_cutlass` and + `torch.ops.trtllm.nvfp4_gemm_cublaslt` (single-backend entry points with + their own tuning caches), `torch.ops.trtllm.cuda_core_nvfp4_gemm` and + `torch.ops.trtllm.marlin_nvfp4_gemm` and + `torch.ops.trtllm.cute_dsl_nvfp4_gemm_blackwell` (the raw per-backend + kernels; the first two take the *unswizzled* activation scales that this op's + `cuda_core`/`marlin` paths derive internally, while the CuTe DSL one takes the + swizzled buffer like this op), `torch.ops.trtllm.fp4_gemm` (the older typed + entry point, also covering + W4A8 MXFP4xMXFP8), `torch.ops.trtllm.fp4_gemm_trtllmgen` and + `torch.ops.trtllm.fp4_bmm`, and `torch.ops.trtllm.nvfp4_gemm_allreduce` + (this GEMM fused with a tensor-parallel allreduce). This entry covers the + unified `nvfp4_gemm` launch only. +- Behaviour under CUDA-graph capture was not exercised (unknown); note that + autotuning and any first-call JIT must complete before capture. diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm.py b/tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm.py new file mode 100644 index 000000000000..e5ef04ba8272 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm.py @@ -0,0 +1,51 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""NVFP4 x NVFP4 dense GEMM in nn.Linear layout via the trtllm unified nvfp4_gemm op.""" + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + + +def nvfp4_gemm( + act_fp4: torch.Tensor, + weight: torch.Tensor, + act_sf: torch.Tensor, + weight_scale: torch.Tensor, + alpha: torch.Tensor, + output_dtype: torch.dtype, + output_buffer_kind: int = 0, + allowed_backends: str = "cutlass,cublaslt,cuda_core", + group: list[int] | None = None, + bias: torch.Tensor | None = None, +) -> torch.Tensor: + """Return `alpha * (act @ weight.T) (+ bias)` over block-scaled NVFP4 operands, one op call.""" + # Every guard is pure metadata, and every guarded violation was observed on + # this machine to be read wrongly *without raising* by at least one backend + # the default `allowed_backends` string can select: the cutlass path checks + # contiguity and raises, cublaslt does not, and none of the three default + # backends rejects an oversized `alpha`. + assert act_fp4.is_contiguous() and weight.is_contiguous(), ( + "act_fp4 [M, K/2] and weight [N, K/2] must be contiguous; the cublaslt " + "backend ignores strides and returns wrong results" + ) + assert act_sf.is_contiguous() and weight_scale.is_contiguous(), ( + "act_sf and weight_scale must be contiguous; the cublaslt backend " + "ignores strides and returns wrong results" + ) + assert alpha.numel() == 1, ( + "alpha must hold exactly one element; extra elements are silently " + "ignored (this build has no per-token alpha)" + ) + return torch.ops.trtllm.nvfp4_gemm( + act_fp4, + weight, + act_sf, + weight_scale, + alpha, + output_dtype, + output_buffer_kind, + allowed_backends, + group, + bias, + ) diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm_test.py b/tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm_test.py new file mode 100644 index 000000000000..3fef8286d096 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm_test.py @@ -0,0 +1,650 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the nvfp4_gemm catalog entry (NVFP4 x NVFP4 dense GEMM).""" + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* +from tensorrt_llm._torch.autotuner import AutoTuner, autotune + +from .nvfp4_gemm import nvfp4_gemm + +assert torch.cuda.is_available(), "nvfp4_gemm requires a CUDA device" +# The reference matmul must be true fp32, never tf32. +torch.backends.cuda.matmul.allow_tf32 = False + +DEV = torch.device("cuda") +VEC = 16 # NVFP4 block size: one e4m3 scale per 16 contiguous elements along K + +# e2m1 code -> value. code = (exponent << 1) | mantissa, bit 3 is the sign. +E2M1_VALUES = torch.tensor( + [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], dtype=torch.float32, device=DEV +) +# e4m3 scale bytes: 0x30..0x40 are the positive finite values 0.5 .. 2.0, whose +# products with e2m1 data stay exactly representable in fp32 (see +# test_wide_dynamic_range for the opposite regime). +SF_LO, SF_HI = 0x30, 0x41 + + +def _pad_up(x: int, m: int) -> int: + return (x + m - 1) // m * m + + +def _swizzle(sf_2d: torch.Tensor, pad_fill: int = 0) -> torch.Tensor: + """[rows, cols] e4m3 scale bytes -> the flat 128x4-swizzled buffer, in native torch. + + Offset of the scale of (row r, block c), the layout every NVFP4 GEMM operand + scale uses: + (c % 4) + (c // 4) * 512 + (r % 32) * 16 + ((r % 128) // 32) * 4 + + (r // 128) * 128 * pad_up(cols, 4) + Buffer length is `pad_up(rows, 128) * pad_up(cols, 4)`; `pad_fill` is written + to every offset no real (r, c) addresses. + """ + rows, cols = sf_2d.shape + padded_cols = _pad_up(cols, 4) + r = torch.arange(rows, device=DEV).view(-1, 1) + c = torch.arange(cols, device=DEV).view(1, -1) + idx = ( + (c % 4) + + (c // 4) * (4 * 128) + + (r % 32) * 16 + + ((r % 128) // 32) * 4 + + (r // 128) * (128 * padded_cols) + ) + out = torch.full((_pad_up(rows, 128) * padded_cols,), pad_fill, dtype=torch.uint8, device=DEV) + out[idx.flatten()] = sf_2d.flatten() + return out + + +def _operand( + rows: int, k: int, seed: int, sf_lo: int = SF_LO, sf_hi: int = SF_HI +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """A random NVFP4 operand: (packed data, linear scales, swizzled scales, fp32 values). + + The operand is built byte-first — random e2m1 codes and random positive + finite e4m3 scale bytes — and the exact real-valued matrix it encodes is + derived from those bytes in native torch. Nothing here depends on how a + quantizer would have produced them. + """ + g = torch.Generator(device="cuda").manual_seed(seed) + codes = torch.randint(0, 16, (rows, k), generator=g, device=DEV, dtype=torch.uint8) + codes = torch.where(codes == 8, torch.zeros_like(codes), codes) # drop -0.0 + data = codes[:, 0::2] | (codes[:, 1::2] << 4) + sf = torch.randint(sf_lo, sf_hi, (rows, k // VEC), generator=g, device=DEV, dtype=torch.uint8) + values = E2M1_VALUES[(codes & 7).long()] + values = torch.where((codes & 8).bool(), -values, values) + values *= sf.view(torch.float8_e4m3fn).float().repeat_interleave(VEC, dim=-1) + return data, sf, _swizzle(sf), values + + +def _reference( + a_values: torch.Tensor, + b_values: torch.Tensor, + alpha: float, + bias: torch.Tensor | None = None, +) -> torch.Tensor: + """alpha * (A @ B.T) (+ bias), fp32 — the kernel accumulates in fp32.""" + out = alpha * (a_values @ b_values.T) + if bias is not None: + out = out + bias.float() + return out + + +def _alpha(value: float) -> torch.Tensor: + return torch.tensor([value], dtype=torch.float32, device=DEV) + + +ALPHA = 0.03125 + + +def test_target_dense_linears() -> None: + """The four NVFP4 dense linears of DeepSeek-V3-Lite, decode through prefill rows. + + (K, N) in nn.Linear [out, in] orientation: layer-0 MLP gate_up [24576, 2560] + and down [2560, 12288], shared-expert gate_up [6144, 2560] and down + [2560, 3072]. + """ + alpha = _alpha(ALPHA) + shapes = [ + (2560, 24576, [1, 2, 8, 64, 1024, 4096]), + (12288, 2560, [1, 8, 64, 1024]), + (2560, 6144, [1, 8, 64, 1024]), + (3072, 2560, [1, 8, 64, 1024]), + ] + for k, n, rows in shapes: + b_data, _, b_sf, b_values = _operand(n, k, seed=1000 + n) + for m in rows: + a_data, _, a_sf, a_values = _operand(m, k, seed=m * 31 + k) + out = nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, torch.bfloat16) + assert out.shape == (m, n) and out.dtype == torch.bfloat16 + ref = _reference(a_values, b_values, ALPHA) + torch.testing.assert_close(out, ref.to(torch.bfloat16)) + del a_data, a_sf, a_values, out, ref + del b_data, b_sf, b_values + torch.cuda.empty_cache() + + +def test_row_counts_around_swizzle_boundaries() -> None: + """M is unconstrained: the activation scale buffer pads rows up to a multiple of 128.""" + alpha = _alpha(ALPHA) + k, n = 2560, 512 + b_data, _, b_sf, b_values = _operand(n, k, seed=7) + for m in [1, 3, 7, 9, 127, 128, 129, 1000]: + a_data, _, a_sf, a_values = _operand(m, k, seed=m) + assert a_sf.numel() == _pad_up(m, 128) * (k // VEC) + out = nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, torch.bfloat16) + torch.testing.assert_close(out, _reference(a_values, b_values, ALPHA).to(torch.bfloat16)) + + +def test_output_dtypes() -> None: + """bf16 / fp16 / fp32 outputs of the same call, each at its own dtype tolerance.""" + alpha = _alpha(ALPHA) + k, n, m = 2560, 512, 8 + a_data, _, a_sf, a_values = _operand(m, k, seed=11) + b_data, _, b_sf, b_values = _operand(n, k, seed=12) + ref = _reference(a_values, b_values, ALPHA) + for dtype in [torch.bfloat16, torch.float16, torch.float32]: + out = nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, dtype) + assert out.dtype == dtype + torch.testing.assert_close(out, ref.to(dtype)) + + +def test_alpha_scales_the_accumulator() -> None: + """alpha multiplies the whole fp32 accumulator, before the output cast.""" + k, n, m = 2560, 512, 8 + a_data, _, a_sf, a_values = _operand(m, k, seed=13) + b_data, _, b_sf, b_values = _operand(n, k, seed=14) + for value in [1.0, 0.5, ALPHA, 1e-4]: + out = nvfp4_gemm(a_data, b_data, a_sf, b_sf, _alpha(value), torch.float32) + torch.testing.assert_close(out, _reference(a_values, b_values, value)) + + +def test_bias_fused() -> None: + """The optional per-column bias is added after alpha, in the output dtype.""" + alpha = _alpha(ALPHA) + k, n, m = 2560, 512, 64 + a_data, _, a_sf, a_values = _operand(m, k, seed=15) + b_data, _, b_sf, b_values = _operand(n, k, seed=16) + torch.manual_seed(0) + bias_bf16 = torch.randn(n, dtype=torch.bfloat16, device=DEV) + out = nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, torch.bfloat16, bias=bias_bf16) + torch.testing.assert_close( + out, _reference(a_values, b_values, ALPHA, bias_bf16).to(torch.bfloat16) + ) + bias_fp32 = torch.randn(n, dtype=torch.float32, device=DEV) + out32 = nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, torch.float32, bias=bias_fp32) + torch.testing.assert_close(out32, _reference(a_values, b_values, ALPHA, bias_fp32)) + + +def test_weight_scale_must_be_swizzled() -> None: + """The weight scale buffer is 128x4-swizzled, exactly like the activation one. + + Pins the load-time obligation for a checkpoint that stores `weight_scale` + row-major linear: the buffer the GEMM needs is `block_scale_interleave` of + that linear tensor, which this test also shows equals the native-torch + swizzle above byte for byte. Feeding the linear buffer instead is accepted + and silently wrong. + """ + alpha = _alpha(ALPHA) + k, n, m = 2560, 512, 16 + a_data, _, a_sf, a_values = _operand(m, k, seed=17) + b_data, b_sf_linear, b_sf, b_values = _operand(n, k, seed=18) + ref = _reference(a_values, b_values, ALPHA) + + interleaved = torch.ops.trtllm.block_scale_interleave(b_sf_linear) + assert torch.equal(interleaved.flatten(), b_sf) + torch.testing.assert_close( + nvfp4_gemm(a_data, b_data, a_sf, interleaved.flatten(), alpha, torch.bfloat16), + ref.to(torch.bfloat16), + ) + + # Same buffer size, linear content: in bounds, accepted, wrong. + linear_padded = torch.zeros_like(b_sf) + linear_padded[: n * (k // VEC)] = b_sf_linear.flatten() + wrong = nvfp4_gemm(a_data, b_data, a_sf, linear_padded, alpha, torch.bfloat16) + assert (wrong.float() - ref).abs().max() > 0.1 * ref.abs().max() + + +def test_weight_rows_not_multiple_of_128() -> None: + """N need only be a multiple of 32; the weight scale buffer still pads rows to 128.""" + alpha = _alpha(ALPHA) + k, n, m = 2560, 160, 8 + a_data, _, a_sf, a_values = _operand(m, k, seed=19) + b_data, _, b_sf, b_values = _operand(n, k, seed=20) + assert b_sf.numel() == 256 * (k // VEC) + out = nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, torch.bfloat16) + torch.testing.assert_close(out, _reference(a_values, b_values, ALPHA).to(torch.bfloat16)) + + +def test_scale_padding_bytes_are_ignored() -> None: + """Whatever fills the row/column padding of a scale buffer never reaches the result.""" + alpha = _alpha(ALPHA) + k, n, m = 2560, 512, 7 + a_data, a_sf_linear, _, a_values = _operand(m, k, seed=21) + b_data, b_sf_linear, _, b_values = _operand(n, k, seed=22) + zero_pad = (_swizzle(a_sf_linear, 0x00), _swizzle(b_sf_linear, 0x00)) + # 0x7E is e4m3's largest finite value (448); 0x7F would be NaN. + loud_pad = (_swizzle(a_sf_linear, 0x7E), _swizzle(b_sf_linear, 0x7E)) + out_zero = nvfp4_gemm(a_data, b_data, *zero_pad, alpha, torch.bfloat16) + out_loud = nvfp4_gemm(a_data, b_data, *loud_pad, alpha, torch.bfloat16) + assert torch.equal(out_zero, out_loud) + torch.testing.assert_close(out_zero, _reference(a_values, b_values, ALPHA).to(torch.bfloat16)) + + +def test_backends() -> None: + """Every selectable backend computes the same GEMM; forcing one is a valid call.""" + alpha = _alpha(ALPHA) + k, n, m = 2560, 512, 8 # m <= 8 so the cuda_core backend is legal too + a_data, _, a_sf, a_values = _operand(m, k, seed=23) + b_data, _, b_sf, b_values = _operand(n, k, seed=24) + ref = _reference(a_values, b_values, ALPHA).to(torch.bfloat16) + for backends in [ + "cutlass,cublaslt,cuda_core", + "cutlass", + "cublaslt", + "cuda_core", + "cutedsl", # JIT-compiled on first use + ]: + out = nvfp4_gemm( + a_data, + b_data, + a_sf, + b_sf, + alpha, + torch.bfloat16, + allowed_backends=backends, + ) + torch.testing.assert_close(out, ref, msg=lambda s, b=backends: f"{b}: {s}") + + +def test_autotuned_selection() -> None: + """A tuner-selected tactic computes the same GEMM as the untuned fallback.""" + alpha = _alpha(ALPHA) + k, n, m = 2560, 6144, 8 + a_data, _, a_sf, a_values = _operand(m, k, seed=25) + b_data, _, b_sf, b_values = _operand(n, k, seed=26) + ref = _reference(a_values, b_values, ALPHA).to(torch.bfloat16) + cache = AutoTuner.get().profiling_cache + cache.clear() + try: + untuned = nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, torch.bfloat16) + torch.testing.assert_close(untuned, ref) + with autotune(): + nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, torch.bfloat16) + assert len(cache.cache) > 0, "autotune recorded no tactic" + tuned = nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, torch.bfloat16) + torch.testing.assert_close(tuned, ref) + finally: + cache.clear() # leave the process-global tuner state as we found it + + +def test_wide_dynamic_range() -> None: + """Block scales spanning 2^-5 .. 2^5, where the fp32 accumulation actually rounds. + + With the narrow scale range used elsewhere every product and partial sum is + exactly representable, so the kernel matches an fp32 reference bit for bit. + Here it cannot, and a fixed tolerance would be a fitted number: the check is + the textbook recursive-summation bound instead, computed elementwise from + the operands. + """ + alpha = _alpha(ALPHA) + k, n, m = 12288, 512, 32 + a_data, _, a_sf, a_values = _operand(m, k, seed=27, sf_lo=0x10, sf_hi=0x61) + b_data, _, b_sf, b_values = _operand(n, k, seed=28, sf_lo=0x10, sf_hi=0x61) + ref64 = ALPHA * (a_values.double() @ b_values.double().T) + out32 = nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, torch.float32) + out = nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, torch.bfloat16) + # fp32 summation of K terms deviates from the exact sum by at most + # K * 2^-24 * sum|a_i b_i| (recursive-summation bound), and the bf16 cast + # adds at most half a bf16 ulp (2^-9 relative). Both are bounds computed + # from the operands, not fitted tolerances. + conditioning = ALPHA * (a_values.abs().double() @ b_values.abs().double().T) + summation_bound = k * 2.0**-24 * conditioning + assert not torch.equal(out32.double(), ref64), "this regime should not be exact" + assert ((out32.double() - ref64).abs() <= summation_bound).all() + assert ((out.double() - ref64).abs() <= summation_bound + 2.0**-9 * ref64.abs()).all() + + +def test_inputs_untouched_and_deterministic() -> None: + """The op writes only its freshly allocated output, and repeats bit for bit.""" + alpha = _alpha(ALPHA) + k, n, m = 2560, 512, 64 + a_data, _, a_sf, _ = _operand(m, k, seed=29) + b_data, _, b_sf, _ = _operand(n, k, seed=30) + snapshots = [t.clone() for t in (a_data, b_data, a_sf, b_sf, alpha)] + first = nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, torch.bfloat16) + second = nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, torch.bfloat16) + for before, after in zip(snapshots, (a_data, b_data, a_sf, b_sf, alpha)): + assert torch.equal(before, after), "an input was modified" + assert torch.equal(first, second), "not bitwise deterministic" + assert first.is_contiguous() and first.data_ptr() != second.data_ptr() + + +# The four NVFP4 dense linears of DeepSeek-R1-0528, in nn.Linear [out, in] +# orientation. Every K and N here exceeds the largest value the shapes above +# reach: K = 18432 (vs 12288) and N = 36864 (vs 24576), and the K = 18432 +# weight-scale buffer is pad_up(N, 128) * pad_up(K/16, 4) with K/16 = 1152 +# columns -- the widest block-scale buffer in this file. +R1_DENSE_LINEARS = [ + (7168, 36864), # dense-MLP gate_up (layers 0-2), N = 2 * 18432 + (18432, 7168), # dense-MLP down + (7168, 4096), # shared-expert gate_up, N = 2 * 2048 + (2048, 7168), # shared-expert down +] +# Decode rows through the serving prefill cap: the dense path is not chunked, +# so M reaches max_num_tokens = 8192 in a single call. +R1_ROWS = [1, 2, 8, 64, 1024, 4096, 8192] + + +def _gate_rejects(out: torch.Tensor, ref_bf16: torch.Tensor) -> bool: + """True when the default-tolerance comparison the tests above use fails.""" + try: + torch.testing.assert_close(out, ref_bf16) + except AssertionError: + return True + return False + + +def test_r1_dense_linears() -> None: + """The four NVFP4 dense linears of DeepSeek-R1-0528, decode through max_num_tokens.""" + alpha = _alpha(ALPHA) + for k, n in R1_DENSE_LINEARS: + b_data, _, b_sf, b_values = _operand(n, k, seed=2000 + k + n) + assert b_sf.numel() == _pad_up(n, 128) * _pad_up(k // VEC, 4) + for m in R1_ROWS: + a_data, _, a_sf, a_values = _operand(m, k, seed=2100 + m + k) + assert a_sf.numel() == _pad_up(m, 128) * _pad_up(k // VEC, 4) + out = nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, torch.bfloat16) + assert out.shape == (m, n) and out.dtype == torch.bfloat16 + ref = _reference(a_values, b_values, ALPHA) + torch.testing.assert_close(out, ref.to(torch.bfloat16)) + del a_data, a_sf, a_values, out, ref + del b_data, b_sf, b_values + torch.cuda.empty_cache() + + +def test_r1_backends_agree() -> None: + """Every backend computes the same GEMM at the R1 shapes too. + + Backend selection is shape-sensitive (cuda_core is only offered for M <= 8, + and the tactic space differs per (K, N)), so agreement is re-established at + each new shape rather than inherited. In this scale band every product and + partial sum is exactly representable in fp32 -- checked directly at K = + 18432 by test_r1_wide_scale_range -- so the backends agree *bitwise*, which + is a far sharper check than the tolerance one. + """ + alpha = _alpha(ALPHA) + for k, n in R1_DENSE_LINEARS: + b_data, _, b_sf, b_values = _operand(n, k, seed=2200 + k + n) + for m in [8, 8192]: # 8: the only band where cuda_core is selectable + a_data, _, a_sf, a_values = _operand(m, k, seed=2300 + m + k) + ref = _reference(a_values, b_values, ALPHA).to(torch.bfloat16) + backends = ["cutlass", "cublaslt", "cutedsl"] + if m <= 8: + backends.append("cuda_core") + outs = [] + for backend in backends: + out = nvfp4_gemm( + a_data, + b_data, + a_sf, + b_sf, + alpha, + torch.bfloat16, + allowed_backends=backend, + ) + torch.testing.assert_close( + out, ref, msg=lambda s, b=backend: f"{b} K={k} N={n} M={m}: {s}" + ) + outs.append(out) + for backend, out in zip(backends[1:], outs[1:]): + assert torch.equal(out, outs[0]), ( + f"{backend} differs from cutlass at K={k} N={n} M={m}" + ) + del a_data, a_sf, a_values, ref, outs + torch.cuda.empty_cache() + del b_data, b_sf, b_values + torch.cuda.empty_cache() + + +def test_r1_autotuned_selection() -> None: + """The tuned path at the R1 shapes, which is the one a serving run executes. + + One tuning call at M = 8192 fills all 14 power-of-2 buckets for that + (K, N), and every cached winner carries a profiled tactic id rather than + the fallback marker -- so warm really is a different execution from cold. + Each M is then run cold, tuned, and run again warm, and the two results + compared **bitwise**: a tactic that changed the result by a single bit + would fail here even though it would sail through the tolerance gate. The M + values land in four different buckets, two of them non-powers of 2. + """ + alpha = _alpha(ALPHA) + rows = [1, 3, 8, 4097, 8192] + cache = AutoTuner.get().profiling_cache + try: + for k, n in R1_DENSE_LINEARS: + b_data, _, b_sf, b_values = _operand(n, k, seed=2400 + k + n) + cache.clear() + cold = {} + for m in rows: + a_data, _, a_sf, a_values = _operand(m, k, seed=2600 + m + k) + cold[m] = nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, torch.bfloat16) + ref = _reference(a_values, b_values, ALPHA) + torch.testing.assert_close(cold[m], ref.to(torch.bfloat16)) + del a_data, a_sf, a_values, ref + + a_data, _, a_sf, _ = _operand(8192, k, seed=2500 + k) + with autotune(): + nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, torch.bfloat16) + del a_data, a_sf + assert len(cache.cache) == 14, ( + f"K={k} N={n}: expected buckets 1..8192, got {len(cache.cache)}" + ) + for value in cache.cache.values(): + backend, sub_tactic = value[1] + assert backend in {"cutlass", "cublaslt", "cuda_core"}, backend + assert sub_tactic >= 0, f"{backend} winner is the fallback marker" + + for m in rows: + a_data, _, a_sf, a_values = _operand(m, k, seed=2600 + m + k) + tuned = nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, torch.bfloat16) + ref = _reference(a_values, b_values, ALPHA) + torch.testing.assert_close(tuned, ref.to(torch.bfloat16)) + assert torch.equal(tuned, cold[m]), ( + f"the tuned tactic changed the bits at K={k} N={n} M={m}" + ) + del a_data, a_sf, a_values, tuned, ref + del b_data, b_sf, b_values, cold + torch.cuda.empty_cache() + finally: + cache.clear() # leave the process-global tuner state as we found it + + +def test_r1_wrong_variants_are_visible() -> None: + """Control: the comparisons above can see wrongness at the R1 widths. + + A clean sweep is only worth what the harness could have caught. Three wrong + variants at the two extreme R1 shapes, each measured against the same + default-tolerance gate the correctness tests use: + - the weight scale fed linear (unswizzled) instead of 128x4-swizzled; + - one e4m3 weight-scale byte moved by one code; + - one e2m1 weight nibble moved by one code. + The last two are the smallest perturbation either operand admits, and at + M = 8192 they still fail the gate -- so a tactic that changed results by + one code anywhere could not pass unnoticed. + """ + alpha = _alpha(ALPHA) + for k, n, m in [(18432, 7168, 8192), (7168, 36864, 8192)]: + a_data, _, a_sf, a_values = _operand(m, k, seed=2700 + k) + b_data, b_sf_linear, b_sf, b_values = _operand(n, k, seed=2800 + k + n) + ref = _reference(a_values, b_values, ALPHA) + ref_bf16 = ref.to(torch.bfloat16) + + # Positive control: the op's own relayout of a row-major [N, K/16] + # scale tensor is byte-identical to the native-torch swizzle at these + # widths (K/16 = 1152 and 448), and gives the reference result. + interleaved = torch.ops.trtllm.block_scale_interleave(b_sf_linear).flatten() + assert torch.equal(interleaved, b_sf) + base = nvfp4_gemm(a_data, b_data, a_sf, interleaved, alpha, torch.bfloat16) + assert not _gate_rejects(base, ref_bf16) + + # Same buffer size, linear content: in bounds, accepted, wrong. + linear_padded = torch.zeros_like(b_sf) + linear_padded[: n * (k // VEC)] = b_sf_linear.flatten() + wrong = nvfp4_gemm(a_data, b_data, a_sf, linear_padded, alpha, torch.bfloat16) + assert (wrong.float() - ref).abs().max() > 0.1 * ref.abs().max() + assert _gate_rejects(wrong, ref_bf16) + del linear_padded, wrong + + # One e4m3 scale byte, one code up. + b_sf_bumped = b_sf.clone() + i = b_sf_bumped.numel() // 3 + b_sf_bumped[i] += 1 + one_byte = nvfp4_gemm(a_data, b_data, a_sf, b_sf_bumped, alpha, torch.bfloat16) + assert _gate_rejects(one_byte, ref_bf16), "a one-code scale change is invisible" + del b_sf_bumped, one_byte + + # One e2m1 data nibble, one code up. + b_bumped = b_data.clone() + flat = b_bumped.view(-1) + j = flat.numel() // 3 + flat[j] = (flat[j] & 0xF0) | (((flat[j] & 0x0F) + 1) & 0x0F) + one_nibble = nvfp4_gemm(a_data, b_bumped, a_sf, b_sf, alpha, torch.bfloat16) + assert _gate_rejects(one_nibble, ref_bf16), "a one-code data change is invisible" + + del a_data, a_sf, a_values, b_data, b_sf, b_sf_linear, b_values + del ref, ref_bf16, base, interleaved, b_bumped, flat, one_nibble + torch.cuda.empty_cache() + + +def test_r1_wide_scale_range() -> None: + """The fp32 accumulator at K = 18432, the largest K certified here. + + Narrow block scales (the band every other test uses) keep every product and + partial sum exactly representable, so the fp32 output is bit-exact against + an fp64 reference at this K too -- which is what licenses the bitwise + cross-backend assertions above. Widening the scales to 2^-5..2^5 makes the + sum round; the deviation is then checked against the recursive-summation + bound computed from the operands, not against a fitted tolerance. + """ + alpha = _alpha(ALPHA) + k, n = 18432, 7168 + a_data, _, a_sf, a_values = _operand(64, k, seed=2900) + b_data, _, b_sf, b_values = _operand(n, k, seed=2901) + out32 = nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, torch.float32) + assert torch.equal(out32.double(), ALPHA * (a_values.double() @ b_values.double().T)) + del a_data, a_sf, a_values, b_data, b_sf, b_values, out32 + torch.cuda.empty_cache() + + m = 32 + a_data, _, a_sf, a_values = _operand(m, k, seed=2902, sf_lo=0x10, sf_hi=0x61) + b_data, _, b_sf, b_values = _operand(n, k, seed=2903, sf_lo=0x10, sf_hi=0x61) + ref64 = ALPHA * (a_values.double() @ b_values.double().T) + out32 = nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, torch.float32) + out = nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, torch.bfloat16) + conditioning = ALPHA * (a_values.abs().double() @ b_values.abs().double().T) + summation_bound = k * 2.0**-24 * conditioning + assert not torch.equal(out32.double(), ref64), "this regime should not be exact" + assert ((out32.double() - ref64).abs() <= summation_bound).all() + assert ((out.double() - ref64).abs() <= summation_bound + 2.0**-9 * ref64.abs()).all() + torch.cuda.empty_cache() + + +def test_rejected_domains() -> None: + """Domains the op rejects loudly (backend-independent unless noted).""" + alpha = _alpha(ALPHA) + k, n, m = 256, 128, 8 + a_data, _, a_sf, _ = _operand(m, k, seed=31) + b_data, _, b_sf, _ = _operand(n, k, seed=32) + bf16 = torch.bfloat16 + + def rejects(fn) -> None: + try: + fn() + except (RuntimeError, ValueError): + return + raise AssertionError("expected the op to raise") + + # K and N alignment (16-byte operand lines): both must be multiples of 32. + a48, _, a48_sf, _ = _operand(m, 48, seed=33) + b48, _, b48_sf, _ = _operand(n, 48, seed=34) + rejects(lambda: nvfp4_gemm(a48, b48, a48_sf, b48_sf, alpha, bf16)) + b_narrow, _, b_narrow_sf, _ = _operand(16, k, seed=35) + rejects(lambda: nvfp4_gemm(a_data, b_narrow, a_sf, b_narrow_sf, alpha, bf16)) + + # dtypes: data and scales are uint8 byte buffers, alpha is a CUDA fp32 scalar + rejects(lambda: nvfp4_gemm(a_data.view(torch.int8), b_data, a_sf, b_sf, alpha, bf16)) + rejects(lambda: nvfp4_gemm(a_data, b_data.view(torch.int8), a_sf, b_sf, alpha, bf16)) + rejects(lambda: nvfp4_gemm(a_data, b_data, a_sf.view(torch.float8_e4m3fn), b_sf, alpha, bf16)) + rejects(lambda: nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha.half(), bf16)) + rejects(lambda: nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha.cpu(), bf16)) + rejects(lambda: nvfp4_gemm(a_data.cpu(), b_data, a_sf, b_sf, alpha, bf16)) + + # rank and K agreement + rejects(lambda: nvfp4_gemm(a_data.reshape(2, 4, k // 2), b_data, a_sf, b_sf, alpha, bf16)) + rejects(lambda: nvfp4_gemm(a_data, b_data[:, : k // 4].contiguous(), a_sf, b_sf, alpha, bf16)) + + # zero rows + rejects(lambda: nvfp4_gemm(a_data[:0], b_data, a_sf, b_sf, alpha, bf16)) + + # bias must be 1-D [N] in the output dtype + bias = torch.zeros(n, dtype=torch.bfloat16, device=DEV) + rejects(lambda: nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, bf16, bias=bias.float())) + rejects(lambda: nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, bf16, bias=bias.view(1, n))) + rejects(lambda: nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, bf16, bias=bias[: n - 32])) + + # allowed_backends parsing, and backends this arch cannot run + rejects(lambda: nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, bf16, 0, "")) + rejects(lambda: nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, bf16, 0, "cutlas")) + rejects(lambda: nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, bf16, 0, "marlin")) + rejects(lambda: nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, torch.float16, 0, "cutedsl")) + # the cuda_core kernel tops out at M = 16 (the tuner only ever picks it for + # M <= 8; forcing it is what can reach the kernel's own limit) + a_big, _, a_big_sf, _ = _operand(32, k, seed=36) + rejects(lambda: nvfp4_gemm(a_big, b_data, a_big_sf, b_sf, alpha, bf16, 0, "cuda_core")) + # userbuffers output without a userbuffers workspace + rejects(lambda: nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, bf16, 1)) + + +def test_wrapper_guards_silent_domains() -> None: + """Each wrapper assert stands where the op itself is silently wrong.""" + alpha = _alpha(ALPHA) + k, n, m = 2560, 512, 8 + a_data, _, a_sf, a_values = _operand(m, k, seed=37) + b_data, _, b_sf, b_values = _operand(n, k, seed=38) + ref = _reference(a_values, b_values, ALPHA) + + # A non-contiguous view whose *values* equal the contiguous operand: any + # deviation is the kernel ignoring the stride. + a_nc = torch.cat([a_data, a_data], 1)[:, : k // 2] + b_nc = torch.cat([b_data, b_data], 1)[:, : k // 2] + a_sf_nc = torch.stack([a_sf, torch.zeros_like(a_sf)], 1).flatten()[::2] + b_sf_nc = torch.stack([b_sf, torch.zeros_like(b_sf)], 1).flatten()[::2] + assert torch.equal(a_nc, a_data) and not a_nc.is_contiguous() + assert torch.equal(a_sf_nc, a_sf) and not a_sf_nc.is_contiguous() + + def silently_wrong(*args) -> None: + out = torch.ops.trtllm.nvfp4_gemm(*args, alpha, torch.bfloat16, 0, "cublaslt", None, None) + assert (out.float() - ref).abs().max() > 0.1 * ref.abs().max() + + def guarded(fn) -> None: + try: + fn() + except AssertionError: + return + raise AssertionError("expected the wrapper to reject this") + + silently_wrong(a_nc, b_data, a_sf, b_sf) + guarded(lambda: nvfp4_gemm(a_nc, b_data, a_sf, b_sf, alpha, torch.bfloat16)) + silently_wrong(a_data, b_nc, a_sf, b_sf) + guarded(lambda: nvfp4_gemm(a_data, b_nc, a_sf, b_sf, alpha, torch.bfloat16)) + silently_wrong(a_data, b_data, a_sf_nc, b_sf) + guarded(lambda: nvfp4_gemm(a_data, b_data, a_sf_nc, b_sf, alpha, torch.bfloat16)) + silently_wrong(a_data, b_data, a_sf, b_sf_nc) + guarded(lambda: nvfp4_gemm(a_data, b_data, a_sf, b_sf_nc, alpha, torch.bfloat16)) + + # A multi-element alpha is accepted and every element past the first ignored. + fat_alpha = torch.tensor([ALPHA, 99.0], dtype=torch.float32, device=DEV) + out = torch.ops.trtllm.nvfp4_gemm(a_data, b_data, a_sf, b_sf, fat_alpha, torch.bfloat16) + torch.testing.assert_close(out, ref.to(torch.bfloat16)) + guarded(lambda: nvfp4_gemm(a_data, b_data, a_sf, b_sf, fat_alpha, torch.bfloat16)) diff --git a/tensorrt_llm/_torch/staircase/catalog/index.yaml b/tensorrt_llm/_torch/staircase/catalog/index.yaml new file mode 100644 index 000000000000..fefe0dd9fb54 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/index.yaml @@ -0,0 +1,200 @@ +# Catalog Index +# +# Quick-reference for entry discovery. The path (relative to catalog/) is the +# unique identifier; the directory encodes the category. Synced manually after +# catalog changes. +# +# The catalog is the complete vocabulary of target modeling code: every call +# that creates or transforms a tensor must be a catalog entry. Only tensor +# metadata reads (.shape/.dtype/.device) and Python control flow live outside. +# +# AN ENTRY SPANS TWO TREES. The contract (.md) and the wrapper (.py) are here; +# the GPU test is at +# tests/unittest/_torch/staircase//test_staircase_.py, because +# that is the tree this repo's CI collects from -- a test inside the package is +# not picked up by any list. The receipt rule is unchanged in substance and +# only gains a second directory to look in: a receipt is valid only if it +# post-dates the last write to *every* file of the entry, contract and wrapper +# and test alike. Check it mechanically against both paths. +# +# The tests are named test_staircase_* rather than test_: three of them +# would otherwise collide with an upstream test of the same basename, and more +# importantly they are not duplicates of those -- upstream covers the modules +# that wrap these ops, these cover the op itself cell by cell. +# +# torch/ entries are hand-added thin mirrors of torch callables — upstream +# owns their correctness, so they carry no contract doc, no test, and no +# arch/precision certification (absence = unconstrained by design; real dtype +# limits, e.g. no float8 arithmetic, are noted in the entry docstring). +# Entries wrapping trtllm ops carry all three. +# +# ── RECEIPT STATUS: all 19 entries certified on sm_103 ──────────────────── +# +# A receipt says this entry's test passed, on a stated GPU architecture, over +# files no newer than the run. Both anchors moved at once here -- the +# migration rewrote every test file, and the targets moved from sm_100 (B200) +# to sm_103 (GB300) -- so every sm_100 receipt was voided and the whole set +# was re-run on GB300 under 1.3.0rc26. **All 19 pass**; per-entry counts are +# in each contract's `receipts:` frontmatter. +# +# Getting there surfaced four real differences, none of which was fixed by +# widening a tolerance: +# +# * thop_attention, mla_rope_generation, mla_rope_append_paged_kv_assign_q +# -- op schema drift rc21 -> rc26 (renamed and added parameters). The +# wrappers now mirror their schemas argument for argument, so the next +# drift fails loudly instead of shifting a positional list. +# +# * mxe4m3_mxe2m1_block_scale_moe_runner -- the FC1 epilogue's MXFP8 block +# scale is `floor(log2(amax))-8` on sm_100 and `ceil(log2(amax/448))` on +# sm_103, bit-exactly. The reference is now architecture-keyed and each +# architecture refutes the other's recipe. See that contract. +# +# * cublas_mm -- torch 2.12 made fp32 matmul default to TF32, so the +# *reference* was the imprecise side. The op is bit-identical to a +# TF32-disabled torch product. See that contract. +# +# The "measured on sm_100" statements throughout the contracts are left +# exactly as written. They are true records of what was observed on another +# device, and rewriting them would manufacture GB300 evidence that does not +# exist. Read them as provenance; the frontmatter is the certification. +# ────────────────────────────────────────────────────────────────────────── +# +# 30 of the source catalog's 43 entries are here — the union of what the two +# migrated targets call. The 13 left behind (silu_and_mul, attn_custom_op_inplace, +# create_attn_outputs, create_mla_outputs, load_chunked_kv_cache_for_mla, +# merge_chunked_attention_for_mla, cute_dsl_bf16_gemm_blackwell, +# flashinfer_fused_add_rmsnorm_quant, flashinfer_gemma_rmsnorm, +# flashinfer_gemma_fused_add_rmsnorm, bf16_mxe2m1_block_scale_moe_runner, +# renorm_moe_routing_op, allreduce) must arrive with their contract AND their +# receipt when a later target needs them: a wrapper on its own would let that +# target consume an op with no certification record at all. + +entries: + # ─── torch (native callables, trtllm-free) ──────────────────── + - path: torch/embedding.py + impl: torch.nn.functional.embedding + summary: "Input-embedding lookup: token ids -> rows of the embedding weight" + + - path: torch/concat.py + impl: torch.cat + summary: "Concatenate a list of tensors along a dim" + + - path: torch/split.py + impl: torch.split + summary: "Split a tensor along a dim into equal chunks or explicit sizes" + + - path: torch/add.py + impl: torch.add + summary: "Elementwise add (residual); torch has no float8 arithmetic" + + - path: torch/reshape.py + impl: torch.reshape + summary: "Reshape (covers view: returns a view when layout allows)" + + - path: torch/empty.py + impl: torch.empty + summary: "Uninitialized tensor allocation (scratch/output buffers whose sizing the caller owns)" + + - path: torch/expand.py + impl: torch.Tensor.expand + summary: "Zero-copy broadcast view over singleton dims (stride-0; unsafe to write through)" + + - path: torch/transpose.py + impl: torch.transpose + summary: "Zero-copy swap of two dims (strided view; presents [tokens, heads, dim] to a head-batched gemm)" + + - path: torch/view_dtype.py + impl: torch.Tensor.view + summary: "Bitwise dtype reinterpretation of the same bytes (uint8 scale buffer -> float8_e4m3fn view); no conversion, no copy" + + - path: torch/copy_.py + impl: torch.Tensor.copy_ + summary: "In-place elementwise copy into a tensor or view (canonical slice fill; destroys dst contents)" + + - path: torch/pad.py + impl: torch.nn.functional.pad + summary: "Constant padding of trailing dims (widening an activation to a kernel's operand width)" + + # ─── norm ────────────────────────────────────────────────────── + - path: norm/flashinfer_rmsnorm.py + impl: torch.ops.trtllm.flashinfer_rmsnorm + summary: "RMS normalization over the last dim with elementwise weight scaling" + + - path: norm/flashinfer_fused_add_rmsnorm.py + impl: torch.ops.trtllm.flashinfer_fused_add_rmsnorm + summary: "In-place fused residual add + RMS norm: residual += x; x = rmsnorm(residual) * weight" + + # ─── activation ──────────────────────────────────────────────── + - path: activation/flashinfer_silu_and_mul.py + impl: torch.ops.trtllm.flashinfer_silu_and_mul + summary: "SwiGLU activation: silu(x[..., :d]) * x[..., d:] over a packed gate/up last dim" + + # ─── gemm ────────────────────────────────────────────────────── + - path: gemm/cublas_mm.py + impl: torch.ops.trtllm.cublas_mm + summary: "General matmul via one cuBLASLt call with optional fused bias: [M,K] @ [K,N] col-major (+ bias), fp32 accumulation" + + - path: gemm/bmm_out.py + impl: torch.ops.trtllm.bmm_out + summary: "Batched matmul into a caller-provided (possibly strided-view) out buffer: [B,M,K] @ [B,K,N] -> [B,M,N], fp32 accumulation" + + - path: gemm/nvfp4_gemm.py + impl: torch.ops.trtllm.nvfp4_gemm + summary: "NVFP4 x NVFP4 dense GEMM in nn.Linear layout with in-kernel block-scale dequantization: alpha * ([M,K/2] e2m1 + 128x4-swizzled e4m3 scales) @ ([N,K/2] e2m1 + swizzled scales)^T (+ per-column bias) -> fresh [M,N] bf16/fp16/fp32 buffer, fp32 accumulation, backend auto-selected (cutlass/cublaslt/cuda_core/cutedsl)" + + # ─── attention ───────────────────────────────────────────────── + - path: attention/fused_qk_norm_rope.py + impl: torch.ops.trtllm.fused_qk_norm_rope + summary: "In-place fused per-head QK RMS norm + RoPE (neox/interleaved, YaRN, interleaved mRoPE) on a packed bf16 QKV tensor; v heads untouched" + + - path: attention/mla_rope_generation.py + impl: torch.ops.trtllm.mla_rope_generation + summary: "MLA decode preprocessing: GPT-J RoPE of q_pe + latent [compressed_kv | rope(k_pe)] paged-KV-cache append + decode-FMHA scheduler-buffer fill; bf16 pool writes fused_q's tail, fp8-e4m3 pool writes quant_q_buffer and the two folded bmm scales instead" + + - path: attention/load_paged_kv_cache_for_mla.py + impl: torch.ops.trtllm.load_paged_kv_cache_for_mla + summary: "Gather each context sequence's full [past + new] MLA latent KV from the paged cache into two new contiguous tensors (compressed_kv [T,C], k_pe [T,R]) for KV-reuse context prefill; an fp8-e4m3 pool is dequantized to bf16/fp16/fp32 by kv_scale_quant_orig" + + - path: attention/mla_rope_append_paged_kv_assign_q.py + impl: torch.ops.trtllm.mla_rope_append_paged_kv_assign_q + summary: "MLA context-prefill preprocessing: in-place GPT-J RoPE of q_pe (in q) and k_pe (in latent_cache) at each new token's absolute position + [compressed_kv | rope(k_pe)] paged-latent-cache append; an fp8-e4m3 pool quantizes only the appended row, by kv_scale_orig_quant" + + - path: attention/thop_attention.py + impl: tensorrt_llm.bindings.internal.thop.attention + summary: "Full attention core with fully explicit state (pybind binding, approved exception): paged KV-cache append + causal/padding-masked GQA FMHA over a caller-owned pool (bf16, or fp8-e4m3 with per-tensor kv scales) and explicit length/offset tensors, written into a caller buffer, on either context execution path (use_paged_context_fmha selects packed-QKV context FMHA, or paged-KV context FMHA so a context call may run over a cached prefix — KV-cache reuse and chunked prefill), with optional per-query-head attention sinks (one extra softmax-denominator logit, dropped from the output) and an optional per-call sliding window (attention_window_size keys ending at the query's absolute position; a pure mask — the append stays at absolute positions, so cyclic pool reuse is the caller's page mapping); MLA mode runs context prefill (in-kernel RoPE + latent append), no-append context over explicit K/V (latent_cache=None: cached-KV prefixes and chunked partial passes with softmax-stats output), and generation latent-MQA decode as separate calls over a paged latent pool (bf16 or fp8-e4m3)" + + # ─── moe ─────────────────────────────────────────────────────── + - path: moe/noaux_tc_op.py + impl: torch.ops.trtllm.noaux_tc_op + summary: "DeepSeek-V3 style MoE routing (noaux_tc): in-kernel sigmoid of the router logits + per-expert bias correction for selection only + optional group-limited top-k (n_group/topk_group; identity at n_group=1) + combine weights gathered from the unbiased sigmoid, renormalized and scaled by routed_scaling_factor, returning weights in the logits' dtype and int32 expert ids" + + - path: moe/fused_moe.py + impl: torch.ops.trtllm.fused_moe + summary: "Full MoE layer over pre-routed tokens: expert permutation + grouped FC1 GEMM over stacked [up; gate] expert weights + gated activation (SwiGLU with optional per-expert alpha/beta/limit, or GeGLU) + grouped FC2 GEMM + routing-weighted fp32 combine, into a fresh or caller-provided [tokens, hidden] buffer (bf16/fp16 unquantized path; expert-parallel slot range selected by ep_size/ep_rank)" + + - path: moe/mxe4m3_mxe2m1_block_scale_moe_runner.py + impl: torch.ops.trtllm.mxe4m3_mxe2m1_block_scale_moe_runner + summary: "Full MXFP4-weight / MXFP8-activation MoE layer (trtllm-gen W4A8): optional top-k routing over router logits (softmax-then-topk or topk-then-softmax) or caller-supplied top-k ids/weights + grouped FC1 GEMM over pre-shuffled block-scaled [up; gate] expert weights + clamped gated activation + MXFP8 requantization of that activation + grouped FC2 GEMM + routing-weighted fp32 combine, into a fresh or caller-provided [tokens, valid_hidden_size] bf16 buffer (e4m3 hidden states + linear UE8M0 block scales in, expert-parallel window selected by local_expert_offset/local_num_experts)" + + - path: moe/fp4_block_scale_moe_runner.py + impl: torch.ops.trtllm.fp4_block_scale_moe_runner + summary: "Full NVFP4-weight / NVFP4-activation MoE layer (trtllm-gen W4A4): caller-supplied top-k ids/weights + grouped FC1 GEMM over pre-shuffled block-scaled [up; gate] expert weights + SwiGLU + NVFP4 requantization of that activation + grouped FC2 GEMM + routing-weighted fp32 combine, into a fresh or caller-provided [tokens, hidden] bf16 buffer; the checkpoint's per-tensor global scales enter as three per-expert fp32 scalars, hidden must be a multiple of 256, expert-parallel window selected by local_expert_offset/local_num_experts, do_finalize=False returns the per-slot rows plus their permuted-row map instead of the combine" + + # ─── quantization ────────────────────────────────────────────── + - path: quantization/mxfp8_quantize.py + impl: torch.ops.trtllm.mxfp8_quantize + summary: "Dynamic MXFP8 activation quantization: bf16/fp16 -> e4m3 data (last dim zero-padded up to `alignment`) + one UE8M0 power-of-two scale per 32 contiguous elements, emitted either row-major linear or 128x4 swizzled" + + - path: quantization/fp4_quantize.py + impl: torch.ops.trtllm.fp4_quantize + summary: "Static-scale NVFP4 activation quantization: bf16/fp16 -> packed e2m1 data (two nibbles per byte, no width padding) + one e4m3 scale per 16 contiguous elements, scale = e4m3(global_scale * blockmax / 6) and data = e2m1(x * global_scale / scale), emitted either row-major linear (trtllm-gen block-scale MoE) or 128x4 swizzled (nvfp4 dense GEMM); MXFP4 mode (32-element blocks, UE8M0 scales) accepted but not certified" + + # ─── comm ────────────────────────────────────────────────────── + - path: comm/allgather.py + impl: torch.ops.trtllm.allgather + summary: "All-gather over a group of MPI-session ranks into a fresh tensor: every rank's rows concatenated along dim 0 in ascending rank order, uniform (sizes=None) or ragged (per-rank row counts in `sizes`, as attention data parallelism produces); a byte-exact move in every dtype, no workspace and no strategy, capturable in a CUDA graph once the group's communicator exists (a captured ragged split is frozen at capture); inside a trtllm serving engine under attention DP the communicator carries the caller's calls only — the runtime synchronizes its ranks with host-side MPI on a different one — and calls pair by position, so ranks that disagree on call order get silently wrong data rather than a hang whenever the mispaired calls move the same number of bytes" + + - path: comm/reducescatter.py + impl: torch.ops.trtllm.reducescatter + summary: "Summing reduce-scatter over a group of MPI-session ranks into a fresh tensor: every rank's tensor summed elementwise, then split along dim 0 in ascending rank order so each rank keeps its own rows, even (sizes=None) or uneven (per-rank row counts in `sizes`, as attention data parallelism produces); the sum is accumulated in the input dtype in a destination-dependent ring order — deterministic and replay-reproducible, but not the correctly rounded fp32 sum — and float8_e4m3fn is summed as raw bytes rather than as floats; calls pair by position on the communicator whichever `sizes` form they take, and because this one computes, ranks that disagree on call order make every rank wrong nearly everywhere rather than a quarter of the way — silently, and bitwise reproducibly, whenever the mispaired calls move the same number of bytes, hanging only when they do not" diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/__init__.py b/tensorrt_llm/_torch/staircase/catalog/moe/__init__.py new file mode 100644 index 000000000000..5e898b4e1168 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/moe/__init__.py @@ -0,0 +1,3 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Mixture-of-experts entries.""" diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner.md b/tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner.md new file mode 100644 index 000000000000..ef65be84bc04 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner.md @@ -0,0 +1,660 @@ +--- +receipts: + sm_100: {status: passed, trtllm: 1.3.0rc21} + sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 16} +--- + +# fp4_block_scale_moe_runner + +**Wraps** `torch.ops.trtllm.fp4_block_scale_moe_runner` (one call). + +## Semantics + +One complete mixture-of-experts layer for **NVFP4 weights over NVFP4 +activations** (E2M1 elements with one E4M3 scale per 16 contiguous K elements +on both operands) — the trtllm-gen "block scale MoE" family's W4A4 member. In +a single call: optional top-k routing over router logits, expert permutation, +the grouped FC1 GEMM, the gated activation, **requantization of that +activation to NVFP4**, the grouped FC2 GEMM, and the routing-weighted combine +back to one row per token. + +Let `T = hidden_states.shape[0]`, `H` = the model hidden size, +`I = intermediate_size`, `K = top_k`, `E = local_num_experts`, +`off = local_expert_offset`. + +### The three scale scalars, and the two global scales behind them + +The kernel's two GEMMs consume raw E2M1 codes and E4M3 block scales; the +**per-tensor global scales** that a modelopt NVFP4 checkpoint carries are not +inside those operands, so they arrive as the three `[E]` fp32 tensors. Write + +- `g1` = the FC1 **activation** global scale (`hidden_states` was quantized + with it), +- `g2` = the FC2 **activation** global scale (the FC1 epilogue quantizes with + it), +- `gw1[e]`, `gw2[e]` = the per-expert weight global scales of the FC1 and FC2 + weights. + +Then, exactly: + +``` +output1_scale_gate_scalar[e] = 1 / (g1 * gw1[e]) # "alpha" of FC1 +output1_scale_scalar[e] = g2 / (g1 * gw1[e]) # that, times g2 +output2_scale_scalar[e] = 1 / (g2 * gw2[e]) # "alpha" of FC2 +``` + +A modelopt checkpoint stores **reciprocals**: its `input_scale` and +`weight_scale_2` scalars are `amax / (448*6)`, so `g1 = 1 / input_scale_fc1`, +`g2 = 1 / input_scale_fc2` (the `down_proj`'s `input_scale`), and + +``` +output1_scale_gate_scalar[e] = input_scale_fc1 * weight_scale_2_fc1[e] +output1_scale_scalar[e] = output1_scale_gate_scalar[e] / input_scale_fc2 +output2_scale_scalar[e] = input_scale_fc2 * weight_scale_2_fc2[e] +``` + +Getting the two `output1_*` scalars the wrong way round is finite and +plausible — measured 177 bf16 ulp off at `H = 256`, `I = 128` and only **48** +(40 ulp RMS) at the R1 shape `H = 7168`, `I = 2048`, in both cases with no +error. It is the *least* loud of the operand mistakes catalogued here, and it +gets quieter at the larger geometry, so a caller cannot expect it to announce +itself. + +### Per token `t` and slot `j < K` + +With `A[t]` = the fp32 value of `hidden_states` **after** dividing out `g1` +(i.e. `code * blockscale / g1`), and `W_up[e]`, `W_gate[e]` (`[I, H]`), +`W_down[e]` (`[H, I]`) the fp32 values of this rank's dequantized NVFP4 +weights (`code * blockscale`, with the weight global scale `gw` divided out — +it lives in the scalars above): + +``` +gid = topk_ids[t, j] # a GLOBAL expert id +skip this slot unless off <= gid < off + E +e = gid - off # index into this rank's weights + +up = A[t] @ W_up[e].T +gate = A[t] @ W_gate[e].T + +# act_type 0 (SwiGlu), the only value certified here: +act = up * gate * sigmoid(gate) # fp32 + +act = nvfp4_requantize(act) # see below — the W4A4 step +y = act @ W_down[e].T +out[t] += topk_weights[t, j] * y # fp32 accumulation +``` + +and `out` is stored as bf16. Note the sigmoid is applied to the **gate** half +and `up` is the plain linear one. + +**The intermediate is NVFP4.** FC2 is an NVFP4 x NVFP4 GEMM, so the FC1 +epilogue quantizes its post-activation output before FC2 reads it. Per token +row and per **16 consecutive intermediate columns** (natural column order, +blocks starting at column 0): + +``` +sf = e4m3_round_to_nearest_even(clamp(g2 * max|act| over the block / 6, max 448)) +act' = e2m1_round_to_nearest_even(act * g2 / sf) * sf / g2 +``` + +— i.e. exactly what `torch.ops.trtllm.fp4_quantize(act, g2, 16, ...)` would +emit, dequantized. `6` is e2m1's largest magnitude, `448` e4m3's largest +finite value; e2m1 rounding is ties-to-even-code (`0.25 -> 0`, `0.75 -> 1`, +`1.25 -> 1`, `1.75 -> 2`, `2.5 -> 2`, `3.5 -> 4`, `5.0 -> 4`), saturating at +`±6`. This entry's test reads the requantized values straight out of the +kernel (a down projection set to the identity, `output2_scale_scalar = 1`) and +matches the formula above **bit-exactly on all 16384 elements**; on the same +data an unrounded (exact) block scale reproduces 23.5% of the elements and the +unquantized activation 0.0%. Modelling this layer without the +intermediate requantization deviates 43 bf16 ulp element-wise / 23 ulp RMS +from the kernel at `H = 256`, `I = 128`, and 28 / 23 at the R1 shape +`H = 7168`, `I = 2048` — it is a real perturbation of the layer output, not a +rounding detail, and the RMS figure barely moves with geometry. An +intermediate element below `1/24` of its block's largest +magnitude rounds to zero (e2m1's smallest nonzero code is `0.5` against a +block top of `6`). + +**FC1 half order is `[up | gate]`.** Before the kernel's row interleave (§ +*Preconditions*) the first `I` rows of the FC1 operand are the up projection +(trtllm's `w3`, HF's `up_proj`) and the last `I` rows the gate projection +(trtllm's `w1`, HF's `gate_proj`). Swapping them produces a plausible-looking, +entirely different result; nothing detects it — this entry's test checks that +a gate/up-swapped reference lands 263 ulp away at `H = 256`, `I = 128` and +278 ulp away at the R1 shape `H = 7168`, `I = 2048`. + +**Fusion boundary.** Inside the call: routing (when driven by logits), +permutation, both GEMMs with on-the-fly NVFP4 dequantization, the gated +activation, the NVFP4 requantization between the GEMMs, and the weighted +combine. Outside, and the caller's job: the router GEMM that produces the +logits and the routing itself when pre-routing (e.g. +`torch.ops.trtllm.noaux_tc_op`), **quantizing the bf16 hidden states to +NVFP4** (`torch.ops.trtllm.fp4_quantize`, § *Preconditions* — there is no +padding step inside this op), all weight preprocessing, any shared/dense +expert branch, the TP all-reduce or EP gather of this call's output, and the +residual add. + +**Nothing is renormalized inside.** `topk_weights` is used exactly as given +(it need not sum to 1). `routed_scaling_factor` is accepted but had no effect +on the certified path. + +**Output.** `do_finalize=True` with `output=None` returns a **one-element +list** whose tensor is a fresh contiguous `[T, H]` **bf16** result. With +`output` given the result is written there (every element overwritten, nothing +outside its rows touched) and the list holds an **empty `[0]` bf16 tensor** — +the caller must read its own buffer. `do_finalize=False` returns three +tensors; see § *Signature*. No input is mutated; two identical calls are +bitwise equal. + +## Signature + +```python +def fp4_block_scale_moe_runner( + routing_logits: Optional[torch.Tensor], + routing_bias: Optional[torch.Tensor], + hidden_states: torch.Tensor, + hidden_states_scale: torch.Tensor, + gemm1_weights: torch.Tensor, + gemm1_weights_scale: torch.Tensor, + gemm1_bias: Optional[torch.Tensor], + gemm1_alpha: Optional[torch.Tensor], + gemm1_beta: Optional[torch.Tensor], + gemm1_clamp_limit: Optional[torch.Tensor], + gemm2_weights: torch.Tensor, + gemm2_weights_scale: torch.Tensor, + gemm2_bias: Optional[torch.Tensor], + output1_scale_scalar: torch.Tensor, + output1_scale_gate_scalar: torch.Tensor, + output2_scale_scalar: torch.Tensor, + num_experts: int, + top_k: int, + n_group: Optional[int], + topk_group: Optional[int], + intermediate_size: int, + local_expert_offset: int, + local_num_experts: int, + routed_scaling_factor: Optional[float], + routing_method_type: int, + do_finalize: bool, + act_type: int = 0, + topk_weights: Optional[torch.Tensor] = None, + topk_ids: Optional[torch.Tensor] = None, + output: Optional[torch.Tensor] = None, + tune_max_num_tokens: int = 8192, + use_dp: bool = False, +) -> list[torch.Tensor] +``` + +Unlike the MXFP4 members of this family the signature carries **no +`valid_hidden_size` / `valid_intermediate_size`**: this op has no padded-vs- +true size distinction at all (§ *Preconditions*), and it gains the three scale +scalars and `do_finalize`. + +**`gemm1_bias`, `gemm1_alpha`, `gemm1_beta`, `gemm1_clamp_limit` and +`gemm2_bias` accept `None`** even though the registered schema spells them +`Tensor`, not `Tensor?`. The schema is generated from the Python custom-op's +annotations; the torch dispatcher turns a `None` argument into an undefined +tensor, the C++ entry point takes `std::optional` for exactly +these five (visible in the exported `FP4BlockScaleMoeRunner::run_moe` symbol), +and it arrives there as an empty optional. Verified at runtime: all five at +`None` are bitwise identical to their neutral values (zero bias, `alpha = 1`, +`beta = 0`, an effectively infinite clamp limit), and each slot is live — a +non-neutral value changes the result. A DeepSeek-V3 style MoE, which has no +expert bias and no clamp, passes `None` for all five. + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `routing_logits` | `None` — the only certified value; see *Notes* | — | — | — | +| `routing_bias` | `None`, or `[num_experts]` and inert | bf16 | contiguous | CUDA | +| `hidden_states` | `[T, H/2]` | **uint8** (2 E2M1 codes per byte, element `2i` in the low nibble) | contiguous | CUDA | +| `hidden_states_scale` | **1-D**, exactly `T * H/16` elements | **float8_e4m3fn** | contiguous, **linear** (row-major `[T, H/16]` flattened) — *not* swizzled | CUDA | +| `gemm1_weights` | `[E, 2*I, H/2]` | **uint8** | contiguous, pre-shuffled | CUDA | +| `gemm1_weights_scale` | `[E, 2*I, H/16]` | **float8_e4m3fn** | contiguous, pre-shuffled + 128x4 swizzled | CUDA | +| `gemm1_bias` / `gemm1_alpha` / `gemm1_beta` / `gemm1_clamp_limit` | `None` | — | — | — | +| `gemm2_weights` | `[E, H, I/2]` | **uint8** | contiguous, pre-shuffled | CUDA | +| `gemm2_weights_scale` | `[E, H, I/16]` | **float8_e4m3fn** | contiguous, pre-shuffled + 128x4 swizzled | CUDA | +| `gemm2_bias` | `None` | — | — | — | +| `output1_scale_scalar` | `[E]` (**local** experts) | **fp32** | contiguous | CUDA | +| `output1_scale_gate_scalar` | `[E]` | **fp32** | contiguous | CUDA | +| `output2_scale_scalar` | `[E]` | **fp32** | contiguous | CUDA | +| `num_experts` | scalar | int, the global routing space; `> top_k` | — | — | +| `top_k` | scalar | int, `0 < top_k < num_experts` | — | — | +| `n_group` / `topk_group` | `None`, `(1,1)`, or (with `routing_method_type = 2`) `(4,2)` / `(8,4)` — all inert | int | — | — | +| `intermediate_size` | scalar | int, **must equal `gemm1_weights.shape[1] // 2`** | — | — | +| `local_expert_offset` / `local_num_experts` | scalars | int, `E = local_num_experts >= 1` | — | — | +| `routed_scaling_factor` | `None`, `1.0`, `2.5` — all inert | float | — | — | +| `routing_method_type` | scalar | int: `0/1/2/4/5/6` all inert on this path | — | — | +| `do_finalize` | scalar | bool | — | — | +| `act_type` | scalar | int: **0** (SwiGlu) only | — | — | +| `topk_weights` | `[T, top_k]` | **bf16** | contiguous | CUDA | +| `topk_ids` | `[T, top_k]` | **int32** | contiguous | CUDA | +| `output` | `[T, H]` or `None` | bf16 | contiguous | CUDA | +| return, `do_finalize=True` | `[[T, H]]`, or `[[0]]` when `output` is given | bf16 | contiguous, newly allocated | same as `hidden_states` | +| return, `do_finalize=False` | see below | | | | + +`torch.ops.trtllm.noaux_tc_op` fed **bf16** router logits produces a directly +consumable `(topk_weights, topk_ids)` pair — bf16 weights and int32 ids, in +that order. Catalog membership of that op is `index.yaml`'s fact alone. + +### `do_finalize=False` + +Returns three tensors: + +0. `[P, H]` bf16 — one row per (token, slot) expert output **before** the + routing-weighted combine, in the kernel's internal permuted/padded order. + `P` depends on `T`, `top_k`, `num_experts` and an internal tile size; it is + always `>= T * top_k`. +1. `[T, top_k]` bf16 — **not written on the pre-routed path.** Observed to + hold uninitialized memory (values of order `1e24`, and NaN once combined). + It is the slot for the routing stage's own weights; do not read it. +2. `[T, top_k]` int32 — the expanded-index -> permuted-row map: token `t`'s + slot `j` expert output is row `out2[t, j]` of output 0. + +Recombining with the caller's own weights, +`out[t] = sum_j topk_weights[t, j] * out0[out2[t, j]]` accumulated in fp32 and +stored as bf16, reproduces the `do_finalize=True` result **bitwise** (verified +here). `output=` is rejected in this mode +(`out_tensor is only supported when do_finalize=true`). **A plain, non +expert-parallel target wants `do_finalize=True`**: the mode exists so an +all-to-all deployment can move the per-slot rows before combining them. + +## Metadata consumed + +None. The op reads no attention metadata, no KV cache, no registered layer and +no module state — every tensor and every scalar it uses is an argument. + +Two process-global caches sit behind it, neither of which changed the result +on any path measured here: + +- an **autotuner profiling cache** that picks the two GEMM tactics, keyed by + `(top_k, intermediate_size, local_num_experts, act_type)` together with the + call's input shapes and a token bucket (the row count rounded down to a + power of two, capped at 8192). Two consequences of that key were measured + here, by counting the profiling sweeps the tuner performed rather than + reading the key: **`local_num_experts` separates entries** — going from 256 + to 64 with everything else fixed forced a fresh sweep, so a cache warmed at + one expert-parallel window size says nothing about another — while + **`local_expert_offset` does not**: after warming the offset-0 64-wide + window, the offset-64 / 128 / 192 windows profiled *nothing*, reusing its + entries. Warming one window of a split therefore warms all four. + + **Cold is the state every call compared against a reference here ran in** — + the tuner then returns the fallback tactic. A serving engine is *not* cold: + the PyTorch runtime profiles this op during warm-up whenever + `enable_autotuner` is set (its default), so a deployed call runs on a tuned + tactic instead. Two direct cold-vs-warm comparisons were made, both by + warming with `tensorrt_llm._torch.autotuner.autotune()` and re-running the + identical operands. At `H = 2560`, `I = 1536`, top-6, `T = 8192`: **bitwise + identical** for `local_num_experts = 72` and all four 18-wide windows. At + the **R1 geometry** (`H = 7168`, `I = 2048`, top-8) the comparison is part + of this entry's test and is broader: warming wrote **28 cache entries** (14 + token buckets x the two `unique_id`s `(8, 2048, 256, 0)` and + `(8, 2048, 64, 0)`), every one of them a **non-fallback** tactic, and all + 85 (expert layout, token count) results — `local_num_experts = 256` plus + the four 64-wide windows, over the full certified `T` column — were + **bitwise unchanged**. *Which* tactics those are is not reproducible: the + tuner selects by measured time, so it is the one figure this test prints + that varies run to run — four runs of the comparison on the same device + recorded between 12 and 15 distinct tactics across the same 28 entries. + The bitwise result therefore held across four different tactic assignments, + which is stronger than one — but it is still a finite sample of tuner picks + on one machine, not a guarantee across every tactic the tuner might select. +- a cache of C++ `FP4BlockScaleMoERunner` objects keyed by `act_type`, each + owning a GPU workspace allocated on first use. + +## Preconditions + +### Sizes + +- **`H` must be a multiple of 256.** With the prepared block-scale layout + below, a hidden size that is only a multiple of 128 makes the kernel read + the FC1 weight scales for the wrong blocks: no error, 300-3400 bf16 ulp + wrong. Measured wrong at `H` = 384, 640 and 896; correct at 256, 512, 768, + 2560 and 7168. The wrapper asserts it. (`H` enters the FC1 GEMM's K axis; + the FC2 K axis, `I`, carries no such rule — every multiple of 64 worked.) +- `I` must be a multiple of **64**, so that `2*I` is a multiple of 128 and + `I/16` a multiple of 4 — the alignment the 128x4 scale swizzle needs on + both operands. Certified `I`: 64, 128, 192, 256, 512, 1536, 2048. +- There is **no padded-vs-true size distinction**. `gemm1_weights.shape[-1] * + 2`, `gemm2_weights.shape[1]` and the width of the returned output are all + the same `H`; a mismatch is rejected (`gemm2_weights_scale has incorrect dim + 1`). `intermediate_size` must equal `gemm1_weights.shape[1] // 2`; anything + else is rejected (`gemm2_weights_scale has incorrect dim 2`, or `No valid + config found for the given problem shape`). +- `0 < top_k < num_experts`. `top_k = 0` is rejected (`only supports + top_k<=32 && top_k>0`), `num_experts == top_k` too (`num_experts must be + greater than top_k`). Certified `top_k`: 1, 2, 4, 6, 8. +- `num_experts` is only the routing space; it need not equal + `local_num_experts`. Certified `num_experts`: 2, 3, 4, 8, 16, 72, 256; + certified `local_num_experts`: 2, 3, 4, 8, 16, 18, 64, 72, 256. Certified + `(local_expert_offset, local_num_experts)` pairs: `(0, E)` for each of + those `E`; `(4, 4)` of `num_experts = 8`; the four-way expert-parallel + split of 72 — `(0, 18)`, `(18, 18)`, `(36, 18)`, `(54, 18)`; and the + four-way split of 256 — `(0, 64)`, `(64, 64)`, `(128, 64)`, `(192, 64)` — + see § *Routing and expert ids*. +- `T >= 1`. Certified `T`: 1, 2, 3, 5, 7, 8, 16, 24, 32, 33, 40, 64, 128, + 256, 1024, 4096, 8192. `T` is the **row count of this call** — + `hidden_states.shape[0]`. It is not `tune_max_num_tokens`, which is an + autotuner bucket hint whose default is also 8192 and which was measured + inert (§ *Notes*); a call passing 8192 rows and a call passing 8 rows with + `tune_max_num_tokens = 8192` are different things, and only the first + certifies this row count. 8192 is trtllm's default `max_num_tokens` + (`llm_args.py`), so it is the widest prefill a stock engine hands this op + in one unchunked call, and it is the top of this list: **a caller whose own + token count exceeds 8192 must chunk into calls this column covers.** The + whole column is certified at the **DeepSeek-R1 routed geometry** + (`H = 7168`, `I = 2048`, `num_experts = 256`, `top_k = 8`, `act_type = 0`, + `do_finalize = True`, pre-routed), for the full 256-expert stack *and* + independently for each of the four 64-wide expert-parallel windows. The + DeepSeek-V3-Lite routed geometry (`H = 2560`, `I = 1536`, top-6) is + certified over a subset of it topping out at the same 8192 rows, for + `local_num_experts = 72` at offset 0 and for the four 18-wide windows — see + § *Routing and expert ids*. Other geometries in the *Certified* lists above + cover only the smaller counts. + +### Activation preparation — NVFP4, linear scales + +`hidden_states` / `hidden_states_scale` are exactly the two outputs of + +```python +# global sf_vec ue8m0 swizzled +data, sf = torch.ops.trtllm.fp4_quantize(x, g1, 16, False, False) +``` + +on the bf16 `[T, H]` hidden states, with `sf` re-viewed as `float8_e4m3fn`. +Concretely, whatever produces them: + +- `hidden_states` is 2-D **uint8** (`hidden_states must be byte` — a + `float8_e4m3fn` view of the same bytes is rejected), contiguous, CUDA, of + width exactly `gemm1_weights.shape[-1]`. +- `hidden_states_scale` is **1-D** (`hidden_states_scale must be 1D`), + **float8_e4m3fn** (`must be fp8` — a uint8 view is rejected), contiguous, + CUDA, and holds exactly `T * H/16` elements (`hidden_states_scale has + incorrect size` otherwise). Element `t * (H/16) + b` is the e4m3 scale of + columns `[16b, 16b+16)` of token `t`. +- The **128x4 swizzled** scale order is a silent wrong answer whenever it + happens to have the same element count (`T % 128 == 0`, `H/16 % 4 == 0` — + true at any graph-friendly batch size) — see *Notes*. Pass + `is_sf_swizzled_layout=False`. +- This op never pads: `H` is used as given, so the quantizer is called on the + full `[T, H]` hidden with no `F.pad` before it. + +### Weight preparation — the caller owns all of it + +Start from the checkpoint's per-expert NVFP4 tensors, laid out **row-major +over the output dim**, K packed two E2M1 codes per byte with the **low nibble +holding the even K index**, and one E4M3 byte per 16 consecutive K elements: + +``` +up (w3 / up_proj) packed [E, I, H/2] uint8 scale [E, I, H/16] e4m3 +gate (w1 / gate_proj) packed [E, I, H/2] uint8 scale [E, I, H/16] e4m3 +down (w2 / down_proj) packed [E, H, I/2] uint8 scale [E, H, I/16] e4m3 +``` + +Then, **per expert**: + +1. **FC1 concat.** Stack the two halves on the row axis as `[up ; gate]`, + giving `[2*I, ...]` for both the packed bytes and the scale bytes. +2. **FC1 row permute** = interleave, then block shuffle: + - interleave: destination row `2i` takes the up half's row `i`, + destination row `2i+1` takes the gate half's row `i`; + - block shuffle (also applied to FC2, which skips the interleave): + within each aligned block of 32 rows, source row `4u + v` + (`0 <= u < 8`, `0 <= v < 4`) moves to destination row `8v + u`. + + Apply the **same** permutation to the weight bytes and the scale bytes, so + row `i` of both still describes the same output channel. Skipping it is a + silent wrong answer (measured 508 ulp at `H = 256`, `I = 128`; 432 at + `H = 7168`, `I = 2048`). +3. **Scale swizzle.** After the row permute, each expert's scale matrix + `[M, C]` (`M % 128 == 0`, `C % 4 == 0`) is rewritten into the trtllm-gen + 128x4 layout: the byte at `(m, c)` moves to flat offset + + ``` + (m // 128) * 512 * (C // 4) + (c // 4) * 512 + + (m % 32) * 16 + ((m % 128) // 32) * 4 + (c % 4) + ``` + + and the flat result is handed to the kernel with its nominal `[E, M, C]` + shape. Weight bytes are **not** swizzled — only scales. Skipping it is a + silent wrong answer (measured 408 ulp at `H = 256`, `I = 128`; 357 at + `H = 7168`, `I = 2048`). + +**No padding is needed anywhere at either certified routed geometry**, and the +swizzle's alignment must be checked per geometry rather than assumed — it is a +property of `H` and `I`, not of the family: + +| geometry | FC1 weights + scales | FC2 weights + scales | `2I % 128`, `H/16 % 4`, `H % 128`, `I/16 % 4` | +|---|---|---|---| +| `H = 2560`, `I = 1536` | `[E, 3072, 1280]` + `[E, 3072, 160]` | `[E, 2560, 768]` + `[E, 2560, 96]` | 0, 0, 0, 0 | +| `H = 7168`, `I = 2048` | `[E, 4096, 3584]` + `[E, 4096, 448]` | `[E, 7168, 1024]` + `[E, 7168, 128]` | 0, 0, 0, 0 | + +At the R1 shapes `torch.ops.trtllm.block_scale_interleave` was checked here to +return exactly `E*M*C` bytes for an `[E, M, C]` scale tensor — measured at +`4096 x 448` and at `7168 x 128` — so the swizzled buffer reshapes straight +back to its nominal `[E, M, C]` shape with nothing appended, and the operand +handed to the kernel is the same size as the one the checkpoint carries. + +The two trtllm helpers `torch.ops.trtllm.shuffle_matrix(x, perm)` (a plain row +gather, `out[i] = x[perm[i]]`) and `torch.ops.trtllm.block_scale_interleave(x)` +(the 128x4 swizzle over a `[E, M, C]` uint8 tensor, returning a flat buffer of +`E * pad_up(M, 128) * pad_up(C, 4)` bytes) produce byte-identical results to +steps 2 and 3; this entry's test asserts that equivalence. + +### Layout and dtypes + +- **Every tensor argument must be contiguous.** A strided view is accepted + silently and read as if dense — see *Notes*. The wrapper asserts this. +- `topk_ids` is `[T, top_k]` **int32** (`topk_ids must be int` — int64 is + rejected) and `topk_weights` is `[T, top_k]` **bf16** (`topk_weights must be + bfloat16` — fp32 and fp16 are both rejected). The two must be given together + (`routing_logits or (topk_ids and topk_weights) must be provided`). +- Both weight-scale tensors must be **float8_e4m3fn** (a uint8 view is + rejected with `must be fp8`). +- The three scale scalars must be **fp32** (`must be float`) with exactly + **`local_num_experts`** elements — not `num_experts` (`has incorrect dim + 0`). +- `output`, when given, must be a contiguous CUDA bf16 tensor of shape exactly + `[T, H]` (`out_tensor must be bfloat16`, `out_tensor dim0 must match + num_tokens`). A leading row-slice of a taller contiguous buffer is valid — + rows past `T` are left bitwise untouched. + +### Routing and expert ids + +- Expert ids outside `[local_expert_offset, local_expert_offset + + local_num_experts)` — including negative ids and ids `>= num_experts` — are + silently dropped; the token's other slots still combine normally. Verified + two ways: a `local_expert_offset = 4`, `local_num_experts = 4` window of + `num_experts = 8` against a reference that masks the same slots, and ids of + `-1`, `num_experts` and `num_experts + 5` all giving the same output as a + dropped slot. +- **The window is a range test plus an index shift, not a validated + partition.** Local weight slot `i` answers for global id + `local_expert_offset + i`; nothing checks the window against + `num_experts`. `local_expert_offset + local_num_experts > num_experts` is + **accepted** — measured at `(60, 18)` of `num_experts = 72`, where the + result matched a reference mapping ids 60..71 onto local slots 0..11 and + the surplus six local slots were simply never addressed — and a window + entirely past the routing space (`local_expert_offset = 72`) returns an + all-zero output. Making the windows tile the routing space is the caller's + arithmetic, not something this op enforces. +- **The four-way expert-parallel split of 72 experts is certified**: + `num_experts = 72` on every call (the routing space is not sharded), + `top_k = 6`, `local_num_experts = 18`, `local_expert_offset` 0 / 18 / 36 / + 54, `H = 2560`, `I = 1536`, `act_type = 0`, `do_finalize = True`, + pre-routed, at T = 1, 8, 256, 1024, 4096 and 8192. Per window — with + `gemm1_weights`, `gemm1_weights_scale`, `gemm2_weights`, + `gemm2_weights_scale` and all three scale scalars sliced to + `[off, off+18)` along the expert axis and made contiguous — the result + matches a torch reference that masks every slot outside the window (worst + 9.97 ulp element-wise, 0.74 ulp RMS; at T = 8192 alone, 9.97 and 0.72), + and every token with **no** slot in the window comes back **bitwise + zero** — at T = 8192 that is ~1300 rows per window (1265-1406 measured + across the four windows: top-6 of 72 lands nothing in a given 18-wide + window for about a sixth of the batch). The four bf16 outputs summed + reproduce the 72-expert result under the same gate (worst 4.86 ulp + element-wise, 0.88 ulp RMS against the fp32 reference; at T = 8192 alone, + 4.86 and 0.84); they are **not** bitwise equal to a single 72-expert + call — measured up to 2.0 ulp element-wise apart at every certified token + count, 8192 included, because each window rounds its own partial to bf16 + before the add. Leaving any one window out of that sum lands >= 124 ulp + RMS away (>= 128 at T = 8192). When every routed slot of a batch happens + to fall inside one window, that window's call **is** bitwise equal to the + 72-expert call and the other three return exactly zero. +- **The four-way expert-parallel split of 256 experts is certified**, the + DeepSeek-R1 routed layout: `num_experts = 256` on every call (the routing + space is not sharded), `top_k = 8`, `local_num_experts = 64`, + `local_expert_offset` 0 / 64 / 128 / 192, `H = 7168`, `I = 2048`, + `act_type = 0`, `do_finalize = True`, pre-routed, over the **whole** + certified `T` column (1 through 8192, § *Sizes*). Sliced the same way — all + four weight/scale operands and all three scale scalars restricted to + `[off, off+64)` along the expert axis and made contiguous — each window + matches a torch reference masking every slot outside it (worst 9.62 ulp + element-wise, 0.74 ulp RMS; at T = 8192 alone, 6.35 and 0.73), and every + token with **no** slot in the window comes back **bitwise zero** — at + T = 8192 that is 764-804 rows per window, i.e. top-8 of 256 lands nothing + in a given 64-wide window for ~10% of the batch (`(3/4)^8`). The four bf16 + outputs summed reproduce the 256-expert result under the same gate (worst + 5.02 ulp element-wise, 0.85 ulp RMS against the fp32 reference; identical + at T = 8192 alone); as with the 18-wide split they are **not** bitwise + equal to a single 256-expert call — up to 2.0 ulp element-wise apart at + every certified token count, because each window rounds its own partial to + bf16 before the add. Leaving any one window out of that sum lands >= 125 + ulp RMS away (>= 128 at T = 8192); a window left at + `local_expert_offset = 0`, a window fed the neighbouring window's weights, + and all four windows issued at offset 0 land 310-363 ulp RMS away. + When every routed slot of a batch falls inside one window, that + window's call **is** bitwise equal to the 256-expert call and the other + three return exactly zero. +- **A repeated expert id inside one token's row is not reliably + deduplicated.** With one geometry it contributed once (the first slot's + weight) at 16 tokens and *twice* at 24 tokens. Supply distinct ids per row; + every routing op in this build does. +- `n_group > 1` requires `routing_method_type = 2` (`Routing kernel with + groups implies DeepSeekV3 routing method`), but on the pre-routed path the + whole grouped-routing configuration is inert: `(n_group, topk_group)` of + `(1,1)`, `(4,2)` and `(8,4)` under `routing_method_type = 2` are all bitwise + identical to `None/None` under `routing_method_type = 1`. +- `routing_method_type = 3` (Llama4) is rejected for `top_k > 1` (`Current + routing kernel (no groups, Llama4) only supports top_k=1`). + +A caller violating none of the above gets the result described under +*Semantics*, to within the bound in *Notes*. + +## Notes + +- **Numerics.** Against a native-torch reference that consumes bit-identical + operands (exact NVFP4 weight and activation dequantization, fp32 GEMM + accumulation, the NVFP4 requantization of the FC1 output modelled exactly, + fp32 combine), the kernel's worst element-wise deviation over every + configuration in this entry's test is **9.97 ulp of the token row's largest + magnitude** (bf16 ulp = 2^-8; worst case: an 18-wide expert-parallel window + of 72 experts, 8192 tokens, `H = 2560`, `I = 1536`), against the test's + 16-ulp element gate; a full-window call sits at 4.75 (3.20 at T = 8192). + Its worst relative RMS deviation over the whole test file is + **0.88 ulp** (those four windows summed), against a 2-ulp aggregate gate. + The **R1 routed geometry** (`H = 7168`, `I = 2048`, 256 experts, top-8) + sits just inside both, over the whole certified token column: 9.62 elt / + 0.74 RMS for a 64-wide window, 4.68 / 0.70 for the 256-expert stack, 5.02 / + 0.85 for the four windows summed. + The element-wise outliers are single intermediate values landing on opposite + sides of an e2m1 rounding boundary — an e2m1 step is coarse (2 mantissa + bits), so one flipped element moves an output row by a visible fraction. + Two things inflate the element figure and not the RMS one: it is an extreme + order statistic over every element compared, and it is normalized by the + reference row's own magnitude, so a window carrying a quarter of the routed + slots divides a similar absolute error by a smaller row. At `H = 2560`, + re-drawn over 144 independent (seed, window) comparisons at T = 1024 / 4096 + it reached **12.0 ulp** while the RMS figure stayed at 0.86; re-drawn again + over 16 fresh seeds at **T = 8192** (16 full-stack + 64 window comparisons) + it reached **16.33 ulp** for a window (p90 9.82, mean 7.15) and 10.41 for + the full stack, while the RMS figure was unchanged from T = 4096 (window + 0.74, full 0.70, four windows summed 0.85). So the element metric's upper + tail *crosses* the 16-ulp gate at the top of the certified token range — it + is a max over `T * H` elements and therefore grows with `T` at constant + accuracy, which is a property of the statistic, not of the kernel. At + `H = 7168` / `I = 2048` the same re-draw (12 fresh weight + activation + seeds, 12 full-stack + 48 window comparisons per token count) has a + **lower** tail: max **12.74** for a window at T = 8192 (p90 8.40, mean + 6.73) and 12.15 at T = 4096, 7.11 / 6.42 for the 256-expert stack, with the + RMS figures flat at 0.74 (window) / 0.70 (full) / 0.86 (sum) at both token + counts — so the element gate keeps ~1.3x headroom on the *distribution* + there, not ~1.0x. The shipped test is seeded and deterministic — two + independent processes on two devices printed byte-identical numbers — and + its T = 8192 draws land at 9.97 / 3.20 (`H = 2560`) and 6.35 / 4.11 + (`H = 7168`). **The RMS figure is the scale-free one and the one a caller + must size its own tolerance from**: it is identical at 4096 and 8192 at both + geometries, and every wrong variant this entry's test constructs sits + 22.5-370 ulp RMS away — 40-370 for every layout, scalar-role or + expert-window mistake (the 40 is an `output1_*` scalar-role swap at the R1 + shape; the same swap is 124 at `H = 256`), and ~23 for the one variant that + is not a mistake in the operands at all (a reference that skips the + FC1-output requantization). +- **Only the pre-routed entry point is certified here.** `routing_logits` was + left at `None` on every certified call and the routing was done outside (see + `torch.ops.trtllm.noaux_tc_op` for the DeepSeek-V3 formula). The in-kernel + routing path — `routing_logits` given, `topk_ids`/`topk_weights` `None`, + `routing_method_type` selecting softmax / renormalize / DeepSeekV3 — is + **not certified by this entry**; nor is `routing_bias`, which was only + observed to be inert on the pre-routed path. +- **A swizzled activation-scale buffer is a silent wrong answer.** + `fp4_quantize(x, g, 16, False, True)` returns the same data bytes and a + scale buffer of `pad_up(T,128) * pad_up(H/16,4)` bytes in 128x4 order. + Whenever that count coincides with the linear one (`T` a multiple of 128 — + exactly the CUDA-graph batch sizes) the size check passes and the kernel + reads scales for the wrong blocks: measured 276 ulp element-wise off at + `T = 128`, `H = 256`. Nothing in the metadata distinguishes the two layouts, + so no guard can catch it — pass `is_sf_swizzled_layout=False`, in a named + constant. +- **Non-contiguous tensors are a silent wrong answer.** The kernel takes raw + data pointers and assumes a dense row-major layout. A strided view of + `hidden_states`, `topk_weights` or any weight/scale operand is accepted + without complaint and reads the wrong elements. The wrapper asserts + contiguity on every tensor argument. +- **`act_type` other than 0 is not certified.** `1` (Relu2) and `2` (Silu) — + the two non-gated activations, which take an `[E, I, H/2]` FC1 operand + rather than `[E, 2*I, H/2]` — run without error when handed a gated-shaped + operand and return something different; so does the undefined value `3`. + Only `0` (SwiGlu) is characterized here. +- **Inert on the certified path** (bitwise identical results): + `routing_method_type` `0/1/2/4/5/6`; `n_group`/`topk_group` as above; + `routed_scaling_factor` `None`/`1.0`/`2.5`; `tune_max_num_tokens` `8192` + and `128` (an autotuner bucket cap); `use_dp` `False` and `True` (an + autotuner token-bucket deflation hint); a zero `routing_bias`. +- **Multi-rank tensor parallelism is not exercised.** `local_expert_offset` / + `local_num_experts` (expert parallelism) **are** certified, including the + four-way splits of 72 and of 256 experts that a 4-rank EP deployment uses — + but every certified call ran in **one process on one device**, the four + windows issued in sequence. This op reads no rank state, holds no + communicator and + launches no collective, so a window is an argument pair rather than a + multi-rank behaviour; summing the per-rank outputs is the caller's + all-reduce and is outside this entry. TP, which instead splits `I` across + ranks and reduces this call's output afterwards, is not certified. +- sm_100 only. The receipt covers sm_100 (B200), the only arch available here; + the installed build carries the assertion `Only SM100f is supported by FP4 + block scale MOE` in `libth_common.so`, so other Blackwell variants are + expected to raise rather than compute something wrong. +- The registered fake (meta) function disagrees with the kernel on **both** + dimensions of the `do_finalize=False` first output: it returns + `[1152, 128]` where the kernel returned `[256, 256]` (`T = 16`, + `top_k = 2`, `num_experts = 8`, `H = 256`). The rows come from a fixed + internal tile of 128; the columns are the **packed** width, because the + helper unpacks only under + `isinstance(hidden_states, Fp4QuantizedTensor)` and the op's schema hands + it a plain `Tensor`, so that branch can never be taken. The + `do_finalize=True` fake path doubles unconditionally and is correct. Eager + callers are unaffected; this is an upstream defect in this build, not + behaviour to rely on. +- Sibling ops in this build address the same job with other operand dtypes: + `torch.ops.trtllm.bf16_mxe2m1_block_scale_moe_runner` (bf16 activations, + MXFP4 weights), `torch.ops.trtllm.mxe4m3_mxe2m1_block_scale_moe_runner` + (MXFP8 activations, MXFP4 weights), + `torch.ops.trtllm.e4m3_mxe2m1_block_scale_moe_runner`, + `torch.ops.trtllm.fp8_block_scale_moe_runner` and + `torch.ops.trtllm.fp8_fp4_block_scale_moe_runner` (fp8 activations, NVFP4 + weights). The MXFP4 members take an `[E, 2*I, H/32]` **uint8** UE8M0 scale + and carry `valid_hidden_size` / `valid_intermediate_size` parameters this op + does not have, so their prepared weights are **not** interchangeable with + this one's. Catalog membership is `index.yaml`'s fact alone. diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner.py b/tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner.py new file mode 100644 index 000000000000..b2ec6994853f --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner.py @@ -0,0 +1,129 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""NVFP4-weight / NVFP4-activation mixture-of-experts layer: caller-supplied +top-k + grouped FC1 GEMM + gated activation + NVFP4 requantization + grouped +FC2 GEMM + routing-weighted combine, in one trtllm-gen call.""" + +from typing import Optional + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + + +def fp4_block_scale_moe_runner( + routing_logits: Optional[torch.Tensor], + routing_bias: Optional[torch.Tensor], + hidden_states: torch.Tensor, + hidden_states_scale: torch.Tensor, + gemm1_weights: torch.Tensor, + gemm1_weights_scale: torch.Tensor, + gemm1_bias: Optional[torch.Tensor], + gemm1_alpha: Optional[torch.Tensor], + gemm1_beta: Optional[torch.Tensor], + gemm1_clamp_limit: Optional[torch.Tensor], + gemm2_weights: torch.Tensor, + gemm2_weights_scale: torch.Tensor, + gemm2_bias: Optional[torch.Tensor], + output1_scale_scalar: torch.Tensor, + output1_scale_gate_scalar: torch.Tensor, + output2_scale_scalar: torch.Tensor, + num_experts: int, + top_k: int, + n_group: Optional[int], + topk_group: Optional[int], + intermediate_size: int, + local_expert_offset: int, + local_num_experts: int, + routed_scaling_factor: Optional[float], + routing_method_type: int, + do_finalize: bool, + act_type: int = 0, + topk_weights: Optional[torch.Tensor] = None, + topk_ids: Optional[torch.Tensor] = None, + output: Optional[torch.Tensor] = None, + tune_max_num_tokens: int = 8192, + use_dp: bool = False, +) -> list[torch.Tensor]: + """Run one NVFP4-weight MoE layer over NVFP4 (e2m1 + e4m3 block scale) activations. + + With `do_finalize=True` returns a one-element list holding the combined + `[num_tokens, hidden]` bf16 result — or an empty `[0]` tensor when `output` + was given, the result having been written there. With `do_finalize=False` + returns three tensors (per-slot expert rows, an unwritten scale buffer, and + the expanded-index -> permuted-row map). + """ + # Pure-metadata guard: the kernel takes raw data pointers and assumes a + # dense row-major layout for every tensor. A strided view is accepted + # without complaint and silently reads (or writes) the wrong elements — + # observed on this machine for hidden_states, topk_weights and the weight + # operands. + for name, tensor in ( + ("routing_logits", routing_logits), + ("routing_bias", routing_bias), + ("hidden_states", hidden_states), + ("hidden_states_scale", hidden_states_scale), + ("gemm1_weights", gemm1_weights), + ("gemm1_weights_scale", gemm1_weights_scale), + ("gemm1_bias", gemm1_bias), + ("gemm1_alpha", gemm1_alpha), + ("gemm1_beta", gemm1_beta), + ("gemm1_clamp_limit", gemm1_clamp_limit), + ("gemm2_weights", gemm2_weights), + ("gemm2_weights_scale", gemm2_weights_scale), + ("gemm2_bias", gemm2_bias), + ("output1_scale_scalar", output1_scale_scalar), + ("output1_scale_gate_scalar", output1_scale_gate_scalar), + ("output2_scale_scalar", output2_scale_scalar), + ("topk_weights", topk_weights), + ("topk_ids", topk_ids), + ("output", output), + ): + if tensor is not None: + assert tensor.is_contiguous(), ( + f"{name} must be contiguous; a strided view is read as if dense " + "and silently produces wrong results" + ) + # Pure-metadata guard: with the block-scale layout this entry contracts + # (row shuffle + 128x4 scale swizzle), a hidden size that is not a multiple + # of 256 makes the kernel read the FC1 weight scales for the wrong blocks. + # Observed on this machine at hidden 384 / 640 / 896: no error, 300-3400 + # bf16 ulp wrong. + assert gemm1_weights.shape[-1] * 2 % 256 == 0, ( + "hidden size (gemm1_weights.shape[-1] * 2) must be a multiple of 256; " + "the prepared block-scale layout is silently misread otherwise" + ) + return torch.ops.trtllm.fp4_block_scale_moe_runner( + routing_logits, + routing_bias, + hidden_states, + hidden_states_scale, + gemm1_weights, + gemm1_weights_scale, + gemm1_bias, + gemm1_alpha, + gemm1_beta, + gemm1_clamp_limit, + gemm2_weights, + gemm2_weights_scale, + gemm2_bias, + output1_scale_scalar, + output1_scale_gate_scalar, + output2_scale_scalar, + num_experts, + top_k, + n_group, + topk_group, + intermediate_size, + local_expert_offset, + local_num_experts, + routed_scaling_factor, + routing_method_type, + do_finalize, + act_type, + topk_weights=topk_weights, + topk_ids=topk_ids, + output=output, + tune_max_num_tokens=tune_max_num_tokens, + use_dp=use_dp, + ) diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner_test.py b/tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner_test.py new file mode 100644 index 000000000000..a0c2ef90fe67 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner_test.py @@ -0,0 +1,1926 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the fp4_block_scale_moe_runner catalog entry.""" + +import torch + +from tensorrt_llm._torch.autotuner import AutoTuner, autotune + +from .fp4_block_scale_moe_runner import fp4_block_scale_moe_runner as moe + +assert torch.cuda.is_available(), "fp4_block_scale_moe_runner requires a CUDA device" + +DEV = "cuda" +# The reference GEMMs must be true fp32; TF32 would leave the reference with +# 10 mantissa bits, coarser than the bf16 output it is meant to bound. +torch.backends.cuda.matmul.allow_tf32 = False + +# Relative distance between neighbouring bf16 values (1 + 7 stored mantissa +# bits -> one ulp is 2^-8 of the binade top). +ULP = 2.0**-8 + +# The 16 e2m1 code points, in code order: sign bit 3, exponent bits 2:1, +# mantissa bit 0. +E2M1 = torch.tensor( + [ + 0.0, + 0.5, + 1.0, + 1.5, + 2.0, + 3.0, + 4.0, + 6.0, + -0.0, + -0.5, + -1.0, + -1.5, + -2.0, + -3.0, + -4.0, + -6.0, + ], + dtype=torch.float32, + device=DEV, +) +SV = 16 # NVFP4 block size: one e4m3 scale per 16 elements along K +E2M1_MAX = 6.0 # largest e2m1 magnitude +E4M3_MAX = 448.0 # largest finite e4m3 magnitude + + +# ── NVFP4 arithmetic (pure torch) ───────────────────────────────────────── + + +def _e2m1_rne(a: torch.Tensor) -> torch.Tensor: + """Round-to-nearest-even onto the e2m1 grid, saturating at +-6. + + The grid step is 0.5 below 2, 1 below 4 and 2 above, i.e. the binade step + of e2m1's three exponents (the subnormal binade shares the step of the + first normal one). `torch.round` is banker's rounding, which reproduces + the format's ties-to-even-code rule: 0.25 -> 0, 0.75 -> 1, 1.25 -> 1, + 1.75 -> 2, 2.5 -> 2, 3.5 -> 4, 5.0 -> 4. + """ + step = torch.where(a.abs() < 2.0, 0.5, torch.where(a.abs() < 4.0, 1.0, 2.0)) + return torch.sign(a) * torch.clamp(torch.round(a.abs() / step) * step, max=E2M1_MAX) + + +def _q_nvfp4(x: torch.Tensor, g: float) -> torch.Tensor: + """NVFP4 quantize-dequantize of `x` under global scale `g`. + + Per 16 consecutive columns: `sf = e4m3(g * blockmax / 6)`, + `data = e2m1(x * g / sf)`. Returns `data * sf`, i.e. the value a + downstream block-scaled MMA multiplies — `g` times the reconstruction of + `x`. Pinned bit-exactly against the kernel's FC1 epilogue by + `test_intermediate_is_nvfp4_requantized`. + """ + rows, cols = x.shape + b = x.reshape(rows, cols // SV, SV) + sf = ( + (g * b.abs().amax(dim=-1, keepdim=True) / E2M1_MAX) + .clamp(max=E4M3_MAX) + .to(torch.float8_e4m3fn) + .float() + ) + out_scale = torch.where(sf == 0, torch.zeros_like(sf), g / sf) + return (_e2m1_rne(b * out_scale) * sf).reshape(rows, cols) + + +def _rand_nvfp4(e: int, n: int, k: int, gen: torch.Generator): + """Random NVFP4 expert stack. + + Returns `(packed [E, N, K/2] uint8, scales [E, N, K/16] float8_e4m3fn, + codes [E, N, K] uint8)`. Codes and (power-of-two, hence exact) e4m3 scales + are drawn directly, so the fp32 value of every weight is exact — the + reference never has to model a quantizer. + """ + codes = torch.randint(0, 16, (e, n, k), dtype=torch.uint8, device=DEV, generator=gen) + exps = torch.randint( + -2, 1, (e, n, k // SV), device=DEV, generator=gen, dtype=torch.int32 + ).float() + sf = torch.exp2(exps).to(torch.float8_e4m3fn) + packed = (codes[..., 0::2] | (codes[..., 1::2] << 4)).contiguous() + return packed, sf, codes + + +def _dequant(codes_e: torch.Tensor, sf_e: torch.Tensor) -> torch.Tensor: + """One expert's `[N, K]` fp32 weight from its codes and e4m3 block scales.""" + return E2M1[codes_e.long()] * sf_e.float().repeat_interleave(SV, dim=1) + + +def _unpack_activation(data: torch.Tensor, sf: torch.Tensor, hidden: int, g: float): + """Exact fp32 value of an NVFP4 activation pair, undoing the global scale. + + `data` is `[T, hidden/2]` uint8 (element 2i in the low nibble), `sf` a 1-D + linear `[T * hidden/16]` e4m3 buffer. + """ + rows = data.shape[0] + codes = torch.stack([data & 0x0F, data >> 4], dim=-1).reshape(rows, hidden).long() + scale = sf.view(torch.float8_e4m3fn).reshape(rows, hidden // SV).float() + return E2M1[codes] * scale.repeat_interleave(SV, dim=1) / g + + +def _rand_activation(rows: int, hidden: int, g: float, gen: torch.Generator): + """Random NVFP4 activation in the layout this op consumes. + + Returns `(data [rows, hidden/2] uint8, sf [rows * hidden/16] uint8, x + [rows, hidden] fp32)`, `x` being the exact dequantization. Codes and + (power-of-two) e4m3 scales are drawn directly, so no quantizer is modelled. + """ + codes = torch.randint(0, 16, (rows, hidden), dtype=torch.uint8, device=DEV, generator=gen) + exps = torch.randint( + -2, 1, (rows, hidden // SV), device=DEV, generator=gen, dtype=torch.int32 + ).float() + sf = torch.exp2(exps).to(torch.float8_e4m3fn) + data = (codes[:, 0::2] | (codes[:, 1::2] << 4)).contiguous() + sf_flat = sf.reshape(-1).view(torch.uint8).contiguous() + x = E2M1[codes.long()] * sf.float().repeat_interleave(SV, dim=1) / g + return data, sf_flat, x + + +# ── kernel weight layout (pure torch) ───────────────────────────────────── + + +def _blk32_perm(m: int) -> torch.Tensor: + """Gather index of the 32-row block shuffle: within each block of 32 rows, + source row `4u + v` lands at destination row `8v + u`.""" + assert m % 32 == 0 + j = torch.arange(32) + dst = (j % 4) * 8 + j // 4 + idx = torch.empty(32, dtype=torch.long) + idx[dst] = j + return (idx.repeat(m // 32) + torch.arange(m // 32).repeat_interleave(32) * 32).to(DEV) + + +def _gate_interleave_perm(m: int) -> torch.Tensor: + """Gather index that interleaves the `[up | gate]` halves of a `2*I` row + stack into `up0, gate0, up1, gate1, ...`.""" + p = torch.empty(m, dtype=torch.long) + p[0::2] = torch.arange(0, m // 2) + p[1::2] = torch.arange(m // 2, m) + return p.to(DEV) + + +def _swizzle_scales(s: torch.Tensor) -> torch.Tensor: + """Per-expert 128x4 block-scale swizzle, result viewed back as `[E, M, C]`. + + Flat destination of scale `(e, m, c)` is + `e*M*C + (m//128)*512*(C//4) + (c//4)*512 + (m%32)*16 + ((m%128)//32)*4 + (c%4)`. + """ + e, m, c = s.shape + assert m % 128 == 0 and c % 4 == 0 + v = s.reshape(e, m // 128, 4, 32, c // 4, 4) + v = v.permute(0, 1, 4, 3, 2, 5) + return v.reshape(e, m, c).contiguous() + + +def _fc1_perm(rows: int) -> torch.Tensor: + return _gate_interleave_perm(rows)[_blk32_perm(rows)] + + +def _prep_fc1(up_p, gt_p, up_s, gt_s, swizzle=True, shuffle=True): + """Concat as `[up | gate]`, interleave + block-shuffle rows, swizzle scales.""" + w = torch.cat([up_p, gt_p], dim=1) + s = torch.cat([up_s, gt_s], dim=1) + if shuffle: + perm = _fc1_perm(w.shape[1]) + w = torch.index_select(w, 1, perm) + s = torch.index_select(s, 1, perm) + s = s.view(torch.uint8) + s = _swizzle_scales(s) if swizzle else s.contiguous() + return w.contiguous(), s.view(torch.float8_e4m3fn).contiguous() + + +def _prep_fc2(dn_p, dn_s, swizzle=True, shuffle=True): + """Block-shuffle the down projection's rows, swizzle its scales.""" + w, s = dn_p, dn_s + if shuffle: + perm = _blk32_perm(w.shape[1]) + w = torch.index_select(w, 1, perm) + s = torch.index_select(s, 1, perm) + s = s.view(torch.uint8) + s = _swizzle_scales(s) if swizzle else s.contiguous() + return w.contiguous(), s.view(torch.float8_e4m3fn).contiguous() + + +def _build(num_experts: int, hidden: int, inter: int, seed: int): + """Build one MoE layer: kernel-ready tensors plus the reference operands.""" + gen = torch.Generator(device=DEV).manual_seed(seed) + up_p, up_s, up_c = _rand_nvfp4(num_experts, inter, hidden, gen) + gt_p, gt_s, gt_c = _rand_nvfp4(num_experts, inter, hidden, gen) + dn_p, dn_s, dn_c = _rand_nvfp4(num_experts, hidden, inter, gen) + w1, s1 = _prep_fc1(up_p, gt_p, up_s, gt_s) + w2, s2 = _prep_fc2(dn_p, dn_s) + args = dict( + gemm1_weights=w1, + gemm1_weights_scale=s1, + gemm2_weights=w2, + gemm2_weights_scale=s2, + intermediate_size=inter, + ) + ref = dict( + up=(up_c, up_s), + gate=(gt_c, gt_s), + down=(dn_c, dn_s), + raw=(up_p, gt_p, up_s, gt_s, dn_p, dn_s), + hidden=hidden, + inter=inter, + num_experts=num_experts, + gen=gen, + ) + return args, ref + + +def _scalars(num_local: int, g1: float, g2: float): + """The three per-expert fp32 scalars, for weight global scales of 1. + + `output1_scale_gate_scalar = alpha1` dequantizes the FC1 accumulator; + `output1_scale_scalar = g2 * alpha1` additionally applies the FC2-input + global scale to the linear half; `output2_scale_scalar = alpha2` + dequantizes the FC2 accumulator. + """ + a1 = torch.full((num_local,), 1.0 / g1, dtype=torch.float32, device=DEV) + a2 = torch.full((num_local,), 1.0 / g2, dtype=torch.float32, device=DEV) + return a1 * g2, a1, a2 + + +def _call(data, sf, args, num_experts, top_k, scale_as_is=False, **kw): + """Invoke the wrapper with this layer's tensors and size scalars. + + `sf` is handed over as `float8_e4m3fn`, the dtype the op demands, unless + `scale_as_is` asks for the raw tensor (used by the dtype negative tests). + """ + sf = sf if scale_as_is else sf.view(torch.float8_e4m3fn) + kw.setdefault("routing_logits", None) + kw.setdefault("routing_bias", None) + kw.setdefault("gemm1_bias", None) + kw.setdefault("gemm1_alpha", None) + kw.setdefault("gemm1_beta", None) + kw.setdefault("gemm1_clamp_limit", None) + kw.setdefault("gemm2_bias", None) + kw.setdefault("n_group", None) + kw.setdefault("topk_group", None) + kw.setdefault("local_expert_offset", 0) + kw.setdefault("local_num_experts", num_experts) + kw.setdefault("routed_scaling_factor", None) + kw.setdefault("routing_method_type", 1) + kw.setdefault("do_finalize", True) + kw.setdefault("act_type", 0) + kw.setdefault("intermediate_size", args["intermediate_size"]) + for k in ( + "gemm1_weights", + "gemm1_weights_scale", + "gemm2_weights", + "gemm2_weights_scale", + ): + kw.setdefault(k, args[k]) + return moe( + kw.pop("routing_logits"), + kw.pop("routing_bias"), + data, + sf, + kw.pop("gemm1_weights"), + kw.pop("gemm1_weights_scale"), + kw.pop("gemm1_bias"), + kw.pop("gemm1_alpha"), + kw.pop("gemm1_beta"), + kw.pop("gemm1_clamp_limit"), + kw.pop("gemm2_weights"), + kw.pop("gemm2_weights_scale"), + kw.pop("gemm2_bias"), + kw.pop("output1_scale_scalar"), + kw.pop("output1_scale_gate_scalar"), + kw.pop("output2_scale_scalar"), + num_experts, + top_k, + kw.pop("n_group"), + kw.pop("topk_group"), + kw.pop("intermediate_size"), + kw.pop("local_expert_offset"), + kw.pop("local_num_experts"), + kw.pop("routed_scaling_factor"), + kw.pop("routing_method_type"), + kw.pop("do_finalize"), + kw.pop("act_type"), + **kw, + ) + + +# ── reference ───────────────────────────────────────────────────────────── + + +def _ref_moe( + x, + ids, + scales, + ref, + g2: float, + offset: int = 0, + num_local: int | None = None, + swap_gate_up: bool = False, + swap_scalars: bool = False, + quantize_intermediate: bool = True, +): + """Native-torch MoE over dequantized NVFP4 weights, fp32 throughout. + + `ids` carries global expert ids; this rank answers for + `[offset, offset + num_local)`. `scales[t, j]` multiplies slot `j`'s expert + output; nothing is renormalized. The FC1 activation is requantized to + NVFP4 under global scale `g2` before FC2, which is what the kernel does. + """ + num_tokens, hidden = x.shape + up_c, up_s = ref["up"] + gt_c, gt_s = ref["gate"] + dn_c, dn_s = ref["down"] + if swap_gate_up: + (up_c, up_s), (gt_c, gt_s) = (gt_c, gt_s), (up_c, up_s) + num_local = up_c.shape[0] if num_local is None else num_local + out = torch.zeros(num_tokens, hidden, dtype=torch.float32, device=x.device) + for local_e in range(num_local): + tok, slot = (ids == offset + local_e).nonzero(as_tuple=True) + if tok.numel() == 0: + continue + xe = x[tok] + up = xe @ _dequant(up_c[local_e], up_s[local_e]).t() + gate = xe @ _dequant(gt_c[local_e], gt_s[local_e]).t() + if swap_scalars: + up, gate = up / g2, gate * g2 + act = up * gate * torch.sigmoid(gate) + # the FC1 epilogue emits g2 * act as NVFP4; FC2's alpha divides it back + act = _q_nvfp4(act, g2) / g2 if quantize_intermediate else act + y = act @ _dequant(dn_c[local_e], dn_s[local_e]).t() + out.index_add_(0, tok, y * scales[tok, slot].float().unsqueeze(1)) + return out + + +def _calibrate_g2(x, ids, ref) -> float: + """`448*6 / amax` of the true FC1 activation — the FC2 input global scale a + checkpoint's `down_proj.input_scale` encodes.""" + up_c, up_s = ref["up"] + gt_c, gt_s = ref["gate"] + amax = 0.0 + for e in range(up_c.shape[0]): + tok = (ids == e).nonzero(as_tuple=True)[0].unique() + if tok.numel() == 0: + continue + xe = x[tok] + up = xe @ _dequant(up_c[e], up_s[e]).t() + gate = xe @ _dequant(gt_c[e], gt_s[e]).t() + amax = max(amax, (up * gate * torch.sigmoid(gate)).abs().max().item()) + return E4M3_MAX * E2M1_MAX / amax + + +def _dev(y: torch.Tensor, ref: torch.Tensor): + """(worst element deviation, relative RMS deviation), both in bf16 ulp.""" + o, r = y.float(), ref.float() + row = r.abs().amax(dim=1, keepdim=True).clamp_min(1e-9) + elt = ((o - r).abs() / row).max().item() / ULP + rms = ((o - r).pow(2).mean().sqrt() / r.pow(2).mean().sqrt().clamp_min(1e-9)).item() / ULP + return elt, rms + + +def _assert_moe_close(y: torch.Tensor, ref: torch.Tensor) -> None: + """Two gates: per-element, row-scaled; and aggregate relative RMS. + + Kernel and reference consume bit-identical NVFP4 weights and activations + and model the same FC1-output NVFP4 requantization, so they differ only in + accumulation order and in whether a marginal FC1 value rounds to the same + e2m1 code — and an e2m1 code is a *coarse* step (2 mantissa bits), so one + flipped intermediate element moves the output row by a visible fraction. + Default `assert_close` tolerances cannot express that: their bf16 + `atol=1e-5` sits three orders of magnitude below one output ulp of a + two-GEMM chain, and per-element `rtol` is meaningless where cancellation + drives `|ref|` to zero. So the element gate is 16 ulp of the row's largest + magnitude and the aggregate gate is 2 ulp of relative RMS. Worst values + measured over every configuration covered here: 9.97 ulp element-wise (an + 18-wide expert-parallel window at 8192 tokens, H=2560, I=1536) and 0.88 + ulp RMS (those four windows summed); a full-window call sits at 4.75 / + 0.72. The DeepSeek-R1 routed geometry (H=7168, I=2048, 256 experts, + top-8) lands just inside both: 9.62 / 0.74 for a 64-wide expert-parallel + window, 4.68 / 0.70 for the 256-expert stack, 5.02 / 0.85 for the four + windows summed. Two things inflate the element figure and not the RMS + one: it is an extreme order statistic over every element compared, and it + is normalized by the *reference's own* row magnitude, so a window carrying + a quarter of the routed slots divides a similar absolute error by a + smaller row. + Re-drawn over 144 independent (seed, window) comparisons at T=1024/4096 it + reached 12.0 while the RMS figure stayed at 0.86 — headroom on this + element gate is ~1.3x for a window, ~2.3x on the RMS gate. Re-drawn over + 16 fresh seeds at T=8192 (16 full-stack + 64 window comparisons) it + reached 16.33 for a window (p90 9.82) and 10.41 for the full stack, while + the RMS figure was unchanged from T=4096 (window 0.74, full 0.70, sum + 0.85): the element metric is a max over T*H elements, so its tail grows + with T at constant accuracy and *crosses* this gate at the top of the + certified range. Both of those re-draws are the H=2560 / I=1536 geometry; + re-drawn the same way at H=7168 / I=2048 over 12 fresh weight+activation + seeds the tail is *lower* — max 12.74 for a 64-wide window at T=8192 (p90 + 8.40, mean 6.73), 12.15 at T=4096, and 7.11 / 6.42 for the 256-expert + stack, with the RMS figures flat at 0.74 (window) / 0.70 (full) / 0.86 + (sum) at both token counts. The seeded draws shipped here land at 9.97 / + 3.20 (H=2560) and 9.62 / 4.68 (H=7168), but only the RMS gate is + scale-free — read it, not the element one, when sizing a caller's + tolerance. + Both gates bite — `test_reference_discriminates` shows, at H=256/I=128 + and at H=7168/I=2048 respectively, a gate/up swap (263 / 278 ulp), a + scalar-role swap (177 / 48), a missing scale swizzle (408 / 357), a + missing row shuffle (508 / 432), and a reference that skips the FC1-output + requantization (43 / 28 ulp element, 23 / 23 ulp RMS); the smallest RMS + distance any of those reaches is 39.8, at the R1 shape's scalar-role swap. + `test_ep_window_18_of_72` and `test_ep_window_64_of_256` add a mis-set + expert-parallel window (304-363 ulp RMS) and one window left out of the + sum (>=124 ulp RMS overall, >=128 at T=8192). + """ + assert y.dtype == ref.dtype == torch.bfloat16, (y.dtype, ref.dtype) + assert y.shape == ref.shape, (y.shape, ref.shape) + row = ref.float().abs().amax(dim=1, keepdim=True).clamp_min(1e-9) + torch.testing.assert_close(y.float() / row, ref.float() / row, rtol=0.0, atol=16 * ULP) + _, rms = _dev(y, ref) + assert rms <= 2.0, f"relative RMS {rms:.2f} ulp > 2 ulp" + + +def _routing(num_tokens, num_experts, top_k, gen): + """Random ids/weights for the pre-routed entry point.""" + ids = torch.stack( + [torch.randperm(num_experts, device=DEV, generator=gen)[:top_k] for _ in range(num_tokens)] + ).to(torch.int32) + wts = torch.rand(num_tokens, top_k, device=DEV, generator=gen).to(torch.bfloat16) + return ids, wts + + +# DeepSeek-V3-Lite routed-expert geometry, built once and shared. +_DSV3 = None + + +def _dsv3(): + global _DSV3 + if _DSV3 is None: + _DSV3 = _build(72, 2560, 1536, seed=100) + return _DSV3 + + +# DeepSeek-R1 routed-expert geometry, built once and shared: 256 experts, +# H = 7168, I = 2048. The kernel-ready stack alone is ~6.3 GiB (FC1 +# [256, 4096, 3584] + [256, 4096, 448], FC2 [256, 7168, 1024] + +# [256, 7168, 128]) and the reference operands take it to ~22 GiB. +_DSR1 = None + + +def _dsr1(): + global _DSR1 + if _DSR1 is None: + _DSR1 = _build(256, 7168, 2048, seed=256) + return _DSR1 + + +# The token column certified at the R1 geometry. 8192 is trtllm's default +# `max_num_tokens`; a dep4 target gathers four ranks' tokens before the expert +# call, so its own T reaches 4 * 8192 and it must chunk down into calls this +# column covers. +_R1_TOKENS = (1, 2, 3, 5, 7, 8, 16, 24, 32, 33, 40, 64, 128, 256, 1024, 4096, 8192) + + +# ── tests ───────────────────────────────────────────────────────────────── + + +def test_layout_helpers_match(): + """The trtllm preprocessing ops reproduce the pure-torch layout exactly. + + The last two scale shapes are the DeepSeek-R1 routed geometry's, at + H = 7168 / I = 2048: FC1 `[E, 2*I, H/16] = [E, 4096, 448]` and FC2 + `[E, H, I/16] = [E, 7168, 128]`. Both satisfy the swizzle's `M % 128 == 0` + and `C % 4 == 0`, so `block_scale_interleave` returns exactly `E*M*C` bytes + and reshapes back — no padding enters either operand at this geometry. + """ + gen = torch.Generator(device=DEV).manual_seed(201) + for rows, cols in ((256, 12), (4096, 448), (7168, 128)): + x = torch.randint(0, 256, (4, rows, cols), dtype=torch.uint8, device=DEV, generator=gen) + perm = _fc1_perm(rows) + for e in range(x.shape[0]): + assert torch.equal( + torch.ops.trtllm.shuffle_matrix(x[e].contiguous(), perm), + torch.index_select(x[e], 0, perm), + ), f"shuffle_matrix is not a plain row gather at {rows}x{cols}" + swz = torch.ops.trtllm.block_scale_interleave(x) + assert swz.numel() == x.numel(), ( + f"block_scale_interleave padded {rows}x{cols}: {swz.numel()} bytes " + f"for {x.numel()} scales" + ) + assert torch.equal(swz.reshape(x.shape), _swizzle_scales(x)), ( + f"block_scale_interleave does not match the documented 128x4 swizzle at {rows}x{cols}" + ) + print(" test_layout_helpers_match OK") + + +def test_deepseek_geometry(): + """E=72, H=2560, I=1536, top-6: decode- through prefill-sized batches. + + The top of the sweep is 8192 rows in one call — trtllm's default + `max_num_tokens`, i.e. the widest prefill a stock engine hands this op + without chunking. Token counts are appended in ascending order, so every + count below the top draws exactly the values it drew before. + """ + args, ref = _dsv3() + gen = ref["gen"] + worst = (0.0, 0.0) + token_counts = (1, 2, 8, 64, 256, 1024, 4096, 8192) + for num_tokens in token_counts: + g1 = 137.0 + data, sf, x = _rand_activation(num_tokens, 2560, g1, gen) + ids, wts = _routing(num_tokens, 72, 6, gen) + g2 = _calibrate_g2(x, ids, ref) + s1, sg, s2 = _scalars(72, g1, g2) + out = _call( + data, + sf, + args, + 72, + 6, + topk_ids=ids, + topk_weights=wts, + output1_scale_scalar=s1, + output1_scale_gate_scalar=sg, + output2_scale_scalar=s2, + ) + assert len(out) == 1, len(out) + y = out[0] + assert y.shape == (num_tokens, 2560) and y.dtype == torch.bfloat16, y.shape + exp = _ref_moe(x, ids, wts, ref, g2).to(torch.bfloat16) + _assert_moe_close(y, exp) + top = _dev(y, exp) # after the loop: the largest token count's numbers + worst = tuple(max(a, b) for a, b in zip(worst, top)) + print( + f" test_deepseek_geometry OK (worst {worst[0]:.2f} elt / {worst[1]:.2f} rms ulp; " + f"at T={token_counts[-1]} {top[0]:.2f} elt / {top[1]:.2f} rms ulp)" + ) + + +def test_other_geometries(): + """Smaller expert counts, hidden and intermediate sizes, and top_k values.""" + worst = (0.0, 0.0) + for num_experts, hidden, inter, top_k, tokens in ( + (8, 256, 128, 2, (1, 3, 16, 128)), + (4, 512, 256, 1, (1, 7, 64)), + (16, 256, 64, 4, (2, 33)), + (2, 256, 512, 1, (5,)), + (3, 768, 192, 2, (1, 40)), + ): + args, ref = _build(num_experts, hidden, inter, seed=7 * num_experts + hidden) + gen = ref["gen"] + for num_tokens in tokens: + g1 = 71.0 + data, sf, x = _rand_activation(num_tokens, hidden, g1, gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + g2 = _calibrate_g2(x, ids, ref) + s1, sg, s2 = _scalars(num_experts, g1, g2) + y = _call( + data, + sf, + args, + num_experts, + top_k, + topk_ids=ids, + topk_weights=wts, + output1_scale_scalar=s1, + output1_scale_gate_scalar=sg, + output2_scale_scalar=s2, + )[0] + assert y.shape == (num_tokens, hidden), y.shape + exp = _ref_moe(x, ids, wts, ref, g2).to(torch.bfloat16) + _assert_moe_close(y, exp) + worst = tuple(max(a, b) for a, b in zip(worst, _dev(y, exp))) + print(f" test_other_geometries OK (worst {worst[0]:.2f} elt / {worst[1]:.2f} rms ulp)") + + +def test_intermediate_is_nvfp4_requantized(): + """Read the FC1 epilogue's output straight out of the kernel. + + An identity down projection with unit block scales and + `output2_scale_scalar = 1` makes the combined output equal the requantized + intermediate value element for element, and e2m1 x e4m3 products are exact + in bf16 — so the comparison is bit-exact, not toleranced. + """ + num_experts, hidden, top_k, num_tokens = 4, 256, 1, 64 + gen = torch.Generator(device=DEV).manual_seed(77) + up_p, up_s, up_c = _rand_nvfp4(num_experts, hidden, hidden, gen) + gt_p, gt_s, gt_c = _rand_nvfp4(num_experts, hidden, hidden, gen) + w1, s1 = _prep_fc1(up_p, gt_p, up_s, gt_s) + eye = torch.zeros(num_experts, hidden, hidden, dtype=torch.uint8, device=DEV) + d = torch.arange(hidden, device=DEV) + eye[:, d, d] = 2 # e2m1 code 2 == +1.0 + dn_p = (eye[..., 0::2] | (eye[..., 1::2] << 4)).contiguous() + dn_s = torch.ones(num_experts, hidden, hidden // SV, device=DEV).to(torch.float8_e4m3fn) + w2, s2 = _prep_fc2(dn_p, dn_s) + + g1 = 137.0 + data, sf, x = _rand_activation(num_tokens, hidden, g1, gen) + ids = torch.randint( + 0, num_experts, (num_tokens, 1), dtype=torch.int32, device=DEV, generator=gen + ) + wts = torch.ones(num_tokens, 1, dtype=torch.bfloat16, device=DEV) + + act = torch.zeros(num_tokens, hidden, dtype=torch.float32, device=DEV) + for e in range(num_experts): + tok = (ids[:, 0] == e).nonzero(as_tuple=True)[0] + if tok.numel() == 0: + continue + xe = x[tok] + up = xe @ _dequant(up_c[e], up_s[e]).t() + gate = xe @ _dequant(gt_c[e], gt_s[e]).t() + act[tok] = up * gate * torch.sigmoid(gate) + g2 = E4M3_MAX * E2M1_MAX / act.abs().max().item() + a1 = torch.full((num_experts,), 1.0 / g1, dtype=torch.float32, device=DEV) + y = _call( + data, + sf, + dict( + gemm1_weights=w1, + gemm1_weights_scale=s1, + gemm2_weights=w2, + gemm2_weights_scale=s2, + intermediate_size=hidden, + ), + num_experts, + top_k, + topk_ids=ids, + topk_weights=wts, + output1_scale_scalar=a1 * g2, + output1_scale_gate_scalar=a1, + output2_scale_scalar=torch.ones(num_experts, dtype=torch.float32, device=DEV), + )[0].float() + + exact = _q_nvfp4(act, g2) + assert torch.equal(y, exact), ( + f"FC1 epilogue is not sf=e4m3(g*amax/6), data=e2m1(g*x/sf): " + f"{(y != exact).sum().item()} of {y.numel()} elements differ" + ) + # an unrounded block scale and the unquantized activation both fail here + b = (act * g2).reshape(num_tokens, hidden // SV, SV) + sfx = b.abs().amax(-1, keepdim=True) / E2M1_MAX + unrounded = (_e2m1_rne(b / sfx.clamp_min(1e-30)) * sfx).reshape(num_tokens, hidden) + fracs = {} + for name, cand in ( + ("no requantization", act * g2), + ("exact (unrounded) scale", unrounded), + ): + fracs[name] = (y == cand).float().mean().item() + assert fracs[name] < 0.9, f"{name} also matches {fracs[name]:.2%} — no discrimination" + print( + f" test_intermediate_is_nvfp4_requantized OK (bit-exact on all " + f"{y.numel()} elements; unrounded scale {fracs['exact (unrounded) scale']:.1%}, " + f"no requantization {fracs['no requantization']:.1%})" + ) + + +def test_reference_discriminates(): + """Every layout / scalar-role mistake lands far outside the numeric gate. + + Run at two hidden/intermediate pairs, because both operand layouts are + functions of `H` and `I`: the small one, and the DeepSeek-R1 routed shape + (`H = 7168`, `I = 2048`) at a reduced expert count — the row shuffle and + the 128x4 scale swizzle are per-expert, so `E` is free to be small while + the layout question is entirely `H`/`I`. + """ + for num_experts, hidden, inter, top_k, num_tokens, seed in ( + (8, 256, 128, 2, 32, 303), + (16, 7168, 2048, 8, 32, 7168), + ): + args, ref = _build(num_experts, hidden, inter, seed=seed) + gen = ref["gen"] + g1 = 137.0 + data, sf, x = _rand_activation(num_tokens, hidden, g1, gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + g2 = _calibrate_g2(x, ids, ref) + s1, sg, s2 = _scalars(num_experts, g1, g2) + base = dict( + topk_ids=ids, + topk_weights=wts, + output1_scale_scalar=s1, + output1_scale_gate_scalar=sg, + output2_scale_scalar=s2, + ) + y = _call(data, sf, args, num_experts, top_k, **base)[0] + good = _ref_moe(x, ids, wts, ref, g2).to(torch.bfloat16) + _assert_moe_close(y, good) + print(f" E={num_experts} H={hidden} I={inter} top-{top_k}, T={num_tokens}:") + + up_p, gt_p, up_s, gt_s, dn_p, dn_s = ref["raw"] + variants = { + "[gate|up] instead of [up|gate]": _dev( + y, _ref_moe(x, ids, wts, ref, g2, swap_gate_up=True).to(torch.bfloat16) + ), + "output1 scalars swapped": _dev( + y, _ref_moe(x, ids, wts, ref, g2, swap_scalars=True).to(torch.bfloat16) + ), + "no intermediate requantization": _dev( + y, + _ref_moe(x, ids, wts, ref, g2, quantize_intermediate=False).to(torch.bfloat16), + ), + } + for name, (elt, rms) in variants.items(): + assert elt > 25.0, f"{name} only {elt:.1f} ulp away — not discriminated" + print(f" {name:34s} {elt:8.1f} elt / {rms:6.2f} rms ulp") + + # operand-preparation mistakes: the kernel is fed the wrong bytes + for name, kw in ( + ( + "weight scales not 128x4 swizzled", + dict( + gemm1_weights_scale=_prep_fc1(up_p, gt_p, up_s, gt_s, swizzle=False)[1], + gemm2_weights_scale=_prep_fc2(dn_p, dn_s, swizzle=False)[1], + ), + ), + ( + "rows not shuffled", + dict( + gemm1_weights=_prep_fc1(up_p, gt_p, up_s, gt_s, shuffle=False)[0], + gemm1_weights_scale=_prep_fc1(up_p, gt_p, up_s, gt_s, shuffle=False)[1], + gemm2_weights=_prep_fc2(dn_p, dn_s, shuffle=False)[0], + gemm2_weights_scale=_prep_fc2(dn_p, dn_s, shuffle=False)[1], + ), + ), + ): + bad = _call(data, sf, args, num_experts, top_k, **base, **kw)[0] + elt, rms = _dev(bad, good.float()) + assert elt > 25.0, f"{name} only {elt:.1f} ulp away — not discriminated" + print(f" {name:34s} {elt:8.1f} elt / {rms:6.2f} rms ulp") + del args, ref, y, good, data, sf, x + torch.cuda.empty_cache() + print(" test_reference_discriminates OK") + + +def test_fp4_quantize_pairing(): + """`fp4_quantize(x, g, 16, False, False)` is the activation pairing. + + The linear scale layout is required; the swizzled one is rejected by size + except when the byte counts coincide, where it is silently wrong. + """ + num_experts, hidden, inter, top_k = 8, 256, 128, 2 + args, ref = _build(num_experts, hidden, inter, seed=404) + gen = ref["gen"] + for num_tokens in (1, 7, 128): + xb = torch.randn(num_tokens, hidden, device=DEV, dtype=torch.bfloat16, generator=gen) + g1 = E4M3_MAX * E2M1_MAX / xb.abs().max().float().item() + gs = torch.tensor([g1], dtype=torch.float32, device=DEV) + data, sf = torch.ops.trtllm.fp4_quantize(xb, gs, SV, False, False) + assert data.shape == (num_tokens, hidden // 2) and data.dtype == torch.uint8 + assert sf.shape == (num_tokens * hidden // SV,) and sf.dtype == torch.uint8 + x = _unpack_activation(data, sf, hidden, g1) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + g2 = _calibrate_g2(x, ids, ref) + s1, sg, s2 = _scalars(num_experts, g1, g2) + kw = dict( + topk_ids=ids, + topk_weights=wts, + output1_scale_scalar=s1, + output1_scale_gate_scalar=sg, + output2_scale_scalar=s2, + ) + y = _call(data, sf, args, num_experts, top_k, **kw)[0] + _assert_moe_close(y, _ref_moe(x, ids, wts, ref, g2).to(torch.bfloat16)) + + _, swz = torch.ops.trtllm.fp4_quantize(xb, gs, SV, False, True) + if swz.numel() == sf.numel(): + assert num_tokens % 128 == 0, num_tokens + bad = _call(data, swz, args, num_experts, top_k, **kw)[0] + elt, _ = _dev(bad, y.float()) + assert elt > 25.0, f"swizzled scales only {elt:.1f} ulp off" + print(f" T={num_tokens}: swizzled scale buffer accepted, {elt:.0f} ulp wrong") + else: + try: + _call(data, swz, args, num_experts, top_k, **kw) + raise AssertionError("swizzled scale buffer of a different size accepted") + except RuntimeError as e: + assert "hidden_states_scale has incorrect size" in str(e), e + print(" test_fp4_quantize_pairing OK") + + +def test_scalars_and_optional_tensors(): + """`None` is accepted for the five schema-mandatory bias/GLU tensors, and is + exactly their neutral value; the three scale scalars are per local expert.""" + num_experts, hidden, inter, top_k, num_tokens = 8, 256, 128, 2, 16 + args, ref = _build(num_experts, hidden, inter, seed=505) + gen = ref["gen"] + g1, g2 = 137.0, 1.0 / 32.0 + data, sf, x = _rand_activation(num_tokens, hidden, g1, gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + s1, sg, s2 = _scalars(num_experts, g1, g2) + kw = dict( + topk_ids=ids, + topk_weights=wts, + output1_scale_scalar=s1, + output1_scale_gate_scalar=sg, + output2_scale_scalar=s2, + ) + base = _call(data, sf, args, num_experts, top_k, **kw)[0] + + def f32(shape, value): + return torch.full(shape, value, dtype=torch.float32, device=DEV) + + neutral = ( + ("gemm1_bias", f32((num_experts, 2 * inter), 0.0)), + ("gemm2_bias", f32((num_experts, hidden), 0.0)), + ("gemm1_alpha", f32((num_experts,), 1.0)), + ("gemm1_beta", f32((num_experts,), 0.0)), + ("gemm1_clamp_limit", f32((num_experts,), 1e9)), + ) + for name, t in neutral: + y = _call(data, sf, args, num_experts, top_k, **kw, **{name: t})[0] + assert torch.equal(y, base), f"{name} at its neutral value differs from None" + # ...and each slot is genuinely live: a non-neutral value changes the result + live = ( + ("gemm1_bias", f32((num_experts, 2 * inter), 1e4)), + ("gemm2_bias", f32((num_experts, hidden), 1e2)), + ("gemm1_alpha", f32((num_experts,), 8.0)), + ("gemm1_beta", f32((num_experts,), 50.0)), + ("gemm1_clamp_limit", f32((num_experts,), 1e-3)), + ) + for name, t in live: + y = _call(data, sf, args, num_experts, top_k, **kw, **{name: t})[0] + assert not torch.equal(y, base), f"{name} is silently ignored" + + for name in ( + "output1_scale_scalar", + "output1_scale_gate_scalar", + "output2_scale_scalar", + ): + for bad, msg in ( + (kw[name][:-1].contiguous(), "incorrect dim 0"), + (kw[name].double(), "must be float"), + ): + try: + _call(data, sf, args, num_experts, top_k, **{**kw, name: bad}) + raise AssertionError(f"{name} {msg} accepted") + except RuntimeError as e: + assert "scalar" in str(e) and msg in str(e), (name, msg, str(e)) + print(" test_scalars_and_optional_tensors OK") + + +def test_expert_window_and_ids(): + """local_expert_offset / local_num_experts, out-of-window ids, duplicates.""" + num_experts, hidden, inter, top_k, num_tokens = 8, 256, 128, 2, 24 + args, ref = _build(num_experts, hidden, inter, seed=606) + gen = ref["gen"] + g1, g2 = 137.0, 1.0 / 32.0 + data, sf, x = _rand_activation(num_tokens, hidden, g1, gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + s1, sg, s2 = _scalars(num_experts, g1, g2) + + off, loc = 4, 4 + y = _call( + data, + sf, + args, + num_experts, + top_k, + topk_ids=ids, + topk_weights=wts, + local_expert_offset=off, + local_num_experts=loc, + gemm1_weights=args["gemm1_weights"][off:].contiguous(), + gemm1_weights_scale=args["gemm1_weights_scale"][off:].contiguous(), + gemm2_weights=args["gemm2_weights"][off:].contiguous(), + gemm2_weights_scale=args["gemm2_weights_scale"][off:].contiguous(), + output1_scale_scalar=s1[off:].contiguous(), + output1_scale_gate_scalar=sg[off:].contiguous(), + output2_scale_scalar=s2[off:].contiguous(), + )[0] + sub = dict(ref) + sub["up"] = (ref["up"][0][off:], ref["up"][1][off:]) + sub["gate"] = (ref["gate"][0][off:], ref["gate"][1][off:]) + sub["down"] = (ref["down"][0][off:], ref["down"][1][off:]) + exp = _ref_moe(x, ids, wts, sub, g2, offset=off, num_local=loc).to(torch.bfloat16) + _assert_moe_close(y, exp) + + kw = dict( + topk_weights=wts, + output1_scale_scalar=s1, + output1_scale_gate_scalar=sg, + output2_scale_scalar=s2, + ) + # ids outside [offset, offset+local) are dropped, negative and >= num_experts too + for bad_id in (-1, num_experts, num_experts + 5): + mixed = ids.clone() + mixed[:, 0] = bad_id + dropped = ids.clone() + dropped[:, 0] = -1 + a = _call(data, sf, args, num_experts, top_k, topk_ids=mixed, **kw)[0] + b = _call(data, sf, args, num_experts, top_k, topk_ids=dropped, **kw)[0] + assert torch.equal(a, b), f"expert id {bad_id} was not dropped" + + # A repeated id inside one row is accepted, but whether it contributes once + # or twice is not stable: with this geometry it collapsed to the first slot + # at 16 tokens and contributed both slots at 24. Pin the disjunction, and + # which branch each token count took. + seen = set() + for tokens in (16, 24): + d16, s16, x16 = _rand_activation(tokens, hidden, g1, gen) + i16, w16 = _routing(tokens, num_experts, top_k, gen) + dup = i16.clone() + dup[:, 1] = dup[:, 0] + drop = torch.full_like(dup[:, 0], -1) + kw16 = dict(kw, topk_weights=w16) + y = _call(d16, s16, args, num_experts, top_k, topk_ids=dup, **kw16)[0] + slot0 = _call( + d16, + s16, + args, + num_experts, + top_k, + topk_ids=torch.stack([dup[:, 0], drop], 1).contiguous(), + **kw16, + )[0] + slot1 = _call( + d16, + s16, + args, + num_experts, + top_k, + topk_ids=torch.stack([drop, dup[:, 1]], 1).contiguous(), + **kw16, + )[0] + both = (slot0.float() + slot1.float()).to(torch.bfloat16) + branch = ( + "once" + if torch.equal(y, slot0) + else ("twice" if _dev(y, both.float())[0] < 2.0 else None) + ) + assert branch is not None, ( + f"a duplicated expert id at {tokens} tokens matched neither one nor two " + f"contributions ({_dev(y, slot0.float())[0]:.1f} / {_dev(y, both.float())[0]:.1f} ulp)" + ) + seen.add((tokens, branch)) + print(f" test_expert_window_and_ids OK (duplicate ids: {sorted(seen)})") + + +def _window_args(args, off: int, num_local: int): + """The kernel-ready weight operands of one expert-parallel window.""" + return dict( + gemm1_weights=args["gemm1_weights"][off : off + num_local].contiguous(), + gemm1_weights_scale=args["gemm1_weights_scale"][off : off + num_local].contiguous(), + gemm2_weights=args["gemm2_weights"][off : off + num_local].contiguous(), + gemm2_weights_scale=args["gemm2_weights_scale"][off : off + num_local].contiguous(), + intermediate_size=args["intermediate_size"], + ) + + +def _window_ref(ref, off: int, num_local: int): + """The reference operands of one expert-parallel window.""" + sub = dict(ref) + for k in ("up", "gate", "down"): + sub[k] = (ref[k][0][off : off + num_local], ref[k][1][off : off + num_local]) + return sub + + +def _window_scalars(s1, sg, s2, off: int, num_local: int): + """The three per-expert fp32 scalars restricted to one window.""" + return dict( + output1_scale_scalar=s1[off : off + num_local].contiguous(), + output1_scale_gate_scalar=sg[off : off + num_local].contiguous(), + output2_scale_scalar=s2[off : off + num_local].contiguous(), + ) + + +def test_ep_window_18_of_72(): + """The tep4 expert-parallel split: 72 routed experts as four 18-wide windows. + + `num_experts` stays 72 on every call — routing happens outside and every + rank sees the whole space — while `local_num_experts = 18` and + `local_expert_offset` runs 0 / 18 / 36 / 54 over the DeepSeek-V3-Lite + routed geometry (H=2560, I=1536, top-6). Three things are checked: each + window answers for exactly the slots routed into it, a token routed + entirely elsewhere comes back bitwise zero, and the four bf16 outputs sum + back to the 72-expert result (the value the tep4 all-reduce adds up). + + The sweep tops out at 8192 rows per call — trtllm's default + `max_num_tokens`, the widest prefill a stock engine hands each rank + unchunked. Counts are appended in ascending order, so every count below + the top draws exactly the values it drew before. + """ + args, ref = _dsv3() + num_experts, hidden, top_k, shard = 72, 2560, 6, 18 + windows = [(r * shard, shard) for r in range(4)] + win_args = [_window_args(args, off, n) for off, n in windows] + win_ref = [_window_ref(ref, off, n) for off, n in windows] + gen = torch.Generator(device=DEV).manual_seed(1818) + g1 = 137.0 + worst_win = (0.0, 0.0) + worst_sum = (0.0, 0.0) + worst_vs_full = (0.0, 0.0) + worst_drop = 1e9 + token_counts = (1, 8, 256, 1024, 4096, 8192) + + for num_tokens in token_counts: + data, sf, x = _rand_activation(num_tokens, hidden, g1, gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + g2 = _calibrate_g2(x, ids, ref) + s1, sg, s2 = _scalars(num_experts, g1, g2) + full = _call( + data, + sf, + args, + num_experts, + top_k, + topk_ids=ids, + topk_weights=wts, + **_window_scalars(s1, sg, s2, 0, num_experts), + )[0] + exp_full = _ref_moe(x, ids, wts, ref, g2).to(torch.bfloat16) + _assert_moe_close(full, exp_full) + + parts = [] + top_win = (0.0, 0.0) # after the loop: the largest token count's numbers + for (off, n), wa, wr in zip(windows, win_args, win_ref): + y = _call( + data, + sf, + wa, + num_experts, + top_k, + local_expert_offset=off, + local_num_experts=n, + topk_ids=ids, + topk_weights=wts, + **_window_scalars(s1, sg, s2, off, n), + )[0] + assert y.shape == (num_tokens, hidden), y.shape + exp = _ref_moe(x, ids, wts, wr, g2, offset=off, num_local=n).to(torch.bfloat16) + _assert_moe_close(y, exp) + outside = ~((ids >= off) & (ids < off + n)).any(dim=1) + assert torch.equal(y[outside], torch.zeros_like(y[outside])), ( + f"window (off={off}, n={n}) wrote into rows whose token routes " + f"entirely elsewhere ({int(outside.sum())} such rows at T={num_tokens})" + ) + top_win = tuple(max(a, b) for a, b in zip(top_win, _dev(y, exp))) + worst_win = tuple(max(a, b) for a, b in zip(worst_win, _dev(y, exp))) + parts.append(y) + + summed = sum(p.float() for p in parts).to(torch.bfloat16) + _assert_moe_close(summed, exp_full) + top_sum = _dev(summed, exp_full) + worst_sum = tuple(max(a, b) for a, b in zip(worst_sum, _dev(summed, exp_full))) + worst_vs_full = tuple(max(a, b) for a, b in zip(worst_vs_full, _dev(summed, full.float()))) + # The sum gate bites: leaving any one window out lands far outside it, + # so passing it is not something four near-arbitrary partials could do. + # Only meaningful once every window carries a settled share of the + # slots — at a handful of tokens a window can legitimately hold none. + if num_tokens >= 256: + top_drop = 1e9 # after the loop: the largest token count's figure + for drop in range(4): + partial = sum(p.float() for i, p in enumerate(parts) if i != drop) + _, rms = _dev(partial.to(torch.bfloat16), exp_full) + assert rms > 25.0, ( + f"dropping window {drop} at T={num_tokens} is only {rms:.1f} rms ulp away" + ) + worst_drop = min(worst_drop, rms) + top_drop = min(top_drop, rms) + + # Every routed slot inside one window: that window alone reproduces the + # 72-expert call bitwise, and the other three return exactly zero. + num_tokens = 64 + data, sf, x = _rand_activation(num_tokens, hidden, g1, gen) + ids = torch.stack( + [torch.randperm(shard, device=DEV, generator=gen)[:top_k] for _ in range(num_tokens)] + ).to(torch.int32) + wts = torch.rand(num_tokens, top_k, device=DEV, generator=gen).to(torch.bfloat16) + g2 = _calibrate_g2(x, ids, ref) + s1, sg, s2 = _scalars(num_experts, g1, g2) + full = _call( + data, + sf, + args, + num_experts, + top_k, + topk_ids=ids, + topk_weights=wts, + **_window_scalars(s1, sg, s2, 0, num_experts), + )[0] + for (off, n), wa in zip(windows, win_args): + y = _call( + data, + sf, + wa, + num_experts, + top_k, + local_expert_offset=off, + local_num_experts=n, + topk_ids=ids, + topk_weights=wts, + **_window_scalars(s1, sg, s2, off, n), + )[0] + if off == 0: + assert torch.equal(y, full), ( + "the window holding every routed slot did not reproduce the 72-expert call bitwise" + ) + else: + assert torch.equal(y, torch.zeros_like(y)), ( + f"window off={off} holds no routed slot but returned nonzero rows" + ) + + # Window mistakes a caller can make, against the correct window-1 result. + num_tokens = 256 + data, sf, x = _rand_activation(num_tokens, hidden, g1, gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + g2 = _calibrate_g2(x, ids, ref) + s1, sg, s2 = _scalars(num_experts, g1, g2) + off, n = windows[1] + good = _call( + data, + sf, + win_args[1], + num_experts, + top_k, + local_expert_offset=off, + local_num_experts=n, + topk_ids=ids, + topk_weights=wts, + **_window_scalars(s1, sg, s2, off, n), + )[0] + good_ref = _ref_moe(x, ids, wts, win_ref[1], g2, offset=off, num_local=n).to(torch.bfloat16) + _assert_moe_close(good, good_ref) + wrong_offset = _call( + data, + sf, + win_args[1], + num_experts, + top_k, + local_expert_offset=0, + local_num_experts=n, + topk_ids=ids, + topk_weights=wts, + **_window_scalars(s1, sg, s2, off, n), + )[0] + wrong_weights = _call( + data, + sf, + win_args[0], + num_experts, + top_k, + local_expert_offset=off, + local_num_experts=n, + topk_ids=ids, + topk_weights=wts, + **_window_scalars(s1, sg, s2, off, n), + )[0] + all_at_zero = sum( + _call( + data, + sf, + wa, + num_experts, + top_k, + local_expert_offset=0, + local_num_experts=w[1], + topk_ids=ids, + topk_weights=wts, + **_window_scalars(s1, sg, s2, w[0], w[1]), + )[0].float() + for w, wa in zip(windows, win_args) + ).to(torch.bfloat16) + full_ref = _ref_moe(x, ids, wts, ref, g2).to(torch.bfloat16) + # The element metric is degenerate for the first two — they write into rows + # the correct window leaves at exactly zero — so gate on relative RMS, + # which the correct window holds under 1 ulp. + for name, bad, against in ( + ("offset left at 0", wrong_offset, good_ref), + ("neighbouring window's weights", wrong_weights, good_ref), + ("all four windows at offset 0", all_at_zero, full_ref), + ): + elt, rms = _dev(bad, against) + assert rms > 25.0, f"{name} only {rms:.1f} rms ulp away — not discriminated" + print(f" {name:31s} {elt:9.3g} elt / {rms:6.1f} rms ulp") + + # The window is a plain range test, not a validated partition: local slot i + # answers for global id off+i, and off + local_num_experts > num_experts is + # accepted — the surplus local slots are simply never addressed. + over = 60 + y = _call( + data, + sf, + win_args[3], + num_experts, + top_k, + local_expert_offset=over, + local_num_experts=shard, + topk_ids=ids, + topk_weights=wts, + **_window_scalars(s1, sg, s2, windows[3][0], shard), + )[0] + _assert_moe_close( + y, + _ref_moe(x, ids, wts, win_ref[3], g2, offset=over, num_local=num_experts - over).to( + torch.bfloat16 + ), + ) + past = _call( + data, + sf, + win_args[3], + num_experts, + top_k, + local_expert_offset=num_experts, + local_num_experts=shard, + topk_ids=ids, + topk_weights=wts, + **_window_scalars(s1, sg, s2, windows[3][0], shard), + )[0] + assert torch.equal(past, torch.zeros_like(past)), ( + "a window entirely past the routing space returned nonzero rows" + ) + print( + f" test_ep_window_18_of_72 OK (window {worst_win[0]:.2f} elt / " + f"{worst_win[1]:.2f} rms, sum {worst_sum[0]:.2f} elt / {worst_sum[1]:.2f} rms, " + f"sum vs 72-expert call {worst_vs_full[0]:.2f} elt, " + f"one window dropped >= {worst_drop:.0f} rms ulp; " + f"at T={token_counts[-1]} window {top_win[0]:.2f} elt / {top_win[1]:.2f} rms, " + f"sum {top_sum[0]:.2f} elt / {top_sum[1]:.2f} rms, " + f"one window dropped >= {top_drop:.0f} rms ulp)" + ) + + +def test_r1_geometry(): + """E=256, H=7168, I=2048, top-8: the DeepSeek-R1 routed expert layer. + + The whole certified token column in one sweep, decode-sized (1 row) through + the widest unchunked prefill a stock engine issues (8192 = trtllm's default + `max_num_tokens`). Nothing about this geometry needs padding: FC1 is + `[256, 4096, 3584]` + `[256, 4096, 448]` and FC2 `[256, 7168, 1024]` + + `[256, 7168, 128]`, and `4096 % 128`, `448 % 4`, `7168 % 128`, `128 % 4` + are all zero, so the operands go in at their nominal shapes. + + Token counts are swept in ascending order, so every count below the top + draws exactly the values it drew before. + """ + args, ref = _dsr1() + gen = ref["gen"] + num_experts, hidden, top_k = 256, 7168, 8 + worst = (0.0, 0.0) + for num_tokens in _R1_TOKENS: + g1 = 137.0 + data, sf, x = _rand_activation(num_tokens, hidden, g1, gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + g2 = _calibrate_g2(x, ids, ref) + s1, sg, s2 = _scalars(num_experts, g1, g2) + out = _call( + data, + sf, + args, + num_experts, + top_k, + topk_ids=ids, + topk_weights=wts, + output1_scale_scalar=s1, + output1_scale_gate_scalar=sg, + output2_scale_scalar=s2, + ) + assert len(out) == 1, len(out) + y = out[0] + assert y.shape == (num_tokens, hidden) and y.dtype == torch.bfloat16, y.shape + exp = _ref_moe(x, ids, wts, ref, g2).to(torch.bfloat16) + _assert_moe_close(y, exp) + top = _dev(y, exp) # after the loop: the largest token count's numbers + worst = tuple(max(a, b) for a, b in zip(worst, top)) + print( + f" test_r1_geometry OK (worst {worst[0]:.2f} elt / {worst[1]:.2f} rms ulp; " + f"at T={_R1_TOKENS[-1]} {top[0]:.2f} elt / {top[1]:.2f} rms ulp)" + ) + + +def test_ep_window_64_of_256(): + """The dep4 expert-parallel split: 256 routed experts as four 64-wide windows. + + `num_experts` stays 256 on every call — routing happens outside and every + rank sees the whole space — while `local_num_experts = 64` and + `local_expert_offset` runs 0 / 64 / 128 / 192 over the DeepSeek-R1 routed + geometry (H=7168, I=2048, top-8), across the whole certified token column. + Three things are checked: each window answers for exactly the slots routed + into it, a token routed entirely elsewhere comes back bitwise zero, and the + four bf16 outputs sum back to the 256-expert result (the value the dep4 + all-reduce adds up). + """ + args, ref = _dsr1() + num_experts, hidden, top_k, shard = 256, 7168, 8, 64 + windows = [(r * shard, shard) for r in range(4)] + win_args = [_window_args(args, off, n) for off, n in windows] + win_ref = [_window_ref(ref, off, n) for off, n in windows] + gen = torch.Generator(device=DEV).manual_seed(6464) + g1 = 137.0 + worst_win = (0.0, 0.0) + worst_sum = (0.0, 0.0) + worst_vs_full = (0.0, 0.0) + worst_drop = 1e9 + worst_zero_rows = 0 + + for num_tokens in _R1_TOKENS: + data, sf, x = _rand_activation(num_tokens, hidden, g1, gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + g2 = _calibrate_g2(x, ids, ref) + s1, sg, s2 = _scalars(num_experts, g1, g2) + full = _call( + data, + sf, + args, + num_experts, + top_k, + topk_ids=ids, + topk_weights=wts, + **_window_scalars(s1, sg, s2, 0, num_experts), + )[0] + exp_full = _ref_moe(x, ids, wts, ref, g2).to(torch.bfloat16) + _assert_moe_close(full, exp_full) + + parts = [] + top_win = (0.0, 0.0) # after the loop: the largest token count's numbers + for (off, n), wa, wr in zip(windows, win_args, win_ref): + y = _call( + data, + sf, + wa, + num_experts, + top_k, + local_expert_offset=off, + local_num_experts=n, + topk_ids=ids, + topk_weights=wts, + **_window_scalars(s1, sg, s2, off, n), + )[0] + assert y.shape == (num_tokens, hidden), y.shape + exp = _ref_moe(x, ids, wts, wr, g2, offset=off, num_local=n).to(torch.bfloat16) + _assert_moe_close(y, exp) + outside = ~((ids >= off) & (ids < off + n)).any(dim=1) + assert torch.equal(y[outside], torch.zeros_like(y[outside])), ( + f"window (off={off}, n={n}) wrote into rows whose token routes " + f"entirely elsewhere ({int(outside.sum())} such rows at T={num_tokens})" + ) + worst_zero_rows = max(worst_zero_rows, int(outside.sum())) + top_win = tuple(max(a, b) for a, b in zip(top_win, _dev(y, exp))) + worst_win = tuple(max(a, b) for a, b in zip(worst_win, _dev(y, exp))) + parts.append(y) + + summed = sum(p.float() for p in parts).to(torch.bfloat16) + _assert_moe_close(summed, exp_full) + top_sum = _dev(summed, exp_full) + worst_sum = tuple(max(a, b) for a, b in zip(worst_sum, top_sum)) + worst_vs_full = tuple(max(a, b) for a, b in zip(worst_vs_full, _dev(summed, full.float()))) + # The sum gate bites: leaving any one window out lands far outside it, + # so passing it is not something four near-arbitrary partials could do. + # Only meaningful once every window carries a settled share of the + # slots — at a handful of tokens a window can legitimately hold none. + if num_tokens >= 256: + top_drop = 1e9 # after the loop: the largest token count's figure + for drop in range(4): + partial = sum(p.float() for i, p in enumerate(parts) if i != drop) + _, rms = _dev(partial.to(torch.bfloat16), exp_full) + assert rms > 25.0, ( + f"dropping window {drop} at T={num_tokens} is only {rms:.1f} rms ulp away" + ) + worst_drop = min(worst_drop, rms) + top_drop = min(top_drop, rms) + + # Every routed slot inside one window: that window alone reproduces the + # 256-expert call bitwise, and the other three return exactly zero. + num_tokens = 64 + data, sf, x = _rand_activation(num_tokens, hidden, g1, gen) + ids = torch.stack( + [torch.randperm(shard, device=DEV, generator=gen)[:top_k] for _ in range(num_tokens)] + ).to(torch.int32) + wts = torch.rand(num_tokens, top_k, device=DEV, generator=gen).to(torch.bfloat16) + g2 = _calibrate_g2(x, ids, ref) + s1, sg, s2 = _scalars(num_experts, g1, g2) + full = _call( + data, + sf, + args, + num_experts, + top_k, + topk_ids=ids, + topk_weights=wts, + **_window_scalars(s1, sg, s2, 0, num_experts), + )[0] + for (off, n), wa in zip(windows, win_args): + y = _call( + data, + sf, + wa, + num_experts, + top_k, + local_expert_offset=off, + local_num_experts=n, + topk_ids=ids, + topk_weights=wts, + **_window_scalars(s1, sg, s2, off, n), + )[0] + if off == 0: + assert torch.equal(y, full), ( + "the window holding every routed slot did not reproduce the 256-expert call bitwise" + ) + else: + assert torch.equal(y, torch.zeros_like(y)), ( + f"window off={off} holds no routed slot but returned nonzero rows" + ) + + # Window mistakes a caller can make, against the correct window-1 result. + num_tokens = 256 + data, sf, x = _rand_activation(num_tokens, hidden, g1, gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + g2 = _calibrate_g2(x, ids, ref) + s1, sg, s2 = _scalars(num_experts, g1, g2) + off, n = windows[1] + good = _call( + data, + sf, + win_args[1], + num_experts, + top_k, + local_expert_offset=off, + local_num_experts=n, + topk_ids=ids, + topk_weights=wts, + **_window_scalars(s1, sg, s2, off, n), + )[0] + good_ref = _ref_moe(x, ids, wts, win_ref[1], g2, offset=off, num_local=n).to(torch.bfloat16) + _assert_moe_close(good, good_ref) + wrong_offset = _call( + data, + sf, + win_args[1], + num_experts, + top_k, + local_expert_offset=0, + local_num_experts=n, + topk_ids=ids, + topk_weights=wts, + **_window_scalars(s1, sg, s2, off, n), + )[0] + wrong_weights = _call( + data, + sf, + win_args[0], + num_experts, + top_k, + local_expert_offset=off, + local_num_experts=n, + topk_ids=ids, + topk_weights=wts, + **_window_scalars(s1, sg, s2, off, n), + )[0] + all_at_zero = sum( + _call( + data, + sf, + wa, + num_experts, + top_k, + local_expert_offset=0, + local_num_experts=w[1], + topk_ids=ids, + topk_weights=wts, + **_window_scalars(s1, sg, s2, w[0], w[1]), + )[0].float() + for w, wa in zip(windows, win_args) + ).to(torch.bfloat16) + full_ref = _ref_moe(x, ids, wts, ref, g2).to(torch.bfloat16) + # The element metric is degenerate for the first two — they write into rows + # the correct window leaves at exactly zero — so gate on relative RMS, + # which the correct window holds under 1 ulp. + for name, bad, against in ( + ("offset left at 0", wrong_offset, good_ref), + ("neighbouring window's weights", wrong_weights, good_ref), + ("all four windows at offset 0", all_at_zero, full_ref), + ): + elt, rms = _dev(bad, against) + assert rms > 25.0, f"{name} only {rms:.1f} rms ulp away — not discriminated" + print(f" {name:31s} {elt:9.3g} elt / {rms:6.1f} rms ulp") + print( + f" test_ep_window_64_of_256 OK (window {worst_win[0]:.2f} elt / " + f"{worst_win[1]:.2f} rms, sum {worst_sum[0]:.2f} elt / {worst_sum[1]:.2f} rms, " + f"sum vs 256-expert call {worst_vs_full[0]:.2f} elt, " + f"one window dropped >= {worst_drop:.0f} rms ulp, " + f"<= {worst_zero_rows} bitwise-zero rows per window; " + f"at T={_R1_TOKENS[-1]} window {top_win[0]:.2f} elt / {top_win[1]:.2f} rms, " + f"sum {top_sum[0]:.2f} elt / {top_sum[1]:.2f} rms, " + f"one window dropped >= {top_drop:.0f} rms ulp)" + ) + + +def test_r1_autotuner_cache_is_inert(): + """A warm autotuner profiling cache does not change the R1-geometry result. + + Every other call in this file runs with a **cold** cache, where the tuner + returns the fallback tactic. A serving engine is not cold — the PyTorch + runtime profiles this op during warm-up whenever `enable_autotuner` is set. + So this test drives the other state: it warms the cache through trtllm's + own `autotune()` context at `local_num_experts` 256 and 64, then re-runs + the whole token column for all five expert layouts and demands **bitwise** + equality with the cold results. Bitwise is the right gate here — the two + runs consume identical operands, so any tactic-dependent difference at all + is visible. + + It must run last: it leaves the tuner warm until the `clear_cache()` at the + end, and the entry's receipt is a cold-cache receipt. + """ + args, ref = _dsr1() + num_experts, hidden, top_k, shard = 256, 7168, 8, 64 + layouts = [(0, num_experts, args)] + [ + (r * shard, shard, _window_args(args, r * shard, shard)) for r in range(4) + ] + gen = torch.Generator(device=DEV).manual_seed(4242) + g1 = 137.0 + cases = {} + for num_tokens in _R1_TOKENS: + data, sf, x = _rand_activation(num_tokens, hidden, g1, gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + cases[num_tokens] = ( + data, + sf, + ids, + wts, + _scalars(num_experts, g1, _calibrate_g2(x, ids, ref)), + ) + + def run(off, n, wa, num_tokens): + data, sf, ids, wts, (s1, sg, s2) = cases[num_tokens] + return _call( + data, + sf, + wa, + num_experts, + top_k, + local_expert_offset=off, + local_num_experts=n, + topk_ids=ids, + topk_weights=wts, + **_window_scalars(s1, sg, s2, off, n), + )[0] + + tuner = AutoTuner.get() + op = "trtllm::fp4_block_scale_moe_runner" + + def entries(): + return [k for k in tuner.profiling_cache.cache if k[0] == op] + + assert not entries(), ( + "the profiling cache already holds entries for this op — the cold-cache " + "receipt every other test in this file records would be void" + ) + cold = {(off, n, t): run(off, n, wa, t) for off, n, wa in layouts for t in _R1_TOKENS} + assert not entries(), ( + "an ordinary call wrote a profiling-cache entry — calls outside " + "autotune() are not cold after all" + ) + + profiled = [] + with autotune(): + for off, n, wa in layouts: + run(off, n, wa, _R1_TOKENS[-1]) + profiled.append(tuner.stats.tuned_op_profiled_configs.get(op, 0)) + # What the tuner's key does and does not separate, counted rather than + # timed: switching from local_num_experts 256 to 64 forces a fresh sweep, + # while the three remaining local_expert_offsets profile nothing at all — + # they hit the entries the offset-0 window just wrote. + assert profiled[1] > profiled[0], ( + f"local_num_experts 64 reused the 256-expert tuning: {profiled}" + ) + assert profiled[2:] == [profiled[1]] * 3, ( + f"a 64-wide window at a nonzero local_expert_offset re-profiled: {profiled}" + ) + keys = entries() + assert keys, "autotune() recorded nothing for this op — nothing was warmed" + tactics = {str(tuner.profiling_cache.cache[k][1]) for k in keys} + assert "-1" not in tactics, ( + f"the tuner recorded the fallback tactic, so warm and cold are the same " + f"execution and this comparison proves nothing: {sorted(tactics)}" + ) + unique_ids = sorted({str(k[2]) for k in keys}) + + for off, n, wa in layouts: + for num_tokens in _R1_TOKENS: + warm = run(off, n, wa, num_tokens) + ref_cold = cold[(off, n, num_tokens)] + assert torch.equal(warm, ref_cold), ( + f"a warm autotuner cache changed the result at " + f"local_expert_offset={off}, local_num_experts={n}, " + f"T={num_tokens}: {_dev(warm, ref_cold.float())} (elt / rms ulp)" + ) + tuner.clear_cache() + print( + f" test_r1_autotuner_cache_is_inert OK ({len(keys)} tuned entries over " + f"unique_ids {unique_ids}, {len(tactics)} distinct tactics, " + f"profiled-config counts {profiled} across the five layouts; all " + f"{len(layouts) * len(_R1_TOKENS)} (expert layout, T) results bitwise " + f"unchanged)" + ) + + +def test_do_finalize_and_output_buffer(): + """Both output modes, and what `do_finalize=False` actually returns.""" + num_experts, hidden, inter, top_k, num_tokens = 8, 256, 128, 2, 16 + args, ref = _build(num_experts, hidden, inter, seed=707) + gen = ref["gen"] + g1, g2 = 137.0, 1.0 / 32.0 + data, sf, x = _rand_activation(num_tokens, hidden, g1, gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + s1, sg, s2 = _scalars(num_experts, g1, g2) + kw = dict( + topk_ids=ids, + topk_weights=wts, + output1_scale_scalar=s1, + output1_scale_gate_scalar=sg, + output2_scale_scalar=s2, + ) + y = _call(data, sf, args, num_experts, top_k, **kw)[0] + + # caller-provided buffer: written in place, [0] returned, tail untouched + buf = torch.full((num_tokens + 3, hidden), 7.0, dtype=torch.bfloat16, device=DEV) + r = _call(data, sf, args, num_experts, top_k, **kw, output=buf[:num_tokens]) + assert len(r) == 1 and r[0].shape == (0,) and r[0].dtype == torch.bfloat16 + assert torch.equal(buf[:num_tokens], y) + assert bool((buf[num_tokens:] == 7.0).all()), "wrote past the output rows" + + for bad, msg in ( + (torch.empty(num_tokens, hidden, dtype=torch.float32, device=DEV), "bfloat16"), + (torch.empty(num_tokens + 1, hidden, dtype=torch.bfloat16, device=DEV), "dim0"), + ): + try: + _call(data, sf, args, num_experts, top_k, **kw, output=bad) + raise AssertionError(f"bad output buffer ({msg}) accepted") + except RuntimeError as e: + assert msg in str(e), e + + part, scales, permuted = _call(data, sf, args, num_experts, top_k, **kw, do_finalize=False) + assert part.dtype == torch.bfloat16 and part.shape[1] == hidden + assert part.shape[0] >= num_tokens * top_k + assert scales.shape == (num_tokens, top_k) and scales.dtype == torch.bfloat16 + assert permuted.shape == (num_tokens, top_k) and permuted.dtype == torch.int32 + assert int(permuted.min()) >= 0 and int(permuted.max()) < part.shape[0] + comb = torch.zeros(num_tokens, hidden, dtype=torch.float32, device=DEV) + for j in range(top_k): + comb += part[permuted[:, j].long()].float() * wts[:, j].float().unsqueeze(1) + assert torch.equal(comb.to(torch.bfloat16), y), ( + "gathering per-slot rows by the third output and weighting them with the " + "caller's topk_weights did not reproduce the finalized result" + ) + try: + _call( + data, + sf, + args, + num_experts, + top_k, + **kw, + do_finalize=False, + output=buf[:num_tokens], + ) + raise AssertionError("output= with do_finalize=False accepted") + except RuntimeError as e: + assert "only supported when do_finalize=true" in str(e), e + print(" test_do_finalize_and_output_buffer OK") + + +def test_inert_and_rejected_arguments(): + """Knobs that change nothing on the pre-routed path, and hard rejections.""" + num_experts, hidden, inter, top_k, num_tokens = 8, 256, 128, 2, 16 + args, ref = _build(num_experts, hidden, inter, seed=808) + gen = ref["gen"] + g1, g2 = 137.0, 1.0 / 32.0 + data, sf, x = _rand_activation(num_tokens, hidden, g1, gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + s1, sg, s2 = _scalars(num_experts, g1, g2) + kw = dict( + topk_ids=ids, + topk_weights=wts, + output1_scale_scalar=s1, + output1_scale_gate_scalar=sg, + output2_scale_scalar=s2, + ) + base = _call(data, sf, args, num_experts, top_k, **kw)[0] + assert torch.equal(_call(data, sf, args, num_experts, top_k, **kw)[0], base), ( + "two identical calls are not bitwise equal" + ) + + inert = ( + [dict(routing_method_type=r) for r in (0, 1, 2, 4, 5, 6)] + + [ + dict(routing_method_type=2, n_group=n, topk_group=t) + for n, t in ((1, 1), (8, 4), (4, 2)) + ] + + [dict(n_group=None, topk_group=None), dict(n_group=1, topk_group=1)] + + [dict(routed_scaling_factor=r) for r in (None, 1.0, 2.5)] + + [dict(tune_max_num_tokens=t) for t in (8192, 128)] + + [dict(use_dp=u) for u in (False, True)] + + [dict(routing_bias=torch.zeros(num_experts, dtype=torch.bfloat16, device=DEV))] + ) + for over in inert: + y = _call(data, sf, args, num_experts, top_k, **kw, **over)[0] + assert torch.equal(y, base), f"{over} changed the pre-routed result" + + # n_group > 1 needs routing_method_type 2 even though it is then inert + try: + _call(data, sf, args, num_experts, top_k, **kw, n_group=8, topk_group=4) + raise AssertionError("n_group>1 with routing_method_type=1 accepted") + except RuntimeError as e: + assert "DeepSeekV3 routing method" in str(e), e + + rejections = ( + (dict(topk_weights=wts.float()), "topk_weights must be bfloat16"), + (dict(topk_weights=wts.half()), "topk_weights must be bfloat16"), + (dict(topk_ids=ids.long()), "topk_ids must be int"), + (dict(topk_ids=None, topk_weights=None), "must be provided"), + (dict(intermediate_size=inter // 2), "incorrect dim 2"), + (dict(num_experts_override=top_k), "num_experts must be greater than top_k"), + ) + for over, msg in rejections: + ne = over.pop("num_experts_override", num_experts) + try: + _call(data, sf, args, ne, top_k, **{**kw, **over}) + raise AssertionError(f"expected rejection: {msg}") + except RuntimeError as e: + assert msg in str(e), (msg, str(e)) + + e4m3 = sf.view(torch.float8_e4m3fn) + for bad, msg in ( + (sf.view(torch.uint8), "must be fp8"), + (e4m3.reshape(num_tokens, -1), "must be 1D"), + (e4m3[:-1], "incorrect size"), + ): + try: + _call(data, bad, args, num_experts, top_k, scale_as_is=True, **kw) + raise AssertionError(f"expected rejection: {msg}") + except RuntimeError as e: + assert msg in str(e), (msg, str(e)) + for name in ("gemm1_weights_scale", "gemm2_weights_scale"): + try: + _call( + data, + sf, + args, + num_experts, + top_k, + **kw, + **{name: args[name].view(torch.uint8)}, + ) + raise AssertionError(f"{name} as uint8 accepted") + except RuntimeError as e: + assert "must be fp8" in str(e), e + try: + _call(data.view(torch.float8_e4m3fn), sf, args, num_experts, top_k, **kw) + raise AssertionError("hidden_states as float8_e4m3fn accepted") + except RuntimeError as e: + assert "must be byte" in str(e), e + print(" test_inert_and_rejected_arguments OK") + + +def test_hidden_must_be_a_multiple_of_256(): + """A hidden size that is not a multiple of 256 is silently wrong. + + The fault is in the weight path, not the activation one: it survives an + activation whose block scales are all exactly 1.0. The wrapper guards it. + """ + num_experts, inter, top_k, num_tokens = 8, 128, 2, 8 + for hidden in (256, 384, 512, 640, 768, 896): + args, ref = _build(num_experts, hidden, inter, seed=1000 + hidden) + gen = ref["gen"] + g1 = 71.0 + data, sf, x = _rand_activation(num_tokens, hidden, g1, gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + g2 = _calibrate_g2(x, ids, ref) + s1, sg, s2 = _scalars(num_experts, g1, g2) + exp = _ref_moe(x, ids, wts, ref, g2).to(torch.bfloat16) + if hidden % 256 == 0: + y = _call( + data, + sf, + args, + num_experts, + top_k, + topk_ids=ids, + topk_weights=wts, + output1_scale_scalar=s1, + output1_scale_gate_scalar=sg, + output2_scale_scalar=s2, + )[0] + _assert_moe_close(y, exp) + continue + try: + _call( + data, + sf, + args, + num_experts, + top_k, + topk_ids=ids, + topk_weights=wts, + output1_scale_scalar=s1, + output1_scale_gate_scalar=sg, + output2_scale_scalar=s2, + ) + raise AssertionError(f"wrapper accepted hidden={hidden}") + except AssertionError as e: + assert "multiple of 256" in str(e), e + raw = torch.ops.trtllm.fp4_block_scale_moe_runner( + None, + None, + data, + sf.view(torch.float8_e4m3fn), + args["gemm1_weights"], + args["gemm1_weights_scale"], + None, + None, + None, + None, + args["gemm2_weights"], + args["gemm2_weights_scale"], + None, + s1, + sg, + s2, + num_experts, + top_k, + None, + None, + inter, + 0, + num_experts, + None, + 1, + True, + 0, + topk_weights=wts, + topk_ids=ids, + )[0] + elt, _ = _dev(raw, exp) + assert elt > 25.0, f"hidden={hidden} unexpectedly correct ({elt:.1f} ulp)" + print(f" hidden={hidden}: accepted by the op, {elt:.0f} ulp wrong") + print(" test_hidden_must_be_a_multiple_of_256 OK") + + +def test_strided_inputs_are_silently_wrong(): + """Justifies the wrapper's contiguity asserts: the kernel ignores strides.""" + num_experts, hidden, inter, top_k, num_tokens = 8, 256, 128, 2, 16 + args, ref = _build(num_experts, hidden, inter, seed=909) + gen = ref["gen"] + g1, g2 = 137.0, 1.0 / 32.0 + data, sf, x = _rand_activation(num_tokens, hidden, g1, gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + s1, sg, s2 = _scalars(num_experts, g1, g2) + kw = dict( + topk_ids=ids, + topk_weights=wts, + output1_scale_scalar=s1, + output1_scale_gate_scalar=sg, + output2_scale_scalar=s2, + ) + base = _call(data, sf, args, num_experts, top_k, **kw)[0] + for name, strided in ( + ("hidden_states", torch.stack([data, data], -1)[..., 0]), + ("topk_weights", torch.stack([wts, wts], -1)[..., 0]), + ("gemm1_weights", torch.stack([args["gemm1_weights"]] * 2, -1)[..., 0]), + ): + assert not strided.is_contiguous() + try: + moe_kw = dict(kw) + if name == "hidden_states": + y = _call(strided, sf, args, num_experts, top_k, **moe_kw)[0] + else: + moe_kw[name] = strided + y = _call(data, sf, args, num_experts, top_k, **moe_kw)[0] + except AssertionError: + continue # the wrapper's guard fired, which is the point + raise AssertionError( + f"wrapper did not reject a strided {name} (result equal: {torch.equal(y, base)})" + ) + + # and the raw op does accept them, silently + raw = torch.ops.trtllm.fp4_block_scale_moe_runner( + None, + None, + torch.stack([data, data], -1)[..., 0], + sf.view(torch.float8_e4m3fn), + args["gemm1_weights"], + args["gemm1_weights_scale"], + None, + None, + None, + None, + args["gemm2_weights"], + args["gemm2_weights_scale"], + None, + s1, + sg, + s2, + num_experts, + top_k, + None, + None, + inter, + 0, + num_experts, + None, + 1, + True, + 0, + topk_weights=wts, + topk_ids=ids, + )[0] + assert not torch.equal(raw, base), ( + "a strided hidden_states no longer changes the result — re-check the guard" + ) + print(" test_strided_inputs_are_silently_wrong OK") diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/fused_moe.md b/tensorrt_llm/_torch/staircase/catalog/moe/fused_moe.md new file mode 100644 index 000000000000..c125b64334d2 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/moe/fused_moe.md @@ -0,0 +1,427 @@ +--- +receipts: + sm_100: {status: passed, trtllm: 1.3.0rc21} + sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 23} +--- + +# fused_moe + +**Wraps** `torch.ops.trtllm.fused_moe` (one call). + +## Semantics + +One complete mixture-of-experts layer over already-routed tokens: expert +permutation, the grouped FC1 GEMM, the gated activation, the grouped FC2 +GEMM, and the routing-weighted combine back to one row per token — all +inside a single call. + +Let `T = input.shape[0]`, `H = hidden size`, `I = intermediate size per +expert`, `K = top-k`, `E = fc1_expert_weights.shape[0]` (the experts this +rank holds) and `first = ep_rank * E` (this rank's first *global* expert +id). For every token `t` and slot `j < K`: + +``` +g = token_selected_experts[t, j] # a GLOBAL expert id +skip this slot unless first <= g < first + E +e = g - first # index into the weight tensors + +up = input[t] @ fc1_expert_weights[e, :I, :].T (+ fc1_expert_biases[e, :I]) +gate = input[t] @ fc1_expert_weights[e, I:, :].T (+ fc1_expert_biases[e, I:]) + +# activation_type 5 (Swiglu) and 7 (SwigluBias) — identical on this path: +gate = min(gate, swiglu_limit[e]) # only when swiglu_limit given +up = clamp(up, -swiglu_limit[e], swiglu_limit[e]) +act = gate * sigmoid(swiglu_alpha[e] * gate) * (up + swiglu_beta[e]) + # swiglu_alpha defaults to 1.0, swiglu_beta to 0.0, no clamp by default, + # so the default is plain SwiGLU: silu(gate) * up +# activation_type 6 (Geglu): +act = gelu(gate) * up + +act = round_to(input.dtype)(act) # materialized between the GEMMs +y = act @ fc2_expert_weights[e].T (+ fc2_expert_biases[e]) +out[t] += token_final_scales[t, j] * y # fp32 accumulation +``` + +and finally `out` is cast to `output_dtype`. Both GEMMs accumulate in +fp32; only the FC1 activation result and the final store are in the +low-precision dtype. `token_final_scales=None` means every selected slot +combines with weight `1.0`. + +**FC1 row order is `[up | gate]`.** The first `I` rows of +`fc1_expert_weights[e]` are the up projection (trtllm's `w3`, HF's +`up_proj`) and the last `I` rows are the gate projection (trtllm's `w1`, +HF's `gate_proj`) — the activation's sigmoid is applied to the *second* +half. `fc1_expert_biases[e]` is split the same way. Swapping the halves +produces a plausible-looking, entirely different result; nothing detects +it. + +**Expert parallelism.** `ep_size`/`ep_rank` do not slice anything: they +only shift which global ids this rank answers for. This rank owns +`[ep_rank * E, (ep_rank + 1) * E)`; slots routed anywhere else contribute +nothing to its output, so summing the outputs of all `ep_size` ranks +reproduces the single-rank result. Certified at `ep_size = 2` with +`E = 8`, and at `ep_size = 4` with `E = 64` over a 256-expert global stack +(all four windows driven, each against its own reference, and their sum +against a 256-expert native-torch reference — see *Notes* for the +numbers). A rank whose window catches no slot at all returns exactly +zero. With `ep_size = 1` (the default) the ids are plain indices into the +weight tensors. `tp_size`/`tp_rank` are a separate axis and are **not** +certified here (held at `1`/`0`): under tensor parallelism a rank passes +its own `I/tp_size` slice of both weight tensors and the reduction across +ranks happens outside this call. + +**No routing, no reduction.** The router GEMM and the top-k selection that +produce `token_selected_experts` / `token_final_scales` are the caller's +(a sibling op, `torch.ops.trtllm.renorm_moe_routing_op`, produces exactly +that pair), as are any shared/dense expert branch, the TP all-reduce or EP +gather of this call's output, and the residual add. Nothing is +renormalized inside: the combine weights are used exactly as given (they +need not sum to 1, and may be negative). + +**Output.** The call returns a Python list. With `out_tensor=None` it is +`[out]`, one freshly allocated contiguous `[T, H]` tensor in +`output_dtype`. With `out_tensor` given the result is written into that +buffer (every element overwritten, nothing outside its rows touched) and +the returned list is **empty** — the caller uses its own buffer. Inputs +are never mutated. + +## Signature + +```python +def fused_moe( + input: torch.Tensor, + token_selected_experts: torch.Tensor, + token_final_scales: Optional[torch.Tensor], + fc1_expert_weights: torch.Tensor, + fc1_expert_biases: Optional[torch.Tensor], + fc2_expert_weights: torch.Tensor, + fc2_expert_biases: Optional[torch.Tensor], + output_dtype: torch.dtype, + quant_scales: List[torch.Tensor], + input_sf: Optional[torch.Tensor] = None, + swizzled_input_sf: bool = True, + swiglu_alpha: Optional[torch.Tensor] = None, + swiglu_beta: Optional[torch.Tensor] = None, + swiglu_limit: Optional[torch.Tensor] = None, + tp_size: int = 1, tp_rank: int = 0, + ep_size: int = 1, ep_rank: int = 0, + cluster_size: int = 1, cluster_rank: int = 0, + enable_alltoall: bool = False, + use_deepseek_fp8_block_scale: bool = False, + use_w4_group_scaling: bool = False, + use_int8_woq_per_channel: bool = False, + use_mxfp8_act_scaling: bool = False, + min_latency_mode: bool = False, + use_fused_finalize: bool = True, + tune_max_num_tokens: int = 8192, + tuner_num_tokens: Optional[int] = None, + tuner_top_k: Optional[int] = None, + activation_type: int = 5, + unpadded_hidden_size: Optional[int] = None, + out_tensor: Optional[torch.Tensor] = None, + use_dynamic_fc2_scale: bool = False, + use_mxfp8_weight_scaling: bool = False, + # routed-expert LoRA, per-request and slot-indexed families + fc1_lora_ranks=None, fc1_lora_weight_ptrs=None, + fc2_lora_ranks=None, fc2_lora_weight_ptrs=None, + gated_lora_ranks=None, gated_lora_weight_ptrs=None, + host_request_types=None, host_context_lengths=None, + lora_max_low_rank: int = 0, + fc1_slot_lora_ranks=None, fc1_slot_lora_weight_ptrs=None, + fc2_slot_lora_ranks=None, fc2_slot_lora_weight_ptrs=None, + gated_slot_lora_ranks=None, gated_slot_lora_weight_ptrs=None, + token_to_slot=None, +) -> List[torch.Tensor] +``` + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `input` | `[T, H]` | bf16 or fp16 | contiguous | CUDA | +| `token_selected_experts` | `[T, K]` | **int32** | contiguous | CUDA | +| `token_final_scales` | `[T, K]` or `None` | **fp32** | contiguous | CUDA | +| `fc1_expert_weights` | `[E, 2I, H]`, rows `[up; gate]` | = `input.dtype` | contiguous | CUDA | +| `fc1_expert_biases` | `[E, 2I]` or `None` | = `input.dtype` | contiguous | CUDA | +| `fc2_expert_weights` | `[E, H, I]` | = `input.dtype` | contiguous | CUDA | +| `fc2_expert_biases` | `[E, H]` or `None` | = `input.dtype` | contiguous | CUDA | +| `output_dtype` | scalar | must equal `input.dtype` here | — | — | +| `quant_scales` | `[]` (empty list) | — | — | — | +| `swiglu_alpha` / `swiglu_beta` / `swiglu_limit` | `[E]` or `None` | **fp32** | contiguous | CUDA | +| `ep_size` / `ep_rank` | scalars | Python int, `0 <= ep_rank < ep_size` | — | — | +| `use_fused_finalize` | scalar | bool — both values accepted, giving identical bits; not switchable per call (see *Metadata consumed*) | — | — | +| `activation_type` | scalar | int: 5 (Swiglu), 6 (Geglu), 7 (SwigluBias) | — | — | +| `out_tensor` | `[T, H]` or `None` | = `output_dtype` | contiguous | CUDA | +| returns `[0]` (absent when `out_tensor` given) | `[T, H]` | `output_dtype` | contiguous, newly allocated | same as `input` | + +Biases are all-or-nothing: pass both or neither (see *Preconditions*). + +### Arguments held inert (not certified) + +`input_sf`, `swizzled_input_sf`, `use_deepseek_fp8_block_scale`, +`use_w4_group_scaling`, `use_int8_woq_per_channel`, +`use_mxfp8_act_scaling`, `use_mxfp8_weight_scaling`, +`use_dynamic_fc2_scale` — the quantized-activation/weight paths (fp8 +per-tensor and block scale, NVFP4, MXFP4/MXFP8, INT8 weight-only), which +also populate `quant_scales`. `min_latency_mode` — NVFP4-only (it is +rejected here) and returns a different tensor list. +`enable_alltoall`, `tuner_num_tokens`, +`tuner_top_k`, `cluster_size`, `cluster_rank` — multi-rank +dispatch/smart-router paths. `tune_max_num_tokens` — an autotuner bucket +cap; left at its default. `unpadded_hidden_size` — for callers that padded +`H`; left at `None`. The whole LoRA family (`fc1_lora_*`, `fc2_lora_*`, +`gated_lora_*`, `host_request_types`, `host_context_lengths`, +`lora_max_low_rank`, `*_slot_lora_*`, `token_to_slot`). `tp_size` / +`tp_rank` are forwarded but were only exercised at `1` / `0`. + +## Metadata consumed + +None. The op reads no attention metadata, no KV cache, no registered +layer, and no module state — every tensor and every scalar it uses is an +argument. + +Two process-global caches sit behind it. The runner cache is result-neutral; +the tactic cache is not: + +- an **autotuner profiling cache** that picks the two GEMM tactics. Cold — + the state most calls certified here run in — both tactics are the `-1` + fallback, and nothing is printed (the cache-miss message is a + `warning_once`, below trtllm's default `error` severity). + **Tactic choice is not bitwise neutral.** After one `autotune()` + pass the same inputs at `T=256, H=512, I=256, E=32, K=4` came back + different in 83447 of 131072 elements (2.96 ulp of the token row's largest + magnitude), and clearing the cache restored the cold bits exactly. It is a + different accumulation order, not a different computation — cold and hot + both pass this entry's gate against the torch reference (2.4 and 2.0 ulp). + Almost all of the spread is the FC2 GEMM: over the tactic space at that + shape, `gemm1`'s 209 tactics yield 2 distinct outputs, `gemm2`'s 309 yield + 101. **Serving runs on the hot side**, so a target exercises tactics a cold + receipt never saw: `PyExecutor` warmup calls `_run_autotuner_warmup` + (`_torch/pyexecutor/model_engine.py:1103,1278`), which wraps a real forward + in `autotune()` (`:1293`) unless `enable_autotuner=False` (default `True`, + `llmapi/llm_args.py:4669`). + + **The hot side is certified at `E=64, H=7168, I=2048, K=8`**, two ways: + + - The *whole tactic space*, driven directly. The tactic population is a + property of the build, not of the geometry — the count comes from + `get_tactic_num(gemm_idx)`, which takes no shape argument — and that was + measured rather than assumed: the tactic space captured at + `E=64, H=7168, I=2048` has exactly the same size as one captured at + `E=32, H=512, I=256`, 209 x 309 = 64581 configurations. Every one of + `gemm2`'s 309 tactics and `gemm1`'s 209 passes this entry's gate at + `T=256`, and all 309 `gemm2` tactics pass at `T=8192`; none failed to + run. The distinct-output counts match the small shape's exactly — 101 + from `gemm2`, 2 from `gemm1` — and the worst deviation over the whole + space is **3.71 ulp** element-wise at `T=256` and **3.98 ulp** at + `T=8192` (relative RMS 1.42 both), so the 8 / 4 ulp bounds below hold at + this cell with 2.0x and 2.8x margin. + - One real `autotune()` pass at that geometry: it fills exactly **28** + cache entries (2 tunable GEMMs x 14 power-of-2 token buckets, 1 … 8192), + every one carrying a profiled tactic rather than the `-1` fallback. The + warm results move the bits (36172254 of 58720256 elements at `T=8192`) + and still pass at 3.88 ulp; clearing the cache restores the cold bits + exactly at every token count. The pass costs ~70 s at this geometry. +- a cache of C++ `FusedMoeRunner` objects keyed by + `(activation dtype, weight dtype, output dtype, quant flags)`, each + holding a GPU **workspace** buffer that is allocated on first use and + grown as needed. It is never freed by the op; + `tensorrt_llm._torch.custom_ops.torch_custom_ops.MoERunner.clear_all_workspaces()` + releases the workspaces of every cached runner. **`use_fused_finalize` is + baked into that object** and is not part of the key, so the first call in + a process for a given dtype set fixes it for every later call — a caller + cannot flip it per call, and the entry's own test therefore runs on a + runner built with `True`. On this bf16 path the value makes no difference + either way: two fresh processes driving `True` then `False`, and `False` + then `True`, produced four bitwise identical outputs at + `T=2048, H=512, I=256, E=32, K=8` (a probe, not part of the test), and + `True` and `False` agree bitwise inside the test at + `T=2048, E=64, H=7168, I=2048, K=8`. + +Two identical calls return bitwise identical results **within one tuner +state** (verified at `T=256, H=512, I=256, E=32, K=4`, cold and hot +separately; and at `E=64, H=7168, I=2048, K=8` for `T` in {1, 256, 8192}, +where a cold call repeated after an intervening tuning pass and a cache +clear returned the same bits); across a cold/hot change they differ by the +ulp above. + +## Preconditions + +Shapes and layout: + +- `input` is a **2D contiguous CUDA** tensor with `T >= 1`. 3D input raises + `input must be 2D`; a strided view raises `input must be contiguous`; a + CPU tensor raises `input must be a CUDA tensor`; `T == 0` fails a C++ + assertion (`Assertion failed: input_activations`). +- `H % 8 == 0` and `I % 8 == 0`, else `hidden_size must be divisible + by 8 for weights` / `inter_size must be divisible by 8 for weights`. + `H = 8` and `I = 8` work. +- `fc1_expert_weights` is `[E, 2I, H]` and `fc2_expert_weights` is + `[E, H, I]`, both 3D, contiguous, CUDA, same dtype as `input`, same `E`. + Violations raise (`must be 3D`, `must be contiguous`, + `must be a CUDA tensor`, `fc1_expert_weights and fc2_expert_weights must + have the same number of experts.`). Certified: `E` in + {1, 2, 7, 8, 16, 32, **64**, 128, 256} and `(H, I)` in {(8,8), (16,8), + (128,64), (192,96), (256,128), (512,256), (1024,512), (2048,768), + **(7168,2048)**}. The two enumerations are not a product: `(7168, 2048)` + was driven at `E = 64` only, and `E = 64` only at `(7168, 2048)`. Nothing + in the permutation or gather path keys on the expert count, or on which + index a given expert sits at: permuting the 64-expert stack and + relabelling the ids to match reproduces the output **bitwise**. +- `token_selected_experts` is `[T, K]` **int32** contiguous CUDA; + `token_final_scales`, when given, is `[T, K]` **fp32** contiguous CUDA + with the same `K`. Other dtypes raise (`token_selected_experts dtype is + Long, while Int is expected`, `token_final_scales.value() dtype is + BFloat16, while Float is expected`), as do mismatched `T` or `K`. + Certified `K`: 1, 2, 3, 4, 8, 16. **A token's `K` ids must be distinct + when `T > 256`** — a repeat there is accepted and then computed wrong, or + faults; see *Notes*. At `T <= 256` a repeat is honoured once per slot, so + that expert's output is added twice with its two combine weights. +- Biases are optional but **paired**: passing exactly one of + `fc1_expert_biases` / `fc2_expert_biases` raises + `RuntimeError: bad optional access`. Both must be `[E, 2I]` and `[E, H]` + contiguous CUDA tensors in `input.dtype` (fp32/fp16 bias against a bf16 + input raises). +- `out_tensor`, when given, must be a contiguous CUDA tensor of shape + exactly `[T, H]` in `output_dtype`; wrong dtype, wrong shape, or a + strided view each raise. A leading row-slice of a larger contiguous + buffer is a valid `out_tensor` — rows past `T` are left bitwise + untouched. + +Dtypes and modes: + +- Activations and weights must be the **same** dtype: bf16 or fp16. + fp32 activations, or weights whose dtype differs from `input`, raise + `Could not construct fused moe op with the requested input combination + ...`. +- On this unquantized path `output_dtype` **must equal** `input.dtype`. + A mismatch does **not** raise — see *Notes*. The wrapper asserts it. +- `quant_scales` must carry the scale set of whatever quantization is + requested; on the unquantized path pass `[]` (a non-empty list is + accepted and ignored). fp8-e4m3 inputs with `[]` raise `Expecting 4 + quant scales for fp8 quantization`. +- `swiglu_alpha` / `swiglu_beta` / `swiglu_limit`, when given, are fp32 + CUDA tensors with exactly `E` elements (`swiglu_alpha must have + num_experts_on_rank elements.`, `... must be a CUDA tensor`, `... dtype + is BFloat16, while Float is expected`). +- `activation_type` must be a gated type (5 Swiglu, 6 Geglu, 7 + SwigluBias) for the `[E, 2I, H]` FC1 layout. The non-gated values of the + same enum (1 Identity, 2 Gelu, 3 Relu, 4 Silu, 8 Relu2) — and any + unrecognised value — expect `[E, I, H]` instead, so with a gated layout + they raise `fc1_expert_weights inter size must be equal to + fc2_expert_weights inter size.` (verified for 1, 2, 4, 8 and 99). +- `0 <= ep_rank < ep_size` (`Assertion failed: ep_rank < ep_size`). +- `min_latency_mode=True` is NVFP4-only and fails here with `Assertion + failed: use_fp4 == true`. `cluster_size > 1` raises `smart_router is + supported in min_latency mode`. +- `tuner_num_tokens` and `tuner_top_k` must both be `None` unless + `enable_alltoall=True`, and both must be set when it is; the op's own + Python `assert` fires otherwise (bare `AssertionError`). +- Any LoRA tensor implies `lora_max_low_rank > 0` (`MoE LoRA requires + lora_max_low_rank > 0; got 0`), and the per-request and slot-indexed + families are mutually exclusive. + +A caller violating none of the above gets the result described under +*Semantics*, to within a few low-precision ulps (see *Notes*). + +## Notes + +- **`output_dtype != input.dtype` is a silent wrong answer.** On the + unquantized path the output buffer is allocated with `output_dtype` + while the epilogue stores elements in the activation dtype, so the + caller reads reinterpreted bits: bf16 input with `output_dtype=fp32` + returns finite, plausibly-scaled values that are pairs of bf16 results + glued into fp32 words (element `j` carries roughly the true element + `2j+1`). No error, no NaN. `use_fused_finalize=False` behaves the same. + The wrapper carries the one guard this entry needs: when + `input.dtype in (bf16, fp16)` and `fc1_expert_weights.dtype == + input.dtype`, `output_dtype` must equal `input.dtype`. The guard is + scoped so it never fires on the quantized paths, where a differing + `output_dtype` is the norm. +- **Expert ids outside this rank's range are silently dropped**, not + clamped and not rejected: with `ep_size=1, E=8`, ids `8`, `100` and `-1` + each contribute nothing, and the token's other slots still combine + normally. A routing bug therefore shows up as a quietly weaker token, + never as an error. Checking ids costs a device-to-host sync, so the + wrapper does not. +- **Three quantization flags are accepted and ignored** on the bf16 path: + `use_w4_group_scaling`, `use_mxfp8_act_scaling` and + `use_dynamic_fc2_scale` return bitwise the unquantized result. The other + three (`use_deepseek_fp8_block_scale`, `use_int8_woq_per_channel`, + `use_mxfp8_weight_scaling`) raise. +- **Numerics.** Against a native-torch reference that consumes the same + low-precision operands (fp32 GEMM accumulation, the FC1 activation + rounded to the input dtype, fp32 combine), the kernel's worst + element-wise deviation over the covered configurations was **2.8 ulp of + the token row's largest magnitude** (bf16 ulp = 2^-8, fp16 ulp = 2^-11) + and its relative RMS deviation **1.2 ulp**. Both were calibrated on the + fallback tactic; walking the *whole* tactic space reaches **3.98 ulp** + element-wise and **1.43 ulp** RMS (`gemm2`, `T = 8192`, at + `E=64, H=7168, I=2048`) — still inside the test's 8 ulp / 4 ulp gates, + with 2.0x and 2.8x margin. That is rounding noise, not + a different computation: the same comparison against a gate/up-swapped + reference lands at 176+ ulp element-wise and ~190 ulp RMS at the small + shapes and 284 / 197 at `E=64, H=7168, I=2048`, and against a + one-expert-dropped reference at 997 / 212 there (measured while + calibrating the bound; the entry's test asserts only that such a + reference is rejected). Default `torch.testing.assert_close` tolerances + do not fit this op — their bf16 `atol` of 1e-5 is three orders of + magnitude below one output ulp of a two-GEMM chain. +- **Scale is inert.** `(H, I) = (7168, 2048)` at `E = 64` is 3.5x the + hidden size and 2.7x the intermediate size of anything else certified + here, and behaves exactly like the small shapes: same tactic population, + same distinct-output counts, same ulp band, no tactic failing to run. + The deviation grows only with token count, and only slightly — 1.12 ulp + at `T = 1` to 2.54 at `T = 8192`, the same trend the smaller shapes show. +- **Four-way expert parallelism**, measured at `E = 64` per rank over a + 256-expert global stack with global ids in `[0, 256)`, at + `T` in {1, 1024, 8192}: each rank's output matches its own reference to + at most 2.69 ulp, and the four outputs summed in fp32 match a 256-expert + native-torch reference to 3.53 ulp element-wise / 1.34 ulp RMS. The sum + sits slightly wider than any single rank because each rank rounds its own + partial to bf16 before the add. A rank whose window catches no slot + returns exactly zero, and the same weights driven under the wrong + `ep_rank` land 90x outside the RMS gate — so the tiling check is a real + check and not a formality. +- **Repeated expert ids in one token's row are a defect past `T = 256`, + not a supported input.** At `T <= 256` a repeat is honoured once per + slot, as *Preconditions* says. At `T >= 257` the op accepts the input and + then either returns values ~200-460 ulp wrong on essentially every row, + or dies with `CUDA error: an illegal memory access was encountered` from + inside the call. Which of the two happens depends on whether the + out-of-bounds address is mapped; both were seen at the same shape. + `compute-sanitizer` names the fault: an invalid 4-byte global read in + `tensorrt_llm::_v1::kernels::cutlass_kernels::finalizeMoeRoutingKernel`, + launched from `finalizeMoeRoutingKernelLauncher` inside + `CutlassMoeFCRunner<...>::gemm2`. The boundary is exactly 256 tokens + (`T = 256` clean, `T = 257` faults) and matches the largest block size the + single-CTA expert-map builder is instantiated for — + `fusedBuildExpertMapsSortFirstTokenDispatch<256, …>` in + `libth_common.so`, whose fallback is + `threeStepBuildExpertMapsSortFirstToken`. `use_fused_finalize` does not + change it. This is **not specific to the new geometry**: reproduced at + `(E,H,I,K) = (8,256,128,4)`, `(32,512,256,8)` and `(128,2048,768,8)` — + cells this entry already certified — as well as at `(64,7168,2048,8)`. + Top-k selection returns distinct ids by construction, so a caller taking + its ids from a selection op cannot reach this; a caller synthesizing + `token_selected_experts` itself can. +- `activation_type=7` (SwigluBias) computed bitwise the same result as + `5` (Swiglu) for every input tried, with and without + `swiglu_alpha/beta/limit`. +- Geglu (6) matches a `gelu(gate) * up` reference; the tanh-approximate + and exact gelu variants are indistinguishable at bf16 output resolution, + so which one the kernel uses is **not** pinned by this entry. +- `T = 8192` (one token past the default `tune_max_num_tokens` bucket cap) + is certified, at `(2048, 768, E=128)` and at `(7168, 2048, E=64)`; + `T = 8193` was also observed correct during an earlier run's probing but + is not in the entry's test. +- **What `(7168, 2048)` at `E = 64` does *not* cover.** bf16 only (no + fp16), no biases, no `swiglu_alpha`/`beta`/`limit`, `activation_type = 5` + only (no Geglu, no SwigluBias), `tp_size = 1`, and `K = 8` only. Those + axes are certified at the smaller shapes and were not re-driven here. +- Sibling ops in this build address the same job differently: + `torch.ops.trtllm.moe_custom_op` takes a registered layer by string id + instead of explicit weights, and `torch.ops.trtllm.fp8_block_scale_moe_runner` + / `torch.ops.trtllm.fp4_block_scale_moe_runner` are the trtllm-gen + block-scaled MoE entry points. Catalog membership is `index.yaml`'s fact + alone. diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/fused_moe.py b/tensorrt_llm/_torch/staircase/catalog/moe/fused_moe.py new file mode 100644 index 000000000000..ca528377a573 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/moe/fused_moe.py @@ -0,0 +1,133 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Fused mixture-of-experts layer: expert permutation + grouped FC1/FC2 GEMMs +with a gated activation between them + routing-weighted combine, in one call.""" + +from typing import List, Optional + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + + +def fused_moe( + input: torch.Tensor, + token_selected_experts: torch.Tensor, + token_final_scales: Optional[torch.Tensor], + fc1_expert_weights: torch.Tensor, + fc1_expert_biases: Optional[torch.Tensor], + fc2_expert_weights: torch.Tensor, + fc2_expert_biases: Optional[torch.Tensor], + output_dtype: torch.dtype, + quant_scales: List[torch.Tensor], + input_sf: Optional[torch.Tensor] = None, + swizzled_input_sf: bool = True, + swiglu_alpha: Optional[torch.Tensor] = None, + swiglu_beta: Optional[torch.Tensor] = None, + swiglu_limit: Optional[torch.Tensor] = None, + tp_size: int = 1, + tp_rank: int = 0, + ep_size: int = 1, + ep_rank: int = 0, + cluster_size: int = 1, + cluster_rank: int = 0, + enable_alltoall: bool = False, + use_deepseek_fp8_block_scale: bool = False, + use_w4_group_scaling: bool = False, + use_int8_woq_per_channel: bool = False, + use_mxfp8_act_scaling: bool = False, + min_latency_mode: bool = False, + use_fused_finalize: bool = True, + tune_max_num_tokens: int = 8192, + tuner_num_tokens: Optional[int] = None, + tuner_top_k: Optional[int] = None, + activation_type: int = 5, # ActivationType.Swiglu + unpadded_hidden_size: Optional[int] = None, + out_tensor: Optional[torch.Tensor] = None, + use_dynamic_fc2_scale: bool = False, + use_mxfp8_weight_scaling: bool = False, + fc1_lora_ranks: Optional[torch.Tensor] = None, + fc1_lora_weight_ptrs: Optional[torch.Tensor] = None, + fc2_lora_ranks: Optional[torch.Tensor] = None, + fc2_lora_weight_ptrs: Optional[torch.Tensor] = None, + gated_lora_ranks: Optional[torch.Tensor] = None, + gated_lora_weight_ptrs: Optional[torch.Tensor] = None, + host_request_types: Optional[torch.Tensor] = None, + host_context_lengths: Optional[torch.Tensor] = None, + lora_max_low_rank: int = 0, + fc1_slot_lora_ranks: Optional[torch.Tensor] = None, + fc1_slot_lora_weight_ptrs: Optional[torch.Tensor] = None, + fc2_slot_lora_ranks: Optional[torch.Tensor] = None, + fc2_slot_lora_weight_ptrs: Optional[torch.Tensor] = None, + gated_slot_lora_ranks: Optional[torch.Tensor] = None, + gated_slot_lora_weight_ptrs: Optional[torch.Tensor] = None, + token_to_slot: Optional[torch.Tensor] = None, +) -> List[torch.Tensor]: + """Run one MoE layer over pre-routed tokens. + + Returns `[out]` with `out` a fresh `[num_tokens, hidden_size]` tensor in + `output_dtype`, or `[]` when `out_tensor` is given (written in place). + """ + # Pure-metadata guard: on the unquantized high-precision path the store + # is done in the activation dtype while the output tensor is allocated + # with output_dtype, so a mismatch reinterprets the written bits — observed + # on this machine to return plausible-looking wrong values, never to raise. + if input.dtype in (torch.bfloat16, torch.float16) and (fc1_expert_weights.dtype == input.dtype): + assert output_dtype == input.dtype, ( + f"output_dtype ({output_dtype}) must equal input.dtype ({input.dtype}) " + "on the unquantized path; a mismatch is written as raw activation-dtype " + "bits into an output_dtype buffer and silently gives wrong values" + ) + return torch.ops.trtllm.fused_moe( + input, + token_selected_experts, + token_final_scales, + fc1_expert_weights, + fc1_expert_biases, + fc2_expert_weights, + fc2_expert_biases, + output_dtype, + quant_scales, + input_sf=input_sf, + swizzled_input_sf=swizzled_input_sf, + swiglu_alpha=swiglu_alpha, + swiglu_beta=swiglu_beta, + swiglu_limit=swiglu_limit, + tp_size=tp_size, + tp_rank=tp_rank, + ep_size=ep_size, + ep_rank=ep_rank, + cluster_size=cluster_size, + cluster_rank=cluster_rank, + enable_alltoall=enable_alltoall, + use_deepseek_fp8_block_scale=use_deepseek_fp8_block_scale, + use_w4_group_scaling=use_w4_group_scaling, + use_int8_woq_per_channel=use_int8_woq_per_channel, + use_mxfp8_act_scaling=use_mxfp8_act_scaling, + min_latency_mode=min_latency_mode, + use_fused_finalize=use_fused_finalize, + tune_max_num_tokens=tune_max_num_tokens, + tuner_num_tokens=tuner_num_tokens, + tuner_top_k=tuner_top_k, + activation_type=activation_type, + unpadded_hidden_size=unpadded_hidden_size, + out_tensor=out_tensor, + use_dynamic_fc2_scale=use_dynamic_fc2_scale, + use_mxfp8_weight_scaling=use_mxfp8_weight_scaling, + fc1_lora_ranks=fc1_lora_ranks, + fc1_lora_weight_ptrs=fc1_lora_weight_ptrs, + fc2_lora_ranks=fc2_lora_ranks, + fc2_lora_weight_ptrs=fc2_lora_weight_ptrs, + gated_lora_ranks=gated_lora_ranks, + gated_lora_weight_ptrs=gated_lora_weight_ptrs, + host_request_types=host_request_types, + host_context_lengths=host_context_lengths, + lora_max_low_rank=lora_max_low_rank, + fc1_slot_lora_ranks=fc1_slot_lora_ranks, + fc1_slot_lora_weight_ptrs=fc1_slot_lora_weight_ptrs, + fc2_slot_lora_ranks=fc2_slot_lora_ranks, + fc2_slot_lora_weight_ptrs=fc2_slot_lora_weight_ptrs, + gated_slot_lora_ranks=gated_slot_lora_ranks, + gated_slot_lora_weight_ptrs=gated_slot_lora_weight_ptrs, + token_to_slot=token_to_slot, + ) diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/fused_moe_test.py b/tensorrt_llm/_torch/staircase/catalog/moe/fused_moe_test.py new file mode 100644 index 000000000000..331e9d395d97 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/moe/fused_moe_test.py @@ -0,0 +1,990 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the fused_moe catalog entry.""" + +import torch +import torch.nn.functional as F + +from tensorrt_llm._torch.autotuner import AutoTuner, autotune + +from .fused_moe import fused_moe + +assert torch.cuda.is_available(), "fused_moe requires a CUDA device" + +DEV = "cuda" +# The reference GEMMs must be true fp32; TF32 would leave the reference with +# 10 mantissa bits, coarser than the bf16 output it is meant to bound. +torch.backends.cuda.matmul.allow_tf32 = False + +# Relative distance between neighbouring representable values (1 + 2^-m for m +# stored mantissa bits): bf16 has 7, fp16 has 10. +_ULP = {torch.bfloat16: 2.0**-8, torch.float16: 2.0**-11} + + +def _ref_moe( + x: torch.Tensor, + ids: torch.Tensor, + scales: torch.Tensor | None, + w31: torch.Tensor, + w2: torch.Tensor, + b1: torch.Tensor | None = None, + b2: torch.Tensor | None = None, + act: str = "swiglu", + alpha: torch.Tensor | None = None, + beta: torch.Tensor | None = None, + limit: torch.Tensor | None = None, + ep_size: int = 1, + ep_rank: int = 0, +) -> torch.Tensor: + """Native-torch MoE: fp32 GEMM accumulation, activation rounded to x's dtype. + + `w31[e]` is `[up | gate]` stacked on dim 0; `ids` carries global expert ids + and this rank owns `[ep_rank * E_local, (ep_rank + 1) * E_local)`. + """ + num_tokens, hidden = x.shape + num_local = w31.shape[0] + out = torch.zeros(num_tokens, hidden, dtype=torch.float32, device=x.device) + xf = x.float() + for local_e in range(num_local): + mask = ids == ep_rank * num_local + local_e + tok, slot = mask.nonzero(as_tuple=True) + if tok.numel() == 0: + continue + xe = xf[tok] + w_up, w_gate = w31[local_e].float().chunk(2, dim=0) + h_up = xe @ w_up.t() + h_gate = xe @ w_gate.t() + if b1 is not None: + b_up, b_gate = b1[local_e].float().chunk(2, dim=0) + h_up, h_gate = h_up + b_up, h_gate + b_gate + if limit is not None: + lim = float(limit[local_e]) + h_gate = h_gate.clamp(max=lim) + h_up = h_up.clamp(min=-lim, max=lim) + if act == "swiglu": + a = float(alpha[local_e]) if alpha is not None else 1.0 + gated = h_gate * torch.sigmoid(a * h_gate) + elif act == "geglu": + gated = F.gelu(h_gate, approximate="tanh") + else: + raise ValueError(act) + if beta is not None: + h_up = h_up + float(beta[local_e]) + # the kernel materializes the FC1 activation in the input dtype + inter = (gated * h_up).to(x.dtype).float() + y = inter @ w2[local_e].float().t() + if b2 is not None: + y = y + b2[local_e].float() + weight = ( + torch.ones(tok.numel(), device=x.device) + if scales is None + else scales[tok, slot].float() + ) + out.index_add_(0, tok, y * weight.unsqueeze(1)) + return out.to(x.dtype) + + +def _assert_moe_close(y: torch.Tensor, ref: torch.Tensor) -> None: + """Two gates: per-element, row-scaled; and aggregate relative RMS. + + Kernel and reference consume bit-identical low-precision operands and + differ only in GEMM accumulation order and in which side of a rounding + boundary each intermediate lands — one flipped intermediate moves an + output element by about one ulp of that row's scale. Default assert_close + tolerances cannot express this: their bf16 `atol=1e-5` sits three orders + of magnitude below one output ulp of a two-GEMM chain, and per-element + `rtol` is meaningless where cancellation makes `|ref|` near zero. So the + element gate is 8 ulp of the row's largest magnitude and the aggregate gate + is 4 ulp of relative RMS. Measured worst case across everything this file + drives: 2.8 ulp / 1.2 ulp on the cold fallback tactic, 3.98 ulp / 1.43 ulp + once test_r1_mtp_tactic_space walks the whole tactic space -- 2.0x and 2.8x + margin. Both gates bite: test_reference_discriminates shows a wrong + computation landing at 176+ ulp per element and ~190 ulp RMS, and + test_r1_mtp_reference_discriminates at 284-997 / 197-212 ulp. + """ + assert y.dtype == ref.dtype, (y.dtype, ref.dtype) + assert y.shape == ref.shape, (y.shape, ref.shape) + ulp = _ULP[ref.dtype] + row_scale = ref.float().abs().amax(dim=1, keepdim=True).clamp_min(1e-9) + torch.testing.assert_close( + y.float() / row_scale, ref.float() / row_scale, rtol=0.0, atol=8 * ulp + ) + rel_rms = ( + (y.float() - ref.float()).pow(2).mean().sqrt() + / ref.float().pow(2).mean().sqrt().clamp_min(1e-9) + ).item() + assert rel_rms <= 4 * ulp, f"relative RMS {rel_rms:.3e} > {4 * ulp:.3e}" + + +def _make( + num_tokens: int, + hidden: int, + inter: int, + num_experts: int, + top_k: int, + seed: int = 0, + dtype: torch.dtype = torch.bfloat16, +): + """(x, ids, scales, w31, w2) with realistic magnitudes and top-k routing.""" + g = torch.Generator(device=DEV).manual_seed(seed) + x = torch.randn(num_tokens, hidden, device=DEV, generator=g).to(dtype) + w31 = (torch.randn(num_experts, 2 * inter, hidden, device=DEV, generator=g) / hidden**0.5).to( + dtype + ) + w2 = (torch.randn(num_experts, hidden, inter, device=DEV, generator=g) / inter**0.5).to(dtype) + logits = torch.randn(num_tokens, num_experts, device=DEV, generator=g) + vals, ids = torch.topk(logits, top_k, dim=-1) + return x, ids.to(torch.int32), torch.softmax(vals, dim=-1), w31, w2 + + +def test_qwen3_moe_shape_bf16() -> None: + # Qwen3-30B-A3B MoE block: hidden 2048, moe_intermediate 768, 128 experts, + # top-8. Decode-like through prefill-like token counts. + # 8192 is one token past the op's default tune_max_num_tokens bucket cap. + for num_tokens in [1, 2, 37, 256, 2048, 8192]: + x, ids, scales, w31, w2 = _make(num_tokens, 2048, 768, 128, 8, seed=num_tokens) + y = fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, [])[0] + assert y.shape == (num_tokens, 2048) and y.is_contiguous() + _assert_moe_close(y, _ref_moe(x, ids, scales, w31, w2)) + del x, ids, scales, w31, w2, y + torch.cuda.empty_cache() + + +def test_dtypes() -> None: + for dtype in [torch.bfloat16, torch.float16]: + for num_tokens in [1, 64, 512]: + x, ids, scales, w31, w2 = _make( + num_tokens, 1024, 512, 32, 4, seed=num_tokens, dtype=dtype + ) + y = fused_moe(x, ids, scales, w31, None, w2, None, dtype, [])[0] + _assert_moe_close(y, _ref_moe(x, ids, scales, w31, w2)) + + +def test_shape_sweep() -> None: + # hidden and inter must be multiples of 8; expert count and top-k are free. + cases = [ + (8, 8, 8, 2, 1), + (4, 16, 8, 1, 1), + (33, 192, 96, 7, 3), + (64, 128, 64, 8, 8), + (16, 256, 128, 256, 2), + (9, 512, 256, 32, 16), + (64, 2048, 768, 128, 1), + ] + for num_tokens, hidden, inter, num_experts, top_k in cases: + x, ids, scales, w31, w2 = _make( + num_tokens, hidden, inter, num_experts, top_k, seed=hidden + num_experts + ) + y = fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, [])[0] + _assert_moe_close(y, _ref_moe(x, ids, scales, w31, w2)) + + +def test_expert_biases() -> None: + g = torch.Generator(device=DEV).manual_seed(77) + x, ids, scales, w31, w2 = _make(64, 512, 256, 16, 4, seed=77) + b1 = (torch.randn(16, 512, device=DEV, generator=g) * 0.3).to(torch.bfloat16) + b2 = (torch.randn(16, 512, device=DEV, generator=g) * 0.3).to(torch.bfloat16) + # fc1 bias is split into [up | gate] halves exactly like fc1's rows + y = fused_moe(x, ids, scales, w31, b1, w2, b2, torch.bfloat16, [])[0] + _assert_moe_close(y, _ref_moe(x, ids, scales, w31, w2, b1=b1, b2=b2)) + # a zero fc2 bias reproduces the bias-free result + y = fused_moe(x, ids, scales, w31, b1, w2, torch.zeros_like(b2), torch.bfloat16, [])[0] + _assert_moe_close(y, _ref_moe(x, ids, scales, w31, w2, b1=b1)) + + +def test_out_tensor_is_written_in_place() -> None: + x, ids, scales, w31, w2 = _make(37, 512, 256, 16, 4, seed=5) + ref = _ref_moe(x, ids, scales, w31, w2) + # a NaN prefill proves every element is written, not accumulated into + pool = torch.full((41, 512), float("nan"), dtype=torch.bfloat16, device=DEV) + out = pool[:37] + ret = fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, [], out_tensor=out) + assert ret == [], "out_tensor form must return an empty list" + _assert_moe_close(out, ref) + # nothing was written past the requested rows + assert bool(torch.isnan(pool[37:]).all()), "wrote outside out_tensor's rows" + + +def test_none_scales_and_arbitrary_scales() -> None: + x, ids, scales, w31, w2 = _make(16, 256, 128, 16, 4, seed=7) + # token_final_scales=None combines the selected experts with weight 1.0 + y = fused_moe(x, ids, None, w31, None, w2, None, torch.bfloat16, [])[0] + _assert_moe_close(y, _ref_moe(x, ids, None, w31, w2)) + # nothing is renormalized inside: any fp32 weights are used as given + g = torch.Generator(device=DEV).manual_seed(8) + free = (torch.rand(16, 4, device=DEV, generator=g) * 4 - 2).contiguous() + y = fused_moe(x, ids, free, w31, None, w2, None, torch.bfloat16, [])[0] + _assert_moe_close(y, _ref_moe(x, ids, free, w31, w2)) + # all-zero weights zero the output exactly + zeros = torch.zeros_like(free) + y = fused_moe(x, ids, zeros, w31, None, w2, None, torch.bfloat16, [])[0] + assert bool((y == 0).all()), "zero combine weights did not zero the output" + + +def test_expert_parallel_split() -> None: + # ep_rank r owns global expert ids [r * E_local, (r + 1) * E_local); tokens + # routed elsewhere contribute nothing to that rank's output, so the two + # rank outputs sum to the single-rank result. + x, ids, scales, w31, w2 = _make(64, 256, 128, 16, 4, seed=13) + full = fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, [])[0] + _assert_moe_close(full, _ref_moe(x, ids, scales, w31, w2)) + parts = [] + for rank in (0, 1): + lo, hi = rank * 8, (rank + 1) * 8 + w31_r, w2_r = w31[lo:hi].contiguous(), w2[lo:hi].contiguous() + y = fused_moe( + x, + ids, + scales, + w31_r, + None, + w2_r, + None, + torch.bfloat16, + [], + ep_size=2, + ep_rank=rank, + )[0] + _assert_moe_close(y, _ref_moe(x, ids, scales, w31_r, w2_r, ep_size=2, ep_rank=rank)) + parts.append(y.float()) + _assert_moe_close((parts[0] + parts[1]).to(torch.bfloat16), full) + + +def test_swiglu_alpha_beta_limit() -> None: + # act = clamp(gate) * sigmoid(alpha * clamp(gate)) * (clamp(up) + beta), + # with gate clamped above by limit and up clamped to [-limit, limit]. + num_experts = 16 + x, ids, scales, w31, w2 = _make(32, 256, 128, num_experts, 4, seed=17) + g = torch.Generator(device=DEV).manual_seed(18) + alpha = 1.0 + torch.rand(num_experts, device=DEV, generator=g) + beta = torch.rand(num_experts, device=DEV, generator=g) + limit = 1.0 + 4.0 * torch.rand(num_experts, device=DEV, generator=g) + for a, b, lim in [ + (alpha, None, None), + (None, beta, None), + (None, None, limit), + (alpha, beta, limit), + ]: + y = fused_moe( + x, + ids, + scales, + w31, + None, + w2, + None, + torch.bfloat16, + [], + swiglu_alpha=a, + swiglu_beta=b, + swiglu_limit=lim, + )[0] + _assert_moe_close(y, _ref_moe(x, ids, scales, w31, w2, alpha=a, beta=b, limit=lim)) + # activation_type 7 (SwigluBias) computes the same thing on this path + y5 = fused_moe( + x, + ids, + scales, + w31, + None, + w2, + None, + torch.bfloat16, + [], + swiglu_alpha=alpha, + swiglu_beta=beta, + swiglu_limit=limit, + activation_type=5, + )[0] + y7 = fused_moe( + x, + ids, + scales, + w31, + None, + w2, + None, + torch.bfloat16, + [], + swiglu_alpha=alpha, + swiglu_beta=beta, + swiglu_limit=limit, + activation_type=7, + )[0] + assert torch.equal(y5, y7), "activation_type 7 differed from 5" + + +def test_geglu_activation() -> None: + x, ids, scales, w31, w2 = _make(32, 256, 128, 16, 4, seed=19) + y = fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, [], activation_type=6)[0] + _assert_moe_close(y, _ref_moe(x, ids, scales, w31, w2, act="geglu")) + + +def test_use_fused_finalize_false() -> None: + x, ids, scales, w31, w2 = _make(64, 512, 256, 16, 4, seed=23) + y = fused_moe( + x, + ids, + scales, + w31, + None, + w2, + None, + torch.bfloat16, + [], + use_fused_finalize=False, + )[0] + _assert_moe_close(y, _ref_moe(x, ids, scales, w31, w2)) + + +def test_out_of_range_expert_ids_are_dropped() -> None: + # Ids outside this rank's slot range are silently skipped, not clamped. + x, ids, scales, w31, w2 = _make(32, 256, 128, 8, 4, seed=29) + for bad_id in [8, 100, -1]: + bad = ids.clone() + bad[:, 1] = bad_id + y = fused_moe(x, bad, scales, w31, None, w2, None, torch.bfloat16, [])[0] + kept = torch.zeros_like(bad, dtype=torch.bool) + kept[:, 0] = True + kept[:, 2:] = True + dropped_scales = scales * kept + _assert_moe_close(y, _ref_moe(x, bad, dropped_scales, w31, w2)) + + +def test_repeated_expert_ids() -> None: + # A row may name the same expert twice; each slot is combined separately. + # This holds only for num_tokens <= 256, which is what this case drives -- + # past that the op takes a different expert-map path and reads out of + # bounds on a repeated id (see the contract's Notes). + x, ids, scales, w31, w2 = _make(32, 256, 128, 8, 4, seed=30) + dup = ids.clone() + dup[:, 1] = dup[:, 0] + y = fused_moe(x, dup, scales, w31, None, w2, None, torch.bfloat16, [])[0] + _assert_moe_close(y, _ref_moe(x, dup, scales, w31, w2)) + + +def test_inputs_untouched_and_deterministic() -> None: + x, ids, scales, w31, w2 = _make(256, 512, 256, 32, 4, seed=31) + snap = [t.clone() for t in (x, ids, scales, w31, w2)] + y_a = fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, [])[0] + y_b = fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, [])[0] + for t, s in zip((x, ids, scales, w31, w2), snap): + assert torch.equal(t, s), "an input tensor was mutated" + assert y_a.data_ptr() != y_b.data_ptr(), "two calls shared an output buffer" + assert torch.equal(y_a, y_b), "two identical calls disagreed" + + +def test_reference_discriminates() -> None: + # The gates must reject a computation that only differs in which half of + # fc1 is the gate: without this control the tolerances prove nothing. + x, ids, scales, w31, w2 = _make(64, 256, 128, 16, 4, seed=37) + y = fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, [])[0] + swapped = torch.cat([w31[:, 128:], w31[:, :128]], dim=1).contiguous() + try: + _assert_moe_close(y, _ref_moe(x, ids, scales, swapped, w2)) + except AssertionError: + pass + else: + raise AssertionError("tolerance accepted a gate/up-swapped reference") + # ...and a dropped expert must be caught too + dropped = scales.clone() + dropped[:, 0] = 0.0 + try: + _assert_moe_close(y, _ref_moe(x, ids, dropped, w31, w2)) + except AssertionError: + pass + else: + raise AssertionError("tolerance accepted a reference missing one expert") + + +def test_wrapper_rejects_output_dtype_mismatch() -> None: + # Observed silent failure: the store happens in the activation dtype while + # the buffer is allocated as output_dtype, so the bits are reinterpreted. + x, ids, scales, w31, w2 = _make(8, 128, 64, 8, 2, seed=41) + for bad in [torch.float16, torch.float32]: + try: + fused_moe(x, ids, scales, w31, None, w2, None, bad, []) + except AssertionError: + continue + raise AssertionError(f"wrapper accepted output_dtype={bad} for a bf16 input") + # positive control: the matching dtype still works + y = fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, [])[0] + _assert_moe_close(y, _ref_moe(x, ids, scales, w31, w2)) + + +def test_op_rejects_unsupported_domains() -> None: + x, ids, scales, w31, w2 = _make(8, 128, 64, 8, 2, seed=43) + args = (x, ids, scales, w31, None, w2, None, torch.bfloat16, []) + wide_x = torch.randn(8, 256, dtype=torch.bfloat16, device=DEV) + wide_ids = torch.zeros(8, 4, dtype=torch.int32, device=DEV) + wide_scales = torch.zeros(8, 4, dtype=torch.float32, device=DEV) + x_bad, ids_bad, sc_bad, w31_bad, w2_bad = _make(8, 136, 72, 8, 2, seed=44) + x_0, ids_0, sc_0, w31_0, w2_0 = _make(0, 128, 64, 8, 2, seed=45) + cases = [ + ( + "zero tokens", + lambda: fused_moe(x_0, ids_0, sc_0, w31_0, None, w2_0, None, torch.bfloat16, []), + ), + ("cpu input", lambda: fused_moe(x.cpu(), *args[1:])), + ( + "cpu weights", + lambda: fused_moe(x, ids, scales, w31.cpu(), None, w2, None, torch.bfloat16, []), + ), + ("3D input", lambda: fused_moe(x.unsqueeze(0), *args[1:])), + ( + "2D fc1", + lambda: fused_moe(x, ids, scales, w31[0], None, w2, None, torch.bfloat16, []), + ), + ("non-contiguous input", lambda: fused_moe(wide_x[:, :128], *args[1:])), + ("non-contiguous ids", lambda: fused_moe(x, wide_ids[:, :2], *args[2:])), + ( + "non-contiguous scales", + lambda: fused_moe(x, ids, wide_scales[:, :2], *args[3:]), + ), + ( + "non-contiguous fc1", + lambda: fused_moe( + x, + ids, + scales, + w31.transpose(1, 2).contiguous().transpose(1, 2), + None, + w2, + None, + torch.bfloat16, + [], + ), + ), + ("int64 ids", lambda: fused_moe(x, ids.long(), *args[2:])), + ("bf16 scales", lambda: fused_moe(x, ids, scales.bfloat16(), *args[3:])), + ( + "fp32 activations", + lambda: fused_moe( + x.float(), + ids, + scales, + w31.float(), + None, + w2.float(), + None, + torch.float32, + [], + ), + ), + ( + "mismatched weight dtype", + lambda: fused_moe( + x, ids, scales, w31.half(), None, w2.half(), None, torch.bfloat16, [] + ), + ), + ( + "fp32 bias", + lambda: fused_moe( + x, + ids, + scales, + w31, + torch.zeros(8, 128, device=DEV), + w2, + torch.zeros(8, 128, device=DEV), + torch.bfloat16, + [], + ), + ), + ( + "fc1 bias without fc2 bias", + lambda: fused_moe( + x, + ids, + scales, + w31, + torch.zeros(8, 128, dtype=torch.bfloat16, device=DEV), + w2, + None, + torch.bfloat16, + [], + ), + ), + ( + "fc2 bias without fc1 bias", + lambda: fused_moe( + x, + ids, + scales, + w31, + None, + w2, + torch.zeros(8, 128, dtype=torch.bfloat16, device=DEV), + torch.bfloat16, + [], + ), + ), + ( + "expert count mismatch", + lambda: fused_moe( + x, ids, scales, w31, None, w2[:4].contiguous(), None, torch.bfloat16, [] + ), + ), + ( + "token count mismatch", + lambda: fused_moe(x, ids[:4].contiguous(), scales[:4].contiguous(), *args[3:]), + ), + ( + "top-k mismatch", + lambda: fused_moe(x, ids, scales[:, :1].contiguous(), *args[3:]), + ), + ( + "hidden_size not a multiple of 8", + lambda: fused_moe( + x_bad[:, :132].contiguous(), + ids_bad, + sc_bad, + w31_bad[:, :, :132].contiguous(), + None, + w2_bad[:, :132].contiguous(), + None, + torch.bfloat16, + [], + ), + ), + ("min_latency_mode on bf16", lambda: fused_moe(*args, min_latency_mode=True)), + ( + "deepseek fp8 block scale on bf16", + lambda: fused_moe(*args, use_deepseek_fp8_block_scale=True), + ), + ("int8 woq on bf16", lambda: fused_moe(*args, use_int8_woq_per_channel=True)), + ( + "mxfp8 weight scaling on bf16", + lambda: fused_moe(*args, use_mxfp8_weight_scaling=True), + ), + ( + "swiglu_alpha with wrong length", + lambda: fused_moe(*args, swiglu_alpha=torch.ones(1, device=DEV)), + ), + ( + "swiglu_alpha in bf16", + lambda: fused_moe(*args, swiglu_alpha=torch.ones(8, dtype=torch.bfloat16, device=DEV)), + ), + # non-gated activation types want an [E, I, H] fc1 instead + ("Identity activation", lambda: fused_moe(*args, activation_type=1)), + ("Gelu activation", lambda: fused_moe(*args, activation_type=2)), + ("Silu activation", lambda: fused_moe(*args, activation_type=4)), + ("Relu2 activation", lambda: fused_moe(*args, activation_type=8)), + ("unknown activation", lambda: fused_moe(*args, activation_type=99)), + ("ep_rank >= ep_size", lambda: fused_moe(*args, ep_size=2, ep_rank=2)), + ( + "cluster_size without min-latency", + lambda: fused_moe(*args, cluster_size=2, cluster_rank=0), + ), + ( + "tuner_num_tokens without alltoall", + lambda: fused_moe(*args, tuner_num_tokens=8), + ), + ("alltoall without tuner args", lambda: fused_moe(*args, enable_alltoall=True)), + ( + "lora without max low rank", + lambda: fused_moe(*args, fc1_lora_ranks=torch.zeros(1, dtype=torch.int32)), + ), + ( + "out_tensor with wrong dtype", + lambda: fused_moe( + *args, out_tensor=torch.empty(8, 128, dtype=torch.float32, device=DEV) + ), + ), + ( + "out_tensor with wrong shape", + lambda: fused_moe( + *args, out_tensor=torch.empty(4, 128, dtype=torch.bfloat16, device=DEV) + ), + ), + ( + "non-contiguous out_tensor", + lambda: fused_moe(*args, out_tensor=wide_x[:, :128]), + ), + ( + "fp8 input without quant scales", + lambda: fused_moe( + x.to(torch.float8_e4m3fn), + ids, + scales, + w31.to(torch.float8_e4m3fn), + None, + w2.to(torch.float8_e4m3fn), + None, + torch.bfloat16, + [], + ), + ), + ] + for tag, call in cases: + try: + call() + except (RuntimeError, ValueError, AssertionError): + continue + raise AssertionError(f"op accepted an unsupported domain: {tag}") + + +def test_quant_flags_ignored_on_the_unquantized_path() -> None: + # These three neither raise nor change the result for bf16 weights; a + # caller cannot rely on them being honoured. + x, ids, scales, w31, w2 = _make(32, 256, 128, 16, 4, seed=47) + base = fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, [])[0] + _assert_moe_close(base, _ref_moe(x, ids, scales, w31, w2)) + args = (x, ids, scales, w31, None, w2, None, torch.bfloat16, []) + flagged = [ + ("use_w4_group_scaling", fused_moe(*args, use_w4_group_scaling=True)), + ("use_mxfp8_act_scaling", fused_moe(*args, use_mxfp8_act_scaling=True)), + ("use_dynamic_fc2_scale", fused_moe(*args, use_dynamic_fc2_scale=True)), + ] + for flag, ret in flagged: + assert torch.equal(ret[0], base), f"{flag}=True changed the result" + # a non-empty quant_scales list is likewise ignored here + y = fused_moe( + x, + ids, + scales, + w31, + None, + w2, + None, + torch.bfloat16, + [torch.ones(1, device=DEV)], + )[0] + assert torch.equal(y, base), "quant_scales changed the unquantized result" + + +# --------------------------------------------------------------------------- +# DeepSeek-R1-0528 MTP layer (model.layers.61) routed-expert geometry. +# hidden 7168, moe_intermediate 2048, 256 routed experts, top-8, bf16 weights +# and activations. Under ep_size 4 each rank holds E = 64 of those experts and +# is handed global ids in [0, 256). +R1_H, R1_I, R1_K = 7168, 2048, 8 +R1_E_LOCAL, R1_EP = 64, 4 +R1_E_GLOBAL = R1_E_LOCAL * R1_EP + + +def _make_r1_weights(num_experts: int, seed: int): + """One rank's [E, 2I, H] / [E, H, I] bf16 weight pair, built in chunks. + + 5.6 GB at E = 64. Materializing the fp32 randn for the whole stack at once + would transiently need another 11 GB; chunking caps the temporary at 64 MB. + """ + g = torch.Generator(device=DEV).manual_seed(seed) + w31 = torch.empty(num_experts, 2 * R1_I, R1_H, dtype=torch.bfloat16, device=DEV) + w2 = torch.empty(num_experts, R1_H, R1_I, dtype=torch.bfloat16, device=DEV) + chunk = max(1, (1 << 26) // (2 * R1_I * R1_H)) + for lo in range(0, num_experts, chunk): + hi = min(lo + chunk, num_experts) + w31[lo:hi] = (torch.randn(hi - lo, 2 * R1_I, R1_H, device=DEV, generator=g) / R1_H**0.5).to( + torch.bfloat16 + ) + w2[lo:hi] = (torch.randn(hi - lo, R1_H, R1_I, device=DEV, generator=g) / R1_I**0.5).to( + torch.bfloat16 + ) + return w31, w2 + + +def _make_r1_routing(num_tokens: int, num_experts: int, seed: int): + """(x, ids, scales) with top-8 routing over `num_experts` global experts.""" + g = torch.Generator(device=DEV).manual_seed(seed) + x = torch.randn(num_tokens, R1_H, device=DEV, generator=g).to(torch.bfloat16) + logits = torch.randn(num_tokens, num_experts, device=DEV, generator=g) + vals, ids = torch.topk(logits, R1_K, dim=-1) + return ( + x, + ids.to(torch.int32).contiguous(), + torch.softmax(vals, dim=-1).contiguous(), + ) + + +def _rejects(y: torch.Tensor, ref: torch.Tensor, tag: str) -> None: + """Assert the gates reject `ref` as a description of `y`.""" + try: + _assert_moe_close(y, ref) + except AssertionError: + return + raise AssertionError(f"tolerance accepted a wrong reference: {tag}") + + +def test_r1_mtp_moe_shape_bf16() -> None: + # 3.5x the largest hidden and 2.7x the largest intermediate size covered by + # the shapes above, at the expert count one ep_size-4 rank holds. + # Measured worst deviation over this function's nine comparisons: 2.54 ulp + # element-wise, 1.16 ulp relative RMS -- inside the 8/4 ulp gates. + w31, w2 = _make_r1_weights(R1_E_LOCAL, seed=1234) + for num_tokens in [1, 2, 37, 256, 2048, 8192]: + x, ids, scales = _make_r1_routing(num_tokens, R1_E_LOCAL, seed=num_tokens) + y = fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, [])[0] + assert y.shape == (num_tokens, R1_H) and y.is_contiguous() + _assert_moe_close(y, _ref_moe(x, ids, scales, w31, w2)) + del x, ids, scales, y + torch.cuda.empty_cache() + + # the two output-side switches, at this geometry + x, ids, scales = _make_r1_routing(2048, R1_E_LOCAL, seed=7) + ref = _ref_moe(x, ids, scales, w31, w2) + y = fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, [])[0] + _assert_moe_close(y, ref) + y_unfused = fused_moe( + x, + ids, + scales, + w31, + None, + w2, + None, + torch.bfloat16, + [], + use_fused_finalize=False, + )[0] + _assert_moe_close(y_unfused, ref) + # the flag changes nothing on this path: the cached C++ runner is keyed on + # the dtypes only, so whichever value the process's first call passed is + # the one in force -- and both values give the same bits anyway, measured + # in two fresh processes with the call order swapped + assert torch.equal(y, y_unfused), "use_fused_finalize changed the result" + pool = torch.full((2051, R1_H), float("nan"), dtype=torch.bfloat16, device=DEV) + out = pool[:2048] + assert fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, [], out_tensor=out) == [] + _assert_moe_close(out, ref) + assert bool(torch.isnan(pool[2048:]).all()), "wrote outside out_tensor's rows" + del w31, w2, x, ids, scales, ref, y, y_unfused, pool, out + torch.cuda.empty_cache() + + +def test_r1_mtp_expert_parallel_4way() -> None: + # ep_size 4 over a 256-expert global stack: each rank passes its own 64 + # weights and the *global* ids, and answers only for + # [64*ep_rank, 64*ep_rank + 64). The four rank outputs must sum to the + # 256-expert result, which is what the caller's correctness rests on. + # The four windows are never held at once: each is built, driven, folded + # into the accumulators and freed, so this needs 5.6 GB rather than 22.5. + token_counts = [1, 1024, 8192] + acts = {t: _make_r1_routing(t, R1_E_GLOBAL, seed=1000 + t) for t in token_counts} + got = {t: torch.zeros(t, R1_H, dtype=torch.float32, device=DEV) for t in token_counts} + want = {t: torch.zeros(t, R1_H, dtype=torch.float32, device=DEV) for t in token_counts} + for rank in range(R1_EP): + w31_r, w2_r = _make_r1_weights(R1_E_LOCAL, seed=900 + rank) + for t in token_counts: + x, ids, scales = acts[t] + y = fused_moe( + x, + ids, + scales, + w31_r, + None, + w2_r, + None, + torch.bfloat16, + [], + ep_size=R1_EP, + ep_rank=rank, + )[0] + ref = _ref_moe(x, ids, scales, w31_r, w2_r, ep_size=R1_EP, ep_rank=rank) + _assert_moe_close(y, ref) + got[t] += y.float() + want[t] += ref.float() + del y, ref + if rank == R1_EP - 1: + # a window that catches nothing contributes exactly zero, and the + # same weights read under the wrong window are a different answer + x, ids, scales = acts[1024] + inside_rank0 = (ids % R1_E_LOCAL).to(torch.int32).contiguous() + y = fused_moe( + x, + inside_rank0, + scales, + w31_r, + None, + w2_r, + None, + torch.bfloat16, + [], + ep_size=R1_EP, + ep_rank=rank, + )[0] + assert bool((y == 0).all()), "an out-of-window rank wrote nonzero output" + y = fused_moe( + x, + ids, + scales, + w31_r, + None, + w2_r, + None, + torch.bfloat16, + [], + ep_size=R1_EP, + ep_rank=0, + )[0] + _rejects( + y, + _ref_moe(x, ids, scales, w31_r, w2_r, ep_size=R1_EP, ep_rank=rank), + "rank 3's weights driven at ep_rank=0", + ) + del inside_rank0, y + del w31_r, w2_r + torch.cuda.empty_cache() + for t in token_counts: + # each rank rounds its own partial to bf16 before this add, so the sum + # sits slightly wider than any single rank does (measured 3.53 ulp + # element-wise at 8192 tokens against 2.69 for the widest single rank) + _assert_moe_close(got[t].to(torch.bfloat16), want[t].to(torch.bfloat16)) + del acts, got, want + torch.cuda.empty_cache() + + +def test_r1_mtp_expert_relabeling_is_bitwise() -> None: + # Nothing in the permutation/gather path may key on which expert index a + # weight sits at: permuting the 64-expert stack and relabelling the ids to + # match must reproduce the same output. Measured bitwise identical. + w31, w2 = _make_r1_weights(R1_E_LOCAL, seed=2222) + x, ids, scales = _make_r1_routing(512, R1_E_LOCAL, seed=31337) + y = fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, [])[0] + _assert_moe_close(y, _ref_moe(x, ids, scales, w31, w2)) + g = torch.Generator(device=DEV).manual_seed(4) + perm = torch.randperm(R1_E_LOCAL, device=DEV, generator=g) + inv = torch.empty_like(perm) + inv[perm] = torch.arange(R1_E_LOCAL, device=DEV) + y_perm = fused_moe( + x, + inv[ids.long()].to(torch.int32).contiguous(), + scales, + w31[perm].contiguous(), + None, + w2[perm].contiguous(), + None, + torch.bfloat16, + [], + )[0] + assert torch.equal(y, y_perm), "relabelling the experts changed the output" + del w31, w2, x, ids, scales, y, y_perm, perm, inv + torch.cuda.empty_cache() + + +def test_r1_mtp_reference_discriminates() -> None: + # The 8/4 ulp gates must reject wrong computations at this geometry too, + # not just at the small shapes. Measured against the correct result's 2.09 + # ulp: 284 / 197 ulp for the swapped halves and 997 / 212 for a dropped + # expert, i.e. 36x and 125x the element gate, 49x and 53x the RMS gate. + w31, w2 = _make_r1_weights(R1_E_LOCAL, seed=555) + x, ids, scales = _make_r1_routing(512, R1_E_LOCAL, seed=556) + y = fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, [])[0] + _assert_moe_close(y, _ref_moe(x, ids, scales, w31, w2)) + swapped = torch.cat([w31[:, R1_I:], w31[:, :R1_I]], dim=1).contiguous() + _rejects(y, _ref_moe(x, ids, scales, swapped, w2), "gate/up halves swapped") + del swapped + torch.cuda.empty_cache() + dropped = scales.clone() + dropped[:, 0] = 0.0 + _rejects(y, _ref_moe(x, ids, dropped, w31, w2), "one expert dropped") + del w31, w2, x, ids, scales, y, dropped + torch.cuda.empty_cache() + + +def test_r1_mtp_autotuned_tactics() -> None: + # Serving runs on the hot side of the tuner, which a cold receipt never + # sees. One autotune() pass at this geometry fills 2 tunable GEMMs x 14 + # power-of-2 token buckets, moves the bits, and still lands inside the + # gates; clearing the cache restores the cold bits exactly. + tuner = AutoTuner.get() + tuner.clear_cache() + w31, w2 = _make_r1_weights(R1_E_LOCAL, seed=1234) + cases = {} + for t in (1, 256, 8192): + x, ids, scales = _make_r1_routing(t, R1_E_LOCAL, seed=t) + cold = fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, [])[0] + cases[t] = (x, ids, scales, cold, _ref_moe(x, ids, scales, w31, w2)) + _assert_moe_close(cold, cases[t][4]) + try: + x, ids, scales = cases[8192][:3] + with autotune(): + fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, []) + assert len(tuner.profiling_cache) == 28, len(tuner.profiling_cache) + for key, (_, tactic, _) in tuner.profiling_cache.cache.items(): + assert tactic >= 0, f"{key} kept the fallback tactic after tuning" + moved = 0 + for t, (x, ids, scales, cold, ref) in cases.items(): + hot = fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, [])[0] + _assert_moe_close(hot, ref) + moved += int(not torch.equal(hot, cold)) + # the control this whole test rests on: if tuning had not changed a + # single bit, a clean pass here would mean nothing + assert moved > 0, "no warm result differed from its cold counterpart" + finally: + tuner.clear_cache() + for t, (x, ids, scales, cold, _) in cases.items(): + again = fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, [])[0] + assert torch.equal(again, cold), f"clearing the cache did not restore T={t}" + del w31, w2, cases + torch.cuda.empty_cache() + + +def _sweep_tactics(combos, args, ref, count_distinct: bool) -> int: + """Drive the op once per (runner, tactic) config, gating every output.""" + distinct = set() + for cfg in combos: + with AutoTuner.get().replay(cfg): + y = fused_moe(*args)[0] + _assert_moe_close(y, ref) + if count_distinct: + distinct.add(y.view(torch.int16).cpu().numpy().tobytes()) + del y + return len(distinct) + + +def _split_index(combos) -> int: + """Number of gemm2 tactics: itertools.product varies that context fastest.""" + return next(i for i in range(1, len(combos)) if combos[i][0][1] != combos[0][0][1]) + + +def test_r1_mtp_tactic_space() -> None: + # get_valid_tactics takes no shape, so the tactic population is a property + # of the build, not of the geometry -- checked here by capturing it at a + # certified small cell and at this one and comparing the sizes. Then every + # tactic in it is driven against the torch reference, which makes the + # receipt cover the warm path as a whole rather than one tuner outcome. + tuner = AutoTuner.get() + x_s, ids_s, scales_s, w31_s, w2_s = _make(64, 512, 256, 32, 4, seed=61) + with tuner.capture() as cap_small: + fused_moe(x_s, ids_s, scales_s, w31_s, None, w2_s, None, torch.bfloat16, []) + small = list(cap_small) + del x_s, ids_s, scales_s, w31_s, w2_s + torch.cuda.empty_cache() + + w31, w2 = _make_r1_weights(R1_E_LOCAL, seed=1234) + x, ids, scales = _make_r1_routing(256, R1_E_LOCAL, seed=99) + args = (x, ids, scales, w31, None, w2, None, torch.bfloat16, []) + ref = _ref_moe(x, ids, scales, w31, w2) + with tuner.capture() as cap: + fused_moe(*args) + combos = list(cap) + assert len(combos) == len(small), (len(combos), len(small)) + n2 = _split_index(combos) + n1 = len(combos) // n2 + assert n1 > 1 and n2 > 1 and n1 * n2 == len(combos), (n1, n2) + + # gemm2 carries nearly all of the spread: measured 101 distinct outputs + # from its 309 tactics here against 2 from gemm1's 209, the same counts the + # small shapes give. Worst deviation over the whole space: 3.71 ulp + # element-wise here and 3.98 at 8192 tokens below -- inside the 8 ulp gate + # with 2.0x margin. + d2 = _sweep_tactics(combos[:n2], args, ref, count_distinct=True) + d1 = _sweep_tactics(combos[::n2], args, ref, count_distinct=True) + # the blindness control: a harness that could not tell one tactic from + # another would report the same clean sweep with nothing measured + assert d2 > 1, d2 + assert d1 < d2, (d1, d2) + del x, ids, scales, args, ref, combos + torch.cuda.empty_cache() + + # and at the token count the caller chunks to + x, ids, scales = _make_r1_routing(8192, R1_E_LOCAL, seed=8192) + args = (x, ids, scales, w31, None, w2, None, torch.bfloat16, []) + ref = _ref_moe(x, ids, scales, w31, w2) + with tuner.capture() as cap: + fused_moe(*args) + combos = list(cap) + _sweep_tactics(combos[: _split_index(combos)], args, ref, count_distinct=False) + del w31, w2, x, ids, scales, args, ref, combos + torch.cuda.empty_cache() diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md b/tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md new file mode 100644 index 000000000000..1cee050f77be --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md @@ -0,0 +1,513 @@ +--- +receipts: + sm_100: {status: passed, trtllm: 1.3.0rc21} + sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 19} +--- + +# mxe4m3_mxe2m1_block_scale_moe_runner + +**Wraps** `torch.ops.trtllm.mxe4m3_mxe2m1_block_scale_moe_runner` (one call). + +## Semantics + +One complete mixture-of-experts layer for **MXFP4 (E2M1 + per-32 E8M0 block +scale) weights over MXFP8 (E4M3 + per-32 E8M0 block scale) activations** — the +trtllm-gen "block scale MoE" family's W4A8 member. In a single call: optional +top-k routing over router logits, expert permutation, the grouped FC1 GEMM, the +clamped gated activation, **requantization of that activation to MXFP8**, the +grouped FC2 GEMM, and the routing-weighted combine back to one row per token. + +Let `T = hidden_states.shape[0]`, `H = valid_hidden_size` (the model's true +hidden size), `I = valid_intermediate_size` (the true per-rank intermediate +size), `K = top_k`, `E = local_num_experts`, `off = local_expert_offset`, and +let `W_up[e]`, `W_gate[e]` (`[I, H]`) and `W_down[e]` (`[H, I]`) be the fp32 +values of this rank's dequantized MXFP4 weights (§ *Preconditions* fixes how +they are packed). + +The activations arrive **already quantized**: `hidden_states` holds e4m3 +elements and `hidden_states_scale` one UE8M0 byte per 32 consecutive columns, +so the value the kernel computes with is + +``` +hidden[t, k] = hidden_states[t, k].float() * 2^(sf[t, k // 32] - 127) +``` + +where `sf` is `hidden_states_scale` read as `[T, hidden_width / 32]` row-major +(§ *Preconditions*). + +**Routing** — two mutually exclusive entry points, both certified: + +- *Routed here*: pass `routing_logits` `[T, num_experts]`, leave + `topk_weights`/`topk_ids` at `None`. `routing_method_type` selects + - `0` (Default): `p = softmax_fp32(logits)` over **all** experts, then + `(w, id) = top_k(p)` — the combine weights are those softmax + probabilities and do **not** sum to 1; + - `1` (Renormalize): `(v, id) = top_k(logits)` first, then + `w = softmax_fp32(v)` over the `K` selected logits — the weights **do** + sum to 1. This is the gpt-oss / Qwen3 style. +- *Pre-routed*: pass `topk_ids` `[T, K]` int32 and `topk_weights` `[T, K]` + **bf16**, leave `routing_logits` at `None`. The kernel's routing stage is + bypassed entirely; `routing_method_type` is still validated but never + used (values `0/1/2/4/5/6` all give bitwise identical output). + +Passing both is not an error: the op announces it uses the pre-routed pair. + +**Per token `t` and slot `j < K`:** + +``` +g = topk_ids[t, j] # a GLOBAL expert id +skip this slot unless off <= g < off + E +e = g - off # index into this rank's weights + +up = hidden[t] @ W_up[e].T (+ gemm1_bias_up[e]) +gate = hidden[t] @ W_gate[e].T (+ gemm1_bias_gate[e]) + +# gemm1_clamp_limit (per expert), applied only when the tensor is given: +gate = min(gate, limit[e]) +up = clamp(up, -limit[e], limit[e]) + +# act_type 0 (SwiGlu) — the only kernel that exists on this path: +act = (up + beta[e]) * gate * sigmoid(alpha[e] * gate) + # alpha defaults to 1.0 when gemm1_alpha is None, + # beta defaults to 0.0 when gemm1_beta is None, + # so the default is plain SwiGLU: up * silu(gate) + +act = mx_requantize(act) # see below — the W4A8 step +y = act @ W_down[e].T (+ gemm2_bias[e]) +out[t] += topk_weights[t, j] * y # fp32 accumulation +``` + +and `out` is stored as bf16. The gpt-oss clamped GLU is exactly +`alpha = 1.702`, `beta = 1.0`, `limit = 7.0`. + +**The intermediate is MXFP8, on the OCP scale.** FC2 is an MXFP8 x MXFP4 GEMM, +so the FC1 epilogue quantizes its post-activation output before FC2 reads it. +Per token row and per **32 consecutive intermediate columns** (natural column +order, blocks starting at column 0 of the padded intermediate): + +``` +amax = max |act| over the 32-element block +e = floor(log2(amax)) - 8 # E8M0 byte = e + 127; e = -127 if amax == 0 +act' = round_to_nearest_even(clamp(act / 2^e, -448, +448)) * 2^e +``` + +This is the OCP MX scale — `floor` of the block max's exponent, with the block +max itself **saturating** to `448 * 2^e` whenever its mantissa exceeds 1.75 +(18–21% of blocks measured, clipping that one element by up to 12.5%). It is +**not** the round-up scale `torch.ops.trtllm.mxfp8_quantize` applies to +activations, and modelling it as such is wrong for exactly those blocks. An +intermediate element below roughly `2^-18` of its block's largest magnitude +rounds to zero. This entry's test reads the requantized values straight out of +the kernel (a down projection set to the identity) and matches the formula +above bit-exactly on all 16384 + 46080 elements of two geometries +(`H = I = 512` and `H = I = 2880`); on the same data the round-up scale +matches 99.2% of elements and the *un*quantized activation 6%. + +**FC1 half order is `[up | gate]`.** Before the kernel's row interleave (§ +*Preconditions*) the first `I_pad` rows of the FC1 operand are the up +projection (trtllm's `w3`, HF's `up_proj`) and the last `I_pad` rows the gate +projection (trtllm's `w1`, HF's `gate_proj`). The sigmoid is applied to the +**gate** half. Swapping them produces a plausible-looking, entirely different +result; nothing detects it — this entry's test checks that a gate/up-swapped +reference lands far outside the numerical gate (measured 246 ulp per element). + +**Fusion boundary.** Inside the call: routing (when driven by logits), +permutation, both GEMMs with on-the-fly MXFP4/MXFP8 dequantization, the clamped +gated activation, the MXFP8 requantization between the GEMMs, and the weighted +combine. Outside, and the caller's job: the router GEMM that produces +`routing_logits`, **adding the router bias to those logits** (see *Notes* — +`routing_bias` is a no-op here), **quantizing the bf16 hidden states to MXFP8** +(`torch.ops.trtllm.mxfp8_quantize`, § *Preconditions* — the hidden-width +padding happens inside that call, not here), all weight preprocessing, any +shared/dense expert branch, the TP all-reduce or EP gather of this call's +output, and the residual add. + +**Nothing is renormalized inside.** `topk_weights` is used exactly as given +(need not sum to 1, may be negative). `routed_scaling_factor` is accepted but +had no effect on any certified path. + +**Output.** With `output=None` the call returns a fresh contiguous +`[T, valid_hidden_size]` **bf16** tensor (the output is bf16 regardless of the +fp8 input). With `output` given, the result is written into that buffer (every +element overwritten, nothing outside its rows touched) and the call returns an +**empty `[0]` bf16 tensor** — the caller must read its own buffer. No input is +mutated; two identical calls are bitwise equal. + +## Signature + +```python +def mxe4m3_mxe2m1_block_scale_moe_runner( + routing_logits: Optional[torch.Tensor], + routing_bias: Optional[torch.Tensor], + hidden_states: torch.Tensor, + hidden_states_scale: torch.Tensor, + gemm1_weights: torch.Tensor, + gemm1_weights_scale: torch.Tensor, + gemm1_bias: Optional[torch.Tensor], + gemm1_alpha: Optional[torch.Tensor], + gemm1_beta: Optional[torch.Tensor], + gemm1_clamp_limit: Optional[torch.Tensor], + gemm2_weights: torch.Tensor, + gemm2_weights_scale: torch.Tensor, + gemm2_bias: Optional[torch.Tensor], + num_experts: int, + top_k: int, + n_group: Optional[int], + topk_group: Optional[int], + intermediate_size: int, + valid_hidden_size: Optional[int], + valid_intermediate_size: Optional[int], + local_expert_offset: int, + local_num_experts: int, + routed_scaling_factor: Optional[float], + routing_method_type: int, + act_type: int, + topk_weights: Optional[torch.Tensor] = None, + topk_ids: Optional[torch.Tensor] = None, + output: Optional[torch.Tensor] = None, + tune_max_num_tokens: int = 8192, + use_dp: bool = False, +) -> torch.Tensor +``` + +The signature is the bf16-activation sibling op's plus `hidden_states_scale` +in position 4; every other argument means the same thing. + +Derived sizes used below (`pad_up(x, a) = ceil(x / a) * a`): + +``` +I_pad = pad_up(I, 128) # FC1 rows per half, FC2 K axis +H1_pad = pad_up(H, 512) # FC1 K axis == hidden_states width +H2_pad = pad_up(H, 128) # FC2 rows +``` + +For gpt-oss-120b (`H = I = 2880`): `I_pad = 2944`, `H1_pad = 3072`, +`H2_pad = 2944`. + +### Certified arguments + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `routing_logits` | `[T, num_experts]` or `None` | bf16 or fp32 (bitwise identical results) | contiguous | CUDA | +| `routing_bias` | must be `None` — see *Notes* | — | — | — | +| `hidden_states` | `[T, H1_pad]` | **float8_e4m3fn** | contiguous | CUDA | +| `hidden_states_scale` | **1-D**, exactly `T * H1_pad/32` elements | **uint8** (E8M0) | contiguous, **linear** (row-major `[T, H1_pad/32]` flattened) — *not* swizzled | CUDA | +| `gemm1_weights` | `[E, 2*I_pad, H1_pad/2]` | **uint8** (2 E2M1 codes per byte) | contiguous, pre-shuffled | CUDA | +| `gemm1_weights_scale` | `[E, 2*I_pad, H1_pad/32]` | **uint8** (E8M0) | contiguous, pre-shuffled + swizzled | CUDA | +| `gemm1_bias` | `[E, 2*I_pad]` or `None` | **fp32** | contiguous, pre-shuffled | CUDA | +| `gemm1_alpha` / `gemm1_beta` / `gemm1_clamp_limit` | `[E]` or `None` | **fp32** | contiguous | CUDA | +| `gemm2_weights` | `[E, H2_pad, I_pad/2]` | **uint8** | contiguous, pre-shuffled | CUDA | +| `gemm2_weights_scale` | `[E, H2_pad, I_pad/32]` | **uint8** (E8M0) | contiguous, pre-shuffled + swizzled | CUDA | +| `gemm2_bias` | `[E, H2_pad]` or `None` | **fp32** | contiguous, pre-shuffled | CUDA | +| `num_experts` | scalar | int, the global routing space; `> top_k` | — | — | +| `top_k` | scalar | int, `0 < top_k < num_experts` | — | — | +| `n_group` / `topk_group` | `None` | — | — | — | +| `intermediate_size` | scalar | int, **must equal `I_pad`** | — | — | +| `valid_hidden_size` | scalar (not `None`) | int, `H`; multiple of 32 with `pad_up(H, 128) == H2_pad` | — | — | +| `valid_intermediate_size` | scalar or `None` | int, multiple of 32, `<= I_pad` | — | — | +| `local_expert_offset` / `local_num_experts` | scalars | int, `E = local_num_experts >= 1` | — | — | +| `routed_scaling_factor` | `None` | — | — | — | +| `routing_method_type` | scalar | int: **0** (Default) or **1** (Renormalize) | — | — | +| `act_type` | scalar | int: **0** (SwiGlu) only | — | — | +| `topk_weights` | `[T, top_k]` or `None` | **bf16** | contiguous | CUDA | +| `topk_ids` | `[T, top_k]` or `None` | **int32** | contiguous | CUDA | +| `output` | `[T, valid_hidden_size]` or `None` | bf16 | contiguous | CUDA | +| return | `[T, valid_hidden_size]`, or `[0]` when `output` is given | bf16 | contiguous, newly allocated | same as `hidden_states` | + +Biases are independently optional: `gemm1_bias` and `gemm2_bias` may each be +`None`, in any combination. + +### Arguments held inert (not certified) + +- `routing_bias` — see *Notes*; the wrapper rejects it on every path where it + was observed to be a no-op. +- `n_group` / `topk_group` — the grouped (DeepSeek-V3 style) routing path, + reachable only with `routing_method_type = 2`. +- `routing_method_type` values other than `0` and `1`. `4` (RenormalizeNaive) + and `6` (SigmoidRenorm) do run on the logits entry point here, but their + weight formulas are not pinned; `2` (DeepSeekV3), `3` (Llama4) and + `5` (MiniMax2) were not exercised on that entry point at all. On the + pre-routed entry point every one of `0/1/2/4/5/6` is inert. +- `routed_scaling_factor` — changed nothing on either certified entry point + (`None`, `1.0` and `2.5` give bitwise identical results); it belongs to the + grouped-routing path. +- `tune_max_num_tokens` — an autotuner bucket cap; `8192` and `128` give + bitwise identical results. Left at its default. +- `use_dp` — an autotuner token-bucket deflation hint for data-parallel + deployments; `True`/`False` give bitwise identical results at `ep_size = 1`. + Left at `False`. +- Multi-rank tensor parallelism. `local_expert_offset` / `local_num_experts` + (expert parallelism) **are** certified; TP, which instead splits `I` across + ranks and reduces this call's output afterwards, is not exercised here. + +## Metadata consumed + +None. The op reads no attention metadata, no KV cache, no registered layer +and no module state — every tensor and every scalar it uses is an argument. + +Two process-global caches sit behind it, neither of which changes the result: + +- an **autotuner profiling cache** that picks the two GEMM tactics. Cold — + the state every call certified here ran in — it takes a fallback tactic; + tactic choice is a performance knob only. +- a cache of C++ `MxE4m3MxE2m1BlockScaleMoERunner` objects keyed by + `act_type`, each owning a GPU workspace allocated on first use. + +## Preconditions + +### Activation preparation — MXFP8, linear scales, padded hidden + +`hidden_states` / `hidden_states_scale` are exactly the two outputs of + +```python +# swizzled_layout alignment +data, sf = torch.ops.trtllm.mxfp8_quantize(x, False, 512) +``` + +on the bf16 `[T, H]` hidden states, and that pairing is certified: for +`H = 2880` it yields `data [T, 3072] float8_e4m3fn` and a 1-D +`sf [T * 96] uint8`, which is what this op wants. Concretely, whatever +produces them: + +- `hidden_states` is 2-D, **float8_e4m3fn** (a uint8 view of the same bytes is + rejected: `hidden_states must be Float8_e4m3fn`), contiguous, CUDA, with + width exactly `gemm1_weights.shape[-1] * 2 == H1_pad`. `alignment` must + therefore be the FC1 K alignment (512 for this weight family); quantizing + with `alignment=32` leaves the tensor at width `H` and is rejected (`the + third dimension of weights must be equal to hidden_size`). The op never + pads the hidden itself — the widening `2880 -> 3072` happens inside + `mxfp8_quantize`, and the caller does no `F.pad` at all. +- `hidden_states_scale` is **1-D** (`hidden_states_scale must be 1D`), + **uint8** (`must be UInt8`), contiguous, CUDA, and holds exactly + `T * H1_pad/32` bytes (`hidden_states_scale has incorrect size` otherwise — + one byte too few or too many both raise). Byte `t * (H1_pad/32) + b` is the + E8M0 scale (`value = 2^(byte - 127)`) of columns `[32b, 32b+32)` of token + `t`. The **128x4 swizzled** scale order is a silent wrong answer, not an + error, whenever it happens to have the same byte count (`T % 128 == 0`) — + see *Notes*. +- Columns `[H, H1_pad)` multiply zero-valued padded weights; `mxfp8_quantize` + writes zero data bytes and zero scale bytes there, but any **finite** + content is inert (verified with e4m3 `400.0` under a live scale byte). A + NaN there propagates into the output. +- `T >= 1`; `T = 0` is rejected. Certified `T`: 1, 2, 4, 6, 8, 10, 12, 16, 17, + 24, 32, 128, 256, 1024, 8192. + +### Weight preparation — the caller owns all of it + +The activation dtype changes none of it: this is the same prepared weight +layout `torch.ops.trtllm.bf16_mxe2m1_block_scale_moe_runner` consumes, and in +this build both ops are fed by the same trtllm weight-loading class +(`MXFP4WeightTRTLLMGenFusedMoEMethod`, which the W4A16 and W4A8-MXFP8 methods +inherit without overriding a single padding, shuffle or swizzle step). One +prepared expert stack therefore serves both. + +Start from the checkpoint's per-expert MXFP4 tensors, laid out **row-major +over the output dim**, K packed two E2M1 codes per byte with the **low nibble +holding the even K index**, and one E8M0 byte (`value = 2^(byte - 127)`) per +32 consecutive K elements: + +``` +up (w3 / up_proj) packed [E, I, H/2] uint8 scale [E, I, H/32] uint8 +gate (w1 / gate_proj) packed [E, I, H/2] uint8 scale [E, I, H/32] uint8 +down (w2 / down_proj) packed [E, H, I/2] uint8 scale [E, H, I/32] uint8 +up/gate bias [E, I] , down bias [E, H] -> both must reach the kernel as fp32 +``` + +(Checkpoints that store gate and up interleaved on one `2I` axis must +de-interleave first.) Then, **per expert**: + +1. **Zero-pad.** `up`/`gate`: rows `I -> I_pad`, byte columns + `H/2 -> H1_pad/2`, scale columns `H/32 -> H1_pad/32`. `down`: rows + `H -> H2_pad`, byte columns `I/2 -> I_pad/2`, scale columns + `I/32 -> I_pad/32`. Biases: `I -> I_pad` and `H -> H2_pad`. +2. **FC1 concat.** Stack the padded halves on the row axis as + `[up ; gate]`, giving `[2*I_pad, ...]`. Same for the bias. +3. **FC1 row permute** = interleave, then block shuffle: + - interleave: destination row `2i` takes the up half's row `i`, + destination row `2i+1` takes the gate half's row `i`; + - block shuffle (also applied to FC2, which skips step 2/3's interleave): + within each aligned block of 32 rows, source row `4u + v` + (`0 <= u < 8`, `0 <= v < 4`) moves to destination row `8v + u`. + + Apply the **same** permutation to the weight bytes, the scale bytes and + the fp32 bias, so row `i` of all three still describe the same output + channel. A bias not permuted with its weights is silently wrong. +4. **Scale swizzle.** After the row permute, each expert's scale matrix + `[M, C]` (`M % 128 == 0`, `C % 4 == 0`) is rewritten into the trtllm-gen + 128x4 layout: the byte at `(m, c)` moves to flat offset + + ``` + (m // 128) * 512 * (C // 4) + (c // 4) * 512 + + (m % 32) * 16 + ((m % 128) // 32) * 4 + (c % 4) + ``` + + and the flat result is handed to the kernel with its nominal `[E, M, C]` + shape. Weight bytes and biases are **not** swizzled — only scales. + +The two trtllm helpers `torch.ops.trtllm.shuffle_matrix(x, perm)` (a plain +row gather, `out[i] = x[perm[i]]`) and +`torch.ops.trtllm.block_scale_interleave(x)` (the 128x4 swizzle over a +`[E, M, C]` uint8 tensor, returning a flat buffer of +`E * pad_up(M, 128) * pad_up(C, 4)` bytes) produce byte-identical results to +steps 3 and 4; this entry's test asserts that equivalence. + +### Shapes and layout + +- `intermediate_size` must equal `gemm1_weights.shape[1] // 2`; anything else + is rejected (`No valid config found for the given problem shape`). +- `valid_hidden_size` is the **output width**, unrelated to the widened + activation: it stays at the model's true hidden `H` (2880 for gpt-oss) even + though `hidden_states` is `H1_pad` (3072) wide. It must be a multiple of 32 + satisfying `pad_up(valid_hidden_size, 128) == gemm2_weights.shape[1]`. + `None` means "use the `hidden_states` width", which now fails whenever + `H1_pad != H2_pad` (`gemm2_weights_scale has incorrect dim 1`) — so on this + op it must always be passed explicitly. Setting it to `H2_pad` yields the + wider output whose tail columns are all zero; other values are rejected + (`gemm2_weights_scale has incorrect dim 1` or `No valid config found`). +- `valid_intermediate_size` must be a multiple of 32 and at most + `intermediate_size`; `None` means `intermediate_size`. It is a bandwidth + hint: the kernel only reads the first `valid_intermediate_size` intermediate + columns. Since padded columns carry zero weight, every value `>= I` gives + the same result — but a **smaller** value silently truncates the layer. +- `topk_ids` is `[T, top_k]` **int32** contiguous CUDA and `topk_weights` is + `[T, top_k]` **bf16** contiguous CUDA; the two must be given together + (`routing_logits or (topk_ids and topk_weights) must be provided`), their + row count must match `hidden_states` and their column count must match + `top_k`. int64 ids and fp32 weights are both rejected. +- `routing_logits` is `[T, num_experts]`, bf16 or fp32, contiguous CUDA. A + column count other than `num_experts` is rejected. +- `gemm1_alpha` / `gemm1_beta` / `gemm1_clamp_limit`, when given, are fp32 + CUDA tensors with exactly **`local_num_experts`** elements (not + `num_experts`), indexed by local expert. +- Both bias tensors, when given, must be **fp32** (bf16 is rejected); both + weight-scale tensors must be **uint8** (an int8 view is rejected). +- `output`, when given, must be a contiguous CUDA bf16 tensor of shape + exactly `[T, valid_hidden_size]`; wrong dtype and wrong shape are both + rejected. A leading row-slice of a taller contiguous buffer is valid — rows + past `T` are left bitwise untouched. +- **Every tensor argument must be contiguous.** A strided view is accepted + silently and read as if dense — see *Notes*. The wrapper asserts this. + +### Sizes and counts + +- `0 < top_k < num_experts` (`num_experts must be greater than top_k`; + `top_k = 0` is rejected). Certified `top_k`: 1, 2, 3, 4. +- `num_experts` is only the routing space; it need not equal + `local_num_experts` (certified at `num_experts = 8`, + `local_num_experts = 4`, `local_expert_offset = 4`). Certified + `num_experts`: 2, 3, 4, 5, 8, 16, 128; certified `local_num_experts`: + 2, 3, 4, 5, 8, 16, 128. +- Certified geometries `(H, I)`: (512, 128), (512, 256), (512, 512), + (640, 128), (1024, 512), (2048, 512), (2880, 1024), (2880, 2880). `H` and + `I` need not be multiples of the kernel's alignments — that is what the + padding is for — but both must be multiples of 32, since they are passed as + `valid_hidden_size` / `valid_intermediate_size` (a `valid_hidden_size` of + 500 was rejected with `No valid config found`). +- `local_expert_offset + local_num_experts > num_experts` is **not** checked. +- Expert ids outside `[local_expert_offset, local_expert_offset + + local_num_experts)` — including negative ids and ids `>= num_experts` — are + silently dropped; the token's other slots still combine normally. +- A repeated expert id inside one token's row contributes **once**, with the + weight of its **first** occurrence; the duplicate slot's weight is + discarded. + +### Dtypes and modes + +- `hidden_states` must be float8_e4m3fn; bf16 is rejected + (`hidden_states must be Float8_e4m3fn`). bf16 hidden states belong to the + sibling op `torch.ops.trtllm.bf16_mxe2m1_block_scale_moe_runner`, whose + signature has no `hidden_states_scale` parameter at all. +- `act_type` must be `0`. `1` (Relu2) and `2` (Silu) fail with + `No kernel found for the given options: mDtypeA: MxE4m3, mDtypeB: MxE2m1 + ...` — the non-gated activations have no cubin in this MXFP8 x MXFP4 family. + `SwigluBias` is not a separate value: the per-expert `alpha`/`beta`/ + `clamp_limit` tensors turn `act_type = 0` into it. +- `hidden_states` must be 2-D; a 3-D `[1, T, H1_pad]` view is rejected. +- sm_100 only. The receipt covers sm_100 (B200), the only arch available + here; the installed build carries the assertion `Only SM100f is supported + by MXFP4 block scale MOE` in this op's C++ entry point, so other Blackwell + variants are expected to raise rather than compute something wrong. + +A caller violating none of the above gets the result described under +*Semantics*, to within the bound in *Notes*. + +## Notes + +- **Numerics.** Against a native-torch reference that consumes bit-identical + operands (exact MXFP4 weight and MXFP8 activation dequantization, fp32 GEMM + accumulation, the MXFP8 requantization of the FC1 output modelled exactly, + fp32 combine), the kernel's worst element-wise deviation over every + configuration in this entry's test was **2.0 ulp of the token row's largest + magnitude** (bf16 ulp = 2^-8, worst case: 128 experts, 8192 tokens, + `H = I = 2880`) and its worst relative RMS deviation **0.87 ulp**. Dropping + the intermediate requantization from that reference — i.e. modelling this + layer the way a bf16-intermediate MoE would be modelled — costs a factor of + seven: 13.5 ulp element-wise and 8.5 ulp RMS (~3% relative) at the gpt-oss + geometry. That gap is a property of the kernel, not of the test: it is what + the extra e4m3 rounding between the two GEMMs does to the layer's output. +- **`routing_bias` is silently ignored** for `routing_method_type` 0, 1, 4 + and 6, and on the pre-routed entry point for every routing method: a bias + of `+1e3 / -1e3` on two experts leaves the output bitwise unchanged. + A checkpoint whose router carries a bias (gpt-oss does) must therefore have + it **added into `routing_logits`** before the call — that is, use the + router `nn.Linear`'s own bias and pass `routing_bias=None`. The wrapper + asserts `routing_bias is None` on exactly those paths, because the failure + is otherwise invisible: expert selection quietly ignores the bias and the + model degrades without any error. The argument is presumably live for the + grouped (`routing_method_type = 2`) path, which is not certified here. +- **A swizzled activation-scale buffer is a silent wrong answer.** + `mxfp8_quantize(x, True, 512)` returns the same data bytes and a scale + buffer of `pad_up(T,128) * pad_up(cols,4)` bytes in 128x4 order. Whenever + that count coincides with the linear one (`T` a multiple of 128, `cols` a + multiple of 4 — true for gpt-oss at any graph-friendly batch size) the size + check passes and the kernel reads scales for the wrong blocks: measured + 260 ulp element-wise off at `T = 128`, `H = I = 2880`. Nothing in the + metadata distinguishes the two layouts, so no guard can catch it — pass + `swizzled_layout=False`. +- **Non-contiguous tensors are a silent wrong answer.** The kernel takes raw + data pointers and assumes a dense row-major layout. A strided view of + `hidden_states` (measured 479 ulp off), of `hidden_states_scale` (311 ulp, + using a same-length stride-2 1-D view), of any of the six + weight/scale/bias tensors, of `routing_logits`, `topk_weights`, `topk_ids` + or `output` is accepted without complaint and reads (or writes) the wrong + elements. The wrapper asserts contiguity on every tensor argument. +- **The FC1 accumulation precision is not pinned by this entry**, and cannot + be: the epilogue's e4m3 requantization (3 mantissa bits) erases any + difference between an fp32 and a bf16 intermediate accumulator. What *is* + pinned is the value FC2 consumes — bit-exactly, per the formula in + *Semantics*. +- Sibling ops in this build address the same job with other operand dtypes: + `torch.ops.trtllm.bf16_mxe2m1_block_scale_moe_runner` (bf16 activations, + same prepared weight layout, no `hidden_states_scale` parameter), + `torch.ops.trtllm.fp4_block_scale_moe_runner` (NVFP4) and + `torch.ops.trtllm.fp8_block_scale_moe_runner`. Catalog membership is + `index.yaml`'s fact alone. + + +## The FC1 epilogue's block-scale recipe is architecture-specific + +trtllm-gen ships one cubin per architecture, and the two differ **bit-exactly** +in how the FC1 epilogue picks the e8m0 scale when it requantizes its activation +output to MXFP8 for FC2. Measured with an identity down-projection, which reads +the intermediate out element by element rather than inferring it from output +noise (`test_intermediate_is_mxfp8_quantized`): + +| Arch | e8m0 exponent | Name | +|---|---|---| +| sm_100 | `floor(log2(amax)) - 8` | OCP scale | +| sm_103 | `ceil(log2(amax / 448))` | round-up scale | + +Both were verified bit-exact on their own architecture (0 mismatched elements) +and each *refutes* the other's, so the cases genuinely separate the recipes. +The round-up form is what `torch.ops.trtllm.mxfp8_quantize` has always used, so +sm_103 brings the MoE epilogue into agreement with the standalone quantizer. + +This is the whole of the sm_100 -> sm_103 numerical difference for this entry. +Before the reference was made architecture-aware, 10 of 19 cells failed -- +relative RMS ~5.5 ulp against a 4 ulp gate, max abs ~0.07 against 0.031. With +the correct recipe in force all 19 pass **at the original tolerances**, which +is what identifies the recipe as the sole cause rather than one contributor. + +A future cubin that changes recipe again will fail the bit-exact test rather +than drift quietly; record the new recipe in `_SCALE_RECIPE_BY_SM`, and do not +widen a tolerance instead. diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.py b/tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.py new file mode 100644 index 000000000000..52e3056cdb35 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.py @@ -0,0 +1,129 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""MXFP4-weight / MXFP8-activation mixture-of-experts layer: routing (or given +top-k) + grouped FC1 GEMM + clamped gated activation + MXFP8 requantization + +grouped FC2 GEMM + routing-weighted combine, in one trtllm-gen call.""" + +from typing import Optional + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + +# routing_method_type values for which the kernel never reads routing_bias: +# 0 Default, 1 Renormalize, 4 RenormalizeNaive, 6 SigmoidRenorm. Observed on +# this machine: a +-1e3 bias leaves the result bitwise unchanged. +_ROUTING_BIAS_IGNORED = (0, 1, 4, 6) + + +def mxe4m3_mxe2m1_block_scale_moe_runner( + routing_logits: Optional[torch.Tensor], + routing_bias: Optional[torch.Tensor], + hidden_states: torch.Tensor, + hidden_states_scale: torch.Tensor, + gemm1_weights: torch.Tensor, + gemm1_weights_scale: torch.Tensor, + gemm1_bias: Optional[torch.Tensor], + gemm1_alpha: Optional[torch.Tensor], + gemm1_beta: Optional[torch.Tensor], + gemm1_clamp_limit: Optional[torch.Tensor], + gemm2_weights: torch.Tensor, + gemm2_weights_scale: torch.Tensor, + gemm2_bias: Optional[torch.Tensor], + num_experts: int, + top_k: int, + n_group: Optional[int], + topk_group: Optional[int], + intermediate_size: int, + valid_hidden_size: Optional[int], + valid_intermediate_size: Optional[int], + local_expert_offset: int, + local_num_experts: int, + routed_scaling_factor: Optional[float], + routing_method_type: int, + act_type: int, + topk_weights: Optional[torch.Tensor] = None, + topk_ids: Optional[torch.Tensor] = None, + output: Optional[torch.Tensor] = None, + tune_max_num_tokens: int = 8192, + use_dp: bool = False, +) -> torch.Tensor: + """Run one MXFP4-weight MoE layer over MXFP8 (e4m3 + UE8M0) activations. + + Returns a fresh `[num_tokens, valid_hidden_size]` bf16 tensor, or — when + `output` is given — an empty `[0]` tensor, the result having been written + into `output`. + """ + # Pure-metadata guard: the kernel takes raw data pointers and assumes a + # dense row-major layout for every tensor. A strided view is accepted + # without complaint and silently reads the wrong elements (observed on + # this machine for every tensor argument listed here, hidden_states_scale + # included). + for name, tensor in ( + ("routing_logits", routing_logits), + ("routing_bias", routing_bias), + ("hidden_states", hidden_states), + ("hidden_states_scale", hidden_states_scale), + ("gemm1_weights", gemm1_weights), + ("gemm1_weights_scale", gemm1_weights_scale), + ("gemm1_bias", gemm1_bias), + ("gemm1_alpha", gemm1_alpha), + ("gemm1_beta", gemm1_beta), + ("gemm1_clamp_limit", gemm1_clamp_limit), + ("gemm2_weights", gemm2_weights), + ("gemm2_weights_scale", gemm2_weights_scale), + ("gemm2_bias", gemm2_bias), + ("topk_weights", topk_weights), + ("topk_ids", topk_ids), + ("output", output), + ): + if tensor is not None: + assert tensor.is_contiguous(), ( + f"{name} must be contiguous; a strided view is read as if dense " + "and silently produces wrong results" + ) + # Pure-metadata guard: routing_bias is a no-op for every non-grouped + # routing method, and for any routing method once topk_ids/topk_weights + # carry the routing. A caller expecting the bias to shift expert selection + # gets a silently different model. + if routing_bias is not None: + assert topk_ids is None, ( + "routing_bias is ignored when topk_ids/topk_weights are given " + "(routing has already happened); fold it into the router logits" + ) + assert routing_method_type not in _ROUTING_BIAS_IGNORED, ( + f"routing_bias is silently ignored for routing_method_type=" + f"{routing_method_type}; add it to routing_logits before the call" + ) + return torch.ops.trtllm.mxe4m3_mxe2m1_block_scale_moe_runner( + routing_logits, + routing_bias, + hidden_states, + hidden_states_scale, + gemm1_weights, + gemm1_weights_scale, + gemm1_bias, + gemm1_alpha, + gemm1_beta, + gemm1_clamp_limit, + gemm2_weights, + gemm2_weights_scale, + gemm2_bias, + num_experts, + top_k, + n_group, + topk_group, + intermediate_size, + valid_hidden_size, + valid_intermediate_size, + local_expert_offset, + local_num_experts, + routed_scaling_factor, + routing_method_type, + act_type, + topk_weights=topk_weights, + topk_ids=topk_ids, + output=output, + tune_max_num_tokens=tune_max_num_tokens, + use_dp=use_dp, + ) diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner_test.py b/tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner_test.py new file mode 100644 index 000000000000..2e39a2dea56f --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner_test.py @@ -0,0 +1,1604 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the mxe4m3_mxe2m1_block_scale_moe_runner catalog entry.""" + +import math + +import torch + +from .mxe4m3_mxe2m1_block_scale_moe_runner import mxe4m3_mxe2m1_block_scale_moe_runner as moe + +assert torch.cuda.is_available(), "mxe4m3_mxe2m1_block_scale_moe_runner requires a CUDA device" + +DEV = "cuda" +# The reference GEMMs must be true fp32; TF32 would leave the reference with +# 10 mantissa bits, coarser than the bf16 output it is meant to bound. +torch.backends.cuda.matmul.allow_tf32 = False + +# Relative distance between neighbouring bf16 values (1 + 2^-7 stored mantissa +# bits -> one ulp is 2^-8 of the binade top). +ULP = 2.0**-8 + +# The 16 e2m1 code points, in code order: sign bit 3, exponent bits 2:1, +# mantissa bit 0. +E2M1 = torch.tensor( + [ + 0.0, + 0.5, + 1.0, + 1.5, + 2.0, + 3.0, + 4.0, + 6.0, + -0.0, + -0.5, + -1.0, + -1.5, + -2.0, + -3.0, + -4.0, + -6.0, + ], + dtype=torch.float32, + device=DEV, +) +SV = 32 # mx block size: one e8m0 scale per 32 elements along K +WEIGHT_ALIGN = 128 # TMA row alignment for both GEMMs +HIDDEN_ALIGN = 512 # extra alignment on FC1's K axis +E4M3_MAX = 448.0 # largest finite e4m3 magnitude + + +def _pad_up(x: int, a: int) -> int: + return (x + a - 1) // a * a + + +# ── operand construction ────────────────────────────────────────────────── + + +def _exp_base(k: int) -> int: + """Base of the 6-wide e8m0 exponent window for a `[*, k]` weight. + + Chosen so a `k`-long dot product against the activations built by + `_rand_mxfp8` lands at a standard deviation near 3: large enough that the + clamp limit of 7 bites on a few percent of the elements, small enough that + the gated activation is not a constant +-56 everywhere. A saturated + activation would make every candidate intermediate-quantization recipe fit + equally well and would leave `test_intermediate_is_mxfp8_quantized` blind. + """ + return 127 + round(0.5 * math.log2(0.01057 / k)) + + +def _rand_mxfp4(e: int, n: int, k: int, gen: torch.Generator): + """Random mxfp4 expert stack. + + Returns `(packed [E, N, K/2] uint8, scales [E, N, K/32] uint8, codes + [E, N, K] uint8)`. Codes and e8m0 exponents are drawn directly, so the + fp32 value of every weight is exact — the reference never has to model a + quantizer. + """ + base = _exp_base(k) + codes = torch.randint(0, 16, (e, n, k), dtype=torch.uint8, device=DEV, generator=gen) + exps = torch.randint( + base, base + 6, (e, n, k // SV), dtype=torch.uint8, device=DEV, generator=gen + ) + packed = (codes[..., 0::2] | (codes[..., 1::2] << 4)).contiguous() + return packed, exps, codes + + +def _dequant(codes_e: torch.Tensor, exps_e: torch.Tensor) -> torch.Tensor: + """One expert's `[N, K]` fp32 weight from its codes and e8m0 exponents.""" + scale = torch.exp2(exps_e.float() - 127.0).repeat_interleave(SV, dim=1) + return E2M1[codes_e.long()] * scale + + +def _rand_mxfp8(rows: int, valid_k: int, padded_k: int, gen: torch.Generator): + """Random MXFP8 activation in the layout this op consumes. + + Returns `(data [rows, padded_k] float8_e4m3fn, sf [rows * padded_k/32] + uint8, x [rows, valid_k] fp32)`, where `x` is the exact dequantization of + the first `valid_k` columns. Columns past `valid_k` carry the zero byte and + a zero scale byte, exactly as `torch.ops.trtllm.mxfp8_quantize` writes its + column padding. Drawing the e4m3 elements and the e8m0 exponents directly + keeps the reference operand exact — no quantizer is modelled here. + """ + assert valid_k % SV == 0 and padded_k % SV == 0 + data = torch.zeros(rows, padded_k, dtype=torch.float32, device=DEV) + data[:, :valid_k] = torch.randn(rows, valid_k, device=DEV, generator=gen) + data = data.to(torch.float8_e4m3fn) + sf = torch.zeros(rows, padded_k // SV, dtype=torch.uint8, device=DEV) + sf[:, : valid_k // SV] = torch.randint( + 125, + 128, + (rows, valid_k // SV), + dtype=torch.uint8, + device=DEV, + generator=gen, + ) + x = data.float() * torch.exp2(sf.float() - 127.0).repeat_interleave(SV, dim=1) + return data, sf.reshape(-1).contiguous(), x[:, :valid_k].contiguous() + + +def _dequant_mxfp8(data: torch.Tensor, sf: torch.Tensor) -> torch.Tensor: + """Exact fp32 value of an `[rows, k]` e4m3 tensor with a linear scale buffer.""" + rows, k = data.shape + scale = torch.exp2(sf.view(rows, k // SV).float() - 127.0) + return data.float() * scale.repeat_interleave(SV, dim=1) + + +# ── kernel weight layout (pure torch) ───────────────────────────────────── + + +def _pad3(t: torch.Tensor, rows: int, cols: int) -> torch.Tensor: + out = torch.zeros(t.shape[0], rows, cols, dtype=t.dtype, device=t.device) + out[:, : t.shape[1], : t.shape[2]] = t + return out + + +def _pad2(t: torch.Tensor, cols: int) -> torch.Tensor: + out = torch.zeros(t.shape[0], cols, dtype=t.dtype, device=t.device) + out[:, : t.shape[1]] = t + return out + + +def _blk32_perm(m: int) -> torch.Tensor: + """Gather index of the 32-row block shuffle: within each block of 32 rows, + source row `4u + v` lands at destination row `8v + u`.""" + assert m % 32 == 0 + j = torch.arange(32) + dst = (j % 4) * 8 + j // 4 + idx = torch.empty(32, dtype=torch.long) + idx[dst] = j + return (idx.repeat(m // 32) + torch.arange(m // 32).repeat_interleave(32) * 32).to(DEV) + + +def _gate_interleave_perm(m: int) -> torch.Tensor: + """Gather index that interleaves the `[up | gate]` halves of a `2*I` row + stack into `up0, gate0, up1, gate1, ...`.""" + p = torch.empty(m, dtype=torch.long) + p[0::2] = torch.arange(0, m // 2) + p[1::2] = torch.arange(m // 2, m) + return p.to(DEV) + + +def _swizzle_scales(s: torch.Tensor) -> torch.Tensor: + """Per-expert 128x4 block-scale swizzle, result viewed back as `[E, M, C]`. + + Flat destination of scale `(e, m, c)` is + `e*M*C + (m//128)*512*(C//4) + (c//4)*512 + (m%32)*16 + ((m%128)//32)*4 + (c%4)`. + """ + e, m, c = s.shape + assert m % 128 == 0 and c % 4 == 0 + v = s.reshape(e, m // 128, 4, 32, c // 4, 4) # (e, m/128, (m%128)/32, m%32, c/4, c%4) + v = v.permute(0, 1, 4, 3, 2, 5) + return v.reshape(e, m, c).contiguous() + + +def _prep_fc1(up_p, gt_p, up_s, gt_s, up_b, gt_b, i_pad, h1_pad): + """Pad both halves, concat as `[up | gate]`, interleave + block-shuffle rows.""" + w = torch.cat([_pad3(up_p, i_pad, h1_pad // 2), _pad3(gt_p, i_pad, h1_pad // 2)], dim=1) + s = torch.cat([_pad3(up_s, i_pad, h1_pad // SV), _pad3(gt_s, i_pad, h1_pad // SV)], dim=1) + b = torch.cat([_pad2(up_b, i_pad), _pad2(gt_b, i_pad)], dim=1).float() + m = w.shape[1] + perm = _gate_interleave_perm(m)[_blk32_perm(m)] + return ( + torch.index_select(w, 1, perm).contiguous(), + _swizzle_scales(torch.index_select(s, 1, perm)), + torch.index_select(b, 1, perm).contiguous(), + ) + + +def _prep_fc2(dn_p, dn_s, dn_b, h2_pad, i_pad): + """Pad the down projection and block-shuffle its rows.""" + w = _pad3(dn_p, h2_pad, i_pad // 2) + s = _pad3(dn_s, h2_pad, i_pad // SV) + b = _pad2(dn_b, h2_pad).float() + perm = _blk32_perm(h2_pad) + return ( + torch.index_select(w, 1, perm).contiguous(), + _swizzle_scales(torch.index_select(s, 1, perm)), + torch.index_select(b, 1, perm).contiguous(), + ) + + +def _build(num_experts: int, hidden: int, inter: int, seed: int): + """Build one MoE layer: kernel-ready tensors plus the reference operands.""" + gen = torch.Generator(device=DEV).manual_seed(seed) + i_pad = _pad_up(inter, WEIGHT_ALIGN) + h1_pad = _pad_up(hidden, HIDDEN_ALIGN) + h2_pad = _pad_up(hidden, WEIGHT_ALIGN) + up_p, up_s, up_c = _rand_mxfp4(num_experts, inter, hidden, gen) + gt_p, gt_s, gt_c = _rand_mxfp4(num_experts, inter, hidden, gen) + dn_p, dn_s, dn_c = _rand_mxfp4(num_experts, hidden, inter, gen) + up_b = torch.randn(num_experts, inter, device=DEV, generator=gen) + gt_b = torch.randn(num_experts, inter, device=DEV, generator=gen) + dn_b = torch.randn(num_experts, hidden, device=DEV, generator=gen) * 0.05 + w1, s1, b1 = _prep_fc1(up_p, gt_p, up_s, gt_s, up_b, gt_b, i_pad, h1_pad) + w2, s2, b2 = _prep_fc2(dn_p, dn_s, dn_b, h2_pad, i_pad) + args = dict( + gemm1_weights=w1, + gemm1_weights_scale=s1, + gemm1_bias=b1, + gemm2_weights=w2, + gemm2_weights_scale=s2, + gemm2_bias=b2, + intermediate_size=i_pad, + valid_hidden_size=hidden, + valid_intermediate_size=inter, + ) + ref = dict( + up=(up_c, up_s, up_b), + gate=(gt_c, gt_s, gt_b), + down=(dn_c, dn_s, dn_b), + h1_pad=h1_pad, + gen=gen, + ) + return args, ref + + +def _call(data, sf, args, num_experts, top_k, **kw): + """Invoke the wrapper with this layer's tensors and size scalars.""" + kw.setdefault("routing_logits", None) + kw.setdefault("routing_bias", None) + kw.setdefault("n_group", None) + kw.setdefault("topk_group", None) + kw.setdefault("local_expert_offset", 0) + kw.setdefault("local_num_experts", num_experts) + kw.setdefault("routed_scaling_factor", None) + kw.setdefault("routing_method_type", 1) + kw.setdefault("act_type", 0) + kw.setdefault("gemm1_alpha", None) + kw.setdefault("gemm1_beta", None) + kw.setdefault("gemm1_clamp_limit", None) + kw.setdefault("valid_hidden_size", args["valid_hidden_size"]) + kw.setdefault("valid_intermediate_size", args["valid_intermediate_size"]) + kw.setdefault("intermediate_size", args["intermediate_size"]) + return moe( + kw.pop("routing_logits"), + kw.pop("routing_bias"), + data, + sf, + args["gemm1_weights"], + args["gemm1_weights_scale"], + args["gemm1_bias"], + kw.pop("gemm1_alpha"), + kw.pop("gemm1_beta"), + kw.pop("gemm1_clamp_limit"), + args["gemm2_weights"], + args["gemm2_weights_scale"], + args["gemm2_bias"], + num_experts, + top_k, + kw.pop("n_group"), + kw.pop("topk_group"), + kw.pop("intermediate_size"), + kw.pop("valid_hidden_size"), + kw.pop("valid_intermediate_size"), + kw.pop("local_expert_offset"), + kw.pop("local_num_experts"), + kw.pop("routed_scaling_factor"), + kw.pop("routing_method_type"), + kw.pop("act_type"), + **kw, + ) + + +# ── reference ───────────────────────────────────────────────────────────── + + +# The FC1 epilogue's block-scale recipe is architecture-specific, and the +# difference is bit-exact rather than a tolerance: trtllm-gen ships one cubin +# per architecture. Measured on each, with an identity down-projection reading +# the intermediate out element by element (test_intermediate_is_mxfp8_quantized): +# +# sm_100 e8m0 = floor(log2(amax)) - 8 "OCP scale" +# sm_103 e8m0 = ceil(log2(amax / 448)) "round-up scale" +# +# The round-up form is the one `torch.ops.trtllm.mxfp8_quantize` has always +# used, so sm_103 makes the MoE epilogue and the standalone quantizer agree. +# Both are named here, and each architecture's test refutes the other's recipe, +# so a future cubin that switches back cannot pass silently. +_OCP_SCALE, _ROUND_UP_SCALE = "ocp", "round_up" + +_SCALE_RECIPE_BY_SM = { + (10, 0): _OCP_SCALE, + (10, 3): _ROUND_UP_SCALE, +} + + +def _scale_recipe() -> str: + sm = torch.cuda.get_device_capability() + recipe = _SCALE_RECIPE_BY_SM.get(sm) + assert recipe is not None, ( + f"the FC1 epilogue's requantization recipe is not certified on sm_{sm[0]}{sm[1]}; " + "run the identity-down-projection probe and record it before trusting this entry" + ) + return recipe + + +def _block_scale(amax: torch.Tensor, recipe: str) -> torch.Tensor: + """The per-32-column e8m0 scale, under the named recipe.""" + if recipe == _OCP_SCALE: + exp = torch.floor(torch.log2(amax)) - 8.0 + else: + exp = torch.ceil(torch.log2(amax / E4M3_MAX)) + exp = torch.where(amax == 0, torch.full_like(amax, -127.0), exp) + return torch.exp2(exp.clamp(-127.0, 127.0)) + + +def _q_intermediate(act: torch.Tensor, recipe: str | None = None) -> torch.Tensor: + """MXFP8 requantization the FC1 epilogue applies to its activation output. + + Per 32 consecutive intermediate columns: the architecture's block scale + (see above), then round-to-nearest-even into e4m3 with saturation at + +-448. Returns the dequantized fp32 value FC2 actually consumes. Pinned + bit-exactly by `test_intermediate_is_mxfp8_quantized`. + """ + rows, cols = act.shape + blocks = act.reshape(rows, cols // SV, SV) + amax = blocks.abs().amax(dim=-1, keepdim=True) + scale = _block_scale(amax, recipe or _scale_recipe()) + q = (blocks / scale).clamp(-E4M3_MAX, E4M3_MAX).to(torch.float8_e4m3fn).float() + return (q * scale).reshape(rows, cols) + + +def _ref_moe( + x_valid, + ids, + scales, + ref, + alpha=None, + beta=None, + limit=None, + offset: int = 0, + num_local: int | None = None, + swap_gate_up: bool = False, + quantize_intermediate: bool = True, +): + """Native-torch MoE over dequantized mxfp4 weights, fp32 throughout. + + `ids` carries global expert ids; this rank answers for + `[offset, offset + num_local)`. `scales[t, j]` multiplies slot `j`'s + expert output; nothing is renormalized. The FC1 activation is requantized + to MXFP8 before FC2, which is what the kernel does. + """ + num_tokens, hidden = x_valid.shape + up_c, up_s, up_b = ref["up"] + gt_c, gt_s, gt_b = ref["gate"] + dn_c, dn_s, dn_b = ref["down"] + if swap_gate_up: + (up_c, up_s, up_b), (gt_c, gt_s, gt_b) = (gt_c, gt_s, gt_b), (up_c, up_s, up_b) + num_local = up_c.shape[0] if num_local is None else num_local + xf = x_valid.float() + out = torch.zeros(num_tokens, hidden, dtype=torch.float32, device=x_valid.device) + for local_e in range(num_local): + mask = ids == offset + local_e + tok, slot = mask.nonzero(as_tuple=True) + if tok.numel() == 0: + continue + xe = xf[tok] + up = xe @ _dequant(up_c[local_e], up_s[local_e]).t() + up_b[local_e] + gate = xe @ _dequant(gt_c[local_e], gt_s[local_e]).t() + gt_b[local_e] + if limit is not None: + lim = float(limit[local_e]) + gate = gate.clamp(max=lim) + up = up.clamp(-lim, lim) + a = 1.0 if alpha is None else float(alpha[local_e]) + b = 0.0 if beta is None else float(beta[local_e]) + act = (up + b) * gate * torch.sigmoid(a * gate) + if quantize_intermediate: + act = _q_intermediate(act) + y = act @ _dequant(dn_c[local_e], dn_s[local_e]).t() + dn_b[local_e] + out.index_add_(0, tok, y * scales[tok, slot].float().unsqueeze(1)) + return out + + +def _dev(y: torch.Tensor, ref: torch.Tensor): + """(worst element deviation, relative RMS deviation), both in bf16 ulp.""" + o, r = y.float(), ref.float() + row = r.abs().amax(dim=1, keepdim=True).clamp_min(1e-9) + elt = ((o - r).abs() / row).max().item() / ULP + rms = ((o - r).pow(2).mean().sqrt() / r.pow(2).mean().sqrt().clamp_min(1e-9)).item() / ULP + return elt, rms + + +def _assert_moe_close(y: torch.Tensor, ref: torch.Tensor) -> None: + """Two gates: per-element, row-scaled; and aggregate relative RMS. + + Kernel and reference consume bit-identical mxfp4 weights and MXFP8 + activations and model the same FC1-output MXFP8 requantization, so they + differ only in accumulation order and in whether a marginal FC1 value + rounds to the same e4m3 code. Default `assert_close` tolerances cannot + express that: their bf16 `atol=1e-5` sits three orders of magnitude below + one output ulp of a two-GEMM chain, and per-element `rtol` is meaningless + where cancellation drives `|ref|` to zero. So the element gate is 8 ulp of + the row's largest magnitude and the aggregate gate is 4 ulp of relative + RMS. Worst values measured over every configuration covered here: 2.0 ulp + element-wise (128 experts, 8192 tokens, H=I=2880) and 0.87 ulp RMS. Both + gates bite — `test_reference_discriminates` shows a reference that skips + the FC1-output requantization landing at 19.6 / 10.6 ulp and a + gate/up-swapped or unshuffled operand at 100+ ulp. + """ + assert y.dtype == ref.dtype == torch.bfloat16, (y.dtype, ref.dtype) + assert y.shape == ref.shape, (y.shape, ref.shape) + row = ref.float().abs().amax(dim=1, keepdim=True).clamp_min(1e-9) + torch.testing.assert_close(y.float() / row, ref.float() / row, rtol=0.0, atol=8 * ULP) + _, rms = _dev(y, ref) + assert rms <= 4.0, f"relative RMS {rms:.2f} ulp > 4 ulp" + + +def _routing(num_tokens, num_experts, top_k, gen): + """Random ids/weights for the pre-routed entry point.""" + ids = torch.stack( + [torch.randperm(num_experts, device=DEV, generator=gen)[:top_k] for _ in range(num_tokens)] + ).to(torch.int32) + wts = torch.rand(num_tokens, top_k, device=DEV, generator=gen).to(torch.bfloat16) + return ids, wts + + +def _logits(num_tokens, num_experts, gen): + """Router logits whose per-token values are all distinct in bf16, so the + kernel's top-k and torch.topk cannot disagree on a tie.""" + ladder = torch.linspace(-4.0, 4.0, num_experts, device=DEV) + rows = [ + ladder[torch.randperm(num_experts, device=DEV, generator=gen)] for _ in range(num_tokens) + ] + lg = torch.stack(rows).to(torch.bfloat16) + assert all(len(set(r.tolist())) == num_experts for r in lg), ( + "logit ties would make topk ambiguous" + ) + return lg + + +# gpt-oss-120b MoE geometry, built once and shared by the tests that use it. +_GPT_OSS = None + + +def _gpt_oss(): + global _GPT_OSS + if _GPT_OSS is None: + _GPT_OSS = _build(128, 2880, 2880, seed=100) + return _GPT_OSS + + +def _swiglu_params(num_local): + """gpt-oss clamped-GLU constants, one entry per local expert.""" + return ( + torch.full((num_local,), 1.702, dtype=torch.float32, device=DEV), + torch.full((num_local,), 1.0, dtype=torch.float32, device=DEV), + torch.full((num_local,), 7.0, dtype=torch.float32, device=DEV), + ) + + +# ── tests ───────────────────────────────────────────────────────────────── + + +def test_layout_helpers_match(): + """The trtllm preprocessing ops reproduce the pure-torch layout exactly.""" + gen = torch.Generator(device=DEV).manual_seed(201) + rows, cols = 256, 12 + x = torch.randint(0, 256, (4, rows, cols), dtype=torch.uint8, device=DEV, generator=gen) + perm = _gate_interleave_perm(rows)[_blk32_perm(rows)] + for e in range(x.shape[0]): + assert torch.equal( + torch.ops.trtllm.shuffle_matrix(x[e].contiguous(), perm), + torch.index_select(x[e], 0, perm), + ), "shuffle_matrix is not a plain row gather" + assert torch.equal( + torch.ops.trtllm.block_scale_interleave(x).reshape(x.shape), + _swizzle_scales(x), + ), "block_scale_interleave does not match the documented 128x4 swizzle" + print(" test_layout_helpers_match OK") + + +def test_gpt_oss_pre_routed(): + """E=128, H=I=2880, top-4, clamped GLU: decode- through prefill-sized batches.""" + args, ref = _gpt_oss() + alpha, beta, limit = _swiglu_params(128) + gen = torch.Generator(device=DEV).manual_seed(11) + for num_tokens in (1, 2, 4, 17, 256, 1024, 8192): + data, sf, xv = _rand_mxfp8(num_tokens, 2880, ref["h1_pad"], gen) + ids, wts = _routing(num_tokens, 128, 4, gen) + out = _call( + data, + sf, + args, + 128, + 4, + topk_ids=ids, + topk_weights=wts, + gemm1_alpha=alpha, + gemm1_beta=beta, + gemm1_clamp_limit=limit, + ) + assert out.shape == (num_tokens, 2880), out.shape + exp = _ref_moe(xv, ids, wts, ref, alpha, beta, limit).to(torch.bfloat16) + _assert_moe_close(out, exp) + print(" test_gpt_oss_pre_routed OK") + + +def test_gpt_oss_router_entry_point(): + """Routing inside the call: Renormalize (top-k then softmax) and Default.""" + args, ref = _gpt_oss() + alpha, beta, limit = _swiglu_params(128) + gen = torch.Generator(device=DEV).manual_seed(12) + for num_tokens in (4, 128): + data, sf, xv = _rand_mxfp8(num_tokens, 2880, ref["h1_pad"], gen) + lg = _logits(num_tokens, 128, gen) + kw = dict(gemm1_alpha=alpha, gemm1_beta=beta, gemm1_clamp_limit=limit) + + out = _call(data, sf, args, 128, 4, routing_logits=lg, **kw) + top_v, top_i = torch.topk(lg.float(), 4, dim=-1) + exp = _ref_moe( + xv, + top_i.to(torch.int32), + torch.softmax(top_v, dim=-1), + ref, + alpha, + beta, + limit, + ).to(torch.bfloat16) + _assert_moe_close(out, exp) + + out0 = _call(data, sf, args, 128, 4, routing_logits=lg, routing_method_type=0, **kw) + sm_v, sm_i = torch.topk(torch.softmax(lg.float(), dim=-1), 4, dim=-1) + exp0 = _ref_moe(xv, sm_i.to(torch.int32), sm_v, ref, alpha, beta, limit).to(torch.bfloat16) + _assert_moe_close(out0, exp0) + + # fp32 logits are accepted and give the same answer as their bf16 form + out32 = _call(data, sf, args, 128, 4, routing_logits=lg.float().contiguous(), **kw) + assert torch.equal(out32, out) + print(" test_gpt_oss_router_entry_point OK") + + +def test_mxfp8_quantize_pairing(): + """`mxfp8_quantize(x, swizzled_layout=False, alignment=512)` is the pairing. + + Shape, dtype, dimensionality and scale order of that op's two outputs are + exactly what this op consumes; every violation of the pairing is pinned. + """ + args, ref = _gpt_oss() + alpha, beta, limit = _swiglu_params(128) + gen = torch.Generator(device=DEV).manual_seed(13) + kw = dict(gemm1_alpha=alpha, gemm1_beta=beta, gemm1_clamp_limit=limit) + for num_tokens in (1, 128): + xb = torch.randn(num_tokens, 2880, device=DEV, dtype=torch.bfloat16, generator=gen) + data, sf = torch.ops.trtllm.mxfp8_quantize(xb, False, 512) + assert data.shape == (num_tokens, 3072) and data.dtype == torch.float8_e4m3fn + assert sf.shape == (num_tokens * 96,) and sf.dtype == torch.uint8 + assert data.shape[1] == args["gemm1_weights"].shape[-1] * 2 + ids, wts = _routing(num_tokens, 128, 4, gen) + out = _call(data, sf, args, 128, 4, topk_ids=ids, topk_weights=wts, **kw) + xv = _dequant_mxfp8(data, sf)[:, :2880].contiguous() + exp = _ref_moe(xv, ids, wts, ref, alpha, beta, limit).to(torch.bfloat16) + _assert_moe_close(out, exp) + + # the swizzled scale buffer has the same byte count whenever the token + # count is a multiple of 128 -- and is then taken without complaint + sw_data, sw_sf = torch.ops.trtllm.mxfp8_quantize(xb, True, 512) + assert torch.equal(sw_data, data), "swizzling must not change the data bytes" + if num_tokens % 128 == 0: + assert sw_sf.numel() == sf.numel() + bad = _call(data, sw_sf, args, 128, 4, topk_ids=ids, topk_weights=wts, **kw) + elt, rms = _dev(bad, exp) + assert elt > 50.0 and rms > 20.0, ( + f"a swizzled scale buffer was not detectably wrong: {elt:.1f} ulp" + ) + else: + assert sw_sf.numel() != sf.numel() + + def rejected(name, fn): + try: + fn() + except (RuntimeError, AssertionError): + return + raise AssertionError(f"{name} was accepted") + + xb = torch.randn(8, 2880, device=DEV, dtype=torch.bfloat16, generator=gen) + data, sf = torch.ops.trtllm.mxfp8_quantize(xb, False, 512) + ids, wts = _routing(8, 128, 4, gen) + kw2 = dict(topk_ids=ids, topk_weights=wts, **kw) + # alignment 32 leaves the hidden un-padded: rejected, never zero-extended + d32, s32 = torch.ops.trtllm.mxfp8_quantize(xb, False, 32) + assert d32.shape == (8, 2880) and s32.shape == (8 * 90,) + rejected("alignment=32 hidden", lambda: _call(d32, s32, args, 128, 4, **kw2)) + rejected("2-D scale buffer", lambda: _call(data, sf.view(8, 96), args, 128, 4, **kw2)) + rejected( + "fp8-typed scale buffer", + lambda: _call(data, sf.view(torch.float8_e4m3fn), args, 128, 4, **kw2), + ) + rejected( + "short scale buffer", + lambda: _call(data, sf[:-32].contiguous(), args, 128, 4, **kw2), + ) + rejected( + "long scale buffer", + lambda: _call(data, torch.cat([sf, sf[:32]]), args, 128, 4, **kw2), + ) + rejected("bf16 hidden_states", lambda: _call(xb, sf, args, 128, 4, **kw2)) + rejected( + "uint8-typed hidden_states", + lambda: _call(data.view(torch.uint8), sf, args, 128, 4, **kw2), + ) + print(" test_mxfp8_quantize_pairing OK") + + +def test_intermediate_is_mxfp8_quantized(): + """The FC1 activation reaches FC2 as MXFP8, on the OCP scale, bit-exactly. + + A down projection set to the identity turns the returned rows into the + kernel's own post-activation intermediate, so the requantization can be + read out element by element instead of inferred from output noise. + """ + for num_experts, hidden, num_tokens, seed in ((4, 512, 32, 71), (2, 2880, 16, 72)): + inter = hidden + gen = torch.Generator(device=DEV).manual_seed(seed) + i_pad = _pad_up(inter, WEIGHT_ALIGN) + h1_pad = _pad_up(hidden, HIDDEN_ALIGN) + h2_pad = _pad_up(hidden, WEIGHT_ALIGN) + up_p, up_s, up_c = _rand_mxfp4(num_experts, inter, hidden, gen) + gt_p, gt_s, gt_c = _rand_mxfp4(num_experts, inter, hidden, gen) + # down = identity: e2m1 code 2 is 1.0, e8m0 byte 127 is 2^0 + dn_c = torch.zeros(num_experts, hidden, inter, dtype=torch.uint8, device=DEV) + dn_c[:, torch.arange(hidden), torch.arange(inter)] = 2 + dn_p = (dn_c[..., 0::2] | (dn_c[..., 1::2] << 4)).contiguous() + dn_s = torch.full((num_experts, hidden, inter // SV), 127, dtype=torch.uint8, device=DEV) + zero1 = torch.zeros(num_experts, inter, device=DEV) + zero2 = torch.zeros(num_experts, hidden, device=DEV) + w1, s1, b1 = _prep_fc1(up_p, gt_p, up_s, gt_s, zero1, zero1, i_pad, h1_pad) + w2, s2, b2 = _prep_fc2(dn_p, dn_s, zero2, h2_pad, i_pad) + args = dict( + gemm1_weights=w1, + gemm1_weights_scale=s1, + gemm1_bias=b1, + gemm2_weights=w2, + gemm2_weights_scale=s2, + gemm2_bias=b2, + intermediate_size=i_pad, + valid_hidden_size=hidden, + valid_intermediate_size=inter, + ) + data, sf, xv = _rand_mxfp8(num_tokens, hidden, h1_pad, gen) + ids = torch.zeros(num_tokens, 1, dtype=torch.int32, device=DEV) + wts = torch.ones(num_tokens, 1, dtype=torch.bfloat16, device=DEV) + alpha, beta, limit = _swiglu_params(num_experts) + out = _call( + data, + sf, + args, + num_experts, + 1, + topk_ids=ids, + topk_weights=wts, + gemm1_alpha=alpha, + gemm1_beta=beta, + gemm1_clamp_limit=limit, + ).float() + + up = xv @ _dequant(up_c[0], up_s[0]).t() + gate = xv @ _dequant(gt_c[0], gt_s[0]).t() + act = ( + (up.clamp(-7.0, 7.0) + 1.0) + * gate.clamp(max=7.0) + * torch.sigmoid(1.702 * gate.clamp(max=7.0)) + ) + blk_amax = act.reshape(num_tokens, -1, SV).abs().amax(dim=-1) + mantissa = blk_amax / torch.exp2(torch.floor(torch.log2(blk_amax))) + saturating = (mantissa > 1.75).float().mean().item() + assert saturating > 0.05, f"only {saturating:.3f} of blocks exercise the saturating branch" + recipe = _scale_recipe() + assert torch.equal(out, _q_intermediate(act, recipe)), ( + f"the FC1 epilogue's MXFP8 requantization is not the {recipe!r} " + "block scale with round-to-nearest-even and +-448 saturation. If " + "this architecture's cubin changed recipe, probe it with an " + "identity down-projection and record the new one in " + "_SCALE_RECIPE_BY_SM -- do not widen a tolerance, the two recipes " + "differ bit-exactly and every other cell's reference depends on " + "which one is in force" + ) + # The other architecture's recipe must NOT also fit, or this case does + # not actually separate them and the bit-exact claim is vacuous. + other = _OCP_SCALE if recipe == _ROUND_UP_SCALE else _ROUND_UP_SCALE + assert not torch.equal(out, _q_intermediate(act, other)), ( + f"the {other!r} scale fits too — the case that separates the recipes is not covered" + ) + assert not torch.equal(out, act.to(torch.bfloat16).float()), ( + "the intermediate was not requantized at all" + ) + print(" test_intermediate_is_mxfp8_quantized OK") + + +def test_activation_variants(): + """alpha/beta/clamp_limit are independently optional and per-expert.""" + num_experts, hidden, inter, top_k, num_tokens = 8, 1024, 512, 3, 24 + args, ref = _build(num_experts, hidden, inter, seed=13) + gen = ref["gen"] + data, sf, xv = _rand_mxfp8(num_tokens, hidden, ref["h1_pad"], gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + one = torch.ones(num_experts, dtype=torch.float32, device=DEV) + cases = [ + (None, None, None), + (None, None, torch.full((num_experts,), 7.0, device=DEV)), + (one * 1.702, one, None), + ( + torch.linspace(1.0, 2.0, num_experts, device=DEV), + torch.linspace(0.0, 1.5, num_experts, device=DEV), + torch.linspace(2.0, 9.0, num_experts, device=DEV), + ), + ] + for alpha, beta, limit in cases: + out = _call( + data, + sf, + args, + num_experts, + top_k, + topk_ids=ids, + topk_weights=wts, + gemm1_alpha=alpha, + gemm1_beta=beta, + gemm1_clamp_limit=limit, + ) + exp = _ref_moe(xv, ids, wts, ref, alpha, beta, limit).to(torch.bfloat16) + _assert_moe_close(out, exp) + print(" test_activation_variants OK") + + +def test_bias_combinations(): + """Each bias tensor is independently optional.""" + num_experts, hidden, inter, top_k, num_tokens = 8, 512, 256, 2, 16 + args, ref = _build(num_experts, hidden, inter, seed=14) + gen = ref["gen"] + data, sf, xv = _rand_mxfp8(num_tokens, hidden, ref["h1_pad"], gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + alpha, beta, limit = _swiglu_params(num_experts) + for want1, want2 in ((True, True), (True, False), (False, True), (False, False)): + sub = dict(args) + sub_ref = dict(ref) + if not want1: + sub["gemm1_bias"] = None + for role in ("up", "gate"): + c, s, b = ref[role] + sub_ref[role] = (c, s, torch.zeros_like(b)) + if not want2: + sub["gemm2_bias"] = None + c, s, b = ref["down"] + sub_ref["down"] = (c, s, torch.zeros_like(b)) + out = _call( + data, + sf, + sub, + num_experts, + top_k, + topk_ids=ids, + topk_weights=wts, + gemm1_alpha=alpha, + gemm1_beta=beta, + gemm1_clamp_limit=limit, + ) + exp = _ref_moe(xv, ids, wts, sub_ref, alpha, beta, limit).to(torch.bfloat16) + _assert_moe_close(out, exp) + print(" test_bias_combinations OK") + + +def test_expert_counts_and_top_k(): + """Small expert counts and top_k = 1.""" + hidden, inter, num_tokens = 512, 128, 12 + for num_experts, top_k in ((2, 1), (3, 2), (5, 1), (16, 4)): + args, ref = _build(num_experts, hidden, inter, seed=200 + num_experts) + gen = ref["gen"] + data, sf, xv = _rand_mxfp8(num_tokens, hidden, ref["h1_pad"], gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + alpha, beta, limit = _swiglu_params(num_experts) + out = _call( + data, + sf, + args, + num_experts, + top_k, + topk_ids=ids, + topk_weights=wts, + gemm1_alpha=alpha, + gemm1_beta=beta, + gemm1_clamp_limit=limit, + ) + exp = _ref_moe(xv, ids, wts, ref, alpha, beta, limit).to(torch.bfloat16) + _assert_moe_close(out, exp) + print(" test_expert_counts_and_top_k OK") + + +def test_expert_parallel_window(): + """local_expert_offset/local_num_experts select a global-id window; slots + routed outside it contribute nothing.""" + num_experts, hidden, inter, top_k, num_tokens = 8, 512, 256, 4, 12 + args, ref = _build(num_experts, hidden, inter, seed=15) + gen = ref["gen"] + data, sf, xv = _rand_mxfp8(num_tokens, hidden, ref["h1_pad"], gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + offset, num_local = 4, 4 + sub = dict(args) + for key in ( + "gemm1_weights", + "gemm1_weights_scale", + "gemm1_bias", + "gemm2_weights", + "gemm2_weights_scale", + "gemm2_bias", + ): + sub[key] = args[key][offset : offset + num_local].contiguous() + sub_ref = dict(ref) + for role in ("up", "gate", "down"): + c, s, b = ref[role] + sub_ref[role] = ( + c[offset : offset + num_local], + s[offset : offset + num_local], + b[offset : offset + num_local], + ) + alpha, beta, limit = _swiglu_params(num_local) + kw = dict( + local_expert_offset=offset, + local_num_experts=num_local, + gemm1_alpha=alpha, + gemm1_beta=beta, + gemm1_clamp_limit=limit, + ) + out = _call(data, sf, sub, num_experts, top_k, topk_ids=ids, topk_weights=wts, **kw) + exp = _ref_moe( + xv, ids, wts, sub_ref, alpha, beta, limit, offset=offset, num_local=num_local + ).to(torch.bfloat16) + _assert_moe_close(out, exp) + + # ids outside the window (including negative and >= num_experts) are dropped + stray = ids.clone() + stray[0, 0] = 999 + stray[1, 1] = -1 + out = _call(data, sf, sub, num_experts, top_k, topk_ids=stray, topk_weights=wts, **kw) + exp = _ref_moe( + xv, stray, wts, sub_ref, alpha, beta, limit, offset=offset, num_local=num_local + ).to(torch.bfloat16) + _assert_moe_close(out, exp) + print(" test_expert_parallel_window OK") + + +def test_output_buffer(): + """`output=` writes in place, returns an empty tensor, touches no extra row.""" + num_experts, hidden, inter, top_k, num_tokens = 8, 512, 256, 2, 10 + args, ref = _build(num_experts, hidden, inter, seed=16) + gen = ref["gen"] + data, sf, _ = _rand_mxfp8(num_tokens, hidden, ref["h1_pad"], gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + alpha, beta, limit = _swiglu_params(num_experts) + kw = dict( + topk_ids=ids, + topk_weights=wts, + gemm1_alpha=alpha, + gemm1_beta=beta, + gemm1_clamp_limit=limit, + ) + fresh = _call(data, sf, args, num_experts, top_k, **kw) + + buf = torch.full((num_tokens, hidden), -3.0, device=DEV, dtype=torch.bfloat16) + ret = _call(data, sf, args, num_experts, top_k, output=buf, **kw) + assert ret.numel() == 0 and ret.dtype == torch.bfloat16, (ret.shape, ret.dtype) + assert torch.equal(buf, fresh) + + big = torch.full((num_tokens + 3, hidden), -3.0, device=DEV, dtype=torch.bfloat16) + snap = big.clone() + view = big[:num_tokens] + assert view.is_contiguous() + _call(data, sf, args, num_experts, top_k, output=view, **kw) + assert torch.equal(big[:num_tokens], fresh) + assert torch.equal(big[num_tokens:], snap[num_tokens:]), "wrote past the requested rows" + print(" test_output_buffer OK") + + +def test_shapes(): + """Hidden/intermediate sizes on and off the padding boundaries.""" + for hidden, inter in ((512, 128), (640, 128), (2048, 512), (2880, 1024)): + num_experts, top_k, num_tokens = 4, 2, 8 + args, ref = _build(num_experts, hidden, inter, seed=17) + gen = ref["gen"] + data, sf, xv = _rand_mxfp8(num_tokens, hidden, ref["h1_pad"], gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + alpha, beta, limit = _swiglu_params(num_experts) + out = _call( + data, + sf, + args, + num_experts, + top_k, + topk_ids=ids, + topk_weights=wts, + gemm1_alpha=alpha, + gemm1_beta=beta, + gemm1_clamp_limit=limit, + ) + assert out.shape == (num_tokens, hidden), out.shape + exp = _ref_moe(xv, ids, wts, ref, alpha, beta, limit).to(torch.bfloat16) + _assert_moe_close(out, exp) + + # token counts, including one past the default autotuner bucket cap + num_experts, hidden, inter, top_k = 8, 512, 256, 4 + args, ref = _build(num_experts, hidden, inter, seed=18) + gen = ref["gen"] + alpha, beta, limit = _swiglu_params(num_experts) + for num_tokens in (1, 2, 17, 1024, 8192): + data, sf, xv = _rand_mxfp8(num_tokens, hidden, ref["h1_pad"], gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + out = _call( + data, + sf, + args, + num_experts, + top_k, + topk_ids=ids, + topk_weights=wts, + gemm1_alpha=alpha, + gemm1_beta=beta, + gemm1_clamp_limit=limit, + ) + exp = _ref_moe(xv, ids, wts, ref, alpha, beta, limit).to(torch.bfloat16) + _assert_moe_close(out, exp) + print(" test_shapes OK") + + +def test_valid_sizes(): + """valid_hidden_size is the output width — it does not follow the widened + activation; valid_intermediate_size bounds the intermediate columns read.""" + num_experts, hidden, inter, top_k, num_tokens = 4, 2880, 2880, 2, 8 + args, ref = _build(num_experts, hidden, inter, seed=19) + gen = ref["gen"] + data, sf, xv = _rand_mxfp8(num_tokens, hidden, ref["h1_pad"], gen) + assert data.shape[1] == 3072, data.shape + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + alpha, beta, limit = _swiglu_params(num_experts) + kw = dict( + topk_ids=ids, + topk_weights=wts, + gemm1_alpha=alpha, + gemm1_beta=beta, + gemm1_clamp_limit=limit, + ) + exp = _ref_moe(xv, ids, wts, ref, alpha, beta, limit).to(torch.bfloat16) + + out = _call(data, sf, args, num_experts, top_k, **kw) + assert out.shape == (num_tokens, 2880) + _assert_moe_close(out, exp) + + # the padded row band [2880, 2944) is all-zero weight, so widening the + # output to the padded hidden must reproduce the same values plus zeros + wide = _call(data, sf, args, num_experts, top_k, valid_hidden_size=2944, **kw) + assert wide.shape == (num_tokens, 2944) + _assert_moe_close(wide[:, :2880].contiguous(), exp) + assert torch.count_nonzero(wide[:, 2880:]) == 0 + + # ... but it never follows the widened activation: None means "use the + # hidden_states width" (3072 here), and that width is not an output width + for bad in (None, 3072): + try: + _call(data, sf, args, num_experts, top_k, valid_hidden_size=bad, **kw) + except RuntimeError: + pass + else: + raise AssertionError(f"valid_hidden_size={bad} was accepted") + + # padded intermediate columns carry zero weight, so any value at or above + # the true intermediate size is equivalent + for vis in (2880, 2944, None): + same = _call(data, sf, args, num_experts, top_k, valid_intermediate_size=vis, **kw) + _assert_moe_close(same, exp) + + # below the true intermediate size the kernel silently drops columns + short = _call(data, sf, args, num_experts, top_k, valid_intermediate_size=1024, **kw) + elt, _ = _dev(short, exp) + assert elt > 50.0, f"truncating the intermediate should change the result, got {elt:.2f} ulp" + print(" test_valid_sizes OK") + + +def test_hidden_padding_is_inert(): + """Columns of hidden_states past valid_hidden_size multiply zero weight.""" + num_experts, hidden, inter, top_k, num_tokens = 4, 2880, 1024, 2, 8 + args, ref = _build(num_experts, hidden, inter, seed=20) + gen = ref["gen"] + data, sf, _ = _rand_mxfp8(num_tokens, hidden, ref["h1_pad"], gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + alpha, beta, limit = _swiglu_params(num_experts) + kw = dict( + topk_ids=ids, + topk_weights=wts, + gemm1_alpha=alpha, + gemm1_beta=beta, + gemm1_clamp_limit=limit, + ) + base = _call(data, sf, args, num_experts, top_k, **kw) + + # arbitrary finite junk in the padded columns, with a live scale byte + junk = data.clone().float() + junk[:, hidden:] = 400.0 + junk = junk.to(torch.float8_e4m3fn) + junk_sf = sf.view(num_tokens, -1).clone() + junk_sf[:, hidden // SV :] = 133 + assert torch.equal( + _call(junk, junk_sf.reshape(-1).contiguous(), args, num_experts, top_k, **kw), + base, + ) + + # ... but a NaN there still poisons the result + nan = data.clone().float() + nan[:, hidden:] = float("nan") + nan = nan.to(torch.float8_e4m3fn) + assert not torch.isfinite( + _call(nan, junk_sf.reshape(-1).contiguous(), args, num_experts, top_k, **kw) + ).all() + print(" test_hidden_padding_is_inert OK") + + +def test_duplicate_expert_id_counts_once(): + """A repeated id in one token's row contributes once, at the first slot's weight.""" + num_experts, hidden, inter, top_k, num_tokens = 8, 512, 256, 4, 6 + args, ref = _build(num_experts, hidden, inter, seed=21) + gen = ref["gen"] + data, sf, xv = _rand_mxfp8(num_tokens, hidden, ref["h1_pad"], gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + alpha, beta, limit = _swiglu_params(num_experts) + ids[2, 1] = ids[2, 0] + out = _call( + data, + sf, + args, + num_experts, + top_k, + topk_ids=ids, + topk_weights=wts, + gemm1_alpha=alpha, + gemm1_beta=beta, + gemm1_clamp_limit=limit, + ) + first_only = wts.clone() + first_only[2, 1] = 0.0 + _assert_moe_close( + out, _ref_moe(xv, ids, first_only, ref, alpha, beta, limit).to(torch.bfloat16) + ) + elt, _ = _dev(out, _ref_moe(xv, ids, wts, ref, alpha, beta, limit).to(torch.bfloat16)) + assert elt > 10.0, f"counted-twice reference should be rejected, got {elt:.2f} ulp" + print(" test_duplicate_expert_id_counts_once OK") + + +def test_inert_arguments(): + """Knobs the contract holds inert really are inert on the certified paths.""" + num_experts, hidden, inter, top_k, num_tokens = 8, 512, 256, 4, 12 + args, ref = _build(num_experts, hidden, inter, seed=202) + gen = ref["gen"] + data, sf, _ = _rand_mxfp8(num_tokens, hidden, ref["h1_pad"], gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + lg = _logits(num_tokens, num_experts, gen) + alpha, beta, limit = _swiglu_params(num_experts) + kw = dict( + topk_ids=ids, + topk_weights=wts, + gemm1_alpha=alpha, + gemm1_beta=beta, + gemm1_clamp_limit=limit, + ) + base = _call(data, sf, args, num_experts, top_k, **kw) + + # routing_method_type is not read once routing is already done + for rmt in (0, 1, 2, 4, 5, 6): + assert torch.equal( + _call(data, sf, args, num_experts, top_k, routing_method_type=rmt, **kw), + base, + ), f"routing_method_type={rmt} changed the pre-routed result" + + # given both entry points, the pre-routed pair wins and the logits are dead + assert torch.equal(_call(data, sf, args, num_experts, top_k, routing_logits=lg, **kw), base), ( + "routing_logits changed the result although topk_ids/topk_weights were given" + ) + + # routed_scaling_factor does nothing on either certified entry point + for rsf in (1.0, 2.5): + assert torch.equal( + _call(data, sf, args, num_experts, top_k, routed_scaling_factor=rsf, **kw), + base, + ) + routed_kw = dict(gemm1_alpha=alpha, gemm1_beta=beta, gemm1_clamp_limit=limit, routing_logits=lg) + routed = _call(data, sf, args, num_experts, top_k, **routed_kw) + assert torch.equal( + _call( + data, + sf, + args, + num_experts, + top_k, + routed_scaling_factor=2.5, + **routed_kw, + ), + routed, + ) + + # autotuner bucket knobs never move the numbers + for tmax in (128, 8192): + assert torch.equal( + _call(data, sf, args, num_experts, top_k, tune_max_num_tokens=tmax, **kw), + base, + ) + for dp in (False, True): + assert torch.equal(_call(data, sf, args, num_experts, top_k, use_dp=dp, **kw), base) + print(" test_inert_arguments OK") + + +def test_reference_discriminates(): + """The tolerance gate rejects a swapped-half or unshuffled operand.""" + num_experts, hidden, inter, top_k, num_tokens = 8, 512, 256, 4, 12 + args, ref = _build(num_experts, hidden, inter, seed=22) + gen = ref["gen"] + data, sf, xv = _rand_mxfp8(num_tokens, hidden, ref["h1_pad"], gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + alpha, beta, limit = _swiglu_params(num_experts) + kw = dict( + topk_ids=ids, + topk_weights=wts, + gemm1_alpha=alpha, + gemm1_beta=beta, + gemm1_clamp_limit=limit, + ) + out = _call(data, sf, args, num_experts, top_k, **kw) + swapped = _ref_moe(xv, ids, wts, ref, alpha, beta, limit, swap_gate_up=True).to(torch.bfloat16) + elt, rms = _dev(out, swapped) + assert elt > 50.0 and rms > 20.0, f"gate/up swap not detected: {elt:.2f}/{rms:.2f} ulp" + + # a reference that skips the FC1-output requantization is also rejected + unq = _ref_moe(xv, ids, wts, ref, alpha, beta, limit, quantize_intermediate=False).to( + torch.bfloat16 + ) + elt, rms = _dev(out, unq) + assert elt > 8.0 and rms > 4.0, ( + f"unquantized intermediate not detected: {elt:.2f}/{rms:.2f} ulp" + ) + + # feeding the un-permuted (merely padded and concatenated) FC1 operand + up_c, up_s, _ = ref["up"] + gt_c, gt_s, _ = ref["gate"] + i_pad = args["intermediate_size"] + h1 = ref["h1_pad"] + up_p = (up_c[..., 0::2] | (up_c[..., 1::2] << 4)).contiguous() + gt_p = (gt_c[..., 0::2] | (gt_c[..., 1::2] << 4)).contiguous() + raw = dict(args) + raw["gemm1_weights"] = torch.cat( + [_pad3(up_p, i_pad, h1 // 2), _pad3(gt_p, i_pad, h1 // 2)], dim=1 + ).contiguous() + raw["gemm1_weights_scale"] = torch.cat( + [_pad3(up_s, i_pad, h1 // SV), _pad3(gt_s, i_pad, h1 // SV)], dim=1 + ).contiguous() + bad = _call(data, sf, raw, num_experts, top_k, **kw) + exp = _ref_moe(xv, ids, wts, ref, alpha, beta, limit).to(torch.bfloat16) + elt, rms = _dev(bad, exp) + assert elt > 50.0 and rms > 20.0, f"unshuffled FC1 not detected: {elt:.2f}/{rms:.2f} ulp" + print(" test_reference_discriminates OK") + + +def test_rejects_unsupported(): + """Domains the contract declares unsupported are rejected loudly.""" + num_experts, hidden, inter, top_k, num_tokens = 8, 512, 256, 4, 6 + args, ref = _build(num_experts, hidden, inter, seed=23) + gen = ref["gen"] + data, sf, _ = _rand_mxfp8(num_tokens, hidden, ref["h1_pad"], gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + alpha, beta, limit = _swiglu_params(num_experts) + act_kw = dict(gemm1_alpha=alpha, gemm1_beta=beta, gemm1_clamp_limit=limit) + kw = dict(topk_ids=ids, topk_weights=wts, **act_kw) + + def rejected(name, fn): + try: + fn() + except (RuntimeError, AssertionError): + return + raise AssertionError(f"{name} was accepted") + + # only the gated SwiGlu kernel family exists for MXFP8 x MXFP4 + for act in (1, 2): + rejected( + f"act_type={act}", + lambda a=act: _call(data, sf, args, num_experts, top_k, act_type=a, **kw), + ) + rejected( + "int64 topk_ids", + lambda: _call( + data, + sf, + args, + num_experts, + top_k, + topk_ids=ids.long(), + topk_weights=wts, + **act_kw, + ), + ) + rejected( + "fp32 topk_weights", + lambda: _call( + data, + sf, + args, + num_experts, + top_k, + topk_ids=ids, + topk_weights=wts.float(), + **act_kw, + ), + ) + rejected( + "no routing input", + lambda: _call(data, sf, args, num_experts, top_k, **act_kw), + ) + rejected( + "top_k == num_experts", + lambda: _call( + data, + sf, + args, + num_experts, + num_experts, + topk_ids=torch.stack( + [torch.arange(num_experts, device=DEV, dtype=torch.int32)] * num_tokens + ), + topk_weights=torch.ones(num_tokens, num_experts, device=DEV, dtype=torch.bfloat16), + **act_kw, + ), + ) + rejected( + "top_k = 0", + lambda: _call( + data, + sf, + args, + num_experts, + 0, + topk_ids=ids[:, :0].contiguous(), + topk_weights=wts[:, :0].contiguous(), + **act_kw, + ), + ) + rejected( + "bf16 gemm1_bias", + lambda: _call( + data, + sf, + {**args, "gemm1_bias": args["gemm1_bias"].bfloat16()}, + num_experts, + top_k, + **kw, + ), + ) + rejected( + "int8-typed weight scale", + lambda: _call( + data, + sf, + { + **args, + "gemm1_weights_scale": args["gemm1_weights_scale"].view(torch.int8), + }, + num_experts, + top_k, + **kw, + ), + ) + rejected( + "alpha sized 1 instead of local_num_experts", + lambda: _call( + data, + sf, + args, + num_experts, + top_k, + topk_ids=ids, + topk_weights=wts, + gemm1_alpha=alpha[:1].contiguous(), + gemm1_beta=beta, + gemm1_clamp_limit=limit, + ), + ) + rejected( + "fp32 output buffer", + lambda: _call( + data, + sf, + args, + num_experts, + top_k, + output=torch.zeros(num_tokens, hidden, device=DEV, dtype=torch.float32), + **kw, + ), + ) + rejected( + "wrong-shape output buffer", + lambda: _call( + data, + sf, + args, + num_experts, + top_k, + output=torch.zeros(num_tokens, hidden + 8, device=DEV, dtype=torch.bfloat16), + **kw, + ), + ) + rejected( + "valid_hidden_size not a multiple of 32", + lambda: _call(data, sf, args, num_experts, top_k, valid_hidden_size=500, **kw), + ) + rejected( + "intermediate_size != gemm1_weights.shape[1] // 2", + lambda: _call( + data, + sf, + args, + num_experts, + top_k, + intermediate_size=args["intermediate_size"] // 2, + **kw, + ), + ) + rejected( + "routing_logits column count != num_experts", + lambda: _call( + data, + sf, + args, + num_experts, + top_k, + routing_logits=torch.zeros( + num_tokens, num_experts + 4, device=DEV, dtype=torch.bfloat16 + ), + **act_kw, + ), + ) + rejected( + "3-D hidden_states", + lambda: _call(data.view(1, num_tokens, -1), sf, args, num_experts, top_k, **kw), + ) + rejected( + "zero tokens", + lambda: _call( + data[:0].contiguous(), + sf[:0].contiguous(), + args, + num_experts, + top_k, + topk_ids=torch.zeros(0, top_k, device=DEV, dtype=torch.int32), + topk_weights=torch.zeros(0, top_k, device=DEV, dtype=torch.bfloat16), + **act_kw, + ), + ) + print(" test_rejects_unsupported OK") + + +def test_wrapper_guards(): + """The wrapper's asserts cover cases the op itself takes silently.""" + num_experts, hidden, inter, top_k, num_tokens = 8, 512, 256, 4, 6 + args, ref = _build(num_experts, hidden, inter, seed=24) + gen = ref["gen"] + data, sf, _ = _rand_mxfp8(num_tokens, hidden, ref["h1_pad"], gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + lg = _logits(num_tokens, num_experts, gen) + alpha, beta, limit = _swiglu_params(num_experts) + kw = dict( + topk_ids=ids, + topk_weights=wts, + gemm1_alpha=alpha, + gemm1_beta=beta, + gemm1_clamp_limit=limit, + ) + base = _call(data, sf, args, num_experts, top_k, **kw) + + def raw(**over): + call = dict( + routing_logits=None, + routing_bias=None, + hidden_states=data, + hidden_states_scale=sf, + gemm1_weights=args["gemm1_weights"], + gemm1_weights_scale=args["gemm1_weights_scale"], + gemm1_bias=args["gemm1_bias"], + gemm1_alpha=alpha, + gemm1_beta=beta, + gemm1_clamp_limit=limit, + gemm2_weights=args["gemm2_weights"], + gemm2_weights_scale=args["gemm2_weights_scale"], + gemm2_bias=args["gemm2_bias"], + topk_weights=wts, + topk_ids=ids, + output=None, + routing_method_type=1, + ) + call.update(over) + return torch.ops.trtllm.mxe4m3_mxe2m1_block_scale_moe_runner( + call["routing_logits"], + call["routing_bias"], + call["hidden_states"], + call["hidden_states_scale"], + call["gemm1_weights"], + call["gemm1_weights_scale"], + call["gemm1_bias"], + call["gemm1_alpha"], + call["gemm1_beta"], + call["gemm1_clamp_limit"], + call["gemm2_weights"], + call["gemm2_weights_scale"], + call["gemm2_bias"], + num_experts, + top_k, + None, + None, + args["intermediate_size"], + args["valid_hidden_size"], + args["valid_intermediate_size"], + 0, + num_experts, + None, + call["routing_method_type"], + 0, + topk_weights=call["topk_weights"], + topk_ids=call["topk_ids"], + output=call["output"], + ) + + def strided(t, extra=4): + fat = torch.zeros(*t.shape[:-1], t.shape[-1] + extra, dtype=t.dtype, device=DEV) + fat[..., : t.shape[-1]] = t + view = fat[..., : t.shape[-1]] + assert not view.is_contiguous() + return view + + # a 1-D scale buffer keeps its element count but reads every other byte + spread = torch.zeros(sf.numel() * 2, dtype=torch.uint8, device=DEV) + spread[::2] = sf + sf_strided = torch.as_strided(spread, (sf.numel(),), (2,)) + assert not sf_strided.is_contiguous() + + # every tensor the wrapper guards: strided is silently honoured by the op + for name in ( + "hidden_states", + "hidden_states_scale", + "gemm1_weights", + "gemm1_weights_scale", + "gemm1_bias", + "gemm2_weights", + "gemm2_weights_scale", + "gemm2_bias", + "topk_weights", + "topk_ids", + ): + src = { + "hidden_states": data, + "hidden_states_scale": sf, + "topk_weights": wts, + "topk_ids": ids, + }.get(name, args.get(name)) + bad = sf_strided if name == "hidden_states_scale" else strided(src) + got = raw(**{name: bad}) + assert not torch.equal(got, base), f"{name}: strided view was NOT silently wrong" + if name == "hidden_states": + guarded = lambda: _call(bad, sf, args, num_experts, top_k, **kw) # noqa: E731 + elif name == "hidden_states_scale": + guarded = lambda: _call(data, bad, args, num_experts, top_k, **kw) # noqa: E731 + elif name in ("topk_weights", "topk_ids"): + guarded = lambda: _call( # noqa: E731 + data, sf, args, num_experts, top_k, **{**kw, name: bad} + ) + else: + guarded = lambda: _call( # noqa: E731 + data, sf, {**args, name: bad}, num_experts, top_k, **kw + ) + try: + guarded() + except AssertionError: + pass + else: + raise AssertionError(f"wrapper did not reject a strided {name}") + + # a strided routing_logits is likewise silently honoured + got = raw(routing_logits=strided(lg), topk_weights=None, topk_ids=None) + ok = raw(routing_logits=lg, topk_weights=None, topk_ids=None) + assert not torch.equal(got, ok) + + # a strided output buffer is written as if dense + buf = torch.zeros(num_tokens, hidden + 4, device=DEV, dtype=torch.bfloat16) + raw(output=buf[:, :hidden]) + assert not torch.equal(buf[:, :hidden], base) + + # routing_bias is a no-op for every non-grouped routing method + bias = torch.zeros(num_experts, dtype=torch.float32, device=DEV) + bias[0], bias[num_experts - 1] = 1.0e3, -1.0e3 + for rmt in (0, 1, 4, 6): + with_bias = raw( + routing_logits=lg, + routing_bias=bias, + topk_weights=None, + topk_ids=None, + routing_method_type=rmt, + ) + without = raw( + routing_logits=lg, + routing_bias=None, + topk_weights=None, + topk_ids=None, + routing_method_type=rmt, + ) + assert torch.equal(with_bias, without), f"routing_bias was honoured at rmt={rmt}" + try: + _call( + data, + sf, + args, + num_experts, + top_k, + routing_logits=lg, + routing_bias=bias, + routing_method_type=rmt, + gemm1_alpha=alpha, + gemm1_beta=beta, + gemm1_clamp_limit=limit, + ) + except AssertionError: + pass + else: + raise AssertionError(f"wrapper did not reject routing_bias at rmt={rmt}") + # ... and on the pre-routed entry point, whatever the routing method + try: + _call( + data, + sf, + args, + num_experts, + top_k, + routing_bias=bias, + routing_method_type=2, + **kw, + ) + except AssertionError: + pass + else: + raise AssertionError("wrapper did not reject routing_bias on the pre-routed path") + + # folding the bias into the logits is the caller's job and does change the result + shifted = _call( + data, + sf, + args, + num_experts, + top_k, + routing_logits=(lg.float() + bias).to(torch.bfloat16).contiguous(), + gemm1_alpha=alpha, + gemm1_beta=beta, + gemm1_clamp_limit=limit, + ) + assert not torch.equal(shifted, ok) + print(" test_wrapper_guards OK") + + +def test_determinism_and_purity(): + """Two identical calls agree bitwise; no input is mutated.""" + num_experts, hidden, inter, top_k, num_tokens = 8, 512, 256, 4, 32 + args, ref = _build(num_experts, hidden, inter, seed=25) + gen = ref["gen"] + data, sf, _ = _rand_mxfp8(num_tokens, hidden, ref["h1_pad"], gen) + ids, wts = _routing(num_tokens, num_experts, top_k, gen) + alpha, beta, limit = _swiglu_params(num_experts) + kw = dict( + topk_ids=ids, + topk_weights=wts, + gemm1_alpha=alpha, + gemm1_beta=beta, + gemm1_clamp_limit=limit, + ) + snaps = {k: v.clone() for k, v in args.items() if torch.is_tensor(v)} + data_snap, sf_snap = data.clone(), sf.clone() + ids_snap, wts_snap = ids.clone(), wts.clone() + a = _call(data, sf, args, num_experts, top_k, **kw) + b = _call(data, sf, args, num_experts, top_k, **kw) + assert torch.equal(a, b) + for k, v in snaps.items(): + assert torch.equal(args[k], v), f"{k} was mutated" + assert torch.equal(data, data_snap) and torch.equal(sf, sf_snap) + assert torch.equal(ids, ids_snap) and torch.equal(wts, wts_snap) + print(" test_determinism_and_purity OK") diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op.md b/tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op.md new file mode 100644 index 000000000000..1ec9ce5f3ce4 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op.md @@ -0,0 +1,244 @@ +--- +receipts: + sm_100: {status: passed, trtllm: 1.3.0rc21} + sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 16} +--- + +# noaux_tc_op + +**Wraps** `torch.ops.trtllm.noaux_tc_op` (one call). + +## Semantics + +DeepSeek-V3 style MoE routing (`noaux_tc` is DeepSeek's `topk_method` name for +auxiliary-loss-free routing with a bias correction term): per-token expert +selection by a *bias-corrected* sigmoid score, with combine weights taken from +the *uncorrected* sigmoid. + +Let `L` be one row of `router_logits` (`num_experts` values) and `b` the +per-expert correction bias. For every token, with the precision of each step +annotated: + +``` +s = sigmoid(L) # fp32; see "The kernel's sigmoid" +c = s + b # fp32; SELECTION SCORE ONLY + +# group stage — active only when n_group > 1 (identity at n_group == 1): +# split c into n_group contiguous groups of num_experts/n_group experts +# gscore[g] = sum of the two largest c inside group g +# keep the topk_group groups with the largest gscore, set c to -inf elsewhere + +ids = indices of the topk largest c, ordered by decreasing c # int32 +w = s[ids] # from BARE s — the bias is NOT in here +tot = sum(w) # fp32 accumulation +w = fp64(w) / (fp64(tot) + eps) * routed_scaling_factor # fp64 +``` + +The call returns `(w, ids)` — **weights first, ids second**. + +The renormalization and the scaling are evaluated in **fp64** — pinned by +observation: an all-fp32 reference disagrees with the kernel on a third of the +elements at `routed_scaling_factor = 2.5`, an fp64 one is bit-exact. Only the +sum `tot` is accumulated in fp32. `eps` is a guard term that keeps a row whose +selected scores are all exactly zero at exactly zero instead of `NaN`; HF uses +`1e-20` and that value reproduces the kernel everywhere, but the constant +itself is not observable from outside (this kernel's sigmoid returns either +exactly `0` or at least `2^-25`, so any `eps <= 1e-20` behaves identically). + +Three consequences that a reader of the argument name `scores` would get +wrong, each verified here: + +- The first argument is **raw router logits**, not scores. The sigmoid is + applied *inside* the call. Passing pre-sigmoided values silently routes on + `sigmoid(sigmoid(x))`. +- The bias enters **selection only**. The returned weights are gathered from + the bias-free sigmoid, so they are not the values that were ranked. +- The renormalization (`norm_topk_prob: true` in HF configs) and the + `routed_scaling_factor` multiply are both **done inside**. There is no flag + to skip either; a model with `norm_topk_prob: false` cannot use this op. + Each returned row sums to `routed_scaling_factor` — except a row whose + selected scores all saturated to zero, which sums to 0 (see below). + +**The kernel's sigmoid.** The kernel evaluates `0.5 * tanh(0.5 * x) + 0.5` in +fp32, not `1 / (1 + exp(-x))`. The two are algebraically identical, and over +the logit range a router actually produces (`|x| <= 8`) they agree well inside +dtype tolerance — a `torch.sigmoid` reference reproduces this op there, which +is what makes it a drop-in for the HF DeepSeek-V3 gate. But the tanh form +saturates early: it returns **exactly 1.0** for `x >= ~17` and **exactly 0.0** +for `x <= ~-18.5`, and its relative error against the exponential form grows +through the negative tail (measured 4.3e-5 at `x = -8`, 2.6e-2 at `x = -15`). +A row whose *selected* experts all sit below `~-18.5` therefore comes back as +all-zero weights (the `eps` turns `0/0` into `0`, not `NaN`) where the +exponential form would return `routed_scaling_factor / topk` each. A torch +reference built on `0.5 * torch.tanh(0.5 * x) + 0.5` reproduced this op at +every shape, dtype and logit range in the test — out to `|logits| ~ 40`. + +**Ordering and ties.** Within a row the slots are sorted by decreasing +selection score `c`, and equal `c` values are ordered by *increasing* expert +index — i.e. the result equals a stable descending sort of `c` truncated to +`topk`. This differs from `torch.topk`, which on this machine emits the +*larger* index first among equals; the difference is visible whenever the +logits are bf16/fp16 (exact ties are common) or the bias is coarse. Note the +returned *weights* are therefore **not** monotone: they are `s` read out in +`c` order. + +**Fusion boundary.** The call is routing only. The router GEMM +(`hidden_states @ gate.weight.T`) that produces `router_logits` is the +caller's, and so is everything downstream: no expert permutation, no expert +GEMMs, no combine, no shared-expert branch. The two returned tensors are +shaped and typed as the `topk_weights` / `topk_ids` pair that trtllm's +pre-routed MoE runners take (see *Notes* for the dtype those runners demand). + +## Signature + +```python +def noaux_tc_op( + router_logits: torch.Tensor, + bias: torch.Tensor, + n_group: int, + topk_group: int, + topk: int, + routed_scaling_factor: float, +) -> tuple[torch.Tensor, torch.Tensor] +``` + +The wrapper renames the op's first two schema arguments (`scores`, `bias`) to +`router_logits`, `bias`; they are passed positionally and unchanged. + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `router_logits` | `[num_tokens, num_experts]` | bf16 / fp16 / fp32 | **contiguous** (see Preconditions) | CUDA | +| `bias` | `[num_experts]` | bf16 / fp16 / fp32, independent of `router_logits` except one combination (see Preconditions) | **contiguous** | CUDA, same device as `router_logits` | +| `n_group` | scalar | Python int, `>= 1` | — | — | +| `topk_group` | scalar | Python int | — | — | +| `topk` | scalar | Python int, `0 <= topk <= min(32, num_experts)` | — | — | +| `routed_scaling_factor` | scalar | Python float (any finite value, including 0 and negative) | — | — | +| returns `[0]` (`topk_weights`) | `[num_tokens, topk]` | **`router_logits.dtype`** | contiguous, newly allocated | CUDA (same device as `router_logits`) | +| returns `[1]` (`topk_ids`) | `[num_tokens, topk]` | **int32** | contiguous, newly allocated | CUDA (same device as `router_logits`) | + +The weight dtype follows `router_logits`, **not** `bias` — a bf16 logits / +fp32 bias call returns bf16 weights. Neither input is mutated (verified +bitwise) and neither is retained; both outputs are fresh allocations. + +## Metadata consumed + +None. Stateless — no attention metadata, no cache, no workspace, no +process-global tuner state. Two identical calls return bitwise identical +tensors, and the call runs on the ambient CUDA stream (verified: an +alternate-stream call returns the same result). + +## Preconditions + +- `router_logits` is a **2D** CUDA tensor. 1D and 3D inputs raise + `RuntimeError: scores must be a 2D Tensor`; a CPU tensor raises + `NotImplementedError` from the dispatcher. +- `bias` is a **1D** CUDA tensor with exactly `num_experts` elements, on the + same device as `router_logits`. Anything else raises `RuntimeError: bias + must be 1D with length == number of experts`; a CPU bias with CUDA logits + raises `RuntimeError: scores and bias must be CUDA tensors`. The bias is + mandatory — there is no "no bias" mode; pass a zero tensor. +- **Both tensors must be contiguous. The kernel ignores strides**: it + addresses `router_logits` as a dense `[num_tokens, num_experts]` row-major + buffer and `bias` as a dense `[num_experts]` buffer, each from `data_ptr()`. + A column slice of a wider buffer, a transposed view, or a strided bias is + silently routed on the wrong elements and returns a plausible-looking wrong + answer — it never raises. The wrapper asserts both; call `.contiguous()` + first (e.g. when the router logits are a slice of a fused projection). +- Dtypes: `router_logits` and `bias` are each bf16, fp16 or fp32. Anything + else raises `ValueError: Invalid dtype, only supports float16, float32, and + bfloat16` (or the `Invalid bias dtype` variant). The two are independent + **except** that `router_logits` fp16 + `bias` bf16 is rejected with `Invalid + bias dtype`, even though the corresponding kernel is compiled into this + build — an op-level dispatch gap, not a kernel limit. The other eight + combinations all work. +- `0 <= topk <= 32`. `topk > 32` raises `RuntimeError: topk should be smaller + than or equal to 32 for now`; a negative `topk` raises from the output + allocation. `topk == 0` is accepted and returns two `[num_tokens, 0]` + tensors. +- **`topk <= num_experts` is not enforced by the op.** With `topk > + num_experts` the kernel reads past the end of each row and emits expert ids + `>= num_experts` at non-zero weight (an out-of-bounds read; the values are + meaningless but **finite and deterministic** — at `num_experts = 3, + topk = 16` eight independent processes, three of them churning the caching + allocator first, each returned 104 of 128 ids out of range, identical + weights, and no NaN, so the read stays inside the logits tensor's own + buffer rather than picking up allocator residue). It never raises. The + wrapper asserts this. +- `num_experts % n_group == 0`, else `RuntimeError: num_experts should be + divisible by n_group`. `n_group <= 32`, else `RuntimeError: n_group should + be smaller than or equal to 32 for now`. +- **`n_group >= 1` is not checked.** `n_group == 0` reaches a host-side + `num_experts % n_group` and raises **SIGFPE, killing the process** — not a + Python exception, so it cannot be caught. The wrapper does not guard it + (guarding would require a non-metadata policy call); callers must never + pass 0. +- Configuration support, enforced inside the kernel with `RuntimeError: + [TensorRT-LLM][ERROR] Assertion failed: invokeNoAuxTc: unsupported + configuration (n_group=..., num_experts=..., topk_group=..., topk=...)`: + - `n_group == 1` (the ungrouped case): requires only `num_experts <= 1024`. + `topk_group` is then **completely ignored** — values `0, 1, 2, 5, 32, 100` + all give bitwise identical results. `num_experts = 1025` is rejected. + - `n_group > 1`: requires `1 <= topk_group <= n_group`, `topk <= 8`, + `num_experts <= 256`, and `num_experts / n_group <= 32` (experts per + group). `topk_group == 0` is *accepted* without error at `n_group > 1` but + is not characterized by this entry — do not use it. + - `topk_group == n_group` keeps every group and is bitwise identical to + `n_group = 1`. +- `num_tokens == 0` is accepted and returns empty `[0, topk]` tensors. + +A caller violating none of the above gets the result described under +*Semantics*. + +## Notes + +- **Certified surface.** Passing on sm_100 / trtllm 1.3.0rc21: `num_tokens` + in `{0, 1, 2, 4, 7, 8, 16, 64, 128, 256, 512, 1024, 2048, 4096, 8192}`; + `num_experts` in `{1, 2, 7, 8, 16, 32, 64, 72, 100, 128, 256, 257, 512, + 1024}`; `topk` in `{0, 1, 2, 3, 4, 6, 8, 16, 31, 32}`; all eight accepted + (logits, bias) dtype pairs; `routed_scaling_factor` in `{-1, 0, 0.5, 1, 2, + 2.5, 3, 1000}`; ungrouped and the grouped configurations `(num_experts, + n_group, topk_group, topk)` = `(256,8,4,8)`, `(128,4,2,6)`, `(72,8,2,6)`, + `(72,4,2,6)`, `(64,8,4,8)`. +- **Numerics: the reference of *Semantics* is not an approximation, it is the + kernel.** Expert ids matched exactly everywhere. Weights were bit-identical + except for a last-bit rounding step that appears when the kernel's fp32 + accumulation order for `tot` differs from `torch.sum`'s — measured at 1 + element in 49152 (8192 tokens x 72 experts, bf16) and 4 in ~10^7 over a + soak. Every element stayed within **one ulp** of the reference, so this + entry's test gates at `rtol = finfo(dtype).eps, atol = 0` (an order of + magnitude tighter than `assert_close`'s defaults for fp32) plus a + bit-exactness budget. Getting the fp64 detail of *Semantics* wrong is not a + rounding difference: an all-fp32 normalization misses a third of the + elements at `routed_scaling_factor = 2.5`. +- **The `register_fake` meta function disagrees with the kernel.** Under + `FakeTensorMode` (torch.compile / export / AD tracing) the first output is + given `bias.dtype`; eager execution gives `router_logits.dtype`. They differ + whenever the two input dtypes differ — e.g. fp32 logits with a bf16 bias + traces as bf16 and runs as fp32. Eager callers are unaffected. This is an + upstream defect in this build, not a behaviour to rely on. +- **Downstream dtype match.** `torch.ops.trtllm.fp4_block_scale_moe_runner`, + the NVFP4 pre-routed MoE runner, was observed here to reject a fp32 + `topk_weights` (`RuntimeError: topk_weights must be bfloat16.`) and an int64 + `topk_ids` (`RuntimeError: topk_ids must be int`), and to accept the + bf16/int32 pair. Since this op's weight dtype follows `router_logits`, + feeding it **bf16** router logits produces a directly consumable pair with + no cast. (Catalog membership of that runner is `index.yaml`'s fact alone.) +- Sigmoid, bias correction, group scoring and selection are computed in fp32 + regardless of the input dtype (the normalization then steps up to fp64); + only the final weights are cast back to `router_logits.dtype`. +- `-inf` logits map to score `0` exactly, like any logit below `~-18.5`, so a + `-inf`-masked expert loses to every expert with a positive score (verified: + none was selected while `topk` fit in the unmasked experts). It is *not* a + hard mask — with more slots than positively-scored experts the masked ones + come back at weight 0. `+inf` saturates to 1 like any logit above `~17`. + `NaN` logits *are* selected and yield `NaN` weights. +- A hand-written PyTorch path with the same semantics exists in this build + (`tensorrt_llm._torch.modules.fused_moe.routing.Deepseekv3RoutingImpl`) and + is what trtllm falls back to outside the supported configurations above; it + uses `torch.sigmoid`, so it differs from this op in the saturating tails and + in tie-breaking. +- Related routing ops in this build: `torch.ops.trtllm.renorm_moe_routing_op` + and `torch.ops.trtllm.default_moe_routing_op` (softmax routing, no bias), + and the block-scale MoE runners' built-in `routing_method_type = 2` + (DeepSeekV3) path, which is a different code path from this op. Catalog + membership is `index.yaml`'s fact alone. diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op.py b/tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op.py new file mode 100644 index 000000000000..904495aee2c2 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op.py @@ -0,0 +1,49 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""DeepSeek-V3 style MoE routing: sigmoid + bias-corrected group-limited top-k.""" + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + + +def noaux_tc_op( + router_logits: torch.Tensor, + bias: torch.Tensor, + n_group: int, + topk_group: int, + topk: int, + routed_scaling_factor: float, +) -> tuple[torch.Tensor, torch.Tensor]: + """Route each token to `topk` experts by biased sigmoid score. + + Selection uses `sigmoid(router_logits) + bias`; the combine weights are + gathered from the *unbiased* `sigmoid(router_logits)`, renormalized to sum + to 1 and scaled by `routed_scaling_factor`. + + Returns `(topk_weights [T, topk] in router_logits.dtype, topk_ids + [T, topk] int32)`, both newly allocated and contiguous. + """ + # Pure-metadata guards for three domains the op does not police. All three + # were observed on this machine to return a plausible-looking wrong answer + # instead of raising: + # - the kernel addresses `router_logits` as a dense [T, num_experts] + # row-major buffer and `bias` as a dense [num_experts] buffer, both from + # data_ptr(), ignoring strides; + # - topk > num_experts makes the kernel read past the end of a row and + # emit expert ids >= num_experts at non-zero weight. + assert router_logits.is_contiguous(), ( + "router_logits must be contiguous; a strided view is read as a dense " + "[num_tokens, num_experts] buffer and silently gives wrong routing" + ) + assert bias.is_contiguous(), ( + "bias must be contiguous; a strided view is read as a dense " + "[num_experts] buffer and silently gives wrong routing" + ) + assert topk <= router_logits.shape[-1], ( + f"topk ({topk}) exceeds num_experts ({router_logits.shape[-1]}); the " + "kernel reads out of bounds and emits out-of-range expert ids" + ) + return torch.ops.trtllm.noaux_tc_op( + router_logits, bias, n_group, topk_group, topk, routed_scaling_factor + ) diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op_test.py b/tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op_test.py new file mode 100644 index 000000000000..7aee4956c92f --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op_test.py @@ -0,0 +1,499 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the noaux_tc_op catalog entry.""" + +import torch + +from .noaux_tc_op import noaux_tc_op + +assert torch.cuda.is_available(), "noaux_tc_op requires a CUDA device" + + +def _kernel_sigmoid(x: torch.Tensor) -> torch.Tensor: + """The sigmoid the kernel actually evaluates: 0.5 * tanh(x/2) + 0.5, fp32. + + Algebraically identical to 1/(1+exp(-x)) but numerically different in the + tails: it saturates to exactly 1.0 at x >= ~17 and to exactly 0.0 at + x <= ~-18.5. Pinned by test_sigmoid_is_the_tanh_form. + """ + return 0.5 * torch.tanh(0.5 * x.float()) + 0.5 + + +def _ref( + router_logits: torch.Tensor, + bias: torch.Tensor, + n_group: int, + topk_group: int, + topk: int, + routed_scaling_factor: float, + sigmoid=_kernel_sigmoid, +) -> tuple[torch.Tensor, torch.Tensor]: + """fp32 reference built from native torch ops only. + + A *stable* descending sort keeps equal selection scores in ascending expert + order, which is the kernel's observed tie-break. + """ + num_experts = router_logits.shape[-1] + scores = sigmoid(router_logits) + choice = scores + bias.float() + if n_group > 1: + grouped = choice.view(-1, n_group, num_experts // n_group) + group_score = torch.topk(grouped, k=2, dim=-1).values.sum(-1) + keep = torch.sort(group_score, dim=-1, descending=True, stable=True).indices[:, :topk_group] + mask = torch.zeros_like(group_score).scatter_(-1, keep, 1.0) + mask = mask.unsqueeze(-1).expand_as(grouped).reshape(choice.shape) + choice = torch.where(mask.bool(), choice, torch.tensor(float("-inf"), device=choice.device)) + ids = torch.sort(choice, dim=-1, descending=True, stable=True).indices[:, :topk] + weights = torch.gather(scores, 1, ids) + # The scores are summed in fp32, but the division and the scaling are + # evaluated in fp64 (pinned by test_routed_scaling_factor: an all-fp32 + # normalization disagrees on a third of the elements at + # routed_scaling_factor = 2.5, an fp64 one is bit-exact). + total = weights.sum(-1, keepdim=True).double() + weights = weights.double() / (total + 1e-20) * routed_scaling_factor + return weights.to(router_logits.dtype), ids.to(torch.int32) + + +def _check( + router_logits: torch.Tensor, + bias: torch.Tensor, + n_group: int, + topk_group: int, + topk: int, + routed_scaling_factor: float, +) -> tuple[torch.Tensor, torch.Tensor]: + weights, ids = noaux_tc_op( + router_logits, bias, n_group, topk_group, topk, routed_scaling_factor + ) + ref_weights, ref_ids = _ref( + router_logits, bias, n_group, topk_group, topk, routed_scaling_factor + ) + num_tokens = router_logits.shape[0] + assert weights.shape == (num_tokens, topk), weights.shape + assert ids.shape == (num_tokens, topk), ids.shape + assert weights.dtype == router_logits.dtype, weights.dtype + assert ids.dtype == torch.int32, ids.dtype + assert weights.is_contiguous() and ids.is_contiguous() + assert weights.device == router_logits.device + assert ids.device == router_logits.device + torch.testing.assert_close(ids, ref_ids) + # Dual gate. The fp32 reference is not an approximation of this kernel, it + # is the kernel — the only observed disagreement is the final rounding of + # the renormalized weight, which can step by one last-bit unit because the + # kernel adds up the topk scores in a different order than torch.sum + # (measured: 1 element in 49152 at 8192x72 bf16). So the tolerance is + # exactly one ulp (rtol = the dtype's eps is the largest relative size one + # ulp can have; atol = 0), and on top of that almost every element must be + # bit-identical. + torch.testing.assert_close( + weights, + ref_weights, + rtol=torch.finfo(router_logits.dtype).eps, + atol=0.0, + ) + mismatched = int((weights != ref_weights).sum()) + budget = 1 + weights.numel() // 10000 + assert mismatched <= budget, ( + f"{mismatched} of {weights.numel()} weights are not bit-exact (budget " + f"{budget}): dtype={router_logits.dtype} shape={tuple(router_logits.shape)} " + f"n_group={n_group} topk={topk}" + ) + return weights, ids + + +def test_deepseek_v3_lite_config() -> None: + # 72 routed experts, top-6, n_group=1, topk_group=1, scaling 2.0, bf16 + # router logits and a bf16 [72] correction bias. Decode-like through + # prefill-like token counts. + torch.manual_seed(0) + bias = torch.randn(72, dtype=torch.bfloat16, device="cuda") * 0.1 + for num_tokens in [1, 2, 7, 64, 512, 2048, 4096, 8192]: + logits = torch.randn(num_tokens, 72, dtype=torch.bfloat16, device="cuda") + weights, _ = _check(logits, bias, 1, 1, 6, 2.0) + # renormalized then scaled: every row sums to routed_scaling_factor + torch.testing.assert_close( + weights.float().sum(-1), + torch.full((num_tokens,), 2.0, device="cuda"), + rtol=8e-3, # 6 bf16 addends, each rounded to 8 mantissa bits + atol=0.0, + ) + + +def test_matches_hf_deepseek_sigmoid_reference() -> None: + # The HF DeepSeek-V3 gate written with the textbook sigmoid. Over the + # logit range a real router produces (|logit| <= 8) the kernel's tanh-form + # sigmoid and torch.sigmoid agree to well inside dtype tolerance. + torch.manual_seed(1) + for dtype in [torch.bfloat16, torch.float32]: + for num_tokens in [1, 256, 2048]: + logits = (torch.randn(num_tokens, 72, device="cuda") * 2.5).to(dtype) + bias = (torch.randn(72, device="cuda") * 0.1).to(dtype) + weights, ids = noaux_tc_op(logits, bias, 1, 1, 6, 2.0) + ref_weights, ref_ids = _ref( + logits, bias, 1, 1, 6, 2.0, sigmoid=lambda x: torch.sigmoid(x.float()) + ) + torch.testing.assert_close(ids, ref_ids) + torch.testing.assert_close(weights, ref_weights) + + +def test_selection_uses_bias_but_weights_do_not() -> None: + # The trap this entry exists to pin: the bias enters selection only. + # A bias large enough to reorder the experts must move the ids, and the + # returned weights must come from the *unbiased* sigmoid. + torch.manual_seed(2) + logits = torch.randn(512, 72, dtype=torch.float32, device="cuda") + bias = torch.randn(72, dtype=torch.float32, device="cuda") + weights, ids = noaux_tc_op(logits, bias, 1, 1, 6, 2.0) + + unbiased_ids = torch.sort( + _kernel_sigmoid(logits), dim=-1, descending=True, stable=True + ).indices[:, :6] + assert not torch.equal(ids, unbiased_ids.to(torch.int32)), ( + "the bias did not affect selection — test input too weak" + ) + + scores = _kernel_sigmoid(logits) + gathered = torch.gather(scores, 1, ids.long()) + expect = gathered / (gathered.sum(-1, keepdim=True) + 1e-20) * 2.0 + torch.testing.assert_close(weights, expect) + + # negative control: gathering the *biased* score instead is a different, + # plausible-looking answer that the kernel does not produce + biased = torch.gather(scores + bias, 1, ids.long()) + wrong = biased / (biased.sum(-1, keepdim=True) + 1e-20) * 2.0 + assert (weights - wrong).abs().max().item() > 1e-2, ( + "biased-weight reference is indistinguishable — test input too weak" + ) + + +def test_slot_order_is_descending_by_biased_score() -> None: + torch.manual_seed(3) + logits = torch.randn(256, 72, dtype=torch.bfloat16, device="cuda") + bias = torch.randn(72, dtype=torch.bfloat16, device="cuda") + _, ids = noaux_tc_op(logits, bias, 1, 1, 6, 2.0) + choice = _kernel_sigmoid(logits) + bias.float() + selected = torch.gather(choice, 1, ids.long()) + assert bool((selected[:, :-1] >= selected[:, 1:]).all()), ( + "slots are not ordered by decreasing selection score" + ) + # the same row is *not* sorted by the unbiased score, i.e. the ordering + # key really is the biased one + unbiased = torch.gather(_kernel_sigmoid(logits), 1, ids.long()) + assert not bool((unbiased[:, :-1] >= unbiased[:, 1:]).all()) + + +def test_ties_break_toward_lower_expert_index() -> None: + # Opposite to torch.topk, which emits the larger index first among equals. + logits = torch.zeros(4, 16, dtype=torch.bfloat16, device="cuda") + bias = torch.zeros(16, dtype=torch.bfloat16, device="cuda") + logits[0, :] = 1.0 + logits[1, 3] = 5.0 + logits[1, 7] = 5.0 + logits[2, ::2] = 2.0 + logits[3, 0] = 1.0 + logits[3, 15] = 1.0 + _, ids = noaux_tc_op(logits, bias, 1, 1, 4, 2.0) + expected = torch.tensor( + [[0, 1, 2, 3], [3, 7, 0, 1], [0, 2, 4, 6], [0, 15, 1, 2]], + dtype=torch.int32, + device="cuda", + ) + torch.testing.assert_close(ids, expected) + _check(logits, bias, 1, 1, 4, 2.0) + + # tie-saturated inputs: ties from the logits, and ties created by a coarse + # bias on top of continuous logits + torch.manual_seed(4) + for num_tokens, num_experts, topk in [(64, 128, 8), (512, 72, 6), (256, 16, 4)]: + coarse = torch.randint(0, 3, (num_tokens, num_experts), device="cuda").to(torch.bfloat16) + zero = torch.zeros(num_experts, dtype=torch.bfloat16, device="cuda") + _check(coarse, zero, 1, 1, topk, 2.0) + smooth = torch.randn(num_tokens, num_experts, dtype=torch.bfloat16, device="cuda") + coarse_bias = torch.randint(0, 2, (num_experts,), device="cuda").to(torch.bfloat16) + _check(smooth, coarse_bias, 1, 1, topk, 2.0) + + +def test_dtype_matrix() -> None: + # weights come back in the *router_logits* dtype; the bias dtype is + # independent, except that fp16 logits reject a bf16 bias. + torch.manual_seed(5) + base = torch.randn(128, 72, device="cuda") + bias_base = torch.randn(72, device="cuda") * 0.1 + dtypes = [torch.bfloat16, torch.float16, torch.float32] + for logits_dtype in dtypes: + for bias_dtype in dtypes: + if logits_dtype is torch.float16 and bias_dtype is torch.bfloat16: + continue # rejected; covered by test_op_rejects_unsupported_domains + logits = base.to(logits_dtype) + bias = bias_base.to(bias_dtype) + weights, _ = _check(logits, bias, 1, 1, 6, 2.0) + assert weights.dtype == logits_dtype, ( + logits_dtype, + bias_dtype, + weights.dtype, + ) + + +def test_routed_scaling_factor() -> None: + torch.manual_seed(6) + logits = torch.randn(256, 72, dtype=torch.float32, device="cuda") + bias = torch.randn(72, dtype=torch.float32, device="cuda") * 0.1 + for scaling in [1.0, 2.0, 2.5, 3.0, 0.5, 0.0, -1.0, 1000.0]: + weights, _ = _check(logits, bias, 1, 1, 6, scaling) + torch.testing.assert_close(weights.sum(-1), torch.full((256,), scaling, device="cuda")) + # negative control for the fp64 normalization: an all-fp32 reference + # is a different answer for scaling factors that are not a power of two + scores = _kernel_sigmoid(logits) + _, ids = noaux_tc_op(logits, bias, 1, 1, 6, scaling) + gathered = torch.gather(scores, 1, ids.long()) + fp32_form = gathered / (gathered.sum(-1, keepdim=True) + 1e-20) * scaling + differs = int((weights != fp32_form).sum()) + if scaling in (2.5, 3.0, 1000.0): + assert differs > 0.1 * weights.numel(), ( + f"fp32 and fp64 normalization are indistinguishable at {scaling}" + ) + + +def test_sigmoid_is_the_tanh_form() -> None: + # Pins the kernel's sigmoid out to |logits| ~ 40, where the tanh form has + # long saturated. _check's reference is the tanh form; the exp form is + # ruled out separately in test_saturation_floor_and_ceiling. + torch.manual_seed(7) + for dtype in [torch.bfloat16, torch.float16, torch.float32]: + for scale in [1.0, 4.0, 10.0, 20.0, 40.0]: + logits = (torch.randn(256, 72, device="cuda") * scale).to(dtype) + bias = (torch.randn(72, device="cuda") * 0.1).to(dtype) + weights, _ = _check(logits, bias, 1, 1, 6, 2.0) + assert bool(torch.isfinite(weights.float()).all()) + + +def test_saturation_floor_and_ceiling() -> None: + # The tail behaviour that follows from the tanh form. + bias = torch.zeros(16, dtype=torch.float32, device="cuda") + high = torch.full((2, 16), 20.0, dtype=torch.float32, device="cuda") + weights, _ = noaux_tc_op(high, bias, 1, 1, 4, 2.0) + # every selected score saturates to exactly 1.0 -> equal weights + assert torch.equal(weights, torch.full_like(weights, 0.5)), weights + + low = torch.full((2, 16), -20.0, dtype=torch.float32, device="cuda") + weights, _ = noaux_tc_op(low, bias, 1, 1, 4, 2.0) + # every selected score saturates to exactly 0.0; the guard term in the + # denominator turns 0/0 into exactly 0 instead of NaN + assert torch.equal(weights, torch.zeros_like(weights)), weights + + # Negative control that separates the two sigmoid forms: with every + # selected logit inside the tanh form's lossy band, a 1/(1+exp(-x)) + # reference lands three orders of magnitude outside the one-ulp gate + # _check applies (fp32 eps * 0.33 ~ 4e-8). + torch.manual_seed(8) + band = (torch.rand(128, 72, device="cuda") * 12.0 - 20.0).to(torch.float32) + zero = torch.zeros(72, dtype=torch.float32, device="cuda") + weights, _ = _check(band, zero, 1, 1, 6, 2.0) + exp_weights, _ = _ref(band, zero, 1, 1, 6, 2.0, sigmoid=lambda x: torch.sigmoid(x.float())) + assert (weights - exp_weights).abs().max().item() > 1e-5, ( + "exp-form and tanh-form references are indistinguishable here" + ) + + +def test_grouped_routing() -> None: + # n_group > 1: group score = sum of the two largest biased scores in the + # group; the topk_group best groups survive, the rest are masked out. + torch.manual_seed(9) + for num_experts, n_group, topk_group, topk in [ + (256, 8, 4, 8), # DeepSeek-V3 full + (128, 4, 2, 6), + (72, 8, 2, 6), + (72, 4, 2, 6), + (64, 8, 4, 8), + ]: + for num_tokens in [1, 64, 1024]: + logits = torch.randn(num_tokens, num_experts, dtype=torch.bfloat16, device="cuda") + bias = torch.randn(num_experts, dtype=torch.bfloat16, device="cuda") * 0.1 + _check(logits, bias, n_group, topk_group, topk, 2.0) + # tie-saturated grouping + coarse = torch.randint(0, 3, (64, num_experts), device="cuda").to(torch.bfloat16) + zero = torch.zeros(num_experts, dtype=torch.bfloat16, device="cuda") + _check(coarse, zero, n_group, topk_group, topk, 2.0) + + # keeping every group is the same as not grouping at all + logits = torch.randn(64, 256, dtype=torch.bfloat16, device="cuda") + bias = torch.randn(256, dtype=torch.bfloat16, device="cuda") * 0.1 + grouped = noaux_tc_op(logits, bias, 8, 8, 8, 2.0) + flat = noaux_tc_op(logits, bias, 1, 1, 8, 2.0) + assert torch.equal(grouped[0], flat[0]) and torch.equal(grouped[1], flat[1]) + + +def test_topk_group_ignored_when_n_group_is_one() -> None: + torch.manual_seed(10) + logits = torch.randn(64, 72, dtype=torch.bfloat16, device="cuda") + bias = torch.randn(72, dtype=torch.bfloat16, device="cuda") * 0.1 + reference = noaux_tc_op(logits, bias, 1, 1, 6, 2.0) + for topk_group in [0, 1, 2, 5, 32, 100]: + weights, ids = noaux_tc_op(logits, bias, 1, topk_group, 6, 2.0) + assert torch.equal(weights, reference[0]) and torch.equal(ids, reference[1]), ( + f"topk_group={topk_group} changed the result at n_group=1" + ) + + +def test_num_experts_and_topk_sweep() -> None: + torch.manual_seed(11) + for num_experts in [1, 2, 7, 32, 64, 72, 100, 128, 256, 257, 512, 1024]: + logits = torch.randn(8, num_experts, dtype=torch.bfloat16, device="cuda") + bias = torch.randn(num_experts, dtype=torch.bfloat16, device="cuda") + _check(logits, bias, 1, 1, min(6, num_experts), 2.0) + logits = torch.randn(64, 64, dtype=torch.bfloat16, device="cuda") + bias = torch.randn(64, dtype=torch.bfloat16, device="cuda") + for topk in [1, 2, 3, 6, 8, 16, 31, 32]: + _check(logits, bias, 1, 1, topk, 2.0) + # topk == num_experts returns the full ranking + small = torch.randn(16, 8, dtype=torch.bfloat16, device="cuda") + small_bias = torch.randn(8, dtype=torch.bfloat16, device="cuda") + _check(small, small_bias, 1, 1, 8, 2.0) + + +def test_degenerate_shapes() -> None: + logits = torch.randn(8, 72, dtype=torch.bfloat16, device="cuda") + bias = torch.randn(72, dtype=torch.bfloat16, device="cuda") + weights, ids = noaux_tc_op(logits, bias, 1, 1, 0, 2.0) + assert weights.shape == (8, 0) and ids.shape == (8, 0) + empty = torch.randn(0, 72, dtype=torch.bfloat16, device="cuda") + weights, ids = noaux_tc_op(empty, bias, 1, 1, 6, 2.0) + assert weights.shape == (0, 6) and ids.shape == (0, 6) + assert weights.dtype == torch.bfloat16 and ids.dtype == torch.int32 + + +def test_input_not_mutated_and_deterministic() -> None: + torch.manual_seed(12) + logits = torch.randn(4096, 72, dtype=torch.bfloat16, device="cuda") + bias = torch.randn(72, dtype=torch.bfloat16, device="cuda") + logits_snapshot, bias_snapshot = logits.clone(), bias.clone() + weights_a, ids_a = noaux_tc_op(logits, bias, 1, 1, 6, 2.0) + weights_b, ids_b = noaux_tc_op(logits, bias, 1, 1, 6, 2.0) + assert torch.equal(logits, logits_snapshot), "router_logits was mutated" + assert torch.equal(bias, bias_snapshot), "bias was mutated" + assert torch.equal(weights_a, weights_b) and torch.equal(ids_a, ids_b), ( + "two identical calls disagreed" + ) + assert weights_a.data_ptr() != weights_b.data_ptr() + assert weights_a.data_ptr() != logits.data_ptr() + + +def test_op_rejects_unsupported_domains() -> None: + logits = torch.randn(8, 72, dtype=torch.bfloat16, device="cuda") + bias = torch.randn(72, dtype=torch.bfloat16, device="cuda") + big = torch.randn(8, 256, dtype=torch.bfloat16, device="cuda") + big_bias = torch.randn(256, dtype=torch.bfloat16, device="cuda") + cases = [ + ("1D router_logits", lambda: noaux_tc_op(logits[0], bias, 1, 1, 6, 2.0)), + ( + "3D router_logits", + lambda: noaux_tc_op(logits.unsqueeze(0), bias, 1, 1, 6, 2.0), + ), + ( + "fp64 router_logits", + lambda: noaux_tc_op(logits.double(), bias, 1, 1, 6, 2.0), + ), + ("int32 router_logits", lambda: noaux_tc_op(logits.int(), bias, 1, 1, 6, 2.0)), + ("fp64 bias", lambda: noaux_tc_op(logits, bias.double(), 1, 1, 6, 2.0)), + ("int32 bias", lambda: noaux_tc_op(logits, bias.int(), 1, 1, 6, 2.0)), + ( + "fp16 logits + bf16 bias", + lambda: noaux_tc_op(logits.half(), bias, 1, 1, 6, 2.0), + ), + ("2D bias", lambda: noaux_tc_op(logits, bias.unsqueeze(0), 1, 1, 6, 2.0)), + ("short bias", lambda: noaux_tc_op(logits, bias[:71], 1, 1, 6, 2.0)), + ("cpu tensors", lambda: noaux_tc_op(logits.cpu(), bias.cpu(), 1, 1, 6, 2.0)), + ("bias on cpu", lambda: noaux_tc_op(logits, bias.cpu(), 1, 1, 6, 2.0)), + ("negative topk", lambda: noaux_tc_op(logits, bias, 1, 1, -1, 2.0)), + ("topk 33", lambda: noaux_tc_op(big, big_bias, 1, 1, 33, 2.0)), + ( + "num_experts 1025", + lambda: noaux_tc_op( + torch.randn(8, 1025, dtype=torch.bfloat16, device="cuda"), + torch.randn(1025, dtype=torch.bfloat16, device="cuda"), + 1, + 1, + 6, + 2.0, + ), + ), + ( + "n_group does not divide num_experts", + lambda: noaux_tc_op(logits, bias, 5, 1, 6, 2.0), + ), + ( + "n_group 33", + lambda: noaux_tc_op( + torch.randn(8, 132, dtype=torch.bfloat16, device="cuda"), + torch.randn(132, dtype=torch.bfloat16, device="cuda"), + 33, + 1, + 6, + 2.0, + ), + ), + ( + "grouped topk_group > n_group", + lambda: noaux_tc_op(big, big_bias, 8, 9, 8, 2.0), + ), + ("grouped topk > 8", lambda: noaux_tc_op(big, big_bias, 8, 4, 9, 2.0)), + ( + "grouped experts_per_group > 32", + lambda: noaux_tc_op(big, big_bias, 4, 2, 8, 2.0), + ), + ( + "grouped num_experts > 256", + lambda: noaux_tc_op( + torch.randn(8, 512, dtype=torch.bfloat16, device="cuda"), + torch.randn(512, dtype=torch.bfloat16, device="cuda"), + 8, + 4, + 8, + 2.0, + ), + ), + ] + for tag, call in cases: + try: + call() + except (RuntimeError, ValueError, NotImplementedError): + continue + raise AssertionError(f"op accepted an unsupported domain: {tag}") + + +def test_wrapper_rejects_silently_wrong_inputs() -> None: + torch.manual_seed(13) + bias = torch.randn(72, dtype=torch.bfloat16, device="cuda") + wide = torch.randn(8, 144, dtype=torch.bfloat16, device="cuda") + logits = torch.randn(8, 72, dtype=torch.bfloat16, device="cuda") + wide_bias = torch.randn(144, dtype=torch.bfloat16, device="cuda") + + # A column slice is not routed on its own elements: the kernel reads the + # first 8*72 elements of the underlying storage as a dense [8, 72] buffer. + sliced = wide[:, :72] + assert not sliced.is_contiguous() + reinterpreted = wide.reshape(-1)[: 8 * 72].reshape(8, 72).contiguous() + raw = torch.ops.trtllm.noaux_tc_op(sliced, bias, 1, 1, 6, 2.0) + dense = torch.ops.trtllm.noaux_tc_op(reinterpreted, bias, 1, 1, 6, 2.0) + correct = torch.ops.trtllm.noaux_tc_op(sliced.contiguous(), bias, 1, 1, 6, 2.0) + assert torch.equal(raw[1], dense[1]), "strided read is not a dense reinterpretation" + assert not torch.equal(raw[1], correct[1]), "strided view happened to be correct" + + bad_inputs = [ + ("strided router_logits", lambda: noaux_tc_op(sliced, bias, 1, 1, 6, 2.0)), + ( + "transposed router_logits", + lambda: noaux_tc_op(wide.t()[:72, :], bias, 1, 1, 6, 2.0), + ), + ("strided bias", lambda: noaux_tc_op(logits, wide_bias[::2], 1, 1, 6, 2.0)), + ("topk > num_experts", lambda: noaux_tc_op(logits, bias, 1, 1, 73, 2.0)), + ] + for tag, call in bad_inputs: + try: + call() + except AssertionError: + continue + raise AssertionError(f"wrapper accepted a silently-wrong input: {tag}") + + # the same data made contiguous routes correctly + _check(sliced.contiguous(), bias, 1, 1, 6, 2.0) diff --git a/tensorrt_llm/_torch/staircase/catalog/norm/__init__.py b/tensorrt_llm/_torch/staircase/catalog/norm/__init__.py new file mode 100644 index 000000000000..13d44653e838 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/norm/__init__.py @@ -0,0 +1,3 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Normalization entries.""" diff --git a/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm.md b/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm.md new file mode 100644 index 000000000000..3ccab4df4b45 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm.md @@ -0,0 +1,100 @@ +--- +receipts: + sm_100: {status: passed, trtllm: 1.3.0rc21} + sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 5} +--- + +# flashinfer_fused_add_rmsnorm + +**Wraps** `torch.ops.trtllm.flashinfer_fused_add_rmsnorm` (one call). + +## Semantics + +Fused residual add + RMS normalization, in-place on both tensor arguments. +One call computes, per row: + +``` +h = fp32(x) + fp32(residual) # fp32 accumulation +residual = cast(h, dtype) # overwritten in place +x = cast(h / sqrt(mean(h^2) + eps) * weight, dtype) # overwritten +``` + +The normalization reads the fp32 `h`, not the rounded `residual` output: +the add, the squared-mean reduction, and the weight scaling all happen in +fp32 inside the kernel before the final cast back to the input dtype. + +Fusion boundary: residual add + rmsnorm + weight scaling, nothing else. +No `(1 + weight)` gemma-style scaling (see +`flashinfer_gemma_fused_add_rmsnorm`) and no quantization (see +`flashinfer_fused_add_rmsnorm_quant`). The caller owns everything before +(producing `x`, e.g. an attention/MLP output) and after (consuming the +normed `x` and the updated `residual`). + +## Signature + +```python +def flashinfer_fused_add_rmsnorm( + x: torch.Tensor, residual: torch.Tensor, weight: torch.Tensor, eps: float +) -> None +``` + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `x` | `[num_tokens, hidden]` | fp16 / bf16 / fp32 | last-dim stride must be 1; row stride may exceed `hidden` | CUDA | +| `residual` | `[num_tokens, hidden]` | same as `x` | same constraint as `x` | CUDA (same device) | +| `weight` | `[hidden]` | same as `x` | contiguous | CUDA (same device) | +| `eps` | scalar | Python float | — | — | +| returns | — | — | — | — | + +Returns `None`. Both `x` and `residual` are mutated in place: on return, +`residual` holds the pre-norm sum `x + residual` (the residual stream for +the next layer) and `x` holds the normalized, weight-scaled output. + +## Metadata consumed + +None. Stateless. + +## Preconditions + +- `x` and `residual` are 2D on the same CUDA device, same shape, same + dtype. The kernel is compiled for `[M, H]`; do not pass 3D tensors (it + would read `M = shape[0]`), but the guard is **loud, not silent**: the + CuTe compiled-kernel argument check raises `ValueError: Mismatched + Tensor on argument #0 ... expected ndim=2` before launch. Nothing is + silently skipped (measured 2026-07-28); the precondition stands, its + earlier justification did not. +- `weight.shape == (hidden,)`, `weight.dtype == x.dtype`, contiguous. +- Dtype is one of fp16, bf16, fp32. `float64`, `int8`, `uint8` and + `float8_e5m2` raise `KeyError`, but **`float8_e4m3fn` does not** — the + CuTe DSL path's dtype table maps it and returns a plausible fp8 + rmsnorm, so a caller who forgets to dequantize gets a result instead of + an error (measured on the certified path, 2026-07-28). The uncertified + CUDA-JIT path (`FLASHINFER_USE_CUDA_NORM=1`) does reject it — this + bullet described that path. +- `x.stride(-1) == 1` and `residual.stride(-1) == 1`. If either tensor is + non-contiguous (row stride > `hidden`), a strided kernel variant is + selected; its symbolic row stride is declared divisible by the kernel + vector size (up to 8 elements for bf16/fp16, 4 for fp32) and data + pointers are assumed 16-byte aligned, so arbitrary odd row strides or + misaligned slices are outside the contract. +- `hidden` has no divisibility requirement: sizes not divisible by the + 128-bit vector width (111, 1152) were verified correct on this machine. +- The caller must not need the pre-call contents of `x` or `residual` + afterward — both are destroyed. + +## Notes + +- The op is registered only when flashinfer is importable + (`IS_FLASHINFER_AVAILABLE`); this pinned install ships + flashinfer-python 0.6.14, which routes to the CuTe DSL kernel + (`fused_add_rmsnorm_cute`); a CUDA JIT fallback exists behind + `FLASHINFER_USE_CUDA_NORM=1` but is not what these receipts certify. +- Programmatic dependent launch (PDL) is controlled by the env var + `TRTLLM_ENABLE_PDL` (default enabled) inside the trtllm custom op; + it affects scheduling only, not results. +- The kernel uses fast-math `rsqrt`; results matched the fp32 torch + reference at default `assert_close` tolerances for all three dtypes. +- TRT-LLM's own `RMSNorm` module routes to this op only for fp16/bf16 + inputs, but fp32 works and passed the test on sm_100. +- Contiguous inputs with `num_tokens * hidden > 2^31 - 1` are + transparently routed to the strided (int64-offset) kernel variant. diff --git a/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm.py b/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm.py new file mode 100644 index 000000000000..92fa4acffc5f --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm.py @@ -0,0 +1,14 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""In-place fused residual add + RMS normalization via the flashinfer kernel.""" + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + + +def flashinfer_fused_add_rmsnorm( + x: torch.Tensor, residual: torch.Tensor, weight: torch.Tensor, eps: float +) -> None: + """In-place: residual += x; x = rmsnorm(residual) * weight. Returns None.""" + torch.ops.trtllm.flashinfer_fused_add_rmsnorm(x, residual, weight, eps) diff --git a/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm_test.py b/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm_test.py new file mode 100644 index 000000000000..0560b9e23477 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm_test.py @@ -0,0 +1,85 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the flashinfer_fused_add_rmsnorm catalog entry.""" + +import torch + +from .flashinfer_fused_add_rmsnorm import flashinfer_fused_add_rmsnorm + +assert torch.cuda.is_available(), "flashinfer_fused_add_rmsnorm requires a CUDA device" + + +def _ref( + x: torch.Tensor, residual: torch.Tensor, weight: torch.Tensor, eps: float +) -> tuple[torch.Tensor, torch.Tensor]: + """fp32-accumulated reference for the fused op. + + h = fp32(x) + fp32(residual); the norm reads the fp32 h, not the + rounded residual output, matching the kernel. + """ + h = x.float() + residual.float() + normed = h * torch.rsqrt(h.pow(2).mean(dim=-1, keepdim=True) + eps) + x_out = (normed * weight.float()).to(x.dtype) + residual_out = h.to(residual.dtype) + return x_out, residual_out + + +def _check(x: torch.Tensor, residual: torch.Tensor, weight: torch.Tensor, eps: float) -> None: + ref_x, ref_res = _ref(x, residual, weight, eps) + flashinfer_fused_add_rmsnorm(x, residual, weight, eps) + torch.testing.assert_close(residual, ref_res) + torch.testing.assert_close(x, ref_x) + + +def test_bf16_2d() -> None: + torch.manual_seed(0) + # decode-like (few tokens) and prefill-like (many tokens) shapes + for num_tokens, hidden in [(1, 4096), (4, 5120), (2048, 4096)]: + x = torch.randn(num_tokens, hidden, dtype=torch.bfloat16, device="cuda") + r = torch.randn(num_tokens, hidden, dtype=torch.bfloat16, device="cuda") + w = torch.randn(hidden, dtype=torch.bfloat16, device="cuda") + _check(x, r, w, 1e-6) + + +def test_bf16_unaligned_hidden() -> None: + # hidden sizes not divisible by the 128-bit vector width + torch.manual_seed(1) + for hidden in [111, 1152]: + x = torch.randn(16, hidden, dtype=torch.bfloat16, device="cuda") + r = torch.randn(16, hidden, dtype=torch.bfloat16, device="cuda") + w = torch.randn(hidden, dtype=torch.bfloat16, device="cuda") + _check(x, r, w, 1e-6) + + +def test_bf16_strided_rows() -> None: + # last-dim contiguous slices of wider buffers (row stride != hidden); + # row stride 8192 is divisible by the kernel vector size + torch.manual_seed(2) + x_buf = torch.randn(8, 8192, dtype=torch.bfloat16, device="cuda") + r_buf = torch.randn(8, 8192, dtype=torch.bfloat16, device="cuda") + x, r = x_buf[:, :4096], r_buf[:, :4096] + assert not x.is_contiguous() and x.stride(-1) == 1 + w = torch.randn(4096, dtype=torch.bfloat16, device="cuda") + # snapshot the untouched right halves before the in-place call + right_before = torch.cat([x_buf[:, 4096:], r_buf[:, 4096:]]).clone() + _check(x, r, w, 1e-6) + # mutation must land in the parent buffers' left halves only + assert torch.equal(torch.cat([x_buf[:, 4096:], r_buf[:, 4096:]]), right_before) + + +def test_fp16_2d() -> None: + torch.manual_seed(3) + for num_tokens, hidden in [(2, 4096), (1024, 2048)]: + x = torch.randn(num_tokens, hidden, dtype=torch.float16, device="cuda") + r = torch.randn(num_tokens, hidden, dtype=torch.float16, device="cuda") + w = torch.randn(hidden, dtype=torch.float16, device="cuda") + _check(x, r, w, 1e-6) + + +def test_fp32_2d() -> None: + torch.manual_seed(4) + for num_tokens, hidden in [(2, 4096), (1024, 2048)]: + x = torch.randn(num_tokens, hidden, dtype=torch.float32, device="cuda") + r = torch.randn(num_tokens, hidden, dtype=torch.float32, device="cuda") + w = torch.randn(hidden, dtype=torch.float32, device="cuda") + _check(x, r, w, 1e-6) diff --git a/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm.md b/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm.md new file mode 100644 index 000000000000..4b1bdaab9f78 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm.md @@ -0,0 +1,74 @@ +--- +receipts: + sm_100: {status: passed, trtllm: 1.3.0rc21} + sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 6} +--- + +# flashinfer_rmsnorm + +**Wraps** `torch.ops.trtllm.flashinfer_rmsnorm` (one call). + +## Semantics + +Root-mean-square normalization over the last dimension, with elementwise +weight scaling: + +``` +out[..., i] = x[..., i] / sqrt(mean(x[..., :]^2) + eps) * weight[i] +``` + +The squared-mean reduction and normalization are accumulated in fp32 +inside the kernel; the result is cast back to the input dtype. + +Fusion boundary: the single call computes normalization and weight scaling +only. There is no residual add (see `flashinfer_fused_add_rmsnorm` for +that), no `(1 + weight)` gemma-style scaling (see `flashinfer_gemma_rmsnorm`), +and no quantization. The caller owns everything else. + +## Signature + +```python +def flashinfer_rmsnorm(x: torch.Tensor, weight: torch.Tensor, eps: float) -> torch.Tensor +``` + +| Argument | Shape | Dtype | Layout | Device | +|---|---|---|---|---| +| `x` | `[num_tokens, hidden]` or `[batch, num_heads, head_dim]` | fp16 / bf16 / fp32 | last-dim stride must be 1; non-contiguous row strides are handled | CUDA | +| `weight` | `[hidden]` (or `[head_dim]` for 3D input) | same as `x` | contiguous | CUDA (same device as `x`) | +| `eps` | scalar | Python float | — | — | +| returns | same shape as `x` | same as `x` | newly allocated | CUDA | + +The output is a new tensor (`x` is not mutated). + +## Metadata consumed + +None. Stateless. + +## Preconditions + +- `x` is 2D or 3D on a CUDA device; the normalized dim is the last one. +- `weight.shape == (x.shape[-1],)` and `weight.dtype == x.dtype`. +- `x.stride(-1) == 1` (rows may be strided, e.g. a column slice of a wider + buffer; the kernel selects a strided code path in that case). +- Dtype is one of fp16, bf16, fp32. `float64`, `int8`, `uint8` and + `float8_e5m2` raise `KeyError`, but **`float8_e4m3fn` does not** — the + CuTe DSL path's dtype table maps it and returns a plausible fp8 + rmsnorm, so a caller who forgets to dequantize gets a result instead of + an error (measured on the certified path, 2026-07-28). The uncertified + CUDA-JIT path (`FLASHINFER_USE_CUDA_NORM=1`) does reject it — this + bullet described that path. +- `hidden` (last dim) has no divisibility requirement: sizes not divisible + by the 128-bit vector width (e.g. 111, 1152) fall back to a smaller + vector size and were verified correct on this machine. + +## Notes + +- The op is registered only when flashinfer is importable + (`IS_FLASHINFER_AVAILABLE`); this pinned install ships + flashinfer-python 0.6.14. +- Programmatic dependent launch (PDL) is controlled by the env var + `TRTLLM_ENABLE_PDL` (default enabled) inside the trtllm custom op; + it affects scheduling only, not results. +- TRT-LLM's own `RMSNorm` module routes to this op only for fp16/bf16 + inputs, but fp32 works and passed the test on sm_100. +- Weight scaling is plain `w * normed(x)`, not the gemma `(1 + w)` form. diff --git a/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm.py b/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm.py new file mode 100644 index 000000000000..6082c216094d --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm.py @@ -0,0 +1,12 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""RMS normalization over the last dim via the flashinfer rmsnorm kernel.""" + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + + +def flashinfer_rmsnorm(x: torch.Tensor, weight: torch.Tensor, eps: float) -> torch.Tensor: + """Return `x / sqrt(mean(x^2, dim=-1) + eps) * weight` as a new tensor.""" + return torch.ops.trtllm.flashinfer_rmsnorm(x, weight, eps) diff --git a/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm_test.py b/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm_test.py new file mode 100644 index 000000000000..18274f9189e2 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm_test.py @@ -0,0 +1,74 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the flashinfer_rmsnorm catalog entry.""" + +import torch + +from .flashinfer_rmsnorm import flashinfer_rmsnorm + +assert torch.cuda.is_available(), "flashinfer_rmsnorm requires a CUDA device" + + +def _ref_rmsnorm(x: torch.Tensor, weight: torch.Tensor, eps: float) -> torch.Tensor: + """fp32-accumulated reference: x / sqrt(mean(x^2, -1) + eps) * weight.""" + xf = x.float() + normed = xf * torch.rsqrt(xf.pow(2).mean(dim=-1, keepdim=True) + eps) + return (normed * weight.float()).to(x.dtype) + + +def _check(x: torch.Tensor, weight: torch.Tensor, eps: float) -> None: + out = flashinfer_rmsnorm(x, weight, eps) + ref = _ref_rmsnorm(x, weight, eps) + assert out.shape == x.shape and out.dtype == x.dtype + torch.testing.assert_close(out, ref) + + +def test_bf16_2d() -> None: + torch.manual_seed(0) + # decode-like (few tokens) and prefill-like (many tokens) shapes + for num_tokens, hidden in [(1, 4096), (4, 5120), (2048, 4096)]: + x = torch.randn(num_tokens, hidden, dtype=torch.bfloat16, device="cuda") + w = torch.randn(hidden, dtype=torch.bfloat16, device="cuda") + _check(x, w, 1e-6) + + +def test_bf16_3d_qk_norm_shape() -> None: + torch.manual_seed(1) + x = torch.randn(16, 32, 128, dtype=torch.bfloat16, device="cuda") + w = torch.randn(128, dtype=torch.bfloat16, device="cuda") + _check(x, w, 1e-5) + + +def test_bf16_unaligned_hidden() -> None: + # hidden sizes not divisible by the 128-bit vector width + torch.manual_seed(2) + for hidden in [111, 1152]: + x = torch.randn(16, hidden, dtype=torch.bfloat16, device="cuda") + w = torch.randn(hidden, dtype=torch.bfloat16, device="cuda") + _check(x, w, 1e-6) + + +def test_bf16_strided_rows() -> None: + # last-dim contiguous slice of a wider buffer (row stride != hidden) + torch.manual_seed(3) + buf = torch.randn(8, 8192, dtype=torch.bfloat16, device="cuda") + x = buf[:, :4096] + assert not x.is_contiguous() and x.stride(-1) == 1 + w = torch.randn(4096, dtype=torch.bfloat16, device="cuda") + _check(x, w, 1e-6) + + +def test_fp16_2d() -> None: + torch.manual_seed(4) + for num_tokens, hidden in [(2, 4096), (1024, 2048)]: + x = torch.randn(num_tokens, hidden, dtype=torch.float16, device="cuda") + w = torch.randn(hidden, dtype=torch.float16, device="cuda") + _check(x, w, 1e-6) + + +def test_fp32_2d() -> None: + torch.manual_seed(5) + for num_tokens, hidden in [(2, 4096), (1024, 2048)]: + x = torch.randn(num_tokens, hidden, dtype=torch.float32, device="cuda") + w = torch.randn(hidden, dtype=torch.float32, device="cuda") + _check(x, w, 1e-6) diff --git a/tensorrt_llm/_torch/staircase/catalog/quantization/__init__.py b/tensorrt_llm/_torch/staircase/catalog/quantization/__init__.py new file mode 100644 index 000000000000..7d0c28eefd77 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/quantization/__init__.py @@ -0,0 +1,3 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Quantization entries.""" diff --git a/tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize.md b/tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize.md new file mode 100644 index 000000000000..76cb1ca311d1 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize.md @@ -0,0 +1,341 @@ +--- +receipts: + sm_100: {status: passed, trtllm: 1.3.0rc21} + sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 21} +--- + +# fp4_quantize + +**Wraps** `torch.ops.trtllm.fp4_quantize` (one call). + +## Semantics + +Block-scaled FP4 quantization of an activation tensor: one call converts +values to 4-bit **e2m1** elements (two packed per output byte) plus one +scale byte per `sf_vec_size` **contiguous elements along the last dim**. +Nothing else is fused in — no normalization, no activation, no transpose, +and (unlike `mxfp8_quantize`) **no width padding**: the last dim is used +exactly as given, so any padding a consumer needs is the caller's own +`pad` before the call. The caller owns producing the input and owns both +outputs afterwards. + +The op has two modes, selected by `sf_use_ue8m0`: + +| Mode | `sf_vec_size` | `sf_use_ue8m0` | `global_scale` | Scale byte | +|---|---|---|---|---| +| **NVFP4** | 16 | `False` | required | e4m3 (UE4M3 in practice) | +| MXFP4 | 32 | `True` | optional | UE8M0 power of two | + +Any other combination raises. **Only the NVFP4 mode is certified here**; +what follows describes it. + +Let `M` = product of all leading dims, `K` = last dim, `cols = K / 16`. +For row `m` and block `b` of 16 elements, with `g` = the `global_scale` +value: + +``` +vecmax = max |x[m, 16b : 16b+16]| # exact in the input dtype +sf[m, b] = e4m3_round_to_nearest_even( g * vecmax / 6 ) # saturates at 448 +out_scale = g / float(sf[m, b]) # 0 when vecmax == 0 +data[m, i] = e2m1_round_to_nearest_even( x[m, i] * out_scale ) # saturates at +-6 +``` + +`6` is e2m1's largest magnitude and `448` is e4m3's largest finite value, +so the canonical `g = 448 * 6 / amax(x)` keeps every block's `sf` inside +e4m3's finite range. + +**Mind the direction.** A modelopt NVFP4 checkpoint stores its per-tensor +`input_scale` as `amax / (448*6)` — the **reciprocal** of this `g` — so a +caller passes `1 / input_scale` here, never `input_scale`. Passing it +straight is accepted and silently destructive: measured on `64 x 7168` +bf16 with `amax = 5.1875`, it drove **96.4 % of the scale bytes to zero** +(those blocks read back as all-zero, the `g`-underflow path in +*Preconditions*) while the surviving blocks saturated — max reconstructed +magnitude 6.07 against a true 5.1875. The result is neither an error nor +an obviously dead tensor. + +**Dequantization is exactly `data * sf / g`.** A consumer GEMM folds +`1/g` into its own `alpha` (typically `alpha = (amax_input / (448*6)) * +weight_scale_2`); this op emits `data` and `sf` only, and never `g` or +`1/g`. + +e2m1 encodes 8 magnitudes — `0, 0.5, 1, 1.5, 2, 3, 4, 6` — as +`code = (exponent << 1) | mantissa`, with bit 3 the sign; `-0.0` keeps +its sign bit (code 8). Rounding is **round-to-nearest, ties to even +code**, saturating: `0.25 -> 0`, `0.75 -> 1.0`, `1.25 -> 1.0`, +`1.75 -> 2.0`, `2.5 -> 2.0`, `3.5 -> 4.0`, `5.0 -> 4.0`, and anything +above `6` clamps to `6`. + +Because `sf` has only a 3-bit mantissa, `sf` may round *below* `g*vecmax/6`, +which makes `out_scale` slightly larger than `6/vecmax` and lets the +block's largest element clip to `6`. That is the format, not an error: it +happened in about half the blocks of random bf16 input, and the +reconstruction still obeys the bound below. + +**Tie non-determinism (the one place this op is not bit-reproducible from +a plain fp32 model).** The kernel forms `out_scale` through two +`rcp.approx.ftz.f32` reciprocals (~2^-23 relative error each) rather than +an exact division. An element whose exact `x * out_scale` sits within a +few fp32 ulps of an e2m1 midpoint (`0.25, 0.75, 1.25, 1.75, 2.5, 3.5, +5.0`) can therefore land on either neighbouring code. Every other element +matches the formula above bit for bit, and the scale bytes are always +exact. The op itself is deterministic: repeating the same call gives the +same bytes. + +**What sets how many such elements there are: the global scale, not the +shape.** A near-tie needs the block's `sf` to come out *exactly* +`g * vecmax / 6` — then `out_scale` is exactly `6 / vecmax` and +`6 * x / vecmax` is an exact small rational, which lands on a midpoint +for a sizeable minority of the block's lanes. Measured on this machine +(bf16, `g = 448*6/amax`, the canonical convention): + +- **the width is inert.** One 16,515,072-element sample reshaped to + `K` = 7168, 18432 and 2048 gave the *identical* counts — 34,405 + near-tie elements and 12,091 kernel/reference disagreements in all + three. Scale blocks are 16 *contiguous* elements, so a reshape never + moves a block boundary; `M` and `K` reach the arithmetic only through + `amax`, hence through `g`. +- **`g` is the whole story.** With that `g`, 0.68 % of blocks had an + exactly representable `sf`. Multiplying `g` by 1.001 left **no** block + with one, and the same call then matched the reference **bit for bit**, + zero near-ties — while multiplying by 2 (a power of two, so `sf` scales + exactly and `out_scale` is unchanged) kept the population at 33,113. +- across the 24 bf16 shapes of the R1 sweep, near-ties were **0.03 % to + 0.60 %** of all elements and the kernel disagreed with an exact-fp32 + reference on **0.00 % to 0.14 %** — always by one adjacent e2m1 code + with the sign unchanged. + +So the population is a property of the *exactly calibrated* case. A +serving call passes a static per-tensor `input_scale` rather than +`448*6/amax` of the tensor in hand, which is the regime measured above as +bit-exact — but that was one `g` offset (0.1 %), not a proof for every +static scale, and no run here used a real checkpoint's scale. + +### Output layouts + +`data` is **identical byte for byte under both `is_sf_swizzled_layout` +values** — only the scale buffer's size and element order change. + +- `is_sf_swizzled_layout=False` (**linear**): `M * cols` bytes, row-major. + Byte `b` of row `m` is at flat offset `m * cols + b`; `sf.view(M, cols)` + is exactly the per-block scale matrix. +- `is_sf_swizzled_layout=True` (**128x4 swizzled**): + `pad_up(M, 128) * pad_up(cols, 4)` bytes. The scale of `(m, b)` sits at + flat offset + + ``` + (b % 4) + + (b // 4) * 512 # 4 cols x 128 rows per column group + + (m % 32) * 16 + + ((m % 128) // 32) * 4 + + (m // 128) * 128 * pad_up(cols, 4) + ``` + + Every offset not addressed by a real `(m, b)` pair — row padding up to + `pad_up(M,128)`, column padding up to `pad_up(cols,4)` — is `0x00`. + +## Signature + +```python +def fp4_quantize( + input: torch.Tensor, + global_scale: Optional[torch.Tensor], + sf_vec_size: int, + sf_use_ue8m0: bool = False, + is_sf_swizzled_layout: bool = True, +) -> tuple[torch.Tensor, torch.Tensor] +``` + +| Argument | Shape | Dtype | Layout / value | Device | +|---|---|---|---|---| +| `input` | `[..., K]`, rank >= 2 | bfloat16 or float16 (e4m3 accepted, uncertified) | contiguous | CUDA | +| `global_scale` | exactly 1 element (`[1]` or 0-D) | float32 | positive, finite; `448*6/amax` by convention | CUDA | +| `sf_vec_size` | scalar | Python int | `16` (NVFP4). `32` only with `sf_use_ue8m0=True` | — | +| `sf_use_ue8m0` | scalar | Python bool | `False` for NVFP4 e4m3 scales | — | +| `is_sf_swizzled_layout` | scalar | Python bool | `True` = 128x4 swizzled scale order, `False` = linear | — | +| returns `data` | `[..., K/2]` (leading dims preserved) | uint8 | contiguous, freshly allocated; element `2i` in the low nibble, `2i+1` in the high nibble | same as `input` | +| returns `sf` | 1-D: `[M*cols]` (linear) or `[pad_up(M,128)*pad_up(cols,4)]` (swizzled) | uint8 | contiguous, freshly allocated; reinterpret as `float8_e4m3fn` to read the values | same as `input` | + +`M` = product of `input`'s leading dims, `cols = K / sf_vec_size`. +`input` is not modified. Leading dims are collapsed for the scale buffer +in both layouts (a `[B, S, K]` input yields a scale buffer laid out for +`B*S` rows), while `data` keeps `input`'s leading dims. + +The op's own schema spells the parameters `globalScale`, `sfVecSize`, +`sfUseUE8M0`, `isSfSwizzledLayout`; all five are positional here. + +## Metadata consumed + +None. Stateless — no runtime, attention metadata, or workspace. + +## Preconditions + +- `input` is on CUDA, contiguous, rank >= 2, dtype bfloat16 / float16 / + float8_e4m3fn. Violations raise: a CPU tensor, a non-contiguous view + (column slice or transpose), a 1-D tensor and an fp32 tensor were each + observed to raise on this machine rather than compute something wrong. +- `K % sf_vec_size == 0`. Raises otherwise. This is the whole width + domain, and it was re-measured as such: `K` = 2048, 7168 and 18432 — + the last 1.5x wider than any previously certified width — behave + exactly like 2560 / 3072 / 12288, at every row count from 1 to 8192. + Nothing about the op is a function of `K` beyond `cols = K / 16` (see + *Notes*: one kernel, no width-keyed dispatch). +- `sf_vec_size == 16` with `sf_use_ue8m0=False`, or `sf_vec_size == 32` + with `sf_use_ue8m0=True`. Every other pairing (including + `sf_vec_size=8` or `0`) raises. +- `global_scale` is required when `sf_use_ue8m0=False` (passing `None` + raises), must be a **float32 CUDA** tensor (a half tensor and a CPU + tensor each raise), and must hold **exactly one element**. A tensor + with more elements is *not* rejected by the op: it silently reads + element 0 and applies it to every row — this build has no per-token + global scale. The wrapper asserts the element count for that reason. +- `global_scale` must be **positive and finite**: + - `g = 0` yields an all-zero scale buffer and every data lane at `+6`; + a consumer dividing by `g` gets NaN. + - `g < 0` yields **negative** e4m3 scale bytes (e.g. `0xB0`), which + every consumer reads as unsigned UE4M3 — silently wrong. +- Every block must satisfy `2^-10 < g * vecmax / 6 <= 448` (or + `vecmax == 0`). Both ends fail silently: + - **Above 448** (a stale/mis-calibrated static `input_scale`, or a + runtime `amax` above the calibration one): `sf` clamps to `0x7E` + (448) and `out_scale` collapses to `g/448`. Only lanes above + `5 * 448 / g` then clip to `+-6`; the rest quantize normally against + the collapsed scale and still dequantize close to their true value — + at 1.17x overshoot a hand-built block came back + `[6, 4, 2, 1, 0.5, 0, 0, 0]`, not all-`+-6` (measured 2026-07-28). + **Do not look for an all-`+-6`, signs-only block as the signature**: + a mis-calibrated block looks ordinary, and only ~100x overshoot + produces the saturated form. + - **At or below `2^-10`** (`g` far too small): `sf` rounds to `0x00`, + the reciprocal of that zero scale is `+inf`, so every lane saturates + to `+6` (exact zeros included, via `0 * inf = NaN` which the e2m1 + cast saturates). Because `sf` is zero, a consumer reads the whole + block back as zeros. One binade higher (`g*vecmax/6 = 2^-9`, the + e4m3 minimum subnormal) is already the normal, exact path. + - An all-zero block is fine: scale byte `0x00`, all-zero data bytes. +- `input` must be **finite**. Non-finite lanes are accepted silently and + are destructive: + - a `+inf`/`-inf` lane drives the block max to infinity, so the scale + saturates at 448 and every other lane in the block is divided by 448 + — small lanes round to zero; the inf lane itself becomes `+6`; + - a NaN lane is excluded from the block max, so the scale stays correct + and the NaN lane saturates to `+6` without disturbing its neighbours. +- `M == 0` (zero rows) is accepted and returns empty `data` and `sf`. +- The caller owns both returned buffers; nothing else writes them. + +## Notes + +- Certified on sm_100 (B200) only, and only for the **NVFP4** mode + (`sf_vec_size=16`, `sf_use_ue8m0=False`) with bfloat16 and float16 + inputs. The kernel behind this op is Blackwell-targeted, and the + installed binary shows it directly: the shipped library carries cubins + of this kernel for sm_80, sm_86, sm_89, sm_90a, sm_100f and sm_120f, + but only the sm_100f and sm_120f ones have a body (32 and 36 registers, + 1024 B shared). The sm_80/86/89 cubins are **4 registers and no shared + memory** and the sm_90a one 4 registers — i.e. the conversion is + compiled away below Blackwell, leaving a kernel that launches, writes + nothing and returns success. A pre-Blackwell arch is therefore a + silent-garbage risk rather than a raise. No receipt is claimed there, + and sm_120 was not run. +- **Not certified** (accepted by the op, no observed-behaviour claims + here): the MXFP4 mode (`sf_vec_size=32`, `sf_use_ue8m0=True`, UE8M0 + scale bytes) beyond its argument validation, and float8_e4m3fn input + (which routes to a different in-kernel conversion). Also uncertified, + though the shape rule covers them: `M > 8192`; float16 and rank-3 + inputs at `K` in {2048, 7168, 18432} (only bf16, rank-2 ran there); + and CUDA-graph capture, which nothing here exercises. +- Against a native-torch reference built from the formula above, the + kernel was **bit-exact on every scale byte** and on every data nibble + outside the near-tie window described in Semantics; every element + inside that window differed by at most one adjacent e2m1 code with the + sign unchanged. Shapes run: + - **linear** layout, bf16: `T x K` for every `T` in + {1, 2, 7, 64, 1023, 1024, 2048} crossed with `K` in + {2560, 3072, 12288} except `2048 x 12288`; plus `T` in + {1, 2, 7, 64, 1023, 1024, 2048, **8192**} crossed with `K` in + {**2048**, **7168**, **18432**}; plus one 16,515,072-element sample + reshaped to `8064 x 2048`, `2304 x 7168` and `896 x 18432`. + - **swizzled** layout, bf16: `1 x 2560`, `7 x 2560`, `129 x 2560`, + `1024 x 2560`, `200 x 3072`, `3 x 112` (`cols = 7`, the only case + exercising column padding), and the full `T` x `K` grid above at + `K` in {2048, 7168, 18432}. Every call listed here was checked + offset-by-offset against the same input's linear call, and its data + bytes asserted equal to the linear call's. + - fp16: `T` in {1, 512, 1024} at `K = 2560`, linear only. A 3-D + `2 x 5 x 2560` bf16 input: linear offset-checked, swizzled checked for + buffer size and data equality only. + The bf16 `K` in {2048, 7168, 18432} grid is the DeepSeek-R1 activation + set — model hidden, dense-MLP intermediate, shared-expert intermediate + — with `T` reaching `max_num_tokens = 8192`. At those widths + `cols = K/16` is 128 / 448 / 1152, all already multiples of 4, so the + swizzled buffer there carries row padding but no column padding. +- Reconstruction bound, verified elementwise on `1024 x 2560` bf16 input + in fp32 and on `1024 x 18432` bf16 input in **float64**: + `|data*sf/g - x| <= max(0.25 * sf/g, |x|/4)`. The two terms are the + e2m1 subnormal half-step carried back through the block scale, and the + half-step of the coarsest e2m1 binade (which also covers the clipping + of a block max when `sf` rounds down). Every operand is exact in + float64, so that form of the check is the mathematical bound; a + consumer reconstructing in **fp32** can land up to half an fp32 ulp + outside it (measured 3.0e-8 at `1024 x 18432`) purely from its own + rounding. +- **One kernel — nothing about rows or width selects a different one.** + Every call of this op in this build launches exactly one CUDA kernel, + `tensorrt_llm::_v1::kernels::quantize_with_block_size`, and it is the same instantiation at 7 rows and at 8192, at + widths 2048 / 7168 / 18432 (multiples of 512) and at 112 (not), in both + scale layouts and for bf16 and fp16 alike — measured with the CUDA + profiler in this entry's test. Consistent with the library's symbol + table, inspected separately, which carries 14 instantiations of that one + kernel template and no TMA variant at all. + A TMA high-throughput variant selected on `M >= 1024 && width % + 512 == 0` **does** exist in FlashInfer's copy of this kernel family, + which ships in this install as readable source — but it is reachable + only through `torch.ops.trtllm.tunable_fp4_quantize`, never through + this op. Treat that vendored source as a *newer* upstream revision, not + as this build: its `invokeFP4Quantization` takes `enable_pdl`, + `use_row_wise_scale` and `inverse_scale` parameters that the installed + symbol does not have, and it accepts a per-token `globalScale` that + this build silently ignores (see *Preconditions*). +- Consumer pairing, read from the TensorRT-LLM call paths in this install + (each consumer's own contract is authoritative; this entry certifies + only what each layout *contains*): + - the NVFP4 dense GEMM family (`torch.ops.trtllm.nvfp4_gemm` and the + backends it dispatches to) is fed the **swizzled** buffer — + `is_sf_swizzled_layout=True`, the op's default; its CUDA-core backend + explicitly un-swizzles before use; + - the trtllm-gen NVFP4 block-scale MoE runner + (`torch.ops.trtllm.fp4_block_scale_moe_runner`) is fed the **linear** + buffer — `is_sf_swizzled_layout=False`, then viewed as + `[num_tokens, cols]` — matching how `mxfp8_quantize` feeds the + trtllm-gen W4A8 runner. + So a target running an NVFP4 dense GEMM and an NVFP4 trtllm-gen MoE + needs two different calls to this op, not one shared result. Both + layouts are certified at `K = 7168`, which is where a DeepSeek-R1-shaped + target needs exactly that pair of calls. +- **How discriminating the certification above is**, measured rather than + assumed — the check tolerates one adjacent e2m1 code inside the + near-tie window and nothing else, so the window is its only blind spot, + and it is 0.03–0.60 % of elements. Driven through the same comparison at + `1024 x 7168`: a single data nibble moved one code *outside* the window + fails it; two codes *inside* the window fails it; a sign flip inside the + window fails it; a single scale byte moved one code fails it (the scale + bytes have no carve-out at all); one code inside the window passes, and + is visible only as one extra disagreement in the count. A swizzled scale + buffer handed over as if it were linear is a different byte string even + when the two have the same length (rows a multiple of 128), so the + layout check is not vacuous either. +- Sibling quantizers with the same scale-layout vocabulary exist + (`torch.ops.trtllm.mxfp8_quantize` for fixed-32-block MXFP8, + `torch.ops.trtllm.fp8_quantize_1x128` for 1x128 block-scaled fp8), as + do fp4 variants that fuse extra work into the same launch + (`torch.ops.trtllm.fp4_quantize_with_residual`, + `torch.ops.trtllm.fp4_batched_quantize`) and a helper that converts an + existing linear scale buffer to the swizzled order + (`torch.ops.trtllm.block_scale_interleave`). + `torch.ops.trtllm.tunable_fp4_quantize` is a different thing again: an + autotuned chooser between this op's kernel and FlashInfer's, so it can + reach the TMA kernel described above. This entry covers the plain + `fp4_quantize` launch only. +- The op allocates both outputs itself: it is functional, with no `out=` + parameter and no in-place mode. diff --git a/tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize.py b/tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize.py new file mode 100644 index 000000000000..058aa6997c70 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize.py @@ -0,0 +1,37 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Block-scaled FP4 activation quantization: packed e2m1 data + one scale byte per +sf_vec_size contiguous elements (NVFP4: 16-element blocks, e4m3 scales).""" + +from typing import Optional + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + + +def fp4_quantize( + input: torch.Tensor, + global_scale: Optional[torch.Tensor], + sf_vec_size: int, + sf_use_ue8m0: bool = False, + is_sf_swizzled_layout: bool = True, +) -> tuple[torch.Tensor, torch.Tensor]: + """Quantize to FP4: returns (packed e2m1 data [..., K/2] uint8, block scales uint8, 1-D). + + One scale per `sf_vec_size` contiguous elements along the last dim. + NVFP4 is `sf_vec_size=16, sf_use_ue8m0=False` (e4m3 scale bytes, scaled by + `global_scale`); `is_sf_swizzled_layout` selects the 128x4 swizzled scale + order (True) or the row-major linear order (False). + """ + # Pure-metadata guard: the kernel loads one scalar from global_scale and + # ignores every element past the first, so a per-token [num_tokens] tensor + # is silently applied as global_scale[0] to every row -- observed on this + # machine to return a valid-looking result, never to raise. + assert global_scale is None or global_scale.numel() == 1, ( + "global_scale must hold exactly one element; extra elements are " + "silently ignored and global_scale[0] is applied to every row" + ) + return torch.ops.trtllm.fp4_quantize( + input, global_scale, sf_vec_size, sf_use_ue8m0, is_sf_swizzled_layout + ) diff --git a/tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize_test.py b/tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize_test.py new file mode 100644 index 000000000000..c184b30853cb --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize_test.py @@ -0,0 +1,730 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the fp4_quantize catalog entry (NVFP4: sf_vec_size=16, e4m3 scales).""" + +import torch +from torch.profiler import ProfilerActivity, profile + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + +from .fp4_quantize import fp4_quantize + +assert torch.cuda.is_available(), "fp4_quantize requires a CUDA device" + +DEV = torch.device("cuda") +VEC = 16 +E4M3_MAX = 448.0 +E2M1_MAX = 6.0 + +# e2m1 code -> value. code = (exponent << 1) | mantissa, bit 3 is the sign. +E2M1_VALUES = torch.tensor( + [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], dtype=torch.float32, device=DEV +) +# Midpoints between consecutive e2m1 magnitudes: a value landing exactly here is +# a rounding tie. `_TIE_UP[i]` is True when rounding *up* at midpoint i yields the +# even code (round-to-nearest-even keeps it) and False when it yields the odd one +# (round-to-nearest-even goes down instead). +E2M1_MIDPOINTS = torch.tensor( + [0.25, 0.75, 1.25, 1.75, 2.5, 3.5, 5.0], dtype=torch.float32, device=DEV +) +_TIE_UP = torch.tensor([False, True, False, True, False, True, False], dtype=torch.bool, device=DEV) + +# The kernel builds its per-block output scale through two `rcp.approx.ftz.f32` +# reciprocals (~2^-23 relative error each), so a value sitting within a few fp32 +# ulps of an e2m1 midpoint can round to either neighbour. Everything further away +# than this window must match the reference bit for bit. The largest deviation +# observed on this machine was 7.9e-8 relative, ~12x inside this window. +TIE_WINDOW = 2.0**-20 + + +def _e2m1_codes(v: torch.Tensor) -> torch.Tensor: + """fp32 -> 4-bit e2m1 codes (bit 3 = sign), round-to-nearest-even, saturating to +-6.""" + a = v.abs().unsqueeze(-1) + gt = (a > E2M1_MIDPOINTS).sum(-1) # strict: ties provisionally round down + tie_up = ((a == E2M1_MIDPOINTS) & _TIE_UP).any(-1) # ties whose even side is up + code = (gt + tie_up).to(torch.uint8) + neg = (v < 0) | ((v == 0) & torch.signbit(v)) + return code | (neg.to(torch.uint8) << 3) + + +def _ref(x: torch.Tensor, gs: float) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Independent native-torch reference: (packed data [M, K/2], sf bytes [M, K/16], scaled). + + Per 16-element block along the last dim (rows = product of the leading dims): + `vecmax = max|x|`; the block scale is `e4m3(gs * vecmax / 6)` (6 = e2m1 max, + e4m3 conversion saturating at 448); the data is the fp32 input times + `gs / scale`, rounded to e2m1. `scaled` is that pre-rounding product, returned + so the caller can identify rounding ties. + """ + k = x.shape[-1] + m = x.numel() // k + xr = x.reshape(m, k) + vecmax = xr.abs().reshape(m, k // VEC, VEC).amax(-1).float() + # torch's fp32 -> e4m3 cast emits NaN above 464 while the kernel's cast + # saturates; clamping first makes the two agree (448 < v <= 464 rounds to + # 448 either way). + sf = (gs * (vecmax / E2M1_MAX)).clamp(max=E4M3_MAX).to(torch.float8_e4m3fn) + out_scale = torch.where(vecmax != 0, gs / sf.float(), torch.zeros_like(vecmax)) + scaled = (xr.float().reshape(m, k // VEC, VEC) * out_scale.unsqueeze(-1)).reshape(m, k) + codes = _e2m1_codes(scaled) + return codes[:, 0::2] | (codes[:, 1::2] << 4), sf.view(torch.uint8), scaled + + +def _unpack(packed: torch.Tensor) -> torch.Tensor: + """[M, K/2] packed bytes -> [M, K] e2m1 codes (element 2i in the low nibble).""" + m, kh = packed.shape + codes = torch.empty(m, kh * 2, dtype=torch.uint8, device=packed.device) + codes[:, 0::2] = packed & 0xF + codes[:, 1::2] = packed >> 4 + return codes + + +def _dequant(packed: torch.Tensor, sf: torch.Tensor, gs: float) -> torch.Tensor: + """data * scale / global_scale -- the value an NVFP4 consumer reconstructs. + + `sf` must be the linear-layout scale buffer of the same call. + """ + codes = _unpack(packed) + v = E2M1_VALUES[(codes & 7).long()] + v = torch.where((codes & 8).bool(), -v, v) + rows, k = codes.shape + scale = sf.view(rows, k // VEC).view(torch.float8_e4m3fn).float() + return v * scale.repeat_interleave(VEC, dim=-1) / gs + + +def _near_tie(scaled: torch.Tensor) -> torch.Tensor: + """Elements whose pre-rounding value sits within TIE_WINDOW (relative) of an e2m1 midpoint.""" + a = scaled.abs().unsqueeze(-1) + return ((a - E2M1_MIDPOINTS).abs() <= TIE_WINDOW * a).any(-1) + + +def _assert_data( + got: torch.Tensor, ref: torch.Tensor, scaled: torch.Tensor +) -> tuple[int, int, int]: + """Bit-exact away from e2m1 rounding ties; one adjacent code at a tie. + + Returns `(near_tie_count, mismatch_count, total)`. The two counts are + different quantities: the first is how many elements sit inside the tie + window at all (the check's blind spot), the second how many of those the + kernel actually resolved the other way. + """ + g = _unpack(got) + r = _unpack(ref) + near_tie = _near_tie(scaled) + torch.testing.assert_close(g[~near_tie], r[~near_tie]) + diff = (g != r) & near_tie + if diff.any(): + # at a tie the kernel may pick either neighbour, never anything else + assert torch.equal(g[diff] >> 3, r[diff] >> 3), "sign flipped at a tie" + delta = (g[diff] & 7).int() - (r[diff] & 7).int() + assert int(delta.abs().max()) == 1, "non-adjacent e2m1 code at a tie" + return int(near_tie.sum()), int(diff.sum()), diff.numel() + + +def _swizzle_index(rows: int, cols: int) -> torch.Tensor: + """[rows, cols] scale coordinates -> flat offsets in the 128x4-swizzled buffer.""" + padded_cols = _pad_up(cols, 4) + r = torch.arange(rows, device=DEV).view(-1, 1) + c = torch.arange(cols, device=DEV).view(1, -1) + return ( + (c % 4) + + (c // 4) * (4 * 128) + + (r % 32) * 16 + + ((r % 128) // 32) * 4 + + (r // 128) * (128 * padded_cols) + ) + + +def _pad_up(x: int, m: int) -> int: + return (x + m - 1) // m * m + + +def _global_scale(x: torch.Tensor) -> tuple[torch.Tensor, float]: + """The canonical NVFP4 activation global scale, 448*6/amax, as a [1] fp32 tensor.""" + gs = (E4M3_MAX * E2M1_MAX / x.abs().max().float()).reshape(1) + return gs, gs.item() + + +def test_linear_layout_bf16() -> None: + """Linear scale order, DeepSeek-V3-Lite widths, decode- through prefill-sized rows. + + Rows straddle 1024 and widths straddle a multiple of 512 -- the two + conditions a TMA high-throughput variant would key on. This build has no + such variant (see test_r1_one_kernel_at_every_width); the coverage is kept + because a future one would. + """ + torch.manual_seed(0) + worst = 0.0 + for t in (1, 2, 7, 64, 1023, 1024, 2048): + for k in (2560, 3072, 12288): + if t * k > 16 * 1024 * 1024: # keeps the reference's [M, K, 7] temporaries bounded + continue + x = torch.randn(t, k, dtype=torch.bfloat16, device=DEV) + gs, gsf = _global_scale(x) + data, sf = fp4_quantize(x, gs, VEC, False, False) + assert data.shape == (t, k // 2) and data.dtype == torch.uint8 + assert data.is_contiguous() + assert sf.shape == (t * k // VEC,) and sf.dtype == torch.uint8 + ref_data, ref_sf, scaled = _ref(x, gsf) + torch.testing.assert_close(sf.view(t, k // VEC), ref_sf) + _, mismatch, total = _assert_data(data, ref_data, scaled) + worst = max(worst, mismatch / total) + # the kernel resolves a near-tie the other way from the exact-fp32 reference + # on at most a fraction of a percent of elements; how many is set by the + # global scale rather than by the shape + assert worst < 0.005, worst + + +def test_swizzled_layout_bf16() -> None: + """Swizzled scales carry the same bytes at the 128x4 offsets; the data is layout-independent.""" + torch.manual_seed(1) + for t, k in ( + (1, 2560), + (7, 2560), + (129, 2560), + (1024, 2560), + (200, 3072), + (3, 112), + ): + x = torch.randn(t, k, dtype=torch.bfloat16, device=DEV) + gs, gsf = _global_scale(x) + cols = k // VEC + data_sw, sf_sw = fp4_quantize(x, gs, VEC, False, True) + data_li, sf_li = fp4_quantize(x, gs, VEC, False, False) + assert sf_sw.shape == (_pad_up(t, 128) * _pad_up(cols, 4),) + assert torch.equal(data_sw, data_li), "data must not depend on the scale layout" + ref_data, ref_sf, scaled = _ref(x, gsf) + torch.testing.assert_close(sf_li.view(t, cols), ref_sf) + _assert_data(data_sw, ref_data, scaled) + idx = _swizzle_index(t, cols).reshape(-1) + torch.testing.assert_close(sf_sw[idx].view(t, cols), ref_sf) + # every offset not addressed by a real (row, col) pair is row/column padding + rest = torch.ones_like(sf_sw, dtype=torch.bool) + rest[idx] = False + assert (sf_sw[rest] == 0).all(), "swizzled padding must be zero" + + +def test_fp16_input() -> None: + torch.manual_seed(2) + for t in (1, 512, 1024): + x = torch.randn(t, 2560, dtype=torch.float16, device=DEV) + gs, gsf = _global_scale(x) + data, sf = fp4_quantize(x, gs, VEC, False, False) + ref_data, ref_sf, scaled = _ref(x, gsf) + assert data.shape == (t, 1280) + torch.testing.assert_close(sf.view(t, 160), ref_sf) + _assert_data(data, ref_data, scaled) + + +def test_3d_input_collapses_leading_dims() -> None: + """data keeps the leading dims; the scale buffer is laid out for their product.""" + torch.manual_seed(3) + x = torch.randn(2, 5, 2560, dtype=torch.bfloat16, device=DEV) + gs, gsf = _global_scale(x) + data, sf = fp4_quantize(x, gs, VEC, False, False) + assert data.shape == (2, 5, 1280) + assert sf.shape == (10 * 160,) + ref_data, ref_sf, scaled = _ref(x, gsf) + torch.testing.assert_close(sf.view(10, 160), ref_sf) + _assert_data(data.reshape(10, 1280), ref_data, scaled) + data_sw, sf_sw = fp4_quantize(x, gs, VEC, False, True) + assert sf_sw.shape == (128 * 160,) + assert torch.equal(data_sw, data) + + +def test_e2m1_ties_round_to_nearest_even() -> None: + """global_scale=1 with a block max of 6 makes the output scale exactly 1, exposing the tie rule.""" + gs = torch.ones(1, dtype=torch.float32, device=DEV) + mids = [0.25, 0.75, 1.25, 1.75, 2.5, 3.5, 5.0] + # the even-code neighbour of each midpoint + expect = [0.0, 1.0, 1.0, 2.0, 2.0, 4.0, 4.0] + x = torch.zeros(2, 16, dtype=torch.bfloat16, device=DEV) + x[0, 0] = E2M1_MAX + x[1, 0] = -E2M1_MAX + for i, mid in enumerate(mids): + x[0, 1 + i] = mid + x[1, 1 + i] = -mid + data, sf = fp4_quantize(x, gs, VEC, False, False) + assert sf.tolist() == [0x38, 0x38], sf.tolist() # e4m3 0x38 == 1.0 + deq = _dequant(data, sf, 1.0) + for i, want in enumerate(expect): + assert deq[0, 1 + i].item() == want, (mids[i], deq[0, 1 + i].item()) + assert deq[1, 1 + i].item() == -want, (mids[i], deq[1, 1 + i].item()) + assert deq[0, 0].item() == E2M1_MAX and deq[1, 0].item() == -E2M1_MAX + + +def test_dequant_error_bound() -> None: + """data * sf / global_scale reconstructs the input within the nvfp4 rounding bound.""" + torch.manual_seed(4) + x = torch.randn(1024, 2560, dtype=torch.bfloat16, device=DEV) * 3.0 + gs, gsf = _global_scale(x) + data, sf = fp4_quantize(x, gs, VEC, False, False) + deq = _dequant(data, sf, gsf) + scale = sf.view(1024, 160).view(torch.float8_e4m3fn).float() + scale = scale.repeat_interleave(VEC, dim=-1) + # In the block-scaled domain s = x*gs/sf the e2m1 grid has spacing 0.5 below + # 2, 1.0 on [2,4), 2.0 on [4,6] and saturates at 6, so the rounding error is + # at most max(0.25, |s|/4). Dividing back by gs/sf gives the bound below; + # |s|/4 / (gs/sf) is exactly |x|/4. + bound = torch.maximum(0.25 * scale / gsf, x.float().abs() / 4) + err = (deq - x.float()).abs() + assert (err <= bound).all(), (err - bound).max().item() + # hard gate: the quantization is the reference, not merely inside the bound + ref_data, ref_sf, scaled = _ref(x, gsf) + torch.testing.assert_close(sf.view(1024, 160), ref_sf) + _, mismatch, total = _assert_data(data, ref_data, scaled) + assert mismatch / total < 0.005, mismatch / total + # a block whose scale rounds down lets its max element clip to 6 -- that is + # the format, and it is inside the bound above + assert (scaled.abs() > E2M1_MAX).any() + + +def test_scale_saturation() -> None: + """A global_scale calibrated on a smaller amax saturates the scale and clips every lane.""" + x = torch.zeros(1, 16, dtype=torch.bfloat16, device=DEV) + x[0, :4] = torch.tensor([1.0, -2.0, 0.5, 3.0]) + gs = torch.tensor([E4M3_MAX * E2M1_MAX / 0.001], dtype=torch.float32, device=DEV) + data, sf = fp4_quantize(x, gs, VEC, False, False) + assert sf.tolist() == [0x7E] # e4m3 0x7E == 448, the finite max + # every nonzero lane clips to +-6 regardless of its magnitude + codes = _unpack(data) + assert codes[0, :4].tolist() == [7, 15, 7, 7] + assert (codes[0, 4:] == 0).all() + deq = _dequant(data, sf, gs.item()) + clipped = torch.full((4,), E2M1_MAX * E4M3_MAX / gs.item(), device=DEV) + clipped[1] = -clipped[1] + torch.testing.assert_close(deq[0, :4], clipped) + assert (deq[0, 4:] == 0.0).all() + + +def test_scale_underflow() -> None: + """A global_scale so small that the block scale rounds to 0 zeroes the whole block.""" + x = torch.zeros(1, 16, dtype=torch.bfloat16, device=DEV) + x[0, 0] = 1.0 + x[0, 1] = -0.5 + # e4m3's smallest subnormal is 2^-9; gs*vecmax/6 = 2^-11 rounds to zero + gs = torch.tensor([E2M1_MAX * 2.0**-11], dtype=torch.float32, device=DEV) + data, sf = fp4_quantize(x, gs, VEC, False, False) + assert sf.tolist() == [0x00] + # the reciprocal of the zero scale is +inf: every lane saturates to +6, + # including the exact zeros (0 * inf = NaN, which the e2m1 cast saturates) + codes = _unpack(data) + assert codes[0, 0].item() == 7 and codes[0, 1].item() == 15 + assert (codes[0, 2:] == 7).all() + # with a zero scale the consumer reads the block back as all zeros + assert (_dequant(data, sf, gs.item()) == 0.0).all() + # one binade higher the scale is the e4m3 minimum subnormal and the block is exact + gs2 = torch.tensor([E2M1_MAX * 2.0**-9], dtype=torch.float32, device=DEV) + data2, sf2 = fp4_quantize(x, gs2, VEC, False, False) + assert sf2.tolist() == [0x01] + deq2 = _dequant(data2, sf2, gs2.item()) + assert deq2[0, 0].item() == 1.0 and deq2[0, 1].item() == -0.5 + + +def test_degenerate_global_scale() -> None: + """A zero or negative global_scale is accepted and silently unusable.""" + x = torch.zeros(1, 16, dtype=torch.bfloat16, device=DEV) + x[0, :4] = torch.tensor([1.0, -2.0, 0.5, 3.0]) + # g = 0: the block scale is 0 and every lane saturates, so a consumer + # dividing by g reads NaN + data, sf = fp4_quantize(x, torch.zeros(1, device=DEV), VEC, False, False) + assert sf.tolist() == [0x00] + assert (_unpack(data)[0] == 7).all() + # g < 0: the scale byte carries e4m3's sign bit, which every consumer reads + # as an unsigned UE4M3 magnitude; the data bytes are those of |g| + data_n, sf_n = fp4_quantize(x, -torch.ones(1, device=DEV), VEC, False, False) + data_p, sf_p = fp4_quantize(x, torch.ones(1, device=DEV), VEC, False, False) + assert sf_n.tolist() == [0xB0] and sf_p.tolist() == [0x30] # -0.5 and +0.5 + assert torch.equal(data_n, data_p) + + +def test_reciprocal_global_scale_is_silently_destructive() -> None: + """A modelopt checkpoint stores amax/(448*6); passing it unreciprocated is accepted. + + The mistake costs a factor of `(448*6/amax)^2`, which for normed activations + puts almost every block under the e4m3 scale floor -- but not all of them, so + the result is neither an error nor an obviously dead tensor. + """ + torch.manual_seed(12) + x = torch.randn(64, 7168, dtype=torch.bfloat16, device=DEV) + amax = x.abs().max().float() + stored = (amax / (E4M3_MAX * E2M1_MAX)).reshape(1) # what is on disk + data, sf = fp4_quantize(x, stored, VEC, False, False) + zero_scales = (sf == 0).float().mean().item() + assert zero_scales > 0.9, zero_scales + assert zero_scales < 1.0, zero_scales + deq = _dequant(data, sf, stored.item()) + # the blocks that survive saturate: reconstructed magnitudes above the true + # amax, next to blocks that read back as exact zeros + assert deq.abs().max().item() > amax.item() + assert (deq == 0).any() + + +def test_zero_block_and_signed_zero() -> None: + """An all-zero block emits a zero scale and zero data; -0.0 keeps its sign bit.""" + z = torch.zeros(1, 32, dtype=torch.bfloat16, device=DEV) + gs = torch.tensor([100.0], dtype=torch.float32, device=DEV) + data, sf = fp4_quantize(z, gs, VEC, False, False) + assert sf.tolist() == [0x00, 0x00] + assert (data == 0).all() + y = torch.zeros(1, 16, dtype=torch.bfloat16, device=DEV) + y[0, 0] = 6.0 + y[0, 1] = -0.0 + y[0, 2] = 0.0 + codes = _unpack(fp4_quantize(y, torch.ones(1, device=DEV), VEC, False, False)[0]) + assert codes[0, 1].item() == 8, "negative zero must keep its sign bit" + assert codes[0, 2].item() == 0 + + +def test_nonfinite_inputs() -> None: + """Non-finite lanes are accepted silently; both outcomes are destructive.""" + gs = torch.ones(1, dtype=torch.float32, device=DEV) + # +inf drives the block max to inf, so the scale saturates at 448 and every + # finite lane is scaled by 1/448 -- small lanes round to zero. + x = torch.zeros(1, 16, dtype=torch.bfloat16, device=DEV) + x[0, 0] = float("inf") + x[0, 1] = 1.0 + data, sf = fp4_quantize(x, gs, VEC, False, False) + assert sf.tolist() == [0x7E] + codes = _unpack(data) + assert codes[0, 0].item() == 7, "inf saturates to +6" + assert (codes[0, 1:] == 0).all() + # NaN is excluded from the block max, so the scale comes from the finite + # lanes; the NaN lane itself saturates to +6 and stays confined. + y = torch.zeros(1, 16, dtype=torch.bfloat16, device=DEV) + y[0, 0] = float("nan") + y[0, 1] = 1.0 + data, sf = fp4_quantize(y, gs, VEC, False, False) + assert sf.tolist() == [0x23] # e4m3 nearest to 1/6 + codes = _unpack(data) + assert codes[0, 0].item() == 7 + assert codes[0, 1].item() == 7 # 1.0 * (1/0.171875) = 5.82 -> 6 + assert (codes[0, 2:] == 0).all() + + +def test_zero_rows() -> None: + """An empty batch is accepted, not rejected: both outputs come back empty.""" + gs = torch.ones(1, dtype=torch.float32, device=DEV) + x = torch.randn(0, 2560, dtype=torch.bfloat16, device=DEV) + for swizzled in (False, True): + data, sf = fp4_quantize(x, gs, VEC, False, swizzled) + assert data.shape == (0, 1280) and data.dtype == torch.uint8 + assert sf.shape == (0,) and sf.dtype == torch.uint8 + + +def test_input_not_mutated_and_deterministic() -> None: + torch.manual_seed(5) + x = torch.randn(64, 2560, dtype=torch.bfloat16, device=DEV) + gs, _ = _global_scale(x) + before = x.clone() + a = fp4_quantize(x, gs, VEC, False, True) + b = fp4_quantize(x, gs, VEC, False, True) + assert torch.equal(x, before) + assert torch.equal(a[0], b[0]) and torch.equal(a[1], b[1]) + + +def test_multi_element_global_scale() -> None: + """The op silently uses global_scale[0] for every row; the wrapper rejects that shape.""" + torch.manual_seed(6) + x = torch.randn(8, 256, dtype=torch.bfloat16, device=DEV) + per_row = torch.tensor( + [1.0, 2.0, 4.0, 8.0, 16.0, 32.0, 64.0, 128.0], dtype=torch.float32, device=DEV + ) + # raw op: accepted, and identical to broadcasting the first element + got = torch.ops.trtllm.fp4_quantize(x, per_row, VEC, False, False)[1] + want = torch.ops.trtllm.fp4_quantize(x, per_row[:1], VEC, False, False)[1] + assert torch.equal(got, want), "multi-element global_scale is not applied per row" + try: + fp4_quantize(x, per_row, VEC, False, False) + except AssertionError: + pass + else: + raise AssertionError("wrapper must reject a multi-element global_scale") + + +# DeepSeek-R1 activation widths: model hidden (dense-MLP / shared-expert / MoE +# input), dense-MLP intermediate, shared-expert intermediate. All three are +# multiples of 512, the width half of the condition a TMA high-throughput +# variant would key on if this build had one. +R1_WIDTHS = (7168, 18432, 2048) +# The receipt's existing row range plus max_num_tokens = 8192, which is also the +# MoE path's chunk bound. +R1_ROWS = (1, 2, 7, 64, 1023, 1024, 2048, 8192) +# Row-chunk budget for the reference: its [rows, K, 7] midpoint comparison is +# the memory ceiling, and the reference is row-independent so chunking is exact. +REF_CHUNK_ELEMS = 16 * 1024 * 1024 + + +def _assert_call_matches_ref( + x: torch.Tensor, data: torch.Tensor, sf_linear: torch.Tensor, gsf: float +) -> tuple[int, int, int]: + """Row-chunked bit-exact check of one linear-layout call. + + Returns `(near_tie_count, mismatch_count, total)` summed over the chunks. + """ + rows, k = x.shape + cols = k // VEC + step = max(1, REF_CHUNK_ELEMS // k) + n_tie = n_mis = n_all = 0 + for lo in range(0, rows, step): + hi = min(rows, lo + step) + ref_data, ref_sf, scaled = _ref(x[lo:hi], gsf) + torch.testing.assert_close(sf_linear.view(rows, cols)[lo:hi], ref_sf) + tie, mismatch, total = _assert_data(data[lo:hi], ref_data, scaled) + n_tie += tie + n_mis += mismatch + n_all += total + return n_tie, n_mis, n_all + + +def test_r1_widths_both_layouts_bf16() -> None: + """DeepSeek-R1 activation widths in both scale layouts, decode- through 8192-row prefill. + + 7168 is the model hidden, taken swizzled by the dense NVFP4 GEMM and linear + by the trtllm-gen MoE runner; 18432 is the dense-MLP intermediate and 2048 + the shared-expert intermediate, both taken swizzled. Every width is a + multiple of 16 with `K/16` already a multiple of 4, so the swizzled buffer + carries row padding but no column padding. + """ + worst_tie = worst_mis = 0.0 + seed = 100 + for k in R1_WIDTHS: + cols = k // VEC + for t in R1_ROWS: + torch.manual_seed(seed) + seed += 1 + x = torch.randn(t, k, dtype=torch.bfloat16, device=DEV) + gs, gsf = _global_scale(x) + data_li, sf_li = fp4_quantize(x, gs, VEC, False, False) + data_sw, sf_sw = fp4_quantize(x, gs, VEC, False, True) + assert data_li.shape == (t, k // 2) and data_li.dtype == torch.uint8 + assert data_li.is_contiguous() and sf_li.is_contiguous() + assert sf_li.shape == (t * cols,) and sf_li.dtype == torch.uint8 + assert sf_sw.shape == (_pad_up(t, 128) * _pad_up(cols, 4),) + # the same bytes out of two separate calls: layout-independent data + assert torch.equal(data_sw, data_li), "data must not depend on the layout" + tie, mismatch, total = _assert_call_matches_ref(x, data_li, sf_li, gsf) + worst_tie = max(worst_tie, tie / total) + worst_mis = max(worst_mis, mismatch / total) + idx = _swizzle_index(t, cols).reshape(-1) + torch.testing.assert_close(sf_sw[idx], sf_li) + rest = torch.ones_like(sf_sw, dtype=torch.bool) + rest[idx] = False + assert (sf_sw[rest] == 0).all(), "swizzled padding must be zero" + del x, data_li, sf_li, data_sw, sf_sw, idx, rest + # Ceilings on the check's blind spot, not correctness gates -- correctness is + # the bit-exactness above. Both quantities are set by the global scale rather + # than by the shape (see the next test), so they are ceilings with room, not + # fits: the worst of these 24 shapes was 0.60 % near-tie / 0.14 % mismatch. + assert worst_mis < 0.005, worst_mis + assert worst_tie < 0.02, worst_tie + + +def test_r1_one_kernel_at_every_width() -> None: + """One call launches exactly one CUDA kernel, and it is the same one at every shape. + + Rows (7 vs 8192) and width (multiple of 512 or not) do not select a + different kernel in this build: there is no TMA high-throughput variant on + this op's path to switch to. + """ + torch.manual_seed(8) + for t, k, dtype in ( + (7, 7168, torch.bfloat16), + (1024, 7168, torch.bfloat16), + (8192, 7168, torch.bfloat16), + (8192, 18432, torch.bfloat16), + (8192, 2048, torch.bfloat16), + (1024, 112, torch.bfloat16), + (8192, 7168, torch.float16), + ): + x = torch.randn(t, k, dtype=dtype, device=DEV) + gs, _ = _global_scale(x) + for swizzled in (False, True): + fp4_quantize(x, gs, VEC, False, swizzled) # warm up the launch path + torch.cuda.synchronize() + with profile(activities=[ProfilerActivity.CUDA]) as prof: + fp4_quantize(x, gs, VEC, False, swizzled) + torch.cuda.synchronize() + launched = [ + e.name + for e in prof.events() + if e.device_type == torch.autograd.DeviceType.CUDA and e.self_device_time_total > 0 + ] + assert len(launched) == 1, (t, k, dtype, swizzled, launched) + assert "quantize_with_block_size" in launched[0], launched[0] + assert "tma" not in launched[0].lower(), launched[0] + del x + + +def test_r1_reference_check_is_discriminating() -> None: + """The tie carve-out tolerates exactly one adjacent code at a near-tie, nothing else. + + Runs the wrong variants through the same comparison the R1 sweep uses, so + the sweep's clean result is a measurement rather than a blind spot. + """ + torch.manual_seed(9) + x = torch.randn(1024, 7168, dtype=torch.bfloat16, device=DEV) + gs, gsf = _global_scale(x) + data, sf = fp4_quantize(x, gs, VEC, False, False) + ref_data, ref_sf, scaled = _ref(x, gsf) + near = _near_tie(scaled) + # the correct result passes, and the carve-out is a small fraction of the tensor + n_tie, _, total = _assert_data(data, ref_data, scaled) + assert 0 < n_tie / total < 0.02, n_tie / total + codes = _unpack(data) + # perturb only elements the kernel already got right, and only where the + # magnitude code leaves room to move two steps up + agree = (codes == _unpack(ref_data)) & ((codes & 7) <= 5) + inside = (near & agree).nonzero() + outside = (~near & agree & ((codes & 7) > 0)).nonzero() + assert len(inside) > 0 and len(outside) > 0 + + def bumped(row: int, col: int, delta: int, flip_sign: bool = False) -> torch.Tensor: + c = codes.clone() + code = int(c[row, col]) + mag = (code & 7) + delta + assert 0 <= mag <= 7, mag + c[row, col] = mag | ((code & 8) ^ (8 if flip_sign else 0)) + return c[:, 0::2] | (c[:, 1::2] << 4) + + def raises(fn) -> bool: + try: + fn() + except AssertionError: + return True + return False + + ro, co = (int(v) for v in outside[len(outside) // 2]) + ri, ci = (int(v) for v in inside[len(inside) // 2]) + assert raises(lambda: _assert_data(bumped(ro, co, 1), ref_data, scaled)), ( + "one code away from a tie must be caught" + ) + assert raises(lambda: _assert_data(bumped(ri, ci, 2), ref_data, scaled)), ( + "two codes at a tie must be caught" + ) + assert raises(lambda: _assert_data(bumped(ri, ci, 0, flip_sign=True), ref_data, scaled)), ( + "a sign flip at a tie must be caught" + ) + # the documented blind spot: one adjacent code at a near-tie is accepted, + # and shows up only as one extra mismatch + _, mis_before, _ = _assert_data(data, ref_data, scaled) + _, mis_after, _ = _assert_data(bumped(ri, ci, 1), ref_data, scaled) + assert mis_after == mis_before + 1, (mis_before, mis_after) + # and the scale bytes carry no such carve-out: one code is caught + bad_sf = sf.clone() + bad_sf[0] = (int(bad_sf[0]) + 1) & 0xFF + assert raises(lambda: torch.testing.assert_close(bad_sf.view(1024, 448), ref_sf)), ( + "one scale code must be caught" + ) + # a swizzled buffer handed over as if it were linear is a different byte + # string even when the two have the same length (rows a multiple of 128) + _, sf_sw = fp4_quantize(x, gs, VEC, False, True) + assert sf_sw.shape == sf.shape and not torch.equal(sf_sw, sf) + + +def test_r1_ties_track_the_global_scale_not_the_shape() -> None: + """The near-tie population is a property of the global scale, not of M or K. + + One sample is reshaped to all three R1 widths. Scale blocks are 16 + *contiguous* elements, so the row width never moves a block boundary and the + only thing a reshape can change is `amax`; holding that fixed, all three + widths must agree on the exact count. Nudging `global_scale` off the + canonical `448*6/amax` then removes the population altogether, because a + near-tie needs `sf` to come out exactly `g*vecmax/6` -- which makes + `out_scale` exactly `6/vecmax` and puts `6*x/vecmax` on an e2m1 midpoint. + """ + # 129024 = lcm(7168, 18432, 2048), so one buffer reshapes to all three + n = 129024 * 128 + torch.manual_seed(11) + flat = torch.randn(n, dtype=torch.bfloat16, device=DEV) + gs, gsf = _global_scale(flat) + counts = [] + for k in R1_WIDTHS: + x = flat.view(n // k, k) + data, sf = fp4_quantize(x, gs, VEC, False, False) + tie, mismatch, total = _assert_call_matches_ref(x, data, sf, gsf) + assert total == n + counts.append((tie, mismatch)) + del data, sf + assert counts[0] == counts[1] == counts[2], counts + assert counts[0][0] > 0, counts + # a global scale 0.1 % off the canonical one: no block keeps an exact sf, so + # the kernel is bit-exact against the reference with no carve-out at all + x = flat.view(n // 7168, 7168) + off = gs * 1.001 + data, sf = fp4_quantize(x, off, VEC, False, False) + tie, mismatch, _ = _assert_call_matches_ref(x, data, sf, off.item()) + assert (tie, mismatch) == (0, 0), (tie, mismatch) + # a power-of-two multiple instead scales sf exactly and leaves out_scale + # unchanged, so the population survives: it is g's mantissa that matters + pow2 = gs * 2.0 + data, sf = fp4_quantize(x, pow2, VEC, False, False) + tie, _, _ = _assert_call_matches_ref(x, data, sf, pow2.item()) + assert tie > 0, tie + + +def test_r1_dequant_error_bound() -> None: + """The nvfp4 reconstruction bound holds at the widest R1 width. + + Evaluated in float64: every operand (e2m1 value, e4m3 scale, bf16 input, + fp32 global scale) is exact there, so the comparison is the mathematical + bound rather than a tolerance. In fp32 -- what a consumer actually computes + -- an element sitting exactly on the bound lands up to half an fp32 ulp + outside it; measured 3.0e-8 on this shape. + """ + torch.manual_seed(10) + x = torch.randn(1024, 18432, dtype=torch.bfloat16, device=DEV) * 3.0 + gs, gsf = _global_scale(x) + data, sf = fp4_quantize(x, gs, VEC, False, False) + codes = _unpack(data) + v = E2M1_VALUES[(codes & 7).long()].double() + v = torch.where((codes & 8).bool(), -v, v) + scale = sf.view(1024, 1152).view(torch.float8_e4m3fn).double() + scale = scale.repeat_interleave(VEC, dim=-1) + # same bound as test_dequant_error_bound: the e2m1 subnormal half-step + # carried back through the block scale, or the half-step of the coarsest + # e2m1 binade (which also covers a block max clipping to 6) + bound = torch.maximum(0.25 * scale / gsf, x.double().abs() / 4) + err = (v * scale / gsf - x.double()).abs() + assert (err <= bound).all(), (err - bound).max().item() + + +def test_rejects_unsupported() -> None: + """Domains the contract declares unsupported must be rejected, not silently wrong.""" + x = torch.randn(8, 256, dtype=torch.bfloat16, device=DEV) + gs = torch.ones(1, dtype=torch.float32, device=DEV) + cases = { + "no global_scale on the nvfp4 path": lambda: fp4_quantize(x, None, 16, False, True), + "sf_vec_size 32 without ue8m0": lambda: fp4_quantize(x, gs, 32, False, True), + "sf_vec_size 16 with ue8m0": lambda: fp4_quantize(x, gs, 16, True, True), + "sf_vec_size 8": lambda: fp4_quantize(x, gs, 8, False, True), + "sf_vec_size 0": lambda: fp4_quantize(x, gs, 0, False, True), + "k not a multiple of 16": lambda: fp4_quantize( + torch.randn(4, 24, dtype=torch.bfloat16, device=DEV), gs, 16, False, True + ), + "fp32 input": lambda: fp4_quantize(x.float(), gs, 16, False, True), + "1-D input": lambda: fp4_quantize(x.reshape(-1), gs, 16, False, True), + "fp16 global_scale": lambda: fp4_quantize(x, gs.half(), 16, False, True), + "cpu global_scale": lambda: fp4_quantize(x, gs.cpu(), 16, False, True), + "cpu input": lambda: fp4_quantize(x.cpu(), gs.cpu(), 16, False, True), + "non-contiguous column slice": lambda: fp4_quantize( + torch.randn(8, 512, dtype=torch.bfloat16, device=DEV)[:, :256], + gs, + 16, + False, + True, + ), + "non-contiguous transpose": lambda: fp4_quantize( + torch.randn(256, 8, dtype=torch.bfloat16, device=DEV).t(), + gs, + 16, + False, + True, + ), + } + for name, fn in cases.items(): + try: + fn() + except (RuntimeError, NotImplementedError): + continue + raise AssertionError(f"{name}: expected a raise, got a result") diff --git a/tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize.md b/tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize.md new file mode 100644 index 000000000000..339212adb107 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize.md @@ -0,0 +1,170 @@ +--- +receipts: + sm_100: {status: passed, trtllm: 1.3.0rc21} + sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 13} +--- + +# mxfp8_quantize + +**Wraps** `torch.ops.trtllm.mxfp8_quantize` (one call). + +## Semantics + +Dynamic MXFP8 (OCP microscaling FP8) quantization of an activation tensor: +one call converts bf16/fp16 values to e4m3 elements plus one UE8M0 +(power-of-two) scale per **32 contiguous elements along the last dim**. +Nothing else is fused in — no normalization, no activation, no transpose. +The caller owns producing the input (e.g. bf16 hidden states) and owns +both outputs afterwards. + +Let `M` = product of all leading dims, `K` = last dim, and +`padded_k = ceil(K / alignment) * alignment`. For row `m` and block +`b` of 32 elements: + +``` +amax = max |x[m, 32b : 32b+32]| # fp32; exact for bf16/fp16 inputs +scale = smallest power of two >= amax / 448 # 448 = e4m3 max finite +sf[m, b] = log2(scale) + 127 # UE8M0 byte, so scale == 2^(sf-127) +data[m,i] = e4m3_round_to_nearest_even( x[m, i] / scale ) for i in the block +``` + +The scale is rounded **up** (toward +inf) to the next power of two, so +the block max always lands at or below e4m3's 448 and no element +saturates. Dequantization is exactly `data.float() * 2^(sf - 127)`; +per element the reconstruction error is at most `2^-4` relative (e4m3's +3-bit mantissa) plus `2^-10 * scale` absolute (e4m3 subnormals). + +Padding is part of the op, not the caller's job: + +- **Column padding** (`padded_k > K`, i.e. `alignment` wider than `K`): + columns `K .. padded_k` of `data` are written as the zero byte `0x00` + and their scale bytes are `0x00`. The valid columns are bit-identical + to what an unpadded call produces. +- **Row padding** (swizzled layout only): the scale buffer is sized for + `pad_up(M, 128)` rows; the bytes belonging to rows `M .. pad_up(M,128)` + are `0x00`. The `data` tensor is never row-padded. + +`data` is identical byte-for-byte under both `swizzled_layout` values — +only the scale buffer's size and element order change: + +- `swizzled_layout=False` (**linear**): `M * cols` bytes, row-major, where + `cols = padded_k / 32`. Byte `b` of row `m` is at flat offset + `m * cols + b`; `sf.view(M, cols)` is exactly the per-block scale matrix. +- `swizzled_layout=True` (**128x4 swizzled**): `pad_up(M,128) * pad_up(cols,4)` + bytes. The scale of `(m, b)` sits at flat offset + + ``` + (b % 4) + + (b // 4) * 512 # 4 cols x 128 rows per column group + + (m % 32) * 16 + + ((m % 128) // 32) * 4 + + (m // 128) * 128 * pad_up(cols, 4) + ``` + + Every offset not addressed by a real `(m, b)` pair is `0x00` padding. + +## Signature + +```python +def mxfp8_quantize( + input: torch.Tensor, + swizzled_layout: bool = True, + alignment: int = 32, +) -> tuple[torch.Tensor, torch.Tensor] +``` + +| Argument | Shape | Dtype | Layout / value | Device | +|---|---|---|---|---| +| `input` | `[..., K]`, rank >= 2 | bfloat16 or float16 | contiguous | CUDA | +| `swizzled_layout` | scalar | Python bool | `False` = linear scale order, `True` = 128x4 swizzled | — | +| `alignment` | scalar | Python int | positive multiple of 32; pads `K` up to `ceil(K/alignment)*alignment` | — | +| returns `data` | `[..., padded_k]` (leading dims preserved) | float8_e4m3fn | contiguous, freshly allocated | same as `input` | +| returns `sf` | 1-D: `[M*cols]` (linear) or `[pad_up(M,128)*pad_up(cols,4)]` (swizzled) | uint8 | contiguous, freshly allocated | same as `input` | + +`M` = product of `input`'s leading dims, `cols = padded_k / 32`. `input` +is not modified. Leading dims are collapsed for the scale buffer in both +layouts (a `[B, S, K]` input yields a scale buffer laid out for `B*S` +rows), while `data` keeps `input`'s leading dims. + +The op's own schema names the second parameter `swizzedLayout` (upstream +spelling); it is positional here. + +## Metadata consumed + +None. Stateless — no runtime, attention metadata, or workspace. + +## Preconditions + +- `input` is on CUDA, contiguous, rank >= 2, dtype bfloat16 or float16. + Violations raise: a CPU tensor, a non-contiguous view (column slice or + transpose), a 1-D tensor and an fp32 tensor were each observed to raise + on this machine rather than compute something wrong. +- `K % 32 == 0` (the block size). Raises otherwise. +- `alignment % 32 == 0` and `alignment > 0`. A non-multiple raises; + `alignment = 0` is **not** an exception — it kills the process with + SIGFPE (integer division by zero in the padding computation). There is + no upper bound: `alignment` may exceed `K` (`K=2880, alignment=512` + gives `padded_k=3072`) or equal the block size (`alignment=32`, no pad). + A **negative** multiple of 32 is the quiet one: it passes the + `% 32 == 0` guard, raises nothing, and **silently truncates** — at + `K = 2880` the values `-32 / -64 / -512` return `padded_k` + `2816 / 2752 / 2048`, matching `((K + a - 1) / a) * a` exactly, i.e. the + round-up above turning into a round-*down* under truncating division. + The columns + that survive are bit-exact against the reference and the scale tensor is + self-consistent, so the result looks entirely well-formed; the tail of + every row is simply gone. Check the sign before the modulus. +- `input` must be **finite**. Non-finite lanes are accepted silently and + are destructive: + - a `+inf`/`-inf` lane drives the block scale to the E8M0 finite max + (byte 254 = 2^127); the inf lane becomes NaN and **every other lane in + that 32-element block is zeroed**; + - a NaN lane is excluded from the block max, so the scale stays correct + and the NaN stays confined to its own lane. +- Every 32-element block must have `amax == 0` or `amax > 448 * 2^-127` + (~2.63e-36). An all-zero block is fine (scale byte `0x00`, all-zero + data). A block whose max is nonzero but at or below that threshold + silently produces garbage: scale byte `0x00`, every nonzero lane + saturated to +-448, and **every exact-zero lane turned into NaN**. Real + bf16 hidden states never reach this range; a block of scaled-down + denormals does. +- `M == 0` (zero rows) is accepted and returns empty `data` and `sf`. +- The caller owns both returned buffers; nothing else writes them. + +## Notes + +- Certified on sm_100 (B200) only. The kernel behind this op is + Blackwell-targeted: the same TensorRT-LLM quantization kernel is + vendored as readable source elsewhere in this install (flashinfer's + bundled `nv_internal` copy) and compiles its body only under + `__CUDA_ARCH__ >= 1000`, the empty branch returning without writing the + output buffers. A pre-Blackwell arch is therefore a silent-garbage risk + rather than a raise. No receipt is claimed there. +- Against a native-torch reference built from the formula above, the + kernel was **bit-exact** on every element and every scale byte across + `T x 2880` and `T x 3072` bf16 inputs for `T` in {1, 2, 7, 64, 1024, + 8192} and `alignment` in {32, 128, 512}, plus fp16 inputs and a 3-D + input. The scale math runs through approximate reciprocals (the + flush-to-zero surface above is their footprint), yet the result did not + shift even when `amax` is exactly `448 * 2^k` — the only value where a + 1-ulp difference could flip the round-up (verified for `k = -2, 0, 1, 3`). +- Scale-byte dynamic range from bf16/fp16 inputs: the largest finite bf16 + magnitude (3.39e38) yields byte 247, so the E8M0 finite max (254) is + reachable only from a non-finite input. +- Consumer pairing. The trtllm-gen W4A8 MXFP4xMXFP8 MoE runner + `torch.ops.trtllm.mxe4m3_mxe2m1_block_scale_moe_runner` takes this op's + two outputs as its `hidden_states` / `hidden_states_scale` and expects + the **linear** scale buffer (`swizzled_layout=False`), sized exactly + `M * padded_k/32` bytes, with `alignment` equal to that kernel's + input-hidden alignment — 512 for the trtllm-gen MXFP4 weight path, which + is what pads gpt-oss's 2880 hidden to 3072 inside this call (so the + caller does not pre-pad). The swizzled layout is what the CUTLASS-side + mxfp8 GEMM/MoE consumers take instead. Those pairings are fixed by the + consuming ops' own contracts; this entry certifies only what each layout + contains. +- Sibling quantizers with the same scale-layout vocabulary exist + (`torch.ops.trtllm.fp4_quantize` for NVFP4/MXFP4 with a 16- or + 32-element vector, `torch.ops.trtllm.fp8_quantize_1x128` for 1x128 + block-scaled fp8); this entry covers the fixed-32-block mxfp8 path only. +- The op allocates both outputs itself: it is functional, with no `out=` + parameter and no in-place mode. diff --git a/tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize.py b/tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize.py new file mode 100644 index 000000000000..511d3051fef8 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize.py @@ -0,0 +1,21 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Dynamic MXFP8 quantization: bf16/fp16 -> e4m3 data + per-32-element UE8M0 block scales.""" + +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* + + +def mxfp8_quantize( + input: torch.Tensor, + swizzled_layout: bool = True, + alignment: int = 32, +) -> tuple[torch.Tensor, torch.Tensor]: + """Quantize to MXFP8: returns (e4m3 data [..., pad_up(K, alignment)], uint8 UE8M0 scales, 1D). + + One scale per 32 contiguous elements along the last dim; `swizzled_layout` + selects the 128x4 swizzled scale order (True) or the row-major linear + order (False). + """ + return torch.ops.trtllm.mxfp8_quantize(input, swizzled_layout, alignment) diff --git a/tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize_test.py b/tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize_test.py new file mode 100644 index 000000000000..b43301ef5490 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize_test.py @@ -0,0 +1,289 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GPU test for the mxfp8_quantize catalog entry.""" + +import torch + +from .mxfp8_quantize import mxfp8_quantize + +assert torch.cuda.is_available(), "mxfp8_quantize requires a CUDA device" + +E4M3_MAX = 448.0 +BLOCK = 32 + + +def _ref(x: torch.Tensor, padded_k: int) -> tuple[torch.Tensor, torch.Tensor]: + """Independent native-torch reference: (e4m3 data [M, padded_k], UE8M0 bytes [M, padded_k/32]). + + Per 32-element block along the last dim (rows = product of the leading dims): + amax = max|x|; the block scale is the smallest power of two >= amax/448, + encoded as the E8M0 byte `exp2_of_scale + 127`; the data is the fp32 input + divided by that scale, rounded to e4m3. Columns beyond the input's K (up to + `padded_k`) are treated as zero. + + Exact for every block with amax > 448 * 2^-127; below that the kernel takes + a flush-to-zero path that `test_zero_and_denormal_blocks` pins directly. + """ + k = x.shape[-1] + m = x.numel() // k + xp = torch.zeros(m, padded_k, dtype=torch.float32, device=x.device) + xp[:, :k] = x.reshape(m, k).float() + blk = xp.view(m, padded_k // BLOCK, BLOCK) + amax = blk.abs().amax(dim=-1) + # E8M0 of amax/448, rounded toward +inf: take the fp32 biased exponent and + # bump it whenever any mantissa bit is set (i.e. the value is not already a + # power of two). amax is exact in fp32 for bf16/fp16 inputs, so this is the + # exact round-up. + u = (amax / E4M3_MAX).view(torch.int32) + byte = (((u >> 23) & 0xFF) + ((u & 0x7FFFFF) > 0).to(torch.int32)).clamp(0, 254) + scale = torch.exp2(byte.to(torch.float32) - 127.0).unsqueeze(-1) + data = (blk / scale).view(m, padded_k).to(torch.float8_e4m3fn) + return data, byte.to(torch.uint8) + + +def _swizzle_index(rows: int, cols: int, device: torch.device) -> torch.Tensor: + """[rows, cols] scale coordinates -> flat offsets in the 128x4-swizzled buffer.""" + padded_cols = (cols + 3) // 4 * 4 + r = torch.arange(rows, device=device).view(-1, 1) + c = torch.arange(cols, device=device).view(1, -1) + return ( + (c % 4) + + (c // 4) * (4 * 128) + + (r % 32) * 16 + + ((r % 128) // 32) * 4 + + (r // 128) * (128 * padded_cols) + ) + + +def _pad_up(x: int, m: int) -> int: + return (x + m - 1) // m * m + + +def test_linear_layout_bf16() -> None: + """Bit-exact match against the reference, decode- through prefill-sized rows.""" + torch.manual_seed(0) + # gpt-oss-120b MoE hidden (2880) and its 512-padded width (3072) + for t in (1, 2, 7, 64, 1024, 8192): + for k, alignment in ((2880, 32), (2880, 512), (3072, 512), (3072, 128)): + x = torch.randn(t, k, dtype=torch.bfloat16, device="cuda") + data, sf = mxfp8_quantize(x, False, alignment) + padded_k = _pad_up(k, alignment) + ref_data, ref_sf = _ref(x, padded_k) + assert data.shape == (t, padded_k), (t, k, alignment, data.shape) + assert data.dtype == torch.float8_e4m3fn and data.is_contiguous() + assert sf.shape == (t * padded_k // BLOCK,) and sf.dtype == torch.uint8 + assert torch.equal(data.view(torch.uint8), ref_data.view(torch.uint8)) + assert torch.equal(sf.view(t, padded_k // BLOCK), ref_sf) + + +def test_swizzled_layout_bf16() -> None: + """Swizzled scales carry the same bytes at the 128x4 offsets; padding is zero.""" + torch.manual_seed(1) + for t, k, alignment in ( + (1, 2880, 512), + (3, 128, 32), + (129, 96, 32), + (200, 3072, 512), + ): + x = torch.randn(t, k, dtype=torch.bfloat16, device="cuda") + data_sw, sf_sw = mxfp8_quantize(x, True, alignment) + padded_k = _pad_up(k, alignment) + cols = padded_k // BLOCK + ref_data, ref_sf = _ref(x, padded_k) + assert sf_sw.shape == (_pad_up(t, 128) * _pad_up(cols, 4),) + # data is layout-independent + assert torch.equal(data_sw.view(torch.uint8), ref_data.view(torch.uint8)) + idx = _swizzle_index(t, cols, x.device).reshape(-1) + assert torch.equal(sf_sw[idx].view(t, cols), ref_sf) + # every offset not addressed by a real (row, col) is row/column padding: zero + rest = torch.ones_like(sf_sw, dtype=torch.bool) + rest[idx] = False + assert (sf_sw[rest] == 0).all() + + +def test_alignment_padding() -> None: + """alignment pads K with zeros: zero data bytes, zero scale bytes, valid part unchanged.""" + torch.manual_seed(2) + x = torch.randn(64, 2880, dtype=torch.bfloat16, device="cuda") + tight_data, tight_sf = mxfp8_quantize(x, False, 32) + padded_data, padded_sf = mxfp8_quantize(x, False, 512) + assert padded_data.shape == (64, 3072) and padded_sf.shape == (64 * 96,) + assert torch.equal( + padded_data[:, :2880].reshape(-1).view(torch.uint8), + tight_data.reshape(-1).view(torch.uint8), + ) + assert torch.equal(padded_sf.view(64, 96)[:, :90], tight_sf.view(64, 90)) + assert (padded_data[:, 2880:].view(torch.uint8) == 0).all() + assert (padded_sf.view(64, 96)[:, 90:] == 0).all() + + +def test_fp16_input() -> None: + torch.manual_seed(3) + for t in (1, 512): + x = torch.randn(t, 2880, dtype=torch.float16, device="cuda") + data, sf = mxfp8_quantize(x, False, 512) + ref_data, ref_sf = _ref(x, 3072) + assert data.dtype == torch.float8_e4m3fn + assert torch.equal(data.view(torch.uint8), ref_data.view(torch.uint8)) + assert torch.equal(sf.view(t, 96), ref_sf) + + +def test_3d_input_collapses_leading_dims() -> None: + torch.manual_seed(4) + x = torch.randn(2, 5, 2880, dtype=torch.bfloat16, device="cuda") + data, sf = mxfp8_quantize(x, False, 512) + assert data.shape == (2, 5, 3072) + assert sf.shape == (2 * 5 * 96,) + ref_data, ref_sf = _ref(x, 3072) + assert torch.equal(data.reshape(10, 3072).view(torch.uint8), ref_data.view(torch.uint8)) + assert torch.equal(sf.view(10, 96), ref_sf) + # swizzled treats the collapsed rows as one 128-row-padded matrix + _, sf_sw = mxfp8_quantize(x, True, 512) + assert sf_sw.shape == (128 * 96,) + + +def test_dequant_error_bound() -> None: + """data * 2^(sf-127) reconstructs the input within the mxfp8 rounding bound.""" + torch.manual_seed(5) + x = torch.randn(256, 3072, dtype=torch.bfloat16, device="cuda") * 3.0 + data, sf = mxfp8_quantize(x, False, 32) + scale = torch.exp2(sf.view(256, 96).to(torch.float32) - 127.0) + deq = (data.float().view(256, 96, BLOCK) * scale.unsqueeze(-1)).view(256, 3072) + # e4m3 has a 3-bit mantissa: a normal value carries at most 2^-4 relative + # rounding error; a subnormal one at most half of the 2^-9 subnormal step, + # i.e. 2^-10 absolute in units of the block scale (per-block atol below). + atol = (2**-10) * scale.repeat_interleave(BLOCK, dim=-1) + err = (deq - x.float()).abs() + assert (err <= 2**-4 * x.float().abs() + atol).all() + # hard gate: the quantization is the exact reference, not merely close + ref_data, ref_sf = _ref(x, 3072) + assert torch.equal(data.view(torch.uint8), ref_data.view(torch.uint8)) + assert torch.equal(sf.view(256, 96), ref_sf) + + +def test_power_of_two_amax_boundary() -> None: + """amax exactly 448*2^k: the scale byte lands on 127+k, the max element on +-448.""" + for k in (-2, 0, 1, 3): + x = torch.zeros(1, 32, dtype=torch.bfloat16, device="cuda") + x[0, 0] = 448.0 * 2.0**k + x[0, 1] = -112.0 * 2.0**k + data, sf = mxfp8_quantize(x, False, 32) + assert sf.tolist() == [127 + k], (k, sf.tolist()) + assert data[0, 0].float().item() == 448.0 + assert data[0, 1].float().item() == -112.0 + + +def test_extreme_magnitude_block() -> None: + """The largest finite bf16 magnitude still quantizes exactly, no inf/NaN.""" + big = torch.finfo(torch.bfloat16).max + x = torch.zeros(1, 32, dtype=torch.bfloat16, device="cuda") + x[0, 0] = big + x[0, 1] = -big / 256.0 + data, sf = mxfp8_quantize(x, False, 32) + ref_data, ref_sf = _ref(x, 32) + # amax/448 = 2^119.19 -> scale 2^120 -> byte 247, well below the E8M0 max + assert sf.tolist() == [247] + assert torch.equal(data.view(torch.uint8), ref_data.view(torch.uint8)) + assert torch.equal(sf, ref_sf.reshape(-1)) + assert not torch.isnan(data.float()).any() + + +def test_nonfinite_inputs() -> None: + """Non-finite lanes are pinned here because both are destructive and silent.""" + # +inf: the block scale saturates to the E8M0 finite max (2^127), whose + # reciprocal flushes to zero -- the inf lane becomes NaN and every finite + # lane in the same block is zeroed. + x = torch.zeros(1, 32, dtype=torch.bfloat16, device="cuda") + x[0, 0] = float("inf") + x[0, 1] = 1.0 + data, sf = mxfp8_quantize(x, False, 32) + assert sf.tolist() == [254] + assert torch.isnan(data[0, 0].float()) + assert (data[0, 1:].float() == 0.0).all() + + # NaN: excluded from the block max, so the scale comes from the finite + # lanes; the NaN stays confined to its own lane. + y = torch.zeros(1, 32, dtype=torch.bfloat16, device="cuda") + y[0, 0] = float("nan") + y[0, 1] = 1.0 + data, sf = mxfp8_quantize(y, False, 32) + assert sf.tolist() == [119] # amax = 1.0 -> scale 2^-8 + assert torch.isnan(data[0, 0].float()) + assert data[0, 1].float().item() == 256.0 + assert (data[0, 2:].float() == 0.0).all() + + +def test_zero_and_denormal_blocks() -> None: + """Pins the two degenerate block-scale paths, including the NaN one.""" + # all-zero block: scale byte 0x00 and all-zero data bytes + x = torch.zeros(1, 64, dtype=torch.bfloat16, device="cuda") + data, sf = mxfp8_quantize(x, False, 32) + assert sf.tolist() == [0, 0] + assert (data.view(torch.uint8) == 0).all() + + # amax = 2^-118 > 448*2^-127: still the normal path, exact reference match + y = torch.zeros(1, 32, dtype=torch.bfloat16, device="cuda") + y[0, 0] = 2.0**-118 + y[0, 1] = 2.0**-119 + data, sf = mxfp8_quantize(y, False, 32) + ref_data, ref_sf = _ref(y, 32) + assert sf.tolist() == [1] + assert torch.equal(data.view(torch.uint8), ref_data.view(torch.uint8)) + assert not torch.isnan(data.float()).any() + + # amax = 2^-119 <= 448*2^-127: the block scale underflows to the E8M0 minimum + # (byte 0x00 = 2^-127, a denormal fp32) and the kernel's flush-to-zero + # reciprocal turns it into +inf -- nonzero lanes saturate to +-448, exact-zero + # lanes become NaN. Wrong answers, no error raised. + z = torch.zeros(1, 32, dtype=torch.bfloat16, device="cuda") + z[0, 0] = 2.0**-119 + z[0, 1] = -(2.0**-120) + data, sf = mxfp8_quantize(z, False, 32) + assert sf.tolist() == [0] + assert data[0, 0].float().item() == 448.0 + assert data[0, 1].float().item() == -448.0 + assert torch.isnan(data[0, 2:].float()).all() + + +def test_zero_rows() -> None: + """An empty batch is accepted, not rejected: both outputs come back empty.""" + x = torch.randn(0, 2880, dtype=torch.bfloat16, device="cuda") + data, sf = mxfp8_quantize(x, False, 512) + assert data.shape == (0, 3072) and data.dtype == torch.float8_e4m3fn + assert sf.shape == (0,) and sf.dtype == torch.uint8 + + +def test_input_not_mutated() -> None: + torch.manual_seed(6) + x = torch.randn(32, 2880, dtype=torch.bfloat16, device="cuda") + before = x.clone() + mxfp8_quantize(x, False, 512) + mxfp8_quantize(x, True, 32) + assert torch.equal(x, before) + + +def test_rejects_unsupported() -> None: + """Domains the contract declares unsupported must be rejected, not silently wrong.""" + x = torch.randn(4, 128, dtype=torch.bfloat16, device="cuda") + cases = { + "k not a multiple of 32": lambda: mxfp8_quantize( + torch.randn(4, 112, dtype=torch.bfloat16, device="cuda"), False, 32 + ), + "alignment not a multiple of 32": lambda: mxfp8_quantize(x, False, 48), + "alignment below the block size": lambda: mxfp8_quantize(x, False, 16), + "fp32 input": lambda: mxfp8_quantize(x.float(), False, 32), + "1-D input": lambda: mxfp8_quantize(x.reshape(-1), False, 32), + "non-contiguous column slice": lambda: mxfp8_quantize( + torch.randn(8, 256, dtype=torch.bfloat16, device="cuda")[:, :128], False, 32 + ), + "non-contiguous transpose": lambda: mxfp8_quantize( + torch.randn(128, 8, dtype=torch.bfloat16, device="cuda").t(), False, 32 + ), + "cpu input": lambda: mxfp8_quantize(torch.randn(4, 128, dtype=torch.bfloat16), False, 32), + } + for name, fn in cases.items(): + try: + fn() + except (RuntimeError, NotImplementedError): + continue + raise AssertionError(f"{name}: expected a raise, got a result") diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/__init__.py b/tensorrt_llm/_torch/staircase/catalog/torch/__init__.py new file mode 100644 index 000000000000..9c12fb0b373a --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/torch/__init__.py @@ -0,0 +1,7 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Thin mirrors of torch callables. + +Upstream owns their correctness, so these carry no contract, no test and +no receipt by design. They exist so a target's forward can satisfy the +closed-vocabulary rule without leaving the catalog.""" diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/add.py b/tensorrt_llm/_torch/staircase/catalog/torch/add.py new file mode 100644 index 000000000000..9eb889284740 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/torch/add.py @@ -0,0 +1,14 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Elementwise addition (thin mirror of torch.add). + +torch does not implement float8 arithmetic; inputs are expected in the +high-precision residual dtypes (fp32/bf16/fp16). +""" + +import torch + + +def add(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: + """Return `a + b`; identical semantics to torch.add.""" + return torch.add(a, b) diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/concat.py b/tensorrt_llm/_torch/staircase/catalog/torch/concat.py new file mode 100644 index 000000000000..8b60d8219c04 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/torch/concat.py @@ -0,0 +1,10 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Tensor concatenation (thin mirror of torch.cat).""" + +import torch + + +def concat(tensors: list[torch.Tensor], dim: int = 0) -> torch.Tensor: + """Concatenate `tensors` along `dim`; identical semantics to torch.cat.""" + return torch.cat(tensors, dim=dim) diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/copy_.py b/tensorrt_llm/_torch/staircase/catalog/torch/copy_.py new file mode 100644 index 000000000000..0c32b729b4a9 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/torch/copy_.py @@ -0,0 +1,16 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""In-place copy into a tensor or view (thin mirror of torch.Tensor.copy_). + +Writes `src` into `dst` element-wise (broadcasting and dtype conversion follow +torch semantics) and returns `dst`. The canonical way to fill a slice of a +larger buffer (e.g. `copy_(k[..., :nope], k_nope)`); the destination's +pre-call contents are destroyed. +""" + +import torch + + +def copy_(dst: torch.Tensor, src: torch.Tensor) -> torch.Tensor: + """Copy `src` into `dst` in place; identical semantics to Tensor.copy_.""" + return dst.copy_(src) diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/embedding.py b/tensorrt_llm/_torch/staircase/catalog/torch/embedding.py new file mode 100644 index 000000000000..d9057a48f7e2 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/torch/embedding.py @@ -0,0 +1,11 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Input-embedding lookup (thin mirror of torch.nn.functional.embedding).""" + +import torch +import torch.nn.functional as F + + +def embedding(input_ids: torch.Tensor, weight: torch.Tensor) -> torch.Tensor: + """Look up rows of `weight` by token id: [*] ids -> [*, H] embeddings.""" + return F.embedding(input_ids, weight) diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/empty.py b/tensorrt_llm/_torch/staircase/catalog/torch/empty.py new file mode 100644 index 000000000000..4f0c5c805193 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/torch/empty.py @@ -0,0 +1,15 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Uninitialized tensor allocation (thin mirror of torch.empty). + +The general-purpose allocation primitive for scratch and output buffers +whose sizing the caller owns; layer-aware output allocation has dedicated +entries (e.g. create_attn_outputs). Contents are garbage until written. +""" + +import torch + + +def empty(shape: list[int], dtype: torch.dtype, device: torch.device | str) -> torch.Tensor: + """Return an uninitialized tensor; identical semantics to torch.empty.""" + return torch.empty(shape, dtype=dtype, device=device) diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/expand.py b/tensorrt_llm/_torch/staircase/catalog/torch/expand.py new file mode 100644 index 000000000000..7265fe27216e --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/torch/expand.py @@ -0,0 +1,15 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Broadcast view over singleton dims (thin mirror of torch.Tensor.expand). + +Returns a zero-copy view with stride 0 on the expanded dims; only singleton +dims can be expanded, and -1 keeps a dim unchanged. Writing through the view +is unsafe (aliased elements) — copy first if mutation is needed. +""" + +import torch + + +def expand(x: torch.Tensor, sizes: list[int]) -> torch.Tensor: + """Return `x` broadcast to `sizes`; identical semantics to Tensor.expand.""" + return x.expand(sizes) diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/pad.py b/tensorrt_llm/_torch/staircase/catalog/torch/pad.py new file mode 100644 index 000000000000..e02e16737c0c --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/torch/pad.py @@ -0,0 +1,16 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Constant padding of a tensor's trailing dims (thin mirror of torch.nn.functional.pad). + +Layout glue: widens a tensor to a kernel's required operand width by +appending constant-valued columns (the padded region multiplies zero-valued +padded weights in block-scaled GEMMs, so the fill value must be finite). +""" + +import torch +import torch.nn.functional as F + + +def pad(x: torch.Tensor, padding: list[int], value: float = 0.0) -> torch.Tensor: + """Return `x` padded with `value` per the last-dim-first `padding` pairs.""" + return F.pad(x, padding, mode="constant", value=value) diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/reshape.py b/tensorrt_llm/_torch/staircase/catalog/torch/reshape.py new file mode 100644 index 000000000000..f167787c3b52 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/torch/reshape.py @@ -0,0 +1,14 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Tensor reshape (thin mirror of torch.reshape). + +Covers view use cases too: returns a view when the layout allows, copies +otherwise — unlike Tensor.view, it never fails on non-contiguous input. +""" + +import torch + + +def reshape(x: torch.Tensor, shape: list[int]) -> torch.Tensor: + """Reshape `x` to `shape`; identical semantics to torch.reshape.""" + return torch.reshape(x, shape) diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/split.py b/tensorrt_llm/_torch/staircase/catalog/torch/split.py new file mode 100644 index 000000000000..6653e4e6b1f0 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/torch/split.py @@ -0,0 +1,12 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Tensor split (thin mirror of torch.split).""" + +import torch + + +def split( + tensor: torch.Tensor, split_size_or_sections: int | list[int], dim: int = 0 +) -> tuple[torch.Tensor, ...]: + """Split `tensor` along `dim`; identical semantics to torch.split.""" + return torch.split(tensor, split_size_or_sections, dim=dim) diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/transpose.py b/tensorrt_llm/_torch/staircase/catalog/torch/transpose.py new file mode 100644 index 000000000000..1a6a2235a94a --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/torch/transpose.py @@ -0,0 +1,15 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Zero-copy swap of two dimensions (thin mirror of torch.transpose). + +Returns a strided view sharing storage with the input; the canonical way to +present a `[tokens, heads, dim]` activation to a batched-gemm entry whose +batch axis is the head axis. +""" + +import torch + + +def transpose(x: torch.Tensor, dim0: int, dim1: int) -> torch.Tensor: + """Return `x` with `dim0` and `dim1` swapped; identical semantics to torch.transpose.""" + return torch.transpose(x, dim0, dim1) diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/view_dtype.py b/tensorrt_llm/_torch/staircase/catalog/torch/view_dtype.py new file mode 100644 index 000000000000..59d33774d4d7 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/catalog/torch/view_dtype.py @@ -0,0 +1,16 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Bitwise dtype reinterpretation of a tensor (thin mirror of torch.Tensor.view(dtype)). + +Reinterprets the same bytes under another dtype of equal itemsize — no +conversion, no copy. The canonical bridge between a producer that emits a +byte buffer (e.g. packed block scales as uint8) and a consumer that demands +the typed view of those bytes (float8_e4m3fn). +""" + +import torch + + +def view_dtype(x: torch.Tensor, dtype: torch.dtype) -> torch.Tensor: + """Reinterpret `x`'s bytes as `dtype`; identical semantics to Tensor.view(dtype).""" + return x.view(dtype) diff --git a/tensorrt_llm/_torch/staircase/docs/models/expert-weight-packing.md b/tensorrt_llm/_torch/staircase/docs/models/expert-weight-packing.md new file mode 100644 index 000000000000..3d642df6b457 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/docs/models/expert-weight-packing.md @@ -0,0 +1,553 @@ +# Expert weight packing (sparse MoE) + +What a target assembled on a sparse-MoE checkpoint had to establish by +observation, because no contract or reference stated it. Written from the +qwen3-30b-a3b/sm_100/tp1 onboard (Qwen3-30B-A3B: 48 layers, all MoE, 128 +experts, top-8, `moe_intermediate_size` 768, bf16). The mechanism — a +router selecting k of E stacked expert MLPs — recurs across model +families; this file is about the mechanism, not that checkpoint. + +**Every shape below is a whole-stack shape, measured at world size 1.** +Under expert parallelism a rank holds `E / ep_size` of the stack and the +first dimension of every weight and scale shrinks accordingly. + +**What survives the split — established by the first expert-parallel +target** (`tep4`, 72 experts over 4 ranks, 18 per rank at offset +`18 * ep_rank`): + +* **Every per-expert preparation step survives unchanged.** The + `[up; gate]` concat, the interleave, the 32-row block shuffle and the + 128x4 scale swizzle all happen *within* one expert, and EP splits the + stack on the expert axis only. So every transform below applies + verbatim, with the first dimension `E → E/ep_size` and the destination + index `e - local_expert_offset`. Nothing about the shuffle or the + swizzle is a function of the stack height. +* **The three per-expert scale scalars must be `[local_num_experts]`, but + the checkpoint's scalars must still be read whole.** The runner rejects + `num_experts`-long scalars, so the window slice is real — but the + *validity* argument for the shared-expert activation scale is global + (`shared.input_scale == max` over all `E` routed `input_scale`), and a + target that windows at load time can no longer assert it. Load all `E` + (six fp32 per expert — negligible), assert globally, slice in the + post-load derivation. +* **`leftover == {}` stops being the coverage assert.** Under EP each rank + legitimately leaves the off-window experts' weight and scale keys + unconsumed. Replace it with an explicitly predicted leftover set + computed from the rank's window — that keeps the assert bidirectional + rather than relaxing it to a warning, which is the whole value of having + it. + +One cross-rank invariant a single rank cannot check, recorded because the +assembly rests on it: **the four windows must tile the routing space +exactly once.** They do because the router GEMM and the routing op are +replicated and deterministic and every rank sees identical all-reduced +inputs — so each token's top-k ids are the same on every rank, and each +id falls in exactly one window. + +**That justification is topology-specific, and it does not survive +attention data parallelism.** Under a `depN` segment the ranks hold +*different* tokens, so "every rank sees identical inputs" is false by +construction and the invariant loses its support — while remaining just as +load-bearing, because EP still requires each expert id to belong to exactly +one rank's window. Established by the first attention-DP MoE target +(`dep4`): what restores it is **placing the token all-gather before the +router GEMM, not merely before the expert call.** The router then runs on a +byte-identical full token set on every rank, the original argument applies +again verbatim, and the four windows tile as before. + +The failure this rules out is worth naming, because it is silent. Gathering +*after* the router — the arrangement that looks equivalent, and is cheaper +by the width of the routing tensors — leaves every rank routing only its +own tokens, so a token's expert ids exist on one rank alone and the other +three windows never see it. There is no error and no hang; the output is +simply missing most of its expert contributions. Both gates would be the +only thing that catches it. + +Stated as the rule: **on any topology where the ranks' token sets differ, +the collective that reconstitutes the token set belongs upstream of the +router, and the tiling invariant is what decides that placement** — not the +expert call's input requirements, which the later placement also satisfies. + +## The shape of the vocabulary + +A bf16 MoE block is three catalog calls, not one: + +``` +router_logits = cublas_mm(x, router_weight_view) # [T, E] +expert_ids, scales = renorm_moe_routing_op(router_logits, topk) +moe_out = fused_moe(x, expert_ids, scales, fc1, None, fc2, None, ...)[0] +``` + +`fused_moe` contains the permutation, both grouped GEMMs, the gated +activation and the routing-weighted combine. It does **not** contain the +router GEMM or the selection — those stay the caller's, which is why the +routing op exists as a separate entry. + +Consequence worth stating because it looks like an omission: a +fully-MoE checkpoint's forward references **no `activation/*` entry**. The +SwiGLU lives inside `fused_moe` (`activation_type=5`). An audit that +expects `silu_and_mul` in every model's forward is wrong for this class. + +## The `[up | gate]` order — the one silent-wrong-answer trap + +`fused_moe`'s `fc1_expert_weights` is `[E, 2I, H]` with **up rows first**: + +- `fc1[e, :I]` is HF `up_proj` (trtllm `w3`) +- `fc1[e, I:]` is HF `gate_proj` (trtllm `w1`) — the sigmoid is applied to + this half + +`fc2_expert_weights` is `[E, H, I]`, plain `down_proj`. + +Getting the halves backwards produces a plausible, entirely different +result: finite, correctly scaled, no error, no NaN. Nothing in the stack +detects it — not the op, not a smoke test. Only an accuracy gate catches +it, and only because it lands ~50 points low. + +**This is the opposite of the dense convention.** A dense target packs +`gate_up` with gate rows first, because that is what +`flashinfer_silu_and_mul` expects over a packed last dim. The two orders +must never be copied across; each is correct for its own consumer. + +## Checkpoint layout, and where the key names actually live + +A 4.51-era checkpoint stores experts **unstacked**: 128 × 3 separate 2D +tensors per layer, +`model.layers.{i}.mlp.experts.{e}.{gate,up,down}_proj.weight`. Recent +transformers models the same block with stacked 3D `gate_up_proj` / +`down_proj` parameters — so **the installed modeling file is not the +source of truth for checkpoint key names**. The safetensors index is. +Read `model.safetensors.index.json`, not `modeling_*.py`, when writing +the manifest. + +The unstacked on-disk layout reaches `fused_moe`'s stacked layout with +**zero transforms** — every copy is layout-preserving: + +```python +(f"{p}.mlp.experts.{e}.up_proj.weight", (e, slice(0, inter))), +(f"{p}.mlp.experts.{e}.gate_proj.weight", (e, slice(inter, 2 * inter))), +(f"{p}.mlp.experts.{e}.down_proj.weight", (e,)), +``` + +This needs one generalization of the manifest convention: the destination +widens from a 2D `(row_start, row_end)` range to an index tuple into +`param.data`, so one table serves both fused-2D parameters (qkv) and +stacked-3D ones (experts). It is a one-line change in `fill`. + +3D parameters need no special handling anywhere else: they pass through +`MetaInitMode` and engine materialization unchanged. Measured at this +scale — 434 `ParameterDict` entries, 58 GB of expert stacks, 18432 +per-expert copies of ~3.1 MB — the whole load sits inside a 14.9 s model +init. + +## What the structure dictates, and what it makes irrelevant + +Established by measurement during the perf campaign; these are properties +of the mechanism, so the next MoE target can skip the sweeps. + +**The MoE is HBM-roofline-bound at serving batch sizes.** Above a decode +batch of roughly `E / topk × small factor` (~40 here) essentially every +expert is active, so a decode step reads the **entire** expert stack once +— 58 GB for this geometry. Measured 7.350 ms/step for the two grouped +GEMMs at concurrency 256 ⇒ **7.89 TB/s**, i.e. the B200 HBM roofline. No +config knob and no scheduling change moves this. Only fewer weight bytes +(quantization — a modeling change) would. + +**The kernel-count floor is dominated by the MoE.** `fused_moe` plus +`renorm_moe_routing_op` contribute 7 of the 16 kernels per layer +(`customMoeRouting`, `fusedBuildExpertMapsSortFirstToken`, +`expandInputRows`, grouped GEMM 1, `doActivation`, grouped GEMM 2, +`computeStridesTmaWarpSpecialized`). Across 48 layers that is 336 of 768 +kernels per decode step. At concurrency 1 the measured 3.87 ms TPOT over +794 kernels is 4.9 µs each — the same order as the smallest kernels +measured at batch 256, i.e. a fixed per-kernel floor rather than work. The +single-stream latency point is structural until the vocabulary gains a +coarser op. + +**KV capacity does not bind.** A 3B-active/30B-total model at tp1 leaves +the KV pool enormous relative to any realistic concurrency (measured: +1,168,096 pool tokens = 570 requests at ISL+OSL 2048, against a 256-request +ceiling). `kv_cache_config.free_gpu_memory_fraction` and the capacity +scheduler policy are inert — do not spend sweeps on them. + +**The autotuner is not a cold-cache risk inside a served engine.** +`fused_moe`'s contract warns that a cold tuning cache silently falls back +to a default tactic. In serving that does not apply: the runtime's +`_run_autotuner_warmup` wraps the forward in `autotune()` before capture +(observed `Cache size after warmup is 28` = 2 tunable GEMMs × 14 +power-of-2 token buckets). Worth knowing before anyone spends an iteration +on it — and worth knowing in the other direction too: the tuned tactic +changes the *bits*, not just the speed (measured 2.96 ulp against the cold +fallback on the same input), so a catalog receipt taken cold certifies a +tactic the served engine never runs. Both states sit inside the entry's +accuracy gate; the point is that "same inputs, same outputs" holds only +within one tuner state. + +## Routing traps + +Both from `renorm_moe_routing_op`'s certification: + +- **Ties break toward the lower expert index**, the opposite of + `torch.topk`. On bf16 logits exact ties are common, so a + `torch.topk`-based reference disagrees on indices while the weights + still match. Do not validate routing against `torch.topk`. +- **The kernel ignores strides**, reading `router_logits` as a dense + row-major buffer from `data_ptr()`. A strided view routes silently + wrong. The wrapper guards it; a caller building logits as a slice of a + wider buffer must materialize them contiguous. + +And from `fused_moe`: expert ids outside the rank's slot range are +**silently dropped**, not clamped and not rejected, so a routing bug +surfaces as a quietly weaker token rather than an error. + +## Deriving "is every layer MoE?" + +Do not assume uniformity, and do not assume `intermediate_size` is live. +The dense-vs-sparse branch per layer is +`layer_idx not in mlp_only_layers and num_experts > 0 and (layer_idx + 1) +% decoder_sparse_step == 0`. With `decoder_sparse_step: 1` and +`mlp_only_layers: []` every layer is sparse, which makes +`intermediate_size` **dead config** that no layer reads — while +`moe_intermediate_size` is the live one. A checkpoint with a nonzero +`decoder_sparse_step` or a non-empty `mlp_only_layers` needs both branches +and both sets of weights. Assert the derivation at construction rather +than trusting it. + +## MXFP4 expert stacks — the same trap, in interleaved form + +From the gpt-oss-120b/sm_100/tp1 onboard (36 layers, all MoE, 128 +experts, top-4, hidden = intermediate = 2880, experts MXFP4 while router, +attention, embedding and lm_head stay bf16). Everything above about the +`[up | gate]` order still holds; a block-scale-quantized stack adds three +things. + +**The parity is inverted relative to the half-split convention.** The +checkpoint stores `gate_up_proj_blocks` as `[E, 2I, K/32, 16]` uint8 — +already in `nn.Linear` `[out, in]` orientation, transposed relative to +HF's `[E, hidden, 2*inter]` parameter — and HF reads `gate = +gate_up[..., ::2]`, `up = gate_up[..., 1::2]`. So the stored row order +along the `2I` axis is (gate, up, gate, up, …), while the trtllm-gen +kernel's interleave wants destination row `2i` = **up** `i`, `2i+1` = +**gate** `i`. The split is therefore `up = t[:, 1::2]`, `gate = +t[:, 0::2]`, re-concatenated `[up ; gate]` before the row permutation. +Measured discrimination on real layer-0 tensors: 1.12 bf16 ulp against a +correct pure-torch reference, 146 ulp against the swapped one — a 130x +separation. Worth running that check *before* the first engine boot: it +covers nibble order, the parity split, the concat, the interleave, the +32-row block shuffle, the 128x4 scale swizzle, both padded axes, both +fp32 biases and the alpha/beta/limit triple in one shot, and the accuracy +gate that would otherwise catch it costs a full evaluation. + +**Where the relayout lives decides whether the model fits.** Expressing +pad → concat → interleave → shuffle → swizzle as *manifest source +transforms* (the sharding convention's `src` slot, generalized to any +callable) keeps peak memory at one layer of scratch. Declaring +checkpoint-shaped parameters and transforming them in `post_load_weights` +instead needs both forms resident — 63 GB checkpoint plus 66 GB +kernel-ready, on a 183 GB card that also wants a KV pool. The streaming +form loaded 63 GB including the on-device relayout in 14.1-14.5 s. The +kernel-ready operands at this geometry are `[128, 5888, 1536]` + +`[128, 5888, 96]` + fp32 `[128, 5888]` (FC1) and `[128, 2944, 1472]` + +`[128, 2944, 92]` + fp32 `[128, 2944]` (FC2) per layer, ~1.85 GB/layer, ++4% over the checkpoint. + +**Dtypes on disk are not the dtypes the kernels want.** Expert biases and +per-head attention sinks are bf16 in the checkpoint; the trtllm-gen MoE +requires fp32 biases and the attention op an fp32 sink tensor. Promote at +load. The same row permutation applied to the weight bytes must also be +applied to the scale bytes **and** the fp32 bias — a non-permuted bias is +a silent wrong answer. + +## W4A16 vs W4A8: the same weights, a different activation path + +Both members of the trtllm-gen block-scale MoE family consume the +**byte-identical** prepared expert stack — confirmed from the quant-method +source (both inherit one base, overriding only `create_weights` / +`load_quant_scales`, which call `super()`) and by measurement across every +certified geometry. Moving between them is a forward change only; +`weights.py` does not move. + +| | W4A16 | W4A8 | +|---|---|---| +| activations | bf16 straight in | `mxfp8_quantize(x, swizzled_layout=False, alignment=512)` first | +| hidden padding | caller pads 2880 -> 3072 | the quantizer does it — **the pad call disappears** | +| `valid_hidden_size` | 2880 | 2880 (output width; unrelated to the widening, and `None` is rejected) | + +Two consequences that are not obvious from the signatures: + +- **The W4A8 kernel requantizes the FC1 activation to MXFP8 between the + GEMMs**, on the OCP scale (`e = floor(log2 amax) - 8`), *not* the + round-up scale the standalone quantizer uses. Skipping that step in a + reference deviates 13.5 ulp element-wise / 8.5 ulp RMS (~3% relative) + from the kernel — so it is a real perturbation of the layer output, and + it is the mechanism behind the accuracy difference between the two + members. On gpt-oss-120b the measured cost was 1.44 GSM8K points + (90.5989 W4A16 -> 89.1585 W4A8), reproducible bit-identically on both + sides, i.e. the recipe's price rather than sampling. +- **`swizzled_layout=True` is silently accepted** by the W4A8 MoE + whenever the byte counts coincide (`T % 128 == 0` — exactly the CUDA + graph batch sizes), and is 260 ulp wrong. No metadata distinguishes the + two layouts, so no wrapper can guard it. Put the `False` in a named + constant. + +The perf reason to pay that accuracy: on gpt-oss-120b the W4A16 MoE ran +5.6x slower than the W4A8 one on byte-identical weights (FC1 711.8 vs +126.7 µs per layer per step at concurrency 128), which was the target's +*entire* gap to stock trtllm. The swap moved peak throughput +88% and put +the step within 0.5% of the reference. + +## The dtype chain, end to end + +Worth writing out because the names invite a wrong guess: **W4A8 +quantizes activations to fp8, never to fp4.** The op name spells both +operands — `mxe4m3_mxe2m1_...` is e4m3 activations against e2m1 weights. +The `mx` prefix is OCP micro-scaling: 32 elements share one E8M0 +(power-of-two) scale, so MXFP8 is 32 e4m3 values plus one scale byte and +MXFP4 is 32 e2m1 values plus one. + +| step | format | +|---|---| +| hidden states entering the block | bf16 `[T, H]` | +| `mxfp8_quantize(x, False, 512)` | e4m3 `[T, pad_up(H, 512)]` + E8M0, one byte per 32 | +| FC1 GEMM | e4m3 x e2m1, **fp32 accumulate** | +| clamped GLU (FC1 epilogue) | fp32 | +| requantization of the intermediate | e4m3 + E8M0 per 32 columns | +| FC2 GEMM | e4m3 x e2m1, **fp32 accumulate** | +| routing-weighted combine | fp32 | +| store | **bf16** `[T, valid_hidden_size]` | + +Two quantization points, both to fp8, with fp32 everywhere between them. +Only the operand *storage* formats are narrow — the arithmetic never drops +to fp4. The output is bf16 at the model's true hidden width, not the +padded one, and `valid_hidden_size` has to be passed explicitly to get it. + +## Why the W4A16 member is slow, and the W4A8 one is not + +The two paths read the **same weight bytes**, so the 5.6x is not +bandwidth. It is what each kernel does with them, and the kernel names say +it outright: + +| | operand fields in the kernel name | conversion | +|---|---|---| +| W4A16 | `bmm_Bfloat16_MxE2m1Bfloat16_castBfloat16_...` | `castBfloat16` — MXFP4 expanded to bf16, then a bf16 MMA | +| W4A8 | `bmm_MxE4m3_MxE2m1MxE4m3_...` | none — MXFP4 goes straight into a block-scaled MMA | + +The expansion's cost is not mainly the conversion arithmetic. A +dequantized weight block occupies **4x the on-chip bytes** (4-bit e2m1 -> +16-bit bf16), so fewer of them fit in registers and shared memory, which +shrinks the tile and shortens the software pipeline — visible in the same +kernel names: + +- W4A16: `m128x8x16`, 3 pipeline stages, 1-CTA clusters +- W4A8: `m256x16x32`, 6 stages, 2-CTA clusters + +Fewer weight bytes in flight means HBM latency stops being hidden, and the +GEMM lands far below the memory roofline instead of at it. At concurrency +128 on gpt-oss-120b, where the batch touches essentially every expert so a +decode step reads the whole stack once (~61 GB at valid sizes): + +| | MoE per step | implied rate | +|---|---|---| +| W4A16 | 38.89 ms | ~1.6 TB/s | +| W4A8 | 8.34 ms | ~7.3 TB/s | +| stock trtllm | 7.51 ms | ~8.1 TB/s | + +against a roofline this repo measured at 7.89 TB/s on the same machine. So +the W4A16 member was not paying more bandwidth — it was failing to use the +bandwidth it had, at roughly a fifth of the achievable rate. The residual +W4A8-vs-reference gap is autotuner tactic choice (~11%), not the recipe. + +*Measured:* the kernel names with their tile / stage / cluster fields, the +per-GEMM times (FC1 711.8 vs 126.7 µs, FC2 352.9 vs 65.6 µs per layer per +step), the MoE step times, and that the weights are byte-identical. +*Derived:* the byte figure behind the implied rates, and the chain from +on-chip footprint to tile size to unhidden latency — consistent with every +number above, but not isolated experimentally. + +## Blackwell only, and it fails quietly elsewhere + +This kernel family is sm_100 / sm_103. Three independent signals: + +- TensorRT-LLM gates its own tests on it — `get_sm_version() not in (100, + 103)`, reason "TRTLLM Gen MoE supports SM100 and SM103 only". +- The quantization kernel's body compiles only under `__CUDA_ARCH__ >= + 1000`; below that the empty branch **returns without writing the + outputs**, so a pre-Blackwell run produces garbage rather than failing. + Both catalog entries carry sm_100-only receipts for that reason — a + statement about the hardware, not caution. +- The performance signature itself. A software emulation of the mixed + fp8 x fp4 MMA — unpack, widen, then a wide MMA — would look exactly like + the W4A16 path measured above, and nothing can be 5.6x faster than the + thing it emulates. + +The hardware feature is Blackwell's block-scaled MMA: the instruction +takes narrow operands plus their per-block scale factors and applies the +scaling inside the tensor-core datapath, so no widened operand is ever +materialized. The E8M0 scale being a power of two is what makes that +nearly free — it is an exponent adjustment. (The arch gating and the +performance signature are measured here; the datapath description is the +standard account of the feature and was not verified at the ISA level.) + +## NVFP4 expert stacks — the same traps, plus per-expert scales + +From the deepseek-v3-lite-nvfp4/sm_100/tp1 onboard (DeepSeek-V3-Lite: 30 +layers, layer 0 dense, layers 1-29 MoE with 72 routed experts at top-6, +`moe_intermediate_size` 1536, plus 2 shared experts fused into one +`[3072, ...]` dense pair; experts NVFP4 while attention, router, embed +and lm_head stay bf16). Everything above about the `[up | gate]` order +still holds. NVFP4 differs from the MXFP4 case in three ways that matter. + +**The consumer is a different op with a different operand shape.** The +trtllm-gen NVFP4 MoE takes `[E, 2I, H/16]` **`float8_e4m3fn`** scales, +where the MXFP4 members take `[E, 2I, H/32]` uint8 UE8M0. The prepared +stacks are **not** interchangeable between the two families, even though +both are "block-scale MoE runners". + +| operand | shape | dtype | +|---|---|---| +| `gemm1_weights` | `[E, 2I, H/2]` | uint8 (two e2m1 codes per byte) | +| `gemm1_weights_scale` | `[E, 2I, H/16]` | float8_e4m3fn | +| `gemm2_weights` | `[E, H, I/2]` | uint8 | +| `gemm2_weights_scale` | `[E, H, I/16]` | float8_e4m3fn | + +Per-expert preparation is: concat `[up ; gate]` (**up rows first**) -> +interleave (dst row `2i` = up `i`, `2i+1` = gate `i`) -> 32-row block +shuffle (src `4u+v` -> dst `8v+u`) applied to weight bytes **and** scale +bytes -> 128x4 swizzle of the scales only. +`torch.ops.trtllm.shuffle_matrix` and +`torch.ops.trtllm.block_scale_interleave` perform the last two; both are +load-time transforms, outside the closed-vocabulary rule. + +Measured discrimination against a correct pure-torch reference, on real +checkpoint tensors — the reason to run this check before the first engine +boot rather than after an accuracy gate: + +| variant | bf16 ulp | +|---|---| +| **correct** | **4.75** | +| `[gate\|up]` instead of `[up\|gate]` | 263 | +| weight scales not 128x4 swizzled | 408 | +| expert rows not 32-row shuffled | 508 | +| `output1_scale_scalar` / `_gate_scalar` swapped | 177 | +| no intermediate requantization | 42.7 | + +**The half-order trap appears twice in one layer on a shared-expert +model.** The dense and shared-expert paths pack `[gate ; up]` for +`flashinfer_silu_and_mul`; the MoE FC1 packs `[up ; gate]`. Both live in +the same decoder layer, three lines apart. The two orders must never be +copied across — each is correct for its own consumer. + +### modelopt stores reciprocals + +A modelopt NVFP4 checkpoint stores `input_scale` and `weight_scale_2` as +`amax / (448 * 6)` — the **reciprocal** of the global scale a quantizer +divides by. So: + +``` +global_scale (for fp4_quantize) = 1 / input_scale +dense GEMM alpha = input_scale * weight_scale_2 # both off disk +``` + +Reading them the other way is wrong by `global_scale^2`, finite, and +invisible short of an accuracy gate. Settle it by arithmetic rather than +by convention: under the reciprocal reading the implied amax values are +0.2-3.0, physically sensible for normed activations and NVFP4 weights; +under the other reading they are ~1e7. `gate_proj` and `up_proj` share +both scalars exactly (verified across all experts of every MoE layer) — +assert it at load rather than assuming it. + +The trtllm-gen runner's three `[E]` fp32 arrays, in checkpoint terms: + +``` +output1_scale_gate_scalar[e] = input_scale_gate_up * weight_scale_2_gate_up[e] +output1_scale_scalar[e] = output1_scale_gate_scalar[e] / input_scale_down +output2_scale_scalar[e] = input_scale_down * weight_scale_2_down[e] +``` + +The gate scalar dequantizes the half feeding the sigmoid; the other +carries the extra FC2-input global scale because the FC1 epilogue +re-quantizes its output to NVFP4. + +### One activation scale, but per-expert weight scales + +Two facts that look like they should match and do not. + +**The activation global scale `g1` is one number, and the shared expert +carries it.** The runner quantizes the hidden states once, so it takes a +single `g1`. Which one? On this checkpoint the shared expert's +`gate_proj.input_scale` is **exactly** the max over all routed experts' +— exact equality in every MoE layer. That is not a coincidence worth +guessing at: the shared expert sees every token, so its amax *is* the +global amax. `g1 = 1 / shared.gate_proj.input_scale` therefore serves +both the dense and MoE quantization calls. Routed experts' own values sit +up to 1.4-2.3x below it. + +**The FC2 input scale `g2` is genuinely per-expert** and does not +collapse. It spans **54.5x** across experts in a single layer. The +runner's three `[E]` arrays exist precisely to carry it; feeding a scalar +would be wrong by that factor on the extreme experts. + +### Two quantization calls, not one + +`nvfp4_gemm` (dense and shared-expert linears) consumes the **swizzled** +scale-factor buffer; the trtllm-gen MoE runner consumes the **linear** +one, viewed as `float8_e4m3fn`. The same hidden states must therefore be +quantized **twice**, once in each layout — they cannot share a result. + +The failure mode is nasty: at `T % 128 == 0` the swizzled buffer happens +to be the *right size* for the runner, so it is silently accepted and is +276 ulp wrong. At every other token count it is rejected on size. A smoke +test at a round batch size will not catch it, and a smoke test at an odd +one will look like a size bug rather than a layout bug. + +### Two smaller obligations + +- **`hidden` must be a multiple of 256** for the trtllm-gen NVFP4 runner. + At a hidden that is only a multiple of 128 the kernel reads FC1 weight + scales for the wrong blocks: no error, 301-3370 bf16 ulp wrong. This is + an upstream hazard too — trtllm's own weight creation only raises + alignment to 256 when `hidden_size > 1024`. +- **Per-tensor quantization scalars are stored 0-dim.** `entry[:]` on + them raises `IndexError: slice() cannot be applied to a 0-dim tensor`, + so a manifest's materialize step needs a rank check that the bf16 path + never needed. + +## Group-limited routing, and where the grouping lives + +From `deepseek-r1-0528-nvfp4/sm_100/dep4` (256 routed experts, top-8, +`n_group 8` / `topk_group 4`), the first target on a grouped checkpoint — +its siblings all run `n_group: 1`. + +**The grouping is entirely the routing op's.** `noaux_tc_op` takes +`n_group`/`topk_group` and does the group-limited top-k; the block-scale +MoE runner's own `n_group`/`topk_group` arguments stay **`None`** at +`routing_method_type = 1`, because the runner is on its pre-routed path +and consumes finished `topk_ids`/`topk_weights`. Passing the grouping +twice is not how it composes. + +Two constraints the grouped path adds that the ungrouped one does not, and +that the routing op reports as one opaque "unsupported configuration": at +`n_group > 1` it requires `1 <= topk_group <= n_group`, `topk <= 8`, +`num_experts <= 256`, and `num_experts / n_group <= 32`. Assert them +separately or a violation is unattributable. + +**Router bias dtype.** The combine weights follow the **logits'** dtype, +not the bias's. A bf16 router GEMM with the fp32 `e_score_correction_bias` +this checkpoint ships returns bf16 weights — which is what the MoE runner +demands (it rejects fp32 `topk_weights`). No cast is needed and adding one +is a mistake. + +## A four-way expert-parallel split at 256 experts + +The invariant is the one this file already records — the windows must tile +the routing space exactly once — now measured at 64 wide: + +* the four 64-wide windows **sum to the whole-layer result** (0.85 ulp RMS + against an fp32 reference), but are **not bitwise equal** to a single + 256-expert call: up to 2.00 ulp apart, because each window rounds its own + partial to bf16 before the add. Dropping any one window lands ≥ 125 ulp + RMS away, so the sum is a real check rather than a formality; +* the **autotuner's key holds `local_num_experts` but not + `local_expert_offset`** — warming one window of a split warms all four, + and a 256 → 64 change forces a fresh sweep. Cold and warm agree bitwise + across 85 configurations here, but do not generalize that: the bf16 + `fused_moe` entry records the same check *failing*. diff --git a/tensorrt_llm/_torch/staircase/docs/models/multi-token-prediction.md b/tensorrt_llm/_torch/staircase/docs/models/multi-token-prediction.md new file mode 100644 index 000000000000..9c9c8e5e40ab --- /dev/null +++ b/tensorrt_llm/_torch/staircase/docs/models/multi-token-prediction.md @@ -0,0 +1,299 @@ +# Multi-token prediction (MTP) + +What a target assembled on a checkpoint that ships an MTP module had to +establish about **what the extra layer computes**, because no contract +states it and the checkpoint ships no reference implementation for it. +Written from the deepseek-r1-0528-nvfp4/sm_100/dep4 MTP increment +(DeepSeek-R1-0528, 61 trunk layers, hidden 7168, 256 routed experts top-8, +one MTP module at layer 61). The mechanism — an extra decoder layer that +consumes *both* the trunk's hidden state and the next token's embedding, and +is replayed to produce a draft — recurs across the DeepSeek family and its +derivatives; this file is about the mechanism, not that checkpoint. + +The runtime half of the picture — how the engine finds the drafter, what +calls the layer, and how the batch shape changes between draft steps — is +`docs/references/trtllm-runtime-integration.md` §13. This file is only the +computation and the weights. + +**The checkpoint is the whole specification here, and it is a partial one.** +DeepSeek's published `modeling_deepseek.py` does **not** implement the MTP +module (verified across four R1/R1-0528 checkpoints on disk, quantized and +not), and `transformers`' own `deepseek_v3` does not either. So the graph +below was reconstructed from the checkpoint's key names, shapes and weight +statistics. Every claim that rests on the reconstruction rather than on a +shape is marked. + +## The mechanism + +The module is one extra decoder layer, structurally identical to a trunk MoE +layer, with a front end bolted on that mixes in the embedding of the token +being predicted: + +``` +e = embed_tokens(input_ids) # embedding of the NEXT token +x = eh_proj( concat( enorm(e), hnorm(h) ) ) # [T, 2*hidden] -> [T, hidden] +x = x + MLA( input_layernorm(x) ) # same block as a trunk layer +x = x + MoE( post_attention_layernorm(x) ) # same structure, different dtype +# shared_head, called separately by the runtime: +logits = lm_head( shared_head.norm(x) ) +``` + +`h` is the trunk's final hidden state for the same position — under +one-model MTP-Eagle the runtime hands the layer the target model's own +output, with no projection in between. + +**"How many MTP layers do we enable" is not the knob.** A checkpoint with +one MTP module (`num_nextn_predict_layers: 1`) is replayed autoregressively: +the same layer runs `max_draft_len` times, each step feeding it the previous +step's output token and hidden state. The draft length is a serving knob, +not a checkpoint property. §13 has the mode-selection rule. + +## The concat order — the one silent-wrong-answer trap + +`eh_proj` is `[hidden, 2*hidden]`. Which half multiplies the embedding +branch and which multiplies the hidden branch is **not** determined by +anything in the checkpoint's metadata, and getting it wrong is a pure +numerical error: no shape mismatch, no assert, no crash. Under rejection +sampling it does not even corrupt output — the drafts are simply always +rejected (see the last section). + +**Two sources disagree, and the naming is the one that is right.** + +* The DeepSeek-V3 report writes the projection as `M_k [ RMSNorm(h) ; + RMSNorm(Emb(t)) ]` — **hidden first**. +* The parameter is named `eh_proj`, with its two gains named `enorm` and + `hnorm` — **embedding first**. + +**Measured: the embedding block is first.** Two independent statistics over +the checkpoint's own weights, each controlled: + +*Column-norm profile.* Take the per-input-dimension L2 norm of each half of +`eh_proj` and correlate it with each RMSNorm gain. The pairing is exclusive: + +| | vs `\|enorm\|` | vs `\|hnorm\|` | vs `\|model.norm\|` | +|---|---|---|---| +| first half `eh_proj[:, :hidden]` | **+0.904** | +0.037 | +0.026 | +| second half `eh_proj[:, hidden:]` | −0.036 | −0.207 | **+0.544** | + +`model.norm` is the trunk's final RMSNorm, which gates the same residual +stream `hnorm` does, so the second half tracking it is the same statement as +the second half being the hidden branch. + +*Functional response.* Feed the embedding branch's actual input, +`enorm(Emb(t))`, through each half and out through `shared_head`, and +measure the entropy of the resulting distribution. **Compare each half +against its own controls, not against the other half** — the first half is +uniformly sharper by ~0.9 nats on *any* input, so a raw cross-half +comparison is confounded: + +| input | via first half | via second half | +|---|---|---| +| `enorm(Emb(t))` — the real embedding-branch input | **2.62 nats** | 6.22 nats | +| `hnorm(gaussian)` | 5.33 | 6.24 | +| gaussian, no gain (control) | 5.36 | 6.23 | +| `enorm(gaussian)` — right gain, wrong vector (control) | 5.09 | 6.17 | + +The first half drops **2.5 nats** on the real embedding and on nothing else. +The second half does not move at all: through it, a real token embedding is +indistinguishable from noise. (Uniform over a 129,280 vocabulary is 11.77 +nats.) + +So: **`eh_proj` consumes `concat(enorm(e), hnorm(h))`.** This is a +reconstruction from weight statistics, not a reading of reference code — it +is strong enough to implement against, and the acceptance rate is what +confirms it end to end. Keep the order as a single named constant at one +place in the layer so flipping it is a one-line experiment. + +## Checkpoint layout + +The module is stored as one more entry in the layer list, at index +`num_hidden_layers`. On this checkpoint that is 790 keys under +`model.layers.61.`: + +| key | count | dtype | shape | load? | +|---|---|---|---|---| +| `enorm.weight`, `hnorm.weight` | 2 | bf16 | `[7168]` | yes | +| `eh_proj.weight` | 1 | bf16 | `[7168, 14336]` | yes | +| `input_layernorm`, `post_attention_layernorm` | 2 | bf16 | `[7168]` | yes | +| `self_attn.*` | 7 | bf16 | **identical to a trunk layer's** | yes | +| `self_attn.{k_proj.k_scale, v_proj.v_scale}` | 2 | fp32 | `[]` | value is **1.0**, as in the trunk | +| `mlp.gate.weight` | 1 | bf16 | `[256, 7168]` | yes | +| `mlp.gate.e_score_correction_bias` | 1 | **fp32** | `[256]` | yes | +| `mlp.experts.{0..255}.{gate,up,down}_proj` | 768 | bf16 | `[2048, 7168]` ×2, `[7168, 2048]` | yes, EP-windowed | +| `mlp.shared_experts.{gate,up,down}_proj` | 3 | bf16 | same | yes | +| `shared_head.norm.weight` | 1 | bf16 | `[7168]` | yes — **a distinct norm**, not `model.norm` | +| `embed_tokens.weight` | 1 | bf16 | `[129280, 7168]` | **no** | +| `shared_head.head.weight` | 1 | bf16 | `[129280, 7168]` | **no** | + +**The last two are bit-identical copies of the trunk's** (`torch.equal` +against `model.embed_tokens.weight` and `lm_head.weight`: both True). The +runtime hands the layer whichever `embed_tokens` and `lm_head` the target's +draft-model container exposes, so pointing them at the trunk's is correct +and saves 1.85 GB per rank. They stay in the weight manifest's *predicted +non-load* set even with MTP on; the other 788 keys flip from non-load to +consumed. + +**The MTP layer's attention geometry is byte-identical to a trunk layer's** +— same `q_a_proj` / `q_b_proj` / `kv_a_proj_with_mqa` / `kv_b_proj` / +`o_proj` shapes, same two LayerNorms, same `q_lora_rank`. Whatever load-time +derivation the trunk's attention needs (absorption operands, row regrouping, +the rope table) applies unchanged. + +## The MTP layer is not quantized, and that is deliberate + +On a quantized export the MTP module can be excluded from quantization +wholesale. Here `hf_quant_config.json` carries `model.layers.61*` as one +wildcard entry in a 63-entry `exclude_modules` list, so **every weight in +the module is bf16** while the trunk's MLP path is NVFP4. + +**Do not re-quantize it at load time to reuse the trunk's expert +vocabulary.** The exclusion is the export's choice about where accuracy is +worth the bytes; quantizing it anyway changes what the checkpoint means, and +the acceptance rate — the only signal that can see the difference — would +absorb the damage silently as a lower draft quality rather than reporting +it. + +The consequence is a vocabulary consequence: the MTP layer's routed experts +need a **bf16** grouped-expert entry, at whatever `(local_experts, hidden, +intermediate)` the parallel split produces, while the trunk's use the NVFP4 +runner. Everything else in the graph — both norms, the concat, the +projection GEMM, the whole MLA block, the router, the shared expert, +`shared_head` — maps onto entries a dense-plus-MoE target already carries. + +## The HBM arithmetic, and the break-even acceptance rate + +A decode step is weight-bandwidth-bound: each rank reads every weight byte +it holds, once. So the cost of drafting is exactly the MTP layer's byte +count, times the number of draft steps. + +Per rank on this checkpoint at `dep4` (64 local experts of 256): + +| | bytes per draft step | +|---|---| +| routed experts, **bf16**: `3 x 2048 x 7168 x 64 x 2` | **5.637 GB** | +| attention, bf16: 187.1 M params x 2 | 374 MB | +| `eh_proj`: 102.8 M x 2 | 206 MB | +| shared expert: `3 x 2048 x 7168 x 2` | 88 MB | +| router | 4 MB | +| **total** | **≈ 6.31 GB** | + +Against a trunk step of ≈ 115 GB per rank (58 NVFP4 MoE layers at 1.585 GB + +61 bf16 attention blocks at 374 MB), that is **+5.5% per draft step**. Note +what the first row means: **the bf16 MTP layer's experts cost 3.56 times +what one NVFP4 trunk MoE layer costs**, purely from the dtype. + +A step now produces `acceptance_length` tokens instead of 1, so drafting +pays for itself when + +``` +acceptance_length >= 1 + max_draft_len * 0.055 +``` + +| `max_draft_len` | extra HBM/step | break-even `acceptance_length` | +|---|---|---| +| 1 | +5.5% | **1.055** | +| 2 | +11.0% | 1.110 | +| 3 | +16.4% | **1.164** | +| 4 | +21.9% | 1.219 | + +**That table is the weight-bytes model, and measurement says it is a lower +bound that stops holding as the batch grows.** Measured on this checkpoint +at `max_draft_len 3`, decomposing each engine step against a non-drafting +one: + +| con | acceptance | step cost — predicted | step cost — **measured** | decode speedup | +|---|---|---|---|---| +| 1 | 3.4904 | 1.164 | **1.511** | 2.311x | +| 32 | 3.3356 | 1.164 | **2.198** | 1.518x | +| 256 | 3.4033 | 1.164 | **2.600** | 1.309x | + +The marginal cost of each successive draft step at con=256 is **+0.645, ++0.569, +0.386**, against con=1's **+0.200, +0.173, +0.137**. The growing +term tracks the **rows** a step carries, not the layer's fixed weight bytes +— it falls off in the same shape the marginal row count does (1→2 rows is ++100%, 2→3 is +50%, 3→4 is +33%). **That is the +memory-bound-to-compute-bound crossover this file calls the genuinely +uncertain end, now measured rather than predicted.** + +So read the weight-bytes table as the floor it is: right at low concurrency, +where a step is bandwidth-bound and drafting really is nearly free, and +roughly 2.2x optimistic at the throughput end. What survives is the +conclusion, because the cost never grows fast enough to catch acceptance: +**at every concurrency measured, acceptance stayed at 2.9x or more of even +the measured break-even**, and no draft length in the certified range lost. + +**The break-even is still very low against the predicted cost** — at +`max_draft_len` 3 it takes only 0.164 extra tokens per step on average — and +two costs that look like they should matter do not: + +* **Collectives.** The MTP layer is an MoE layer, so it adds two per draft + step, and the trunk's grow by `draft_len + 1` in bytes. But at decode + sizes these calls are latency-bound rather than bandwidth-bound (measured + on this target: 2.753 MB in 25.472 µs = 108 GB/s, an order of magnitude + under NVLink), so multiplying the byte count barely moves the time. +* **KV capacity.** One extra layer on a 61-layer pool is **+1.6%**, plus + `max_draft_len - 1` extra tokens per sequence. + +The end that is genuinely uncertain was the **high-concurrency** end — once +a rank's batch already activates all of its local experts and the routed +GEMM is near the HBM roofline, `draft_len + 1`× the rows means the same +weight bytes with several times the arithmetic. It is measured above: the +crossover is real, it costs 2.2x the predicted step cost at con=256, and it +still does not overtake acceptance. + +**Acceptance itself is near-flat in batch size on this mechanism, which is +the other half of why no draft length loses.** Measured across three draft +lengths and nine concurrencies: `max_draft_len` 1 stayed in 1.9446–1.9673 of +a 2.0 ceiling, 2 in 2.6945–2.7707 of 3.0, and 3 in 3.3356–3.4904 of 4.0. +**So a "shorten the draft as the batch grows" schedule has no +acceptance-side motivation here** — any case for one has to be made on cost, +and on this checkpoint the cost side does not invert. (It is also unusable +under attention DP for a structural reason — see +`docs/references/trtllm-runtime-integration.md` §13.) + +## What only the acceptance rate can tell you + +**An MTP layer that computes the wrong thing does not produce wrong +output.** Rejection sampling guarantees the emitted distribution is the +target model's regardless of draft quality. A layer with the concat order +reversed, a norm on the wrong operand, or the expert stack packed in the +wrong interleave produces **bit-correct text, more slowly**, with every +draft rejected. An accuracy benchmark cannot see it. Neither can a smoke +gate. + +`acceptance_length` — the mean tokens emitted per step by requests that +carried a draft, `1.0` meaning total rejection — is the only detector, and +it needs a reference to be read against: + +* **≈ 1.0** at `max_draft_len` 3: the layer is wrong outright. +* **clearly above 1.0 but well below the reference**: the layer is *subtly* + wrong. This is the band the concat order lands in, and nothing else + distinguishes it from "this model just drafts poorly". +* **at the reference**: the layer computes what the checkpoint says. + +The reference is the same checkpoint under stock in-tree modeling, at the +same load and the same speculative configuration. Acceptance is a property +of the model, so a reference forced onto different memory knobs to boot is +still comparable. + +**One trap when choosing the workload, and it does not point the way it +looks like it should.** A serving harness that builds prompts from uniformly +random token ids seems like it must *understate* acceptance — the +continuation of a nonsense prefix is unpredictable, so the drafts should +miss. Measured, it runs the other way, and strongly. + +Only the **prompt** is random. Generation then runs for a fixed output +length with EOS ignored, so what MTP is drafting is the model's own +continuation of a nonsense prefix — which degenerates into repetition, and +repetition is the easiest thing in the world to draft. Measured on this +checkpoint at `max_draft_len 3` (ceiling 4.0): **`acceptance_length` 3.93 at +concurrency 1 and 3.46 at 32**, i.e. 97.6% and 82.1% of proposed draft +tokens accepted. Real text does not do that. + +So the random-prompt harness **flatters** MTP, and two things follow. A +throughput curve measured on it overstates the gain, so it answers "did +anything get slower" and not "is MTP worth enabling" — that one needs a +real-text workload, and an accuracy-gate run already produces one. And a +future run that sees a *low* acceptance number here must **not** excuse it +as "well, the workload is random": on this harness a low number means the +layer is wrong. diff --git a/tensorrt_llm/_torch/staircase/docs/references/trtllm-runtime-integration.md b/tensorrt_llm/_torch/staircase/docs/references/trtllm-runtime-integration.md new file mode 100644 index 000000000000..62603d2f1312 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/docs/references/trtllm-runtime-integration.md @@ -0,0 +1,1053 @@ +# Target integration with the TensorRT-LLM runtime + +How a staircase target plugs into the TensorRT-LLM engine: registration, +construction, weight loading, the per-step forward duties, KV-cache +boundaries, CUDA-graph discipline, and the fail-fast ladder. Together +with the catalog contracts and `thop-attention-step-args.md`, this is the +complete integration knowledge for writing a target — no TensorRT-LLM +source reading is needed. + +Everything below was written against `tensorrt-llm == 1.3.0rc21` and +verified by the qwen3-8b/sm_100/tp1 gates (that target was the readable +exemplar of every pattern named here; it is not one of the two migrated in +this batch). + +> **Written before staircase moved in-tree.** The *runtime* content — the +> engine's construction order, the per-step forward duties, KV-cache +> boundaries, CUDA-graph discipline, the fail-fast ladder, and §13's MTP +> shell specification — is what this file is for and still holds. The +> *packaging* content does not: `uv run`, `bench/`, `model_dir/`, +> `llm_args.yaml`, `scripts/env.sh`, `STAIRCASE_TARGET` and the multi-rank +> registration ritual are all gone. Sections 1, 9 and 12 are marked where +> they describe the old shape; `../../README.md` is the current one. + +## 1. Directory shape + +**Superseded — see `../../README.md`.** In-tree a target is a package under +`models//targets////`; the identity triple is +unchanged, but `llm_args.yaml`, `model_dir/` and `perf/data` are gone (the +topology is the caller's input, the checkpoint is read unpatched). The +pre-move shape, kept because the rest of this file refers to it: + +``` +targets//// # identity = the path triple +├── modeling.py # registration shell + core + step-args + fail-fast +├── weights.py # MANIFEST + load loop + post-load derivations +├── smoke.py # self-locating keyword-assert gate (uv run ) +├── llm_args.yaml # LLM-API overrides; empty {} = trtllm defaults; +│ # parallel topology lives here for multi-rank targets +├── TARGET.md # identity, checkpoint hashes, pins, vocabulary, +│ # verification records, and — added by the tuner — +│ # the Performance section +├── configs/ # tuner knob variants (tracked); added by the perf +│ # campaign, absent on a freshly assembled target +├── perf/ # data/ machine-local and ignored; figures/ tracked +└── model_dir/ # what trtllm loads + ├── config.json # committed; architectures[0] = registered class name; + │ # all other fields = the upstream checkpoint config + └── *.safetensors, tokenizer*, generation_config.json + # machine-local symlinks into the real checkpoint — + # untracked; relink against the TARGET.md hashes +``` + +`model_dir/` is a stub HF checkpoint directory: trtllm resolves the model +class from its `config.json`, reads weights and tokenizer from the +symlinks, and never sees the rest of the target. Tools pass +`model=/model_dir` and read `llm_args.yaml` next to it. + +## 2. Registration and resolution + +- `@register_auto_model("StaircaseForCausalLM")` on the shell class + writes the class into trtllm's process-global registry at import time. + Every target registers the **same** name (identity lives in the + directory path), so one process hosts one target; comparative runs use + separate processes. +- Tools import `modeling.py` **by file path** (the target dir name + contains a hyphen and is not a package): + + ```python + spec = importlib.util.spec_from_file_location( + "staircase_target_modeling", target / "modeling.py") + module = importlib.util.module_from_spec(spec) + sys.modules["staircase_target_modeling"] = module # §12: pickle needs it + spec.loader.exec_module(module) # registration side effect + ``` + + The import must precede engine construction in **every process that + builds the model** — which above world size 1 is not this one. The + `sys.modules` line binds the loaded module to the name the class will + carry in `__module__`; without it a multi-rank run dies in the launcher + before any rank starts. §12 has the mechanism and the rest of what + changes there. +- `modeling.py` bootstraps `sys.path` from its own `__file__` (the repo + root is the target directory's `parents[3]` — the path depth is frozen + by the `targets////` shape) before its + `catalog.*` imports — this is why E402 is disabled for `targets/**` in + pyproject. +- `LLM(model=/model_dir)` then resolves `architectures[0]` + against the registry and construction begins. + +## 3. The shell and the core + +```python +class StaircaseCore(DecoderModel): # the computation + def __init__(self, model_config): ... # geometry asserts + weights + def forward(self, attn_metadata, input_ids=None, position_ids=None, + inputs_embeds=None, lora_params=None, **kwargs): ... + +@register_auto_model("StaircaseForCausalLM") +class StaircaseForCausalLM(DecoderModelForCausalLM[StaircaseCore, ...]): + def __init__(self, model_config): + super().__init__(StaircaseCore(model_config), config=model_config, + hidden_size=..., vocab_size=...) + def load_weights(self, weights, *args, **kwargs): ... + def post_load_weights(self): ... +``` + +- The shell (`DecoderModelForCausalLM`) is composition: it stores the + core as `self.model`, creates `self.lm_head`, and owns the logits path + — its forward calls the core, then gathers exactly the rows that need + logits (last token per context sequence + every generation token) and + applies lm_head. **The core returns final-normed hidden states + `[num_tokens, hidden]` and never touches logits.** +- Inputs are **packed**: first dimension = total tokens of the batch, + context sequences first. `position_ids` arrives as an int32 `[1, T]` + view of an engine buffer — flatten with `reshape(..., [-1])`. +- Construction asserts the target's identity — geometry, dtype, and every + config property the assembly depends on (tie_word_embeddings, qk-norm + presence, rope scaling, sliding window, attention bias). Assert, never + adapt: a mismatched checkpoint is a different target. +- Inputs the target does not implement (`lora_params`, `spec_metadata` + in kwargs) must **assert loudly** — silently ignoring them produces + wrong output with no signal. This is distinct from runtime-owned + feature flags, which pass through (see §6). A target that *does* + implement speculative decoding stops asserting on `spec_metadata` and + grows a second forward it owns; §13 is that axis. + +## 4. The weight lifecycle + +| stage | who | what happens | +|---|---|---| +| t0 construct | engine (`MetaInitMode`) + your `__init__` | `torch.empty` is intercepted to the meta device: parameters are shape/dtype shadows, zero memory. Declare with `torch.empty` only (`zeros`/`full` are NOT intercepted) and do no tensor math in `__init__` (meta tensors raise on compute). | +| t1 materialize | engine | walks **registered** parameters (`nn.Parameter` / `ParameterDict` / module tree — what `named_parameters()` can see) and reallocates each on CUDA, contents garbage. Registration is what requests the allocation; anything held in a plain dict/list stays a dead shadow. | +| t2 read | engine | reads every safetensors in `model_dir/` into a dict `{ckpt_key: tensor}`. At this pin the values arrive **already materialized** as `torch.Tensor`, not as lazy slices — measured at world size 4, where every rank is handed the whole dict (25,687 keys on all four). So the engine pre-shards nothing: a multi-rank split is entirely the manifest's, and a `src` transform's benefit is that only this rank's bytes cross to the device, not that fewer bytes are read off disk. | +| t3 load | your `load_weights(weights)` | full delegation, never validated by the engine — the manifest loop copies bytes into the materialized storage (`param.data[a:b].copy_(src)`). | +| t4 derive | your `post_load_weights()` | the designated home for derived state: `.t()` GEMM views, per-layer tuples, and (future) checkpoint-calibrated call tensors such as fp8 KV scales. Real tensors may be created here — meta is over. | +| t5 serve | your `forward` | reads weights by reference; zero copies or transforms on the hot path. | + +Storage layout: declare parameters in **HF `[out, in]` row-major** so t3 +copies are layout-preserving, and consume GEMMs through `.t()` views +(row-major `[N, K]` transposed is exactly the dense column-major `[K, N]` +that `cublas_mm`'s contract requires; the view is zero-copy and shares +storage, so reloading weights in place keeps views valid). + +## 5. The weight manifest convention + +`weights.py` owns a pure data table: + +```python +MANIFEST[param_key] = [(ckpt_key, dst_slice_or_None), ...] +# tp1: no src transform. Sharded targets add one: +# (ckpt_key, src=col_shard(rank, tp), dst_slice) +# and declare per-rank shapes in modeling. +``` + +Rules: fused parameters (qkv, gate_up) fill by destination row slices — +no intermediate concat buffers; assert shape and dtype per copy; assert +**bidirectional coverage** (manifest keys == declared parameters before +the loop; consumed ckpt keys == the whole checkpoint after it). The +shell-registered exceptions (`lm_head.weight`, and `embed_tokens` when +the base class owns it for tied checkpoints) are fed explicitly. + +## 6. Forward duties, KV cache, and feature flags + +The engine drives everything around the forward: scheduling, KV block +bookkeeping (allocation at admission, per-token growth, prefix reuse), +metadata filling and `prepare()` — all before your forward runs. Your +duties per step: + +1. Build the attention step-args once from the prepared metadata and + share them across layers (`thop-attention-step-args.md` is the + normative mapping). +2. Keep the loop flat: catalog calls, tensor-metadata reads, Python + control flow — nothing else (the closed-vocabulary rule; audit = + collect the calls in `forward` **and the private methods it reaches** + and match them against `catalog/index.yaml`. Scope to that closure: a + whole-file scan flags the rope table's `arange`/`cos`/`sin` and + load-time `.t()`/`.to()`, which run at init and are outside the rule). +3. Let one attention call serve any batch composition + (`attention_input_type=0` for the standard configuration — no phase + branches; MLA targets dispatch per phase by metadata reads). + +The KV-ownership red line does not move: the pool is sized by the engine from the +`config.json` declarations (`num_hidden_layers`, `num_key_value_heads`, +`head_dim`, `torch_dtype`, `vocab_size`, `max_position_embeddings` — +declare them honestly), pages are assigned by the C++ manager, and the +new tokens' K/V are written **inside** the attention op. The target only +forwards the address book. + +Runtime-owned feature flags (block reuse, CUDA graphs, beam, spec-dec +plumbing) **pass through** from metadata exactly as prepared: the target +runs trtllm defaults, feature behavior is upstream's responsibility, and +the gates validate the result. Kernel-surface certification accounting +lives in the catalog contracts, not in target config. + +## 7. CUDA-graph discipline + +Decode-only steps are captured and replayed with **no Python executing**. +Every value the attention call consumes must therefore fall into one of +three classes: + +- **reference**: engine-owned persistent buffers, refreshed in place each + step — replays read fresh contents automatically (all step-args + tensors); +- **GPU-derived**: recomputed by captured kernels on replay; +- **host-derived**: Python scalars frozen at capture — legal only when + they are per-capture constants (`num_contexts == 0` in a decode graph). + +The step-args builder is the single audit point: classify every entry. +Per-forward `empty(...)` output buffers are fine (the caching allocator +serves stable blocks; upstream captures the same pattern). + +## 8. The fail-fast ladder + +| fuse | when | checks | catches | +|---|---|---|---| +| static contract | import | trtllm version == pin; every `torch.ops.trtllm.*` symbol the forward calls; the thop binding | version drift (voids receipts and gate records), missing/renamed ops. torch mirrors (`embedding`, `empty`, `reshape`, ...) are exempt — core PyTorch API, upstream-owned | +| step contract | first forward | metadata is `TrtllmAttentionMetadata`; every `_STEP_FIELDS` name exists; `position_ids` is int32 | private-surface layout drift within a same-version build. Everything checked is fixed at engine construction — once per model instance is sound | +| gates | explicit runs | smoke keywords; accuracy score vs `bench/references/accuracy.yaml` | the computation itself | + +Keep `_STEP_FIELDS` exactly equal to the set of metadata attributes the +code reads, and the static-contract op list exactly equal to the trtllm +ops the forward calls — the lists are dependency declarations. + +**One metadata surface is conditional, and a flat existence check over it +turns a legal config into a hard failure.** MLA's cached-KV fields — +`enable_context_mla_with_cached_kv`, `ctx_cached_token_indptr`, +`ctx_kv_indptr`, `ctx_uncached_token_indptr`, `max_ctx_seq_len`, +`max_ctx_kv_len`, `num_ctx_cached_tokens` — exist only when block reuse +is on. With `enable_block_reuse: false` they are **absent, not False**, so +a `_STEP_FIELDS` tuple containing them fails the first forward on a +configuration the target is supposed to support. Split the tuple: the +unconditional fields keep the existence check, and the conditional ones +select the context flavor by their *presence* rather than being asserted. +See `docs/models/latent-kv-cache.md`. + +## 9. Gates and launch + +**Superseded — see `../../README.md` for the current commands.** The +substance below (what each gate is for, and how to author smoke cases) +carries forward; the invocations do not. + +- `source scripts/env.sh` before any GPU run (exports the + single-process-worker flag; consumed upstream only at world_size == 1, + ignored by multi-rank runs). Claim the devices with + `CUDA_VISIBLE_DEVICES` — one for a `tp1` target, `N` for a `tepN`/`depN` + one. +- Smoke: `uv run targets//smoke.py` — self-locating, boots from its + own directory, greedy keyword asserts, exit code is the verdict. The + cases are target-authored: pick prompts whose greedy continuation is + high-confidence for this model, and **verify every keyword on the real + model before freezing it** — an unverified keyword bakes a false + failure into the gate. +- Release: `uv run bench/accuracy.py --target targets/` — one-sided gate, + measured >= reference − tol; after a first pass on an external anchor, + write the measured score back (see the references file header). + +## 10. What varies per model — the re-derive list + +The exemplar target shows the pattern; every model-specific value must be +re-derived from the new checkpoint's `config.json` and the catalog +contracts, never copied. The known variation axes: + +**attention structure** — GQA/MHA vs **MLA**, the axis that reshapes the +most: MLA replaces per-head K/V with one compressed latent, needs two +attention calls per forward (context and generation are separate call +shapes, mixed batches rejected upstream), and brings its own load-time +obligations. It is not a variant of the geometry row below; see +`docs/references/mla-custom-op-decomposition.md` for the vocabulary and +`docs/models/latent-kv-cache.md` for what a target had to derive. +Then: geometry (layers, heads, kv-heads, head_dim, intermediate, vocab); +qk-norm presence and its eps; rope kind, theta, scaling, `is_neox`, +partial-rotary factor; `tie_word_embeddings` (changes the lm_head/embed +feeding and the shell's tying path); attention bias (needs cublas_mm's +fused-bias argument); sliding window; activation function; checkpoint +key names and fusion grouping; dtype and quantization scheme — including +*which modules it excludes*, since a checkpoint may quantize only its MLP +path and leave attention bf16; parallel topology (llm_args.yaml + +per-rank shapes + manifest src transforms + collective-communication +catalog entries — §12 has the launch substrate, the rest is the first +multi-rank target's to derive); KV-cache dtype (turns the kv-scale +constants into post-load tensors — certification extension first). + +A construction-time assert exists for each axis the assembly depends on: +if the new config violates one, the right response is to re-derive that +part of the assembly, not to delete the assert. + +### Read the axes off the engine's config object, not off AutoConfig + +`model_config.pretrained_config` — the object the engine hands +`__init__` — is the **un-migrated** config: fields sit where the +checkpoint's own `config.json` put them. A standalone +`AutoConfig.from_pretrained(model_dir)` probe can return a *different* +shape of the same information, because transformers migrates fields +across versions. + +Observed at transformers 5.5.4 on a checkpoint written by 4.51: through +AutoConfig, rope had been migrated into +`rope_parameters = {'rope_theta': ..., 'rope_type': 'default'}` with +`rope_scaling` an alias of that same dict, and `cfg.rope_theta` raising +`AttributeError`. Through the engine, `rope_parameters is None`, +`rope_scaling is None`, and `rope_theta` is a plain instance attribute. +Deriving the axis from the AutoConfig surface produced a target that +failed at construction with `TypeError: 'NoneType' object is not +subscriptable`. + +So: an AutoConfig probe is **not** a valid oracle for what +`pretrained_config` will look like. Read the checkpoint's `config.json` +directly to learn what the model *is*, and use the flat +`cfg.` idiom inside the target. If a probe is needed, instrument +the target's own `__init__`. + +### A missing `dtype` is not inherited — declare it in the stub config + +A checkpoint may declare **no** `dtype` and no `torch_dtype` at all +(observed on gpt-oss-120b, whose weights are all BF16 outside the +quantized expert blocks). The two config surfaces then disagree: +`model_config.pretrained_config.torch_dtype` is `None` while the engine's +own `ModelConfig.torch_dtype` resolves `torch.bfloat16`. The shell reads +the **pretrained** one, so `DecoderModelForCausalLM` materializes its +shell-owned parameters — `lm_head.weight` — as **fp32**, and a target with +a dtype-checking weight manifest fails there with a cause two layers +upstream of the symptom. + +The fix belongs in `model_dir/config.json`, which is a stub the target +owns: declare `"dtype": "bfloat16"` alongside the `architectures` patch. +That makes the stub carry two deliberate divergences from the checkpoint's +config rather than one, which is correct — the field states what the +checkpoint *is*, and it is the field the engine sizes the KV pool from. +Assert both surfaces in `__init__` so a future drift is loud. + +### `model_dir/` linking is not a `tokenizer*` glob + +An accuracy protocol that applies a chat template needs the template +files, and a recent checkpoint may keep **no** `chat_template` inside +`tokenizer_config.json` — it lives in `chat_template.jinja` (plus +`chat_template.json`). Those names, and `special_tokens_map.json`, match +no `tokenizer*` pattern. Link by inspecting what the checkpoint actually +ships, not by a fixed glob, or the gate fails at template application +after a full engine build. + +## 11. Serving-time costs a target inherits + +The engine owns these, not the forward, but they land on the target's +Pareto curve and the first two are worth checking on every new target +before any modeling work. + +**Decode CUDA-graph coverage defaults to batch 128.** Above it the whole +decode phase runs eager. The deeper the model the worse this is: measured +on a 48-layer target, concurrency 256 lost 19.7% throughput and 151% TPOT +against its own graph-covered configuration, and raising +`cuda_graph_config.max_batch_size` to 256 with `enable_padding: true` was +worth +56.5% at that point with every other point unchanged. This should +be the first config experiment on any new target — but raise the ceiling +and A/B the padding **separately**: on a later target `enable_padding: +true` measured 11% *worse* at concurrency 256 against the same raised +ceiling, because it swaps the fine `[1..32]` batch grid for a coarse one. + +**`max_seq_len` is a per-step cost, not only an admission cap.** It sizes +`max_blocks_per_seq = max_seq_len / tokens_per_block`, which sizes a +pinned `[1, num_seqs, 2, max_blocks_per_seq]` int32 block-offset staging +buffer that the resource manager allocates, memcpys and pushes H2D **every +step**. Left at a model cap far above the served workload (40960 vs a +1024/1024 benchmark) that is 2.62 MB per step at 256 in flight. Capping it +is a real knob — but size the expectation: a measured 10× shrink moved +`_prepare_inputs` only 5.17 → 4.44 ms/step, so the staging is a minority +of that cost and the residual is `O(num_seqs)` per-request executor +Python. Capping also **restricts capability** (longer requests are +rejected), and the floor is set by the accuracy gate's own prompts, not by +the benchmark — 5-shot MMLU prompts reach 2687 tokens, so a cap at the +benchmark's exact ISL+OSL makes the gate unrunnable. + +## 12. Multi-rank launch — what changes above world size 1 + +Everything here was measured at the pinned version on B200 hosts at world +size 4. Two layers, and they answer different questions: the **launch +substrate** (process model, registration, the failure modes that exist only +above world size 1) applies to every multi-rank segment, while **what a +rank does differently** and **what attention DP changes** are per-topology +and were each established by the first target on that topology. See "Still +not established here" at the end for what no target has reached yet. + +### The process picture + +`LLM(..., tensor_parallel_size=N)` with `N > 1` does not build the model +in the calling process. It spawns one MPI worker per rank +(`MpiPoolSession` → `mpi4py.futures.MPIPoolExecutor` → `MPI_Comm_spawn`) +and the engine is constructed there; the launcher becomes a proxy. The +single-process-worker flag is read only after that branch, so at world +size > 1 it neither helps nor hurts. + +### Registration reaches a rank through exactly one hook + +mpi4py re-imports the launcher's **main module** in every spawned worker, +under `__name__ == "__worker__"`. Measured consequences: + +| | launcher | worker rank | +|---|---|---| +| module-scope code | runs | **runs** | +| `if __name__ == "__main__":` body | runs | does not run | +| `sys.argv` | full | **script path only** | +| `os.environ` | — | **inherited** | + +**This whole problem is gone in-tree, and the table above is now only an +explanation of why the old code looked the way it did.** A target class is +an ordinary member of an installed package, so a spawned rank resolves it +through the normal registry — no module bound into `sys.modules` before +pickling, no `__worker__` re-import to time correctly, and no +`STAIRCASE_TARGET`, which existed solely because argv does not reach the +workers and `LlmArgs` does. `bench/register.py` and its module-scope +`from_env()` were deleted rather than migrated. + +What the table still explains: any *tool* that has to influence worker +ranks must do it through `LlmArgs` or the environment, never argv. + +### Two failure modes that exist only above world size 1 + +**The loaded module must be bound in `sys.modules`.** The engine resolves +`architectures[0]` to a class object in the launcher and ships it to the +ranks by pickle, which stores a class *by reference* — `__module__` plus +`__qualname__` — and both sides resolve those names through +`sys.modules`. Loading by path without binding the name fails in the +launcher, before any rank starts: + +``` +PicklingError: Can't pickle : + import of module 'staircase_target_modeling' failed +``` + +Binding it *twice* is equally fatal — the second load leaves a different +class object under the same name and pickle's identity check rejects it +(`it's not the same object as ...`) — so the load has to be idempotent. + +**An environment variable set after MPI initializes never reaches a +rank.** Importing `tensorrt_llm` initializes MPI, and OpenMPI hands a +spawned process the environment as it stood at that moment. A tool that +exports into its own environment must do so *before* that import; one +that hands a freshly spawned server a prepared environment satisfies this +by construction. + +### The parallel segment and its knobs + +| segment | `llm_args.yaml` | +|---|---| +| `tp` | `tensor_parallel_size: N` | +| `tep` | the above plus `moe_expert_parallel_size: N` | +| `dep` | the above plus `enable_attention_dp: true` | + +`tep4` and `dep4` were both observed building a serving engine and +generating correct greedy text from the stock DeepSeek-V3-Lite NVFP4 +checkpoint (254 s and 142 s to a ready engine). That is a statement about +the LLM API on this host — not about any staircase target. + +### What the engine hands the model + +`model_config.mapping` carries the rank's place in the topology. Fields +observed on `Mapping(world_size=4, tp_size=4, moe_ep_size=4, rank=1)`: + +| field | value | +|---|---| +| `rank` / `world_size` | 1 / 4 | +| `tp_size` / `tp_rank` | 4 / 1 | +| `pp_size` | 1 | +| `moe_ep_size` / `moe_ep_rank` | 4 / 1 | +| `moe_tp_size` / `moe_tp_rank` | 1 / 0 | +| `enable_attention_dp` | False | +| `tp_group` | `[0, 1, 2, 3]` | + +Read the topology off this object, never off the segment string: the +segment names the intent, `mapping` is what the engine actually built. + +### What a rank actually has to do differently + +Established by the first multi-rank target (`tep4`, world size 4). These +are the four things this section used to list as unobserved. + +**`lm_head` is topology-aware, and that is a manifest obligation §4/§5 do +not mention.** At `tp_size = 4` `DecoderModelForCausalLM` builds an +`LMHead` of `[vocab/tp, hidden]` — vocab-parallel — so the manifest must +feed *this rank's contiguous block of vocabulary rows*, and the logits +gather stays upstream's. Consequence: `vocab % tp_size == 0` belongs in +the construction asserts, because nothing else checks it. + +**Where the collectives belong: after every row/column-sharded producer, +and nowhere else.** For a TP attention + EP MoE layer that is two per +layer — after `o_proj`, and after `routed_window + shared_partial` +(summed locally first, so one collective serves the whole MLP). The +entry is `comm/allreduce`; the workspace-free certified path is +`strategy=0` (NCCL) or `8`, `workspace=None`. Keeping `op=0` (plain sum) +and leaving the existing fused-add-rmsnorm in place preserves the +single-rank residual structure, which is also the CUDA-graph-friendly +shape. + +**Which transport, though, is worth about 2x at decode sizes, and the +workspace-free default is the slow one.** A decode message is three orders +of magnitude below the point where NCCL's ring algorithm starts to pay for +itself, and a blocking collective serializes far more than its own time. +`docs/references/collective-allreduce-transport.md` has the algorithms, +the measured crossover, and the two profiling traps that make this easy to +get wrong. + +**The manifest's `src` transforms: a fourth column, applied before the +relayout.** `(ckpt_key | tuple, src, dst_index, transform)`, with `src` +selecting this rank's slice per key of a multi-key row. Order is +load-bearing: relayout transforms are functions of the *per-rank* row +count, so a transform that ran on the whole tensor and sharded afterwards +produces a different byte order. Column shards are strided views, so +densify before any transform reinterprets bytes. + +**Per-rank parameter shapes** follow the topology mechanically — heads, +dense/shared intermediates and the expert-axis window all divide — with +one trap worth stating: re-derive every kernel alignment rule at the +*per-rank* width rather than inheriting the single-rank conclusion. + +**Coverage asserts change shape.** `leftover == {}` no longer holds: under +EP each rank legitimately leaves the off-window experts' keys unconsumed. +Replace it with an explicitly predicted leftover set computed from the +rank's window, and add a parameter-side assert that the shell-registered +parameters are exactly `{lm_head.weight}`. + +How MLA's latent cache behaves under attention TP is a mechanism fact +rather than a runtime one, and lives in `docs/models/latent-kv-cache.md`: +it is **replicated, not sharded** — attention TP buys zero KV memory on +MLA. + +### What attention data parallelism changes on top of that + +Established by the first `depN` target (`dep4`, world size 4, an MLA + EP +MoE checkpoint). **Three of the four TP bullets above come out +differently**, so read this section as replacing them rather than adding to +them whenever `enable_attention_dp` is on. + +**The split is over requests, not over heads.** Attention is *replicated*: +each rank builds the full head count (32, not `32/tp`), holds the full +attention weights, and serves its own subset of the batch. Consequently +**the post-`o_proj` all-reduce disappears** — a rank's attention output is +already complete for its own tokens. + +**`lm_head` is replicated, not vocab-parallel.** At `tp_size = 4` *without* +attention DP the shell builds an `LMHead` of `[vocab/tp, hidden]`; with it, +the shell builds the full `[vocab, hidden]` on every rank. The manifest +obligation reverses, and `vocab % tp_size == 0` stops being a construction +requirement. + +**The manifest's `src` transform column is a TP artifact.** Nothing outside +the routed expert stacks is sharded, so a `depN` manifest is the tp1 +three-column form with the expert loop windowed. + +**The collectives move to the MoE and change identity.** Not two +all-reduces per layer, but one `comm/allgather` + one `comm/reducescatter` +per **MoE** layer (a dense layer has none). The gather **belongs before the +router GEMM, not merely before the expert call** — see +`docs/models/expert-weight-packing.md` for why the EP tiling invariant +depends on that placement and fails silently otherwise. + +**How a rank learns the other ranks' token counts:** +`attn_metadata.all_rank_num_tokens`, a host int list, identical on every +rank. Every rank pads to `max(...)` so both collectives run in their +uniform (`sizes=None`) form — which is also the only form a CUDA graph can +replay, since `sizes` is a host argument frozen at capture. + +**That padding creates an obligation.** Collectives pair **by position** on +the communicator, and at equal byte counts an ordering divergence is +**silent** — every rank wrong in 98-99% of elements, bitwise reproducibly, +no hang. Padding guarantees equal byte counts, so a `depN` forward must +issue an identical call sequence on every rank and must not wait for a hang +to detect that it did not. Certified in `catalog/comm/allgather.md` and +`catalog/comm/reducescatter.md`. + +**The runtime is not a second party on that communicator.** Traced live: +one NCCL communicator per rank, 14,036 collectives over 242 forwards, +**zero** issued by the runtime. The engine's own attention-DP +synchronization is host-side MPI on a *different* communicator +(`MPIDist.tp_comm`, built by `MPI_Comm_create_group`). So a target's +collectives cannot be mispaired against the engine's. + +**The CUDA-graph gate that looks like a landmine and is not.** +`cuda_graph_runner.py` replays a graph under `enable_attention_dp` only if +*all* ranks are generation-only **and** their batch sizes are exactly +equal; the padding that would force equality is off by default. Measured at +steady-state serving load, it **never binds**: `cudaGraphLaunch` = 1.00 per +rank per decode step, i.e. **100% replay coverage** without +`enable_padding`. It does bind on a draining workload — an accuracy-gate +run observed captures but no replays across 334 forwards, where the batch +shrinks monotonically. Both observations are correct about different loads. +Consequences measured on `dep4`: `cuda_graph_config.enable_padding: true` +bought nothing and cost **-2.55%** at concurrency 256 (it can only coarsen +the grid), and `attention_dp_config.enable_balance: true` cost **-13.10%** +at concurrency 128 with mean TTFT doubling, because its `batching_wait_iters` +hold costs more than the alignment buys once random arrival already +balances the ranks. + +**KV cache.** No rank factor: see `docs/models/latent-kv-cache.md` — on +MLA the per-rank pool is identical across tp1/tepN/depN, and what `dep` +buys is `world_size x` *aggregate* capacity from holding disjoint requests, +not a narrower pool. + +### Still not established here + +Pipeline parallelism (`pp_size > 1`); multi-node, where the loopback +pinning `scripts/env.sh` applies for single-host spawn must be overridden; +and `moe_tp_size > 1`, i.e. splitting experts a second way along the +intermediate dimension. + +## 13. One-engine speculative decoding — what changes when the target drafts + +Everything above describes a target that emits **one** token per generation +sequence per step. A speculative target emits `1 + draft_len`, and the draft +tokens come from a second forward that the target itself owns. This section +is the axis that brings: what the engine looks for on the model object, what +the drafting loop calls, where the ownership line falls, and what a rank has +to do differently inside its forward. + +Read like §12, in two layers. The **binding surface** applies to every +one-engine speculative mode. The **per-step consequences** below it were +established for **MTP-Eagle one-model** (`MTP_EAGLE_ONE_MODEL`) on an MLA + +EP-MoE checkpoint at `dep4`, and are marked where another mode differs. +`docs/models/multi-token-prediction.md` carries the other half — what an MTP +layer computes, which is a checkpoint fact rather than a runtime one. + +### The checkpoint picks the mode; the target does not + +`update_spec_config_from_model_config` runs **before the model is built** +and reads the MTP layer count out of the pretrained config +(`num_nextn_predict_layers`, or `mtp_num_hidden_layers` on Qwen3Next-style +configs; 1 if neither is present). `MTPDecodingConfig`'s defaults are +`use_mtp_vanilla=False` and `mtp_eagle_one_model=True`, so: + +| checkpoint layer count | resulting mode | +|---|---| +| `n == 1` | **`MTP_EAGLE_ONE_MODEL`** — one layer of MTP weights, replayed | +| `n > 1` | `MTP` (vanilla) — one distinct layer per draft position | + +**Under MTP-Eagle, `max_draft_len` is not bounded by the checkpoint's layer +count.** That bound belongs to vanilla MTP. The single MTP layer is replayed +autoregressively `max_draft_len` times, so "how many MTP layers do we turn +on" is the wrong question and "what is `max_draft_len`" is the right one. + +**Spell `max_draft_len` out in every config variant.** Left unset on the +MTP-Eagle path it resolves to **1**, not to anything derived from the +workload. `max_total_draft_tokens` is then mirrored from it (linear tree), +and `tokens_per_gen_step = 1 + max_total_draft_tokens`. + +### The engine finds the drafter through exactly one getattr + +```python +# _torch/pyexecutor/model_engine.py +def _get_spec_worker(self): + return getattr(self.model, 'spec_worker', None) +``` + +That is the whole registration. Everything else the runtime touches on the +model side, it reaches through the worker's own arguments: + +| attribute / callable | type | what reads it | +|---|---|---| +| `model.spec_worker` | `SpecWorkerBase` | the engine's one getattr | +| `model.config` | pretrained config | `update_spec_config_from_loaded_model` (the base shell already provides it) | +| `model.draft_config` | — | read with `getattr(..., None)`; **absent is correct** for a single-checkpoint MTP target | +| `draft_model.mtp_layers` | `nn.ModuleList` | only `[0]` is ever indexed — MTP-Eagle replays one layer | +| `draft_model.embed_tokens` | module | passed to the layer as a kwarg | +| `draft_model.lm_head` | module | passed to `shared_head` | +| `draft_model.model.d2t` | — | read with nested `getattr(..., None)`; **absent is correct** (draft and target share a vocabulary) | +| `mtp_layers[0](...)` | callable | the draft loop, once per draft step | +| `mtp_layers[0].shared_head(h, lm_head, attn_metadata, True)` | method | returns draft logits | + +`draft_model` is a container the target defines and hands to the worker; the +runtime never constructs it and never inspects it beyond the four names +above. + +The layer is called by keyword, with `inputs` splatted in: + +```python +hidden_states = draft_model.mtp_layers[0]( + embed_tokens=draft_model.embed_tokens, + all_rank_num_tokens=, + input_ids=..., position_ids=..., hidden_states=..., + attn_metadata=..., spec_metadata=..., +) +``` + +It returns **one tensor**, `[rows_this_step, hidden]`, unpadded — the same +row count its `input_ids` carried. The loop slices it with its own +`gather_ids` afterwards. (Eagle3 one-model returns a second tensor here; +MTP-Eagle does not.) + +### Where the ownership line falls + +| responsibility | owner | +|---|---| +| accept/reject, rejection sampling, the golden token | runtime | +| KV rewind, `attn_metadata` rewrite between draft steps and its restore | runtime | +| the draft loop, `gather_ids`, position shifting, sampling draft tokens | runtime | +| `runtime_draft_len` scheduling and padding, `(bs, draft_len)` graph capture | runtime | +| the KV pool's extra layer and extra tokens | runtime | +| **the MTP layer's forward** | target | +| **`shared_head`** | target | +| **the `draft_model` container** | target | +| **the shell's speculative branch** | target | +| **loading the MTP layer's weights** | target | +| **attention-DP padding inside the MTP layer** | target | + +### Inheriting the in-tree one-engine shell is not the shortcut it looks like + +`SpecDecOneEngineForCausalLM.__init__` builds its drafter by calling +`get_draft_model(...)`, which is a **module-level function, not a method** — +a subclass has no override point. It dispatches on the config's +`model_type`, and a staircase stub config patches `architectures` only, so +`model_type` still names the upstream family and the call returns +**trtllm's own MTP layer**. Inheriting therefore hands the one computation +this project exists to write to the engine instead. + +This is a statement about what gets constructed, not a rule against +inheriting: a shell already inherits `DecoderModelForCausalLM` from the same +package, and the isolation hook gates *reading* whole-model definitions, not +importing them. Writing the branch out (about 40 lines) keeps the forward +readable end to end, which is what the Vocabulary table and the +closed-vocabulary audit both rest on. + +### The shell's shape, and the four things it has to get right + +```python +class StaircaseForCausalLM(DecoderModelForCausalLM[StaircaseCore, ...]): + def __init__(self, model_config): + ... # unchanged + self.spec_config = getattr(model_config, "spec_config", None) + self.draft_model = None + self.spec_worker = None + if self.spec_config is not None: + assert self.spec_config.spec_dec_mode.is_mtp_eagle_one_model() + self.draft_model = (...) + self.spec_worker = get_spec_worker(self.spec_config, model_config, + model_config.mapping) + + def forward(self, attn_metadata, **kw): + if self.spec_worker is None: + assert kw.get("spec_metadata") is None + return super().forward(attn_metadata, **kw) # bit-identical + spec_metadata = kw["spec_metadata"] + hidden = self.model(attn_metadata=attn_metadata, **kw) + logits = self.logits_processor.forward( + hidden[spec_metadata.gather_ids], self.lm_head, attn_metadata, True) + return self.spec_worker( + input_ids=kw["input_ids"], position_ids=kw["position_ids"], + hidden_states=hidden, logits=logits, + attn_metadata=attn_metadata, spec_metadata=spec_metadata, + draft_model=self.draft_model, + resource_manager=kw.get("resource_manager")) +``` + +`get_spec_worker` is imported from `tensorrt_llm._torch.speculative` — the +**runtime**, the part of the stack this project reuses, not modeling. + +Four things this shape is load-bearing about: + +**The non-speculative path must be a delegation, not a reimplementation.** +A target's release criterion was measured on the inherited base forward; the +only way to keep it bit-identical is to call it. Declaring +`resource_manager` as a named parameter would silently drop it from `**kw`, +so leave it in the dict and pull it with `.get` in the speculative branch +only — then the base receives exactly what it receives today. + +**The shell gathers the logits; the engine does not.** For every one-model +mode `without_logits` is True, so `_forward_step` returns the model's dict +verbatim and applies no second gather. `spec_metadata.gather_ids` holds one +row per context request (its last token) and `runtime_draft_len + 1` rows +per generation request — pass `hidden` **ungathered** to the worker and the +gathered logits alongside it. + +**`position_ids` reaches the worker in the engine's `[1, T]` shape.** The +worker does `position_ids.squeeze(0)` itself. A shell that flattens before +handing it over produces a silently wrong draft position sequence. + +**Nothing needs adding to `epilogue`, and nothing needs a `layer_idx`.** +`epilogue` is only consulted by `__pp_init__`'s `skip_forward`, so at +`pp_size == 1` there is nothing to register. And on this mode +`Eagle3OneModelSpecMetadata` sets `layers_to_capture = ()`, which makes +`is_layer_capture()` False at every layer and leaves `hidden_states` +unallocated — the trunk owes the runtime **no hidden-state capture hook** +(that is Eagle3's requirement, not MTP-Eagle's). Measured: nothing in +`_torch/pyexecutor/` or `_torch/speculative/` reads `model.layer_idx`. + +### Draft length is per iteration, not per request + +`_handle_dynamic_draft_len` runs **before** `prepare_resources`, so KV +allocation already knows the answer: + +1. `draft_len_schedule` maps a batch-size threshold to a draft length. +2. The current `scheduled_batch.batch_size` selects `runtime_draft_len`. +3. Every generation request's `py_draft_tokens` is padded or truncated to + **exactly** that length — the source comment names CUDA-graph replay and + the attention kernel as the reasons. +4. It lands on `spec_metadata.runtime_draft_len`; + `runtime_tokens_per_gen_step = 1 + runtime_draft_len`. +5. `runtime_draft_len == 0` takes `skip_drafting`, i.e. speculation is off + for that iteration only. + +**For modeling this means there is no ragged draft tree to handle.** The +draft length is a host int, constant across the batch, constant within a +capture. `cuda_graph_runner.get_graph_key` asserts it directly: *"All draft +lengths must be the same"*. + +Without a schedule, `runtime_draft_len` is simply `max_draft_len` every +step. + +**`draft_len_schedule` deadlocks under attention data parallelism, and +nothing rejects the combination.** Step 2 above reads +`scheduled_batch.batch_size` — **each rank's own local batch** — with no +cross-rank reduction, and attention DP does not equalize batch sizes: +`_pad_attention_dp_dummy_request` only tops a rank up from zero to one, and +`attention_dp_config.enable_balance` is off unless asked for. Two ranks +either side of a schedule threshold therefore run **different numbers of +draft replays**, hence different numbers of MoE collectives — and the +collectives pair by position, so the job hangs. + +Measured at world size 4: boot, CUDA-graph capture and warmup all pass, and +the hang needs real traffic to make the rank batch sizes diverge. trtllm's +own `HangDetector` fired at 300 s and hard-killed all four ranks, whose +stacks sat at three different points of one forward. **On an attention-DP +target, skip this knob.** With `tp`/`ep` alone the ranks share one batch and +the mechanism is sound. + +Setting it also **silently turns on `cuda_graph_config.enable_padding`** — +logged at INFO only — so it is never a single-variable experiment. + +### What changes in the trunk's own forward + +**One argument, and it is the whole of it on Blackwell.** A generation +request arrives with `runtime_draft_len + 1` query tokens instead of 1, and +that is expressed to the attention op through **`predicted_tokens_per_seq`** +alone — the value the total-token arithmetic uses +(`num_ctx_tokens + (num_seqs - num_contexts) * predicted_tokens_per_seq`). + +**Those extra query tokens stay on the generation call.** A one-engine mode +returns False from the runtime's `extend_ctx` predicate — "1-model has +separate logic for handling draft tokens" — so a generation request carrying +drafts is *not* re-shaped into a chunked context request the way two-model +speculation does it. The batch keeps its `[context | generation]` split and +the generation call simply gets a taller query block. What that block is +allowed to attend to — each draft position seeing the cache plus the earlier +positions of its own block, and no later one — is a **kernel** fact, so the +catalog contract for the attention entry is its authority, not this file. +Do not assume it from `mask_type` alone. + +**The spec-dec mask machinery is forced off at sm_100 and stays inert.** + +```python +# _torch/attention/backends/trtllm.py (was _torch/attention_backend/) +# Blackwell trtllm-gen spec-dec is enabled only for dynamic-tree masks. +self.is_spec_decoding_enabled = is_spec_decoding_enabled and ( + not self.is_sm_version_trtllm_gen_kernel(sm=get_sm_version()) + or is_spec_dec_dynamic_tree) +``` + +`is_sm_version_trtllm_gen_kernel(sm)` is `not (sm < 100 or sm in [120, 121])`, +so it is True on sm_100; a linear-tree MTP has `is_spec_dec_dynamic_tree` +False; the conjunction is **False**. `is_spec_decoding_enabled`, +`use_spec_decoding` and `is_spec_dec_tree` are all False and every +`spec_decoding_*` tensor is None — the same inert values the MLA columns of +`catalog/attention/thop_attention.md` already certify. **On a pre-Blackwell +arch this is not true** and the mask surface would need certifying first. + +The context path is unchanged: a context request still contributes its +prompt, one row of logits, and one attention call of the same shape. + +### The draft loop rewrites the batch between step 0 and step 1+ + +The single most important fact for writing the layer. After the first draft +step the worker mutates `attn_metadata` in place: + +| | step 0 | step 1+ | +|---|---|---| +| `_seq_lens` / `_seq_lens_cuda` | real | **filled with 1** | +| `num_contexts` | real | **0** (when a KV cache manager is present) | +| `num_ctx_tokens` | real | **0** (recomputed) | +| `host_request_types` | mixed | context entries overwritten to generation | +| tokens per generation request | `runtime_draft_len + 1` | **1** | +| a context request | the whole prompt | **1 token** | +| `use_spec_decoding` | as the engine set it | False | +| `kv_lens_cuda` | as the engine set it | rewound, then `+1` per step | + +**So the layer must read its phase from `attn_metadata` on every call.** It +is invoked N times inside one forward, and a value computed on the first +call is wrong on the rest. This is not the trunk's situation — the trunk +builds step-args once per forward because there is only one step in it. + +The read is cheap and sync-free. `on_update()` recomputes `_num_ctx_tokens`, +`_num_generations` and `_num_tokens` from `_seq_lens`, which is a **pinned +host** tensor, and the loop calls it (and triggers it again through the +`num_contexts` setter) at exactly the step-0/step-1+ boundary. Therefore + +``` +tokens_per_gen_seq = (rows - md.num_ctx_tokens) // (md.num_seqs - md.num_contexts) +``` + +evaluates to `runtime_draft_len + 1` on step 0 and to exactly `1` on step +1+, with no step counter threaded through and no device read. That is the +`predicted_tokens_per_seq` the layer's own generation call needs. Guard the +all-context case (`num_seqs == num_contexts`) rather than dividing by zero. + +**What a context request is fed on step 0** is `prompt[1:]` with the +request's first accepted (golden) token written at its last position — +`_prepare_context_input_ids`, shared by both MTP flavours. Its hidden states +are the trunk's, in full, at full prompt length. So step 0 runs a real +context-phase attention for those rows and the layer needs both phases, +exactly like the trunk. + +### Attention DP: the padding basis changes source, and reading the wrong one is silent + +Under `enable_attention_dp`, §12 established that every rank pads its token +block to `max(attn_metadata.all_rank_num_tokens)` so both collectives run in +their uniform form. **Inside the draft loop that list is the wrong one from +step 1 onwards.** + +The worker passes the right basis in **as a keyword argument** and leaves +`attn_metadata.all_rank_num_tokens` holding the trunk's value for the whole +loop (it saves and restores it around the loop, but does not maintain it +during it): + +| draft step | `all_rank_num_tokens` kwarg | +|---|---| +| 0 | `spec_metadata.all_rank_num_tokens` — the trunk's token counts | +| 1+ | `spec_metadata.subseq_all_rank_num_tokens` — the per-rank **sequence** counts | + +`subseq_all_rank_num_tokens` is set to `all_rank_num_seqs` by the engine for +the one-model modes, which is semantically right: from step 1 every sequence +contributes exactly one token. + +**Pad from `md.all_rank_num_tokens` inside the MTP layer and the collectives +mispair silently.** `catalog/comm/allgather.md` and +`catalog/comm/reducescatter.md` both certify that calls pair by *position* +on the communicator and that at equal byte counts a divergence produces no +hang — every rank wrong in 98–99% of elements, bitwise reproducibly. The +defence is structural: give the layer's padding helper a signature that +**only accepts a passed-in list**, so it has no way to reach the metadata. + +(A `_dp_rows` that asserts `all_rank_num_tokens[rank] == rows` would in fact +fire here rather than corrupt, since the two counts differ from step 1. Do +not rely on it: the assert is a property of one target's helper, not of the +rule, and it is exactly the kind of guard a later edit removes.) + +### KV cache: the engine adds the layer, and the target declares nothing + +Under a one-model MTP mode `ModelConfig` raises the pool's layer count +itself: + +```python +num_layers += spec_config.num_nextn_predict_layers +num_attention_layers += spec_config.num_nextn_predict_layers +``` + +so the pool is sized for `num_hidden_layers + n` layers and the MTP layer's +own attention addresses layer index `num_hidden_layers`. **Nothing in the +target's config stub or manifest declares any of this** — but the layer's +attention calls must pass the right `layer_idx`, and the pool-addressing +surface they use has to be certified at the new layer count. + +**A `speculative_config` also raises `max_seq_len`, and not by the amount the +extra-KV-token helper suggests.** Three separate terms are added to the model +engine's `max_seq_len`, in `py_executor_creator`: + +```python +if not disable_overlap_scheduler and spec_config is not None: + max_seq_len += spec_config.tokens_per_gen_step - 1 +if spec_config is not None: + max_seq_len += get_num_extra_kv_tokens(spec_config) # max_draft_len - 1 + max_seq_len += spec_config.tokens_per_gen_step - 1 +``` + +With a linear tree (`tokens_per_gen_step = 1 + max_draft_len`) and the overlap +scheduler at its default (**on**), that is `3 * max_draft_len - 1` — 2, 5 and +8 at `max_draft_len` 1, 2 and 3. Measured: a `163840` model cap becomes +**`163848`** at `max_draft_len: 3`. + +**So read `max_seq_len` off the metadata, never compute it.** A target that +derives anything from the config's own `max_position_embeddings` — a rope +table sized to it, a bound asserted against it — is off by that amount the +moment speculation is switched on, and by a different amount per +`max_draft_len`. The failure is a silent out-of-bounds read on a rope table, +or an assert that fires on a legal config. + +### One more construction-time trap, from §4's weight lifecycle + +**A reference to another module's parameter, captured in `__init__`, stays +bound to the meta-device shadow.** The engine materializes at t1 by +*replacing* tensor objects, not by filling them in place, so a drafter +container that does `self.embed_tokens = core.w["embed"]` at construction +holds a meta tensor forever and fails at the first draft step. §4 already says +an unregistered parameter "stays a dead shadow"; this is the adjacent case — +the parameter *is* registered, on the trunk, and the copy of the reference is +what goes stale. Resolve it lazily (a `property` that reads through to the +trunk on each access) rather than caching it. + +### CUDA graphs: the whole draft loop is inside the capture + +Capture wraps `_forward_step`, which calls the shell's forward, which calls +the worker — so every MTP-layer invocation is captured. The key is +`(batch_size, draft_len, is_first_draft, short_seq_len_mode, +is_all_greedy_sample)`, and the capture set becomes `(bs, draft_len)` pairs +rather than bare batch sizes; with a `draft_len_schedule` the runner also +captures one extra `(max_bs, original_max_draft_len)` graph, whose stated +purpose is to keep a later graph from resizing the shared attention +workspace and invalidating pointers baked into earlier ones. + +Consequences for the layer, all of them §7's discipline applied one level +down: + +* The step-0 / step-1+ divergence is **re-traced per draft step at capture + time** and frozen at each step's position in the loop. Reading the phase + from `attn_metadata` every call is what makes that correct — a cached + first-call decision would be baked in at all N positions. +* Host ints read from `attn_metadata` (`num_contexts`, `num_ctx_tokens`) and + the `all_rank_num_tokens` kwarg are per-capture constants, which is the + host-derived class §7 permits. +* `attn_metadata.padded_num_tokens` is **`None`** unless torch-compile + piecewise CUDA graphs are configured; without a `torch_compile_config` the + padding path is not reachable and the base shell's slice never applies. +* The attention-DP replay gate of §12 is unchanged: all ranks + generation-only and equal batch sizes, or the step runs eager — and eager + is where the layer meets a real mixed batch. + +### The acceptance rate is the only correctness signal + +`stats.specdec_stats.acceptance_length` is computed per iteration over the +generation requests that carried draft tokens: + +``` +acceptance_length = (accepted_draft_tokens + requests_with_draft) / requests_with_draft +``` + +i.e. the mean number of tokens a drafting request produces per step, `1.0` +meaning every draft was rejected. + +**This matters more than it looks.** Rejection sampling guarantees the +output distribution is unchanged, so a *miscomputed* MTP layer does not +produce wrong text — it produces correct text more slowly, with every draft +rejected. An accuracy gate cannot see it. A subtly wrong layer (a +transposed concatenation, a norm applied to the wrong operand) lands +somewhere above 1.0 and well below the reference, which no other measurement +distinguishes from "this model is just hard to draft for". The reference is +the same checkpoint under stock in-tree modeling at the same load and the +same `speculative_config`; acceptance is a property of the model, so a +reference forced onto a different `max_num_tokens` / `max_seq_len` to boot +is still a valid comparison. + +### Still not established here + +Vanilla MTP (`n > 1`, one layer per draft position) and its per-layer +sequential loop; Eagle3 in either form, including the hidden-state capture +hook and `apply_eagle3_fc`; tree drafting of any kind (static or dynamic +`eagle_choices`), which is also the only way the Blackwell spec-dec mask +surface becomes reachable; two-model speculation; and speculative decoding +combined with pipeline parallelism, where `skip_forward` and the `epilogue` +path start to matter. diff --git a/tensorrt_llm/_torch/staircase/explain.py b/tensorrt_llm/_torch/staircase/explain.py new file mode 100644 index 000000000000..054eb8d01d56 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/explain.py @@ -0,0 +1,111 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Say which target a configuration routes to, and why. + + python -m tensorrt_llm._torch.staircase.explain \ + --model /path/to/DeepSeek-R1-0528-NVFP4 --tp 4 --ep 4 --attention-dp + +Prints the routing module's decision tree as it was actually evaluated, one +line per criterion, ending either in the target's class name and directory or +in the criterion that did not match. This is what a forward-reading decision +tree buys that a set of reverse predicates cannot: an answer to "why did I not +get the target I expected". + +``--sm`` defaults to the local device but can be given explicitly, so a +configuration can be explained from a machine that has no GPU. +""" + +from __future__ import annotations + +import argparse +import sys +from typing import Optional, Tuple + +from tensorrt_llm.mapping import Mapping + +from ._router_index import STAIRCASE_ROUTERS, StaircaseContext, Trace, routing_module + + +def _sm(value: Optional[str]) -> Tuple[int, int]: + if value is not None: + major, _, minor = value.partition(".") + return (int(major), int(minor)) + import torch + + assert torch.cuda.is_available(), ( + "no CUDA device visible; pass --sm (e.g. --sm 10.3) to explain a " + "configuration from a host without one" + ) + return torch.cuda.get_device_capability() + + +def build_parser() -> argparse.ArgumentParser: + p = argparse.ArgumentParser( + prog="python -m tensorrt_llm._torch.staircase.explain", + description=__doc__, + formatter_class=argparse.RawDescriptionHelpFormatter, + ) + p.add_argument("--model", required=True, help="checkpoint directory") + p.add_argument("--tp", type=int, default=1, help="tensor_parallel_size") + p.add_argument("--pp", type=int, default=1, help="pipeline_parallel_size") + p.add_argument("--ep", type=int, default=-1, help="moe_expert_parallel_size") + p.add_argument("--moe-tp", type=int, default=-1, help="moe_tensor_parallel_size") + p.add_argument("--attention-dp", action="store_true", help="enable_attention_dp") + p.add_argument("--sm", default=None, help="SM version as major.minor; defaults to this device") + return p + + +def main(argv: Optional[list] = None) -> int: + args = build_parser().parse_args(argv) + + from tensorrt_llm._torch.pyexecutor.config_utils import load_pretrained_config + + pretrained_config = load_pretrained_config(args.model) + world_size = args.tp * args.pp + mapping = Mapping( + world_size=world_size, + tp_size=args.tp, + pp_size=args.pp, + moe_ep_size=args.ep, + moe_tp_size=args.moe_tp, + enable_attention_dp=args.attention_dp, + ) + + ctx = StaircaseContext( + pretrained_config=pretrained_config, + mapping=mapping, + sm=_sm(args.sm), + quant_config=None, + spec_config=None, + is_disagg=False, + ) + + arch = (pretrained_config.architectures or ["(none)"])[0] + routing = routing_module(arch) + if routing is None: + print(f"{arch} -> no staircase routing module") + print(" routed architectures: " + (", ".join(sorted(STAIRCASE_ROUTERS)) or "(none)")) + return 1 + + family = routing.__name__.rpartition(".")[0].rpartition(".")[2] + print(f"{arch} -> models/{family}/routing.py") + + trace = Trace() + target = routing.route(ctx, trace) + for label, value, outcome in trace.steps: + mark = "no match" if outcome is None else ("ok" if outcome is True else f"-> {outcome}") + print(f" {label:<10}{str(value):<52}{mark}") + + if target is None: + print(" => no target") + return 1 + + module = routing.TARGET_MODULES[target] + directory = module.rpartition(".")[0].replace(".", "/") + print(f" => {target}") + print(f" {directory}/") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tensorrt_llm/_torch/staircase/models/__init__.py b/tensorrt_llm/_torch/staircase/models/__init__.py new file mode 100644 index 000000000000..9da2ae758216 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/__init__.py @@ -0,0 +1,8 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""One package per upstream architecture family. + +``models//`` corresponds one-to-one with the built-in zoo's +``modeling_.py``: the same architecture, split into one self-contained +codebase per deployment target instead of one class serving every +checkpoint.""" diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/__init__.py b/tensorrt_llm/_torch/staircase/models/deepseek_v3/__init__.py new file mode 100644 index 000000000000..53e5e688b62f --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/deepseek_v3/__init__.py @@ -0,0 +1,3 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""DeepseekV3ForCausalLM targets.""" diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/routing.py b/tensorrt_llm/_torch/staircase/models/deepseek_v3/routing.py new file mode 100644 index 000000000000..f6c1b26c512b --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/deepseek_v3/routing.py @@ -0,0 +1,76 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Where a DeepseekV3ForCausalLM config lands. Read this file and you know. + +One forward-reading decision tree per architecture family: the criteria are +evaluated in the order a reader would ask them, and every branch that does not +end in a target returns None (in ``auto`` the engine then uses the built-in +DeepseekV3 implementation; in ``require`` it raises, quoting the trace below). +""" + +from __future__ import annotations + +from typing import Optional + +from ..._router_index import NULL_TRACE, StaircaseContext, Trace + +# The one GPU architecture these targets are written for. sm is part of a +# target's identity, not a knob: a different SM is a different target. +_SM = (10, 3) + +# Config-shape fingerprint -> checkpoint identity; see the note in the sibling +# gpt_oss routing module for what shape-sniffing does and does not pin. +# +# (num_hidden_layers, hidden_size, n_routed_experts, q_lora_rank). q_lora_rank +# is in the fingerprint because it changes the attention weight layout -- a +# checkpoint that matched the first three and not this one would be a +# different assembly, not a variant. +_CHECKPOINTS = { + (61, 7168, 256, 1536): "r1_0528_nvfp4", + # (30, 2560, 72, 0): "v3_lite_nvfp4", <- second batch +} + +_TARGETS = { + ("r1_0528_nvfp4", "dep4"): "StaircaseDeepseekR10528Nvfp4Sm103Dep4", +} + +# Synthetic architecture name -> the module whose import registers it. +TARGET_MODULES = { + "StaircaseDeepseekR10528Nvfp4Sm103Dep4": "models.deepseek_v3.targets.r1_0528_nvfp4.sm_103.dep4.modeling", +} + + +def _parallel(m) -> Optional[str]: + """Name the parallel topology, or None if no target implements it. + + This is a weight-layout question, which is why it selects a target rather + than a runtime branch: dep4 replicates attention across ranks and shards + only the experts, while tep4 (not in this batch) shards the attention + heads. The two load different weights into different shapes. + """ + if m.world_size == 4 and m.moe_ep_size == 4 and m.moe_tp_size == 1 and m.enable_attention_dp: + return "dep4" + return None + + +def route(ctx: StaircaseContext, trace: Trace = NULL_TRACE) -> Optional[str]: + c, m = ctx.pretrained_config, ctx.mapping + + if not trace.check("sm", ctx.sm, ctx.sm == _SM): + return None + + shape = (c.num_hidden_layers, c.hidden_size, c.n_routed_experts, c.q_lora_rank) + ckpt = trace.resolve("shape", shape, _CHECKPOINTS.get(shape)) + if ckpt is None: + return None + + parallel = trace.resolve( + "parallel", + f"ws={m.world_size} ep={m.moe_ep_size} " + f"moe_tp={m.moe_tp_size} attention_dp={m.enable_attention_dp}", + _parallel(m), + ) + if parallel is None: + return None + + return _TARGETS.get((ckpt, parallel)) diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/__init__.py b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/__init__.py new file mode 100644 index 000000000000..d7eb0af07fc9 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/__init__.py @@ -0,0 +1,3 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Targets, keyed by the // identity path.""" diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/__init__.py b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/__init__.py new file mode 100644 index 000000000000..c2bed1f3a52a --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/__init__.py @@ -0,0 +1,3 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""DeepSeek-R1-0528 NVFP4 targets.""" diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/__init__.py b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/__init__.py new file mode 100644 index 000000000000..cc5208da61da --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/__init__.py @@ -0,0 +1,3 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""DeepSeek-R1-0528-NVFP4 on sm_103 (GB300).""" diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/TARGET.md b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/TARGET.md new file mode 100644 index 000000000000..ce5def4aa247 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/TARGET.md @@ -0,0 +1,1631 @@ +# Target: deepseek-r1-0528-nvfp4 / sm_103 / dep4 + +## Identity + +| | | +|---|---| +| Checkpoint | DeepSeek-R1-0528, modelopt NVFP4 export (**v1**, `DeepSeek-R1-0528-FP4`): 61 layers, hidden 7168, 128 query heads, `q_lora_rank` 1536, vocab 129280, untied embeddings; MLA attention in bf16; layers 0-2 dense (intermediate 18432), layers 3-60 MoE with 256 routed experts at top-8 (group-limited: `n_group` 8, `topk_group` 4) plus one shared expert (intermediate 2048); YaRN rope, factor 40 over an original 4096-position window; **NVFP4 MLP weights with an fp8-e4m3 KV cache** (`hf_quant_config.json`: `quant_algo: NVFP4`, `kv_cache_quant_algo: FP8`), attention / router / embedding / lm_head bf16. The checkpoint also ships a bf16 MTP module at layer 61 (`num_nextn_predict_layers: 1`), which the identity config does not load and `configs/mtp{1,2,3}.yaml` do | +| GPU arch | sm_103 (GB300) | +| Parallel | dep4 — `tensor_parallel_size: 4` + `moe_expert_parallel_size: 4` + `enable_attention_dp: true`; world size 4, one rank per GPU. The engine builds `tp_size=4`, `moe_ep_size=4`, `moe_tp_size=1`, `pp_size=1`, `enable_attention_dp=True`, and construction asserts every one of those | +| Registered class | `StaircaseDeepseekR10528Nvfp4Sm103Dep4` — a synthetic architecture name no checkpoint declares. `models/deepseek_v3/routing.py` rewrites `DeepseekV3ForCausalLM` into it when the config shape, SM and topology all match; the checkpoint is read unpatched. Per-target names mean one process can hold every target at once | + +> **NO GATE RECORD HOLDS FOR THIS TARGET.** Every result below — boot, +> gsm8k, the MTP acceptance comparison, the whole *Performance* section — +> was measured on **sm_100 (B200)** through the pre-move standalone harness. +> This target is **sm_103 (GB300)**, and certification is per architecture. +> The numbers are kept as provenance, but this target is **ungated** until +> they are reproduced on GB300. Read every "passed" below as "passed, on +> sm_100, before the move". The MTP variant needs all three of its gates +> repeated, acceptance included — see *The MTP variant*. + +Checkpoint sha256 — the checkpoint directory passed to `--model` should +resolve to files with these digests. **Routing does not check them**: it +fingerprints the config's shape +(`num_hidden_layers`, `hidden_size`, `n_routed_experts`, `q_lora_rank`), so +a fine-tune or a re-export of this checkpoint routes here silently. That is +the deliberate trade in `models/deepseek_v3/routing.py`, and it changes what +a gate record means — not "this target passed" but "this modeling code +passed *on the checkpoint with these digests*". Run it on another one and +the result is ungated (recorded at: +`umbriel-b200-027:/home/scratch.trt_llm_data/llm-models/DeepSeek-R1/DeepSeek-R1-0528-FP4`). +The source `config.json` this stub was patched from hashes +`a80de10ccf70e7e98dcc6730b45872298937b232d3e8eb0321cfd2bce4cb40e0`; it is +the one checkpoint file that is copied rather than linked. + +
+169 linked files (163 safetensors shards + index + 5 aux) + +``` +c139e0cd9aa4c418ebd38ebae9aff73b4eaacc7723e25d2490dfe41c02bb8708 configuration_deepseek.py +0ef9febae6b6087f4822b02bc9a1c03a83263dabaa931fb9155d061c8951ca07 generation_config.json +36dc07886afcd1679ebf5328a325f9f9629d7a0ddfa644ee6427b10874b2c295 hf_quant_config.json +991aaf2f2c7a5b21f6c6708bd736285640ebec92790d7882958dcc26b0c9f6f6 model.safetensors.index.json +ecb6f9fc369894346f0511f4074ca75cee5cd5f3b06d02f1ba35fcd39f8e121d tokenizer.json +a58700120f68f96faf27a8921876e5e3dbb95f66355c6b331312354cbbf792f9 tokenizer_config.json +bbf04b49f4cdb3c2cb139d84e9701089cff14d08b95b969379aa5b57a63a9242 model-00001-of-000163.safetensors +ac6cb6c4a30d18ee8d24c5237c1b0f9cb30ab24aa9f4e4758a8beb5125360720 model-00002-of-000163.safetensors +c280f0fc3241c685b033c26c92fd4c3fca34285f288922e3d8119ba7d779aa32 model-00003-of-000163.safetensors +ca5f80f8d62248366271495d0a9a95953dd5fb06237231a3f8465207d475edd9 model-00004-of-000163.safetensors +5a454d4577381ec2fb9511015c9a06024fe592949badd45d36d2efe24dc1a0f2 model-00005-of-000163.safetensors +9bb8873103ba89dd16efad52e7846abe7a0ee40af47d4739f3503a82160b1001 model-00006-of-000163.safetensors +48a5a4d26db3e4a11d1bacdfd5c94e27bfc90b14750c165ee7ebca2d81b924e4 model-00007-of-000163.safetensors +b52d09d1147f1704bb39096221afbd992c6a1fe1114358975ddca79ea214361e model-00008-of-000163.safetensors +0f0274233de53fe22c50262c36d6154579af3cc29b83d3b2d601761b0475c810 model-00009-of-000163.safetensors +2658cef0619d4912da43964d312fa1448eac38f5e44e0313218367b888f8e7ae model-00010-of-000163.safetensors +360189a2c32647c16efbf465a9f2bb3f607e47f0848c6dbfca1eae455d8fa6dd model-00011-of-000163.safetensors +8ae7554ea1a6d65c9e194dd99254bc0797c9fd848a7be23fca94e2defc56ce9d model-00012-of-000163.safetensors +dc6ed60abb1be1223979884405aad63f35e5eec42c28b73773bd19f834c90065 model-00013-of-000163.safetensors +7d359af3e3ccf70794d83235e5d022392cb76ade1a34d4fd8797c588c212056a model-00014-of-000163.safetensors +afbdb1f5da08dc38ecaf900ae8866e4687a3ec9e2b21b16d0cba7fa3d529b74a model-00015-of-000163.safetensors +b0ed219a52ec2d617ecc750963519df23d5e27160621d2fd964b506ba017a5bc model-00016-of-000163.safetensors +bbe5cb864267a27a55add115d27f9fdec6d47405f626de9774e0f79ed271de8e model-00017-of-000163.safetensors +ca0e49de986c19b866bf2c330ca91f8c0a092ee56e19ce832267f375174cc496 model-00018-of-000163.safetensors +0b1db3ad16aea8afdcc667011dabc11addf41c20b4c10234508117578278ba21 model-00019-of-000163.safetensors +c016097de09ca54826380158d531ef7121ccc6f5aa44dbbe255dd4c80e817e7d model-00020-of-000163.safetensors +4929c6acd11fa3f87465aef13e5ce7e3d24c83a1babac6d1c99c0201bba7985b model-00021-of-000163.safetensors +05f55fd3fa47fcebc0c4910204672fb602b30f364cf088544c6ec1d67f7d33ba model-00022-of-000163.safetensors +8f861d34dfd9299c043ba0c351505bef3a107d8f95ea8a46ceda8a3ebf6f2625 model-00023-of-000163.safetensors +95362bd84e97a23460933f630b69d9ff720b7d1da15c544b97ade3481bcca444 model-00024-of-000163.safetensors +c1db194b1d6a0f2c6a1935e5105622174007db387c79e6c1b5430bc3273c778b model-00025-of-000163.safetensors +d181f80f2aaf9cb7bdf68e421b950d64a583a52420ce5fd0d3b411ed33c3ea66 model-00026-of-000163.safetensors +04ee66ce0f3e51f407d5a68e4b29c0b63e9bda0a5e88896c39a097aaba3a2dd2 model-00027-of-000163.safetensors +9167d91670023079e629de55d02be50a06a720b125da70b9556481b9f23e49d3 model-00028-of-000163.safetensors +ecf62e772486af899afb440540803550480343b66e5dc096e1e4a8dea6a44154 model-00029-of-000163.safetensors +18dfa7d3ed08561b4109e83e49fe4a073ce5b1a23d374bce98dc5bab832f4948 model-00030-of-000163.safetensors +70e471763de9f5ad8b12cc26be6f44d48c8ebcbb7f1c4a57f4eaabf112f7f57f model-00031-of-000163.safetensors +fd86c19b1c61cd171d5f573cd62a85d09baae34b54b678ceed921c522412a406 model-00032-of-000163.safetensors +e4a2457ca0b979830ba5bdff1b9e17be8ec227c866a8adeeb455abc87a780761 model-00033-of-000163.safetensors +a69fad78c8fce390011d98b140789010e2fcb83129e7ca17ce1bce64bfd2b759 model-00034-of-000163.safetensors +a3898dfd91f59bec65c82ae04b707575501e32f43d2f1dfb26174912f01d0fd5 model-00035-of-000163.safetensors +2e0fe82025efc32f27a290438b2e607c89ac394716e97b3c9b6cfa27e2480623 model-00036-of-000163.safetensors +77afb4330256683ef7962bf43cae8762f80fb0edf889be55073f421c3ad24930 model-00037-of-000163.safetensors +8e2498c311a925d54212cee7d41ce4f98c345d788adcddcc749053292511e4bd model-00038-of-000163.safetensors +0ee6089eb7bfef59dde3b154b1e133f0034347b391b7ed3c903d6e6525b20ea9 model-00039-of-000163.safetensors +54bee64d4e307768da9a6bf437e00eec91e3a4986a2358d2d5eb41d915f4dd0e model-00040-of-000163.safetensors +e8fc96ae38097b90d1b4424f7610186ff7102cfa67930823a5500906f0398b7f model-00041-of-000163.safetensors +2adea35e960e5e35f95c62c0bc97d883248aa79d186da815535336fb0ea914c3 model-00042-of-000163.safetensors +bad155fdce361b40b328edbe35e3995cf6e149e1f798eba415db58ff998880fb model-00043-of-000163.safetensors +17e353f93d66c07708586a8374012a8cb68939a831931bb215508553adcdbeb4 model-00044-of-000163.safetensors +215c3ee090030d9f6a12cac238e1cdc14c1bf7d33a45eb5a84114986a65e988c model-00045-of-000163.safetensors +ba7377f46d611140338dd6efde6ffbb47fce763a70ffe2d0746d5175ad799eea model-00046-of-000163.safetensors +d1ba05c9c00dffcfe86e08ff0109d83d68694179ef45dbcc3766c32ab0316922 model-00047-of-000163.safetensors +171a534fa7797a90cfcfc5b340a25013359750ef37f637b89d7ce5665b13d191 model-00048-of-000163.safetensors +d4c05d603481809dd5a6910461e08ed2387c6b7459438e93f201d00c23222690 model-00049-of-000163.safetensors +558e86f0977914e39665e2a4756082cc5a9ea1d2d99ef4a863c69adc84febf66 model-00050-of-000163.safetensors +5e07cadc4d02a7783e0ea6edfeeff88bbcf02b78679ad4efe32592c54d154208 model-00051-of-000163.safetensors +83c7f0c70e1bb1b0ff6e0915f40d1ca91d98c7c63aebeae50265fa214b8005ad model-00052-of-000163.safetensors +96fa3cc110f476e1414000da8431ddac767c694143a20f07cfb14de62bdbe419 model-00053-of-000163.safetensors +7b170e72c34df97980ed4988a261f1732f2b573244f3345456fefde9874609bb model-00054-of-000163.safetensors +a5b7b9001f1ad50039e63989eb90be1f1e5d9dc0654544a87073627b73289eae model-00055-of-000163.safetensors +98d711b4214ef67598ac60887c12b1f2186e7aa3ec97d95858389c799f8b46f7 model-00056-of-000163.safetensors +3ec5916a49938b634be79548cbe6944660fdcacbdb255a6393932fd404e796ee model-00057-of-000163.safetensors +0d60c7c3f78cd237164c0ed466bed5e73e53328a1db0605aabfa430920210003 model-00058-of-000163.safetensors +0eb7d1c45c8444dbcfc152b44280357543fc1b798f9dc38cd7cb555204f55d99 model-00059-of-000163.safetensors +92bc2a13fd190ecbc25dfab34ceb210026a6674b984609e5642c6c152bb3c6ba model-00060-of-000163.safetensors +4d00fd183118a24cbd5ebb357238505fc252b117e0cbe4e469605639d10b0b3f model-00061-of-000163.safetensors +14ac6babc7b5ec03944a0253aa3276d114e07e36221d797fca50292ab4f371da model-00062-of-000163.safetensors +5d3a9d7c49fe341f8b7f5e21c3ae63a0718f1a3227b18dac393e2565d3154e09 model-00063-of-000163.safetensors +49a1ddf499f4129111ffb50edac498fc23312f98ebc1b69e8841e853ac609355 model-00064-of-000163.safetensors +5341f93faa5c992d684f234c3a479aab87eb3b4576b1f978eb119476131a8ef2 model-00065-of-000163.safetensors +72d183baf3c7c2c7b816e286867e4f1e8398685b6c76110f778dbda4d85ce6fc model-00066-of-000163.safetensors +5ec23f22dcae3d5d883dc925ce70d463a6c9bba3c1f9b5150d447ccfee451e77 model-00067-of-000163.safetensors +205d88536e4a52d83ebaa62005665d531e6f0f8440c280772e7ae61cad23d89f model-00068-of-000163.safetensors +dde4b037231af225ed529148ce14e6dcf10ea3bcb5f79c47aa5c7fde35770d9e model-00069-of-000163.safetensors +7f7ee508718627657ef87966a885a9956f3734674cb988fcb0bbb9e22de005d1 model-00070-of-000163.safetensors +1824edecf1bbce9246038693adf729d0484a7148aff8ce7c4c2f302184874cf1 model-00071-of-000163.safetensors +633dcca915de5aa3629f2b53345a86c5eb4b4bcd4731aa3fbf5e3b5ee4c25bef model-00072-of-000163.safetensors +b53bc434459993f27ce202bd913640f64129c6d91cea6445fddaf6869ea3e618 model-00073-of-000163.safetensors +86680218bc3719271bc7c47f74810b469fbdeb7cb4370433121cc2c8e1416dca model-00074-of-000163.safetensors +0eac1ff5b435996ca17677429909307081eced38c7aa8543f54569e509a18b7c model-00075-of-000163.safetensors +6b49ec351aa02e4d152f95d2d7fb25fefd537a48ab12941ef1ede7b8f2879601 model-00076-of-000163.safetensors +36d338fba2c4a1543a691c72ed2297aad49c46230219c3e570c41a30b6430671 model-00077-of-000163.safetensors +1611afd72d0e702e4edd0e519525f0fd70763039b18db539d1d1a465d1262732 model-00078-of-000163.safetensors +c2cbd1c7ac0762ae69a453fdca414ff3b41d6a831e36f33a0e365eeb3a4f7796 model-00079-of-000163.safetensors +5968523bd66250a4250e41fa01fd1868d558efc4df271d17e44d8a3931bb9f4f model-00080-of-000163.safetensors +9f1bf3ae7b459e3471ead763cacddaf77b030c1bd153ae7c6c330fbe9309a75a model-00081-of-000163.safetensors +b6e47522fcbcd15ee2d4c8a7921954d48403b3daaf45d1fe0b49f3ab6e7c7bd7 model-00082-of-000163.safetensors +a641e0c5932548bc1f93d94bdcf435898656123298425c04c5cfa7f17fc756d3 model-00083-of-000163.safetensors +f43e8dbbfbdbf808a11d8c42668aed50855b668e0f3c1a3cbd007e6b03652e78 model-00084-of-000163.safetensors +894fe387005e9b73fbe8483c5f4948f127527fd90dcf07bfb1dbdedf74357b3e model-00085-of-000163.safetensors +217e2a040495a5a3825c308638806b750508d23ce78f996d53459d83c848d50e model-00086-of-000163.safetensors +9eb75ba7b09bffab0a4d45fe739ab339db47352c227ca0ee15beed80cc1e7da5 model-00087-of-000163.safetensors +f4497ac9f44495704458aa0542395dc29abf97d8a3abbc16e36b4ef522068725 model-00088-of-000163.safetensors +1cde75047119b80d9bb27333d9d52340b56c5c3879cb326f7418dd589dfd73d8 model-00089-of-000163.safetensors +e3b9df1600cabbb33a0963412a8a919fdeb75472938a46346fb9207876940da9 model-00090-of-000163.safetensors +2aa0a431652ac12273e3a0d38870baf9810f5d748be16e1ca37cccf1381a817f model-00091-of-000163.safetensors +a1c78001c5e2ac6ae5a8212e87440091dd56142c0f8829da3c463e0f717725bb model-00092-of-000163.safetensors +9bec57d3138b23e9c9e5dc53e9b6ac63adfb27e35584232b88ab4bcea0653dfd model-00093-of-000163.safetensors +b1bcc0b137739cacff96dd6c2fe0861af0102244be4865ddd50ff9fe71bc5c0a model-00094-of-000163.safetensors +76c031439199fdba5c01bc8ad571a0f8655c3ce1a04bb466338c25ca5256cd6c model-00095-of-000163.safetensors +2f10fce8ee1187960f53cf5bbc5750d4033688f0e6bb24fb92447dd1bffb9ac4 model-00096-of-000163.safetensors +ceebab7ff4022aaa21e15ff05e6a3ebd4a334e10049bf721d3e2519b98514737 model-00097-of-000163.safetensors +87ef9077d9212db15cb8fe89342501c9cea08f3c5c96a31235db683d8751221f model-00098-of-000163.safetensors +cac29cbf8e08520c89bcf27a08a77343a5f6c2aea1771ee7a1bbdf64474e90d1 model-00099-of-000163.safetensors +5ee6438bc3414412a7ab279b3e0b72d631245a7e27975a895d0e94f5aeef90d8 model-00100-of-000163.safetensors +d99c26d1874005fee504b44ac96c496a496ff25cb87ff755832e926541a48492 model-00101-of-000163.safetensors +dff5cc926dc3c1684fbec174508b5ba5abc67c65e1acf6f20354c3acce74b05a model-00102-of-000163.safetensors +babcf9affea89d453f2023379a34b32d0ecc4f9b2787255c0fcb8889cabc093a model-00103-of-000163.safetensors +45e2f2dae475a13c6d10d7a2bfde28b8d65765d421e73d08c527d8c1f5b9392a model-00104-of-000163.safetensors +1b276c9186ee2f335e0930840d406a59776d695433081bbabf4c5ba25e79bd05 model-00105-of-000163.safetensors +a9a5527aede2cb7f292671d86f2864d6fc843396466d16f60e74743d7fc0a6ee model-00106-of-000163.safetensors +a950aedf442134db181e4ac440c9007178c36b0ef57245f7f49a3ee5bf50540f model-00107-of-000163.safetensors +e7a7bf9ae8f6ea8ecc8750cd44ccb3814eafbb2f0b88fc7482e6ee39e9d8dc1b model-00108-of-000163.safetensors +cb9ac1d529250b0524bd86864725144d45d5e91c70e77375ff0bf30434633250 model-00109-of-000163.safetensors +a6999a13f265c55995ba5c1e0260d64469dc714949a4df300b37af6ae9a1e61a model-00110-of-000163.safetensors +7ada750ccc74bcc4c283332ce7ca573c74254ed61931737b1d739e9b633d0e78 model-00111-of-000163.safetensors +3fbfd40aa863cb1ae17eb175958b95e168325c1cae9c8affe639d88c09e1960a model-00112-of-000163.safetensors +61dea7ce4e66bf886e2ede5ff2a9df65660b7492eb820116b4bac9d7419b1e85 model-00113-of-000163.safetensors +12fa6385c3932d28de5300199e53aadacc8a0d3638d03c42f5cbac71f2d8574b model-00114-of-000163.safetensors +22e0a03279a7d2ebe6b0bfb6d5e2ce98b5da39f518cb9a82c04b4d5b52096fa1 model-00115-of-000163.safetensors +c1c9fc85b7cd9d403a2b73098e1d81a10e12fdffd9549c669a2ae515d851052a model-00116-of-000163.safetensors +4dc24b7fcbc6517c62222ca7da18bc664d5f1da5e2efcc8bcb8d5fd590652b78 model-00117-of-000163.safetensors +794d6d03b6447e7f31882ab070bbcede362becd464139cbcca5970d2f41f12ab model-00118-of-000163.safetensors +dae3a9aceb361372feef019635c17a1bda7c826dba8cacbcb3881ea7de2139db model-00119-of-000163.safetensors +93d8435bb4af72314030be201f0f5e8f431fe6e48f79e7918e46d7ef92f2cb59 model-00120-of-000163.safetensors +9aa69352d752be9438a99776732b4e0b62ff589151bcf8d8cfb5e9dfecc45a0a model-00121-of-000163.safetensors +8f7bc0398f6417d9d60567b0687406a951a6190bc8df46be675b2a90c69aa4cf model-00122-of-000163.safetensors +1aa442507c1ee8a26afef5950aa87e6834724231f0c2852b642254c3afb79eaf model-00123-of-000163.safetensors +6e4185179e96b94dbe61de9db3fc68386f524bd5119ef5025c67608f852479be model-00124-of-000163.safetensors +e03dae8708f37980fea4471fe9cf88d9ac703fc2db7b128985fe759cd9922fa8 model-00125-of-000163.safetensors +7c910c972cbfa6f8afcd6cd55e97a90688429037e429b0c67b048fb48bff53c7 model-00126-of-000163.safetensors +bae19ab02d0fb95a32fa08867741e437592c9857e39b4cdb9279da627948a647 model-00127-of-000163.safetensors +fa3c946d0c051bfd44077f497971dc21b93750a3412f541bc34b742a0b504398 model-00128-of-000163.safetensors +8d1afa3ef67315e898c3fcca58e522b278c52ee6db14e3beac8d3f3da9aace23 model-00129-of-000163.safetensors +4a0db8ab07b13372e598790ccb5222e0ebc345e514575169c98ce30bf1b54f8d model-00130-of-000163.safetensors +95e59b883e43b7259e2f061c5ae47071076f8d9ea26d025d12e86c3979f15007 model-00131-of-000163.safetensors +e7310c0a55bd9126ae844d8b1c9074c476d5ff31b9d8902f27114dbde16fce23 model-00132-of-000163.safetensors +184d0abd9be8e5b966cbbf17264b204872450c0fae4c29a82083e39e273432cb model-00133-of-000163.safetensors +4f0d6beb59421c0bc1560495dc89324de1f2e7a794b65ffb854534a21b17b50a model-00134-of-000163.safetensors +ec9fc4a523c9aabe818356a72a8f032540ebd22273a03976bfa8f9af11a5d8b6 model-00135-of-000163.safetensors +9862e7589f2288d86e45d9377203935621b982a0e376101996196fdfc4fd8cad model-00136-of-000163.safetensors +395a89c3c93071004c8d73810f73c23954449acc03efbe220a6b4cf576afdb60 model-00137-of-000163.safetensors +a747eb61204cb7d1f78acfcae182c75781a7839ca185132a0421acd72668eafc model-00138-of-000163.safetensors +27651e88d2fbfc582e0de685ec318c6e9af6e8ad7c13233a2bb3c33ea14ca634 model-00139-of-000163.safetensors +eaa5f027070aed1ad02c4295218c99f3c91c1ad067c4bf60f56ec69219fe0949 model-00140-of-000163.safetensors +6c924b6d4b1fc90bd0aa18a9421125a6c54d317516fe58ad8b8d07d08f12c7b7 model-00141-of-000163.safetensors +4e3cd496cb1af08ea801b831982573883b158bb60a1a9f3576f3986a4cff3467 model-00142-of-000163.safetensors +37716bff9e48a87fd4801e644527b2fd31468c4f6866bda70ff9db409902c593 model-00143-of-000163.safetensors +760f7034122d6ef1738bb74fcf4c9230e5a7ee9024b499931e05760c66aff3ac model-00144-of-000163.safetensors +fa1203747f79017c0324b5564ea70dd611b237928054dbe3a8f9cf4217cea66a model-00145-of-000163.safetensors +c05235067f189ade00c6d0ce7faf2ba728012ed78863352cc9d613da2cf959ba model-00146-of-000163.safetensors +248b214fbd8455664ca1183cbc8f1ccd8f6fc3a9e3b2898798b5b1a716225811 model-00147-of-000163.safetensors +738de1a2b1844982c6f29317e508a9c8b57c218dd4aac058ca113e90f265ba96 model-00148-of-000163.safetensors +c35cd109b7895cf7a74b6f561f47145a99daccdead56fce6b256a0914aceaf25 model-00149-of-000163.safetensors +a08febee45f832a99e4029bc73c0093a40b4f3d7ef2f1690088982ba65fbe00c model-00150-of-000163.safetensors +eb659a6d3d6dcce0c0a051ea5b2184fc7c190c6037859d197637abd6ee305f6c model-00151-of-000163.safetensors +d19f2846145b89a456e71ef2e6e3f00a4ede37c29f8e4ba6cef55159a3230922 model-00152-of-000163.safetensors +801cc06d203ea5e4cbe9097622e7b971174249c5b1b0298b3560cb831946dfb6 model-00153-of-000163.safetensors +42e9ce66a658f447f16c3bca1cb9d843a66df989201ef7fe2696138309d2ef25 model-00154-of-000163.safetensors +ed61b343d9f632a15d1b8b98a103e7ca9077a8e57c92c798f82472d4f7e9776c model-00155-of-000163.safetensors +8d2689d216da51e5a3d507a9a1eec347200fc1146a9c84ebd1e30d9d93ec779a model-00156-of-000163.safetensors +4413b79d84ec94cb627440e096b36012184cce599151355691213348b6e4b795 model-00157-of-000163.safetensors +fe09a5c54db66630ceea47091df108d8235736abe679700786bcf7294f0e9cbc model-00158-of-000163.safetensors +79b6b4bb774569c433cf9b4ac391b4fae11a6ebcb235b213e4c23b2b66d3c477 model-00159-of-000163.safetensors +243cb503cc88e16170bad7a216742da5262032bc640730fc0b3e348a7b14f451 model-00160-of-000163.safetensors +4981af366b2bedb6ec8d6cd2eeae8bd2c027c3ccaf001ab54af207883da0cdfd model-00161-of-000163.safetensors +783d3e45422f562d9dc7b11495dc7a44834f0cfab59e0235432d1188170565f2 model-00162-of-000163.safetensors +8ad1e6011ac0ebd811cbf00edb62ad095be42b01266cf4638492fecc1e9020fa model-00163-of-000163.safetensors +``` + +
+ +`model_dir/config.json` is the checkpoint's own config with two deliberate +divergences: `architectures` patched to `["StaircaseForCausalLM"]`, and a +`dtype: "bfloat16"` field added beside the checkpoint's `torch_dtype` +(transformers 5.x renamed the field, and the shell materializes its own +`lm_head` from whichever surface resolves — both were observed present and +`torch.bfloat16` on this build). `hf_quant_config.json` is linked because the +engine reads it: it sets `quant_algo=NVFP4` **and** +`kv_cache_quant_algo=FP8`, and the target asserts the latter — the fp8 latent +pool is what selects `quant_mode` on every MLA call, and a checkpoint +declaring no KV quantization is a different assembly. Nothing in the stub +carries the topology; that lives in `llm_args.yaml`. + +## Version + +| | | +|---|---| +| tensorrt_llm | in-tree — the target moves with the trunk, so there is no version to pin and none is asserted. What *is* asserted at construction is the SM version (`_SM = (10, 3)`), which the pin used to stand in for. The gate records below name the commit they were taken at | +| torch | 2.11.0+cu130 | +| transformers | 5.5.4 (the config surface the engine hands the target) | +| Attention metadata fact source | `TrtllmAttentionMetadata` (TRTLLM backend) | + +## Vocabulary + +**Two audit roots.** This target carries two forward paths: the trunk's, which +runs on every engine step, and the MTP layer's, which runs `max_draft_len` +times per step under a `configs/mtp*.yaml` variant and not at all without one. +Auditing only the first would leave the second free to use any op unseen. + +### Trunk — the `StaircaseCore.forward` closure + +`forward` plus the private methods it reaches: `_dp_rows`, `_dense_mlp`, +`_check_step_contract`, `_build_step_args`, `_moe_chunk_sizes`, and — through +one first-forward branch only — `_rope_tables`. + +| call | catalog entry | +|---|---| +| `embedding` | `torch/embedding.py` | +| `flashinfer_rmsnorm` | `norm/flashinfer_rmsnorm.py` | +| `flashinfer_fused_add_rmsnorm` | `norm/flashinfer_fused_add_rmsnorm.py` | +| `cublas_mm` | `gemm/cublas_mm.py` | +| `bmm_out` | `gemm/bmm_out.py` | +| `mla_rope_append_paged_kv_assign_q` | `attention/mla_rope_append_paged_kv_assign_q.py` | +| `load_paged_kv_cache_for_mla` | `attention/load_paged_kv_cache_for_mla.py` | +| `mla_rope_generation` | `attention/mla_rope_generation.py` | +| `thop_attention` | `attention/thop_attention.py` | +| `allgather` | `comm/allgather.py` | +| `reducescatter` | `comm/reducescatter.py` | +| `fp4_quantize` | `quantization/fp4_quantize.py` | +| `nvfp4_gemm` | `gemm/nvfp4_gemm.py` | +| `flashinfer_silu_and_mul` | `activation/flashinfer_silu_and_mul.py` | +| `noaux_tc_op` | `moe/noaux_tc_op.py` | +| `fp4_block_scale_moe_runner` | `moe/fp4_block_scale_moe_runner.py` | +| `empty`, `reshape`, `split`, `concat`, `copy_`, `expand`, `transpose`, `view_dtype`, `add`, `pad` | `torch/*.py` | + +`_rope_tables` enters the closure through exactly one branch, taken on the +first forward and never again: when the engine's admitted `max_seq_len` +exceeds the config's `max_position_embeddings` — which a `speculative_config` +causes, measured **163840 -> 163848** at `max_draft_len: 3` — the constant rope +table is rebuilt at the larger row count. Its `torch.arange` / `.cos()` / +`.sin()` / `torch.empty` are host-side table construction, not step math: they +create no activation, run once before any CUDA-graph capture, and every table +row depends only on its own position, so the rows the identity config uses are +bit-identical whether the table was built at 163840 or 163848. Recorded here +rather than left for a closure scan to trip over. + +Everything else in the closure is a builtin or a metadata read (`getattr`, +`hasattr`, `isinstance`, `int`, `bool`, `all`, `len`, `max`, `range`, +`sorted`, `divmod`, `dict`, `kwargs.get`, `list.append`, +`host_kv_cache_pool_mapping.tolist()`). + +### MTP layer — the `MTPLayer.forward` closure, plus `.shared_head` + +`forward` plus `_shared_mlp`, `_routed_experts`, `_mtp_dp_rows`, +`_build_step_args`, `_moe_chunk_sizes`; and `shared_head`, which the runtime +calls separately for the draft logits. + +| call | catalog entry | +|---|---| +| `embedding` | `torch/embedding.py` | +| `flashinfer_rmsnorm` | `norm/flashinfer_rmsnorm.py` | +| `flashinfer_fused_add_rmsnorm` | `norm/flashinfer_fused_add_rmsnorm.py` | +| `cublas_mm` | `gemm/cublas_mm.py` | +| `bmm_out` | `gemm/bmm_out.py` | +| `mla_rope_append_paged_kv_assign_q` | `attention/mla_rope_append_paged_kv_assign_q.py` | +| `load_paged_kv_cache_for_mla` | `attention/load_paged_kv_cache_for_mla.py` | +| `mla_rope_generation` | `attention/mla_rope_generation.py` | +| `thop_attention` | `attention/thop_attention.py` | +| `allgather` | `comm/allgather.py` | +| `reducescatter` | `comm/reducescatter.py` | +| `flashinfer_silu_and_mul` | `activation/flashinfer_silu_and_mul.py` | +| `noaux_tc_op` | `moe/noaux_tc_op.py` | +| `fused_moe` | `moe/fused_moe.py` | +| `empty`, `reshape`, `split`, `concat`, `copy_`, `expand`, `transpose`, `add`, `pad` | `torch/*.py` | + +The two tables differ in exactly one place, and it is a dtype consequence: the +checkpoint excludes `model.layers.61*` from NVFP4 wholesale, so this module is +bf16 throughout. It therefore uses **`moe/fused_moe.py`**, the unquantized +grouped-expert runner, where the trunk uses `fp4_block_scale_moe_runner` — and +correspondingly the trunk's `fp4_quantize`, `nvfp4_gemm` and `view_dtype` do +not appear here at all. Everything else is the same vocabulary. + +`shared_head` adds `flashinfer_rmsnorm` and one `logits_processor.forward`. + +### The shell's speculative branch + +`StaircaseForCausalLM.forward` is a plain delegation to the inherited base +when there is no spec worker. With one it contains exactly one tensor +expression — the row gather of the trunk's hidden states at +`spec_metadata.gather_ids`, which is **`torch/embedding.py`** +(`torch.nn.functional.embedding` is a row lookup, used here as one) — plus a +`logits_processor.forward`. That projection is the inherited shell's own +logits path, the same one the non-speculative forward runs internally, so it +is runtime rather than modeling; it is named here rather than left implicit. + +**No catalog entry was added by this target**, torch mirror included. The +trunk's entry set is exactly the `deepseek-v3-lite-nvfp4/sm_100/dep4` +sibling's; the MTP increment consumed one further **existing** entry, +`moe/fused_moe.py`, which the catalog owner certified at this checkpoint's MTP +routed geometry rather than this target adding anything. + +Load time (outside the closed-vocabulary rule, `weights.py`): +`torch.ops.trtllm.block_scale_interleave` for the 128x4 scale swizzle, plus +`torch.cat` / `torch.index_select` for the `[up; gate]` concat and the +interleave + 32-row block shuffle of the expert stacks, and the `kv_b_proj` +row regrouping. The MTP module adds only plain `torch.cat` (its `[up; gate]` +FC1 stack and its shared-expert `[gate; up]` pair) — no swizzle, no shuffle, +because none of it is quantized. `derive_after_load` builds the YaRN rope +table, the `.t()` GEMM views, the two MLA absorption operands, every NVFP4 +call scalar, and the expert-window slices of the three MoE scale scalars — and +asserts the checkpoint's fp8 KV scales are all exactly 1.0 (122 tensors, 124 +with the MTP module loaded). + +Audit is mechanical: collect the calls in each root and the private methods it +reaches, then match against `catalog/index.yaml`. + +## Verification + +**The gate records below are the identity config's.** A `configs/mtp*.yaml` +variant is a different forward and a different weight load, so nothing here +speaks for it; its own records are in *The MTP variant* at the end of this +section. + +Both gates run under the target's identity config (`llm_args.yaml` = +`tensor_parallel_size: 4` + `moe_expert_parallel_size: 4` + +`enable_attention_dp: true`, everything else trtllm defaults — block reuse +and CUDA graphs on, page size 32), 2026-08-01, **GPUs 3-6** of +`umbriel-b200-027` (8x NVIDIA B200, driver 595.58.03, 224-core), +`source scripts/env.sh` then `CUDA_VISIBLE_DEVICES=3,4,5,6`. + +### Required on sm_103 — not yet run + +Every row needs 4 GB300 GPUs and `trtllm-llmapi-launch` over a 4-task srun +allocation. `CFG=` this target's `configs/` directory: each file there is a +complete `--extra_llm_api_options`, carrying the dep4 topology that selects +this target plus the knobs of its own variant. + +Records below that speak of "smoke" were measured with a per-target +`smoke.py` — a bespoke CLI that asserted a keyword in each of ten greedy +continuations. It was removed once this package moved in-tree: the generic +script in the `boot` row below starts the engine the same way and prints +the same continuations, and +unlike a module nothing ran, the accuracy rows below are wired into CI. + +| Gate | Command | Result | +|---|---|---| +| boot | `TRTLLM_STAIRCASE=require trtllm-llmapi-launch python examples/llm-api/quickstart_advanced.py --model_dir --tp_size 4 --moe_ep_size 4 --enable_attention_dp --max_tokens 16 --prompt "The capital of France is" "The chemical symbol for gold is" "1, 2, 3, 4, 5,"` | **passed, 10/10** greedy keyword asserts, 2026-09-10, 4x GB300 on nvl72d199-T07, trtllm 1.3.0rc26, 5m18s. Run through `trtllm-llmapi-launch` over a 4-task srun allocation; the engine built `tensor_parallel_size=4`, `moe_expert_parallel_size=4`, `enable_attention_dp=True` with `TRTLLM_STAIRCASE=require` exported | + +**The identity assembly and the MTP variant are both gated on sm_103.** Everything up to the +weight load is driven by the checkpoint's config alone, and that part +was exercised on a GB300 against the config *as published* (no +target-owned stub): routing resolved this target, the module imported, +`StaircaseCore.__init__` passed every geometry, topology and dtype +assert, and 1980 parameters declared, 61 layers, hidden 7168. `lm_head.weight` +came out **bfloat16**, which is the specific thing removing the stub +put at risk -- the shell sizes it from the pretrained dtype, and a +regression there materializes fp32 two layers from its cause. + +The weight path is verified by the two gates above, over four ranks: +the manifest load fills every declared parameter, the post-load +derivations run, and the forward is exercised through prefill, the +CUDA-graph decode path, and the expert-parallel MoE round trip. + +The checkpoint they ran on is the one this file records. Every digest +that can be checked was re-verified after download -- `config.json`, +`generation_config.json`, `hf_quant_config.json`, `tokenizer_config.json` +and `tokenizer.json` by hand, the 163 safetensors shards by git-lfs +whose object id *is* the sha256 -- so these records and the sm_100 ones +below were measured on byte-identical weights. + +`configs/mtp3.yaml` carries its own three gate records above -- it selects a +second forward path *and* a second weight-loading path, so the identity +records do not speak for it. `mtp1.yaml` and `mtp2.yaml` remain **ungated**; +they are the dominated low end of the measured draft-length axis, and +`mtp3.yaml` is the one to serve. + +| boot, MTP | as above plus `--spec_decode_algo MTP --spec_decode_max_draft_len 3 --kv_cache_fraction 0.75`, i.e. what `$CFG/mtp3.yaml` declares | **passed, 10/10**, 2026-09-10, 4x GB300, 6m16s. `MTPDecodingConfig(max_draft_len=3)` in the run's LLM Args and 124.65 GiB of weights loaded against the identity path's ~118.8 GiB -- the layer-61 bf16 MTP module, i.e. the variant's second weight-loading path, really ran | +| gsm8k full | `TRTLLM_STAIRCASE=require trtllm-llmapi-launch trtllm-eval --model --extra_llm_api_options $CFG/identity.yaml gsm8k --output_path ` | **passed, 95.0720** (`exact_match,flexible-extract`, +-0.5962, full 1319 questions) against threshold **89.9962** (anchor `deepseek-ai/DeepSeek-R1-0528` = 94.9962, tol 5.0) -- pass by 5.08 points. 2026-09-10, 4x GB300, 6m47s. `strict-match` on the same run: 94.7688 | +| gsm8k, **paired in one session** | `$CFG/identity.yaml` and `$CFG/mtp3.yaml`, back to back in one allocation; in CI as `accuracy/test_staircase.py::TestStaircaseDeepseekR10528Nvfp4Sm103Dep4::test_gsm8k_identity_vs_mtp3` | **passed.** identity **94.9962** (+-0.6005) / strict 94.7688; mtp3 **94.6171** (+-0.6216) / strict 94.4655. **delta = -0.3791 flexible, -0.3033 strict** against a `|delta| < 1.2` criterion (2 sigma at sigma = 0.60) -- 0.63 sigma, five questions of 1319, and both filters move the same way | +| acceptance vs stock | `trtllm-bench throughput` twice in one allocation over one fixed-seed dataset, `TRTLLM_STAIRCASE=require` against `=off`; in CI as the two `accuracy/test_staircase.py::TestStaircaseDeepseekR10528Nvfp4Sm103Dep4::test_mtp3_acceptance[...]` legs | **passed.** `acceptance_length` **3.3514** (ours) against **3.2752** (stock) -- ratio **1.023**; draft acceptance 78.38% against 75.84%. 2026-09-10, 4x GB300, 64 requests at ISL=OSL=1024, concurrency 32 | + +`TRTLLM_STAIRCASE=require` is what makes these gates at all: under `auto` a +configuration that missed this target would measure the built-in DeepseekV3 +implementation and report it as this target's score. The old `--trtllm` flag +of the perf harness, which worked by unsetting an environment variable, is +now just `staircase: off` — one config key selecting between the two +systems. + + + +#### The MTP variant is gated, and why the third gate was the one that mattered + +Rejection sampling holds the emitted distribution to the target model's, so a +*miscomputed* draft layer produces correct text more slowly rather than wrong +text: the boot gate passes, accuracy passes, and only speed moves. `acceptance_length` +is the sole detector, and it is only readable against a reference -- 1.0 would +mean every draft rejected, and "clearly above 1.0 but below the reference" +would mean subtly wrong. + +Measured **3.3514** against stock's **3.2752** on the same 64-request +fixed-seed dataset in the same allocation, a ratio of **1.023**, against a +ceiling of 4.0 at `max_draft_len: 3`. The sm_100 record below measured ratios +of 1.0215 and 1.0175 at concurrency 1 and 32, so this sits in the same family. +The draft path computes what the checkpoint says. + +Two things about the comparison worth stating rather than leaving implicit. +`TRTLLM_CAN_USE_DEEP_EP=0` was exported for **both** sides: stock cannot boot +this checkpoint at dep4 with MTP without it (the MoE communication factory +lands on DeepEPLowLatency, whose dispatch takes only NVFP4 uint8 hidden states, +and the MTP layer is bf16), and it is inert for the staircase target, which +implements the all-gather/reduce-scatter round trip by hand rather than +through that factory. So it makes the two systems more comparable, not less. +And the throughput on the same runs -- 2955.2 against 2732.7 output tok/s -- is +**recorded, not claimed**: it is one concurrency of one synthetic workload, and +this project measures perf rather than gating it. + +#### On the flexible-extract number landing exactly on the sm_100 one + +This run measured **95.0720 (+-0.5962)**, and the sm_100 identity run +recorded below measured **95.0720 (+-0.5962)** -- the same 1254 of 1319 +questions. That is agreement at the score level, not proof of identical +generations: `strict-match` differs on the same pair of runs (94.7688 here +against 94.9204 there), so the underlying text does move, as it should +across two architectures whose MoE epilogues use different block-scale +recipes. Read the flexible-extract match as a strong reproduction, and not +as evidence that the two forwards are bit-identical -- they are not. + +### Prior record — sm_100 (B200), pre-move harness, does not gate this target + +| Gate | Result | +|---|---| +| smoke — `uv run targets/deepseek-r1-0528-nvfp4/sm_100/dep4/smoke.py` | **passed, 10/10** greedy keyword asserts, first run, no iteration. Keywords were authored provisionally and frozen against the continuations this model actually produced | +| gsm8k full — `uv run bench/accuracy.py --target targets/deepseek-r1-0528-nvfp4/sm_100/dep4` | **passed, 94.9962** (`exact_match,flexible-extract`, 1319 questions, 5-shot completion, no chat template, 256 output tokens) against `reference: deepseek-ai/DeepSeek-R1-0528 = 94.24 (trtllm); gate: accuracy >= 89.24`. First run, stock defaults, no instrumentation; TRTLLM execution 37.695 s, engine init 197.663 s | + +Both lm-eval filters of the passing run: `flexible-extract` **94.9962** +(±0.6005) and `strict-match` **94.6171** (±0.6216). The gap is 5 questions of +1319 — under few-shot completion this checkpoint mostly answers in the strict +`#### N` form, and flexible extraction finds a few more. + +**Read the +0.76 over the anchor as agreement, not as an improvement.** The +anchor was measured on a **different NVFP4 export of the same base model** +(`DeepSeek-R1-0528-FP4-v2`, whose modelopt `exclude_modules` list differs +substantially from this v1 export's), which is what `tol: 5.0` absorbs; and +one filter's stderr alone is ±0.60. Nothing at this distance is a result in +either direction. + +The reference entry is still at its external anchor (`source: trtllm`). Its +header asks for a write-back after a target's first passing run; `bench/` is +read-only for the assembler, so that edit is left to whoever owns the file. + +### The MTP variant + +`configs/mtp{1,2,3}.yaml`, 2026-08-02, same host and same **GPUs 3-6**. + +| Gate | Result | +|---|---| +| identity smoke, re-run after the MTP increment — `uv run .../smoke.py` | **passed, 10/10**, and all 10 continuations **byte-identical** to the pre-increment run (`md5` of the captured case lines equal). This is the invariant the increment is held to: no `speculative_config`, no MTP parameter declared, `forward` a plain `super().forward(...)`, so the gsm8k 94.9962 record stands unmoved | +| MTP smoke — `uv run .../smoke.py --config .../configs/mtp3.yaml` | **passed, 10/10**. Engine built at `max_draft_len: 3`, `max_seq_len` 163848, KV pool 32.20 GiB (968,288 tokens, 62 layers), 12 fp8 MLA decode JIT compiles per rank | +| acceptance probe — `uv run bench/perf.py --target ... --label probe-mtp3-acceptance --config .../configs/mtp3.yaml --acceptance --concurrency 1,32 --rounds 2` | **`acceptance_length` 3.9274 at con=1 and 3.4635 at con=32**, against a ceiling of 4.0 at `max_draft_len: 3` and a break-even of 1.164. Draft acceptance rate 97.58% / 82.12% (5849 of 5994 and 39152 of 47679 draft tokens accepted) | + +**Two of the ten MTP continuations differ in wording from the identity run's, +and that is expected rather than a defect.** Rejection sampling makes the +emitted token the *target* model's argmax, so a correct MTP layer cannot +change what is emitted — but the trunk's own generation call now runs four +query rows per sequence instead of one, on a different decode kernel +(`HVPerCta256` rather than `HVPerCta128`) with a different accumulation order. +Greedy decoding is chaotic under a 1-ulp logit difference at a near-tie: cases +2 and 3 diverge a few tokens in and stay grammatical and correct, and all ten +keywords pass. **Byte-identity is the right bar for the identity config and +the wrong one for a variant that changes the attention tile.** + +### The MTP variant's own gate records + +The two gates the increment is released on were run after the assembly, on +the same host and the same **GPUs 3-6**, 2026-08-02. Like everything else in +this file they are **sm_100 records and do not gate the sm_103 target**; the +acceptance row in particular has to be repeated, because it is the *only* +detector of a miscomputed draft path — rejection sampling keeps the emitted +distribution correct, so boot and accuracy both pass while a wrong draft +layer merely runs slower. + +| Gate | Result | +|---|---| +| gsm8k, **paired in one session** — `uv run bench/accuracy.py --target ...` with and without `--config .../configs/mtp3.yaml` | **passed.** identity **95.0720** (±0.5962) / strict 94.9204; mtp3 **95.2237** (±0.5874) / strict 94.8446. **Δ = +0.1517 flexible, −0.0758 strict** against a `\|Δ\| < 1.2` criterion (2σ at σ = 0.6005) — 0.25σ, two questions of 1319, and the two filters move in opposite directions | +| acceptance vs stock trtllm — `uv run bench/perf.py --target ... --label trtllm-mtp3-acceptance --trtllm --config .../configs/trtllm-ref-mtp3.yaml --acceptance --concurrency 1,32 --rounds 2`, with `TRTLLM_CAN_USE_DEEP_EP=0` | **passed.** `acceptance_length` **3.9274 / 3.4635** (ours) against **3.8447 / 3.4040** (stock) at con 1 / 32 — ratios 1.0215 and 1.0175. Draft acceptance 97.58% / 82.12% against 94.82% / 80.13% | + +**The paired accuracy run is the one to quote, not a cross-session +comparison.** The same identity code path measured 94.7688 on 2026-08-01 and +95.0720 on 2026-08-02 — a 0.30 spread on a bit-identical forward, which is +the scale any accuracy claim about this variant has to be read at. Both runs +of the pair were back to back on the same devices. + +**One incidental number worth keeping**, because it is the only *real-text* +measurement of what MTP buys here: gsm8k execution time fell from **37.025 s +to 31.440 s, −15.1% (1.178x)**, on 1319 genuine prompts. It is not a Pareto +point — lm-eval drives its own concurrency rather than a swept one — so it +does not belong on the perf curve, but it is a far better answer to "is MTP +worth enabling" than the flattered random-prompt throughput. Engine init +rose 193.6 s → 278.4 s (layer 61's weights plus the extra decode JIT). + +**Read the ~2% acceptance lead as agreement, not as an improvement.** The two +systems do not run the same knobs: the reference is forced to `max_num_tokens` +/ `max_seq_len` 2048 to boot at all and to a different MoE transport (below), +so scheduling, batching and block reuse differ and each system verifies a +slightly different token stream. What the comparison establishes is that this +target's MTP layer computes what the checkpoint says — a miscomputed one sits +at 1.0, a subtly miscomputed one clearly below the reference. + +**Stock trtllm cannot serve this checkpoint at `dep4` with MTP under its +default MoE communication strategy**, and that is why the reference carries an +environment variable as well as a config. All four ranks die in +`Failed to initialize executor` on +`deep_ep_low_latency.py:238`'s `assert hidden_states.dtype == torch.uint8`: +the communication factory falls through `NVLinkOneSided` / `NVLinkTwoSided` / +`DeepEP` (each `not available: Invalid Argument`) to `DeepEPLowLatency`, whose +dispatch accepts only NVFP4 hidden states — and **the MTP layer is bf16**, +by the export's own `exclude_modules`. The existing `trtllm` Pareto curve went +through DeepEPLowLatency happily because without MTP every MoE layer is NVFP4. +`TRTLLM_CAN_USE_DEEP_EP=0` lands the reference on `AllGatherReduceScatter`, +which is the strategy *this target implements by hand*, so it makes the two +systems more numerically comparable rather than less. +`configs/trtllm-ref-mtp3.yaml` carries the whole ledger. + +Consequently `trtllm-mtp3-acceptance` is an **acceptance measurement only**. +Its throughput column (180.6 tok/s/user at con=1, 1668.4 tok/s at con=32) is +not comparable to the `trtllm` Pareto curve: different boot knobs, a different +transport, and a draft length that curve does not carry. + +### Why the accuracy gate alone could not have released this + +Rejection sampling means a miscomputed draft layer produces correct text more +slowly, so an accuracy score cannot separate a good MTP layer from a broken +one — the Δ above would have been just as small with the `eh_proj` halves +swapped. `acceptance_length` against a reference is the only measurement that +can, which is why both gates are listed and why neither is optional. + +**What the acceptance numbers do and do not establish.** They establish that +the MTP layer computes something the target model agrees with: 97.6% of draft +tokens accepted at con=1 is a ceiling-adjacent 3.9274 of a possible 4.0, and a +layer with the `eh_proj` halves swapped, a norm on the wrong operand or the +expert stack packed wrong would sit near 1.0. On their own they did **not** +establish that this is the *best* achievable draft quality — that needs the +same checkpoint under stock in-tree modeling at the same load and the same +`speculative_config`, which the section below measures. + +**And read the workload the other way round from the usual warning.** The +harness generates uniformly random prompt token ids with `ignore_eos`, and +`docs/models/multi-token-prediction.md` warns that such a workload gives +acceptance *far below* real text. That warning does not hold here, and the +reason is that only the **prompt** is random: the 1024 output tokens are the +model's own continuation of nonsense, which is highly repetitive, and +repetition is exactly what an MTP layer drafts perfectly. So this probe +**flatters** MTP rather than penalizing it. It is a correctness instrument +here and nothing more; the value of enabling MTP has to be judged on a +real-text workload. + +Two numbers from the probe that are *not* results and must not be quoted as +Pareto points: 211.17 tok/s/user at con=1 and 1942.9 tok/s at con=32. They are +measured under a different config (`free_gpu_memory_fraction: 0.75`, and +`--acceptance` turns on `enable_iter_perf_stats`, which the harness's own +header says makes a label non-comparable to one without it), in a two-round +probe rather than a sweep, on the flattering workload above. The perf campaign +for the variant belongs to the tuner. + +### The parallel split, and what it rests on + +| part | split | per-rank shape | +|---|---|---| +| `q_a_proj`, `q_a_layernorm`, `q_b_proj`, `kv_a_proj_with_mqa`, `kv_a_layernorm`, `kv_b_proj`, `o_proj` | **replicated**, all 128 query heads | `[1536, 7168]`, `[1536]`, `[24576, 1536]`, `[576, 7168]`, `[512]`, `[32768, 512]`, `[7168, 16384]` | +| layers 0-2 dense MLP | **replicated**, intermediate 18432 | `gate_up [36864, 3584]`, `down [7168, 9216]` | +| shared expert | **replicated**, intermediate 2048 | `gate_up [4096, 3584]`, `down [7168, 1024]` | +| routed experts | EP window, 64 of 256 at offset `64 * moe_ep_rank` | `fc1 [64, 4096, 3584]`, `fc2 [64, 7168, 1024]` | +| per-expert NVFP4 scalars | **replicated over all 256** (deliberate) | `[256]` each | +| fp8 KV scales, router, both norms, embedding | **replicated** | unchanged | +| `lm_head` (shell-owned) | **replicated** — the shell builds the whole matrix under attention DP | `[129280, 7168]` | + +Two collectives per **MoE** layer, none in layers 0-2 (dense) and none on the +attention or residual path — 116 per forward: + +* `comm/allgather` on the post-attention normed hidden states, **before** the + router GEMM; +* `comm/reducescatter` on the expert window's output over the whole gathered + token set, handing each rank back exactly its own rows. + +Three things the correctness rests on, none of them checkable inside one rank: + +* **the routing agrees across ranks, and the gather placement is what makes + it so.** The four 64-wide windows must tile the routing space exactly once; + under attention DP the ranks hold different tokens, so that is restored by + gathering **before** the router GEMM — the router and `noaux_tc_op` then run + on byte-identical full token sets on all four ranks and each token's top-8 + ids agree. Routing locally and gathering afterwards would break it with no + error (`docs/models/expert-weight-packing.md`); +* **one activation quantization feeds every window.** The routed FC1 global + scale is `1 / input_scale` of the *shared* expert, which the checkpoint sets + to the max over all 256 routed experts — `derive_after_load` asserts that + over the full 256, which is why the per-expert scalars are loaded whole on + every rank; +* **the reduce-scatter is crossed in bf16.** The op sums, and it sums + `float8_e4m3fn` as raw bytes rather than as floats; the gather is a byte + move and would survive fp8, so the asymmetry is a live trap on the return + leg only. + +`attn_metadata.all_rank_num_tokens` is where the engine publishes the per-rank +row counts; `_dp_rows` reads it, asserts `all_rank_num_tokens[rank] == +hidden_states.shape[0]`, and returns `max(counts)` — every rank pads its token +block to the group-wide maximum, so both collectives run in their uniform form +and the ragged form is never used. Padded rows are sliced off after the +reduce-scatter and never reach the residual stream. + +### What the fp8 latent pool moves, and why each choice is what it is + +The KV cache is fp8-e4m3 per the checkpoint, and **nothing validates the fp8 +round trip at any layer**: the write scale, the read scale and the two folded +FMHA scales are independent roles with no relation checked anywhere in the +chain. Getting one wrong is silently mis-scaled output, never an error. + +* **`quant_mode` is derived from the checkpoint's quant config**, not + hard-coded: `kv_cache_quant_algo: FP8` selects `QuantMode`'s fp8-KV bit + (128), the value every MLA entry certifies. The engine's own + `quant_config.quant_mode` is a `QuantModeWrapper` rather than an int at this + pin, and bits outside the KV-cache group ride along unread anyway (`1152` + and `384` measured bit-identical to a bare `128` on every MLA flavor), so the + target maps the declared algo onto the bit itself. +* **`s = 1.0`, and it is checked rather than assumed.** All 122 per-layer + `k_scale`/`v_scale` tensors are loaded (they are 0-dim fp32 scalars) and + `derive_after_load` asserts each is exactly 1.0. Both scale arguments are + then passed as `None`, which every op reads as exactly 1.0 and which is what + the engine's own call sites pass. This is not a convenience: the fp8 MLA + *context* path is internally inconsistent at `s != 1.0` in **both** context + flavors (it quantizes q/k/v at 1.0 while applying `s^2`/`s` as if it had + not), so 1.0 is the only correct value and the assert is the defence. +* **The decode producers are ordered, not concurrent.** + `mla_rope_generation` does not write `fused_q` over an fp8 pool — it *reads* + `fused_q[..., :C]` to build `quant_q_buffer`, which is the query the decode + FMHA consumes. The absorbed-q BMM is issued before it on the ambient stream, + which is what makes that safe; the bf16 reading (disjoint halves, free to + overlap) is a silent race here. +* **The two phases divide the quantization labour oppositely.** Context: + `mla_rope_append_paged_kv_assign_q` leaves `q` plain bf16 and + `thop_attention`'s context call quantizes q/k/v to e4m3 itself, so `q` is + quantized exactly once. Decode: `mla_rope_generation` produces the quantized + query and the two folded FMHA scales, and the generation call reads + `quant_q_buffer`, `mla_bmm1_scale[1]` and `mla_bmm2_scale[0]` — **ignoring + both kv scale tensors entirely**. Nothing in either signature says so. +* **A cached prefix pays fp8 twice.** The engine dequantizes cached latent + rows off the pool, `kv_b_proj` up-projects them, and `thop_attention` + quantizes the result straight back to e4m3. Not incorrect, but a + cached-prefix context call carries strictly more quantization error than a + fresh prefill of the same tokens — do not attribute an accuracy gap to block + reuse without accounting for it. +* **fp8 raises the attention workspace requirement.** Measured here: + **2,126,512,128 B (1.98 GiB) per rank**, resized on the first call, against + 531,628,032 B for the 32-head bf16 MLA shape at the same + `max_num_tokens=8192`. Sizing from bf16 MLA figures under-budgets. + +### The rope table is the whole of the rope configuration + +`thop_attention` reads the table's **content** and `q_scaling`; the seven +scalar rope arguments and `rotary_inv_freq` beside them are measured inert on +the MLA path. So the YaRN blend lives entirely in the table this target +builds, and the model's YaRN attention temperature lives entirely in +`q_scaling`: + +* table `inv_freq(d) = ramp(d)/(factor*freq(d)) + (1-ramp(d))/freq(d)` with + `low = 10`, `high = 23` at this config (`R` 64, `theta` 10000, `factor` 40, + `original_max_position_embeddings` 4096, `beta_fast` 32, `beta_slow` 1) — + cross-checked against the HF reference's own YaRN rope init to **1.05e-8 + max abs / 1.27e-7 max relative** on `inv_freq` (the HF value is fp32, so + that is its rounding); +* table amplitude `m(mscale)/m(mscale_all_dim)` = **exactly 1.0** (both are + 1.0 here), matching the same reference's `attention_factor`; +* `q_scaling = 1/m(mscale_all_dim)^2` = **0.5336594470450011**, which the + construction asserts lands on `thop_attention`'s certified fp8-pool value + (1.0 and 0.53366 are the two certified). + +**The table is built for the full `max_position_embeddings` = 163840 +positions** (84 MB fp32 per rank). A short table is read out of bounds with no +check — `rope_max_positions`, the argument that looks like it bounds this, is +one of the inert seven — so the first forward additionally asserts +`max_seq_len <= max_pos`. + +### Envelope + +* Context sequences with a cached prefix are served; chunked prefill is not + implemented (`AttentionRuntimeFeatures.chunked_prefill=False` under + defaults, and the target holds `chunked_prefill_buffer_batch_size=1`). +* **MTP is a `configs/` variant, never the identity.** The checkpoint declares + `num_nextn_predict_layers: 1` and ships the whole of layer 61 for it — its + own 256 experts, `embed_tokens`, `eh_proj`, two extra norms and a + `shared_head.head`, 790 keys in all. + + Under `llm_args.yaml` alone those 790 keys are a **predicted non-load** in + the weight manifest (read off the checkpoint by layer index rather than + enumerated), no MTP parameter is declared, and the shell's `forward` is a + plain delegation to the inherited base — nothing about this identity changed + when MTP was added, which is a measured claim and not a design intention: + after the increment, smoke's 10 greedy continuations are **byte-identical** + to the pre-increment run's. + + Under `configs/mtp{1,2,3}.yaml` the module is loaded and a second forward + path runs. **Those variants cross the line a config variant normally + respects, and that is stated rather than left to be inferred**: a variant is + supposed to move a knob, and these change the **weight-loading path** — layer + 61 goes from non-load to loaded, 212 keys and ~6.3 GB per rank — add a + second forward (`MTPLayer`, replayed `max_draft_len` times per step), and + change what the runtime allocates (a 62-layer KV pool, `max_draft_len - 1` + extra tokens per sequence, and a `max_seq_len` the engine raises by 8 at + `max_draft_len: 3`). Only the gates run against a variant can speak for it; + the identity gate records below are not evidence about it, and vice versa. +* fp8-e4m3 latent pool only (`quant_mode` = the fp8-KV bit, KV scaling factor + 1.0); beam width 1; no LoRA, no cross attention, no FlashMLA layout — each + asserted at the first forward or in the forward itself. +* `tokens_per_block == 32` is asserted: every MLA entry's **fp8** column is + certified at page 32 only (their bf16 columns also carry 64). 32 is what a + default `KvCacheConfig` produces, so this binds only if someone tunes the + page size — which would need the certification extended first. +* Pipeline parallelism and a second (tensor) split of the routed experts are + asserted off: `pp_size == 1`, `moe_tp_size == 1`. Attention DP is asserted + **on** — this target is the DP assembly. +* A rank with **zero** tokens was never observed (idle ranks get a 1-token + dummy) and is not something this forward was exercised on. + +Certification coverage this target consumes, and where it sits relative to +what the entries measured: + +* **`thop_attention`** runs MLA at `H = 128` over an **fp8-e4m3 latent pool**, + page 32, `q_lora_rank = 1536`, at the complete DeepSeek-R1-0528 rope/scale + cell — the exact configuration the entry certifies for all three MLA call + flavors and their mixed-batch pairing. `H = 128` is the only head count + certified over an fp8 pool, and page 32 the only page size. +* **`mla_rope_generation`**, **`mla_rope_append_paged_kv_assign_q`** and + **`load_paged_kv_cache_for_mla`** run their fp8 columns at the same cell + (`H = 128`, page 32, `C/R/nope/v = 512/64/128/128`, scale omitted). +* **`predicted_tokens_per_seq` (`P`) is 1 on every context call and the + generation call's own query-tokens-per-sequence.** Without a + `speculative_config` that is 1 everywhere; under `configs/mtp{1,2,3}.yaml` the + MLA **generation** calls run at `P = max_draft_len + 1` (2, 3, 4) on the + trunk and on the MTP layer's draft step 0, and at `P = 1` on draft steps 1+. + `thop_attention` and `mla_rope_generation` certify `P` at **1, 2, 3 and 4** + over this exact fp8 cell and no further, which is why `mtp3` is the largest + variant here: `max_draft_len: 4` would need the certification extended + first, not just a bigger number in the yaml. Two preconditions ride with it, + both held structurally rather than checked (a device read would cost a sync + and is illegal under capture): + * **`L_g >= P` for every generation sequence** — draft row 0 attends to + `[0, L_g - P]`, so a shorter KV length leaves it no keys. `L_g` counts all + `P` of the step's tokens, so it cannot be smaller than `P`. + * the `G*P` cache rows one `mla_rope_generation` call writes must address + **distinct physical slots**, which a `KVCacheManager`-allocated batch gives + automatically. +* **`fused_moe`** (the MTP layer's routed experts, bf16) runs at + `E = 64`, `(H, I) = (7168, 2048)`, `K = 8`, `ep_size = 4`, + `ep_rank ∈ {0,1,2,3}` — the cell the entry certifies, including all four + windows against a 256-expert reference — with `T <= 2048`. Its + **distinct-expert-ids precondition** matters here and is worth stating + rather than leaving to inference: above 256 tokens a repeated id in one + token's row reads out of bounds in `finalizeMoeRoutingKernel` (an illegal + access, or ~200-460 ulp of silent garbage), and this call drives `T` to 2048. + It holds **structurally**: the ids come from `noaux_tc_op`, whose semantics + are the indices of the top-k largest corrected scores, and a top-k over + expert indices cannot select one twice. + + The chunk bound is 2048 rather than the trunk's 8192 for a memory reason, + not a certification one — see *What MTP costs in memory* below. +* **`fp4_block_scale_moe_runner`** runs at the R1 routed geometry + (`H = 7168`, `I = 2048`, `num_experts = 256`, `top_k = 8`) with + `local_num_experts = 64`, `local_expert_offset ∈ {0, 64, 128, 192}` — the + certified four-way split — and at **`T <= 8192`**, the top of the entry's + certified token column, which the whole column covers at this geometry. The + gathered token set reaches `4 * max_num_tokens = 32768`, so the expert call + is **chunked** (`_MOE_MAX_T` in `modeling.py`). **That bound is a + certification boundary, not a tuning knob.** Routing and the activation + quantization are chunked with it, which keeps `noaux_tc_op` inside its own + enumerated column (up to `num_tokens` 8192) as well. +* **`noaux_tc_op`** runs the grouped configuration + `(num_experts, n_group, topk_group, topk) = (256, 8, 4, 8)`, which the entry + enumerates explicitly. +* **`comm/allgather` and `comm/reducescatter`** are driven at world size 4, + group `[0,1,2,3]`, bf16, hidden 7168, in their uniform form only, at 58 + sites per step — 58 + `max_draft_len` under an MTP variant, the MTP layer + being one more MoE layer per draft step. Both contracts certify the + call-order surface: calls pair by **position** on the communicator, so + issuing the same sequence on every rank is this forward's obligation — + discharged structurally, since every rank runs the same layer loop and pads + to `max(all_rank_num_tokens)`. + + **Inside the draft loop the padding basis is a different list, and reading + the wrong one is silent.** The worker leaves `attn_metadata.all_rank_num_tokens` + holding the trunk's counts for the whole loop and passes the correct basis in + as the `all_rank_num_tokens` **keyword** — + `spec_metadata.all_rank_num_tokens` at draft step 0, then + `spec_metadata.subseq_all_rank_num_tokens`, which is the per-rank *sequence* + count. Both contracts certify that at equal byte counts a mispairing does not + hang: every rank comes back wrong in 98-99% of elements, bitwise + reproducibly. The defence is structural rather than vigilant — the MTP + layer's padding helper (`_mtp_dp_rows`) takes the list as a parameter and has + **no metadata argument at all**, so it cannot reach the wrong one; the + trunk's `_dp_rows` keeps its own shape and the two are not interchangeable. + + **Both collectives stay on the ambient stream**, which since the + `iter1-shared-side-stream` iteration is not the only stream the forward + uses: the shared-expert branch is forked onto a side stream and joined + before the add. That is inside what both contracts certify — "the side + stream joined to the current one on both ends, which is what the engine and + a target's forward both do" — and it changes neither collective's arguments, + its call order, nor its uniform form. +* **`nvfp4_gemm` and `fp4_quantize` are certified at this target's shapes**, + on receipts rather than on a domain rule. They were assembled as "inside the + stated domain, outside the enumerated list" and closed afterwards, because + R1's widths *exceed* every previously-run value rather than falling between + them. `nvfp4_gemm` at `(K, N)` = (7168, 36864), (18432, 7168), (7168, 4096), + (2048, 7168) with `M` to 8192; `fp4_quantize` at `K` = 7168 / 18432 / 2048 + through `T = 8192`, in **both** scale layouts (swizzled for the GEMM, linear + for the MoE runner — `K = 7168` is taken in both). + + Both runs came back "the rule was sufficient after all", and both found the + contract wrong about *why*. `nvfp4_gemm`: the tactic space does move with + shape, but every R1 count lands inside the range the smaller shapes produce + and all four backends agree bitwise. `fp4_quantize`: the kernel-selection + axis its contract described (a TMA variant above 1024 rows) **does not exist + in this build** — that text came from flashinfer's vendored source, which is + a newer revision than the installed binary; the profiler sees one kernel at + every shape. + +### What MTP costs in memory, and why the variants carry a second knob + +The MTP module adds **5.87 GiB (6.30 GB) of declared weights per rank** — 5.637 +GB of it the bf16 routed expert stacks alone, which cost **3.56x what one NVFP4 +trunk MoE layer costs** purely from the dtype. `nvidia-smi` read **122,076 MiB +(119.2 GiB) per rank** after model init with MTP on, against the 117,924 MiB +(115.2 GiB) recorded above for the identity config; the two readings were taken +at different points of the load, so treat them as two figures rather than as a +clean difference. The engine then adds one KV-pool layer (62 instead of 61, ++1.6%) and `max_draft_len - 1` extra tokens per sequence, neither of which the +target declares. + +**That is not what makes it tight.** The engine sizes the KV pool from free +memory *after* the weights, at `free_gpu_memory_fraction`, and the drafting +forward's transient demand after that point is larger than the identity +config's. At the default 0.9 the pool takes 38.95 GiB (1,171,200 tokens per +rank) and boot then reaches CUDA-graph capture at **182.4 of 183.4 GiB and +livelocks** — three of four ranks spinning in `cudaFree` inside the CUDA +caching allocator's `release_cached_blocks`, reached from an ordinary +`empty_cuda` in the *trunk's* MoE runner, while the fourth waits at an +`MPI_Barrier`. There is no OOM exception and no error: the run simply stops +making progress and has to be killed. **Read the signature — one rank idle at a +barrier, the rest at 100% GPU with a frozen log — as memory exhaustion, not as +a mispaired collective.** + +**Freeing memory before the pool is sized does not help**, and that is worth +recording because it is the obvious first move: chunking the MTP layer's +expert call to 2048 rows (from the trunk's 8192, both inside `fused_moe`'s +certified column) moved pool sizing by **0.31 GiB**, 38.64 -> 38.95, and the +boot failed identically — the pool grows into exactly whatever is freed. The +chunk was reverted to the trunk's `_MOE_MAX_T`; it is not the lever. + +So `configs/mtp{1,2,3}.yaml` each carry +`kv_cache_config.free_gpu_memory_fraction: 0.75` as a **boot requirement**, +documented in the files themselves, in the same spirit as +`configs/trtllm-ref-boot.yaml`. It hands the pool ~32.5 GiB and leaves ~28 GiB +of headroom, and it constrains nothing: ~975k KV tokens per rank is an order +of magnitude past what any concurrency measured on this target uses. + +### Engine-side facts observed on this checkpoint + +* **The latent pool is fp8 and the engine sized it that way.** 44.77 GiB for + 1,368,032 tokens per rank = **35,136 B/token** = + `num_layers * (kv_lora_rank + qk_rope_head_dim) * 1` = `61 * 576 * 1`. One + byte per element, i.e. half the bf16 width — no target-side declaration was + needed beyond keeping the checkpoint's `hf_quant_config.json` linked. +* **The engine allocates the KV pool twice**: a small profiling pool (5.50 + GiB, 167,936 tokens) is created and released before the real one. +* **Per-rank weights are ~115.2 GiB** (`nvidia-smi` read 117,924 MiB on three + of the four devices after model init and before the pool allocation; the + fourth carried an unrelated 810 MiB from another user's process), against a + 183 GiB card. `Model init total` is 51-55 s per rank with the checkpoint + warm in the host page cache — the 163 shards load in ~10 s per rank there. +* **CUDA-graph capture is the default decode grid**: `batch_sizes = [1..32, + 64, 128]`, `max_batch_size 128`, `enable_padding False` — 34 sizes, all + inside the enumerated CUDA-graph coverage both collective entries carry + (`1..32, 64, 128, 256`). +* **First boot JIT-compiles 8 fp8 MLA decode kernels per rank** at ~6.0-7.0 s + each — `fmhaSm100aKernel_QkvE4m3OBfloat16HQk576HV512...P32VarSeqQ16Kv128StaticSwapsAbForGen` + and its `MultiCtasKvCga` / `HVPerCta256` siblings. `QkvE4m3` in the name is + the fp8 pool: q, K and V are all e4m3 in the decode MMAs. No MLA context + call triggers a compile. + + **Each `max_draft_len` costs its own decode compiles**, because + `predicted_tokens_per_seq` becomes the kernel's `maxSeqLenQ` and moves the + `HVPerCta` split in the name — `P` 1 and 2 take `HVPerCta128`, `P` 3 and 4 + take `HVPerCta256`, and `P` 1 and 2 pay separate compiles despite sharing a + name. Measured on `configs/mtp3.yaml`: **12 compiles per rank** against the + identity config's 8, i.e. **+4 at ~5.4-5.7 s each**, all four ranks + compiling in parallel. A campaign sweeping `max_draft_len` should budget one + first-boot compile set per value, not one for the model. +* `max_seq_len` is the config's 163840 under the identity config, so + `max_blocks_per_seq` is 5120 at page 32. **A `speculative_config` raises it**: + 163848 at `max_draft_len: 3` (5121 blocks per sequence), which is more than + the `max_draft_len - 1` extra KV tokens per sequence the runtime reference + documents. The rope table is sized from the engine's own number rather than + from a formula, on the first forward. + +## Performance + +![Serving Pareto](perf/figures/pareto.png) + +Environment: `umbriel-b200-027`, 224-core, 8x NVIDIA B200 (178.34 GiB each), +**GPUs 3-6** — the campaign device set, held for every label — driver +**595.58.03**, `tensorrt_llm 1.3.0rc21`, `torch 2.11.0+cu130`. +`source scripts/env.sh` then `CUDA_VISIBLE_DEVICES=3,4,5,6`. ISL=OSL=1024, +concurrency 1..256. **Two sessions on the same host and the same device set**: +the first three curves 2026-08-01 10:45-15:18 UTC, the three MTP curves plus +the two new stock references 2026-08-02 13:09-17:35 UTC. The second session +opened by re-measuring `iter1` (`probe-anchor-iter1`), which reproduced it to +**−0.14% at con=1 and +0.15% at con=256**, and the `trtllm` reference was +spot-checked at the end of it to **−0.01% / +0.20%** — that is what licenses +one figure across the two days. + +`con=1 tok/s/user` is `1000 / mean_tpot_ms` at con=1; `peak tok/s/GPU` is +`max(output_throughput) / 4`. + +| label | config | commit | accuracy | con=1 tok/s/user | peak tok/s/GPU | change | +|---|---|---|---|---|---|---| +| `baseline` | `llm_args.yaml` (identity) | `70b1ecd` | gsm8k 94.9962 | 76.09 | 1906.86 (con=256) | trtllm defaults | +| `iter1-shared-side-stream` | `llm_args.yaml` (identity) | `75a471a` | gsm8k **94.7688** | 82.43 | 1948.78 (con=256) | shared expert forked onto a side stream | +| `mtp1` | identity + `configs/mtp1.yaml` | `bdbb1ab` | covered by the mtp3 pair, below | 135.11 | 2184.44 (con=256) | MTP, `max_draft_len: 1` | +| `mtp2` | identity + `configs/mtp2.yaml` | `bdbb1ab` | covered by the mtp3 pair, below | 164.94 | 2190.04 (con=256) | MTP, `max_draft_len: 2` | +| `mtp3` | identity + `configs/mtp3.yaml` | `bdbb1ab` | gsm8k **95.2237** (paired, +0.1517) | **190.45** | **2262.06** (con=256) | MTP, `max_draft_len: 3` — the frontier | +| `trtllm` | identity + `configs/trtllm-ref-boot.yaml` | `70b1ecd` | ungated | 73.94 | 1522.73 (con=256) | stock in-tree modeling at its **default** MoE transport (DeepEPLowLatency); config is boot-forced, see below | +| `trtllm-nodeepep` | identity + `configs/trtllm-ref-boot.yaml`, `TRTLLM_CAN_USE_DEEP_EP=0` | `471e6a3` | ungated | 79.20 | 1768.06 (con=256) | the same, on **AllGatherReduceScatter** — the transport control, 3 points | +| `trtllm-mtp3` | identity + `configs/trtllm-ref-mtp3.yaml`, `TRTLLM_CAN_USE_DEEP_EP=0` | `471e6a3` | ungated | 178.66 | 1780.70 (con=256) | **stock modeling with MTP at the same `max_draft_len: 3`** — the reference `mtp3` should be read against | +| `trtllm-tuned` | — | — | — | — | — | superseded by `trtllm-mtp3`, which measures exactly this; see below | +| **`trtllm-aligned`** | **`llm_args.yaml` (identity), `TRTLLM_CAN_USE_DEEP_EP=0`** | `471e6a3` | ungated | **77.82** | **1887.78** (con=256) | **stock modeling at this target's own config — no boot knobs at all. Supersedes `trtllm` / `trtllm-nodeepep`** | +| **`trtllm-aligned-mtp3`** | **identity + `configs/mtp3.yaml`, `TRTLLM_CAN_USE_DEEP_EP=0`** | `471e6a3` | ungated | **181.74** | **1728.86** (con=256) | **stock modeling with MTP at the same config file `mtp3` uses. Supersedes `trtllm-mtp3`** | + +Per point, output tok/s (whole 4-GPU node): + +| con | `trtllm` | `baseline` | `iter1` | iter1/trtllm | t TPOT | i TPOT | t TTFT | i TTFT | +|---|---|---|---|---|---|---|---|---| +| 1 | 73.28 | 75.63 | 81.90 | **1.118x** | 13.525 | 12.132 | 138.6 | 92.4 | +| 2 | 140.11 | 145.44 | 157.59 | 1.125x | 14.059 | 12.554 | 234.0 | 151.9 | +| 4 | 270.94 | 285.41 | 308.40 | 1.138x | 14.525 | 12.831 | 258.2 | 155.2 | +| 8 | 474.91 | 513.02 | 552.59 | 1.164x | 16.486 | 14.279 | 383.2 | 216.2 | +| 16 | 818.94 | 913.71 | 972.02 | 1.187x | 19.091 | 16.131 | 464.4 | 352.0 | +| 32 | 1360.02 | 1584.57 | 1673.50 | 1.230x | 23.033 | 18.514 | 518.1 | 637.5 | +| 64 | 2335.92 | 2724.10 | 2817.55 | 1.206x | 26.745 | 21.639 | 665.4 | 1116.6 | +| 128 | 3835.72 | 4507.69 | 4566.30 | 1.190x | 32.373 | 26.663 | 930.3 | 1406.9 | +| 256 | 6090.94 | 7627.45 | 7795.10 | **1.280x** | 40.129 | 31.092 | 1427.8 | 1762.4 | + +**That 11.8-28.0% is measured against a reference that could not run this +target's own config, and a later session showed it did not have to be that +way.** Read the next subsection before quoting any number from the table +above. Mean TPOT is lower at every point, and the mean-TTFT crossover at +con=32..256 is a property of the reference's forced config rather than of the +modeling layer. + +### The reference's boot config turned out to be avoidable, and it cost 24% + +`configs/trtllm-ref-boot.yaml` exists because stock trtllm OOMs at engine +construction on this checkpoint: 113.9 GiB per rank allocated before any weight +is materialized, `= 1.96 GiB x 58 MoE layers`, which is DeepEP low-latency's +communication workspace sized from `max_num_tokens`. Capping `max_num_tokens` +and `max_seq_len` to 2048 treats that symptom. **`TRTLLM_CAN_USE_DEEP_EP=0` +removes the cause** — the MoE communication factory then falls through to +`AllGatherReduceScatter`, which its own source comment calls "always works", +and which is the transport *this target implements by hand*. Stock trtllm then +boots at the **identity config**, `max_num_tokens` 8192 and `max_seq_len` +163840, peak 179,188 MiB with ~4 GB to spare. That switch was not known when +the boot config was written. + +Re-measured on 2026-08-03, GPUs 3-6, both sides byte-identical config +(`trtllm-aligned` and `trtllm-aligned-mtp3`): + +| con | 1 | 8 | 32 | 64 | 128 | 256 | +|---|---|---|---|---|---|---| +| the boot config was costing the reference | +6.2% | +8.2% | +10.9% | +11.6% | +15.9% | **+24.0%** | +| **`iter1` / `trtllm-aligned`** (no MTP) | 1.052x | 1.075x | **1.109x** | 1.081x | 1.027x | **1.032x** | +| **`mtp3` / `trtllm-aligned-mtp3`** | 1.045x | 1.313x | 1.289x | 1.339x | **1.406x** | **1.308x** | + +**So the honest no-MTP lead is 2.7-10.9%, not 11.8-28.0%**, and the honest MTP +lead is 1.308x. Both older reference lines and the whole +*MoE transport is worth 7-16%* subsection below are superseded by this: the +transport is no longer a variable to isolate, because both systems now run the +same one. + +**And the MTP lead is not a better MTP layer.** Per system, at matched config, +MTP is worth: + +| con | 1 | 8 | 32 | 128 | 256 | +|---|---|---|---|---|---| +| to `staircase` | 2.285x | 1.930x | 1.301x | 1.332x | **1.161x** | +| to stock trtllm | **2.300x** | 1.580x | 1.119x | **0.973x** | **0.916x** | + +They are level at con=1 — stock is fractionally ahead — and stock goes +*negative* from con=128. The 1.308x is that divergence, not a drafting-quality +difference; acceptance agrees to a few percent everywhere. +`workbench/docs/2026-08-03-r1-dep4-staircase-vs-trtllm.md` carries the +kernel-level attribution: it is one family, MoE expert GEMM, and the driver is +**row count, not a re-run draft layer**. MTP gives every generation sequence +`max_draft_len + 1 = 4` query tokens, so the *trunk*'s MoE gathers 4x the rows +(256 -> 1024 per rank at con=256). The two implementations scale differently +under that: CUTLASS grouped GEMM goes from 5 to 8 kernels per MoE layer, while +trtllm-gen's `bmm_E2m1_*` barely moves. Expert-GEMM-family launches per rank +per step: **174 -> 192 (+10%) on our side, 290 -> 482 (+66%) on theirs** — and +of their +192, the MTP layer itself accounts for only 15. **The growth is in +the trunk, not in the draft layer.** + +The MTP curves, against `iter1` (all four are full sweeps; `mtp*` from the +2026-08-02 session). `accept` is mean tokens emitted per engine step, ceiling +`max_draft_len + 1`: + +| con | `iter1` | `mtp1` | `mtp2` | `mtp3` | mtp3/iter1 | i TPOT | m3 TPOT | i TTFT | m3 TTFT | m3 accept | +|---|---|---|---|---|---|---|---|---|---|---| +| 1 | 81.90 | 133.54 | 162.51 | 187.12 | **2.285x** | 12.132 | 5.251 | 92.4 | 100.9 | 3.4904 | +| 2 | 157.59 | 252.37 | 313.21 | 341.41 | 2.166x | 12.554 | 5.611 | 151.9 | 124.7 | 3.4040 | +| 4 | 308.40 | 501.37 | 622.24 | 677.82 | 2.198x | 12.831 | 5.689 | 155.2 | 123.8 | 3.3923 | +| 8 | 552.59 | 803.10 | 928.24 | 1066.23 | 1.930x | 14.279 | 7.145 | 216.2 | 166.6 | 3.4151 | +| 16 | 972.02 | 1316.50 | 1313.32 | 1431.51 | 1.473x | 16.131 | 9.394 | 352.0 | 217.9 | 3.3905 | +| 32 | 1673.50 | 2135.72 | 2091.70 | 2176.90 | 1.301x | 18.514 | 12.199 | 637.5 | 282.6 | 3.3356 | +| 64 | 2817.55 | 3341.71 | 3627.89 | 3723.24 | 1.321x | 21.639 | 13.618 | 1116.6 | 415.4 | 3.4407 | +| 128 | 4566.30 | 5877.32 | 5733.13 | 6080.38 | 1.332x | 26.663 | 17.524 | 1406.9 | 612.1 | 3.4173 | +| 256 | 7795.10 | 8737.78 | 8760.17 | 9048.23 | **1.161x** | 31.092 | 23.758 | 1762.4 | 923.9 | 3.4033 | + +**`mtp3` is ahead of `mtp1` and `mtp2` at every one of the nine +concurrencies**, and ahead of `iter1` at every one, so the `max_draft_len` axis +never turns over inside the certified range. Mean TTFT also falls at every +point except con=1, where it rises 92.4 -> 100.9 ms. + +**Do not read the `mtp3` column against the `trtllm` row.** That row has no +speculative decoding at all, so the ratio between them is mostly "one system +drafts and the other does not", not a modeling-layer delta. The comparison +this target should be quoted on is `mtp3` against **`trtllm-mtp3`** — stock +in-tree modeling, same checkpoint, same `max_draft_len: 3`, same MoE +transport: + +| con | `trtllm-mtp3` | `mtp3` | **mtp3 / trtllm-mtp3** | their accept | our accept | tm TPOT | m3 TPOT | tm TTFT | m3 TTFT | +|---|---|---|---|---|---|---|---|---|---| +| 1 | 175.44 | 187.12 | **1.067x** | 3.5108 | 3.4904 | 5.597 | 5.251 | 110.6 | 100.9 | +| 2 | 342.16 | 341.41 | **0.998x** | 3.5415 | 3.4040 | 5.709 | 5.611 | 128.4 | 124.7 | +| 4 | 621.75 | 677.82 | 1.090x | 3.3315 | 3.3923 | 6.183 | 5.689 | 126.3 | 123.8 | +| 8 | 845.69 | 1066.23 | 1.261x | 3.3058 | 3.4151 | 8.788 | 7.145 | 225.3 | 166.6 | +| 16 | 1044.72 | 1431.51 | **1.370x** | 3.3305 | 3.3905 | 12.020 | 9.394 | 270.3 | 217.9 | +| 32 | 1768.44 | 2176.90 | 1.231x | 3.3458 | 3.3356 | 14.690 | 12.199 | 335.1 | 282.6 | +| 64 | 2797.56 | 3723.24 | 1.331x | 3.2835 | 3.4407 | 18.454 | 13.618 | 446.2 | 415.4 | +| 128 | 4510.25 | 6080.38 | 1.348x | 3.3125 | 3.4173 | 23.658 | 17.524 | 688.2 | 612.1 | +| 256 | 7122.80 | 9048.23 | **1.270x** | 3.3145 | 3.4033 | 29.746 | 23.758 | 1134.4 | 923.9 | + +**The honest headline is 1.27x at con=256 and 1.00-1.09x at con=1..4**, not +the 2.29x the no-MTP row invites. Two things follow, and they are the point of +this reference: + +* **Both systems draft equally well.** Acceptance agrees to within a few + percent at every concurrency (theirs 3.28-3.54, ours 3.40-3.49), which is + the same conclusion the level-3 acceptance gate reached and is what makes + the throughput ratio a *speed* comparison rather than a draft-quality one. +* **The modeling-layer delta survives MTP essentially unchanged.** At con=256 + it is **1.270x with MTP against 1.280x without** — and the gap decomposition + below prices the reference's forced config at 1.035x on our side at that + point, so both reduce to a ~1.23x modeling delta. MTP moved the whole + frontier; it did not move the distance between the two systems. + +At con=1..4 the two are level. Stock trtllm's low-concurrency MTP path is as +good as ours; our lead only opens up from con=8, which is where `iter1`'s +side-stream overlap and the rest of the trunk work start to matter. + +#### The MoE transport is worth 7-16% to stock trtllm, and it is not noise + +**Superseded by `trtllm-aligned`.** This subsection isolated the transport as +a variable because the reference was stuck on `configs/trtllm-ref-boot.yaml`. +It no longer is: with `TRTLLM_CAN_USE_DEEP_EP=0` stock trtllm boots at the +identity config, so both systems now run `AllGatherReduceScatter` and there +is nothing left to isolate. The measurement below stands on its own. + +`trtllm-mtp3` cannot run on DeepEPLowLatency — that transport's dispatch +accepts only NVFP4 hidden states and the MTP layer is bf16 — so it is forced +onto `AllGatherReduceScatter`, while the existing `trtllm` curve was measured +on DeepEPLowLatency. `trtllm-nodeepep` isolates that one variable: stock +modeling, **no** MTP, the same boot config, `TRTLLM_CAN_USE_DEEP_EP=0`. + +| con | `trtllm` (DeepEPLowLatency) | `trtllm-nodeepep` (AllGatherReduceScatter) | delta | +|---|---|---|---| +| 1 | 73.28 | 78.64 | **+7.32%** | +| 32 | 1360.02 | 1491.47 | **+9.67%** | +| 256 | 6090.94 | 7072.22 | **+16.11%** | + +This does **not** collapse into the existing reference — it is 7 to 16 times +this session's measured con=256 floor of 0.01%. Two corrections follow, and +both cut against this target: + +* **Stock trtllm's default transport is the slower one here**, so the main + table's `trtllm` row understates what stock trtllm can do on this host. + Against the faster stock configuration, `iter1`'s no-MTP lead is + **1.041x / 1.122x / 1.102x** at con 1/32/256 — not the 1.118x / 1.230x / + 1.280x measured against the default. The 11.8-28.0% claim above is against + `trtllm` **as stock defaults configure it**, which is the honest definition + of that reference line but is not the best stock can do. +* **MTP is worth far less to stock trtllm at scale than to us, once the + transport is held fixed.** `trtllm-mtp3 / trtllm-nodeepep` is **2.231x at + con=1, 1.186x at con=32 and 1.007x at con=256** — at the throughput end, + drafting buys stock trtllm essentially nothing. Ours, `mtp3 / iter1`, is + **2.285x / 1.301x / 1.161x**. That divergence at high concurrency is the + clearest single statement of what this target's modeling layer is worth + under MTP: the two systems draft equally well and both pay a step-cost + penalty that grows with batch, but ours stays ahead of the penalty and + stock's does not. + +`trtllm-nodeepep` is three points, not a swept curve — con 1/32/256 at the +same rounds the full sweeps use, so it pairs point-to-point with `trtllm`. +con=128 was deliberately skipped: this session measured 4.2% run-to-run spread +there, so it could not have carried the claim. + +**Read these curves as "did anything get slower, and where does it stop +paying" — not as "is MTP worth enabling".** The harness builds prompts from +uniformly random token ids and then generates with `--ignore-eos`, so what the +draft layer predicts is the model's own continuation of a nonsense prefix, +which degenerates into repetition — and repetition drafts near-perfectly. The +`accept` column above, 3.39-3.49 of a ceiling of 4.0, is that artefact. **The +value number for MTP on this checkpoint is the real-text one**, from the paired +gsm8k runs recorded under *The MTP variant's own gate records*: execution +**37.025 s -> 31.440 s, 1.178x** over 1319 genuine prompts. + +### Reference lines + +* **`trtllm`** = the original checkpoint under stock in-tree modeling, this + target's `llm_args.yaml` (identity config: `tensor_parallel_size: 4`, + `moe_expert_parallel_size: 4`, `enable_attention_dp: true`), **plus + `configs/trtllm-ref-boot.yaml`, which is boot-forced rather than tuned.** + Ungated. `tensorrt_llm 1.3.0rc21`. + + **Stock trtllm cannot boot this checkpoint at `dep4` on a 178.34 GiB B200 + under its own defaults.** Measured, per rank: at the default + `max_num_tokens: 8192`, **113.9 GiB is allocated at engine construction + outside the torch allocator, before any weight is materialized** — an + `nvidia-smi` sampler at 5 s intervals caught it going 20 MiB -> + 116,662 MiB inside one sample, immediately after the MoE communication + factory logged `Selected communication strategy: DeepEPLowLatency` once per + MoE layer (`NVLinkOneSided` / `NVLinkTwoSided` / `DeepEP` all reported + `not available: Invalid Argument`). Model init then dies in + `init_meta_tensor` with 64.17 GiB held by PyTorch and 113.8 GiB outside it. + 113.9 GiB / 58 MoE layers = 1.96 GiB per layer. `NVSHMEM_SYMMETRIC_SIZE=1g` + does not change it. + + The allocation is proportional to `max_num_tokens`: at 2048 the weights + load (135.85 GiB torch, including 1.53 GiB of CUDA-graph pools) and the + failure moves to `configure_kv_cache_capacity`, which asks for 7.00 GiB + with 5.61 GiB free. `kv_cache_config.free_gpu_memory_fraction` does **not** + move that 7.00 GiB (identical failure at 0.9 and at 0.6). The failure's own + memory ledger names the remaining lever, and `max_seq_len: 2048` — exactly + the harness's ISL+OSL, and `max_blocks_per_seq` 5120 -> 64 — is what makes + it boot. Both knobs together are the reference's config. + +* **`trtllm-nodeepep`** = the `trtllm` line with one variable moved: + `TRTLLM_CAN_USE_DEEP_EP=0`, which disables DeepEP and DeepEPLowLatency + together and lands the MoE communication factory on + `AllGatherReduceScatter` — the strategy this target implements by hand. + Stock in-tree modeling, no MTP, same boot config, three points. Ungated. + It exists to keep the transport from being confounded with MTP, and it is + worth 7-16%; see above. + +* **`trtllm-mtp3`** = the original checkpoint under stock in-tree modeling + with **the same speculative config this target's kept variant uses** + (`configs/trtllm-ref-mtp3.yaml` = the boot knobs + `decoding_type: MTP`, + `max_draft_len: 3`, `free_gpu_memory_fraction: 0.75`), plus + `TRTLLM_CAN_USE_DEEP_EP=0`. Ungated. Full nine-point sweep, no + `--acceptance`. **This is the reference `mtp3` is quoted against**, and it + replaces what the `trtllm-tuned` slot was for: it *is* stock modeling under + this campaign's final best config. + + It carries three forced deviations from `configs/mtp3.yaml` — + `max_num_tokens` / `max_seq_len: 2048` to boot at all, and the transport + variable — so it is not a *portable-config* split in the usual sense. Both + are priced rather than waved away: the boot knobs are worth `+1.5% / −0.6% / + −3.4%` on our side (the gap decomposition below), and the transport is worth + `+7.3% / +9.7% / +16.1%` on stock's side and is applied to `trtllm-mtp3` + already. Netting the boot knobs out of the con=256 ratio leaves ~1.23x, the + same modeling delta the no-MTP decomposition finds. + +* **`trtllm-tuned`** — **retired, superseded by `trtllm-mtp3`.** The slot means + "stock modeling plus the campaign's final best config", and that is now + measured rather than argued: the final best config is `configs/mtp3.yaml` + and `trtllm-mtp3` runs stock modeling under it. (For the identity campaign + the slot was genuinely degenerate — no config variant was kept, and stock + trtllm cannot boot at the identity config at all.) + +### The gap decomposition + +Because the reference carries a forced config, the honest split is measured +rather than asserted: `probe-refcfg` ran **this target's `iter1` code under +the reference's own config**, so both systems can be compared at identical +knobs. + +| con | staircase @ identity | staircase @ ref config | trtllm @ ref config | config-matched delta | end-to-end | +|---|---|---|---|---|---| +| 1 | 81.90 | 83.13 (+1.51%) | 73.28 | **1.134x** | 1.118x | +| 32 | 1682.70* | 1673.15 (−0.57%) | 1360.02 | **1.230x** | 1.230x | +| 256 | 7802.74* | 7535.64 (−3.42%) | 6090.94 | **1.237x** | 1.280x | + +\* paired two-round probes, the shape `probe-refcfg` used; the full-sweep +numbers are in the table above. + +**0% of the end-to-end gap is portable config and 100% of it is the +modeling-layer delta** — the campaign kept no config variant, so there is no +portable config gain to transfer. The reference's forced config is worth +`+1.5% / −0.6% / −3.4%` on our side, i.e. at con=256 the end-to-end 1.280x is +a 1.237x modeling delta plus a 1.035x config difference that happens to +favour the identity config on throughput. At con=1 and con=32 the +config-matched and end-to-end numbers agree to within the session spread. + +That forced config is also a real **TTFT-vs-throughput trade on this target**, +which is what the TTFT crossover in the main table is: at con=256 it costs +−3.42% output throughput and buys **mean TTFT 2269.5 -> 1645.7 ms (−27.5%)**; +at con=32, −0.57% for 627.8 -> 384.0 ms (−38.8%). It is correctly not kept — +the Pareto axes are tok/s/user and tok/s/GPU, and it loses on one and is flat +on the other — but a latency-sensitive deployment of this target should reach +for `max_num_tokens: 2048` first. + +### Iterations + +**Two kept: one modeling change (`iter1`), one config axis (`mtp3`).** + +**`iter1-shared-side-stream` — the shared expert overlaps the MoE round +trip.** *Evidence.* An nsys window of 100 executor iterations at con=256 +(`TLLM_PROFILE_START_STOP=1200-1300`, `-c cudaProfilerApi`, +`--cuda-graph-trace node`) caught all 4 ranks, 778,000 kernels, and **100 +pure-decode steps** — no context FMHA kernel appears at all, and +`cudaGraphLaunch` = 400 against 4 ranks x 100 steps = **1.00 graph replay per +rank per step, 100% coverage**. Per rank the GPU was **98.5% busy** (union +2822.3 ms against a 2865.1 ms window, idle **1.49%**), so the target is +GPU-bound, not host-bound. Walking the intervals per device — kernels on +different GPUs genuinely run in parallel — gave exclusive time (the interval +covered by no other kernel): + +| family | sum µs/rank-step | exclusive µs/rank-step | exclusive/union | % of GPU busy | +|---|---|---|---|---| +| routed expert GEMMs (`bmm_E2m1*`, `bmm_Bfloat16*`) | 13488.2 | **13078.7** | 98.9% | **46.3%** | +| cuBLAS bf16 GEMMs (`nvjet_*`, attention path) | 6828.4 | 5665.5 | 86.6% | 20.1% | +| NCCL collectives (`AllGather`/`ReduceScatter` `RING_LL`) | 3572.7 | **3369.7** | **94.3%** | **11.9%** | +| MLA decode FMHA | 1079.7 | 1015.2 | 94.0% | 3.6% | +| `splitKreduce_kernel` | 1062.6 | 324.3 | 30.5% | 1.1% | + +GPU busy is 28.22 ms per rank-step. The collectives are **94.3% exclusive** — +for almost all of the time one is running, the device is running nothing else +— which is a 3.37 ms per-step window with nothing in it. + +*Change.* The shared expert and the routed round trip both read the +post-attention `o` and meet only at the final `add`, but on one stream the +shared expert's five kernels sat *in front of* the all-gather. `forward` now +forks that branch onto a side stream (`side.wait_stream(main)` -> +`with torch.cuda.stream(side)` -> `main.wait_stream(side)` before the add), +so it runs inside the collective window. Both collective contracts certify +"a side stream joined to the current one on both ends", which is exactly this +shape; the same fork/join pair is what propagates a CUDA-graph capture into +the branch and back, and the captured decode graph keeps working (smoke's 10 +greedy continuations are byte-identical to the baseline's). + +*Effect.* Output throughput rises at **every** concurrency — +8.28, +8.35, ++8.05, +7.71, +6.38, +5.61, +3.43, +1.30, +2.20 % at con 1..256 — and mean +TPOT falls at every point. Frontier: con=1 tok/s/user 76.09 -> **82.43** +(+8.3%), peak tok/s/GPU 1906.86 -> **1948.78** (+2.2%). The gain is largest +at low concurrency, which the profile predicts: the collectives are +latency-bound and roughly batch-independent, so the window they leave is a +larger share of a smaller step, while at con=256 the shared expert competes +for HBM bandwidth with the routed expert GEMMs and only part of it hides. +Accuracy: **gsm8k 94.7688** `exact_match,flexible-extract` (strict-match +94.6171) against reference 94.9962, gate >= 90.0 — passed. The 0.23 delta is +3 questions of 1319, well inside the run's own ±0.6133 stderr, and +strict-match is bit-for-bit the same score as the assembly run. + +**This iteration makes two statements elsewhere in this file stale**, and +they are gate-record prose that a tuner does not edit: the *parallel split* +section and the *Vocabulary* section both say the two collectives are driven +"in their uniform form only" and that every rank "pads to +`max(all_rank_num_tokens)`". Both are still true — `iter1` did not change +either collective's arguments — but the surrounding claim that the forward is +single-stream no longer is. The certification consumed is unchanged: same +entries, same arguments, same uniform form, and both contracts certify side +streams explicitly. + +**`mtp3` — multi-token prediction at `max_draft_len: 3`, the whole certified +range swept.** *Evidence.* The MTP layer's own arithmetic +(`docs/models/multi-token-prediction.md`) puts the break-even acceptance at +1.055 / 1.110 / 1.164 for draft lengths 1 / 2 / 3, and predicts the step cost +to be flat in concurrency because a decode step is weight-bandwidth-bound. It +also flags the high-concurrency end as the genuinely uncertain one: once a +rank's batch already activates all 64 of its local experts, `draft_len + 1` +times the rows means the same weight bytes with several times the arithmetic, +and the layer can cross from memory-bound to compute-bound. + +*Change.* Config only — `configs/mtp{1,2,3}.yaml`, one full Pareto sweep each, +no modeling edit. Nothing else moved: `modeling.py` is byte-identical across +all three, and `--acceptance` was deliberately **not** used (below). + +*Effect.* `mtp3` takes the frontier on both axes: con=1 tok/s/user +**82.43 -> 190.45 (+131.1%)** and peak tok/s/GPU **1948.78 -> 2262.06 +(+16.1%)**, ahead of `iter1` at all nine concurrencies and ahead of `mtp1` and +`mtp2` at all nine. Against the matched-draft-length reference `trtllm-mtp3` +it leads by **1.270x at con=256** and is level (1.00-1.09x) at con=1..4 — +that, not the ratio to the non-drafting `trtllm` row, is what this iteration +is worth against stock. Decomposing throughput into the two terms it is made +of — the engine step now emits `accept` tokens instead of 1, and costs more — + +| con | accept | step cost vs a non-drafting step | decode speedup | end-to-end | +|---|---|---|---|---| +| 1 | 3.4904 | 1.511 | 2.311x | 2.285x | +| 32 | 3.3356 | 2.198 | 1.518x | 1.301x | +| 256 | 3.4033 | 2.600 | 1.309x | 1.161x | + +**Acceptance is essentially flat in concurrency and the whole decay is on the +cost side.** Across con 1..256 acceptance moves only 3.49 -> 3.40 (and +`mtp1` 1.967 -> 1.949 of 2.0, `mtp2` 2.748 -> 2.754 of 3.0) — never below 2.86x +the 1.164 break-even. The step cost, which the weight-bytes model says +should sit flat at 1.164, instead climbs **1.511 -> 2.600**. Reading the +marginal cost of each successive draft step at con=256 — **+0.645, +0.569, ++0.386** — against con=1's **+0.200, +0.173, +0.137** shows what it is: the +term that grows is proportional to the *rows* a step carries, not to the MTP +layer's fixed weight bytes, and the marginal cost falls with each further draft +step in the same way the marginal row count does (1 -> 2 rows is +100%, 2 -> 3 +is +50%, 3 -> 4 is +33%). That is the memory-bound-to-compute-bound crossover +the mechanism doc predicted, measured. It never becomes steep enough to +overtake acceptance, which is why the curve is still climbing at 3. + +*Gate.* No new accuracy run was needed and none is claimed: the paired gsm8k +record under *The MTP variant's own gate records* was measured on +**`configs/mtp3.yaml` itself**, and `modeling.py` has not changed since +(`4194dcb`, and this campaign added no code). identity **95.0720** / mtp3 +**95.2237**, Δ **+0.1517** against `|Δ| < 1.2`. `mtp1` and `mtp2` are not +separately gated and are not the recommended setting; they are the two +measured points that establish the axis is monotone, and the note in each file +says so. What the gate certifies is that the variant did not break the model — +rejection sampling makes even a miscomputed draft layer emit correct text, so +`acceptance_length` against stock trtllm, not the score, is what says drafting +works. + +### Rejected, with what was actually measured + +* **`max_seq_len: 2048` — rejected, and it retires a blocked knob.** + `kv_cache_config.tokens_per_block: 64` is uncertified over this target's + fp8 latent pool, and it was worth **+22% at con=256** on the + `deepseek-v3-lite-nvfp4/sm_100/tp1` sibling, so it is the knob a campaign + here would want most. `max_seq_len` reaches the *same* + `max_blocks_per_seq` reduction by the other route (5120 -> 64, an 80x cut + where page 64 gives 2x), and measured **−0.13% at con=32 and +4.08% at + con=256** against a paired probe — where a clean re-measure puts the same + point at +0.10%. The profile says why: **GPU idle is 1.49%**, so there is + no exposed host bookkeeping to recover. The tp1 sibling's win came from a + host-bound decode (GPU idle 0.401); this target is 40x larger per step and + the same host work is entirely hidden. **A vocabulary request to certify + page 64 over an fp8 pool would not pay for itself here.** +* **Ragged collectives on the non-captured steps — rejected.** Under + attention DP a mixed step (one rank prefilling beside three decoding) pads + every rank to `max(counts)`, so the expert call runs over `4 x max` rows of + which ~70% are zeros. `_dp_rows` was changed to return a `sizes` vector + whenever the counts disagree and to pass it to both collectives (the + uniform form is unreachable-by-construction under capture, where counts are + equal). The engine does publish true per-rank counts — + `_get_all_rank_num_tokens` is a plain `tp_allgather` of + `attn_metadata.num_tokens`, no padding — so the branch is reachable, and + smoke passed. It measured **+0.23% at con=32 and +4.08% at con=256** + against the same perturbed baseline, i.e. **+0.10% against the clean + cluster**: 7634.91 against 7635.26 (`max_seq_len` probe) and 7627.45 + (full baseline). The padded rows are real but are not a material share of + GPU time on this workload; the assembler's uniform-only choice stands. +* **`stream_interval`, `cuda_graph_config.*`, `kv_cache_config.*`, + `scheduler_config.capacity_scheduler_policy` — not swept, on profile + evidence.** GPU idle 1.49% leaves nothing for the host-side response path + to recover; graph replay coverage is already 100% at `enable_padding: + false`, so padding could only coarsen the grid (−2.55% at con=256 on the + `dep4` sibling); and the KV pool is 44.76 GiB = **1,368,000 tokens per + rank** against the 64 requests x 2048 tokens = 131,072 a rank holds at + con=256, a 10.4x margin, so capacity cannot bind. +* **`speculative_config.draft_len_schedule` — rejected twice over: the premise + is refuted, and on this target it deadlocks.** What was measured: + `configs/mtp3-sched.yaml`, `max_draft_len: 3` with + `draft_len_schedule: {2: 3, 16: 2, 128: 1}` (thresholds are per-*rank* batch + size, so under attention DP at dep4 they map to con 1..8 / 16..64 / + 128..256), full sweep attempted on GPUs 3-6. + + *The premise.* The schedule's case is "long drafts at low concurrency, + shorter at high", which needs the best draft length to fall as batch grows. + The three full sweeps say it does not: `mtp3` leads at **all nine** + concurrencies, and acceptance decays by only 3.49 -> 3.40 over the whole + range. There is no crossover inside the certified range, so a schedule that + shortens the draft at high batch can only give throughput away. + + *The deadlock.* The run never produced a point. Engine build, CUDA-graph + capture (one graph per `(batch_size, draft_len)` pair — the log shows + `batch size=3..16, draft_len=2` and `batch size=1,2, draft_len=3`) and warmup + all succeeded; roughly a minute into serving, **all four ranks hung and + trtllm's own HangDetector fired at 300 s and hard-killed via `MPI_Abort`**. + The stacks put the ranks at three *different* points of one forward — the q + up-projection, the MLA generation call, and the MoE `fp4_quantize` — i.e. not + in lockstep. The mechanism is structural: `_handle_dynamic_draft_len` + resolves `runtime_draft_len` from **`scheduled_batch.batch_size`, which is + each rank's own local batch**, with no cross-rank reduction, and attention DP + does not equalize batch sizes — `_pad_attention_dp_dummy_request` only tops a + rank up from zero to one. Two ranks either side of a threshold therefore run + different numbers of MTP replays, hence different numbers of MoE collectives, + and the job wedges. Nothing validates the combination at config time. + + *And it is not a single-variable comparison anyway*: setting the schedule + makes the runtime log `Automatically enabling cuda_graph_config.enable_padding + because draft_len_schedule is set`, flipping a knob this target measured as + costly (padding can only coarsen an already 100%-covered replay grid). The + variant file was deleted; this entry is its record. +* **`--acceptance` on the Pareto sweeps — rejected, and not needed.** The flag + sets `enable_iter_perf_stats: true` and lifts `iter_stats_max_iterations`, so + a label carrying it is not comparable to `baseline` / `iter1`. Priced in a + paired probe on one `mtp3` config, back to back, con 1/32/256: + **−1.39% / −3.52% / −3.96%** (212.34 -> 209.38, 2056.91 -> 1984.44, + 8229.25 -> 7903.35 tok/s). That is far outside the sub-1% floor, so it was + carried on no sweep. It is also unnecessary here: the benchmark client + already reports `avg_decoded_tokens_per_iter` per request, sourced from the + **response body** rather than the `/metrics` iteration-stats stream, so it + survives with the flag off — at con=1 the two probes recorded a bit-identical + `3.8863`, and against the engine's own `acceptance_length` the client proxy + agrees to **0.99-1.02x**. Every `accept` number in this section is that free + proxy, measured at zero config cost on the same sweeps as the throughput. +* **`max_draft_len` 1 and 2 — measured and dominated.** `mtp1` and `mtp2` are + full sweeps in the tables above; `mtp3` beats both at every concurrency. They + are kept as configs because they are the evidence that the axis is monotone, + not as recommended settings. + +### Caveats + +* **Session variance is measured, and the dominant term is co-tenancy, not + sampling.** Three unchanged-config measurements of con=256 landed at + 7627.45 (full sweep, 1280 requests), 7635.26 and 7634.91 (two-round probes, + 512 requests) — a **0.10% spread across two different request counts**. A + fourth, `probe-baseline`, read 7335.94 (**−3.9%**); an 8-GPU neighbour + holding 1.9 GiB on GPUs 0-6 — including this campaign's device set — was + caught in `nvidia-smi` minutes later and was gone within two. That probe is + the only contaminated measurement in the campaign and no kept result rests + on it. **Read the noise floor as well under 1% between clean runs, with a + transient co-tenant worth ~4%.** +* **Probes and full sweeps are not interchangeable at the throughput end for + TTFT.** The same config measured mean TTFT 1765.6 ms over 1280 requests and + 2279.7 ms over 512: with only two waves of 256 the first wave's queue + dominates the mean. Throughput is insensitive to this (0.10% above); TTFT + is not. Every TTFT comparison above is probe-to-probe or sweep-to-sweep. +* Host compute-apps and load average were recorded before and after every + label (`logs/hoststate/`); GPUs 3-6 carried no other tenant for any + measured curve. +* Both identity-config curves are single full sweeps; `iter1`'s two endpoints + were additionally reproduced in a paired probe before the sweep was run. +* `perf/data/` is machine-local and not tracked; the figure is. +* **The two sessions are stitched by measurement, not by assumption.** Opening + the 2026-08-02 session, `probe-anchor-iter1` re-measured the `iter1` config + and landed within **−0.14% (con=1) and +0.15% (con=256)** of the 2026-08-01 + full sweep. At wrap-up the `trtllm` reference was spot-checked the same way: + **−0.01% (con=1, TPOT bit-identical at 13.525 ms) and +0.20% (con=256)**. + Neither curve was re-swept. +* **This session's noise floor, measured rather than inherited.** `mtp1` at + con=256 read **8737.78** in its full sweep and **8736.8** in a re-measure + 2.5 hours later at the same request count (`probe-recheck-mtp1`) — + **0.01%**. The same pair at **con=128 differs by 4.2%** (5877.32 vs 5630.1), + so con=128 is this target's noisy point and no claim above rests on it + alone. The `mtp3`-over-`mtp2` margin at con=256 is +3.29%, comfortably + outside the con=256 floor. +* **Probes and full sweeps are much further apart under MTP than without it, + and it is throughput this time, not just TTFT.** The same `mtp3` config + measured con=256 at **8229.25** over 512 requests and **9048.23** over 1280 — + **+9.95%** — where the identity config's probe-to-sweep gap at that point is + 0.15%. A two-round probe under-reports MTP because the ramp and drain, where + batches are small and drafting is least profitable, are a much larger share + of a 512-request point. **Every MTP comparison in this section is + sweep-to-sweep**; an early probe-vs-sweep reading of these curves inverted + the `mtp1`/`mtp3` ordering at con=256 before the full sweeps corrected it. +* **Acceptance at con=1 is prompt-dependent and needs more than a two-request + probe.** The `mtp3` full sweep (20 requests at con=1) reads **3.4904**, while + the two-request probes recorded under *The MTP variant's own gate records* + read 3.8863/3.9274. Both are far above the 1.164 break-even, so the + correctness reading there is unaffected, but the sweep number is the better + estimate of acceptance. +* **A cold host page cache can push the stock-trtllm boot past the harness's + 900 s limit, and it looks like a failure rather than a slow start.** The + first `trtllm-mtp3` attempt returned `[perf] FAILED: server not healthy + after 900s`. It was not hung — it had reached `[Autotuner] Autotuning + process starts` and was still progressing. The time went into reading the + 397 GB checkpoint: the log shows 101 s, 57 s and 38 s gaps between + `Finished prefetching ...` lines, and boot took **837 s to reach the + autotuner** against **220 s** for the same config earlier the same day. + Re-reading the 163 shards took 7 s (~56 GB/s, i.e. already resident, the + failed boot having warmed them), and the retry reached the autotuner in + **219 s** and completed the full sweep. **Read a 900 s boot timeout on this + checkpoint as a page-cache miss first**; warm the shards and retry before + suspecting the model. Nothing in the harness was changed. +* **Host CPU load moved a lot and the measurement did not.** Load average over + the campaign ranged 0.31 to 13.72, and a neighbour ran on GPUs 1-2 at 100% + during wrap-up; the `trtllm` spot-check taken under the *heaviest* load still + reproduced to 0.20%. The previous campaign's −3.9% contamination came from a + co-tenant **on the device set itself**, which is the case to keep watching. + +### Regenerating the figure + +The committed figure plots exactly seven curves — the two **aligned** +stock-trtllm references and the five staircase sweeps — in this order: + +``` +uv run utils/plot_pareto.py \ + targets/deepseek-r1-0528-nvfp4/sm_100/dep4/perf/data/trtllm-aligned \ + targets/deepseek-r1-0528-nvfp4/sm_100/dep4/perf/data/trtllm-aligned-mtp3 \ + targets/deepseek-r1-0528-nvfp4/sm_100/dep4/perf/data/baseline \ + targets/deepseek-r1-0528-nvfp4/sm_100/dep4/perf/data/iter1-shared-side-stream \ + targets/deepseek-r1-0528-nvfp4/sm_100/dep4/perf/data/mtp1 \ + targets/deepseek-r1-0528-nvfp4/sm_100/dep4/perf/data/mtp2 \ + targets/deepseek-r1-0528-nvfp4/sm_100/dep4/perf/data/mtp3 \ + -o targets/deepseek-r1-0528-nvfp4/sm_100/dep4/perf/figures/pareto.png +``` + +**Explicit paths, not the `perf/data` directory.** That directory also holds +the two acceptance-only gate-record labels and the three superseded +boot-forced reference lines; expanding it would put all of them on the figure. +The committed figure plots the two aligned references plus the five staircase +curves. + +**`trtllm-nodeepep` is on the figure on purpose even though it is only three +points.** It is the fastest stock configuration measured here, and it sits +almost on top of `baseline` / `iter1`; leaving it off would let the figure +imply a larger no-MTP lead over stock trtllm than exists. + +**The paths are explicit on purpose: pointing the plotter at `perf/data` +wholesale is wrong here.** That directory also holds `probe-mtp3-acceptance` +and `trtllm-mtp3-acceptance` — a two-point probe and an acceptance-only +reference from the MTP gate records, neither of them a Pareto curve — and the +plotter picks up every subdirectory carrying a `meta.json`, which would put a +second grey reference line on the figure and blow past its categorical slots. + +The `trtllm` curve needs its boot config: +`uv run bench/perf.py --target targets/deepseek-r1-0528-nvfp4/sm_100/dep4 +--label trtllm --trtllm --config +targets/deepseek-r1-0528-nvfp4/sm_100/dep4/configs/trtllm-ref-boot.yaml`. +The three MTP curves are +`uv run bench/perf.py --target --label mtp --config /configs/mtp.yaml` +— no `--acceptance`. + +### Remaining headroom + +* **The `max_draft_len` curve is still climbing where the certification + stops.** `mtp3` is the best of the three at every concurrency and the gain + from 2 to 3 is still positive everywhere (con=1 162.51 -> 187.12, con=256 + 8760.17 -> 9048.23), while acceptance holds at 3.40 of a 4.0 ceiling — 2.9x + the 1.164 break-even. Nothing in the *data* says 3 is the optimum; 3 is the + top of what `catalog/attention/mla_rope_generation` and + `catalog/attention/thop_attention` certify, which is + `predicted_tokens_per_seq` ∈ {1,2,3,4} and hence `max_draft_len` ≤ 3. **A + campaign that wants 4 or beyond needs those two entries certified at + `predicted_tokens_per_seq` 5+ first — a vocabulary request, not a bigger + number in a config.** The step-cost series says the return is decelerating + (marginal cost per draft step at con=256: +0.645, +0.569, +0.386 against + acceptance gains of +0.95, +0.80, +0.65 tokens), so the crossover is probably + not far past 3, but it is unmeasured. +* **The `mtp3` cost decomposition is arithmetic, not a timeline.** The + memory-bound-to-compute-bound reading above is derived from the step-cost + series across four draft lengths and nine concurrencies; no nsys capture was + taken under MTP. A timeline at con=256 with `max_draft_len: 3` would say + *which* kernel family absorbed the growth — the routed expert GEMMs at 4x + rows, the bf16 MTP expert stack, or the extra collectives — and that is the + next thing to profile if the MTP path is tuned further. +* **The collectives are the one identified lever and they are blocked on + vocabulary.** **Corrected 2026-08-03:** the 3369.7 µs / 94.3%-exclusive + figure below is the **pre-`iter1`** profile — it is the measurement that + motivated the side stream ("a 3.37 ms window with nothing in it"), not the + state after it. A post-`iter1` capture at the same window and phase measures + **1930.7 µs per rank per decode step at 55.2% exclusive**: the side stream + moved the shared expert into that window, so the recoverable time is 57% of + what this paragraph claimed. Sizing a transport change off 3369.7 overstates + it. The pre-`iter1` numbers, kept because the rest of the analysis rests on + them: **3369.7 µs per rank per decode step of exclusive GPU time, 11.9% of + the step, at 94.3% exclusive**, in two + `RING_LL` NCCL calls per MoE layer. `comm/allgather` and + `comm/reducescatter` expose **no strategy argument and no workspace**, so + the ONESHOT swap that won 2.03x per call on the `tep4` sibling does not + exist here. Closing it needs new catalog entries: a strategy-carrying + gather/scatter, or the `moe_a2a_dispatch` / `moe_a2a_combine` family (a + stateful workspace pair, so `moe_a2a_initialize` and + `moe_a2a_get_combine_payload_tensor` come with it). Note `iter1` has + already spent part of this window — a transport change would now compete + with the shared expert for it, so the two do not simply add. +* **Reducing collective *bytes* is not the lever.** At con=256 each rank's + decode batch is 64 rows, so one all-gather receives `3 x 64 x 7168 x 2` = + **2.753 MB** in its measured 25.472 µs = **108 GB/s**, an order of + magnitude under NVLink: `RING_LL` at these sizes is latency-bound, not + bandwidth-bound. Gathering NVFP4 activations instead of bf16 (4032 vs + 14336 B/token, **3.56x** fewer) would buy nothing and would cost extra + calls. +* **The routed expert GEMMs are near the HBM roofline and are the floor.** + 46.3% of GPU busy, 98.9% exclusive. With top-8 of 256 experts over 256 + gathered tokens every one of a rank's 64 local experts is active, so a step + reads the whole local stack: FC1 (gate+up) is `2 x 2048 x 7168 x 64` at + 0.5625 B/element (NVFP4 data plus the fp8 block scale) = **1.057 GB in + 155.25 µs = 6.81 TB/s**; FC2 (down) is 0.528 GB in 77.97 µs = **6.78 + TB/s** — ~85% of a ~8 TB/s B200 roofline. No config knob and no scheduling + change moves this; only fewer weight bytes would. +* **The bf16 attention GEMMs are the second floor**: 20.1% of GPU busy. The + largest, at 45.21 µs x 61 layers per rank-step, reads `o_proj`'s + 16384x7168 bf16 (235 MB) at **5.19 TB/s** — the one family with visible + daylight to the roofline, but it is cuBLAS's tactic choice, not ours. + These weights are bf16 by checkpoint design (`hf_quant_config.json` + quantizes the MLP, not attention), so this is a checkpoint property rather + than a target one. diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/__init__.py b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/__init__.py new file mode 100644 index 000000000000..6319f720487b --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/__init__.py @@ -0,0 +1,3 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""DeepSeek-R1-0528-NVFP4 / sm_103 / dep4.""" diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/identity.yaml b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/identity.yaml new file mode 100644 index 000000000000..d800d2e26e16 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/identity.yaml @@ -0,0 +1,21 @@ +# The identity assembly: the target with no variant on top. +# +# This is what the gsm8k record in TARGET.md was measured on. It carries only +# the parallel topology, and the topology is not a tuning knob here -- routing +# derives this target *from* it, so a run that drops it does not get a +# differently-tuned staircase, it gets the built-in DeepseekV3 implementation +# (or, under `TRTLLM_STAIRCASE=require`, an error naming the criterion it +# missed). +# +# Every other file in this directory is this one plus the knobs under +# experiment, so any of them can be passed directly: +# +# TRTLLM_STAIRCASE=require trtllm-eval --model \ +# --extra_llm_api_options /identity.yaml \ +# gsm8k --apply_chat_template --fewshot_as_multiturn +# +# The switch has to be exported before the ranks start; see +# `_router_index.STAIRCASE_ENV` for why. +tensor_parallel_size: 4 +moe_expert_parallel_size: 4 +enable_attention_dp: true diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp1.yaml b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp1.yaml new file mode 100644 index 000000000000..ad119802debc --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp1.yaml @@ -0,0 +1,69 @@ +# Multi-token prediction, 1 draft token per step. +# +# This variant crosses the line a config variant normally respects: it does +# not just move a knob, it selects a **second forward path and a different +# weight-loading path**. With it, the checkpoint's bf16 MTP module at layer 61 +# stops being a predicted non-load and is loaded (212 keys per rank, ~6.3 GB), +# the engine raises the KV pool to 62 layers, and every engine step runs the +# trunk plus `max_draft_len` replays of the MTP layer. Without it the target is +# the identity assembly the gsm8k record was measured on, bit for bit. +# +# decoding_type: MTP + `num_nextn_predict_layers: 1` in the checkpoint config +# resolves to MTP_EAGLE_ONE_MODEL — one MTP layer, replayed autoregressively. +# The layer count is therefore NOT the draft-length knob. +# +# max_draft_len is spelled out on purpose: left unset on the MTP-Eagle path it +# resolves to 1, not to anything workload-derived. It is also the axis that +# costs first-boot time — it becomes the decode kernel's maxSeqLenQ, so each +# value JIT-compiles its own MLA decode kernel per rank (1 and 2 pay separate +# ~5.5 s compiles despite sharing a kernel name; 3 and 4 share another). +# +# Break-even: one draft step adds ~5.5% to a decode step's HBM traffic on this +# checkpoint at dep4, so drafting pays for itself at an acceptance_length of +# about 1.055 here. +# +# **Measured, and dominated — this is not the recommended setting.** A full +# Pareto sweep of all three draft lengths in one session (perf/data/mtp{1,2,3}, +# GPUs 3-6) puts `max_draft_len: 3` ahead at every one of the nine +# concurrencies. This file's curve is 1.63x iter1 at con=1 and 1.12x at +# con=256, against mtp3's 2.29x and 1.16x. Acceptance here is ~1.95 of a 2.0 +# ceiling and barely moves with batch size (1.967 at con=1, 1.949 at con=256). +# Keep it as the low end of the measured axis; use configs/mtp3.yaml to serve. + +# The parallel topology. It is what selects this target at all -- routing +# derives the target from it -- so every variant carries it and these files +# are usable as-is: `trtllm-eval --extra_llm_api_options `. +tensor_parallel_size: 4 +moe_expert_parallel_size: 4 +enable_attention_dp: true + +speculative_config: + decoding_type: MTP + max_draft_len: 1 + +# **This variant does not boot at trtllm's default KV-cache sizing.** Measured +# on this checkpoint at dep4 on a 183.4 GiB B200, with max_draft_len 3: +# +# The MTP module adds 5.87 GiB of declared weights per rank (nvidia-smi read +# 122,076 MiB after model init), and the engine then sizes the KV pool from +# what is free, at +# the default free_gpu_memory_fraction 0.9: 38.95 GiB (1,171,200 tokens, +# 62 layers). Boot then reaches CUDA-graph capture at 182.4 of 183.4 GiB and +# livelocks: three of four ranks spin in cudaFree inside the CUDA caching +# allocator (release_cached_blocks, reached from a plain empty_cuda in the +# *trunk's* MoE runner) while the fourth waits at an MPI_Barrier. No error, +# no OOM exception, no progress -- the run has to be killed. +# +# The drafting forward's post-pool transient is simply larger than the ~19.6 +# GiB the identity config leaves and lives inside. Freeing memory *before* +# the pool is sized does not help: chunking the MTP layer's expert call to +# 2048 rows moved pool sizing by 0.31 GiB (38.64 -> 38.95) and the boot +# failed identically, because the pool grows into whatever is freed. +# +# 0.75 hands the pool ~32.5 GiB instead, leaving ~28 GiB of headroom. That is +# still ~975k KV tokens per rank -- an order of magnitude more than any +# concurrency this target is measured at needs -- so nothing about the +# workload is constrained by it. It is a boot requirement of the variant, not +# a tuning choice, and it is the reason these files carry a second key at all. +kv_cache_config: + free_gpu_memory_fraction: 0.75 diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp2.yaml b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp2.yaml new file mode 100644 index 000000000000..704879fd3b76 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp2.yaml @@ -0,0 +1,69 @@ +# Multi-token prediction, 2 draft tokens per step. +# +# This variant crosses the line a config variant normally respects: it does +# not just move a knob, it selects a **second forward path and a different +# weight-loading path**. With it, the checkpoint's bf16 MTP module at layer 61 +# stops being a predicted non-load and is loaded (212 keys per rank, ~6.3 GB), +# the engine raises the KV pool to 62 layers, and every engine step runs the +# trunk plus `max_draft_len` replays of the MTP layer. Without it the target is +# the identity assembly the gsm8k record was measured on, bit for bit. +# +# decoding_type: MTP + `num_nextn_predict_layers: 1` in the checkpoint config +# resolves to MTP_EAGLE_ONE_MODEL — one MTP layer, replayed autoregressively +# `max_draft_len` times. The checkpoint's layer count does not bound this. +# +# max_draft_len is spelled out on purpose: left unset on the MTP-Eagle path it +# resolves to 1, not to anything workload-derived. It also sets the generation +# call's predicted_tokens_per_seq (= max_draft_len + 1 = 3 here), which the MLA +# entries certify at 1..4 and which becomes the decode kernel's maxSeqLenQ — +# one extra JIT compile per rank on first boot. +# +# Break-even: two draft steps add ~11% to a decode step's HBM traffic on this +# checkpoint at dep4, so drafting pays for itself at an acceptance_length of +# about 1.110 here. +# +# **Measured, and dominated — this is not the recommended setting.** A full +# Pareto sweep of all three draft lengths in one session (perf/data/mtp{1,2,3}, +# GPUs 3-6) puts `max_draft_len: 3` ahead at every one of the nine +# concurrencies. This file's curve is 1.98x iter1 at con=1 and 1.12x at +# con=256, against mtp3's 2.29x and 1.16x. Acceptance here is ~2.75 of a 3.0 +# ceiling and barely moves with batch size (2.748 at con=1, 2.754 at con=256). +# Keep it as the middle of the measured axis; use configs/mtp3.yaml to serve. + +# The parallel topology. It is what selects this target at all -- routing +# derives the target from it -- so every variant carries it and these files +# are usable as-is: `trtllm-eval --extra_llm_api_options `. +tensor_parallel_size: 4 +moe_expert_parallel_size: 4 +enable_attention_dp: true + +speculative_config: + decoding_type: MTP + max_draft_len: 2 + +# **This variant does not boot at trtllm's default KV-cache sizing.** Measured +# on this checkpoint at dep4 on a 183.4 GiB B200, with max_draft_len 3: +# +# The MTP module adds 5.87 GiB of declared weights per rank (nvidia-smi read +# 122,076 MiB after model init), and the engine then sizes the KV pool from +# what is free, at +# the default free_gpu_memory_fraction 0.9: 38.95 GiB (1,171,200 tokens, +# 62 layers). Boot then reaches CUDA-graph capture at 182.4 of 183.4 GiB and +# livelocks: three of four ranks spin in cudaFree inside the CUDA caching +# allocator (release_cached_blocks, reached from a plain empty_cuda in the +# *trunk's* MoE runner) while the fourth waits at an MPI_Barrier. No error, +# no OOM exception, no progress -- the run has to be killed. +# +# The drafting forward's post-pool transient is simply larger than the ~19.6 +# GiB the identity config leaves and lives inside. Freeing memory *before* +# the pool is sized does not help: chunking the MTP layer's expert call to +# 2048 rows moved pool sizing by 0.31 GiB (38.64 -> 38.95) and the boot +# failed identically, because the pool grows into whatever is freed. +# +# 0.75 hands the pool ~32.5 GiB instead, leaving ~28 GiB of headroom. That is +# still ~975k KV tokens per rank -- an order of magnitude more than any +# concurrency this target is measured at needs -- so nothing about the +# workload is constrained by it. It is a boot requirement of the variant, not +# a tuning choice, and it is the reason these files carry a second key at all. +kv_cache_config: + free_gpu_memory_fraction: 0.75 diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp3.yaml b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp3.yaml new file mode 100644 index 000000000000..eb17e13071f3 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp3.yaml @@ -0,0 +1,79 @@ +# Multi-token prediction, 3 draft tokens per step. +# +# This variant crosses the line a config variant normally respects: it does +# not just move a knob, it selects a **second forward path and a different +# weight-loading path**. With it, the checkpoint's bf16 MTP module at layer 61 +# stops being a predicted non-load and is loaded (212 keys per rank, ~6.3 GB), +# the engine raises the KV pool to 62 layers, and every engine step runs the +# trunk plus `max_draft_len` replays of the MTP layer. Without it the target is +# the identity assembly the gsm8k record was measured on, bit for bit. +# +# decoding_type: MTP + `num_nextn_predict_layers: 1` in the checkpoint config +# resolves to MTP_EAGLE_ONE_MODEL — one MTP layer, replayed autoregressively +# `max_draft_len` times. The checkpoint's layer count does not bound this. +# +# max_draft_len is spelled out on purpose: left unset on the MTP-Eagle path it +# resolves to 1, not to anything workload-derived. It also sets the generation +# call's predicted_tokens_per_seq (= max_draft_len + 1 = 4 here), which is the +# top of what the MLA entries certify (1..4) — **going past 3 needs the +# certification extended first**, not just a bigger number here. It becomes the +# decode kernel's maxSeqLenQ too, so it costs one extra JIT compile per rank on +# first boot. +# +# Break-even: three draft steps add ~16.4% to a decode step's HBM traffic on +# this checkpoint at dep4, so drafting pays for itself at an acceptance_length +# of about 1.164 here. +# +# **This is the recommended MTP setting, and the measured best of the three.** +# A full Pareto sweep of all three draft lengths in one session +# (perf/data/mtp{1,2,3}, GPUs 3-6) puts this file ahead at every one of the +# nine concurrencies, on both frontier axes: con=1 tok/s/user 82.43 -> 190.45 +# and peak tok/s/GPU 1948.78 -> 2262.06 against the iter1 frontier. Measured +# acceptance is 3.39-3.49 of a 4.0 ceiling and essentially flat in batch size, +# so the gain decays with concurrency (2.29x at con=1, 1.16x at con=256) +# entirely because the step gets more expensive, not because drafting gets +# worse. **The curve is still climbing at 3** — 3 is the top of what the two +# MLA entries certify (predicted_tokens_per_seq <= 4), not an optimum. +# +# Note the throughput numbers above are measured on the harness's random-prompt +# workload, which *flatters* MTP: only the prompt is random, and the model's +# own continuation of it is repetitive and drafts near-perfectly. The real-text +# figure, from the paired gsm8k gate runs, is 1.178x. + +# The parallel topology. It is what selects this target at all -- routing +# derives the target from it -- so every variant carries it and these files +# are usable as-is: `trtllm-eval --extra_llm_api_options `. +tensor_parallel_size: 4 +moe_expert_parallel_size: 4 +enable_attention_dp: true + +speculative_config: + decoding_type: MTP + max_draft_len: 3 + +# **This variant does not boot at trtllm's default KV-cache sizing.** Measured +# on this checkpoint at dep4 on a 183.4 GiB B200, with max_draft_len 3: +# +# The MTP module adds 5.87 GiB of declared weights per rank (nvidia-smi read +# 122,076 MiB after model init), and the engine then sizes the KV pool from +# what is free, at +# the default free_gpu_memory_fraction 0.9: 38.95 GiB (1,171,200 tokens, +# 62 layers). Boot then reaches CUDA-graph capture at 182.4 of 183.4 GiB and +# livelocks: three of four ranks spin in cudaFree inside the CUDA caching +# allocator (release_cached_blocks, reached from a plain empty_cuda in the +# *trunk's* MoE runner) while the fourth waits at an MPI_Barrier. No error, +# no OOM exception, no progress -- the run has to be killed. +# +# The drafting forward's post-pool transient is simply larger than the ~19.6 +# GiB the identity config leaves and lives inside. Freeing memory *before* +# the pool is sized does not help: chunking the MTP layer's expert call to +# 2048 rows moved pool sizing by 0.31 GiB (38.64 -> 38.95) and the boot +# failed identically, because the pool grows into whatever is freed. +# +# 0.75 hands the pool ~32.5 GiB instead, leaving ~28 GiB of headroom. That is +# still ~975k KV tokens per rank -- an order of magnitude more than any +# concurrency this target is measured at needs -- so nothing about the +# workload is constrained by it. It is a boot requirement of the variant, not +# a tuning choice, and it is the reason these files carry a second key at all. +kv_cache_config: + free_gpu_memory_fraction: 0.75 diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/trtllm-ref-boot.yaml b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/trtllm-ref-boot.yaml new file mode 100644 index 000000000000..9641a0700523 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/trtllm-ref-boot.yaml @@ -0,0 +1,30 @@ +# Boot-forcing config for the stock-trtllm reference curve ONLY. +# +# Stock trtllm (in-tree DeepseekV3 modeling) cannot boot this checkpoint at +# dep4 on a 178.34 GiB B200 under its own defaults. Measured, per rank: +# +# max_num_tokens 8192 (default): 113.9 GiB is allocated at engine +# construction outside the torch allocator, before any weight is +# materialized (nvidia-smi 20 MiB -> 116,662 MiB inside one 5 s sample), +# and model init then OOMs in init_meta_tensor. +# max_num_tokens 2048: weights load (135.85 GiB torch, incl. 1.53 GiB of +# CUDA-graph pools), then configure_kv_cache_capacity asks for 7.00 GiB +# with 5.61 GiB free. kv_cache_config.free_gpu_memory_fraction does not +# move that 7.00 GiB (measured identical at 0.9 and 0.6). +# +# max_seq_len is the remaining lever the failure's own memory ledger names: +# it sets max_blocks_per_seq = max_seq_len / tokens_per_block, 5120 at the +# checkpoint's 163840 and 64 here, and 2048 is exactly the harness's ISL+OSL. +# +# Not a tuning variant for this target: staircase boots at stock defaults, and +# both knobs measured inert on it (max_seq_len 2048: -0.13% / +0.10%). + +# The parallel topology. It is what selects this target at all -- routing +# derives the target from it -- so every variant carries it and these files +# are usable as-is: `trtllm-eval --extra_llm_api_options `. +tensor_parallel_size: 4 +moe_expert_parallel_size: 4 +enable_attention_dp: true + +max_num_tokens: 2048 +max_seq_len: 2048 diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/trtllm-ref-mtp3.yaml b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/trtllm-ref-mtp3.yaml new file mode 100644 index 000000000000..9e726cc9df75 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/trtllm-ref-mtp3.yaml @@ -0,0 +1,75 @@ +# Stock-trtllm reference line for the MTP variant. NOT a tuning variant. +# +# This is `trtllm-ref-boot.yaml` + `mtp3.yaml`, in one file because the perf +# harness takes a single --config. It exists to answer one question that +# nothing else can: **is this target's MTP layer computing what the +# checkpoint says?** +# +# Rejection sampling makes that question invisible to every other gate. A +# miscomputed draft layer still emits the target model's exact distribution — +# it just has every draft rejected, so the boot gate passes, the accuracy gate passes, +# and the only thing that moves is speed. The acceptance rate is the sole +# detector, and a rate is only readable against a reference. That reference is +# stock in-tree modeling on the same checkpoint, at the same workload and the +# same speculative_config. +# +# Why the boot knobs ride along: stock trtllm cannot boot this checkpoint at +# dep4 under its own defaults (see trtllm-ref-boot.yaml for the measured +# ledger), and MTP adds ~5.9 GiB of layer-61 weights per rank on top of that. +# The knobs change batching capacity, not what the model predicts — +# **acceptance is a property of the model**, so a reference forced onto +# different memory knobs is still the right comparison. Only the workload has +# to match, and it does: bench/perf.py seeds prompts by concurrency value, so +# both systems see byte-identical requests. + +# The parallel topology. It is what selects this target at all -- routing +# derives the target from it -- so every variant carries it and these files +# are usable as-is: `trtllm-eval --extra_llm_api_options `. +tensor_parallel_size: 4 +moe_expert_parallel_size: 4 +enable_attention_dp: true + +max_num_tokens: 2048 +max_seq_len: 2048 + +# Must match configs/mtp3.yaml exactly — a reference at a different draft +# length answers a different question. +speculative_config: + decoding_type: MTP + max_draft_len: 3 + +# The staircase side needs 0.75 to boot with MTP at all (configs/mtp3.yaml has +# the livelock ledger). Stock trtllm carries strictly more memory pressure +# here, so it starts from the same value rather than from the default. +kv_cache_config: + free_gpu_memory_fraction: 0.75 + +# ONE MORE THING IS NEEDED AND IT IS NOT A CONFIG KEY. +# +# Export `TRTLLM_CAN_USE_DEEP_EP=0` before launching, or this does not boot: +# +# _torch/modules/fused_moe/communication/deep_ep_low_latency.py:238, dispatch +# assert hidden_states.dtype == torch.uint8 +# AssertionError -- all four ranks, "Failed to initialize executor" +# +# Under attention DP + EP the MoE communication factory picks a strategy by +# priority. Measured here: NVLinkOneSided, NVLinkTwoSided and DeepEP each log +# `not available: Invalid Argument`, leaving DeepEPLowLatency -- whose dispatch +# accepts only NVFP4 (uint8-packed) hidden states. The MTP layer is **bf16**, +# because modelopt excludes `model.layers.61*` from quantization wholesale. +# TARGET.md's existing `trtllm` curve went through DeepEPLowLatency happily, +# because without MTP every MoE layer is NVFP4. +# +# The variable disables DeepEP and DeepEPLowLatency together and lands on +# AllGatherReduceScatter, which the selector's own comment calls "always +# works" -- and which is the strategy this target implements by hand, so it +# makes the two systems more numerically comparable rather than less. +# +# It must be exported **before** the server subprocess is spawned: OpenMPI +# hands a spawned process the environment as it stood when MPI initialized. +# bench/perf.py prepares the server's environment and then spawns, so an +# export in the calling shell reaches every rank. +# +# Consequence to state wherever this label is used: its **throughput** is not +# comparable to the `trtllm` Pareto curve, which carries neither this +# transport nor this draft length. This config measures acceptance. diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py new file mode 100644 index 000000000000..97bc07b765cd --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py @@ -0,0 +1,2229 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Staircase target: deepseek-r1-0528-nvfp4 / sm_103 / dep4 — self-contained modeling code. + +Flat single-entry forward assembled from catalog entries only; every call +that creates or transforms a tensor is a catalog entry, everything else is +tensor-metadata reads and Python control flow. + +DeepSeek-R1-0528, modelopt NVFP4 export: 61 layers, hidden 7168, 128 query +heads, `q_lora_rank` 1536, layers 0-2 dense (intermediate 18432), layers 3-60 +MoE with 256 routed experts at top-8 plus one shared expert (intermediate +2048). The quantization is `nvfp4_moe_only`-shaped — every `self_attn*` and +`lm_head` is excluded from NVFP4, so attention, router, embedding and lm_head +are bf16 and only the MLP path is NVFP4 — **and the KV cache is fp8-e4m3** +(`hf_quant_config.json`: `kv_cache_quant_algo: FP8`). + +Against the `deepseek-v3-lite-nvfp4/dep4` sibling (not in this batch) — same architecture +family, same topology, same quantization shape — five things differ, and each +is re-derived here rather than inherited: + +* **the query path is a LoRA pair.** `q_lora_rank: 1536`, so the query is + `q_b_proj(q_a_layernorm(q_a_proj(x)))` — two GEMMs and a norm — where the + lite checkpoint's `q_lora_rank: null` projects directly. Nothing downstream + changes: `q_b_proj` produces the same `[T, H*(nope+rope)]` rows; +* **routing is group-limited.** `n_group 8`, `topk_group 4`: `noaux_tc_op` + scores 8 contiguous groups of 32 experts by their two best bias-corrected + scores, keeps the best 4 groups, and takes top-8 inside them. The lite + sibling runs the ungrouped case. The MoE runner's own `n_group`/`topk_group` + stay inert — routing is done before the call, on either checkpoint; +* **the rope table is YaRN-scaled** (factor 40 over an original 4096-position + window, `beta_fast` 32, `beta_slow` 1), and the model's YaRN attention + temperature `mscale = 0.1*ln(40) + 1` rides in **`q_scaling = 1/mscale^2`** + rather than in the table: the table's own amplitude is + `m(mscale)/m(mscale_all_dim)` = exactly 1.0 because this config sets both to + 1.0. The op reads the table's *content* and `q_scaling`; the seven scalar + rope arguments beside them are inert (measured — see thop_attention.md); +* **the latent KV pool is fp8-e4m3**, which changes what every one of the five + MLA-family calls a layer issues does — see below; +* **scale.** 128 query heads (attention DP replicates them, so every rank runs + all 128), 64 routed experts per rank, and a 163840-position rope table. + +**Layer 61 — the checkpoint's bf16 MTP module — is a second forward path this +file also carries, and it exists only when a `configs/` variant turns it on.** +Under the target's identity config (`llm_args.yaml`, no `speculative_config`) +nothing below MTP is declared, layer 61's 790 keys stay a predicted non-load in +the weight manifest, and the shell's forward is a plain call to the inherited +base — bit-identical to the assembly the accuracy gate was measured on. With +`configs/mtp{1,2,3}.yaml` the engine resolves a `spec_config` onto the model +config before the model is built, the core declares the module's parameters, +and the shell builds a draft-model container plus the runtime's spec worker. +`MTPLayer` below is the module's forward; `docs/models/multi-token-prediction.md` +is its semantics (the checkpoint ships no reference implementation of it and +neither does transformers), and `docs/references/trtllm-runtime-integration.md` +§13 is the runtime binding. Two consequences the rest of this file carries: +`predicted_tokens_per_seq` is per call site rather than inert, because a +generation request under MTP arrives carrying its whole draft chain; and the +`spec_decoding_*` group stays inert, because on a trtllm-gen arch a +linear-tree draft has +its mask machinery forced off and drafting reaches the attention ops through +`predicted_tokens_per_seq` alone. + +**The fp8 latent pool, and the four things it moves.** Nothing validates the +fp8 round trip at any layer: the write scale, the read scale and the two +folded FMHA scales are independent roles with no relation checked anywhere, so +a mistake here is silently mis-scaled output rather than an error. + +* **the context append and the cache gather** take the KV scaling factor `s` + as a write-side `1/s` and a read-side `s`. This checkpoint's 122 per-layer + `k_scale`/`v_scale` tensors are all exactly 1.0 (loaded and asserted in + `derive_after_load`), and `s = 1.0` is the **only** correct value on the fp8 + MLA context path — both context flavors quantize q/k/v at 1.0 while applying + `s^2`/`s` as if they had not, so any other `s` is silently wrong. Both scale + arguments are therefore passed as `None`, which the ops read as exactly 1.0 + and which is what the engine's own call sites pass; +* **the context FMHA quantizes its own q/k/v to e4m3**, in both flavors. So + prefill accuracy *is* affected by cache quantization, and a cached-prefix + context call pays fp8 twice — the gather dequantizes the cached latent rows + off the pool and the FMHA quantizes the up-projected result straight back; +* **the decode producers are ordered, not concurrent.** `mla_rope_generation` + does not write `fused_q` here — it **reads** `fused_q[..., :C]` to build + `quant_q_buffer`, the query the decode FMHA actually consumes. The absorbed-q + BMM must have finished before it. This forward issues both on the ambient + stream in that order, which is what makes it safe; the bf16 reading (the two + producers write disjoint halves and may overlap) is a silent race here; +* **the decode FMHA reads neither kv scale tensor.** Its query comes from + `quant_q_buffer` and its two scales from `mla_bmm1_scale[1]` and + `mla_bmm2_scale[0]`, both written by `mla_rope_generation` from `q_scaling`, + the MLA dims and the read-side factor. The caller owns the dequantization + entirely. + +**The parallel segment.** `dep4` is `tensor_parallel_size: 4` plus +`moe_expert_parallel_size: 4` plus `enable_attention_dp: true`: the requests +are split, not the heads. + +* **attention, the q-LoRA pair, both dense MLP shapes, the shared expert, the + router, the embedding, the norms and the residual stream are replicated**, + and each rank runs them over its own tokens only — 128 query heads per rank. + A rank's `o_proj` output is complete, so there is no attention-side + collective at all; +* **the MoE stays expert-parallel** (`moe_tp_size == 1`): rank `r` holds + experts `[64r, 64r+64)` whole. Every token must reach every window, so the + rank's tokens are **gathered** before the router and the four windows' + partials are **reduce-scattered** back — one `comm/allgather` and one + `comm/reducescatter` per MoE layer, 116 per forward, none in layers 0-2; +* **the gather goes before the router GEMM, and that placement is + load-bearing.** Expert parallelism rests on the four windows tiling the + routing space exactly once, which needs every rank to select the same experts + for the same token. Ranks hold different tokens here, so the invariant is + restored by routing on the gathered full token set: identical bytes in, + replicated deterministic router GEMM and `noaux_tc_op`, identical top-8 out. + Routing locally and gathering afterwards would break it with no error; +* **the reduce-scatter is crossed in bf16**: the op sums, and it sums + `float8_e4m3fn` as raw bytes rather than as floats, so a post-quantization + return trip is silently wrong (the gather is a byte move and would survive + it — the asymmetry is the trap); +* **`lm_head` is replicated** — the shell builds the whole `[vocab, hidden]` on + every rank under attention DP, because a rank's logits rows are its own + tokens' and no other rank computed them. + +The expert call is **chunked** to `_MOE_MAX_T` rows: the gathered token set +reaches `4 * max_num_tokens` = 32768 and both `fp4_block_scale_moe_runner` and +`noaux_tc_op` are certified to 8192. That bound is a certification boundary, +not a tuning knob. + +Weights are target-owned: a flat ParameterDict declared here (HF [out, in] +storage so checkpoint rows copy in unchanged, plus the kernel-ready expert +stacks weights.py builds during the load), with the column-major GEMM views, +the MLA absorption operands, the rope table and every NVFP4 call scalar +derived once after load. The registration shell inherits +DecoderModelForCausalLM for lm_head, packed-batch logits gathering, and the +meta-init/load/post-load hooks. + +The import-time and first-forward contract checks below fail fast on drift. +""" + +import math +from typing import cast + +import torch +from torch import nn +from transformers import PretrainedConfig + +from tensorrt_llm._torch.attention.backends.interface import AttentionMetadata +from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata +from tensorrt_llm._torch.model_config import ModelConfig +from tensorrt_llm._torch.models.modeling_utils import ( + DecoderModel, + DecoderModelForCausalLM, + register_auto_model, +) +from tensorrt_llm._torch.speculative import get_spec_worker +from tensorrt_llm._torch.staircase.catalog.activation.flashinfer_silu_and_mul import ( # noqa: E501 + flashinfer_silu_and_mul, +) +from tensorrt_llm._torch.staircase.catalog.attention.load_paged_kv_cache_for_mla import ( # noqa: E501 + load_paged_kv_cache_for_mla, +) +from tensorrt_llm._torch.staircase.catalog.attention.mla_rope_append_paged_kv_assign_q import ( # noqa: E501 + mla_rope_append_paged_kv_assign_q, +) +from tensorrt_llm._torch.staircase.catalog.attention.mla_rope_generation import mla_rope_generation +from tensorrt_llm._torch.staircase.catalog.attention.thop_attention import thop_attention +from tensorrt_llm._torch.staircase.catalog.comm.allgather import allgather +from tensorrt_llm._torch.staircase.catalog.comm.reducescatter import reducescatter +from tensorrt_llm._torch.staircase.catalog.gemm.bmm_out import bmm_out +from tensorrt_llm._torch.staircase.catalog.gemm.cublas_mm import cublas_mm +from tensorrt_llm._torch.staircase.catalog.gemm.nvfp4_gemm import nvfp4_gemm +from tensorrt_llm._torch.staircase.catalog.moe.fp4_block_scale_moe_runner import ( # noqa: E501 + fp4_block_scale_moe_runner, +) +from tensorrt_llm._torch.staircase.catalog.moe.fused_moe import fused_moe +from tensorrt_llm._torch.staircase.catalog.moe.noaux_tc_op import noaux_tc_op +from tensorrt_llm._torch.staircase.catalog.norm.flashinfer_fused_add_rmsnorm import ( # noqa: E501 + flashinfer_fused_add_rmsnorm, +) +from tensorrt_llm._torch.staircase.catalog.norm.flashinfer_rmsnorm import flashinfer_rmsnorm +from tensorrt_llm._torch.staircase.catalog.quantization.fp4_quantize import fp4_quantize +from tensorrt_llm._torch.staircase.catalog.torch.add import add +from tensorrt_llm._torch.staircase.catalog.torch.concat import concat +from tensorrt_llm._torch.staircase.catalog.torch.copy_ import copy_ +from tensorrt_llm._torch.staircase.catalog.torch.embedding import embedding +from tensorrt_llm._torch.staircase.catalog.torch.empty import empty +from tensorrt_llm._torch.staircase.catalog.torch.expand import expand +from tensorrt_llm._torch.staircase.catalog.torch.pad import pad +from tensorrt_llm._torch.staircase.catalog.torch.reshape import reshape +from tensorrt_llm._torch.staircase.catalog.torch.split import split +from tensorrt_llm._torch.staircase.catalog.torch.transpose import transpose +from tensorrt_llm._torch.staircase.catalog.torch.view_dtype import view_dtype + +from . import weights as _weights + +# The GPU architecture this target IS. Routing will not send another one here, +# but a direct instantiation could, and the certification is per arch: this +# assert is what the version pin used to be. In-tree the version moves with +# the code, so pinning it is meaningless; the architecture does not. +_SM = (10, 3) + + +def _check_static_contract() -> None: + """Import-time fail-fast: op symbol existence. The list is the forward's + trtllm call set plus `block_scale_interleave`, which the load-time + expert/scale relayout in weights.py depends on.""" + for op in ( + "cublas_mm", + "bmm_out", + "nvfp4_gemm", + "flashinfer_rmsnorm", + "flashinfer_fused_add_rmsnorm", + "flashinfer_silu_and_mul", + "fp4_quantize", + "noaux_tc_op", + "fp4_block_scale_moe_runner", + "fused_moe", + "mla_rope_generation", + "mla_rope_append_paged_kv_assign_q", + "load_paged_kv_cache_for_mla", + "allgather", + "reducescatter", + "block_scale_interleave", + ): + assert hasattr(torch.ops.trtllm, op), f"missing op trtllm::{op}" + from tensorrt_llm.bindings.internal import thop + + assert hasattr(thop, "attention"), "missing pybind thop.attention" + + +_check_static_contract() + +# Metadata fields consumed each step (sourcing mirrors the in-tree +# FallbackFmha for this trtllm version; existence checked at first forward). +# The tail group is read only by the first-forward contract check: this +# target holds those features at the MLA columns' inert values, and the +# asserts make that honest instead of silently dropping an enabled feature. +_STEP_FIELDS = ( + "all_rank_num_tokens", + "kv_lens_cuda_runtime", + "kv_lens_runtime", + "host_total_kv_lens", + "prompt_lens_cuda_runtime", + "prompt_lens_cpu_runtime", + "host_request_types_runtime", + "kv_cache_block_offsets", + "host_kv_cache_pool_pointers", + "host_kv_cache_pool_mapping", + "effective_workspace", + "tokens_per_block", + "max_num_requests", + "max_context_length", + "max_seq_len", + "num_contexts", + "num_ctx_tokens", + "num_seqs", + "trtllm_gen_jit_warmup", + "effective_beam_width", + "cache_indirection", + "block_ids_per_seq", + "is_cross", + "is_spec_decoding_enabled", + "use_spec_decoding", + "flash_mla_tile_scheduler_metadata", + "flash_mla_num_splits", + # Added between 1.3.0rc21 and 1.3.0rc26. Both are engine-prepared + # per-instance constants (max_num_sequences defaults to max_num_requests; + # the tree-mask flag is set from is_spec_dec_dynamic_tree, and this + # target's MTP is a linear tree), so they project like the rest. + "max_num_sequences", + "force_prepare_spec_dec_tree_mask", +) + +# The cached-prefix context group. The engine only creates these attributes +# when it prepares the metadata for MLA context over reused blocks — under +# trtllm's default kv_cache_config that is on, and with block reuse disabled +# they are absent entirely — so their existence selects the context flavor +# rather than being a hard requirement. +_CACHED_CTX_FIELDS = ( + "enable_context_mla_with_cached_kv", + "ctx_cached_token_indptr", + "ctx_kv_indptr", + "max_ctx_seq_len", + "max_ctx_kv_len", +) + + +def _build_step_args(md: TrtllmAttentionMetadata) -> dict: + """Project the prepared metadata onto the batch state both MLA attention + calls of one forward share; built once per forward. CUDA-graph classes: + every tensor here is an engine-owned persistent buffer refreshed in place + (reference class), and every Python int is an engine-construction or + per-capture constant (host-derived class — a decode-only graph always + sees num_contexts == 0). Phase-specific arguments (q/k/v, output, head + geometry, the scheduler buffers, the fp8 decode buffers) are passed at the + call sites.""" + return dict( + sequence_length=md.kv_lens_cuda_runtime, + host_past_key_value_lengths=md.kv_lens_runtime, + host_total_kv_lens=md.host_total_kv_lens, + context_lengths=md.prompt_lens_cuda_runtime, + host_context_lengths=md.prompt_lens_cpu_runtime, + host_request_types=md.host_request_types_runtime, + kv_cache_block_offsets=md.kv_cache_block_offsets, + host_kv_cache_pool_pointers=md.host_kv_cache_pool_pointers, + host_kv_cache_pool_mapping=md.host_kv_cache_pool_mapping, + workspace_=md.effective_workspace, + tokens_per_block=md.tokens_per_block, + max_num_requests=md.max_num_requests, + max_context_length=md.max_context_length, + max_seq_len=md.max_seq_len, + attention_window_size=md.max_seq_len, + rope_max_positions=md.max_seq_len, + rope_original_max_positions=md.max_seq_len, + num_contexts=md.num_contexts, + num_ctx_tokens=md.num_ctx_tokens, + trtllm_gen_jit_warmup=md.trtllm_gen_jit_warmup, + max_num_sequences=md.max_num_sequences, + force_prepare_spec_dec_tree_mask=md.force_prepare_spec_dec_tree_mask, + ) + + +# Feature groups this target holds inert, shared by both MLA calls: the +# contract's MLA columns plus every not-certified group at its listed inert +# value. Causal mask, beam width 1, no sinks, no spec-dec mask / sparse / +# cross / relative-bias / mRoPE / helix / FlashMLA. `quant_mode` and +# `q_scaling` are *not* here — both are derived from this checkpoint's config +# in __init__ and joined on in `self._call`. Neither is +# `predicted_tokens_per_seq`: it is 1 on the context call and the generation +# call's own query-tokens-per-sequence, which is 1 for an ordinary decode step +# and `runtime_draft_len + 1` under MTP, so it is passed per call site. +# +# Both kv scale tensors stay None over the fp8 latent pool: the generation +# call reads neither, and both context flavors are correct only at s = 1.0, +# which None is read as exactly. This checkpoint's k_scale/v_scale are 1.0, +# asserted after load. +# +# position_embedding_type=8 selects the in-kernel GPT-J rope of the MLA path; +# rope_dim / rope_base / the two tables are per-instance. The seven scalars +# below (rope_scale_type, rope_scale, the two m-scales, and the two position +# windows in the step args) are measured inert on the MLA path — the table's +# content is the only rope input the op reads, and the whole YaRN blend lives +# there. They are held at neutral values rather than at the config's `yarn` +# names, whose trtllm enum coding is not derivable from this target's sources. +# The MLA KV pool stores no residual tail: the op accepts 0 or rope_size +# and rejects non-zero unless the pool is FP4, and this checkpoint's is +# fp8-e4m3. The in-tree caller passes a literal 0 on the same path. +_KV_RESIDUAL_DIM = 0 + +_CALL_INERT = dict( + output_sf=None, + out_scale=None, + kv_scale_orig_quant=None, + kv_scale_quant_orig=None, + attention_sinks=None, + update_kv_cache=True, + beam_width=1, + mask_type=1, + use_paged_context_fmha=False, + is_mla_enable=True, + rope_append=True, + position_embedding_type=8, + rope_scale_type=0, + rope_scale=1.0, + rope_short_m_scale=1.0, + rope_long_m_scale=1.0, + chunked_prefill_buffer_batch_size=1, + attention_chunk_size=None, + softmax_stats_tensor=None, + cache_indirection=None, + block_ids_per_seq=None, + max_context_q_len_override=None, + is_cross=False, + cross_kv=None, + relative_attention_bias=None, + relative_attention_max_distance=0, + mrope_rotary_cos_sin=None, + mrope_position_deltas=None, + helix_position_offsets=None, + helix_is_inactive_rank=None, + is_spec_decoding_enabled=False, + use_spec_decoding=False, + is_spec_dec_tree=False, + spec_decoding_generation_lengths=None, + spec_decoding_position_offsets_for_cpp=None, + spec_decoding_packed_mask=None, + spec_decoding_bl_tree_mask_offset=None, + spec_decoding_bl_tree_mask=None, + spec_decoding_target_max_draft_tokens=None, + spec_bl_tree_first_sparse_mask_offset_kv=None, + sparse_kv_indices=None, + sparse_kv_offsets=None, + sparse_attn_indices=None, + sparse_attn_offsets=None, + sparse_attn_indices_block_size=0, + num_sparse_topk=None, + flash_mla_tile_scheduler_metadata=None, + flash_mla_num_splits=None, + # Added between 1.3.0rc21 and 1.3.0rc26; held at the values the op had + # before they existed, which are also what the in-tree caller passes on + # this path. kv_norm_weight is not merely a default: non-None would fold + # the kv_a_layernorm into the KV kernel and make it read latent_cache + # RAW, and this target normalizes that itself -- passing the weight would + # normalize twice. skip_correction is a lossy trtllm-gen MLA option + # (SM100/SM103, off by default upstream); enabling it is a configs/ + # variant's business, not the identity assembly's. + kv_norm_weight=None, + kv_norm_eps=1e-6, + skip_correction_threshold=0.0, +) + +# QuantMode's fp8-KV-cache bit — the value every MLA entry certifies for an +# fp8-e4m3 latent pool. Bits outside the KV-cache group ride along unread +# (`1152` = | FP8_1x128_128x128 and `384` = | FP8_QDQ were measured +# bit-identical to a bare 128 on every MLA flavor and on the gather), so the +# derivation below maps the checkpoint's declared KV algo onto this bit rather +# than reconstructing the engine's full bitmask. +_QUANT_MODE_FP8_KV = 128 + +# NVFP4 block size: one e4m3 scale per 16 contiguous elements along K. +_SF_VEC = 16 +# The MoE runner reads the activation scales as a *linear* (row-major) +# buffer; the dense GEMM reads the 128x4-swizzled one. The two have the same +# byte count whenever num_tokens is a multiple of 128 — every decode +# CUDA-graph batch — and the wrong one is then accepted silently, so both +# layouts are spelled out rather than left to the quantizer's default. +_SF_LINEAR = False +_SF_SWIZZLED = True +# Routing already happened outside (noaux_tc_op), so the runner's routing +# method and group configuration are inert on this pre-routed path — including +# this checkpoint's (n_group 8, topk_group 4), measured bitwise identical to +# None/None at routing_method_type 1. +_ROUTING_METHOD_INERT = 1 +# Largest token count the expert call is allowed to see in one invocation. +# `fp4_block_scale_moe_runner` is certified over its whole token column — up +# to 8192 — at this checkpoint's routed geometry (H 7168, I 2048, 256 experts +# top-8) for the full stack and for each of the four 64-wide expert-parallel +# windows; `noaux_tc_op` is certified to num_tokens 8192. Under dep4 the +# gathered token set reaches `4 * max_num_tokens` = 32768, so the call is +# chunked rather than run outside those columns. Routing and the activation +# quantization ride along per chunk — both are per-token, so chunking is a +# no-op for them mathematically. +_MOE_MAX_T = 8192 +_ACT_TYPE_SWIGLU = 0 +# `fused_moe`'s activation enum is a *different* enum from the trtllm-gen +# runner's above: 5 is Swiglu there, 0 is not a gated type at all. +_FUSED_MOE_ACT_SWIGLU = 5 +# The MTP layer's `eh_proj` operand order, at one place so flipping it is a +# one-line experiment. True = `concat(enorm(e), hnorm(h))`, the embedding block +# first — measured on this checkpoint's own weights by two independent +# statistics (docs/models/multi-token-prediction.md), which contradicts the +# DeepSeek-V3 report's `M_k[RMSNorm(h); RMSNorm(e)]` notation and agrees with +# the parameter's name. Nothing in the checkpoint's metadata pins it and +# nothing downstream detects a flip: the drafts are simply rejected, so the +# acceptance rate is what confirms it end to end. +_MTP_EMBED_BLOCK_FIRST = True +# thop_attention's certified q_scaling values over an fp8 latent pool are 1.0 +# and DeepSeek-R1's YaRN temperature; the config-derived value is checked +# against the latter so a config change lands outside the certified column +# loudly rather than silently. +_CERTIFIED_YARN_Q_SCALING = 0.5336594 + + +def _moe_chunk_sizes(total: int) -> list[int]: + """Split `total` gathered rows into expert-call chunks of at most + `_MOE_MAX_T` rows. Host arithmetic only — no tensor is created.""" + full, rest = divmod(total, _MOE_MAX_T) + return [_MOE_MAX_T] * full + ([rest] if rest else []) + + +def _mtp_dp_rows(all_rank_num_tokens, rank: int, dp_size: int, rows: int) -> int: + """The uniform row count the MTP layer pads its token block to before its + MoE round trip: `max` over the group's per-rank counts. + + **The list is a parameter and there is no metadata argument, on purpose.** + Inside the draft loop `attn_metadata.all_rank_num_tokens` is the *trunk's* + list from draft step 1 onwards — the worker leaves it in place for the + whole loop and passes the correct basis in as the `all_rank_num_tokens` + keyword instead (`spec_metadata.all_rank_num_tokens` at step 0, then + `spec_metadata.subseq_all_rank_num_tokens`, which is the per-rank + *sequence* count). Both collective contracts certify that calls pair by + position and that at equal byte counts a divergence is silent — every rank + wrong in 98-99% of elements, bitwise reproducibly, with no hang — so a + helper that *could* reach the metadata is the whole hazard. This one + cannot. `_dp_rows` on the trunk keeps its own shape; the two are not + interchangeable.""" + counts = [int(n) for n in all_rank_num_tokens] + assert len(counts) == dp_size, (counts, dp_size) + assert counts[rank] == rows, ( + f"rank {rank} holds {rows} MTP rows but the worker's " + f"all_rank_num_tokens says {counts[rank]} ({counts}); the collectives " + "would gather the wrong split" + ) + return max(counts) + + +def _yarn_mscale(factor: float, mscale: float) -> float: + """DeepSeek's YaRN magnitude scaling `m(x)`: 1.0 at or below factor 1, + `0.1 * x * ln(factor) + 1` above it.""" + if factor <= 1.0: + return 1.0 + return 0.1 * mscale * math.log(factor) + 1.0 + + +class StaircaseCore(DecoderModel): + def __init__(self, model_config: ModelConfig): + super().__init__(model_config) + cfg = model_config.pretrained_config + assert cfg is not None + + # This target IS this topology and this geometry — assert, never + # adapt. Read the topology off the mapping the engine built, never + # off the path segment: the segment names the intent. + assert torch.cuda.get_device_capability() == _SM, ( + f"target certified on sm_{_SM[0]}{_SM[1]}, running on " + f"sm_{''.join(map(str, torch.cuda.get_device_capability()))}" + ) + mapping = model_config.mapping + assert mapping.world_size == 4, f"dep4 target, world size {mapping.world_size}" + assert mapping.tp_size == 4, f"dep4 target, tp_size {mapping.tp_size}" + assert mapping.pp_size == 1, "pipeline parallelism is not implemented" + assert mapping.moe_ep_size == 4, f"dep4 target, ep {mapping.moe_ep_size}" + assert mapping.moe_tp_size == 1, ( + "the routed experts are split by expert parallelism only; a " + f"moe_tp_size of {mapping.moe_tp_size} splits them a second way" + ) + assert mapping.enable_attention_dp, ( + "attention data parallelism is what the dep4 segment declares; " + "without it attention would be head-split (that is tep4)" + ) + self.rank = mapping.rank + self.ep_rank = mapping.moe_ep_rank + self.ep_size = mapping.moe_ep_size + # The collective group: at pp_size 1 with one expert-parallel split, + # every rank participates in both halves of the MoE round trip, so + # the group is the whole world. + self.dp_group = list(range(mapping.world_size)) + self.dp_size = mapping.world_size + + dt = model_config.torch_dtype + assert dt == torch.bfloat16, f"bf16 target, engine resolved {dt}" + assert cfg.torch_dtype == torch.bfloat16, cfg.torch_dtype + assert not cfg.tie_word_embeddings, "untied lm_head" + assert not cfg.attention_bias, "no q/kv/o bias anywhere" + assert cfg.hidden_act == "silu", "SwiGLU over a silu gate" + # The NVFP4 recipe covers the MLP only; the latent KV pool is fp8-e4m3 + # per the checkpoint's own quant config, which is what selects + # quant_mode on every MLA call. Derived, not hard-coded: a checkpoint + # declaring no KV quantization is a bf16-pool target and a different + # assembly. + kv_algo = model_config.quant_config.kv_cache_quant_algo + assert kv_algo is not None and str(kv_algo).upper().endswith("FP8"), ( + f"this target is the fp8-e4m3 latent-pool assembly; the checkpoint " + f"declares kv_cache_quant_algo {kv_algo!r}" + ) + self.quant_mode = _QUANT_MODE_FP8_KV + + self.num_layers = cfg.num_hidden_layers + self.hidden = cfg.hidden_size + self.eps = cfg.rms_norm_eps + self.vocab = cfg.vocab_size + # The checkpoint ships `num_nextn_predict_layers` extra decoder layers + # past `num_hidden_layers` for multi-token prediction — on this one a + # single bf16 layer 61 with its own 256 experts, embedding, eh_proj, + # two extra norms and an output head. + self.mtp_layers = int(getattr(cfg, "num_nextn_predict_layers", 0) or 0) + # Whether this engine drafts. The checkpoint decides the *mode* + # (`num_nextn_predict_layers: 1` -> MTP-Eagle one-model, one layer + # replayed `max_draft_len` times); a `configs/` variant's + # `speculative_config` decides whether it runs at all, and the engine + # has already resolved that onto `model_config.spec_config` by the time + # the model is built. With it absent — the target's identity config — + # layer 61 is not part of this model at all: its keys stay a predicted + # non-load in the weight manifest and nothing below is declared. + spec_config = getattr(model_config, "spec_config", None) + self.mtp_enabled = spec_config is not None + if self.mtp_enabled: + assert self.mtp_layers == 1, ( + "this target implements the one-layer MTP-Eagle module the " + f"checkpoint ships; num_nextn_predict_layers is {self.mtp_layers}" + ) + + # MLA geometry. num_key_value_heads is 128 in this config but MLA has + # no separate KV heads: the context call runs Hq == Hkv == heads over + # head_size nope+rope, the generation call runs Hq == heads against a + # single latent KV head of width kv_lora+rope. Attention is + # replicated under DP, so every rank runs the whole head set. + self.heads = cfg.num_attention_heads + self.nope = cfg.qk_nope_head_dim + self.rope = cfg.qk_rope_head_dim + self.v_dim = cfg.v_head_dim + self.kv_lora = cfg.kv_lora_rank + self.q_lora = cfg.q_lora_rank + self.qk_dim = self.nope + self.rope + self.lat_dim = self.kv_lora + self.rope + # The query is a LoRA pair here: q_a_proj -> q_a_layernorm -> q_b_proj. + # A checkpoint with q_lora_rank null projects directly and needs the + # single-q_proj path instead (that is the deepseek-v3-lite sibling). + assert isinstance(self.q_lora, int) and self.q_lora > 0, ( + "this target implements the q-LoRA query path; a checkpoint with " + f"q_lora_rank {self.q_lora!r} needs the direct q_proj path" + ) + assert cfg.num_key_value_heads == cfg.num_attention_heads, "MLA: Hkv == Hq" + # thop_attention aborts the process on a head_size its FMHA kernels + # do not carry, so both call shapes are pinned here. + assert self.qk_dim == 192 and self.lat_dim == 576 and self.v_dim == 128, ( + f"certified MLA head dims are 192 (context) / 576 (generation) / " + f"128 (v), got {self.qk_dim} / {self.lat_dim} / {self.v_dim}" + ) + # The MLA generation phase compiles a decode kernel per head count, so + # the count is a certification axis rather than a free shape — and over + # an fp8-e4m3 latent pool 128 is the *only* certified count. + assert self.heads == 128, ( + f"thop_attention certifies MLA over an fp8 latent pool at 128 query " + f"heads only; this topology yields {self.heads}" + ) + + # RoPE: GPT-J interleaved pairs over the rope slice, YaRN-scaled. The + # engine hands either the flat pre-migration fields or the + # transformers-5.x rope_parameters dict; both shapes are resolved and + # every scalar the table depends on is asserted rather than defaulted. + rope_cfg = getattr(cfg, "rope_scaling", None) or getattr(cfg, "rope_parameters", None) + assert isinstance(rope_cfg, dict), f"YaRN rope config expected, got {rope_cfg!r}" + rope_kind = rope_cfg.get("type", rope_cfg.get("rope_type")) + assert rope_kind == "yarn", f"this target builds a YaRN table, got {rope_kind!r}" + theta = rope_cfg.get("rope_theta", getattr(cfg, "rope_theta", None)) + assert theta is not None and float(theta) > 0.0, "rope theta" + self.theta = float(theta) + self.rope_factor = float(rope_cfg["factor"]) + self.rope_orig_max = int(rope_cfg["original_max_position_embeddings"]) + self.beta_fast = float(rope_cfg["beta_fast"]) + self.beta_slow = float(rope_cfg["beta_slow"]) + mscale = float(rope_cfg["mscale"]) + mscale_all_dim = float(rope_cfg["mscale_all_dim"]) + assert self.rope_factor > 1.0 and self.rope_orig_max > 0, rope_cfg + assert getattr(cfg, "rope_interleave", True), ( + "the MLA ops apply GPT-J (interleaved-pair) rope in kernel" + ) + self.max_pos = cfg.max_position_embeddings + # The table's amplitude and the softmax temperature are the two halves + # of YaRN's magnitude correction, and they go to different places: the + # amplitude multiplies cos/sin (exactly 1.0 whenever the config's two + # m-scales agree), while the temperature — built from mscale_all_dim, + # as the reference model does — is folded into the op's q_scaling as + # 1/m^2. Putting either in the other's place is silently wrong. + self.rope_amplitude = _yarn_mscale(self.rope_factor, mscale) / _yarn_mscale( + self.rope_factor, mscale_all_dim + ) + temperature = _yarn_mscale(self.rope_factor, mscale_all_dim) + self.q_scaling = 1.0 / (temperature * temperature) + assert abs(self.q_scaling - _CERTIFIED_YARN_Q_SCALING) < 1e-6, ( + f"q_scaling {self.q_scaling} is outside thop_attention's certified " + f"fp8-pool column (1.0 and {_CERTIFIED_YARN_Q_SCALING})" + ) + # Per-call constants: the inert groups above plus the two values this + # checkpoint derives. + self._call = dict(_CALL_INERT, quant_mode=self.quant_mode, q_scaling=self.q_scaling) + + # MLP structure: the first `first_k_dense_replace` layers are dense, + # the rest are routed MoE plus a shared-expert pair fused into one + # dense linear pair of width n_shared * moe_intermediate_size. Both + # dense shapes are replicated under attention DP and run over this + # rank's own tokens. + self.dense_layers = cfg.first_k_dense_replace + assert 0 < self.dense_layers < self.num_layers, "mixed dense/MoE stack" + assert cfg.moe_layer_freq == 1, "every layer past the dense prefix is MoE" + self.num_experts = cfg.n_routed_experts + self.topk = cfg.num_experts_per_tok + self.moe_inter = cfg.moe_intermediate_size + self.shared_inter = cfg.moe_intermediate_size * cfg.n_shared_experts + self.dense_inter = cfg.intermediate_size + assert 0 < self.topk < self.num_experts, "MoE top-k bound" + # noaux_tc_op is the whole gate: in-kernel sigmoid, bias correction + # for selection only, group-limited selection, renormalization and the + # routed_scaling_factor multiply. A checkpoint with norm_topk_prob + # false cannot use it. + assert cfg.topk_method == "noaux_tc", cfg.topk_method + assert cfg.scoring_func == "sigmoid", cfg.scoring_func + assert cfg.norm_topk_prob, "noaux_tc_op always renormalizes" + self.n_group = cfg.n_group + self.topk_group = cfg.topk_group + # Group-limited routing, which this checkpoint uses and the lite + # sibling does not. noaux_tc_op's grouped configuration carries four + # hard limits of its own; each is checked here because the op reports + # a violation as one opaque "unsupported configuration". + assert self.n_group > 1 and 1 <= self.topk_group <= self.n_group, ( + f"grouped routing needs 1 <= topk_group <= n_group, got " + f"{self.topk_group} / {self.n_group}" + ) + assert self.num_experts % self.n_group == 0, "experts per routing group" + assert self.num_experts <= 256 and self.num_experts // self.n_group <= 32, ( + "noaux_tc_op's grouped path takes at most 256 experts and 32 per " + f"group; this config has {self.num_experts} in {self.n_group} groups" + ) + assert self.topk <= 8, "noaux_tc_op's grouped path takes at most top-8" + self.routed_scale = float(cfg.routed_scaling_factor) + # Expert parallelism: the routing space stays global and every rank + # runs the full top-k over the *gathered* token set, but a rank holds + # only its own window of experts and the kernel drops every slot + # outside it. The four windows' outputs sum to the whole layer — that + # is what the reduce-scatter completes. + assert self.num_experts % mapping.moe_ep_size == 0, "experts per rank" + self.local_experts = self.num_experts // mapping.moe_ep_size + self.expert_offset = self.local_experts * self.ep_rank + assert self.local_experts == 64, ( + f"fp4_block_scale_moe_runner certifies the four-way split of 256 " + f"experts (windows of 64); this topology yields {self.local_experts}" + ) + # fp4_block_scale_moe_runner's block-scale layout rules. + assert self.hidden % 256 == 0, "MoE hidden must be a multiple of 256" + assert self.moe_inter % 64 == 0, "MoE intermediate must be a multiple of 64" + # nvfp4_gemm needs K and N multiples of 32 on both dense linears + # (gate_up: K=hidden, N=2*inter; down: K=inter, N=hidden), and + # flashinfer_silu_and_mul vectorizes the up half from element + # `inter`, so that offset must be 16-element aligned too. The 128x4 + # scale swizzle adds two more: the gate_up scale matrix has 2*inter + # rows (a multiple of 128) and the down one inter/16 columns (a + # multiple of 4) — together, inter must be a multiple of 64. + assert self.hidden % 32 == 0, "nvfp4_gemm operand width" + for inter in (self.dense_inter, self.shared_inter): + assert inter % 64 == 0, ( + "MLP intermediate must be a multiple of 64: nvfp4_gemm " + "operand width, the silu_and_mul half offset, and the 128x4 " + "scale swizzle's row/column alignment" + ) + + # Weight declaration. HF [out, in] row-major so checkpoint rows copy + # in unchanged, except kv_b_proj (row-regrouped at load, see + # weights.py) and the expert stacks (declared in the MoE runner's + # kernel-ready shuffled/swizzled layout, at the rank's window). + # Meta-init intercepts torch.empty here — real CUDA storage arrives + # when the engine materializes the registry. + def P(*shape, dtype=dt): + return nn.Parameter(torch.empty(*shape, dtype=dtype), requires_grad=False) + + u8, f32 = torch.uint8, torch.float32 + w = nn.ParameterDict() + for i in range(self.num_layers): + w[f"l{i}_norm1"] = P(self.hidden) + w[f"l{i}_qa"] = P(self.q_lora, self.hidden) + w[f"l{i}_q_norm"] = P(self.q_lora) + w[f"l{i}_qb"] = P(self.heads * self.qk_dim, self.q_lora) + w[f"l{i}_kva"] = P(self.lat_dim, self.hidden) + w[f"l{i}_kv_norm"] = P(self.kv_lora) + w[f"l{i}_kvb"] = P(self.heads * (self.nope + self.v_dim), self.kv_lora) + w[f"l{i}_o"] = P(self.hidden, self.heads * self.v_dim) + # The checkpoint's calibrated fp8 KV-cache scales. They are loaded + # rather than skipped so `derive_after_load` can assert the value + # the whole fp8 MLA path depends on. + w[f"l{i}_k_scale"] = P(1, dtype=f32) + w[f"l{i}_v_scale"] = P(1, dtype=f32) + w[f"l{i}_norm2"] = P(self.hidden) + inter = self.dense_inter if i < self.dense_layers else self.shared_inter + w[f"l{i}_mlp_gu_w"] = P(2 * inter, self.hidden // 2, dtype=u8) + w[f"l{i}_mlp_gu_s"] = P(2 * inter * (self.hidden // _SF_VEC), dtype=u8) + w[f"l{i}_mlp_dn_w"] = P(self.hidden, inter // 2, dtype=u8) + w[f"l{i}_mlp_dn_s"] = P(self.hidden * (inter // _SF_VEC), dtype=u8) + for name in ( + "isc1", + "isc1_up", + "ws2_1", + "ws2_1_up", + "isc2", + "ws2_2", + ): + w[f"l{i}_mlp_{name}"] = P(1, dtype=f32) + if i < self.dense_layers: + continue + e, mi = self.local_experts, self.moe_inter + w[f"l{i}_router"] = P(self.num_experts, self.hidden) + # fp32 on this checkpoint (the reference model keeps the + # correction bias in fp32 whatever the rest of the weights are); + # noaux_tc_op takes bf16 logits against an fp32 bias and returns + # weights in the *logits* dtype, which is what the MoE runner + # demands. + w[f"l{i}_router_bias"] = P(self.num_experts, dtype=f32) + w[f"l{i}_fc1_w"] = P(e, 2 * mi, self.hidden // 2, dtype=u8) + w[f"l{i}_fc1_s"] = P(e, 2 * mi, self.hidden // _SF_VEC, dtype=u8) + w[f"l{i}_fc2_w"] = P(e, self.hidden, mi // 2, dtype=u8) + w[f"l{i}_fc2_s"] = P(e, self.hidden, mi // _SF_VEC, dtype=u8) + # The per-expert NVFP4 scalars stay whole on every rank: they + # cost 6 floats per expert, and the shared-expert activation + # scale is asserted against the max over *all* routed experts, + # which a window could not see. The window is sliced out in + # derive_after_load, where the kernel's [local_num_experts] + # operands are built. + for name in ( + "isc1", + "isc1_up", + "ws2_1", + "ws2_1_up", + "isc2", + "ws2_2", + ): + w[f"l{i}_e_{name}"] = P(self.num_experts, dtype=f32) + w["final_norm"] = P(self.hidden) + w["embed"] = P(self.vocab, self.hidden) + # The MTP module at layer index `num_hidden_layers`, declared only when + # a configs/ variant turned drafting on. Its attention block is + # byte-identical in geometry to a trunk layer's; its MLP path is the + # same structure at a different **dtype** — `hf_quant_config.json` + # carries `model.layers.61*` as one wildcard entry in its + # `exclude_modules` list, so every weight here is bf16 while the + # trunk's MLP is NVFP4. That is the export's choice about where + # accuracy is worth the bytes, so the stacks are declared bf16 and fed + # to the unquantized `fused_moe` rather than re-quantized at load to + # reuse the trunk's expert vocabulary. `embed_tokens` and + # `shared_head.head` are *not* declared: both are bitwise copies of the + # trunk's `model.embed_tokens` / `lm_head`, and the draft-model + # container points at those instead — 1.85 GB per rank saved. + if self.mtp_enabled: + e, mi = self.local_experts, self.moe_inter + w["mtp_enorm"] = P(self.hidden) + w["mtp_hnorm"] = P(self.hidden) + w["mtp_eh"] = P(self.hidden, 2 * self.hidden) + w["mtp_norm1"] = P(self.hidden) + w["mtp_qa"] = P(self.q_lora, self.hidden) + w["mtp_q_norm"] = P(self.q_lora) + w["mtp_qb"] = P(self.heads * self.qk_dim, self.q_lora) + w["mtp_kva"] = P(self.lat_dim, self.hidden) + w["mtp_kv_norm"] = P(self.kv_lora) + w["mtp_kvb"] = P(self.heads * (self.nope + self.v_dim), self.kv_lora) + w["mtp_o"] = P(self.hidden, self.heads * self.v_dim) + w["mtp_k_scale"] = P(1, dtype=f32) + w["mtp_v_scale"] = P(1, dtype=f32) + w["mtp_norm2"] = P(self.hidden) + w["mtp_router"] = P(self.num_experts, self.hidden) + w["mtp_router_bias"] = P(self.num_experts, dtype=f32) + # `fused_moe`'s stacked layout: `[E, 2I, H]` with the **up** rows + # first and the gate rows last (the opposite half order from the + # dense gate_up linear below, which flashinfer_silu_and_mul reads + # gate-first), and `[E, H, I]` for FC2. No interleave, no 32-row + # block shuffle, no swizzle — those belong to the trtllm-gen + # block-scale runner the trunk uses, not to this one. + w["mtp_fc1"] = P(e, 2 * mi, self.hidden) + w["mtp_fc2"] = P(e, self.hidden, mi) + w["mtp_sh_gu"] = P(2 * self.shared_inter, self.hidden) + w["mtp_sh_dn"] = P(self.hidden, self.shared_inter) + w["mtp_head_norm"] = P(self.hidden) + self.w = w + + self._attn: list | None = None + self._mlp: list | None = None + self._moe: list | None = None + self._mtp: dict | None = None + self._next_norm: list | None = None + self._rope: dict | None = None + self._rope_positions = 0 + self._side_stream: torch.cuda.Stream | None = None + self._cached_ctx = False + self._step_contract_checked = False + + def _rope_tables(self, device, positions: int) -> dict: + """The duplicated-layout GPT-J rope table the MLA ops read: per + position, `rope` (cos, sin) pairs whose second half duplicates the + first, flattened to `[1, positions * rope * 2]` fp32, plus the + `[rope/2]` inverse-frequency vector from the same construction. + + `positions` is a row count, not a model property: every row depends + only on its own index, so a longer table is the same table with more + rows and rebuilding one at a larger size changes no existing row. + + The inverse frequencies carry this checkpoint's **YaRN** blend: the + interpolated (factor-divided) frequency for the low-frequency half of + the spectrum, the original one for the high-frequency half, ramped + between the two correction dimensions. Everything else about the rope + configuration is inert on the MLA path — the table's content is the + only rope input the ops read — so this construction is the whole of + it, and a table built from the unscaled theta is a silently wrong + model that diverges with position. Built in fp64 on the host, rounded + once.""" + half = self.rope // 2 + d = torch.arange(half, dtype=torch.float64) + freq = self.theta ** (2.0 * d / self.rope) + two_pi = 2.0 * math.pi + log_theta = math.log(self.theta) + low = max( + 0.0, + math.floor( + self.rope + * math.log(self.rope_orig_max / (self.beta_fast * two_pi)) + / (2.0 * log_theta) + ), + ) + high = min( + self.rope - 1.0, + math.ceil( + self.rope + * math.log(self.rope_orig_max / (self.beta_slow * two_pi)) + / (2.0 * log_theta) + ), + ) + ramp = ((d - low) / max(high - low, 0.001)).clamp(0.0, 1.0) + inv = ramp / (self.rope_factor * freq) + (1.0 - ramp) / freq + ang = torch.arange(positions, dtype=torch.float64)[:, None] * inv[None, :] + cos, sin = ang.cos() * self.rope_amplitude, ang.sin() * self.rope_amplitude + table = torch.empty(positions, self.rope, 2, dtype=torch.float64) + table[:, :half, 0] = cos + table[:, half:, 0] = cos + table[:, :half, 1] = sin + table[:, half:, 1] = sin + return { + "rotary_cos_sin": table.reshape(1, positions * self.rope * 2) + .float() + .to(device) + .contiguous(), + "rotary_inv_freq": inv.float().to(device).contiguous(), + "rope_dim": self.rope, + "rope_base": self.theta, + } + + def derive_after_load(self) -> None: + """Post-load derivation: column-major GEMM views (`.t()` is + zero-copy), the two MLA absorption operands split out of the + row-regrouped kv_b_proj, the rope table, and every NVFP4 call scalar + folded from the checkpoint's per-tensor `input_scale` / + `weight_scale_2` pairs — the routed ones sliced to this rank's expert + window, which is where the kernel's `[local_num_experts]` operands + come from. Meta is over here, so real tensors may be created.""" + w = self.w + device = w["final_norm"].device + self._rope_positions = self.max_pos + self._rope = self._rope_tables(device, self._rope_positions) + window = slice(self.expert_offset, self.expert_offset + self.local_experts) + + attn, mlp, moe, nxt = [], [], [], [] + hn = self.heads * self.nope + for i in range(self.num_layers): + # The fp8 latent pool is written at 1/s and read at s, and the two + # roles live in different ops with no relation checked anywhere in + # the chain. This target passes None for both, which every op reads + # as exactly 1.0 — correct only for a checkpoint calibrated at 1.0, + # and additionally the only value the fp8 MLA *context* path is + # self-consistent at (it quantizes q/k/v at 1.0 while applying + # s^2/s regardless). So the checkpoint's own scales are checked + # rather than assumed. + for role in ("k_scale", "v_scale"): + s = w[f"l{i}_{role}"] + assert torch.equal(s, torch.ones_like(s)), ( + f"layer {i}: {role} is {s.item()}, not 1.0; the fp8 MLA " + "context path is only correct at a KV scaling factor of " + "1.0, and this assembly passes no scale tensors" + ) + kvb = w[f"l{i}_kvb"] + attn.append( + ( + w[f"l{i}_qa"].t(), + w[f"l{i}_q_norm"], + w[f"l{i}_qb"].t(), + w[f"l{i}_kva"].t(), + w[f"l{i}_kv_norm"], + # k_b [H, nope, C] absorbs into q_nope; v_b_t [H, C, v] + # expands the latent attention output. Both are views of + # the row-regrouped kv_b_proj. + kvb[:hn].reshape(self.heads, self.nope, self.kv_lora), + transpose(kvb[hn:].reshape(self.heads, self.v_dim, self.kv_lora), 1, 2), + kvb.t(), + w[f"l{i}_o"].t(), + w[f"l{i}_norm2"], + ) + ) + nxt.append(w[f"l{i + 1}_norm1"] if i + 1 < self.num_layers else w["final_norm"]) + # The fused gate_up GEMM assumes one activation scale and one + # weight global scale for both halves; the checkpoint stores them + # per projection, so the equality the fusion rests on is asserted. + assert torch.equal(w[f"l{i}_mlp_isc1"], w[f"l{i}_mlp_isc1_up"]), ( + f"layer {i}: gate/up input_scale differ; the fused gate_up " + "GEMM needs one activation scale" + ) + assert torch.equal(w[f"l{i}_mlp_ws2_1"], w[f"l{i}_mlp_ws2_1_up"]), ( + f"layer {i}: gate/up weight_scale_2 differ; the fused gate_up GEMM needs one alpha" + ) + # The checkpoint stores reciprocals: `input_scale = amax/(448*6) + # = 1/g_act` and `weight_scale_2 = 1/g_w`, so the quantizer's + # global scale is `1/input_scale` and the GEMM's alpha is their + # product — both straight off disk, no further reciprocal. + isc1 = w[f"l{i}_mlp_isc1"] + isc2 = w[f"l{i}_mlp_isc2"] + mlp.append( + ( + w[f"l{i}_mlp_gu_w"], + w[f"l{i}_mlp_gu_s"], + (isc1 * w[f"l{i}_mlp_ws2_1"]).contiguous(), + (1.0 / isc1).contiguous(), + w[f"l{i}_mlp_dn_w"], + w[f"l{i}_mlp_dn_s"], + (isc2 * w[f"l{i}_mlp_ws2_2"]).contiguous(), + (1.0 / isc2).contiguous(), + ) + ) + if i < self.dense_layers: + moe.append(None) + continue + assert torch.equal(w[f"l{i}_e_isc1"], w[f"l{i}_e_isc1_up"]), ( + f"layer {i}: expert gate/up input_scale differ" + ) + assert torch.equal(w[f"l{i}_e_ws2_1"], w[f"l{i}_e_ws2_1_up"]), ( + f"layer {i}: expert gate/up weight_scale_2 differ" + ) + # One quantization of the gathered hidden states feeds every + # expert on every rank, so the routed FC1 activation scale must + # be a single value — and the same value everywhere, or the four + # windows would not sum to the whole layer. The shared expert + # sees every token where each routed expert sees only its own + # subset, and the checkpoint's shared-expert input_scale is + # exactly the max over *all* routed ones — the conservative + # choice that cannot saturate an activation block. Asserted over + # the full 256, which is why the scalars are loaded whole. + e_isc1 = w[f"l{i}_e_isc1"] + assert torch.equal(e_isc1.max(), w[f"l{i}_mlp_isc1"][0]), ( + f"layer {i}: shared-expert input_scale is not the max over " + "the routed experts; the shared activation quantization " + "would saturate" + ) + e_isc2 = w[f"l{i}_e_isc2"][window] + gate1 = (w[f"l{i}_mlp_isc1"][0] * w[f"l{i}_e_ws2_1"][window]).contiguous() + moe.append( + ( + w[f"l{i}_router"].t(), + w[f"l{i}_router_bias"], + w[f"l{i}_fc1_w"], + view_dtype(w[f"l{i}_fc1_s"], torch.float8_e4m3fn), + w[f"l{i}_fc2_w"], + view_dtype(w[f"l{i}_fc2_s"], torch.float8_e4m3fn), + # output1_scale_scalar, output1_scale_gate_scalar, + # output2_scale_scalar — the FC1 alpha, that alpha times + # the per-expert FC2 activation global scale, and the FC2 + # alpha, each [local_num_experts]. Swapping the first two + # is finite and silent. + (gate1 / e_isc2).contiguous(), + gate1, + (e_isc2 * w[f"l{i}_e_ws2_2"][window]).contiguous(), + # The routed FC1 activation global scale: the same + # 1/input_scale the shared-expert GEMM quantizes with, + # in the linear scale layout the MoE runner reads. + (1.0 / w[f"l{i}_mlp_isc1"]).contiguous(), + ) + ) + self._attn, self._mlp, self._moe, self._next_norm = attn, mlp, moe, nxt + if self.mtp_enabled: + self._mtp = self._derive_mtp() + # The side stream the shared-expert / dense-MLP branch runs on. Built + # here because a stream must exist before the first forward: creating + # one inside a CUDA-graph capture is not a capturable operation, and + # the runtime's first eager forwards are already past this point. + self._side_stream = torch.cuda.Stream(device=device) + + def _derive_mtp(self) -> dict: + """The MTP layer's operand set, derived exactly as a trunk layer's is: + column-major GEMM views (`.t()` is zero-copy), the two MLA absorption + operands split out of the row-regrouped kv_b_proj, and the same fp8 KV + scale check. No NVFP4 scalars — this module is bf16 throughout, so its + expert stacks go to `fused_moe` as they are stored.""" + w = self.w + hn = self.heads * self.nope + for role in ("k_scale", "v_scale"): + s = w[f"mtp_{role}"] + assert torch.equal(s, torch.ones_like(s)), ( + f"MTP layer: {role} is {s.item()}, not 1.0; the fp8 MLA " + "context path is only correct at a KV scaling factor of 1.0, " + "and this assembly passes no scale tensors" + ) + kvb = w["mtp_kvb"] + return { + "enorm": w["mtp_enorm"], + "hnorm": w["mtp_hnorm"], + "eh": w["mtp_eh"].t(), + "norm1": w["mtp_norm1"], + "qa": w["mtp_qa"].t(), + "q_norm": w["mtp_q_norm"], + "qb": w["mtp_qb"].t(), + "kva": w["mtp_kva"].t(), + "kv_norm": w["mtp_kv_norm"], + "k_b": kvb[:hn].reshape(self.heads, self.nope, self.kv_lora), + "v_b_t": transpose(kvb[hn:].reshape(self.heads, self.v_dim, self.kv_lora), 1, 2), + "kvb": kvb.t(), + "o": w["mtp_o"].t(), + "norm2": w["mtp_norm2"], + "router": w["mtp_router"].t(), + "router_bias": w["mtp_router_bias"], + "fc1": w["mtp_fc1"], + "fc2": w["mtp_fc2"], + "sh_gu": w["mtp_sh_gu"].t(), + "sh_dn": w["mtp_sh_dn"].t(), + "head_norm": w["mtp_head_norm"], + } + + def _check_step_contract(self, md, position_ids) -> None: + """First-forward fail-fast: the metadata fields this target consumes + must exist (private trtllm surface, pinned by version), the paged + latent pool must be the single pool the MLA entries are certified + over at the page size their fp8 column covers, the rope table must + cover every position the engine admits, and every feature this target + holds inert must actually be off. Everything checked is fixed at + engine construction — once per model instance is sound.""" + missing = [name for name in _STEP_FIELDS if not hasattr(md, name)] + assert not missing, f"metadata fields missing: {missing}" + # position_ids is not consumed: the MLA ops derive each context + # token's position from its row index within the sequence and each + # generation token's from sequence_length - 1. Checked anyway so a + # layout change upstream is loud rather than silent. + assert position_ids.dtype == torch.int32, position_ids.dtype + pools = {row[0] for row in md.host_kv_cache_pool_mapping.tolist()} + assert pools == {0}, ( + f"multi-pool KV addressing is not certified; layer->pool ids {sorted(pools)}" + ) + # Every MLA entry's fp8-e4m3 column is page 32 only (the bf16 columns + # also carry 64). 32 is what a default KvCacheConfig produces. + assert md.tokens_per_block == 32, ( + f"the fp8 latent-pool column of every MLA entry is certified at " + f"tokens_per_block 32; this engine built {md.tokens_per_block}" + ) + # The rope table must cover every position the engine admits: a short + # table is read out of bounds with no check, and `rope_max_positions`, + # the argument that looks like it bounds this, is one of the inert + # seven. `max_position_embeddings` is the right size under the identity + # config, but **a speculative_config inflates the engine's max_seq_len + # past it** — measured 163840 -> 163848 at `max_draft_len: 3`, which is + # more than the `max_draft_len - 1` extra KV tokens per sequence the + # runtime reference documents, so the engine's own number is taken + # rather than a formula. Growing here is exact: every row of the table + # depends only on its own position, so the rows the identity config + # uses are bit-identical either way. This runs on the first forward, + # before any CUDA-graph capture. + if md.max_seq_len > self._rope_positions: + self._rope_positions = md.max_seq_len + self._rope = self._rope_tables(self.w["final_norm"].device, self._rope_positions) + # Context flavor, fixed at engine construction: with the cached-KV + # surface present the target runs append -> gather -> up-project -> + # explicit-K/V FMHA, which serves reused and fresh sequences alike; + # without it (block reuse off) no context sequence can carry a + # prefix, and the fresh-prefill flavor with the in-kernel rope and + # append is the whole context path. The two cache ops reject any + # index dtype but int64. + self._cached_ctx = all(hasattr(md, n) for n in _CACHED_CTX_FIELDS) and bool( + md.enable_context_mla_with_cached_kv + ) + if self._cached_ctx: + for name in ("ctx_cached_token_indptr", "ctx_kv_indptr"): + t = getattr(md, name) + assert t.dtype == torch.int64 and t.is_cuda, (name, t.dtype, t.device) + assert md.effective_beam_width == 1, "beam search is not implemented" + assert md.cache_indirection is None, "beam search is not implemented" + assert md.block_ids_per_seq is None, "the FlashMLA layout is not implemented" + assert md.flash_mla_tile_scheduler_metadata is None, "FlashMLA is not implemented" + assert md.flash_mla_num_splits is None, "FlashMLA is not implemented" + assert not md.is_cross, "cross attention is not implemented" + # The spec-dec **mask** machinery, which is a different thing from + # speculative decoding being on. Under a linear-tree MTP on a trtllm-gen + # arch the + # backend computes `is_spec_decoding_enabled and (not trtllm_gen_arch + # or is_spec_dec_dynamic_tree)` and gets False, so every + # `spec_decoding_*` tensor stays None — which is exactly the inert + # group `_CALL_INERT` holds and every MLA column certifies. Drafting + # itself reaches the attention ops through `predicted_tokens_per_seq` + # alone. This assert is the precondition those inert values rest on: it + # fires on a tree draft, or on a pre-Blackwell arch where the gating + # does not force the mask off, and either would need the mask surface + # certified first. + assert not md.is_spec_decoding_enabled and not md.use_spec_decoding, ( + "the spec-decoding mask surface is live; this target holds the " + "whole spec_decoding_* group at its inert values, which is only " + "valid while the Blackwell linear-tree gating keeps it off" + ) + self._step_contract_checked = True + + def _dp_rows(self, md, num_tokens: int) -> int: + """The uniform row count every rank pads its token block to before the + MoE round trip: `max` over the group's per-rank token counts. + + Attention DP gives each rank a different batch by construction, so the + split is engine state, not something a rank can derive locally: + `attn_metadata.all_rank_num_tokens` is where the engine publishes it, + as a plain list of host ints identical on every rank. Every rank + therefore computes the same maximum, and both collectives run in their + uniform form (`sizes=None`). + + **Both collectives are run uniform on purpose, and the ragged form is + not used at all.** The even split is cheaper for the reduce-scatter at + these row counts, and `sizes` is a host argument baked into a + CUDA-graph capture, so only the uniform form is replayable. + + CUDA-graph classification: this returns a host int, but under capture + the counts are uniform (the engine pads the decode batch to a captured + size on every rank), so the padding is zero rows and the graph + contains no padding at all — and no host value reaches either + collective, whose row counts come from tensor shapes the graph + fixes.""" + counts = [int(n) for n in md.all_rank_num_tokens] + assert len(counts) == self.dp_size, (counts, self.dp_size) + assert counts[self.rank] == num_tokens, ( + f"rank {self.rank} holds {num_tokens} rows but the engine's " + f"all_rank_num_tokens says {counts[self.rank]} ({counts}); the " + "collectives would gather the wrong split" + ) + return max(counts) + + def _dense_mlp(self, x, params, dt): + """One NVFP4 SwiGLU MLP over this rank's own tokens: fused gate_up + GEMM, silu_and_mul, down GEMM. Replicated weights, so the result is + complete — nothing to reduce. Both quantizations emit the + 128x4-swizzled scale buffer nvfp4_gemm consumes.""" + gu_w, gu_s, gu_alpha, gu_g, dn_w, dn_s, dn_alpha, dn_g = params + xq, xsf = fp4_quantize(x, gu_g, _SF_VEC, False, _SF_SWIZZLED) + gu = nvfp4_gemm(xq, gu_w, xsf, gu_s, gu_alpha, dt) + act = flashinfer_silu_and_mul(gu) + aq, asf = fp4_quantize(act, dn_g, _SF_VEC, False, _SF_SWIZZLED) + return nvfp4_gemm(aq, dn_w, asf, dn_s, dn_alpha, dt) + + def forward( + self, + attn_metadata: AttentionMetadata, + input_ids: torch.IntTensor | None = None, + position_ids: torch.IntTensor | None = None, + inputs_embeds: torch.FloatTensor | None = None, + lora_params: dict | None = None, + **kwargs, + ) -> torch.Tensor: + attn_w, mlp_w, moe_w = self._attn, self._mlp, self._moe + next_norm = self._next_norm + assert attn_w is not None and mlp_w is not None and moe_w is not None, ( + "load_weights must run before forward" + ) + assert next_norm is not None, "load_weights must run before forward" + assert position_ids is not None + assert isinstance(attn_metadata, TrtllmAttentionMetadata) + # Inputs this target does not implement must fail loudly, not be + # silently dropped (unlike runtime-owned features, which pass through). + assert lora_params is None, "LoRA is not implemented by this target" + # The trunk does not read `spec_metadata` even when the engine drafts: + # on MTP_EAGLE_ONE_MODEL the runtime sets `layers_to_capture = ()`, so + # no hidden-state capture hook is owed and the shell keeps the argument + # to itself. Keeping this assert is a stronger guarantee than deleting + # it. + assert kwargs.get("spec_metadata") is None, ( + "spec_metadata reached the trunk; the shell owns the draft loop and must not forward it" + ) + md = attn_metadata + if not self._step_contract_checked: + self._check_step_contract(md, position_ids) + # Read after the contract check: that is where a table too short for + # the engine's admitted max_seq_len is regrown. + rope = self._rope + assert rope is not None, "load_weights must run before forward" + + step = _build_step_args(md) + num_ctx = md.num_contexts + tc = md.num_ctx_tokens + # The paged-cache address book, shared by the two MLA preprocessing + # ops and both attention calls. Under attention DP each rank owns a + # pool holding only its own requests' latent rows. + tokens_per_block = md.tokens_per_block + block_offsets = md.kv_cache_block_offsets + pool_ptrs = md.host_kv_cache_pool_pointers + pool_map = md.host_kv_cache_pool_mapping + assert tokens_per_block is not None, "paged KV cache is required" + assert block_offsets is not None and pool_ptrs is not None, "paged KV cache is required" + assert pool_map is not None, "paged KV cache is required" + + if inputs_embeds is None: + assert input_ids is not None + h = embedding(input_ids, self.w["embed"]) + else: + h = inputs_embeds + num_tokens = h.shape[0] + gen = num_tokens - tc + dt = h.dtype + dev = h.device + # Query tokens per generation sequence, which is the MLA generation + # call's `predicted_tokens_per_seq`. 1 for an ordinary decode step; + # under MTP a generation request arrives carrying its whole draft + # chain, `runtime_draft_len + 1` rows, and that count is the *only* + # thing expressing the taller query block to the attention ops. Derived + # from the metadata rather than from a spec object, so it costs nothing + # when MTP is off and is a per-capture host constant when it is on. + gen_seqs = md.num_seqs - num_ctx + gen_p = gen // gen_seqs if gen_seqs else 1 + assert gen_p * gen_seqs == gen, ( + f"{gen} generation rows do not divide over {gen_seqs} generation " + "sequences; a ragged per-sequence draft length cannot be expressed " + "to the MLA generation call" + ) + # The context sequences' cached+new latent-KV row count: what the + # cache gather returns and what the context FMHA attends over. + ctx_kv_tokens = int(md.host_total_kv_lens[0]) if tc else 0 + # Every rank pads its token block to the group-wide maximum, so both + # MoE collectives run in their uniform form; the padded rows are + # sliced off after the reduce-scatter and never reach the residual + # stream. Zero rows quantize to an all-zero NVFP4 block (scale byte + # 0x00, zero data) and route like any other token, so they perturb + # nothing but wasted expert work. + dp_rows = self._dp_rows(md, num_tokens) + pad_rows = dp_rows - num_tokens + + attn_out = empty([num_tokens, self.heads * self.v_dim], dt, dev) + attn_ctx, attn_gen = split(attn_out, [tc, gen], 0) + x = flashinfer_rmsnorm(h, self.w["l0_norm1"], self.eps) + residual = h + for i in range(self.num_layers): + ( + w_qa, + w_q_norm, + w_qb, + w_kva, + w_kv_norm, + k_b, + v_b_t, + w_kvb, + w_o, + w_n2, + ) = attn_w[i] + # The q-LoRA pair: down-project to q_lora_rank, RMS-norm the + # latent, up-project to the per-head [nope | rope] rows. + q = cublas_mm(flashinfer_rmsnorm(cublas_mm(x, w_qa), w_q_norm, self.eps), w_qb) + kva = cublas_mm(x, w_kva) + ckv_raw, k_pe = split(kva, [self.kv_lora, self.rope], -1) + ckv = flashinfer_rmsnorm(ckv_raw, w_kv_norm, self.eps) + latent = concat([ckv, k_pe], -1) + q_ctx, q_gen = split(q, [tc, gen], 0) + latent_ctx, latent_gen = split(latent, [tc, gen], 0) + + if tc: + if self._cached_ctx: + # Reuse-capable context: rotate q_pe/k_pe in place and + # append this step's latent rows (quantized to e4m3 on the + # way into the pool, at the write-side scale — None = 1.0), + # read each sequence's whole [cached + new] latent range + # back out (dequantized at the read-side scale — None = + # 1.0), up-project it, and attend over explicit K/V. One + # call serves a batch mixing reused and fresh sequences — + # a fresh one is the cached_s == 0 case. A cached prefix + # therefore pays fp8 twice: the gather widens the pool's + # e4m3 rows to bf16 and the FMHA below quantizes the + # up-projected result straight back. + mla_rope_append_paged_kv_assign_q( + q_ctx, + latent_ctx, + num_ctx, + md.ctx_cached_token_indptr, + md.ctx_kv_indptr, + int(md.max_ctx_seq_len), + rope["rotary_cos_sin"], + self.heads, + self.nope, + self.rope, + self.kv_lora, + block_offsets, + pool_ptrs, + pool_map, + None, + _KV_RESIDUAL_DIM, + i, + tokens_per_block, + md.max_seq_len, + 1, + self.quant_mode, + ) + ckv_full, k_pe_full = load_paged_kv_cache_for_mla( + dt, + num_ctx, + ctx_kv_tokens, + int(md.max_ctx_kv_len), + md.ctx_kv_indptr, + block_offsets, + pool_ptrs, + pool_map, + None, + i, + self.kv_lora, + self.rope, + tokens_per_block, + md.max_seq_len, + 1, + self.quant_mode, + ) + latent_arg = None + else: + # Block reuse off: no context sequence can carry a + # prefix, so the fresh-prefill flavor does the rope and + # the append inside the attention call and the latent + # rows never leave registers. + assert ctx_kv_tokens == tc, ( + "a context sequence arrived with a cached KV prefix " + "while the cached-KV metadata surface is off" + ) + ckv_full, _ = split(ckv, [tc, gen], 0) + k_pe_full = None + latent_arg = latent_ctx + tkv = ctx_kv_tokens + # kv is packed [all heads' k_nope | all heads' v]: the context + # FMHA hard-codes V's row stride as the full packed width and + # reads the column block, so V must stay this split view. + kv = cublas_mm(ckv_full, w_kvb) + k_nope, v_view = split(kv, [self.heads * self.nope, self.heads * self.v_dim], -1) + k = empty([tkv, self.heads, self.qk_dim], dt, dev) + k_nope_dst, k_pe_dst = split(k, [self.nope, self.rope], -1) + copy_(k_nope_dst, reshape(k_nope, [tkv, self.heads, self.nope])) + if k_pe_full is not None: + # k_pe came back from the pool already rotated; every + # query head shares it. On the fresh-prefill flavor the + # rope slice is left uninitialized instead — that call + # overwrites it in place from latent_cache. + copy_( + k_pe_dst, + expand( + reshape(k_pe_full, [tkv, 1, self.rope]), + [tkv, self.heads, self.rope], + ), + ) + thop_attention( + q=q_ctx, + k=reshape(k, [tkv, self.heads * self.qk_dim]), + v=v_view, + output=attn_ctx, + latent_cache=latent_arg, + q_pe=None, + local_layer_idx=i, + is_fused_qkv=False, + attention_input_type=1, + num_heads=self.heads, + num_kv_heads=self.heads, + head_size=self.qk_dim, + q_lora_rank=self.q_lora, + kv_lora_rank=self.kv_lora, + qk_nope_head_dim=self.nope, + qk_rope_head_dim=self.rope, + v_head_dim=self.v_dim, + cu_q_seqlens=None, + cu_kv_seqlens=None, + fmha_scheduler_counter=None, + # The context call's row count is `num_ctx_tokens` + # whatever the generation phase carries, and the entry + # certifies the MLA context flavors at 1 only. + predicted_tokens_per_seq=1, + **step, + **rope, + **self._call, + ) + + if gen: + q3 = reshape(q_gen, [gen, self.heads, self.qk_dim]) + q_nope, q_pe = split(q3, [self.nope, self.rope], -1) + fused_q = empty([gen, self.heads, self.lat_dim], dt, dev) + fq_nope, _ = split(fused_q, [self.kv_lora, self.rope], -1) + # Absorbed q: (q_nope @ W_k_nope) is what the latent-space + # dot product needs. Over an fp8 pool the next call **reads** + # this half to build the quantized query, so this BMM must + # have finished first — the two are issued in this order on + # one stream, which is what makes that safe. (On a bf16 pool + # they write disjoint halves and may overlap; assembling from + # that reading and switching to an fp8 cache is a silent race.) + bmm_out(transpose(q_nope, 0, 1), k_b, transpose(fq_nope, 0, 1)) + cu_q = empty([gen + 1], torch.int32, dev) + cu_kv = empty([gen + 1], torch.int32, dev) + counter = empty([1], torch.uint32, dev) + # The fp8 decode triple: the quantized query the FMHA reads + # instead of fused_q, and the two folded softmax/output scales + # it takes instead of either kv scale tensor. All three are + # written by the call below from q_scaling, the MLA dims and + # the read-side factor (None = 1.0). + quant_q = empty([gen, self.heads, self.lat_dim], torch.float8_e4m3fn, dev) + bmm1_scale = empty([2], torch.float32, dev) + bmm2_scale = empty([1], torch.float32, dev) + mla_rope_generation( + fused_q, + q_pe, + latent_gen, + rope["rotary_cos_sin"], + cu_q, + cu_kv, + counter, + bmm1_scale, + bmm2_scale, + quant_q, + md.kv_lens_cuda_runtime, + md.kv_lens_runtime, + md.prompt_lens_cpu_runtime, + num_ctx, + block_offsets, + pool_ptrs, + pool_map, + None, + None, + None, + None, + None, + [None, None], + gen_p, + i, + self.heads, + 1, + self.lat_dim, + _KV_RESIDUAL_DIM, + tokens_per_block, + md.max_seq_len, + 1, + self.quant_mode, + self.q_scaling, + self.q_lora, + self.kv_lora, + self.nope, + self.rope, + self.v_dim, + True, + ) + lat_out = empty([gen, self.heads * self.kv_lora], dt, dev) + thop_attention( + q=reshape(fused_q, [gen, self.heads * self.lat_dim]), + k=None, + v=None, + output=lat_out, + latent_cache=latent_gen, + q_pe=q_pe, + local_layer_idx=i, + is_fused_qkv=True, + attention_input_type=2, + num_heads=self.heads, + num_kv_heads=1, + head_size=self.lat_dim, + q_lora_rank=self.q_lora, + kv_lora_rank=self.kv_lora, + qk_nope_head_dim=self.nope, + qk_rope_head_dim=self.rope, + v_head_dim=self.kv_lora, + cu_q_seqlens=cu_q, + cu_kv_seqlens=cu_kv, + fmha_scheduler_counter=counter, + mla_bmm1_scale=bmm1_scale, + mla_bmm2_scale=bmm2_scale, + quant_q_buffer=quant_q, + # The whole of drafting, as far as this op is concerned: + # the query block is `gen_p` rows per generation sequence, + # token-major, and `gen_p` is what produces the + # bottom-right-aligned within-block causal mask (draft row + # t sees [0, L_g - gen_p + t] and none of its later + # siblings). Certified at 1..4 over this fp8 cell. + predicted_tokens_per_seq=gen_p, + **step, + **rope, + **self._call, + ) + bmm_out( + transpose(reshape(lat_out, [gen, self.heads, self.kv_lora]), 0, 1), + v_b_t, + transpose(reshape(attn_gen, [gen, self.heads, self.v_dim]), 0, 1), + ) + + # Replicated o_proj over this rank's own tokens: complete as it + # stands, so the residual stream is updated with no collective. + o = cublas_mm(attn_out, w_o) + flashinfer_fused_add_rmsnorm(o, residual, w_n2, self.eps) + if moe_w[i] is None: + mlp_out = self._dense_mlp(o, mlp_w[i], dt) + else: + # The shared expert and the routed round trip read the same `o` + # and meet only at the add below, so they are independent — but + # on one stream the shared expert's five kernels sit in front of + # a gather that is 94% exclusive on the device. Forking it onto + # a side stream lets it run inside that window. Both collective + # contracts certify a side stream joined to the current one on + # both ends, which is exactly the shape here; the fork/join pair + # is also what propagates a CUDA-graph capture into the branch + # and back, so a decode capture records both streams. + side = self._side_stream + main = torch.cuda.current_stream() + assert side is not None, "derive_after_load must run before forward" + side.wait_stream(main) + with torch.cuda.stream(side): + shared = self._dense_mlp(o, mlp_w[i], dt) + ( + w_router, + router_bias, + fc1_w, + fc1_s, + fc2_w, + fc2_s, + o1, + o1_gate, + o2, + routed_g, + ) = moe_w[i] + # The expert-parallel round trip. Gathering *before* the + # router is what keeps the four windows tiling the routing + # space exactly once: every rank routes the identical full + # token set, so a token's top-8 ids agree across ranks and + # each id falls in exactly one window. + o_pad = o if not pad_rows else pad(o, [0, 0, 0, pad_rows]) + o_all = allgather(o_pad, None, self.dp_group) + # The expert call is chunked to `_MOE_MAX_T` rows: the + # gathered set reaches 4 * max_num_tokens and both the runner's + # and the routing op's certified token columns stop at 8192. + # Every token is independent through routing, quantization and + # the expert GEMMs, so the chunks are the whole call, re-joined. + parts = [] + for chunk in split(o_all, _moe_chunk_sizes(o_all.shape[0]), 0): + # Raw bf16 logits: noaux_tc_op applies the sigmoid + # itself, and its weight dtype follows the logits, so + # bf16 in means the MoE runner's bf16 topk_weights need + # no cast (the fp32 correction bias does not change that). + logits = cublas_mm(chunk, w_router) + topk_w, topk_ids = noaux_tc_op( + logits, + router_bias, + self.n_group, + self.topk_group, + self.topk, + self.routed_scale, + ) + # A second quantization of the same hidden states: the + # MoE runner reads the linear scale buffer as + # float8_e4m3fn, never the swizzled one the shared GEMM + # above consumed. + xq, xsf = fp4_quantize(chunk, routed_g, _SF_VEC, False, _SF_LINEAR) + parts.append( + fp4_block_scale_moe_runner( + None, + None, + xq, + view_dtype(xsf, torch.float8_e4m3fn), + fc1_w, + fc1_s, + None, + None, + None, + None, + fc2_w, + fc2_s, + None, + o1, + o1_gate, + o2, + self.num_experts, + self.topk, + None, + None, + self.moe_inter, + self.expert_offset, + self.local_experts, + None, + _ROUTING_METHOD_INERT, + True, + _ACT_TYPE_SWIGLU, + topk_w, + topk_ids, + )[0] + ) + routed_all = parts[0] if len(parts) == 1 else concat(parts, 0) + # The four windows' partials over the whole token set sum to + # the layer's routed output; the scatter hands this rank back + # exactly its own rows. bf16 in: the op sums, and it sums + # float8 as raw bytes. + routed_pad = reducescatter(routed_all, None, self.dp_group) + routed = ( + routed_pad if not pad_rows else split(routed_pad, [num_tokens, pad_rows], 0)[0] + ) + # Join: the add is the first reader of `shared` on this stream, + # and `shared` stays referenced until then, so the branch's + # allocations cannot be recycled underneath it. + main.wait_stream(side) + mlp_out = add(routed, shared) + flashinfer_fused_add_rmsnorm(mlp_out, residual, next_norm[i], self.eps) + x = mlp_out + return x + + +class MTPLayer: + """The checkpoint's multi-token-prediction module, at layer index + `num_hidden_layers`, replayed once per draft step. + + Structurally one more decoder layer with a front end bolted on that mixes + in the embedding of the token being predicted: + + e = embed_tokens(input_ids) # the NEXT token + x = eh_proj(concat(enorm(e), hnorm(h))) # [T, 2H] -> [T, H] + x = x + MLA(input_layernorm(x)) + x = x + MoE(post_attention_layernorm(x)) + logits = lm_head(shared_head.norm(x)) # `shared_head`, below + + `h` is the hidden state the runtime hands in: the trunk's own output at + draft step 0, and this layer's own output at every step after. The layer + returns `x` **un-normalized** — `shared_head` is where the module's own + norm is applied — and the runtime slices that with its `gather_ids` for the + logits and feeds it forward as the next step's `h`. + + A deliberate copy of the trunk's layer body rather than a shared helper: + the trunk's loop is inlined and tuner-specialized, and three things differ + here anyway — the front end, the layer index (`num_hidden_layers`, the + extra pool layer the engine adds under a one-model MTP mode), and the MoE + dtype. The checkpoint excludes `model.layers.61*` from NVFP4 wholesale, so + this module is bf16 throughout and its routed experts go to `fused_moe` + where the trunk's go to `fp4_block_scale_moe_runner`. + + Not an `nn.Module`, and `mtp_layers` is a plain list: every weight this + layer reads lives in the core's ParameterDict, so registering the layer + again would put the trunk's parameters on a second path through the shell's + module tree. The runtime never inspects the container beyond `mtp_layers`, + `embed_tokens`, `lm_head` and an absent `model.d2t`. + + Two arguments of the documented calling convention are accepted and unused, + for the same reasons the trunk ignores them: `position_ids` (the MLA ops + take every position from `sequence_length`, not from this tensor) and + `spec_metadata` (the one thing the layer needs off it, the DP padding + basis, arrives as the `all_rank_num_tokens` keyword instead).""" + + def __init__(self, core: StaircaseCore, logits_processor) -> None: + self.core = core + self.logits_processor = logits_processor + + def __call__(self, *args, **kwargs) -> torch.Tensor: + return self.forward(*args, **kwargs) + + def shared_head( + self, hidden_states, lm_head, attn_metadata, return_context_logits + ) -> torch.Tensor: + """The module's own output head: its own RMS norm — a **distinct** + parameter from the trunk's `model.norm` — then the trunk's `lm_head`, + which the checkpoint's `shared_head.head` is a bitwise copy of. The + projection itself is the inherited shell's `logits_processor`, exactly + as on the non-speculative path.""" + mw = self.core._mtp + assert mw is not None, "load_weights must run before the draft loop" + return self.logits_processor.forward( + flashinfer_rmsnorm(hidden_states, mw["head_norm"], self.core.eps), + lm_head, + attn_metadata, + return_context_logits, + ) + + def _shared_mlp(self, x, mw): + """The module's shared expert: one bf16 SwiGLU MLP, replicated, over + this rank's own tokens — the same structure as the trunk's shared + expert at the same width, without the NVFP4 quantize/dequantize pair. + `sh_gu` holds the gate rows first, the half order silu_and_mul reads.""" + return cublas_mm(flashinfer_silu_and_mul(cublas_mm(x, mw["sh_gu"])), mw["sh_dn"]) + + def _routed_experts(self, x, mw, all_rank_num_tokens, rows, dt): + """The module's expert-parallel round trip, the trunk's shape at a + different dtype: pad to the group-wide row count, gather **before** the + router so all four ranks route the identical token set, run this rank's + 64-expert window, reduce-scatter the partials back. + + Two things differ from the trunk's. The router GEMM emits **fp32** + logits: `fused_moe` demands fp32 `token_final_scales` where the + trtllm-gen runner demands bf16, and `noaux_tc_op`'s weight dtype + follows its logits, so the dtype is chosen here rather than cast later. + And the expert call takes the global expert ids directly with + `ep_size`/`ep_rank` shifting this rank's window, where the trtllm-gen + runner takes an explicit offset/count pair. + + `fused_moe` requires a token's `topk` ids to be **distinct** above 256 + tokens — a repeat reads out of bounds in `finalizeMoeRoutingKernel`, + faulting or returning ~200-460 ulp of silent garbage — and this call + drives `T` to 8192. `noaux_tc_op` returns the indices of the top-k + largest corrected scores, and a top-k over expert indices cannot repeat + one, so the precondition holds structurally.""" + core = self.core + dp_rows = _mtp_dp_rows(all_rank_num_tokens, core.rank, core.dp_size, rows) + pad_rows = dp_rows - rows + x_pad = x if not pad_rows else pad(x, [0, 0, 0, pad_rows]) + x_all = allgather(x_pad, None, core.dp_group) + parts = [] + for chunk in split(x_all, _moe_chunk_sizes(x_all.shape[0]), 0): + logits = cublas_mm(chunk, mw["router"], None, torch.float32) + topk_w, topk_ids = noaux_tc_op( + logits, + mw["router_bias"], + core.n_group, + core.topk_group, + core.topk, + core.routed_scale, + ) + parts.append( + fused_moe( + chunk, + topk_ids, + topk_w, + mw["fc1"], + None, + mw["fc2"], + None, + dt, + [], + ep_size=core.ep_size, + ep_rank=core.ep_rank, + activation_type=_FUSED_MOE_ACT_SWIGLU, + )[0] + ) + routed_all = parts[0] if len(parts) == 1 else concat(parts, 0) + routed_pad = reducescatter(routed_all, None, core.dp_group) + return routed_pad if not pad_rows else split(routed_pad, [rows, pad_rows], 0)[0] + + def forward( + self, + embed_tokens: torch.Tensor, + all_rank_num_tokens, + input_ids: torch.Tensor, + hidden_states: torch.Tensor, + attn_metadata: AttentionMetadata, + **kwargs, + ) -> torch.Tensor: + core = self.core + mw, rope = core._mtp, core._rope + assert mw is not None and rope is not None, "load_weights must run before the draft loop" + assert isinstance(attn_metadata, TrtllmAttentionMetadata) + assert core._step_contract_checked, ( + "the trunk's first-forward contract check has not run; the shell " + "calls the core before the worker, so this cannot be reached first" + ) + md = attn_metadata + step = _build_step_args(md) + rows = hidden_states.shape[0] + assert input_ids.shape[0] == rows, ( + f"the draft step's {input_ids.shape[0]} token ids and " + f"{rows} hidden-state rows must describe the same tokens" + ) + dt = hidden_states.dtype + dev = hidden_states.device + num_ctx = md.num_contexts + tc = md.num_ctx_tokens + gen = rows - tc + # **Read the phase from the metadata on every call.** The worker + # rewrites `attn_metadata` in place between draft step 0 and step 1+ — + # `_seq_lens` filled with 1, `num_contexts` and `num_ctx_tokens` to 0, + # one token per generation request — and this layer is invoked N times + # inside one forward, so a value computed on the first call is wrong on + # the rest (and, under capture, would be frozen wrong at all N + # positions). The read is host-side and sync-free: `on_update()` + # recomputes these from the pinned-host `_seq_lens` and the loop + # triggers it at exactly that boundary. The quotient is + # `runtime_draft_len + 1` on step 0 and exactly 1 afterwards. + gen_seqs = md.num_seqs - num_ctx + gen_p = gen // gen_seqs if gen_seqs else 1 + assert gen_p * gen_seqs == gen, ( + f"{gen} generation rows do not divide over {gen_seqs} generation " + "sequences; a ragged per-sequence draft length cannot be expressed " + "to the MLA generation call" + ) + tokens_per_block = md.tokens_per_block + block_offsets = md.kv_cache_block_offsets + pool_ptrs = md.host_kv_cache_pool_pointers + pool_map = md.host_kv_cache_pool_mapping + assert tokens_per_block is not None, "paged KV cache is required" + assert block_offsets is not None and pool_ptrs is not None, "paged KV cache is required" + assert pool_map is not None, "paged KV cache is required" + ctx_kv_tokens = int(md.host_total_kv_lens[0]) if tc else 0 + # The engine raises the KV pool's layer count by + # `num_nextn_predict_layers` under a one-model MTP mode, so this + # module's attention addresses layer index `num_hidden_layers` in the + # same single pool the trunk's 61 layers use. Nothing in the target's + # config stub or manifest declares that. + layer_idx = core.num_layers + + e = embedding(input_ids, embed_tokens) + en = flashinfer_rmsnorm(e, mw["enorm"], core.eps) + hn = flashinfer_rmsnorm(hidden_states, mw["hnorm"], core.eps) + halves = [en, hn] if _MTP_EMBED_BLOCK_FIRST else [hn, en] + x = cublas_mm(concat(halves, -1), mw["eh"]) + + residual = x + xn = flashinfer_rmsnorm(x, mw["norm1"], core.eps) + attn_out = empty([rows, core.heads * core.v_dim], dt, dev) + attn_ctx, attn_gen = split(attn_out, [tc, gen], 0) + q = cublas_mm( + flashinfer_rmsnorm(cublas_mm(xn, mw["qa"]), mw["q_norm"], core.eps), + mw["qb"], + ) + kva = cublas_mm(xn, mw["kva"]) + ckv_raw, k_pe = split(kva, [core.kv_lora, core.rope], -1) + ckv = flashinfer_rmsnorm(ckv_raw, mw["kv_norm"], core.eps) + latent = concat([ckv, k_pe], -1) + q_ctx, q_gen = split(q, [tc, gen], 0) + latent_ctx, latent_gen = split(latent, [tc, gen], 0) + + if tc: + # A context request reaches the draft loop at step 0 only, fed + # `prompt[1:]` with its first accepted token written at the last + # position — the same row count and the same per-sequence lengths + # the trunk saw, so the context flavor the trunk settled at its + # first forward applies unchanged here. + if core._cached_ctx: + mla_rope_append_paged_kv_assign_q( + q_ctx, + latent_ctx, + num_ctx, + md.ctx_cached_token_indptr, + md.ctx_kv_indptr, + int(md.max_ctx_seq_len), + rope["rotary_cos_sin"], + core.heads, + core.nope, + core.rope, + core.kv_lora, + block_offsets, + pool_ptrs, + pool_map, + None, + _KV_RESIDUAL_DIM, + layer_idx, + tokens_per_block, + md.max_seq_len, + 1, + core.quant_mode, + ) + ckv_full, k_pe_full = load_paged_kv_cache_for_mla( + dt, + num_ctx, + ctx_kv_tokens, + int(md.max_ctx_kv_len), + md.ctx_kv_indptr, + block_offsets, + pool_ptrs, + pool_map, + None, + layer_idx, + core.kv_lora, + core.rope, + tokens_per_block, + md.max_seq_len, + 1, + core.quant_mode, + ) + latent_arg = None + else: + assert ctx_kv_tokens == tc, ( + "a context sequence arrived with a cached KV prefix while " + "the cached-KV metadata surface is off" + ) + ckv_full, _ = split(ckv, [tc, gen], 0) + k_pe_full = None + latent_arg = latent_ctx + tkv = ctx_kv_tokens + kv = cublas_mm(ckv_full, mw["kvb"]) + k_nope, v_view = split(kv, [core.heads * core.nope, core.heads * core.v_dim], -1) + k = empty([tkv, core.heads, core.qk_dim], dt, dev) + k_nope_dst, k_pe_dst = split(k, [core.nope, core.rope], -1) + copy_(k_nope_dst, reshape(k_nope, [tkv, core.heads, core.nope])) + if k_pe_full is not None: + copy_( + k_pe_dst, + expand( + reshape(k_pe_full, [tkv, 1, core.rope]), + [tkv, core.heads, core.rope], + ), + ) + thop_attention( + q=q_ctx, + k=reshape(k, [tkv, core.heads * core.qk_dim]), + v=v_view, + output=attn_ctx, + latent_cache=latent_arg, + q_pe=None, + local_layer_idx=layer_idx, + is_fused_qkv=False, + attention_input_type=1, + num_heads=core.heads, + num_kv_heads=core.heads, + head_size=core.qk_dim, + q_lora_rank=core.q_lora, + kv_lora_rank=core.kv_lora, + qk_nope_head_dim=core.nope, + qk_rope_head_dim=core.rope, + v_head_dim=core.v_dim, + cu_q_seqlens=None, + cu_kv_seqlens=None, + fmha_scheduler_counter=None, + predicted_tokens_per_seq=1, + **step, + **rope, + **core._call, + ) + + if gen: + q3 = reshape(q_gen, [gen, core.heads, core.qk_dim]) + q_nope, q_pe = split(q3, [core.nope, core.rope], -1) + fused_q = empty([gen, core.heads, core.lat_dim], dt, dev) + fq_nope, _ = split(fused_q, [core.kv_lora, core.rope], -1) + bmm_out(transpose(q_nope, 0, 1), mw["k_b"], transpose(fq_nope, 0, 1)) + cu_q = empty([gen + 1], torch.int32, dev) + cu_kv = empty([gen + 1], torch.int32, dev) + counter = empty([1], torch.uint32, dev) + quant_q = empty([gen, core.heads, core.lat_dim], torch.float8_e4m3fn, dev) + bmm1_scale = empty([2], torch.float32, dev) + bmm2_scale = empty([1], torch.float32, dev) + mla_rope_generation( + fused_q, + q_pe, + latent_gen, + rope["rotary_cos_sin"], + cu_q, + cu_kv, + counter, + bmm1_scale, + bmm2_scale, + quant_q, + md.kv_lens_cuda_runtime, + md.kv_lens_runtime, + md.prompt_lens_cpu_runtime, + num_ctx, + block_offsets, + pool_ptrs, + pool_map, + None, + None, + None, + None, + None, + [None, None], + gen_p, + layer_idx, + core.heads, + 1, + core.lat_dim, + _KV_RESIDUAL_DIM, + tokens_per_block, + md.max_seq_len, + 1, + core.quant_mode, + core.q_scaling, + core.q_lora, + core.kv_lora, + core.nope, + core.rope, + core.v_dim, + True, + ) + lat_out = empty([gen, core.heads * core.kv_lora], dt, dev) + thop_attention( + q=reshape(fused_q, [gen, core.heads * core.lat_dim]), + k=None, + v=None, + output=lat_out, + latent_cache=latent_gen, + q_pe=q_pe, + local_layer_idx=layer_idx, + is_fused_qkv=True, + attention_input_type=2, + num_heads=core.heads, + num_kv_heads=1, + head_size=core.lat_dim, + q_lora_rank=core.q_lora, + kv_lora_rank=core.kv_lora, + qk_nope_head_dim=core.nope, + qk_rope_head_dim=core.rope, + v_head_dim=core.kv_lora, + cu_q_seqlens=cu_q, + cu_kv_seqlens=cu_kv, + fmha_scheduler_counter=counter, + mla_bmm1_scale=bmm1_scale, + mla_bmm2_scale=bmm2_scale, + quant_q_buffer=quant_q, + predicted_tokens_per_seq=gen_p, + **step, + **rope, + **core._call, + ) + bmm_out( + transpose(reshape(lat_out, [gen, core.heads, core.kv_lora]), 0, 1), + mw["v_b_t"], + transpose(reshape(attn_gen, [gen, core.heads, core.v_dim]), 0, 1), + ) + + o = cublas_mm(attn_out, mw["o"]) + flashinfer_fused_add_rmsnorm(o, residual, mw["norm2"], core.eps) + shared = self._shared_mlp(o, mw) + routed = self._routed_experts(o, mw, all_rank_num_tokens, rows, dt) + # The layer's output is the residual stream itself, un-normalized: + # `shared_head` owns the module's norm, and the runtime feeds this + # tensor straight back in as the next draft step's `h`. + return add(residual, add(routed, shared)) + + +class DraftModel: + """The container the spec worker reaches this target's drafter through. + + The runtime never constructs it and inspects exactly four names on it: + `mtp_layers` (only `[0]` is ever indexed — MTP-Eagle replays one layer), + `embed_tokens` and `lm_head`, which it hands to the layer and to + `shared_head`, and `model.d2t`, read with a nested `getattr(..., None)` and + correctly absent here because draft and target share one vocabulary. + + `embed_tokens` is the **trunk's** embedding weight and `lm_head` the + trunk's head. The checkpoint's `model.layers.61.embed_tokens.weight` and + `.shared_head.head.weight` are `torch.equal` to those two, so pointing at + them is exact and saves 1.85 GB per rank; they stay a predicted non-load in + the weight manifest even with MTP on. + + `embed_tokens` resolves through the core on every read rather than being + snapshotted here: this container is built in the shell's `__init__`, where + every parameter is still a **meta** tensor, and the engine materializes the + registry by replacing those tensor objects. A reference captured at + construction stays on meta and fails at the first draft step.""" + + def __init__(self, core: StaircaseCore, lm_head, logits_processor) -> None: + self.core = core + self.mtp_layers = [MTPLayer(core, logits_processor)] + self.lm_head = lm_head + + @property + def embed_tokens(self) -> torch.Tensor: + return self.core.w["embed"] + + +@register_auto_model("StaircaseDeepseekR10528Nvfp4Sm103Dep4") +class StaircaseDeepseekR10528Nvfp4Sm103Dep4( + DecoderModelForCausalLM[StaircaseCore, PretrainedConfig] +): + def __init__(self, model_config: ModelConfig): + cfg = model_config.pretrained_config + assert cfg is not None + super().__init__( + StaircaseCore(model_config), + config=model_config, + hidden_size=cfg.hidden_size, + vocab_size=cfg.vocab_size, + ) + # The speculative branch, built only when a configs/ variant asked for + # it. `spec_config` is already resolved on the model config by the time + # the model is built, and the checkpoint — not the target — picked the + # mode: `num_nextn_predict_layers: 1` gives MTP-Eagle one-model, where + # `max_draft_len` is a serving knob rather than a checkpoint property. + # + # The branch is written out rather than inherited from the in-tree + # one-engine shell on purpose: that shell builds its drafter through a + # module-level function with no override point, which dispatches on the + # config's `model_type` — still the upstream family name, since the + # stub config patches `architectures` only — and would construct + # trtllm's own MTP layer instead of this target's. + self.spec_config = getattr(model_config, "spec_config", None) + self.draft_model = None + self.spec_worker = None + if self.spec_config is not None: + mode = self.spec_config.spec_dec_mode + assert mode.is_mtp_eagle_one_model(), ( + f"this target implements the one-model MTP-Eagle draft loop; " + f"the engine resolved spec_dec_mode {mode!r}" + ) + assert self.model.mtp_enabled, "core built without the MTP module" + assert hasattr(self, "logits_processor"), ( + "the inherited shell no longer exposes logits_processor; the " + "draft head and the shell's own gather both project through it" + ) + self.draft_model = DraftModel(self.model, self.lm_head, self.logits_processor) + self.spec_worker = get_spec_worker(self.spec_config, model_config, model_config.mapping) + + def forward( + self, + attn_metadata: AttentionMetadata, + input_ids: torch.IntTensor | None = None, + position_ids: torch.IntTensor | None = None, + inputs_embeds: torch.FloatTensor | None = None, + return_context_logits: bool = False, + spec_metadata=None, + lora_params: dict | None = None, + **kwargs, + ): + """One engine step. + + Without a spec worker this **delegates to the inherited base forward** + with every argument it was handed, rather than reimplementing it: the + release criterion was measured on that base forward, and calling it is + the only way to keep this path bit-identical. The parameter list + mirrors the base's exactly for that reason, and `resource_manager` + deliberately stays inside `**kwargs` — naming it here would drop it + from what the base receives. + + With one, the shell owns the logits gather: every one-model mode sets + `without_logits`, so the engine applies no second gather and the worker + needs the trunk's hidden states **ungathered** beside the gathered + logits. `position_ids` is handed on in the engine's `[1, T]` shape — + the worker squeezes it itself, and flattening here would produce a + silently wrong draft position sequence.""" + if self.spec_worker is None: + assert spec_metadata is None, ( + "spec_metadata arrived without a spec worker; this engine was " + "built without a speculative_config" + ) + return super().forward( + attn_metadata, + # A typing no-op: the base declares `input_ids: torch.IntTensor + # = None`, a non-Optional annotation with a None default, so + # the value is forwarded exactly as received — None included. + cast(torch.IntTensor, input_ids), + position_ids, + inputs_embeds, + return_context_logits, + spec_metadata, + lora_params, + **kwargs, + ) + assert spec_metadata is not None, ( + "the spec worker is built but the engine passed no spec_metadata" + ) + # `spec_metadata` is not forwarded into the core: the trunk genuinely + # does not read it — on MTP_EAGLE_ONE_MODEL the runtime sets + # `layers_to_capture = ()`, so `is_layer_capture()` is False everywhere + # and no hidden-state capture hook is owed (that is Eagle3's + # requirement, not this mode's) — and the core's assert that it is + # absent is a stronger guarantee than deleting the assert would be. + hidden = self.model( + attn_metadata=attn_metadata, + input_ids=input_ids, + position_ids=position_ids, + inputs_embeds=inputs_embeds, + lora_params=lora_params, + **kwargs, + ) + # `gather_ids` holds one row per context request (its last token) and + # `runtime_draft_len + 1` per generation request. `embedding` is the + # catalog's row-lookup entry (torch.nn.functional.embedding), used here + # for what it is — `hidden[gather_ids]`. + logits = self.logits_processor.forward( + embedding(spec_metadata.gather_ids, hidden), + self.lm_head, + attn_metadata, + True, + ) + return self.spec_worker( + input_ids=input_ids, + position_ids=position_ids, + hidden_states=hidden, + logits=logits, + attn_metadata=attn_metadata, + spec_metadata=spec_metadata, + draft_model=self.draft_model, + resource_manager=kwargs.get("resource_manager"), + ) + + def load_weights(self, weights, *args, **kwargs): + _weights.load(self, weights) + + def post_load_weights(self): + super().post_load_weights() + self.model.derive_after_load() diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/weights.py b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/weights.py new file mode 100644 index 000000000000..78e1ecdce657 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/weights.py @@ -0,0 +1,448 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Weight manifest and loader: deepseek-r1-0528-nvfp4 / sm_103 / dep4. + +MANIFEST is a data table: target parameter -> the checkpoint key(s) that fill +it, each with an optional destination index into the parameter and an +optional source transform. + +Every rank is handed the **whole** checkpoint dict, so any split is entirely +this table's — and under attention data parallelism there is only one: + +* **expert windows** — rank `r` lists only routed experts `[64r, 64r+64)` and + writes each at local index `e - offset`. The off-window experts' six + weight/scale keys are one of the two families a rank does not consume, and + the coverage assert names them exactly rather than being relaxed; +* **replicated** — everything else, and that is the point of the segment: the + topology divides the *requests*, not the weights. Both norms per layer, the + whole 128-head MLA attention block including the q-LoRA pair, the fp8 KV + scales, the dense MLPs of layers 0-2 and the shared-expert pair at full + width, the router and its bias, the embedding, the final norm, the shell's + `lm_head`, and **all 256 experts' per-tensor NVFP4 scalars** (6 fp32 values + per expert, kept whole so the shared-expert activation-scale assert in + `modeling.derive_after_load` still sees the max over every routed expert; + the window is sliced there). + +So this table has no per-rank row or column slices anywhere, which is exactly +what makes the `src`-transform column a tensor-parallel target needs +unnecessary here. + +**The second predicted non-load is layer 61, and how much of it is non-load +depends on the config.** The checkpoint ships a bf16 multi-token-prediction +module at layer index `num_hidden_layers` — its own 256 experts, +`embed_tokens`, `eh_proj`, two extra norms and a `shared_head.head`, 790 keys. + +* Under the target's identity config all 790 are an expected non-load, listed + rather than swept under a relaxed assert. +* Under a `configs/mtp*.yaml` variant the module is loaded, and this table + gains its rows: **212 keys per rank** are consumed (the whole front end and + attention block, the router, the shared expert, and this rank's 64-expert + window), leaving 578 — the 192 off-window experts' 576 weight tensors, which + join the first non-load family above, plus exactly two that stay non-load on + *every* rank. `embed_tokens.weight` and `shared_head.head.weight` are + `torch.equal` to `model.embed_tokens.weight` and `lm_head.weight`, so the + draft-model container points at the trunk's and saves 1.85 GB per rank. + +The module's expert stacks land in `fused_moe`'s layout — `[E, 2I, H]` with +the **up** rows first, `[E, H, I]` for FC2 — which needs a plain concatenation +and none of the interleave / 32-row shuffle / swizzle the trtllm-gen runner's +NVFP4 stacks need: `hf_quant_config.json` excludes `model.layers.61*` from +quantization wholesale, so every weight here is bf16. + +Storage is HF `[out, in]` row-major, so attention/router/embedding copies and +every NVFP4 weight-byte copy are layout-preserving. Three families need a +transform, all of them relayouts a kernel demands and no checkpoint stores: + +* **`kv_b_proj` row reorder.** The checkpoint interleaves per head + `[k_nope | v]` (row `h*(nope+v)+j`). The MLA context FMHA hard-codes V's + row stride as the full packed width `H*(nope+v)` and reads V's per-head + stride as `v_head_dim`, i.e. it addresses the `[.., H*nope:]` **column + block** of the projection output. So the rows are re-grouped once at load + into `[all heads' k_nope ; all heads' v]`; the two absorption operands + (`k_b [H, nope, C]`, `v_b [H, v, C]`) are then plain views of the halves. + +* **NVFP4 dense scales.** `nvfp4_gemm` consumes the 128x4-swizzled scale + order; the checkpoint stores row-major `[N, K/16]`. For the fused + `gate_up` linear the two halves' scale tensors are concatenated **before** + the single `block_scale_interleave` — interleaving them separately and + concatenating after produces a different byte order. + +* **NVFP4 expert stacks.** Per expert, per the MoE runner's contract: + concat `[up ; gate]` (up rows first), interleave (dest row `2i` = up `i`, + `2i+1` = gate `i`), 32-row block shuffle (src `4u+v` -> dest `8v+u`) of + the weight *and* scale bytes, then the 128x4 swizzle of the scales. + FC2 skips the interleave. Expert parallelism splits the stack whole, so + each expert's preparation is exactly the single-rank one — at this geometry + (H=7168, I=2048) no padding is needed anywhere. + +The per-tensor NVFP4 scalars (`input_scale`, `weight_scale_2`) are loaded +raw, one parameter per checkpoint key including the `up_proj` duplicates the +fused GEMM assumes equal to `gate_proj`'s; `modeling.derive_after_load` +asserts that equality and folds them into the kernels' `alpha` / `g` / +three-scalar forms. The per-layer `k_scale` / `v_scale` are loaded for the +same reason: the fp8 latent pool is only correct at a KV scaling factor of +1.0 and this target passes no scale tensors, so the value is checked there +rather than assumed here. + +Loading contract: `load(model, weights)` consumes the engine-provided +checkpoint dict, fills every declared parameter exactly once, and asserts +bidirectional coverage — every target parameter written, and every checkpoint +key either consumed or in the two predicted non-load sets above. +""" + +import torch + +SF_BLOCK = 16 # NVFP4: one e4m3 scale per 16 elements along K + + +def _block32_perm(rows: int, device) -> torch.Tensor: + """Gather index of the 32-row block shuffle: inside each aligned block of + 32 rows, source row `4u + v` (0<=u<8, 0<=v<4) lands at destination row + `8v + u`.""" + assert rows % 32 == 0, rows + src = torch.arange(32) + dst = (src % 4) * 8 + src // 4 + idx = torch.empty(32, dtype=torch.long) + idx[dst] = src + blocks = rows // 32 + return (idx.repeat(blocks) + torch.arange(blocks).repeat_interleave(32) * 32).to(device) + + +def _interleave_perm(rows: int, device) -> torch.Tensor: + """Gather index that interleaves the `[up | gate]` halves of a `2*I` row + stack into up0, gate0, up1, gate1, ...""" + p = torch.empty(rows, dtype=torch.long) + p[0::2] = torch.arange(0, rows // 2) + p[1::2] = torch.arange(rows // 2, rows) + return p.to(device) + + +def _fc1_perm(rows: int, device) -> torch.Tensor: + """Interleave then 32-row block shuffle, composed into one gather.""" + return _interleave_perm(rows, device)[_block32_perm(rows, device)] + + +def _swizzle(x: torch.Tensor) -> torch.Tensor: + """128x4 block-scale swizzle of a `[M, C]` uint8 scale matrix, returned + flat (`M % 128 == 0` and `C % 4 == 0` hold at this geometry, so the + result is exactly `M * C` bytes).""" + return torch.ops.trtllm.block_scale_interleave(x.contiguous()) + + +def _reorder_kv_b(t: torch.Tensor, core) -> torch.Tensor: + """`[H*(nope+v), C]` per-head-interleaved -> `[H*nope ; H*v]` blocks.""" + h, n, v, c = core.heads, core.nope, core.v_dim, core.kv_lora + t3 = t.reshape(h, n + v, c) + return torch.cat( + [t3[:, :n, :].reshape(h * n, c), t3[:, n:, :].reshape(h * v, c)], dim=0 + ).contiguous() + + +def _cat_rows(t: tuple[torch.Tensor, ...], core) -> torch.Tensor: + """Fused `gate_up` weight bytes: gate rows first (the half order + `flashinfer_silu_and_mul` reads).""" + return torch.cat(t, dim=0).contiguous() + + +def _cat_rows_swizzle(t: tuple[torch.Tensor, ...], core) -> torch.Tensor: + """Fused `gate_up` scale bytes: concatenate, then one swizzle.""" + return _swizzle(torch.cat([x.view(torch.uint8) for x in t], dim=0)) + + +def _swizzle_only(t: torch.Tensor, core) -> torch.Tensor: + """`down_proj` scale bytes: swizzle in place of the row-major order.""" + return _swizzle(t.view(torch.uint8)) + + +def _expert_fc1_w(t: tuple[torch.Tensor, ...], core) -> torch.Tensor: + """(up, gate) packed nibbles -> one expert's FC1 operand.""" + up, gate = t + w = torch.cat([up, gate], dim=0) + return torch.index_select(w, 0, _fc1_perm(w.shape[0], w.device)).contiguous() + + +def _expert_fc1_s(t: tuple[torch.Tensor, ...], core) -> torch.Tensor: + """(up, gate) e4m3 scales -> one expert's FC1 scale operand.""" + up, gate = t + s = torch.cat([up.view(torch.uint8), gate.view(torch.uint8)], dim=0) + rows, cols = s.shape + s = torch.index_select(s, 0, _fc1_perm(rows, s.device)) + return _swizzle(s).reshape(rows, cols) + + +def _expert_fc2_w(t: torch.Tensor, core) -> torch.Tensor: + """`down_proj` packed nibbles -> one expert's FC2 operand (no interleave).""" + return torch.index_select(t, 0, _block32_perm(t.shape[0], t.device)).contiguous() + + +def _expert_fc2_s(t: torch.Tensor, core) -> torch.Tensor: + """`down_proj` e4m3 scales -> one expert's FC2 scale operand.""" + s = t.view(torch.uint8) + rows, cols = s.shape + s = torch.index_select(s, 0, _block32_perm(rows, s.device)) + return _swizzle(s).reshape(rows, cols) + + +def _mtp_fc1(t: tuple[torch.Tensor, ...], core) -> torch.Tensor: + """One MTP expert's FC1 operand: `[up ; gate]` rows, up first — the half + order `fused_moe` reads, the opposite of the dense gate_up linear's.""" + return torch.cat(t, dim=0).contiguous() + + +def _mtp_rows(rows: dict, core) -> None: + """The MTP module's manifest rows, at layer index `num_hidden_layers`. + Everything here is bf16 and stored HF `[out, in]`, so the only transforms + are the same `kv_b_proj` row regroup the trunk's attention needs and two + plain row concatenations.""" + p = f"model.layers.{core.num_layers}" + rows["mtp_enorm"] = [(f"{p}.enorm.weight", None, None)] + rows["mtp_hnorm"] = [(f"{p}.hnorm.weight", None, None)] + rows["mtp_eh"] = [(f"{p}.eh_proj.weight", None, None)] + rows["mtp_norm1"] = [(f"{p}.input_layernorm.weight", None, None)] + rows["mtp_qa"] = [(f"{p}.self_attn.q_a_proj.weight", None, None)] + rows["mtp_q_norm"] = [(f"{p}.self_attn.q_a_layernorm.weight", None, None)] + rows["mtp_qb"] = [(f"{p}.self_attn.q_b_proj.weight", None, None)] + rows["mtp_kva"] = [(f"{p}.self_attn.kv_a_proj_with_mqa.weight", None, None)] + rows["mtp_kv_norm"] = [(f"{p}.self_attn.kv_a_layernorm.weight", None, None)] + rows["mtp_kvb"] = [(f"{p}.self_attn.kv_b_proj.weight", None, _reorder_kv_b)] + rows["mtp_o"] = [(f"{p}.self_attn.o_proj.weight", None, None)] + rows["mtp_k_scale"] = [(f"{p}.self_attn.k_proj.k_scale", (0,), None)] + rows["mtp_v_scale"] = [(f"{p}.self_attn.v_proj.v_scale", (0,), None)] + rows["mtp_norm2"] = [(f"{p}.post_attention_layernorm.weight", None, None)] + rows["mtp_head_norm"] = [(f"{p}.shared_head.norm.weight", None, None)] + rows["mtp_router"] = [(f"{p}.mlp.gate.weight", None, None)] + rows["mtp_router_bias"] = [(f"{p}.mlp.gate.e_score_correction_bias", None, None)] + s = f"{p}.mlp.shared_experts" + rows["mtp_sh_gu"] = [((f"{s}.gate_proj.weight", f"{s}.up_proj.weight"), None, _cat_rows)] + rows["mtp_sh_dn"] = [(f"{s}.down_proj.weight", None, None)] + fc1, fc2 = [], [] + for local in range(core.local_experts): + q = f"{p}.mlp.experts.{core.expert_offset + local}" + fc1.append(((f"{q}.up_proj.weight", f"{q}.gate_proj.weight"), (local,), _mtp_fc1)) + fc2.append((f"{q}.down_proj.weight", (local,), None)) + rows["mtp_fc1"] = fc1 + rows["mtp_fc2"] = fc2 + + +def _dense_mlp_rows(rows: dict, key: str, prefix: str) -> None: + """Manifest rows shared by the dense MLPs of layers 0-2 and the shared + experts: one fused gate_up NVFP4 linear plus one down NVFP4 linear. Both + are replicated and run over this rank's own tokens.""" + rows[f"{key}_gu_w"] = [ + ((f"{prefix}.gate_proj.weight", f"{prefix}.up_proj.weight"), None, _cat_rows) + ] + rows[f"{key}_gu_s"] = [ + ( + (f"{prefix}.gate_proj.weight_scale", f"{prefix}.up_proj.weight_scale"), + None, + _cat_rows_swizzle, + ) + ] + rows[f"{key}_dn_w"] = [(f"{prefix}.down_proj.weight", None, None)] + rows[f"{key}_dn_s"] = [(f"{prefix}.down_proj.weight_scale", None, _swizzle_only)] + rows[f"{key}_isc1"] = [(f"{prefix}.gate_proj.input_scale", (0,), None)] + rows[f"{key}_isc1_up"] = [(f"{prefix}.up_proj.input_scale", (0,), None)] + rows[f"{key}_ws2_1"] = [(f"{prefix}.gate_proj.weight_scale_2", (0,), None)] + rows[f"{key}_ws2_1_up"] = [(f"{prefix}.up_proj.weight_scale_2", (0,), None)] + rows[f"{key}_isc2"] = [(f"{prefix}.down_proj.input_scale", (0,), None)] + rows[f"{key}_ws2_2"] = [(f"{prefix}.down_proj.weight_scale_2", (0,), None)] + + +def _materialize(entry) -> torch.Tensor: + """Realize one checkpoint entry: `[:]` on a lazy safetensors slice, but a + 0-dim tensor (every NVFP4 per-tensor scalar and every KV scale) rejects + that index.""" + return entry if getattr(entry, "ndim", 1) == 0 else entry[:] + + +def _manifest(core) -> dict: + """target param key -> list of (ckpt key or key tuple, index into the + param | None, source transform | None).""" + rows: dict = {} + for i in range(core.num_layers): + p = f"model.layers.{i}" + rows[f"l{i}_norm1"] = [(f"{p}.input_layernorm.weight", None, None)] + rows[f"l{i}_qa"] = [(f"{p}.self_attn.q_a_proj.weight", None, None)] + rows[f"l{i}_q_norm"] = [(f"{p}.self_attn.q_a_layernorm.weight", None, None)] + rows[f"l{i}_qb"] = [(f"{p}.self_attn.q_b_proj.weight", None, None)] + rows[f"l{i}_kva"] = [(f"{p}.self_attn.kv_a_proj_with_mqa.weight", None, None)] + rows[f"l{i}_kv_norm"] = [(f"{p}.self_attn.kv_a_layernorm.weight", None, None)] + rows[f"l{i}_kvb"] = [(f"{p}.self_attn.kv_b_proj.weight", None, _reorder_kv_b)] + rows[f"l{i}_o"] = [(f"{p}.self_attn.o_proj.weight", None, None)] + # The fp8 KV-cache scales sit under the k/v projections the MLA + # checkpoint does not otherwise have. + rows[f"l{i}_k_scale"] = [(f"{p}.self_attn.k_proj.k_scale", (0,), None)] + rows[f"l{i}_v_scale"] = [(f"{p}.self_attn.v_proj.v_scale", (0,), None)] + rows[f"l{i}_norm2"] = [(f"{p}.post_attention_layernorm.weight", None, None)] + if i < core.dense_layers: + _dense_mlp_rows(rows, f"l{i}_mlp", f"{p}.mlp") + continue + _dense_mlp_rows(rows, f"l{i}_mlp", f"{p}.mlp.shared_experts") + rows[f"l{i}_router"] = [(f"{p}.mlp.gate.weight", None, None)] + rows[f"l{i}_router_bias"] = [(f"{p}.mlp.gate.e_score_correction_bias", None, None)] + fc1_w, fc1_s, fc2_w, fc2_s = [], [], [], [] + for local in range(core.local_experts): + q = f"{p}.mlp.experts.{core.expert_offset + local}" + fc1_w.append( + ( + (f"{q}.up_proj.weight", f"{q}.gate_proj.weight"), + (local,), + _expert_fc1_w, + ) + ) + fc1_s.append( + ( + (f"{q}.up_proj.weight_scale", f"{q}.gate_proj.weight_scale"), + (local,), + _expert_fc1_s, + ) + ) + fc2_w.append((f"{q}.down_proj.weight", (local,), _expert_fc2_w)) + fc2_s.append((f"{q}.down_proj.weight_scale", (local,), _expert_fc2_s)) + rows[f"l{i}_fc1_w"] = fc1_w + rows[f"l{i}_fc1_s"] = fc1_s + rows[f"l{i}_fc2_w"] = fc2_w + rows[f"l{i}_fc2_s"] = fc2_s + # The per-tensor scalars are replicated over the whole routing space, + # not windowed: derive_after_load asserts the shared-expert + # activation scale against the max over all 256 routed experts. + isc1, isc1_up, ws2_1, ws2_1_up, isc2, ws2_2 = [], [], [], [], [], [] + for e in range(core.num_experts): + q = f"{p}.mlp.experts.{e}" + isc1.append((f"{q}.gate_proj.input_scale", (e,), None)) + isc1_up.append((f"{q}.up_proj.input_scale", (e,), None)) + ws2_1.append((f"{q}.gate_proj.weight_scale_2", (e,), None)) + ws2_1_up.append((f"{q}.up_proj.weight_scale_2", (e,), None)) + isc2.append((f"{q}.down_proj.input_scale", (e,), None)) + ws2_2.append((f"{q}.down_proj.weight_scale_2", (e,), None)) + rows[f"l{i}_e_isc1"] = isc1 + rows[f"l{i}_e_isc1_up"] = isc1_up + rows[f"l{i}_e_ws2_1"] = ws2_1 + rows[f"l{i}_e_ws2_1_up"] = ws2_1_up + rows[f"l{i}_e_isc2"] = isc2 + rows[f"l{i}_e_ws2_2"] = ws2_2 + rows["final_norm"] = [("model.norm.weight", None, None)] + rows["embed"] = [("model.embed_tokens.weight", None, None)] + if core.mtp_enabled: + _mtp_rows(rows, core) + return rows + + +def _offwindow_expert_keys(core) -> set: + """The first predicted non-load: the weight/scale tensors of every routed + expert outside this rank's EP window. Their per-tensor scalars are consumed + on every rank, so nothing else of theirs is left over. + + The trunk's MoE layers store six such tensors per expert (NVFP4 data plus + block scales); the MTP module's are bf16, so its off-window experts leave + three each and no `weight_scale` at all.""" + lo, hi = core.expert_offset, core.expert_offset + core.local_experts + layers = list(range(core.dense_layers, core.num_layers)) + if core.mtp_enabled: + layers += list(range(core.num_layers, core.num_layers + core.mtp_layers)) + keys = set() + for i in layers: + quantized = i < core.num_layers + for e in range(core.num_experts): + if lo <= e < hi: + continue + q = f"model.layers.{i}.mlp.experts.{e}" + for proj in ("gate_proj", "up_proj", "down_proj"): + keys.add(f"{q}.{proj}.weight") + if quantized: + keys.add(f"{q}.{proj}.weight_scale") + return keys + + +def _mtp_keys(core, weights) -> set: + """The second predicted non-load: what the multi-token-prediction layers + the checkpoint ships past `num_hidden_layers` leave behind. Read off the + checkpoint by layer index rather than enumerated, so a key belonging to a + real layer can never land here. + + With MTP off that is those layers whole — on this checkpoint layer 61's 790 + keys: its own 256 experts, embedding, `eh_proj`, norms and output head. + With MTP on the manifest consumes all but two per rank: `embed_tokens` and + `shared_head.head` are bitwise copies of the trunk's embedding and + `lm_head`, which the draft-model container points at instead of loading + them twice. (The off-window experts are left over too, but they belong to + the family above and are named there.)""" + prefixes = tuple( + f"model.layers.{i}." for i in range(core.num_layers, core.num_layers + core.mtp_layers) + ) + if not prefixes: + return set() + keys = {k for k in weights if k.startswith(prefixes)} + if not core.mtp_enabled: + return keys + aliased = (".embed_tokens.weight", ".shared_head.head.weight") + return {k for k in keys if k.endswith(aliased)} + + +def load(model, weights) -> None: + core = model.model + manifest = _manifest(core) + consumed: set = set() + + def fill(param: torch.nn.Parameter, ckpt_key, index, transform) -> None: + keys = ckpt_key if isinstance(ckpt_key, tuple) else (ckpt_key,) + for key in keys: + assert key in weights, f"checkpoint key missing: {key}" + dst = param.data if index is None else param.data[index] + # Materialize the checkpoint entries where the destination lives: the + # expert relayouts are row gathers over ~15 MB per expert and the + # concatenations allocate scratch of the same order. The per-tensor + # scalars are stored 0-dim, which `[:]` rejects. + src_all = tuple( + _materialize(weights[key]).to(dst.device, non_blocking=True) for key in keys + ) + if transform is not None: + src = transform(src_all if isinstance(ckpt_key, tuple) else src_all[0], core) + else: + assert len(src_all) == 1, "a multi-key row needs a transform" + src = src_all[0] + assert dst.shape == src.shape, (ckpt_key, tuple(dst.shape), tuple(src.shape)) + assert src.dtype == dst.dtype, (ckpt_key, src.dtype, dst.dtype) + dst.copy_(src, non_blocking=True) + consumed.update(keys) + + assert set(manifest.keys()) == set(core.w.keys()), ( + "manifest/parameter drift", + set(manifest.keys()) ^ set(core.w.keys()), + ) + for param_key, sources in manifest.items(): + for ckpt_key, index, transform in sources: + fill(core.w[param_key], ckpt_key, index, transform) + # The expert relayouts allocate per-expert scratch 64 times per + # layer; release it before the next parameter so peak load memory + # stays one layer deep. + if param_key.endswith("_fc2_s"): + torch.cuda.empty_cache() + + # Shell-registered exception: the base class owns lm_head (untied), and + # under attention DP it builds it **whole** — `[vocab, hidden]` on every + # rank. That is the opposite of a tensor-parallel target, where the same + # shell builds a vocab-parallel `[vocab/tp, hidden]`, and it is the only + # consistent choice here: each rank holds different tokens, so a vocab + # shard could not be completed by a collective over rows no other rank + # computed. Asserted rather than adapted — a shape change is a different + # logits path. + assert tuple(model.lm_head.weight.shape) == (core.vocab, core.hidden), ( + f"attention DP replicates lm_head; the shell built " + f"{tuple(model.lm_head.weight.shape)} instead of {(core.vocab, core.hidden)}" + ) + fill(model.lm_head.weight, "lm_head.weight", None, None) + # Parameter-side coverage, the other direction: the shell's construction + # is topology-dependent, so name every registered parameter outside the + # target's own ParameterDict instead of assuming lm_head is the only one. + shell = {n for n, _ in model.named_parameters() if not n.startswith("model.w.")} + assert shell == {"lm_head.weight"}, f"unfed shell parameters: {sorted(shell)}" + + torch.cuda.synchronize() + leftover = set(weights.keys()) - consumed + expected = _offwindow_expert_keys(core) | _mtp_keys(core, weights) + assert leftover == expected, ( + "checkpoint coverage: leftover keys are not exactly this rank's " + f"off-window experts plus the MTP layers; unexpected " + f"{sorted(leftover - expected)[:8]}, missed {sorted(expected - leftover)[:8]}" + ) diff --git a/tensorrt_llm/_torch/staircase/models/gpt_oss/__init__.py b/tensorrt_llm/_torch/staircase/models/gpt_oss/__init__.py new file mode 100644 index 000000000000..6b4453f48f69 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/gpt_oss/__init__.py @@ -0,0 +1,3 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""GptOssForCausalLM targets.""" diff --git a/tensorrt_llm/_torch/staircase/models/gpt_oss/routing.py b/tensorrt_llm/_torch/staircase/models/gpt_oss/routing.py new file mode 100644 index 000000000000..8a1407d0af19 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/gpt_oss/routing.py @@ -0,0 +1,59 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Where a GptOssForCausalLM config lands. Read this file and you know. + +One forward-reading decision tree per architecture family: the criteria are +evaluated in the order a reader would ask them, and every branch that does not +end in a target returns None (in ``auto`` the engine then uses the built-in +GptOss implementation; in ``require`` it raises, quoting the trace below). +""" + +from __future__ import annotations + +from typing import Optional + +from ..._router_index import NULL_TRACE, StaircaseContext, Trace + +# The one GPU architecture these targets are written for. sm is part of a +# target's identity, not a knob: a different SM is a different target. +_SM = (10, 3) + +# Config-shape fingerprint -> checkpoint identity. Sniffing the shape is the +# upstream idiom (``is_mla``, ``is_nemotron_hybrid`` do the same). It buys +# automatic routing at a stated cost: a *fine-tune* of this checkpoint has the +# same shape and is routed here silently. See TARGET.md -- the gate record is +# what pins the identity, and it says which checkpoint it was measured on. +# +# (num_hidden_layers, hidden_size, num_local_experts). Layer count alone +# separates 120b from 20b, but the expert count is what makes the MoE operand +# geometry this target declares correct, so it is part of the fingerprint. +_CHECKPOINTS = { + (36, 2880, 128): "gpt_oss_120b", +} + +_TARGETS = { + ("gpt_oss_120b", "tp1"): "StaircaseGptOss120bSm103Tp1", +} + +# Synthetic architecture name -> the module whose import registers it. +TARGET_MODULES = { + "StaircaseGptOss120bSm103Tp1": "models.gpt_oss.targets.gpt_oss_120b.sm_103.tp1.modeling", +} + + +def route(ctx: StaircaseContext, trace: Trace = NULL_TRACE) -> Optional[str]: + c, m = ctx.pretrained_config, ctx.mapping + + if not trace.check("sm", ctx.sm, ctx.sm == _SM): + return None + + shape = (c.num_hidden_layers, c.hidden_size, c.num_local_experts) + ckpt = trace.resolve("shape", shape, _CHECKPOINTS.get(shape)) + if ckpt is None: + return None + + parallel = trace.resolve("parallel", f"ws={m.world_size}", "tp1" if m.world_size == 1 else None) + if parallel is None: + return None + + return _TARGETS.get((ckpt, parallel)) diff --git a/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/__init__.py b/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/__init__.py new file mode 100644 index 000000000000..d7eb0af07fc9 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/__init__.py @@ -0,0 +1,3 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Targets, keyed by the // identity path.""" diff --git a/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/__init__.py b/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/__init__.py new file mode 100644 index 000000000000..b2a58d2d0fae --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/__init__.py @@ -0,0 +1,3 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""gpt-oss-120b targets.""" diff --git a/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/__init__.py b/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/__init__.py new file mode 100644 index 000000000000..4ea8999b15d3 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/__init__.py @@ -0,0 +1,3 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""gpt-oss-120b on sm_103 (GB300).""" diff --git a/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/TARGET.md b/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/TARGET.md new file mode 100644 index 000000000000..ebbf5a0caf70 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/TARGET.md @@ -0,0 +1,480 @@ +# Target: gpt-oss-120b / sm_103 / tp1 + +## Identity + +| | | +|---|---| +| Checkpoint | gpt-oss-120b (HF safetensors; bf16 attention/router/embedding/lm_head, **MXFP4 experts** — E2M1 blocks + per-32 E8M0 scales; untied embeddings; sparse MoE on every layer: 128 experts, top-4, renormalized; **attention sinks**; alternating sliding-window / full attention) | +| GPU arch | sm_103 (GB300) | +| Parallel | tp1 | +| Registered class | `StaircaseGptOss120bSm103Tp1` — a synthetic architecture name no checkpoint declares. `models/gpt_oss/routing.py` rewrites `GptOssForCausalLM` into it when the config, SM and topology all match; the checkpoint is read unpatched. Per-target names mean one process can hold every target at once | + +> **NO GATE RECORD HOLDS FOR THIS TARGET.** Every result below was +> measured on **sm_100 (B200)** through the pre-move standalone harness. This +> target is **sm_103 (GB300)**, and certification is per architecture. The +> numbers are kept as provenance — they are true records of what the same +> modeling code did on another device — but this target is **ungated** until +> the boot and gsm8k runs in *Verification* are repeated on GB300 and their +> results replace those rows. Read every "passed" below as "passed, on +> sm_100, before the move". + +Checkpoint sha256 — the checkpoint directory passed to `--model` should +resolve to files with these digests. **Routing does not check them**: it +fingerprints the config's shape, so a fine-tune of this checkpoint routes +here silently. That is the deliberate trade in `models/gpt_oss/routing.py`, +and it changes what a gate record means — not "this target passed" but +"this modeling code passed *on the checkpoint with these digests*". Run it +on another one and the result is ungated (recorded at: +`umbriel-b200-027:/home/scratch.trt_llm_data/llm-models/gpt_oss/gpt-oss-120b`): + +``` +695218884684c611fe08a74751ee443f971e9bd9bc062edba822da3fe45969b7 model-00000-of-00014.safetensors +a881aa5f561b26a22b14a8262aa61849ace349ffd73d74769e030ac90a1fcf8a model-00001-of-00014.safetensors +022478dd04398c5bdb545a5be0a6437ecc2eb53d1dbd29edafcfff4b3ddf0a41 model-00002-of-00014.safetensors +47aee9e7b9d5bedb215042c01ccededd9bd9c30b0dddea862dc2506b9d6c74de model-00003-of-00014.safetensors +f6c2752acda607b1d5ca52df9e75c1b9b2761e6875ff10c9bd6ddac473c0262e model-00004-of-00014.safetensors +0c8dd401544c31cb93b8459eee7da20ea2a07626a59455d7d92b85257df9b46c model-00005-of-00014.safetensors +28d839f2e027985a8b14e45f2323798862eddb7770ee9800ea6b7c803abee489 model-00006-of-00014.safetensors +c8958c5f183c04f6ea959cfd90562b5128124154b2bbf979b8a22b9405b30ed8 model-00007-of-00014.safetensors +bf1f2a88868ffc37d520dcf77d26f0e823710b5e682d473ff10f6974fa3b7517 model-00008-of-00014.safetensors +f72d34a4004241b45c332b61f8ffa124e9a913bc1ab442b66e717d3e94e741ce model-00009-of-00014.safetensors +f48c867c2cb0a44bfc2f8768cb98e4aec9a350946fceacfebdcad5d32ad4a471 model-00010-of-00014.safetensors +a06851b2cfd35f48722f823bc1ab8f7bcb4a878a5b8e975f4d3544f230454eeb model-00011-of-00014.safetensors +3af33667c307e20ae2a7648ea52653de46dd0171601ec5c696e47a2f5d5bf1e4 model-00012-of-00014.safetensors +bcbcb74b043e071d1e05471d500d74dcf661175e00878ed302ccdf1801a75aef model-00013-of-00014.safetensors +54b1be1609696c307cc5ca117b1fa54feaddebffa04e9c2db117652a01964230 model-00014-of-00014.safetensors +ede2655fdc05008561983b6e0829c600727c28d591e071077377059f03a6c00e model.safetensors.index.json +0614fe83cadab421296e664e1f48f4261fa8fef6e03e63bb75c20f38e37d07d3 tokenizer.json +9279e942392b742d633c7adbb89ebe002c98399db8926a7af5125c726f404070 tokenizer_config.json +dd5e191d20c12d2fee1da5bae14ca1db0f5f4215300af691f23cdee97120a293 special_tokens_map.json +f8d9255777615591a7cc1a7c932f5a69e181128902295e1b81221d20d983cac7 chat_template.jinja +7bfd294f3e29b53db1e126d5cec050a12dc27adad4445fb5eab540e1cad74ea1 chat_template.json +199566674b96510c3b9a1141b494223a86a3ff83097e2c2f259c4d94fafc5847 generation_config.json +``` + +Two link-set notes specific to this checkpoint. Its chat template lives in +`chat_template.jinja` (`tokenizer_config.json` carries none), and the +accuracy gate's protocol applies that template — so the template files and +`special_tokens_map.json` are linked alongside `tokenizer*`. And its +`config.json` declares **no dtype at all**, which used to be patched around +with a target-owned `model_dir/config.json` declaring `dtype: bfloat16`. + +That stub is gone — the checkpoint is now read exactly as published — so the +divergence it papered over is live code. `DecoderModelForCausalLM` sizes +`lm_head` from `pretrained_config.torch_dtype`, which is `None` here, and +would materialize it in the torch default fp32 while every other tensor is +bf16; the failure surfaces two layers from its cause. The shell fills that +gap explicitly before `super().__init__`, adopting the dtype the engine +already resolved (`ModelConfig.torch_dtype`, bf16) and only when the +checkpoint declares none. Every safetensors tensor outside the MXFP4 expert +blocks is bf16, so that is what the checkpoint is. + +## Version + +| | | +|---|---| +| tensorrt_llm | in-tree — the target moves with the trunk, so there is no version to pin and none is asserted. What *is* asserted at construction is the SM version (`_SM = (10, 3)`), which the pin used to stand in for. The gate records below name the commit they were taken at | +| torch | 2.11.0+cu130 | +| transformers | 5.5.4 (the config surface the engine hands the target; see the rope note in `modeling.py`) | +| Attention metadata fact source | `TrtllmAttentionMetadata` (TRTLLM backend) | + +## Vocabulary + +Forward: `flashinfer_rmsnorm`, `flashinfer_fused_add_rmsnorm`, +`cublas_mm` (fused bias — qkv, o, router all carry one), +`fused_qk_norm_rope` (`is_qk_norm=False`: YaRN RoPE only), +`thop_attention` (per-layer `attention_sinks` and `attention_window_size`), +`mxfp8_quantize`, `mxe4m3_mxe2m1_block_scale_moe_runner`, +`torch/embedding`, `torch/empty`, `torch/reshape`. + +The MoE call is the whole expert block — routing, both grouped GEMMs, the +clamped GLU, the MXFP8 requantization between the GEMMs and the combine — +so no `activation/*` and no `moe/*routing*` entry appears: this checkpoint +has no dense MLP, and the router bias rides the router GEMM's fused-bias +epilogue because the MoE op silently ignores `routing_bias` at +`routing_method_type=1`. + +**Changed by the perf campaign (`iter1-mxfp8-moe`, see *Performance*).** +The assembly shipped the W4A16 member of this kernel family +(`bf16_mxe2m1_block_scale_moe_runner`) with `torch/pad` widening hidden +2880 → 3072 in front of it; the target now runs the W4A8 member over +MXFP8 activations, and `mxfp8_quantize` does that widening inside itself, +so the `torch/pad` call is gone. Weight preparation is unchanged — the two +ops consume the identical prepared expert stack. + +Weight loading fills parameters via the manifest loop (`torch/copy_` +semantics); `.t()` views are cublas_mm's column-major consumption form, +derived once post-load. The expert operands are the exception: the +manifest's source transforms rebuild the checkpoint's MXFP4 blocks and +E8M0 scales into the MoE op's kernel-ready layout (pad → `[up ; gate]` +concat → row interleave → 32-row block shuffle → 128x4 scale swizzle), +promote both expert biases and the attention sinks from bf16 to fp32, and +run on device one layer at a time. + +Audit is mechanical: grep the forward's calls against `catalog/index.yaml`. + +## Verification + +### Required on sm_103 — not yet run + +| Gate | Command | Result | +|---|---|---| +| boot | `TRTLLM_STAIRCASE=require python examples/llm-api/quickstart_advanced.py --model_dir --max_tokens 16 --prompt "The capital of France is" "The chemical symbol for gold is" "1, 2, 3, 4, 5,"` | **passed, 10/10** greedy keyword asserts, 2026-09-09, GPU 0 of nvl72d001-T18 (NVIDIA GB300, 284208 MiB, sm_103), trtllm 1.3.0rc26. 65.70 GiB of weights loaded; decode CUDA graphs captured at batch sizes 1-32, 64, 128; engine boot to last token 3m29s. The run had the switch set to `require`, so the built-in GptOss implementation could not have been substituted. Measured with a per-target `smoke.py` that asserted a keyword in each of ten greedy continuations; that file was removed in favour of the generic script above, which boots the same way and prints the same continuations for a human to read | + +**Execution is verified on sm_103; the accuracy gate is not yet.** Everything up to the +weight load is driven by the checkpoint's config alone, and that part +was exercised on a GB300 against the config *as published* (no +target-owned stub): routing resolved this target, the module imported, +`StaircaseCore.__init__` passed every geometry, topology and dtype +assert, and 543 parameters declared, 36 layers, hidden 2880. `lm_head.weight` +came out **bfloat16**, which is the specific thing removing the stub +put at risk -- the shell sizes it from the pretrained dtype, and a +regression there materializes fp32 two layers from its cause. + +The weight path is verified too, by the boot run above: the manifest +load fills every declared parameter (its coverage asserts are part of +the gate), the post-load derivations run, and the forward produces +coherent greedy continuations through both the prefill and the +CUDA-graph decode path. + +The checkpoint it ran on is the one this file records. Every digest +above was re-verified after download -- the six small files by hand and +the fifteen safetensors shards by git-lfs, whose object id *is* the +sha256 -- so this gate record and the sm_100 records below were measured +on byte-identical weights. + +| gsm8k full | `TRTLLM_STAIRCASE=require trtllm-eval --model gsm8k --output_path --apply_chat_template --fewshot_as_multiturn --max_output_length 8192`, or in CI as `accuracy/test_staircase.py::TestStaircaseGptOss120bSm103Tp1::test_gsm8k` | **passed, 90.6748** (`exact_match,flexible-extract`, +-0.8010, full 1319 questions) against threshold **85.5989** (anchor `openai/gpt-oss-120b` = 90.5989, tol 5.0) -- pass by 5.08 points. 2026-09-09, GPU 0 of nvl72d173-T18 (GB300), trtllm 1.3.0rc26, 2m54s wall. `strict-match` on the same run: 27.0660 | + +`TRTLLM_STAIRCASE=require` is what makes the second command a gate at all: under +`auto` a configuration that missed this target would measure the built-in +GptOss implementation and report it as this target's score. + + +#### Two notes on reading the gsm8k number + +**`--check_accuracy` is not used, and the filter is read by hand.** The gate +filter is `exact_match,flexible-extract`, and `trtllm-eval` exposes no CLI flag +for `scores_filter` -- it is a keyword of the evaluator's `evaluate()`. Left +unset the evaluator *averages* the filters, which for this checkpoint mixes +90.6748 with a `strict-match` of 27.0660 and reports 58.87. That average is +meaningless here: this is a reasoning model, its answer never arrives in the +strict `#### N` form, so the strict filter scores it near the floor. Read the +flexible-extract row of the logged table, which is also saved to +`--output_path`. + +**The anchor in `references/accuracy.yaml` was deliberately not written back.** +The rule in that file is to write a target's first passing score back -- but it +also says to skip the write-back when doing so would confuse what the anchor +means. It would here: the anchor is keyed by *checkpoint* (`openai/gpt-oss-120b`) +and currently holds an sm_100 measurement of the W4A16 assembly, deliberately +left high so the gate stays the stricter of the two. This measurement is a third +thing -- the W4A8 forward on sm_103 -- and folding it into a checkpoint-keyed +anchor would make that anchor architecture-dependent without saying so. The +number lives here, where gate records are per target and therefore per +architecture. + +For the record, the three measurements of this checkpoint through this harness: +sm_100 W4A16 **90.5989** (the anchor), sm_100 W4A8 **89.16**, sm_103 W4A8 +**90.6748**. The last is +1.51 over the sm_100 W4A8 record, which is ~1.9 sigma +on this run's +-0.80 stderr and is **not** claimed as an improvement -- the MoE +FC1 epilogue genuinely uses a different block-scale recipe on sm_103 (see +`catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md`), so a difference of this +size has a plausible mechanism, but separating it from session variance would +need repeated runs on both architectures. + +### Prior record — sm_100 (B200), pre-move harness, does not gate this target + +Both gates passed under trtllm defaults (block reuse and CUDA graphs on, +`llm_args.yaml` empty, single-process worker via `scripts/env.sh`), +2026-07-26/27, GPU 0 of umbriel-b200-027 (NVIDIA B200, driver 595.58.03): + +| Gate | Result | +|---|---| +| smoke — `uv run targets/gpt-oss-120b/sm_100/tp1/smoke.py` | **passed**, 10/10 greedy keyword asserts (keywords frozen against continuations observed on this model; the arithmetic case is Q/A-framed because a bare `2 + 2 =` is genuinely ambiguous here) | +| gsm8k full — `uv run bench/accuracy.py --target targets/gpt-oss-120b/sm_100/tp1` | **passed**, `measured 89.16 >= 85.6`, exit 0 (1319 questions, 5-shot, chat template + few-shot-as-multiturn + 8192 output tokens, filter `exact_match,flexible-extract`) | + +**The row above is the shipped forward** — the W4A8 MoE the perf campaign +left in place. The gate it clears is the written-back anchor +`openai/gpt-oss-120b` = 90.5989 (`source: trtllm-eval`, tol 5.0 ⇒ +threshold 85.6): **pass by 3.56 points**. + +The anchor itself was measured on the *assembly's* W4A16 forward, and is +deliberately not re-written down to the W4A8 number — leaving it high +keeps the gate the stricter of the two, and it remains a real measurement +of this checkpoint through this harness. The W4A16 record it came from, +kept because the anchor derives from it: `measured 90.60 >= 85.3` against +`reference: openai/gpt-oss-120b = 90.3 (trtllm)`, exit 0, same protocol — +reproduced bit-identically by two independent invocations, the +assembler's and the orchestrator's verification re-run. + +Measured GSM8K on that W4A16 forward, full 1319 questions (2026-07-27): + +| filter | score | +|---|---| +| **flexible-extract (gated)** | **90.5989 ± 0.8039** | +| strict-match (not gated) | 25.4738 ± 1.2002 | + +Reference `openai/gpt-oss-120b` = 90.3 (`source: trtllm`), tol 5.0 ⇒ +threshold 85.3: **pass by 5.30 points**, and 0.30 *above* the anchor +itself. Stock trtllm on this checkpoint under the identical protocol +measured 90.2199 here during the onboard's anchor phase. The strict-match +filter is not the gated metric and is not comparable: it scores the +`#### N` surface form a harmony-format model never emits. + +**Rerun noise measured on this target, protocol-specific.** The same build +under the identical protocol scored 90.2199 in a first run and 90.5989 in +the gating run — 1190 vs 1195 of 1319, a 0.38-point spread with no code +change between them (strict-match swung further, 22.59 → 25.47). GSM8K +with 8192-token reasoning generations is therefore noisier than the +~0.1-0.3 the repo's guidance cites for completion MMLU: treat anything +under ~0.5 points on this gate as unresolved. + +Engine facts observed on those runs: model init 14.1-14.5 s (63 GB of +checkpoint, including the on-device expert relayout); one KV pool sized +for full attention (99.27 GiB, 1,445,664 tokens, `tokens_per_block` 32, +`window size=131072`), `host_kv_cache_pool_mapping` `[36, 2]` with +identity rows `[0, l]` — the runtime is not told about the 128-token +sliding layers, which is exactly the certified single-pool route, and the +alternation therefore saves no memory. Both FMHA kernel families appear at +warmup (`...H64PagedKvDense...` for the full layers, +`...H64PagedKvSlidingOrChunkedCausal...` for the sliding ones), which is +the per-layer window taking effect. Under trtllm defaults +`cache_reuse=True`, so the engine prepares `use_paged_context_fmha=True` +and the target passes it through as documented; both gates above ran with it +True. That value was **not certified** when this target merged — the +contract's certified column said `False`, and this line was the only +record of the mismatch anywhere. It was closed on 2026-07-28 by a +certification extension, which also established that the target's +behaviour was the correct one: at `False` a context call with a cached +prefix returns normally and lands 51x outside the tolerance band, while +with nothing cached the flag is bitwise inert. The paged read does add a +caller obligation — the pages in range must be valid and distinct — which +this target satisfies by passing the engine's own offsets through. + +Pre-gate check of the highest-risk axis, run before the first engine boot +(scratch script, not a repo product): layer 0's real MXFP4 expert tensors +prepared by `weights.py` and fed to the catalog MoE call land **1.12 bf16 +ulp** (of the token row's largest magnitude) from a pure-torch HF +reference — dequantized with the transformers mxfp4 semantics, gate/up +split as `[..., ::2]` / `[..., 1::2]`, clamped GLU, top-4 renormalized +routing — while the gate/up-swapped reference lands **146.36 ulp**. That +is the interleave-parity trap, pinned by measurement rather than by +reading. + +## Performance + +![Serving Pareto](perf/figures/pareto.png) + +Environment — every curve below: GPU 0 of `umbriel-b200-027` (NVIDIA B200, +driver 595.58.03, CUDA 13.0), one device for the whole campaign, +2026-07-27; `tensorrt_llm 1.3.0rc21`, `torch 2.11.0+cu130`, +`transformers 5.5.4`, single-process worker via `scripts/env.sh`. Sweeps: +`bench/perf.py`, ISL=OSL=1024, concurrency 1→256, requests per point = +concurrency × rounds (20 for con≤8, else 5). `baseline` 07:30-08:10 UTC, +`trtllm` 08:10-08:38 UTC (back to back), `iter1-mxfp8-moe` 13:29-13:55 UTC. +Perf is recorded, never gated. + +### Curves + +`peak tok/s/gpu` is the maximum over the sweep; the concurrency that +reaches it is in parentheses. + +| label | config | commit | accuracy | con=1 tok/s/user | peak tok/s/gpu | change | +|---|---|---|---|---|---|---| +| `trtllm` | `llm_args.yaml` only (`{}`) | `c5e6d53` | ungated reference | 394.66 | 7962.9 (con=128) | stock trtllm modeling, stock defaults | +| `baseline` | `llm_args.yaml` only (`{}`) | `c5e6d53` | 90.5989 | 228.74 | 4476.7 (con=256) | staircase at trtllm defaults, W4A16 MoE | +| `iter1-mxfp8-moe` | `llm_args.yaml` only (`{}`) | `41f955a` | **89.1585** | 380.35 | **8421.1 (con=128)** | MoE swapped to the W4A8 (MXFP8-activation) member of the same kernel family | + +No config variant was kept, so every curve runs the identity config and +`configs/` does not exist. Both accuracy cells are full 1319-question +gsm8k under the target's own protocol; each reproduced bit-identically +across two independent `bench/accuracy.py` invocations. + +Full point-by-point (tok/s/gpu, ISL=OSL=1024): + +| con | trtllm | baseline | iter1 | iter1 vs baseline | iter1 vs trtllm | iter1 tpot ms | iter1 ttft ms | +|---|---|---|---|---|---|---|---| +| 1 | 394.6 | 228.7 | 380.3 | +66.3% | −3.6% | 2.61 | 22.3 | +| 2 | 690.0 | 324.0 | 669.7 | +106.7% | −3.0% | 2.96 | 35.0 | +| 4 | 1181.2 | 527.8 | 1146.2 | +117.2% | −3.0% | 3.45 | 47.6 | +| 8 | 1856.8 | 749.9 | 1866.2 | +148.9% | +0.5% | 4.22 | 76.6 | +| 16 | 2799.8 | 1099.6 | 2872.8 | +161.3% | +2.6% | 5.47 | 109.2 | +| 32 | 4028.0 | 1434.6 | 4186.2 | +191.8% | +3.9% | 7.50 | 155.9 | +| 64 | 5851.3 | 2015.0 | 6066.8 | +201.1% | +3.7% | 10.34 | 218.0 | +| 128 | 7962.9 | 2760.4 | 8421.1 | +205.1% | +5.8% | 14.89 | 306.6 | +| 256 | 5509.8 | 4476.7 | 7267.3 | +62.3% | +31.9% | 34.74 | 455.2 | + +### Kept iterations + +**iter1 — the MoE call from W4A16 to W4A8 (MXFP8 activations).** + +*Evidence.* Four `nsys` steady-state windows (100 executor iterations, +every one pure decode) said the W4A16 forward was GPU-saturated at every +concurrency — idle −10.7% / −1.1% / −0.8% at con=1 / 128 / 256 (negative = +slight cross-stream overlap) — and that 60.8% / 91.9% / 89.8% of that GPU +time was the MoE call. Per layer per step at con=128 its two grouped GEMMs +cost 711.8 µs and 352.9 µs against the stock reference's 126.7 µs and +65.6 µs — **5.62× and 5.38× on byte-identical expert weights**. The +reference resolves this checkpoint's `quant_method: mxfp4` to +`W4A8_MXFP4_MXFP8`; our kernel's name carries `castBfloat16` (MXFP4 +expanded to bf16 for a `m128x8x16` bf16 MMA at 3 pipeline stages, 1-CTA +clusters) while the reference's feeds MXFP4 straight into a block-scaled +`m256x16x32` MMA at 6 stages, 2-CTA clusters. Non-MoE GPU work already +matched (3.41 vs 3.10 ms/step), so the MoE was the entire gap. + +*Change.* One catalog call swapped for one: `mxfp8_quantize(o, False, 512)` +produces e4m3 data plus per-32 UE8M0 **linear** block scales, and +`mxe4m3_mxe2m1_block_scale_moe_runner` consumes the pair. The quantizer +performs the hidden widening 2880 → 3072 itself, so the `torch/pad` that +fed the W4A16 call is gone — the forward issues 530 kernels per decode step +instead of 566. `weights.py` is untouched: both ops read the identical +prepared expert stack. The scale layout is spelled out rather than +defaulted, because the 128×4 swizzled buffer has the same byte count +whenever `num_tokens % 128 == 0` — every decode CUDA graph of 128 or 256 — +and is then accepted as a silently wrong answer. + +*Effect.* Every point improves, from +62.3% (con=256) to +205.1% +(con=128); peak throughput +88.1% (4476.7 → 8421.1 tok/s/gpu) and con=1 ++66.3% (228.74 → 380.35 tok/s/user). TTFT at con=1 falls 54.1 → 22.3 ms. +The curve also changes shape: it now peaks at con=128 and turns over at +256, exactly as the stock reference does. + +*Accuracy cost — a result, not noise.* gsm8k **89.1585**, gate passed by +3.56 points (threshold 85.6), but **1.44 points below** the W4A16 +measurement of 90.5989 recorded above. Two independent invocations of +`bench/accuracy.py` returned 89.1585 bit-identically (stderr 0.8564 both +times), as the W4A16 path reproduces 90.5989 bit-identically — so this is +the recipe's price, well outside the ~0.5-point band this gate leaves +unresolved. The mechanism is documented in the op's contract: FC1's +epilogue requantizes the activation to MXFP8 on the OCP scale (`floor` of +the block exponent, block max saturating) before FC2 reads it. For +reference, stock trtllm running this same recipe measured 90.2199 here +during the onboard. + +### Gap decomposition + +`trtllm-tuned` is **absent by construction**: no config variant was kept, +so it would be byte-identical to `trtllm` and the portable-config share of +the gap is 0. The decomposition below is therefore kernel-level, from +`nsys` windows of 100 pure-decode iterations at con=128: + +| | baseline (W4A16) | iter1 (W4A8) | trtllm | +|---|---|---|---| +| step (wall) | 41.84 ms | **16.10 ms** | 16.18 ms | +| GPU active | 42.30 ms | 11.60 ms | 10.60 ms | +| GPU idle | −1.1% | **+27.9%** | +34.5% | +| MoE total | 38.89 ms | 8.34 ms | 7.51 ms | +| non-MoE | 3.41 ms | 3.26 ms | 3.10 ms | +| kernels / step | 566 | 530 | 502 | + +The target's decode step is now **16.10 ms against the reference's +16.18 ms**, and the bottleneck has moved: this forward was 100% GPU-bound +and is now 27.9% idle, i.e. host-bound in the same regime as the stock +reference. What remains at the kernel level is an 11% MoE-GEMM difference +from tactic selection alone — the autotuner picks `t128x8x512_s3` / +`m128x8x32` / 1-CTA for us (139.3 and 74.1 µs per layer) against the +reference's `t128x16x256u2_s6` / `m256x16x32` / 2-CTA (126.7 and 65.6 µs) +on the same op family. The new quantize call costs 143.12 µs per step +(36 calls, 3.98 µs each) — 1.2% of GPU-active time. + +### Reference lines + +- `trtllm` = the original HF checkpoint under **stock in-tree trtllm + modeling**, sharing this target's `llm_args.yaml` (identity config, `{}` + at tp1) and otherwise trtllm defaults — `bench/perf.py --trtllm`. Since + `iter1`, both systems run the same W4A8 numerical recipe, so this is now + a like-for-like comparison; against `baseline` it was not (that curve is + W4A16). +- `trtllm-tuned` is absent by construction (no kept config variant). +- The reference is **ungated** — no accuracy gate is run against in-tree + modeling. Pinned to `tensorrt_llm 1.3.0rc21`. + +### Measurement caveats + +- The host CPU of `umbriel-b200-027` is shared and other tenants ran GPU + work on neighbouring devices during the campaign; GPU 0 was reserved + throughout, but host and chassis contention is an error bar on every + point. +- **Session variance at mid-curve concurrencies exceeds the 3% rule of + thumb.** Under the W4A16 MoE, con=128 at rounds=2 measured 3097.2, + 3101.0 and 2725.2 tok/s across three sessions of configurations later + shown equivalent (12.1% spread) against 0.35% at con=1 and 0.77% at + con=256. Under the W4A8 MoE the same point was far steadier: two control + probes 26 minutes apart measured 8615.3 and 8603.4 tok/s (0.14%). Probe + A/Bs in this campaign therefore always carry a same-session control. +- Wrap-up spot-check of the reference (`--trtllm`, con=1/32/256, default + rounds, 5 h after the label): 394.0 (−0.15%), 4267.2 (+5.94%), 5607.6 + (+1.77%). The `trtllm` label was **not** re-swept: the +5.94% sits inside + the session variance above, and re-measuring one curve alone would break + the back-to-back pairing with `baseline`. +- Profiled runs are never comparable to clean ones: under `nsys` the + W4A16 con=128 point measured 2468.4 tok/s against 2760.4 clean. + +Figure regeneration: + +``` +uv run utils/plot_pareto.py targets/gpt-oss-120b/sm_100/tp1/perf/data \ + -o targets/gpt-oss-120b/sm_100/tp1/perf/figures/pareto.png +``` + +### Tried and rejected (with the number that rejected it) + +The first four were measured **against the W4A16 forward**, when the GPU +was saturated; that scope matters, because the first one changed verdict +after iter1 and had to be re-tested. + +| hypothesis | axis | measured | verdict | +|---|---|---|---| +| decode CUDA-graph coverage 256 (`cuda_graph_config.max_batch_size: 256`, `enable_padding: true`) — **under W4A16** | config | con=1 +0.04%, con=128 +0.12%, con=256 +0.77% | inert: no idle existed to recover | +| `max_seq_len: 12288` (blocks/seq 4096 → 384, confirmed in `server.log`) | config | con=1 −0.3%, con=256 +0.07% | inert: the per-step H2D block-offset staging measures 4.19 MB / 87 µs = 0.2% of a con=128 step | +| `cute_dsl_bf16_gemm_blackwell` for the qkv / o / router GEMMs | modeling | graph-replay kernel time at M=1: 7.12 vs cublas+bias 5.95 µs (qkv), 7.30 vs 6.28 (o), 5.61 vs 5.18 (router) | rejected: slower at every shape before the bias cublas fuses | +| FC1 K-padding 3072 → 2944 (−4.17% of FC1 weight bytes) | modeling | MoE µs/layer: T=128 1176.3 vs 1151.4 (worse), T=256 1177.0 vs 1191.9 | rejected: no consistent gain | +| decode CUDA-graph coverage 256 — **re-tested under W4A8** | config | see below | **trade-off, not shipped** | + +**The graph-coverage trade-off, re-measured after iter1.** Once the MoE +stopped saturating the GPU, the con=256 point began to collapse the way the +stock reference's does (8421.1 at con=128 → 7267.3 at con=256). Capturing a +batch-256 decode graph recovers it, but costs the con=128 point. Same +session, control probed twice before and after (rounds=2, con=128/256): + +| config | con=128 | con=256 | +|---|---|---| +| identity (control, 14:04) | 8615.3 | 7221.6 | +| identity (control, 14:30) | 8603.4 | 7516.3 | +| `max_batch_size: 256`, padding on | 7897.7 (−8.3%) | 9805.3 (+33.1%) | +| `max_batch_size: 256`, padding off | 7753.2 (−10.0%) | 9792.2 (+32.9%) | + +Padding is not the mechanism — exact-hit coverage regresses con=128 just as +much — and neither is memory: the KV pool moves 99.31 → 99.15 GiB (0.16%) +and graph memory 10.83 → 11.28 GiB. Decode TPOT at con=128 rises 14.35 → +15.97 ms with the extra graph present. The keep rule is regress-nowhere, so +this is **not shipped**; a deployment that only ever serves con≥256 should +adopt it deliberately, since its con=256 point (9805.3 tok/s at 38.6 +tok/s/user) is outside anything the shipped config reaches. + +### Remaining headroom + +- **The target is now host-bound at the throughput end** (27.9% GPU idle at + con=128), in the same regime as stock trtllm (34.5%) and for the same + reason: stock executor Python between steps. No knob in the tuner's table + reaches it. +- **~11% of the MoE GEMM time is tactic selection**, not recipe: same op, + same weights, different autotuner choice than the in-tree path makes + (`t128x8x512_s3`/1-CTA vs `t128x16x256u2_s6`/2-CTA). Worth a look at what + drives the autotuner's bucket set. +- **The con=256 graph-coverage trade-off above** is a real +33% at the + throughput end blocked by a −9% at con=128 whose mechanism is not + established. Establishing it would unlock the point. +- **Decode cost still saturates with batch**, measured on the W4A16 op per + layer: T=1 67.3 µs, T=64 986.9, T=128 1151.4, T=256 1191.9, T=1024 + 1307.3 — +13.5% for 8× the tokens above T=128. The expert-weight read is + a fixed per-step cost once the batch touches all 128 experts. +- **Per-window KV pools are not the win** and that thread is retired: the + pool is sized for full attention (99.27 GiB, 1,445,664 tokens = 705 + requests at ISL+OSL 2048) but capacity never binds at con≤256 (36.3% + used), the per-layer window already limits what the FMHA kernel reads, + and stock trtllm sizes the same single pool on this checkpoint. +- **The accuracy cost of W4A8 (−1.44 points) is the price of the frontier + above.** A deployment that needs the last point of gsm8k accuracy should + run the W4A16 forward (`baseline`, commit `c5e6d53`) and accept a third + of the throughput. diff --git a/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/__init__.py b/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/__init__.py new file mode 100644 index 000000000000..12aed53d4a69 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/__init__.py @@ -0,0 +1,3 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""gpt-oss-120b / sm_103 / tp1.""" diff --git a/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py b/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py new file mode 100644 index 000000000000..66f4184dcf2e --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py @@ -0,0 +1,677 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Staircase target: gpt-oss-120b / sm_103 / tp1 — self-contained modeling code. + +Flat single-entry forward assembled from catalog entries only; every call +that creates or transforms a tensor is a catalog entry, everything else is +tensor-metadata reads and Python control flow. Attention consumes runtime +state fully explicitly through thop_attention: per-step arguments are +projected from the engine-prepared TrtllmAttentionMetadata once per forward +in _build_step_args and shared by all layers. Two attention arguments are +per-layer here rather than per-step: the fp32 sink logits (one extra softmax +denominator column per query head) and attention_window_size — this +checkpoint alternates sliding_attention (window 128) and full_attention +layers, and the window is a pure mask, so one shared pool and one +block-offset table serve both kinds. + +Every layer's MLP is an MXFP4 sparse mixture of experts (128 experts, top-4, +renormalized, clamped GLU) run W4A8: mxfp8_quantize turns the bf16 hidden +states into e4m3 data plus per-32 UE8M0 block scales — widening hidden 2880 +to the FC1 K alignment 3072 inside that call — and one +mxe4m3_mxe2m1_block_scale_moe_runner call per layer covers routing, both +grouped GEMMs, the clamped activation, the MXFP8 requantization between them +and the combine. The router bias is folded into the router GEMM's fused bias +because the MoE op silently ignores routing_bias on this routing method. The +expert weights are declared in the kernel-ready padded/shuffled/swizzled +layout — identical for the W4A16 and W4A8 members of this kernel family — and +the manifest loop in weights.py transforms the checkpoint's block/scale +tensors into it at load time. + +Weights are target-owned: a flat ParameterDict declared here (HF [out, in] +storage so checkpoint rows copy in unchanged), loaded by the manifest loop +in the sibling weights.py, with column-major GEMM views derived once after +load. The registration shell inherits DecoderModelForCausalLM for lm_head, +packed-batch logits gathering, and the meta-init/load/post-load hooks. + +The import-time and first-forward contract checks below fail fast on drift. +""" + +import math + +import torch +from torch import nn +from transformers import PretrainedConfig + +from tensorrt_llm._torch.attention.backends.interface import AttentionMetadata +from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata +from tensorrt_llm._torch.model_config import ModelConfig +from tensorrt_llm._torch.models.modeling_utils import ( + DecoderModel, + DecoderModelForCausalLM, + register_auto_model, +) +from tensorrt_llm._torch.staircase.catalog.attention.fused_qk_norm_rope import fused_qk_norm_rope +from tensorrt_llm._torch.staircase.catalog.attention.thop_attention import thop_attention +from tensorrt_llm._torch.staircase.catalog.gemm.cublas_mm import cublas_mm +from tensorrt_llm._torch.staircase.catalog.moe.mxe4m3_mxe2m1_block_scale_moe_runner import ( # noqa: E501 + mxe4m3_mxe2m1_block_scale_moe_runner, +) +from tensorrt_llm._torch.staircase.catalog.norm.flashinfer_fused_add_rmsnorm import ( # noqa: E501 + flashinfer_fused_add_rmsnorm, +) +from tensorrt_llm._torch.staircase.catalog.norm.flashinfer_rmsnorm import flashinfer_rmsnorm +from tensorrt_llm._torch.staircase.catalog.quantization.mxfp8_quantize import mxfp8_quantize +from tensorrt_llm._torch.staircase.catalog.torch.embedding import embedding +from tensorrt_llm._torch.staircase.catalog.torch.empty import empty +from tensorrt_llm._torch.staircase.catalog.torch.reshape import reshape + +from . import weights as _weights + +# The GPU architecture this target IS. Routing will not send another one here, +# but a direct instantiation could, and the certification is per arch: this +# assert is what the version pin used to be. In-tree the version moves with +# the code, so pinning it is meaningless; the architecture does not. +_SM = (10, 3) + + +def _check_static_contract() -> None: + """Import-time fail-fast: op symbol existence.""" + for op in ( + "fused_qk_norm_rope", + "flashinfer_rmsnorm", + "flashinfer_fused_add_rmsnorm", + "cublas_mm", + "mxfp8_quantize", + "mxe4m3_mxe2m1_block_scale_moe_runner", + ): + assert hasattr(torch.ops.trtllm, op), f"missing op trtllm::{op}" + from tensorrt_llm.bindings.internal import thop + + assert hasattr(thop, "attention"), "missing pybind thop.attention" + + +_check_static_contract() + +# Metadata fields consumed each step (sourcing mirrors the in-tree +# FallbackFmha for this trtllm version; existence checked at first forward). +_STEP_FIELDS = ( + "kv_lens_cuda_runtime", + "kv_lens_runtime", + "host_total_kv_lens", + "prompt_lens_cuda_runtime", + "prompt_lens_cpu_runtime", + "host_request_types_runtime", + "kv_cache_block_offsets", + "host_kv_cache_pool_pointers", + "host_kv_cache_pool_mapping", + "effective_workspace", + "tokens_per_block", + "max_num_requests", + "max_context_length", + "max_seq_len", + "num_contexts", + "num_ctx_tokens", + "trtllm_gen_jit_warmup", + "use_paged_context_fmha", + "effective_beam_width", + "cache_indirection", + "block_ids_per_seq", + "max_context_q_len_override", + "is_cross", + "is_spec_decoding_enabled", + "use_spec_decoding", + "is_spec_dec_tree", + "spec_decoding_generation_lengths", + "spec_decoding_position_offsets_for_cpp", + "spec_decoding_packed_mask", + "spec_decoding_bl_tree_mask_offset", + "spec_decoding_bl_tree_mask", + "max_total_draft_tokens", + "spec_bl_tree_first_sparse_mask_offset_kv", + "num_sparse_topk", + "flash_mla_tile_scheduler_metadata", + "flash_mla_num_splits", + # Added between 1.3.0rc21 and 1.3.0rc26. Both are engine-prepared + # per-instance constants (max_num_sequences defaults to + # max_num_requests; the tree-mask flag is set from + # is_spec_dec_dynamic_tree), so they project like the rest. + "max_num_sequences", + "force_prepare_spec_dec_tree_mask", +) + + +def _build_step_args(md: TrtllmAttentionMetadata) -> dict: + """Project the prepared metadata onto thop_attention's explicit batch + state, once per forward; every runtime-owned value passes through as + the engine prepared it. CUDA-graph classes: tensors are engine-owned + persistent buffers refreshed in place (reference class); Python ints + are per-capture constants (host-derived class — a decode-only graph + always sees num_contexts == 0). attention_window_size is absent here + on purpose: it varies per layer and is passed at the call site.""" + return dict( + sequence_length=md.kv_lens_cuda_runtime, + host_past_key_value_lengths=md.kv_lens_runtime, + host_total_kv_lens=md.host_total_kv_lens, + context_lengths=md.prompt_lens_cuda_runtime, + host_context_lengths=md.prompt_lens_cpu_runtime, + host_request_types=md.host_request_types_runtime, + kv_cache_block_offsets=md.kv_cache_block_offsets, + host_kv_cache_pool_pointers=md.host_kv_cache_pool_pointers, + host_kv_cache_pool_mapping=md.host_kv_cache_pool_mapping, + workspace_=md.effective_workspace, + tokens_per_block=md.tokens_per_block, + max_num_requests=md.max_num_requests, + max_context_length=md.max_context_length, + max_seq_len=md.max_seq_len, + num_contexts=md.num_contexts, + num_ctx_tokens=md.num_ctx_tokens, + trtllm_gen_jit_warmup=md.trtllm_gen_jit_warmup, + use_paged_context_fmha=md.use_paged_context_fmha, + beam_width=md.effective_beam_width, + cache_indirection=md.cache_indirection, + block_ids_per_seq=md.block_ids_per_seq, + max_context_q_len_override=md.max_context_q_len_override, + is_cross=md.is_cross, + is_spec_decoding_enabled=md.is_spec_decoding_enabled, + use_spec_decoding=md.use_spec_decoding, + is_spec_dec_tree=md.is_spec_dec_tree, + spec_decoding_generation_lengths=md.spec_decoding_generation_lengths, + spec_decoding_position_offsets_for_cpp=md.spec_decoding_position_offsets_for_cpp, + spec_decoding_packed_mask=md.spec_decoding_packed_mask, + spec_decoding_bl_tree_mask_offset=md.spec_decoding_bl_tree_mask_offset, + spec_decoding_bl_tree_mask=md.spec_decoding_bl_tree_mask, + spec_decoding_target_max_draft_tokens=md.max_total_draft_tokens, + spec_bl_tree_first_sparse_mask_offset_kv=md.spec_bl_tree_first_sparse_mask_offset_kv, + num_sparse_topk=md.num_sparse_topk, + flash_mla_tile_scheduler_metadata=md.flash_mla_tile_scheduler_metadata, + flash_mla_num_splits=md.flash_mla_num_splits, + max_num_sequences=md.max_num_sequences, + force_prepare_spec_dec_tree_mask=md.force_prepare_spec_dec_tree_mask, + ) + + +# Per-call constants of this target's call shape — the values the in-tree +# path sources from the attention module and forward args: packed-QKV +# causal GQA, RoPE applied outside (fused_qk_norm_rope), bf16 activations +# over a bf16 KV pool, no MLA / mRoPE / cross / relative-bias / sparse +# features. attention_sinks and attention_window_size are per-layer, not +# constants, and are passed at the call site. +_CALL_CONSTANTS = dict( + output_sf=None, + k=None, + v=None, + is_fused_qkv=True, + update_kv_cache=True, + predicted_tokens_per_seq=1, + attention_input_type=0, + is_mla_enable=False, + mask_type=1, + q_scaling=1.0, + quant_mode=0, + kv_scale_orig_quant=None, + kv_scale_quant_orig=None, + # In-kernel RoPE is disabled (position_embedding_type=0): rotation + # happens outside in fused_qk_norm_rope. The rope_* values below are + # the contract's inert placeholders, not this model's rope config. + position_embedding_type=0, + rotary_inv_freq=None, + rotary_cos_sin=None, + rope_dim=0, + rope_base=10000.0, + rope_scale_type=0, + rope_scale=1.0, + rope_short_m_scale=1.0, + rope_long_m_scale=1.0, + rope_max_positions=1024, + rope_original_max_positions=1024, + out_scale=None, + latent_cache=None, + q_pe=None, + q_lora_rank=None, + kv_lora_rank=None, + qk_nope_head_dim=None, + qk_rope_head_dim=None, + v_head_dim=None, + rope_append=None, + chunked_prefill_buffer_batch_size=1, + attention_chunk_size=None, + softmax_stats_tensor=None, + sparse_kv_indices=None, + sparse_kv_offsets=None, + sparse_attn_indices=None, + sparse_attn_offsets=None, + sparse_attn_indices_block_size=0, + mrope_rotary_cos_sin=None, + mrope_position_deltas=None, + helix_position_offsets=None, + helix_is_inactive_rank=None, + cross_kv=None, + relative_attention_bias=None, + relative_attention_max_distance=0, + # Added between 1.3.0rc21 and 1.3.0rc26; held at the values the op had + # before they existed. All three are MLA-only surface that this target + # does not use -- skip_correction is forced to 0.0 for a non-MLA layer by + # the engine's own resolver, and kv_norm_* folds an MLA kv_a_layernorm + # that does not exist here. + kv_norm_weight=None, + kv_norm_eps=1e-6, + skip_correction_threshold=0.0, +) + +# The MoE op's clamped gated activation: (up + beta) * gate * +# sigmoid(alpha * gate) after gate.clamp(max=limit), up.clamp(+-limit) — +# exactly this checkpoint's expert activation. alpha/beta are fixed in the +# HF reference; limit is config (swiglu_limit). +_GLU_ALPHA = 1.702 +_GLU_BETA = 1.0 +# Renormalize routing (top-k first, then fp32 softmax over the selected +# logits) and the only gated-activation kernel in this dtype family. +_ROUTING_METHOD_RENORMALIZE = 1 +_ACT_TYPE_SWIGLU = 0 +# FC1's K alignment for the trtllm-gen MXFP4 weight family: it sizes the +# declared expert operands and is the alignment mxfp8_quantize pads the +# hidden states up to, so the two always agree. +_FC1_K_ALIGN = 512 +# The MoE op reads the activation scales as a linear (row-major) buffer. +# The 128x4 swizzled order has the same byte count whenever num_tokens is a +# multiple of 128 — which every decode CUDA-graph batch of 128 or 256 is — +# and is then accepted silently as a wrong answer, so this is spelled out +# rather than left to the quantizer's default. +_LINEAR_SCALE_LAYOUT = False + + +def _pad_up(x: int, align: int) -> int: + return (x + align - 1) // align * align + + +class StaircaseCore(DecoderModel): + def __init__(self, model_config: ModelConfig): + super().__init__(model_config) + cfg = model_config.pretrained_config + assert cfg is not None + + # This target IS this geometry — assert, never adapt. + assert torch.cuda.get_device_capability() == _SM, ( + f"target certified on sm_{_SM[0]}{_SM[1]}, running on " + f"sm_{''.join(map(str, torch.cuda.get_device_capability()))}" + ) + assert model_config.mapping.tp_size == 1, "tp1 target" + assert model_config.mapping.pp_size == 1, "tp1 target" + # This checkpoint's config.json declares no dtype at all, so + # `pretrained_config.torch_dtype` is None. The engine resolves bf16 + # regardless (ModelConfig.torch_dtype defaults it), but the shell + # sizes lm_head from the *pretrained* value and would materialize it + # in the torch default fp32 — two layers away from where it is read. + # The shell normalizes that before super().__init__ (see below), so + # by the time the core is built both surfaces agree, and both are + # asserted: the declaration the shell and the KV pool are sized from, + # and the value quant_mode=0 must agree with. + dt = model_config.torch_dtype + assert cfg.torch_dtype == torch.bfloat16, cfg.torch_dtype + assert dt == torch.bfloat16, f"bf16 target, engine resolved {dt}" + assert not cfg.tie_word_embeddings, "untied lm_head" + assert cfg.attention_bias, "q/k/v/o carry bias" + assert cfg.hidden_act == "silu", "clamped GLU over a silu-shaped gate" + + # bf16 KV pool only: an fp8 pool needs quant_mode=128 plus the two + # fp32 scale tensors, which this target does not build. The expert + # quantization (mxfp4 blocks + e8m0 scales) is consumed directly by + # the MoE op, so quant_config.quant_algo — which the engine reads + # off quantization_config as W4A8_MXFP4_MXFP8, an mxfp8-activation + # recipe this target does not implement — is inert here. + kv_algo = model_config.quant_config.kv_cache_quant_algo + assert kv_algo is None, f"unsupported KV algo {kv_algo}; bf16 pool only" + quant = getattr(cfg, "quantization_config", None) + assert isinstance(quant, dict) and quant.get("quant_method") == "mxfp4", ( + "expert weights must be mxfp4 blocks + e8m0 scales" + ) + + self.num_layers = cfg.num_hidden_layers + self.hidden = cfg.hidden_size + self.heads_q = cfg.num_attention_heads + self.heads_kv = cfg.num_key_value_heads + self.head_dim = cfg.head_dim + self.eps = cfg.rms_norm_eps + + # RoPE: YaRN over the full head_dim, half-split (neox) pairs. The + # engine hands this checkpoint the transformers-5.x migrated rope + # dict (rope_theta lives inside it and cfg.rope_theta is absent); + # a checkpoint written before that migration keeps the flat field, + # so both shapes are resolved here and every scalar the ramp + # depends on is asserted rather than defaulted. + rope = getattr(cfg, "rope_scaling", None) or getattr(cfg, "rope_parameters", None) + assert isinstance(rope, dict), "rope parameters must be a dict" + assert rope.get("rope_type") == "yarn", "YaRN rope scaling" + assert rope.get("partial_rotary_factor", 1.0) == 1.0, "full-width rotation" + assert rope.get("attention_factor") is None, "attention_factor is derived" + assert rope.get("mscale") is None and rope.get("mscale_all_dim") is None, ( + "the derived attention_factor assumes no mscale pair" + ) + self.theta = float(rope.get("rope_theta", getattr(cfg, "rope_theta", 0.0))) + assert self.theta > 0.0, "rope theta" + self.rotary_dim = self.head_dim + self.yarn_factor = float(rope["factor"]) + beta_fast = float(rope["beta_fast"]) + beta_slow = float(rope["beta_slow"]) + orig_max = float(rope["original_max_position_embeddings"]) + truncate = bool(rope.get("truncate", True)) + + def correction_dim(rotations: float) -> float: + return ( + self.rotary_dim + * math.log(orig_max / (rotations * 2.0 * math.pi)) + / (2.0 * math.log(self.theta)) + ) + + low = correction_dim(beta_fast) + high = correction_dim(beta_slow) + if truncate: + low, high = math.floor(low), math.ceil(high) + self.yarn_low = max(low, 0.0) + self.yarn_high = min(high, self.rotary_dim - 1.0) + self.yarn_attn_factor = ( + 0.1 * math.log(self.yarn_factor) + 1.0 if self.yarn_factor > 1.0 else 1.0 + ) + + # Sliding window: half the layers mask to the newest `window` keys, + # the rest are plain causal. The window is a mask only — the op + # appends at absolute positions — so both kinds share one pool, one + # layer->pool mapping and one block-offset table. + layer_types = list(cfg.layer_types) + assert len(layer_types) == self.num_layers + assert set(layer_types) <= {"sliding_attention", "full_attention"}, layer_types + self.sliding = [t == "sliding_attention" for t in layer_types] + self.window = cfg.sliding_window + assert any(self.sliding) and self.window > 0, "alternating sliding window" + + # Sparse MoE on every layer; no dense MLP branch exists. + self.num_experts = cfg.num_local_experts + self.topk = cfg.num_experts_per_tok + self.inter = cfg.intermediate_size + self.glu_limit = float(cfg.swiglu_limit) + assert 0 < self.topk < self.num_experts, "MoE top-k bound" + + # Kernel bounds the geometry must fit: fused_qk_norm_rope's head_dim + # set, thop_attention's GQA rule, and the MoE op's padded operand + # widths (the padded intermediate is what `intermediate_size` means + # to that call, and hidden_states reach it widened to fc1_k_pad by + # the quantizer). + assert self.head_dim in (64, 128, 256), "fused_qk_norm_rope head_dim set" + assert self.heads_q % self.heads_kv == 0, "thop_attention GQA rule" + self.inter_pad = _pad_up(self.inter, 128) + self.fc1_k_pad = _pad_up(self.hidden, _FC1_K_ALIGN) + self.fc2_rows_pad = _pad_up(self.hidden, 128) + assert self.hidden % 32 == 0, "valid_hidden_size must be a multiple of 32" + assert self.inter % 32 == 0, "valid_intermediate_size must be a multiple of 32" + + q_width = self.heads_q * self.head_dim + kv_width = self.heads_kv * self.head_dim + + # Weight declaration: HF [out, in] row-major storage so checkpoint + # rows copy in unchanged; GEMM consumes .t() column-major views + # built after load. The expert tensors are the exception — they are + # declared in the MoE op's kernel-ready layout (padded, row-permuted + # weights/scales/biases), which weights.py builds from the + # checkpoint's block/scale tensors during the load. Meta-init + # intercepts torch.empty here — real CUDA storage arrives when the + # engine materializes the registry. + def P(*shape, dtype=dt): + return nn.Parameter(torch.empty(*shape, dtype=dtype), requires_grad=False) + + fc1_rows = 2 * self.inter_pad + w = nn.ParameterDict() + for i in range(self.num_layers): + w[f"l{i}_norm1"] = P(self.hidden) + w[f"l{i}_qkv"] = P(q_width + 2 * kv_width, self.hidden) + w[f"l{i}_qkv_bias"] = P(q_width + 2 * kv_width) + w[f"l{i}_sinks"] = P(self.heads_q, dtype=torch.float32) + w[f"l{i}_o"] = P(self.hidden, q_width) + w[f"l{i}_o_bias"] = P(self.hidden) + w[f"l{i}_norm2"] = P(self.hidden) + w[f"l{i}_router"] = P(self.num_experts, self.hidden) + w[f"l{i}_router_bias"] = P(self.num_experts) + w[f"l{i}_fc1_w"] = P(self.num_experts, fc1_rows, self.fc1_k_pad // 2, dtype=torch.uint8) + w[f"l{i}_fc1_s"] = P( + self.num_experts, fc1_rows, self.fc1_k_pad // 32, dtype=torch.uint8 + ) + w[f"l{i}_fc1_b"] = P(self.num_experts, fc1_rows, dtype=torch.float32) + w[f"l{i}_fc2_w"] = P( + self.num_experts, + self.fc2_rows_pad, + self.inter_pad // 2, + dtype=torch.uint8, + ) + w[f"l{i}_fc2_s"] = P( + self.num_experts, + self.fc2_rows_pad, + self.inter_pad // 32, + dtype=torch.uint8, + ) + w[f"l{i}_fc2_b"] = P(self.num_experts, self.fc2_rows_pad, dtype=torch.float32) + w["final_norm"] = P(self.hidden) + w["embed"] = P(cfg.vocab_size, self.hidden) + self.w = w + + self._layers: list | None = None + self._call_tensors: dict | None = None + self._step_contract_checked = False + + def build_layer_views(self) -> None: + """Post-load derivation: per-layer tuples of column-major GEMM views, + norm weights and expert operands (kills hot-path dict lookups; .t() + is zero-copy), plus the constant fp32 call tensors the MoE and RoPE + calls require. Meta is over here, so real tensors may be created.""" + w = self.w + device = w["final_norm"].device + # Per-expert activation scalars, sized by local_num_experts. + self._call_tensors = { + "alpha": torch.full( + (self.num_experts,), _GLU_ALPHA, dtype=torch.float32, device=device + ), + "beta": torch.full((self.num_experts,), _GLU_BETA, dtype=torch.float32, device=device), + "limit": torch.full( + (self.num_experts,), self.glu_limit, dtype=torch.float32, device=device + ), + # fused_qk_norm_rope requires valid q/k norm weight tensors even + # with is_qk_norm=False, where their values are unused. + "no_qk_norm": torch.zeros(self.head_dim, dtype=w["final_norm"].dtype, device=device), + } + layers = [] + for i in range(self.num_layers): + next_norm = w[f"l{i + 1}_norm1"] if i + 1 < self.num_layers else w["final_norm"] + layers.append( + ( + w[f"l{i}_qkv"].t(), + w[f"l{i}_qkv_bias"], + w[f"l{i}_sinks"], + w[f"l{i}_o"].t(), + w[f"l{i}_o_bias"], + w[f"l{i}_norm2"], + w[f"l{i}_router"].t(), + w[f"l{i}_router_bias"], + w[f"l{i}_fc1_w"], + w[f"l{i}_fc1_s"], + w[f"l{i}_fc1_b"], + w[f"l{i}_fc2_w"], + w[f"l{i}_fc2_s"], + w[f"l{i}_fc2_b"], + next_norm, + ) + ) + self._layers = layers + + def _check_step_contract(self, md, position_ids) -> None: + """First-forward fail-fast: the metadata fields this target consumes + must exist (they are private trtllm surface, pinned by version), and + the KV pool must be the single shared pool the sliding-window + surface is certified over. Everything checked is fixed at engine + construction — once per model instance is sound.""" + missing = [name for name in _STEP_FIELDS if not hasattr(md, name)] + assert not missing, f"metadata fields missing: {missing}" + assert position_ids.dtype == torch.int32 + pools = {row[0] for row in md.host_kv_cache_pool_mapping.tolist()} + assert pools == {0}, ( + f"multi-pool KV addressing is not certified; layer->pool ids {sorted(pools)}" + ) + self._step_contract_checked = True + + def forward( + self, + attn_metadata: AttentionMetadata, + input_ids: torch.IntTensor | None = None, + position_ids: torch.IntTensor | None = None, + inputs_embeds: torch.FloatTensor | None = None, + lora_params: dict | None = None, + **kwargs, + ) -> torch.Tensor: + assert self._layers is not None and self._call_tensors is not None, ( + "load_weights must run before forward" + ) + assert position_ids is not None + assert isinstance(attn_metadata, TrtllmAttentionMetadata) + # Inputs this target does not implement must fail loudly, not be + # silently dropped (unlike runtime-owned features, which pass through). + assert lora_params is None, "LoRA is not implemented by this target" + assert kwargs.get("spec_metadata") is None, ( + "speculative decoding is not implemented by this target" + ) + if not self._step_contract_checked: + self._check_step_contract(attn_metadata, position_ids) + + step = _build_step_args(attn_metadata) + # Full-attention layers take the no-window value; sliding layers take + # the checkpoint's window. Both are host-derived per-capture constants + # (max_seq_len is an engine-construction constant). + full_window = attn_metadata.max_seq_len + const = self._call_tensors + no_qk_norm = const["no_qk_norm"] + pos = reshape(position_ids, [-1]) + + if inputs_embeds is None: + assert input_ids is not None + h = embedding(input_ids, self.w["embed"]) + else: + h = inputs_embeds + num_tokens = h.shape[0] + dt = h.dtype + attn_out = empty([num_tokens, self.heads_q * self.head_dim], dt, h.device) + + x = flashinfer_rmsnorm(h, self.w["l0_norm1"], self.eps) + residual = h + for i in range(self.num_layers): + ( + w_qkv, + b_qkv, + sinks, + w_o, + b_o, + w_n2, + w_rt, + b_rt, + fc1_w, + fc1_s, + fc1_b, + fc2_w, + fc2_s, + fc2_b, + w_next, + ) = self._layers[i] + qkv = cublas_mm(x, w_qkv, b_qkv) + fused_qk_norm_rope( + qkv, + num_heads_q=self.heads_q, + num_heads_k=self.heads_kv, + num_heads_v=self.heads_kv, + head_dim=self.head_dim, + rotary_dim=self.rotary_dim, + eps=self.eps, + q_weight=no_qk_norm, + k_weight=no_qk_norm, + base=self.theta, + is_neox=True, + position_ids=pos, + factor=self.yarn_factor, + low=self.yarn_low, + high=self.yarn_high, + attention_factor=self.yarn_attn_factor, + is_qk_norm=False, + ) + thop_attention( + q=qkv, + output=attn_out, + local_layer_idx=i, + num_heads=self.heads_q, + num_kv_heads=self.heads_kv, + head_size=self.head_dim, + attention_sinks=sinks, + attention_window_size=self.window if self.sliding[i] else full_window, + **step, + **_CALL_CONSTANTS, + ) + o = cublas_mm(attn_out, w_o, b_o) + flashinfer_fused_add_rmsnorm(o, residual, w_n2, self.eps) + # The router bias rides the GEMM epilogue: the MoE op silently + # ignores routing_bias on this routing method. + router_logits = cublas_mm(o, w_rt, b_rt) + # The quantizer owns the hidden widening 2880 -> fc1_k_pad: it + # zero-fills the padded columns and their scale bytes, which + # multiply zero-valued padded weights. + hidden_fp8, hidden_sf = mxfp8_quantize(o, _LINEAR_SCALE_LAYOUT, _FC1_K_ALIGN) + moe = mxe4m3_mxe2m1_block_scale_moe_runner( + router_logits, + None, + hidden_fp8, + hidden_sf, + fc1_w, + fc1_s, + fc1_b, + const["alpha"], + const["beta"], + const["limit"], + fc2_w, + fc2_s, + fc2_b, + self.num_experts, + self.topk, + None, + None, + self.inter_pad, + self.hidden, + self.inter, + 0, + self.num_experts, + None, + _ROUTING_METHOD_RENORMALIZE, + _ACT_TYPE_SWIGLU, + ) + flashinfer_fused_add_rmsnorm(moe, residual, w_next, self.eps) + x = moe + return x + + +@register_auto_model("StaircaseGptOss120bSm103Tp1") +class StaircaseGptOss120bSm103Tp1(DecoderModelForCausalLM[StaircaseCore, PretrainedConfig]): + def __init__(self, model_config: ModelConfig): + cfg = model_config.pretrained_config + assert cfg is not None + # DecoderModelForCausalLM sizes lm_head from the *pretrained* dtype, + # which this checkpoint's config.json does not declare — leaving + # lm_head.weight fp32 while every other tensor is bf16, and failing + # two layers away from here. The engine has already resolved the + # dtype it will run at; adopt it, rather than let the default of a + # missing field decide. Only fills the gap: a declared dtype wins. + if cfg.torch_dtype is None: + cfg.torch_dtype = model_config.torch_dtype + super().__init__( + StaircaseCore(model_config), + config=model_config, + hidden_size=cfg.hidden_size, + vocab_size=cfg.vocab_size, + ) + + def load_weights(self, weights, *args, **kwargs): + _weights.load(self, weights) + + def post_load_weights(self): + super().post_load_weights() + self.model.build_layer_views() diff --git a/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/weights.py b/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/weights.py new file mode 100644 index 000000000000..23b87584d6cd --- /dev/null +++ b/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/weights.py @@ -0,0 +1,246 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Weight manifest and loader: gpt-oss-120b / sm_103 / tp1. + +MANIFEST is a data table: target parameter -> the checkpoint keys that fill +it, each with an optional destination index into the parameter and an +optional source transform. tp1: no sharding — every checkpoint tensor +reaches exactly one parameter whole. + +Storage layout is HF [out, in] row-major, so every attention/router/embed +copy is layout-preserving; the GEMM-side column-major views are derived +after load in modeling.build_layer_views(). Two families need a transform: + +* **fp32 promotions.** The expert biases and the attention sink logits are + bf16 on disk, but the MoE op rejects a bf16 bias and thop_attention reads + the sink buffer as raw fp32. + +* **The MXFP4 expert operands.** The checkpoint stores each layer's experts + as `gate_up_proj_blocks` [E, 2I, H/32, 16] / `_scales` [E, 2I, H/32] and + `down_proj_blocks` [E, H, I/32, 16] / `_scales` [E, H, I/32] — E2M1 codes + two per byte (low nibble = even K index) with one E8M0 exponent per 32 K + elements, already in [out, in] orientation. The MoE op wants them padded, + row-permuted and (scales only) swizzled; `_prep_fc1` / `_prep_fc2` below + are that recipe, applied on device, one layer at a time. + + The parity trap: this checkpoint's `2I` axis runs (gate, up, gate, up, + ...) — HF reads `gate = gate_up[..., ::2]`, `up = gate_up[..., 1::2]` — + while the kernel's interleave wants destination row `2i` = **up** `i` and + `2i+1` = gate `i`. The halves are therefore split by parity and + re-concatenated as [up ; gate] before the permutation. Getting it + backwards is finite, plausibly scaled and invisible to a boot check. + +Loading contract: `load(model, weights)` consumes the engine-provided +per-rank checkpoint dict (safetensors lazy slices), fills every declared +parameter exactly once, and asserts full bidirectional coverage — every +target parameter written, every checkpoint key consumed. +""" + +import torch + +SCALE_BLOCK = 32 # mxfp4: one E8M0 exponent per 32 elements along K + + +def _pad_rows_cols(t: torch.Tensor, rows: int, cols: int) -> torch.Tensor: + """Zero-pad a `[E, R, C]` stack to `[E, rows, cols]`.""" + out = torch.zeros(t.shape[0], rows, cols, dtype=t.dtype, device=t.device) + out[:, : t.shape[1], : t.shape[2]] = t + return out + + +def _pad_cols(t: torch.Tensor, cols: int) -> torch.Tensor: + """Zero-pad a `[E, C]` stack to `[E, cols]`.""" + out = torch.zeros(t.shape[0], cols, dtype=t.dtype, device=t.device) + out[:, : t.shape[1]] = t + return out + + +def _block32_perm(rows: int, device) -> torch.Tensor: + """Gather index of the 32-row block shuffle: inside each aligned block of + 32 rows, source row `4u + v` (0<=u<8, 0<=v<4) lands at destination row + `8v + u`.""" + assert rows % 32 == 0, rows + src = torch.arange(32) + dst = (src % 4) * 8 + src // 4 + idx = torch.empty(32, dtype=torch.long) + idx[dst] = src + blocks = rows // 32 + return (idx.repeat(blocks) + torch.arange(blocks).repeat_interleave(32) * 32).to(device) + + +def _interleave_perm(rows: int, device) -> torch.Tensor: + """Gather index that interleaves the `[up | gate]` halves of a `2*I_pad` + row stack into up0, gate0, up1, gate1, ...""" + p = torch.empty(rows, dtype=torch.long) + p[0::2] = torch.arange(0, rows // 2) + p[1::2] = torch.arange(rows // 2, rows) + return p.to(device) + + +def _swizzle_scales(s: torch.Tensor) -> torch.Tensor: + """Per-expert 128x4 block-scale swizzle, viewed back as `[E, M, C]`: the + byte at `(m, c)` moves to flat offset `(m//128)*512*(C//4) + (c//4)*512 + + (m%32)*16 + ((m%128)//32)*4 + (c%4)`.""" + e, m, c = s.shape + assert m % 128 == 0 and c % 4 == 0, (m, c) + v = s.reshape(e, m // 128, 4, 32, c // 4, 4) # (e, m/128, (m%128)/32, m%32, c/4, c%4) + return v.permute(0, 1, 4, 3, 2, 5).reshape(e, m, c).contiguous() + + +def _split_gate_up(t: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """Split the checkpoint's interleaved `2I` output axis into (up, gate).""" + return t[:, 1::2], t[:, 0::2] + + +def _fc1_perm(rows: int, device) -> torch.Tensor: + """Interleave then 32-row block shuffle, composed into one gather.""" + return _interleave_perm(rows, device)[_block32_perm(rows, device)] + + +def _prep_fc1_weight(blocks: torch.Tensor, core) -> torch.Tensor: + """[E, 2I, H/32, 16] uint8 blocks -> kernel FC1 operand.""" + b = blocks.reshape(core.num_experts, 2 * core.inter, core.hidden // 2) + up, gate = _split_gate_up(b) + cols = core.fc1_k_pad // 2 + w = torch.cat( + [ + _pad_rows_cols(up, core.inter_pad, cols), + _pad_rows_cols(gate, core.inter_pad, cols), + ], + dim=1, + ) + return torch.index_select(w, 1, _fc1_perm(w.shape[1], w.device)).contiguous() + + +def _prep_fc1_scale(scales: torch.Tensor, core) -> torch.Tensor: + """[E, 2I, H/32] uint8 E8M0 scales -> kernel FC1 scale operand.""" + up, gate = _split_gate_up(scales) + cols = core.fc1_k_pad // SCALE_BLOCK + s = torch.cat( + [ + _pad_rows_cols(up, core.inter_pad, cols), + _pad_rows_cols(gate, core.inter_pad, cols), + ], + dim=1, + ) + return _swizzle_scales(torch.index_select(s, 1, _fc1_perm(s.shape[1], s.device))) + + +def _prep_fc1_bias(bias: torch.Tensor, core) -> torch.Tensor: + """[E, 2I] bf16 biases -> fp32 kernel FC1 bias (same row permutation).""" + up, gate = _split_gate_up(bias.float()) + b = torch.cat([_pad_cols(up, core.inter_pad), _pad_cols(gate, core.inter_pad)], dim=1) + return torch.index_select(b, 1, _fc1_perm(b.shape[1], b.device)).contiguous() + + +def _prep_fc2_weight(blocks: torch.Tensor, core) -> torch.Tensor: + """[E, H, I/32, 16] uint8 blocks -> kernel FC2 operand (no interleave).""" + w = _pad_rows_cols( + blocks.reshape(core.num_experts, core.hidden, core.inter // 2), + core.fc2_rows_pad, + core.inter_pad // 2, + ) + return torch.index_select(w, 1, _block32_perm(core.fc2_rows_pad, w.device)).contiguous() + + +def _prep_fc2_scale(scales: torch.Tensor, core) -> torch.Tensor: + """[E, H, I/32] uint8 E8M0 scales -> kernel FC2 scale operand.""" + s = _pad_rows_cols(scales, core.fc2_rows_pad, core.inter_pad // SCALE_BLOCK) + return _swizzle_scales(torch.index_select(s, 1, _block32_perm(core.fc2_rows_pad, s.device))) + + +def _prep_fc2_bias(bias: torch.Tensor, core) -> torch.Tensor: + """[E, H] bf16 biases -> fp32 kernel FC2 bias (same row permutation).""" + b = _pad_cols(bias.float(), core.fc2_rows_pad) + return torch.index_select(b, 1, _block32_perm(core.fc2_rows_pad, b.device)).contiguous() + + +def _to_fp32(t: torch.Tensor, core) -> torch.Tensor: + """bf16 -> fp32 promotion for tensors the kernels require in fp32.""" + return t.float() + + +def _manifest(core) -> dict: + """target param key -> list of (ckpt key, index into the param | None, + source transform | None).""" + q_width = core.heads_q * core.head_dim + kv_width = core.heads_kv * core.head_dim + rows: dict = {} + for i in range(core.num_layers): + p = f"model.layers.{i}" + rows[f"l{i}_norm1"] = [(f"{p}.input_layernorm.weight", None, None)] + rows[f"l{i}_qkv"] = [ + (f"{p}.self_attn.q_proj.weight", (slice(0, q_width),), None), + ( + f"{p}.self_attn.k_proj.weight", + (slice(q_width, q_width + kv_width),), + None, + ), + ( + f"{p}.self_attn.v_proj.weight", + (slice(q_width + kv_width, q_width + 2 * kv_width),), + None, + ), + ] + rows[f"l{i}_qkv_bias"] = [ + (f"{p}.self_attn.q_proj.bias", (slice(0, q_width),), None), + (f"{p}.self_attn.k_proj.bias", (slice(q_width, q_width + kv_width),), None), + ( + f"{p}.self_attn.v_proj.bias", + (slice(q_width + kv_width, q_width + 2 * kv_width),), + None, + ), + ] + rows[f"l{i}_sinks"] = [(f"{p}.self_attn.sinks", None, _to_fp32)] + rows[f"l{i}_o"] = [(f"{p}.self_attn.o_proj.weight", None, None)] + rows[f"l{i}_o_bias"] = [(f"{p}.self_attn.o_proj.bias", None, None)] + rows[f"l{i}_norm2"] = [(f"{p}.post_attention_layernorm.weight", None, None)] + rows[f"l{i}_router"] = [(f"{p}.mlp.router.weight", None, None)] + rows[f"l{i}_router_bias"] = [(f"{p}.mlp.router.bias", None, None)] + rows[f"l{i}_fc1_w"] = [(f"{p}.mlp.experts.gate_up_proj_blocks", None, _prep_fc1_weight)] + rows[f"l{i}_fc1_s"] = [(f"{p}.mlp.experts.gate_up_proj_scales", None, _prep_fc1_scale)] + rows[f"l{i}_fc1_b"] = [(f"{p}.mlp.experts.gate_up_proj_bias", None, _prep_fc1_bias)] + rows[f"l{i}_fc2_w"] = [(f"{p}.mlp.experts.down_proj_blocks", None, _prep_fc2_weight)] + rows[f"l{i}_fc2_s"] = [(f"{p}.mlp.experts.down_proj_scales", None, _prep_fc2_scale)] + rows[f"l{i}_fc2_b"] = [(f"{p}.mlp.experts.down_proj_bias", None, _prep_fc2_bias)] + rows["final_norm"] = [("model.norm.weight", None, None)] + rows["embed"] = [("model.embed_tokens.weight", None, None)] + return rows + + +def load(model, weights) -> None: + core = model.model + manifest = _manifest(core) + consumed: set = set() + + def fill(param: torch.nn.Parameter, ckpt_key: str, index, transform) -> None: + assert ckpt_key in weights, f"checkpoint key missing: {ckpt_key}" + src = weights[ckpt_key][:] # materialize the lazy slice + dst = param.data if index is None else param.data[index] + if transform is not None: + # The expert transforms are heavy row gathers over ~1 GB of + # blocks; run them where the destination lives. + src = transform(src.to(dst.device, non_blocking=True), core) + assert dst.shape == src.shape, (ckpt_key, tuple(dst.shape), tuple(src.shape)) + assert src.dtype == dst.dtype, (ckpt_key, src.dtype, dst.dtype) + dst.copy_(src, non_blocking=True) + consumed.add(ckpt_key) + + assert set(manifest.keys()) == set(core.w.keys()), ( + "manifest/parameter drift", + set(manifest.keys()) ^ set(core.w.keys()), + ) + for param_key, sources in manifest.items(): + for ckpt_key, index, transform in sources: + fill(core.w[param_key], ckpt_key, index, transform) + # The transforms allocate several GB of scratch per layer; release it + # before the next one so peak load memory stays one layer deep. + if param_key.endswith("_fc2_b"): + torch.cuda.empty_cache() + + # Shell-registered exception: the base class owns lm_head (untied). + fill(model.lm_head.weight, "lm_head.weight", None, None) + + torch.cuda.synchronize() + leftover = set(weights.keys()) - consumed + assert not leftover, f"unconsumed checkpoint keys: {sorted(leftover)[:8]}" diff --git a/tensorrt_llm/_torch/staircase/references/accuracy.yaml b/tensorrt_llm/_torch/staircase/references/accuracy.yaml new file mode 100644 index 000000000000..57c1a3c18842 --- /dev/null +++ b/tensorrt_llm/_torch/staircase/references/accuracy.yaml @@ -0,0 +1,190 @@ +# Accuracy reference scores, keyed by Hugging Face repo name; target_ckpt +# joins an entry to the segment of a target's identity path, +# models//targets////. +# +# Gate: measured accuracy >= score - tol, one-sided, enforced by +# trtllm-eval --check_accuracy. tol is 5.0 everywhere — a reference score +# drifts a few points across checkpoint variants, harnesses and +# quantizations, while an assembly catastrophe sits 20+ points away +# (random = 25 on MMLU). +# +# source, preferred first: +# `trtllm` — TensorRT-LLM's accuracy suite +# (tests/integration/defs/accuracy/). Best source: the +# harness we delegate to, and the test that produced the +# score also states its protocol (see `eval`). A +# checkpoint lists one entry per quantization — take the +# one this target matches, not the first line. +# `model-card` — the checkpoint's published number. +# `trtllm-eval` — measured here, e.g. stock trtllm on the same +# checkpoint when nothing is published. +# After a target's first passing run, write its measured score back and +# set source to `trtllm-eval`; tol stays 5.0. Skip the write-back when the +# target was accepted with a protocol caveat (no passing value exists), or +# when the anchor is already this repo's own stock-trtllm measurement — +# replacing that with the target's number makes the gate self-referential. +# +# `eval` (optional) — the protocol, so our number and the anchor measure +# the same thing. Absent = trtllm-eval defaults: 5-shot completion MMLU +# scored on the first generated token. That fits a base model; an instruct +# model without its chat template, or a reasoning model with no room to +# reason, scores far below its ability and the gate then reads as an +# assembly defect. Copy the protocol from the test that gave the score. +# +# eval: +# task: gsm8k # trtllm-eval subcommand; default mmlu +# args: # passed through by name to that subcommand +# apply_chat_template: true +# fewshot_as_multiturn: true +# max_output_length: 8192 +# +# args pass through by name: any option the subcommand accepts (booleans +# become flags, mappings are JSON-encoded). +# +# `scores_filter` is a sibling of `args`, not one of them: it names which +# lm-eval filter the gate reads (e.g. `exact_match,flexible-extract`). +# trtllm-eval exposes no CLI flag for it — it is a keyword of the +# evaluator's evaluate() — so a caller has to bind it before delegating. +# Set it whenever the anchor's protocol names one. Unset, the evaluator +# averages every metric the task reports, which mixes filters measuring +# different things and, on a task whose result dict carries a string +# alias, raises instead of returning a score. A reasoning model needs it: +# its answer never arrives in the strict `#### N` form, so the strict +# filter scores it near the floor while flexible extraction reads what it +# actually answered. +# +# ── THE SCORES BELOW ARE ANCHORS, NOT GATE RECORDS ──────────────────────── +# +# Every `score` here that says `source: trtllm-eval` was written back from a +# passing run on **sm_100 (B200)**, through a standalone harness that did not +# survive the move in-tree. The migrated targets are sm_103 (GB300) and have +# no passing run yet. These values are still the right bar to gate against — +# an accuracy anchor is a property of the checkpoint and the protocol, not of +# the GPU — but no target in this tree has cleared one. See each TARGET.md. +# +# The replacement invocation is trtllm-eval directly, with the protocol below +# spelled out on the command line and `TRTLLM_STAIRCASE=require` exported, so +# that a configuration which does not match a target fails loudly instead of +# quietly measuring the built-in implementation. +# ────────────────────────────────────────────────────────────────────────── + +# TensorRT-LLM's accuracy suite, references/gsm8k.yaml key +# `openai/gpt-oss-120b`, read at commit 1662a877f374ee944d1907e8efa735d35ff2abf6 +# whose tensorrt_llm/version.py is 1.3.0rc21 — the pin exactly. The +# quantization-keyed twin `GPT-OSS/120B-MXFP4` carries the same 90.3, +# including its `quant_algo: W4A16_MXFP4` row, which is this checkpoint's +# quantization. The suite has no MMLU entry for gpt-oss at all and gates +# the family on GSM8K only: the model reasons before it answers, and +# 5-shot completion MMLU scored on the first generated token measures +# none of that. +# Protocol copied from TestGPTOSS (test_llm_api_pytorch.py), whose +# MODEL_PATH is this exact checkpoint: chat template, few-shot as +# multiturn, MAX_OUTPUT_LEN patched to 8192, and flexible answer +# extraction. Stock trtllm measured here under precisely that protocol +# scored 90.22 (strict-match, for contrast: 28.13) — the anchor +# reproduces on this machine, so it was gated against directly rather +# than through a stock-trtllm proxy. +# +# The score below is the target's own first passing measurement, 90.5989 +# over the full 1319 questions, which replaces the 90.3 it was gated +# against. Measured twice by separate invocations of `bench/accuracy.py` +# — the assembler's gating run and an independent orchestrator re-run — +# reproducing bit-identically, down to the non-gated strict-match filter +# (25.4738). Read the +0.30 over the external anchor and the +0.38 over +# the stock-trtllm number as agreement, not as an improvement: the same +# build measured 90.22 through a differently-invoked driver, so +# differences of this size on this benchmark are not results. +# +# That measurement is of the assembly's W4A16 MoE forward. The perf +# campaign then shipped the W4A8 (MXFP8-activation) member of the same +# kernel family, which gates at 89.1585 — also bit-identical across two +# invocations, so the 1.44-point drop is the recipe's price and not +# sampling. The anchor is deliberately left at the higher W4A16 value: +# lowering it would only loosen the gate, and 90.5989 remains a real +# measurement of this checkpoint through this harness. Open thread for +# whoever revisits: stock trtllm runs the same W4A8 recipe and measured +# 90.2199 here, so our W4A8 sits 1.06 below it — larger than this gate's +# 0.38-point cross-driver spread. +openai/gpt-oss-120b: + target_ckpt: gpt-oss-120b + score: 90.5989 + n_samples: 1319 + source: trtllm-eval + tol: 5.0 + eval: + task: gsm8k + scores_filter: exact_match,flexible-extract + args: + apply_chat_template: true + fewshot_as_multiturn: true + max_output_length: 8192 + +# TensorRT-LLM's accuracy suite, references/gsm8k.yaml key +# `deepseek-ai/DeepSeek-R1-0528`, the `quant_algo: NVFP4` + +# `kv_cache_quant_algo: FP8` row = 94.24, read at commit +# 1662a877f374ee944d1907e8efa735d35ff2abf6 whose tensorrt_llm/version.py +# is 1.3.0rc21 — the pin exactly. That row is the checkpoint's exact +# quantization pair: this target's hf_quant_config.json declares +# `quant_algo: NVFP4` with `kv_cache_quant_algo: FP8`, so it is the row +# to take rather than the block's first line by position (they coincide +# here). +# +# Why GSM8K and not MMLU. references/mmlu.yaml does carry a +# `deepseek-ai/DeepSeek-R1-0528` block, but it holds **only +# FP8_BLOCK_SCALES rows — there is no NVFP4 row at all**, so this +# checkpoint's quantization has no MMLU anchor to be gated against. GSM8K +# is also the right task on its own terms: R1-0528 is a reasoning model, +# and 5-shot completion MMLU scored on the first generated token measures +# none of that. +# +# Protocol: the consuming test is TestModelRegistryAccuracy's +# `test_autodeploy_from_registry` (test_llm_api_autodeploy.py), whose +# `nvidia/DeepSeek-R1-0528-NVFP4-v2` parameter is aliased onto this key +# and whose task list is `[GSM8K]`. It passes `evaluate_kwargs = {}` for +# every non-MMLU task, so the measurement is accuracy_core.GSM8K's +# defaults — full 1319 questions, random_seed 0, **no chat template**, no +# fewshot-as-multiturn, max_input_length 4096, max_output_length 256. +# trtllm-eval's `gsm8k` subcommand defaults to exactly those, so no +# `args` block is needed. Being a reasoning model does not change that: +# under plain few-shot completion the model continues the exemplars' +# form rather than opening a think block, which is what makes 94.24 +# reachable inside 256 output tokens. +# +# `scores_filter` is this repo's one deliberate deviation, for the reason +# the file header gives — the suite leaves it None and averages every +# metric, which this harness cannot do. Same choice as the gpt-oss-120b +# and DeepSeek-V3-Lite entries. Report both filters on the first run. +# +# CAVEAT, and the reason this entry is worth re-reading at write-back +# time: the anchor was measured on a **different NVFP4 export of the same +# base model**. The test's `nvidia/DeepSeek-R1-0528-NVFP4-v2` resolves +# (tests/test_common/llm_data.py:44) to +# `DeepSeek-R1/DeepSeek-R1-0528-FP4-v2`, while this target is built from +# `DeepSeek-R1/DeepSeek-R1-0528-FP4` — the v1 export, which sits beside +# it on disk. Both declare NVFP4 + FP8 KV, but their modelopt +# `exclude_modules` lists differ substantially (63 entries here against +# 246 in v2), so they are not the same weights. tol 5.0 is what absorbs +# an export-revision difference of this kind; a measured score a point or +# two off 94.24 is not evidence of an assembly defect. +# +# WRITE-BACK. The score below is the target's own first passing +# measurement, 94.9962 over the full 1319 questions, replacing the 94.24 it +# was gated against. `strict-match` on the same run was 94.6171; the 0.38 +# between the filters is 5 questions, so under few-shot completion this +# checkpoint mostly answers in the strict `#### N` form and flexible +# extraction finds a few more. +# +# Read the +0.76 over the external anchor as agreement, not as an +# improvement. One filter's stderr alone is +-0.60, and the anchor was +# measured on the other export — the caveat above is exactly why a +# difference of this size carries no information. Nothing in the target's +# record rests on it. +deepseek-ai/DeepSeek-R1-0528: + target_ckpt: deepseek-r1-0528-nvfp4 + score: 94.9962 + n_samples: 1319 + source: trtllm-eval + tol: 5.0 + eval: + task: gsm8k + scores_filter: exact_match,flexible-extract From 8faab78f873cc3bdd876a2f2fd622cfa80d75d56 Mon Sep 17 00:00:00 2001 From: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> Date: Thu, 10 Sep 2026 08:15:09 -0700 Subject: [PATCH 02/19] [TRTLLM-16304][test] Wire the staircase catalog tests into the unit test suite The catalog's GPU tests were living beside their contracts and wrappers inside tensorrt_llm/_torch/staircase/catalog/, which is where the upstreaming plan put them -- but nothing collects tests from the package. The two in-tree precedents for a test inside tensorrt_llm/ (cute_dsl_kernels/test_argmax.py, kv_cache_manager_v2/rawref/test_rawref.py) are on no list either. So they move to tests/unittest/_torch/staircase/, which is the tree CI collects from. Named test_staircase_.py rather than test_.py for two reasons. Three would otherwise collide with an upstream test of the same basename, and under pytest's default import mode a duplicate basename is an import error. More importantly they are not duplicates of those tests: upstream covers the modules that wrap these ops, these cover the op itself cell by cell -- the strategy x operation matrix, which strategies are bitwise identical, which combinations are silently wrong rather than loud. Distinct names keep review from reading one as a copy of the other. The two collective entries keep their rank bodies in catalog/comm/. The launcher re-execs them as `python -m` and the ranks need that package context for their relative imports; only the collected shells moved. Test list entries, and why they are on these two lists: l0_gb300.yml unittest/_torch/staircase, comm excluded l0_gb300_multi_gpus.yml unittest/_torch/staircase/comm The catalog's certification is per GPU architecture -- receipts are keyed by sm -- and these are the sm_103 lists. The collectives are on the 4-GPU list because they are certified at world size 4, which is dep4's topology. An entry now spans two trees, so the receipt rule gains a second directory to look in: a receipt is valid only if it post-dates the last write to every file of the entry, contract and wrapper and test alike. index.yaml and the README say so. Verified on GB300 (sm_103), run from tests/ with the exact path strings the lists carry: 343 passed for the single-GPU entry, 2 passed for the collective entry (each spawning its own 4-rank job). Signed-off-by: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> --- .../test_lists/test-db/l0_gb300.yml | 9 +++++++ .../test-db/l0_gb300_multi_gpus.yml | 7 +++++ .../test_staircase_flashinfer_silu_and_mul.py | 4 ++- .../test_staircase_fused_qk_norm_rope.py | 2 +- ...t_staircase_load_paged_kv_cache_for_mla.py | 5 ++-- ...rcase_mla_rope_append_paged_kv_assign_q.py | 5 ++-- .../test_staircase_mla_rope_generation.py | 3 +-- .../test_staircase_thop_attention.py | 3 +-- .../test_staircase_allgather_op_matrix.py | 2 +- .../test_staircase_reducescatter_op_matrix.py | 2 +- .../staircase/gemm/test_staircase_bmm_out.py | 2 +- .../gemm/test_staircase_cublas_mm.py | 2 +- .../gemm/test_staircase_nvfp4_gemm.py | 3 +-- ...st_staircase_fp4_block_scale_moe_runner.py | 5 ++-- .../staircase/moe/test_staircase_fused_moe.py | 3 +-- ...se_mxe4m3_mxe2m1_block_scale_moe_runner.py | 4 ++- .../moe/test_staircase_noaux_tc_op.py | 2 +- ..._staircase_flashinfer_fused_add_rmsnorm.py | 4 ++- .../norm/test_staircase_flashinfer_rmsnorm.py | 2 +- .../test_staircase_fp4_quantize.py | 3 +-- .../test_staircase_mxfp8_quantize.py | 2 +- .../_torch/staircase/test_staircase_claims.py | 27 +++++++++++++++---- .../staircase/test_staircase_routing.py | 2 +- 23 files changed, 70 insertions(+), 33 deletions(-) rename tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul_test.py => tests/unittest/_torch/staircase/activation/test_staircase_flashinfer_silu_and_mul.py (94%) rename tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope_test.py => tests/unittest/_torch/staircase/attention/test_staircase_fused_qk_norm_rope.py (98%) rename tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla_test.py => tests/unittest/_torch/staircase/attention/test_staircase_load_paged_kv_cache_for_mla.py (99%) rename tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q_test.py => tests/unittest/_torch/staircase/attention/test_staircase_mla_rope_append_paged_kv_assign_q.py (99%) rename tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation_test.py => tests/unittest/_torch/staircase/attention/test_staircase_mla_rope_generation.py (99%) rename tensorrt_llm/_torch/staircase/catalog/attention/thop_attention_test.py => tests/unittest/_torch/staircase/attention/test_staircase_thop_attention.py (99%) rename tensorrt_llm/_torch/staircase/catalog/comm/test_allgather_op_matrix.py => tests/unittest/_torch/staircase/comm/test_staircase_allgather_op_matrix.py (93%) rename tensorrt_llm/_torch/staircase/catalog/comm/test_reducescatter_op_matrix.py => tests/unittest/_torch/staircase/comm/test_staircase_reducescatter_op_matrix.py (93%) rename tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out_test.py => tests/unittest/_torch/staircase/gemm/test_staircase_bmm_out.py (97%) rename tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm_test.py => tests/unittest/_torch/staircase/gemm/test_staircase_cublas_mm.py (98%) rename tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm_test.py => tests/unittest/_torch/staircase/gemm/test_staircase_nvfp4_gemm.py (99%) rename tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner_test.py => tests/unittest/_torch/staircase/moe/test_staircase_fp4_block_scale_moe_runner.py (99%) rename tensorrt_llm/_torch/staircase/catalog/moe/fused_moe_test.py => tests/unittest/_torch/staircase/moe/test_staircase_fused_moe.py (99%) rename tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner_test.py => tests/unittest/_torch/staircase/moe/test_staircase_mxe4m3_mxe2m1_block_scale_moe_runner.py (99%) rename tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op_test.py => tests/unittest/_torch/staircase/moe/test_staircase_noaux_tc_op.py (99%) rename tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm_test.py => tests/unittest/_torch/staircase/norm/test_staircase_flashinfer_fused_add_rmsnorm.py (96%) rename tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm_test.py => tests/unittest/_torch/staircase/norm/test_staircase_flashinfer_rmsnorm.py (96%) rename tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize_test.py => tests/unittest/_torch/staircase/quantization/test_staircase_fp4_quantize.py (99%) rename tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize_test.py => tests/unittest/_torch/staircase/quantization/test_staircase_mxfp8_quantize.py (99%) rename tensorrt_llm/_torch/staircase/_claim_test.py => tests/unittest/_torch/staircase/test_staircase_claims.py (84%) rename tensorrt_llm/_torch/staircase/_routing_test.py => tests/unittest/_torch/staircase/test_staircase_routing.py (99%) diff --git a/tests/integration/test_lists/test-db/l0_gb300.yml b/tests/integration/test_lists/test-db/l0_gb300.yml index 40c08cb821a6..325fa815dd0e 100644 --- a/tests/integration/test_lists/test-db/l0_gb300.yml +++ b/tests/integration/test_lists/test-db/l0_gb300.yml @@ -26,3 +26,12 @@ l0_gb300: - accuracy/test_disaggregated_serving.py::TestQwen3_8_Flash_Next::test_fp8_nixl_python[prefix_cache] - unittest/_torch/thop/parallel TIMEOUT (90) - unittest/_torch/visual_gen/kernels/parallel + # Staircase catalog: the certification matrix for each op a staircase target + # calls, plus the two consistency tests that guard the routing tables against + # the targets they name. Deliberately NOT a duplicate of the upstream tests + # for the same ops -- those cover the modules that wrap them, these cover the + # op itself cell by cell -- which is why every file is named test_staircase_*. + # These entries are the receipts: the catalog's certification is per GPU + # architecture, and this list is the sm_103 one. comm/ is excluded here and + # carried by l0_gb300_multi_gpus.yml, since those two need 4 ranks. + - unittest/_torch/staircase --ignore=unittest/_torch/staircase/comm diff --git a/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml b/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml index 0a869488e0ad..c54ce17f30a0 100644 --- a/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml +++ b/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml @@ -15,6 +15,13 @@ l0_gb300_multi_gpus: stage: post_merge backend: pytorch tests: + # Staircase collective entries. Each spawns its own 4-rank job rather than + # using the mpi_pool_executor fixture: their cases read module-global rank + # state and several assert on communicator state the previous case left + # behind, so they are one ordered sequence inside one job, and a broken + # collective hangs rather than raising -- the launcher's deadline is what + # keeps that from wedging the run, and the fixture has none. + - unittest/_torch/staircase/comm # Covers tests/unittest/_torch/attention/. Two sub-trees moved in here from elsewhere # under tests/unittest/_torch/ and have no entry of their own on any list, so this entry # is what picks them up: diff --git a/tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul_test.py b/tests/unittest/_torch/staircase/activation/test_staircase_flashinfer_silu_and_mul.py similarity index 94% rename from tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul_test.py rename to tests/unittest/_torch/staircase/activation/test_staircase_flashinfer_silu_and_mul.py index af19e2ef7d1e..5b2c9c5302ef 100644 --- a/tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul_test.py +++ b/tests/unittest/_torch/staircase/activation/test_staircase_flashinfer_silu_and_mul.py @@ -5,7 +5,9 @@ import torch import torch.nn.functional as F -from .flashinfer_silu_and_mul import flashinfer_silu_and_mul +from tensorrt_llm._torch.staircase.catalog.activation.flashinfer_silu_and_mul import ( + flashinfer_silu_and_mul, +) assert torch.cuda.is_available(), "flashinfer_silu_and_mul requires a CUDA device" diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope_test.py b/tests/unittest/_torch/staircase/attention/test_staircase_fused_qk_norm_rope.py similarity index 98% rename from tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope_test.py rename to tests/unittest/_torch/staircase/attention/test_staircase_fused_qk_norm_rope.py index 66126930e796..62e8c5462490 100644 --- a/tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope_test.py +++ b/tests/unittest/_torch/staircase/attention/test_staircase_fused_qk_norm_rope.py @@ -4,7 +4,7 @@ import torch -from .fused_qk_norm_rope import fused_qk_norm_rope +from tensorrt_llm._torch.staircase.catalog.attention.fused_qk_norm_rope import fused_qk_norm_rope assert torch.cuda.is_available(), "fused_qk_norm_rope requires a CUDA device" diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla_test.py b/tests/unittest/_torch/staircase/attention/test_staircase_load_paged_kv_cache_for_mla.py similarity index 99% rename from tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla_test.py rename to tests/unittest/_torch/staircase/attention/test_staircase_load_paged_kv_cache_for_mla.py index 26405494c7bf..9eea002561ae 100644 --- a/tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla_test.py +++ b/tests/unittest/_torch/staircase/attention/test_staircase_load_paged_kv_cache_for_mla.py @@ -35,13 +35,14 @@ from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata from tensorrt_llm._torch.metadata import KVCacheParams from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager +from tensorrt_llm._torch.staircase.catalog.attention.load_paged_kv_cache_for_mla import ( + load_paged_kv_cache_for_mla, +) from tensorrt_llm.bindings import DataType from tensorrt_llm.bindings.internal.batch_manager import CacheType from tensorrt_llm.llmapi.llm_args import KvCacheConfig from tensorrt_llm.mapping import Mapping -from .load_paged_kv_cache_for_mla import load_paged_kv_cache_for_mla - assert torch.cuda.is_available(), "load_paged_kv_cache_for_mla requires a CUDA device" # DeepSeek-V3 MLA latent geometry. diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q_test.py b/tests/unittest/_torch/staircase/attention/test_staircase_mla_rope_append_paged_kv_assign_q.py similarity index 99% rename from tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q_test.py rename to tests/unittest/_torch/staircase/attention/test_staircase_mla_rope_append_paged_kv_assign_q.py index 6077938c440c..1b60f2bc1f97 100644 --- a/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q_test.py +++ b/tests/unittest/_torch/staircase/attention/test_staircase_mla_rope_append_paged_kv_assign_q.py @@ -40,13 +40,14 @@ from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata from tensorrt_llm._torch.metadata import KVCacheParams from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager +from tensorrt_llm._torch.staircase.catalog.attention.mla_rope_append_paged_kv_assign_q import ( + mla_rope_append_paged_kv_assign_q, +) from tensorrt_llm.bindings import DataType from tensorrt_llm.bindings.internal.batch_manager import CacheType from tensorrt_llm.llmapi.llm_args import KvCacheConfig from tensorrt_llm.mapping import Mapping -from .mla_rope_append_paged_kv_assign_q import mla_rope_append_paged_kv_assign_q - assert torch.cuda.is_available(), "mla_rope_append_paged_kv_assign_q requires a CUDA device" # DeepSeek-V3 MLA head geometry (num_heads reduced to a TP-slice-like 16). diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation_test.py b/tests/unittest/_torch/staircase/attention/test_staircase_mla_rope_generation.py similarity index 99% rename from tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation_test.py rename to tests/unittest/_torch/staircase/attention/test_staircase_mla_rope_generation.py index cead752aad1d..1e2bd8055d7c 100644 --- a/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation_test.py +++ b/tests/unittest/_torch/staircase/attention/test_staircase_mla_rope_generation.py @@ -52,13 +52,12 @@ from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata from tensorrt_llm._torch.metadata import KVCacheParams from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager +from tensorrt_llm._torch.staircase.catalog.attention.mla_rope_generation import mla_rope_generation from tensorrt_llm.bindings import DataType from tensorrt_llm.bindings.internal.batch_manager import CacheType from tensorrt_llm.llmapi.llm_args import KvCacheConfig from tensorrt_llm.mapping import Mapping -from .mla_rope_generation import mla_rope_generation - assert torch.cuda.is_available(), "mla_rope_generation requires a CUDA device" # DeepSeek-V3 MLA head geometry (num_heads reduced to a TP-slice-like 16). diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/thop_attention_test.py b/tests/unittest/_torch/staircase/attention/test_staircase_thop_attention.py similarity index 99% rename from tensorrt_llm/_torch/staircase/catalog/attention/thop_attention_test.py rename to tests/unittest/_torch/staircase/attention/test_staircase_thop_attention.py index 7fc5375d5d29..23fcaecaf3b9 100644 --- a/tensorrt_llm/_torch/staircase/catalog/attention/thop_attention_test.py +++ b/tests/unittest/_torch/staircase/attention/test_staircase_thop_attention.py @@ -174,12 +174,11 @@ from tensorrt_llm._torch.attention.backends.interface import RopeParams from tensorrt_llm._torch.pyexecutor.resource_manager import CacheTypeCpp, DataType, KVCacheManager +from tensorrt_llm._torch.staircase.catalog.attention.thop_attention import thop_attention from tensorrt_llm.functional import RotaryScalingType from tensorrt_llm.llmapi.llm_args import KvCacheConfig from tensorrt_llm.mapping import Mapping -from .thop_attention import thop_attention - assert torch.cuda.is_available(), "thop_attention requires a CUDA device" # The output is a softmax-weighted combination of unit-scale bf16 v rows; diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/test_allgather_op_matrix.py b/tests/unittest/_torch/staircase/comm/test_staircase_allgather_op_matrix.py similarity index 93% rename from tensorrt_llm/_torch/staircase/catalog/comm/test_allgather_op_matrix.py rename to tests/unittest/_torch/staircase/comm/test_staircase_allgather_op_matrix.py index a1af97045f6d..04aec3e02270 100644 --- a/tensorrt_llm/_torch/staircase/catalog/comm/test_allgather_op_matrix.py +++ b/tests/unittest/_torch/staircase/comm/test_staircase_allgather_op_matrix.py @@ -15,7 +15,7 @@ import torch -from . import _rank_job +from tensorrt_llm._torch.staircase.catalog.comm import _rank_job assert torch.cuda.is_available(), "allgather requires CUDA devices" diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/test_reducescatter_op_matrix.py b/tests/unittest/_torch/staircase/comm/test_staircase_reducescatter_op_matrix.py similarity index 93% rename from tensorrt_llm/_torch/staircase/catalog/comm/test_reducescatter_op_matrix.py rename to tests/unittest/_torch/staircase/comm/test_staircase_reducescatter_op_matrix.py index b4db607ac5b7..ce4121bd68a7 100644 --- a/tensorrt_llm/_torch/staircase/catalog/comm/test_reducescatter_op_matrix.py +++ b/tests/unittest/_torch/staircase/comm/test_staircase_reducescatter_op_matrix.py @@ -15,7 +15,7 @@ import torch -from . import _rank_job +from tensorrt_llm._torch.staircase.catalog.comm import _rank_job assert torch.cuda.is_available(), "reducescatter requires CUDA devices" diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out_test.py b/tests/unittest/_torch/staircase/gemm/test_staircase_bmm_out.py similarity index 97% rename from tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out_test.py rename to tests/unittest/_torch/staircase/gemm/test_staircase_bmm_out.py index af91668f9622..87f000b87c68 100644 --- a/tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out_test.py +++ b/tests/unittest/_torch/staircase/gemm/test_staircase_bmm_out.py @@ -4,7 +4,7 @@ import torch -from .bmm_out import bmm_out +from tensorrt_llm._torch.staircase.catalog.gemm.bmm_out import bmm_out assert torch.cuda.is_available(), "bmm_out requires a CUDA device" diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm_test.py b/tests/unittest/_torch/staircase/gemm/test_staircase_cublas_mm.py similarity index 98% rename from tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm_test.py rename to tests/unittest/_torch/staircase/gemm/test_staircase_cublas_mm.py index eb1c13784950..a39191e8f59b 100644 --- a/tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm_test.py +++ b/tests/unittest/_torch/staircase/gemm/test_staircase_cublas_mm.py @@ -6,7 +6,7 @@ import torch -from .cublas_mm import cublas_mm +from tensorrt_llm._torch.staircase.catalog.gemm.cublas_mm import cublas_mm assert torch.cuda.is_available(), "cublas_mm requires a CUDA device" diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm_test.py b/tests/unittest/_torch/staircase/gemm/test_staircase_nvfp4_gemm.py similarity index 99% rename from tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm_test.py rename to tests/unittest/_torch/staircase/gemm/test_staircase_nvfp4_gemm.py index 3fef8286d096..abcc7cd6132a 100644 --- a/tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm_test.py +++ b/tests/unittest/_torch/staircase/gemm/test_staircase_nvfp4_gemm.py @@ -6,8 +6,7 @@ import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* from tensorrt_llm._torch.autotuner import AutoTuner, autotune - -from .nvfp4_gemm import nvfp4_gemm +from tensorrt_llm._torch.staircase.catalog.gemm.nvfp4_gemm import nvfp4_gemm assert torch.cuda.is_available(), "nvfp4_gemm requires a CUDA device" # The reference matmul must be true fp32, never tf32. diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner_test.py b/tests/unittest/_torch/staircase/moe/test_staircase_fp4_block_scale_moe_runner.py similarity index 99% rename from tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner_test.py rename to tests/unittest/_torch/staircase/moe/test_staircase_fp4_block_scale_moe_runner.py index a0c2ef90fe67..d0c488c034d8 100644 --- a/tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner_test.py +++ b/tests/unittest/_torch/staircase/moe/test_staircase_fp4_block_scale_moe_runner.py @@ -5,8 +5,9 @@ import torch from tensorrt_llm._torch.autotuner import AutoTuner, autotune - -from .fp4_block_scale_moe_runner import fp4_block_scale_moe_runner as moe +from tensorrt_llm._torch.staircase.catalog.moe.fp4_block_scale_moe_runner import ( + fp4_block_scale_moe_runner as moe, +) assert torch.cuda.is_available(), "fp4_block_scale_moe_runner requires a CUDA device" diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/fused_moe_test.py b/tests/unittest/_torch/staircase/moe/test_staircase_fused_moe.py similarity index 99% rename from tensorrt_llm/_torch/staircase/catalog/moe/fused_moe_test.py rename to tests/unittest/_torch/staircase/moe/test_staircase_fused_moe.py index 331e9d395d97..078236f81523 100644 --- a/tensorrt_llm/_torch/staircase/catalog/moe/fused_moe_test.py +++ b/tests/unittest/_torch/staircase/moe/test_staircase_fused_moe.py @@ -6,8 +6,7 @@ import torch.nn.functional as F from tensorrt_llm._torch.autotuner import AutoTuner, autotune - -from .fused_moe import fused_moe +from tensorrt_llm._torch.staircase.catalog.moe.fused_moe import fused_moe assert torch.cuda.is_available(), "fused_moe requires a CUDA device" diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner_test.py b/tests/unittest/_torch/staircase/moe/test_staircase_mxe4m3_mxe2m1_block_scale_moe_runner.py similarity index 99% rename from tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner_test.py rename to tests/unittest/_torch/staircase/moe/test_staircase_mxe4m3_mxe2m1_block_scale_moe_runner.py index 2e39a2dea56f..6cba8a7058f1 100644 --- a/tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner_test.py +++ b/tests/unittest/_torch/staircase/moe/test_staircase_mxe4m3_mxe2m1_block_scale_moe_runner.py @@ -6,7 +6,9 @@ import torch -from .mxe4m3_mxe2m1_block_scale_moe_runner import mxe4m3_mxe2m1_block_scale_moe_runner as moe +from tensorrt_llm._torch.staircase.catalog.moe.mxe4m3_mxe2m1_block_scale_moe_runner import ( + mxe4m3_mxe2m1_block_scale_moe_runner as moe, +) assert torch.cuda.is_available(), "mxe4m3_mxe2m1_block_scale_moe_runner requires a CUDA device" diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op_test.py b/tests/unittest/_torch/staircase/moe/test_staircase_noaux_tc_op.py similarity index 99% rename from tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op_test.py rename to tests/unittest/_torch/staircase/moe/test_staircase_noaux_tc_op.py index 7aee4956c92f..e9e22f159df9 100644 --- a/tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op_test.py +++ b/tests/unittest/_torch/staircase/moe/test_staircase_noaux_tc_op.py @@ -4,7 +4,7 @@ import torch -from .noaux_tc_op import noaux_tc_op +from tensorrt_llm._torch.staircase.catalog.moe.noaux_tc_op import noaux_tc_op assert torch.cuda.is_available(), "noaux_tc_op requires a CUDA device" diff --git a/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm_test.py b/tests/unittest/_torch/staircase/norm/test_staircase_flashinfer_fused_add_rmsnorm.py similarity index 96% rename from tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm_test.py rename to tests/unittest/_torch/staircase/norm/test_staircase_flashinfer_fused_add_rmsnorm.py index 0560b9e23477..139b05951c68 100644 --- a/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm_test.py +++ b/tests/unittest/_torch/staircase/norm/test_staircase_flashinfer_fused_add_rmsnorm.py @@ -4,7 +4,9 @@ import torch -from .flashinfer_fused_add_rmsnorm import flashinfer_fused_add_rmsnorm +from tensorrt_llm._torch.staircase.catalog.norm.flashinfer_fused_add_rmsnorm import ( + flashinfer_fused_add_rmsnorm, +) assert torch.cuda.is_available(), "flashinfer_fused_add_rmsnorm requires a CUDA device" diff --git a/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm_test.py b/tests/unittest/_torch/staircase/norm/test_staircase_flashinfer_rmsnorm.py similarity index 96% rename from tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm_test.py rename to tests/unittest/_torch/staircase/norm/test_staircase_flashinfer_rmsnorm.py index 18274f9189e2..c9bb7cad58c8 100644 --- a/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm_test.py +++ b/tests/unittest/_torch/staircase/norm/test_staircase_flashinfer_rmsnorm.py @@ -4,7 +4,7 @@ import torch -from .flashinfer_rmsnorm import flashinfer_rmsnorm +from tensorrt_llm._torch.staircase.catalog.norm.flashinfer_rmsnorm import flashinfer_rmsnorm assert torch.cuda.is_available(), "flashinfer_rmsnorm requires a CUDA device" diff --git a/tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize_test.py b/tests/unittest/_torch/staircase/quantization/test_staircase_fp4_quantize.py similarity index 99% rename from tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize_test.py rename to tests/unittest/_torch/staircase/quantization/test_staircase_fp4_quantize.py index c184b30853cb..3ee5d5fc884b 100644 --- a/tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize_test.py +++ b/tests/unittest/_torch/staircase/quantization/test_staircase_fp4_quantize.py @@ -6,8 +6,7 @@ from torch.profiler import ProfilerActivity, profile import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* - -from .fp4_quantize import fp4_quantize +from tensorrt_llm._torch.staircase.catalog.quantization.fp4_quantize import fp4_quantize assert torch.cuda.is_available(), "fp4_quantize requires a CUDA device" diff --git a/tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize_test.py b/tests/unittest/_torch/staircase/quantization/test_staircase_mxfp8_quantize.py similarity index 99% rename from tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize_test.py rename to tests/unittest/_torch/staircase/quantization/test_staircase_mxfp8_quantize.py index b43301ef5490..5187ae28c55c 100644 --- a/tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize_test.py +++ b/tests/unittest/_torch/staircase/quantization/test_staircase_mxfp8_quantize.py @@ -4,7 +4,7 @@ import torch -from .mxfp8_quantize import mxfp8_quantize +from tensorrt_llm._torch.staircase.catalog.quantization.mxfp8_quantize import mxfp8_quantize assert torch.cuda.is_available(), "mxfp8_quantize requires a CUDA device" diff --git a/tensorrt_llm/_torch/staircase/_claim_test.py b/tests/unittest/_torch/staircase/test_staircase_claims.py similarity index 84% rename from tensorrt_llm/_torch/staircase/_claim_test.py rename to tests/unittest/_torch/staircase/test_staircase_claims.py index fba88c277ce3..5afe0499c857 100644 --- a/tensorrt_llm/_torch/staircase/_claim_test.py +++ b/tests/unittest/_torch/staircase/test_staircase_claims.py @@ -7,6 +7,10 @@ extensions. That is deliberate: the failure mode being guarded against is a rename or a move, and those are visible without executing anything. +The paths it walks are the *package's*, resolved from the imported module +rather than from this file, because the tests live under tests/ and the tree +they guard lives under tensorrt_llm/. + What is *not* guarded here is whether a target is correct; that is what its gate records are for. """ @@ -18,9 +22,12 @@ import pytest -from ._router_index import STAIRCASE_ROUTERS, routing_module +import tensorrt_llm._torch.staircase as _staircase +from tensorrt_llm._torch.staircase._router_index import STAIRCASE_ROUTERS, routing_module -_ROOT = Path(__file__).resolve().parent +# The package, not this file: these paths address the tree under test, and this +# test lives in tests/ while that tree lives in tensorrt_llm/. +_ROOT = Path(_staircase.__file__).resolve().parent _ARCHS = sorted(STAIRCASE_ROUTERS) @@ -137,13 +144,23 @@ def test_no_routing_module_reads_an_unplumbed_dimension(arch): ) -def test_every_target_ships_the_four_products(): - """modeling.py, weights.py, smoke.py and TARGET.md travel together.""" +def test_every_target_ships_the_three_products(): + """modeling.py, weights.py and TARGET.md travel together. + + The forward, the weights it expects, and the record of what that pair was + measured to do. A target missing the third is one nobody can check. + + There used to be a fourth, a per-target ``smoke.py``. In-tree there is no + reason to carry a bespoke keyword-assert CLI: the same boot-and-generate + check is ``examples/llm-api/quickstart_advanced.py``, and the accuracy + gates are ``trtllm-eval`` and ``accuracy/test_staircase.py`` -- which, + unlike a module that nothing ever ran, CI actually runs. + """ for arch in _ARCHS: routing = routing_module(arch) for name, dotted in routing.TARGET_MODULES.items(): target_dir = _module_path(dotted).parent - for product in ("modeling.py", "weights.py", "smoke.py", "TARGET.md"): + for product in ("modeling.py", "weights.py", "TARGET.md"): assert (target_dir / product).is_file(), f"{name}: missing {product}" diff --git a/tensorrt_llm/_torch/staircase/_routing_test.py b/tests/unittest/_torch/staircase/test_staircase_routing.py similarity index 99% rename from tensorrt_llm/_torch/staircase/_routing_test.py rename to tests/unittest/_torch/staircase/test_staircase_routing.py index 41ee7e59bd16..af9c0a985b39 100644 --- a/tensorrt_llm/_torch/staircase/_routing_test.py +++ b/tests/unittest/_torch/staircase/test_staircase_routing.py @@ -29,7 +29,7 @@ _is_builtin_model_class, get_registered_model_class, ) -from ._router_index import ( +from tensorrt_llm._torch.staircase._router_index import ( STAIRCASE_ENV, StaircaseMode, staircase_resolve, From 4175d7a98d603ec4d75c7e67e529e916274cc34d Mon Sep 17 00:00:00 2001 From: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> Date: Thu, 10 Sep 2026 22:52:05 -0700 Subject: [PATCH 03/19] [TRTLLM-16304][test] Add the staircase whole-model gates to the accuracy suite The op-level entries certify the vocabulary; these certify the assemblies that call it. A separate file rather than entries in test_llm_api_pytorch.py, for the same reason the targets are separate codebases: reading them beside the built-in model's tests would invite treating one as a variant of the other. l0_gb300.yml TestStaircaseGptOss120bSm103Tp1::test_gsm8k test_staircase_off_is_the_default l0_gb300_multi_gpus.yml TestStaircaseDeepseekR10528Nvfp4Sm103Dep4:: test_gsm8k test_gsm8k_identity_vs_mtp3 test_mtp3_acceptance[staircase] test_mtp3_acceptance[stock] Two of these are shaped unlike an ordinary accuracy test, for reasons worth stating. **identity vs mtp3 is one test, not two.** Both measurements are taken in a single case on purpose. The same identity forward measured 94.7688 and 95.0720 on consecutive days -- 0.30 apart on bit-identical code -- so a delta against a score recorded in an earlier session carries that session's variance into the judgement. Paired, the variance is common to both and cancels. Three runs here: +0.3412, +0.9477, +0.3033, against a |delta| < 1.2 criterion. **Acceptance is two independent cases sharing one anchor.** Rejection sampling makes a miscomputed draft path slower rather than wrong, so smoke and accuracy are both blind to it and acceptance length is the only detector. The two legs are read against the same recorded minimum, which is what makes the pair informative: staircase failing while stock passes means the draft path regressed; both failing means the anchor is stale and should be re-derived rather than the target blamed. The anchor is populated from the stock leg -- taking it from the target's own number would make the gate self-referential, the same rule the accuracy anchors follow. min_al is hand-set to 2.5 rather than the automatic 95% of the reference. Measured over three runs each on GB300: stock 3.020/2.998/3.009/3.014, staircase 2.921/2.938/2.933. Each side is stable to ~0.7% and the ranges do not overlap, so staircase sits about 2.6% under stock consistently. The automatic floor of 2.869 would have left the staircase leg 1.8% of headroom against a 0.7% spread -- a flaky test rather than a gate. 2.5 is a collapse tripwire, the role the TestKimiK3DSpark entry already documents. Also needed: - a gsm8k.yaml reference for R1-0528 NVFP4 + FP8 KV + MTP. Same accuracy as the non-MTP entry: MTP is distribution-preserving, so it must not move the score, which is what the paired test asserts directly. - scores_filter bound to exact_match,flexible-extract. Unset, the evaluator averages the filters, and for gpt-oss that means the mean of ~90 flexible and ~25 strict -- 56.1, which reads as catastrophic failure of a model answering correctly. - an output budget of 8192 for gpt-oss. The stock 256 truncates this reasoning model mid-chain-of-thought, before it reaches an answer. Verified on GB300 (sm_103), run from tests/ with the list path strings: 1-GPU 2 passed (gsm8k 89.083 against the 90.300 reference), 4-GPU 4 passed (acceptance staircase 2.933 / stock 3.014). Signed-off-by: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> --- .../references/acceptance_length.yaml | 20 ++ .../defs/accuracy/references/gsm8k.yaml | 4 + .../defs/accuracy/test_staircase.py | 260 ++++++++++++++++++ .../test_lists/test-db/l0_gb300.yml | 4 + .../test-db/l0_gb300_multi_gpus.yml | 10 + 5 files changed, 298 insertions(+) create mode 100644 tests/integration/defs/accuracy/test_staircase.py diff --git a/tests/integration/defs/accuracy/references/acceptance_length.yaml b/tests/integration/defs/accuracy/references/acceptance_length.yaml index f9795e8ac4d7..c206436db209 100644 --- a/tests/integration/defs/accuracy/references/acceptance_length.yaml +++ b/tests/integration/defs/accuracy/references/acceptance_length.yaml @@ -96,3 +96,23 @@ disagg::TestNemotron3Super120B::test_auto_dtype: TestQwen3_8_27B::test_dflash2: ref_al: 5.899181450085933 min_al: 5.604222377581636 +# Shared by both legs of TestStaircaseDeepseekR10528Nvfp4Sm103Dep4:: +# test_mtp3_acceptance. Keyed by target and variant rather than by test +# function, because that is what the number is a property of, and because the +# staircase and stock legs are read against the same minimum on purpose: +# staircase failing while stock passes means the draft path regressed; both +# failing means this anchor is stale. Populated from the STOCK leg -- taking it +# from the target's own number would make the gate self-referential. +StaircaseDeepseekR10528Nvfp4Sm103Dep4::mtp3: + ref_al: 3.020 + # Hand-set, not the automatic 95%. Measured on GB300 over three runs each: + # stock 3.020 / 2.998 / 3.009 / 3.014, staircase 2.921 / 2.938 / 2.933. Each + # side is stable to ~0.7%, and the two ranges do not overlap -- staircase + # runs about 2.6% under stock on this workload, consistently rather than as + # noise. The automatic 95% floor (2.869) would therefore have left the + # staircase leg 1.8% of headroom against a 0.7% spread: a flaky test, not a + # gate. 2.5 is a collapse tripwire instead, the same role the + # TestKimiK3DSpark entry above states -- a collapsed draft path scores ~1.0 + # and a subtly wrong one ~2.1, both far below this, while every accuracy + # gate stays green. It leaves staircase 17% of headroom and stock 20%. + min_al: 2.5 diff --git a/tests/integration/defs/accuracy/references/gsm8k.yaml b/tests/integration/defs/accuracy/references/gsm8k.yaml index cc96711cf306..bf3bc2bb1431 100644 --- a/tests/integration/defs/accuracy/references/gsm8k.yaml +++ b/tests/integration/defs/accuracy/references/gsm8k.yaml @@ -69,6 +69,10 @@ deepseek-ai/DeepSeek-R1-0528: - quant_algo: NVFP4 kv_cache_quant_algo: FP8 accuracy: 94.24 + - quant_algo: NVFP4 + kv_cache_quant_algo: FP8 + spec_dec_algo: MTP + accuracy: 94.24 - quant_algo: FP8_BLOCK_SCALES accuracy: 92.722 - quant_algo: FP8_BLOCK_SCALES diff --git a/tests/integration/defs/accuracy/test_staircase.py b/tests/integration/defs/accuracy/test_staircase.py new file mode 100644 index 000000000000..4520a51e82f0 --- /dev/null +++ b/tests/integration/defs/accuracy/test_staircase.py @@ -0,0 +1,260 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Whole-model gates for the staircase targets. + +Separate file rather than entries in test_llm_api_pytorch.py, for the same +reason the targets are separate codebases: these gate a parallel +implementation, and reading them next to the built-in model's tests would +invite treating one as a variant of the other. + +Every test here needs ``TRTLLM_STAIRCASE=require``. Under ``"auto"`` a +configuration that missed a target's criteria would quietly fall back to the +built-in implementation, pass, and report the built-in's numbers as the +target's -- which is the one failure this whole system exists to prevent. The +one exception is the stock leg of the acceptance gate, which asks for +``"off"`` on purpose. + +The switch is an environment variable, and worker ranks read it as it stood +when they started. At world size 1 that is this process. Above it the ranks +are already running by the time a test body executes, so a multi-rank case +cannot choose its own mode -- it can only assert that the environment it was +given is the one it needs, which is what ``_require_mode`` does. +""" + +import os + +import pytest + +from tensorrt_llm import LLM +from tensorrt_llm._torch.staircase import STAIRCASE_ENV +from tensorrt_llm._utils import get_sm_version +from tensorrt_llm.llmapi import CudaGraphConfig, KvCacheConfig, MTPDecodingConfig + +from ..conftest import llm_models_root +from .accuracy_core import ( + GSM8K, + LlmapiAccuracyTestHarness, + assert_acceptance_length, + compute_acceptance_length, +) + +# The targets assert their own SM at construction: certification is per GPU +# architecture, and a receipt from another one says nothing here. +skip_not_sm103 = pytest.mark.skipif( + get_sm_version() != 103, reason="staircase targets in this batch are certified on sm_103 only" +) + + +def _require_mode(expected: str) -> None: + """Skip unless the ranks were started with the mode this case needs. + + Not a failure: which mode a multi-rank job runs under is a property of how + it was launched, so a case that wants the other one has nothing to say. It + must not silently measure the wrong system either, which is what reading + the variable here rules out. + """ + actual = os.environ.get(STAIRCASE_ENV, "off") + if actual != expected: + pytest.skip(f"{STAIRCASE_ENV}={actual!r}, this case needs {expected!r}") + + +class _StaircaseGSM8K(GSM8K): + """GSM8K reading the filter the staircase anchors were measured on. + + Unset, the evaluator averages every metric the task reports, which for + GSM8K means the mean of ``strict-match`` and ``flexible-extract``. Those + measure different things here: neither of these checkpoints answers purely + in the strict ``#### N`` form, so the average is a number no reference was + ever taken at -- gpt-oss scores ~90 flexible, ~25 strict, and the mean of + 56 reads as a catastrophic failure of a model that is answering correctly. + """ + + EVALUATE_KWARGS = {"scores_filter": "exact_match,flexible-extract"} + + +class _GSM8KWithRoomToReason(_StaircaseGSM8K): + """The above, with the output budget a reasoning model needs. + + The stock 256 tokens truncate this checkpoint mid-chain-of-thought, before + it ever reaches an answer, and the gate then reads as an assembly defect + rather than as the protocol being wrong for the model. + """ + + MAX_OUTPUT_LEN = 8192 + + +class TestStaircaseGptOss120bSm103Tp1(LlmapiAccuracyTestHarness): + """gpt-oss-120b / sm_103 / tp1.""" + + # The registry key upstream uses for this checkpoint; it carries the + # W4A8_MXFP4_MXFP8 entry the engine resolves from its quantization_config. + MODEL_NAME = "GPT-OSS/120B-MXFP4" + MODEL_PATH = f"{llm_models_root()}/gpt_oss/gpt-oss-120b" + + # This checkpoint is gated as a reasoning model: its answer never arrives + # in the strict "#### N" form, so the protocol applies the chat template + # and gives the model room to reason. Matches the protocol recorded in + # _torch/staircase/references/accuracy.yaml. + extra_evaluator_kwargs = { + "apply_chat_template": True, + "fewshot_as_multiturn": True, + } + + @skip_not_sm103 + def test_gsm8k(self): + _require_mode("require") + with LLM(self.MODEL_PATH) as llm: + task = _GSM8KWithRoomToReason(self.MODEL_NAME) + task.evaluate(llm, extra_evaluator_kwargs=self.extra_evaluator_kwargs) + + +class TestStaircaseDeepseekR10528Nvfp4Sm103Dep4(LlmapiAccuracyTestHarness): + """deepseek-r1-0528-nvfp4 / sm_103 / dep4, identity and the mtp3 variant.""" + + MODEL_NAME = "deepseek-ai/DeepSeek-R1-0528" + MODEL_PATH = f"{llm_models_root()}/DeepSeek-R1/DeepSeek-R1-0528-FP4" + + # The parallel topology the path's dep4 segment declares. Routing derives + # the target *from* these, and the target then asserts every one of them + # against the mapping the engine actually built. + DEP4 = dict(tensor_parallel_size=4, moe_expert_parallel_size=4, enable_attention_dp=True) + + # configs/mtp3.yaml. The kv-cache fraction is a boot requirement of the + # variant rather than a tuning choice: the drafting forward's post-pool + # transient does not fit what the default 0.9 leaves. + MTP3 = MTPDecodingConfig(max_draft_len=3) + MTP3_KV = KvCacheConfig(free_gpu_memory_fraction=0.75) + + # One anchor shared by both legs of the acceptance gate below. It names the + # target and variant, not a test function, because that is what the number + # is a property of. + ACCEPTANCE_KEY = "StaircaseDeepseekR10528Nvfp4Sm103Dep4::mtp3" + + # There is deliberately no standalone identity gsm8k case. The paired test + # below evaluates the identity config as its first leg, and + # ``task.evaluate`` asserts accuracy against the reference on the way past, + # so a separate one would gate nothing new and would cost a fifth engine + # boot of a 61-layer, 4-rank model. + + @skip_not_sm103 + @pytest.mark.skip_less_device(4) + def test_gsm8k_identity_vs_mtp3(self): + """The identity accuracy gate, and the gate on MTP not moving it. + + Turning MTP on must not move the answers. + + Rejection sampling holds the emitted distribution to the target + model's, so the two scores should differ only by sampling noise. + + Both legs assert accuracy against the registered reference as they + run -- this test is therefore the identity gate as well as the + comparison. + + Both measurements are taken **in this one test** on purpose. The same + identity forward measured 94.7688 and 95.0720 on consecutive days -- + 0.30 apart, on bit-identical code -- so a delta against a score + recorded in some earlier session carries that session's variance into + the judgement. Paired, the variance is common to both and cancels. + """ + task = _StaircaseGSM8K(self.MODEL_NAME) + + _require_mode("require") + with LLM(self.MODEL_PATH, **self.DEP4) as llm: + identity = task.evaluate(llm) + + with LLM( + self.MODEL_PATH, + speculative_config=self.MTP3, + kv_cache_config=self.MTP3_KV, + **self.DEP4, + ) as llm: + mtp3 = task.evaluate(llm) + + delta = mtp3 - identity + print(f"[staircase] gsm8k identity={identity:.4f} mtp3={mtp3:.4f} delta={delta:+.4f}") + # 2 sigma at the ~0.6 stderr this benchmark reports at n=1319. + assert abs(delta) < 1.2, ( + f"MTP moved gsm8k by {delta:+.4f} (identity={identity:.4f}, " + f"mtp3={mtp3:.4f}); rejection sampling should have held the " + f"distribution, so this is not sampling noise" + ) + + @skip_not_sm103 + @pytest.mark.skip_less_device(4) + @pytest.mark.parametrize("mode", ["require", "off"], ids=["staircase", "stock"]) + def test_mtp3_acceptance(self, mode): + """The only gate that can see a miscomputed draft layer. + + Rejection sampling makes a wrong draft path *slower*, not wrong: every + draft is rejected, the text stays correct, and the boot and accuracy + gates both pass. Acceptance length is the sole detector. + + Two independent cases rather than one that compares them in-session. + They share ``ACCEPTANCE_KEY``, so both are read against the same + recorded minimum, which is what makes the pair informative: + + staircase fails, stock passes -> the draft path regressed + both fail -> the anchor is stale; re-derive it + rather than blaming the target + + The anchor is populated from the **stock** leg. Populating it from the + target's own number would make the gate self-referential, which is the + same rule references/accuracy.yaml states for accuracy anchors. + """ + _require_mode(mode) + if mode == "off": + # Stock cannot boot this checkpoint at dep4 with MTP otherwise: + # under attention DP + EP the MoE communication factory lands on + # DeepEPLowLatency, whose dispatch takes only NVFP4 uint8 hidden + # states, and the MTP layer is bf16 because modelopt excludes + # model.layers.61* from quantization. Disabling DeepEP lands on + # AllGatherReduceScatter -- which is the strategy the staircase + # target implements by hand, so it makes the two comparable rather + # than less so. Set in the launching environment, like the switch. + assert os.environ.get("TRTLLM_CAN_USE_DEEP_EP") == "0", ( + "the stock leg needs TRTLLM_CAN_USE_DEEP_EP=0 exported; without " + "it stock cannot boot this checkpoint at dep4 with MTP" + ) + + with LLM( + self.MODEL_PATH, + speculative_config=self.MTP3, + kv_cache_config=self.MTP3_KV, + cuda_graph_config=CudaGraphConfig(), + enable_iter_perf_stats=True, + **self.DEP4, + ) as llm: + task = _StaircaseGSM8K(self.MODEL_NAME) + task.evaluate(llm) + acceptance_length = compute_acceptance_length(llm) + print(f"[AL] {mode} acceptance_length = {acceptance_length:.3f}") + assert_acceptance_length(self.ACCEPTANCE_KEY, acceptance_length) + + +def test_staircase_off_is_the_default(monkeypatch): + """Unset means off, on the code path the engine actually takes. + + Cheap, GPU-free, and the thing most worth never regressing: everything in + this file rests on staircase being opt-in. + """ + from tensorrt_llm._torch.staircase import StaircaseMode + + monkeypatch.delenv(STAIRCASE_ENV, raising=False) + assert StaircaseMode.from_env() is StaircaseMode.OFF + assert os.environ.get("STAIRCASE_TARGET") is None, ( + "STAIRCASE_TARGET was retired with the move in-tree; it named a " + "target, where TRTLLM_STAIRCASE names only a mode and lets routing " + "pick the target from the configuration" + ) diff --git a/tests/integration/test_lists/test-db/l0_gb300.yml b/tests/integration/test_lists/test-db/l0_gb300.yml index 325fa815dd0e..7288ae3e6811 100644 --- a/tests/integration/test_lists/test-db/l0_gb300.yml +++ b/tests/integration/test_lists/test-db/l0_gb300.yml @@ -35,3 +35,7 @@ l0_gb300: # architecture, and this list is the sm_103 one. comm/ is excluded here and # carried by l0_gb300_multi_gpus.yml, since those two need 4 ranks. - unittest/_torch/staircase --ignore=unittest/_torch/staircase/comm + # Staircase whole-model gate. The op-level entry above certifies the + # vocabulary; this certifies the assembly that calls it. + - accuracy/test_staircase.py::TestStaircaseGptOss120bSm103Tp1::test_gsm8k + - accuracy/test_staircase.py::test_staircase_off_is_the_default diff --git a/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml b/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml index c54ce17f30a0..57a7040a1d17 100644 --- a/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml +++ b/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml @@ -22,6 +22,16 @@ l0_gb300_multi_gpus: # collective hangs rather than raising -- the launcher's deadline is what # keeps that from wedging the run, and the fixture has none. - unittest/_torch/staircase/comm + # Staircase whole-model gates for the dep4 target. The mtp3 variant selects a + # second forward path AND a second weight-loading path, so the identity gate + # does not speak for it: it carries its own accuracy pairing and its own + # acceptance gate. The two acceptance legs are independent cases read against + # one shared minimum -- staircase failing while stock passes means the draft + # path regressed; both failing means the anchor is stale. + - accuracy/test_staircase.py::TestStaircaseDeepseekR10528Nvfp4Sm103Dep4::test_gsm8k + - accuracy/test_staircase.py::TestStaircaseDeepseekR10528Nvfp4Sm103Dep4::test_gsm8k_identity_vs_mtp3 + - accuracy/test_staircase.py::TestStaircaseDeepseekR10528Nvfp4Sm103Dep4::test_mtp3_acceptance[staircase] + - accuracy/test_staircase.py::TestStaircaseDeepseekR10528Nvfp4Sm103Dep4::test_mtp3_acceptance[stock] # Covers tests/unittest/_torch/attention/. Two sub-trees moved in here from elsewhere # under tests/unittest/_torch/ and have no entry of their own on any list, so this entry # is what picks them up: From c9fb8146e7a7ffb426bb0f3ba5fae167e97e6408 Mon Sep 17 00:00:00 2001 From: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> Date: Thu, 10 Sep 2026 23:58:51 -0700 Subject: [PATCH 04/19] [TRTLLM-16304][test] Keep the staircase collectives off xdist, drop a redundant gate Two follow-ups on the staircase CI wiring. **The collective entries are marked no_xdist.** Each case starts its own 4-rank mpirun over every visible device, so several xdist workers would fight for the same GPUs. The runner does retry failures serially, so this would likely still go green -- but only after a confusing parallel failure. The marker is what the other collective entry in this repo (_torch/thop/serial/test_moe_alltoall.py) uses for the same reason. Note on the parallel config: agg_unit_mem_df.csv gets no entries here. Its header says it is generated from execution results by the infra team and not manually edited, and it currently carries no GB300 rows at all, so every GB300 unittest entry already falls back to serial with a warning. That is a pre-existing condition for that list rather than something these entries introduce, and hand-writing rows would contradict the file's stated policy. **The standalone R1 identity gsm8k case is gone.** test_gsm8k_identity_vs_mtp3 evaluates the identity config as its first leg, and task.evaluate asserts accuracy against the registered reference on the way past, so the separate case gated nothing new while costing an engine boot of a 61-layer model on 4 ranks. Both the docstring and a comment where the case used to be now say that the paired test carries the identity gate. Verified on GB300 after both changes: unit 343 passed (1 GPU) + 2 passed (4 GPU); the marker does not change collection whole 2 passed (1 GPU); 3 passed (4 GPU), 21:47 -> 19:07 model paired delta +0.0000 this run; acceptance staircase 2.937 against stock 3.014, floor 2.5 Signed-off-by: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> --- .../integration/test_lists/test-db/l0_gb300_multi_gpus.yml | 1 - .../staircase/comm/test_staircase_allgather_op_matrix.py | 6 ++++++ .../comm/test_staircase_reducescatter_op_matrix.py | 6 ++++++ 3 files changed, 12 insertions(+), 1 deletion(-) diff --git a/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml b/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml index 57a7040a1d17..466d010790c8 100644 --- a/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml +++ b/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml @@ -28,7 +28,6 @@ l0_gb300_multi_gpus: # acceptance gate. The two acceptance legs are independent cases read against # one shared minimum -- staircase failing while stock passes means the draft # path regressed; both failing means the anchor is stale. - - accuracy/test_staircase.py::TestStaircaseDeepseekR10528Nvfp4Sm103Dep4::test_gsm8k - accuracy/test_staircase.py::TestStaircaseDeepseekR10528Nvfp4Sm103Dep4::test_gsm8k_identity_vs_mtp3 - accuracy/test_staircase.py::TestStaircaseDeepseekR10528Nvfp4Sm103Dep4::test_mtp3_acceptance[staircase] - accuracy/test_staircase.py::TestStaircaseDeepseekR10528Nvfp4Sm103Dep4::test_mtp3_acceptance[stock] diff --git a/tests/unittest/_torch/staircase/comm/test_staircase_allgather_op_matrix.py b/tests/unittest/_torch/staircase/comm/test_staircase_allgather_op_matrix.py index 04aec3e02270..703a7eef3c88 100644 --- a/tests/unittest/_torch/staircase/comm/test_staircase_allgather_op_matrix.py +++ b/tests/unittest/_torch/staircase/comm/test_staircase_allgather_op_matrix.py @@ -13,6 +13,7 @@ launcher; see ``_rank_job`` for why that is left intact. """ +import pytest import torch from tensorrt_llm._torch.staircase.catalog.comm import _rank_job @@ -20,5 +21,10 @@ assert torch.cuda.is_available(), "allgather requires CUDA devices" +# Each case starts its own 4-rank mpirun over every visible device. Under +# xdist several workers would fight for the same GPUs, so this must run +# alone -- the same reason the other collective entry in this repo +# (_torch/thop/serial/test_moe_alltoall.py) carries the marker. +@pytest.mark.no_xdist def test_allgather_op_matrix() -> None: _rank_job.run("allgather") diff --git a/tests/unittest/_torch/staircase/comm/test_staircase_reducescatter_op_matrix.py b/tests/unittest/_torch/staircase/comm/test_staircase_reducescatter_op_matrix.py index ce4121bd68a7..5b51e8a56432 100644 --- a/tests/unittest/_torch/staircase/comm/test_staircase_reducescatter_op_matrix.py +++ b/tests/unittest/_torch/staircase/comm/test_staircase_reducescatter_op_matrix.py @@ -13,6 +13,7 @@ cannot be a normal test, because the job that runs it never reports). """ +import pytest import torch from tensorrt_llm._torch.staircase.catalog.comm import _rank_job @@ -20,5 +21,10 @@ assert torch.cuda.is_available(), "reducescatter requires CUDA devices" +# Each case starts its own 4-rank mpirun over every visible device. Under +# xdist several workers would fight for the same GPUs, so this must run +# alone -- the same reason the other collective entry in this repo +# (_torch/thop/serial/test_moe_alltoall.py) carries the marker. +@pytest.mark.no_xdist def test_reducescatter_op_matrix() -> None: _rank_job.run("reducescatter") From a566f1b801a5ca0f9484102985b93e3692d90d7d Mon Sep 17 00:00:00 2001 From: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> Date: Mon, 14 Sep 2026 06:28:51 -0700 Subject: [PATCH 05/19] [TRTLLM-16304][test] Rename the staircase collective rank bodies off the test naming The two collective catalog entries kept their rank bodies in `catalog/comm/` as `allgather_test.py` and `reducescatter_test.py`. A `*_test.py` name inside the package tree reads as a stray test file under `tensorrt_llm/`, which is the one thing the earlier move was meant to end -- every other catalog entry's test lives in `tests/unittest/_torch/staircase//`. **The files are renamed to `_allgather_op_matrix.py` and `_reducescatter_op_matrix.py`**, after the collected entry points that drive them (`test_staircase__op_matrix.py`), so the two halves of each entry carry the same noun and a reader who greps `op_matrix` finds both. The underscore is the same "not a catalog entry" signal `_rank_job.py` already uses; a catalog entry is `.py` plus `.md`. Their bodies go from `test_*` to `check_*`, and `TESTS` to `CHECKS`. **`catalog/comm/conftest.py` is deleted.** Its `collect_ignore` existed only because the old file names matched pytest's `python_files`. The new names match neither `python_files` nor, for the bodies, `python_functions`, and that is the stronger guard of the two: measured on pytest 9.0.3, `collect_ignore` does *not* suppress a path named explicitly on the command line, while a file that matches neither pattern collects nothing under any invocation. The conftest's account of why these two files stay in the package -- the launcher re-execs them as `python -m`, and the ranks need the package context for their relative imports -- moves into their module docstrings, next to the rest of the reasoning it belongs with. No behavioural change: the launcher, the deadline it enforces, the fixed rank sequence and reducescatter's separately capped wedge sub-job are untouched, and `_rank_job.run` derives the new name the same way it derived the old one. Verified without GPUs: no stale references remain repo-wide; all five touched files compile; 20/20 and 28/28 `check_` definitions are listed in `CHECKS` with none orphaned or dangling; `pytest` over `catalog/comm/` collects nothing; ruff 0.9.4 check and format are clean and leave the files unchanged. The 4-rank matrices themselves re-run on GB300 to restore the entries' certification receipts, which this rename voids by rewriting their files. Signed-off-by: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> --- tensorrt_llm/_torch/staircase/README.md | 9 +- ...gather_test.py => _allgather_op_matrix.py} | 112 +++++++------ .../staircase/catalog/comm/_rank_job.py | 6 +- ...er_test.py => _reducescatter_op_matrix.py} | 156 ++++++++++-------- .../_torch/staircase/catalog/comm/conftest.py | 19 --- .../test_staircase_allgather_op_matrix.py | 4 +- .../test_staircase_reducescatter_op_matrix.py | 9 +- 7 files changed, 164 insertions(+), 151 deletions(-) rename tensorrt_llm/_torch/staircase/catalog/comm/{allgather_test.py => _allgather_op_matrix.py} (91%) rename tensorrt_llm/_torch/staircase/catalog/comm/{reducescatter_test.py => _reducescatter_op_matrix.py} (92%) delete mode 100644 tensorrt_llm/_torch/staircase/catalog/comm/conftest.py diff --git a/tensorrt_llm/_torch/staircase/README.md b/tensorrt_llm/_torch/staircase/README.md index df4396ca5476..9ab4f08c83b5 100644 --- a/tensorrt_llm/_torch/staircase/README.md +++ b/tensorrt_llm/_torch/staircase/README.md @@ -123,9 +123,12 @@ is valid only if it post-dates the last write to *every* file of its entry, so that check now has to look in both trees. The two collective entries keep their rank bodies in `catalog/comm/` -(`allgather_test.py`, `reducescatter_test.py`) because the launcher re-execs -them as `python -m` and the ranks need the package context; only the collected -shells moved. +(`_allgather_op_matrix.py`, `_reducescatter_op_matrix.py`) because the launcher +re-execs them as `python -m` and the ranks need the package context; only the +collected shells moved. Neither those file names nor their `check_*` bodies +match pytest's collection patterns: a package tree is no place for a +collectable test, and each of those two is one fixed 4-rank sequence that +cannot run as independent cases anyway. Identity is the path. `targets/` keeps all three segments rather than flattening them, and the class name carries the same triple; diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/allgather_test.py b/tensorrt_llm/_torch/staircase/catalog/comm/_allgather_op_matrix.py similarity index 91% rename from tensorrt_llm/_torch/staircase/catalog/comm/allgather_test.py rename to tensorrt_llm/_torch/staircase/catalog/comm/_allgather_op_matrix.py index ad375cbd2702..be61839ad1d1 100644 --- a/tensorrt_llm/_torch/staircase/catalog/comm/allgather_test.py +++ b/tensorrt_llm/_torch/staircase/catalog/comm/_allgather_op_matrix.py @@ -1,6 +1,6 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""GPU test for the allgather catalog entry. +"""GPU certification matrix for the allgather catalog entry. A collective cannot be exercised in one process, so this script is its own launcher: run it plainly and it re-executes itself under `mpirun` with one @@ -15,7 +15,21 @@ synchronisation interleaved with the op, the engine's stream switching, and what disagreeing on call order actually does. - CUDA_VISIBLE_DEVICES=0,1,2,3 uv run python catalog/comm/allgather_test.py + CUDA_VISIBLE_DEVICES=0,1,2,3 uv run python catalog/comm/_allgather_op_matrix.py + +Not a pytest module, despite the `check_*` bodies. They are one fixed +sequence inside a single 4-rank job rather than independent cases: each reads +module-global rank state that only `_run_one_rank` binds, and several assert +on communicator state the previous one left behind. Collected as tests they +would run at world size 1 against unbound globals — so neither this file's +name nor its function names match pytest's collection patterns, which is what +keeps it uncollectable however pytest is pointed at this tree. + +The collected entry point is +`tests/unittest/_torch/staircase/comm/test_staircase_allgather_op_matrix.py`: +it starts this job and turns its exit code into an assertion. This half stays +in the package because the launcher re-execs it as `python -m`, and the ranks +need the package context for their relative imports. """ import os @@ -175,7 +189,7 @@ def _warm_up_off_capture_stream(body, reps: int = 2) -> None: Two things have to be done before a capture and cannot be done inside one: the group's NCCL communicator has to exist (see - test_cuda_graph_capture_of_a_first_call_raises), and torch wants the work + check_cuda_graph_capture_of_a_first_call_raises), and torch wants the work warmed on a non-default stream. """ side = torch.cuda.Stream() @@ -261,7 +275,7 @@ def _verify_replay( _assert_bitwise(out, _gather_ref([rows] * WORLD, seed + 13 * site), f"{where} site={site}") -def test_cuda_graph_capture_of_a_first_call_raises() -> None: +def check_cuda_graph_capture_of_a_first_call_raises() -> None: """A group's first-ever call cannot be captured; the build inside fails. Must run before anything else touches GROUP — the failure is specifically @@ -318,7 +332,7 @@ def test_cuda_graph_capture_of_a_first_call_raises() -> None: COMM.Barrier() -def test_uniform_gather() -> None: +def check_uniform_gather() -> None: """sizes=None: every rank holds `rows` rows, the result is WORLD * rows. Decode-like (1, 2, 8, 32) through prefill-like (2048) row counts, which is @@ -335,7 +349,7 @@ def test_uniform_gather() -> None: COMM.Barrier() -def test_ragged_gather() -> None: +def check_ragged_gather() -> None: """sizes=[...]: per-rank row counts differ, the result is sum(sizes). Attention data parallelism produces exactly this — each rank owns its own @@ -352,7 +366,7 @@ def test_ragged_gather() -> None: COMM.Barrier() -def test_output_is_fresh_and_input_is_untouched() -> None: +def check_output_is_fresh_and_input_is_untouched() -> None: """The op allocates its result; the caller keeps ownership of `input`.""" rows, seed = 8, 4100 x = _payload(RANK, rows, seed) @@ -369,7 +383,7 @@ def test_output_is_fresh_and_input_is_untouched() -> None: COMM.Barrier() -def test_trailing_dims_are_preserved() -> None: +def check_trailing_dims_are_preserved() -> None: """dim 0 is the gather axis; every other dim is carried through untouched.""" cases: List[Tuple[Tuple[int, ...], Optional[List[int]]]] = [ ((2, HIDDEN // 2), None), @@ -392,7 +406,7 @@ def test_trailing_dims_are_preserved() -> None: COMM.Barrier() -def test_dtypes_move_bitwise() -> None: +def check_dtypes_move_bitwise() -> None: """The payload dtypes an attention-DP dispatch moves, in both forms. bf16 hidden states, plus what a post-quantization dispatch carries next to @@ -423,7 +437,7 @@ def test_dtypes_move_bitwise() -> None: COMM.Barrier() -def test_group_selects_a_rank_subset() -> None: +def check_group_selects_a_rank_subset() -> None: """`group` names MPI session ranks; ranks outside it must not call.""" subset = [0, 1] rows, seed = 4, 7100 @@ -434,7 +448,7 @@ def test_group_selects_a_rank_subset() -> None: COMM.Barrier() -def test_group_order_does_not_change_the_output_order() -> None: +def check_group_order_does_not_change_the_output_order() -> None: """The result is ordered by ascending rank, whatever order `group` lists.""" rows, seed = 4, 7200 x = _payload(RANK, rows, seed) @@ -443,7 +457,7 @@ def test_group_order_does_not_change_the_output_order() -> None: COMM.Barrier() -def test_cuda_graph_at_every_engine_batch_size() -> None: +def check_cuda_graph_at_every_engine_batch_size() -> None: """Captured at all 35 engine batch sizes, then replayed interleaved. Capturing one row count proves nothing about an engine, which holds every @@ -494,7 +508,7 @@ def test_cuda_graph_at_every_engine_batch_size() -> None: COMM.Barrier() -def test_cuda_graph_one_site_per_moe_layer_in_every_batch_size_graph() -> None: +def check_cuda_graph_one_site_per_moe_layer_in_every_batch_size_graph() -> None: """The decode graph this checkpoint captures, at all 35 batch sizes. 1015 captured gathers, one memory pool. The sites are independent — each @@ -528,7 +542,7 @@ def test_cuda_graph_one_site_per_moe_layer_in_every_batch_size_graph() -> None: COMM.Barrier() -def test_cuda_graph_replays_survive_eager_calls_of_other_shapes() -> None: +def check_cuda_graph_replays_survive_eager_calls_of_other_shapes() -> None: """Between two decode replays, a server runs work no graph holds. A prefill runs eagerly at a row count far outside the captured set, and a @@ -560,7 +574,7 @@ def test_cuda_graph_replays_survive_eager_calls_of_other_shapes() -> None: COMM.Barrier() -def test_cuda_graph_holds_the_sizes_vector_it_captured() -> None: +def check_cuda_graph_holds_the_sizes_vector_it_captured() -> None: """`sizes` is a host argument: a replay re-runs the split it was captured with. Two graphs with different sizes vectors are captured into one pool and @@ -600,7 +614,7 @@ def test_cuda_graph_holds_the_sizes_vector_it_captured() -> None: COMM.Barrier() -def test_the_gate_discriminates_a_wrong_gather() -> None: +def check_the_gate_discriminates_a_wrong_gather() -> None: """The bitwise gate rejects every plausible wrong gather, by a wide margin. A gate of exactly 0 cannot be too loose, but it can be blind: if every @@ -660,7 +674,7 @@ def test_the_gate_discriminates_a_wrong_gather() -> None: COMM.Barrier() -def test_the_engines_own_cross_rank_step_is_not_on_this_communicator() -> None: +def check_the_engines_own_cross_rank_step_is_not_on_this_communicator() -> None: """What a serving engine synchronises with, next to what this op uses. Under attention data parallelism the engine agrees the per-rank token @@ -682,7 +696,7 @@ def test_the_engines_own_cross_rank_step_is_not_on_this_communicator() -> None: COMM.Barrier() -def test_interleaved_with_the_engines_attention_dp_synchronisation() -> None: +def check_interleaved_with_the_engines_attention_dp_synchronisation() -> None: """The op inside a forward, between the engine's own cross-rank steps. The shape a served forward has: once per step the engine agrees the @@ -720,7 +734,7 @@ def test_interleaved_with_the_engines_attention_dp_synchronisation() -> None: COMM.Barrier() -def test_the_stream_the_call_lands_on_is_not_part_of_the_match() -> None: +def check_the_stream_the_call_lands_on_is_not_part_of_the_match() -> None: """The op runs on whatever stream is current, and ranks need not agree. A serving engine moves the current stream under the model: the same @@ -748,7 +762,7 @@ def test_the_stream_the_call_lands_on_is_not_part_of_the_match() -> None: COMM.Barrier() -def test_wrapper_guards_a_non_contiguous_input() -> None: +def check_wrapper_guards_a_non_contiguous_input() -> None: """The wrapper's assert stands where the op itself is silently wrong.""" rows, seed = 8, 92000 values = _payload(RANK, rows, seed) @@ -776,7 +790,7 @@ def test_wrapper_guards_a_non_contiguous_input() -> None: COMM.Barrier() -def test_wrapper_guards_a_sizes_list_of_the_wrong_length() -> None: +def check_wrapper_guards_a_sizes_list_of_the_wrong_length() -> None: """A short `sizes` list silently drops the trailing ranks.""" rows, seed = 4, 93000 short = [rows] * (WORLD - 1) @@ -796,7 +810,7 @@ def test_wrapper_guards_a_sizes_list_of_the_wrong_length() -> None: COMM.Barrier() -def test_wrapper_guards_a_zero_dim_input() -> None: +def check_wrapper_guards_a_zero_dim_input() -> None: """A 0-d input segfaults inside the op, so the wrapper stops it first. The op is deliberately not called here: the crash is in @@ -811,7 +825,7 @@ def test_wrapper_guards_a_zero_dim_input() -> None: COMM.Barrier() -def test_call_order_disagreement_corrupts_silently() -> None: +def check_call_order_disagreement_corrupts_silently() -> None: """Ranks that disagree on the *order* of two equal-sized gathers get wrong data back, with no error and no hang. @@ -866,32 +880,32 @@ def test_call_order_disagreement_corrupts_silently() -> None: COMM.Barrier() -TESTS = ( +CHECKS = ( # Stays first: it is the only test that can observe GROUP's first-ever # call, and every later test needs the communicator it builds. - test_cuda_graph_capture_of_a_first_call_raises, - test_uniform_gather, - test_ragged_gather, - test_output_is_fresh_and_input_is_untouched, - test_trailing_dims_are_preserved, - test_dtypes_move_bitwise, - test_group_selects_a_rank_subset, - test_group_order_does_not_change_the_output_order, - test_cuda_graph_at_every_engine_batch_size, - test_cuda_graph_one_site_per_moe_layer_in_every_batch_size_graph, - test_cuda_graph_replays_survive_eager_calls_of_other_shapes, - test_cuda_graph_holds_the_sizes_vector_it_captured, - test_the_gate_discriminates_a_wrong_gather, - test_the_engines_own_cross_rank_step_is_not_on_this_communicator, - test_interleaved_with_the_engines_attention_dp_synchronisation, - test_the_stream_the_call_lands_on_is_not_part_of_the_match, - test_wrapper_guards_a_non_contiguous_input, - test_wrapper_guards_a_sizes_list_of_the_wrong_length, - test_wrapper_guards_a_zero_dim_input, + check_cuda_graph_capture_of_a_first_call_raises, + check_uniform_gather, + check_ragged_gather, + check_output_is_fresh_and_input_is_untouched, + check_trailing_dims_are_preserved, + check_dtypes_move_bitwise, + check_group_selects_a_rank_subset, + check_group_order_does_not_change_the_output_order, + check_cuda_graph_at_every_engine_batch_size, + check_cuda_graph_one_site_per_moe_layer_in_every_batch_size_graph, + check_cuda_graph_replays_survive_eager_calls_of_other_shapes, + check_cuda_graph_holds_the_sizes_vector_it_captured, + check_the_gate_discriminates_a_wrong_gather, + check_the_engines_own_cross_rank_step_is_not_on_this_communicator, + check_interleaved_with_the_engines_attention_dp_synchronisation, + check_the_stream_the_call_lands_on_is_not_part_of_the_match, + check_wrapper_guards_a_non_contiguous_input, + check_wrapper_guards_a_sizes_list_of_the_wrong_length, + check_wrapper_guards_a_zero_dim_input, # Stays last: it deliberately disagrees on call order, and although a # swapped pair realigns the communicator (its final assertion proves it), # nothing after it should depend on that. - test_call_order_disagreement_corrupts_silently, + check_call_order_disagreement_corrupts_silently, ) @@ -931,13 +945,13 @@ def _run_one_rank() -> int: ) DIST = engine_dist - for test in TESTS: + for check in CHECKS: try: - test() + check() except BaseException: import traceback - print(f"[rank {RANK}] FAILED {test.__name__}", flush=True) + print(f"[rank {RANK}] FAILED {check.__name__}", flush=True) traceback.print_exc() sys.stdout.flush() sys.stderr.flush() @@ -945,7 +959,7 @@ def _run_one_rank() -> int: # wedges every other rank in it. COMM.Abort(1) COMM.Barrier() - print(f"[rank {RANK}] {len(TESTS)} tests passed", flush=True) + print(f"[rank {RANK}] {len(CHECKS)} checks passed", flush=True) return 0 @@ -966,7 +980,7 @@ def _spawn_ranks() -> None: str(world_size), sys.executable, "-m", - "tensorrt_llm._torch.staircase.catalog.comm.allgather_test", + "tensorrt_llm._torch.staircase.catalog.comm._allgather_op_matrix", _WORKER_FLAG, ] print(f"[launcher] {' '.join(command)}", flush=True) diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/_rank_job.py b/tensorrt_llm/_torch/staircase/catalog/comm/_rank_job.py index c4a9510746e0..837c0234596e 100644 --- a/tensorrt_llm/_torch/staircase/catalog/comm/_rank_job.py +++ b/tensorrt_llm/_torch/staircase/catalog/comm/_rank_job.py @@ -57,9 +57,9 @@ def _devices() -> str: def run(entry: str) -> None: - """Run ``_test``'s launcher over ``WORLD_SIZE`` devices.""" - module = f"{__package__}.{entry}_test" - launcher = Path(__file__).with_name(f"{entry}_test.py") + """Run ``__op_matrix``'s launcher over ``WORLD_SIZE`` devices.""" + module = f"{__package__}._{entry}_op_matrix" + launcher = Path(__file__).with_name(f"_{entry}_op_matrix.py") env = dict(os.environ, CUDA_VISIBLE_DEVICES=_devices()) # The launcher re-execs itself per rank and needs this package importable # from the ranks; by path it has no package context of its own to inherit. diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/reducescatter_test.py b/tensorrt_llm/_torch/staircase/catalog/comm/_reducescatter_op_matrix.py similarity index 92% rename from tensorrt_llm/_torch/staircase/catalog/comm/reducescatter_test.py rename to tensorrt_llm/_torch/staircase/catalog/comm/_reducescatter_op_matrix.py index 0b74117f01db..296554ba2059 100644 --- a/tensorrt_llm/_torch/staircase/catalog/comm/reducescatter_test.py +++ b/tensorrt_llm/_torch/staircase/catalog/comm/_reducescatter_op_matrix.py @@ -1,6 +1,6 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""GPU test for the reducescatter catalog entry. +"""GPU certification matrix for the reducescatter catalog entry. A collective cannot be exercised in one process, so this script is its own launcher: run it plainly and it re-executes itself under `mpirun` with one @@ -10,10 +10,10 @@ disagree about the split — so the deadline is what keeps a broken kernel from taking the calling run down with it. - CUDA_VISIBLE_DEVICES=0,1,2,3 uv run python catalog/comm/reducescatter_test.py + CUDA_VISIBLE_DEVICES=0,1,2,3 uv run python catalog/comm/_reducescatter_op_matrix.py -That runs two jobs, in this order. The first is the test proper: every -`TESTS` entry on `world_size` ranks, and it must exit 0. The second is a +That runs two jobs, in this order. The first is the matrix proper: every +`CHECKS` entry on `world_size` ranks, and it must exit 0. The second is a four-rank sub-job of its own that deliberately mispairs this op against an all-gather of a *different* byte count, which is the one call-order divergence that hangs instead of returning wrong data — it cannot live in @@ -22,6 +22,20 @@ the pair, none marks it afterwards, and the job is ended by its own watchdog. That sub-job doubles as the harness's positive control, since it is a real collective deadlock this launcher has to survive. + +Not a pytest module, despite the `check_*` bodies. They are one fixed +sequence inside a single 4-rank job rather than independent cases: each reads +module-global rank state that only `_run_one_rank` binds, and several assert +on communicator state the previous one left behind. Collected as tests they +would run at world size 1 against unbound globals — so neither this file's +name nor its function names match pytest's collection patterns, which is what +keeps it uncollectable however pytest is pointed at this tree. + +The collected entry point is +`tests/unittest/_torch/staircase/comm/test_staircase_reducescatter_op_matrix.py`: +it starts this job and turns its exit code into an assertion. This half stays +in the package because the launcher re-execs it as `python -m`, and the ranks +need the package context for their relative imports. """ import os @@ -104,7 +118,7 @@ def _block( fp32 all represent k/8 exactly for |k| <= 255, and a sum of WORLD of them reaches |k| <= 4 * 31 = 124, so **every partial sum is exact in every summation order**. That is what lets the assertions be bitwise despite the - op summing in the input dtype (see test_reduction_is_deterministic_and_ + op summing in the input dtype (see check_reduction_is_deterministic_and_ accumulates_in_the_input_dtype for what happens when they are not). """ gen = torch.Generator(device="cuda").manual_seed((seed * 977 + rank) * 131 + dest + 1) @@ -208,7 +222,7 @@ def _ring_chain(rows: int, pos: int, seeds: Sequence[int]) -> torch.Tensor: """Ring-order sequential sum for group position `pos`, rounded every step. The order this op reduces in (certified by - test_reduction_is_deterministic_and_accumulates_in_the_input_dtype): + check_reduction_is_deterministic_and_accumulates_in_the_input_dtype): `x_{pos+1} + x_{pos+2} + ... + x_{pos+G-1} + x_pos`, indices mod WORLD, where `x_r` is rank `r`'s block for `pos` drawn from `seeds[r]`. Per-rank seeds, so a mispaired call's exact bits can be predicted too — on payloads @@ -259,7 +273,7 @@ def _warm_up_off_capture_stream(body, reps: int = 2) -> None: Two things have to be done before a capture and cannot be done inside one: the group's NCCL communicator has to exist (see - test_cuda_graph_capture_of_a_first_call_raises), and torch wants the work + check_cuda_graph_capture_of_a_first_call_raises), and torch wants the work warmed on a non-default stream. """ side = torch.cuda.Stream() @@ -345,7 +359,7 @@ def _verify_replay( _assert_bitwise(out, _ref([rows] * WORLD, RANK, seed + 13 * site), f"{where} site={site}") -def test_cuda_graph_capture_of_a_first_call_raises() -> None: +def check_cuda_graph_capture_of_a_first_call_raises() -> None: """A group's first-ever call cannot be captured; the build inside fails. Must run before anything else touches GROUP — the failure is specifically @@ -404,7 +418,7 @@ def test_cuda_graph_capture_of_a_first_call_raises() -> None: COMM.Barrier() -def test_uniform_reduce_scatter() -> None: +def check_uniform_reduce_scatter() -> None: """sizes=None: every rank sends WORLD * rows rows and keeps `rows` of them. Decode-like (1, 2, 8, 32) through prefill-like (2048) per-rank row counts, @@ -424,7 +438,7 @@ def test_uniform_reduce_scatter() -> None: COMM.Barrier() -def test_ragged_reduce_scatter() -> None: +def check_ragged_reduce_scatter() -> None: """sizes=[...]: the split is uneven, this rank keeps sizes[my position]. Attention data parallelism produces exactly this — each rank owns its own @@ -442,7 +456,7 @@ def test_ragged_reduce_scatter() -> None: COMM.Barrier() -def test_output_is_fresh_and_input_is_untouched() -> None: +def check_output_is_fresh_and_input_is_untouched() -> None: """The op allocates its result; the caller keeps ownership of `input`.""" rows, seed = 8, 4100 sizes = [rows] * WORLD @@ -460,7 +474,7 @@ def test_output_is_fresh_and_input_is_untouched() -> None: COMM.Barrier() -def test_trailing_dims_are_preserved() -> None: +def check_trailing_dims_are_preserved() -> None: """dim 0 is the scatter axis; every other dim is carried through untouched.""" cases: List[Tuple[Tuple[int, ...], bool]] = [ ((2, HIDDEN // 2), False), @@ -482,7 +496,7 @@ def test_trailing_dims_are_preserved() -> None: COMM.Barrier() -def test_dtypes_reduce_arithmetically() -> None: +def check_dtypes_reduce_arithmetically() -> None: """The dtypes whose sum this op actually computes, in both forms. fp16 and fp32 alongside bf16, and the integer widths, which reduce as @@ -513,7 +527,7 @@ def test_dtypes_reduce_arithmetically() -> None: COMM.Barrier() -def test_group_selects_a_rank_subset() -> None: +def check_group_selects_a_rank_subset() -> None: """`group` names MPI session ranks; the slice index is the position in it. The second subset is the discriminating one: it excludes rank 0, so a rank @@ -539,7 +553,7 @@ def test_group_selects_a_rank_subset() -> None: COMM.Barrier() -def test_group_order_does_not_change_the_output() -> None: +def check_group_order_does_not_change_the_output() -> None: """The split is ordered by ascending rank, whatever order `group` lists.""" rows, seed = 4, 7200 sizes = [rows] * WORLD @@ -549,7 +563,7 @@ def test_group_order_does_not_change_the_output() -> None: COMM.Barrier() -def test_reduction_is_deterministic_and_accumulates_in_the_input_dtype() -> None: +def check_reduction_is_deterministic_and_accumulates_in_the_input_dtype() -> None: """What the sum is, exactly — the fact a target's accuracy gate rests on. On payloads with no exactness property (standard normal, cast to bf16): @@ -625,7 +639,7 @@ def test_reduction_is_deterministic_and_accumulates_in_the_input_dtype() -> None COMM.Barrier() -def test_uniform_form_and_an_explicit_even_split_agree_bitwise() -> None: +def check_uniform_form_and_an_explicit_even_split_agree_bitwise() -> None: """`sizes=None` and an explicit even `sizes` vector return the same bits. Worth pinning because the two are not obliged to take the same path — and @@ -644,7 +658,7 @@ def test_uniform_form_and_an_explicit_even_split_agree_bitwise() -> None: COMM.Barrier() -def test_round_trip_with_a_gather_returns_each_rank_its_own_rows() -> None: +def check_round_trip_with_a_gather_returns_each_rank_its_own_rows() -> None: """The attention-DP MoE round trip, end to end, against a local reference. Gather every rank's tokens, let each rank apply its own expert window to @@ -684,7 +698,7 @@ def test_round_trip_with_a_gather_returns_each_rank_its_own_rows() -> None: COMM.Barrier() -def test_cuda_graph_at_every_engine_batch_size() -> None: +def check_cuda_graph_at_every_engine_batch_size() -> None: """Captured at all 35 engine batch sizes, then replayed interleaved. Capturing one row count proves nothing about an engine, which holds every @@ -732,7 +746,7 @@ def test_cuda_graph_at_every_engine_batch_size() -> None: COMM.Barrier() -def test_cuda_graph_one_site_per_moe_layer_in_every_batch_size_graph() -> None: +def check_cuda_graph_one_site_per_moe_layer_in_every_batch_size_graph() -> None: """The decode graph this checkpoint captures, at all 35 batch sizes. 1015 captured reduce-scatters, one memory pool. The sites are independent — @@ -766,7 +780,7 @@ def test_cuda_graph_one_site_per_moe_layer_in_every_batch_size_graph() -> None: COMM.Barrier() -def test_cuda_graph_replays_survive_eager_calls_of_other_shapes() -> None: +def check_cuda_graph_replays_survive_eager_calls_of_other_shapes() -> None: """Between two decode replays, a server runs work no graph holds. A prefill runs eagerly at a row count far outside the captured set, and a @@ -799,7 +813,7 @@ def test_cuda_graph_replays_survive_eager_calls_of_other_shapes() -> None: COMM.Barrier() -def test_cuda_graph_holds_the_sizes_vector_it_captured() -> None: +def check_cuda_graph_holds_the_sizes_vector_it_captured() -> None: """`sizes` is a host argument: a replay re-runs the split it was captured with. Two graphs with different sizes vectors are captured into one pool and @@ -839,7 +853,7 @@ def test_cuda_graph_holds_the_sizes_vector_it_captured() -> None: COMM.Barrier() -def test_the_gate_discriminates_a_wrong_reduce_scatter() -> None: +def check_the_gate_discriminates_a_wrong_reduce_scatter() -> None: """The bitwise gate rejects every plausible wrong result, by a wide margin. A gate of exactly 0 cannot be too loose, but it can be blind: if every @@ -885,7 +899,7 @@ def test_the_gate_discriminates_a_wrong_reduce_scatter() -> None: COMM.Barrier() -def test_calls_pair_by_position_and_a_swapped_pair_realigns() -> None: +def check_calls_pair_by_position_and_a_swapped_pair_realigns() -> None: """Calls pair by their **position** on the communicator, not by intent. One rank issuing two same-shaped calls in the opposite order to everybody @@ -953,7 +967,7 @@ def test_calls_pair_by_position_and_a_swapped_pair_realigns() -> None: COMM.Barrier() -def test_a_mispaired_result_is_deterministic_rather_than_noise() -> None: +def check_a_mispaired_result_is_deterministic_rather_than_noise() -> None: """Because this op computes, "wrong" could have meant "unreproducible". It does not. On payloads with no exactness property, a mispaired call is @@ -1007,7 +1021,7 @@ def test_a_mispaired_result_is_deterministic_rather_than_noise() -> None: COMM.Barrier() -def test_an_extra_call_on_one_rank_misaligns_until_the_counts_match() -> None: +def check_an_extra_call_on_one_rank_misaligns_until_the_counts_match() -> None: """An odd number of extra calls does not realign; a swapped pair does. Rank 0 issues one call the others never issue, then all ranks issue four @@ -1061,7 +1075,7 @@ def test_an_extra_call_on_one_rank_misaligns_until_the_counts_match() -> None: COMM.Barrier() -def test_mispaired_against_an_all_gather_of_equal_byte_count_corrupts_silently() -> None: +def check_mispaired_against_an_all_gather_of_equal_byte_count_corrupts_silently() -> None: """A different collective at the same byte count is not detected either. What pairs is position, not the identity of the op: with rank 0 issuing @@ -1125,7 +1139,7 @@ def test_mispaired_against_an_all_gather_of_equal_byte_count_corrupts_silently() COMM.Barrier() -def test_the_stream_the_call_lands_on_is_not_part_of_the_match() -> None: +def check_the_stream_the_call_lands_on_is_not_part_of_the_match() -> None: """The op runs on whatever stream is current, and ranks need not agree. A serving engine moves the current stream under the model: the same forward @@ -1175,7 +1189,7 @@ def test_the_stream_the_call_lands_on_is_not_part_of_the_match() -> None: COMM.Barrier() -def test_float8_is_summed_as_raw_bytes() -> None: +def check_float8_is_summed_as_raw_bytes() -> None: """float8_e4m3fn is accepted and reduced as unsigned bytes, not as floats. Measured, and the reason the wrapper rejects the dtype: an all-gather moves @@ -1213,7 +1227,7 @@ def test_float8_is_summed_as_raw_bytes() -> None: COMM.Barrier() -def test_unsupported_dtypes_raise_and_poison_every_later_collective() -> None: +def check_unsupported_dtypes_raise_and_poison_every_later_collective() -> None: """fp64 and the other float8 formats raise — and the raise is terminal. Runs last, and has to: the raise leaves NCCL's group state unbalanced (the @@ -1257,7 +1271,7 @@ def test_unsupported_dtypes_raise_and_poison_every_later_collective() -> None: COMM.Barrier() -def test_wrapper_guards_a_non_contiguous_input() -> None: +def check_wrapper_guards_a_non_contiguous_input() -> None: """The wrapper's assert stands where the op itself is silently wrong.""" rows, seed = 8, 92000 sizes = [rows] * WORLD @@ -1286,7 +1300,7 @@ def test_wrapper_guards_a_non_contiguous_input() -> None: COMM.Barrier() -def test_wrapper_guards_a_zero_dim_input() -> None: +def check_wrapper_guards_a_zero_dim_input() -> None: """A 0-d input segfaults inside the op, so the wrapper stops it first. The op is deliberately not called here: the crash is in @@ -1301,7 +1315,7 @@ def test_wrapper_guards_a_zero_dim_input() -> None: COMM.Barrier() -def test_wrapper_guards_a_sizes_list_of_the_wrong_length() -> None: +def check_wrapper_guards_a_sizes_list_of_the_wrong_length() -> None: """Neither wrong length is survivable, so the wrapper stops both. The raw op is deliberately not called with either. A list one entry short @@ -1326,7 +1340,7 @@ def test_wrapper_guards_a_sizes_list_of_the_wrong_length() -> None: COMM.Barrier() -def test_wrapper_guards_a_split_that_does_not_cover_the_input() -> None: +def check_wrapper_guards_a_split_that_does_not_cover_the_input() -> None: """Rows the split does not reach are silently dropped, not flagged. Two ways to get there, both exercised on the raw op because both are @@ -1361,7 +1375,7 @@ def test_wrapper_guards_a_split_that_does_not_cover_the_input() -> None: COMM.Barrier() -def test_the_group_still_works_after_the_negative_tests() -> None: +def check_the_group_still_works_after_the_negative_tests() -> None: """The raises above leave the communicator usable — checked, not assumed.""" rows, seed = 8, 95000 sizes = [rows] * WORLD @@ -1370,42 +1384,42 @@ def test_the_group_still_works_after_the_negative_tests() -> None: COMM.Barrier() -TESTS = ( +CHECKS = ( # Stays first: it is the only test that can observe GROUP's first-ever # call, and every later test needs the communicator it builds. - test_cuda_graph_capture_of_a_first_call_raises, - test_uniform_reduce_scatter, - test_ragged_reduce_scatter, - test_output_is_fresh_and_input_is_untouched, - test_trailing_dims_are_preserved, - test_dtypes_reduce_arithmetically, - test_group_selects_a_rank_subset, - test_group_order_does_not_change_the_output, - test_reduction_is_deterministic_and_accumulates_in_the_input_dtype, - test_uniform_form_and_an_explicit_even_split_agree_bitwise, - test_round_trip_with_a_gather_returns_each_rank_its_own_rows, - test_cuda_graph_at_every_engine_batch_size, - test_cuda_graph_one_site_per_moe_layer_in_every_batch_size_graph, - test_cuda_graph_replays_survive_eager_calls_of_other_shapes, - test_cuda_graph_holds_the_sizes_vector_it_captured, - test_the_gate_discriminates_a_wrong_reduce_scatter, + check_cuda_graph_capture_of_a_first_call_raises, + check_uniform_reduce_scatter, + check_ragged_reduce_scatter, + check_output_is_fresh_and_input_is_untouched, + check_trailing_dims_are_preserved, + check_dtypes_reduce_arithmetically, + check_group_selects_a_rank_subset, + check_group_order_does_not_change_the_output, + check_reduction_is_deterministic_and_accumulates_in_the_input_dtype, + check_uniform_form_and_an_explicit_even_split_agree_bitwise, + check_round_trip_with_a_gather_returns_each_rank_its_own_rows, + check_cuda_graph_at_every_engine_batch_size, + check_cuda_graph_one_site_per_moe_layer_in_every_batch_size_graph, + check_cuda_graph_replays_survive_eager_calls_of_other_shapes, + check_cuda_graph_holds_the_sizes_vector_it_captured, + check_the_gate_discriminates_a_wrong_reduce_scatter, # The call-order block. Each of these deliberately disagrees about call # order and each restores alignment before it returns — the plain call # every one of them ends on is what proves it. - test_calls_pair_by_position_and_a_swapped_pair_realigns, - test_a_mispaired_result_is_deterministic_rather_than_noise, - test_an_extra_call_on_one_rank_misaligns_until_the_counts_match, - test_mispaired_against_an_all_gather_of_equal_byte_count_corrupts_silently, - test_the_stream_the_call_lands_on_is_not_part_of_the_match, - test_float8_is_summed_as_raw_bytes, - test_wrapper_guards_a_non_contiguous_input, - test_wrapper_guards_a_zero_dim_input, - test_wrapper_guards_a_sizes_list_of_the_wrong_length, - test_wrapper_guards_a_split_that_does_not_cover_the_input, - test_the_group_still_works_after_the_negative_tests, + check_calls_pair_by_position_and_a_swapped_pair_realigns, + check_a_mispaired_result_is_deterministic_rather_than_noise, + check_an_extra_call_on_one_rank_misaligns_until_the_counts_match, + check_mispaired_against_an_all_gather_of_equal_byte_count_corrupts_silently, + check_the_stream_the_call_lands_on_is_not_part_of_the_match, + check_float8_is_summed_as_raw_bytes, + check_wrapper_guards_a_non_contiguous_input, + check_wrapper_guards_a_zero_dim_input, + check_wrapper_guards_a_sizes_list_of_the_wrong_length, + check_wrapper_guards_a_split_that_does_not_cover_the_input, + check_the_group_still_works_after_the_negative_tests, # Stays last: the raise it asserts leaves every later collective in the # process returning garbage, so nothing can run after it. - test_unsupported_dtypes_raise_and_poison_every_later_collective, + check_unsupported_dtypes_raise_and_poison_every_later_collective, ) @@ -1424,13 +1438,13 @@ def _run_one_rank() -> int: assert WORLD >= 2, f"a collective needs at least 2 ranks, got {WORLD}" torch.cuda.set_device(RANK) - for test in TESTS: + for check in CHECKS: try: - test() + check() except BaseException: import traceback - print(f"[rank {RANK}] FAILED {test.__name__}", flush=True) + print(f"[rank {RANK}] FAILED {check.__name__}", flush=True) traceback.print_exc() sys.stdout.flush() sys.stderr.flush() @@ -1438,7 +1452,7 @@ def _run_one_rank() -> int: # wedges every other rank in it. COMM.Abort(1) COMM.Barrier() - print(f"[rank {RANK}] {len(TESTS)} tests passed", flush=True) + print(f"[rank {RANK}] {len(CHECKS)} checks passed", flush=True) return 0 @@ -1469,7 +1483,7 @@ def _run_wedge_rank() -> int: """Body of one rank of the sub-job that certifies the wedge. Same mispairing as - test_mispaired_against_an_all_gather_of_equal_byte_count_corrupts_silently, + check_mispaired_against_an_all_gather_of_equal_byte_count_corrupts_silently, with one difference: the two calls carry **different** element counts (`rows*HIDDEN` against `5*HIDDEN`). That is the case NCCL cannot serve out of the buffers it was given, and it hangs rather than returning wrong data. @@ -1529,7 +1543,7 @@ def _mpirun(world_size: int, flag: str, env: Optional[Dict[str, str]] = None) -> str(world_size), sys.executable, "-m", - "tensorrt_llm._torch.staircase.catalog.comm.reducescatter_test", + "tensorrt_llm._torch.staircase.catalog.comm._reducescatter_op_matrix", flag, ] print(f"[launcher] {' '.join(command)}", flush=True) @@ -1539,7 +1553,7 @@ def _mpirun(world_size: int, flag: str, env: Optional[Dict[str, str]] = None) -> def _certify_the_unequal_byte_count_wedge(world_size: int) -> None: """Second job: the one call-order divergence that hangs instead of lying. - It cannot be a `TESTS` entry, because the job that runs it never reports. + It cannot be a `CHECKS` entry, because the job that runs it never reports. So it runs on its own, and the evidence is the marks its ranks leave: all of them entered the mispaired pair, none came out of it within WEDGE_GRACE_S, and the job ended by its own watchdog rather than by diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/conftest.py b/tensorrt_llm/_torch/staircase/catalog/comm/conftest.py deleted file mode 100644 index 88bd8941aa50..000000000000 --- a/tensorrt_llm/_torch/staircase/catalog/comm/conftest.py +++ /dev/null @@ -1,19 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -"""Keep pytest from collecting the collective entries' rank bodies. - -``allgather_test.py`` and ``reducescatter_test.py`` are their own launchers: -their ``test_*`` functions read module-global rank state that only -``_run_one_rank`` binds, and several of them assert on communicator state the -previous test left behind, so they are a fixed sequence inside one 4-rank job -rather than independent cases. Collected directly they would run at world size -1 against unbound globals. - -``tests/unittest/_torch/staircase/comm/test_staircase_*_op_matrix.py`` are the -collected entry points; each starts the 4-rank job and reports its result. -These two files stay here rather than moving with them because the launcher -re-execs them as ``python -m`` and the ranks need this package context for -their relative imports. -""" - -collect_ignore = ["allgather_test.py", "reducescatter_test.py"] diff --git a/tests/unittest/_torch/staircase/comm/test_staircase_allgather_op_matrix.py b/tests/unittest/_torch/staircase/comm/test_staircase_allgather_op_matrix.py index 703a7eef3c88..952f3c350841 100644 --- a/tests/unittest/_torch/staircase/comm/test_staircase_allgather_op_matrix.py +++ b/tests/unittest/_torch/staircase/comm/test_staircase_allgather_op_matrix.py @@ -9,8 +9,8 @@ are silently wrong rather than loud. The names are kept apart so review does not read one as a copy of the other. -The matrix itself lives in ``allgather_test.py``, which is its own 4-rank -launcher; see ``_rank_job`` for why that is left intact. +The matrix itself lives in ``catalog/comm/_allgather_op_matrix.py``, which is +its own 4-rank launcher; see ``_rank_job`` for why that is left intact. """ import pytest diff --git a/tests/unittest/_torch/staircase/comm/test_staircase_reducescatter_op_matrix.py b/tests/unittest/_torch/staircase/comm/test_staircase_reducescatter_op_matrix.py index 5b51e8a56432..407cd19f7d50 100644 --- a/tests/unittest/_torch/staircase/comm/test_staircase_reducescatter_op_matrix.py +++ b/tests/unittest/_torch/staircase/comm/test_staircase_reducescatter_op_matrix.py @@ -7,10 +7,11 @@ ``AllReduce`` module, while this covers the op itself cell by cell. The names are kept apart so review does not read one as a copy of the other. -The matrix lives in ``reducescatter_test.py``, which is its own launcher and -runs two jobs: the ordered test sequence, then a separately capped job that -certifies the one call-order divergence that wedges instead of lying (it -cannot be a normal test, because the job that runs it never reports). +The matrix lives in ``catalog/comm/_reducescatter_op_matrix.py``, which is its +own launcher and runs two jobs: the ordered check sequence, then a separately +capped job that certifies the one call-order divergence that wedges instead of +lying (it cannot be a normal check, because the job that runs it never +reports). """ import pytest From b52ec33be749a461988f5bf89d94008f025e6afa Mon Sep 17 00:00:00 2001 From: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> Date: Tue, 15 Sep 2026 08:23:05 -0700 Subject: [PATCH 06/19] [TRTLLM-16304][chore] Address the in-tree migration review The out-of-tree tree came in largely intact. This takes out what only made sense there, and closes the gaps review found. Records that had no reader in tree (-4549 lines). Each target carried four products; only modeling.py and weights.py are code. TARGET.md was the bring-up flow's own log -- checkpoint digests, machine names, run tables -- and configs/ was six knob variants whose whole payload is thirty lines of LLM API arguments the accuracy gates already carry inline. docs/ and references/accuracy.yaml restated what the code and the suite's own anchors say, with nothing keeping either in step. Accuracy anchors live where every other model keeps them, tests/integration/defs/accuracy/references/. Deleting these takes the staircase package_data block out of setup.py with it, and the internal host names and checkpoint digests go with them. Pinned versions. Out of tree a receipt was keyed by (arch, trtllm version) and a target asserted its pin at import, because a separate repo against a fixed dependency really was voided by a bump. In tree there is no external pin to drift against and the catalog tests run in pre-merge, so a version written into a contract is a number that is wrong the next day with nothing to notice. The trtllm and flashinfer pins, the measurement dates and the "added between rc21 and rc26" notes are gone; what those notes said about the code is kept without the version. The arch key stays and is now the whole key -- the mxe4m3_mxe2m1 FC1 epilogue computes bit-different block scales on sm_100 and sm_103, so one architecture's receipt refutes the other's. The sm_100 receipts are dropped rather than stripped: each predates every file of the entry it sat in, and no sm_100 machine is in CI to re-take them on. test_staircase_no_stale_claims.py is what stops the next pin being written. Test layout. test_staircase.py held two unrelated models; it is now one file per family, matching the package. The two GSM8K subclasses go with it -- pinning scores_filter and the output budget is what mocker.patch.dict(GSM8K.EVALUATE_KWARGS, ...) already does at twenty-odd call sites, including the TestGPTOSS case that gates this exact checkpoint. The collectives' rank bodies (2637 lines nothing under tensorrt_llm/ imports) move under tests/: ranks now start by file path and reach the catalog by absolute import, so neither half needs a package context. Gaps review found. TRTLLM_STAIRCASE=require on a backend that never reaches the resolver now raises instead of silently running something else. explain builds a real ModelConfig, so quant_config is the value the engine will route on rather than a hardcoded None, and the claim test now also forbids routing from reading what explain still cannot fill. The precedence between the staircase rewrite and the two above it in _resolve_class is stated. The l0 entry declares a measured TIMEOUT. nvfp4_gemm rejects a scale buffer shorter than the padded rectangle the kernel indexes, and flashinfer_silu_and_mul rejects a half that is not 16-byte aligned -- an out-of-bounds read and a context-poisoning misaligned launch respectively, both of which the ops accept. dep4's 131-line module docstring is down to the identity, the shape, the second forward path and what dep4 means; every trap it listed is stated at the call site that depends on it. _check_static_contract stops running at import: the op list stays as REQUIRED_TRTLLM_OPS and a test asserts it. CODEOWNERS gains the subtree, which was falling through to the runtime team. Signed-off-by: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> --- .github/CODEOWNERS | 13 + setup.py | 8 - tensorrt_llm/_torch/models/modeling_auto.py | 9 + tensorrt_llm/_torch/staircase/README.md | 92 +- tensorrt_llm/_torch/staircase/__init__.py | 4 + .../_torch/staircase/_router_index.py | 65 +- .../activation/flashinfer_silu_and_mul.md | 6 +- .../activation/flashinfer_silu_and_mul.py | 18 + .../catalog/attention/fused_qk_norm_rope.md | 6 +- .../attention/load_paged_kv_cache_for_mla.md | 3 +- .../mla_rope_append_paged_kv_assign_q.md | 17 +- .../mla_rope_append_paged_kv_assign_q.py | 2 +- .../catalog/attention/mla_rope_generation.md | 10 +- .../catalog/attention/mla_rope_generation.py | 10 +- .../catalog/attention/thop_attention.md | 3 +- .../catalog/attention/thop_attention.py | 3 +- .../staircase/catalog/comm/allgather.md | 3 +- .../staircase/catalog/comm/reducescatter.md | 3 +- .../_torch/staircase/catalog/gemm/bmm_out.md | 21 +- .../staircase/catalog/gemm/cublas_mm.md | 5 +- .../staircase/catalog/gemm/nvfp4_gemm.md | 3 +- .../staircase/catalog/gemm/nvfp4_gemm.py | 37 + .../_torch/staircase/catalog/index.yaml | 26 +- .../catalog/moe/fp4_block_scale_moe_runner.md | 3 +- .../_torch/staircase/catalog/moe/fused_moe.md | 3 +- .../mxe4m3_mxe2m1_block_scale_moe_runner.md | 3 +- .../staircase/catalog/moe/noaux_tc_op.md | 5 +- .../norm/flashinfer_fused_add_rmsnorm.md | 8 +- .../catalog/norm/flashinfer_rmsnorm.md | 6 +- .../catalog/quantization/fp4_quantize.md | 5 +- .../catalog/quantization/mxfp8_quantize.md | 3 +- .../docs/models/expert-weight-packing.md | 553 ------ .../docs/models/multi-token-prediction.md | 299 --- .../references/trtllm-runtime-integration.md | 1053 ----------- tensorrt_llm/_torch/staircase/explain.py | 18 +- .../r1_0528_nvfp4/sm_103/dep4/TARGET.md | 1631 ----------------- .../sm_103/dep4/configs/identity.yaml | 21 - .../sm_103/dep4/configs/mtp1.yaml | 69 - .../sm_103/dep4/configs/mtp2.yaml | 69 - .../sm_103/dep4/configs/mtp3.yaml | 79 - .../sm_103/dep4/configs/trtllm-ref-boot.yaml | 30 - .../sm_103/dep4/configs/trtllm-ref-mtp3.yaml | 75 - .../r1_0528_nvfp4/sm_103/dep4/modeling.py | 221 +-- .../r1_0528_nvfp4/sm_103/dep4/weights.py | 2 +- .../staircase/models/gpt_oss/routing.py | 4 +- .../targets/gpt_oss_120b/sm_103/tp1/TARGET.md | 480 ----- .../gpt_oss_120b/sm_103/tp1/modeling.py | 40 +- .../_torch/staircase/references/accuracy.yaml | 190 -- tensorrt_llm/llmapi/llm.py | 7 + ...rcase.py => test_staircase_deepseek_v3.py} | 115 +- .../defs/accuracy/test_staircase_gpt_oss.py | 91 + .../test_lists/test-db/l0_gb300.yml | 8 +- .../test-db/l0_gb300_multi_gpus.yml | 6 +- .../test_staircase_flashinfer_silu_and_mul.py | 30 + ...rcase_mla_rope_append_paged_kv_assign_q.py | 6 +- .../staircase}/comm/_allgather_op_matrix.py | 16 +- .../_torch/staircase}/comm/_rank_job.py | 47 +- .../comm/_reducescatter_op_matrix.py | 17 +- .../test_staircase_allgather_op_matrix.py | 5 +- .../test_staircase_reducescatter_op_matrix.py | 5 +- .../gemm/test_staircase_nvfp4_gemm.py | 42 + .../_torch/staircase/test_staircase_claims.py | 58 +- .../test_staircase_no_stale_claims.py | 78 + .../staircase/test_staircase_routing.py | 13 +- .../test_staircase_target_contract.py | 73 + 65 files changed, 812 insertions(+), 5042 deletions(-) delete mode 100644 tensorrt_llm/_torch/staircase/docs/models/expert-weight-packing.md delete mode 100644 tensorrt_llm/_torch/staircase/docs/models/multi-token-prediction.md delete mode 100644 tensorrt_llm/_torch/staircase/docs/references/trtllm-runtime-integration.md delete mode 100644 tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/TARGET.md delete mode 100644 tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/identity.yaml delete mode 100644 tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp1.yaml delete mode 100644 tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp2.yaml delete mode 100644 tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp3.yaml delete mode 100644 tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/trtllm-ref-boot.yaml delete mode 100644 tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/trtllm-ref-mtp3.yaml delete mode 100644 tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/TARGET.md delete mode 100644 tensorrt_llm/_torch/staircase/references/accuracy.yaml rename tests/integration/defs/accuracy/{test_staircase.py => test_staircase_deepseek_v3.py} (67%) create mode 100644 tests/integration/defs/accuracy/test_staircase_gpt_oss.py rename {tensorrt_llm/_torch/staircase/catalog => tests/unittest/_torch/staircase}/comm/_allgather_op_matrix.py (98%) rename {tensorrt_llm/_torch/staircase/catalog => tests/unittest/_torch/staircase}/comm/_rank_job.py (62%) rename {tensorrt_llm/_torch/staircase/catalog => tests/unittest/_torch/staircase}/comm/_reducescatter_op_matrix.py (99%) create mode 100644 tests/unittest/_torch/staircase/test_staircase_no_stale_claims.py create mode 100644 tests/unittest/_torch/staircase/test_staircase_target_contract.py diff --git a/.github/CODEOWNERS b/.github/CODEOWNERS index 98806120e518..2d93e55f7e20 100644 --- a/.github/CODEOWNERS +++ b/.github/CODEOWNERS @@ -385,6 +385,19 @@ /tensorrt_llm/scaffolding @WeiHaocheng @dc3671 /tests/unittest/scaffolding @WeiHaocheng @dc3671 +# ===== STAIRCASE ===== +# Overrides the /tensorrt_llm/_torch runtime-devs rule above for this subtree. +# Individual handles rather than a team, like SCAFFOLDING: this is one bounded +# experiment with named owners, not a standing domain. +/tensorrt_llm/_torch/staircase @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju +/tests/unittest/_torch/staircase @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju +/tests/integration/defs/accuracy/test_staircase_deepseek_v3.py @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju +/tests/integration/defs/accuracy/test_staircase_gpt_oss.py @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju +# The two catalog categories whose contracts state kernel behaviour the +# attention and MoE owners are the authority on; co-owned rather than reassigned. +/tensorrt_llm/_torch/staircase/catalog/attention @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju @xxi-nv @yuxianq +/tensorrt_llm/_torch/staircase/catalog/moe @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju @xxi-nv @yuxianq + ## TensorRT-LLM LLM Disaggregated /examples/disaggregated @NVIDIA/trt-llm-disagg-devs @NVIDIA/trt-llm-doc-owners /examples/disaggregated/slurm/benchmark @NVIDIA/trt-llm-disagg-devs @NVIDIA/trtllm-bench-reviewers diff --git a/setup.py b/setup.py index 2176f4f8166d..a4150a691261 100644 --- a/setup.py +++ b/setup.py @@ -209,14 +209,6 @@ def has_ext_modules(self): 'bindings/**/*.pyi', 'evaluate/lm_eval_tasks/**/*', 'usage/schemas/*.json', - # Two staircase patterns, both load-bearing. A target's configs/ are - # `--extra_llm_api_options` files that TARGET.md's verification commands - # pass to trtllm-eval by path, and the claim test asserts TARGET.md sits - # beside the modeling.py it vouches for; both read the installed tree. The - # contracts, catalog index and docs are read by people in a checkout and - # by no code, so they stay out of the wheel. - '_torch/staircase/models/**/*.md', - '_torch/staircase/models/**/configs/*.yaml', ] diff --git a/tensorrt_llm/_torch/models/modeling_auto.py b/tensorrt_llm/_torch/models/modeling_auto.py index cfe12520247f..97fe74e71f3f 100644 --- a/tensorrt_llm/_torch/models/modeling_auto.py +++ b/tensorrt_llm/_torch/models/modeling_auto.py @@ -39,6 +39,15 @@ def _resolve_class(config: ModelConfig) -> Optional[Type]: # Returns None unless `staircase` is on and a target claims this exact # (checkpoint, GPU arch, parallel topology), so the default path is # byte-for-byte unchanged. + # + # Precedence, since this runs last and would override the rewrite + # above: staircase wins. It reads the *un-rewritten* architectures[0], + # so it decides on the checkpoint rather than on what that rewrite made + # of it, and a target that claims a configuration carries that + # configuration's draft path itself. Not reachable today -- Eagle3 + # needs draft_vocab_size, and no draft checkpoint matches a target's + # shape fingerprint -- so this note is the contract, not a description + # of observed behaviour. if (staircase_arch := staircase_resolve(config)) is not None: model_arch = staircase_arch diff --git a/tensorrt_llm/_torch/staircase/README.md b/tensorrt_llm/_torch/staircase/README.md index 9ab4f08c83b5..9aa1473c3bff 100644 --- a/tensorrt_llm/_torch/staircase/README.md +++ b/tensorrt_llm/_torch/staircase/README.md @@ -79,8 +79,8 @@ A quantity may be a routing criterion only if it is known when `_resolve_class` runs **and** constant for the engine's life — and only if it changes the *structure of the forward*. Config shape, SM version, parallel topology qualify. `max_num_tokens` and `cuda_graph_config.max_batch_size` pass -the first test and fail the second: they are tuning knobs and belong in a -target's `configs/` variant. Batch composition and `num_contexts` fail the +the first test and fail the second: they are tuning knobs — LLM API arguments, +not identity. Batch composition and `num_contexts` fail the first outright — those move every step, and a target chosen from them would be chosen once and then be wrong. Per-step specialization is a separate mechanism (dispatch inside a target's `forward`), not a routing dimension. @@ -91,8 +91,8 @@ Routing recognizes a checkpoint by `(num_hidden_layers, hidden_size, ...)`, the same sniffing idiom as `is_mla` / `is_nemotron_hybrid` upstream. It cannot tell a fine-tune from the original. That is a deliberate trade, and it changes what a gate record means: not "this target passed" but "this modeling code -passed **on the checkpoint whose digests TARGET.md records**". Running it on -any other checkpoint of the same shape is ungated. Each `TARGET.md` says so. +passed **on the checkpoint the accuracy gate names**". Running it on any other +checkpoint of the same shape is ungated. ## Layout @@ -102,12 +102,13 @@ explain.py why a configuration routed where it did models// routing.py one forward-reading decision tree per architecture family targets//// - modeling.py weights.py TARGET.md [configs/] + modeling.py weights.py catalog/ the kernel vocabulary: contract .md + wrapper .py -references/ accuracy anchors, keyed by HF repo name -docs/ cross-target mechanism and runtime notes ``` +Accuracy anchors are not kept here. They live where every other model's do, +`tests/integration/defs/accuracy/references/`, read by the gates CI runs. + The third piece of a catalog entry, its GPU test, lives in the tests tree: ``` @@ -115,6 +116,8 @@ tests/unittest/_torch/staircase/ test_staircase_claims.py routing tables vs the targets they name (no GPU) test_staircase_routing.py what staircase_resolve does (no GPU) /test_staircase_.py + comm/__op_matrix.py the two collectives' 4-rank rank bodies + comm/_rank_job.py starts one of those and asserts on its exit code ``` That split is not a preference; it is where this repo's CI collects from, and @@ -122,13 +125,13 @@ an in-package test is on no list. It costs one thing worth stating: a receipt is valid only if it post-dates the last write to *every* file of its entry, so that check now has to look in both trees. -The two collective entries keep their rank bodies in `catalog/comm/` -(`_allgather_op_matrix.py`, `_reducescatter_op_matrix.py`) because the launcher -re-execs them as `python -m` and the ranks need the package context; only the -collected shells moved. Neither those file names nor their `check_*` bodies -match pytest's collection patterns: a package tree is no place for a -collectable test, and each of those two is one fixed 4-rank sequence that -cannot run as independent cases anyway. +The two collectives' rank bodies sit in the tests tree with everything else +that only tests run. Both halves are started by file path -- the launcher must +not import `tensorrt_llm`, because that calls `MPI_Init` and an +MPI-initialized process cannot start `mpirun`, and the ranks reach the catalog +by absolute import -- so neither needs a package to live in. Neither those +file names nor their `check_*` bodies match pytest's collection patterns: +each is one fixed 4-rank sequence that cannot run as independent cases. Identity is the path. `targets/` keeps all three segments rather than flattening them, and the class name carries the same triple; @@ -154,12 +157,12 @@ checkpoints", which is the opposite of what this package is for. 1. **Boot** — minutes, binary. `examples/llm-api/quickstart_advanced.py` with the target's topology flags and `TRTLLM_STAIRCASE=require` exported. Engine cold start, weight-manifest coverage, a handful of greedy continuations. - Catches catastrophes, not accuracy. A target whose `configs/` holds a - variant that changes the forward has its own file there, so that path gets - a minutes-scale gate too. + Catches catastrophes, not accuracy. A variant that changes the forward gets + its own minutes-scale gate on the same footing. 2. **Accuracy** — the release criterion. `trtllm-eval` with the protocol from - `references/accuracy.yaml` and `TRTLLM_STAIRCASE=require` exported. One-sided: measured >= reference − tol. - This is the gate CI runs, as `accuracy/test_staircase.py`. + `tests/integration/defs/accuracy/references/` and `TRTLLM_STAIRCASE=require` + exported. One-sided: measured >= reference − tol. This is the gate CI runs, + in the accuracy suite's staircase files. 3. **Acceptance** — required whenever a variant is distribution-preserving by construction, speculative decoding above all. Rejection sampling holds the emitted distribution to the target model's, so a *miscomputed* draft path @@ -183,13 +186,16 @@ never executed.** Two things voided every receipt in the move: each catalog test file was rewritten, and the targets moved from sm_100 (B200) to sm_103 (GB300), where certification is per architecture. The whole catalog was therefore re-run on -GB300 under 1.3.0rc26 -- **19 of 19 entries pass, 312 certified cells**, plus -both 4-rank collective matrices. +GB300 -- **19 of 19 entries pass, 312 certified cells**, plus both 4-rank +collective matrices. The sm_100 receipts were dropped rather than carried: +they predate every file of the entries they sat in, and no sm_100 machine is +in CI to re-take them on. A missing arch key reads as unknown, which is the +honest state. Getting there surfaced four real differences. None was resolved by widening a tolerance, and each is written up in its own contract: -* **Op schema drift rc21 -> rc26** (`thop_attention`, `mla_rope_generation`, +* **Op schema drift** (`thop_attention`, `mla_rope_generation`, `mla_rope_append_paged_kv_assign_q`). Parameters were renamed and added. The wrappers now mirror their schemas argument for argument, so the next drift fails loudly rather than shifting a positional list silently. @@ -201,10 +207,8 @@ tolerance, and each is written up in its own contract: * **The MLA append op now accepts NVFP4 latent pools** as well as fp8 (accepted by the op, not certified here). -**gpt-oss-120b / sm_103 / tp1 is gated on GB300.** Its checkpoint was -downloaded and every digest re-verified against `TARGET.md`, so the new records -and the old sm_100 ones were measured on byte-identical weights. Both gates -pass with `TRTLLM_STAIRCASE=require` in force, which is what rules out the built-in +**gpt-oss-120b / sm_103 / tp1 is gated on GB300.** Both gates pass with +`TRTLLM_STAIRCASE=require` in force, which is what rules out the built-in implementation having been measured instead: | Gate | Result | @@ -219,40 +223,42 @@ implementation having been measured instead: |---|---| | boot, 10 greedy continuations | **10/10** | | gsm8k, full 1319 | **95.0720** vs threshold 89.9962 -- pass by 5.08 | -| boot, `configs/mtp3.yaml` | **10/10**, with the layer-61 MTP module loaded | +| boot, mtp3 | **10/10**, with the layer-61 MTP module loaded | | gsm8k, identity vs mtp3 **paired in one session** | delta **-0.3791** against a `\|delta\| < 1.2` criterion -- 0.63 sigma | | acceptance vs stock | `acceptance_length` **3.3514** vs **3.2752**, ratio **1.023** | That last row is the one that matters for a speculative variant: rejection sampling makes a miscomputed draft layer *slower*, not wrong, so boot and accuracy are blind to it and only the acceptance rate against a reference can -see it. `configs/mtp{1,2}.yaml` remain ungated -- they are the dominated end of +see it. Draft lengths 1 and 2 remain ungated -- they are the dominated end of the measured draft-length axis. -Measured statements throughout the contracts are left exactly as written. -They are true records of what was observed on sm_100, and rewriting them would -manufacture GB300 evidence that does not exist. - ### What replaced the version pin -Out of tree each target hard-asserted `tensorrt_llm == 1.3.0rc21` at import, -because a target reads private engine surface and a drifting engine silently -voids its records. In tree that assert is meaningless — the target moves with -the trunk — so it is gone, replaced by an SM assert, which is the part of the -identity that does *not* move. +Out of tree each target hard-asserted a pinned `tensorrt_llm` version at +import, because a target reads private engine surface and a drifting engine +silently voids its records. In tree that assert is meaningless — the target +moves with the trunk — so it is gone, replaced by an SM assert, which is the +part of the identity that does *not* move. + +The same reasoning retires the version from the records themselves. A receipt +is keyed by GPU architecture and nothing else: there is no external pin left to +drift against, and a version written into a contract is a number that is wrong +the next day with nothing to notice. What a receipt is still fresh against is +its own entry's files, and that CI re-proves on every commit. -The pin was earning its keep, though, and the migration paid the bill -immediately: between rc21 and rc26 the attention backends moved from +The pin was earning its keep out of tree, though, and the migration paid the +bill immediately: the attention backends moved from `_torch/attention_backend/{interface,trtllm}.py` to `_torch/attention/backends/`. The compatibility shim left behind re-exports the names but not the submodule paths, so all five files importing them failed at import. They now use the canonical path. The lesson generalizes: with no pin, a target's contact with private engine -surface is checked only by running it. The import-time -`_check_static_contract` op-existence loop and the first-forward metadata -field check are what turn that from a wrong answer into a loud failure, which -is why both survived the move. +surface is checked only by running it. Each target declares that surface as +`REQUIRED_TRTLLM_OPS`, `test_staircase_target_contract.py` asserts every name +in it exists, and the first-forward metadata field check catches the rest -- +which turns a drifting engine from a wrong answer into a loud failure. ## Two facts this migration surfaced about upstream diff --git a/tensorrt_llm/_torch/staircase/__init__.py b/tensorrt_llm/_torch/staircase/__init__.py index 155310cbcb67..ddf929d3e050 100644 --- a/tensorrt_llm/_torch/staircase/__init__.py +++ b/tensorrt_llm/_torch/staircase/__init__.py @@ -29,17 +29,21 @@ """ from ._router_index import ( + ROUTING_BACKENDS, STAIRCASE_ENV, STAIRCASE_ROUTERS, StaircaseContext, StaircaseMode, + assert_backend_can_route, staircase_resolve, ) __all__ = [ + "ROUTING_BACKENDS", "STAIRCASE_ENV", "STAIRCASE_ROUTERS", "StaircaseContext", "StaircaseMode", + "assert_backend_can_route", "staircase_resolve", ] diff --git a/tensorrt_llm/_torch/staircase/_router_index.py b/tensorrt_llm/_torch/staircase/_router_index.py index 3c62838bb6d0..4bb8b8569e64 100644 --- a/tensorrt_llm/_torch/staircase/_router_index.py +++ b/tensorrt_llm/_torch/staircase/_router_index.py @@ -59,6 +59,13 @@ "DeepseekV3ForCausalLM": "models.deepseek_v3.routing", } +#: Backends whose model construction reaches ``staircase_resolve``. The +#: resolver is called from ``AutoModelForCausalLM._resolve_class``, so a +#: backend that builds its model some other way never consults it -- AutoDeploy +#: goes through ``ADEngine.build_from_config`` and nothing under +#: ``_torch/auto_deploy/`` mentions staircase at all. +ROUTING_BACKENDS = frozenset({"pytorch"}) + class StaircaseMode(str, enum.Enum): """What to do when a config reaches the staircase resolver.""" @@ -90,6 +97,32 @@ def from_env(cls) -> "StaircaseMode": ) from None +def assert_backend_can_route(backend: str) -> None: + """Refuse ``require`` on a backend that never reaches the resolver. + + ``require`` is a promise that the run measured a staircase target. A + backend outside ``ROUTING_BACKENDS`` cannot keep it: nothing raises, + nothing routes, and the run reports another implementation's numbers as + the target's. That is the same misattribution ``from_env`` already refuses + for a typo'd value, arriving by a different door. + + ``auto`` is left alone on purpose -- it licenses the non-staircase path by + definition, so taking it is the documented outcome rather than a silent + one. + """ + if StaircaseMode.from_env() is not StaircaseMode.REQUIRE: + return + if backend in ROUTING_BACKENDS: + return + raise ValueError( + f"{STAIRCASE_ENV}=require, but the {backend!r} backend never reaches " + f"the staircase resolver, so no target can be selected and nothing " + f"would report that. Backends that route: " + f"{', '.join(sorted(ROUTING_BACKENDS))}. Use one of those, or unset " + f"{STAIRCASE_ENV}." + ) + + @dataclass(frozen=True) class StaircaseContext: """Everything a routing decision is allowed to depend on. @@ -104,8 +137,8 @@ class StaircaseContext: That is a necessary condition, not a sufficient one. ``max_num_tokens`` and ``cuda_graph_config.max_batch_size`` are per-instance constants too, - but they are tuning knobs: they belong in a ``configs/`` variant. Only an - instance constant that changes the *structure of the forward* earns a + but they are tuning knobs: they are LLM API arguments, not identity. Only + an instance constant that changes the *structure of the forward* earns a target of its own. ``is_disagg`` is declared but **not yet plumbed**: nothing sets it on @@ -127,17 +160,29 @@ class StaircaseContext: is_disagg: bool @classmethod - def from_model_config(cls, config: "ModelConfig") -> "StaircaseContext": - import torch - - assert torch.cuda.is_available(), ( - "staircase routes on the SM version of the device it will run on; " - "no CUDA device is visible" - ) + def from_model_config( + cls, config: "ModelConfig", sm: Optional[Tuple[int, int]] = None + ) -> "StaircaseContext": + """Build the context the resolver routes on. + + ``sm`` defaults to the device this process will run on, which is what + the engine wants. ``explain`` passes it explicitly so a configuration + can be explained from a host with no GPU -- every other field it reads + off the same ``ModelConfig`` the engine built, rather than restating + one, so the two cannot disagree about what a checkpoint is. + """ + if sm is None: + import torch + + assert torch.cuda.is_available(), ( + "staircase routes on the SM version of the device it will run on; " + "no CUDA device is visible" + ) + sm = torch.cuda.get_device_capability() return cls( pretrained_config=config.pretrained_config, mapping=config.mapping, - sm=torch.cuda.get_device_capability(), + sm=sm, quant_config=config.quant_config, spec_config=config.spec_config, is_disagg=getattr(config, "is_disagg", False), diff --git a/tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul.md b/tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul.md index 4121d98883d4..2bc9c9720d67 100644 --- a/tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul.md +++ b/tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul.md @@ -1,7 +1,6 @@ --- receipts: - sm_100: {status: passed, trtllm: 1.3.0rc21} - sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 4} + sm_103: {status: passed, tests: 4} --- # flashinfer_silu_and_mul @@ -67,8 +66,7 @@ None. Stateless. ## Notes - The op is registered only when flashinfer is importable - (`IS_FLASHINFER_AVAILABLE`); this pinned install ships - flashinfer-python 0.6.14. + (`IS_FLASHINFER_AVAILABLE`). - Alignment trap: the op's Python-side check only validates `x.shape[-1] * itemsize % 16 == 0` (raising `ValueError`), but the vectorized load of the up half starts at element offset `d`, so `d` diff --git a/tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul.py b/tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul.py index 611c999e0338..7817e4effc85 100644 --- a/tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul.py +++ b/tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul.py @@ -9,4 +9,22 @@ def flashinfer_silu_and_mul(x: torch.Tensor) -> torch.Tensor: """Return `silu(x[..., :d]) * x[..., d:]` with `d = x.shape[-1] // 2` as a new tensor.""" + # The op's own check is on the *row*, `x.shape[-1] * itemsize % 16 == 0`, + # which a final dimension of 24 passes. The kernel vectorizes over the two + # halves, so what has to be 16-byte aligned is the *half*: at 24 the + # second half starts mid-vector and the launch dies with `CUDA misaligned + # address`, poisoning the context rather than raising. This is the + # stricter precondition the contract states. + width = x.shape[-1] + assert width % 2 == 0, f"x.shape[-1] must be even to split in half; got {width}" + half_bytes = (width // 2) * x.element_size() + assert half_bytes % 16 == 0, ( + f"x.shape[-1] // 2 must be a whole number of 16-byte vectors; " + f"{width} halves to {half_bytes} bytes, which is not a multiple of 16 " + f"-- the kernel would fault on the misaligned second half" + ) + assert half_bytes >= 16, ( + f"x.shape[-1] // 2 must hold at least one 16-byte vector; {width} " + f"gives {half_bytes} bytes, for which the computed block size is 0" + ) return torch.ops.trtllm.flashinfer_silu_and_mul(x) diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope.md b/tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope.md index 156153e4bf79..33a26545df88 100644 --- a/tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope.md +++ b/tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope.md @@ -1,7 +1,6 @@ --- receipts: - sm_100: {status: passed, trtllm: 1.3.0rc21} - sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 9} + sm_103: {status: passed, tests: 9} --- # fused_qk_norm_rope @@ -46,8 +45,7 @@ then irrelevant). `attention_factor` scales cos/sin — but **not at `factor == 1.0` exactly**, where any `attention_factor != 1.0` raises `Assertion failed: attention_factor == 1.0f` (`fusedQKNormRopeKernel.cu:322`). `factor = 1.0000001` accepts it. -Measured 2026-07-28; the gate is on `factor` alone and is not otherwise -documented. +The gate is on `factor` alone and is not otherwise documented. Interleaved mRoPE (`use_mrope=True`) takes 3 position rows (temporal/height/width) per token. Frequency index `j` reads its position diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla.md b/tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla.md index df9ca8b988ca..e51fee753fe6 100644 --- a/tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla.md +++ b/tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla.md @@ -1,7 +1,6 @@ --- receipts: - sm_100: {status: passed, trtllm: 1.3.0rc21} - sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 18} + sm_103: {status: passed, tests: 18} --- # load_paged_kv_cache_for_mla diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q.md b/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q.md index 8b31d6ad40bb..4be4de772b8e 100644 --- a/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q.md +++ b/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q.md @@ -1,7 +1,6 @@ --- receipts: - sm_100: {status: passed, trtllm: 1.3.0rc21} - sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 13} + sm_103: {status: passed, tests: 13} --- # mla_rope_append_paged_kv_assign_q @@ -177,7 +176,7 @@ above: | `host_kv_cache_pool_pointers` | `[num_pools, 2]`: (primary ptr, secondary ptr=0) | int64 | contiguous | CPU | | `host_kv_cache_pool_mapping` | `[num_layers, 2]`: (pool index, layer-within-pool) row per layer | int32 | contiguous | CPU | | `kv_scale_orig_quant` | `quant_mode=0`: `None`. `quant_mode=128`: `[1]` holding the **write-side** factor `w` — everything this call quantizes is multiplied by it. `None` = 1.0, which is what the engine's own MLA call site passes. Only element `[0]` is read (a longer tensor is accepted) | fp32 (fp16 rejected, see *Notes*) | contiguous | CUDA | -| `residual_dim` | — | `int` | **rc26 addition.** Must be `0` or `rope_size`; the op rejects non-zero unless the KV pool is FP4. `0` on every path this entry certifies (bf16 and fp8-e4m3 pools), which is also what the in-tree caller passes | — | +| `residual_dim` | — | `int` | Must be `0` or `rope_size`; the op rejects non-zero unless the KV pool is FP4. `0` on every path this entry certifies (bf16 and fp8-e4m3 pools), which is also what the in-tree caller passes | — | | `layer_idx` | row into `host_kv_cache_pool_mapping` | Python int | — | — | | `tokens_per_block` | pool page size; 32 and 64 certified on the matching-dtype pool, 32 on the fp8 pool (32 is what a default `KvCacheConfig` produces) | Python int | — | — | | `attention_window_size` | `>= max(kv_s)`; production passes the manager's `max_seq_len` (smaller values imply cyclic-cache addressing, not certified) | Python int | — | — | @@ -365,11 +364,9 @@ arguments. The length/addressing tensors are exactly what a *content* is not an axis this op's gate exercises). -## rc26 change to the accepted KV-cache formats +## Accepted KV-cache formats -The op accepted only an fp8-e4m3 latent pool in 1.3.0rc21 and now also -accepts NVFP4; its rejection message changed accordingly from -`Only FP8 KV cache is supported for now` to `Only FP8 and NVFP4 KV -caches are supported for now`. An int8 pool (`quant_mode=64`) is still -rejected. **NVFP4 is accepted by the op but not certified here** — no -cell in this entry's test drives it, so it stays outside the envelope. +The op takes an fp8-e4m3 latent pool and an NVFP4 one; an int8 pool +(`quant_mode=64`) is rejected with `Only FP8 and NVFP4 KV caches are +supported for now`. **NVFP4 is accepted by the op but not certified here** — +no cell in this entry's test drives it, so it stays outside the envelope. diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q.py b/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q.py index d88e22599ebc..e7f39b6182bd 100644 --- a/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q.py +++ b/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q.py @@ -25,7 +25,7 @@ def mla_rope_append_paged_kv_assign_q( host_kv_cache_pool_pointers: torch.Tensor, host_kv_cache_pool_mapping: torch.Tensor, kv_scale_orig_quant: Optional[torch.Tensor], - # ``residual_dim`` (rc26; absent in rc21) must be 0 or ``rope_size``, + # ``residual_dim`` must be 0 or ``rope_size``, # and the op rejects non-zero unless the KV pool is FP4. Every caller # here runs a bf16 or fp8-e4m3 pool, so 0 is the only legal value. residual_dim: int, diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation.md b/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation.md index e234b5a642ef..90f743a2b19f 100644 --- a/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation.md +++ b/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation.md @@ -1,7 +1,6 @@ --- receipts: - sm_100: {status: passed, trtllm: 1.3.0rc21} - sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 22} + sm_103: {status: passed, tests: 22} --- # mla_rope_generation @@ -410,11 +409,10 @@ arguments. The length/addressing tensors are exactly what a prepared other than bf16. -## rc26 additions +## Parameters this entry pins to their defaults -Parameters that did not exist in 1.3.0rc21. Every value this entry certifies -reproduces the op's pre-rc26 behaviour, and matches what the in-tree caller -passes on the same path. +Every value certified below leaves the op on its default behaviour, and +matches what the in-tree caller passes on the same path. | Parameter | Certified value | Why | |---|---|---| diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation.py b/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation.py index 20865dc2ba53..ba5d77856847 100644 --- a/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation.py +++ b/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation.py @@ -29,8 +29,8 @@ def mla_rope_generation( host_kv_cache_pool_mapping: Optional[torch.Tensor], kv_scale_orig_quant: Optional[torch.Tensor], kv_scale_quant_orig: Optional[torch.Tensor], - # rc26: when None the op falls back to kv_scale_orig_quant, which is what - # it did before this parameter existed (dsv3RopeOp.cpp:280). + # When None the op falls back to kv_scale_orig_quant + # (dsv3RopeOp.cpp:280). kv_cache_scale_orig_quant: Optional[torch.Tensor], out_scale: Optional[torch.Tensor], block_ids_per_seq: Optional[torch.Tensor], @@ -40,7 +40,7 @@ def mla_rope_generation( num_heads: int, num_kv_heads: int, head_size: int, - # rc26: 0 or rope_size, and non-zero requires an FP4 KV pool. + # 0 or rope_size, and non-zero requires an FP4 KV pool. residual_dim: int, tokens_per_block: int, attention_window_size: int, @@ -53,8 +53,8 @@ def mla_rope_generation( qk_rope_head_dim: int, v_head_dim: int, rope_append: bool, - # Added in rc26; every default below reproduces the op's pre-rc26 - # behaviour. kv_norm_weight non-None would fold the kv_a_layernorm into + # Every default below leaves the op on its default behaviour. + # kv_norm_weight non-None would fold the kv_a_layernorm into # this kernel, which then reads latent_cache RAW -- a caller that already # normalized would be normalizing twice. kv_norm_weight: Optional[torch.Tensor] = None, diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/thop_attention.md b/tensorrt_llm/_torch/staircase/catalog/attention/thop_attention.md index 534da795ac24..b5852e6791fb 100644 --- a/tensorrt_llm/_torch/staircase/catalog/attention/thop_attention.md +++ b/tensorrt_llm/_torch/staircase/catalog/attention/thop_attention.md @@ -1,7 +1,6 @@ --- receipts: - sm_100: {status: passed, trtllm: 1.3.0rc21} - sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 96} + sm_103: {status: passed, tests: 96} --- # thop_attention diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/thop_attention.py b/tensorrt_llm/_torch/staircase/catalog/attention/thop_attention.py index 9ee461204103..17815fdf2c96 100644 --- a/tensorrt_llm/_torch/staircase/catalog/attention/thop_attention.py +++ b/tensorrt_llm/_torch/staircase/catalog/attention/thop_attention.py @@ -127,8 +127,7 @@ def thop_attention( quant_scale_qkv: Optional[torch.Tensor] = None, dsv4_inv_rope_cos_sin_cache: Optional[torch.Tensor] = None, enable_dsv4_epilogue_fusion: bool = False, - # Added between 1.3.0rc21 and 1.3.0rc26. Defaults reproduce the behaviour - # the op had before they existed, and match what the in-tree caller + # Defaults match what the in-tree caller # (attention/backends/fmha/fallback.py) passes on a dense, non-sparse, # non-folded path -- which is the path both migrated targets are on. max_num_sequences: Optional[int] = None, diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/allgather.md b/tensorrt_llm/_torch/staircase/catalog/comm/allgather.md index dd289ecc9652..82072d3843a6 100644 --- a/tensorrt_llm/_torch/staircase/catalog/comm/allgather.md +++ b/tensorrt_llm/_torch/staircase/catalog/comm/allgather.md @@ -1,7 +1,6 @@ --- receipts: - sm_100: {status: passed, trtllm: 1.3.0rc21, world_size: 4} - sm_103: {status: passed, trtllm: 1.3.0rc26, world_size: 4} + sm_103: {status: passed, world_size: 4} --- # allgather diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/reducescatter.md b/tensorrt_llm/_torch/staircase/catalog/comm/reducescatter.md index 3f5ca642d0b5..c2512fb0a0b4 100644 --- a/tensorrt_llm/_torch/staircase/catalog/comm/reducescatter.md +++ b/tensorrt_llm/_torch/staircase/catalog/comm/reducescatter.md @@ -1,7 +1,6 @@ --- receipts: - sm_100: {status: passed, trtllm: 1.3.0rc21, world_size: 4} - sm_103: {status: passed, trtllm: 1.3.0rc26, world_size: 4} + sm_103: {status: passed, world_size: 4} --- # reducescatter diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out.md b/tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out.md index e7a5d35ec814..39e4685606d7 100644 --- a/tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out.md +++ b/tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out.md @@ -1,7 +1,6 @@ --- receipts: - sm_100: {status: passed, trtllm: 1.3.0rc21} - sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 5} + sm_103: {status: passed, tests: 5} --- # bmm_out @@ -66,19 +65,19 @@ None. Stateless. - `out.shape == (B, M, N)` exactly. **The op does not validate this**: a wrong-shaped `out` is silently resized (with only a deprecation warning). The hazard is real but is **shape/stride corruption, not a - lost write**: measured 2026-07-28, an aliased view with room to grow - keeps its `data_ptr` and the result *does* land in the caller's buffer, - now under a silently rewritten shape; a genuine reallocation moves the - shared storage, so sibling views follow it rather than being detached. - The wrapper asserts the shape. + lost write**: an aliased view with room to grow keeps its `data_ptr` and + the result *does* land in the caller's buffer, now under a silently + rewritten shape; a genuine reallocation moves the shared storage, so + sibling views follow it rather than being detached. The wrapper asserts + the shape. - One dtype across `a`, `b`, `out`. **Mixed input dtypes always raise** — there is no promotion path and no silent hazard here. The meta check demands `out` in `b.dtype` while the kernel demands `a.dtype`, so when `a.dtype != b.dtype` the two can never both be satisfied: all six combinations over {bf16, fp16, fp32} raise, including the bf16-`a` / fp32-`b` / fp32-`out` case an earlier revision of this bullet described - as working (measured 2026-07-28). The wrapper's single-dtype assert is - therefore redundant rather than load-bearing. + as working. The wrapper's single-dtype assert is therefore redundant + rather than load-bearing. - `out.dtype` must equal the input dtype; a mismatch raises (`Expected out tensor to have dtype ...`). - Dtypes verified: bf16, fp16, fp32. float8_e4m3fn raises @@ -97,5 +96,5 @@ None. Stateless. - Not arch-gated; TRT-LLM uses it as the bf16 batched-gemm path in MLA weight-absorption and output projections, with the batch dim carrying head groups. -- Behavior above was established empirically under trtllm 1.3.0rc21 / - torch 2.11.0 on sm_100. +- Behavior above was established empirically, not read off the op's + documentation; this entry's test is what holds it. diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm.md b/tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm.md index de74e6a9b5a1..d2ca81837268 100644 --- a/tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm.md +++ b/tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm.md @@ -1,7 +1,6 @@ --- receipts: - sm_100: {status: passed, trtllm: 1.3.0rc21} - sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 7} + sm_103: {status: passed, tests: 7} --- # cublas_mm @@ -105,7 +104,7 @@ a plain allocation). The enum is importable as single-pass cluster-mode kernels for small-M (decode) gemms there; the op itself is not gated on arch. - All silently-wrong-result behaviors listed under Preconditions were - observed under trtllm 1.3.0rc21 on sm_100. + observed directly, not inferred from the op's documentation. ## The fp32 cell needs its reference pinned, not its tolerance widened diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm.md b/tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm.md index b326f8d640f4..919bcfbca639 100644 --- a/tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm.md +++ b/tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm.md @@ -1,7 +1,6 @@ --- receipts: - sm_100: {status: passed, trtllm: 1.3.0rc21} - sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 19} + sm_103: {status: passed, tests: 19} --- # nvfp4_gemm diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm.py b/tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm.py index e5ef04ba8272..52f439ac64cb 100644 --- a/tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm.py +++ b/tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm.py @@ -7,6 +7,19 @@ import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* +def _pad_up(value: int, multiple: int) -> int: + return -(-value // multiple) * multiple + + +def _swizzled_scale_numel(rows: int, k: int) -> int: + """Bytes a 128x4 swizzled block-scale buffer holds for `rows` x `k`. + + The swizzle addresses the padded rectangle, not the real one, so the + buffer is this size whatever `rows` and `k` are -- see nvfp4_gemm.md. + """ + return _pad_up(rows, 128) * _pad_up(k // 16, 4) + + def nvfp4_gemm( act_fp4: torch.Tensor, weight: torch.Tensor, @@ -37,6 +50,30 @@ def nvfp4_gemm( "alpha must hold exactly one element; extra elements are silently " "ignored (this build has no per-token alpha)" ) + # Contiguity above says nothing about length, and a short scale buffer is + # read past its end rather than rejected: the kernel indexes the *padded* + # rectangle. Both operands are checked because M and N pad independently. + # + # Only for 2-D operands, and only for *under*-length. Anything else about + # the shapes -- wrong rank, zero rows, K disagreement -- the op rejects + # itself and loudly, and pre-empting that here would swap its RuntimeError + # for an AssertionError that says less. Over-length is not the hazard + # either: the kernel never reads past what the swizzle addresses, and the + # contract already says the padding bytes may hold anything. + if act_fp4.dim() == 2 and weight.dim() == 2: + k = act_fp4.shape[1] * 2 # act_fp4 packs two 4-bit values per byte + floor_act_sf = _swizzled_scale_numel(act_fp4.shape[0], k) + floor_weight_scale = _swizzled_scale_numel(weight.shape[0], k) + assert act_sf.numel() >= floor_act_sf, ( + f"act_sf must hold at least pad_up(M,128) * pad_up(K/16,4) = " + f"{floor_act_sf} bytes for M={act_fp4.shape[0]}, K={k}; got " + f"{act_sf.numel()} -- the kernel would read past its end" + ) + assert weight_scale.numel() >= floor_weight_scale, ( + f"weight_scale must hold at least pad_up(N,128) * pad_up(K/16,4) = " + f"{floor_weight_scale} bytes for N={weight.shape[0]}, K={k}; got " + f"{weight_scale.numel()} -- the kernel would read past its end" + ) return torch.ops.trtllm.nvfp4_gemm( act_fp4, weight, diff --git a/tensorrt_llm/_torch/staircase/catalog/index.yaml b/tensorrt_llm/_torch/staircase/catalog/index.yaml index fefe0dd9fb54..594f3dc42482 100644 --- a/tensorrt_llm/_torch/staircase/catalog/index.yaml +++ b/tensorrt_llm/_torch/staircase/catalog/index.yaml @@ -30,20 +30,28 @@ # # ── RECEIPT STATUS: all 19 entries certified on sm_103 ──────────────────── # -# A receipt says this entry's test passed, on a stated GPU architecture, over -# files no newer than the run. Both anchors moved at once here -- the -# migration rewrote every test file, and the targets moved from sm_100 (B200) -# to sm_103 (GB300) -- so every sm_100 receipt was voided and the whole set -# was re-run on GB300 under 1.3.0rc26. **All 19 pass**; per-entry counts are -# in each contract's `receipts:` frontmatter. +# A receipt says this entry's test passed on a stated GPU architecture. The +# architecture is the whole key: it is a real axis -- see the +# mxe4m3_mxe2m1_block_scale_moe_runner note below, where two architectures +# compute bit-different results -- and CI cannot stand in for it unless CI +# runs on every one. Staying current with the trunk is not a receipt's job: +# there is no external pin left to drift against, and the catalog tests run +# in pre-merge, so the trunk proves itself on every commit. +# +# The migration rewrote every test file and moved the targets from sm_100 +# (B200) to sm_103 (GB300), so the whole set was re-run on GB300. **All 19 +# pass**; per-entry counts are in each contract's `receipts:` frontmatter. +# The sm_100 receipts were dropped rather than carried -- they predate every +# file of the entries they sat in, and no sm_100 machine is in CI to re-take +# them on. A missing arch key reads as unknown, which is the honest state. # # Getting there surfaced four real differences, none of which was fixed by # widening a tolerance: # # * thop_attention, mla_rope_generation, mla_rope_append_paged_kv_assign_q -# -- op schema drift rc21 -> rc26 (renamed and added parameters). The -# wrappers now mirror their schemas argument for argument, so the next -# drift fails loudly instead of shifting a positional list. +# -- op schema drift (renamed and added parameters). The wrappers now +# mirror their schemas argument for argument, so the next drift fails +# loudly instead of shifting a positional list. # # * mxe4m3_mxe2m1_block_scale_moe_runner -- the FC1 epilogue's MXFP8 block # scale is `floor(log2(amax))-8` on sm_100 and `ceil(log2(amax/448))` on diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner.md b/tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner.md index ef65be84bc04..a13c6255ddc4 100644 --- a/tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner.md +++ b/tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner.md @@ -1,7 +1,6 @@ --- receipts: - sm_100: {status: passed, trtllm: 1.3.0rc21} - sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 16} + sm_103: {status: passed, tests: 16} --- # fp4_block_scale_moe_runner diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/fused_moe.md b/tensorrt_llm/_torch/staircase/catalog/moe/fused_moe.md index c125b64334d2..2446edb87003 100644 --- a/tensorrt_llm/_torch/staircase/catalog/moe/fused_moe.md +++ b/tensorrt_llm/_torch/staircase/catalog/moe/fused_moe.md @@ -1,7 +1,6 @@ --- receipts: - sm_100: {status: passed, trtllm: 1.3.0rc21} - sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 23} + sm_103: {status: passed, tests: 23} --- # fused_moe diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md b/tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md index 1cee050f77be..56fd242811b3 100644 --- a/tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md +++ b/tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md @@ -1,7 +1,6 @@ --- receipts: - sm_100: {status: passed, trtllm: 1.3.0rc21} - sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 19} + sm_103: {status: passed, tests: 19} --- # mxe4m3_mxe2m1_block_scale_moe_runner diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op.md b/tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op.md index 1ec9ce5f3ce4..a70d6556364c 100644 --- a/tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op.md +++ b/tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op.md @@ -1,7 +1,6 @@ --- receipts: - sm_100: {status: passed, trtllm: 1.3.0rc21} - sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 16} + sm_103: {status: passed, tests: 16} --- # noaux_tc_op @@ -191,7 +190,7 @@ A caller violating none of the above gets the result described under ## Notes -- **Certified surface.** Passing on sm_100 / trtllm 1.3.0rc21: `num_tokens` +- **Certified surface.** Passing: `num_tokens` in `{0, 1, 2, 4, 7, 8, 16, 64, 128, 256, 512, 1024, 2048, 4096, 8192}`; `num_experts` in `{1, 2, 7, 8, 16, 32, 64, 72, 100, 128, 256, 257, 512, 1024}`; `topk` in `{0, 1, 2, 3, 4, 6, 8, 16, 31, 32}`; all eight accepted diff --git a/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm.md b/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm.md index 3ccab4df4b45..f1809bcb9d2b 100644 --- a/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm.md +++ b/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm.md @@ -1,7 +1,6 @@ --- receipts: - sm_100: {status: passed, trtllm: 1.3.0rc21} - sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 5} + sm_103: {status: passed, tests: 5} --- # flashinfer_fused_add_rmsnorm @@ -61,7 +60,7 @@ None. Stateless. would read `M = shape[0]`), but the guard is **loud, not silent**: the CuTe compiled-kernel argument check raises `ValueError: Mismatched Tensor on argument #0 ... expected ndim=2` before launch. Nothing is - silently skipped (measured 2026-07-28); the precondition stands, its + silently skipped; the precondition stands, its earlier justification did not. - `weight.shape == (hidden,)`, `weight.dtype == x.dtype`, contiguous. - Dtype is one of fp16, bf16, fp32. `float64`, `int8`, `uint8` and @@ -85,8 +84,7 @@ None. Stateless. ## Notes - The op is registered only when flashinfer is importable - (`IS_FLASHINFER_AVAILABLE`); this pinned install ships - flashinfer-python 0.6.14, which routes to the CuTe DSL kernel + (`IS_FLASHINFER_AVAILABLE`), and routes to the CuTe DSL kernel (`fused_add_rmsnorm_cute`); a CUDA JIT fallback exists behind `FLASHINFER_USE_CUDA_NORM=1` but is not what these receipts certify. - Programmatic dependent launch (PDL) is controlled by the env var diff --git a/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm.md b/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm.md index 4b1bdaab9f78..aa067e75e405 100644 --- a/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm.md +++ b/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm.md @@ -1,7 +1,6 @@ --- receipts: - sm_100: {status: passed, trtllm: 1.3.0rc21} - sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 6} + sm_103: {status: passed, tests: 6} --- # flashinfer_rmsnorm @@ -64,8 +63,7 @@ None. Stateless. ## Notes - The op is registered only when flashinfer is importable - (`IS_FLASHINFER_AVAILABLE`); this pinned install ships - flashinfer-python 0.6.14. + (`IS_FLASHINFER_AVAILABLE`). - Programmatic dependent launch (PDL) is controlled by the env var `TRTLLM_ENABLE_PDL` (default enabled) inside the trtllm custom op; it affects scheduling only, not results. diff --git a/tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize.md b/tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize.md index 76cb1ca311d1..f82b0b2eceaf 100644 --- a/tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize.md +++ b/tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize.md @@ -1,7 +1,6 @@ --- receipts: - sm_100: {status: passed, trtllm: 1.3.0rc21} - sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 21} + sm_103: {status: passed, tests: 21} --- # fp4_quantize @@ -202,7 +201,7 @@ None. Stateless — no runtime, attention metadata, or workspace. `5 * 448 / g` then clip to `+-6`; the rest quantize normally against the collapsed scale and still dequantize close to their true value — at 1.17x overshoot a hand-built block came back - `[6, 4, 2, 1, 0.5, 0, 0, 0]`, not all-`+-6` (measured 2026-07-28). + `[6, 4, 2, 1, 0.5, 0, 0, 0]`, not all-`+-6`. **Do not look for an all-`+-6`, signs-only block as the signature**: a mis-calibrated block looks ordinary, and only ~100x overshoot produces the saturated form. diff --git a/tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize.md b/tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize.md index 339212adb107..221d652b0c97 100644 --- a/tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize.md +++ b/tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize.md @@ -1,7 +1,6 @@ --- receipts: - sm_100: {status: passed, trtllm: 1.3.0rc21} - sm_103: {status: passed, trtllm: 1.3.0rc26, tests: 13} + sm_103: {status: passed, tests: 13} --- # mxfp8_quantize diff --git a/tensorrt_llm/_torch/staircase/docs/models/expert-weight-packing.md b/tensorrt_llm/_torch/staircase/docs/models/expert-weight-packing.md deleted file mode 100644 index 3d642df6b457..000000000000 --- a/tensorrt_llm/_torch/staircase/docs/models/expert-weight-packing.md +++ /dev/null @@ -1,553 +0,0 @@ -# Expert weight packing (sparse MoE) - -What a target assembled on a sparse-MoE checkpoint had to establish by -observation, because no contract or reference stated it. Written from the -qwen3-30b-a3b/sm_100/tp1 onboard (Qwen3-30B-A3B: 48 layers, all MoE, 128 -experts, top-8, `moe_intermediate_size` 768, bf16). The mechanism — a -router selecting k of E stacked expert MLPs — recurs across model -families; this file is about the mechanism, not that checkpoint. - -**Every shape below is a whole-stack shape, measured at world size 1.** -Under expert parallelism a rank holds `E / ep_size` of the stack and the -first dimension of every weight and scale shrinks accordingly. - -**What survives the split — established by the first expert-parallel -target** (`tep4`, 72 experts over 4 ranks, 18 per rank at offset -`18 * ep_rank`): - -* **Every per-expert preparation step survives unchanged.** The - `[up; gate]` concat, the interleave, the 32-row block shuffle and the - 128x4 scale swizzle all happen *within* one expert, and EP splits the - stack on the expert axis only. So every transform below applies - verbatim, with the first dimension `E → E/ep_size` and the destination - index `e - local_expert_offset`. Nothing about the shuffle or the - swizzle is a function of the stack height. -* **The three per-expert scale scalars must be `[local_num_experts]`, but - the checkpoint's scalars must still be read whole.** The runner rejects - `num_experts`-long scalars, so the window slice is real — but the - *validity* argument for the shared-expert activation scale is global - (`shared.input_scale == max` over all `E` routed `input_scale`), and a - target that windows at load time can no longer assert it. Load all `E` - (six fp32 per expert — negligible), assert globally, slice in the - post-load derivation. -* **`leftover == {}` stops being the coverage assert.** Under EP each rank - legitimately leaves the off-window experts' weight and scale keys - unconsumed. Replace it with an explicitly predicted leftover set - computed from the rank's window — that keeps the assert bidirectional - rather than relaxing it to a warning, which is the whole value of having - it. - -One cross-rank invariant a single rank cannot check, recorded because the -assembly rests on it: **the four windows must tile the routing space -exactly once.** They do because the router GEMM and the routing op are -replicated and deterministic and every rank sees identical all-reduced -inputs — so each token's top-k ids are the same on every rank, and each -id falls in exactly one window. - -**That justification is topology-specific, and it does not survive -attention data parallelism.** Under a `depN` segment the ranks hold -*different* tokens, so "every rank sees identical inputs" is false by -construction and the invariant loses its support — while remaining just as -load-bearing, because EP still requires each expert id to belong to exactly -one rank's window. Established by the first attention-DP MoE target -(`dep4`): what restores it is **placing the token all-gather before the -router GEMM, not merely before the expert call.** The router then runs on a -byte-identical full token set on every rank, the original argument applies -again verbatim, and the four windows tile as before. - -The failure this rules out is worth naming, because it is silent. Gathering -*after* the router — the arrangement that looks equivalent, and is cheaper -by the width of the routing tensors — leaves every rank routing only its -own tokens, so a token's expert ids exist on one rank alone and the other -three windows never see it. There is no error and no hang; the output is -simply missing most of its expert contributions. Both gates would be the -only thing that catches it. - -Stated as the rule: **on any topology where the ranks' token sets differ, -the collective that reconstitutes the token set belongs upstream of the -router, and the tiling invariant is what decides that placement** — not the -expert call's input requirements, which the later placement also satisfies. - -## The shape of the vocabulary - -A bf16 MoE block is three catalog calls, not one: - -``` -router_logits = cublas_mm(x, router_weight_view) # [T, E] -expert_ids, scales = renorm_moe_routing_op(router_logits, topk) -moe_out = fused_moe(x, expert_ids, scales, fc1, None, fc2, None, ...)[0] -``` - -`fused_moe` contains the permutation, both grouped GEMMs, the gated -activation and the routing-weighted combine. It does **not** contain the -router GEMM or the selection — those stay the caller's, which is why the -routing op exists as a separate entry. - -Consequence worth stating because it looks like an omission: a -fully-MoE checkpoint's forward references **no `activation/*` entry**. The -SwiGLU lives inside `fused_moe` (`activation_type=5`). An audit that -expects `silu_and_mul` in every model's forward is wrong for this class. - -## The `[up | gate]` order — the one silent-wrong-answer trap - -`fused_moe`'s `fc1_expert_weights` is `[E, 2I, H]` with **up rows first**: - -- `fc1[e, :I]` is HF `up_proj` (trtllm `w3`) -- `fc1[e, I:]` is HF `gate_proj` (trtllm `w1`) — the sigmoid is applied to - this half - -`fc2_expert_weights` is `[E, H, I]`, plain `down_proj`. - -Getting the halves backwards produces a plausible, entirely different -result: finite, correctly scaled, no error, no NaN. Nothing in the stack -detects it — not the op, not a smoke test. Only an accuracy gate catches -it, and only because it lands ~50 points low. - -**This is the opposite of the dense convention.** A dense target packs -`gate_up` with gate rows first, because that is what -`flashinfer_silu_and_mul` expects over a packed last dim. The two orders -must never be copied across; each is correct for its own consumer. - -## Checkpoint layout, and where the key names actually live - -A 4.51-era checkpoint stores experts **unstacked**: 128 × 3 separate 2D -tensors per layer, -`model.layers.{i}.mlp.experts.{e}.{gate,up,down}_proj.weight`. Recent -transformers models the same block with stacked 3D `gate_up_proj` / -`down_proj` parameters — so **the installed modeling file is not the -source of truth for checkpoint key names**. The safetensors index is. -Read `model.safetensors.index.json`, not `modeling_*.py`, when writing -the manifest. - -The unstacked on-disk layout reaches `fused_moe`'s stacked layout with -**zero transforms** — every copy is layout-preserving: - -```python -(f"{p}.mlp.experts.{e}.up_proj.weight", (e, slice(0, inter))), -(f"{p}.mlp.experts.{e}.gate_proj.weight", (e, slice(inter, 2 * inter))), -(f"{p}.mlp.experts.{e}.down_proj.weight", (e,)), -``` - -This needs one generalization of the manifest convention: the destination -widens from a 2D `(row_start, row_end)` range to an index tuple into -`param.data`, so one table serves both fused-2D parameters (qkv) and -stacked-3D ones (experts). It is a one-line change in `fill`. - -3D parameters need no special handling anywhere else: they pass through -`MetaInitMode` and engine materialization unchanged. Measured at this -scale — 434 `ParameterDict` entries, 58 GB of expert stacks, 18432 -per-expert copies of ~3.1 MB — the whole load sits inside a 14.9 s model -init. - -## What the structure dictates, and what it makes irrelevant - -Established by measurement during the perf campaign; these are properties -of the mechanism, so the next MoE target can skip the sweeps. - -**The MoE is HBM-roofline-bound at serving batch sizes.** Above a decode -batch of roughly `E / topk × small factor` (~40 here) essentially every -expert is active, so a decode step reads the **entire** expert stack once -— 58 GB for this geometry. Measured 7.350 ms/step for the two grouped -GEMMs at concurrency 256 ⇒ **7.89 TB/s**, i.e. the B200 HBM roofline. No -config knob and no scheduling change moves this. Only fewer weight bytes -(quantization — a modeling change) would. - -**The kernel-count floor is dominated by the MoE.** `fused_moe` plus -`renorm_moe_routing_op` contribute 7 of the 16 kernels per layer -(`customMoeRouting`, `fusedBuildExpertMapsSortFirstToken`, -`expandInputRows`, grouped GEMM 1, `doActivation`, grouped GEMM 2, -`computeStridesTmaWarpSpecialized`). Across 48 layers that is 336 of 768 -kernels per decode step. At concurrency 1 the measured 3.87 ms TPOT over -794 kernels is 4.9 µs each — the same order as the smallest kernels -measured at batch 256, i.e. a fixed per-kernel floor rather than work. The -single-stream latency point is structural until the vocabulary gains a -coarser op. - -**KV capacity does not bind.** A 3B-active/30B-total model at tp1 leaves -the KV pool enormous relative to any realistic concurrency (measured: -1,168,096 pool tokens = 570 requests at ISL+OSL 2048, against a 256-request -ceiling). `kv_cache_config.free_gpu_memory_fraction` and the capacity -scheduler policy are inert — do not spend sweeps on them. - -**The autotuner is not a cold-cache risk inside a served engine.** -`fused_moe`'s contract warns that a cold tuning cache silently falls back -to a default tactic. In serving that does not apply: the runtime's -`_run_autotuner_warmup` wraps the forward in `autotune()` before capture -(observed `Cache size after warmup is 28` = 2 tunable GEMMs × 14 -power-of-2 token buckets). Worth knowing before anyone spends an iteration -on it — and worth knowing in the other direction too: the tuned tactic -changes the *bits*, not just the speed (measured 2.96 ulp against the cold -fallback on the same input), so a catalog receipt taken cold certifies a -tactic the served engine never runs. Both states sit inside the entry's -accuracy gate; the point is that "same inputs, same outputs" holds only -within one tuner state. - -## Routing traps - -Both from `renorm_moe_routing_op`'s certification: - -- **Ties break toward the lower expert index**, the opposite of - `torch.topk`. On bf16 logits exact ties are common, so a - `torch.topk`-based reference disagrees on indices while the weights - still match. Do not validate routing against `torch.topk`. -- **The kernel ignores strides**, reading `router_logits` as a dense - row-major buffer from `data_ptr()`. A strided view routes silently - wrong. The wrapper guards it; a caller building logits as a slice of a - wider buffer must materialize them contiguous. - -And from `fused_moe`: expert ids outside the rank's slot range are -**silently dropped**, not clamped and not rejected, so a routing bug -surfaces as a quietly weaker token rather than an error. - -## Deriving "is every layer MoE?" - -Do not assume uniformity, and do not assume `intermediate_size` is live. -The dense-vs-sparse branch per layer is -`layer_idx not in mlp_only_layers and num_experts > 0 and (layer_idx + 1) -% decoder_sparse_step == 0`. With `decoder_sparse_step: 1` and -`mlp_only_layers: []` every layer is sparse, which makes -`intermediate_size` **dead config** that no layer reads — while -`moe_intermediate_size` is the live one. A checkpoint with a nonzero -`decoder_sparse_step` or a non-empty `mlp_only_layers` needs both branches -and both sets of weights. Assert the derivation at construction rather -than trusting it. - -## MXFP4 expert stacks — the same trap, in interleaved form - -From the gpt-oss-120b/sm_100/tp1 onboard (36 layers, all MoE, 128 -experts, top-4, hidden = intermediate = 2880, experts MXFP4 while router, -attention, embedding and lm_head stay bf16). Everything above about the -`[up | gate]` order still holds; a block-scale-quantized stack adds three -things. - -**The parity is inverted relative to the half-split convention.** The -checkpoint stores `gate_up_proj_blocks` as `[E, 2I, K/32, 16]` uint8 — -already in `nn.Linear` `[out, in]` orientation, transposed relative to -HF's `[E, hidden, 2*inter]` parameter — and HF reads `gate = -gate_up[..., ::2]`, `up = gate_up[..., 1::2]`. So the stored row order -along the `2I` axis is (gate, up, gate, up, …), while the trtllm-gen -kernel's interleave wants destination row `2i` = **up** `i`, `2i+1` = -**gate** `i`. The split is therefore `up = t[:, 1::2]`, `gate = -t[:, 0::2]`, re-concatenated `[up ; gate]` before the row permutation. -Measured discrimination on real layer-0 tensors: 1.12 bf16 ulp against a -correct pure-torch reference, 146 ulp against the swapped one — a 130x -separation. Worth running that check *before* the first engine boot: it -covers nibble order, the parity split, the concat, the interleave, the -32-row block shuffle, the 128x4 scale swizzle, both padded axes, both -fp32 biases and the alpha/beta/limit triple in one shot, and the accuracy -gate that would otherwise catch it costs a full evaluation. - -**Where the relayout lives decides whether the model fits.** Expressing -pad → concat → interleave → shuffle → swizzle as *manifest source -transforms* (the sharding convention's `src` slot, generalized to any -callable) keeps peak memory at one layer of scratch. Declaring -checkpoint-shaped parameters and transforming them in `post_load_weights` -instead needs both forms resident — 63 GB checkpoint plus 66 GB -kernel-ready, on a 183 GB card that also wants a KV pool. The streaming -form loaded 63 GB including the on-device relayout in 14.1-14.5 s. The -kernel-ready operands at this geometry are `[128, 5888, 1536]` + -`[128, 5888, 96]` + fp32 `[128, 5888]` (FC1) and `[128, 2944, 1472]` + -`[128, 2944, 92]` + fp32 `[128, 2944]` (FC2) per layer, ~1.85 GB/layer, -+4% over the checkpoint. - -**Dtypes on disk are not the dtypes the kernels want.** Expert biases and -per-head attention sinks are bf16 in the checkpoint; the trtllm-gen MoE -requires fp32 biases and the attention op an fp32 sink tensor. Promote at -load. The same row permutation applied to the weight bytes must also be -applied to the scale bytes **and** the fp32 bias — a non-permuted bias is -a silent wrong answer. - -## W4A16 vs W4A8: the same weights, a different activation path - -Both members of the trtllm-gen block-scale MoE family consume the -**byte-identical** prepared expert stack — confirmed from the quant-method -source (both inherit one base, overriding only `create_weights` / -`load_quant_scales`, which call `super()`) and by measurement across every -certified geometry. Moving between them is a forward change only; -`weights.py` does not move. - -| | W4A16 | W4A8 | -|---|---|---| -| activations | bf16 straight in | `mxfp8_quantize(x, swizzled_layout=False, alignment=512)` first | -| hidden padding | caller pads 2880 -> 3072 | the quantizer does it — **the pad call disappears** | -| `valid_hidden_size` | 2880 | 2880 (output width; unrelated to the widening, and `None` is rejected) | - -Two consequences that are not obvious from the signatures: - -- **The W4A8 kernel requantizes the FC1 activation to MXFP8 between the - GEMMs**, on the OCP scale (`e = floor(log2 amax) - 8`), *not* the - round-up scale the standalone quantizer uses. Skipping that step in a - reference deviates 13.5 ulp element-wise / 8.5 ulp RMS (~3% relative) - from the kernel — so it is a real perturbation of the layer output, and - it is the mechanism behind the accuracy difference between the two - members. On gpt-oss-120b the measured cost was 1.44 GSM8K points - (90.5989 W4A16 -> 89.1585 W4A8), reproducible bit-identically on both - sides, i.e. the recipe's price rather than sampling. -- **`swizzled_layout=True` is silently accepted** by the W4A8 MoE - whenever the byte counts coincide (`T % 128 == 0` — exactly the CUDA - graph batch sizes), and is 260 ulp wrong. No metadata distinguishes the - two layouts, so no wrapper can guard it. Put the `False` in a named - constant. - -The perf reason to pay that accuracy: on gpt-oss-120b the W4A16 MoE ran -5.6x slower than the W4A8 one on byte-identical weights (FC1 711.8 vs -126.7 µs per layer per step at concurrency 128), which was the target's -*entire* gap to stock trtllm. The swap moved peak throughput +88% and put -the step within 0.5% of the reference. - -## The dtype chain, end to end - -Worth writing out because the names invite a wrong guess: **W4A8 -quantizes activations to fp8, never to fp4.** The op name spells both -operands — `mxe4m3_mxe2m1_...` is e4m3 activations against e2m1 weights. -The `mx` prefix is OCP micro-scaling: 32 elements share one E8M0 -(power-of-two) scale, so MXFP8 is 32 e4m3 values plus one scale byte and -MXFP4 is 32 e2m1 values plus one. - -| step | format | -|---|---| -| hidden states entering the block | bf16 `[T, H]` | -| `mxfp8_quantize(x, False, 512)` | e4m3 `[T, pad_up(H, 512)]` + E8M0, one byte per 32 | -| FC1 GEMM | e4m3 x e2m1, **fp32 accumulate** | -| clamped GLU (FC1 epilogue) | fp32 | -| requantization of the intermediate | e4m3 + E8M0 per 32 columns | -| FC2 GEMM | e4m3 x e2m1, **fp32 accumulate** | -| routing-weighted combine | fp32 | -| store | **bf16** `[T, valid_hidden_size]` | - -Two quantization points, both to fp8, with fp32 everywhere between them. -Only the operand *storage* formats are narrow — the arithmetic never drops -to fp4. The output is bf16 at the model's true hidden width, not the -padded one, and `valid_hidden_size` has to be passed explicitly to get it. - -## Why the W4A16 member is slow, and the W4A8 one is not - -The two paths read the **same weight bytes**, so the 5.6x is not -bandwidth. It is what each kernel does with them, and the kernel names say -it outright: - -| | operand fields in the kernel name | conversion | -|---|---|---| -| W4A16 | `bmm_Bfloat16_MxE2m1Bfloat16_castBfloat16_...` | `castBfloat16` — MXFP4 expanded to bf16, then a bf16 MMA | -| W4A8 | `bmm_MxE4m3_MxE2m1MxE4m3_...` | none — MXFP4 goes straight into a block-scaled MMA | - -The expansion's cost is not mainly the conversion arithmetic. A -dequantized weight block occupies **4x the on-chip bytes** (4-bit e2m1 -> -16-bit bf16), so fewer of them fit in registers and shared memory, which -shrinks the tile and shortens the software pipeline — visible in the same -kernel names: - -- W4A16: `m128x8x16`, 3 pipeline stages, 1-CTA clusters -- W4A8: `m256x16x32`, 6 stages, 2-CTA clusters - -Fewer weight bytes in flight means HBM latency stops being hidden, and the -GEMM lands far below the memory roofline instead of at it. At concurrency -128 on gpt-oss-120b, where the batch touches essentially every expert so a -decode step reads the whole stack once (~61 GB at valid sizes): - -| | MoE per step | implied rate | -|---|---|---| -| W4A16 | 38.89 ms | ~1.6 TB/s | -| W4A8 | 8.34 ms | ~7.3 TB/s | -| stock trtllm | 7.51 ms | ~8.1 TB/s | - -against a roofline this repo measured at 7.89 TB/s on the same machine. So -the W4A16 member was not paying more bandwidth — it was failing to use the -bandwidth it had, at roughly a fifth of the achievable rate. The residual -W4A8-vs-reference gap is autotuner tactic choice (~11%), not the recipe. - -*Measured:* the kernel names with their tile / stage / cluster fields, the -per-GEMM times (FC1 711.8 vs 126.7 µs, FC2 352.9 vs 65.6 µs per layer per -step), the MoE step times, and that the weights are byte-identical. -*Derived:* the byte figure behind the implied rates, and the chain from -on-chip footprint to tile size to unhidden latency — consistent with every -number above, but not isolated experimentally. - -## Blackwell only, and it fails quietly elsewhere - -This kernel family is sm_100 / sm_103. Three independent signals: - -- TensorRT-LLM gates its own tests on it — `get_sm_version() not in (100, - 103)`, reason "TRTLLM Gen MoE supports SM100 and SM103 only". -- The quantization kernel's body compiles only under `__CUDA_ARCH__ >= - 1000`; below that the empty branch **returns without writing the - outputs**, so a pre-Blackwell run produces garbage rather than failing. - Both catalog entries carry sm_100-only receipts for that reason — a - statement about the hardware, not caution. -- The performance signature itself. A software emulation of the mixed - fp8 x fp4 MMA — unpack, widen, then a wide MMA — would look exactly like - the W4A16 path measured above, and nothing can be 5.6x faster than the - thing it emulates. - -The hardware feature is Blackwell's block-scaled MMA: the instruction -takes narrow operands plus their per-block scale factors and applies the -scaling inside the tensor-core datapath, so no widened operand is ever -materialized. The E8M0 scale being a power of two is what makes that -nearly free — it is an exponent adjustment. (The arch gating and the -performance signature are measured here; the datapath description is the -standard account of the feature and was not verified at the ISA level.) - -## NVFP4 expert stacks — the same traps, plus per-expert scales - -From the deepseek-v3-lite-nvfp4/sm_100/tp1 onboard (DeepSeek-V3-Lite: 30 -layers, layer 0 dense, layers 1-29 MoE with 72 routed experts at top-6, -`moe_intermediate_size` 1536, plus 2 shared experts fused into one -`[3072, ...]` dense pair; experts NVFP4 while attention, router, embed -and lm_head stay bf16). Everything above about the `[up | gate]` order -still holds. NVFP4 differs from the MXFP4 case in three ways that matter. - -**The consumer is a different op with a different operand shape.** The -trtllm-gen NVFP4 MoE takes `[E, 2I, H/16]` **`float8_e4m3fn`** scales, -where the MXFP4 members take `[E, 2I, H/32]` uint8 UE8M0. The prepared -stacks are **not** interchangeable between the two families, even though -both are "block-scale MoE runners". - -| operand | shape | dtype | -|---|---|---| -| `gemm1_weights` | `[E, 2I, H/2]` | uint8 (two e2m1 codes per byte) | -| `gemm1_weights_scale` | `[E, 2I, H/16]` | float8_e4m3fn | -| `gemm2_weights` | `[E, H, I/2]` | uint8 | -| `gemm2_weights_scale` | `[E, H, I/16]` | float8_e4m3fn | - -Per-expert preparation is: concat `[up ; gate]` (**up rows first**) -> -interleave (dst row `2i` = up `i`, `2i+1` = gate `i`) -> 32-row block -shuffle (src `4u+v` -> dst `8v+u`) applied to weight bytes **and** scale -bytes -> 128x4 swizzle of the scales only. -`torch.ops.trtllm.shuffle_matrix` and -`torch.ops.trtllm.block_scale_interleave` perform the last two; both are -load-time transforms, outside the closed-vocabulary rule. - -Measured discrimination against a correct pure-torch reference, on real -checkpoint tensors — the reason to run this check before the first engine -boot rather than after an accuracy gate: - -| variant | bf16 ulp | -|---|---| -| **correct** | **4.75** | -| `[gate\|up]` instead of `[up\|gate]` | 263 | -| weight scales not 128x4 swizzled | 408 | -| expert rows not 32-row shuffled | 508 | -| `output1_scale_scalar` / `_gate_scalar` swapped | 177 | -| no intermediate requantization | 42.7 | - -**The half-order trap appears twice in one layer on a shared-expert -model.** The dense and shared-expert paths pack `[gate ; up]` for -`flashinfer_silu_and_mul`; the MoE FC1 packs `[up ; gate]`. Both live in -the same decoder layer, three lines apart. The two orders must never be -copied across — each is correct for its own consumer. - -### modelopt stores reciprocals - -A modelopt NVFP4 checkpoint stores `input_scale` and `weight_scale_2` as -`amax / (448 * 6)` — the **reciprocal** of the global scale a quantizer -divides by. So: - -``` -global_scale (for fp4_quantize) = 1 / input_scale -dense GEMM alpha = input_scale * weight_scale_2 # both off disk -``` - -Reading them the other way is wrong by `global_scale^2`, finite, and -invisible short of an accuracy gate. Settle it by arithmetic rather than -by convention: under the reciprocal reading the implied amax values are -0.2-3.0, physically sensible for normed activations and NVFP4 weights; -under the other reading they are ~1e7. `gate_proj` and `up_proj` share -both scalars exactly (verified across all experts of every MoE layer) — -assert it at load rather than assuming it. - -The trtllm-gen runner's three `[E]` fp32 arrays, in checkpoint terms: - -``` -output1_scale_gate_scalar[e] = input_scale_gate_up * weight_scale_2_gate_up[e] -output1_scale_scalar[e] = output1_scale_gate_scalar[e] / input_scale_down -output2_scale_scalar[e] = input_scale_down * weight_scale_2_down[e] -``` - -The gate scalar dequantizes the half feeding the sigmoid; the other -carries the extra FC2-input global scale because the FC1 epilogue -re-quantizes its output to NVFP4. - -### One activation scale, but per-expert weight scales - -Two facts that look like they should match and do not. - -**The activation global scale `g1` is one number, and the shared expert -carries it.** The runner quantizes the hidden states once, so it takes a -single `g1`. Which one? On this checkpoint the shared expert's -`gate_proj.input_scale` is **exactly** the max over all routed experts' -— exact equality in every MoE layer. That is not a coincidence worth -guessing at: the shared expert sees every token, so its amax *is* the -global amax. `g1 = 1 / shared.gate_proj.input_scale` therefore serves -both the dense and MoE quantization calls. Routed experts' own values sit -up to 1.4-2.3x below it. - -**The FC2 input scale `g2` is genuinely per-expert** and does not -collapse. It spans **54.5x** across experts in a single layer. The -runner's three `[E]` arrays exist precisely to carry it; feeding a scalar -would be wrong by that factor on the extreme experts. - -### Two quantization calls, not one - -`nvfp4_gemm` (dense and shared-expert linears) consumes the **swizzled** -scale-factor buffer; the trtllm-gen MoE runner consumes the **linear** -one, viewed as `float8_e4m3fn`. The same hidden states must therefore be -quantized **twice**, once in each layout — they cannot share a result. - -The failure mode is nasty: at `T % 128 == 0` the swizzled buffer happens -to be the *right size* for the runner, so it is silently accepted and is -276 ulp wrong. At every other token count it is rejected on size. A smoke -test at a round batch size will not catch it, and a smoke test at an odd -one will look like a size bug rather than a layout bug. - -### Two smaller obligations - -- **`hidden` must be a multiple of 256** for the trtllm-gen NVFP4 runner. - At a hidden that is only a multiple of 128 the kernel reads FC1 weight - scales for the wrong blocks: no error, 301-3370 bf16 ulp wrong. This is - an upstream hazard too — trtllm's own weight creation only raises - alignment to 256 when `hidden_size > 1024`. -- **Per-tensor quantization scalars are stored 0-dim.** `entry[:]` on - them raises `IndexError: slice() cannot be applied to a 0-dim tensor`, - so a manifest's materialize step needs a rank check that the bf16 path - never needed. - -## Group-limited routing, and where the grouping lives - -From `deepseek-r1-0528-nvfp4/sm_100/dep4` (256 routed experts, top-8, -`n_group 8` / `topk_group 4`), the first target on a grouped checkpoint — -its siblings all run `n_group: 1`. - -**The grouping is entirely the routing op's.** `noaux_tc_op` takes -`n_group`/`topk_group` and does the group-limited top-k; the block-scale -MoE runner's own `n_group`/`topk_group` arguments stay **`None`** at -`routing_method_type = 1`, because the runner is on its pre-routed path -and consumes finished `topk_ids`/`topk_weights`. Passing the grouping -twice is not how it composes. - -Two constraints the grouped path adds that the ungrouped one does not, and -that the routing op reports as one opaque "unsupported configuration": at -`n_group > 1` it requires `1 <= topk_group <= n_group`, `topk <= 8`, -`num_experts <= 256`, and `num_experts / n_group <= 32`. Assert them -separately or a violation is unattributable. - -**Router bias dtype.** The combine weights follow the **logits'** dtype, -not the bias's. A bf16 router GEMM with the fp32 `e_score_correction_bias` -this checkpoint ships returns bf16 weights — which is what the MoE runner -demands (it rejects fp32 `topk_weights`). No cast is needed and adding one -is a mistake. - -## A four-way expert-parallel split at 256 experts - -The invariant is the one this file already records — the windows must tile -the routing space exactly once — now measured at 64 wide: - -* the four 64-wide windows **sum to the whole-layer result** (0.85 ulp RMS - against an fp32 reference), but are **not bitwise equal** to a single - 256-expert call: up to 2.00 ulp apart, because each window rounds its own - partial to bf16 before the add. Dropping any one window lands ≥ 125 ulp - RMS away, so the sum is a real check rather than a formality; -* the **autotuner's key holds `local_num_experts` but not - `local_expert_offset`** — warming one window of a split warms all four, - and a 256 → 64 change forces a fresh sweep. Cold and warm agree bitwise - across 85 configurations here, but do not generalize that: the bf16 - `fused_moe` entry records the same check *failing*. diff --git a/tensorrt_llm/_torch/staircase/docs/models/multi-token-prediction.md b/tensorrt_llm/_torch/staircase/docs/models/multi-token-prediction.md deleted file mode 100644 index 9c9c8e5e40ab..000000000000 --- a/tensorrt_llm/_torch/staircase/docs/models/multi-token-prediction.md +++ /dev/null @@ -1,299 +0,0 @@ -# Multi-token prediction (MTP) - -What a target assembled on a checkpoint that ships an MTP module had to -establish about **what the extra layer computes**, because no contract -states it and the checkpoint ships no reference implementation for it. -Written from the deepseek-r1-0528-nvfp4/sm_100/dep4 MTP increment -(DeepSeek-R1-0528, 61 trunk layers, hidden 7168, 256 routed experts top-8, -one MTP module at layer 61). The mechanism — an extra decoder layer that -consumes *both* the trunk's hidden state and the next token's embedding, and -is replayed to produce a draft — recurs across the DeepSeek family and its -derivatives; this file is about the mechanism, not that checkpoint. - -The runtime half of the picture — how the engine finds the drafter, what -calls the layer, and how the batch shape changes between draft steps — is -`docs/references/trtllm-runtime-integration.md` §13. This file is only the -computation and the weights. - -**The checkpoint is the whole specification here, and it is a partial one.** -DeepSeek's published `modeling_deepseek.py` does **not** implement the MTP -module (verified across four R1/R1-0528 checkpoints on disk, quantized and -not), and `transformers`' own `deepseek_v3` does not either. So the graph -below was reconstructed from the checkpoint's key names, shapes and weight -statistics. Every claim that rests on the reconstruction rather than on a -shape is marked. - -## The mechanism - -The module is one extra decoder layer, structurally identical to a trunk MoE -layer, with a front end bolted on that mixes in the embedding of the token -being predicted: - -``` -e = embed_tokens(input_ids) # embedding of the NEXT token -x = eh_proj( concat( enorm(e), hnorm(h) ) ) # [T, 2*hidden] -> [T, hidden] -x = x + MLA( input_layernorm(x) ) # same block as a trunk layer -x = x + MoE( post_attention_layernorm(x) ) # same structure, different dtype -# shared_head, called separately by the runtime: -logits = lm_head( shared_head.norm(x) ) -``` - -`h` is the trunk's final hidden state for the same position — under -one-model MTP-Eagle the runtime hands the layer the target model's own -output, with no projection in between. - -**"How many MTP layers do we enable" is not the knob.** A checkpoint with -one MTP module (`num_nextn_predict_layers: 1`) is replayed autoregressively: -the same layer runs `max_draft_len` times, each step feeding it the previous -step's output token and hidden state. The draft length is a serving knob, -not a checkpoint property. §13 has the mode-selection rule. - -## The concat order — the one silent-wrong-answer trap - -`eh_proj` is `[hidden, 2*hidden]`. Which half multiplies the embedding -branch and which multiplies the hidden branch is **not** determined by -anything in the checkpoint's metadata, and getting it wrong is a pure -numerical error: no shape mismatch, no assert, no crash. Under rejection -sampling it does not even corrupt output — the drafts are simply always -rejected (see the last section). - -**Two sources disagree, and the naming is the one that is right.** - -* The DeepSeek-V3 report writes the projection as `M_k [ RMSNorm(h) ; - RMSNorm(Emb(t)) ]` — **hidden first**. -* The parameter is named `eh_proj`, with its two gains named `enorm` and - `hnorm` — **embedding first**. - -**Measured: the embedding block is first.** Two independent statistics over -the checkpoint's own weights, each controlled: - -*Column-norm profile.* Take the per-input-dimension L2 norm of each half of -`eh_proj` and correlate it with each RMSNorm gain. The pairing is exclusive: - -| | vs `\|enorm\|` | vs `\|hnorm\|` | vs `\|model.norm\|` | -|---|---|---|---| -| first half `eh_proj[:, :hidden]` | **+0.904** | +0.037 | +0.026 | -| second half `eh_proj[:, hidden:]` | −0.036 | −0.207 | **+0.544** | - -`model.norm` is the trunk's final RMSNorm, which gates the same residual -stream `hnorm` does, so the second half tracking it is the same statement as -the second half being the hidden branch. - -*Functional response.* Feed the embedding branch's actual input, -`enorm(Emb(t))`, through each half and out through `shared_head`, and -measure the entropy of the resulting distribution. **Compare each half -against its own controls, not against the other half** — the first half is -uniformly sharper by ~0.9 nats on *any* input, so a raw cross-half -comparison is confounded: - -| input | via first half | via second half | -|---|---|---| -| `enorm(Emb(t))` — the real embedding-branch input | **2.62 nats** | 6.22 nats | -| `hnorm(gaussian)` | 5.33 | 6.24 | -| gaussian, no gain (control) | 5.36 | 6.23 | -| `enorm(gaussian)` — right gain, wrong vector (control) | 5.09 | 6.17 | - -The first half drops **2.5 nats** on the real embedding and on nothing else. -The second half does not move at all: through it, a real token embedding is -indistinguishable from noise. (Uniform over a 129,280 vocabulary is 11.77 -nats.) - -So: **`eh_proj` consumes `concat(enorm(e), hnorm(h))`.** This is a -reconstruction from weight statistics, not a reading of reference code — it -is strong enough to implement against, and the acceptance rate is what -confirms it end to end. Keep the order as a single named constant at one -place in the layer so flipping it is a one-line experiment. - -## Checkpoint layout - -The module is stored as one more entry in the layer list, at index -`num_hidden_layers`. On this checkpoint that is 790 keys under -`model.layers.61.`: - -| key | count | dtype | shape | load? | -|---|---|---|---|---| -| `enorm.weight`, `hnorm.weight` | 2 | bf16 | `[7168]` | yes | -| `eh_proj.weight` | 1 | bf16 | `[7168, 14336]` | yes | -| `input_layernorm`, `post_attention_layernorm` | 2 | bf16 | `[7168]` | yes | -| `self_attn.*` | 7 | bf16 | **identical to a trunk layer's** | yes | -| `self_attn.{k_proj.k_scale, v_proj.v_scale}` | 2 | fp32 | `[]` | value is **1.0**, as in the trunk | -| `mlp.gate.weight` | 1 | bf16 | `[256, 7168]` | yes | -| `mlp.gate.e_score_correction_bias` | 1 | **fp32** | `[256]` | yes | -| `mlp.experts.{0..255}.{gate,up,down}_proj` | 768 | bf16 | `[2048, 7168]` ×2, `[7168, 2048]` | yes, EP-windowed | -| `mlp.shared_experts.{gate,up,down}_proj` | 3 | bf16 | same | yes | -| `shared_head.norm.weight` | 1 | bf16 | `[7168]` | yes — **a distinct norm**, not `model.norm` | -| `embed_tokens.weight` | 1 | bf16 | `[129280, 7168]` | **no** | -| `shared_head.head.weight` | 1 | bf16 | `[129280, 7168]` | **no** | - -**The last two are bit-identical copies of the trunk's** (`torch.equal` -against `model.embed_tokens.weight` and `lm_head.weight`: both True). The -runtime hands the layer whichever `embed_tokens` and `lm_head` the target's -draft-model container exposes, so pointing them at the trunk's is correct -and saves 1.85 GB per rank. They stay in the weight manifest's *predicted -non-load* set even with MTP on; the other 788 keys flip from non-load to -consumed. - -**The MTP layer's attention geometry is byte-identical to a trunk layer's** -— same `q_a_proj` / `q_b_proj` / `kv_a_proj_with_mqa` / `kv_b_proj` / -`o_proj` shapes, same two LayerNorms, same `q_lora_rank`. Whatever load-time -derivation the trunk's attention needs (absorption operands, row regrouping, -the rope table) applies unchanged. - -## The MTP layer is not quantized, and that is deliberate - -On a quantized export the MTP module can be excluded from quantization -wholesale. Here `hf_quant_config.json` carries `model.layers.61*` as one -wildcard entry in a 63-entry `exclude_modules` list, so **every weight in -the module is bf16** while the trunk's MLP path is NVFP4. - -**Do not re-quantize it at load time to reuse the trunk's expert -vocabulary.** The exclusion is the export's choice about where accuracy is -worth the bytes; quantizing it anyway changes what the checkpoint means, and -the acceptance rate — the only signal that can see the difference — would -absorb the damage silently as a lower draft quality rather than reporting -it. - -The consequence is a vocabulary consequence: the MTP layer's routed experts -need a **bf16** grouped-expert entry, at whatever `(local_experts, hidden, -intermediate)` the parallel split produces, while the trunk's use the NVFP4 -runner. Everything else in the graph — both norms, the concat, the -projection GEMM, the whole MLA block, the router, the shared expert, -`shared_head` — maps onto entries a dense-plus-MoE target already carries. - -## The HBM arithmetic, and the break-even acceptance rate - -A decode step is weight-bandwidth-bound: each rank reads every weight byte -it holds, once. So the cost of drafting is exactly the MTP layer's byte -count, times the number of draft steps. - -Per rank on this checkpoint at `dep4` (64 local experts of 256): - -| | bytes per draft step | -|---|---| -| routed experts, **bf16**: `3 x 2048 x 7168 x 64 x 2` | **5.637 GB** | -| attention, bf16: 187.1 M params x 2 | 374 MB | -| `eh_proj`: 102.8 M x 2 | 206 MB | -| shared expert: `3 x 2048 x 7168 x 2` | 88 MB | -| router | 4 MB | -| **total** | **≈ 6.31 GB** | - -Against a trunk step of ≈ 115 GB per rank (58 NVFP4 MoE layers at 1.585 GB + -61 bf16 attention blocks at 374 MB), that is **+5.5% per draft step**. Note -what the first row means: **the bf16 MTP layer's experts cost 3.56 times -what one NVFP4 trunk MoE layer costs**, purely from the dtype. - -A step now produces `acceptance_length` tokens instead of 1, so drafting -pays for itself when - -``` -acceptance_length >= 1 + max_draft_len * 0.055 -``` - -| `max_draft_len` | extra HBM/step | break-even `acceptance_length` | -|---|---|---| -| 1 | +5.5% | **1.055** | -| 2 | +11.0% | 1.110 | -| 3 | +16.4% | **1.164** | -| 4 | +21.9% | 1.219 | - -**That table is the weight-bytes model, and measurement says it is a lower -bound that stops holding as the batch grows.** Measured on this checkpoint -at `max_draft_len 3`, decomposing each engine step against a non-drafting -one: - -| con | acceptance | step cost — predicted | step cost — **measured** | decode speedup | -|---|---|---|---|---| -| 1 | 3.4904 | 1.164 | **1.511** | 2.311x | -| 32 | 3.3356 | 1.164 | **2.198** | 1.518x | -| 256 | 3.4033 | 1.164 | **2.600** | 1.309x | - -The marginal cost of each successive draft step at con=256 is **+0.645, -+0.569, +0.386**, against con=1's **+0.200, +0.173, +0.137**. The growing -term tracks the **rows** a step carries, not the layer's fixed weight bytes -— it falls off in the same shape the marginal row count does (1→2 rows is -+100%, 2→3 is +50%, 3→4 is +33%). **That is the -memory-bound-to-compute-bound crossover this file calls the genuinely -uncertain end, now measured rather than predicted.** - -So read the weight-bytes table as the floor it is: right at low concurrency, -where a step is bandwidth-bound and drafting really is nearly free, and -roughly 2.2x optimistic at the throughput end. What survives is the -conclusion, because the cost never grows fast enough to catch acceptance: -**at every concurrency measured, acceptance stayed at 2.9x or more of even -the measured break-even**, and no draft length in the certified range lost. - -**The break-even is still very low against the predicted cost** — at -`max_draft_len` 3 it takes only 0.164 extra tokens per step on average — and -two costs that look like they should matter do not: - -* **Collectives.** The MTP layer is an MoE layer, so it adds two per draft - step, and the trunk's grow by `draft_len + 1` in bytes. But at decode - sizes these calls are latency-bound rather than bandwidth-bound (measured - on this target: 2.753 MB in 25.472 µs = 108 GB/s, an order of magnitude - under NVLink), so multiplying the byte count barely moves the time. -* **KV capacity.** One extra layer on a 61-layer pool is **+1.6%**, plus - `max_draft_len - 1` extra tokens per sequence. - -The end that is genuinely uncertain was the **high-concurrency** end — once -a rank's batch already activates all of its local experts and the routed -GEMM is near the HBM roofline, `draft_len + 1`× the rows means the same -weight bytes with several times the arithmetic. It is measured above: the -crossover is real, it costs 2.2x the predicted step cost at con=256, and it -still does not overtake acceptance. - -**Acceptance itself is near-flat in batch size on this mechanism, which is -the other half of why no draft length loses.** Measured across three draft -lengths and nine concurrencies: `max_draft_len` 1 stayed in 1.9446–1.9673 of -a 2.0 ceiling, 2 in 2.6945–2.7707 of 3.0, and 3 in 3.3356–3.4904 of 4.0. -**So a "shorten the draft as the batch grows" schedule has no -acceptance-side motivation here** — any case for one has to be made on cost, -and on this checkpoint the cost side does not invert. (It is also unusable -under attention DP for a structural reason — see -`docs/references/trtllm-runtime-integration.md` §13.) - -## What only the acceptance rate can tell you - -**An MTP layer that computes the wrong thing does not produce wrong -output.** Rejection sampling guarantees the emitted distribution is the -target model's regardless of draft quality. A layer with the concat order -reversed, a norm on the wrong operand, or the expert stack packed in the -wrong interleave produces **bit-correct text, more slowly**, with every -draft rejected. An accuracy benchmark cannot see it. Neither can a smoke -gate. - -`acceptance_length` — the mean tokens emitted per step by requests that -carried a draft, `1.0` meaning total rejection — is the only detector, and -it needs a reference to be read against: - -* **≈ 1.0** at `max_draft_len` 3: the layer is wrong outright. -* **clearly above 1.0 but well below the reference**: the layer is *subtly* - wrong. This is the band the concat order lands in, and nothing else - distinguishes it from "this model just drafts poorly". -* **at the reference**: the layer computes what the checkpoint says. - -The reference is the same checkpoint under stock in-tree modeling, at the -same load and the same speculative configuration. Acceptance is a property -of the model, so a reference forced onto different memory knobs to boot is -still comparable. - -**One trap when choosing the workload, and it does not point the way it -looks like it should.** A serving harness that builds prompts from uniformly -random token ids seems like it must *understate* acceptance — the -continuation of a nonsense prefix is unpredictable, so the drafts should -miss. Measured, it runs the other way, and strongly. - -Only the **prompt** is random. Generation then runs for a fixed output -length with EOS ignored, so what MTP is drafting is the model's own -continuation of a nonsense prefix — which degenerates into repetition, and -repetition is the easiest thing in the world to draft. Measured on this -checkpoint at `max_draft_len 3` (ceiling 4.0): **`acceptance_length` 3.93 at -concurrency 1 and 3.46 at 32**, i.e. 97.6% and 82.1% of proposed draft -tokens accepted. Real text does not do that. - -So the random-prompt harness **flatters** MTP, and two things follow. A -throughput curve measured on it overstates the gain, so it answers "did -anything get slower" and not "is MTP worth enabling" — that one needs a -real-text workload, and an accuracy-gate run already produces one. And a -future run that sees a *low* acceptance number here must **not** excuse it -as "well, the workload is random": on this harness a low number means the -layer is wrong. diff --git a/tensorrt_llm/_torch/staircase/docs/references/trtllm-runtime-integration.md b/tensorrt_llm/_torch/staircase/docs/references/trtllm-runtime-integration.md deleted file mode 100644 index 62603d2f1312..000000000000 --- a/tensorrt_llm/_torch/staircase/docs/references/trtllm-runtime-integration.md +++ /dev/null @@ -1,1053 +0,0 @@ -# Target integration with the TensorRT-LLM runtime - -How a staircase target plugs into the TensorRT-LLM engine: registration, -construction, weight loading, the per-step forward duties, KV-cache -boundaries, CUDA-graph discipline, and the fail-fast ladder. Together -with the catalog contracts and `thop-attention-step-args.md`, this is the -complete integration knowledge for writing a target — no TensorRT-LLM -source reading is needed. - -Everything below was written against `tensorrt-llm == 1.3.0rc21` and -verified by the qwen3-8b/sm_100/tp1 gates (that target was the readable -exemplar of every pattern named here; it is not one of the two migrated in -this batch). - -> **Written before staircase moved in-tree.** The *runtime* content — the -> engine's construction order, the per-step forward duties, KV-cache -> boundaries, CUDA-graph discipline, the fail-fast ladder, and §13's MTP -> shell specification — is what this file is for and still holds. The -> *packaging* content does not: `uv run`, `bench/`, `model_dir/`, -> `llm_args.yaml`, `scripts/env.sh`, `STAIRCASE_TARGET` and the multi-rank -> registration ritual are all gone. Sections 1, 9 and 12 are marked where -> they describe the old shape; `../../README.md` is the current one. - -## 1. Directory shape - -**Superseded — see `../../README.md`.** In-tree a target is a package under -`models//targets////`; the identity triple is -unchanged, but `llm_args.yaml`, `model_dir/` and `perf/data` are gone (the -topology is the caller's input, the checkpoint is read unpatched). The -pre-move shape, kept because the rest of this file refers to it: - -``` -targets//// # identity = the path triple -├── modeling.py # registration shell + core + step-args + fail-fast -├── weights.py # MANIFEST + load loop + post-load derivations -├── smoke.py # self-locating keyword-assert gate (uv run ) -├── llm_args.yaml # LLM-API overrides; empty {} = trtllm defaults; -│ # parallel topology lives here for multi-rank targets -├── TARGET.md # identity, checkpoint hashes, pins, vocabulary, -│ # verification records, and — added by the tuner — -│ # the Performance section -├── configs/ # tuner knob variants (tracked); added by the perf -│ # campaign, absent on a freshly assembled target -├── perf/ # data/ machine-local and ignored; figures/ tracked -└── model_dir/ # what trtllm loads - ├── config.json # committed; architectures[0] = registered class name; - │ # all other fields = the upstream checkpoint config - └── *.safetensors, tokenizer*, generation_config.json - # machine-local symlinks into the real checkpoint — - # untracked; relink against the TARGET.md hashes -``` - -`model_dir/` is a stub HF checkpoint directory: trtllm resolves the model -class from its `config.json`, reads weights and tokenizer from the -symlinks, and never sees the rest of the target. Tools pass -`model=/model_dir` and read `llm_args.yaml` next to it. - -## 2. Registration and resolution - -- `@register_auto_model("StaircaseForCausalLM")` on the shell class - writes the class into trtllm's process-global registry at import time. - Every target registers the **same** name (identity lives in the - directory path), so one process hosts one target; comparative runs use - separate processes. -- Tools import `modeling.py` **by file path** (the target dir name - contains a hyphen and is not a package): - - ```python - spec = importlib.util.spec_from_file_location( - "staircase_target_modeling", target / "modeling.py") - module = importlib.util.module_from_spec(spec) - sys.modules["staircase_target_modeling"] = module # §12: pickle needs it - spec.loader.exec_module(module) # registration side effect - ``` - - The import must precede engine construction in **every process that - builds the model** — which above world size 1 is not this one. The - `sys.modules` line binds the loaded module to the name the class will - carry in `__module__`; without it a multi-rank run dies in the launcher - before any rank starts. §12 has the mechanism and the rest of what - changes there. -- `modeling.py` bootstraps `sys.path` from its own `__file__` (the repo - root is the target directory's `parents[3]` — the path depth is frozen - by the `targets////` shape) before its - `catalog.*` imports — this is why E402 is disabled for `targets/**` in - pyproject. -- `LLM(model=/model_dir)` then resolves `architectures[0]` - against the registry and construction begins. - -## 3. The shell and the core - -```python -class StaircaseCore(DecoderModel): # the computation - def __init__(self, model_config): ... # geometry asserts + weights - def forward(self, attn_metadata, input_ids=None, position_ids=None, - inputs_embeds=None, lora_params=None, **kwargs): ... - -@register_auto_model("StaircaseForCausalLM") -class StaircaseForCausalLM(DecoderModelForCausalLM[StaircaseCore, ...]): - def __init__(self, model_config): - super().__init__(StaircaseCore(model_config), config=model_config, - hidden_size=..., vocab_size=...) - def load_weights(self, weights, *args, **kwargs): ... - def post_load_weights(self): ... -``` - -- The shell (`DecoderModelForCausalLM`) is composition: it stores the - core as `self.model`, creates `self.lm_head`, and owns the logits path - — its forward calls the core, then gathers exactly the rows that need - logits (last token per context sequence + every generation token) and - applies lm_head. **The core returns final-normed hidden states - `[num_tokens, hidden]` and never touches logits.** -- Inputs are **packed**: first dimension = total tokens of the batch, - context sequences first. `position_ids` arrives as an int32 `[1, T]` - view of an engine buffer — flatten with `reshape(..., [-1])`. -- Construction asserts the target's identity — geometry, dtype, and every - config property the assembly depends on (tie_word_embeddings, qk-norm - presence, rope scaling, sliding window, attention bias). Assert, never - adapt: a mismatched checkpoint is a different target. -- Inputs the target does not implement (`lora_params`, `spec_metadata` - in kwargs) must **assert loudly** — silently ignoring them produces - wrong output with no signal. This is distinct from runtime-owned - feature flags, which pass through (see §6). A target that *does* - implement speculative decoding stops asserting on `spec_metadata` and - grows a second forward it owns; §13 is that axis. - -## 4. The weight lifecycle - -| stage | who | what happens | -|---|---|---| -| t0 construct | engine (`MetaInitMode`) + your `__init__` | `torch.empty` is intercepted to the meta device: parameters are shape/dtype shadows, zero memory. Declare with `torch.empty` only (`zeros`/`full` are NOT intercepted) and do no tensor math in `__init__` (meta tensors raise on compute). | -| t1 materialize | engine | walks **registered** parameters (`nn.Parameter` / `ParameterDict` / module tree — what `named_parameters()` can see) and reallocates each on CUDA, contents garbage. Registration is what requests the allocation; anything held in a plain dict/list stays a dead shadow. | -| t2 read | engine | reads every safetensors in `model_dir/` into a dict `{ckpt_key: tensor}`. At this pin the values arrive **already materialized** as `torch.Tensor`, not as lazy slices — measured at world size 4, where every rank is handed the whole dict (25,687 keys on all four). So the engine pre-shards nothing: a multi-rank split is entirely the manifest's, and a `src` transform's benefit is that only this rank's bytes cross to the device, not that fewer bytes are read off disk. | -| t3 load | your `load_weights(weights)` | full delegation, never validated by the engine — the manifest loop copies bytes into the materialized storage (`param.data[a:b].copy_(src)`). | -| t4 derive | your `post_load_weights()` | the designated home for derived state: `.t()` GEMM views, per-layer tuples, and (future) checkpoint-calibrated call tensors such as fp8 KV scales. Real tensors may be created here — meta is over. | -| t5 serve | your `forward` | reads weights by reference; zero copies or transforms on the hot path. | - -Storage layout: declare parameters in **HF `[out, in]` row-major** so t3 -copies are layout-preserving, and consume GEMMs through `.t()` views -(row-major `[N, K]` transposed is exactly the dense column-major `[K, N]` -that `cublas_mm`'s contract requires; the view is zero-copy and shares -storage, so reloading weights in place keeps views valid). - -## 5. The weight manifest convention - -`weights.py` owns a pure data table: - -```python -MANIFEST[param_key] = [(ckpt_key, dst_slice_or_None), ...] -# tp1: no src transform. Sharded targets add one: -# (ckpt_key, src=col_shard(rank, tp), dst_slice) -# and declare per-rank shapes in modeling. -``` - -Rules: fused parameters (qkv, gate_up) fill by destination row slices — -no intermediate concat buffers; assert shape and dtype per copy; assert -**bidirectional coverage** (manifest keys == declared parameters before -the loop; consumed ckpt keys == the whole checkpoint after it). The -shell-registered exceptions (`lm_head.weight`, and `embed_tokens` when -the base class owns it for tied checkpoints) are fed explicitly. - -## 6. Forward duties, KV cache, and feature flags - -The engine drives everything around the forward: scheduling, KV block -bookkeeping (allocation at admission, per-token growth, prefix reuse), -metadata filling and `prepare()` — all before your forward runs. Your -duties per step: - -1. Build the attention step-args once from the prepared metadata and - share them across layers (`thop-attention-step-args.md` is the - normative mapping). -2. Keep the loop flat: catalog calls, tensor-metadata reads, Python - control flow — nothing else (the closed-vocabulary rule; audit = - collect the calls in `forward` **and the private methods it reaches** - and match them against `catalog/index.yaml`. Scope to that closure: a - whole-file scan flags the rope table's `arange`/`cos`/`sin` and - load-time `.t()`/`.to()`, which run at init and are outside the rule). -3. Let one attention call serve any batch composition - (`attention_input_type=0` for the standard configuration — no phase - branches; MLA targets dispatch per phase by metadata reads). - -The KV-ownership red line does not move: the pool is sized by the engine from the -`config.json` declarations (`num_hidden_layers`, `num_key_value_heads`, -`head_dim`, `torch_dtype`, `vocab_size`, `max_position_embeddings` — -declare them honestly), pages are assigned by the C++ manager, and the -new tokens' K/V are written **inside** the attention op. The target only -forwards the address book. - -Runtime-owned feature flags (block reuse, CUDA graphs, beam, spec-dec -plumbing) **pass through** from metadata exactly as prepared: the target -runs trtllm defaults, feature behavior is upstream's responsibility, and -the gates validate the result. Kernel-surface certification accounting -lives in the catalog contracts, not in target config. - -## 7. CUDA-graph discipline - -Decode-only steps are captured and replayed with **no Python executing**. -Every value the attention call consumes must therefore fall into one of -three classes: - -- **reference**: engine-owned persistent buffers, refreshed in place each - step — replays read fresh contents automatically (all step-args - tensors); -- **GPU-derived**: recomputed by captured kernels on replay; -- **host-derived**: Python scalars frozen at capture — legal only when - they are per-capture constants (`num_contexts == 0` in a decode graph). - -The step-args builder is the single audit point: classify every entry. -Per-forward `empty(...)` output buffers are fine (the caching allocator -serves stable blocks; upstream captures the same pattern). - -## 8. The fail-fast ladder - -| fuse | when | checks | catches | -|---|---|---|---| -| static contract | import | trtllm version == pin; every `torch.ops.trtllm.*` symbol the forward calls; the thop binding | version drift (voids receipts and gate records), missing/renamed ops. torch mirrors (`embedding`, `empty`, `reshape`, ...) are exempt — core PyTorch API, upstream-owned | -| step contract | first forward | metadata is `TrtllmAttentionMetadata`; every `_STEP_FIELDS` name exists; `position_ids` is int32 | private-surface layout drift within a same-version build. Everything checked is fixed at engine construction — once per model instance is sound | -| gates | explicit runs | smoke keywords; accuracy score vs `bench/references/accuracy.yaml` | the computation itself | - -Keep `_STEP_FIELDS` exactly equal to the set of metadata attributes the -code reads, and the static-contract op list exactly equal to the trtllm -ops the forward calls — the lists are dependency declarations. - -**One metadata surface is conditional, and a flat existence check over it -turns a legal config into a hard failure.** MLA's cached-KV fields — -`enable_context_mla_with_cached_kv`, `ctx_cached_token_indptr`, -`ctx_kv_indptr`, `ctx_uncached_token_indptr`, `max_ctx_seq_len`, -`max_ctx_kv_len`, `num_ctx_cached_tokens` — exist only when block reuse -is on. With `enable_block_reuse: false` they are **absent, not False**, so -a `_STEP_FIELDS` tuple containing them fails the first forward on a -configuration the target is supposed to support. Split the tuple: the -unconditional fields keep the existence check, and the conditional ones -select the context flavor by their *presence* rather than being asserted. -See `docs/models/latent-kv-cache.md`. - -## 9. Gates and launch - -**Superseded — see `../../README.md` for the current commands.** The -substance below (what each gate is for, and how to author smoke cases) -carries forward; the invocations do not. - -- `source scripts/env.sh` before any GPU run (exports the - single-process-worker flag; consumed upstream only at world_size == 1, - ignored by multi-rank runs). Claim the devices with - `CUDA_VISIBLE_DEVICES` — one for a `tp1` target, `N` for a `tepN`/`depN` - one. -- Smoke: `uv run targets//smoke.py` — self-locating, boots from its - own directory, greedy keyword asserts, exit code is the verdict. The - cases are target-authored: pick prompts whose greedy continuation is - high-confidence for this model, and **verify every keyword on the real - model before freezing it** — an unverified keyword bakes a false - failure into the gate. -- Release: `uv run bench/accuracy.py --target targets/` — one-sided gate, - measured >= reference − tol; after a first pass on an external anchor, - write the measured score back (see the references file header). - -## 10. What varies per model — the re-derive list - -The exemplar target shows the pattern; every model-specific value must be -re-derived from the new checkpoint's `config.json` and the catalog -contracts, never copied. The known variation axes: - -**attention structure** — GQA/MHA vs **MLA**, the axis that reshapes the -most: MLA replaces per-head K/V with one compressed latent, needs two -attention calls per forward (context and generation are separate call -shapes, mixed batches rejected upstream), and brings its own load-time -obligations. It is not a variant of the geometry row below; see -`docs/references/mla-custom-op-decomposition.md` for the vocabulary and -`docs/models/latent-kv-cache.md` for what a target had to derive. -Then: geometry (layers, heads, kv-heads, head_dim, intermediate, vocab); -qk-norm presence and its eps; rope kind, theta, scaling, `is_neox`, -partial-rotary factor; `tie_word_embeddings` (changes the lm_head/embed -feeding and the shell's tying path); attention bias (needs cublas_mm's -fused-bias argument); sliding window; activation function; checkpoint -key names and fusion grouping; dtype and quantization scheme — including -*which modules it excludes*, since a checkpoint may quantize only its MLP -path and leave attention bf16; parallel topology (llm_args.yaml + -per-rank shapes + manifest src transforms + collective-communication -catalog entries — §12 has the launch substrate, the rest is the first -multi-rank target's to derive); KV-cache dtype (turns the kv-scale -constants into post-load tensors — certification extension first). - -A construction-time assert exists for each axis the assembly depends on: -if the new config violates one, the right response is to re-derive that -part of the assembly, not to delete the assert. - -### Read the axes off the engine's config object, not off AutoConfig - -`model_config.pretrained_config` — the object the engine hands -`__init__` — is the **un-migrated** config: fields sit where the -checkpoint's own `config.json` put them. A standalone -`AutoConfig.from_pretrained(model_dir)` probe can return a *different* -shape of the same information, because transformers migrates fields -across versions. - -Observed at transformers 5.5.4 on a checkpoint written by 4.51: through -AutoConfig, rope had been migrated into -`rope_parameters = {'rope_theta': ..., 'rope_type': 'default'}` with -`rope_scaling` an alias of that same dict, and `cfg.rope_theta` raising -`AttributeError`. Through the engine, `rope_parameters is None`, -`rope_scaling is None`, and `rope_theta` is a plain instance attribute. -Deriving the axis from the AutoConfig surface produced a target that -failed at construction with `TypeError: 'NoneType' object is not -subscriptable`. - -So: an AutoConfig probe is **not** a valid oracle for what -`pretrained_config` will look like. Read the checkpoint's `config.json` -directly to learn what the model *is*, and use the flat -`cfg.` idiom inside the target. If a probe is needed, instrument -the target's own `__init__`. - -### A missing `dtype` is not inherited — declare it in the stub config - -A checkpoint may declare **no** `dtype` and no `torch_dtype` at all -(observed on gpt-oss-120b, whose weights are all BF16 outside the -quantized expert blocks). The two config surfaces then disagree: -`model_config.pretrained_config.torch_dtype` is `None` while the engine's -own `ModelConfig.torch_dtype` resolves `torch.bfloat16`. The shell reads -the **pretrained** one, so `DecoderModelForCausalLM` materializes its -shell-owned parameters — `lm_head.weight` — as **fp32**, and a target with -a dtype-checking weight manifest fails there with a cause two layers -upstream of the symptom. - -The fix belongs in `model_dir/config.json`, which is a stub the target -owns: declare `"dtype": "bfloat16"` alongside the `architectures` patch. -That makes the stub carry two deliberate divergences from the checkpoint's -config rather than one, which is correct — the field states what the -checkpoint *is*, and it is the field the engine sizes the KV pool from. -Assert both surfaces in `__init__` so a future drift is loud. - -### `model_dir/` linking is not a `tokenizer*` glob - -An accuracy protocol that applies a chat template needs the template -files, and a recent checkpoint may keep **no** `chat_template` inside -`tokenizer_config.json` — it lives in `chat_template.jinja` (plus -`chat_template.json`). Those names, and `special_tokens_map.json`, match -no `tokenizer*` pattern. Link by inspecting what the checkpoint actually -ships, not by a fixed glob, or the gate fails at template application -after a full engine build. - -## 11. Serving-time costs a target inherits - -The engine owns these, not the forward, but they land on the target's -Pareto curve and the first two are worth checking on every new target -before any modeling work. - -**Decode CUDA-graph coverage defaults to batch 128.** Above it the whole -decode phase runs eager. The deeper the model the worse this is: measured -on a 48-layer target, concurrency 256 lost 19.7% throughput and 151% TPOT -against its own graph-covered configuration, and raising -`cuda_graph_config.max_batch_size` to 256 with `enable_padding: true` was -worth +56.5% at that point with every other point unchanged. This should -be the first config experiment on any new target — but raise the ceiling -and A/B the padding **separately**: on a later target `enable_padding: -true` measured 11% *worse* at concurrency 256 against the same raised -ceiling, because it swaps the fine `[1..32]` batch grid for a coarse one. - -**`max_seq_len` is a per-step cost, not only an admission cap.** It sizes -`max_blocks_per_seq = max_seq_len / tokens_per_block`, which sizes a -pinned `[1, num_seqs, 2, max_blocks_per_seq]` int32 block-offset staging -buffer that the resource manager allocates, memcpys and pushes H2D **every -step**. Left at a model cap far above the served workload (40960 vs a -1024/1024 benchmark) that is 2.62 MB per step at 256 in flight. Capping it -is a real knob — but size the expectation: a measured 10× shrink moved -`_prepare_inputs` only 5.17 → 4.44 ms/step, so the staging is a minority -of that cost and the residual is `O(num_seqs)` per-request executor -Python. Capping also **restricts capability** (longer requests are -rejected), and the floor is set by the accuracy gate's own prompts, not by -the benchmark — 5-shot MMLU prompts reach 2687 tokens, so a cap at the -benchmark's exact ISL+OSL makes the gate unrunnable. - -## 12. Multi-rank launch — what changes above world size 1 - -Everything here was measured at the pinned version on B200 hosts at world -size 4. Two layers, and they answer different questions: the **launch -substrate** (process model, registration, the failure modes that exist only -above world size 1) applies to every multi-rank segment, while **what a -rank does differently** and **what attention DP changes** are per-topology -and were each established by the first target on that topology. See "Still -not established here" at the end for what no target has reached yet. - -### The process picture - -`LLM(..., tensor_parallel_size=N)` with `N > 1` does not build the model -in the calling process. It spawns one MPI worker per rank -(`MpiPoolSession` → `mpi4py.futures.MPIPoolExecutor` → `MPI_Comm_spawn`) -and the engine is constructed there; the launcher becomes a proxy. The -single-process-worker flag is read only after that branch, so at world -size > 1 it neither helps nor hurts. - -### Registration reaches a rank through exactly one hook - -mpi4py re-imports the launcher's **main module** in every spawned worker, -under `__name__ == "__worker__"`. Measured consequences: - -| | launcher | worker rank | -|---|---|---| -| module-scope code | runs | **runs** | -| `if __name__ == "__main__":` body | runs | does not run | -| `sys.argv` | full | **script path only** | -| `os.environ` | — | **inherited** | - -**This whole problem is gone in-tree, and the table above is now only an -explanation of why the old code looked the way it did.** A target class is -an ordinary member of an installed package, so a spawned rank resolves it -through the normal registry — no module bound into `sys.modules` before -pickling, no `__worker__` re-import to time correctly, and no -`STAIRCASE_TARGET`, which existed solely because argv does not reach the -workers and `LlmArgs` does. `bench/register.py` and its module-scope -`from_env()` were deleted rather than migrated. - -What the table still explains: any *tool* that has to influence worker -ranks must do it through `LlmArgs` or the environment, never argv. - -### Two failure modes that exist only above world size 1 - -**The loaded module must be bound in `sys.modules`.** The engine resolves -`architectures[0]` to a class object in the launcher and ships it to the -ranks by pickle, which stores a class *by reference* — `__module__` plus -`__qualname__` — and both sides resolve those names through -`sys.modules`. Loading by path without binding the name fails in the -launcher, before any rank starts: - -``` -PicklingError: Can't pickle : - import of module 'staircase_target_modeling' failed -``` - -Binding it *twice* is equally fatal — the second load leaves a different -class object under the same name and pickle's identity check rejects it -(`it's not the same object as ...`) — so the load has to be idempotent. - -**An environment variable set after MPI initializes never reaches a -rank.** Importing `tensorrt_llm` initializes MPI, and OpenMPI hands a -spawned process the environment as it stood at that moment. A tool that -exports into its own environment must do so *before* that import; one -that hands a freshly spawned server a prepared environment satisfies this -by construction. - -### The parallel segment and its knobs - -| segment | `llm_args.yaml` | -|---|---| -| `tp` | `tensor_parallel_size: N` | -| `tep` | the above plus `moe_expert_parallel_size: N` | -| `dep` | the above plus `enable_attention_dp: true` | - -`tep4` and `dep4` were both observed building a serving engine and -generating correct greedy text from the stock DeepSeek-V3-Lite NVFP4 -checkpoint (254 s and 142 s to a ready engine). That is a statement about -the LLM API on this host — not about any staircase target. - -### What the engine hands the model - -`model_config.mapping` carries the rank's place in the topology. Fields -observed on `Mapping(world_size=4, tp_size=4, moe_ep_size=4, rank=1)`: - -| field | value | -|---|---| -| `rank` / `world_size` | 1 / 4 | -| `tp_size` / `tp_rank` | 4 / 1 | -| `pp_size` | 1 | -| `moe_ep_size` / `moe_ep_rank` | 4 / 1 | -| `moe_tp_size` / `moe_tp_rank` | 1 / 0 | -| `enable_attention_dp` | False | -| `tp_group` | `[0, 1, 2, 3]` | - -Read the topology off this object, never off the segment string: the -segment names the intent, `mapping` is what the engine actually built. - -### What a rank actually has to do differently - -Established by the first multi-rank target (`tep4`, world size 4). These -are the four things this section used to list as unobserved. - -**`lm_head` is topology-aware, and that is a manifest obligation §4/§5 do -not mention.** At `tp_size = 4` `DecoderModelForCausalLM` builds an -`LMHead` of `[vocab/tp, hidden]` — vocab-parallel — so the manifest must -feed *this rank's contiguous block of vocabulary rows*, and the logits -gather stays upstream's. Consequence: `vocab % tp_size == 0` belongs in -the construction asserts, because nothing else checks it. - -**Where the collectives belong: after every row/column-sharded producer, -and nowhere else.** For a TP attention + EP MoE layer that is two per -layer — after `o_proj`, and after `routed_window + shared_partial` -(summed locally first, so one collective serves the whole MLP). The -entry is `comm/allreduce`; the workspace-free certified path is -`strategy=0` (NCCL) or `8`, `workspace=None`. Keeping `op=0` (plain sum) -and leaving the existing fused-add-rmsnorm in place preserves the -single-rank residual structure, which is also the CUDA-graph-friendly -shape. - -**Which transport, though, is worth about 2x at decode sizes, and the -workspace-free default is the slow one.** A decode message is three orders -of magnitude below the point where NCCL's ring algorithm starts to pay for -itself, and a blocking collective serializes far more than its own time. -`docs/references/collective-allreduce-transport.md` has the algorithms, -the measured crossover, and the two profiling traps that make this easy to -get wrong. - -**The manifest's `src` transforms: a fourth column, applied before the -relayout.** `(ckpt_key | tuple, src, dst_index, transform)`, with `src` -selecting this rank's slice per key of a multi-key row. Order is -load-bearing: relayout transforms are functions of the *per-rank* row -count, so a transform that ran on the whole tensor and sharded afterwards -produces a different byte order. Column shards are strided views, so -densify before any transform reinterprets bytes. - -**Per-rank parameter shapes** follow the topology mechanically — heads, -dense/shared intermediates and the expert-axis window all divide — with -one trap worth stating: re-derive every kernel alignment rule at the -*per-rank* width rather than inheriting the single-rank conclusion. - -**Coverage asserts change shape.** `leftover == {}` no longer holds: under -EP each rank legitimately leaves the off-window experts' keys unconsumed. -Replace it with an explicitly predicted leftover set computed from the -rank's window, and add a parameter-side assert that the shell-registered -parameters are exactly `{lm_head.weight}`. - -How MLA's latent cache behaves under attention TP is a mechanism fact -rather than a runtime one, and lives in `docs/models/latent-kv-cache.md`: -it is **replicated, not sharded** — attention TP buys zero KV memory on -MLA. - -### What attention data parallelism changes on top of that - -Established by the first `depN` target (`dep4`, world size 4, an MLA + EP -MoE checkpoint). **Three of the four TP bullets above come out -differently**, so read this section as replacing them rather than adding to -them whenever `enable_attention_dp` is on. - -**The split is over requests, not over heads.** Attention is *replicated*: -each rank builds the full head count (32, not `32/tp`), holds the full -attention weights, and serves its own subset of the batch. Consequently -**the post-`o_proj` all-reduce disappears** — a rank's attention output is -already complete for its own tokens. - -**`lm_head` is replicated, not vocab-parallel.** At `tp_size = 4` *without* -attention DP the shell builds an `LMHead` of `[vocab/tp, hidden]`; with it, -the shell builds the full `[vocab, hidden]` on every rank. The manifest -obligation reverses, and `vocab % tp_size == 0` stops being a construction -requirement. - -**The manifest's `src` transform column is a TP artifact.** Nothing outside -the routed expert stacks is sharded, so a `depN` manifest is the tp1 -three-column form with the expert loop windowed. - -**The collectives move to the MoE and change identity.** Not two -all-reduces per layer, but one `comm/allgather` + one `comm/reducescatter` -per **MoE** layer (a dense layer has none). The gather **belongs before the -router GEMM, not merely before the expert call** — see -`docs/models/expert-weight-packing.md` for why the EP tiling invariant -depends on that placement and fails silently otherwise. - -**How a rank learns the other ranks' token counts:** -`attn_metadata.all_rank_num_tokens`, a host int list, identical on every -rank. Every rank pads to `max(...)` so both collectives run in their -uniform (`sizes=None`) form — which is also the only form a CUDA graph can -replay, since `sizes` is a host argument frozen at capture. - -**That padding creates an obligation.** Collectives pair **by position** on -the communicator, and at equal byte counts an ordering divergence is -**silent** — every rank wrong in 98-99% of elements, bitwise reproducibly, -no hang. Padding guarantees equal byte counts, so a `depN` forward must -issue an identical call sequence on every rank and must not wait for a hang -to detect that it did not. Certified in `catalog/comm/allgather.md` and -`catalog/comm/reducescatter.md`. - -**The runtime is not a second party on that communicator.** Traced live: -one NCCL communicator per rank, 14,036 collectives over 242 forwards, -**zero** issued by the runtime. The engine's own attention-DP -synchronization is host-side MPI on a *different* communicator -(`MPIDist.tp_comm`, built by `MPI_Comm_create_group`). So a target's -collectives cannot be mispaired against the engine's. - -**The CUDA-graph gate that looks like a landmine and is not.** -`cuda_graph_runner.py` replays a graph under `enable_attention_dp` only if -*all* ranks are generation-only **and** their batch sizes are exactly -equal; the padding that would force equality is off by default. Measured at -steady-state serving load, it **never binds**: `cudaGraphLaunch` = 1.00 per -rank per decode step, i.e. **100% replay coverage** without -`enable_padding`. It does bind on a draining workload — an accuracy-gate -run observed captures but no replays across 334 forwards, where the batch -shrinks monotonically. Both observations are correct about different loads. -Consequences measured on `dep4`: `cuda_graph_config.enable_padding: true` -bought nothing and cost **-2.55%** at concurrency 256 (it can only coarsen -the grid), and `attention_dp_config.enable_balance: true` cost **-13.10%** -at concurrency 128 with mean TTFT doubling, because its `batching_wait_iters` -hold costs more than the alignment buys once random arrival already -balances the ranks. - -**KV cache.** No rank factor: see `docs/models/latent-kv-cache.md` — on -MLA the per-rank pool is identical across tp1/tepN/depN, and what `dep` -buys is `world_size x` *aggregate* capacity from holding disjoint requests, -not a narrower pool. - -### Still not established here - -Pipeline parallelism (`pp_size > 1`); multi-node, where the loopback -pinning `scripts/env.sh` applies for single-host spawn must be overridden; -and `moe_tp_size > 1`, i.e. splitting experts a second way along the -intermediate dimension. - -## 13. One-engine speculative decoding — what changes when the target drafts - -Everything above describes a target that emits **one** token per generation -sequence per step. A speculative target emits `1 + draft_len`, and the draft -tokens come from a second forward that the target itself owns. This section -is the axis that brings: what the engine looks for on the model object, what -the drafting loop calls, where the ownership line falls, and what a rank has -to do differently inside its forward. - -Read like §12, in two layers. The **binding surface** applies to every -one-engine speculative mode. The **per-step consequences** below it were -established for **MTP-Eagle one-model** (`MTP_EAGLE_ONE_MODEL`) on an MLA + -EP-MoE checkpoint at `dep4`, and are marked where another mode differs. -`docs/models/multi-token-prediction.md` carries the other half — what an MTP -layer computes, which is a checkpoint fact rather than a runtime one. - -### The checkpoint picks the mode; the target does not - -`update_spec_config_from_model_config` runs **before the model is built** -and reads the MTP layer count out of the pretrained config -(`num_nextn_predict_layers`, or `mtp_num_hidden_layers` on Qwen3Next-style -configs; 1 if neither is present). `MTPDecodingConfig`'s defaults are -`use_mtp_vanilla=False` and `mtp_eagle_one_model=True`, so: - -| checkpoint layer count | resulting mode | -|---|---| -| `n == 1` | **`MTP_EAGLE_ONE_MODEL`** — one layer of MTP weights, replayed | -| `n > 1` | `MTP` (vanilla) — one distinct layer per draft position | - -**Under MTP-Eagle, `max_draft_len` is not bounded by the checkpoint's layer -count.** That bound belongs to vanilla MTP. The single MTP layer is replayed -autoregressively `max_draft_len` times, so "how many MTP layers do we turn -on" is the wrong question and "what is `max_draft_len`" is the right one. - -**Spell `max_draft_len` out in every config variant.** Left unset on the -MTP-Eagle path it resolves to **1**, not to anything derived from the -workload. `max_total_draft_tokens` is then mirrored from it (linear tree), -and `tokens_per_gen_step = 1 + max_total_draft_tokens`. - -### The engine finds the drafter through exactly one getattr - -```python -# _torch/pyexecutor/model_engine.py -def _get_spec_worker(self): - return getattr(self.model, 'spec_worker', None) -``` - -That is the whole registration. Everything else the runtime touches on the -model side, it reaches through the worker's own arguments: - -| attribute / callable | type | what reads it | -|---|---|---| -| `model.spec_worker` | `SpecWorkerBase` | the engine's one getattr | -| `model.config` | pretrained config | `update_spec_config_from_loaded_model` (the base shell already provides it) | -| `model.draft_config` | — | read with `getattr(..., None)`; **absent is correct** for a single-checkpoint MTP target | -| `draft_model.mtp_layers` | `nn.ModuleList` | only `[0]` is ever indexed — MTP-Eagle replays one layer | -| `draft_model.embed_tokens` | module | passed to the layer as a kwarg | -| `draft_model.lm_head` | module | passed to `shared_head` | -| `draft_model.model.d2t` | — | read with nested `getattr(..., None)`; **absent is correct** (draft and target share a vocabulary) | -| `mtp_layers[0](...)` | callable | the draft loop, once per draft step | -| `mtp_layers[0].shared_head(h, lm_head, attn_metadata, True)` | method | returns draft logits | - -`draft_model` is a container the target defines and hands to the worker; the -runtime never constructs it and never inspects it beyond the four names -above. - -The layer is called by keyword, with `inputs` splatted in: - -```python -hidden_states = draft_model.mtp_layers[0]( - embed_tokens=draft_model.embed_tokens, - all_rank_num_tokens=, - input_ids=..., position_ids=..., hidden_states=..., - attn_metadata=..., spec_metadata=..., -) -``` - -It returns **one tensor**, `[rows_this_step, hidden]`, unpadded — the same -row count its `input_ids` carried. The loop slices it with its own -`gather_ids` afterwards. (Eagle3 one-model returns a second tensor here; -MTP-Eagle does not.) - -### Where the ownership line falls - -| responsibility | owner | -|---|---| -| accept/reject, rejection sampling, the golden token | runtime | -| KV rewind, `attn_metadata` rewrite between draft steps and its restore | runtime | -| the draft loop, `gather_ids`, position shifting, sampling draft tokens | runtime | -| `runtime_draft_len` scheduling and padding, `(bs, draft_len)` graph capture | runtime | -| the KV pool's extra layer and extra tokens | runtime | -| **the MTP layer's forward** | target | -| **`shared_head`** | target | -| **the `draft_model` container** | target | -| **the shell's speculative branch** | target | -| **loading the MTP layer's weights** | target | -| **attention-DP padding inside the MTP layer** | target | - -### Inheriting the in-tree one-engine shell is not the shortcut it looks like - -`SpecDecOneEngineForCausalLM.__init__` builds its drafter by calling -`get_draft_model(...)`, which is a **module-level function, not a method** — -a subclass has no override point. It dispatches on the config's -`model_type`, and a staircase stub config patches `architectures` only, so -`model_type` still names the upstream family and the call returns -**trtllm's own MTP layer**. Inheriting therefore hands the one computation -this project exists to write to the engine instead. - -This is a statement about what gets constructed, not a rule against -inheriting: a shell already inherits `DecoderModelForCausalLM` from the same -package, and the isolation hook gates *reading* whole-model definitions, not -importing them. Writing the branch out (about 40 lines) keeps the forward -readable end to end, which is what the Vocabulary table and the -closed-vocabulary audit both rest on. - -### The shell's shape, and the four things it has to get right - -```python -class StaircaseForCausalLM(DecoderModelForCausalLM[StaircaseCore, ...]): - def __init__(self, model_config): - ... # unchanged - self.spec_config = getattr(model_config, "spec_config", None) - self.draft_model = None - self.spec_worker = None - if self.spec_config is not None: - assert self.spec_config.spec_dec_mode.is_mtp_eagle_one_model() - self.draft_model = (...) - self.spec_worker = get_spec_worker(self.spec_config, model_config, - model_config.mapping) - - def forward(self, attn_metadata, **kw): - if self.spec_worker is None: - assert kw.get("spec_metadata") is None - return super().forward(attn_metadata, **kw) # bit-identical - spec_metadata = kw["spec_metadata"] - hidden = self.model(attn_metadata=attn_metadata, **kw) - logits = self.logits_processor.forward( - hidden[spec_metadata.gather_ids], self.lm_head, attn_metadata, True) - return self.spec_worker( - input_ids=kw["input_ids"], position_ids=kw["position_ids"], - hidden_states=hidden, logits=logits, - attn_metadata=attn_metadata, spec_metadata=spec_metadata, - draft_model=self.draft_model, - resource_manager=kw.get("resource_manager")) -``` - -`get_spec_worker` is imported from `tensorrt_llm._torch.speculative` — the -**runtime**, the part of the stack this project reuses, not modeling. - -Four things this shape is load-bearing about: - -**The non-speculative path must be a delegation, not a reimplementation.** -A target's release criterion was measured on the inherited base forward; the -only way to keep it bit-identical is to call it. Declaring -`resource_manager` as a named parameter would silently drop it from `**kw`, -so leave it in the dict and pull it with `.get` in the speculative branch -only — then the base receives exactly what it receives today. - -**The shell gathers the logits; the engine does not.** For every one-model -mode `without_logits` is True, so `_forward_step` returns the model's dict -verbatim and applies no second gather. `spec_metadata.gather_ids` holds one -row per context request (its last token) and `runtime_draft_len + 1` rows -per generation request — pass `hidden` **ungathered** to the worker and the -gathered logits alongside it. - -**`position_ids` reaches the worker in the engine's `[1, T]` shape.** The -worker does `position_ids.squeeze(0)` itself. A shell that flattens before -handing it over produces a silently wrong draft position sequence. - -**Nothing needs adding to `epilogue`, and nothing needs a `layer_idx`.** -`epilogue` is only consulted by `__pp_init__`'s `skip_forward`, so at -`pp_size == 1` there is nothing to register. And on this mode -`Eagle3OneModelSpecMetadata` sets `layers_to_capture = ()`, which makes -`is_layer_capture()` False at every layer and leaves `hidden_states` -unallocated — the trunk owes the runtime **no hidden-state capture hook** -(that is Eagle3's requirement, not MTP-Eagle's). Measured: nothing in -`_torch/pyexecutor/` or `_torch/speculative/` reads `model.layer_idx`. - -### Draft length is per iteration, not per request - -`_handle_dynamic_draft_len` runs **before** `prepare_resources`, so KV -allocation already knows the answer: - -1. `draft_len_schedule` maps a batch-size threshold to a draft length. -2. The current `scheduled_batch.batch_size` selects `runtime_draft_len`. -3. Every generation request's `py_draft_tokens` is padded or truncated to - **exactly** that length — the source comment names CUDA-graph replay and - the attention kernel as the reasons. -4. It lands on `spec_metadata.runtime_draft_len`; - `runtime_tokens_per_gen_step = 1 + runtime_draft_len`. -5. `runtime_draft_len == 0` takes `skip_drafting`, i.e. speculation is off - for that iteration only. - -**For modeling this means there is no ragged draft tree to handle.** The -draft length is a host int, constant across the batch, constant within a -capture. `cuda_graph_runner.get_graph_key` asserts it directly: *"All draft -lengths must be the same"*. - -Without a schedule, `runtime_draft_len` is simply `max_draft_len` every -step. - -**`draft_len_schedule` deadlocks under attention data parallelism, and -nothing rejects the combination.** Step 2 above reads -`scheduled_batch.batch_size` — **each rank's own local batch** — with no -cross-rank reduction, and attention DP does not equalize batch sizes: -`_pad_attention_dp_dummy_request` only tops a rank up from zero to one, and -`attention_dp_config.enable_balance` is off unless asked for. Two ranks -either side of a schedule threshold therefore run **different numbers of -draft replays**, hence different numbers of MoE collectives — and the -collectives pair by position, so the job hangs. - -Measured at world size 4: boot, CUDA-graph capture and warmup all pass, and -the hang needs real traffic to make the rank batch sizes diverge. trtllm's -own `HangDetector` fired at 300 s and hard-killed all four ranks, whose -stacks sat at three different points of one forward. **On an attention-DP -target, skip this knob.** With `tp`/`ep` alone the ranks share one batch and -the mechanism is sound. - -Setting it also **silently turns on `cuda_graph_config.enable_padding`** — -logged at INFO only — so it is never a single-variable experiment. - -### What changes in the trunk's own forward - -**One argument, and it is the whole of it on Blackwell.** A generation -request arrives with `runtime_draft_len + 1` query tokens instead of 1, and -that is expressed to the attention op through **`predicted_tokens_per_seq`** -alone — the value the total-token arithmetic uses -(`num_ctx_tokens + (num_seqs - num_contexts) * predicted_tokens_per_seq`). - -**Those extra query tokens stay on the generation call.** A one-engine mode -returns False from the runtime's `extend_ctx` predicate — "1-model has -separate logic for handling draft tokens" — so a generation request carrying -drafts is *not* re-shaped into a chunked context request the way two-model -speculation does it. The batch keeps its `[context | generation]` split and -the generation call simply gets a taller query block. What that block is -allowed to attend to — each draft position seeing the cache plus the earlier -positions of its own block, and no later one — is a **kernel** fact, so the -catalog contract for the attention entry is its authority, not this file. -Do not assume it from `mask_type` alone. - -**The spec-dec mask machinery is forced off at sm_100 and stays inert.** - -```python -# _torch/attention/backends/trtllm.py (was _torch/attention_backend/) -# Blackwell trtllm-gen spec-dec is enabled only for dynamic-tree masks. -self.is_spec_decoding_enabled = is_spec_decoding_enabled and ( - not self.is_sm_version_trtllm_gen_kernel(sm=get_sm_version()) - or is_spec_dec_dynamic_tree) -``` - -`is_sm_version_trtllm_gen_kernel(sm)` is `not (sm < 100 or sm in [120, 121])`, -so it is True on sm_100; a linear-tree MTP has `is_spec_dec_dynamic_tree` -False; the conjunction is **False**. `is_spec_decoding_enabled`, -`use_spec_decoding` and `is_spec_dec_tree` are all False and every -`spec_decoding_*` tensor is None — the same inert values the MLA columns of -`catalog/attention/thop_attention.md` already certify. **On a pre-Blackwell -arch this is not true** and the mask surface would need certifying first. - -The context path is unchanged: a context request still contributes its -prompt, one row of logits, and one attention call of the same shape. - -### The draft loop rewrites the batch between step 0 and step 1+ - -The single most important fact for writing the layer. After the first draft -step the worker mutates `attn_metadata` in place: - -| | step 0 | step 1+ | -|---|---|---| -| `_seq_lens` / `_seq_lens_cuda` | real | **filled with 1** | -| `num_contexts` | real | **0** (when a KV cache manager is present) | -| `num_ctx_tokens` | real | **0** (recomputed) | -| `host_request_types` | mixed | context entries overwritten to generation | -| tokens per generation request | `runtime_draft_len + 1` | **1** | -| a context request | the whole prompt | **1 token** | -| `use_spec_decoding` | as the engine set it | False | -| `kv_lens_cuda` | as the engine set it | rewound, then `+1` per step | - -**So the layer must read its phase from `attn_metadata` on every call.** It -is invoked N times inside one forward, and a value computed on the first -call is wrong on the rest. This is not the trunk's situation — the trunk -builds step-args once per forward because there is only one step in it. - -The read is cheap and sync-free. `on_update()` recomputes `_num_ctx_tokens`, -`_num_generations` and `_num_tokens` from `_seq_lens`, which is a **pinned -host** tensor, and the loop calls it (and triggers it again through the -`num_contexts` setter) at exactly the step-0/step-1+ boundary. Therefore - -``` -tokens_per_gen_seq = (rows - md.num_ctx_tokens) // (md.num_seqs - md.num_contexts) -``` - -evaluates to `runtime_draft_len + 1` on step 0 and to exactly `1` on step -1+, with no step counter threaded through and no device read. That is the -`predicted_tokens_per_seq` the layer's own generation call needs. Guard the -all-context case (`num_seqs == num_contexts`) rather than dividing by zero. - -**What a context request is fed on step 0** is `prompt[1:]` with the -request's first accepted (golden) token written at its last position — -`_prepare_context_input_ids`, shared by both MTP flavours. Its hidden states -are the trunk's, in full, at full prompt length. So step 0 runs a real -context-phase attention for those rows and the layer needs both phases, -exactly like the trunk. - -### Attention DP: the padding basis changes source, and reading the wrong one is silent - -Under `enable_attention_dp`, §12 established that every rank pads its token -block to `max(attn_metadata.all_rank_num_tokens)` so both collectives run in -their uniform form. **Inside the draft loop that list is the wrong one from -step 1 onwards.** - -The worker passes the right basis in **as a keyword argument** and leaves -`attn_metadata.all_rank_num_tokens` holding the trunk's value for the whole -loop (it saves and restores it around the loop, but does not maintain it -during it): - -| draft step | `all_rank_num_tokens` kwarg | -|---|---| -| 0 | `spec_metadata.all_rank_num_tokens` — the trunk's token counts | -| 1+ | `spec_metadata.subseq_all_rank_num_tokens` — the per-rank **sequence** counts | - -`subseq_all_rank_num_tokens` is set to `all_rank_num_seqs` by the engine for -the one-model modes, which is semantically right: from step 1 every sequence -contributes exactly one token. - -**Pad from `md.all_rank_num_tokens` inside the MTP layer and the collectives -mispair silently.** `catalog/comm/allgather.md` and -`catalog/comm/reducescatter.md` both certify that calls pair by *position* -on the communicator and that at equal byte counts a divergence produces no -hang — every rank wrong in 98–99% of elements, bitwise reproducibly. The -defence is structural: give the layer's padding helper a signature that -**only accepts a passed-in list**, so it has no way to reach the metadata. - -(A `_dp_rows` that asserts `all_rank_num_tokens[rank] == rows` would in fact -fire here rather than corrupt, since the two counts differ from step 1. Do -not rely on it: the assert is a property of one target's helper, not of the -rule, and it is exactly the kind of guard a later edit removes.) - -### KV cache: the engine adds the layer, and the target declares nothing - -Under a one-model MTP mode `ModelConfig` raises the pool's layer count -itself: - -```python -num_layers += spec_config.num_nextn_predict_layers -num_attention_layers += spec_config.num_nextn_predict_layers -``` - -so the pool is sized for `num_hidden_layers + n` layers and the MTP layer's -own attention addresses layer index `num_hidden_layers`. **Nothing in the -target's config stub or manifest declares any of this** — but the layer's -attention calls must pass the right `layer_idx`, and the pool-addressing -surface they use has to be certified at the new layer count. - -**A `speculative_config` also raises `max_seq_len`, and not by the amount the -extra-KV-token helper suggests.** Three separate terms are added to the model -engine's `max_seq_len`, in `py_executor_creator`: - -```python -if not disable_overlap_scheduler and spec_config is not None: - max_seq_len += spec_config.tokens_per_gen_step - 1 -if spec_config is not None: - max_seq_len += get_num_extra_kv_tokens(spec_config) # max_draft_len - 1 - max_seq_len += spec_config.tokens_per_gen_step - 1 -``` - -With a linear tree (`tokens_per_gen_step = 1 + max_draft_len`) and the overlap -scheduler at its default (**on**), that is `3 * max_draft_len - 1` — 2, 5 and -8 at `max_draft_len` 1, 2 and 3. Measured: a `163840` model cap becomes -**`163848`** at `max_draft_len: 3`. - -**So read `max_seq_len` off the metadata, never compute it.** A target that -derives anything from the config's own `max_position_embeddings` — a rope -table sized to it, a bound asserted against it — is off by that amount the -moment speculation is switched on, and by a different amount per -`max_draft_len`. The failure is a silent out-of-bounds read on a rope table, -or an assert that fires on a legal config. - -### One more construction-time trap, from §4's weight lifecycle - -**A reference to another module's parameter, captured in `__init__`, stays -bound to the meta-device shadow.** The engine materializes at t1 by -*replacing* tensor objects, not by filling them in place, so a drafter -container that does `self.embed_tokens = core.w["embed"]` at construction -holds a meta tensor forever and fails at the first draft step. §4 already says -an unregistered parameter "stays a dead shadow"; this is the adjacent case — -the parameter *is* registered, on the trunk, and the copy of the reference is -what goes stale. Resolve it lazily (a `property` that reads through to the -trunk on each access) rather than caching it. - -### CUDA graphs: the whole draft loop is inside the capture - -Capture wraps `_forward_step`, which calls the shell's forward, which calls -the worker — so every MTP-layer invocation is captured. The key is -`(batch_size, draft_len, is_first_draft, short_seq_len_mode, -is_all_greedy_sample)`, and the capture set becomes `(bs, draft_len)` pairs -rather than bare batch sizes; with a `draft_len_schedule` the runner also -captures one extra `(max_bs, original_max_draft_len)` graph, whose stated -purpose is to keep a later graph from resizing the shared attention -workspace and invalidating pointers baked into earlier ones. - -Consequences for the layer, all of them §7's discipline applied one level -down: - -* The step-0 / step-1+ divergence is **re-traced per draft step at capture - time** and frozen at each step's position in the loop. Reading the phase - from `attn_metadata` every call is what makes that correct — a cached - first-call decision would be baked in at all N positions. -* Host ints read from `attn_metadata` (`num_contexts`, `num_ctx_tokens`) and - the `all_rank_num_tokens` kwarg are per-capture constants, which is the - host-derived class §7 permits. -* `attn_metadata.padded_num_tokens` is **`None`** unless torch-compile - piecewise CUDA graphs are configured; without a `torch_compile_config` the - padding path is not reachable and the base shell's slice never applies. -* The attention-DP replay gate of §12 is unchanged: all ranks - generation-only and equal batch sizes, or the step runs eager — and eager - is where the layer meets a real mixed batch. - -### The acceptance rate is the only correctness signal - -`stats.specdec_stats.acceptance_length` is computed per iteration over the -generation requests that carried draft tokens: - -``` -acceptance_length = (accepted_draft_tokens + requests_with_draft) / requests_with_draft -``` - -i.e. the mean number of tokens a drafting request produces per step, `1.0` -meaning every draft was rejected. - -**This matters more than it looks.** Rejection sampling guarantees the -output distribution is unchanged, so a *miscomputed* MTP layer does not -produce wrong text — it produces correct text more slowly, with every draft -rejected. An accuracy gate cannot see it. A subtly wrong layer (a -transposed concatenation, a norm applied to the wrong operand) lands -somewhere above 1.0 and well below the reference, which no other measurement -distinguishes from "this model is just hard to draft for". The reference is -the same checkpoint under stock in-tree modeling at the same load and the -same `speculative_config`; acceptance is a property of the model, so a -reference forced onto a different `max_num_tokens` / `max_seq_len` to boot -is still a valid comparison. - -### Still not established here - -Vanilla MTP (`n > 1`, one layer per draft position) and its per-layer -sequential loop; Eagle3 in either form, including the hidden-state capture -hook and `apply_eagle3_fc`; tree drafting of any kind (static or dynamic -`eagle_choices`), which is also the only way the Blackwell spec-dec mask -surface becomes reachable; two-model speculation; and speculative decoding -combined with pipeline parallelism, where `skip_forward` and the `epilogue` -path start to matter. diff --git a/tensorrt_llm/_torch/staircase/explain.py b/tensorrt_llm/_torch/staircase/explain.py index 054eb8d01d56..87fb96f8ee14 100644 --- a/tensorrt_llm/_torch/staircase/explain.py +++ b/tensorrt_llm/_torch/staircase/explain.py @@ -58,9 +58,8 @@ def build_parser() -> argparse.ArgumentParser: def main(argv: Optional[list] = None) -> int: args = build_parser().parse_args(argv) - from tensorrt_llm._torch.pyexecutor.config_utils import load_pretrained_config + from tensorrt_llm._torch.model_config import ModelConfig - pretrained_config = load_pretrained_config(args.model) world_size = args.tp * args.pp mapping = Mapping( world_size=world_size, @@ -71,14 +70,17 @@ def main(argv: Optional[list] = None) -> int: enable_attention_dp=args.attention_dp, ) - ctx = StaircaseContext( - pretrained_config=pretrained_config, + # The engine's own loader, not a second reading of the checkpoint. It is + # what fills `quant_config` from hf_quant_config.json, so a tree that gates + # on quantization is explained against the same value the engine will route + # on. It touches no CUDA, which is what keeps `--sm` usable off-GPU. + model_config = ModelConfig.from_pretrained( + args.model, mapping=mapping, - sm=_sm(args.sm), - quant_config=None, - spec_config=None, - is_disagg=False, + moe_backend="AUTO", ) + ctx = StaircaseContext.from_model_config(model_config, sm=_sm(args.sm)) + pretrained_config = model_config.pretrained_config arch = (pretrained_config.architectures or ["(none)"])[0] routing = routing_module(arch) diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/TARGET.md b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/TARGET.md deleted file mode 100644 index ce5def4aa247..000000000000 --- a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/TARGET.md +++ /dev/null @@ -1,1631 +0,0 @@ -# Target: deepseek-r1-0528-nvfp4 / sm_103 / dep4 - -## Identity - -| | | -|---|---| -| Checkpoint | DeepSeek-R1-0528, modelopt NVFP4 export (**v1**, `DeepSeek-R1-0528-FP4`): 61 layers, hidden 7168, 128 query heads, `q_lora_rank` 1536, vocab 129280, untied embeddings; MLA attention in bf16; layers 0-2 dense (intermediate 18432), layers 3-60 MoE with 256 routed experts at top-8 (group-limited: `n_group` 8, `topk_group` 4) plus one shared expert (intermediate 2048); YaRN rope, factor 40 over an original 4096-position window; **NVFP4 MLP weights with an fp8-e4m3 KV cache** (`hf_quant_config.json`: `quant_algo: NVFP4`, `kv_cache_quant_algo: FP8`), attention / router / embedding / lm_head bf16. The checkpoint also ships a bf16 MTP module at layer 61 (`num_nextn_predict_layers: 1`), which the identity config does not load and `configs/mtp{1,2,3}.yaml` do | -| GPU arch | sm_103 (GB300) | -| Parallel | dep4 — `tensor_parallel_size: 4` + `moe_expert_parallel_size: 4` + `enable_attention_dp: true`; world size 4, one rank per GPU. The engine builds `tp_size=4`, `moe_ep_size=4`, `moe_tp_size=1`, `pp_size=1`, `enable_attention_dp=True`, and construction asserts every one of those | -| Registered class | `StaircaseDeepseekR10528Nvfp4Sm103Dep4` — a synthetic architecture name no checkpoint declares. `models/deepseek_v3/routing.py` rewrites `DeepseekV3ForCausalLM` into it when the config shape, SM and topology all match; the checkpoint is read unpatched. Per-target names mean one process can hold every target at once | - -> **NO GATE RECORD HOLDS FOR THIS TARGET.** Every result below — boot, -> gsm8k, the MTP acceptance comparison, the whole *Performance* section — -> was measured on **sm_100 (B200)** through the pre-move standalone harness. -> This target is **sm_103 (GB300)**, and certification is per architecture. -> The numbers are kept as provenance, but this target is **ungated** until -> they are reproduced on GB300. Read every "passed" below as "passed, on -> sm_100, before the move". The MTP variant needs all three of its gates -> repeated, acceptance included — see *The MTP variant*. - -Checkpoint sha256 — the checkpoint directory passed to `--model` should -resolve to files with these digests. **Routing does not check them**: it -fingerprints the config's shape -(`num_hidden_layers`, `hidden_size`, `n_routed_experts`, `q_lora_rank`), so -a fine-tune or a re-export of this checkpoint routes here silently. That is -the deliberate trade in `models/deepseek_v3/routing.py`, and it changes what -a gate record means — not "this target passed" but "this modeling code -passed *on the checkpoint with these digests*". Run it on another one and -the result is ungated (recorded at: -`umbriel-b200-027:/home/scratch.trt_llm_data/llm-models/DeepSeek-R1/DeepSeek-R1-0528-FP4`). -The source `config.json` this stub was patched from hashes -`a80de10ccf70e7e98dcc6730b45872298937b232d3e8eb0321cfd2bce4cb40e0`; it is -the one checkpoint file that is copied rather than linked. - -
-169 linked files (163 safetensors shards + index + 5 aux) - -``` -c139e0cd9aa4c418ebd38ebae9aff73b4eaacc7723e25d2490dfe41c02bb8708 configuration_deepseek.py -0ef9febae6b6087f4822b02bc9a1c03a83263dabaa931fb9155d061c8951ca07 generation_config.json -36dc07886afcd1679ebf5328a325f9f9629d7a0ddfa644ee6427b10874b2c295 hf_quant_config.json -991aaf2f2c7a5b21f6c6708bd736285640ebec92790d7882958dcc26b0c9f6f6 model.safetensors.index.json -ecb6f9fc369894346f0511f4074ca75cee5cd5f3b06d02f1ba35fcd39f8e121d tokenizer.json -a58700120f68f96faf27a8921876e5e3dbb95f66355c6b331312354cbbf792f9 tokenizer_config.json -bbf04b49f4cdb3c2cb139d84e9701089cff14d08b95b969379aa5b57a63a9242 model-00001-of-000163.safetensors -ac6cb6c4a30d18ee8d24c5237c1b0f9cb30ab24aa9f4e4758a8beb5125360720 model-00002-of-000163.safetensors -c280f0fc3241c685b033c26c92fd4c3fca34285f288922e3d8119ba7d779aa32 model-00003-of-000163.safetensors -ca5f80f8d62248366271495d0a9a95953dd5fb06237231a3f8465207d475edd9 model-00004-of-000163.safetensors -5a454d4577381ec2fb9511015c9a06024fe592949badd45d36d2efe24dc1a0f2 model-00005-of-000163.safetensors -9bb8873103ba89dd16efad52e7846abe7a0ee40af47d4739f3503a82160b1001 model-00006-of-000163.safetensors -48a5a4d26db3e4a11d1bacdfd5c94e27bfc90b14750c165ee7ebca2d81b924e4 model-00007-of-000163.safetensors -b52d09d1147f1704bb39096221afbd992c6a1fe1114358975ddca79ea214361e model-00008-of-000163.safetensors -0f0274233de53fe22c50262c36d6154579af3cc29b83d3b2d601761b0475c810 model-00009-of-000163.safetensors -2658cef0619d4912da43964d312fa1448eac38f5e44e0313218367b888f8e7ae model-00010-of-000163.safetensors -360189a2c32647c16efbf465a9f2bb3f607e47f0848c6dbfca1eae455d8fa6dd model-00011-of-000163.safetensors -8ae7554ea1a6d65c9e194dd99254bc0797c9fd848a7be23fca94e2defc56ce9d model-00012-of-000163.safetensors -dc6ed60abb1be1223979884405aad63f35e5eec42c28b73773bd19f834c90065 model-00013-of-000163.safetensors -7d359af3e3ccf70794d83235e5d022392cb76ade1a34d4fd8797c588c212056a model-00014-of-000163.safetensors -afbdb1f5da08dc38ecaf900ae8866e4687a3ec9e2b21b16d0cba7fa3d529b74a model-00015-of-000163.safetensors -b0ed219a52ec2d617ecc750963519df23d5e27160621d2fd964b506ba017a5bc model-00016-of-000163.safetensors -bbe5cb864267a27a55add115d27f9fdec6d47405f626de9774e0f79ed271de8e model-00017-of-000163.safetensors -ca0e49de986c19b866bf2c330ca91f8c0a092ee56e19ce832267f375174cc496 model-00018-of-000163.safetensors -0b1db3ad16aea8afdcc667011dabc11addf41c20b4c10234508117578278ba21 model-00019-of-000163.safetensors -c016097de09ca54826380158d531ef7121ccc6f5aa44dbbe255dd4c80e817e7d model-00020-of-000163.safetensors -4929c6acd11fa3f87465aef13e5ce7e3d24c83a1babac6d1c99c0201bba7985b model-00021-of-000163.safetensors -05f55fd3fa47fcebc0c4910204672fb602b30f364cf088544c6ec1d67f7d33ba model-00022-of-000163.safetensors -8f861d34dfd9299c043ba0c351505bef3a107d8f95ea8a46ceda8a3ebf6f2625 model-00023-of-000163.safetensors -95362bd84e97a23460933f630b69d9ff720b7d1da15c544b97ade3481bcca444 model-00024-of-000163.safetensors -c1db194b1d6a0f2c6a1935e5105622174007db387c79e6c1b5430bc3273c778b model-00025-of-000163.safetensors -d181f80f2aaf9cb7bdf68e421b950d64a583a52420ce5fd0d3b411ed33c3ea66 model-00026-of-000163.safetensors -04ee66ce0f3e51f407d5a68e4b29c0b63e9bda0a5e88896c39a097aaba3a2dd2 model-00027-of-000163.safetensors -9167d91670023079e629de55d02be50a06a720b125da70b9556481b9f23e49d3 model-00028-of-000163.safetensors -ecf62e772486af899afb440540803550480343b66e5dc096e1e4a8dea6a44154 model-00029-of-000163.safetensors -18dfa7d3ed08561b4109e83e49fe4a073ce5b1a23d374bce98dc5bab832f4948 model-00030-of-000163.safetensors -70e471763de9f5ad8b12cc26be6f44d48c8ebcbb7f1c4a57f4eaabf112f7f57f model-00031-of-000163.safetensors -fd86c19b1c61cd171d5f573cd62a85d09baae34b54b678ceed921c522412a406 model-00032-of-000163.safetensors -e4a2457ca0b979830ba5bdff1b9e17be8ec227c866a8adeeb455abc87a780761 model-00033-of-000163.safetensors -a69fad78c8fce390011d98b140789010e2fcb83129e7ca17ce1bce64bfd2b759 model-00034-of-000163.safetensors -a3898dfd91f59bec65c82ae04b707575501e32f43d2f1dfb26174912f01d0fd5 model-00035-of-000163.safetensors -2e0fe82025efc32f27a290438b2e607c89ac394716e97b3c9b6cfa27e2480623 model-00036-of-000163.safetensors -77afb4330256683ef7962bf43cae8762f80fb0edf889be55073f421c3ad24930 model-00037-of-000163.safetensors -8e2498c311a925d54212cee7d41ce4f98c345d788adcddcc749053292511e4bd model-00038-of-000163.safetensors -0ee6089eb7bfef59dde3b154b1e133f0034347b391b7ed3c903d6e6525b20ea9 model-00039-of-000163.safetensors -54bee64d4e307768da9a6bf437e00eec91e3a4986a2358d2d5eb41d915f4dd0e model-00040-of-000163.safetensors -e8fc96ae38097b90d1b4424f7610186ff7102cfa67930823a5500906f0398b7f model-00041-of-000163.safetensors -2adea35e960e5e35f95c62c0bc97d883248aa79d186da815535336fb0ea914c3 model-00042-of-000163.safetensors -bad155fdce361b40b328edbe35e3995cf6e149e1f798eba415db58ff998880fb model-00043-of-000163.safetensors -17e353f93d66c07708586a8374012a8cb68939a831931bb215508553adcdbeb4 model-00044-of-000163.safetensors -215c3ee090030d9f6a12cac238e1cdc14c1bf7d33a45eb5a84114986a65e988c model-00045-of-000163.safetensors -ba7377f46d611140338dd6efde6ffbb47fce763a70ffe2d0746d5175ad799eea model-00046-of-000163.safetensors -d1ba05c9c00dffcfe86e08ff0109d83d68694179ef45dbcc3766c32ab0316922 model-00047-of-000163.safetensors -171a534fa7797a90cfcfc5b340a25013359750ef37f637b89d7ce5665b13d191 model-00048-of-000163.safetensors -d4c05d603481809dd5a6910461e08ed2387c6b7459438e93f201d00c23222690 model-00049-of-000163.safetensors -558e86f0977914e39665e2a4756082cc5a9ea1d2d99ef4a863c69adc84febf66 model-00050-of-000163.safetensors -5e07cadc4d02a7783e0ea6edfeeff88bbcf02b78679ad4efe32592c54d154208 model-00051-of-000163.safetensors -83c7f0c70e1bb1b0ff6e0915f40d1ca91d98c7c63aebeae50265fa214b8005ad model-00052-of-000163.safetensors -96fa3cc110f476e1414000da8431ddac767c694143a20f07cfb14de62bdbe419 model-00053-of-000163.safetensors -7b170e72c34df97980ed4988a261f1732f2b573244f3345456fefde9874609bb model-00054-of-000163.safetensors -a5b7b9001f1ad50039e63989eb90be1f1e5d9dc0654544a87073627b73289eae model-00055-of-000163.safetensors -98d711b4214ef67598ac60887c12b1f2186e7aa3ec97d95858389c799f8b46f7 model-00056-of-000163.safetensors -3ec5916a49938b634be79548cbe6944660fdcacbdb255a6393932fd404e796ee model-00057-of-000163.safetensors -0d60c7c3f78cd237164c0ed466bed5e73e53328a1db0605aabfa430920210003 model-00058-of-000163.safetensors -0eb7d1c45c8444dbcfc152b44280357543fc1b798f9dc38cd7cb555204f55d99 model-00059-of-000163.safetensors -92bc2a13fd190ecbc25dfab34ceb210026a6674b984609e5642c6c152bb3c6ba model-00060-of-000163.safetensors -4d00fd183118a24cbd5ebb357238505fc252b117e0cbe4e469605639d10b0b3f model-00061-of-000163.safetensors -14ac6babc7b5ec03944a0253aa3276d114e07e36221d797fca50292ab4f371da model-00062-of-000163.safetensors -5d3a9d7c49fe341f8b7f5e21c3ae63a0718f1a3227b18dac393e2565d3154e09 model-00063-of-000163.safetensors -49a1ddf499f4129111ffb50edac498fc23312f98ebc1b69e8841e853ac609355 model-00064-of-000163.safetensors -5341f93faa5c992d684f234c3a479aab87eb3b4576b1f978eb119476131a8ef2 model-00065-of-000163.safetensors -72d183baf3c7c2c7b816e286867e4f1e8398685b6c76110f778dbda4d85ce6fc model-00066-of-000163.safetensors -5ec23f22dcae3d5d883dc925ce70d463a6c9bba3c1f9b5150d447ccfee451e77 model-00067-of-000163.safetensors -205d88536e4a52d83ebaa62005665d531e6f0f8440c280772e7ae61cad23d89f model-00068-of-000163.safetensors -dde4b037231af225ed529148ce14e6dcf10ea3bcb5f79c47aa5c7fde35770d9e model-00069-of-000163.safetensors -7f7ee508718627657ef87966a885a9956f3734674cb988fcb0bbb9e22de005d1 model-00070-of-000163.safetensors -1824edecf1bbce9246038693adf729d0484a7148aff8ce7c4c2f302184874cf1 model-00071-of-000163.safetensors -633dcca915de5aa3629f2b53345a86c5eb4b4bcd4731aa3fbf5e3b5ee4c25bef model-00072-of-000163.safetensors -b53bc434459993f27ce202bd913640f64129c6d91cea6445fddaf6869ea3e618 model-00073-of-000163.safetensors -86680218bc3719271bc7c47f74810b469fbdeb7cb4370433121cc2c8e1416dca model-00074-of-000163.safetensors -0eac1ff5b435996ca17677429909307081eced38c7aa8543f54569e509a18b7c model-00075-of-000163.safetensors -6b49ec351aa02e4d152f95d2d7fb25fefd537a48ab12941ef1ede7b8f2879601 model-00076-of-000163.safetensors -36d338fba2c4a1543a691c72ed2297aad49c46230219c3e570c41a30b6430671 model-00077-of-000163.safetensors -1611afd72d0e702e4edd0e519525f0fd70763039b18db539d1d1a465d1262732 model-00078-of-000163.safetensors -c2cbd1c7ac0762ae69a453fdca414ff3b41d6a831e36f33a0e365eeb3a4f7796 model-00079-of-000163.safetensors -5968523bd66250a4250e41fa01fd1868d558efc4df271d17e44d8a3931bb9f4f model-00080-of-000163.safetensors -9f1bf3ae7b459e3471ead763cacddaf77b030c1bd153ae7c6c330fbe9309a75a model-00081-of-000163.safetensors -b6e47522fcbcd15ee2d4c8a7921954d48403b3daaf45d1fe0b49f3ab6e7c7bd7 model-00082-of-000163.safetensors -a641e0c5932548bc1f93d94bdcf435898656123298425c04c5cfa7f17fc756d3 model-00083-of-000163.safetensors -f43e8dbbfbdbf808a11d8c42668aed50855b668e0f3c1a3cbd007e6b03652e78 model-00084-of-000163.safetensors -894fe387005e9b73fbe8483c5f4948f127527fd90dcf07bfb1dbdedf74357b3e model-00085-of-000163.safetensors -217e2a040495a5a3825c308638806b750508d23ce78f996d53459d83c848d50e model-00086-of-000163.safetensors -9eb75ba7b09bffab0a4d45fe739ab339db47352c227ca0ee15beed80cc1e7da5 model-00087-of-000163.safetensors -f4497ac9f44495704458aa0542395dc29abf97d8a3abbc16e36b4ef522068725 model-00088-of-000163.safetensors -1cde75047119b80d9bb27333d9d52340b56c5c3879cb326f7418dd589dfd73d8 model-00089-of-000163.safetensors -e3b9df1600cabbb33a0963412a8a919fdeb75472938a46346fb9207876940da9 model-00090-of-000163.safetensors -2aa0a431652ac12273e3a0d38870baf9810f5d748be16e1ca37cccf1381a817f model-00091-of-000163.safetensors -a1c78001c5e2ac6ae5a8212e87440091dd56142c0f8829da3c463e0f717725bb model-00092-of-000163.safetensors -9bec57d3138b23e9c9e5dc53e9b6ac63adfb27e35584232b88ab4bcea0653dfd model-00093-of-000163.safetensors -b1bcc0b137739cacff96dd6c2fe0861af0102244be4865ddd50ff9fe71bc5c0a model-00094-of-000163.safetensors -76c031439199fdba5c01bc8ad571a0f8655c3ce1a04bb466338c25ca5256cd6c model-00095-of-000163.safetensors -2f10fce8ee1187960f53cf5bbc5750d4033688f0e6bb24fb92447dd1bffb9ac4 model-00096-of-000163.safetensors -ceebab7ff4022aaa21e15ff05e6a3ebd4a334e10049bf721d3e2519b98514737 model-00097-of-000163.safetensors -87ef9077d9212db15cb8fe89342501c9cea08f3c5c96a31235db683d8751221f model-00098-of-000163.safetensors -cac29cbf8e08520c89bcf27a08a77343a5f6c2aea1771ee7a1bbdf64474e90d1 model-00099-of-000163.safetensors -5ee6438bc3414412a7ab279b3e0b72d631245a7e27975a895d0e94f5aeef90d8 model-00100-of-000163.safetensors -d99c26d1874005fee504b44ac96c496a496ff25cb87ff755832e926541a48492 model-00101-of-000163.safetensors -dff5cc926dc3c1684fbec174508b5ba5abc67c65e1acf6f20354c3acce74b05a model-00102-of-000163.safetensors -babcf9affea89d453f2023379a34b32d0ecc4f9b2787255c0fcb8889cabc093a model-00103-of-000163.safetensors -45e2f2dae475a13c6d10d7a2bfde28b8d65765d421e73d08c527d8c1f5b9392a model-00104-of-000163.safetensors -1b276c9186ee2f335e0930840d406a59776d695433081bbabf4c5ba25e79bd05 model-00105-of-000163.safetensors -a9a5527aede2cb7f292671d86f2864d6fc843396466d16f60e74743d7fc0a6ee model-00106-of-000163.safetensors -a950aedf442134db181e4ac440c9007178c36b0ef57245f7f49a3ee5bf50540f model-00107-of-000163.safetensors -e7a7bf9ae8f6ea8ecc8750cd44ccb3814eafbb2f0b88fc7482e6ee39e9d8dc1b model-00108-of-000163.safetensors -cb9ac1d529250b0524bd86864725144d45d5e91c70e77375ff0bf30434633250 model-00109-of-000163.safetensors -a6999a13f265c55995ba5c1e0260d64469dc714949a4df300b37af6ae9a1e61a model-00110-of-000163.safetensors -7ada750ccc74bcc4c283332ce7ca573c74254ed61931737b1d739e9b633d0e78 model-00111-of-000163.safetensors -3fbfd40aa863cb1ae17eb175958b95e168325c1cae9c8affe639d88c09e1960a model-00112-of-000163.safetensors -61dea7ce4e66bf886e2ede5ff2a9df65660b7492eb820116b4bac9d7419b1e85 model-00113-of-000163.safetensors -12fa6385c3932d28de5300199e53aadacc8a0d3638d03c42f5cbac71f2d8574b model-00114-of-000163.safetensors -22e0a03279a7d2ebe6b0bfb6d5e2ce98b5da39f518cb9a82c04b4d5b52096fa1 model-00115-of-000163.safetensors -c1c9fc85b7cd9d403a2b73098e1d81a10e12fdffd9549c669a2ae515d851052a model-00116-of-000163.safetensors -4dc24b7fcbc6517c62222ca7da18bc664d5f1da5e2efcc8bcb8d5fd590652b78 model-00117-of-000163.safetensors -794d6d03b6447e7f31882ab070bbcede362becd464139cbcca5970d2f41f12ab model-00118-of-000163.safetensors -dae3a9aceb361372feef019635c17a1bda7c826dba8cacbcb3881ea7de2139db model-00119-of-000163.safetensors -93d8435bb4af72314030be201f0f5e8f431fe6e48f79e7918e46d7ef92f2cb59 model-00120-of-000163.safetensors -9aa69352d752be9438a99776732b4e0b62ff589151bcf8d8cfb5e9dfecc45a0a model-00121-of-000163.safetensors -8f7bc0398f6417d9d60567b0687406a951a6190bc8df46be675b2a90c69aa4cf model-00122-of-000163.safetensors -1aa442507c1ee8a26afef5950aa87e6834724231f0c2852b642254c3afb79eaf model-00123-of-000163.safetensors -6e4185179e96b94dbe61de9db3fc68386f524bd5119ef5025c67608f852479be model-00124-of-000163.safetensors -e03dae8708f37980fea4471fe9cf88d9ac703fc2db7b128985fe759cd9922fa8 model-00125-of-000163.safetensors -7c910c972cbfa6f8afcd6cd55e97a90688429037e429b0c67b048fb48bff53c7 model-00126-of-000163.safetensors -bae19ab02d0fb95a32fa08867741e437592c9857e39b4cdb9279da627948a647 model-00127-of-000163.safetensors -fa3c946d0c051bfd44077f497971dc21b93750a3412f541bc34b742a0b504398 model-00128-of-000163.safetensors -8d1afa3ef67315e898c3fcca58e522b278c52ee6db14e3beac8d3f3da9aace23 model-00129-of-000163.safetensors -4a0db8ab07b13372e598790ccb5222e0ebc345e514575169c98ce30bf1b54f8d model-00130-of-000163.safetensors -95e59b883e43b7259e2f061c5ae47071076f8d9ea26d025d12e86c3979f15007 model-00131-of-000163.safetensors -e7310c0a55bd9126ae844d8b1c9074c476d5ff31b9d8902f27114dbde16fce23 model-00132-of-000163.safetensors -184d0abd9be8e5b966cbbf17264b204872450c0fae4c29a82083e39e273432cb model-00133-of-000163.safetensors -4f0d6beb59421c0bc1560495dc89324de1f2e7a794b65ffb854534a21b17b50a model-00134-of-000163.safetensors -ec9fc4a523c9aabe818356a72a8f032540ebd22273a03976bfa8f9af11a5d8b6 model-00135-of-000163.safetensors -9862e7589f2288d86e45d9377203935621b982a0e376101996196fdfc4fd8cad model-00136-of-000163.safetensors -395a89c3c93071004c8d73810f73c23954449acc03efbe220a6b4cf576afdb60 model-00137-of-000163.safetensors -a747eb61204cb7d1f78acfcae182c75781a7839ca185132a0421acd72668eafc model-00138-of-000163.safetensors -27651e88d2fbfc582e0de685ec318c6e9af6e8ad7c13233a2bb3c33ea14ca634 model-00139-of-000163.safetensors -eaa5f027070aed1ad02c4295218c99f3c91c1ad067c4bf60f56ec69219fe0949 model-00140-of-000163.safetensors -6c924b6d4b1fc90bd0aa18a9421125a6c54d317516fe58ad8b8d07d08f12c7b7 model-00141-of-000163.safetensors -4e3cd496cb1af08ea801b831982573883b158bb60a1a9f3576f3986a4cff3467 model-00142-of-000163.safetensors -37716bff9e48a87fd4801e644527b2fd31468c4f6866bda70ff9db409902c593 model-00143-of-000163.safetensors -760f7034122d6ef1738bb74fcf4c9230e5a7ee9024b499931e05760c66aff3ac model-00144-of-000163.safetensors -fa1203747f79017c0324b5564ea70dd611b237928054dbe3a8f9cf4217cea66a model-00145-of-000163.safetensors -c05235067f189ade00c6d0ce7faf2ba728012ed78863352cc9d613da2cf959ba model-00146-of-000163.safetensors -248b214fbd8455664ca1183cbc8f1ccd8f6fc3a9e3b2898798b5b1a716225811 model-00147-of-000163.safetensors -738de1a2b1844982c6f29317e508a9c8b57c218dd4aac058ca113e90f265ba96 model-00148-of-000163.safetensors -c35cd109b7895cf7a74b6f561f47145a99daccdead56fce6b256a0914aceaf25 model-00149-of-000163.safetensors -a08febee45f832a99e4029bc73c0093a40b4f3d7ef2f1690088982ba65fbe00c model-00150-of-000163.safetensors -eb659a6d3d6dcce0c0a051ea5b2184fc7c190c6037859d197637abd6ee305f6c model-00151-of-000163.safetensors -d19f2846145b89a456e71ef2e6e3f00a4ede37c29f8e4ba6cef55159a3230922 model-00152-of-000163.safetensors -801cc06d203ea5e4cbe9097622e7b971174249c5b1b0298b3560cb831946dfb6 model-00153-of-000163.safetensors -42e9ce66a658f447f16c3bca1cb9d843a66df989201ef7fe2696138309d2ef25 model-00154-of-000163.safetensors -ed61b343d9f632a15d1b8b98a103e7ca9077a8e57c92c798f82472d4f7e9776c model-00155-of-000163.safetensors -8d2689d216da51e5a3d507a9a1eec347200fc1146a9c84ebd1e30d9d93ec779a model-00156-of-000163.safetensors -4413b79d84ec94cb627440e096b36012184cce599151355691213348b6e4b795 model-00157-of-000163.safetensors -fe09a5c54db66630ceea47091df108d8235736abe679700786bcf7294f0e9cbc model-00158-of-000163.safetensors -79b6b4bb774569c433cf9b4ac391b4fae11a6ebcb235b213e4c23b2b66d3c477 model-00159-of-000163.safetensors -243cb503cc88e16170bad7a216742da5262032bc640730fc0b3e348a7b14f451 model-00160-of-000163.safetensors -4981af366b2bedb6ec8d6cd2eeae8bd2c027c3ccaf001ab54af207883da0cdfd model-00161-of-000163.safetensors -783d3e45422f562d9dc7b11495dc7a44834f0cfab59e0235432d1188170565f2 model-00162-of-000163.safetensors -8ad1e6011ac0ebd811cbf00edb62ad095be42b01266cf4638492fecc1e9020fa model-00163-of-000163.safetensors -``` - -
- -`model_dir/config.json` is the checkpoint's own config with two deliberate -divergences: `architectures` patched to `["StaircaseForCausalLM"]`, and a -`dtype: "bfloat16"` field added beside the checkpoint's `torch_dtype` -(transformers 5.x renamed the field, and the shell materializes its own -`lm_head` from whichever surface resolves — both were observed present and -`torch.bfloat16` on this build). `hf_quant_config.json` is linked because the -engine reads it: it sets `quant_algo=NVFP4` **and** -`kv_cache_quant_algo=FP8`, and the target asserts the latter — the fp8 latent -pool is what selects `quant_mode` on every MLA call, and a checkpoint -declaring no KV quantization is a different assembly. Nothing in the stub -carries the topology; that lives in `llm_args.yaml`. - -## Version - -| | | -|---|---| -| tensorrt_llm | in-tree — the target moves with the trunk, so there is no version to pin and none is asserted. What *is* asserted at construction is the SM version (`_SM = (10, 3)`), which the pin used to stand in for. The gate records below name the commit they were taken at | -| torch | 2.11.0+cu130 | -| transformers | 5.5.4 (the config surface the engine hands the target) | -| Attention metadata fact source | `TrtllmAttentionMetadata` (TRTLLM backend) | - -## Vocabulary - -**Two audit roots.** This target carries two forward paths: the trunk's, which -runs on every engine step, and the MTP layer's, which runs `max_draft_len` -times per step under a `configs/mtp*.yaml` variant and not at all without one. -Auditing only the first would leave the second free to use any op unseen. - -### Trunk — the `StaircaseCore.forward` closure - -`forward` plus the private methods it reaches: `_dp_rows`, `_dense_mlp`, -`_check_step_contract`, `_build_step_args`, `_moe_chunk_sizes`, and — through -one first-forward branch only — `_rope_tables`. - -| call | catalog entry | -|---|---| -| `embedding` | `torch/embedding.py` | -| `flashinfer_rmsnorm` | `norm/flashinfer_rmsnorm.py` | -| `flashinfer_fused_add_rmsnorm` | `norm/flashinfer_fused_add_rmsnorm.py` | -| `cublas_mm` | `gemm/cublas_mm.py` | -| `bmm_out` | `gemm/bmm_out.py` | -| `mla_rope_append_paged_kv_assign_q` | `attention/mla_rope_append_paged_kv_assign_q.py` | -| `load_paged_kv_cache_for_mla` | `attention/load_paged_kv_cache_for_mla.py` | -| `mla_rope_generation` | `attention/mla_rope_generation.py` | -| `thop_attention` | `attention/thop_attention.py` | -| `allgather` | `comm/allgather.py` | -| `reducescatter` | `comm/reducescatter.py` | -| `fp4_quantize` | `quantization/fp4_quantize.py` | -| `nvfp4_gemm` | `gemm/nvfp4_gemm.py` | -| `flashinfer_silu_and_mul` | `activation/flashinfer_silu_and_mul.py` | -| `noaux_tc_op` | `moe/noaux_tc_op.py` | -| `fp4_block_scale_moe_runner` | `moe/fp4_block_scale_moe_runner.py` | -| `empty`, `reshape`, `split`, `concat`, `copy_`, `expand`, `transpose`, `view_dtype`, `add`, `pad` | `torch/*.py` | - -`_rope_tables` enters the closure through exactly one branch, taken on the -first forward and never again: when the engine's admitted `max_seq_len` -exceeds the config's `max_position_embeddings` — which a `speculative_config` -causes, measured **163840 -> 163848** at `max_draft_len: 3` — the constant rope -table is rebuilt at the larger row count. Its `torch.arange` / `.cos()` / -`.sin()` / `torch.empty` are host-side table construction, not step math: they -create no activation, run once before any CUDA-graph capture, and every table -row depends only on its own position, so the rows the identity config uses are -bit-identical whether the table was built at 163840 or 163848. Recorded here -rather than left for a closure scan to trip over. - -Everything else in the closure is a builtin or a metadata read (`getattr`, -`hasattr`, `isinstance`, `int`, `bool`, `all`, `len`, `max`, `range`, -`sorted`, `divmod`, `dict`, `kwargs.get`, `list.append`, -`host_kv_cache_pool_mapping.tolist()`). - -### MTP layer — the `MTPLayer.forward` closure, plus `.shared_head` - -`forward` plus `_shared_mlp`, `_routed_experts`, `_mtp_dp_rows`, -`_build_step_args`, `_moe_chunk_sizes`; and `shared_head`, which the runtime -calls separately for the draft logits. - -| call | catalog entry | -|---|---| -| `embedding` | `torch/embedding.py` | -| `flashinfer_rmsnorm` | `norm/flashinfer_rmsnorm.py` | -| `flashinfer_fused_add_rmsnorm` | `norm/flashinfer_fused_add_rmsnorm.py` | -| `cublas_mm` | `gemm/cublas_mm.py` | -| `bmm_out` | `gemm/bmm_out.py` | -| `mla_rope_append_paged_kv_assign_q` | `attention/mla_rope_append_paged_kv_assign_q.py` | -| `load_paged_kv_cache_for_mla` | `attention/load_paged_kv_cache_for_mla.py` | -| `mla_rope_generation` | `attention/mla_rope_generation.py` | -| `thop_attention` | `attention/thop_attention.py` | -| `allgather` | `comm/allgather.py` | -| `reducescatter` | `comm/reducescatter.py` | -| `flashinfer_silu_and_mul` | `activation/flashinfer_silu_and_mul.py` | -| `noaux_tc_op` | `moe/noaux_tc_op.py` | -| `fused_moe` | `moe/fused_moe.py` | -| `empty`, `reshape`, `split`, `concat`, `copy_`, `expand`, `transpose`, `add`, `pad` | `torch/*.py` | - -The two tables differ in exactly one place, and it is a dtype consequence: the -checkpoint excludes `model.layers.61*` from NVFP4 wholesale, so this module is -bf16 throughout. It therefore uses **`moe/fused_moe.py`**, the unquantized -grouped-expert runner, where the trunk uses `fp4_block_scale_moe_runner` — and -correspondingly the trunk's `fp4_quantize`, `nvfp4_gemm` and `view_dtype` do -not appear here at all. Everything else is the same vocabulary. - -`shared_head` adds `flashinfer_rmsnorm` and one `logits_processor.forward`. - -### The shell's speculative branch - -`StaircaseForCausalLM.forward` is a plain delegation to the inherited base -when there is no spec worker. With one it contains exactly one tensor -expression — the row gather of the trunk's hidden states at -`spec_metadata.gather_ids`, which is **`torch/embedding.py`** -(`torch.nn.functional.embedding` is a row lookup, used here as one) — plus a -`logits_processor.forward`. That projection is the inherited shell's own -logits path, the same one the non-speculative forward runs internally, so it -is runtime rather than modeling; it is named here rather than left implicit. - -**No catalog entry was added by this target**, torch mirror included. The -trunk's entry set is exactly the `deepseek-v3-lite-nvfp4/sm_100/dep4` -sibling's; the MTP increment consumed one further **existing** entry, -`moe/fused_moe.py`, which the catalog owner certified at this checkpoint's MTP -routed geometry rather than this target adding anything. - -Load time (outside the closed-vocabulary rule, `weights.py`): -`torch.ops.trtllm.block_scale_interleave` for the 128x4 scale swizzle, plus -`torch.cat` / `torch.index_select` for the `[up; gate]` concat and the -interleave + 32-row block shuffle of the expert stacks, and the `kv_b_proj` -row regrouping. The MTP module adds only plain `torch.cat` (its `[up; gate]` -FC1 stack and its shared-expert `[gate; up]` pair) — no swizzle, no shuffle, -because none of it is quantized. `derive_after_load` builds the YaRN rope -table, the `.t()` GEMM views, the two MLA absorption operands, every NVFP4 -call scalar, and the expert-window slices of the three MoE scale scalars — and -asserts the checkpoint's fp8 KV scales are all exactly 1.0 (122 tensors, 124 -with the MTP module loaded). - -Audit is mechanical: collect the calls in each root and the private methods it -reaches, then match against `catalog/index.yaml`. - -## Verification - -**The gate records below are the identity config's.** A `configs/mtp*.yaml` -variant is a different forward and a different weight load, so nothing here -speaks for it; its own records are in *The MTP variant* at the end of this -section. - -Both gates run under the target's identity config (`llm_args.yaml` = -`tensor_parallel_size: 4` + `moe_expert_parallel_size: 4` + -`enable_attention_dp: true`, everything else trtllm defaults — block reuse -and CUDA graphs on, page size 32), 2026-08-01, **GPUs 3-6** of -`umbriel-b200-027` (8x NVIDIA B200, driver 595.58.03, 224-core), -`source scripts/env.sh` then `CUDA_VISIBLE_DEVICES=3,4,5,6`. - -### Required on sm_103 — not yet run - -Every row needs 4 GB300 GPUs and `trtllm-llmapi-launch` over a 4-task srun -allocation. `CFG=` this target's `configs/` directory: each file there is a -complete `--extra_llm_api_options`, carrying the dep4 topology that selects -this target plus the knobs of its own variant. - -Records below that speak of "smoke" were measured with a per-target -`smoke.py` — a bespoke CLI that asserted a keyword in each of ten greedy -continuations. It was removed once this package moved in-tree: the generic -script in the `boot` row below starts the engine the same way and prints -the same continuations, and -unlike a module nothing ran, the accuracy rows below are wired into CI. - -| Gate | Command | Result | -|---|---|---| -| boot | `TRTLLM_STAIRCASE=require trtllm-llmapi-launch python examples/llm-api/quickstart_advanced.py --model_dir --tp_size 4 --moe_ep_size 4 --enable_attention_dp --max_tokens 16 --prompt "The capital of France is" "The chemical symbol for gold is" "1, 2, 3, 4, 5,"` | **passed, 10/10** greedy keyword asserts, 2026-09-10, 4x GB300 on nvl72d199-T07, trtllm 1.3.0rc26, 5m18s. Run through `trtllm-llmapi-launch` over a 4-task srun allocation; the engine built `tensor_parallel_size=4`, `moe_expert_parallel_size=4`, `enable_attention_dp=True` with `TRTLLM_STAIRCASE=require` exported | - -**The identity assembly and the MTP variant are both gated on sm_103.** Everything up to the -weight load is driven by the checkpoint's config alone, and that part -was exercised on a GB300 against the config *as published* (no -target-owned stub): routing resolved this target, the module imported, -`StaircaseCore.__init__` passed every geometry, topology and dtype -assert, and 1980 parameters declared, 61 layers, hidden 7168. `lm_head.weight` -came out **bfloat16**, which is the specific thing removing the stub -put at risk -- the shell sizes it from the pretrained dtype, and a -regression there materializes fp32 two layers from its cause. - -The weight path is verified by the two gates above, over four ranks: -the manifest load fills every declared parameter, the post-load -derivations run, and the forward is exercised through prefill, the -CUDA-graph decode path, and the expert-parallel MoE round trip. - -The checkpoint they ran on is the one this file records. Every digest -that can be checked was re-verified after download -- `config.json`, -`generation_config.json`, `hf_quant_config.json`, `tokenizer_config.json` -and `tokenizer.json` by hand, the 163 safetensors shards by git-lfs -whose object id *is* the sha256 -- so these records and the sm_100 ones -below were measured on byte-identical weights. - -`configs/mtp3.yaml` carries its own three gate records above -- it selects a -second forward path *and* a second weight-loading path, so the identity -records do not speak for it. `mtp1.yaml` and `mtp2.yaml` remain **ungated**; -they are the dominated low end of the measured draft-length axis, and -`mtp3.yaml` is the one to serve. - -| boot, MTP | as above plus `--spec_decode_algo MTP --spec_decode_max_draft_len 3 --kv_cache_fraction 0.75`, i.e. what `$CFG/mtp3.yaml` declares | **passed, 10/10**, 2026-09-10, 4x GB300, 6m16s. `MTPDecodingConfig(max_draft_len=3)` in the run's LLM Args and 124.65 GiB of weights loaded against the identity path's ~118.8 GiB -- the layer-61 bf16 MTP module, i.e. the variant's second weight-loading path, really ran | -| gsm8k full | `TRTLLM_STAIRCASE=require trtllm-llmapi-launch trtllm-eval --model --extra_llm_api_options $CFG/identity.yaml gsm8k --output_path ` | **passed, 95.0720** (`exact_match,flexible-extract`, +-0.5962, full 1319 questions) against threshold **89.9962** (anchor `deepseek-ai/DeepSeek-R1-0528` = 94.9962, tol 5.0) -- pass by 5.08 points. 2026-09-10, 4x GB300, 6m47s. `strict-match` on the same run: 94.7688 | -| gsm8k, **paired in one session** | `$CFG/identity.yaml` and `$CFG/mtp3.yaml`, back to back in one allocation; in CI as `accuracy/test_staircase.py::TestStaircaseDeepseekR10528Nvfp4Sm103Dep4::test_gsm8k_identity_vs_mtp3` | **passed.** identity **94.9962** (+-0.6005) / strict 94.7688; mtp3 **94.6171** (+-0.6216) / strict 94.4655. **delta = -0.3791 flexible, -0.3033 strict** against a `|delta| < 1.2` criterion (2 sigma at sigma = 0.60) -- 0.63 sigma, five questions of 1319, and both filters move the same way | -| acceptance vs stock | `trtllm-bench throughput` twice in one allocation over one fixed-seed dataset, `TRTLLM_STAIRCASE=require` against `=off`; in CI as the two `accuracy/test_staircase.py::TestStaircaseDeepseekR10528Nvfp4Sm103Dep4::test_mtp3_acceptance[...]` legs | **passed.** `acceptance_length` **3.3514** (ours) against **3.2752** (stock) -- ratio **1.023**; draft acceptance 78.38% against 75.84%. 2026-09-10, 4x GB300, 64 requests at ISL=OSL=1024, concurrency 32 | - -`TRTLLM_STAIRCASE=require` is what makes these gates at all: under `auto` a -configuration that missed this target would measure the built-in DeepseekV3 -implementation and report it as this target's score. The old `--trtllm` flag -of the perf harness, which worked by unsetting an environment variable, is -now just `staircase: off` — one config key selecting between the two -systems. - - - -#### The MTP variant is gated, and why the third gate was the one that mattered - -Rejection sampling holds the emitted distribution to the target model's, so a -*miscomputed* draft layer produces correct text more slowly rather than wrong -text: the boot gate passes, accuracy passes, and only speed moves. `acceptance_length` -is the sole detector, and it is only readable against a reference -- 1.0 would -mean every draft rejected, and "clearly above 1.0 but below the reference" -would mean subtly wrong. - -Measured **3.3514** against stock's **3.2752** on the same 64-request -fixed-seed dataset in the same allocation, a ratio of **1.023**, against a -ceiling of 4.0 at `max_draft_len: 3`. The sm_100 record below measured ratios -of 1.0215 and 1.0175 at concurrency 1 and 32, so this sits in the same family. -The draft path computes what the checkpoint says. - -Two things about the comparison worth stating rather than leaving implicit. -`TRTLLM_CAN_USE_DEEP_EP=0` was exported for **both** sides: stock cannot boot -this checkpoint at dep4 with MTP without it (the MoE communication factory -lands on DeepEPLowLatency, whose dispatch takes only NVFP4 uint8 hidden states, -and the MTP layer is bf16), and it is inert for the staircase target, which -implements the all-gather/reduce-scatter round trip by hand rather than -through that factory. So it makes the two systems more comparable, not less. -And the throughput on the same runs -- 2955.2 against 2732.7 output tok/s -- is -**recorded, not claimed**: it is one concurrency of one synthetic workload, and -this project measures perf rather than gating it. - -#### On the flexible-extract number landing exactly on the sm_100 one - -This run measured **95.0720 (+-0.5962)**, and the sm_100 identity run -recorded below measured **95.0720 (+-0.5962)** -- the same 1254 of 1319 -questions. That is agreement at the score level, not proof of identical -generations: `strict-match` differs on the same pair of runs (94.7688 here -against 94.9204 there), so the underlying text does move, as it should -across two architectures whose MoE epilogues use different block-scale -recipes. Read the flexible-extract match as a strong reproduction, and not -as evidence that the two forwards are bit-identical -- they are not. - -### Prior record — sm_100 (B200), pre-move harness, does not gate this target - -| Gate | Result | -|---|---| -| smoke — `uv run targets/deepseek-r1-0528-nvfp4/sm_100/dep4/smoke.py` | **passed, 10/10** greedy keyword asserts, first run, no iteration. Keywords were authored provisionally and frozen against the continuations this model actually produced | -| gsm8k full — `uv run bench/accuracy.py --target targets/deepseek-r1-0528-nvfp4/sm_100/dep4` | **passed, 94.9962** (`exact_match,flexible-extract`, 1319 questions, 5-shot completion, no chat template, 256 output tokens) against `reference: deepseek-ai/DeepSeek-R1-0528 = 94.24 (trtllm); gate: accuracy >= 89.24`. First run, stock defaults, no instrumentation; TRTLLM execution 37.695 s, engine init 197.663 s | - -Both lm-eval filters of the passing run: `flexible-extract` **94.9962** -(±0.6005) and `strict-match` **94.6171** (±0.6216). The gap is 5 questions of -1319 — under few-shot completion this checkpoint mostly answers in the strict -`#### N` form, and flexible extraction finds a few more. - -**Read the +0.76 over the anchor as agreement, not as an improvement.** The -anchor was measured on a **different NVFP4 export of the same base model** -(`DeepSeek-R1-0528-FP4-v2`, whose modelopt `exclude_modules` list differs -substantially from this v1 export's), which is what `tol: 5.0` absorbs; and -one filter's stderr alone is ±0.60. Nothing at this distance is a result in -either direction. - -The reference entry is still at its external anchor (`source: trtllm`). Its -header asks for a write-back after a target's first passing run; `bench/` is -read-only for the assembler, so that edit is left to whoever owns the file. - -### The MTP variant - -`configs/mtp{1,2,3}.yaml`, 2026-08-02, same host and same **GPUs 3-6**. - -| Gate | Result | -|---|---| -| identity smoke, re-run after the MTP increment — `uv run .../smoke.py` | **passed, 10/10**, and all 10 continuations **byte-identical** to the pre-increment run (`md5` of the captured case lines equal). This is the invariant the increment is held to: no `speculative_config`, no MTP parameter declared, `forward` a plain `super().forward(...)`, so the gsm8k 94.9962 record stands unmoved | -| MTP smoke — `uv run .../smoke.py --config .../configs/mtp3.yaml` | **passed, 10/10**. Engine built at `max_draft_len: 3`, `max_seq_len` 163848, KV pool 32.20 GiB (968,288 tokens, 62 layers), 12 fp8 MLA decode JIT compiles per rank | -| acceptance probe — `uv run bench/perf.py --target ... --label probe-mtp3-acceptance --config .../configs/mtp3.yaml --acceptance --concurrency 1,32 --rounds 2` | **`acceptance_length` 3.9274 at con=1 and 3.4635 at con=32**, against a ceiling of 4.0 at `max_draft_len: 3` and a break-even of 1.164. Draft acceptance rate 97.58% / 82.12% (5849 of 5994 and 39152 of 47679 draft tokens accepted) | - -**Two of the ten MTP continuations differ in wording from the identity run's, -and that is expected rather than a defect.** Rejection sampling makes the -emitted token the *target* model's argmax, so a correct MTP layer cannot -change what is emitted — but the trunk's own generation call now runs four -query rows per sequence instead of one, on a different decode kernel -(`HVPerCta256` rather than `HVPerCta128`) with a different accumulation order. -Greedy decoding is chaotic under a 1-ulp logit difference at a near-tie: cases -2 and 3 diverge a few tokens in and stay grammatical and correct, and all ten -keywords pass. **Byte-identity is the right bar for the identity config and -the wrong one for a variant that changes the attention tile.** - -### The MTP variant's own gate records - -The two gates the increment is released on were run after the assembly, on -the same host and the same **GPUs 3-6**, 2026-08-02. Like everything else in -this file they are **sm_100 records and do not gate the sm_103 target**; the -acceptance row in particular has to be repeated, because it is the *only* -detector of a miscomputed draft path — rejection sampling keeps the emitted -distribution correct, so boot and accuracy both pass while a wrong draft -layer merely runs slower. - -| Gate | Result | -|---|---| -| gsm8k, **paired in one session** — `uv run bench/accuracy.py --target ...` with and without `--config .../configs/mtp3.yaml` | **passed.** identity **95.0720** (±0.5962) / strict 94.9204; mtp3 **95.2237** (±0.5874) / strict 94.8446. **Δ = +0.1517 flexible, −0.0758 strict** against a `\|Δ\| < 1.2` criterion (2σ at σ = 0.6005) — 0.25σ, two questions of 1319, and the two filters move in opposite directions | -| acceptance vs stock trtllm — `uv run bench/perf.py --target ... --label trtllm-mtp3-acceptance --trtllm --config .../configs/trtllm-ref-mtp3.yaml --acceptance --concurrency 1,32 --rounds 2`, with `TRTLLM_CAN_USE_DEEP_EP=0` | **passed.** `acceptance_length` **3.9274 / 3.4635** (ours) against **3.8447 / 3.4040** (stock) at con 1 / 32 — ratios 1.0215 and 1.0175. Draft acceptance 97.58% / 82.12% against 94.82% / 80.13% | - -**The paired accuracy run is the one to quote, not a cross-session -comparison.** The same identity code path measured 94.7688 on 2026-08-01 and -95.0720 on 2026-08-02 — a 0.30 spread on a bit-identical forward, which is -the scale any accuracy claim about this variant has to be read at. Both runs -of the pair were back to back on the same devices. - -**One incidental number worth keeping**, because it is the only *real-text* -measurement of what MTP buys here: gsm8k execution time fell from **37.025 s -to 31.440 s, −15.1% (1.178x)**, on 1319 genuine prompts. It is not a Pareto -point — lm-eval drives its own concurrency rather than a swept one — so it -does not belong on the perf curve, but it is a far better answer to "is MTP -worth enabling" than the flattered random-prompt throughput. Engine init -rose 193.6 s → 278.4 s (layer 61's weights plus the extra decode JIT). - -**Read the ~2% acceptance lead as agreement, not as an improvement.** The two -systems do not run the same knobs: the reference is forced to `max_num_tokens` -/ `max_seq_len` 2048 to boot at all and to a different MoE transport (below), -so scheduling, batching and block reuse differ and each system verifies a -slightly different token stream. What the comparison establishes is that this -target's MTP layer computes what the checkpoint says — a miscomputed one sits -at 1.0, a subtly miscomputed one clearly below the reference. - -**Stock trtllm cannot serve this checkpoint at `dep4` with MTP under its -default MoE communication strategy**, and that is why the reference carries an -environment variable as well as a config. All four ranks die in -`Failed to initialize executor` on -`deep_ep_low_latency.py:238`'s `assert hidden_states.dtype == torch.uint8`: -the communication factory falls through `NVLinkOneSided` / `NVLinkTwoSided` / -`DeepEP` (each `not available: Invalid Argument`) to `DeepEPLowLatency`, whose -dispatch accepts only NVFP4 hidden states — and **the MTP layer is bf16**, -by the export's own `exclude_modules`. The existing `trtllm` Pareto curve went -through DeepEPLowLatency happily because without MTP every MoE layer is NVFP4. -`TRTLLM_CAN_USE_DEEP_EP=0` lands the reference on `AllGatherReduceScatter`, -which is the strategy *this target implements by hand*, so it makes the two -systems more numerically comparable rather than less. -`configs/trtllm-ref-mtp3.yaml` carries the whole ledger. - -Consequently `trtllm-mtp3-acceptance` is an **acceptance measurement only**. -Its throughput column (180.6 tok/s/user at con=1, 1668.4 tok/s at con=32) is -not comparable to the `trtllm` Pareto curve: different boot knobs, a different -transport, and a draft length that curve does not carry. - -### Why the accuracy gate alone could not have released this - -Rejection sampling means a miscomputed draft layer produces correct text more -slowly, so an accuracy score cannot separate a good MTP layer from a broken -one — the Δ above would have been just as small with the `eh_proj` halves -swapped. `acceptance_length` against a reference is the only measurement that -can, which is why both gates are listed and why neither is optional. - -**What the acceptance numbers do and do not establish.** They establish that -the MTP layer computes something the target model agrees with: 97.6% of draft -tokens accepted at con=1 is a ceiling-adjacent 3.9274 of a possible 4.0, and a -layer with the `eh_proj` halves swapped, a norm on the wrong operand or the -expert stack packed wrong would sit near 1.0. On their own they did **not** -establish that this is the *best* achievable draft quality — that needs the -same checkpoint under stock in-tree modeling at the same load and the same -`speculative_config`, which the section below measures. - -**And read the workload the other way round from the usual warning.** The -harness generates uniformly random prompt token ids with `ignore_eos`, and -`docs/models/multi-token-prediction.md` warns that such a workload gives -acceptance *far below* real text. That warning does not hold here, and the -reason is that only the **prompt** is random: the 1024 output tokens are the -model's own continuation of nonsense, which is highly repetitive, and -repetition is exactly what an MTP layer drafts perfectly. So this probe -**flatters** MTP rather than penalizing it. It is a correctness instrument -here and nothing more; the value of enabling MTP has to be judged on a -real-text workload. - -Two numbers from the probe that are *not* results and must not be quoted as -Pareto points: 211.17 tok/s/user at con=1 and 1942.9 tok/s at con=32. They are -measured under a different config (`free_gpu_memory_fraction: 0.75`, and -`--acceptance` turns on `enable_iter_perf_stats`, which the harness's own -header says makes a label non-comparable to one without it), in a two-round -probe rather than a sweep, on the flattering workload above. The perf campaign -for the variant belongs to the tuner. - -### The parallel split, and what it rests on - -| part | split | per-rank shape | -|---|---|---| -| `q_a_proj`, `q_a_layernorm`, `q_b_proj`, `kv_a_proj_with_mqa`, `kv_a_layernorm`, `kv_b_proj`, `o_proj` | **replicated**, all 128 query heads | `[1536, 7168]`, `[1536]`, `[24576, 1536]`, `[576, 7168]`, `[512]`, `[32768, 512]`, `[7168, 16384]` | -| layers 0-2 dense MLP | **replicated**, intermediate 18432 | `gate_up [36864, 3584]`, `down [7168, 9216]` | -| shared expert | **replicated**, intermediate 2048 | `gate_up [4096, 3584]`, `down [7168, 1024]` | -| routed experts | EP window, 64 of 256 at offset `64 * moe_ep_rank` | `fc1 [64, 4096, 3584]`, `fc2 [64, 7168, 1024]` | -| per-expert NVFP4 scalars | **replicated over all 256** (deliberate) | `[256]` each | -| fp8 KV scales, router, both norms, embedding | **replicated** | unchanged | -| `lm_head` (shell-owned) | **replicated** — the shell builds the whole matrix under attention DP | `[129280, 7168]` | - -Two collectives per **MoE** layer, none in layers 0-2 (dense) and none on the -attention or residual path — 116 per forward: - -* `comm/allgather` on the post-attention normed hidden states, **before** the - router GEMM; -* `comm/reducescatter` on the expert window's output over the whole gathered - token set, handing each rank back exactly its own rows. - -Three things the correctness rests on, none of them checkable inside one rank: - -* **the routing agrees across ranks, and the gather placement is what makes - it so.** The four 64-wide windows must tile the routing space exactly once; - under attention DP the ranks hold different tokens, so that is restored by - gathering **before** the router GEMM — the router and `noaux_tc_op` then run - on byte-identical full token sets on all four ranks and each token's top-8 - ids agree. Routing locally and gathering afterwards would break it with no - error (`docs/models/expert-weight-packing.md`); -* **one activation quantization feeds every window.** The routed FC1 global - scale is `1 / input_scale` of the *shared* expert, which the checkpoint sets - to the max over all 256 routed experts — `derive_after_load` asserts that - over the full 256, which is why the per-expert scalars are loaded whole on - every rank; -* **the reduce-scatter is crossed in bf16.** The op sums, and it sums - `float8_e4m3fn` as raw bytes rather than as floats; the gather is a byte - move and would survive fp8, so the asymmetry is a live trap on the return - leg only. - -`attn_metadata.all_rank_num_tokens` is where the engine publishes the per-rank -row counts; `_dp_rows` reads it, asserts `all_rank_num_tokens[rank] == -hidden_states.shape[0]`, and returns `max(counts)` — every rank pads its token -block to the group-wide maximum, so both collectives run in their uniform form -and the ragged form is never used. Padded rows are sliced off after the -reduce-scatter and never reach the residual stream. - -### What the fp8 latent pool moves, and why each choice is what it is - -The KV cache is fp8-e4m3 per the checkpoint, and **nothing validates the fp8 -round trip at any layer**: the write scale, the read scale and the two folded -FMHA scales are independent roles with no relation checked anywhere in the -chain. Getting one wrong is silently mis-scaled output, never an error. - -* **`quant_mode` is derived from the checkpoint's quant config**, not - hard-coded: `kv_cache_quant_algo: FP8` selects `QuantMode`'s fp8-KV bit - (128), the value every MLA entry certifies. The engine's own - `quant_config.quant_mode` is a `QuantModeWrapper` rather than an int at this - pin, and bits outside the KV-cache group ride along unread anyway (`1152` - and `384` measured bit-identical to a bare `128` on every MLA flavor), so the - target maps the declared algo onto the bit itself. -* **`s = 1.0`, and it is checked rather than assumed.** All 122 per-layer - `k_scale`/`v_scale` tensors are loaded (they are 0-dim fp32 scalars) and - `derive_after_load` asserts each is exactly 1.0. Both scale arguments are - then passed as `None`, which every op reads as exactly 1.0 and which is what - the engine's own call sites pass. This is not a convenience: the fp8 MLA - *context* path is internally inconsistent at `s != 1.0` in **both** context - flavors (it quantizes q/k/v at 1.0 while applying `s^2`/`s` as if it had - not), so 1.0 is the only correct value and the assert is the defence. -* **The decode producers are ordered, not concurrent.** - `mla_rope_generation` does not write `fused_q` over an fp8 pool — it *reads* - `fused_q[..., :C]` to build `quant_q_buffer`, which is the query the decode - FMHA consumes. The absorbed-q BMM is issued before it on the ambient stream, - which is what makes that safe; the bf16 reading (disjoint halves, free to - overlap) is a silent race here. -* **The two phases divide the quantization labour oppositely.** Context: - `mla_rope_append_paged_kv_assign_q` leaves `q` plain bf16 and - `thop_attention`'s context call quantizes q/k/v to e4m3 itself, so `q` is - quantized exactly once. Decode: `mla_rope_generation` produces the quantized - query and the two folded FMHA scales, and the generation call reads - `quant_q_buffer`, `mla_bmm1_scale[1]` and `mla_bmm2_scale[0]` — **ignoring - both kv scale tensors entirely**. Nothing in either signature says so. -* **A cached prefix pays fp8 twice.** The engine dequantizes cached latent - rows off the pool, `kv_b_proj` up-projects them, and `thop_attention` - quantizes the result straight back to e4m3. Not incorrect, but a - cached-prefix context call carries strictly more quantization error than a - fresh prefill of the same tokens — do not attribute an accuracy gap to block - reuse without accounting for it. -* **fp8 raises the attention workspace requirement.** Measured here: - **2,126,512,128 B (1.98 GiB) per rank**, resized on the first call, against - 531,628,032 B for the 32-head bf16 MLA shape at the same - `max_num_tokens=8192`. Sizing from bf16 MLA figures under-budgets. - -### The rope table is the whole of the rope configuration - -`thop_attention` reads the table's **content** and `q_scaling`; the seven -scalar rope arguments and `rotary_inv_freq` beside them are measured inert on -the MLA path. So the YaRN blend lives entirely in the table this target -builds, and the model's YaRN attention temperature lives entirely in -`q_scaling`: - -* table `inv_freq(d) = ramp(d)/(factor*freq(d)) + (1-ramp(d))/freq(d)` with - `low = 10`, `high = 23` at this config (`R` 64, `theta` 10000, `factor` 40, - `original_max_position_embeddings` 4096, `beta_fast` 32, `beta_slow` 1) — - cross-checked against the HF reference's own YaRN rope init to **1.05e-8 - max abs / 1.27e-7 max relative** on `inv_freq` (the HF value is fp32, so - that is its rounding); -* table amplitude `m(mscale)/m(mscale_all_dim)` = **exactly 1.0** (both are - 1.0 here), matching the same reference's `attention_factor`; -* `q_scaling = 1/m(mscale_all_dim)^2` = **0.5336594470450011**, which the - construction asserts lands on `thop_attention`'s certified fp8-pool value - (1.0 and 0.53366 are the two certified). - -**The table is built for the full `max_position_embeddings` = 163840 -positions** (84 MB fp32 per rank). A short table is read out of bounds with no -check — `rope_max_positions`, the argument that looks like it bounds this, is -one of the inert seven — so the first forward additionally asserts -`max_seq_len <= max_pos`. - -### Envelope - -* Context sequences with a cached prefix are served; chunked prefill is not - implemented (`AttentionRuntimeFeatures.chunked_prefill=False` under - defaults, and the target holds `chunked_prefill_buffer_batch_size=1`). -* **MTP is a `configs/` variant, never the identity.** The checkpoint declares - `num_nextn_predict_layers: 1` and ships the whole of layer 61 for it — its - own 256 experts, `embed_tokens`, `eh_proj`, two extra norms and a - `shared_head.head`, 790 keys in all. - - Under `llm_args.yaml` alone those 790 keys are a **predicted non-load** in - the weight manifest (read off the checkpoint by layer index rather than - enumerated), no MTP parameter is declared, and the shell's `forward` is a - plain delegation to the inherited base — nothing about this identity changed - when MTP was added, which is a measured claim and not a design intention: - after the increment, smoke's 10 greedy continuations are **byte-identical** - to the pre-increment run's. - - Under `configs/mtp{1,2,3}.yaml` the module is loaded and a second forward - path runs. **Those variants cross the line a config variant normally - respects, and that is stated rather than left to be inferred**: a variant is - supposed to move a knob, and these change the **weight-loading path** — layer - 61 goes from non-load to loaded, 212 keys and ~6.3 GB per rank — add a - second forward (`MTPLayer`, replayed `max_draft_len` times per step), and - change what the runtime allocates (a 62-layer KV pool, `max_draft_len - 1` - extra tokens per sequence, and a `max_seq_len` the engine raises by 8 at - `max_draft_len: 3`). Only the gates run against a variant can speak for it; - the identity gate records below are not evidence about it, and vice versa. -* fp8-e4m3 latent pool only (`quant_mode` = the fp8-KV bit, KV scaling factor - 1.0); beam width 1; no LoRA, no cross attention, no FlashMLA layout — each - asserted at the first forward or in the forward itself. -* `tokens_per_block == 32` is asserted: every MLA entry's **fp8** column is - certified at page 32 only (their bf16 columns also carry 64). 32 is what a - default `KvCacheConfig` produces, so this binds only if someone tunes the - page size — which would need the certification extended first. -* Pipeline parallelism and a second (tensor) split of the routed experts are - asserted off: `pp_size == 1`, `moe_tp_size == 1`. Attention DP is asserted - **on** — this target is the DP assembly. -* A rank with **zero** tokens was never observed (idle ranks get a 1-token - dummy) and is not something this forward was exercised on. - -Certification coverage this target consumes, and where it sits relative to -what the entries measured: - -* **`thop_attention`** runs MLA at `H = 128` over an **fp8-e4m3 latent pool**, - page 32, `q_lora_rank = 1536`, at the complete DeepSeek-R1-0528 rope/scale - cell — the exact configuration the entry certifies for all three MLA call - flavors and their mixed-batch pairing. `H = 128` is the only head count - certified over an fp8 pool, and page 32 the only page size. -* **`mla_rope_generation`**, **`mla_rope_append_paged_kv_assign_q`** and - **`load_paged_kv_cache_for_mla`** run their fp8 columns at the same cell - (`H = 128`, page 32, `C/R/nope/v = 512/64/128/128`, scale omitted). -* **`predicted_tokens_per_seq` (`P`) is 1 on every context call and the - generation call's own query-tokens-per-sequence.** Without a - `speculative_config` that is 1 everywhere; under `configs/mtp{1,2,3}.yaml` the - MLA **generation** calls run at `P = max_draft_len + 1` (2, 3, 4) on the - trunk and on the MTP layer's draft step 0, and at `P = 1` on draft steps 1+. - `thop_attention` and `mla_rope_generation` certify `P` at **1, 2, 3 and 4** - over this exact fp8 cell and no further, which is why `mtp3` is the largest - variant here: `max_draft_len: 4` would need the certification extended - first, not just a bigger number in the yaml. Two preconditions ride with it, - both held structurally rather than checked (a device read would cost a sync - and is illegal under capture): - * **`L_g >= P` for every generation sequence** — draft row 0 attends to - `[0, L_g - P]`, so a shorter KV length leaves it no keys. `L_g` counts all - `P` of the step's tokens, so it cannot be smaller than `P`. - * the `G*P` cache rows one `mla_rope_generation` call writes must address - **distinct physical slots**, which a `KVCacheManager`-allocated batch gives - automatically. -* **`fused_moe`** (the MTP layer's routed experts, bf16) runs at - `E = 64`, `(H, I) = (7168, 2048)`, `K = 8`, `ep_size = 4`, - `ep_rank ∈ {0,1,2,3}` — the cell the entry certifies, including all four - windows against a 256-expert reference — with `T <= 2048`. Its - **distinct-expert-ids precondition** matters here and is worth stating - rather than leaving to inference: above 256 tokens a repeated id in one - token's row reads out of bounds in `finalizeMoeRoutingKernel` (an illegal - access, or ~200-460 ulp of silent garbage), and this call drives `T` to 2048. - It holds **structurally**: the ids come from `noaux_tc_op`, whose semantics - are the indices of the top-k largest corrected scores, and a top-k over - expert indices cannot select one twice. - - The chunk bound is 2048 rather than the trunk's 8192 for a memory reason, - not a certification one — see *What MTP costs in memory* below. -* **`fp4_block_scale_moe_runner`** runs at the R1 routed geometry - (`H = 7168`, `I = 2048`, `num_experts = 256`, `top_k = 8`) with - `local_num_experts = 64`, `local_expert_offset ∈ {0, 64, 128, 192}` — the - certified four-way split — and at **`T <= 8192`**, the top of the entry's - certified token column, which the whole column covers at this geometry. The - gathered token set reaches `4 * max_num_tokens = 32768`, so the expert call - is **chunked** (`_MOE_MAX_T` in `modeling.py`). **That bound is a - certification boundary, not a tuning knob.** Routing and the activation - quantization are chunked with it, which keeps `noaux_tc_op` inside its own - enumerated column (up to `num_tokens` 8192) as well. -* **`noaux_tc_op`** runs the grouped configuration - `(num_experts, n_group, topk_group, topk) = (256, 8, 4, 8)`, which the entry - enumerates explicitly. -* **`comm/allgather` and `comm/reducescatter`** are driven at world size 4, - group `[0,1,2,3]`, bf16, hidden 7168, in their uniform form only, at 58 - sites per step — 58 + `max_draft_len` under an MTP variant, the MTP layer - being one more MoE layer per draft step. Both contracts certify the - call-order surface: calls pair by **position** on the communicator, so - issuing the same sequence on every rank is this forward's obligation — - discharged structurally, since every rank runs the same layer loop and pads - to `max(all_rank_num_tokens)`. - - **Inside the draft loop the padding basis is a different list, and reading - the wrong one is silent.** The worker leaves `attn_metadata.all_rank_num_tokens` - holding the trunk's counts for the whole loop and passes the correct basis in - as the `all_rank_num_tokens` **keyword** — - `spec_metadata.all_rank_num_tokens` at draft step 0, then - `spec_metadata.subseq_all_rank_num_tokens`, which is the per-rank *sequence* - count. Both contracts certify that at equal byte counts a mispairing does not - hang: every rank comes back wrong in 98-99% of elements, bitwise - reproducibly. The defence is structural rather than vigilant — the MTP - layer's padding helper (`_mtp_dp_rows`) takes the list as a parameter and has - **no metadata argument at all**, so it cannot reach the wrong one; the - trunk's `_dp_rows` keeps its own shape and the two are not interchangeable. - - **Both collectives stay on the ambient stream**, which since the - `iter1-shared-side-stream` iteration is not the only stream the forward - uses: the shared-expert branch is forked onto a side stream and joined - before the add. That is inside what both contracts certify — "the side - stream joined to the current one on both ends, which is what the engine and - a target's forward both do" — and it changes neither collective's arguments, - its call order, nor its uniform form. -* **`nvfp4_gemm` and `fp4_quantize` are certified at this target's shapes**, - on receipts rather than on a domain rule. They were assembled as "inside the - stated domain, outside the enumerated list" and closed afterwards, because - R1's widths *exceed* every previously-run value rather than falling between - them. `nvfp4_gemm` at `(K, N)` = (7168, 36864), (18432, 7168), (7168, 4096), - (2048, 7168) with `M` to 8192; `fp4_quantize` at `K` = 7168 / 18432 / 2048 - through `T = 8192`, in **both** scale layouts (swizzled for the GEMM, linear - for the MoE runner — `K = 7168` is taken in both). - - Both runs came back "the rule was sufficient after all", and both found the - contract wrong about *why*. `nvfp4_gemm`: the tactic space does move with - shape, but every R1 count lands inside the range the smaller shapes produce - and all four backends agree bitwise. `fp4_quantize`: the kernel-selection - axis its contract described (a TMA variant above 1024 rows) **does not exist - in this build** — that text came from flashinfer's vendored source, which is - a newer revision than the installed binary; the profiler sees one kernel at - every shape. - -### What MTP costs in memory, and why the variants carry a second knob - -The MTP module adds **5.87 GiB (6.30 GB) of declared weights per rank** — 5.637 -GB of it the bf16 routed expert stacks alone, which cost **3.56x what one NVFP4 -trunk MoE layer costs** purely from the dtype. `nvidia-smi` read **122,076 MiB -(119.2 GiB) per rank** after model init with MTP on, against the 117,924 MiB -(115.2 GiB) recorded above for the identity config; the two readings were taken -at different points of the load, so treat them as two figures rather than as a -clean difference. The engine then adds one KV-pool layer (62 instead of 61, -+1.6%) and `max_draft_len - 1` extra tokens per sequence, neither of which the -target declares. - -**That is not what makes it tight.** The engine sizes the KV pool from free -memory *after* the weights, at `free_gpu_memory_fraction`, and the drafting -forward's transient demand after that point is larger than the identity -config's. At the default 0.9 the pool takes 38.95 GiB (1,171,200 tokens per -rank) and boot then reaches CUDA-graph capture at **182.4 of 183.4 GiB and -livelocks** — three of four ranks spinning in `cudaFree` inside the CUDA -caching allocator's `release_cached_blocks`, reached from an ordinary -`empty_cuda` in the *trunk's* MoE runner, while the fourth waits at an -`MPI_Barrier`. There is no OOM exception and no error: the run simply stops -making progress and has to be killed. **Read the signature — one rank idle at a -barrier, the rest at 100% GPU with a frozen log — as memory exhaustion, not as -a mispaired collective.** - -**Freeing memory before the pool is sized does not help**, and that is worth -recording because it is the obvious first move: chunking the MTP layer's -expert call to 2048 rows (from the trunk's 8192, both inside `fused_moe`'s -certified column) moved pool sizing by **0.31 GiB**, 38.64 -> 38.95, and the -boot failed identically — the pool grows into exactly whatever is freed. The -chunk was reverted to the trunk's `_MOE_MAX_T`; it is not the lever. - -So `configs/mtp{1,2,3}.yaml` each carry -`kv_cache_config.free_gpu_memory_fraction: 0.75` as a **boot requirement**, -documented in the files themselves, in the same spirit as -`configs/trtllm-ref-boot.yaml`. It hands the pool ~32.5 GiB and leaves ~28 GiB -of headroom, and it constrains nothing: ~975k KV tokens per rank is an order -of magnitude past what any concurrency measured on this target uses. - -### Engine-side facts observed on this checkpoint - -* **The latent pool is fp8 and the engine sized it that way.** 44.77 GiB for - 1,368,032 tokens per rank = **35,136 B/token** = - `num_layers * (kv_lora_rank + qk_rope_head_dim) * 1` = `61 * 576 * 1`. One - byte per element, i.e. half the bf16 width — no target-side declaration was - needed beyond keeping the checkpoint's `hf_quant_config.json` linked. -* **The engine allocates the KV pool twice**: a small profiling pool (5.50 - GiB, 167,936 tokens) is created and released before the real one. -* **Per-rank weights are ~115.2 GiB** (`nvidia-smi` read 117,924 MiB on three - of the four devices after model init and before the pool allocation; the - fourth carried an unrelated 810 MiB from another user's process), against a - 183 GiB card. `Model init total` is 51-55 s per rank with the checkpoint - warm in the host page cache — the 163 shards load in ~10 s per rank there. -* **CUDA-graph capture is the default decode grid**: `batch_sizes = [1..32, - 64, 128]`, `max_batch_size 128`, `enable_padding False` — 34 sizes, all - inside the enumerated CUDA-graph coverage both collective entries carry - (`1..32, 64, 128, 256`). -* **First boot JIT-compiles 8 fp8 MLA decode kernels per rank** at ~6.0-7.0 s - each — `fmhaSm100aKernel_QkvE4m3OBfloat16HQk576HV512...P32VarSeqQ16Kv128StaticSwapsAbForGen` - and its `MultiCtasKvCga` / `HVPerCta256` siblings. `QkvE4m3` in the name is - the fp8 pool: q, K and V are all e4m3 in the decode MMAs. No MLA context - call triggers a compile. - - **Each `max_draft_len` costs its own decode compiles**, because - `predicted_tokens_per_seq` becomes the kernel's `maxSeqLenQ` and moves the - `HVPerCta` split in the name — `P` 1 and 2 take `HVPerCta128`, `P` 3 and 4 - take `HVPerCta256`, and `P` 1 and 2 pay separate compiles despite sharing a - name. Measured on `configs/mtp3.yaml`: **12 compiles per rank** against the - identity config's 8, i.e. **+4 at ~5.4-5.7 s each**, all four ranks - compiling in parallel. A campaign sweeping `max_draft_len` should budget one - first-boot compile set per value, not one for the model. -* `max_seq_len` is the config's 163840 under the identity config, so - `max_blocks_per_seq` is 5120 at page 32. **A `speculative_config` raises it**: - 163848 at `max_draft_len: 3` (5121 blocks per sequence), which is more than - the `max_draft_len - 1` extra KV tokens per sequence the runtime reference - documents. The rope table is sized from the engine's own number rather than - from a formula, on the first forward. - -## Performance - -![Serving Pareto](perf/figures/pareto.png) - -Environment: `umbriel-b200-027`, 224-core, 8x NVIDIA B200 (178.34 GiB each), -**GPUs 3-6** — the campaign device set, held for every label — driver -**595.58.03**, `tensorrt_llm 1.3.0rc21`, `torch 2.11.0+cu130`. -`source scripts/env.sh` then `CUDA_VISIBLE_DEVICES=3,4,5,6`. ISL=OSL=1024, -concurrency 1..256. **Two sessions on the same host and the same device set**: -the first three curves 2026-08-01 10:45-15:18 UTC, the three MTP curves plus -the two new stock references 2026-08-02 13:09-17:35 UTC. The second session -opened by re-measuring `iter1` (`probe-anchor-iter1`), which reproduced it to -**−0.14% at con=1 and +0.15% at con=256**, and the `trtllm` reference was -spot-checked at the end of it to **−0.01% / +0.20%** — that is what licenses -one figure across the two days. - -`con=1 tok/s/user` is `1000 / mean_tpot_ms` at con=1; `peak tok/s/GPU` is -`max(output_throughput) / 4`. - -| label | config | commit | accuracy | con=1 tok/s/user | peak tok/s/GPU | change | -|---|---|---|---|---|---|---| -| `baseline` | `llm_args.yaml` (identity) | `70b1ecd` | gsm8k 94.9962 | 76.09 | 1906.86 (con=256) | trtllm defaults | -| `iter1-shared-side-stream` | `llm_args.yaml` (identity) | `75a471a` | gsm8k **94.7688** | 82.43 | 1948.78 (con=256) | shared expert forked onto a side stream | -| `mtp1` | identity + `configs/mtp1.yaml` | `bdbb1ab` | covered by the mtp3 pair, below | 135.11 | 2184.44 (con=256) | MTP, `max_draft_len: 1` | -| `mtp2` | identity + `configs/mtp2.yaml` | `bdbb1ab` | covered by the mtp3 pair, below | 164.94 | 2190.04 (con=256) | MTP, `max_draft_len: 2` | -| `mtp3` | identity + `configs/mtp3.yaml` | `bdbb1ab` | gsm8k **95.2237** (paired, +0.1517) | **190.45** | **2262.06** (con=256) | MTP, `max_draft_len: 3` — the frontier | -| `trtllm` | identity + `configs/trtllm-ref-boot.yaml` | `70b1ecd` | ungated | 73.94 | 1522.73 (con=256) | stock in-tree modeling at its **default** MoE transport (DeepEPLowLatency); config is boot-forced, see below | -| `trtllm-nodeepep` | identity + `configs/trtllm-ref-boot.yaml`, `TRTLLM_CAN_USE_DEEP_EP=0` | `471e6a3` | ungated | 79.20 | 1768.06 (con=256) | the same, on **AllGatherReduceScatter** — the transport control, 3 points | -| `trtllm-mtp3` | identity + `configs/trtllm-ref-mtp3.yaml`, `TRTLLM_CAN_USE_DEEP_EP=0` | `471e6a3` | ungated | 178.66 | 1780.70 (con=256) | **stock modeling with MTP at the same `max_draft_len: 3`** — the reference `mtp3` should be read against | -| `trtllm-tuned` | — | — | — | — | — | superseded by `trtllm-mtp3`, which measures exactly this; see below | -| **`trtllm-aligned`** | **`llm_args.yaml` (identity), `TRTLLM_CAN_USE_DEEP_EP=0`** | `471e6a3` | ungated | **77.82** | **1887.78** (con=256) | **stock modeling at this target's own config — no boot knobs at all. Supersedes `trtllm` / `trtllm-nodeepep`** | -| **`trtllm-aligned-mtp3`** | **identity + `configs/mtp3.yaml`, `TRTLLM_CAN_USE_DEEP_EP=0`** | `471e6a3` | ungated | **181.74** | **1728.86** (con=256) | **stock modeling with MTP at the same config file `mtp3` uses. Supersedes `trtllm-mtp3`** | - -Per point, output tok/s (whole 4-GPU node): - -| con | `trtllm` | `baseline` | `iter1` | iter1/trtllm | t TPOT | i TPOT | t TTFT | i TTFT | -|---|---|---|---|---|---|---|---|---| -| 1 | 73.28 | 75.63 | 81.90 | **1.118x** | 13.525 | 12.132 | 138.6 | 92.4 | -| 2 | 140.11 | 145.44 | 157.59 | 1.125x | 14.059 | 12.554 | 234.0 | 151.9 | -| 4 | 270.94 | 285.41 | 308.40 | 1.138x | 14.525 | 12.831 | 258.2 | 155.2 | -| 8 | 474.91 | 513.02 | 552.59 | 1.164x | 16.486 | 14.279 | 383.2 | 216.2 | -| 16 | 818.94 | 913.71 | 972.02 | 1.187x | 19.091 | 16.131 | 464.4 | 352.0 | -| 32 | 1360.02 | 1584.57 | 1673.50 | 1.230x | 23.033 | 18.514 | 518.1 | 637.5 | -| 64 | 2335.92 | 2724.10 | 2817.55 | 1.206x | 26.745 | 21.639 | 665.4 | 1116.6 | -| 128 | 3835.72 | 4507.69 | 4566.30 | 1.190x | 32.373 | 26.663 | 930.3 | 1406.9 | -| 256 | 6090.94 | 7627.45 | 7795.10 | **1.280x** | 40.129 | 31.092 | 1427.8 | 1762.4 | - -**That 11.8-28.0% is measured against a reference that could not run this -target's own config, and a later session showed it did not have to be that -way.** Read the next subsection before quoting any number from the table -above. Mean TPOT is lower at every point, and the mean-TTFT crossover at -con=32..256 is a property of the reference's forced config rather than of the -modeling layer. - -### The reference's boot config turned out to be avoidable, and it cost 24% - -`configs/trtllm-ref-boot.yaml` exists because stock trtllm OOMs at engine -construction on this checkpoint: 113.9 GiB per rank allocated before any weight -is materialized, `= 1.96 GiB x 58 MoE layers`, which is DeepEP low-latency's -communication workspace sized from `max_num_tokens`. Capping `max_num_tokens` -and `max_seq_len` to 2048 treats that symptom. **`TRTLLM_CAN_USE_DEEP_EP=0` -removes the cause** — the MoE communication factory then falls through to -`AllGatherReduceScatter`, which its own source comment calls "always works", -and which is the transport *this target implements by hand*. Stock trtllm then -boots at the **identity config**, `max_num_tokens` 8192 and `max_seq_len` -163840, peak 179,188 MiB with ~4 GB to spare. That switch was not known when -the boot config was written. - -Re-measured on 2026-08-03, GPUs 3-6, both sides byte-identical config -(`trtllm-aligned` and `trtllm-aligned-mtp3`): - -| con | 1 | 8 | 32 | 64 | 128 | 256 | -|---|---|---|---|---|---|---| -| the boot config was costing the reference | +6.2% | +8.2% | +10.9% | +11.6% | +15.9% | **+24.0%** | -| **`iter1` / `trtllm-aligned`** (no MTP) | 1.052x | 1.075x | **1.109x** | 1.081x | 1.027x | **1.032x** | -| **`mtp3` / `trtllm-aligned-mtp3`** | 1.045x | 1.313x | 1.289x | 1.339x | **1.406x** | **1.308x** | - -**So the honest no-MTP lead is 2.7-10.9%, not 11.8-28.0%**, and the honest MTP -lead is 1.308x. Both older reference lines and the whole -*MoE transport is worth 7-16%* subsection below are superseded by this: the -transport is no longer a variable to isolate, because both systems now run the -same one. - -**And the MTP lead is not a better MTP layer.** Per system, at matched config, -MTP is worth: - -| con | 1 | 8 | 32 | 128 | 256 | -|---|---|---|---|---|---| -| to `staircase` | 2.285x | 1.930x | 1.301x | 1.332x | **1.161x** | -| to stock trtllm | **2.300x** | 1.580x | 1.119x | **0.973x** | **0.916x** | - -They are level at con=1 — stock is fractionally ahead — and stock goes -*negative* from con=128. The 1.308x is that divergence, not a drafting-quality -difference; acceptance agrees to a few percent everywhere. -`workbench/docs/2026-08-03-r1-dep4-staircase-vs-trtllm.md` carries the -kernel-level attribution: it is one family, MoE expert GEMM, and the driver is -**row count, not a re-run draft layer**. MTP gives every generation sequence -`max_draft_len + 1 = 4` query tokens, so the *trunk*'s MoE gathers 4x the rows -(256 -> 1024 per rank at con=256). The two implementations scale differently -under that: CUTLASS grouped GEMM goes from 5 to 8 kernels per MoE layer, while -trtllm-gen's `bmm_E2m1_*` barely moves. Expert-GEMM-family launches per rank -per step: **174 -> 192 (+10%) on our side, 290 -> 482 (+66%) on theirs** — and -of their +192, the MTP layer itself accounts for only 15. **The growth is in -the trunk, not in the draft layer.** - -The MTP curves, against `iter1` (all four are full sweeps; `mtp*` from the -2026-08-02 session). `accept` is mean tokens emitted per engine step, ceiling -`max_draft_len + 1`: - -| con | `iter1` | `mtp1` | `mtp2` | `mtp3` | mtp3/iter1 | i TPOT | m3 TPOT | i TTFT | m3 TTFT | m3 accept | -|---|---|---|---|---|---|---|---|---|---|---| -| 1 | 81.90 | 133.54 | 162.51 | 187.12 | **2.285x** | 12.132 | 5.251 | 92.4 | 100.9 | 3.4904 | -| 2 | 157.59 | 252.37 | 313.21 | 341.41 | 2.166x | 12.554 | 5.611 | 151.9 | 124.7 | 3.4040 | -| 4 | 308.40 | 501.37 | 622.24 | 677.82 | 2.198x | 12.831 | 5.689 | 155.2 | 123.8 | 3.3923 | -| 8 | 552.59 | 803.10 | 928.24 | 1066.23 | 1.930x | 14.279 | 7.145 | 216.2 | 166.6 | 3.4151 | -| 16 | 972.02 | 1316.50 | 1313.32 | 1431.51 | 1.473x | 16.131 | 9.394 | 352.0 | 217.9 | 3.3905 | -| 32 | 1673.50 | 2135.72 | 2091.70 | 2176.90 | 1.301x | 18.514 | 12.199 | 637.5 | 282.6 | 3.3356 | -| 64 | 2817.55 | 3341.71 | 3627.89 | 3723.24 | 1.321x | 21.639 | 13.618 | 1116.6 | 415.4 | 3.4407 | -| 128 | 4566.30 | 5877.32 | 5733.13 | 6080.38 | 1.332x | 26.663 | 17.524 | 1406.9 | 612.1 | 3.4173 | -| 256 | 7795.10 | 8737.78 | 8760.17 | 9048.23 | **1.161x** | 31.092 | 23.758 | 1762.4 | 923.9 | 3.4033 | - -**`mtp3` is ahead of `mtp1` and `mtp2` at every one of the nine -concurrencies**, and ahead of `iter1` at every one, so the `max_draft_len` axis -never turns over inside the certified range. Mean TTFT also falls at every -point except con=1, where it rises 92.4 -> 100.9 ms. - -**Do not read the `mtp3` column against the `trtllm` row.** That row has no -speculative decoding at all, so the ratio between them is mostly "one system -drafts and the other does not", not a modeling-layer delta. The comparison -this target should be quoted on is `mtp3` against **`trtllm-mtp3`** — stock -in-tree modeling, same checkpoint, same `max_draft_len: 3`, same MoE -transport: - -| con | `trtllm-mtp3` | `mtp3` | **mtp3 / trtllm-mtp3** | their accept | our accept | tm TPOT | m3 TPOT | tm TTFT | m3 TTFT | -|---|---|---|---|---|---|---|---|---|---| -| 1 | 175.44 | 187.12 | **1.067x** | 3.5108 | 3.4904 | 5.597 | 5.251 | 110.6 | 100.9 | -| 2 | 342.16 | 341.41 | **0.998x** | 3.5415 | 3.4040 | 5.709 | 5.611 | 128.4 | 124.7 | -| 4 | 621.75 | 677.82 | 1.090x | 3.3315 | 3.3923 | 6.183 | 5.689 | 126.3 | 123.8 | -| 8 | 845.69 | 1066.23 | 1.261x | 3.3058 | 3.4151 | 8.788 | 7.145 | 225.3 | 166.6 | -| 16 | 1044.72 | 1431.51 | **1.370x** | 3.3305 | 3.3905 | 12.020 | 9.394 | 270.3 | 217.9 | -| 32 | 1768.44 | 2176.90 | 1.231x | 3.3458 | 3.3356 | 14.690 | 12.199 | 335.1 | 282.6 | -| 64 | 2797.56 | 3723.24 | 1.331x | 3.2835 | 3.4407 | 18.454 | 13.618 | 446.2 | 415.4 | -| 128 | 4510.25 | 6080.38 | 1.348x | 3.3125 | 3.4173 | 23.658 | 17.524 | 688.2 | 612.1 | -| 256 | 7122.80 | 9048.23 | **1.270x** | 3.3145 | 3.4033 | 29.746 | 23.758 | 1134.4 | 923.9 | - -**The honest headline is 1.27x at con=256 and 1.00-1.09x at con=1..4**, not -the 2.29x the no-MTP row invites. Two things follow, and they are the point of -this reference: - -* **Both systems draft equally well.** Acceptance agrees to within a few - percent at every concurrency (theirs 3.28-3.54, ours 3.40-3.49), which is - the same conclusion the level-3 acceptance gate reached and is what makes - the throughput ratio a *speed* comparison rather than a draft-quality one. -* **The modeling-layer delta survives MTP essentially unchanged.** At con=256 - it is **1.270x with MTP against 1.280x without** — and the gap decomposition - below prices the reference's forced config at 1.035x on our side at that - point, so both reduce to a ~1.23x modeling delta. MTP moved the whole - frontier; it did not move the distance between the two systems. - -At con=1..4 the two are level. Stock trtllm's low-concurrency MTP path is as -good as ours; our lead only opens up from con=8, which is where `iter1`'s -side-stream overlap and the rest of the trunk work start to matter. - -#### The MoE transport is worth 7-16% to stock trtllm, and it is not noise - -**Superseded by `trtllm-aligned`.** This subsection isolated the transport as -a variable because the reference was stuck on `configs/trtllm-ref-boot.yaml`. -It no longer is: with `TRTLLM_CAN_USE_DEEP_EP=0` stock trtllm boots at the -identity config, so both systems now run `AllGatherReduceScatter` and there -is nothing left to isolate. The measurement below stands on its own. - -`trtllm-mtp3` cannot run on DeepEPLowLatency — that transport's dispatch -accepts only NVFP4 hidden states and the MTP layer is bf16 — so it is forced -onto `AllGatherReduceScatter`, while the existing `trtllm` curve was measured -on DeepEPLowLatency. `trtllm-nodeepep` isolates that one variable: stock -modeling, **no** MTP, the same boot config, `TRTLLM_CAN_USE_DEEP_EP=0`. - -| con | `trtllm` (DeepEPLowLatency) | `trtllm-nodeepep` (AllGatherReduceScatter) | delta | -|---|---|---|---| -| 1 | 73.28 | 78.64 | **+7.32%** | -| 32 | 1360.02 | 1491.47 | **+9.67%** | -| 256 | 6090.94 | 7072.22 | **+16.11%** | - -This does **not** collapse into the existing reference — it is 7 to 16 times -this session's measured con=256 floor of 0.01%. Two corrections follow, and -both cut against this target: - -* **Stock trtllm's default transport is the slower one here**, so the main - table's `trtllm` row understates what stock trtllm can do on this host. - Against the faster stock configuration, `iter1`'s no-MTP lead is - **1.041x / 1.122x / 1.102x** at con 1/32/256 — not the 1.118x / 1.230x / - 1.280x measured against the default. The 11.8-28.0% claim above is against - `trtllm` **as stock defaults configure it**, which is the honest definition - of that reference line but is not the best stock can do. -* **MTP is worth far less to stock trtllm at scale than to us, once the - transport is held fixed.** `trtllm-mtp3 / trtllm-nodeepep` is **2.231x at - con=1, 1.186x at con=32 and 1.007x at con=256** — at the throughput end, - drafting buys stock trtllm essentially nothing. Ours, `mtp3 / iter1`, is - **2.285x / 1.301x / 1.161x**. That divergence at high concurrency is the - clearest single statement of what this target's modeling layer is worth - under MTP: the two systems draft equally well and both pay a step-cost - penalty that grows with batch, but ours stays ahead of the penalty and - stock's does not. - -`trtllm-nodeepep` is three points, not a swept curve — con 1/32/256 at the -same rounds the full sweeps use, so it pairs point-to-point with `trtllm`. -con=128 was deliberately skipped: this session measured 4.2% run-to-run spread -there, so it could not have carried the claim. - -**Read these curves as "did anything get slower, and where does it stop -paying" — not as "is MTP worth enabling".** The harness builds prompts from -uniformly random token ids and then generates with `--ignore-eos`, so what the -draft layer predicts is the model's own continuation of a nonsense prefix, -which degenerates into repetition — and repetition drafts near-perfectly. The -`accept` column above, 3.39-3.49 of a ceiling of 4.0, is that artefact. **The -value number for MTP on this checkpoint is the real-text one**, from the paired -gsm8k runs recorded under *The MTP variant's own gate records*: execution -**37.025 s -> 31.440 s, 1.178x** over 1319 genuine prompts. - -### Reference lines - -* **`trtllm`** = the original checkpoint under stock in-tree modeling, this - target's `llm_args.yaml` (identity config: `tensor_parallel_size: 4`, - `moe_expert_parallel_size: 4`, `enable_attention_dp: true`), **plus - `configs/trtllm-ref-boot.yaml`, which is boot-forced rather than tuned.** - Ungated. `tensorrt_llm 1.3.0rc21`. - - **Stock trtllm cannot boot this checkpoint at `dep4` on a 178.34 GiB B200 - under its own defaults.** Measured, per rank: at the default - `max_num_tokens: 8192`, **113.9 GiB is allocated at engine construction - outside the torch allocator, before any weight is materialized** — an - `nvidia-smi` sampler at 5 s intervals caught it going 20 MiB -> - 116,662 MiB inside one sample, immediately after the MoE communication - factory logged `Selected communication strategy: DeepEPLowLatency` once per - MoE layer (`NVLinkOneSided` / `NVLinkTwoSided` / `DeepEP` all reported - `not available: Invalid Argument`). Model init then dies in - `init_meta_tensor` with 64.17 GiB held by PyTorch and 113.8 GiB outside it. - 113.9 GiB / 58 MoE layers = 1.96 GiB per layer. `NVSHMEM_SYMMETRIC_SIZE=1g` - does not change it. - - The allocation is proportional to `max_num_tokens`: at 2048 the weights - load (135.85 GiB torch, including 1.53 GiB of CUDA-graph pools) and the - failure moves to `configure_kv_cache_capacity`, which asks for 7.00 GiB - with 5.61 GiB free. `kv_cache_config.free_gpu_memory_fraction` does **not** - move that 7.00 GiB (identical failure at 0.9 and at 0.6). The failure's own - memory ledger names the remaining lever, and `max_seq_len: 2048` — exactly - the harness's ISL+OSL, and `max_blocks_per_seq` 5120 -> 64 — is what makes - it boot. Both knobs together are the reference's config. - -* **`trtllm-nodeepep`** = the `trtllm` line with one variable moved: - `TRTLLM_CAN_USE_DEEP_EP=0`, which disables DeepEP and DeepEPLowLatency - together and lands the MoE communication factory on - `AllGatherReduceScatter` — the strategy this target implements by hand. - Stock in-tree modeling, no MTP, same boot config, three points. Ungated. - It exists to keep the transport from being confounded with MTP, and it is - worth 7-16%; see above. - -* **`trtllm-mtp3`** = the original checkpoint under stock in-tree modeling - with **the same speculative config this target's kept variant uses** - (`configs/trtllm-ref-mtp3.yaml` = the boot knobs + `decoding_type: MTP`, - `max_draft_len: 3`, `free_gpu_memory_fraction: 0.75`), plus - `TRTLLM_CAN_USE_DEEP_EP=0`. Ungated. Full nine-point sweep, no - `--acceptance`. **This is the reference `mtp3` is quoted against**, and it - replaces what the `trtllm-tuned` slot was for: it *is* stock modeling under - this campaign's final best config. - - It carries three forced deviations from `configs/mtp3.yaml` — - `max_num_tokens` / `max_seq_len: 2048` to boot at all, and the transport - variable — so it is not a *portable-config* split in the usual sense. Both - are priced rather than waved away: the boot knobs are worth `+1.5% / −0.6% / - −3.4%` on our side (the gap decomposition below), and the transport is worth - `+7.3% / +9.7% / +16.1%` on stock's side and is applied to `trtllm-mtp3` - already. Netting the boot knobs out of the con=256 ratio leaves ~1.23x, the - same modeling delta the no-MTP decomposition finds. - -* **`trtllm-tuned`** — **retired, superseded by `trtllm-mtp3`.** The slot means - "stock modeling plus the campaign's final best config", and that is now - measured rather than argued: the final best config is `configs/mtp3.yaml` - and `trtllm-mtp3` runs stock modeling under it. (For the identity campaign - the slot was genuinely degenerate — no config variant was kept, and stock - trtllm cannot boot at the identity config at all.) - -### The gap decomposition - -Because the reference carries a forced config, the honest split is measured -rather than asserted: `probe-refcfg` ran **this target's `iter1` code under -the reference's own config**, so both systems can be compared at identical -knobs. - -| con | staircase @ identity | staircase @ ref config | trtllm @ ref config | config-matched delta | end-to-end | -|---|---|---|---|---|---| -| 1 | 81.90 | 83.13 (+1.51%) | 73.28 | **1.134x** | 1.118x | -| 32 | 1682.70* | 1673.15 (−0.57%) | 1360.02 | **1.230x** | 1.230x | -| 256 | 7802.74* | 7535.64 (−3.42%) | 6090.94 | **1.237x** | 1.280x | - -\* paired two-round probes, the shape `probe-refcfg` used; the full-sweep -numbers are in the table above. - -**0% of the end-to-end gap is portable config and 100% of it is the -modeling-layer delta** — the campaign kept no config variant, so there is no -portable config gain to transfer. The reference's forced config is worth -`+1.5% / −0.6% / −3.4%` on our side, i.e. at con=256 the end-to-end 1.280x is -a 1.237x modeling delta plus a 1.035x config difference that happens to -favour the identity config on throughput. At con=1 and con=32 the -config-matched and end-to-end numbers agree to within the session spread. - -That forced config is also a real **TTFT-vs-throughput trade on this target**, -which is what the TTFT crossover in the main table is: at con=256 it costs -−3.42% output throughput and buys **mean TTFT 2269.5 -> 1645.7 ms (−27.5%)**; -at con=32, −0.57% for 627.8 -> 384.0 ms (−38.8%). It is correctly not kept — -the Pareto axes are tok/s/user and tok/s/GPU, and it loses on one and is flat -on the other — but a latency-sensitive deployment of this target should reach -for `max_num_tokens: 2048` first. - -### Iterations - -**Two kept: one modeling change (`iter1`), one config axis (`mtp3`).** - -**`iter1-shared-side-stream` — the shared expert overlaps the MoE round -trip.** *Evidence.* An nsys window of 100 executor iterations at con=256 -(`TLLM_PROFILE_START_STOP=1200-1300`, `-c cudaProfilerApi`, -`--cuda-graph-trace node`) caught all 4 ranks, 778,000 kernels, and **100 -pure-decode steps** — no context FMHA kernel appears at all, and -`cudaGraphLaunch` = 400 against 4 ranks x 100 steps = **1.00 graph replay per -rank per step, 100% coverage**. Per rank the GPU was **98.5% busy** (union -2822.3 ms against a 2865.1 ms window, idle **1.49%**), so the target is -GPU-bound, not host-bound. Walking the intervals per device — kernels on -different GPUs genuinely run in parallel — gave exclusive time (the interval -covered by no other kernel): - -| family | sum µs/rank-step | exclusive µs/rank-step | exclusive/union | % of GPU busy | -|---|---|---|---|---| -| routed expert GEMMs (`bmm_E2m1*`, `bmm_Bfloat16*`) | 13488.2 | **13078.7** | 98.9% | **46.3%** | -| cuBLAS bf16 GEMMs (`nvjet_*`, attention path) | 6828.4 | 5665.5 | 86.6% | 20.1% | -| NCCL collectives (`AllGather`/`ReduceScatter` `RING_LL`) | 3572.7 | **3369.7** | **94.3%** | **11.9%** | -| MLA decode FMHA | 1079.7 | 1015.2 | 94.0% | 3.6% | -| `splitKreduce_kernel` | 1062.6 | 324.3 | 30.5% | 1.1% | - -GPU busy is 28.22 ms per rank-step. The collectives are **94.3% exclusive** — -for almost all of the time one is running, the device is running nothing else -— which is a 3.37 ms per-step window with nothing in it. - -*Change.* The shared expert and the routed round trip both read the -post-attention `o` and meet only at the final `add`, but on one stream the -shared expert's five kernels sat *in front of* the all-gather. `forward` now -forks that branch onto a side stream (`side.wait_stream(main)` -> -`with torch.cuda.stream(side)` -> `main.wait_stream(side)` before the add), -so it runs inside the collective window. Both collective contracts certify -"a side stream joined to the current one on both ends", which is exactly this -shape; the same fork/join pair is what propagates a CUDA-graph capture into -the branch and back, and the captured decode graph keeps working (smoke's 10 -greedy continuations are byte-identical to the baseline's). - -*Effect.* Output throughput rises at **every** concurrency — +8.28, +8.35, -+8.05, +7.71, +6.38, +5.61, +3.43, +1.30, +2.20 % at con 1..256 — and mean -TPOT falls at every point. Frontier: con=1 tok/s/user 76.09 -> **82.43** -(+8.3%), peak tok/s/GPU 1906.86 -> **1948.78** (+2.2%). The gain is largest -at low concurrency, which the profile predicts: the collectives are -latency-bound and roughly batch-independent, so the window they leave is a -larger share of a smaller step, while at con=256 the shared expert competes -for HBM bandwidth with the routed expert GEMMs and only part of it hides. -Accuracy: **gsm8k 94.7688** `exact_match,flexible-extract` (strict-match -94.6171) against reference 94.9962, gate >= 90.0 — passed. The 0.23 delta is -3 questions of 1319, well inside the run's own ±0.6133 stderr, and -strict-match is bit-for-bit the same score as the assembly run. - -**This iteration makes two statements elsewhere in this file stale**, and -they are gate-record prose that a tuner does not edit: the *parallel split* -section and the *Vocabulary* section both say the two collectives are driven -"in their uniform form only" and that every rank "pads to -`max(all_rank_num_tokens)`". Both are still true — `iter1` did not change -either collective's arguments — but the surrounding claim that the forward is -single-stream no longer is. The certification consumed is unchanged: same -entries, same arguments, same uniform form, and both contracts certify side -streams explicitly. - -**`mtp3` — multi-token prediction at `max_draft_len: 3`, the whole certified -range swept.** *Evidence.* The MTP layer's own arithmetic -(`docs/models/multi-token-prediction.md`) puts the break-even acceptance at -1.055 / 1.110 / 1.164 for draft lengths 1 / 2 / 3, and predicts the step cost -to be flat in concurrency because a decode step is weight-bandwidth-bound. It -also flags the high-concurrency end as the genuinely uncertain one: once a -rank's batch already activates all 64 of its local experts, `draft_len + 1` -times the rows means the same weight bytes with several times the arithmetic, -and the layer can cross from memory-bound to compute-bound. - -*Change.* Config only — `configs/mtp{1,2,3}.yaml`, one full Pareto sweep each, -no modeling edit. Nothing else moved: `modeling.py` is byte-identical across -all three, and `--acceptance` was deliberately **not** used (below). - -*Effect.* `mtp3` takes the frontier on both axes: con=1 tok/s/user -**82.43 -> 190.45 (+131.1%)** and peak tok/s/GPU **1948.78 -> 2262.06 -(+16.1%)**, ahead of `iter1` at all nine concurrencies and ahead of `mtp1` and -`mtp2` at all nine. Against the matched-draft-length reference `trtllm-mtp3` -it leads by **1.270x at con=256** and is level (1.00-1.09x) at con=1..4 — -that, not the ratio to the non-drafting `trtllm` row, is what this iteration -is worth against stock. Decomposing throughput into the two terms it is made -of — the engine step now emits `accept` tokens instead of 1, and costs more — - -| con | accept | step cost vs a non-drafting step | decode speedup | end-to-end | -|---|---|---|---|---| -| 1 | 3.4904 | 1.511 | 2.311x | 2.285x | -| 32 | 3.3356 | 2.198 | 1.518x | 1.301x | -| 256 | 3.4033 | 2.600 | 1.309x | 1.161x | - -**Acceptance is essentially flat in concurrency and the whole decay is on the -cost side.** Across con 1..256 acceptance moves only 3.49 -> 3.40 (and -`mtp1` 1.967 -> 1.949 of 2.0, `mtp2` 2.748 -> 2.754 of 3.0) — never below 2.86x -the 1.164 break-even. The step cost, which the weight-bytes model says -should sit flat at 1.164, instead climbs **1.511 -> 2.600**. Reading the -marginal cost of each successive draft step at con=256 — **+0.645, +0.569, -+0.386** — against con=1's **+0.200, +0.173, +0.137** shows what it is: the -term that grows is proportional to the *rows* a step carries, not to the MTP -layer's fixed weight bytes, and the marginal cost falls with each further draft -step in the same way the marginal row count does (1 -> 2 rows is +100%, 2 -> 3 -is +50%, 3 -> 4 is +33%). That is the memory-bound-to-compute-bound crossover -the mechanism doc predicted, measured. It never becomes steep enough to -overtake acceptance, which is why the curve is still climbing at 3. - -*Gate.* No new accuracy run was needed and none is claimed: the paired gsm8k -record under *The MTP variant's own gate records* was measured on -**`configs/mtp3.yaml` itself**, and `modeling.py` has not changed since -(`4194dcb`, and this campaign added no code). identity **95.0720** / mtp3 -**95.2237**, Δ **+0.1517** against `|Δ| < 1.2`. `mtp1` and `mtp2` are not -separately gated and are not the recommended setting; they are the two -measured points that establish the axis is monotone, and the note in each file -says so. What the gate certifies is that the variant did not break the model — -rejection sampling makes even a miscomputed draft layer emit correct text, so -`acceptance_length` against stock trtllm, not the score, is what says drafting -works. - -### Rejected, with what was actually measured - -* **`max_seq_len: 2048` — rejected, and it retires a blocked knob.** - `kv_cache_config.tokens_per_block: 64` is uncertified over this target's - fp8 latent pool, and it was worth **+22% at con=256** on the - `deepseek-v3-lite-nvfp4/sm_100/tp1` sibling, so it is the knob a campaign - here would want most. `max_seq_len` reaches the *same* - `max_blocks_per_seq` reduction by the other route (5120 -> 64, an 80x cut - where page 64 gives 2x), and measured **−0.13% at con=32 and +4.08% at - con=256** against a paired probe — where a clean re-measure puts the same - point at +0.10%. The profile says why: **GPU idle is 1.49%**, so there is - no exposed host bookkeeping to recover. The tp1 sibling's win came from a - host-bound decode (GPU idle 0.401); this target is 40x larger per step and - the same host work is entirely hidden. **A vocabulary request to certify - page 64 over an fp8 pool would not pay for itself here.** -* **Ragged collectives on the non-captured steps — rejected.** Under - attention DP a mixed step (one rank prefilling beside three decoding) pads - every rank to `max(counts)`, so the expert call runs over `4 x max` rows of - which ~70% are zeros. `_dp_rows` was changed to return a `sizes` vector - whenever the counts disagree and to pass it to both collectives (the - uniform form is unreachable-by-construction under capture, where counts are - equal). The engine does publish true per-rank counts — - `_get_all_rank_num_tokens` is a plain `tp_allgather` of - `attn_metadata.num_tokens`, no padding — so the branch is reachable, and - smoke passed. It measured **+0.23% at con=32 and +4.08% at con=256** - against the same perturbed baseline, i.e. **+0.10% against the clean - cluster**: 7634.91 against 7635.26 (`max_seq_len` probe) and 7627.45 - (full baseline). The padded rows are real but are not a material share of - GPU time on this workload; the assembler's uniform-only choice stands. -* **`stream_interval`, `cuda_graph_config.*`, `kv_cache_config.*`, - `scheduler_config.capacity_scheduler_policy` — not swept, on profile - evidence.** GPU idle 1.49% leaves nothing for the host-side response path - to recover; graph replay coverage is already 100% at `enable_padding: - false`, so padding could only coarsen the grid (−2.55% at con=256 on the - `dep4` sibling); and the KV pool is 44.76 GiB = **1,368,000 tokens per - rank** against the 64 requests x 2048 tokens = 131,072 a rank holds at - con=256, a 10.4x margin, so capacity cannot bind. -* **`speculative_config.draft_len_schedule` — rejected twice over: the premise - is refuted, and on this target it deadlocks.** What was measured: - `configs/mtp3-sched.yaml`, `max_draft_len: 3` with - `draft_len_schedule: {2: 3, 16: 2, 128: 1}` (thresholds are per-*rank* batch - size, so under attention DP at dep4 they map to con 1..8 / 16..64 / - 128..256), full sweep attempted on GPUs 3-6. - - *The premise.* The schedule's case is "long drafts at low concurrency, - shorter at high", which needs the best draft length to fall as batch grows. - The three full sweeps say it does not: `mtp3` leads at **all nine** - concurrencies, and acceptance decays by only 3.49 -> 3.40 over the whole - range. There is no crossover inside the certified range, so a schedule that - shortens the draft at high batch can only give throughput away. - - *The deadlock.* The run never produced a point. Engine build, CUDA-graph - capture (one graph per `(batch_size, draft_len)` pair — the log shows - `batch size=3..16, draft_len=2` and `batch size=1,2, draft_len=3`) and warmup - all succeeded; roughly a minute into serving, **all four ranks hung and - trtllm's own HangDetector fired at 300 s and hard-killed via `MPI_Abort`**. - The stacks put the ranks at three *different* points of one forward — the q - up-projection, the MLA generation call, and the MoE `fp4_quantize` — i.e. not - in lockstep. The mechanism is structural: `_handle_dynamic_draft_len` - resolves `runtime_draft_len` from **`scheduled_batch.batch_size`, which is - each rank's own local batch**, with no cross-rank reduction, and attention DP - does not equalize batch sizes — `_pad_attention_dp_dummy_request` only tops a - rank up from zero to one. Two ranks either side of a threshold therefore run - different numbers of MTP replays, hence different numbers of MoE collectives, - and the job wedges. Nothing validates the combination at config time. - - *And it is not a single-variable comparison anyway*: setting the schedule - makes the runtime log `Automatically enabling cuda_graph_config.enable_padding - because draft_len_schedule is set`, flipping a knob this target measured as - costly (padding can only coarsen an already 100%-covered replay grid). The - variant file was deleted; this entry is its record. -* **`--acceptance` on the Pareto sweeps — rejected, and not needed.** The flag - sets `enable_iter_perf_stats: true` and lifts `iter_stats_max_iterations`, so - a label carrying it is not comparable to `baseline` / `iter1`. Priced in a - paired probe on one `mtp3` config, back to back, con 1/32/256: - **−1.39% / −3.52% / −3.96%** (212.34 -> 209.38, 2056.91 -> 1984.44, - 8229.25 -> 7903.35 tok/s). That is far outside the sub-1% floor, so it was - carried on no sweep. It is also unnecessary here: the benchmark client - already reports `avg_decoded_tokens_per_iter` per request, sourced from the - **response body** rather than the `/metrics` iteration-stats stream, so it - survives with the flag off — at con=1 the two probes recorded a bit-identical - `3.8863`, and against the engine's own `acceptance_length` the client proxy - agrees to **0.99-1.02x**. Every `accept` number in this section is that free - proxy, measured at zero config cost on the same sweeps as the throughput. -* **`max_draft_len` 1 and 2 — measured and dominated.** `mtp1` and `mtp2` are - full sweeps in the tables above; `mtp3` beats both at every concurrency. They - are kept as configs because they are the evidence that the axis is monotone, - not as recommended settings. - -### Caveats - -* **Session variance is measured, and the dominant term is co-tenancy, not - sampling.** Three unchanged-config measurements of con=256 landed at - 7627.45 (full sweep, 1280 requests), 7635.26 and 7634.91 (two-round probes, - 512 requests) — a **0.10% spread across two different request counts**. A - fourth, `probe-baseline`, read 7335.94 (**−3.9%**); an 8-GPU neighbour - holding 1.9 GiB on GPUs 0-6 — including this campaign's device set — was - caught in `nvidia-smi` minutes later and was gone within two. That probe is - the only contaminated measurement in the campaign and no kept result rests - on it. **Read the noise floor as well under 1% between clean runs, with a - transient co-tenant worth ~4%.** -* **Probes and full sweeps are not interchangeable at the throughput end for - TTFT.** The same config measured mean TTFT 1765.6 ms over 1280 requests and - 2279.7 ms over 512: with only two waves of 256 the first wave's queue - dominates the mean. Throughput is insensitive to this (0.10% above); TTFT - is not. Every TTFT comparison above is probe-to-probe or sweep-to-sweep. -* Host compute-apps and load average were recorded before and after every - label (`logs/hoststate/`); GPUs 3-6 carried no other tenant for any - measured curve. -* Both identity-config curves are single full sweeps; `iter1`'s two endpoints - were additionally reproduced in a paired probe before the sweep was run. -* `perf/data/` is machine-local and not tracked; the figure is. -* **The two sessions are stitched by measurement, not by assumption.** Opening - the 2026-08-02 session, `probe-anchor-iter1` re-measured the `iter1` config - and landed within **−0.14% (con=1) and +0.15% (con=256)** of the 2026-08-01 - full sweep. At wrap-up the `trtllm` reference was spot-checked the same way: - **−0.01% (con=1, TPOT bit-identical at 13.525 ms) and +0.20% (con=256)**. - Neither curve was re-swept. -* **This session's noise floor, measured rather than inherited.** `mtp1` at - con=256 read **8737.78** in its full sweep and **8736.8** in a re-measure - 2.5 hours later at the same request count (`probe-recheck-mtp1`) — - **0.01%**. The same pair at **con=128 differs by 4.2%** (5877.32 vs 5630.1), - so con=128 is this target's noisy point and no claim above rests on it - alone. The `mtp3`-over-`mtp2` margin at con=256 is +3.29%, comfortably - outside the con=256 floor. -* **Probes and full sweeps are much further apart under MTP than without it, - and it is throughput this time, not just TTFT.** The same `mtp3` config - measured con=256 at **8229.25** over 512 requests and **9048.23** over 1280 — - **+9.95%** — where the identity config's probe-to-sweep gap at that point is - 0.15%. A two-round probe under-reports MTP because the ramp and drain, where - batches are small and drafting is least profitable, are a much larger share - of a 512-request point. **Every MTP comparison in this section is - sweep-to-sweep**; an early probe-vs-sweep reading of these curves inverted - the `mtp1`/`mtp3` ordering at con=256 before the full sweeps corrected it. -* **Acceptance at con=1 is prompt-dependent and needs more than a two-request - probe.** The `mtp3` full sweep (20 requests at con=1) reads **3.4904**, while - the two-request probes recorded under *The MTP variant's own gate records* - read 3.8863/3.9274. Both are far above the 1.164 break-even, so the - correctness reading there is unaffected, but the sweep number is the better - estimate of acceptance. -* **A cold host page cache can push the stock-trtllm boot past the harness's - 900 s limit, and it looks like a failure rather than a slow start.** The - first `trtllm-mtp3` attempt returned `[perf] FAILED: server not healthy - after 900s`. It was not hung — it had reached `[Autotuner] Autotuning - process starts` and was still progressing. The time went into reading the - 397 GB checkpoint: the log shows 101 s, 57 s and 38 s gaps between - `Finished prefetching ...` lines, and boot took **837 s to reach the - autotuner** against **220 s** for the same config earlier the same day. - Re-reading the 163 shards took 7 s (~56 GB/s, i.e. already resident, the - failed boot having warmed them), and the retry reached the autotuner in - **219 s** and completed the full sweep. **Read a 900 s boot timeout on this - checkpoint as a page-cache miss first**; warm the shards and retry before - suspecting the model. Nothing in the harness was changed. -* **Host CPU load moved a lot and the measurement did not.** Load average over - the campaign ranged 0.31 to 13.72, and a neighbour ran on GPUs 1-2 at 100% - during wrap-up; the `trtllm` spot-check taken under the *heaviest* load still - reproduced to 0.20%. The previous campaign's −3.9% contamination came from a - co-tenant **on the device set itself**, which is the case to keep watching. - -### Regenerating the figure - -The committed figure plots exactly seven curves — the two **aligned** -stock-trtllm references and the five staircase sweeps — in this order: - -``` -uv run utils/plot_pareto.py \ - targets/deepseek-r1-0528-nvfp4/sm_100/dep4/perf/data/trtllm-aligned \ - targets/deepseek-r1-0528-nvfp4/sm_100/dep4/perf/data/trtllm-aligned-mtp3 \ - targets/deepseek-r1-0528-nvfp4/sm_100/dep4/perf/data/baseline \ - targets/deepseek-r1-0528-nvfp4/sm_100/dep4/perf/data/iter1-shared-side-stream \ - targets/deepseek-r1-0528-nvfp4/sm_100/dep4/perf/data/mtp1 \ - targets/deepseek-r1-0528-nvfp4/sm_100/dep4/perf/data/mtp2 \ - targets/deepseek-r1-0528-nvfp4/sm_100/dep4/perf/data/mtp3 \ - -o targets/deepseek-r1-0528-nvfp4/sm_100/dep4/perf/figures/pareto.png -``` - -**Explicit paths, not the `perf/data` directory.** That directory also holds -the two acceptance-only gate-record labels and the three superseded -boot-forced reference lines; expanding it would put all of them on the figure. -The committed figure plots the two aligned references plus the five staircase -curves. - -**`trtllm-nodeepep` is on the figure on purpose even though it is only three -points.** It is the fastest stock configuration measured here, and it sits -almost on top of `baseline` / `iter1`; leaving it off would let the figure -imply a larger no-MTP lead over stock trtllm than exists. - -**The paths are explicit on purpose: pointing the plotter at `perf/data` -wholesale is wrong here.** That directory also holds `probe-mtp3-acceptance` -and `trtllm-mtp3-acceptance` — a two-point probe and an acceptance-only -reference from the MTP gate records, neither of them a Pareto curve — and the -plotter picks up every subdirectory carrying a `meta.json`, which would put a -second grey reference line on the figure and blow past its categorical slots. - -The `trtllm` curve needs its boot config: -`uv run bench/perf.py --target targets/deepseek-r1-0528-nvfp4/sm_100/dep4 ---label trtllm --trtllm --config -targets/deepseek-r1-0528-nvfp4/sm_100/dep4/configs/trtllm-ref-boot.yaml`. -The three MTP curves are -`uv run bench/perf.py --target --label mtp --config /configs/mtp.yaml` -— no `--acceptance`. - -### Remaining headroom - -* **The `max_draft_len` curve is still climbing where the certification - stops.** `mtp3` is the best of the three at every concurrency and the gain - from 2 to 3 is still positive everywhere (con=1 162.51 -> 187.12, con=256 - 8760.17 -> 9048.23), while acceptance holds at 3.40 of a 4.0 ceiling — 2.9x - the 1.164 break-even. Nothing in the *data* says 3 is the optimum; 3 is the - top of what `catalog/attention/mla_rope_generation` and - `catalog/attention/thop_attention` certify, which is - `predicted_tokens_per_seq` ∈ {1,2,3,4} and hence `max_draft_len` ≤ 3. **A - campaign that wants 4 or beyond needs those two entries certified at - `predicted_tokens_per_seq` 5+ first — a vocabulary request, not a bigger - number in a config.** The step-cost series says the return is decelerating - (marginal cost per draft step at con=256: +0.645, +0.569, +0.386 against - acceptance gains of +0.95, +0.80, +0.65 tokens), so the crossover is probably - not far past 3, but it is unmeasured. -* **The `mtp3` cost decomposition is arithmetic, not a timeline.** The - memory-bound-to-compute-bound reading above is derived from the step-cost - series across four draft lengths and nine concurrencies; no nsys capture was - taken under MTP. A timeline at con=256 with `max_draft_len: 3` would say - *which* kernel family absorbed the growth — the routed expert GEMMs at 4x - rows, the bf16 MTP expert stack, or the extra collectives — and that is the - next thing to profile if the MTP path is tuned further. -* **The collectives are the one identified lever and they are blocked on - vocabulary.** **Corrected 2026-08-03:** the 3369.7 µs / 94.3%-exclusive - figure below is the **pre-`iter1`** profile — it is the measurement that - motivated the side stream ("a 3.37 ms window with nothing in it"), not the - state after it. A post-`iter1` capture at the same window and phase measures - **1930.7 µs per rank per decode step at 55.2% exclusive**: the side stream - moved the shared expert into that window, so the recoverable time is 57% of - what this paragraph claimed. Sizing a transport change off 3369.7 overstates - it. The pre-`iter1` numbers, kept because the rest of the analysis rests on - them: **3369.7 µs per rank per decode step of exclusive GPU time, 11.9% of - the step, at 94.3% exclusive**, in two - `RING_LL` NCCL calls per MoE layer. `comm/allgather` and - `comm/reducescatter` expose **no strategy argument and no workspace**, so - the ONESHOT swap that won 2.03x per call on the `tep4` sibling does not - exist here. Closing it needs new catalog entries: a strategy-carrying - gather/scatter, or the `moe_a2a_dispatch` / `moe_a2a_combine` family (a - stateful workspace pair, so `moe_a2a_initialize` and - `moe_a2a_get_combine_payload_tensor` come with it). Note `iter1` has - already spent part of this window — a transport change would now compete - with the shared expert for it, so the two do not simply add. -* **Reducing collective *bytes* is not the lever.** At con=256 each rank's - decode batch is 64 rows, so one all-gather receives `3 x 64 x 7168 x 2` = - **2.753 MB** in its measured 25.472 µs = **108 GB/s**, an order of - magnitude under NVLink: `RING_LL` at these sizes is latency-bound, not - bandwidth-bound. Gathering NVFP4 activations instead of bf16 (4032 vs - 14336 B/token, **3.56x** fewer) would buy nothing and would cost extra - calls. -* **The routed expert GEMMs are near the HBM roofline and are the floor.** - 46.3% of GPU busy, 98.9% exclusive. With top-8 of 256 experts over 256 - gathered tokens every one of a rank's 64 local experts is active, so a step - reads the whole local stack: FC1 (gate+up) is `2 x 2048 x 7168 x 64` at - 0.5625 B/element (NVFP4 data plus the fp8 block scale) = **1.057 GB in - 155.25 µs = 6.81 TB/s**; FC2 (down) is 0.528 GB in 77.97 µs = **6.78 - TB/s** — ~85% of a ~8 TB/s B200 roofline. No config knob and no scheduling - change moves this; only fewer weight bytes would. -* **The bf16 attention GEMMs are the second floor**: 20.1% of GPU busy. The - largest, at 45.21 µs x 61 layers per rank-step, reads `o_proj`'s - 16384x7168 bf16 (235 MB) at **5.19 TB/s** — the one family with visible - daylight to the roofline, but it is cuBLAS's tactic choice, not ours. - These weights are bf16 by checkpoint design (`hf_quant_config.json` - quantizes the MLP, not attention), so this is a checkpoint property rather - than a target one. diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/identity.yaml b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/identity.yaml deleted file mode 100644 index d800d2e26e16..000000000000 --- a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/identity.yaml +++ /dev/null @@ -1,21 +0,0 @@ -# The identity assembly: the target with no variant on top. -# -# This is what the gsm8k record in TARGET.md was measured on. It carries only -# the parallel topology, and the topology is not a tuning knob here -- routing -# derives this target *from* it, so a run that drops it does not get a -# differently-tuned staircase, it gets the built-in DeepseekV3 implementation -# (or, under `TRTLLM_STAIRCASE=require`, an error naming the criterion it -# missed). -# -# Every other file in this directory is this one plus the knobs under -# experiment, so any of them can be passed directly: -# -# TRTLLM_STAIRCASE=require trtllm-eval --model \ -# --extra_llm_api_options /identity.yaml \ -# gsm8k --apply_chat_template --fewshot_as_multiturn -# -# The switch has to be exported before the ranks start; see -# `_router_index.STAIRCASE_ENV` for why. -tensor_parallel_size: 4 -moe_expert_parallel_size: 4 -enable_attention_dp: true diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp1.yaml b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp1.yaml deleted file mode 100644 index ad119802debc..000000000000 --- a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp1.yaml +++ /dev/null @@ -1,69 +0,0 @@ -# Multi-token prediction, 1 draft token per step. -# -# This variant crosses the line a config variant normally respects: it does -# not just move a knob, it selects a **second forward path and a different -# weight-loading path**. With it, the checkpoint's bf16 MTP module at layer 61 -# stops being a predicted non-load and is loaded (212 keys per rank, ~6.3 GB), -# the engine raises the KV pool to 62 layers, and every engine step runs the -# trunk plus `max_draft_len` replays of the MTP layer. Without it the target is -# the identity assembly the gsm8k record was measured on, bit for bit. -# -# decoding_type: MTP + `num_nextn_predict_layers: 1` in the checkpoint config -# resolves to MTP_EAGLE_ONE_MODEL — one MTP layer, replayed autoregressively. -# The layer count is therefore NOT the draft-length knob. -# -# max_draft_len is spelled out on purpose: left unset on the MTP-Eagle path it -# resolves to 1, not to anything workload-derived. It is also the axis that -# costs first-boot time — it becomes the decode kernel's maxSeqLenQ, so each -# value JIT-compiles its own MLA decode kernel per rank (1 and 2 pay separate -# ~5.5 s compiles despite sharing a kernel name; 3 and 4 share another). -# -# Break-even: one draft step adds ~5.5% to a decode step's HBM traffic on this -# checkpoint at dep4, so drafting pays for itself at an acceptance_length of -# about 1.055 here. -# -# **Measured, and dominated — this is not the recommended setting.** A full -# Pareto sweep of all three draft lengths in one session (perf/data/mtp{1,2,3}, -# GPUs 3-6) puts `max_draft_len: 3` ahead at every one of the nine -# concurrencies. This file's curve is 1.63x iter1 at con=1 and 1.12x at -# con=256, against mtp3's 2.29x and 1.16x. Acceptance here is ~1.95 of a 2.0 -# ceiling and barely moves with batch size (1.967 at con=1, 1.949 at con=256). -# Keep it as the low end of the measured axis; use configs/mtp3.yaml to serve. - -# The parallel topology. It is what selects this target at all -- routing -# derives the target from it -- so every variant carries it and these files -# are usable as-is: `trtllm-eval --extra_llm_api_options `. -tensor_parallel_size: 4 -moe_expert_parallel_size: 4 -enable_attention_dp: true - -speculative_config: - decoding_type: MTP - max_draft_len: 1 - -# **This variant does not boot at trtllm's default KV-cache sizing.** Measured -# on this checkpoint at dep4 on a 183.4 GiB B200, with max_draft_len 3: -# -# The MTP module adds 5.87 GiB of declared weights per rank (nvidia-smi read -# 122,076 MiB after model init), and the engine then sizes the KV pool from -# what is free, at -# the default free_gpu_memory_fraction 0.9: 38.95 GiB (1,171,200 tokens, -# 62 layers). Boot then reaches CUDA-graph capture at 182.4 of 183.4 GiB and -# livelocks: three of four ranks spin in cudaFree inside the CUDA caching -# allocator (release_cached_blocks, reached from a plain empty_cuda in the -# *trunk's* MoE runner) while the fourth waits at an MPI_Barrier. No error, -# no OOM exception, no progress -- the run has to be killed. -# -# The drafting forward's post-pool transient is simply larger than the ~19.6 -# GiB the identity config leaves and lives inside. Freeing memory *before* -# the pool is sized does not help: chunking the MTP layer's expert call to -# 2048 rows moved pool sizing by 0.31 GiB (38.64 -> 38.95) and the boot -# failed identically, because the pool grows into whatever is freed. -# -# 0.75 hands the pool ~32.5 GiB instead, leaving ~28 GiB of headroom. That is -# still ~975k KV tokens per rank -- an order of magnitude more than any -# concurrency this target is measured at needs -- so nothing about the -# workload is constrained by it. It is a boot requirement of the variant, not -# a tuning choice, and it is the reason these files carry a second key at all. -kv_cache_config: - free_gpu_memory_fraction: 0.75 diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp2.yaml b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp2.yaml deleted file mode 100644 index 704879fd3b76..000000000000 --- a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp2.yaml +++ /dev/null @@ -1,69 +0,0 @@ -# Multi-token prediction, 2 draft tokens per step. -# -# This variant crosses the line a config variant normally respects: it does -# not just move a knob, it selects a **second forward path and a different -# weight-loading path**. With it, the checkpoint's bf16 MTP module at layer 61 -# stops being a predicted non-load and is loaded (212 keys per rank, ~6.3 GB), -# the engine raises the KV pool to 62 layers, and every engine step runs the -# trunk plus `max_draft_len` replays of the MTP layer. Without it the target is -# the identity assembly the gsm8k record was measured on, bit for bit. -# -# decoding_type: MTP + `num_nextn_predict_layers: 1` in the checkpoint config -# resolves to MTP_EAGLE_ONE_MODEL — one MTP layer, replayed autoregressively -# `max_draft_len` times. The checkpoint's layer count does not bound this. -# -# max_draft_len is spelled out on purpose: left unset on the MTP-Eagle path it -# resolves to 1, not to anything workload-derived. It also sets the generation -# call's predicted_tokens_per_seq (= max_draft_len + 1 = 3 here), which the MLA -# entries certify at 1..4 and which becomes the decode kernel's maxSeqLenQ — -# one extra JIT compile per rank on first boot. -# -# Break-even: two draft steps add ~11% to a decode step's HBM traffic on this -# checkpoint at dep4, so drafting pays for itself at an acceptance_length of -# about 1.110 here. -# -# **Measured, and dominated — this is not the recommended setting.** A full -# Pareto sweep of all three draft lengths in one session (perf/data/mtp{1,2,3}, -# GPUs 3-6) puts `max_draft_len: 3` ahead at every one of the nine -# concurrencies. This file's curve is 1.98x iter1 at con=1 and 1.12x at -# con=256, against mtp3's 2.29x and 1.16x. Acceptance here is ~2.75 of a 3.0 -# ceiling and barely moves with batch size (2.748 at con=1, 2.754 at con=256). -# Keep it as the middle of the measured axis; use configs/mtp3.yaml to serve. - -# The parallel topology. It is what selects this target at all -- routing -# derives the target from it -- so every variant carries it and these files -# are usable as-is: `trtllm-eval --extra_llm_api_options `. -tensor_parallel_size: 4 -moe_expert_parallel_size: 4 -enable_attention_dp: true - -speculative_config: - decoding_type: MTP - max_draft_len: 2 - -# **This variant does not boot at trtllm's default KV-cache sizing.** Measured -# on this checkpoint at dep4 on a 183.4 GiB B200, with max_draft_len 3: -# -# The MTP module adds 5.87 GiB of declared weights per rank (nvidia-smi read -# 122,076 MiB after model init), and the engine then sizes the KV pool from -# what is free, at -# the default free_gpu_memory_fraction 0.9: 38.95 GiB (1,171,200 tokens, -# 62 layers). Boot then reaches CUDA-graph capture at 182.4 of 183.4 GiB and -# livelocks: three of four ranks spin in cudaFree inside the CUDA caching -# allocator (release_cached_blocks, reached from a plain empty_cuda in the -# *trunk's* MoE runner) while the fourth waits at an MPI_Barrier. No error, -# no OOM exception, no progress -- the run has to be killed. -# -# The drafting forward's post-pool transient is simply larger than the ~19.6 -# GiB the identity config leaves and lives inside. Freeing memory *before* -# the pool is sized does not help: chunking the MTP layer's expert call to -# 2048 rows moved pool sizing by 0.31 GiB (38.64 -> 38.95) and the boot -# failed identically, because the pool grows into whatever is freed. -# -# 0.75 hands the pool ~32.5 GiB instead, leaving ~28 GiB of headroom. That is -# still ~975k KV tokens per rank -- an order of magnitude more than any -# concurrency this target is measured at needs -- so nothing about the -# workload is constrained by it. It is a boot requirement of the variant, not -# a tuning choice, and it is the reason these files carry a second key at all. -kv_cache_config: - free_gpu_memory_fraction: 0.75 diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp3.yaml b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp3.yaml deleted file mode 100644 index eb17e13071f3..000000000000 --- a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/mtp3.yaml +++ /dev/null @@ -1,79 +0,0 @@ -# Multi-token prediction, 3 draft tokens per step. -# -# This variant crosses the line a config variant normally respects: it does -# not just move a knob, it selects a **second forward path and a different -# weight-loading path**. With it, the checkpoint's bf16 MTP module at layer 61 -# stops being a predicted non-load and is loaded (212 keys per rank, ~6.3 GB), -# the engine raises the KV pool to 62 layers, and every engine step runs the -# trunk plus `max_draft_len` replays of the MTP layer. Without it the target is -# the identity assembly the gsm8k record was measured on, bit for bit. -# -# decoding_type: MTP + `num_nextn_predict_layers: 1` in the checkpoint config -# resolves to MTP_EAGLE_ONE_MODEL — one MTP layer, replayed autoregressively -# `max_draft_len` times. The checkpoint's layer count does not bound this. -# -# max_draft_len is spelled out on purpose: left unset on the MTP-Eagle path it -# resolves to 1, not to anything workload-derived. It also sets the generation -# call's predicted_tokens_per_seq (= max_draft_len + 1 = 4 here), which is the -# top of what the MLA entries certify (1..4) — **going past 3 needs the -# certification extended first**, not just a bigger number here. It becomes the -# decode kernel's maxSeqLenQ too, so it costs one extra JIT compile per rank on -# first boot. -# -# Break-even: three draft steps add ~16.4% to a decode step's HBM traffic on -# this checkpoint at dep4, so drafting pays for itself at an acceptance_length -# of about 1.164 here. -# -# **This is the recommended MTP setting, and the measured best of the three.** -# A full Pareto sweep of all three draft lengths in one session -# (perf/data/mtp{1,2,3}, GPUs 3-6) puts this file ahead at every one of the -# nine concurrencies, on both frontier axes: con=1 tok/s/user 82.43 -> 190.45 -# and peak tok/s/GPU 1948.78 -> 2262.06 against the iter1 frontier. Measured -# acceptance is 3.39-3.49 of a 4.0 ceiling and essentially flat in batch size, -# so the gain decays with concurrency (2.29x at con=1, 1.16x at con=256) -# entirely because the step gets more expensive, not because drafting gets -# worse. **The curve is still climbing at 3** — 3 is the top of what the two -# MLA entries certify (predicted_tokens_per_seq <= 4), not an optimum. -# -# Note the throughput numbers above are measured on the harness's random-prompt -# workload, which *flatters* MTP: only the prompt is random, and the model's -# own continuation of it is repetitive and drafts near-perfectly. The real-text -# figure, from the paired gsm8k gate runs, is 1.178x. - -# The parallel topology. It is what selects this target at all -- routing -# derives the target from it -- so every variant carries it and these files -# are usable as-is: `trtllm-eval --extra_llm_api_options `. -tensor_parallel_size: 4 -moe_expert_parallel_size: 4 -enable_attention_dp: true - -speculative_config: - decoding_type: MTP - max_draft_len: 3 - -# **This variant does not boot at trtllm's default KV-cache sizing.** Measured -# on this checkpoint at dep4 on a 183.4 GiB B200, with max_draft_len 3: -# -# The MTP module adds 5.87 GiB of declared weights per rank (nvidia-smi read -# 122,076 MiB after model init), and the engine then sizes the KV pool from -# what is free, at -# the default free_gpu_memory_fraction 0.9: 38.95 GiB (1,171,200 tokens, -# 62 layers). Boot then reaches CUDA-graph capture at 182.4 of 183.4 GiB and -# livelocks: three of four ranks spin in cudaFree inside the CUDA caching -# allocator (release_cached_blocks, reached from a plain empty_cuda in the -# *trunk's* MoE runner) while the fourth waits at an MPI_Barrier. No error, -# no OOM exception, no progress -- the run has to be killed. -# -# The drafting forward's post-pool transient is simply larger than the ~19.6 -# GiB the identity config leaves and lives inside. Freeing memory *before* -# the pool is sized does not help: chunking the MTP layer's expert call to -# 2048 rows moved pool sizing by 0.31 GiB (38.64 -> 38.95) and the boot -# failed identically, because the pool grows into whatever is freed. -# -# 0.75 hands the pool ~32.5 GiB instead, leaving ~28 GiB of headroom. That is -# still ~975k KV tokens per rank -- an order of magnitude more than any -# concurrency this target is measured at needs -- so nothing about the -# workload is constrained by it. It is a boot requirement of the variant, not -# a tuning choice, and it is the reason these files carry a second key at all. -kv_cache_config: - free_gpu_memory_fraction: 0.75 diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/trtllm-ref-boot.yaml b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/trtllm-ref-boot.yaml deleted file mode 100644 index 9641a0700523..000000000000 --- a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/trtllm-ref-boot.yaml +++ /dev/null @@ -1,30 +0,0 @@ -# Boot-forcing config for the stock-trtllm reference curve ONLY. -# -# Stock trtllm (in-tree DeepseekV3 modeling) cannot boot this checkpoint at -# dep4 on a 178.34 GiB B200 under its own defaults. Measured, per rank: -# -# max_num_tokens 8192 (default): 113.9 GiB is allocated at engine -# construction outside the torch allocator, before any weight is -# materialized (nvidia-smi 20 MiB -> 116,662 MiB inside one 5 s sample), -# and model init then OOMs in init_meta_tensor. -# max_num_tokens 2048: weights load (135.85 GiB torch, incl. 1.53 GiB of -# CUDA-graph pools), then configure_kv_cache_capacity asks for 7.00 GiB -# with 5.61 GiB free. kv_cache_config.free_gpu_memory_fraction does not -# move that 7.00 GiB (measured identical at 0.9 and 0.6). -# -# max_seq_len is the remaining lever the failure's own memory ledger names: -# it sets max_blocks_per_seq = max_seq_len / tokens_per_block, 5120 at the -# checkpoint's 163840 and 64 here, and 2048 is exactly the harness's ISL+OSL. -# -# Not a tuning variant for this target: staircase boots at stock defaults, and -# both knobs measured inert on it (max_seq_len 2048: -0.13% / +0.10%). - -# The parallel topology. It is what selects this target at all -- routing -# derives the target from it -- so every variant carries it and these files -# are usable as-is: `trtllm-eval --extra_llm_api_options `. -tensor_parallel_size: 4 -moe_expert_parallel_size: 4 -enable_attention_dp: true - -max_num_tokens: 2048 -max_seq_len: 2048 diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/trtllm-ref-mtp3.yaml b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/trtllm-ref-mtp3.yaml deleted file mode 100644 index 9e726cc9df75..000000000000 --- a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/configs/trtllm-ref-mtp3.yaml +++ /dev/null @@ -1,75 +0,0 @@ -# Stock-trtllm reference line for the MTP variant. NOT a tuning variant. -# -# This is `trtllm-ref-boot.yaml` + `mtp3.yaml`, in one file because the perf -# harness takes a single --config. It exists to answer one question that -# nothing else can: **is this target's MTP layer computing what the -# checkpoint says?** -# -# Rejection sampling makes that question invisible to every other gate. A -# miscomputed draft layer still emits the target model's exact distribution — -# it just has every draft rejected, so the boot gate passes, the accuracy gate passes, -# and the only thing that moves is speed. The acceptance rate is the sole -# detector, and a rate is only readable against a reference. That reference is -# stock in-tree modeling on the same checkpoint, at the same workload and the -# same speculative_config. -# -# Why the boot knobs ride along: stock trtllm cannot boot this checkpoint at -# dep4 under its own defaults (see trtllm-ref-boot.yaml for the measured -# ledger), and MTP adds ~5.9 GiB of layer-61 weights per rank on top of that. -# The knobs change batching capacity, not what the model predicts — -# **acceptance is a property of the model**, so a reference forced onto -# different memory knobs is still the right comparison. Only the workload has -# to match, and it does: bench/perf.py seeds prompts by concurrency value, so -# both systems see byte-identical requests. - -# The parallel topology. It is what selects this target at all -- routing -# derives the target from it -- so every variant carries it and these files -# are usable as-is: `trtllm-eval --extra_llm_api_options `. -tensor_parallel_size: 4 -moe_expert_parallel_size: 4 -enable_attention_dp: true - -max_num_tokens: 2048 -max_seq_len: 2048 - -# Must match configs/mtp3.yaml exactly — a reference at a different draft -# length answers a different question. -speculative_config: - decoding_type: MTP - max_draft_len: 3 - -# The staircase side needs 0.75 to boot with MTP at all (configs/mtp3.yaml has -# the livelock ledger). Stock trtllm carries strictly more memory pressure -# here, so it starts from the same value rather than from the default. -kv_cache_config: - free_gpu_memory_fraction: 0.75 - -# ONE MORE THING IS NEEDED AND IT IS NOT A CONFIG KEY. -# -# Export `TRTLLM_CAN_USE_DEEP_EP=0` before launching, or this does not boot: -# -# _torch/modules/fused_moe/communication/deep_ep_low_latency.py:238, dispatch -# assert hidden_states.dtype == torch.uint8 -# AssertionError -- all four ranks, "Failed to initialize executor" -# -# Under attention DP + EP the MoE communication factory picks a strategy by -# priority. Measured here: NVLinkOneSided, NVLinkTwoSided and DeepEP each log -# `not available: Invalid Argument`, leaving DeepEPLowLatency -- whose dispatch -# accepts only NVFP4 (uint8-packed) hidden states. The MTP layer is **bf16**, -# because modelopt excludes `model.layers.61*` from quantization wholesale. -# TARGET.md's existing `trtllm` curve went through DeepEPLowLatency happily, -# because without MTP every MoE layer is NVFP4. -# -# The variable disables DeepEP and DeepEPLowLatency together and lands on -# AllGatherReduceScatter, which the selector's own comment calls "always -# works" -- and which is the strategy this target implements by hand, so it -# makes the two systems more numerically comparable rather than less. -# -# It must be exported **before** the server subprocess is spawned: OpenMPI -# hands a spawned process the environment as it stood when MPI initialized. -# bench/perf.py prepares the server's environment and then spawns, so an -# export in the calling shell reaches every rank. -# -# Consequence to state wherever this label is used: its **throughput** is not -# comparable to the `trtllm` Pareto curve, which carries neither this -# transport nor this draft length. This config measures acceptance. diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py index 97bc07b765cd..664546b53113 100644 --- a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py +++ b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py @@ -14,113 +14,24 @@ are bf16 and only the MLP path is NVFP4 — **and the KV cache is fp8-e4m3** (`hf_quant_config.json`: `kv_cache_quant_algo: FP8`). -Against the `deepseek-v3-lite-nvfp4/dep4` sibling (not in this batch) — same architecture -family, same topology, same quantization shape — five things differ, and each -is re-derived here rather than inherited: - -* **the query path is a LoRA pair.** `q_lora_rank: 1536`, so the query is - `q_b_proj(q_a_layernorm(q_a_proj(x)))` — two GEMMs and a norm — where the - lite checkpoint's `q_lora_rank: null` projects directly. Nothing downstream - changes: `q_b_proj` produces the same `[T, H*(nope+rope)]` rows; -* **routing is group-limited.** `n_group 8`, `topk_group 4`: `noaux_tc_op` - scores 8 contiguous groups of 32 experts by their two best bias-corrected - scores, keeps the best 4 groups, and takes top-8 inside them. The lite - sibling runs the ungrouped case. The MoE runner's own `n_group`/`topk_group` - stay inert — routing is done before the call, on either checkpoint; -* **the rope table is YaRN-scaled** (factor 40 over an original 4096-position - window, `beta_fast` 32, `beta_slow` 1), and the model's YaRN attention - temperature `mscale = 0.1*ln(40) + 1` rides in **`q_scaling = 1/mscale^2`** - rather than in the table: the table's own amplitude is - `m(mscale)/m(mscale_all_dim)` = exactly 1.0 because this config sets both to - 1.0. The op reads the table's *content* and `q_scaling`; the seven scalar - rope arguments beside them are inert (measured — see thop_attention.md); -* **the latent KV pool is fp8-e4m3**, which changes what every one of the five - MLA-family calls a layer issues does — see below; -* **scale.** 128 query heads (attention DP replicates them, so every rank runs - all 128), 64 routed experts per rank, and a 163840-position rope table. +The query path is a LoRA pair (`q_lora_rank: 1536`), routing is group-limited +(`n_group 8`, `topk_group 4`), and the rope table is YaRN-scaled with the +attention temperature carried in `q_scaling` rather than in the table. **Layer 61 — the checkpoint's bf16 MTP module — is a second forward path this -file also carries, and it exists only when a `configs/` variant turns it on.** -Under the target's identity config (`llm_args.yaml`, no `speculative_config`) -nothing below MTP is declared, layer 61's 790 keys stay a predicted non-load in -the weight manifest, and the shell's forward is a plain call to the inherited -base — bit-identical to the assembly the accuracy gate was measured on. With -`configs/mtp{1,2,3}.yaml` the engine resolves a `spec_config` onto the model -config before the model is built, the core declares the module's parameters, -and the shell builds a draft-model container plus the runtime's spec worker. -`MTPLayer` below is the module's forward; `docs/models/multi-token-prediction.md` -is its semantics (the checkpoint ships no reference implementation of it and -neither does transformers), and `docs/references/trtllm-runtime-integration.md` -§13 is the runtime binding. Two consequences the rest of this file carries: -`predicted_tokens_per_seq` is per call site rather than inert, because a -generation request under MTP arrives carrying its whole draft chain; and the -`spec_decoding_*` group stays inert, because on a trtllm-gen arch a -linear-tree draft has -its mask machinery forced off and drafting reaches the attention ops through -`predicted_tokens_per_seq` alone. - -**The fp8 latent pool, and the four things it moves.** Nothing validates the -fp8 round trip at any layer: the write scale, the read scale and the two -folded FMHA scales are independent roles with no relation checked anywhere, so -a mistake here is silently mis-scaled output rather than an error. - -* **the context append and the cache gather** take the KV scaling factor `s` - as a write-side `1/s` and a read-side `s`. This checkpoint's 122 per-layer - `k_scale`/`v_scale` tensors are all exactly 1.0 (loaded and asserted in - `derive_after_load`), and `s = 1.0` is the **only** correct value on the fp8 - MLA context path — both context flavors quantize q/k/v at 1.0 while applying - `s^2`/`s` as if they had not, so any other `s` is silently wrong. Both scale - arguments are therefore passed as `None`, which the ops read as exactly 1.0 - and which is what the engine's own call sites pass; -* **the context FMHA quantizes its own q/k/v to e4m3**, in both flavors. So - prefill accuracy *is* affected by cache quantization, and a cached-prefix - context call pays fp8 twice — the gather dequantizes the cached latent rows - off the pool and the FMHA quantizes the up-projected result straight back; -* **the decode producers are ordered, not concurrent.** `mla_rope_generation` - does not write `fused_q` here — it **reads** `fused_q[..., :C]` to build - `quant_q_buffer`, the query the decode FMHA actually consumes. The absorbed-q - BMM must have finished before it. This forward issues both on the ambient - stream in that order, which is what makes it safe; the bf16 reading (the two - producers write disjoint halves and may overlap) is a silent race here; -* **the decode FMHA reads neither kv scale tensor.** Its query comes from - `quant_q_buffer` and its two scales from `mla_bmm1_scale[1]` and - `mla_bmm2_scale[0]`, both written by `mla_rope_generation` from `q_scaling`, - the MLA dims and the read-side factor. The caller owns the dequantization - entirely. - -**The parallel segment.** `dep4` is `tensor_parallel_size: 4` plus -`moe_expert_parallel_size: 4` plus `enable_attention_dp: true`: the requests -are split, not the heads. - -* **attention, the q-LoRA pair, both dense MLP shapes, the shared expert, the - router, the embedding, the norms and the residual stream are replicated**, - and each rank runs them over its own tokens only — 128 query heads per rank. - A rank's `o_proj` output is complete, so there is no attention-side - collective at all; -* **the MoE stays expert-parallel** (`moe_tp_size == 1`): rank `r` holds - experts `[64r, 64r+64)` whole. Every token must reach every window, so the - rank's tokens are **gathered** before the router and the four windows' - partials are **reduce-scattered** back — one `comm/allgather` and one - `comm/reducescatter` per MoE layer, 116 per forward, none in layers 0-2; -* **the gather goes before the router GEMM, and that placement is - load-bearing.** Expert parallelism rests on the four windows tiling the - routing space exactly once, which needs every rank to select the same experts - for the same token. Ranks hold different tokens here, so the invariant is - restored by routing on the gathered full token set: identical bytes in, - replicated deterministic router GEMM and `noaux_tc_op`, identical top-8 out. - Routing locally and gathering afterwards would break it with no error; -* **the reduce-scatter is crossed in bf16**: the op sums, and it sums - `float8_e4m3fn` as raw bytes rather than as floats, so a post-quantization - return trip is silently wrong (the gather is a byte move and would survive - it — the asymmetry is the trap); -* **`lm_head` is replicated** — the shell builds the whole `[vocab, hidden]` on - every rank under attention DP, because a rank's logits rows are its own - tokens' and no other rank computed them. - -The expert call is **chunked** to `_MOE_MAX_T` rows: the gathered token set -reaches `4 * max_num_tokens` = 32768 and both `fp4_block_scale_moe_runner` and -`noaux_tc_op` are certified to 8192. That bound is a certification boundary, -not a tuning knob. +file also carries, and it exists only when `speculative_config` turns it on.** +Under the target's identity config nothing below MTP is declared, layer 61's +790 keys stay a predicted non-load in the weight manifest, and the shell's +forward is a plain call to the inherited base. `MTPLayer` below is the +module's forward; the checkpoint ships no reference implementation of it and +neither does transformers, so what it computes is stated where it is built. + +`dep4` is `tensor_parallel_size: 4` plus `moe_expert_parallel_size: 4` plus +`enable_attention_dp: true`: the requests are split, not the heads. Everything +but the MoE is replicated and runs over a rank's own tokens, so there is no +attention-side collective; the MoE stays expert-parallel, rank `r` holding +experts `[64r, 64r+64)` whole, which costs one `comm/allgather` and one +`comm/reducescatter` per MoE layer. Weights are target-owned: a flat ParameterDict declared here (HF [out, in] storage so checkpoint rows copy in unchanged, plus the kernel-ready expert @@ -130,7 +41,11 @@ DecoderModelForCausalLM for lm_head, packed-batch logits gathering, and the meta-init/load/post-load hooks. -The import-time and first-forward contract checks below fail fast on drift. +Everything this forward relies on that is not visible in the call itself is +stated at the call: which scales are read and which are inert, why the two +decode producers are ordered rather than concurrent, why the gather precedes +the router GEMM, why the reduce-scatter is crossed in bf16, and which bounds +are certification boundaries rather than tuning knobs. """ import math @@ -196,35 +111,31 @@ _SM = (10, 3) -def _check_static_contract() -> None: - """Import-time fail-fast: op symbol existence. The list is the forward's - trtllm call set plus `block_scale_interleave`, which the load-time - expert/scale relayout in weights.py depends on.""" - for op in ( - "cublas_mm", - "bmm_out", - "nvfp4_gemm", - "flashinfer_rmsnorm", - "flashinfer_fused_add_rmsnorm", - "flashinfer_silu_and_mul", - "fp4_quantize", - "noaux_tc_op", - "fp4_block_scale_moe_runner", - "fused_moe", - "mla_rope_generation", - "mla_rope_append_paged_kv_assign_q", - "load_paged_kv_cache_for_mla", - "allgather", - "reducescatter", - "block_scale_interleave", - ): - assert hasattr(torch.ops.trtllm, op), f"missing op trtllm::{op}" - from tensorrt_llm.bindings.internal import thop - - assert hasattr(thop, "attention"), "missing pybind thop.attention" - - -_check_static_contract() +#: Every trtllm op this target reaches for, in its forward and in the weight +#: load. Declared here, asserted in tests/unittest/_torch/staircase: a symbol +#: that does not exist is a fact of the build, and the place to find that out +#: is a machine with the extension built rather than every import of this +#: module. +REQUIRED_TRTLLM_OPS = ( + "cublas_mm", + "bmm_out", + "nvfp4_gemm", + "flashinfer_rmsnorm", + "flashinfer_fused_add_rmsnorm", + "flashinfer_silu_and_mul", + "fp4_quantize", + "noaux_tc_op", + "fp4_block_scale_moe_runner", + "fused_moe", + "mla_rope_generation", + "mla_rope_append_paged_kv_assign_q", + "load_paged_kv_cache_for_mla", + "allgather", + "reducescatter", + # Not a forward call: the load-time expert/scale relayout in weights.py + # depends on it, so it belongs to the same contract. + "block_scale_interleave", +) # Metadata fields consumed each step (sourcing mirrors the in-tree # FallbackFmha for this trtllm version; existence checked at first forward). @@ -259,10 +170,10 @@ def _check_static_contract() -> None: "use_spec_decoding", "flash_mla_tile_scheduler_metadata", "flash_mla_num_splits", - # Added between 1.3.0rc21 and 1.3.0rc26. Both are engine-prepared - # per-instance constants (max_num_sequences defaults to max_num_requests; - # the tree-mask flag is set from is_spec_dec_dynamic_tree, and this - # target's MTP is a linear tree), so they project like the rest. + # Both are engine-prepared per-instance constants (max_num_sequences + # defaults to max_num_requests; the tree-mask flag is set from + # is_spec_dec_dynamic_tree, and this target's MTP is a linear tree), so + # they project like the rest. "max_num_sequences", "force_prepare_spec_dec_tree_mask", ) @@ -392,14 +303,13 @@ def _build_step_args(md: TrtllmAttentionMetadata) -> dict: num_sparse_topk=None, flash_mla_tile_scheduler_metadata=None, flash_mla_num_splits=None, - # Added between 1.3.0rc21 and 1.3.0rc26; held at the values the op had - # before they existed, which are also what the in-tree caller passes on - # this path. kv_norm_weight is not merely a default: non-None would fold + # Held at the op's defaults, which are also what the in-tree caller passes + # on this path. kv_norm_weight is not merely a default: non-None would fold # the kv_a_layernorm into the KV kernel and make it read latent_cache # RAW, and this target normalizes that itself -- passing the weight would # normalize twice. skip_correction is a lossy trtllm-gen MLA option - # (SM100/SM103, off by default upstream); enabling it is a configs/ - # variant's business, not the identity assembly's. + # (SM100/SM103, off by default upstream); enabling it is a caller's + # business, not the identity assembly's. kv_norm_weight=None, kv_norm_eps=1e-6, skip_correction_threshold=0.0, @@ -444,7 +354,7 @@ def _build_step_args(md: TrtllmAttentionMetadata) -> dict: # The MTP layer's `eh_proj` operand order, at one place so flipping it is a # one-line experiment. True = `concat(enorm(e), hnorm(h))`, the embedding block # first — measured on this checkpoint's own weights by two independent -# statistics (docs/models/multi-token-prediction.md), which contradicts the +# statistics, which contradicts the # DeepSeek-V3 report's `M_k[RMSNorm(h); RMSNorm(e)]` notation and agrees with # the parameter's name. Nothing in the checkpoint's metadata pins it and # nothing downstream detects a flip: the drafts are simply rejected, so the @@ -562,7 +472,7 @@ def __init__(self, model_config: ModelConfig): self.mtp_layers = int(getattr(cfg, "num_nextn_predict_layers", 0) or 0) # Whether this engine drafts. The checkpoint decides the *mode* # (`num_nextn_predict_layers: 1` -> MTP-Eagle one-model, one layer - # replayed `max_draft_len` times); a `configs/` variant's + # replayed `max_draft_len` times); the caller's # `speculative_config` decides whether it runs at all, and the engine # has already resolved that onto `model_config.spec_config` by the time # the model is built. With it absent — the target's identity config — @@ -666,19 +576,18 @@ def __init__(self, model_config: ModelConfig): self.shared_inter = cfg.moe_intermediate_size * cfg.n_shared_experts self.dense_inter = cfg.intermediate_size assert 0 < self.topk < self.num_experts, "MoE top-k bound" - # noaux_tc_op is the whole gate: in-kernel sigmoid, bias correction - # for selection only, group-limited selection, renormalization and the - # routed_scaling_factor multiply. A checkpoint with norm_topk_prob - # false cannot use it. + # noaux_tc_op does the whole gate in-kernel and is not configurable: + # sigmoid, bias correction for selection only, renormalization, the + # routed_scaling_factor multiply. A checkpoint that wants any of those + # differently cannot use it. assert cfg.topk_method == "noaux_tc", cfg.topk_method assert cfg.scoring_func == "sigmoid", cfg.scoring_func assert cfg.norm_topk_prob, "noaux_tc_op always renormalizes" self.n_group = cfg.n_group self.topk_group = cfg.topk_group - # Group-limited routing, which this checkpoint uses and the lite - # sibling does not. noaux_tc_op's grouped configuration carries four - # hard limits of its own; each is checked here because the op reports - # a violation as one opaque "unsupported configuration". + # Four hard limits of the op's grouped path, each checked here because + # the op reports any violation as one opaque "unsupported + # configuration". assert self.n_group > 1 and 1 <= self.topk_group <= self.n_group, ( f"grouped routing needs 1 <= topk_group <= n_group, got " f"{self.topk_group} / {self.n_group}" @@ -792,7 +701,7 @@ def P(*shape, dtype=dt): w["final_norm"] = P(self.hidden) w["embed"] = P(self.vocab, self.hidden) # The MTP module at layer index `num_hidden_layers`, declared only when - # a configs/ variant turned drafting on. Its attention block is + # `speculative_config` turned drafting on. Its attention block is # byte-identical in geometry to a trunk layer's; its MLP path is the # same structure at a different **dtype** — `hf_quant_config.json` # carries `model.layers.61*` as one wildcard entry in its @@ -2109,7 +2018,7 @@ def __init__(self, model_config: ModelConfig): hidden_size=cfg.hidden_size, vocab_size=cfg.vocab_size, ) - # The speculative branch, built only when a configs/ variant asked for + # The speculative branch, built only when `speculative_config` asked for # it. `spec_config` is already resolved on the model config by the time # the model is built, and the checkpoint — not the target — picked the # mode: `num_nextn_predict_layers: 1` gives MTP-Eagle one-model, where diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/weights.py b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/weights.py index 78e1ecdce657..ed36ef94bc0e 100644 --- a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/weights.py +++ b/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/weights.py @@ -34,7 +34,7 @@ * Under the target's identity config all 790 are an expected non-load, listed rather than swept under a relaxed assert. -* Under a `configs/mtp*.yaml` variant the module is loaded, and this table +* Under an MTP variant (`speculative_config` set) the module is loaded, and this table gains its rows: **212 keys per rank** are consumed (the whole front end and attention block, the router, the shared expert, and this rank's 64-expert window), leaving 578 — the 192 off-window experts' 576 weight tensors, which diff --git a/tensorrt_llm/_torch/staircase/models/gpt_oss/routing.py b/tensorrt_llm/_torch/staircase/models/gpt_oss/routing.py index 8a1407d0af19..08c2be6c0fe0 100644 --- a/tensorrt_llm/_torch/staircase/models/gpt_oss/routing.py +++ b/tensorrt_llm/_torch/staircase/models/gpt_oss/routing.py @@ -21,8 +21,8 @@ # Config-shape fingerprint -> checkpoint identity. Sniffing the shape is the # upstream idiom (``is_mla``, ``is_nemotron_hybrid`` do the same). It buys # automatic routing at a stated cost: a *fine-tune* of this checkpoint has the -# same shape and is routed here silently. See TARGET.md -- the gate record is -# what pins the identity, and it says which checkpoint it was measured on. +# same shape and is routed here silently. The accuracy gate is what pins the +# identity: it names the checkpoint this target was measured on. # # (num_hidden_layers, hidden_size, num_local_experts). Layer count alone # separates 120b from 20b, but the expert count is what makes the MoE operand diff --git a/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/TARGET.md b/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/TARGET.md deleted file mode 100644 index ebbf5a0caf70..000000000000 --- a/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/TARGET.md +++ /dev/null @@ -1,480 +0,0 @@ -# Target: gpt-oss-120b / sm_103 / tp1 - -## Identity - -| | | -|---|---| -| Checkpoint | gpt-oss-120b (HF safetensors; bf16 attention/router/embedding/lm_head, **MXFP4 experts** — E2M1 blocks + per-32 E8M0 scales; untied embeddings; sparse MoE on every layer: 128 experts, top-4, renormalized; **attention sinks**; alternating sliding-window / full attention) | -| GPU arch | sm_103 (GB300) | -| Parallel | tp1 | -| Registered class | `StaircaseGptOss120bSm103Tp1` — a synthetic architecture name no checkpoint declares. `models/gpt_oss/routing.py` rewrites `GptOssForCausalLM` into it when the config, SM and topology all match; the checkpoint is read unpatched. Per-target names mean one process can hold every target at once | - -> **NO GATE RECORD HOLDS FOR THIS TARGET.** Every result below was -> measured on **sm_100 (B200)** through the pre-move standalone harness. This -> target is **sm_103 (GB300)**, and certification is per architecture. The -> numbers are kept as provenance — they are true records of what the same -> modeling code did on another device — but this target is **ungated** until -> the boot and gsm8k runs in *Verification* are repeated on GB300 and their -> results replace those rows. Read every "passed" below as "passed, on -> sm_100, before the move". - -Checkpoint sha256 — the checkpoint directory passed to `--model` should -resolve to files with these digests. **Routing does not check them**: it -fingerprints the config's shape, so a fine-tune of this checkpoint routes -here silently. That is the deliberate trade in `models/gpt_oss/routing.py`, -and it changes what a gate record means — not "this target passed" but -"this modeling code passed *on the checkpoint with these digests*". Run it -on another one and the result is ungated (recorded at: -`umbriel-b200-027:/home/scratch.trt_llm_data/llm-models/gpt_oss/gpt-oss-120b`): - -``` -695218884684c611fe08a74751ee443f971e9bd9bc062edba822da3fe45969b7 model-00000-of-00014.safetensors -a881aa5f561b26a22b14a8262aa61849ace349ffd73d74769e030ac90a1fcf8a model-00001-of-00014.safetensors -022478dd04398c5bdb545a5be0a6437ecc2eb53d1dbd29edafcfff4b3ddf0a41 model-00002-of-00014.safetensors -47aee9e7b9d5bedb215042c01ccededd9bd9c30b0dddea862dc2506b9d6c74de model-00003-of-00014.safetensors -f6c2752acda607b1d5ca52df9e75c1b9b2761e6875ff10c9bd6ddac473c0262e model-00004-of-00014.safetensors -0c8dd401544c31cb93b8459eee7da20ea2a07626a59455d7d92b85257df9b46c model-00005-of-00014.safetensors -28d839f2e027985a8b14e45f2323798862eddb7770ee9800ea6b7c803abee489 model-00006-of-00014.safetensors -c8958c5f183c04f6ea959cfd90562b5128124154b2bbf979b8a22b9405b30ed8 model-00007-of-00014.safetensors -bf1f2a88868ffc37d520dcf77d26f0e823710b5e682d473ff10f6974fa3b7517 model-00008-of-00014.safetensors -f72d34a4004241b45c332b61f8ffa124e9a913bc1ab442b66e717d3e94e741ce model-00009-of-00014.safetensors -f48c867c2cb0a44bfc2f8768cb98e4aec9a350946fceacfebdcad5d32ad4a471 model-00010-of-00014.safetensors -a06851b2cfd35f48722f823bc1ab8f7bcb4a878a5b8e975f4d3544f230454eeb model-00011-of-00014.safetensors -3af33667c307e20ae2a7648ea52653de46dd0171601ec5c696e47a2f5d5bf1e4 model-00012-of-00014.safetensors -bcbcb74b043e071d1e05471d500d74dcf661175e00878ed302ccdf1801a75aef model-00013-of-00014.safetensors -54b1be1609696c307cc5ca117b1fa54feaddebffa04e9c2db117652a01964230 model-00014-of-00014.safetensors -ede2655fdc05008561983b6e0829c600727c28d591e071077377059f03a6c00e model.safetensors.index.json -0614fe83cadab421296e664e1f48f4261fa8fef6e03e63bb75c20f38e37d07d3 tokenizer.json -9279e942392b742d633c7adbb89ebe002c98399db8926a7af5125c726f404070 tokenizer_config.json -dd5e191d20c12d2fee1da5bae14ca1db0f5f4215300af691f23cdee97120a293 special_tokens_map.json -f8d9255777615591a7cc1a7c932f5a69e181128902295e1b81221d20d983cac7 chat_template.jinja -7bfd294f3e29b53db1e126d5cec050a12dc27adad4445fb5eab540e1cad74ea1 chat_template.json -199566674b96510c3b9a1141b494223a86a3ff83097e2c2f259c4d94fafc5847 generation_config.json -``` - -Two link-set notes specific to this checkpoint. Its chat template lives in -`chat_template.jinja` (`tokenizer_config.json` carries none), and the -accuracy gate's protocol applies that template — so the template files and -`special_tokens_map.json` are linked alongside `tokenizer*`. And its -`config.json` declares **no dtype at all**, which used to be patched around -with a target-owned `model_dir/config.json` declaring `dtype: bfloat16`. - -That stub is gone — the checkpoint is now read exactly as published — so the -divergence it papered over is live code. `DecoderModelForCausalLM` sizes -`lm_head` from `pretrained_config.torch_dtype`, which is `None` here, and -would materialize it in the torch default fp32 while every other tensor is -bf16; the failure surfaces two layers from its cause. The shell fills that -gap explicitly before `super().__init__`, adopting the dtype the engine -already resolved (`ModelConfig.torch_dtype`, bf16) and only when the -checkpoint declares none. Every safetensors tensor outside the MXFP4 expert -blocks is bf16, so that is what the checkpoint is. - -## Version - -| | | -|---|---| -| tensorrt_llm | in-tree — the target moves with the trunk, so there is no version to pin and none is asserted. What *is* asserted at construction is the SM version (`_SM = (10, 3)`), which the pin used to stand in for. The gate records below name the commit they were taken at | -| torch | 2.11.0+cu130 | -| transformers | 5.5.4 (the config surface the engine hands the target; see the rope note in `modeling.py`) | -| Attention metadata fact source | `TrtllmAttentionMetadata` (TRTLLM backend) | - -## Vocabulary - -Forward: `flashinfer_rmsnorm`, `flashinfer_fused_add_rmsnorm`, -`cublas_mm` (fused bias — qkv, o, router all carry one), -`fused_qk_norm_rope` (`is_qk_norm=False`: YaRN RoPE only), -`thop_attention` (per-layer `attention_sinks` and `attention_window_size`), -`mxfp8_quantize`, `mxe4m3_mxe2m1_block_scale_moe_runner`, -`torch/embedding`, `torch/empty`, `torch/reshape`. - -The MoE call is the whole expert block — routing, both grouped GEMMs, the -clamped GLU, the MXFP8 requantization between the GEMMs and the combine — -so no `activation/*` and no `moe/*routing*` entry appears: this checkpoint -has no dense MLP, and the router bias rides the router GEMM's fused-bias -epilogue because the MoE op silently ignores `routing_bias` at -`routing_method_type=1`. - -**Changed by the perf campaign (`iter1-mxfp8-moe`, see *Performance*).** -The assembly shipped the W4A16 member of this kernel family -(`bf16_mxe2m1_block_scale_moe_runner`) with `torch/pad` widening hidden -2880 → 3072 in front of it; the target now runs the W4A8 member over -MXFP8 activations, and `mxfp8_quantize` does that widening inside itself, -so the `torch/pad` call is gone. Weight preparation is unchanged — the two -ops consume the identical prepared expert stack. - -Weight loading fills parameters via the manifest loop (`torch/copy_` -semantics); `.t()` views are cublas_mm's column-major consumption form, -derived once post-load. The expert operands are the exception: the -manifest's source transforms rebuild the checkpoint's MXFP4 blocks and -E8M0 scales into the MoE op's kernel-ready layout (pad → `[up ; gate]` -concat → row interleave → 32-row block shuffle → 128x4 scale swizzle), -promote both expert biases and the attention sinks from bf16 to fp32, and -run on device one layer at a time. - -Audit is mechanical: grep the forward's calls against `catalog/index.yaml`. - -## Verification - -### Required on sm_103 — not yet run - -| Gate | Command | Result | -|---|---|---| -| boot | `TRTLLM_STAIRCASE=require python examples/llm-api/quickstart_advanced.py --model_dir --max_tokens 16 --prompt "The capital of France is" "The chemical symbol for gold is" "1, 2, 3, 4, 5,"` | **passed, 10/10** greedy keyword asserts, 2026-09-09, GPU 0 of nvl72d001-T18 (NVIDIA GB300, 284208 MiB, sm_103), trtllm 1.3.0rc26. 65.70 GiB of weights loaded; decode CUDA graphs captured at batch sizes 1-32, 64, 128; engine boot to last token 3m29s. The run had the switch set to `require`, so the built-in GptOss implementation could not have been substituted. Measured with a per-target `smoke.py` that asserted a keyword in each of ten greedy continuations; that file was removed in favour of the generic script above, which boots the same way and prints the same continuations for a human to read | - -**Execution is verified on sm_103; the accuracy gate is not yet.** Everything up to the -weight load is driven by the checkpoint's config alone, and that part -was exercised on a GB300 against the config *as published* (no -target-owned stub): routing resolved this target, the module imported, -`StaircaseCore.__init__` passed every geometry, topology and dtype -assert, and 543 parameters declared, 36 layers, hidden 2880. `lm_head.weight` -came out **bfloat16**, which is the specific thing removing the stub -put at risk -- the shell sizes it from the pretrained dtype, and a -regression there materializes fp32 two layers from its cause. - -The weight path is verified too, by the boot run above: the manifest -load fills every declared parameter (its coverage asserts are part of -the gate), the post-load derivations run, and the forward produces -coherent greedy continuations through both the prefill and the -CUDA-graph decode path. - -The checkpoint it ran on is the one this file records. Every digest -above was re-verified after download -- the six small files by hand and -the fifteen safetensors shards by git-lfs, whose object id *is* the -sha256 -- so this gate record and the sm_100 records below were measured -on byte-identical weights. - -| gsm8k full | `TRTLLM_STAIRCASE=require trtllm-eval --model gsm8k --output_path --apply_chat_template --fewshot_as_multiturn --max_output_length 8192`, or in CI as `accuracy/test_staircase.py::TestStaircaseGptOss120bSm103Tp1::test_gsm8k` | **passed, 90.6748** (`exact_match,flexible-extract`, +-0.8010, full 1319 questions) against threshold **85.5989** (anchor `openai/gpt-oss-120b` = 90.5989, tol 5.0) -- pass by 5.08 points. 2026-09-09, GPU 0 of nvl72d173-T18 (GB300), trtllm 1.3.0rc26, 2m54s wall. `strict-match` on the same run: 27.0660 | - -`TRTLLM_STAIRCASE=require` is what makes the second command a gate at all: under -`auto` a configuration that missed this target would measure the built-in -GptOss implementation and report it as this target's score. - - -#### Two notes on reading the gsm8k number - -**`--check_accuracy` is not used, and the filter is read by hand.** The gate -filter is `exact_match,flexible-extract`, and `trtllm-eval` exposes no CLI flag -for `scores_filter` -- it is a keyword of the evaluator's `evaluate()`. Left -unset the evaluator *averages* the filters, which for this checkpoint mixes -90.6748 with a `strict-match` of 27.0660 and reports 58.87. That average is -meaningless here: this is a reasoning model, its answer never arrives in the -strict `#### N` form, so the strict filter scores it near the floor. Read the -flexible-extract row of the logged table, which is also saved to -`--output_path`. - -**The anchor in `references/accuracy.yaml` was deliberately not written back.** -The rule in that file is to write a target's first passing score back -- but it -also says to skip the write-back when doing so would confuse what the anchor -means. It would here: the anchor is keyed by *checkpoint* (`openai/gpt-oss-120b`) -and currently holds an sm_100 measurement of the W4A16 assembly, deliberately -left high so the gate stays the stricter of the two. This measurement is a third -thing -- the W4A8 forward on sm_103 -- and folding it into a checkpoint-keyed -anchor would make that anchor architecture-dependent without saying so. The -number lives here, where gate records are per target and therefore per -architecture. - -For the record, the three measurements of this checkpoint through this harness: -sm_100 W4A16 **90.5989** (the anchor), sm_100 W4A8 **89.16**, sm_103 W4A8 -**90.6748**. The last is +1.51 over the sm_100 W4A8 record, which is ~1.9 sigma -on this run's +-0.80 stderr and is **not** claimed as an improvement -- the MoE -FC1 epilogue genuinely uses a different block-scale recipe on sm_103 (see -`catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md`), so a difference of this -size has a plausible mechanism, but separating it from session variance would -need repeated runs on both architectures. - -### Prior record — sm_100 (B200), pre-move harness, does not gate this target - -Both gates passed under trtllm defaults (block reuse and CUDA graphs on, -`llm_args.yaml` empty, single-process worker via `scripts/env.sh`), -2026-07-26/27, GPU 0 of umbriel-b200-027 (NVIDIA B200, driver 595.58.03): - -| Gate | Result | -|---|---| -| smoke — `uv run targets/gpt-oss-120b/sm_100/tp1/smoke.py` | **passed**, 10/10 greedy keyword asserts (keywords frozen against continuations observed on this model; the arithmetic case is Q/A-framed because a bare `2 + 2 =` is genuinely ambiguous here) | -| gsm8k full — `uv run bench/accuracy.py --target targets/gpt-oss-120b/sm_100/tp1` | **passed**, `measured 89.16 >= 85.6`, exit 0 (1319 questions, 5-shot, chat template + few-shot-as-multiturn + 8192 output tokens, filter `exact_match,flexible-extract`) | - -**The row above is the shipped forward** — the W4A8 MoE the perf campaign -left in place. The gate it clears is the written-back anchor -`openai/gpt-oss-120b` = 90.5989 (`source: trtllm-eval`, tol 5.0 ⇒ -threshold 85.6): **pass by 3.56 points**. - -The anchor itself was measured on the *assembly's* W4A16 forward, and is -deliberately not re-written down to the W4A8 number — leaving it high -keeps the gate the stricter of the two, and it remains a real measurement -of this checkpoint through this harness. The W4A16 record it came from, -kept because the anchor derives from it: `measured 90.60 >= 85.3` against -`reference: openai/gpt-oss-120b = 90.3 (trtllm)`, exit 0, same protocol — -reproduced bit-identically by two independent invocations, the -assembler's and the orchestrator's verification re-run. - -Measured GSM8K on that W4A16 forward, full 1319 questions (2026-07-27): - -| filter | score | -|---|---| -| **flexible-extract (gated)** | **90.5989 ± 0.8039** | -| strict-match (not gated) | 25.4738 ± 1.2002 | - -Reference `openai/gpt-oss-120b` = 90.3 (`source: trtllm`), tol 5.0 ⇒ -threshold 85.3: **pass by 5.30 points**, and 0.30 *above* the anchor -itself. Stock trtllm on this checkpoint under the identical protocol -measured 90.2199 here during the onboard's anchor phase. The strict-match -filter is not the gated metric and is not comparable: it scores the -`#### N` surface form a harmony-format model never emits. - -**Rerun noise measured on this target, protocol-specific.** The same build -under the identical protocol scored 90.2199 in a first run and 90.5989 in -the gating run — 1190 vs 1195 of 1319, a 0.38-point spread with no code -change between them (strict-match swung further, 22.59 → 25.47). GSM8K -with 8192-token reasoning generations is therefore noisier than the -~0.1-0.3 the repo's guidance cites for completion MMLU: treat anything -under ~0.5 points on this gate as unresolved. - -Engine facts observed on those runs: model init 14.1-14.5 s (63 GB of -checkpoint, including the on-device expert relayout); one KV pool sized -for full attention (99.27 GiB, 1,445,664 tokens, `tokens_per_block` 32, -`window size=131072`), `host_kv_cache_pool_mapping` `[36, 2]` with -identity rows `[0, l]` — the runtime is not told about the 128-token -sliding layers, which is exactly the certified single-pool route, and the -alternation therefore saves no memory. Both FMHA kernel families appear at -warmup (`...H64PagedKvDense...` for the full layers, -`...H64PagedKvSlidingOrChunkedCausal...` for the sliding ones), which is -the per-layer window taking effect. Under trtllm defaults -`cache_reuse=True`, so the engine prepares `use_paged_context_fmha=True` -and the target passes it through as documented; both gates above ran with it -True. That value was **not certified** when this target merged — the -contract's certified column said `False`, and this line was the only -record of the mismatch anywhere. It was closed on 2026-07-28 by a -certification extension, which also established that the target's -behaviour was the correct one: at `False` a context call with a cached -prefix returns normally and lands 51x outside the tolerance band, while -with nothing cached the flag is bitwise inert. The paged read does add a -caller obligation — the pages in range must be valid and distinct — which -this target satisfies by passing the engine's own offsets through. - -Pre-gate check of the highest-risk axis, run before the first engine boot -(scratch script, not a repo product): layer 0's real MXFP4 expert tensors -prepared by `weights.py` and fed to the catalog MoE call land **1.12 bf16 -ulp** (of the token row's largest magnitude) from a pure-torch HF -reference — dequantized with the transformers mxfp4 semantics, gate/up -split as `[..., ::2]` / `[..., 1::2]`, clamped GLU, top-4 renormalized -routing — while the gate/up-swapped reference lands **146.36 ulp**. That -is the interleave-parity trap, pinned by measurement rather than by -reading. - -## Performance - -![Serving Pareto](perf/figures/pareto.png) - -Environment — every curve below: GPU 0 of `umbriel-b200-027` (NVIDIA B200, -driver 595.58.03, CUDA 13.0), one device for the whole campaign, -2026-07-27; `tensorrt_llm 1.3.0rc21`, `torch 2.11.0+cu130`, -`transformers 5.5.4`, single-process worker via `scripts/env.sh`. Sweeps: -`bench/perf.py`, ISL=OSL=1024, concurrency 1→256, requests per point = -concurrency × rounds (20 for con≤8, else 5). `baseline` 07:30-08:10 UTC, -`trtllm` 08:10-08:38 UTC (back to back), `iter1-mxfp8-moe` 13:29-13:55 UTC. -Perf is recorded, never gated. - -### Curves - -`peak tok/s/gpu` is the maximum over the sweep; the concurrency that -reaches it is in parentheses. - -| label | config | commit | accuracy | con=1 tok/s/user | peak tok/s/gpu | change | -|---|---|---|---|---|---|---| -| `trtllm` | `llm_args.yaml` only (`{}`) | `c5e6d53` | ungated reference | 394.66 | 7962.9 (con=128) | stock trtllm modeling, stock defaults | -| `baseline` | `llm_args.yaml` only (`{}`) | `c5e6d53` | 90.5989 | 228.74 | 4476.7 (con=256) | staircase at trtllm defaults, W4A16 MoE | -| `iter1-mxfp8-moe` | `llm_args.yaml` only (`{}`) | `41f955a` | **89.1585** | 380.35 | **8421.1 (con=128)** | MoE swapped to the W4A8 (MXFP8-activation) member of the same kernel family | - -No config variant was kept, so every curve runs the identity config and -`configs/` does not exist. Both accuracy cells are full 1319-question -gsm8k under the target's own protocol; each reproduced bit-identically -across two independent `bench/accuracy.py` invocations. - -Full point-by-point (tok/s/gpu, ISL=OSL=1024): - -| con | trtllm | baseline | iter1 | iter1 vs baseline | iter1 vs trtllm | iter1 tpot ms | iter1 ttft ms | -|---|---|---|---|---|---|---|---| -| 1 | 394.6 | 228.7 | 380.3 | +66.3% | −3.6% | 2.61 | 22.3 | -| 2 | 690.0 | 324.0 | 669.7 | +106.7% | −3.0% | 2.96 | 35.0 | -| 4 | 1181.2 | 527.8 | 1146.2 | +117.2% | −3.0% | 3.45 | 47.6 | -| 8 | 1856.8 | 749.9 | 1866.2 | +148.9% | +0.5% | 4.22 | 76.6 | -| 16 | 2799.8 | 1099.6 | 2872.8 | +161.3% | +2.6% | 5.47 | 109.2 | -| 32 | 4028.0 | 1434.6 | 4186.2 | +191.8% | +3.9% | 7.50 | 155.9 | -| 64 | 5851.3 | 2015.0 | 6066.8 | +201.1% | +3.7% | 10.34 | 218.0 | -| 128 | 7962.9 | 2760.4 | 8421.1 | +205.1% | +5.8% | 14.89 | 306.6 | -| 256 | 5509.8 | 4476.7 | 7267.3 | +62.3% | +31.9% | 34.74 | 455.2 | - -### Kept iterations - -**iter1 — the MoE call from W4A16 to W4A8 (MXFP8 activations).** - -*Evidence.* Four `nsys` steady-state windows (100 executor iterations, -every one pure decode) said the W4A16 forward was GPU-saturated at every -concurrency — idle −10.7% / −1.1% / −0.8% at con=1 / 128 / 256 (negative = -slight cross-stream overlap) — and that 60.8% / 91.9% / 89.8% of that GPU -time was the MoE call. Per layer per step at con=128 its two grouped GEMMs -cost 711.8 µs and 352.9 µs against the stock reference's 126.7 µs and -65.6 µs — **5.62× and 5.38× on byte-identical expert weights**. The -reference resolves this checkpoint's `quant_method: mxfp4` to -`W4A8_MXFP4_MXFP8`; our kernel's name carries `castBfloat16` (MXFP4 -expanded to bf16 for a `m128x8x16` bf16 MMA at 3 pipeline stages, 1-CTA -clusters) while the reference's feeds MXFP4 straight into a block-scaled -`m256x16x32` MMA at 6 stages, 2-CTA clusters. Non-MoE GPU work already -matched (3.41 vs 3.10 ms/step), so the MoE was the entire gap. - -*Change.* One catalog call swapped for one: `mxfp8_quantize(o, False, 512)` -produces e4m3 data plus per-32 UE8M0 **linear** block scales, and -`mxe4m3_mxe2m1_block_scale_moe_runner` consumes the pair. The quantizer -performs the hidden widening 2880 → 3072 itself, so the `torch/pad` that -fed the W4A16 call is gone — the forward issues 530 kernels per decode step -instead of 566. `weights.py` is untouched: both ops read the identical -prepared expert stack. The scale layout is spelled out rather than -defaulted, because the 128×4 swizzled buffer has the same byte count -whenever `num_tokens % 128 == 0` — every decode CUDA graph of 128 or 256 — -and is then accepted as a silently wrong answer. - -*Effect.* Every point improves, from +62.3% (con=256) to +205.1% -(con=128); peak throughput +88.1% (4476.7 → 8421.1 tok/s/gpu) and con=1 -+66.3% (228.74 → 380.35 tok/s/user). TTFT at con=1 falls 54.1 → 22.3 ms. -The curve also changes shape: it now peaks at con=128 and turns over at -256, exactly as the stock reference does. - -*Accuracy cost — a result, not noise.* gsm8k **89.1585**, gate passed by -3.56 points (threshold 85.6), but **1.44 points below** the W4A16 -measurement of 90.5989 recorded above. Two independent invocations of -`bench/accuracy.py` returned 89.1585 bit-identically (stderr 0.8564 both -times), as the W4A16 path reproduces 90.5989 bit-identically — so this is -the recipe's price, well outside the ~0.5-point band this gate leaves -unresolved. The mechanism is documented in the op's contract: FC1's -epilogue requantizes the activation to MXFP8 on the OCP scale (`floor` of -the block exponent, block max saturating) before FC2 reads it. For -reference, stock trtllm running this same recipe measured 90.2199 here -during the onboard. - -### Gap decomposition - -`trtllm-tuned` is **absent by construction**: no config variant was kept, -so it would be byte-identical to `trtllm` and the portable-config share of -the gap is 0. The decomposition below is therefore kernel-level, from -`nsys` windows of 100 pure-decode iterations at con=128: - -| | baseline (W4A16) | iter1 (W4A8) | trtllm | -|---|---|---|---| -| step (wall) | 41.84 ms | **16.10 ms** | 16.18 ms | -| GPU active | 42.30 ms | 11.60 ms | 10.60 ms | -| GPU idle | −1.1% | **+27.9%** | +34.5% | -| MoE total | 38.89 ms | 8.34 ms | 7.51 ms | -| non-MoE | 3.41 ms | 3.26 ms | 3.10 ms | -| kernels / step | 566 | 530 | 502 | - -The target's decode step is now **16.10 ms against the reference's -16.18 ms**, and the bottleneck has moved: this forward was 100% GPU-bound -and is now 27.9% idle, i.e. host-bound in the same regime as the stock -reference. What remains at the kernel level is an 11% MoE-GEMM difference -from tactic selection alone — the autotuner picks `t128x8x512_s3` / -`m128x8x32` / 1-CTA for us (139.3 and 74.1 µs per layer) against the -reference's `t128x16x256u2_s6` / `m256x16x32` / 2-CTA (126.7 and 65.6 µs) -on the same op family. The new quantize call costs 143.12 µs per step -(36 calls, 3.98 µs each) — 1.2% of GPU-active time. - -### Reference lines - -- `trtllm` = the original HF checkpoint under **stock in-tree trtllm - modeling**, sharing this target's `llm_args.yaml` (identity config, `{}` - at tp1) and otherwise trtllm defaults — `bench/perf.py --trtllm`. Since - `iter1`, both systems run the same W4A8 numerical recipe, so this is now - a like-for-like comparison; against `baseline` it was not (that curve is - W4A16). -- `trtllm-tuned` is absent by construction (no kept config variant). -- The reference is **ungated** — no accuracy gate is run against in-tree - modeling. Pinned to `tensorrt_llm 1.3.0rc21`. - -### Measurement caveats - -- The host CPU of `umbriel-b200-027` is shared and other tenants ran GPU - work on neighbouring devices during the campaign; GPU 0 was reserved - throughout, but host and chassis contention is an error bar on every - point. -- **Session variance at mid-curve concurrencies exceeds the 3% rule of - thumb.** Under the W4A16 MoE, con=128 at rounds=2 measured 3097.2, - 3101.0 and 2725.2 tok/s across three sessions of configurations later - shown equivalent (12.1% spread) against 0.35% at con=1 and 0.77% at - con=256. Under the W4A8 MoE the same point was far steadier: two control - probes 26 minutes apart measured 8615.3 and 8603.4 tok/s (0.14%). Probe - A/Bs in this campaign therefore always carry a same-session control. -- Wrap-up spot-check of the reference (`--trtllm`, con=1/32/256, default - rounds, 5 h after the label): 394.0 (−0.15%), 4267.2 (+5.94%), 5607.6 - (+1.77%). The `trtllm` label was **not** re-swept: the +5.94% sits inside - the session variance above, and re-measuring one curve alone would break - the back-to-back pairing with `baseline`. -- Profiled runs are never comparable to clean ones: under `nsys` the - W4A16 con=128 point measured 2468.4 tok/s against 2760.4 clean. - -Figure regeneration: - -``` -uv run utils/plot_pareto.py targets/gpt-oss-120b/sm_100/tp1/perf/data \ - -o targets/gpt-oss-120b/sm_100/tp1/perf/figures/pareto.png -``` - -### Tried and rejected (with the number that rejected it) - -The first four were measured **against the W4A16 forward**, when the GPU -was saturated; that scope matters, because the first one changed verdict -after iter1 and had to be re-tested. - -| hypothesis | axis | measured | verdict | -|---|---|---|---| -| decode CUDA-graph coverage 256 (`cuda_graph_config.max_batch_size: 256`, `enable_padding: true`) — **under W4A16** | config | con=1 +0.04%, con=128 +0.12%, con=256 +0.77% | inert: no idle existed to recover | -| `max_seq_len: 12288` (blocks/seq 4096 → 384, confirmed in `server.log`) | config | con=1 −0.3%, con=256 +0.07% | inert: the per-step H2D block-offset staging measures 4.19 MB / 87 µs = 0.2% of a con=128 step | -| `cute_dsl_bf16_gemm_blackwell` for the qkv / o / router GEMMs | modeling | graph-replay kernel time at M=1: 7.12 vs cublas+bias 5.95 µs (qkv), 7.30 vs 6.28 (o), 5.61 vs 5.18 (router) | rejected: slower at every shape before the bias cublas fuses | -| FC1 K-padding 3072 → 2944 (−4.17% of FC1 weight bytes) | modeling | MoE µs/layer: T=128 1176.3 vs 1151.4 (worse), T=256 1177.0 vs 1191.9 | rejected: no consistent gain | -| decode CUDA-graph coverage 256 — **re-tested under W4A8** | config | see below | **trade-off, not shipped** | - -**The graph-coverage trade-off, re-measured after iter1.** Once the MoE -stopped saturating the GPU, the con=256 point began to collapse the way the -stock reference's does (8421.1 at con=128 → 7267.3 at con=256). Capturing a -batch-256 decode graph recovers it, but costs the con=128 point. Same -session, control probed twice before and after (rounds=2, con=128/256): - -| config | con=128 | con=256 | -|---|---|---| -| identity (control, 14:04) | 8615.3 | 7221.6 | -| identity (control, 14:30) | 8603.4 | 7516.3 | -| `max_batch_size: 256`, padding on | 7897.7 (−8.3%) | 9805.3 (+33.1%) | -| `max_batch_size: 256`, padding off | 7753.2 (−10.0%) | 9792.2 (+32.9%) | - -Padding is not the mechanism — exact-hit coverage regresses con=128 just as -much — and neither is memory: the KV pool moves 99.31 → 99.15 GiB (0.16%) -and graph memory 10.83 → 11.28 GiB. Decode TPOT at con=128 rises 14.35 → -15.97 ms with the extra graph present. The keep rule is regress-nowhere, so -this is **not shipped**; a deployment that only ever serves con≥256 should -adopt it deliberately, since its con=256 point (9805.3 tok/s at 38.6 -tok/s/user) is outside anything the shipped config reaches. - -### Remaining headroom - -- **The target is now host-bound at the throughput end** (27.9% GPU idle at - con=128), in the same regime as stock trtllm (34.5%) and for the same - reason: stock executor Python between steps. No knob in the tuner's table - reaches it. -- **~11% of the MoE GEMM time is tactic selection**, not recipe: same op, - same weights, different autotuner choice than the in-tree path makes - (`t128x8x512_s3`/1-CTA vs `t128x16x256u2_s6`/2-CTA). Worth a look at what - drives the autotuner's bucket set. -- **The con=256 graph-coverage trade-off above** is a real +33% at the - throughput end blocked by a −9% at con=128 whose mechanism is not - established. Establishing it would unlock the point. -- **Decode cost still saturates with batch**, measured on the W4A16 op per - layer: T=1 67.3 µs, T=64 986.9, T=128 1151.4, T=256 1191.9, T=1024 - 1307.3 — +13.5% for 8× the tokens above T=128. The expert-weight read is - a fixed per-step cost once the batch touches all 128 experts. -- **Per-window KV pools are not the win** and that thread is retired: the - pool is sized for full attention (99.27 GiB, 1,445,664 tokens = 705 - requests at ISL+OSL 2048) but capacity never binds at con≤256 (36.3% - used), the per-layer window already limits what the FMHA kernel reads, - and stock trtllm sizes the same single pool on this checkpoint. -- **The accuracy cost of W4A8 (−1.44 points) is the price of the frontier - above.** A deployment that needs the last point of gsm8k accuracy should - run the W4A16 forward (`baseline`, commit `c5e6d53`) and accept a third - of the throughput. diff --git a/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py b/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py index 66f4184dcf2e..b9f456a00239 100644 --- a/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py +++ b/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py @@ -74,23 +74,19 @@ _SM = (10, 3) -def _check_static_contract() -> None: - """Import-time fail-fast: op symbol existence.""" - for op in ( - "fused_qk_norm_rope", - "flashinfer_rmsnorm", - "flashinfer_fused_add_rmsnorm", - "cublas_mm", - "mxfp8_quantize", - "mxe4m3_mxe2m1_block_scale_moe_runner", - ): - assert hasattr(torch.ops.trtllm, op), f"missing op trtllm::{op}" - from tensorrt_llm.bindings.internal import thop - - assert hasattr(thop, "attention"), "missing pybind thop.attention" - - -_check_static_contract() +#: Every trtllm op this target reaches for, in its forward and in the weight +#: load. Declared here, asserted in tests/unittest/_torch/staircase: a symbol +#: that does not exist is a fact of the build, and the place to find that out +#: is a machine with the extension built rather than every import of this +#: module. +REQUIRED_TRTLLM_OPS = ( + "fused_qk_norm_rope", + "flashinfer_rmsnorm", + "flashinfer_fused_add_rmsnorm", + "cublas_mm", + "mxfp8_quantize", + "mxe4m3_mxe2m1_block_scale_moe_runner", +) # Metadata fields consumed each step (sourcing mirrors the in-tree # FallbackFmha for this trtllm version; existence checked at first forward). @@ -131,9 +127,8 @@ def _check_static_contract() -> None: "num_sparse_topk", "flash_mla_tile_scheduler_metadata", "flash_mla_num_splits", - # Added between 1.3.0rc21 and 1.3.0rc26. Both are engine-prepared - # per-instance constants (max_num_sequences defaults to - # max_num_requests; the tree-mask flag is set from + # Both are engine-prepared per-instance constants (max_num_sequences + # defaults to max_num_requests; the tree-mask flag is set from # is_spec_dec_dynamic_tree), so they project like the rest. "max_num_sequences", "force_prepare_spec_dec_tree_mask", @@ -248,9 +243,8 @@ def _build_step_args(md: TrtllmAttentionMetadata) -> dict: cross_kv=None, relative_attention_bias=None, relative_attention_max_distance=0, - # Added between 1.3.0rc21 and 1.3.0rc26; held at the values the op had - # before they existed. All three are MLA-only surface that this target - # does not use -- skip_correction is forced to 0.0 for a non-MLA layer by + # Held at the op's defaults. All three are MLA-only surface that this + # target does not use -- skip_correction is forced to 0.0 for a non-MLA layer by # the engine's own resolver, and kv_norm_* folds an MLA kv_a_layernorm # that does not exist here. kv_norm_weight=None, diff --git a/tensorrt_llm/_torch/staircase/references/accuracy.yaml b/tensorrt_llm/_torch/staircase/references/accuracy.yaml deleted file mode 100644 index 57c1a3c18842..000000000000 --- a/tensorrt_llm/_torch/staircase/references/accuracy.yaml +++ /dev/null @@ -1,190 +0,0 @@ -# Accuracy reference scores, keyed by Hugging Face repo name; target_ckpt -# joins an entry to the segment of a target's identity path, -# models//targets////. -# -# Gate: measured accuracy >= score - tol, one-sided, enforced by -# trtllm-eval --check_accuracy. tol is 5.0 everywhere — a reference score -# drifts a few points across checkpoint variants, harnesses and -# quantizations, while an assembly catastrophe sits 20+ points away -# (random = 25 on MMLU). -# -# source, preferred first: -# `trtllm` — TensorRT-LLM's accuracy suite -# (tests/integration/defs/accuracy/). Best source: the -# harness we delegate to, and the test that produced the -# score also states its protocol (see `eval`). A -# checkpoint lists one entry per quantization — take the -# one this target matches, not the first line. -# `model-card` — the checkpoint's published number. -# `trtllm-eval` — measured here, e.g. stock trtllm on the same -# checkpoint when nothing is published. -# After a target's first passing run, write its measured score back and -# set source to `trtllm-eval`; tol stays 5.0. Skip the write-back when the -# target was accepted with a protocol caveat (no passing value exists), or -# when the anchor is already this repo's own stock-trtllm measurement — -# replacing that with the target's number makes the gate self-referential. -# -# `eval` (optional) — the protocol, so our number and the anchor measure -# the same thing. Absent = trtllm-eval defaults: 5-shot completion MMLU -# scored on the first generated token. That fits a base model; an instruct -# model without its chat template, or a reasoning model with no room to -# reason, scores far below its ability and the gate then reads as an -# assembly defect. Copy the protocol from the test that gave the score. -# -# eval: -# task: gsm8k # trtllm-eval subcommand; default mmlu -# args: # passed through by name to that subcommand -# apply_chat_template: true -# fewshot_as_multiturn: true -# max_output_length: 8192 -# -# args pass through by name: any option the subcommand accepts (booleans -# become flags, mappings are JSON-encoded). -# -# `scores_filter` is a sibling of `args`, not one of them: it names which -# lm-eval filter the gate reads (e.g. `exact_match,flexible-extract`). -# trtllm-eval exposes no CLI flag for it — it is a keyword of the -# evaluator's evaluate() — so a caller has to bind it before delegating. -# Set it whenever the anchor's protocol names one. Unset, the evaluator -# averages every metric the task reports, which mixes filters measuring -# different things and, on a task whose result dict carries a string -# alias, raises instead of returning a score. A reasoning model needs it: -# its answer never arrives in the strict `#### N` form, so the strict -# filter scores it near the floor while flexible extraction reads what it -# actually answered. -# -# ── THE SCORES BELOW ARE ANCHORS, NOT GATE RECORDS ──────────────────────── -# -# Every `score` here that says `source: trtllm-eval` was written back from a -# passing run on **sm_100 (B200)**, through a standalone harness that did not -# survive the move in-tree. The migrated targets are sm_103 (GB300) and have -# no passing run yet. These values are still the right bar to gate against — -# an accuracy anchor is a property of the checkpoint and the protocol, not of -# the GPU — but no target in this tree has cleared one. See each TARGET.md. -# -# The replacement invocation is trtllm-eval directly, with the protocol below -# spelled out on the command line and `TRTLLM_STAIRCASE=require` exported, so -# that a configuration which does not match a target fails loudly instead of -# quietly measuring the built-in implementation. -# ────────────────────────────────────────────────────────────────────────── - -# TensorRT-LLM's accuracy suite, references/gsm8k.yaml key -# `openai/gpt-oss-120b`, read at commit 1662a877f374ee944d1907e8efa735d35ff2abf6 -# whose tensorrt_llm/version.py is 1.3.0rc21 — the pin exactly. The -# quantization-keyed twin `GPT-OSS/120B-MXFP4` carries the same 90.3, -# including its `quant_algo: W4A16_MXFP4` row, which is this checkpoint's -# quantization. The suite has no MMLU entry for gpt-oss at all and gates -# the family on GSM8K only: the model reasons before it answers, and -# 5-shot completion MMLU scored on the first generated token measures -# none of that. -# Protocol copied from TestGPTOSS (test_llm_api_pytorch.py), whose -# MODEL_PATH is this exact checkpoint: chat template, few-shot as -# multiturn, MAX_OUTPUT_LEN patched to 8192, and flexible answer -# extraction. Stock trtllm measured here under precisely that protocol -# scored 90.22 (strict-match, for contrast: 28.13) — the anchor -# reproduces on this machine, so it was gated against directly rather -# than through a stock-trtllm proxy. -# -# The score below is the target's own first passing measurement, 90.5989 -# over the full 1319 questions, which replaces the 90.3 it was gated -# against. Measured twice by separate invocations of `bench/accuracy.py` -# — the assembler's gating run and an independent orchestrator re-run — -# reproducing bit-identically, down to the non-gated strict-match filter -# (25.4738). Read the +0.30 over the external anchor and the +0.38 over -# the stock-trtllm number as agreement, not as an improvement: the same -# build measured 90.22 through a differently-invoked driver, so -# differences of this size on this benchmark are not results. -# -# That measurement is of the assembly's W4A16 MoE forward. The perf -# campaign then shipped the W4A8 (MXFP8-activation) member of the same -# kernel family, which gates at 89.1585 — also bit-identical across two -# invocations, so the 1.44-point drop is the recipe's price and not -# sampling. The anchor is deliberately left at the higher W4A16 value: -# lowering it would only loosen the gate, and 90.5989 remains a real -# measurement of this checkpoint through this harness. Open thread for -# whoever revisits: stock trtllm runs the same W4A8 recipe and measured -# 90.2199 here, so our W4A8 sits 1.06 below it — larger than this gate's -# 0.38-point cross-driver spread. -openai/gpt-oss-120b: - target_ckpt: gpt-oss-120b - score: 90.5989 - n_samples: 1319 - source: trtllm-eval - tol: 5.0 - eval: - task: gsm8k - scores_filter: exact_match,flexible-extract - args: - apply_chat_template: true - fewshot_as_multiturn: true - max_output_length: 8192 - -# TensorRT-LLM's accuracy suite, references/gsm8k.yaml key -# `deepseek-ai/DeepSeek-R1-0528`, the `quant_algo: NVFP4` + -# `kv_cache_quant_algo: FP8` row = 94.24, read at commit -# 1662a877f374ee944d1907e8efa735d35ff2abf6 whose tensorrt_llm/version.py -# is 1.3.0rc21 — the pin exactly. That row is the checkpoint's exact -# quantization pair: this target's hf_quant_config.json declares -# `quant_algo: NVFP4` with `kv_cache_quant_algo: FP8`, so it is the row -# to take rather than the block's first line by position (they coincide -# here). -# -# Why GSM8K and not MMLU. references/mmlu.yaml does carry a -# `deepseek-ai/DeepSeek-R1-0528` block, but it holds **only -# FP8_BLOCK_SCALES rows — there is no NVFP4 row at all**, so this -# checkpoint's quantization has no MMLU anchor to be gated against. GSM8K -# is also the right task on its own terms: R1-0528 is a reasoning model, -# and 5-shot completion MMLU scored on the first generated token measures -# none of that. -# -# Protocol: the consuming test is TestModelRegistryAccuracy's -# `test_autodeploy_from_registry` (test_llm_api_autodeploy.py), whose -# `nvidia/DeepSeek-R1-0528-NVFP4-v2` parameter is aliased onto this key -# and whose task list is `[GSM8K]`. It passes `evaluate_kwargs = {}` for -# every non-MMLU task, so the measurement is accuracy_core.GSM8K's -# defaults — full 1319 questions, random_seed 0, **no chat template**, no -# fewshot-as-multiturn, max_input_length 4096, max_output_length 256. -# trtllm-eval's `gsm8k` subcommand defaults to exactly those, so no -# `args` block is needed. Being a reasoning model does not change that: -# under plain few-shot completion the model continues the exemplars' -# form rather than opening a think block, which is what makes 94.24 -# reachable inside 256 output tokens. -# -# `scores_filter` is this repo's one deliberate deviation, for the reason -# the file header gives — the suite leaves it None and averages every -# metric, which this harness cannot do. Same choice as the gpt-oss-120b -# and DeepSeek-V3-Lite entries. Report both filters on the first run. -# -# CAVEAT, and the reason this entry is worth re-reading at write-back -# time: the anchor was measured on a **different NVFP4 export of the same -# base model**. The test's `nvidia/DeepSeek-R1-0528-NVFP4-v2` resolves -# (tests/test_common/llm_data.py:44) to -# `DeepSeek-R1/DeepSeek-R1-0528-FP4-v2`, while this target is built from -# `DeepSeek-R1/DeepSeek-R1-0528-FP4` — the v1 export, which sits beside -# it on disk. Both declare NVFP4 + FP8 KV, but their modelopt -# `exclude_modules` lists differ substantially (63 entries here against -# 246 in v2), so they are not the same weights. tol 5.0 is what absorbs -# an export-revision difference of this kind; a measured score a point or -# two off 94.24 is not evidence of an assembly defect. -# -# WRITE-BACK. The score below is the target's own first passing -# measurement, 94.9962 over the full 1319 questions, replacing the 94.24 it -# was gated against. `strict-match` on the same run was 94.6171; the 0.38 -# between the filters is 5 questions, so under few-shot completion this -# checkpoint mostly answers in the strict `#### N` form and flexible -# extraction finds a few more. -# -# Read the +0.76 over the external anchor as agreement, not as an -# improvement. One filter's stderr alone is +-0.60, and the anchor was -# measured on the other export — the caveat above is exactly why a -# difference of this size carries no information. Nothing in the target's -# record rests on it. -deepseek-ai/DeepSeek-R1-0528: - target_ckpt: deepseek-r1-0528-nvfp4 - score: 94.9962 - n_samples: 1319 - source: trtllm-eval - tol: 5.0 - eval: - task: gsm8k - scores_filter: exact_match,flexible-extract diff --git a/tensorrt_llm/llmapi/llm.py b/tensorrt_llm/llmapi/llm.py index fdda58cbbf31..04ede393189f 100644 --- a/tensorrt_llm/llmapi/llm.py +++ b/tensorrt_llm/llmapi/llm.py @@ -393,6 +393,13 @@ def __init__(self, f"Unknown backend: {backend!r}. Supported backends are " "'pytorch'.") + # TRTLLM_STAIRCASE=require promises the run measured a staircase + # target. Only the pytorch backend reaches the resolver that could + # select one, so on any other backend the promise would be broken + # silently -- the one failure that mode exists to prevent. + from .._torch.staircase import assert_backend_can_route + assert_backend_can_route(backend) + # check the kwargs and raise ValueError directly valid_keys = set( list(llm_args_cls.model_fields.keys()) + diff --git a/tests/integration/defs/accuracy/test_staircase.py b/tests/integration/defs/accuracy/test_staircase_deepseek_v3.py similarity index 67% rename from tests/integration/defs/accuracy/test_staircase.py rename to tests/integration/defs/accuracy/test_staircase_deepseek_v3.py index 4520a51e82f0..c0b56df14680 100644 --- a/tests/integration/defs/accuracy/test_staircase.py +++ b/tests/integration/defs/accuracy/test_staircase_deepseek_v3.py @@ -12,25 +12,25 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. -"""Whole-model gates for the staircase targets. +"""Whole-model gates for the deepseek-v3 staircase targets. -Separate file rather than entries in test_llm_api_pytorch.py, for the same -reason the targets are separate codebases: these gate a parallel -implementation, and reading them next to the built-in model's tests would -invite treating one as a variant of the other. +One file per model family, beside the other accuracy suites rather than inside +test_llm_api_pytorch.py: these gate a parallel implementation, and reading them +next to the built-in model's tests would invite treating one as a variant of +the other. Every test here needs ``TRTLLM_STAIRCASE=require``. Under ``"auto"`` a configuration that missed a target's criteria would quietly fall back to the built-in implementation, pass, and report the built-in's numbers as the target's -- which is the one failure this whole system exists to prevent. The -one exception is the stock leg of the acceptance gate, which asks for -``"off"`` on purpose. +one exception is the stock leg of the acceptance gate, which asks for ``"off"`` +on purpose. The switch is an environment variable, and worker ranks read it as it stood -when they started. At world size 1 that is this process. Above it the ranks -are already running by the time a test body executes, so a multi-rank case -cannot choose its own mode -- it can only assert that the environment it was -given is the one it needs, which is what ``_require_mode`` does. +when they started. These targets are multi-rank, so the ranks are already +running by the time a test body executes: a case cannot choose its own mode, it +can only assert that the environment it was given is the one it needs, which is +what ``_require_mode`` does. """ import os @@ -56,6 +56,12 @@ get_sm_version() != 103, reason="staircase targets in this batch are certified on sm_103 only" ) +# The filter the staircase anchors were measured on. Unset, the evaluator +# averages every metric GSM8K reports, which means the mean of strict-match and +# flexible-extract -- two numbers measuring different things on a checkpoint +# that does not answer purely in the strict "#### N" form. +_SCORES_FILTER = {"scores_filter": "exact_match,flexible-extract"} + def _require_mode(expected: str) -> None: """Skip unless the ranks were started with the mode this case needs. @@ -70,56 +76,6 @@ def _require_mode(expected: str) -> None: pytest.skip(f"{STAIRCASE_ENV}={actual!r}, this case needs {expected!r}") -class _StaircaseGSM8K(GSM8K): - """GSM8K reading the filter the staircase anchors were measured on. - - Unset, the evaluator averages every metric the task reports, which for - GSM8K means the mean of ``strict-match`` and ``flexible-extract``. Those - measure different things here: neither of these checkpoints answers purely - in the strict ``#### N`` form, so the average is a number no reference was - ever taken at -- gpt-oss scores ~90 flexible, ~25 strict, and the mean of - 56 reads as a catastrophic failure of a model that is answering correctly. - """ - - EVALUATE_KWARGS = {"scores_filter": "exact_match,flexible-extract"} - - -class _GSM8KWithRoomToReason(_StaircaseGSM8K): - """The above, with the output budget a reasoning model needs. - - The stock 256 tokens truncate this checkpoint mid-chain-of-thought, before - it ever reaches an answer, and the gate then reads as an assembly defect - rather than as the protocol being wrong for the model. - """ - - MAX_OUTPUT_LEN = 8192 - - -class TestStaircaseGptOss120bSm103Tp1(LlmapiAccuracyTestHarness): - """gpt-oss-120b / sm_103 / tp1.""" - - # The registry key upstream uses for this checkpoint; it carries the - # W4A8_MXFP4_MXFP8 entry the engine resolves from its quantization_config. - MODEL_NAME = "GPT-OSS/120B-MXFP4" - MODEL_PATH = f"{llm_models_root()}/gpt_oss/gpt-oss-120b" - - # This checkpoint is gated as a reasoning model: its answer never arrives - # in the strict "#### N" form, so the protocol applies the chat template - # and gives the model room to reason. Matches the protocol recorded in - # _torch/staircase/references/accuracy.yaml. - extra_evaluator_kwargs = { - "apply_chat_template": True, - "fewshot_as_multiturn": True, - } - - @skip_not_sm103 - def test_gsm8k(self): - _require_mode("require") - with LLM(self.MODEL_PATH) as llm: - task = _GSM8KWithRoomToReason(self.MODEL_NAME) - task.evaluate(llm, extra_evaluator_kwargs=self.extra_evaluator_kwargs) - - class TestStaircaseDeepseekR10528Nvfp4Sm103Dep4(LlmapiAccuracyTestHarness): """deepseek-r1-0528-nvfp4 / sm_103 / dep4, identity and the mtp3 variant.""" @@ -131,9 +87,9 @@ class TestStaircaseDeepseekR10528Nvfp4Sm103Dep4(LlmapiAccuracyTestHarness): # against the mapping the engine actually built. DEP4 = dict(tensor_parallel_size=4, moe_expert_parallel_size=4, enable_attention_dp=True) - # configs/mtp3.yaml. The kv-cache fraction is a boot requirement of the - # variant rather than a tuning choice: the drafting forward's post-pool - # transient does not fit what the default 0.9 leaves. + # The kv-cache fraction is a boot requirement of the variant rather than a + # tuning choice: the drafting forward's post-pool transient does not fit + # what the default 0.9 leaves. MTP3 = MTPDecodingConfig(max_draft_len=3) MTP3_KV = KvCacheConfig(free_gpu_memory_fraction=0.75) @@ -150,7 +106,7 @@ class TestStaircaseDeepseekR10528Nvfp4Sm103Dep4(LlmapiAccuracyTestHarness): @skip_not_sm103 @pytest.mark.skip_less_device(4) - def test_gsm8k_identity_vs_mtp3(self): + def test_gsm8k_identity_vs_mtp3(self, mocker): """The identity accuracy gate, and the gate on MTP not moving it. Turning MTP on must not move the answers. @@ -168,7 +124,8 @@ def test_gsm8k_identity_vs_mtp3(self): recorded in some earlier session carries that session's variance into the judgement. Paired, the variance is common to both and cancels. """ - task = _StaircaseGSM8K(self.MODEL_NAME) + mocker.patch.dict(GSM8K.EVALUATE_KWARGS, _SCORES_FILTER) + task = GSM8K(self.MODEL_NAME) _require_mode("require") with LLM(self.MODEL_PATH, **self.DEP4) as llm: @@ -194,7 +151,7 @@ def test_gsm8k_identity_vs_mtp3(self): @skip_not_sm103 @pytest.mark.skip_less_device(4) @pytest.mark.parametrize("mode", ["require", "off"], ids=["staircase", "stock"]) - def test_mtp3_acceptance(self, mode): + def test_mtp3_acceptance(self, mode, mocker): """The only gate that can see a miscomputed draft layer. Rejection sampling makes a wrong draft path *slower*, not wrong: every @@ -210,8 +167,7 @@ def test_mtp3_acceptance(self, mode): rather than blaming the target The anchor is populated from the **stock** leg. Populating it from the - target's own number would make the gate self-referential, which is the - same rule references/accuracy.yaml states for accuracy anchors. + target's own number would make the gate self-referential. """ _require_mode(mode) if mode == "off": @@ -228,6 +184,8 @@ def test_mtp3_acceptance(self, mode): "it stock cannot boot this checkpoint at dep4 with MTP" ) + mocker.patch.dict(GSM8K.EVALUATE_KWARGS, _SCORES_FILTER) + with LLM( self.MODEL_PATH, speculative_config=self.MTP3, @@ -236,25 +194,8 @@ def test_mtp3_acceptance(self, mode): enable_iter_perf_stats=True, **self.DEP4, ) as llm: - task = _StaircaseGSM8K(self.MODEL_NAME) + task = GSM8K(self.MODEL_NAME) task.evaluate(llm) acceptance_length = compute_acceptance_length(llm) print(f"[AL] {mode} acceptance_length = {acceptance_length:.3f}") assert_acceptance_length(self.ACCEPTANCE_KEY, acceptance_length) - - -def test_staircase_off_is_the_default(monkeypatch): - """Unset means off, on the code path the engine actually takes. - - Cheap, GPU-free, and the thing most worth never regressing: everything in - this file rests on staircase being opt-in. - """ - from tensorrt_llm._torch.staircase import StaircaseMode - - monkeypatch.delenv(STAIRCASE_ENV, raising=False) - assert StaircaseMode.from_env() is StaircaseMode.OFF - assert os.environ.get("STAIRCASE_TARGET") is None, ( - "STAIRCASE_TARGET was retired with the move in-tree; it named a " - "target, where TRTLLM_STAIRCASE names only a mode and lets routing " - "pick the target from the configuration" - ) diff --git a/tests/integration/defs/accuracy/test_staircase_gpt_oss.py b/tests/integration/defs/accuracy/test_staircase_gpt_oss.py new file mode 100644 index 000000000000..2d0292e5dffe --- /dev/null +++ b/tests/integration/defs/accuracy/test_staircase_gpt_oss.py @@ -0,0 +1,91 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Whole-model gate for the gpt-oss staircase targets. + +One file per model family, beside the other accuracy suites rather than inside +test_llm_api_pytorch.py: these gate a parallel implementation, and reading them +next to the built-in model's tests would invite treating one as a variant of +the other. + +Every test here needs ``TRTLLM_STAIRCASE=require``. Under ``"auto"`` a +configuration that missed a target's criteria would quietly fall back to the +built-in implementation, pass, and report the built-in's numbers as the +target's -- which is the one failure this whole system exists to prevent. +""" + +import os + +import pytest + +from tensorrt_llm import LLM +from tensorrt_llm._torch.staircase import STAIRCASE_ENV +from tensorrt_llm._utils import get_sm_version + +from ..conftest import llm_models_root +from .accuracy_core import GSM8K, LlmapiAccuracyTestHarness + +# The targets assert their own SM at construction: certification is per GPU +# architecture, and a receipt from another one says nothing here. +skip_not_sm103 = pytest.mark.skipif( + get_sm_version() != 103, reason="staircase targets in this batch are certified on sm_103 only" +) + + +def _require_mode(expected: str) -> None: + """Skip unless the ranks were started with the mode this case needs. + + Not a failure: which mode a multi-rank job runs under is a property of how + it was launched, so a case that wants the other one has nothing to say. It + must not silently measure the wrong system either, which is what reading + the variable here rules out. + """ + actual = os.environ.get(STAIRCASE_ENV, "off") + if actual != expected: + pytest.skip(f"{STAIRCASE_ENV}={actual!r}, this case needs {expected!r}") + + +class TestStaircaseGptOss120bSm103Tp1(LlmapiAccuracyTestHarness): + """gpt-oss-120b / sm_103 / tp1.""" + + # The registry key upstream uses for this checkpoint; it carries the + # W4A8_MXFP4_MXFP8 entry the engine resolves from its quantization_config. + MODEL_NAME = "GPT-OSS/120B-MXFP4" + MODEL_PATH = f"{llm_models_root()}/gpt_oss/gpt-oss-120b" + + # This checkpoint is gated as a reasoning model: its answer never arrives + # in the strict "#### N" form, so the protocol applies the chat template + # and gives the model room to reason. Same protocol as TestGPTOSS in + # test_llm_api_pytorch.py, which gates this exact checkpoint. + extra_evaluator_kwargs = { + "apply_chat_template": True, + "fewshot_as_multiturn": True, + } + + @skip_not_sm103 + def test_gsm8k(self, mocker): + # Both patches are the protocol the anchor was measured under, and both + # are what TestGPTOSS applies to this same checkpoint. The stock 256 + # tokens truncate it mid-chain-of-thought, before it ever reaches an + # answer; and unfiltered, the evaluator averages strict-match with + # flexible-extract, which measure different things here -- this model + # scores ~90 flexible and ~25 strict, so the mean of 56 reads as a + # catastrophic failure of a model that is answering correctly. + mocker.patch.object(GSM8K, "MAX_OUTPUT_LEN", 8192) + mocker.patch.dict(GSM8K.EVALUATE_KWARGS, {"scores_filter": "exact_match,flexible-extract"}) + + _require_mode("require") + with LLM(self.MODEL_PATH) as llm: + task = GSM8K(self.MODEL_NAME) + task.evaluate(llm, extra_evaluator_kwargs=self.extra_evaluator_kwargs) diff --git a/tests/integration/test_lists/test-db/l0_gb300.yml b/tests/integration/test_lists/test-db/l0_gb300.yml index 7288ae3e6811..eb405d87e626 100644 --- a/tests/integration/test_lists/test-db/l0_gb300.yml +++ b/tests/integration/test_lists/test-db/l0_gb300.yml @@ -34,8 +34,10 @@ l0_gb300: # These entries are the receipts: the catalog's certification is per GPU # architecture, and this list is the sm_103 one. comm/ is excluded here and # carried by l0_gb300_multi_gpus.yml, since those two need 4 ranks. - - unittest/_torch/staircase --ignore=unittest/_torch/staircase/comm + # TIMEOUT measured, not guessed: the whole entry is 348 cases in 5m34s on one + # GB300. 30 gives five times that -- enough for a loaded node, short enough + # that a hung case does not hold this stage for its full budget. + - unittest/_torch/staircase --ignore=unittest/_torch/staircase/comm TIMEOUT (30) # Staircase whole-model gate. The op-level entry above certifies the # vocabulary; this certifies the assembly that calls it. - - accuracy/test_staircase.py::TestStaircaseGptOss120bSm103Tp1::test_gsm8k - - accuracy/test_staircase.py::test_staircase_off_is_the_default + - accuracy/test_staircase_gpt_oss.py::TestStaircaseGptOss120bSm103Tp1::test_gsm8k diff --git a/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml b/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml index 466d010790c8..17239f6fd984 100644 --- a/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml +++ b/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml @@ -28,9 +28,9 @@ l0_gb300_multi_gpus: # acceptance gate. The two acceptance legs are independent cases read against # one shared minimum -- staircase failing while stock passes means the draft # path regressed; both failing means the anchor is stale. - - accuracy/test_staircase.py::TestStaircaseDeepseekR10528Nvfp4Sm103Dep4::test_gsm8k_identity_vs_mtp3 - - accuracy/test_staircase.py::TestStaircaseDeepseekR10528Nvfp4Sm103Dep4::test_mtp3_acceptance[staircase] - - accuracy/test_staircase.py::TestStaircaseDeepseekR10528Nvfp4Sm103Dep4::test_mtp3_acceptance[stock] + - accuracy/test_staircase_deepseek_v3.py::TestStaircaseDeepseekR10528Nvfp4Sm103Dep4::test_gsm8k_identity_vs_mtp3 + - accuracy/test_staircase_deepseek_v3.py::TestStaircaseDeepseekR10528Nvfp4Sm103Dep4::test_mtp3_acceptance[staircase] + - accuracy/test_staircase_deepseek_v3.py::TestStaircaseDeepseekR10528Nvfp4Sm103Dep4::test_mtp3_acceptance[stock] # Covers tests/unittest/_torch/attention/. Two sub-trees moved in here from elsewhere # under tests/unittest/_torch/ and have no entry of their own on any list, so this entry # is what picks them up: diff --git a/tests/unittest/_torch/staircase/activation/test_staircase_flashinfer_silu_and_mul.py b/tests/unittest/_torch/staircase/activation/test_staircase_flashinfer_silu_and_mul.py index 5b2c9c5302ef..5beefaf8cdcb 100644 --- a/tests/unittest/_torch/staircase/activation/test_staircase_flashinfer_silu_and_mul.py +++ b/tests/unittest/_torch/staircase/activation/test_staircase_flashinfer_silu_and_mul.py @@ -56,3 +56,33 @@ def test_fp16_2d() -> None: for num_tokens, two_d in [(2, 8192), (1024, 4096)]: x = torch.randn(num_tokens, two_d, dtype=torch.float16, device="cuda") _check(x) + + +def test_a_misaligned_half_is_rejected_before_dispatch() -> None: + """A final dim of 24 passes the op's own check and faults the kernel. + + The op validates the *row* (`shape[-1] * itemsize % 16 == 0`, which 24 + satisfies at 48 bytes), but the kernel vectorizes over each half + separately, and at 24 the second half starts 24 bytes in -- mid-vector. + The launch then dies with `CUDA misaligned address`, which poisons the + context for everything after it rather than raising something a caller + could catch. The wrapper's precondition is what keeps that off the device, + so this asserts it raises *without* calling the op. + """ + for width in (24, 40, 56): # d * 2 bytes = 24, 40, 56 -- none a multiple of 16 + x = torch.randn(4, width, dtype=torch.bfloat16, device="cuda") + try: + flashinfer_silu_and_mul(x) + except AssertionError: + continue + raise AssertionError(f"shape[-1]={width} should have been rejected before dispatch") + + +def test_an_odd_width_is_rejected() -> None: + """There is no half to split at an odd width; the op would truncate.""" + x = torch.randn(4, 33, dtype=torch.bfloat16, device="cuda") + try: + flashinfer_silu_and_mul(x) + except AssertionError: + return + raise AssertionError("an odd shape[-1] should have been rejected") diff --git a/tests/unittest/_torch/staircase/attention/test_staircase_mla_rope_append_paged_kv_assign_q.py b/tests/unittest/_torch/staircase/attention/test_staircase_mla_rope_append_paged_kv_assign_q.py index 1b60f2bc1f97..a3953ee49fcc 100644 --- a/tests/unittest/_torch/staircase/attention/test_staircase_mla_rope_append_paged_kv_assign_q.py +++ b/tests/unittest/_torch/staircase/attention/test_staircase_mla_rope_append_paged_kv_assign_q.py @@ -946,9 +946,9 @@ def test_fp8_kv_rejects_int8_kv_cache_quant_mode() -> None: torch.cuda.synchronize() except RuntimeError as exc: raised = str(exc) - # rc26 added NVFP4 latent pools, so the rejection message now - # enumerates two accepted formats. int8 is still rejected -- - # what this test certifies -- only the wording widened. + # The op accepts fp8 and NVFP4 latent pools, so the rejection + # message enumerates two formats. int8 is rejected -- that is what + # this test certifies. assert "Only FP8 and NVFP4 KV caches are supported for now" in raised, ( f"quant_mode={QUANT_MODE_INT8_KV_CACHE} was not rejected as expected; got: {raised!r}" ) diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/_allgather_op_matrix.py b/tests/unittest/_torch/staircase/comm/_allgather_op_matrix.py similarity index 98% rename from tensorrt_llm/_torch/staircase/catalog/comm/_allgather_op_matrix.py rename to tests/unittest/_torch/staircase/comm/_allgather_op_matrix.py index be61839ad1d1..0e37668615eb 100644 --- a/tensorrt_llm/_torch/staircase/catalog/comm/_allgather_op_matrix.py +++ b/tests/unittest/_torch/staircase/comm/_allgather_op_matrix.py @@ -15,7 +15,7 @@ synchronisation interleaved with the op, the engine's stream switching, and what disagreeing on call order actually does. - CUDA_VISIBLE_DEVICES=0,1,2,3 uv run python catalog/comm/_allgather_op_matrix.py + CUDA_VISIBLE_DEVICES=0,1,2,3 python _allgather_op_matrix.py Not a pytest module, despite the `check_*` bodies. They are one fixed sequence inside a single 4-rank job rather than independent cases: each reads @@ -27,9 +27,10 @@ The collected entry point is `tests/unittest/_torch/staircase/comm/test_staircase_allgather_op_matrix.py`: -it starts this job and turns its exit code into an assertion. This half stays -in the package because the launcher re-execs it as `python -m`, and the ranks -need the package context for their relative imports. +it starts this job and turns its exit code into an assertion. Both halves are +started by file path: the launcher must not import `tensorrt_llm` (that calls +`MPI_Init`, and an MPI-initialized process cannot start `mpirun`), and the +ranks reach the catalog by absolute import, so neither needs a package. """ import os @@ -37,6 +38,7 @@ import signal import subprocess import sys +from pathlib import Path from typing import Any, Dict, List, Optional, Sequence, Tuple import torch @@ -915,10 +917,9 @@ def _run_one_rank() -> int: from mpi4py import MPI as _MPI from tensorrt_llm._torch.distributed import Distributed + from tensorrt_llm._torch.staircase.catalog.comm import allgather as entry from tensorrt_llm.mapping import Mapping - from . import allgather as entry - allgather = entry.allgather MPI = _MPI COMM = MPI.COMM_WORLD @@ -979,8 +980,7 @@ def _spawn_ranks() -> None: "-n", str(world_size), sys.executable, - "-m", - "tensorrt_llm._torch.staircase.catalog.comm._allgather_op_matrix", + str(Path(__file__).resolve()), _WORKER_FLAG, ] print(f"[launcher] {' '.join(command)}", flush=True) diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/_rank_job.py b/tests/unittest/_torch/staircase/comm/_rank_job.py similarity index 62% rename from tensorrt_llm/_torch/staircase/catalog/comm/_rank_job.py rename to tests/unittest/_torch/staircase/comm/_rank_job.py index 837c0234596e..271f976f75d0 100644 --- a/tensorrt_llm/_torch/staircase/catalog/comm/_rank_job.py +++ b/tests/unittest/_torch/staircase/comm/_rank_job.py @@ -15,23 +15,20 @@ initialized MPI, and a pytest process that has imported ``tensorrt_llm`` already has. -For the same reason the launcher is started **by file path, not by -``-m``**. ``-m`` on a module inside this package imports every parent -package on the way to it, ``tensorrt_llm`` included, and that import calls -``MPI_Init``; the launcher would then fail to start ``mpirun`` at all -(measured: it exits 1 with no output from any rank). Run by path the file -has no package context and imports nothing but torch, which is all its -launcher half needs -- the entry keeps its relative imports inside -``_run_one_rank``, and the ranks it spawns *are* started with ``-m`` so -they get one. +Everything here is started **by file path**. The launcher must not import +``tensorrt_llm`` -- that calls ``MPI_Init``, and an MPI-initialized process +cannot start ``mpirun`` at all (measured: it exits 1 with no output from any +rank) -- and the ranks it spawns reach the catalog by absolute import, so +neither half needs a package context. This tree does not have one to give: +it is tests/, not a package. """ from __future__ import annotations +import importlib.util import os import subprocess import sys -from importlib import import_module from pathlib import Path WORLD_SIZE = 4 @@ -56,22 +53,40 @@ def _devices() -> str: return ",".join(devices[:WORLD_SIZE]) +def _load_launcher_constants(launcher: Path): + """Read the entry's module-level budget without importing it by name. + + There is no dotted name to import it by -- this directory is not a + package -- and only its constants are wanted. Executing the file at module + scope is cheap and pulls in nothing but torch; the entry keeps its + tensorrt_llm imports inside the rank body for exactly that reason. + """ + spec = importlib.util.spec_from_file_location(launcher.stem, launcher) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + def run(entry: str) -> None: """Run ``__op_matrix``'s launcher over ``WORLD_SIZE`` devices.""" - module = f"{__package__}._{entry}_op_matrix" - launcher = Path(__file__).with_name(f"_{entry}_op_matrix.py") + launcher = Path(__file__).resolve().with_name(f"_{entry}_op_matrix.py") env = dict(os.environ, CUDA_VISIBLE_DEVICES=_devices()) - # The launcher re-execs itself per rank and needs this package importable - # from the ranks; by path it has no package context of its own to inherit. + + # The ranks import tensorrt_llm absolutely, and a source checkout is not + # necessarily installed. tests/unittest/_torch/staircase/comm -> repo root. repo_root = Path(__file__).resolve().parents[5] + assert (repo_root / "tensorrt_llm").is_dir(), ( + f"expected the repo root at {repo_root}, found no tensorrt_llm/ there; " + f"this file moved without its parents[] index following" + ) env["PYTHONPATH"] = os.pathsep.join(p for p in (str(repo_root), env.get("PYTHONPATH", "")) if p) # Read the budget off the entry rather than restating it: an entry that # raises its own deadline would otherwise be killed by this one first, # and the message a reader needs ("wedged") would be lost. reducescatter # runs a second, separately capped job after the main one. - entry_module = import_module(module) - timeout = entry_module.DEADLINE_S + getattr(entry_module, "WEDGE_CAP_S", 0) + _LAUNCHER_GRACE_S + constants = _load_launcher_constants(launcher) + timeout = constants.DEADLINE_S + getattr(constants, "WEDGE_CAP_S", 0) + _LAUNCHER_GRACE_S completed = subprocess.run( [sys.executable, str(launcher)], diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/_reducescatter_op_matrix.py b/tests/unittest/_torch/staircase/comm/_reducescatter_op_matrix.py similarity index 99% rename from tensorrt_llm/_torch/staircase/catalog/comm/_reducescatter_op_matrix.py rename to tests/unittest/_torch/staircase/comm/_reducescatter_op_matrix.py index 296554ba2059..3e76415d85b8 100644 --- a/tensorrt_llm/_torch/staircase/catalog/comm/_reducescatter_op_matrix.py +++ b/tests/unittest/_torch/staircase/comm/_reducescatter_op_matrix.py @@ -10,7 +10,7 @@ disagree about the split — so the deadline is what keeps a broken kernel from taking the calling run down with it. - CUDA_VISIBLE_DEVICES=0,1,2,3 uv run python catalog/comm/_reducescatter_op_matrix.py + CUDA_VISIBLE_DEVICES=0,1,2,3 python _reducescatter_op_matrix.py That runs two jobs, in this order. The first is the matrix proper: every `CHECKS` entry on `world_size` ranks, and it must exit 0. The second is a @@ -33,9 +33,10 @@ The collected entry point is `tests/unittest/_torch/staircase/comm/test_staircase_reducescatter_op_matrix.py`: -it starts this job and turns its exit code into an assertion. This half stays -in the package because the launcher re-execs it as `python -m`, and the ranks -need the package context for their relative imports. +it starts this job and turns its exit code into an assertion. Both halves are +started by file path: the launcher must not import `tensorrt_llm` (that calls +`MPI_Init`, and an MPI-initialized process cannot start `mpirun`), and the +ranks reach the catalog by absolute import, so neither needs a package. """ import os @@ -46,6 +47,7 @@ import tempfile import threading import time +from pathlib import Path from typing import Any, Dict, List, Optional, Sequence, Tuple import torch @@ -1428,7 +1430,7 @@ def _run_one_rank() -> int: global COMM, RANK, WORLD, GROUP, reducescatter from mpi4py import MPI - from . import reducescatter as entry + from tensorrt_llm._torch.staircase.catalog.comm import reducescatter as entry reducescatter = entry.reducescatter COMM = MPI.COMM_WORLD @@ -1495,7 +1497,7 @@ def _run_wedge_rank() -> int: global COMM, RANK, WORLD, GROUP, reducescatter from mpi4py import MPI - from . import reducescatter as entry + from tensorrt_llm._torch.staircase.catalog.comm import reducescatter as entry reducescatter = entry.reducescatter COMM = MPI.COMM_WORLD @@ -1542,8 +1544,7 @@ def _mpirun(world_size: int, flag: str, env: Optional[Dict[str, str]] = None) -> "-n", str(world_size), sys.executable, - "-m", - "tensorrt_llm._torch.staircase.catalog.comm._reducescatter_op_matrix", + str(Path(__file__).resolve()), flag, ] print(f"[launcher] {' '.join(command)}", flush=True) diff --git a/tests/unittest/_torch/staircase/comm/test_staircase_allgather_op_matrix.py b/tests/unittest/_torch/staircase/comm/test_staircase_allgather_op_matrix.py index 952f3c350841..bca298e5aa43 100644 --- a/tests/unittest/_torch/staircase/comm/test_staircase_allgather_op_matrix.py +++ b/tests/unittest/_torch/staircase/comm/test_staircase_allgather_op_matrix.py @@ -9,15 +9,14 @@ are silently wrong rather than loud. The names are kept apart so review does not read one as a copy of the other. -The matrix itself lives in ``catalog/comm/_allgather_op_matrix.py``, which is +The matrix itself is ``_allgather_op_matrix.py`` beside this file, which is its own 4-rank launcher; see ``_rank_job`` for why that is left intact. """ +import _rank_job import pytest import torch -from tensorrt_llm._torch.staircase.catalog.comm import _rank_job - assert torch.cuda.is_available(), "allgather requires CUDA devices" diff --git a/tests/unittest/_torch/staircase/comm/test_staircase_reducescatter_op_matrix.py b/tests/unittest/_torch/staircase/comm/test_staircase_reducescatter_op_matrix.py index 407cd19f7d50..3b1ac2bba531 100644 --- a/tests/unittest/_torch/staircase/comm/test_staircase_reducescatter_op_matrix.py +++ b/tests/unittest/_torch/staircase/comm/test_staircase_reducescatter_op_matrix.py @@ -7,18 +7,17 @@ ``AllReduce`` module, while this covers the op itself cell by cell. The names are kept apart so review does not read one as a copy of the other. -The matrix lives in ``catalog/comm/_reducescatter_op_matrix.py``, which is its +The matrix is ``_reducescatter_op_matrix.py`` beside this file, which is its own launcher and runs two jobs: the ordered check sequence, then a separately capped job that certifies the one call-order divergence that wedges instead of lying (it cannot be a normal check, because the job that runs it never reports). """ +import _rank_job import pytest import torch -from tensorrt_llm._torch.staircase.catalog.comm import _rank_job - assert torch.cuda.is_available(), "reducescatter requires CUDA devices" diff --git a/tests/unittest/_torch/staircase/gemm/test_staircase_nvfp4_gemm.py b/tests/unittest/_torch/staircase/gemm/test_staircase_nvfp4_gemm.py index abcc7cd6132a..ed388b240e2a 100644 --- a/tests/unittest/_torch/staircase/gemm/test_staircase_nvfp4_gemm.py +++ b/tests/unittest/_torch/staircase/gemm/test_staircase_nvfp4_gemm.py @@ -647,3 +647,45 @@ def guarded(fn) -> None: out = torch.ops.trtllm.nvfp4_gemm(a_data, b_data, a_sf, b_sf, fat_alpha, torch.bfloat16) torch.testing.assert_close(out, ref.to(torch.bfloat16)) guarded(lambda: nvfp4_gemm(a_data, b_data, a_sf, b_sf, fat_alpha, torch.bfloat16)) + + +def test_a_short_scale_buffer_is_rejected_before_dispatch() -> None: + """A contiguous but undersized scale buffer reads past its end. + + The swizzle addresses the *padded* rectangle, so the length the kernel + indexes is `pad_up(rows,128) * pad_up(K/16,4)` whatever the real shape is. + Contiguity says nothing about that, and the op does not check it: a short + buffer is read out of bounds, which is either a wrong result or a dead + worker depending on what follows it in the allocation. + + M and N pad independently, so both operands are exercised, and both are + sized so that the padding is not a no-op: 200 and 160 each pad to 256. + Picking a row count already on the 128 boundary would make the "unpadded + length" case below identical to the correct one and assert nothing. + """ + m, n, k = 200, 160, 512 + a_data, _, a_sf, a_values = _operand(m, k, seed=61) + b_data, _, b_sf, b_values = _operand(n, k, seed=62) + alpha = _alpha(ALPHA) + + # The full-length buffers are accepted and correct, so the rejection below + # is about length alone and not about these operands. + ref = _reference(a_values, b_values, ALPHA) + out = nvfp4_gemm(a_data, b_data, a_sf, b_sf, alpha, torch.bfloat16) + torch.testing.assert_close(out, ref.to(torch.bfloat16)) + assert a_sf.numel() == _pad_up(m, 128) * _pad_up(k // VEC, 4) + assert b_sf.numel() == _pad_up(n, 128) * _pad_up(k // VEC, 4) + + def rejected(*args) -> None: + try: + nvfp4_gemm(*args, alpha, torch.bfloat16) + except AssertionError: + return + raise AssertionError("expected the wrapper to reject a short scale buffer") + + # One byte short, and the un-padded length -- the two ways to get it wrong: + # truncation, and sizing the buffer from the real rectangle. + rejected(a_data, b_data, a_sf[:-1].contiguous(), b_sf) + rejected(a_data, b_data, a_sf[: m * (k // VEC)].contiguous(), b_sf) + rejected(a_data, b_data, a_sf, b_sf[:-1].contiguous()) + rejected(a_data, b_data, a_sf, b_sf[: n * (k // VEC)].contiguous()) diff --git a/tests/unittest/_torch/staircase/test_staircase_claims.py b/tests/unittest/_torch/staircase/test_staircase_claims.py index 5afe0499c857..78d2812f14df 100644 --- a/tests/unittest/_torch/staircase/test_staircase_claims.py +++ b/tests/unittest/_torch/staircase/test_staircase_claims.py @@ -130,37 +130,59 @@ def test_checkpoint_fingerprints_are_distinct(arch): ) +# Context fields a routing tree may not branch on yet, and what has to happen +# before it can. Both would otherwise decide silently on a value that is not +# the deployment's -- the failure ``require`` exists to prevent. Delete a row +# once its prerequisite is met. +_UNREADABLE_CONTEXT_FIELDS = { + "is_disagg": ( + "no caller sets it, so it reads False in every deployment, " + "disaggregated or not; plumb it onto ModelConfig first" + ), + "spec_config": ( + "explain.py has no flag for a speculative config, so it would report " + "the wrong target for every drafting configuration; give explain a " + "way to name one first" + ), +} + + @pytest.mark.parametrize("arch", _ARCHS) -def test_no_routing_module_reads_an_unplumbed_dimension(arch): - """``ctx.is_disagg`` is declared but nothing sets it, so it is always - False. A branch on it would take the wrong side in a disaggregated - deployment and say nothing -- exactly the silent-wrong-system failure - ``require`` exists to prevent. Delete this test when it is plumbed. +@pytest.mark.parametrize("field", sorted(_UNREADABLE_CONTEXT_FIELDS)) +def test_no_routing_module_reads_an_unplumbed_dimension(arch, field): + """A routing tree may only read what both the engine and explain can fill. + + ``explain`` replays the same tree to answer "why did I not get the target + I expected". A criterion it cannot evaluate makes that answer wrong on + exactly the configurations someone would ask about -- so the set of + readable fields is bounded by the weaker of the two callers, not the + engine alone. """ source = _module_path(STAIRCASE_ROUTERS[arch]).read_text() - assert "is_disagg" not in source, ( - f"{arch}: routing reads ctx.is_disagg, which no caller sets yet; " - f"plumb it onto ModelConfig first" + assert field not in source, ( + f"{arch}: routing reads ctx.{field}, but {_UNREADABLE_CONTEXT_FIELDS[field]}" ) -def test_every_target_ships_the_three_products(): - """modeling.py, weights.py and TARGET.md travel together. +def test_every_target_ships_both_products(): + """modeling.py and weights.py travel together. - The forward, the weights it expects, and the record of what that pair was - measured to do. A target missing the third is one nobody can check. + The forward and the weights it expects. A target missing either is not a + target. - There used to be a fourth, a per-target ``smoke.py``. In-tree there is no - reason to carry a bespoke keyword-assert CLI: the same boot-and-generate - check is ``examples/llm-api/quickstart_advanced.py``, and the accuracy - gates are ``trtllm-eval`` and ``accuracy/test_staircase.py`` -- which, - unlike a module that nothing ever ran, CI actually runs. + Two products, not the four the out-of-tree tree carried. The per-target + ``smoke.py`` and ``configs/`` went with the move: the boot-and-generate + check is ``examples/llm-api/quickstart_advanced.py``, the knob variants are + LLM API arguments the accuracy gates pass directly, and what a target was + measured to do is recorded where every other model records it -- + ``tests/integration/defs/accuracy/references/``, read by the gates CI runs + rather than by a reader. """ for arch in _ARCHS: routing = routing_module(arch) for name, dotted in routing.TARGET_MODULES.items(): target_dir = _module_path(dotted).parent - for product in ("modeling.py", "weights.py", "TARGET.md"): + for product in ("modeling.py", "weights.py"): assert (target_dir / product).is_file(), f"{name}: missing {product}" diff --git a/tests/unittest/_torch/staircase/test_staircase_no_stale_claims.py b/tests/unittest/_torch/staircase/test_staircase_no_stale_claims.py new file mode 100644 index 000000000000..6a70dcbee41a --- /dev/null +++ b/tests/unittest/_torch/staircase/test_staircase_no_stale_claims.py @@ -0,0 +1,78 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""The audit that keeps pinned versions and dates out of the staircase tree. + +Out of tree, staircase was a separate repo against a pinned `tensorrt_llm`, so +writing that version into a contract was meaningful: the pin was the thing a +receipt was valid against, and a bump really did void what had been measured. + +In tree there is no external pin to drift against. The catalog and the ops it +wraps move with the trunk together, and the catalog tests run in pre-merge, so +the trunk proves itself on every commit. A version written into a contract is +then a number that is wrong the next day with nothing to notice -- worse than +absent, because every release bump makes a receipt CI is actively keeping +green look stale and reopens a question already answered. + +Removing them once was the easy half; this is what stops the next one. A +scan rather than a review checklist, for the reason the tree prefers +everywhere: a rule nothing enforces is a rule that decays. +""" + +from __future__ import annotations + +import re +from pathlib import Path + +import tensorrt_llm._torch.staircase as _staircase + +# The package, not this file: the tree under audit lives under tensorrt_llm/ +# while this test lives under tests/. +_ROOT = Path(_staircase.__file__).resolve().parent + +# Each pattern, and what to say instead of it. +_STALE_CLAIM_PATTERNS = ( + ( + re.compile(r"\b\d+\.\d+\.\d+rc\d+\b"), + "a pinned tensorrt_llm release. The catalog moves with the trunk, so " + "state the behaviour, not the build it was seen on", + ), + ( + re.compile(r"(?i)\bflashinfer[-\w]*\s+v?\d+\.\d+\.\d+"), + "a pinned flashinfer release. Name the kernel or the code path, not " + "the version that happened to be installed", + ), + ( + re.compile(r"(?i)\btensorrt[-_]llm\s*[=<>!]=\s*\d"), + "a pinned tensorrt_llm requirement. A target moves with the trunk; " + "the SM assert is the part of its identity that does not move", + ), + ( + re.compile(r"(?i)\bmeasured\s+(?:on\s+)?20\d\d-\d\d-\d\d"), + "a measurement date, which says nothing about whether the claim still " + "holds. Say that it was measured rather than when", + ), +) + +#: Scanned for the above. The catalog contracts are the point -- they are where +#: a version pin last accumulated -- but a target's modeling code carries the +#: same kind of prose, so the whole package is in. +_PROSE_SUFFIXES = {".py", ".md", ".yaml", ".yml"} + + +def _package_files(): + for path in sorted(_ROOT.rglob("*")): + if path.suffix in _PROSE_SUFFIXES and "__pycache__" not in path.parts: + yield path + + +def test_no_file_pins_a_version_or_a_date(): + """Report every stale claim at once, with the line and what to write instead.""" + offences = [] + for path in _package_files(): + text = path.read_text(encoding="utf-8", errors="ignore") + for lineno, line in enumerate(text.splitlines(), 1): + for pattern, why in _STALE_CLAIM_PATTERNS: + if (hit := pattern.search(line)) is not None: + rel = path.relative_to(_ROOT) + offences.append(f" {rel}:{lineno}: {hit.group(0)!r} is {why}") + assert not offences, "staircase files must not pin a version or a date:\n" + "\n".join(offences) diff --git a/tests/unittest/_torch/staircase/test_staircase_routing.py b/tests/unittest/_torch/staircase/test_staircase_routing.py index af9c0a985b39..70c07d7b92df 100644 --- a/tests/unittest/_torch/staircase/test_staircase_routing.py +++ b/tests/unittest/_torch/staircase/test_staircase_routing.py @@ -209,10 +209,19 @@ def test_an_unrouted_architecture_raises_under_require(monkeypatch): @pytest.mark.parametrize( - "raw,expected", [("off", "off"), ("AUTO", "auto"), (" require ", "require"), ("", "off")] + "raw,expected", + [(None, "off"), ("off", "off"), ("AUTO", "auto"), (" require ", "require"), ("", "off")], ) def test_the_env_var_is_read_leniently(monkeypatch, raw, expected): - monkeypatch.setenv(STAIRCASE_ENV, raw) + """``None`` is the unset case, and it is the one that must never drift. + + Everything in the accuracy suite rests on staircase being opt-in: unset has + to read as off on the code path the engine actually takes. + """ + if raw is None: + monkeypatch.delenv(STAIRCASE_ENV, raising=False) + else: + monkeypatch.setenv(STAIRCASE_ENV, raw) assert StaircaseMode.from_env().value == expected diff --git a/tests/unittest/_torch/staircase/test_staircase_target_contract.py b/tests/unittest/_torch/staircase/test_staircase_target_contract.py new file mode 100644 index 000000000000..933a84e704fa --- /dev/null +++ b/tests/unittest/_torch/staircase/test_staircase_target_contract.py @@ -0,0 +1,73 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Every op a target names has to exist in the build it is running against. + +A target reads private engine surface: custom ops registered under +``torch.ops.trtllm`` and one pybind entry point. With no version pin, nothing +declares which build it was written for -- so the only thing that can tell you +the surface moved is running against it. + +This used to be an ``assert`` loop that each target ran at import. That made +every import of a modeling module pay for the check and put a test in the +product tree; the missing symbol is a fact of the build, so the place to find +it out is a test on a machine that has the extension built. Each target now +declares ``REQUIRED_TRTLLM_OPS`` and this asserts it. + +Distinct from ``test_staircase_claims.py``, which is deliberately import-free +and runs anywhere: this one imports the targets, so it needs a real build. +""" + +from __future__ import annotations + +import importlib + +import pytest +import torch + +import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* +from tensorrt_llm._torch.staircase._router_index import STAIRCASE_ROUTERS, routing_module + +_PACKAGE = "tensorrt_llm._torch.staircase" +_ARCHS = sorted(STAIRCASE_ROUTERS) + + +def _target_modules(): + """(target name, imported modeling module) for every routed target.""" + for arch in _ARCHS: + routing = routing_module(arch) + for name, dotted in routing.TARGET_MODULES.items(): + yield name, importlib.import_module(f"{_PACKAGE}.{dotted}") + + +def _target_ids(): + return [name for name, _ in _target_modules()] + + +@pytest.mark.parametrize("name,module", list(_target_modules()), ids=_target_ids()) +def test_every_declared_op_exists(name, module): + """A named op that is not registered is a build the target cannot run on.""" + declared = getattr(module, "REQUIRED_TRTLLM_OPS", None) + assert declared, f"{name}: modeling.py declares no REQUIRED_TRTLLM_OPS" + + missing = [op for op in declared if not hasattr(torch.ops.trtllm, op)] + assert not missing, ( + f"{name} names {len(missing)} op(s) this build does not register: " + f"{', '.join(missing)}. Either the op was renamed upstream and the " + f"target has not followed, or this build predates it." + ) + + +@pytest.mark.parametrize("name,module", list(_target_modules()), ids=_target_ids()) +def test_the_pybind_attention_entry_point_exists(name, module): + """``thop.attention`` is reached through the bindings, not torch.ops. + + Separate from the loop above because a missing pybind symbol fails in a + different way -- an ImportError or an AttributeError on the module object + rather than a missing torch.ops entry -- and both targets go through it. + """ + from tensorrt_llm.bindings.internal import thop + + assert hasattr(thop, "attention"), ( + f"{name} calls the attention op through tensorrt_llm.bindings.internal" + f".thop.attention, which this build does not expose" + ) From 16a5e072a950ccc4a46842b285cf286671eaede9 Mon Sep 17 00:00:00 2001 From: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> Date: Tue, 15 Sep 2026 08:50:32 -0700 Subject: [PATCH 07/19] [TRTLLM-16304][chore] Rename staircase to modeling_v2 "Staircase" named the approach, not the thing, so every reader needed it explained before they could read anything else. modeling_v2 says what the subtree is: a second modeling path, beside the zoo rather than inside it. Mechanical throughout. tensorrt_llm/_torch/staircase -> tensorrt_llm/_torch/modeling_v2, the test tree and its 25 files with it, TRTLLM_STAIRCASE -> TRTLLM_MODELING_V2, and the identifiers, synthetic architecture names and accuracy anchor key follow the same substitution. The package still sits beside _torch/models/ rather than inside it, which is what makes its registrations count as external and keeps them out of the zoo's static index. Two consequences of the move that needed a real change rather than a substitution: _rank_job resolves the repo root by parent index, which the depth change would have silently broken, and the l0 entries, CODEOWNERS paths and the acceptance anchor key all had to stay in step with the new names. Verified on GB300 against this tree: 358 passed on the 1-GPU l0 entry and 2 on the 4-GPU one, both identical to the pre-rename counts, plus the backend guard, explain against a real checkpoint, and collection of all four accuracy node IDs. check_test_list.py --validate passes. Left alone deliberately: "staircase" as a graphics term in the cosmos3 negative prompts, which is the word for aliasing and nothing to do with this. Signed-off-by: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> --- .github/CODEOWNERS | 14 +-- .../{staircase => modeling_v2}/README.md | 52 +++++------ .../{staircase => modeling_v2}/__init__.py | 26 +++--- .../_router_index.py | 84 +++++++++--------- .../catalog/__init__.py | 0 .../catalog/activation/__init__.py | 0 .../activation/flashinfer_silu_and_mul.md | 0 .../activation/flashinfer_silu_and_mul.py | 0 .../catalog/attention/__init__.py | 0 .../catalog/attention/fused_qk_norm_rope.md | 0 .../catalog/attention/fused_qk_norm_rope.py | 0 .../attention/load_paged_kv_cache_for_mla.md | 0 .../attention/load_paged_kv_cache_for_mla.py | 0 .../mla_rope_append_paged_kv_assign_q.md | 0 .../mla_rope_append_paged_kv_assign_q.py | 0 .../catalog/attention/mla_rope_generation.md | 0 .../catalog/attention/mla_rope_generation.py | 0 .../catalog/attention/thop_attention.md | 0 .../catalog/attention/thop_attention.py | 0 .../catalog/comm/__init__.py | 0 .../catalog/comm/allgather.md | 0 .../catalog/comm/allgather.py | 0 .../catalog/comm/reducescatter.md | 0 .../catalog/comm/reducescatter.py | 0 .../catalog/gemm/__init__.py | 0 .../catalog/gemm/bmm_out.md | 0 .../catalog/gemm/bmm_out.py | 0 .../catalog/gemm/cublas_mm.md | 0 .../catalog/gemm/cublas_mm.py | 0 .../catalog/gemm/nvfp4_gemm.md | 0 .../catalog/gemm/nvfp4_gemm.py | 0 .../catalog/index.yaml | 4 +- .../catalog/moe/__init__.py | 0 .../catalog/moe/fp4_block_scale_moe_runner.md | 0 .../catalog/moe/fp4_block_scale_moe_runner.py | 0 .../catalog/moe/fused_moe.md | 0 .../catalog/moe/fused_moe.py | 0 .../mxe4m3_mxe2m1_block_scale_moe_runner.md | 0 .../mxe4m3_mxe2m1_block_scale_moe_runner.py | 0 .../catalog/moe/noaux_tc_op.md | 0 .../catalog/moe/noaux_tc_op.py | 0 .../catalog/norm/__init__.py | 0 .../norm/flashinfer_fused_add_rmsnorm.md | 0 .../norm/flashinfer_fused_add_rmsnorm.py | 0 .../catalog/norm/flashinfer_rmsnorm.md | 0 .../catalog/norm/flashinfer_rmsnorm.py | 0 .../catalog/quantization/__init__.py | 0 .../catalog/quantization/fp4_quantize.md | 0 .../catalog/quantization/fp4_quantize.py | 0 .../catalog/quantization/mxfp8_quantize.md | 0 .../catalog/quantization/mxfp8_quantize.py | 0 .../catalog/torch/__init__.py | 0 .../catalog/torch/add.py | 0 .../catalog/torch/concat.py | 0 .../catalog/torch/copy_.py | 0 .../catalog/torch/embedding.py | 0 .../catalog/torch/empty.py | 0 .../catalog/torch/expand.py | 0 .../catalog/torch/pad.py | 0 .../catalog/torch/reshape.py | 0 .../catalog/torch/split.py | 0 .../catalog/torch/transpose.py | 0 .../catalog/torch/view_dtype.py | 0 .../{staircase => modeling_v2}/explain.py | 12 +-- .../models/__init__.py | 0 .../models/deepseek_v3/__init__.py | 0 .../models/deepseek_v3/routing.py | 8 +- .../models/deepseek_v3/targets/__init__.py | 0 .../targets/r1_0528_nvfp4/__init__.py | 0 .../targets/r1_0528_nvfp4/sm_103/__init__.py | 0 .../r1_0528_nvfp4/sm_103/dep4/__init__.py | 0 .../r1_0528_nvfp4/sm_103/dep4/modeling.py | 86 ++++++++++--------- .../r1_0528_nvfp4/sm_103/dep4/weights.py | 0 .../models/gpt_oss/__init__.py | 0 .../models/gpt_oss/routing.py | 8 +- .../models/gpt_oss/targets/__init__.py | 0 .../gpt_oss/targets/gpt_oss_120b/__init__.py | 0 .../targets/gpt_oss_120b/sm_103/__init__.py | 0 .../gpt_oss_120b/sm_103/tp1/__init__.py | 0 .../gpt_oss_120b/sm_103/tp1/modeling.py | 40 ++++----- .../gpt_oss_120b/sm_103/tp1/weights.py | 0 tensorrt_llm/_torch/models/modeling_auto.py | 28 +++--- tensorrt_llm/llmapi/llm.py | 4 +- .../references/acceptance_length.yaml | 16 ++-- ..._v3.py => test_modeling_v2_deepseek_v3.py} | 26 +++--- ...gpt_oss.py => test_modeling_v2_gpt_oss.py} | 14 +-- .../test_lists/test-db/l0_gb300.yml | 10 +-- .../test-db/l0_gb300_multi_gpus.yml | 14 +-- ...st_modeling_v2_flashinfer_silu_and_mul.py} | 2 +- .../test_modeling_v2_fused_qk_norm_rope.py} | 2 +- ...odeling_v2_load_paged_kv_cache_for_mla.py} | 4 +- ...g_v2_mla_rope_append_paged_kv_assign_q.py} | 4 +- .../test_modeling_v2_mla_rope_generation.py} | 4 +- .../test_modeling_v2_thop_attention.py} | 2 +- .../comm/_allgather_op_matrix.py | 4 +- .../comm/_rank_job.py | 2 +- .../comm/_reducescatter_op_matrix.py | 6 +- .../test_modeling_v2_allgather_op_matrix.py} | 0 ...st_modeling_v2_reducescatter_op_matrix.py} | 0 .../gemm/test_modeling_v2_bmm_out.py} | 2 +- .../gemm/test_modeling_v2_cublas_mm.py} | 2 +- .../gemm/test_modeling_v2_nvfp4_gemm.py} | 2 +- ...modeling_v2_fp4_block_scale_moe_runner.py} | 2 +- .../moe/test_modeling_v2_fused_moe.py} | 2 +- ...2_mxe4m3_mxe2m1_block_scale_moe_runner.py} | 2 +- .../moe/test_modeling_v2_noaux_tc_op.py} | 2 +- ...deling_v2_flashinfer_fused_add_rmsnorm.py} | 2 +- .../test_modeling_v2_flashinfer_rmsnorm.py} | 2 +- .../test_modeling_v2_fp4_quantize.py} | 2 +- .../test_modeling_v2_mxfp8_quantize.py} | 2 +- .../test_modeling_v2_claims.py} | 16 ++-- .../test_modeling_v2_no_stale_claims.py} | 12 +-- .../test_modeling_v2_routing.py} | 60 ++++++------- .../test_modeling_v2_target_contract.py} | 8 +- 114 files changed, 300 insertions(+), 294 deletions(-) rename tensorrt_llm/_torch/{staircase => modeling_v2}/README.md (89%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/__init__.py (78%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/_router_index.py (80%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/__init__.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/activation/__init__.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/activation/flashinfer_silu_and_mul.md (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/activation/flashinfer_silu_and_mul.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/attention/__init__.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/attention/fused_qk_norm_rope.md (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/attention/fused_qk_norm_rope.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/attention/load_paged_kv_cache_for_mla.md (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/attention/load_paged_kv_cache_for_mla.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/attention/mla_rope_append_paged_kv_assign_q.md (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/attention/mla_rope_append_paged_kv_assign_q.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/attention/mla_rope_generation.md (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/attention/mla_rope_generation.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/attention/thop_attention.md (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/attention/thop_attention.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/comm/__init__.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/comm/allgather.md (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/comm/allgather.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/comm/reducescatter.md (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/comm/reducescatter.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/gemm/__init__.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/gemm/bmm_out.md (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/gemm/bmm_out.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/gemm/cublas_mm.md (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/gemm/cublas_mm.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/gemm/nvfp4_gemm.md (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/gemm/nvfp4_gemm.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/index.yaml (99%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/moe/__init__.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/moe/fp4_block_scale_moe_runner.md (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/moe/fp4_block_scale_moe_runner.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/moe/fused_moe.md (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/moe/fused_moe.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/moe/noaux_tc_op.md (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/moe/noaux_tc_op.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/norm/__init__.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/norm/flashinfer_fused_add_rmsnorm.md (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/norm/flashinfer_fused_add_rmsnorm.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/norm/flashinfer_rmsnorm.md (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/norm/flashinfer_rmsnorm.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/quantization/__init__.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/quantization/fp4_quantize.md (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/quantization/fp4_quantize.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/quantization/mxfp8_quantize.md (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/quantization/mxfp8_quantize.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/torch/__init__.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/torch/add.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/torch/concat.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/torch/copy_.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/torch/embedding.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/torch/empty.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/torch/expand.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/torch/pad.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/torch/reshape.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/torch/split.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/torch/transpose.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/catalog/torch/view_dtype.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/explain.py (89%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/models/__init__.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/models/deepseek_v3/__init__.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/models/deepseek_v3/routing.py (88%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/models/deepseek_v3/targets/__init__.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/models/deepseek_v3/targets/r1_0528_nvfp4/__init__.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/__init__.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/__init__.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py (97%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/weights.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/models/gpt_oss/__init__.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/models/gpt_oss/routing.py (86%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/models/gpt_oss/targets/__init__.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/models/gpt_oss/targets/gpt_oss_120b/__init__.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/models/gpt_oss/targets/gpt_oss_120b/sm_103/__init__.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/__init__.py (100%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py (95%) rename tensorrt_llm/_torch/{staircase => modeling_v2}/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/weights.py (100%) rename tests/integration/defs/accuracy/{test_staircase_deepseek_v3.py => test_modeling_v2_deepseek_v3.py} (90%) rename tests/integration/defs/accuracy/{test_staircase_gpt_oss.py => test_modeling_v2_gpt_oss.py} (87%) rename tests/unittest/_torch/{staircase/activation/test_staircase_flashinfer_silu_and_mul.py => modeling_v2/activation/test_modeling_v2_flashinfer_silu_and_mul.py} (97%) rename tests/unittest/_torch/{staircase/attention/test_staircase_fused_qk_norm_rope.py => modeling_v2/attention/test_modeling_v2_fused_qk_norm_rope.py} (98%) rename tests/unittest/_torch/{staircase/attention/test_staircase_load_paged_kv_cache_for_mla.py => modeling_v2/attention/test_modeling_v2_load_paged_kv_cache_for_mla.py} (99%) rename tests/unittest/_torch/{staircase/attention/test_staircase_mla_rope_append_paged_kv_assign_q.py => modeling_v2/attention/test_modeling_v2_mla_rope_append_paged_kv_assign_q.py} (99%) rename tests/unittest/_torch/{staircase/attention/test_staircase_mla_rope_generation.py => modeling_v2/attention/test_modeling_v2_mla_rope_generation.py} (99%) rename tests/unittest/_torch/{staircase/attention/test_staircase_thop_attention.py => modeling_v2/attention/test_modeling_v2_thop_attention.py} (99%) rename tests/unittest/_torch/{staircase => modeling_v2}/comm/_allgather_op_matrix.py (99%) rename tests/unittest/_torch/{staircase => modeling_v2}/comm/_rank_job.py (98%) rename tests/unittest/_torch/{staircase => modeling_v2}/comm/_reducescatter_op_matrix.py (99%) rename tests/unittest/_torch/{staircase/comm/test_staircase_allgather_op_matrix.py => modeling_v2/comm/test_modeling_v2_allgather_op_matrix.py} (100%) rename tests/unittest/_torch/{staircase/comm/test_staircase_reducescatter_op_matrix.py => modeling_v2/comm/test_modeling_v2_reducescatter_op_matrix.py} (100%) rename tests/unittest/_torch/{staircase/gemm/test_staircase_bmm_out.py => modeling_v2/gemm/test_modeling_v2_bmm_out.py} (97%) rename tests/unittest/_torch/{staircase/gemm/test_staircase_cublas_mm.py => modeling_v2/gemm/test_modeling_v2_cublas_mm.py} (98%) rename tests/unittest/_torch/{staircase/gemm/test_staircase_nvfp4_gemm.py => modeling_v2/gemm/test_modeling_v2_nvfp4_gemm.py} (99%) rename tests/unittest/_torch/{staircase/moe/test_staircase_fp4_block_scale_moe_runner.py => modeling_v2/moe/test_modeling_v2_fp4_block_scale_moe_runner.py} (99%) rename tests/unittest/_torch/{staircase/moe/test_staircase_fused_moe.py => modeling_v2/moe/test_modeling_v2_fused_moe.py} (99%) rename tests/unittest/_torch/{staircase/moe/test_staircase_mxe4m3_mxe2m1_block_scale_moe_runner.py => modeling_v2/moe/test_modeling_v2_mxe4m3_mxe2m1_block_scale_moe_runner.py} (99%) rename tests/unittest/_torch/{staircase/moe/test_staircase_noaux_tc_op.py => modeling_v2/moe/test_modeling_v2_noaux_tc_op.py} (99%) rename tests/unittest/_torch/{staircase/norm/test_staircase_flashinfer_fused_add_rmsnorm.py => modeling_v2/norm/test_modeling_v2_flashinfer_fused_add_rmsnorm.py} (97%) rename tests/unittest/_torch/{staircase/norm/test_staircase_flashinfer_rmsnorm.py => modeling_v2/norm/test_modeling_v2_flashinfer_rmsnorm.py} (96%) rename tests/unittest/_torch/{staircase/quantization/test_staircase_fp4_quantize.py => modeling_v2/quantization/test_modeling_v2_fp4_quantize.py} (99%) rename tests/unittest/_torch/{staircase/quantization/test_staircase_mxfp8_quantize.py => modeling_v2/quantization/test_modeling_v2_mxfp8_quantize.py} (99%) rename tests/unittest/_torch/{staircase/test_staircase_claims.py => modeling_v2/test_modeling_v2_claims.py} (94%) rename tests/unittest/_torch/{staircase/test_staircase_no_stale_claims.py => modeling_v2/test_modeling_v2_no_stale_claims.py} (88%) rename tests/unittest/_torch/{staircase/test_staircase_routing.py => modeling_v2/test_modeling_v2_routing.py} (82%) rename tests/unittest/_torch/{staircase/test_staircase_target_contract.py => modeling_v2/test_modeling_v2_target_contract.py} (91%) diff --git a/.github/CODEOWNERS b/.github/CODEOWNERS index 2d93e55f7e20..4bad576cf8d3 100644 --- a/.github/CODEOWNERS +++ b/.github/CODEOWNERS @@ -385,18 +385,18 @@ /tensorrt_llm/scaffolding @WeiHaocheng @dc3671 /tests/unittest/scaffolding @WeiHaocheng @dc3671 -# ===== STAIRCASE ===== +# ===== MODELING_V2 ===== # Overrides the /tensorrt_llm/_torch runtime-devs rule above for this subtree. # Individual handles rather than a team, like SCAFFOLDING: this is one bounded # experiment with named owners, not a standing domain. -/tensorrt_llm/_torch/staircase @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju -/tests/unittest/_torch/staircase @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju -/tests/integration/defs/accuracy/test_staircase_deepseek_v3.py @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju -/tests/integration/defs/accuracy/test_staircase_gpt_oss.py @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju +/tensorrt_llm/_torch/modeling_v2 @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju +/tests/unittest/_torch/modeling_v2 @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju +/tests/integration/defs/accuracy/test_modeling_v2_deepseek_v3.py @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju +/tests/integration/defs/accuracy/test_modeling_v2_gpt_oss.py @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju # The two catalog categories whose contracts state kernel behaviour the # attention and MoE owners are the authority on; co-owned rather than reassigned. -/tensorrt_llm/_torch/staircase/catalog/attention @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju @xxi-nv @yuxianq -/tensorrt_llm/_torch/staircase/catalog/moe @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju @xxi-nv @yuxianq +/tensorrt_llm/_torch/modeling_v2/catalog/attention @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju @xxi-nv @yuxianq +/tensorrt_llm/_torch/modeling_v2/catalog/moe @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju @xxi-nv @yuxianq ## TensorRT-LLM LLM Disaggregated /examples/disaggregated @NVIDIA/trt-llm-disagg-devs @NVIDIA/trt-llm-doc-owners diff --git a/tensorrt_llm/_torch/staircase/README.md b/tensorrt_llm/_torch/modeling_v2/README.md similarity index 89% rename from tensorrt_llm/_torch/staircase/README.md rename to tensorrt_llm/_torch/modeling_v2/README.md index 9aa1473c3bff..98b9000b145f 100644 --- a/tensorrt_llm/_torch/staircase/README.md +++ b/tensorrt_llm/_torch/modeling_v2/README.md @@ -1,11 +1,11 @@ -# Staircase +# ModelingV2 One self-contained modeling codebase per deployment target, beside the built-in model zoo rather than inside it. Where `_torch/models/modeling_deepseekv3.py` is one class serving V3, V3-Lite, R1 and V3.2 across every GPU generation and parallel topology, -`_torch/staircase/models/deepseek_v3/` is one flat forward per (checkpoint, +`_torch/modeling_v2/models/deepseek_v3/` is one flat forward per (checkpoint, GPU architecture, parallel topology) triple — assembled only from `catalog/` entries, sharing nothing with its siblings, and trusted through accuracy gates instead of shared abstractions. The one-to-one correspondence between @@ -14,19 +14,19 @@ instead of shared abstractions. The one-to-one correspondence between ## Using it ```bash -export TRTLLM_STAIRCASE=require # off | auto | require +export TRTLLM_MODELING_V2=require # off | auto | require ``` -| `TRTLLM_STAIRCASE` | Behaviour | +| `TRTLLM_MODELING_V2` | Behaviour | |---|---| | unset or `off` (default) | The resolver returns immediately. Nothing in this package is imported and behaviour is byte-for-byte what it is today. | | `auto` | Uses a target when one matches this exact configuration; falls back to the built-in implementation when none does. | | `require` | Raises instead of falling back, naming the criterion that did not match. | -**Use `require` for anything you will attribute to staircase.** Under `auto`, +**Use `require` for anything you will attribute to modeling_v2.** Under `auto`, a configuration that misses a target's criteria silently gets the built-in implementation — and a performance curve measured that way reads as -staircase's. That is the single most expensive mistake available here. +modeling_v2's. That is the single most expensive mistake available here. **Export it before the ranks start**, not merely before `LLM(...)`. Worker ranks receive the environment as it stood when MPI initialized, and long-lived @@ -36,12 +36,12 @@ its workers resolve the built-in is the silent split this package exists to prevent. An environment variable rather than an LLM-API field is a deliberate trade: it -keeps the entire concept inside this package, so the only change staircase +keeps the entire concept inside this package, so the only change modeling_v2 needs anywhere else is the `_resolve_class` hook. The cost is that the switch does not appear in a run's recorded `llm.args` and cannot be set through `--extra_llm_api_options`. -The predecessor `STAIRCASE_TARGET` is gone. It named a *target*; this names +The predecessor `MODELING_V2_TARGET` is gone. It named a *target*; this names only a mode, and routing picks the target from the configuration. The checkpoint is read exactly as published, with no target-owned @@ -52,7 +52,7 @@ The checkpoint is read exactly as published, with no target-owned ``` LLM(model=...) -> ModelLoader -> AutoModelForCausalLM._resolve_class | - +-- staircase_resolve(config) + +-- modeling_v2_resolve(config) | _router_index.py +-- architectures[0] -> routing module models//routing.py +-- one forward-reading decision tree @@ -61,7 +61,7 @@ LLM(model=...) -> ModelLoader -> AutoModelForCausalLM._resolve_class get_registered_model_class(name) ``` -The synthetic name (`StaircaseGptOss120bSm103Tp1`) is a registry key that no +The synthetic name (`ModelingV2GptOss120bSm103Tp1`) is a registry key that no checkpoint declares. Upstream already does exactly this for `MTPDraftModelForCausalLM`, which also exists only as a `_resolve_class` rewrite. @@ -69,7 +69,7 @@ rewrite. To ask why a configuration landed where it did: ``` -python -m tensorrt_llm._torch.staircase.explain \ +python -m tensorrt_llm._torch.modeling_v2.explain \ --model /path/to/DeepSeek-R1-0528-NVFP4 --tp 4 --ep 4 --attention-dp ``` @@ -112,10 +112,10 @@ Accuracy anchors are not kept here. They live where every other model's do, The third piece of a catalog entry, its GPU test, lives in the tests tree: ``` -tests/unittest/_torch/staircase/ - test_staircase_claims.py routing tables vs the targets they name (no GPU) - test_staircase_routing.py what staircase_resolve does (no GPU) - /test_staircase_.py +tests/unittest/_torch/modeling_v2/ + test_modeling_v2_claims.py routing tables vs the targets they name (no GPU) + test_modeling_v2_routing.py what modeling_v2_resolve does (no GPU) + /test_modeling_v2_.py comm/__op_matrix.py the two collectives' 4-rank rank bodies comm/_rank_job.py starts one of those and asserts on its exit code ``` @@ -135,7 +135,7 @@ each is one fixed 4-rank sequence that cannot run as independent cases. Identity is the path. `targets/` keeps all three segments rather than flattening them, and the class name carries the same triple; -`test_staircase_claims.py` asserts they agree. +`test_modeling_v2_claims.py` asserts they agree. ### Why beside `_torch/models/`, not inside it @@ -143,7 +143,7 @@ Two upstream mechanisms, and together they are the worst combination — inheriting the built-in constraints without inheriting the built-in guards: * `is_builtin_zoo_module` matches on the zoo's package prefix. Inside it, - staircase registrations would count as built-in and only fill *empty* + modeling_v2 registrations would count as built-in and only fill *empty* registry slots. Outside it they are external and always win their slot. * `test_lazy_model_zoo.py`'s scan of the zoo directory is **non-recursive**, so decorators in a subpackage are invisible to it — putting synthetic names @@ -155,21 +155,21 @@ checkpoints", which is the opposite of what this package is for. ## Gates 1. **Boot** — minutes, binary. `examples/llm-api/quickstart_advanced.py` with - the target's topology flags and `TRTLLM_STAIRCASE=require` exported. Engine + the target's topology flags and `TRTLLM_MODELING_V2=require` exported. Engine cold start, weight-manifest coverage, a handful of greedy continuations. Catches catastrophes, not accuracy. A variant that changes the forward gets its own minutes-scale gate on the same footing. 2. **Accuracy** — the release criterion. `trtllm-eval` with the protocol from - `tests/integration/defs/accuracy/references/` and `TRTLLM_STAIRCASE=require` + `tests/integration/defs/accuracy/references/` and `TRTLLM_MODELING_V2=require` exported. One-sided: measured >= reference − tol. This is the gate CI runs, - in the accuracy suite's staircase files. + in the accuracy suite's modeling_v2 files. 3. **Acceptance** — required whenever a variant is distribution-preserving by construction, speculative decoding above all. Rejection sampling holds the emitted distribution to the target model's, so a *miscomputed* draft path produces correct text more slowly: boot passes, accuracy passes, only speed moves. The detector is `acceptance_length` against the same checkpoint under stock in-tree modeling at the same workload and the same - `speculative_config` — i.e. `TRTLLM_STAIRCASE=off` versus `=require`, + `speculative_config` — i.e. `TRTLLM_MODELING_V2=off` versus `=require`, which is now one variable rather than two harnesses. Compare a variant against an identity run **in the same session**: one target @@ -208,7 +208,7 @@ tolerance, and each is written up in its own contract: by the op, not certified here). **gpt-oss-120b / sm_103 / tp1 is gated on GB300.** Both gates pass with -`TRTLLM_STAIRCASE=require` in force, which is what rules out the built-in +`TRTLLM_MODELING_V2=require` in force, which is what rules out the built-in implementation having been measured instead: | Gate | Result | @@ -256,7 +256,7 @@ failed at import. They now use the canonical path. The lesson generalizes: with no pin, a target's contact with private engine surface is checked only by running it. Each target declares that surface as -`REQUIRED_TRTLLM_OPS`, `test_staircase_target_contract.py` asserts every name +`REQUIRED_TRTLLM_OPS`, `test_modeling_v2_target_contract.py` asserts every name in it exists, and the first-forward metadata field check catches the rest -- which turns a drifting engine from a wrong answer into a loud failure. @@ -264,15 +264,15 @@ which turns a drifting engine from a wrong answer into a loud failure. * **`pip` is a runtime dependency of TensorRT-LLM itself.** `libtensorrt_llm.so` locates its bundled kernel headers on the NVRTC JIT path by running - `pip show tensorrt_llm`. Staircase was merely the first thing to write that + `pip show tensorrt_llm`. ModelingV2 was merely the first thing to write that down (a `uv`-managed venv does not install `pip` by default), and it is not - a staircase dependency. + a modeling_v2 dependency. * **The catalog contracts are now documentation of upstream ops.** For example `_torch/modules/linear.py` calls `cublas_mm(input, module.weight.t(), ...)` correctly, but nothing there says why the `.t()` is mandatory or that a non-contiguous weight computes a silently wrong answer — `catalog/gemm/cublas_mm.md` does. The reverse duty comes with it: if one of these ops changes and its - contract does not follow, the misleading is no longer confined to staircase. + contract does not follow, the misleading is no longer confined to modeling_v2. ## Not migrated in this batch diff --git a/tensorrt_llm/_torch/staircase/__init__.py b/tensorrt_llm/_torch/modeling_v2/__init__.py similarity index 78% rename from tensorrt_llm/_torch/staircase/__init__.py rename to tensorrt_llm/_torch/modeling_v2/__init__.py index ddf929d3e050..d1b0f9599ba8 100644 --- a/tensorrt_llm/_torch/staircase/__init__.py +++ b/tensorrt_llm/_torch/modeling_v2/__init__.py @@ -1,9 +1,9 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""Staircase: one self-contained modeling codebase per deployment target. +"""ModelingV2: one self-contained modeling codebase per deployment target. Where the built-in zoo has one class per architecture serving every -checkpoint, parallel topology and GPU generation, staircase has one flat, +checkpoint, parallel topology and GPU generation, modeling_v2 has one flat, self-contained forward per (checkpoint, GPU arch, parallel) triple, assembled only from ``catalog/`` entries and trusted through accuracy gates rather than shared abstractions. The two live side by side: ``models//`` here @@ -12,7 +12,7 @@ Entry is a single environment variable:: - TRTLLM_STAIRCASE = require + TRTLLM_MODELING_V2 = require ``off`` (unset, the default) is byte-for-byte today's behaviour: the resolver returns immediately and nothing in this package is imported. ``auto`` uses a target @@ -29,21 +29,21 @@ """ from ._router_index import ( + MODELING_V2_ENV, + MODELING_V2_ROUTERS, ROUTING_BACKENDS, - STAIRCASE_ENV, - STAIRCASE_ROUTERS, - StaircaseContext, - StaircaseMode, + ModelingV2Context, + ModelingV2Mode, assert_backend_can_route, - staircase_resolve, + modeling_v2_resolve, ) __all__ = [ "ROUTING_BACKENDS", - "STAIRCASE_ENV", - "STAIRCASE_ROUTERS", - "StaircaseContext", - "StaircaseMode", + "MODELING_V2_ENV", + "MODELING_V2_ROUTERS", + "ModelingV2Context", + "ModelingV2Mode", "assert_backend_can_route", - "staircase_resolve", + "modeling_v2_resolve", ] diff --git a/tensorrt_llm/_torch/staircase/_router_index.py b/tensorrt_llm/_torch/modeling_v2/_router_index.py similarity index 80% rename from tensorrt_llm/_torch/staircase/_router_index.py rename to tensorrt_llm/_torch/modeling_v2/_router_index.py index 4bb8b8569e64..919e33132d67 100644 --- a/tensorrt_llm/_torch/staircase/_router_index.py +++ b/tensorrt_llm/_torch/modeling_v2/_router_index.py @@ -2,10 +2,10 @@ # SPDX-License-Identifier: Apache-2.0 """Architecture name -> routing module, and the resolver that reads the table. -Staircase targets are keyed by a *synthetic* architecture name that no +ModelingV2 targets are keyed by a *synthetic* architecture name that no checkpoint declares. The checkpoint's own ``architectures[0]`` stays what it always was (``GptOssForCausalLM``, ``DeepseekV3ForCausalLM``), so it cannot -also be the target selector; ``staircase_resolve`` rewrites it into the +also be the target selector; ``modeling_v2_resolve`` rewrites it into the synthetic name, and ``AutoModelForCausalLM._resolve_class`` then looks that up in the ordinary registry. @@ -38,61 +38,61 @@ if TYPE_CHECKING: from tensorrt_llm._torch.model_config import ModelConfig -_PACKAGE = "tensorrt_llm._torch.staircase" +_PACKAGE = "tensorrt_llm._torch.modeling_v2" #: The switch. An environment variable rather than an LLM-API field, so that #: nothing outside this package has to carry the concept: the only upstream -#: change staircase needs is the ``_resolve_class`` hook itself. +#: change modeling_v2 needs is the ``_resolve_class`` hook itself. #: #: It has to be exported **before the ranks start**, not merely before #: ``LLM(...)``. Worker ranks receive the environment as it stood when MPI #: initialized, and long-lived ranks under ``trtllm-llmapi-launch`` receive it #: once at launch, so a value set later reaches the driver and not them -- and -#: a driver that resolves a staircase target while its workers resolve the +#: a driver that resolves a modeling_v2 target while its workers resolve the #: built-in is exactly the silent split this package exists to prevent. Export #: it in the shell, or before ``import tensorrt_llm``. -STAIRCASE_ENV = "TRTLLM_STAIRCASE" +MODELING_V2_ENV = "TRTLLM_MODELING_V2" # architectures[0] -> routing module, relative to this package. -STAIRCASE_ROUTERS = { +MODELING_V2_ROUTERS = { "GptOssForCausalLM": "models.gpt_oss.routing", "DeepseekV3ForCausalLM": "models.deepseek_v3.routing", } -#: Backends whose model construction reaches ``staircase_resolve``. The +#: Backends whose model construction reaches ``modeling_v2_resolve``. The #: resolver is called from ``AutoModelForCausalLM._resolve_class``, so a #: backend that builds its model some other way never consults it -- AutoDeploy #: goes through ``ADEngine.build_from_config`` and nothing under -#: ``_torch/auto_deploy/`` mentions staircase at all. +#: ``_torch/auto_deploy/`` mentions modeling_v2 at all. ROUTING_BACKENDS = frozenset({"pytorch"}) -class StaircaseMode(str, enum.Enum): - """What to do when a config reaches the staircase resolver.""" +class ModelingV2Mode(str, enum.Enum): + """What to do when a config reaches the modeling_v2 resolver.""" OFF = "off" AUTO = "auto" REQUIRE = "require" @classmethod - def from_env(cls) -> "StaircaseMode": - """Read ``TRTLLM_STAIRCASE``; unset means off. + def from_env(cls) -> "ModelingV2Mode": + """Read ``TRTLLM_MODELING_V2``; unset means off. An unknown value raises rather than falling back. Silently reading a typo as "off" would hand back the built-in implementation while the - caller believed they had asked for a staircase target, and a number + caller believed they had asked for a modeling_v2 target, and a number measured that way is attributed to the wrong system -- the failure the ``require`` mode below exists to prevent, arriving through the door instead of the window. """ - raw = os.environ.get(STAIRCASE_ENV) + raw = os.environ.get(MODELING_V2_ENV) if raw is None or raw == "": return cls.OFF try: return cls(raw.strip().lower()) except ValueError: raise ValueError( - f"{STAIRCASE_ENV}={raw!r} is not a staircase mode; expected " + f"{MODELING_V2_ENV}={raw!r} is not a modeling_v2 mode; expected " f"one of {', '.join(m.value for m in cls)}" ) from None @@ -100,31 +100,31 @@ def from_env(cls) -> "StaircaseMode": def assert_backend_can_route(backend: str) -> None: """Refuse ``require`` on a backend that never reaches the resolver. - ``require`` is a promise that the run measured a staircase target. A + ``require`` is a promise that the run measured a modeling_v2 target. A backend outside ``ROUTING_BACKENDS`` cannot keep it: nothing raises, nothing routes, and the run reports another implementation's numbers as the target's. That is the same misattribution ``from_env`` already refuses for a typo'd value, arriving by a different door. - ``auto`` is left alone on purpose -- it licenses the non-staircase path by + ``auto`` is left alone on purpose -- it licenses the non-modeling_v2 path by definition, so taking it is the documented outcome rather than a silent one. """ - if StaircaseMode.from_env() is not StaircaseMode.REQUIRE: + if ModelingV2Mode.from_env() is not ModelingV2Mode.REQUIRE: return if backend in ROUTING_BACKENDS: return raise ValueError( - f"{STAIRCASE_ENV}=require, but the {backend!r} backend never reaches " - f"the staircase resolver, so no target can be selected and nothing " + f"{MODELING_V2_ENV}=require, but the {backend!r} backend never reaches " + f"the modeling_v2 resolver, so no target can be selected and nothing " f"would report that. Backends that route: " f"{', '.join(sorted(ROUTING_BACKENDS))}. Use one of those, or unset " - f"{STAIRCASE_ENV}." + f"{MODELING_V2_ENV}." ) @dataclass(frozen=True) -class StaircaseContext: +class ModelingV2Context: """Everything a routing decision is allowed to depend on. The admission rule is one line: a quantity may live here only if it is @@ -162,7 +162,7 @@ class StaircaseContext: @classmethod def from_model_config( cls, config: "ModelConfig", sm: Optional[Tuple[int, int]] = None - ) -> "StaircaseContext": + ) -> "ModelingV2Context": """Build the context the resolver routes on. ``sm`` defaults to the device this process will run on, which is what @@ -175,7 +175,7 @@ def from_model_config( import torch assert torch.cuda.is_available(), ( - "staircase routes on the SM version of the device it will run on; " + "modeling_v2 routes on the SM version of the device it will run on; " "no CUDA device is visible" ) sm = torch.cuda.get_device_capability() @@ -231,49 +231,49 @@ def resolve(self, label: str, value: Any, outcome: Any) -> Any: def routing_module(arch: str): """Import and return the routing module for ``arch``, or None.""" - name = STAIRCASE_ROUTERS.get(arch) + name = MODELING_V2_ROUTERS.get(arch) if name is None: return None return import_module(f"{_PACKAGE}.{name}") -def staircase_resolve(config: "ModelConfig") -> Optional[str]: - """Rewrite ``architectures[0]`` into a staircase target's class name. +def modeling_v2_resolve(config: "ModelConfig") -> Optional[str]: + """Rewrite ``architectures[0]`` into a modeling_v2 target's class name. - Driven by ``TRTLLM_STAIRCASE``; see ``STAIRCASE_ENV`` for why it is an + Driven by ``TRTLLM_MODELING_V2``; see ``MODELING_V2_ENV`` for why it is an environment variable and when it has to be set. - Returns None when staircase is off, when no routing module claims the + Returns None when modeling_v2 is off, when no routing module claims the architecture, or when the routing tree finds no matching target -- in ``auto`` the caller then falls back to the built-in implementation. In ``require`` a non-match raises instead, because the failure this mode exists to prevent is silent: asking for a target that does not exist, getting the in-tree implementation, and reading the resulting curve as - staircase's. + modeling_v2's. """ - mode = StaircaseMode.from_env() - if mode is StaircaseMode.OFF: + mode = ModelingV2Mode.from_env() + if mode is ModelingV2Mode.OFF: return None pretrained_config = config.pretrained_config if not getattr(pretrained_config, "architectures", None): return None - ctx = StaircaseContext.from_model_config(config) + ctx = ModelingV2Context.from_model_config(config) arch = pretrained_config.architectures[0] - trace = Trace() if mode is StaircaseMode.REQUIRE else NULL_TRACE + trace = Trace() if mode is ModelingV2Mode.REQUIRE else NULL_TRACE routing = routing_module(arch) target = None if routing is None else routing.route(ctx, trace) if target is None: - if mode is StaircaseMode.REQUIRE: + if mode is ModelingV2Mode.REQUIRE: raise ValueError(explain_no_match(arch, routing, trace)) return None # The synthetic name is only a registry key: importing the target module # is what puts the class behind it. Nothing else would -- the built-in - # static index does not, and must not, carry staircase names. + # static index does not, and must not, carry modeling_v2 names. import_module(f"{_PACKAGE}.{routing.TARGET_MODULES[target]}") return target @@ -281,15 +281,15 @@ def staircase_resolve(config: "ModelConfig") -> Optional[str]: def explain_no_match(arch: str, routing, trace: Trace) -> str: """Say which criterion the configuration failed, not just that it did.""" if routing is None: - known = ", ".join(sorted(STAIRCASE_ROUTERS)) or "(none)" + known = ", ".join(sorted(MODELING_V2_ROUTERS)) or "(none)" return ( - f"staircase is set to 'require' but no staircase target " - f"exists for architecture {arch!r}; routed architectures: " + f"{MODELING_V2_ENV}=require, but no target exists for " + f"architecture {arch!r}; routed architectures: " f"{known}" ) lines = [ - f"staircase is set to 'require' but no target matched " - f"architecture {arch!r}. The decision tree in " + f"{MODELING_V2_ENV}=require, but no target matched architecture " + f"{arch!r}. The decision tree in " f"{routing.__name__.rpartition('.')[0].replace('.', '/')}/routing.py " f"got as far as:" ] diff --git a/tensorrt_llm/_torch/staircase/catalog/__init__.py b/tensorrt_llm/_torch/modeling_v2/catalog/__init__.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/__init__.py rename to tensorrt_llm/_torch/modeling_v2/catalog/__init__.py diff --git a/tensorrt_llm/_torch/staircase/catalog/activation/__init__.py b/tensorrt_llm/_torch/modeling_v2/catalog/activation/__init__.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/activation/__init__.py rename to tensorrt_llm/_torch/modeling_v2/catalog/activation/__init__.py diff --git a/tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul.md b/tensorrt_llm/_torch/modeling_v2/catalog/activation/flashinfer_silu_and_mul.md similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul.md rename to tensorrt_llm/_torch/modeling_v2/catalog/activation/flashinfer_silu_and_mul.md diff --git a/tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul.py b/tensorrt_llm/_torch/modeling_v2/catalog/activation/flashinfer_silu_and_mul.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/activation/flashinfer_silu_and_mul.py rename to tensorrt_llm/_torch/modeling_v2/catalog/activation/flashinfer_silu_and_mul.py diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/__init__.py b/tensorrt_llm/_torch/modeling_v2/catalog/attention/__init__.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/attention/__init__.py rename to tensorrt_llm/_torch/modeling_v2/catalog/attention/__init__.py diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope.md b/tensorrt_llm/_torch/modeling_v2/catalog/attention/fused_qk_norm_rope.md similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope.md rename to tensorrt_llm/_torch/modeling_v2/catalog/attention/fused_qk_norm_rope.md diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope.py b/tensorrt_llm/_torch/modeling_v2/catalog/attention/fused_qk_norm_rope.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/attention/fused_qk_norm_rope.py rename to tensorrt_llm/_torch/modeling_v2/catalog/attention/fused_qk_norm_rope.py diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla.md b/tensorrt_llm/_torch/modeling_v2/catalog/attention/load_paged_kv_cache_for_mla.md similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla.md rename to tensorrt_llm/_torch/modeling_v2/catalog/attention/load_paged_kv_cache_for_mla.md diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla.py b/tensorrt_llm/_torch/modeling_v2/catalog/attention/load_paged_kv_cache_for_mla.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/attention/load_paged_kv_cache_for_mla.py rename to tensorrt_llm/_torch/modeling_v2/catalog/attention/load_paged_kv_cache_for_mla.py diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q.md b/tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_append_paged_kv_assign_q.md similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q.md rename to tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_append_paged_kv_assign_q.md diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q.py b/tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_append_paged_kv_assign_q.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_append_paged_kv_assign_q.py rename to tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_append_paged_kv_assign_q.py diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation.md b/tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_generation.md similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation.md rename to tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_generation.md diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation.py b/tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_generation.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/attention/mla_rope_generation.py rename to tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_generation.py diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/thop_attention.md b/tensorrt_llm/_torch/modeling_v2/catalog/attention/thop_attention.md similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/attention/thop_attention.md rename to tensorrt_llm/_torch/modeling_v2/catalog/attention/thop_attention.md diff --git a/tensorrt_llm/_torch/staircase/catalog/attention/thop_attention.py b/tensorrt_llm/_torch/modeling_v2/catalog/attention/thop_attention.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/attention/thop_attention.py rename to tensorrt_llm/_torch/modeling_v2/catalog/attention/thop_attention.py diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/__init__.py b/tensorrt_llm/_torch/modeling_v2/catalog/comm/__init__.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/comm/__init__.py rename to tensorrt_llm/_torch/modeling_v2/catalog/comm/__init__.py diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/allgather.md b/tensorrt_llm/_torch/modeling_v2/catalog/comm/allgather.md similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/comm/allgather.md rename to tensorrt_llm/_torch/modeling_v2/catalog/comm/allgather.md diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/allgather.py b/tensorrt_llm/_torch/modeling_v2/catalog/comm/allgather.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/comm/allgather.py rename to tensorrt_llm/_torch/modeling_v2/catalog/comm/allgather.py diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/reducescatter.md b/tensorrt_llm/_torch/modeling_v2/catalog/comm/reducescatter.md similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/comm/reducescatter.md rename to tensorrt_llm/_torch/modeling_v2/catalog/comm/reducescatter.md diff --git a/tensorrt_llm/_torch/staircase/catalog/comm/reducescatter.py b/tensorrt_llm/_torch/modeling_v2/catalog/comm/reducescatter.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/comm/reducescatter.py rename to tensorrt_llm/_torch/modeling_v2/catalog/comm/reducescatter.py diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/__init__.py b/tensorrt_llm/_torch/modeling_v2/catalog/gemm/__init__.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/gemm/__init__.py rename to tensorrt_llm/_torch/modeling_v2/catalog/gemm/__init__.py diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out.md b/tensorrt_llm/_torch/modeling_v2/catalog/gemm/bmm_out.md similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out.md rename to tensorrt_llm/_torch/modeling_v2/catalog/gemm/bmm_out.md diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out.py b/tensorrt_llm/_torch/modeling_v2/catalog/gemm/bmm_out.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/gemm/bmm_out.py rename to tensorrt_llm/_torch/modeling_v2/catalog/gemm/bmm_out.py diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm.md b/tensorrt_llm/_torch/modeling_v2/catalog/gemm/cublas_mm.md similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm.md rename to tensorrt_llm/_torch/modeling_v2/catalog/gemm/cublas_mm.md diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm.py b/tensorrt_llm/_torch/modeling_v2/catalog/gemm/cublas_mm.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/gemm/cublas_mm.py rename to tensorrt_llm/_torch/modeling_v2/catalog/gemm/cublas_mm.py diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm.md b/tensorrt_llm/_torch/modeling_v2/catalog/gemm/nvfp4_gemm.md similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm.md rename to tensorrt_llm/_torch/modeling_v2/catalog/gemm/nvfp4_gemm.md diff --git a/tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm.py b/tensorrt_llm/_torch/modeling_v2/catalog/gemm/nvfp4_gemm.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/gemm/nvfp4_gemm.py rename to tensorrt_llm/_torch/modeling_v2/catalog/gemm/nvfp4_gemm.py diff --git a/tensorrt_llm/_torch/staircase/catalog/index.yaml b/tensorrt_llm/_torch/modeling_v2/catalog/index.yaml similarity index 99% rename from tensorrt_llm/_torch/staircase/catalog/index.yaml rename to tensorrt_llm/_torch/modeling_v2/catalog/index.yaml index 594f3dc42482..7382f31d99a1 100644 --- a/tensorrt_llm/_torch/staircase/catalog/index.yaml +++ b/tensorrt_llm/_torch/modeling_v2/catalog/index.yaml @@ -10,14 +10,14 @@ # # AN ENTRY SPANS TWO TREES. The contract (.md) and the wrapper (.py) are here; # the GPU test is at -# tests/unittest/_torch/staircase//test_staircase_.py, because +# tests/unittest/_torch/modeling_v2//test_modeling_v2_.py, because # that is the tree this repo's CI collects from -- a test inside the package is # not picked up by any list. The receipt rule is unchanged in substance and # only gains a second directory to look in: a receipt is valid only if it # post-dates the last write to *every* file of the entry, contract and wrapper # and test alike. Check it mechanically against both paths. # -# The tests are named test_staircase_* rather than test_: three of them +# The tests are named test_modeling_v2_* rather than test_: three of them # would otherwise collide with an upstream test of the same basename, and more # importantly they are not duplicates of those -- upstream covers the modules # that wrap these ops, these cover the op itself cell by cell. diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/__init__.py b/tensorrt_llm/_torch/modeling_v2/catalog/moe/__init__.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/moe/__init__.py rename to tensorrt_llm/_torch/modeling_v2/catalog/moe/__init__.py diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner.md b/tensorrt_llm/_torch/modeling_v2/catalog/moe/fp4_block_scale_moe_runner.md similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner.md rename to tensorrt_llm/_torch/modeling_v2/catalog/moe/fp4_block_scale_moe_runner.md diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner.py b/tensorrt_llm/_torch/modeling_v2/catalog/moe/fp4_block_scale_moe_runner.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/moe/fp4_block_scale_moe_runner.py rename to tensorrt_llm/_torch/modeling_v2/catalog/moe/fp4_block_scale_moe_runner.py diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/fused_moe.md b/tensorrt_llm/_torch/modeling_v2/catalog/moe/fused_moe.md similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/moe/fused_moe.md rename to tensorrt_llm/_torch/modeling_v2/catalog/moe/fused_moe.md diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/fused_moe.py b/tensorrt_llm/_torch/modeling_v2/catalog/moe/fused_moe.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/moe/fused_moe.py rename to tensorrt_llm/_torch/modeling_v2/catalog/moe/fused_moe.py diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md b/tensorrt_llm/_torch/modeling_v2/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md rename to tensorrt_llm/_torch/modeling_v2/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.py b/tensorrt_llm/_torch/modeling_v2/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.py rename to tensorrt_llm/_torch/modeling_v2/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.py diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op.md b/tensorrt_llm/_torch/modeling_v2/catalog/moe/noaux_tc_op.md similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op.md rename to tensorrt_llm/_torch/modeling_v2/catalog/moe/noaux_tc_op.md diff --git a/tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op.py b/tensorrt_llm/_torch/modeling_v2/catalog/moe/noaux_tc_op.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/moe/noaux_tc_op.py rename to tensorrt_llm/_torch/modeling_v2/catalog/moe/noaux_tc_op.py diff --git a/tensorrt_llm/_torch/staircase/catalog/norm/__init__.py b/tensorrt_llm/_torch/modeling_v2/catalog/norm/__init__.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/norm/__init__.py rename to tensorrt_llm/_torch/modeling_v2/catalog/norm/__init__.py diff --git a/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm.md b/tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_fused_add_rmsnorm.md similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm.md rename to tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_fused_add_rmsnorm.md diff --git a/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm.py b/tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_fused_add_rmsnorm.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_fused_add_rmsnorm.py rename to tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_fused_add_rmsnorm.py diff --git a/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm.md b/tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_rmsnorm.md similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm.md rename to tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_rmsnorm.md diff --git a/tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm.py b/tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_rmsnorm.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/norm/flashinfer_rmsnorm.py rename to tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_rmsnorm.py diff --git a/tensorrt_llm/_torch/staircase/catalog/quantization/__init__.py b/tensorrt_llm/_torch/modeling_v2/catalog/quantization/__init__.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/quantization/__init__.py rename to tensorrt_llm/_torch/modeling_v2/catalog/quantization/__init__.py diff --git a/tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize.md b/tensorrt_llm/_torch/modeling_v2/catalog/quantization/fp4_quantize.md similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize.md rename to tensorrt_llm/_torch/modeling_v2/catalog/quantization/fp4_quantize.md diff --git a/tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize.py b/tensorrt_llm/_torch/modeling_v2/catalog/quantization/fp4_quantize.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/quantization/fp4_quantize.py rename to tensorrt_llm/_torch/modeling_v2/catalog/quantization/fp4_quantize.py diff --git a/tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize.md b/tensorrt_llm/_torch/modeling_v2/catalog/quantization/mxfp8_quantize.md similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize.md rename to tensorrt_llm/_torch/modeling_v2/catalog/quantization/mxfp8_quantize.md diff --git a/tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize.py b/tensorrt_llm/_torch/modeling_v2/catalog/quantization/mxfp8_quantize.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/quantization/mxfp8_quantize.py rename to tensorrt_llm/_torch/modeling_v2/catalog/quantization/mxfp8_quantize.py diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/__init__.py b/tensorrt_llm/_torch/modeling_v2/catalog/torch/__init__.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/torch/__init__.py rename to tensorrt_llm/_torch/modeling_v2/catalog/torch/__init__.py diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/add.py b/tensorrt_llm/_torch/modeling_v2/catalog/torch/add.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/torch/add.py rename to tensorrt_llm/_torch/modeling_v2/catalog/torch/add.py diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/concat.py b/tensorrt_llm/_torch/modeling_v2/catalog/torch/concat.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/torch/concat.py rename to tensorrt_llm/_torch/modeling_v2/catalog/torch/concat.py diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/copy_.py b/tensorrt_llm/_torch/modeling_v2/catalog/torch/copy_.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/torch/copy_.py rename to tensorrt_llm/_torch/modeling_v2/catalog/torch/copy_.py diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/embedding.py b/tensorrt_llm/_torch/modeling_v2/catalog/torch/embedding.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/torch/embedding.py rename to tensorrt_llm/_torch/modeling_v2/catalog/torch/embedding.py diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/empty.py b/tensorrt_llm/_torch/modeling_v2/catalog/torch/empty.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/torch/empty.py rename to tensorrt_llm/_torch/modeling_v2/catalog/torch/empty.py diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/expand.py b/tensorrt_llm/_torch/modeling_v2/catalog/torch/expand.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/torch/expand.py rename to tensorrt_llm/_torch/modeling_v2/catalog/torch/expand.py diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/pad.py b/tensorrt_llm/_torch/modeling_v2/catalog/torch/pad.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/torch/pad.py rename to tensorrt_llm/_torch/modeling_v2/catalog/torch/pad.py diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/reshape.py b/tensorrt_llm/_torch/modeling_v2/catalog/torch/reshape.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/torch/reshape.py rename to tensorrt_llm/_torch/modeling_v2/catalog/torch/reshape.py diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/split.py b/tensorrt_llm/_torch/modeling_v2/catalog/torch/split.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/torch/split.py rename to tensorrt_llm/_torch/modeling_v2/catalog/torch/split.py diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/transpose.py b/tensorrt_llm/_torch/modeling_v2/catalog/torch/transpose.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/torch/transpose.py rename to tensorrt_llm/_torch/modeling_v2/catalog/torch/transpose.py diff --git a/tensorrt_llm/_torch/staircase/catalog/torch/view_dtype.py b/tensorrt_llm/_torch/modeling_v2/catalog/torch/view_dtype.py similarity index 100% rename from tensorrt_llm/_torch/staircase/catalog/torch/view_dtype.py rename to tensorrt_llm/_torch/modeling_v2/catalog/torch/view_dtype.py diff --git a/tensorrt_llm/_torch/staircase/explain.py b/tensorrt_llm/_torch/modeling_v2/explain.py similarity index 89% rename from tensorrt_llm/_torch/staircase/explain.py rename to tensorrt_llm/_torch/modeling_v2/explain.py index 87fb96f8ee14..ad35eccb9f03 100644 --- a/tensorrt_llm/_torch/staircase/explain.py +++ b/tensorrt_llm/_torch/modeling_v2/explain.py @@ -2,7 +2,7 @@ # SPDX-License-Identifier: Apache-2.0 """Say which target a configuration routes to, and why. - python -m tensorrt_llm._torch.staircase.explain \ + python -m tensorrt_llm._torch.modeling_v2.explain \ --model /path/to/DeepSeek-R1-0528-NVFP4 --tp 4 --ep 4 --attention-dp Prints the routing module's decision tree as it was actually evaluated, one @@ -23,7 +23,7 @@ from tensorrt_llm.mapping import Mapping -from ._router_index import STAIRCASE_ROUTERS, StaircaseContext, Trace, routing_module +from ._router_index import MODELING_V2_ROUTERS, ModelingV2Context, Trace, routing_module def _sm(value: Optional[str]) -> Tuple[int, int]: @@ -41,7 +41,7 @@ def _sm(value: Optional[str]) -> Tuple[int, int]: def build_parser() -> argparse.ArgumentParser: p = argparse.ArgumentParser( - prog="python -m tensorrt_llm._torch.staircase.explain", + prog="python -m tensorrt_llm._torch.modeling_v2.explain", description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter, ) @@ -79,14 +79,14 @@ def main(argv: Optional[list] = None) -> int: mapping=mapping, moe_backend="AUTO", ) - ctx = StaircaseContext.from_model_config(model_config, sm=_sm(args.sm)) + ctx = ModelingV2Context.from_model_config(model_config, sm=_sm(args.sm)) pretrained_config = model_config.pretrained_config arch = (pretrained_config.architectures or ["(none)"])[0] routing = routing_module(arch) if routing is None: - print(f"{arch} -> no staircase routing module") - print(" routed architectures: " + (", ".join(sorted(STAIRCASE_ROUTERS)) or "(none)")) + print(f"{arch} -> no modeling_v2 routing module") + print(" routed architectures: " + (", ".join(sorted(MODELING_V2_ROUTERS)) or "(none)")) return 1 family = routing.__name__.rpartition(".")[0].rpartition(".")[2] diff --git a/tensorrt_llm/_torch/staircase/models/__init__.py b/tensorrt_llm/_torch/modeling_v2/models/__init__.py similarity index 100% rename from tensorrt_llm/_torch/staircase/models/__init__.py rename to tensorrt_llm/_torch/modeling_v2/models/__init__.py diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/__init__.py b/tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/__init__.py similarity index 100% rename from tensorrt_llm/_torch/staircase/models/deepseek_v3/__init__.py rename to tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/__init__.py diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/routing.py b/tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/routing.py similarity index 88% rename from tensorrt_llm/_torch/staircase/models/deepseek_v3/routing.py rename to tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/routing.py index f6c1b26c512b..cfecce68b334 100644 --- a/tensorrt_llm/_torch/staircase/models/deepseek_v3/routing.py +++ b/tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/routing.py @@ -12,7 +12,7 @@ from typing import Optional -from ..._router_index import NULL_TRACE, StaircaseContext, Trace +from ..._router_index import NULL_TRACE, ModelingV2Context, Trace # The one GPU architecture these targets are written for. sm is part of a # target's identity, not a knob: a different SM is a different target. @@ -31,12 +31,12 @@ } _TARGETS = { - ("r1_0528_nvfp4", "dep4"): "StaircaseDeepseekR10528Nvfp4Sm103Dep4", + ("r1_0528_nvfp4", "dep4"): "ModelingV2DeepseekR10528Nvfp4Sm103Dep4", } # Synthetic architecture name -> the module whose import registers it. TARGET_MODULES = { - "StaircaseDeepseekR10528Nvfp4Sm103Dep4": "models.deepseek_v3.targets.r1_0528_nvfp4.sm_103.dep4.modeling", + "ModelingV2DeepseekR10528Nvfp4Sm103Dep4": "models.deepseek_v3.targets.r1_0528_nvfp4.sm_103.dep4.modeling", } @@ -53,7 +53,7 @@ def _parallel(m) -> Optional[str]: return None -def route(ctx: StaircaseContext, trace: Trace = NULL_TRACE) -> Optional[str]: +def route(ctx: ModelingV2Context, trace: Trace = NULL_TRACE) -> Optional[str]: c, m = ctx.pretrained_config, ctx.mapping if not trace.check("sm", ctx.sm, ctx.sm == _SM): diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/__init__.py b/tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/__init__.py similarity index 100% rename from tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/__init__.py rename to tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/__init__.py diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/__init__.py b/tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/__init__.py similarity index 100% rename from tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/__init__.py rename to tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/__init__.py diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/__init__.py b/tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/__init__.py similarity index 100% rename from tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/__init__.py rename to tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/__init__.py diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/__init__.py b/tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/__init__.py similarity index 100% rename from tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/__init__.py rename to tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/__init__.py diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py b/tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py similarity index 97% rename from tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py rename to tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py index 664546b53113..b9db02c748f6 100644 --- a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py +++ b/tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py @@ -1,6 +1,6 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""Staircase target: deepseek-r1-0528-nvfp4 / sm_103 / dep4 — self-contained modeling code. +"""ModelingV2 target: deepseek-r1-0528-nvfp4 / sm_103 / dep4 — self-contained modeling code. Flat single-entry forward assembled from catalog entries only; every call that creates or transforms a tensor is a catalog entry, everything else is @@ -58,49 +58,51 @@ from tensorrt_llm._torch.attention.backends.interface import AttentionMetadata from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata from tensorrt_llm._torch.model_config import ModelConfig -from tensorrt_llm._torch.models.modeling_utils import ( - DecoderModel, - DecoderModelForCausalLM, - register_auto_model, -) -from tensorrt_llm._torch.speculative import get_spec_worker -from tensorrt_llm._torch.staircase.catalog.activation.flashinfer_silu_and_mul import ( # noqa: E501 +from tensorrt_llm._torch.modeling_v2.catalog.activation.flashinfer_silu_and_mul import ( # noqa: E501 flashinfer_silu_and_mul, ) -from tensorrt_llm._torch.staircase.catalog.attention.load_paged_kv_cache_for_mla import ( # noqa: E501 +from tensorrt_llm._torch.modeling_v2.catalog.attention.load_paged_kv_cache_for_mla import ( # noqa: E501 load_paged_kv_cache_for_mla, ) -from tensorrt_llm._torch.staircase.catalog.attention.mla_rope_append_paged_kv_assign_q import ( # noqa: E501 +from tensorrt_llm._torch.modeling_v2.catalog.attention.mla_rope_append_paged_kv_assign_q import ( # noqa: E501 mla_rope_append_paged_kv_assign_q, ) -from tensorrt_llm._torch.staircase.catalog.attention.mla_rope_generation import mla_rope_generation -from tensorrt_llm._torch.staircase.catalog.attention.thop_attention import thop_attention -from tensorrt_llm._torch.staircase.catalog.comm.allgather import allgather -from tensorrt_llm._torch.staircase.catalog.comm.reducescatter import reducescatter -from tensorrt_llm._torch.staircase.catalog.gemm.bmm_out import bmm_out -from tensorrt_llm._torch.staircase.catalog.gemm.cublas_mm import cublas_mm -from tensorrt_llm._torch.staircase.catalog.gemm.nvfp4_gemm import nvfp4_gemm -from tensorrt_llm._torch.staircase.catalog.moe.fp4_block_scale_moe_runner import ( # noqa: E501 +from tensorrt_llm._torch.modeling_v2.catalog.attention.mla_rope_generation import ( + mla_rope_generation, +) +from tensorrt_llm._torch.modeling_v2.catalog.attention.thop_attention import thop_attention +from tensorrt_llm._torch.modeling_v2.catalog.comm.allgather import allgather +from tensorrt_llm._torch.modeling_v2.catalog.comm.reducescatter import reducescatter +from tensorrt_llm._torch.modeling_v2.catalog.gemm.bmm_out import bmm_out +from tensorrt_llm._torch.modeling_v2.catalog.gemm.cublas_mm import cublas_mm +from tensorrt_llm._torch.modeling_v2.catalog.gemm.nvfp4_gemm import nvfp4_gemm +from tensorrt_llm._torch.modeling_v2.catalog.moe.fp4_block_scale_moe_runner import ( # noqa: E501 fp4_block_scale_moe_runner, ) -from tensorrt_llm._torch.staircase.catalog.moe.fused_moe import fused_moe -from tensorrt_llm._torch.staircase.catalog.moe.noaux_tc_op import noaux_tc_op -from tensorrt_llm._torch.staircase.catalog.norm.flashinfer_fused_add_rmsnorm import ( # noqa: E501 +from tensorrt_llm._torch.modeling_v2.catalog.moe.fused_moe import fused_moe +from tensorrt_llm._torch.modeling_v2.catalog.moe.noaux_tc_op import noaux_tc_op +from tensorrt_llm._torch.modeling_v2.catalog.norm.flashinfer_fused_add_rmsnorm import ( # noqa: E501 flashinfer_fused_add_rmsnorm, ) -from tensorrt_llm._torch.staircase.catalog.norm.flashinfer_rmsnorm import flashinfer_rmsnorm -from tensorrt_llm._torch.staircase.catalog.quantization.fp4_quantize import fp4_quantize -from tensorrt_llm._torch.staircase.catalog.torch.add import add -from tensorrt_llm._torch.staircase.catalog.torch.concat import concat -from tensorrt_llm._torch.staircase.catalog.torch.copy_ import copy_ -from tensorrt_llm._torch.staircase.catalog.torch.embedding import embedding -from tensorrt_llm._torch.staircase.catalog.torch.empty import empty -from tensorrt_llm._torch.staircase.catalog.torch.expand import expand -from tensorrt_llm._torch.staircase.catalog.torch.pad import pad -from tensorrt_llm._torch.staircase.catalog.torch.reshape import reshape -from tensorrt_llm._torch.staircase.catalog.torch.split import split -from tensorrt_llm._torch.staircase.catalog.torch.transpose import transpose -from tensorrt_llm._torch.staircase.catalog.torch.view_dtype import view_dtype +from tensorrt_llm._torch.modeling_v2.catalog.norm.flashinfer_rmsnorm import flashinfer_rmsnorm +from tensorrt_llm._torch.modeling_v2.catalog.quantization.fp4_quantize import fp4_quantize +from tensorrt_llm._torch.modeling_v2.catalog.torch.add import add +from tensorrt_llm._torch.modeling_v2.catalog.torch.concat import concat +from tensorrt_llm._torch.modeling_v2.catalog.torch.copy_ import copy_ +from tensorrt_llm._torch.modeling_v2.catalog.torch.embedding import embedding +from tensorrt_llm._torch.modeling_v2.catalog.torch.empty import empty +from tensorrt_llm._torch.modeling_v2.catalog.torch.expand import expand +from tensorrt_llm._torch.modeling_v2.catalog.torch.pad import pad +from tensorrt_llm._torch.modeling_v2.catalog.torch.reshape import reshape +from tensorrt_llm._torch.modeling_v2.catalog.torch.split import split +from tensorrt_llm._torch.modeling_v2.catalog.torch.transpose import transpose +from tensorrt_llm._torch.modeling_v2.catalog.torch.view_dtype import view_dtype +from tensorrt_llm._torch.models.modeling_utils import ( + DecoderModel, + DecoderModelForCausalLM, + register_auto_model, +) +from tensorrt_llm._torch.speculative import get_spec_worker from . import weights as _weights @@ -112,7 +114,7 @@ #: Every trtllm op this target reaches for, in its forward and in the weight -#: load. Declared here, asserted in tests/unittest/_torch/staircase: a symbol +#: load. Declared here, asserted in tests/unittest/_torch/modeling_v2: a symbol #: that does not exist is a fact of the build, and the place to find that out #: is a machine with the extension built rather than every import of this #: module. @@ -408,7 +410,7 @@ def _yarn_mscale(factor: float, mscale: float) -> float: return 0.1 * mscale * math.log(factor) + 1.0 -class StaircaseCore(DecoderModel): +class ModelingV2Core(DecoderModel): def __init__(self, model_config: ModelConfig): super().__init__(model_config) cfg = model_config.pretrained_config @@ -1607,7 +1609,7 @@ class MTPLayer: `spec_metadata` (the one thing the layer needs off it, the DP padding basis, arrives as the `all_rank_num_tokens` keyword instead).""" - def __init__(self, core: StaircaseCore, logits_processor) -> None: + def __init__(self, core: ModelingV2Core, logits_processor) -> None: self.core = core self.logits_processor = logits_processor @@ -1995,7 +1997,7 @@ class DraftModel: registry by replacing those tensor objects. A reference captured at construction stays on meta and fails at the first draft step.""" - def __init__(self, core: StaircaseCore, lm_head, logits_processor) -> None: + def __init__(self, core: ModelingV2Core, lm_head, logits_processor) -> None: self.core = core self.mtp_layers = [MTPLayer(core, logits_processor)] self.lm_head = lm_head @@ -2005,15 +2007,15 @@ def embed_tokens(self) -> torch.Tensor: return self.core.w["embed"] -@register_auto_model("StaircaseDeepseekR10528Nvfp4Sm103Dep4") -class StaircaseDeepseekR10528Nvfp4Sm103Dep4( - DecoderModelForCausalLM[StaircaseCore, PretrainedConfig] +@register_auto_model("ModelingV2DeepseekR10528Nvfp4Sm103Dep4") +class ModelingV2DeepseekR10528Nvfp4Sm103Dep4( + DecoderModelForCausalLM[ModelingV2Core, PretrainedConfig] ): def __init__(self, model_config: ModelConfig): cfg = model_config.pretrained_config assert cfg is not None super().__init__( - StaircaseCore(model_config), + ModelingV2Core(model_config), config=model_config, hidden_size=cfg.hidden_size, vocab_size=cfg.vocab_size, diff --git a/tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/weights.py b/tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/weights.py similarity index 100% rename from tensorrt_llm/_torch/staircase/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/weights.py rename to tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/weights.py diff --git a/tensorrt_llm/_torch/staircase/models/gpt_oss/__init__.py b/tensorrt_llm/_torch/modeling_v2/models/gpt_oss/__init__.py similarity index 100% rename from tensorrt_llm/_torch/staircase/models/gpt_oss/__init__.py rename to tensorrt_llm/_torch/modeling_v2/models/gpt_oss/__init__.py diff --git a/tensorrt_llm/_torch/staircase/models/gpt_oss/routing.py b/tensorrt_llm/_torch/modeling_v2/models/gpt_oss/routing.py similarity index 86% rename from tensorrt_llm/_torch/staircase/models/gpt_oss/routing.py rename to tensorrt_llm/_torch/modeling_v2/models/gpt_oss/routing.py index 08c2be6c0fe0..d3476f008165 100644 --- a/tensorrt_llm/_torch/staircase/models/gpt_oss/routing.py +++ b/tensorrt_llm/_torch/modeling_v2/models/gpt_oss/routing.py @@ -12,7 +12,7 @@ from typing import Optional -from ..._router_index import NULL_TRACE, StaircaseContext, Trace +from ..._router_index import NULL_TRACE, ModelingV2Context, Trace # The one GPU architecture these targets are written for. sm is part of a # target's identity, not a knob: a different SM is a different target. @@ -32,16 +32,16 @@ } _TARGETS = { - ("gpt_oss_120b", "tp1"): "StaircaseGptOss120bSm103Tp1", + ("gpt_oss_120b", "tp1"): "ModelingV2GptOss120bSm103Tp1", } # Synthetic architecture name -> the module whose import registers it. TARGET_MODULES = { - "StaircaseGptOss120bSm103Tp1": "models.gpt_oss.targets.gpt_oss_120b.sm_103.tp1.modeling", + "ModelingV2GptOss120bSm103Tp1": "models.gpt_oss.targets.gpt_oss_120b.sm_103.tp1.modeling", } -def route(ctx: StaircaseContext, trace: Trace = NULL_TRACE) -> Optional[str]: +def route(ctx: ModelingV2Context, trace: Trace = NULL_TRACE) -> Optional[str]: c, m = ctx.pretrained_config, ctx.mapping if not trace.check("sm", ctx.sm, ctx.sm == _SM): diff --git a/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/__init__.py b/tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/__init__.py similarity index 100% rename from tensorrt_llm/_torch/staircase/models/gpt_oss/targets/__init__.py rename to tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/__init__.py diff --git a/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/__init__.py b/tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/__init__.py similarity index 100% rename from tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/__init__.py rename to tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/__init__.py diff --git a/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/__init__.py b/tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/__init__.py similarity index 100% rename from tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/__init__.py rename to tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/__init__.py diff --git a/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/__init__.py b/tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/__init__.py similarity index 100% rename from tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/__init__.py rename to tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/__init__.py diff --git a/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py b/tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py similarity index 95% rename from tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py rename to tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py index b9f456a00239..44cc94500137 100644 --- a/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py +++ b/tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py @@ -1,6 +1,6 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""Staircase target: gpt-oss-120b / sm_103 / tp1 — self-contained modeling code. +"""ModelingV2 target: gpt-oss-120b / sm_103 / tp1 — self-contained modeling code. Flat single-entry forward assembled from catalog entries only; every call that creates or transforms a tensor is a catalog entry, everything else is @@ -45,25 +45,25 @@ from tensorrt_llm._torch.attention.backends.interface import AttentionMetadata from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata from tensorrt_llm._torch.model_config import ModelConfig +from tensorrt_llm._torch.modeling_v2.catalog.attention.fused_qk_norm_rope import fused_qk_norm_rope +from tensorrt_llm._torch.modeling_v2.catalog.attention.thop_attention import thop_attention +from tensorrt_llm._torch.modeling_v2.catalog.gemm.cublas_mm import cublas_mm +from tensorrt_llm._torch.modeling_v2.catalog.moe.mxe4m3_mxe2m1_block_scale_moe_runner import ( # noqa: E501 + mxe4m3_mxe2m1_block_scale_moe_runner, +) +from tensorrt_llm._torch.modeling_v2.catalog.norm.flashinfer_fused_add_rmsnorm import ( # noqa: E501 + flashinfer_fused_add_rmsnorm, +) +from tensorrt_llm._torch.modeling_v2.catalog.norm.flashinfer_rmsnorm import flashinfer_rmsnorm +from tensorrt_llm._torch.modeling_v2.catalog.quantization.mxfp8_quantize import mxfp8_quantize +from tensorrt_llm._torch.modeling_v2.catalog.torch.embedding import embedding +from tensorrt_llm._torch.modeling_v2.catalog.torch.empty import empty +from tensorrt_llm._torch.modeling_v2.catalog.torch.reshape import reshape from tensorrt_llm._torch.models.modeling_utils import ( DecoderModel, DecoderModelForCausalLM, register_auto_model, ) -from tensorrt_llm._torch.staircase.catalog.attention.fused_qk_norm_rope import fused_qk_norm_rope -from tensorrt_llm._torch.staircase.catalog.attention.thop_attention import thop_attention -from tensorrt_llm._torch.staircase.catalog.gemm.cublas_mm import cublas_mm -from tensorrt_llm._torch.staircase.catalog.moe.mxe4m3_mxe2m1_block_scale_moe_runner import ( # noqa: E501 - mxe4m3_mxe2m1_block_scale_moe_runner, -) -from tensorrt_llm._torch.staircase.catalog.norm.flashinfer_fused_add_rmsnorm import ( # noqa: E501 - flashinfer_fused_add_rmsnorm, -) -from tensorrt_llm._torch.staircase.catalog.norm.flashinfer_rmsnorm import flashinfer_rmsnorm -from tensorrt_llm._torch.staircase.catalog.quantization.mxfp8_quantize import mxfp8_quantize -from tensorrt_llm._torch.staircase.catalog.torch.embedding import embedding -from tensorrt_llm._torch.staircase.catalog.torch.empty import empty -from tensorrt_llm._torch.staircase.catalog.torch.reshape import reshape from . import weights as _weights @@ -75,7 +75,7 @@ #: Every trtllm op this target reaches for, in its forward and in the weight -#: load. Declared here, asserted in tests/unittest/_torch/staircase: a symbol +#: load. Declared here, asserted in tests/unittest/_torch/modeling_v2: a symbol #: that does not exist is a fact of the build, and the place to find that out #: is a machine with the extension built rather than every import of this #: module. @@ -278,7 +278,7 @@ def _pad_up(x: int, align: int) -> int: return (x + align - 1) // align * align -class StaircaseCore(DecoderModel): +class ModelingV2Core(DecoderModel): def __init__(self, model_config: ModelConfig): super().__init__(model_config) cfg = model_config.pretrained_config @@ -643,8 +643,8 @@ def forward( return x -@register_auto_model("StaircaseGptOss120bSm103Tp1") -class StaircaseGptOss120bSm103Tp1(DecoderModelForCausalLM[StaircaseCore, PretrainedConfig]): +@register_auto_model("ModelingV2GptOss120bSm103Tp1") +class ModelingV2GptOss120bSm103Tp1(DecoderModelForCausalLM[ModelingV2Core, PretrainedConfig]): def __init__(self, model_config: ModelConfig): cfg = model_config.pretrained_config assert cfg is not None @@ -657,7 +657,7 @@ def __init__(self, model_config: ModelConfig): if cfg.torch_dtype is None: cfg.torch_dtype = model_config.torch_dtype super().__init__( - StaircaseCore(model_config), + ModelingV2Core(model_config), config=model_config, hidden_size=cfg.hidden_size, vocab_size=cfg.vocab_size, diff --git a/tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/weights.py b/tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/weights.py similarity index 100% rename from tensorrt_llm/_torch/staircase/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/weights.py rename to tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/weights.py diff --git a/tensorrt_llm/_torch/models/modeling_auto.py b/tensorrt_llm/_torch/models/modeling_auto.py index 97fe74e71f3f..96f2d55a430c 100644 --- a/tensorrt_llm/_torch/models/modeling_auto.py +++ b/tensorrt_llm/_torch/models/modeling_auto.py @@ -1,7 +1,7 @@ from typing import Generic, Optional, Type from ..model_config import ModelConfig -from ..staircase import staircase_resolve +from ..modeling_v2 import modeling_v2_resolve from ..utils import model_extra_attrs from .modeling_utils import (DecoderModelForCausalLM, TConfig, TModel, get_registered_model_class, @@ -34,22 +34,22 @@ def _resolve_class(config: ModelConfig) -> Optional[Type]: "") # Strip the appended EAGLE3 model_arch = "EAGLE3" + model_arch - # Staircase targets are keyed by a synthetic architecture name that no + # ModelingV2 targets are keyed by a synthetic architecture name that no # checkpoint declares -- the same shape as the Eagle3 rewrite above. - # Returns None unless `staircase` is on and a target claims this exact - # (checkpoint, GPU arch, parallel topology), so the default path is - # byte-for-byte unchanged. + # Returns None unless `modeling_v2` is on and a target claims this + # exact (checkpoint, GPU arch, parallel topology), so the default path + # is byte-for-byte unchanged. # # Precedence, since this runs last and would override the rewrite - # above: staircase wins. It reads the *un-rewritten* architectures[0], - # so it decides on the checkpoint rather than on what that rewrite made - # of it, and a target that claims a configuration carries that - # configuration's draft path itself. Not reachable today -- Eagle3 - # needs draft_vocab_size, and no draft checkpoint matches a target's - # shape fingerprint -- so this note is the contract, not a description - # of observed behaviour. - if (staircase_arch := staircase_resolve(config)) is not None: - model_arch = staircase_arch + # above: modeling_v2 wins. It reads the *un-rewritten* + # architectures[0], so it decides on the checkpoint rather than on what + # that rewrite made of it, and a target that claims a configuration + # carries that configuration's draft path itself. Not reachable today + # -- Eagle3 needs draft_vocab_size, and no draft checkpoint matches a + # target's shape fingerprint -- so this note is the contract, not a + # description of observed behaviour. + if (modeling_v2_arch := modeling_v2_resolve(config)) is not None: + model_arch = modeling_v2_arch return get_registered_model_class(model_arch) diff --git a/tensorrt_llm/llmapi/llm.py b/tensorrt_llm/llmapi/llm.py index 04ede393189f..00bdad277a1f 100644 --- a/tensorrt_llm/llmapi/llm.py +++ b/tensorrt_llm/llmapi/llm.py @@ -393,11 +393,11 @@ def __init__(self, f"Unknown backend: {backend!r}. Supported backends are " "'pytorch'.") - # TRTLLM_STAIRCASE=require promises the run measured a staircase + # TRTLLM_MODELING_V2=require promises the run measured a modeling_v2 # target. Only the pytorch backend reaches the resolver that could # select one, so on any other backend the promise would be broken # silently -- the one failure that mode exists to prevent. - from .._torch.staircase import assert_backend_can_route + from .._torch.modeling_v2 import assert_backend_can_route assert_backend_can_route(backend) # check the kwargs and raise ValueError directly diff --git a/tests/integration/defs/accuracy/references/acceptance_length.yaml b/tests/integration/defs/accuracy/references/acceptance_length.yaml index c206436db209..7a5fce89a575 100644 --- a/tests/integration/defs/accuracy/references/acceptance_length.yaml +++ b/tests/integration/defs/accuracy/references/acceptance_length.yaml @@ -96,23 +96,23 @@ disagg::TestNemotron3Super120B::test_auto_dtype: TestQwen3_8_27B::test_dflash2: ref_al: 5.899181450085933 min_al: 5.604222377581636 -# Shared by both legs of TestStaircaseDeepseekR10528Nvfp4Sm103Dep4:: +# Shared by both legs of TestModelingV2DeepseekR10528Nvfp4Sm103Dep4:: # test_mtp3_acceptance. Keyed by target and variant rather than by test # function, because that is what the number is a property of, and because the -# staircase and stock legs are read against the same minimum on purpose: -# staircase failing while stock passes means the draft path regressed; both +# modeling_v2 and stock legs are read against the same minimum on purpose: +# modeling_v2 failing while stock passes means the draft path regressed; both # failing means this anchor is stale. Populated from the STOCK leg -- taking it # from the target's own number would make the gate self-referential. -StaircaseDeepseekR10528Nvfp4Sm103Dep4::mtp3: +ModelingV2DeepseekR10528Nvfp4Sm103Dep4::mtp3: ref_al: 3.020 # Hand-set, not the automatic 95%. Measured on GB300 over three runs each: - # stock 3.020 / 2.998 / 3.009 / 3.014, staircase 2.921 / 2.938 / 2.933. Each - # side is stable to ~0.7%, and the two ranges do not overlap -- staircase + # stock 3.020 / 2.998 / 3.009 / 3.014, modeling_v2 2.921 / 2.938 / 2.933. Each + # side is stable to ~0.7%, and the two ranges do not overlap -- modeling_v2 # runs about 2.6% under stock on this workload, consistently rather than as # noise. The automatic 95% floor (2.869) would therefore have left the - # staircase leg 1.8% of headroom against a 0.7% spread: a flaky test, not a + # modeling_v2 leg 1.8% of headroom against a 0.7% spread: a flaky test, not a # gate. 2.5 is a collapse tripwire instead, the same role the # TestKimiK3DSpark entry above states -- a collapsed draft path scores ~1.0 # and a subtly wrong one ~2.1, both far below this, while every accuracy - # gate stays green. It leaves staircase 17% of headroom and stock 20%. + # gate stays green. It leaves modeling_v2 17% of headroom and stock 20%. min_al: 2.5 diff --git a/tests/integration/defs/accuracy/test_staircase_deepseek_v3.py b/tests/integration/defs/accuracy/test_modeling_v2_deepseek_v3.py similarity index 90% rename from tests/integration/defs/accuracy/test_staircase_deepseek_v3.py rename to tests/integration/defs/accuracy/test_modeling_v2_deepseek_v3.py index c0b56df14680..ccc58e2830f9 100644 --- a/tests/integration/defs/accuracy/test_staircase_deepseek_v3.py +++ b/tests/integration/defs/accuracy/test_modeling_v2_deepseek_v3.py @@ -12,14 +12,14 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. -"""Whole-model gates for the deepseek-v3 staircase targets. +"""Whole-model gates for the deepseek-v3 modeling_v2 targets. One file per model family, beside the other accuracy suites rather than inside test_llm_api_pytorch.py: these gate a parallel implementation, and reading them next to the built-in model's tests would invite treating one as a variant of the other. -Every test here needs ``TRTLLM_STAIRCASE=require``. Under ``"auto"`` a +Every test here needs ``TRTLLM_MODELING_V2=require``. Under ``"auto"`` a configuration that missed a target's criteria would quietly fall back to the built-in implementation, pass, and report the built-in's numbers as the target's -- which is the one failure this whole system exists to prevent. The @@ -38,7 +38,7 @@ import pytest from tensorrt_llm import LLM -from tensorrt_llm._torch.staircase import STAIRCASE_ENV +from tensorrt_llm._torch.modeling_v2 import MODELING_V2_ENV from tensorrt_llm._utils import get_sm_version from tensorrt_llm.llmapi import CudaGraphConfig, KvCacheConfig, MTPDecodingConfig @@ -53,10 +53,10 @@ # The targets assert their own SM at construction: certification is per GPU # architecture, and a receipt from another one says nothing here. skip_not_sm103 = pytest.mark.skipif( - get_sm_version() != 103, reason="staircase targets in this batch are certified on sm_103 only" + get_sm_version() != 103, reason="modeling_v2 targets in this batch are certified on sm_103 only" ) -# The filter the staircase anchors were measured on. Unset, the evaluator +# The filter the modeling_v2 anchors were measured on. Unset, the evaluator # averages every metric GSM8K reports, which means the mean of strict-match and # flexible-extract -- two numbers measuring different things on a checkpoint # that does not answer purely in the strict "#### N" form. @@ -71,12 +71,12 @@ def _require_mode(expected: str) -> None: must not silently measure the wrong system either, which is what reading the variable here rules out. """ - actual = os.environ.get(STAIRCASE_ENV, "off") + actual = os.environ.get(MODELING_V2_ENV, "off") if actual != expected: - pytest.skip(f"{STAIRCASE_ENV}={actual!r}, this case needs {expected!r}") + pytest.skip(f"{MODELING_V2_ENV}={actual!r}, this case needs {expected!r}") -class TestStaircaseDeepseekR10528Nvfp4Sm103Dep4(LlmapiAccuracyTestHarness): +class TestModelingV2DeepseekR10528Nvfp4Sm103Dep4(LlmapiAccuracyTestHarness): """deepseek-r1-0528-nvfp4 / sm_103 / dep4, identity and the mtp3 variant.""" MODEL_NAME = "deepseek-ai/DeepSeek-R1-0528" @@ -96,7 +96,7 @@ class TestStaircaseDeepseekR10528Nvfp4Sm103Dep4(LlmapiAccuracyTestHarness): # One anchor shared by both legs of the acceptance gate below. It names the # target and variant, not a test function, because that is what the number # is a property of. - ACCEPTANCE_KEY = "StaircaseDeepseekR10528Nvfp4Sm103Dep4::mtp3" + ACCEPTANCE_KEY = "ModelingV2DeepseekR10528Nvfp4Sm103Dep4::mtp3" # There is deliberately no standalone identity gsm8k case. The paired test # below evaluates the identity config as its first leg, and @@ -140,7 +140,7 @@ def test_gsm8k_identity_vs_mtp3(self, mocker): mtp3 = task.evaluate(llm) delta = mtp3 - identity - print(f"[staircase] gsm8k identity={identity:.4f} mtp3={mtp3:.4f} delta={delta:+.4f}") + print(f"[modeling_v2] gsm8k identity={identity:.4f} mtp3={mtp3:.4f} delta={delta:+.4f}") # 2 sigma at the ~0.6 stderr this benchmark reports at n=1319. assert abs(delta) < 1.2, ( f"MTP moved gsm8k by {delta:+.4f} (identity={identity:.4f}, " @@ -150,7 +150,7 @@ def test_gsm8k_identity_vs_mtp3(self, mocker): @skip_not_sm103 @pytest.mark.skip_less_device(4) - @pytest.mark.parametrize("mode", ["require", "off"], ids=["staircase", "stock"]) + @pytest.mark.parametrize("mode", ["require", "off"], ids=["modeling_v2", "stock"]) def test_mtp3_acceptance(self, mode, mocker): """The only gate that can see a miscomputed draft layer. @@ -162,7 +162,7 @@ def test_mtp3_acceptance(self, mode, mocker): They share ``ACCEPTANCE_KEY``, so both are read against the same recorded minimum, which is what makes the pair informative: - staircase fails, stock passes -> the draft path regressed + modeling_v2 fails, stock passes -> the draft path regressed both fail -> the anchor is stale; re-derive it rather than blaming the target @@ -176,7 +176,7 @@ def test_mtp3_acceptance(self, mode, mocker): # DeepEPLowLatency, whose dispatch takes only NVFP4 uint8 hidden # states, and the MTP layer is bf16 because modelopt excludes # model.layers.61* from quantization. Disabling DeepEP lands on - # AllGatherReduceScatter -- which is the strategy the staircase + # AllGatherReduceScatter -- which is the strategy the modeling_v2 # target implements by hand, so it makes the two comparable rather # than less so. Set in the launching environment, like the switch. assert os.environ.get("TRTLLM_CAN_USE_DEEP_EP") == "0", ( diff --git a/tests/integration/defs/accuracy/test_staircase_gpt_oss.py b/tests/integration/defs/accuracy/test_modeling_v2_gpt_oss.py similarity index 87% rename from tests/integration/defs/accuracy/test_staircase_gpt_oss.py rename to tests/integration/defs/accuracy/test_modeling_v2_gpt_oss.py index 2d0292e5dffe..1905a698409f 100644 --- a/tests/integration/defs/accuracy/test_staircase_gpt_oss.py +++ b/tests/integration/defs/accuracy/test_modeling_v2_gpt_oss.py @@ -12,14 +12,14 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. -"""Whole-model gate for the gpt-oss staircase targets. +"""Whole-model gate for the gpt-oss modeling_v2 targets. One file per model family, beside the other accuracy suites rather than inside test_llm_api_pytorch.py: these gate a parallel implementation, and reading them next to the built-in model's tests would invite treating one as a variant of the other. -Every test here needs ``TRTLLM_STAIRCASE=require``. Under ``"auto"`` a +Every test here needs ``TRTLLM_MODELING_V2=require``. Under ``"auto"`` a configuration that missed a target's criteria would quietly fall back to the built-in implementation, pass, and report the built-in's numbers as the target's -- which is the one failure this whole system exists to prevent. @@ -30,7 +30,7 @@ import pytest from tensorrt_llm import LLM -from tensorrt_llm._torch.staircase import STAIRCASE_ENV +from tensorrt_llm._torch.modeling_v2 import MODELING_V2_ENV from tensorrt_llm._utils import get_sm_version from ..conftest import llm_models_root @@ -39,7 +39,7 @@ # The targets assert their own SM at construction: certification is per GPU # architecture, and a receipt from another one says nothing here. skip_not_sm103 = pytest.mark.skipif( - get_sm_version() != 103, reason="staircase targets in this batch are certified on sm_103 only" + get_sm_version() != 103, reason="modeling_v2 targets in this batch are certified on sm_103 only" ) @@ -51,12 +51,12 @@ def _require_mode(expected: str) -> None: must not silently measure the wrong system either, which is what reading the variable here rules out. """ - actual = os.environ.get(STAIRCASE_ENV, "off") + actual = os.environ.get(MODELING_V2_ENV, "off") if actual != expected: - pytest.skip(f"{STAIRCASE_ENV}={actual!r}, this case needs {expected!r}") + pytest.skip(f"{MODELING_V2_ENV}={actual!r}, this case needs {expected!r}") -class TestStaircaseGptOss120bSm103Tp1(LlmapiAccuracyTestHarness): +class TestModelingV2GptOss120bSm103Tp1(LlmapiAccuracyTestHarness): """gpt-oss-120b / sm_103 / tp1.""" # The registry key upstream uses for this checkpoint; it carries the diff --git a/tests/integration/test_lists/test-db/l0_gb300.yml b/tests/integration/test_lists/test-db/l0_gb300.yml index eb405d87e626..12790fb33ad0 100644 --- a/tests/integration/test_lists/test-db/l0_gb300.yml +++ b/tests/integration/test_lists/test-db/l0_gb300.yml @@ -26,18 +26,18 @@ l0_gb300: - accuracy/test_disaggregated_serving.py::TestQwen3_8_Flash_Next::test_fp8_nixl_python[prefix_cache] - unittest/_torch/thop/parallel TIMEOUT (90) - unittest/_torch/visual_gen/kernels/parallel - # Staircase catalog: the certification matrix for each op a staircase target + # ModelingV2 catalog: the certification matrix for each op a modeling_v2 target # calls, plus the two consistency tests that guard the routing tables against # the targets they name. Deliberately NOT a duplicate of the upstream tests # for the same ops -- those cover the modules that wrap them, these cover the - # op itself cell by cell -- which is why every file is named test_staircase_*. + # op itself cell by cell -- which is why every file is named test_modeling_v2_*. # These entries are the receipts: the catalog's certification is per GPU # architecture, and this list is the sm_103 one. comm/ is excluded here and # carried by l0_gb300_multi_gpus.yml, since those two need 4 ranks. # TIMEOUT measured, not guessed: the whole entry is 348 cases in 5m34s on one # GB300. 30 gives five times that -- enough for a loaded node, short enough # that a hung case does not hold this stage for its full budget. - - unittest/_torch/staircase --ignore=unittest/_torch/staircase/comm TIMEOUT (30) - # Staircase whole-model gate. The op-level entry above certifies the + - unittest/_torch/modeling_v2 --ignore=unittest/_torch/modeling_v2/comm TIMEOUT (30) + # ModelingV2 whole-model gate. The op-level entry above certifies the # vocabulary; this certifies the assembly that calls it. - - accuracy/test_staircase_gpt_oss.py::TestStaircaseGptOss120bSm103Tp1::test_gsm8k + - accuracy/test_modeling_v2_gpt_oss.py::TestModelingV2GptOss120bSm103Tp1::test_gsm8k diff --git a/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml b/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml index 17239f6fd984..88d382b928d5 100644 --- a/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml +++ b/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml @@ -15,22 +15,22 @@ l0_gb300_multi_gpus: stage: post_merge backend: pytorch tests: - # Staircase collective entries. Each spawns its own 4-rank job rather than + # ModelingV2 collective entries. Each spawns its own 4-rank job rather than # using the mpi_pool_executor fixture: their cases read module-global rank # state and several assert on communicator state the previous case left # behind, so they are one ordered sequence inside one job, and a broken # collective hangs rather than raising -- the launcher's deadline is what # keeps that from wedging the run, and the fixture has none. - - unittest/_torch/staircase/comm - # Staircase whole-model gates for the dep4 target. The mtp3 variant selects a + - unittest/_torch/modeling_v2/comm + # ModelingV2 whole-model gates for the dep4 target. The mtp3 variant selects a # second forward path AND a second weight-loading path, so the identity gate # does not speak for it: it carries its own accuracy pairing and its own # acceptance gate. The two acceptance legs are independent cases read against - # one shared minimum -- staircase failing while stock passes means the draft + # one shared minimum -- modeling_v2 failing while stock passes means the draft # path regressed; both failing means the anchor is stale. - - accuracy/test_staircase_deepseek_v3.py::TestStaircaseDeepseekR10528Nvfp4Sm103Dep4::test_gsm8k_identity_vs_mtp3 - - accuracy/test_staircase_deepseek_v3.py::TestStaircaseDeepseekR10528Nvfp4Sm103Dep4::test_mtp3_acceptance[staircase] - - accuracy/test_staircase_deepseek_v3.py::TestStaircaseDeepseekR10528Nvfp4Sm103Dep4::test_mtp3_acceptance[stock] + - accuracy/test_modeling_v2_deepseek_v3.py::TestModelingV2DeepseekR10528Nvfp4Sm103Dep4::test_gsm8k_identity_vs_mtp3 + - accuracy/test_modeling_v2_deepseek_v3.py::TestModelingV2DeepseekR10528Nvfp4Sm103Dep4::test_mtp3_acceptance[modeling_v2] + - accuracy/test_modeling_v2_deepseek_v3.py::TestModelingV2DeepseekR10528Nvfp4Sm103Dep4::test_mtp3_acceptance[stock] # Covers tests/unittest/_torch/attention/. Two sub-trees moved in here from elsewhere # under tests/unittest/_torch/ and have no entry of their own on any list, so this entry # is what picks them up: diff --git a/tests/unittest/_torch/staircase/activation/test_staircase_flashinfer_silu_and_mul.py b/tests/unittest/_torch/modeling_v2/activation/test_modeling_v2_flashinfer_silu_and_mul.py similarity index 97% rename from tests/unittest/_torch/staircase/activation/test_staircase_flashinfer_silu_and_mul.py rename to tests/unittest/_torch/modeling_v2/activation/test_modeling_v2_flashinfer_silu_and_mul.py index 5beefaf8cdcb..ea2dfc783cda 100644 --- a/tests/unittest/_torch/staircase/activation/test_staircase_flashinfer_silu_and_mul.py +++ b/tests/unittest/_torch/modeling_v2/activation/test_modeling_v2_flashinfer_silu_and_mul.py @@ -5,7 +5,7 @@ import torch import torch.nn.functional as F -from tensorrt_llm._torch.staircase.catalog.activation.flashinfer_silu_and_mul import ( +from tensorrt_llm._torch.modeling_v2.catalog.activation.flashinfer_silu_and_mul import ( flashinfer_silu_and_mul, ) diff --git a/tests/unittest/_torch/staircase/attention/test_staircase_fused_qk_norm_rope.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_fused_qk_norm_rope.py similarity index 98% rename from tests/unittest/_torch/staircase/attention/test_staircase_fused_qk_norm_rope.py rename to tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_fused_qk_norm_rope.py index 62e8c5462490..bfa85580bb2c 100644 --- a/tests/unittest/_torch/staircase/attention/test_staircase_fused_qk_norm_rope.py +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_fused_qk_norm_rope.py @@ -4,7 +4,7 @@ import torch -from tensorrt_llm._torch.staircase.catalog.attention.fused_qk_norm_rope import fused_qk_norm_rope +from tensorrt_llm._torch.modeling_v2.catalog.attention.fused_qk_norm_rope import fused_qk_norm_rope assert torch.cuda.is_available(), "fused_qk_norm_rope requires a CUDA device" diff --git a/tests/unittest/_torch/staircase/attention/test_staircase_load_paged_kv_cache_for_mla.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_load_paged_kv_cache_for_mla.py similarity index 99% rename from tests/unittest/_torch/staircase/attention/test_staircase_load_paged_kv_cache_for_mla.py rename to tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_load_paged_kv_cache_for_mla.py index 9eea002561ae..b06197c1a7f5 100644 --- a/tests/unittest/_torch/staircase/attention/test_staircase_load_paged_kv_cache_for_mla.py +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_load_paged_kv_cache_for_mla.py @@ -34,10 +34,10 @@ from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata from tensorrt_llm._torch.metadata import KVCacheParams -from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager -from tensorrt_llm._torch.staircase.catalog.attention.load_paged_kv_cache_for_mla import ( +from tensorrt_llm._torch.modeling_v2.catalog.attention.load_paged_kv_cache_for_mla import ( load_paged_kv_cache_for_mla, ) +from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager from tensorrt_llm.bindings import DataType from tensorrt_llm.bindings.internal.batch_manager import CacheType from tensorrt_llm.llmapi.llm_args import KvCacheConfig diff --git a/tests/unittest/_torch/staircase/attention/test_staircase_mla_rope_append_paged_kv_assign_q.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_mla_rope_append_paged_kv_assign_q.py similarity index 99% rename from tests/unittest/_torch/staircase/attention/test_staircase_mla_rope_append_paged_kv_assign_q.py rename to tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_mla_rope_append_paged_kv_assign_q.py index a3953ee49fcc..2b7560a11e98 100644 --- a/tests/unittest/_torch/staircase/attention/test_staircase_mla_rope_append_paged_kv_assign_q.py +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_mla_rope_append_paged_kv_assign_q.py @@ -39,10 +39,10 @@ from tensorrt_llm._torch.attention.backends.interface import RopeParams from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata from tensorrt_llm._torch.metadata import KVCacheParams -from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager -from tensorrt_llm._torch.staircase.catalog.attention.mla_rope_append_paged_kv_assign_q import ( +from tensorrt_llm._torch.modeling_v2.catalog.attention.mla_rope_append_paged_kv_assign_q import ( mla_rope_append_paged_kv_assign_q, ) +from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager from tensorrt_llm.bindings import DataType from tensorrt_llm.bindings.internal.batch_manager import CacheType from tensorrt_llm.llmapi.llm_args import KvCacheConfig diff --git a/tests/unittest/_torch/staircase/attention/test_staircase_mla_rope_generation.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_mla_rope_generation.py similarity index 99% rename from tests/unittest/_torch/staircase/attention/test_staircase_mla_rope_generation.py rename to tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_mla_rope_generation.py index 1e2bd8055d7c..b8ed67b35184 100644 --- a/tests/unittest/_torch/staircase/attention/test_staircase_mla_rope_generation.py +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_mla_rope_generation.py @@ -51,8 +51,10 @@ from tensorrt_llm._torch.attention.backends.interface import RopeParams from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata from tensorrt_llm._torch.metadata import KVCacheParams +from tensorrt_llm._torch.modeling_v2.catalog.attention.mla_rope_generation import ( + mla_rope_generation, +) from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager -from tensorrt_llm._torch.staircase.catalog.attention.mla_rope_generation import mla_rope_generation from tensorrt_llm.bindings import DataType from tensorrt_llm.bindings.internal.batch_manager import CacheType from tensorrt_llm.llmapi.llm_args import KvCacheConfig diff --git a/tests/unittest/_torch/staircase/attention/test_staircase_thop_attention.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_thop_attention.py similarity index 99% rename from tests/unittest/_torch/staircase/attention/test_staircase_thop_attention.py rename to tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_thop_attention.py index 23fcaecaf3b9..19f04f8e931c 100644 --- a/tests/unittest/_torch/staircase/attention/test_staircase_thop_attention.py +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_thop_attention.py @@ -173,8 +173,8 @@ import torch.nn.functional as F from tensorrt_llm._torch.attention.backends.interface import RopeParams +from tensorrt_llm._torch.modeling_v2.catalog.attention.thop_attention import thop_attention from tensorrt_llm._torch.pyexecutor.resource_manager import CacheTypeCpp, DataType, KVCacheManager -from tensorrt_llm._torch.staircase.catalog.attention.thop_attention import thop_attention from tensorrt_llm.functional import RotaryScalingType from tensorrt_llm.llmapi.llm_args import KvCacheConfig from tensorrt_llm.mapping import Mapping diff --git a/tests/unittest/_torch/staircase/comm/_allgather_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_allgather_op_matrix.py similarity index 99% rename from tests/unittest/_torch/staircase/comm/_allgather_op_matrix.py rename to tests/unittest/_torch/modeling_v2/comm/_allgather_op_matrix.py index 0e37668615eb..94e8e0add9f6 100644 --- a/tests/unittest/_torch/staircase/comm/_allgather_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_allgather_op_matrix.py @@ -26,7 +26,7 @@ keeps it uncollectable however pytest is pointed at this tree. The collected entry point is -`tests/unittest/_torch/staircase/comm/test_staircase_allgather_op_matrix.py`: +`tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_allgather_op_matrix.py`: it starts this job and turns its exit code into an assertion. Both halves are started by file path: the launcher must not import `tensorrt_llm` (that calls `MPI_Init`, and an MPI-initialized process cannot start `mpirun`), and the @@ -917,7 +917,7 @@ def _run_one_rank() -> int: from mpi4py import MPI as _MPI from tensorrt_llm._torch.distributed import Distributed - from tensorrt_llm._torch.staircase.catalog.comm import allgather as entry + from tensorrt_llm._torch.modeling_v2.catalog.comm import allgather as entry from tensorrt_llm.mapping import Mapping allgather = entry.allgather diff --git a/tests/unittest/_torch/staircase/comm/_rank_job.py b/tests/unittest/_torch/modeling_v2/comm/_rank_job.py similarity index 98% rename from tests/unittest/_torch/staircase/comm/_rank_job.py rename to tests/unittest/_torch/modeling_v2/comm/_rank_job.py index 271f976f75d0..a24b842d6f45 100644 --- a/tests/unittest/_torch/staircase/comm/_rank_job.py +++ b/tests/unittest/_torch/modeling_v2/comm/_rank_job.py @@ -73,7 +73,7 @@ def run(entry: str) -> None: env = dict(os.environ, CUDA_VISIBLE_DEVICES=_devices()) # The ranks import tensorrt_llm absolutely, and a source checkout is not - # necessarily installed. tests/unittest/_torch/staircase/comm -> repo root. + # necessarily installed. tests/unittest/_torch/modeling_v2/comm -> repo root. repo_root = Path(__file__).resolve().parents[5] assert (repo_root / "tensorrt_llm").is_dir(), ( f"expected the repo root at {repo_root}, found no tensorrt_llm/ there; " diff --git a/tests/unittest/_torch/staircase/comm/_reducescatter_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_reducescatter_op_matrix.py similarity index 99% rename from tests/unittest/_torch/staircase/comm/_reducescatter_op_matrix.py rename to tests/unittest/_torch/modeling_v2/comm/_reducescatter_op_matrix.py index 3e76415d85b8..aea196658056 100644 --- a/tests/unittest/_torch/staircase/comm/_reducescatter_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_reducescatter_op_matrix.py @@ -32,7 +32,7 @@ keeps it uncollectable however pytest is pointed at this tree. The collected entry point is -`tests/unittest/_torch/staircase/comm/test_staircase_reducescatter_op_matrix.py`: +`tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_reducescatter_op_matrix.py`: it starts this job and turns its exit code into an assertion. Both halves are started by file path: the launcher must not import `tensorrt_llm` (that calls `MPI_Init`, and an MPI-initialized process cannot start `mpirun`), and the @@ -1430,7 +1430,7 @@ def _run_one_rank() -> int: global COMM, RANK, WORLD, GROUP, reducescatter from mpi4py import MPI - from tensorrt_llm._torch.staircase.catalog.comm import reducescatter as entry + from tensorrt_llm._torch.modeling_v2.catalog.comm import reducescatter as entry reducescatter = entry.reducescatter COMM = MPI.COMM_WORLD @@ -1497,7 +1497,7 @@ def _run_wedge_rank() -> int: global COMM, RANK, WORLD, GROUP, reducescatter from mpi4py import MPI - from tensorrt_llm._torch.staircase.catalog.comm import reducescatter as entry + from tensorrt_llm._torch.modeling_v2.catalog.comm import reducescatter as entry reducescatter = entry.reducescatter COMM = MPI.COMM_WORLD diff --git a/tests/unittest/_torch/staircase/comm/test_staircase_allgather_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_allgather_op_matrix.py similarity index 100% rename from tests/unittest/_torch/staircase/comm/test_staircase_allgather_op_matrix.py rename to tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_allgather_op_matrix.py diff --git a/tests/unittest/_torch/staircase/comm/test_staircase_reducescatter_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_reducescatter_op_matrix.py similarity index 100% rename from tests/unittest/_torch/staircase/comm/test_staircase_reducescatter_op_matrix.py rename to tests/unittest/_torch/modeling_v2/comm/test_modeling_v2_reducescatter_op_matrix.py diff --git a/tests/unittest/_torch/staircase/gemm/test_staircase_bmm_out.py b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_bmm_out.py similarity index 97% rename from tests/unittest/_torch/staircase/gemm/test_staircase_bmm_out.py rename to tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_bmm_out.py index 87f000b87c68..fefc1067e6d5 100644 --- a/tests/unittest/_torch/staircase/gemm/test_staircase_bmm_out.py +++ b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_bmm_out.py @@ -4,7 +4,7 @@ import torch -from tensorrt_llm._torch.staircase.catalog.gemm.bmm_out import bmm_out +from tensorrt_llm._torch.modeling_v2.catalog.gemm.bmm_out import bmm_out assert torch.cuda.is_available(), "bmm_out requires a CUDA device" diff --git a/tests/unittest/_torch/staircase/gemm/test_staircase_cublas_mm.py b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_cublas_mm.py similarity index 98% rename from tests/unittest/_torch/staircase/gemm/test_staircase_cublas_mm.py rename to tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_cublas_mm.py index a39191e8f59b..4c88737300d8 100644 --- a/tests/unittest/_torch/staircase/gemm/test_staircase_cublas_mm.py +++ b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_cublas_mm.py @@ -6,7 +6,7 @@ import torch -from tensorrt_llm._torch.staircase.catalog.gemm.cublas_mm import cublas_mm +from tensorrt_llm._torch.modeling_v2.catalog.gemm.cublas_mm import cublas_mm assert torch.cuda.is_available(), "cublas_mm requires a CUDA device" diff --git a/tests/unittest/_torch/staircase/gemm/test_staircase_nvfp4_gemm.py b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_nvfp4_gemm.py similarity index 99% rename from tests/unittest/_torch/staircase/gemm/test_staircase_nvfp4_gemm.py rename to tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_nvfp4_gemm.py index ed388b240e2a..08b2d5db8342 100644 --- a/tests/unittest/_torch/staircase/gemm/test_staircase_nvfp4_gemm.py +++ b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_nvfp4_gemm.py @@ -6,7 +6,7 @@ import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* from tensorrt_llm._torch.autotuner import AutoTuner, autotune -from tensorrt_llm._torch.staircase.catalog.gemm.nvfp4_gemm import nvfp4_gemm +from tensorrt_llm._torch.modeling_v2.catalog.gemm.nvfp4_gemm import nvfp4_gemm assert torch.cuda.is_available(), "nvfp4_gemm requires a CUDA device" # The reference matmul must be true fp32, never tf32. diff --git a/tests/unittest/_torch/staircase/moe/test_staircase_fp4_block_scale_moe_runner.py b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_fp4_block_scale_moe_runner.py similarity index 99% rename from tests/unittest/_torch/staircase/moe/test_staircase_fp4_block_scale_moe_runner.py rename to tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_fp4_block_scale_moe_runner.py index d0c488c034d8..5c6ba3b33468 100644 --- a/tests/unittest/_torch/staircase/moe/test_staircase_fp4_block_scale_moe_runner.py +++ b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_fp4_block_scale_moe_runner.py @@ -5,7 +5,7 @@ import torch from tensorrt_llm._torch.autotuner import AutoTuner, autotune -from tensorrt_llm._torch.staircase.catalog.moe.fp4_block_scale_moe_runner import ( +from tensorrt_llm._torch.modeling_v2.catalog.moe.fp4_block_scale_moe_runner import ( fp4_block_scale_moe_runner as moe, ) diff --git a/tests/unittest/_torch/staircase/moe/test_staircase_fused_moe.py b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_fused_moe.py similarity index 99% rename from tests/unittest/_torch/staircase/moe/test_staircase_fused_moe.py rename to tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_fused_moe.py index 078236f81523..7b56c3c0f93d 100644 --- a/tests/unittest/_torch/staircase/moe/test_staircase_fused_moe.py +++ b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_fused_moe.py @@ -6,7 +6,7 @@ import torch.nn.functional as F from tensorrt_llm._torch.autotuner import AutoTuner, autotune -from tensorrt_llm._torch.staircase.catalog.moe.fused_moe import fused_moe +from tensorrt_llm._torch.modeling_v2.catalog.moe.fused_moe import fused_moe assert torch.cuda.is_available(), "fused_moe requires a CUDA device" diff --git a/tests/unittest/_torch/staircase/moe/test_staircase_mxe4m3_mxe2m1_block_scale_moe_runner.py b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_mxe4m3_mxe2m1_block_scale_moe_runner.py similarity index 99% rename from tests/unittest/_torch/staircase/moe/test_staircase_mxe4m3_mxe2m1_block_scale_moe_runner.py rename to tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_mxe4m3_mxe2m1_block_scale_moe_runner.py index 6cba8a7058f1..f3c98329c999 100644 --- a/tests/unittest/_torch/staircase/moe/test_staircase_mxe4m3_mxe2m1_block_scale_moe_runner.py +++ b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_mxe4m3_mxe2m1_block_scale_moe_runner.py @@ -6,7 +6,7 @@ import torch -from tensorrt_llm._torch.staircase.catalog.moe.mxe4m3_mxe2m1_block_scale_moe_runner import ( +from tensorrt_llm._torch.modeling_v2.catalog.moe.mxe4m3_mxe2m1_block_scale_moe_runner import ( mxe4m3_mxe2m1_block_scale_moe_runner as moe, ) diff --git a/tests/unittest/_torch/staircase/moe/test_staircase_noaux_tc_op.py b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_noaux_tc_op.py similarity index 99% rename from tests/unittest/_torch/staircase/moe/test_staircase_noaux_tc_op.py rename to tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_noaux_tc_op.py index e9e22f159df9..e4d2e14e0694 100644 --- a/tests/unittest/_torch/staircase/moe/test_staircase_noaux_tc_op.py +++ b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_noaux_tc_op.py @@ -4,7 +4,7 @@ import torch -from tensorrt_llm._torch.staircase.catalog.moe.noaux_tc_op import noaux_tc_op +from tensorrt_llm._torch.modeling_v2.catalog.moe.noaux_tc_op import noaux_tc_op assert torch.cuda.is_available(), "noaux_tc_op requires a CUDA device" diff --git a/tests/unittest/_torch/staircase/norm/test_staircase_flashinfer_fused_add_rmsnorm.py b/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_flashinfer_fused_add_rmsnorm.py similarity index 97% rename from tests/unittest/_torch/staircase/norm/test_staircase_flashinfer_fused_add_rmsnorm.py rename to tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_flashinfer_fused_add_rmsnorm.py index 139b05951c68..55f873117e37 100644 --- a/tests/unittest/_torch/staircase/norm/test_staircase_flashinfer_fused_add_rmsnorm.py +++ b/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_flashinfer_fused_add_rmsnorm.py @@ -4,7 +4,7 @@ import torch -from tensorrt_llm._torch.staircase.catalog.norm.flashinfer_fused_add_rmsnorm import ( +from tensorrt_llm._torch.modeling_v2.catalog.norm.flashinfer_fused_add_rmsnorm import ( flashinfer_fused_add_rmsnorm, ) diff --git a/tests/unittest/_torch/staircase/norm/test_staircase_flashinfer_rmsnorm.py b/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_flashinfer_rmsnorm.py similarity index 96% rename from tests/unittest/_torch/staircase/norm/test_staircase_flashinfer_rmsnorm.py rename to tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_flashinfer_rmsnorm.py index c9bb7cad58c8..2c2c61d31dff 100644 --- a/tests/unittest/_torch/staircase/norm/test_staircase_flashinfer_rmsnorm.py +++ b/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_flashinfer_rmsnorm.py @@ -4,7 +4,7 @@ import torch -from tensorrt_llm._torch.staircase.catalog.norm.flashinfer_rmsnorm import flashinfer_rmsnorm +from tensorrt_llm._torch.modeling_v2.catalog.norm.flashinfer_rmsnorm import flashinfer_rmsnorm assert torch.cuda.is_available(), "flashinfer_rmsnorm requires a CUDA device" diff --git a/tests/unittest/_torch/staircase/quantization/test_staircase_fp4_quantize.py b/tests/unittest/_torch/modeling_v2/quantization/test_modeling_v2_fp4_quantize.py similarity index 99% rename from tests/unittest/_torch/staircase/quantization/test_staircase_fp4_quantize.py rename to tests/unittest/_torch/modeling_v2/quantization/test_modeling_v2_fp4_quantize.py index 3ee5d5fc884b..02687d657f22 100644 --- a/tests/unittest/_torch/staircase/quantization/test_staircase_fp4_quantize.py +++ b/tests/unittest/_torch/modeling_v2/quantization/test_modeling_v2_fp4_quantize.py @@ -6,7 +6,7 @@ from torch.profiler import ProfilerActivity, profile import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* -from tensorrt_llm._torch.staircase.catalog.quantization.fp4_quantize import fp4_quantize +from tensorrt_llm._torch.modeling_v2.catalog.quantization.fp4_quantize import fp4_quantize assert torch.cuda.is_available(), "fp4_quantize requires a CUDA device" diff --git a/tests/unittest/_torch/staircase/quantization/test_staircase_mxfp8_quantize.py b/tests/unittest/_torch/modeling_v2/quantization/test_modeling_v2_mxfp8_quantize.py similarity index 99% rename from tests/unittest/_torch/staircase/quantization/test_staircase_mxfp8_quantize.py rename to tests/unittest/_torch/modeling_v2/quantization/test_modeling_v2_mxfp8_quantize.py index 5187ae28c55c..a7d638ee37e4 100644 --- a/tests/unittest/_torch/staircase/quantization/test_staircase_mxfp8_quantize.py +++ b/tests/unittest/_torch/modeling_v2/quantization/test_modeling_v2_mxfp8_quantize.py @@ -4,7 +4,7 @@ import torch -from tensorrt_llm._torch.staircase.catalog.quantization.mxfp8_quantize import mxfp8_quantize +from tensorrt_llm._torch.modeling_v2.catalog.quantization.mxfp8_quantize import mxfp8_quantize assert torch.cuda.is_available(), "mxfp8_quantize requires a CUDA device" diff --git a/tests/unittest/_torch/staircase/test_staircase_claims.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py similarity index 94% rename from tests/unittest/_torch/staircase/test_staircase_claims.py rename to tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py index 78d2812f14df..84a41c344a4c 100644 --- a/tests/unittest/_torch/staircase/test_staircase_claims.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py @@ -22,13 +22,13 @@ import pytest -import tensorrt_llm._torch.staircase as _staircase -from tensorrt_llm._torch.staircase._router_index import STAIRCASE_ROUTERS, routing_module +import tensorrt_llm._torch.modeling_v2 as _modeling_v2 +from tensorrt_llm._torch.modeling_v2._router_index import MODELING_V2_ROUTERS, routing_module # The package, not this file: these paths address the tree under test, and this # test lives in tests/ while that tree lives in tensorrt_llm/. -_ROOT = Path(_staircase.__file__).resolve().parent -_ARCHS = sorted(STAIRCASE_ROUTERS) +_ROOT = Path(_modeling_v2.__file__).resolve().parent +_ARCHS = sorted(MODELING_V2_ROUTERS) def _module_path(dotted: str) -> Path: @@ -43,10 +43,10 @@ def test_every_routed_architecture_has_an_importable_routing_module(): def test_no_routing_module_is_orphaned(): """A routing.py the index does not name would never be consulted.""" - indexed = {_module_path(m) for m in STAIRCASE_ROUTERS.values()} + indexed = {_module_path(m) for m in MODELING_V2_ROUTERS.values()} on_disk = set((_ROOT / "models").glob("*/routing.py")) assert on_disk == indexed, ( - f"routing modules on disk but not in STAIRCASE_ROUTERS: " + f"routing modules on disk but not in MODELING_V2_ROUTERS: " f"{sorted(p.relative_to(_ROOT) for p in on_disk - indexed)}" ) @@ -68,7 +68,7 @@ def test_targets_table_and_target_modules_agree(arch): def test_each_target_module_registers_its_own_name(arch): """The synthetic name is only a registry key -- the module must fill it. - Nothing else can: the built-in static index does not carry staircase + Nothing else can: the built-in static index does not carry modeling_v2 names, so a target whose decorator says something different resolves to None and the engine reports an unknown architecture. """ @@ -158,7 +158,7 @@ def test_no_routing_module_reads_an_unplumbed_dimension(arch, field): readable fields is bounded by the weaker of the two callers, not the engine alone. """ - source = _module_path(STAIRCASE_ROUTERS[arch]).read_text() + source = _module_path(MODELING_V2_ROUTERS[arch]).read_text() assert field not in source, ( f"{arch}: routing reads ctx.{field}, but {_UNREADABLE_CONTEXT_FIELDS[field]}" ) diff --git a/tests/unittest/_torch/staircase/test_staircase_no_stale_claims.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_no_stale_claims.py similarity index 88% rename from tests/unittest/_torch/staircase/test_staircase_no_stale_claims.py rename to tests/unittest/_torch/modeling_v2/test_modeling_v2_no_stale_claims.py index 6a70dcbee41a..98286ab3f99a 100644 --- a/tests/unittest/_torch/staircase/test_staircase_no_stale_claims.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_no_stale_claims.py @@ -1,8 +1,8 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""The audit that keeps pinned versions and dates out of the staircase tree. +"""The audit that keeps pinned versions and dates out of the modeling_v2 tree. -Out of tree, staircase was a separate repo against a pinned `tensorrt_llm`, so +Out of tree, modeling_v2 was a separate repo against a pinned `tensorrt_llm`, so writing that version into a contract was meaningful: the pin was the thing a receipt was valid against, and a bump really did void what had been measured. @@ -23,11 +23,11 @@ import re from pathlib import Path -import tensorrt_llm._torch.staircase as _staircase +import tensorrt_llm._torch.modeling_v2 as _modeling_v2 # The package, not this file: the tree under audit lives under tensorrt_llm/ # while this test lives under tests/. -_ROOT = Path(_staircase.__file__).resolve().parent +_ROOT = Path(_modeling_v2.__file__).resolve().parent # Each pattern, and what to say instead of it. _STALE_CLAIM_PATTERNS = ( @@ -75,4 +75,6 @@ def test_no_file_pins_a_version_or_a_date(): if (hit := pattern.search(line)) is not None: rel = path.relative_to(_ROOT) offences.append(f" {rel}:{lineno}: {hit.group(0)!r} is {why}") - assert not offences, "staircase files must not pin a version or a date:\n" + "\n".join(offences) + assert not offences, "modeling_v2 files must not pin a version or a date:\n" + "\n".join( + offences + ) diff --git a/tests/unittest/_torch/staircase/test_staircase_routing.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py similarity index 82% rename from tests/unittest/_torch/staircase/test_staircase_routing.py rename to tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py index 70c07d7b92df..85d85438ffc1 100644 --- a/tests/unittest/_torch/staircase/test_staircase_routing.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py @@ -1,6 +1,6 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""What ``staircase_resolve`` actually does, driven by synthetic configs. +"""What ``modeling_v2_resolve`` actually does, driven by synthetic configs. No checkpoint and no weights: routing reads config *shape*, the mapping and the SM version, all of which can be stated directly. The SM version is @@ -24,16 +24,16 @@ from transformers import PretrainedConfig from tensorrt_llm._torch.model_config import ModelConfig +from tensorrt_llm._torch.modeling_v2._router_index import ( + MODELING_V2_ENV, + ModelingV2Mode, + modeling_v2_resolve, +) from tensorrt_llm._torch.models.modeling_auto import AutoModelForCausalLM from tensorrt_llm._torch.models.modeling_utils import ( _is_builtin_model_class, get_registered_model_class, ) -from tensorrt_llm._torch.staircase._router_index import ( - STAIRCASE_ENV, - StaircaseMode, - staircase_resolve, -) from tensorrt_llm.mapping import Mapping _SM103 = (10, 3) @@ -80,14 +80,14 @@ def _mode(monkeypatch, request): The switch is an environment variable, so the tests set one too -- that is the surface under test. """ - monkeypatch.setenv(STAIRCASE_ENV, "auto") + monkeypatch.setenv(MODELING_V2_ENV, "auto") def _set_mode(monkeypatch, mode): if mode is None: - monkeypatch.delenv(STAIRCASE_ENV, raising=False) + monkeypatch.delenv(MODELING_V2_ENV, raising=False) else: - monkeypatch.setenv(STAIRCASE_ENV, mode) + monkeypatch.setenv(MODELING_V2_ENV, mode) def _model_config(pretrained_config, **mapping_kwargs): @@ -100,14 +100,14 @@ def _model_config(pretrained_config, **mapping_kwargs): @pytest.mark.parametrize("unset", [True, False], ids=["env-unset", "env-off"]) def test_off_resolves_nothing(monkeypatch, unset): - """The default must be indistinguishable from staircase not existing. + """The default must be indistinguishable from modeling_v2 not existing. Unset and an explicit "off" have to behave identically: the common case is that nobody has heard of this package. """ _set_mode(monkeypatch, None if unset else "off") config = _model_config(_gpt_oss_config()) - assert staircase_resolve(config) is None + assert modeling_v2_resolve(config) is None def test_off_still_reaches_the_builtin_implementation(monkeypatch): @@ -122,25 +122,25 @@ def test_off_still_reaches_the_builtin_implementation(monkeypatch): def test_gpt_oss_tp1_matches(monkeypatch, mode): _set_mode(monkeypatch, mode) config = _model_config(_gpt_oss_config()) - assert staircase_resolve(config) == "StaircaseGptOss120bSm103Tp1" + assert modeling_v2_resolve(config) == "ModelingV2GptOss120bSm103Tp1" @pytest.mark.parametrize("mode", ["auto", "require"]) def test_r1_dep4_matches(monkeypatch, mode): _set_mode(monkeypatch, mode) config = _model_config(_r1_config(), **_DEP4) - assert staircase_resolve(config) == "StaircaseDeepseekR10528Nvfp4Sm103Dep4" + assert modeling_v2_resolve(config) == "ModelingV2DeepseekR10528Nvfp4Sm103Dep4" def test_resolving_registers_the_target_class(): """The synthetic name is a key; the import behind it is what fills it.""" config = _model_config(_gpt_oss_config()) - name = staircase_resolve(config) + name = modeling_v2_resolve(config) cls = get_registered_model_class(name) assert cls is not None, f"{name} resolved to no class" assert cls.__name__ == name assert cls.__module__.endswith( - "staircase.models.gpt_oss.targets.gpt_oss_120b.sm_103.tp1.modeling" + "modeling_v2.models.gpt_oss.targets.gpt_oss_120b.sm_103.tp1.modeling" ) @@ -149,14 +149,14 @@ def test_the_target_registration_counts_as_external(): empty ones. Living beside the zoo rather than inside it is what buys this, and a move into _torch/models/ would silently reverse it.""" config = _model_config(_gpt_oss_config()) - cls = get_registered_model_class(staircase_resolve(config)) + cls = get_registered_model_class(modeling_v2_resolve(config)) assert not _is_builtin_model_class(cls) def test_resolve_class_rewrites_the_architecture_end_to_end(): config = _model_config(_gpt_oss_config()) resolved = AutoModelForCausalLM._resolve_class(config) - assert resolved.__name__ == "StaircaseGptOss120bSm103Tp1" + assert resolved.__name__ == "ModelingV2GptOss120bSm103Tp1" @pytest.mark.parametrize( @@ -170,11 +170,11 @@ def test_resolve_class_rewrites_the_architecture_end_to_end(): ) def test_gpt_oss_near_misses_do_not_match(monkeypatch, config_kwargs, mapping_kwargs, missed): config = _model_config(_gpt_oss_config(**config_kwargs), **mapping_kwargs) - assert staircase_resolve(config) is None + assert modeling_v2_resolve(config) is None _set_mode(monkeypatch, "require") with pytest.raises(ValueError, match=missed): - staircase_resolve(config) + modeling_v2_resolve(config) def test_r1_without_attention_dp_does_not_match(): @@ -182,7 +182,7 @@ def test_r1_without_attention_dp_does_not_match(): attention DP is an identity criterion rather than a knob.""" mapping_kwargs = dict(_DEP4, enable_attention_dp=False) config = _model_config(_r1_config(), **mapping_kwargs) - assert staircase_resolve(config) is None + assert modeling_v2_resolve(config) is None def test_require_names_the_criterion_that_missed(monkeypatch): @@ -190,7 +190,7 @@ def test_require_names_the_criterion_that_missed(monkeypatch): _set_mode(monkeypatch, "require") config = _model_config(_gpt_oss_config()) with pytest.raises(ValueError) as excinfo: - staircase_resolve(config) + modeling_v2_resolve(config) message = str(excinfo.value) assert "sm" in message and "(10, 0)" in message assert "no match" in message @@ -198,14 +198,14 @@ def test_require_names_the_criterion_that_missed(monkeypatch): def test_an_unrouted_architecture_is_not_an_error_under_auto(): config = _model_config(PretrainedConfig(architectures=["LlamaForCausalLM"])) - assert staircase_resolve(config) is None + assert modeling_v2_resolve(config) is None def test_an_unrouted_architecture_raises_under_require(monkeypatch): _set_mode(monkeypatch, "require") config = _model_config(PretrainedConfig(architectures=["LlamaForCausalLM"])) with pytest.raises(ValueError, match="LlamaForCausalLM"): - staircase_resolve(config) + modeling_v2_resolve(config) @pytest.mark.parametrize( @@ -215,14 +215,14 @@ def test_an_unrouted_architecture_raises_under_require(monkeypatch): def test_the_env_var_is_read_leniently(monkeypatch, raw, expected): """``None`` is the unset case, and it is the one that must never drift. - Everything in the accuracy suite rests on staircase being opt-in: unset has + Everything in the accuracy suite rests on modeling_v2 being opt-in: unset has to read as off on the code path the engine actually takes. """ if raw is None: - monkeypatch.delenv(STAIRCASE_ENV, raising=False) + monkeypatch.delenv(MODELING_V2_ENV, raising=False) else: - monkeypatch.setenv(STAIRCASE_ENV, raw) - assert StaircaseMode.from_env().value == expected + monkeypatch.setenv(MODELING_V2_ENV, raw) + assert ModelingV2Mode.from_env().value == expected def test_an_unknown_mode_raises_rather_than_falling_back(monkeypatch): @@ -234,6 +234,6 @@ def test_an_unknown_mode_raises_rather_than_falling_back(monkeypatch): """ # "yes" rather than a misspelling: it is what someone reaching for a # boolean would write, and it is the reading that must not be invented. - monkeypatch.setenv(STAIRCASE_ENV, "yes") - with pytest.raises(ValueError, match="not a staircase mode"): - StaircaseMode.from_env() + monkeypatch.setenv(MODELING_V2_ENV, "yes") + with pytest.raises(ValueError, match="not a modeling_v2 mode"): + ModelingV2Mode.from_env() diff --git a/tests/unittest/_torch/staircase/test_staircase_target_contract.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_target_contract.py similarity index 91% rename from tests/unittest/_torch/staircase/test_staircase_target_contract.py rename to tests/unittest/_torch/modeling_v2/test_modeling_v2_target_contract.py index 933a84e704fa..50e40c22e9fc 100644 --- a/tests/unittest/_torch/staircase/test_staircase_target_contract.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_target_contract.py @@ -13,7 +13,7 @@ it out is a test on a machine that has the extension built. Each target now declares ``REQUIRED_TRTLLM_OPS`` and this asserts it. -Distinct from ``test_staircase_claims.py``, which is deliberately import-free +Distinct from ``test_modeling_v2_claims.py``, which is deliberately import-free and runs anywhere: this one imports the targets, so it needs a real build. """ @@ -25,10 +25,10 @@ import torch import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* -from tensorrt_llm._torch.staircase._router_index import STAIRCASE_ROUTERS, routing_module +from tensorrt_llm._torch.modeling_v2._router_index import MODELING_V2_ROUTERS, routing_module -_PACKAGE = "tensorrt_llm._torch.staircase" -_ARCHS = sorted(STAIRCASE_ROUTERS) +_PACKAGE = "tensorrt_llm._torch.modeling_v2" +_ARCHS = sorted(MODELING_V2_ROUTERS) def _target_modules(): From 62c771c94a499261187ec1aadbafe5774cb7e6a4 Mon Sep 17 00:00:00 2001 From: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> Date: Tue, 15 Sep 2026 21:01:04 -0700 Subject: [PATCH 08/19] [TRTLLM-16304][chore] Follow upstream's removal of MTPDraftModelForCausalLM main removed the two-model speculative decoding path (#18721), taking MTPDraftModelForCausalLM with it -- the class, its registration and its _resolve_class rewrite. Three places here cited that rewrite as the upstream precedent for a synthetic architecture name, which is now a claim about code that no longer exists. The Eagle3 rewrite is the surviving instance of the same pattern and builds EAGLE3 the same way, so it takes over as the example. The precedence note in _resolve_class drops to one rewrite for the same reason, along with the clause about MTPDecodingConfig's validator, which was explaining why a branch that is now gone was unreachable. Nothing in the targets changes: they were always on the one-model MTP path, which is what survived. Signed-off-by: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> --- tensorrt_llm/_torch/modeling_v2/README.md | 5 ++--- tensorrt_llm/_torch/modeling_v2/_router_index.py | 5 +++-- 2 files changed, 5 insertions(+), 5 deletions(-) diff --git a/tensorrt_llm/_torch/modeling_v2/README.md b/tensorrt_llm/_torch/modeling_v2/README.md index 98b9000b145f..9b20a9a6d701 100644 --- a/tensorrt_llm/_torch/modeling_v2/README.md +++ b/tensorrt_llm/_torch/modeling_v2/README.md @@ -62,9 +62,8 @@ LLM(model=...) -> ModelLoader -> AutoModelForCausalLM._resolve_class ``` The synthetic name (`ModelingV2GptOss120bSm103Tp1`) is a registry key that no -checkpoint declares. Upstream already does exactly this for -`MTPDraftModelForCausalLM`, which also exists only as a `_resolve_class` -rewrite. +checkpoint declares. Upstream already does exactly this for `EAGLE3`, +which also exists only as a `_resolve_class` rewrite. To ask why a configuration landed where it did: diff --git a/tensorrt_llm/_torch/modeling_v2/_router_index.py b/tensorrt_llm/_torch/modeling_v2/_router_index.py index 919e33132d67..e9437214a31a 100644 --- a/tensorrt_llm/_torch/modeling_v2/_router_index.py +++ b/tensorrt_llm/_torch/modeling_v2/_router_index.py @@ -9,8 +9,9 @@ synthetic name, and ``AutoModelForCausalLM._resolve_class`` then looks that up in the ordinary registry. -The pattern is upstream's own: ``MTPDraftModelForCausalLM`` is likewise a name -no config.json declares, reached only through a rewrite in ``_resolve_class``. +The pattern is upstream's own: the Eagle3 rewrite in ``_resolve_class`` builds +``EAGLE3`` the same way -- a name no config.json declares, reached only +through that rewrite. Two levels, on purpose: From 48dec0d2c678c4c26277afd475e177ca85474762 Mon Sep 17 00:00:00 2001 From: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> Date: Wed, 16 Sep 2026 01:03:52 -0700 Subject: [PATCH 09/19] [TRTLLM-16304][fix] Certify thop_attention over two KV pools gpt-oss alternates sliding-window and full attention, and KVCacheManagerV2 gives each attention-window class its own layer group -- so this checkpoint now arrives with a pool mapping carrying ids {0, 1}. The target refused it: its first-forward contract check required a single pool, which is what the catalog certified. The refusal was right. Nothing in the forward needed changing -- it passes the pool tables through and the op reads the pool column itself -- but "the op accepts it" is not "we measured it", and the entry had no cell where the pool column was ever non-zero. So the cell comes first. Two managers stand in for two pools, the tables are composed across them, and every layer is driven through prefill and decode. The assertion that carries it is the negative one: a call for a layer in one pool must leave the other bitwise untouched. Collapsing pool selection to 0 -- exactly the failure worth catching -- still produces plausible outputs, since both pools hold validly shaped pages; only the sibling-pool comparison sees it. call_op grows three overrides for this, since one manager cannot produce a two-pool layout on its own. With that measured, the gpt-oss bound widens to match the contract, and the contract says what is now certified and what is not (three pools, and pools differing in geometry, are not). The deepseek target keeps the tighter bound. It reaches its KV cache through the MLA entries instead, none of which has a multi-pool cell, and its layers are all one attention-window class so a manager gives them one group. Widening it would be borrowing this evidence for ops it does not cover. Verified on GB300: the new cell passes, the whole entry passes at 97 cases, and the gpt-oss gate scores 89.917 against a threshold of 87.097 -- with layer_group_id 0 and 1 both present in the run, so the relaxed bound is what let it execute rather than a quiet return to one pool. Signed-off-by: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> --- .../catalog/attention/thop_attention.md | 22 ++- .../r1_0528_nvfp4/sm_103/dep4/modeling.py | 11 +- .../gpt_oss_120b/sm_103/tp1/modeling.py | 20 ++- .../test_modeling_v2_thop_attention.py | 134 +++++++++++++++++- 4 files changed, 170 insertions(+), 17 deletions(-) diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/attention/thop_attention.md b/tensorrt_llm/_torch/modeling_v2/catalog/attention/thop_attention.md index b5852e6791fb..ebee2724729a 100644 --- a/tensorrt_llm/_torch/modeling_v2/catalog/attention/thop_attention.md +++ b/tensorrt_llm/_torch/modeling_v2/catalog/attention/thop_attention.md @@ -1048,18 +1048,26 @@ exactly the geometry above. | Argument | Shape / value | Dtype | Device | |---|---|---|---| | `kv_cache_block_offsets` | `[1, max_num_requests, 2, max_blocks_per_seq]`; row `[0, s, 0, j]` = K-slab offset of sequence `s`'s `j`-th page, `[0, s, 1, j]` = V-slab offset. Offsets count slabs and are **layer-agnostic** — every layer of the pool shares one offsets row per sequence; layer selection comes from the pool mapping, not the offsets. Standard `L`-layer pool → page `p` has K `p * 2L`, V `p * 2L + 1` (single layer: K `2p`, V `2p + 1`); MLA pool → raw block id `p` in **both** rows. Only the first `ceil(kv_s / tokens_per_block)` entries per row are read | int32 | CUDA | -| `host_kv_cache_pool_pointers` | `[1, 2]`: `[0, 0]` = pool base address (`pool.data_ptr()`), `[0, 1]` = secondary-pool address (0 = none) | int64 | CPU | -| `host_kv_cache_pool_mapping` | `[num_layers, 2]`, one row per layer: row `local_layer_idx` = (pool index, layer index within pool). The row's **layer column drives the pool-base shift** — the call addresses slabs starting `layer_in_pool * kv_factor` slabs from the pool base (see *Notes*). Certified: the identity rows a real manager produces — `[[0, 0] .. [0, 3]]` for the 4-layer standard pool; `[1, 2]` zeros for single-layer pools | int32 | CPU | -| `local_layer_idx` (int) | row of `host_kv_cache_pool_mapping` for this layer (production: the layer's index among the rank's local layers); 0-3 certified (standard), 0 (MLA) | — | — | +| `host_kv_cache_pool_pointers` | `[num_pools, 2]`: row `p` is pool `p`'s (base address = `pool.data_ptr()`, secondary-pool address; 0 = none). Certified at `num_pools` 1 and 2 | int64 | CPU | +| `host_kv_cache_pool_mapping` | `[num_layers, 2]`, one row per layer: row `local_layer_idx` = (pool index, layer index within pool). The row's **layer column drives the pool-base shift** — the call addresses slabs starting `layer_in_pool * kv_factor` slabs from the pool base (see *Notes*). The **pool column selects which base** to shift from. Certified: the identity rows a real manager produces — `[[0, 0] .. [0, 3]]` for the 4-layer standard pool, `[1, 2]` zeros for single-layer pools, and `[[0, 0], [0, 1], [1, 0], [1, 1]]` for two pools of two layers each, which is the shape KVCacheManagerV2 produces when a checkpoint's layers fall into more than one attention-window class | int32 | CPU | +| `local_layer_idx` (int) | row of `host_kv_cache_pool_mapping` for this layer (production: the layer's index among the rank's local layers); 0-3 certified (standard, one pool and two), 0 (MLA) | — | — | | `tokens_per_block` (int) | pool page size; must be a power of two; 32 certified (standard), 32 and 64 certified (MLA — see *MLA page size*) | — | — | | `update_kv_cache` (bool) | `True` (also for the MLA generation and no-append context calls, which nevertheless write nothing) | — | — | | `cache_indirection` | `None` (beam search only) | — | — | | `block_ids_per_seq` | `None` (FlashMLA path only) | — | — | -A sequence's page set is shared by every layer of the pool: the per-layer -calls of one batch pass identical offsets and pool pointers and differ -only in `local_layer_idx`. Multiple pools (`num_pools > 1`: extra -offsets/pointer rows, mapping rows with pool index > 0) are not certified. +A sequence's page set is shared by every layer of *its own* pool: the +per-layer calls of one batch pass identical offsets and pool pointers and +differ only in `local_layer_idx`. + +Two pools are certified. A layer's pool comes from the mapping row's pool +column, which selects both the pointer row and the offsets slot; layers in +different pools share nothing. The certified case drives every layer of a +two-pool, two-layers-each layout through prefill and decode and asserts that +each call leaves the *other* pool bitwise untouched -- collapsing the pool +selection to 0 still produces plausible outputs, since both pools hold validly +shaped pages, so only that comparison detects it. `num_pools > 2` is +uncertified, as are pools that differ from each other in geometry. ### Paged KV cache addressing under a sliding window diff --git a/tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py b/tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py index b9db02c748f6..da0aab9e4e21 100644 --- a/tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py +++ b/tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py @@ -1007,9 +1007,18 @@ def _check_step_contract(self, md, position_ids) -> None: # generation token's from sequence_length - 1. Checked anyway so a # layout change upstream is loud rather than silent. assert position_ids.dtype == torch.int32, position_ids.dtype + # Narrower than the gpt-oss target's bound on purpose. That one reaches + # its KV cache through thop_attention, which is certified at one pool + # and two; this one goes through the MLA family -- + # load_paged_kv_cache_for_mla, mla_rope_generation, + # mla_rope_append_paged_kv_assign_q -- and none of those entries has a + # multi-pool cell. Every layer here is the same attention-window class, + # so a manager gives them one group and this holds; if that ever + # changes, certify the MLA entries before widening it. pools = {row[0] for row in md.host_kv_cache_pool_mapping.tolist()} assert pools == {0}, ( - f"multi-pool KV addressing is not certified; layer->pool ids {sorted(pools)}" + f"multi-pool KV addressing is not certified for the MLA entries; " + f"layer->pool ids {sorted(pools)}" ) # Every MLA entry's fp8-e4m3 column is page 32 only (the bf16 columns # also carry 64). 32 is what a default KvCacheConfig produces. diff --git a/tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py b/tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py index 44cc94500137..db6c43bc476e 100644 --- a/tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py +++ b/tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py @@ -496,16 +496,24 @@ def build_layer_views(self) -> None: def _check_step_contract(self, md, position_ids) -> None: """First-forward fail-fast: the metadata fields this target consumes - must exist (they are private trtllm surface, pinned by version), and - the KV pool must be the single shared pool the sliding-window - surface is certified over. Everything checked is fixed at engine - construction — once per model instance is sound.""" + must exist (they are private trtllm surface), and the KV pool layout + must be one `thop_attention` is certified over. Everything checked is + fixed at engine construction — once per model instance is sound. + + This checkpoint alternates sliding-window and full attention, and + KVCacheManagerV2 gives each attention-window class its own layer + group, so its mapping carries two pool ids rather than one. The bound + is the catalog's: `thop_attention.md` certifies one pool and two, and + this forward passes the mapping through untouched -- the op reads the + pool column itself. A third pool would be outside what was measured. + """ missing = [name for name in _STEP_FIELDS if not hasattr(md, name)] assert not missing, f"metadata fields missing: {missing}" assert position_ids.dtype == torch.int32 pools = {row[0] for row in md.host_kv_cache_pool_mapping.tolist()} - assert pools == {0}, ( - f"multi-pool KV addressing is not certified; layer->pool ids {sorted(pools)}" + assert pools <= {0, 1}, ( + f"KV addressing over {len(pools)} pools is not certified " + f"(thop_attention.md certifies one and two); layer->pool ids {sorted(pools)}" ) self._step_contract_checked = True diff --git a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_thop_attention.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_thop_attention.py index 19f04f8e931c..1eb1740f37a9 100644 --- a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_thop_attention.py +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_thop_attention.py @@ -893,6 +893,9 @@ def call_op( attention_window_size: Optional[int] = None, use_paged_context_fmha: bool = False, ctx_lens: Optional[List[int]] = None, + pool_pointers: Optional[torch.Tensor] = None, + block_offsets: Optional[torch.Tensor] = None, + local_layer_idx: Optional[int] = None, ) -> torch.Tensor: """One thop_attention call for one layer of the shared pool. @@ -904,6 +907,18 @@ def call_op( use_paged_context_fmha=True selects the paged-context execution path; ctx_lens then carries this call's per-context-row new-token count, which a cached prefix makes differ from the registered prompt length. + + pool_pointers and block_offsets override the manager's, which is what + lets one env stand in for one *pool* of a several-pool layout: the + caller composes the tables across envs and this call addresses its own + rows through them. A single manager cannot produce that -- it owns one + pool -- and the pool-index column of the mapping is unexercised + without it. + + local_layer_idx defaults to layer_idx and separates the two meanings + that coincide in a single-pool layout: layer_idx is this env's own + layer (its K/V history and buffer view), while local_layer_idx is the + row of the composed mapping. Across pools they differ. """ ns = len(request_ids) kv_lens = [self.cached_len(layer_idx, rid) + new for rid, new in zip(request_ids, seq_lens)] @@ -913,6 +928,12 @@ def call_op( num_ctx_tokens = sum(seq_lens[:num_contexts]) if pool_mapping is None: pool_mapping = self.pool_mapping + if pool_pointers is None: + pool_pointers = self.mgr.kv_cache_pool_pointers + if block_offsets is None: + block_offsets = self.block_offsets + if local_layer_idx is None: + local_layer_idx = layer_idx output = torch.empty( sum(seq_lens), @@ -937,8 +958,8 @@ def call_op( host_context_lengths=torch.tensor(ctx_lens, dtype=torch.int32), host_request_types=torch.tensor(req_types, dtype=torch.int32), max_context_q_len_override=None, - kv_cache_block_offsets=self.block_offsets, - host_kv_cache_pool_pointers=self.mgr.kv_cache_pool_pointers, + kv_cache_block_offsets=block_offsets, + host_kv_cache_pool_pointers=pool_pointers, host_kv_cache_pool_mapping=pool_mapping, cache_indirection=None, kv_scale_orig_quant=self.kv_scale_orig_quant, @@ -953,7 +974,7 @@ def call_op( is_fused_qkv=True, update_kv_cache=True, predicted_tokens_per_seq=1, - local_layer_idx=layer_idx, + local_layer_idx=local_layer_idx, num_heads=self.num_heads, num_kv_heads=self.num_kv_heads, head_size=self.head_dim, @@ -1171,6 +1192,113 @@ def test_bf16_multilayer_shared_pool_gqa_d128() -> None: assert torch.equal(env.layer_views[2], views_before[2]) +def _compose_two_pools(env_a, env_b): + """Pool pointers, layer->pool mapping and block offsets for two pools. + + Row i of the mapping is what `local_layer_idx=i` selects. Rows 0..La-1 + address pool 0 (env_a), rows La.. address pool 1 (env_b), each with its + own layer-in-pool column. This is the shape a KVCacheManagerV2 produces + when a checkpoint's layers fall into more than one attention-window class + -- gpt-oss alternates sliding-window and full attention, so its layers + land in two groups and its mapping carries pool ids {0, 1}. + """ + ptr_a = env_a.mgr.kv_cache_pool_pointers + ptr_b = env_b.mgr.kv_cache_pool_pointers + pool_pointers = torch.zeros(2, 2, dtype=torch.int64, device="cpu") + pool_pointers[0] = ptr_a[0] + pool_pointers[1] = ptr_b[0] + + rows = [[0, i] for i in range(env_a.num_layers)] + rows += [[1, i] for i in range(env_b.num_layers)] + pool_mapping = torch.tensor(rows, dtype=torch.int32, device="cpu") + + assert env_a.block_offsets.shape == env_b.block_offsets.shape + block_offsets = torch.cat([env_a.block_offsets, env_b.block_offsets], dim=0) + return pool_pointers, pool_mapping, block_offsets + + +def test_bf16_two_pool_layer_routing_gqa_d128() -> None: + """Two paged pools, layers routed between them by the mapping's pool column. + + Everything above this point certifies pool id 0 only: one manager owns one + pool, so the first column of every mapping row is 0 and the op's pool + selection is never exercised. A checkpoint whose layers differ in + attention-window class does not get that layout -- KVCacheManagerV2 puts + each class in its own group, and the mapping then carries ids {0, 1}. + + Two managers stand in for the two pools. Each layer is driven through the + composed tables, and the check that matters is the negative one: a call + for a layer in pool 1 must leave pool 0 bitwise untouched and vice versa. + Collapsing pool selection to 0 -- the failure this cell exists to catch -- + would still produce plausible outputs, because both pools hold validly + shaped pages; only the sibling-pool comparison sees it. + """ + torch.manual_seed(11) + kw = dict(num_heads=8, num_kv_heads=2, head_dim=128) + env_a = _MultiLayerPagedAttnEnv(num_layers=2, **kw) + env_b = _MultiLayerPagedAttnEnv(num_layers=2, **kw) + envs = [(env_a, 0), (env_a, 1), (env_b, 0), (env_b, 1)] # (env, layer in that env) + + for env in (env_a, env_b): + env.add_request(0, 64) # two whole pages; decode crosses into a third + env.add_request(1, 17) + env.refresh_offsets([0, 1], num_contexts=2) + ptrs, mapping, offsets = _compose_two_pools(env_a, env_b) + assert mapping.tolist() == [[0, 0], [0, 1], [1, 0], [1, 1]] + assert len({int(ptrs[0, 0]), int(ptrs[1, 0])}) == 2, "the two pools must have distinct bases" + + def other(env): + return env_b if env is env_a else env_a + + for row, (env, layer) in enumerate(envs): + qkv = env.random_qkv(81) + foreign_before = [v.clone() for v in other(env).layer_views] + out = env.call_op( + layer, + qkv, + [64, 17], + 2, + [0, 1], + pool_mapping=mapping, + pool_pointers=ptrs, + block_offsets=offsets, + local_layer_idx=row, + ) + ref = env.reference(layer, qkv, [64, 17], [0, 1], [0, 0]) + torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) + for before, after in zip(foreign_before, other(env).layer_views): + assert torch.equal(before, after), ( + f"row {row} addressed pool {mapping[row, 0].item()} but wrote the other pool" + ) + for env in (env_a, env_b): + env.check_caches([0, 1]) + + for _ in range(2): # decode, so the routing holds past the prefill pages + for env in (env_a, env_b): + env.add_decode_token(0) + env.add_decode_token(1) + env.refresh_offsets([0, 1], num_contexts=0) + _, _, offsets = _compose_two_pools(env_a, env_b) + for row, (env, layer) in enumerate(envs): + cached = [env.cached_len(layer, 0), env.cached_len(layer, 1)] + qkv = env.random_qkv(2) + out = env.call_op( + layer, + qkv, + [1, 1], + 0, + [0, 1], + pool_mapping=mapping, + pool_pointers=ptrs, + block_offsets=offsets, + local_layer_idx=row, + ) + ref = env.reference(layer, qkv, [1, 1], [0, 1], cached) + torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) + for env in (env_a, env_b): + env.check_caches([0, 1]) + + # ─── Standard configuration, fp8-e4m3 paged KV pool (quant_mode 128) ─── From 2d82d7a11891765ef828e1158ff90ae5e8fce0ab Mon Sep 17 00:00:00 2001 From: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> Date: Wed, 16 Sep 2026 02:02:09 -0700 Subject: [PATCH 10/19] [TRTLLM-16304][infra] Add a CBTS rule for modeling_v2 Without a rule, any change under tensorrt_llm/_torch/modeling_v2/ is an unclaimed residual and CBTS falls back to the baseline filter chain -- the whole pre-merge set, for a subtree whose tests are six entries in two blocks. The rule is the plainest of its family because the subtree was laid out that way. Its entries are matched on `unittest/_torch/modeling_v2/` and `test_modeling_v2_`, and both are exact rather than lucky: every test file in the subtree carries that prefix precisely so it cannot collide with the upstream test of the same op. SpecDecRule needs an `mtp_nextn=0` carve-out because its markers are substrings of entries it does not own; nothing here has that problem. No outward-facing fallback list either, for the same kind of reason. Nothing imports the subtree unless TRTLLM_MODELING_V2 is set, so there is no eagerly imported file whose edit has to force a full run. The one caller outside it, _torch/models/modeling_auto.py, is deliberately left unclaimed: a change to the shared resolver should fall back to baseline. .md is excluded, and that matters more here than elsewhere -- every catalog entry ships a contract document, so a fifth of the subtree is Markdown, and claiming it would let a docs-only PR pull in multi-GPU GB300 stages. Verified by driving main.py over five inputs: source-only narrows to modelingv2only (2 blocks, 4 stages, all six entries and nothing else); contract-only lands on noop; modeling_auto.py falls back; source plus an unrelated file falls back; source plus its own tests combines through the testsonly family. The other rules were compared against the pre-change main.py on four inputs and are byte-identical in outcome. Signed-off-by: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> --- jenkins/scripts/cbts/README.md | 6 +- jenkins/scripts/cbts/main.py | 4 + jenkins/scripts/cbts/rules/README.md | 51 ++++++ .../scripts/cbts/rules/modeling_v2_rule.py | 166 ++++++++++++++++++ 4 files changed, 225 insertions(+), 2 deletions(-) create mode 100644 jenkins/scripts/cbts/rules/modeling_v2_rule.py diff --git a/jenkins/scripts/cbts/README.md b/jenkins/scripts/cbts/README.md index f9416a477a62..d5e751bb2641 100644 --- a/jenkins/scripts/cbts/README.md +++ b/jenkins/scripts/cbts/README.md @@ -37,7 +37,7 @@ filter chain. ## Rules -Nine rules, registered in `main.py::RULE_CLASSES`: +Ten rules, registered in `main.py::RULE_CLASSES`: | Rule | Scope | Files | |---|---|---| @@ -46,6 +46,7 @@ Nine rules, registered in `main.py::RULE_CLASSES`: | `TestListRule` | `testlistonly` | `tests/integration/test_lists/test-db/*.yml` | | `VisualGenRule` | `visualgenonly` | `examples/visual_gen/**`, `scripts/visualgen_eval/**`, `tensorrt_llm/_torch/visual_gen/**`, `tensorrt_llm/media/**`, `tensorrt_llm/visual_gen/**` (excl. `.md`; reference images such as `cat_piano.png` ARE test fixtures and stay claimed; outward-facing files force fallback) | | `SpecDecRule` | `specdeconly` | `tensorrt_llm/_torch/speculative/**`, `tensorrt_llm/models/{eagle,medusa,redrafter}/**`, `examples/{eagle,medusa,redrafter,draft_target_model,ngram}/**`, `examples/llm-api/llm_speculative_decoding.py` (excl. `.md`; other suffixes incl. images kept as potential test fixtures) | +| `ModelingV2Rule` | `modelingv2only` | `tensorrt_llm/_torch/modeling_v2/**` (excl. `.md`) | | `AgentFlowRule` | `agentflowonly` | `agent-flow/**` (excl. `.md`) | | `OpenEngineRule` | `openengineonly` | `tensorrt_llm/grpc/openengine/**` (excl. `.md`) | | `OutOfScopeRule` | `noop` | QA / dev test lists, `.test_durations`, `microbenchmarks/`, `**/*.md` (image suffixes intentionally not claimed — image fixtures cannot be distinguished from doc diagrams by location, so image edits fall back to baseline) | @@ -61,9 +62,10 @@ See `rules/README.md` for per-rule logic. | `testlistonly` | `TestListRule` fired solo: PR only adds entries under `tests/integration/test_lists/test-db/*.yml`. | | `visualgenonly` | `VisualGenRule` fired solo: PR only touches VisualGen internal source paths (`examples/visual_gen/**`, `scripts/visualgen_eval/**`, `tensorrt_llm/_torch/visual_gen/**`; excl. `.md`; image fixtures like `cat_piano.png` are claimed). Narrows to blocks containing VG test entries. Outward-facing files under `tensorrt_llm/visual_gen/**` and `tensorrt_llm/media/**` (eagerly imported by `trtllm-serve`) force `null` fallback. | | `specdeconly` | `SpecDecRule` fired solo: PR only touches speculative-decoding source paths (`tensorrt_llm/_torch/speculative/**`, `tensorrt_llm/models/{eagle,medusa,redrafter}/**`, `examples/{eagle,medusa,redrafter,draft_target_model,ngram}/**`, `examples/llm-api/llm_speculative_decoding.py`; excl. `.md`). Narrows to blocks containing spec-dec test entries (eagle / medusa / redrafter / ngram / draft-target-model / MTP). | +| `modelingv2only` | `ModelingV2Rule` fired solo: PR only touches the modeling_v2 subtree (`tensorrt_llm/_torch/modeling_v2/**`; excl. `.md`, which is a fifth of the subtree — every catalog entry carries a contract document). Narrows to blocks containing modeling_v2 test entries (`unittest/_torch/modeling_v2/`, `test_modeling_v2_*`). No outward-facing fallback is needed: nothing imports the subtree unless `TRTLLM_MODELING_V2` is set, and its one caller outside the subtree (`_torch/models/modeling_auto.py`) is left unclaimed, so touching the shared resolver falls back to baseline. | | `agentflowonly` | `AgentFlowRule` fired solo: PR only touches `agent-flow/**` source or test files (excl. `.md`). Runs `CPU-AgentFlow-UnitTest`. | | `openengineonly` | `OpenEngineRule` fired solo: PR only touches `tensorrt_llm/grpc/openengine/**` source files (excl. `.md`). Narrows to the registered OpenEngine unit tests: the stub-based ones on the always-run `CPU-Generic-*` stages, plus `test_capability_conformance.py` on `A10-PyTorch-*`, which needs a GPU. | -| `testsonly` | Multiple rules from the testsonly family fired (`waiveonly`, `testdefonly`, `testlistonly`, `visualgenonly`, `specdeconly`, `agentflowonly`, `openengineonly`); their narrows union. | +| `testsonly` | Multiple rules from the testsonly family fired (`waiveonly`, `testdefonly`, `testlistonly`, `visualgenonly`, `specdeconly`, `modelingv2only`, `agentflowonly`, `openengineonly`); their narrows union. | | `noop` | Rule(s) fired but determined no test stages need to run (QA-only path, removals-only test list, all-miss waives, in-namespace .py with no covering YAML entry, docs-only edits). Layer 2 still applies. | | `null` (fallback) | A rule cannot decide, scopes don't combine, or there are unhandled files. Groovy defers to baseline filter chain. | diff --git a/jenkins/scripts/cbts/main.py b/jenkins/scripts/cbts/main.py index d4fece092f01..cfdb291886ce 100644 --- a/jenkins/scripts/cbts/main.py +++ b/jenkins/scripts/cbts/main.py @@ -61,6 +61,7 @@ from rules._helpers import strip_noop_diff_lines # noqa: E402 from rules.agent_flow_rule import AgentFlowRule # noqa: E402 from rules.base import PRInputs, Rule, RuleResult, format_reason # noqa: E402 +from rules.modeling_v2_rule import ModelingV2Rule # noqa: E402 from rules.openengine_rule import OpenEngineRule # noqa: E402 from rules.out_of_scope_rule import OutOfScopeRule # noqa: E402 from rules.spec_dec_rule import SpecDecRule # noqa: E402 @@ -81,6 +82,7 @@ TestListRule, VisualGenRule, SpecDecRule, + ModelingV2Rule, AgentFlowRule, OpenEngineRule, OutOfScopeRule, @@ -98,6 +100,7 @@ def build_rules( TestListRule(yaml_index, stages, repo_root=repo_root), VisualGenRule(yaml_index, stages), SpecDecRule(yaml_index, stages), + ModelingV2Rule(yaml_index, stages), AgentFlowRule(yaml_index, stages), OpenEngineRule(yaml_index, stages), OutOfScopeRule(yaml_index, stages), @@ -217,6 +220,7 @@ def _rule_reason(rule, r) -> dict: "testlistonly", "visualgenonly", "specdeconly", + "modelingv2only", "agentflowonly", "openengineonly", } diff --git a/jenkins/scripts/cbts/rules/README.md b/jenkins/scripts/cbts/rules/README.md index dac37edebe0e..d2038abe5bbf 100644 --- a/jenkins/scripts/cbts/rules/README.md +++ b/jenkins/scripts/cbts/rules/README.md @@ -13,6 +13,7 @@ for the overall CBTS architecture. | `test_list_rule.py` | `TestListRule` | `testlistonly` | `tests/integration/test_lists/test-db/*.yml` | | `visual_gen_rule.py` | `VisualGenRule` | `visualgenonly` | `examples/visual_gen/**`, `scripts/visualgen_eval/**`, `tensorrt_llm/_torch/visual_gen/**`, `tensorrt_llm/media/**`, `tensorrt_llm/visual_gen/**` (each excl. `.md`) | | `spec_dec_rule.py` | `SpecDecRule` | `specdeconly` | `tensorrt_llm/_torch/speculative/**`, `tensorrt_llm/models/{eagle,medusa,redrafter}/**`, `examples/{eagle,medusa,redrafter,draft_target_model,ngram}/**`, `examples/llm-api/llm_speculative_decoding.py` (each excl. `.md`) | +| `modeling_v2_rule.py` | `ModelingV2Rule` | `modelingv2only` | `tensorrt_llm/_torch/modeling_v2/**` (excl. `.md`) | | `agent_flow_rule.py` | `AgentFlowRule` | `agentflowonly` | `agent-flow/**` (excl. `.md`) → the single `CPU-AgentFlow-UnitTest` stage; not test-db-driven | | `openengine_rule.py` | `OpenEngineRule` | `openengineonly` | `tensorrt_llm/grpc/openengine/**` (excl. `.md`) → the `l0_cpu` block containing `unittest/grpc/openengine/` | | `out_of_scope_rule.py` | `OutOfScopeRule` | `noop` | `tests/integration/test_lists/{qa,dev}/**`, `tests/integration/defs/.test_durations*`, `tests/microbenchmarks/**`, `**/*.md` (image suffixes intentionally not claimed — fall back to baseline since fixtures and doc diagrams are indistinguishable by location) | @@ -254,6 +255,56 @@ Outcomes: - Spec-dec source touched but no spec-dec block found anywhere (defensive) → `scope=None` (fallback). +## ModelingV2Rule + +Path-only rule. Claims non-documentation source changes under +`tensorrt_llm/_torch/modeling_v2/`, the self-contained second modeling +path (one flat forward per checkpoint/arch/parallel triple, assembled +from a catalog of op wrappers). + +`.md` exclusion carries more weight here than elsewhere: every catalog +entry ships a contract document, so roughly a fifth of the subtree is +Markdown. Claiming those would let a documentation-only PR pull in +multi-GPU GB300 stages. Other suffixes are NOT excluded — a data file +under the subtree could be a fixture, so the rule keeps claiming it +(safe over-run). + +Block selection — entry-pattern based only: +modeling_v2 has no `condition.terms.backend` of its own; its entries sit +in `backend: pytorch` blocks beside everything else. A block belongs to +modeling_v2 iff one of its `tests:` entries matches +`_MV2_ENTRY_PATTERNS`: + +- `unittest/_torch/modeling_v2/` — the op-level catalog matrix, carried + as whole-directory entries (one on `l0_gb300`, one on + `l0_gb300_multi_gpus` for the 4-rank collectives). +- `test_modeling_v2_` — the accuracy gates, and any future unit file. + +Both markers are exact by construction rather than by luck: every test +file in the subtree is named `test_modeling_v2_*` precisely so it cannot +collide with the upstream test of the same op. So unlike `SpecDecRule`'s +`mtp_nextn`, no substring here can claim an unrelated entry and no +carve-out is needed. + +Outward fallback: not needed, and by design rather than by accident. +Nothing imports the subtree unless `TRTLLM_MODELING_V2` is set — +`AutoModelForCausalLM._resolve_class` calls `modeling_v2_resolve`, which +returns immediately when the switch is off, and the routing modules are +imported lazily behind it. The one caller outside the subtree, +`tensorrt_llm/_torch/models/modeling_auto.py`, is deliberately left +unclaimed: a change to the shared resolver falls back to baseline, which +is what it deserves. + +`sanity_relevant=False` — the subtree ships no user-facing entry point +and is not imported by `trtllm-serve` or by `import tensorrt_llm`, so +none of it is what PackageSanityCheck verifies about the wheel. +`perfsanity_relevant` is dynamic (True only if a matched block lives in +a `*_perf_sanity*` yaml); there are no modeling_v2 perf-sanity entries +today, so it aggregates to False. + +Source changed but no modeling_v2 block in any yaml (defensive) → +`scope=None` (fallback). + ## OpenEngineRule Claims non-documentation source changes under `tensorrt_llm/grpc/openengine/` and keeps only test-db diff --git a/jenkins/scripts/cbts/rules/modeling_v2_rule.py b/jenkins/scripts/cbts/rules/modeling_v2_rule.py new file mode 100644 index 000000000000..68e05d963be7 --- /dev/null +++ b/jenkins/scripts/cbts/rules/modeling_v2_rule.py @@ -0,0 +1,166 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""ModelingV2Rule — narrows CI when the modeling_v2 subtree changes. + +modeling_v2 is a second modeling path living entirely under +`tensorrt_llm/_torch/modeling_v2/`: one self-contained forward per +(checkpoint, GPU arch, parallel topology), assembled from a catalog of +op wrappers. + +Block selection — entry-pattern based only: +It has no `condition.terms.backend` of its own; its entries sit in +`backend: pytorch` blocks beside everything else. A block belongs to +modeling_v2 iff one of its `tests:` entries matches a marker in +`_MV2_ENTRY_PATTERNS`. Those markers are exact by construction rather +than by luck: every unit test file in the subtree is named +`test_modeling_v2_*` precisely so it cannot collide with the upstream +test of the same op, and the accuracy files follow the same prefix. So +there is no substring that could claim an unrelated entry, and no +`mtp_nextn=0`-style carve-out is needed. + +Outward fallback: not needed, and that is a property of the design +rather than an accident. Nothing imports this subtree unless +`TRTLLM_MODELING_V2` is set: `AutoModelForCausalLM._resolve_class` calls +`modeling_v2_resolve`, which returns immediately when the switch is off, +and the routing modules are imported lazily behind it. The one caller +outside the subtree is `tensorrt_llm/_torch/models/modeling_auto.py`, +which this rule does not claim -- a PR touching it falls back to +baseline, which is what a change to the shared resolver deserves. + +`.md` exclusion matters more here than for most rules: the catalog +carries a contract document per entry, so roughly a fifth of the files +in the subtree are Markdown. Claiming them would make a +documentation-only PR pull in multi-GPU GB300 stages. + +PerfSanity policy: `perfsanity_relevant` is dynamic, True only when a +matched block lives in a `*_perf_sanity*` yaml -- same as AutoDeployRule +/ VisualGenRule / SpecDecRule. modeling_v2 has no perf-sanity entries +today, so this aggregates to False and Groovy Layer 2 drops the +force-keep of `*-PerfSanity-*` stages. + +Sanity policy: `sanity_relevant=False`. The subtree ships no +user-facing entry point and is not imported by `trtllm-serve` or by +`import tensorrt_llm`, so nothing it contains is what PackageSanityCheck +verifies about the wheel. +""" + +from __future__ import annotations + +from typing import Optional + +from blocks import Stage, YAMLIndex, _entry_target + +from ._helpers import resolve_affected_stages, stages_by_yaml_stem +from .base import PRInputs, Rule, RuleResult + +# Source-path prefixes the rule may claim. Tests under tests/** are left +# to TestsDefRule; the two scopes combine via _TESTSONLY_FAMILY. +_MV2_SRC_PREFIXES: tuple[str, ...] = ("tensorrt_llm/_torch/modeling_v2/",) + +# Substrings that mark a test entry as modeling_v2. Both are unambiguous: +# - "unittest/_torch/modeling_v2/" → the op-level catalog matrix, taken +# as whole-directory entries (one on l0_gb300, one on +# l0_gb300_multi_gpus for the 4-rank collectives) +# - "test_modeling_v2_" → the accuracy gates, and any future unit file +# named by the subtree's own convention +_MV2_ENTRY_PATTERNS: tuple[str, ...] = ( + "unittest/_torch/modeling_v2/", + "test_modeling_v2_", +) + + +def _is_mv2_claim(path: str) -> bool: + """Decide whether ModelingV2Rule claims `path`. + + `*.md` is excluded so a contract-only edit does not force GPU stages + -- `OutOfScopeRule` claims those as noop instead. Other suffixes are + NOT excluded: a data file under this subtree could be a fixture, so + the rule keeps claiming it and re-runs the stages (safe over-run). + """ + if not path.startswith(_MV2_SRC_PREFIXES): + return False + if path.endswith(".md"): + return False + return True + + +def _entry_is_mv2(entry: str) -> bool: + return any(p in entry for p in _MV2_ENTRY_PATTERNS) + + +def _mv2_entries(block) -> list[str]: + return [t for t in block.tests if _entry_is_mv2(t)] + + +def _is_perf_sanity_stem(stem: str) -> bool: + """True for perf-sanity yaml stems (`l0_*_perf_sanity*`).""" + return "perf_sanity" in stem + + +class ModelingV2Rule(Rule): + name = "modelingv2" + needs_diff_for: tuple[str, ...] = () + + def __init__(self, yaml_index: YAMLIndex, stages: dict[str, Stage]) -> None: + self.yaml_index = yaml_index + self._stages_by_yaml = stages_by_yaml_stem(stages) + + def apply(self, pr: PRInputs) -> Optional[RuleResult]: + claimed = {f for f in pr.changed_files if _is_mv2_claim(f)} + if not claimed: + return None + + block_filters: dict[tuple[str, int], dict[str, set[str]]] = {} + for block in self.yaml_index.blocks: + entries = _mv2_entries(block) + if not entries: + continue + key = (block.yaml_stem, block.block_index) + prefix_dict = block_filters.setdefault(key, {}) + for entry in entries: + target = _entry_target(entry) + if target: + prefix_dict.setdefault(target, set()).add(entry) + + if not block_filters: + # Defensive: modeling_v2 source changed but no modeling_v2 block + # exists in any yaml. Do not fabricate stages -- fall back to + # baseline so the change still gets coverage. Reachable if the + # subtree's entries are ever removed from the test-db without + # the subtree going with them. + return RuleResult( + handled_files=claimed, + affected_stages=set(), + scope=None, + reason=( + f"modelingv2: {len(claimed)} modeling_v2 source file(s); " + "no modeling_v2 block matched in any test-db yaml — fallback" + ), + ) + + affected = resolve_affected_stages(block_filters, self.yaml_index, self._stages_by_yaml) + perfsanity_relevant = any(_is_perf_sanity_stem(stem) for stem, _ in block_filters) + + return RuleResult( + handled_files=claimed, + affected_stages=affected, + scope="modelingv2only", + block_filters=block_filters, + sanity_relevant=False, + perfsanity_relevant=perfsanity_relevant, + reason=( + f"modelingv2: {len(claimed)} modeling_v2 source file(s) → " + f"{len(block_filters)} modeling_v2 block(s), {len(affected)} stage(s)" + ), + ) From 6ae72b55ef5325c022b393a086a07300e2435523 Mon Sep 17 00:00:00 2001 From: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> Date: Wed, 16 Sep 2026 06:40:35 -0700 Subject: [PATCH 11/19] [TRTLLM-16304][test] Cut the catalog tests down to what the targets run A catalog entry's test exists to catch a regression in the op the shipped targets call. It was covering a good deal more than that: 344 cases, of which 92 drove configurations no target reaches. What the two targets actually run is narrow and checkable. gpt-oss is bf16 activations, GQA 64q/8kv at head_dim 64, sinks, a 128-wide sliding window alternating with full attention, and a bf16 KV pool. R1-0528 is MLA at 128 heads with an fp8 KV pool on every call, q_lora_rank 1536, and tokens_per_block 32 -- the last three asserted in the target itself, so the alternatives are not merely unused but refused. Measured against that: bf16 MLA (30 cases) when every MLA call is quant_mode=128; head counts h8/h16/h32 (14) when the target is h128; page-size sweeps (14) against an assert pinning 32; fp16 and fp32 (15) when both targets are bf16; GQA at d128 (6) when gpt-oss is d64; and the reference-discriminating controls (5), which check the reference rather than the op. Two corrections while doing it, both worth stating because the first would have been silent. `test_fp8_kv_*` means different things in different files: in thop_attention it is a non-MLA fp8 GQA path no target runs, but in the three MLA entries it IS the production path. A name-only rule deletes the latter. And the two-pool routing cell was caught by the d128 in its name though multi-pool is exactly what gpt-oss's window split produces -- it stays, re-parameterized to the real d64 geometry. 18 helpers left with no caller go too. The two autouse fixtures in test_modeling_v2_routing.py do not: they are reachable without a name reference, which is what an autouse fixture is for, and removing them dropped the sm_103 stub and the default mode so routing quietly returned the un-rewritten architecture. Restored. Receipt counts are updated, and the run that re-earns them is green: 267 cases on the 1-GPU entry, 2 on the 4-GPU one. Signed-off-by: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> --- .../activation/flashinfer_silu_and_mul.md | 2 +- .../attention/load_paged_kv_cache_for_mla.md | 2 +- .../mla_rope_append_paged_kv_assign_q.md | 2 +- .../catalog/attention/mla_rope_generation.md | 2 +- .../catalog/attention/thop_attention.md | 2 +- .../modeling_v2/catalog/gemm/bmm_out.md | 2 +- .../modeling_v2/catalog/gemm/cublas_mm.md | 2 +- .../modeling_v2/catalog/gemm/nvfp4_gemm.md | 2 +- .../catalog/moe/fp4_block_scale_moe_runner.md | 2 +- .../modeling_v2/catalog/moe/fused_moe.md | 2 +- .../mxe4m3_mxe2m1_block_scale_moe_runner.md | 2 +- .../norm/flashinfer_fused_add_rmsnorm.md | 2 +- .../catalog/norm/flashinfer_rmsnorm.md | 2 +- .../catalog/quantization/fp4_quantize.md | 2 +- .../catalog/quantization/mxfp8_quantize.md | 2 +- ...est_modeling_v2_flashinfer_silu_and_mul.py | 7 - ...modeling_v2_load_paged_kv_cache_for_mla.py | 235 --- ...ng_v2_mla_rope_append_paged_kv_assign_q.py | 169 -- .../test_modeling_v2_mla_rope_generation.py | 286 --- .../test_modeling_v2_thop_attention.py | 1535 +---------------- .../gemm/test_modeling_v2_bmm_out.py | 12 - .../gemm/test_modeling_v2_cublas_mm.py | 17 - ..._modeling_v2_fp4_block_scale_moe_runner.py | 84 +- .../moe/test_modeling_v2_fused_moe.py | 111 +- ...v2_mxe4m3_mxe2m1_block_scale_moe_runner.py | 56 +- ...odeling_v2_flashinfer_fused_add_rmsnorm.py | 18 - .../test_modeling_v2_flashinfer_rmsnorm.py | 16 - .../test_modeling_v2_fp4_quantize.py | 78 - .../test_modeling_v2_mxfp8_quantize.py | 11 - 29 files changed, 37 insertions(+), 2628 deletions(-) diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/activation/flashinfer_silu_and_mul.md b/tensorrt_llm/_torch/modeling_v2/catalog/activation/flashinfer_silu_and_mul.md index 2bc9c9720d67..71e421866607 100644 --- a/tensorrt_llm/_torch/modeling_v2/catalog/activation/flashinfer_silu_and_mul.md +++ b/tensorrt_llm/_torch/modeling_v2/catalog/activation/flashinfer_silu_and_mul.md @@ -1,6 +1,6 @@ --- receipts: - sm_103: {status: passed, tests: 4} + sm_103: {status: passed, tests: 5} --- # flashinfer_silu_and_mul diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/attention/load_paged_kv_cache_for_mla.md b/tensorrt_llm/_torch/modeling_v2/catalog/attention/load_paged_kv_cache_for_mla.md index e51fee753fe6..d6ed03181e5a 100644 --- a/tensorrt_llm/_torch/modeling_v2/catalog/attention/load_paged_kv_cache_for_mla.md +++ b/tensorrt_llm/_torch/modeling_v2/catalog/attention/load_paged_kv_cache_for_mla.md @@ -1,6 +1,6 @@ --- receipts: - sm_103: {status: passed, tests: 18} + sm_103: {status: passed, tests: 10} --- # load_paged_kv_cache_for_mla diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_append_paged_kv_assign_q.md b/tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_append_paged_kv_assign_q.md index 4be4de772b8e..768f3c875a12 100644 --- a/tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_append_paged_kv_assign_q.md +++ b/tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_append_paged_kv_assign_q.md @@ -1,6 +1,6 @@ --- receipts: - sm_103: {status: passed, tests: 13} + sm_103: {status: passed, tests: 6} --- # mla_rope_append_paged_kv_assign_q diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_generation.md b/tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_generation.md index 90f743a2b19f..a28e160aa6c0 100644 --- a/tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_generation.md +++ b/tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_generation.md @@ -1,6 +1,6 @@ --- receipts: - sm_103: {status: passed, tests: 22} + sm_103: {status: passed, tests: 15} --- # mla_rope_generation diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/attention/thop_attention.md b/tensorrt_llm/_torch/modeling_v2/catalog/attention/thop_attention.md index ebee2724729a..94f95df61644 100644 --- a/tensorrt_llm/_torch/modeling_v2/catalog/attention/thop_attention.md +++ b/tensorrt_llm/_torch/modeling_v2/catalog/attention/thop_attention.md @@ -1,6 +1,6 @@ --- receipts: - sm_103: {status: passed, tests: 96} + sm_103: {status: passed, tests: 47} --- # thop_attention diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/gemm/bmm_out.md b/tensorrt_llm/_torch/modeling_v2/catalog/gemm/bmm_out.md index 39e4685606d7..9730280f57a7 100644 --- a/tensorrt_llm/_torch/modeling_v2/catalog/gemm/bmm_out.md +++ b/tensorrt_llm/_torch/modeling_v2/catalog/gemm/bmm_out.md @@ -1,6 +1,6 @@ --- receipts: - sm_103: {status: passed, tests: 5} + sm_103: {status: passed, tests: 3} --- # bmm_out diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/gemm/cublas_mm.md b/tensorrt_llm/_torch/modeling_v2/catalog/gemm/cublas_mm.md index d2ca81837268..7b32f301dc29 100644 --- a/tensorrt_llm/_torch/modeling_v2/catalog/gemm/cublas_mm.md +++ b/tensorrt_llm/_torch/modeling_v2/catalog/gemm/cublas_mm.md @@ -1,6 +1,6 @@ --- receipts: - sm_103: {status: passed, tests: 7} + sm_103: {status: passed, tests: 5} --- # cublas_mm diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/gemm/nvfp4_gemm.md b/tensorrt_llm/_torch/modeling_v2/catalog/gemm/nvfp4_gemm.md index 919bcfbca639..98d9f02bec2b 100644 --- a/tensorrt_llm/_torch/modeling_v2/catalog/gemm/nvfp4_gemm.md +++ b/tensorrt_llm/_torch/modeling_v2/catalog/gemm/nvfp4_gemm.md @@ -1,6 +1,6 @@ --- receipts: - sm_103: {status: passed, tests: 19} + sm_103: {status: passed, tests: 20} --- # nvfp4_gemm diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/moe/fp4_block_scale_moe_runner.md b/tensorrt_llm/_torch/modeling_v2/catalog/moe/fp4_block_scale_moe_runner.md index a13c6255ddc4..a8e1f3efdb97 100644 --- a/tensorrt_llm/_torch/modeling_v2/catalog/moe/fp4_block_scale_moe_runner.md +++ b/tensorrt_llm/_torch/modeling_v2/catalog/moe/fp4_block_scale_moe_runner.md @@ -1,6 +1,6 @@ --- receipts: - sm_103: {status: passed, tests: 16} + sm_103: {status: passed, tests: 15} --- # fp4_block_scale_moe_runner diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/moe/fused_moe.md b/tensorrt_llm/_torch/modeling_v2/catalog/moe/fused_moe.md index 2446edb87003..2fa0f34c4814 100644 --- a/tensorrt_llm/_torch/modeling_v2/catalog/moe/fused_moe.md +++ b/tensorrt_llm/_torch/modeling_v2/catalog/moe/fused_moe.md @@ -1,6 +1,6 @@ --- receipts: - sm_103: {status: passed, tests: 23} + sm_103: {status: passed, tests: 17} --- # fused_moe diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md b/tensorrt_llm/_torch/modeling_v2/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md index 56fd242811b3..168e1448fb77 100644 --- a/tensorrt_llm/_torch/modeling_v2/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md +++ b/tensorrt_llm/_torch/modeling_v2/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md @@ -1,6 +1,6 @@ --- receipts: - sm_103: {status: passed, tests: 19} + sm_103: {status: passed, tests: 18} --- # mxe4m3_mxe2m1_block_scale_moe_runner diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_fused_add_rmsnorm.md b/tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_fused_add_rmsnorm.md index f1809bcb9d2b..b8e61d3bc7e1 100644 --- a/tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_fused_add_rmsnorm.md +++ b/tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_fused_add_rmsnorm.md @@ -1,6 +1,6 @@ --- receipts: - sm_103: {status: passed, tests: 5} + sm_103: {status: passed, tests: 3} --- # flashinfer_fused_add_rmsnorm diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_rmsnorm.md b/tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_rmsnorm.md index aa067e75e405..490f1d9f71ac 100644 --- a/tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_rmsnorm.md +++ b/tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_rmsnorm.md @@ -1,6 +1,6 @@ --- receipts: - sm_103: {status: passed, tests: 6} + sm_103: {status: passed, tests: 4} --- # flashinfer_rmsnorm diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/quantization/fp4_quantize.md b/tensorrt_llm/_torch/modeling_v2/catalog/quantization/fp4_quantize.md index f82b0b2eceaf..f8e81a788774 100644 --- a/tensorrt_llm/_torch/modeling_v2/catalog/quantization/fp4_quantize.md +++ b/tensorrt_llm/_torch/modeling_v2/catalog/quantization/fp4_quantize.md @@ -1,6 +1,6 @@ --- receipts: - sm_103: {status: passed, tests: 21} + sm_103: {status: passed, tests: 19} --- # fp4_quantize diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/quantization/mxfp8_quantize.md b/tensorrt_llm/_torch/modeling_v2/catalog/quantization/mxfp8_quantize.md index 221d652b0c97..83275c688373 100644 --- a/tensorrt_llm/_torch/modeling_v2/catalog/quantization/mxfp8_quantize.md +++ b/tensorrt_llm/_torch/modeling_v2/catalog/quantization/mxfp8_quantize.md @@ -1,6 +1,6 @@ --- receipts: - sm_103: {status: passed, tests: 13} + sm_103: {status: passed, tests: 12} --- # mxfp8_quantize diff --git a/tests/unittest/_torch/modeling_v2/activation/test_modeling_v2_flashinfer_silu_and_mul.py b/tests/unittest/_torch/modeling_v2/activation/test_modeling_v2_flashinfer_silu_and_mul.py index ea2dfc783cda..02384727943c 100644 --- a/tests/unittest/_torch/modeling_v2/activation/test_modeling_v2_flashinfer_silu_and_mul.py +++ b/tests/unittest/_torch/modeling_v2/activation/test_modeling_v2_flashinfer_silu_and_mul.py @@ -51,13 +51,6 @@ def test_bf16_edge_sizes() -> None: _check(x) -def test_fp16_2d() -> None: - torch.manual_seed(3) - for num_tokens, two_d in [(2, 8192), (1024, 4096)]: - x = torch.randn(num_tokens, two_d, dtype=torch.float16, device="cuda") - _check(x) - - def test_a_misaligned_half_is_rejected_before_dispatch() -> None: """A final dim of 24 passes the op's own check and faults the kernel. diff --git a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_load_paged_kv_cache_for_mla.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_load_paged_kv_cache_for_mla.py index b06197c1a7f5..2554428a6cae 100644 --- a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_load_paged_kv_cache_for_mla.py +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_load_paged_kv_cache_for_mla.py @@ -245,41 +245,6 @@ def _check_outputs( assert compressed_kv.is_contiguous() and k_pe.is_contiguous() -def _run_and_check( - env: _MlaCacheEnv, - request_ids: List[int], - seq_lens: List[int], - num_contexts: int, - cached_lens: List[int], - layer_idx: int = 0, -) -> None: - """Fill the context sequences' cache rows, run the op, and verify the - gathered compressed_kv / k_pe against the written rows.""" - kv_lens = [c + s for c, s in zip(cached_lens, seq_lens)] - metadata = env.prepare_metadata(request_ids, seq_lens, num_contexts, cached_lens) - ctx_kv_lens = kv_lens[:num_contexts] - stored, _ = env.fill_layer(layer_idx, request_ids[:num_contexts], ctx_kv_lens) - - total_ctx_kv = int(metadata.num_ctx_cached_tokens + metadata.num_ctx_tokens) - assert total_ctx_kv == sum(ctx_kv_lens) - assert int(metadata.max_ctx_kv_len) == max(ctx_kv_lens) - - compressed_kv, k_pe = _call( - env, - metadata, - num_contexts, - env.torch_dtype, - None, # kv_scale_quant_orig: fp8-KV-cache path only - layer_idx=layer_idx, - ) - - expected = torch.cat(stored, dim=0) # [total_ctx_kv, HEAD_SIZE] - _check_outputs(compressed_kv, k_pe, total_ctx_kv, env.torch_dtype) - # Dtype-preserving gather: outputs must be bitwise equal to the cache rows. - torch.testing.assert_close(compressed_kv, expected[:, :KV_LORA_RANK], rtol=0.0, atol=0.0) - torch.testing.assert_close(k_pe, expected[:, KV_LORA_RANK:], rtol=0.0, atol=0.0) - - def _dequant_mirror( stored: torch.Tensor, read_scale: Optional[torch.Tensor], @@ -360,171 +325,6 @@ def _raises(fn) -> str: # ───────────────────────── matching-dtype pool ────────────────────────── -def test_bf16_mixed_lengths_layer1() -> None: - """Prefill-like batch (~1k gathered tokens): cached lengths of 0, exactly - one block, and multi-block; addressed through pool-mapping row 1 of a - two-layer pool.""" - torch.manual_seed(0) - env = _MlaCacheEnv(DataType.BF16, num_layers=2) - try: - cached = [0, 64, 100, 511] - new = [37, 64, 300, 1] - rids = [0, 1, 2, 3] - env.kv_cache_manager.add_dummy_requests( - rids, token_nums=[c + n for c, n in zip(cached, new)] - ) - _run_and_check( - env, - request_ids=rids, - seq_lens=new, - num_contexts=4, - cached_lens=cached, - layer_idx=1, - ) - finally: - env.shutdown() - - -def test_bf16_single_tiny_sequence() -> None: - """Smallest cached-context case: one sequence, one cached token plus one - new token (two gathered rows).""" - torch.manual_seed(1) - env = _MlaCacheEnv(DataType.BF16) - try: - env.kv_cache_manager.add_dummy_requests([0], token_nums=[2]) - _run_and_check( - env, - request_ids=[0], - seq_lens=[1], - num_contexts=1, - cached_lens=[1], - ) - finally: - env.shutdown() - - -def test_bf16_trailing_generation_seqs_ignored() -> None: - """Mixed batch: two context sequences followed by two generation - sequences. The op must gather exactly the context sequences' rows and - index per-seq tensors over [0, num_contexts) only.""" - torch.manual_seed(2) - env = _MlaCacheEnv(DataType.BF16) - try: - cached = [128, 3, 200, 77] - new = [40, 60, 1, 1] - rids = [0, 1, 2, 3] - env.kv_cache_manager.add_dummy_requests( - rids, token_nums=[c + n for c, n in zip(cached, new)] - ) - _run_and_check( - env, - request_ids=rids, - seq_lens=new, - num_contexts=2, - cached_lens=cached, - ) - finally: - env.shutdown() - - -def test_fp16_mixed_lengths() -> None: - """fp16 latent cache with fp16 out_dtype over block-crossing lengths.""" - torch.manual_seed(3) - env = _MlaCacheEnv(DataType.HALF) - try: - cached = [65, 640, 1] - new = [63, 128, 6] - rids = [0, 1, 2] - env.kv_cache_manager.add_dummy_requests( - rids, token_nums=[c + n for c, n in zip(cached, new)] - ) - _run_and_check( - env, - request_ids=rids, - seq_lens=new, - num_contexts=3, - cached_lens=cached, - ) - finally: - env.shutdown() - - -def test_bf16_page32_mixed_lengths_layer1() -> None: - """Page size 32 (the engine default). Gathered ranges of 37, 96, 400 and - 512 rows: 96 and 512 fill 3 and 16 pages exactly at 32 (neither is a - page multiple at 64), 400 spans 12 full pages plus 16 rows, 37 spans two - — every sequence crosses at least one boundary and the deepest one walks - 16 offsets-row entries. Addressed through pool-mapping row 1 of a - two-layer pool.""" - torch.manual_seed(4) - env = _MlaCacheEnv(DataType.BF16, num_layers=2, tokens_per_block=PAGE32) - try: - cached = [0, 32, 100, 511] - new = [37, 64, 300, 1] - rids = [0, 1, 2, 3] - env.kv_cache_manager.add_dummy_requests( - rids, token_nums=[c + n for c, n in zip(cached, new)] - ) - _run_and_check( - env, - request_ids=rids, - seq_lens=new, - num_contexts=4, - cached_lens=cached, - layer_idx=1, - ) - finally: - env.shutdown() - - -def test_bf16_page32_trailing_generation_seqs_ignored() -> None: - """Page size 32, mixed batch: two context sequences followed by two - generation sequences. The gathered lengths (128 = 4 exact pages, 63 = - two pages minus one row) sit either side of a page boundary, and the - ignored generation sequences own pages of their own.""" - torch.manual_seed(5) - env = _MlaCacheEnv(DataType.BF16, tokens_per_block=PAGE32) - try: - cached = [96, 32, 200, 77] - new = [32, 31, 1, 1] - rids = [0, 1, 2, 3] - env.kv_cache_manager.add_dummy_requests( - rids, token_nums=[c + n for c, n in zip(cached, new)] - ) - _run_and_check( - env, - request_ids=rids, - seq_lens=new, - num_contexts=2, - cached_lens=cached, - ) - finally: - env.shutdown() - - -def test_fp16_page32_mixed_lengths() -> None: - """fp16 latent cache at page size 32: gathered ranges of 64 (2 exact - pages), 608 (19 exact pages) and 7 (a partial first page).""" - torch.manual_seed(6) - env = _MlaCacheEnv(DataType.HALF, tokens_per_block=PAGE32) - try: - cached = [33, 512, 0] - new = [31, 96, 7] - rids = [0, 1, 2] - env.kv_cache_manager.add_dummy_requests( - rids, token_nums=[c + n for c, n in zip(cached, new)] - ) - _run_and_check( - env, - request_ids=rids, - seq_lens=new, - num_contexts=3, - cached_lens=cached, - ) - finally: - env.shutdown() - - # ───────────────────────── fp8-e4m3 latent pool ───────────────────────── @@ -776,41 +576,6 @@ def test_fp8_two_layer_pool_layer1() -> None: env.shutdown() -def test_fp8_page64_prefill_scale() -> None: - """fp8 at page size 64 and prefill scale: 2183 gathered rows across - three sequences (16 exact pages, 16 exact pages, three partial), fp32 - output, write scale 0.5 undone by read scale 2.0.""" - torch.manual_seed(46) - env = _fp8_env(tokens_per_block=TOKENS_PER_BLOCK) - try: - cached = [512, 0, 128] - new = [512, 1024, 7] - rids = [0, 1, 2] - ckv, kpe, stored, source, scale_t = _run_fp8( - env, - request_ids=rids, - seq_lens=new, - num_contexts=3, - cached_lens=cached, - out_dtype=torch.float32, - read_scale=2.0, - write_scale=0.5, - ) - assert ckv.shape[0] == 2183, ckv.shape - _assert_fp8_exact(ckv, kpe, stored, scale_t, torch.float32) - # The consistent pair recovers the pre-quantization rows to within - # one e4m3 rounding. - got = torch.cat([ckv, kpe], dim=1) - torch.testing.assert_close( - got, - source, - rtol=E4M3_HALF_ULP_REL, - atol=0.5 * E4M3_MIN_SUBNORMAL * 2.0, - ) - finally: - env.shutdown() - - def test_fp8_byte_domain_and_out_dtype_range() -> None: """The dequantization is a plain conversion over the whole e4m3 domain: all 256 byte values (including +-448, +-0, subnormals and both NaN diff --git a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_mla_rope_append_paged_kv_assign_q.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_mla_rope_append_paged_kv_assign_q.py index 2b7560a11e98..2a9c7ccc4ffa 100644 --- a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_mla_rope_append_paged_kv_assign_q.py +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_mla_rope_append_paged_kv_assign_q.py @@ -501,175 +501,6 @@ def _run_and_check( ) -def test_bf16_mixed_cached_lengths_layer1() -> None: - """Prefill-like batch (~400 new tokens): cached prefixes of 0 (fresh - prefill), exactly one block, mid-block, and near max_seq_len; addressed - through pool-mapping row 1 of a two-layer pool.""" - torch.manual_seed(0) - env = _MlaCtxEnv(num_layers=2) - try: - cached = [0, 64, 100, 511] - new = [37, 64, 300, 1] - rids = [0, 1, 2, 3] - env.kv_cache_manager.add_dummy_requests( - rids, token_nums=[c + n for c, n in zip(cached, new)] - ) - _run_and_check( - env, - request_ids=rids, - seq_lens=new, - num_contexts=4, - cached_lens=cached, - layer_idx=1, - ) - finally: - env.shutdown() - - -def test_bf16_trailing_generation_seqs_ignored() -> None: - """Mixed batch: two context sequences followed by two generation - sequences. The op must touch only the context sequences' new rows and - index per-seq tensors over [0, num_contexts).""" - torch.manual_seed(1) - env = _MlaCtxEnv() - try: - cached = [128, 3, 200, 77] - new = [40, 60, 1, 1] - rids = [0, 1, 2, 3] - env.kv_cache_manager.add_dummy_requests( - rids, token_nums=[c + n for c, n in zip(cached, new)] - ) - _run_and_check( - env, - request_ids=rids, - seq_lens=new, - num_contexts=2, - cached_lens=cached, - ) - finally: - env.shutdown() - - -def test_bf16_single_token_append() -> None: - """Smallest cached-context case: one sequence, one cached token plus one - new token (decode-like single-row call, appended at position 1).""" - torch.manual_seed(2) - env = _MlaCtxEnv() - try: - env.kv_cache_manager.add_dummy_requests([0], token_nums=[2]) - _run_and_check( - env, - request_ids=[0], - seq_lens=[1], - num_contexts=1, - cached_lens=[1], - ) - finally: - env.shutdown() - - -def test_fp16_mixed_cached_lengths() -> None: - """fp16 activations with an fp16 latent cache over block-crossing - cached/new lengths.""" - torch.manual_seed(3) - env = _MlaCtxEnv(cache_dtype=DataType.HALF) - try: - cached = [65, 640, 1] - new = [63, 128, 6] - rids = [0, 1, 2] - env.kv_cache_manager.add_dummy_requests( - rids, token_nums=[c + n for c, n in zip(cached, new)] - ) - _run_and_check( - env, - request_ids=rids, - seq_lens=new, - num_contexts=3, - cached_lens=cached, - ) - finally: - env.shutdown() - - -def test_bf16_page32_mixed_cached_lengths_layer1() -> None: - """Page size 32 (the engine default). The four cached prefixes place the - first new token at every alignment a 32-token page has: 0 (fresh - prefill), 32 (page 1 slot 0, an exact boundary), 100 (page 3 slot 4, - mid-page) and 511 (page 15 slot 31, a page's last slot). The 300-token - sequence then walks ten pages in one call. Addressed through - pool-mapping row 1 of a two-layer pool.""" - torch.manual_seed(4) - env = _MlaCtxEnv(num_layers=2, tokens_per_block=PAGE32) - try: - cached = [0, 32, 100, 511] - new = [37, 64, 300, 1] - rids = [0, 1, 2, 3] - env.kv_cache_manager.add_dummy_requests( - rids, token_nums=[c + n for c, n in zip(cached, new)] - ) - _run_and_check( - env, - request_ids=rids, - seq_lens=new, - num_contexts=4, - cached_lens=cached, - layer_idx=1, - ) - finally: - env.shutdown() - - -def test_bf16_page32_trailing_generation_seqs_ignored() -> None: - """Page size 32, mixed batch: two context sequences followed by two - generation sequences. Sequence 0's new tokens fill page 3 exactly - (96..127); sequence 1 starts on page 0's last slot and crosses on its - very first token (31..63). The op must still touch only the context - rows and index per-seq tensors over [0, num_contexts).""" - torch.manual_seed(5) - env = _MlaCtxEnv(tokens_per_block=PAGE32) - try: - cached = [96, 31, 200, 77] - new = [32, 33, 1, 1] - rids = [0, 1, 2, 3] - env.kv_cache_manager.add_dummy_requests( - rids, token_nums=[c + n for c, n in zip(cached, new)] - ) - _run_and_check( - env, - request_ids=rids, - seq_lens=new, - num_contexts=2, - cached_lens=cached, - ) - finally: - env.shutdown() - - -def test_fp16_page32_mixed_cached_lengths() -> None: - """fp16 activations and fp16 latent cache at page size 32: a prefix - ending mid-page (33 + 31 closes page 1 exactly), a 16-page prefix - followed by three full pages of new tokens (512 + 96), and a - single-page tail.""" - torch.manual_seed(6) - env = _MlaCtxEnv(cache_dtype=DataType.HALF, tokens_per_block=PAGE32) - try: - cached = [33, 512, 1] - new = [31, 96, 6] - rids = [0, 1, 2] - env.kv_cache_manager.add_dummy_requests( - rids, token_nums=[c + n for c, n in zip(cached, new)] - ) - _run_and_check( - env, - request_ids=rids, - seq_lens=new, - num_contexts=3, - cached_lens=cached, - ) - finally: - env.shutdown() - - def _fp8_env(max_batch_size: int = 8, orig_quant: Optional[float] = None) -> _MlaCtxEnv: """DeepSeek-R1-0528 cell over an fp8-e4m3 latent pool: H = 128, page 32, bf16 activations, single-layer pool.""" diff --git a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_mla_rope_generation.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_mla_rope_generation.py index b8ed67b35184..b526e4c54323 100644 --- a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_mla_rope_generation.py +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_mla_rope_generation.py @@ -300,123 +300,6 @@ def _positions( ] -def _run_and_check( - env: _MlaEnv, - request_ids: List[int], - seq_lens: List[int], - num_contexts: int, - cached_lens: List[int], - q_pe_contiguous: bool, - predicted_tokens_per_seq: int = 1, -) -> None: - """One generation step over a bf16 latent pool: allocate P slots per - generation sequence, run the op over the generation tokens, and verify - every kernel effect.""" - gen_ids = request_ids[num_contexts:] - num_gen = len(gen_ids) - p = predicted_tokens_per_seq - assert seq_lens[num_contexts:] == [p] * num_gen - num_heads = env.num_heads - for rid in gen_ids: - for _ in range(p): - env.kv_cache_manager.impl.add_token(rid) - metadata = env.prepare_metadata(request_ids, seq_lens, num_contexts, cached_lens) - kv_lens = [c + s for c, s in zip(cached_lens, seq_lens)] - positions = _positions(kv_lens, num_contexts, num_gen, p) - rows = num_gen * p - - fused_q = torch.randn(rows, num_heads, GEN_HEAD_SIZE, dtype=torch.bfloat16, device="cuda") - q_pe = _make_q_pe(rows, num_heads, q_pe_contiguous) - latent_cache = torch.randn(rows, GEN_HEAD_SIZE, dtype=torch.bfloat16, device="cuda") - fused_q_orig = fused_q.clone() - q_pe_orig = q_pe.clone() - cu_q_seqlens = torch.full((num_gen + 1,), -1, dtype=torch.int32, device="cuda") - cu_kv_seqlens = torch.full((num_gen + 1,), -1, dtype=torch.int32, device="cuda") - fmha_scheduler_counter = torch.full((1,), 7, dtype=torch.uint32, device="cuda") - - mla_rope_generation( - fused_q, - q_pe, - latent_cache, - env.rotary_cos_sin, - cu_q_seqlens, - cu_kv_seqlens, - fmha_scheduler_counter, - None, # mla_bmm1_scale: fp8-KV-cache path only - None, # mla_bmm2_scale - None, # quant_q_buffer - metadata.kv_lens_cuda_runtime, - metadata.kv_lens_runtime, - metadata.prompt_lens_cpu_runtime, - num_contexts, - metadata.kv_cache_block_offsets, - env.kv_cache_manager.kv_cache_pool_pointers, - env.kv_cache_manager.kv_cache_pool_mapping, - None, # kv_scale_orig_quant - None, # kv_scale_quant_orig - None, # kv_cache_scale_orig_quant - None, # out_scale - None, # block_ids_per_seq - [None, None], # helix_tensor_params - p, # predicted_tokens_per_seq - 0, # layer_idx - num_heads, - 1, # num_kv_heads - GEN_HEAD_SIZE, - 0, # residual_dim - env.tokens_per_block, - MAX_SEQ_LEN, # attention_window_size - 1, # beam_width - 0, # quant_mode: bf16 KV cache - 1.0, # q_scaling - 0, # q_lora_rank - KV_LORA_RANK, - QK_NOPE_HEAD_DIM, - QK_ROPE_HEAD_DIM, - V_HEAD_DIM, - True, # rope_append - ) - torch.cuda.synchronize() - - # 1. fused_q tail = rope(q_pe) at that row's own position; the absorbed-q - # head slice is untouched. - for r in range(rows): - ref = env.rope_ref(q_pe_orig[r], positions[r]) - torch.testing.assert_close(fused_q[r, :, KV_LORA_RANK:], ref) - torch.testing.assert_close( - fused_q[..., :KV_LORA_RANK], - fused_q_orig[..., :KV_LORA_RANK], - rtol=0.0, - atol=0.0, # caller-owned region: must be bitwise untouched - ) - # q_pe is an input only (mutable in the schema, not mutated in practice). - torch.testing.assert_close(q_pe, q_pe_orig, rtol=0.0, atol=0.0) - - # 2. Cache append: [compressed_kv | rope(k_pe)] at each row's own slot. - for r in range(rows): - rid = gen_ids[r // p] - row = env.cache_row(rid, positions[r]) - torch.testing.assert_close( - row[:KV_LORA_RANK], - latent_cache[r, :KV_LORA_RANK], - rtol=0.0, - atol=0.0, # dtype-preserving copy: must be bitwise equal - ) - ref_k = env.rope_ref(latent_cache[r, KV_LORA_RANK:], positions[r]) - torch.testing.assert_close(row[KV_LORA_RANK:], ref_k) - - # 3. Scheduler buffers over generation sequences only. - _assert_scheduler_buffers( - cu_q_seqlens, - cu_kv_seqlens, - fmha_scheduler_counter, - num_gen, - num_heads, - kv_lens[num_contexts:], - p, - ) - - def _run_and_check_fp8( env: _MlaEnv, request_ids: List[int], @@ -653,135 +536,6 @@ def _assert_scheduler_buffers( assert fmha_scheduler_counter.item() == 0 -def test_bf16_decode_batch_strided_q_pe() -> None: - """Pure-decode batch; one sequence's new slot crosses a block boundary - (cached 64 = one full 64-token block); q_pe is a packed-q strided view.""" - torch.manual_seed(0) - env = _MlaEnv() - try: - env.kv_cache_manager.add_dummy_requests([0, 1], token_nums=[64, 32]) - _run_and_check( - env, - request_ids=[0, 1], - seq_lens=[1, 1], - num_contexts=0, - cached_lens=[64, 32], - q_pe_contiguous=False, - ) - finally: - env.shutdown() - - -def test_bf16_mixed_batch_skips_context() -> None: - """Context sequence leads the batch; the op consumes only the generation - tokens and indexes length/block tensors starting at num_contexts.""" - torch.manual_seed(1) - env = _MlaEnv() - try: - env.kv_cache_manager.add_dummy_requests([0, 1, 2], token_nums=[40, 100, 7]) - _run_and_check( - env, - request_ids=[0, 1, 2], - seq_lens=[40, 1, 1], - num_contexts=1, - cached_lens=[0, 100, 7], - q_pe_contiguous=True, - ) - finally: - env.shutdown() - - -def test_bf16_large_decode_batch_multi_step() -> None: - """64-sequence decode batch (many tokens per call), two consecutive - steps so the second call appends after the first call's tokens.""" - torch.manual_seed(2) - env = _MlaEnv(max_batch_size=64) - try: - rids = list(range(64)) - cached = [(37 * (i + 1)) % 800 + 1 for i in range(64)] - env.kv_cache_manager.add_dummy_requests(rids, token_nums=cached) - for step in range(2): - _run_and_check( - env, - request_ids=rids, - seq_lens=[1] * 64, - num_contexts=0, - cached_lens=[c + step for c in cached], - q_pe_contiguous=True, - ) - finally: - env.shutdown() - - -def test_bf16_page32_decode_batch_alignments() -> None: - """Page size 32 (the engine default): a pure-decode batch whose three - new slots land at every alignment a 32-token page has — position 32 - (page 1 slot 0, a fresh page after one full page), 31 (page 0's last - slot) and 64 (page 2 slot 0, after two full pages). q_pe is a packed-q - strided view.""" - torch.manual_seed(3) - env = _MlaEnv(tokens_per_block=PAGE32) - try: - env.kv_cache_manager.add_dummy_requests([0, 1, 2], token_nums=[32, 31, 64]) - _run_and_check( - env, - request_ids=[0, 1, 2], - seq_lens=[1, 1, 1], - num_contexts=0, - cached_lens=[32, 31, 64], - q_pe_contiguous=False, - ) - finally: - env.shutdown() - - -def test_bf16_page32_mixed_batch_skips_context() -> None: - """Page size 32, context sequence leading the batch: the op consumes - only the generation tokens and indexes length/block tensors from - num_contexts. The generation slots sit at position 96 (page 3 slot 0) - and 7 (mid first page).""" - torch.manual_seed(4) - env = _MlaEnv(tokens_per_block=PAGE32) - try: - env.kv_cache_manager.add_dummy_requests([0, 1, 2], token_nums=[40, 96, 7]) - _run_and_check( - env, - request_ids=[0, 1, 2], - seq_lens=[40, 1, 1], - num_contexts=1, - cached_lens=[0, 96, 7], - q_pe_contiguous=True, - ) - finally: - env.shutdown() - - -def test_bf16_page32_large_decode_batch_multi_step() -> None: - """Page size 32, 64-sequence decode batch over two consecutive steps. - The cached lengths spread the new slots over 25 pages; two sequences - (i = 5, 37) sit on a page's last slot at step 0 and cross into the next - page at step 1, and two (i = 18, 50) start a fresh page at step 0.""" - torch.manual_seed(5) - env = _MlaEnv(max_batch_size=64, tokens_per_block=PAGE32) - try: - rids = list(range(64)) - cached = [(37 * (i + 1)) % 800 + 1 for i in range(64)] - assert [i for i in range(64) if cached[i] % PAGE32 == PAGE32 - 1] == [5, 37] - assert [i for i in range(64) if cached[i] % PAGE32 == 0] == [18, 50] - env.kv_cache_manager.add_dummy_requests(rids, token_nums=cached) - for step in range(2): - _run_and_check( - env, - request_ids=rids, - seq_lens=[1] * 64, - num_contexts=0, - cached_lens=[c + step for c in cached], - q_pe_contiguous=True, - ) - finally: - env.shutdown() - - def _fp8_env( max_batch_size: int = 8, orig_quant: Optional[float] = None, @@ -1418,43 +1172,3 @@ def test_fp8_kv_mtp_append_race_control() -> None: step.assert_rows_at_sequence_length_slots() finally: env.shutdown() - - -def test_bf16_mtp_decode_batch() -> None: - """bf16 latent pool at P = 2 and 3 over the engine-default page size: the - roped q lands in fused_q's tail at each row's own position, and each - sequence's P cache rows land at consecutive slots — one sequence per case - inside a page, one straddling the boundary. The last case leads with a - context sequence.""" - for p, q_pe_contiguous in ((2, True), (3, False)): - torch.manual_seed(40 + p) - env = _MlaEnv(tokens_per_block=PAGE32) - try: - cached = [PAGE32 - p, PAGE32 - 1, 2 * PAGE32 - 1] - env.kv_cache_manager.add_dummy_requests([0, 1, 2], token_nums=cached) - _run_and_check( - env, - request_ids=[0, 1, 2], - seq_lens=[p] * 3, - num_contexts=0, - cached_lens=cached, - q_pe_contiguous=q_pe_contiguous, - predicted_tokens_per_seq=p, - ) - finally: - env.shutdown() - torch.manual_seed(44) - env = _MlaEnv(tokens_per_block=PAGE32) - try: - env.kv_cache_manager.add_dummy_requests([0, 1, 2], token_nums=[40, 31, 63]) - _run_and_check( - env, - request_ids=[0, 1, 2], - seq_lens=[40, 2, 2], - num_contexts=1, - cached_lens=[0, 31, 63], - q_pe_contiguous=True, - predicted_tokens_per_seq=2, - ) - finally: - env.shutdown() diff --git a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_thop_attention.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_thop_attention.py index 1eb1740f37a9..bd3c814d44a0 100644 --- a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_thop_attention.py +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_thop_attention.py @@ -636,50 +636,6 @@ def _run_and_check( torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) -def test_bf16_context_prefill_gqa_d128() -> None: - """Prefill-like: pure-context batch, mixed lengths crossing a block - boundary, GQA 8q/2kv, head_dim 128; cache append checked bit-exactly.""" - torch.manual_seed(0) - env = _PagedAttnEnv(num_heads=8, num_kv_heads=2, head_dim=128) - env.add_request(0, 48) - env.add_request(1, 17) - _run_and_check(env, [48, 17], 2, [0, 1]) - env.check_cache(0) - env.check_cache(1) - - -def test_bf16_decode_reads_cached_kv_gqa_d128() -> None: - """Decode-like: prefill writes the cache, then two gen-only steps read it. - - ctx len 64 exactly fills two 32-token blocks, so the first decode token - lands in a freshly allocated block (block-boundary crossing). - """ - torch.manual_seed(1) - env = _PagedAttnEnv(num_heads=8, num_kv_heads=2, head_dim=128) - env.add_request(0, 64) - env.add_request(1, 17) - _run_and_check(env, [64, 17], 2, [0, 1]) - for _ in range(2): # two decode steps: one token per sequence each - _run_and_check(env, [1, 1], 0, [0, 1]) - env.check_cache(0) - env.check_cache(1) - - -def test_bf16_mixed_batch_gqa_d128() -> None: - """One batch mixing a context-phase sequence and a generation-phase one. - - Context sequences must precede generation sequences in the batch. - """ - torch.manual_seed(2) - env = _PagedAttnEnv(num_heads=8, num_kv_heads=2, head_dim=128) - env.add_request(0, 40) - _run_and_check(env, [40], 1, [0]) - env.add_request(1, 23) - _run_and_check(env, [23, 1], 1, [1, 0]) - env.check_cache(0) - env.check_cache(1) - - def test_bf16_padding_mask_context() -> None: """mask_type=padding: bidirectional attention over a context batch.""" torch.manual_seed(3) @@ -724,26 +680,6 @@ def test_bf16_shipped_target_geometries() -> None: env.check_unwritten_pool_zero() -def test_bf16_head_count_axis() -> None: - """Head counts are a free axis, not an enumerated list: any (Hq, Hkv) - with Hq % Hkv == 0 works, and head_size selects the FMHA kernel. - - Covered here: MQA (16q/1kv), a non-power-of-2 GQA ratio (28q/4kv, - ratio 7), non-power-of-2 head counts (12q/3kv), and head_size 256 — - none of them a power-of-2 grouping the earlier cases already pinned. - """ - geometries = [(16, 1, 128), (28, 4, 128), (12, 3, 128), (8, 2, 256)] - for i, (num_heads, num_kv_heads, head_dim) in enumerate(geometries): - torch.manual_seed(40 + i) - env = _PagedAttnEnv(num_heads=num_heads, num_kv_heads=num_kv_heads, head_dim=head_dim) - env.add_request(0, 40) - env.add_request(1, 7) - _run_and_check(env, [40, 7], 2, [0, 1]) - _run_and_check(env, [1, 1], 0, [0, 1]) - env.check_cache(0) - env.check_cache(1) - - def test_rejects_non_divisible_head_counts() -> None: """Hq must be an integer multiple of Hkv, and the wrapper must be the one to say so: a context-only call at 6q/4kv d128 returns without @@ -1134,64 +1070,6 @@ def random_qkv(self, num_tokens: int) -> torch.Tensor: return torch.randn(num_tokens, width, dtype=torch.bfloat16, device="cuda") -def test_bf16_multilayer_shared_pool_gqa_d128() -> None: - """One paged pool shared by 4 layers, addressed as a real multi-layer - KVCacheManager lays it out: pages interleave every layer's K/V slabs, - the block-offset table is layer-agnostic, and each call selects its - layer via local_layer_idx -> pool-mapping row. Prefill plus two decode - steps per layer (decode crosses a page boundary), GQA 8q/2kv d128: - per-layer outputs vs fp32 references, per-layer appends bit-exact, - sibling layers bitwise untouched by each call. A final doctored-mapping - call pins the pool-base shift to the mapping row's layer-in-pool column - (local_layer_idx only selects the row).""" - torch.manual_seed(10) - num_layers = 4 - env = _MultiLayerPagedAttnEnv(num_layers=num_layers, num_heads=8, num_kv_heads=2, head_dim=128) - # The real manager maps layer l to (pool 0, layer-in-pool l): one - # mapping row per layer, rows beyond 0 shifting the pool base. - assert env.pool_mapping.tolist() == [[0, layer] for layer in range(num_layers)] - - env.add_request(0, 64) # exactly two pages: decode crosses into a third - env.add_request(1, 17) - env.refresh_offsets([0, 1], num_contexts=2) - for layer in range(num_layers): - qkv = env.random_qkv(81) - siblings_before = [env.layer_views[m].clone() for m in range(num_layers) if m != layer] - out = env.call_op(layer, qkv, [64, 17], 2, [0, 1]) - ref = env.reference(layer, qkv, [64, 17], [0, 1], [0, 0]) - torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) - siblings_after = [env.layer_views[m] for m in range(num_layers) if m != layer] - for before, after in zip(siblings_before, siblings_after): - assert torch.equal(before, after) # no cross-layer write - env.check_caches([0, 1]) - - for _ in range(2): # decode: one token per sequence per layer per step - env.add_decode_token(0) - env.add_decode_token(1) - env.refresh_offsets([0, 1], num_contexts=0) - for layer in range(num_layers): - cached = [env.cached_len(layer, 0), env.cached_len(layer, 1)] - qkv = env.random_qkv(2) - out = env.call_op(layer, qkv, [1, 1], 0, [0, 1]) - ref = env.reference(layer, qkv, [1, 1], [0, 1], cached) - torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) - env.check_caches([0, 1]) - - # Doctored mapping: local_layer_idx=1 whose row claims layer-in-pool 3. - # The append must land in layer 3's slabs — the shift is driven by the - # mapping row's layer column, not by local_layer_idx itself. (Appends - # one token past each sequence's history, inside already-allocated - # pages; run after all correctness checks since it plants garbage.) - doctored = env.pool_mapping.clone() - doctored[1, 1] = 3 - views_before = [env.layer_views[m].clone() for m in range(num_layers)] - env.call_op(1, env.random_qkv(2), [1, 1], 0, [0, 1], pool_mapping=doctored, record=False) - assert torch.equal(env.layer_views[1], views_before[1]) # row idx not the shift - assert not torch.equal(env.layer_views[3], views_before[3]) - assert torch.equal(env.layer_views[0], views_before[0]) - assert torch.equal(env.layer_views[2], views_before[2]) - - def _compose_two_pools(env_a, env_b): """Pool pointers, layer->pool mapping and block offsets for two pools. @@ -1217,7 +1095,7 @@ def _compose_two_pools(env_a, env_b): return pool_pointers, pool_mapping, block_offsets -def test_bf16_two_pool_layer_routing_gqa_d128() -> None: +def test_bf16_two_pool_layer_routing_gqa_d64() -> None: """Two paged pools, layers routed between them by the mapping's pool column. Everything above this point certifies pool id 0 only: one manager owns one @@ -1234,7 +1112,7 @@ def test_bf16_two_pool_layer_routing_gqa_d128() -> None: shaped pages; only the sibling-pool comparison sees it. """ torch.manual_seed(11) - kw = dict(num_heads=8, num_kv_heads=2, head_dim=128) + kw = dict(num_heads=64, num_kv_heads=8, head_dim=64) env_a = _MultiLayerPagedAttnEnv(num_layers=2, **kw) env_b = _MultiLayerPagedAttnEnv(num_layers=2, **kw) envs = [(env_a, 0), (env_a, 1), (env_b, 0), (env_b, 1)] # (env, layer in that env) @@ -1302,228 +1180,6 @@ def other(env): # ─── Standard configuration, fp8-e4m3 paged KV pool (quant_mode 128) ─── -def _fp8_env(kv_scaling_factor: float = 1.0) -> _PagedAttnEnv: - """GQA 32q/8kv d128 tpb 32 (qwen3-8b-like geometry), fp8-e4m3 pool.""" - return _PagedAttnEnv( - num_heads=32, - num_kv_heads=8, - head_dim=128, - pool_dtype=torch.float8_e4m3fn, - quant_mode=QUANT_MODE_FP8_KV_CACHE, - kv_scaling_factor=kv_scaling_factor, - ) - - -def _e4m3_roundtrip(env: "_PagedAttnEnv | _MultiLayerPagedAttnEnv"): - """bf16 -> e4m3 (scaled by orig_quant) -> bf16 (scaled by quant_orig): - what a value fed into the fp8 pool looks like when read back out.""" - orig_quant, quant_orig = env.kv_scale_orig_quant, env.kv_scale_quant_orig - assert orig_quant is not None and quant_orig is not None - - def transform(t: torch.Tensor) -> torch.Tensor: - q = (t.float() * orig_quant).to(torch.float8_e4m3fn) - return (q.float() * quant_orig).to(t.dtype) - - return transform - - -def test_fp8_kv_context_prefill_gqa32_8_d128() -> None: - """fp8 pool, context prefill: the context FMHA computes over the bf16 - packed QKV — accuracy identical to the bf16-pool surface (an fp8-KV - reference is ~1.4e-1 off; the pool plays no part in context math) — - while the in-op append writes e4m3(K * orig_quant) into the pool, - bit-exactly, touching nothing else.""" - torch.manual_seed(20) - env = _fp8_env() - env.add_request(0, 48) # crosses a 32-token page boundary - env.add_request(1, 17) - qkv = env.random_qkv(65) - out = env.call_op(qkv, [48, 17], 2, [0, 1], MASK_CAUSAL) - ref = env.reference(qkv, [48, 17], [0, 1], [0, 0], MASK_CAUSAL) - torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) - env.check_cache(0) - env.check_cache(1) - env.check_unwritten_pool_zero() - - -def test_fp8_kv_decode_reads_fp8_cache() -> None: - """Decode over the fp8 pool: prefill [64, 17] (64 fills exactly two - pages, so the first decode token opens a fresh page), then two decode - steps checked against the quantized-KV reference; appends stay bit-exact - through decode. A 200-token-history decode covers prefill-like KV - extents (observed err there is ~4x smaller than the short-history max).""" - torch.manual_seed(21) - env = _fp8_env() - env.add_request(0, 64) - env.add_request(1, 17) - env.call_op(env.random_qkv(81), [64, 17], 2, [0, 1], MASK_CAUSAL) - rt = _e4m3_roundtrip(env) - for _ in range(2): - cached = [env.cached_len(0), env.cached_len(1)] - qkv_d = env.random_qkv(2) - out = env.call_op(qkv_d, [1, 1], 0, [0, 1], MASK_CAUSAL) - ref = env.reference( - qkv_d, [1, 1], [0, 1], cached, MASK_CAUSAL, q_transform=rt, kv_transform=rt - ) - torch.testing.assert_close(out, ref, rtol=FP8_DECODE_RTOL, atol=FP8_DECODE_ATOL) - env.check_cache(0) - env.check_cache(1) - env.check_unwritten_pool_zero() - - torch.manual_seed(22) - env = _fp8_env() - env.add_request(0, 200) - env.call_op(env.random_qkv(200), [200], 1, [0], MASK_CAUSAL) - rt = _e4m3_roundtrip(env) - qkv_d = env.random_qkv(1) - out = env.call_op(qkv_d, [1], 0, [0], MASK_CAUSAL) - ref = env.reference(qkv_d, [1], [0], [200], MASK_CAUSAL, q_transform=rt, kv_transform=rt) - torch.testing.assert_close(out, ref, rtol=FP8_DECODE_RTOL, atol=FP8_DECODE_ATOL) - - -def test_fp8_kv_mixed_batch() -> None: - """One call mixing a context and a generation sequence over the fp8 - pool: context rows match the bf16-KV reference (bf16 context FMHA), the - generation row matches the quantized-KV reference (fp8 decode kernel).""" - torch.manual_seed(23) - env = _fp8_env() - env.add_request(0, 40) - env.call_op(env.random_qkv(40), [40], 1, [0], MASK_CAUSAL) - env.add_request(1, 23) - cached_gen = env.cached_len(0) - qkv = env.random_qkv(24) - out = env.call_op(qkv, [23, 1], 1, [1, 0], MASK_CAUSAL) - ref_ctx = env.reference(qkv[:23], [23], [1], [0], MASK_CAUSAL) - torch.testing.assert_close(out[:23], ref_ctx, rtol=RTOL, atol=ATOL) - rt = _e4m3_roundtrip(env) - ref_gen = env.reference( - qkv[23:], [1], [0], [cached_gen], MASK_CAUSAL, q_transform=rt, kv_transform=rt - ) - torch.testing.assert_close(out[23:], ref_gen, rtol=FP8_DECODE_RTOL, atol=FP8_DECODE_ATOL) - env.check_cache(0) - env.check_cache(1) - - -def test_fp8_kv_scale_semantics() -> None: - """Non-1.0 kv scales. s=2.0 (a power of two: scaling is an exact - exponent shift): the append stays bit-exact vs the e4m3(K * orig_quant) - mirror — orig_quant is consumed on write — and decode matches the - scale-aware reference — quant_orig is consumed on read (ignoring it - shows as ~4.4e-1). s=1.5 (not a power of two, so the e4m3 round-trip - genuinely depends on the scale value): the kernel's quantization - arithmetic differs from fp32-multiply-then-round-to-nearest on a small - fraction of elements (observed 0.6%), each within one e4m3 ulp; decode - matches within the same fp8 tolerance.""" - torch.manual_seed(24) - env = _fp8_env(kv_scaling_factor=2.0) - env.add_request(0, 40) - qkv = env.random_qkv(40) - out = env.call_op(qkv, [40], 1, [0], MASK_CAUSAL) - ref = env.reference(qkv, [40], [0], [0], MASK_CAUSAL) - torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) # bf16 context - env.check_cache(0) - rt = _e4m3_roundtrip(env) - qkv_d = env.random_qkv(1) - out = env.call_op(qkv_d, [1], 0, [0], MASK_CAUSAL) - ref = env.reference(qkv_d, [1], [0], [40], MASK_CAUSAL, q_transform=rt, kv_transform=rt) - torch.testing.assert_close(out, ref, rtol=FP8_DECODE_RTOL, atol=FP8_DECODE_ATOL) - - torch.manual_seed(25) - env = _fp8_env(kv_scaling_factor=1.5) - env.add_request(0, 64) - env.call_op(env.random_qkv(64), [64], 1, [0], MASK_CAUSAL) - got_k, got_v, exp_k, exp_v = env.cache_pages_content(0) - for got, exp in ((got_k, exp_k), (got_v, exp_v)): - exact = (got.view(torch.uint8) == exp.view(torch.uint8)).float().mean().item() - assert exact >= 0.99, f"append bit-exact fraction only {exact:.4f}" - diff = (got.float() - exp.float()).abs() - # One e4m3 ulp at the element's magnitude: 2^(floor(log2 |x|) - 3) - # for normals, 2^-9 below the min normal 2^-6. - mag = torch.maximum(got.float().abs(), exp.float().abs()).clamp(min=2**-6) - ulp = 2.0 ** (torch.floor(torch.log2(mag)) - 3) - assert bool((diff <= ulp).all()), "append off by more than one e4m3 ulp" - rt = _e4m3_roundtrip(env) - qkv_d = env.random_qkv(1) - out = env.call_op(qkv_d, [1], 0, [0], MASK_CAUSAL) - ref = env.reference(qkv_d, [1], [0], [64], MASK_CAUSAL, q_transform=rt, kv_transform=rt) - torch.testing.assert_close(out, ref, rtol=FP8_DECODE_RTOL, atol=FP8_DECODE_ATOL) - - -def test_fp8_kv_multilayer_shared_pool_gqa32_8_d128() -> None: - """One fp8-e4m3 paged pool shared by 4 layers (quant_mode 128, s=1.0), - manager state consumed as-is — the production serving shape of the fp8 - cache, at the production GQA 32q/8kv d128 tpb 32 geometry. The op sizes - slabs from quant_mode alone, so the layer-base shift must be computed in - e4m3 slab units for the append to land where the fp8 manager laid the - layer out. Prefill plus two decode steps per layer (decode crosses a - page boundary): per-layer context outputs vs bf16 references (context - FMHA reads the packed bf16 q rows, not the pool), per-layer decode - outputs vs quantized-KV references over that layer's own history, - per-layer e4m3 appends bit-exact, and sibling layers bitwise untouched - by every call — context and generation append paths both.""" - torch.manual_seed(26) - num_layers = 4 - env = _MultiLayerPagedAttnEnv( - num_layers=num_layers, - num_heads=32, - num_kv_heads=8, - head_dim=128, - dtype=DataType.FP8, - quant_mode=QUANT_MODE_FP8_KV_CACHE, - kv_scaling_factor=1.0, - ) - # A DataType.FP8 manager allocates a real e4m3 pool with identity - # mapping rows, exactly like the bf16 one. - assert env.layer_views[0].dtype == torch.float8_e4m3fn - assert env.pool_mapping.tolist() == [[0, layer] for layer in range(num_layers)] - - def check_sibling_isolation(layer: int, fn) -> None: - before = [env.layer_views[m].clone() for m in range(num_layers) if m != layer] - fn() - after = [env.layer_views[m] for m in range(num_layers) if m != layer] - for b, a in zip(before, after): - assert _bitwise_equal(b, a), f"call on layer {layer} wrote a sibling" - - env.add_request(0, 64) # exactly two pages: decode crosses into a third - env.add_request(1, 17) - env.refresh_offsets([0, 1], num_contexts=2) - for layer in range(num_layers): - qkv = env.random_qkv(81) - - def ctx_call(layer=layer, qkv=qkv): - out = env.call_op(layer, qkv, [64, 17], 2, [0, 1]) - ref = env.reference(layer, qkv, [64, 17], [0, 1], [0, 0]) - torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) - - check_sibling_isolation(layer, ctx_call) - env.check_caches([0, 1]) - - rt = _e4m3_roundtrip(env) - for _ in range(2): # decode: one token per sequence per layer per step - env.add_decode_token(0) - env.add_decode_token(1) - env.refresh_offsets([0, 1], num_contexts=0) - for layer in range(num_layers): - cached = [env.cached_len(layer, 0), env.cached_len(layer, 1)] - qkv = env.random_qkv(2) - - def gen_call(layer=layer, qkv=qkv, cached=cached): - out = env.call_op(layer, qkv, [1, 1], 0, [0, 1]) - ref = env.reference( - layer, - qkv, - [1, 1], - [0, 1], - cached, - q_transform=rt, - kv_transform=rt, - ) - torch.testing.assert_close(out, ref, rtol=FP8_DECODE_RTOL, atol=FP8_DECODE_ATOL) - - check_sibling_isolation(layer, gen_call) - env.check_caches([0, 1]) - - # ─── Attention sinks (standard configuration, bf16 pool) ─────────────── # gpt-oss-120b tp1 attention geometry: 64 q heads, 8 kv heads, head_size 64. SINK_HQ, SINK_HKV, SINK_D = 64, 8, 64 @@ -3835,283 +3491,6 @@ def _random_context_inputs( return q, k, v, latent -def _mla_context_prefill_case( - num_heads: int, - seed: int, - tokens_per_block: int, - seq_lens: List[int], - q_scaling: float = 1.0, - rope: Optional[RopeParams] = None, - rope_scalars: Optional[_RopeScalars] = None, -) -> None: - """MLA context_only prefill: two fresh sequences, at least one crossing - a page boundary. Verifies the FMHA output against a rope-aware fp32 - reference, the latent-cache append page by page (check_cache), that - latent_cache is read-only, and that only the q/k rope slices are - clobbered.""" - torch.manual_seed(seed) - env = _MlaPagedEnv( - num_heads=num_heads, - max_blocks_per_seq=MLA_MAX_SEQ_LEN // tokens_per_block, - tokens_per_block=tokens_per_block, - q_scaling=q_scaling, - rope=rope, - rope_scalars=rope_scalars, - ) - num_tokens = sum(seq_lens) - for rid, ln in enumerate(seq_lens): - env.add_request(rid, ln) - q, k, v, latent = _random_context_inputs(num_tokens, num_heads) - q_orig, k_orig, latent_orig = q.clone(), k.clone(), latent.clone() - - rids = list(range(len(seq_lens))) - out = env.call_context(rids, seq_lens, q, k, v, latent) - ref = env.context_reference(seq_lens, q_orig, k_orig, v, latent_orig) - torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) - - # The append this case checks must really walk more than one page. - assert max(len(env.pages[rid]) for rid in rids) > 1 - for rid in rids: - env.check_cache(rid) - # latent_cache is an input only. - assert torch.equal(latent, latent_orig) - # The op ropes q_pe/k_pe in place: nope slices intact, rope slices clobbered. - q3 = q.view(num_tokens, num_heads, QK_HEAD_DIM) - k3 = k.view(num_tokens, num_heads, QK_HEAD_DIM) - q3_orig = q_orig.view(num_tokens, num_heads, QK_HEAD_DIM) - k3_orig = k_orig.view(num_tokens, num_heads, QK_HEAD_DIM) - assert torch.equal(q3[..., :QK_NOPE_HEAD_DIM], q3_orig[..., :QK_NOPE_HEAD_DIM]) - assert torch.equal(k3[..., :QK_NOPE_HEAD_DIM], k3_orig[..., :QK_NOPE_HEAD_DIM]) - assert not torch.equal(q3[..., QK_NOPE_HEAD_DIM:], q3_orig[..., QK_NOPE_HEAD_DIM:]) - assert not torch.equal(k3[..., QK_NOPE_HEAD_DIM:], k3_orig[..., QK_NOPE_HEAD_DIM:]) - - -def test_bf16_mla_context_prefill() -> None: - """Fresh-prefill MLA context at 16 query heads, page size 64 (100 > 64 - crosses a boundary).""" - _mla_context_prefill_case(MLA_NUM_HEADS, 5, MLA_TOKENS_PER_BLOCK, [100, 17]) - - -def test_bf16_mla_context_prefill_h32() -> None: - """The same fresh-prefill case at 32 query heads (32/32/192).""" - _mla_context_prefill_case(MLA_NUM_HEADS_H32, 105, MLA_TOKENS_PER_BLOCK, [100, 17]) - - -def test_bf16_mla_context_prefill_h8() -> None: - """The same fresh-prefill case at 8 query heads (8/8/192) — the tep4 - slice of a 32-head checkpoint.""" - _mla_context_prefill_case(MLA_NUM_HEADS_H8, 805, MLA_TOKENS_PER_BLOCK, [100, 17]) - - -def test_bf16_mla_page32_context_prefill_h32() -> None: - """Fresh-prefill MLA context at page size 32 (the engine default), 32 - query heads. 96 fills three 32-token pages exactly — a 64-token page - never ends there — and 33 crosses into a second page by one token, so - the append's page/slot arithmetic is exercised at both a page-aligned - end and a one-token spill.""" - _mla_context_prefill_case(MLA_NUM_HEADS_H32, 205, MLA_PAGE32, [96, 33]) - - -def test_bf16_mla_page32_context_prefill_h8() -> None: - """The same page-32 fresh-prefill case at 8 query heads.""" - _mla_context_prefill_case(MLA_NUM_HEADS_H8, 815, MLA_PAGE32, [96, 33]) - - -def test_bf16_mla_page32_context_prefill_h128() -> None: - """The same page-32 fresh-prefill case at 128 query heads (128/128/192) — - the DeepSeek-R1-0528 layer shape, which attention DP replicates whole onto - every rank rather than slicing. Baseline rope/scale so this case isolates - the head count; the R1 cell adds the other two axes further down.""" - _mla_context_prefill_case(MLA_NUM_HEADS_H128, 905, MLA_PAGE32, [96, 33]) - - -def _mla_generation_decode_case( - num_heads: int, - seed: int, - tokens_per_block: int, - prefill_lens: List[int], - q_scaling: float = 1.0, - rope: Optional[RopeParams] = None, - rope_scalars: Optional[_RopeScalars] = None, -) -> None: - """MLA generation_only decode over cache written by the context call. - - The first prefill length is an exact multiple of the page size, so the - first decode token lands in a fresh page. Two decode steps; each step - the test appends the new latent row (the op does not append in - generation) and checks the FMHA output against an fp32 latent-MQA - reference. The garbage latent_cache/q_pe arguments plus cache/fused_q - invariance pin the fusion boundary: the generation call only reads the - paged pool. - """ - torch.manual_seed(seed) - env = _MlaPagedEnv( - num_heads=num_heads, - max_blocks_per_seq=MLA_MAX_SEQ_LEN // tokens_per_block, - tokens_per_block=tokens_per_block, - q_scaling=q_scaling, - rope=rope, - rope_scalars=rope_scalars, - ) - assert prefill_lens[0] % tokens_per_block == 0 - rids = list(range(len(prefill_lens))) - for rid, ln in zip(rids, prefill_lens): - env.add_request(rid, ln) - q, k, v, latent = _random_context_inputs(sum(prefill_lens), num_heads) - env.call_context(rids, prefill_lens, q, k, v, latent) - for rid in rids: - env.check_cache(rid) - pages_after_prefill = {rid: len(env.pages[rid]) for rid in rids} - - for _ in range(2): - for rid in rids: - env.append_decode_latent( - rid, torch.randn(LATENT_DIM, dtype=torch.bfloat16, device="cuda") - ) - fused_q = torch.randn( - len(rids), num_heads * LATENT_DIM, dtype=torch.bfloat16, device="cuda" - ) - fused_q_orig = fused_q.clone() - pool_before = env.pool.clone() - out = env.call_generation(rids, fused_q) - ref = env.generation_reference(rids, fused_q_orig) - torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) - assert torch.equal(fused_q, fused_q_orig) - assert torch.equal(env.pool, pool_before) # generation never writes - # The decode steps must really have opened a new page on some sequence, - # i.e. the read range crossed a page boundary the prefill had not. - assert any(len(env.pages[rid]) > pages_after_prefill[rid] for rid in rids) - - -def test_bf16_mla_generation_decode() -> None: - """Latent-MQA MLA decode at 16 query heads, page size 64.""" - _mla_generation_decode_case(MLA_NUM_HEADS, 6, MLA_TOKENS_PER_BLOCK, [64, 30]) - - -def test_bf16_mla_generation_decode_h32() -> None: - """The same decode case at 32 query heads (32/1/576).""" - _mla_generation_decode_case(MLA_NUM_HEADS_H32, 106, MLA_TOKENS_PER_BLOCK, [64, 30]) - - -def test_bf16_mla_generation_decode_h8() -> None: - """The same decode case at 8 query heads (8/1/576). Unlike 16 and 32, - which share the ...VarSeqQ16... decode kernel, 8 heads map to a - ...VarSeqQ8... one — a differently named compiled variant, so this is the - flavor where the head count actually changes the kernel.""" - _mla_generation_decode_case(MLA_NUM_HEADS_H8, 806, MLA_TOKENS_PER_BLOCK, [64, 30]) - - -def test_bf16_mla_page32_generation_decode_h32() -> None: - """Latent-MQA MLA decode at page size 32 (the engine default), 32 query - heads — the decode path is where the page size selects a different - compiled trtllm-gen kernel (...PagedKvDenseP32... rather than ...P64...). - Sequence 0's 64-token prefill fills two pages exactly, so its first - decode token opens page 2; sequence 1's 31-token prefill leaves its - first decode token on page 0's last slot and its second one opens page - 1, so a decode read range crosses a boundary mid-case.""" - _mla_generation_decode_case(MLA_NUM_HEADS_H32, 206, MLA_PAGE32, [64, 31]) - - -def test_bf16_mla_page32_generation_decode_h16() -> None: - """The same page-32 decode case at 16 query heads. The decode kernel is - JIT-compiled once per head count even though its name does not carry the - count, so the page-32 variant is exercised at both certified counts.""" - _mla_generation_decode_case(MLA_NUM_HEADS, 216, MLA_PAGE32, [64, 31]) - - -def test_bf16_mla_page32_generation_decode_h8() -> None: - """The same page-32 decode case at 8 query heads. Both axes that reach - the compiled decode kernel move here at once: the page size is in the - kernel name (...P32...) and 8 heads take the ...VarSeqQ8... q-tile.""" - _mla_generation_decode_case(MLA_NUM_HEADS_H8, 816, MLA_PAGE32, [64, 31]) - - -def test_bf16_mla_page32_generation_decode_h128() -> None: - """The same page-32 decode case at 128 query heads (128/1/576). Decode is - the phase where the head count reaches kernel selection: 128 reports the - same ...P32VarSeqQ16Kv128... name 16 and 32 do, yet pays its own compile - (the cache is keyed more finely than the name).""" - _mla_generation_decode_case(MLA_NUM_HEADS_H128, 906, MLA_PAGE32, [64, 31]) - - -def _mla_mixed_batch_case( - num_heads: int, - seed: int, - tokens_per_block: int, - first_len: int, - second_len: int, - q_scaling: float = 1.0, - rope: Optional[RopeParams] = None, - rope_scalars: Optional[_RopeScalars] = None, -) -> None: - """A mixed batch is two calls sharing full-batch metadata: context_only - over the leading context rows, generation_only over the trailing - generation rows (indexed from num_contexts).""" - torch.manual_seed(seed) - env = _MlaPagedEnv( - num_heads=num_heads, - max_blocks_per_seq=MLA_MAX_SEQ_LEN // tokens_per_block, - tokens_per_block=tokens_per_block, - q_scaling=q_scaling, - rope=rope, - rope_scalars=rope_scalars, - ) - # Prefill request 0 alone, then decode it alongside a new context request. - env.add_request(0, first_len) - q, k, v, latent = _random_context_inputs(first_len, num_heads) - env.call_context([0], [first_len], q, k, v, latent) - - env.add_request(1, second_len) - env.append_decode_latent(0, torch.randn(LATENT_DIM, dtype=torch.bfloat16, device="cuda")) - q, k, v, latent = _random_context_inputs(second_len, num_heads) - q_orig, k_orig, latent_orig = q.clone(), k.clone(), latent.clone() - out_ctx = env.call_context([1], [second_len], q, k, v, latent, gen_request_ids=[0]) - ref_ctx = env.context_reference([second_len], q_orig, k_orig, v, latent_orig) - torch.testing.assert_close(out_ctx, ref_ctx, rtol=RTOL, atol=ATOL) - - fused_q = torch.randn(1, num_heads * LATENT_DIM, dtype=torch.bfloat16, device="cuda") - out_gen = env.call_generation([0], fused_q, ctx_request_ids=[1], ctx_seq_lens=[second_len]) - ref_gen = env.generation_reference([0], fused_q) - torch.testing.assert_close(out_gen, ref_gen, rtol=RTOL, atol=ATOL) - env.check_cache(0) - env.check_cache(1) - - -def test_bf16_mla_mixed_batch_split_calls() -> None: - """Mixed-batch MLA (two phase calls) at 16 query heads, page size 64.""" - _mla_mixed_batch_case(MLA_NUM_HEADS, 7, MLA_TOKENS_PER_BLOCK, 40, 23) - - -def test_bf16_mla_mixed_batch_split_calls_h32() -> None: - """The same mixed-batch case at 32 query heads.""" - _mla_mixed_batch_case(MLA_NUM_HEADS_H32, 107, MLA_TOKENS_PER_BLOCK, 40, 23) - - -def test_bf16_mla_mixed_batch_split_calls_h8() -> None: - """The same mixed-batch case at 8 query heads: the two phase calls of one - batch run at 8/8/192 and 8/1/576 off the same full-batch state tensors.""" - _mla_mixed_batch_case(MLA_NUM_HEADS_H8, 807, MLA_TOKENS_PER_BLOCK, 40, 23) - - -def test_bf16_mla_page32_mixed_batch_split_calls_h32() -> None: - """Mixed-batch MLA at page size 32, 32 query heads: the decoding - sequence's 64-token history fills two pages exactly, so its appended - token opens page 2 and the decode call reads across three pages, while - the context sequence sharing the batch spans two.""" - _mla_mixed_batch_case(MLA_NUM_HEADS_H32, 207, MLA_PAGE32, 64, 33) - - -def test_bf16_mla_page32_mixed_batch_split_calls_h8() -> None: - """The same page-32 mixed-batch case at 8 query heads.""" - _mla_mixed_batch_case(MLA_NUM_HEADS_H8, 817, MLA_PAGE32, 64, 33) - - -def test_bf16_mla_page32_mixed_batch_split_calls_h128() -> None: - """The same page-32 mixed-batch case at 128 query heads: the two phase - calls run at 128/128/192 and 128/1/576 off one set of batch tensors.""" - _mla_mixed_batch_case(MLA_NUM_HEADS_H128, 907, MLA_PAGE32, 64, 33) - - # ─── MLA context with latent_cache=None (cached KV / chunked prefill) ─── @@ -4197,226 +3576,6 @@ def _explicit_kv_reference( return out.to(torch.bfloat16), stats -def _assert_valid_rows_close( - actual: torch.Tensor, ref: torch.Tensor, ref_stats: torch.Tensor -) -> None: - """Compare only rows of sequences that had KV in this pass (non-NaN in - the reference); zero-KV rows are undefined op output.""" - valid = ~torch.isnan(ref_stats[:, 0, 0]) - torch.testing.assert_close(actual[valid], ref[valid], rtol=RTOL, atol=ATOL) - - -def _assert_stats_close(actual: torch.Tensor, ref_stats: torch.Tensor) -> None: - valid = ~torch.isnan(ref_stats[:, 0, 0]) - # Both sides are fp32 reductions over the same bf16-rounded q/k, so they - # differ only by accumulation order: observed max abs err 2e-6 (max stat) - # and rel err 1.2e-6 (sum stat) on sm_100, on stat magnitudes ~1-40. - # 1e-4 gives ~50x margin while still catching wrong-domain (log2) or - # unscaled-logit stats outright. - torch.testing.assert_close(actual[valid], ref_stats[valid], rtol=1e-4, atol=1e-4) - - -def _mla_context_cached_kv_case( - num_heads: int, - seed: int, - tokens_per_block: int, - cached_lens: List[int], - new_lens: List[int], - q_scaling: float = 1.0, - rope: Optional[RopeParams] = None, - rope_scalars: Optional[_RopeScalars] = None, -) -> None: - """MLA context over a cached KV prefix: latent_cache=None, q pre-rotated - upstream, K/V supplied for the full [cached + new] range so KV length - exceeds q length. Causal masking is bottom-right aligned. Verifies the - output against an fp32 reference and that the call mutates nothing but - output: q, k, v, and the paged pool stay bitwise intact (no in-kernel - RoPE, no append).""" - torch.manual_seed(seed) - env = _MlaPagedEnv( - num_heads=num_heads, - max_blocks_per_seq=MLA_MAX_SEQ_LEN // tokens_per_block, - tokens_per_block=tokens_per_block, - q_scaling=q_scaling, - rope=rope, - rope_scalars=rope_scalars, - ) - kv_lens = [c + n for c, n in zip(cached_lens, new_lens)] - for rid, total in enumerate(kv_lens): - env.add_request(rid, new_lens[rid]) - env.reserve_cache_pages(rid, total) - env.pool.normal_() # op must not read or write the pool in this mode - pool_before = env.pool.clone() - - q = torch.randn(sum(new_lens), num_heads * QK_HEAD_DIM, dtype=torch.bfloat16, device="cuda") - k, packed_kv = _random_explicit_kv(sum(kv_lens), num_heads) - v = _v_split_view(packed_kv, num_heads) - q_orig, k_orig, v_orig = q.clone(), k.clone(), v.clone() - - out = env.call_context_no_append([0, 1, 2], new_lens, kv_lens, q, k, v) - ref, _ = _explicit_kv_reference( - q, k, v, new_lens, kv_lens, MASK_CAUSAL, num_heads, env.softmax_scale - ) - torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) - - # Fusion boundary: everything is an input; only output rows are written. - assert torch.equal(q, q_orig) - assert torch.equal(k, k_orig) - assert torch.equal(v, v_orig) - assert torch.equal(env.pool, pool_before) - - -def test_bf16_mla_context_cached_kv_no_append() -> None: - """Cached-KV (no-append) MLA context at 16 query heads, page size 64. - Prefixes: page-crossing, mid-page, and empty.""" - _mla_context_cached_kv_case(MLA_NUM_HEADS, 8, MLA_TOKENS_PER_BLOCK, [80, 33, 0], [40, 7, 25]) - - -def test_bf16_mla_context_cached_kv_no_append_h32() -> None: - """The same cached-KV context case at 32 query heads.""" - _mla_context_cached_kv_case( - MLA_NUM_HEADS_H32, 108, MLA_TOKENS_PER_BLOCK, [80, 33, 0], [40, 7, 25] - ) - - -def test_bf16_mla_context_cached_kv_no_append_h8() -> None: - """The same cached-KV context case at 8 query heads.""" - _mla_context_cached_kv_case( - MLA_NUM_HEADS_H8, 808, MLA_TOKENS_PER_BLOCK, [80, 33, 0], [40, 7, 25] - ) - - -def test_bf16_mla_page32_context_cached_kv_no_append_h32() -> None: - """Cached-KV (no-append) MLA context at page size 32, 32 query heads. - The pages are reserved as production does even though this flavor never - touches the pool, so the batch carries a page-32 offsets table: prefixes - of 96 (three exact pages), 31 (one short of a page) and 0, reaching KV - lengths of 128 (four exact pages), 40 and 25.""" - _mla_context_cached_kv_case(MLA_NUM_HEADS_H32, 208, MLA_PAGE32, [96, 31, 0], [32, 9, 25]) - - -def test_bf16_mla_page32_context_cached_kv_no_append_h8() -> None: - """The same page-32 cached-KV context case at 8 query heads.""" - _mla_context_cached_kv_case(MLA_NUM_HEADS_H8, 818, MLA_PAGE32, [96, 31, 0], [32, 9, 25]) - - -def test_bf16_mla_page32_context_cached_kv_no_append_h128() -> None: - """The same page-32 cached-KV context case at 128 query heads. This is the - flavor an engine with block reuse on runs for every context request that - hits a cached prefix, so it carries the head count as much as the fresh - one does.""" - _mla_context_cached_kv_case(MLA_NUM_HEADS_H128, 908, MLA_PAGE32, [96, 31, 0], [32, 9, 25]) - - -def test_bf16_mla_context_chunked_prefill_with_merge() -> None: - """MLA chunked context: one padding-masked partial pass per cached-KV - chunk with softmax_stats_tensor emitted, a final causal pass over the - new tokens, each folded into the running output by the downstream - trtllm merge op (the production consumer of the emitted stats — the - reference is still built from torch alone). Every pass's output and - stats are checked against fp32 partial-attention references, and the - fully merged output/stats against a single-pass full-range reference.""" - torch.manual_seed(9) - env = _MlaPagedEnv() - cached_lens = [96, 48, 0] - new_lens = [32, 17, 23] - kv_lens = [c + n for c, n in zip(cached_lens, new_lens)] - num_ctx_tokens = sum(new_lens) - for rid, total in enumerate(kv_lens): - env.add_request(rid, new_lens[rid]) - env.reserve_cache_pages(rid, total) - - # Per-sequence full KV timeline; chunk passes slice the cached region. - # V buffers are sliced and concatenated as packed rows so every pass's V - # keeps the required H*(nope+v) row stride. - k_full = [] - packed_full = [] - for total in kv_lens: - k_seq, packed_seq = _random_explicit_kv(total, MLA_NUM_HEADS) - k_full.append(k_seq) - packed_full.append(packed_seq) - q = torch.randn( - num_ctx_tokens, MLA_NUM_HEADS * QK_HEAD_DIM, dtype=torch.bfloat16, device="cuda" - ) - - # Production chunk plan (greedy, 64-token buffer over the cached - # regions 96/48/0): loop 0 takes 64 of seq 0; loop 1 the remaining 32 of - # seq 0 plus 32 of seq 1; loop 2 the remaining 16 of seq 1. Merge ops: - # 2 = copy on a sequence's first pass, 1 = merge, 0 = skip; the final - # new-token pass merges (or copies for the cache-less seq 2). - chunk_lens = [[64, 0, 0], [32, 32, 0], [0, 16, 0]] - chunk_offsets = [[0, 0, 0], [64, 0, 0], [96, 32, 0]] - merge_ops = [[2, 0, 0], [1, 2, 0], [0, 1, 0], [1, 1, 2]] - - merged = torch.empty( - num_ctx_tokens, MLA_NUM_HEADS * V_HEAD_DIM, dtype=torch.bfloat16, device="cuda" - ) - merged_stats = torch.empty(num_ctx_tokens, MLA_NUM_HEADS, 2, dtype=torch.float32, device="cuda") - temp_stats = torch.empty_like(merged_stats) - cu_q = torch.tensor( - [0, *torch.tensor(new_lens).cumsum(0).tolist()], - dtype=torch.int64, - device="cuda", - ) - - def run_pass( - pass_kv_lens: List[int], - offsets: List[int], - mask_type: int, - ops: List[int], - ) -> None: - slices = list(enumerate(zip(offsets, pass_kv_lens))) - k_buf = torch.cat([k_full[s][o : o + n] for s, (o, n) in slices]) - v_buf = _v_split_view( - torch.cat([packed_full[s][o : o + n] for s, (o, n) in slices]), - MLA_NUM_HEADS, - ) - temp_out = env.call_context_no_append( - [0, 1, 2], - new_lens, - pass_kv_lens, - q, - k_buf, - v_buf, - mask_type=mask_type, - softmax_stats_tensor=temp_stats, - ) - ref_out, ref_stats = _explicit_kv_reference( - q, k_buf, v_buf, new_lens, pass_kv_lens, mask_type, MLA_NUM_HEADS - ) - _assert_valid_rows_close(temp_out, ref_out, ref_stats) - _assert_stats_close(temp_stats, ref_stats) - # The downstream merge op consumes the emitted stats directly. - torch.ops.trtllm.merge_chunked_attention_for_mla( - merged, - temp_out, - merged_stats, - temp_stats, - len(new_lens), - cu_q, - max(new_lens), - torch.tensor(ops, dtype=torch.int64, device="cuda"), - MLA_NUM_HEADS, - V_HEAD_DIM, - ) - torch.cuda.synchronize() - - for loop_idx, (lens, offs) in enumerate(zip(chunk_lens, chunk_offsets)): - run_pass(lens, offs, MASK_PADDING, merge_ops[loop_idx]) - # Final pass: causal attention of the new tokens over themselves only. - run_pass(new_lens, cached_lens, MASK_CAUSAL, merge_ops[-1]) - - # The merged result must equal single-pass attention over the full - # [cached + new] range (bottom-right-aligned causal). - k_all = torch.cat(k_full) - v_all = _v_split_view(torch.cat(packed_full), MLA_NUM_HEADS) - ref_full, ref_full_stats = _explicit_kv_reference( - q, k_all, v_all, new_lens, kv_lens, MASK_CAUSAL, MLA_NUM_HEADS - ) - torch.testing.assert_close(merged, ref_full, rtol=RTOL, atol=ATOL) - _assert_stats_close(merged_stats, ref_full_stats) - - # ─── MLA at q_lora_rank = 0 (checkpoints with no q-LoRA) ─── # The values swept per MLA call flavor. Index 0 is the reference run and @@ -4433,230 +3592,6 @@ def run_pass( KV_SCALE_SWEEP = [1.0, 1.5, 2.0] -def _mla_page32_env(num_heads: int, q_lora_rank: int) -> _MlaPagedEnv: - """A page-32 MLA env at one head count and one q_lora_rank. Head count and - page size are certified axes of their own; fixing both within a sweep - leaves q_lora_rank the only variable. The two counts swept are the shipped - ones: 32 (deepseek-v3-lite tp1) and 8 (its tep4 slice), which is the count - that takes its own compiled decode kernel (...VarSeqQ8...).""" - return _MlaPagedEnv( - num_heads=num_heads, - max_blocks_per_seq=MLA_MAX_SEQ_LEN // MLA_PAGE32, - tokens_per_block=MLA_PAGE32, - q_lora_rank=q_lora_rank, - ) - - -def _assert_q_lora_rank_inert( - flavor: str, - runs: List[Tuple[int, Dict[str, torch.Tensor], int]], -) -> None: - """Every observable of the reference run — outputs, the whole paged pool, - the in-place-written inputs, and the size the op grew the workspace to — - must come back bitwise identical at every other q_lora_rank.""" - base_rank, base, base_ws = runs[0] - for rank, cur, ws in runs[1:]: - for key, expected in base.items(): - got = cur[key] - assert torch.equal(expected, got), ( - f"{flavor}: {key} differs at q_lora_rank={rank} vs " - f"{base_rank}: {int((expected != got).sum())}/{expected.numel()} " - f"elements, max abs " - f"{(expected.float() - got.float()).abs().max().item():.3e}" - ) - assert ws == base_ws, ( - f"{flavor}: workspace grew to {ws} bytes at q_lora_rank={rank}, " - f"{base_ws} at {base_rank}" - ) - - -def _mla_qlora_context_prefill_case(num_heads: int, seed: int) -> None: - """Fresh-prefill MLA context swept over q_lora_rank (page 32). - - This is the flavor with the most machinery behind the MLA meta params — - in-kernel GPT-J RoPE of q_pe/k_pe plus the paged latent append — so it is - where a q_lora_rank-dependent layout would surface. The zero run is - checked against the fp32 reference and its append against the mirrored - latent rows; the whole sweep is then checked bitwise against the - q_lora_rank=1536 run on identical inputs. - """ - seq_lens = [96, 33] - num_tokens = sum(seq_lens) - h = num_heads - torch.manual_seed(seed) - q_src, k_src, v, latent_src = _random_context_inputs(num_tokens, h) - v_src = v.clone() # v is shared across the sweep: it must stay an input - - runs: List[Tuple[int, Dict[str, torch.Tensor], int]] = [] - for rank in Q_LORA_RANK_SWEEP: - env = _mla_page32_env(h, rank) - for rid, ln in enumerate(seq_lens): - env.add_request(rid, ln) - q, k, latent = q_src.clone(), k_src.clone(), latent_src.clone() - out = env.call_context([0, 1], seq_lens, q, k, v, latent) - - if rank == Q_LORA_RANK_ZERO: - # The surface under test has to be right, not merely reproducible. - ref = env.context_reference(seq_lens, q_src, k_src, v, latent_src) - torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) - assert max(len(env.pages[rid]) for rid in (0, 1)) > 1 - for rid in (0, 1): - env.check_cache(rid) - assert torch.equal(latent, latent_src) # latent_cache is read-only - q3 = q.view(num_tokens, h, QK_HEAD_DIM) - k3 = k.view(num_tokens, h, QK_HEAD_DIM) - q3_src = q_src.view(num_tokens, h, QK_HEAD_DIM) - k3_src = k_src.view(num_tokens, h, QK_HEAD_DIM) - nope = QK_NOPE_HEAD_DIM - assert torch.equal(q3[..., :nope], q3_src[..., :nope]) - assert torch.equal(k3[..., :nope], k3_src[..., :nope]) - assert not torch.equal(q3[..., nope:], q3_src[..., nope:]) - assert not torch.equal(k3[..., nope:], k3_src[..., nope:]) - - runs.append( - ( - rank, - { - "output": out, - "pool": env.pool.clone(), - "q_after_call": q, - "k_after_call": k, - }, - env.workspace.numel(), - ) - ) - assert torch.equal(v, v_src), "the sweep must have run on identical inputs" - _assert_q_lora_rank_inert(f"fresh-prefill context H={h}", runs) - - -def test_bf16_mla_qlora0_context_prefill_h32() -> None: - _mla_qlora_context_prefill_case(MLA_NUM_HEADS_H32, 305) - - -def test_bf16_mla_qlora0_context_prefill_h8() -> None: - _mla_qlora_context_prefill_case(MLA_NUM_HEADS_H8, 315) - - -def _mla_qlora_generation_decode_case(num_heads: int, seed: int) -> None: - """Latent-MQA MLA decode swept over q_lora_rank (page 32). - - The generation phase is the one that JIT-compiles its FMHA kernel, so - this is where q_lora_rank would show up as a kernel-selection axis: the - whole sweep runs in one process against a single compiled decode kernel - (the compile cache is keyed by head count and page size, both fixed - here). Each variant builds its own history through its own context call, - so a rank that changed the append would already separate the pools. - """ - prefill_lens = [64, 31] - h = num_heads - torch.manual_seed(seed) - q_src, k_src, v, latent_src = _random_context_inputs(sum(prefill_lens), h) - v_src = v.clone() - # Decode rows and fused q are drawn once and replayed per variant. - decode_rows = [ - [torch.randn(LATENT_DIM, dtype=torch.bfloat16, device="cuda") for _ in prefill_lens] - for _ in range(2) - ] - fused_qs = [ - torch.randn(len(prefill_lens), h * LATENT_DIM, dtype=torch.bfloat16, device="cuda") - for _ in range(2) - ] - - runs: List[Tuple[int, Dict[str, torch.Tensor], int]] = [] - for rank in Q_LORA_RANK_SWEEP: - env = _mla_page32_env(h, rank) - for rid, ln in enumerate(prefill_lens): - env.add_request(rid, ln) - env.call_context([0, 1], prefill_lens, q_src.clone(), k_src.clone(), v, latent_src.clone()) - observed = {"pool_after_prefill": env.pool.clone()} - for step in range(2): - for rid in range(len(prefill_lens)): - env.append_decode_latent(rid, decode_rows[step][rid]) - # call_generation draws the (unconsumed) latent_cache/q_pe - # garbage from the global RNG: reseed so every variant is fed the - # same garbage and q_lora_rank stays the only difference. - torch.manual_seed(seed * 10 + step) - fused_q = fused_qs[step].clone() - out = env.call_generation([0, 1], fused_q) - if rank == Q_LORA_RANK_ZERO: - ref = env.generation_reference([0, 1], fused_qs[step]) - torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) - assert torch.equal(fused_q, fused_qs[step]) - observed[f"decode_out_{step}"] = out - observed[f"pool_after_decode_{step}"] = env.pool.clone() - runs.append((rank, observed, env.workspace.numel())) - assert torch.equal(v, v_src), "the sweep must have run on identical inputs" - _assert_q_lora_rank_inert(f"generation decode H={h}", runs) - - -def test_bf16_mla_qlora0_generation_decode_h32() -> None: - _mla_qlora_generation_decode_case(MLA_NUM_HEADS_H32, 306) - - -def test_bf16_mla_qlora0_generation_decode_h8() -> None: - """H=8 decode compiles ...VarSeqQ8Kv128... where 16 and 32 both take - ...VarSeqQ16...: a genuinely different kernel, so the H=32 sweep says - nothing about it. This is the tep4 deepseek-v3-lite cell, whose config - carries "q_lora_rank": null.""" - _mla_qlora_generation_decode_case(MLA_NUM_HEADS_H8, 316) - - -def _mla_qlora_context_cached_kv_no_append_case(num_heads: int, seed: int) -> None: - """Cached-KV (no-append) MLA context swept over q_lora_rank (page 32). - - latent_cache=None skips the in-kernel RoPE and the append, so this flavor - reaches a different context kernel from the fresh-prefill one; the paged - pool is pre-filled with noise and must come back bitwise untouched at - every rank. - """ - cached_lens, new_lens = [96, 31, 0], [32, 9, 25] - kv_lens = [c + n for c, n in zip(cached_lens, new_lens)] - h = num_heads - torch.manual_seed(seed) - q_src = torch.randn(sum(new_lens), h * QK_HEAD_DIM, dtype=torch.bfloat16, device="cuda") - k_src, packed_kv = _random_explicit_kv(sum(kv_lens), h) - v = _v_split_view(packed_kv, h) - v_src = v.clone() - pool_fill: Optional[torch.Tensor] = None - - runs: List[Tuple[int, Dict[str, torch.Tensor], int]] = [] - for rank in Q_LORA_RANK_SWEEP: - env = _mla_page32_env(h, rank) - for rid, total in enumerate(kv_lens): - env.add_request(rid, new_lens[rid]) - env.reserve_cache_pages(rid, total) - if pool_fill is None: - pool_fill = torch.randn_like(env.pool) - env.pool.copy_(pool_fill) # the op must neither read nor write it - q, k = q_src.clone(), k_src.clone() - out = env.call_context_no_append([0, 1, 2], new_lens, kv_lens, q, k, v) - - if rank == Q_LORA_RANK_ZERO: - ref, _ = _explicit_kv_reference(q, k, v, new_lens, kv_lens, MASK_CAUSAL, h) - torch.testing.assert_close(out, ref, rtol=RTOL, atol=ATOL) - assert torch.equal(q, q_src) - assert torch.equal(k, k_src) - assert torch.equal(env.pool, pool_fill) - - runs.append( - ( - rank, - {"output": out, "pool": env.pool.clone()}, - env.workspace.numel(), - ) - ) - assert torch.equal(v, v_src), "the sweep must have run on identical inputs" - _assert_q_lora_rank_inert(f"no-append context H={h}", runs) - - -def test_bf16_mla_qlora0_context_cached_kv_no_append_h32() -> None: - _mla_qlora_context_cached_kv_no_append_case(MLA_NUM_HEADS_H32, 307) - - -def test_bf16_mla_qlora0_context_cached_kv_no_append_h8() -> None: - _mla_qlora_context_cached_kv_no_append_case(MLA_NUM_HEADS_H8, 317) - - # ─── The DeepSeek-R1-0528 cell: H=128 + YaRN rope table + q_scaling != 1 ─── # # Three axes move at once relative to the deepseek-v3-lite cells above: 128 @@ -4667,80 +3602,6 @@ def test_bf16_mla_qlora0_context_cached_kv_no_append_h8() -> None: # separately, by the two axis sweeps after them. -def _assert_outside_band( - out: torch.Tensor, rival: torch.Tensor, threshold: float, label: str -) -> float: - """A rival hypothesis must land this far outside the tolerance band around - the observed output. Without it, matching the positive reference proves - little: an op that silently ignored the argument under test would pass the - positive comparison too wherever the argument's effect happens to be - small. Returns the separation in units of the allowance.""" - diff = (out.float() - rival.float()).abs().max().item() - allowance = ATOL + RTOL * rival.float().abs().max().item() - ratio = diff / allowance - assert ratio > threshold, f"{label}: only {ratio:.4g}x outside the tolerance band" - return ratio - - -def test_bf16_mla_r1_cell_context_prefill_h128() -> None: - """Fresh-prefill MLA context at the full R1 cell: 128/128/192, page 32, - the checkpoint's YaRN table, q_scaling = 1/mscale^2.""" - _mla_context_prefill_case( - MLA_NUM_HEADS_H128, - 925, - MLA_PAGE32, - [96, 33], - q_scaling=R1_Q_SCALING, - rope=_r1_rope_params(MLA_MAX_SEQ_LEN), - rope_scalars=R1_ROPE_SCALARS, - ) - - -def test_bf16_mla_r1_cell_generation_decode_h128() -> None: - """Latent-MQA decode at the full R1 cell: 128/1/576, page 32. The rope - table is passed but nothing is rotated here — the scale is what carries.""" - _mla_generation_decode_case( - MLA_NUM_HEADS_H128, - 926, - MLA_PAGE32, - [64, 31], - q_scaling=R1_Q_SCALING, - rope=_r1_rope_params(MLA_MAX_SEQ_LEN), - rope_scalars=R1_ROPE_SCALARS, - ) - - -def test_bf16_mla_r1_cell_mixed_batch_h128() -> None: - """Mixed batch at the full R1 cell: the two phase calls share one set of - batch tensors, one YaRN table and one q_scaling.""" - _mla_mixed_batch_case( - MLA_NUM_HEADS_H128, - 927, - MLA_PAGE32, - 64, - 33, - q_scaling=R1_Q_SCALING, - rope=_r1_rope_params(MLA_MAX_SEQ_LEN), - rope_scalars=R1_ROPE_SCALARS, - ) - - -def test_bf16_mla_r1_cell_context_cached_kv_no_append_h128() -> None: - """Cached-KV (no-append) context at the full R1 cell. This flavor applies - no rope at all — q arrives pre-rotated — so it is where q_scaling is the - only one of the three op-side axes with an effect.""" - _mla_context_cached_kv_case( - MLA_NUM_HEADS_H128, - 928, - MLA_PAGE32, - [96, 31, 0], - [32, 9, 25], - q_scaling=R1_Q_SCALING, - rope=_r1_rope_params(MLA_MAX_SEQ_LEN), - rope_scalars=R1_ROPE_SCALARS, - ) - - # q_scaling values swept in both phases. 1.0 is the value certified before # this run, R1_Q_SCALING (~0.53366) is what the checkpoint's YaRN config # wants, and 2.0 / 0.25 bracket it from both sides so the argument is @@ -4760,138 +3621,6 @@ def test_bf16_mla_r1_cell_context_cached_kv_no_append_h128() -> None: Q_SCALING_MIN_SEPARATION = 5.0 -def test_bf16_mla_q_scaling_axis_h128() -> None: - """q_scaling is the softmax-scale axis of both MLA phases: QK^T is scaled - by 1 / (q_scaling * sqrt(nope + R)) in the context call and in the - generation call alike, the latter despite its head_size being C + R. - - Four values on identical inputs. Each run is checked against a reference - built at its own value and gated outside the references of the other - three, and the paged latent append is asserted bitwise identical across - the sweep — q_scaling moves the softmax scale and nothing else. - """ - h = MLA_NUM_HEADS_H128 - prefill_lens = [96, 33] - torch.manual_seed(935) - q_src, k_src, v, latent_src = _random_context_inputs(sum(prefill_lens), h) - decode_rows = [ - torch.randn(LATENT_DIM, dtype=torch.bfloat16, device="cuda") for _ in prefill_lens - ] - fused_q = torch.randn(len(prefill_lens), h * LATENT_DIM, dtype=torch.bfloat16, device="cuda") - - runs: List[Tuple[float, _MlaPagedEnv, torch.Tensor, torch.Tensor]] = [] - for q_scaling in Q_SCALING_SWEEP: - env = _MlaPagedEnv( - num_heads=h, - max_blocks_per_seq=MLA_MAX_SEQ_LEN // MLA_PAGE32, - tokens_per_block=MLA_PAGE32, - q_scaling=q_scaling, - rope=_r1_rope_params(MLA_MAX_SEQ_LEN), - rope_scalars=R1_ROPE_SCALARS, - ) - for rid, ln in enumerate(prefill_lens): - env.add_request(rid, ln) - out_ctx = env.call_context( - [0, 1], prefill_lens, q_src.clone(), k_src.clone(), v, latent_src.clone() - ) - for rid, row in enumerate(decode_rows): - env.append_decode_latent(rid, row) - out_gen = env.call_generation([0, 1], fused_q) - runs.append((q_scaling, env, out_ctx, out_gen)) - - for q_scaling, env, out_ctx, out_gen in runs: - for rival in Q_SCALING_SWEEP: - scale = 1.0 / (rival * math.sqrt(QK_HEAD_DIM)) - ref_ctx = env.context_reference( - prefill_lens, q_src, k_src, v, latent_src, softmax_scale=scale - ) - ref_gen = env.generation_reference([0, 1], fused_q, softmax_scale=scale) - if rival == q_scaling: - torch.testing.assert_close(out_ctx, ref_ctx, rtol=RTOL, atol=ATOL) - torch.testing.assert_close(out_gen, ref_gen, rtol=RTOL, atol=ATOL) - else: - _assert_outside_band( - out_ctx, - ref_ctx, - Q_SCALING_MIN_SEPARATION, - f"context at q_scaling={q_scaling} vs a reference at {rival}", - ) - _assert_outside_band( - out_gen, - ref_gen, - Q_SCALING_MIN_SEPARATION, - f"decode at q_scaling={q_scaling} vs a reference at {rival}", - ) - - base_pool = runs[0][1].pool - for q_scaling, env, _, _ in runs[1:]: - assert torch.equal(env.pool, base_pool), ( - f"the paged latent pool moved at q_scaling={q_scaling}" - ) - - -def _yarn_cos_sin_table( - num_positions: int, - dim: int, - theta: float, - factor: float, - original_max_positions: int, - beta_fast: int, - beta_slow: int, - mscale: float, - mscale_all_dim: float, -) -> torch.Tensor: - """The duplicated-layout fp32 (cos, sin) table a YaRN rope config - produces, built from the published YaRN formula with plain torch. - - A caller that cannot reach TensorRT-LLM's own table builder has to - reproduce this; asserting it against the builder's output is what makes - the formula usable as a contract statement rather than a description. - Layout: dim (cos, sin) pairs per position, the second dim/2 duplicating - the first, flattened to [1, num_positions * dim * 2]. - """ - half = dim // 2 - - def correction_dim(rotations: float) -> float: - return ( - dim - * math.log(original_max_positions / (rotations * 2 * math.pi)) - / (2 * math.log(theta)) - ) - - def attention_mscale(cfg_mscale: float) -> float: - return 1.0 if factor <= 1 else 0.1 * cfg_mscale * math.log(factor) + 1.0 - - low = max(0, math.floor(correction_dim(beta_fast))) - high = min(dim - 1, math.ceil(correction_dim(beta_slow))) - # The table's own amplitude factor — 1.0 whenever the two config mscales - # agree, which is where a model folds mscale into the softmax scale - # instead (see R1_Q_SCALING). - amplitude = attention_mscale(mscale) / attention_mscale(mscale_all_dim) - - pos_freqs = theta ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim) - ramp = torch.clamp( - (torch.arange(half, dtype=torch.float32) - low) / max(high - low, 0.001), 0, 1 - ) - # Interpolated (position-stretched) frequencies below the ramp, the - # original ones above it, blended across it. - inv_freq = ramp / (factor * pos_freqs) + (1 - ramp) / pos_freqs - angles = torch.outer(torch.arange(num_positions, dtype=torch.float32), inv_freq) - angles = torch.cat([angles, angles], dim=-1) # duplicate_data layout - table = torch.stack([torch.cos(angles) * amplitude, torch.sin(angles) * amplitude], dim=-1) - return table.reshape(1, -1).cuda() - - -def _unscaled_cos_sin_table() -> torch.Tensor: - """The plain theta-10000 table every MLA case above this section uses.""" - return RopeParams( - dim=QK_ROPE_HEAD_DIM, - theta=10000.0, - max_positions=MLA_MAX_SEQ_LEN, - duplicate_data=True, - ).create_rope_const_params()[1] - - # A YaRN-table run must land this far outside the tolerance band of a # reference built from the unscaled theta-10000 table. The separation grows # with position, because YaRN only rescales the low-frequency half of the @@ -4903,201 +3632,6 @@ def _unscaled_cos_sin_table() -> torch.Tensor: YARN_TABLE_MIN_SEPARATION = 10.0 -def test_bf16_mla_yarn_rope_table_h128() -> None: - """The in-kernel RoPE of the fresh-prefill context call reads the - rotary_cos_sin table's *content*. - - Pinned here: the YaRN table a DeepSeek-R1-0528 config produces matches a - plain-torch build of the published formula; truncating it to max_seq_len - rows is the leading slice of the 163840-row table an engine allocates; - and the op driven from the torch-built table reproduces a table-aware - fp32 reference while landing far outside a reference built from the - unscaled theta-10000 table (and vice versa), so the table content is - established as load-bearing rather than assumed. - """ - trtllm_inv_freq, trtllm_table = _r1_rope_params(MLA_MAX_SEQ_LEN).create_rope_const_params() - torch_table = _yarn_cos_sin_table( - MLA_MAX_SEQ_LEN, - QK_ROPE_HEAD_DIM, - R1_ROPE_THETA, - R1_ROPE_FACTOR, - R1_ROPE_ORIGINAL_MAX_POSITIONS, - R1_ROPE_BETA_FAST, - R1_ROPE_BETA_SLOW, - R1_ROPE_MSCALE, - R1_ROPE_MSCALE_ALL_DIM, - ) - # Not bitwise: the two evaluate the same blend in different fp32 orders - # ((1 - (1 - ramp)) against ramp). Observed max abs diff 1.9e-6 on cos/sin - # values in [-1, 1] — a few fp32 ulp, 19% of this gate — while formula - # slips land 4-5 orders of magnitude past it: dropping the YaRN - # interpolation entirely (the unscaled table) or moving beta_fast 32 -> 64 - # both differ by 2.0, and taking the amplitude as mscale rather than the - # mscale/mscale_all_dim ratio differs by 3.7e-1. - torch.testing.assert_close(torch_table, trtllm_table, rtol=0.0, atol=1e-5) - assert trtllm_inv_freq is not None and trtllm_inv_freq.numel() == (QK_ROPE_HEAD_DIM // 2) - - # The unscaled table every MLA case above this section uses is the same - # construction at factor = 1: the ramp drops out and inv_freq(d) becomes - # 1/theta^(2d/R). Gated here so the formula is certified for both - # contents rather than only for the YaRN one (observed 3.8e-6, 38% of - # the gate). - unscaled_table = _unscaled_cos_sin_table() - torch.testing.assert_close( - _yarn_cos_sin_table( - MLA_MAX_SEQ_LEN, - QK_ROPE_HEAD_DIM, - R1_ROPE_THETA, - 1.0, - R1_ROPE_ORIGINAL_MAX_POSITIONS, - R1_ROPE_BETA_FAST, - R1_ROPE_BETA_SLOW, - R1_ROPE_MSCALE, - R1_ROPE_MSCALE_ALL_DIM, - ), - unscaled_table, - rtol=0.0, - atol=1e-5, - ) - - # Row content does not depend on the row count, so a table truncated to - # max_seq_len is the production table's leading slice. - _, full_table = _r1_rope_params(R1_MAX_POSITION_EMBEDDINGS).create_rope_const_params() - row = QK_ROPE_HEAD_DIM * 2 - assert torch.equal(full_table.view(-1, row)[:MLA_MAX_SEQ_LEN], trtllm_table.view(-1, row)) - del full_table - torch.cuda.empty_cache() - - h = MLA_NUM_HEADS_H128 - seq_lens = [512, 33] # far enough past the ramp for YaRN to bite - - def prefill(table: torch.Tensor) -> Tuple[_MlaPagedEnv, torch.Tensor, tuple]: - torch.manual_seed(936) - env = _MlaPagedEnv( - num_heads=h, - max_blocks_per_seq=MLA_MAX_SEQ_LEN // MLA_PAGE32, - tokens_per_block=MLA_PAGE32, - q_scaling=R1_Q_SCALING, - rope=_r1_rope_params(MLA_MAX_SEQ_LEN), - rope_scalars=R1_ROPE_SCALARS, - ) - env.rotary_cos_sin = table # op and reference both read it here - for rid, ln in enumerate(seq_lens): - env.add_request(rid, ln) - q, k, v, latent = _random_context_inputs(sum(seq_lens), h) - pre = (list(seq_lens), q.clone(), k.clone(), v, latent.clone()) - out = env.call_context([0, 1], seq_lens, q, k, v, latent) - return env, out, pre - - for driven, rival, label in ( - (torch_table, unscaled_table, "YaRN table"), - (unscaled_table, torch_table, "unscaled table"), - ): - env, out, pre = prefill(driven) - torch.testing.assert_close(out, env.context_reference(*pre), rtol=RTOL, atol=ATOL) - env.rotary_cos_sin = rival - _assert_outside_band( - out, - env.context_reference(*pre), - YARN_TABLE_MIN_SEPARATION, - f"{label} run against a reference built from the other table", - ) - env.rotary_cos_sin = driven - for rid in (0, 1): - env.check_cache(rid) # rope(k_pe) in the append follows the table - - -def _r1_prefill_observables( - rope_scalars: _RopeScalars, - cos_sin: Optional[torch.Tensor] = None, - inv_freq_mode: str = "keep", -) -> Dict[str, torch.Tensor]: - """One R1-cell fresh-prefill context call on fixed inputs, returning every - observable a rope argument could move: the output rows, the whole paged - latent pool (which receives rope(k_pe)), and the in-place-roped q/k.""" - seq_lens = [96, 33] - torch.manual_seed(936) - env = _MlaPagedEnv( - num_heads=MLA_NUM_HEADS_H128, - max_blocks_per_seq=MLA_MAX_SEQ_LEN // MLA_PAGE32, - tokens_per_block=MLA_PAGE32, - q_scaling=R1_Q_SCALING, - rope=_r1_rope_params(MLA_MAX_SEQ_LEN), - rope_scalars=rope_scalars, - ) - if cos_sin is not None: - env.rotary_cos_sin = cos_sin - if inv_freq_mode == "zeros": - env.rotary_inv_freq = torch.zeros_like(env.rotary_inv_freq) - elif inv_freq_mode == "none": - env.rotary_inv_freq = None - for rid, ln in enumerate(seq_lens): - env.add_request(rid, ln) - q, k, v, latent = _random_context_inputs(sum(seq_lens), MLA_NUM_HEADS_H128) - out = env.call_context([0, 1], seq_lens, q, k, v, latent) - return { - "output": out, - "pool": env.pool.clone(), - "q_after_call": q, - "k_after_call": k, - } - - -def test_bf16_mla_rope_scalars_inert_h128() -> None: - """With the table held fixed, the seven scalar rope arguments and - rotary_inv_freq move no observable of an MLA call. - - That is what makes the YaRN certification portable: a caller supplies the - scaled table and may pass whatever scalars its config carries. The last - check is the control that gives the bitwise comparisons meaning — the one - rope input the op does read (the table) is swapped, and the same - comparison has to see it. - """ - base = _r1_prefill_observables(R1_ROPE_SCALARS) - variants = { - # The set every MLA case above this section passes. - "unscaled-config scalars": _r1_prefill_observables( - _RopeScalars( - rope_max_positions=MLA_MAX_SEQ_LEN, - rope_original_max_positions=MLA_MAX_SEQ_LEN, - ) - ), - # Run-to-run determinism control for the comparisons around it. - "R1 scalars again": _r1_prefill_observables(R1_ROPE_SCALARS), - # Nothing a config would produce: a different theta, a different - # scaling family, m-scales far from 1, and position windows shorter - # than the sequences in flight. - "out-of-range scalars": _r1_prefill_observables( - _RopeScalars( - rope_base=500000.0, - rope_scale_type=3, - rope_scale=7.5, - rope_short_m_scale=3.0, - rope_long_m_scale=9.0, - rope_max_positions=77, - rope_original_max_positions=13, - ) - ), - "rotary_inv_freq zeroed": _r1_prefill_observables(R1_ROPE_SCALARS, inv_freq_mode="zeros"), - "rotary_inv_freq=None": _r1_prefill_observables(R1_ROPE_SCALARS, inv_freq_mode="none"), - } - for label, observed in variants.items(): - for key, expected in base.items(): - got = observed[key] - assert torch.equal(expected, got), ( - f"{label}: {key} is not bitwise identical — " - f"{int((expected != got).sum())}/{expected.numel()} elements, " - f"max abs {(expected.float() - got.float()).abs().max().item():.3e}" - ) - - swapped = _r1_prefill_observables(R1_ROPE_SCALARS, cos_sin=_unscaled_cos_sin_table()) - for key in ("output", "pool", "q_after_call", "k_after_call"): - assert not torch.equal(base[key], swapped[key]), ( - f"swapping the rope table left {key} bitwise unchanged — the " - f"inertness comparisons above cannot see a rope change at all" - ) - - # ─── MLA over an fp8-e4m3 latent pool (quant_mode 128) ───────────────── # The DeepSeek-R1-0528-FP4 cell: H = 128, page 32, C/R/nope/v = # 512/64/128/128, q_lora_rank 1536, one latent row of C+R e4m3 bytes per @@ -6643,71 +5177,6 @@ def test_fp8_mla_mtp_mixed_batch_h128() -> None: env.check_unwritten_pool_zero([0, 1]) -def test_bf16_mla_mtp_decode_h128() -> None: - """The same MTP generation surface over a bf16 latent pool. - - Two things this adds over the fp8 cases. The bf16 band is 8x tighter, so - the full-mask rival separates properly here — observed 17-23x outside, - against 0.22 for the matching model, which is what makes "the within-block - mask is causal" a gated result on realistic inputs and not only a readout. - And it pins that P > 1 is a property of the MLA decode kernel rather than - of the fp8 decode path, which is the only one the R1 target runs. - """ - lens = [33, 50] - offsets = [0, KV_LORA_RANK // 2] - for p in (2, 4): - torch.manual_seed(1700 + p) - env = _mtp_readout_env(lens, offsets, fp8=False) - fused_q = _mtp_readout_q(len(lens) * p, env.num_heads) - out = env.call_generation([0, 1], fused_q, predicted_tokens_per_seq=p) - _assert_mtp_causal_readout( - env, out, lens, offsets, p, f"bf16 P={p}", rtol=MTP_READOUT_BF16_RTOL - ) - - torch.manual_seed(1750 + p) - real = _MlaPagedEnv( - num_heads=MLA_NUM_HEADS_H128, - max_blocks_per_seq=MLA_MAX_SEQ_LEN // MLA_PAGE32, - tokens_per_block=MLA_PAGE32, - q_lora_rank=Q_LORA_RANK_DSV3, - ) - for rid, ln in enumerate(lens): - real.add_request(rid, ln - p) - for _ in range(ln): - real.append_decode_latent( - rid, - torch.randn(LATENT_DIM, dtype=torch.bfloat16, device="cuda") * 0.5, - ) - fq = ( - torch.randn( - len(lens) * p, - real.num_heads * LATENT_DIM, - dtype=torch.bfloat16, - device="cuda", - ) - * 0.3 - ) - pool_before = real.pool.clone() - got = real.call_generation([0, 1], fq, predicted_tokens_per_seq=p) - torch.testing.assert_close( - got, - real.generation_reference([0, 1], fq, predicted_tokens_per_seq=p), - rtol=RTOL, - atol=ATOL, - ) - _assert_far_outside_band( - got, - real.generation_reference([0, 1], fq, predicted_tokens_per_seq=p, full_mask=True), - 5.0, - f"bf16 MTP decode at P={p} against a full-mask rival", - ATOL, - RTOL, - ) - assert _bitwise_equal(pool_before, real.pool), ( - "the MLA generation call wrote to the bf16 pool" - ) - - def test_fp8_mla_mtp_decode_sees_a_torn_pool_h128() -> None: """Positive control on the decode-side comparison every claim above rests on: it must be able to see a pool whose content moved under it. diff --git a/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_bmm_out.py b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_bmm_out.py index fefc1067e6d5..a614806ba5db 100644 --- a/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_bmm_out.py +++ b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_bmm_out.py @@ -61,15 +61,3 @@ def test_bf16_unaligned_shapes() -> None: torch.manual_seed(2) for batch, m, k, n in [(3, 5, 100, 60), (7, 13, 333, 129)]: _check(*_make(batch, m, k, n, torch.bfloat16)) - - -def test_fp16() -> None: - torch.manual_seed(3) - for batch, m, k, n in [(64, 1, 512, 128), (8, 1024, 512, 256)]: - _check(*_make(batch, m, k, n, torch.float16)) - - -def test_fp32() -> None: - torch.manual_seed(4) - for batch, m, k, n in [(32, 2, 256, 128), (4, 1024, 512, 256)]: - _check(*_make(batch, m, k, n, torch.float32)) diff --git a/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_cublas_mm.py b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_cublas_mm.py index 4c88737300d8..51a9fcad2c49 100644 --- a/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_cublas_mm.py +++ b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_cublas_mm.py @@ -111,23 +111,6 @@ def test_bf16_unaligned_shapes() -> None: _check(mat_a, mat_b, bias) -def test_fp16() -> None: - torch.manual_seed(4) - for m, k, n in [(1, 4096, 4096), (1024, 2048, 2048)]: - mat_a, mat_b, bias = _make(m, k, n, torch.float16) - _check(mat_a, mat_b) - _check(mat_a, mat_b, bias) - - -def test_fp32() -> None: - # no bias: the op silently ignores bias when inputs are fp32 - # (contract precondition; guarded by an assert in the wrapper) - torch.manual_seed(5) - for m, k, n in [(2, 1024, 1024), (256, 2048, 1024)]: - mat_a, mat_b, _ = _make(m, k, n, torch.float32) - _check(mat_a, mat_b) - - def test_fp8_e4m3_to_bf16() -> None: # fp8 inputs require an explicit out_dtype; no scales are applied (alpha=1) torch.manual_seed(6) diff --git a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_fp4_block_scale_moe_runner.py b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_fp4_block_scale_moe_runner.py index 5c6ba3b33468..849a658aa9dd 100644 --- a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_fp4_block_scale_moe_runner.py +++ b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_fp4_block_scale_moe_runner.py @@ -420,9 +420,10 @@ def _assert_moe_close(y: torch.Tensor, ref: torch.Tensor) -> None: 3.20 (H=2560) and 9.62 / 4.68 (H=7168), but only the RMS gate is scale-free — read it, not the element one, when sizing a caller's tolerance. - Both gates bite — `test_reference_discriminates` shows, at H=256/I=128 - and at H=7168/I=2048 respectively, a gate/up swap (263 / 278 ulp), a - scalar-role swap (177 / 48), a missing scale swizzle (408 / 357), a + Both gates were sized against deliberately wrong computations when they + were set. At H=256/I=128 and at H=7168/I=2048 respectively: a gate/up + swap (263 / 278 ulp), a scalar-role swap (177 / 48), a missing scale + swizzle (408 / 357), a missing row shuffle (508 / 432), and a reference that skips the FC1-output requantization (43 / 28 ulp element, 23 / 23 ulp RMS); the smallest RMS distance any of those reaches is 39.8, at the R1 shape's scalar-role swap. @@ -671,83 +672,6 @@ def test_intermediate_is_nvfp4_requantized(): ) -def test_reference_discriminates(): - """Every layout / scalar-role mistake lands far outside the numeric gate. - - Run at two hidden/intermediate pairs, because both operand layouts are - functions of `H` and `I`: the small one, and the DeepSeek-R1 routed shape - (`H = 7168`, `I = 2048`) at a reduced expert count — the row shuffle and - the 128x4 scale swizzle are per-expert, so `E` is free to be small while - the layout question is entirely `H`/`I`. - """ - for num_experts, hidden, inter, top_k, num_tokens, seed in ( - (8, 256, 128, 2, 32, 303), - (16, 7168, 2048, 8, 32, 7168), - ): - args, ref = _build(num_experts, hidden, inter, seed=seed) - gen = ref["gen"] - g1 = 137.0 - data, sf, x = _rand_activation(num_tokens, hidden, g1, gen) - ids, wts = _routing(num_tokens, num_experts, top_k, gen) - g2 = _calibrate_g2(x, ids, ref) - s1, sg, s2 = _scalars(num_experts, g1, g2) - base = dict( - topk_ids=ids, - topk_weights=wts, - output1_scale_scalar=s1, - output1_scale_gate_scalar=sg, - output2_scale_scalar=s2, - ) - y = _call(data, sf, args, num_experts, top_k, **base)[0] - good = _ref_moe(x, ids, wts, ref, g2).to(torch.bfloat16) - _assert_moe_close(y, good) - print(f" E={num_experts} H={hidden} I={inter} top-{top_k}, T={num_tokens}:") - - up_p, gt_p, up_s, gt_s, dn_p, dn_s = ref["raw"] - variants = { - "[gate|up] instead of [up|gate]": _dev( - y, _ref_moe(x, ids, wts, ref, g2, swap_gate_up=True).to(torch.bfloat16) - ), - "output1 scalars swapped": _dev( - y, _ref_moe(x, ids, wts, ref, g2, swap_scalars=True).to(torch.bfloat16) - ), - "no intermediate requantization": _dev( - y, - _ref_moe(x, ids, wts, ref, g2, quantize_intermediate=False).to(torch.bfloat16), - ), - } - for name, (elt, rms) in variants.items(): - assert elt > 25.0, f"{name} only {elt:.1f} ulp away — not discriminated" - print(f" {name:34s} {elt:8.1f} elt / {rms:6.2f} rms ulp") - - # operand-preparation mistakes: the kernel is fed the wrong bytes - for name, kw in ( - ( - "weight scales not 128x4 swizzled", - dict( - gemm1_weights_scale=_prep_fc1(up_p, gt_p, up_s, gt_s, swizzle=False)[1], - gemm2_weights_scale=_prep_fc2(dn_p, dn_s, swizzle=False)[1], - ), - ), - ( - "rows not shuffled", - dict( - gemm1_weights=_prep_fc1(up_p, gt_p, up_s, gt_s, shuffle=False)[0], - gemm1_weights_scale=_prep_fc1(up_p, gt_p, up_s, gt_s, shuffle=False)[1], - gemm2_weights=_prep_fc2(dn_p, dn_s, shuffle=False)[0], - gemm2_weights_scale=_prep_fc2(dn_p, dn_s, shuffle=False)[1], - ), - ), - ): - bad = _call(data, sf, args, num_experts, top_k, **base, **kw)[0] - elt, rms = _dev(bad, good.float()) - assert elt > 25.0, f"{name} only {elt:.1f} ulp away — not discriminated" - print(f" {name:34s} {elt:8.1f} elt / {rms:6.2f} rms ulp") - del args, ref, y, good, data, sf, x - torch.cuda.empty_cache() - print(" test_reference_discriminates OK") - - def test_fp4_quantize_pairing(): """`fp4_quantize(x, g, 16, False, False)` is the activation pairing. diff --git a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_fused_moe.py b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_fused_moe.py index 7b56c3c0f93d..94ec6e153bb2 100644 --- a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_fused_moe.py +++ b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_fused_moe.py @@ -97,9 +97,10 @@ def _assert_moe_close(y: torch.Tensor, ref: torch.Tensor) -> None: is 4 ulp of relative RMS. Measured worst case across everything this file drives: 2.8 ulp / 1.2 ulp on the cold fallback tactic, 3.98 ulp / 1.43 ulp once test_r1_mtp_tactic_space walks the whole tactic space -- 2.0x and 2.8x - margin. Both gates bite: test_reference_discriminates shows a wrong - computation landing at 176+ ulp per element and ~190 ulp RMS, and - test_r1_mtp_reference_discriminates at 284-997 / 197-212 ulp. + margin. Both gates were sized against deliberately wrong computations when + they were set: those land two orders of magnitude out (176+ ulp per element + and ~190 ulp RMS; 284-997 / 197-212 on the R1 geometry), so the bound + separates a real defect from accumulation order by a wide margin. """ assert y.dtype == ref.dtype, (y.dtype, ref.dtype) assert y.shape == ref.shape, (y.shape, ref.shape) @@ -136,48 +137,6 @@ def _make( return x, ids.to(torch.int32), torch.softmax(vals, dim=-1), w31, w2 -def test_qwen3_moe_shape_bf16() -> None: - # Qwen3-30B-A3B MoE block: hidden 2048, moe_intermediate 768, 128 experts, - # top-8. Decode-like through prefill-like token counts. - # 8192 is one token past the op's default tune_max_num_tokens bucket cap. - for num_tokens in [1, 2, 37, 256, 2048, 8192]: - x, ids, scales, w31, w2 = _make(num_tokens, 2048, 768, 128, 8, seed=num_tokens) - y = fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, [])[0] - assert y.shape == (num_tokens, 2048) and y.is_contiguous() - _assert_moe_close(y, _ref_moe(x, ids, scales, w31, w2)) - del x, ids, scales, w31, w2, y - torch.cuda.empty_cache() - - -def test_dtypes() -> None: - for dtype in [torch.bfloat16, torch.float16]: - for num_tokens in [1, 64, 512]: - x, ids, scales, w31, w2 = _make( - num_tokens, 1024, 512, 32, 4, seed=num_tokens, dtype=dtype - ) - y = fused_moe(x, ids, scales, w31, None, w2, None, dtype, [])[0] - _assert_moe_close(y, _ref_moe(x, ids, scales, w31, w2)) - - -def test_shape_sweep() -> None: - # hidden and inter must be multiples of 8; expert count and top-k are free. - cases = [ - (8, 8, 8, 2, 1), - (4, 16, 8, 1, 1), - (33, 192, 96, 7, 3), - (64, 128, 64, 8, 8), - (16, 256, 128, 256, 2), - (9, 512, 256, 32, 16), - (64, 2048, 768, 128, 1), - ] - for num_tokens, hidden, inter, num_experts, top_k in cases: - x, ids, scales, w31, w2 = _make( - num_tokens, hidden, inter, num_experts, top_k, seed=hidden + num_experts - ) - y = fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, [])[0] - _assert_moe_close(y, _ref_moe(x, ids, scales, w31, w2)) - - def test_expert_biases() -> None: g = torch.Generator(device=DEV).manual_seed(77) x, ids, scales, w31, w2 = _make(64, 512, 256, 16, 4, seed=77) @@ -313,12 +272,6 @@ def test_swiglu_alpha_beta_limit() -> None: assert torch.equal(y5, y7), "activation_type 7 differed from 5" -def test_geglu_activation() -> None: - x, ids, scales, w31, w2 = _make(32, 256, 128, 16, 4, seed=19) - y = fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, [], activation_type=6)[0] - _assert_moe_close(y, _ref_moe(x, ids, scales, w31, w2, act="geglu")) - - def test_use_fused_finalize_false() -> None: x, ids, scales, w31, w2 = _make(64, 512, 256, 16, 4, seed=23) y = fused_moe( @@ -373,29 +326,6 @@ def test_inputs_untouched_and_deterministic() -> None: assert torch.equal(y_a, y_b), "two identical calls disagreed" -def test_reference_discriminates() -> None: - # The gates must reject a computation that only differs in which half of - # fc1 is the gate: without this control the tolerances prove nothing. - x, ids, scales, w31, w2 = _make(64, 256, 128, 16, 4, seed=37) - y = fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, [])[0] - swapped = torch.cat([w31[:, 128:], w31[:, :128]], dim=1).contiguous() - try: - _assert_moe_close(y, _ref_moe(x, ids, scales, swapped, w2)) - except AssertionError: - pass - else: - raise AssertionError("tolerance accepted a gate/up-swapped reference") - # ...and a dropped expert must be caught too - dropped = scales.clone() - dropped[:, 0] = 0.0 - try: - _assert_moe_close(y, _ref_moe(x, ids, dropped, w31, w2)) - except AssertionError: - pass - else: - raise AssertionError("tolerance accepted a reference missing one expert") - - def test_wrapper_rejects_output_dtype_mismatch() -> None: # Observed silent failure: the store happens in the activation dtype while # the buffer is allocated as output_dtype, so the bits are reinterpreted. @@ -861,31 +791,16 @@ def test_r1_mtp_expert_relabeling_is_bitwise() -> None: torch.cuda.empty_cache() -def test_r1_mtp_reference_discriminates() -> None: - # The 8/4 ulp gates must reject wrong computations at this geometry too, - # not just at the small shapes. Measured against the correct result's 2.09 - # ulp: 284 / 197 ulp for the swapped halves and 997 / 212 for a dropped - # expert, i.e. 36x and 125x the element gate, 49x and 53x the RMS gate. - w31, w2 = _make_r1_weights(R1_E_LOCAL, seed=555) - x, ids, scales = _make_r1_routing(512, R1_E_LOCAL, seed=556) - y = fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, [])[0] - _assert_moe_close(y, _ref_moe(x, ids, scales, w31, w2)) - swapped = torch.cat([w31[:, R1_I:], w31[:, :R1_I]], dim=1).contiguous() - _rejects(y, _ref_moe(x, ids, scales, swapped, w2), "gate/up halves swapped") - del swapped - torch.cuda.empty_cache() - dropped = scales.clone() - dropped[:, 0] = 0.0 - _rejects(y, _ref_moe(x, ids, dropped, w31, w2), "one expert dropped") - del w31, w2, x, ids, scales, y, dropped - torch.cuda.empty_cache() - - def test_r1_mtp_autotuned_tactics() -> None: # Serving runs on the hot side of the tuner, which a cold receipt never - # sees. One autotune() pass at this geometry fills 2 tunable GEMMs x 14 - # power-of-2 token buckets, moves the bits, and still lands inside the - # gates; clearing the cache restores the cold bits exactly. + # sees. One autotune() pass at this geometry moves the bits and still + # lands inside the gates; clearing the cache restores the cold bits + # exactly. + # + # How many entries that pass leaves in the cache is the tuner's business, + # not this op's -- `misc/test_autotuner.py::test_bucket_mapping` owns the + # bucketing rule. Asserting the count here would fail on a tuner change + # that says nothing about fused_moe. tuner = AutoTuner.get() tuner.clear_cache() w31, w2 = _make_r1_weights(R1_E_LOCAL, seed=1234) @@ -899,7 +814,7 @@ def test_r1_mtp_autotuned_tactics() -> None: x, ids, scales = cases[8192][:3] with autotune(): fused_moe(x, ids, scales, w31, None, w2, None, torch.bfloat16, []) - assert len(tuner.profiling_cache) == 28, len(tuner.profiling_cache) + assert tuner.profiling_cache, "the autotune pass recorded nothing" for key, (_, tactic, _) in tuner.profiling_cache.cache.items(): assert tactic >= 0, f"{key} kept the fallback tactic after tuning" moved = 0 diff --git a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_mxe4m3_mxe2m1_block_scale_moe_runner.py b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_mxe4m3_mxe2m1_block_scale_moe_runner.py index f3c98329c999..da117aa8c9de 100644 --- a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_mxe4m3_mxe2m1_block_scale_moe_runner.py +++ b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_mxe4m3_mxe2m1_block_scale_moe_runner.py @@ -417,9 +417,9 @@ def _assert_moe_close(y: torch.Tensor, ref: torch.Tensor) -> None: the row's largest magnitude and the aggregate gate is 4 ulp of relative RMS. Worst values measured over every configuration covered here: 2.0 ulp element-wise (128 experts, 8192 tokens, H=I=2880) and 0.87 ulp RMS. Both - gates bite — `test_reference_discriminates` shows a reference that skips - the FC1-output requantization landing at 19.6 / 10.6 ulp and a - gate/up-swapped or unshuffled operand at 100+ ulp. + gates were sized against deliberately wrong computations when they were + set: a reference that skips the FC1-output requantization lands at + 19.6 / 10.6 ulp, and a gate/up-swapped or unshuffled operand at 100+ ulp. """ assert y.dtype == ref.dtype == torch.bfloat16, (y.dtype, ref.dtype) assert y.shape == ref.shape, (y.shape, ref.shape) @@ -1128,56 +1128,6 @@ def test_inert_arguments(): print(" test_inert_arguments OK") -def test_reference_discriminates(): - """The tolerance gate rejects a swapped-half or unshuffled operand.""" - num_experts, hidden, inter, top_k, num_tokens = 8, 512, 256, 4, 12 - args, ref = _build(num_experts, hidden, inter, seed=22) - gen = ref["gen"] - data, sf, xv = _rand_mxfp8(num_tokens, hidden, ref["h1_pad"], gen) - ids, wts = _routing(num_tokens, num_experts, top_k, gen) - alpha, beta, limit = _swiglu_params(num_experts) - kw = dict( - topk_ids=ids, - topk_weights=wts, - gemm1_alpha=alpha, - gemm1_beta=beta, - gemm1_clamp_limit=limit, - ) - out = _call(data, sf, args, num_experts, top_k, **kw) - swapped = _ref_moe(xv, ids, wts, ref, alpha, beta, limit, swap_gate_up=True).to(torch.bfloat16) - elt, rms = _dev(out, swapped) - assert elt > 50.0 and rms > 20.0, f"gate/up swap not detected: {elt:.2f}/{rms:.2f} ulp" - - # a reference that skips the FC1-output requantization is also rejected - unq = _ref_moe(xv, ids, wts, ref, alpha, beta, limit, quantize_intermediate=False).to( - torch.bfloat16 - ) - elt, rms = _dev(out, unq) - assert elt > 8.0 and rms > 4.0, ( - f"unquantized intermediate not detected: {elt:.2f}/{rms:.2f} ulp" - ) - - # feeding the un-permuted (merely padded and concatenated) FC1 operand - up_c, up_s, _ = ref["up"] - gt_c, gt_s, _ = ref["gate"] - i_pad = args["intermediate_size"] - h1 = ref["h1_pad"] - up_p = (up_c[..., 0::2] | (up_c[..., 1::2] << 4)).contiguous() - gt_p = (gt_c[..., 0::2] | (gt_c[..., 1::2] << 4)).contiguous() - raw = dict(args) - raw["gemm1_weights"] = torch.cat( - [_pad3(up_p, i_pad, h1 // 2), _pad3(gt_p, i_pad, h1 // 2)], dim=1 - ).contiguous() - raw["gemm1_weights_scale"] = torch.cat( - [_pad3(up_s, i_pad, h1 // SV), _pad3(gt_s, i_pad, h1 // SV)], dim=1 - ).contiguous() - bad = _call(data, sf, raw, num_experts, top_k, **kw) - exp = _ref_moe(xv, ids, wts, ref, alpha, beta, limit).to(torch.bfloat16) - elt, rms = _dev(bad, exp) - assert elt > 50.0 and rms > 20.0, f"unshuffled FC1 not detected: {elt:.2f}/{rms:.2f} ulp" - print(" test_reference_discriminates OK") - - def test_rejects_unsupported(): """Domains the contract declares unsupported are rejected loudly.""" num_experts, hidden, inter, top_k, num_tokens = 8, 512, 256, 4, 6 diff --git a/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_flashinfer_fused_add_rmsnorm.py b/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_flashinfer_fused_add_rmsnorm.py index 55f873117e37..a6742c988aec 100644 --- a/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_flashinfer_fused_add_rmsnorm.py +++ b/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_flashinfer_fused_add_rmsnorm.py @@ -67,21 +67,3 @@ def test_bf16_strided_rows() -> None: _check(x, r, w, 1e-6) # mutation must land in the parent buffers' left halves only assert torch.equal(torch.cat([x_buf[:, 4096:], r_buf[:, 4096:]]), right_before) - - -def test_fp16_2d() -> None: - torch.manual_seed(3) - for num_tokens, hidden in [(2, 4096), (1024, 2048)]: - x = torch.randn(num_tokens, hidden, dtype=torch.float16, device="cuda") - r = torch.randn(num_tokens, hidden, dtype=torch.float16, device="cuda") - w = torch.randn(hidden, dtype=torch.float16, device="cuda") - _check(x, r, w, 1e-6) - - -def test_fp32_2d() -> None: - torch.manual_seed(4) - for num_tokens, hidden in [(2, 4096), (1024, 2048)]: - x = torch.randn(num_tokens, hidden, dtype=torch.float32, device="cuda") - r = torch.randn(num_tokens, hidden, dtype=torch.float32, device="cuda") - w = torch.randn(hidden, dtype=torch.float32, device="cuda") - _check(x, r, w, 1e-6) diff --git a/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_flashinfer_rmsnorm.py b/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_flashinfer_rmsnorm.py index 2c2c61d31dff..bc2a1925b89c 100644 --- a/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_flashinfer_rmsnorm.py +++ b/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_flashinfer_rmsnorm.py @@ -56,19 +56,3 @@ def test_bf16_strided_rows() -> None: assert not x.is_contiguous() and x.stride(-1) == 1 w = torch.randn(4096, dtype=torch.bfloat16, device="cuda") _check(x, w, 1e-6) - - -def test_fp16_2d() -> None: - torch.manual_seed(4) - for num_tokens, hidden in [(2, 4096), (1024, 2048)]: - x = torch.randn(num_tokens, hidden, dtype=torch.float16, device="cuda") - w = torch.randn(hidden, dtype=torch.float16, device="cuda") - _check(x, w, 1e-6) - - -def test_fp32_2d() -> None: - torch.manual_seed(5) - for num_tokens, hidden in [(2, 4096), (1024, 2048)]: - x = torch.randn(num_tokens, hidden, dtype=torch.float32, device="cuda") - w = torch.randn(hidden, dtype=torch.float32, device="cuda") - _check(x, w, 1e-6) diff --git a/tests/unittest/_torch/modeling_v2/quantization/test_modeling_v2_fp4_quantize.py b/tests/unittest/_torch/modeling_v2/quantization/test_modeling_v2_fp4_quantize.py index 02687d657f22..d49c0b662686 100644 --- a/tests/unittest/_torch/modeling_v2/quantization/test_modeling_v2_fp4_quantize.py +++ b/tests/unittest/_torch/modeling_v2/quantization/test_modeling_v2_fp4_quantize.py @@ -203,18 +203,6 @@ def test_swizzled_layout_bf16() -> None: assert (sf_sw[rest] == 0).all(), "swizzled padding must be zero" -def test_fp16_input() -> None: - torch.manual_seed(2) - for t in (1, 512, 1024): - x = torch.randn(t, 2560, dtype=torch.float16, device=DEV) - gs, gsf = _global_scale(x) - data, sf = fp4_quantize(x, gs, VEC, False, False) - ref_data, ref_sf, scaled = _ref(x, gsf) - assert data.shape == (t, 1280) - torch.testing.assert_close(sf.view(t, 160), ref_sf) - _assert_data(data, ref_data, scaled) - - def test_3d_input_collapses_leading_dims() -> None: """data keeps the leading dims; the scale buffer is laid out for their product.""" torch.manual_seed(3) @@ -555,72 +543,6 @@ def test_r1_one_kernel_at_every_width() -> None: del x -def test_r1_reference_check_is_discriminating() -> None: - """The tie carve-out tolerates exactly one adjacent code at a near-tie, nothing else. - - Runs the wrong variants through the same comparison the R1 sweep uses, so - the sweep's clean result is a measurement rather than a blind spot. - """ - torch.manual_seed(9) - x = torch.randn(1024, 7168, dtype=torch.bfloat16, device=DEV) - gs, gsf = _global_scale(x) - data, sf = fp4_quantize(x, gs, VEC, False, False) - ref_data, ref_sf, scaled = _ref(x, gsf) - near = _near_tie(scaled) - # the correct result passes, and the carve-out is a small fraction of the tensor - n_tie, _, total = _assert_data(data, ref_data, scaled) - assert 0 < n_tie / total < 0.02, n_tie / total - codes = _unpack(data) - # perturb only elements the kernel already got right, and only where the - # magnitude code leaves room to move two steps up - agree = (codes == _unpack(ref_data)) & ((codes & 7) <= 5) - inside = (near & agree).nonzero() - outside = (~near & agree & ((codes & 7) > 0)).nonzero() - assert len(inside) > 0 and len(outside) > 0 - - def bumped(row: int, col: int, delta: int, flip_sign: bool = False) -> torch.Tensor: - c = codes.clone() - code = int(c[row, col]) - mag = (code & 7) + delta - assert 0 <= mag <= 7, mag - c[row, col] = mag | ((code & 8) ^ (8 if flip_sign else 0)) - return c[:, 0::2] | (c[:, 1::2] << 4) - - def raises(fn) -> bool: - try: - fn() - except AssertionError: - return True - return False - - ro, co = (int(v) for v in outside[len(outside) // 2]) - ri, ci = (int(v) for v in inside[len(inside) // 2]) - assert raises(lambda: _assert_data(bumped(ro, co, 1), ref_data, scaled)), ( - "one code away from a tie must be caught" - ) - assert raises(lambda: _assert_data(bumped(ri, ci, 2), ref_data, scaled)), ( - "two codes at a tie must be caught" - ) - assert raises(lambda: _assert_data(bumped(ri, ci, 0, flip_sign=True), ref_data, scaled)), ( - "a sign flip at a tie must be caught" - ) - # the documented blind spot: one adjacent code at a near-tie is accepted, - # and shows up only as one extra mismatch - _, mis_before, _ = _assert_data(data, ref_data, scaled) - _, mis_after, _ = _assert_data(bumped(ri, ci, 1), ref_data, scaled) - assert mis_after == mis_before + 1, (mis_before, mis_after) - # and the scale bytes carry no such carve-out: one code is caught - bad_sf = sf.clone() - bad_sf[0] = (int(bad_sf[0]) + 1) & 0xFF - assert raises(lambda: torch.testing.assert_close(bad_sf.view(1024, 448), ref_sf)), ( - "one scale code must be caught" - ) - # a swizzled buffer handed over as if it were linear is a different byte - # string even when the two have the same length (rows a multiple of 128) - _, sf_sw = fp4_quantize(x, gs, VEC, False, True) - assert sf_sw.shape == sf.shape and not torch.equal(sf_sw, sf) - - def test_r1_ties_track_the_global_scale_not_the_shape() -> None: """The near-tie population is a property of the global scale, not of M or K. diff --git a/tests/unittest/_torch/modeling_v2/quantization/test_modeling_v2_mxfp8_quantize.py b/tests/unittest/_torch/modeling_v2/quantization/test_modeling_v2_mxfp8_quantize.py index a7d638ee37e4..2b804fe39047 100644 --- a/tests/unittest/_torch/modeling_v2/quantization/test_modeling_v2_mxfp8_quantize.py +++ b/tests/unittest/_torch/modeling_v2/quantization/test_modeling_v2_mxfp8_quantize.py @@ -117,17 +117,6 @@ def test_alignment_padding() -> None: assert (padded_sf.view(64, 96)[:, 90:] == 0).all() -def test_fp16_input() -> None: - torch.manual_seed(3) - for t in (1, 512): - x = torch.randn(t, 2880, dtype=torch.float16, device="cuda") - data, sf = mxfp8_quantize(x, False, 512) - ref_data, ref_sf = _ref(x, 3072) - assert data.dtype == torch.float8_e4m3fn - assert torch.equal(data.view(torch.uint8), ref_data.view(torch.uint8)) - assert torch.equal(sf.view(t, 96), ref_sf) - - def test_3d_input_collapses_leading_dims() -> None: torch.manual_seed(4) x = torch.randn(2, 5, 2880, dtype=torch.bfloat16, device="cuda") From f9b71e8eba1cb33be8afd5be2d72d01805076d00 Mon Sep 17 00:00:00 2001 From: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> Date: Wed, 16 Sep 2026 20:30:29 -0700 Subject: [PATCH 12/19] [TRTLLM-16304][chore] Move modeling_v2 under _experimental Review asked for the subtree to sit under `_experimental/` so the path says what the API stability tests already imply: nothing here is covered by them, and it may change shape or be removed without a deprecation cycle. The new `_experimental/__init__.py` re-exports nothing, so reaching a subtree has to be written at the import site. Source package only. The unit tests stay at `tests/unittest/_torch/modeling_v2/`: `_experimental` marks importable surface, and mirroring the move there would churn the l0 entry strings, the CBTS match patterns and the rank-job `parents[]` index for no signal. Mechanical apart from three places the new path forced. `_experimental` sorts ahead of `attention`, so the two targets' import blocks reorder; the longer module path pushes several of those imports past 100 columns, so they wrap; and llm.py's lazy import crosses the 80 that file is held to, so it splits. CODEOWNERS and the CBTS rule's source prefix follow the path. Signed-off-by: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> --- .github/CODEOWNERS | 6 +- jenkins/scripts/cbts/README.md | 4 +- jenkins/scripts/cbts/rules/README.md | 4 +- .../scripts/cbts/rules/modeling_v2_rule.py | 4 +- tensorrt_llm/_torch/_experimental/__init__.py | 15 +++++ .../{ => _experimental}/modeling_v2/README.md | 4 +- .../modeling_v2/__init__.py | 0 .../modeling_v2/_router_index.py | 2 +- .../modeling_v2/catalog/__init__.py | 0 .../catalog/activation/__init__.py | 0 .../activation/flashinfer_silu_and_mul.md | 0 .../activation/flashinfer_silu_and_mul.py | 0 .../modeling_v2/catalog/attention/__init__.py | 0 .../catalog/attention/fused_qk_norm_rope.md | 0 .../catalog/attention/fused_qk_norm_rope.py | 0 .../attention/load_paged_kv_cache_for_mla.md | 0 .../attention/load_paged_kv_cache_for_mla.py | 0 .../mla_rope_append_paged_kv_assign_q.md | 0 .../mla_rope_append_paged_kv_assign_q.py | 0 .../catalog/attention/mla_rope_generation.md | 0 .../catalog/attention/mla_rope_generation.py | 0 .../catalog/attention/thop_attention.md | 0 .../catalog/attention/thop_attention.py | 0 .../modeling_v2/catalog/comm/__init__.py | 0 .../modeling_v2/catalog/comm/allgather.md | 0 .../modeling_v2/catalog/comm/allgather.py | 0 .../modeling_v2/catalog/comm/reducescatter.md | 0 .../modeling_v2/catalog/comm/reducescatter.py | 0 .../modeling_v2/catalog/gemm/__init__.py | 0 .../modeling_v2/catalog/gemm/bmm_out.md | 0 .../modeling_v2/catalog/gemm/bmm_out.py | 0 .../modeling_v2/catalog/gemm/cublas_mm.md | 0 .../modeling_v2/catalog/gemm/cublas_mm.py | 0 .../modeling_v2/catalog/gemm/nvfp4_gemm.md | 0 .../modeling_v2/catalog/gemm/nvfp4_gemm.py | 0 .../modeling_v2/catalog/index.yaml | 0 .../modeling_v2/catalog/moe/__init__.py | 0 .../catalog/moe/fp4_block_scale_moe_runner.md | 0 .../catalog/moe/fp4_block_scale_moe_runner.py | 0 .../modeling_v2/catalog/moe/fused_moe.md | 0 .../modeling_v2/catalog/moe/fused_moe.py | 0 .../mxe4m3_mxe2m1_block_scale_moe_runner.md | 0 .../mxe4m3_mxe2m1_block_scale_moe_runner.py | 0 .../modeling_v2/catalog/moe/noaux_tc_op.md | 0 .../modeling_v2/catalog/moe/noaux_tc_op.py | 0 .../modeling_v2/catalog/norm/__init__.py | 0 .../norm/flashinfer_fused_add_rmsnorm.md | 0 .../norm/flashinfer_fused_add_rmsnorm.py | 0 .../catalog/norm/flashinfer_rmsnorm.md | 0 .../catalog/norm/flashinfer_rmsnorm.py | 0 .../catalog/quantization/__init__.py | 0 .../catalog/quantization/fp4_quantize.md | 0 .../catalog/quantization/fp4_quantize.py | 0 .../catalog/quantization/mxfp8_quantize.md | 0 .../catalog/quantization/mxfp8_quantize.py | 0 .../modeling_v2/catalog/torch/__init__.py | 0 .../modeling_v2/catalog/torch/add.py | 0 .../modeling_v2/catalog/torch/concat.py | 0 .../modeling_v2/catalog/torch/copy_.py | 0 .../modeling_v2/catalog/torch/embedding.py | 0 .../modeling_v2/catalog/torch/empty.py | 0 .../modeling_v2/catalog/torch/expand.py | 0 .../modeling_v2/catalog/torch/pad.py | 0 .../modeling_v2/catalog/torch/reshape.py | 0 .../modeling_v2/catalog/torch/split.py | 0 .../modeling_v2/catalog/torch/transpose.py | 0 .../modeling_v2/catalog/torch/view_dtype.py | 0 .../modeling_v2/explain.py | 4 +- .../modeling_v2/models/__init__.py | 0 .../models/deepseek_v3/__init__.py | 0 .../modeling_v2/models/deepseek_v3/routing.py | 0 .../models/deepseek_v3/targets/__init__.py | 0 .../targets/r1_0528_nvfp4/__init__.py | 0 .../targets/r1_0528_nvfp4/sm_103/__init__.py | 0 .../r1_0528_nvfp4/sm_103/dep4/__init__.py | 0 .../r1_0528_nvfp4/sm_103/dep4/modeling.py | 66 ++++++++++--------- .../r1_0528_nvfp4/sm_103/dep4/weights.py | 0 .../modeling_v2/models/gpt_oss/__init__.py | 0 .../modeling_v2/models/gpt_oss/routing.py | 0 .../models/gpt_oss/targets/__init__.py | 0 .../gpt_oss/targets/gpt_oss_120b/__init__.py | 0 .../targets/gpt_oss_120b/sm_103/__init__.py | 0 .../gpt_oss_120b/sm_103/tp1/__init__.py | 0 .../gpt_oss_120b/sm_103/tp1/modeling.py | 34 ++++++---- .../gpt_oss_120b/sm_103/tp1/weights.py | 0 tensorrt_llm/_torch/models/modeling_auto.py | 2 +- tensorrt_llm/llmapi/llm.py | 3 +- .../accuracy/test_modeling_v2_deepseek_v3.py | 2 +- .../defs/accuracy/test_modeling_v2_gpt_oss.py | 2 +- ...est_modeling_v2_flashinfer_silu_and_mul.py | 2 +- .../test_modeling_v2_fused_qk_norm_rope.py | 4 +- ...modeling_v2_load_paged_kv_cache_for_mla.py | 6 +- ...ng_v2_mla_rope_append_paged_kv_assign_q.py | 6 +- .../test_modeling_v2_mla_rope_generation.py | 6 +- .../test_modeling_v2_thop_attention.py | 4 +- .../modeling_v2/comm/_allgather_op_matrix.py | 2 +- .../comm/_reducescatter_op_matrix.py | 4 +- .../gemm/test_modeling_v2_bmm_out.py | 2 +- .../gemm/test_modeling_v2_cublas_mm.py | 2 +- .../gemm/test_modeling_v2_nvfp4_gemm.py | 2 +- ..._modeling_v2_fp4_block_scale_moe_runner.py | 4 +- .../moe/test_modeling_v2_fused_moe.py | 2 +- ...v2_mxe4m3_mxe2m1_block_scale_moe_runner.py | 2 +- .../moe/test_modeling_v2_noaux_tc_op.py | 2 +- ...odeling_v2_flashinfer_fused_add_rmsnorm.py | 2 +- .../test_modeling_v2_flashinfer_rmsnorm.py | 4 +- .../test_modeling_v2_fp4_quantize.py | 4 +- .../test_modeling_v2_mxfp8_quantize.py | 4 +- .../modeling_v2/test_modeling_v2_claims.py | 7 +- .../test_modeling_v2_no_stale_claims.py | 2 +- .../modeling_v2/test_modeling_v2_routing.py | 4 +- .../test_modeling_v2_target_contract.py | 7 +- 112 files changed, 141 insertions(+), 95 deletions(-) create mode 100644 tensorrt_llm/_torch/_experimental/__init__.py rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/README.md (98%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/__init__.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/_router_index.py (99%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/__init__.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/activation/__init__.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/activation/flashinfer_silu_and_mul.md (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/activation/flashinfer_silu_and_mul.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/attention/__init__.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/attention/fused_qk_norm_rope.md (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/attention/fused_qk_norm_rope.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/attention/load_paged_kv_cache_for_mla.md (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/attention/load_paged_kv_cache_for_mla.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/attention/mla_rope_append_paged_kv_assign_q.md (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/attention/mla_rope_append_paged_kv_assign_q.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/attention/mla_rope_generation.md (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/attention/mla_rope_generation.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/attention/thop_attention.md (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/attention/thop_attention.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/comm/__init__.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/comm/allgather.md (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/comm/allgather.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/comm/reducescatter.md (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/comm/reducescatter.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/gemm/__init__.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/gemm/bmm_out.md (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/gemm/bmm_out.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/gemm/cublas_mm.md (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/gemm/cublas_mm.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/gemm/nvfp4_gemm.md (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/gemm/nvfp4_gemm.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/index.yaml (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/moe/__init__.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/moe/fp4_block_scale_moe_runner.md (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/moe/fp4_block_scale_moe_runner.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/moe/fused_moe.md (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/moe/fused_moe.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/moe/noaux_tc_op.md (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/moe/noaux_tc_op.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/norm/__init__.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/norm/flashinfer_fused_add_rmsnorm.md (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/norm/flashinfer_fused_add_rmsnorm.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/norm/flashinfer_rmsnorm.md (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/norm/flashinfer_rmsnorm.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/quantization/__init__.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/quantization/fp4_quantize.md (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/quantization/fp4_quantize.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/quantization/mxfp8_quantize.md (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/quantization/mxfp8_quantize.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/torch/__init__.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/torch/add.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/torch/concat.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/torch/copy_.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/torch/embedding.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/torch/empty.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/torch/expand.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/torch/pad.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/torch/reshape.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/torch/split.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/torch/transpose.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/catalog/torch/view_dtype.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/explain.py (96%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/models/__init__.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/models/deepseek_v3/__init__.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/models/deepseek_v3/routing.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/models/deepseek_v3/targets/__init__.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/__init__.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/__init__.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/__init__.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py (97%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/weights.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/models/gpt_oss/__init__.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/models/gpt_oss/routing.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/models/gpt_oss/targets/__init__.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/__init__.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/__init__.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/__init__.py (100%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py (96%) rename tensorrt_llm/_torch/{ => _experimental}/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/weights.py (100%) diff --git a/.github/CODEOWNERS b/.github/CODEOWNERS index 4bad576cf8d3..7faf5a6e8a3c 100644 --- a/.github/CODEOWNERS +++ b/.github/CODEOWNERS @@ -389,14 +389,14 @@ # Overrides the /tensorrt_llm/_torch runtime-devs rule above for this subtree. # Individual handles rather than a team, like SCAFFOLDING: this is one bounded # experiment with named owners, not a standing domain. -/tensorrt_llm/_torch/modeling_v2 @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju +/tensorrt_llm/_torch/_experimental/modeling_v2 @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju /tests/unittest/_torch/modeling_v2 @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju /tests/integration/defs/accuracy/test_modeling_v2_deepseek_v3.py @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju /tests/integration/defs/accuracy/test_modeling_v2_gpt_oss.py @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju # The two catalog categories whose contracts state kernel behaviour the # attention and MoE owners are the authority on; co-owned rather than reassigned. -/tensorrt_llm/_torch/modeling_v2/catalog/attention @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju @xxi-nv @yuxianq -/tensorrt_llm/_torch/modeling_v2/catalog/moe @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju @xxi-nv @yuxianq +/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju @xxi-nv @yuxianq +/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe @tianyuxbear @Wanli-Jiang @WeiHaocheng @litaotju @xxi-nv @yuxianq ## TensorRT-LLM LLM Disaggregated /examples/disaggregated @NVIDIA/trt-llm-disagg-devs @NVIDIA/trt-llm-doc-owners diff --git a/jenkins/scripts/cbts/README.md b/jenkins/scripts/cbts/README.md index d5e751bb2641..f5de608724c7 100644 --- a/jenkins/scripts/cbts/README.md +++ b/jenkins/scripts/cbts/README.md @@ -46,7 +46,7 @@ Ten rules, registered in `main.py::RULE_CLASSES`: | `TestListRule` | `testlistonly` | `tests/integration/test_lists/test-db/*.yml` | | `VisualGenRule` | `visualgenonly` | `examples/visual_gen/**`, `scripts/visualgen_eval/**`, `tensorrt_llm/_torch/visual_gen/**`, `tensorrt_llm/media/**`, `tensorrt_llm/visual_gen/**` (excl. `.md`; reference images such as `cat_piano.png` ARE test fixtures and stay claimed; outward-facing files force fallback) | | `SpecDecRule` | `specdeconly` | `tensorrt_llm/_torch/speculative/**`, `tensorrt_llm/models/{eagle,medusa,redrafter}/**`, `examples/{eagle,medusa,redrafter,draft_target_model,ngram}/**`, `examples/llm-api/llm_speculative_decoding.py` (excl. `.md`; other suffixes incl. images kept as potential test fixtures) | -| `ModelingV2Rule` | `modelingv2only` | `tensorrt_llm/_torch/modeling_v2/**` (excl. `.md`) | +| `ModelingV2Rule` | `modelingv2only` | `tensorrt_llm/_torch/_experimental/modeling_v2/**` (excl. `.md`) | | `AgentFlowRule` | `agentflowonly` | `agent-flow/**` (excl. `.md`) | | `OpenEngineRule` | `openengineonly` | `tensorrt_llm/grpc/openengine/**` (excl. `.md`) | | `OutOfScopeRule` | `noop` | QA / dev test lists, `.test_durations`, `microbenchmarks/`, `**/*.md` (image suffixes intentionally not claimed — image fixtures cannot be distinguished from doc diagrams by location, so image edits fall back to baseline) | @@ -62,7 +62,7 @@ See `rules/README.md` for per-rule logic. | `testlistonly` | `TestListRule` fired solo: PR only adds entries under `tests/integration/test_lists/test-db/*.yml`. | | `visualgenonly` | `VisualGenRule` fired solo: PR only touches VisualGen internal source paths (`examples/visual_gen/**`, `scripts/visualgen_eval/**`, `tensorrt_llm/_torch/visual_gen/**`; excl. `.md`; image fixtures like `cat_piano.png` are claimed). Narrows to blocks containing VG test entries. Outward-facing files under `tensorrt_llm/visual_gen/**` and `tensorrt_llm/media/**` (eagerly imported by `trtllm-serve`) force `null` fallback. | | `specdeconly` | `SpecDecRule` fired solo: PR only touches speculative-decoding source paths (`tensorrt_llm/_torch/speculative/**`, `tensorrt_llm/models/{eagle,medusa,redrafter}/**`, `examples/{eagle,medusa,redrafter,draft_target_model,ngram}/**`, `examples/llm-api/llm_speculative_decoding.py`; excl. `.md`). Narrows to blocks containing spec-dec test entries (eagle / medusa / redrafter / ngram / draft-target-model / MTP). | -| `modelingv2only` | `ModelingV2Rule` fired solo: PR only touches the modeling_v2 subtree (`tensorrt_llm/_torch/modeling_v2/**`; excl. `.md`, which is a fifth of the subtree — every catalog entry carries a contract document). Narrows to blocks containing modeling_v2 test entries (`unittest/_torch/modeling_v2/`, `test_modeling_v2_*`). No outward-facing fallback is needed: nothing imports the subtree unless `TRTLLM_MODELING_V2` is set, and its one caller outside the subtree (`_torch/models/modeling_auto.py`) is left unclaimed, so touching the shared resolver falls back to baseline. | +| `modelingv2only` | `ModelingV2Rule` fired solo: PR only touches the modeling_v2 subtree (`tensorrt_llm/_torch/_experimental/modeling_v2/**`; excl. `.md`, which is a fifth of the subtree — every catalog entry carries a contract document). Narrows to blocks containing modeling_v2 test entries (`unittest/_torch/modeling_v2/`, `test_modeling_v2_*`). No outward-facing fallback is needed: nothing imports the subtree unless `TRTLLM_MODELING_V2` is set, and its one caller outside the subtree (`_torch/models/modeling_auto.py`) is left unclaimed, so touching the shared resolver falls back to baseline. | | `agentflowonly` | `AgentFlowRule` fired solo: PR only touches `agent-flow/**` source or test files (excl. `.md`). Runs `CPU-AgentFlow-UnitTest`. | | `openengineonly` | `OpenEngineRule` fired solo: PR only touches `tensorrt_llm/grpc/openengine/**` source files (excl. `.md`). Narrows to the registered OpenEngine unit tests: the stub-based ones on the always-run `CPU-Generic-*` stages, plus `test_capability_conformance.py` on `A10-PyTorch-*`, which needs a GPU. | | `testsonly` | Multiple rules from the testsonly family fired (`waiveonly`, `testdefonly`, `testlistonly`, `visualgenonly`, `specdeconly`, `modelingv2only`, `agentflowonly`, `openengineonly`); their narrows union. | diff --git a/jenkins/scripts/cbts/rules/README.md b/jenkins/scripts/cbts/rules/README.md index d2038abe5bbf..5870dd140427 100644 --- a/jenkins/scripts/cbts/rules/README.md +++ b/jenkins/scripts/cbts/rules/README.md @@ -13,7 +13,7 @@ for the overall CBTS architecture. | `test_list_rule.py` | `TestListRule` | `testlistonly` | `tests/integration/test_lists/test-db/*.yml` | | `visual_gen_rule.py` | `VisualGenRule` | `visualgenonly` | `examples/visual_gen/**`, `scripts/visualgen_eval/**`, `tensorrt_llm/_torch/visual_gen/**`, `tensorrt_llm/media/**`, `tensorrt_llm/visual_gen/**` (each excl. `.md`) | | `spec_dec_rule.py` | `SpecDecRule` | `specdeconly` | `tensorrt_llm/_torch/speculative/**`, `tensorrt_llm/models/{eagle,medusa,redrafter}/**`, `examples/{eagle,medusa,redrafter,draft_target_model,ngram}/**`, `examples/llm-api/llm_speculative_decoding.py` (each excl. `.md`) | -| `modeling_v2_rule.py` | `ModelingV2Rule` | `modelingv2only` | `tensorrt_llm/_torch/modeling_v2/**` (excl. `.md`) | +| `modeling_v2_rule.py` | `ModelingV2Rule` | `modelingv2only` | `tensorrt_llm/_torch/_experimental/modeling_v2/**` (excl. `.md`) | | `agent_flow_rule.py` | `AgentFlowRule` | `agentflowonly` | `agent-flow/**` (excl. `.md`) → the single `CPU-AgentFlow-UnitTest` stage; not test-db-driven | | `openengine_rule.py` | `OpenEngineRule` | `openengineonly` | `tensorrt_llm/grpc/openengine/**` (excl. `.md`) → the `l0_cpu` block containing `unittest/grpc/openengine/` | | `out_of_scope_rule.py` | `OutOfScopeRule` | `noop` | `tests/integration/test_lists/{qa,dev}/**`, `tests/integration/defs/.test_durations*`, `tests/microbenchmarks/**`, `**/*.md` (image suffixes intentionally not claimed — fall back to baseline since fixtures and doc diagrams are indistinguishable by location) | @@ -258,7 +258,7 @@ Outcomes: ## ModelingV2Rule Path-only rule. Claims non-documentation source changes under -`tensorrt_llm/_torch/modeling_v2/`, the self-contained second modeling +`tensorrt_llm/_torch/_experimental/modeling_v2/`, the self-contained second modeling path (one flat forward per checkpoint/arch/parallel triple, assembled from a catalog of op wrappers). diff --git a/jenkins/scripts/cbts/rules/modeling_v2_rule.py b/jenkins/scripts/cbts/rules/modeling_v2_rule.py index 68e05d963be7..1702c16f1416 100644 --- a/jenkins/scripts/cbts/rules/modeling_v2_rule.py +++ b/jenkins/scripts/cbts/rules/modeling_v2_rule.py @@ -14,7 +14,7 @@ """ModelingV2Rule — narrows CI when the modeling_v2 subtree changes. modeling_v2 is a second modeling path living entirely under -`tensorrt_llm/_torch/modeling_v2/`: one self-contained forward per +`tensorrt_llm/_torch/_experimental/modeling_v2/`: one self-contained forward per (checkpoint, GPU arch, parallel topology), assembled from a catalog of op wrappers. @@ -66,7 +66,7 @@ # Source-path prefixes the rule may claim. Tests under tests/** are left # to TestsDefRule; the two scopes combine via _TESTSONLY_FAMILY. -_MV2_SRC_PREFIXES: tuple[str, ...] = ("tensorrt_llm/_torch/modeling_v2/",) +_MV2_SRC_PREFIXES: tuple[str, ...] = ("tensorrt_llm/_torch/_experimental/modeling_v2/",) # Substrings that mark a test entry as modeling_v2. Both are unambiguous: # - "unittest/_torch/modeling_v2/" → the op-level catalog matrix, taken diff --git a/tensorrt_llm/_torch/_experimental/__init__.py b/tensorrt_llm/_torch/_experimental/__init__.py new file mode 100644 index 000000000000..5222f5d23a8f --- /dev/null +++ b/tensorrt_llm/_torch/_experimental/__init__.py @@ -0,0 +1,15 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Subpackages here are not covered by the API stability tests. + +Anything under this package may change shape or be removed without a +deprecation cycle. Import it from outside `tensorrt_llm` at your own risk; +in-tree callers should reach it through a switch that stays off by default, +the way `modeling_v2` is reached through `TRTLLM_MODELING_V2`. + +Nothing is re-exported here on purpose, so this module itself pulls in no +subpackage. That is not the same as saying a subtree here is unreachable at +startup -- `modeling_v2` is imported eagerly by `_torch/models/modeling_auto.py` +-- only that reaching one has to be written down at the import site rather than +happening as a side effect of this package. +""" diff --git a/tensorrt_llm/_torch/modeling_v2/README.md b/tensorrt_llm/_torch/_experimental/modeling_v2/README.md similarity index 98% rename from tensorrt_llm/_torch/modeling_v2/README.md rename to tensorrt_llm/_torch/_experimental/modeling_v2/README.md index 9b20a9a6d701..0019e050785f 100644 --- a/tensorrt_llm/_torch/modeling_v2/README.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/README.md @@ -5,7 +5,7 @@ built-in model zoo rather than inside it. Where `_torch/models/modeling_deepseekv3.py` is one class serving V3, V3-Lite, R1 and V3.2 across every GPU generation and parallel topology, -`_torch/modeling_v2/models/deepseek_v3/` is one flat forward per (checkpoint, +`_torch/_experimental/modeling_v2/models/deepseek_v3/` is one flat forward per (checkpoint, GPU architecture, parallel topology) triple — assembled only from `catalog/` entries, sharing nothing with its siblings, and trusted through accuracy gates instead of shared abstractions. The one-to-one correspondence between @@ -68,7 +68,7 @@ which also exists only as a `_resolve_class` rewrite. To ask why a configuration landed where it did: ``` -python -m tensorrt_llm._torch.modeling_v2.explain \ +python -m tensorrt_llm._torch._experimental.modeling_v2.explain \ --model /path/to/DeepSeek-R1-0528-NVFP4 --tp 4 --ep 4 --attention-dp ``` diff --git a/tensorrt_llm/_torch/modeling_v2/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/__init__.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/__init__.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/__init__.py diff --git a/tensorrt_llm/_torch/modeling_v2/_router_index.py b/tensorrt_llm/_torch/_experimental/modeling_v2/_router_index.py similarity index 99% rename from tensorrt_llm/_torch/modeling_v2/_router_index.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/_router_index.py index e9437214a31a..3a4094bb1632 100644 --- a/tensorrt_llm/_torch/modeling_v2/_router_index.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/_router_index.py @@ -39,7 +39,7 @@ if TYPE_CHECKING: from tensorrt_llm._torch.model_config import ModelConfig -_PACKAGE = "tensorrt_llm._torch.modeling_v2" +_PACKAGE = "tensorrt_llm._torch._experimental.modeling_v2" #: The switch. An environment variable rather than an LLM-API field, so that #: nothing outside this package has to carry the concept: the only upstream diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/__init__.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/__init__.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/__init__.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/activation/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/activation/__init__.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/activation/__init__.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/activation/__init__.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/activation/flashinfer_silu_and_mul.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/activation/flashinfer_silu_and_mul.md similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/activation/flashinfer_silu_and_mul.md rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/activation/flashinfer_silu_and_mul.md diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/activation/flashinfer_silu_and_mul.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/activation/flashinfer_silu_and_mul.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/activation/flashinfer_silu_and_mul.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/activation/flashinfer_silu_and_mul.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/attention/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/__init__.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/attention/__init__.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/__init__.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/attention/fused_qk_norm_rope.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/fused_qk_norm_rope.md similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/attention/fused_qk_norm_rope.md rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/fused_qk_norm_rope.md diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/attention/fused_qk_norm_rope.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/fused_qk_norm_rope.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/attention/fused_qk_norm_rope.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/fused_qk_norm_rope.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/attention/load_paged_kv_cache_for_mla.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/load_paged_kv_cache_for_mla.md similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/attention/load_paged_kv_cache_for_mla.md rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/load_paged_kv_cache_for_mla.md diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/attention/load_paged_kv_cache_for_mla.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/load_paged_kv_cache_for_mla.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/attention/load_paged_kv_cache_for_mla.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/load_paged_kv_cache_for_mla.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_append_paged_kv_assign_q.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/mla_rope_append_paged_kv_assign_q.md similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_append_paged_kv_assign_q.md rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/mla_rope_append_paged_kv_assign_q.md diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_append_paged_kv_assign_q.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/mla_rope_append_paged_kv_assign_q.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_append_paged_kv_assign_q.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/mla_rope_append_paged_kv_assign_q.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_generation.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/mla_rope_generation.md similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_generation.md rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/mla_rope_generation.md diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_generation.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/mla_rope_generation.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/attention/mla_rope_generation.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/mla_rope_generation.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/attention/thop_attention.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/thop_attention.md similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/attention/thop_attention.md rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/thop_attention.md diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/attention/thop_attention.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/thop_attention.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/attention/thop_attention.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/attention/thop_attention.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/comm/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/__init__.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/comm/__init__.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/__init__.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/comm/allgather.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/allgather.md similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/comm/allgather.md rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/allgather.md diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/comm/allgather.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/allgather.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/comm/allgather.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/allgather.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/comm/reducescatter.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/reducescatter.md similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/comm/reducescatter.md rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/reducescatter.md diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/comm/reducescatter.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/reducescatter.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/comm/reducescatter.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/comm/reducescatter.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/gemm/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/__init__.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/gemm/__init__.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/__init__.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/gemm/bmm_out.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/bmm_out.md similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/gemm/bmm_out.md rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/bmm_out.md diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/gemm/bmm_out.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/bmm_out.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/gemm/bmm_out.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/bmm_out.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/gemm/cublas_mm.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/cublas_mm.md similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/gemm/cublas_mm.md rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/cublas_mm.md diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/gemm/cublas_mm.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/cublas_mm.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/gemm/cublas_mm.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/cublas_mm.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/gemm/nvfp4_gemm.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/nvfp4_gemm.md similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/gemm/nvfp4_gemm.md rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/nvfp4_gemm.md diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/gemm/nvfp4_gemm.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/nvfp4_gemm.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/gemm/nvfp4_gemm.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/nvfp4_gemm.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/index.yaml b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/index.yaml rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/index.yaml diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/moe/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/__init__.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/moe/__init__.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/__init__.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/moe/fp4_block_scale_moe_runner.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/fp4_block_scale_moe_runner.md similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/moe/fp4_block_scale_moe_runner.md rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/fp4_block_scale_moe_runner.md diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/moe/fp4_block_scale_moe_runner.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/fp4_block_scale_moe_runner.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/moe/fp4_block_scale_moe_runner.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/fp4_block_scale_moe_runner.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/moe/fused_moe.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/fused_moe.md similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/moe/fused_moe.md rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/fused_moe.md diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/moe/fused_moe.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/fused_moe.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/moe/fused_moe.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/fused_moe.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.md diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/mxe4m3_mxe2m1_block_scale_moe_runner.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/moe/noaux_tc_op.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/noaux_tc_op.md similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/moe/noaux_tc_op.md rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/noaux_tc_op.md diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/moe/noaux_tc_op.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/noaux_tc_op.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/moe/noaux_tc_op.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/moe/noaux_tc_op.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/norm/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/__init__.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/norm/__init__.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/__init__.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_fused_add_rmsnorm.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/flashinfer_fused_add_rmsnorm.md similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_fused_add_rmsnorm.md rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/flashinfer_fused_add_rmsnorm.md diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_fused_add_rmsnorm.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/flashinfer_fused_add_rmsnorm.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_fused_add_rmsnorm.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/flashinfer_fused_add_rmsnorm.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_rmsnorm.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/flashinfer_rmsnorm.md similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_rmsnorm.md rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/flashinfer_rmsnorm.md diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_rmsnorm.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/flashinfer_rmsnorm.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/norm/flashinfer_rmsnorm.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/norm/flashinfer_rmsnorm.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/quantization/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/quantization/__init__.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/quantization/__init__.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/quantization/__init__.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/quantization/fp4_quantize.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/quantization/fp4_quantize.md similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/quantization/fp4_quantize.md rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/quantization/fp4_quantize.md diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/quantization/fp4_quantize.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/quantization/fp4_quantize.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/quantization/fp4_quantize.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/quantization/fp4_quantize.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/quantization/mxfp8_quantize.md b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/quantization/mxfp8_quantize.md similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/quantization/mxfp8_quantize.md rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/quantization/mxfp8_quantize.md diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/quantization/mxfp8_quantize.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/quantization/mxfp8_quantize.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/quantization/mxfp8_quantize.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/quantization/mxfp8_quantize.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/torch/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/__init__.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/torch/__init__.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/__init__.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/torch/add.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/add.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/torch/add.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/add.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/torch/concat.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/concat.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/torch/concat.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/concat.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/torch/copy_.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/copy_.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/torch/copy_.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/copy_.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/torch/embedding.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/embedding.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/torch/embedding.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/embedding.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/torch/empty.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/empty.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/torch/empty.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/empty.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/torch/expand.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/expand.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/torch/expand.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/expand.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/torch/pad.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/pad.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/torch/pad.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/pad.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/torch/reshape.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/reshape.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/torch/reshape.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/reshape.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/torch/split.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/split.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/torch/split.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/split.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/torch/transpose.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/transpose.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/torch/transpose.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/transpose.py diff --git a/tensorrt_llm/_torch/modeling_v2/catalog/torch/view_dtype.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/view_dtype.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/catalog/torch/view_dtype.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/catalog/torch/view_dtype.py diff --git a/tensorrt_llm/_torch/modeling_v2/explain.py b/tensorrt_llm/_torch/_experimental/modeling_v2/explain.py similarity index 96% rename from tensorrt_llm/_torch/modeling_v2/explain.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/explain.py index ad35eccb9f03..683a11400d38 100644 --- a/tensorrt_llm/_torch/modeling_v2/explain.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/explain.py @@ -2,7 +2,7 @@ # SPDX-License-Identifier: Apache-2.0 """Say which target a configuration routes to, and why. - python -m tensorrt_llm._torch.modeling_v2.explain \ + python -m tensorrt_llm._torch._experimental.modeling_v2.explain \ --model /path/to/DeepSeek-R1-0528-NVFP4 --tp 4 --ep 4 --attention-dp Prints the routing module's decision tree as it was actually evaluated, one @@ -41,7 +41,7 @@ def _sm(value: Optional[str]) -> Tuple[int, int]: def build_parser() -> argparse.ArgumentParser: p = argparse.ArgumentParser( - prog="python -m tensorrt_llm._torch.modeling_v2.explain", + prog="python -m tensorrt_llm._torch._experimental.modeling_v2.explain", description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter, ) diff --git a/tensorrt_llm/_torch/modeling_v2/models/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/__init__.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/models/__init__.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/models/__init__.py diff --git a/tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/__init__.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/__init__.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/__init__.py diff --git a/tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/routing.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/routing.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/routing.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/routing.py diff --git a/tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/__init__.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/__init__.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/__init__.py diff --git a/tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/__init__.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/__init__.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/__init__.py diff --git a/tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/__init__.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/__init__.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/__init__.py diff --git a/tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/__init__.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/__init__.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/__init__.py diff --git a/tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py similarity index 97% rename from tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py index da0aab9e4e21..65e903151a7a 100644 --- a/tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py @@ -55,48 +55,54 @@ from torch import nn from transformers import PretrainedConfig -from tensorrt_llm._torch.attention.backends.interface import AttentionMetadata -from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata -from tensorrt_llm._torch.model_config import ModelConfig -from tensorrt_llm._torch.modeling_v2.catalog.activation.flashinfer_silu_and_mul import ( # noqa: E501 +from tensorrt_llm._torch._experimental.modeling_v2.catalog.activation.flashinfer_silu_and_mul import ( # noqa: E501 flashinfer_silu_and_mul, ) -from tensorrt_llm._torch.modeling_v2.catalog.attention.load_paged_kv_cache_for_mla import ( # noqa: E501 +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.load_paged_kv_cache_for_mla import ( # noqa: E501 load_paged_kv_cache_for_mla, ) -from tensorrt_llm._torch.modeling_v2.catalog.attention.mla_rope_append_paged_kv_assign_q import ( # noqa: E501 +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.mla_rope_append_paged_kv_assign_q import ( # noqa: E501 mla_rope_append_paged_kv_assign_q, ) -from tensorrt_llm._torch.modeling_v2.catalog.attention.mla_rope_generation import ( +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.mla_rope_generation import ( mla_rope_generation, ) -from tensorrt_llm._torch.modeling_v2.catalog.attention.thop_attention import thop_attention -from tensorrt_llm._torch.modeling_v2.catalog.comm.allgather import allgather -from tensorrt_llm._torch.modeling_v2.catalog.comm.reducescatter import reducescatter -from tensorrt_llm._torch.modeling_v2.catalog.gemm.bmm_out import bmm_out -from tensorrt_llm._torch.modeling_v2.catalog.gemm.cublas_mm import cublas_mm -from tensorrt_llm._torch.modeling_v2.catalog.gemm.nvfp4_gemm import nvfp4_gemm -from tensorrt_llm._torch.modeling_v2.catalog.moe.fp4_block_scale_moe_runner import ( # noqa: E501 +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.thop_attention import ( + thop_attention, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.allgather import allgather +from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm.reducescatter import reducescatter +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.bmm_out import bmm_out +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.cublas_mm import cublas_mm +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.nvfp4_gemm import nvfp4_gemm +from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.fp4_block_scale_moe_runner import ( # noqa: E501 fp4_block_scale_moe_runner, ) -from tensorrt_llm._torch.modeling_v2.catalog.moe.fused_moe import fused_moe -from tensorrt_llm._torch.modeling_v2.catalog.moe.noaux_tc_op import noaux_tc_op -from tensorrt_llm._torch.modeling_v2.catalog.norm.flashinfer_fused_add_rmsnorm import ( # noqa: E501 +from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.fused_moe import fused_moe +from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.noaux_tc_op import noaux_tc_op +from tensorrt_llm._torch._experimental.modeling_v2.catalog.norm.flashinfer_fused_add_rmsnorm import ( # noqa: E501 flashinfer_fused_add_rmsnorm, ) -from tensorrt_llm._torch.modeling_v2.catalog.norm.flashinfer_rmsnorm import flashinfer_rmsnorm -from tensorrt_llm._torch.modeling_v2.catalog.quantization.fp4_quantize import fp4_quantize -from tensorrt_llm._torch.modeling_v2.catalog.torch.add import add -from tensorrt_llm._torch.modeling_v2.catalog.torch.concat import concat -from tensorrt_llm._torch.modeling_v2.catalog.torch.copy_ import copy_ -from tensorrt_llm._torch.modeling_v2.catalog.torch.embedding import embedding -from tensorrt_llm._torch.modeling_v2.catalog.torch.empty import empty -from tensorrt_llm._torch.modeling_v2.catalog.torch.expand import expand -from tensorrt_llm._torch.modeling_v2.catalog.torch.pad import pad -from tensorrt_llm._torch.modeling_v2.catalog.torch.reshape import reshape -from tensorrt_llm._torch.modeling_v2.catalog.torch.split import split -from tensorrt_llm._torch.modeling_v2.catalog.torch.transpose import transpose -from tensorrt_llm._torch.modeling_v2.catalog.torch.view_dtype import view_dtype +from tensorrt_llm._torch._experimental.modeling_v2.catalog.norm.flashinfer_rmsnorm import ( + flashinfer_rmsnorm, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.quantization.fp4_quantize import ( + fp4_quantize, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.torch.add import add +from tensorrt_llm._torch._experimental.modeling_v2.catalog.torch.concat import concat +from tensorrt_llm._torch._experimental.modeling_v2.catalog.torch.copy_ import copy_ +from tensorrt_llm._torch._experimental.modeling_v2.catalog.torch.embedding import embedding +from tensorrt_llm._torch._experimental.modeling_v2.catalog.torch.empty import empty +from tensorrt_llm._torch._experimental.modeling_v2.catalog.torch.expand import expand +from tensorrt_llm._torch._experimental.modeling_v2.catalog.torch.pad import pad +from tensorrt_llm._torch._experimental.modeling_v2.catalog.torch.reshape import reshape +from tensorrt_llm._torch._experimental.modeling_v2.catalog.torch.split import split +from tensorrt_llm._torch._experimental.modeling_v2.catalog.torch.transpose import transpose +from tensorrt_llm._torch._experimental.modeling_v2.catalog.torch.view_dtype import view_dtype +from tensorrt_llm._torch.attention.backends.interface import AttentionMetadata +from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata +from tensorrt_llm._torch.model_config import ModelConfig from tensorrt_llm._torch.models.modeling_utils import ( DecoderModel, DecoderModelForCausalLM, diff --git a/tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/weights.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/weights.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/weights.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/weights.py diff --git a/tensorrt_llm/_torch/modeling_v2/models/gpt_oss/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/__init__.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/models/gpt_oss/__init__.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/__init__.py diff --git a/tensorrt_llm/_torch/modeling_v2/models/gpt_oss/routing.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/routing.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/models/gpt_oss/routing.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/routing.py diff --git a/tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/__init__.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/__init__.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/__init__.py diff --git a/tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/__init__.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/__init__.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/__init__.py diff --git a/tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/__init__.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/__init__.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/__init__.py diff --git a/tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/__init__.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/__init__.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/__init__.py diff --git a/tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py similarity index 96% rename from tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py index db6c43bc476e..6adc9e2b0798 100644 --- a/tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py @@ -42,23 +42,31 @@ from torch import nn from transformers import PretrainedConfig -from tensorrt_llm._torch.attention.backends.interface import AttentionMetadata -from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata -from tensorrt_llm._torch.model_config import ModelConfig -from tensorrt_llm._torch.modeling_v2.catalog.attention.fused_qk_norm_rope import fused_qk_norm_rope -from tensorrt_llm._torch.modeling_v2.catalog.attention.thop_attention import thop_attention -from tensorrt_llm._torch.modeling_v2.catalog.gemm.cublas_mm import cublas_mm -from tensorrt_llm._torch.modeling_v2.catalog.moe.mxe4m3_mxe2m1_block_scale_moe_runner import ( # noqa: E501 +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.fused_qk_norm_rope import ( + fused_qk_norm_rope, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.thop_attention import ( + thop_attention, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.cublas_mm import cublas_mm +from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.mxe4m3_mxe2m1_block_scale_moe_runner import ( # noqa: E501 mxe4m3_mxe2m1_block_scale_moe_runner, ) -from tensorrt_llm._torch.modeling_v2.catalog.norm.flashinfer_fused_add_rmsnorm import ( # noqa: E501 +from tensorrt_llm._torch._experimental.modeling_v2.catalog.norm.flashinfer_fused_add_rmsnorm import ( # noqa: E501 flashinfer_fused_add_rmsnorm, ) -from tensorrt_llm._torch.modeling_v2.catalog.norm.flashinfer_rmsnorm import flashinfer_rmsnorm -from tensorrt_llm._torch.modeling_v2.catalog.quantization.mxfp8_quantize import mxfp8_quantize -from tensorrt_llm._torch.modeling_v2.catalog.torch.embedding import embedding -from tensorrt_llm._torch.modeling_v2.catalog.torch.empty import empty -from tensorrt_llm._torch.modeling_v2.catalog.torch.reshape import reshape +from tensorrt_llm._torch._experimental.modeling_v2.catalog.norm.flashinfer_rmsnorm import ( + flashinfer_rmsnorm, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.quantization.mxfp8_quantize import ( + mxfp8_quantize, +) +from tensorrt_llm._torch._experimental.modeling_v2.catalog.torch.embedding import embedding +from tensorrt_llm._torch._experimental.modeling_v2.catalog.torch.empty import empty +from tensorrt_llm._torch._experimental.modeling_v2.catalog.torch.reshape import reshape +from tensorrt_llm._torch.attention.backends.interface import AttentionMetadata +from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata +from tensorrt_llm._torch.model_config import ModelConfig from tensorrt_llm._torch.models.modeling_utils import ( DecoderModel, DecoderModelForCausalLM, diff --git a/tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/weights.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/weights.py similarity index 100% rename from tensorrt_llm/_torch/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/weights.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/weights.py diff --git a/tensorrt_llm/_torch/models/modeling_auto.py b/tensorrt_llm/_torch/models/modeling_auto.py index 96f2d55a430c..9fee868223a6 100644 --- a/tensorrt_llm/_torch/models/modeling_auto.py +++ b/tensorrt_llm/_torch/models/modeling_auto.py @@ -1,7 +1,7 @@ from typing import Generic, Optional, Type +from .._experimental.modeling_v2 import modeling_v2_resolve from ..model_config import ModelConfig -from ..modeling_v2 import modeling_v2_resolve from ..utils import model_extra_attrs from .modeling_utils import (DecoderModelForCausalLM, TConfig, TModel, get_registered_model_class, diff --git a/tensorrt_llm/llmapi/llm.py b/tensorrt_llm/llmapi/llm.py index 00bdad277a1f..19a412b3e343 100644 --- a/tensorrt_llm/llmapi/llm.py +++ b/tensorrt_llm/llmapi/llm.py @@ -397,7 +397,8 @@ def __init__(self, # target. Only the pytorch backend reaches the resolver that could # select one, so on any other backend the promise would be broken # silently -- the one failure that mode exists to prevent. - from .._torch.modeling_v2 import assert_backend_can_route + from .._torch._experimental.modeling_v2 import \ + assert_backend_can_route assert_backend_can_route(backend) # check the kwargs and raise ValueError directly diff --git a/tests/integration/defs/accuracy/test_modeling_v2_deepseek_v3.py b/tests/integration/defs/accuracy/test_modeling_v2_deepseek_v3.py index ccc58e2830f9..8fdaaa199604 100644 --- a/tests/integration/defs/accuracy/test_modeling_v2_deepseek_v3.py +++ b/tests/integration/defs/accuracy/test_modeling_v2_deepseek_v3.py @@ -38,7 +38,7 @@ import pytest from tensorrt_llm import LLM -from tensorrt_llm._torch.modeling_v2 import MODELING_V2_ENV +from tensorrt_llm._torch._experimental.modeling_v2 import MODELING_V2_ENV from tensorrt_llm._utils import get_sm_version from tensorrt_llm.llmapi import CudaGraphConfig, KvCacheConfig, MTPDecodingConfig diff --git a/tests/integration/defs/accuracy/test_modeling_v2_gpt_oss.py b/tests/integration/defs/accuracy/test_modeling_v2_gpt_oss.py index 1905a698409f..746335586630 100644 --- a/tests/integration/defs/accuracy/test_modeling_v2_gpt_oss.py +++ b/tests/integration/defs/accuracy/test_modeling_v2_gpt_oss.py @@ -30,7 +30,7 @@ import pytest from tensorrt_llm import LLM -from tensorrt_llm._torch.modeling_v2 import MODELING_V2_ENV +from tensorrt_llm._torch._experimental.modeling_v2 import MODELING_V2_ENV from tensorrt_llm._utils import get_sm_version from ..conftest import llm_models_root diff --git a/tests/unittest/_torch/modeling_v2/activation/test_modeling_v2_flashinfer_silu_and_mul.py b/tests/unittest/_torch/modeling_v2/activation/test_modeling_v2_flashinfer_silu_and_mul.py index 02384727943c..54338b915cb5 100644 --- a/tests/unittest/_torch/modeling_v2/activation/test_modeling_v2_flashinfer_silu_and_mul.py +++ b/tests/unittest/_torch/modeling_v2/activation/test_modeling_v2_flashinfer_silu_and_mul.py @@ -5,7 +5,7 @@ import torch import torch.nn.functional as F -from tensorrt_llm._torch.modeling_v2.catalog.activation.flashinfer_silu_and_mul import ( +from tensorrt_llm._torch._experimental.modeling_v2.catalog.activation.flashinfer_silu_and_mul import ( flashinfer_silu_and_mul, ) diff --git a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_fused_qk_norm_rope.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_fused_qk_norm_rope.py index bfa85580bb2c..b71a64912ef6 100644 --- a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_fused_qk_norm_rope.py +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_fused_qk_norm_rope.py @@ -4,7 +4,9 @@ import torch -from tensorrt_llm._torch.modeling_v2.catalog.attention.fused_qk_norm_rope import fused_qk_norm_rope +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.fused_qk_norm_rope import ( + fused_qk_norm_rope, +) assert torch.cuda.is_available(), "fused_qk_norm_rope requires a CUDA device" diff --git a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_load_paged_kv_cache_for_mla.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_load_paged_kv_cache_for_mla.py index 2554428a6cae..690d54221ec5 100644 --- a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_load_paged_kv_cache_for_mla.py +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_load_paged_kv_cache_for_mla.py @@ -32,11 +32,11 @@ import torch -from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata -from tensorrt_llm._torch.metadata import KVCacheParams -from tensorrt_llm._torch.modeling_v2.catalog.attention.load_paged_kv_cache_for_mla import ( +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.load_paged_kv_cache_for_mla import ( load_paged_kv_cache_for_mla, ) +from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata +from tensorrt_llm._torch.metadata import KVCacheParams from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager from tensorrt_llm.bindings import DataType from tensorrt_llm.bindings.internal.batch_manager import CacheType diff --git a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_mla_rope_append_paged_kv_assign_q.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_mla_rope_append_paged_kv_assign_q.py index 2a9c7ccc4ffa..1158d946fcc9 100644 --- a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_mla_rope_append_paged_kv_assign_q.py +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_mla_rope_append_paged_kv_assign_q.py @@ -36,12 +36,12 @@ import torch +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.mla_rope_append_paged_kv_assign_q import ( + mla_rope_append_paged_kv_assign_q, +) from tensorrt_llm._torch.attention.backends.interface import RopeParams from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata from tensorrt_llm._torch.metadata import KVCacheParams -from tensorrt_llm._torch.modeling_v2.catalog.attention.mla_rope_append_paged_kv_assign_q import ( - mla_rope_append_paged_kv_assign_q, -) from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager from tensorrt_llm.bindings import DataType from tensorrt_llm.bindings.internal.batch_manager import CacheType diff --git a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_mla_rope_generation.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_mla_rope_generation.py index b526e4c54323..10837d336854 100644 --- a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_mla_rope_generation.py +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_mla_rope_generation.py @@ -48,12 +48,12 @@ import torch +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.mla_rope_generation import ( + mla_rope_generation, +) from tensorrt_llm._torch.attention.backends.interface import RopeParams from tensorrt_llm._torch.attention.backends.trtllm import TrtllmAttentionMetadata from tensorrt_llm._torch.metadata import KVCacheParams -from tensorrt_llm._torch.modeling_v2.catalog.attention.mla_rope_generation import ( - mla_rope_generation, -) from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager from tensorrt_llm.bindings import DataType from tensorrt_llm.bindings.internal.batch_manager import CacheType diff --git a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_thop_attention.py b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_thop_attention.py index bd3c814d44a0..cb0a1f090d4c 100644 --- a/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_thop_attention.py +++ b/tests/unittest/_torch/modeling_v2/attention/test_modeling_v2_thop_attention.py @@ -172,8 +172,10 @@ import torch import torch.nn.functional as F +from tensorrt_llm._torch._experimental.modeling_v2.catalog.attention.thop_attention import ( + thop_attention, +) from tensorrt_llm._torch.attention.backends.interface import RopeParams -from tensorrt_llm._torch.modeling_v2.catalog.attention.thop_attention import thop_attention from tensorrt_llm._torch.pyexecutor.resource_manager import CacheTypeCpp, DataType, KVCacheManager from tensorrt_llm.functional import RotaryScalingType from tensorrt_llm.llmapi.llm_args import KvCacheConfig diff --git a/tests/unittest/_torch/modeling_v2/comm/_allgather_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_allgather_op_matrix.py index 94e8e0add9f6..bf10463742ee 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_allgather_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_allgather_op_matrix.py @@ -916,8 +916,8 @@ def _run_one_rank() -> int: global COMM, RANK, WORLD, GROUP, allgather, DIST, MPI from mpi4py import MPI as _MPI + from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm import allgather as entry from tensorrt_llm._torch.distributed import Distributed - from tensorrt_llm._torch.modeling_v2.catalog.comm import allgather as entry from tensorrt_llm.mapping import Mapping allgather = entry.allgather diff --git a/tests/unittest/_torch/modeling_v2/comm/_reducescatter_op_matrix.py b/tests/unittest/_torch/modeling_v2/comm/_reducescatter_op_matrix.py index aea196658056..fd5cc5444bd8 100644 --- a/tests/unittest/_torch/modeling_v2/comm/_reducescatter_op_matrix.py +++ b/tests/unittest/_torch/modeling_v2/comm/_reducescatter_op_matrix.py @@ -1430,7 +1430,7 @@ def _run_one_rank() -> int: global COMM, RANK, WORLD, GROUP, reducescatter from mpi4py import MPI - from tensorrt_llm._torch.modeling_v2.catalog.comm import reducescatter as entry + from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm import reducescatter as entry reducescatter = entry.reducescatter COMM = MPI.COMM_WORLD @@ -1497,7 +1497,7 @@ def _run_wedge_rank() -> int: global COMM, RANK, WORLD, GROUP, reducescatter from mpi4py import MPI - from tensorrt_llm._torch.modeling_v2.catalog.comm import reducescatter as entry + from tensorrt_llm._torch._experimental.modeling_v2.catalog.comm import reducescatter as entry reducescatter = entry.reducescatter COMM = MPI.COMM_WORLD diff --git a/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_bmm_out.py b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_bmm_out.py index a614806ba5db..28c57009b5ee 100644 --- a/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_bmm_out.py +++ b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_bmm_out.py @@ -4,7 +4,7 @@ import torch -from tensorrt_llm._torch.modeling_v2.catalog.gemm.bmm_out import bmm_out +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.bmm_out import bmm_out assert torch.cuda.is_available(), "bmm_out requires a CUDA device" diff --git a/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_cublas_mm.py b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_cublas_mm.py index 51a9fcad2c49..6bd55d915a4c 100644 --- a/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_cublas_mm.py +++ b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_cublas_mm.py @@ -6,7 +6,7 @@ import torch -from tensorrt_llm._torch.modeling_v2.catalog.gemm.cublas_mm import cublas_mm +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.cublas_mm import cublas_mm assert torch.cuda.is_available(), "cublas_mm requires a CUDA device" diff --git a/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_nvfp4_gemm.py b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_nvfp4_gemm.py index 08b2d5db8342..a0f739f36531 100644 --- a/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_nvfp4_gemm.py +++ b/tests/unittest/_torch/modeling_v2/gemm/test_modeling_v2_nvfp4_gemm.py @@ -5,8 +5,8 @@ import torch import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* +from tensorrt_llm._torch._experimental.modeling_v2.catalog.gemm.nvfp4_gemm import nvfp4_gemm from tensorrt_llm._torch.autotuner import AutoTuner, autotune -from tensorrt_llm._torch.modeling_v2.catalog.gemm.nvfp4_gemm import nvfp4_gemm assert torch.cuda.is_available(), "nvfp4_gemm requires a CUDA device" # The reference matmul must be true fp32, never tf32. diff --git a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_fp4_block_scale_moe_runner.py b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_fp4_block_scale_moe_runner.py index 849a658aa9dd..ae56f8a4e4ac 100644 --- a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_fp4_block_scale_moe_runner.py +++ b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_fp4_block_scale_moe_runner.py @@ -4,10 +4,10 @@ import torch -from tensorrt_llm._torch.autotuner import AutoTuner, autotune -from tensorrt_llm._torch.modeling_v2.catalog.moe.fp4_block_scale_moe_runner import ( +from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.fp4_block_scale_moe_runner import ( fp4_block_scale_moe_runner as moe, ) +from tensorrt_llm._torch.autotuner import AutoTuner, autotune assert torch.cuda.is_available(), "fp4_block_scale_moe_runner requires a CUDA device" diff --git a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_fused_moe.py b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_fused_moe.py index 94ec6e153bb2..a66387a68479 100644 --- a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_fused_moe.py +++ b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_fused_moe.py @@ -5,8 +5,8 @@ import torch import torch.nn.functional as F +from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.fused_moe import fused_moe from tensorrt_llm._torch.autotuner import AutoTuner, autotune -from tensorrt_llm._torch.modeling_v2.catalog.moe.fused_moe import fused_moe assert torch.cuda.is_available(), "fused_moe requires a CUDA device" diff --git a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_mxe4m3_mxe2m1_block_scale_moe_runner.py b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_mxe4m3_mxe2m1_block_scale_moe_runner.py index da117aa8c9de..71795816de8b 100644 --- a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_mxe4m3_mxe2m1_block_scale_moe_runner.py +++ b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_mxe4m3_mxe2m1_block_scale_moe_runner.py @@ -6,7 +6,7 @@ import torch -from tensorrt_llm._torch.modeling_v2.catalog.moe.mxe4m3_mxe2m1_block_scale_moe_runner import ( +from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.mxe4m3_mxe2m1_block_scale_moe_runner import ( mxe4m3_mxe2m1_block_scale_moe_runner as moe, ) diff --git a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_noaux_tc_op.py b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_noaux_tc_op.py index e4d2e14e0694..f66e797524fc 100644 --- a/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_noaux_tc_op.py +++ b/tests/unittest/_torch/modeling_v2/moe/test_modeling_v2_noaux_tc_op.py @@ -4,7 +4,7 @@ import torch -from tensorrt_llm._torch.modeling_v2.catalog.moe.noaux_tc_op import noaux_tc_op +from tensorrt_llm._torch._experimental.modeling_v2.catalog.moe.noaux_tc_op import noaux_tc_op assert torch.cuda.is_available(), "noaux_tc_op requires a CUDA device" diff --git a/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_flashinfer_fused_add_rmsnorm.py b/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_flashinfer_fused_add_rmsnorm.py index a6742c988aec..4b5858a26502 100644 --- a/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_flashinfer_fused_add_rmsnorm.py +++ b/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_flashinfer_fused_add_rmsnorm.py @@ -4,7 +4,7 @@ import torch -from tensorrt_llm._torch.modeling_v2.catalog.norm.flashinfer_fused_add_rmsnorm import ( +from tensorrt_llm._torch._experimental.modeling_v2.catalog.norm.flashinfer_fused_add_rmsnorm import ( flashinfer_fused_add_rmsnorm, ) diff --git a/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_flashinfer_rmsnorm.py b/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_flashinfer_rmsnorm.py index bc2a1925b89c..8492fdc49eeb 100644 --- a/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_flashinfer_rmsnorm.py +++ b/tests/unittest/_torch/modeling_v2/norm/test_modeling_v2_flashinfer_rmsnorm.py @@ -4,7 +4,9 @@ import torch -from tensorrt_llm._torch.modeling_v2.catalog.norm.flashinfer_rmsnorm import flashinfer_rmsnorm +from tensorrt_llm._torch._experimental.modeling_v2.catalog.norm.flashinfer_rmsnorm import ( + flashinfer_rmsnorm, +) assert torch.cuda.is_available(), "flashinfer_rmsnorm requires a CUDA device" diff --git a/tests/unittest/_torch/modeling_v2/quantization/test_modeling_v2_fp4_quantize.py b/tests/unittest/_torch/modeling_v2/quantization/test_modeling_v2_fp4_quantize.py index d49c0b662686..86881d82e12c 100644 --- a/tests/unittest/_torch/modeling_v2/quantization/test_modeling_v2_fp4_quantize.py +++ b/tests/unittest/_torch/modeling_v2/quantization/test_modeling_v2_fp4_quantize.py @@ -6,7 +6,9 @@ from torch.profiler import ProfilerActivity, profile import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* -from tensorrt_llm._torch.modeling_v2.catalog.quantization.fp4_quantize import fp4_quantize +from tensorrt_llm._torch._experimental.modeling_v2.catalog.quantization.fp4_quantize import ( + fp4_quantize, +) assert torch.cuda.is_available(), "fp4_quantize requires a CUDA device" diff --git a/tests/unittest/_torch/modeling_v2/quantization/test_modeling_v2_mxfp8_quantize.py b/tests/unittest/_torch/modeling_v2/quantization/test_modeling_v2_mxfp8_quantize.py index 2b804fe39047..3dc7b5a1a8cd 100644 --- a/tests/unittest/_torch/modeling_v2/quantization/test_modeling_v2_mxfp8_quantize.py +++ b/tests/unittest/_torch/modeling_v2/quantization/test_modeling_v2_mxfp8_quantize.py @@ -4,7 +4,9 @@ import torch -from tensorrt_llm._torch.modeling_v2.catalog.quantization.mxfp8_quantize import mxfp8_quantize +from tensorrt_llm._torch._experimental.modeling_v2.catalog.quantization.mxfp8_quantize import ( + mxfp8_quantize, +) assert torch.cuda.is_available(), "mxfp8_quantize requires a CUDA device" diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py index 84a41c344a4c..8d533738a2c7 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py @@ -22,8 +22,11 @@ import pytest -import tensorrt_llm._torch.modeling_v2 as _modeling_v2 -from tensorrt_llm._torch.modeling_v2._router_index import MODELING_V2_ROUTERS, routing_module +import tensorrt_llm._torch._experimental.modeling_v2 as _modeling_v2 +from tensorrt_llm._torch._experimental.modeling_v2._router_index import ( + MODELING_V2_ROUTERS, + routing_module, +) # The package, not this file: these paths address the tree under test, and this # test lives in tests/ while that tree lives in tensorrt_llm/. diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_no_stale_claims.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_no_stale_claims.py index 98286ab3f99a..f251846bcc71 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_no_stale_claims.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_no_stale_claims.py @@ -23,7 +23,7 @@ import re from pathlib import Path -import tensorrt_llm._torch.modeling_v2 as _modeling_v2 +import tensorrt_llm._torch._experimental.modeling_v2 as _modeling_v2 # The package, not this file: the tree under audit lives under tensorrt_llm/ # while this test lives under tests/. diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py index 85d85438ffc1..920c14c027b9 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py @@ -23,12 +23,12 @@ import torch from transformers import PretrainedConfig -from tensorrt_llm._torch.model_config import ModelConfig -from tensorrt_llm._torch.modeling_v2._router_index import ( +from tensorrt_llm._torch._experimental.modeling_v2._router_index import ( MODELING_V2_ENV, ModelingV2Mode, modeling_v2_resolve, ) +from tensorrt_llm._torch.model_config import ModelConfig from tensorrt_llm._torch.models.modeling_auto import AutoModelForCausalLM from tensorrt_llm._torch.models.modeling_utils import ( _is_builtin_model_class, diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_target_contract.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_target_contract.py index 50e40c22e9fc..e04a4e01ac4f 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_target_contract.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_target_contract.py @@ -25,9 +25,12 @@ import torch import tensorrt_llm._torch.custom_ops # noqa: F401 — registers torch.ops.trtllm.* -from tensorrt_llm._torch.modeling_v2._router_index import MODELING_V2_ROUTERS, routing_module +from tensorrt_llm._torch._experimental.modeling_v2._router_index import ( + MODELING_V2_ROUTERS, + routing_module, +) -_PACKAGE = "tensorrt_llm._torch.modeling_v2" +_PACKAGE = "tensorrt_llm._torch._experimental.modeling_v2" _ARCHS = sorted(MODELING_V2_ROUTERS) From d97d768713a6f935e11254e24734cd06a3404327 Mon Sep 17 00:00:00 2001 From: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> Date: Wed, 16 Sep 2026 20:46:13 -0700 Subject: [PATCH 13/19] [TRTLLM-16304][chore] Fix the codespell finding blocking pre-commit codespell rejects `pre-empting`. The hook has been failing since the first commit; the login node has no pre-commit, so it only showed up in CI. Signed-off-by: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> --- .../_torch/_experimental/modeling_v2/catalog/gemm/nvfp4_gemm.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/nvfp4_gemm.py b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/nvfp4_gemm.py index 52f439ac64cb..3e29958228c6 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/nvfp4_gemm.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/catalog/gemm/nvfp4_gemm.py @@ -56,7 +56,7 @@ def nvfp4_gemm( # # Only for 2-D operands, and only for *under*-length. Anything else about # the shapes -- wrong rank, zero rows, K disagreement -- the op rejects - # itself and loudly, and pre-empting that here would swap its RuntimeError + # itself and loudly, and preempting that here would swap its RuntimeError # for an AssertionError that says less. Over-length is not the hazard # either: the kernel never reads past what the swizzle addresses, and the # contract already says the padding bytes may hold anything. From 4c9b605edf269948979c185e9b7c8443e8898b9f Mon Sep 17 00:00:00 2001 From: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> Date: Mon, 21 Sep 2026 23:13:58 -0700 Subject: [PATCH 14/19] [TRTLLM-16304][test] Carry the modeling_v2 switch on the case, not the shell The whole-model gates read TRTLLM_MODELING_V2 and skipped unless the launching shell had exported it. Nothing in CI exports it, so every one of them skipped: the stage reported green without ever measuring a modeling_v2 target. That is the same misattribution `require` mode exists to prevent, reached through the test harness rather than through the model. Setting os.environ in the test process would not have fixed it. The variable is read by modeling_v2_resolve, called from AutoModelForCausalLM._resolve_class, which runs in a worker rank -- and MPI caches the environment at import time and spawns from that snapshot, which is why worker_main re-applies env_overrides by hand at its top. So the switch travels as an LLM(...) field instead: worker_main applies it inside each rank before the model is built, and LLM.__init__ applies the same overrides to the calling process before it checks the mode itself. The helper takes monkeypatch, which is not decoration: LLM applies its overrides to the calling process and never puts them back, so with no teardown the first case asking for "require" would leave every later case in the session asking for it -- and `require` raises on an architecture with no target, so unrelated tests downstream would fail with none of these files in the traceback. The stock leg of the acceptance gate asserted TRTLLM_CAN_USE_DEEP_EP=0 had been exported. It rides along the same route now, since it is read on the ranks too. Signed-off-by: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> --- .../modeling_v2/_router_index.py | 15 ++-- .../defs/accuracy/modeling_v2_env.py | 61 +++++++++++++++ .../accuracy/test_modeling_v2_deepseek_v3.py | 76 ++++++++----------- .../defs/accuracy/test_modeling_v2_gpt_oss.py | 31 ++------ 4 files changed, 107 insertions(+), 76 deletions(-) create mode 100644 tests/integration/defs/accuracy/modeling_v2_env.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/_router_index.py b/tensorrt_llm/_torch/_experimental/modeling_v2/_router_index.py index 3a4094bb1632..d1b2a7840899 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/_router_index.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/_router_index.py @@ -45,13 +45,14 @@ #: nothing outside this package has to carry the concept: the only upstream #: change modeling_v2 needs is the ``_resolve_class`` hook itself. #: -#: It has to be exported **before the ranks start**, not merely before -#: ``LLM(...)``. Worker ranks receive the environment as it stood when MPI -#: initialized, and long-lived ranks under ``trtllm-llmapi-launch`` receive it -#: once at launch, so a value set later reaches the driver and not them -- and -#: a driver that resolves a modeling_v2 target while its workers resolve the -#: built-in is exactly the silent split this package exists to prevent. Export -#: it in the shell, or before ``import tensorrt_llm``. +#: Assigning to ``os.environ`` from a script does not set it. Worker ranks +#: receive the environment as it stood when MPI initialized, and long-lived +#: ranks under ``trtllm-llmapi-launch`` receive it once at launch, so a value +#: set later reaches the driver and not them -- and a driver that resolves a +#: modeling_v2 target while its workers resolve the built-in is exactly the +#: silent split this package exists to prevent. Two routes do reach the ranks: +#: export it in the shell before they start, or pass ``LLM(env_overrides={...})``, +#: which every rank re-applies to its own environment before it builds a model. MODELING_V2_ENV = "TRTLLM_MODELING_V2" # architectures[0] -> routing module, relative to this package. diff --git a/tests/integration/defs/accuracy/modeling_v2_env.py b/tests/integration/defs/accuracy/modeling_v2_env.py new file mode 100644 index 000000000000..fdb5dd9c677d --- /dev/null +++ b/tests/integration/defs/accuracy/modeling_v2_env.py @@ -0,0 +1,61 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Put the modeling_v2 switch in front of the ranks that resolve the model. + +Shared by the whole-model gates rather than copied into each of them: what +follows is easy to get subtly wrong, and two copies would drift. + +``TRTLLM_MODELING_V2`` is read by ``modeling_v2_resolve``, called from +``AutoModelForCausalLM._resolve_class`` -- which runs in a worker rank, not in +the process running the test. Assigning to ``os.environ`` here does not +reliably reach that rank: MPI caches the environment at import time and spawns +from that snapshot, which is why ``worker_main`` re-applies ``env_overrides`` +by hand at its top. So the switch travels as an ``LLM(...)`` field rather than +as a shell export, and it then reaches both sides -- ``LLM.__init__`` applies +the same overrides to the calling process before it checks the mode itself. + +These cases used to read the variable instead and skip unless the launching +shell had exported it. Nothing in CI exports it, so all of them skipped and the +stage reported green without ever measuring a modeling_v2 target -- the same +misattribution ``require`` mode exists to prevent, reached through the test +harness rather than through the model. +""" + +from typing import Any, Dict + +import pytest + +from tensorrt_llm._torch._experimental.modeling_v2 import MODELING_V2_ENV + + +def modeling_v2_llm_args( + mode: str, monkeypatch: pytest.MonkeyPatch, **extra: str +) -> Dict[str, Any]: + """``LLM(...)`` kwargs that build the model under ``TRTLLM_MODELING_V2=mode``. + + ``extra`` carries any further environment switch the case needs the ranks + to see; it travels by the same route. + + ``monkeypatch`` is not decoration. ``LLM`` applies its overrides to the + calling process as well and never puts them back, so with no teardown the + first case asking for ``"require"`` would leave every later case in the + session asking for it too -- and ``require`` raises on an architecture with + no target, so unrelated tests downstream would fail with this file's name + nowhere in the traceback. + """ + overrides = {MODELING_V2_ENV: mode, **extra} + for key, value in overrides.items(): + monkeypatch.setenv(key, value) + return {"env_overrides": overrides} diff --git a/tests/integration/defs/accuracy/test_modeling_v2_deepseek_v3.py b/tests/integration/defs/accuracy/test_modeling_v2_deepseek_v3.py index 8fdaaa199604..1b28045609dd 100644 --- a/tests/integration/defs/accuracy/test_modeling_v2_deepseek_v3.py +++ b/tests/integration/defs/accuracy/test_modeling_v2_deepseek_v3.py @@ -19,26 +19,21 @@ next to the built-in model's tests would invite treating one as a variant of the other. -Every test here needs ``TRTLLM_MODELING_V2=require``. Under ``"auto"`` a -configuration that missed a target's criteria would quietly fall back to the -built-in implementation, pass, and report the built-in's numbers as the -target's -- which is the one failure this whole system exists to prevent. The -one exception is the stock leg of the acceptance gate, which asks for ``"off"`` -on purpose. - -The switch is an environment variable, and worker ranks read it as it stood -when they started. These targets are multi-rank, so the ranks are already -running by the time a test body executes: a case cannot choose its own mode, it -can only assert that the environment it was given is the one it needs, which is -what ``_require_mode`` does. +Every test here runs under ``TRTLLM_MODELING_V2=require``, which the case puts +in place itself -- see ``modeling_v2_env``. Under ``"auto"`` a configuration +that missed a target's criteria would quietly fall back to the built-in +implementation, pass, and report the built-in's numbers as the target's -- +which is the one failure this whole system exists to prevent. The one exception +is the stock leg of the acceptance gate, which asks for ``"off"`` on purpose. + +These targets are multi-rank, and the mode has to hold on every rank rather +than only on the one running the test body; ``modeling_v2_env`` explains why +that rules out setting the variable directly. """ -import os - import pytest from tensorrt_llm import LLM -from tensorrt_llm._torch._experimental.modeling_v2 import MODELING_V2_ENV from tensorrt_llm._utils import get_sm_version from tensorrt_llm.llmapi import CudaGraphConfig, KvCacheConfig, MTPDecodingConfig @@ -49,6 +44,7 @@ assert_acceptance_length, compute_acceptance_length, ) +from .modeling_v2_env import modeling_v2_llm_args # The targets assert their own SM at construction: certification is per GPU # architecture, and a receipt from another one says nothing here. @@ -63,19 +59,6 @@ _SCORES_FILTER = {"scores_filter": "exact_match,flexible-extract"} -def _require_mode(expected: str) -> None: - """Skip unless the ranks were started with the mode this case needs. - - Not a failure: which mode a multi-rank job runs under is a property of how - it was launched, so a case that wants the other one has nothing to say. It - must not silently measure the wrong system either, which is what reading - the variable here rules out. - """ - actual = os.environ.get(MODELING_V2_ENV, "off") - if actual != expected: - pytest.skip(f"{MODELING_V2_ENV}={actual!r}, this case needs {expected!r}") - - class TestModelingV2DeepseekR10528Nvfp4Sm103Dep4(LlmapiAccuracyTestHarness): """deepseek-r1-0528-nvfp4 / sm_103 / dep4, identity and the mtp3 variant.""" @@ -106,7 +89,7 @@ class TestModelingV2DeepseekR10528Nvfp4Sm103Dep4(LlmapiAccuracyTestHarness): @skip_not_sm103 @pytest.mark.skip_less_device(4) - def test_gsm8k_identity_vs_mtp3(self, mocker): + def test_gsm8k_identity_vs_mtp3(self, mocker, monkeypatch): """The identity accuracy gate, and the gate on MTP not moving it. Turning MTP on must not move the answers. @@ -127,14 +110,18 @@ def test_gsm8k_identity_vs_mtp3(self, mocker): mocker.patch.dict(GSM8K.EVALUATE_KWARGS, _SCORES_FILTER) task = GSM8K(self.MODEL_NAME) - _require_mode("require") - with LLM(self.MODEL_PATH, **self.DEP4) as llm: + # One dict for both legs: the comparison only means anything if the two + # engines were built under the same mode. + modeling_v2 = modeling_v2_llm_args("require", monkeypatch) + + with LLM(self.MODEL_PATH, **modeling_v2, **self.DEP4) as llm: identity = task.evaluate(llm) with LLM( self.MODEL_PATH, speculative_config=self.MTP3, kv_cache_config=self.MTP3_KV, + **modeling_v2, **self.DEP4, ) as llm: mtp3 = task.evaluate(llm) @@ -151,7 +138,7 @@ def test_gsm8k_identity_vs_mtp3(self, mocker): @skip_not_sm103 @pytest.mark.skip_less_device(4) @pytest.mark.parametrize("mode", ["require", "off"], ids=["modeling_v2", "stock"]) - def test_mtp3_acceptance(self, mode, mocker): + def test_mtp3_acceptance(self, mode, mocker, monkeypatch): """The only gate that can see a miscomputed draft layer. Rejection sampling makes a wrong draft path *slower*, not wrong: every @@ -169,20 +156,16 @@ def test_mtp3_acceptance(self, mode, mocker): The anchor is populated from the **stock** leg. Populating it from the target's own number would make the gate self-referential. """ - _require_mode(mode) - if mode == "off": - # Stock cannot boot this checkpoint at dep4 with MTP otherwise: - # under attention DP + EP the MoE communication factory lands on - # DeepEPLowLatency, whose dispatch takes only NVFP4 uint8 hidden - # states, and the MTP layer is bf16 because modelopt excludes - # model.layers.61* from quantization. Disabling DeepEP lands on - # AllGatherReduceScatter -- which is the strategy the modeling_v2 - # target implements by hand, so it makes the two comparable rather - # than less so. Set in the launching environment, like the switch. - assert os.environ.get("TRTLLM_CAN_USE_DEEP_EP") == "0", ( - "the stock leg needs TRTLLM_CAN_USE_DEEP_EP=0 exported; without " - "it stock cannot boot this checkpoint at dep4 with MTP" - ) + # Stock cannot boot this checkpoint at dep4 with MTP otherwise: under + # attention DP + EP the MoE communication factory lands on + # DeepEPLowLatency, whose dispatch takes only NVFP4 uint8 hidden states, + # and the MTP layer is bf16 because modelopt excludes model.layers.61* + # from quantization. Disabling DeepEP lands on AllGatherReduceScatter -- + # which is the strategy the modeling_v2 target implements by hand, so it + # makes the two comparable rather than less so. Only this leg needs it: + # the target builds that path itself and never consults the factory. + # It rides along with the switch because it is read on the ranks too. + deep_ep = {"TRTLLM_CAN_USE_DEEP_EP": "0"} if mode == "off" else {} mocker.patch.dict(GSM8K.EVALUATE_KWARGS, _SCORES_FILTER) @@ -192,6 +175,7 @@ def test_mtp3_acceptance(self, mode, mocker): kv_cache_config=self.MTP3_KV, cuda_graph_config=CudaGraphConfig(), enable_iter_perf_stats=True, + **modeling_v2_llm_args(mode, monkeypatch, **deep_ep), **self.DEP4, ) as llm: task = GSM8K(self.MODEL_NAME) diff --git a/tests/integration/defs/accuracy/test_modeling_v2_gpt_oss.py b/tests/integration/defs/accuracy/test_modeling_v2_gpt_oss.py index 746335586630..03708a9037db 100644 --- a/tests/integration/defs/accuracy/test_modeling_v2_gpt_oss.py +++ b/tests/integration/defs/accuracy/test_modeling_v2_gpt_oss.py @@ -19,22 +19,21 @@ next to the built-in model's tests would invite treating one as a variant of the other. -Every test here needs ``TRTLLM_MODELING_V2=require``. Under ``"auto"`` a -configuration that missed a target's criteria would quietly fall back to the -built-in implementation, pass, and report the built-in's numbers as the -target's -- which is the one failure this whole system exists to prevent. +Every test here runs under ``TRTLLM_MODELING_V2=require``, which the case puts +in place itself -- see ``modeling_v2_env``. Under ``"auto"`` a configuration +that missed a target's criteria would quietly fall back to the built-in +implementation, pass, and report the built-in's numbers as the target's -- +which is the one failure this whole system exists to prevent. """ -import os - import pytest from tensorrt_llm import LLM -from tensorrt_llm._torch._experimental.modeling_v2 import MODELING_V2_ENV from tensorrt_llm._utils import get_sm_version from ..conftest import llm_models_root from .accuracy_core import GSM8K, LlmapiAccuracyTestHarness +from .modeling_v2_env import modeling_v2_llm_args # The targets assert their own SM at construction: certification is per GPU # architecture, and a receipt from another one says nothing here. @@ -43,19 +42,6 @@ ) -def _require_mode(expected: str) -> None: - """Skip unless the ranks were started with the mode this case needs. - - Not a failure: which mode a multi-rank job runs under is a property of how - it was launched, so a case that wants the other one has nothing to say. It - must not silently measure the wrong system either, which is what reading - the variable here rules out. - """ - actual = os.environ.get(MODELING_V2_ENV, "off") - if actual != expected: - pytest.skip(f"{MODELING_V2_ENV}={actual!r}, this case needs {expected!r}") - - class TestModelingV2GptOss120bSm103Tp1(LlmapiAccuracyTestHarness): """gpt-oss-120b / sm_103 / tp1.""" @@ -74,7 +60,7 @@ class TestModelingV2GptOss120bSm103Tp1(LlmapiAccuracyTestHarness): } @skip_not_sm103 - def test_gsm8k(self, mocker): + def test_gsm8k(self, mocker, monkeypatch): # Both patches are the protocol the anchor was measured under, and both # are what TestGPTOSS applies to this same checkpoint. The stock 256 # tokens truncate it mid-chain-of-thought, before it ever reaches an @@ -85,7 +71,6 @@ def test_gsm8k(self, mocker): mocker.patch.object(GSM8K, "MAX_OUTPUT_LEN", 8192) mocker.patch.dict(GSM8K.EVALUATE_KWARGS, {"scores_filter": "exact_match,flexible-extract"}) - _require_mode("require") - with LLM(self.MODEL_PATH) as llm: + with LLM(self.MODEL_PATH, **modeling_v2_llm_args("require", monkeypatch)) as llm: task = GSM8K(self.MODEL_NAME) task.evaluate(llm, extra_evaluator_kwargs=self.extra_evaluator_kwargs) From d43692f85eeb321ee4585a9010e33d3c42f11d31 Mon Sep 17 00:00:00 2001 From: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> Date: Mon, 21 Sep 2026 23:14:26 -0700 Subject: [PATCH 15/19] [TRTLLM-16304][test] Move the single-GPU modeling_v2 entries off the 4-GPU stage The catalog matrix runs 348 cases on one GPU and the gpt-oss gate is tp1, but both sat on l0_gb300.yml, and every stage reading that list is declared with four GPUs. The sub-job split routes by the "N_GPUs" pattern in the stage name rather than by what a test actually needs, so those two rode into L0_Test-SBSA-Multi-GPU, which pre-merge is gated on the 'ci: full pre-merge approved' label. Without the label the sub-job is blocked and nothing here runs at all -- build 60719 has no record of a single test_modeling_v2_* case. A single-GPU list on a single-GPU stage, named without "N_GPUs" so the split leaves it alone. It asks for one GPU rather than four, so it is also cheaper than what it replaces on hardware whose x4 capacity is the binding constraint. Shape copied from DGX_B200-PyTorch-1, the existing single-GPU SLURM stage on a -flex label. The two entries that genuinely need four ranks are untouched on l0_gb300_multi_gpus.yml, as are the nine tests left on l0_gb300.yml. The new list is deliberately not added to L0_MergeRequest's multi-GPU relatedFileList: everything there is a multi-GPU list, and editing a single-GPU one should not force the multi-GPU dispatch. Signed-off-by: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> --- jenkins/L0_Test.groovy | 4 ++ .../test_lists/test-db/l0_gb300.yml | 19 ++------- .../test-db/l0_gb300_single_gpu.yml | 40 +++++++++++++++++++ 3 files changed, 48 insertions(+), 15 deletions(-) create mode 100644 tests/integration/test_lists/test-db/l0_gb300_single_gpu.yml diff --git a/jenkins/L0_Test.groovy b/jenkins/L0_Test.groovy index f5a5ae6b722d..878fa45609ba 100644 --- a/jenkins/L0_Test.groovy +++ b/jenkins/L0_Test.groovy @@ -6439,6 +6439,10 @@ def launchTestJobs(pipeline, testFilter, globalVars) "GB200-4_GPUs-PyTorch-5": ["auto:gb200-x4", "l0_gb200_multi_gpus", 5, 5, 4, 1, false, true], "GB200-4_GPUs-PyTorch-Post-Merge-1": ["auto:gb200-x4", "l0_gb200_multi_gpus", 1, 1, 4, 1, false, true], "GB10-PyTorch-Post-Merge-1": ["gb10x-single", "l0_gb10", 1, 1], + // One GPU, so the name carries no "N_GPUs": the split at the bottom of + // this file routes by that pattern, and a single-GPU list has no reason + // to wait on the multi-GPU dispatch's approval label. + "GB300-PyTorch-1": ["auto:gb300-flex", "l0_gb300_single_gpu", 1, 1, 1, 1, true, false], "GB300-4_GPUs-PyTorch-1": ["auto:gb300-x4", "l0_gb300", 1, 1, 4, 1, true, false], "GB300-4_GPUs-PyTorch-Post-Merge-1": ["auto:gb300-x4", "l0_gb300_multi_gpus", 1, 3, 4, 1, true, false], "GB300-4_GPUs-PyTorch-Post-Merge-2": ["auto:gb300-x4", "l0_gb300_multi_gpus", 2, 3, 4, 1, true, false], diff --git a/tests/integration/test_lists/test-db/l0_gb300.yml b/tests/integration/test_lists/test-db/l0_gb300.yml index 12790fb33ad0..d1be35a0c14a 100644 --- a/tests/integration/test_lists/test-db/l0_gb300.yml +++ b/tests/integration/test_lists/test-db/l0_gb300.yml @@ -26,18 +26,7 @@ l0_gb300: - accuracy/test_disaggregated_serving.py::TestQwen3_8_Flash_Next::test_fp8_nixl_python[prefix_cache] - unittest/_torch/thop/parallel TIMEOUT (90) - unittest/_torch/visual_gen/kernels/parallel - # ModelingV2 catalog: the certification matrix for each op a modeling_v2 target - # calls, plus the two consistency tests that guard the routing tables against - # the targets they name. Deliberately NOT a duplicate of the upstream tests - # for the same ops -- those cover the modules that wrap them, these cover the - # op itself cell by cell -- which is why every file is named test_modeling_v2_*. - # These entries are the receipts: the catalog's certification is per GPU - # architecture, and this list is the sm_103 one. comm/ is excluded here and - # carried by l0_gb300_multi_gpus.yml, since those two need 4 ranks. - # TIMEOUT measured, not guessed: the whole entry is 348 cases in 5m34s on one - # GB300. 30 gives five times that -- enough for a loaded node, short enough - # that a hung case does not hold this stage for its full budget. - - unittest/_torch/modeling_v2 --ignore=unittest/_torch/modeling_v2/comm TIMEOUT (30) - # ModelingV2 whole-model gate. The op-level entry above certifies the - # vocabulary; this certifies the assembly that calls it. - - accuracy/test_modeling_v2_gpt_oss.py::TestModelingV2GptOss120bSm103Tp1::test_gsm8k + # The modeling_v2 sm_103 entries live on l0_gb300_single_gpu.yml: they need one + # GPU, and every stage reading this list declares four, which puts them behind + # the multi-GPU approval label for no reason. The two that do need 4 ranks are + # on l0_gb300_multi_gpus.yml. diff --git a/tests/integration/test_lists/test-db/l0_gb300_single_gpu.yml b/tests/integration/test_lists/test-db/l0_gb300_single_gpu.yml new file mode 100644 index 000000000000..cd991cf00e6b --- /dev/null +++ b/tests/integration/test_lists/test-db/l0_gb300_single_gpu.yml @@ -0,0 +1,40 @@ +version: 0.0.1 +l0_gb300_single_gpu: +- condition: + ranges: + system_gpu_count: + gte: 1 + lte: 1 + wildcards: + gpu: + - '*gb110*' + - '*gb300*' + linux_distribution_name: ubuntu* + cpu: aarch64 + terms: + stage: pre_merge + backend: pytorch + tests: + # The sm_103 entries that need one GPU and nothing else. They sat on + # l0_gb300.yml until it became clear what that costs: every GB300 stage is + # declared with four GPUs, the multi-GPU dispatch is gated on the + # 'ci: full pre-merge approved' label, and so a single-GPU test list was + # waiting on an approval it has no reason to need. Nothing here is a + # multi-rank test; the two that are live on l0_gb300_multi_gpus.yml. + # + # ModelingV2 catalog: the certification matrix for each op a modeling_v2 target + # calls, plus the two consistency tests that guard the routing tables against + # the targets they name. Deliberately NOT a duplicate of the upstream tests + # for the same ops -- those cover the modules that wrap them, these cover the + # op itself cell by cell -- which is why every file is named test_modeling_v2_*. + # These entries are the receipts: the catalog's certification is per GPU + # architecture, and this list is the sm_103 one. comm/ is excluded here and + # carried by l0_gb300_multi_gpus.yml, since those two need 4 ranks. + # TIMEOUT measured, not guessed: the whole entry is 348 cases in 5m34s on one + # GB300. 30 gives five times that -- enough for a loaded node, short enough + # that a hung case does not hold this stage for its full budget. + - unittest/_torch/modeling_v2 --ignore=unittest/_torch/modeling_v2/comm TIMEOUT (30) + # ModelingV2 whole-model gate. The op-level entry above certifies the + # vocabulary; this certifies the assembly that calls it. tp1, so it belongs + # here rather than beside the dep4 gates. + - accuracy/test_modeling_v2_gpt_oss.py::TestModelingV2GptOss120bSm103Tp1::test_gsm8k From ff842472498442e336e9466b12213747449bb221 Mon Sep 17 00:00:00 2001 From: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> Date: Tue, 22 Sep 2026 18:22:21 -0700 Subject: [PATCH 16/19] [TRTLLM-16304][test] Put the single-GPU modeling_v2 entries on a stage that accepts them The GB300-PyTorch-1 stage added in 341dbc1a6 never ran. Slurm rejects it before pytest starts: #SBATCH --nodes=1 #SBATCH --gpus-per-node=1 sbatch: error: QOSMinGRES sbatch: error: Batch job submission failed: Job violates accounting/QOS policy The GB300 QOS has a minimum GRES and a one-GPU job is under it, which is why every GB300 stage in this file asks for four. There is no single-GPU GB300 stage to add; the premise of that commit was wrong. The routing half of it did work -- the stage landed in L0_Test-SBSA-Single-GPU unblocked -- but the allocation never happened, and CBTS had narrowed builds 61754 and 61819 to a stage set in which this was the only GPU stage, so those runs tested nothing. l0_b300.yml is the single-GPU list for the same silicon: it and l0_gb300.yml both match '*gb110*', the Blackwell Ultra die, so both are sm_103 and a receipt from either is the same architecture's. Its stages ask for one GPU on an x86 cluster whose QOS accepts that, and they run today. So the two single-GPU entries go there instead, and both the new list and the new stage are dropped. The four entries that genuinely need four ranks are untouched on l0_gb300_multi_gpus.yml. Checked with scripts/test_to_stage_mapping.py: both entries now resolve to B300-PyTorch-1/2 rather than to a stage that cannot start. What this cannot check is that they pass on an x86 host -- the kernels are the same sm_103 cubins, but that is an expectation until CI runs it. Signed-off-by: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> --- jenkins/L0_Test.groovy | 4 -- .../test_lists/test-db/l0_b300.yml | 24 +++++++++++ .../test_lists/test-db/l0_gb300.yml | 10 +++-- .../test-db/l0_gb300_single_gpu.yml | 40 ------------------- 4 files changed, 30 insertions(+), 48 deletions(-) delete mode 100644 tests/integration/test_lists/test-db/l0_gb300_single_gpu.yml diff --git a/jenkins/L0_Test.groovy b/jenkins/L0_Test.groovy index 878fa45609ba..f5a5ae6b722d 100644 --- a/jenkins/L0_Test.groovy +++ b/jenkins/L0_Test.groovy @@ -6439,10 +6439,6 @@ def launchTestJobs(pipeline, testFilter, globalVars) "GB200-4_GPUs-PyTorch-5": ["auto:gb200-x4", "l0_gb200_multi_gpus", 5, 5, 4, 1, false, true], "GB200-4_GPUs-PyTorch-Post-Merge-1": ["auto:gb200-x4", "l0_gb200_multi_gpus", 1, 1, 4, 1, false, true], "GB10-PyTorch-Post-Merge-1": ["gb10x-single", "l0_gb10", 1, 1], - // One GPU, so the name carries no "N_GPUs": the split at the bottom of - // this file routes by that pattern, and a single-GPU list has no reason - // to wait on the multi-GPU dispatch's approval label. - "GB300-PyTorch-1": ["auto:gb300-flex", "l0_gb300_single_gpu", 1, 1, 1, 1, true, false], "GB300-4_GPUs-PyTorch-1": ["auto:gb300-x4", "l0_gb300", 1, 1, 4, 1, true, false], "GB300-4_GPUs-PyTorch-Post-Merge-1": ["auto:gb300-x4", "l0_gb300_multi_gpus", 1, 3, 4, 1, true, false], "GB300-4_GPUs-PyTorch-Post-Merge-2": ["auto:gb300-x4", "l0_gb300_multi_gpus", 2, 3, 4, 1, true, false], diff --git a/tests/integration/test_lists/test-db/l0_b300.yml b/tests/integration/test_lists/test-db/l0_b300.yml index 3a0e50c09bae..2c293f535d9c 100644 --- a/tests/integration/test_lists/test-db/l0_b300.yml +++ b/tests/integration/test_lists/test-db/l0_b300.yml @@ -119,6 +119,30 @@ l0_b300: - accuracy/test_llm_api_pytorch.py::TestDeepSeekV3Lite::test_chunked_prefill[quant_dtype=fp8-kv_cache_reuse=True-fp8kv=True-overlap_scheduler=True] # Cover nvbugs 6084445 - accuracy/test_llm_api_pytorch.py::TestDeepSeekV3Lite::test_bfloat16_flashinfer[enable_chunked_prefill=False] - accuracy/test_llm_api_pytorch.py::TestDeepSeekV3Lite::test_bfloat16_flashinfer[enable_chunked_prefill=True] + # ModelingV2 catalog: the certification matrix for each op a modeling_v2 target + # calls, plus the two consistency tests that guard the routing tables against + # the targets they name. Deliberately NOT a duplicate of the upstream tests + # for the same ops -- those cover the modules that wrap them, these cover the + # op itself cell by cell -- which is why every file is named test_modeling_v2_*. + # + # These entries are the receipts, and the catalog's certification is per GPU + # architecture rather than per board: this list and l0_gb300.yml both match + # '*gb110*', the Blackwell Ultra die, so both are sm_103 and this is where the + # sm_103 receipt can actually be collected. They sit here rather than beside + # the dep4 entries on l0_gb300.yml because they need one GPU and the GB300 QOS + # rejects a single-GPU job outright (`sbatch: error: QOSMinGRES`); every GB300 + # stage is four GPUs, which would put a single-GPU test list behind the + # 'ci: full pre-merge approved' label for no reason. The entries that do need + # four ranks stay on l0_gb300_multi_gpus.yml. + # + # TIMEOUT measured, not guessed: the whole entry is 348 cases in 5m34s on one + # GB300. 30 gives five times that -- enough for a loaded node, short enough + # that a hung case does not hold this stage for its full budget. + - unittest/_torch/modeling_v2 --ignore=unittest/_torch/modeling_v2/comm TIMEOUT (30) + # ModelingV2 whole-model gate. The op-level entry above certifies the + # vocabulary; this certifies the assembly that calls it. tp1, so it belongs on + # a single-GPU list rather than beside the dep4 gates. + - accuracy/test_modeling_v2_gpt_oss.py::TestModelingV2GptOss120bSm103Tp1::test_gsm8k - condition: ranges: system_gpu_count: diff --git a/tests/integration/test_lists/test-db/l0_gb300.yml b/tests/integration/test_lists/test-db/l0_gb300.yml index d1be35a0c14a..82e360bfe426 100644 --- a/tests/integration/test_lists/test-db/l0_gb300.yml +++ b/tests/integration/test_lists/test-db/l0_gb300.yml @@ -26,7 +26,9 @@ l0_gb300: - accuracy/test_disaggregated_serving.py::TestQwen3_8_Flash_Next::test_fp8_nixl_python[prefix_cache] - unittest/_torch/thop/parallel TIMEOUT (90) - unittest/_torch/visual_gen/kernels/parallel - # The modeling_v2 sm_103 entries live on l0_gb300_single_gpu.yml: they need one - # GPU, and every stage reading this list declares four, which puts them behind - # the multi-GPU approval label for no reason. The two that do need 4 ranks are - # on l0_gb300_multi_gpus.yml. + # The single-GPU modeling_v2 sm_103 entries live on l0_b300.yml. Same + # Blackwell Ultra die ('*gb110*' matches on both lists), so the receipt is the + # same architecture's, and that list runs on a stage that accepts a one-GPU + # job -- the GB300 QOS does not (`sbatch: error: QOSMinGRES`), so every stage + # reading this list declares four. The entries that genuinely need four ranks + # are on l0_gb300_multi_gpus.yml. diff --git a/tests/integration/test_lists/test-db/l0_gb300_single_gpu.yml b/tests/integration/test_lists/test-db/l0_gb300_single_gpu.yml deleted file mode 100644 index cd991cf00e6b..000000000000 --- a/tests/integration/test_lists/test-db/l0_gb300_single_gpu.yml +++ /dev/null @@ -1,40 +0,0 @@ -version: 0.0.1 -l0_gb300_single_gpu: -- condition: - ranges: - system_gpu_count: - gte: 1 - lte: 1 - wildcards: - gpu: - - '*gb110*' - - '*gb300*' - linux_distribution_name: ubuntu* - cpu: aarch64 - terms: - stage: pre_merge - backend: pytorch - tests: - # The sm_103 entries that need one GPU and nothing else. They sat on - # l0_gb300.yml until it became clear what that costs: every GB300 stage is - # declared with four GPUs, the multi-GPU dispatch is gated on the - # 'ci: full pre-merge approved' label, and so a single-GPU test list was - # waiting on an approval it has no reason to need. Nothing here is a - # multi-rank test; the two that are live on l0_gb300_multi_gpus.yml. - # - # ModelingV2 catalog: the certification matrix for each op a modeling_v2 target - # calls, plus the two consistency tests that guard the routing tables against - # the targets they name. Deliberately NOT a duplicate of the upstream tests - # for the same ops -- those cover the modules that wrap them, these cover the - # op itself cell by cell -- which is why every file is named test_modeling_v2_*. - # These entries are the receipts: the catalog's certification is per GPU - # architecture, and this list is the sm_103 one. comm/ is excluded here and - # carried by l0_gb300_multi_gpus.yml, since those two need 4 ranks. - # TIMEOUT measured, not guessed: the whole entry is 348 cases in 5m34s on one - # GB300. 30 gives five times that -- enough for a loaded node, short enough - # that a hung case does not hold this stage for its full budget. - - unittest/_torch/modeling_v2 --ignore=unittest/_torch/modeling_v2/comm TIMEOUT (30) - # ModelingV2 whole-model gate. The op-level entry above certifies the - # vocabulary; this certifies the assembly that calls it. tp1, so it belongs - # here rather than beside the dep4 gates. - - accuracy/test_modeling_v2_gpt_oss.py::TestModelingV2GptOss120bSm103Tp1::test_gsm8k From 652be0a40ca29828dd943eb9da938c5d9ef745f7 Mon Sep 17 00:00:00 2001 From: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> Date: Wed, 23 Sep 2026 01:38:48 -0700 Subject: [PATCH 17/19] [TRTLLM-16304][test] Take the prose back out of the test-db lists Forty lines of commentary for six entries, in files whose other entries are one line each. Per review: a test list is an index, and the explanation belongs with the test. It was also duplication rather than documentation. The collective entry's reason for spawning its own 4-rank job is already in `_rank_job` and in the comment at the top of the matrix files that use it; the acceptance pair's shared-anchor logic is already in `test_mtp3_acceptance`'s own docstring, in more detail than the copy here. A second copy in a file nobody edits when the test changes is a copy that goes stale. What is left is the one thing that is about the list entry and is not recoverable from the test: where the TIMEOUT number came from. l0_gb300.yml's block is deleted outright rather than shortened. It explained which entries are *not* on that list, which is not something a list should carry. Comments only -- no entry added, removed or edited, and scripts/test_to_stage_mapping.py resolves the same stages as before. Signed-off-by: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> --- .../test_lists/test-db/l0_b300.yml | 23 +------------------ .../test_lists/test-db/l0_gb300.yml | 6 ----- .../test-db/l0_gb300_multi_gpus.yml | 13 +---------- 3 files changed, 2 insertions(+), 40 deletions(-) diff --git a/tests/integration/test_lists/test-db/l0_b300.yml b/tests/integration/test_lists/test-db/l0_b300.yml index 2c293f535d9c..14307de1b9dc 100644 --- a/tests/integration/test_lists/test-db/l0_b300.yml +++ b/tests/integration/test_lists/test-db/l0_b300.yml @@ -119,29 +119,8 @@ l0_b300: - accuracy/test_llm_api_pytorch.py::TestDeepSeekV3Lite::test_chunked_prefill[quant_dtype=fp8-kv_cache_reuse=True-fp8kv=True-overlap_scheduler=True] # Cover nvbugs 6084445 - accuracy/test_llm_api_pytorch.py::TestDeepSeekV3Lite::test_bfloat16_flashinfer[enable_chunked_prefill=False] - accuracy/test_llm_api_pytorch.py::TestDeepSeekV3Lite::test_bfloat16_flashinfer[enable_chunked_prefill=True] - # ModelingV2 catalog: the certification matrix for each op a modeling_v2 target - # calls, plus the two consistency tests that guard the routing tables against - # the targets they name. Deliberately NOT a duplicate of the upstream tests - # for the same ops -- those cover the modules that wrap them, these cover the - # op itself cell by cell -- which is why every file is named test_modeling_v2_*. - # - # These entries are the receipts, and the catalog's certification is per GPU - # architecture rather than per board: this list and l0_gb300.yml both match - # '*gb110*', the Blackwell Ultra die, so both are sm_103 and this is where the - # sm_103 receipt can actually be collected. They sit here rather than beside - # the dep4 entries on l0_gb300.yml because they need one GPU and the GB300 QOS - # rejects a single-GPU job outright (`sbatch: error: QOSMinGRES`); every GB300 - # stage is four GPUs, which would put a single-GPU test list behind the - # 'ci: full pre-merge approved' label for no reason. The entries that do need - # four ranks stay on l0_gb300_multi_gpus.yml. - # - # TIMEOUT measured, not guessed: the whole entry is 348 cases in 5m34s on one - # GB300. 30 gives five times that -- enough for a loaded node, short enough - # that a hung case does not hold this stage for its full budget. + # TIMEOUT measured: 348 cases in 5m34s on one GB300, so 30 is 5x headroom. - unittest/_torch/modeling_v2 --ignore=unittest/_torch/modeling_v2/comm TIMEOUT (30) - # ModelingV2 whole-model gate. The op-level entry above certifies the - # vocabulary; this certifies the assembly that calls it. tp1, so it belongs on - # a single-GPU list rather than beside the dep4 gates. - accuracy/test_modeling_v2_gpt_oss.py::TestModelingV2GptOss120bSm103Tp1::test_gsm8k - condition: ranges: diff --git a/tests/integration/test_lists/test-db/l0_gb300.yml b/tests/integration/test_lists/test-db/l0_gb300.yml index 82e360bfe426..40c08cb821a6 100644 --- a/tests/integration/test_lists/test-db/l0_gb300.yml +++ b/tests/integration/test_lists/test-db/l0_gb300.yml @@ -26,9 +26,3 @@ l0_gb300: - accuracy/test_disaggregated_serving.py::TestQwen3_8_Flash_Next::test_fp8_nixl_python[prefix_cache] - unittest/_torch/thop/parallel TIMEOUT (90) - unittest/_torch/visual_gen/kernels/parallel - # The single-GPU modeling_v2 sm_103 entries live on l0_b300.yml. Same - # Blackwell Ultra die ('*gb110*' matches on both lists), so the receipt is the - # same architecture's, and that list runs on a stage that accepts a one-GPU - # job -- the GB300 QOS does not (`sbatch: error: QOSMinGRES`), so every stage - # reading this list declares four. The entries that genuinely need four ranks - # are on l0_gb300_multi_gpus.yml. diff --git a/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml b/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml index 88d382b928d5..25474411ecb9 100644 --- a/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml +++ b/tests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml @@ -15,19 +15,8 @@ l0_gb300_multi_gpus: stage: post_merge backend: pytorch tests: - # ModelingV2 collective entries. Each spawns its own 4-rank job rather than - # using the mpi_pool_executor fixture: their cases read module-global rank - # state and several assert on communicator state the previous case left - # behind, so they are one ordered sequence inside one job, and a broken - # collective hangs rather than raising -- the launcher's deadline is what - # keeps that from wedging the run, and the fixture has none. + # One ordered sequence in one job: the cases share module-global rank state. - unittest/_torch/modeling_v2/comm - # ModelingV2 whole-model gates for the dep4 target. The mtp3 variant selects a - # second forward path AND a second weight-loading path, so the identity gate - # does not speak for it: it carries its own accuracy pairing and its own - # acceptance gate. The two acceptance legs are independent cases read against - # one shared minimum -- modeling_v2 failing while stock passes means the draft - # path regressed; both failing means the anchor is stale. - accuracy/test_modeling_v2_deepseek_v3.py::TestModelingV2DeepseekR10528Nvfp4Sm103Dep4::test_gsm8k_identity_vs_mtp3 - accuracy/test_modeling_v2_deepseek_v3.py::TestModelingV2DeepseekR10528Nvfp4Sm103Dep4::test_mtp3_acceptance[modeling_v2] - accuracy/test_modeling_v2_deepseek_v3.py::TestModelingV2DeepseekR10528Nvfp4Sm103Dep4::test_mtp3_acceptance[stock] From e1c23f4de287c550eb1fb0c257828cad472b6c5c Mon Sep 17 00:00:00 2001 From: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> Date: Wed, 23 Sep 2026 03:07:03 -0700 Subject: [PATCH 18/19] [TRTLLM-16304][chore] Flatten a target's path into one directory name `models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/` was six levels, and three of them had exactly one child and an `__init__.py` whose only job was to exist. It read as a hierarchy that was never one: there is no sense in which targets group under a checkpoint, then under an SM. The triple is flat, and so is the name now -- models/gpt_oss/gpt_oss_120b__sm_103__tp1/ Identity is still the whole triple, which is the part worth keeping: a target is one self-contained forward per (checkpoint, GPU arch, parallel), and the directory says which. It is just one segment instead of three nested ones. `targets/` goes with them. It separated `routing.py` from the targets at a level that held one file, which is not a separation worth a directory. `test_modeling_v2_claims.py` reads the triple back out of the module path to check it against the class name; it now splits one segment on `__` rather than walking three parents. That check is why this is a rename and not a rewrite -- it caught nothing here because it was updated in the same commit, but it is what keeps the path and the class name from drifting apart later. Verified on GB300: the three no-GPU gates pass (42), and both target modules import through their new paths. `_TARGETS` keys are (checkpoint, parallel) and carry no path, and CODEOWNERS and the CBTS rule both match on the `modeling_v2/` prefix, so none of the three needed touching. Signed-off-by: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> --- .../_torch/_experimental/modeling_v2/README.md | 12 +++++++----- .../__init__.py | 0 .../modeling.py | 0 .../dep4 => r1_0528_nvfp4__sm_103__dep4}/weights.py | 0 .../modeling_v2/models/deepseek_v3/routing.py | 2 +- .../models/deepseek_v3/targets/__init__.py | 3 --- .../deepseek_v3/targets/r1_0528_nvfp4/__init__.py | 3 --- .../targets/r1_0528_nvfp4/sm_103/__init__.py | 3 --- .../tp1 => gpt_oss_120b__sm_103__tp1}/__init__.py | 0 .../tp1 => gpt_oss_120b__sm_103__tp1}/modeling.py | 0 .../tp1 => gpt_oss_120b__sm_103__tp1}/weights.py | 0 .../modeling_v2/models/gpt_oss/routing.py | 2 +- .../modeling_v2/models/gpt_oss/targets/__init__.py | 3 --- .../models/gpt_oss/targets/gpt_oss_120b/__init__.py | 3 --- .../gpt_oss/targets/gpt_oss_120b/sm_103/__init__.py | 3 --- .../_torch/modeling_v2/test_modeling_v2_claims.py | 13 +++++++++---- .../_torch/modeling_v2/test_modeling_v2_routing.py | 2 +- 17 files changed, 19 insertions(+), 30 deletions(-) rename tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/{targets/r1_0528_nvfp4/sm_103/dep4 => r1_0528_nvfp4__sm_103__dep4}/__init__.py (100%) rename tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/{targets/r1_0528_nvfp4/sm_103/dep4 => r1_0528_nvfp4__sm_103__dep4}/modeling.py (100%) rename tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/{targets/r1_0528_nvfp4/sm_103/dep4 => r1_0528_nvfp4__sm_103__dep4}/weights.py (100%) delete mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/__init__.py delete mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/__init__.py delete mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/__init__.py rename tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/{targets/gpt_oss_120b/sm_103/tp1 => gpt_oss_120b__sm_103__tp1}/__init__.py (100%) rename tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/{targets/gpt_oss_120b/sm_103/tp1 => gpt_oss_120b__sm_103__tp1}/modeling.py (100%) rename tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/{targets/gpt_oss_120b/sm_103/tp1 => gpt_oss_120b__sm_103__tp1}/weights.py (100%) delete mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/__init__.py delete mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/__init__.py delete mode 100644 tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/__init__.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/README.md b/tensorrt_llm/_torch/_experimental/modeling_v2/README.md index 0019e050785f..61be49f5be0d 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/README.md +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/README.md @@ -100,8 +100,8 @@ _router_index.py architectures[0] -> routing module. Small, stable, test-guar explain.py why a configuration routed where it did models// routing.py one forward-reading decision tree per architecture family - targets//// - modeling.py weights.py + ____/ + modeling.py weights.py catalog/ the kernel vocabulary: contract .md + wrapper .py ``` @@ -132,9 +132,11 @@ by absolute import -- so neither needs a package to live in. Neither those file names nor their `check_*` bodies match pytest's collection patterns: each is one fixed 4-rank sequence that cannot run as independent cases. -Identity is the path. `targets/` keeps all three segments rather than -flattening them, and the class name carries the same triple; -`test_modeling_v2_claims.py` asserts they agree. +Identity is the directory name, and it carries all three segments: +`gpt_oss_120b__sm_103__tp1`. They were three nested directories once, which +read as a hierarchy that was never one -- every level had exactly one child, +and each needed an `__init__.py` whose only job was to exist. The class name +carries the same triple; `test_modeling_v2_claims.py` asserts they agree. ### Why beside `_torch/models/`, not inside it diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/r1_0528_nvfp4__sm_103__dep4/__init__.py similarity index 100% rename from tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/__init__.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/r1_0528_nvfp4__sm_103__dep4/__init__.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/r1_0528_nvfp4__sm_103__dep4/modeling.py similarity index 100% rename from tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/modeling.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/r1_0528_nvfp4__sm_103__dep4/modeling.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/weights.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/r1_0528_nvfp4__sm_103__dep4/weights.py similarity index 100% rename from tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/dep4/weights.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/r1_0528_nvfp4__sm_103__dep4/weights.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/routing.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/routing.py index cfecce68b334..d386db19bd9e 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/routing.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/routing.py @@ -36,7 +36,7 @@ # Synthetic architecture name -> the module whose import registers it. TARGET_MODULES = { - "ModelingV2DeepseekR10528Nvfp4Sm103Dep4": "models.deepseek_v3.targets.r1_0528_nvfp4.sm_103.dep4.modeling", + "ModelingV2DeepseekR10528Nvfp4Sm103Dep4": "models.deepseek_v3.r1_0528_nvfp4__sm_103__dep4.modeling", } diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/__init__.py deleted file mode 100644 index d7eb0af07fc9..000000000000 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -"""Targets, keyed by the // identity path.""" diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/__init__.py deleted file mode 100644 index c2bed1f3a52a..000000000000 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -"""DeepSeek-R1-0528 NVFP4 targets.""" diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/__init__.py deleted file mode 100644 index cc5208da61da..000000000000 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/deepseek_v3/targets/r1_0528_nvfp4/sm_103/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -"""DeepSeek-R1-0528-NVFP4 on sm_103 (GB300).""" diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/gpt_oss_120b__sm_103__tp1/__init__.py similarity index 100% rename from tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/__init__.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/gpt_oss_120b__sm_103__tp1/__init__.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/gpt_oss_120b__sm_103__tp1/modeling.py similarity index 100% rename from tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/modeling.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/gpt_oss_120b__sm_103__tp1/modeling.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/weights.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/gpt_oss_120b__sm_103__tp1/weights.py similarity index 100% rename from tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/tp1/weights.py rename to tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/gpt_oss_120b__sm_103__tp1/weights.py diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/routing.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/routing.py index d3476f008165..9855495c22f6 100644 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/routing.py +++ b/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/routing.py @@ -37,7 +37,7 @@ # Synthetic architecture name -> the module whose import registers it. TARGET_MODULES = { - "ModelingV2GptOss120bSm103Tp1": "models.gpt_oss.targets.gpt_oss_120b.sm_103.tp1.modeling", + "ModelingV2GptOss120bSm103Tp1": "models.gpt_oss.gpt_oss_120b__sm_103__tp1.modeling", } diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/__init__.py deleted file mode 100644 index d7eb0af07fc9..000000000000 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -"""Targets, keyed by the // identity path.""" diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/__init__.py deleted file mode 100644 index b2a58d2d0fae..000000000000 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -"""gpt-oss-120b targets.""" diff --git a/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/__init__.py b/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/__init__.py deleted file mode 100644 index 4ea8999b15d3..000000000000 --- a/tensorrt_llm/_torch/_experimental/modeling_v2/models/gpt_oss/targets/gpt_oss_120b/sm_103/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -"""gpt-oss-120b on sm_103 (GB300).""" diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py index 8d533738a2c7..b2eff05da849 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py @@ -87,12 +87,12 @@ def test_each_target_module_registers_its_own_name(arch): @pytest.mark.parametrize("arch", _ARCHS) def test_target_identity_matches_its_path(arch): - """Identity is the path: //. + """Identity is the directory name: ____. The class name encodes the same triple, and the routing module's ``_SM`` has to be the arch segment those targets actually live under -- a target - moved to a new SM directory without its routing constant following is the - one drift that would still route, and route wrong. + moved to a new SM without its routing constant following is the one drift + that would still route, and route wrong. """ routing = routing_module(arch) major, minor = routing._SM @@ -101,7 +101,12 @@ def test_target_identity_matches_its_path(arch): for name, dotted in routing.TARGET_MODULES.items(): parts = dotted.split(".") assert parts[-1] == "modeling", dotted - parallel, sm_segment, checkpoint = parts[-2], parts[-3], parts[-4] + segments = parts[-2].split("__") + assert len(segments) == 3, ( + f"{arch}: {parts[-2]!r} is not a ____ " + f"directory name" + ) + checkpoint, sm_segment, parallel = segments assert sm_segment == expected_segment, ( f"{arch}: {name} lives under {sm_segment} but its routing module " diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py index 920c14c027b9..3bf3423bfc33 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py @@ -140,7 +140,7 @@ def test_resolving_registers_the_target_class(): assert cls is not None, f"{name} resolved to no class" assert cls.__name__ == name assert cls.__module__.endswith( - "modeling_v2.models.gpt_oss.targets.gpt_oss_120b.sm_103.tp1.modeling" + "modeling_v2.models.gpt_oss.gpt_oss_120b__sm_103__tp1.modeling" ) From 8622232bae85fe75b7a9e12593ca838792313232 Mon Sep 17 00:00:00 2001 From: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> Date: Wed, 23 Sep 2026 19:06:42 -0700 Subject: [PATCH 19/19] [TRTLLM-16304][chore] Rewrap two lines the flattened path made short enough to join ruff-format, not a behaviour change. The flattened directory name is shorter than the three nested ones it replaced, so two strings that needed wrapping before now fit inside the 100-column limit, and the hook joins them. Missed because the previous commit was checked with `ruff check` and not `ruff format --check`. The two are separate hooks and only the first was run. Signed-off-by: Fred Wei <20514172+WeiHaocheng@users.noreply.github.com> --- tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py | 3 +-- tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py | 4 +--- 2 files changed, 2 insertions(+), 5 deletions(-) diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py index b2eff05da849..1f7e0435d5b3 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_claims.py @@ -103,8 +103,7 @@ def test_target_identity_matches_its_path(arch): assert parts[-1] == "modeling", dotted segments = parts[-2].split("__") assert len(segments) == 3, ( - f"{arch}: {parts[-2]!r} is not a ____ " - f"directory name" + f"{arch}: {parts[-2]!r} is not a ____ directory name" ) checkpoint, sm_segment, parallel = segments diff --git a/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py b/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py index 3bf3423bfc33..f40cb6428713 100644 --- a/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py +++ b/tests/unittest/_torch/modeling_v2/test_modeling_v2_routing.py @@ -139,9 +139,7 @@ def test_resolving_registers_the_target_class(): cls = get_registered_model_class(name) assert cls is not None, f"{name} resolved to no class" assert cls.__name__ == name - assert cls.__module__.endswith( - "modeling_v2.models.gpt_oss.gpt_oss_120b__sm_103__tp1.modeling" - ) + assert cls.__module__.endswith("modeling_v2.models.gpt_oss.gpt_oss_120b__sm_103__tp1.modeling") def test_the_target_registration_counts_as_external():