Skip to content
Open
Show file tree
Hide file tree
Changes from 13 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
163 changes: 163 additions & 0 deletions csrc/module_inject/decode_loop.cu
Original file line number Diff line number Diff line change
@@ -0,0 +1,163 @@
// SPDX-License-Identifier: Apache-2.0
Comment thread
hwchen2017 marked this conversation as resolved.
Outdated
// DeepSpeed Team

// C++ decode loop: runs the full generation loop in a single C++ call.
// The replay callable is a std::function<void()> — it can wrap a CUDA
// graph replay, a SYCL submission, an eager PyTorch forward, or any
// backend-specific mechanism. The C++ loop is agnostic to the backend.
//
// Per-step overhead: replay_fn() ≈ 2-5μs (pybind11 dispatch)
// vs Python for-loop ≈ 30-50μs.
// The fused step-update kernel launch is pure C++ (zero Python).

#include <ATen/cuda/CUDAContext.h>
#include <cuda_runtime.h>
#include <torch/extension.h>
#include <functional>
Comment thread
delock marked this conversation as resolved.
Outdated

__global__ void ds_loop_step_kernel(const __nv_bfloat16* __restrict__ logits,
int64_t* __restrict__ token_out,
int64_t* __restrict__ write_pos,
bool* __restrict__ mask,
int64_t* __restrict__ out_buf,
int step,
int vocab_size,
int max_len)
{
__shared__ int s_idx[1024];
__shared__ float s_val[1024];

int tid = threadIdx.x;
int local_idx = 0;
float local_val = -INFINITY;

for (int v = tid; v < vocab_size; v += blockDim.x) {
float val = __bfloat162float(logits[v]);
if (val > local_val) {
local_val = val;
local_idx = v;
}
}
s_idx[tid] = local_idx;
s_val[tid] = local_val;
__syncthreads();

for (int stride = blockDim.x / 2; stride > 0; stride >>= 1) {
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];
}
}
__syncthreads();
}

if (tid == 0) {
int64_t best = (int64_t)s_idx[0];
token_out[0] = best;
out_buf[step] = best;
int64_t new_pos = write_pos[0] + 1;
write_pos[0] = new_pos;
// new_pos is the slot the token just written occupies during the
// next replay; revealing new_pos+1 masks the current slot for
// layers still on native SDPA.
if (new_pos < max_len) { mask[new_pos] = true; }
}
}

// Fills buf[from, len) with pad on the device — used after an early EOS exit.
__global__ void pad_tail_kernel(int64_t* __restrict__ buf, int64_t from, int64_t len, int64_t pad)
{
int64_t i = from + (int64_t)blockIdx.x * blockDim.x + threadIdx.x;
if (i < len) { buf[i] = pad; }
}

int64_t ds_decode_loop(std::function<void()> replay_fn,
at::Tensor logits,
at::Tensor token_out,
at::Tensor write_pos,
at::Tensor mask,
at::Tensor out_buf,
int64_t max_steps,
int64_t eos_token_id,
int64_t pad_token_id,
int64_t eos_check_every)
Comment thread
delock marked this conversation as resolved.
Outdated
{
TORCH_CHECK(logits.is_cuda() && logits.scalar_type() == at::ScalarType::BFloat16,
"logits must be CUDA bf16");
TORCH_CHECK(logits.is_contiguous(), "logits must be contiguous");

// Accept [batch, vocab] 2D tensors. The kernel
// currently supports batch_size=1 only; select(0,0) is a zero-copy view.
// When the kernel gains batch support, this check will be lifted without
// any interface change on the Python side.
TORCH_CHECK(logits.dim() == 2 && logits.size(0) == 1,
"decode_loop currently supports batch_size=1, got batch=",
logits.dim() == 2 ? logits.size(0) : -1);
auto logits_1d = logits.select(0, 0);
auto token_1d = token_out.dim() == 2 ? token_out.select(0, 0) : token_out;
auto out_buf_1d = out_buf.dim() == 2 ? out_buf.select(0, 0) : out_buf;

auto stream = at::cuda::getCurrentCUDAStream().stream();
int vocab = (int)logits_1d.numel();
int max_len = (int)mask.size(-1);

const __nv_bfloat16* log_ptr =
reinterpret_cast<const __nv_bfloat16*>(logits_1d.data_ptr<at::BFloat16>());
int64_t* tok_ptr = token_1d.data_ptr<int64_t>();
int64_t* wp_ptr = write_pos.data_ptr<int64_t>();
bool* mask_ptr = mask.data_ptr<bool>();
int64_t* buf_ptr = out_buf_1d.data_ptr<int64_t>();

int64_t steps_done = 0;

for (int64_t step = 1; step < max_steps; step++) {
// 1) Invoke the replay callable (graph replay, eager forward, or
// any backend-specific mechanism — the C++ loop is agnostic)
replay_fn();

// 2) Fused step-update kernel (pure C++, zero Python)
ds_loop_step_kernel<<<1, 1024, 0, stream>>>(
log_ptr, tok_ptr, wp_ptr, mask_ptr, buf_ptr, (int)step, vocab, max_len);

steps_done = step + 1;

// 3) Periodic EOS check (D2H sync, amortized over N steps)
if (eos_token_id >= 0 && step % eos_check_every == 0) {
int64_t token;
cudaMemcpyAsync(&token, tok_ptr, sizeof(int64_t), cudaMemcpyDeviceToHost, stream);
cudaStreamSynchronize(stream);
if (token == eos_token_id) {
// buf_ptr is a device pointer: pad the remaining slots with a
// kernel on the same stream instead of writing it from host.
int64_t remaining = max_steps - (step + 1);
if (remaining > 0) {
int threads = 128;
int64_t blocks = (remaining + threads - 1) / threads;
pad_tail_kernel<<<(int)blocks, threads, 0, stream>>>(
buf_ptr, step + 1, max_steps, pad_token_id);
}
break;
}
}
}

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

}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)
{
m.def("decode_loop",
&ds_decode_loop,
"C++ decode loop: graph replay + step update, zero Python per step",
py::arg("replay_fn"),
py::arg("logits"),
py::arg("token_out"),
py::arg("write_pos"),
py::arg("mask"),
py::arg("out_buf"),
py::arg("max_steps"),
py::arg("eos_token_id"),
py::arg("pad_token_id"),
py::arg("eos_check_every"));
}
Loading
Loading