[AMD] fix(jit): port kv_canary write/verify/plan kernels to ROCm (#28357)
This commit is contained in:
@@ -10,7 +10,6 @@
|
|||||||
#include <tvm/ffi/container/tensor.h>
|
#include <tvm/ffi/container/tensor.h>
|
||||||
|
|
||||||
#include <cstdint>
|
#include <cstdint>
|
||||||
#include <cuda_runtime.h>
|
|
||||||
|
|
||||||
namespace {
|
namespace {
|
||||||
|
|
||||||
@@ -96,7 +95,11 @@ __global__ void plan_entries_persistent_kernel(
|
|||||||
static_cast<long long>(total_verify),
|
static_cast<long long>(total_verify),
|
||||||
static_cast<long long>(params.verify_capacity));
|
static_cast<long long>(params.verify_capacity));
|
||||||
}
|
}
|
||||||
|
#ifndef USE_ROCM
|
||||||
__trap();
|
__trap();
|
||||||
|
#else
|
||||||
|
__builtin_trap(); // HIP/clang device-side trap; __trap() is CUDA-only.
|
||||||
|
#endif
|
||||||
}
|
}
|
||||||
|
|
||||||
const int64_t tid_start = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
|
const int64_t tid_start = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
|
||||||
@@ -195,31 +198,31 @@ struct PlanEntriesKernel {
|
|||||||
SymbolicSize Npool_rows = {"req_to_verify_expected_tokens_rows"};
|
SymbolicSize Npool_rows = {"req_to_verify_expected_tokens_rows"};
|
||||||
SymbolicSize Npool_cols = {"req_to_verify_expected_tokens_cols"};
|
SymbolicSize Npool_cols = {"req_to_verify_expected_tokens_cols"};
|
||||||
SymbolicDevice device_;
|
SymbolicDevice device_;
|
||||||
device_.set_options<kDLCUDA>();
|
device_.set_options<kDLGPU>();
|
||||||
|
|
||||||
TensorMatcher({Nbs}) //
|
TensorMatcher({Nbs}) //
|
||||||
.with_dtype<int64_t>()
|
.with_dtype<int64_t>()
|
||||||
.with_device<kDLCUDA>(device_)
|
.with_device<kDLGPU>(device_)
|
||||||
.verify(req_pool_indices)
|
.verify(req_pool_indices)
|
||||||
.verify(prefix_lens);
|
.verify(prefix_lens);
|
||||||
TensorMatcher({Nscratch}) //
|
TensorMatcher({Nscratch}) //
|
||||||
.with_dtype<int64_t>()
|
.with_dtype<int64_t>()
|
||||||
.with_device<kDLCUDA>(device_)
|
.with_device<kDLGPU>(device_)
|
||||||
.verify(verify_offsets_scratch);
|
.verify(verify_offsets_scratch);
|
||||||
TensorMatcher({1}) //
|
TensorMatcher({1}) //
|
||||||
.with_dtype<int32_t>()
|
.with_dtype<int32_t>()
|
||||||
.with_device<kDLCUDA>(device_)
|
.with_device<kDLGPU>(device_)
|
||||||
.verify(verify_enable);
|
.verify(verify_enable);
|
||||||
TensorMatcher({Ncap}) //
|
TensorMatcher({Ncap}) //
|
||||||
.with_dtype<int64_t>()
|
.with_dtype<int64_t>()
|
||||||
.with_device<kDLCUDA>(device_)
|
.with_device<kDLGPU>(device_)
|
||||||
.verify(out_verify_slot_indices)
|
.verify(out_verify_slot_indices)
|
||||||
.verify(out_verify_expected_tokens)
|
.verify(out_verify_expected_tokens)
|
||||||
.verify(out_verify_expected_positions)
|
.verify(out_verify_expected_positions)
|
||||||
.verify(out_verify_prev_slot_indices);
|
.verify(out_verify_prev_slot_indices);
|
||||||
TensorMatcher({Nmax_reqs, Nmax_seq_len}) //
|
TensorMatcher({Nmax_reqs, Nmax_seq_len}) //
|
||||||
.with_dtype<int32_t>()
|
.with_dtype<int32_t>()
|
||||||
.with_device<kDLCUDA>(device_)
|
.with_device<kDLGPU>(device_)
|
||||||
.verify(req_to_token);
|
.verify(req_to_token);
|
||||||
RuntimeCheck(
|
RuntimeCheck(
|
||||||
full_to_swa_index_mapping.has_value() == HAS_SWA_LUT,
|
full_to_swa_index_mapping.has_value() == HAS_SWA_LUT,
|
||||||
@@ -235,17 +238,17 @@ struct PlanEntriesKernel {
|
|||||||
if constexpr (HAS_SWA_LUT) {
|
if constexpr (HAS_SWA_LUT) {
|
||||||
TensorMatcher({Nlut}) //
|
TensorMatcher({Nlut}) //
|
||||||
.with_dtype<int64_t>()
|
.with_dtype<int64_t>()
|
||||||
.with_device<kDLCUDA>(device_)
|
.with_device<kDLGPU>(device_)
|
||||||
.verify(full_to_swa_index_mapping.value());
|
.verify(full_to_swa_index_mapping.value());
|
||||||
}
|
}
|
||||||
if constexpr (HAS_VERIFY_EXPECTED_TOKEN_POOL) {
|
if constexpr (HAS_VERIFY_EXPECTED_TOKEN_POOL) {
|
||||||
TensorMatcher({Npool_rows, Npool_cols}) //
|
TensorMatcher({Npool_rows, Npool_cols}) //
|
||||||
.with_dtype<int32_t>()
|
.with_dtype<int32_t>()
|
||||||
.with_device<kDLCUDA>(device_)
|
.with_device<kDLGPU>(device_)
|
||||||
.verify(req_to_verify_expected_tokens.value());
|
.verify(req_to_verify_expected_tokens.value());
|
||||||
TensorMatcher({Nbs}) //
|
TensorMatcher({Nbs}) //
|
||||||
.with_dtype<int64_t>()
|
.with_dtype<int64_t>()
|
||||||
.with_device<kDLCUDA>(device_)
|
.with_device<kDLGPU>(device_)
|
||||||
.verify(req_to_verify_expected_tokens_valid_lens.value());
|
.verify(req_to_verify_expected_tokens_valid_lens.value());
|
||||||
}
|
}
|
||||||
RuntimeCheck(Nscratch.unwrap() >= Nbs.unwrap() + 1, "verify_offsets_scratch length must be >= bs_padded + 1");
|
RuntimeCheck(Nscratch.unwrap() >= Nbs.unwrap() + 1, "verify_offsets_scratch length must be >= bs_padded + 1");
|
||||||
|
|||||||
@@ -4,6 +4,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, SGL_DEVICE
|
#include <sgl_kernel/utils.cuh> // For LaunchKernel, SGL_DEVICE
|
||||||
|
#include <sgl_kernel/warp.cuh> // For device::warp::reduce_sum
|
||||||
|
|
||||||
#include <dlpack/dlpack.h>
|
#include <dlpack/dlpack.h>
|
||||||
#include <tvm/ffi/container/tensor.h>
|
#include <tvm/ffi/container/tensor.h>
|
||||||
@@ -124,11 +125,13 @@ __global__ void canary_verify_kernel(const VerifyKernelParams __grid_constant__
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
uint32_t warp_active_count = local_active_count;
|
// Warp-size-agnostic reduction: device::warp::reduce_sum performs a width-kWarpThreads (32) xor
|
||||||
for (int offset = 16; offset > 0; offset >>= 1) {
|
// all-reduce, so each 32-lane sub-group gets its own total with its own lane-0 leader. Correct on
|
||||||
warp_active_count += __shfl_down_sync(0xFFFFFFFFu, warp_active_count, offset);
|
// both CUDA (warp=32) and ROCm wave64 (two 32-lane sub-groups per wavefront); a hand-rolled
|
||||||
}
|
// __shfl_down over the full wavefront would cross the 32-lane boundary on wave64.
|
||||||
if ((threadIdx.x & 31u) == 0u && warp_active_count != 0u) {
|
const uint32_t warp_active_count = device::warp::reduce_sum<device::kWarpThreads>(local_active_count);
|
||||||
|
const bool is_subgroup_leader = (threadIdx.x % device::kWarpThreads) == 0u;
|
||||||
|
if (is_subgroup_leader && warp_active_count != 0u) {
|
||||||
atomicAdd(
|
atomicAdd(
|
||||||
reinterpret_cast<unsigned long long*>(p.slot_run_counter), static_cast<unsigned long long>(warp_active_count));
|
reinterpret_cast<unsigned long long*>(p.slot_run_counter), static_cast<unsigned long long>(warp_active_count));
|
||||||
}
|
}
|
||||||
@@ -173,29 +176,29 @@ struct CanaryVerifyKernel {
|
|||||||
SymbolicSize N_stride = {"slot_stride_bytes"};
|
SymbolicSize N_stride = {"slot_stride_bytes"};
|
||||||
SymbolicSize N_verify = {"verify_capacity"};
|
SymbolicSize N_verify = {"verify_capacity"};
|
||||||
SymbolicDevice device_;
|
SymbolicDevice device_;
|
||||||
device_.set_options<kDLCUDA>();
|
device_.set_options<kDLGPU>();
|
||||||
|
|
||||||
TensorMatcher({N_slots, N_stride}).with_dtype<uint8_t>().with_device<kDLCUDA>(device_).verify(canary_buf);
|
TensorMatcher({N_slots, N_stride}).with_dtype<uint8_t>().with_device<kDLGPU>(device_).verify(canary_buf);
|
||||||
|
|
||||||
TensorMatcher({N_verify})
|
TensorMatcher({N_verify})
|
||||||
.with_dtype<int64_t>()
|
.with_dtype<int64_t>()
|
||||||
.with_device<kDLCUDA>(device_)
|
.with_device<kDLGPU>(device_)
|
||||||
.verify(verify_slot_indices)
|
.verify(verify_slot_indices)
|
||||||
.verify(verify_expected_tokens)
|
.verify(verify_expected_tokens)
|
||||||
.verify(verify_expected_positions)
|
.verify(verify_expected_positions)
|
||||||
.verify(verify_prev_slot_indices);
|
.verify(verify_prev_slot_indices);
|
||||||
TensorMatcher({1}).with_dtype<int32_t>().with_device<kDLCUDA>(device_).verify(verify_num_valid);
|
TensorMatcher({1}).with_dtype<int32_t>().with_device<kDLGPU>(device_).verify(verify_num_valid);
|
||||||
TensorMatcher({1}).with_dtype<int32_t>().with_device<kDLCUDA>(device_).verify(verify_enable);
|
TensorMatcher({1}).with_dtype<int32_t>().with_device<kDLGPU>(device_).verify(verify_enable);
|
||||||
|
|
||||||
TensorMatcher({1}).with_dtype<int32_t>().with_device<kDLCUDA>(device_).verify(violation_write_index);
|
TensorMatcher({1}).with_dtype<int32_t>().with_device<kDLGPU>(device_).verify(violation_write_index);
|
||||||
SymbolicSize N_ring = {"ring_capacity"};
|
SymbolicSize N_ring = {"ring_capacity"};
|
||||||
TensorMatcher({N_ring, static_cast<int64_t>(kViolationFields)})
|
TensorMatcher({N_ring, static_cast<int64_t>(kViolationFields)})
|
||||||
.with_dtype<int64_t>()
|
.with_dtype<int64_t>()
|
||||||
.with_device<kDLCUDA>(device_)
|
.with_device<kDLGPU>(device_)
|
||||||
.verify(violation_ring);
|
.verify(violation_ring);
|
||||||
TensorMatcher({1})
|
TensorMatcher({1})
|
||||||
.with_dtype<int64_t>()
|
.with_dtype<int64_t>()
|
||||||
.with_device<kDLCUDA>(device_)
|
.with_device<kDLGPU>(device_)
|
||||||
.verify(slot_run_counter)
|
.verify(slot_run_counter)
|
||||||
.verify(kernel_run_counter);
|
.verify(kernel_run_counter);
|
||||||
|
|
||||||
@@ -205,25 +208,25 @@ struct CanaryVerifyKernel {
|
|||||||
SymbolicSize N_real_kv_cols_0 = {"real_kv_cols_0"};
|
SymbolicSize N_real_kv_cols_0 = {"real_kv_cols_0"};
|
||||||
TensorMatcher({N_real_kv_rows_0, N_real_kv_cols_0})
|
TensorMatcher({N_real_kv_rows_0, N_real_kv_cols_0})
|
||||||
.with_dtype<uint8_t>()
|
.with_dtype<uint8_t>()
|
||||||
.with_device<kDLCUDA>(device_)
|
.with_device<kDLGPU>(device_)
|
||||||
.verify(real_kv_buf_0);
|
.verify(real_kv_buf_0);
|
||||||
SymbolicSize N_real_kv_rows_1 = {"real_kv_rows_1"};
|
SymbolicSize N_real_kv_rows_1 = {"real_kv_rows_1"};
|
||||||
SymbolicSize N_real_kv_cols_1 = {"real_kv_cols_1"};
|
SymbolicSize N_real_kv_cols_1 = {"real_kv_cols_1"};
|
||||||
TensorMatcher({N_real_kv_rows_1, N_real_kv_cols_1})
|
TensorMatcher({N_real_kv_rows_1, N_real_kv_cols_1})
|
||||||
.with_dtype<uint8_t>()
|
.with_dtype<uint8_t>()
|
||||||
.with_device<kDLCUDA>(device_)
|
.with_device<kDLGPU>(device_)
|
||||||
.verify(real_kv_buf_1);
|
.verify(real_kv_buf_1);
|
||||||
SymbolicSize N_real_kv_rows_2 = {"real_kv_rows_2"};
|
SymbolicSize N_real_kv_rows_2 = {"real_kv_rows_2"};
|
||||||
SymbolicSize N_real_kv_cols_2 = {"real_kv_cols_2"};
|
SymbolicSize N_real_kv_cols_2 = {"real_kv_cols_2"};
|
||||||
TensorMatcher({N_real_kv_rows_2, N_real_kv_cols_2})
|
TensorMatcher({N_real_kv_rows_2, N_real_kv_cols_2})
|
||||||
.with_dtype<uint8_t>()
|
.with_dtype<uint8_t>()
|
||||||
.with_device<kDLCUDA>(device_)
|
.with_device<kDLGPU>(device_)
|
||||||
.verify(real_kv_buf_2);
|
.verify(real_kv_buf_2);
|
||||||
SymbolicSize N_real_kv_rows_3 = {"real_kv_rows_3"};
|
SymbolicSize N_real_kv_rows_3 = {"real_kv_rows_3"};
|
||||||
SymbolicSize N_real_kv_cols_3 = {"real_kv_cols_3"};
|
SymbolicSize N_real_kv_cols_3 = {"real_kv_cols_3"};
|
||||||
TensorMatcher({N_real_kv_rows_3, N_real_kv_cols_3})
|
TensorMatcher({N_real_kv_rows_3, N_real_kv_cols_3})
|
||||||
.with_dtype<uint8_t>()
|
.with_dtype<uint8_t>()
|
||||||
.with_device<kDLCUDA>(device_)
|
.with_device<kDLGPU>(device_)
|
||||||
.verify(real_kv_buf_3);
|
.verify(real_kv_buf_3);
|
||||||
TensorMatcher({static_cast<int64_t>(kMaxRealKvSources), static_cast<int64_t>(kRealKvSourceFieldsPerEntry)})
|
TensorMatcher({static_cast<int64_t>(kMaxRealKvSources), static_cast<int64_t>(kRealKvSourceFieldsPerEntry)})
|
||||||
.with_dtype<int32_t>()
|
.with_dtype<int32_t>()
|
||||||
|
|||||||
@@ -205,20 +205,20 @@ inline void canary_write_step_cuda(
|
|||||||
SymbolicSize N_write_reqs = {"write_req_capacity"};
|
SymbolicSize N_write_reqs = {"write_req_capacity"};
|
||||||
SymbolicSize N_tokens = {"num_tokens_padded"};
|
SymbolicSize N_tokens = {"num_tokens_padded"};
|
||||||
SymbolicDevice device_;
|
SymbolicDevice device_;
|
||||||
device_.set_options<kDLCUDA>();
|
device_.set_options<kDLGPU>();
|
||||||
|
|
||||||
TensorMatcher({N_slots, N_stride}).with_dtype<uint8_t>().with_device<kDLCUDA>(device_).verify(canary_buf);
|
TensorMatcher({N_slots, N_stride}).with_dtype<uint8_t>().with_device<kDLGPU>(device_).verify(canary_buf);
|
||||||
|
|
||||||
// write_offsets has shape [write_req_capacity + 1]; the length relationship is pinned by the
|
// write_offsets has shape [write_req_capacity + 1]; the length relationship is pinned by the
|
||||||
// RuntimeCheck below, this matcher pins dtype + device.
|
// RuntimeCheck below, this matcher pins dtype + device.
|
||||||
SymbolicSize N_write_offsets = {"write_offsets_len"};
|
SymbolicSize N_write_offsets = {"write_offsets_len"};
|
||||||
TensorMatcher({N_write_offsets}).with_dtype<int64_t>().with_device<kDLCUDA>(device_).verify(write_offsets);
|
TensorMatcher({N_write_offsets}).with_dtype<int64_t>().with_device<kDLGPU>(device_).verify(write_offsets);
|
||||||
TensorMatcher({N_write_reqs}).with_dtype<int64_t>().with_device<kDLCUDA>(device_).verify(write_seed_slot_indices);
|
TensorMatcher({N_write_reqs}).with_dtype<int64_t>().with_device<kDLGPU>(device_).verify(write_seed_slot_indices);
|
||||||
TensorMatcher({1}).with_dtype<int32_t>().with_device<kDLCUDA>(device_).verify(write_num_valid_reqs);
|
TensorMatcher({1}).with_dtype<int32_t>().with_device<kDLGPU>(device_).verify(write_num_valid_reqs);
|
||||||
|
|
||||||
TensorMatcher({N_tokens})
|
TensorMatcher({N_tokens})
|
||||||
.with_dtype<int64_t>()
|
.with_dtype<int64_t>()
|
||||||
.with_device<kDLCUDA>(device_)
|
.with_device<kDLGPU>(device_)
|
||||||
.verify(input_ids)
|
.verify(input_ids)
|
||||||
.verify(positions)
|
.verify(positions)
|
||||||
.verify(out_cache_loc);
|
.verify(out_cache_loc);
|
||||||
@@ -232,47 +232,47 @@ inline void canary_write_step_cuda(
|
|||||||
if (enable_write_input_assert_bool) {
|
if (enable_write_input_assert_bool) {
|
||||||
TensorMatcher({N_tokens})
|
TensorMatcher({N_tokens})
|
||||||
.with_dtype<int64_t>()
|
.with_dtype<int64_t>()
|
||||||
.with_device<kDLCUDA>(device_)
|
.with_device<kDLGPU>(device_)
|
||||||
.verify(expected_input_tokens.value())
|
.verify(expected_input_tokens.value())
|
||||||
.verify(expected_input_positions.value());
|
.verify(expected_input_positions.value());
|
||||||
}
|
}
|
||||||
|
|
||||||
TensorMatcher({1}).with_dtype<int32_t>().with_device<kDLCUDA>(device_).verify(violation_write_index);
|
TensorMatcher({1}).with_dtype<int32_t>().with_device<kDLGPU>(device_).verify(violation_write_index);
|
||||||
SymbolicSize N_ring = {"ring_capacity"};
|
SymbolicSize N_ring = {"ring_capacity"};
|
||||||
TensorMatcher({N_ring, static_cast<int64_t>(kViolationFields)})
|
TensorMatcher({N_ring, static_cast<int64_t>(kViolationFields)})
|
||||||
.with_dtype<int64_t>()
|
.with_dtype<int64_t>()
|
||||||
.with_device<kDLCUDA>(device_)
|
.with_device<kDLGPU>(device_)
|
||||||
.verify(violation_ring);
|
.verify(violation_ring);
|
||||||
TensorMatcher({1})
|
TensorMatcher({1})
|
||||||
.with_dtype<int64_t>()
|
.with_dtype<int64_t>()
|
||||||
.with_device<kDLCUDA>(device_)
|
.with_device<kDLGPU>(device_)
|
||||||
.verify(slot_run_counter)
|
.verify(slot_run_counter)
|
||||||
.verify(kernel_run_counter);
|
.verify(kernel_run_counter);
|
||||||
TensorMatcher({1}).with_dtype<int32_t>().with_device<kDLCUDA>(device_).verify(enable_chain_position_assert);
|
TensorMatcher({1}).with_dtype<int32_t>().with_device<kDLGPU>(device_).verify(enable_chain_position_assert);
|
||||||
|
|
||||||
SymbolicSize N_real_kv_rows_0 = {"real_kv_rows_0"};
|
SymbolicSize N_real_kv_rows_0 = {"real_kv_rows_0"};
|
||||||
SymbolicSize N_real_kv_cols_0 = {"real_kv_cols_0"};
|
SymbolicSize N_real_kv_cols_0 = {"real_kv_cols_0"};
|
||||||
TensorMatcher({N_real_kv_rows_0, N_real_kv_cols_0})
|
TensorMatcher({N_real_kv_rows_0, N_real_kv_cols_0})
|
||||||
.with_dtype<uint8_t>()
|
.with_dtype<uint8_t>()
|
||||||
.with_device<kDLCUDA>(device_)
|
.with_device<kDLGPU>(device_)
|
||||||
.verify(real_kv_buf_0);
|
.verify(real_kv_buf_0);
|
||||||
SymbolicSize N_real_kv_rows_1 = {"real_kv_rows_1"};
|
SymbolicSize N_real_kv_rows_1 = {"real_kv_rows_1"};
|
||||||
SymbolicSize N_real_kv_cols_1 = {"real_kv_cols_1"};
|
SymbolicSize N_real_kv_cols_1 = {"real_kv_cols_1"};
|
||||||
TensorMatcher({N_real_kv_rows_1, N_real_kv_cols_1})
|
TensorMatcher({N_real_kv_rows_1, N_real_kv_cols_1})
|
||||||
.with_dtype<uint8_t>()
|
.with_dtype<uint8_t>()
|
||||||
.with_device<kDLCUDA>(device_)
|
.with_device<kDLGPU>(device_)
|
||||||
.verify(real_kv_buf_1);
|
.verify(real_kv_buf_1);
|
||||||
SymbolicSize N_real_kv_rows_2 = {"real_kv_rows_2"};
|
SymbolicSize N_real_kv_rows_2 = {"real_kv_rows_2"};
|
||||||
SymbolicSize N_real_kv_cols_2 = {"real_kv_cols_2"};
|
SymbolicSize N_real_kv_cols_2 = {"real_kv_cols_2"};
|
||||||
TensorMatcher({N_real_kv_rows_2, N_real_kv_cols_2})
|
TensorMatcher({N_real_kv_rows_2, N_real_kv_cols_2})
|
||||||
.with_dtype<uint8_t>()
|
.with_dtype<uint8_t>()
|
||||||
.with_device<kDLCUDA>(device_)
|
.with_device<kDLGPU>(device_)
|
||||||
.verify(real_kv_buf_2);
|
.verify(real_kv_buf_2);
|
||||||
SymbolicSize N_real_kv_rows_3 = {"real_kv_rows_3"};
|
SymbolicSize N_real_kv_rows_3 = {"real_kv_rows_3"};
|
||||||
SymbolicSize N_real_kv_cols_3 = {"real_kv_cols_3"};
|
SymbolicSize N_real_kv_cols_3 = {"real_kv_cols_3"};
|
||||||
TensorMatcher({N_real_kv_rows_3, N_real_kv_cols_3})
|
TensorMatcher({N_real_kv_rows_3, N_real_kv_cols_3})
|
||||||
.with_dtype<uint8_t>()
|
.with_dtype<uint8_t>()
|
||||||
.with_device<kDLCUDA>(device_)
|
.with_device<kDLGPU>(device_)
|
||||||
.verify(real_kv_buf_3);
|
.verify(real_kv_buf_3);
|
||||||
TensorMatcher({static_cast<int64_t>(kMaxRealKvSources), static_cast<int64_t>(kRealKvSourceFieldsPerEntry)})
|
TensorMatcher({static_cast<int64_t>(kMaxRealKvSources), static_cast<int64_t>(kRealKvSourceFieldsPerEntry)})
|
||||||
.with_dtype<int32_t>()
|
.with_dtype<int32_t>()
|
||||||
|
|||||||
@@ -5,9 +5,10 @@ from pathlib import Path
|
|||||||
|
|
||||||
import sglang.jit_kernel
|
import sglang.jit_kernel
|
||||||
from sglang.jit_kernel.kv_canary import consts
|
from sglang.jit_kernel.kv_canary import consts
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||||
|
|
||||||
register_cuda_ci(est_time=5, suite="base-b-kernel-unit-1-gpu-large")
|
register_cuda_ci(est_time=5, suite="base-b-kernel-unit-1-gpu-large")
|
||||||
|
register_amd_ci(est_time=5, suite="jit-kernel-unit-test-amd")
|
||||||
|
|
||||||
|
|
||||||
# Resolve the kernel source against the installed jit_kernel package rather
|
# Resolve the kernel source against the installed jit_kernel package rather
|
||||||
|
|||||||
@@ -33,9 +33,10 @@ from sglang.jit_kernel.tests.kv_canary._fixtures import (
|
|||||||
dummy_pseudo_tensors,
|
dummy_pseudo_tensors,
|
||||||
empty_extras,
|
empty_extras,
|
||||||
)
|
)
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||||
|
|
||||||
register_cuda_ci(est_time=60, suite="base-b-kernel-unit-1-gpu-large")
|
register_cuda_ci(est_time=60, suite="base-b-kernel-unit-1-gpu-large")
|
||||||
|
register_amd_ci(est_time=60, suite="jit-kernel-unit-test-amd")
|
||||||
|
|
||||||
_DEVICE = torch.device("cuda")
|
_DEVICE = torch.device("cuda")
|
||||||
|
|
||||||
|
|||||||
@@ -38,9 +38,10 @@ from sglang.jit_kernel.tests.kv_canary._fixtures import (
|
|||||||
empty_extras,
|
empty_extras,
|
||||||
make_req_to_token,
|
make_req_to_token,
|
||||||
)
|
)
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||||
|
|
||||||
register_cuda_ci(est_time=30, suite="base-b-kernel-unit-1-gpu-large")
|
register_cuda_ci(est_time=30, suite="base-b-kernel-unit-1-gpu-large")
|
||||||
|
register_amd_ci(est_time=30, suite="jit-kernel-unit-test-amd")
|
||||||
|
|
||||||
|
|
||||||
_DEVICE = torch.device("cuda")
|
_DEVICE = torch.device("cuda")
|
||||||
|
|||||||
@@ -20,9 +20,10 @@ from sglang.jit_kernel.tests.kv_canary._fuzz_driver import (
|
|||||||
run_fuzz_combo,
|
run_fuzz_combo,
|
||||||
)
|
)
|
||||||
from sglang.jit_kernel.tests.kv_canary._invariants import PlanInvariants
|
from sglang.jit_kernel.tests.kv_canary._invariants import PlanInvariants
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||||
|
|
||||||
register_cuda_ci(est_time=30, suite="base-b-kernel-unit-1-gpu-large")
|
register_cuda_ci(est_time=30, suite="base-b-kernel-unit-1-gpu-large")
|
||||||
|
register_amd_ci(est_time=30, suite="jit-kernel-unit-test-amd")
|
||||||
|
|
||||||
|
|
||||||
_DEVICE = torch.device("cuda")
|
_DEVICE = torch.device("cuda")
|
||||||
|
|||||||
@@ -20,9 +20,10 @@ from sglang.jit_kernel.tests.kv_canary._fixtures import (
|
|||||||
make_req_to_token,
|
make_req_to_token,
|
||||||
)
|
)
|
||||||
from sglang.jit_kernel.tests.kv_canary._invariants import PlanInvariants
|
from sglang.jit_kernel.tests.kv_canary._invariants import PlanInvariants
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||||
|
|
||||||
register_cuda_ci(est_time=30, suite="base-b-kernel-unit-1-gpu-large")
|
register_cuda_ci(est_time=30, suite="base-b-kernel-unit-1-gpu-large")
|
||||||
|
register_amd_ci(est_time=30, suite="jit-kernel-unit-test-amd")
|
||||||
|
|
||||||
|
|
||||||
_DEVICE = torch.device("cuda")
|
_DEVICE = torch.device("cuda")
|
||||||
|
|||||||
@@ -10,10 +10,11 @@ from sglang.jit_kernel.kv_canary.scatter_req_token_ids import (
|
|||||||
launch_scatter_req_token_ids_kernel,
|
launch_scatter_req_token_ids_kernel,
|
||||||
scatter_req_token_ids_torch_reference,
|
scatter_req_token_ids_torch_reference,
|
||||||
)
|
)
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
register_cuda_ci(est_time=10, suite="base-b-kernel-unit-1-gpu-large")
|
register_cuda_ci(est_time=10, suite="base-b-kernel-unit-1-gpu-large")
|
||||||
|
register_amd_ci(est_time=10, suite="jit-kernel-unit-test-amd")
|
||||||
|
|
||||||
|
|
||||||
_DEVICE = torch.device("cuda")
|
_DEVICE = torch.device("cuda")
|
||||||
|
|||||||
@@ -29,9 +29,10 @@ from sglang.jit_kernel.tests.kv_canary._fuzz_driver import (
|
|||||||
run_fuzz_combo,
|
run_fuzz_combo,
|
||||||
)
|
)
|
||||||
from sglang.jit_kernel.tests.kv_canary._invariants import VerifyInvariants
|
from sglang.jit_kernel.tests.kv_canary._invariants import VerifyInvariants
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||||
|
|
||||||
register_cuda_ci(est_time=30, suite="base-b-kernel-unit-1-gpu-large")
|
register_cuda_ci(est_time=30, suite="base-b-kernel-unit-1-gpu-large")
|
||||||
|
register_amd_ci(est_time=30, suite="jit-kernel-unit-test-amd")
|
||||||
|
|
||||||
|
|
||||||
_DEVICE = torch.device("cuda")
|
_DEVICE = torch.device("cuda")
|
||||||
|
|||||||
@@ -51,9 +51,10 @@ from sglang.jit_kernel.tests.kv_canary._hand_oracle import (
|
|||||||
_hand_fold_all,
|
_hand_fold_all,
|
||||||
_hand_fold_partial,
|
_hand_fold_partial,
|
||||||
)
|
)
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||||
|
|
||||||
register_cuda_ci(est_time=30, suite="base-b-kernel-unit-1-gpu-large")
|
register_cuda_ci(est_time=30, suite="base-b-kernel-unit-1-gpu-large")
|
||||||
|
register_amd_ci(est_time=30, suite="jit-kernel-unit-test-amd")
|
||||||
|
|
||||||
|
|
||||||
_DEVICE = torch.device("cuda")
|
_DEVICE = torch.device("cuda")
|
||||||
|
|||||||
@@ -29,9 +29,10 @@ from sglang.jit_kernel.tests.kv_canary._fuzz_driver import (
|
|||||||
run_fuzz_combo,
|
run_fuzz_combo,
|
||||||
)
|
)
|
||||||
from sglang.jit_kernel.tests.kv_canary._invariants import WriteInvariants
|
from sglang.jit_kernel.tests.kv_canary._invariants import WriteInvariants
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||||
|
|
||||||
register_cuda_ci(est_time=30, suite="base-b-kernel-unit-1-gpu-large")
|
register_cuda_ci(est_time=30, suite="base-b-kernel-unit-1-gpu-large")
|
||||||
|
register_amd_ci(est_time=30, suite="jit-kernel-unit-test-amd")
|
||||||
|
|
||||||
|
|
||||||
_DEVICE = torch.device("cuda")
|
_DEVICE = torch.device("cuda")
|
||||||
|
|||||||
@@ -49,9 +49,10 @@ from sglang.jit_kernel.tests.kv_canary._hand_oracle import (
|
|||||||
_hand_fold_all,
|
_hand_fold_all,
|
||||||
_hand_fold_partial,
|
_hand_fold_partial,
|
||||||
)
|
)
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||||
|
|
||||||
register_cuda_ci(est_time=30, suite="base-b-kernel-unit-1-gpu-large")
|
register_cuda_ci(est_time=30, suite="base-b-kernel-unit-1-gpu-large")
|
||||||
|
register_amd_ci(est_time=30, suite="jit-kernel-unit-test-amd")
|
||||||
|
|
||||||
|
|
||||||
_DEVICE = torch.device("cuda")
|
_DEVICE = torch.device("cuda")
|
||||||
|
|||||||
Reference in New Issue
Block a user