Skip to content
KernelIndex
Search⌘K

submission 808705

trxonphoenix · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-808705?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
6.29ms
#218 of 515
2026-06-17

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:7be935bf0b5e65cec3d01deebec0cb01b05ee7ae19a3f94fcadd24cb32d7aea9
license declaredunknown
license concludedunknown
authorstrxonphoenix
imported2026-08-26

Techniques

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

shared-memory__shared__ float a[MAX_N * MAX_N];

Kernel source

submission.py1523 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200

import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t


CPP_SRC = """
void qr32(torch::Tensor a, torch::Tensor h, torch::Tensor tau);
void qr_shared_176(torch::Tensor a, torch::Tensor h, torch::Tensor tau);
void qr_global_256(torch::Tensor a, torch::Tensor h, torch::Tensor tau);
void qr_global_512(torch::Tensor a, torch::Tensor h, torch::Tensor tau);
void qr_blocked_wy(torch::Tensor a, torch::Tensor h, torch::Tensor tau, int nb, int custom_small);
"""


CUDA_SRC = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cublas_v2.h>
#include <cusolverDn.h>
#include <cmath>
#include <stdexcept>
#include <string>

namespace {

constexpr int MAX_N = 32;
constexpr int PANEL_THREADS = 512;

inline void check_cuda(cudaError_t status, const char* where) {
    if (status != cudaSuccess) {
        throw std::runtime_error(std::string(where) + ": " + cudaGetErrorString(status));
    }
}

inline void check_blas(cublasStatus_t status, const char* where) {
    if (status != CUBLAS_STATUS_SUCCESS) {
        throw std::runtime_error(std::string(where) + ": cuBLAS status " + std::to_string(status));
    }
}

inline void check_solver(cusolverStatus_t status, const char* where) {
    if (status != CUSOLVER_STATUS_SUCCESS) {
        throw std::runtime_error(std::string(where) + ": cuSOLVER status " + std::to_string(status));
    }
}

cublasHandle_t get_blas_handle() {
    static cublasHandle_t handle = nullptr;
    static bool initialized = false;
    if (!initialized) {
        check_blas(cublasCreate(&handle), "cublasCreate");
#ifdef CUBLAS_TF32_TENSOR_OP_MATH
        check_blas(cublasSetMathMode(handle, CUBLAS_TF32_TENSOR_OP_MATH), "cublasSetMathMode");
#endif
        initialized = true;
    }
    return handle;
}

template <int QR_THREADS>
__global__ void qr32_kernel(const float* __restrict__ a_in,
                            float* __restrict__ h_out,
                            float* __restrict__ tau_out,
                            int batch,
                            int n) {
    const int b = blockIdx.x;
    const int tid = threadIdx.x;

    __shared__ float a[MAX_N * MAX_N];
    __shared__ float red[QR_THREADS];
    __shared__ float tau_s;
    __shared__ float scale_s;

    const float* src = a_in + static_cast<long long>(b) * n * n;
    float* dst = h_out + static_cast<long long>(b) * n * n;
    float* tau_dst = tau_out + static_cast<long long>(b) * n;

    for (int idx = tid; idx < n * n; idx += blockDim.x) {
        a[idx] = src[idx];
    }
    __syncthreads();

    for (int k = 0; k < n; ++k) {
        float local = 0.0f;
        for (int i = k + 1 + tid; i < n; i += blockDim.x) {
            const float x = a[i * n + k];
            local = fmaf(x, x, local);
        }
        red[tid] = local;
        __syncthreads();

        for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
            if (tid < stride) {
                red[tid] += red[tid + stride];
            }
            __syncthreads();
        }

        if (tid == 0) {
            const float alpha = a[k * n + k];
            const float xnorm2 = red[0];
            if (xnorm2 == 0.0f) {
                tau_s = 0.0f;
                scale_s = 0.0f;
                tau_dst[k] = 0.0f;
            } else {
                const float norm = sqrtf(fmaf(alpha, alpha, xnorm2));
                const float beta = (alpha >= 0.0f) ? -norm : norm;
                const float tau = (beta - alpha) / beta;
                tau_s = tau;
                scale_s = 1.0f / (alpha - beta);
                a[k * n + k] = beta;
                tau_dst[k] = tau;
            }
        }
        __syncthreads();

        if (tau_s != 0.0f) {
            for (int i = k + 1 + tid; i < n; i += blockDim.x) {
                a[i * n + k] *= scale_s;
            }
        }
        __syncthreads();

        const float tau = tau_s;
        for (int j = k + 1 + tid; j < n; j += blockDim.x) {
            float dot = a[k * n + j];
            for (int i = k + 1; i < n; ++i) {
                dot = fmaf(a[i * n + k], a[i * n + j], dot);
            }
            const float w = tau * dot;
            a[k * n + j] -= w;
            for (int i = k + 1; i < n; ++i) {
                a[i * n + j] = fmaf(-a[i * n + k], w, a[i * n + j]);
            }
        }
        __syncthreads();
    }

    for (int idx = tid; idx < n * n; idx += blockDim.x) {
        dst[idx] = a[idx];
    }
}

template <int QR_THREADS>
__global__ void qr_shared_kernel(const float* __restrict__ a_in,
                                 float* __restrict__ h_out,
                                 float* __restrict__ tau_out,
                                 int n) {
    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    const long long matrix_size = static_cast<long long>(n) * n;
    const float* src = a_in + static_cast<long long>(b) * matrix_size;
    float* dst = h_out + static_cast<long long>(b) * matrix_size;
    float* tau_dst = tau_out + static_cast<long long>(b) * n;

    extern __shared__ float a[];
    __shared__ float red[QR_THREADS];
    __shared__ float tau_s;
    __shared__ float scale_s;

    for (long long idx = tid; idx < matrix_size; idx += blockDim.x) {
        a[idx] = src[idx];
    }
    __syncthreads();

    for (int k = 0; k < n; ++k) {
        float local = 0.0f;
        for (int i = k + 1 + tid; i < n; i += blockDim.x) {
            const float x = a[static_cast<long long>(i) * n + k];
            local = fmaf(x, x, local);
        }
        red[tid] = local;
        __syncthreads();

        for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
            if (tid < stride) {
                red[tid] += red[tid + stride];
            }
            __syncthreads();
        }

        if (tid == 0) {
            const long long diag = static_cast<long long>(k) * n + k;
            const float alpha = a[diag];
            const float xnorm2 = red[0];
            if (xnorm2 == 0.0f) {
                tau_s = 0.0f;
                scale_s = 0.0f;
                tau_dst[k] = 0.0f;
            } else {
                const float norm = sqrtf(fmaf(alpha, alpha, xnorm2));
                const float beta = (alpha >= 0.0f) ? -norm : norm;
                const float tau = (beta - alpha) / beta;
                tau_s = tau;
                scale_s = 1.0f / (alpha - beta);
                a[diag] = beta;
                tau_dst[k] = tau;
            }
        }
        __syncthreads();

        if (tau_s != 0.0f) {
            const float scale = scale_s;
            for (int i = k + 1 + tid; i < n; i += blockDim.x) {
                a[static_cast<long long>(i) * n + k] *= scale;
            }
        }
        __syncthreads();

        const float tau = tau_s;
        for (int j = k + 1 + tid; j < n; j += blockDim.x) {
            float dot = a[static_cast<long long>(k) * n + j];
            for (int i = k + 1; i < n; ++i) {
                dot = fmaf(a[static_cast<long long>(i) * n + k],
                           a[static_cast<long long>(i) * n + j],
                           dot);
            }
            const float w = tau * dot;
            a[static_cast<long long>(k) * n + j] -= w;
            for (int i = k + 1; i < n; ++i) {
                const long long idx = static_cast<long long>(i) * n + j;
                a[idx] = fmaf(-a[static_cast<long long>(i) * n + k], w, a[idx]);
            }
        }
        __syncthreads();
    }

    for (long long idx = tid; idx < matrix_size; idx += blockDim.x) {
        dst[idx] = a[idx];
    }
}

template <int QR_THREADS>
__global__ void qr_global_kernel(const float* __restrict__ a_in,
                                 float* __restrict__ h_out,
                                 float* __restrict__ tau_out,
                                 int n) {
    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    const long long matrix_size = static_cast<long long>(n) * n;
    const float* src = a_in + static_cast<long long>(b) * matrix_size;
    float* a = h_out + static_cast<long long>(b) * matrix_size;
    float* tau_dst = tau_out + static_cast<long long>(b) * n;

    __shared__ float red[QR_THREADS];
    __shared__ float tau_s;
    __shared__ float scale_s;

    for (long long idx = tid; idx < matrix_size; idx += blockDim.x) {
        a[idx] = src[idx];
    }
    __syncthreads();

    for (int k = 0; k < n; ++k) {
        float local = 0.0f;
        for (int i = k + 1 + tid; i < n; i += blockDim.x) {
            const float x = a[static_cast<long long>(i) * n + k];
            local = fmaf(x, x, local);
        }
        red[tid] = local;
        __syncthreads();

        for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
            if (tid < stride) {
                red[tid] += red[tid + stride];
            }
            __syncthreads();
        }

        if (tid == 0) {
            const long long diag = static_cast<long long>(k) * n + k;
            const float alpha = a[diag];
            const float xnorm2 = red[0];
            if (xnorm2 == 0.0f) {
                tau_s = 0.0f;
                scale_s = 0.0f;
                tau_dst[k] = 0.0f;
            } else {
                const float norm = sqrtf(fmaf(alpha, alpha, xnorm2));
                const float beta = (alpha >= 0.0f) ? -norm : norm;
                const float tau = (beta - alpha) / beta;
                tau_s = tau;
                scale_s = 1.0f / (alpha - beta);
                a[diag] = beta;
                tau_dst[k] = tau;
            }
        }
        __syncthreads();

        if (tau_s != 0.0f) {
            for (int i = k + 1 + tid; i < n; i += blockDim.x) {
                a[static_cast<long long>(i) * n + k] *= scale_s;
            }
        }
        __syncthreads();

        const float tau = tau_s;
        for (int j = k + 1 + tid; j < n; j += blockDim.x) {
            float dot = a[static_cast<long long>(k) * n + j];
            for (int i = k + 1; i < n; ++i) {
                dot = fmaf(a[static_cast<long long>(i) * n + k],
                           a[static_cast<long long>(i) * n + j],
                           dot);
            }
            const float w = tau * dot;
            a[static_cast<long long>(k) * n + j] -= w;
            for (int i = k + 1; i < n; ++i) {
                const long long idx = static_cast<long long>(i) * n + j;
                a[idx] = fmaf(-a[static_cast<long long>(i) * n + k], w, a[idx]);
            }
        }
        __syncthreads();
    }
}

__global__ void row_to_col_major_kernel(const float* __restrict__ row,
                                        float* __restrict__ col,
                                        int n,
                                        long long total) {
    const long long idx = static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
    if (idx >= total) {
        return;
    }
    const long long matrix_size = static_cast<long long>(n) * n;
    const long long b = idx / matrix_size;
    const long long rem = idx - b * matrix_size;
    const int i = static_cast<int>(rem / n);
    const int j = static_cast<int>(rem - static_cast<long long>(i) * n);
    col[b * matrix_size + i + static_cast<long long>(j) * n] =
        row[b * matrix_size + static_cast<long long>(i) * n + j];
}

__global__ void col_to_row_major_kernel(const float* __restrict__ col,
                                        float* __restrict__ row,
                                        int n,
                                        long long total) {
    const long long idx = static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
    if (idx >= total) {
        return;
    }
    const long long matrix_size = static_cast<long long>(n) * n;
    const long long b = idx / matrix_size;
    const long long rem = idx - b * matrix_size;
    const int i = static_cast<int>(rem / n);
    const int j = static_cast<int>(rem - static_cast<long long>(i) * n);
    row[b * matrix_size + static_cast<long long>(i) * n + j] =
        col[b * matrix_size + i + static_cast<long long>(j) * n];
}

__global__ void transpose_tiled_kernel(const float* __restrict__ src,
                                       float* __restrict__ dst,
                                       int n) {
    constexpr int TILE_DIM = 32;
    constexpr int BLOCK_ROWS = 8;
    __shared__ float tile[TILE_DIM][TILE_DIM + 1];

    const int b = blockIdx.z;
    const long long matrix_size = static_cast<long long>(n) * n;
    const float* src_b = src + static_cast<long long>(b) * matrix_size;
    float* dst_b = dst + static_cast<long long>(b) * matrix_size;

    int x = blockIdx.x * TILE_DIM + threadIdx.x;
    int y = blockIdx.y * TILE_DIM + threadIdx.y;
    for (int j = 0; j < TILE_DIM; j += BLOCK_ROWS) {
        if (x < n && y + j < n) {
            tile[threadIdx.y + j][threadIdx.x] =
                src_b[static_cast<long long>(y + j) * n + x];
        }
    }
    __syncthreads();

    x = blockIdx.y * TILE_DIM + threadIdx.x;
    y = blockIdx.x * TILE_DIM + threadIdx.y;
    for (int j = 0; j < TILE_DIM; j += BLOCK_ROWS) {
        if (x < n && y + j < n) {
            dst_b[static_cast<long long>(y + j) * n + x] =
                tile[threadIdx.x][threadIdx.y + j];
        }
    }
}

