[JIT kernel] hd=512,1024 in JIT QK norm (cta based) (#17515)
Signed-off-by: vincentzed <207368749+vincentzed@users.noreply.github.com>
This commit is contained in:
@@ -74,8 +74,6 @@ def torch_impl_qknorm(
|
|||||||
k.copy_(k.float() * k_norm * k_weight.float())
|
k.copy_(k.float() * k_norm * k_weight.float())
|
||||||
|
|
||||||
|
|
||||||
HEAD_DIM = 128
|
|
||||||
|
|
||||||
BS_RANGE = get_benchmark_range(
|
BS_RANGE = get_benchmark_range(
|
||||||
full_range=[2**n for n in range(0, 14)],
|
full_range=[2**n for n in range(0, 14)],
|
||||||
ci_range=[16],
|
ci_range=[16],
|
||||||
@@ -88,17 +86,21 @@ KV_HEAD_RANGE = get_benchmark_range(
|
|||||||
full_range=[1, 2, 4, 8],
|
full_range=[1, 2, 4, 8],
|
||||||
ci_range=[1],
|
ci_range=[1],
|
||||||
)
|
)
|
||||||
|
HEAD_DIM_RANGE = get_benchmark_range(
|
||||||
|
full_range=[128, 256, 512, 1024],
|
||||||
|
ci_range=[128],
|
||||||
|
)
|
||||||
|
|
||||||
LINE_VALS = ["aot", "jit", "fi", "torch"]
|
LINE_VALS = ["aot", "jit", "fi", "torch"]
|
||||||
LINE_NAMES = ["SGL AOT Kernel", "SGL JIT Kernel", "FlashInfer", "PyTorch"]
|
LINE_NAMES = ["SGL AOT Kernel", "SGL JIT Kernel", "FlashInfer", "PyTorch"]
|
||||||
STYLES = [("orange", "-"), ("blue", "--"), ("green", "-."), ("red", ":")]
|
STYLES = [("orange", "-"), ("blue", "--"), ("green", "-."), ("red", ":")]
|
||||||
|
|
||||||
configs = list(itertools.product(GQA_RANGE, KV_HEAD_RANGE, BS_RANGE))
|
configs = list(itertools.product(HEAD_DIM_RANGE, GQA_RANGE, KV_HEAD_RANGE, BS_RANGE))
|
||||||
|
|
||||||
|
|
||||||
@triton.testing.perf_report(
|
@triton.testing.perf_report(
|
||||||
triton.testing.Benchmark(
|
triton.testing.Benchmark(
|
||||||
x_names=["GQA", "num_kv_heads", "batch_size"],
|
x_names=["head_dim", "GQA", "num_kv_heads", "batch_size"],
|
||||||
x_vals=configs,
|
x_vals=configs,
|
||||||
line_arg="provider",
|
line_arg="provider",
|
||||||
line_vals=LINE_VALS,
|
line_vals=LINE_VALS,
|
||||||
@@ -109,16 +111,18 @@ configs = list(itertools.product(GQA_RANGE, KV_HEAD_RANGE, BS_RANGE))
|
|||||||
args={},
|
args={},
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
def benchmark(batch_size: int, GQA: int, num_kv_heads: int, provider: str):
|
def benchmark(
|
||||||
|
head_dim: int, GQA: int, num_kv_heads: int, batch_size: int, provider: str
|
||||||
|
):
|
||||||
num_qo_heads = GQA * num_kv_heads
|
num_qo_heads = GQA * num_kv_heads
|
||||||
q = torch.randn(
|
q = torch.randn(
|
||||||
(batch_size, num_qo_heads, HEAD_DIM), dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE
|
(batch_size, num_qo_heads, head_dim), dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE
|
||||||
)
|
)
|
||||||
k = torch.randn(
|
k = torch.randn(
|
||||||
(batch_size, num_kv_heads, HEAD_DIM), dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE
|
(batch_size, num_kv_heads, head_dim), dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE
|
||||||
)
|
)
|
||||||
q_weight = torch.randn(HEAD_DIM, dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE)
|
q_weight = torch.randn(head_dim, dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE)
|
||||||
k_weight = torch.randn(HEAD_DIM, dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE)
|
k_weight = torch.randn(head_dim, dtype=DEFAULT_DTYPE, device=DEFAULT_DEVICE)
|
||||||
FN_MAP = {
|
FN_MAP = {
|
||||||
"aot": sglang_aot_qknorm,
|
"aot": sglang_aot_qknorm,
|
||||||
"jit": sglang_jit_qknorm,
|
"jit": sglang_jit_qknorm,
|
||||||
|
|||||||
@@ -33,8 +33,9 @@ struct QKNormParams {
|
|||||||
constexpr uint32_t kWarpsPerBlock = 4;
|
constexpr uint32_t kWarpsPerBlock = 4;
|
||||||
constexpr uint32_t kThreadsPerBlock = kWarpsPerBlock * device::kWarpThreads;
|
constexpr uint32_t kThreadsPerBlock = kWarpsPerBlock * device::kWarpThreads;
|
||||||
|
|
||||||
|
// Warp-level kernel for head_dim <= 256
|
||||||
template <int64_t kHeadDim, bool kUsePDL, typename Float>
|
template <int64_t kHeadDim, bool kUsePDL, typename Float>
|
||||||
__global__ void fused_qknorm(const QKNormParams __grid_constant__ params) {
|
__global__ void fused_qknorm_warp(const QKNormParams __grid_constant__ params) {
|
||||||
using namespace device;
|
using namespace device;
|
||||||
using Storage = norm::StorageType<Float, kHeadDim>;
|
using Storage = norm::StorageType<Float, kHeadDim>;
|
||||||
|
|
||||||
@@ -66,11 +67,47 @@ __global__ void fused_qknorm(const QKNormParams __grid_constant__ params) {
|
|||||||
PDLTriggerSecondary<kUsePDL>(); // launch secondary kernel
|
PDLTriggerSecondary<kUsePDL>(); // launch secondary kernel
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// For CTA level, used for head_dim > 256 (512,1024)
|
||||||
|
template <int64_t kHeadDim, bool kUsePDL, typename Float>
|
||||||
|
__global__ void fused_qknorm_cta(const QKNormParams __grid_constant__ params) {
|
||||||
|
using namespace device;
|
||||||
|
using Storage = norm::StorageType<Float, kHeadDim>;
|
||||||
|
|
||||||
|
constexpr auto kNumThreads = host::norm::get_cta_threads<Float, kHeadDim>();
|
||||||
|
constexpr auto kNumWarps = kNumThreads / kWarpThreads;
|
||||||
|
|
||||||
|
static_assert(sizeof(Float) == 2, "Only support FP16/BF16");
|
||||||
|
const auto& [q, k, q_stride, k_stride, num_qo_heads, num_kv_heads, eps, q_weight, k_weight, num_tokens] = params;
|
||||||
|
|
||||||
|
const auto num_q_and_k_heads = num_qo_heads + num_kv_heads;
|
||||||
|
const auto num_works = num_q_and_k_heads * num_tokens;
|
||||||
|
const auto gmem = tile::Memory<Storage>::cta(kNumThreads);
|
||||||
|
__shared__ float smem[norm::kSmemBufferSize];
|
||||||
|
|
||||||
|
PDLWaitPrimary<kUsePDL>(); // wait for primary kernel
|
||||||
|
|
||||||
|
for (auto idx = blockIdx.x; idx < num_works; idx += gridDim.x) {
|
||||||
|
const int64_t token_id = idx / num_q_and_k_heads;
|
||||||
|
const int64_t head_id = idx % num_q_and_k_heads;
|
||||||
|
const auto load_q = head_id < num_qo_heads;
|
||||||
|
const auto input = load_q ? pointer::offset(q, 2 * (token_id * q_stride + head_id * kHeadDim))
|
||||||
|
: pointer::offset(k, 2 * (token_id * k_stride + head_id * kHeadDim));
|
||||||
|
const auto weight = load_q ? q_weight : k_weight;
|
||||||
|
const auto input_vec = gmem.load(input);
|
||||||
|
const auto weight_vec = gmem.load(weight);
|
||||||
|
const auto output_vec = norm::apply_norm_cta<kHeadDim>(input_vec, weight_vec, eps, smem, kNumWarps);
|
||||||
|
gmem.store(input, output_vec);
|
||||||
|
}
|
||||||
|
|
||||||
|
PDLTriggerSecondary<kUsePDL>(); // launch secondary kernel
|
||||||
|
}
|
||||||
|
|
||||||
|
// Warp-level kernel struct for head_dim <= 256
|
||||||
template <int64_t kHeadDim, bool kUsePDL, typename DType>
|
template <int64_t kHeadDim, bool kUsePDL, typename DType>
|
||||||
struct QKNormKernel {
|
struct QKNormKernelWarp {
|
||||||
static_assert(std::is_same_v<DType, fp16_t> || std::is_same_v<DType, bf16_t>);
|
static_assert(std::is_same_v<DType, fp16_t> || std::is_same_v<DType, bf16_t>);
|
||||||
static_assert(!host::norm::should_use_cta<DType, kHeadDim>(), "Head dim too large for QKNorm");
|
static_assert(!host::norm::should_use_cta<DType, kHeadDim>(), "Use QKNormKernelCTA for head_dim > 256");
|
||||||
static constexpr auto kernel = fused_qknorm<kHeadDim, kUsePDL, DType>;
|
static constexpr auto kernel = fused_qknorm_warp<kHeadDim, kUsePDL, DType>;
|
||||||
|
|
||||||
static void
|
static void
|
||||||
run(const tvm::ffi::TensorView q,
|
run(const tvm::ffi::TensorView q,
|
||||||
@@ -138,4 +175,83 @@ struct QKNormKernel {
|
|||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
|
// This goes with fused_qknorm_cta
|
||||||
|
template <int64_t kHeadDim, bool kUsePDL, typename DType>
|
||||||
|
struct QKNormKernelCTA {
|
||||||
|
static_assert(std::is_same_v<DType, fp16_t> || std::is_same_v<DType, bf16_t>);
|
||||||
|
static_assert(host::norm::should_use_cta<DType, kHeadDim>(), "Use QKNormKernelWarp for head_dim <= 256");
|
||||||
|
static constexpr auto kernel = fused_qknorm_cta<kHeadDim, kUsePDL, DType>;
|
||||||
|
static constexpr auto kNumThreads = host::norm::get_cta_threads<DType, kHeadDim>();
|
||||||
|
|
||||||
|
static void
|
||||||
|
run(const tvm::ffi::TensorView q,
|
||||||
|
const tvm::ffi::TensorView k,
|
||||||
|
const tvm::ffi::TensorView q_weight,
|
||||||
|
const tvm::ffi::TensorView k_weight,
|
||||||
|
float eps) {
|
||||||
|
using namespace host;
|
||||||
|
|
||||||
|
auto N = SymbolicSize{"num_tokens"};
|
||||||
|
auto Q = SymbolicSize{"num_qo_heads"};
|
||||||
|
auto K = SymbolicSize{"num_kv_heads"};
|
||||||
|
auto D = SymbolicSize{"head_dim"};
|
||||||
|
auto Sq = SymbolicSize{"q_stride"};
|
||||||
|
auto Sk = SymbolicSize{"k_stride"};
|
||||||
|
auto device = SymbolicDevice{};
|
||||||
|
D.set_value(kHeadDim);
|
||||||
|
device.set_options<kDLCUDA>();
|
||||||
|
|
||||||
|
TensorMatcher({N, Q, D}) // q input
|
||||||
|
.with_strides({Sq, D, 1})
|
||||||
|
.with_dtype<DType>()
|
||||||
|
.with_device(device)
|
||||||
|
.verify(q);
|
||||||
|
TensorMatcher({N, K, D}) // k input
|
||||||
|
.with_strides({Sk, D, 1})
|
||||||
|
.with_dtype<DType>()
|
||||||
|
.with_device(device)
|
||||||
|
.verify(k);
|
||||||
|
TensorMatcher({D}) // weight
|
||||||
|
.with_dtype<DType>()
|
||||||
|
.with_device(device)
|
||||||
|
.verify(q_weight)
|
||||||
|
.verify(k_weight);
|
||||||
|
|
||||||
|
const auto num_tokens = static_cast<uint32_t>(N.unwrap());
|
||||||
|
const auto num_qo_heads = static_cast<uint32_t>(Q.unwrap());
|
||||||
|
const auto num_kv_heads = static_cast<uint32_t>(K.unwrap());
|
||||||
|
|
||||||
|
// NOTE: we offset the k here to reduce computation cost in the kernel
|
||||||
|
const auto params = QKNormParams{
|
||||||
|
.q = q.data_ptr(),
|
||||||
|
.k = pointer::offset(k.data_ptr(), -2 * static_cast<int64_t>(num_qo_heads) * kHeadDim),
|
||||||
|
.q_stride = static_cast<int64_t>(Sq.unwrap()),
|
||||||
|
.k_stride = static_cast<int64_t>(Sk.unwrap()),
|
||||||
|
.num_qo_heads = num_qo_heads,
|
||||||
|
.num_kv_heads = num_kv_heads,
|
||||||
|
.eps = eps,
|
||||||
|
.q_weight = q_weight.data_ptr(),
|
||||||
|
.k_weight = k_weight.data_ptr(),
|
||||||
|
.num_tokens = num_tokens,
|
||||||
|
};
|
||||||
|
|
||||||
|
static const uint32_t max_occupancy = runtime::get_blocks_per_sm(kernel, kNumThreads);
|
||||||
|
static const uint32_t kNumSM = runtime::get_sm_count(device.unwrap().device_id);
|
||||||
|
|
||||||
|
const auto num_works = (num_qo_heads + num_kv_heads) * num_tokens;
|
||||||
|
|
||||||
|
// we use persistent kernel, which limit the number of blocks to reduce overhead
|
||||||
|
const auto num_blocks = std::min<uint32_t>(num_works, max_occupancy * kNumSM);
|
||||||
|
LaunchKernel(num_blocks, kNumThreads, device.unwrap()) //
|
||||||
|
.enable_pdl(kUsePDL)(kernel, params);
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
// Unified dispatch: select warp or CTA kernel based on head_dim
|
||||||
|
template <int64_t kHeadDim, bool kUsePDL, typename DType>
|
||||||
|
using QKNormKernel = std::conditional_t<
|
||||||
|
host::norm::should_use_cta<DType, kHeadDim>(),
|
||||||
|
QKNormKernelCTA<kHeadDim, kUsePDL, DType>,
|
||||||
|
QKNormKernelWarp<kHeadDim, kUsePDL, DType>>;
|
||||||
|
|
||||||
} // namespace
|
} // namespace
|
||||||
|
|||||||
@@ -63,7 +63,7 @@ BS_LIST = [2**n for n in range(0, 14)]
|
|||||||
BS_LIST += [x + 1 + i for i, x in enumerate(BS_LIST)]
|
BS_LIST += [x + 1 + i for i, x in enumerate(BS_LIST)]
|
||||||
N_K_LIST = [2, 4]
|
N_K_LIST = [2, 4]
|
||||||
N_Q_LIST = [8, 16]
|
N_Q_LIST = [8, 16]
|
||||||
HEAD_DIM_LIST = [64, 128, 256]
|
HEAD_DIM_LIST = [64, 128, 256, 512, 1024]
|
||||||
DEVICE = "cuda"
|
DEVICE = "cuda"
|
||||||
DTYPE = torch.bfloat16
|
DTYPE = torch.bfloat16
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user