Skip to content

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

Open
delock wants to merge 40 commits into
masterfrom
gma/qwen_ki_expr
Open

delock wants to merge 40 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>
…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>
@delock
delock marked this pull request as ready for review September 28, 2026 06:27

@chatgpt-codex-connector chatgpt-codex-connector Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

💡 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".

Comment thread csrc/module_inject/decode_loop.cu Outdated
Comment on lines +184 to +185
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))

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

P1 Badge 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 👍 / 👎.

Comment thread deepspeed/module_inject/segment_ki.py Outdated
Comment on lines +123 to +124
grad_gate_w = torch.matmul(grad_gate_coef.t(), hidden)
grad_up_w = torch.matmul(grad_up_coef.t(), hidden)

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

P1 Badge 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 👍 / 👎.

Comment on lines +233 to +236
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)

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

P1 Badge 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 👍 / 👎.

Comment on lines +590 to +594
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))

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

P1 Badge 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 👍 / 👎.

Comment on lines +916 to +917
graph_op.decode_step_graph(static_logits[:, -1, :].contiguous(), static_token.view(batch_size, 1),
write_pos, static_attn, full_token_buf)

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

P1 Badge 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 👍 / 👎.

Comment thread csrc/module_inject/fused_glu.cu Outdated
Comment on lines +198 to +201
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; }

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

P1 Badge 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 👍 / 👎.

Comment on lines +273 to +276
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

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

P1 Badge 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 👍 / 👎.

Comment thread csrc/module_inject/decode_loop.cu Outdated
Comment on lines +125 to +127
}

return steps_done;

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

P2 Badge 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>

Copilot AI left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

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 High severity · 6 Medium severity · 2 Low severity

Open (18)
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.

Comment thread csrc/module_inject/decode_loop.cu Outdated
Comment thread deepspeed/module_inject/segment_ki.py Outdated
Comment thread deepspeed/module_inject/segment_ki.py Outdated
Comment thread deepspeed/module_inject/segment_ki.py
Comment thread deepspeed/module_inject/segment_ki.py
Comment thread deepspeed/ops/module_inject/decode_loop.py Outdated
Comment thread deepspeed/ops/module_inject/fused_glu.py Outdated
Comment thread tests/unit/module_inject/test_segment_ki_kernels.py Outdated
Comment thread csrc/module_inject/decode_loop.cu Outdated
Comment thread docs/code-docs/source/inference-engine.rst Outdated
delock added 5 commits October 5, 2026 18:54
…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>
delock added 15 commits October 5, 2026 18:54
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>
@nathon-lee

Copy link
Copy Markdown
Contributor

Hi @delock, thanks for submitting this PR. I have left a few comments and questions.

Copilot AI left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

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 High severity · 1 Medium severity · 2 Low severity

Open (6)
Resolved since last review (18)

Comment thread deepspeed/module_inject/segment_ki.py Outdated
Comment on lines +374 to +376
kw.pop("use_cache", None)
kw.pop("cu_seqlens", None)
kw.pop("cu_seq_lens_q", None)
Comment thread deepspeed/module_inject/segment_ki.py Outdated
Comment on lines +562 to +563
module._ki_qkv_op = kernel_op if (hasattr(kernel_op, "triple_gemv")
and os.environ.get("DS_TIER2", "1") == "1") else None
Comment on lines +1036 to +1039
attention_mask={
"full_attention": None,
"linear_attention": None
},
Comment on lines +109 to +114
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];
}
}
Comment on lines +731 to +733
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)");
Comment on lines +275 to +277
batch2 = rollout.generate(req, greedy)
assert not torch.equal(batch1.input_ids, batch2.input_ids), \
"rollout output unchanged after training — kernels read stale weights"
@delock

delock commented Oct 8, 2026

Copy link
Copy Markdown
Collaborator Author

Hi @delock, thanks for submitting this PR. I have left a few comments and questions.

Hi @nathon-lee , kindly reminder in case you didn't submit your review comments. Thanks!

delock added 7 commits October 8, 2026 09:11
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>

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