submission 503034
J · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 584 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-503034?include=source"interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp8_e4m3, nvfp4
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:585444581b32f6a01da000d86c396a82ce284a5a1ea39ecc248630f9819c8e78
license declaredunknown
license concludedunknown
authorsJ
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
cluster
using ClusterShape = Shape<int32_t,int32_t,_1>;fp4
using ElementInput = cutlass::float_e2m1_t;fused-epilogue
using ES1 = cutlass::epilogue::PtrArrayTmaWarpSpecialized1Sm;warp-specialization
using ES1 = cutlass::epilogue::PtrArrayTmaWarpSpecialized1Sm;Kernel source
submission.py584 lines
#!POPCORN leaderboard nvfp4_group_gemm
#!POPCORN gpu NVIDIA
from typing import Dict, List, Tuple
import weakref
import sys
import time
import torch
import torch.utils.cpp_extension
from task import input_t, output_t
from utils import make_match_reference
sf_vec_size = 16
def ceil_div(a, b):
return (a + b - 1) // b
def to_blocked(input_matrix):
rows, cols = input_matrix.shape
n_row_blocks = ceil_div(rows, 128)
n_col_blocks = ceil_div(cols, 4)
padded_rows = n_row_blocks * 128
padded_cols = n_col_blocks * 4
if padded_rows != rows or padded_cols != cols:
padded = torch.nn.functional.pad(input_matrix, (0, padded_cols - cols, 0, padded_rows - rows))
else:
padded = input_matrix
blocks = padded.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
rearranged = blocks.reshape(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)
return rearranged.flatten()
def _unpack_input(data):
if len(data) == 3:
return data[0], data[1], None, data[2]
if len(data) == 4:
return data[0], data[1], data[2], data[3]
raise ValueError(f"Unexpected input format with {len(data)} elements")
CUDA_SOURCE = r"""
#include "cutlass/cutlass.h"
#include "cute/tensor.hpp"
#include "cutlass/tensor_ref.h"
#include "cutlass/epilogue/collective/default_epilogue.hpp"
#include "cutlass/epilogue/thread/linear_combination.h"
#include "cutlass/gemm/dispatch_policy.hpp"
#include "cutlass/gemm/group_array_problem_shape.hpp"
#include "cutlass/gemm/collective/collective_builder.hpp"
#include "cutlass/epilogue/collective/collective_builder.hpp"
#include "cutlass/gemm/device/gemm_universal_adapter.h"
#include "cutlass/gemm/kernel/gemm_universal.hpp"
#include "cutlass/util/packed_stride.hpp"
using namespace cute;
using ProblemShape = cutlass::gemm::GroupProblemShape<Shape<int,int,int>>;
using ElementInput = cutlass::float_e2m1_t;
using ElementSF = cutlass::float_ue4m3_t;
using ElementC = cutlass::half_t;
using ElementA = cutlass::nv_float4_t<ElementInput>;
using LayoutA = cutlass::layout::RowMajor;
constexpr int AlignmentA = 32;
using ElementB = cutlass::nv_float4_t<ElementInput>;
using LayoutB = cutlass::layout::ColumnMajor;
constexpr int AlignmentB = 32;
using ElementD = ElementC;
using LayoutC = cutlass::layout::RowMajor;
constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
using ElementAccumulator = float;
using ArchTag = cutlass::arch::Sm100;
using OperatorClass = cutlass::arch::OpClassBlockScaledTensorOp;
using ClusterShape = Shape<int32_t,int32_t,_1>;
// Mode 1: 1SM
using KS1 = cutlass::gemm::KernelPtrArrayTmaWarpSpecialized1SmNvf4Sm100;
using ES1 = cutlass::epilogue::PtrArrayTmaWarpSpecialized1Sm;
using Tile1 = Shape<_128,_256,_256>;
// Mode 2: 2SM N256
using KS2 = cutlass::gemm::KernelPtrArrayTmaWarpSpecialized2SmNvf4Sm100;
using ES2 = cutlass::epilogue::PtrArrayTmaWarpSpecialized2Sm;
using Tile2 = Shape<_256,_256,_256>;
// Mode 3: 2SM N128
using Tile3 = Shape<_256,_128,_256>;
#define BUILD_GEMM(Name, TileShape, KSched, ESched) \
using Epi_##Name = typename cutlass::epilogue::collective::CollectiveBuilder< \
ArchTag, OperatorClass, TileShape, ClusterShape, Shape<_128,_64>, \
ElementAccumulator, ElementAccumulator, \
ElementC, LayoutC *, AlignmentC, ElementD, LayoutC *, AlignmentD, ESched \
>::CollectiveOp; \
using Main_##Name = typename cutlass::gemm::collective::CollectiveBuilder< \
ArchTag, OperatorClass, ElementA, LayoutA *, AlignmentA, \
ElementB, LayoutB *, AlignmentB, ElementAccumulator, TileShape, ClusterShape, \
cutlass::gemm::collective::StageCountAutoCarveout< \
static_cast<int>(sizeof(typename Epi_##Name::SharedStorage))>, KSched \
>::CollectiveOp; \
using Kernel_##Name = cutlass::gemm::kernel::GemmUniversal<ProblemShape, Main_##Name, Epi_##Name>; \
using Name = cutlass::gemm::device::GemmUniversalAdapter<Kernel_##Name>;
BUILD_GEMM(Gemm1SM, Tile1, KS1, ES1)
BUILD_GEMM(Gemm2SM, Tile2, KS2, ES2)
BUILD_GEMM(Gemm2SM_N128, Tile3, KS2, ES2)
using Gemm = Gemm1SM;
using StrideA = typename Gemm::GemmKernel::InternalStrideA;
using StrideB = typename Gemm::GemmKernel::InternalStrideB;
using StrideC = typename Gemm::GemmKernel::InternalStrideC;
using StrideD = typename Gemm::GemmKernel::InternalStrideD;
using InternalLayoutSFA = typename Gemm::GemmKernel::CollectiveMainloop::InternalLayoutSFA;
using InternalLayoutSFB = typename Gemm::GemmKernel::CollectiveMainloop::InternalLayoutSFB;
using Sm1xxBlkScaledConfig = typename Gemm::GemmKernel::CollectiveMainloop::Sm1xxBlkScaledConfig;
template <typename GemmT>
int run_impl(int ng, void* dps,
void* pa, void* sa, void* pb, void* sb,
void* psfa, void* lsfa, void* psfb, void* lsfb,
void* pc, void* sc, void* pd, void* sd,
void* ws, size_t ws_sz, void* hps,
int cx, int cy, int fcx, int fcy,
bool run_kernel, bool use_pdl
) {
using SA = typename GemmT::GemmKernel::InternalStrideA;
using SB = typename GemmT::GemmKernel::InternalStrideB;
using SC = typename GemmT::GemmKernel::InternalStrideC;
using SD = typename GemmT::GemmKernel::InternalStrideD;
using AA = typename GemmT::GemmKernel::CollectiveMainloop::ArrayElementA;
using AB = typename GemmT::GemmKernel::CollectiveMainloop::ArrayElementB;
using LSFA = typename GemmT::GemmKernel::CollectiveMainloop::InternalLayoutSFA;
using LSFB = typename GemmT::GemmKernel::CollectiveMainloop::InternalLayoutSFB;
cutlass::KernelHardwareInfo hw;
hw.device_id = 0;
hw.sm_count = cutlass::KernelHardwareInfo::query_device_multiprocessor_count(hw.device_id);
hw.cluster_shape = dim3(cx, cy, 1);
hw.cluster_shape_fallback = dim3(fcx, fcy, 1);
typename GemmT::GemmKernel::TileSchedulerArguments sched;
typename GemmT::Arguments args;
decltype(args.epilogue.thread) fargs;
fargs.alpha = 1.0f; fargs.beta = 0.0f;
fargs.alpha_ptr = nullptr; fargs.beta_ptr = nullptr;
fargs.alpha_ptr_array = nullptr; fargs.beta_ptr_array = nullptr;
fargs.dAlpha = {_0{}, _0{}, 0}; fargs.dBeta = {_0{}, _0{}, 0};
args = typename GemmT::Arguments {
cutlass::gemm::GemmUniversalMode::kGrouped,
{ng, reinterpret_cast<typename ProblemShape::UnderlyingProblemShape*>(dps),
reinterpret_cast<typename ProblemShape::UnderlyingProblemShape*>(hps)},
{reinterpret_cast<const AA**>(pa), reinterpret_cast<SA*>(sa),
reinterpret_cast<const AB**>(pb), reinterpret_cast<SB*>(sb),
reinterpret_cast<const ElementSF**>(psfa), reinterpret_cast<LSFA*>(lsfa),
reinterpret_cast<const ElementSF**>(psfb), reinterpret_cast<LSFB*>(lsfb)},
{fargs, reinterpret_cast<const ElementC**>(pc), reinterpret_cast<SC*>(sc),
reinterpret_cast<ElementD**>(pd), reinterpret_cast<SD*>(sd)},
hw, sched
};
size_t needed = GemmT::get_workspace_size(args);
if (needed > ws_sz) return -1;
static GemmT op;
auto st = op.can_implement(args);
if (st != cutlass::Status::kSuccess) return -2;
if (!run_kernel) return 0;
st = op.initialize(args, ws, nullptr);
if (st != cutlass::Status::kSuccess) return -3;
st = op.run(nullptr, nullptr, use_pdl);
if (st != cutlass::Status::kSuccess) return -4;
return 0;
}
extern "C" int run_mode(int ng, void* dps,
void* pa, void* sa, void* pb, void* sb,
void* psfa, void* lsfa, void* psfb, void* lsfb,
void* pc, void* sc, void* pd, void* sd,
void* ws, size_t ws_sz, void* hps,
int mode, int cx, int cy, int fcx, int fcy,
bool run_kernel, bool use_pdl
) {
switch (mode) {
case 1: return run_impl<Gemm1SM>(ng,dps,pa,sa,pb,sb,psfa,lsfa,psfb,lsfb,pc,sc,pd,sd,ws,ws_sz,hps,cx,cy,fcx,fcy,run_kernel,use_pdl);
case 2: return run_impl<Gemm2SM>(ng,dps,pa,sa,pb,sb,psfa,lsfa,psfb,lsfb,pc,sc,pd,sd,ws,ws_sz,hps,cx,cy,fcx,fcy,run_kernel,use_pdl);
case 3: return run_impl<Gemm2SM_N128>(ng,dps,pa,sa,pb,sb,psfa,lsfa,psfb,lsfb,pc,sc,pd,sd,ws,ws_sz,hps,cx,cy,fcx,fcy,run_kernel,use_pdl);
}
return -9;
}
extern "C" void get_type_sizes(int* s) {
s[0]=sizeof(StrideA); s[1]=sizeof(StrideB); s[2]=sizeof(StrideC); s[3]=sizeof(StrideD);
s[4]=sizeof(InternalLayoutSFA); s[5]=sizeof(InternalLayoutSFB);
s[6]=sizeof(cute::Shape<int,int,int>);
}
template <typename GemmT>
void fill_meta_impl(int ng, int* mnk,
void* ps, void* sa, void* sb, void* sc, void* sd, void* lsfa, void* lsfb
) {
using SA = typename GemmT::GemmKernel::InternalStrideA;
using SB = typename GemmT::GemmKernel::InternalStrideB;
using SC = typename GemmT::GemmKernel::InternalStrideC;
using SD = typename GemmT::GemmKernel::InternalStrideD;
using LSFA = typename GemmT::GemmKernel::CollectiveMainloop::InternalLayoutSFA;
using LSFB = typename GemmT::GemmKernel::CollectiveMainloop::InternalLayoutSFB;
using Cfg = typename GemmT::GemmKernel::CollectiveMainloop::Sm1xxBlkScaledConfig;
auto* p=(cute::Shape<int,int,int>*)ps; auto* a=(SA*)sa; auto* b=(SB*)sb;
auto* c=(SC*)sc; auto* d=(SD*)sd; auto* la=(LSFA*)lsfa; auto* lb=(LSFB*)lsfb;
for(int i=0;i<ng;i++){
int M=mnk[i*3],N=mnk[i*3+1],K=mnk[i*3+2];
p[i]=cute::make_shape(M,N,K);
a[i]=cutlass::make_cute_packed_stride(SA{},{M,K,1});
b[i]=cutlass::make_cute_packed_stride(SB{},{N,K,1});
c[i]=cutlass::make_cute_packed_stride(SC{},{M,N,1});
d[i]=cutlass::make_cute_packed_stride(SD{},{M,N,1});
la[i]=Cfg::tile_atom_to_shape_SFA(cute::make_shape(M,N,K,1));
lb[i]=Cfg::tile_atom_to_shape_SFB(cute::make_shape(M,N,K,1));
}
}
extern "C" int fill_meta_mode(int mode, int ng, int* mnk,
void* ps, void* sa, void* sb, void* sc, void* sd, void* lsfa, void* lsfb
) {
switch(mode){
case 1: fill_meta_impl<Gemm1SM>(ng,mnk,ps,sa,sb,sc,sd,lsfa,lsfb); return 0;
case 2: fill_meta_impl<Gemm2SM>(ng,mnk,ps,sa,sb,sc,sd,lsfa,lsfb); return 0;
case 3: fill_meta_impl<Gemm2SM_N128>(ng,mnk,ps,sa,sb,sc,sd,lsfa,lsfb); return 0;
} return -1;
}
"""
CPP_SOURCE = r"""
#include <torch/extension.h>
#include <vector>
extern "C" int run_mode(int,void*,void*,void*,void*,void*,void*,void*,void*,void*,
void*,void*,void*,void*,void*,size_t,void*,int,int,int,int,int,bool,bool);
extern "C" void get_type_sizes(int*);
extern "C" int fill_meta_mode(int,int,int*,void*,void*,void*,void*,void*,void*,void*);
std::vector<int> get_sizes() { std::vector<int> s(7); get_type_sizes(s.data()); return s; }
int run_grouped(int ng,
at::Tensor dps, at::Tensor pa, at::Tensor sa, at::Tensor pb, at::Tensor sb,
at::Tensor psfa, at::Tensor lsfa, at::Tensor psfb, at::Tensor lsfb,
at::Tensor pc, at::Tensor sc, at::Tensor pd, at::Tensor sd,
at::Tensor ws, at::Tensor hps,
int mode, int cx, int cy, int fcx, int fcy, bool pdl
) {
return run_mode(ng,dps.data_ptr(),pa.data_ptr(),sa.data_ptr(),
pb.data_ptr(),sb.data_ptr(),psfa.data_ptr(),lsfa.data_ptr(),
psfb.data_ptr(),lsfb.data_ptr(),pc.data_ptr(),sc.data_ptr(),
pd.data_ptr(),sd.data_ptr(),ws.data_ptr(),ws.nbytes(),hps.data_ptr(),
mode,cx,cy,fcx,fcy,true,pdl);
}
int can_impl(int ng,
at::Tensor dps, at::Tensor pa, at::Tensor sa, at::Tensor pb, at::Tensor sb,
at::Tensor psfa, at::Tensor lsfa, at::Tensor psfb, at::Tensor lsfb,
at::Tensor pc, at::Tensor sc, at::Tensor pd, at::Tensor sd,
at::Tensor ws, at::Tensor hps,
int mode, int cx, int cy, int fcx, int fcy
) {
return run_mode(ng,dps.data_ptr(),pa.data_ptr(),sa.data_ptr(),
pb.data_ptr(),sb.data_ptr(),psfa.data_ptr(),lsfa.data_ptr(),
psfb.data_ptr(),lsfb.data_ptr(),pc.data_ptr(),sc.data_ptr(),
pd.data_ptr(),sd.data_ptr(),ws.data_ptr(),ws.nbytes(),hps.data_ptr(),
mode,cx,cy,fcx,fcy,false,false);
}
int fill_meta(int mode, int ng, at::Tensor mnk,
at::Tensor ps, at::Tensor sa, at::Tensor sb, at::Tensor sc, at::Tensor sd,
at::Tensor lsfa, at::Tensor lsfb
) {
return fill_meta_mode(mode,ng,mnk.data_ptr<int>(),
ps.data_ptr(),sa.data_ptr(),sb.data_ptr(),sc.data_ptr(),sd.data_ptr(),
lsfa.data_ptr(),lsfb.data_ptr());
}
struct Plan {
int ng, mode, cx, cy, fcx, fcy;
bool pdl;
at::Tensor dps,pa,sa,pb,sb,psfa,lsfa,psfb,lsfb,pc,sc,pd,sd,ws,hps;
};
static std::vector<Plan> g_plans;
int register_plan(int ng,
at::Tensor dps, at::Tensor pa, at::Tensor sa, at::Tensor pb, at::Tensor sb,
at::Tensor psfa, at::Tensor lsfa, at::Tensor psfb, at::Tensor lsfb,
at::Tensor pc, at::Tensor sc, at::Tensor pd, at::Tensor sd,
at::Tensor ws, at::Tensor hps,
int mode, int cx, int cy, int fcx, int fcy, bool pdl
) {
int ret = run_mode(ng,dps.data_ptr(),pa.data_ptr(),sa.data_ptr(),
pb.data_ptr(),sb.data_ptr(),psfa.data_ptr(),lsfa.data_ptr(),
psfb.data_ptr(),lsfb.data_ptr(),pc.data_ptr(),sc.data_ptr(),
pd.data_ptr(),sd.data_ptr(),ws.data_ptr(),ws.nbytes(),hps.data_ptr(),
mode,cx,cy,fcx,fcy,false,false);
if(ret!=0) return ret;
Plan p{ng,mode,cx,cy,fcx,fcy,pdl,dps,pa,sa,pb,sb,psfa,lsfa,psfb,lsfb,pc,sc,pd,sd,ws,hps};
int idx = g_plans.size();
g_plans.push_back(p);
return idx;
}
int run_plan(int idx) {
if(idx<0||idx>=(int)g_plans.size()) return -100;
auto& p=g_plans[idx];
return run_mode(p.ng,p.dps.data_ptr(),p.pa.data_ptr(),p.sa.data_ptr(),
p.pb.data_ptr(),p.sb.data_ptr(),p.psfa.data_ptr(),p.lsfa.data_ptr(),
p.psfb.data_ptr(),p.lsfb.data_ptr(),p.pc.data_ptr(),p.sc.data_ptr(),
p.pd.data_ptr(),p.sd.data_ptr(),p.ws.data_ptr(),p.ws.nbytes(),p.hps.data_ptr(),
p.mode,p.cx,p.cy,p.fcx,p.fcy,true,p.pdl);
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("get_sizes", &get_sizes);
m.def("run", &run_grouped);
m.def("can_impl", &can_impl);
m.def("fill_meta", &fill_meta);
m.def("register_plan", ®ister_plan);
m.def("run_plan", &run_plan);
}
"""
print("[BUILD] Starting build...", file=sys.stderr, flush=True)
_MODULE = torch.utils.cpp_extension.load_inline(
name="nvfp4_group_pdl_v2",
cpp_sources=CPP_SOURCE,
cuda_sources=CUDA_SOURCE,
extra_cuda_cflags=[
"-O3", "--use_fast_math",
"-gencode=arch=compute_100a,code=sm_100a",
"-I/opt/cutlass/4.3.0/include",
"-I/opt/cutlass/4.3.0/build/include",
"-std=c++17", "--expt-relaxed-constexpr",
],
extra_cflags=["-O3"],
extra_include_paths=["/usr/local/cuda/include"],
verbose=True,
)
print("[BUILD] Complete!", file=sys.stderr, flush=True)
_TYPE_SIZES = _MODULE.get_sizes()
_WORKSPACE = torch.empty(32 * 1024 * 1024, dtype=torch.uint8, device="cuda")
class ScaleCache:
def __init__(self):
self._cache = {}
def get_or_create(self, s3d, l_idx, dev):
key = (id(s3d), l_idx, s3d._version, str(dev))
c = self._cache.get(key)
if c and c[0]() is s3d: return c[1]
b = to_blocked(s3d[:,:,l_idx]).to(dev).contiguous()
self._cache[key] = (weakref.ref(s3d), b)
return b
class ReorderedScaleCache:
def __init__(self):
self._cache = {}
def get_or_create(self, s6d, l_idx, dev):
key = (id(s6d), l_idx, s6d._version, str(dev))
c = self._cache.get(key)
if c and c[0]() is s6d: return c[1]
s = s6d[..., l_idx]
if s.device != dev: s = s.to(dev)
b = s.permute(2,4,0,1,3).contiguous().view(-1)
self._cache[key] = (weakref.ref(s6d), b)
return b
_SC = ScaleCache()
_RSC = ReorderedScaleCache()
def _build_meta(ps_list, a_l, b_l, sfa_l, sfb_l, out_l, mode):
ng = len(ps_list)
sz = _TYPE_SIZES
mnk = torch.zeros(ng*3, dtype=torch.int32)
for i,(m,n,k) in enumerate(ps_list):
mnk[i*3], mnk[i*3+1], mnk[i*3+2] = m, n, k
h_ps = torch.zeros(ng*sz[6], dtype=torch.uint8)
h_sa = torch.zeros(ng*sz[0], dtype=torch.uint8)
h_sb = torch.zeros(ng*sz[1], dtype=torch.uint8)
h_sc = torch.zeros(ng*sz[2], dtype=torch.uint8)
h_sd = torch.zeros(ng*sz[3], dtype=torch.uint8)
h_lsfa = torch.zeros(ng*sz[4], dtype=torch.uint8)
h_lsfb = torch.zeros(ng*sz[5], dtype=torch.uint8)
ret = _MODULE.fill_meta(mode, ng, mnk, h_ps, h_sa, h_sb, h_sc, h_sd, h_lsfa, h_lsfb)
if ret != 0: raise RuntimeError(f"fill_meta failed mode={mode} ret={ret}")
dev = a_l[0].device
return {
"d_ps": h_ps.to(dev), "h_ps": h_ps,
"d_sa": h_sa.to(dev), "d_sb": h_sb.to(dev),
"d_sc": h_sc.to(dev), "d_sd": h_sd.to(dev),
"d_lsfa": h_lsfa.to(dev), "d_lsfb": h_lsfb.to(dev),
"ptr_a": torch.tensor([t.data_ptr() for t in a_l], dtype=torch.int64, device=dev),
"ptr_b": torch.tensor([t.data_ptr() for t in b_l], dtype=torch.int64, device=dev),
"ptr_sfa": torch.tensor([t.data_ptr() for t in sfa_l], dtype=torch.int64, device=dev),
"ptr_sfb": torch.tensor([t.data_ptr() for t in sfb_l], dtype=torch.int64, device=dev),
"ptr_c": torch.tensor([t.data_ptr() for t in out_l], dtype=torch.int64, device=dev),
"ptr_d": torch.tensor([t.data_ptr() for t in out_l], dtype=torch.int64, device=dev),
}
_META_CACHE = {}
_MODE_CACHE = {}
_PLAN_CACHE = {}
def _try_config(mode, cx, cy, fcx, fcy, ps_list, a_l, b_l, sfa_l, sfb_l, out_l):
try:
meta = _build_meta(ps_list, a_l, b_l, sfa_l, sfb_l, out_l, mode)
except RuntimeError:
return False, None
ret = _MODULE.can_impl(
len(ps_list), meta["d_ps"],
meta["ptr_a"], meta["d_sa"], meta["ptr_b"], meta["d_sb"],
meta["ptr_sfa"], meta["d_lsfa"], meta["ptr_sfb"], meta["d_lsfb"],
meta["ptr_c"], meta["d_sc"], meta["ptr_d"], meta["d_sd"],
_WORKSPACE, meta["h_ps"], mode, cx, cy, fcx, fcy,
)
return ret == 0, meta
def _select_mode(ps_list, a_l, b_l, sfa_l, sfb_l, out_l):
key = tuple(ps_list)
if key in _MODE_CACHE: return _MODE_CACHE[key]
ng = len(ps_list)
min_k = min(k for _,_,k in ps_list)
configs = []
if min_k >= 256:
if ng <= 4:
configs = [
(3, 2, 2, 2, 1, "2SM_N128_c22"),
(2, 2, 2, 2, 1, "2SM_N256_c22"),
(3, 2, 1, 2, 1, "2SM_N128"),
(2, 2, 1, 2, 1, "2SM_N256"),
]
else:
configs = [
(2, 2, 2, 2, 1, "2SM_N256_c22"),
(3, 2, 2, 2, 1, "2SM_N128_c22"),
(2, 2, 1, 2, 1, "2SM_N256"),
(3, 2, 1, 2, 1, "2SM_N128"),
]
configs.append((1, 1, 1, 1, 1, "1SM"))
for mode, cx, cy, fcx, fcy, label in configs:
ok, _ = _try_config(mode, cx, cy, fcx, fcy, ps_list, a_l, b_l, sfa_l, sfb_l, out_l)
if ok:
print(f"[MODE] {label} (mode={mode} cluster=({cx},{cy}) fallback=({fcx},{fcy}))", file=sys.stderr, flush=True)
_MODE_CACHE[key] = (mode, cx, cy, fcx, fcy)
return mode, cx, cy, fcx, fcy
print(f"[MODE] {label}: unavailable", file=sys.stderr, flush=True)
raise RuntimeError("No valid mode")
def _build_plan(abc, sfasfb, sfasfb_r, ps):
a_l, b_l, sfa_l, sfb_l, out_l, ps_l = [], [], [], [], [], []
use_r = sfasfb_r is not None
if use_r:
for (a,b,c),(_,_),(sr_a,sr_b),(m,n,k,l) in zip(abc, sfasfb, sfasfb_r, ps):
for li in range(l):
a_l.append(a.view(torch.uint8)[:,:,li].contiguous())
b_l.append(b.view(torch.uint8)[:,:,li].contiguous())
sfa_l.append(_RSC.get_or_create(sr_a, li, a_l[-1].device))
sfb_l.append(_RSC.get_or_create(sr_b, li, a_l[-1].device))
out_l.append(c[:,:,li])
ps_l.append((int(m),int(n),int(k)))
else:
for (a,b,c),(sfa,sfb),(m,n,k,l) in zip(abc, sfasfb, ps):
for li in range(l):
a_l.append(a.view(torch.uint8)[:,:,li].contiguous())
b_l.append(b.view(torch.uint8)[:,:,li].contiguous())
sfa_l.append(_SC.get_or_create(sfa, li, a_l[-1].device))
sfb_l.append(_SC.get_or_create(sfb, li, a_l[-1].device))
out_l.append(c[:,:,li])
ps_l.append((int(m),int(n),int(k)))
mode, cx, cy, fcx, fcy = _select_mode(ps_l, a_l, b_l, sfa_l, sfb_l, out_l)
meta = _build_meta(ps_l, a_l, b_l, sfa_l, sfb_l, out_l, mode)
use_pdl = True
plan_id = _MODULE.register_plan(
len(ps_l), meta["d_ps"],
meta["ptr_a"], meta["d_sa"], meta["ptr_b"], meta["d_sb"],
meta["ptr_sfa"], meta["d_lsfa"], meta["ptr_sfb"], meta["d_lsfb"],
meta["ptr_c"], meta["d_sc"], meta["ptr_d"], meta["d_sd"],
_WORKSPACE, meta["h_ps"], mode, cx, cy, fcx, fcy, use_pdl,
)
if plan_id < 0: raise RuntimeError(f"register_plan failed: {plan_id}")
return {"plan_id": plan_id, "mode": mode}
def _get_plan(abc, sfasfb, sfasfb_r, ps):
use_r = sfasfb_r is not None
kp = [use_r]
if use_r:
for (a,b,c),(sr_a,sr_b),(m,n,k,l) in zip(abc, sfasfb_r, ps):
kp.append((id(a),a._version,id(b),b._version,id(c),
id(sr_a),sr_a._version,id(sr_b),sr_b._version,int(m),int(n),int(k),int(l)))
else:
for (a,b,c),(sfa,sfb),(m,n,k,l) in zip(abc, sfasfb, ps):
kp.append((id(a),a._version,id(b),b._version,id(c),
id(sfa),sfa._version,id(sfb),sfb._version,int(m),int(n),int(k),int(l)))
key = tuple(kp)
if key in _PLAN_CACHE: return _PLAN_CACHE[key]
plan = _build_plan(abc, sfasfb, sfasfb_r, ps)
_PLAN_CACHE[key] = plan
return plan
def ref_kernel(data):
abc, sfasfb, _, ps = _unpack_input(data)
result = []
for (a,b,c),(sfa,sfb),(_,_,_,l) in zip(abc, sfasfb, ps):
for li in range(l):
av = a[:,:,li].view(torch.float4_e2m1fn_x2)
bv = b[:,:,li].transpose(0,1).view(torch.float4_e2m1fn_x2)
sa = _SC.get_or_create(sfa, li, av.device)
sb = _SC.get_or_create(sfb, li, av.device)
torch._scaled_mm(av, bv, sa, sb, bias=None, out_dtype=torch.float16, out=c[:,:,li])
result.append(c)
return result
def custom_kernel(data):
abc, sfasfb, sfasfb_r, ps = _unpack_input(data)
plan = _get_plan(abc, sfasfb, sfasfb_r, ps)
ret = _MODULE.run_plan(plan["plan_id"])
if ret != 0: raise RuntimeError(f"CUTLASS failed code={ret}")
return [c for _,_,c in abc]
def create_reordered_scale_factor_tensor(l, mn, k, ref_f8_tensor):
sf_k = ceil_div(k, sf_vec_size)
atom_m = (32, 4)
atom_k = 4
mma_shape = (l, ceil_div(mn, atom_m[0]*atom_m[1]), ceil_div(sf_k, atom_k), atom_m[0], atom_m[1], atom_k)
rand_int = torch.randint(1, 3, mma_shape, dtype=torch.int8, device="cuda")
reordered = rand_int.to(dtype=torch.float8_e4m3fn).permute(3,4,1,5,2,0)
if ref_f8_tensor.device.type == "cpu": ref_f8_tensor = ref_f8_tensor.cuda()
i_idx = torch.arange(mn, device="cuda")
j_idx = torch.arange(sf_k, device="cuda")
b_idx = torch.arange(l, device="cuda")
i_grid, j_grid, b_grid = torch.meshgrid(i_idx, j_idx, b_idx, indexing="ij")
mm = i_grid // (atom_m[0]*atom_m[1])
mm32 = i_grid % atom_m[0]
mm4 = (i_grid % 128) // atom_m[0]
kk = j_grid // atom_k
kk4 = j_grid % atom_k
reordered[mm32, mm4, mm, kk4, kk, b_grid] = ref_f8_tensor[i_grid, j_grid, b_grid]
return reordered
def _create_fp4_tensors(l, mn, k):
ref_i8 = torch.randint(255, size=(l, mn, k//2), dtype=torch.uint8, device="cuda")
return (ref_i8 & 0b1011_1011).permute(1,2,0).view(torch.float4_e2m1fn_x2)
def generate_input(m, n, k, g, seed):
torch.manual_seed(seed)
abc, sfasfb, sfasfb_r, ps = [], [], [], []
l = 1
for gi in range(g):
mi, ni, ki = m[gi], n[gi], k[gi]
a = _create_fp4_tensors(l, mi, ki)
b = _create_fp4_tensors(l, ni, ki)
c = torch.randn((l, mi, ni), dtype=torch.float16, device="cuda").permute(1,2,0)
sf_k = ceil_div(ki, sf_vec_size)
sfa_cpu = torch.randint(1,3,(l,mi,sf_k),dtype=torch.int8).to(dtype=torch.float8_e4m3fn).permute(1,2,0)
sfb_cpu = torch.randint(1,3,(l,ni,sf_k),dtype=torch.int8).to(dtype=torch.float8_e4m3fn).permute(1,2,0)
sfa_r = create_reordered_scale_factor_tensor(l, mi, ki, sfa_cpu)
sfb_r = create_reordered_scale_factor_tensor(l, ni, ki, sfb_cpu)
abc.append((a,b,c)); sfasfb.append((sfa_cpu,sfb_cpu)); sfasfb_r.append((sfa_r,sfb_r))
ps.append((mi,ni,ki,l))
return (abc, sfasfb, sfasfb_r, ps)
check_implementation = make_match_reference(ref_kernel, rtol=1e-03, atol=1e-03)
scrolls · 584 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 502078.
⋯ diff truncated: revisions differ almost entirely
Best evidence level for this revision: reported
JSON