Skip to content
KernelIndex
Search⌘K

submission 843960

vvrp · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 1417 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-843960?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
2.85ms
#67 of 515
2026-06-29

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f8426060fd8d345f5cd088df9ea63e81ee3a7ac047d3904d7758075ca93de732
license declaredunknown
license concludedunknown
authorsvvrp
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

shared-memoryconstexpr int Warps = Threads / 32; __shared__ unsigned int warp_maxima[2 * Warps]; const int b = blockIdx.x; const float* a_b = a + static_cast<long long>(b) * N * N; const int tid = threadIdx.x;
vector-width = float4__device__ __forceinline__ float4 ptx_ld_global_v4_f32(const float* ptr) {

Kernel source

submission.py1417 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200
import torch
import tempfile
from pathlib import Path
from task import input_t, output_t
torch.backends.cuda.matmul.allow_tf32 = True
CPP_SRC = r"""
#include <torch/extension.h>
std::vector<torch::Tensor> qr_small_cuda(torch::Tensor a);
std::vector<torch::Tensor> qr_small_prefix_cuda(torch::Tensor a, int64_t factor_cols);
int64_t detect_tiny_suffix_512_cuda(torch::Tensor a);
std::vector<torch::Tensor> qr_2048_cuda(torch::Tensor a);
std::vector<torch::Tensor> qr_4096_cuda(torch::Tensor a);
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("qr_small", &qr_small_cuda, "batched Householder QR");
    m.def("qr_small_prefix", &qr_small_prefix_cuda, "batched prefix Householder QR");
    m.def("detect_tiny_suffix_512", &detect_tiny_suffix_512_cuda, "certified n512 tiny-suffix detector");
    m.def("qr_2048", &qr_2048_cuda, "batched n=2048 Householder QR");
    m.def("qr_4096", &qr_4096_cuda, "batched n=4096 Householder QR");
}
"""
CUDA_SRC = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <cublasLt.h>
#include <cublas_v2.h>
#include <cuda_runtime.h>
namespace {
constexpr int kN = 32; constexpr int kLD = 33; constexpr int kThreads = 32; constexpr int kN176 = 176; constexpr int kLD176Resident = 177; constexpr int kPanelThreads176 = 256; constexpr int kN352 = 352;
constexpr int kN512 = 512; constexpr int kN1024 = 1024; constexpr int kN2048 = 2048; constexpr int kN4096 = 4096; constexpr int kPanel352 = 8; constexpr int kPanelThreads352 = 256; constexpr int kTileUpdate352 = 32;
constexpr int kPanel512 = 16; constexpr int kPanelThreads512 = 128; constexpr int kPanel1024 = 16; constexpr int kPanelThreads1024 = 512; constexpr int kPanel2048 = 16; constexpr int kPanelThreads2048 = 512;
constexpr int kPanel4096 = 8; constexpr int kPanelThreads4096 = 256; constexpr int kPanel4096Late = 16; constexpr int kPanelThreads4096Late = 512; constexpr int kPanel4096LateMaxRows = 3584;
#define SET_KERNEL_ATTRS(kernel, shared_bytes)                                      \
    do {                                                                           \
        C10_CUDA_CHECK(cudaFuncSetAttribute(                                       \
            kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, shared_bytes));   \
        C10_CUDA_CHECK(cudaFuncSetAttribute(                                       \
            kernel, cudaFuncAttributePreferredSharedMemoryCarveout, 100));         \
    } while (0)
#define LT_CHECK(expr, message) TORCH_CHECK((expr) == CUBLAS_STATUS_SUCCESS, message)
#define LT_SET_LAYOUT(layout, attr, value, message) \
    LT_CHECK(cublasLtMatrixLayoutSetAttribute(layout, attr, &(value), sizeof(value)), message)
inline void cublas_strided_nt( float* c, const float* a, const float* b, int m, int n, int k, int lda, long long stride_a, int ldb, long long stride_b, int ldc, long long stride_c, int batch, bool allow_tf32, const char* message) {
    cublasHandle_t handle = at::cuda::getCurrentCUDABlasHandle(); const float alpha = 1.0f, beta = 0.0f; const cublasComputeType_t compute_type = allow_tf32 ? CUBLAS_COMPUTE_32F_FAST_TF32 : CUBLAS_COMPUTE_32F;
    const cublasGemmAlgo_t algo = allow_tf32 ? CUBLAS_GEMM_DEFAULT_TENSOR_OP : CUBLAS_GEMM_DEFAULT;
    cublasStatus_t status = cublasGemmStridedBatchedEx(
        handle, CUBLAS_OP_N, CUBLAS_OP_T, m, n, k, &alpha, a, CUDA_R_32F, lda, stride_a,
        b, CUDA_R_32F, ldb, stride_b, &beta, c, CUDA_R_32F, ldc, stride_c, batch, compute_type, algo);
    TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, message);
}
inline void cublas_leaf0_gram_from_macro( float* gram, long long gram_stride0, const float* v_macro, long long v_stride0, int macro_cols, int active_rows, int batch, bool allow_tf32) {
    cublas_strided_nt(gram, v_macro, v_macro, kPanel512, kPanel512, active_rows, macro_cols, v_stride0,
                      macro_cols, v_stride0, kPanel512, gram_stride0, batch, allow_tf32, "leaf0 Gram cuBLAS call failed");
}
inline void cublas_leaf1_cross_gram_from_macro( float* s_macro, long long s_stride0, const float* v_macro, long long v_stride0, int macro_cols, int active_rows, int batch, bool allow_tf32) {
    cublas_strided_nt(s_macro, v_macro + kPanel512, v_macro, kPanel512, macro_cols, active_rows, macro_cols, v_stride0,
                      macro_cols, v_stride0, kPanel512, s_stride0, batch, allow_tf32, "leaf1 cross Gram cuBLAS call failed");
}
inline void cublas_w_from_vt_c( float* w, const float* c_tail, const float* v, long long v_stride0, int matrix_n, int panel_cols, int v_ld_cols, int active_rows, int trailing_cols, int batch, bool allow_tf32) {
    cublas_strided_nt(w, c_tail, v, trailing_cols, panel_cols, active_rows, matrix_n, static_cast<long long>(matrix_n) * matrix_n,
                      v_ld_cols, v_stride0, trailing_cols, static_cast<long long>(panel_cols) * trailing_cols,
                      batch, allow_tf32, "W=V^T C cuBLAS call failed");
}
inline void cublaslt_tail_update_out_of_place( float* d_tail, const float* c_tail, const float* v_macro, long long v_stride0, const float* z, int n, int macro_cols, int active_rows, int trailing_cols, int batch, bool allow_tf32, void* workspace, size_t workspace_bytes) {
    cublasLtHandle_t handle = at::cuda::getCurrentCUDABlasLtHandle();
    const cublasComputeType_t compute_type = allow_tf32
        ? CUBLAS_COMPUTE_32F_FAST_TF32
        : CUBLAS_COMPUTE_32F;
    cublasLtMatmulDesc_t op_desc = nullptr; cublasLtMatrixLayout_t a_desc = nullptr; cublasLtMatrixLayout_t b_desc = nullptr; cublasLtMatrixLayout_t c_desc = nullptr; cublasLtMatrixLayout_t d_desc = nullptr; cublasLtMatmulPreference_t pref = nullptr;
    LT_CHECK(cublasLtMatmulDescCreate(&op_desc, compute_type, CUDA_R_32F),
             "tail update cuBLASLt desc create failed");
    cublasOperation_t trans = CUBLAS_OP_N;
    LT_CHECK(cublasLtMatmulDescSetAttribute(op_desc, CUBLASLT_MATMUL_DESC_TRANSA, &trans, sizeof(trans)),
             "tail update cuBLASLt transa set failed");
    LT_CHECK(cublasLtMatmulDescSetAttribute(op_desc, CUBLASLT_MATMUL_DESC_TRANSB, &trans, sizeof(trans)),
             "tail update cuBLASLt transb set failed");
    LT_CHECK(cublasLtMatrixLayoutCreate(&a_desc, CUDA_R_32F, trailing_cols, macro_cols, trailing_cols),
             "tail update cuBLASLt A layout create failed");
    LT_CHECK(cublasLtMatrixLayoutCreate(&b_desc, CUDA_R_32F, macro_cols, active_rows, macro_cols),
             "tail update cuBLASLt B layout create failed");
    LT_CHECK(cublasLtMatrixLayoutCreate(&c_desc, CUDA_R_32F, trailing_cols, active_rows, n),
             "tail update cuBLASLt C layout create failed");
    LT_CHECK(cublasLtMatrixLayoutCreate(&d_desc, CUDA_R_32F, trailing_cols, active_rows, n),
             "tail update cuBLASLt D layout create failed");
    const int batch_count = batch; const int64_t a_stride = static_cast<int64_t>(macro_cols) * trailing_cols; const int64_t b_stride = static_cast<int64_t>(v_stride0); const int64_t cd_stride = static_cast<int64_t>(n) * n;
    LT_SET_LAYOUT(a_desc, CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT, batch_count,
                  "tail update cuBLASLt A batch set failed");
    LT_SET_LAYOUT(b_desc, CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT, batch_count,
                  "tail update cuBLASLt B batch set failed");
    LT_SET_LAYOUT(c_desc, CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT, batch_count,
                  "tail update cuBLASLt C batch set failed");
    LT_SET_LAYOUT(d_desc, CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT, batch_count,
                  "tail update cuBLASLt D batch set failed");
    LT_SET_LAYOUT(a_desc, CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET, a_stride,
                  "tail update cuBLASLt A stride set failed");
    LT_SET_LAYOUT(b_desc, CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET, b_stride,
                  "tail update cuBLASLt B stride set failed");
    LT_SET_LAYOUT(c_desc, CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET, cd_stride,
                  "tail update cuBLASLt C stride set failed");
    LT_SET_LAYOUT(d_desc, CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET, cd_stride,
                  "tail update cuBLASLt D stride set failed");
    LT_CHECK(cublasLtMatmulPreferenceCreate(&pref), "tail update cuBLASLt preference create failed");
    LT_CHECK(cublasLtMatmulPreferenceSetAttribute( pref, CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES, &workspace_bytes, sizeof(workspace_bytes)), "tail update cuBLASLt workspace set failed"); cublasLtMatmulHeuristicResult_t heuristic; int returned = 0;
    LT_CHECK(cublasLtMatmulAlgoGetHeuristic( handle, op_desc, a_desc, b_desc, c_desc, d_desc, pref, 1, &heuristic, &returned), "tail update cuBLASLt heuristic query failed");
    TORCH_CHECK(returned > 0, "tail update cuBLASLt returned no heuristic");
    const float alpha = -1.0f; const float beta = 1.0f;
    const cublasStatus_t status = cublasLtMatmul(
        handle,
        op_desc,
        &alpha,
        z,
        a_desc,
        v_macro,
        b_desc,
        &beta,
        c_tail,
        c_desc,
        d_tail,
        d_desc,
        &heuristic.algo,
        workspace,
        workspace_bytes,
        nullptr);
    TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, "tail update cuBLASLt matmul failed");
}
__device__ __forceinline__ float warp_sum(float value) {
    value += __shfl_down_sync(0xffffffff, value, 16); value += __shfl_down_sync(0xffffffff, value, 8); value += __shfl_down_sync(0xffffffff, value, 4); value += __shfl_down_sync(0xffffffff, value, 2);
    value += __shfl_down_sync(0xffffffff, value, 1);
    return value;
}
__device__ __forceinline__ float4 ptx_ld_global_v4_f32(const float* ptr) {
    float4 value;
    asm volatile( "ld.global.v4.f32 {%0, %1, %2, %3}, [%4];" : "=f"(value.x), "=f"(value.y), "=f"(value.z), "=f"(value.w) : "l"(ptr));
    return value;
}
__device__ __forceinline__ void ptx_st_global_v4_f32(float* ptr, float4 value) {
    asm volatile( "st.global.v4.f32 [%0], {%1, %2, %3, %4};" : : "l"(ptr), "f"(value.x), "f"(value.y), "f"(value.z), "f"(value.w));
}
template <int Threads>
__global__ __launch_bounds__(Threads, 4) void copy_matrix_v4_kernel( const float* __restrict__ a, float* __restrict__ h, long long total) {
    const long long vec_total = total >> 2;
    for (long long idx = static_cast<long long>(blockIdx.x) * Threads + threadIdx.x;
         idx < vec_total;
         idx += static_cast<long long>(gridDim.x) * Threads) {
        const long long base = idx << 2; const float4 vals = ptx_ld_global_v4_f32(a + base); ptx_st_global_v4_f32(h + base, vals);
    }
}
template <int N, int Cols, int Threads>
__global__ __launch_bounds__(Threads, 4) void copy_first_cols_v4_kernel( const float* __restrict__ a, float* __restrict__ h, int batch) {
    constexpr int VecCols = Cols / 4; const long long total = static_cast<long long>(batch) * N * VecCols;
    for (long long idx = static_cast<long long>(blockIdx.x) * Threads + threadIdx.x;
         idx < total;
         idx += static_cast<long long>(gridDim.x) * Threads) {
        const int vec_col = static_cast<int>(idx % VecCols); const long long row_tmp = idx / VecCols; const int row = static_cast<int>(row_tmp % N); const int b = static_cast<int>(row_tmp / N);
        const long long base =
            (static_cast<long long>(b) * N + row) * N + vec_col * 4;
        const float4 vals = ptx_ld_global_v4_f32(a + base); ptx_st_global_v4_f32(h + base, vals);
    }
}
template <int N, int Threads>
__global__ __launch_bounds__(Threads, 4) void zero_suffix_cols_v4_kernel( float* __restrict__ h, int start_col, int batch) {
    const int vec_cols = (N - start_col) / 4; const long long total = static_cast<long long>(batch) * N * vec_cols;
    for (long long idx = static_cast<long long>(blockIdx.x) * Threads + threadIdx.x;
         idx < total;
         idx += static_cast<long long>(gridDim.x) * Threads) {
        const int vec_col = static_cast<int>(idx % vec_cols); const long long row_tmp = idx / vec_cols; const int row = static_cast<int>(row_tmp % N); const int b = static_cast<int>(row_tmp / N);
        float* dst =
            h + (static_cast<long long>(b) * N + row) * N + start_col + vec_col * 4;
        ptx_st_global_v4_f32(dst, make_float4(0.0f, 0.0f, 0.0f, 0.0f));
    }
}
__device__ __forceinline__ unsigned int warp_max_u32(unsigned int value) {
    value = max(value, __shfl_down_sync(0xffffffff, value, 16)); value = max(value, __shfl_down_sync(0xffffffff, value, 8)); value = max(value, __shfl_down_sync(0xffffffff, value, 4));
    value = max(value, __shfl_down_sync(0xffffffff, value, 2)); value = max(value, __shfl_down_sync(0xffffffff, value, 1));
    return value;
}
template <int N, int Threads>
__global__ __launch_bounds__(Threads, 4) void suffix_sample_reject_kernel( const float* __restrict__ a, int* __restrict__ reject) {
    constexpr int Warps = Threads / 32; __shared__ unsigned int warp_maxima[2 * Warps]; const int b = blockIdx.x; const float* a_b = a + static_cast<long long>(b) * N * N; const int tid = threadIdx.x;
    const int row0 = (tid * 73) & (N - 1); const int col0 = tid & 15; const int row1 = (tid * 37) & (N - 1); const int col1 = (3 * N) / 4 + ((tid * 29) & (N / 4 - 1));
    unsigned int prefix = __float_as_uint(fabsf(a_b[row0 * N + col0])); unsigned int tail = __float_as_uint(fabsf(a_b[row1 * N + col1])); prefix = warp_max_u32(prefix); tail = warp_max_u32(tail); const int lane = tid & 31;
    const int warp = tid >> 5;
    if (lane == 0) {
        warp_maxima[warp] = prefix; warp_maxima[Warps + warp] = tail;
    }
    __syncthreads();
    if (warp == 0) {
        unsigned int p = (lane < Warps) ? warp_maxima[lane] : 0; unsigned int t = (lane < Warps) ? warp_maxima[Warps + lane] : 0; p = warp_max_u32(p); t = warp_max_u32(t);
        if (lane == 0) {
            const float pv = __uint_as_float(p); const float tv = __uint_as_float(t);
            if (pv > 0.0f && tv > 1.0e-3f * pv) {
                atomicExch(reject, 1);
            }
        }
    }
}
template <int N, int Panel, int Threads>
__global__ __launch_bounds__(Threads, 4) void suffix_factor_cols_kernel( const float* __restrict__ a, int* __restrict__ factors, int k0, int k1, int k2, int k3, int k4) {
    constexpr int Warps = Threads / 32; constexpr int RowVecs = N / 4; constexpr int MatrixVecs = N * RowVecs; __shared__ unsigned int warp_maxima[6 * Warps]; const int b = blockIdx.x;
    const float* a_b = a + static_cast<long long>(b) * N * N; const int tid = threadIdx.x; unsigned int local0 = 0; unsigned int local1 = 0; unsigned int local2 = 0; unsigned int local3 = 0; unsigned int local4 = 0;
    unsigned int local5 = 0;
    for (int vec = tid; vec < MatrixVecs; vec += Threads) {
        const int col4 = (vec % RowVecs) * 4; const float4 vals = ptx_ld_global_v4_f32(a_b + static_cast<long long>(vec) * 4); unsigned int v = __float_as_uint(fabsf(vals.x)); v = max(v, __float_as_uint(fabsf(vals.y)));
        v = max(v, __float_as_uint(fabsf(vals.z))); v = max(v, __float_as_uint(fabsf(vals.w))); local0 = max(local0, v);
        if (col4 >= k0) local1 = max(local1, v);
        if (col4 >= k1) local2 = max(local2, v);
        if (col4 >= k2) local3 = max(local3, v);
        if (col4 >= k3) local4 = max(local4, v);
        if (col4 >= k4) local5 = max(local5, v);
    }
    local0 = warp_max_u32(local0); local1 = warp_max_u32(local1); local2 = warp_max_u32(local2); local3 = warp_max_u32(local3); local4 = warp_max_u32(local4); local5 = warp_max_u32(local5); const int lane = tid & 31;
    const int warp = tid >> 5;
    if (lane == 0) {
        warp_maxima[0 * Warps + warp] = local0; warp_maxima[1 * Warps + warp] = local1; warp_maxima[2 * Warps + warp] = local2; warp_maxima[3 * Warps + warp] = local3; warp_maxima[4 * Warps + warp] = local4;
        warp_maxima[5 * Warps + warp] = local5;
    }
    __syncthreads();
    if (warp == 0) {
        unsigned int block0 = (lane < Warps) ? warp_maxima[0 * Warps + lane] : 0; unsigned int block1 = (lane < Warps) ? warp_maxima[1 * Warps + lane] : 0;
        unsigned int block2 = (lane < Warps) ? warp_maxima[2 * Warps + lane] : 0; unsigned int block3 = (lane < Warps) ? warp_maxima[3 * Warps + lane] : 0;
        unsigned int block4 = (lane < Warps) ? warp_maxima[4 * Warps + lane] : 0; unsigned int block5 = (lane < Warps) ? warp_maxima[5 * Warps + lane] : 0; block0 = warp_max_u32(block0); block1 = warp_max_u32(block1);
        block2 = warp_max_u32(block2); block3 = warp_max_u32(block3); block4 = warp_max_u32(block4); block5 = warp_max_u32(block5);
        if (lane == 0) {
            constexpr float eps32 = 1.1920928955078125e-7f; constexpr float route_budget = 6.0f; const float all = __uint_as_float(block0); const float tail0 = __uint_as_float(block1); const float tail1 = __uint_as_float(block2);
            const float tail2 = __uint_as_float(block3); const float tail3 = __uint_as_float(block4); const float tail4 = __uint_as_float(block5); int matrix_cols = 0;
            if (all == 0.0f) {
                matrix_cols = 2 * Panel;
            } else {
                const float limit = route_budget * eps32 * all;
                if (tail0 <= limit) matrix_cols = k0;
                else if (tail1 <= limit) matrix_cols = k1;
                else if (tail2 <= limit) matrix_cols = k2;
                else if (tail3 <= limit) matrix_cols = k3;
                else if (tail4 <= limit) matrix_cols = k4;
            }
            factors[b] = matrix_cols;
        }
    }
}
template <int Threads>
__global__ __launch_bounds__(Threads, 1) void reduce_factor_cols_kernel( const int* __restrict__ factors, int* __restrict__ result, int batch) {
    constexpr int Warps = Threads / 32; __shared__ int warp_values[2 * Warps]; const int tid = threadIdx.x; int local_max = 0; int local_bad = 0;
    for (int idx = tid; idx < batch; idx += Threads) {
        const int value = factors[idx];
        if (value == 0) {
            local_bad = 1;
        } else {
            local_max = max(local_max, value);
        }
    }
    local_max = max(local_max, __shfl_down_sync(0xffffffff, local_max, 16)); local_max = max(local_max, __shfl_down_sync(0xffffffff, local_max, 8)); local_max = max(local_max, __shfl_down_sync(0xffffffff, local_max, 4));
    local_max = max(local_max, __shfl_down_sync(0xffffffff, local_max, 2)); local_max = max(local_max, __shfl_down_sync(0xffffffff, local_max, 1)); local_bad = max(local_bad, __shfl_down_sync(0xffffffff, local_bad, 16));
    local_bad = max(local_bad, __shfl_down_sync(0xffffffff, local_bad, 8)); local_bad = max(local_bad, __shfl_down_sync(0xffffffff, local_bad, 4)); local_bad = max(local_bad, __shfl_down_sync(0xffffffff, local_bad, 2));
    local_bad = max(local_bad, __shfl_down_sync(0xffffffff, local_bad, 1)); const int lane = tid & 31; const int warp = tid >> 5;
    if (lane == 0) {
        warp_values[warp] = local_max; warp_values[Warps + warp] = local_bad;
    }
    __syncthreads();
    if (warp == 0) {
        int block_max = (lane < Warps) ? warp_values[lane] : 0; int block_bad = (lane < Warps) ? warp_values[Warps + lane] : 0; block_max = max(block_max, __shfl_down_sync(0xffffffff, block_max, 16));
        block_max = max(block_max, __shfl_down_sync(0xffffffff, block_max, 8)); block_max = max(block_max, __shfl_down_sync(0xffffffff, block_max, 4)); block_max = max(block_max, __shfl_down_sync(0xffffffff, block_max, 2));
        block_max = max(block_max, __shfl_down_sync(0xffffffff, block_max, 1)); block_bad = max(block_bad, __shfl_down_sync(0xffffffff, block_bad, 16)); block_bad = max(block_bad, __shfl_down_sync(0xffffffff, block_bad, 8));
        block_bad = max(block_bad, __shfl_down_sync(0xffffffff, block_bad, 4)); block_bad = max(block_bad, __shfl_down_sync(0xffffffff, block_bad, 2)); block_bad = max(block_bad, __shfl_down_sync(0xffffffff, block_bad, 1));
        if (lane == 0) {
            result[0] = (block_bad != 0) ? 0 : block_max;
        }
    }
}
template <int N, int Threads>
__global__ __launch_bounds__(Threads, 4) void zero_tau_suffix_kernel( float* __restrict__ tau, int factor_cols, int batch) {
    const int suffix = N - factor_cols; const long long total = static_cast<long long>(batch) * suffix;
    for (long long idx = static_cast<long long>(blockIdx.x) * Threads + threadIdx.x;
         idx < total;
         idx += static_cast<long long>(gridDim.x) * Threads) {
        const int b = static_cast<int>(idx / suffix); const int j = factor_cols + static_cast<int>(idx - static_cast<long long>(b) * suffix); tau[static_cast<long long>(b) * N + j] = 0.0f;
    }
}
template <int Threads>
__device__ __forceinline__ float block_sum_thread0(float value, float* work) {
    constexpr int Warps = Threads / 32; const int lane = threadIdx.x & 31; const int warp = threadIdx.x >> 5; value = warp_sum(value);
    if (lane == 0) {
        work[warp] = value;
    }
    __syncthreads(); float total = 0.0f;
    if (threadIdx.x < 32) {
        total = (threadIdx.x < Warps) ? work[threadIdx.x] : 0.0f; total = warp_sum(total);
    }
    return total;
}
__global__ __launch_bounds__(kThreads, 16) void qr32_kernel( const float* __restrict__ a, float* __restrict__ h, float* __restrict__ tau) {
    const int b = blockIdx.x; const int tid = threadIdx.x; const float* a_b = a + static_cast<long long>(b) * kN * kN; float* h_b = h + static_cast<long long>(b) * kN * kN;
    float* tau_b = tau + static_cast<long long>(b) * kN; __shared__ float s[kN * kLD]; constexpr int RowVecs = kN / 4;
    for (int idx = tid; idx < kN * RowVecs; idx += kThreads) {
        const int row = idx / RowVecs; const int col = (idx - row * RowVecs) * 4; const float4 vals = ptx_ld_global_v4_f32(a_b + row * kN + col); s[row * kLD + col + 0] = vals.x; s[row * kLD + col + 1] = vals.y;
        s[row * kLD + col + 2] = vals.z; s[row * kLD + col + 3] = vals.w;
    }
    __syncwarp();
    #pragma unroll 32
    for (int k = 0; k < kN; ++k) {
        float local = 0.0f;
        for (int i = k + 1 + tid; i < kN; i += kThreads) {
            const float x = s[i * kLD + k]; local = fmaf(x, x, local);
        }
        const float sigma = warp_sum(local); float tau_k = 0.0f; float inv = 0.0f;
        if (tid == 0) {
            const float alpha = s[k * kLD + k];
            if (sigma == 0.0f) {
                tau_b[k] = 0.0f;
            } else {
                const float norm = sqrtf(fmaf(alpha, alpha, sigma)); const float beta = (alpha < 0.0f) ? norm : -norm; tau_k = (beta - alpha) / beta; inv = 1.0f / (alpha - beta); tau_b[k] = tau_k; s[k * kLD + k] = beta;
            }
        }
        tau_k = __shfl_sync(0xffffffff, tau_k, 0); inv = __shfl_sync(0xffffffff, inv, 0);
        for (int i = k + 1 + tid; i < kN; i += kThreads) {
            s[i * kLD + k] *= inv;
        }
        __syncwarp();
        for (int j = k + 1 + tid; j < kN; j += kThreads) {
            float dot = s[k * kLD + j];
            #pragma unroll 4
            for (int i = k + 1; i < kN; ++i) {
                dot = fmaf(s[i * kLD + k], s[i * kLD + j], dot);
            }
            dot *= tau_k; s[k * kLD + j] -= dot;
            #pragma unroll 4
            for (int i = k + 1; i < kN; ++i) {
                s[i * kLD + j] = fmaf(-s[i * kLD + k], dot, s[i * kLD + j]);
            }
        }
        __syncwarp();
    }
    for (int idx = tid; idx < kN * RowVecs; idx += kThreads) {
        const int row = idx / RowVecs; const int col = (idx - row * RowVecs) * 4;
        ptx_st_global_v4_f32( h_b + row * kN + col, make_float4( s[row * kLD + col + 0], s[row * kLD + col + 1], s[row * kLD + col + 2], s[row * kLD + col + 3]));
    }
}
template <int Threads>
__global__ __launch_bounds__(Threads, 1) void qr176_resident_kernel( const float* __restrict__ a, float* __restrict__ h, float* __restrict__ tau) {
    constexpr int Warps = Threads / 32; const int b = blockIdx.x; const int tid = threadIdx.x; const int lane = tid & 31; const int warp = tid >> 5; const float* a_b = a + static_cast<long long>(b) * kN176 * kN176;
    float* h_b = h + static_cast<long long>(b) * kN176 * kN176; float* tau_b = tau + static_cast<long long>(b) * kN176; extern __shared__ float s[]; __shared__ float reduce[Warps]; __shared__ float params[2];
    for (int idx = tid; idx < kN176 * kN176; idx += Threads) {
        const int row = idx / kN176; const int col = idx - row * kN176; s[row * kLD176Resident + col] = a_b[idx];
    }
    __syncthreads();
    for (int k = 0; k < kN176; ++k) {
        float local = 0.0f;
        for (int i = k + 1 + tid; i < kN176; i += Threads) {
            const float x = s[i * kLD176Resident + k]; local = fmaf(x, x, local);
        }
        const float sigma = block_sum_thread0<Threads>(local, reduce);
        if (tid == 0) {
            const float alpha = s[k * kLD176Resident + k];
            if (sigma == 0.0f) {
                tau_b[k] = 0.0f; params[0] = 0.0f; params[1] = 0.0f;
            } else {
                const float norm = sqrtf(fmaf(alpha, alpha, sigma)); const float beta = (alpha < 0.0f) ? norm : -norm; const float tau_k = (beta - alpha) / beta; tau_b[k] = tau_k; s[k * kLD176Resident + k] = beta; params[0] = tau_k;
                params[1] = 1.0f / (alpha - beta);
            }
        }
        __syncthreads(); const float tau_k = params[0]; const float inv = params[1];
        for (int i = k + 1 + tid; i < kN176; i += Threads) {
            s[i * kLD176Resident + k] *= inv;
        }
        __syncthreads();
        for (int j_base = k + 1; j_base < kN176; j_base += Warps) {
            const int j = j_base + warp; float dot = 0.0f;
            if (j < kN176) {
                dot = (lane == 0) ? s[k * kLD176Resident + j] : 0.0f;
                for (int i = k + 1 + lane; i < kN176; i += 32) {
                    dot = fmaf( s[i * kLD176Resident + k], s[i * kLD176Resident + j], dot);
                }
            }
            dot = warp_sum(dot); dot = __shfl_sync(0xffffffff, dot, 0) * tau_k;
            if (j < kN176) {
                if (lane == 0) {
                    s[k * kLD176Resident + j] -= dot;
                }
                for (int i = k + 1 + lane; i < kN176; i += 32) {
                    const int offset = i * kLD176Resident + j; s[offset] = fmaf(-s[i * kLD176Resident + k], dot, s[offset]);
                }
            }
        }
        __syncthreads();
    }
    for (int idx = tid; idx < kN176 * kN176; idx += Threads) {
        const int row = idx / kN176; const int col = idx - row * kN176; h_b[idx] = s[row * kLD176Resident + col];
    }
}
template <
    int N,
    int Panel,
    int Threads,
    bool PackV,
    bool DynamicStride = false,
    bool BuildT = true,
    bool FusedT = true,
    bool WriteMacroV = false>
