Skip to content
KernelIndex
Search⌘K

submission 833276

ptxv · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-833276?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.87ms
#69 of 515
2026-06-24

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c4ca7060bfabce7b27eb0dd0efc9de7d2d68b8504af41dc91dd6bc37dccdc125
license declaredunknown
license concludedunknown
authorsptxv
imported2026-08-26

Techniques

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

shared-memory__shared__ unsigned int warp_maxima[2 * Warps];
vector-width = float4__device__ __forceinline__ float4 ptx_ld_global_v4_f32(const float* ptr) {

Kernel source

submission.py3814 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);
std::vector<torch::Tensor> qr_cholqr_hr512_cuda(torch::Tensor a, torch::Tensor r, torch::Tensor info);
int64_t detect_tiny_suffix_512_cuda(torch::Tensor a);
int64_t detect_upper_512_cuda(torch::Tensor a);
int64_t detect_upper_1024_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("qr_cholqr_hr512", &qr_cholqr_hr512_cuda, "n512 CholeskyQR direct Householder reconstruction");
    m.def("detect_tiny_suffix_512", &detect_tiny_suffix_512_cuda, "certified n512 tiny-suffix detector");
    m.def("detect_upper_512", &detect_upper_512_cuda, "certified n512 upper-triangular detector");
    m.def("detect_upper_1024", &detect_upper_1024_cuda, "certified n1024 upper-triangular 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 kPanel176 = 8;
constexpr int kPanelThreads176 = 256;
constexpr int kUpdateThreads176 = 128;
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 kUpdateThreads352 = 128;
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;
// B200 / compute capability 10.0 permits a 16-column panel up to 3584 active
// rows within the large dynamic shared-memory limit.
constexpr int kPanel4096LateMaxRows = 3584;

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) {
    cublasHandle_t handle = at::cuda::getCurrentCUDABlasHandle();
    const float alpha = 1.0f;
    const float 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,
        kPanel512,
        kPanel512,
        active_rows,
        &alpha,
        v_macro,
        CUDA_R_32F,
        macro_cols,
        v_stride0,
        v_macro,
        CUDA_R_32F,
        macro_cols,
        v_stride0,
        &beta,
        gram,
        CUDA_R_32F,
        kPanel512,
        gram_stride0,
        batch,
        compute_type,
        algo);
    TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, "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) {
    cublasHandle_t handle = at::cuda::getCurrentCUDABlasHandle();
    const float alpha = 1.0f;
    const float 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,
        kPanel512,
        macro_cols,
        active_rows,
        &alpha,
        v_macro + kPanel512,
        CUDA_R_32F,
        macro_cols,
        v_stride0,
        v_macro,
        CUDA_R_32F,
        macro_cols,
        v_stride0,
        &beta,
        s_macro,
        CUDA_R_32F,
        kPanel512,
        s_stride0,
        batch,
        compute_type,
        algo);
    TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, "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 n,
    int panel_cols,
    int v_ld_cols,
    int active_rows,
    int trailing_cols,
    int batch,
    bool allow_tf32) {
    cublasHandle_t handle = at::cuda::getCurrentCUDABlasHandle();
    const float alpha = 1.0f;
    const float 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,
        trailing_cols,
        panel_cols,
        active_rows,
        &alpha,
        c_tail,
        CUDA_R_32F,
        n,
        static_cast<long long>(n) * n,
        v,
        CUDA_R_32F,
        v_ld_cols,
        v_stride0,
        &beta,
        w,
        CUDA_R_32F,
        trailing_cols,
        static_cast<long long>(panel_cols) * trailing_cols,
        batch,
        compute_type,
        algo);
    TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, "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();
    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;

    const cublasComputeType_t compute_type = allow_tf32
        ? CUBLAS_COMPUTE_32F_FAST_TF32
        : CUBLAS_COMPUTE_32F;
    TORCH_CHECK(
        cublasLtMatmulDescCreate(&op_desc, compute_type, CUDA_R_32F) ==
            CUBLAS_STATUS_SUCCESS,
        "tail update cuBLASLt desc create failed");
    cublasOperation_t trans = CUBLAS_OP_N;
    TORCH_CHECK(
        cublasLtMatmulDescSetAttribute(
            op_desc, CUBLASLT_MATMUL_DESC_TRANSA, &trans, sizeof(trans)) ==
            CUBLAS_STATUS_SUCCESS,
        "tail update cuBLASLt transa set failed");
    TORCH_CHECK(
        cublasLtMatmulDescSetAttribute(
            op_desc, CUBLASLT_MATMUL_DESC_TRANSB, &trans, sizeof(trans)) ==
            CUBLAS_STATUS_SUCCESS,
        "tail update cuBLASLt transb set failed");

    TORCH_CHECK(
        cublasLtMatrixLayoutCreate(&a_desc, CUDA_R_32F, trailing_cols, macro_cols, trailing_cols) ==
            CUBLAS_STATUS_SUCCESS,
        "tail update cuBLASLt A layout create failed");
    TORCH_CHECK(
        cublasLtMatrixLayoutCreate(&b_desc, CUDA_R_32F, macro_cols, active_rows, macro_cols) ==
            CUBLAS_STATUS_SUCCESS,
        "tail update cuBLASLt B layout create failed");
    TORCH_CHECK(
        cublasLtMatrixLayoutCreate(&c_desc, CUDA_R_32F, trailing_cols, active_rows, n) ==
            CUBLAS_STATUS_SUCCESS,
        "tail update cuBLASLt C layout create failed");
    TORCH_CHECK(
        cublasLtMatrixLayoutCreate(&d_desc, CUDA_R_32F, trailing_cols, active_rows, n) ==
            CUBLAS_STATUS_SUCCESS,
        "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;
    TORCH_CHECK(
        cublasLtMatrixLayoutSetAttribute(
            a_desc, CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT, &batch_count, sizeof(batch_count)) ==
            CUBLAS_STATUS_SUCCESS,
        "tail update cuBLASLt A batch set failed");
    TORCH_CHECK(
        cublasLtMatrixLayoutSetAttribute(
            b_desc, CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT, &batch_count, sizeof(batch_count)) ==
            CUBLAS_STATUS_SUCCESS,
        "tail update cuBLASLt B batch set failed");
    TORCH_CHECK(
        cublasLtMatrixLayoutSetAttribute(
            c_desc, CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT, &batch_count, sizeof(batch_count)) ==
            CUBLAS_STATUS_SUCCESS,
        "tail update cuBLASLt C batch set failed");
    TORCH_CHECK(
        cublasLtMatrixLayoutSetAttribute(
            d_desc, CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT, &batch_count, sizeof(batch_count)) ==
            CUBLAS_STATUS_SUCCESS,
        "tail update cuBLASLt D batch set failed");
    TORCH_CHECK(
        cublasLtMatrixLayoutSetAttribute(
            a_desc, CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET, &a_stride, sizeof(a_stride)) ==
            CUBLAS_STATUS_SUCCESS,
        "tail update cuBLASLt A stride set failed");
    TORCH_CHECK(
        cublasLtMatrixLayoutSetAttribute(
            b_desc, CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET, &b_stride, sizeof(b_stride)) ==
            CUBLAS_STATUS_SUCCESS,
        "tail update cuBLASLt B stride set failed");
    TORCH_CHECK(
        cublasLtMatrixLayoutSetAttribute(
            c_desc, CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET, &cd_stride, sizeof(cd_stride)) ==
            CUBLAS_STATUS_SUCCESS,
        "tail update cuBLASLt C stride set failed");
    TORCH_CHECK(
        cublasLtMatrixLayoutSetAttribute(
            d_desc, CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET, &cd_stride, sizeof(cd_stride)) ==
            CUBLAS_STATUS_SUCCESS,
        "tail update cuBLASLt D stride set failed");

    TORCH_CHECK(
        cublasLtMatmulPreferenceCreate(&pref) == CUBLAS_STATUS_SUCCESS,
        "tail update cuBLASLt preference create failed");
    TORCH_CHECK(
        cublasLtMatmulPreferenceSetAttribute(
            pref,
            CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES,
            &workspace_bytes,
            sizeof(workspace_bytes)) == CUBLAS_STATUS_SUCCESS,
        "tail update cuBLASLt workspace set failed");

    cublasLtMatmulHeuristicResult_t heuristic;
    int returned = 0;
    TORCH_CHECK(
        cublasLtMatmulAlgoGetHeuristic(
            handle,
            op_desc,
            a_desc,
            b_desc,
            c_desc,
            d_desc,
            pref,
            1,
            &heuristic,
            &returned) == CUBLAS_STATUS_SUCCESS,
        "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);

    cublasLtMatmulPreferenceDestroy(pref);
    cublasLtMatrixLayoutDestroy(d_desc);
    cublasLtMatrixLayoutDestroy(c_desc);
    cublasLtMatrixLayoutDestroy(b_desc);
    cublasLtMatrixLayoutDestroy(a_desc);
    cublasLtMatmulDescDestroy(op_desc);
    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_kernel(
    const float* __restrict__ a,
    float* __restrict__ h,
    long long total) {
    for (long long idx = static_cast<long long>(blockIdx.x) * Threads + threadIdx.x;
         idx < total;
         idx += static_cast<long long>(gridDim.x) * Threads) {
        h[idx] = a[idx];
    }
}

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;

    // This stage can only reject. It never authorizes a shortcut.
    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) {
    constexpr int Warps = Threads / 32;
    constexpr int RowVecs = N / 4;
    constexpr int MatrixVecs = N * RowVecs;
    __shared__ unsigned int warp_maxima[4 * 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;

    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);
    }

    local0 = warp_max_u32(local0);
    local1 = warp_max_u32(local1);
    local2 = warp_max_u32(local2);
    local3 = warp_max_u32(local3);

    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;
    }
    __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;
        block0 = warp_max_u32(block0);
        block1 = warp_max_u32(block1);
        block2 = warp_max_u32(block2);
        block3 = warp_max_u32(block3);
        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);
            int matrix_cols = 0;
            if (all == 0.0f) {
                matrix_cols = 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;
            }
            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 copy_prefix_zero_suffix_v4_kernel(
    const float* __restrict__ a,
    float* __restrict__ h,
    long long total,
    int factor_cols) {
    const long long vec_total = total >> 2;
    constexpr int RowVecs = N / 4;
    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 int col4 = static_cast<int>(idx % RowVecs) * 4;
        const long long base = idx << 2;
        const float4 vals = (col4 < factor_cols)
            ? ptx_ld_global_v4_f32(a + base)
            : make_float4(0.0f, 0.0f, 0.0f, 0.0f);
        ptx_st_global_v4_f32(h + base, vals);
    }
}

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;
}

