Repository navigation
feat(rollout): segment-based kernel injection for hybrid-engine greedy decode #8602
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: master
Are you sure you want to change the base?
Changes from 13 commits
d41c914
1ae2d48
2bbb133
3f0e824
f6ce9ce
2a0cfc9
9c0177b
97115ce
b31fadd
415d3de
87fddf9
8d99be3
f4b9393
5ae7bd7
e555539
6946a91
c3c8c1f
1121aaa
3b17495
229b5fa
11d6241
4f79316
51a791a
b9872af
efba6c8
3459bda
77121e4
c2f9b5c
6239052
2340597
0f1d1b0
ff93f09
e484d2d
0db6caf
5adfb20
e9d812b
537d040
4c3371e
dff71e0
1f1371b
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,163 @@ | ||
| // SPDX-License-Identifier: Apache-2.0 | ||
| // 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> | ||
|
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) | ||
|
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; | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
When an EOS is found at a periodic check, 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")); | ||
| } | ||
Uh oh!
There was an error while loading. Please reload this page.