__global__ __launch_bounds__(Threads, 1) void qr_panel_cached_kernel( float* __restrict__ h, float* __restrict__ tau, float* __restrict__ t_scratch, float* __restrict__ v_pack, long long v_stride0, int panel_start, float* __restrict__ v_macro = nullptr, long long v_macro_stride0 = 0, int macro_cols = 0, int macro_row_offset = 0, int macro_col_offset = 0) {
    const int b = blockIdx.x; const int tid = threadIdx.x; const long long matrix_stride = static_cast<long long>(N) * N; float* h_b = h + static_cast<long long>(b) * matrix_stride;
    float* tau_b = tau + static_cast<long long>(b) * N; float* t_b = nullptr;
    if constexpr (BuildT) {
        t_b = t_scratch + static_cast<long long>(b) * Panel * Panel;
    }
    float* v_b = nullptr;
    if constexpr (PackV) {
        v_b = v_pack + static_cast<long long>(b) * v_stride0;
    }
    float* v_macro_b = nullptr;
    if constexpr (WriteMacroV) {
        v_macro_b = v_macro + static_cast<long long>(b) * v_macro_stride0;
    }
    extern __shared__ float smem[]; float* panel = smem; const int active_rows = N - panel_start; const int panel_stride = DynamicStride ? (active_rows + 1) : (N + 1); float* work = smem + panel_stride * Panel;
    constexpr int Warps = Threads / 32; constexpr int PanelVec = Panel / 4;
    for (int idx = tid; idx < active_rows * PanelVec; idx += Threads) {
        const int rel = idx / PanelVec; const int t = (idx - rel * PanelVec) * 4;
        const float4 vals = ptx_ld_global_v4_f32(
            h_b + static_cast<long long>(panel_start + rel) * N + panel_start + t);
        panel[(t + 0) * panel_stride + rel] = vals.x; panel[(t + 1) * panel_stride + rel] = vals.y; panel[(t + 2) * panel_stride + rel] = vals.z; panel[(t + 3) * panel_stride + rel] = vals.w;
    }
    __syncthreads();
    if constexpr (BuildT) {
        for (int idx = tid; idx < Panel * Panel; idx += Threads) {
            t_b[idx] = 0.0f;
        }
    }
    __syncthreads();
    #pragma unroll
    for (int col = 0; col < Panel; ++col) {
        const int k = panel_start + col; float local = 0.0f;
        for (int rel = col + 1 + tid; rel < active_rows; rel += Threads) {
            const float x = panel[col * panel_stride + rel]; local = fmaf(x, x, local);
        }
        const float sigma = block_sum_thread0<Threads>(local, work);
        if (tid == 0) {
            const float alpha = panel[col * panel_stride + col];
            if (sigma == 0.0f) {
                tau_b[k] = 0.0f; work[0] = 0.0f; work[1] = 0.0f;
            } else {
                const float norm = sqrtf(fmaf(alpha, alpha, sigma)); const float beta = (alpha < 0.0f) ? norm : -norm; const float tau_k = (beta - alpha) / beta; tau_b[k] = tau_k; panel[col * panel_stride + col] = beta;
                work[0] = tau_k; work[1] = 1.0f / (alpha - beta);
            }
        }
        __syncthreads(); const float tau_k = work[0]; const float inv = work[1];
        if constexpr (BuildT && FusedT) {
            const int lane = tid & 31; const int warp = tid >> 5; float t_dots[Panel];
            #pragma unroll
            for (int prev = 0; prev < Panel; ++prev) {
                t_dots[prev] = 0.0f;
            }
            for (int rel = col + 1 + tid; rel < active_rows; rel += Threads) {
                const float v_cur = panel[col * panel_stride + rel] * inv; panel[col * panel_stride + rel] = v_cur;
                #pragma unroll
                for (int prev = 0; prev < Panel; ++prev) {
                    if (prev < col) {
                        t_dots[prev] = fmaf(
                            panel[prev * panel_stride + rel],
                            v_cur,
                            t_dots[prev]);
                    }
                }
            }
            if (tid == 0) {
                #pragma unroll
                for (int prev = 0; prev < Panel; ++prev) {
                    if (prev < col) {
                        t_dots[prev] += panel[prev * panel_stride + col];
                    }
                }
            }
            #pragma unroll
            for (int prev = 0; prev < Panel; ++prev) {
                const float partial = (prev < col) ? warp_sum(t_dots[prev]) : 0.0f;
                if (lane == 0) {
                    work[warp * Panel + prev] = partial;
                }
            }
            __syncthreads();
            if (tid < Panel && tid < col) {
                float dot = 0.0f;
                #pragma unroll
                for (int w = 0; w < Warps; ++w) {
                    dot += work[w * Panel + tid];
                }
                work[Warps * Panel + tid] = -tau_k * dot;
            }
            __syncthreads();
            if (tid < col) {
                float accum = 0.0f;
                #pragma unroll
                for (int inner = 0; inner < Panel; ++inner) {
                    if (inner < col) {
                        accum = fmaf( t_b[tid * Panel + inner], work[Warps * Panel + inner], accum);
                    }
                }
                t_b[tid * Panel + col] = accum;
            }
            if (tid == col) {
                t_b[col * Panel + col] = tau_k;
            }
        } else {
            for (int rel = col + 1 + tid; rel < active_rows; rel += Threads) {
                panel[col * panel_stride + rel] *= inv;
            }
        }
        __syncthreads(); const int lane = tid & 31; const int warp = tid >> 5;
        for (int j_base = col + 1; j_base < Panel; j_base += Warps) {
            const int j = j_base + warp; float local_dot = 0.0f;
            if (j < Panel) {
                local_dot = (lane == 0) ? panel[j * panel_stride + col] : 0.0f;
                for (int rel = col + 1 + lane; rel < active_rows; rel += 32) {
                    local_dot = fmaf( panel[col * panel_stride + rel], panel[j * panel_stride + rel], local_dot);
                }
            }
            float dot = warp_sum(local_dot); dot = __shfl_sync(0xffffffff, dot, 0);
            if (j < Panel) {
                dot *= tau_k;
                if (lane == 0) {
                    panel[j * panel_stride + col] -= dot;
                }
                for (int rel = col + 1 + lane; rel < active_rows; rel += 32) {
                    const int offset = j * panel_stride + rel;
                    panel[offset] = fmaf(
                        -panel[col * panel_stride + rel],
                        dot,
                        panel[offset]);
                }
            }
        }
        __syncthreads();
    }
    if constexpr (BuildT && !FusedT) {
        #pragma unroll
        for (int j = 0; j < Panel; ++j) {
            const float tau_j = tau_b[panel_start + j];
            #pragma unroll
            for (int i = 0; i < Panel; ++i) {
                if (i < j) {
                    float local = 0.0f;
                    for (int rel = j + 1 + tid; rel < active_rows; rel += Threads) {
                        local = fmaf( panel[i * panel_stride + rel], panel[j * panel_stride + rel], local);
                    }
                    if (tid == 0) {
                        local += panel[i * panel_stride + j];
                    }
                    const float dot = block_sum_thread0<Threads>(local, work);
                    if (tid == 0) {
                        work[Warps + i] = -tau_j * dot;
                    }
                    __syncthreads();
                }
            }
            if (tid < j) {
                float accum = 0.0f;
                #pragma unroll
                for (int inner = 0; inner < Panel; ++inner) {
                    if (inner < j) {
                        accum = fmaf( t_b[tid * Panel + inner], work[Warps + inner], accum);
                    }
                }
                t_b[tid * Panel + j] = accum;
            }
            if (tid == j) {
                t_b[j * Panel + j] = tau_j;
            }
            __syncthreads();
        }
    }
    if constexpr (WriteMacroV) {
        for (int idx = tid; idx < macro_row_offset * PanelVec; idx += Threads) {
            const int rel = idx / PanelVec; const int t = (idx - rel * PanelVec) * 4; ptx_st_global_v4_f32( v_macro_b + static_cast<long long>(rel) * macro_cols + macro_col_offset + t, make_float4(0.0f, 0.0f, 0.0f, 0.0f));
        }
    }
    for (int idx = tid; idx < active_rows * PanelVec; idx += Threads) {
        const int rel = idx / PanelVec; const int t = (idx - rel * PanelVec) * 4; const float h0 = panel[(t + 0) * panel_stride + rel]; const float h1 = panel[(t + 1) * panel_stride + rel];
        const float h2 = panel[(t + 2) * panel_stride + rel]; const float h3 = panel[(t + 3) * panel_stride + rel];
        ptx_st_global_v4_f32( h_b + static_cast<long long>(panel_start + rel) * N + panel_start + t, make_float4(h0, h1, h2, h3));
        if constexpr (PackV) {
            const float v0 = (rel == t + 0) ? 1.0f : ((rel > t + 0) ? h0 : 0.0f); const float v1 = (rel == t + 1) ? 1.0f : ((rel > t + 1) ? h1 : 0.0f); const float v2 = (rel == t + 2) ? 1.0f : ((rel > t + 2) ? h2 : 0.0f);
            const float v3 = (rel == t + 3) ? 1.0f : ((rel > t + 3) ? h3 : 0.0f); const float4 v_vals = make_float4(v0, v1, v2, v3); ptx_st_global_v4_f32( v_b + static_cast<long long>(rel) * Panel + t, v_vals);
            if constexpr (WriteMacroV) {
                ptx_st_global_v4_f32( v_macro_b + static_cast<long long>(macro_row_offset + rel) * macro_cols + macro_col_offset + t, v_vals);
            }
        } else if constexpr (WriteMacroV) {
            const float v0 = (rel == t + 0) ? 1.0f : ((rel > t + 0) ? h0 : 0.0f); const float v1 = (rel == t + 1) ? 1.0f : ((rel > t + 1) ? h1 : 0.0f); const float v2 = (rel == t + 2) ? 1.0f : ((rel > t + 2) ? h2 : 0.0f);
            const float v3 = (rel == t + 3) ? 1.0f : ((rel > t + 3) ? h3 : 0.0f); ptx_st_global_v4_f32( v_macro_b + static_cast<long long>(macro_row_offset + rel) * macro_cols + macro_col_offset + t, make_float4(v0, v1, v2, v3));
        }
    }
}
template <int N, int Panel, int TileCols, int Threads>
__global__ __launch_bounds__(Threads, 1) void qr_panel_tile_update_kernel( float* __restrict__ h, const float* __restrict__ t_scratch, int panel_start) {
    const int tile = blockIdx.x; const int b = blockIdx.y; const int tid = threadIdx.x; const int panel_end = panel_start + Panel; const int active_rows = N - panel_start; const int tile_start = panel_end + tile * TileCols;
    if (tile_start >= N) {
        return;
    }
    const int tile_cols = min(TileCols, N - tile_start); const long long matrix_stride = static_cast<long long>(N) * N; float* h_b = h + static_cast<long long>(b) * matrix_stride;
    const float* t_b = t_scratch + static_cast<long long>(b) * Panel * Panel; extern __shared__ float smem[]; float* c_tile = smem; float* v_panel = c_tile + active_rows * TileCols;
    float* w_tile = v_panel + active_rows * Panel; float* z_tile = w_tile + Panel * TileCols;
    for (int idx = tid; idx < active_rows * TileCols; idx += Threads) {
        const int rel = idx / TileCols; const int col = idx - rel * TileCols; float value = 0.0f;
        if (col < tile_cols) {
            value = h_b[
                static_cast<long long>(panel_start + rel) * N +
                tile_start + col];
        }
        c_tile[idx] = value;
    }
    for (int idx = tid; idx < active_rows * Panel; idx += Threads) {
        const int rel = idx / Panel; const int t = idx - rel * Panel; float value = 0.0f;
        if (rel == t) {
            value = 1.0f;
        } else if (rel > t) {
            value = h_b[
                static_cast<long long>(panel_start + rel) * N +
                panel_start + t];
        }
        v_panel[idx] = value;
    }
    __syncthreads();
    if (tid < TileCols) {
        const int col = tid; float w0 = 0.0f; float w1 = 0.0f; float w2 = 0.0f; float w3 = 0.0f; float w4 = 0.0f; float w5 = 0.0f; float w6 = 0.0f; float w7 = 0.0f;
        if (col < tile_cols) {
            #pragma unroll 24
            for (int rel = 0; rel < active_rows; ++rel) {
                const float c = c_tile[rel * TileCols + col]; const float* v = v_panel + rel * Panel; w0 = fmaf(v[0], c, w0); w1 = fmaf(v[1], c, w1); w2 = fmaf(v[2], c, w2); w3 = fmaf(v[3], c, w3); w4 = fmaf(v[4], c, w4);
                w5 = fmaf(v[5], c, w5); w6 = fmaf(v[6], c, w6); w7 = fmaf(v[7], c, w7);
            }
        }
        w_tile[0 * TileCols + col] = w0; w_tile[1 * TileCols + col] = w1; w_tile[2 * TileCols + col] = w2; w_tile[3 * TileCols + col] = w3; w_tile[4 * TileCols + col] = w4; w_tile[5 * TileCols + col] = w5;
        w_tile[6 * TileCols + col] = w6; w_tile[7 * TileCols + col] = w7; z_tile[0 * TileCols + col] = t_b[0] * w0; z_tile[1 * TileCols + col] = fmaf(t_b[1], w0, t_b[Panel + 1] * w1);
        z_tile[2 * TileCols + col] =
            fmaf(t_b[2], w0, fmaf(t_b[Panel + 2], w1, t_b[2 * Panel + 2] * w2));
        z_tile[3 * TileCols + col] =
            fmaf(t_b[3], w0, fmaf(t_b[Panel + 3], w1, fmaf(t_b[2 * Panel + 3], w2, t_b[3 * Panel + 3] * w3)));
        z_tile[4 * TileCols + col] =
            fmaf(t_b[4], w0, fmaf(t_b[Panel + 4], w1, fmaf(t_b[2 * Panel + 4], w2, fmaf(t_b[3 * Panel + 4], w3, t_b[4 * Panel + 4] * w4))));
        z_tile[5 * TileCols + col] =
            fmaf(t_b[5], w0, fmaf(t_b[Panel + 5], w1, fmaf(t_b[2 * Panel + 5], w2, fmaf(t_b[3 * Panel + 5], w3, fmaf(t_b[4 * Panel + 5], w4, t_b[5 * Panel + 5] * w5)))));
        z_tile[6 * TileCols + col] =
            fmaf(t_b[6], w0, fmaf(t_b[Panel + 6], w1, fmaf(t_b[2 * Panel + 6], w2, fmaf(t_b[3 * Panel + 6], w3, fmaf(t_b[4 * Panel + 6], w4, fmaf(t_b[5 * Panel + 6], w5, t_b[6 * Panel + 6] * w6))))));
        z_tile[7 * TileCols + col] =
            fmaf(t_b[7], w0, fmaf(t_b[Panel + 7], w1, fmaf(t_b[2 * Panel + 7], w2, fmaf(t_b[3 * Panel + 7], w3, fmaf(t_b[4 * Panel + 7], w4, fmaf(t_b[5 * Panel + 7], w5, fmaf(t_b[6 * Panel + 7], w6, t_b[7 * Panel + 7] * w7)))))));
    }
    __syncthreads();
    for (int idx = tid; idx < active_rows * TileCols; idx += Threads) {
        const int rel = idx / TileCols; const int col = idx - rel * TileCols;
        if (col < tile_cols) {
            const float* v = v_panel + rel * Panel; const float* z = z_tile + col; float delta = v[0] * z[0 * TileCols]; delta = fmaf(v[1], z[1 * TileCols], delta); delta = fmaf(v[2], z[2 * TileCols], delta);
            delta = fmaf(v[3], z[3 * TileCols], delta); delta = fmaf(v[4], z[4 * TileCols], delta); delta = fmaf(v[5], z[5 * TileCols], delta); delta = fmaf(v[6], z[6 * TileCols], delta); delta = fmaf(v[7], z[7 * TileCols], delta);
            c_tile[idx] -= delta;
        }
    }
    __syncthreads();
    for (int idx = tid; idx < active_rows * TileCols; idx += Threads) {
        const int rel = idx / TileCols; const int col = idx - rel * TileCols;
        if (col < tile_cols) {
            h_b[
                static_cast<long long>(panel_start + rel) * N +
                tile_start + col] = c_tile[idx];
        }
    }
}
template <int Panel, int Threads>
__global__ __launch_bounds__(Threads, 4) void apply_t_transpose_kernel( const float* __restrict__ t_scratch, const float* __restrict__ w, float* __restrict__ z, int trailing_cols) {
    const int b = blockIdx.y; const int col = blockIdx.x * Threads + threadIdx.x; const float* t_b = t_scratch + static_cast<long long>(b) * Panel * Panel;
    const float* w_b = w + static_cast<long long>(b) * Panel * trailing_cols; float* z_b = z + static_cast<long long>(b) * Panel * trailing_cols; __shared__ float t_shared[Panel * Panel];
    for (int idx = threadIdx.x; idx < Panel * Panel; idx += Threads) {
        t_shared[idx] = t_b[idx];
    }
    __syncthreads();
    if (col >= trailing_cols) {
        return;
    }
    float wv[Panel];
    #pragma unroll
    for (int i = 0; i < Panel; ++i) {
        wv[i] = w_b[i * trailing_cols + col];
    }
    #pragma unroll
    for (int row = 0; row < Panel; ++row) {
        float accum = 0.0f;
        #pragma unroll
        for (int inner = 0; inner < Panel; ++inner) {
            if (inner <= row) {
                accum = fmaf(t_shared[inner * Panel + row], wv[inner], accum);
            }
        }
        z_b[row * trailing_cols + col] = accum;
    }
}
template <int N, int Panel, int Threads>
__global__ __launch_bounds__(Threads, 4) void panel_rank_update_kernel( float* __restrict__ h, const float* __restrict__ v, long long v_stride0, int v_ld_cols, const float* __restrict__ z, int row_start, int col_start, int active_rows, int local_cols) {
    const int b = blockIdx.y; const int idx = blockIdx.x * Threads + threadIdx.x; const int total = active_rows * local_cols;
    if (idx >= total) {
        return;
    }
    const int col = idx % local_cols; const int row = idx / local_cols; const float* v_b = v + static_cast<long long>(b) * v_stride0 + static_cast<long long>(row) * v_ld_cols; const float* z_b = z + static_cast<long long>(b) * Panel * local_cols + col;
    float accum = 0.0f;
    #pragma unroll
    for (int k = 0; k < Panel; ++k) {
        accum = fmaf(v_b[k], z_b[k * local_cols], accum);
    }
    float* h_b = h + static_cast<long long>(b) * N * N + static_cast<long long>(row_start + row) * N + col_start + col;
    *h_b -= accum;
}
template <int Panel, int Threads>
__global__ __launch_bounds__(Threads, 4) void solve_inverse_wy_from_gram_kernel( const float* __restrict__ gram_scratch, const float* __restrict__ tau, const float* __restrict__ w, float* __restrict__ z, int n, int panel_start, int trailing_cols) {
    const int b = blockIdx.y; const int col = blockIdx.x * Threads + threadIdx.x; const float* g_b = gram_scratch + static_cast<long long>(b) * Panel * Panel;
    const float* tau_b = tau + static_cast<long long>(b) * n + panel_start; const float* w_b = w + static_cast<long long>(b) * Panel * trailing_cols; float* z_b = z + static_cast<long long>(b) * Panel * trailing_cols;
    __shared__ float g_shared[Panel * Panel]; __shared__ float tau_shared[Panel];
    for (int idx = threadIdx.x; idx < Panel * Panel; idx += Threads) {
        g_shared[idx] = g_b[idx];
    }
    for (int idx = threadIdx.x; idx < Panel; idx += Threads) {
        tau_shared[idx] = tau_b[idx];
    }
    __syncthreads();
    if (col >= trailing_cols) {
        return;
    }
    float zv[Panel];
    #pragma unroll
    for (int row = 0; row < Panel; ++row) {
        float accum = 0.0f;
        #pragma unroll
        for (int inner = 0; inner < Panel; ++inner) {
            if (inner < row) {
                accum = fmaf(g_shared[inner * Panel + row], zv[inner], accum);
            }
        }
        const float zi = tau_shared[row] * (w_b[row * trailing_cols + col] - accum); zv[row] = zi; z_b[row * trailing_cols + col] = zi;
    }
}
template <int Panel, int Block, int Threads>
__global__ __launch_bounds__(Threads, 2) void solve_inverse_wy_from_gram_blocked_kernel( const float* __restrict__ gram_scratch, const float* __restrict__ tau, const float* __restrict__ w, float* __restrict__ z, int n, int panel_start, int trailing_cols) {
    const int b = blockIdx.y; const int col = blockIdx.x * Threads + threadIdx.x; const float* g_b = gram_scratch + static_cast<long long>(b) * Panel * Panel;
    const float* tau_b = tau + static_cast<long long>(b) * n + panel_start; const float* w_b = w + static_cast<long long>(b) * Panel * trailing_cols; float* z_b = z + static_cast<long long>(b) * Panel * trailing_cols;
    __shared__ float g_shared[Panel * Panel]; __shared__ float tau_shared[Panel]; __shared__ float z_shared[Panel * Threads];
    for (int idx = threadIdx.x; idx < Panel * Panel; idx += Threads) {
        g_shared[idx] = g_b[idx];
    }
    for (int idx = threadIdx.x; idx < Panel; idx += Threads) {
        tau_shared[idx] = tau_b[idx];
    }
    __syncthreads();
    #pragma unroll
    for (int block_start = 0; block_start < Panel; block_start += Block) {
        float zv[Block];
        #pragma unroll
        for (int r = 0; r < Block; ++r) {
            zv[r] = 0.0f;
        }
        if (col < trailing_cols) {
            #pragma unroll
            for (int r = 0; r < Block; ++r) {
                const int row = block_start + r; float accum = 0.0f;
                for (int inner = 0; inner < Panel; ++inner) {
                    if (inner < block_start) {
                        accum = fmaf( g_shared[inner * Panel + row], z_shared[inner * Threads + threadIdx.x], accum);
                    }
                }
                #pragma unroll
                for (int inner = 0; inner < Block; ++inner) {
                    if (inner < r) {
                        accum = fmaf( g_shared[(block_start + inner) * Panel + row], zv[inner], accum);
                    }
                }
                const float zi = tau_shared[row] * (w_b[row * trailing_cols + col] - accum); zv[r] = zi; z_shared[row * Threads + threadIdx.x] = zi; z_b[row * trailing_cols + col] = zi;
            }
        }
        __syncthreads();
    }
}
template <int Panel, int Block, int Threads>
__global__ __launch_bounds__(Threads, 1) void solve_inverse_wy_from_gram_blocked_dynamic_kernel( const float* __restrict__ gram_scratch, const float* __restrict__ tau, const float* __restrict__ w, float* __restrict__ z, int n, int panel_start, int trailing_cols) {
    const int b = blockIdx.y; const int col = blockIdx.x * Threads + threadIdx.x; const float* g_b = gram_scratch + static_cast<long long>(b) * Panel * Panel;
    const float* tau_b = tau + static_cast<long long>(b) * n + panel_start; const float* w_b = w + static_cast<long long>(b) * Panel * trailing_cols; float* z_b = z + static_cast<long long>(b) * Panel * trailing_cols;
    extern __shared__ float smem[]; float* g_shared = smem; float* tau_shared = g_shared + Panel * Panel; float* z_shared = tau_shared + Panel;
    for (int idx = threadIdx.x; idx < Panel * Panel; idx += Threads) {
        g_shared[idx] = g_b[idx];
    }
    for (int idx = threadIdx.x; idx < Panel; idx += Threads) {
        tau_shared[idx] = tau_b[idx];
    }
    __syncthreads();
    #pragma unroll
    for (int block_start = 0; block_start < Panel; block_start += Block) {
        float zv[Block];
        #pragma unroll
        for (int r = 0; r < Block; ++r) {
            zv[r] = 0.0f;
        }
        if (col < trailing_cols) {
            #pragma unroll
            for (int r = 0; r < Block; ++r) {
                const int row = block_start + r; float accum = 0.0f;
                for (int inner = 0; inner < Panel; ++inner) {
                    if (inner < block_start) {
                        accum = fmaf( g_shared[inner * Panel + row], z_shared[inner * Threads + threadIdx.x], accum);
                    }
                }
                #pragma unroll
                for (int inner = 0; inner < Block; ++inner) {
                    if (inner < r) {
                        accum = fmaf( g_shared[(block_start + inner) * Panel + row], zv[inner], accum);
                    }
                }
                const float zi = tau_shared[row] * (w_b[row * trailing_cols + col] - accum); zv[r] = zi; z_shared[row * Threads + threadIdx.x] = zi; z_b[row * trailing_cols + col] = zi;
            }
        }
        __syncthreads();
    }
}
template <int Panel, int Block, int Threads>
__global__ __launch_bounds__(Threads, 2) void solve_inverse_wy_from_split_gram_blocked_kernel( const float* __restrict__ g11_scratch, long long g11_stride0, const float* __restrict__ s_scratch, const float* __restrict__ tau, const float* __restrict__ w, float* __restrict__ z, int n, int panel_start, int trailing_cols) {
    const int b = blockIdx.y; const int col = blockIdx.x * Threads + threadIdx.x; const float* g11_b = g11_scratch + static_cast<long long>(b) * g11_stride0;
    const float* s_b = s_scratch + static_cast<long long>(b) * Panel * Block; const float* tau_b = tau + static_cast<long long>(b) * n + panel_start; const float* w_b = w + static_cast<long long>(b) * Panel * trailing_cols;
    float* z_b = z + static_cast<long long>(b) * Panel * trailing_cols; __shared__ float g11_shared[Block * Block]; __shared__ float s_shared[Panel * Block]; __shared__ float tau_shared[Panel];
    __shared__ float z_shared[Panel * Threads];
    for (int idx = threadIdx.x; idx < Block * Block; idx += Threads) {
        g11_shared[idx] = g11_b[idx];
    }
    for (int idx = threadIdx.x; idx < Panel * Block; idx += Threads) {
        s_shared[idx] = s_b[idx];
    }
    for (int idx = threadIdx.x; idx < Panel; idx += Threads) {
        tau_shared[idx] = tau_b[idx];
    }
    __syncthreads();
    #pragma unroll
    for (int block_start = 0; block_start < Panel; block_start += Block) {
        float zv[Block];
        #pragma unroll
        for (int r = 0; r < Block; ++r) {
            zv[r] = 0.0f;
        }
        if (col < trailing_cols) {
            #pragma unroll
            for (int r = 0; r < Block; ++r) {
                const int row = block_start + r; float accum = 0.0f;
                for (int inner = 0; inner < Panel; ++inner) {
                    if (inner < block_start) {
                        float gij = 0.0f;
                        if (row < Block) {
                            gij = g11_shared[inner * Block + row];
                        } else if (inner < Block) {
                            gij = s_shared[inner * Block + (row - Block)];
                        } else {
                            gij = s_shared[row * Block + (inner - Block)];
                        }
                        accum = fmaf( gij, z_shared[inner * Threads + threadIdx.x], accum);
                    }
                }
                #pragma unroll
                for (int inner = 0; inner < Block; ++inner) {
                    if (inner < r) {
                        const int gram_inner = block_start + inner; float gij = 0.0f;
                        if (row < Block) {
                            gij = g11_shared[gram_inner * Block + row];
                        } else if (gram_inner < Block) {
                            gij = s_shared[gram_inner * Block + (row - Block)];
                        } else {
                            gij = s_shared[row * Block + (gram_inner - Block)];
                        }
                        accum = fmaf(gij, zv[inner], accum);
                    }
                }
                const float zi = tau_shared[row] * (w_b[row * trailing_cols + col] - accum); zv[r] = zi; z_shared[row * Threads + threadIdx.x] = zi; z_b[row * trailing_cols + col] = zi;
            }
        }
        __syncthreads();
    }
}
template <int N, int MacroCols, int Panel, int Threads>
__global__ __launch_bounds__(Threads, 2) void apply_prev_leaves_to_next_leaf_kernel( float* __restrict__ h, const float* __restrict__ v_macro, long long v_stride0, const float* __restrict__ leaf_grams, const float* __restrict__ tau, int macro_start, int prev_leaves, int target_offset) {
    const int b = blockIdx.x; const int tid = threadIdx.x; const int active_rows = N - macro_start; float* h_b = h + static_cast<long long>(b) * N * N; const float* v_b = v_macro + static_cast<long long>(b) * v_stride0;
    constexpr int Leaves = MacroCols / Panel; const float* gram_b = leaf_grams + static_cast<long long>(b) * Leaves * Panel * Panel; const float* tau_b = tau + static_cast<long long>(b) * N + macro_start;
    extern __shared__ float smem[]; float* c_tile = smem; float* w_tile = c_tile + active_rows * Panel; float* z_tile = w_tile + Panel * Panel; float* g_tile = z_tile + Panel * Panel;
    float* tau_tile = g_tile + Panel * Panel;
    for (int idx = tid; idx < active_rows * Panel; idx += Threads) {
        const int row = idx / Panel; const int col = idx - row * Panel;
        c_tile[idx] = h_b[
            static_cast<long long>(macro_start + row) * N +
            macro_start + target_offset + col];
    }
    __syncthreads();
    for (int leaf = 0; leaf < prev_leaves; ++leaf) {
        for (int idx = tid; idx < Panel * Panel; idx += Threads) {
            g_tile[idx] = gram_b[leaf * Panel * Panel + idx];
        }
        for (int idx = tid; idx < Panel; idx += Threads) {
            tau_tile[idx] = tau_b[leaf * Panel + idx];
        }
        __syncthreads();
        for (int idx = tid; idx < Panel * Panel; idx += Threads) {
            const int row = idx / Panel; const int col = idx - row * Panel; float accum = 0.0f;
            for (int rel = 0; rel < active_rows; ++rel) {
                accum = fmaf( v_b[static_cast<long long>(rel) * MacroCols + leaf * Panel + row], c_tile[rel * Panel + col], accum);
            }
            w_tile[idx] = accum;
        }
        __syncthreads();
        if (tid < Panel) {
            const int col = tid; float zv[Panel];
            #pragma unroll
            for (int row = 0; row < Panel; ++row) {
                float accum = 0.0f;
                #pragma unroll
                for (int inner = 0; inner < Panel; ++inner) {
                    if (inner < row) {
                        accum = fmaf(g_tile[inner * Panel + row], zv[inner], accum);
                    }
                }
                const float zi = tau_tile[row] * (w_tile[row * Panel + col] - accum); zv[row] = zi; z_tile[row * Panel + col] = zi;
            }
        }
        __syncthreads();
        for (int idx = tid; idx < active_rows * Panel; idx += Threads) {
            const int rel = idx / Panel; const int col = idx - rel * Panel; float value = c_tile[idx];
            #pragma unroll
            for (int k = 0; k < Panel; ++k) {
                value = fmaf( -v_b[static_cast<long long>(rel) * MacroCols + leaf * Panel + k], z_tile[k * Panel + col], value);
            }
            c_tile[idx] = value;
        }
        __syncthreads();
    }
    for (int idx = tid; idx < active_rows * Panel; idx += Threads) {
        const int row = idx / Panel; const int col = idx - row * Panel;
        h_b[
            static_cast<long long>(macro_start + row) * N +
            macro_start + target_offset + col] = c_tile[idx];
    }
}
template <int N, int Panel, int PanelThreads, int PanelWork, bool SetAttrs, bool FirstLtTf32>
void qr_macro64_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, int batch) {
    constexpr int macro32 = 2 * Panel; constexpr int macro64 = 4 * Panel; constexpr int macro_leaves = macro64 / Panel; constexpr int prep_threads = 256;
    constexpr int panel_shared_bytes = ((N + 1) * Panel + PanelWork) * sizeof(float);
    auto leaf_grams = torch::empty({a.size(0), macro_leaves, Panel, Panel}, a.options());
    auto g_local = torch::empty({a.size(0), macro32, macro32}, a.options());
    auto g_macro = torch::empty({a.size(0), macro64, macro64}, a.options());
    auto w_local_workspace = torch::empty({a.size(0), macro32, macro32}, a.options());
    auto w_workspace = torch::empty({a.size(0), macro64, N}, a.options());
    constexpr size_t lt_workspace_bytes = 16 * 1024 * 1024;
    auto lt_workspace = torch::empty({static_cast<long long>(lt_workspace_bytes)}, a.options().dtype(at::kByte));
    if constexpr (SetAttrs) {
        constexpr int prep_shared_bytes = (N * Panel + 3 * Panel * Panel + Panel) * sizeof(float); static bool panel_attrs_set = false;
        if (!panel_attrs_set) {
            SET_KERNEL_ATTRS( (qr_panel_cached_kernel<N, Panel, PanelThreads, false, false, false, false, true>), panel_shared_bytes);
            SET_KERNEL_ATTRS( (apply_prev_leaves_to_next_leaf_kernel<N, macro32, Panel, prep_threads>), prep_shared_bytes);
            SET_KERNEL_ATTRS( (apply_prev_leaves_to_next_leaf_kernel<N, macro64, Panel, prep_threads>), prep_shared_bytes); panel_attrs_set = true;
        }
    }
    constexpr int copy_threads = 256; constexpr int copy_blocks = 4096; copy_first_cols_v4_kernel<N, macro64, copy_threads><<<copy_blocks, copy_threads, 0>>>( a.data_ptr<float>(), h.data_ptr<float>(), batch);
    for (int macro_start = 0; macro_start < N; macro_start += macro64) {
        const int active_rows = N - macro_start; const int second_start = macro_start + macro32; const int macro_end = macro_start + macro64;
        const int prep_active_shared_bytes =
            (active_rows * Panel + 3 * Panel * Panel + Panel) * static_cast<int>(sizeof(float));
        auto v_macro = torch::empty({a.size(0), active_rows, macro64}, a.options());
        for (int leaf = 0; leaf < 2; ++leaf) {
            const int leaf_offset = leaf * Panel; const int panel_start = macro_start + leaf_offset;
            if (leaf == 0) {
                qr_panel_cached_kernel<N, Panel, PanelThreads, false, false, false, false, true>
                    <<<batch, PanelThreads, panel_shared_bytes>>>( h.data_ptr<float>(), tau.data_ptr<float>(), nullptr, nullptr, 0, panel_start, v_macro.data_ptr<float>(), v_macro.stride(0), macro64, leaf_offset, leaf_offset);
                auto g_leaf = leaf_grams.select(1, leaf); cublas_leaf0_gram_from_macro( g_leaf.data_ptr<float>(), g_leaf.stride(0), v_macro.data_ptr<float>(), v_macro.stride(0), macro64, active_rows, batch, true);
                apply_prev_leaves_to_next_leaf_kernel<N, macro64, Panel, prep_threads>
                    <<<batch, prep_threads, prep_active_shared_bytes>>>( h.data_ptr<float>(), v_macro.data_ptr<float>(), v_macro.stride(0), leaf_grams.data_ptr<float>(), tau.data_ptr<float>(), macro_start, leaf + 1, leaf_offset + Panel);
            } else {
                qr_panel_cached_kernel<N, Panel, PanelThreads, false, false, false, false, true>
                    <<<batch, PanelThreads, panel_shared_bytes>>>( h.data_ptr<float>(), tau.data_ptr<float>(), nullptr, nullptr, 0, panel_start, v_macro.data_ptr<float>(), v_macro.stride(0), macro64, leaf_offset, leaf_offset);
            }
        }
        auto v_first = v_macro.as_strided(
            {a.size(0), active_rows, macro32},
            {v_macro.stride(0), macro64, 1});
        auto c_next = h.slice(1, macro_start, N).slice(2, second_start, macro_end); at::bmm_out(g_local, v_first.transpose(1, 2), v_first); at::bmm_out(w_local_workspace, v_first.transpose(1, 2), c_next);
        constexpr int local_apply_threads = 128;
        solve_inverse_wy_from_gram_blocked_kernel<macro32, Panel, local_apply_threads>
            <<<dim3(1, batch), local_apply_threads, 0>>>( g_local.data_ptr<float>(), tau.data_ptr<float>(), w_local_workspace.data_ptr<float>(), w_local_workspace.data_ptr<float>(), N, macro_start, macro32);
        c_next.baddbmm_(v_first, w_local_workspace, 1.0, -1.0); const int second_active_rows = N - second_start;
        const int second_prep_active_shared_bytes =
            (second_active_rows * Panel + 3 * Panel * Panel + Panel) * static_cast<int>(sizeof(float));
        for (int leaf = 0; leaf < 2; ++leaf) {
            const int leaf_offset = leaf * Panel; const int panel_start = second_start + leaf_offset; const int macro_row_offset = macro32 + leaf_offset; const int macro_col_offset = macro32 + leaf_offset;
            float* v_second_base = v_macro.data_ptr<float>() + static_cast<long long>(macro32) * macro64 + macro32; float* g_second_base = leaf_grams.data_ptr<float>() + static_cast<long long>(2) * Panel * Panel;
            if (leaf == 0) {
                qr_panel_cached_kernel<N, Panel, PanelThreads, false, false, false, false, true>
                    <<<batch, PanelThreads, panel_shared_bytes>>>( h.data_ptr<float>(), tau.data_ptr<float>(), nullptr, nullptr, 0, panel_start, v_macro.data_ptr<float>(), v_macro.stride(0), macro64, macro_row_offset, macro_col_offset);
                auto g_leaf = leaf_grams.select(1, 2); cublas_leaf0_gram_from_macro( g_leaf.data_ptr<float>(), g_leaf.stride(0), v_second_base, v_macro.stride(0), macro64, second_active_rows, batch, true);
                apply_prev_leaves_to_next_leaf_kernel<N, macro64, Panel, prep_threads>
                    <<<batch, prep_threads, second_prep_active_shared_bytes>>>( h.data_ptr<float>(), v_second_base, v_macro.stride(0), g_second_base, tau.data_ptr<float>(), second_start, leaf + 1, leaf_offset + Panel);
            } else {
                qr_panel_cached_kernel<N, Panel, PanelThreads, false, false, false, false, true>
                    <<<batch, PanelThreads, panel_shared_bytes>>>( h.data_ptr<float>(), tau.data_ptr<float>(), nullptr, nullptr, 0, panel_start, v_macro.data_ptr<float>(), v_macro.stride(0), macro64, macro_row_offset, macro_col_offset);
            }
        }
        if (macro_end < N) {
            const int trailing_cols = N - macro_end; at::bmm_out(g_macro, v_macro.transpose(1, 2), v_macro);
            auto w = w_workspace.as_strided({a.size(0), macro64, trailing_cols}, {macro64 * trailing_cols, trailing_cols, 1});
            auto c = h.slice(1, macro_start, N).slice(2, macro_end, N);
            if (macro_start == 0) {
                auto c_in = a.slice(1, macro_start, N).slice(2, macro_end, N);
                cublas_w_from_vt_c( w.data_ptr<float>(), c_in.data_ptr<float>(), v_macro.data_ptr<float>(), v_macro.stride(0), N, macro64, macro64, active_rows, trailing_cols, batch, true);
            } else {
                cublas_w_from_vt_c( w.data_ptr<float>(), c.data_ptr<float>(), v_macro.data_ptr<float>(), v_macro.stride(0), N, macro64, macro64, active_rows, trailing_cols, batch, true);
            }
            constexpr int apply_t_threads = 96; const int apply_t_blocks = (trailing_cols + apply_t_threads - 1) / apply_t_threads;
            solve_inverse_wy_from_gram_blocked_kernel<macro64, Panel, apply_t_threads>
                <<<dim3(apply_t_blocks, batch), apply_t_threads, 0>>>( g_macro.data_ptr<float>(), tau.data_ptr<float>(), w.data_ptr<float>(), w.data_ptr<float>(), N, macro_start, trailing_cols);
            if (macro_start == 0) {
                auto c_in = a.slice(1, macro_start, N).slice(2, macro_end, N);
                cublaslt_tail_update_out_of_place( c.data_ptr<float>(), c_in.data_ptr<float>(), v_macro.data_ptr<float>(), v_macro.stride(0), w.data_ptr<float>(), N, macro64, active_rows, trailing_cols, batch, FirstLtTf32, lt_workspace.data_ptr(), lt_workspace_bytes);
            } else {
                c.baddbmm_(v_macro, w, 1.0, -1.0);
            }
        }
    }
}
}  // namespace
std::vector<torch::Tensor> qr_small_cuda(torch::Tensor a) {
    const int n = static_cast<int>(a.size(1)); const int batch = static_cast<int>(a.size(0));
    auto tau = torch::empty({a.size(0), a.size(1)}, a.options());
    auto h = torch::empty({a.size(0), a.size(1), a.size(2)}, a.options());
    if (n == kN) {
        qr32_kernel<<<batch, kThreads, 0>>>( a.data_ptr<float>(), h.data_ptr<float>(), tau.data_ptr<float>());
    } else if (n == kN176) {
        constexpr int threads = 1024; constexpr int shared_bytes = kN176 * kLD176Resident * sizeof(float); static bool attrs_set = false;
        if (!attrs_set) {
            SET_KERNEL_ATTRS((qr176_resident_kernel<threads>), shared_bytes); attrs_set = true;
        }
        qr176_resident_kernel<threads><<<batch, threads, shared_bytes>>>( a.data_ptr<float>(), h.data_ptr<float>(), tau.data_ptr<float>());
    } else if (n == kN352) {
        constexpr int panel_shared_bytes = ((kN352 + 1) * kPanel352 + kPanelThreads352) * sizeof(float);
        constexpr int update_shared_bytes =
            (kN352 * kTileUpdate352 + kN352 * kPanel352 + 2 * kPanel352 * kTileUpdate352) *
            static_cast<int>(sizeof(float));
        constexpr int copy_threads = 256; constexpr int copy_blocks = 1024; static bool update_attrs_set = false;
        if (!update_attrs_set) {
            SET_KERNEL_ATTRS( (qr_panel_tile_update_kernel<kN352, kPanel352, kTileUpdate352, kPanelThreads352>), update_shared_bytes); update_attrs_set = true;
        }
        auto t_scratch = torch::empty(
            {a.size(0), kPanel352, kPanel352},
            a.options());
        copy_matrix_v4_kernel<copy_threads><<<copy_blocks, copy_threads, 0>>>( a.data_ptr<float>(), h.data_ptr<float>(), static_cast<long long>(batch) * kN352 * kN352);
        for (int panel_start = 0; panel_start < kN352; panel_start += kPanel352) {
            qr_panel_cached_kernel<kN352, kPanel352, kPanelThreads352, false>
                <<<batch, kPanelThreads352, panel_shared_bytes>>>( h.data_ptr<float>(), tau.data_ptr<float>(), t_scratch.data_ptr<float>(), nullptr, 0, panel_start);
            const int panel_end = panel_start + kPanel352;
            if (panel_end < kN352) {
                const int trailing_cols = kN352 - panel_end; const int tiles = (trailing_cols + kTileUpdate352 - 1) / kTileUpdate352;
                qr_panel_tile_update_kernel<
                    kN352,
                    kPanel352,
                    kTileUpdate352,
                    kPanelThreads352>
                    <<<dim3(tiles, batch), kPanelThreads352, update_shared_bytes>>>( h.data_ptr<float>(), t_scratch.data_ptr<float>(), panel_start);
            }
        }
    } else if (n == kN512) {
        qr_macro64_cuda<kN512, kPanel512, kPanelThreads512, kPanelThreads512, false, false>( a, h, tau, batch);
    } else if (n == kN1024) {
        constexpr int panel_work1024 = ((kPanelThreads1024 / 32) + 1) * kPanel1024; qr_macro64_cuda<kN1024, kPanel1024, kPanelThreads1024, panel_work1024, true, true>( a, h, tau, batch);
    }
    return {h, tau};
}
int64_t detect_tiny_suffix_512_cuda(torch::Tensor a) {
    TORCH_CHECK(a.is_cuda(), "a must be CUDA"); TORCH_CHECK(a.scalar_type() == at::kFloat, "a must be float32");
    TORCH_CHECK(a.dim() == 3 && a.size(1) == kN512 && a.size(2) == kN512,
                "detector expects [batch,512,512]");
    constexpr int k0 = kN512 / 8; constexpr int k1 = kN512 / 4; constexpr int k2 = kN512 / 2; constexpr int k3 = kN512 / 2 + 2 * kPanel512; constexpr int k4 = (3 * kN512) / 4; constexpr int threads = 256;
    const int batch = static_cast<int>(a.size(0));
    auto reject = torch::empty({1}, a.options().dtype(at::kInt));
    C10_CUDA_CHECK(cudaMemset(reject.data_ptr<int>(), 0, sizeof(int))); suffix_sample_reject_kernel<kN512, 32><<<batch, 32, 0>>>( a.data_ptr<float>(), reject.data_ptr<int>()); int host_reject = 0;
    C10_CUDA_CHECK(cudaMemcpy( &host_reject, reject.data_ptr<int>(), sizeof(int), cudaMemcpyDeviceToHost));
    if (host_reject != 0) return 0;
    auto factors = torch::empty({batch}, a.options().dtype(at::kInt));
    suffix_factor_cols_kernel<kN512, kPanel512, threads>
        <<<batch, threads, 0>>>( a.data_ptr<float>(), factors.data_ptr<int>(), k0, k1, k2, k3, k4);
    auto factor_result = torch::empty({1}, a.options().dtype(at::kInt));
    reduce_factor_cols_kernel<threads>
        <<<1, threads, 0>>>( factors.data_ptr<int>(), factor_result.data_ptr<int>(), batch);
    int factor_cols = 0; C10_CUDA_CHECK(cudaMemcpy( &factor_cols, factor_result.data_ptr<int>(), sizeof(int), cudaMemcpyDeviceToHost));
    return factor_cols;
}
std::vector<torch::Tensor> qr_small_prefix_cuda(torch::Tensor a, int64_t factor_cols_arg) {
    const int n = static_cast<int>(a.size(1)); const int batch = static_cast<int>(a.size(0)); const int factor_cols = static_cast<int>(factor_cols_arg);
    if (n != kN512 || factor_cols <= 0 || factor_cols >= n || (factor_cols & 15) != 0) {
        return qr_small_cuda(a);
    }
    auto tau = torch::empty({a.size(0), a.size(1)}, a.options());
    auto h = torch::empty({a.size(0), a.size(1), a.size(2)}, a.options());
    const bool old_allow_tf32 = at::globalContext().allowTF32CuBLAS(); constexpr int panel_shared_bytes = ((kN512 + 1) * kPanel512 + kPanelThreads512) * sizeof(float); constexpr int copy_threads = 256;
    constexpr int copy_blocks = 4096; constexpr int kMacroPrefix512 = 2 * kPanel512; constexpr int kMacroPrefixLeaves512 = kMacroPrefix512 / kPanel512; constexpr int prep_threads = 256;
    copy_first_cols_v4_kernel<kN512, kMacroPrefix512, copy_threads><<<copy_blocks, copy_threads, 0>>>( a.data_ptr<float>(), h.data_ptr<float>(), batch);
    zero_suffix_cols_v4_kernel<kN512, copy_threads><<<copy_blocks, copy_threads, 0>>>( h.data_ptr<float>(), factor_cols, batch);
    auto leaf_grams = torch::empty(
        {a.size(0), kMacroPrefixLeaves512, kPanel512, kPanel512},
        a.options());
    auto s_macro = torch::empty({a.size(0), kMacroPrefix512, kPanel512}, a.options());
    auto w_workspace = torch::empty({a.size(0), kMacroPrefix512, factor_cols}, a.options());
    constexpr size_t lt_workspace_bytes = 16 * 1024 * 1024;
    auto lt_workspace = torch::empty(
        {static_cast<long long>(lt_workspace_bytes)},
        a.options().dtype(at::kByte));
    for (int macro_start = 0; macro_start < factor_cols; macro_start += kMacroPrefix512) {
        const int active_rows = kN512 - macro_start; const int macro_end = macro_start + kMacroPrefix512;
        const int prep_active_shared_bytes =
            (active_rows * kPanel512 + 3 * kPanel512 * kPanel512 + kPanel512) *
            static_cast<int>(sizeof(float));
        auto v_macro = torch::empty({a.size(0), active_rows, kMacroPrefix512}, a.options());
        for (int leaf = 0; leaf < kMacroPrefixLeaves512; ++leaf) {
            const int leaf_offset = leaf * kPanel512; const int panel_start = macro_start + leaf_offset;
            if (leaf + 1 < kMacroPrefixLeaves512) {
                qr_panel_cached_kernel<kN512, kPanel512, kPanelThreads512, false, false, false, false, true>
                    <<<batch, kPanelThreads512, panel_shared_bytes>>>( h.data_ptr<float>(), tau.data_ptr<float>(), nullptr, nullptr, 0, panel_start, v_macro.data_ptr<float>(), v_macro.stride(0), kMacroPrefix512, leaf_offset, leaf_offset);
                auto g_leaf = leaf_grams.select(1, leaf); cublas_leaf0_gram_from_macro( g_leaf.data_ptr<float>(), g_leaf.stride(0), v_macro.data_ptr<float>(), v_macro.stride(0), kMacroPrefix512, active_rows, batch, false);
                apply_prev_leaves_to_next_leaf_kernel<
                    kN512,
                    kMacroPrefix512,
                    kPanel512,
                    prep_threads>
                    <<<batch, prep_threads, prep_active_shared_bytes>>>( h.data_ptr<float>(), v_macro.data_ptr<float>(), v_macro.stride(0), leaf_grams.data_ptr<float>(), tau.data_ptr<float>(), macro_start, leaf + 1, leaf_offset + kPanel512);
            } else {
                qr_panel_cached_kernel<kN512, kPanel512, kPanelThreads512, false, false, false, false, true>
                    <<<batch, kPanelThreads512, panel_shared_bytes>>>( h.data_ptr<float>(), tau.data_ptr<float>(), nullptr, nullptr, 0, panel_start, v_macro.data_ptr<float>(), v_macro.stride(0), kMacroPrefix512, leaf_offset, leaf_offset);
            }
        }
        if (macro_end < factor_cols) {
            const int trailing_cols = factor_cols - macro_end; auto g_first = leaf_grams.select(1, 0);
            cublas_leaf1_cross_gram_from_macro( s_macro.data_ptr<float>(), s_macro.stride(0), v_macro.data_ptr<float>(), v_macro.stride(0), kMacroPrefix512, active_rows, batch, false);
            auto w = w_workspace.as_strided(
                {a.size(0), kMacroPrefix512, trailing_cols},
                {kMacroPrefix512 * trailing_cols, trailing_cols, 1});
            auto z = w; auto c = h.slice(1, macro_start, kN512).slice(2, macro_end, factor_cols);
            if (macro_start == 0) {
                auto c_in = a.slice(1, macro_start, kN512).slice(2, macro_end, factor_cols);
                cublas_w_from_vt_c( w.data_ptr<float>(), c_in.data_ptr<float>(), v_macro.data_ptr<float>(), v_macro.stride(0), kN512, kMacroPrefix512, kMacroPrefix512, active_rows, trailing_cols, batch, old_allow_tf32);
            } else {
                cublas_w_from_vt_c( w.data_ptr<float>(), c.data_ptr<float>(), v_macro.data_ptr<float>(), v_macro.stride(0), kN512, kMacroPrefix512, kMacroPrefix512, active_rows, trailing_cols, batch, old_allow_tf32);
            }
            constexpr int apply_t_threads = 96; const int apply_t_blocks = (trailing_cols + apply_t_threads - 1) / apply_t_threads;
            solve_inverse_wy_from_split_gram_blocked_kernel<
                kMacroPrefix512,
                kPanel512,
                apply_t_threads>
                <<<dim3(apply_t_blocks, batch), apply_t_threads, 0>>>( g_first.data_ptr<float>(), g_first.stride(0), s_macro.data_ptr<float>(), tau.data_ptr<float>(), w.data_ptr<float>(), z.data_ptr<float>(), kN512, macro_start, trailing_cols);
            if (macro_start == 0) {
                auto c_in = a.slice(1, macro_start, kN512).slice(2, macro_end, factor_cols);
                cublaslt_tail_update_out_of_place( c.data_ptr<float>(), c_in.data_ptr<float>(), v_macro.data_ptr<float>(), v_macro.stride(0), z.data_ptr<float>(), kN512, kMacroPrefix512, active_rows, trailing_cols, batch, old_allow_tf32, lt_workspace.data_ptr(), lt_workspace_bytes);
            } else {
                c.baddbmm_(v_macro, z, 1.0, -1.0);
            }
        }
    }
    constexpr int zero_threads = 256; const long long tau_total = static_cast<long long>(batch) * (kN512 - factor_cols); int zero_blocks = static_cast<int>((tau_total + zero_threads - 1) / zero_threads);
    if (zero_blocks < 1) zero_blocks = 1;
    if (zero_blocks > 1024) zero_blocks = 1024;
    zero_tau_suffix_kernel<kN512, zero_threads><<<zero_blocks, zero_threads, 0>>>( tau.data_ptr<float>(), factor_cols, batch); at::globalContext().setAllowTF32CuBLAS(old_allow_tf32);
    return {h, tau};
}
std::vector<torch::Tensor> qr_2048_cuda(torch::Tensor a) {
    const int batch = static_cast<int>(a.size(0));
    auto h = torch::empty({a.size(0), a.size(1), a.size(2)}, a.options());
    auto tau = torch::empty({a.size(0), a.size(1)}, a.options());
    constexpr int macro_cols = 4 * kPanel2048;
    auto t_scratch = torch::empty({a.size(0), kPanel2048, kPanel2048}, a.options());
    auto g_macro = torch::empty({a.size(0), macro_cols, macro_cols}, a.options());
    auto w_workspace = torch::empty({a.size(0), macro_cols, kN2048}, a.options());
    auto z_workspace = torch::empty({a.size(0), kPanel2048, kN2048}, a.options());
    constexpr size_t lt_workspace_bytes = 16 * 1024 * 1024;
    auto lt_workspace = torch::empty(
        {static_cast<long long>(lt_workspace_bytes)},
        a.options().dtype(at::kByte));
    constexpr int max_panel_shared_bytes = ((kN2048 + 1) * kPanel2048 + kPanelThreads2048) * sizeof(float); constexpr int copy_threads = 256; constexpr int copy_blocks = 2048; static bool panel_attrs_set = false;
    if (!panel_attrs_set) {
        SET_KERNEL_ATTRS( (qr_panel_cached_kernel<kN2048, kPanel2048, kPanelThreads2048, false, true>), max_panel_shared_bytes);
        SET_KERNEL_ATTRS( (qr_panel_cached_kernel<kN2048, kPanel2048, kPanelThreads2048, true, true>), max_panel_shared_bytes);
        SET_KERNEL_ATTRS( (qr_panel_cached_kernel<kN2048, kPanel2048, kPanelThreads2048, true, true, true, true, true>), max_panel_shared_bytes);
        SET_KERNEL_ATTRS( (qr_panel_cached_kernel<kN2048, kPanel2048, kPanelThreads2048, false, true, true, true, true>), max_panel_shared_bytes); panel_attrs_set = true;
    }
    copy_matrix_v4_kernel<copy_threads><<<copy_blocks, copy_threads, 0>>>( a.data_ptr<float>(), h.data_ptr<float>(), static_cast<long long>(batch) * kN2048 * kN2048);
    for (int macro_start = 0; macro_start < kN2048; macro_start += macro_cols) {
        const int active_rows = kN2048 - macro_start;
        const int macro_end = (macro_start + macro_cols < kN2048)
            ? macro_start + macro_cols
            : kN2048;
        const int panel_shared_bytes =
            ((active_rows + 1) * kPanel2048 + kPanelThreads2048) * sizeof(float);
        auto v_macro = torch::empty({a.size(0), active_rows, macro_cols}, a.options());
        for (int panel_start = macro_start; panel_start < macro_end; panel_start += kPanel2048) {
            const int panel_end = panel_start + kPanel2048; const int leaf_active_rows = kN2048 - panel_start; const int macro_offset = panel_start - macro_start;
            qr_panel_cached_kernel<kN2048, kPanel2048, kPanelThreads2048, false, true, true, true, true>
                <<<batch, kPanelThreads2048, panel_shared_bytes>>>( h.data_ptr<float>(), tau.data_ptr<float>(), t_scratch.data_ptr<float>(), nullptr, 0, panel_start, v_macro.data_ptr<float>(), v_macro.stride(0), macro_cols, macro_offset, macro_offset);
            auto v_leaf = v_macro.as_strided(
                {a.size(0), leaf_active_rows, kPanel2048},
                {v_macro.stride(0), macro_cols, 1},
                c10::optional<int64_t>( static_cast<int64_t>(macro_offset) * macro_cols + macro_offset));
            if (panel_end < macro_end) {
                const int local_cols = macro_end - panel_end;
                auto w = w_workspace.as_strided(
                    {a.size(0), kPanel2048, local_cols},
                    {kPanel2048 * local_cols, local_cols, 1});
                auto z = z_workspace.as_strided(
                    {a.size(0), kPanel2048, local_cols},
                    {kPanel2048 * local_cols, local_cols, 1});
                auto c = h.slice(1, panel_start, kN2048).slice(2, panel_end, macro_end);
                cublas_w_from_vt_c( w.data_ptr<float>(), c.data_ptr<float>(), v_leaf.data_ptr<float>(), v_leaf.stride(0), kN2048, kPanel2048, macro_cols, leaf_active_rows, local_cols, batch, true);
                constexpr int local_apply_t_threads = 64;
                const int local_apply_t_blocks =
                    (local_cols + local_apply_t_threads - 1) / local_apply_t_threads;
                apply_t_transpose_kernel<kPanel2048, local_apply_t_threads>
                    <<<dim3(local_apply_t_blocks, batch), local_apply_t_threads, 0>>>( t_scratch.data_ptr<float>(), w.data_ptr<float>(), z.data_ptr<float>(), local_cols);
                constexpr int local_update_threads = 256;
                const int local_update_blocks =
                    (leaf_active_rows * local_cols + local_update_threads - 1) / local_update_threads;
                panel_rank_update_kernel<kN2048, kPanel2048, local_update_threads>
                    <<<dim3(local_update_blocks, batch), local_update_threads, 0>>>(
                        h.data_ptr<float>(),
                        v_leaf.data_ptr<float>(),
                        v_leaf.stride(0),
                        macro_cols,
                        z.data_ptr<float>(),
                        panel_start,
                        panel_end,
                        leaf_active_rows,
                        local_cols);
            }
        }
        if (macro_end < kN2048) {
            const int trailing_cols = kN2048 - macro_end; at::bmm_out(g_macro, v_macro.transpose(1, 2), v_macro);
            auto w = w_workspace.as_strided(
                {a.size(0), macro_cols, trailing_cols},
                {macro_cols * trailing_cols, trailing_cols, 1});
            auto c = h.slice(1, macro_start, kN2048).slice(2, macro_end, kN2048);
            cublas_w_from_vt_c( w.data_ptr<float>(), c.data_ptr<float>(), v_macro.data_ptr<float>(), v_macro.stride(0), kN2048, macro_cols, macro_cols, active_rows, trailing_cols, batch, true); constexpr int apply_t_threads = 96;
            const int apply_t_blocks = (trailing_cols + apply_t_threads - 1) / apply_t_threads;
            solve_inverse_wy_from_gram_blocked_kernel<macro_cols, kPanel2048, apply_t_threads>
                <<<dim3(apply_t_blocks, batch), apply_t_threads, 0>>>( g_macro.data_ptr<float>(), tau.data_ptr<float>(), w.data_ptr<float>(), w.data_ptr<float>(), kN2048, macro_start, trailing_cols);
            cublaslt_tail_update_out_of_place( c.data_ptr<float>(), c.data_ptr<float>(), v_macro.data_ptr<float>(), v_macro.stride(0), w.data_ptr<float>(), kN2048, macro_cols, active_rows, trailing_cols, batch, true, lt_workspace.data_ptr(), lt_workspace_bytes);
        }
    }
    return {h, tau};
}
std::vector<torch::Tensor> qr_4096_cuda(torch::Tensor a) {
    const int batch = static_cast<int>(a.size(0));
    auto h = torch::empty({a.size(0), a.size(1), a.size(2)}, a.options());
    auto tau = torch::empty({a.size(0), a.size(1)}, a.options());
    constexpr int super_panel = 2 * kPanel4096; constexpr int early_macro = 16 * kPanel4096;
    auto t_first = torch::empty({a.size(0), kPanel4096, kPanel4096}, a.options());
    auto t_scratch = torch::empty({a.size(0), super_panel, super_panel}, a.options());
    auto g_early = torch::empty({a.size(0), early_macro, early_macro}, a.options());
    auto w_workspace = torch::empty({a.size(0), early_macro, kN4096}, a.options());
    auto z_workspace = torch::empty({a.size(0), super_panel, kN4096}, a.options());
    constexpr size_t lt_workspace_bytes = 16 * 1024 * 1024;
    auto lt_workspace = torch::empty(
        {static_cast<long long>(lt_workspace_bytes)},
        a.options().dtype(at::kByte));
    constexpr int panel_shared_bytes = ((kN4096 + 1) * kPanel4096 + kPanelThreads4096) * sizeof(float); constexpr int copy_threads = 256; constexpr int copy_blocks = 8192; constexpr int macro_apply_threads = 32;
    constexpr int macro_solve_shared_bytes =
        (early_macro * early_macro + early_macro + early_macro * macro_apply_threads) *
        static_cast<int>(sizeof(float));
    constexpr int late_max_rows = kPanel4096LateMaxRows;
    constexpr int late_panel_shared_bytes =
        ((kPanel4096LateMaxRows + 1) * kPanel4096Late + kPanelThreads4096Late) * sizeof(float);
    static bool panel_attrs_set = false;
    if (!panel_attrs_set) {
        SET_KERNEL_ATTRS( (qr_panel_cached_kernel<kN4096, kPanel4096, kPanelThreads4096, false>), panel_shared_bytes);
        SET_KERNEL_ATTRS( (qr_panel_cached_kernel<kN4096, kPanel4096, kPanelThreads4096, true, false, true, true, true>), panel_shared_bytes);
        SET_KERNEL_ATTRS( (qr_panel_cached_kernel<kN4096, kPanel4096, kPanelThreads4096, false, false, true, true, true>), panel_shared_bytes);
        SET_KERNEL_ATTRS( (qr_panel_cached_kernel<kN4096, kPanel4096Late, kPanelThreads4096Late, false, true>), late_panel_shared_bytes);
        SET_KERNEL_ATTRS( (qr_panel_cached_kernel<kN4096, kPanel4096Late, kPanelThreads4096Late, true, true>), late_panel_shared_bytes);
        SET_KERNEL_ATTRS( (solve_inverse_wy_from_gram_blocked_dynamic_kernel<early_macro, kPanel4096, macro_apply_threads>), macro_solve_shared_bytes); panel_attrs_set = true;
    }
    copy_matrix_v4_kernel<copy_threads><<<copy_blocks, copy_threads, 0>>>( a.data_ptr<float>(), h.data_ptr<float>(), static_cast<long long>(batch) * kN4096 * kN4096);
    for (int panel_start = 0; panel_start < kN4096;) {
        const int active_rows = kN4096 - panel_start;
        int panel_end = (panel_start + super_panel < kN4096)
            ? panel_start + super_panel
            : kN4096;
        if (active_rows <= late_max_rows) {
            const int direct_shared_bytes =
                ((active_rows + 1) * kPanel4096Late + kPanelThreads4096Late) * sizeof(float);
            if (panel_end >= kN4096) {
                qr_panel_cached_kernel<kN4096, kPanel4096Late, kPanelThreads4096Late, false, true>
                    <<<batch, kPanelThreads4096Late, direct_shared_bytes>>>( h.data_ptr<float>(), tau.data_ptr<float>(), t_scratch.data_ptr<float>(), nullptr, 0, panel_start);
                break;
            }
            const int trailing_cols = kN4096 - panel_end;
            auto v = torch::empty({a.size(0), active_rows, super_panel}, a.options());
            qr_panel_cached_kernel<kN4096, kPanel4096Late, kPanelThreads4096Late, true, true>
                <<<batch, kPanelThreads4096Late, direct_shared_bytes>>>( h.data_ptr<float>(), tau.data_ptr<float>(), t_scratch.data_ptr<float>(), v.data_ptr<float>(), v.stride(0), panel_start);
            auto w = w_workspace.as_strided(
                {a.size(0), super_panel, trailing_cols},
                {super_panel * trailing_cols, trailing_cols, 1});
            auto z = z_workspace.as_strided(
                {a.size(0), super_panel, trailing_cols},
                {super_panel * trailing_cols, trailing_cols, 1});
            auto c = h.slice(1, panel_start, kN4096).slice(2, panel_end, kN4096);
            cublas_w_from_vt_c( w.data_ptr<float>(), c.data_ptr<float>(), v.data_ptr<float>(), v.stride(0), kN4096, super_panel, super_panel, active_rows, trailing_cols, batch, true); constexpr int apply_t_threads = 128;
            const int apply_t_blocks = (trailing_cols + apply_t_threads - 1) / apply_t_threads;
            apply_t_transpose_kernel<super_panel, apply_t_threads>
                <<<dim3(apply_t_blocks, batch), apply_t_threads, 0>>>( t_scratch.data_ptr<float>(), w.data_ptr<float>(), z.data_ptr<float>(), trailing_cols);
            c.baddbmm_(v, z, 1.0, -1.0); panel_start += super_panel; continue;
        }
        panel_end = (panel_start + early_macro < kN4096)
            ? panel_start + early_macro
            : kN4096;
        auto v_macro = torch::empty({a.size(0), active_rows, early_macro}, a.options());
        for (int leaf_start = panel_start; leaf_start < panel_end; leaf_start += kPanel4096) {
            const int leaf_end = leaf_start + kPanel4096; const int leaf_active_rows = kN4096 - leaf_start; const int macro_offset = leaf_start - panel_start;
            qr_panel_cached_kernel<kN4096, kPanel4096, kPanelThreads4096, false, false, true, true, true>
                <<<batch, kPanelThreads4096, panel_shared_bytes>>>( h.data_ptr<float>(), tau.data_ptr<float>(), t_first.data_ptr<float>(), nullptr, 0, leaf_start, v_macro.data_ptr<float>(), v_macro.stride(0), early_macro, macro_offset, macro_offset);
            auto v_leaf = v_macro.as_strided(
                {a.size(0), leaf_active_rows, kPanel4096},
                {v_macro.stride(0), early_macro, 1},
                c10::optional<int64_t>( static_cast<int64_t>(macro_offset) * early_macro + macro_offset));
            if (leaf_end < panel_end) {
                const int local_cols = panel_end - leaf_end;
                auto w = w_workspace.as_strided(
                    {a.size(0), kPanel4096, local_cols},
                    {kPanel4096 * local_cols, local_cols, 1});
                auto z = z_workspace.as_strided(
                    {a.size(0), kPanel4096, local_cols},
                    {kPanel4096 * local_cols, local_cols, 1});
                auto c = h.slice(1, leaf_start, kN4096).slice(2, leaf_end, panel_end);
                cublas_w_from_vt_c( w.data_ptr<float>(), c.data_ptr<float>(), v_leaf.data_ptr<float>(), v_leaf.stride(0), kN4096, kPanel4096, early_macro, leaf_active_rows, local_cols, batch, true);
                constexpr int local_apply_t_threads = 64;
                const int local_apply_t_blocks =
                    (local_cols + local_apply_t_threads - 1) / local_apply_t_threads;
                apply_t_transpose_kernel<kPanel4096, local_apply_t_threads>
                    <<<dim3(local_apply_t_blocks, batch), local_apply_t_threads, 0>>>( t_first.data_ptr<float>(), w.data_ptr<float>(), z.data_ptr<float>(), local_cols);
                constexpr int local_update_threads = 256;
                const int local_update_blocks = (leaf_active_rows * local_cols + local_update_threads - 1) / local_update_threads;
                panel_rank_update_kernel<kN4096, kPanel4096, local_update_threads>
                    <<<dim3(local_update_blocks, batch), local_update_threads, 0>>>( h.data_ptr<float>(), v_leaf.data_ptr<float>(), v_leaf.stride(0), early_macro, z.data_ptr<float>(), leaf_start, leaf_end, leaf_active_rows, local_cols);
            }
        }
        if (panel_end < kN4096) {
            const int trailing_cols = kN4096 - panel_end; at::bmm_out(g_early, v_macro.transpose(1, 2), v_macro);
            auto w = w_workspace.as_strided(
                {a.size(0), early_macro, trailing_cols},
                {early_macro * trailing_cols, trailing_cols, 1});
            auto c = h.slice(1, panel_start, kN4096).slice(2, panel_end, kN4096);
            cublas_w_from_vt_c( w.data_ptr<float>(), c.data_ptr<float>(), v_macro.data_ptr<float>(), v_macro.stride(0), kN4096, early_macro, early_macro, active_rows, trailing_cols, batch, true);
            const int apply_t_blocks = (trailing_cols + macro_apply_threads - 1) / macro_apply_threads;
            solve_inverse_wy_from_gram_blocked_dynamic_kernel<
                early_macro,
                kPanel4096,
                macro_apply_threads>
                <<<dim3(apply_t_blocks, batch), macro_apply_threads, macro_solve_shared_bytes>>>( g_early.data_ptr<float>(), tau.data_ptr<float>(), w.data_ptr<float>(), w.data_ptr<float>(), kN4096, panel_start, trailing_cols);
            cublaslt_tail_update_out_of_place( c.data_ptr<float>(), c.data_ptr<float>(), v_macro.data_ptr<float>(), v_macro.stride(0), w.data_ptr<float>(), kN4096, early_macro, active_rows, trailing_cols, batch, true, lt_workspace.data_ptr(), lt_workspace_bytes);
        }
        panel_start += early_macro;
    }
    return {h, tau};
}
"""
EXT = None
def load_ext():
    global EXT
    if EXT is None:
        from torch.utils.cpp_extension import load
        source_dir = Path(tempfile.gettempdir()) / "qr_householder_ext_n512_all_tf32_gram_v12_macro32"
        source_dir.mkdir(parents=True, exist_ok=True)
        cuda_path = source_dir / "qr_householder_all.cu"
        cuda_path.write_text(CPP_SRC + "\n" + CUDA_SRC)
        EXT = load(
            name="qr_householder_ext_n512_all_tf32_gram_v12_macro32",
            sources=[str(cuda_path)],
            extra_cuda_cflags=[
                "-O3",
                "--use_fast_math",
            ],
            verbose=False,
        )
    return EXT
def small_qr(data: torch.Tensor) -> output_t:
    ext = load_ext()
    if ext is None:
        return torch.ops.aten.geqrf.default(data)
    contiguous = data if data.is_contiguous() else data.contiguous()
    if contiguous.shape[-1] == 512 and contiguous.shape[0] <= 32:
        return torch.ops.aten.geqrf.default(contiguous)
    if contiguous.shape[-1] == 512:
        factor_cols = int(ext.detect_tiny_suffix_512(contiguous))
        if 0 < factor_cols < 512:
            h, tau = ext.qr_small_prefix(contiguous, factor_cols)
            return h, tau
    h, tau = ext.qr_small(contiguous)
    return h, tau
def large_qr(data: torch.Tensor) -> output_t:
    ext = load_ext()
    if ext is not None:
        contiguous = data if data.is_contiguous() else data.contiguous()
        if contiguous.shape[-1] == 2048:
            h, tau = ext.qr_2048(contiguous)
            return h, tau
        if contiguous.shape[-1] == 4096:
            h, tau = ext.qr_4096(contiguous)
            return h, tau
    return torch.ops.aten.geqrf.default(data)
def custom_kernel(data: input_t) -> output_t:
    if not (
        data.is_cuda
        and data.dtype == torch.float32
        and data.dim() == 3
        and data.shape[-1] == data.shape[-2]
    ):
        return torch.ops.aten.geqrf.default(data)
    n = data.shape[-1]
    if n in (32, 176, 352, 512, 1024):
        return small_qr(data)
    if n in (2048, 4096):
        return large_qr(data)
    return torch.ops.aten.geqrf.default(data)
scrolls · 1417 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