submission 666540
Jingze · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 2559 lines, June 9 Researcher Reciprocity License v1.0.
submission_gluon.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-666540?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:31ef14c214dae0d9e19b647788533d77babef9d716583cb157d169c3468b2847
license declaredunknown
license concludedunknown
authorsJingze
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
Converts given x (in fp32) to mxfp4 format.num-warps = 1
NUM_WARPS = 1shared-memory
smem_a = gl.allocate_shared_memory(split-k
SPLITK_BLOCK_SIZE: gl.constexpr,stages = 1
NUM_STAGES = 1tile-m = 64
BLOCK_SIZE_M = 64tile-n = 32
BLOCK_SIZE_N = 32Kernel source
submission_gluon.py2559 lines
from typing import Optional, Any
import functools
import json
import math
import os
import time
# os.environ.setdefault("TORCH_EXTENSIONS_DIR", "/tmp/torch_extensions")
# os.environ.setdefault("MAX_JOBS", "16")
import torch
import triton
import triton.language as tl
from triton.experimental import gluon
from triton.experimental.gluon import language as gl
MXFP4_GROUP_SIZE = 32
_HIP_QUANT_INLINE_MODULE = None
_HIP_QUANT_INLINE_LOAD_ERROR = None
def _inline_quant_backend() -> Optional[str]:
if not torch.cuda.is_available():
return None
if torch.version.hip is not None:
try:
device_name = torch.cuda.get_device_name(torch.cuda.current_device()).lower()
except Exception:
return None
if any(token in device_name for token in ("amd", "radeon", "instinct", "mi")):
return "hip"
if torch.version.cuda is not None:
return "cuda"
return None
def _dynamic_mxfp4_quant_inline_sources() -> tuple[str, str]:
cpp_source = """
#include <torch/extension.h>
#include <vector>
std::vector<torch::Tensor> dynamic_mxfp4_quantize(torch::Tensor x);
"""
gpu_source = r"""
#include <ATen/ATen.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/util/Exception.h>
#include <c10/util/BFloat16.h>
#include <cmath>
#include <cstdint>
#include <stdexcept>
#include <vector>
#if defined(__HIP_PLATFORM_AMD__) || defined(USE_ROCM)
#include <hip/hip_runtime.h>
#define MXFP4_GPU_BACKEND_HIP 1
#else
#include <cuda_runtime.h>
#define MXFP4_GPU_BACKEND_CUDA 1
#endif
namespace {
#define MX_CAT2(a, b) a##b
#define MX_CAT3(a, b, c) a##b##c
#if defined(MXFP4_GPU_BACKEND_CUDA)
using mx_queue_t = MX_CAT2(cudaSt, ream_t);
__device__ __forceinline__ float warp_shuffle_down(float value, int offset) {
return __shfl_down_sync(0xFFFFFFFFu, value, offset, 32);
}
__device__ __forceinline__ float warp_shuffle(float value, int src_lane) {
return __shfl_sync(0xFFFFFFFFu, value, src_lane, 32);
}
__device__ __forceinline__ uint32_t warp_shuffle_xor_u32(uint32_t value, int lane_mask) {
return __shfl_xor_sync(0xFFFFFFFFu, value, lane_mask, 32);
}
__device__ __forceinline__ uint8_t warp_shuffle_xor(uint8_t value, int lane_mask) {
return static_cast<uint8_t>(
__shfl_xor_sync(0xFFFFFFFFu, static_cast<unsigned int>(value), lane_mask, 32)
);
}
#else
using mx_queue_t = MX_CAT2(hipSt, ream_t);
__device__ __forceinline__ float warp_shuffle_down(float value, int offset) {
return __shfl_down(value, offset, 32);
}
__device__ __forceinline__ float warp_shuffle(float value, int src_lane) {
return __shfl(value, src_lane, 32);
}
__device__ __forceinline__ uint32_t warp_shuffle_xor_u32(uint32_t value, int lane_mask) {
return static_cast<uint32_t>(__shfl_xor(value, lane_mask, 32));
}
__device__ __forceinline__ uint8_t warp_shuffle_xor(uint8_t value, int lane_mask) {
return static_cast<uint8_t>(__shfl_xor(value, lane_mask, 32));
}
#endif
__device__ __forceinline__ float bf16_to_float(const uint16_t v) {
return __uint_as_float(static_cast<uint32_t>(v) << 16);
}
__device__ __forceinline__ int clamp_int(const int value, const int lo, const int hi) {
return value < lo ? lo : (value > hi ? hi : value);
}
__device__ __forceinline__ uint32_t float_to_bits(float value) {
return __float_as_uint(value);
}
__device__ __forceinline__ float bits_to_float(uint32_t value) {
return __uint_as_float(value);
}
__device__ __forceinline__ float exact_pow2_neg_exp(int scale_unbiased) {
return bits_to_float(static_cast<uint32_t>(127 - scale_unbiased) << 23);
}
__device__ __forceinline__ uint8_t quantize_mxfp4_bits(float value, float quant_scale) {
constexpr uint32_t FP32_SIGN_MASK = 0x80000000u;
constexpr uint32_t FP32_ONE_BITS = 0x3F800000u;
constexpr uint32_t FP32_SIX_BITS = 0x40C00000u;
constexpr int MBITS_F32 = 23;
constexpr int MBITS_FP4 = 1;
constexpr int EBITS_F32 = 8;
constexpr int EBITS_FP4 = 2;
constexpr int EXP_BIAS_FP32 = 127;
constexpr int EXP_BIAS_FP4 = 1;
constexpr int DENORM_EXP = (EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1;
constexpr uint32_t DENORM_MASK_INT = static_cast<uint32_t>(DENORM_EXP) << MBITS_F32;
constexpr uint32_t VAL_TO_ADD =
(static_cast<uint32_t>(EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + ((1u << 21) - 1u);
const float DENORM_MASK_FLOAT = bits_to_float(DENORM_MASK_INT);
const float scaled = value * quant_scale;
const uint32_t scaled_bits = float_to_bits(scaled);
const uint32_t sign = scaled_bits & FP32_SIGN_MASK;
const uint32_t abs_bits = scaled_bits ^ sign;
uint8_t e2m1 = 0x7;
if (abs_bits < FP32_SIX_BITS) {
if (abs_bits < FP32_ONE_BITS) {
const float abs_value = bits_to_float(abs_bits);
uint32_t denormal_x = float_to_bits(abs_value + DENORM_MASK_FLOAT);
denormal_x -= DENORM_MASK_INT;
e2m1 = static_cast<uint8_t>(denormal_x);
} else {
const uint32_t mant_odd = (abs_bits >> (MBITS_F32 - MBITS_FP4)) & 1u;
const uint32_t normal_u32 = abs_bits + VAL_TO_ADD + mant_odd;
e2m1 = static_cast<uint8_t>(normal_u32 >> (MBITS_F32 - MBITS_FP4));
}
}
const uint8_t sign_lp = static_cast<uint8_t>(sign >> (MBITS_F32 + EBITS_F32 - MBITS_FP4 - EBITS_FP4));
return static_cast<uint8_t>(e2m1 | sign_lp);
}
__global__ void dynamic_mxfp4_quantize_kernel(
const uint16_t* __restrict__ x,
uint8_t* __restrict__ x_fp4,
uint8_t* __restrict__ scales,
const int m,
const int n,
const int packed_n,
const int quant_blocks
) {
const int lane = threadIdx.x;
const int warp_in_cta = threadIdx.y;
const int block_n = blockIdx.x * blockDim.y + warp_in_cta;
const int row = blockIdx.y;
if (row >= m || block_n >= quant_blocks) {
return;
}
const int base_n = block_n * 32;
const float value = bf16_to_float(x[row * n + base_n + lane]);
float abs_value = fabsf(value);
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
abs_value = fmaxf(abs_value, warp_shuffle_down(abs_value, offset));
}
const float amax = warp_shuffle(abs_value, 0);
uint32_t amax_bits = float_to_bits(amax);
amax_bits = (amax_bits + 0x200000u) & 0xFF800000u;
int scale_unbiased = -127;
if (amax_bits != 0) {
scale_unbiased = static_cast<int>((amax_bits >> 23) & 0xFFu) - 127 - 2;
scale_unbiased = clamp_int(scale_unbiased, -127, 127);
}
const float quant_scale = exact_pow2_neg_exp(scale_unbiased);
const uint8_t q = quantize_mxfp4_bits(value, quant_scale);
if ((lane & 1) == 0) {
const float partner_value = bf16_to_float(x[row * n + base_n + lane + 1]);
const uint8_t q_hi = quantize_mxfp4_bits(partner_value, quant_scale);
const int packed_col = block_n * 16 + (lane >> 1);
x_fp4[row * packed_n + packed_col] = static_cast<uint8_t>(q | (q_hi << 4));
}
if (lane == 0) {
#if defined(MXFP4_GPU_BACKEND_HIP)
const int stored_scale_unbiased = scale_unbiased;
#else
const int stored_scale_unbiased = scale_unbiased > 0 ? scale_unbiased : 0;
#endif
scales[row * quant_blocks + block_n] = static_cast<uint8_t>(stored_scale_unbiased + 127);
}
}
} // namespace
std::vector<at::Tensor> dynamic_mxfp4_quantize(at::Tensor x) {
const auto m = static_cast<int>(x.size(0));
const auto n = static_cast<int>(x.size(1));
const int packed_n = n / 2;
const int quant_blocks = n / 32;
const int groups_per_cta = quant_blocks >= 8 ? 8 : (quant_blocks >= 4 ? 4 : (quant_blocks >= 2 ? 2 : 1));
auto options = at::TensorOptions().device(x.device()).dtype(at::kByte);
auto x_fp4 = at::empty({m, packed_n}, options);
auto scales = at::empty({m, quant_blocks}, options);
dim3 block(32, groups_per_cta);
dim3 grid((quant_blocks + groups_per_cta - 1) / groups_per_cta, m);
#if defined(MXFP4_GPU_BACKEND_HIP)
auto q = at::cuda::MX_CAT3(getCurrentHIPSt, ream, )();
hipLaunchKernelGGL(
dynamic_mxfp4_quantize_kernel,
grid,
block,
0,
static_cast<mx_queue_t>(q),
reinterpret_cast<const uint16_t*>(x.data_ptr<c10::BFloat16>()),
x_fp4.data_ptr<uint8_t>(),
scales.data_ptr<uint8_t>(),
m,
n,
packed_n,
quant_blocks
);
const auto err = hipGetLastError();
if (err != hipSuccess) {
throw std::runtime_error(hipGetErrorString(err));
}
#else
auto q = at::cuda::MX_CAT3(getCurrentCUDASt, ream, )();
dynamic_mxfp4_quantize_kernel<<<grid, block, 0, static_cast<mx_queue_t>(q)>>>(
reinterpret_cast<const uint16_t*>(x.data_ptr<c10::BFloat16>()),
x_fp4.data_ptr<uint8_t>(),
scales.data_ptr<uint8_t>(),
m,
n,
packed_n,
quant_blocks
);
const auto err = cudaGetLastError();
if (err != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(err));
}
#endif
return {x_fp4, scales};
}
"""
return cpp_source, gpu_source
def _load_dynamic_mxfp4_quant_inline_module():
global _HIP_QUANT_INLINE_MODULE, _HIP_QUANT_INLINE_LOAD_ERROR
if _HIP_QUANT_INLINE_MODULE is not None:
return _HIP_QUANT_INLINE_MODULE
if _HIP_QUANT_INLINE_LOAD_ERROR is not None:
raise RuntimeError(_HIP_QUANT_INLINE_LOAD_ERROR)
backend = _inline_quant_backend()
if backend is None:
_HIP_QUANT_INLINE_LOAD_ERROR = "GPU inline quantization is unavailable in the current environment"
raise RuntimeError(_HIP_QUANT_INLINE_LOAD_ERROR)
try:
from torch.utils.cpp_extension import load_inline
cpp_source, gpu_source = _dynamic_mxfp4_quant_inline_sources()
extra_cuda_cflags = ["-O3", "-std=c++17"]
if backend == "cuda":
extra_cuda_cflags.append("--use_fast_math")
elif backend == "hip":
extra_cuda_cflags.append("-ffast-math")
_HIP_QUANT_INLINE_MODULE = load_inline(
name=f"mxfp4_dynamic_quant_inline_v4_{backend}",
cpp_sources=cpp_source,
cuda_sources=gpu_source,
functions=["dynamic_mxfp4_quantize"],
extra_cflags=["-O3", "-std=c++17"],
extra_cuda_cflags=extra_cuda_cflags,
with_cuda=True,
verbose=False,
)
return _HIP_QUANT_INLINE_MODULE
except Exception as exc:
_HIP_QUANT_INLINE_LOAD_ERROR = str(exc)
raise RuntimeError(_HIP_QUANT_INLINE_LOAD_ERROR) from exc
def dynamic_mxfp4_quant_inline(
x: torch.Tensor, scaling_mode: str = "even"
) -> tuple[torch.Tensor, torch.Tensor]:
if scaling_mode != "even":
raise NotImplementedError(f"Unsupported scaling_mode: {scaling_mode}")
x_input = x.contiguous()
if x_input.dtype != torch.bfloat16:
x_input = x_input.to(torch.bfloat16)
module = _load_dynamic_mxfp4_quant_inline_module()
x_fp4, blockscale_e8m0 = module.dynamic_mxfp4_quantize(x_input)
return x_fp4, blockscale_e8m0
_DEFAULT_GEMM_CONFIGS = {
"M_LEQ_8": {
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 4,
"num_warps": 2,
"num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
},
"M_LEQ_31": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 4,
"num_warps": 2,
"num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
},
"M_LEQ_32": {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 4,
"num_warps": 2,
"num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
},
"M_LEQ_64": {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 4,
"num_warps": 2,
"num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
},
"M_LEQ_128": {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 4,
"num_warps": 2,
"num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
},
"M_LEQ_256": {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 4,
"num_warps": 2,
"num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
},
"any": {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 4,
"num_warps": 2,
"num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
},
}
_BENCHMARK_GEMM_CONFIGS = {
# (4, 2880, 512): {
# "BLOCK_SIZE_M": 4,
# "BLOCK_SIZE_N": 64,
# "BLOCK_SIZE_K": 512,
# "GROUP_SIZE_M": 1,
# "num_warps": 2,
# "num_stages": 2,
# "waves_per_eu": 3,
# "matrix_instr_nonkdim": 16,
# "cache_modifier": None,
# "NUM_KSPLIT": 1,
# }, # 10.1 ± 0.02 µs
(4, 2880, 512): {
"BLOCK_SIZE_M": 4,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 3,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
(16, 2112, 7168): {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 7,
},
(32, 4096, 512): {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
(32, 2880, 512): {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
(64, 7168, 2048): {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 1024,
"GROUP_SIZE_M": 1,
"num_warps": 2,
"num_stages": 2,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
(256, 3072, 1536): {
"BLOCK_SIZE_M": 128,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
}
def _select_default_config_by_m(M: int):
if M <= 8:
return dict(_DEFAULT_GEMM_CONFIGS["M_LEQ_8"])
if M <= 31:
return dict(_DEFAULT_GEMM_CONFIGS["M_LEQ_31"])
if M <= 32:
return dict(_DEFAULT_GEMM_CONFIGS["M_LEQ_32"])
if M <= 64:
return dict(_DEFAULT_GEMM_CONFIGS["M_LEQ_64"])
if M <= 128:
return dict(_DEFAULT_GEMM_CONFIGS["M_LEQ_128"])
if M <= 256:
return dict(_DEFAULT_GEMM_CONFIGS["M_LEQ_256"])
return dict(_DEFAULT_GEMM_CONFIGS["any"])
@triton.jit
def _mxfp4_quant_op(
x,
BLOCK_SIZE_N,
BLOCK_SIZE_M,
MXFP4_QUANT_BLOCK_SIZE,
):
"""
Converts given x (in fp32) to mxfp4 format.
x: [BLOCK_SIZE_M, BLOCK_SIZE_N], fp32
"""
EXP_BIAS_FP32: tl.constexpr = 127
EXP_BIAS_FP4: tl.constexpr = 1
EBITS_F32: tl.constexpr = 8
EBITS_FP4: tl.constexpr = 2
MBITS_F32: tl.constexpr = 23
MBITS_FP4: tl.constexpr = 1
max_normal: tl.constexpr = 6
min_normal: tl.constexpr = 1
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
# Calculate scale
amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
amax = amax.to(tl.int32, bitcast=True)
amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = amax.to(tl.float32, bitcast=True)
scale_e8m0_unbiased = tl.log2(amax).floor() - 2
scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
# blockscale_e8m0
bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127 # in fp32, we have 2&(e - 127)
quant_scale = tl.exp2(-scale_e8m0_unbiased)
# Compute quantized x
qx = x * quant_scale
# Convert quantized fp32 tensor to uint32 before converting to mxfp4 format
# Note: MXFP4 S:1-bit, E:2-bit, M:1-bit
# Zeros: S000 -> +/-0
# Denormal Numbers: S001 -> +/- 0.5
# Normal Numbers:
# S010 -> +/- 1.0
# S011 -> +/- 1.5
# S100 -> +/- 2.0
# S101 -> +/- 3.0
# S110 -> +/- 4.0
# S111 -> +/- 6.0
qx = qx.to(tl.uint32, bitcast=True)
# Extract sign
s = qx & 0x80000000
# Set everything to positive, will add sign back at the end
qx = qx ^ s
qx_fp32 = qx.to(tl.float32, bitcast=True)
saturate_mask = qx_fp32 >= max_normal
denormal_mask = (not saturate_mask) & (qx_fp32 < min_normal)
normal_mask = not (saturate_mask | denormal_mask)
# Denormal numbers
denorm_exp: tl.constexpr = (
(EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1
)
denorm_mask_int: tl.constexpr = denorm_exp << MBITS_F32
denorm_mask_float: tl.constexpr = tl.cast(denorm_mask_int, tl.float32, bitcast=True)
denormal_x = qx_fp32 + denorm_mask_float
denormal_x = denormal_x.to(tl.uint32, bitcast=True)
denormal_x -= denorm_mask_int
denormal_x = denormal_x.to(tl.uint8)
# Normal numbers
normal_x = qx
# resulting mantissa is odd
mant_odd = (normal_x >> (MBITS_F32 - MBITS_FP4)) & 1
# update exponent, rounding bias part 1
val_to_add = ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1 << 21) - 1
normal_x += val_to_add
# rounding bias part 2
normal_x += mant_odd
# take the bits!
normal_x = normal_x >> (MBITS_F32 - MBITS_FP4)
normal_x = normal_x.to(tl.uint8)
# Merge results
e2m1_value = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
e2m1_value = tl.where(normal_mask, normal_x, e2m1_value)
e2m1_value = tl.where(denormal_mask, denormal_x, e2m1_value)
# add sign back
sign_lp = s >> (MBITS_F32 + EBITS_F32 - MBITS_FP4 - EBITS_FP4)
sign_lp = sign_lp.to(tl.uint8)
e2m1_value = e2m1_value | sign_lp
e2m1_value = tl.reshape(
e2m1_value, [BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2]
)
evens, odds = tl.split(e2m1_value)
x_fp4 = evens | (odds << 4)
x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
@triton.jit
def _dynamic_mxfp4_quant_kernel(
x_ptr,
x_fp4_ptr,
bs_ptr,
stride_x_m_in,
stride_x_n_in,
stride_x_fp4_m_in,
stride_x_fp4_n_in,
stride_bs_m_in,
stride_bs_n_in,
M,
N,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
NUM_ITER: tl.constexpr,
NUM_STAGES: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
SCALING_MODE: tl.constexpr,
num_warps: tl.constexpr,
waves_per_eu: tl.constexpr,
num_stages: tl.constexpr,
):
pid_m = tl.program_id(0)
start_n = tl.program_id(1) * NUM_ITER
# cast strides to int64, in case M*N > max int32
stride_x_m = tl.cast(stride_x_m_in, tl.int64)
stride_x_n = tl.cast(stride_x_n_in, tl.int64)
stride_x_fp4_m = tl.cast(stride_x_fp4_m_in, tl.int64)
stride_x_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)
stride_bs_m = tl.cast(stride_bs_m_in, tl.int64)
stride_bs_n = tl.cast(stride_bs_n_in, tl.int64)
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
x = tl.load(x_ptr + x_offs, cache_modifier=".cg").to(tl.float32)
out_tensor, bs_e8m0 = _mxfp4_quant_op(
x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE
)
out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
out_offs = (
out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
)
tl.store(x_fp4_ptr + out_offs, out_tensor)
bs_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
bs_offs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
bs_offs = bs_offs_m[:, None] * stride_bs_m + bs_offs_n[None, :] * stride_bs_n
tl.store(bs_ptr + bs_offs, bs_e8m0)
def dynamic_mxfp4_quant(
x: torch.Tensor, scaling_mode: str = "even"
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Quantize a tensor to MX FP4 format.
Args:
x: The input tensor, typically fp16 or bf16.
scaling_mode: The method to calculate MX block scaling.
- "even" (default): `even_round` in `quark.torch.quantization.utils`.
- etc.
Returns:
A tuple of (x_fp4, blockscale_e8m0).
"""
# Assume x is 2D-Tensor for now
M, N = x.shape
assert (N // 2) % 2 == 0
# This is fixed by spec for MXFP4. Do not tune this.
MXFP4_QUANT_BLOCK_SIZE = 32
x_fp4 = torch.empty((M, N // 2), dtype=torch.uint8, device=x.device)
blockscale_e8m0 = torch.empty(
((N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE, M),
dtype=torch.uint8,
device=x.device,
).T
# for large N values
if M <= 32:
NUM_ITER = 1
BLOCK_SIZE_M = triton.next_power_of_2(M)
BLOCK_SIZE_N = 32
NUM_WARPS = 1
NUM_STAGES = 1
else:
NUM_ITER = 4
BLOCK_SIZE_M = 64
BLOCK_SIZE_N = 64
NUM_WARPS = 4
NUM_STAGES = 2
if N <= 16384:
BLOCK_SIZE_M = 32
BLOCK_SIZE_N = 128
# for small N values
if N <= 1024:
NUM_ITER = 1
NUM_STAGES = 1
NUM_WARPS = 4
BLOCK_SIZE_N = min(256, triton.next_power_of_2(N))
# BLOCK_SIZE_N needs to be multiple of 32
BLOCK_SIZE_N = max(32, BLOCK_SIZE_N)
BLOCK_SIZE_M = min(8, triton.next_power_of_2(M))
grid = (
triton.cdiv(M, BLOCK_SIZE_M),
triton.cdiv(N, BLOCK_SIZE_N * NUM_ITER),
)
_dynamic_mxfp4_quant_kernel[grid](
x,
x_fp4,
blockscale_e8m0,
*x.stride(),
*x_fp4.stride(),
*blockscale_e8m0.stride(),
M=M,
N=N,
MXFP4_QUANT_BLOCK_SIZE=MXFP4_QUANT_BLOCK_SIZE,
SCALING_MODE=0,
NUM_ITER=NUM_ITER,
BLOCK_SIZE_M=BLOCK_SIZE_M,
BLOCK_SIZE_N=BLOCK_SIZE_N,
NUM_STAGES=NUM_STAGES,
num_warps=NUM_WARPS,
waves_per_eu=0,
num_stages=1,
)
return (x_fp4, blockscale_e8m0)
@triton.jit
def pid_grid(pid: int, num_pid_m: int, num_pid_n: int, GROUP_SIZE_M: tl.constexpr = 1):
"""
Maps 1D pid to 2D grid coords (pid_m, pid_n).
Args:
- pid: 1D pid
- num_pid_m: grid m size
- num_pid_n: grid n size
- GROUP_SIZE_M: tl.constexpr: default is 1
"""
if GROUP_SIZE_M == 1:
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
else:
num_pid_in_group = GROUP_SIZE_M * num_pid_n
group_id = pid // num_pid_in_group
first_pid_m = group_id * GROUP_SIZE_M
group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
tl.assume(group_size_m >= 0)
pid_m = first_pid_m + (pid % group_size_m)
pid_n = (pid % num_pid_in_group) // group_size_m
return pid_m, pid_n
@triton.jit
def remap_xcd(pid, GRID_MN, NUM_XCDS: tl.constexpr = 8):
## pid remapping on xcds
# Number of pids per XCD in the new arrangement
pids_per_xcd = (GRID_MN + NUM_XCDS - 1) // NUM_XCDS
# When GRID_MN cannot divide NUM_XCDS, some xcds will have
# pids_per_xcd pids, the other will have pids_per_xcd - 1 pids.
# We calculate the number of xcds that have pids_per_xcd pids as
# tall_xcds
tall_xcds = GRID_MN % NUM_XCDS
tall_xcds = NUM_XCDS if tall_xcds == 0 else tall_xcds
# Compute current XCD and local pid within the XCD
xcd = pid % NUM_XCDS
local_pid = pid // NUM_XCDS
# Calculate new pid based on the new grouping
# Note that we need to consider the following two cases:
# 1. the current pid is on a tall xcd
# 2. the current pid is on a short xcd
if xcd < tall_xcds:
pid = xcd * pids_per_xcd + local_pid
else:
pid = (
tall_xcds * pids_per_xcd
+ (xcd - tall_xcds) * (pids_per_xcd - 1)
+ local_pid
)
return pid
@gluon.jit
def _gemm_afp4wfp4_kernel(
a_ptr,
b_ptr,
c_ptr,
a_scales_ptr,
b_scales_ptr,
M,
N,
K,
stride_am,
stride_ak,
stride_bk,
stride_bn,
stride_ck,
stride_cm,
stride_cn,
stride_asm,
stride_ask,
stride_bsn,
stride_bsk,
# Meta-parameters
BLOCK_SIZE_M: gl.constexpr,
BLOCK_SIZE_N: gl.constexpr,
BLOCK_SIZE_K: gl.constexpr,
GROUP_SIZE_M: gl.constexpr,
NUM_KSPLIT: gl.constexpr,
SPLITK_BLOCK_SIZE: gl.constexpr,
num_warps: gl.constexpr,
num_stages: gl.constexpr,
waves_per_eu: gl.constexpr,
matrix_instr_nonkdim: gl.constexpr,
cache_modifier: gl.constexpr,
):
"""
Kernel for computing the matmul C = A x B.
A and B inputs are in the microscale fp4 (mxfp4) format.
A_scales and B_scales are in e8m0 format.
A has shape (M, K), B has shape (K, N) and C has shape (M, N)
"""
GRID_MN = gl.cdiv(M, BLOCK_SIZE_M) * gl.cdiv(N, BLOCK_SIZE_N)
# Grouped and XCD-remapped launch ordering improves L2 residency.
pid_unified = gl.program_id(axis=0)
pid_unified = remap_xcd(pid_unified, GRID_MN * NUM_KSPLIT, NUM_XCDS=8)
pid_k = pid_unified % NUM_KSPLIT
pid = pid_unified // NUM_KSPLIT
num_pid_m = gl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = gl.cdiv(N, BLOCK_SIZE_N)
if NUM_KSPLIT == 1:
pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
else:
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
SCALE_GROUP_SIZE: gl.constexpr = 32
BLOCK_K_PACKED: gl.constexpr = BLOCK_SIZE_K // 2
BLOCK_K_SCALE: gl.constexpr = BLOCK_SIZE_K // SCALE_GROUP_SIZE
NUM_BUFFERS: gl.constexpr = num_stages if num_stages > 1 else 2
gl.static_assert(num_warps % 2 == 0)
gl.static_assert(BLOCK_SIZE_K % SCALE_GROUP_SIZE == 0)
wmma_layout: gl.constexpr = gl.amd.AMDWMMALayout(
version=3,
transposed=True,
warps_per_cta=[2, num_warps // 2],
instr_shape=[16, 16, 128],
)
wmma_packed_layout: gl.constexpr = gl.amd.AMDWMMALayout(
version=3,
transposed=True,
warps_per_cta=[2, num_warps // 2],
instr_shape=[16, 16, 64],
)
dot_a_layout: gl.constexpr = gl.DotOperandLayout(
operand_index=0, parent=wmma_packed_layout, k_width=16
)
dot_b_layout: gl.constexpr = gl.DotOperandLayout(
operand_index=1, parent=wmma_packed_layout, k_width=16
)
scale_a_layout: gl.constexpr = gl.amd.gfx1250.get_wmma_scale_layout(
dot_a_layout, [BLOCK_SIZE_M, BLOCK_K_SCALE]
)
scale_b_layout: gl.constexpr = gl.amd.gfx1250.get_wmma_scale_layout(
dot_b_layout, [BLOCK_SIZE_N, BLOCK_K_SCALE]
)
PAD_INTERVAL_A: gl.constexpr = 256 if BLOCK_K_PACKED <= 256 else BLOCK_K_PACKED
PAD_INTERVAL_B: gl.constexpr = 256 if BLOCK_K_PACKED <= 256 else BLOCK_K_PACKED
shared_layout_a: gl.constexpr = gl.PaddedSharedLayout.with_identity_for(
[[PAD_INTERVAL_A, 16]], [BLOCK_SIZE_M, BLOCK_K_PACKED], [1, 0]
)
shared_layout_b: gl.constexpr = gl.PaddedSharedLayout.with_identity_for(
[[PAD_INTERVAL_B, 16]], [BLOCK_K_PACKED, BLOCK_SIZE_N], [1, 0]
)
shared_layout_as: gl.constexpr = gl.PaddedSharedLayout.with_identity_for(
[[256, 16]], [BLOCK_SIZE_M, BLOCK_K_SCALE], [1, 0]
)
shared_layout_bs: gl.constexpr = gl.PaddedSharedLayout.with_identity_for(
[[256, 16]], [BLOCK_SIZE_N, BLOCK_K_SCALE], [1, 0]
)
split_k_start = pid_k * (SPLITK_BLOCK_SIZE // 2)
split_k_start_scale = pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE)
if split_k_start < K:
valid_packed_k = K - split_k_start
num_k_iter = gl.cdiv(valid_packed_k, BLOCK_K_PACKED)
a_desc = gl.amd.gfx1250.tdm.make_tensor_descriptor(
base=a_ptr + pid_m * BLOCK_SIZE_M * stride_am + split_k_start * stride_ak,
shape=(M, K),
strides=(stride_am, stride_ak),
block_shape=(BLOCK_SIZE_M, BLOCK_K_PACKED),
layout=shared_layout_a,
)
b_desc = gl.amd.gfx1250.tdm.make_tensor_descriptor(
base=b_ptr + pid_n * BLOCK_SIZE_N * stride_bn + split_k_start * stride_bk,
shape=(K, N),
strides=(stride_bk, stride_bn),
block_shape=(BLOCK_K_PACKED, BLOCK_SIZE_N),
layout=shared_layout_b,
)
packed_scale_k = K // (SCALE_GROUP_SIZE // 2)
a_scale_desc = gl.amd.gfx1250.tdm.make_tensor_descriptor(
base=(
a_scales_ptr
+ pid_m * BLOCK_SIZE_M * stride_asm
+ split_k_start_scale * stride_ask
),
shape=(M, packed_scale_k),
strides=(stride_asm, stride_ask),
block_shape=(BLOCK_SIZE_M, BLOCK_K_SCALE),
layout=shared_layout_as,
)
b_scale_desc = gl.amd.gfx1250.tdm.make_tensor_descriptor(
base=(
b_scales_ptr
+ pid_n * BLOCK_SIZE_N * stride_bsn
+ split_k_start_scale * stride_bsk
),
shape=(N, packed_scale_k),
strides=(stride_bsn, stride_bsk),
block_shape=(BLOCK_SIZE_N, BLOCK_K_SCALE),
layout=shared_layout_bs,
)
a_buffer = gl.allocate_shared_memory(
a_desc.dtype, shape=[NUM_BUFFERS] + a_desc.block_shape, layout=a_desc.layout
)
b_buffer = gl.allocate_shared_memory(
b_desc.dtype, shape=[NUM_BUFFERS] + b_desc.block_shape, layout=b_desc.layout
)
a_scale_buffer = gl.allocate_shared_memory(
a_scale_desc.dtype,
shape=[NUM_BUFFERS] + a_scale_desc.block_shape,
layout=a_scale_desc.layout,
)
b_scale_buffer = gl.allocate_shared_memory(
b_scale_desc.dtype,
shape=[NUM_BUFFERS] + b_scale_desc.block_shape,
layout=b_scale_desc.layout,
)
load_idx = 0
wmma_idx = 0
for _ in gl.static_range(NUM_BUFFERS - 1):
if load_idx < num_k_iter:
gl.amd.gfx1250.tdm.async_load(
a_desc, [0, load_idx * BLOCK_K_PACKED], a_buffer.index(load_idx)
)
gl.amd.gfx1250.tdm.async_load(
b_desc, [load_idx * BLOCK_K_PACKED, 0], b_buffer.index(load_idx)
)
gl.amd.gfx1250.tdm.async_load(
a_scale_desc,
[0, load_idx * BLOCK_K_SCALE],
a_scale_buffer.index(load_idx),
)
gl.amd.gfx1250.tdm.async_load(
b_scale_desc,
[0, load_idx * BLOCK_K_SCALE],
b_scale_buffer.index(load_idx),
)
load_idx += 1
accumulator = gl.zeros(
(BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=gl.float32, layout=wmma_layout
)
for _ in range(0, num_k_iter - (NUM_BUFFERS - 1)):
gl.amd.gfx1250.tdm.async_load(
a_desc,
[0, load_idx * BLOCK_K_PACKED],
a_buffer.index(load_idx % NUM_BUFFERS),
)
gl.amd.gfx1250.tdm.async_load(
b_desc,
[load_idx * BLOCK_K_PACKED, 0],
b_buffer.index(load_idx % NUM_BUFFERS),
)
gl.amd.gfx1250.tdm.async_load(
a_scale_desc,
[0, load_idx * BLOCK_K_SCALE],
a_scale_buffer.index(load_idx % NUM_BUFFERS),
)
gl.amd.gfx1250.tdm.async_load(
b_scale_desc,
[0, load_idx * BLOCK_K_SCALE],
b_scale_buffer.index(load_idx % NUM_BUFFERS),
)
load_idx += 1
gl.amd.gfx1250.tdm.async_wait((NUM_BUFFERS - 1) * 4)
a = a_buffer.index(wmma_idx % NUM_BUFFERS).load(layout=dot_a_layout)
b = b_buffer.index(wmma_idx % NUM_BUFFERS).load(layout=dot_b_layout)
scale_a = a_scale_buffer.index(wmma_idx % NUM_BUFFERS).load(
layout=scale_a_layout
)
scale_b = b_scale_buffer.index(wmma_idx % NUM_BUFFERS).load(
layout=scale_b_layout
)
accumulator = gl.amd.gfx1250.wmma_scaled(
a,
scale_a,
"e2m1",
b,
scale_b,
"e2m1",
accumulator,
)
wmma_idx += 1
for i in gl.static_range(NUM_BUFFERS - 1):
if wmma_idx < num_k_iter:
gl.amd.gfx1250.tdm.async_wait((NUM_BUFFERS - 2 - i) * 4)
a = a_buffer.index(wmma_idx % NUM_BUFFERS).load(layout=dot_a_layout)
b = b_buffer.index(wmma_idx % NUM_BUFFERS).load(layout=dot_b_layout)
scale_a = a_scale_buffer.index(wmma_idx % NUM_BUFFERS).load(
layout=scale_a_layout
)
scale_b = b_scale_buffer.index(wmma_idx % NUM_BUFFERS).load(
layout=scale_b_layout
)
accumulator = gl.amd.gfx1250.wmma_scaled(
a,
scale_a,
"e2m1",
b,
scale_b,
"e2m1",
accumulator,
)
wmma_idx += 1
c = accumulator.to(c_ptr.type.element_ty)
offs_cm = pid_m * BLOCK_SIZE_M + gl.arange(
0, BLOCK_SIZE_M, layout=gl.SliceLayout(1, wmma_layout)
)
offs_cn = pid_n * BLOCK_SIZE_N + gl.arange(
0, BLOCK_SIZE_N, layout=gl.SliceLayout(0, wmma_layout)
)
offs_c = (
stride_cm * offs_cm[:, None]
+ stride_cn * offs_cn[None, :]
+ pid_k * stride_ck
)
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
gl.amd.gfx1250.buffer_store(c, c_ptr, offs_c, c_mask)
@triton.jit
def _gemm_afp4wfp4_triton_kernel(
a_ptr,
b_ptr,
c_ptr,
a_scales_ptr,
b_scales_ptr,
M: tl.constexpr,
N: tl.constexpr,
K: tl.constexpr,
stride_am,
stride_ak,
stride_bk,
stride_bn,
stride_ck,
stride_cm,
stride_cn,
stride_asm,
stride_ask,
stride_bsn,
stride_bsk,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr,
NUM_KSPLIT: tl.constexpr,
SPLITK_BLOCK_SIZE: tl.constexpr,
num_warps: tl.constexpr,
num_stages: tl.constexpr,
waves_per_eu: tl.constexpr,
matrix_instr_nonkdim: tl.constexpr,
cache_modifier: tl.constexpr,
):
"""Triton WMMA-scaled split-K kernel specialized for fixed benchmark shapes."""
GRID_MN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)
pid_unified = tl.program_id(axis=0)
pid_unified = remap_xcd(pid_unified, GRID_MN * NUM_KSPLIT, NUM_XCDS=8)
pid_k = pid_unified % NUM_KSPLIT
pid = pid_unified // NUM_KSPLIT
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
if NUM_KSPLIT == 1:
pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
else:
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
SCALE_GROUP_SIZE: tl.constexpr = 32
BLOCK_K_PACKED: tl.constexpr = BLOCK_SIZE_K // 2
BLOCK_K_SCALE: tl.constexpr = BLOCK_SIZE_K // SCALE_GROUP_SIZE
split_k_start = pid_k * (SPLITK_BLOCK_SIZE // 2)
split_ks_start = pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE)
if split_k_start < K:
num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_K_PACKED)
offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
offs_bn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
offs_am_base = pid_m * BLOCK_SIZE_M
offs_bn_base = pid_n * BLOCK_SIZE_N
offs_cm = offs_am.to(tl.int64)
offs_cn = offs_bn.to(tl.int64)
a_desc = tl.make_tensor_descriptor(
a_ptr,
shape=[M, K],
strides=[stride_am, stride_ak],
block_shape=[BLOCK_SIZE_M, BLOCK_K_PACKED],
)
b_desc = tl.make_tensor_descriptor(
b_ptr,
shape=[N, K],
strides=[stride_bn, stride_bk],
block_shape=[BLOCK_SIZE_N, BLOCK_K_PACKED],
)
acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k_iter in tl.range(0, num_k_iter, num_stages=num_stages):
k_base = split_k_start + k_iter * BLOCK_K_PACKED
ks_base = split_ks_start + k_iter * BLOCK_K_SCALE
a = a_desc.load([offs_am_base, k_base])
b = b_desc.load([offs_bn_base, k_base]).trans(1, 0)
offs_ks = ks_base + tl.arange(0, BLOCK_K_SCALE)
a_scale_ptrs = (
a_scales_ptr
+ offs_am[:, None] * stride_asm
+ offs_ks[None, :] * stride_ask
)
b_scale_ptrs = (
b_scales_ptr
+ offs_bn[:, None] * stride_bsn
+ offs_ks[None, :] * stride_bsk
)
a_scales = tl.load(a_scale_ptrs)
b_scales = tl.load(b_scale_ptrs)
acc = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", acc)
c = acc.to(c_ptr.type.element_ty)
c_ptrs = (
c_ptr
+ stride_cm * offs_cm[:, None]
+ stride_cn * offs_cn[None, :]
+ pid_k * stride_ck
)
tl.store(c_ptrs, c, cache_modifier=".wt")
@triton.jit
def _gemm_afp4wfp4_preshuffle_triton_kernel(
a_ptr,
b_ptr,
c_ptr,
a_scales_ptr,
b_scales_ptr,
M: tl.constexpr,
N: tl.constexpr,
K: tl.constexpr,
stride_am,
stride_ak,
stride_bn,
stride_bk,
stride_ck,
stride_cm,
stride_cn,
stride_asm,
stride_ask,
stride_bsn,
stride_bsk,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr,
NUM_KSPLIT: tl.constexpr,
SPLITK_BLOCK_SIZE: tl.constexpr,
A_SHUFFLED: tl.constexpr,
B_SHUFFLED: tl.constexpr,
ATOMIC_ADD: tl.constexpr,
num_warps: tl.constexpr,
num_stages: tl.constexpr,
waves_per_eu: tl.constexpr,
matrix_instr_nonkdim: tl.constexpr,
cache_modifier: tl.constexpr,
):
"""Triton preshuffle kernel supporting shuffled A/B packed mxfp4 tensors."""
GRID_MN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)
pid_unified = tl.program_id(axis=0)
pid_unified = remap_xcd(pid_unified, GRID_MN * NUM_KSPLIT, NUM_XCDS=8)
pid_k = pid_unified % NUM_KSPLIT
pid = pid_unified // NUM_KSPLIT
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
if NUM_KSPLIT == 1:
pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
else:
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
SCALE_GROUP_SIZE: tl.constexpr = 32
BLOCK_K_PACKED: tl.constexpr = BLOCK_SIZE_K // 2
BLOCK_K_SCALE: tl.constexpr = BLOCK_SIZE_K // SCALE_GROUP_SIZE
packed_k = K
packed_scale_k = K // (SCALE_GROUP_SIZE // 2)
if A_SHUFFLED:
tl.static_assert(BLOCK_SIZE_M % 16 == 0)
if B_SHUFFLED:
tl.static_assert(BLOCK_SIZE_N % 16 == 0)
tl.static_assert(BLOCK_SIZE_K % 256 == 0)
tl.static_assert(BLOCK_K_PACKED % 32 == 0)
split_k_start = pid_k * (SPLITK_BLOCK_SIZE // 2)
split_ks_start = pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE)
if split_k_start < K:
num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_K_PACKED)
offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
offs_bn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
offs_bn_packed = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
offs_bsn = (pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)) % N
offs_am_base = pid_m * BLOCK_SIZE_M
offs_cm = offs_am.to(tl.int64)
offs_cn = offs_bn.to(tl.int64)
a_desc = tl.make_tensor_descriptor(
a_ptr,
shape=[M, K],
strides=[stride_am, stride_ak],
block_shape=[BLOCK_SIZE_M, BLOCK_K_PACKED],
)
acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k_iter in tl.range(0, num_k_iter, num_stages=num_stages):
k_base = split_k_start + k_iter * BLOCK_K_PACKED
ks_base = split_ks_start + k_iter * BLOCK_K_SCALE
a = a_desc.load([offs_am_base, k_base])
if A_SHUFFLED:
a = (
a.reshape(
BLOCK_SIZE_M // 16,
BLOCK_K_PACKED // 32,
2,
16,
16,
)
.permute(0, 3, 1, 2, 4)
.reshape(BLOCK_SIZE_M, BLOCK_K_PACKED)
)
if B_SHUFFLED:
offs_k_shuffle = (k_base * 16) + tl.arange(0, BLOCK_K_PACKED * 16)
b_rows = offs_bn_packed[:, None] * 16 + (
offs_k_shuffle[None, :] // packed_k
)
b_cols = offs_k_shuffle[None, :] % packed_k
b_ptrs = b_ptr + (b_rows * stride_bn + b_cols * stride_bk)
b = (
tl.load(b_ptrs)
.reshape(
1,
BLOCK_SIZE_N // 16,
BLOCK_SIZE_K // 64,
2,
16,
16,
)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_SIZE_N, BLOCK_K_PACKED)
)
else:
b_desc = tl.make_tensor_descriptor(
b_ptr,
shape=[N, K],
strides=[stride_bn, stride_bk],
block_shape=[BLOCK_SIZE_N, BLOCK_K_PACKED],
)
b = b_desc.load([pid_n * BLOCK_SIZE_N, k_base])
b = b.trans(1, 0)
offs_ks = ks_base + tl.arange(0, BLOCK_K_SCALE)
a_scale_ptrs = (
a_scales_ptr
+ offs_am[:, None] * stride_asm
+ offs_ks[None, :] * stride_ask
)
a_scales = tl.load(a_scale_ptrs)
if B_SHUFFLED:
offs_ks_shuffled = (ks_base * 32) + tl.arange(0, BLOCK_K_SCALE * 32)
b_scale_rows = offs_bsn[:, None] * 32 + (
offs_ks_shuffled[None, :] // packed_scale_k
)
b_scale_cols = offs_ks_shuffled[None, :] % packed_scale_k
b_scale_ptrs = (
b_scales_ptr
+ b_scale_rows * stride_bsn
+ b_scale_cols * stride_bsk
)
b_scales = (
tl.load(b_scale_ptrs)
.reshape(
BLOCK_SIZE_N // 32,
BLOCK_K_SCALE // 8,
4,
16,
2,
2,
1,
)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_SIZE_N, BLOCK_K_SCALE)
)
else:
b_scale_ptrs = (
b_scales_ptr
+ offs_bn[:, None] * stride_bsn
+ offs_ks[None, :] * stride_bsk
)
b_scales = tl.load(b_scale_ptrs)
acc = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", acc)
c_ptrs = (
c_ptr
+ stride_cm * offs_cm[:, None]
+ stride_cn * offs_cn[None, :]
+ pid_k * stride_ck
)
if ATOMIC_ADD:
tl.atomic_add(c_ptrs, acc, sem="relaxed")
else:
tl.store(c_ptrs, acc.to(c_ptr.type.element_ty), cache_modifier=".wt")
@gluon.jit
def _gemm_afp4wfp4_preshuffle_kernel(
a_ptr,
b_ptr,
c_ptr,
a_scales_ptr,
b_scales_ptr,
M,
N,
K,
stride_am,
stride_ak,
stride_bn,
stride_bk,
stride_ck,
stride_cm,
stride_cn,
stride_asm,
stride_ask,
stride_bsn,
stride_bsk,
# Meta-parameters
BLOCK_SIZE_M: gl.constexpr,
BLOCK_SIZE_N: gl.constexpr,
BLOCK_SIZE_K: gl.constexpr,
GROUP_SIZE_M: gl.constexpr,
NUM_KSPLIT: gl.constexpr,
SPLITK_BLOCK_SIZE: gl.constexpr,
num_warps: gl.constexpr,
num_stages: gl.constexpr,
waves_per_eu: gl.constexpr,
matrix_instr_nonkdim: gl.constexpr,
cache_modifier: gl.constexpr,
):
"""
Kernel for computing the matmul C = A x B.
A and B inputs are in the microscale fp4 (mxfp4) format.
A_scales and B_scales are in e8m0 format.
A has shape (M, K), B and B_scales are loaded from preshuffled storage,
and C has shape (M, N)
"""
GRID_MN = gl.cdiv(M, BLOCK_SIZE_M) * gl.cdiv(N, BLOCK_SIZE_N)
# -----------------------------------------------------------
# Map program ids `pid` to the block of C it should compute.
# This is done in a grouped ordering to promote L2 data reuse.
pid_unified = gl.program_id(axis=0)
# remap so that XCDs get continous chunks of pids (of CHUNK_SIZE).
pid_unified = remap_xcd(pid_unified, GRID_MN * NUM_KSPLIT, NUM_XCDS=8)
pid_k = pid_unified % NUM_KSPLIT
pid = pid_unified // NUM_KSPLIT
num_pid_m = gl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = gl.cdiv(N, BLOCK_SIZE_N)
if NUM_KSPLIT == 1:
pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
else:
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
# We assume 32 elements along K share the same scale.
SCALE_GROUP_SIZE: gl.constexpr = 32
blocked_mk: gl.constexpr = gl.BlockedLayout(
size_per_thread=[1, 16],
threads_per_warp=[8, 8],
warps_per_cta=[num_warps, 1],
order=[1, 0],
)
blocked_scales: gl.constexpr = gl.BlockedLayout(
size_per_thread=[4, 1],
threads_per_warp=[8, 8],
warps_per_cta=[1, num_warps],
order=[0, 1],
)
blocked_b_preshuffle: gl.constexpr = gl.BlockedLayout(
size_per_thread=[1, 16],
threads_per_warp=[8, 8],
warps_per_cta=[1, num_warps],
order=[1, 0],
)
blocked_shuffle_scales: gl.constexpr = gl.BlockedLayout(
size_per_thread=[1, 4],
threads_per_warp=[8, 8],
warps_per_cta=[1, num_warps],
order=[1, 0],
)
shared_a: gl.constexpr = gl.SwizzledSharedLayout(
vec=16, per_phase=2, max_phase=8, order=[1, 0]
)
shared_b: gl.constexpr = gl.SwizzledSharedLayout(
vec=16, per_phase=2, max_phase=8, order=[0, 1]
)
shared_scales: gl.constexpr = gl.SwizzledSharedLayout(
vec=1, per_phase=1, max_phase=1, order=[0, 1]
)
mfma_layout: gl.constexpr = gl.amd.AMDMFMALayout(
version=4,
instr_shape=[32, 32, 32],
transposed=True,
warps_per_cta=[2, num_warps // 2],
)
dot_a_layout: gl.constexpr = gl.DotOperandLayout(
operand_index=0, parent=mfma_layout, k_width=16
)
dot_b_layout: gl.constexpr = gl.DotOperandLayout(
operand_index=1, parent=mfma_layout, k_width=16
)
scale_a_layout: gl.constexpr = gl.amd.cdna4.get_mfma_scale_layout(
dot_a_layout, [BLOCK_SIZE_M, BLOCK_SIZE_K // SCALE_GROUP_SIZE]
)
scale_b_layout: gl.constexpr = gl.amd.cdna4.get_mfma_scale_layout(
dot_b_layout, [BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE]
)
if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
num_k_iter = gl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
packed_k = K
packed_scale_k = K // (SCALE_GROUP_SIZE // 2)
packed_bn = N // 16
packed_bsn = N // 32
# Create pointers for first block of A and B input matrices
# A stays in packed [M, K/2] form, while B is read back from the
# preshuffled storage used by the Triton submission path.
offs_ak = gl.arange(0, BLOCK_SIZE_K // 2, layout=gl.SliceLayout(0, blocked_mk))
offs_bk = gl.arange(
0,
(BLOCK_SIZE_K // 2) * 16,
layout=gl.SliceLayout(0, blocked_b_preshuffle),
)
offs_ks_shuffle = gl.arange(
0,
(BLOCK_SIZE_K // SCALE_GROUP_SIZE) * 32,
layout=gl.SliceLayout(0, blocked_shuffle_scales),
)
offs_ks = gl.arange(
0,
BLOCK_SIZE_K // SCALE_GROUP_SIZE,
layout=gl.SliceLayout(0, blocked_scales),
)
offs_am = (
pid_m * BLOCK_SIZE_M
+ gl.arange(0, BLOCK_SIZE_M, layout=gl.SliceLayout(1, blocked_mk))
) % M
offs_bn_packed = (
pid_n * (BLOCK_SIZE_N // 16)
+ gl.arange(
0,
BLOCK_SIZE_N // 16,
layout=gl.SliceLayout(1, blocked_b_preshuffle),
)
) % packed_bn
offs_asm = (
pid_m * BLOCK_SIZE_M
+ gl.arange(0, BLOCK_SIZE_M, layout=gl.SliceLayout(1, blocked_scales))
) % M
offs_bsn_packed = (
pid_n * (BLOCK_SIZE_N // 32)
+ gl.arange(
0,
BLOCK_SIZE_N // 32,
layout=gl.SliceLayout(1, blocked_shuffle_scales),
)
) % packed_bsn
# Create shared memories
smem_a = gl.allocate_shared_memory(
a_ptr.type.element_ty, [BLOCK_SIZE_M, BLOCK_SIZE_K // 2], layout=shared_a
)
smem_b = gl.allocate_shared_memory(
b_ptr.type.element_ty, [BLOCK_SIZE_K // 2, BLOCK_SIZE_N], layout=shared_b
)
smem_as = gl.allocate_shared_memory(
a_scales_ptr.type.element_ty,
[BLOCK_SIZE_M, BLOCK_SIZE_K // SCALE_GROUP_SIZE],
layout=shared_scales,
)
smem_bs = gl.allocate_shared_memory(
b_scales_ptr.type.element_ty,
[BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE],
layout=shared_scales,
)
accumulator = gl.zeros(
(BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=gl.float32, layout=mfma_layout
)
# Load first blocks of A and B input matrices
offs_ak_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_ak
offs_a = offs_am[:, None] * stride_am + offs_ak_split[None, :] * stride_ak
a = gl.amd.cdna4.buffer_load(
ptr=a_ptr,
offsets=offs_a,
)
# Create pointers for the first block of A and B scales.
offs_ks_split = pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) + offs_ks
offs_as = offs_asm[:, None] * stride_asm + offs_ks_split[None, :] * stride_ask
a_scales = gl.amd.cdna4.buffer_load(
ptr=a_scales_ptr,
offsets=offs_as,
)
offs_bk_split = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_bk
b_row_idx = offs_bn_packed[:, None] * 16 + (
offs_bk_split[None, :] // packed_k
)
b_col_idx = offs_bk_split[None, :] % packed_k
offs_b = b_row_idx * stride_bn + b_col_idx * stride_bk
b = gl.amd.cdna4.buffer_load(
ptr=b_ptr,
offsets=offs_b,
cache=cache_modifier,
)
# B scales are N x K even though B operand is K x N.
offs_ks_split = (
pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32 + offs_ks_shuffle
)
b_scale_row_idx = offs_bsn_packed[:, None] * 32 + (
offs_ks_split[None, :] // packed_scale_k
)
b_scale_col_idx = offs_ks_split[None, :] % packed_scale_k
offs_bs = b_scale_row_idx * stride_bsn + b_scale_col_idx * stride_bsk
b_scales = (
gl.amd.cdna4.buffer_load(
ptr=b_scales_ptr,
offsets=offs_bs,
cache=cache_modifier,
)
.reshape(
BLOCK_SIZE_N // 32,
BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8,
4,
16,
2,
2,
1,
)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
)
# Reconstruct B from preshuffled storage into the layout consumed by LDS/MFMA.
b = (
b.reshape(
1,
BLOCK_SIZE_N // 16,
BLOCK_SIZE_K // 64,
2,
16,
16,
)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
.trans(1, 0)
)
smem_a.store(a)
smem_as.store(a_scales)
# num_stages:2
for k in range(0, num_k_iter - 1):
# Load next block of A.
offs_ak_split = (
pid_k * (SPLITK_BLOCK_SIZE // 2)
+ (k + 1) * (BLOCK_SIZE_K // 2)
+ offs_ak
)
offs_a = offs_am[:, None] * stride_am + offs_ak_split[None, :] * stride_ak
a = gl.amd.cdna4.buffer_load(
ptr=a_ptr,
offsets=offs_a,
)
# LDS write current blocks of B and B scales.
smem_b.store(b)
smem_bs.store(b_scales)
curr_a = smem_a.load(layout=dot_a_layout)
curr_a_scales = smem_as.load(layout=scale_a_layout)
# Load next block of A scales.
offs_ks_split = (
pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE)
+ (k + 1) * (BLOCK_SIZE_K // SCALE_GROUP_SIZE)
+ offs_ks
)
offs_as = (
offs_asm[:, None] * stride_asm + offs_ks_split[None, :] * stride_ask
)
a_scales = gl.amd.cdna4.buffer_load(
ptr=a_scales_ptr,
offsets=offs_as,
)
curr_b_scales = smem_bs.load(layout=scale_b_layout)
# Load next block of B from preshuffled storage.
offs_bk_split = (
pid_k * (SPLITK_BLOCK_SIZE // 2) * 16
+ (k + 1) * (BLOCK_SIZE_K // 2) * 16
+ offs_bk
)
b_row_idx = offs_bn_packed[:, None] * 16 + (
offs_bk_split[None, :] // packed_k
)
b_col_idx = offs_bk_split[None, :] % packed_k
offs_b = b_row_idx * stride_bn + b_col_idx * stride_bk
b = gl.amd.cdna4.buffer_load(
ptr=b_ptr,
offsets=offs_b,
cache=cache_modifier,
)
# Load next block of B scales from preshuffled storage.
offs_ks_split = (
pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32
+ (k + 1) * BLOCK_SIZE_K
+ offs_ks_shuffle
)
b_scale_row_idx = offs_bsn_packed[:, None] * 32 + (
offs_ks_split[None, :] // packed_scale_k
)
b_scale_col_idx = offs_ks_split[None, :] % packed_scale_k
offs_bs = b_scale_row_idx * stride_bsn + b_scale_col_idx * stride_bsk
b_scales = (
gl.amd.cdna4.buffer_load(
ptr=b_scales_ptr,
offsets=offs_bs,
cache=cache_modifier,
)
.reshape(
BLOCK_SIZE_N // 32,
BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8,
4,
16,
2,
2,
1,
)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
)
# Read current block of B from LDS.
curr_b = smem_b.load(layout=dot_b_layout)
accumulator = gl.amd.cdna4.mfma_scaled(
a=curr_a,
a_scale=curr_a_scales,
a_format="e2m1",
b=curr_b,
b_scale=curr_b_scales,
b_format="e2m1",
acc=accumulator,
)
# Reconstruct next block of B into the layout consumed by LDS/MFMA.
b = (
b.reshape(
1,
BLOCK_SIZE_N // 16,
BLOCK_SIZE_K // 64,
2,
16,
16,
)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
.trans(1, 0)
)
# LDS write next block of A and A scales.
smem_a.store(a)
smem_as.store(a_scales)
# ======= Epilogue ========
smem_b.store(b)
smem_bs.store(b_scales)
curr_a = smem_a.load(layout=dot_a_layout)
curr_b = smem_b.load(layout=dot_b_layout)
curr_a_scales = smem_as.load(layout=scale_a_layout)
curr_b_scales = smem_bs.load(layout=scale_b_layout)
accumulator = gl.amd.cdna4.mfma_scaled(
a=curr_a,
a_scale=curr_a_scales,
a_format="e2m1",
b=curr_b,
b_scale=curr_b_scales,
b_format="e2m1",
acc=accumulator,
)
c = accumulator.to(c_ptr.type.element_ty)
offs_cm = pid_m * BLOCK_SIZE_M + gl.arange(
0, BLOCK_SIZE_M, layout=gl.SliceLayout(1, mfma_layout)
)
offs_cn = pid_n * BLOCK_SIZE_N + gl.arange(
0, BLOCK_SIZE_N, layout=gl.SliceLayout(0, mfma_layout)
)
offs_c = (
stride_cm * offs_cm[:, None]
+ stride_cn * offs_cn[None, :]
+ pid_k * stride_ck
)
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
gl.amd.cdna4.buffer_store(c, c_ptr, offs_c, c_mask)
@gluon.jit
def _gemm_afp4wfp4_reduce_kernel(
c_in_ptr,
c_out_ptr,
M,
N,
stride_c_in_k,
stride_c_in_m,
stride_c_in_n,
stride_c_out_m,
stride_c_out_n,
BLOCK_SIZE_M: gl.constexpr,
BLOCK_SIZE_N: gl.constexpr,
ACTUAL_KSPLIT: gl.constexpr,
MAX_KSPLIT: gl.constexpr,
):
pid_m = gl.program_id(axis=0)
pid_n = gl.program_id(axis=1)
blocked_kmn: gl.constexpr = gl.BlockedLayout(
size_per_thread=[1, 1, 4],
threads_per_warp=[2, 2, 16],
warps_per_cta=[1, 4, 1],
order=[2, 0, 1],
)
blocked_mn: gl.constexpr = gl.BlockedLayout(
size_per_thread=[1, 4],
threads_per_warp=[4, 16],
warps_per_cta=[4, 1],
order=[1, 0],
)
offs_m = (
pid_m * BLOCK_SIZE_M
+ gl.arange(
0, BLOCK_SIZE_M, layout=gl.SliceLayout(0, gl.SliceLayout(2, blocked_kmn))
)
) % M
offs_n = (
pid_n * BLOCK_SIZE_N
+ gl.arange(
0, BLOCK_SIZE_N, layout=gl.SliceLayout(0, gl.SliceLayout(1, blocked_kmn))
)
) % N
offs_k = gl.arange(
0, MAX_KSPLIT, layout=gl.SliceLayout(1, gl.SliceLayout(2, blocked_kmn))
)
c_in_ptrs = (
c_in_ptr
+ (offs_k[:, None, None] * stride_c_in_k)
+ (offs_m[None, :, None] * stride_c_in_m)
+ (offs_n[None, None, :] * stride_c_in_n)
)
if ACTUAL_KSPLIT == MAX_KSPLIT:
c = gl.load(c_in_ptrs)
else:
c = gl.load(c_in_ptrs, mask=offs_k[:, None, None] < ACTUAL_KSPLIT)
c = gl.sum(c, axis=0)
c = c.to(c_out_ptr.type.element_ty)
offs_m = (
pid_m * BLOCK_SIZE_M
+ gl.arange(0, BLOCK_SIZE_M, layout=gl.SliceLayout(1, blocked_mn))
) % M
offs_n = (
pid_n * BLOCK_SIZE_N
+ gl.arange(0, BLOCK_SIZE_N, layout=gl.SliceLayout(0, blocked_mn))
) % N
c_out_ptrs = (
c_out_ptr
+ (offs_m[:, None] * stride_c_out_m)
+ (offs_n[None, :] * stride_c_out_n)
)
c = gl.convert_layout(c, layout=blocked_mn, assert_trivial=False)
gl.store(c_out_ptrs, c)
@triton.jit
def _gemm_afp4wfp4_reduce_triton_kernel(
c_in_ptr,
c_out_ptr,
M,
N,
stride_c_in_k,
stride_c_in_m,
stride_c_in_n,
stride_c_out_m,
stride_c_out_n,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
ACTUAL_KSPLIT: tl.constexpr,
MAX_KSPLIT: tl.constexpr,
):
pid_m = tl.program_id(axis=0)
pid_n = tl.program_id(axis=1)
offs_m = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
offs_n = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
offs_k = tl.arange(0, MAX_KSPLIT)
c_in_ptrs = (
c_in_ptr
+ (offs_k[:, None, None] * stride_c_in_k)
+ (offs_m[None, :, None] * stride_c_in_m)
+ (offs_n[None, None, :] * stride_c_in_n)
)
if ACTUAL_KSPLIT == MAX_KSPLIT:
c = tl.load(c_in_ptrs)
else:
c = tl.load(c_in_ptrs, mask=offs_k[:, None, None] < ACTUAL_KSPLIT)
c = tl.sum(c, axis=0)
c = c.to(c_out_ptr.type.element_ty)
c_out_ptrs = (
c_out_ptr
+ (offs_m[:, None] * stride_c_out_m)
+ (offs_n[None, :] * stride_c_out_n)
)
tl.store(c_out_ptrs, c)
@functools.lru_cache(maxsize=1024)
def _get_config(
M: int,
N: int,
K: int,
):
# K in this API is packed K/2, benchmark table uses full K.
full_k = 2 * K
key = (M, N, full_k)
if key in _BENCHMARK_GEMM_CONFIGS:
return dict(_BENCHMARK_GEMM_CONFIGS[key])
return _select_default_config_by_m(M)
def gemm_afp4wfp4_preshuffle(
x: torch.Tensor,
w: torch.Tensor,
x_scales: torch.Tensor,
w_scales: torch.Tensor,
dtype: Optional[torch.dtype] = torch.bfloat16,
y: Optional[torch.Tensor] = None,
config: Optional[dict] = None,
skip_reduce: Optional[bool] = False,
a_shuffled: Optional[bool] = False,
b_shuffled: Optional[bool] = True,
) -> torch.Tensor:
"""
Computes matrix multiplication Y = X @ W with FP4 activations and FP4 weights.
Args:
x (torch.Tensor): FP4 E2M1 input matrix with shape (M, K//2).
w (torch.Tensor): FP4 E2M1 preshuffled weight storage.
x_scales (torch.Tensor): E8M0 per-group scale for x with shape (M, K//32).
One scale per 32 elements in K dimension.
w_scales (torch.Tensor): E8M0 per-group scale for w with shape (N, K//32).
One scale per 32 elements in K dimension.
dtype (Optional[torch.dtype]): Output datatype (BF16 or FP16).
y (Optional[torch.Tensor]): Pre-allocated output tensor with shape (M, N).
config (Optional[dict]): Kernel tuning parameters (BLOCK_SIZE_M, BLOCK_SIZE_N,
BLOCK_SIZE_K, GROUP_SIZE_M, NUM_KSPLIT, SPLITK_BLOCK_SIZE).
skip_reduce (Optional[bool]): skip reduction, y becomes (SPK, M, N) where SPK is determined by config
Returns:
y (torch.Tensor): Output with shape (M, N) or (SPK, M, N).
"""
atomic_add = False
M, K = x.shape
N, K = w.shape
if config is None:
config = _get_config(M, N, K)
if config["BLOCK_SIZE_K"] >= K * 2:
config["NUM_KSPLIT"] = 1
if config["NUM_KSPLIT"] > 1:
SPLITK_BLOCK_SIZE = (
triton.cdiv(
(2 * triton.cdiv(K, config["NUM_KSPLIT"])), config["BLOCK_SIZE_K"]
)
* config["BLOCK_SIZE_K"]
)
else:
SPLITK_BLOCK_SIZE = 2 * K
config["SPLITK_BLOCK_SIZE"] = SPLITK_BLOCK_SIZE
if config["NUM_KSPLIT"] > 1 and not atomic_add:
y_pp = torch.empty(
(config["NUM_KSPLIT"], M, N), dtype=torch.float32, device=x.device
)
else:
y_pp = None
y_atomic = None
if atomic_add:
y_atomic = torch.zeros((M, N), dtype=torch.float32, device=x.device)
if y is None and (config["NUM_KSPLIT"] == 1 or not skip_reduce):
y = torch.empty((M, N), dtype=dtype, device=x.device)
grid = lambda META: ( # noqa: E731
(
META["NUM_KSPLIT"]
* triton.cdiv(M, META["BLOCK_SIZE_M"])
* triton.cdiv(N, META["BLOCK_SIZE_N"])
),
)
def alloc_fn(size: int, align: int, sm: Optional[int]):
return torch.empty(size, device=x.device, dtype=torch.int8)
triton.set_allocator(alloc_fn)
# _gemm_afp4wfp4_preshuffle_kernel[grid](
_gemm_afp4wfp4_preshuffle_triton_kernel[grid](
x,
w,
y_atomic if atomic_add else (y if y_pp is None else y_pp),
x_scales,
w_scales,
M,
N,
K,
x.stride(0),
x.stride(1),
w.stride(0),
w.stride(1),
0 if (atomic_add or y_pp is None) else y_pp.stride(0),
(y_atomic.stride(0) if atomic_add else (y.stride(0) if y_pp is None else y_pp.stride(1))),
(y_atomic.stride(1) if atomic_add else (y.stride(1) if y_pp is None else y_pp.stride(2))),
x_scales.stride(0),
x_scales.stride(1),
w_scales.stride(0),
w_scales.stride(1),
A_SHUFFLED=a_shuffled,
B_SHUFFLED=b_shuffled,
ATOMIC_ADD=atomic_add,
**config,
)
if config["NUM_KSPLIT"] > 1 and not atomic_add:
if skip_reduce:
return y_pp
REDUCE_BLOCK_SIZE_M = 16
REDUCE_BLOCK_SIZE_N = 64
ACTUAL_KSPLIT = triton.cdiv(K, (config["SPLITK_BLOCK_SIZE"] // 2))
grid_reduce = (
triton.cdiv(M, REDUCE_BLOCK_SIZE_M),
triton.cdiv(N, REDUCE_BLOCK_SIZE_N),
)
# _gemm_afp4wfp4_reduce_kernel[grid_reduce](
_gemm_afp4wfp4_reduce_triton_kernel[grid_reduce](
y_pp,
y,
M,
N,
y_pp.stride(0),
y_pp.stride(1),
y_pp.stride(2),
y.stride(0),
y.stride(1),
REDUCE_BLOCK_SIZE_M,
REDUCE_BLOCK_SIZE_N,
ACTUAL_KSPLIT,
triton.next_power_of_2(config["NUM_KSPLIT"]),
)
if atomic_add:
y.copy_(y_atomic.to(y.dtype))
return y
def gemm_afp4wfp4(
x: torch.Tensor,
w: torch.Tensor,
x_scales: torch.Tensor,
w_scales: torch.Tensor,
dtype: Optional[torch.dtype] = torch.bfloat16,
y: Optional[torch.Tensor] = None,
config: Optional[dict] = None,
skip_reduce: Optional[bool] = False,
) -> torch.Tensor:
"""
Computes matrix multiplication Y = X @ W with FP4 activations and FP4 weights.
This entrypoint expects the original non-preshuffled weight layout.
"""
M, K = x.shape
N, _ = w.shape
if config is None:
config = _get_config(M, N, K)
if config["BLOCK_SIZE_K"] >= K * 2:
config["NUM_KSPLIT"] = 1
if config["NUM_KSPLIT"] > 1:
SPLITK_BLOCK_SIZE = (
triton.cdiv(
(2 * triton.cdiv(K, config["NUM_KSPLIT"])), config["BLOCK_SIZE_K"]
)
* config["BLOCK_SIZE_K"]
)
else:
SPLITK_BLOCK_SIZE = 2 * K
config["SPLITK_BLOCK_SIZE"] = SPLITK_BLOCK_SIZE
if config["NUM_KSPLIT"] > 1:
y_pp = torch.empty(
(config["NUM_KSPLIT"], M, N), dtype=torch.float32, device=x.device
)
else:
y_pp = None
if y is None and (config["NUM_KSPLIT"] == 1 or not skip_reduce):
y = torch.empty((M, N), dtype=dtype, device=x.device)
grid = lambda META: ( # noqa: E731
(
META["NUM_KSPLIT"]
* triton.cdiv(M, META["BLOCK_SIZE_M"])
* triton.cdiv(N, META["BLOCK_SIZE_N"])
),
)
def alloc_fn(size: int, align: int, sm: Optional[int]):
return torch.empty(size, device=x.device, dtype=torch.int8)
triton.set_allocator(alloc_fn)
# _gemm_afp4wfp4_kernel[grid](
_gemm_afp4wfp4_triton_kernel[grid](
x,
w,
y if config["NUM_KSPLIT"] == 1 else y_pp,
x_scales,
w_scales,
M,
N,
K,
x.stride(0),
x.stride(1),
w.stride(1),
w.stride(0),
0 if config["NUM_KSPLIT"] == 1 else y_pp.stride(0),
y.stride(0) if config["NUM_KSPLIT"] == 1 else y_pp.stride(1),
y.stride(1) if config["NUM_KSPLIT"] == 1 else y_pp.stride(2),
x_scales.stride(0),
x_scales.stride(1),
w_scales.stride(0),
w_scales.stride(1),
**config,
)
if config["NUM_KSPLIT"] > 1:
if skip_reduce:
return y_pp
REDUCE_BLOCK_SIZE_M = 16
REDUCE_BLOCK_SIZE_N = 64
ACTUAL_KSPLIT = triton.cdiv(K, (config["SPLITK_BLOCK_SIZE"] // 2))
grid_reduce = (
triton.cdiv(M, REDUCE_BLOCK_SIZE_M),
triton.cdiv(N, REDUCE_BLOCK_SIZE_N),
)
# _gemm_afp4wfp4_reduce_kernel[grid_reduce](
_gemm_afp4wfp4_reduce_triton_kernel[grid_reduce](
y_pp,
y,
M,
N,
y_pp.stride(0),
y_pp.stride(1),
y_pp.stride(2),
y.stride(0),
y.stride(1),
REDUCE_BLOCK_SIZE_M,
REDUCE_BLOCK_SIZE_N,
ACTUAL_KSPLIT,
triton.next_power_of_2(config["NUM_KSPLIT"]),
)
return y
def e8m0_shuffle(scale):
if scale is None:
return scale
if scale.dtype == torch.float32:
return scale
assert scale.ndim == 2, "scale must be a 2D tensor"
m, n = scale.shape
scale_padded = torch.empty(
(m + 255) // 256 * 256,
(n + 7) // 8 * 8,
dtype=scale.dtype,
device=scale.device,
)
scale_padded[:m, :n] = scale
scale = scale_padded
sm, sn = scale.shape
scale = scale.view(sm // 32, 2, 16, sn // 8, 2, 4)
scale = scale.permute(0, 3, 5, 2, 4, 1).contiguous()
scale = scale.view(sm, sn)
return scale
def _quant_mxfp4(x, shuffle=True):
# x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
x_fp4, bs_e8m0 = dynamic_mxfp4_quant_inline(x)
if shuffle:
bs_e8m0 = e8m0_shuffle(bs_e8m0)
return x_fp4, bs_e8m0
def _as_uint8_storage(x: torch.Tensor) -> torch.Tensor:
if x.dtype == torch.uint8:
return x
return x.view(torch.uint8)
# import time
def custom_kernel(data: Any) -> Any:
A, B, B_q, B_shuffle, B_scale_sh = data
B_shuffle = _as_uint8_storage(B_shuffle)
B_scale_sh = _as_uint8_storage(B_scale_sh)
start = time.time()
A_q, A_scale_sh = _quant_mxfp4(A, shuffle=False)
end = time.time()
print(f"Quantization time: {(end - start) * 1e6:.2f} us")
start = time.time()
out_gemm = gemm_afp4wfp4_preshuffle(
A_q,
B_shuffle,
A_scale_sh,
B_scale_sh,
dtype=torch.bfloat16,
a_shuffled=False,
b_shuffled=True,
)
end = time.time()
print(f"GEMM time: {(end - start) * 1e6:.2f} us")
# A_q, A_scale_sh = _quant_mxfp4(A, shuffle=True)
# B_q, B_scale_sh = _quant_mxfp4(B, shuffle=True)
# out_gemm = gemm_afp4wfp4(
# A_q,
# B_q,
# A_scale_sh,
# B_scale_sh,
# dtype=torch.bfloat16,
# )
return out_gemm
# def shuffle_weight(x: torch.Tensor, layout=(16, 16), use_int4=False) -> torch.Tensor:
# # Hardcode BLOCK_K and BLOCK_N
# x_type = x.dtype
# if hasattr(torch, "float4_e2m1fn_x2") and x_type == torch.float4_e2m1fn_x2:
# x = x.view(torch.uint8)
# IN, IK = layout
# BK = IK * 2
# K = 16 // x.element_size() if not use_int4 else 32
# BN = IN
# assert x.shape[-2] % BN == 0, f"{x.shape[-2]} % {BN} == {x.shape[-2] % BN }"
# assert x.shape[-1] % BK == 0, f"{x.shape[-1]} % {BK} == {x.shape[-1] % BK }"
# x_ = x
# x_ = x_.view(-1, x.shape[-2] // BN, BN, x.shape[-1] // BK, BK // K, K)
# x_ = x_.permute(0, 1, 3, 4, 2, 5)
# x_ = x_.contiguous()
# x_ = x_.view(*x.shape)
# x_ = x_.view(x_type)
# x_.is_shuffled = True
# return x_
# def generate_input(m: int, n: int, k: int, seed: int) -> tuple[torch.Tensor, ...]:
# assert k % 64 == 0, "k must be divisible by 64"
# gen = torch.Generator(device="cuda")
# gen.manual_seed(seed)
# a = torch.randn((m, k), dtype=torch.bfloat16, device="cuda", generator=gen)
# b = torch.randn((n, k), dtype=torch.bfloat16, device="cuda", generator=gen)
# b_q, b_scale_sh = _quant_mxfp4(b, shuffle=True)
# b_shuffle = _as_uint8_storage(shuffle_weight(b_q, layout=(16, 16)))
# return a, b, b_q, b_shuffle, _as_uint8_storage(b_scale_sh)
# def mxfp4_to_f32(x):
# if x.dtype == torch.float4_e2m1fn_x2:
# x = x.view(torch.uint8)
# # 2 because we pack fp4 in uint8.
# x = x.repeat_interleave(2, dim=-1)
# x[..., ::2] = x[..., ::2] & 0xF
# x[..., 1::2] = x[..., 1::2] >> 4
# mxfp4_list = [
# 0.0,
# 0.5,
# 1.0,
# 1.5,
# 2.0,
# 3.0,
# 4.0,
# 6.0,
# -0.0,
# -0.5,
# -1.0,
# -1.5,
# -2.0,
# -3.0,
# -4.0,
# -6.0,
# ]
# mxfp4_in_f32 = torch.tensor(mxfp4_list, dtype=torch.float32, device=x.device)
# return mxfp4_in_f32[x.long()]
# def e8m0_to_f32(scale_e8m0_biased):
# scale_e8m0_biased = scale_e8m0_biased.view(torch.uint8)
# zero_case = scale_e8m0_biased == 0
# nan_case = scale_e8m0_biased == 0xFF
# scale_f32 = scale_e8m0_biased.to(torch.int32) << 23
# scale_f32[zero_case] = 0x00400000
# scale_f32[nan_case] = 0x7F800001
# scale_f32 = scale_f32.view(torch.float32)
# return scale_f32
# def run_torch_fp4_mm(
# x: torch.Tensor,
# w: torch.Tensor,
# x_scales: torch.Tensor,
# w_scales: torch.Tensor,
# dtype: torch.dtype = torch.bfloat16,
# ) -> torch.Tensor:
# """
# PyTorch reference: dequant MXFP4 + E8M0 scale -> f32 -> mm -> dtype.
# Same logic as aiter op_tests/test_gemm_a4w4.run_torch.
# x: [m, k//2] fp4 packed, w: [n, k//2] fp4 packed
# x_scales: [m, k//32] E8M0, w_scales: [n, k//32] E8M0
# Returns: [m, n] in dtype
# """
# m, _ = x.shape
# n, _ = w.shape
# # fp4 packed -> f32
# x_f32 = mxfp4_to_f32(x)
# w_f32 = mxfp4_to_f32(w)
# # E8M0 scale: [*, k//32] -> repeat 32 along k -> f32
# x_scales = x_scales[:m].repeat_interleave(MXFP4_GROUP_SIZE, dim=1)
# x_scales_f32 = e8m0_to_f32(x_scales)
# x_f32 = x_f32 * x_scales_f32
# w_scales = w_scales[:n].repeat_interleave(MXFP4_GROUP_SIZE, dim=1)
# w_scales_f32 = e8m0_to_f32(w_scales)
# w_f32 = w_f32 * w_scales_f32
# return torch.mm(x_f32, w_f32.T).to(dtype)[:m, :n]
# def _calculate_stats(durations_ns: list[float]) -> dict[str, float]:
# runs = len(durations_ns)
# mean = sum(durations_ns) / runs
# best = min(durations_ns)
# worst = max(durations_ns)
# if runs > 1:
# variance = sum((x - mean) ** 2 for x in durations_ns) / (runs - 1)
# std = math.sqrt(variance)
# err = std / math.sqrt(runs)
# else:
# std = 0.0
# err = 0.0
# return {
# "runs": float(runs),
# "mean_ns": mean,
# "std_ns": std,
# "err_ns": err,
# "best_ns": best,
# "worst_ns": worst,
# }
# def benchmark_custom_kernel(
# benchmarks: list[dict[str, int]],
# warmup: int = 0,
# max_repeats: int = 1,
# max_time_ns: float = 30e9,
# ) -> None:
# if not torch.cuda.is_available():
# print("CUDA/ROCm device not available, skip benchmark")
# return
# def _measure_latency_ns(fn, warmup_count: int, max_repeat_count: int) -> dict[str, float]:
# for _ in range(warmup_count):
# _ = fn()
# torch.cuda.synchronize()
# durations_ns: list[float] = []
# bm_start = time.perf_counter_ns()
# for i in range(max_repeat_count):
# start_event = torch.cuda.Event(enable_timing=True)
# end_event = torch.cuda.Event(enable_timing=True)
# start_event.record()
# _ = fn()
# end_event.record()
# torch.cuda.synchronize()
# durations_ns.append(start_event.elapsed_time(end_event) * 1e6)
# if i > 1:
# stats = _calculate_stats(durations_ns)
# total_bm_duration = time.perf_counter_ns() - bm_start
# if (
# stats["err_ns"] / max(stats["mean_ns"], 1.0) < 0.001
# or stats["mean_ns"] * len(durations_ns) > max_time_ns
# or total_bm_duration > 120e9
# ):
# break
# return _calculate_stats(durations_ns)
# print(f"benchmark-count: {len(benchmarks)}")
# for idx, case in enumerate(benchmarks):
# m, n, k, seed = case["m"], case["n"], case["k"], case["seed"]
# spec = f"m:{m};n:{n};k:{k};seed:{seed}"
# print(f"benchmark.{idx}.spec: {spec}")
# data = generate_input(m, n, k, seed)
# a, b, _b_q_shuffled_scale, _b_shuffle, _b_scale_sh = data
# out = custom_kernel(data)
# torch.cuda.synchronize()
# if out.shape != (m, n):
# print(f"benchmark.{idx}.status: fail")
# print(f"benchmark.{idx}.error: shape mismatch got={tuple(out.shape)} expected={(m, n)}")
# continue
# a_q_ref, a_scale_ref = _quant_mxfp4(a, shuffle=False)
# b_q_ref, b_scale_ref = _quant_mxfp4(b, shuffle=False)
# out_ref = run_torch_fp4_mm(a_q_ref, b_q_ref, a_scale_ref, b_scale_ref, dtype=out.dtype)
# out_f32 = out.float()
# out_ref_f32 = out_ref.float()
# abs_diff = (out_f32 - out_ref_f32).abs()
# max_abs_diff = abs_diff.max().item()
# mean_abs_diff = abs_diff.mean().item()
# ref_abs_max = out_ref_f32.abs().max().item()
# rel_max_diff = max_abs_diff / max(ref_abs_max, 1e-6)
# custom_stats = _measure_latency_ns(lambda: custom_kernel(data), warmup, max_repeats)
# # Torch reference is much slower, use fewer repeats to keep benchmark time practical.
# ref_repeats = min(max_repeats, 20)
# ref_stats = _measure_latency_ns(
# lambda: run_torch_fp4_mm(a_q_ref, b_q_ref, a_scale_ref, b_scale_ref, dtype=out.dtype),
# min(warmup, 5),
# ref_repeats,
# )
# speedup = ref_stats["mean_ns"] / max(custom_stats["mean_ns"], 1.0)
# print(f"benchmark.{idx}.correctness.max_abs_diff: {max_abs_diff:.6f}")
# print(f"benchmark.{idx}.correctness.mean_abs_diff: {mean_abs_diff:.6f}")
# print(f"benchmark.{idx}.correctness.rel_max_diff: {rel_max_diff:.6f}")
# print(f"benchmark.{idx}.perf.custom_mean_us: {custom_stats['mean_ns'] / 1e3:.3f}")
# print(f"benchmark.{idx}.perf.reference_mean_us: {ref_stats['mean_ns'] / 1e3:.3f}")
# print(f"benchmark.{idx}.perf.speedup_vs_reference: {speedup:.2f}x")
# if __name__ == "__main__":
# benchmarks = [
# {"m": 4, "n": 2880, "k": 512, "seed": 4565},
# {"m": 16, "n": 2112, "k": 7168, "seed": 15},
# {"m": 32, "n": 4096, "k": 512, "seed": 457},
# {"m": 32, "n": 2880, "k": 512, "seed": 54},
# {"m": 64, "n": 7168, "k": 2048, "seed": 687},
# {"m": 256, "n": 3072, "k": 1536, "seed": 7856},
# ]
# benchmark_custom_kernel(benchmarks)scrolls · 2559 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON