Skip to content
Open
Show file tree
Hide file tree
Changes from 33 commits
Commits
Show all changes
40 commits
Select commit Hold shift + click to select a range
d41c914
feat: segment-based kernel injection for hybrid-engine rollout
delock Sep 26, 2026
1ae2d48
style: apply yapf and clang-format (CI formatting checks)
delock Sep 26, 2026
2bbb133
fix: gate decode_attn/fused_norm injection on use_segki
delock Sep 26, 2026
3f0e824
feat: GDN input projections read original weights (co-location safe)
delock Sep 27, 2026
f6ce9ce
cleanup: drop redundant noqa (covered by per-file-ignores in .flake8)
delock Sep 27, 2026
2a0cfc9
test: kernel-reference consistency + rollout-train-rollout integration
delock Sep 27, 2026
9c0177b
fix: align static_cache tests with write-position convention; drop to…
delock Sep 28, 2026
97115ce
fix: drop .cuda() and hardcoded cuda device from kernel tests
delock Sep 28, 2026
b31fadd
docs: document use_segki segment-based kernel injection
delock Sep 28, 2026
415d3de
fix(segment_ki): SiLU-only GLU detection, 3-D-safe backward, bf16-gat…
delock Sep 28, 2026
87fddf9
fix(rollout): dtype-gate graph step kernel; hybrid-cache-safe continu…
delock Sep 28, 2026
8d99be3
fix(kernels): reveal current KV slot; pad output tail on the device
delock Sep 28, 2026
f4b9393
fix(segment_ki): recognize transformers' SiLUActivation in GLU detection
delock Sep 28, 2026
5ae7bd7
fix(decode_loop): include cuda_bf16.h and pybind11/functional.h expli…
delock Oct 5, 2026
e555539
fix(segment_ki): keep no-autograd kernels on no-grad paths only
delock Oct 5, 2026
6946a91
fix(segment_ki): do not truth-test the cu_seq_lens tensor
delock Oct 5, 2026
c3c8c1f
fix(segment_ki): fall back to SDPA attention for padded prompts
delock Oct 5, 2026
1121aaa
fix(rollout): restore GDN states after capture regardless of step kernel
delock Oct 5, 2026
3b17495
fix(rollout): gate full-step-graph decode loop on actual kernel mounting
delock Oct 5, 2026
229b5fa
fix(rollout): pad the token buffer in absolute coordinates after EOS
delock Oct 5, 2026
11d6241
fix(segment_ki): route 3-D single-token GLU activations to the fused …
delock Oct 5, 2026
4f79316
fix(segment_ki): fall back for non-bf16 decode attention inputs
delock Oct 5, 2026
51a791a
refactor: move segment-KI op builders into op_builder
delock Oct 5, 2026
b9872af
test(segment_ki): gate kernel tests on op-builder registration
delock Oct 5, 2026
efba6c8
docs: drop Gemma from the seg-KI family list
delock Oct 5, 2026
3459bda
refactor: drop the C++ decode-loop fallback layer
delock Oct 5, 2026
77121e4
fix(segment_ki): gate gdn_gates native dispatch on bf16 too
delock Oct 5, 2026
c2f9b5c
fix(rollout): route padded prompts away from the graph path
delock Oct 5, 2026
6239052
fix(segment_ki): best-effort native op load keyed on builder registra…
delock Oct 5, 2026
2340597
fix(rollout): reject continuous batching on hybrid models up front
delock Oct 5, 2026
0f1d1b0
style: apply yapf to the review-fix changes
delock Oct 5, 2026
ff93f09
Merge master into gma/qwen_ki_expr
delock Oct 5, 2026
e484d2d
style: reword a comment codespell flags
delock Oct 5, 2026
0db6caf
fix(segment_ki): keep packed-sequence boundaries in the GDN conv/scan
delock Oct 8, 2026
5adfb20
fix(segment_ki): keep fused QKV off for biased attention projections
delock Oct 8, 2026
e9d812b
fix(rollout): pass a tensor mask to non-hybrid models on the graph path
delock Oct 8, 2026
537d040
fix(kernels): match torch.argmax first-index tie rule in step kernels
delock Oct 8, 2026
4c3371e
test(segment_ki): direct contract test for decode_step_graph
delock Oct 8, 2026
dff71e0
test(segment_ki): assert the live-weight contract deterministically
delock Oct 8, 2026
1f1371b
refactor(segment_ki): drop the DS_TIER1/DS_TIER2 env kill switches
delock Oct 8, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
745 changes: 745 additions & 0 deletions csrc/module_inject/fused_glu.cu

