refactor: hgemm_blocktiling.cu — strict 1:1 from siboehm kernel 6
Only 3 changes from upstream_ref/sgemm_cuda/6_kernel_vectorize.cuh: 1. float → __half for A/B/C data and shared memory 2. float4 vectorized load → 4 scalar half loads (float4 needs 16-byte align) 3. threadResults accumulator stays float (FP32 accumulation) Everything else identical: same shared mem layout, same indexing, same A-transpose-while-loading, same thread tile computation. No WARPSIZE. No cooperative_groups. No cuda::barrier.
This commit is contained in:
@@ -1,262 +1,167 @@
|
||||
// hgemm_blocktiling.cu — FP16 GEMM kernel for BI-V100 (ivcore10)
|
||||
// hgemm_blocktiling.cu — FP16 GEMM for BI-V100
|
||||
//
|
||||
// Adapted from siboehm/SGEMM_CUDA kernel 6 (sgemmVectorize)
|
||||
// and wangzyon/NVIDIA_SGEMM_PRACTICE kernel 6 (mysgemm_v6).
|
||||
//
|
||||
// Key adaptations for BI-V100:
|
||||
// - FP16 (__half) data type with FP32 accumulation
|
||||
// - No WARPSIZE dependency (kernels 1-9 don't use it)
|
||||
// - Uses half2 vectorized loads (4 bytes) instead of float4 (16 bytes)
|
||||
// - Shared memory: BI-V100 has 128KB per block (vs 48KB on V100)
|
||||
// - Boundary checks for non-aligned M/N/K (MoE expert sizes vary)
|
||||
//
|
||||
// This kernel is used for MoE expert GEMM where each expert has different
|
||||
// token counts (non-uniform M). cublas batched GEMM requires uniform M
|
||||
// across the batch, so we need a custom kernel for the prefill path.
|
||||
//
|
||||
// For decode path (M=1 per expert), use cublasHgemmStridedBatched instead.
|
||||
// 1:1 from siboehm/SGEMM_CUDA kernel 6 (sgemmVectorize).
|
||||
// Changes: float→__half, float4→load 4 halfs, FP32 accumulator.
|
||||
// No WARPSIZE usage. No cooperative_groups. CUDA 10.2 safe.
|
||||
|
||||
#include <cuda_fp16.h>
|
||||
#include <cuda_runtime.h>
|
||||
#include <cstdint>
|
||||
|
||||
#define CEIL_DIV(M, N) (((M) + (N)-1) / (N))
|
||||
#define OFFSET(row, col, ld) ((row)*(ld)+(col))
|
||||
|
||||
// ============================================================================
|
||||
// Kernel: FP16 2D block tiling with A transpose and vectorized loads
|
||||
// ============================================================================
|
||||
// Based on siboehm kernel 6 / wangzyon kernel 6.
|
||||
// FP32 accumulation to avoid FP16 precision loss.
|
||||
//
|
||||
// Template params:
|
||||
// BM, BN: block tile size (rows of C, cols of C)
|
||||
// BK: block tile K dimension
|
||||
// TM, TN: per-thread tile size
|
||||
template<const int BM, const int BN, const int BK, const int TM, const int TN>
|
||||
__global__ void hgemm_blocktiling_v6(
|
||||
int M, int N, int K,
|
||||
__half alpha_h,
|
||||
const __half* __restrict__ A, // (M, K) row-major
|
||||
const __half* __restrict__ B, // (K, N) row-major
|
||||
__half beta_h,
|
||||
__half* __restrict__ C // (M, N) row-major
|
||||
) {
|
||||
int bx = blockIdx.x;
|
||||
int by = blockIdx.y;
|
||||
template <const int BM, const int BN, const int BK, const int TM, const int TN>
|
||||
__global__ void hgemmVectorize(int M, int N, int K, float alpha,
|
||||
const __half *A, const __half *B,
|
||||
float beta, __half *C) {
|
||||
const uint cRow = blockIdx.y;
|
||||
const uint cCol = blockIdx.x;
|
||||
|
||||
const int block_row_thread = BN / TN;
|
||||
const int block_col_thread = BM / TM;
|
||||
const int thread_num = block_row_thread * block_col_thread;
|
||||
// BN/TN are the number of threads to span a column
|
||||
const int threadCol = threadIdx.x % (BN / TN);
|
||||
const int threadRow = threadIdx.x / (BN / TN);
|
||||
|
||||
int tx = (threadIdx.x % block_row_thread) * TN;
|
||||
int ty = (threadIdx.x / block_row_thread) * TM;
|
||||
// allocate space for the current blocktile in smem
|
||||
// A stored transposed: As[BK][BM], B normal: Bs[BK][BN]
|
||||
__shared__ __half As[BM * BK];
|
||||
__shared__ __half Bs[BK * BN];
|
||||
|
||||
// Shared memory: A is stored transposed for vectorized reads
|
||||
__shared__ __half As[BK * BM]; // transposed: As[k][m]
|
||||
__shared__ __half Bs[BK * BN]; // normal: Bs[k][n]
|
||||
// Move blocktile to beginning of A's row and B's column
|
||||
A += cRow * BM * K;
|
||||
B += cCol * BN;
|
||||
C += cRow * BM * N + cCol * BN;
|
||||
|
||||
// Each thread loads multiple elements per round
|
||||
// For FP16, we load 4 halfs (8 bytes) at a time via half2 pairs
|
||||
const int ldg_a_num = BK * BM / thread_num / 4;
|
||||
const int ldg_b_num = BK * BN / thread_num / 4;
|
||||
// calculating the indices that this thread will load into SMEM
|
||||
// FP16: load 4 halfs (8 bytes) per step. 4 halfs per thread.
|
||||
// siboehm: float4 = 4 floats = 128bit. We do 4 halfs = 64bit.
|
||||
const uint innerRowA = threadIdx.x / (BK / 4);
|
||||
const uint innerColA = threadIdx.x % (BK / 4);
|
||||
const uint innerRowB = threadIdx.x / (BN / 4);
|
||||
const uint innerColB = threadIdx.x % (BN / 4);
|
||||
|
||||
int a_tile_row = threadIdx.x / (BK / 4);
|
||||
int a_tile_col = threadIdx.x % (BK / 4) * 4;
|
||||
int a_tile_stride = BM / ldg_a_num;
|
||||
// allocate thread-local cache for results in registerfile
|
||||
// FP32 accumulation to avoid FP16 precision loss
|
||||
float threadResults[TM * TN] = {0.0f};
|
||||
__half regM[TM];
|
||||
__half regN[TN];
|
||||
|
||||
int b_tile_row = threadIdx.x / (BN / 4);
|
||||
int b_tile_col = threadIdx.x % (BN / 4) * 4;
|
||||
int b_tile_stride = BK / ldg_b_num;
|
||||
// outer-most loop over block tiles
|
||||
for (uint bkIdx = 0; bkIdx < K; bkIdx += BK) {
|
||||
// populate the SMEM caches
|
||||
// transpose A while loading it (same as siboehm)
|
||||
// Load 4 halfs from A
|
||||
__half a0 = A[innerRowA * K + innerColA * 4 + 0];
|
||||
__half a1 = A[innerRowA * K + innerColA * 4 + 1];
|
||||
__half a2 = A[innerRowA * K + innerColA * 4 + 2];
|
||||
__half a3 = A[innerRowA * K + innerColA * 4 + 3];
|
||||
As[(innerColA * 4 + 0) * BM + innerRowA] = a0;
|
||||
As[(innerColA * 4 + 1) * BM + innerRowA] = a1;
|
||||
As[(innerColA * 4 + 2) * BM + innerRowA] = a2;
|
||||
As[(innerColA * 4 + 3) * BM + innerRowA] = a3;
|
||||
|
||||
// FP32 accumulators to avoid precision loss
|
||||
float accum[TM][TN] = {0.0f};
|
||||
// Load 4 halfs from B (no transpose)
|
||||
Bs[innerRowB * BN + innerColB * 4 + 0] = B[innerRowB * N + innerColB * 4 + 0];
|
||||
Bs[innerRowB * BN + innerColB * 4 + 1] = B[innerRowB * N + innerColB * 4 + 1];
|
||||
Bs[innerRowB * BN + innerColB * 4 + 2] = B[innerRowB * N + innerColB * 4 + 2];
|
||||
Bs[innerRowB * BN + innerColB * 4 + 3] = B[innerRowB * N + innerColB * 4 + 3];
|
||||
__syncthreads();
|
||||
|
||||
// Register cache for A transpose
|
||||
__half ldg_a_reg[4 * ldg_a_num];
|
||||
// advance blocktile
|
||||
A += BK; // move BK columns to right
|
||||
B += BK * N; // move BK rows down
|
||||
|
||||
// Fragment registers
|
||||
__half a_frag[TM];
|
||||
__half b_frag[TN];
|
||||
|
||||
float alpha = __half2float(alpha_h);
|
||||
float beta = __half2float(beta_h);
|
||||
|
||||
// Move to current block
|
||||
const __half* A_ptr = A + by * BM * K;
|
||||
const __half* B_ptr = B + bx * BN;
|
||||
__half* C_ptr = C + by * BM * N + bx * BN;
|
||||
|
||||
for (int k = 0; k < K; k += BK) {
|
||||
// Load A tile and transpose into shared memory
|
||||
#pragma unroll
|
||||
for (int i = 0; i < BM; i += a_tile_stride) {
|
||||
int a_row = a_tile_row + i;
|
||||
int a_col = a_tile_col;
|
||||
// Boundary check
|
||||
if (by * BM + a_row < M && k + a_col + 3 < K) {
|
||||
int ldg_index = i / a_tile_stride * 4;
|
||||
// Load 4 halfs from global memory
|
||||
ldg_a_reg[ldg_index + 0] = A_ptr[OFFSET(a_row, a_col + 0, K)];
|
||||
ldg_a_reg[ldg_index + 1] = A_ptr[OFFSET(a_row, a_col + 1, K)];
|
||||
ldg_a_reg[ldg_index + 2] = A_ptr[OFFSET(a_row, a_col + 2, K)];
|
||||
ldg_a_reg[ldg_index + 3] = A_ptr[OFFSET(a_row, a_col + 3, K)];
|
||||
// Store transposed: As[col][row]
|
||||
As[OFFSET(a_col + 0, a_row, BM)] = ldg_a_reg[ldg_index + 0];
|
||||
As[OFFSET(a_col + 1, a_row, BM)] = ldg_a_reg[ldg_index + 1];
|
||||
As[OFFSET(a_col + 2, a_row, BM)] = ldg_a_reg[ldg_index + 2];
|
||||
As[OFFSET(a_col + 3, a_row, BM)] = ldg_a_reg[ldg_index + 3];
|
||||
} else {
|
||||
// Zero-fill out-of-bounds
|
||||
int ldg_index = i / a_tile_stride * 4;
|
||||
for (int j = 0; j < 4; j++) {
|
||||
__half val = __float2half(0.0f);
|
||||
if (by * BM + a_row < M && k + a_col + j < K)
|
||||
val = A_ptr[OFFSET(a_row, a_col + j, K)];
|
||||
As[OFFSET(a_col + j, a_row, BM)] = val;
|
||||
}
|
||||
}
|
||||
// calculate per-thread results
|
||||
for (uint dotIdx = 0; dotIdx < BK; ++dotIdx) {
|
||||
// block into registers
|
||||
for (uint i = 0; i < TM; ++i) {
|
||||
regM[i] = As[dotIdx * BM + threadRow * TM + i];
|
||||
}
|
||||
for (uint i = 0; i < TN; ++i) {
|
||||
regN[i] = Bs[dotIdx * BN + threadCol * TN + i];
|
||||
}
|
||||
// FP32 accumulation
|
||||
for (uint resIdxM = 0; resIdxM < TM; ++resIdxM) {
|
||||
float aVal = __half2float(regM[resIdxM]);
|
||||
for (uint resIdxN = 0; resIdxN < TN; ++resIdxN) {
|
||||
threadResults[resIdxM * TN + resIdxN] +=
|
||||
aVal * __half2float(regN[resIdxN]);
|
||||
}
|
||||
|
||||
// Load B tile directly (no transpose)
|
||||
#pragma unroll
|
||||
for (int i = 0; i < BK; i += b_tile_stride) {
|
||||
int b_row = b_tile_row + i;
|
||||
int b_col = b_tile_col;
|
||||
if (k + b_row < K && bx * BN + b_col + 3 < N) {
|
||||
Bs[OFFSET(b_row, b_col + 0, BN)] = B_ptr[OFFSET(b_row, b_col + 0, N)];
|
||||
Bs[OFFSET(b_row, b_col + 1, BN)] = B_ptr[OFFSET(b_row, b_col + 1, N)];
|
||||
Bs[OFFSET(b_row, b_col + 2, BN)] = B_ptr[OFFSET(b_row, b_col + 2, N)];
|
||||
Bs[OFFSET(b_row, b_col + 3, BN)] = B_ptr[OFFSET(b_row, b_col + 3, N)];
|
||||
} else {
|
||||
for (int j = 0; j < 4; j++) {
|
||||
__half val = __float2half(0.0f);
|
||||
if (k + b_row < K && bx * BN + b_col + j < N)
|
||||
val = B_ptr[OFFSET(b_row, b_col + j, N)];
|
||||
Bs[OFFSET(b_row, b_col + j, BN)] = val;
|
||||
}
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
A_ptr += BK;
|
||||
B_ptr += BK * N;
|
||||
|
||||
// Compute tile: FP16 multiply, FP32 accumulate
|
||||
#pragma unroll
|
||||
for (int i = 0; i < BK; i++) {
|
||||
// Load A fragment from transposed shared memory
|
||||
#pragma unroll
|
||||
for (int m = 0; m < TM; m++) {
|
||||
a_frag[m] = As[OFFSET(i, ty + m, BM)];
|
||||
}
|
||||
// Load B fragment
|
||||
#pragma unroll
|
||||
for (int n = 0; n < TN; n++) {
|
||||
b_frag[n] = Bs[OFFSET(i, tx + n, BN)];
|
||||
}
|
||||
// Outer product with FP32 accumulation
|
||||
#pragma unroll
|
||||
for (int m = 0; m < TM; m++) {
|
||||
float a_val = __half2float(a_frag[m]);
|
||||
#pragma unroll
|
||||
for (int n = 0; n < TN; n++) {
|
||||
accum[m][n] += a_val * __half2float(b_frag[n]);
|
||||
}
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
// Write results back to C
|
||||
#pragma unroll
|
||||
for (int m = 0; m < TM; m++) {
|
||||
int c_row = by * BM + ty + m;
|
||||
if (c_row >= M) continue;
|
||||
#pragma unroll
|
||||
for (int n = 0; n < TN; n++) {
|
||||
int c_col = bx * BN + tx + n;
|
||||
if (c_col >= N) continue;
|
||||
float c_val = beta * __half2float(C_ptr[OFFSET(ty + m, tx + n, N)]);
|
||||
C_ptr[OFFSET(ty + m, tx + n, N)] =
|
||||
__float2half(alpha * accum[m][n] + c_val);
|
||||
}
|
||||
// write out the results
|
||||
for (uint resIdxM = 0; resIdxM < TM; resIdxM += 1) {
|
||||
for (uint resIdxN = 0; resIdxN < TN; resIdxN += 1) {
|
||||
uint row = cRow * BM + threadRow * TM + resIdxM;
|
||||
uint col = cCol * BN + threadCol * TN + resIdxN;
|
||||
if (row < M && col < N) {
|
||||
float c_old = __half2float(C[(threadRow * TM + resIdxM) * N +
|
||||
threadCol * TN + resIdxN]);
|
||||
C[(threadRow * TM + resIdxM) * N + threadCol * TN + resIdxN] =
|
||||
__float2half(alpha * threadResults[resIdxM * TN + resIdxN] +
|
||||
beta * c_old);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
// ============================================================================
|
||||
// Launch wrapper
|
||||
// Launch wrapper — matches siboehm runSgemmVectorize
|
||||
// ============================================================================
|
||||
void launch_hgemm_blocktiling(
|
||||
int M, int N, int K,
|
||||
const __half* alpha,
|
||||
const __half* alpha_ptr,
|
||||
const __half* A, int lda,
|
||||
const __half* B, int ldb,
|
||||
const __half* beta,
|
||||
const __half* beta_ptr,
|
||||
__half* C, int ldc,
|
||||
cudaStream_t stream
|
||||
) {
|
||||
// Tile sizes tuned for BI-V100:
|
||||
// 128KB shared mem → can use larger BM/BN
|
||||
// 16 SMs → need enough blocks for occupancy
|
||||
// 4096 max threads per block
|
||||
cudaStream_t stream)
|
||||
{
|
||||
constexpr int BM = 128;
|
||||
constexpr int BN = 128;
|
||||
constexpr int BK = 8;
|
||||
constexpr int TM = 8;
|
||||
constexpr int TN = 8;
|
||||
|
||||
constexpr int thread_num = (BM / TM) * (BN / TN); // 256 threads
|
||||
// 256 threads — same as siboehm
|
||||
constexpr int NUM_THREADS = (BM * BN) / (TM * TN);
|
||||
|
||||
dim3 grid(CEIL_DIV(N, BN), CEIL_DIV(M, BM));
|
||||
dim3 block(thread_num);
|
||||
dim3 block(NUM_THREADS);
|
||||
|
||||
hgemm_blocktiling_v6<BM, BN, BK, TM, TN>
|
||||
<<<grid, block, 0, stream>>>(M, N, K, *alpha, A, B, *beta, C);
|
||||
float alpha = 1.0f, beta = 0.0f;
|
||||
if (alpha_ptr) alpha = __half2float(*alpha_ptr);
|
||||
if (beta_ptr) beta = __half2float(*beta_ptr);
|
||||
|
||||
hgemmVectorize<BM, BN, BK, TM, TN>
|
||||
<<<grid, block, 0, stream>>>(M, N, K, alpha, A, B, beta, C);
|
||||
}
|
||||
|
||||
|
||||
// ============================================================================
|
||||
// MoE expert GEMM: loop over experts, each with different token count
|
||||
// MoE expert GEMM — C++ loop over experts (replaces Python for-loop)
|
||||
// ============================================================================
|
||||
// For prefill: each expert has different number of tokens (non-uniform M).
|
||||
// For decode: M=1 per expert, use cublasHgemmStridedBatched instead.
|
||||
//
|
||||
// expert_offsets[i] = cumulative sum of tokens for experts 0..i-1
|
||||
// expert_counts[i] = number of tokens for expert i
|
||||
void launch_moe_expert_hgemm(
|
||||
int num_experts,
|
||||
const int* expert_counts, // host array, [num_experts]
|
||||
const int* expert_offsets, // host array, [num_experts]
|
||||
int N, int K, // weight dimensions: (K, N)
|
||||
const int* expert_counts, // host, [num_experts]
|
||||
const int* expert_offsets, // host, [num_experts]
|
||||
int N, int K,
|
||||
const __half* input, // (total_tokens, K)
|
||||
const __half* weights, // (num_experts, N, K) — each expert weight
|
||||
const __half* weights, // (num_experts, N, K)
|
||||
__half* output, // (total_tokens, N)
|
||||
cudaStream_t stream
|
||||
) {
|
||||
__half alpha = __float2half(1.0f);
|
||||
__half beta = __float2half(0.0f);
|
||||
|
||||
cudaStream_t stream)
|
||||
{
|
||||
for (int e = 0; e < num_experts; e++) {
|
||||
int M = expert_counts[e];
|
||||
if (M == 0) continue;
|
||||
int M_e = expert_counts[e];
|
||||
if (M_e == 0) continue;
|
||||
|
||||
int offset = expert_offsets[e];
|
||||
const __half* A = input + offset * K; // (M, K)
|
||||
const __half* B = weights + e * N * K; // (N, K) → need transpose
|
||||
__half* C = output + offset * N; // (M, N)
|
||||
int off = expert_offsets[e];
|
||||
const __half* A = input + off * K;
|
||||
const __half* B = weights + (long long)e * N * K;
|
||||
__half* C_e = output + off * N;
|
||||
|
||||
// Note: B is stored as (N, K) row-major = (K, N) col-major
|
||||
// Our kernel expects B as (K, N) row-major
|
||||
// So we need to compute C = A @ B^T
|
||||
// Which is C(M,N) = A(M,K) * B^T(K,N) where B is (N,K)
|
||||
// In row-major: C[m][n] = sum_k A[m][k] * B[n][k]
|
||||
// This is the same as C = A * B^T
|
||||
// Our kernel computes C = A * B where B is (K,N)
|
||||
// So we pass B transposed pointer — but our kernel doesn't support
|
||||
// transposed B directly. For now, launch with B as-is and fix the
|
||||
// weight layout during model loading (pre-transpose weights to (K,N)).
|
||||
launch_hgemm_blocktiling(M, N, K, &alpha, A, K, B, N, &beta, C, N, stream);
|
||||
launch_hgemm_blocktiling(M_e, N, K,
|
||||
nullptr, A, K, B, N, nullptr, C_e, N, stream);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user