Skip to content
KernelIndex
Search⌘K

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
NVIDIA B200
40.9ms
#383 of 515
2026-06-29

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.

mbarrierasm volatile("cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.cta_group::%6 "
mmaconst float hi = nvcuda::wmma::__float_to_tf32(x);
num-warps = 4constexpr int NUM_WARPS = 4;
shared-memorystatic __shared__ float warp_sums[NUM_WARP];
tcgen05"tcgen05.mma.cta_group::%5.kind::tf32 [%0], %1 , %2 , %3, p;\n\t"
tmaasm 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