Skip to content
KernelIndex
Search⌘K

submission 801995

Álvaro Borrás Fernández · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-801995?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
18.5ms
#340 of 515
2026-06-16

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c8bf54bba67abc34d9d09cbb83557c519313de308d5ffe717db2fdad45a02ae7
license declaredunknown
license concludedunknown
authorsÁlvaro Borrás Fernández
imported2026-08-26

Techniques

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

shared-memory__shared__ float mat[32 * 32];

Kernel source

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

import gc
import os

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

gc.disable()
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
try:
    torch.set_float32_matmul_precision('high')
except Exception:
    pass


CUDA_SRC = r"""
#include <cuda_runtime.h>
#include <mma.h>
#include <stdexcept>
#include <string>

#define THREADS 256
#define WARP_SIZE 32

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

__device__ __forceinline__ float warp_reduce_sum(float val) {
    #pragma unroll
    for (int offset = 16; offset > 0; offset >>= 1) {
        val += __shfl_xor_sync(0xffffffffu, val, offset);
    }
    return val;
}

__device__ __forceinline__ float block_reduce_sum_64(float val, float* buf) {
    int tid = threadIdx.x;
    buf[tid] = val;
    __syncthreads();
    #pragma unroll
    for (int offset = 32; offset > 0; offset >>= 1) {
        if (tid < offset) buf[tid] += buf[tid + offset];
        __syncthreads();
    }
    return buf[0];
}

__global__ void copy_input_kernel_512(
    const float* __restrict__ a,
    float* __restrict__ h,
    float* __restrict__ tau,
    int batch,
    long long total
) {
    long long stride = (long long)blockDim.x * gridDim.x;
    for (long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
         idx < total; idx += stride) {
        h[idx] = a[idx];
    }
    for (int b = blockIdx.x * blockDim.x + threadIdx.x;
         b < batch; b += blockDim.x * gridDim.x) {
        tau[(long long)b * 512 + 511] = 0.0f;
    }
}

__global__ void qr32_kernel(
    const float* __restrict__ a,
    float* __restrict__ h,
    float* __restrict__ tau
) {
    __shared__ float mat[32 * 32];
    __shared__ float reduce_buf[64];
    __shared__ float tau_s;
    __shared__ float inv_s;
    __shared__ float dot_s;
    __shared__ int active_s;

    int b = blockIdx.x;
    int tid = threadIdx.x;
    long long base = (long long)b * 32 * 32;

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

    for (int k = 0; k < 31; ++k) {
        float ss = 0.0f;
        for (int i = k + tid; i < 32; i += blockDim.x) {
            float x = mat[i * 32 + k];
            ss = fmaf(x, x, ss);
        }
        float norm_sq = block_reduce_sum_64(ss, reduce_buf);

        if (tid == 0) {
            float norm = sqrtf(norm_sq);
            float diag = mat[k * 32 + k];
            if (norm <= 1e-20f) {
                tau_s = 0.0f;
                inv_s = 0.0f;
                active_s = 0;
                tau[(long long)b * 32 + k] = 0.0f;
            } else {
                float alpha = (diag >= 0.0f) ? -norm : norm;
                float denom = diag - alpha;
                inv_s = 1.0f / denom;
                tau_s = (alpha - diag) / alpha;
                active_s = 1;
                mat[k * 32 + k] = alpha;
                tau[(long long)b * 32 + k] = tau_s;
            }
        }
        __syncthreads();

        if (active_s == 0) continue;

        for (int i = k + 1 + tid; i < 32; i += blockDim.x) {
            mat[i * 32 + k] *= inv_s;
        }
        __syncthreads();

        for (int j = k + 1; j < 32; ++j) {
            float dot_part = (tid == 0) ? mat[k * 32 + j] : 0.0f;
            for (int i = k + 1 + tid; i < 32; i += blockDim.x) {
                dot_part = fmaf(mat[i * 32 + k], mat[i * 32 + j], dot_part);
            }
            float dot = block_reduce_sum_64(dot_part, reduce_buf);
            if (tid == 0) dot_s = dot;
            __syncthreads();

            if (tid == 0) {
                mat[k * 32 + j] -= tau_s * dot_s;
            }
            for (int i = k + 1 + tid; i < 32; i += blockDim.x) {
                mat[i * 32 + j] -= tau_s * mat[i * 32 + k] * dot_s;
            }
            __syncthreads();
        }
    }

    if (tid == 0) tau[(long long)b * 32 + 31] = 0.0f;
    for (int idx = tid; idx < 32 * 32; idx += blockDim.x) {
        h[base + idx] = mat[idx];
    }
}

template <int BLOCK_SIZE, int NUM_WARPS>
__global__ void qr512_panel_factor(
    float* __restrict__ h,
    float* __restrict__ tau,
    int kk,
    int ib
) {
    extern __shared__ float smem[];
    float* v = smem;
    float* warp_buf = smem + 512;

    int b = blockIdx.x;
    int tid = threadIdx.x;
    int wid = tid / WARP_SIZE;
    int lid = tid % WARP_SIZE;

    float* __restrict__ h_b = h + (long long)b * 512 * 512;
    float* __restrict__ tau_b = tau + (long long)b * 512;

    __shared__ float tau_s;
    __shared__ float inv_s;
    __shared__ float dot_s;
    __shared__ int active_s;

    for (int local_k = 0; local_k < ib && kk + local_k < 511; ++local_k) {
        int k = kk + local_k;
        int m = 512 - k;
        int panel_end = kk + ib;

        float s = 0.0f;
        for (int i = tid; i < m; i += BLOCK_SIZE) {
            float x = h_b[(k + i) * 512 + k];
            s = fmaf(x, x, s);
        }

        s = warp_reduce_sum(s);
        if (lid == 0) warp_buf[wid] = s;
        __syncthreads();

        if (wid == 0) {
            s = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
            s = warp_reduce_sum(s);
            if (lid == 0) warp_buf[0] = s;
        }
        __syncthreads();

        if (tid == 0) {
            float norm = sqrtf(warp_buf[0]);
            float diag = h_b[k * 512 + k];
            if (norm <= 1e-20f) {
                tau_s = 0.0f;
                inv_s = 0.0f;
                active_s = 0;
                tau_b[k] = 0.0f;
            } else {
                float alpha = (diag >= 0.0f) ? -norm : norm;
                float denom = diag - alpha;
                inv_s = 1.0f / denom;
                tau_s = (alpha - diag) / alpha;
                active_s = 1;
                h_b[k * 512 + k] = alpha;
                tau_b[k] = tau_s;
            }
        }
        __syncthreads();

        if (active_s != 0) {
            if (tid == 0) v[0] = 1.0f;
            for (int i = tid; i < m - 1; i += BLOCK_SIZE) {
                float vi = h_b[(k + 1 + i) * 512 + k] * inv_s;
                v[i + 1] = vi;
                h_b[(k + 1 + i) * 512 + k] = vi;
            }
            __syncthreads();

            for (int j = k + 1; j < panel_end; ++j) {
                float dot_part = (tid == 0) ? h_b[k * 512 + j] : 0.0f;
                for (int i = k + 1 + tid; i < 512; i += BLOCK_SIZE) {
                    dot_part = fmaf(h_b[i * 512 + k], h_b[i * 512 + j], dot_part);
                }

                dot_part = warp_reduce_sum(dot_part);
                if (lid == 0) warp_buf[wid] = dot_part;
                __syncthreads();

                if (wid == 0) {
                    dot_part = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
                    dot_part = warp_reduce_sum(dot_part);
                    if (lid == 0) dot_s = dot_part;
                }
                __syncthreads();

                if (tid == 0) {
                    h_b[k * 512 + j] -= tau_s * dot_s;
                }
                for (int i = k + 1 + tid; i < 512; i += BLOCK_SIZE) {
                    h_b[i * 512 + j] -= tau_s * h_b[i * 512 + k] * dot_s;
                }
                __syncthreads();
            }
        }
        __syncthreads();
    }
}

template <int NB, int BLOCK_SIZE>
__global__ void qr512_build_t(
    const float* __restrict__ h,
    const float* __restrict__ tau,
    float* __restrict__ workspace,
    int kk,
    int ib,
    int workspace_stride
) {
    __shared__ float reduce_buf[BLOCK_SIZE];
    __shared__ float tmp[NB];

    int b = blockIdx.x;
    int tid = threadIdx.x;

    const float* __restrict__ h_b = h + (long long)b * 512 * 512;
    const float* __restrict__ tau_b = tau + (long long)b * 512;
    float* __restrict__ t = workspace + (long long)b * workspace_stride;

    for (int idx = tid; idx < NB * NB; idx += BLOCK_SIZE) {
        t[idx] = 0.0f;
    }
    __syncthreads();

    for (int j = 0; j < ib; ++j) {
        float tau_j = tau_b[kk + j];
        if (tau_j == 0.0f) {
            if (tid == 0) t[j * NB + j] = 0.0f;
            __syncthreads();
            continue;
        }

        for (int idx = tid; idx < NB; idx += BLOCK_SIZE) {
            tmp[idx] = 0.0f;
        }
        __syncthreads();

        for (int i = 0; i < j; ++i) {
            float dot = (tid == 0) ? h_b[(kk + j) * 512 + (kk + i)] : 0.0f;
            for (int r = j + 1 + tid; r < 512 - kk; r += BLOCK_SIZE) {
                dot = fmaf(
                    h_b[(kk + r) * 512 + (kk + i)],
                    h_b[(kk + r) * 512 + (kk + j)],
                    dot
                );
            }
            reduce_buf[tid] = dot;
            __syncthreads();
            for (int offset = BLOCK_SIZE / 2; offset > 0; offset >>= 1) {
                if (tid < offset) {
                    reduce_buf[tid] += reduce_buf[tid + offset];
                }
                __syncthreads();
            }
            if (tid == 0) {
                tmp[i] = -tau_j * reduce_buf[0];
            }
            __syncthreads();
        }

        if (tid == 0) {
            for (int row = 0; row < j; ++row) {
                float val = 0.0f;
                for (int col = row; col < j; ++col) {
                    val = fmaf(t[row * NB + col], tmp[col], val);
                }
                t[row * NB + j] = val;
            }
            t[j * NB + j] = tau_j;
        }
        __syncthreads();
    }
}

template <int BLOCK_SIZE, int NUM_WARPS>
__global__ void qr512_panel_build_t_fused(
    float* __restrict__ h,
    float* __restrict__ tau,
    float* __restrict__ workspace,
    int kk,
    int ib,
    int workspace_stride
) {
    constexpr int NB = 8;
    extern __shared__ float smem[];
    float* v = smem;
    float* warp_buf = smem + 512;

    int b = blockIdx.x;
    int tid = threadIdx.x;
    int wid = tid / WARP_SIZE;
    int lid = tid % WARP_SIZE;

    float* __restrict__ h_b = h + (long long)b * 512 * 512;
    float* __restrict__ tau_b = tau + (long long)b * 512;
    float* __restrict__ t = workspace + (long long)b * workspace_stride;

    __shared__ float tau_s;
    __shared__ float inv_s;
    __shared__ float dot_s;
    __shared__ int active_s;

    for (int local_k = 0; local_k < ib && kk + local_k < 511; ++local_k) {
        int k = kk + local_k;
        int m = 512 - k;
        int panel_end = kk + ib;

        float s = 0.0f;
        for (int i = tid; i < m; i += BLOCK_SIZE) {
            float x = h_b[(k + i) * 512 + k];
            s = fmaf(x, x, s);
        }

        s = warp_reduce_sum(s);
        if (lid == 0) warp_buf[wid] = s;
        __syncthreads();

        if (wid == 0) {
            s = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
            s = warp_reduce_sum(s);
            if (lid == 0) warp_buf[0] = s;
        }
        __syncthreads();

        if (tid == 0) {
            float norm = sqrtf(warp_buf[0]);
            float diag = h_b[k * 512 + k];
            if (norm <= 1e-20f) {
                tau_s = 0.0f;
                inv_s = 0.0f;
                active_s = 0;
                tau_b[k] = 0.0f;
            } else {
                float alpha = (diag >= 0.0f) ? -norm : norm;
                float denom = diag - alpha;
                inv_s = 1.0f / denom;
                tau_s = (alpha - diag) / alpha;
                active_s = 1;
                h_b[k * 512 + k] = alpha;
                tau_b[k] = tau_s;
            }
        }
        __syncthreads();

        if (active_s != 0) {
            if (tid == 0) v[0] = 1.0f;
            for (int i = tid; i < m - 1; i += BLOCK_SIZE) {
                float vi = h_b[(k + 1 + i) * 512 + k] * inv_s;
                v[i + 1] = vi;
                h_b[(k + 1 + i) * 512 + k] = vi;
            }
            __syncthreads();

            for (int j = k + 1; j < panel_end; ++j) {
                float dot_part = (tid == 0) ? h_b[k * 512 + j] : 0.0f;
                for (int i = k + 1 + tid; i < 512; i += BLOCK_SIZE) {
                    dot_part = fmaf(v[i - k], h_b[i * 512 + j], dot_part);
                }

                dot_part = warp_reduce_sum(dot_part);
                if (lid == 0) warp_buf[wid] = dot_part;
                __syncthreads();

                if (wid == 0) {
                    dot_part = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
                    dot_part = warp_reduce_sum(dot_part);
                    if (lid == 0) dot_s = dot_part;
                }
                __syncthreads();

                if (tid == 0) {
                    h_b[k * 512 + j] -= tau_s * dot_s;
                }
                for (int i = k + 1 + tid; i < 512; i += BLOCK_SIZE) {
                    h_b[i * 512 + j] -= tau_s * v[i - k] * dot_s;
                }
                __syncthreads();
            }
        }
        __syncthreads();
    }

    if (kk + ib >= 512) return;

    float* tmp = smem;

    for (int idx = tid; idx < NB * NB; idx += BLOCK_SIZE) {
        t[idx] = 0.0f;
    }
    __syncthreads();

    for (int j = 0; j < ib; ++j) {
        float tau_j = tau_b[kk + j];
        if (tau_j == 0.0f) {
            if (tid == 0) t[j * NB + j] = 0.0f;
            __syncthreads();
            continue;
        }

        for (int idx = tid; idx < NB; idx += BLOCK_SIZE) {
            tmp[idx] = 0.0f;
        }
        __syncthreads();

        for (int i = 0; i < j; ++i) {
            float dot_part = (tid == 0) ? h_b[(kk + j) * 512 + (kk + i)] : 0.0f;
            for (int r = j + 1 + tid; r < 512 - kk; r += BLOCK_SIZE) {
                dot_part = fmaf(
                    h_b[(kk + r) * 512 + (kk + i)],
                    h_b[(kk + r) * 512 + (kk + j)],
                    dot_part
                );
            }

            dot_part = warp_reduce_sum(dot_part);
            if (lid == 0) warp_buf[wid] = dot_part;
            __syncthreads();

            if (wid == 0) {
                dot_part = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
                dot_part = warp_reduce_sum(dot_part);
                if (lid == 0) warp_buf[0] = dot_part;
            }
            __syncthreads();

            if (tid == 0) {
                tmp[i] = -tau_j * warp_buf[0];
            }
            __syncthreads();
        }

        if (tid == 0) {
            for (int row = 0; row < j; ++row) {
                float val = 0.0f;
                for (int col = row; col < j; ++col) {
                    val = fmaf(t[row * NB + col], tmp[col], val);
                }
                t[row * NB + j] = val;
            }
            t[j * NB + j] = tau_j;
        }
        __syncthreads();
    }
}

template <int NB, int COL_TILE, int RSPLIT>
__global__ void qr512_compute_y(
    const float* __restrict__ h,
    const float* __restrict__ workspace,
    float* __restrict__ y_workspace,
    int kk,
    int ib,
    int workspace_stride
) {
    __shared__ float partial_s[NB * COL_TILE * RSPLIT];
    __shared__ float w_s[NB * COL_TILE];

    int b = blockIdx.x;
    int col_tile = blockIdx.y;
    int idx = threadIdx.x;
    int split = idx % RSPLIT;
    int c = (idx / RSPLIT) % COL_TILE;
    int q = idx / (RSPLIT * COL_TILE);
    int trail_col = col_tile * COL_TILE + c;
    int global_col = kk + ib + trail_col;
    int pair = q * COL_TILE + c;

    const float* __restrict__ h_b = h + (long long)b * 512 * 512;
    const float* __restrict__ t = workspace + (long long)b * workspace_stride;
    float* __restrict__ y = y_workspace + (long long)b * NB * 512;

    if (q < ib && c < COL_TILE && global_col < 512) {
        float sum = (split == 0) ? h_b[(kk + q) * 512 + global_col] : 0.0f;
        for (int r = q + 1 + split; r < 512 - kk; r += RSPLIT) {
            sum = fmaf(
                h_b[(kk + r) * 512 + (kk + q)],
                h_b[(kk + r) * 512 + global_col],
                sum
            );
        }
        partial_s[pair * RSPLIT + split] = sum;
    } else if (q < NB && c < COL_TILE) {
        partial_s[pair * RSPLIT + split] = 0.0f;
    }
    __syncthreads();

    if (split == 0 && q < ib && c < COL_TILE && global_col < 512) {
        float w = 0.0f;
        #pragma unroll
        for (int s = 0; s < RSPLIT; ++s) {
            w += partial_s[pair * RSPLIT + s];
        }
        w_s[pair] = w;
    }
    __syncthreads();

    if (split == 0 && q < ib && c < COL_TILE && global_col < 512) {
        float sum = 0.0f;
        for (int l = 0; l <= q; ++l) {
            sum = fmaf(t[l * NB + q], w_s[l * COL_TILE + c], sum);
        }
        y[q * 512 + global_col] = sum;
    }
}

template <int NB, int ROW_TILE, int COL_TILE>
__global__ void qr512_apply_block(
    float* __restrict__ h,
    const float* __restrict__ y_workspace,
    int kk,
    int ib
) {
    int b = blockIdx.x;
    int row_tile = blockIdx.y;
    int col_tile = blockIdx.z;
    int idx = threadIdx.x;
    int r_local = row_tile * ROW_TILE + idx / COL_TILE;
    int c_local = col_tile * COL_TILE + idx % COL_TILE;
    int global_col = kk + ib + c_local;

    if (r_local >= 512 - kk || global_col >= 512) return;

    float* __restrict__ h_b = h + (long long)b * 512 * 512;
    const float* __restrict__ y = y_workspace + (long long)b * NB * 512;

    float correction = 0.0f;
    int nq = (r_local < ib) ? r_local : ib;
    #pragma unroll
    for (int q = 0; q < nq; ++q) {
        float vq = h_b[(kk + r_local) * 512 + (kk + q)];
        correction = fmaf(vq, y[q * 512 + global_col], correction);
    }
    if (r_local < ib) {
        correction = fmaf(1.0f, y[r_local * 512 + global_col], correction);
    }

    h_b[(kk + r_local) * 512 + global_col] -= correction;
}

// Tensor-core (WMMA m16n16k16 tf32) bottom-apply pass for n=512.
// Computes H[bottom, bottom] -= V[bottom, NB] @ Y[NB, bottom] where
// bottom rows = [kk+ib, 512), NB=8.  NB is padded to 16 with zeros.
//
// Note: the nvcuda::wmma tf32 fragment template is not exposed for sm_100
// (B200) in CUDA 12.8's mma.h, so this kernel uses a scalar fallback
// (manual matmul) that compiles on all arches. The dispatcher on sm_100
// falls through to the CuTe path, so this scalar kernel is not on the hot
// path there.
__global__ void qr512_apply_block_wmma(
    float* __restrict__ h,
    const float* __restrict__ y_workspace,
    int kk,
    int ib
) {
    constexpr int NB = 8;
    constexpr int NB_PAD = 16;
    constexpr int ROW_TILE = 16;
    constexpr int COL_TILE = 16;
    constexpr int N = 512;
    constexpr int LDH = N;
    constexpr int LDY = N;

    int b = blockIdx.x;
    int row_tile = blockIdx.y;
    int col_tile = blockIdx.z;
    int tid = threadIdx.x;

    int global_row_base = kk + ib + row_tile * ROW_TILE;
    int global_col_base = kk + ib + col_tile * COL_TILE;
    if (global_row_base >= N || global_col_base >= N) return;

    float* __restrict__ h_b = h + (long long)b * N * N;
    const float* __restrict__ y = y_workspace + (long long)b * NB * N;

    __shared__ float V_smem[ROW_TILE * NB_PAD];
    __shared__ float Y_smem[NB_PAD * COL_TILE];

    #pragma unroll
    for (int i = tid; i < ROW_TILE * NB_PAD; i += 32) {
        int r = i / NB_PAD;
        int c = i - r * NB_PAD;
        float v = 0.0f;
        if (c < NB) {
            int gr = kk + ib + row_tile * ROW_TILE + r;
            if (gr < N) {
                v = h_b[gr * LDH + (kk + c)];
            }
        }
        V_smem[r * NB_PAD + c] = v;
    }

    #pragma unroll
    for (int i = tid; i < NB_PAD * COL_TILE; i += 32) {
        int r = i / COL_TILE;
        int c = i - r * COL_TILE;
        float v = 0.0f;
        if (r < NB) {
            int gc = kk + ib + col_tile * COL_TILE + c;
            if (gc < N) {
                v = y[r * LDY + gc];
            }
        }
        Y_smem[r * COL_TILE + c] = v;
    }

    __syncwarp();

    float acc[ROW_TILE * COL_TILE / 32] = {0.0f};
    constexpr int PER_THREAD = (ROW_TILE * COL_TILE) / 32;
    #pragma unroll
    for (int k = 0; k < NB; ++k) {
        #pragma unroll
        for (int j = 0; j < COL_TILE; ++j) {
            int slot = (j / 1);
            int col_in_tile = (tid + j * 32 / COL_TILE) % COL_TILE;
            (void)slot; (void)col_in_tile;
        }
        #pragma unroll
        for (int i = 0; i < PER_THREAD; ++i) {
            int linear = i * 32 + tid;
            int r = linear / COL_TILE;
            int c = linear - r * COL_TILE;
            acc[i] = fmaf(V_smem[r * NB_PAD + k], Y_smem[k * COL_TILE + c], acc[i]);
        }
    }

    #pragma unroll
    for (int i = 0; i < PER_THREAD; ++i) {
        int linear = i * 32 + tid;
        int r = linear / COL_TILE;
        int c = linear - r * COL_TILE;
        int gr = kk + ib + row_tile * ROW_TILE + r;
        int gc = kk + ib + col_tile * COL_TILE + c;
        if (gr < N && gc < N) {
            h_b[gr * LDH + gc] -= acc[i];
        }
    }
}

void qr512_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace) {
    constexpr int NB = 8;
    constexpr int COL_TILE = 32;
    constexpr int ROW_TILE = 16;
    constexpr int WORKSPACE_STRIDE = NB * NB + NB * 512;
    int batch = (int)a.size(0);
    long long total = (long long)batch * 512 * 512;
    int copy_blocks = (int)((total + THREADS - 1) / THREADS);
    if (copy_blocks < 1) copy_blocks = 1;
    if (copy_blocks > 8192) copy_blocks = 8192;

    copy_input_kernel_512<<<copy_blocks, THREADS>>>(
        a.data_ptr<float>(), h.data_ptr<float>(), tau.data_ptr<float>(), batch, total
    );

    float* workspace_ptr = workspace.data_ptr<float>();
    float* y_ptr = workspace_ptr + (long long)batch * WORKSPACE_STRIDE;

    for (int kk = 0; kk < 511; kk += NB) {
        int ib = min(NB, 512 - kk);
        qr512_panel_build_t_fused<256, 8><<<batch, 256, (512 + 8) * (int)sizeof(float)>>>(
            h.data_ptr<float>(), tau.data_ptr<float>(), workspace_ptr, kk, ib, WORKSPACE_STRIDE
        );

        if (kk + ib < 512) {
            int col_tiles = (512 - kk - ib + COL_TILE - 1) / COL_TILE;
            dim3 y_grid((unsigned int)batch, (unsigned int)col_tiles);
            qr512_compute_y<NB, COL_TILE, 2><<<y_grid, NB * COL_TILE * 2>>>(
                h.data_ptr<float>(), workspace_ptr, y_ptr, kk, ib, WORKSPACE_STRIDE
            );

            int row_tiles = (512 - kk + ROW_TILE - 1) / ROW_TILE;
            dim3 update_grid((unsigned int)batch, (unsigned int)row_tiles, (unsigned int)col_tiles);
            qr512_apply_block<NB, ROW_TILE, COL_TILE><<<update_grid, ROW_TILE * COL_TILE>>>(
                h.data_ptr<float>(), y_ptr, kk, ib
            );
        }
    }

    check_cuda(cudaGetLastError(), "qr512_blocked_cuda");
}

void qr512_panel_build_t_cuda(torch::Tensor h, torch::Tensor tau, torch::Tensor workspace, int kk, int batch) {
    constexpr int NB = 8;
    constexpr int WORKSPACE_STRIDE = NB * NB + NB * 512;
    int ib = min(NB, 512 - kk);
    qr512_panel_build_t_fused<256, 8><<<batch, 256, (512 + 8) * (int)sizeof(float)>>>(
        h.data_ptr<float>(), tau.data_ptr<float>(), workspace.data_ptr<float>(), kk, ib, WORKSPACE_STRIDE
    );
}

void qr32_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau) {
    int batch = (int)a.size(0);
    qr32_kernel<<<batch, 64>>>(
        a.data_ptr<float>(),
        h.data_ptr<float>(),
        tau.data_ptr<float>()
    );
    check_cuda(cudaGetLastError(), "qr32_cuda");
}

// WMMA tf32 fragments are not exposed by mma.h for sm_100 (B200), so the
// Host-side launcher for the scalar-fallback apply kernel. Compiles on all
// arches; on sm_100 the dispatcher falls through to the CuTe path so this
// is a dead code path.
void qr512_apply_block_wmma_cuda(torch::Tensor h, torch::Tensor y, int kk, int ib, int batch) {
    constexpr int NB = 8;
    constexpr int ROW_TILE = 16;
    constexpr int COL_TILE = 16;
    int row_tiles = (512 - kk - ib + ROW_TILE - 1) / ROW_TILE;
    int col_tiles = (512 - kk - ib + COL_TILE - 1) / COL_TILE;
    if (row_tiles < 0) row_tiles = 0;
    if (col_tiles < 0) col_tiles = 0;
    if (row_tiles == 0 || col_tiles == 0) return;
    dim3 grid((unsigned int)batch, (unsigned int)row_tiles, (unsigned int)col_tiles);
    qr512_apply_block_wmma<<<grid, 32>>>(
        h.data_ptr<float>(), y.data_ptr<float>(), kk, ib
    );
}

__global__ void copy_input_kernel_1024(
    const float* __restrict__ a,
    float* __restrict__ h,
    float* __restrict__ tau,
    int batch,
    long long total
) {
    long long stride = (long long)blockDim.x * gridDim.x;
    for (long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
         idx < total; idx += stride) {
        h[idx] = a[idx];
    }
    for (int b = blockIdx.x * blockDim.x + threadIdx.x;
         b < batch; b += blockDim.x * gridDim.x) {
        tau[(long long)b * 1024 + 1023] = 0.0f;
    }
}

template <int BLOCK_SIZE, int NUM_WARPS>
__global__ void qr1024_panel_build_t_fused(
    float* __restrict__ h,
    float* __restrict__ tau,
    float* __restrict__ workspace,
    int kk,
    int ib,
    int workspace_stride
) {
    constexpr int NB = 8;
    extern __shared__ float smem[];
    float* v = smem;
    float* warp_buf = smem + 1024;

    int b = blockIdx.x;
    int tid = threadIdx.x;
    int wid = tid / WARP_SIZE;
    int lid = tid % WARP_SIZE;

    float* __restrict__ h_b = h + (long long)b * 1024 * 1024;
    float* __restrict__ tau_b = tau + (long long)b * 1024;
    float* __restrict__ t = workspace + (long long)b * workspace_stride;

    __shared__ float tau_s;
    __shared__ float inv_s;
    __shared__ float dot_s;
    __shared__ int active_s;

    for (int local_k = 0; local_k < ib && kk + local_k < 1023; ++local_k) {
        int k = kk + local_k;
        int m = 1024 - k;
        int panel_end = kk + ib;

        float s = 0.0f;
        for (int i = tid; i < m; i += BLOCK_SIZE) {
            float x = h_b[(k + i) * 1024 + k];
            s = fmaf(x, x, s);
        }

        s = warp_reduce_sum(s);
        if (lid == 0) warp_buf[wid] = s;
        __syncthreads();

        if (wid == 0) {
            s = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
            s = warp_reduce_sum(s);
            if (lid == 0) warp_buf[0] = s;
        }
        __syncthreads();

        if (tid == 0) {
            float norm = sqrtf(warp_buf[0]);
            float diag = h_b[k * 1024 + k];
            if (norm <= 1e-20f) {
                tau_s = 0.0f;
                inv_s = 0.0f;
                active_s = 0;
                tau_b[k] = 0.0f;
            } else {
                float alpha = (diag >= 0.0f) ? -norm : norm;
                float denom = diag - alpha;
                inv_s = 1.0f / denom;
                tau_s = (alpha - diag) / alpha;
                active_s = 1;
                h_b[k * 1024 + k] = alpha;
                tau_b[k] = tau_s;
            }
        }
        __syncthreads();

        if (active_s != 0) {
            if (tid == 0) v[0] = 1.0f;
            for (int i = tid; i < m - 1; i += BLOCK_SIZE) {
                float vi = h_b[(k + 1 + i) * 1024 + k] * inv_s;
                v[i + 1] = vi;
                h_b[(k + 1 + i) * 1024 + k] = vi;
            }
            __syncthreads();

            for (int j = k + 1; j < panel_end; ++j) {
                float dot_part = (tid == 0) ? h_b[k * 1024 + j] : 0.0f;
                for (int i = k + 1 + tid; i < 1024; i += BLOCK_SIZE) {
                    dot_part = fmaf(v[i - k], h_b[i * 1024 + j], dot_part);
                }

                dot_part = warp_reduce_sum(dot_part);
                if (lid == 0) warp_buf[wid] = dot_part;
                __syncthreads();

                if (wid == 0) {
                    dot_part = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
                    dot_part = warp_reduce_sum(dot_part);
                    if (lid == 0) dot_s = dot_part;
                }
                __syncthreads();

                if (tid == 0) {
                    h_b[k * 1024 + j] -= tau_s * dot_s;
                }
                for (int i = k + 1 + tid; i < 1024; i += BLOCK_SIZE) {
                    h_b[i * 1024 + j] -= tau_s * v[i - k] * dot_s;
                }
                __syncthreads();
            }
        }
        __syncthreads();
    }

    if (kk + ib >= 1024) return;

    float* tmp = smem;

    for (int idx = tid; idx < NB * NB; idx += BLOCK_SIZE) {
        t[idx] = 0.0f;
    }
    __syncthreads();

    for (int j = 0; j < ib; ++j) {
        float tau_j = tau_b[kk + j];
        if (tau_j == 0.0f) {
            if (tid == 0) t[j * NB + j] = 0.0f;
            __syncthreads();
            continue;
        }

        for (int idx = tid; idx < NB; idx += BLOCK_SIZE) {
            tmp[idx] = 0.0f;
        }
        __syncthreads();

        for (int i = 0; i < j; ++i) {
            float dot_part = (tid == 0) ? h_b[(kk + j) * 1024 + (kk + i)] : 0.0f;
            for (int r = j + 1 + tid; r < 1024 - kk; r += BLOCK_SIZE) {
                dot_part = fmaf(
                    h_b[(kk + r) * 1024 + (kk + i)],
                    h_b[(kk + r) * 1024 + (kk + j)],
                    dot_part
                );
            }

            dot_part = warp_reduce_sum(dot_part);
            if (lid == 0) warp_buf[wid] = dot_part;
            __syncthreads();

            if (wid == 0) {
                dot_part = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
                dot_part = warp_reduce_sum(dot_part);
                if (lid == 0) warp_buf[0] = dot_part;
            }
            __syncthreads();

            if (tid == 0) {
                tmp[i] = -tau_j * warp_buf[0];
            }
            __syncthreads();
        }

        if (tid == 0) {
            for (int row = 0; row < j; ++row) {
                float val = 0.0f;
                for (int col = row; col < j; ++col) {
                    val = fmaf(t[row * NB + col], tmp[col], val);
                }
                t[row * NB + j] = val;
            }
            t[j * NB + j] = tau_j;
        }
        __syncthreads();
    }}

template <int NB, int COL_TILE, int RSPLIT>
__global__ void qr1024_compute_y(
    const float* __restrict__ h,
    const float* __restrict__ workspace,
    float* __restrict__ y_workspace,
    int kk,
    int ib,
    int workspace_stride
) {
    __shared__ float partial_s[NB * COL_TILE * RSPLIT];
    __shared__ float w_s[NB * COL_TILE];

    int b = blockIdx.x;
    int col_tile = blockIdx.y;
    int idx = threadIdx.x;
    int split = idx % RSPLIT;
    int c = (idx / RSPLIT) % COL_TILE;
    int q = idx / (RSPLIT * COL_TILE);
    int trail_col = col_tile * COL_TILE + c;
    int global_col = kk + ib + trail_col;
    int pair = q * COL_TILE + c;

    const float* __restrict__ h_b = h + (long long)b * 1024 * 1024;
    const float* __restrict__ t = workspace + (long long)b * workspace_stride;
    float* __restrict__ y = y_workspace + (long long)b * NB * 1024;

    if (q < ib && c < COL_TILE && global_col < 1024) {
        float sum = (split == 0) ? h_b[(kk + q) * 1024 + global_col] : 0.0f;
        for (int r = q + 1 + split; r < 1024 - kk; r += RSPLIT) {
            sum = fmaf(
                h_b[(kk + r) * 1024 + (kk + q)],
                h_b[(kk + r) * 1024 + global_col],
                sum
            );
        }
        partial_s[pair * RSPLIT + split] = sum;
    } else if (q < NB && c < COL_TILE) {
        partial_s[pair * RSPLIT + split] = 0.0f;
    }
    __syncthreads();

    if (split == 0 && q < ib && c < COL_TILE && global_col < 1024) {
        float w = 0.0f;
        #pragma unroll
        for (int s = 0; s < RSPLIT; ++s) {
            w += partial_s[pair * RSPLIT + s];
        }
        w_s[pair] = w;
    }
    __syncthreads();

    if (split == 0 && q < ib && c < COL_TILE && global_col < 1024) {
        float sum = 0.0f;
        for (int l = 0; l <= q; ++l) {
            sum = fmaf(t[l * NB + q], w_s[l * COL_TILE + c], sum);
        }
        y[q * 1024 + global_col] = sum;
    }
}

template <int NB, int ROW_TILE, int COL_TILE>
__global__ void qr1024_apply_block(
    float* __restrict__ h,
    const float* __restrict__ y_workspace,
    int kk,
    int ib
) {
    __shared__ float V_smem[ROW_TILE * NB];

    int b = blockIdx.x;
    int row_tile = blockIdx.y;
    int col_tile = blockIdx.z;
    int idx = threadIdx.x;

    if (idx < ROW_TILE * NB) {
        int r = idx / NB;
        int q = idx % NB;
        int global_r = kk + row_tile * ROW_TILE + r;
        if (global_r < 1024) {
            float* __restrict__ h_b = h + (long long)b * 1024 * 1024;
            V_smem[idx] = h_b[global_r * 1024 + (kk + q)];
        } else {
            V_smem[idx] = 0.0f;
        }
    }
    __syncthreads();

    int r_local = row_tile * ROW_TILE + idx / COL_TILE;
    int c_local = col_tile * COL_TILE + idx % COL_TILE;
    int global_col = kk + ib + c_local;

    if (r_local >= 1024 - kk || global_col >= 1024) return;

    float* __restrict__ h_b = h + (long long)b * 1024 * 1024;
    const float* __restrict__ y = y_workspace + (long long)b * NB * 1024;

    int local_r = idx / COL_TILE;
    int nq = (r_local < ib) ? r_local : ib;
    float correction = 0.0f;
    #pragma unroll
    for (int q = 0; q < nq; ++q) {
        float vq = V_smem[local_r * NB + q];
        correction = fmaf(vq, y[q * 1024 + global_col], correction);
    }
    if (r_local < ib) {
        correction = fmaf(1.0f, y[r_local * 1024 + global_col], correction);
    }

    h_b[(kk + r_local) * 1024 + global_col] -= correction;
}

void qr1024_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace) {
    constexpr int NB = 8;
    constexpr int COL_TILE = 32;
    constexpr int ROW_TILE = 16;
    constexpr int WORKSPACE_STRIDE = NB * NB + NB * 1024;
    int batch = (int)a.size(0);
    long long total = (long long)batch * 1024 * 1024;
    int copy_blocks = (int)((total + THREADS - 1) / THREADS);
    if (copy_blocks < 1) copy_blocks = 1;
    if (copy_blocks > 8192) copy_blocks = 8192;

    copy_input_kernel_1024<<<copy_blocks, THREADS>>>(
        a.data_ptr<float>(), h.data_ptr<float>(), tau.data_ptr<float>(), batch, total
    );

    float* workspace_ptr = workspace.data_ptr<float>();
    float* y_ptr = workspace_ptr + (long long)batch * WORKSPACE_STRIDE;

    for (int kk = 0; kk < 1023; kk += NB) {
        int ib = min(NB, 1024 - kk);
        qr1024_panel_build_t_fused<256, 8><<<batch, 256, (1024 + 8) * (int)sizeof(float)>>>(
            h.data_ptr<float>(), tau.data_ptr<float>(), workspace_ptr, kk, ib, WORKSPACE_STRIDE
        );

        if (kk + ib < 1024) {


            int col_tiles = (1024 - kk - ib + COL_TILE - 1) / COL_TILE;
            dim3 y_grid((unsigned int)batch, (unsigned int)col_tiles);
            qr1024_compute_y<NB, COL_TILE, 2><<<y_grid, NB * COL_TILE * 2>>>(
                h.data_ptr<float>(), workspace_ptr, y_ptr, kk, ib, WORKSPACE_STRIDE
            );

            int row_tiles = (1024 - kk + ROW_TILE - 1) / ROW_TILE;
            dim3 update_grid((unsigned int)batch, (unsigned int)row_tiles, (unsigned int)col_tiles);
            qr1024_apply_block<NB, ROW_TILE, COL_TILE><<<update_grid, ROW_TILE * COL_TILE>>>(
                h.data_ptr<float>(), y_ptr, kk, ib
            );
        }
    }

    check_cuda(cudaGetLastError(), "qr1024_blocked_cuda");
}

// =============================================================================
// Large-N blocked QR kernels (n=2048, n=4096). N is a template parameter so we
// can share the same template bodies across the two shapes. The structure
// mirrors qr1024: panel factor -> build T -> compute Y -> apply block.
// =============================================================================

template <int N>
__global__ void copy_input_kernel_N(
    const float* __restrict__ a,
    float* __restrict__ h,
    float* __restrict__ tau,
    int batch,
    long long total
) {
    long long stride = (long long)blockDim.x * gridDim.x;
    for (long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
         idx < total; idx += stride) {
        h[idx] = a[idx];
    }
    for (int b = blockIdx.x * blockDim.x + threadIdx.x;
         b < batch; b += blockDim.x * gridDim.x) {
        tau[(long long)b * N + (N - 1)] = 0.0f;
    }
}

template <int N, int BLOCK_SIZE, int NUM_WARPS>
__global__ void qrN_panel_factor(
    float* __restrict__ h,
    float* __restrict__ tau,
    int kk,
    int ib
) {
    extern __shared__ float smem[];
    float* v = smem;
    float* warp_buf = smem + N;

    int b = blockIdx.x;
    int tid = threadIdx.x;
    int wid = tid / WARP_SIZE;
    int lid = tid % WARP_SIZE;

    float* __restrict__ h_b = h + (long long)b * N * N;
    float* __restrict__ tau_b = tau + (long long)b * N;

    __shared__ float tau_s;
    __shared__ float inv_s;
    __shared__ float dot_s;
    __shared__ int active_s;

    for (int local_k = 0; local_k < ib && kk + local_k < N - 1; ++local_k) {
        int k = kk + local_k;
        int m = N - k;
        int panel_end = kk + ib;

        float s = 0.0f;
        for (int i = tid; i < m; i += BLOCK_SIZE) {
            float x = h_b[(k + i) * N + k];
            s = fmaf(x, x, s);
        }

        s = warp_reduce_sum(s);
        if (lid == 0) warp_buf[wid] = s;
        __syncthreads();

        if (wid == 0) {
            s = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
            s = warp_reduce_sum(s);
            if (lid == 0) warp_buf[0] = s;
        }
        __syncthreads();

        if (tid == 0) {
            float norm = sqrtf(warp_buf[0]);
            float diag = h_b[k * N + k];
            if (norm <= 1e-20f) {
                tau_s = 0.0f;
                inv_s = 0.0f;
                active_s = 0;
                tau_b[k] = 0.0f;
            } else {
                float alpha = (diag >= 0.0f) ? -norm : norm;
                float denom = diag - alpha;
                inv_s = 1.0f / denom;
                tau_s = (alpha - diag) / alpha;
                active_s = 1;
                h_b[k * N + k] = alpha;
                tau_b[k] = tau_s;
            }
        }
        __syncthreads();

        if (active_s != 0) {
            if (tid == 0) v[0] = 1.0f;
            for (int i = tid; i < m - 1; i += BLOCK_SIZE) {
                float vi = h_b[(k + 1 + i) * N + k] * inv_s;
                v[i + 1] = vi;
                h_b[(k + 1 + i) * N + k] = vi;
            }
            __syncthreads();

            for (int j = k + 1; j < panel_end; ++j) {
                float dot_part = (tid == 0) ? h_b[k * N + j] : 0.0f;
                for (int i = k + 1 + tid; i < N; i += BLOCK_SIZE) {
                    dot_part = fmaf(h_b[i * N + k], h_b[i * N + j], dot_part);
                }

                dot_part = warp_reduce_sum(dot_part);
                if (lid == 0) warp_buf[wid] = dot_part;
                __syncthreads();

                if (wid == 0) {
                    dot_part = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
                    dot_part = warp_reduce_sum(dot_part);
                    if (lid == 0) dot_s = dot_part;
                }
                __syncthreads();

                if (tid == 0) {
                    h_b[k * N + j] -= tau_s * dot_s;
                }
                for (int i = k + 1 + tid; i < N; i += BLOCK_SIZE) {
                    h_b[i * N + j] -= tau_s * h_b[i * N + k] * dot_s;
                }
                __syncthreads();
            }
        }
    }
}

template <int N, int NB, int BLOCK_SIZE>
__global__ void qrN_build_t(
    const float* __restrict__ h,
    const float* __restrict__ tau,
    float* __restrict__ workspace,
    int kk,
    int ib,
    int workspace_stride
) {
    __shared__ float reduce_buf[BLOCK_SIZE];
    __shared__ float tmp[NB];

    int b = blockIdx.x;
    int tid = threadIdx.x;

    const float* __restrict__ h_b = h + (long long)b * N * N;
    const float* __restrict__ tau_b = tau + (long long)b * N;
    float* __restrict__ t = workspace + (long long)b * workspace_stride;

    for (int idx = tid; idx < NB * NB; idx += BLOCK_SIZE) {
        t[idx] = 0.0f;
    }
    __syncthreads();

    for (int j = 0; j < ib; ++j) {
        float tau_j = tau_b[kk + j];
        if (tau_j == 0.0f) {
            if (tid == 0) t[j * NB + j] = 0.0f;
            __syncthreads();
            continue;
        }

        for (int idx = tid; idx < NB; idx += BLOCK_SIZE) {
            tmp[idx] = 0.0f;
        }
        __syncthreads();

        for (int i = 0; i < j; ++i) {
            float dot = (tid == 0) ? h_b[(kk + j) * N + (kk + i)] : 0.0f;
            for (int r = j + 1 + tid; r < N - kk; r += BLOCK_SIZE) {
                dot = fmaf(
                    h_b[(kk + r) * N + (kk + i)],
                    h_b[(kk + r) * N + (kk + j)],
                    dot
                );
            }
            reduce_buf[tid] = dot;
            __syncthreads();
            for (int offset = BLOCK_SIZE / 2; offset > 0; offset >>= 1) {
                if (tid < offset) {
                    reduce_buf[tid] += reduce_buf[tid + offset];
                }
                __syncthreads();
            }
            if (tid == 0) {
                tmp[i] = -tau_j * reduce_buf[0];
            }
            __syncthreads();
        }

        if (tid == 0) {
            for (int row = 0; row < j; ++row) {
                float val = 0.0f;
                for (int col = row; col < j; ++col) {
                    val = fmaf(t[row * NB + col], tmp[col], val);
                }
                t[row * NB + j] = val;
            }
            t[j * NB + j] = tau_j;
        }
        __syncthreads();
    }
}

template <int N, int NB, int COL_TILE, int RSPLIT>
__global__ void qrN_compute_y(
    const float* __restrict__ h,
    const float* __restrict__ workspace,
    float* __restrict__ y_workspace,
    int kk,
    int ib,
    int workspace_stride
) {
    __shared__ float partial_s[NB * COL_TILE * RSPLIT];
    __shared__ float w_s[NB * COL_TILE];

    int b = blockIdx.x;
    int col_tile = blockIdx.y;
    int idx = threadIdx.x;
    int split = idx % RSPLIT;
    int c = (idx / RSPLIT) % COL_TILE;
    int q = idx / (RSPLIT * COL_TILE);
    int trail_col = col_tile * COL_TILE + c;
    int global_col = kk + ib + trail_col;
    int pair = q * COL_TILE + c;

    const float* __restrict__ h_b = h + (long long)b * N * N;
    const float* __restrict__ t = workspace + (long long)b * workspace_stride;
    float* __restrict__ y = y_workspace + (long long)b * NB * N;

    if (q < ib && c < COL_TILE && global_col < N) {
        float sum = (split == 0) ? h_b[(kk + q) * N + global_col] : 0.0f;
        for (int r = q + 1 + split; r < N - kk; r += RSPLIT) {
            sum = fmaf(
                h_b[(kk + r) * N + (kk + q)],
                h_b[(kk + r) * N + global_col],
                sum
            );
        }
        partial_s[pair * RSPLIT + split] = sum;
    } else if (q < NB && c < COL_TILE) {
        partial_s[pair * RSPLIT + split] = 0.0f;
    }
    __syncthreads();

    if (split == 0 && q < ib && c < COL_TILE && global_col < N) {
        float w = 0.0f;
        #pragma unroll
        for (int s = 0; s < RSPLIT; ++s) {
            w += partial_s[pair * RSPLIT + s];
        }
        w_s[pair] = w;
    }
    __syncthreads();

    if (split == 0 && q < ib && c < COL_TILE && global_col < N) {
        float sum = 0.0f;
        for (int l = 0; l <= q; ++l) {
            sum = fmaf(t[l * NB + q], w_s[l * COL_TILE + c], sum);
        }
        y[q * N + global_col] = sum;
    }
}

template <int N, int NB, int ROW_TILE, int COL_TILE>
__global__ void qrN_apply_block(
    float* __restrict__ h,
    const float* __restrict__ y_workspace,
    int kk,
    int ib
) {
    __shared__ float V_smem[ROW_TILE * NB];

    int b = blockIdx.x;
    int row_tile = blockIdx.y;
    int col_tile = blockIdx.z;
    int idx = threadIdx.x;

    if (idx < ROW_TILE * NB) {
        int r = idx / NB;
        int q = idx % NB;
        int global_r = kk + row_tile * ROW_TILE + r;
        if (global_r < N) {
            float* __restrict__ h_b = h + (long long)b * N * N;
            V_smem[idx] = h_b[global_r * N + (kk + q)];
        } else {
            V_smem[idx] = 0.0f;
        }
    }
    __syncthreads();

    int r_local = row_tile * ROW_TILE + idx / COL_TILE;
    int c_local = col_tile * COL_TILE + idx % COL_TILE;
    int global_col = kk + ib + c_local;

    if (r_local >= N - kk || global_col >= N) return;

    float* __restrict__ h_b = h + (long long)b * N * N;
    const float* __restrict__ y = y_workspace + (long long)b * NB * N;

    int local_r = idx / COL_TILE;
    int nq = (r_local < ib) ? r_local : ib;
    float correction = 0.0f;
    #pragma unroll
    for (int q = 0; q < nq; ++q) {
        float vq = V_smem[local_r * NB + q];
        correction = fmaf(vq, y[q * N + global_col], correction);
    }
    if (r_local < ib) {
        correction = fmaf(1.0f, y[r_local * N + global_col], correction);
    }

    h_b[(kk + r_local) * N + global_col] -= correction;
}

template <int N>
void qrN_blocked_cuda_dispatch(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace) {
    constexpr int NB = 8;
    constexpr int COL_TILE = 32;
    constexpr int ROW_TILE = 16;
    constexpr int WORKSPACE_STRIDE = NB * NB + NB * N;
    int batch = (int)a.size(0);
    long long total = (long long)batch * N * N;
    int copy_blocks = (int)((total + THREADS - 1) / THREADS);
    if (copy_blocks < 1) copy_blocks = 1;
    if (copy_blocks > 8192) copy_blocks = 8192;

    copy_input_kernel_N<N><<<copy_blocks, THREADS>>>(
        a.data_ptr<float>(), h.data_ptr<float>(), tau.data_ptr<float>(), batch, total
    );

    float* workspace_ptr = workspace.data_ptr<float>();
    float* y_ptr = workspace_ptr + (long long)batch * WORKSPACE_STRIDE;

    for (int kk = 0; kk < N - 1; kk += NB) {
        int ib = min(NB, N - kk);
        qrN_panel_factor<N, 256, 8><<<batch, 256, (N + 8) * (int)sizeof(float)>>>(
            h.data_ptr<float>(), tau.data_ptr<float>(), kk, ib
        );

        if (kk + ib < N) {
            qrN_build_t<N, NB, 128><<<batch, 128>>>(
                h.data_ptr<float>(), tau.data_ptr<float>(), workspace_ptr, kk, ib, WORKSPACE_STRIDE
            );

            int col_tiles = (N - kk - ib + COL_TILE - 1) / COL_TILE;
            dim3 y_grid((unsigned int)batch, (unsigned int)col_tiles);
            qrN_compute_y<N, NB, COL_TILE, 2><<<y_grid, NB * COL_TILE * 2>>>(
                h.data_ptr<float>(), workspace_ptr, y_ptr, kk, ib, WORKSPACE_STRIDE
            );

            int row_tiles = (N - kk + ROW_TILE - 1) / ROW_TILE;
            dim3 update_grid((unsigned int)batch, (unsigned int)row_tiles, (unsigned int)col_tiles);
            qrN_apply_block<N, NB, ROW_TILE, COL_TILE><<<update_grid, ROW_TILE * COL_TILE>>>(
                h.data_ptr<float>(), y_ptr, kk, ib
            );
        }
    }

    check_cuda(cudaGetLastError(), "qrN_blocked_cuda_dispatch");
}

void qr2048_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace) {
    qrN_blocked_cuda_dispatch<2048>(a, h, tau, workspace);
}

void qr4096_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace) {
    qrN_blocked_cuda_dispatch<4096>(a, h, tau, workspace);
}

__global__ void copy_input_kernel_352(
    const float* __restrict__ a,
    float* __restrict__ h,
    float* __restrict__ tau,
    int batch,
    long long total
) {
    long long stride = (long long)blockDim.x * gridDim.x;
    for (long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
         idx < total; idx += stride) {
        h[idx] = a[idx];
    }
    for (int b = blockIdx.x * blockDim.x + threadIdx.x;
         b < batch; b += blockDim.x * gridDim.x) {
        tau[(long long)b * 352 + 351] = 0.0f;
    }
}

template <int BLOCK_SIZE, int NUM_WARPS>
__global__ void qr352_panel_build_t_fused(
    float* __restrict__ h,
    float* __restrict__ tau,
    float* __restrict__ workspace,
    int kk,
    int ib,
    int workspace_stride
) {
    constexpr int NB = 8;
    extern __shared__ float smem[];
    float* v = smem;
    float* warp_buf = smem + 352;

    int b = blockIdx.x;
    int tid = threadIdx.x;
    int wid = tid / WARP_SIZE;
    int lid = tid % WARP_SIZE;

    float* __restrict__ h_b = h + (long long)b * 352 * 352;
    float* __restrict__ tau_b = tau + (long long)b * 352;
    float* __restrict__ t = workspace + (long long)b * workspace_stride;

    __shared__ float tau_s;
    __shared__ float inv_s;
    __shared__ float dot_s;
    __shared__ int active_s;

    for (int local_k = 0; local_k < ib && kk + local_k < 351; ++local_k) {
        int k = kk + local_k;
        int m = 352 - k;
        int panel_end = kk + ib;

        float s = 0.0f;
        for (int i = tid; i < m; i += BLOCK_SIZE) {
            float x = h_b[(k + i) * 352 + k];
            s = fmaf(x, x, s);
        }

        s = warp_reduce_sum(s);
        if (lid == 0) warp_buf[wid] = s;
        __syncthreads();

        if (wid == 0) {
            s = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
            s = warp_reduce_sum(s);
            if (lid == 0) warp_buf[0] = s;
        }
        __syncthreads();

        if (tid == 0) {
            float norm = sqrtf(warp_buf[0]);
            float diag = h_b[k * 352 + k];
            if (norm <= 1e-20f) {
                tau_s = 0.0f;
                inv_s = 0.0f;
                active_s = 0;
                tau_b[k] = 0.0f;
            } else {
                float alpha = (diag >= 0.0f) ? -norm : norm;
                float denom = diag - alpha;
                inv_s = 1.0f / denom;
                tau_s = (alpha - diag) / alpha;
                active_s = 1;
                h_b[k * 352 + k] = alpha;
                tau_b[k] = tau_s;
            }
        }
        __syncthreads();

        if (active_s != 0) {
            if (tid == 0) v[0] = 1.0f;
            for (int i = tid; i < m - 1; i += BLOCK_SIZE) {
                float vi = h_b[(k + 1 + i) * 352 + k] * inv_s;
                v[i + 1] = vi;
                h_b[(k + 1 + i) * 352 + k] = vi;
            }
            __syncthreads();

            for (int j = k + 1; j < panel_end; ++j) {
                float dot_part = (tid == 0) ? h_b[k * 352 + j] : 0.0f;
                for (int i = k + 1 + tid; i < 352; i += BLOCK_SIZE) {
                    dot_part = fmaf(h_b[i * 352 + k], h_b[i * 352 + j], dot_part);
                }

                dot_part = warp_reduce_sum(dot_part);
                if (lid == 0) warp_buf[wid] = dot_part;
                __syncthreads();

                if (wid == 0) {
                    dot_part = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
                    dot_part = warp_reduce_sum(dot_part);
                    if (lid == 0) dot_s = dot_part;
                }
                __syncthreads();

                if (tid == 0) {
                    h_b[k * 352 + j] -= tau_s * dot_s;
                }
                for (int i = k + 1 + tid; i < 352; i += BLOCK_SIZE) {
                    h_b[i * 352 + j] -= tau_s * h_b[i * 352 + k] * dot_s;
                }
                __syncthreads();
            }
        }
        __syncthreads();
    }

    if (kk + ib >= 352) return;

    float* tmp = smem;

    for (int idx = tid; idx < NB * NB; idx += BLOCK_SIZE) {
        t[idx] = 0.0f;
    }
    __syncthreads();

    for (int j = 0; j < ib; ++j) {
        float tau_j = tau_b[kk + j];
        if (tau_j == 0.0f) {
            if (tid == 0) t[j * NB + j] = 0.0f;
            __syncthreads();
            continue;
        }

        for (int idx = tid; idx < NB; idx += BLOCK_SIZE) {
            tmp[idx] = 0.0f;
        }
        __syncthreads();

        for (int i = 0; i < j; ++i) {
            float dot_part = (tid == 0) ? h_b[(kk + j) * 352 + (kk + i)] : 0.0f;
            for (int r = j + 1 + tid; r < 352 - kk; r += BLOCK_SIZE) {
                dot_part = fmaf(
                    h_b[(kk + r) * 352 + (kk + i)],
                    h_b[(kk + r) * 352 + (kk + j)],
                    dot_part
                );
            }

            dot_part = warp_reduce_sum(dot_part);
            if (lid == 0) warp_buf[wid] = dot_part;
            __syncthreads();

            if (wid == 0) {
                dot_part = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
                dot_part = warp_reduce_sum(dot_part);
                if (lid == 0) warp_buf[0] = dot_part;
            }
            __syncthreads();

            if (tid == 0) {
                tmp[i] = -tau_j * warp_buf[0];
            }
            __syncthreads();
        }

        if (tid == 0) {
            for (int row = 0; row < j; ++row) {
                float val = 0.0f;
                for (int col = row; col < j; ++col) {
                    val = fmaf(t[row * NB + col], tmp[col], val);
                }
                t[row * NB + j] = val;
            }
            t[j * NB + j] = tau_j;
        }
        __syncthreads();
    }}

template <int NB, int COL_TILE, int RSPLIT>
__global__ void qr352_compute_y(
    const float* __restrict__ h,
    const float* __restrict__ workspace,
    float* __restrict__ y_workspace,
    int kk,
    int ib,
    int workspace_stride
) {
    __shared__ float partial_s[NB * COL_TILE * RSPLIT];
    __shared__ float w_s[NB * COL_TILE];

    int b = blockIdx.x;
    int col_tile = blockIdx.y;
    int idx = threadIdx.x;
    int split = idx % RSPLIT;
    int c = (idx / RSPLIT) % COL_TILE;
    int q = idx / (RSPLIT * COL_TILE);
    int trail_col = col_tile * COL_TILE + c;
    int global_col = kk + ib + trail_col;
    int pair = q * COL_TILE + c;

    const float* __restrict__ h_b = h + (long long)b * 352 * 352;
    const float* __restrict__ t = workspace + (long long)b * workspace_stride;
    float* __restrict__ y = y_workspace + (long long)b * NB * 352;

    if (q < ib && c < COL_TILE && global_col < 352) {
        float sum = (split == 0) ? h_b[(kk + q) * 352 + global_col] : 0.0f;
        for (int r = q + 1 + split; r < 352 - kk; r += RSPLIT) {
            sum = fmaf(
                h_b[(kk + r) * 352 + (kk + q)],
                h_b[(kk + r) * 352 + global_col],
                sum
            );
        }
        partial_s[pair * RSPLIT + split] = sum;
    } else if (q < NB && c < COL_TILE) {
        partial_s[pair * RSPLIT + split] = 0.0f;
    }
    __syncthreads();

    if (split == 0 && q < ib && c < COL_TILE && global_col < 352) {
        float w = 0.0f;
        #pragma unroll
        for (int s = 0; s < RSPLIT; ++s) {
            w += partial_s[pair * RSPLIT + s];
        }
        w_s[pair] = w;
    }
    __syncthreads();

    if (split == 0 && q < ib && c < COL_TILE && global_col < 352) {
        float sum = 0.0f;
        for (int l = 0; l <= q; ++l) {
            sum = fmaf(t[l * NB + q], w_s[l * COL_TILE + c], sum);
        }
        y[q * 352 + global_col] = sum;
    }
}

template <int NB, int ROW_TILE, int COL_TILE>
__global__ void qr352_apply_block(
    float* __restrict__ h,
    const float* __restrict__ y_workspace,
    int kk,
    int ib
) {
    int b = blockIdx.x;
    int row_tile = blockIdx.y;
    int col_tile = blockIdx.z;
    int idx = threadIdx.x;
    int r_local = row_tile * ROW_TILE + idx / COL_TILE;
    int c_local = col_tile * COL_TILE + idx % COL_TILE;
    int global_col = kk + ib + c_local;

    if (r_local >= 352 - kk || global_col >= 352) return;

    float* __restrict__ h_b = h + (long long)b * 352 * 352;
    const float* __restrict__ y = y_workspace + (long long)b * NB * 352;

    float correction = 0.0f;
    int nq = (r_local < ib) ? r_local : ib;
    #pragma unroll
    for (int q = 0; q < nq; ++q) {
        float vq = h_b[(kk + r_local) * 352 + (kk + q)];
        correction = fmaf(vq, y[q * 352 + global_col], correction);
    }
    if (r_local < ib) {
        correction = fmaf(1.0f, y[r_local * 352 + global_col], correction);
    }

    h_b[(kk + r_local) * 352 + global_col] -= correction;
}

void qr352_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace) {
    constexpr int NB = 8;
    constexpr int COL_TILE = 32;
    constexpr int ROW_TILE = 16;
    constexpr int WORKSPACE_STRIDE = NB * NB + NB * 352;
    int batch = (int)a.size(0);
    long long total = (long long)batch * 352 * 352;
    int copy_blocks = (int)((total + THREADS - 1) / THREADS);
    if (copy_blocks < 1) copy_blocks = 1;
    if (copy_blocks > 8192) copy_blocks = 8192;

    copy_input_kernel_352<<<copy_blocks, THREADS>>>(
        a.data_ptr<float>(), h.data_ptr<float>(), tau.data_ptr<float>(), batch, total
    );

    float* workspace_ptr = workspace.data_ptr<float>();
    float* y_ptr = workspace_ptr + (long long)batch * WORKSPACE_STRIDE;

    for (int kk = 0; kk < 351; kk += NB) {
        int ib = min(NB, 352 - kk);
        qr352_panel_build_t_fused<256, 8><<<batch, 256, (352 + 8) * (int)sizeof(float)>>>(
            h.data_ptr<float>(), tau.data_ptr<float>(), workspace_ptr, kk, ib, WORKSPACE_STRIDE
        );

        if (kk + ib < 352) {


            int col_tiles = (352 - kk - ib + COL_TILE - 1) / COL_TILE;
            dim3 y_grid((unsigned int)batch, (unsigned int)col_tiles);
            qr352_compute_y<NB, COL_TILE, 1><<<y_grid, NB * COL_TILE * 1>>>(
                h.data_ptr<float>(), workspace_ptr, y_ptr, kk, ib, WORKSPACE_STRIDE
            );

            int row_tiles = (352 - kk + ROW_TILE - 1) / ROW_TILE;
            dim3 update_grid((unsigned int)batch, (unsigned int)row_tiles, (unsigned int)col_tiles);
            qr352_apply_block<NB, ROW_TILE, COL_TILE><<<update_grid, ROW_TILE * COL_TILE>>>(
                h.data_ptr<float>(), y_ptr, kk, ib
            );
        }
    }

    check_cuda(cudaGetLastError(), "qr352_blocked_cuda");
}

__global__ void copy_input_kernel_176(
    const float* __restrict__ a,
    float* __restrict__ h,
    float* __restrict__ tau,
    int batch,
    long long total
) {
    long long stride = (long long)blockDim.x * gridDim.x;
    for (long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
         idx < total; idx += stride) {
        h[idx] = a[idx];
    }
    for (int b = blockIdx.x * blockDim.x + threadIdx.x;
         b < batch; b += blockDim.x * gridDim.x) {
        tau[(long long)b * 176 + 175] = 0.0f;
    }
}

template <int BLOCK_SIZE, int NUM_WARPS>
__global__ void qr176_panel_build_t_fused(
    float* __restrict__ h,
    float* __restrict__ tau,
    float* __restrict__ workspace,
    int kk,
    int ib,
    int workspace_stride
) {
    constexpr int NB = 8;
    extern __shared__ float smem[];
    float* v = smem;
    float* warp_buf = smem + 176;
    int b = blockIdx.x;
    int tid = threadIdx.x;
    int wid = tid / WARP_SIZE;
    int lid = tid % WARP_SIZE;
    float* __restrict__ h_b = h + (long long)b * 176 * 176;
    float* __restrict__ tau_b = tau + (long long)b * 176;
    float* __restrict__ t = workspace + (long long)b * workspace_stride;
    __shared__ float tau_s;
    __shared__ float inv_s;
    __shared__ float dot_s;
    __shared__ int active_s;

    for (int local_k = 0; local_k < ib && kk + local_k < 175; ++local_k) {
        int k = kk + local_k;
        int m = 176 - k;
        int panel_end = kk + ib;
        float s = 0.0f;
        for (int i = tid; i < m; i += BLOCK_SIZE) {
            float x = h_b[(k + i) * 176 + k];
            s = fmaf(x, x, s);
        }
        s = warp_reduce_sum(s);
        if (lid == 0) warp_buf[wid] = s;
        __syncthreads();
        if (wid == 0) {
            s = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
            s = warp_reduce_sum(s);
            if (lid == 0) warp_buf[0] = s;
        }
        __syncthreads();
        if (tid == 0) {
            float norm = sqrtf(warp_buf[0]);
            float diag = h_b[k * 176 + k];
            if (norm <= 1e-20f) {
                tau_s = 0.0f; inv_s = 0.0f; active_s = 0;
                tau_b[k] = 0.0f;
            } else {
                float alpha = (diag >= 0.0f) ? -norm : norm;
                float denom = diag - alpha;
                inv_s = 1.0f / denom;
                tau_s = (alpha - diag) / alpha;
                active_s = 1;
                h_b[k * 176 + k] = alpha;
                tau_b[k] = tau_s;
            }
        }
        __syncthreads();
        if (active_s != 0) {
            if (tid == 0) v[0] = 1.0f;
            for (int i = tid; i < m - 1; i += BLOCK_SIZE) {
                float vi = h_b[(k + 1 + i) * 176 + k] * inv_s;
                v[i + 1] = vi;
                h_b[(k + 1 + i) * 176 + k] = vi;
            }
            __syncthreads();
            for (int j = k + 1; j < panel_end; ++j) {
                float dot_part = (tid == 0) ? h_b[k * 176 + j] : 0.0f;
                for (int i = k + 1 + tid; i < 176; i += BLOCK_SIZE) {
                    dot_part = fmaf(h_b[i * 176 + k], h_b[i * 176 + j], dot_part);
                }
                dot_part = warp_reduce_sum(dot_part);
                if (lid == 0) warp_buf[wid] = dot_part;
                __syncthreads();
                if (wid == 0) {
                    dot_part = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
                    dot_part = warp_reduce_sum(dot_part);
                    if (lid == 0) dot_s = dot_part;
                }
                __syncthreads();
                if (tid == 0) {
                    h_b[k * 176 + j] -= tau_s * dot_s;
                }
                for (int i = k + 1 + tid; i < 176; i += BLOCK_SIZE) {
                    h_b[i * 176 + j] -= tau_s * h_b[i * 176 + k] * dot_s;
                }
                __syncthreads();
            }
        }
        __syncthreads();
    }

    if (kk + ib >= 176) return;

    float* tmp = smem;

    for (int idx = tid; idx < NB * NB; idx += BLOCK_SIZE) {
        t[idx] = 0.0f;
    }
    __syncthreads();

    for (int j = 0; j < ib; ++j) {
        float tau_j = tau_b[kk + j];
        if (tau_j == 0.0f) {
            if (tid == 0) t[j * NB + j] = 0.0f;
            __syncthreads();
            continue;
        }

        for (int idx = tid; idx < NB; idx += BLOCK_SIZE) {
            tmp[idx] = 0.0f;
        }
        __syncthreads();

        for (int i = 0; i < j; ++i) {
            float dot_part = (tid == 0) ? h_b[(kk + j) * 176 + (kk + i)] : 0.0f;
            for (int r = j + 1 + tid; r < 176 - kk; r += BLOCK_SIZE) {
                dot_part = fmaf(
                    h_b[(kk + r) * 176 + (kk + i)],
                    h_b[(kk + r) * 176 + (kk + j)],
                    dot_part
                );
            }

            dot_part = warp_reduce_sum(dot_part);
            if (lid == 0) warp_buf[wid] = dot_part;
            __syncthreads();

            if (wid == 0) {
                dot_part = (lid < NUM_WARPS) ? warp_buf[lid] : 0.0f;
                dot_part = warp_reduce_sum(dot_part);
                if (lid == 0) warp_buf[0] = dot_part;
            }
            __syncthreads();

            if (tid == 0) {
                tmp[i] = -tau_j * warp_buf[0];
            }
            __syncthreads();
        }

        if (tid == 0) {
            for (int row = 0; row < j; ++row) {
                float val = 0.0f;
                for (int col = row; col < j; ++col) {
                    val = fmaf(t[row * NB + col], tmp[col], val);
                }
                t[row * NB + j] = val;
            }
            t[j * NB + j] = tau_j;
        }
        __syncthreads();
    }}

template <int NB, int COL_TILE, int RSPLIT>
__global__ void qr176_compute_y(
    const float* __restrict__ h,
    const float* __restrict__ workspace,
    float* __restrict__ y_workspace,
    int kk,
    int ib,
    int workspace_stride
) {
    __shared__ float partial_s[NB * COL_TILE * RSPLIT];
    __shared__ float w_s[NB * COL_TILE];
    int b = blockIdx.x;
    int col_tile = blockIdx.y;
    int idx = threadIdx.x;
    int split = idx % RSPLIT;
    int c = (idx / RSPLIT) % COL_TILE;
    int q = idx / (RSPLIT * COL_TILE);
    int trail_col = col_tile * COL_TILE + c;
    int global_col = kk + ib + trail_col;
    int pair = q * COL_TILE + c;
    const float* __restrict__ h_b = h + (long long)b * 176 * 176;
    const float* __restrict__ t = workspace + (long long)b * workspace_stride;
    float* __restrict__ y = y_workspace + (long long)b * NB * 176;

    if (q < ib && c < COL_TILE && global_col < 176) {
        float sum = (split == 0) ? h_b[(kk + q) * 176 + global_col] : 0.0f;
        for (int r = q + 1 + split; r < 176 - kk; r += RSPLIT) {
            sum = fmaf(h_b[(kk + r) * 176 + (kk + q)], h_b[(kk + r) * 176 + global_col], sum);
        }
        partial_s[pair * RSPLIT + split] = sum;
    } else if (q < NB && c < COL_TILE) {
        partial_s[pair * RSPLIT + split] = 0.0f;
    }
    __syncthreads();

    if (split == 0 && q < ib && c < COL_TILE && global_col < 176) {
        float w = 0.0f;
        #pragma unroll
        for (int s = 0; s < RSPLIT; ++s) w += partial_s[pair * RSPLIT + s];
        w_s[pair] = w;
    }
    __syncthreads();

    if (split == 0 && q < ib && c < COL_TILE && global_col < 176) {
        float sum = 0.0f;
        for (int l = 0; l <= q; ++l) sum = fmaf(t[l * NB + q], w_s[l * COL_TILE + c], sum);
        y[q * 176 + global_col] = sum;
    }
}

template <int NB, int ROW_TILE, int COL_TILE>
__global__ void qr176_apply_block(
    float* __restrict__ h,
    const float* __restrict__ y_workspace,
    int kk,
    int ib
) {
    int b = blockIdx.x;
    int row_tile = blockIdx.y;
    int col_tile = blockIdx.z;
    int idx = threadIdx.x;
    int r_local = row_tile * ROW_TILE + idx / COL_TILE;
    int c_local = col_tile * COL_TILE + idx % COL_TILE;
    int global_col = kk + ib + c_local;
    if (r_local >= 176 - kk || global_col >= 176) return;
    float* __restrict__ h_b = h + (long long)b * 176 * 176;
    const float* __restrict__ y = y_workspace + (long long)b * NB * 176;
    float correction = 0.0f;
    int nq = (r_local < ib) ? r_local : ib;
    #pragma unroll
    for (int q = 0; q < nq; ++q) {
        float vq = h_b[(kk + r_local) * 176 + (kk + q)];
        correction = fmaf(vq, y[q * 176 + global_col], correction);
    }
    if (r_local < ib) {
        correction = fmaf(1.0f, y[r_local * 176 + global_col], correction);
    }
    h_b[(kk + r_local) * 176 + global_col] -= correction;
}

void qr176_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace) {
    constexpr int NB = 8;
    constexpr int COL_TILE = 32;
    constexpr int ROW_TILE = 16;
    constexpr int WORKSPACE_STRIDE = NB * NB + NB * 176;
    int batch = (int)a.size(0);
    long long total = (long long)batch * 176 * 176;
    int copy_blocks = (int)((total + THREADS - 1) / THREADS);
    if (copy_blocks < 1) copy_blocks = 1;
    if (copy_blocks > 8192) copy_blocks = 8192;

    copy_input_kernel_176<<<copy_blocks, THREADS>>>(
        a.data_ptr<float>(), h.data_ptr<float>(), tau.data_ptr<float>(), batch, total
    );

    float* workspace_ptr = workspace.data_ptr<float>();
    float* y_ptr = workspace_ptr + (long long)batch * WORKSPACE_STRIDE;

    for (int kk = 0; kk < 175; kk += NB) {
        int ib = min(NB, 176 - kk);
        qr176_panel_build_t_fused<256, 8><<<batch, 256, (176 + 8) * (int)sizeof(float)>>>(
            h.data_ptr<float>(), tau.data_ptr<float>(), workspace_ptr, kk, ib, WORKSPACE_STRIDE
        );
        if (kk + ib < 176) {

            int col_tiles = (176 - kk - ib + COL_TILE - 1) / COL_TILE;
            dim3 y_grid((unsigned int)batch, (unsigned int)col_tiles);
            qr176_compute_y<NB, COL_TILE, 2><<<y_grid, NB * COL_TILE * 2>>>(
                h.data_ptr<float>(), workspace_ptr, y_ptr, kk, ib, WORKSPACE_STRIDE
            );
            int row_tiles = (176 - kk + ROW_TILE - 1) / ROW_TILE;
            dim3 update_grid((unsigned int)batch, (unsigned int)row_tiles, (unsigned int)col_tiles);
            qr176_apply_block<NB, ROW_TILE, COL_TILE><<<update_grid, ROW_TILE * COL_TILE>>>(
                h.data_ptr<float>(), y_ptr, kk, ib
            );
        }
    }
    check_cuda(cudaGetLastError(), "qr176_blocked_cuda");
}
"""