template <int N, int Threads>
__global__ __launch_bounds__(Threads, 1) void cholqr_hr_lu_kernel(
    const float* __restrict__ r,
    const int* __restrict__ info,
    float* __restrict__ h,
    float* __restrict__ tau,
    int* __restrict__ ok) {
    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    constexpr long long MatrixElems = static_cast<long long>(N) * N;
    const float* r_b = r + static_cast<long long>(b) * MatrixElems;
    float* h_b = h + static_cast<long long>(b) * MatrixElems;
    float* tau_b = tau + static_cast<long long>(b) * N;

    __shared__ float reduce[32];
    __shared__ float row_sign;
    __shared__ float pivot;
    __shared__ float diag_scale;
    __shared__ int bad;

    if (tid == 0) {
        bad = (info[b] != 0);
        float min_diag = 3.4028234663852886e38f;
        float max_diag = 0.0f;
        if (!bad) {
            for (int j = 0; j < N; ++j) {
                const float d = fabsf(r_b[j * N + j]);
                if (!(isfinite(d) && d > 0.0f)) {
                    bad = 1;
                    break;
                }
                min_diag = fminf(min_diag, d);
                max_diag = fmaxf(max_diag, d);
            }
        }
        diag_scale = fmaxf(max_diag, 1.0e-30f);
        if (!bad && min_diag < 1.0e-7f * diag_scale) {
            bad = 1;
        }
    }
    __syncthreads();

    if (bad) {
        for (int j = tid; j < N; j += Threads) {
            tau_b[j] = 0.0f;
        }
        if (tid == 0) ok[b] = 0;
        return;
    }

    for (int k = 0; k < N; ++k) {
        if (tid == 0) {
            const float x = h_b[k * N + k];
            if (isfinite(x)) {
                row_sign = (x >= 0.0f) ? -1.0f : 1.0f;
            } else {
                bad = 1;
                row_sign = -1.0f;
            }
        }
        __syncthreads();
        if (bad) {
            if (tid == 0) ok[b] = 0;
            return;
        }

        for (int j = k + tid; j < N; j += Threads) {
            const float shifted = h_b[k * N + j] - row_sign * r_b[k * N + j];
            h_b[k * N + j] = shifted;
            if (!isfinite(shifted)) {
                atomicExch(&bad, 1);
            }
        }
        __syncthreads();
        if (bad) {
            if (tid == 0) ok[b] = 0;
            return;
        }

        if (tid == 0) {
            pivot = h_b[k * N + k];
            const float pivot_floor = fmaxf(1.0e-20f, 1.0e-7f * diag_scale);
            if (!(isfinite(pivot) && fabsf(pivot) > pivot_floor)) {
                bad = 1;
            }
        }
        __syncthreads();
        if (bad) {
            if (tid == 0) ok[b] = 0;
            return;
        }

        const float inv_pivot = 1.0f / pivot;
        float local_norm = 0.0f;
        for (int i = k + 1 + tid; i < N; i += Threads) {
            const float v = h_b[i * N + k] * inv_pivot;
            h_b[i * N + k] = v;
            local_norm = fmaf(v, v, local_norm);
            if (!isfinite(v)) {
                atomicExch(&bad, 1);
            }
        }
        const float tail_norm = block_sum_thread0<Threads>(local_norm, reduce);
        if (tid == 0) {
            const float tau_k = 2.0f / (1.0f + tail_norm);
            tau_b[k] = tau_k;
            if (!(isfinite(tau_k) && tau_k >= 0.0f && tau_k <= 2.0f)) {
                bad = 1;
            }
        }
        __syncthreads();
        if (bad) {
            if (tid == 0) ok[b] = 0;
            return;
        }

        const int trailing = N - k - 1;
        const int update_total = trailing * trailing;
        for (int idx = tid; idx < update_total; idx += Threads) {
            const int rel_i = idx / trailing;
            const int rel_j = idx - rel_i * trailing;
            const int i = k + 1 + rel_i;
            const int j = k + 1 + rel_j;
            const float updated = fmaf(-h_b[i * N + k], h_b[k * N + j], h_b[i * N + j]);
            h_b[i * N + j] = updated;
            if (!isfinite(updated)) {
                atomicExch(&bad, 1);
            }
        }
        __syncthreads();
        if (bad) {
            if (tid == 0) ok[b] = 0;
            return;
        }

        for (int j = k + tid; j < N; j += Threads) {
            const float out = row_sign * r_b[k * N + j];
            h_b[k * N + j] = out;
            if (!isfinite(out)) {
                atomicExch(&bad, 1);
            }
        }
        __syncthreads();
        if (bad) {
            if (tid == 0) ok[b] = 0;
            return;
        }
    }

    if (tid == 0) ok[b] = 1;
}

template <int N, int Threads>
__global__ __launch_bounds__(Threads, 4) void upper_max_classifier_kernel(
    const float* __restrict__ a,
    unsigned int* __restrict__ maxima) {
    constexpr int Warps = Threads / 32;
    __shared__ unsigned int warp_maxima[2 * Warps];

    const int b = blockIdx.y;
    constexpr int MatrixElems = N * N;
    const float* a_b = a + static_cast<long long>(b) * MatrixElems;

    unsigned int local_all = 0;
    unsigned int local_lower = 0;
    for (int idx = blockIdx.x * Threads + threadIdx.x;
         idx < MatrixElems;
         idx += gridDim.x * Threads) {
        const unsigned int bits = __float_as_uint(fabsf(a_b[idx]));
        local_all = max(local_all, bits);
        const int row = idx / N;
        const int col = idx - row * N;
        if (row > col) {
            local_lower = max(local_lower, bits);
        }
    }

    local_all = warp_max_u32(local_all);
    local_lower = warp_max_u32(local_lower);

    const int lane = threadIdx.x & 31;
    const int warp = threadIdx.x >> 5;
    if (lane == 0) {
        warp_maxima[warp] = local_all;
        warp_maxima[Warps + warp] = local_lower;
    }
    __syncthreads();

    if (warp == 0) {
        unsigned int block_all = (lane < Warps) ? warp_maxima[lane] : 0;
        unsigned int block_lower = (lane < Warps) ? warp_maxima[Warps + lane] : 0;
        block_all = warp_max_u32(block_all);
        block_lower = warp_max_u32(block_lower);
        if (lane == 0) {
            unsigned int* out = maxima + static_cast<long long>(b) * 2;
            atomicMax(out + 0, block_all);
            atomicMax(out + 1, block_lower);
        }
    }
}

