Skip to content
KernelIndex
Search⌘K

submission 843485

mertdonmez1453 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c3d9fb911362d385ae2772b9828771b561581205a66fcb416e7fe9adccb138af
license declaredunknown
license concludedunknown
authorsmertdonmez1453
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.py1828 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200

from pathlib import Path

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 mode);
void qr_n512_custom(torch::Tensor a, torch::Tensor h, torch::Tensor tau);
void qr_n512_human(torch::Tensor a, torch::Tensor h, torch::Tensor tau);
"""


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;
constexpr int MODE_CUBLAS_GRAM_CUBLAS_T = 8;   // n512 large batch
constexpr int MODE_CUBLAS_GRAM_FUSED_T = 9;    // n352 / n1024
constexpr int MODE_CUSTOM_GRAM_FUSED_T = 10;   // n2048
constexpr int MODE_N512_CUSTOM_UPDATE = 11;    // n512 scaffold: custom C -= V @ W2

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

// n176 fits in B200 dynamic shared memory, so the whole matrix can be factored
// inside one CTA without repeated global-memory panel traffic.
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];
    }
}

// Simple one-CTA Householder QR for mid-size shapes where launch count and
// batch-level parallelism beat cuSOLVER overhead.
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();
    }
}

// Tiled row-major <-> column-major transform. The blocked-WY path stores its
// working matrix in column-major order but must return row-major compact H.
__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];
        }
    }
}

// Final transpose plus delayed restore of the panel R blocks that were
// temporarily overwritten while using the panel storage as explicit V.
__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;
        }
    }
}

// Build the triangular WY T factor from G = V.T @ V and tau. This path is kept
// for n512, where cuBLAS handles the following T.T @ W better than the fused
// custom kernel.
__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;
    }
}

// n2048 has a small batch; a custom dot kernel wins over cuBLAS for the tiny
// per-panel Gram matrix.
__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];
    }
}

// Shared-memory panel factorization. One CTA owns one matrix in the batch; each
// warp updates one target panel column after a reflector is formed.
__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;
    }
}

// Fused small operation for modes where building T and then launching cuBLAS for
// T.T @ W costs more than doing both in one custom kernel.
__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;
    }
}

// Human-editable replacement for the n512 cuBLAS W = V.T @ C call.
//
// cuBLAS sees W = V.T @ C as a strided-batched column-major SGEMM:
//     V : m x ib,       lda = n,  stride = n*n
//     C : m x trailing, ldc = n,  stride = n*n
//     W : ib x trailing, ldw = ib, stride = ib*trailing
//
// The production library uses a tiled SIMT GEMM. This starter version computes
// one tile of output columns per CTA and loops over all ib panel rows inside the
// CTA. It is intentionally regular and boring, so it can be rewritten by hand.
__global__ void n512_vt_c_tile_kernel(const float* __restrict__ v,
                                      const float* __restrict__ c,
                                      float* __restrict__ w,
                                      int n,
                                      int m,
                                      int ib,
                                      int trailing,
                                      long long matrix_stride,
                                      long long w_stride) {
    constexpr int COL_TILE = 8;
    constexpr int THREADS = 256;
    __shared__ float red[COL_TILE][THREADS];

    const int b = blockIdx.x;
    const int tile_col = blockIdx.y * COL_TILE;
    const int tid = threadIdx.x;

    const float* vb = v + static_cast<long long>(b) * matrix_stride;
    const float* cb = c + static_cast<long long>(b) * matrix_stride;
    float* wb = w + static_cast<long long>(b) * w_stride;

    for (int p = 0; p < ib; ++p) {
        #pragma unroll
        for (int tc = 0; tc < COL_TILE; ++tc) {
            const int col = tile_col + tc;
            float sum = 0.0f;
            if (col < trailing) {
                for (int r = tid; r < m; r += blockDim.x) {
                    sum = fmaf(vb[r + static_cast<long long>(p) * n],
                               cb[r + static_cast<long long>(col) * n],
                               sum);
                }
            }
            red[tc][tid] = sum;
        }
        __syncthreads();

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

        if (tid == 0) {
            #pragma unroll
            for (int tc = 0; tc < COL_TILE; ++tc) {
                const int col = tile_col + tc;
                if (col < trailing) {
                    wb[p + static_cast<long long>(col) * ib] = red[tc][0];
                }
            }
        }
        __syncthreads();
    }
}

// Human-editable replacement for the n512 cuBLAS SGEMM update:
//
//     C = C - V @ W2
//
// cuBLAS sees this as a strided-batched column-major SGEMM:
//     A = V   : m x ib,       lda = n,  stride = n*n
//     B = W2  : ib x trailing, ldb = ib, stride = ib*trailing
//     C = C   : m x trailing,  ldc = n,  stride = n*n
//
// The library implementation is a tiled SIMT SGEMM: load a tile of V and W2,
// compute many C elements per CTA, and write back beta*C + alpha*A*B. This
// version is intentionally simple instead: one thread computes one C element
// and the small K dimension (ib <= 16 here) is unrolled. It is not trying to be
// cuBLAS-fast yet; it is a compact, correct place to start rewriting.
__global__ void n512_update_trailing_elementwise_kernel(const float* __restrict__ v,
                                                        const float* __restrict__ w2,
                                                        float* __restrict__ c,
                                                        int n,
                                                        int m,
                                                        int ib,
                                                        int trailing,
                                                        long long matrix_stride,
                                                        long long w_stride) {
    const int col = blockIdx.x * blockDim.x + threadIdx.x;
    const int row = blockIdx.y * blockDim.y + threadIdx.y;
    const int b = blockIdx.z;
    if (row >= m || col >= trailing) {
        return;
    }

    const float* vb = v + static_cast<long long>(b) * matrix_stride;
    const float* w2b = w2 + static_cast<long long>(b) * w_stride;
    float* cb = c + static_cast<long long>(b) * matrix_stride;

    float acc = 0.0f;
    #pragma unroll
    for (int q = 0; q < 16; ++q) {
        if (q < ib) {
            acc = fmaf(vb[row + static_cast<long long>(q) * n],
                       w2b[q + static_cast<long long>(col) * ib],
                       acc);
        }
    }
    cb[row + static_cast<long long>(col) * n] -= acc;
}

}  // namespace

void qr_global_256(torch::Tensor a, torch::Tensor h, torch::Tensor tau);

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));
    cudaError_t err = cudaFuncSetAttribute(
        qr_shared_kernel<256>,
        cudaFuncAttributeMaxDynamicSharedMemorySize,
        shmem);
    if (err != cudaSuccess) {
        cudaGetLastError(); // Clear the error
        qr_global_256(a, h, tau);
        return;
    }
    qr_shared_kernel<256><<<batch, 256, shmem>>>(
        a.data_ptr<float>(),
        h.data_ptr<float>(),
        tau.data_ptr<float>(),
        n);

    cudaError_t launch_err = cudaGetLastError();
    if (launch_err != cudaSuccess) {
        throw std::runtime_error(cudaGetErrorString(launch_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 mode) {
    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;

    int device_id = 0;
    check_cuda(cudaGetDevice(&device_id), "cudaGetDevice");
    cudaDeviceProp prop;
    check_cuda(cudaGetDeviceProperties(&prop, device_id), "cudaGetDeviceProperties");
    int max_shmem = prop.sharedMemPerBlockOptin;
    if (max_shmem <= 0) {
        max_shmem = prop.sharedMemPerBlock;
    }
    while (nb > 1 && n * nb * static_cast<int>(sizeof(float)) > max_shmem) {
        nb /= 2;
    }

    if (nb <= 0 || nb > 64) {
        throw std::runtime_error("unsupported blocked-WY panel width");
    }
    if (mode != MODE_CUBLAS_GRAM_CUBLAS_T &&
        mode != MODE_CUBLAS_GRAM_FUSED_T &&
        mode != MODE_CUSTOM_GRAM_FUSED_T &&
        mode != MODE_N512_CUSTOM_UPDATE) {
        throw std::runtime_error("unsupported blocked-WY mode");
    }

    // Work in column-major layout so panel columns and cuBLAS operands are natural.
    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);
    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);
    // Panel storage is temporarily overwritten with explicit V. Save the upper
    // triangular R blocks here and restore them during the final transpose.
    auto saved = torch::empty({static_cast<long long>(batch) * n * 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;

        // Factor A[k:n, k:k+ib] in shared memory. The kernel leaves the panel
        // as explicit V in-place and writes compact Householder tau.
        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,
            saved.data_ptr<float>(),
            n,
            k,
            ib,
            nb,
            1);
        check_cuda(cudaGetLastError(), "panel_factor_colmajor_shared_warpcols_kernel");

        const long long stride_small = static_cast<long long>(ib) * ib;
        float* v_ptr = col_ptr + static_cast<long long>(k) + static_cast<long long>(k) * n;
        const int ldv = n;
        const long long stride_v_use = matrix_size;

        // Gram is tiny (ib x ib) but repeated per panel and per batch. n2048
        // prefers the custom strided dot kernel; n352/n512/n1024 use cuBLAS.
        if (mode == MODE_CUSTOM_GRAM_FUSED_T) {
            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 = (mode == MODE_CUBLAS_GRAM_FUSED_T ||
                                    mode == MODE_CUSTOM_GRAM_FUSED_T);
        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 {
                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");
            }

            if (mode == MODE_N512_CUSTOM_UPDATE) {
                const dim3 update_block(16, 16);
                const dim3 update_grid((trailing + update_block.x - 1) / update_block.x,
                                       (m + update_block.y - 1) / update_block.y,
                                       batch);
                n512_update_trailing_elementwise_kernel<<<update_grid, update_block>>>(
                    v_ptr,
                    w2.data_ptr<float>(),
                    c_panel,
                    n,
                    m,
                    ib,
                    trailing,
                    matrix_size,
                    stride_w);
                check_cuda(cudaGetLastError(), "n512_update_trailing_elementwise_kernel");
            } else {
                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");
            }
        }
    }

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

}

void qr_n512_custom(torch::Tensor a, torch::Tensor h, torch::Tensor tau) {
    const int batch = static_cast<int>(a.size(0));
    constexpr int n = 512;
    constexpr int nb = 16;
    constexpr long long matrix_size = static_cast<long long>(n) * n;

    // Same layout and panel QR as the cuBLAS-backed n512 path. The difference is
    // that every trailing-matrix operation below is now an editable CUDA kernel.
    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 n512_custom");

    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 n512_custom");

    auto opts = torch::TensorOptions().device(a.device()).dtype(torch::kFloat32);
    auto gram = 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);
    auto saved = torch::empty({static_cast<long long>(batch) * n * nb}, opts);

    float* col_ptr = col.data_ptr<float>();
    float* tau_ptr = tau.data_ptr<float>();

    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 long long stride_small = static_cast<long long>(ib) * ib;
        const int panel_shmem = m * ib * static_cast<int>(sizeof(float));

        panel_factor_colmajor_shared_warpcols_kernel<<<batch, PANEL_THREADS, panel_shmem>>>(
            col_ptr,
            tau_ptr,
            saved.data_ptr<float>(),
            n,
            k,
            ib,
            nb,
            1);
        check_cuda(cudaGetLastError(), "panel_factor_colmajor_shared_warpcols_kernel n512_custom");

        float* v_ptr = col_ptr + static_cast<long long>(k) + static_cast<long long>(k) * n;

        const dim3 gram_grid(batch, ib, ib);
        build_gram_strided_small_kernel<<<gram_grid, PANEL_THREADS>>>(
            v_ptr,
            gram.data_ptr<float>(),
            m,
            ib,
            n,
            matrix_size,
            stride_small);
        check_cuda(cudaGetLastError(), "build_gram_strided_small_kernel n512_custom");

        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;

            constexpr int vt_cols = 8;
            const dim3 vt_grid(batch, (trailing + vt_cols - 1) / vt_cols);
            n512_vt_c_tile_kernel<<<vt_grid, 256>>>(
                v_ptr,
                c_panel,
                w.data_ptr<float>(),
                n,
                m,
                ib,
                trailing,
                matrix_size,
                stride_w);
            check_cuda(cudaGetLastError(), "n512_vt_c_tile_kernel");

            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 n512_custom");

            const dim3 update_block(16, 16);
            const dim3 update_grid((trailing + update_block.x - 1) / update_block.x,
                                   (m + update_block.y - 1) / update_block.y,
                                   batch);
            n512_update_trailing_elementwise_kernel<<<update_grid, update_block>>>(
                v_ptr,
                w2.data_ptr<float>(),
                c_panel,
                n,
                m,
                ib,
                trailing,
                matrix_size,
                stride_w);
            check_cuda(cudaGetLastError(), "n512_update_trailing_elementwise_kernel");
        }
    }

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


# Popcorn uploads only submission.py. During local development we compile the
# neighboring human.cu directly; this embedded copy keeps the uploaded file
# self-contained. Run studies/embed_human_cuda.py after editing human.cu.
# BEGIN HUMAN_CUDA_EMBEDDED
HUMAN_CUDA_EMBEDDED = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <stdexcept>

#define N 512
#define THREADS_PER_COL 32
#define BLOCK_SIZE 32
#define APPLY_THREADS 256
#define COLS_PER_TRAILING_BLOCK 32

#ifndef HUMAN_N512_IMPLEMENTED
#define HUMAN_N512_IMPLEMENTED 1
#endif

//=============================================================================
// Kernel 1: Panel Factorization
// Her panel sutunu icin: norm -> alpha/beta/tau/scale -> v scale -> panel-ici update
// Grid: <<<batch, 512>>>
//=============================================================================
__global__ void __launch_bounds__(512, 3) kernel_factor_panel(float* __restrict__ h,
                                     float* __restrict__ tau,
                                     float* __restrict__ T_workspace,
                                     int k, int ib, int batch)
{
    if (blockIdx.x >= batch) return;

    int tid = threadIdx.x;
    int base = blockIdx.x * N * N;

    __shared__ float warp_sums[16]; // 512 / 32
    __shared__ float v_col[N];
    __shared__ float alpha;
    __shared__ float beta;
    __shared__ float scale;
    __shared__ float tau_col;
    __shared__ float T_shared[BLOCK_SIZE][BLOCK_SIZE + 1];
    __shared__ float w_shared[BLOCK_SIZE];

    // Initialize T_shared to 0
    if (tid < BLOCK_SIZE * BLOCK_SIZE) {
        int r = tid % BLOCK_SIZE;
        int c = tid / BLOCK_SIZE;
        T_shared[r][c] = 0.0f;
    }


    for (int j = 0; j < ib; ++j) {
        int col = k + j;

        v_col[tid] = h[base + col * N + tid];
        //__syncthreads(); gerek var mi

        // Norm reduction: once her warp kendi toplamini hesaplar.. i think this approach
        //  is not much necessary
        float norm_sq = 0.0f;

        if (tid >= col) {
            float val = v_col[tid];
            norm_sq = val * val;
        }

        for (int offset = 16; offset > 0; offset >>= 1) {
            norm_sq += __shfl_down_sync(0xffffffff, norm_sq, offset);
        }

        int norm_lane = tid & 31;
        int norm_warp = tid >> 5;

        // 16 warp'in sonuclarini shared memory'ye yaz
        if (norm_lane == 0) {
            warp_sums[norm_warp] = norm_sq;
        }

        __syncthreads();

        // ilk warp, 16 ara sonucu toplar
        if (norm_warp == 0) {
            norm_sq = norm_lane < 16 ? warp_sums[norm_lane] : 0.0f;

            for (int offset = 16; offset > 0; offset >>= 1) {
                norm_sq += __shfl_down_sync(0xffffffff, norm_sq, offset);
            }

            if (norm_lane == 0) {
                warp_sums[0] = norm_sq;
            }
        }


        // Compute scalars (thread 0 only)
        if (tid == 0) {
            alpha = v_col[col];
            float norm = sqrtf(warp_sums[0]);
            if (norm < 1e-20f) {
                beta = alpha;
                tau[blockIdx.x * N + col] = 0.0f;
                tau_col = 0.0f;
                scale = 0.0f;
            } else {
                beta = (alpha >= 0.0f) ? -norm : norm;
                float tau_val = (beta - alpha) / beta;
                tau[blockIdx.x * N + col] = tau_val;
                tau_col = tau_val;
                scale = 1.0f / (alpha - beta);
            }
            h[base + col * N + col] = beta;
            v_col[col] = 1.0f;
        }
        __syncthreads();

        // Scale Householder vector and write back
        if (tid > col) {
            v_col[tid] *= scale;
        }
        if (tid > col) {
            //burada globale yazmak yerine başka biryerde yazilabilrimi
            h[base + col * N + tid] = v_col[tid]; 
        }
        __syncthreads();

        // Update remaining columns in the current panel
        int local_tid = tid % THREADS_PER_COL;
        int col_group = tid / THREADS_PER_COL;
        int cols_per_block = blockDim.x / THREADS_PER_COL;

        for (int panel_col = col + 1; panel_col < k + ib; panel_col += cols_per_block) {
            int j_panel = panel_col + col_group;
            float sum = 0.0f;

            if (j_panel < k + ib) {
                int row = (col & ~31) + local_tid;

            // Yalnızca ilk, kısmi warp parçasında kontrol gerekli
            if (row >= col) {
                sum = fmaf(
                    v_col[row],
                    h[base + j_panel * N + row],
                    sum
                );
            }

            // Bundan sonraki bütün parçalar tam ve hizalı
            for (row += 32; row < N; row += 32) {
                sum = fmaf(
                    v_col[row],
                    h[base + j_panel * N + row],
                    sum
                );
            }
            }

            // Warp reduction
            for (int offset = THREADS_PER_COL / 2; offset > 0; offset /= 2) {
                sum += __shfl_down_sync(0xffffffff, sum, offset);
            }
            float dot = __shfl_sync(0xffffffff, sum, 0);

            if (j_panel < k + ib) {
                dot *= tau_col;

                int row = (col & ~31) + local_tid;

                // The first partial warp includes the diagonal. v_col[col] is
                // already 1, so the same FMA also performs C[col] -= dot.
                if (row >= col) {
                    h[base + j_panel * N + row] =
                        fmaf(-v_col[row], dot,
                             h[base + j_panel * N + row]);
                }

                // All subsequent warp accesses start on a 128-byte boundary.
                for (row += THREADS_PER_COL; row < N;
                     row += THREADS_PER_COL) {
                    h[base + j_panel * N + row] =
                        fmaf(-v_col[row], dot,
                             h[base + j_panel * N + row]);
                }
            }
        }
        //__syncthreads();

        // Compute w_shared[p] = -tau_col * dot(v_p, v_j) for p < j in parallel
        int lane_id = tid % 16;
        int group_id = tid / 16;
        if (group_id < j) {
            int p = group_id;
            float sum = 0.0f;

            int row = (col & ~15) + lane_id;

            // The first partial half-warp includes v_j's unit diagonal.
            if (row >= col) {
                float val_p = h[base + (k + p) * N + row];
                sum = fmaf(val_p, v_col[row], sum);
            }

            // Remaining half-warp accesses are aligned to 64-byte boundaries.
            for (row += 16; row < N; row += 16) {
                float val_p = h[base + (k + p) * N + row];
                sum = fmaf(val_p, v_col[row], sum);
            }
            unsigned int mask = (tid % 32 < 16) ? 0x0000ffff : 0xffff0000;
            for (int offset = 8; offset > 0; offset /= 2) {
                sum += __shfl_down_sync(mask, sum, offset, 16);
            }
            if (lane_id == 0) {
                w_shared[p] = -tau_col * sum;
            }
        }
        __syncthreads();

        // Compute T_shared[i][j]
        if (tid < j) {
            float sum_T = 0.0f;
            for (int p = tid; p < j; ++p) {
                sum_T = fmaf(T_shared[tid][p], w_shared[p], sum_T);
            }
            T_shared[tid][j] = sum_T;
        }
        if (tid == j) {
            T_shared[j][j] = tau_col;
        }
        __syncthreads();
    }

    // Write T_shared to global workspace
    for (int idx = tid; idx < ib * ib; idx += blockDim.x) {
        int r = idx % ib;
        int c = idx / ib;
        T_workspace[blockIdx.x * BLOCK_SIZE * BLOCK_SIZE + c * BLOCK_SIZE + r] = T_shared[r][c];
    }
}

//=============================================================================
// Kernel 2: Load T + Apply to Trailing Matrix
// Grid: <<<batch * num_col_blocks, 512>>>
//=============================================================================
__global__ void __launch_bounds__(APPLY_THREADS) kernel_build_T_apply(float* __restrict__ h,
                                      const float* __restrict__ tau,
                                      const float* __restrict__ T_workspace,
                                      int k, int ib, int batch,
                                      int num_col_blocks)
{
    int batch_id = blockIdx.x / num_col_blocks;
    int col_block_id = blockIdx.x % num_col_blocks;

    if (batch_id >= batch) return;

    int tid = threadIdx.x;
    int base = batch_id * N * N;
    int trailingStart = k + ib;

    __shared__ float T[BLOCK_SIZE][BLOCK_SIZE + 1];
    constexpr int TILE_ROWS = 64;
    __shared__ float V_shared[TILE_ROWS][BLOCK_SIZE + 1];

    // Load T from global workspace
    for (int idx = tid; idx < ib * ib; idx += blockDim.x) {
        int r = idx % ib;
        int c = idx / ib;
        T[r][c] = T_workspace[batch_id * BLOCK_SIZE * BLOCK_SIZE + c * BLOCK_SIZE + r];
    }

    // == Step 2-4: Apply block reflector to this block's trailing columns ==
    int col_start = trailingStart + col_block_id * COLS_PER_TRAILING_BLOCK;
    int col_end = col_start + COLS_PER_TRAILING_BLOCK;
    if (col_end > N) col_end = N;
    if (col_start >= N) return;

    int col_group = tid / 32;   // 0..7 (8 warps)
    int local_tid = tid % 32;   // lane within warp

    // Warp processes up to 4 columns in parallel (unrolled)
    int c0 = col_start + col_group + 0 * 8;
    int c1 = col_start + col_group + 1 * 8;
    int c2 = col_start + col_group + 2 * 8;
    int c3 = col_start + col_group + 3 * 8;

    float w_val0 = 0.0f;
    float w_val1 = 0.0f;
    float w_val2 = 0.0f;
    float w_val3 = 0.0f;

    // Pass 1: Dot Product (Compute W = V^T * C)
    for (int tile_row_start = k; tile_row_start < N; tile_row_start += TILE_ROWS) {
        __syncthreads();
        for (int load_idx = tid; load_idx < TILE_ROWS * ib; load_idx += APPLY_THREADS) {
            int r = load_idx % TILE_ROWS;
            int c_V = load_idx / TILE_ROWS;
            int global_row = tile_row_start + r;
            float val = 0.0f;
            if (global_row < N && c_V < ib) {
                if (global_row == k + c_V) {
                    val = 1.0f;
                } else if (global_row > k + c_V) {
                    val = h[base + (k + c_V) * N + global_row];
                }
            }
            V_shared[r][c_V] = val;
        }
        __syncthreads();

        if (local_tid < ib) {
            int limit = min(TILE_ROWS, N - tile_row_start);
            for (int r = 0; r < limit; ++r) {
                int global_row = tile_row_start + r;
                float v_val = V_shared[r][local_tid];

                if (c0 < col_end) w_val0 = fmaf(v_val, h[base + c0 * N + global_row], w_val0);
                if (c1 < col_end) w_val1 = fmaf(v_val, h[base + c1 * N + global_row], w_val1);
                if (c2 < col_end) w_val2 = fmaf(v_val, h[base + c2 * N + global_row], w_val2);
                if (c3 < col_end) w_val3 = fmaf(v_val, h[base + c3 * N + global_row], w_val3);
            }
        }
    }

    __syncthreads();

    // Pass 2: Compute Y = T * W
    float y_val0 = 0.0f;
    float y_val1 = 0.0f;
    float y_val2 = 0.0f;
    float y_val3 = 0.0f;

    for (int p = 0; p < ib; ++p) {
        float wp0 = __shfl_sync(0xffffffff, w_val0, p);
        float wp1 = __shfl_sync(0xffffffff, w_val1, p);
        float wp2 = __shfl_sync(0xffffffff, w_val2, p);
        float wp3 = __shfl_sync(0xffffffff, w_val3, p);

        if (local_tid < ib && p <= local_tid) {
            y_val0 = fmaf(T[p][local_tid], wp0, y_val0);
            y_val1 = fmaf(T[p][local_tid], wp1, y_val1);
            y_val2 = fmaf(T[p][local_tid], wp2, y_val2);
            y_val3 = fmaf(T[p][local_tid], wp3, y_val3);
        }
    }

    // Pass 3: Apply Update (C -= V * Y)
    for (int tile_row_start = k; tile_row_start < N; tile_row_start += TILE_ROWS) {
        __syncthreads();
        for (int load_idx = tid; load_idx < TILE_ROWS * ib; load_idx += APPLY_THREADS) {
            int r = load_idx % TILE_ROWS;
            int c_V = load_idx / TILE_ROWS;
            int global_row = tile_row_start + r;
            float val = 0.0f;
            if (global_row < N && c_V < ib) {
                if (global_row == k + c_V) {
                    val = 1.0f;
                } else if (global_row > k + c_V) {
                    val = h[base + (k + c_V) * N + global_row];
                }
            }
            V_shared[r][c_V] = val;
        }
        __syncthreads();

        int limit = min(TILE_ROWS, N - tile_row_start);
        for (int r = local_tid; r < limit; r += 32) {
            int global_row = tile_row_start + r;
            float sum0 = 0.0f;
            float sum1 = 0.0f;
            float sum2 = 0.0f;
            float sum3 = 0.0f;

            for (int i = 0; i < ib; ++i) {
                float v_val = V_shared[r][i];
                float yi0 = __shfl_sync(0xffffffff, y_val0, i);
                float yi1 = __shfl_sync(0xffffffff, y_val1, i);
                float yi2 = __shfl_sync(0xffffffff, y_val2, i);
                float yi3 = __shfl_sync(0xffffffff, y_val3, i);

                sum0 = fmaf(v_val, yi0, sum0);
                sum1 = fmaf(v_val, yi1, sum1);
                sum2 = fmaf(v_val, yi2, sum2);
                sum3 = fmaf(v_val, yi3, sum3);
            }

            if (c0 < col_end) h[base + c0 * N + global_row] -= sum0;
            if (c1 < col_end) h[base + c1 * N + global_row] -= sum1; 
            if (c2 < col_end) h[base + c2 * N + global_row] -= sum2;
            if (c3 < col_end) h[base + c3 * N + global_row] -= sum3;
        }
    }
}

//=============================================================================
// Host Function
//=============================================================================
void qr_n512_human(torch::Tensor a, torch::Tensor h, torch::Tensor tau)
{
    TORCH_CHECK(a.is_cuda() && h.is_cuda() && tau.is_cuda(),
                "human n512 expects CUDA tensors");

    TORCH_CHECK(a.scalar_type() == torch::kFloat32 &&
                h.scalar_type() == torch::kFloat32 &&
                tau.scalar_type() == torch::kFloat32,
                "human n512 expects FP32 tensors");

    TORCH_CHECK(a.dim() == 3 && a.size(1) == N && a.size(2) == N,
                "human n512 input must have shape [batch, 512, 512]");

    TORCH_CHECK(h.sizes() == a.sizes(),
                "human n512 H must match input shape");

    TORCH_CHECK(tau.dim() == 2 &&
                tau.size(0) == a.size(0) &&
                tau.size(1) == N,
                "human n512 tau must have shape [batch, 512]");

    TORCH_CHECK(a.is_contiguous() && h.is_contiguous() && tau.is_contiguous(),
                "human n512 expects contiguous tensors");

#if HUMAN_N512_IMPLEMENTED
    constexpr int threads = 512;
    int batch = static_cast<int>(a.size(0));

    // Transpose input
    torch::Tensor a_T = a.transpose(1, 2).contiguous();
    torch::Tensor h_T = torch::empty_like(a_T);

    // Copy a_T -> h_T (replaces the in-kernel copy)
    h_T.copy_(a_T);

    float* h_ptr = h_T.data_ptr<float>();
    float* tau_ptr = tau.data_ptr<float>();

    // Allocate workspace for T matrix
    auto options = torch::TensorOptions().device(a.device()).dtype(torch::kFloat32);
    torch::Tensor work_T = torch::empty({batch, BLOCK_SIZE, BLOCK_SIZE}, options);
    float* work_T_ptr = work_T.data_ptr<float>();

    for (int k = 0; k < N; k += BLOCK_SIZE) {
        int ib = (N - k < BLOCK_SIZE) ? (N - k) : BLOCK_SIZE;
        int trailingCols = N - k - ib;

        // Kernel 1: Panel factorization
        kernel_factor_panel<<<batch, threads>>>(h_ptr, tau_ptr, work_T_ptr, k, ib, batch);

        // Kernel 2: Build T + apply to trailing matrix (multi-block per batch)
        if (trailingCols > 0) {
            int num_col_blocks = (trailingCols + COLS_PER_TRAILING_BLOCK - 1)
                                 / COLS_PER_TRAILING_BLOCK;
            kernel_build_T_apply<<<batch * num_col_blocks, APPLY_THREADS>>>(
                h_ptr, tau_ptr, work_T_ptr, k, ib, batch, num_col_blocks);
        }
    }

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

    // Transpose back
    h.copy_(h_T.transpose(1, 2));
#else
    constexpr int MODE_CUBLAS_GRAM_CUBLAS_T = 8;
    qr_blocked_wy(a, h, tau, 16, MODE_CUBLAS_GRAM_CUBLAS_T);
#endif
}
"""
# END HUMAN_CUDA_EMBEDDED