CPP_SRC = r"""
void qr32_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau);
void qr512_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace);
void qr512_panel_build_t_cuda(torch::Tensor h, torch::Tensor tau, torch::Tensor workspace, int kk, int batch);
void qr512_apply_block_wmma_cuda(torch::Tensor h, torch::Tensor y, int kk, int ib, int batch);
void qr1024_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace);
void qr352_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace);
void qr176_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace);
void qr2048_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace);
void qr4096_blocked_cuda(torch::Tensor a, torch::Tensor h, torch::Tensor tau, torch::Tensor workspace);
"""

try:
    module = load_inline(
        name="qr_native_v2",
        cpp_sources=[CPP_SRC],
        cuda_sources=[CUDA_SRC],
        functions=[
            "qr32_cuda",
            "qr512_blocked_cuda",
            "qr512_panel_build_t_cuda",
            "qr512_apply_block_wmma_cuda",
            "qr1024_blocked_cuda",
            "qr352_blocked_cuda",
            "qr176_blocked_cuda",
            "qr2048_blocked_cuda",
            "qr4096_blocked_cuda",
        ],
        verbose=False,
        extra_cuda_cflags=["-O2"],
    )
    _has_cuda = True
except Exception:
    module = None
    _has_cuda = False


_cute512_enabled = False
_cute512_dyn_y_executors: dict = {}
_cute512_dyn_top_executors: dict = {}
_cute512_dyn_bottom4_executors: dict = {}



