submission 845171
drillyb · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1579 lines, June 9 Researcher Reciprocity License v1.0.
new_sub.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-845171?include=source"interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp32
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:510c799090f7aca89fed2d2197b883b2d0f5677547a0183d72580798a9acbb4c
license declaredunknown
license concludedunknown
authorsdrillyb
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mbarrier
asm volatile("cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.cta_group::%6 "mma
const float hi = nvcuda::wmma::__float_to_tf32(x);num-warps = 4
constexpr int NUM_WARPS = 4;shared-memory
static __shared__ float warp_sums[NUM_WARP];tcgen05
"tcgen05.mma.cta_group::%5.kind::tf32 [%0], %1 , %2 , %3, p;\n\t"tma
asm volatile("cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.cta_group::%6 "Kernel source
new_sub.py1579 lines
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline
import os
os.environ['CUDA_LAUNCH_BLOCKING'] = "1"
_FACTOR_RTOL_FACTOR = 20.0
_ORTH_RTOL_FACTOR = 100.0
def _apply_column_scaling(a: torch.Tensor, cond: int) -> torch.Tensor:
# `cond` is a deterministic dynamic-range knob, not an exact condition number.
if cond:
n = a.shape[-1]
scales = torch.logspace(0.0, -float(cond), n, device=a.device, dtype=torch.float32)
return a * scales
return a.contiguous()
def _band_mask(n: int, bandwidth: int, device: torch.device) -> torch.Tensor:
idx = torch.arange(n, device=device)
return (idx[:, None] - idx[None, :]).abs() <= bandwidth
def generate_input(batch: int, n: int, cond: int, seed: int, case: str = "dense") -> input_t:
assert batch > 0, "batch must be positive"
assert n > 0, "n must be positive"
assert cond >= 0, "cond must be non-negative"
device = "cuda" if torch.cuda.is_available() else "cpu"
gen = torch.Generator(device=device)
gen.manual_seed(seed)
case = case.lower()
a = torch.randn((batch, n, n), device=device, dtype=torch.float32, generator=gen)
if case == "dense":
a = _apply_column_scaling(a, cond)
elif case == "upper":
diag_boost = torch.linspace(1.0, 0.25, n, device=device, dtype=torch.float32)
a = torch.triu(a)
a.diagonal(dim1=-2, dim2=-1).add_(diag_boost)
a = _apply_column_scaling(a, cond)
elif case == "diagonal":
diag = torch.randn((batch, n), device=device, dtype=torch.float32, generator=gen)
diag = diag.sign().clamp(min=0.0).mul(2.0).sub(1.0) * torch.logspace(
0.0, -float(max(cond, 2)), n, device=device, dtype=torch.float32
)
a = torch.diag_embed(diag)
elif case == "rankdef":
rank = max(1, (3 * n) // 4)
a[:, :, rank:] = 0.0
a = _apply_column_scaling(a, cond)
elif case == "nearrank":
rank = max(1, (3 * n) // 4)
tail = n - rank
if tail > 0:
noise = torch.randn(
(batch, n, tail), device=device, dtype=torch.float32, generator=gen
)
a[:, :, rank:] = a[:, :, :tail] + 1.0e-5 * noise
a = _apply_column_scaling(a, cond)
elif case == "clustered":
scales = torch.ones((n,), device=device, dtype=torch.float32)
scales[n // 2 :] = 4.0 * torch.finfo(torch.float32).eps
if n >= 8:
lo = max(0, n // 2 - 2)
hi = min(n, n // 2 + 2)
scales[lo:hi] = torch.sqrt(torch.tensor(torch.finfo(torch.float32).eps, device=device))
a = a * scales
elif case == "band":
bandwidth = max(2, min(32, n // 32))
a = a * _band_mask(n, bandwidth, device)
diag_boost = torch.linspace(1.0, 0.5, n, device=device, dtype=torch.float32)
a.diagonal(dim1=-2, dim2=-1).add_(diag_boost)
a = _apply_column_scaling(a, cond)
elif case == "nearcollinear":
base = torch.randn((batch, n, 1), device=device, dtype=torch.float32, generator=gen)
noise = torch.randn((batch, n, n), device=device, dtype=torch.float32, generator=gen)
a = base.expand(batch, n, n) + 1.0e-4 * noise
a = _apply_column_scaling(a, cond)
elif case == "rowscale":
row_cond = max(cond, 4)
scales = torch.logspace(0.0, -float(row_cond), n, device=device, dtype=torch.float32)
a = scales.reshape(1, n, 1) * a
else:
raise ValueError(f"unknown QR test case: {case}")
return a.contiguous()
def ref_kernel(data: input_t) -> output_t:
# Starter/reference path: correctness first; submissions compete on speed.
return torch.geqrf(data)
CUDA_SOURCE = r"""
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <mma.h>
#include <algorithm>
#include <cstdio>
#include <vector>
#include <tuple>
namespace {
constexpr int cdiv(int x, int y) {
return (x + y - 1) / y;
}
constexpr int kPanelWidth = 32;
constexpr int kPanelThreads = 256;
constexpr int kPackVThreads = 256;
constexpr int kUpdateThreads = 256;
constexpr int kReflectorTileRows = 128;
constexpr int kPanelColumnTile = 8;
constexpr int kTileRows = 32;
constexpr int kTileCols = 32;
constexpr bool kUseReferencePath = false;
constexpr bool kEnableKernelTiming = false;
constexpr int NUM_WARPS = 4;
constexpr int WARP_SIZE = 32;
constexpr int TB_SIZE = NUM_WARPS * WARP_SIZE;
constexpr int MMA_K = 8;
constexpr int kGemmBlockM = 128;
constexpr int kGemmBlockN = 64;
constexpr int kGemmBlockK = 16;
constexpr int kYGemmBlockM = 128;
constexpr int kYGemmBlockN = 64;
constexpr int kYGemmBlockK = 16;
template<int WARP_SIZE = 32>
__device__ __forceinline__ float warp_reduce_sum(float value) {
#pragma unroll
for (int offset = WARP_SIZE >> 1; offset >= 1; offset >>= 1) {
value += __shfl_xor_sync(0xffffffffu, value, offset);
}
return value;
}
template<int BLOCK_SIZE = 256>
__device__ __forceinline__ float block_reduce_sum(float value) {
constexpr int WARP_SIZE = 32;
constexpr int NUM_WARP = (BLOCK_SIZE + WARP_SIZE - 1)/ WARP_SIZE;
const int lane_id = threadIdx.x % WARP_SIZE;
const int warp_id = threadIdx.x / WARP_SIZE;
static __shared__ float warp_sums[NUM_WARP];
__shared__ float block_sum;
const float warp_sum = warp_reduce_sum<WARP_SIZE>(value);
if (lane_id == 0) {
warp_sums[warp_id] = warp_sum;
}
__syncthreads();
if ( warp_id ==0) {
float out = (lane_id < NUM_WARP) ? warp_sums[lane_id] : 0.0f;
const float total_sum = warp_reduce_sum<WARP_SIZE>(out);
if(lane_id == 0) {
block_sum = total_sum;
}
}
__syncthreads();
return block_sum;
}
__device__ __forceinline__ float block_broadcast(float value, int src_thread = 0) {
__shared__ float shared_value;
if (threadIdx.x == src_thread) {
shared_value = value;
}
__syncthreads();
return shared_value;
}
template <int CTA_GROUP = 1>
__device__ inline
void tma_3d_gmem2smem(int dst, const void *tmap_ptr, int x, int y, int z, int mbar_addr) {
asm volatile("cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.cta_group::%6 "
"[%0], [%1, {%2, %3, %4}], [%5];"
:: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "n"(CTA_GROUP)
: "memory");
}
__device__ __forceinline__ uint32_t elect_thr(){
uint32_t pred = 0;
asm volatile(
"{\n\t"
".reg .pred %%px;\n\t"
"elect.sync _|%%px, %1;\n\t"
"@%%px mov.s32 %0, 1;\n\t"
"}"
: "+r"(pred)
: "r"(0xFFFFFFFFu)
);
return pred;
}
template<int CTA_NUM = 1>
__device__ __inline__ void tcgen05mma_tf32(int addr, uint64_t a_desc , uint64_t b_desc, uint32_t i_desc , int enable_input_d){
asm volatile (
"{\n\t"
".reg .pred p;\n\t"
"setp.ne.b32 p , %4,0; \n\t"
"tcgen05.mma.cta_group::%5.kind::tf32 [%0], %1 , %2 , %3, p;\n\t"
"}"
:: "r"(addr), "l"(a_desc), "l"(b_desc), "r"(i_desc), "r"(enable_input_d), "n"(CTA_NUM)
);
}
__device__ __inline__ constexpr uint64_t desc_encode(uint64_t x) {return (x & 0x3'FFFFULL) >> 4ULL ;};
__device__ __forceinline__ int64_t offset3(int batch_idx, int row, int col, int n) {
return (static_cast<int64_t>(batch_idx) * n + row) * n + col;
}
__device__ __forceinline__ int64_t tau_offset(int batch_idx, int col, int n) {
return static_cast<int64_t>(batch_idx) * n + col;
}
__device__ __forceinline__ int64_t offset_v(int batch_idx, int row, int col, int rows, int cols) {
return (static_cast<int64_t>(batch_idx) * rows + row) * cols + col;
}
__device__ __forceinline__ void tma_2d_gmem2smem(int dst , const void *tmp_ptr , int x , int y , int mbar_addr){
asm volatile ("cp.async.bulk.tensor.2d.shared::cta.global.mbarrier::complete_tx::bytes [%0], [%1 , {%2 , %3}], [%4];"
:: "r"(dst) , "l"(tmp_ptr), "r"(x) , "r"(y), "r"(mbar_addr) : "memory");
}
__device__ __inline__ void mbarrier_wait(int mbar_adder , int phase){
uint32_t ticks = 0x989680;
asm volatile (
"{\n\t"
".reg .pred P1;\n\t"
"LAB_WAIT:\n\t"
"mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P1 , [%0] , %1 , %2; \n\t"
"@P1 bra.uni DONE;\n\t"
"bra.uni LAB_WAIT;\n\t"
"DONE:\n\t"
"}"
:: "r"(mbar_adder), "r"(phase), "r"(ticks)
);
}
inline void check_cu(CUresult err) {
TORCH_CHECK(err == CUDA_SUCCESS, "cuTensorMapEncodeTiled failed with code ", static_cast<int>(err));
}
inline void init_tmap_2d_simple(
CUtensorMap *tmap,
const float *ptr,
uint64_t global_height, uint64_t global_width,
uint32_t shared_height, uint32_t shared_width,
CUtensorMapSwizzle swizzle
) {
constexpr uint32_t rank = 2;
uint64_t globalDim[rank] = {global_width, global_height};
uint64_t globalStrides[rank-1] = {global_width * sizeof(float)}; // in bytes
uint32_t boxDim[rank] = {shared_width, shared_height};
uint32_t elementStrides[rank] = {1, 1};
auto err = cuTensorMapEncodeTiled(
tmap,
CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_FLOAT32,
rank,
(void *)ptr,
globalDim,
globalStrides,
boxDim,
elementStrides,
CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE,
swizzle,
CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE,
CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE
);
check_cu(err);
}
inline void init_tmap_3d_batched(
CUtensorMap *tmap,
const float *ptr,
uint64_t batch,
uint64_t global_height, uint64_t global_width,
uint32_t shared_height, uint32_t shared_width,
CUtensorMapSwizzle swizzle
) {
constexpr uint32_t rank = 3;
// The source tensors are contiguous [batch, height, width].
// TMA dimensions are ordered from fastest to slowest: [width, height, batch].
uint64_t globalDim[rank] = {
global_width,
global_height,
batch
};
uint64_t globalStrides[rank - 1] = {
global_width * sizeof(float),
global_height * global_width * sizeof(float)
};
uint32_t boxDim[rank] = {
shared_width,
shared_height,
1
};
uint32_t elementStrides[rank] = {1, 1, 1};
auto err = cuTensorMapEncodeTiled(
tmap,
CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_FLOAT32,
rank,
(void *)ptr,
globalDim,
globalStrides,
boxDim,
elementStrides,
CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE,
swizzle,
CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE,
CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE
);
check_cu(err);
}
inline void init_tmap_3d_strided(
CUtensorMap *tmap,
const float *ptr,
uint64_t batch,
uint64_t global_height,
uint64_t global_width,
uint64_t row_stride_elements,
uint64_t batch_stride_elements,
uint32_t shared_height,
uint32_t shared_width,
CUtensorMapSwizzle swizzle
) {
constexpr uint32_t rank = 3;
uint64_t globalDim[rank] = {
global_width,
global_height,
batch
};
uint64_t globalStrides[rank - 1] = {
row_stride_elements * sizeof(float),
batch_stride_elements * sizeof(float)
};
uint32_t boxDim[rank] = {
shared_width,
shared_height,
1
};
uint32_t elementStrides[rank] = {1, 1, 1};
auto err = cuTensorMapEncodeTiled(
tmap,
CUtensorMapDataType::CU_TENSOR_MAP_DATA_TYPE_FLOAT32,
rank,
(void *)ptr,
globalDim,
globalStrides,
boxDim,
elementStrides,
CUtensorMapInterleave::CU_TENSOR_MAP_INTERLEAVE_NONE,
swizzle,
CUtensorMapL2promotion::CU_TENSOR_MAP_L2_PROMOTION_NONE,
CUtensorMapFloatOOBfill::CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE
);
check_cu(err);
}
template<bool FUSE_SUBTRACT, int blockM ,int blockN, int blockK>
__global__
__launch_bounds__(TB_SIZE)
void GEMM_kernel(
const __grid_constant__ CUtensorMap A_tmap,
const __grid_constant__ CUtensorMap B_tmap,
float * C_ptr,
int batch,
int M, int N, int K,
int64_t output_batch_stride,
int output_row_stride,
int64_t output_base_offset
){
const int tid = threadIdx.x;
const int warp_id = tid / WARP_SIZE;
const int grid_m = cdiv(M, blockM);
const int grid_n = cdiv(N, blockN);
const int tiles_per_batch = grid_m * grid_n;
const int linear_block = static_cast<int>(blockIdx.x);
const int batch_idx = linear_block / tiles_per_batch;
const int tile_idx = linear_block - batch_idx * tiles_per_batch;
const int bid_m = tile_idx / grid_n;
const int bid_n = tile_idx % grid_n;
if (batch_idx >= batch) {
return;
}
const int off_m = bid_m * blockM;
const int off_n = bid_n * blockN;
extern __shared__ __align__(1024) char smem[];
constexpr int A_elems = blockM * blockK;
constexpr int B_elems = blockN * blockK;
constexpr int A_bytes = A_elems * static_cast<int>(sizeof(float));
constexpr int B_bytes = B_elems * static_cast<int>(sizeof(float));
// TMA first loads the original FP32 values into the *_hi regions.
// After the transfer, all CTA threads split each value in-place:
// x_hi = round_to_tf32(x)
// x_lo = x - x_hi
// The three tcgen05 products are:
// A_hi * B_hi + A_hi * B_lo + A_lo * B_hi.
const int A_hi_smem = static_cast<int>(__cvta_generic_to_shared(smem));
const int B_hi_smem = A_hi_smem + A_bytes;
const int A_lo_smem = B_hi_smem + B_bytes;
const int B_lo_smem = A_lo_smem + A_bytes;
float* const A_hi_ptr = reinterpret_cast<float*>(smem);
float* const B_hi_ptr = A_hi_ptr + A_elems;
float* const A_lo_ptr = B_hi_ptr + B_elems;
float* const B_lo_ptr = A_lo_ptr + A_elems;
__shared__ uint64_t mbars[1];
__shared__ int tmem_adder[1];
const int mbar_addr = static_cast<int>(__cvta_generic_to_shared(mbars));
if(warp_id == 0 && elect_thr()){
asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(mbar_addr) , "r"(1));
asm volatile("fence.mbarrier_init.release.cluster;");
} else if (warp_id == 1){
const int addr = static_cast<int>(__cvta_generic_to_shared(tmem_adder));
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(addr) , "r"(blockN));
}
__syncthreads();
int phase = 0;
const int taddr = tmem_adder[0];
constexpr uint32_t i_desc = (1U << 4U)
| (2U << 7U)
| (2U << 10U)
| ((uint32_t)blockN >> 3U << 17U)
| ((uint32_t)blockM >> 4U << 24U)
;
const int iters = cdiv(K, blockK);
__syncthreads();
for(int i = 0 ; i < iters; i++ ){
if(warp_id == 0 && elect_thr()){
for(int strid_id = 0 ; strid_id < blockK/4; strid_id++){
const int off_set = i * blockK + strid_id * 4;
tma_3d_gmem2smem<1>(
A_hi_smem + strid_id * blockM * 16,
&A_tmap,
off_set,
off_m,
batch_idx,
mbar_addr);
tma_3d_gmem2smem<1>(
B_hi_smem + strid_id * blockN * 16,
&B_tmap,
off_set,
off_n,
batch_idx,
mbar_addr);
}
constexpr int cp_size = (blockM + blockN) * blockK * sizeof(float);
asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0] , %1;"
:: "r"(mbar_addr) , "r"(cp_size) : "memory");
}
mbarrier_wait(mbar_addr , phase);
phase ^= 1;
// Split the TMA-loaded FP32 tiles into TF32-high and FP32 residual
// components. The canonical tcgen05 shared-memory layout is already
// contiguous in these buffers, so an elementwise split preserves it.
for (int idx = tid; idx < A_elems; idx += TB_SIZE) {
const float x = A_hi_ptr[idx];
const float hi = nvcuda::wmma::__float_to_tf32(x);
A_hi_ptr[idx] = hi;
A_lo_ptr[idx] = x - hi;
}
for (int idx = tid; idx < B_elems; idx += TB_SIZE) {
const float x = B_hi_ptr[idx];
const float hi = nvcuda::wmma::__float_to_tf32(x);
B_hi_ptr[idx] = hi;
B_lo_ptr[idx] = x - hi;
}
__syncthreads();
// Publish the thread-written shared-memory operands to tcgen05.
asm volatile ("tcgen05.fence::after_thread_sync;");
if(warp_id == 0 && elect_thr()){
auto desc = [](int addr , int height) -> uint64_t {
const int LBO = height * 16;
const int SBO = 8 * 16;
return desc_encode(addr) | (desc_encode(LBO) << 16ULL) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
};
#pragma unroll
for(int id = 0 ; id < blockK / MMA_K ; id++){
const int a_offset =
id * blockM * MMA_K * static_cast<int>(sizeof(float));
const int b_offset =
id * blockN * MMA_K * static_cast<int>(sizeof(float));
const uint64_t a_hi_desc = desc(A_hi_smem + a_offset, blockM);
const uint64_t b_hi_desc = desc(B_hi_smem + b_offset, blockN);
const uint64_t a_lo_desc = desc(A_lo_smem + a_offset, blockM);
const uint64_t b_lo_desc = desc(B_lo_smem + b_offset, blockN);
// Only the first high-high product of the first outer K tile
// initializes TMEM. Every later product accumulates into it.
const int accumulate_hi_hi = (i != 0 || id != 0) ? 1 : 0;
tcgen05mma_tf32(
taddr, a_hi_desc, b_hi_desc, i_desc, accumulate_hi_hi);
// Compensated TF32 correction terms. The omitted A_lo * B_lo
// term is second order in the TF32 rounding residual.
tcgen05mma_tf32(
taddr, a_hi_desc, b_lo_desc, i_desc, 1);
tcgen05mma_tf32(
taddr, a_lo_desc, b_hi_desc, i_desc, 1);
}
asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
:: "r"(mbar_addr) : "memory");
}
mbarrier_wait(mbar_addr, phase);
phase ^=1;
}
asm volatile ("tcgen05.fence::after_thread_sync;");
for(int n = 0 ; n < blockN/ 8 ; n++){
float tmp[8];
const int addr = taddr + ((warp_id * 32) << 16) + (n * 8);
asm volatile ("tcgen05.ld.sync.aligned.32x32b.x8.b32 {%0, %1 , %2, %3 , %4, %5, %6, %7}, [%8];"
: "=f"(tmp[0]), "=f"(tmp[1]), "=f"(tmp[2]), "=f"(tmp[3]),
"=f"(tmp[4]),"=f"(tmp[5]),"=f"(tmp[6]),"=f"(tmp[7])
: "r"(addr));
asm volatile("tcgen05.wait::ld.sync.aligned;");
const int row = off_m + tid;
const int col0 = off_n + n * 8;
if (row < M) {
#pragma unroll
for (int x = 0; x < 8; ++x) {
const int col = col0 + x;
if (col < N) {
const int64_t out_idx =
static_cast<int64_t>(batch_idx) * output_batch_stride
+ output_base_offset
+ static_cast<int64_t>(row) * output_row_stride
+ col;
if constexpr (FUSE_SUBTRACT) {
C_ptr[out_idx] -= tmp[x];
} else {
C_ptr[out_idx] = tmp[x];
}
}
}
}
}
__syncthreads();
if(warp_id == 0){
asm volatile ("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0 , %1;" :: "r"(taddr), "r"(blockN));
}
}
__global__ void panel_factor_kernel(
float* H,
float* tau,
int n,
int j,
int pb) {
const int batch_idx = blockIdx.x;
const int tid = threadIdx.x;
extern __shared__ float shared[];
float* v_tile = shared; // [kReflectorTileRows]
float* w_tile = v_tile + kReflectorTileRows; // [kPanelColumnTile]
__shared__ float reflector_tau;
__shared__ float reflector_inv;
for (int k = 0; k < pb; ++k) {
const int col_idx = j + k;
const int active_len = n - col_idx;
const float alpha_local =
tid == 0 ? H[offset3(batch_idx, col_idx, col_idx, n)] : 0.0f;
const float alpha = block_broadcast(alpha_local);
float sigma_local = 0.0f;
for (int row = 1 + tid; row < active_len; row += blockDim.x) {
const float v = H[offset3(batch_idx, col_idx + row, col_idx, n)];
sigma_local += v * v;
}
const float sigma = block_reduce_sum(sigma_local);
if (tid == 0) {
float tau_k = 0.0f;
float inv = 0.0f;
if (sigma != 0.0f) {
const float norm = sqrtf(alpha * alpha + sigma);
const float beta = alpha <= 0.0f ? norm : -norm;
inv = 1.0f / (alpha - beta);
tau_k = (beta - alpha) / beta;
H[offset3(batch_idx, col_idx, col_idx, n)] = beta;
}
reflector_tau = tau_k;
reflector_inv = inv;
tau[tau_offset(batch_idx, col_idx, n)] = tau_k;
}
__syncthreads();
const float tau_k = reflector_tau;
const float inv = reflector_inv;
// Scale the Householder tail cooperatively instead of making thread 0
// walk the complete active column serially.
if (tau_k != 0.0f) {
for (int row = 1 + tid; row < active_len; row += blockDim.x) {
H[offset3(batch_idx, col_idx + row, col_idx, n)] *= inv;
}
}
__syncthreads();
if (tau_k == 0.0f) {
continue;
}
for (int panel_col0 = k + 1; panel_col0 < pb; panel_col0 += kPanelColumnTile) {
const int cols_this_tile = min(kPanelColumnTile, pb - panel_col0);
float dot_accum[kPanelColumnTile] = {};
// Pass 1: accumulate dot products v^T * A_tile.
for (int row0 = 0; row0 < active_len; row0 += kReflectorTileRows) {
const int rows_this_tile = min(kReflectorTileRows, active_len - row0);
for (int t = tid; t < rows_this_tile; t += blockDim.x) {
if (row0 + t == 0) {
v_tile[t] = 1.0f;
} else {
v_tile[t] =
H[offset3(batch_idx, col_idx + row0 + t, col_idx, n)];
}
}
__syncthreads();
for (int t = tid; t < rows_this_tile; t += blockDim.x) {
const float v = v_tile[t];
const int global_row = col_idx + row0 + t;
#pragma unroll
for (int c = 0; c < kPanelColumnTile; ++c) {
if (c < cols_this_tile) {
const int target_col = j + panel_col0 + c;
dot_accum[c] +=
v * H[offset3(batch_idx, global_row, target_col, n)];
}
}
}
__syncthreads();
}
#pragma unroll
for (int c = 0; c < kPanelColumnTile; ++c) {
if (c < cols_this_tile) {
const float dot = block_reduce_sum(dot_accum[c]);
if (tid == 0) {
w_tile[c] = tau_k * dot;
}
__syncthreads();
}
}
// Pass 2: apply A_tile -= v * w_tile.
for (int row0 = 0; row0 < active_len; row0 += kReflectorTileRows) {
const int rows_this_tile = min(kReflectorTileRows, active_len - row0);
for (int t = tid; t < rows_this_tile; t += blockDim.x) {
if (row0 + t == 0) {
v_tile[t] = 1.0f;
} else {
v_tile[t] =
H[offset3(batch_idx, col_idx + row0 + t, col_idx, n)];
}
}
__syncthreads();
for (int t = tid; t < rows_this_tile; t += blockDim.x) {
const float v = v_tile[t];
const int global_row = col_idx + row0 + t;
#pragma unroll
for (int c = 0; c < kPanelColumnTile; ++c) {
if (c < cols_this_tile) {
const int target_col = j + panel_col0 + c;
H[offset3(batch_idx, global_row, target_col, n)] -=
v * w_tile[c];
}
}
}
__syncthreads();
}
}
}
}
__global__ void pack_v_kernel(
const float* H,
float* V,
float* Vt,
int n,
int j,
int pb) {
const int batch_idx = blockIdx.x;
const int tid = threadIdx.x;
const int m = n - j;
for (int idx = tid; idx < m * pb; idx += blockDim.x) {
const int row = idx / pb;
const int col = idx % pb;
float value = 0.0f;
if (row == col) {
value = 1.0f;
} else if (row > col) {
value = H[offset3(batch_idx, j + row, j + col, n)];
}
V[(static_cast<int64_t>(batch_idx) * n + row) * kPanelWidth + col] = value;
Vt[(static_cast<int64_t>(batch_idx) * kPanelWidth + col) * n + row] = value;
}
}
__global__ void build_t_kernel(
const float* H,
const float* tau,
float* T,
int n,
int j,
int pb) {
const int batch_idx = blockIdx.x;
const int tid = threadIdx.x;
const int lane_id = tid & (WARP_SIZE - 1);
const int warp_id = tid / WARP_SIZE;
constexpr int WARPS_PER_BLOCK = kPanelThreads / WARP_SIZE;
extern __shared__ float shared[];
float* t_local = shared; // [pb, pb]
float* tmp = t_local + pb * pb; // [pb]
for (int idx = tid; idx < pb * pb; idx += blockDim.x) {
t_local[idx] = 0.0f;
}
__syncthreads();
for (int k = 0; k < pb; ++k) {
const float tau_k = tau[tau_offset(batch_idx, j + k, n)];
if (tid == 0) {
t_local[k * pb + k] = tau_k;
}
__syncthreads();
if (tau_k == 0.0f || k == 0) {
continue;
}
// One warp computes each V_i^T V_k dot product. Lanes split the long
// row dimension and reduce locally with shuffles.
for (int i = warp_id; i < k; i += WARPS_PER_BLOCK) {
float dot = 0.0f;
for (int row = j + k + 1 + lane_id;
row < n;
row += WARP_SIZE) {
const float v_i = H[offset3(batch_idx, row, j + i, n)];
const float v_k = H[offset3(batch_idx, row, j + k, n)];
dot += v_i * v_k;
}
dot = warp_reduce_sum<WARP_SIZE>(dot);
if (lane_id == 0) {
const float diagonal_term =
H[offset3(batch_idx, j + k, j + i, n)];
tmp[i] = -tau_k * (diagonal_term + dot);
}
}
__syncthreads();
// T(0:k-1,k) = T(0:k-1,0:k-1) * tmp.
for (int i = tid; i < k; i += blockDim.x) {
float sum = 0.0f;
for (int p = 0; p < k; ++p) {
sum += t_local[i * pb + p] * tmp[p];
}
t_local[i * pb + k] = sum;
}
__syncthreads();
}
for (int idx = tid; idx < pb * pb; idx += blockDim.x) {
const int row = idx / pb;
const int col = idx % pb;
T[(static_cast<int64_t>(batch_idx) * kPanelWidth + row)
* kPanelWidth + col] = t_local[idx];
}
}
__global__ void apply_t_kernel(
const float* T,
const float* Y,
float* Zt,
int workspace_n,
int trailing_cols,
int pb) {
const int batch_idx = blockIdx.x;
const int tile_col = blockIdx.y;
const int tid = threadIdx.x;
const int col0 = tile_col * kTileCols;
if (col0 >= trailing_cols) {
return;
}
// Z = T^T Y, stored directly as Zt[batch, trailing_cols, pb].
for (int p = 0; p < pb; ++p) {
for (int col = tid;
col < kTileCols && (col0 + col) < trailing_cols;
col += blockDim.x) {
float sum = 0.0f;
for (int q = 0; q <= p; ++q) {
const float tqp =
T[(static_cast<int64_t>(batch_idx) * kPanelWidth + q)
* kPanelWidth + p];
const float y_val =
Y[(static_cast<int64_t>(batch_idx) * kPanelWidth + q)
* workspace_n + (col0 + col)];
sum += tqp * y_val;
}
Zt[(static_cast<int64_t>(batch_idx) * workspace_n + (col0 + col))
* kPanelWidth + p] = sum;
}
}
}
template <bool FUSE_SUBTRACT, int BLOCK_M, int BLOCK_N, int BLOCK_K>
void matmul_batched_launch_impl(
const float *A_ptr,
const float *B_ptr,
float *C_ptr,
int batch,
int M, int N, int K,
int64_t A_row_stride,
int64_t A_batch_stride,
int64_t B_row_stride,
int64_t B_batch_stride,
int64_t output_batch_stride,
int output_row_stride,
int64_t output_base_offset
) {
CUtensorMap A_tmap, B_tmap;
init_tmap_3d_strided(
&A_tmap, A_ptr, batch, M, K,
A_row_stride, A_batch_stride,
BLOCK_M, 4,
CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_NONE);
init_tmap_3d_strided(
&B_tmap, B_ptr, batch, N, K,
B_row_stride, B_batch_stride,
BLOCK_N, 4,
CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_NONE);
const int grid_m = cdiv(M, BLOCK_M);
const int grid_n = cdiv(N, BLOCK_N);
const int grid = batch * grid_m * grid_n;
const int size_AB = 2 * (BLOCK_M + BLOCK_N) * BLOCK_K;
const int smem_size = size_AB * static_cast<int>(sizeof(float));
auto this_kernel = GEMM_kernel<FUSE_SUBTRACT, BLOCK_M, BLOCK_N, BLOCK_K>;
if (smem_size > 48'000) {
C10_CUDA_CHECK(cudaFuncSetAttribute(
this_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
smem_size));
}
this_kernel<<<grid, TB_SIZE, smem_size>>>(
A_tmap, B_tmap, C_ptr, batch, M, N, K,
output_batch_stride, output_row_stride, output_base_offset);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
template<int blockM, int blockN, int blockK>
__global__
__launch_bounds__(TB_SIZE)
void Y_GEMM_kernel(
const __grid_constant__ CUtensorMap A_tmap,
const float* H,
float* Y,
int batch,
int n,
int j,
int pb,
int m,
int trailing_cols
) {
const int tid = threadIdx.x;
const int warp_id = tid / WARP_SIZE;
const int grid_m = cdiv(pb, blockM);
const int grid_n = cdiv(trailing_cols, blockN);
const int tiles_per_batch = grid_m * grid_n;
const int linear_block = static_cast<int>(blockIdx.x);
const int batch_idx = linear_block / tiles_per_batch;
const int tile_idx = linear_block - batch_idx * tiles_per_batch;
const int bid_m = tile_idx / grid_n;
const int bid_n = tile_idx % grid_n;
if (batch_idx >= batch) return;
const int off_m = bid_m * blockM;
const int off_n = bid_n * blockN;
extern __shared__ __align__(1024) char smem[];
constexpr int A_elems = blockM * blockK;
constexpr int B_elems = blockN * blockK;
constexpr int A_bytes = A_elems * static_cast<int>(sizeof(float));
constexpr int B_bytes = B_elems * static_cast<int>(sizeof(float));
const int A_hi_smem = static_cast<int>(__cvta_generic_to_shared(smem));
const int B_hi_smem = A_hi_smem + A_bytes;
const int A_lo_smem = B_hi_smem + B_bytes;
const int B_lo_smem = A_lo_smem + A_bytes;
float* const A_hi_ptr = reinterpret_cast<float*>(smem);
float* const B_hi_ptr = A_hi_ptr + A_elems;
float* const A_lo_ptr = B_hi_ptr + B_elems;
float* const B_lo_ptr = A_lo_ptr + A_elems;
__shared__ uint64_t mbars[1];
__shared__ int tmem_adder[1];
const int mbar_addr = static_cast<int>(__cvta_generic_to_shared(mbars));
if (warp_id == 0 && elect_thr()) {
asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(mbar_addr), "r"(1));
asm volatile("fence.mbarrier_init.release.cluster;");
} else if (warp_id == 1) {
const int addr = static_cast<int>(__cvta_generic_to_shared(tmem_adder));
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
:: "r"(addr), "r"(blockN));
}
__syncthreads();
int phase = 0;
const int taddr = tmem_adder[0];
constexpr uint32_t i_desc = (1U << 4U)
| (2U << 7U)
| (2U << 10U)
| ((uint32_t)blockN >> 3U << 17U)
| ((uint32_t)blockM >> 4U << 24U);
const int iters = cdiv(m, blockK);
for (int i = 0; i < iters; ++i) {
if (warp_id == 0 && elect_thr()) {
#pragma unroll
for (int stride_id = 0; stride_id < blockK / 4; ++stride_id) {
const int off_k = i * blockK + stride_id * 4;
tma_3d_gmem2smem<1>(
A_hi_smem + stride_id * blockM * 16,
&A_tmap,
off_k,
off_m,
batch_idx,
mbar_addr);
}
constexpr int cp_size = blockM * blockK * sizeof(float);
asm volatile(
"mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
:: "r"(mbar_addr), "r"(cp_size) : "memory");
}
// Repack the original trailing matrix C directly into the same shared
// layout previously produced by materializing C^T and loading it by TMA.
for (int idx = tid; idx < B_elems; idx += TB_SIZE) {
const int chunk = idx / (blockN * 4);
const int rem = idx - chunk * blockN * 4;
const int local_n = rem / 4;
const int local_k4 = rem & 3;
const int local_k = chunk * 4 + local_k4;
const int global_k = i * blockK + local_k;
const int global_n = off_n + local_n;
float value = 0.0f;
if (global_k < m && global_n < trailing_cols) {
value = H[
(static_cast<int64_t>(batch_idx) * n + (j + global_k)) * n
+ (j + pb + global_n)];
}
B_hi_ptr[idx] = value;
}
mbarrier_wait(mbar_addr, phase);
phase ^= 1;
__syncthreads();
for (int idx = tid; idx < A_elems; idx += TB_SIZE) {
const float x = A_hi_ptr[idx];
const float hi = nvcuda::wmma::__float_to_tf32(x);
A_hi_ptr[idx] = hi;
A_lo_ptr[idx] = x - hi;
}
for (int idx = tid; idx < B_elems; idx += TB_SIZE) {
const float x = B_hi_ptr[idx];
const float hi = nvcuda::wmma::__float_to_tf32(x);
B_hi_ptr[idx] = hi;
B_lo_ptr[idx] = x - hi;
}
__syncthreads();
asm volatile("tcgen05.fence::after_thread_sync;");
if (warp_id == 0 && elect_thr()) {
auto desc = [](int addr, int height) -> uint64_t {
const int LBO = height * 16;
const int SBO = 8 * 16;
return desc_encode(addr)
| (desc_encode(LBO) << 16ULL)
| (desc_encode(SBO) << 32ULL)
| (1ULL << 46ULL);
};
#pragma unroll
for (int id = 0; id < blockK / MMA_K; ++id) {
const int a_offset = id * blockM * MMA_K * static_cast<int>(sizeof(float));
const int b_offset = id * blockN * MMA_K * static_cast<int>(sizeof(float));
const uint64_t a_hi_desc = desc(A_hi_smem + a_offset, blockM);
const uint64_t b_hi_desc = desc(B_hi_smem + b_offset, blockN);
const uint64_t a_lo_desc = desc(A_lo_smem + a_offset, blockM);
const uint64_t b_lo_desc = desc(B_lo_smem + b_offset, blockN);
const int accumulate_hi_hi = (i != 0 || id != 0) ? 1 : 0;
tcgen05mma_tf32(taddr, a_hi_desc, b_hi_desc, i_desc, accumulate_hi_hi);
tcgen05mma_tf32(taddr, a_hi_desc, b_lo_desc, i_desc, 1);
tcgen05mma_tf32(taddr, a_lo_desc, b_hi_desc, i_desc, 1);
}
asm volatile(
"tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
:: "r"(mbar_addr) : "memory");
}
mbarrier_wait(mbar_addr, phase);
phase ^= 1;
}
asm volatile("tcgen05.fence::after_thread_sync;");
for (int n8 = 0; n8 < blockN / 8; ++n8) {
float tmp[8];
const int addr = taddr + ((warp_id * 32) << 16) + (n8 * 8);
asm volatile(
"tcgen05.ld.sync.aligned.32x32b.x8.b32 {%0, %1, %2, %3, %4, %5, %6, %7}, [%8];"
: "=f"(tmp[0]), "=f"(tmp[1]), "=f"(tmp[2]), "=f"(tmp[3]),
"=f"(tmp[4]), "=f"(tmp[5]), "=f"(tmp[6]), "=f"(tmp[7])
: "r"(addr));
asm volatile("tcgen05.wait::ld.sync.aligned;");
const int row = off_m + tid;
const int col0 = off_n + n8 * 8;
if (row < pb) {
#pragma unroll
for (int x = 0; x < 8; ++x) {
const int col = col0 + x;
if (col < trailing_cols) {
Y[(static_cast<int64_t>(batch_idx) * kPanelWidth + row) * n + col] = tmp[x];
}
}
}
}
__syncthreads();
if (warp_id == 0) {
asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;"
:: "r"(taddr), "r"(blockN));
}
}
template<int BLOCK_M, int BLOCK_N, int BLOCK_K>
void y_matmul_batched_direct_c(
const float* Vt_ptr,
const float* H_ptr,
float* Y_ptr,
int batch,
int n,
int j,
int pb,
int m,
int trailing_cols) {
CUtensorMap A_tmap;
init_tmap_3d_strided(
&A_tmap,
Vt_ptr,
batch,
pb,
m,
n,
static_cast<uint64_t>(kPanelWidth) * n,
BLOCK_M,
4,
CUtensorMapSwizzle::CU_TENSOR_MAP_SWIZZLE_NONE);
const int grid = batch * cdiv(pb, BLOCK_M) * cdiv(trailing_cols, BLOCK_N);
const int smem_size =
2 * (BLOCK_M + BLOCK_N) * BLOCK_K * static_cast<int>(sizeof(float));
auto kernel = Y_GEMM_kernel<BLOCK_M, BLOCK_N, BLOCK_K>;
if (smem_size > 48'000) {
C10_CUDA_CHECK(cudaFuncSetAttribute(
kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size));
}
kernel<<<grid, TB_SIZE, smem_size>>>(
A_tmap, H_ptr, Y_ptr, batch, n, j, pb, m, trailing_cols);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
template <int BLOCK_M, int BLOCK_N, int BLOCK_K>
void matmul_batched_subtract(
const float *A_ptr,
const float *B_ptr,
float *H_ptr,
int batch,
int n,
int M, int N, int K,
int row_base,
int col_base
) {
matmul_batched_launch_impl<true, BLOCK_M, BLOCK_N, BLOCK_K>(
A_ptr, B_ptr, H_ptr,
batch, M, N, K,
kPanelWidth,
static_cast<int64_t>(n) * kPanelWidth,
kPanelWidth,
static_cast<int64_t>(n) * kPanelWidth,
static_cast<int64_t>(n) * n,
n,
static_cast<int64_t>(row_base) * n + col_base);
}
std::tuple<torch::Tensor , torch::Tensor> compact_qr_reference(torch::Tensor a) {
auto [H,tau] = at::geqrf(a);
return {
H,
tau
};
}
std::tuple<torch::Tensor, torch::Tensor> compact_qr_stitched(torch::Tensor a) {
const c10::cuda::CUDAGuard device_guard(a.device());
TORCH_CHECK(a.is_cuda(), "compact_qr expects a CUDA tensor");
TORCH_CHECK(a.dtype() == torch::kFloat32, "compact_qr expects float32 input");
TORCH_CHECK(a.dim() == 3, "compact_qr expects a [batch, n, n] tensor");
TORCH_CHECK(a.size(1) == a.size(2), "compact_qr expects square matrices");
TORCH_CHECK(a.is_contiguous(), "compact_qr expects contiguous input");
auto H = a.clone();
auto tau = torch::zeros({a.size(0), a.size(1)}, a.options());
const int batch = static_cast<int>(H.size(0));
const int n = static_cast<int>(H.size(1));
float* H_ptr = H.data_ptr<float>();
float* tau_ptr = tau.data_ptr<float>();
// Reusable fixed-stride workspaces. Every panel writes only its active
// [m, pb] or [pb, trailing_cols] region.
auto T_workspace = torch::empty(
{batch, kPanelWidth, kPanelWidth}, a.options());
auto V_workspace = torch::empty(
{batch, n, kPanelWidth}, a.options());
auto Vt_workspace = torch::empty(
{batch, kPanelWidth, n}, a.options());
auto Y_workspace = torch::empty(
{batch, kPanelWidth, n}, a.options());
auto Zt_workspace = torch::empty(
{batch, n, kPanelWidth}, a.options());
float* T_ptr = T_workspace.data_ptr<float>();
float* V_ptr = V_workspace.data_ptr<float>();
float* Vt_ptr = Vt_workspace.data_ptr<float>();
float* Y_ptr = Y_workspace.data_ptr<float>();
float* Zt_ptr = Zt_workspace.data_ptr<float>();
cudaEvent_t panel_start;
cudaEvent_t panel_stop;
cudaEvent_t pack_v_start;
cudaEvent_t pack_v_stop;
cudaEvent_t t_start;
cudaEvent_t t_stop;
cudaEvent_t update_start;
cudaEvent_t update_stop;
cudaEvent_t y_start;
cudaEvent_t y_stop;
cudaEvent_t z_start;
cudaEvent_t z_stop;
cudaEvent_t c_start;
cudaEvent_t c_stop;
float panel_ms_total = 0.0f;
float pack_v_ms_total = 0.0f;
float t_ms_total = 0.0f;
float update_ms_total = 0.0f;
float y_ms_total = 0.0f;
float z_ms_total = 0.0f;
float c_ms_total = 0.0f;
if constexpr (kEnableKernelTiming) {
C10_CUDA_CHECK(cudaEventCreate(&panel_start));
C10_CUDA_CHECK(cudaEventCreate(&panel_stop));
C10_CUDA_CHECK(cudaEventCreate(&pack_v_start));
C10_CUDA_CHECK(cudaEventCreate(&pack_v_stop));
C10_CUDA_CHECK(cudaEventCreate(&t_start));
C10_CUDA_CHECK(cudaEventCreate(&t_stop));
C10_CUDA_CHECK(cudaEventCreate(&update_start));
C10_CUDA_CHECK(cudaEventCreate(&update_stop));
C10_CUDA_CHECK(cudaEventCreate(&y_start));
C10_CUDA_CHECK(cudaEventCreate(&y_stop));
C10_CUDA_CHECK(cudaEventCreate(&z_start));
C10_CUDA_CHECK(cudaEventCreate(&z_stop));
C10_CUDA_CHECK(cudaEventCreate(&c_start));
C10_CUDA_CHECK(cudaEventCreate(&c_stop));
}
for (int j = 0; j < n; j += kPanelWidth) {
const int pb = std::min(kPanelWidth, n - j); //works because its a square matrix
const size_t panel_shared_bytes =
static_cast<size_t>(kReflectorTileRows + kPanelColumnTile) * sizeof(float);
if constexpr (kEnableKernelTiming) {
C10_CUDA_CHECK(cudaEventRecord(panel_start));
}
panel_factor_kernel<<<batch, kPanelThreads, panel_shared_bytes>>>(
H_ptr, tau_ptr, n, j, pb);
C10_CUDA_KERNEL_LAUNCH_CHECK();
if constexpr (kEnableKernelTiming) {
float elapsed_ms = 0.0f;
C10_CUDA_CHECK(cudaEventRecord(panel_stop));
C10_CUDA_CHECK(cudaEventSynchronize(panel_stop));
C10_CUDA_CHECK(cudaEventElapsedTime(&elapsed_ms, panel_start, panel_stop));
panel_ms_total += elapsed_ms;
}
if constexpr (kEnableKernelTiming) {
C10_CUDA_CHECK(cudaEventRecord(pack_v_start));
}
pack_v_kernel<<<batch, kPackVThreads>>>(
H_ptr, V_ptr, Vt_ptr, n, j, pb);
C10_CUDA_KERNEL_LAUNCH_CHECK();
if constexpr (kEnableKernelTiming) {
float elapsed_ms = 0.0f;
C10_CUDA_CHECK(cudaEventRecord(pack_v_stop));
C10_CUDA_CHECK(cudaEventSynchronize(pack_v_stop));
C10_CUDA_CHECK(cudaEventElapsedTime(&elapsed_ms, pack_v_start, pack_v_stop));
pack_v_ms_total += elapsed_ms;
}
const size_t t_shared_bytes =
static_cast<size_t>(pb * pb + pb) * sizeof(float);
if constexpr (kEnableKernelTiming) {
C10_CUDA_CHECK(cudaEventRecord(t_start));
}
build_t_kernel<<<batch, kPanelThreads, t_shared_bytes>>>(
H_ptr, tau_ptr, T_ptr, n, j, pb);
C10_CUDA_KERNEL_LAUNCH_CHECK();
if constexpr (kEnableKernelTiming) {
float elapsed_ms = 0.0f;
C10_CUDA_CHECK(cudaEventRecord(t_stop));
C10_CUDA_CHECK(cudaEventSynchronize(t_stop));
C10_CUDA_CHECK(cudaEventElapsedTime(&elapsed_ms, t_start, t_stop));
t_ms_total += elapsed_ms;
}
if (j + pb < n) {
const int trailing_cols = n - (j + pb);
const int m = n - j;
dim3 grid(batch,
(trailing_cols + kTileCols - 1) / kTileCols);
if constexpr (kEnableKernelTiming) {
C10_CUDA_CHECK(cudaEventRecord(update_start));
C10_CUDA_CHECK(cudaEventRecord(y_start));
}
// Y = V^T C. Vt is produced directly by pack_v_kernel, while C is
// read from H and repacked CTA-locally inside the Y GEMM kernel.
y_matmul_batched_direct_c<kYGemmBlockM, kYGemmBlockN, kYGemmBlockK>(
Vt_ptr,
H_ptr,
Y_ptr,
batch,
n,
j,
pb,
m,
trailing_cols);
if constexpr (kEnableKernelTiming) {
float elapsed_ms = 0.0f;
C10_CUDA_CHECK(cudaEventRecord(y_stop));
C10_CUDA_CHECK(cudaEventSynchronize(y_stop));
C10_CUDA_CHECK(cudaEventElapsedTime(&elapsed_ms, y_start, y_stop));
y_ms_total += elapsed_ms;
}
if constexpr (kEnableKernelTiming) {
C10_CUDA_CHECK(cudaEventRecord(z_start));
}
apply_t_kernel<<<grid, kUpdateThreads>>>(
T_ptr, Y_ptr, Zt_ptr, n, trailing_cols, pb);
C10_CUDA_KERNEL_LAUNCH_CHECK();
if constexpr (kEnableKernelTiming) {
float elapsed_ms = 0.0f;
C10_CUDA_CHECK(cudaEventRecord(z_stop));
C10_CUDA_CHECK(cudaEventSynchronize(z_stop));
C10_CUDA_CHECK(cudaEventElapsedTime(&elapsed_ms, z_start, z_stop));
z_ms_total += elapsed_ms;
}
if constexpr (kEnableKernelTiming) {
C10_CUDA_CHECK(cudaEventRecord(c_start));
}
// Compute VZ and fuse the epilogue directly into the trailing matrix:
// H[:, j:, j+pb:] -= V @ Z.
matmul_batched_subtract<kGemmBlockM, kGemmBlockN, kGemmBlockK>(
V_ptr,
Zt_ptr,
H_ptr,
batch,
n,
m,
trailing_cols,
pb,
j,
j + pb);
//C10_CUDA_CHECK(cudaDeviceSynchronize());
// std::printf("[debug] tcgen completed j=%d\n", j);
// std::fflush(stdout);
C10_CUDA_KERNEL_LAUNCH_CHECK();
if constexpr (kEnableKernelTiming) {
float elapsed_ms = 0.0f;
C10_CUDA_CHECK(cudaEventRecord(c_stop));
C10_CUDA_CHECK(cudaEventSynchronize(c_stop));
C10_CUDA_CHECK(cudaEventElapsedTime(&elapsed_ms, c_start, c_stop));
c_ms_total += elapsed_ms;
}
if constexpr (kEnableKernelTiming) {
float elapsed_ms = 0.0f;
C10_CUDA_CHECK(cudaEventRecord(update_stop));
C10_CUDA_CHECK(cudaEventSynchronize(update_stop));
C10_CUDA_CHECK(cudaEventElapsedTime(&elapsed_ms, update_start, update_stop));
update_ms_total += elapsed_ms;
}
}
}
if constexpr (kEnableKernelTiming) {
// std::printf(
// "[compact_qr] n=%d batch=%d panel_ms=%.3f pack_v_ms=%.3f build_t_ms=%.3f update_ms=%.3f (y=%.3f z=%.3f c=%.3f) total_ms=%.3f\n",
// n,
// batch,
// panel_ms_total,
// pack_v_ms_total,
// t_ms_total,
// update_ms_total,
// y_ms_total,
// z_ms_total,
// c_ms_total,
// panel_ms_total + pack_v_ms_total + t_ms_total + update_ms_total);
C10_CUDA_CHECK(cudaEventDestroy(panel_start));
C10_CUDA_CHECK(cudaEventDestroy(panel_stop));
C10_CUDA_CHECK(cudaEventDestroy(pack_v_start));
C10_CUDA_CHECK(cudaEventDestroy(pack_v_stop));
C10_CUDA_CHECK(cudaEventDestroy(t_start));
C10_CUDA_CHECK(cudaEventDestroy(t_stop));
C10_CUDA_CHECK(cudaEventDestroy(update_start));
C10_CUDA_CHECK(cudaEventDestroy(update_stop));
C10_CUDA_CHECK(cudaEventDestroy(y_start));
C10_CUDA_CHECK(cudaEventDestroy(y_stop));
C10_CUDA_CHECK(cudaEventDestroy(z_start));
C10_CUDA_CHECK(cudaEventDestroy(z_stop));
C10_CUDA_CHECK(cudaEventDestroy(c_start));
C10_CUDA_CHECK(cudaEventDestroy(c_stop));
}
return std::make_tuple(H, tau);
}
std::tuple<torch::Tensor , torch::Tensor> compact_qr_cuda(torch::Tensor a) {
TORCH_CHECK(a.is_cuda(), "compact_qr expects a CUDA tensor");
TORCH_CHECK(a.dtype() == torch::kFloat32, "compact_qr expects float32 input");
TORCH_CHECK(a.dim() == 3, "compact_qr expects a [batch, n, n] tensor");
TORCH_CHECK(a.size(1) == a.size(2), "compact_qr expects square matrices");
if constexpr (kUseReferencePath) {
// Keep the extension usable while the stitched kernel path is under
// development. Flip this constant once the custom path is ready to test.
return compact_qr_reference(a);
}
return compact_qr_stitched(a.contiguous());
}
} // namespace
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("compact_qr", &compact_qr_cuda, "Batched compact Householder QR (CUDA)");
}
// Paste the complete contents of your qr.cu file here.
//
// Keep this binding at the bottom:
//
// PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
// m.def(
// "compact_qr",
// &compact_qr_cuda,
// "Batched compact Householder QR (CUDA)"
// );
// }
"""
_QR_EXT = None
def _load_qr_extension():
global _QR_EXT
if _QR_EXT is None:
_QR_EXT = load_inline(
name="compact_qr_inline_ext_v2_compensated_tf32",
cpp_sources="",
cuda_sources=CUDA_SOURCE,
functions=None,
with_cuda=True,
extra_cuda_cflags=["-O3",
"-std=c++17",
"-gencode=arch=compute_100a,code=sm_100a",],
extra_ldflags=["-lcuda",],
verbose=True,
)
return _QR_EXT
def custom_kernel(data: input_t) -> output_t:
if not data.is_cuda:
return ref_kernel(data)
ext = _load_qr_extension()
return ext.compact_qr(data)
def _property_rtol(n: int, factor: float) -> float:
eps = torch.finfo(torch.float32).eps
return factor * max(n, 1) * eps
def _scaled_residual(
residual: torch.Tensor,
scale: torch.Tensor,
n: int,
) -> torch.Tensor:
eps = torch.finfo(torch.float32).eps
return residual / (eps * max(n, 1) * scale.clamp_min(1e-30))
def _matrix_l1_norm(value: torch.Tensor) -> torch.Tensor:
return torch.linalg.matrix_norm(value.double(), ord=1, dim=(-2, -1))
def _check_tensor(name: str, value: torch.Tensor, shape: tuple[int, ...], device: torch.device) -> str | None:
if not isinstance(value, torch.Tensor):
return f"{name} must be a torch.Tensor"
if value.shape != shape:
return f"{name} shape must be {shape}, got {tuple(value.shape)}"
if value.dtype != torch.float32:
return f"{name} dtype must be torch.float32, got {value.dtype}"
if value.device != device:
return f"{name} must be on {device}, got {value.device}"
if not torch.isfinite(value).all().item():
return f"{name} contains NaN or Inf"
return None
def check_implementation(data: input_t, output: output_t) -> tuple[bool, str]:
a = data
batch, n, _ = a.shape
factor_rtol = _property_rtol(n, _FACTOR_RTOL_FACTOR)
orth_rtol = _property_rtol(n, _ORTH_RTOL_FACTOR)
if not isinstance(output, tuple) or len(output) != 2:
return False, "output must be a tuple `(H, tau)`"
h, tau = output
error = _check_tensor("H", h, (batch, n, n), a.device)
if error is not None:
return False, error
error = _check_tensor("tau", tau, (batch, n), a.device)
if error is not None:
return False, error
q = torch.linalg.householder_product(h, tau)
r = torch.triu(h)
a_check = a.double()
q_check = q.double()
r_check = r.double()
projected = q_check.transpose(-1, -2) @ a_check
factor_residual = _matrix_l1_norm(r_check - projected).amax()
factor_scale = _matrix_l1_norm(a_check).amax()
factor_allowed = factor_rtol * factor_scale
factor_scaled = _scaled_residual(factor_residual, factor_scale, n)
if factor_residual.item() > factor_allowed.item():
return False, (
"R - Q.T @ A is too large: "
f"residual={factor_residual.item():.3g}, allowed={factor_allowed.item():.3g}, "
f"scaled={factor_scaled.item():.3g}"
)
eye = torch.eye(n, device=a.device, dtype=torch.float64).expand(batch, n, n)
qtq = q_check.transpose(-1, -2) @ q_check
orth_residual = _matrix_l1_norm(qtq - eye).amax()
orth_scale = _matrix_l1_norm(eye).amax()
orth_allowed = orth_rtol * orth_scale
orth_scaled = _scaled_residual(orth_residual, orth_scale, n)
if orth_residual.item() > orth_allowed.item():
return False, (
"Q is not orthogonal enough: "
f"residual={orth_residual.item():.3g}, allowed={orth_allowed.item():.3g}, "
f"scaled={orth_scaled.item():.3g}"
)
lower = torch.tril(projected, diagonal=-1)
tri_residual = _matrix_l1_norm(lower).amax()
tri_scale = _matrix_l1_norm(a_check).amax()
tri_scaled = _scaled_residual(tri_residual, tri_scale, n)
recon = q_check @ r_check
recon_residual = _matrix_l1_norm(recon - a_check).amax()
recon_scale = _matrix_l1_norm(a_check).amax()
recon_scaled = _scaled_residual(recon_residual, recon_scale, n)
return True, (
f"factor_rtol={factor_rtol:.3g}; "
f"orth_rtol={orth_rtol:.3g}; "
f"scaled_factor_residual={factor_scaled.item():.3g}; "
f"scaled_reconstruction_residual={recon_scaled.item():.3g}; "
f"scaled_triangular_residual={tri_scaled.item():.3g}; "
f"scaled_orthogonality_residual={orth_scaled.item():.3g}; "
f"batch={batch}; n={n}"
)
scrolls · 1579 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