__global__ void col_to_row_major_restore_v_kernel(const float* __restrict__ col,
                                                  float* __restrict__ row,
                                                  const float* __restrict__ saved,
                                                  int n,
                                                  int nb,
                                                  long long total) {
    const long long idx = static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
    if (idx >= total) {
        return;
    }
    const long long matrix_size = static_cast<long long>(n) * n;
    const long long b = idx / matrix_size;
    const long long rem = idx - b * matrix_size;
    const int i = static_cast<int>(rem / n);
    const int j = static_cast<int>(rem - static_cast<long long>(i) * n);

    float value = col[b * matrix_size + i + static_cast<long long>(j) * n];
    const int panel_k = (j / nb) * nb;
    const int r = i - panel_k;
    const int s = j - panel_k;
    if (r >= 0 && r <= s) {
        const long long save_stride = static_cast<long long>(n) * nb;
        value = saved[b * save_stride + static_cast<long long>(panel_k) * nb +
                      r + static_cast<long long>(s) * nb];
    }
    row[b * matrix_size + static_cast<long long>(i) * n + j] = value;
}

__global__ void transpose_tiled_restore_v_kernel(const float* __restrict__ col,
                                                 float* __restrict__ row,
                                                 const float* __restrict__ saved,
                                                 int n,
                                                 int nb) {
    constexpr int TILE_DIM = 32;
    constexpr int BLOCK_ROWS = 8;
    __shared__ float tile[TILE_DIM][TILE_DIM + 1];

    const int b = blockIdx.z;
    const long long matrix_size = static_cast<long long>(n) * n;
    const float* col_b = col + static_cast<long long>(b) * matrix_size;
    float* row_b = row + static_cast<long long>(b) * matrix_size;

    int x = blockIdx.x * TILE_DIM + threadIdx.x;
    int y = blockIdx.y * TILE_DIM + threadIdx.y;
    for (int j = 0; j < TILE_DIM; j += BLOCK_ROWS) {
        if (x < n && y + j < n) {
            tile[threadIdx.y + j][threadIdx.x] =
                col_b[static_cast<long long>(y + j) * n + x];
        }
    }
    __syncthreads();

    x = blockIdx.y * TILE_DIM + threadIdx.x;
    y = blockIdx.x * TILE_DIM + threadIdx.y;
    for (int j = 0; j < TILE_DIM; j += BLOCK_ROWS) {
        const int row_idx = y + j;
        const int col_idx = x;
        if (row_idx < n && col_idx < n) {
            float value = tile[threadIdx.x][threadIdx.y + j];
            const int panel_k = (col_idx / nb) * nb;
            const int r = row_idx - panel_k;
            const int s = col_idx - panel_k;
            if (r >= 0 && r <= s) {
                const long long save_stride = static_cast<long long>(n) * nb;
                value = saved[static_cast<long long>(b) * save_stride +
                              static_cast<long long>(panel_k) * nb +
                              r + static_cast<long long>(s) * nb];
            }
            row_b[static_cast<long long>(row_idx) * n + col_idx] = value;
        }
    }
}

__global__ void build_explicit_v_kernel(const float* __restrict__ col,
                                        float* __restrict__ v,
                                        int n,
                                        int k,
                                        int m,
                                        int ib,
                                        long long stride_v,
                                        long long total) {
    const long long idx = static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
    if (idx >= total) {
        return;
    }
    const long long per_batch = static_cast<long long>(m) * ib;
    const long long b = idx / per_batch;
    const long long rem = idx - b * per_batch;
    const int r = static_cast<int>(rem % m);
    const int s = static_cast<int>(rem / m);
    const long long matrix_size = static_cast<long long>(n) * n;
    const float* mat = col + b * matrix_size;
    float value = 0.0f;
    if (r == s) {
        value = 1.0f;
    } else if (r > s) {
        value = mat[static_cast<long long>(k + r) + static_cast<long long>(k + s) * n];
    }
    v[b * stride_v + r + static_cast<long long>(s) * m] = value;
}

__global__ void prepare_panel_v_inplace_kernel(float* __restrict__ col,
                                               float* __restrict__ saved,
                                               int n,
                                               int k,
                                               int ib,
                                               long long stride_small) {
    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    const long long matrix_size = static_cast<long long>(n) * n;
    float* mat = col + static_cast<long long>(b) * matrix_size;
    float* sb = saved + static_cast<long long>(b) * stride_small;

    for (int idx = tid; idx < ib * ib; idx += blockDim.x) {
        const int r = idx % ib;
        const int s = idx / ib;
        const long long mat_idx = static_cast<long long>(k + r) +
                                  static_cast<long long>(k + s) * n;
        const float original = mat[mat_idx];
        if (r <= s) {
            sb[idx] = original;
        }
        if (r < s) {
            mat[mat_idx] = 0.0f;
        } else if (r == s) {
            mat[mat_idx] = 1.0f;
        }
    }
}

__global__ void restore_panel_r_kernel(float* __restrict__ col,
                                       const float* __restrict__ saved,
                                       int n,
                                       int k,
                                       int ib,
                                       long long stride_small) {
    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    const long long matrix_size = static_cast<long long>(n) * n;
    float* mat = col + static_cast<long long>(b) * matrix_size;
    const float* sb = saved + static_cast<long long>(b) * stride_small;

    for (int idx = tid; idx < ib * ib; idx += blockDim.x) {
        const int r = idx % ib;
        const int s = idx / ib;
        const long long mat_idx = static_cast<long long>(k + r) +
                                  static_cast<long long>(k + s) * n;
        if (r <= s) {
            mat[mat_idx] = sb[idx];
        }
    }
}

__global__ void build_t_from_gram_kernel(const float* __restrict__ gram,
                                         float* __restrict__ t,
                                         const float* __restrict__ tau,
                                         int n,
                                         int k,
                                         int ib,
                                         long long stride_small) {
    const int b = blockIdx.x;
    const float* g = gram + static_cast<long long>(b) * stride_small;
    float* tb = t + static_cast<long long>(b) * stride_small;
    const float* tau_b = tau + static_cast<long long>(b) * n + k;

    for (int idx = 0; idx < ib * ib; ++idx) {
        tb[idx] = 0.0f;
    }

    float work[64];
    for (int i = 0; i < ib; ++i) {
        const float tau_i = tau_b[i];
        if (tau_i == 0.0f) {
            tb[i + static_cast<long long>(i) * ib] = 0.0f;
            continue;
        }

        for (int j = 0; j < i; ++j) {
            work[j] = -tau_i * g[j + static_cast<long long>(i) * ib];
        }

        for (int row = 0; row < i; ++row) {
            float sum = 0.0f;
            for (int col = row; col < i; ++col) {
                sum = fmaf(tb[row + static_cast<long long>(col) * ib], work[col], sum);
            }
            tb[row + static_cast<long long>(i) * ib] = sum;
        }
        tb[i + static_cast<long long>(i) * ib] = tau_i;
    }
}

