[Spec] Fuse small kenrels under gather_spec_extras (#27233)

This commit is contained in:
Khoa Pham
2026-06-08 15:02:04 -07:00
committed by GitHub
parent b5c64b94d5
commit c95179bc85
3 changed files with 334 additions and 21 deletions
+3 -21
View File
@@ -1,11 +1,12 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Optional, Sequence, Union
from typing import TYPE_CHECKING, Sequence, Union
import torch
from sglang.srt.environ import envs
from sglang.srt.speculative.spec_utils import spec_need_hidden_states
from sglang.srt.speculative.triton_ops.gather_spec_extras import gather_spec_extras
from sglang.srt.utils import is_cuda, is_hip, is_npu
if TYPE_CHECKING:
@@ -59,25 +60,6 @@ def _assert_nonneg_and_invalidate(
buf[indices] = -1
@torch.compile(dynamic=True, disable=_is_npu)
def _gather_spec_extras(
indices: torch.Tensor,
topk_p_buf: torch.Tensor,
topk_index_buf: torch.Tensor,
output_tokens_buf: torch.Tensor,
hidden_states_buf: Optional[torch.Tensor],
):
"""Compiled gather of spec extras. `hidden_states_buf` is None when the
build does not capture hidden states."""
topk_p = topk_p_buf[indices]
topk_index = topk_index_buf[indices]
bonus_tokens = output_tokens_buf[indices]
hidden_states = (
hidden_states_buf[indices] if hidden_states_buf is not None else None
)
return topk_p, topk_index, bonus_tokens, hidden_states
def resolve_forward_inputs(batch: ScheduleBatch, future_map: FutureMap) -> None:
"""Materialize input_ids at forward entry. Two sources:
@@ -200,7 +182,7 @@ class FutureMap:
draft_input.topk_index,
draft_input.bonus_tokens,
hidden_states,
) = _gather_spec_extras(
) = gather_spec_extras(
indices,
self.topk_p_buf,
self.topk_index_buf,
@@ -0,0 +1,117 @@
from __future__ import annotations
from typing import Optional
import torch
import triton
import triton.language as tl
@triton.jit
def _gather_rows_kernel(
idx_ptr,
s0,
d0,
n0,
s1,
d1,
n1,
s2,
d2,
n2,
s3,
d3,
n3,
HAS3: tl.constexpr,
BLOCK: tl.constexpr,
):
# One program == one (output row, column block). All buffers share the
# same gather index, so a single launch copies every buffer's row and
# the per-kernel launch bubbles between the old separate gathers vanish.
row = tl.program_id(0)
cb = tl.program_id(1)
src = tl.load(idx_ptr + row).to(tl.int64)
cols = cb * BLOCK + tl.arange(0, BLOCK)
m0 = cols < n0
tl.store(d0 + row * n0 + cols, tl.load(s0 + src * n0 + cols, mask=m0), mask=m0)
m1 = cols < n1
tl.store(d1 + row * n1 + cols, tl.load(s1 + src * n1 + cols, mask=m1), mask=m1)
m2 = cols < n2
tl.store(d2 + row * n2 + cols, tl.load(s2 + src * n2 + cols, mask=m2), mask=m2)
if HAS3:
m3 = cols < n3
tl.store(d3 + row * n3 + cols, tl.load(s3 + src * n3 + cols, mask=m3), mask=m3)
def _row_width(buf: torch.Tensor) -> int:
"""Flattened per-row element count (trailing dims), 1 for a 1-D buffer."""
return buf[0].numel() if buf.dim() > 1 else 1
def _empty_like_rows(buf: torch.Tensor, m: int) -> torch.Tensor:
"""Output buffer for `m` gathered rows of `buf` (same trailing dims/dtype/device)."""
return torch.empty((m, *buf.shape[1:]), dtype=buf.dtype, device=buf.device)
def gather_spec_extras(
indices: torch.Tensor,
topk_p_buf: torch.Tensor,
topk_index_buf: torch.Tensor,
output_tokens_buf: torch.Tensor,
hidden_states_buf: Optional[torch.Tensor],
):
"""Gather spec extras (topk_p / topk_index / bonus_tokens / optional hidden
states) by a shared row index in a single fused Triton launch (one kernel
for all buffers) instead of one advanced-index gather per buffer.
`hidden_states_buf` is None when the build does not capture hidden states."""
# Source buffers are allocated once (torch.empty/full) and only ever mutated
# in place, so they are guaranteed row-contiguous. `indices` flows from
# several producers (req_pool_indices, filtered/merged future_indices); the
# kernel addresses it linearly, so normalize layout here (no-op when already
# contiguous) to avoid a silent wrong-result on a strided index tensor.
indices = indices.contiguous()
m = indices.shape[0]
has_hidden = hidden_states_buf is not None
topk_p = _empty_like_rows(topk_p_buf, m)
topk_index = _empty_like_rows(topk_index_buf, m)
bonus_tokens = _empty_like_rows(output_tokens_buf, m)
hidden_states = _empty_like_rows(hidden_states_buf, m) if has_hidden else None
if m == 0:
return topk_p, topk_index, bonus_tokens, hidden_states
n0 = _row_width(topk_p_buf)
n1 = _row_width(topk_index_buf)
n2 = _row_width(output_tokens_buf)
n3 = _row_width(hidden_states_buf) if has_hidden else 1
max_n = max(n0, n1, n2, n3)
# Dummy operands for the disabled hidden-states slot: the pointers must be
# valid even though the kernel never dereferences them (gated off by HAS3).
s3 = hidden_states_buf if has_hidden else indices
d3 = hidden_states if has_hidden else indices
block = min(1024, triton.next_power_of_2(max_n))
grid = (m, triton.cdiv(max_n, block))
_gather_rows_kernel[grid](
indices,
topk_p_buf,
topk_p,
n0,
topk_index_buf,
topk_index,
n1,
output_tokens_buf,
bonus_tokens,
n2,
s3,
d3,
n3,
HAS3=has_hidden,
BLOCK=block,
)
return topk_p, topk_index, bonus_tokens, hidden_states
@@ -0,0 +1,214 @@
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=10, stage="base-b", runner_config="1-gpu-small")
import unittest
import torch
from sglang.srt.speculative.triton_ops.gather_spec_extras import gather_spec_extras
from sglang.test.test_utils import CustomTestCase
_OUTPUT_NAMES = ("topk_p", "topk_index", "bonus_tokens", "hidden_states")
def _ref_gather(
indices, topk_p_buf, topk_index_buf, output_tokens_buf, hidden_states_buf
):
"""Reference oracle: the exact torch.compile'd advanced-index gather that the
fused Triton kernel replaced (see overlap_utils._gather_spec_extras pre-fusion).
A gather is a pure copy, so the kernel must match this bit-for-bit."""
topk_p = topk_p_buf[indices]
topk_index = topk_index_buf[indices]
bonus_tokens = output_tokens_buf[indices]
hidden_states = (
hidden_states_buf[indices] if hidden_states_buf is not None else None
)
return topk_p, topk_index, bonus_tokens, hidden_states
def _make_buffers(
pool_size,
topk,
hidden_dim,
*,
with_hidden,
hidden_dtype=torch.bfloat16,
device="cuda",
seed=0,
):
"""Build FutureMap-shaped relay buffers.
Mirrors overlap_utils.FutureMap: topk_p / topk_index / hidden_states are
2-D (pool_size, width) while output_tokens is 1-D (pool_size,). The width
mix (incl. the 1-D buffer -> row width 1) exercises the kernel's per-buffer
masking.
"""
g = torch.Generator(device=device).manual_seed(seed)
topk_p_buf = torch.rand(
(pool_size, topk), dtype=torch.float32, device=device, generator=g
)
topk_index_buf = torch.randint(
0, 32000, (pool_size, topk), dtype=torch.int64, device=device, generator=g
)
output_tokens_buf = torch.randint(
0, 32000, (pool_size,), dtype=torch.int64, device=device, generator=g
)
hidden_states_buf = (
torch.randn(
(pool_size, hidden_dim), dtype=hidden_dtype, device=device, generator=g
)
if with_hidden
else None
)
return topk_p_buf, topk_index_buf, output_tokens_buf, hidden_states_buf
@unittest.skipUnless(torch.cuda.is_available(), "CUDA is required for this test.")
class TestGatherSpecExtras(CustomTestCase):
@classmethod
def setUpClass(cls):
super().setUpClass()
cls.device = torch.device("cuda")
def _assert_matches_reference(self, indices, bufs):
"""Run fused kernel + reference on the same inputs and assert every
output is identical (dtype, shape, exact values) and that the source
buffers are never mutated."""
src_snapshots = [None if b is None else b.clone() for b in bufs]
ref = _ref_gather(indices, *bufs)
got = gather_spec_extras(indices, *bufs)
self.assertEqual(len(got), len(ref))
for name, r, o in zip(_OUTPUT_NAMES, ref, got):
if r is None:
self.assertIsNone(o, f"{name} should be None when buffer is None")
continue
self.assertIsNotNone(o, f"{name} unexpectedly None")
self.assertEqual(o.dtype, r.dtype, f"{name} dtype mismatch")
self.assertEqual(tuple(o.shape), tuple(r.shape), f"{name} shape mismatch")
self.assertEqual(o.device.type, r.device.type, f"{name} device mismatch")
# Pure gather == bit-exact copy, so require zero tolerance.
torch.testing.assert_close(
o, r, rtol=0, atol=0, msg=f"{name} value mismatch"
)
# The kernel only reads sources; it must not scribble into them.
for name, before, buf in zip(_OUTPUT_NAMES, src_snapshots, bufs):
if before is None:
continue
torch.testing.assert_close(
buf, before, rtol=0, atol=0, msg=f"source buffer {name} was mutated"
)
def test_matches_reference_across_shapes(self):
# (pool_size, m, topk, hidden_dim). Covers: m<pool, m==pool, topk==1,
# 1-column blocks, exact-1024-width boundary, wide multi-block widths,
# and wide non-power-of-2 widths (partial trailing column block).
configs = [
(16, 8, 1, 7),
(64, 33, 4, 128),
(128, 128, 8, 1024),
(100, 50, 2, 4096),
(257, 200, 16, 4097),
(2048, 777, 8, 5120),
]
for pool_size, m, topk, hidden_dim in configs:
for with_hidden in (True, False):
with self.subTest(
pool_size=pool_size,
m=m,
topk=topk,
hidden_dim=hidden_dim,
with_hidden=with_hidden,
):
bufs = _make_buffers(
pool_size,
topk,
hidden_dim,
with_hidden=with_hidden,
device=self.device,
)
indices = torch.randint(
0, pool_size, (m,), dtype=torch.int64, device=self.device
)
self._assert_matches_reference(indices, bufs)
def test_empty_indices_returns_empty_rows(self):
# m == 0 hits the early-return path; outputs must still carry the right
# trailing dims / dtypes so downstream concatenation stays valid.
for with_hidden in (True, False):
with self.subTest(with_hidden=with_hidden):
bufs = _make_buffers(
32, 4, 256, with_hidden=with_hidden, device=self.device
)
indices = torch.empty(0, dtype=torch.int64, device=self.device)
self._assert_matches_reference(indices, bufs)
def test_non_contiguous_indices(self):
# indices flows from filtered/merged producers and can be strided; the
# kernel addresses it linearly and relies on the internal .contiguous().
pool_size, m = 256, 64
bufs = _make_buffers(pool_size, 8, 512, with_hidden=True, device=self.device)
pairs = torch.randint(
0, pool_size, (m, 2), dtype=torch.int64, device=self.device
)
indices = pairs[:, 0]
self.assertFalse(indices.is_contiguous(), "test setup: indices must be strided")
self._assert_matches_reference(indices, bufs)
def test_duplicate_indices(self):
# Gather (not scatter): repeated source rows are well-defined and must
# each produce an identical copy.
pool_size, m = 8, 64
bufs = _make_buffers(pool_size, 4, 333, with_hidden=True, device=self.device)
indices = torch.randint(
0, 3, (m,), dtype=torch.int64, device=self.device
) # tiny range -> many duplicates
self._assert_matches_reference(indices, bufs)
def test_index_dtype_variants(self):
pool_size, m = 128, 50
bufs = _make_buffers(pool_size, 8, 1024, with_hidden=True, device=self.device)
base = torch.randint(0, pool_size, (m,), device=self.device)
for idx_dtype in (torch.int32, torch.int64):
with self.subTest(idx_dtype=idx_dtype):
self._assert_matches_reference(base.to(idx_dtype), bufs)
def test_hidden_dtype_variants(self):
pool_size, m = 96, 40
indices = torch.randint(
0, pool_size, (m,), dtype=torch.int64, device=self.device
)
for hidden_dtype in (torch.bfloat16, torch.float16, torch.float32):
with self.subTest(hidden_dtype=hidden_dtype):
bufs = _make_buffers(
pool_size,
8,
2048,
with_hidden=True,
hidden_dtype=hidden_dtype,
device=self.device,
)
self._assert_matches_reference(indices, bufs)
def test_outputs_do_not_alias_source_buffers(self):
pool_size, m = 64, 32
bufs = _make_buffers(pool_size, 8, 512, with_hidden=True, device=self.device)
indices = torch.randint(
0, pool_size, (m,), dtype=torch.int64, device=self.device
)
outputs = gather_spec_extras(indices, *bufs)
for name, out, buf in zip(_OUTPUT_NAMES, outputs, bufs):
if out is None or buf is None:
continue
self.assertNotEqual(
out.data_ptr(),
buf.data_ptr(),
f"{name} output aliases its source buffer",
)
if __name__ == "__main__":
unittest.main()