Conversation
hchings
left a comment
There was a problem hiding this comment.
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.pyimports onlymath/re/torchand 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 thevllm/quantization/fp8.pyanalog, 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_5MoeHfWeightMapperexpects, the[ceil(N/128), ceil(K/128)]FP32 scale orientation anddequant = fp8 * scale_invmatchcreate_weights, the MXFP8 plain-uint8[..., K/32]layout matchesMXFP8CutlassFusedMoEMethod(which does its own packing and swizzle, so not pre-swizzling is correct), and all 13modules_to_not_convertpatterns 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.pyis a genuine improvement — the oldresults[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
0fbfddf to
9d6eb42
Compare
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.
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>
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>
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>
43efd59 to
0807be9
Compare
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
# Add a code snippet demonstrating how to use thisBefore your PR is "Ready for review"
Pre checks:
Additional Information