Skip to content

feat(trtllm): Qwen3.5 routed-expert block-FP8/MXFP8 refit - #4122

Open
hchings wants to merge 11 commits into
mainfrom
erinh/trtllm-mlperf-fp8
Open

hchings wants to merge 11 commits into
mainfrom
erinh/trtllm-mlperf-fp8

Conversation

@hchings

@hchings hchings commented Sep 14, 2026

Copy link
Copy Markdown
Contributor

What does this PR do ?

Add a one line overview of what this PR aims to accomplish.

Issues

List issues that this PR closes (syntax):

Usage

  • You can potentially add a usage example below
# Add a code snippet demonstrating how to use this

Before your PR is "Ready for review"

Pre checks:

  • Make sure you read and followed Contributor guidelines
  • Did you write any new necessary tests?
  • Did you run the unit tests and functional tests locally? Visit our Testing Guide for how to run tests
  • Did you add or update any necessary documentation? Visit our Document Development Guide for how to write, build and test the docs.

Additional Information

  • ...

@hchings hchings self-assigned this Sep 14, 2026
@copy-pr-bot

copy-pr-bot Bot commented Sep 14, 2026

Copy link
Copy Markdown

This pull request requires additional validation before any workflows can run on NVIDIA's runners.

Pull request vetters can view their responsibilities here.

Contributors can view more details about this message here.

@github-actions github-actions Bot added the Documentation Improvements or additions to documentation label Sep 14, 2026

@hchings hchings left a comment

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Team review: Qwen3.5 routed-expert block-FP8 / MXFP8 refit

Reviewed by a team of 4 agents plus a leader re-verification pass over the full diff, with every upstream claim checked against the TensorRT-LLM commit this repo pins rather than main. 8 comments below.

The quantization module itself is good, and that's the bulk of the work here. Worth saying explicitly:

  • quantization/fp8.py imports only math/re/torch and takes TRT-LLM types as parameters (configure_fp8_moe_backend(llm_kwargs, moe_config_type, ...)) rather than importing them. That's why the 223-line test file exercises the real quantization math on CPU with a 6-line stub and no engine mocking — a strictly better shape than the vllm/quantization/fp8.py analog, which routes the same decisions through a module-level mutable global.
  • We checked the numerics against the pinned consumer and they're right: the VANILLA per-expert split is what Qwen3_5MoeHfWeightMapper expects, the [ceil(N/128), ceil(K/128)] FP32 scale orientation and dequant = fp8 * scale_inv match create_weights, the MXFP8 plain-uint8 [..., K/32] layout matches MXFP8CutlassFusedMoEMethod (which does its own packing and swizzle, so not pre-swizzling is correct), and all 13 modules_to_not_convert patterns translate correctly through upstream's Qwen3.5 normalizer. The block-FP8 caster was validated on CPU across aligned, unaligned, batched and non-contiguous inputs.
  • The multi-worker result aggregation fix in trtllm_worker_async.py is a genuine improvement — the old results[0] masked failures on ranks > 0, and the updated tests cover it.

The findings share one root cause, which is why several read similarly: this was developed against a newer TensorRT-LLM than the repo pins, and the branch appears not to have been run against the pinned build. Three of the comments below are consequences of that (recompute_active_requests, the FP8 refit hooks, the torch.compile hooks); one is an unrelated rebase artifact that breaks all TRT-LLM runs (engine_cfg); and three are test/doc drift. None of them touch the design.

Generated by Claude Code

Comment thread nemo_rl/models/generation/trtllm/trtllm_worker_async.py Outdated
Comment thread nemo_rl/models/generation/trtllm/trtllm_backend.py Outdated
Comment thread nemo_rl/models/generation/trtllm/trtllm_backend.py
Comment thread tests/unit/models/generation/trtllm/test_trtllm_fp8.py Outdated
Comment thread tests/unit/models/generation/trtllm/test_trtllm_backend.py Outdated
Comment thread tests/unit/models/generation/trtllm/test_trtllm_fp8.py
Comment thread nemo_rl/models/generation/trtllm/trtllm_worker_async.py Outdated
Comment thread docs/guides/refit.md Outdated
@hchings hchings removed the Documentation Improvements or additions to documentation label Sep 15, 2026
@github-actions github-actions Bot added the Documentation Improvements or additions to documentation label Sep 15, 2026
@hchings
hchings force-pushed the erinh/trtllm-mlperf-fp8 branch 2 times, most recently from 0fbfddf to 9d6eb42 Compare September 15, 2026 21:50
@hchings
hchings marked this pull request as ready for review September 16, 2026 20:36
@hchings
hchings requested review from a team as code owners September 16, 2026 20:36
hchings added a commit that referenced this pull request Sep 22, 2026
Retested on GB200 (SLURM job 3126032) after the fixes: all 44 trtllm
unit tests pass (was 41 passed / 3 failed).

