perf: use fp8 tensor-core gemm in linear backward
- dX/dW run as fp8 cublasLt gemms via fused transpose-cast - shared (m,k,n) algo cache for fwd/bwd, mutex-protected - bias add in-place on bf16 output, drop output copy
This commit is contained in:
+163
-82
@@ -13,6 +13,10 @@
|
|||||||
#include <cublasLt.h>
|
#include <cublasLt.h>
|
||||||
#include <cuda_fp8.h>
|
#include <cuda_fp8.h>
|
||||||
#include <cstdint>
|
#include <cstdint>
|
||||||
|
#include <mutex>
|
||||||
|
#include <unordered_map>
|
||||||
|
|
||||||
|
static std::recursive_mutex g_mutex;
|
||||||
|
|
||||||
static cublasLtHandle_t g_handle = nullptr;
|
static cublasLtHandle_t g_handle = nullptr;
|
||||||
static cublasLtMatmulDesc_t g_desc = nullptr;
|
static cublasLtMatmulDesc_t g_desc = nullptr;
|
||||||
@@ -23,22 +27,54 @@ static cublasLtMatmulPreference_t g_pref = nullptr;
|
|||||||
static void* g_workspace = nullptr;
|
static void* g_workspace = nullptr;
|
||||||
static size_t g_ws_size = 0;
|
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<int64_t>()(s.m);
|
||||||
|
h ^= std::hash<int64_t>()(s.k) + 0x9e3779b9 + (h << 6) + (h >> 2);
|
||||||
|
h ^= std::hash<int64_t>()(s.n) + 0x9e3779b9 + (h << 6) + (h >> 2);
|
||||||
|
return h;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
using AlgoCache = std::unordered_map<ShapeKey, cublasLtMatmulAlgo_t, ShapeKeyHash>;
|
||||||
|
|
||||||
|
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() {
|
static void ensure_cublas_lt() {
|
||||||
|
std::lock_guard<std::recursive_mutex> lock(g_mutex);
|
||||||
if (g_handle) {
|
if (g_handle) {
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
TORCH_CHECK(cublasLtCreate(&g_handle) == CUBLAS_STATUS_SUCCESS);
|
TORCH_CHECK(cublasLtCreate(&g_handle) == CUBLAS_STATUS_SUCCESS);
|
||||||
TORCH_CHECK(cublasLtMatmulDescCreate(&g_desc, CUBLAS_COMPUTE_32F, CUDA_R_32F) ==
|
create_matmul_config(&g_desc, &g_layout_a, &g_layout_b, &g_layout_c);
|
||||||
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);
|
|
||||||
TORCH_CHECK(cublasLtMatmulPreferenceCreate(&g_pref) == CUBLAS_STATUS_SUCCESS);
|
TORCH_CHECK(cublasLtMatmulPreferenceCreate(&g_pref) == CUBLAS_STATUS_SUCCESS);
|
||||||
size_t ws = 16 * 1024 * 1024;
|
size_t ws = 16 * 1024 * 1024;
|
||||||
TORCH_CHECK(cublasLtMatmulPreferenceSetAttribute(
|
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,
|
static cublasStatus_t get_algo_cached(int64_t m, int64_t k, int64_t n,
|
||||||
|
AlgoCache* cache,
|
||||||
cublasLtMatmulAlgo_t* algo);
|
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,
|
static void set_layout(cublasLtMatrixLayout_t layout, int64_t rows, int64_t cols,
|
||||||
int64_t ld) {
|
int64_t ld) {
|
||||||
TORCH_CHECK(cublasLtMatrixLayoutSetAttribute(layout, CUBLASLT_MATRIX_LAYOUT_ROWS,
|
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);
|
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");
|
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));
|
auto buf = torch::empty({m, n}, a_c.options().dtype(torch::kBFloat16));
|
||||||
|
|
||||||
ensure_cublas_lt();
|
ensure_cublas_lt();
|
||||||
set_layout(g_layout_a, k, n, k); // A col-major [K,N] (b row-major, op=T)
|
fp8_gemm_into(a_c, b_c, buf, m, k, n, stream.stream());
|
||||||
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));
|
|
||||||
return buf;
|
return buf;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
// ---------------------------------------------------------------------------
|
// ---------------------------------------------------------------------------
|
||||||
// Fused FP8 linear forward: one call = scale cast x8/w8 -> cublasLt GEMM
|
// 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(
|
__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]));
|
dst[i] = __nv_fp8_e4m3(__bfloat162float(src[i]));
|
||||||
}
|
}
|
||||||
|
|
||||||
__global__ void bias_cast_kernel(
|
__global__ void transpose_cast_bf16_to_fp8_kernel(
|
||||||
const __nv_bfloat16* __restrict__ src, __nv_bfloat16* __restrict__ dst,
|
const __nv_bfloat16* __restrict__ src, __nv_fp8_e4m3* __restrict__ dst,
|
||||||
const float* __restrict__ bias, int64_t total, int64_t n) {
|
int64_t rows, int64_t cols) {
|
||||||
// Same layout both sides (row-major [M,N]); bias added per column.
|
__shared__ __nv_fp8_e4m3 tile[32][33];
|
||||||
int64_t idx = blockIdx.x * (int64_t)blockDim.x + threadIdx.x;
|
int64_t x = blockIdx.x * 32 + threadIdx.x;
|
||||||
if (idx >= total) return;
|
int64_t y = blockIdx.y * 32 + threadIdx.y;
|
||||||
float v = __bfloat162float(src[idx]);
|
for (int j = 0; j < 32; j += 8) {
|
||||||
if (bias) v += bias[idx % n];
|
if (x < cols && y + j < rows) {
|
||||||
dst[idx] = __float2bfloat16(v);
|
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;
|
__global__ void bias_add_bf16_kernel(
|
||||||
static cublasLtMatmulAlgo_t g_last_algo;
|
__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,
|
static cublasStatus_t get_algo_cached(int64_t m, int64_t k, int64_t n,
|
||||||
|
AlgoCache* cache,
|
||||||
cublasLtMatmulAlgo_t* algo) {
|
cublasLtMatmulAlgo_t* algo) {
|
||||||
if (m == g_last_m && k == g_last_k && n == g_last_n) {
|
std::lock_guard<std::recursive_mutex> lock(g_mutex);
|
||||||
*algo = g_last_algo;
|
ShapeKey key{m, k, n};
|
||||||
|
auto it = cache->find(key);
|
||||||
|
if (it != cache->end()) {
|
||||||
|
*algo = it->second;
|
||||||
return CUBLAS_STATUS_SUCCESS;
|
return CUBLAS_STATUS_SUCCESS;
|
||||||
}
|
}
|
||||||
cublasLtMatmulHeuristicResult_t heur;
|
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);
|
TORCH_CHECK(cudaMalloc(&g_workspace, heur.workspaceSize) == cudaSuccess);
|
||||||
g_ws_size = heur.workspaceSize;
|
g_ws_size = heur.workspaceSize;
|
||||||
}
|
}
|
||||||
g_last_algo = heur.algo;
|
cache->emplace(key, heur.algo);
|
||||||
g_last_m = m; g_last_k = k; g_last_n = n;
|
|
||||||
*algo = heur.algo;
|
*algo = heur.algo;
|
||||||
return CUBLAS_STATUS_SUCCESS;
|
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<std::recursive_mutex> 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 fp8_linear_forward(torch::Tensor x, torch::Tensor w,
|
||||||
torch::Tensor bias) {
|
torch::Tensor bias) {
|
||||||
TORCH_CHECK(x.is_cuda() && w.is_cuda(), "CUDA tensors required");
|
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();
|
auto w_c = w.contiguous();
|
||||||
int64_t m = x_c.size(0), k = x_c.size(1), n = w_c.size(0);
|
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");
|
TORCH_CHECK(w_c.size(1) == k, "inner dim mismatch");
|
||||||
auto out = torch::empty({m, n}, x_c.options());
|
|
||||||
|
|
||||||
ensure_cublas_lt();
|
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 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));
|
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);
|
reinterpret_cast<__nv_fp8_e4m3*>(w8.data_ptr()), n * k);
|
||||||
C10_CUDA_CHECK(cudaGetLastError());
|
C10_CUDA_CHECK(cudaGetLastError());
|
||||||
|
|
||||||
auto buf = torch::empty({m, n}, out.options()); // row-major C[M,N] direct
|
// A/B swap makes the col-major [N,M] storage directly represent the
|
||||||
float alpha = 1.0f, beta = 0.0f;
|
// row-major output [M,N]. Keep this buffer as the public output so the
|
||||||
cublasLtMatmulAlgo_t algo;
|
// bias path does not need a second allocation or a copy kernel.
|
||||||
cublasStatus_t st = get_algo_cached(m, k, n, &algo);
|
auto out = torch::empty({m, n}, x_c.options());
|
||||||
TORCH_CHECK(st == CUBLAS_STATUS_SUCCESS,
|
fp8_gemm_into(x8, w8, out, m, k, n, stream.stream());
|
||||||
"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));
|
|
||||||
|
|
||||||
float* bias_ptr = nullptr;
|
|
||||||
auto bias_f = torch::Tensor();
|
|
||||||
if (bias.defined() && bias.numel() > 0) {
|
if (bias.defined() && bias.numel() > 0) {
|
||||||
bias_f = bias.to(torch::kFloat32).contiguous();
|
TORCH_CHECK(bias.scalar_type() == torch::kBFloat16 && bias.numel() == n,
|
||||||
bias_ptr = bias_f.data_ptr<float>();
|
"bias must be bf16 with shape [N]");
|
||||||
}
|
bias_add_bf16_kernel<<<(unsigned)((m * n + block - 1) / block), block, 0, stream>>>(
|
||||||
bias_cast_kernel<<<(unsigned)((m * n + block - 1) / block), block, 0, stream>>>(
|
reinterpret_cast<__nv_bfloat16*>(out.data_ptr()),
|
||||||
reinterpret_cast<const __nv_bfloat16*>(buf.data_ptr()),
|
reinterpret_cast<const __nv_bfloat16*>(bias.data_ptr()), m * n, n);
|
||||||
reinterpret_cast<__nv_bfloat16*>(out.data_ptr()), bias_ptr, m * n, n);
|
|
||||||
C10_CUDA_CHECK(cudaGetLastError());
|
C10_CUDA_CHECK(cudaGetLastError());
|
||||||
|
}
|
||||||
|
|
||||||
std::vector<int64_t> shape(x.sizes().begin(), x.sizes().end() - 1);
|
std::vector<int64_t> shape(x.sizes().begin(), x.sizes().end() - 1);
|
||||||
shape.push_back(n);
|
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).
|
// Fused FP8 linear backward: dX = g @ W, dW = g^T @ X, dB = sum(g).
|
||||||
// Scales are recomputed from x/w (identical to forward, no state needed).
|
|
||||||
// ---------------------------------------------------------------------------
|
// ---------------------------------------------------------------------------
|
||||||
|
|
||||||
std::tuple<torch::Tensor, torch::Tensor, torch::Tensor> fp8_linear_backward(
|
std::tuple<torch::Tensor, torch::Tensor, torch::Tensor> fp8_linear_backward(
|
||||||
torch::Tensor g, torch::Tensor x, torch::Tensor w,
|
torch::Tensor g, torch::Tensor x, torch::Tensor w,
|
||||||
std::vector<int64_t> masks) {
|
std::vector<int64_t> masks) {
|
||||||
const at::cuda::OptionalCUDAGuard guard(g.device());
|
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 g_c = g.reshape({-1, w.size(0)}).contiguous();
|
||||||
auto x_c = x.reshape({-1, x.size(-1)}).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 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_input = torch::empty_like(x);
|
||||||
auto grad_weight = torch::empty_like(w);
|
auto grad_weight = torch::empty_like(w);
|
||||||
auto grad_bias = torch::empty({0}, g_c.options().dtype(g.dtype()));
|
auto grad_bias = torch::empty({0}, g_c.options().dtype(g.dtype()));
|
||||||
// Compute dtype follows the input tensor (bf16 model -> bf16 GEMMs,
|
ensure_cublas_lt();
|
||||||
// fp32 input -> fp32); w is cast to match, no branch needed.
|
|
||||||
auto dtype = x_c.dtype();
|
auto fp8_options = g_c.options().dtype(torch::kFloat8_e4m3fn);
|
||||||
auto g_w = g_c.to(dtype);
|
auto g8 = torch::empty({m, n}, fp8_options);
|
||||||
auto w_w = w.to(dtype);
|
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<const __nv_bfloat16*>(g_c.data_ptr()),
|
||||||
|
reinterpret_cast<__nv_fp8_e4m3*>(g8.data_ptr()), m * n);
|
||||||
|
dim3 threads(32, 8);
|
||||||
if (masks[0]) {
|
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<<<blocks, threads, 0, stream>>>(
|
||||||
|
reinterpret_cast<const __nv_bfloat16*>(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]) {
|
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<<<g_blocks, threads, 0, stream>>>(
|
||||||
|
reinterpret_cast<const __nv_bfloat16*>(g_c.data_ptr()),
|
||||||
|
reinterpret_cast<__nv_fp8_e4m3*>(gt8.data_ptr()), m, n);
|
||||||
|
transpose_cast_bf16_to_fp8_kernel<<<x_blocks, threads, 0, stream>>>(
|
||||||
|
reinterpret_cast<const __nv_bfloat16*>(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]) {
|
if (masks[2]) {
|
||||||
grad_bias = g.sum(0).to(g.dtype());
|
grad_bias = g_c.sum(0).to(g.dtype());
|
||||||
}
|
}
|
||||||
return std::tuple<torch::Tensor, torch::Tensor, torch::Tensor>(
|
return std::tuple<torch::Tensor, torch::Tensor, torch::Tensor>(
|
||||||
grad_input, grad_weight, grad_bias);
|
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)");
|
"FP8 e4m3 GEMM: a[M,K] x b[N,K] -> bf16[M,N] (pre-scaled inputs)");
|
||||||
m.def("fp8_linear_forward", &fp8_linear_forward,
|
m.def("fp8_linear_forward", &fp8_linear_forward,
|
||||||
py::arg("x"), py::arg("w"), py::arg("bias"),
|
py::arg("x"), py::arg("w"), py::arg("bias"),
|
||||||
"Fused FP8 linear forward: scale cast + cublasLt GEMM + unscale "
|
"Fused FP8 linear forward: scale cast + cublasLt GEMM + bias "
|
||||||
"+ bias + transpose -> bf16, single call");
|
"-> bf16, single call");
|
||||||
m.def("fp8_linear_backward", &fp8_linear_backward,
|
m.def("fp8_linear_backward", &fp8_linear_backward,
|
||||||
py::arg("g"), py::arg("x"), py::arg("w"), py::arg("masks"),
|
py::arg("g"), py::arg("x"), py::arg("w"), py::arg("masks"),
|
||||||
"Fused linear backward: dX = g*sw @ W, dW = (g*sx)^T @ X, "
|
"Fused linear backward: dX = g*sw @ W, dW = (g*sx)^T @ X, "
|
||||||
|
|||||||
Reference in New Issue
Block a user