template <int N>
bool all_matrices_upper_certified(torch::Tensor a) {
    constexpr int threads = 256;
    constexpr int blocks_per_matrix = (N >= 4096) ? 512 : 128;
    constexpr float eps32 = 1.1920928955078125e-7f;
    constexpr float route_budget = 2.0f;
    const int batch = static_cast<int>(a.size(0));

    auto maxima = torch::empty({a.size(0), 2}, a.options().dtype(at::kInt));
    C10_CUDA_CHECK(cudaMemset(
        maxima.data_ptr<int>(), 0, static_cast<size_t>(batch) * 2 * sizeof(int)));
    upper_max_classifier_kernel<N, threads>
        <<<dim3(blocks_per_matrix, batch), threads, 0>>>(
            a.data_ptr<float>(),
            reinterpret_cast<unsigned int*>(maxima.data_ptr<int>()));

    std::vector<unsigned int> host(static_cast<size_t>(batch) * 2);
    C10_CUDA_CHECK(cudaMemcpy(
        host.data(),
        maxima.data_ptr<int>(),
        host.size() * sizeof(unsigned int),
        cudaMemcpyDeviceToHost));

    union BitsFloat { unsigned int u; float f; };
    for (int b = 0; b < batch; ++b) {
        BitsFloat all{host[static_cast<size_t>(b) * 2 + 0]};
        BitsFloat lower{host[static_cast<size_t>(b) * 2 + 1]};
        if (all.f != 0.0f && lower.f > route_budget * eps32 * all.f) {
            return false;
        }
    }
    return true;
}

__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 Threads>
__global__ __launch_bounds__(Threads, 1) void rebuild_t_from_panel_kernel(
    const float* __restrict__ h,
    const float* __restrict__ tau,
    float* __restrict__ t_scratch,
    int panel_start) {
    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    constexpr int Warps = Threads / 32;

    const int active_rows = N - panel_start;
    const long long matrix_stride = static_cast<long long>(N) * N;
    const float* h_b = h + static_cast<long long>(b) * matrix_stride;
    const float* tau_b = tau + static_cast<long long>(b) * N;
    float* t_b = t_scratch + static_cast<long long>(b) * Panel * Panel;

    __shared__ float reduce[Warps];
    __shared__ float y[Panel];

    for (int idx = tid; idx < Panel * Panel; idx += Threads) {
        t_b[idx] = 0.0f;
    }
    __syncthreads();

    #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) {
                    const float vi = h_b[
                        static_cast<long long>(panel_start + rel) * N + panel_start + i];
                    const float vj = h_b[
                        static_cast<long long>(panel_start + rel) * N + panel_start + j];
                    local = fmaf(vi, vj, local);
                }
                if (tid == 0) {
                    local += h_b[
                        static_cast<long long>(panel_start + j) * N + panel_start + i];
                }
                const float dot = block_sum_thread0<Threads>(local, reduce);
                if (tid == 0) {
                    y[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], y[inner], accum);
                }
            }
            t_b[tid * Panel + j] = accum;
        }
        if (tid == j) {
            t_b[j * Panel + j] = tau_j;
        }
        __syncthreads();
    }
}

template <int Panel, int Threads>
__global__ __launch_bounds__(Threads, 4) void build_t_from_gram_kernel(
    float* __restrict__ t_scratch,
    const float* __restrict__ tau,
    int n,
    int panel_start) {
    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    float* t_b = t_scratch + static_cast<long long>(b) * Panel * Panel;
    const float* tau_b = tau + static_cast<long long>(b) * n;

    __shared__ float y[Panel];

    #pragma unroll
    for (int j = 0; j < Panel; ++j) {
        const float tau_j = tau_b[panel_start + j];
        if (tid < j) {
            y[tid] = -tau_j * t_b[tid * Panel + j];
        }
        __syncthreads();

        if (tid < j) {
            float accum = 0.0f;
            #pragma unroll
            for (int inner = 0; inner < Panel; ++inner) {
                if (inner < j && tid <= inner) {
                    accum = fmaf(t_b[tid * Panel + inner], y[inner], accum);
                }
            }
            t_b[tid * Panel + j] = accum;
        }
        if (tid == j) {
            t_b[j * Panel + j] = tau_j;
        }
        __syncthreads();
    }
}