- tests/.../test_trtllm_backend.py: force the manual finalize fallback
  in test_collective_refit_runs_at_async_engine_boundary via
  `monkeypatch.setattr(backend.WorkerExtension, "finalize_weight_update",
  None, raising=False)`. rc28's WorkerExtension now has
  finalize_weight_update (rc21 didn't), so trtllm_backend.py's
  _finalize_weight_update silently switched to that fast path and the
  test's hardcoded call-order assertion (written only for the fallback)
  started failing. Ported verbatim from the equivalent fix already
  landing in #4122 (which adds the real
  finalize_weight_update-aware FP8 refit logic later) — this commit
  only makes the test deterministic across TRT-LLM versions, it does
  not implement anything new.
- nemo_rl/models/generation/trtllm/trtllm_backend.py: import
  control_action_decorator from tensorrt_llm.executor.ray.utils (the
  rc28 location) with a fallback to the old tensorrt_llm._ray_utils
  path. The old path still works but emits "will be removed in a
  future release"; fixing now avoids a silent break on the next bump.
- uv.lock: regenerated against the bumped ref (was still rc21-era).
  Resolves cleanly (529 packages), confirms tensorrt/tensorrt-cu13/
  tensorrt-cu13-bindings/tensorrt-cu13-libs are no longer pulled in
  (matches upstream dropping its own tensorrt pin). Note: it resolves
  torch and nvidia-cudnn-frontend via this repo's override-dependencies
  (torch==2.11.0+cu130, nvidia-cudnn-frontend==1.25.0) rather than what
  TRT-LLM's own dependencies actually declare (torch>=2.13.0,
  nvidia-cudnn-frontend>=1.27.0) — uv's override mechanism forces the
  version project-wide without validating the sub-dependency is
  actually satisfied. This is the same conflict already flagged in
  3rdparty/TensorRT-LLM-workspace/pyproject.toml's "KNOWN UNRESOLVED
  CONFLICTS" comment; still unresolved, not addressed by this commit.
hchings added a commit that referenced this pull request Sep 24, 2026
Retested on GB200 (SLURM job 3126032) after the fixes: all 44 trtllm
unit tests pass (was 41 passed / 3 failed).

- tests/.../test_trtllm_backend.py: force the manual finalize fallback
  in test_collective_refit_runs_at_async_engine_boundary via
  `monkeypatch.setattr(backend.WorkerExtension, "finalize_weight_update",
  None, raising=False)`. rc28's WorkerExtension now has
  finalize_weight_update (rc21 didn't), so trtllm_backend.py's
  _finalize_weight_update silently switched to that fast path and the
  test's hardcoded call-order assertion (written only for the fallback)
  started failing. Ported verbatim from the equivalent fix already
  landing in #4122 (which adds the real
  finalize_weight_update-aware FP8 refit logic later) — this commit
  only makes the test deterministic across TRT-LLM versions, it does
  not implement anything new.
- nemo_rl/models/generation/trtllm/trtllm_backend.py: import
  control_action_decorator from tensorrt_llm.executor.ray.utils (the
  rc28 location) with a fallback to the old tensorrt_llm._ray_utils
  path. The old path still works but emits "will be removed in a
  future release"; fixing now avoids a silent break on the next bump.