Large diffs are not rendered by default.

96 changes: 96 additions & 0 deletions deepspeed/module_inject/kernel_reference.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,96 @@
# SPDX-License-Identifier: Apache-2.0
# DeepSpeed Team
"""Reference implementations (executable specifications) for segment-KI kernels.

Each function here defines the mathematically correct output for its
corresponding custom kernel. Custom kernels (CUDA today, SYCL/XPU in the
future) must produce results equivalent to these functions within bf16
tolerance. Three roles in one:

1. **Specification**: an agent porting kernels to a new backend reads these
to understand what the kernel must compute.
2. **Runtime fallback**: the forward paths in segment_ki call these when no
custom kernel is available (e.g. CPU-only, XPU without compiled kernels).
3. **Test oracle**: unit tests compare custom kernel output against these
functions on identical inputs.

All implementations use only standard PyTorch ops — no custom kernels — so
they are correct by construction and run on any device.
"""

import math

import torch
import torch.nn.functional as F


def dual_gemv_silu_mul(hidden, gate_w, up_w):
"""silu(hidden @ gate_w.T) * (hidden @ up_w.T).

Reads gate and up weight matrices directly (no concatenation), so
gradients flow to the original Parameters. Works for any batch size.
"""
gate_out = torch.matmul(hidden, gate_w.t())
up_out = torch.matmul(hidden, up_w.t())
return F.silu(gate_out) * up_out


def gdn_gates(a, b, a_log, dt_bias):
"""GDN gating: beta = sigmoid(b), g = -exp(a_log) * softplus(a + dt_bias)."""
beta = torch.sigmoid(b)
g = -torch.exp(a_log.float()) * F.softplus(a.float() + dt_bias.float())
return beta, g.to(b.dtype)


def gdn_input_proj(hidden, w_qkv, w_z, w_b, w_a):
"""GDN input projections: cat(h @ Wqkv.T, h @ Wz.T, h @ Wb.T, h @ Wa.T)."""
parts = [torch.matmul(hidden, w.t()) for w in (w_qkv, w_z, w_b, w_a)]
return torch.cat(parts, dim=-1)


def triple_gemv(hidden, q_w, k_w, v_w):
"""QKV projection: (hidden @ q_w.T, hidden @ k_w.T, hidden @ v_w.T)."""
q = torch.matmul(hidden, q_w.t())
k = torch.matmul(hidden, k_w.t())
v = torch.matmul(hidden, v_w.t())
return q, k, v


def decode_attn(q, K, V, pos):
"""Single-query attention over valid KV positions.

q: [num_heads, head_dim]
K, V: [num_kv_heads, max_len, head_dim]
pos: int (number of valid positions, 0-indexed inclusive)

Returns [num_heads, head_dim] with GQA (K/V heads shared across query heads).
"""
num_q_heads = q.shape[0]
num_kv_heads = K.shape[0]
head_dim = q.shape[1]
n_rep = num_q_heads // num_kv_heads

K_valid = K[:, :pos + 1, :] # [nkv, pos+1, hd]
V_valid = V[:, :pos + 1, :]

if n_rep > 1:
K_valid = K_valid.repeat_interleave(n_rep, dim=0)
V_valid = V_valid.repeat_interleave(n_rep, dim=0)

scale = 1.0 / math.sqrt(head_dim)
scores = torch.einsum('hd,hld->hl', q.float(), K_valid.float()) * scale
weights = torch.softmax(scores, dim=-1)
out = torch.einsum('hl,hld->hd', weights, V_valid.float())
return out.to(q.dtype)


def fused_add_norm(hidden, residual, weight, eps=1e-5):
"""Residual add + RMSNorm: rms_norm(hidden + residual) * weight."""
combined = hidden.float() + residual.float()
rms = torch.rsqrt(combined.pow(2).mean(dim=-1, keepdim=True) + eps)
return (combined * rms * weight.float()).to(hidden.dtype)


def decode_step(logits, vocab_size):
"""Argmax over logits (greedy decode)."""
return torch.argmax(logits).item()
Loading
Loading