template <int N, int Panel, int Threads, bool PackedV>
__global__ __launch_bounds__(Threads, 1) void qr_panel_wy_update_kernel(
    float* __restrict__ h,
    const float* __restrict__ t_scratch,
    int panel_start,
    const float* __restrict__ v_pack,
    long long v_stride0) {
    const int tile = blockIdx.x;
    const int b = blockIdx.y;
    const int tid = threadIdx.x;

    const int panel_end = (panel_start + Panel < N) ? (panel_start + Panel) : N;
    const int panel_cols = panel_end - panel_start;
    const int active_rows = N - panel_start;
    const int j = panel_end + tile * Threads + tid;

    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;
    const float* v_b = nullptr;
    if constexpr (PackedV) {
        v_b = v_pack + static_cast<long long>(b) * v_stride0;
    }

    extern __shared__ float v_panel[];
    if constexpr (!PackedV && Panel == 8) {
        for (int rel_load = tid; rel_load < active_rows; rel_load += Threads) {
            const int row = panel_start + rel_load;
            const float4 lo = ptx_ld_global_v4_f32(
                h_b + static_cast<long long>(row) * N + panel_start);
            const float4 hi = ptx_ld_global_v4_f32(
                h_b + static_cast<long long>(row) * N + panel_start + 4);
            float* dst = v_panel + rel_load * Panel;
            dst[0] = (rel_load == 0) ? 1.0f : ((rel_load > 0) ? lo.x : 0.0f);
            dst[1] = (rel_load == 1) ? 1.0f : ((rel_load > 1) ? lo.y : 0.0f);
            dst[2] = (rel_load == 2) ? 1.0f : ((rel_load > 2) ? lo.z : 0.0f);
            dst[3] = (rel_load == 3) ? 1.0f : ((rel_load > 3) ? lo.w : 0.0f);
            dst[4] = (rel_load == 4) ? 1.0f : ((rel_load > 4) ? hi.x : 0.0f);
            dst[5] = (rel_load == 5) ? 1.0f : ((rel_load > 5) ? hi.y : 0.0f);
            dst[6] = (rel_load == 6) ? 1.0f : ((rel_load > 6) ? hi.z : 0.0f);
            dst[7] = (rel_load == 7) ? 1.0f : ((rel_load > 7) ? hi.w : 0.0f);
        }
    } else {
        constexpr int RowThreads = Threads / Panel;
        const int t_load = tid - (tid / Panel) * Panel;
        for (int rel_load = tid / Panel; rel_load < active_rows; rel_load += RowThreads) {
            const int row = panel_start + rel_load;
            const int k = panel_start + t_load;
            float value = 0.0f;
            if constexpr (PackedV) {
                value = v_b[static_cast<long long>(rel_load) * Panel + t_load];
            } else if (t_load < panel_cols) {
                if (rel_load == t_load) {
                    value = 1.0f;
                } else if (rel_load > t_load) {
                    value = h_b[static_cast<long long>(row) * N + k];
                }
            }
            v_panel[rel_load * Panel + t_load] = value;
        }
    }
    __syncthreads();

    if (j >= N) {
        return;
    }

    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;
    long long offset = static_cast<long long>(panel_start) * N + j;
    int rel = 0;
    #pragma unroll 24
    for (int row = panel_start; row < N; ++row, ++rel, offset += N) {
        const float c = h_b[offset];
        const int v_offset = rel * Panel;
        w0 = fmaf(v_panel[v_offset], c, w0);
        w1 = fmaf(v_panel[v_offset + 1], c, w1);
        w2 = fmaf(v_panel[v_offset + 2], c, w2);
        w3 = fmaf(v_panel[v_offset + 3], c, w3);
        w4 = fmaf(v_panel[v_offset + 4], c, w4);
        w5 = fmaf(v_panel[v_offset + 5], c, w5);
        w6 = fmaf(v_panel[v_offset + 6], c, w6);
        w7 = fmaf(v_panel[v_offset + 7], c, w7);
    }

    const float z0 = t_b[0] * w0;
    const float z1 = fmaf(t_b[1], w0, t_b[Panel + 1] * w1);
    const float z2 = fmaf(t_b[2], w0, fmaf(t_b[Panel + 2], w1, t_b[2 * Panel + 2] * w2));
    const float z3 = fmaf(t_b[3], w0, fmaf(t_b[Panel + 3], w1, fmaf(t_b[2 * Panel + 3], w2, t_b[3 * Panel + 3] * w3)));
    const float z4 = 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))));
    const float z5 = 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)))));
    const float z6 = 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))))));
    const float z7 = 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)))))));

    offset = static_cast<long long>(panel_start) * N + j;
    rel = 0;
    #pragma unroll 24
    for (int row = panel_start; row < N; ++row, ++rel, offset += N) {
        const int v_offset = rel * Panel;
        float delta = v_panel[v_offset] * z0;
        delta = fmaf(v_panel[v_offset + 1], z1, delta);
        delta = fmaf(v_panel[v_offset + 2], z2, delta);
        delta = fmaf(v_panel[v_offset + 3], z3, delta);
        delta = fmaf(v_panel[v_offset + 4], z4, delta);
        delta = fmaf(v_panel[v_offset + 5], z5, delta);
        delta = fmaf(v_panel[v_offset + 6], z6, delta);
        delta = fmaf(v_panel[v_offset + 7], z7, delta);
        h_b[offset] -= delta;
    }
}

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 N, int Panel, int Cols, int Threads>
__global__ __launch_bounds__(Threads, 1) void apply_panel_to_next_cols_kernel(
    float* __restrict__ h,
    const float* __restrict__ t_scratch,
    int panel_start) {
    const int target_col = blockIdx.x;
    const int b = blockIdx.y;
    const int tid = threadIdx.x;
    const int lane = tid & 31;
    const int warp = tid >> 5;
    constexpr int Warps = Threads / 32;

    const int active_rows = N - panel_start;
    const int j = panel_start + Panel + target_col;
    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* v_panel = smem;
    float* partials = smem + Panel * active_rows;
    float* z = partials + Warps * Panel;

    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[t * active_rows + rel] = value;
    }
    __syncthreads();

    if (j >= N || target_col >= Cols) {
        return;
    }

    float local[Panel];
    #pragma unroll
    for (int t = 0; t < Panel; ++t) {
        local[t] = 0.0f;
    }

    for (int rel = tid; rel < active_rows; rel += Threads) {
        const float c = h_b[static_cast<long long>(panel_start + rel) * N + j];
        #pragma unroll
        for (int t = 0; t < Panel; ++t) {
            local[t] = fmaf(v_panel[t * active_rows + rel], c, local[t]);
        }
    }

    #pragma unroll
    for (int t = 0; t < Panel; ++t) {
        const float sum = warp_sum(local[t]);
        if (lane == 0) {
            partials[warp * Panel + t] = sum;
        }
    }
    __syncthreads();

    if (tid < Panel) {
        float w = 0.0f;
        #pragma unroll
        for (int r = 0; r < Warps; ++r) {
            w += partials[r * Panel + tid];
        }
        partials[Warps * Panel + tid] = w;
    }
    __syncthreads();

    if (tid < Panel) {
        float value = 0.0f;
        #pragma unroll
        for (int r = 0; r < Panel; ++r) {
            value = fmaf(t_b[r * Panel + tid], partials[Warps * Panel + r], value);
        }
        z[tid] = value;
    }
    __syncthreads();

    for (int rel = tid; rel < active_rows; rel += Threads) {
        float delta = 0.0f;
        #pragma unroll
        for (int t = 0; t < Panel; ++t) {
            delta = fmaf(v_panel[t * active_rows + rel], z[t], delta);
        }
        h_b[static_cast<long long>(panel_start + rel) * N + j] -= delta;
    }
}

template <int N, int Threads>
__global__ __launch_bounds__(Threads, 1) void build_t16_from_two_t8_singleblock_kernel(
    const float* __restrict__ h,
    const float* __restrict__ t_first,
    const float* __restrict__ t_second,
    float* __restrict__ t_out,
    float* __restrict__ v_pack,
    long long v_stride,
    int panel_start) {
    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    const int lane = tid & 31;
    const int warp = tid >> 5;
    constexpr int Half = 8;
    constexpr int Super = 16;
    constexpr int Warps = Threads / 32;

    const int active_rows = N - panel_start;
    const long long matrix_stride = static_cast<long long>(N) * N;
    const float* h_b = h + static_cast<long long>(b) * matrix_stride;
    const float* t1 = t_first + static_cast<long long>(b) * Half * Half;
    const float* t2 = t_second + static_cast<long long>(b) * Half * Half;
    float* tout = t_out + static_cast<long long>(b) * Super * Super;
    float* v_b = v_pack + static_cast<long long>(b) * v_stride;

    constexpr int SuperVec = Super / 4;
    for (int idx = tid; idx < active_rows * SuperVec; idx += Threads) {
        const int rel = idx / SuperVec;
        const int t = (idx - rel * SuperVec) * 4;
        float4 values;
        if (rel > t + 3) {
            values = ptx_ld_global_v4_f32(
                h_b + static_cast<long long>(panel_start + rel) * N + panel_start + t);
        } else {
            values.x = (rel == t) ? 1.0f :
                ((rel > t) ? h_b[static_cast<long long>(panel_start + rel) * N + panel_start + t] : 0.0f);
            values.y = (rel == t + 1) ? 1.0f :
                ((rel > t + 1) ? h_b[static_cast<long long>(panel_start + rel) * N + panel_start + t + 1] : 0.0f);
            values.z = (rel == t + 2) ? 1.0f :
                ((rel > t + 2) ? h_b[static_cast<long long>(panel_start + rel) * N + panel_start + t + 2] : 0.0f);
            values.w = (rel == t + 3) ? 1.0f :
                ((rel > t + 3) ? h_b[static_cast<long long>(panel_start + rel) * N + panel_start + t + 3] : 0.0f);
        }
        ptx_st_global_v4_f32(v_b + static_cast<long long>(rel) * Super + t, values);
    }
    __syncthreads();

    __shared__ float s_shared[Half * Half];
    __shared__ float middle[Half * Half];
    for (int group = warp; group < Half * 2; group += Warps) {
        const int i = group >> 1;
        const int j_base = (group & 1) << 2;
        float local0 = 0.0f;
        float local1 = 0.0f;
        float local2 = 0.0f;
        float local3 = 0.0f;
        for (int rel = lane; rel < active_rows; rel += 32) {
            const float* v_row = v_b + static_cast<long long>(rel) * Super;
            const float v1 = v_row[i];
            const float4 vals = ptx_ld_global_v4_f32(v_row + Half + j_base);
            const float v20 = vals.x;
            const float v21 = vals.y;
            const float v22 = vals.z;
            const float v23 = vals.w;
            local0 = fmaf(v1, v20, local0);
            local1 = fmaf(v1, v21, local1);
            local2 = fmaf(v1, v22, local2);
            local3 = fmaf(v1, v23, local3);
        }
        const float total0 = warp_sum(local0);
        const float total1 = warp_sum(local1);
        const float total2 = warp_sum(local2);
        const float total3 = warp_sum(local3);
        if (lane == 0) {
            s_shared[i * Half + j_base] = total0;
            s_shared[i * Half + j_base + 1] = total1;
            s_shared[i * Half + j_base + 2] = total2;
            s_shared[i * Half + j_base + 3] = total3;
        }
    }
    __syncthreads();

    for (int idx = tid; idx < Super * Super; idx += Threads) {
        tout[idx] = 0.0f;
    }
    __syncthreads();

    for (int idx = tid; idx < Half * Half; idx += Threads) {
        const int row = idx / Half;
        const int col = idx - row * Half;
        tout[row * Super + col] = t1[idx];
        tout[(Half + row) * Super + Half + col] = t2[idx];
    }
    for (int idx = tid; idx < Half * Half; idx += Threads) {
        const int row = idx / Half;
        const int col = idx - row * Half;
        float value = 0.0f;
        #pragma unroll
        for (int k = 0; k < Half; ++k) {
            value = fmaf(t1[row * Half + k], s_shared[k * Half + col], value);
        }
        middle[idx] = value;
    }
    __syncthreads();

    for (int idx = tid; idx < Half * Half; idx += Threads) {
        const int row = idx / Half;
        const int col = idx - row * Half;
        float value = 0.0f;
        #pragma unroll
        for (int k = 0; k < Half; ++k) {
            value = fmaf(middle[row * Half + k], t2[k * Half + col], value);
        }
        tout[row * Super + Half + col] = -value;
    }
}

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 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, 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];
    }
}

}  // namespace

