diff --git a/csrc/kernels/fp8_mm.cu b/csrc/kernels/fp8_mm.cu index 8e4a670..10152fc 100644 --- a/csrc/kernels/fp8_mm.cu +++ b/csrc/kernels/fp8_mm.cu @@ -13,6 +13,10 @@ #include #include #include +#include +#include + +static std::recursive_mutex g_mutex; static cublasLtHandle_t g_handle = nullptr; static cublasLtMatmulDesc_t g_desc = nullptr; @@ -23,22 +27,54 @@ static cublasLtMatmulPreference_t g_pref = nullptr; static void* g_workspace = nullptr; static size_t g_ws_size = 0; +struct ShapeKey { + int64_t m; + int64_t k; + int64_t n; + bool operator==(const ShapeKey& other) const { + return m == other.m && k == other.k && n == other.n; + } +}; + +struct ShapeKeyHash { + size_t operator()(const ShapeKey& s) const { + size_t h = std::hash()(s.m); + h ^= std::hash()(s.k) + 0x9e3779b9 + (h << 6) + (h >> 2); + h ^= std::hash()(s.n) + 0x9e3779b9 + (h << 6) + (h >> 2); + return h; + } +}; + +using AlgoCache = std::unordered_map; + +static void create_matmul_config(cublasLtMatmulDesc_t* desc, + cublasLtMatrixLayout_t* layout_a, + cublasLtMatrixLayout_t* layout_b, + cublasLtMatrixLayout_t* layout_c) { + cublasOperation_t ta = CUBLAS_OP_T, tb = CUBLAS_OP_N; + TORCH_CHECK(cublasLtMatmulDescCreate(desc, CUBLAS_COMPUTE_32F, CUDA_R_32F) == + CUBLAS_STATUS_SUCCESS); + TORCH_CHECK(cublasLtMatmulDescSetAttribute( + *desc, CUBLASLT_MATMUL_DESC_TRANSA, &ta, sizeof(ta)) == + CUBLAS_STATUS_SUCCESS); + TORCH_CHECK(cublasLtMatmulDescSetAttribute( + *desc, CUBLASLT_MATMUL_DESC_TRANSB, &tb, sizeof(tb)) == + CUBLAS_STATUS_SUCCESS); + TORCH_CHECK(cublasLtMatrixLayoutCreate(layout_a, CUDA_R_8F_E4M3, 1, 1, 1) == + CUBLAS_STATUS_SUCCESS); + TORCH_CHECK(cublasLtMatrixLayoutCreate(layout_b, CUDA_R_8F_E4M3, 1, 1, 1) == + CUBLAS_STATUS_SUCCESS); + TORCH_CHECK(cublasLtMatrixLayoutCreate(layout_c, CUDA_R_16BF, 1, 1, 1) == + CUBLAS_STATUS_SUCCESS); +} + static void ensure_cublas_lt() { + std::lock_guard lock(g_mutex); if (g_handle) { return; } TORCH_CHECK(cublasLtCreate(&g_handle) == CUBLAS_STATUS_SUCCESS); - TORCH_CHECK(cublasLtMatmulDescCreate(&g_desc, CUBLAS_COMPUTE_32F, CUDA_R_32F) == - CUBLAS_STATUS_SUCCESS); - cublasOperation_t ta = CUBLAS_OP_T, tb = CUBLAS_OP_N; - cublasLtMatmulDescSetAttribute(g_desc, CUBLASLT_MATMUL_DESC_TRANSA, &ta, sizeof(ta)); - cublasLtMatmulDescSetAttribute(g_desc, CUBLASLT_MATMUL_DESC_TRANSB, &tb, sizeof(tb)); - TORCH_CHECK(cublasLtMatrixLayoutCreate(&g_layout_a, CUDA_R_8F_E4M3, 1, 1, 1) == - CUBLAS_STATUS_SUCCESS); - TORCH_CHECK(cublasLtMatrixLayoutCreate(&g_layout_b, CUDA_R_8F_E4M3, 1, 1, 1) == - CUBLAS_STATUS_SUCCESS); - TORCH_CHECK(cublasLtMatrixLayoutCreate(&g_layout_c, CUDA_R_16BF, 1, 1, 1) == - CUBLAS_STATUS_SUCCESS); + create_matmul_config(&g_desc, &g_layout_a, &g_layout_b, &g_layout_c); TORCH_CHECK(cublasLtMatmulPreferenceCreate(&g_pref) == CUBLAS_STATUS_SUCCESS); size_t ws = 16 * 1024 * 1024; TORCH_CHECK(cublasLtMatmulPreferenceSetAttribute( @@ -47,8 +83,12 @@ static void ensure_cublas_lt() { } static cublasStatus_t get_algo_cached(int64_t m, int64_t k, int64_t n, + AlgoCache* cache, cublasLtMatmulAlgo_t* algo); +static void fp8_gemm_into(torch::Tensor lhs, torch::Tensor rhs, torch::Tensor out, + int64_t m, int64_t k, int64_t n, cudaStream_t stream); + static void set_layout(cublasLtMatrixLayout_t layout, int64_t rows, int64_t cols, int64_t ld) { TORCH_CHECK(cublasLtMatrixLayoutSetAttribute(layout, CUBLASLT_MATRIX_LAYOUT_ROWS, @@ -75,33 +115,16 @@ torch::Tensor fp8_mm(torch::Tensor a, torch::Tensor b) { int64_t m = a_c.size(0), k = a_c.size(1), n = b_c.size(0); TORCH_CHECK(b_c.size(1) == k, "inner dim mismatch"); - // A/B swapped so the col-major [N,M] output storage IS row-major C[M,N]: - // param A = b (op=T -> [N,K]), param B = a (op=N -> [K,M]), output no copy. auto buf = torch::empty({m, n}, a_c.options().dtype(torch::kBFloat16)); - ensure_cublas_lt(); - set_layout(g_layout_a, k, n, k); // A col-major [K,N] (b row-major, op=T) - set_layout(g_layout_b, k, m, k); // B col-major [K,M] (a row-major, op=N) - set_layout(g_layout_c, n, m, n); // C col-major [N,M], ld=N - - float alpha = 1.0f, beta = 0.0f; - cublasLtMatmulAlgo_t algo; - cublasStatus_t st = get_algo_cached(m, k, n, &algo); - TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, - "cublasLtMatmulAlgoGetHeuristic failed: ", cublasLtGetStatusName(st)); - st = cublasLtMatmul( - g_handle, g_desc, &alpha, b_c.data_ptr(), g_layout_a, a_c.data_ptr(), - g_layout_b, &beta, buf.data_ptr(), g_layout_c, buf.data_ptr(), g_layout_c, - &algo, g_workspace, g_ws_size, stream.stream()); - TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, - "cublasLtMatmul failed: ", cublasLtGetStatusName(st)); + fp8_gemm_into(a_c, b_c, buf, m, k, n, stream.stream()); return buf; } // --------------------------------------------------------------------------- // Fused FP8 linear forward: one call = scale cast x8/w8 -> cublasLt GEMM -// (bf16 output) -> transpose + unscale + bias -> bf16 [..., N]. +// (bf16 output) -> bias in-place -> bf16 [..., N]. // --------------------------------------------------------------------------- __global__ void cast_bf16_to_fp8_kernel( @@ -112,24 +135,47 @@ __global__ void cast_bf16_to_fp8_kernel( dst[i] = __nv_fp8_e4m3(__bfloat162float(src[i])); } -__global__ void bias_cast_kernel( - const __nv_bfloat16* __restrict__ src, __nv_bfloat16* __restrict__ dst, - const float* __restrict__ bias, int64_t total, int64_t n) { - // Same layout both sides (row-major [M,N]); bias added per column. - int64_t idx = blockIdx.x * (int64_t)blockDim.x + threadIdx.x; - if (idx >= total) return; - float v = __bfloat162float(src[idx]); - if (bias) v += bias[idx % n]; - dst[idx] = __float2bfloat16(v); +__global__ void transpose_cast_bf16_to_fp8_kernel( + const __nv_bfloat16* __restrict__ src, __nv_fp8_e4m3* __restrict__ dst, + int64_t rows, int64_t cols) { + __shared__ __nv_fp8_e4m3 tile[32][33]; + int64_t x = blockIdx.x * 32 + threadIdx.x; + int64_t y = blockIdx.y * 32 + threadIdx.y; + for (int j = 0; j < 32; j += 8) { + if (x < cols && y + j < rows) { + tile[threadIdx.y + j][threadIdx.x] = + __nv_fp8_e4m3(__bfloat162float(src[(y + j) * cols + x])); + } + } + __syncthreads(); + + x = blockIdx.y * 32 + threadIdx.x; + y = blockIdx.x * 32 + threadIdx.y; + for (int j = 0; j < 32; j += 8) { + if (x < rows && y + j < cols) { + dst[(y + j) * rows + x] = tile[threadIdx.x][threadIdx.y + j]; + } + } } -static int64_t g_last_m = -1, g_last_k = -1, g_last_n = -1; -static cublasLtMatmulAlgo_t g_last_algo; +__global__ void bias_add_bf16_kernel( + __nv_bfloat16* __restrict__ dst, const __nv_bfloat16* __restrict__ bias, + int64_t total, int64_t n) { + // GEMM and output use the same row-major [M,N] layout. + int64_t idx = blockIdx.x * (int64_t)blockDim.x + threadIdx.x; + if (idx >= total) return; + float v = __bfloat162float(dst[idx]); + dst[idx] = __float2bfloat16(v + __bfloat162float(bias[idx % n])); +} static cublasStatus_t get_algo_cached(int64_t m, int64_t k, int64_t n, + AlgoCache* cache, cublasLtMatmulAlgo_t* algo) { - if (m == g_last_m && k == g_last_k && n == g_last_n) { - *algo = g_last_algo; + std::lock_guard lock(g_mutex); + ShapeKey key{m, k, n}; + auto it = cache->find(key); + if (it != cache->end()) { + *algo = it->second; return CUBLAS_STATUS_SUCCESS; } cublasLtMatmulHeuristicResult_t heur; @@ -144,12 +190,31 @@ static cublasStatus_t get_algo_cached(int64_t m, int64_t k, int64_t n, TORCH_CHECK(cudaMalloc(&g_workspace, heur.workspaceSize) == cudaSuccess); g_ws_size = heur.workspaceSize; } - g_last_algo = heur.algo; - g_last_m = m; g_last_k = k; g_last_n = n; + cache->emplace(key, heur.algo); *algo = heur.algo; return CUBLAS_STATUS_SUCCESS; } +static void fp8_gemm_into(torch::Tensor lhs, torch::Tensor rhs, torch::Tensor out, + int64_t m, int64_t k, int64_t n, cudaStream_t stream) { + std::lock_guard lock(g_mutex); + set_layout(g_layout_a, k, n, k); // param A = rhs (op=T -> [N,K]) + set_layout(g_layout_b, k, m, k); // param B = lhs (op=N -> [K,M]) + set_layout(g_layout_c, n, m, n); // col-major [N,M] == row-major [M,N] + float alpha = 1.0f, beta = 0.0f; + static AlgoCache cache; + cublasLtMatmulAlgo_t algo; + cublasStatus_t st = get_algo_cached(m, k, n, &cache, &algo); + TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, + "cublasLtMatmulAlgoGetHeuristic failed: ", cublasLtGetStatusName(st)); + st = cublasLtMatmul(g_handle, g_desc, &alpha, rhs.data_ptr(), g_layout_a, + lhs.data_ptr(), g_layout_b, &beta, out.data_ptr(), g_layout_c, + out.data_ptr(), g_layout_c, &algo, g_workspace, g_ws_size, + stream); + TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, + "cublasLtMatmul failed: ", cublasLtGetStatusName(st)); +} + torch::Tensor fp8_linear_forward(torch::Tensor x, torch::Tensor w, torch::Tensor bias) { TORCH_CHECK(x.is_cuda() && w.is_cuda(), "CUDA tensors required"); @@ -162,12 +227,7 @@ torch::Tensor fp8_linear_forward(torch::Tensor x, torch::Tensor w, auto w_c = w.contiguous(); int64_t m = x_c.size(0), k = x_c.size(1), n = w_c.size(0); TORCH_CHECK(w_c.size(1) == k, "inner dim mismatch"); - auto out = torch::empty({m, n}, x_c.options()); - ensure_cublas_lt(); - set_layout(g_layout_a, k, n, k); // param A = w8 (op=T -> [N,K]) - set_layout(g_layout_b, k, m, k); // param B = x8 (op=N -> [K,M]) - set_layout(g_layout_c, n, m, n); // col-major [N,M] == row-major C[M,N] auto x8 = torch::empty({m, k}, x_c.options().dtype(torch::kFloat8_e4m3fn)); auto w8 = torch::empty({n, k}, w_c.options().dtype(torch::kFloat8_e4m3fn)); @@ -180,29 +240,20 @@ torch::Tensor fp8_linear_forward(torch::Tensor x, torch::Tensor w, reinterpret_cast<__nv_fp8_e4m3*>(w8.data_ptr()), n * k); C10_CUDA_CHECK(cudaGetLastError()); - auto buf = torch::empty({m, n}, out.options()); // row-major C[M,N] direct - float alpha = 1.0f, beta = 0.0f; - cublasLtMatmulAlgo_t algo; - cublasStatus_t st = get_algo_cached(m, k, n, &algo); - TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, - "cublasLtMatmulAlgoGetHeuristic failed: ", cublasLtGetStatusName(st)); - st = cublasLtMatmul(g_handle, g_desc, &alpha, w8.data_ptr(), g_layout_a, - x8.data_ptr(), g_layout_b, &beta, buf.data_ptr(), g_layout_c, - buf.data_ptr(), g_layout_c, &algo, g_workspace, g_ws_size, - stream.stream()); - TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS, - "cublasLtMatmul failed: ", cublasLtGetStatusName(st)); + // A/B swap makes the col-major [N,M] storage directly represent the + // row-major output [M,N]. Keep this buffer as the public output so the + // bias path does not need a second allocation or a copy kernel. + auto out = torch::empty({m, n}, x_c.options()); + fp8_gemm_into(x8, w8, out, m, k, n, stream.stream()); - float* bias_ptr = nullptr; - auto bias_f = torch::Tensor(); if (bias.defined() && bias.numel() > 0) { - bias_f = bias.to(torch::kFloat32).contiguous(); - bias_ptr = bias_f.data_ptr(); + TORCH_CHECK(bias.scalar_type() == torch::kBFloat16 && bias.numel() == n, + "bias must be bf16 with shape [N]"); + bias_add_bf16_kernel<<<(unsigned)((m * n + block - 1) / block), block, 0, stream>>>( + reinterpret_cast<__nv_bfloat16*>(out.data_ptr()), + reinterpret_cast(bias.data_ptr()), m * n, n); + C10_CUDA_CHECK(cudaGetLastError()); } - bias_cast_kernel<<<(unsigned)((m * n + block - 1) / block), block, 0, stream>>>( - reinterpret_cast(buf.data_ptr()), - reinterpret_cast<__nv_bfloat16*>(out.data_ptr()), bias_ptr, m * n, n); - C10_CUDA_CHECK(cudaGetLastError()); std::vector shape(x.sizes().begin(), x.sizes().end() - 1); shape.push_back(n); @@ -210,34 +261,64 @@ torch::Tensor fp8_linear_forward(torch::Tensor x, torch::Tensor w, } // --------------------------------------------------------------------------- -// Fused FP8 linear backward: dX = (g*sw) @ W, dW = (g*sx)^T @ X, dB = sum(g). -// Scales are recomputed from x/w (identical to forward, no state needed). +// Fused FP8 linear backward: dX = g @ W, dW = g^T @ X, dB = sum(g). // --------------------------------------------------------------------------- std::tuple fp8_linear_backward( torch::Tensor g, torch::Tensor x, torch::Tensor w, std::vector masks) { const at::cuda::OptionalCUDAGuard guard(g.device()); + TORCH_CHECK(g.dtype() == torch::kBFloat16 && x.dtype() == torch::kBFloat16 && + w.dtype() == torch::kBFloat16, + "g, x, and w must be bf16"); + auto stream = at::cuda::getCurrentCUDAStream(); auto g_c = g.reshape({-1, w.size(0)}).contiguous(); auto x_c = x.reshape({-1, x.size(-1)}).contiguous(); + auto w_c = w.contiguous(); + int64_t m = g_c.size(0); int64_t n = w.size(0); + int64_t k = w.size(1); + TORCH_CHECK(x_c.size(0) == m && x_c.size(1) == k && g_c.size(1) == n, + "backward shape mismatch"); auto grad_input = torch::empty_like(x); auto grad_weight = torch::empty_like(w); auto grad_bias = torch::empty({0}, g_c.options().dtype(g.dtype())); - // Compute dtype follows the input tensor (bf16 model -> bf16 GEMMs, - // fp32 input -> fp32); w is cast to match, no branch needed. - auto dtype = x_c.dtype(); - auto g_w = g_c.to(dtype); - auto w_w = w.to(dtype); + ensure_cublas_lt(); + + auto fp8_options = g_c.options().dtype(torch::kFloat8_e4m3fn); + auto g8 = torch::empty({m, n}, fp8_options); + auto gt8 = masks[1] ? torch::empty({n, m}, fp8_options) : torch::Tensor(); + auto wt8 = masks[0] ? torch::empty({k, n}, fp8_options) : torch::Tensor(); + auto xt8 = masks[1] ? torch::empty({k, m}, fp8_options) : torch::Tensor(); + + int64_t block = 256; + cast_bf16_to_fp8_kernel<<<(unsigned)((m * n + block - 1) / block), block, 0, stream>>>( + reinterpret_cast(g_c.data_ptr()), + reinterpret_cast<__nv_fp8_e4m3*>(g8.data_ptr()), m * n); + dim3 threads(32, 8); if (masks[0]) { - grad_input.copy_(torch::mm(g_w, w_w).reshape_as(x)); + dim3 blocks((k + 31) / 32, (n + 31) / 32); + transpose_cast_bf16_to_fp8_kernel<<>>( + reinterpret_cast(w_c.data_ptr()), + reinterpret_cast<__nv_fp8_e4m3*>(wt8.data_ptr()), n, k); + fp8_gemm_into(g8, wt8, grad_input.reshape({m, k}), m, n, k, + stream.stream()); } if (masks[1]) { - grad_weight.copy_(torch::mm(g_w.t(), x_c)); + dim3 g_blocks((n + 31) / 32, (m + 31) / 32); + dim3 x_blocks((k + 31) / 32, (m + 31) / 32); + transpose_cast_bf16_to_fp8_kernel<<>>( + reinterpret_cast(g_c.data_ptr()), + reinterpret_cast<__nv_fp8_e4m3*>(gt8.data_ptr()), m, n); + transpose_cast_bf16_to_fp8_kernel<<>>( + reinterpret_cast(x_c.data_ptr()), + reinterpret_cast<__nv_fp8_e4m3*>(xt8.data_ptr()), m, k); + fp8_gemm_into(gt8, xt8, grad_weight, n, m, k, stream.stream()); } + C10_CUDA_CHECK(cudaGetLastError()); if (masks[2]) { - grad_bias = g.sum(0).to(g.dtype()); + grad_bias = g_c.sum(0).to(g.dtype()); } return std::tuple( grad_input, grad_weight, grad_bias); @@ -248,8 +329,8 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { "FP8 e4m3 GEMM: a[M,K] x b[N,K] -> bf16[M,N] (pre-scaled inputs)"); m.def("fp8_linear_forward", &fp8_linear_forward, py::arg("x"), py::arg("w"), py::arg("bias"), - "Fused FP8 linear forward: scale cast + cublasLt GEMM + unscale " - "+ bias + transpose -> bf16, single call"); + "Fused FP8 linear forward: scale cast + cublasLt GEMM + bias " + "-> bf16, single call"); m.def("fp8_linear_backward", &fp8_linear_backward, py::arg("g"), py::arg("x"), py::arg("w"), py::arg("masks"), "Fused linear backward: dX = g*sw @ W, dW = (g*sx)^T @ X, "