Repository navigation
Conversation
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>
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>
…rch.cuda The PR moved DeepSpeedStaticLayer to the 'write_position = next slot to write == cached-token count' convention (the graph-path kernels write at the position then increment, so get_seq_length must return the position itself, not position + 1). Continuous batching is unaffected: it passes cache_position explicitly and never reads get_seq_length. The three master tests pinned the old +1 convention and cleared-row 0; updated to the new contract (-1 is the empty-row marker, matching cache initialization). Also replace torch.cuda.is_available() with get_accelerator() in the new kernel test file (check-torchcuda hook). Local: static_cache 6/6, segment_ki CPU 4/4 pass. Signed-off-by: Guokai Ma <guokai.ma@intel.com>
check-torchcuda also rejects tensor.cuda() calls; use get_accelerator().device_name() consistently (module-level DEV) so the tests are accelerator-portable alongside the GPU_AVAILABLE gate. Signed-off-by: Guokai Ma <guokai.ma@intel.com>
Signed-off-by: Ma, Guokai <guokai.ma@intel.com>
There was a problem hiding this comment.
💡 Codex Review
Here are some automated review suggestions for this pull request.
Reviewed commit: b31faddfc4
ℹ️ About Codex in GitHub
Codex has been enabled to automatically review pull requests in this repo. Reviews are triggered when you
- Open a pull request for review
- Mark a draft as ready
- Comment "@codex review".
If Codex has suggestions, it will comment; otherwise it will react with 👍.
When you sign up for Codex through ChatGPT, Codex can also answer questions or update the PR, like "@codex address that feedback".
| from deepspeed.module_inject.kernel_reference import dual_gemv_silu_mul | ||
| return self.down_proj(dual_gemv_silu_mul(input, self.gate_proj.weight, self.up_proj.weight)) |
There was a problem hiding this comment.
Preserve the original MLP activation
The structural selector accepts every bias-free module with these three projection names, but this replacement always computes SiLU. A GELU-gated MLP such as Gemma satisfies that selector and will therefore silently produce different logits whenever use_segki=True; constrain detection to SiLU MLPs or retain/call the module's configured activation.
Useful? React with 👍 / 👎.
| grad_gate_w = torch.matmul(grad_gate_coef.t(), hidden) | ||
| grad_up_w = torch.matmul(grad_up_coef.t(), hidden) |
There was a problem hiding this comment.
Flatten GLU activations before computing gradients
For normal decoder MLP inputs shaped [batch, sequence, hidden], this branch is used during backward, but Tensor.t() only accepts tensors with at most two dimensions. Consequently, a training step through any injected GLU with a 3-D activation fails at this expression; flatten the leading dimensions as the GDN backward path does before forming the weight gradients.
Useful? React with 👍 / 👎.
| if batch_size == 1 and seq_len == 1: | ||
| fused = GDNInputProj.apply( | ||
| hidden_states.squeeze(0).squeeze(0), self.in_proj_qkv.weight, self.in_proj_z.weight, self.in_proj_b.weight, | ||
| self.in_proj_a.weight, gdn_kernel_op).view(1, 1, -1) |
There was a problem hiding this comment.
Fall back when segment-KI inputs are not bf16
A CUDA use_segki=True rollout installs the native GDN op regardless of model dtype, and the b=1 decode path passes it here. quad_gemv unconditionally obtains BFloat16 data pointers, so ordinary fp16 Qwen3.5 inference raises instead of using the composite path; gate native dispatch on bf16 (and contiguous compatible weights) or provide an fp16 kernel.
Useful? React with 👍 / 👎.
| if getattr(module, "_ki_orig_layer_forward", None) is None: | ||
| module._ki_orig_layer_forward = module.forward | ||
| module._ki_norm_op = kernel_op | ||
| module._ki_postattn_scale = torch.add(1.0, norms[1].weight.detach()).contiguous() | ||
| module.forward = _fused_norm_layer_forward.__get__(module, type(module)) |
There was a problem hiding this comment.
Keep fused RMSNorm scale tied to its Parameter
This creates a detached, independent copy of the RMSNorm scale at injection time. After a co-located training step updates post_attention_layernorm.weight, graph decode still multiplies by the stale copy, so rollout no longer reflects the trained model; compute the scale from the live parameter in the forward path or refresh it after updates.
Useful? React with 👍 / 👎.
| graph_op.decode_step_graph(static_logits[:, -1, :].contiguous(), static_token.view(batch_size, 1), | ||
| write_pos, static_attn, full_token_buf) |
There was a problem hiding this comment.
Do not invoke the bf16 graph kernel for fp16 logits
The graph-step op is loaded and used for every batch-one graph capture, even when use_segki is false, but decode_step_graph rejects logits that are not CUDA bf16. Existing fp16 graph-capture rollouts therefore fail during capture rather than reaching a fallback; only add this node for bf16 logits or retain the PyTorch argmax/update path.
Useful? React with 👍 / 👎.
| int64_t new_pos = write_pos[0] + 1; | ||
| write_pos[0] = new_pos; | ||
| out_buf[new_pos] = best; | ||
| if (new_pos + 1 < max_len) { mask[new_pos + 1] = true; } |
There was a problem hiding this comment.
Reveal the current KV slot after graph replay
After the first captured decode forward, new_pos is the slot that will hold the next input token, but this reveals new_pos + 1. Thus any full-attention layer that remains on native SDPA sees its current KV slot masked and an unwritten future slot unmasked on the next replay, corrupting generation for models not patched by install_decode_attention; reveal new_pos instead.
Useful? React with 👍 / 👎.
| layer_types = getattr(text_config, "layer_types", None) or ["full_attention"] * num_layers | ||
| self._layers = [ | ||
| DSStaticGDNSlot(max_cache_len) if t == "linear_attention" else DeepSpeedStaticLayer(max_cache_len) | ||
| for t in layer_types |
There was a problem hiding this comment.
Keep hybrid cache slots compatible with continuous batching
For hybrid configurations this now substitutes DSStaticGDNSlot into the cache layer list, but continuous batching still iterates every layer and unconditionally copies target_layer.keys and target_layer.values in _continuous_prefill. The first admitted request on a model with a linear_attention layer therefore raises AttributeError; skip/pass through linear slots in that prefill and compaction path.
Useful? React with 👍 / 👎.
| } | ||
|
|
||
| return steps_done; |
There was a problem hiding this comment.
Pad the CUDA output buffer from device code
When an EOS is found at a periodic check, buf_ptr is a CUDA device pointer, yet this host-side loop dereferences it directly. On normal CUDA allocations this is an invalid host memory access (and can crash) precisely when EOS is encountered; launch a padding kernel or use a device tensor operation instead.
Useful? React with 👍 / 👎.
…ed dispatch, live norm scale Codex review follow-ups: - only fuse gated MLPs whose activation is positively SiLU; the fused GLU hardcodes SiLU and would silently change logits for GELU-gated families - flatten leading dims in the GLU backward weight-gradient GEMMs; .t() rejects the [batch, seq, hidden] training activations - gate native dual_gemv/quad_gemv dispatch on bf16 + contiguous inputs so fp16 models fall back to the composite reference instead of raising - recompute the fused RMSNorm 1+w scale from the live Parameter in the forward path; the injected snapshot went stale after optimizer steps Signed-off-by: Guokai Ma <guokai.ma@intel.com>
…ous batching - only mount decode_step_graph for CUDA bf16 logits; fp16 rollouts keep the PyTorch argmax path instead of failing during graph capture - _continuous_prefill skips pass-through GDN slots (no keys/values to copy) - DeepSpeedStaticCache.compact only compacts KV layers; GDN slots keep no per-row rows, and layer[0] may be a GDN slot Signed-off-by: Guokai Ma <guokai.ma@intel.com>
- reveal new_pos (not new_pos+1) after a decode step: new_pos is the slot the written token occupies on the next replay; the off-by-one masked the current slot and unmasked an unwritten one for native-SDPA layers - decode_loop pads the post-EOS tail with a device kernel; writing the device pointer from host on EOS was an illegal host memory access Signed-off-by: Guokai Ma <guokai.ma@intel.com>
Qwen3.5 MLPs carry transformers.activations.SiLUActivation, not torch.nn.SiLU, so the newly added SiLU gate silently skipped every GLU/GDN segment on Qwen3.5 (caught by the rollout-train-rollout integration test asserting segments_replaced > 0). Verified on RTX 4080 SUPER: 12/12 tests pass (incl. integration), E2E Qwen3.5-4B text correct at 69.0 tok/s with segment fusion restored. Signed-off-by: Guokai Ma <guokai.ma@intel.com>
There was a problem hiding this comment.
Copilot review overview
🟡 Changes recommended
Several paths currently break gradients, padded prompts, fallback decoding, hybrid continuous batching, and native-op compilation or packaging.
Review effort: Balanced
Findings: 10
Open (18)
Missing CUDA bf16 and pybind functional headers · New Native gdn_gates path cuts off training gradients · New Tensor truth-value check breaks packed-sequence calls · New Custom attention ignores padding masks · New Fused add-norm loses gradients during training · New Unbound GDN slots are skipped without state initialization · New Prompt attention mask is dropped during prefill · New Fallback graph reuses already-consumed recurrent state · New Forward-only graph replay uses uncaptured token buffers · New EOS padding slice ignores prompt-length offset · New Dual-GEMV excludes normal 3D greedy-decode inputs · New Non-CUDA backends incorrectly load CUDA builder · New FP16 models incorrectly enter bf16-only decode kernels · New Decode loop builder is not registered for wheel builds · New Fused GLU builder is not registered for wheel builds · New CUDA-only tests run on unsupported accelerators · New Missing direct coverage for native decode fallback · New Gemma is documented as supported despite GELU rejection · New
What changed in this PR
Adds segment-based kernel injection and native decode acceleration for hybrid-engine greedy rollout.
Changes:
- Adds fused CUDA kernels, reference implementations, and op builders.
- Extends static-cache support for hybrid attention models.
- Integrates graph capture, native decode loops, documentation, and tests.
| File | Description |
|---|---|
csrc/module_inject/fused_glu.cu |
Implements fused CUDA kernels. |
csrc/module_inject/decode_loop.cu |
Implements the native decode loop. |
deepspeed/module_inject/segment_ki.py |
Detects and injects fused model segments. |
deepspeed/module_inject/kernel_reference.py |
Provides PyTorch kernel references. |
deepspeed/ops/module_inject/__init__.py |
Exports native-op builders. |
deepspeed/ops/module_inject/fused_glu.py |
Adds the fused-kernel builder. |
deepspeed/ops/module_inject/decode_loop.py |
Adds the decode-loop builder. |
deepspeed/utils/static_cache.py |
Adds hybrid cache slots and position semantics. |
deepspeed/runtime/rollout/hybrid_engine_rollout.py |
Integrates injection and accelerated decoding. |
tests/unit/module_inject/test_segment_ki_kernels.py |
Tests kernels and injection behavior. |
tests/unit/utils/test_static_cache.py |
Updates static-cache expectations. |
docs/code-docs/source/inference-engine.rst |
Documents segment-based injection. |
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
…citly __nv_bfloat16 and the std::function<void()> pybind type caster were only reachable through transitive includes; declare both directly so the extension does not depend on include order. Signed-off-by: Guokai Ma <guokai.ma@intel.com>
The native gdn_gates and fused_add_norm ops have no autograd registration, so a training step through their fused branches silently cut gradients (gates: a/b projections, A_log, dt_bias; fused_add_norm: the residual/norm path). Route grad-enabled forwards through the composite/native ops; the kernels still serve the launch-bound no-grad decode path. Signed-off-by: Guokai Ma <guokai.ma@intel.com>
kwargs.get('cu_seq_lens_q') or kwargs.get('cu_seqlens') evaluates the
tensor's boolean value, which raises for any multi-element packed-sequence
input. Use dict.get's default argument instead.
Signed-off-by: Guokai Ma <guokai.ma@intel.com>
The decode_attn kernel attends every cached slot up to write_pos unconditionally and never sees attention_mask, so a padded prompt leaks pad positions into the softmax and yields logits that differ from the original path. Rollout flags padded prompts once per generate call and the injected forward rides the SDPA path for them. Signed-off-by: Guokai Ma <guokai.ma@intel.com>
The capture forward consumes static_token and advances the GDN conv/recurrent buffers whether or not decode_step_graph was mounted, so gating the restore on graph_op left the fallback graph (native op unavailable) starting from a state that already consumed the first decode token. Signed-off-by: Guokai Ma <guokai.ma@intel.com>
Mounting decode_step_graph is conditional (b=1, CUDA bf16 logits), but the loop selector only checked that the op module loaded — a b=1 fp16 rollout would replay a forward-only graph and read a token buffer the kernel never wrote. Track mounting with a flag and branch on it. Signed-off-by: Guokai Ma <guokai.ma@intel.com>
step is response-relative while full_token_buf carries the prompt prefix, so [step + 2:] on a short prompt padded from inside the prompt and erased valid tokens including the detected EOS. Offset by prompt_len and start after the EOS slot. Signed-off-by: Guokai Ma <guokai.ma@intel.com>
…kernel HF decode activations reach the MLP as [1, 1, hidden], but the kernel gate required dim()==2, so the dual-weight GEMV never ran in the greedy decode path it was built for. Accept any input that is exactly one hidden vector and reshape around the kernel call. Signed-off-by: Guokai Ma <guokai.ma@intel.com>
The decode_attn guard checked shapes and cache identity but not dtype, so a pure-fp16 model passed the guard and crashed in the kernel's bf16 TORCH_CHECK during warmup/capture instead of riding the original attention path. Check dtype and contiguity like the fused-norm gate. Signed-off-by: Guokai Ma <guokai.ma@intel.com>
The builders lived under deepspeed/ops/module_inject/, which setup.py's op_builder scan never sees, so DS_BUILD_FUSED_GLU/DS_BUILD_DECODE_LOOP were silently ignored at install time and the ops could only JIT-build on hosts with a compiler toolchain. Move the builder classes next to the other CUDAOpBuilders and keep the ops/module_inject modules as thin runtime loaders, matching how fused_adam and friends are wired. Signed-off-by: Guokai Ma <guokai.ma@intel.com>
GPU_AVAILABLE keyed on device_name != 'cpu', so XPU/HPU/NPU/MPS CI ran the CUDA-only kernel tests and failed at JIT load instead of skipping. Ask the accelerator whether it registers the segment-KI builder; backends without the op skip, and one that implements it later runs them. Signed-off-by: Guokai Ma <guokai.ma@intel.com>
The GLU replacement hardcodes SiLU and the selector now positively verifies the module's activation, so GELU-gated families such as Gemma are passed through. State that in the docs instead of listing Gemma as picked up. Signed-off-by: Guokai Ma <guokai.ma@intel.com>
With the fused_glu op loaded, b=1 greedy decode rides the full-step graph (decode_step_graph is captured inside the graph) and b>1 uses graph+Python argmax, so the standalone C++ loop was only reachable when fused_glu failed to JIT while decode_loop compiled — a corner that also crashed on fp16 logits. The layer duplicated the step-update kernel (reveal/pad fixes had to land twice) and had no test coverage. Remove the op, its builder, and the rollout branch; decode now falls back directly to the Python loop. Signed-off-by: Guokai Ma <guokai.ma@intel.com>
An fp16 model crashed inside gdn_gates ('expected BFloat16 but found
Half') because only the grad-enabled case was gated; the kernel reads
raw BFloat16 pointers. Route non-bf16 inputs through the composite ops
like the other native dispatch points.
Signed-off-by: Guokai Ma <guokai.ma@intel.com>
The graph prefill passes per-type None attention masks, so pad positions leak into the KV cache and GDN recurrent states and the first token is argmaxed off a pad position — a padded prompt came back as echoed EOS tokens. Fall back to module.generate, which owns the full mask semantics, whenever the prompt mask has padding. Signed-off-by: Guokai Ma <guokai.ma@intel.com>
…tion apply_segment_ki loaded the CUDA op bare whenever device_name() != 'cpu', so a CUDA host without a usable JIT toolchain raised out of rollout construction instead of falling back, and non-CUDA backends (XPU/HPU/ NPU/MPS) attempted a doomed CUDA JIT build. Ask the accelerator whether it registers the builder (NotImplemented backends skip), and swallow load failures into the composite fallback — matching the kernel-test gating. Signed-off-by: Guokai Ma <guokai.ma@intel.com>
GDN layers keep per-row conv/recurrent state that the continuous- batching cache neither scatters on admission nor compacts on retirement, so a hybrid run crashed mid-decode on an unbound pass-through slot. Raise a clear ValueError during input validation instead. Signed-off-by: Guokai Ma <guokai.ma@intel.com>
Signed-off-by: Guokai Ma <guokai.ma@intel.com>
Conflicts in hybrid_engine_rollout.py resolved as: - config fields and continuous-batching validations: keep both sides (use_segki + align_decode_fronts/enable_cache_trimming/ continuous_cache_capacity; hybrid rejection + shared-prefill rejection) - _continuous_prefill: take master's refactored structure (_continuous_prefill_cache/_cache_layer_count/_cache_layer_tensors); the hybrid GDN skip is unnecessary there because _validate_continuous_ inputs now rejects hybrid models up front - _generate_graph KV copy: master's helper-based loop plus the GDN slot bind for hybrid caches Signed-off-by: Guokai Ma <guokai.ma@intel.com>
Signed-off-by: Guokai Ma <guokai.ma@intel.com>
|
Hi @delock, thanks for submitting this PR. I have left a few comments and questions. |
There was a problem hiding this comment.
Copilot review overview
🟡 Changes recommended
Graph compatibility, packed GDN execution, projection bias handling, and deterministic argmax behavior have unresolved correctness issues.
Review effort: Balanced
Findings: 3
Open (6)
Preserve packed-sequence boundaries through causal convolution · New Disable fused QKV projection when any bias is present · New Pass tensor masks for ordinary models during graph capture · New Match torch.argmax first-index tie breaking · New Add direct CUDA coverage for decode_step_graph · New Make live-weight integration test deterministically observe weight changes · New
Resolved since last review (18)
EOS padding slice ignores prompt-length offset Forward-only graph replay uses uncaptured token buffers Fallback graph reuses already-consumed recurrent state Prompt attention mask is dropped during prefill Unbound GDN slots are skipped without state initialization Fused add-norm loses gradients during training Custom attention ignores padding masks Tensor truth-value check breaks packed-sequence calls Native gdn_gates path cuts off training gradients Missing CUDA bf16 and pybind functional headers CUDA-only tests run on unsupported accelerators Fused GLU builder is not registered for wheel builds Decode loop builder is not registered for wheel builds FP16 models incorrectly enter bf16-only decode kernels Non-CUDA backends incorrectly load CUDA builder Dual-GEMV excludes normal 3D greedy-decode inputs Gemma is documented as supported despite GELU rejection Missing direct coverage for native decode fallback
| kw.pop("use_cache", None) | ||
| kw.pop("cu_seqlens", None) | ||
| kw.pop("cu_seq_lens_q", None) |
| module._ki_qkv_op = kernel_op if (hasattr(kernel_op, "triple_gemv") | ||
| and os.environ.get("DS_TIER2", "1") == "1") else None |
| attention_mask={ | ||
| "full_attention": None, | ||
| "linear_attention": None | ||
| }, |
| if (tid < stride) { | ||
| if (s_val[tid + stride] > s_val[tid]) { | ||
| s_val[tid] = s_val[tid + stride]; | ||
| s_idx[tid] = s_idx[tid + stride]; | ||
| } | ||
| } |
| m.def("decode_attn", &decode_attn, "b=1 decode attention with GQA, graph-compatible (CUDA)"); | ||
| m.def("decode_step", &decode_step, "fused decode step update (CUDA)"); | ||
| m.def("decode_step_graph", &decode_step_graph, "graph-capturable decode step (CUDA)"); |
| batch2 = rollout.generate(req, greedy) | ||
| assert not torch.equal(batch1.input_ids, batch2.input_ids), \ | ||
| "rollout output unchanged after training — kernels read stale weights" |
Hi @nathon-lee , kindly reminder in case you didn't submit your review comments. Thanks! |
The conv and scan closures popped cu_seq_lens_q/cu_seqlens before calling causal_conv1d_fn and the FLA scan, letting state bleed across packed sequences; upstream passes them through. Filter only unrelated generation kwargs (use_cache). Signed-off-by: Guokai Ma <guokai.ma@intel.com>
triple_gemv consumes weights only, so enabling it for projections with biases (Qwen3.5 attention_bias) silently dropped all three biases and produced different logits. Require bias-free q/k/v for the fused path; the composite nn.Linear forward keeps honoring them. Signed-off-by: Guokai Ma <guokai.ma@intel.com>
The graph path passed the per-type attention-mask dict unconditionally, but the dict form is a hybrid-model contract: standard full-attention models (LLaMA and families) expect a tensor/None and regress under use_graph_capture even with use_segki=False. Build the per-type mapping only when layer_types contains linear_attention and keep the tensor form (2-D prompt mask at prefill, 4-D static mask at warmup/capture) otherwise. Signed-off-by: Guokai Ma <guokai.ma@intel.com>
The warp reduction kept the lower thread ID on equal bf16 maxima, but thread strides do not follow vocabulary order, so tied logits (common at bf16 with large vocabularies) could argmax to a different token than torch.argmax. Break ties on the smaller vocabulary index. Signed-off-by: Guokai Ma <guokai.ma@intel.com>
The b=1 full-step kernel had no direct coverage: verify argmax against torch.argmax including a tied-logits case (first-index rule), write_pos advancement, one-slot-per-step mask revelation, and output-buffer updates across repeated calls. Signed-off-by: Guokai Ma <guokai.ma@intel.com>
The rollout-train-rollout test asserted the greedy sequence changes after an SGD step, but changed logits need not change any argmax, so a correct live-weight implementation could fail it. Probe the injected projection directly: its output on fixed input must differ once the optimizer updates the weights it reads. The post-training generate still runs for path coverage. Signed-off-by: Guokai Ma <guokai.ma@intel.com>
Both tiers default to enabled; a second escape hatch beyond the config flags adds env lookups without a user. Kernel injection is already gated by use_segki and the per-dispatch checks. Signed-off-by: Guokai Ma <guokai.ma@intel.com>



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.