__global__ void build_gram_small_kernel(const float* __restrict__ v,
                                        float* __restrict__ gram,
                                        int m,
                                        int ib,
                                        long long stride_v,
                                        long long stride_small) {
    const int b = blockIdx.x;
    const int row = blockIdx.y;
    const int col = blockIdx.z;
    const int tid = threadIdx.x;

    __shared__ float red[PANEL_THREADS];
    const float* vb = v + static_cast<long long>(b) * stride_v;
    float local = 0.0f;
    for (int r = tid; r < m; r += blockDim.x) {
        local = fmaf(vb[r + static_cast<long long>(row) * m],
                     vb[r + static_cast<long long>(col) * m],
                     local);
    }
    red[tid] = local;
    __syncthreads();

    for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
        if (tid < stride) {
            red[tid] += red[tid + stride];
        }
        __syncthreads();
    }

    if (tid == 0) {
        gram[static_cast<long long>(b) * stride_small + row + static_cast<long long>(col) * ib] = red[0];
    }
}

__global__ void build_gram_strided_small_kernel(const float* __restrict__ v,
                                                float* __restrict__ gram,
                                                int m,
                                                int ib,
                                                int ldv,
                                                long long stride_v,
                                                long long stride_small) {
    const int b = blockIdx.x;
    const int row = blockIdx.y;
    const int col = blockIdx.z;
    const int tid = threadIdx.x;

    __shared__ float red[PANEL_THREADS];
    const float* vb = v + static_cast<long long>(b) * stride_v;
    float local = 0.0f;
    for (int r = tid; r < m; r += blockDim.x) {
        local = fmaf(vb[r + static_cast<long long>(row) * ldv],
                     vb[r + static_cast<long long>(col) * ldv],
                     local);
    }
    red[tid] = local;
    __syncthreads();

    for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
        if (tid < stride) {
            red[tid] += red[tid + stride];
        }
        __syncthreads();
    }

    if (tid == 0) {
        gram[static_cast<long long>(b) * stride_small + row + static_cast<long long>(col) * ib] = red[0];
    }
}

__global__ void panel_factor_colmajor_kernel(float* __restrict__ col,
                                             float* __restrict__ tau,
                                             int n,
                                             int k0,
                                             int ib) {
    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    const long long matrix_size = static_cast<long long>(n) * n;
    float* mat = col + static_cast<long long>(b) * matrix_size;
    float* tau_b = tau + static_cast<long long>(b) * n;

    __shared__ float red[PANEL_THREADS];
    __shared__ float tau_s;
    __shared__ float scale_s;
    __shared__ float w_s;

    for (int s = 0; s < ib; ++s) {
        const int k = k0 + s;

        float local = 0.0f;
        for (int i = k + 1 + tid; i < n; i += blockDim.x) {
            const float x = mat[i + static_cast<long long>(k) * n];
            local = fmaf(x, x, local);
        }
        red[tid] = local;
        __syncthreads();

        for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
            if (tid < stride) {
                red[tid] += red[tid + stride];
            }
            __syncthreads();
        }

        if (tid == 0) {
            const long long diag = static_cast<long long>(k) + static_cast<long long>(k) * n;
            const float alpha = mat[diag];
            const float xnorm2 = red[0];
            if (xnorm2 == 0.0f) {
                tau_s = 0.0f;
                scale_s = 0.0f;
                tau_b[k] = 0.0f;
            } else {
                const float norm = sqrtf(fmaf(alpha, alpha, xnorm2));
                const float beta = (alpha >= 0.0f) ? -norm : norm;
                const float tau_val = (beta - alpha) / beta;
                tau_s = tau_val;
                scale_s = 1.0f / (alpha - beta);
                mat[diag] = beta;
                tau_b[k] = tau_val;
            }
        }
        __syncthreads();

        if (tau_s != 0.0f) {
            const float scale = scale_s;
            for (int i = k + 1 + tid; i < n; i += blockDim.x) {
                mat[i + static_cast<long long>(k) * n] *= scale;
            }
        }
        __syncthreads();

        for (int jj = s + 1; jj < ib; ++jj) {
            const int j = k0 + jj;
            float dot_local = 0.0f;
            for (int i = k + 1 + tid; i < n; i += blockDim.x) {
                dot_local = fmaf(mat[i + static_cast<long long>(k) * n],
                                 mat[i + static_cast<long long>(j) * n],
                                 dot_local);
            }
            red[tid] = dot_local;
            __syncthreads();

            for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
                if (tid < stride) {
                    red[tid] += red[tid + stride];
                }
                __syncthreads();
            }

            if (tid == 0) {
                const float dot = mat[static_cast<long long>(k) + static_cast<long long>(j) * n] + red[0];
                const float w = tau_s * dot;
                w_s = w;
                mat[static_cast<long long>(k) + static_cast<long long>(j) * n] -= w;
            }
            __syncthreads();

            const float w = w_s;
            for (int i = k + 1 + tid; i < n; i += blockDim.x) {
                const long long idx = static_cast<long long>(i) + static_cast<long long>(j) * n;
                mat[idx] = fmaf(-mat[i + static_cast<long long>(k) * n], w, mat[idx]);
            }
            __syncthreads();
        }
    }
}

