Conversation
delock
force-pushed
the
gma/qwen_ki_expr
branch
from
September 25, 2026 07:43
f1f9beb to
b67e55c
Compare
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
force-pushed
the
gma/qwen_ki_expr
branch
from
September 26, 2026 07:00
11231a7 to
d41c914
Compare
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
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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 thein_proj_{qkv,z,b,a}set is picked up without per-model code. Projections carrying TP collectives are never fused (a defensiveisinstancecheck 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
csrc/module_inject/fused_glu.cucsrc/module_inject/decode_loop.custd::functionreplay 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 stepsdeepspeed/module_inject/segment_ki.pyapply_segment_ki)deepspeed/module_inject/kernel_reference.pydeepspeed/ops/module_inject/{fused_glu,decode_loop}.pydeepspeed/utils/static_cache.pydeepspeed/runtime/rollout/hybrid_engine_rollout.pyPerformance
Qwen3.5-4B, RTX 4080 SUPER, greedy, b=1:
Tests
tests/unit/module_inject/test_segment_ki_kernels.py— 11 tests, all passing (executed on RTX 4080 SUPER):kernel_reference.pyoracle on identical inputs (dual_gemv, quad_gemv, gdn_gates, triple_gemv, fused_add_norm, decode_attn)Notes
kernel_reference.pyserves as the executable spec;decode_loop'sstd::functioninterface is backend-agnostic (graph replay or eager forward); SYCL kernels would follow the same op-builder pattern.