Skip to content

feat(pytorch): add glm5.3-flash support - #4968

Open
qescccczmr wants to merge 40 commits into
InternLM:mainfrom
qescccczmr:feat/glm5.3-flash-lmdeploy-reuse
Open

qescccczmr wants to merge 40 commits into
InternLM:mainfrom
qescccczmr:feat/glm5.3-flash-lmdeploy-reuse

Conversation

@qescccczmr

@qescccczmr qescccczmr commented Sep 14, 2026 •

Copy link
Copy Markdown
Collaborator

Summary

Add GLM-5.3 Flash text/image/video support to the PyTorch engine, including MTP and combined data/expert parallelism. Extend existing LMDeploy components for model loading, hybrid cache layout, multimodal preprocessing, and CUDA Graph execution.

  • Reuse DeepSeek MLA/MoE loaders, standard TP projections, FA3, HcPrePost, sparse-index Top-K, and blocked-FP8 MoE. GLM sparse attention uses the common FlashMLA backend.
  • Keep KDA as a thin backend adapter around public FLA operations and the shared recurrent/causal-convolution kernels. Preserve channelwise gates and per-token speculative state checkpoints.
  • Reuse common vision LayerNorm/RoPE operations with FP32 computation, and preserve clamped SwiGLU and routed-expert scaling.
  • Reuse the MTP predictor, rejection sampler, cache planner, and target-to-draft multimodal embedding handoff.

Optimizations

The original fused-serving optimization stage modified 12 existing production files. The latest addition, 19307597, modifies only the existing nn/hc_prepost.py (+9/−1 lines); it adds no files to the PR. Benchmark scripts and reports remain outside the repository diff. Follow-up head 932285b0 only wraps the existing MoE reduction docstring to satisfy CI; the measured runtime implementation is unchanged, and the PR still changes 66 files.

  • KPool verification: move the per-token Python loop into one request-parallel kernel, read the persistent state ring directly, and write every intermediate checkpoint. Reuse public compression and indexed cache writes. Rejection can still resume from any accepted prefix.
  • Sparse indices: skip score/selection when the entire short-prefill history fits the index budget; reuse compiled causal expansion. During decode, broadcast selected group IDs and expand them on each TP rank, reducing the index payload from 2,051 to 512 columns for this model.
  • KPool query rotation: share the existing FP32 Hadamard butterfly with compression, eliminating the query path's repeated add/sub/cat tensors and launches while retaining the output cast.
  • mHC/RMSNorm: fuse public HC reduction with public RMSNorm math, retaining the original intermediate BF16/FP16 rounding. Use wider post-expand tiles for four-stream prefills.
  • mHC input conversion: compile only the FP32 input conversion for contiguous inputs with at least 8,192 token rows. Keep the existing FP32 GEMM, RMS statistics and rounding order. Short requests retain the original conversion path; the threshold is based on the tested H200 2K/8K workloads, not a universal crossover.
  • KDA: make the common convolution update accept input strides, removing a layout copy. Parallelize independent value tiles in the channelwise recurrent kernel while preserving sequential time updates and all checkpoints.
  • TP reduction: retain FP32 reduction and reconnect explicit-dtype reductions to the shared communicator. Optional FlashInfer support handles FP32 and caches workspaces by dtype/hidden size using the correct process group. NCCL remains the default; oversized inputs retain the existing fallback.

FlashKDA, masked-MHA prefill, and FlashInfer NoPE MLA are not enabled by this patch.

Measured optimization gains

H200, the same GLM-5.3 Flash FP8 checkpoint, TP4/MTP5, input 2K/8K, output 5 tokens, BS1 unless specified. The original staged optimization results are retained below, followed by the 2026-09-28/29 additions. Timings are measured without active profiling; Torch Profiler runs are separate. Speedup = before / after; reduction = (before − after) / before. These are independent staged A/B experiments, not additive or multiplicative cumulative gains.

End-to-end staged measurements

Change / workload Before After Speedup Latency reduction
KPool verification fusion, 2K total request, 10 samples 315.26 ms 254.91 ms 1.24× 19.15%
KPool verification fusion, 2K TPOT 17.60 ms 9.09 ms 1.94× 48.34%
Index/mHC/KDA/shared-communicator stage, 2K total, 30 samples 245.21 ms 234.28 ms 1.05× 4.46%
Query-rotation/post-expand prefill stage, 8K TTFT, 60 samples per variant 477.13 ms 441.15 ms 1.08× 7.54%
New: sparse MLA remap + padding, 8K TTFT, 30 paired samples 440.30 ms 438.90 ms 1.0032× 0.32%
New: sparse MLA remap + padding, 8K total request, 30 paired samples 470.82 ms 469.91 ms 1.0019× 0.19%
New (2026-09-29): mHC input conversion, 8K TTFT, 30 paired samples 446.23 ms 433.35 ms 1.030× 2.89%
New (2026-09-29): mHC input conversion, 8K total request, 30 paired samples 475.91 ms 462.68 ms 1.029× 2.78%