int64_t detect_upper_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,
                "upper detector expects [batch,512,512]");
    TORCH_CHECK(a.is_contiguous(), "a must be contiguous");
    return all_matrices_upper_certified<kN512>(a) ? 1 : 0;
}

int64_t detect_upper_1024_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) == kN1024 && a.size(2) == kN1024,
                "upper detector expects [batch,1024,1024]");
    TORCH_CHECK(a.is_contiguous(), "a must be contiguous");
    return all_matrices_upper_certified<kN1024>(a) ? 1 : 0;
}

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) {
            C10_CUDA_CHECK(cudaFuncSetAttribute(
                qr176_resident_kernel<threads>,
                cudaFuncAttributeMaxDynamicSharedMemorySize,
                shared_bytes));
            C10_CUDA_CHECK(cudaFuncSetAttribute(
                qr176_resident_kernel<threads>,
                cudaFuncAttributePreferredSharedMemoryCarveout,
                100));
            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) {
            C10_CUDA_CHECK(cudaFuncSetAttribute(
                qr_panel_tile_update_kernel<
                    kN352,
                    kPanel352,
                    kTileUpdate352,
                    kPanelThreads352>,
                cudaFuncAttributeMaxDynamicSharedMemorySize,
                update_shared_bytes));
            C10_CUDA_CHECK(cudaFuncSetAttribute(
                qr_panel_tile_update_kernel<
                    kN352,
                    kPanel352,
                    kTileUpdate352,
                    kPanelThreads352>,
                cudaFuncAttributePreferredSharedMemoryCarveout,
                100));
            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) {
        const bool old_allow_tf32 = at::globalContext().allowTF32CuBLAS();
        constexpr int kMacro512 = 2 * kPanel512;
        constexpr int kMacro64_512 = 4 * kPanel512;
        constexpr int kMacroLeaves64_512 = kMacro64_512 / kPanel512;
        constexpr int panel_shared_bytes = ((kN512 + 1) * kPanel512 + kPanelThreads512) * sizeof(float);
        constexpr int prep_threads = 256;
        constexpr int prep_shared_bytes =
            (kN512 * kPanel512 + 3 * kPanel512 * kPanel512 + kPanel512) * sizeof(float);
        auto leaf_grams = torch::empty(
            {a.size(0), kMacroLeaves64_512, kPanel512, kPanel512},
            a.options());
        auto g_local = torch::empty({a.size(0), kMacro512, kMacro512}, a.options());
        auto g_macro = torch::empty({a.size(0), kMacro64_512, kMacro64_512}, a.options());
        auto w_local_workspace = torch::empty({a.size(0), kMacro512, kMacro512}, a.options());
        auto w_workspace = torch::empty({a.size(0), kMacro64_512, kN512}, a.options());
        constexpr size_t lt_workspace_bytes = 4 * 1024 * 1024;
        auto lt_workspace = torch::empty(
            {static_cast<long long>(lt_workspace_bytes)},
            a.options().dtype(at::kByte));
        constexpr int copy_threads = 256;
        constexpr int copy_blocks = 4096;
        copy_first_cols_v4_kernel<kN512, kMacro64_512, copy_threads><<<copy_blocks, copy_threads, 0>>>(
            a.data_ptr<float>(),
            h.data_ptr<float>(),
            batch);

        for (int macro_start = 0; macro_start < kN512; macro_start += kMacro64_512) {
            const int active_rows = kN512 - macro_start;
            const int second_start = macro_start + kMacro512;
            const int macro_end = macro_start + kMacro64_512;
            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, kMacro64_512}, a.options());

            for (int leaf = 0; leaf < 2; ++leaf) {
                const int leaf_offset = leaf * kPanel512;
                const int panel_start = macro_start + leaf_offset;
                if (leaf == 0) {
                    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),
                            kMacro64_512,
                            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),
                        kMacro64_512,
                        active_rows,
                        batch,
                        true);

                    apply_prev_leaves_to_next_leaf_kernel<
                        kN512,
                        kMacro64_512,
                        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),
                            kMacro64_512,
                            leaf_offset,
                            leaf_offset);
                }
            }

            auto v_first = v_macro.as_strided(
                {a.size(0), active_rows, kMacro512},
                {v_macro.stride(0), kMacro64_512, 1});
            auto c_next = h.slice(1, macro_start, kN512).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<
                kMacro512,
                kPanel512,
                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>(),
                    kN512,
                    macro_start,
                    kMacro512);
            c_next.baddbmm_(v_first, w_local_workspace, 1.0, -1.0);

            const int second_active_rows = kN512 - second_start;
            const int second_prep_active_shared_bytes =
                (second_active_rows * kPanel512 + 3 * kPanel512 * kPanel512 + kPanel512) *
                static_cast<int>(sizeof(float));
            for (int leaf = 0; leaf < 2; ++leaf) {
                const int leaf_offset = leaf * kPanel512;
                const int panel_start = second_start + leaf_offset;
                const int macro_row_offset = kMacro512 + leaf_offset;
                const int macro_col_offset = kMacro512 + leaf_offset;
                float* v_second_base =
                    v_macro.data_ptr<float>() +
                    static_cast<long long>(kMacro512) * kMacro64_512 +
                    kMacro512;
                float* g_second_base =
                    leaf_grams.data_ptr<float>() +
                    static_cast<long long>(2) * kPanel512 * kPanel512;
                if (leaf == 0) {
                    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),
                            kMacro64_512,
                            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),
                        kMacro64_512,
                        second_active_rows,
                        batch,
                        true);

                    apply_prev_leaves_to_next_leaf_kernel<
                        kN512,
                        kMacro64_512,
                        kPanel512,
                        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 + 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),
                            kMacro64_512,
                            macro_row_offset,
                            macro_col_offset);
                }
            }

            if (macro_end < kN512) {
                const int trailing_cols = kN512 - macro_end;
                at::bmm_out(g_macro, v_macro.transpose(1, 2), v_macro);

                auto w = w_workspace.as_strided(
                    {a.size(0), kMacro64_512, trailing_cols},
                    {kMacro64_512 * trailing_cols, trailing_cols, 1});
                auto z = w;
                auto c = h.slice(1, macro_start, kN512).slice(2, macro_end, kN512);
                if (macro_start == 0) {
                    auto c_in = a.slice(1, macro_start, kN512).slice(2, macro_end, kN512);
                    at::bmm_out(w, v_macro.transpose(1, 2), c_in);
                } else {
                    at::bmm_out(w, v_macro.transpose(1, 2), c);
                }

                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<kMacro64_512, kPanel512, 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>(),
                        z.data_ptr<float>(),
                        kN512,
                        macro_start,
                        trailing_cols);
                if (macro_start < 32) {
                    auto c_in = a.slice(1, macro_start, kN512).slice(2, macro_end, kN512);
                    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,
                        kMacro64_512,
                        active_rows,
                        trailing_cols,
                        batch,
                        false,
                        lt_workspace.data_ptr(),
                        lt_workspace_bytes);
                } else {
                    c.baddbmm_(v_macro, z, 1.0, -1.0);
                }
            }
        }
        at::globalContext().setAllowTF32CuBLAS(old_allow_tf32);
    } else if (n == kN1024) {
        const bool old_allow_tf32 = at::globalContext().allowTF32CuBLAS();
        constexpr int panel_work1024 = ((kPanelThreads1024 / 32) + 1) * kPanel1024;
        constexpr int panel_shared_bytes = ((kN1024 + 1) * kPanel1024 + panel_work1024) * sizeof(float);
        constexpr int kMacro1024 = 2 * kPanel1024;
        constexpr int kMacro64_1024 = 4 * kPanel1024;
        constexpr int kMacroLeaves64_1024 = kMacro64_1024 / kPanel1024;
        constexpr int prep_threads = 256;
        constexpr int prep_shared_bytes =
            (kN1024 * kPanel1024 + 3 * kPanel1024 * kPanel1024 + kPanel1024) * sizeof(float);
        auto leaf_grams = torch::empty(
            {a.size(0), kMacroLeaves64_1024, kPanel1024, kPanel1024},
            a.options());
        auto g_local = torch::empty({a.size(0), kMacro1024, kMacro1024}, a.options());
        auto g_macro = torch::empty({a.size(0), kMacro64_1024, kMacro64_1024}, a.options());
        auto w_local_workspace = torch::empty({a.size(0), kMacro1024, kMacro1024}, a.options());
        auto w_workspace = torch::empty({a.size(0), kMacro64_1024, kN1024}, a.options());
        constexpr size_t lt_workspace_bytes = 4 * 1024 * 1024;
        auto lt_workspace = torch::empty(
            {static_cast<long long>(lt_workspace_bytes)},
            a.options().dtype(at::kByte));
        constexpr int copy_threads = 256;
        constexpr int copy_blocks = 4096;
        static bool panel_attrs_set = false;
        if (!panel_attrs_set) {
            C10_CUDA_CHECK(cudaFuncSetAttribute(
                qr_panel_cached_kernel<kN1024, kPanel1024, kPanelThreads1024, false, false, false, false, true>,
                cudaFuncAttributeMaxDynamicSharedMemorySize,
                panel_shared_bytes));
            C10_CUDA_CHECK(cudaFuncSetAttribute(
                qr_panel_cached_kernel<kN1024, kPanel1024, kPanelThreads1024, false, false, false, false, true>,
                cudaFuncAttributePreferredSharedMemoryCarveout,
                100));
            C10_CUDA_CHECK(cudaFuncSetAttribute(
                apply_prev_leaves_to_next_leaf_kernel<
                    kN1024,
                    kMacro1024,
                    kPanel1024,
                    prep_threads>,
                cudaFuncAttributeMaxDynamicSharedMemorySize,
                prep_shared_bytes));
            C10_CUDA_CHECK(cudaFuncSetAttribute(
                apply_prev_leaves_to_next_leaf_kernel<
                    kN1024,
                    kMacro1024,
                    kPanel1024,
                    prep_threads>,
                cudaFuncAttributePreferredSharedMemoryCarveout,
                100));
            C10_CUDA_CHECK(cudaFuncSetAttribute(
                apply_prev_leaves_to_next_leaf_kernel<
                    kN1024,
                    kMacro64_1024,
                    kPanel1024,
                    prep_threads>,
                cudaFuncAttributeMaxDynamicSharedMemorySize,
                prep_shared_bytes));
            C10_CUDA_CHECK(cudaFuncSetAttribute(
                apply_prev_leaves_to_next_leaf_kernel<
                    kN1024,
                    kMacro64_1024,
                    kPanel1024,
                    prep_threads>,
                cudaFuncAttributePreferredSharedMemoryCarveout,
                100));
            panel_attrs_set = true;
        }
        copy_first_cols_v4_kernel<kN1024, kMacro64_1024, copy_threads><<<copy_blocks, copy_threads, 0>>>(
            a.data_ptr<float>(),
            h.data_ptr<float>(),
            batch);
        for (int macro_start = 0; macro_start < kN1024; macro_start += kMacro64_1024) {
            const int active_rows = kN1024 - macro_start;
            const int second_start = macro_start + kMacro1024;
            const int macro_end = macro_start + kMacro64_1024;
            const int prep_active_shared_bytes =
                (active_rows * kPanel1024 + 3 * kPanel1024 * kPanel1024 + kPanel1024) *
                static_cast<int>(sizeof(float));
            auto v_macro = torch::empty({a.size(0), active_rows, kMacro64_1024}, a.options());

            for (int leaf = 0; leaf < 2; ++leaf) {
                const int leaf_offset = leaf * kPanel1024;
                const int panel_start = macro_start + leaf_offset;
                if (leaf == 0) {
                    qr_panel_cached_kernel<kN1024, kPanel1024, kPanelThreads1024, false, false, false, false, true>
                        <<<batch, kPanelThreads1024, 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),
                            kMacro64_1024,
                            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),
                        kMacro64_1024,
                        active_rows,
                        batch,
                        true);

                    apply_prev_leaves_to_next_leaf_kernel<
                        kN1024,
                        kMacro64_1024,
                        kPanel1024,
                        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 + kPanel1024);
                } else {
                    qr_panel_cached_kernel<kN1024, kPanel1024, kPanelThreads1024, false, false, false, false, true>
                        <<<batch, kPanelThreads1024, 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),
                            kMacro64_1024,
                            leaf_offset,
                            leaf_offset);
                }
            }

            auto v_first = v_macro.as_strided(
                {a.size(0), active_rows, kMacro1024},
                {v_macro.stride(0), kMacro64_1024, 1});
            auto c_next = h.slice(1, macro_start, kN1024).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<
                kMacro1024,
                kPanel1024,
                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>(),
                    kN1024,
                    macro_start,
                    kMacro1024);
            c_next.baddbmm_(v_first, w_local_workspace, 1.0, -1.0);

            const int second_active_rows = kN1024 - second_start;
            const int second_prep_active_shared_bytes =
                (second_active_rows * kPanel1024 + 3 * kPanel1024 * kPanel1024 + kPanel1024) *
                static_cast<int>(sizeof(float));
            for (int leaf = 0; leaf < 2; ++leaf) {
                const int leaf_offset = leaf * kPanel1024;
                const int panel_start = second_start + leaf_offset;
                const int macro_row_offset = kMacro1024 + leaf_offset;
                const int macro_col_offset = kMacro1024 + leaf_offset;
                float* v_second_base =
                    v_macro.data_ptr<float>() +
                    static_cast<long long>(kMacro1024) * kMacro64_1024 +
                    kMacro1024;
                float* g_second_base =
                    leaf_grams.data_ptr<float>() +
                    static_cast<long long>(2) * kPanel1024 * kPanel1024;
                if (leaf == 0) {
                    qr_panel_cached_kernel<kN1024, kPanel1024, kPanelThreads1024, false, false, false, false, true>
                        <<<batch, kPanelThreads1024, 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),
                            kMacro64_1024,
                            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),
                        kMacro64_1024,
                        second_active_rows,
                        batch,
                        true);

                    apply_prev_leaves_to_next_leaf_kernel<
                        kN1024,
                        kMacro64_1024,
                        kPanel1024,
                        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 + kPanel1024);
                } else {
                    qr_panel_cached_kernel<kN1024, kPanel1024, kPanelThreads1024, false, false, false, false, true>
                        <<<batch, kPanelThreads1024, 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),
                            kMacro64_1024,
                            macro_row_offset,
                            macro_col_offset);
                }
            }

            if (macro_end < kN1024) {
                const int trailing_cols = kN1024 - macro_end;
                at::bmm_out(g_macro, v_macro.transpose(1, 2), v_macro);
                auto w = w_workspace.as_strided(
                    {a.size(0), kMacro64_1024, trailing_cols},
                    {kMacro64_1024 * trailing_cols, trailing_cols, 1});
                auto c = h.slice(1, macro_start, kN1024).slice(2, macro_end, kN1024);
                if (macro_start == 0) {
                    auto c_in = a.slice(1, macro_start, kN1024).slice(2, macro_end, kN1024);
                    cublas_w_from_vt_c(
                        w.data_ptr<float>(),
                        c_in.data_ptr<float>(),
                        v_macro.data_ptr<float>(),
                        v_macro.stride(0),
                        kN1024,
                        kMacro64_1024,
                        kMacro64_1024,
                        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),
                        kN1024,
                        kMacro64_1024,
                        kMacro64_1024,
                        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<
                    kMacro64_1024,
                    kPanel1024,
                    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>(),
                        kN1024,
                        macro_start,
                        trailing_cols);

                if (macro_start == 0) {
                    auto c_in = a.slice(1, macro_start, kN1024).slice(2, macro_end, kN1024);
                    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>(),
                        kN1024,
                        kMacro64_1024,
                        active_rows,
                        trailing_cols,
                        batch,
                        true,
                        lt_workspace.data_ptr(),
                        lt_workspace_bytes);
                } else {
                    c.baddbmm_(v_macro, w, 1.0, -1.0);
                }
            }
        }
        at::globalContext().setAllowTF32CuBLAS(old_allow_tf32);
    }
    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 / 2;
    constexpr int k1 = kN512 / 2 + 2 * kPanel512;
    constexpr int k2 = (3 * kN512) / 4;
    constexpr int threads = 256;
    const int batch = static_cast<int>(a.size(0));

    // Cheap deterministic rejector. A rejection only selects the safe path;
    // acceptance still requires the complete per-matrix bound below.
    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);
    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_cholqr_hr512_cuda(
    torch::Tensor a,
    torch::Tensor r,
    torch::Tensor info) {
    TORCH_CHECK(a.is_cuda(), "a must be CUDA");
    TORCH_CHECK(r.is_cuda(), "r must be CUDA");
    TORCH_CHECK(info.is_cuda(), "info must be CUDA");
    TORCH_CHECK(a.scalar_type() == at::kFloat, "a must be float32");
    TORCH_CHECK(r.scalar_type() == at::kFloat, "r must be float32");
    TORCH_CHECK(info.scalar_type() == at::kInt, "info must be int32");
    TORCH_CHECK(a.dim() == 3 && a.size(1) == kN512 && a.size(2) == kN512,
                "CholQR-HR expects a [batch,512,512] input");
    TORCH_CHECK(r.dim() == 3 && r.size(0) == a.size(0) && r.size(1) == kN512 && r.size(2) == kN512,
                "CholQR-HR expects r [batch,512,512]");
    TORCH_CHECK(info.dim() == 1 && info.size(0) == a.size(0),
                "CholQR-HR expects info [batch]");
    TORCH_CHECK(a.is_contiguous(), "a must be contiguous");
    TORCH_CHECK(r.is_contiguous(), "r must be contiguous");
    TORCH_CHECK(info.is_contiguous(), "info must be contiguous");

    auto h = a.clone();
    auto tau = torch::empty({a.size(0), kN512}, a.options());
    auto ok = torch::empty({a.size(0)}, a.options().dtype(at::kInt));
    constexpr int threads = 256;
    const int batch = static_cast<int>(a.size(0));
    cholqr_hr_lu_kernel<kN512, threads><<<batch, threads, 0>>>(
        r.data_ptr<float>(),
        info.data_ptr<int>(),
        h.data_ptr<float>(),
        tau.data_ptr<float>(),
        ok.data_ptr<int>());
    C10_CUDA_CHECK(cudaGetLastError());
    return {h, tau, ok};
}

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;
    constexpr int prep_shared_bytes =
        (kN512 * kPanel512 + 3 * kPanel512 * kPanel512 + kPanel512) * sizeof(float);
    if (factor_cols >= kMacroPrefix512) {
        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);
    } else {
        copy_prefix_zero_suffix_v4_kernel<kN512, copy_threads><<<copy_blocks, copy_threads, 0>>>(
            a.data_ptr<float>(),
            h.data_ptr<float>(),
            static_cast<long long>(batch) * kN512 * kN512,
            factor_cols);
    }
    if ((factor_cols % kMacroPrefix512) == 0) {
        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 = 4 * 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;
                const int leaf_active_rows = kN512 - panel_start;
                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);
                    at::bmm_out(w, v_macro.transpose(1, 2), c_in);
                } else {
                    at::bmm_out(w, v_macro.transpose(1, 2), c);
                }

                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);
                }
            }
        }
    } else {
        auto t_scratch = torch::empty({a.size(0), kPanel512, kPanel512}, a.options());
        auto w_workspace = torch::empty({a.size(0), kPanel512, factor_cols}, a.options());
        for (int panel_start = 0; panel_start < factor_cols; panel_start += kPanel512) {
            const int panel_end = panel_start + kPanel512;
            if (panel_end < factor_cols) {
                const int active_rows = kN512 - panel_start;
                const int trailing_cols = factor_cols - panel_end;
                // Keep V compact. A max-stride slab measurably regressed the batched GEMMs.
                auto v = torch::empty({a.size(0), active_rows, kPanel512}, a.options());
                qr_panel_cached_kernel<kN512, kPanel512, kPanelThreads512, true, false, false, false>
                    <<<batch, kPanelThreads512, panel_shared_bytes>>>(
                        h.data_ptr<float>(),
                        tau.data_ptr<float>(),
                        t_scratch.data_ptr<float>(),
                        v.data_ptr<float>(),
                        v.stride(0),
                        panel_start);
                at::globalContext().setAllowTF32CuBLAS(false);
                at::bmm_out(t_scratch, v.transpose(1, 2), v);
                at::globalContext().setAllowTF32CuBLAS(old_allow_tf32);

                auto w = w_workspace.as_strided(
                    {a.size(0), kPanel512, trailing_cols},
                    {kPanel512 * trailing_cols, trailing_cols, 1});
                auto z = w;
                auto c = h.slice(1, panel_start, kN512).slice(2, panel_end, factor_cols);
                at::bmm_out(w, v.transpose(1, 2), c);

                constexpr int apply_t_threads = 128;
                const int apply_t_blocks = (trailing_cols + apply_t_threads - 1) / apply_t_threads;
                solve_inverse_wy_from_gram_kernel<kPanel512, apply_t_threads>
                    <<<dim3(apply_t_blocks, batch), apply_t_threads, 0>>>(
                        t_scratch.data_ptr<float>(),
                        tau.data_ptr<float>(),
                        w.data_ptr<float>(),
                        z.data_ptr<float>(),
                        kN512,
                        panel_start,
                        trailing_cols);

                c.baddbmm_(v, z, 1.0, -1.0);
            } else {
                qr_panel_cached_kernel<kN512, kPanel512, kPanelThreads512, false, false, false>
                    <<<batch, kPanelThreads512, panel_shared_bytes>>>(
                        h.data_ptr<float>(),
                        tau.data_ptr<float>(),
                        t_scratch.data_ptr<float>(),
                        nullptr,
                        0,
                        panel_start);
            }
        }
    }
    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 int max_panel_shared_bytes = ((kN2048 + 1) * kPanel2048 + kPanelThreads2048) * sizeof(float);
    constexpr int copy_threads = 256;
    constexpr int copy_blocks = 1024;
    static bool panel_attrs_set = false;
    if (!panel_attrs_set) {
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            qr_panel_cached_kernel<kN2048, kPanel2048, kPanelThreads2048, false, true>,
            cudaFuncAttributeMaxDynamicSharedMemorySize,
            max_panel_shared_bytes));
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            qr_panel_cached_kernel<kN2048, kPanel2048, kPanelThreads2048, false, true>,
            cudaFuncAttributePreferredSharedMemoryCarveout,
            100));
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            qr_panel_cached_kernel<kN2048, kPanel2048, kPanelThreads2048, true, true>,
            cudaFuncAttributeMaxDynamicSharedMemorySize,
            max_panel_shared_bytes));
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            qr_panel_cached_kernel<kN2048, kPanel2048, kPanelThreads2048, true, true>,
            cudaFuncAttributePreferredSharedMemoryCarveout,
            100));
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            qr_panel_cached_kernel<kN2048, kPanel2048, kPanelThreads2048, true, true, true, true, true>,
            cudaFuncAttributeMaxDynamicSharedMemorySize,
            max_panel_shared_bytes));
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            qr_panel_cached_kernel<kN2048, kPanel2048, kPanelThreads2048, true, true, true, true, true>,
            cudaFuncAttributePreferredSharedMemoryCarveout,
            100));
        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;
            auto v_leaf = torch::empty({a.size(0), leaf_active_rows, kPanel2048}, a.options());
            qr_panel_cached_kernel<kN2048, kPanel2048, kPanelThreads2048, true, true, true, true, true>
                <<<batch, kPanelThreads2048, panel_shared_bytes>>>(
                    h.data_ptr<float>(),
                    tau.data_ptr<float>(),
                    t_scratch.data_ptr<float>(),
                    v_leaf.data_ptr<float>(),
                    v_leaf.stride(0),
                    panel_start,
                    v_macro.data_ptr<float>(),
                    v_macro.stride(0),
                    macro_cols,
                    macro_offset,
                    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,
                    kPanel2048,
                    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);
                c.baddbmm_(v_leaf, z, 1.0, -1.0);
            }
        }

        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 = 64;
            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);
            c.baddbmm_(v_macro, w, 1.0, -1.0);
        }
    }

    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 = 8 * kPanel4096;
    auto t_first = torch::empty({a.size(0), kPanel4096, kPanel4096}, a.options());
    auto t_second = 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 int panel_shared_bytes = ((kN4096 + 1) * kPanel4096 + kPanelThreads4096) * sizeof(float);
    constexpr int apply_threads = 512;
    constexpr int apply_shared_bytes =
        (kN4096 * kPanel4096 + (apply_threads / 32) * kPanel4096 + kPanel4096) * sizeof(float);
    constexpr int build_threads = 512;
    constexpr int copy_threads = 256;
    constexpr int copy_blocks = 8192;
    static int late_max_rows = kPanel4096LateMaxRows;
    static int late_panel_shared_bytes =
        ((kPanel4096LateMaxRows + 1) * kPanel4096Late + kPanelThreads4096Late) * sizeof(float);
    static bool panel_attrs_set = false;
    if (!panel_attrs_set) {
        int device = 0;
        int max_shared_bytes = late_panel_shared_bytes;
        C10_CUDA_CHECK(cudaGetDevice(&device));
        C10_CUDA_CHECK(cudaDeviceGetAttribute(
            &max_shared_bytes,
            cudaDevAttrMaxSharedMemoryPerBlockOptin,
            device));
        int device_rows =
            ((max_shared_bytes / static_cast<int>(sizeof(float))) - kPanelThreads4096Late) /
            kPanel4096Late - 1;
        if (device_rows > kPanel4096LateMaxRows) {
            device_rows = kPanel4096LateMaxRows;
        }
        if (device_rows < 0) {
            device_rows = 0;
        }
        late_max_rows = (device_rows / super_panel) * super_panel;
        late_panel_shared_bytes =
            ((late_max_rows + 1) * kPanel4096Late + kPanelThreads4096Late) * sizeof(float);

        C10_CUDA_CHECK(cudaFuncSetAttribute(
            qr_panel_cached_kernel<kN4096, kPanel4096, kPanelThreads4096, false>,
            cudaFuncAttributeMaxDynamicSharedMemorySize,
            panel_shared_bytes));
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            qr_panel_cached_kernel<kN4096, kPanel4096, kPanelThreads4096, false>,
            cudaFuncAttributePreferredSharedMemoryCarveout,
            100));
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            qr_panel_cached_kernel<kN4096, kPanel4096, kPanelThreads4096, true, false, true, true, true>,
            cudaFuncAttributeMaxDynamicSharedMemorySize,
            panel_shared_bytes));
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            qr_panel_cached_kernel<kN4096, kPanel4096, kPanelThreads4096, true, false, true, true, true>,
            cudaFuncAttributePreferredSharedMemoryCarveout,
            100));
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            qr_panel_cached_kernel<kN4096, kPanel4096Late, kPanelThreads4096Late, false, true>,
            cudaFuncAttributeMaxDynamicSharedMemorySize,
            late_panel_shared_bytes));
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            qr_panel_cached_kernel<kN4096, kPanel4096Late, kPanelThreads4096Late, false, true>,
            cudaFuncAttributePreferredSharedMemoryCarveout,
            100));
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            qr_panel_cached_kernel<kN4096, kPanel4096Late, kPanelThreads4096Late, true, true>,
            cudaFuncAttributeMaxDynamicSharedMemorySize,
            late_panel_shared_bytes));
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            qr_panel_cached_kernel<kN4096, kPanel4096Late, kPanelThreads4096Late, true, true>,
            cudaFuncAttributePreferredSharedMemoryCarveout,
            100));
        panel_attrs_set = true;
    }
    static bool apply_attrs_set = false;
    if (!apply_attrs_set) {
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            apply_panel_to_next_cols_kernel<kN4096, kPanel4096, kPanel4096, apply_threads>,
            cudaFuncAttributeMaxDynamicSharedMemorySize,
            apply_shared_bytes));
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            apply_panel_to_next_cols_kernel<kN4096, kPanel4096, kPanel4096, apply_threads>,
            cudaFuncAttributePreferredSharedMemoryCarveout,
            100));
        apply_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;
            auto v_leaf = torch::empty({a.size(0), leaf_active_rows, kPanel4096}, a.options());
            qr_panel_cached_kernel<kN4096, kPanel4096, kPanelThreads4096, true, false, true, true, true>
                <<<batch, kPanelThreads4096, panel_shared_bytes>>>(
                    h.data_ptr<float>(),
                    tau.data_ptr<float>(),
                    t_first.data_ptr<float>(),
                    v_leaf.data_ptr<float>(),
                    v_leaf.stride(0),
                    leaf_start,
                    v_macro.data_ptr<float>(),
                    v_macro.stride(0),
                    early_macro,
                    macro_offset,
                    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,
                    kPanel4096,
                    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);
                c.baddbmm_(v_leaf, z, 1.0, -1.0);
            }
        }

        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);
            constexpr int apply_t_threads = 64;
            const int apply_t_blocks = (trailing_cols + apply_t_threads - 1) / apply_t_threads;
            solve_inverse_wy_from_gram_blocked_kernel<
                early_macro,
                kPanel4096,
                apply_t_threads>
                <<<dim3(apply_t_blocks, batch), apply_t_threads, 0>>>(
                    g_early.data_ptr<float>(),
                    tau.data_ptr<float>(),
                    w.data_ptr<float>(),
                    w.data_ptr<float>(),
                    kN4096,
                    panel_start,
                    trailing_cols);
            c.baddbmm_(v_macro, w, 1.0, -1.0);
        }
        panel_start += early_macro;
    }

    return {h, tau};
}

