Skip to content

feat(rollout): segment-based kernel injection for hybrid-engine greedy decode - #8602

Draft
delock wants to merge 6 commits into
masterfrom
gma/qwen_ki_expr
Draft

delock wants to merge 6 commits into
masterfrom
gma/qwen_ki_expr

Conversation

@delock

@delock delock commented Sep 20, 2026 •

Copy link
Copy Markdown
Collaborator

Summary

Adds a segment-KI framework that accelerates greedy decode in the hybrid-engine rollout by injecting native CUDA kernels into comm-free segments of the model, plus a C++ decode loop that runs the full autoregressive generation with near-zero per-step Python overhead.

Segment detection: structural, not per-model

Detection is attribute-name based: any HF family whose gated MLP uses gate_proj/up_proj/down_proj (LLaMA, Qwen, Mistral, DeepSeek, Gemma...) or whose GatedDeltaNet block uses the in_proj_{qkv,z,b,a} set is picked up without per-model code. Projections carrying TP collectives are never fused (a defensive isinstance check against the Allreduce layer hierarchy).

Co-location by construction

Both GLU (dual-weight GEMV) and GDN (quad-weight GEMV) replacements read the original weight Parameters directly through autograd.Functions with exact backward, so training and generation share one forward path: zero weight copies, no inject/eject switching between rollout and training. Verified end-to-end by an integration test (below): generate → train one step → generate, output changes, no sync calls.

Components

File Role
csrc/module_inject/fused_glu.cu dual/quad-weight GEMV (MLP / GDN input projections), GDN gates, decode attention (warp-per-head, KV split-K, online softmax), fused residual+RMSNorm, triple QKV GEMV, graph-capturable decode step
csrc/module_inject/decode_loop.cu C++ generation loop taking a backend-agnostic std::function replay callable (CUDA graph, eager forward, or future SYCL); argmax and buffer updates run as a private kernel; EOS checked via amortized D2H sync every 16 steps
deepspeed/module_inject/segment_ki.py segment detection + forward replacement (apply_segment_ki)
deepspeed/module_inject/kernel_reference.py pure-PyTorch executable specifications for every kernel — test oracle, runtime fallback, porting reference
deepspeed/ops/module_inject/{fused_glu,decode_loop}.py op builders (JIT compile, lazy load)
deepspeed/utils/static_cache.py hybrid-slot static cache: real static KV layers for full attention, pass-through slot reusing HF's cudagraph-safe GDN state buffers by reference
deepspeed/runtime/rollout/hybrid_engine_rollout.py prefill → cache setup → graph capture → C++ decode loop, with layered fallback (full-step graph for b=1, graph+argmax for b>1, C++ loop, Python loop)

Performance

Qwen3.5-4B, RTX 4080 SUPER, greedy, b=1:

tok/s relative
HF eager 22.5 1.0×
DeepSpeed (this PR) 69.4 3.1×
vLLM 74.9 —

Tests

tests/unit/module_inject/test_segment_ki_kernels.py — 11 tests, all passing (executed on RTX 4080 SUPER):

  • Kernel vs reference consistency (GPU, 6 tests): every kernel compared against its kernel_reference.py oracle on identical inputs (dual_gemv, quad_gemv, gdn_gates, triple_gemv, fused_add_norm, decode_attn)
  • Injection invariants (CPU-runnable, 4 tests): injected forward matches un-injected output, no weight copies installed, composite autograd path, backward scatters gradients to the original Parameters
  • Rollout → train → rollout integration (GPU, Qwen3.5-0.8B-Base): generate with graph capture + segKI, one SGD step through the same injected model, generate again with changed output — no inject/eject/sync calls anywhere in between
  • Multi-GPU (AutoTP) validation (planned follow-up)

Notes

  • SYCL/XPU porting: kernel_reference.py serves as the executable spec; decode_loop's std::function interface is backend-agnostic (graph replay or eager forward); SYCL kernels would follow the same op-builder pattern.
  • Related upstream fix split into a standalone PR: fix: attention_unfused alpha missing norm_factor squaring #8666 (attention_unfused alpha squaring).

Add a segment-KI framework that accelerates greedy decode in the
hybrid-engine rollout by injecting native CUDA kernels into
comm-free segments of the model, plus a C++ decode loop that runs
the full autoregressive generation with near-zero per-step Python
overhead.

Segment detection is structural (attribute-name based): any HF
family whose gated MLP uses gate_proj/up_proj/down_proj or whose
GatedDeltaNet block uses the in_proj_{qkv,z,b,a} set is picked up
without per-model code. Projections carrying TP collectives are
never fused. The GLU replacement reads the original weight
Parameters directly through an autograd.Function with exact
backward, so train/generate share one forward path with no weight
copies and no inject/eject switching (rollout-training
co-location).

Components:
- fused_glu.cu: dual-weight GEMV+silu (MLP), GDN gates, decode
  attention (warp-per-head, KV split-K, online softmax), fused
  residual+RMSNorm, triple QKV GEMV, graph-capturable decode step
- decode_loop.cu: C++ generation loop taking a backend-agnostic
  std::function replay callable (CUDA graph, eager forward, or
  future SYCL); argmax and buffer updates run as a private kernel,
  EOS checked via amortized D2H sync every 16 steps
- kernel_reference.py: pure-PyTorch executable specifications for
  every kernel (test oracle, runtime fallback, porting reference)
- DeepSpeedStaticCache: hybrid-slot static cache with a
  pass-through slot for linear-attention (GDN) layers that reuses
  HF's cudagraph-safe state buffers by reference
- rollout integration: prefill -> cache setup -> graph capture ->
  C++ decode loop, layered fallback (full-step graph for b=1,
  graph+argmax for b>1, C++ loop, Python loop)

Verified on Qwen3.5-4B (RTX 4080 SUPER, greedy, b=1): correct text,
69.5 tok/s vs HF eager 22.5 (~3.1x), 92% of vLLM 74.9; co-location
checks (gradient flow, weight freshness, zero copies) all pass.

Signed-off-by: Guokai Ma <guokai.ma@intel.com>
@delock delock changed the title feat(rollout): HybridEngineRollout graph capture + segment-KI for Qwen3.5 hybrid — 67 tok/s (1.14× of vLLM) feat(rollout): segment-based kernel injection for hybrid-engine greedy decode Sep 26, 2026
Signed-off-by: Guokai Ma <guokai.ma@intel.com>
use_segki=False previously still injected the decode attention and
fused norm kernels in the graph-capture path (they were gated on op
availability only), so disabling segki did not restore the fully
native forward. Now the flag disables all kernel injection.

Signed-off-by: Guokai Ma <guokai.ma@intel.com>
The GDN segment previously materialized a concatenated weight copy
(_ki_gdn_fused_weight) at install time — stale after any optimizer
step, so the fusion was effectively inference-only. Replace it with
a quad_gemv kernel (b=1) that reads the four original weight
matrices directly through a GDNInputProj autograd.Function with
exact backward, mirroring the dual-weight GLU pattern:

- b=1 decode: quad_gemv kernel, warp-per-row across qkv|z|b|a
- b>1: on-the-fly concat of four GEMMs (kernel_reference oracle)
- gradients flow to the original in_proj Parameters

The fused weight buffer is gone; weight freshness after training
is now guaranteed by construction for both GLU and GDN segments.

Signed-off-by: Guokai Ma <guokai.ma@intel.com>
Signed-off-by: Guokai Ma <guokai.ma@intel.com>
Three groups covering the segment-KI contracts:

1. Kernel vs reference (GPU): every fused_glu kernel (dual_gemv,
   quad_gemv, gdn_gates, triple_gemv, fused_add_norm, decode_attn)
   compared against its kernel_reference.py oracle on identical inputs.
2. Injection invariants (CPU): apply_segment_ki forward equivalence,
   no weight copies installed, composite autograd path and backward
   gradient scattering to original Parameters.
3. Rollout-train-rollout integration (GPU, Qwen3.5-0.8B-Base): the
   co-location contract end-to-end — generate with graph capture +
   segKI, one SGD step through the same injected model, generate again
   with changed output, no inject/eject/sync anywhere in between.

Executed on RTX 4080 SUPER: 11 passed (6 kernel-consistency, 4 CPU
invariants, 1 integration).

Signed-off-by: Guokai Ma <guokai.ma@intel.com>

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.

1 participant