Historical reduction percentages are retained as originally reported (displayed times are rounded). The query-rotation/post-expand experiment did not improve 2K TTFT: 218.23 → 220.52 ms; its 8K total latency was 509.16 → 472.46 ms. Short-request TPOT is not the time of a single verification kernel.

The new sparse MLA test alternated old/new paths within one loaded service, with 3 warmup pairs and 30 measured pairs. Its 8K paired TTFT median delta is −1.24 ms (bootstrap 95% interval −2.17 to −0.41 ms); the total-latency interval includes zero, so the 0.19% observed total reduction is not an established end-to-end gain. 2K prefill uses dense attention and is a negative control. Five-token paired output texts, and token IDs in separate profile requests, match.

The mHC conversion A/B uses the exact source in 19307597, based on ba4787d5, on GPUs4–7: 3 warmup pairs +30 measured pairs per length, alternating old/new order within one loaded service. Both 8K TTFT and total latency improve in 30/30 pairs. Paired median deltas (new − old) are −12.49 ms TTFT (bootstrap 95% interval −14.90 to −11.38 ms) and −12.90 ms total (−14.36 to −11.65 ms). 2K and TPOT have no established gain: 2K keeps the original conversion path, and its observed total change 267.23 →264.33 ms has an interval crossing zero. The same-session unmodified control differs from the older published timings; only these paired controls are used to attribute the new gain.

Full-model 2K/8K checks (two repeated pairs each) match all five generated token IDs and 96 captured full-vocabulary logits tensors, 24,161,280 elements, bitwise; max absolute error 0. All timed output texts match too. The final dispatch passes 28 numerical/CUDA Graph checks; the existing five HC kernel tests also pass. This is equality against the preceding LMDeploy implementation, not cross-engine token parity or a new task-accuracy evaluation.

Separate four-rank Torch Profiler traces attribute the 8K saving to input conversion: rank0's 90 casts total 35.17 →22.16 ms, with FP32 GEMM and RMS statistics unchanged. Mean mHC GPU kernel sum across ranks is 96.22 →83.22 ms; this is a profiling diagnostic, not an additional serving gain. Source snapshots, raw paired requests, logits and traces are retained in glm5.3flash/reports/mhc_compile_20260929/final.

Operator-level gains

These operator measurements exclude the rest of the model and must not be interpreted as serving speedups. Historical measurements and new measurements have their own controls; for example, the older KPool 8.38 µs and newer 5.76 µs controls are not one continuous A/B experiment.

Optimization / workload Before After Speedup Latency reduction
KPool sequence update/compression/cache write, B1/S6 (805 → 4 kernels) 1130.10 µs 8.38 µs 134.86× 99.26%
KDA adapter, B1/S6 37.86 µs 14.66 µs 2.58× 61.28%
Query rotation, 8K rows 2055.00 µs 45.00 µs 45.67× 97.81%
HC post-expand, 8K rows 323.00 µs 179.00 µs 1.80× 44.58%
New: prefill group broadcast + local metadata/expansion, TP4/4096 rows 0.3953 ms 0.3428 ms 1.15× 13.29%
New: prefill group broadcast + local metadata/expansion, TP4/8192 rows 0.5865 ms 0.4414 ms 1.33× 24.75%
New: valid-aware KPool update/compress/cache write, B1/S1 (no closed pool) 4.088 µs 3.363 µs 1.22× 17.74%
New: valid-aware KPool update/compress/cache write, B1/S6 (no gain) 5.761 µs 5.819 µs 0.99× -1.00%
New: valid-aware KPool update/compress/cache write, B64/S6 6.835 µs 6.603 µs 1.04× 3.39%
New: sparse MLA index remap + padding, 8K/BS1, 7 → 1 kernels 119.66 µs 49.82 µs 2.40× 58.36%
New: sparse MLA index remap + padding, 8K/ragged BS2, 7 → 5 kernels 120.07 µs 66.79 µs 1.80× 44.37%

The valid-aware path still updates/checkpoints every valid token; it only skips closed-buffer work for rows that do not close a pool. The B1/S6 result shows no demonstrated latency gain. Sparse MLA keeps the selected keys and negative-index semantics unchanged; four-rank full-model traces confirm fused remap/padding execution.

Capacity improvement: bounded KPool scoring

Workload / measured resource Before After Memory reduction
8K query rows / 128K history, largest FP32 score payload 1 GiB ≤512 MiB ≥50%
8K query rows / 256K history, largest FP32 score payload 2 GiB ≤512 MiB ≥75%

This is a capacity/stability improvement, not a claimed latency speedup. KPool reuses the configurable DSA score budget, flattens KV once, immediately selects deterministic Top-K per query chunk, and explicitly reserves score workspace without changing mla_index_topk. The bound accounts for DeepGEMM row/column alignment and rejects budgets below the minimum allocation; it does not bound all operator temporaries or the allocator pool. Full-model 128K/256K requests completed; selected indices match exactly.

The 2026-09-28 additions are 200a5591 (bounded scoring/reservation), fc843a30 (valid-aware compression and prefill group broadcast), and ba4787d5 (sparse MLA preparation). Local validation includes 41 existing executor/NSA tests, 12 existing MLA tests, 32 strict-compiled GPU mapper checks, and targeted state/cache, chunking, ragged-input, graph-replay and TP broadcast checks. No new full task-accuracy run is claimed for these additions.

FlashKDA is not enabled: the vLLM fork's synthetic 8K operator is faster (0.811 → 0.455 ms), but fails the existing numerical gate; both tested FlashKDA variants internally round recurrent state to BF16. Removing explicit KDA Q/K/V contiguous calls also provides no speedup (8K 757.30 → 757.46 µs), because FLA performs the copies internally. The earlier synthetic mHC projection profile (~779.90 µs at 8K, including ~381.27 µs input casting) motivated the input-conversion change measured above. That change preserves the original GEMM/RMS computation and does not reduce the FP32 temporary-buffer size.

LMDeploy vs vLLM

Latest LMDeploy measurements: 19307597, with all accepted optimizations above present, measured on 2026-09-29 using the new-path samples from the paired experiment. LMDeploy uses NCCL and FLA Triton; FlashKDA and optional FlashInfer all-reduce are not enabled. These are actual final-implementation timings, not estimates obtained by adding individual optimization gains.

vLLM was following the official GLM-5.3-Flash recipe, on unchanged source 606d124b (v0.29.1rc1.dev474, FlashInfer0.6.18.post1). The recipe's Hopper requirement is BF16 KV. Native vllm serve uses the default single-node multiprocess executor, collective selection and DeepGEMM scale policy, plus the glm47 tool/reasoning parsers. The earlier Ray/disabled-FlashInfer/disabled-symmetric-memory/forced-FP32-scale overrides are removed. Trace confirms FlashKDA, FlashInfer MNNVL, symmetric-memory2K reductions and NCCL8K large-message reductions.

The request workload matches: H200 TP4/EP1/DP1, MTP5, BF16 KV, FP32 recurrent-state buffers, prefix caching off, BS1, output 5 tokens, max batch16, prefill8192, session270336; identical checkpoint and prompt-token manifest; 3 warmups +30 timed requests, medians without active profiling. The latest LMDeploy and vLLM runs both used GPUs4–7 on the same host, at different times; GPUs0–3 were occupied by other work. vLLM was not rerun for the mHC update, so this is not a contemporaneous cross-engine A/B. The preceding LMDeploy table used the 2026-09-28 run on GPUs0–3 (239.36/469.91 ms total at 2K/8K); the fresh unmodified controls are 267.23/475.91 ms. Use the same-session paired table above to assess the code change, rather than treating cross-run host/card variation as a regression or speedup.

Input Metric LMDeploy vLLM
2,048 TTFT 230.69 ms 132.44 ms
2,048 TPOT 8.49 ms 7.19 ms
2,048 Total request 264.33 ms 160.98 ms
8,192 TTFT 433.35 ms 310.40 ms
8,192 TPOT 7.43 ms 6.67 ms
8,192 Total request 462.68 ms 336.95 ms

The new vLLM total medians are 160.98/336.95 ms versus the earlier163.71/339.73 ms (2K/8K), an observed reduction of1.67%/0.82%. These small cross-run differences are not an isolated communication speedup: executor, scale defaults and physical cards also changed. Raw requests, native counters, launch configuration and four-rank Torch Profiler traces are retained in glm5.3flash/reports/vllm_recipe_20260929.

TTFT includes scheduling and first-output preparation. TPOT is (total − TTFT) / 4; SSE events can contain multiple MTP tokens, and medians are computed separately. Native cache budgets differ: LMDeploy free-memory fraction 0.8 versus vLLM total-memory utilization 0.9. Greedy outputs are not uniformly token-identical across the engines; these five-token latency measurements do not establish acceptance-rate or long-generation throughput parity.

Separate Torch Profiler traces were collected for both engines at 2K/8K, BS1/MTP5/output5, on all four ranks. mHC, KDA, MoE, sparse attention and collective waiting remain in the attribution; kernel-duration sums can overlap and are not request latency. Independent FP32 collective graph probes confirm FlashInfer handles TP4 rows 6/128 and TP8 rows 6/30, while rows 132/36 respectively fall back to NCCL. Those are actual flattened-row probes, not client concurrency or Python replay counters.

Accuracy

MMMU-Pro vision/test

Backend Final score
LMDeploy, NCCL 1333/1730 (77.0520%)
vLLM 1329/1730 (76.8208%)

LMDeploy is +0.1720 percentage points from the 76.88% reference. Both engines finish with 0 request errors and 0 remaining length stops.

Protocol: all 1,730 vision/test items from dataset revision 563f3e84bb3b90893083a1f039cfa13077f2302b; original image bytes; identical image-first prompts; pinned NeMo prompt and answer-extraction helpers (778f31a). Temperature 1.0, top_p 0.95, top_k disabled, maximum 327,680 new tokens, seed 42, reasoning effort max. TP4/EP1/DP1/MTP5, client concurrency 8, maximum batch 16, session length 344,064, prefill chunk 8,192; LMDeploy uses NCCL. Both backends use FP8 checkpoint 3f1971b7b5f7a528c9c4ef6212c8785298a8c24a on disjoint groups of four H200 GPUs. NCCL_NVLS_ENABLE=0 is an explicit benchmark setting for both.

The 76.88% reference is a BF16 baseline reported in the NVIDIA GLM-5.3 Flash model card. This FP8 evaluation is not an exact reproduction of that baseline or its unpublished harness details.

Final scores include one fresh-service retry of the single LMDeploy length-stopped item, using the same sampling parameters and 327,680-token output cap. All length-selected replacements are retained regardless of correctness; incorrect answers alone do not trigger retries.

GPQA Diamond and SciCode

Evaluated revision: the scores below were measured on frozen LMDeploy commit 38779f12. The PR has since advanced to 18a8c917, including MoE-reduction changes. This table does not include a full-model accuracy rerun of those later commits.

Same GLM-5.3 Flash FP8 checkpoint as above; LMDeploy 38779f12, vLLM 606d124b. Both run TP4/EP1/DP1/MTP5 on separate four-H200 groups. LMDeploy uses NCCL; FlashInfer/symmetric-memory all-reduce is disabled. Temperature 1.0, top_p 0.95, top_k disabled, repetition penalty 1, seed 42, reasoning effort max, one sample per item. Client concurrency 8, maximum batch 16, BF16 KV, FP32 recurrent state, prefix caching off, prefill chunk 8,192, NCCL_NVLS_ENABLE=0.

Benchmark LMDeploy NCCL vLLM LMDeploy − vLLM
GPQA Diamond 184/198 (92.9293%) 182/198 (91.9192%) +1.0101 pp
SciCode with background: subproblem 151/288 (52.4306%) 143/288 (49.6528%) +2.7778 pp

Generation starts with a 327,680-token output limit and context 393,216. Length-stopped items are retried once at 655,360 output tokens and context 720,896; SciCode also regenerates their dependent downstream steps. The table includes every selected replacement, regardless of whether it improves the score. Incorrect answers alone do not trigger retries.

SciCode evaluation fix: restore the published Maxwell, Block, and EnlargedBlock class definitions omitted by the pinned helper, then regenerate all 14 dependent subproblems for both engines. The final scores above include these corrected runs. No numerical assertion or tolerance was changed.

GPQA uses all 198 Diamond questions, frozen zero-shot prompts/option permutations validated against the source answers, and the pinned NeMo MCQ extractor. Cross-engine prompt-token count differences: 0 / 198.

SciCode uses all 65 test problems with background, 288 scored subproblems, and the three official prefilled steps excluded from the scored denominator. Main score requires every scored subproblem in that problem to pass. Each next prompt uses that engine’s previous generated code. Dataset revision; NeMo prompt/generation helpers. Original numerical assertions run in an isolated Python 3.11 environment with NumPy 1.26.4, SciPy 1.10.1, SymPy 1.12, h5py 3.11.0; 1,800-second per-step timeout, one BLAS thread, Python/NumPy execution seed 42. Development-reference validation passed 48/50: 70.8 and 78.3 also fail using the original SciCode helpers (large-phase numerical mismatch and wall-clock-dependent trajectory shape respectively). These are development cases; no test case was excluded and no scoring tolerance was relaxed.

These are single-seed stochastic measurements with dynamic batching, not token-parity proof or an exact reproduction of an unpublished official harness. Raw responses, execution logs, provenance, and original/retry scores are retained under glm5.3flash/reports/gpqa_scicode_20260928; evaluation/benchmark scripts are not committed.

GSM8K

Backend Correct / total Accuracy
LMDeploy, NCCL 1,290 / 1,319 97.8014%
LMDeploy, optional FlashInfer 1,286 / 1,319 97.4981%
vLLM 1,288 / 1,319 97.6497%

TP4/MTP5, client concurrency 32, greedy, maximum 2,048 output tokens, identical prompt token IDs and parser. Errors and length stops remain in the denominator; explicit final answers are scored even at a length stop. Score proximity does not establish token parity or absence of regression.

Validation and DP/EP support

Focused GPU checks cover state/cache outputs and per-token checkpoints, invalid/padded requests, ring wrap/rejection, causal indices, strides, BF16/FP16/FP32 HC paths, and CUDA Graph replay with changed inputs. KPool state/cache checks and query/compression checks retain bitwise matching against the reference paths. Shared communicator dispatch tests pass on the merged candidate (9 tests). Validation scripts and reports remain outside the PR.

GLM uses the shared DeepseekV2MoE / blocked-FP8 / DeepEP path. It propagates clamped activation and 2.5 routed-expert scaling through synchronous/asynchronous prefill and decode, leaves shared experts unscaled, avoids a second TP sum of already-combined routed outputs, and reduces the shared projection separately where DP1/EP requires it. DeepEP partial-sum transport remains BF16; non-default routed scaling uses FP32 local prefill reduction. Sparse FlashMLA padding and the current DeepGEMM masked-GEMM symbol are supported.

Existing DP/EP smoke checks at 3f21b23f, before the latest optimization patch: H200, Ray, BF16 KV, 5 output tokens, maximum batch 4, session 9,216, prefill chunk 2,048, prefix caching disabled.

Configuration Completed requests
DP4 / EP4 / Attention TP1 / eager / MTP0 18/18
DP4 / EP4 / Attention TP1 / CUDA Graph / MTP5 18/18
DP2 / EP4 / Attention TP2 / CUDA Graph / MTP5 7/7
Final DP/EP commit: DP4 / EP4 / Attention TP1 / CUDA Graph / MTP5 18/18

These smoke checks are not full accuracy evaluations of DP/EP. The current MMMU-Pro evaluation uses EP1/DP1.

Reuse MLA/MoE loaders, FA3, mHC, sparse Top-K and compact DeepGEMM. Add KDA state adaptation and multimodal GLM configuration/processing. Preserve existing defaults and the positional MoE prefix argument.
@qescccczmr qescccczmr changed the title Feat/glm5.3 flash lmdeploy reuse Feat/glm5.3 flash Sep 14, 2026
@qescccczmr qescccczmr changed the title Feat/glm5.3 flash feat(pytorch): add glm5.3-flash support Sep 14, 2026
Comment thread lmdeploy/pytorch/models/glm5_next.py Outdated
device=device,
is_tp=True,
quant_config=None,
dp_disable_tp=True,

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Keep KV-B sharding consistent with the attention TP group
When dp > 1 and attn_tp > 1, dp_disable_tp=True makes kv_b_proj retain all attention heads, while DeepseekV2BMM still shards kc/vc by attn_tp. The inherited process_weights_after_loading() then copies the full KV-B-derived weights directly into these sharded tensors.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Fixed in 3fcabb7. Removed dp_disable_tp=True from GLM's KV-B projection so it uses the existing attention-TP sharding and weight loader, consistent with DeepseekV2BMM KC/VC absorption.

The focused CPU test exercises the real builder/loader/absorption in nine DP/TP/rank combinations and checks the exact local weight slices. This is loader-contract validation, not an EP+DP distributed-service validation; the unrelated EP+DP experiments are not included.

@RunningLeon

Copy link
Copy Markdown
Collaborator

@qescccczmr Hi, is this PR ready to review?

@RunningLeon
RunningLeon self-requested a review September 21, 2026 06:59
from .step_metadata import register_step_metadata_impl


def _select_state(state: torch.Tensor, metadata: Any) -> torch.Tensor:

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

can we use existing _state_select , _state_scatter in here?

def _state_select(state, state_indices, spec_offsets):

def _state_scatter(state, state_indices, spec_offsets, src):

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Addressed in 3fcabb7. KDA now reuses the existing GDN _state_select / _state_scatter helpers for both AR and MTP. The ordinary AR bank is exposed as a one-slot ring, rather than introducing another state kernel.