- uv.lock: regenerated against the bumped ref (was still rc21-era).
  Resolves cleanly (529 packages), confirms tensorrt/tensorrt-cu13/
  tensorrt-cu13-bindings/tensorrt-cu13-libs are no longer pulled in
  (matches upstream dropping its own tensorrt pin). Note: it resolves
  torch and nvidia-cudnn-frontend via this repo's override-dependencies
  (torch==2.11.0+cu130, nvidia-cudnn-frontend==1.25.0) rather than what
  TRT-LLM's own dependencies actually declare (torch>=2.13.0,
  nvidia-cudnn-frontend>=1.27.0) — uv's override mechanism forces the
  version project-wide without validating the sub-dependency is
  actually satisfied. This is the same conflict already flagged in
  3rdparty/TensorRT-LLM-workspace/pyproject.toml's "KNOWN UNRESOLVED
  CONFLICTS" comment; still unresolved, not addressed by this commit.

Signed-off-by: Erin Ho <14718778+hchings@users.noreply.github.com>
shuyixiong pushed a commit that referenced this pull request Sep 29, 2026
Retested on GB200 (SLURM job 3126032) after the fixes: all 44 trtllm
unit tests pass (was 41 passed / 3 failed).

- tests/.../test_trtllm_backend.py: force the manual finalize fallback
  in test_collective_refit_runs_at_async_engine_boundary via
  `monkeypatch.setattr(backend.WorkerExtension, "finalize_weight_update",
  None, raising=False)`. rc28's WorkerExtension now has
  finalize_weight_update (rc21 didn't), so trtllm_backend.py's
  _finalize_weight_update silently switched to that fast path and the
  test's hardcoded call-order assertion (written only for the fallback)
  started failing. Ported verbatim from the equivalent fix already
  landing in #4122 (which adds the real
  finalize_weight_update-aware FP8 refit logic later) — this commit
  only makes the test deterministic across TRT-LLM versions, it does
  not implement anything new.
- nemo_rl/models/generation/trtllm/trtllm_backend.py: import
  control_action_decorator from tensorrt_llm.executor.ray.utils (the
  rc28 location) with a fallback to the old tensorrt_llm._ray_utils
  path. The old path still works but emits "will be removed in a
  future release"; fixing now avoids a silent break on the next bump.
- uv.lock: regenerated against the bumped ref (was still rc21-era).
  Resolves cleanly (529 packages), confirms tensorrt/tensorrt-cu13/
  tensorrt-cu13-bindings/tensorrt-cu13-libs are no longer pulled in
  (matches upstream dropping its own tensorrt pin). Note: it resolves
  torch and nvidia-cudnn-frontend via this repo's override-dependencies
  (torch==2.11.0+cu130, nvidia-cudnn-frontend==1.25.0) rather than what
  TRT-LLM's own dependencies actually declare (torch>=2.13.0,
  nvidia-cudnn-frontend>=1.27.0) — uv's override mechanism forces the
  version project-wide without validating the sub-dependency is
  actually satisfied. This is the same conflict already flagged in
  3rdparty/TensorRT-LLM-workspace/pyproject.toml's "KNOWN UNRESOLVED
  CONFLICTS" comment; still unresolved, not addressed by this commit.

Signed-off-by: Erin Ho <14718778+hchings@users.noreply.github.com>
shikicloud and others added 11 commits September 29, 2026 06:25
Signed-off-by: Erin Ho <14718778+hchings@users.noreply.github.com>
The collective (non-colocated) refit path called reset_prefix_cache(), which
only drops the reusable prefix blocks. Requests already in flight keep the KV
they computed under the old weights, so the tail of every in-flight
trajectory is generated against a mixed weight state.

recompute_active_requests() — added to PyExecutor by the tekit ref this
branch pins — instead sends the active requests back through context so
their KV is rebuilt with the weights that just landed.

Signed-off-by: Erin Ho <14718778+hchings@users.noreply.github.com>
init_collective calls ncclCommInitRank to build the refit communicator
across all train + inference ranks. Since tekit 97b62625 the executor loop
runs its per-iteration _broadcast_request_count as an NCCL broadcast on the
engine's TP group rather than over CPU/gloo, so the two now contend for the
same device and deadlock: the broadcast waits on peer ranks whose main
thread is inside ncclCommInitRank, which in turn waits for every rank to
arrive. PG5 stalls at work 85 after 84 clean iterations and ProcessGroupNCCL
aborts the engine 600 s later.

Wrap init_collective in control_action so the loop parks at a step boundary
first. This is the mechanism the refit path already uses
(update_weights_from_collective, update_weights_via_ipc_zmq); only
init_collective was missing it.

