diff --git a/ex_engine/xllm_kernels/build_test_hgemm_warp.sh b/ex_engine/xllm_kernels/build_test_hgemm_warp.sh new file mode 100755 index 00000000..a31c041b --- /dev/null +++ b/ex_engine/xllm_kernels/build_test_hgemm_warp.sh @@ -0,0 +1,131 @@ +#!/bin/bash +# build_test_hgemm_warp.sh — Compile and benchmark kernel 10 (warp tiling) +set -eo pipefail + +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +CUDA_DIR="${SCRIPT_DIR}/cuda" + +echo "=== Compile hgemm_warptiling (kernel 10, WARPSIZE=64) ===" +python3 -c " +import torch.utils.cpp_extension as ext +import os, shutil, glob + +name = 'hgemm_warptiling' +build_dir = '${SCRIPT_DIR}/build/tmp_' + name +os.makedirs(build_dir, exist_ok=True) + +try: + mod = ext.load( + name=name, + sources=[ + '${CUDA_DIR}/hgemm_warptiling.cu', + '${CUDA_DIR}/bindings/hgemm_warp_bind.cpp', + ], + extra_include_paths=['${CUDA_DIR}/headers'], + extra_cflags=['-O2', '-std=c++17'], + extra_cuda_cflags=['-O2'], + build_directory=build_dir, + verbose=True, + ) + built = glob.glob(build_dir + '/' + name + '*.so') + if built: + dst = '${SCRIPT_DIR}/build/' + name + '.so' + shutil.copy2(built[0], dst) + print(f'[build] SUCCESS: {dst} ({os.path.getsize(dst)} bytes)') +except Exception as e: + print(f'[build] FAILED: {e}') + import traceback; traceback.print_exc() +" + +echo "" +echo "=== Test ===" +python3 << 'PYTEST' +import torch, sys, os, glob, time + +build_dir = 'ex_engine/xllm_kernels/build' +sys.path.insert(0, build_dir) + +# Load kernel 10 +try: + so = glob.glob(f'{build_dir}/tmp_hgemm_warptiling/hgemm_warptiling*.so') + if so: + import importlib.util + spec = importlib.util.spec_from_file_location("hgemm_warptiling", so[0]) + hw = importlib.util.module_from_spec(spec) + spec.loader.exec_module(hw) + print("kernel 10 (warp tiling) loaded") + else: + print("SKIP: kernel 10 .so not found") + sys.exit(0) +except Exception as e: + print(f"SKIP: {e}") + sys.exit(0) + +# Load kernel 6 for comparison +try: + so6 = glob.glob(f'{build_dir}/tmp_hgemm_blocktiling/hgemm_blocktiling*.so') + if so6: + spec6 = importlib.util.spec_from_file_location("hgemm_blocktiling", so6[0]) + hb = importlib.util.module_from_spec(spec6) + spec6.loader.exec_module(hb) + has_k6 = True + print("kernel 6 (block tiling) loaded") + else: + has_k6 = False +except: + has_k6 = False + +# Correctness +print("\n--- Correctness (128x128 @ 128x128) ---") +M, N, K = 128, 128, 128 +A = torch.randn(M, K, dtype=torch.float16, device='cuda') +B = torch.randn(K, N, dtype=torch.float16, device='cuda') +C_ref = torch.matmul(A.float(), B.float()).half() +C_k10 = hw.hgemm_warp(A, B) +diff = (C_ref.float() - C_k10.float()).abs().max().item() +print(f" Max abs diff: {diff:.6f}") +assert diff < 2.0, f"FAIL diff={diff}" +print(" PASS") + +# Correctness on MoE size +print("\n--- Correctness (256x4096 @ 4096x11008) ---") +M, N, K = 256, 11008, 4096 +A = torch.randn(M, K, dtype=torch.float16, device='cuda') * 0.01 +B = torch.randn(K, N, dtype=torch.float16, device='cuda') * 0.01 +C_ref = torch.matmul(A.float(), B.float()).half() +C_k10 = hw.hgemm_warp(A, B) +diff = (C_ref.float() - C_k10.float()).abs().max().item() +rel = diff / (C_ref.float().abs().max().item() + 1e-8) +print(f" Max abs diff: {diff:.6f}, rel: {rel:.6f}") +print(" PASS" if rel < 0.1 else " WARN: large relative diff") + +# Performance benchmark +print("\n--- Performance (256x4096 @ 4096x11008, 100 iters) ---") +M, N, K = 256, 11008, 4096 +A = torch.randn(M, K, dtype=torch.float16, device='cuda') +B = torch.randn(K, N, dtype=torch.float16, device='cuda') + +def bench(fn, name, iters=100, warmup=10): + for _ in range(warmup): + fn() + torch.cuda.synchronize() + t0 = time.time() + for _ in range(iters): + fn() + torch.cuda.synchronize() + ms = (time.time() - t0) / iters * 1000 + print(f" {name}: {ms:.2f} ms/iter") + return ms + +t_torch = bench(lambda: torch.matmul(A, B), "torch.matmul") +t_k10 = bench(lambda: hw.hgemm_warp(A, B), "kernel 10 (warp)") +if has_k6: + t_k6 = bench(lambda: hb.hgemm(A, B), "kernel 6 (block)") + print(f"\n K10/torch = {t_k10/t_torch:.2f}x") + print(f" K6/torch = {t_k6/t_torch:.2f}x") + print(f" K10/K6 = {t_k10/t_k6:.2f}x (K10 should be faster)") +else: + print(f"\n K10/torch = {t_k10/t_torch:.2f}x") + +print("\n=== DONE ===") +PYTEST diff --git a/ex_engine/xllm_kernels/cuda/bindings/hgemm_warp_bind.cpp b/ex_engine/xllm_kernels/cuda/bindings/hgemm_warp_bind.cpp new file mode 100644 index 00000000..57339247 --- /dev/null +++ b/ex_engine/xllm_kernels/cuda/bindings/hgemm_warp_bind.cpp @@ -0,0 +1,36 @@ +// hgemm_warp_bind.cpp — pybind11 for hgemm_warptiling (kernel 10, warp64) + +#include +#include +#include +#include +#include + +void launch_hgemm_warptiling( + int M, int N, int K, float alpha, + const __half* A, const __half* B, + float beta, __half* C, cudaStream_t stream); + +torch::Tensor hgemm_warp(torch::Tensor A, torch::Tensor B) { + TORCH_CHECK(A.is_cuda() && B.is_cuda(), "Inputs must be CUDA tensors"); + TORCH_CHECK(A.scalar_type() == torch::kHalf, "A must be fp16"); + TORCH_CHECK(B.scalar_type() == torch::kHalf, "B must be fp16"); + TORCH_CHECK(A.size(1) == B.size(0), "Inner dims must match"); + + int M = A.size(0), K = A.size(1), N = B.size(1); + auto C = torch::zeros({M, N}, A.options()); + + cudaStream_t stream = c10::cuda::getCurrentCUDAStream().stream(); + launch_hgemm_warptiling(M, N, K, 1.0f, + reinterpret_cast(A.data_ptr()), + reinterpret_cast(B.data_ptr()), + 0.0f, + reinterpret_cast<__half*>(C.data_ptr()), + stream); + return C; +} + +PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { + m.def("hgemm_warp", &hgemm_warp, + "FP16 GEMM warp-tiling (siboehm K10, WARPSIZE=64 for BI-V100)"); +} diff --git a/ex_engine/xllm_kernels/cuda/hgemm_warptiling.cu b/ex_engine/xllm_kernels/cuda/hgemm_warptiling.cu new file mode 100644 index 00000000..d38f2abe --- /dev/null +++ b/ex_engine/xllm_kernels/cuda/hgemm_warptiling.cu @@ -0,0 +1,196 @@ +// hgemm_warptiling.cu — FP16 warp-tiling GEMM for BI-V100 (warp_size=64) +// +// 1:1 from siboehm/SGEMM_CUDA kernel 10 (sgemmWarptiling). +// Changes from original: +// 1. WARPSIZE = 32 → 64 (BI-V100 confirmed) +// 2. float → __half for A/B/C data and shared memory +// 3. float4 vectorized load → 4 scalar __half loads +// 4. threadResults accumulator stays float (FP32 accumulation) +// 5. C writeback: scalar instead of float4 + +#include +#include + +#define CEIL_DIV(M, N) (((M) + (N)-1) / (N)) +const int WARPSIZE = 64; // BI-V100 confirmed + +namespace wt { +template +__device__ void loadFromGmem(int N, int K, const __half *A, const __half *B, + __half *As, __half *Bs, int innerRowA, int innerColA, + int innerRowB, int innerColB) { + for (uint offset = 0; offset + rowStrideA <= BM; offset += rowStrideA) { + // Load 4 halfs from A, transpose while storing + __half a0 = A[(innerRowA + offset) * K + innerColA * 4 + 0]; + __half a1 = A[(innerRowA + offset) * K + innerColA * 4 + 1]; + __half a2 = A[(innerRowA + offset) * K + innerColA * 4 + 2]; + __half a3 = A[(innerRowA + offset) * K + innerColA * 4 + 3]; + As[(innerColA * 4 + 0) * BM + innerRowA + offset] = a0; + As[(innerColA * 4 + 1) * BM + innerRowA + offset] = a1; + As[(innerColA * 4 + 2) * BM + innerRowA + offset] = a2; + As[(innerColA * 4 + 3) * BM + innerRowA + offset] = a3; + } + + for (uint offset = 0; offset + rowStrideB <= BK; offset += rowStrideB) { + // Load 4 halfs from B, no transpose + Bs[(innerRowB + offset) * BN + innerColB * 4 + 0] = + B[(innerRowB + offset) * N + innerColB * 4 + 0]; + Bs[(innerRowB + offset) * BN + innerColB * 4 + 1] = + B[(innerRowB + offset) * N + innerColB * 4 + 1]; + Bs[(innerRowB + offset) * BN + innerColB * 4 + 2] = + B[(innerRowB + offset) * N + innerColB * 4 + 2]; + Bs[(innerRowB + offset) * BN + innerColB * 4 + 3] = + B[(innerRowB + offset) * N + innerColB * 4 + 3]; + } +} + +template +__device__ void +processFromSmem(float *regM, float *regN, float *threadResults, const __half *As, + const __half *Bs, const uint warpRow, const uint warpCol, + const uint threadRowInWarp, const uint threadColInWarp) { + for (uint dotIdx = 0; dotIdx < BK; ++dotIdx) { + // populate registers for whole warptile + for (uint wSubRowIdx = 0; wSubRowIdx < WMITER; ++wSubRowIdx) { + for (uint i = 0; i < TM; ++i) { + regM[wSubRowIdx * TM + i] = __half2float( + As[(dotIdx * BM) + warpRow * WM + wSubRowIdx * WSUBM + + threadRowInWarp * TM + i]); + } + } + for (uint wSubColIdx = 0; wSubColIdx < WNITER; ++wSubColIdx) { + for (uint i = 0; i < TN; ++i) { + regN[wSubColIdx * TN + i] = __half2float( + Bs[(dotIdx * BN) + warpCol * WN + wSubColIdx * WSUBN + + threadColInWarp * TN + i]); + } + } + + // execute warptile matmul — FP32 accumulation + for (uint wSubRowIdx = 0; wSubRowIdx < WMITER; ++wSubRowIdx) { + for (uint wSubColIdx = 0; wSubColIdx < WNITER; ++wSubColIdx) { + for (uint resIdxM = 0; resIdxM < TM; ++resIdxM) { + for (uint resIdxN = 0; resIdxN < TN; ++resIdxN) { + threadResults[(wSubRowIdx * TM + resIdxM) * (WNITER * TN) + + (wSubColIdx * TN) + resIdxN] += + regM[wSubRowIdx * TM + resIdxM] * + regN[wSubColIdx * TN + resIdxN]; + } + } + } + } + } +} + +} // namespace wt + +template +__global__ void __launch_bounds__(NUM_THREADS) + hgemmWarptiling(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; + + // Placement of the warp in the threadblock tile + const uint warpIdx = threadIdx.x / WARPSIZE; // the warp this thread is in + const uint warpCol = warpIdx % (BN / WN); + const uint warpRow = warpIdx / (BN / WN); + + // size of the warp subtile + constexpr uint WMITER = (WM * WN) / (WARPSIZE * TM * TN * WNITER); + constexpr uint WSUBM = WM / WMITER; + constexpr uint WSUBN = WN / WNITER; + + // Placement of the thread in the warp subtile + const uint threadIdxInWarp = threadIdx.x % WARPSIZE; // [0, 63] + const uint threadColInWarp = threadIdxInWarp % (WSUBN / TN); + const uint threadRowInWarp = threadIdxInWarp / (WSUBN / TN); + + // allocate space for the current blocktile in SMEM + __shared__ __half As[BM * BK]; + __shared__ __half Bs[BK * BN]; + + // Move blocktile to beginning of A's row and B's column + A += cRow * BM * K; + B += cCol * BN; + // Move C_ptr to warp's output tile + C += (cRow * BM + warpRow * WM) * N + cCol * BN + warpCol * WN; + + // calculating the indices that this thread will load into SMEM + // FP16: 4 halfs per thread per step + const uint innerRowA = threadIdx.x / (BK / 4); + const uint innerColA = threadIdx.x % (BK / 4); + constexpr uint rowStrideA = (NUM_THREADS * 4) / BK; + const uint innerRowB = threadIdx.x / (BN / 4); + const uint innerColB = threadIdx.x % (BN / 4); + constexpr uint rowStrideB = NUM_THREADS / (BN / 4); + + // allocate thread-local cache for results in registerfile + float threadResults[WMITER * TM * WNITER * TN] = {0.0f}; + // we cache into registers on the warptile level + float regM[WMITER * TM] = {0.0f}; + float regN[WNITER * TN] = {0.0f}; + + // outer-most loop over block tiles + for (uint bkIdx = 0; bkIdx < K; bkIdx += BK) { + wt::loadFromGmem( + N, K, A, B, As, Bs, innerRowA, innerColA, innerRowB, innerColB); + __syncthreads(); + wt::processFromSmem(regM, regN, threadResults, As, Bs, warpRow, warpCol, + threadRowInWarp, threadColInWarp); + A += BK; // move BK columns to right + B += BK * N; // move BK rows down + __syncthreads(); + } + + // write out the results — scalar writeback (no float4 for __half) + for (uint wSubRowIdx = 0; wSubRowIdx < WMITER; ++wSubRowIdx) { + for (uint wSubColIdx = 0; wSubColIdx < WNITER; ++wSubColIdx) { + __half *C_interim = C + (wSubRowIdx * WSUBM) * N + wSubColIdx * WSUBN; + for (uint resIdxM = 0; resIdxM < TM; resIdxM += 1) { + for (uint resIdxN = 0; resIdxN < TN; resIdxN += 1) { + uint idx = (threadRowInWarp * TM + resIdxM) * N + + threadColInWarp * TN + resIdxN; + float c_old = __half2float(C_interim[idx]); + const int i = (wSubRowIdx * TM + resIdxM) * (WNITER * TN) + + wSubColIdx * TN + resIdxN; + C_interim[idx] = __float2half(alpha * threadResults[i] + beta * c_old); + } + } + } + } +} + + +// ============================================================================ +// Launch wrapper +// ============================================================================ +void launch_hgemm_warptiling( + int M, int N, int K, + float alpha, + const __half* A, + const __half* B, + float beta, + __half* C, + cudaStream_t stream) +{ + // Config for BI-V100 (warp_size=64, 128KB smem, 16 SMs): + // 128 threads = 2 warps of 64 + // probe confirmed: NUM_WARPS=2, WMITER=2, threads_per_warp=64 ✓ + constexpr int NUM_THREADS = 128; + constexpr int BM = 128, BN = 128, BK = 16; + constexpr int WM = 64, WN = 128; + constexpr int WNITER = 4; + constexpr int TM = 4, TN = 4; + + dim3 grid(CEIL_DIV(N, BN), CEIL_DIV(M, BM)); + dim3 block(NUM_THREADS); + + hgemmWarptiling + <<>>(M, N, K, alpha, A, B, beta, C); +}