submission 401132
jon9509 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1645 lines, June 9 Researcher Reciprocity License v1.0.
refine.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-401132?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:761b92693b06da83749795ba08735b8de947c5a213e30dbda6398edb98e6f4e1
license declaredunknown
license concludedunknown
authorsjon9509
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
cluster
"problem_shapes.get_host_problem_shape(), TileShape{}, AtomThrShapeMNK{}, ClusterShape{},\n args.hw_info, args.scheduler, scheduler_workspace",fp4
CUTLASS SM100 (B200) grouped GEMM for NVFP4 with blockwise scaling.fused-epilogue
using EpilogueSchedule = cutlass::epilogue::PtrArrayTmaWarpSpecialized1Sm;warp-specialization
"sm100_gemm_array_tma_warpspecialized.hpp",Kernel source
refine.py1645 lines
import os
import time
import torch
import subprocess
from collections import OrderedDict
from torch.utils.cpp_extension import load_inline, include_paths
def custom_kernel(data):
"""
CUTLASS SM100 (B200) grouped GEMM for NVFP4 with blockwise scaling.
Based on example 81: blackwell_grouped_gemm_blockwise.cu pattern.
"""
# -----------------------------
# 0) Unpack input (robust)
# -----------------------------
if isinstance(data, (list, tuple)) and len(data) == 3:
abc_tensors, sfasfb_tensors, problem_sizes = data
elif isinstance(data, (list, tuple)) and len(data) == 4:
abc_tensors, _, sfasfb_tensors, problem_sizes = data
else:
raise TypeError("Unexpected data format for custom_kernel")
g = len(problem_sizes)
if g == 0:
return []
# -----------------------------
# 0.5) Optional Python cache (shape-keyed; ptr tensors updated in-place)
# -----------------------------
if not hasattr(custom_kernel, "_CACHE"):
custom_kernel._CACHE = OrderedDict()
def _tensor_key(t):
return (
tuple(int(x) for x in t.size()),
tuple(int(x) for x in t.stride()),
str(t.dtype),
str(t.device),
)
def _build_cache_key(abc_tensors, sfasfb_tensors, problem_sizes):
abc_sig = tuple((_tensor_key(a), _tensor_key(b), _tensor_key(c)) for a, b, c in abc_tensors)
sfs_sig = tuple((_tensor_key(sfa), _tensor_key(sfb)) for sfa, sfb in sfasfb_tensors)
ps_sig = tuple((int(m), int(n), int(k), int(l)) for (m, n, k, l) in problem_sizes)
return (abc_sig, sfs_sig, ps_sig)
def _bucket_id(problem_sizes, mod):
# Deterministic FNV-1a over shapes
h = 1469598103934665603
for (m, n, k, l) in problem_sizes:
for v in (m, n, k, l):
h ^= int(v) & 0xFFFFFFFFFFFFFFFF
h *= 1099511628211
h &= 0xFFFFFFFFFFFFFFFF
return int(h % mod)
use_py_cache = int(os.environ.get("NVFP4_PY_CACHE", "1")) > 0
use_echokey = int(os.environ.get("NVFP4_ECHOKEY", "0")) > 0
echokey_debug = int(os.environ.get("NVFP4_ECHOKEY_DEBUG", "0")) > 0
use_warmup = int(os.environ.get("NVFP4_ECHOKEY_WARMUP", "0")) > 0
map_size = int(os.environ.get("NVFP4_ECHOKEY_MAP_SIZE", "16384"))
map_size = max(1, map_size)
cache_key = None
cached_state = None
if not hasattr(custom_kernel, "_ECHO_DUMMY"):
custom_kernel._ECHO_DUMMY = (
torch.empty(0, dtype=torch.int64, device=abc_tensors[0][2].device),
torch.empty(0, dtype=torch.float32, device=abc_tensors[0][2].device),
)
if not hasattr(custom_kernel, "_POTENTIAL_MAP"):
custom_kernel._POTENTIAL_MAP = {}
if not hasattr(custom_kernel, "_SYNERGY_OUT"):
custom_kernel._SYNERGY_OUT = {}
# Defer cache lookup until after padding / L-expansion so ptrs are up to date.
timing_py = int(os.environ.get("NVFP4_TIMING_PY", "0"))
if timing_py:
t_py0 = time.perf_counter()
# -----------------------------
# 1) Lazy-build extension
# -----------------------------
if not hasattr(custom_kernel, "_EXT"):
cuda_root = os.environ.get("CUDA_HOME", "/usr/local/cuda")
cuda_inc = os.path.join(cuda_root, "targets", "x86_64-linux", "include")
cccl_inc = os.path.join(cuda_inc, "cccl")
repo_path = "cutlass_repo"
cutlass_inc = os.path.join(repo_path, "include")
cutlass_util_inc = os.path.join(repo_path, "tools", "util", "include")
header_check = os.path.join(cutlass_inc, "cutlass", "cutlass.h")
if not os.path.exists(header_check):
os.makedirs(repo_path, exist_ok=True)
try:
subprocess.run(
["git", "clone", "--depth", "1", "https://github.com/NVIDIA/cutlass.git", repo_path],
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
)
except subprocess.CalledProcessError as e:
msg = e.stderr.decode(errors="ignore")
raise RuntimeError(f"Could not clone CUTLASS. Details:\n{msg}")
# Pin to the CUTLASS commit observed on the server.
cutlass_commit = "2fafefb7b9c7e233c1cfb258f38139a6406d1511"
try:
subprocess.run(
["git", "-C", repo_path, "checkout", cutlass_commit],
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
)
except subprocess.CalledProcessError:
subprocess.run(
["git", "-C", repo_path, "fetch", "--depth", "1", "origin", cutlass_commit],
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
)
subprocess.run(
["git", "-C", repo_path, "checkout", cutlass_commit],
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
)
# Patch the CUTLASS bug (apply even if repo already existed)
kernel_file = os.path.join(
repo_path,
"include",
"cutlass",
"gemm",
"kernel",
"sm100_gemm_array_tma_warpspecialized.hpp",
)
if not os.path.exists(kernel_file):
raise RuntimeError("CUTLASS kernel header not found to patch.")
with open(kernel_file, "r") as f:
content = f.read()
patched = content.replace(
"args.hw_info, scheduler_workspace",
"args.hw_info, args.scheduler, scheduler_workspace",
)
patched = patched.replace(
"problem_shapes.get_host_problem_shape(), TileShape{}, AtomThrShapeMNK{}, ClusterShape{},\n args.hw_info, args.scheduler, scheduler_workspace",
"problem_shapes.get_host_problem_shape(), TileShape{}, ClusterShape{},\n args.hw_info, args.scheduler, scheduler_workspace",
)
patched = patched.replace(
"grid_shape = TileScheduler::get_grid_shape(\n params.scheduler,\n params.problem_shape.get_host_problem_shape(),\n TileShape{},\n AtomThrShapeMNK{},\n cluster_shape,\n params.hw_info);",
"grid_shape = TileScheduler::get_grid_shape(\n params.scheduler,\n params.problem_shape.get_host_problem_shape(),\n TileShape{},\n cluster_shape,\n params.hw_info,\n TileSchedulerArguments{});",
)
if patched != content:
with open(kernel_file, "w") as f:
f.write(patched)
# Inject in-kernel freework inside the SM100 pipeline consumer wait (optional compile-time)
pipeline_file = os.path.join(
repo_path,
"include",
"cutlass",
"pipeline",
"sm100_pipeline.hpp",
)
if not os.path.exists(pipeline_file):
raise RuntimeError("CUTLASS pipeline header not found to patch.")
with open(pipeline_file, "r") as f:
pcontent = f.read()
old_wait = (
" void consumer_wait(uint32_t stage, uint32_t phase, ConsumerToken barrier_token) {\n"
" detail::pipeline_check_is_consumer(params_.role);\n"
" if (barrier_token == BarrierStatus::WaitAgain) {\n"
" full_barrier_ptr_[stage].wait(phase);\n"
" }\n"
" }\n"
)
new_wait = (
" void consumer_wait(uint32_t stage, uint32_t phase, ConsumerToken barrier_token) {\n"
" detail::pipeline_check_is_consumer(params_.role);\n"
" if (barrier_token == BarrierStatus::WaitAgain) {\n"
"#if defined(NVFP4_INKERNEL_BURN) && (NVFP4_INKERNEL_BURN > 0)\n"
" auto do_freework = [&]() {\n"
" float acc = float(lane_idx_ + 1);\n"
" #pragma unroll 4\n"
" for (int i = 0; i < NVFP4_INKERNEL_BURN; ++i) {\n"
" acc = acc * 1.000001f + 0.000001f;\n"
" }\n"
" asm volatile(\"\" :: \"f\"(acc));\n"
" };\n"
" while (!full_barrier_ptr_[stage].try_wait(phase)) {\n"
" do_freework();\n"
" }\n"
"#if defined(NVFP4_INKERNEL_BURN_MARKER)\n"
" asm volatile(\"/* NVFP4_INKERNEL_BURN_MARKER=%0 */\" :: \"n\"(NVFP4_INKERNEL_BURN_MARKER));\n"
"#endif\n"
"#else\n"
" full_barrier_ptr_[stage].wait(phase);\n"
"#endif\n"
" }\n"
" }\n"
)
if old_wait in pcontent:
pcontent = pcontent.replace(old_wait, new_wait)
with open(pipeline_file, "w") as f:
f.write(pcontent)
if not os.path.exists(header_check):
raise RuntimeError("CUTLASS headers not found after auto-heal.")
# Include paths
extra_includes = []
for p in (cutlass_inc, cutlass_util_inc, cuda_inc, cccl_inc):
if p and os.path.isdir(p):
extra_includes.append(p)
for p in list(dict.fromkeys(include_paths())):
if os.path.isfile(os.path.join(p, "cutlass", "cutlass.h")):
extra_includes.append(p)
seen = set()
extra_includes = [p for p in extra_includes if not (p in seen or seen.add(p))]
cuda_src = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <vector>
#include <stdexcept>
#include <string>
#include <cstdio>
#include <cstdint>
#include <cstring>
#include <cstdlib>
#include <cmath>
#include <unordered_map>
#include <mutex>
#include <chrono>
// Core CUTLASS + CuTe
#include "cutlass/cutlass.h"
#include "cutlass/numeric_types.h"
#include "cutlass/numeric_conversion.h"
#include <cute/arch/config.hpp>
#if !defined(__CUDA_ARCH__)
#undef CUTE_ARCH_TMA_SM90_ENABLED
#undef CUTE_ARCH_TMA_SM100_ENABLED
#endif
#include "cute/tensor.hpp"
// Dispatch policies
#include "cutlass/gemm/dispatch_policy.hpp"
#include "cutlass/epilogue/dispatch_policy.hpp"
// GEMM infrastructure
#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"
// Packed stride helper
#include "cutlass/util/packed_stride.hpp"
using namespace cute;
static inline void check_cuda_tensor(const torch::Tensor& t, const char* name) {
if (!t.defined()) throw std::runtime_error(std::string(name) + " is undefined");
if (!t.is_cuda()) throw std::runtime_error(std::string(name) + " must be CUDA");
if (!t.is_contiguous()) throw std::runtime_error(std::string(name) + " must be contiguous");
}
static inline void check_int32_tensor(const torch::Tensor& t, const char* name) {
if (!t.defined()) throw std::runtime_error(std::string(name) + " is undefined");
if (t.scalar_type() != torch::kInt32) throw std::runtime_error(std::string(name) + " must be int32");
if (!t.is_contiguous()) throw std::runtime_error(std::string(name) + " must be contiguous");
if (!(t.is_cuda() || t.is_cpu())) throw std::runtime_error(std::string(name) + " must be CPU or CUDA");
}
static inline void check_float_tensor_optional(const torch::Tensor& t, const char* name) {
if (!t.defined()) return;
if (!t.is_cuda()) throw std::runtime_error(std::string(name) + " must be CUDA");
if (t.scalar_type() != torch::kFloat32) throw std::runtime_error(std::string(name) + " must be float32");
if (!t.is_contiguous()) throw std::runtime_error(std::string(name) + " must be contiguous");
}
static inline void check_int64_tensor_optional(const torch::Tensor& t, const char* name) {
if (!t.defined()) return;
if (!t.is_cuda()) throw std::runtime_error(std::string(name) + " must be CUDA");
if (t.scalar_type() != torch::kInt64) throw std::runtime_error(std::string(name) + " must be int64");
if (!t.is_contiguous()) throw std::runtime_error(std::string(name) + " must be contiguous");
}
static inline std::string status_to_string(cutlass::Status st) {
switch (st) {
case cutlass::Status::kSuccess: return "kSuccess";
case cutlass::Status::kErrorInvalidProblem: return "kErrorInvalidProblem";
case cutlass::Status::kErrorNotSupported: return "kErrorNotSupported";
default: return "cutlass::Status(" + std::to_string(int(st)) + ")";
}
}
static inline int timing_mode() {
static int mode = -1;
if (mode == -1) {
const char* env = std::getenv("NVFP4_TIMING");
mode = env ? std::atoi(env) : 0;
}
return mode;
}
static inline int timing_abort() {
static int mode = -1;
if (mode == -1) {
const char* env = std::getenv("NVFP4_TIMING_ABORT");
mode = env ? std::atoi(env) : 0;
}
return mode;
}
static inline int timing_detail() {
static int mode = -1;
if (mode == -1) {
const char* env = std::getenv("NVFP4_TIMING_DETAIL");
mode = env ? std::atoi(env) : 0;
}
return mode;
}
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
// Block-scaled tensor core ops
using ArchTag = cutlass::arch::Sm100;
using OperatorClass = cutlass::arch::OpClassBlockScaledTensorOp;
// Problem shape (L is not supported in this kernel; use 3D M,N,K)
using ProblemShape = cutlass::gemm::GroupProblemShape<Shape<int,int,int>>;
using UnderlyingProblemShape = typename ProblemShape::UnderlyingProblemShape;
// Element types
using ElementAccumulator = float;
using ElementCompute = float;
// FROZEN: Element types below are validated. Do not change unless explicitly requested.
// Touching these will loop back on prior fixes.
using ElementInput = cutlass::float_e2m1_t;
using ElementA = cutlass::nv_float4_t<ElementInput>;
using ElementB = cutlass::nv_float4_t<ElementInput>;
using ElementC = cutlass::half_t;
using ElementD = ElementC;
// Layout tags
using LayoutA = cutlass::layout::RowMajor*;
using LayoutB = cutlass::layout::ColumnMajor*;
using LayoutC = cutlass::layout::RowMajor;
using LayoutD = cutlass::layout::RowMajor;
// Alignments
constexpr int AlignmentA = 32;
constexpr int AlignmentB = 32;
constexpr int AlignmentC = 8;
constexpr int AlignmentD = 8;
// Tile + cluster shape
using MmaTileShape = Shape<_128,_128,_128>;
using ClusterShape = Shape<_1,_1,_1>;
// FROZEN: Schedule tags below are validated. Do not change unless explicitly requested.
// Touching these will loop back on prior fixes.
using KernelSchedule = cutlass::gemm::KernelPtrArrayTmaWarpSpecialized1SmNvf4Sm100;
using EpilogueSchedule = cutlass::epilogue::PtrArrayTmaWarpSpecialized1Sm;
// Epilogue tile
using EpilogueTileShape = Shape<_128,_64>;
// Build epilogue first
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
ArchTag, OperatorClass,
MmaTileShape, ClusterShape,
EpilogueTileShape,
ElementAccumulator, ElementCompute,
ElementC, LayoutC*, AlignmentC,
ElementD, LayoutD*, AlignmentD,
EpilogueSchedule
>::CollectiveOp;
// Build mainloop with element-pair tuple for block-scaled GEMM
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
ArchTag, OperatorClass,
ElementA, LayoutA, AlignmentA,
ElementB, LayoutB, AlignmentB,
ElementAccumulator,
MmaTileShape, ClusterShape,
cutlass::gemm::collective::StageCountAutoCarveout<
static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))
>,
KernelSchedule
>::CollectiveOp;
// Get SFVecSize from the mainloop (this is the part that must match the kernel)
static constexpr int SFVecSize = int(CollectiveMainloop::SFVecSize);
using Sm1xxBlkScaledConfig = cutlass::detail::Sm1xxBlockScaledConfig<SFVecSize>;
// ScaleConfig + layouts (match mainloop expected types)
using LayoutSFA = typename cutlass::detail::LayoutSFAType<CollectiveMainloop>::type;
using LayoutSFB = typename cutlass::detail::LayoutSFBType<CollectiveMainloop>::type;
using InternalLayoutSFA = cute::remove_pointer_t<LayoutSFA>;
using InternalLayoutSFB = cute::remove_pointer_t<LayoutSFB>;
using ElementSF = typename cutlass::detail::ElementSFType<CollectiveMainloop>::type;
// Kernel
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
ProblemShape,
CollectiveMainloop,
CollectiveEpilogue
>;
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
// Extract strides
using StrideA = typename CollectiveMainloop::StrideA;
using InternalStrideA = typename CollectiveMainloop::InternalStrideA;
using StrideB = typename CollectiveMainloop::StrideB;
using InternalStrideB = typename CollectiveMainloop::InternalStrideB;
using StrideC = typename Gemm::GemmKernel::InternalStrideC;
using StrideD = typename Gemm::GemmKernel::InternalStrideD;
struct CacheEntry {
at::Tensor dev_buffer;
at::Tensor workspace;
size_t total_bytes = 0;
size_t off_sA = 0;
size_t off_sB = 0;
size_t off_sC = 0;
size_t off_sD = 0;
size_t off_lSFA = 0;
size_t off_lSFB = 0;
size_t off_ps = 0;
int device = -1;
int64_t g = 0;
std::vector<UnderlyingProblemShape> ps_host_shape;
};
static std::unordered_map<std::string, CacheEntry> g_cache;
static std::mutex g_cache_mutex;
static std::mutex g_last_mutex;
static int g_last_device = -1;
static int64_t g_last_g = 0;
static std::vector<int32_t> g_last_ps_host;
static int64_t g_last_ps_ptr = 0;
static std::string g_last_key;
static std::unordered_map<int, cutlass::KernelHardwareInfo> g_hw_cache;
static std::mutex g_hw_mutex;
static inline std::string make_cache_key(int device, int64_t g, const std::vector<int32_t>& ps_host) {
std::string key;
key.reserve(size_t(16 + ps_host.size() * 6));
key += std::to_string(device);
key.push_back(':');
key += std::to_string((long long)g);
for (int32_t v : ps_host) {
key.push_back(',');
key += std::to_string(int(v));
}
return key;
}
static inline std::string make_cache_key_ptr(int device, int64_t g, int64_t ps_ptr, int64_t nvals) {
std::string key;
key.reserve(64);
key += std::to_string(device);
key.push_back(':');
key += std::to_string((long long)g);
key += ":ptr=" + std::to_string((long long)ps_ptr);
key += ":n=" + std::to_string((long long)nvals);
return key;
}
__global__ void sf_float_to_ue4m3_kernel(const float* in, ElementSF* out, int64_t n) {
int64_t idx = int64_t(blockIdx.x) * blockDim.x + threadIdx.x;
if (idx < n) {
cutlass::NumericConverter<ElementSF, float, cutlass::FloatRoundStyle::round_to_nearest> cvt;
out[idx] = cvt(in[idx]);
}
}
__global__ void pad_copy_fp4_kernel(
uint8_t* dst,
const uint8_t* src,
int64_t total) {
int64_t idx = int64_t(blockIdx.x) * blockDim.x + threadIdx.x;
if (idx < total) {
dst[idx] = src[idx];
}
}
__global__ void pad_copy_c_half_kernel(
cutlass::half_t* dst,
const cutlass::half_t* src,
int64_t m,
int64_t n,
int64_t dst_stride_m,
int64_t src_stride_m) {
int64_t idx = int64_t(blockIdx.x) * blockDim.x + threadIdx.x;
int64_t total = m * n;
if (idx < total) {
int64_t row = idx / n;
int64_t col = idx - row * n;
dst[row * dst_stride_m + col] = src[row * src_stride_m + col];
}
}
__global__ void synergy_prologue_kernel(int64_t* timestamps) {
if (threadIdx.x == 0) {
uint64_t t;
asm volatile("mov.u64 %0, %%globaltimer;" : "=l"(t));
timestamps[0] = (int64_t)t;
}
}
__global__ void synergy_epilogue_kernel(
int64_t* timestamps,
float* potential_map,
int64_t map_size,
int64_t bucket_id,
int64_t total_work_items,
int burn) {
float acc0 = 0.0f;
if (burn > 0) {
float acc = float(threadIdx.x + 1);
#pragma unroll 4
for (int i = 0; i < burn; ++i) {
acc = acc * 1.000001f + 0.000001f;
}
if (threadIdx.x == 0) {
acc0 = acc;
if (timestamps) {
timestamps[0] ^= (int64_t)(acc);
}
}
}
if (threadIdx.x == 0) {
uint64_t t1;
asm volatile("mov.u64 %0, %%globaltimer;" : "=l"(t1));
int64_t t0 = timestamps[0];
int64_t dt = (int64_t)t1 - t0;
float resonance = 0.0f;
int64_t cascade = 0;
if (potential_map && map_size > 0 && total_work_items > 0) {
int64_t idx = bucket_id % map_size;
if (idx < 0) idx += map_size;
float old_v = potential_map[idx];
float time_per_item = float(dt) / float(total_work_items);
float base_v = (old_v == 0.0f) ? time_per_item : old_v;
float expected = base_v * float(total_work_items);
float tau_res = 100000.0f;
float dilation = float(dt) - expected;
resonance = tanhf(dilation / tau_res);
if (fabsf(resonance) > 0.95f) {
cascade = 1;
}
float alpha = 0.002f;
float new_v = (old_v == 0.0f) ? time_per_item : ((1.0f - alpha) * old_v + alpha * time_per_item);
potential_map[idx] = new_v;
}
timestamps[1] = (int64_t)t1;
timestamps[2] = dt;
timestamps[3] = (int64_t)(resonance * 1000000.0f);
timestamps[4] = cascade;
if (burn > 0) {
timestamps[4] += burn;
if (potential_map && map_size > 0) {
potential_map[0] += acc0 * 1.0e-7f;
}
}
}
}
torch::Tensor convert_sf_float_to_ue4m3(torch::Tensor in) {
check_cuda_tensor(in, "sf_in");
if (in.scalar_type() != torch::kFloat32) {
throw std::runtime_error("convert_sf_float_to_ue4m3 expects float32 input");
}
auto out = torch::empty_like(in, torch::TensorOptions().dtype(torch::kUInt8).device(in.device()));
int64_t n = in.numel();
if (n == 0) return out;
int threads = 256;
int blocks = int((n + threads - 1) / threads);
sf_float_to_ue4m3_kernel<<<blocks, threads, 0, 0>>>(
in.data_ptr<float>(),
reinterpret_cast<ElementSF*>(out.data_ptr<uint8_t>()),
n);
return out;
}
torch::Tensor pad_a_fp4(torch::Tensor a, int64_t m_pad) {
check_cuda_tensor(a, "a");
auto sizes = a.sizes();
if (sizes.size() != 3) {
throw std::runtime_error("pad_a_fp4 expects a 3D tensor");
}
int64_t m = sizes[0];
int64_t k2 = sizes[1];
int64_t l = sizes[2];
if (m_pad < m) {
throw std::runtime_error("pad_a_fp4 expects m_pad >= m");
}
auto out = torch::empty({m_pad, k2, l}, a.options());
int64_t total = m * k2 * l;
int64_t total_bytes = int64_t(out.numel()) * int64_t(sizeof(ElementA));
int64_t valid_bytes = total * int64_t(sizeof(ElementA));
if (valid_bytes < total_bytes) {
uint8_t* base = reinterpret_cast<uint8_t*>(out.data_ptr());
cudaMemset(base + valid_bytes, 0, size_t(total_bytes - valid_bytes));
}
if (total == 0) return out;
int threads = 256;
int blocks = int((total + threads - 1) / threads);
pad_copy_fp4_kernel<<<blocks, threads, 0, 0>>>(
reinterpret_cast<uint8_t*>(out.data_ptr()),
reinterpret_cast<const uint8_t*>(a.data_ptr()),
total);
return out;
}
torch::Tensor pad_c_half(torch::Tensor c, int64_t m_pad, int64_t n_pad) {
check_cuda_tensor(c, "c");
auto sizes = c.sizes();
if (sizes.size() != 3) {
throw std::runtime_error("pad_c_half expects a 3D tensor");
}
int64_t m = sizes[0];
int64_t n = sizes[1];
int64_t l = sizes[2];
if (l != 1) {
throw std::runtime_error("pad_c_half expects L=1");
}
if (m_pad < m || n_pad < n) {
throw std::runtime_error("pad_c_half expects m_pad >= m and n_pad >= n");
}
auto out = torch::empty({m_pad, n_pad, l}, c.options());
int64_t total_elems = out.numel();
if (total_elems > 0) {
cudaMemset(out.data_ptr(), 0, size_t(total_elems) * sizeof(ElementC));
}
int64_t total = m * n;
if (total == 0) return out;
int threads = 256;
int blocks = int((total + threads - 1) / threads);
pad_copy_c_half_kernel<<<blocks, threads, 0, 0>>>(
reinterpret_cast<cutlass::half_t*>(out.data_ptr()),
reinterpret_cast<const cutlass::half_t*>(c.data_ptr()),
m,
n,
n_pad,
n);
return out;
}
void grouped_nvfp4_tma_gemm(
torch::Tensor problem_sizes, // int32 [g,3]
torch::Tensor ptr_A, // int64 [g]
torch::Tensor ptr_B, // int64 [g]
torch::Tensor ptr_SFA, // int64 [g]
torch::Tensor ptr_SFB, // int64 [g]
torch::Tensor ptr_C, // int64 [g]
torch::Tensor ptr_D, // int64 [g]
torch::Tensor synergy_out, // int64 [>=5] or empty
torch::Tensor potential_map, // float32 [map_size] or empty
int64_t bucket_id, // map index (unused for now)
int64_t total_work_items // scalar (unused for now)
) {
check_int32_tensor(problem_sizes, "problem_sizes");
check_cuda_tensor(ptr_A, "ptr_A");
check_cuda_tensor(ptr_B, "ptr_B");
check_cuda_tensor(ptr_SFA, "ptr_SFA");
check_cuda_tensor(ptr_SFB, "ptr_SFB");
check_cuda_tensor(ptr_C, "ptr_C");
check_cuda_tensor(ptr_D, "ptr_D");
check_int64_tensor_optional(synergy_out, "synergy_out");
check_float_tensor_optional(potential_map, "potential_map");
int64_t g = problem_sizes.size(0);
if (g <= 0) return;
if (synergy_out.defined() && synergy_out.numel() > 0 && synergy_out.numel() < 5) {
throw std::runtime_error("synergy_out must have at least 5 elements");
}
if (synergy_out.defined() && synergy_out.numel() > 0) {
if (!potential_map.defined() || potential_map.numel() == 0) {
throw std::runtime_error("potential_map must be provided when synergy_out is set");
}
}
int tmode = timing_mode();
int tdetail = (tmode > 0) ? timing_detail() : 0;
auto t0 = std::chrono::high_resolution_clock::now();
auto t_ps0 = t0;
auto t_ps1 = t0;
auto t_cache0 = t0;
auto t_cache1 = t0;
auto t_build0 = t0;
auto t_build1 = t0;
auto t_upload0 = t0;
auto t_upload1 = t0;
auto t_can0 = t0;
auto t_can1 = t0;
auto t_ws0 = t0;
auto t_ws1 = t0;
auto t_sync0 = t0;
auto t_sync1 = t0;
// problem_sizes: int32 [g,3] (CPU or CUDA)
bool ps_on_cuda = problem_sizes.is_cuda();
std::vector<int32_t> ps_host;
auto copy_ps_host = [&]() {
if (!ps_host.empty()) return;
ps_host.resize(size_t(g) * 3);
if (ps_on_cuda) {
cudaMemcpy(ps_host.data(),
problem_sizes.data_ptr<int32_t>(),
sizeof(int32_t) * size_t(g) * 3,
cudaMemcpyDeviceToHost);
} else {
std::memcpy(ps_host.data(),
problem_sizes.data_ptr<int32_t>(),
sizeof(int32_t) * size_t(g) * 3);
}
};
if (tmode > 0) t_ps0 = std::chrono::high_resolution_clock::now();
if (!ps_on_cuda) {
copy_ps_host();
}
if (tmode > 0) t_ps1 = std::chrono::high_resolution_clock::now();
// Build per-group strides/layouts/problem-shapes into a single aligned host buffer
auto bytes = torch::TensorOptions().dtype(torch::kUInt8).device(torch::kCUDA);
constexpr size_t kAlign = 128;
auto align_up = [](size_t n, size_t a) { return (n + a - 1) & ~(a - 1); };
auto max_align = [](size_t a, size_t b) { return a > b ? a : b; };
size_t sz_sA = size_t(g) * sizeof(InternalStrideA);
size_t sz_sB = size_t(g) * sizeof(InternalStrideB);
size_t sz_sC = size_t(g) * sizeof(StrideC);
size_t sz_sD = size_t(g) * sizeof(StrideD);
size_t sz_lSFA = size_t(g) * sizeof(InternalLayoutSFA);
size_t sz_lSFB = size_t(g) * sizeof(InternalLayoutSFB);
size_t sz_ps = size_t(g) * sizeof(UnderlyingProblemShape);
size_t base_align = kAlign;
base_align = max_align(base_align, alignof(InternalStrideA));
base_align = max_align(base_align, alignof(InternalStrideB));
base_align = max_align(base_align, alignof(StrideC));
base_align = max_align(base_align, alignof(StrideD));
base_align = max_align(base_align, alignof(InternalLayoutSFA));
base_align = max_align(base_align, alignof(InternalLayoutSFB));
base_align = max_align(base_align, alignof(UnderlyingProblemShape));
int device = -1;
cudaGetDevice(&device);
cutlass::KernelHardwareInfo hw_info;
{
std::lock_guard<std::mutex> lock(g_hw_mutex);
auto it = g_hw_cache.find(device);
if (it != g_hw_cache.end()) {
hw_info = it->second;
} else {
cudaDeviceProp prop{};
cudaGetDeviceProperties(&prop, device);
hw_info.device_id = device;
hw_info.sm_count = prop.multiProcessorCount;
g_hw_cache.emplace(device, hw_info);
}
}
typename GemmKernel::TileScheduler::Arguments scheduler_args{};
if (tmode > 0) t_cache0 = std::chrono::high_resolution_clock::now();
CacheEntry* entry = nullptr;
std::string key;
bool last_hit = false;
int64_t ps_ptr_val = ps_on_cuda ? (int64_t)problem_sizes.data_ptr<int32_t>() : 0;
{
std::lock_guard<std::mutex> lock(g_last_mutex);
if (g_last_device == device && g_last_g == g) {
if (ps_on_cuda) {
if (g_last_ps_ptr == ps_ptr_val) {
key = g_last_key;
last_hit = true;
}
} else {
if (g_last_ps_host == ps_host) {
key = g_last_key;
last_hit = true;
}
}
}
}
if (last_hit) {
std::lock_guard<std::mutex> lock(g_cache_mutex);
auto it = g_cache.find(key);
if (it != g_cache.end()) {
entry = &it->second;
} else {
last_hit = false;
}
}
if (!last_hit) {
if (ps_on_cuda) {
key = make_cache_key_ptr(device, g, ps_ptr_val, g * 3);
} else {
key = make_cache_key(device, g, ps_host);
}
std::lock_guard<std::mutex> lock(g_cache_mutex);
auto it = g_cache.find(key);
if (it != g_cache.end()) {
entry = &it->second;
}
}
if (tmode > 0) t_cache1 = std::chrono::high_resolution_clock::now();
std::vector<uint8_t> host_buffer;
uint8_t* h_ptr = nullptr;
UnderlyingProblemShape* h_ps = nullptr;
if (entry == nullptr) {
if (ps_on_cuda && ps_host.empty()) {
if (tmode > 0) t_ps0 = std::chrono::high_resolution_clock::now();
copy_ps_host();
if (tmode > 0) t_ps1 = std::chrono::high_resolution_clock::now();
}
if (tmode > 0) t_build0 = std::chrono::high_resolution_clock::now();
size_t off = 0;
auto alloc = [&](size_t sz, size_t align) {
off = align_up(off, align);
size_t cur = off;
off += sz;
return cur;
};
size_t off_sA = alloc(sz_sA, max_align(kAlign, alignof(InternalStrideA)));
size_t off_sB = alloc(sz_sB, max_align(kAlign, alignof(InternalStrideB)));
size_t off_sC = alloc(sz_sC, max_align(kAlign, alignof(StrideC)));
size_t off_sD = alloc(sz_sD, max_align(kAlign, alignof(StrideD)));
size_t off_lSFA = alloc(sz_lSFA, max_align(kAlign, alignof(InternalLayoutSFA)));
size_t off_lSFB = alloc(sz_lSFB, max_align(kAlign, alignof(InternalLayoutSFB)));
size_t off_ps = alloc(sz_ps, max_align(kAlign, alignof(UnderlyingProblemShape)));
size_t total_bytes = align_up(off, kAlign);
host_buffer.resize(total_bytes + base_align);
uintptr_t base = reinterpret_cast<uintptr_t>(host_buffer.data());
uintptr_t aligned_base = (base + (base_align - 1)) & ~(uintptr_t(base_align - 1));
h_ptr = reinterpret_cast<uint8_t*>(aligned_base);
auto* h_sA = reinterpret_cast<InternalStrideA*>(h_ptr + off_sA);
auto* h_sB = reinterpret_cast<InternalStrideB*>(h_ptr + off_sB);
auto* h_sC = reinterpret_cast<StrideC*>(h_ptr + off_sC);
auto* h_sD = reinterpret_cast<StrideD*>(h_ptr + off_sD);
auto* h_lSFA = reinterpret_cast<InternalLayoutSFA*>(h_ptr + off_lSFA);
auto* h_lSFB = reinterpret_cast<InternalLayoutSFB*>(h_ptr + off_lSFB);
h_ps = reinterpret_cast<UnderlyingProblemShape*>(h_ptr + off_ps);
for (int64_t i = 0; i < g; ++i) {
int M = ps_host[size_t(i)*3 + 0];
int N = ps_host[size_t(i)*3 + 1];
int K = ps_host[size_t(i)*3 + 2];
// A/B are stored packed as [M, K/2, 1] and [N, K/2, 1] for NVFP4.
// Use K-major packed strides to match reference layout: (m,k,1) stride (k,1,m*k).
h_sA[i] = cutlass::make_cute_packed_stride(InternalStrideA{}, cute::make_shape(M, K, 1));
h_sB[i] = cutlass::make_cute_packed_stride(InternalStrideB{}, cute::make_shape(N, K, 1));
h_sC[i] = cutlass::make_cute_packed_stride(StrideC{}, cute::make_shape(M, N, 1));
h_sD[i] = cutlass::make_cute_packed_stride(StrideD{}, cute::make_shape(M, N, 1));
h_lSFA[i] = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(cute::make_shape(M, N, K, 1));
h_lSFB[i] = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(cute::make_shape(M, N, K, 1));
h_ps[i] = UnderlyingProblemShape{M, N, K};
}
if (tmode > 0) t_build1 = std::chrono::high_resolution_clock::now();
CacheEntry new_entry;
new_entry.total_bytes = total_bytes;
new_entry.off_sA = off_sA;
new_entry.off_sB = off_sB;
new_entry.off_sC = off_sC;
new_entry.off_sD = off_sD;
new_entry.off_lSFA = off_lSFA;
new_entry.off_lSFB = off_lSFB;
new_entry.off_ps = off_ps;
new_entry.device = device;
new_entry.g = g;
new_entry.ps_host_shape.resize(size_t(g));
for (int64_t i = 0; i < g; ++i) {
int M = ps_host[size_t(i)*3 + 0];
int N = ps_host[size_t(i)*3 + 1];
int K = ps_host[size_t(i)*3 + 2];
new_entry.ps_host_shape[size_t(i)] = UnderlyingProblemShape{M, N, K};
}
// One device buffer for strides/layouts/problem shapes
new_entry.dev_buffer = torch::empty({(long long)total_bytes}, bytes);
if (tmode > 0) t_upload0 = std::chrono::high_resolution_clock::now();
cudaMemcpy(new_entry.dev_buffer.data_ptr(), h_ptr, total_bytes, cudaMemcpyHostToDevice);
if (tmode > 0) t_upload1 = std::chrono::high_resolution_clock::now();
{
std::lock_guard<std::mutex> lock(g_cache_mutex);
auto it = g_cache.emplace(key, std::move(new_entry)).first;
entry = &it->second;
}
{
std::lock_guard<std::mutex> lock(g_last_mutex);
g_last_device = device;
g_last_g = g;
if (ps_on_cuda) {
g_last_ps_ptr = ps_ptr_val;
g_last_ps_host.clear();
} else {
g_last_ps_ptr = 0;
g_last_ps_host = ps_host;
}
g_last_key = key;
}
}
auto d_ptr = reinterpret_cast<uint8_t*>(entry->dev_buffer.data_ptr<uint8_t>());
auto ps_ptr = reinterpret_cast<UnderlyingProblemShape*>(d_ptr + entry->off_ps);
auto ps_host_ptr = ps_on_cuda ? nullptr : entry->ps_host_shape.data();
auto pA = reinterpret_cast<typename Gemm::ElementA const**>(ptr_A.data_ptr<int64_t>());
auto pB = reinterpret_cast<typename Gemm::ElementB const**>(ptr_B.data_ptr<int64_t>());
auto pSFA = reinterpret_cast<ElementSF const**>(ptr_SFA.data_ptr<int64_t>());
auto pSFB = reinterpret_cast<ElementSF const**>(ptr_SFB.data_ptr<int64_t>());
auto pC = reinterpret_cast<typename Gemm::ElementC const**>(ptr_C.data_ptr<int64_t>());
auto pD = reinterpret_cast<typename Gemm::EpilogueOutputOp::ElementOutput**>(ptr_D.data_ptr<int64_t>());
auto sA = reinterpret_cast<InternalStrideA*>(d_ptr + entry->off_sA);
auto sB = reinterpret_cast<InternalStrideB*>(d_ptr + entry->off_sB);
auto sC = reinterpret_cast<StrideC*>(d_ptr + entry->off_sC);
auto sD = reinterpret_cast<StrideD*>(d_ptr + entry->off_sD);
auto lSFA = reinterpret_cast<InternalLayoutSFA*>(d_ptr + entry->off_lSFA);
auto lSFB = reinterpret_cast<InternalLayoutSFB*>(d_ptr + entry->off_lSFB);
typename Gemm::Arguments args{
cutlass::gemm::GemmUniversalMode::kGrouped,
{ (int)g, ps_ptr, ps_host_ptr },
{ pA, StrideA(sA), pB, StrideB(sB), pSFA, LayoutSFA(lSFA), pSFB, LayoutSFB(lSFB) },
{ {}, pC, sC, pD, sD },
hw_info,
scheduler_args
};
args.epilogue.thread.alpha = 1.0f;
args.epilogue.thread.beta = 0.0f;
Gemm gemm;
std::vector<InternalStrideA> h_sA_dbg;
std::vector<InternalStrideB> h_sB_dbg;
std::vector<StrideC> h_sC_dbg;
std::vector<StrideD> h_sD_dbg;
std::vector<InternalLayoutSFA> h_lSFA_dbg;
std::vector<InternalLayoutSFB> h_lSFB_dbg;
bool have_debug = false;
auto build_debug = [&]() {
if (have_debug) return;
have_debug = true;
if (ps_host.empty()) {
copy_ps_host();
}
if (ps_host.empty()) {
return;
}
h_sA_dbg.resize(size_t(g));
h_sB_dbg.resize(size_t(g));
h_sC_dbg.resize(size_t(g));
h_sD_dbg.resize(size_t(g));
h_lSFA_dbg.resize(size_t(g));
h_lSFB_dbg.resize(size_t(g));
for (int64_t i = 0; i < g; ++i) {
int M = ps_host[size_t(i)*3 + 0];
int N = ps_host[size_t(i)*3 + 1];
int K = ps_host[size_t(i)*3 + 2];
h_sA_dbg[size_t(i)] = cutlass::make_cute_packed_stride(InternalStrideA{}, cute::make_shape(M, K, 1));
h_sB_dbg[size_t(i)] = cutlass::make_cute_packed_stride(InternalStrideB{}, cute::make_shape(N, K, 1));
h_sC_dbg[size_t(i)] = cutlass::make_cute_packed_stride(StrideC{}, cute::make_shape(M, N, 1));
h_sD_dbg[size_t(i)] = cutlass::make_cute_packed_stride(StrideD{}, cute::make_shape(M, N, 1));
h_lSFA_dbg[size_t(i)] = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(cute::make_shape(M, N, K, 1));
h_lSFB_dbg[size_t(i)] = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(cute::make_shape(M, N, K, 1));
}
};
if (tmode > 0) t_can0 = std::chrono::high_resolution_clock::now();
cutlass::Status st = gemm.can_implement(args);
cutlass::Status st_nohost = st;
cutlass::Status st_host = st;
if (ps_on_cuda) {
auto args_nohost = args;
args_nohost.problem_shape = { (int)g, ps_ptr, nullptr };
st_nohost = gemm.can_implement(args_nohost);
if (st_nohost == cutlass::Status::kSuccess) {
args = args_nohost;
st = st_nohost;
} else if (!entry->ps_host_shape.empty()) {
auto args_host = args;
args_host.problem_shape = { (int)g, ps_ptr, entry->ps_host_shape.data() };
st_host = gemm.can_implement(args_host);
if (st_host == cutlass::Status::kSuccess) {
args = args_host;
st = st_host;
}
}
}
if (tmode > 0) t_can1 = std::chrono::high_resolution_clock::now();
if (st != cutlass::Status::kSuccess) {
// Probe: can_implement with host problem shape pointer cleared (matches reference host_problem_shapes_available=false).
auto args_nohost = args;
args_nohost.problem_shape = { (int)g, ps_ptr, nullptr };
st_nohost = gemm.can_implement(args_nohost);
std::string msg = "can_implement failed: " + status_to_string(st);
msg += " (" + std::string(cutlassGetStatusString(st)) + ")";
// Always include minimal shape/type info to aid debugging without
msg += " | g=" + std::to_string((long long)g);
msg += " SFVecSize=" + std::to_string(SFVecSize);
msg += " ElementSF_size=" + std::to_string(sizeof(ElementSF));
msg += " ElementA_size=" + std::to_string(sizeof(ElementA));
msg += " ElementB_size=" + std::to_string(sizeof(ElementB));
msg += " ElementC_size=" + std::to_string(sizeof(ElementC));
if (g > 0) {
if (ps_host.empty()) {
copy_ps_host();
}
if (!ps_host.empty()) {
int M = ps_host[0];
int N = ps_host[1];
int K = ps_host[2];
msg += " | group0 M=" + std::to_string(M)
+ " N=" + std::to_string(N)
+ " K=" + std::to_string(K);
msg += " K_div2=" + std::to_string(K / 2);
} else {
msg += " | group0_unavailable";
}
}
build_debug();
// Per-group isolation probe to find the first failing group (attach to thrown msg).
std::vector<int64_t> ptrA_host;
std::vector<int64_t> ptrB_host;
std::vector<int64_t> ptrSFA_host;
std::vector<int64_t> ptrSFB_host;
std::vector<int64_t> ptrC_host;
std::vector<int64_t> ptrD_host;
ptrA_host.resize(size_t(g));
ptrB_host.resize(size_t(g));
ptrSFA_host.resize(size_t(g));
ptrSFB_host.resize(size_t(g));
ptrC_host.resize(size_t(g));
ptrD_host.resize(size_t(g));
cudaMemcpy(ptrA_host.data(), ptr_A.data_ptr<int64_t>(), sizeof(int64_t) * size_t(g), cudaMemcpyDeviceToHost);
cudaMemcpy(ptrB_host.data(), ptr_B.data_ptr<int64_t>(), sizeof(int64_t) * size_t(g), cudaMemcpyDeviceToHost);
cudaMemcpy(ptrSFA_host.data(), ptr_SFA.data_ptr<int64_t>(), sizeof(int64_t) * size_t(g), cudaMemcpyDeviceToHost);
cudaMemcpy(ptrSFB_host.data(), ptr_SFB.data_ptr<int64_t>(), sizeof(int64_t) * size_t(g), cudaMemcpyDeviceToHost);
cudaMemcpy(ptrC_host.data(), ptr_C.data_ptr<int64_t>(), sizeof(int64_t) * size_t(g), cudaMemcpyDeviceToHost);
cudaMemcpy(ptrD_host.data(), ptr_D.data_ptr<int64_t>(), sizeof(int64_t) * size_t(g), cudaMemcpyDeviceToHost);
if (ps_host.empty()) {
copy_ps_host();
}
if (ps_host.empty()) {
msg += " | ps_host_unavailable";
throw std::runtime_error(msg);
}
bool found = false;
for (int64_t i = 0; i < g; ++i) {
auto ps_ptr_i = ps_ptr + i;
auto ps_host_ptr_i = ps_host_ptr ? (ps_host_ptr + i) : nullptr;
auto pA_i = pA + i;
auto pB_i = pB + i;
auto pSFA_i = pSFA + i;
auto pSFB_i = pSFB + i;
auto pC_i = pC + i;
auto pD_i = pD + i;
auto sA_i = sA + i;
auto sB_i = sB + i;
auto sC_i = sC + i;
auto sD_i = sD + i;
auto lSFA_i = lSFA + i;
auto lSFB_i = lSFB + i;
typename Gemm::Arguments args1{
cutlass::gemm::GemmUniversalMode::kGrouped,
{ 1, ps_ptr_i, ps_host_ptr_i },
{ pA_i, StrideA(sA_i), pB_i, StrideB(sB_i), pSFA_i, LayoutSFA(lSFA_i), pSFB_i, LayoutSFB(lSFB_i) },
{ {}, pC_i, sC_i, pD_i, sD_i },
hw_info,
scheduler_args
};
args1.epilogue.thread.alpha = 1.0f;
args1.epilogue.thread.beta = 0.0f;
auto st1 = gemm.can_implement(args1);
if (st1 != cutlass::Status::kSuccess) {
int M = ps_host[size_t(i)*3 + 0];
int N = ps_host[size_t(i)*3 + 1];
int K = ps_host[size_t(i)*3 + 2];
msg += " | first_bad_group=" + std::to_string((long long)i);
msg += " M=" + std::to_string(M);
msg += " N=" + std::to_string(N);
msg += " K=" + std::to_string(K);
msg += " Mmod128=" + std::to_string(M % 128);
msg += " Nmod128=" + std::to_string(N % 128);
msg += " Kmod128=" + std::to_string(K % 128);
int64_t reqA = int64_t(AlignmentA) * int64_t(sizeof(ElementA));
int64_t reqB = int64_t(AlignmentB) * int64_t(sizeof(ElementB));
int64_t reqC = int64_t(AlignmentC) * int64_t(sizeof(ElementC));
int64_t reqD = int64_t(AlignmentD) * int64_t(sizeof(ElementD));
msg += " A_ptr_mod16=" + std::to_string((long long)(ptrA_host[size_t(i)] % 16));
msg += " B_ptr_mod16=" + std::to_string((long long)(ptrB_host[size_t(i)] % 16));
msg += " SFA_ptr_mod16=" + std::to_string((long long)(ptrSFA_host[size_t(i)] % 16));
msg += " SFB_ptr_mod16=" + std::to_string((long long)(ptrSFB_host[size_t(i)] % 16));
msg += " | AlignA_bytes=" + std::to_string((long long)reqA)
+ " A_ptr_mod_req=" + std::to_string((long long)(ptrA_host[size_t(i)] % reqA));
msg += " AlignB_bytes=" + std::to_string((long long)reqB)
+ " B_ptr_mod_req=" + std::to_string((long long)(ptrB_host[size_t(i)] % reqB));
msg += " AlignC_bytes=" + std::to_string((long long)reqC)
+ " C_ptr_mod_req=" + std::to_string((long long)(ptrC_host[size_t(i)] % reqC));
msg += " AlignD_bytes=" + std::to_string((long long)reqD)
+ " D_ptr_mod_req=" + std::to_string((long long)(ptrD_host[size_t(i)] % reqD));
msg += " | StrideA=("
+ std::to_string((long long)cute::get<0>(h_sA_dbg[size_t(i)])) + ","
+ std::to_string((long long)cute::get<1>(h_sA_dbg[size_t(i)])) + ","
+ std::to_string((long long)cute::get<2>(h_sA_dbg[size_t(i)])) + ")";
msg += " StrideB=("
+ std::to_string((long long)cute::get<0>(h_sB_dbg[size_t(i)])) + ","
+ std::to_string((long long)cute::get<1>(h_sB_dbg[size_t(i)])) + ","
+ std::to_string((long long)cute::get<2>(h_sB_dbg[size_t(i)])) + ")";
msg += " StrideC=("
+ std::to_string((long long)cute::get<0>(h_sC_dbg[size_t(i)])) + ","
+ std::to_string((long long)cute::get<1>(h_sC_dbg[size_t(i)])) + ","
+ std::to_string((long long)cute::get<2>(h_sC_dbg[size_t(i)])) + ")";
msg += " StrideD=("
+ std::to_string((long long)cute::get<0>(h_sD_dbg[size_t(i)])) + ","
+ std::to_string((long long)cute::get<1>(h_sD_dbg[size_t(i)])) + ","
+ std::to_string((long long)cute::get<2>(h_sD_dbg[size_t(i)])) + ")";
msg += " | ExpStride2(A,B,C,D)=(0,0,0,0)";
auto sfa_layout_i = h_lSFA_dbg[size_t(i)];
auto sfb_layout_i = h_lSFB_dbg[size_t(i)];
msg += " | SFA_size=" + std::to_string((long long)cute::size(sfa_layout_i));
msg += " SFA_compact=" + std::to_string((long long)cute::size(cute::filter_zeros(sfa_layout_i)));
msg += " SFB_size=" + std::to_string((long long)cute::size(sfb_layout_i));
msg += " SFB_compact=" + std::to_string((long long)cute::size(cute::filter_zeros(sfb_layout_i)));
msg += " | TileMNK=("
+ std::to_string((long long)cute::size<0>(MmaTileShape{})) + ","
+ std::to_string((long long)cute::size<1>(MmaTileShape{})) + ","
+ std::to_string((long long)cute::size<2>(MmaTileShape{})) + ")";
msg += " Cluster=("
+ std::to_string((long long)cute::size<0>(ClusterShape{})) + ","
+ std::to_string((long long)cute::size<1>(ClusterShape{})) + ","
+ std::to_string((long long)cute::size<2>(ClusterShape{})) + ")";
found = true;
break;
}
}
if (!found) {
msg += " | per_group_probe=all_pass";
}
msg += " | can_impl_nohost=" + status_to_string(st_nohost);
throw std::runtime_error(msg);
}
if (tmode > 0) t_ws0 = std::chrono::high_resolution_clock::now();
size_t workspace_size = Gemm::get_workspace_size(args);
if (!entry->workspace.defined() || entry->workspace.numel() < (long long)workspace_size) {
entry->workspace = torch::empty({(long long)workspace_size}, bytes);
}
if (tmode > 0) t_ws1 = std::chrono::high_resolution_clock::now();
bool do_synergy = synergy_out.defined() && synergy_out.numel() >= 5 &&
potential_map.defined() && potential_map.numel() > 0;
int64_t map_size = do_synergy ? potential_map.numel() : 0;
float kernel_ms = 0.0f;
int epi_threads = 32;
const char* epi_env = std::getenv("NVFP4_EPILOGUE_THREADS");
if (epi_env) {
int v = std::atoi(epi_env);
if (v > 0) epi_threads = v;
}
int epi_burn = 200;
const char* epi_burn_env = std::getenv("NVFP4_EPILOGUE_BURN");
if (epi_burn_env) {
int v = std::atoi(epi_burn_env);
if (v > 0) epi_burn = v;
}
if (tmode > 0) {
cudaEvent_t ev_start, ev_stop;
cudaEventCreate(&ev_start);
cudaEventCreate(&ev_stop);
if (do_synergy) {
synergy_prologue_kernel<<<1, 1, 0, 0>>>(
synergy_out.data_ptr<int64_t>());
}
cudaEventRecord(ev_start, 0);
st = gemm(args, entry->workspace.data_ptr(), 0);
cudaEventRecord(ev_stop, 0);
cudaEventSynchronize(ev_stop);
cudaEventElapsedTime(&kernel_ms, ev_start, ev_stop);
cudaEventDestroy(ev_start);
cudaEventDestroy(ev_stop);
if (do_synergy) {
synergy_epilogue_kernel<<<1, epi_threads, 0, 0>>>(
synergy_out.data_ptr<int64_t>(),
potential_map.data_ptr<float>(),
map_size,
bucket_id,
total_work_items,
epi_burn);
}
} else {
if (do_synergy) {
synergy_prologue_kernel<<<1, 1, 0, 0>>>(
synergy_out.data_ptr<int64_t>());
}
st = gemm(args, entry->workspace.data_ptr(), 0);
if (do_synergy) {
synergy_epilogue_kernel<<<1, epi_threads, 0, 0>>>(
synergy_out.data_ptr<int64_t>(),
potential_map.data_ptr<float>(),
map_size,
bucket_id,
total_work_items,
epi_burn);
}
}
if (st != cutlass::Status::kSuccess) {
throw std::runtime_error("GEMM launch failed: " + status_to_string(st));
}
if (tmode > 0) {
t_sync0 = std::chrono::high_resolution_clock::now();
cudaDeviceSynchronize();
t_sync1 = std::chrono::high_resolution_clock::now();
auto t1 = std::chrono::high_resolution_clock::now();
double total_ms = std::chrono::duration<double, std::milli>(t1 - t0).count();
double setup_ms = total_ms - double(kernel_ms);
std::fprintf(stderr,
"[nvfp4] g=%lld setup_ms=%.3f kernel_ms=%.3f total_ms=%.3f workspace_bytes=%lld\n",
(long long)g, setup_ms, kernel_ms, total_ms, (long long)workspace_size);
if (tdetail > 0) {
auto ms = [](auto a, auto b) {
return std::chrono::duration<double, std::milli>(b - a).count();
};
double ps_ms = ms(t_ps0, t_ps1);
double cache_ms = ms(t_cache0, t_cache1);
double build_ms = ms(t_build0, t_build1);
double upload_ms = ms(t_upload0, t_upload1);
double can_ms = ms(t_can0, t_can1);
double ws_ms = ms(t_ws0, t_ws1);
double sync_ms = ms(t_sync0, t_sync1);
std::fprintf(stderr,
"[nvfp4_detail] ps_ms=%.3f cache_ms=%.3f build_ms=%.3f upload_ms=%.3f can_ms=%.3f ws_ms=%.3f sync_ms=%.3f\n",
ps_ms, cache_ms, build_ms, upload_ms, can_ms, ws_ms, sync_ms);
}
std::fflush(stderr);
if (timing_abort() > 0) {
throw std::runtime_error("NVFP4_TIMING_ABORT");
}
}
}
#else
void grouped_nvfp4_tma_gemm(
torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor,
torch::Tensor, torch::Tensor, int64_t, int64_t
) {
throw std::runtime_error("SM100 support not enabled");
}
#endif
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("grouped_nvfp4_tma_gemm", &grouped_nvfp4_tma_gemm, "SM100 grouped NVFP4 GEMM");
m.def("convert_sf_float_to_ue4m3", &convert_sf_float_to_ue4m3, "Convert float scale factors to ue4m3");
m.def("pad_a_fp4", &pad_a_fp4, "Pad A for FP4 with zero fill");
m.def("pad_c_half", &pad_c_half, "Pad C for fp16 with zero fill");
}
"""
burn_macro = int(os.environ.get("NVFP4_INKERNEL_BURN", "100000"))
build_name = f"ek_nvfp4_sm100_ex81_pattern_v3_burn{burn_macro}_probe28"
print(f"[nvfp4] build_name={build_name} NVFP4_INKERNEL_BURN={burn_macro}", file=os.sys.stderr, flush=True)
extra_cuda_cflags = [
"-O3",
"--use_fast_math",
"--std=c++17",
"--expt-relaxed-constexpr",
"--expt-extended-lambda",
"-lineinfo",
"-gencode=arch=compute_100a,code=sm_100a",
"-gencode=arch=compute_100a,code=compute_100a",
"-DCUTLASS_ARCH_MMA_SM100_SUPPORTED",
"-DCUTLASS_ARCH_MMA_SM100A_ENABLED",
"-UCUTLASS_ENABLE_SYNCLOG",
]
if burn_macro > 0:
extra_cuda_cflags.append(f"-DNVFP4_INKERNEL_BURN={burn_macro}")
extra_cuda_cflags.append(f"-DNVFP4_INKERNEL_BURN_MARKER={burn_macro}")
extra_cflags = ["-O3", "-std=c++17"]
custom_kernel._EXT = load_inline(
name=build_name,
cpp_sources="",
cuda_sources=cuda_src,
functions=None,
extra_cuda_cflags=extra_cuda_cflags,
extra_cflags=extra_cflags,
with_cuda=True,
verbose=False,
extra_include_paths=extra_includes,
)
ext = custom_kernel._EXT
# -----------------------------
# 2) Build pointer arrays (3D kernel; expand L into groups)
# -----------------------------
dev = abc_tensors[0][2].device
abc_padded = []
sfs_out = []
problem_sizes_out = []
c_pad_info = [] # (orig_c, c_pad, m, n, l_idx)
def _maybe_contig(t):
return t if t.is_contiguous() else t.contiguous()
graph_safe = True
for i in range(g):
a, b, c = abc_tensors[i]
sfa, sfb = sfasfb_tensors[i][0], sfasfb_tensors[i][1]
m, n, k, l = problem_sizes[i]
if l <= 0:
raise RuntimeError(f"Invalid L={l} for group {i}")
if l != 1:
raise RuntimeError(f"Only L=1 is supported in this fast path (got L={l} for group {i})")
if sfa.dtype == torch.float32:
graph_safe = False
sfa = ext.convert_sf_float_to_ue4m3(sfa)
elif sfa.dtype == torch.float8_e4m3fnuz:
pass
elif sfa.dtype == torch.float8_e4m3fn:
pass
elif sfa.dtype != torch.uint8:
raise TypeError("SFA tensors must be float32, float8_e4m3fnuz, float8_e4m3fn, or uint8 (ue4m3 storage).")
if sfb.dtype == torch.float32:
graph_safe = False
sfb = ext.convert_sf_float_to_ue4m3(sfb)
elif sfb.dtype == torch.float8_e4m3fnuz:
pass
elif sfb.dtype == torch.float8_e4m3fn:
pass
elif sfb.dtype != torch.uint8:
raise TypeError("SFB tensors must be float32, float8_e4m3fnuz, float8_e4m3fn, or uint8 (ue4m3 storage).")
if not sfa.is_cuda or not sfb.is_cuda:
raise TypeError("SFA/SFB tensors must be CUDA tensors.")
a_i = _maybe_contig(a)
b_i = _maybe_contig(b)
c_i = _maybe_contig(c)
if a_i.data_ptr() != a.data_ptr() or b_i.data_ptr() != b.data_ptr() or c_i.data_ptr() != c.data_ptr():
graph_safe = False
sfa_i = sfa
sfb_i = sfb
div_sfa = (k // 16)
div_sfb = (k // 16)
if div_sfa <= 0 or (sfa_i.numel() % div_sfa) != 0:
raise RuntimeError(f"SFA size mismatch for group {i}: sfa.numel={sfa_i.numel()}")
if div_sfb <= 0 or (sfb_i.numel() % div_sfb) != 0:
raise RuntimeError(f"SFB size mismatch for group {i}: sfb.numel={sfb_i.numel()}")
m_eff = sfa_i.numel() // div_sfa
n_eff = sfb_i.numel() // div_sfb
if m_eff < m:
raise RuntimeError(f"SFA implies smaller M for group {i}: m_eff={m_eff} m={m}")
if n_eff < n:
raise RuntimeError(f"SFB implies smaller N for group {i}: n_eff={n_eff} n={n}")
skip_pad = int(os.environ.get("NVFP4_SKIP_PAD", "1"))
if m_eff != m:
if skip_pad > 0:
m_eff = m
else:
graph_safe = False
a_i = ext.pad_a_fp4(a_i, m_eff)
if m_eff != m or n_eff != n:
if skip_pad > 0:
n_eff = n
else:
graph_safe = False
c_i = ext.pad_c_half(c_i, m_eff, n_eff)
c_pad_info.append((abc_tensors[i][2], c_i, m, n, None))
abc_padded.append((a_i, b_i, c_i))
sfs_out.append((sfa_i, sfb_i))
problem_sizes_out.append((int(m_eff), int(n_eff), int(k)))
if timing_py:
t_py1 = time.perf_counter()
def _build_cache_key_prepared(abc_tensors, sfasfb_tensors, problem_sizes):
abc_sig = tuple((_tensor_key(a), _tensor_key(b), _tensor_key(c)) for a, b, c in abc_tensors)
sfs_sig = tuple((_tensor_key(sfa), _tensor_key(sfb)) for sfa, sfb in sfasfb_tensors)
ps_sig = tuple((int(m), int(n), int(k)) for (m, n, k) in problem_sizes)
return (abc_sig, sfs_sig, ps_sig)
cache_key = _build_cache_key_prepared(abc_padded, sfs_out, problem_sizes_out) if use_py_cache else None
cached_state = custom_kernel._CACHE.get(cache_key) if use_py_cache else None
if use_py_cache and use_echokey and cached_state is not None and cached_state.get("echokey_pack") is None:
cached_state = None
A_ptrs = [int(abc_padded[i][0].data_ptr()) for i in range(len(problem_sizes_out))]
B_ptrs = [int(abc_padded[i][1].data_ptr()) for i in range(len(problem_sizes_out))]
C_ptrs = [int(abc_padded[i][2].data_ptr()) for i in range(len(problem_sizes_out))]
D_ptrs = C_ptrs
SFA_ptrs = [int(sfs_out[i][0].data_ptr()) for i in range(len(problem_sizes_out))]
SFB_ptrs = [int(sfs_out[i][1].data_ptr()) for i in range(len(problem_sizes_out))]
def _update_ptr_tensor(dst, src_list):
src = torch.as_tensor(src_list, device="cpu", dtype=torch.int64)
dst.copy_(src, non_blocking=True)
if cached_state is not None:
ps = cached_state["ps"]
ptr_A = cached_state["ptr_A"]
ptr_B = cached_state["ptr_B"]
ptr_SFA = cached_state["ptr_SFA"]
ptr_SFB = cached_state["ptr_SFB"]
ptr_C = cached_state["ptr_C"]
ptr_D = cached_state["ptr_D"]
prev_ptrs = cached_state.get("ptr_host", None)
cur_ptrs = (A_ptrs, B_ptrs, SFA_ptrs, SFB_ptrs, C_ptrs, D_ptrs)
if prev_ptrs != cur_ptrs:
_update_ptr_tensor(ptr_A, A_ptrs)
_update_ptr_tensor(ptr_B, B_ptrs)
_update_ptr_tensor(ptr_SFA, SFA_ptrs)
_update_ptr_tensor(ptr_SFB, SFB_ptrs)
_update_ptr_tensor(ptr_C, C_ptrs)
_update_ptr_tensor(ptr_D, D_ptrs)
cached_state["ptr_host"] = cur_ptrs
custom_kernel._CACHE.move_to_end(cache_key)
else:
ptr_A = torch.tensor(A_ptrs, device=dev, dtype=torch.int64)
ptr_B = torch.tensor(B_ptrs, device=dev, dtype=torch.int64)
ptr_SFA = torch.tensor(SFA_ptrs, device=dev, dtype=torch.int64)
ptr_SFB = torch.tensor(SFB_ptrs, device=dev, dtype=torch.int64)
ptr_C = torch.tensor(C_ptrs, device=dev, dtype=torch.int64)
ptr_D = torch.tensor(D_ptrs, device=dev, dtype=torch.int64)
ps = torch.tensor(
[(int(m), int(n), int(k)) for (m, n, k) in problem_sizes_out],
device=dev,
dtype=torch.int32,
)
if timing_py:
t_py2 = time.perf_counter()
echokey_pack = None
if use_echokey:
if cached_state is not None and cached_state.get("echokey_pack") is not None:
echokey_pack = cached_state["echokey_pack"]
else:
bucket_id = _bucket_id(problem_sizes, map_size)
total_work_items = 0
for (m, n, k) in problem_sizes_out:
total_work_items += int(m) * int(n) * int(k)
if dev not in custom_kernel._POTENTIAL_MAP:
custom_kernel._POTENTIAL_MAP[dev] = torch.zeros(map_size, dtype=torch.float32, device=dev)
if dev not in custom_kernel._SYNERGY_OUT:
custom_kernel._SYNERGY_OUT[dev] = torch.zeros(5, dtype=torch.int64, device=dev)
potential_map = custom_kernel._POTENTIAL_MAP[dev]
synergy_out = custom_kernel._SYNERGY_OUT[dev]
echokey_pack = {
"bucket_id": bucket_id,
"total_work_items": total_work_items,
"potential_map": potential_map,
"synergy_out": synergy_out,
"warmed": False,
}
if use_py_cache:
custom_kernel._CACHE[cache_key] = {
"ps": ps,
"ptr_A": ptr_A,
"ptr_B": ptr_B,
"ptr_SFA": ptr_SFA,
"ptr_SFB": ptr_SFB,
"ptr_C": ptr_C,
"ptr_D": ptr_D,
"ptr_host": (A_ptrs, B_ptrs, SFA_ptrs, SFB_ptrs, C_ptrs, D_ptrs),
"echokey_pack": echokey_pack,
}
custom_kernel._CACHE.move_to_end(cache_key)
max_cache = int(os.environ.get("NVFP4_PY_CACHE_SIZE", "8"))
while len(custom_kernel._CACHE) > max_cache:
custom_kernel._CACHE.popitem(last=False)
# -----------------------------
# 3) Run grouped GEMM
# -----------------------------
if timing_py:
torch.cuda.synchronize(dev)
t_call0 = time.perf_counter()
if use_echokey:
synergy_out = echokey_pack["synergy_out"]
potential_map = echokey_pack["potential_map"]
bucket_id = int(echokey_pack["bucket_id"])
total_work_items = int(echokey_pack["total_work_items"])
else:
synergy_out, potential_map = custom_kernel._ECHO_DUMMY
bucket_id = 0
total_work_items = 0
use_graph = int(os.environ.get("NVFP4_USE_GRAPH", "0")) > 0
use_manual_launch = True
if use_graph and use_py_cache and graph_safe:
if not hasattr(custom_kernel, "_GRAPH_CACHE"):
custom_kernel._GRAPH_CACHE = {}
graph_key = (cache_key, "graph", tuple(A_ptrs), tuple(B_ptrs), tuple(C_ptrs), tuple(SFA_ptrs), tuple(SFB_ptrs))
g_entry = custom_kernel._GRAPH_CACHE.get(graph_key)
if g_entry is None:
g_entry = {"graph": None, "armed": False, "disabled": False}
custom_kernel._GRAPH_CACHE[graph_key] = g_entry
if g_entry["disabled"]:
pass
elif g_entry["graph"] is None:
if not g_entry["armed"]:
# Arm on first encounter; capture on next call to avoid skewing first timing.
g_entry["armed"] = True
else:
# Warmup to ensure allocations are done
ext.grouped_nvfp4_tma_gemm(
ps, ptr_A, ptr_B, ptr_SFA, ptr_SFB, ptr_C, ptr_D,
synergy_out, potential_map, bucket_id, total_work_items
)
torch.cuda.synchronize(dev)
graph = torch.cuda.CUDAGraph()
torch.cuda.set_device(dev)
import warnings
with warnings.catch_warnings(record=True) as w:
warnings.simplefilter("always")
with torch.cuda.graph(graph):
ext.grouped_nvfp4_tma_gemm(
ps, ptr_A, ptr_B, ptr_SFA, ptr_SFB, ptr_C, ptr_D,
synergy_out, potential_map, bucket_id, total_work_items
)
if w:
g_entry["disabled"] = True
else:
g_entry["graph"] = graph
graph.replay()
use_manual_launch = False
else:
g_entry["graph"].replay()
use_manual_launch = False
if use_manual_launch:
if use_echokey:
if use_warmup and not echokey_pack.get("warmed", False):
ext.grouped_nvfp4_tma_gemm(
ps, ptr_A, ptr_B, ptr_SFA, ptr_SFB, ptr_C, ptr_D,
synergy_out, potential_map, bucket_id, total_work_items
)
echokey_pack["warmed"] = True
ext.grouped_nvfp4_tma_gemm(
ps, ptr_A, ptr_B, ptr_SFA, ptr_SFB, ptr_C, ptr_D,
synergy_out, potential_map, bucket_id, total_work_items
)
else:
ext.grouped_nvfp4_tma_gemm(
ps, ptr_A, ptr_B, ptr_SFA, ptr_SFB, ptr_C, ptr_D,
synergy_out, potential_map, bucket_id, total_work_items
)
if use_echokey and echokey_debug:
torch.cuda.synchronize(dev)
dbg = echokey_pack["synergy_out"].detach().cpu().tolist()
print(f"[echokey] bucket={echokey_pack['bucket_id']} dt={dbg[2]}", file=os.sys.stderr, flush=True)
if timing_py:
torch.cuda.synchronize(dev)
t_call1 = time.perf_counter()
for orig_c, c_pad, m, n, l_idx in c_pad_info:
if l_idx is None:
orig_c.copy_(c_pad[:m, :n, :])
else:
orig_c[:, :, l_idx].copy_(c_pad[:m, :n, 0])
if timing_py:
t_py3 = time.perf_counter()
prep_ms = (t_py1 - t_py0) * 1e3
ptr_ms = (t_py2 - t_py1) * 1e3
call_ms = (t_call1 - t_call0) * 1e3
post_ms = (t_py3 - t_call1) * 1e3
total_ms = (t_py3 - t_py0) * 1e3
print(
f"[nvfp4_py] prep_ms={prep_ms:.3f} ptr_ms={ptr_ms:.3f} "
f"call_ms={call_ms:.3f} post_ms={post_ms:.3f} total_ms={total_ms:.3f}",
file=os.sys.stderr,
flush=True,
)
return [abc_tensors[i][2] for i in range(g)]
scrolls · 1645 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