Verified on a 16-node 397B MXFP8 run with the NCCL object collectives left
enabled: step 7/30, six refits, the step-5 checkpoint save, zero watchdog
timeouts. Before the change the same configuration died in init_collective
every time.

Signed-off-by: Erin Ho <14718778+hchings@users.noreply.github.com>
Signed-off-by: Erin Ho <14718778+hchings@users.noreply.github.com>
…requests

engine_cfg was a leftover from extracting this branch out of a larger
rubin-pin branch; use trtllm_cfg like the rest of the file. Also fall back to
reset_prefix_cache when recompute_active_requests isn't available on the
installed TRT-LLM.

Signed-off-by: Erin Ho <14718778+hchings@users.noreply.github.com>
…equests

Point the fallback comment at NVIDIA/TensorRT-LLM#17937 (open, not yet
merged) instead of naming an internal branch.

Signed-off-by: Erin Ho <14718778+hchings@users.noreply.github.com>
…pports it

configure_fp8_llm_kwargs takes an optional llm_args_type used only to detect
whether the installed TRT-LLM's LlmArgs declares
use_cute_dsl_blockscaling_mm. Setting it unconditionally raised on TRT-LLM
builds without the field; skipping it entirely left it permanently off on
builds that do have it.

Signed-off-by: Erin Ho <14718778+hchings@users.noreply.github.com>
… mock

The lambda only took weights, so the fp8=True parametrization raised
TypeError on the is_mx kwarg production actually passes, silently converting
into the poisoning RuntimeError instead of exercising the conversion path.

Signed-off-by: Erin Ho <14718778+hchings@users.noreply.github.com>
…llback

The engine fixture was a bare MagicMock, which auto-creates any attribute on
access -- so every existing test silently exercised the "hook present"
branch even though none configured one explicitly. Make "absent" (every
released TRT-LLM today) the honest fixture default, and add a test for the
"present" branch (e.g. a TRT-LLM build carrying
NVIDIA/TensorRT-LLM#17937).

Signed-off-by: Erin Ho <14718778+hchings@users.noreply.github.com>
- trtllm_worker_async.py: reject is_mx=True with precision != "fp8" instead
  of silently starting a plain BF16 engine; use bool(trtllm_cfg.get("is_mx"))
  per config-conventions rather than a non-None default at the access site.
- test_trtllm_fp8.py: add MXFP8 coverage that was entirely missing --
  cast_tensor_to_mxfp8_blockwise round-trip, configure_fp8_moe_backend(is_mx=True)
  requiring CUTLASS, configure_fp8_llm_kwargs(is_mx=True)'s quant contract, and
  a random-data (not periodic-ramp, which can coincidentally survive a
  scramble) block-FP8 round-trip that actually exercises intra-block element
  order, unlike the existing constant-block tests. Verified all four against
  the real implementation, and verified the new block-FP8 test fails on a
  simulated dropped-permute regression while the old constant-block tests do
  not.
- refit.md: the "MXFP8/E8M0 scales are not used" sentence contradicted the
  is_mx=True path this same PR ships; scope it to the default path and
  document the MXFP8 variant, its YAML, and its CUTLASS/SM-constraint limits.

Signed-off-by: Erin Ho <14718778+hchings@users.noreply.github.com>
Both passed only by accident against the sandboxed/old-pin CI environment
and failed when actually run against an installed TRT-LLM carrying the full
RL refit lifecycle:

- test_collective_refit_runs_at_async_engine_boundary asserted a call order
  that only holds when WorkerExtension.finalize_weight_update is absent, but
  never forced that fallback path (unlike the sibling
  test_refit_finalization_falls_back_for_older_trtllm, which does). Force it
  explicitly.
- test_collective_refit_always_resets_prefix_cache used a bare MagicMock
  engine, which auto-creates recompute_active_requests and masks the
  reset_prefix_cache fallback its name asserts. Delete the attribute so
  absence is real.

Verified: all 72 trtllm-marked tests in tests/unit/models/generation/trtllm/
pass against a real tekit-backed TRT-LLM install (job 3074375).

Signed-off-by: Erin Ho <14718778+hchings@users.noreply.github.com>
@hchings
hchings force-pushed the erinh/trtllm-mlperf-fp8 branch from 43efd59 to 0807be9 Compare September 29, 2026 13:26

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

Documentation Improvements or additions to documentation

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants