Support JIT fused A GEMM (MLA down projection) and support GLM-5 hidden size, SM120 (#27397)
Co-authored-by: Brayden Zhong <brayden@radixark.ai>
This commit is contained in:
co-authored by
Brayden Zhong
parent
592f6c849b
commit
e4253b39e2
@@ -0,0 +1,631 @@
|
||||
/*
|
||||
* Adapted from
|
||||
* https://github.com/NVIDIA/TensorRT-LLM/blob/619709fc33bd5dc268f19d6a741fe7ed51c0f8f5/cpp/tensorrt_llm/kernels/dsv3MinLatencyKernels/dsv3FusedAGemm.cu
|
||||
*
|
||||
* Copyright (c) 2019-2024, NVIDIA CORPORATION. All rights reserved.
|
||||
* Copyright (c) 2021, NAVER Corp. Authored by CLOVA.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#include <sgl_kernel/tensor.h>
|
||||
#include <sgl_kernel/utils.h>
|
||||
|
||||
#include <sgl_kernel/utils.cuh>
|
||||
|
||||
#include <tvm/ffi/container/tensor.h>
|
||||
|
||||
#include <algorithm>
|
||||
#include <cuda_bf16.h>
|
||||
#include <utility>
|
||||
|
||||
namespace {
|
||||
|
||||
using bf16_t = __nv_bfloat16;
|
||||
|
||||
__device__ void hmma_16_8_16_f32acc_bf16ab(
|
||||
float (&d_reg)[4], const bf16_t (&a_reg)[8], const bf16_t (&b_reg)[4], float const (&c_reg)[4]) {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
|
||||
uint32_t a0 = *reinterpret_cast<uint32_t const*>(a_reg + 0);
|
||||
uint32_t a1 = *reinterpret_cast<uint32_t const*>(a_reg + 2);
|
||||
uint32_t a2 = *reinterpret_cast<uint32_t const*>(a_reg + 4);
|
||||
uint32_t a3 = *reinterpret_cast<uint32_t const*>(a_reg + 6);
|
||||
uint32_t b0 = *reinterpret_cast<uint32_t const*>(b_reg + 0);
|
||||
uint32_t b1 = *reinterpret_cast<uint32_t const*>(b_reg + 2);
|
||||
asm volatile(
|
||||
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
|
||||
"{%0, %1, %2, %3},"
|
||||
"{%4, %5, %6, %7},"
|
||||
"{%8, %9},"
|
||||
"{%10, %11, %12, %13};\n"
|
||||
: "=f"(d_reg[0]), "=f"(d_reg[1]), "=f"(d_reg[2]), "=f"(d_reg[3])
|
||||
: "r"(a0),
|
||||
"r"(a1),
|
||||
"r"(a2),
|
||||
"r"(a3),
|
||||
"r"(b0),
|
||||
"r"(b1),
|
||||
"f"(d_reg[0]),
|
||||
"f"(d_reg[1]),
|
||||
"f"(d_reg[2]),
|
||||
"f"(d_reg[3]));
|
||||
#endif
|
||||
}
|
||||
|
||||
extern "C" {
|
||||
__device__ uint32_t __nvvm_get_smem_pointer(void*);
|
||||
}
|
||||
|
||||
__device__ void ldgsts_128(void const* gPtr, void* sPtr, uint32_t pred) {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
|
||||
if (pred) {
|
||||
uint32_t smemPtrAsUint32 = __nvvm_get_smem_pointer(sPtr);
|
||||
asm volatile("cp.async.cg.shared.global.L2::128B [%0], [%1], %2;\n" ::"r"(smemPtrAsUint32), "l"(gPtr), "n"(16));
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
__device__ void ldsm_x4(void* smem_ptr, uint32_t* reg_ptr) {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
|
||||
asm volatile("ldmatrix.sync.aligned.x4.m8n8.shared.b16 {%0, %1, %2, %3}, [%4];\n"
|
||||
: "=r"(reg_ptr[0]), "=r"(reg_ptr[1]), "=r"(reg_ptr[2]), "=r"(reg_ptr[3])
|
||||
: "r"(__nvvm_get_smem_pointer(smem_ptr)));
|
||||
#endif
|
||||
}
|
||||
|
||||
template <class Type>
|
||||
__device__ int apply_swizzle_343_on_elem_row_col(int row_idx_, int col_idx_) {
|
||||
uint32_t row_idx = *reinterpret_cast<uint32_t*>(&row_idx_);
|
||||
uint32_t col_idx = *reinterpret_cast<uint32_t*>(&col_idx_);
|
||||
row_idx = row_idx % 8;
|
||||
row_idx = row_idx * (16 / sizeof(Type));
|
||||
col_idx = col_idx ^ row_idx;
|
||||
return *reinterpret_cast<int*>(&col_idx);
|
||||
}
|
||||
|
||||
__device__ void initialize_barrier(uint64_t* smem_barrier, int thread_count = 1) {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
|
||||
uint32_t smem_int_ptr = __nvvm_get_smem_pointer(smem_barrier);
|
||||
asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;\n" ::"r"(smem_int_ptr), "r"(thread_count));
|
||||
#endif
|
||||
}
|
||||
|
||||
__device__ void wait_barrier(uint64_t* smem_barrier, int phase_bit) {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
|
||||
uint32_t smem_int_ptr = __nvvm_get_smem_pointer(smem_barrier);
|
||||
asm volatile(
|
||||
"{\n"
|
||||
".reg .pred P1;\n"
|
||||
"LAB_WAIT:\n"
|
||||
"mbarrier.try_wait.parity.shared::cta.b64 P1, [%0], %1;\n"
|
||||
"@P1 bra DONE;\n"
|
||||
"bra LAB_WAIT;\n"
|
||||
"DONE:\n"
|
||||
"}\n" ::"r"(smem_int_ptr),
|
||||
"r"(phase_bit));
|
||||
#endif
|
||||
}
|
||||
|
||||
__device__ bool try_wait_barrier(uint64_t* smem_ptr, int phase_bit) {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
|
||||
uint32_t wait_complete;
|
||||
uint32_t smem_int_ptr = __nvvm_get_smem_pointer(smem_ptr);
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred P1; \n\t"
|
||||
"mbarrier.try_wait.parity.shared::cta.b64 P1, [%1], %2; \n\t"
|
||||
"selp.b32 %0, 1, 0, P1; \n\t"
|
||||
"}"
|
||||
: "=r"(wait_complete)
|
||||
: "r"(smem_int_ptr), "r"(phase_bit));
|
||||
return static_cast<bool>(wait_complete);
|
||||
#endif
|
||||
return false;
|
||||
}
|
||||
|
||||
__device__ void arrive_barrier(uint64_t* smem_barrier) {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
|
||||
uint32_t smem_int_ptr = __nvvm_get_smem_pointer(smem_barrier);
|
||||
asm volatile(
|
||||
"{\n"
|
||||
".reg .b64 state; \n"
|
||||
"mbarrier.arrive.shared::cta.b64 state, [%0];\n"
|
||||
"}\n" ::"r"(smem_int_ptr));
|
||||
#endif
|
||||
}
|
||||
|
||||
__device__ void ldgsts_arrive(uint64_t* smem_barrier) {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
|
||||
uint32_t smem_int_ptr = __nvvm_get_smem_pointer(smem_barrier);
|
||||
asm volatile("cp.async.mbarrier.arrive.noinc.shared.b64 [%0];" : : "r"(smem_int_ptr));
|
||||
#endif
|
||||
}
|
||||
|
||||
template <int gemm_k, int tile_m, int tile_k, int stage_cnt>
|
||||
struct GmemLoaderA {
|
||||
static constexpr int elem_bytes = 2;
|
||||
static constexpr int vec_bytes = 16;
|
||||
static constexpr int vec_elems = vec_bytes / elem_bytes;
|
||||
static constexpr int thread_cnt = 64;
|
||||
static_assert((tile_m * tile_k) % (vec_elems * thread_cnt) == 0);
|
||||
static constexpr int a_inst_cnt_per_iter = (tile_m * tile_k) / (vec_elems * thread_cnt);
|
||||
static_assert(gemm_k % tile_k == 0);
|
||||
static constexpr int k_iter_cnt = gemm_k / tile_k;
|
||||
|
||||
static constexpr int mma_warp_cnt = 4;
|
||||
static constexpr int per_mma_warp_k = tile_k / mma_warp_cnt;
|
||||
static constexpr int k_each_chunk = gemm_k / mma_warp_cnt;
|
||||
|
||||
private:
|
||||
__device__ int k_project(int tile_k_idx) {
|
||||
return (tile_k_idx / per_mma_warp_k * k_each_chunk) + (tile_k_idx % per_mma_warp_k);
|
||||
}
|
||||
|
||||
public:
|
||||
__device__ GmemLoaderA(bf16_t const* gmem_a_local_, bf16_t* smem_a_, uint64_t* smem_barrier_)
|
||||
: gmem_a(gmem_a_local_), smem_a(smem_a_), smem_barrier(smem_barrier_), local_tid(threadIdx.x % thread_cnt) {}
|
||||
|
||||
__device__ void prepare() {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
|
||||
#pragma unroll
|
||||
for (int i = 0; i < a_inst_cnt_per_iter; i++) {
|
||||
int linear_idx = local_tid * vec_elems + i * thread_cnt * vec_elems;
|
||||
int m_idx = linear_idx / tile_k;
|
||||
int k_idx = linear_idx % tile_k;
|
||||
k_idx = apply_swizzle_343_on_elem_row_col<bf16_t>(m_idx, k_idx);
|
||||
a_smem_offsets[i] = m_idx * tile_k + k_idx;
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
__device__ void issue_mainloop() {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
|
||||
#pragma unroll 1
|
||||
for (int loop_idx = 0; loop_idx < k_iter_cnt; loop_idx++) {
|
||||
if (need_wait) {
|
||||
wait_barrier(smem_barrier + 1 + stage_idx * 2, phase_bit);
|
||||
}
|
||||
int next_stage_idx = stage_idx + 1;
|
||||
int next_phase_bit = next_stage_idx == stage_cnt ? phase_bit ^ 1 : phase_bit;
|
||||
next_stage_idx = next_stage_idx == stage_cnt ? 0 : next_stage_idx;
|
||||
if (loop_idx != k_iter_cnt - 1) {
|
||||
need_wait = !try_wait_barrier(smem_barrier + 1 + next_stage_idx * 2, next_phase_bit);
|
||||
}
|
||||
|
||||
#pragma unroll
|
||||
for (int i = 0; i < a_inst_cnt_per_iter; i++) {
|
||||
int smem_offset = a_smem_offsets[i];
|
||||
bf16_t* smem_ptr_this_iter = smem_a + stage_idx * tile_m * tile_k + smem_offset;
|
||||
int linear_idx = local_tid * vec_elems + i * thread_cnt * vec_elems;
|
||||
int m_idx = linear_idx / tile_k;
|
||||
int k_idx = linear_idx % tile_k;
|
||||
int gmem_offset = m_idx * gemm_k + k_project(k_idx);
|
||||
bf16_t const* gmem_ptr_this_iter = gmem_a + gmem_offset;
|
||||
ldgsts_128(gmem_ptr_this_iter, smem_ptr_this_iter, true);
|
||||
}
|
||||
ldgsts_arrive(smem_barrier + stage_idx * 2);
|
||||
|
||||
stage_idx = next_stage_idx;
|
||||
phase_bit = next_phase_bit;
|
||||
gmem_a += per_mma_warp_k;
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
bf16_t const* gmem_a;
|
||||
bf16_t* smem_a;
|
||||
uint64_t* smem_barrier;
|
||||
int local_tid;
|
||||
int stage_idx = 0;
|
||||
int phase_bit = 1;
|
||||
bool need_wait = true;
|
||||
|
||||
int a_smem_offsets[a_inst_cnt_per_iter];
|
||||
};
|
||||
|
||||
template <int gemm_k, int tile_n, int tile_k, int stage_cnt>
|
||||
struct GmemLoaderB {
|
||||
static constexpr int elem_bytes = 2;
|
||||
static constexpr int vec_bytes = 16;
|
||||
static constexpr int vec_elems = vec_bytes / elem_bytes;
|
||||
static constexpr int thread_cnt = 64;
|
||||
static_assert((tile_n * tile_k) % (vec_elems * thread_cnt) == 0);
|
||||
static constexpr int b_inst_cnt_per_iter = (tile_n * tile_k) / (vec_elems * thread_cnt);
|
||||
static_assert(gemm_k % tile_k == 0);
|
||||
static constexpr int k_iter_cnt = gemm_k / tile_k;
|
||||
|
||||
static constexpr int mma_warp_cnt = 4;
|
||||
static constexpr int per_mma_warp_k = tile_k / mma_warp_cnt;
|
||||
static constexpr int k_each_chunk = gemm_k / mma_warp_cnt;
|
||||
|
||||
private:
|
||||
__device__ int k_project(int tile_k_idx) {
|
||||
return (tile_k_idx / per_mma_warp_k * k_each_chunk) + (tile_k_idx % per_mma_warp_k);
|
||||
}
|
||||
|
||||
public:
|
||||
__device__ GmemLoaderB(bf16_t const* gmem_b_local_, bf16_t* smem_b_, uint64_t* smem_barrier_, int gemm_n_)
|
||||
: gmem_b(gmem_b_local_),
|
||||
smem_b(smem_b_),
|
||||
smem_barrier(smem_barrier_),
|
||||
gemm_n(gemm_n_),
|
||||
local_tid(threadIdx.x % thread_cnt) {}
|
||||
|
||||
__device__ void prepare() {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
|
||||
#pragma unroll
|
||||
for (int i = 0; i < b_inst_cnt_per_iter; i++) {
|
||||
int linear_idx = local_tid * vec_elems + i * thread_cnt * vec_elems;
|
||||
int n_idx = linear_idx / tile_k;
|
||||
int k_idx = linear_idx % tile_k;
|
||||
k_idx = apply_swizzle_343_on_elem_row_col<bf16_t>(n_idx, k_idx);
|
||||
b_smem_offsets[i] = n_idx * tile_k + k_idx;
|
||||
preds[i] = n_idx < gemm_n;
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
__device__ void issue_mainloop() {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
|
||||
cudaGridDependencySynchronize();
|
||||
#pragma unroll 1
|
||||
for (int loop_idx = 0; loop_idx < k_iter_cnt; loop_idx++) {
|
||||
if (need_wait) {
|
||||
wait_barrier(smem_barrier + 1 + stage_idx * 2, phase_bit);
|
||||
}
|
||||
int next_stage_idx = stage_idx + 1;
|
||||
int next_phase_bit = next_stage_idx == stage_cnt ? phase_bit ^ 1 : phase_bit;
|
||||
next_stage_idx = next_stage_idx == stage_cnt ? 0 : next_stage_idx;
|
||||
if (loop_idx != k_iter_cnt - 1) {
|
||||
need_wait = !try_wait_barrier(smem_barrier + 1 + next_stage_idx * 2, next_phase_bit);
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i = 0; i < b_inst_cnt_per_iter; i++) {
|
||||
int smem_offset = b_smem_offsets[i];
|
||||
bf16_t* smem_ptr_this_iter = smem_b + stage_idx * tile_n * tile_k + smem_offset;
|
||||
int linear_idx = local_tid * vec_elems + i * thread_cnt * vec_elems;
|
||||
int n_idx = linear_idx / tile_k;
|
||||
int k_idx = linear_idx % tile_k;
|
||||
int gmem_offset = n_idx * gemm_k + k_project(k_idx);
|
||||
bf16_t const* gmem_ptr_this_iter = gmem_b + gmem_offset;
|
||||
ldgsts_128(gmem_ptr_this_iter, smem_ptr_this_iter, preds[i]);
|
||||
}
|
||||
ldgsts_arrive(smem_barrier + stage_idx * 2);
|
||||
|
||||
stage_idx = next_stage_idx;
|
||||
phase_bit = next_phase_bit;
|
||||
gmem_b += per_mma_warp_k;
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
bf16_t const* gmem_b;
|
||||
bf16_t* smem_b;
|
||||
uint64_t* smem_barrier;
|
||||
int gemm_n;
|
||||
int local_tid;
|
||||
int stage_idx = 0;
|
||||
int phase_bit = 1;
|
||||
bool need_wait = true;
|
||||
|
||||
int b_smem_offsets[b_inst_cnt_per_iter];
|
||||
uint32_t preds[b_inst_cnt_per_iter];
|
||||
};
|
||||
|
||||
template <int gemm_m, int gemm_k, int tile_m, int tile_n, int tile_k, int stage_cnt>
|
||||
struct MmaComputer {
|
||||
static constexpr int elem_bytes = 2;
|
||||
static constexpr int thread_cnt = 128;
|
||||
static_assert(gemm_k % tile_k == 0);
|
||||
static_assert(tile_k % (thread_cnt / 32) == 0);
|
||||
static constexpr int per_warp_tile_k = tile_k / (thread_cnt / 32);
|
||||
static constexpr int k_iter_cnt = gemm_k / tile_k;
|
||||
static constexpr int k_phase_cnt = per_warp_tile_k / 16;
|
||||
static constexpr int m_iter_cnt = (tile_m + 15) / 16;
|
||||
static constexpr int n_iter_cnt = (tile_n + 7) / 8;
|
||||
static_assert(m_iter_cnt == 1);
|
||||
static_assert(n_iter_cnt == 1 || n_iter_cnt == 2);
|
||||
|
||||
__device__ MmaComputer(
|
||||
bf16_t* gmem_c_local_, bf16_t* smem_a_, bf16_t* smem_b_, uint64_t* smem_barrier_, int warp_idx_, int gemm_n_)
|
||||
: gmem_c(gmem_c_local_),
|
||||
smem_a(smem_a_),
|
||||
smem_b(smem_b_),
|
||||
smem_barrier(smem_barrier_),
|
||||
warp_idx(warp_idx_ - (thread_cnt / 32)),
|
||||
gemm_n(gemm_n_) {}
|
||||
|
||||
private:
|
||||
__device__ constexpr int internal_b_atom_func(int tid) {
|
||||
if constexpr (tile_n < 8) {
|
||||
return (tid % tile_n) + ((tid % 8) / tile_n * 0) + tid / 8 * 8 * tile_n;
|
||||
} else {
|
||||
return (tid % 8) + ((tid % 32) / 8 * (tile_n * 8));
|
||||
}
|
||||
}
|
||||
|
||||
public:
|
||||
__device__ void prepare() {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
|
||||
#pragma unroll
|
||||
for (int i = 0; i < k_phase_cnt; i++) {
|
||||
int linear_idx = (lane_idx % 16) + (lane_idx / 16) * 128 + i * 256;
|
||||
int m_idx = linear_idx % tile_m;
|
||||
int k_idx = linear_idx / tile_m + warp_k_offset_in_tile_k;
|
||||
k_idx = apply_swizzle_343_on_elem_row_col<bf16_t>(m_idx, k_idx);
|
||||
a_smem_offsets[0][i] = m_idx * tile_k + k_idx;
|
||||
}
|
||||
#pragma unroll
|
||||
for (int n_iter_idx = 0; n_iter_idx < n_iter_cnt; n_iter_idx++) {
|
||||
#pragma unroll
|
||||
for (int i = 0; i < k_phase_cnt; i += 2) {
|
||||
int linear_idx = internal_b_atom_func(lane_idx) + i * tile_n * 16 + n_iter_idx * 8;
|
||||
int n_idx = linear_idx % tile_n;
|
||||
int k_idx = linear_idx / tile_n + warp_k_offset_in_tile_k;
|
||||
k_idx = apply_swizzle_343_on_elem_row_col<bf16_t>(n_idx, k_idx);
|
||||
b_smem_offsets[n_iter_idx][i] = n_idx * tile_k + k_idx;
|
||||
}
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
__device__ void issue_mainloop() {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
|
||||
#pragma unroll 1
|
||||
for (int loop_idx = 0; loop_idx < k_iter_cnt; loop_idx++) {
|
||||
wait_barrier(smem_barrier + 0 + stage_idx * 2, phase_bit);
|
||||
|
||||
#pragma unroll
|
||||
for (int i = 0; i < k_phase_cnt; i++) {
|
||||
int smem_offset = a_smem_offsets[0][i];
|
||||
bf16_t* smem_ptr_this_iter = smem_a + stage_idx * tile_m * tile_k + smem_offset;
|
||||
ldsm_x4(smem_ptr_this_iter, reinterpret_cast<uint32_t*>(a_reg[0][i]));
|
||||
}
|
||||
|
||||
#pragma unroll
|
||||
for (int n_iter_idx = 0; n_iter_idx < n_iter_cnt; n_iter_idx++) {
|
||||
#pragma unroll
|
||||
for (int i = 0; i < k_phase_cnt; i += 2) {
|
||||
int smem_offset = b_smem_offsets[n_iter_idx][i];
|
||||
bf16_t* smem_ptr_this_iter = smem_b + stage_idx * tile_n * tile_k + smem_offset;
|
||||
ldsm_x4(smem_ptr_this_iter, reinterpret_cast<uint32_t*>(b_reg[n_iter_idx][i]));
|
||||
}
|
||||
}
|
||||
|
||||
#pragma unroll
|
||||
for (int k_iter_idx = 0; k_iter_idx < k_phase_cnt; k_iter_idx++) {
|
||||
#pragma unroll
|
||||
for (int n_iter_idx = 0; n_iter_idx < n_iter_cnt; n_iter_idx++) {
|
||||
hmma_16_8_16_f32acc_bf16ab(
|
||||
acc_reg[0][n_iter_idx], a_reg[0][k_iter_idx], b_reg[n_iter_idx][k_iter_idx], acc_reg[0][n_iter_idx]);
|
||||
}
|
||||
}
|
||||
::arrive_barrier(smem_barrier + 1 + stage_idx * 2);
|
||||
stage_idx += 1;
|
||||
phase_bit = stage_idx == stage_cnt ? phase_bit ^ 1 : phase_bit;
|
||||
stage_idx = stage_idx == stage_cnt ? 0 : stage_idx;
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
__device__ void epi() {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
|
||||
asm volatile("bar.sync %0, %1;" : : "r"(1), "r"(thread_cnt));
|
||||
constexpr int thread_m = 2;
|
||||
constexpr int thread_n = 2 * n_iter_cnt;
|
||||
constexpr int cta_mma_n = n_iter_cnt * 8;
|
||||
float acc_reg_reorg[thread_m][thread_n];
|
||||
|
||||
for (int i = 0; i < thread_m; i++) {
|
||||
for (int j = 0; j < thread_n; j++) {
|
||||
acc_reg_reorg[i][j] = acc_reg[0][j / 2][(j % 2) + (i * 2)];
|
||||
}
|
||||
}
|
||||
|
||||
float* smem_c = reinterpret_cast<float*>(smem_a);
|
||||
auto smem_c_index_func = [&](int m_idx, int n_idx) {
|
||||
int group_rows = 32 / cta_mma_n;
|
||||
int group_cnt = 2;
|
||||
return (m_idx % group_rows * cta_mma_n) + (m_idx / group_rows * (32 + group_cnt)) + n_idx;
|
||||
};
|
||||
constexpr int cosize_smem_c = ((tile_m * cta_mma_n) / 32) * (32 + 2);
|
||||
|
||||
#pragma unroll
|
||||
for (int m_idx_thread = 0; m_idx_thread < thread_m; m_idx_thread++) {
|
||||
#pragma unroll
|
||||
for (int n_idx_thread = 0; n_idx_thread < thread_n; n_idx_thread++) {
|
||||
int m_idx = (lane_idx / 4) + m_idx_thread * 8;
|
||||
int n_idx = ((lane_idx % 4) * 2) + (n_idx_thread % 2) + (n_idx_thread / 2) * 8;
|
||||
smem_c[cosize_smem_c * warp_idx + smem_c_index_func(m_idx, n_idx)] = acc_reg_reorg[m_idx_thread][n_idx_thread];
|
||||
}
|
||||
}
|
||||
asm volatile("bar.sync %0, %1;" : : "r"(1), "r"(thread_cnt));
|
||||
|
||||
if (warp_idx == 0) {
|
||||
constexpr int final_acc_reg_cnt = (tile_m * tile_n + 31) / 32;
|
||||
float acc_final[final_acc_reg_cnt]{};
|
||||
|
||||
#pragma unroll
|
||||
for (int reg_idx = 0; reg_idx < final_acc_reg_cnt; reg_idx++) {
|
||||
int linear_idx = reg_idx * 32 + lane_idx;
|
||||
int m_idx = linear_idx % tile_m;
|
||||
int n_idx = linear_idx / tile_m;
|
||||
acc_final[reg_idx] += smem_c[smem_c_index_func(m_idx, n_idx) + 0 * cosize_smem_c] +
|
||||
smem_c[smem_c_index_func(m_idx, n_idx) + 1 * cosize_smem_c] +
|
||||
smem_c[smem_c_index_func(m_idx, n_idx) + 2 * cosize_smem_c] +
|
||||
smem_c[smem_c_index_func(m_idx, n_idx) + 3 * cosize_smem_c];
|
||||
}
|
||||
|
||||
#pragma unroll
|
||||
for (int reg_idx = 0; reg_idx < final_acc_reg_cnt; reg_idx++) {
|
||||
int linear_idx = reg_idx * 32 + lane_idx;
|
||||
int m_idx = linear_idx % tile_m;
|
||||
int n_idx = linear_idx / tile_m;
|
||||
if (m_idx < tile_m && n_idx < gemm_n) {
|
||||
gmem_c[n_idx * gemm_m + m_idx] = acc_final[reg_idx];
|
||||
}
|
||||
}
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
bf16_t* gmem_c;
|
||||
bf16_t* smem_a;
|
||||
bf16_t* smem_b;
|
||||
uint64_t* smem_barrier;
|
||||
int warp_idx;
|
||||
int gemm_n;
|
||||
int stage_idx = 0;
|
||||
int phase_bit = 0;
|
||||
int lane_idx = threadIdx.x % 32;
|
||||
int warp_k_offset_in_tile_k = warp_idx * per_warp_tile_k;
|
||||
|
||||
int a_smem_offsets[m_iter_cnt][k_phase_cnt];
|
||||
int b_smem_offsets[n_iter_cnt][k_phase_cnt];
|
||||
|
||||
bf16_t a_reg[m_iter_cnt][k_phase_cnt][8];
|
||||
bf16_t b_reg[n_iter_cnt][k_phase_cnt][4];
|
||||
float acc_reg[m_iter_cnt][n_iter_cnt][4]{};
|
||||
};
|
||||
|
||||
template <int batch_size, int gemm_m, int gemm_k, int tile_m, int tile_n, int tile_k, int stage_cnt>
|
||||
__global__ __launch_bounds__(256, 1) void fused_a_gemm_kernel(
|
||||
bf16_t* output, bf16_t const* mat_a, bf16_t const* mat_b, int gemm_n) {
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
|
||||
constexpr int load_thread_cnt = 128;
|
||||
constexpr int compute_thread_cnt = 128;
|
||||
constexpr int thread_cnt = load_thread_cnt + compute_thread_cnt;
|
||||
(void)thread_cnt;
|
||||
static_assert(gemm_m % 16 == 0);
|
||||
static_assert(gemm_k % tile_k == 0);
|
||||
static_assert(gemm_m % tile_m == 0);
|
||||
static_assert(tile_k == 128 || tile_k == 256 || tile_k == 512 || tile_k == 1024);
|
||||
static_assert(tile_m == 16);
|
||||
constexpr int g2s_vec_bytes = 16;
|
||||
constexpr int a_elem_bytes = 2;
|
||||
constexpr int b_elem_bytes = 2;
|
||||
static_assert((tile_m * a_elem_bytes + tile_n * b_elem_bytes) * tile_k * stage_cnt <= 225 * 1024);
|
||||
static_assert((tile_m * tile_k * a_elem_bytes) % (load_thread_cnt * g2s_vec_bytes) == 0);
|
||||
static_assert((tile_n * tile_k * b_elem_bytes) % (load_thread_cnt * g2s_vec_bytes) == 0);
|
||||
|
||||
extern __shared__ char smem[];
|
||||
uint64_t* smem_barrier = reinterpret_cast<uint64_t*>(smem);
|
||||
bf16_t* smem_a = reinterpret_cast<bf16_t*>(smem + (stage_cnt * 8 * 2 + 1024) / 1024 * 1024);
|
||||
bf16_t* smem_b = smem_a + tile_m * tile_k * stage_cnt;
|
||||
|
||||
int cta_m_idx = tile_m * blockIdx.x;
|
||||
int cta_n_idx = tile_n * blockIdx.y;
|
||||
bf16_t const* gmem_a_local = mat_a + cta_m_idx * gemm_k;
|
||||
bf16_t const* gmem_b_local = mat_b + cta_n_idx * gemm_k;
|
||||
bf16_t* gmem_c_local = output + cta_n_idx * gemm_m + cta_m_idx;
|
||||
|
||||
int warp_idx = __shfl_sync(0xffffffff, threadIdx.x / 32, 0);
|
||||
|
||||
if (warp_idx == 4) {
|
||||
for (int i = 0; i < stage_cnt; i++) {
|
||||
initialize_barrier(smem_barrier + i * 2 + 0, load_thread_cnt);
|
||||
initialize_barrier(smem_barrier + i * 2 + 1, compute_thread_cnt);
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
if (warp_idx < 2) {
|
||||
GmemLoaderA<gemm_k, tile_m, tile_k, stage_cnt> a_loader(gmem_a_local, smem_a, smem_barrier);
|
||||
a_loader.prepare();
|
||||
a_loader.issue_mainloop();
|
||||
} else if (warp_idx < 4) {
|
||||
GmemLoaderB<gemm_k, tile_n, tile_k, stage_cnt> b_loader(gmem_b_local, smem_b, smem_barrier, gemm_n);
|
||||
b_loader.prepare();
|
||||
b_loader.issue_mainloop();
|
||||
} else {
|
||||
MmaComputer<gemm_m, gemm_k, tile_m, tile_n, tile_k, stage_cnt> mma_computer(
|
||||
gmem_c_local, smem_a, smem_b, smem_barrier, warp_idx, gemm_n);
|
||||
mma_computer.prepare();
|
||||
mma_computer.issue_mainloop();
|
||||
mma_computer.epi();
|
||||
}
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
#endif
|
||||
}
|
||||
|
||||
template <typename T, int kHdIn, int kHdOut, int kTileN, bool kUsePDL>
|
||||
void invokeFusedAGemm(T* output, T const* mat_a, T const* mat_b, int num_tokens, DLDevice device) {
|
||||
constexpr int gemm_m = kHdOut; // 2112
|
||||
int const gemm_n = num_tokens; // 16
|
||||
constexpr int gemm_k = kHdIn; // 7168
|
||||
constexpr int batch_size = 1;
|
||||
std::swap(mat_a, mat_b);
|
||||
constexpr int tile_m = 16;
|
||||
constexpr int tile_n = kTileN; // 8 or 16
|
||||
constexpr int tile_k = std::max(256, 1024 / tile_n); // 256
|
||||
#if defined(SGL_CUDA_ARCH) && SGL_CUDA_ARCH >= 1200
|
||||
constexpr int smem_stage_budget = 96 * 1024;
|
||||
#else
|
||||
constexpr int smem_stage_budget = 192 * 1024;
|
||||
#endif
|
||||
constexpr int max_stage_cnt = smem_stage_budget / ((tile_m + tile_n) * tile_k * sizeof(bf16_t));
|
||||
constexpr int k_iter_cnt = gemm_k / tile_k;
|
||||
constexpr int stage_cnt = k_iter_cnt > max_stage_cnt ? max_stage_cnt : k_iter_cnt;
|
||||
int cta_m_cnt = gemm_m / tile_m;
|
||||
int cta_n_cnt = (gemm_n + tile_n - 1) / tile_n;
|
||||
constexpr int barrier_bytes = (stage_cnt * 16 + 1023) / 1024 * 1024; // 4096
|
||||
constexpr int smem_bytes = ((tile_m * 2 + tile_n * 2) * tile_k * stage_cnt + barrier_bytes + 1023) / 1024 * 1024;
|
||||
|
||||
dim3 grid(cta_m_cnt, cta_n_cnt, 1);
|
||||
dim3 block_size(256);
|
||||
|
||||
auto kernel = fused_a_gemm_kernel<batch_size, gemm_m, gemm_k, tile_m, tile_n, tile_k, stage_cnt>;
|
||||
if (smem_bytes >= (48 * 1024)) {
|
||||
host::RuntimeDeviceCheck(cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_bytes));
|
||||
}
|
||||
host::LaunchKernel(grid, block_size, device, smem_bytes).enable_pdl(kUsePDL)(kernel, output, mat_a, mat_b, gemm_n);
|
||||
}
|
||||
|
||||
template <int kHdIn, int kHdOut, bool kUsePDL>
|
||||
struct DSV3FusedAGemmKernel {
|
||||
static void
|
||||
run(const tvm::ffi::TensorView mat_a, const tvm::ffi::TensorView mat_b, const tvm::ffi::TensorView output) {
|
||||
using namespace host;
|
||||
|
||||
auto M = SymbolicSize{"num_tokens"};
|
||||
auto K = SymbolicSize{"hd_in"};
|
||||
auto N = SymbolicSize{"hd_out"};
|
||||
auto device = SymbolicDevice{};
|
||||
K.set_value(kHdIn);
|
||||
N.set_value(kHdOut);
|
||||
device.set_options<kDLCUDA>();
|
||||
|
||||
// mat_a: [num_tokens, hd_in] row-major; output: [num_tokens, hd_out] row-major;
|
||||
// mat_b: [hd_in, hd_out] column-major (weight.T), i.e. strides {1, hd_in}.
|
||||
TensorMatcher({M, K}).with_dtype<bf16_t>().with_device(device).verify(mat_a);
|
||||
TensorMatcher({K, N}).with_dtype<bf16_t>().with_device(device).with_strides({1, kHdIn}).verify(mat_b);
|
||||
TensorMatcher({M, N}).with_dtype<bf16_t>().with_device(device).verify(output);
|
||||
|
||||
const int num_tokens = static_cast<int>(M.unwrap());
|
||||
RuntimeCheck(
|
||||
num_tokens >= 1 && num_tokens <= 16, "dsv3_fused_a_gemm: num_tokens must be in [1, 16], got ", num_tokens);
|
||||
|
||||
const DLDevice dev = device.unwrap();
|
||||
auto* out_ptr = static_cast<bf16_t*>(output.data_ptr());
|
||||
auto* a_ptr = static_cast<bf16_t const*>(mat_a.data_ptr());
|
||||
auto* b_ptr = static_cast<bf16_t const*>(mat_b.data_ptr());
|
||||
|
||||
if (num_tokens <= 8) {
|
||||
invokeFusedAGemm<bf16_t, kHdIn, kHdOut, 8, kUsePDL>(out_ptr, a_ptr, b_ptr, num_tokens, dev);
|
||||
} else {
|
||||
invokeFusedAGemm<bf16_t, kHdIn, kHdOut, 16, kUsePDL>(out_ptr, a_ptr, b_ptr, num_tokens, dev);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace
|
||||
@@ -0,0 +1,384 @@
|
||||
# Copyright (c) 2019-2024, NVIDIA CORPORATION. All rights reserved.
|
||||
# Copyright 2023-2024 SGLang Team
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""CuTe DSL DeepSeek-V3 fused-A GEMM (sm90+): out[M, N] = mat_a[M, K] @ weight,
|
||||
N in {2112, 6144}, M = num_tokens in [1, 16], K any multiple of 1024, bf16.
|
||||
|
||||
Adapted from NVIDIA TensorRT-LLM dsv3FusedAGemm.cu
|
||||
(cpp/tensorrt_llm/kernels/dsv3MinLatencyKernels/dsv3FusedAGemm.cu), reimplemented
|
||||
in the CuTe DSL: AB-swap, warp-specialized 4-way split-K, cp.async + mbarrier
|
||||
pipeline, ldmatrix + mma.sync.m16n8k16, 3-4-3 swizzle.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import cuda.bindings.driver as cuda
|
||||
import cutlass
|
||||
import cutlass.cute as cute
|
||||
import torch
|
||||
from cutlass._mlir import ir
|
||||
from cutlass._mlir.dialects import llvm
|
||||
from cutlass.cute.runtime import from_dlpack
|
||||
from cutlass.utils import get_smem_capacity_in_bytes
|
||||
|
||||
from sglang.kernel_api_logging import debug_kernel_api
|
||||
from sglang.srt.utils import get_device_sm
|
||||
from sglang.srt.utils.common import direct_register_custom_op
|
||||
|
||||
TILE_M = 16
|
||||
TILE_K = 256
|
||||
SPLITK = 4
|
||||
LOAD_WARPS = 4
|
||||
MAX_NSTAGE = 16
|
||||
PWK = TILE_K // SPLITK
|
||||
KSTEPS = PWK // 16
|
||||
COMPUTE_THREADS = SPLITK * 32
|
||||
LOADER_THREADS = LOAD_WARPS * 32
|
||||
NTHREADS = COMPUTE_THREADS + LOADER_THREADS
|
||||
KI = TILE_K // 2
|
||||
|
||||
_BAR_I32 = 2 * MAX_NSTAGE * 2
|
||||
|
||||
|
||||
def _stage_i32(tile_n: int) -> int:
|
||||
return (TILE_M + tile_n) * KI
|
||||
|
||||
|
||||
def _cp_async_16b(smem_ptr, gmem_ptr):
|
||||
llvm.inline_asm(
|
||||
None,
|
||||
[smem_ptr.toint().ir_value(), gmem_ptr.toint().ir_value()],
|
||||
"{ .reg .u32 sa; cvt.u32.u64 sa, $0; cp.async.cg.shared.global.L2::128B [sa], [$1], 16; }",
|
||||
"l,l",
|
||||
has_side_effects=True,
|
||||
is_align_stack=False,
|
||||
asm_dialect=0,
|
||||
)
|
||||
|
||||
|
||||
def _cp_async_16b_pred(smem_ptr, gmem_ptr, pred_i32):
|
||||
llvm.inline_asm(
|
||||
None,
|
||||
[smem_ptr.toint().ir_value(), gmem_ptr.toint().ir_value(), pred_i32.ir_value()],
|
||||
"{ .reg .pred p; .reg .u32 sa; setp.ne.s32 p, $2, 0; cvt.u32.u64 sa, $0; "
|
||||
"@p cp.async.cg.shared.global.L2::128B [sa], [$1], 16; }",
|
||||
"l,l,r",
|
||||
has_side_effects=True,
|
||||
is_align_stack=False,
|
||||
asm_dialect=0,
|
||||
)
|
||||
|
||||
|
||||
def _ldmatrix_x4(smem_ptr):
|
||||
i32 = ir.IntegerType.get_signless(32)
|
||||
res = llvm.inline_asm(
|
||||
llvm.StructType.get_literal([i32, i32, i32, i32]),
|
||||
[smem_ptr.toint().ir_value()],
|
||||
"{ .reg .u32 sa; cvt.u32.u64 sa, $4; "
|
||||
"ldmatrix.sync.aligned.x4.m8n8.shared.b16 {$0,$1,$2,$3}, [sa]; }",
|
||||
"=r,=r,=r,=r,l",
|
||||
has_side_effects=True,
|
||||
is_align_stack=False,
|
||||
asm_dialect=0,
|
||||
)
|
||||
return [llvm.extractvalue(i32, res, [i]) for i in range(4)]
|
||||
|
||||
|
||||
def _ldmatrix_x2(smem_ptr):
|
||||
i32 = ir.IntegerType.get_signless(32)
|
||||
res = llvm.inline_asm(
|
||||
llvm.StructType.get_literal([i32, i32]),
|
||||
[smem_ptr.toint().ir_value()],
|
||||
"{ .reg .u32 sa; cvt.u32.u64 sa, $2; "
|
||||
"ldmatrix.sync.aligned.x2.m8n8.shared.b16 {$0,$1}, [sa]; }",
|
||||
"=r,=r,l",
|
||||
has_side_effects=True,
|
||||
is_align_stack=False,
|
||||
asm_dialect=0,
|
||||
)
|
||||
return [llvm.extractvalue(i32, res, [i]) for i in range(2)]
|
||||
|
||||
|
||||
def _mma_m16n8k16(a0, a1, a2, a3, b0, b1, c0, c1, c2, c3):
|
||||
f32 = ir.F32Type.get()
|
||||
res = llvm.inline_asm(
|
||||
llvm.StructType.get_literal([f32, f32, f32, f32]),
|
||||
[a0, a1, a2, a3, b0, b1, c0, c1, c2, c3],
|
||||
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 "
|
||||
"{$0,$1,$2,$3}, {$4,$5,$6,$7}, {$8,$9}, {$10,$11,$12,$13};",
|
||||
"=f,=f,=f,=f,r,r,r,r,r,r,f,f,f,f",
|
||||
has_side_effects=False,
|
||||
is_align_stack=False,
|
||||
asm_dialect=0,
|
||||
)
|
||||
return [llvm.extractvalue(f32, res, [i]) for i in range(4)]
|
||||
|
||||
|
||||
def _swizzle_343(row, col):
|
||||
return col ^ ((row % 8) * 4)
|
||||
|
||||
|
||||
def _global_k_col(tile, col, kgi):
|
||||
kp_warp = KI // SPLITK
|
||||
kp_chunk = kgi // SPLITK
|
||||
return (col // kp_warp) * kp_chunk + tile * kp_warp + (col % kp_warp)
|
||||
|
||||
|
||||
def _load_stage(ltid, feat0, sa, sb, mW, mA, tile, buf, M, kgi, tile_n):
|
||||
for it in range(TILE_M * KI // (LOADER_THREADS * 4)):
|
||||
idx = (it * LOADER_THREADS + ltid) * 4
|
||||
row, col = idx // KI, idx % KI
|
||||
_cp_async_16b(
|
||||
sa.iterator + (buf * TILE_M * KI + row * KI + _swizzle_343(row, col)),
|
||||
mW.iterator + ((feat0 + row) * kgi + _global_k_col(tile, col, kgi)),
|
||||
)
|
||||
for it in range(tile_n * KI // (LOADER_THREADS * 4)):
|
||||
idx = (it * LOADER_THREADS + ltid) * 4
|
||||
row, col = idx // KI, idx % KI
|
||||
pred = (row < M).to(cutlass.Int32)
|
||||
_cp_async_16b_pred(
|
||||
sb.iterator + (buf * tile_n * KI + row * KI + _swizzle_343(row, col)),
|
||||
mA.iterator + (row * pred * kgi + _global_k_col(tile, col, kgi)),
|
||||
pred,
|
||||
)
|
||||
|
||||
|
||||
@cute.kernel
|
||||
def _dsv3_fused_a_gemm_kernel(
|
||||
mW: cute.Tensor,
|
||||
mA: cute.Tensor,
|
||||
mOut: cute.Tensor,
|
||||
M: cutlass.Int32,
|
||||
num_kt: cutlass.Constexpr,
|
||||
nstage: cutlass.Constexpr,
|
||||
tile_n: cutlass.Constexpr,
|
||||
):
|
||||
NB = tile_n // 8
|
||||
tid, _, _ = cute.arch.thread_idx()
|
||||
bid, _, _ = cute.arch.block_idx()
|
||||
warp, lane = tid // 32, tid % 32
|
||||
r0, cc = lane // 4, lane % 4
|
||||
feat0 = bid * TILE_M
|
||||
kgi = num_kt * KI
|
||||
|
||||
base = cute.arch.get_dyn_smem(cutlass.Int32, alignment=16)
|
||||
bar = cute.recast_ptr(base, dtype=cutlass.Int64)
|
||||
full, empty = bar, bar + MAX_NSTAGE
|
||||
sa_off = _BAR_I32
|
||||
sc_warp_stride = TILE_M * tile_n + 2
|
||||
sC = cute.make_tensor(
|
||||
cute.recast_ptr(base + sa_off, dtype=cutlass.Float32),
|
||||
cute.make_layout((SPLITK, TILE_M, tile_n), stride=(sc_warp_stride, tile_n, 1)),
|
||||
)
|
||||
sA = cute.make_tensor(
|
||||
base + sa_off,
|
||||
cute.make_layout((nstage, TILE_M, KI), stride=(TILE_M * KI, KI, 1)),
|
||||
)
|
||||
sB = cute.make_tensor(
|
||||
base + sa_off + nstage * TILE_M * KI,
|
||||
cute.make_layout((nstage, tile_n, KI), stride=(tile_n * KI, KI, 1)),
|
||||
)
|
||||
|
||||
if tid == 0:
|
||||
for s in range(nstage):
|
||||
cute.arch.mbarrier_init(full + s, LOADER_THREADS)
|
||||
cute.arch.mbarrier_init(empty + s, COMPUTE_THREADS)
|
||||
cute.arch.barrier()
|
||||
|
||||
if warp >= SPLITK:
|
||||
cute.arch.griddepcontrol_wait()
|
||||
ltid = tid - COMPUTE_THREADS
|
||||
for kt in cutlass.range_constexpr(num_kt):
|
||||
st = kt % nstage
|
||||
if kt >= nstage:
|
||||
cute.arch.mbarrier_wait(empty + st, ((kt // nstage) & 1) ^ 1)
|
||||
_load_stage(ltid, feat0, sA, sB, mW, mA, kt, st, M, kgi, tile_n)
|
||||
cute.arch.cp_async_mbarrier_arrive_noinc(full + st)
|
||||
else:
|
||||
acc = [[cutlass.Float32(0.0) for _ in range(4)] for _ in range(NB)]
|
||||
ph = 0
|
||||
for kt in cutlass.range_constexpr(num_kt):
|
||||
buf = kt % nstage
|
||||
cute.arch.mbarrier_wait(full + buf, ph)
|
||||
brow_lo = lane % 8
|
||||
boff = ((lane // 8) & 1) * 4
|
||||
for step in cutlass.range_constexpr(KSTEPS):
|
||||
kbh = warp * (PWK // 2) + step * 8
|
||||
arow, aoff = lane % 16, (lane // 16) * 4
|
||||
a = _ldmatrix_x4(
|
||||
sA.iterator
|
||||
+ (buf * TILE_M * KI + arow * KI + _swizzle_343(arow, kbh + aoff))
|
||||
)
|
||||
for nb in cutlass.range_constexpr(NB):
|
||||
brow = nb * 8 + brow_lo
|
||||
bb = _ldmatrix_x2(
|
||||
sB.iterator
|
||||
+ (
|
||||
buf * tile_n * KI
|
||||
+ brow * KI
|
||||
+ _swizzle_343(brow, kbh + boff)
|
||||
)
|
||||
)
|
||||
d = _mma_m16n8k16(
|
||||
a[0],
|
||||
a[1],
|
||||
a[2],
|
||||
a[3],
|
||||
bb[0],
|
||||
bb[1],
|
||||
acc[nb][0].ir_value(),
|
||||
acc[nb][1].ir_value(),
|
||||
acc[nb][2].ir_value(),
|
||||
acc[nb][3].ir_value(),
|
||||
)
|
||||
for i in cutlass.range_constexpr(4):
|
||||
acc[nb][i] = cutlass.Float32(d[i])
|
||||
cute.arch.mbarrier_arrive(empty + buf)
|
||||
ph = (ph ^ 1) if buf == nstage - 1 else ph
|
||||
for nb in cutlass.range_constexpr(NB):
|
||||
for i in cutlass.range_constexpr(4):
|
||||
m = r0 + (8 if i >= 2 else 0)
|
||||
n = nb * 8 + cc * 2 + (i % 2)
|
||||
sC[warp, m, n] = acc[nb][i]
|
||||
|
||||
cute.arch.barrier(barrier_id=1, number_of_threads=COMPUTE_THREADS)
|
||||
nred = TILE_M * tile_n
|
||||
for it in cutlass.range_constexpr(
|
||||
(nred + COMPUTE_THREADS - 1) // COMPUTE_THREADS
|
||||
):
|
||||
e = it * COMPUTE_THREADS + tid
|
||||
if e < nred:
|
||||
m, n = e // tile_n, e % tile_n
|
||||
s = sC[0, m, n]
|
||||
for w in cutlass.range_constexpr(1, SPLITK):
|
||||
s = s + sC[w, m, n]
|
||||
if n < M:
|
||||
mOut[n, feat0 + m] = s.to(cutlass.BFloat16)
|
||||
|
||||
cute.arch.griddepcontrol_launch_dependents()
|
||||
|
||||
|
||||
@cute.jit
|
||||
def _dsv3_fused_a_gemm_host(
|
||||
mW: cute.Tensor,
|
||||
mA: cute.Tensor,
|
||||
mOut: cute.Tensor,
|
||||
M: cutlass.Int32,
|
||||
stream: cuda.CUstream,
|
||||
num_kt: cutlass.Constexpr,
|
||||
gemm_m: cutlass.Constexpr,
|
||||
smem_bytes: cutlass.Constexpr,
|
||||
nstage: cutlass.Constexpr,
|
||||
tile_n: cutlass.Constexpr,
|
||||
):
|
||||
_dsv3_fused_a_gemm_kernel(mW, mA, mOut, M, num_kt, nstage, tile_n).launch(
|
||||
grid=[gemm_m // TILE_M, 1, 1],
|
||||
block=[NTHREADS, 1, 1],
|
||||
max_number_threads=[NTHREADS, 1, 1],
|
||||
min_blocks_per_mp=1,
|
||||
smem=smem_bytes,
|
||||
use_pdl=True,
|
||||
stream=stream,
|
||||
)
|
||||
|
||||
|
||||
_compiled: dict[tuple[int, int, int], object] = {}
|
||||
|
||||
|
||||
def _pick_nstage(num_kt: int, tile_n: int) -> int:
|
||||
nstage = (get_smem_capacity_in_bytes() // 4 - _BAR_I32) // _stage_i32(tile_n)
|
||||
return min(nstage, MAX_NSTAGE, num_kt)
|
||||
|
||||
|
||||
def _pick_tile_n(num_tokens: int) -> int:
|
||||
return 8 if num_tokens <= 8 else 16
|
||||
|
||||
|
||||
def _compiled_kernel(num_kt: int, gemm_m: int, tile_n: int):
|
||||
if get_device_sm() < 90:
|
||||
raise RuntimeError("dsv3_fused_a_gemm requires SM90 (Hopper) or later")
|
||||
if (num_kt, gemm_m, tile_n) not in _compiled:
|
||||
nstage = _pick_nstage(num_kt, tile_n)
|
||||
smem_bytes = (_BAR_I32 + nstage * _stage_i32(tile_n)) * 4
|
||||
k = num_kt * TILE_K
|
||||
w = torch.empty(gemm_m, k, dtype=torch.bfloat16, device="cuda")
|
||||
a = torch.empty(16, k, dtype=torch.bfloat16, device="cuda")
|
||||
o = torch.empty(16, gemm_m, dtype=torch.bfloat16, device="cuda")
|
||||
stream = cuda.CUstream(torch.cuda.current_stream().cuda_stream)
|
||||
_compiled[(num_kt, gemm_m, tile_n)] = cute.compile(
|
||||
_dsv3_fused_a_gemm_host,
|
||||
from_dlpack(w.view(torch.int32)),
|
||||
from_dlpack(a.view(torch.int32)),
|
||||
from_dlpack(o),
|
||||
cutlass.Int32(16),
|
||||
stream,
|
||||
num_kt,
|
||||
gemm_m,
|
||||
smem_bytes,
|
||||
nstage,
|
||||
tile_n,
|
||||
)
|
||||
return _compiled[(num_kt, gemm_m, tile_n)]
|
||||
|
||||
|
||||
def _dsv3_fused_a_gemm_run(mat_a: torch.Tensor, mat_b: torch.Tensor) -> torch.Tensor:
|
||||
M, K = mat_a.shape
|
||||
N = mat_b.shape[1]
|
||||
assert mat_a.dtype == torch.bfloat16 and mat_b.dtype == torch.bfloat16
|
||||
assert K % 1024 == 0, f"K must be a multiple of 1024, got {K}"
|
||||
assert N % TILE_M == 0, f"N must be a multiple of {TILE_M}, got {N}"
|
||||
assert (
|
||||
tuple(mat_b.shape) == (K, N) and mat_b.stride(0) == 1
|
||||
), "mat_b must be [K, N] column-major"
|
||||
assert 1 <= M <= 16, "num_tokens must be in [1, 16]"
|
||||
assert mat_a.stride(1) == 1, "mat_a must be row-major [M, K]"
|
||||
|
||||
weight = mat_b.t()
|
||||
out = torch.empty(M, N, dtype=torch.bfloat16, device=mat_a.device)
|
||||
|
||||
stream = cuda.CUstream(torch.cuda.current_stream().cuda_stream)
|
||||
_compiled_kernel(K // TILE_K, N, _pick_tile_n(M))(
|
||||
from_dlpack(weight.view(torch.int32)),
|
||||
from_dlpack(mat_a.view(torch.int32)),
|
||||
from_dlpack(out),
|
||||
M,
|
||||
stream,
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
def _dsv3_fused_a_gemm_fake(mat_a: torch.Tensor, mat_b: torch.Tensor) -> torch.Tensor:
|
||||
return mat_a.new_empty((mat_a.shape[0], mat_b.shape[1]), dtype=torch.bfloat16)
|
||||
|
||||
|
||||
direct_register_custom_op(
|
||||
op_name="cutedsl_dsv3_fused_a_gemm",
|
||||
op_func=_dsv3_fused_a_gemm_run,
|
||||
mutates_args=[],
|
||||
fake_impl=_dsv3_fused_a_gemm_fake,
|
||||
)
|
||||
|
||||
|
||||
@debug_kernel_api
|
||||
def dsv3_fused_a_gemm(
|
||||
mat_a: torch.Tensor, mat_b: torch.Tensor, output: torch.Tensor | None = None
|
||||
) -> torch.Tensor:
|
||||
"""out[M, N] = mat_a[M, K] @ mat_b, with mat_a row-major [M, K] (M in [1, 16]),
|
||||
mat_b column-major [K, N] (the weight, stride(0) == 1), N a multiple of 16
|
||||
(e.g. 2112, 6144), K a multiple of 1024."""
|
||||
result = torch.ops.sglang.cutedsl_dsv3_fused_a_gemm(mat_a, mat_b)
|
||||
if output is not None:
|
||||
output.copy_(result)
|
||||
return output
|
||||
return result
|
||||
@@ -0,0 +1,90 @@
|
||||
"""
|
||||
JIT kernel for DeepSeek V3 fused QKV-A GEMM (min-latency).
|
||||
|
||||
Replaces the AOT sgl_kernel.dsv3_fused_a_gemm for SM90+ (Hopper) GPUs.
|
||||
Shapes: hd_in a multiple of 256, hd_out a multiple of 16, num_tokens 1-16, bfloat16.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.jit_kernel.utils import (
|
||||
cache_once,
|
||||
is_arch_support_pdl,
|
||||
load_jit,
|
||||
make_cpp_args,
|
||||
)
|
||||
from sglang.kernel_api_logging import debug_kernel_api
|
||||
from sglang.srt.utils.common import direct_register_custom_op
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from tvm_ffi.module import Module
|
||||
|
||||
|
||||
@cache_once
|
||||
def _jit_dsv3_fused_a_gemm_module(hd_in: int, hd_out: int, use_pdl: bool) -> Module:
|
||||
args = make_cpp_args(hd_in, hd_out, use_pdl)
|
||||
return load_jit(
|
||||
"dsv3_fused_a_gemm",
|
||||
*args,
|
||||
cuda_files=["gemm/dsv3_fused_a_gemm.cuh"],
|
||||
cuda_wrappers=[
|
||||
("dsv3_fused_a_gemm", f"DSV3FusedAGemmKernel<{args}>::run"),
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
def _dsv3_fused_a_gemm_run(mat_a: torch.Tensor, mat_b: torch.Tensor) -> torch.Tensor:
|
||||
assert mat_a.stride(1) == 1, "mat_a must be row-major [M, K]"
|
||||
output = torch.empty(
|
||||
(mat_a.shape[0], mat_b.shape[1]),
|
||||
device=mat_a.device,
|
||||
dtype=mat_a.dtype,
|
||||
)
|
||||
module = _jit_dsv3_fused_a_gemm_module(
|
||||
mat_a.shape[1], mat_b.shape[1], is_arch_support_pdl()
|
||||
)
|
||||
module.dsv3_fused_a_gemm(mat_a, mat_b, output)
|
||||
return output
|
||||
|
||||
|
||||
def _dsv3_fused_a_gemm_fake(mat_a: torch.Tensor, mat_b: torch.Tensor) -> torch.Tensor:
|
||||
return mat_a.new_empty((mat_a.shape[0], mat_b.shape[1]), dtype=torch.bfloat16)
|
||||
|
||||
|
||||
direct_register_custom_op(
|
||||
op_name="jit_dsv3_fused_a_gemm",
|
||||
op_func=_dsv3_fused_a_gemm_run,
|
||||
mutates_args=[],
|
||||
fake_impl=_dsv3_fused_a_gemm_fake,
|
||||
)
|
||||
|
||||
|
||||
@debug_kernel_api
|
||||
def dsv3_fused_a_gemm(
|
||||
mat_a: torch.Tensor,
|
||||
mat_b: torch.Tensor,
|
||||
output: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
DeepSeek V3 fused QKV-A GEMM kernel (JIT variant).
|
||||
|
||||
Args:
|
||||
mat_a: Input tensor of shape [num_tokens, hd_in], bfloat16, row-major.
|
||||
hd_in must be a multiple of 256 and num_tokens in [1, 16].
|
||||
mat_b: Weight tensor of shape [hd_in, hd_out], bfloat16, column-major
|
||||
(i.e. ``weight.T`` of a row-major [hd_out, hd_in] weight).
|
||||
hd_out must be a multiple of 16.
|
||||
output: Optional pre-allocated output tensor of shape [num_tokens, hd_out].
|
||||
|
||||
Returns:
|
||||
Output tensor of shape [num_tokens, hd_out].
|
||||
"""
|
||||
result = torch.ops.sglang.jit_dsv3_fused_a_gemm(mat_a, mat_b)
|
||||
if output is not None:
|
||||
output.copy_(result)
|
||||
return output
|
||||
return result
|
||||
@@ -0,0 +1,52 @@
|
||||
"""Unified entry point for the DeepSeek-V3 fused QKV-A GEMM.
|
||||
|
||||
Dispatches to one of three interchangeable implementations via ``backend``:
|
||||
|
||||
- ``"aot"``: prebuilt ``sgl_kernel.dsv3_fused_a_gemm`` (CUDA C++).
|
||||
- ``"jit"``: runtime-compiled CUDA C++ (``sglang.jit_kernel.dsv3_fused_a_gemm``).
|
||||
- ``"cutedsl"``: CuTe DSL (``sglang.jit_kernel.cutedsl_dsv3_fused_a_gemm``).
|
||||
- ``"auto"``: CuTe DSL on SM120+, otherwise the JIT kernel.
|
||||
|
||||
All backends share the signature ``(mat_a, mat_b, output=None) -> Tensor`` with
|
||||
``mat_a`` row-major ``[M, K]`` (M in [1, 16], bf16) and ``mat_b`` the column-major
|
||||
weight ``[K, N]`` (``weight.T``).
|
||||
"""
|
||||
|
||||
from enum import Enum
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.utils.common import is_sm120_supported
|
||||
|
||||
|
||||
class FusedAGemmBackend(str, Enum):
|
||||
AUTO = "auto"
|
||||
AOT = "aot"
|
||||
JIT = "jit"
|
||||
CUTEDSL = "cutedsl"
|
||||
|
||||
|
||||
_AUTO_BACKEND = (
|
||||
FusedAGemmBackend.CUTEDSL if is_sm120_supported() else FusedAGemmBackend.JIT
|
||||
)
|
||||
|
||||
|
||||
def dsv3_fused_a_gemm(
|
||||
mat_a: torch.Tensor,
|
||||
mat_b: torch.Tensor,
|
||||
output: torch.Tensor | None = None,
|
||||
backend: FusedAGemmBackend | str = FusedAGemmBackend.AUTO,
|
||||
) -> torch.Tensor:
|
||||
backend = FusedAGemmBackend(backend)
|
||||
if backend == FusedAGemmBackend.AUTO:
|
||||
backend = _AUTO_BACKEND
|
||||
|
||||
if backend == FusedAGemmBackend.AOT:
|
||||
from sgl_kernel import dsv3_fused_a_gemm as impl
|
||||
elif backend == FusedAGemmBackend.JIT:
|
||||
from sglang.jit_kernel.dsv3_fused_a_gemm import dsv3_fused_a_gemm as impl
|
||||
else:
|
||||
from sglang.jit_kernel.cutedsl_dsv3_fused_a_gemm import (
|
||||
dsv3_fused_a_gemm as impl,
|
||||
)
|
||||
return impl(mat_a, mat_b, output)
|
||||
@@ -201,11 +201,10 @@ if _use_aiter:
|
||||
pass
|
||||
|
||||
if _is_cuda:
|
||||
from sgl_kernel import dsv3_fused_a_gemm
|
||||
|
||||
from sglang.jit_kernel.dsv3_router_gemm import (
|
||||
dsv3_router_gemm as _jit_dsv3_router_gemm,
|
||||
)
|
||||
from sglang.jit_kernel.fused_a_gemm import dsv3_fused_a_gemm
|
||||
elif _is_npu:
|
||||
from sglang.srt.hardware_backend.npu.modules.deepseek_v2_attention_mla_npu import (
|
||||
forward_dsa_core_npu,
|
||||
@@ -1769,11 +1768,12 @@ class DeepseekV2AttentionMLA(
|
||||
self.has_fused_proj
|
||||
and not self.is_packed_weight
|
||||
and self.fused_qkv_a_proj_with_mqa.weight.dtype == torch.bfloat16
|
||||
and self.fused_qkv_a_proj_with_mqa.weight.shape[0] == 2112
|
||||
and self.fused_qkv_a_proj_with_mqa.weight.shape[1] == 7168
|
||||
and self.fused_qkv_a_proj_with_mqa.weight.shape[0] % 16 == 0
|
||||
and self.fused_qkv_a_proj_with_mqa.weight.shape[1] % 256 == 0
|
||||
and _is_cuda
|
||||
and 90 <= _device_sm < 120
|
||||
and _device_sm >= 90
|
||||
)
|
||||
self.fused_a_gemm_backend = "auto"
|
||||
|
||||
self.init_mha_forward()
|
||||
self.init_mla_forward()
|
||||
@@ -1987,7 +1987,9 @@ class DeepseekV2AttentionMLA(
|
||||
and not lora_active
|
||||
):
|
||||
qkv_latent = dsv3_fused_a_gemm(
|
||||
hidden_states, self.fused_qkv_a_proj_with_mqa.weight.T
|
||||
hidden_states,
|
||||
self.fused_qkv_a_proj_with_mqa.weight.T,
|
||||
backend=self.fused_a_gemm_backend,
|
||||
)
|
||||
else:
|
||||
qkv_latent = self.fused_qkv_a_proj_with_mqa(hidden_states)[0]
|
||||
|
||||
Reference in New Issue
Block a user