__global__ void panel_factor_colmajor_shared_kernel(float* __restrict__ col,
                                                    float* __restrict__ tau,
                                                    float* __restrict__ saved,
                                                    int n,
                                                    int k0,
                                                    int ib,
                                                    int nb_storage,
                                                    int make_explicit_v) {
    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    const int m = n - k0;
    const long long matrix_size = static_cast<long long>(n) * n;
    float* mat = col + static_cast<long long>(b) * matrix_size;
    float* tau_b = tau + static_cast<long long>(b) * n;

    extern __shared__ float panel[];
    __shared__ float red[PANEL_THREADS];
    __shared__ float tau_s;
    __shared__ float scale_s;
    __shared__ float w_s;

    const int panel_elems = m * ib;
    for (int idx = tid; idx < panel_elems; idx += blockDim.x) {
        const int r = idx % m;
        const int s = idx / m;
        panel[idx] = mat[static_cast<long long>(k0 + r) + static_cast<long long>(k0 + s) * n];
    }
    __syncthreads();

    for (int s = 0; s < ib; ++s) {
        float local = 0.0f;
        for (int r = s + 1 + tid; r < m; r += blockDim.x) {
            const float x = panel[r + static_cast<long long>(s) * m];
            local = fmaf(x, x, local);
        }
        red[tid] = local;
        __syncthreads();

        for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
            if (tid < stride) {
                red[tid] += red[tid + stride];
            }
            __syncthreads();
        }

        if (tid == 0) {
            const long long diag = static_cast<long long>(s) + static_cast<long long>(s) * m;
            const float alpha = panel[diag];
            const float xnorm2 = red[0];
            if (xnorm2 == 0.0f) {
                tau_s = 0.0f;
                scale_s = 0.0f;
                tau_b[k0 + s] = 0.0f;
            } else {
                const float norm = sqrtf(fmaf(alpha, alpha, xnorm2));
                const float beta = (alpha >= 0.0f) ? -norm : norm;
                const float tau_val = (beta - alpha) / beta;
                tau_s = tau_val;
                scale_s = 1.0f / (alpha - beta);
                panel[diag] = beta;
                tau_b[k0 + s] = tau_val;
            }
        }
        __syncthreads();

        if (tau_s != 0.0f) {
            const float scale = scale_s;
            for (int r = s + 1 + tid; r < m; r += blockDim.x) {
                panel[r + static_cast<long long>(s) * m] *= scale;
            }
        }
        __syncthreads();

        for (int jj = s + 1; jj < ib; ++jj) {
            float dot_local = 0.0f;
            for (int r = s + 1 + tid; r < m; r += blockDim.x) {
                dot_local = fmaf(panel[r + static_cast<long long>(s) * m],
                                 panel[r + static_cast<long long>(jj) * m],
                                 dot_local);
            }
            red[tid] = dot_local;
            __syncthreads();

            for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
                if (tid < stride) {
                    red[tid] += red[tid + stride];
                }
                __syncthreads();
            }

            if (tid == 0) {
                const float dot = panel[s + static_cast<long long>(jj) * m] + red[0];
                const float w = tau_s * dot;
                w_s = w;
                panel[s + static_cast<long long>(jj) * m] -= w;
            }
            __syncthreads();

            const float w = w_s;
            for (int r = s + 1 + tid; r < m; r += blockDim.x) {
                const long long idx = static_cast<long long>(r) + static_cast<long long>(jj) * m;
                panel[idx] = fmaf(-panel[r + static_cast<long long>(s) * m], w, panel[idx]);
            }
            __syncthreads();
        }
    }

    for (int idx = tid; idx < panel_elems; idx += blockDim.x) {
        const int r = idx % m;
        const int s = idx / m;
        float value = panel[idx];
        if (make_explicit_v && r <= s) {
            const long long save_stride = static_cast<long long>(n) * nb_storage;
            saved[static_cast<long long>(b) * save_stride +
                  static_cast<long long>(k0) * nb_storage +
                  r + static_cast<long long>(s) * nb_storage] = value;
            value = (r == s) ? 1.0f : 0.0f;
        }
        mat[static_cast<long long>(k0 + r) + static_cast<long long>(k0 + s) * n] = value;
    }
}

__global__ void panel_factor_colmajor_shared_warpcols_kernel(float* __restrict__ col,
                                                             float* __restrict__ tau,
                                                             float* __restrict__ saved,
                                                             int n,
                                                             int k0,
                                                             int ib,
                                                             int nb_storage,
                                                             int make_explicit_v) {
    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    const int lane = tid & 31;
    const int warp = tid >> 5;
    const int warp_count = blockDim.x >> 5;
    const int m = n - k0;
    const long long matrix_size = static_cast<long long>(n) * n;
    float* mat = col + static_cast<long long>(b) * matrix_size;
    float* tau_b = tau + static_cast<long long>(b) * n;

    extern __shared__ float panel[];
    __shared__ float red[PANEL_THREADS];
    __shared__ float tau_s;
    __shared__ float scale_s;

    const int panel_elems = m * ib;
    for (int idx = tid; idx < panel_elems; idx += blockDim.x) {
        const int r = idx % m;
        const int s = idx / m;
        panel[idx] = mat[static_cast<long long>(k0 + r) + static_cast<long long>(k0 + s) * n];
    }
    __syncthreads();

    for (int s = 0; s < ib; ++s) {
        float local = 0.0f;
        for (int r = s + 1 + tid; r < m; r += blockDim.x) {
            const float x = panel[r + static_cast<long long>(s) * m];
            local = fmaf(x, x, local);
        }
        unsigned mask = 0xffffffffu;
        for (int offset = 16; offset > 0; offset >>= 1) {
            local += __shfl_down_sync(mask, local, offset);
        }
        if (lane == 0) {
            red[warp] = local;
        }
        __syncthreads();

        if (warp == 0) {
            float warp_sum = (lane < warp_count) ? red[lane] : 0.0f;
            for (int offset = 16; offset > 0; offset >>= 1) {
                warp_sum += __shfl_down_sync(mask, warp_sum, offset);
            }
            if (lane == 0) {
                red[0] = warp_sum;
            }
        }
        __syncthreads();

        if (tid == 0) {
            const long long diag = static_cast<long long>(s) + static_cast<long long>(s) * m;
            const float alpha = panel[diag];
            const float xnorm2 = red[0];
            if (xnorm2 == 0.0f) {
                tau_s = 0.0f;
                scale_s = 0.0f;
                tau_b[k0 + s] = 0.0f;
            } else {
                const float norm = sqrtf(fmaf(alpha, alpha, xnorm2));
                const float beta = (alpha >= 0.0f) ? -norm : norm;
                const float tau_val = (beta - alpha) / beta;
                tau_s = tau_val;
                scale_s = 1.0f / (alpha - beta);
                panel[diag] = beta;
                tau_b[k0 + s] = tau_val;
            }
        }
        __syncthreads();

        if (tau_s != 0.0f) {
            const float scale = scale_s;
            for (int r = s + 1 + tid; r < m; r += blockDim.x) {
                panel[r + static_cast<long long>(s) * m] *= scale;
            }
        }
        __syncthreads();

        const int jj = s + 1 + warp;
        if (jj < ib) {
            float dot = 0.0f;
            for (int r = s + 1 + lane; r < m; r += 32) {
                dot = fmaf(panel[r + static_cast<long long>(s) * m],
                           panel[r + static_cast<long long>(jj) * m],
                           dot);
            }
            for (int offset = 16; offset > 0; offset >>= 1) {
                dot += __shfl_down_sync(mask, dot, offset);
            }

            float w = 0.0f;
            if (lane == 0) {
                const float full_dot = panel[s + static_cast<long long>(jj) * m] + dot;
                w = tau_s * full_dot;
                panel[s + static_cast<long long>(jj) * m] -= w;
            }
            w = __shfl_sync(mask, w, 0);

            for (int r = s + 1 + lane; r < m; r += 32) {
                const long long idx = static_cast<long long>(r) + static_cast<long long>(jj) * m;
                panel[idx] = fmaf(-panel[r + static_cast<long long>(s) * m], w, panel[idx]);
            }
        }
        __syncthreads();
    }

    for (int idx = tid; idx < panel_elems; idx += blockDim.x) {
        const int r = idx % m;
        const int s = idx / m;
        float value = panel[idx];
        if (make_explicit_v && r <= s) {
            const long long save_stride = static_cast<long long>(n) * nb_storage;
            saved[static_cast<long long>(b) * save_stride +
                  static_cast<long long>(k0) * nb_storage +
                  r + static_cast<long long>(s) * nb_storage] = value;
            value = (r == s) ? 1.0f : 0.0f;
        }
        mat[static_cast<long long>(k0 + r) + static_cast<long long>(k0 + s) * n] = value;
    }
}