The shared helpers mask invalid state/slot IDs: reads return zero and writes are skipped, preventing padded rows from overwriting a live request. Actual CUDA tests cover FP32/BF16, irregular shapes, negative/out-of-range IDs, initialization, dummy/live-row collisions, and replay after updating graph inputs.

Comment thread lmdeploy/pytorch/backends/cuda/kda.py Outdated
head_dim: int,
lower_bound: float,
) -> torch.Tensor:
if (metadata.spec_state_offsets is not None

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

mtp may need support as well

@qescccczmr qescccczmr Sep 22, 2026 •

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Updated in 3fcabb7. KDA restores state using accepted history and saves every verified token's convolution/recurrent checkpoint, including partial rejection, reordering and dummy rows. MTP now also honors index_share_for_mtp_iteration through the existing DSATopKIndicesBuffer; cache updates still run on every draft step. Growing the shared buffer changes the graph key so old allocations are not replayed.

We also found and fixed a sampling-distribution bug: rejection/recovery did not apply AR's top-k/top-p/min-p filters. The fix reuses FusedLogitsProcessor filtering, preserves greedy fast-path behavior, and leaves source logits intact. Focused CUDA tests verify that filtered-out tokens cannot be accepted or recovered.

Mode LMDeploy vLLM
Greedy, fixed 256 tokens 62.9341% 62.7595%
Sampling, fixed 256 tokens 54.7305% 53.5599%
Greedy, natural EOS (max 256) 62.4700% 61.5550%
Sampling, natural EOS (max 256) 56.9082% 57.5145%

H200, TP4/EP1, CUDA Graph on, prefix cache off, BF16 KV, BS1, 1024-token prefill chunks. Inputs: 20 GSM8K prompts plus four 8436-token long prompts, each run under all four policies; sampling uses temperature=1, top_p=0.95, seed=42 and disabled top-k. Stop IDs/minimum lengths are matched.

Comment thread lmdeploy/pytorch/backends/cuda/kda.py Outdated
inputs = [x.unflatten(1, (batch_size, steps))
for x in (mixed_qkv, raw_gate, raw_beta)]
outputs = []
for step in range(steps):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

this may lead to poor performance. If fla does not provide this verification kernel, we may need to change the tilelang kernel of gated_delta_rule.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Agreed; addressed in 3fcabb7 without a new kernel file. The existing GDN TileLang recurrent kernel's transposed-state path now supports KDA channelwise decay. Convolution windows are batched through FLA, then all verification tokens use one recurrent-kernel call. The kernel is parallel across batch/head/state tiles; causally dependent timesteps remain sequential inside that call, and each state is checkpointed for partial acceptance. The default GDN scalar-gate path is preserved.

Fresh CUDA regressions cover scalar/channelwise gates, AR versus multi-token state equality, dummy rows, graph replay and the original GDN tests. An earlier module-only CUDA Graph benchmark (3 verification tokens, 8 heads, head dim 128) measured 0.190/0.247/0.349 ms for the serial reference versus 0.057/0.059/0.065 ms for the shared path at BS1/4/16. This is KDA-local evidence, not a whole-model or vLLM performance-parity claim.

Reuse GDN state helpers and channelwise TileLang verification; honor MTP index sharing and invalidate graphs after buffer growth. Share AR sampling filters with rejection/recovery and stabilize GLM TP/mHC and sparse-index arithmetic. Keep EP/DP, prefix-cache and performance experiments outside this patch.
@qescccczmr

qescccczmr commented Sep 22, 2026 •

Copy link
Copy Markdown
Collaborator Author

@RunningLeon The MTP correctness/reuse changes are ready for another review in 3fcabb7. I am keeping this PR in Draft while whole-model accuracy/performance validation remains incomplete.

This update modifies 11 existing production files, with no new production module/kernel. It addresses the four review comments: attention-TP KV-B sharding, shared state helpers, accepted-history MTP state recovery, and reuse of GDN's TileLang verification kernel. It also fixes rejection sampling's missing AR filters, MTP index sharing and graph-buffer growth, and includes the GLM opt-in FP32 TP / mHC / stable sparse-index changes used for AR/MTP numerical stability. Existing defaults for other model callers are preserved. EP+DP, prefix-cache, FlashInfer communication and compact-MoE scheduling experiments are excluded.

config.state_cache_specs = [
StateCacheSpec(
GLM5_KDA_CONV_STATE,
(num_linear_layers, *ring_shape, conv_dim, conv_kernel_size),

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

this layout could the be the same as Qwen3.5 and reuse the causal conv kernel when mtp is on.

conv_kernel_size = text_config.linear_conv_kernel_dim + num_spec_tokens
conv_state_shape = (num_delta_layers, conv_dim, conv_kernel_size)

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Updated in 41563c7.

The convolution cache now uses (num_linear_layers, conv_dim, conv_kernel_size + num_spec_tokens), matching Qwen3.5. MTP decode calls the existing causal_conv1d_update with accepted-history cache_seqlens and request state indices. This removes the per-draft convolution-window checkpoints; recurrent/KPool checkpoints remain unchanged. For kernel width 4 and MTP5, convolution-state storage drops from 24 to 9 elements per channel.

The shared kernel now accepts the actual weight/bias dtype, preserving GLM's FP32 convolution weights with BF16 activations. Prefill retains FLA's existing short-chunk handling and maps its chronological window into the compact token ring.

Validation on the isolated patch: 59 CPU checks passed; the CUDA-enabled regression run passed 75 checks, including MTP2/MTP5 partial-acceptance recovery (bit-exact against sequential AR in the tested KDA cases), request reordering/padding, and CUDA Graph replay. The 24 upstream optional Dao-reference cases were skipped because that dependency is unavailable; six additional independent FP16/BF16/FP32 non-circular convolution checks passed. Ruff and git diff --check also passed. These are local operator/contract checks, not a fresh whole-model acceptance or performance benchmark.

prefix: str = '',
*,
fp32_acc: bool = False,
output_scale: float = 1.0,

@RunningLeon RunningLeon Sep 22, 2026 •

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

[P1] Honor or explicitly reject the new MoE options across backend branches
build_fused_moe() now exposes fp32_acc and output_scale, but only the blocked-FP8 branch forwards them. The BF16 and W8A8 branches silently ignore these arguments, so the same public API has different numerical semantics depending on the selected backend.
This causes a concrete correctness issue for GLM’s BF16 path: the router returns unscaled weights and relies on output_scale=2.5 in the expert reduction. Dropping that argument produces an incorrectly scaled routed-expert contribution.
Please propagate these options through supported implementations, preserving the FP32 weighted reduction → output scaling → output cast order. For implementations that do not support this contract yet, explicitly reject non-default values rather than silently ignoring them. Keep False / 1.0 as defaults to preserve existing behavior, and add regression tests for parameter propagation, numerical correctness, and unsupported-backend rejection.
The scale must apply only to the routed-expert contribution, not to the shared-expert output.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Fixed in 96423321.

  • The BF16 fused-MoE path now forwards fp32_acc and output_scale through its build spec and applies them in the CUDA reduction as FP32 weighted reduction → scale → output cast.
  • The scale is applied inside the routed-expert reduction, before the caller combines the result with the shared expert output.
  • W8A8, static-FP8, W4A16, and DeepEP branches now explicitly reject non-default values instead of silently ignoring them.
  • Added propagation, unsupported-backend, and CUDA numerical regression tests in tests/pytorch/nn/test_moe_options.py.

Validation: pytest -q tests/pytorch/nn/test_moe_options.py (4 passed, 1 CUDA test skipped because CUDA is unavailable in this environment); Ruff and compile checks pass.

@qescccczmr
qescccczmr marked this pull request as ready for review September 22, 2026 11:19
Comment thread lmdeploy/pytorch/nn/rotary_embedding.py Outdated


@torch.compile(dynamic=True)
def apply_rotary_pos_emb_fp32(query: Tensor,

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

does glm5.3 need to do this on fp32?

Comment thread lmdeploy/pytorch/nn/norm.py Outdated
return out.to(result_dtype)


class FP32LayerNorm(nn.Module):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Where do we use this module?

Comment thread lmdeploy/pytorch/backends/moe.py Outdated
layer_idx: int
output_dtype: torch.dtype
num_max_dispatch_tokens_per_rank: int
fp32_acc: bool = False

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

consider dlinfer backend for these two arguments

Comment thread lmdeploy/pytorch/config.py Outdated

# Model-specific defaults that must be present before the distributed
# process group is initialized. Explicit process environment values win.
process_group_env_defaults: dict[str, str] = field(default_factory=dict)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

why add this config?

Reject unsupported DLINFER MoE reduction options instead of ignoring them. Leave NCCL NVLS policy to the runtime and remove model-specific environment plumbing. Document GLM FP32 operator usage and cover dtype, reduction, and environment contracts.

Validation: 110 focused tests passed on CPU/H200; Ruff passed for lmdeploy and tests/pytorch. DLINFER rejection regressions failed before the fix and pass afterward.
Apply the configured string and docstring formatters to GLM PR files. Verified executable AST is unchanged; both previously failing hooks and Ruff pass locally.
qescccczmr and others added 28 commits September 23, 2026 11:27
Replace the compiled GLM helper with an opt-in FP32 compute contract through ApplyRotaryEmb and its backend build spec. CUDA retains FP32 arithmetic inside the fused kernel; defaults remain unchanged and unsupported Dlinfer requests fail explicitly.

Validated 73 focused tests and actual vLLM rotary kernel parity with identical inputs and tables. FP32-table outputs can differ from the old compiled helper at rounding boundaries.
Reuse the existing strided TileLang recurrent inputs instead of copying Q/K/V in AR and speculative decode. Preserve contiguous prefill inputs and explicitly allocate contiguous recurrent output for channel-major views.

Validated 25 new stride/state/graph tests, 25 shared GDR tests, and 12 complete KDA adapter comparisons. TP4 MTP-off text regression matches the previous head exactly at 124 full-vocabulary logits positions and all 29 requests.
Propagate GLM activation and routed scaling through DeepEP prefill and decode. Use FP32 local expert reduction for non-default output scaling while preserving the default BF16 path and transport. Avoid reducing combined experts twice and keep shared experts unscaled.

Pad sparse FlashMLA indices with invalid entries for its tile alignment and accept the current DeepGEMM masked-GEMM symbol.
Use an FP32 accumulator in ep_gather independently of the output dtype,
matching vLLM's DeepGEMM unpermute-and-reduce. Store the result directly
as BF16 for DeepEP combine, removing the FP32 output buffer and separate
cast from GLM's normal path.

Remove the private fp32_acc plumbing from the blocked-FP8 DeepEP builder
and normal execution paths. Preserve activation callbacks, FP32 scaled
routing weights, and the existing low-latency combine implementation.
The shared gather now uses FP32 for other callers, including BF16 experts.

Validation: 16 H200 kernel cases matched vLLM 606d124b and the previous
FP32-output-then-cast path exactly; covered BF16/FP16, top-k 1/8, scales
1/2.5, missing experts, row strides, and more than 1024 tokens. CUDA graph
replay and cancellation regression passed. Ruff 0.15.4, Python compile,
and git diff --check passed. Full-model/multi-rank inference not rerun.
Restore fp32_acc=False plumbing through the DeepEP Normal builder and
both execution entry points. Pass the flag into ep_gather as a Triton
constexpr so default callers retain output-dtype accumulation while
GLM's scaled routed-expert path explicitly uses FP32.

Keep the BF16 gather output and cast inside the kernel, avoiding the
previous FP32 temporary and separate output cast. Preserve activation
callbacks, routing-weight scaling and the low-latency combine path.

Validation: 16 H200 cases matched legacy accumulation with the default
and explicit False, and vLLM/previous GLM FP32 results with True. CUDA
graph replay, cancellation, builder/sync/async parameter propagation,
Ruff 0.15.4, Python compile and git diff --check passed. Full-model and
multi-rank inference were not rerun.
…opt-in

Restore moe_reduce's positional/keyword fp32_acc=False argument and the
legacy weighted-product precision for default callers. Keep output_scale
keyword-only and apply it after expert reduction before the output cast.

GLM explicitly sets fused_moe_fp32_acc=True. Propagate it through the
shared model, builder, typed specs, BF16/blocked-FP8 backends and kernels.
DeepEP Normal uses this flag independently of routing scale. Reject
unsupported non-default requests instead of silently dropping them.

Restore compressed_tensors_w4a16.py exactly to main: its existing explicit
FP32 calls are compatible again. The PR now changes 65 files, with no new
production or test files introduced by this update.

Validation: 22 focused H200/CPU checks passed, including exact legacy and
previous GLM reduction parity, positional/keyword API compatibility,
FP32 scaling, CUDA graph replay, parameter propagation, unsupported
backend rejection and actual BF16/blocked-FP8 expert pipelines. Ruff and
git diff --check passed. Full-model and distributed inference not rerun.
Compile the FP32 input conversion for contiguous HC inputs with at least
8192 token rows. Preserve the existing FP32 GEMM and RMS statistics,
rounding boundaries, and short-request execution path.

H200 TP4/BS1/MTP5, output 5, 30 paired requests: 8K TTFT falls from
446.23 to 433.35 ms and total latency from 475.91 to 462.68 ms (2.78%).
Full-model generated tokens and captured logits remain bitwise equal.
28 numeric/CUDA Graph checks pass; 2K shows no established speedup.

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

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants