Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
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
136 changes: 136 additions & 0 deletions python/sglang/kernels/ops/memory/small_copy.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,136 @@
import torch
import triton
import triton.language as tl


@triton.jit(
do_not_specialize=["sizes", "cols", "source_strides", "dest_strides"],
do_not_specialize_on_alignment=["sizes", "cols", "source_strides", "dest_strides"],
)
def _small_copy_kernel(
sources,
destinations,
sizes,
cols,
source_strides,
dest_strides,
BIT_WIDTHS: tl.constexpr,
BLOCK: tl.constexpr,
):
field = tl.program_id(0)
offsets = (tl.program_id(1) * BLOCK + tl.arange(0, BLOCK)).to(tl.int64)
for i in tl.static_range(len(BIT_WIDTHS)):
if field == i:
src_offset = (
offsets // cols[i] * source_strides[i][0]
+ offsets % cols[i] * source_strides[i][1]
)
dst_offset = (
offsets // cols[i] * dest_strides[i][0]
+ offsets % cols[i] * dest_strides[i][1]
)
src = sources[i]
dst = destinations[i]
if BIT_WIDTHS[i] == 8:
src = src.to(tl.pointer_type(tl.uint8))
dst = dst.to(tl.pointer_type(tl.uint8))
elif BIT_WIDTHS[i] == 16:
src = src.to(tl.pointer_type(tl.uint16))
dst = dst.to(tl.pointer_type(tl.uint16))
elif BIT_WIDTHS[i] == 32:
src = src.to(tl.pointer_type(tl.uint32))
dst = dst.to(tl.pointer_type(tl.uint32))
elif BIT_WIDTHS[i] == 64:
src = src.to(tl.pointer_type(tl.uint64))
dst = dst.to(tl.pointer_type(tl.uint64))
values = tl.load(src + src_offset, offsets < sizes[i], other=0)
tl.store(dst + dst_offset, values, offsets < sizes[i])