__global__ void apply_t_transpose_small_kernel(const float* __restrict__ t,
                                               const float* __restrict__ w,
                                               float* __restrict__ w2,
                                               int ib,
                                               int trailing,
                                               long long stride_small,
                                               long long stride_w,
                                               long long total) {
    const long long idx = static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
    if (idx >= total) {
        return;
    }
    const long long per_batch = static_cast<long long>(ib) * trailing;
    const long long b = idx / per_batch;
    const long long rem = idx - b * per_batch;
    const int p = static_cast<int>(rem % ib);
    const int j = static_cast<int>(rem / ib);
    const float* tb = t + b * stride_small;
    const float* wb = w + b * stride_w;
    float* w2b = w2 + b * stride_w;

    float sum = 0.0f;
    for (int q = 0; q < ib; ++q) {
        sum = fmaf(tb[q + static_cast<long long>(p) * ib],
                   wb[q + static_cast<long long>(j) * ib],
                   sum);
    }
    w2b[p + static_cast<long long>(j) * ib] = sum;
}

__global__ void build_t_apply_transpose_fused_kernel(const float* __restrict__ gram,
                                                     const float* __restrict__ tau,
                                                     const float* __restrict__ w,
                                                     float* __restrict__ w2,
                                                     int n,
                                                     int k,
                                                     int ib,
                                                     int trailing,
                                                     long long stride_small,
                                                     long long stride_w) {
    constexpr int COL_TILE = 64;
    const int b = blockIdx.x;
    const int tile_start = blockIdx.y * COL_TILE;
    const int tid = threadIdx.x;
    const int cols_left = trailing - tile_start;
    const int cols = (cols_left < COL_TILE) ? cols_left : COL_TILE;
    if (cols <= 0) {
        return;
    }

    __shared__ float ts[64 * 64];
    if (tid < ib * ib) {
        ts[tid] = 0.0f;
    }
    __syncthreads();

    if (tid == 0) {
        const float* g = gram + static_cast<long long>(b) * stride_small;
        const float* tau_b = tau + static_cast<long long>(b) * n + k;
        float work[64];

        for (int i = 0; i < ib; ++i) {
            const float tau_i = tau_b[i];
            if (tau_i == 0.0f) {
                ts[i + static_cast<long long>(i) * ib] = 0.0f;
                continue;
            }

            for (int j = 0; j < i; ++j) {
                work[j] = -tau_i * g[j + static_cast<long long>(i) * ib];
            }

            for (int row = 0; row < i; ++row) {
                float sum = 0.0f;
                for (int col = row; col < i; ++col) {
                    sum = fmaf(ts[row + static_cast<long long>(col) * ib], work[col], sum);
                }
                ts[row + static_cast<long long>(i) * ib] = sum;
            }
            ts[i + static_cast<long long>(i) * ib] = tau_i;
        }
    }
    __syncthreads();

    const float* wb = w + static_cast<long long>(b) * stride_w;
    float* w2b = w2 + static_cast<long long>(b) * stride_w;
    const int tile_outputs = ib * cols;
    for (int idx = tid; idx < tile_outputs; idx += blockDim.x) {
        const int p = idx % ib;
        const int j = tile_start + idx / ib;
        float sum = 0.0f;
        for (int q = 0; q < ib; ++q) {
            sum = fmaf(ts[q + static_cast<long long>(p) * ib],
                       wb[q + static_cast<long long>(j) * ib],
                       sum);
        }
        w2b[p + static_cast<long long>(j) * ib] = sum;
    }
}

}  // namespace

void qr32(torch::Tensor a, torch::Tensor h, torch::Tensor tau) {
    const int batch = static_cast<int>(a.size(0));
    const int n = static_cast<int>(a.size(1));
    qr32_kernel<128><<<batch, 128>>>(
        a.data_ptr<float>(),
        h.data_ptr<float>(),
        tau.data_ptr<float>(),
        batch,
        n);

    cudaError_t err = cudaGetLastError();
    if (err != cudaSuccess) {
        throw std::runtime_error(cudaGetErrorString(err));
    }
}

void qr_shared_176(torch::Tensor a, torch::Tensor h, torch::Tensor tau) {
    const int batch = static_cast<int>(a.size(0));
    const int n = static_cast<int>(a.size(1));
    const int shmem = n * n * static_cast<int>(sizeof(float));
    check_cuda(
        cudaFuncSetAttribute(
            qr_shared_kernel<256>,
            cudaFuncAttributeMaxDynamicSharedMemorySize,
            shmem),
        "cudaFuncSetAttribute qr_shared_kernel");
    qr_shared_kernel<256><<<batch, 256, shmem>>>(
        a.data_ptr<float>(),
        h.data_ptr<float>(),
        tau.data_ptr<float>(),
        n);

    cudaError_t err = cudaGetLastError();
    if (err != cudaSuccess) {
        throw std::runtime_error(cudaGetErrorString(err));
    }
}

void qr_global_256(torch::Tensor a, torch::Tensor h, torch::Tensor tau) {
    const int batch = static_cast<int>(a.size(0));
    const int n = static_cast<int>(a.size(1));
    qr_global_kernel<256><<<batch, 256>>>(
        a.data_ptr<float>(),
        h.data_ptr<float>(),
        tau.data_ptr<float>(),
        n);

    cudaError_t err = cudaGetLastError();
    if (err != cudaSuccess) {
        throw std::runtime_error(cudaGetErrorString(err));
    }
}

