| |
| #pragma once |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| #include <cuda_bf16.h> |
| #include <cuda_fp8.h> |
| #include <cuda_runtime.h> |
| #include <cstdint> |
|
|
| namespace flash_rt { |
| namespace gemm { |
| namespace block128_sm89 { |
|
|
| __device__ __forceinline__ void mma_m16n8k32_e4m3( |
| float &d0, float &d1, float &d2, float &d3, |
| uint32_t a0, uint32_t a1, uint32_t a2, uint32_t a3, |
| uint32_t b0, uint32_t b1) |
| { |
| |
| asm volatile( |
| "mma.sync.aligned.m16n8k32.row.col.f32.e4m3.e4m3.f32 " |
| "{%0, %1, %2, %3}, {%4, %5, %6, %7}, {%8, %9}, {%0, %1, %2, %3};\n" |
| : "+f"(d0), "+f"(d1), "+f"(d2), "+f"(d3) |
| : "r"(a0), "r"(a1), "r"(a2), "r"(a3), "r"(b0), "r"(b1)); |
| } |
|
|
| __device__ __forceinline__ void cp_async_16(uint32_t smem, const uint8_t* src) { |
| int b = (src == nullptr) ? 0 : 16; |
| asm volatile("cp.async.ca.shared.global [%0], [%1], 16, %2;\n" |
| :: "r"(smem), "l"(src), "r"(b)); |
| } |
|
|
| __device__ __forceinline__ uint32_t to_smem(const void* p) { |
| return static_cast<uint32_t>(__cvta_generic_to_shared(p)); |
| } |
|
|
| |
| |
| |
| __device__ __forceinline__ bool col_pair_ok(int c, int N) { |
| return c + 1 < N; |
| } |
|
|
| |
| |
| |
| __device__ __forceinline__ void ldmatrix_x4_b16( |
| uint32_t &d0, uint32_t &d1, uint32_t &d2, uint32_t &d3, uint32_t smem_addr) |
| { |
| asm volatile( |
| "ldmatrix.sync.aligned.x4.m8n8.shared.b16 {%0, %1, %2, %3}, [%4];\n" |
| : "=r"(d0), "=r"(d1), "=r"(d2), "=r"(d3) |
| : "r"(smem_addr)); |
| } |
|
|
| |
| |
| __device__ __forceinline__ float silu_f32(float x) { |
| return x / (1.0f + expf(-x)); |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| template <int BLOCK_M, int BLOCK_N, int NUM_WARPS, int STAGES, |
| int MIN_BLOCKS_PER_SM, bool RESID = false> |
| __global__ __launch_bounds__(NUM_WARPS * 32, MIN_BLOCKS_PER_SM) |
| void fp8_bs_gemm_kernel( |
| const __nv_fp8_e4m3* __restrict__ A, |
| const __nv_fp8_e4m3* __restrict__ B, |
| const float* __restrict__ act_scale, |
| const float* __restrict__ w_scale, |
| __nv_bfloat16* __restrict__ D, |
| int M, int N, int K, |
| const __nv_bfloat16* __restrict__ resid = nullptr) |
| { |
| constexpr int BLOCK_K = 128; |
| constexpr int THREADS = NUM_WARPS * 32; |
| constexpr int M_ATOMS = BLOCK_M / 16; |
| constexpr int N_ATOMS = BLOCK_N / 8; |
| constexpr int N_ATOMS_PW = N_ATOMS / NUM_WARPS; |
| constexpr int N_PAIRS_PW = N_ATOMS_PW / 2; |
| constexpr int K_ATOMS = BLOCK_K / 32; |
| constexpr int NUM_CHUNKS_PER_ROW = BLOCK_K / 16; |
| |
| |
| |
| constexpr int SWIZZLE_MASK = NUM_CHUNKS_PER_ROW - 1; |
|
|
| static_assert(BLOCK_M % 16 == 0, "BLOCK_M multiple of 16"); |
| static_assert(BLOCK_N % 8 == 0, "BLOCK_N multiple of 8"); |
| static_assert(BLOCK_N <= 128, "one CTA must fit one N scale block"); |
| static_assert((BLOCK_N / 8) % NUM_WARPS == 0, "N-atoms split across warps"); |
| static_assert(N_ATOMS_PW >= 2 && N_ATOMS_PW % 2 == 0, |
| "ldmatrix pairs 2 N-atoms: N_ATOMS_PW must be even >= 2"); |
|
|
| |
| |
| |
| |
| |
| |
| constexpr int SCALE_KTILE = 8; |
| constexpr int A_TILE = BLOCK_M * BLOCK_K; |
| constexpr int B_TILE = BLOCK_N * BLOCK_K; |
|
|
| extern __shared__ uint8_t smem_raw[]; |
| uint8_t* A_smem = smem_raw; |
| uint8_t* B_smem = A_smem + STAGES * A_TILE; |
| float* as_smem = reinterpret_cast<float*>(B_smem + STAGES * B_TILE); |
| float* ws_smem = as_smem + BLOCK_M * SCALE_KTILE; |
|
|
| const int cta_m = blockIdx.x; |
| const int cta_n = blockIdx.y; |
| const int m_base = cta_m * BLOCK_M; |
| const int n_base = cta_n * BLOCK_N; |
|
|
| const int t = threadIdx.x; |
| const int warp_id = t / 32; |
| const int lane = t % 32; |
| const int l = lane % 4; |
| const int h = lane / 4; |
| |
| const int frag_group = lane / 8; |
| const int row_in_frag = lane % 8; |
| const int row_block = frag_group / 2; |
| const int col_block = frag_group % 2; |
|
|
| const int K128 = K >> 7; |
|
|
| |
| auto stage_scales = [&](int kb0) { |
| const int as_total = BLOCK_M * SCALE_KTILE; |
| for (int idx = t; idx < as_total; idx += THREADS) { |
| int r = idx / SCALE_KTILE; |
| int kc = idx - r * SCALE_KTILE; |
| int row = m_base + r; |
| int kb = kb0 + kc; |
| as_smem[idx] = (row < M && kb < K128) |
| ? act_scale[(size_t)row * K128 + kb] : 0.0f; |
| } |
| for (int kc = t; kc < SCALE_KTILE; kc += THREADS) { |
| int kb = kb0 + kc; |
| ws_smem[kc] = (kb < K128) |
| ? w_scale[(size_t)(n_base >> 7) * K128 + kb] : 0.0f; |
| } |
| __syncthreads(); |
| }; |
|
|
| auto issue_load = [&](int stage, int k_base) { |
| constexpr int A_CHUNKS = BLOCK_M * NUM_CHUNKS_PER_ROW; |
| constexpr int A_ITERS = (A_CHUNKS + THREADS - 1) / THREADS; |
| #pragma unroll |
| for (int it = 0; it < A_ITERS; ++it) { |
| int idx = it * THREADS + t; |
| if (idx >= A_CHUNKS) break; |
| int row_a = idx / NUM_CHUNKS_PER_ROW; |
| int chunk_a = idx % NUM_CHUNKS_PER_ROW; |
| int m_glob = m_base + row_a; |
| int k_glob = k_base + chunk_a * 16; |
| const uint8_t* a_src = nullptr; |
| if (m_glob < M && k_glob < K) { |
| a_src = reinterpret_cast<const uint8_t*>(&A[(size_t)m_glob * K + k_glob]); |
| } |
| int csw = chunk_a ^ (row_a & SWIZZLE_MASK); |
| cp_async_16( |
| to_smem(&A_smem[stage * A_TILE + row_a * BLOCK_K + csw * 16]), |
| a_src); |
| } |
| constexpr int B_CHUNKS = BLOCK_N * NUM_CHUNKS_PER_ROW; |
| constexpr int B_ITERS = (B_CHUNKS + THREADS - 1) / THREADS; |
| #pragma unroll |
| for (int it = 0; it < B_ITERS; ++it) { |
| int idx = it * THREADS + t; |
| if (idx >= B_CHUNKS) break; |
| int row_b = idx / NUM_CHUNKS_PER_ROW; |
| int chunk_b = idx % NUM_CHUNKS_PER_ROW; |
| int n_glob = n_base + row_b; |
| int k_glob = k_base + chunk_b * 16; |
| const uint8_t* b_src = nullptr; |
| if (n_glob < N && k_glob < K) { |
| b_src = reinterpret_cast<const uint8_t*>(&B[(size_t)n_glob * K + k_glob]); |
| } |
| int csw = chunk_b ^ (row_b & SWIZZLE_MASK); |
| cp_async_16( |
| to_smem(&B_smem[stage * B_TILE + row_b * BLOCK_K + csw * 16]), |
| b_src); |
| } |
| }; |
|
|
| |
| float acc[M_ATOMS][N_ATOMS_PW][4]; |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) |
| #pragma unroll |
| for (int ni = 0; ni < N_ATOMS_PW; ++ni) |
| #pragma unroll |
| for (int j = 0; j < 4; ++j) acc[mi][ni][j] = 0.0f; |
|
|
| const int K_ITERS = (K + BLOCK_K - 1) / BLOCK_K; |
| #pragma unroll |
| for (int s = 0; s < STAGES - 1; ++s) { |
| int kb = s * BLOCK_K; |
| if (kb < K) issue_load(s, kb); |
| asm volatile("cp.async.commit_group;\n" ::); |
| } |
|
|
| int compute_stage = 0; |
| for (int k_iter = 0; k_iter < K_ITERS; ++k_iter) { |
| int issue_iter = k_iter + (STAGES - 1); |
| int issue_stage = issue_iter % STAGES; |
| if (issue_iter < K_ITERS) issue_load(issue_stage, issue_iter * BLOCK_K); |
| asm volatile("cp.async.commit_group;\n" ::); |
| asm volatile("cp.async.wait_group %0;\n" :: "n"(STAGES - 1)); |
| __syncthreads(); |
|
|
| |
| const int kb = k_iter; |
| |
| if ((kb % SCALE_KTILE) == 0) stage_scales(kb); |
| |
| |
| float tacc[M_ATOMS][N_ATOMS_PW][4]; |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) |
| #pragma unroll |
| for (int ni = 0; ni < N_ATOMS_PW; ++ni) |
| #pragma unroll |
| for (int j = 0; j < 4; ++j) tacc[mi][ni][j] = 0.0f; |
|
|
| uint8_t* A_stage = A_smem + compute_stage * A_TILE; |
| uint8_t* B_stage = B_smem + compute_stage * B_TILE; |
| #pragma unroll |
| for (int ka = 0; ka < K_ATOMS; ++ka) { |
| |
| uint32_t A_regs[M_ATOMS][4]; |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) { |
| int row = mi * 16 + row_block * 8 + row_in_frag; |
| int chunk = 2 * ka + col_block; |
| int csw = chunk ^ (row & SWIZZLE_MASK); |
| ldmatrix_x4_b16(A_regs[mi][0], A_regs[mi][1], A_regs[mi][2], A_regs[mi][3], |
| to_smem(&A_stage[row * BLOCK_K + csw * 16])); |
| } |
| |
| uint32_t B_regs[N_PAIRS_PW][4]; |
| #pragma unroll |
| for (int np = 0; np < N_PAIRS_PW; ++np) { |
| int nrow = warp_id * N_ATOMS_PW * 8 + np * 16 + row_block * 8 + row_in_frag; |
| int chunk = 2 * ka + col_block; |
| int csw = chunk ^ (nrow & SWIZZLE_MASK); |
| ldmatrix_x4_b16(B_regs[np][0], B_regs[np][1], B_regs[np][2], B_regs[np][3], |
| to_smem(&B_stage[nrow * BLOCK_K + csw * 16])); |
| } |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) { |
| #pragma unroll |
| for (int np = 0; np < N_PAIRS_PW; ++np) { |
| int ni0 = np * 2, ni1 = np * 2 + 1; |
| |
| mma_m16n8k32_e4m3( |
| tacc[mi][ni0][0], tacc[mi][ni0][1], tacc[mi][ni0][2], tacc[mi][ni0][3], |
| A_regs[mi][0], A_regs[mi][2], A_regs[mi][1], A_regs[mi][3], |
| B_regs[np][0], B_regs[np][1]); |
| mma_m16n8k32_e4m3( |
| tacc[mi][ni1][0], tacc[mi][ni1][1], tacc[mi][ni1][2], tacc[mi][ni1][3], |
| A_regs[mi][0], A_regs[mi][2], A_regs[mi][1], A_regs[mi][3], |
| B_regs[np][2], B_regs[np][3]); |
| } |
| } |
| } |
|
|
| |
| |
| |
| |
| int kbt = kb % SCALE_KTILE; |
| float ws_cta = ws_smem[kbt]; |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) { |
| int row0 = m_base + mi * 16 + h; |
| int row1 = row0 + 8; |
| float as0 = as_smem[(mi * 16 + h) * SCALE_KTILE + kbt]; |
| float as1 = as_smem[(mi * 16 + h + 8) * SCALE_KTILE + kbt]; |
| #pragma unroll |
| for (int ni = 0; ni < N_ATOMS_PW; ++ni) { |
| acc[mi][ni][0] += tacc[mi][ni][0] * (as0 * ws_cta); |
| acc[mi][ni][1] += tacc[mi][ni][1] * (as0 * ws_cta); |
| acc[mi][ni][2] += tacc[mi][ni][2] * (as1 * ws_cta); |
| acc[mi][ni][3] += tacc[mi][ni][3] * (as1 * ws_cta); |
| } |
| } |
| |
| |
| __syncthreads(); |
| compute_stage = (compute_stage + 1) % STAGES; |
| } |
| asm volatile("cp.async.wait_all;\n" ::); |
|
|
| |
| |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) { |
| int row0 = m_base + mi * 16 + h; |
| int row1 = row0 + 8; |
| #pragma unroll |
| for (int ni = 0; ni < N_ATOMS_PW; ++ni) { |
| int n_pair_base = n_base + warp_id * N_ATOMS_PW * 8 + ni * 8 + 2 * l; |
| |
| |
| |
| |
| if constexpr (RESID) { |
| if (row0 < M && col_pair_ok(n_pair_base, N)) { |
| __nv_bfloat162 r = *reinterpret_cast<const __nv_bfloat162*>( |
| &resid[(size_t)row0 * N + n_pair_base]); |
| *reinterpret_cast<__nv_bfloat162*>(&D[(size_t)row0 * N + n_pair_base]) = |
| __floats2bfloat162_rn(acc[mi][ni][0] + __low2float(r), |
| acc[mi][ni][1] + __high2float(r)); |
| } else if (row0 < M) { |
| if (n_pair_base < N) |
| D[(size_t)row0 * N + n_pair_base] = __float2bfloat16( |
| acc[mi][ni][0] + __bfloat162float(resid[(size_t)row0 * N + n_pair_base])); |
| if (n_pair_base + 1 < N) |
| D[(size_t)row0 * N + n_pair_base + 1] = __float2bfloat16( |
| acc[mi][ni][1] + __bfloat162float(resid[(size_t)row0 * N + n_pair_base + 1])); |
| } |
| if (row1 < M && col_pair_ok(n_pair_base, N)) { |
| __nv_bfloat162 r = *reinterpret_cast<const __nv_bfloat162*>( |
| &resid[(size_t)row1 * N + n_pair_base]); |
| *reinterpret_cast<__nv_bfloat162*>(&D[(size_t)row1 * N + n_pair_base]) = |
| __floats2bfloat162_rn(acc[mi][ni][2] + __low2float(r), |
| acc[mi][ni][3] + __high2float(r)); |
| } else if (row1 < M) { |
| if (n_pair_base < N) |
| D[(size_t)row1 * N + n_pair_base] = __float2bfloat16( |
| acc[mi][ni][2] + __bfloat162float(resid[(size_t)row1 * N + n_pair_base])); |
| if (n_pair_base + 1 < N) |
| D[(size_t)row1 * N + n_pair_base + 1] = __float2bfloat16( |
| acc[mi][ni][3] + __bfloat162float(resid[(size_t)row1 * N + n_pair_base + 1])); |
| } |
| } else { |
| |
| |
| |
| if (row0 < M && col_pair_ok(n_pair_base, N)) { |
| *reinterpret_cast<__nv_bfloat162*>(&D[(size_t)row0 * N + n_pair_base]) = |
| __floats2bfloat162_rn(acc[mi][ni][0], acc[mi][ni][1]); |
| } else if (row0 < M) { |
| if (n_pair_base < N) D[(size_t)row0 * N + n_pair_base] = __float2bfloat16(acc[mi][ni][0]); |
| if (n_pair_base + 1 < N) D[(size_t)row0 * N + n_pair_base+1] = __float2bfloat16(acc[mi][ni][1]); |
| } |
| if (row1 < M && col_pair_ok(n_pair_base, N)) { |
| *reinterpret_cast<__nv_bfloat162*>(&D[(size_t)row1 * N + n_pair_base]) = |
| __floats2bfloat162_rn(acc[mi][ni][2], acc[mi][ni][3]); |
| } else if (row1 < M) { |
| if (n_pair_base < N) D[(size_t)row1 * N + n_pair_base] = __float2bfloat16(acc[mi][ni][2]); |
| if (n_pair_base + 1 < N) D[(size_t)row1 * N + n_pair_base+1] = __float2bfloat16(acc[mi][ni][3]); |
| } |
| } |
| } |
| } |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| template <int BLOCK_M, int BLOCK_N, int NUM_WARPS, int STAGES, |
| int MIN_BLOCKS_PER_SM> |
| __global__ __launch_bounds__(NUM_WARPS * 32, MIN_BLOCKS_PER_SM) |
| void fp8_bs_geglu_silu_fold_kernel( |
| const __nv_fp8_e4m3* __restrict__ A, |
| const __nv_fp8_e4m3* __restrict__ B, |
| const float* __restrict__ act_scale, |
| const float* __restrict__ w_scale, |
| __nv_fp8_e4m3* __restrict__ output, |
| float* __restrict__ out_scale, |
| int M, int N, int K) |
| { |
| static_assert(BLOCK_N == 128, |
| "GeGLU silu-fold requires BLOCK_N==128 (one quant block per CTA)"); |
| constexpr int BLOCK_K = 128; |
| constexpr int THREADS = NUM_WARPS * 32; |
| constexpr int M_ATOMS = BLOCK_M / 16; |
| constexpr int N_ATOMS = BLOCK_N / 8; |
| constexpr int N_ATOMS_PW = N_ATOMS / NUM_WARPS; |
| constexpr int N_PAIRS_PW = N_ATOMS_PW / 2; |
| constexpr int K_ATOMS = BLOCK_K / 32; |
| constexpr int NUM_CHUNKS_PER_ROW = BLOCK_K / 16; |
| constexpr int SWIZZLE_MASK = NUM_CHUNKS_PER_ROW - 1; |
| constexpr int SCALE_KTILE = 8; |
| constexpr int A_TILE = BLOCK_M * BLOCK_K; |
| constexpr int B_TILE = BLOCK_N * BLOCK_K; |
|
|
| static_assert(BLOCK_M % 16 == 0, "BLOCK_M multiple of 16"); |
| static_assert(N_ATOMS_PW >= 2 && N_ATOMS_PW % 2 == 0, |
| "ldmatrix pairs 2 N-atoms: N_ATOMS_PW must be even >= 2"); |
|
|
| extern __shared__ uint8_t smem_raw[]; |
| uint8_t* A_smem = smem_raw; |
| uint8_t* B_smem = A_smem + STAGES * A_TILE; |
| |
| |
| |
| __nv_bfloat16* gate_smem = reinterpret_cast<__nv_bfloat16*>( |
| B_smem + STAGES * B_TILE); |
| float* as_smem = reinterpret_cast<float*>(gate_smem + BLOCK_M * BLOCK_N); |
| float* wsg_smem = as_smem + BLOCK_M * SCALE_KTILE; |
| float* wsu_smem = wsg_smem + SCALE_KTILE; |
| |
| float* amax_smem = wsu_smem + SCALE_KTILE; |
|
|
| const int cta_m = blockIdx.x; |
| const int cta_n = blockIdx.y; |
| const int m_base = cta_m * BLOCK_M; |
| const int n_base = cta_n * BLOCK_N; |
|
|
| const int t = threadIdx.x; |
| const int warp_id = t / 32; |
| const int lane = t % 32; |
| const int l = lane % 4; |
| const int h = lane / 4; |
| const int frag_group = lane / 8; |
| const int row_in_frag = lane % 8; |
| const int row_block = frag_group / 2; |
| const int col_block = frag_group % 2; |
|
|
| const int K128 = K >> 7; |
| const int N128 = N >> 7; |
| |
| const int gate_b_row0 = n_base; |
| const int up_b_row0 = n_base + N; |
| const int gate_ws_row = (n_base >> 7); |
| const int up_ws_row = gate_ws_row + N128; |
|
|
| |
| auto stage_scales = [&](int kb0) { |
| const int as_total = BLOCK_M * SCALE_KTILE; |
| for (int idx = t; idx < as_total; idx += THREADS) { |
| int r = idx / SCALE_KTILE; |
| int kc = idx - r * SCALE_KTILE; |
| int row = m_base + r; |
| int kb = kb0 + kc; |
| as_smem[idx] = (row < M && kb < K128) |
| ? act_scale[(size_t)row * K128 + kb] : 0.0f; |
| } |
| for (int kc = t; kc < SCALE_KTILE; kc += THREADS) { |
| int kb = kb0 + kc; |
| wsg_smem[kc] = (kb < K128) |
| ? w_scale[(size_t)gate_ws_row * K128 + kb] : 0.0f; |
| wsu_smem[kc] = (kb < K128) |
| ? w_scale[(size_t)up_ws_row * K128 + kb] : 0.0f; |
| } |
| __syncthreads(); |
| }; |
|
|
| |
| |
| auto issue_load = [&](int stage, int k_base, int b_row0) { |
| constexpr int A_CHUNKS = BLOCK_M * NUM_CHUNKS_PER_ROW; |
| constexpr int A_ITERS = (A_CHUNKS + THREADS - 1) / THREADS; |
| #pragma unroll |
| for (int it = 0; it < A_ITERS; ++it) { |
| int idx = it * THREADS + t; |
| if (idx >= A_CHUNKS) break; |
| int row_a = idx / NUM_CHUNKS_PER_ROW; |
| int chunk_a = idx % NUM_CHUNKS_PER_ROW; |
| int m_glob = m_base + row_a; |
| int k_glob = k_base + chunk_a * 16; |
| const uint8_t* a_src = nullptr; |
| if (m_glob < M && k_glob < K) { |
| a_src = reinterpret_cast<const uint8_t*>(&A[(size_t)m_glob * K + k_glob]); |
| } |
| int csw = chunk_a ^ (row_a & SWIZZLE_MASK); |
| cp_async_16( |
| to_smem(&A_smem[stage * A_TILE + row_a * BLOCK_K + csw * 16]), |
| a_src); |
| } |
| constexpr int B_CHUNKS = BLOCK_N * NUM_CHUNKS_PER_ROW; |
| constexpr int B_ITERS = (B_CHUNKS + THREADS - 1) / THREADS; |
| #pragma unroll |
| for (int it = 0; it < B_ITERS; ++it) { |
| int idx = it * THREADS + t; |
| if (idx >= B_CHUNKS) break; |
| int row_b = idx / NUM_CHUNKS_PER_ROW; |
| int chunk_b = idx % NUM_CHUNKS_PER_ROW; |
| int n_glob = b_row0 + row_b; |
| int k_glob = k_base + chunk_b * 16; |
| const uint8_t* b_src = nullptr; |
| if (n_glob < 2 * N && k_glob < K) { |
| b_src = reinterpret_cast<const uint8_t*>(&B[(size_t)n_glob * K + k_glob]); |
| } |
| int csw = chunk_b ^ (row_b & SWIZZLE_MASK); |
| cp_async_16( |
| to_smem(&B_smem[stage * B_TILE + row_b * BLOCK_K + csw * 16]), |
| b_src); |
| } |
| }; |
|
|
| |
| |
| auto run_pass = [&](float (*acc)[N_ATOMS_PW][4], int b_row0, |
| const float* ws_smem_pass) { |
| const int K_ITERS = (K + BLOCK_K - 1) / BLOCK_K; |
| #pragma unroll |
| for (int s = 0; s < STAGES - 1; ++s) { |
| int kb = s * BLOCK_K; |
| if (kb < K) issue_load(s, kb, b_row0); |
| asm volatile("cp.async.commit_group;\n" ::); |
| } |
| int compute_stage = 0; |
| for (int k_iter = 0; k_iter < K_ITERS; ++k_iter) { |
| int issue_iter = k_iter + (STAGES - 1); |
| int issue_stage = issue_iter % STAGES; |
| if (issue_iter < K_ITERS) issue_load(issue_stage, issue_iter * BLOCK_K, b_row0); |
| asm volatile("cp.async.commit_group;\n" ::); |
| asm volatile("cp.async.wait_group %0;\n" :: "n"(STAGES - 1)); |
| __syncthreads(); |
|
|
| const int kb = k_iter; |
| if ((kb % SCALE_KTILE) == 0) stage_scales(kb); |
|
|
| float tacc[M_ATOMS][N_ATOMS_PW][4]; |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) |
| #pragma unroll |
| for (int ni = 0; ni < N_ATOMS_PW; ++ni) |
| #pragma unroll |
| for (int j = 0; j < 4; ++j) tacc[mi][ni][j] = 0.0f; |
|
|
| uint8_t* A_stage = A_smem + compute_stage * A_TILE; |
| uint8_t* B_stage = B_smem + compute_stage * B_TILE; |
| #pragma unroll |
| for (int ka = 0; ka < K_ATOMS; ++ka) { |
| uint32_t A_regs[M_ATOMS][4]; |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) { |
| int row = mi * 16 + row_block * 8 + row_in_frag; |
| int chunk = 2 * ka + col_block; |
| int csw = chunk ^ (row & SWIZZLE_MASK); |
| ldmatrix_x4_b16(A_regs[mi][0], A_regs[mi][1], A_regs[mi][2], A_regs[mi][3], |
| to_smem(&A_stage[row * BLOCK_K + csw * 16])); |
| } |
| uint32_t B_regs[N_PAIRS_PW][4]; |
| #pragma unroll |
| for (int np = 0; np < N_PAIRS_PW; ++np) { |
| int nrow = warp_id * N_ATOMS_PW * 8 + np * 16 + row_block * 8 + row_in_frag; |
| int chunk = 2 * ka + col_block; |
| int csw = chunk ^ (nrow & SWIZZLE_MASK); |
| ldmatrix_x4_b16(B_regs[np][0], B_regs[np][1], B_regs[np][2], B_regs[np][3], |
| to_smem(&B_stage[nrow * BLOCK_K + csw * 16])); |
| } |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) { |
| #pragma unroll |
| for (int np = 0; np < N_PAIRS_PW; ++np) { |
| int ni0 = np * 2, ni1 = np * 2 + 1; |
| mma_m16n8k32_e4m3( |
| tacc[mi][ni0][0], tacc[mi][ni0][1], tacc[mi][ni0][2], tacc[mi][ni0][3], |
| A_regs[mi][0], A_regs[mi][2], A_regs[mi][1], A_regs[mi][3], |
| B_regs[np][0], B_regs[np][1]); |
| mma_m16n8k32_e4m3( |
| tacc[mi][ni1][0], tacc[mi][ni1][1], tacc[mi][ni1][2], tacc[mi][ni1][3], |
| A_regs[mi][0], A_regs[mi][2], A_regs[mi][1], A_regs[mi][3], |
| B_regs[np][2], B_regs[np][3]); |
| } |
| } |
| } |
|
|
| int kbt = kb % SCALE_KTILE; |
| float ws_cta = ws_smem_pass[kbt]; |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) { |
| int row0 = m_base + mi * 16 + h; |
| int row1 = row0 + 8; |
| float as0 = as_smem[(mi * 16 + h) * SCALE_KTILE + kbt]; |
| float as1 = as_smem[(mi * 16 + h + 8) * SCALE_KTILE + kbt]; |
| #pragma unroll |
| for (int ni = 0; ni < N_ATOMS_PW; ++ni) { |
| acc[mi][ni][0] += tacc[mi][ni][0] * (as0 * ws_cta); |
| acc[mi][ni][1] += tacc[mi][ni][1] * (as0 * ws_cta); |
| acc[mi][ni][2] += tacc[mi][ni][2] * (as1 * ws_cta); |
| acc[mi][ni][3] += tacc[mi][ni][3] * (as1 * ws_cta); |
| } |
| } |
| __syncthreads(); |
| compute_stage = (compute_stage + 1) % STAGES; |
| } |
| asm volatile("cp.async.wait_all;\n" ::); |
| }; |
|
|
| |
| float gate_acc[M_ATOMS][N_ATOMS_PW][4]; |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) |
| #pragma unroll |
| for (int ni = 0; ni < N_ATOMS_PW; ++ni) |
| #pragma unroll |
| for (int j = 0; j < 4; ++j) gate_acc[mi][ni][j] = 0.0f; |
|
|
| run_pass(gate_acc, gate_b_row0, wsg_smem); |
|
|
| |
| |
| |
| |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) { |
| int row0 = m_base + mi * 16 + h; |
| int row1 = row0 + 8; |
| #pragma unroll |
| for (int ni = 0; ni < N_ATOMS_PW; ++ni) { |
| int n_pair_base = warp_id * N_ATOMS_PW * 8 + ni * 8 + 2 * l; |
| |
| if (row0 < M) { |
| __nv_bfloat162 gs = __floats2bfloat162_rn( |
| silu_f32(gate_acc[mi][ni][0]), silu_f32(gate_acc[mi][ni][1])); |
| *reinterpret_cast<__nv_bfloat162*>( |
| &gate_smem[(row0 - m_base) * BLOCK_N + n_pair_base]) = gs; |
| __nv_bfloat162 gs2 = __floats2bfloat162_rn( |
| silu_f32(gate_acc[mi][ni][2]), silu_f32(gate_acc[mi][ni][3])); |
| *reinterpret_cast<__nv_bfloat162*>( |
| &gate_smem[(row1 - m_base) * BLOCK_N + n_pair_base]) = gs2; |
| } |
| } |
| } |
| __syncthreads(); |
| |
|
|
| |
| float up_acc[M_ATOMS][N_ATOMS_PW][4]; |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) |
| #pragma unroll |
| for (int ni = 0; ni < N_ATOMS_PW; ++ni) |
| #pragma unroll |
| for (int j = 0; j < 4; ++j) up_acc[mi][ni][j] = 0.0f; |
|
|
| run_pass(up_acc, up_b_row0, wsu_smem); |
|
|
| |
| |
| |
| |
| |
| constexpr float kFp8Max = 448.0f; |
| |
| |
| |
| float v[M_ATOMS][N_ATOMS_PW][4]; |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) { |
| int row0 = m_base + mi * 16 + h; |
| int row1 = row0 + 8; |
| int rloc0 = mi * 16 + h; |
| int rloc1 = rloc0 + 8; |
| float amax0 = 0.0f, amax1 = 0.0f; |
| #pragma unroll |
| for (int ni = 0; ni < N_ATOMS_PW; ++ni) { |
| int n_pair_base = warp_id * N_ATOMS_PW * 8 + ni * 8 + 2 * l; |
| |
| if (row0 < M) { |
| __nv_bfloat162 g = *reinterpret_cast<const __nv_bfloat162*>( |
| &gate_smem[rloc0 * BLOCK_N + n_pair_base]); |
| float gf0 = __low2float(g), gf1 = __high2float(g); |
| v[mi][ni][0] = __bfloat162float(__float2bfloat16(gf0 * up_acc[mi][ni][0])); |
| v[mi][ni][1] = __bfloat162float(__float2bfloat16(gf1 * up_acc[mi][ni][1])); |
| amax0 = fmaxf(amax0, fmaxf(fabsf(v[mi][ni][0]), fabsf(v[mi][ni][1]))); |
| } else { |
| v[mi][ni][0] = 0.0f; v[mi][ni][1] = 0.0f; |
| } |
| if (row1 < M) { |
| __nv_bfloat162 g = *reinterpret_cast<const __nv_bfloat162*>( |
| &gate_smem[rloc1 * BLOCK_N + n_pair_base]); |
| float gf0 = __low2float(g), gf1 = __high2float(g); |
| v[mi][ni][2] = __bfloat162float(__float2bfloat16(gf0 * up_acc[mi][ni][2])); |
| v[mi][ni][3] = __bfloat162float(__float2bfloat16(gf1 * up_acc[mi][ni][3])); |
| amax1 = fmaxf(amax1, fmaxf(fabsf(v[mi][ni][2]), fabsf(v[mi][ni][3]))); |
| } else { |
| v[mi][ni][2] = 0.0f; v[mi][ni][3] = 0.0f; |
| } |
| } |
| |
| for (int off = 2; off > 0; off >>= 1) { |
| amax0 = fmaxf(amax0, __shfl_xor_sync(0xffffffff, amax0, off)); |
| amax1 = fmaxf(amax1, __shfl_xor_sync(0xffffffff, amax1, off)); |
| } |
| if (l == 0) { |
| amax_smem[warp_id * BLOCK_M + rloc0] = amax0; |
| amax_smem[warp_id * BLOCK_M + rloc1] = amax1; |
| } |
| } |
| __syncthreads(); |
|
|
| |
| |
| #pragma unroll |
| for (int rloc = t; rloc < BLOCK_M; rloc += THREADS) { |
| int row = m_base + rloc; |
| if (row >= M) continue; |
| float amax = 0.0f; |
| #pragma unroll |
| for (int w = 0; w < NUM_WARPS; ++w) |
| amax = fmaxf(amax, amax_smem[w * BLOCK_M + rloc]); |
| float sc = fmaxf(amax / kFp8Max, 1.0e-12f); |
| amax_smem[rloc] = sc; |
| |
| |
| |
| |
| out_scale[(size_t)row * (N >> 7) + (n_base >> 7)] = sc; |
| } |
| __syncthreads(); |
|
|
| |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) { |
| int row0 = m_base + mi * 16 + h; |
| int row1 = row0 + 8; |
| int rloc0 = mi * 16 + h; |
| int rloc1 = rloc0 + 8; |
| float sc0 = (row0 < M) ? amax_smem[rloc0] : 1.0f; |
| float sc1 = (row1 < M) ? amax_smem[rloc1] : 1.0f; |
| float inv0 = 1.0f / sc0, inv1 = 1.0f / sc1; |
| #pragma unroll |
| for (int ni = 0; ni < N_ATOMS_PW; ++ni) { |
| int n_pair_base = n_base + warp_id * N_ATOMS_PW * 8 + ni * 8 + 2 * l; |
| if (row0 < M && col_pair_ok(n_pair_base, N)) { |
| float q0 = fminf(fmaxf(v[mi][ni][0] * inv0, -kFp8Max), kFp8Max); |
| float q1 = fminf(fmaxf(v[mi][ni][1] * inv0, -kFp8Max), kFp8Max); |
| |
| __nv_fp8_e4m3 p0(q0), p1(q1); |
| uint16_t pack = (uint16_t)(*reinterpret_cast<const uint8_t*>(&p1)) << 8 |
| | (uint16_t)(*reinterpret_cast<const uint8_t*>(&p0)); |
| *reinterpret_cast<uint16_t*>(&output[(size_t)row0 * N + n_pair_base]) = pack; |
| } else if (row0 < M) { |
| if (n_pair_base < N) { |
| float q = fminf(fmaxf(v[mi][ni][0] * inv0, -kFp8Max), kFp8Max); |
| output[(size_t)row0 * N + n_pair_base] = __nv_fp8_e4m3(q); |
| } |
| if (n_pair_base + 1 < N) { |
| float q = fminf(fmaxf(v[mi][ni][1] * inv0, -kFp8Max), kFp8Max); |
| output[(size_t)row0 * N + n_pair_base + 1] = __nv_fp8_e4m3(q); |
| } |
| } |
| if (row1 < M && col_pair_ok(n_pair_base, N)) { |
| float q2 = fminf(fmaxf(v[mi][ni][2] * inv1, -kFp8Max), kFp8Max); |
| float q3 = fminf(fmaxf(v[mi][ni][3] * inv1, -kFp8Max), kFp8Max); |
| __nv_fp8_e4m3 p2(q2), p3(q3); |
| uint16_t pack = (uint16_t)(*reinterpret_cast<const uint8_t*>(&p3)) << 8 |
| | (uint16_t)(*reinterpret_cast<const uint8_t*>(&p2)); |
| *reinterpret_cast<uint16_t*>(&output[(size_t)row1 * N + n_pair_base]) = pack; |
| } else if (row1 < M) { |
| if (n_pair_base < N) { |
| float q = fminf(fmaxf(v[mi][ni][2] * inv1, -kFp8Max), kFp8Max); |
| output[(size_t)row1 * N + n_pair_base] = __nv_fp8_e4m3(q); |
| } |
| if (n_pair_base + 1 < N) { |
| float q = fminf(fmaxf(v[mi][ni][3] * inv1, -kFp8Max), kFp8Max); |
| output[(size_t)row1 * N + n_pair_base + 1] = __nv_fp8_e4m3(q); |
| } |
| } |
| } |
| } |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| template <int BLOCK_M, int BLOCK_N, int NUM_WARPS, int STAGES, |
| int MIN_BLOCKS_PER_SM> |
| __global__ __launch_bounds__(NUM_WARPS * 32, MIN_BLOCKS_PER_SM) |
| void fp8_bs_geglu_silu_fold_apersist_kernel( |
| const __nv_fp8_e4m3* __restrict__ A, |
| const __nv_fp8_e4m3* __restrict__ B, |
| const float* __restrict__ act_scale, |
| const float* __restrict__ w_scale, |
| __nv_fp8_e4m3* __restrict__ output, |
| float* __restrict__ out_scale, |
| int M, int N, int K) |
| { |
| static_assert(BLOCK_N == 128, |
| "GeGLU silu-fold requires BLOCK_N==128 (one quant block per CTA)"); |
| constexpr int BLOCK_K = 128; |
| constexpr int THREADS = NUM_WARPS * 32; |
| constexpr int M_ATOMS = BLOCK_M / 16; |
| constexpr int N_ATOMS = BLOCK_N / 8; |
| constexpr int N_ATOMS_PW = N_ATOMS / NUM_WARPS; |
| constexpr int N_PAIRS_PW = N_ATOMS_PW / 2; |
| constexpr int K_ATOMS = BLOCK_K / 32; |
| constexpr int NUM_CHUNKS_PER_ROW = BLOCK_K / 16; |
| constexpr int SWIZZLE_MASK = NUM_CHUNKS_PER_ROW - 1; |
| constexpr int SCALE_KTILE = 8; |
| constexpr int A_TILE = BLOCK_M * BLOCK_K; |
| constexpr int B_TILE = BLOCK_N * BLOCK_K; |
|
|
| static_assert(BLOCK_M % 16 == 0, "BLOCK_M multiple of 16"); |
| static_assert(N_ATOMS_PW >= 2 && N_ATOMS_PW % 2 == 0, |
| "ldmatrix pairs 2 N-atoms: N_ATOMS_PW must be even >= 2"); |
|
|
| extern __shared__ uint8_t smem_raw[]; |
| uint8_t* A_smem = smem_raw; |
| uint8_t* B_smem = A_smem + STAGES * A_TILE; |
| |
| |
| __nv_bfloat16* gate_smem = reinterpret_cast<__nv_bfloat16*>( |
| B_smem + STAGES * B_TILE); |
| float* as_smem = reinterpret_cast<float*>(gate_smem + BLOCK_M * BLOCK_N); |
| float* wsg_smem = as_smem + BLOCK_M * SCALE_KTILE; |
| float* wsu_smem = wsg_smem + SCALE_KTILE; |
| float* amax_smem = wsu_smem + SCALE_KTILE; |
|
|
| const int cta_m = blockIdx.x; |
| const int cta_n = blockIdx.y; |
| const int m_base = cta_m * BLOCK_M; |
| const int n_base = cta_n * BLOCK_N; |
| const int gate_b_row0 = n_base; |
| const int up_b_row0 = n_base + N; |
|
|
| const int t = threadIdx.x; |
| const int warp_id = t / 32; |
| const int lane = t % 32; |
| const int l = lane % 4; |
| const int h = lane / 4; |
| const int frag_group = lane / 8; |
| const int row_in_frag = lane % 8; |
| const int row_block = frag_group / 2; |
| const int col_block = frag_group % 2; |
|
|
| const int K128 = K >> 7; |
| const int N128 = N >> 7; |
| const int gate_ws_row = (n_base >> 7); |
| const int up_ws_row = gate_ws_row + N128; |
|
|
| auto stage_scales = [&](int kb0) { |
| const int as_total = BLOCK_M * SCALE_KTILE; |
| for (int idx = t; idx < as_total; idx += THREADS) { |
| int r = idx / SCALE_KTILE; |
| int kc = idx - r * SCALE_KTILE; |
| int row = m_base + r; |
| int kb = kb0 + kc; |
| as_smem[idx] = (row < M && kb < K128) |
| ? act_scale[(size_t)row * K128 + kb] : 0.0f; |
| } |
| for (int kc = t; kc < SCALE_KTILE; kc += THREADS) { |
| int kb = kb0 + kc; |
| wsg_smem[kc] = (kb < K128) |
| ? w_scale[(size_t)gate_ws_row * K128 + kb] : 0.0f; |
| wsu_smem[kc] = (kb < K128) |
| ? w_scale[(size_t)up_ws_row * K128 + kb] : 0.0f; |
| } |
| __syncthreads(); |
| }; |
|
|
| |
| auto issue_load = [&](int stage, int k_base, int b_row0) { |
| constexpr int A_CHUNKS = BLOCK_M * NUM_CHUNKS_PER_ROW; |
| constexpr int A_ITERS = (A_CHUNKS + THREADS - 1) / THREADS; |
| #pragma unroll |
| for (int it = 0; it < A_ITERS; ++it) { |
| int idx = it * THREADS + t; |
| if (idx >= A_CHUNKS) break; |
| int row_a = idx / NUM_CHUNKS_PER_ROW; |
| int chunk_a = idx % NUM_CHUNKS_PER_ROW; |
| int m_glob = m_base + row_a; |
| int k_glob = k_base + chunk_a * 16; |
| const uint8_t* a_src = nullptr; |
| if (m_glob < M && k_glob < K) { |
| a_src = reinterpret_cast<const uint8_t*>(&A[(size_t)m_glob * K + k_glob]); |
| } |
| int csw = chunk_a ^ (row_a & SWIZZLE_MASK); |
| cp_async_16( |
| to_smem(&A_smem[stage * A_TILE + row_a * BLOCK_K + csw * 16]), |
| a_src); |
| } |
| constexpr int B_CHUNKS = BLOCK_N * NUM_CHUNKS_PER_ROW; |
| constexpr int B_ITERS = (B_CHUNKS + THREADS - 1) / THREADS; |
| #pragma unroll |
| for (int it = 0; it < B_ITERS; ++it) { |
| int idx = it * THREADS + t; |
| if (idx >= B_CHUNKS) break; |
| int row_b = idx / NUM_CHUNKS_PER_ROW; |
| int chunk_b = idx % NUM_CHUNKS_PER_ROW; |
| int n_glob = b_row0 + row_b; |
| int k_glob = k_base + chunk_b * 16; |
| const uint8_t* b_src = nullptr; |
| if (n_glob < 2 * N && k_glob < K) { |
| b_src = reinterpret_cast<const uint8_t*>(&B[(size_t)n_glob * K + k_glob]); |
| } |
| int csw = chunk_b ^ (row_b & SWIZZLE_MASK); |
| cp_async_16( |
| to_smem(&B_smem[stage * B_TILE + row_b * BLOCK_K + csw * 16]), |
| b_src); |
| } |
| }; |
|
|
| |
| |
| auto mma_tile = [&](float (*acc)[N_ATOMS_PW][4], int compute_stage, |
| const float* ws_smem_pass) { |
| const int kb = (compute_stage); |
| (void)kb; |
| float tacc[M_ATOMS][N_ATOMS_PW][4]; |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) |
| #pragma unroll |
| for (int ni = 0; ni < N_ATOMS_PW; ++ni) |
| #pragma unroll |
| for (int j = 0; j < 4; ++j) tacc[mi][ni][j] = 0.0f; |
|
|
| uint8_t* A_stage = A_smem + compute_stage * A_TILE; |
| uint8_t* B_stage = B_smem + compute_stage * B_TILE; |
| #pragma unroll |
| for (int ka = 0; ka < K_ATOMS; ++ka) { |
| uint32_t A_regs[M_ATOMS][4]; |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) { |
| int row = mi * 16 + row_block * 8 + row_in_frag; |
| int chunk = 2 * ka + col_block; |
| int csw = chunk ^ (row & SWIZZLE_MASK); |
| ldmatrix_x4_b16(A_regs[mi][0], A_regs[mi][1], A_regs[mi][2], A_regs[mi][3], |
| to_smem(&A_stage[row * BLOCK_K + csw * 16])); |
| } |
| uint32_t B_regs[N_PAIRS_PW][4]; |
| #pragma unroll |
| for (int np = 0; np < N_PAIRS_PW; ++np) { |
| int nrow = warp_id * N_ATOMS_PW * 8 + np * 16 + row_block * 8 + row_in_frag; |
| int chunk = 2 * ka + col_block; |
| int csw = chunk ^ (nrow & SWIZZLE_MASK); |
| ldmatrix_x4_b16(B_regs[np][0], B_regs[np][1], B_regs[np][2], B_regs[np][3], |
| to_smem(&B_stage[nrow * BLOCK_K + csw * 16])); |
| } |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) { |
| #pragma unroll |
| for (int np = 0; np < N_PAIRS_PW; ++np) { |
| int ni0 = np * 2, ni1 = np * 2 + 1; |
| mma_m16n8k32_e4m3( |
| tacc[mi][ni0][0], tacc[mi][ni0][1], tacc[mi][ni0][2], tacc[mi][ni0][3], |
| A_regs[mi][0], A_regs[mi][2], A_regs[mi][1], A_regs[mi][3], |
| B_regs[np][0], B_regs[np][1]); |
| mma_m16n8k32_e4m3( |
| tacc[mi][ni1][0], tacc[mi][ni1][1], tacc[mi][ni1][2], tacc[mi][ni1][3], |
| A_regs[mi][0], A_regs[mi][2], A_regs[mi][1], A_regs[mi][3], |
| B_regs[np][2], B_regs[np][3]); |
| } |
| } |
| } |
| return tacc; |
| }; |
|
|
| |
| |
| float gate_acc[M_ATOMS][N_ATOMS_PW][4]; |
| float up_acc[M_ATOMS][N_ATOMS_PW][4]; |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) |
| #pragma unroll |
| for (int ni = 0; ni < N_ATOMS_PW; ++ni) |
| #pragma unroll |
| for (int j = 0; j < 4; ++j) { |
| gate_acc[mi][ni][j] = 0.0f; |
| up_acc[mi][ni][j] = 0.0f; |
| } |
|
|
| const int K_ITERS = (K + BLOCK_K - 1) / BLOCK_K; |
| |
| |
| |
| #pragma unroll |
| for (int s = 0; s < STAGES - 1; ++s) { |
| int kb = s * BLOCK_K; |
| if (kb < K) issue_load(s, kb, gate_b_row0); |
| asm volatile("cp.async.commit_group;\n" ::); |
| } |
|
|
| int compute_stage = 0; |
| for (int k_iter = 0; k_iter < K_ITERS; ++k_iter) { |
| int issue_iter = k_iter + (STAGES - 1); |
| int issue_stage = issue_iter % STAGES; |
| |
| |
| |
| if (issue_iter < K_ITERS) issue_load(issue_stage, issue_iter * BLOCK_K, gate_b_row0); |
| asm volatile("cp.async.commit_group;\n" ::); |
| asm volatile("cp.async.wait_group %0;\n" :: "n"(STAGES - 1)); |
| __syncthreads(); |
|
|
| const int kb = k_iter; |
| if ((kb % SCALE_KTILE) == 0) stage_scales(kb); |
|
|
| |
| { |
| float tacc[M_ATOMS][N_ATOMS_PW][4]; |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) |
| #pragma unroll |
| for (int ni = 0; ni < N_ATOMS_PW; ++ni) |
| #pragma unroll |
| for (int j = 0; j < 4; ++j) tacc[mi][ni][j] = 0.0f; |
| uint8_t* A_stage = A_smem + compute_stage * A_TILE; |
| uint8_t* B_stage = B_smem + compute_stage * B_TILE; |
| #pragma unroll |
| for (int ka = 0; ka < K_ATOMS; ++ka) { |
| uint32_t A_regs[M_ATOMS][4]; |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) { |
| int row = mi * 16 + row_block * 8 + row_in_frag; |
| int chunk = 2 * ka + col_block; |
| int csw = chunk ^ (row & SWIZZLE_MASK); |
| ldmatrix_x4_b16(A_regs[mi][0], A_regs[mi][1], A_regs[mi][2], A_regs[mi][3], |
| to_smem(&A_stage[row * BLOCK_K + csw * 16])); |
| } |
| uint32_t B_regs[N_PAIRS_PW][4]; |
| #pragma unroll |
| for (int np = 0; np < N_PAIRS_PW; ++np) { |
| int nrow = warp_id * N_ATOMS_PW * 8 + np * 16 + row_block * 8 + row_in_frag; |
| int chunk = 2 * ka + col_block; |
| int csw = chunk ^ (nrow & SWIZZLE_MASK); |
| ldmatrix_x4_b16(B_regs[np][0], B_regs[np][1], B_regs[np][2], B_regs[np][3], |
| to_smem(&B_stage[nrow * BLOCK_K + csw * 16])); |
| } |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) { |
| #pragma unroll |
| for (int np = 0; np < N_PAIRS_PW; ++np) { |
| int ni0 = np * 2, ni1 = np * 2 + 1; |
| mma_m16n8k32_e4m3( |
| tacc[mi][ni0][0], tacc[mi][ni0][1], tacc[mi][ni0][2], tacc[mi][ni0][3], |
| A_regs[mi][0], A_regs[mi][2], A_regs[mi][1], A_regs[mi][3], |
| B_regs[np][0], B_regs[np][1]); |
| mma_m16n8k32_e4m3( |
| tacc[mi][ni1][0], tacc[mi][ni1][1], tacc[mi][ni1][2], tacc[mi][ni1][3], |
| A_regs[mi][0], A_regs[mi][2], A_regs[mi][1], A_regs[mi][3], |
| B_regs[np][2], B_regs[np][3]); |
| } |
| } |
| } |
| int kbt = kb % SCALE_KTILE; |
| float ws_cta = wsg_smem[kbt]; |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) { |
| int row0 = m_base + mi * 16 + h; |
| int row1 = row0 + 8; |
| float as0 = as_smem[(mi * 16 + h) * SCALE_KTILE + kbt]; |
| float as1 = as_smem[(mi * 16 + h + 8) * SCALE_KTILE + kbt]; |
| #pragma unroll |
| for (int ni = 0; ni < N_ATOMS_PW; ++ni) { |
| gate_acc[mi][ni][0] += tacc[mi][ni][0] * (as0 * ws_cta); |
| gate_acc[mi][ni][1] += tacc[mi][ni][1] * (as0 * ws_cta); |
| gate_acc[mi][ni][2] += tacc[mi][ni][2] * (as1 * ws_cta); |
| gate_acc[mi][ni][3] += tacc[mi][ni][3] * (as1 * ws_cta); |
| } |
| } |
| } |
| __syncthreads(); |
|
|
| |
| |
| { |
| |
| constexpr int B_CHUNKS = BLOCK_N * NUM_CHUNKS_PER_ROW; |
| constexpr int B_ITERS = (B_CHUNKS + THREADS - 1) / THREADS; |
| int k_base = k_iter * BLOCK_K; |
| #pragma unroll |
| for (int it = 0; it < B_ITERS; ++it) { |
| int idx = it * THREADS + t; |
| if (idx >= B_CHUNKS) break; |
| int row_b = idx / NUM_CHUNKS_PER_ROW; |
| int chunk_b = idx % NUM_CHUNKS_PER_ROW; |
| int n_glob = up_b_row0 + row_b; |
| int k_glob = k_base + chunk_b * 16; |
| const uint8_t* b_src = nullptr; |
| if (n_glob < 2 * N && k_glob < K) { |
| b_src = reinterpret_cast<const uint8_t*>(&B[(size_t)n_glob * K + k_glob]); |
| } |
| int csw = chunk_b ^ (row_b & SWIZZLE_MASK); |
| cp_async_16( |
| to_smem(&B_smem[compute_stage * B_TILE + row_b * BLOCK_K + csw * 16]), |
| b_src); |
| } |
| asm volatile("cp.async.commit_group;\n" ::); |
| asm volatile("cp.async.wait_group %0;\n" :: "n"(0)); |
| __syncthreads(); |
|
|
| float tacc[M_ATOMS][N_ATOMS_PW][4]; |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) |
| #pragma unroll |
| for (int ni = 0; ni < N_ATOMS_PW; ++ni) |
| #pragma unroll |
| for (int j = 0; j < 4; ++j) tacc[mi][ni][j] = 0.0f; |
| uint8_t* A_stage = A_smem + compute_stage * A_TILE; |
| uint8_t* B_stage = B_smem + compute_stage * B_TILE; |
| #pragma unroll |
| for (int ka = 0; ka < K_ATOMS; ++ka) { |
| uint32_t A_regs[M_ATOMS][4]; |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) { |
| int row = mi * 16 + row_block * 8 + row_in_frag; |
| int chunk = 2 * ka + col_block; |
| int csw = chunk ^ (row & SWIZZLE_MASK); |
| ldmatrix_x4_b16(A_regs[mi][0], A_regs[mi][1], A_regs[mi][2], A_regs[mi][3], |
| to_smem(&A_stage[row * BLOCK_K + csw * 16])); |
| } |
| uint32_t B_regs[N_PAIRS_PW][4]; |
| #pragma unroll |
| for (int np = 0; np < N_PAIRS_PW; ++np) { |
| int nrow = warp_id * N_ATOMS_PW * 8 + np * 16 + row_block * 8 + row_in_frag; |
| int chunk = 2 * ka + col_block; |
| int csw = chunk ^ (nrow & SWIZZLE_MASK); |
| ldmatrix_x4_b16(B_regs[np][0], B_regs[np][1], B_regs[np][2], B_regs[np][3], |
| to_smem(&B_stage[nrow * BLOCK_K + csw * 16])); |
| } |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) { |
| #pragma unroll |
| for (int np = 0; np < N_PAIRS_PW; ++np) { |
| int ni0 = np * 2, ni1 = np * 2 + 1; |
| mma_m16n8k32_e4m3( |
| tacc[mi][ni0][0], tacc[mi][ni0][1], tacc[mi][ni0][2], tacc[mi][ni0][3], |
| A_regs[mi][0], A_regs[mi][2], A_regs[mi][1], A_regs[mi][3], |
| B_regs[np][0], B_regs[np][1]); |
| mma_m16n8k32_e4m3( |
| tacc[mi][ni1][0], tacc[mi][ni1][1], tacc[mi][ni1][2], tacc[mi][ni1][3], |
| A_regs[mi][0], A_regs[mi][2], A_regs[mi][1], A_regs[mi][3], |
| B_regs[np][2], B_regs[np][3]); |
| } |
| } |
| } |
| int kbt = kb % SCALE_KTILE; |
| float ws_cta = wsu_smem[kbt]; |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) { |
| int row0 = m_base + mi * 16 + h; |
| int row1 = row0 + 8; |
| float as0 = as_smem[(mi * 16 + h) * SCALE_KTILE + kbt]; |
| float as1 = as_smem[(mi * 16 + h + 8) * SCALE_KTILE + kbt]; |
| #pragma unroll |
| for (int ni = 0; ni < N_ATOMS_PW; ++ni) { |
| up_acc[mi][ni][0] += tacc[mi][ni][0] * (as0 * ws_cta); |
| up_acc[mi][ni][1] += tacc[mi][ni][1] * (as0 * ws_cta); |
| up_acc[mi][ni][2] += tacc[mi][ni][2] * (as1 * ws_cta); |
| up_acc[mi][ni][3] += tacc[mi][ni][3] * (as1 * ws_cta); |
| } |
| } |
| } |
| __syncthreads(); |
| compute_stage = (compute_stage + 1) % STAGES; |
| } |
| asm volatile("cp.async.wait_all;\n" ::); |
|
|
| |
| |
| |
| constexpr float kFp8Max = 448.0f; |
| float v[M_ATOMS][N_ATOMS_PW][4]; |
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) { |
| int row0 = m_base + mi * 16 + h; |
| int row1 = row0 + 8; |
| int rloc0 = mi * 16 + h; |
| int rloc1 = rloc0 + 8; |
| float amax0 = 0.0f, amax1 = 0.0f; |
| #pragma unroll |
| for (int ni = 0; ni < N_ATOMS_PW; ++ni) { |
| if (row0 < M) { |
| float gf0 = __bfloat162float(__float2bfloat16(silu_f32(gate_acc[mi][ni][0]))); |
| float gf1 = __bfloat162float(__float2bfloat16(silu_f32(gate_acc[mi][ni][1]))); |
| v[mi][ni][0] = __bfloat162float(__float2bfloat16(gf0 * up_acc[mi][ni][0])); |
| v[mi][ni][1] = __bfloat162float(__float2bfloat16(gf1 * up_acc[mi][ni][1])); |
| amax0 = fmaxf(amax0, fmaxf(fabsf(v[mi][ni][0]), fabsf(v[mi][ni][1]))); |
| } else { v[mi][ni][0] = 0.0f; v[mi][ni][1] = 0.0f; } |
| if (row1 < M) { |
| float gf0 = __bfloat162float(__float2bfloat16(silu_f32(gate_acc[mi][ni][2]))); |
| float gf1 = __bfloat162float(__float2bfloat16(silu_f32(gate_acc[mi][ni][3]))); |
| v[mi][ni][2] = __bfloat162float(__float2bfloat16(gf0 * up_acc[mi][ni][2])); |
| v[mi][ni][3] = __bfloat162float(__float2bfloat16(gf1 * up_acc[mi][ni][3])); |
| amax1 = fmaxf(amax1, fmaxf(fabsf(v[mi][ni][2]), fabsf(v[mi][ni][3]))); |
| } else { v[mi][ni][2] = 0.0f; v[mi][ni][3] = 0.0f; } |
| } |
| for (int off = 2; off > 0; off >>= 1) { |
| amax0 = fmaxf(amax0, __shfl_xor_sync(0xffffffff, amax0, off)); |
| amax1 = fmaxf(amax1, __shfl_xor_sync(0xffffffff, amax1, off)); |
| } |
| if (l == 0) { |
| amax_smem[warp_id * BLOCK_M + rloc0] = amax0; |
| amax_smem[warp_id * BLOCK_M + rloc1] = amax1; |
| } |
| } |
| __syncthreads(); |
|
|
| #pragma unroll |
| for (int rloc = t; rloc < BLOCK_M; rloc += THREADS) { |
| int row = m_base + rloc; |
| if (row >= M) continue; |
| float amax = 0.0f; |
| #pragma unroll |
| for (int w = 0; w < NUM_WARPS; ++w) |
| amax = fmaxf(amax, amax_smem[w * BLOCK_M + rloc]); |
| float sc = fmaxf(amax / kFp8Max, 1.0e-12f); |
| amax_smem[rloc] = sc; |
| |
| |
| |
| out_scale[(size_t)row * (N >> 7) + (n_base >> 7)] = sc; |
| } |
| __syncthreads(); |
|
|
| #pragma unroll |
| for (int mi = 0; mi < M_ATOMS; ++mi) { |
| int row0 = m_base + mi * 16 + h; |
| int row1 = row0 + 8; |
| int rloc0 = mi * 16 + h; |
| int rloc1 = rloc0 + 8; |
| float sc0 = (row0 < M) ? amax_smem[rloc0] : 1.0f; |
| float sc1 = (row1 < M) ? amax_smem[rloc1] : 1.0f; |
| float inv0 = 1.0f / sc0, inv1 = 1.0f / sc1; |
| #pragma unroll |
| for (int ni = 0; ni < N_ATOMS_PW; ++ni) { |
| int n_pair_base = n_base + warp_id * N_ATOMS_PW * 8 + ni * 8 + 2 * l; |
| if (row0 < M && col_pair_ok(n_pair_base, N)) { |
| float q0 = fminf(fmaxf(v[mi][ni][0] * inv0, -kFp8Max), kFp8Max); |
| float q1 = fminf(fmaxf(v[mi][ni][1] * inv0, -kFp8Max), kFp8Max); |
| __nv_fp8_e4m3 p0(q0), p1(q1); |
| uint16_t pack = (uint16_t)(*reinterpret_cast<const uint8_t*>(&p1)) << 8 |
| | (uint16_t)(*reinterpret_cast<const uint8_t*>(&p0)); |
| *reinterpret_cast<uint16_t*>(&output[(size_t)row0 * N + n_pair_base]) = pack; |
| } else if (row0 < M) { |
| if (n_pair_base < N) output[(size_t)row0 * N + n_pair_base] = __nv_fp8_e4m3(fminf(fmaxf(v[mi][ni][0] * inv0, -kFp8Max), kFp8Max)); |
| if (n_pair_base + 1 < N) output[(size_t)row0 * N + n_pair_base + 1] = __nv_fp8_e4m3(fminf(fmaxf(v[mi][ni][1] * inv0, -kFp8Max), kFp8Max)); |
| } |
| if (row1 < M && col_pair_ok(n_pair_base, N)) { |
| float q2 = fminf(fmaxf(v[mi][ni][2] * inv1, -kFp8Max), kFp8Max); |
| float q3 = fminf(fmaxf(v[mi][ni][3] * inv1, -kFp8Max), kFp8Max); |
| __nv_fp8_e4m3 p2(q2), p3(q3); |
| uint16_t pack = (uint16_t)(*reinterpret_cast<const uint8_t*>(&p3)) << 8 |
| | (uint16_t)(*reinterpret_cast<const uint8_t*>(&p2)); |
| *reinterpret_cast<uint16_t*>(&output[(size_t)row1 * N + n_pair_base]) = pack; |
| } else if (row1 < M) { |
| if (n_pair_base < N) output[(size_t)row1 * N + n_pair_base] = __nv_fp8_e4m3(fminf(fmaxf(v[mi][ni][2] * inv1, -kFp8Max), kFp8Max)); |
| if (n_pair_base + 1 < N) output[(size_t)row1 * N + n_pair_base + 1] = __nv_fp8_e4m3(fminf(fmaxf(v[mi][ni][3] * inv1, -kFp8Max), kFp8Max)); |
| } |
| } |
| } |
| (void)gate_smem; |
| (void)mma_tile; |
| } |
|
|
| } |
| } |
| } |
|
|