def try_small_copy(dsts, srcs):
if len(dsts) < 2 or len(dsts) != len(srcs):
return False
supported = (
torch.bool,
torch.uint8,
torch.int8,
torch.int16,
torch.int32,
torch.int64,
torch.float16,
torch.bfloat16,
torch.float32,
torch.float64,
)
device = dsts[0].device
if device.type != "cuda":
return False
sizes, cols, src_strides, dst_strides, bit_widths = [], [], [], [], []
src_ranges, dst_ranges = [], []
for dst, src in zip(dsts, srcs):
if (
dst.dtype not in supported
or src.dtype not in supported
or dst.device != device
or src.device != device
or dst.shape != src.shape
or dst.ndim not in (1, 2)
or dst.numel() > 8192
or dst.requires_grad
or src.requires_grad
or src.is_neg()
or dst.is_neg()
):
return False
same_dtype = dst.dtype == src.dtype
if not same_dtype and not (
dst.dtype in (torch.int32, torch.int64)
and src.dtype in (torch.int32, torch.int64)
):
return False
n = dst.numel()
c = max(dst.shape[-1], 1)
ds = dst.stride() if dst.ndim == 2 else (0, dst.stride(0))
ss = src.stride() if src.ndim == 2 else (0, src.stride(0))
if ds[1] <= 0 or (dst.ndim == 2 and ds[0] < c * ds[1]):
return False
sizes.append(n)
cols.append(c)
src_strides.append(ss)
dst_strides.append(ds)
bit_widths.append(dst.element_size() * 8 if same_dtype else 0)
for tensor, ranges in ((src, src_ranges), (dst, dst_ranges)):
first = tensor.data_ptr()
extent = (
(sum((d - 1) * s for d, s in zip(tensor.shape, tensor.stride())) + 1)
* tensor.element_size()
if n
else 0
)
ranges.append((first, first + extent))
for i, (first, last) in enumerate(dst_ranges):
for j, (other_first, other_last) in enumerate(src_ranges):
if first < other_last and other_first < last:
if (
i != j
or dsts[i].data_ptr() != srcs[i].data_ptr()
or dsts[i].stride() != srcs[i].stride()
or dsts[i].dtype != srcs[i].dtype
):
return False
for other_first, other_last in dst_ranges[:i]:
if first < other_last and other_first < last:
return False
largest = max(sizes)
if largest:
_small_copy_kernel[(len(dsts), triton.cdiv(largest, 256))](
tuple(srcs),
tuple(dsts),
tuple(sizes),
tuple(cols),
tuple(src_strides),
tuple(dst_strides),
tuple(bit_widths),
BLOCK=256,
)
return True
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,12 @@ def _foreach_copy(
for dst, src in zip(group_dsts, group_srcs):
dst.copy_(src)

if dsts and dsts[0].is_cuda:
from sglang.kernels.ops.memory.small_copy import try_small_copy

if try_small_copy(dsts, srcs):
return

groups: Dict[Tuple[torch.dtype, torch.dtype], Tuple[List, List]] = {}
for dst, src in zip(dsts, srcs):
key = (dst.dtype, src.dtype)
Expand Down
6 changes: 6 additions & 0 deletions python/sglang/srt/model_executor/runner_utils/buffers.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,12 @@ def foreach_copy(dsts: List[torch.Tensor], srcs: List[torch.Tensor]) -> None:
for dst, src in zip(dsts, srcs):
dst.copy_(src)

if dsts and dsts[0].is_cuda:
from sglang.kernels.ops.memory.small_copy import try_small_copy

if try_small_copy(dsts, srcs):
return

groups: Dict[Tuple[torch.dtype, torch.dtype], Tuple[List, List]] = {}
for dst, src in zip(dsts, srcs):
key = (dst.dtype, src.dtype)
Expand Down
146 changes: 146 additions & 0 deletions test/registered/kernels/ops/memory/test_small_copy.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,146 @@
import itertools
import unittest

import torch

from sglang.kernels.ops.memory.small_copy import _small_copy_kernel, try_small_copy
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase

register_cuda_ci(est_time=10, stage="base-b-kernel-unit", runner_config="1-gpu-large")

_DTYPES = (
torch.bool,
torch.uint8,
torch.int8,
torch.int16,
torch.int32,
torch.int64,
torch.float16,
torch.bfloat16,
torch.float32,
torch.float64,
)


def _as_bytes(tensor):
return (
torch.empty(tensor.shape, device=tensor.device, dtype=tensor.dtype)
.copy_(tensor)
.view(torch.uint8)
)


@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA")
class TestSmallCopy(CustomTestCase):
def test_bits(self):
for n, strided in itertools.product((0, 1, 8, 128, 4096), (False, True)):
with self.subTest(n=n, strided=strided):
dsts = []
srcs = []
backings = []
for dtype in _DTYPES:
source = torch.randint(
0,
256,
(2 * n * torch.empty((), dtype=dtype).element_size(),),
device="cuda",
dtype=torch.uint8,
).view(dtype)
src = source[::2] if strided else source[:n]
backing = torch.zeros(n * 3, device="cuda", dtype=dtype)
dst = backing[::3] if strided else backing[:n]
dsts.append(dst)
srcs.append(src)
backings.append(backing)
expected = [s.clone() for s in srcs]
self.assertTrue(try_small_copy(dsts, srcs))
for d, s in zip(dsts, expected):
self.assertTrue(torch.equal(_as_bytes(d), _as_bytes(s)))

def test_cast_and_mrope_graph(self):
for n in (1, 4, 128, 2048):
with self.subTest(n=n):
s0 = torch.randint(
-(2**40), 2**40, (n * 2,), device="cuda", dtype=torch.int64
)[::2]
s1 = torch.arange(n * 3, device="cuda", dtype=torch.int32).reshape(3, n)
d0 = torch.empty(n, device="cuda", dtype=torch.int32)
backing = torch.full((3, n + 5), -7, device="cuda", dtype=torch.int64)
d1 = backing[:, :n]
try_small_copy([d0, d1], [s0, s1])
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
self.assertTrue(try_small_copy([d0, d1], [s0, s1]))
for _ in range(3):
s0.add_(17)
s1.add_(13)
graph.replay()
self.assertTrue(torch.equal(d0, s0.int()))
self.assertTrue(torch.equal(d1, s1.long()))
self.assertTrue(bool((backing[:, n:] == -7).all()))

def test_fallbacks(self):
x = torch.arange(17, device="cuda")
y = torch.empty_like(x)
self.assertFalse(try_small_copy([x[1:], y[:16]], [x[:-1], x[:16]]))
self.assertFalse(try_small_copy([x, y], [y, x]))
self.assertFalse(try_small_copy([x[:8], x[4:12]], [y[:8], y[:8]]))
self.assertFalse(
try_small_copy(
[torch.empty(8193, device="cuda")] * 2,
[torch.empty(8193, device="cuda")] * 2,
)
)
self.assertFalse(
try_small_copy(
[torch.empty(2, dtype=torch.complex64, device="cuda")] * 2,
[torch.empty(2, dtype=torch.complex64, device="cuda")] * 2,
)
)
a = torch.empty(2)
self.assertFalse(try_small_copy([a, a], [a, a]))
self.assertTrue(try_small_copy([x, y], [x, y]))

def test_broadcast_and_negative_views(self):
source = torch.arange(7, device="cuda", dtype=torch.float32)
expanded = source[None, :].expand(3, 7)
destinations = [torch.empty_like(expanded), torch.empty_like(source)]
self.assertTrue(try_small_copy(destinations, [expanded, source]))
self.assertTrue(torch.equal(destinations[0], expanded))
self.assertTrue(torch.equal(destinations[1], source))
self.assertFalse(try_small_copy(destinations, [expanded, source._neg_view()]))

def test_dynamic_metadata_has_bounded_specializations(self):
kernel_cache = _small_copy_kernel.device_caches[torch.cuda.current_device()][0]
initial_count = len(kernel_cache)
specialization_count = None
for n in [*range(1, 130), 255, 256, 511, 512, 1023, 1024, 2048, 2730]:
step = 1 + n % 3
source = torch.arange(n * step, device="cuda", dtype=torch.int64)[::step]
backing = torch.full((n * step,), -7, device="cuda", dtype=torch.int32)
destination = backing[::step]
matrix = torch.arange(
3 * (n + 3), device="cuda", dtype=torch.int32
).reshape(3, n + 3)[:, :n]
matrix_backing = torch.full(
(3, n + 5), -7, device="cuda", dtype=torch.int64
)
matrix_destination = matrix_backing[:, :n]
self.assertTrue(
try_small_copy([destination, matrix_destination], [source, matrix])
)
self.assertTrue(torch.equal(destination, source.int()))
self.assertTrue(torch.equal(matrix_destination, matrix.long()))
self.assertTrue(bool((matrix_backing[:, n:] == -7).all()))
for offset in range(1, step):
self.assertTrue(bool((backing[offset::step] == -7).all()))
if n == 48:
specialization_count = len(kernel_cache)
self.assertLessEqual(specialization_count - initial_count, 16)
elif n > 48:
self.assertEqual(len(kernel_cache), specialization_count)


if __name__ == "__main__":
unittest.main()
Loading