Optimize large add_constant tensors (#24755)
Co-authored-by: Codex <codex@example.com> Co-authored-by: BBuf <xiaoyu.zhang@radixark.net>
This commit is contained in:
co-authored by
Codex
BBuf
parent
b421e60eed
commit
e279b0bf72
@@ -0,0 +1,59 @@
|
|||||||
|
import torch
|
||||||
|
import triton
|
||||||
|
import triton.testing
|
||||||
|
|
||||||
|
from sglang.jit_kernel.add_constant import _jit_add_constant_module, add_constant
|
||||||
|
from sglang.jit_kernel.benchmark.utils import (
|
||||||
|
DEFAULT_DEVICE,
|
||||||
|
get_benchmark_range,
|
||||||
|
run_benchmark_no_cudagraph,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=15, suite="base-b-kernel-benchmark-1-gpu-large")
|
||||||
|
|
||||||
|
CONSTANT = 7
|
||||||
|
SIZE_LIST = get_benchmark_range(
|
||||||
|
full_range=[128, 1024, 1025, 4096, 4097, 65536, 2**20, 2**22, 2**24],
|
||||||
|
ci_range=[4096, 2**20],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.testing.perf_report(
|
||||||
|
triton.testing.Benchmark(
|
||||||
|
x_names=["size"],
|
||||||
|
x_vals=SIZE_LIST,
|
||||||
|
line_arg="provider",
|
||||||
|
line_vals=["jit_module", "jit_wrapper", "torch"],
|
||||||
|
line_names=["JIT module", "JIT wrapper", "PyTorch"],
|
||||||
|
styles=[("blue", "-"), ("orange", "-"), ("green", "--")],
|
||||||
|
ylabel="us",
|
||||||
|
plot_name="add-constant-performance",
|
||||||
|
args={},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
def benchmark(size: int, provider: str):
|
||||||
|
src = torch.arange(size, dtype=torch.int32, device=DEFAULT_DEVICE)
|
||||||
|
|
||||||
|
if provider == "jit_module":
|
||||||
|
dst = torch.empty_like(src)
|
||||||
|
module = _jit_add_constant_module(CONSTANT)
|
||||||
|
|
||||||
|
def fn():
|
||||||
|
module.add_constant(dst, src)
|
||||||
|
|
||||||
|
elif provider == "jit_wrapper":
|
||||||
|
|
||||||
|
def fn():
|
||||||
|
add_constant(src, CONSTANT)
|
||||||
|
|
||||||
|
else:
|
||||||
|
|
||||||
|
def fn():
|
||||||
|
src + CONSTANT
|
||||||
|
|
||||||
|
return run_benchmark_no_cudagraph(fn)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
benchmark.run(print_data=True)
|
||||||
@@ -2,6 +2,7 @@
|
|||||||
#include <sgl_kernel/utils.h> // For div_ceil, RuntimeCheck
|
#include <sgl_kernel/utils.h> // For div_ceil, RuntimeCheck
|
||||||
|
|
||||||
#include <sgl_kernel/utils.cuh> // For LaunchKernel
|
#include <sgl_kernel/utils.cuh> // For LaunchKernel
|
||||||
|
#include <sgl_kernel/vec.cuh>
|
||||||
|
|
||||||
#include <dlpack/dlpack.h>
|
#include <dlpack/dlpack.h>
|
||||||
#include <tvm/ffi/container/tensor.h>
|
#include <tvm/ffi/container/tensor.h>
|
||||||
@@ -11,6 +12,17 @@
|
|||||||
|
|
||||||
namespace {
|
namespace {
|
||||||
|
|
||||||
|
constexpr size_t kBlockSize = 256;
|
||||||
|
constexpr size_t kVectorizedMinElements = 1 << 20;
|
||||||
|
constexpr size_t kVectorBytes = device::kMaxVecBytes;
|
||||||
|
static_assert(kVectorBytes % sizeof(int32_t) == 0, "Vector byte width must contain whole int32_t elements");
|
||||||
|
constexpr size_t kElementsPerVector = kVectorBytes / sizeof(int32_t);
|
||||||
|
|
||||||
|
template <typename Vector>
|
||||||
|
bool is_aligned_for_vector(const int32_t* ptr) {
|
||||||
|
return reinterpret_cast<uintptr_t>(ptr) % alignof(Vector) == 0;
|
||||||
|
}
|
||||||
|
|
||||||
template <int32_t kConstant>
|
template <int32_t kConstant>
|
||||||
__global__ void add_constant_kernel(int32_t* dst, const int32_t* src, size_t length) {
|
__global__ void add_constant_kernel(int32_t* dst, const int32_t* src, size_t length) {
|
||||||
size_t idx = blockIdx.x * blockDim.x + threadIdx.x;
|
size_t idx = blockIdx.x * blockDim.x + threadIdx.x;
|
||||||
@@ -19,7 +31,28 @@ __global__ void add_constant_kernel(int32_t* dst, const int32_t* src, size_t len
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
constexpr size_t kBlockSize = 256;
|
template <int32_t kConstant, size_t kElementsPerVector>
|
||||||
|
__global__ void add_constant_vectorized_kernel(int32_t* dst, const int32_t* src, size_t length) {
|
||||||
|
using Vector = device::AlignedVector<int32_t, kElementsPerVector>;
|
||||||
|
|
||||||
|
const size_t work_idx = blockIdx.x * blockDim.x + threadIdx.x;
|
||||||
|
const size_t vector_count = length / kElementsPerVector;
|
||||||
|
const size_t tail_start = vector_count * kElementsPerVector;
|
||||||
|
|
||||||
|
if (work_idx < vector_count) {
|
||||||
|
auto values = device::load_as<Vector>(src, work_idx);
|
||||||
|
#pragma unroll
|
||||||
|
for (size_t i = 0; i < kElementsPerVector; ++i) {
|
||||||
|
values[i] += kConstant;
|
||||||
|
}
|
||||||
|
device::store_as<Vector>(dst, values, work_idx);
|
||||||
|
} else {
|
||||||
|
const size_t tail_idx = tail_start + work_idx - vector_count;
|
||||||
|
if (tail_idx < length) {
|
||||||
|
dst[tail_idx] = src[tail_idx] + kConstant;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// You can also use struct with static method as an alternative
|
// You can also use struct with static method as an alternative
|
||||||
template <int32_t kConstant>
|
template <int32_t kConstant>
|
||||||
@@ -37,7 +70,6 @@ void add_constant(tvm::ffi::TensorView dst, tvm::ffi::TensorView src) {
|
|||||||
|
|
||||||
// 2. Extract required parameters, prepare for kernel launch
|
// 2. Extract required parameters, prepare for kernel launch
|
||||||
const size_t num_elements = N.unwrap();
|
const size_t num_elements = N.unwrap();
|
||||||
const size_t grid_size = div_ceil(num_elements, kBlockSize);
|
|
||||||
const DLDevice device = device_.unwrap();
|
const DLDevice device = device_.unwrap();
|
||||||
[[maybe_unused]] // optional, can be omitted
|
[[maybe_unused]] // optional, can be omitted
|
||||||
const size_t dynamic_smem = 0;
|
const size_t dynamic_smem = 0;
|
||||||
@@ -46,14 +78,24 @@ void add_constant(tvm::ffi::TensorView dst, tvm::ffi::TensorView src) {
|
|||||||
// some extra runtime checks using host::RuntimeCheck
|
// some extra runtime checks using host::RuntimeCheck
|
||||||
RuntimeCheck(num_elements > 0, "We only support non-empty tensors, got num_elements = ", num_elements);
|
RuntimeCheck(num_elements > 0, "We only support non-empty tensors, got num_elements = ", num_elements);
|
||||||
|
|
||||||
|
const auto* src_ptr = static_cast<const int32_t*>(src.data_ptr());
|
||||||
|
auto* dst_ptr = static_cast<int32_t*>(dst.data_ptr());
|
||||||
|
using Vector = device::AlignedVector<int32_t, kElementsPerVector>;
|
||||||
|
const bool is_vector_aligned = is_aligned_for_vector<Vector>(src_ptr) && is_aligned_for_vector<Vector>(dst_ptr);
|
||||||
|
|
||||||
// 3. Launch the kernel. Error code will be automatically checked.
|
// 3. Launch the kernel. Error code will be automatically checked.
|
||||||
LaunchKernel(grid_size, kBlockSize, device /*, dynamic_smem*/)(
|
if (num_elements >= kVectorizedMinElements && is_vector_aligned) {
|
||||||
// kernel function
|
const size_t vector_count = num_elements / kElementsPerVector;
|
||||||
add_constant_kernel<kConstant>,
|
const size_t tail_count = num_elements - vector_count * kElementsPerVector;
|
||||||
// kernel arguments
|
const size_t work_items = vector_count + tail_count;
|
||||||
static_cast<int32_t*>(dst.data_ptr()),
|
const size_t grid_size = div_ceil(work_items, kBlockSize);
|
||||||
static_cast<int32_t*>(src.data_ptr()),
|
LaunchKernel(grid_size, kBlockSize, device /*, dynamic_smem*/)(
|
||||||
num_elements);
|
add_constant_vectorized_kernel<kConstant, kElementsPerVector>, dst_ptr, src_ptr, num_elements);
|
||||||
|
} else {
|
||||||
|
const size_t grid_size = div_ceil(num_elements, kBlockSize);
|
||||||
|
LaunchKernel(grid_size, kBlockSize, device /*, dynamic_smem*/)(
|
||||||
|
add_constant_kernel<kConstant>, dst_ptr, src_ptr, num_elements);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
} // namespace
|
} // namespace
|
||||||
|
|||||||
@@ -10,7 +10,7 @@ register_cuda_ci(est_time=45, suite="base-b-kernel-unit-1-gpu-large")
|
|||||||
register_cuda_ci(est_time=180, suite="nightly-kernel-1-gpu", nightly=True)
|
register_cuda_ci(est_time=180, suite="nightly-kernel-1-gpu", nightly=True)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("size", [1, 2, 127, 128, 1024, 1025])
|
@pytest.mark.parametrize("size", [1, 2, 127, 128, 1024, 1025, 4096, 4097])
|
||||||
@pytest.mark.parametrize("constant", [0, 1, 7, 1024, -3])
|
@pytest.mark.parametrize("constant", [0, 1, 7, 1024, -3])
|
||||||
def test_add_constant(size: int, constant: int) -> None:
|
def test_add_constant(size: int, constant: int) -> None:
|
||||||
src = torch.arange(0, size, dtype=torch.int32, device="cuda")
|
src = torch.arange(0, size, dtype=torch.int32, device="cuda")
|
||||||
@@ -18,5 +18,24 @@ def test_add_constant(size: int, constant: int) -> None:
|
|||||||
assert torch.all(dst == src + constant)
|
assert torch.all(dst == src + constant)
|
||||||
|
|
||||||
|
|
||||||
|
def test_add_constant_unaligned_input() -> None:
|
||||||
|
src = torch.arange(0, 4098, dtype=torch.int32, device="cuda")[1:]
|
||||||
|
dst = add_constant(src, 7)
|
||||||
|
assert torch.all(dst == src + 7)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("size", [2**20, 2**20 + 3])
|
||||||
|
def test_add_constant_large_aligned_input(size: int) -> None:
|
||||||
|
src = torch.arange(0, size, dtype=torch.int32, device="cuda")
|
||||||
|
dst = add_constant(src, -3)
|
||||||
|
assert torch.all(dst == src - 3)
|
||||||
|
|
||||||
|
|
||||||
|
def test_add_constant_large_unaligned_input() -> None:
|
||||||
|
src = torch.arange(0, 2**20 + 4, dtype=torch.int32, device="cuda")[1:]
|
||||||
|
dst = add_constant(src, 7)
|
||||||
|
assert torch.all(dst == src + 7)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
sys.exit(pytest.main([__file__, "-v", "-s"]))
|
sys.exit(pytest.main([__file__, "-v", "-s"]))
|
||||||
|
|||||||
Reference in New Issue
Block a user