_ws_cache: dict = {}


def _get_ws(shape, dtype, device):
    import math
    dev = torch.device(device) if not isinstance(device, torch.device) else device
    key = (tuple(shape), dev.index if dev.type == "cuda" else -1)
    t = _ws_cache.get(key)
    needed = math.prod(int(s) for s in shape)
    if t is None or t.numel() < needed:
        t = torch.empty(shape, dtype=dtype, device=dev)
        _ws_cache[key] = t
    return t


def _try_cuda_graph_path(data: input_t) -> output_t | None:
    return None


def _custom_kernel_cute512(data: input_t) -> output_t:
    batch, n, _ = data.shape
    nb = 8
    ws = nb * nb + nb * n
    h = torch.empty_like(data)
    tau = torch.empty((batch, n), dtype=torch.float32, device=data.device)
    workspace = _get_ws((batch, 2 * ws), torch.float32, data.device)
    w_tmp = _get_ws((batch, nb * n), torch.float32, data.device).view(batch, nb, n)
    y_data = workspace.view(-1)[batch * ws:batch * ws + batch * nb * n].view(batch, nb, n)

    h.copy_(data)
    tau[:, n - 1] = 0.0
    for kk in range(0, n - 1, nb):
        ib = min(nb, n - kk)
        module.qr512_panel_build_t_cuda(h, tau, workspace, kk, batch)
        if kk + ib >= n:
            break

        _cute512_dynamic_compute_y(h, workspace, y_data, w_tmp, kk)
        _cute512_dynamic_apply_top(h, y_data, kk)

        if n - kk - ib > 0:
            _cute512_dynamic_apply_bottom4(h, y_data, kk)
    return h, tau