"""



_EXT = None
_EXT_FAILED = False
def _load_ext():
    global _EXT, _EXT_FAILED

    if _EXT is None and not _EXT_FAILED:
        try:
            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,
            )
        except Exception as exc:
            _EXT_FAILED = True
            raise RuntimeError("qr_householder extension build failed") from exc
    return _EXT


def _cholqr_hr512_experiment(contiguous: torch.Tensor, ext) -> output_t:
    old_allow_tf32 = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = False
    try:
        gram = torch.bmm(contiguous.transpose(1, 2), contiguous)
        r, info = torch.linalg.cholesky_ex(gram, upper=True, check_errors=False)
    finally:
        torch.backends.cuda.matmul.allow_tf32 = old_allow_tf32

    cand_h, cand_tau, ok = ext.qr_cholqr_hr512(
        contiguous,
        r.contiguous(),
        info.contiguous(),
    )
    ok_mask = ok.to(dtype=torch.bool)
    if bool(ok_mask.all()):
        return cand_h, cand_tau
    safe_h, safe_tau = ext.qr_small(contiguous)
    h = torch.where(ok_mask.view(-1, 1, 1), cand_h, safe_h)
    tau = torch.where(ok_mask.view(-1, 1), cand_tau, safe_tau)
    return h, tau


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 (
        data.is_cuda
        and data.dtype == torch.float32
        and data.dim() == 3
        and data.shape[-1] == data.shape[-2]
        and data.shape[-1] in (32, 176, 352, 512, 1024)
    ):
        return _small_qr(data)
    if (
        data.is_cuda
        and data.dtype == torch.float32
        and data.dim() == 3
        and data.shape[-1] == data.shape[-2]
        and data.shape[-1] in (2048, 4096)
    ):
        return _large_qr(data)
    return torch.ops.aten.geqrf.default(data)
scrolls · 3814 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