void qr_global_512(torch::Tensor a, torch::Tensor h, torch::Tensor tau) {
    const int batch = static_cast<int>(a.size(0));
    const int n = static_cast<int>(a.size(1));
    qr_global_kernel<512><<<batch, 512>>>(
        a.data_ptr<float>(),
        h.data_ptr<float>(),
        tau.data_ptr<float>(),
        n);

    cudaError_t err = cudaGetLastError();
    if (err != cudaSuccess) {
        throw std::runtime_error(cudaGetErrorString(err));
    }
}

void qr_blocked_wy(torch::Tensor a, torch::Tensor h, torch::Tensor tau, int nb, int custom_small) {
    const int batch = static_cast<int>(a.size(0));
    const int n = static_cast<int>(a.size(1));
    const long long matrix_size = static_cast<long long>(n) * n;
    const long long total = static_cast<long long>(batch) * matrix_size;
    if (nb <= 0 || nb > 64) {
        throw std::runtime_error("unsupported blocked-WY panel width");
    }

    auto col = torch::empty_like(a);
    constexpr int TILE_DIM = 32;
    constexpr int BLOCK_ROWS = 8;
    const dim3 transpose_block(TILE_DIM, BLOCK_ROWS);
    const dim3 transpose_grid((n + TILE_DIM - 1) / TILE_DIM,
                              (n + TILE_DIM - 1) / TILE_DIM,
                              batch);
    transpose_tiled_kernel<<<transpose_grid, transpose_block>>>(
        a.data_ptr<float>(),
        col.data_ptr<float>(),
        n);
    check_cuda(cudaGetLastError(), "transpose_tiled_kernel row_to_col");

    cublasHandle_t blas = get_blas_handle();
    const int max_panel_shmem = n * nb * static_cast<int>(sizeof(float));
    check_cuda(
        cudaFuncSetAttribute(
            panel_factor_colmajor_shared_warpcols_kernel,
            cudaFuncAttributeMaxDynamicSharedMemorySize,
            max_panel_shmem),
        "cudaFuncSetAttribute panel_factor_colmajor_shared_warpcols_kernel");

    auto opts = torch::TensorOptions().device(a.device()).dtype(torch::kFloat32);
    const bool delayed_inplace_v_mode = (custom_small == 8 || custom_small == 9 || custom_small == 10);
    const bool inplace_v_mode = (custom_small == 6 || custom_small == 7 || delayed_inplace_v_mode);
    torch::Tensor v;
    if (!inplace_v_mode) {
        v = torch::empty({static_cast<long long>(batch) * n * nb}, opts);
    }
    auto gram = torch::empty({static_cast<long long>(batch) * nb * nb}, opts);
    auto t = torch::empty({static_cast<long long>(batch) * nb * nb}, opts);
    auto w = torch::empty({static_cast<long long>(batch) * nb * n}, opts);
    auto w2 = torch::empty({static_cast<long long>(batch) * nb * n}, opts);
    torch::Tensor saved;
    if (delayed_inplace_v_mode) {
        saved = torch::empty({static_cast<long long>(batch) * n * nb}, opts);
    } else if (inplace_v_mode) {
        saved = torch::empty({static_cast<long long>(batch) * nb * nb}, opts);
    }

    float* col_ptr = col.data_ptr<float>();
    float* tau_ptr = tau.data_ptr<float>();
    const float one = 1.0f;
    const float zero = 0.0f;
    const float minus_one = -1.0f;

    for (int k = 0; k < n; k += nb) {
        const int m = n - k;
        const int ib = (m < nb) ? m : nb;
        const int trailing = n - k - ib;
        const bool delayed_inplace_v = (custom_small == 8 || custom_small == 9 || custom_small == 10);
        const bool inplace_v = (custom_small == 6 || custom_small == 7 || delayed_inplace_v);

        const int panel_shmem = (n - k) * ib * static_cast<int>(sizeof(float));
        panel_factor_colmajor_shared_warpcols_kernel<<<batch, PANEL_THREADS, panel_shmem>>>(
            col_ptr,
            tau_ptr,
            delayed_inplace_v ? saved.data_ptr<float>() : nullptr,
            n,
            k,
            ib,
            nb,
            delayed_inplace_v ? 1 : 0);
        check_cuda(cudaGetLastError(), "panel_factor_colmajor_shared_warpcols_kernel");

        const long long stride_v = static_cast<long long>(m) * ib;
        const long long stride_small = static_cast<long long>(ib) * ib;
        float* v_ptr = nullptr;
        int ldv = m;
        long long stride_v_use = stride_v;
        if (delayed_inplace_v) {
            v_ptr = col_ptr + static_cast<long long>(k) + static_cast<long long>(k) * n;
            ldv = n;
            stride_v_use = matrix_size;
        } else if (inplace_v) {
            prepare_panel_v_inplace_kernel<<<batch, 256>>>(
                col_ptr,
                saved.data_ptr<float>(),
                n,
                k,
                ib,
                stride_small);
            check_cuda(cudaGetLastError(), "prepare_panel_v_inplace_kernel");
            v_ptr = col_ptr + static_cast<long long>(k) + static_cast<long long>(k) * n;
            ldv = n;
            stride_v_use = matrix_size;
        } else {
            v_ptr = v.data_ptr<float>();
            const long long v_total = static_cast<long long>(batch) * stride_v;
            const int v_blocks = static_cast<int>((v_total + 255) / 256);
            build_explicit_v_kernel<<<v_blocks, 256>>>(
                col_ptr,
                v.data_ptr<float>(),
                n,
                k,
                m,
                ib,
                stride_v,
                v_total);
            check_cuda(cudaGetLastError(), "build_explicit_v_kernel");
        }

        if (custom_small == 1 || custom_small == 2) {
            const dim3 gram_grid(batch, ib, ib);
            build_gram_small_kernel<<<gram_grid, PANEL_THREADS>>>(
                v.data_ptr<float>(),
                gram.data_ptr<float>(),
                m,
                ib,
                stride_v,
                stride_small);
            check_cuda(cudaGetLastError(), "build_gram_small_kernel");
        } else if (custom_small == 10) {
            const dim3 gram_grid(batch, ib, ib);
            build_gram_strided_small_kernel<<<gram_grid, PANEL_THREADS>>>(
                v_ptr,
                gram.data_ptr<float>(),
                m,
                ib,
                ldv,
                stride_v_use,
                stride_small);
            check_cuda(cudaGetLastError(), "build_gram_strided_small_kernel");
        } else {
            check_blas(
                cublasSgemmStridedBatched(
                    blas,
                    CUBLAS_OP_T,
                    CUBLAS_OP_N,
                    ib,
                    ib,
                    m,
                    &one,
                    v_ptr,
                    ldv,
                    stride_v_use,
                    v_ptr,
                    ldv,
                    stride_v_use,
                    &zero,
                    gram.data_ptr<float>(),
                    ib,
                    stride_small,
                    batch),
                "cublasSgemmStridedBatched gram");
        }

        const bool fused_t_apply = (custom_small == 2 || custom_small == 5 || custom_small == 9 || custom_small == 10);
        if (!fused_t_apply) {
            build_t_from_gram_kernel<<<batch, 1>>>(
                gram.data_ptr<float>(),
                t.data_ptr<float>(),
                tau_ptr,
                n,
                k,
                ib,
                stride_small);
            check_cuda(cudaGetLastError(), "build_t_from_gram_kernel");
        }

        if (trailing > 0) {
            const long long stride_w = static_cast<long long>(ib) * trailing;
            float* c_panel = col_ptr + static_cast<long long>(k) +
                             static_cast<long long>(k + ib) * n;

            check_blas(
                cublasSgemmStridedBatched(
                    blas,
                    CUBLAS_OP_T,
                    CUBLAS_OP_N,
                    ib,
                    trailing,
                    m,
                    &one,
                    v_ptr,
                    ldv,
                    stride_v_use,
                    c_panel,
                    n,
                    matrix_size,
                    &zero,
                    w.data_ptr<float>(),
                    ib,
                    stride_w,
                    batch),
                "cublasSgemmStridedBatched vt_c");

            if (fused_t_apply) {
                constexpr int fused_cols = 64;
                const dim3 fused_grid(batch, (trailing + fused_cols - 1) / fused_cols);
                build_t_apply_transpose_fused_kernel<<<fused_grid, 256>>>(
                    gram.data_ptr<float>(),
                    tau_ptr,
                    w.data_ptr<float>(),
                    w2.data_ptr<float>(),
                    n,
                    k,
                    ib,
                    trailing,
                    stride_small,
                    stride_w);
                check_cuda(cudaGetLastError(), "build_t_apply_transpose_fused_kernel");
            } else if (custom_small == 1) {
                const long long tw_total = static_cast<long long>(batch) * stride_w;
                const int tw_blocks = static_cast<int>((tw_total + 255) / 256);
                apply_t_transpose_small_kernel<<<tw_blocks, 256>>>(
                    t.data_ptr<float>(),
                    w.data_ptr<float>(),
                    w2.data_ptr<float>(),
                    ib,
                    trailing,
                    stride_small,
                    stride_w,
                    tw_total);
                check_cuda(cudaGetLastError(), "apply_t_transpose_small_kernel");
            } else {
                check_blas(
                    cublasSgemmStridedBatched(
                        blas,
                        CUBLAS_OP_T,
                        CUBLAS_OP_N,
                        ib,
                        trailing,
                        ib,
                        &one,
                        t.data_ptr<float>(),
                        ib,
                        stride_small,
                        w.data_ptr<float>(),
                        ib,
                        stride_w,
                        &zero,
                        w2.data_ptr<float>(),
                        ib,
                        stride_w,
                        batch),
                    "cublasSgemmStridedBatched t_w");
            }

            check_blas(
                cublasSgemmStridedBatched(
                    blas,
                    CUBLAS_OP_N,
                    CUBLAS_OP_N,
                    m,
                    trailing,
                    ib,
                    &minus_one,
                    v_ptr,
                    ldv,
                    stride_v_use,
                    w2.data_ptr<float>(),
                    ib,
                    stride_w,
                    &one,
                    c_panel,
                    n,
                    matrix_size,
                    batch),
                "cublasSgemmStridedBatched update");
        }

        if (inplace_v && !delayed_inplace_v) {
            restore_panel_r_kernel<<<batch, 256>>>(
                col_ptr,
                saved.data_ptr<float>(),
                n,
                k,
                ib,
                stride_small);
            check_cuda(cudaGetLastError(), "restore_panel_r_kernel");
        }
    }

    if (delayed_inplace_v_mode) {
        transpose_tiled_restore_v_kernel<<<transpose_grid, transpose_block>>>(
            col.data_ptr<float>(),
            h.data_ptr<float>(),
            saved.data_ptr<float>(),
            n,
            nb);
        check_cuda(cudaGetLastError(), "transpose_tiled_restore_v_kernel");
    } else {
        transpose_tiled_kernel<<<transpose_grid, transpose_block>>>(
            col.data_ptr<float>(),
            h.data_ptr<float>(),
            n);
        check_cuda(cudaGetLastError(), "transpose_tiled_kernel col_to_row");
    }

}
"""


_module = load_inline(
    name="qr_compact_householder_b200_v99_n352_nb16_retry",
    cpp_sources=[CPP_SRC],
    cuda_sources=[CUDA_SRC],
    functions=["qr32", "qr_shared_176", "qr_global_256", "qr_global_512", "qr_blocked_wy"],
    verbose=False,
    extra_cuda_cflags=["-O3", "--use_fast_math"],
    extra_ldflags=["-lcublas"],
)


def custom_kernel(data: input_t) -> output_t:
    if data.is_cuda and data.dtype == torch.float32 and data.is_contiguous():
        batch, n, m = data.shape
        if m == n and n <= 32:
            h = torch.empty_like(data)
            tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
            _module.qr32(data, h, tau)
            return h, tau
        if m == n and n == 176:
            h = torch.empty_like(data)
            tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
            _module.qr_shared_176(data, h, tau)
            return h, tau
        if m == n and n <= 176:
            h = torch.empty_like(data)
            tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
            _module.qr_global_256(data, h, tau)
            return h, tau
        if m == n and n <= 512 and not (n == 352 or (n == 512 and batch >= 128)):
            h = torch.empty_like(data)
            tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
            _module.qr_global_512(data, h, tau)
            return h, tau
        if m == n and n == 512 and batch >= 128:
            h = torch.empty_like(data)
            tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
            _module.qr_blocked_wy(data, h, tau, 16, 8)
            return h, tau
        if m == n and n == 352:
            h = torch.empty_like(data)
            tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
            _module.qr_blocked_wy(data, h, tau, 16, 9)
            return h, tau
        if m == n and n == 1024:
            h = torch.empty_like(data)
            tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
            _module.qr_blocked_wy(data, h, tau, 16, 9)
            return h, tau
        if m == n and n == 2048:
            h = torch.empty_like(data)
            tau = torch.empty((batch, n), device=data.device, dtype=torch.float32)
            _module.qr_blocked_wy(data, h, tau, 16, 10)
            return h, tau
    return torch.geqrf(data)
scrolls · 1523 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