submission 570719
divc13 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 429 lines, June 9 Researcher Reciprocity License v1.0.
sub60_hybrid_dispatch.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-570719?include=source"interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, mxfp4
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:b522675dbefa84f7cda91ff273cdce2ebb55fe17b880fb9c817f14b91118b99c
license declaredunknown
license concludedunknown
authorsdivc13
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ __align__(16) uint8_t A_lds[BLOCK_M * LDS_ROW];split-k
void launch_n64_splitk(tile-m = 16
constexpr int BLOCK_M = 16;tile-n = 64
Phase 60: Hybrid dispatch — sub53 (BLOCK_N=64) for small-M, sub57 (BLOCK_N=128) for large-M.vector-width = int4
int4 raw[2];Kernel source
sub60_hybrid_dispatch.py429 lines
"""
Phase 60: Hybrid dispatch — sub53 (BLOCK_N=64) for small-M, sub57 (BLOCK_N=128) for large-M.
Benchmarks 5/6 (M>=64) use BLOCK_N=128; benchmarks 1-4 (M<=32) use BLOCK_N=64.
"""
from task import input_t, output_t
import torch
import os
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
CPP_SOURCE = r"""
#include <torch/extension.h>
// BLOCK_N=64 (sub53) path
void launch_n64_nosplit(
torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
torch::Tensor C, int M, int N, int K);
void launch_n64_splitk(
torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
torch::Tensor workspace, int M, int N, int K, int split_k);
// BLOCK_N=128 (sub57) path
void launch_n128_nosplit(
torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
torch::Tensor C, int M, int N, int K);
void launch_n128_splitk(
torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
torch::Tensor workspace, int M, int N, int K, int split_k);
// Reduce
void launch_reduce(torch::Tensor workspace, torch::Tensor C, int M, int N, int split_k);
"""
HIP_SOURCE = r"""
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>
#include <torch/extension.h>
constexpr int WARP_SIZE = 64;
constexpr int MFMA_K = 128;
constexpr int DOUBLE_K = MFMA_K * 2;
constexpr int LDS_ROW = DOUBLE_K >> 1;
constexpr int HALF_K = MFMA_K >> 1;
constexpr int SCALE_GROUP = 32;
constexpr int BLOCK_M = 16;
typedef int __attribute__((ext_vector_type(4))) int4_vec;
typedef float __attribute__((ext_vector_type(4))) float4_vec;
typedef int32_t __attribute__((ext_vector_type(4))) i32x4;
typedef uint32_t __attribute__((address_space(3)))* as3_uint32_ptr;
extern "C" __device__ void llvm_amdgcn_raw_buffer_load_lds(
i32x4 rsrc, as3_uint32_ptr lds_ptr,
int size, int voffset, int soffset, int offset, int aux
) __asm("llvm.amdgcn.raw.buffer.load.lds");
struct buffer_resource { uint64_t ptr; uint32_t range; uint32_t config; };
__device__ __forceinline__ i32x4 make_srsrc(const void* ptr, uint32_t range_bytes) {
buffer_resource rsrc = {reinterpret_cast<uint64_t>(ptr), range_bytes, 0x110000};
return *reinterpret_cast<const i32x4*>(&rsrc);
}
__device__ __forceinline__ float4_vec mfma_fp4_scaled(
int4_vec A, int4_vec B, float4_vec C, int sA, int sB
) {
float4_vec D;
asm volatile(
"v_mfma_scale_f32_16x16x128_f8f6f4 %0, %1, %2, %3, %4, %5 cbsz:4 blgp:4"
: "=v"(D) : "v"(A), "v"(B), "v"(C), "v"(sA), "v"(sB));
return D;
}
__device__ __forceinline__ int lds_swz(int offset) {
return offset ^ (((offset & 2047) >> 8) << 4);
}
__device__ __forceinline__ void compute_scale(float max_abs, uint8_t& sc, float& scale_f) {
if (max_abs > 0.0f) {
uint32_t b = __float_as_uint(max_abs);
b = (b + 0x200000u) & 0xFF800000u;
int su = ((b >> 23) & 0xFF) - 129;
su = su < -127 ? -127 : (su > 127 ? 127 : su);
sc = (uint8_t)(su + 127);
scale_f = __uint_as_float((uint32_t)(su + 127) << 23);
} else { sc = 0; scale_f = 0.0f; }
}
// ============================================================
// Templated kernel — BLOCK_N and NUM_WARPS as template params
// ============================================================
template <int BLOCK_N, int NUM_WARPS>
__global__ __launch_bounds__(NUM_WARPS * 64, (NUM_WARPS == 4 ? 3 : 2))
void gemm_kernel(
const __hip_bfloat16* __restrict__ A_bf16,
const uint8_t* __restrict__ B_q,
const uint8_t* __restrict__ B_scale,
float* __restrict__ workspace,
__hip_bfloat16* __restrict__ C_out,
const int M, const int N, const int K,
const int k_steps_per_split
) {
constexpr int NUM_THREADS = NUM_WARPS * 64;
const int warp_id = threadIdx.x >> 6;
const int lane_id = threadIdx.x & 63;
const int lane_m = lane_id & 15;
const int lane_k = lane_id >> 4;
const int tid = threadIdx.x;
const int block_m = blockIdx.y * BLOCK_M;
const int block_n = blockIdx.x * BLOCK_N;
const int warp_n = block_n + (warp_id << 4);
const int split_id = blockIdx.z;
const int b_stride = K >> 1;
const int sc_stride = K >> 5;
__shared__ __align__(16) uint8_t A_lds[BLOCK_M * LDS_ROW];
__shared__ __align__(16) uint8_t B_lds[BLOCK_N * LDS_ROW];
__shared__ uint8_t A_scale_lds[BLOCK_M * 8];
__shared__ uint8_t B_scale_lds[BLOCK_N * 8];
const i32x4 b_srsrc = make_srsrc(B_q, N * b_stride);
float4_vec acc = {0.0f, 0.0f, 0.0f, 0.0f};
const int ks_start = split_id * k_steps_per_split;
const int ks_end = ks_start + k_steps_per_split;
for (int ks = ks_start; ks < ks_end; ks++) {
const int k_elem = ks * DOUBLE_K;
const int k_byte = ks * LDS_ROW;
// A quant: HW FP4 conversion (only 256 threads needed)
if (tid < 256) {
const int group_id = tid >> 1;
const int half = tid & 1;
const int q_row = group_id >> 3;
const int q_grp = group_id & 7;
const int g_row = block_m + q_row;
const int k_off = k_elem + q_grp * SCALE_GROUP + half * 16;
uint32_t pk_lo = 0, pk_hi = 0;
uint8_t a_scale_val = 0x7f;
if (g_row < M) {
const __hip_bfloat16* src = A_bf16 + g_row * K + k_off;
int4 raw[2];
#pragma unroll
for (int j = 0; j < 2; j++)
raw[j] = reinterpret_cast<const int4*>(src)[j];
const __hip_bfloat16* bf = reinterpret_cast<const __hip_bfloat16*>(raw);
float vals[16];
float local_max = 0.0f;
#pragma unroll
for (int i = 0; i < 16; i++) {
vals[i] = __bfloat162float(bf[i]);
local_max = fmaxf(local_max, fabsf(vals[i]));
}
float global_max = fmaxf(local_max, __shfl_xor(local_max, 1));
float scale_f;
compute_scale(global_max, a_scale_val, scale_f);
pk_lo = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_lo, vals[0], vals[1], scale_f, 0);
pk_lo = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_lo, vals[2], vals[3], scale_f, 1);
pk_lo = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_lo, vals[4], vals[5], scale_f, 2);
pk_lo = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_lo, vals[6], vals[7], scale_f, 3);
pk_hi = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_hi, vals[8], vals[9], scale_f, 0);
pk_hi = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_hi, vals[10], vals[11], scale_f, 1);
pk_hi = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_hi, vals[12], vals[13], scale_f, 2);
pk_hi = __builtin_amdgcn_cvt_scalef32_pk_fp4_f32(pk_hi, vals[14], vals[15], scale_f, 3);
}
const int a_lds_off = lds_swz(q_row * LDS_ROW + q_grp * 16 + half * 8);
int2 packed; packed.x = (int)pk_lo; packed.y = (int)pk_hi;
*reinterpret_cast<int2*>(&A_lds[a_lds_off]) = packed;
if (half == 0) A_scale_lds[q_row * 8 + q_grp] = a_scale_val;
}
// B: buffer_load_lds from B_shuffle with swizzled global source
{
constexpr int B_LOADS = (BLOCK_N * LDS_ROW / 16 + NUM_THREADS - 1) / NUM_THREADS;
#pragma unroll
for (int ld = 0; ld < B_LOADS; ld++) {
const int flat = (ld * NUM_THREADS + tid) << 4;
const int row = flat >> 7;
const int g_row = block_n + row;
if (row < BLOCK_N && g_row < N) {
const int swz_flat = lds_swz(flat);
const int swz_col = swz_flat & 127;
const int abs_col = k_byte + swz_col;
const int tile_n = g_row >> 4;
const int inner_n = g_row & 15;
const int tile_k = abs_col >> 5;
const int inner_k_hi = (abs_col >> 4) & 1;
const int src_off = tile_n * (b_stride << 4) + tile_k * 512 + inner_k_hi * 256 + inner_n * 16;
llvm_amdgcn_raw_buffer_load_lds(b_srsrc,
(as3_uint32_ptr)(reinterpret_cast<uintptr_t>(B_lds) + flat),
16, src_off, 0, 0, 0);
}
}
}
// B_scale from B_scale_sh — precomputed row_base
{
const int sc_off = ks << 3;
if (tid < BLOCK_N) {
const int g_row = block_n + tid;
if (g_row < N) {
const int row_base = (g_row >> 5) * 32 * sc_stride
+ (g_row & 15) * 4
+ ((g_row >> 4) & 1);
#pragma unroll
for (int grp = 0; grp < 8; grp++) {
const int abs_col = sc_off + grp;
const int col_off = (abs_col & 3) * 64
+ ((abs_col & 7) >> 2) * 2
+ (abs_col >> 3) * 256;
B_scale_lds[tid * 8 + grp] = B_scale[row_base + col_off];
}
} else {
#pragma unroll
for (int grp = 0; grp < 8; grp++)
B_scale_lds[tid * 8 + grp] = 0x7f;
}
}
}
asm volatile("s_waitcnt vmcnt(0)");
__syncthreads();
// MFMA
#pragma unroll
for (int half = 0; half < 2; half++) {
const int kh = half * HALF_K;
int4_vec A_reg;
{
const int a_off = lds_swz(lane_m * LDS_ROW + kh + (lane_k << 4));
const int4 tmp = *reinterpret_cast<const int4*>(&A_lds[a_off]);
A_reg.s0 = tmp.x; A_reg.s1 = tmp.y; A_reg.s2 = tmp.z; A_reg.s3 = tmp.w;
}
int4_vec B_reg;
{
const int b_row = (warp_id << 4) + lane_m;
const int b_off = lds_swz(b_row * LDS_ROW + kh + (lane_k << 4));
const int4 tmp = *reinterpret_cast<const int4*>(&B_lds[b_off]);
B_reg.s0 = tmp.x; B_reg.s1 = tmp.y; B_reg.s2 = tmp.z; B_reg.s3 = tmp.w;
}
const int a_sc = (int)A_scale_lds[(lane_m << 3) + (half << 2) + lane_k];
const int b_sc = (int)B_scale_lds[((warp_id << 4) + lane_m) * 8 + (half << 2) + lane_k];
acc = mfma_fp4_scaled(A_reg, B_reg, acc, a_sc, b_sc);
}
__syncthreads();
}
// Store
const int out_row = block_m + (lane_k << 2);
const int out_col = warp_n + lane_m;
if (out_col < N) {
const float* ap = reinterpret_cast<const float*>(&acc);
if (C_out) {
#pragma unroll
for (int r = 0; r < 4; r++) {
const int gm = out_row + r;
if (gm < M) C_out[gm * N + out_col] = __float2bfloat16(ap[r]);
}
} else {
float* ws = workspace + split_id * M * N;
#pragma unroll
for (int r = 0; r < 4; r++) {
const int gm = out_row + r;
if (gm < M) ws[gm * N + out_col] = ap[r];
}
}
}
}
// Reduction kernel
__global__ void reduce_kernel(
const float* __restrict__ workspace,
__hip_bfloat16* __restrict__ C,
const int M, const int N, const int split_k
) {
const int idx = blockIdx.x * blockDim.x + threadIdx.x;
if (idx >= M * N) return;
float sum = 0.0f;
for (int s = 0; s < split_k; s++)
sum += workspace[s * M * N + idx];
C[idx] = __float2bfloat16(sum);
}
// ---- Launch functions for BLOCK_N=64 (4 warps, 256 threads) ----
void launch_n64_nosplit(
torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
torch::Tensor C, int M, int N, int K
) {
const int k_steps = K / (MFMA_K * 2);
dim3 block(256);
dim3 grid((N + 63) / 64, (M + BLOCK_M - 1) / BLOCK_M, 1);
hipLaunchKernelGGL((gemm_kernel<64, 4>), grid, block, 0, 0,
reinterpret_cast<const __hip_bfloat16*>(A_bf16.data_ptr()),
reinterpret_cast<const uint8_t*>(B_q.data_ptr()),
reinterpret_cast<const uint8_t*>(B_scale.data_ptr()),
(float*)nullptr,
reinterpret_cast<__hip_bfloat16*>(C.data_ptr()), M, N, K, k_steps);
}
void launch_n64_splitk(
torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
torch::Tensor workspace, int M, int N, int K, int split_k
) {
const int k_steps = K / (MFMA_K * 2);
dim3 block(256);
dim3 grid((N + 63) / 64, (M + BLOCK_M - 1) / BLOCK_M, split_k);
hipLaunchKernelGGL((gemm_kernel<64, 4>), grid, block, 0, 0,
reinterpret_cast<const __hip_bfloat16*>(A_bf16.data_ptr()),
reinterpret_cast<const uint8_t*>(B_q.data_ptr()),
reinterpret_cast<const uint8_t*>(B_scale.data_ptr()),
reinterpret_cast<float*>(workspace.data_ptr()),
(__hip_bfloat16*)nullptr, M, N, K, k_steps / split_k);
}
// ---- Launch functions for BLOCK_N=128 (8 warps, 512 threads) ----
void launch_n128_nosplit(
torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
torch::Tensor C, int M, int N, int K
) {
const int k_steps = K / (MFMA_K * 2);
dim3 block(512);
dim3 grid((N + 127) / 128, (M + BLOCK_M - 1) / BLOCK_M, 1);
hipLaunchKernelGGL((gemm_kernel<128, 8>), grid, block, 0, 0,
reinterpret_cast<const __hip_bfloat16*>(A_bf16.data_ptr()),
reinterpret_cast<const uint8_t*>(B_q.data_ptr()),
reinterpret_cast<const uint8_t*>(B_scale.data_ptr()),
(float*)nullptr,
reinterpret_cast<__hip_bfloat16*>(C.data_ptr()), M, N, K, k_steps);
}
void launch_n128_splitk(
torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,
torch::Tensor workspace, int M, int N, int K, int split_k
) {
const int k_steps = K / (MFMA_K * 2);
dim3 block(512);
dim3 grid((N + 127) / 128, (M + BLOCK_M - 1) / BLOCK_M, split_k);
hipLaunchKernelGGL((gemm_kernel<128, 8>), grid, block, 0, 0,
reinterpret_cast<const __hip_bfloat16*>(A_bf16.data_ptr()),
reinterpret_cast<const uint8_t*>(B_q.data_ptr()),
reinterpret_cast<const uint8_t*>(B_scale.data_ptr()),
reinterpret_cast<float*>(workspace.data_ptr()),
(__hip_bfloat16*)nullptr, M, N, K, k_steps / split_k);
}
void launch_reduce(torch::Tensor workspace, torch::Tensor C, int M, int N, int split_k) {
const int num = M * N;
hipLaunchKernelGGL(reduce_kernel, dim3((num+255)/256), dim3(256), 0, 0,
reinterpret_cast<const float*>(workspace.data_ptr()),
reinterpret_cast<__hip_bfloat16*>(C.data_ptr()), M, N, split_k);
}
"""
from torch.utils.cpp_extension import load_inline
_module = None
def _get_module():
global _module
if _module is None:
_module = load_inline(
name="hybrid_v60",
cpp_sources=CPP_SOURCE,
cuda_sources=HIP_SOURCE,
functions=[
"launch_n64_nosplit", "launch_n64_splitk",
"launch_n128_nosplit", "launch_n128_splitk",
"launch_reduce",
],
verbose=False,
extra_cuda_cflags=["-O3", "-fno-gpu-rdc", "-ffp-contract=fast"],
)
return _module
def _pick_split_k(m, n, k, block_n):
k_steps = k // 256
blocks_mn = ((n + block_n - 1) // block_n) * ((m + 15) // 16)
if blocks_mn >= 256:
return 1
target_split = max(1, (608 + blocks_mn - 1) // blocks_mn)
best = 1
for s in range(1, k_steps + 1):
if k_steps % s == 0 and s <= target_split:
best = s
while best > 1 and k_steps // best < 2:
best //= 2
return max(1, best)
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
m, k = A.shape
n = B.shape[0]
mod = _get_module()
B_sh_u8 = B_shuffle.contiguous().view(torch.uint8)
B_sc = B_scale_sh.contiguous().view(torch.uint8)
# Dispatch: use BLOCK_N=128 for large-M benchmarks (M>=64)
use_n128 = (m >= 64)
if use_n128:
block_n = 128
split_k = _pick_split_k(m, n, k, block_n)
if split_k == 1:
C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")
mod.launch_n128_nosplit(A, B_sh_u8, B_sc, C, m, n, k)
else:
workspace = torch.empty((split_k, m, n), dtype=torch.float32, device="cuda")
mod.launch_n128_splitk(A, B_sh_u8, B_sc, workspace, m, n, k, split_k)
C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")
mod.launch_reduce(workspace, C, m, n, split_k)
else:
block_n = 64
split_k = _pick_split_k(m, n, k, block_n)
if split_k == 1:
C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")
mod.launch_n64_nosplit(A, B_sh_u8, B_sc, C, m, n, k)
else:
workspace = torch.empty((split_k, m, n), dtype=torch.float32, device="cuda")
mod.launch_n64_splitk(A, B_sh_u8, B_sc, workspace, m, n, k, split_k)
C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")
mod.launch_reduce(workspace, C, m, n, split_k)
return C
scrolls · 429 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 570680.
"""- Phase 59: Vectorized B_scale loading — all 256 threads active, 4 loads each.- Base: sub53_opt_bscale_addr. Change: Replace B_scale loading to use ALL 256 threads- (4 loads each) instead of 64 threads (8 loads each), giving 4x more memory-level- parallelism while keeping sub53's row_base precomputation.+ Phase 60: Hybrid dispatch — sub53 (BLOCK_N=64) for small-M, sub57 (BLOCK_N=128) for large-M.+ Benchmarks 5/6 (M>=64) use BLOCK_N=128; benchmarks 1-4 (M<=32) use BLOCK_N=64."""from task import input_t, output_timport torch⋯ 2 unchanged linesCPP_SOURCE = r"""#include <torch/extension.h>- // Separate B_scale path only (B_shuffle + B_scale_sh)- void launch_sep_gemm_splitk(+ // BLOCK_N=64 (sub53) path+ void launch_n64_nosplit(torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,+ torch::Tensor C, int M, int N, int K);+ void launch_n64_splitk(+ torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,torch::Tensor workspace, int M, int N, int K, int split_k);- void launch_sep_gemm_nosplit(+ // BLOCK_N=128 (sub57) path+ void launch_n128_nosplit(torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,torch::Tensor C, int M, int N, int K);+ void launch_n128_splitk(+ torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,+ torch::Tensor workspace, int M, int N, int K, int split_k);// Reducevoid launch_reduce(torch::Tensor workspace, torch::Tensor C, int M, int N, int split_k);"""⋯ 4 unchanged lines#include <torch/extension.h>constexpr int WARP_SIZE = 64;- constexpr int NUM_WARPS = 4;- constexpr int NUM_THREADS = NUM_WARPS * WARP_SIZE;- constexpr int BLOCK_M = 16;- constexpr int BLOCK_N = 64;constexpr int MFMA_K = 128;constexpr int DOUBLE_K = MFMA_K * 2;constexpr int LDS_ROW = DOUBLE_K >> 1;constexpr int HALF_K = MFMA_K >> 1;constexpr int SCALE_GROUP = 32;+ constexpr int BLOCK_M = 16;typedef int __attribute__((ext_vector_type(4))) int4_vec;typedef float __attribute__((ext_vector_type(4))) float4_vec;⋯ 37 unchanged lines} else { sc = 0; scale_f = 0.0f; }}- __global__ __launch_bounds__(NUM_THREADS, 3)+ // ============================================================+ // Templated kernel — BLOCK_N and NUM_WARPS as template params+ // ============================================================+ template <int BLOCK_N, int NUM_WARPS>+ __global__ __launch_bounds__(NUM_WARPS * 64, (NUM_WARPS == 4 ? 3 : 2))void gemm_kernel(const __hip_bfloat16* __restrict__ A_bf16,- const uint8_t* __restrict__ B_q, // B_shuffle data- const uint8_t* __restrict__ B_scale, // B_scale_sh data+ const uint8_t* __restrict__ B_q,+ const uint8_t* __restrict__ B_scale,float* __restrict__ workspace,__hip_bfloat16* __restrict__ C_out,const int M, const int N, const int K,const int k_steps_per_split) {+ constexpr int NUM_THREADS = NUM_WARPS * 64;const int warp_id = threadIdx.x >> 6;const int lane_id = threadIdx.x & 63;const int lane_m = lane_id & 15;⋯ 23 unchanged linesconst int k_elem = ks * DOUBLE_K;const int k_byte = ks * LDS_ROW;- // A quant: HW FP4 conversion- {+ // A quant: HW FP4 conversion (only 256 threads needed)+ if (tid < 256) {const int group_id = tid >> 1;const int half = tid & 1;const int q_row = group_id >> 3;⋯ 37 unchanged lines}// B: buffer_load_lds from B_shuffle with swizzled global source- // Load swizzled data into LDS so that swizzled LDS reads get correct data{+ constexpr int B_LOADS = (BLOCK_N * LDS_ROW / 16 + NUM_THREADS - 1) / NUM_THREADS;#pragma unroll- for (int ld = 0; ld < 2; ld++) {- const int flat = (ld * NUM_THREADS + tid) << 4; // LDS destination (linear)- const int row = flat >> 7; // which B row in this block (0..63)- const int col = flat & 127; // byte offset within 128-byte LDS row+ for (int ld = 0; ld < B_LOADS; ld++) {+ const int flat = (ld * NUM_THREADS + tid) << 4;+ const int row = flat >> 7;const int g_row = block_n + row;if (row < BLOCK_N && g_row < N) {- // Compute swizzled column: what data does swizzled LDS read expect at this flat pos?const int swz_flat = lds_swz(flat);const int swz_col = swz_flat & 127;-- // Use swz_col instead of col for the tile address computation- const int abs_col = k_byte + swz_col; // absolute byte column (swizzled)+ const int abs_col = k_byte + swz_col;const int tile_n = g_row >> 4;const int inner_n = g_row & 15;const int tile_k = abs_col >> 5;⋯ 7 unchanged lines}}- // B_scale from B_scale_sh (e8m0_shuffled layout)- // Optimized: 128 threads active, each handles 4 groups with precomputed row_base+ // B_scale from B_scale_sh — precomputed row_base{const int sc_off = ks << 3;- if (tid < 128) {- const int row = tid & 63; // 0..63- const int grp_base = (tid >> 6) << 2; // 0 or 4- const int g_row = block_n + row;-+ if (tid < BLOCK_N) {+ const int g_row = block_n + tid;if (g_row < N) {const int row_base = (g_row >> 5) * 32 * sc_stride+ (g_row & 15) * 4+ ((g_row >> 4) & 1);#pragma unroll- for (int g = 0; g < 4; g++) {- const int grp = grp_base + g;+ for (int grp = 0; grp < 8; grp++) {const int abs_col = sc_off + grp;const int col_off = (abs_col & 3) * 64+ ((abs_col & 7) >> 2) * 2+ (abs_col >> 3) * 256;- B_scale_lds[row * 8 + grp] = B_scale[row_base + col_off];+ B_scale_lds[tid * 8 + grp] = B_scale[row_base + col_off];}} else {#pragma unroll- for (int g = 0; g < 4; g++)- B_scale_lds[row * 8 + grp_base + g] = 0x7f;+ for (int grp = 0; grp < 8; grp++)+ B_scale_lds[tid * 8 + grp] = 0x7f;}}}⋯ 1 unchanged linesasm volatile("s_waitcnt vmcnt(0)");__syncthreads();- // MFMA — B reads use swizzle (data was loaded swizzled)+ // MFMA#pragma unrollfor (int half = 0; half < 2; half++) {const int kh = half * HALF_K;⋯ 6 unchanged linesint4_vec B_reg;{const int b_row = (warp_id << 4) + lane_m;- const int b_off = lds_swz(b_row * LDS_ROW + kh + (lane_k << 4)); // swizzled read!+ const int b_off = lds_swz(b_row * LDS_ROW + kh + (lane_k << 4));const int4 tmp = *reinterpret_cast<const int4*>(&B_lds[b_off]);B_reg.s0 = tmp.x; B_reg.s1 = tmp.y; B_reg.s2 = tmp.z; B_reg.s3 = tmp.w;}⋯ 40 unchanged linesC[idx] = __float2bfloat16(sum);}- // ---- Launch functions ----- void launch_sep_gemm_splitk(+ // ---- Launch functions for BLOCK_N=64 (4 warps, 256 threads) ----+ void launch_n64_nosplit(torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,+ torch::Tensor C, int M, int N, int K+ ) {+ const int k_steps = K / (MFMA_K * 2);+ dim3 block(256);+ dim3 grid((N + 63) / 64, (M + BLOCK_M - 1) / BLOCK_M, 1);+ hipLaunchKernelGGL((gemm_kernel<64, 4>), grid, block, 0, 0,+ reinterpret_cast<const __hip_bfloat16*>(A_bf16.data_ptr()),+ reinterpret_cast<const uint8_t*>(B_q.data_ptr()),+ reinterpret_cast<const uint8_t*>(B_scale.data_ptr()),+ (float*)nullptr,+ reinterpret_cast<__hip_bfloat16*>(C.data_ptr()), M, N, K, k_steps);+ }++ void launch_n64_splitk(+ torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,torch::Tensor workspace, int M, int N, int K, int split_k) {const int k_steps = K / (MFMA_K * 2);- dim3 block(NUM_THREADS);- dim3 grid((N + BLOCK_N - 1) / BLOCK_N, (M + BLOCK_M - 1) / BLOCK_M, split_k);- hipLaunchKernelGGL(gemm_kernel, grid, block, 0, 0,+ dim3 block(256);+ dim3 grid((N + 63) / 64, (M + BLOCK_M - 1) / BLOCK_M, split_k);+ hipLaunchKernelGGL((gemm_kernel<64, 4>), grid, block, 0, 0,reinterpret_cast<const __hip_bfloat16*>(A_bf16.data_ptr()),reinterpret_cast<const uint8_t*>(B_q.data_ptr()),reinterpret_cast<const uint8_t*>(B_scale.data_ptr()),⋯ 1 unchanged lines(__hip_bfloat16*)nullptr, M, N, K, k_steps / split_k);}- void launch_sep_gemm_nosplit(+ // ---- Launch functions for BLOCK_N=128 (8 warps, 512 threads) ----+ void launch_n128_nosplit(torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,torch::Tensor C, int M, int N, int K) {const int k_steps = K / (MFMA_K * 2);- dim3 block(NUM_THREADS);- dim3 grid((N + BLOCK_N - 1) / BLOCK_N, (M + BLOCK_M - 1) / BLOCK_M, 1);- hipLaunchKernelGGL(gemm_kernel, grid, block, 0, 0,+ dim3 block(512);+ dim3 grid((N + 127) / 128, (M + BLOCK_M - 1) / BLOCK_M, 1);+ hipLaunchKernelGGL((gemm_kernel<128, 8>), grid, block, 0, 0,reinterpret_cast<const __hip_bfloat16*>(A_bf16.data_ptr()),reinterpret_cast<const uint8_t*>(B_q.data_ptr()),reinterpret_cast<const uint8_t*>(B_scale.data_ptr()),⋯ 1 unchanged linesreinterpret_cast<__hip_bfloat16*>(C.data_ptr()), M, N, K, k_steps);}+ void launch_n128_splitk(+ torch::Tensor A_bf16, torch::Tensor B_q, torch::Tensor B_scale,+ torch::Tensor workspace, int M, int N, int K, int split_k+ ) {+ const int k_steps = K / (MFMA_K * 2);+ dim3 block(512);+ dim3 grid((N + 127) / 128, (M + BLOCK_M - 1) / BLOCK_M, split_k);+ hipLaunchKernelGGL((gemm_kernel<128, 8>), grid, block, 0, 0,+ reinterpret_cast<const __hip_bfloat16*>(A_bf16.data_ptr()),+ reinterpret_cast<const uint8_t*>(B_q.data_ptr()),+ reinterpret_cast<const uint8_t*>(B_scale.data_ptr()),+ reinterpret_cast<float*>(workspace.data_ptr()),+ (__hip_bfloat16*)nullptr, M, N, K, k_steps / split_k);+ }+void launch_reduce(torch::Tensor workspace, torch::Tensor C, int M, int N, int split_k) {const int num = M * N;hipLaunchKernelGGL(reduce_kernel, dim3((num+255)/256), dim3(256), 0, 0,⋯ 10 unchanged linesglobal _moduleif _module is None:_module = load_inline(- name="hybrid_v59",+ name="hybrid_v60",cpp_sources=CPP_SOURCE,cuda_sources=HIP_SOURCE,functions=[- "launch_sep_gemm_splitk", "launch_sep_gemm_nosplit",+ "launch_n64_nosplit", "launch_n64_splitk",+ "launch_n128_nosplit", "launch_n128_splitk","launch_reduce",],verbose=False,⋯ 1 unchanged lines)return _module- def _pick_split_k(m, n, k):+ def _pick_split_k(m, n, k, block_n):k_steps = k // 256- blocks_mn = ((n + 63) // 64) * ((m + 15) // 16)+ blocks_mn = ((n + block_n - 1) // block_n) * ((m + 15) // 16)if blocks_mn >= 256:return 1target_split = max(1, (608 + blocks_mn - 1) // blocks_mn)⋯ 16 unchanged linesB_sh_u8 = B_shuffle.contiguous().view(torch.uint8)B_sc = B_scale_sh.contiguous().view(torch.uint8)- split_k = _pick_split_k(m, n, k)+ # Dispatch: use BLOCK_N=128 for large-M benchmarks (M>=64)+ use_n128 = (m >= 64)- if split_k == 1:- C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")- mod.launch_sep_gemm_nosplit(A, B_sh_u8, B_sc, C, m, n, k)+ if use_n128:+ block_n = 128+ split_k = _pick_split_k(m, n, k, block_n)+ if split_k == 1:+ C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")+ mod.launch_n128_nosplit(A, B_sh_u8, B_sc, C, m, n, k)+ else:+ workspace = torch.empty((split_k, m, n), dtype=torch.float32, device="cuda")+ mod.launch_n128_splitk(A, B_sh_u8, B_sc, workspace, m, n, k, split_k)+ C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")+ mod.launch_reduce(workspace, C, m, n, split_k)else:- workspace = torch.empty((split_k, m, n), dtype=torch.float32, device="cuda")- mod.launch_sep_gemm_splitk(A, B_sh_u8, B_sc, workspace, m, n, k, split_k)- C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")- mod.launch_reduce(workspace, C, m, n, split_k)+ block_n = 64+ split_k = _pick_split_k(m, n, k, block_n)+ if split_k == 1:+ C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")+ mod.launch_n64_nosplit(A, B_sh_u8, B_sc, C, m, n, k)+ else:+ workspace = torch.empty((split_k, m, n), dtype=torch.float32, device="cuda")+ mod.launch_n64_splitk(A, B_sh_u8, B_sc, workspace, m, n, k, split_k)+ C = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")+ mod.launch_reduce(workspace, C, m, n, split_k)return C
scrolls · 316 diff lines total
Best evidence level for this revision: reported
JSON