def _load_human_cuda_source() -> str:
    try:
        human_path = Path(__file__).with_name("human_adaptive.cu")
    except NameError:
        return HUMAN_CUDA_EMBEDDED

    if human_path.is_file():
        return human_path.read_text(encoding="utf-8")
    return HUMAN_CUDA_EMBEDDED


HUMAN_CUDA_SRC = _load_human_cuda_source()


import os
_module = load_inline(
    name="qr_compact_householder_b200_v100_human_n512",
    cpp_sources=[CPP_SRC],
    cuda_sources=[CUDA_SRC, HUMAN_CUDA_SRC],
    functions=[
        "qr32",
        "qr_shared_176",
        "qr_global_256",
        "qr_global_512",
        "qr_blocked_wy",
        "qr_n512_custom",
        "qr_n512_human",
    ],
    verbose=False,
    extra_cuda_cflags=["-O3", "--use_fast_math"],
    extra_ldflags=["cublas.lib"] if os.name == "nt" else ["-lcublas"],
)

MODE_CUBLAS_GRAM_CUBLAS_T = 8
MODE_CUBLAS_GRAM_FUSED_T = 9
MODE_CUSTOM_GRAM_FUSED_T = 10
MODE_N512_CUSTOM_UPDATE = 11


def _empty_output(data: input_t, batch: int, n: int) -> output_t:
    return (
        torch.empty_like(data),
        torch.empty((batch, n), device=data.device, dtype=torch.float32),
    )


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, tau = _empty_output(data, batch, n)
            _module.qr32(data, h, tau)
            return h, tau
        if m == n and n == 176:
            h, tau = _empty_output(data, batch, n)
            _module.qr_shared_176(data, h, tau)
            return h, tau
        if m == n and n <= 176:
            h, tau = _empty_output(data, batch, n)
            _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, tau = _empty_output(data, batch, n)
            _module.qr_global_512(data, h, tau)
            return h, tau
        if m == n and n == 512 and batch >= 128:
            # human.cu owns this route. Its disabled stub safely falls back to
            # the old cuBLAS blocked-WY implementation.
            h, tau = _empty_output(data, batch, n)
            _module.qr_n512_human(data, h, tau)
            return h, tau
        if m == n and n == 352:
            h, tau = _empty_output(data, batch, n)
            _module.qr_blocked_wy(data, h, tau, 16, MODE_CUBLAS_GRAM_FUSED_T)
            return h, tau
        if m == n and n == 1024:
            h, tau = _empty_output(data, batch, n)
            _module.qr_blocked_wy(data, h, tau, 16, MODE_CUBLAS_GRAM_FUSED_T)
            return h, tau
        if m == n and n == 2048:
            h, tau = _empty_output(data, batch, n)
            _module.qr_blocked_wy(data, h, tau, 16, MODE_CUSTOM_GRAM_FUSED_T)
            return h, tau
    return torch.geqrf(data)
scrolls · 1828 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