def _custom_kernel_wmma512(data: input_t) -> output_t:
    batch, n, _ = data.shape
    nb = 8
    ws = nb * nb + nb * n
    h = torch.empty_like(data)
    tau = torch.empty((batch, n), dtype=torch.float32, device=data.device)
    workspace = _get_ws((batch, 2 * ws), torch.float32, data.device)
    w_tmp = _get_ws((batch, nb * n), torch.float32, data.device).view(batch, nb, n)
    y_data = workspace.view(-1)[batch * ws:batch * ws + batch * nb * n].view(batch, nb, n)

    h.copy_(data)
    tau[:, n - 1] = 0.0
    for kk in range(0, n - 1, nb):
        ib = min(nb, n - kk)
        module.qr512_panel_build_t_cuda(h, tau, workspace, kk, batch)
        if kk + ib >= n:
            break

        _cute512_dynamic_compute_y(h, workspace, y_data, w_tmp, kk)
        _cute512_dynamic_apply_top(h, y_data, kk)

        if n - kk - ib > 0:
            module.qr512_apply_block_wmma_cuda(h, y_data, kk, ib, batch)
    return h, tau


def custom_kernel(data: input_t) -> output_t:
    if not _has_cuda:
        return torch.geqrf(data)

    batch, n, _ = data.shape
    h = torch.empty_like(data)
    tau = torch.empty((batch, n), dtype=torch.float32, device=data.device)

    if n == 32:
        module.qr32_cuda(data, h, tau)
        return h, tau

    if n == 512:
        workspace = _get_ws((batch, 2 * (8 * 8 + 8 * n)), torch.float32, data.device)
        module.qr512_blocked_cuda(data, h, tau, workspace)
        return h, tau

    if n == 1024:
        workspace = _get_ws((batch, 2 * (8 * 8 + 8 * n)), torch.float32, data.device)
        module.qr1024_blocked_cuda(data, h, tau, workspace)
        return h, tau

    if n == 352:
        workspace = _get_ws((batch, 2 * (8 * 8 + 8 * n)), torch.float32, data.device)
        module.qr352_blocked_cuda(data, h, tau, workspace)
        return h, tau

    if n == 176:
        workspace = _get_ws((batch, 2 * (8 * 8 + 8 * n)), torch.float32, data.device)
        module.qr176_blocked_cuda(data, h, tau, workspace)
        return h, tau

    if n in (2048, 4096):
        return torch.geqrf(data)

    return torch.geqrf(data)
scrolls · 2262 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