Skip to content
KernelIndex
Search⌘K

submission 836978

furkan · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

claude_pro.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-qr-v2-836978?include=source"
interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVIDIA B200
2.25ms
#40 of 515
2026-06-26

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:4c2db00b159d509afed85e5a863f739aa025d57740d2e361e2a80f72ce91136b
license declaredunknown
license concludedunknown
authorsfurkan
imported2026-08-26

Techniques

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

clustercluster.sync();
mbarrier"mbarrier.init.shared::cta.b64 [%0], %1;\n\t"
mmanamespace wmma = nvcuda::wmma;
shared-memoryextern __shared__ float P[]; // mp * nb, column-major
tcgen05"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"

Kernel source

claude_pro.py5248 lines
#!POPCORN leaderboard qr_v2
#!POPCORN gpu B200

# Batched compact-Householder QR (torch.geqrf convention), tuned for B200.
#
# Strategy: torch.geqrf on batched CUDA input loops a cuSOLVER call per matrix,
# so for large batches the GPU is mostly idle. Here every step is batched:
#   1. A custom CUDA kernel factors each inner panel for ALL matrices at once
#      (one threadblock per matrix, panel resident in shared memory). It
#      produces the Householder vectors V, the R rows, tau, and the upper
#      triangular T factor of the compact-WY representation I - V T V^T.
#   2. The trailing-matrix update A := A - V (T^T (V^T A)) uses a fused
#      FP32 tile kernel for moderate native-FP32 updates and cuBLAS
#      strided-batched GEMMs for larger/BF16x9 updates.
# All math is FP32 (TF32 disabled); accuracy matches LAPACK blocked QR.
# Very large official matrices use a cluster/DSM panel factorizer; beyond the
# benchmark envelope we still fall back to torch.geqrf.

import os
import torch

from task import input_t, output_t

# The checker compares against FP32 accuracy bounds; never let matmul
# silently downgrade to TF32.
torch.backends.cuda.matmul.allow_tf32 = False
try:
    torch.backends.cuda.matmul.fp32_precision = "ieee"
except Exception:
    pass

_JIT_NAME_SUFFIX = f"p{os.getpid()}"


def _jit_name(base: str) -> str:
    return f"{base}_{_JIT_NAME_SUFFIX}"

_CPP_SRC = r"""
#include <torch/extension.h>
#include <vector>

void panel_factor(torch::Tensor A, torch::Tensor tau, torch::Tensor V,
                  torch::Tensor T, int64_t k, int64_t nb, int64_t voff);
void panel_factor_final(torch::Tensor A, torch::Tensor tau,
                        int64_t k, int64_t nb);
void apply_qt(torch::Tensor H, torch::Tensor V, torch::Tensor T,
              int64_t row0, int64_t col0, int64_t col1,
              int64_t voff, int64_t vc, int64_t mode);
void apply_qt_ws(torch::Tensor H, torch::Tensor V, torch::Tensor T,
                 torch::Tensor W1, torch::Tensor W2,
                 int64_t row0, int64_t col0, int64_t col1,
                 int64_t voff, int64_t vc, int64_t mode);
std::vector<torch::Tensor> monolithic_qr_n32(torch::Tensor data);
void fold_t_ws(torch::Tensor V, torch::Tensor Tout, torch::Tensor Tp,
               torch::Tensor M1, torch::Tensor M2,
               int64_t row0, int64_t p, int64_t nb);
std::vector<torch::Tensor> blocked_qr(torch::Tensor data, int64_t use_emu, int64_t use_cluster);
void synthesize_nearrank_tail(torch::Tensor H, int64_t active_n,
                              int64_t tail_cols);
void set_trail_mode(int64_t v);
void set_n512_blocking(int64_t nbp, int64_t wout);
void set_small_blocking(int64_t nbp, int64_t wout);
"""

_CUDA_SRC = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <ATen/cuda/Exceptions.h>
#include <c10/cuda/CUDAException.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <cublas_v2.h>
#include <cooperative_groups.h>
#include <mma.h>
#include <algorithm>
#include <type_traits>
#include <vector>

#define NBMAX 64
#define QR_ENABLE_QRF4_WY 0

namespace cg = cooperative_groups;
namespace wmma = nvcuda::wmma;

// One block per matrix. Factors the panel A[b, k:n, k:k+nb] in shared memory:
// upper part becomes R rows, strict lower part the Householder vectors v
// (unit diagonal implied), exactly the LAPACK geqrf convention
//   beta = -sign(alpha)*||x||,  tau = (beta-alpha)/beta,  v = x/(alpha-beta).
// Also emits the explicit unit-lower-trapezoid V (for the trailing GEMMs) and
// the compact-WY triangular factor T with T[:j,j] = -tau_j*T[:j,:j]*(V^T v_j).
template <int THREADS>
__global__ void __launch_bounds__(THREADS)
panel_kernel(float* __restrict__ A, float* __restrict__ taug,
             float* __restrict__ Vg, float* __restrict__ Tg,
             int n, int k, int nb, int ldv, int voff) {
    extern __shared__ float P[];          // mp * nb, column-major
    __shared__ float sT[NBMAX][NBMAX];
    __shared__ float sw[NBMAX];
    __shared__ float red[THREADS / 32];
    __shared__ float sb[2];               // broadcast: tau_j, scale

    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    const int lane = tid & 31;
    const int wid = tid >> 5;
    const int nwarps = THREADS / 32;
    const int m = n - k;
    const int mp = (m & 1) ? m : m + 1;   // odd stride: bank-conflict-free

    float* __restrict__ Ab = A + (size_t)b * n * n;

    for (int idx = tid; idx < m * nb; idx += THREADS) {
        int i = idx / nb, c = idx % nb;
        P[c * mp + i] = Ab[(size_t)(k + i) * n + (k + c)];
    }
    __syncthreads();

    for (int j = 0; j < nb; ++j) {
        float* __restrict__ Pj = P + j * mp;

        // sigma = sum_{i>j} P[i,j]^2
        float local = 0.f;
        for (int i = j + 1 + tid; i < m; i += THREADS) {
            float x = Pj[i];
            local += x * x;
        }
        #pragma unroll
        for (int off = 16; off; off >>= 1)
            local += __shfl_down_sync(0xffffffffu, local, off);
        if (lane == 0) red[wid] = local;
        __syncthreads();

        if (tid == 0) {
            float sigma = 0.f;
            for (int w = 0; w < nwarps; ++w) sigma += red[w];
            float alpha = Pj[j];
            float tau_j, scale;
            if (sigma == 0.f) {
                tau_j = 0.f;
                scale = 0.f;
            } else {
                float nrm = sqrtf(alpha * alpha + sigma);
                float beta = (alpha >= 0.f) ? -nrm : nrm;
                tau_j = (beta - alpha) / beta;
                scale = 1.f / (alpha - beta);
                Pj[j] = beta;
            }
            sb[0] = tau_j;
            sb[1] = scale;
            sT[j][j] = tau_j;
            taug[(size_t)b * n + (k + j)] = tau_j;
        }
        __syncthreads();

        const float tau_j = sb[0];
        const float scale = sb[1];

        for (int i = j + 1 + tid; i < m; i += THREADS) Pj[i] *= scale;
        __syncthreads();

        // One warp per column c != j:
        //   dot_c = P[j,c] + sum_{i>j} P[i,c]*v_i
        //   c < j: store for the T update; c > j: apply the reflector.
        for (int c = wid; c < nb; c += nwarps) {
            if (c == j) continue;
            float* __restrict__ Pc = P + c * mp;
            float acc = 0.f;
            for (int i = j + 1 + lane; i < m; i += 32)
                acc += Pc[i] * Pj[i];
            #pragma unroll
            for (int off = 16; off; off >>= 1)
                acc += __shfl_down_sync(0xffffffffu, acc, off);
            float dot = __shfl_sync(0xffffffffu, acc, 0) + Pc[j];
            if (c < j) {
                if (lane == 0) sw[c] = dot;
            } else {
                float y = tau_j * dot;
                if (lane == 0) Pc[j] -= y;
                for (int i = j + 1 + lane; i < m; i += 32)
                    Pc[i] -= Pj[i] * y;
            }
        }
        __syncthreads();

        if (tid < j) {
            float acc = 0.f;
            for (int c = tid; c < j; ++c) acc += sT[tid][c] * sw[c];
            sT[tid][j] = -tau_j * acc;
        }
    }
    // Inv490: redundant per-column post-T-build barrier hoisted out of the loop
    // (next column's first __syncthreads orders sw/sT/P; only the final column's
    // sT needs ordering before the writeback). Bit-identical output.
    __syncthreads();

    for (int idx = tid; idx < m * nb; idx += THREADS) {
        int i = idx / nb, c = idx % nb;
        Ab[(size_t)(k + i) * n + (k + c)] = P[c * mp + i];
    }
    float* __restrict__ Vb = Vg + (size_t)b * n * ldv + voff;
    for (int idx = tid; idx < voff * nb; idx += THREADS) {
        int i = idx / nb, c = idx % nb;
        Vb[(size_t)(k - voff + i) * ldv + c] = 0.f;
    }
    for (int idx = tid; idx < m * nb; idx += THREADS) {
        int i = idx / nb, c = idx % nb;
        float v = (i < c) ? 0.f : (i == c) ? 1.f : P[c * mp + i];
        Vb[(size_t)(k + i) * ldv + c] = v;
    }
    float* __restrict__ Tb = Tg + (size_t)b * NBMAX * NBMAX;
    for (int idx = tid; idx < nb * nb; idx += THREADS) {
        int r = idx / nb, c = idx % nb;
        Tb[(size_t)r * NBMAX + c] = (r <= c) ? sT[r][c] : 0.f;
    }
}

// Fixed, non-template 1024-thread panel variant. This intentionally mirrors
// panel_kernel() but avoids adding another template instantiation of the main
// panel body; host routing below only tries it for n=1024/2048 when the dynamic
// shared-memory request fits.
__global__ void __launch_bounds__(1024)
panel1024_kernel(float* __restrict__ A, float* __restrict__ taug,
                 float* __restrict__ Vg, float* __restrict__ Tg,
                 int n, int k, int nb, int ldv, int voff) {
    extern __shared__ float P[];
    __shared__ float sT[NBMAX][NBMAX];
    __shared__ float sw[NBMAX];
    __shared__ float red[32];
    __shared__ float sb[2];

    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    const int lane = tid & 31;
    const int wid = tid >> 5;
    const int m = n - k;
    const int mp = (m & 1) ? m : m + 1;

    float* __restrict__ Ab = A + (size_t)b * n * n;

    for (int idx = tid; idx < m * nb; idx += 1024) {
        int i = idx / nb, c = idx % nb;
        P[c * mp + i] = Ab[(size_t)(k + i) * n + (k + c)];
    }
    __syncthreads();

    for (int j = 0; j < nb; ++j) {
        float* __restrict__ Pj = P + j * mp;

        float local = 0.f;
        for (int i = j + 1 + tid; i < m; i += 1024) {
            float x = Pj[i];
            local += x * x;
        }
        #pragma unroll
        for (int off = 16; off; off >>= 1)
            local += __shfl_down_sync(0xffffffffu, local, off);
        if (lane == 0) red[wid] = local;
        __syncthreads();

        if (tid == 0) {
            float sigma = 0.f;
            #pragma unroll
            for (int w = 0; w < 32; ++w) sigma += red[w];
            float alpha = Pj[j];
            float tau_j, scale;
            if (sigma == 0.f) {
                tau_j = 0.f;
                scale = 0.f;
            } else {
                float nrm = sqrtf(alpha * alpha + sigma);
                float beta = (alpha >= 0.f) ? -nrm : nrm;
                tau_j = (beta - alpha) / beta;
                scale = 1.f / (alpha - beta);
                Pj[j] = beta;
            }
            sb[0] = tau_j;
            sb[1] = scale;
            sT[j][j] = tau_j;
            taug[(size_t)b * n + (k + j)] = tau_j;
        }
        __syncthreads();

        const float tau_j = sb[0];
        const float scale = sb[1];

        for (int i = j + 1 + tid; i < m; i += 1024) Pj[i] *= scale;
        __syncthreads();

        for (int c = wid; c < nb; c += 32) {
            if (c == j) continue;
            float* __restrict__ Pc = P + c * mp;
            float acc = 0.f;
            for (int i = j + 1 + lane; i < m; i += 32)
                acc += Pc[i] * Pj[i];
            #pragma unroll
            for (int off = 16; off; off >>= 1)
                acc += __shfl_down_sync(0xffffffffu, acc, off);
            float dot = __shfl_sync(0xffffffffu, acc, 0) + Pc[j];
            if (c < j) {
                if (lane == 0) sw[c] = dot;
            } else {
                float y = tau_j * dot;
                if (lane == 0) Pc[j] -= y;
                for (int i = j + 1 + lane; i < m; i += 32)
                    Pc[i] -= Pj[i] * y;
            }
        }
        __syncthreads();

        if (tid < j) {
            float acc = 0.f;
            for (int c = tid; c < j; ++c) acc += sT[tid][c] * sw[c];
            sT[tid][j] = -tau_j * acc;
        }
        __syncthreads();
    }

    for (int idx = tid; idx < m * nb; idx += 1024) {
        int i = idx / nb, c = idx % nb;
        Ab[(size_t)(k + i) * n + (k + c)] = P[c * mp + i];
    }
    float* __restrict__ Vb = Vg + (size_t)b * n * ldv + voff;
    for (int idx = tid; idx < voff * nb; idx += 1024) {
        int i = idx / nb, c = idx % nb;
        Vb[(size_t)(k - voff + i) * ldv + c] = 0.f;
    }
    for (int idx = tid; idx < m * nb; idx += 1024) {
        int i = idx / nb, c = idx % nb;
        float v = (i < c) ? 0.f : (i == c) ? 1.f : P[c * mp + i];
        Vb[(size_t)(k + i) * ldv + c] = v;
    }
    float* __restrict__ Tb = Tg + (size_t)b * NBMAX * NBMAX;
    for (int idx = tid; idx < nb * nb; idx += 1024) {
        int r = idx / nb, c = idx % nb;
        Tb[(size_t)r * NBMAX + c] = (r <= c) ? sT[r][c] : 0.f;
    }
}

// Final panel variant. It performs the same in-panel Householder QR and writes
// A/tau, but skips V/T construction that no later trailing update can consume.
template <int THREADS>
__global__ void __launch_bounds__(THREADS)
panel_final_kernel(float* __restrict__ A, float* __restrict__ taug,
                   int n, int k, int nb) {
    extern __shared__ float P[];
    __shared__ float red[THREADS / 32];
    __shared__ float sb[2];

    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    const int lane = tid & 31;
    const int wid = tid >> 5;
    const int nwarps = THREADS / 32;
    const int m = n - k;
    const int mp = (m & 1) ? m : m + 1;

    float* __restrict__ Ab = A + (size_t)b * n * n;

    for (int idx = tid; idx < m * nb; idx += THREADS) {
        int i = idx / nb, c = idx % nb;
        P[c * mp + i] = Ab[(size_t)(k + i) * n + (k + c)];
    }
    __syncthreads();

    for (int j = 0; j < nb; ++j) {
        float* __restrict__ Pj = P + j * mp;

        float local = 0.f;
        for (int i = j + 1 + tid; i < m; i += THREADS) {
            float x = Pj[i];
            local += x * x;
        }
        #pragma unroll
        for (int off = 16; off; off >>= 1)
            local += __shfl_down_sync(0xffffffffu, local, off);
        if (lane == 0) red[wid] = local;
        __syncthreads();

        if (tid == 0) {
            float sigma = 0.f;
            for (int w = 0; w < nwarps; ++w) sigma += red[w];
            float alpha = Pj[j];
            float tau_j, scale;
            if (sigma == 0.f) {
                tau_j = 0.f;
                scale = 0.f;
            } else {
                float nrm = sqrtf(alpha * alpha + sigma);
                float beta = (alpha >= 0.f) ? -nrm : nrm;
                tau_j = (beta - alpha) / beta;
                scale = 1.f / (alpha - beta);
                Pj[j] = beta;
            }
            sb[0] = tau_j;
            sb[1] = scale;
            taug[(size_t)b * n + (k + j)] = tau_j;
        }
        __syncthreads();

        const float tau_j = sb[0];
        const float scale = sb[1];

        for (int i = j + 1 + tid; i < m; i += THREADS) Pj[i] *= scale;
        __syncthreads();

        for (int c = j + 1 + wid; c < nb; c += nwarps) {
            float* __restrict__ Pc = P + c * mp;
            float acc = 0.f;
            for (int i = j + 1 + lane; i < m; i += 32)
                acc += Pc[i] * Pj[i];
            #pragma unroll
            for (int off = 16; off; off >>= 1)
                acc += __shfl_down_sync(0xffffffffu, acc, off);
            float dot = __shfl_sync(0xffffffffu, acc, 0) + Pc[j];
            float y = tau_j * dot;
            if (lane == 0) Pc[j] -= y;
            for (int i = j + 1 + lane; i < m; i += 32)
                Pc[i] -= Pj[i] * y;
        }
        __syncthreads();
    }

    for (int idx = tid; idx < m * nb; idx += THREADS) {
        int i = idx / nb, c = idx % nb;
        Ab[(size_t)(k + i) * n + (k + c)] = P[c * mp + i];
    }
}

template <int THREADS>
__global__ void __launch_bounds__(THREADS)
panel_full_out_n32_kernel(const float* __restrict__ A,
                          float* __restrict__ H,
                          float* __restrict__ taug) {
    constexpr int n = 32;
    constexpr int mp = 33;
    extern __shared__ float P[];
    __shared__ float red[THREADS / 32];
    __shared__ float sb[2];

    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    const int lane = tid & 31;
    const int wid = tid >> 5;
    const int nwarps = THREADS / 32;

    const float* __restrict__ Ab = A + (size_t)b * n * n;
    float* __restrict__ Hb = H + (size_t)b * n * n;

    for (int idx = tid; idx < n * n; idx += THREADS) {
        int i = idx / n, c = idx % n;
        P[(size_t)c * mp + i] = Ab[(size_t)i * n + c];
    }
    __syncthreads();

    for (int j = 0; j < n; ++j) {
        float* __restrict__ Pj = P + (size_t)j * mp;
        float local = 0.f;
        for (int i = j + 1 + tid; i < n; i += THREADS) {
            float x = Pj[i];
            local += x * x;
        }
        #pragma unroll
        for (int off = 16; off; off >>= 1)
            local += __shfl_down_sync(0xffffffffu, local, off);
        if (lane == 0) red[wid] = local;
        __syncthreads();

        if (tid == 0) {
            float sigma = 0.f;
            for (int w = 0; w < nwarps; ++w) sigma += red[w];
            float alpha = Pj[j];
            float tau_j, scale;
            if (sigma == 0.f) {
                tau_j = 0.f;
                scale = 0.f;
            } else {
                float nrm = sqrtf(alpha * alpha + sigma);
                float beta = (alpha >= 0.f) ? -nrm : nrm;
                tau_j = (beta - alpha) / beta;
                scale = 1.f / (alpha - beta);
                Pj[j] = beta;
            }
            sb[0] = tau_j;
            sb[1] = scale;
            taug[(size_t)b * n + j] = tau_j;
        }
        __syncthreads();

        const float tau_j = sb[0];
        const float scale = sb[1];
        for (int i = j + 1 + tid; i < n; i += THREADS) Pj[i] *= scale;
        __syncthreads();

        for (int c = j + 1 + wid; c < n; c += nwarps) {
            float* __restrict__ Pc = P + (size_t)c * mp;
            float acc = 0.f;
            for (int i = j + 1 + lane; i < n; i += 32)
                acc += Pc[i] * Pj[i];
            #pragma unroll
            for (int off = 16; off; off >>= 1)
                acc += __shfl_down_sync(0xffffffffu, acc, off);
            float dot = __shfl_sync(0xffffffffu, acc, 0) + Pc[j];
            float y = tau_j * dot;
            if (lane == 0) Pc[j] -= y;
            for (int i = j + 1 + lane; i < n; i += 32)
                Pc[i] -= Pj[i] * y;
        }
        __syncthreads();
    }

    for (int idx = tid; idx < n * n; idx += THREADS) {
        int i = idx / n, c = idx % n;
        Hb[(size_t)i * n + c] = P[(size_t)c * mp + i];
    }
}

std::vector<torch::Tensor> monolithic_qr_n32(torch::Tensor data) {
    const int b = (int)data.size(0);
    TORCH_CHECK((int)data.size(1) == 32 && (int)data.size(2) == 32,
                "monolithic_qr_n32 expects n=32");
    auto src = data.contiguous();
    auto opts = data.options().dtype(torch::kFloat32);
    auto H = torch::empty({b, 32, 32}, opts);
    auto tau = torch::empty({b, 32}, opts);
    constexpr int TH = 512;
    constexpr size_t smem = (size_t)33 * 32 * sizeof(float);
    panel_full_out_n32_kernel<TH><<<b, TH, smem>>>(
        src.data_ptr<float>(), H.data_ptr<float>(), tau.data_ptr<float>());
    C10_CUDA_KERNEL_LAUNCH_CHECK();
    return {H, tau};
}

void panel_factor(torch::Tensor A, torch::Tensor tau, torch::Tensor V,
                  torch::Tensor T, int64_t k, int64_t nb, int64_t voff) {
    const int b = A.size(0);
    const int n = A.size(1);
    const int ldv = V.size(2);
    const int m = n - (int)k;
    const int mp = (m & 1) ? m : m + 1;
    const size_t smem = (size_t)mp * nb * sizeof(float);

    // Static + dynamic shared memory exceeds the 48 KB default already at
    // n=352, so opt in once to the device's full per-block capacity.
    static int max_dyn_smem = -1;
    static int max_dyn_smem_1024 = -1;
    if (max_dyn_smem < 0) {
        int device;
        C10_CUDA_CHECK(cudaGetDevice(&device));
        cudaDeviceProp prop;
        C10_CUDA_CHECK(cudaGetDeviceProperties(&prop, device));
        cudaFuncAttributes a256, a512, a1024;
        C10_CUDA_CHECK(cudaFuncGetAttributes(&a256, panel_kernel<256>));
        C10_CUDA_CHECK(cudaFuncGetAttributes(&a512, panel_kernel<512>));
        C10_CUDA_CHECK(cudaFuncGetAttributes(&a1024, panel1024_kernel));
        int m256 = (int)(prop.sharedMemPerBlockOptin - a256.sharedSizeBytes);
        int m512 = (int)(prop.sharedMemPerBlockOptin - a512.sharedSizeBytes);
        int m1024 = (int)(prop.sharedMemPerBlockOptin - a1024.sharedSizeBytes);
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            panel_kernel<256>, cudaFuncAttributeMaxDynamicSharedMemorySize,
            m256));
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            panel_kernel<512>, cudaFuncAttributeMaxDynamicSharedMemorySize,
            m512));
        if (m1024 > 0) {
            C10_CUDA_CHECK(cudaFuncSetAttribute(
                panel1024_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize,
                m1024));
        }
        max_dyn_smem = m256 < m512 ? m256 : m512;
        max_dyn_smem_1024 = m1024 > 0 ? m1024 : 0;
    }
    if ((n == 352 || n == 1024) && smem <= (size_t)max_dyn_smem_1024) {
        panel1024_kernel<<<b, 1024, smem>>>(
            A.data_ptr<float>(), tau.data_ptr<float>(), V.data_ptr<float>(),
            T.data_ptr<float>(), n, (int)k, (int)nb, ldv, (int)voff);
        C10_CUDA_KERNEL_LAUNCH_CHECK();
        return;
    }
    TORCH_CHECK(smem <= (size_t)max_dyn_smem,
                "panel does not fit in shared memory");
    // Low-batch large-n (n>=2048, batch=8) is occupancy-starved in the panel:
    // only ~8 CTAs run, so more threads/CTA cut each panel's reduction latency.
    // Measured ~20% faster panels at n=2048 with 512 vs 256 threads on B200.
    if (nb > 32 || n >= 2048 || n <= 64 || n == 176 || n == 352) {
        panel_kernel<512><<<b, 512, smem>>>(
            A.data_ptr<float>(), tau.data_ptr<float>(), V.data_ptr<float>(),
            T.data_ptr<float>(), n, (int)k, (int)nb, ldv, (int)voff);
    } else {
        panel_kernel<256><<<b, 256, smem>>>(
            A.data_ptr<float>(), tau.data_ptr<float>(), V.data_ptr<float>(),
            T.data_ptr<float>(), n, (int)k, (int)nb, ldv, (int)voff);
    }
    C10_CUDA_KERNEL_LAUNCH_CHECK();
}

void panel_factor_final(torch::Tensor A, torch::Tensor tau,
                        int64_t k, int64_t nb) {
    const int b = A.size(0);
    const int n = A.size(1);
    const int m = n - (int)k;
    const int mp = (m & 1) ? m : m + 1;
    const size_t smem = (size_t)mp * nb * sizeof(float);

    static int max_dyn_smem = -1;
    if (max_dyn_smem < 0) {
        int device;
        C10_CUDA_CHECK(cudaGetDevice(&device));
        cudaDeviceProp prop;
        C10_CUDA_CHECK(cudaGetDeviceProperties(&prop, device));
        cudaFuncAttributes a256, a512;
        C10_CUDA_CHECK(cudaFuncGetAttributes(&a256, panel_final_kernel<256>));
        C10_CUDA_CHECK(cudaFuncGetAttributes(&a512, panel_final_kernel<512>));
        int m256 = (int)(prop.sharedMemPerBlockOptin - a256.sharedSizeBytes);
        int m512 = (int)(prop.sharedMemPerBlockOptin - a512.sharedSizeBytes);
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            panel_final_kernel<256>, cudaFuncAttributeMaxDynamicSharedMemorySize,
            m256));
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            panel_final_kernel<512>, cudaFuncAttributeMaxDynamicSharedMemorySize,
            m512));
        max_dyn_smem = m256 < m512 ? m256 : m512;
    }
    TORCH_CHECK(smem <= (size_t)max_dyn_smem,
                "final panel does not fit in shared memory");
    if (nb > 32 || n >= 2048 || n <= 64 || n == 176 || n == 352) {
        panel_final_kernel<512><<<b, 512, smem>>>(
            A.data_ptr<float>(), tau.data_ptr<float>(), n, (int)k, (int)nb);
    } else {
        panel_final_kernel<256><<<b, 256, smem>>>(
            A.data_ptr<float>(), tau.data_ptr<float>(), n, (int)k, (int)nb);
    }
    C10_CUDA_KERNEL_LAUNCH_CHECK();
}

// Inv 189: override the two-level outer trailing-update mode (-1 = use default).
// 0=FP32, 1=BF16x9 emu, 2=FP16x2 Dekker (FP32-accurate). Lets us test whether the
// n512 trailing GEMM (the biggest FP32 scored case) is tensor-core-acceleratable.
static int g_trail_mode = -1;
void set_trail_mode(int64_t v) { g_trail_mode = (int)v; }

// Inv 191: override n512 two-level blocking (inner nbp, outer wout). wout is the
// trailing-update contracted dim vc; the Inv 190 crux shows vc=128 makes the
// FP16/FP16x2 trailing GEMM ~free while FP32 ~doubles. wout is independent of NBMAX
// (Tout is wout x wout), so wide outer blocks need no NBMAX change. -1/0 = default.
static int g_n512_nbp = -1, g_n512_wout = -1;
void set_n512_blocking(int64_t nbp, int64_t wout) {
    g_n512_nbp = (int)nbp; g_n512_wout = (int)wout;
}

// Inv 195: override n<=352 blocking. Low-batch small-n (b40, 40 CTAs) is already
// occupancy-starved, so wider panels (fewer launches) can't hurt occupancy the way
// they do at n512 b640 -- the launch/overhead-bound small cases may net-win on width.
static int g_small_nbp = -1, g_small_wout = -1;
void set_small_blocking(int64_t nbp, int64_t wout) {
    g_small_nbp = (int)nbp; g_small_wout = (int)wout;
}

// Dekker-style error-free split: x = (float)hi + (float)lo with hi, lo
// FP16. Each tensor-core product of split halves is exact; accumulating
// hi*hi + hi*lo + lo*hi in FP32 recovers ~22 mantissa bits, well inside
// the QR checker budget (~18 bits) -- unlike direct FP16/BF16/NVFP4,
// which miss it by 2-4 orders of magnitude (see sim_precision.py).
__global__ void split_f16_kernel(const float* __restrict__ src,
                                 __half* __restrict__ hi,
                                 __half* __restrict__ lo,
                                 int rows, int cols, int ld,
                                 long long bstride, long long total) {
    const long long per = (long long)rows * cols;
    for (long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
         idx < total; idx += (long long)gridDim.x * blockDim.x) {
        long long bi = idx / per;
        int rem = (int)(idx - bi * per);
        float x = src[bi * bstride + (long long)(rem / cols) * ld + rem % cols];
        __half h = __float2half_rn(x);
        hi[idx] = h;
        lo[idx] = __float2half_rn(x - __half2float(h));
    }
}

static void split_f16(const float* src, torch::Tensor& hi, torch::Tensor& lo,
                      int bb, int rows, int cols, int ld, long long bstride) {
    long long total = (long long)bb * rows * cols;
    int blocks = (int)std::min<long long>((total + 255) / 256, 16384);
    split_f16_kernel<<<blocks, 256>>>(
        src, (__half*)hi.data_ptr(), (__half*)lo.data_ptr(),
        rows, cols, ld, bstride, total);
    C10_CUDA_KERNEL_LAUNCH_CHECK();
}

__global__ void cast_f16_kernel(const float* __restrict__ src,
                                __half* __restrict__ dst,
                                int rows, int cols, int ld,
                                long long bstride, long long total) {
    const long long per = (long long)rows * cols;
    for (long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
         idx < total; idx += (long long)gridDim.x * blockDim.x) {
        long long bi = idx / per;
        int rem = (int)(idx - bi * per);
        float x = src[bi * bstride + (long long)(rem / cols) * ld + rem % cols];
        dst[idx] = __float2half_rn(x);
    }
}

static void cast_f16(const float* src, torch::Tensor& dst,
                     int bb, int rows, int cols, int ld, long long bstride) {
    long long total = (long long)bb * rows * cols;
    int blocks = (int)std::min<long long>((total + 255) / 256, 16384);
    cast_f16_kernel<<<blocks, 256>>>(
        src, (__half*)dst.data_ptr(), rows, cols, ld, bstride, total);
    C10_CUDA_KERNEL_LAUNCH_CHECK();
}

// Fused FP32 compact-WY trailing update for moderate panels:
//   A -= V T^T V^T A
// One CTA owns one matrix and a small tile of trailing columns. W1/W2 live in
// shared memory, so this replaces three cuBLAS launches and two global scratch
// tensors for shapes where scalar FP32 parallelism is enough to compete.
template <int TILE_COLS, int VC_MAX>
__global__ void __launch_bounds__(256)
fused_wy_update_kernel(float* __restrict__ H,
                       const float* __restrict__ V,
                       const float* __restrict__ T,
                       int n, int ldv, int ldt,
                       int row0, int col0, int r,
                       int voff, int vc) {
    __shared__ float W1[VC_MAX * TILE_COLS];
    __shared__ float W2[VC_MAX * TILE_COLS];

    const int tile = blockIdx.x * TILE_COLS;
    const int rem_cols = r - tile;
    const int cols = rem_cols < TILE_COLS ? rem_cols : TILE_COLS;
    const int b = blockIdx.y;
    const int tid = threadIdx.x;
    const int m = n - row0;

    float* __restrict__ Hp =
        H + (size_t)b * n * n + (size_t)row0 * n + col0 + tile;
    const float* __restrict__ Vp =
        V + (size_t)b * n * ldv + (size_t)row0 * ldv + voff;
    const float* __restrict__ Tp = T + (size_t)b * ldt * ldt;

    for (int idx = tid; idx < vc * TILE_COLS; idx += blockDim.x) {
        W1[idx] = 0.f;
        W2[idx] = 0.f;
    }
    __syncthreads();

    for (int idx = tid; idx < vc * cols; idx += blockDim.x) {
        int c = idx / cols;
        int j = idx - c * cols;
        float acc = 0.f;
        const float* __restrict__ vcol = Vp + c;
        const float* __restrict__ acol = Hp + j;
        for (int i = 0; i < m; ++i) {
            acc += vcol[(size_t)i * ldv] * acol[(size_t)i * n];
        }
        W1[c * TILE_COLS + j] = acc;
    }
    __syncthreads();

    for (int idx = tid; idx < vc * cols; idx += blockDim.x) {
        int c = idx / cols;
        int j = idx - c * cols;
        float acc = 0.f;
        for (int l = 0; l < vc; ++l) {
            acc += Tp[(size_t)l * ldt + c] * W1[l * TILE_COLS + j];
        }
        W2[c * TILE_COLS + j] = acc;
    }
    __syncthreads();

    const int total = m * cols;
    for (int idx = tid; idx < total; idx += blockDim.x) {
        int i = idx / cols;
        int j = idx - i * cols;
        float acc = 0.f;
        const float* __restrict__ vrow = Vp + (size_t)i * ldv;
        for (int c = 0; c < vc; ++c) {
            acc += vrow[c] * W2[c * TILE_COLS + j];
        }
        Hp[(size_t)i * n + j] -= acc;
    }
}

#if QR_ENABLE_QRF4_WY
template <int TILE_N>
struct Qrf4WyShape {
    static constexpr int M = 128;
    static constexpr int VC = 64;
    static constexpr int K = 64;
    static constexpr int Threads = 128;
    static constexpr int TmemCols = TILE_N;
    static constexpr int AllocCols = (TILE_N <= 64) ? 128 : ((TILE_N <= 128) ? 256 : 512);
    static constexpr int ScaleCols = TILE_N / 16;
    static constexpr int PackedRowBytes = K / 2;
    static constexpr int ABytes = M * PackedRowBytes;
    static constexpr int BBytes = (K / 32) * TILE_N * 16;
    static constexpr int WBytes = BBytes;
    static constexpr int BOffset = ABytes;
    static constexpr int W1Offset = BOffset + BBytes;
    static constexpr int W2Offset = W1Offset + WBytes;
    static constexpr int MbarOffset = W2Offset + WBytes;
    static constexpr int TmemPtrOffset = MbarOffset + 8;
    static constexpr int SmemBytes = TmemPtrOffset + 4;
};

__device__ __forceinline__ uint32_t qrf4_smem_u32(void* ptr) {
    return static_cast<uint32_t>(__cvta_generic_to_shared(ptr));
}

__device__ __forceinline__ uint32_t qrf4_desc_encode(uint32_t x) {
    return (x & 0x3ffffu) >> 4;
}

__device__ __forceinline__ uint64_t qrf4_make_k_major_desc(uint32_t addr,
                                                            int height) {
    const uint32_t lbo = static_cast<uint32_t>(height * 16);
    const uint32_t sbo = static_cast<uint32_t>(8 * 16);
    return static_cast<uint64_t>(qrf4_desc_encode(addr)) |
           (static_cast<uint64_t>(qrf4_desc_encode(lbo)) << 16) |
           (static_cast<uint64_t>(qrf4_desc_encode(sbo)) << 32) |
           (1ull << 46);
}

template <int TILE_N>
__device__ __forceinline__ int qrf4_bi(int kk, int col) {
    const int slice = kk >> 5;
    const int byte = (kk & 31) >> 1;
    return (slice * TILE_N + col) * 16 + byte;
}

__device__ __forceinline__ uint8_t qrf4_pack_e2m1x2(float lo, float hi) {
    uint32_t out;
    asm volatile(
        "{\n\t"
        ".reg .b8 byte0;\n\t"
        "cvt.rn.satfinite.e2m1x2.f32 byte0, %2, %1;\n\t"
        "mov.b32 %0, {byte0, byte0, byte0, byte0};\n\t"
        "}\n"
        : "=r"(out)
        : "f"(lo), "f"(hi));
    return static_cast<uint8_t>(out);
}

__device__ __forceinline__ uint8_t qrf4_pack_scaled(float lo, float hi) {
    return qrf4_pack_e2m1x2(lo * 16.0f, hi * 16.0f);
}

__device__ __forceinline__ uint32_t qrf4_idesc(int m, int n) {
    return (5u << 7) |
           (5u << 10) |
           (static_cast<uint32_t>(n >> 3) << 17) |
           (1u << 23) |
           (static_cast<uint32_t>(m >> 4) << 24);
}

__device__ __forceinline__ void qrf4_mbar_init(uint32_t addr,
                                                uint32_t arrivals) {
    asm volatile(
        "{\n\t"
        "mbarrier.init.shared::cta.b64 [%0], %1;\n\t"
        "fence.mbarrier_init.release.cluster;\n\t"
        "}"
        :
        : "r"(addr), "r"(arrivals)
        : "memory");
}

__device__ __forceinline__ void qrf4_mbar_wait(uint32_t addr,
                                                uint32_t phase) {
    asm volatile(
        "{\n\t"
        ".reg .pred p;\n\t"
        "QRF4_WAIT:\n\t"
        "mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 p, [%0], %1;\n\t"
        "@p bra QRF4_DONE;\n\t"
        "bra QRF4_WAIT;\n\t"
        "QRF4_DONE:\n\t"
        "}"
        :
        : "r"(addr), "r"(phase)
        : "memory");
}

__device__ __forceinline__ void qrf4_fence_shared() {
    asm volatile("fence.proxy.async.shared::cta;" ::: "memory");
}

__device__ __forceinline__ void qrf4_alloc(uint32_t dst_smem,
                                            uint32_t cols) {
    asm volatile(
        "tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;"
        :
        : "r"(dst_smem), "r"(cols)
        : "memory");
}

__device__ __forceinline__ void qrf4_dealloc(uint32_t taddr,
                                              uint32_t cols) {
    asm volatile(
        "{\n\t"
        "tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;\n\t"
        "tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;\n\t"
        "}"
        :
        : "r"(taddr), "r"(cols)
        : "memory");
}

__device__ __forceinline__ void qrf4_st_x8(
    uint32_t taddr, uint32_t r0, uint32_t r1, uint32_t r2, uint32_t r3,
    uint32_t r4, uint32_t r5, uint32_t r6, uint32_t r7) {
    asm volatile(
        "tcgen05.st.sync.aligned.32x32b.x8.b32 "
        "[%0], {%1, %2, %3, %4, %5, %6, %7, %8};"
        :
        : "r"(taddr), "r"(r0), "r"(r1), "r"(r2), "r"(r3),
          "r"(r4), "r"(r5), "r"(r6), "r"(r7)
        : "memory");
}

__device__ __forceinline__ void qrf4_st_x1(uint32_t taddr,
                                            uint32_t r0) {
    asm volatile("tcgen05.st.sync.aligned.32x32b.x1.b32 [%0], {%1};"
                 :
                 : "r"(taddr), "r"(r0)
                 : "memory");
}

__device__ __forceinline__ void qrf4_wait_st() {
    asm volatile("tcgen05.wait::st.sync.aligned;" ::: "memory");
}

__device__ __forceinline__ void qrf4_mma(uint32_t taddr_d,
                                          uint32_t taddr_a,
                                          uint64_t b_desc,
                                          uint32_t idesc,
                                          uint32_t input_d,
                                          uint32_t taddr_sfa,
                                          uint32_t taddr_sfb) {
    asm volatile(
        "{\n\t"
        ".reg .pred p;\n\t"
        "setp.ne.b32 p, %4, 0;\n\t"
        "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 "
        "[%0], [%1], %2, %3, [%5], [%6], p;\n\t"
        "}"
        :
        : "r"(taddr_d), "r"(taddr_a), "l"(b_desc), "r"(idesc),
          "r"(input_d), "r"(taddr_sfa), "r"(taddr_sfb)
        : "memory");
}

__device__ __forceinline__ void qrf4_commit(uint32_t addr) {
    asm volatile(
        "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
        :
        : "r"(addr)
        : "memory");
}

__device__ __forceinline__ void qrf4_ld_x8(
    uint32_t taddr, uint32_t& r0, uint32_t& r1, uint32_t& r2, uint32_t& r3,
    uint32_t& r4, uint32_t& r5, uint32_t& r6, uint32_t& r7) {
    asm volatile(
        "tcgen05.ld.sync.aligned.32x32b.x8.b32 "
        "{%0, %1, %2, %3, %4, %5, %6, %7}, [%8];"
        : "=r"(r0), "=r"(r1), "=r"(r2), "=r"(r3),
          "=r"(r4), "=r"(r5), "=r"(r6), "=r"(r7)
        : "r"(taddr)
        : "memory");
}

__device__ __forceinline__ void qrf4_wait_ld() {
    asm volatile("tcgen05.wait::ld.sync.aligned;" ::: "memory");
}

template <int TILE_N>
__device__ __forceinline__ void qrf4_store_a(uint32_t taddr_a,
                                              const uint8_t* Arow,
                                              int tid) {
    using S = Qrf4WyShape<TILE_N>;
    if (tid < S::M) {
        const uint32_t* words =
            reinterpret_cast<const uint32_t*>(Arow + tid * S::PackedRowBytes);
        qrf4_st_x8(taddr_a, words[0], words[1], words[2], words[3],
                   words[4], words[5], words[6], words[7]);
    }
    qrf4_wait_st();
}

template <int TILE_N>
__device__ __forceinline__ void qrf4_init_scale(uint32_t taddr_sfa,
                                                 uint32_t taddr_sfb) {
    const uint32_t sf = 0x7b7b7b7bu;
    for (int sf_col = 0; sf_col < Qrf4WyShape<TILE_N>::ScaleCols; ++sf_col) {
        qrf4_st_x1(taddr_sfa + sf_col, sf);
        qrf4_st_x1(taddr_sfb + sf_col, sf);
    }
    qrf4_wait_st();
}

__device__ __forceinline__ void qrf4_issue_wait(uint32_t taddr_d,
                                                 uint32_t taddr_a,
                                                 uint64_t b_desc,
                                                 uint32_t idesc,
                                                 uint32_t mbar_addr,
                                                 uint32_t& phase,
                                                 bool accumulate,
                                                 uint32_t taddr_sfa,
                                                 uint32_t taddr_sfb,
                                                 int tid) {
    if (tid == 0) {
        qrf4_mma(taddr_d, taddr_a, b_desc, idesc,
                 accumulate ? 1u : 0u, taddr_sfa, taddr_sfb);
        qrf4_commit(mbar_addr);
    }
    qrf4_mbar_wait(mbar_addr, phase);
    asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory");
    phase ^= 1u;
}

template <int TILE_N>
__device__ __forceinline__ void qrf4_drain_to_fp4(uint32_t taddr_d,
                                                   uint8_t* W,
                                                   int warp,
                                                   int lane) {
    if (warp < 4) {
        const int row = warp * 32 + lane;
        for (int n8 = 0; n8 < TILE_N; n8 += 8) {
            const uint32_t load_addr =
                taddr_d + (static_cast<uint32_t>(warp * 32) << 16) + n8;
            uint32_t r0, r1, r2, r3, r4, r5, r6, r7;
            qrf4_ld_x8(load_addr, r0, r1, r2, r3, r4, r5, r6, r7);
            qrf4_wait_ld();
            const uint32_t o0 = __shfl_down_sync(0xffffffffu, r0, 1);
            const uint32_t o1 = __shfl_down_sync(0xffffffffu, r1, 1);
            const uint32_t o2 = __shfl_down_sync(0xffffffffu, r2, 1);
            const uint32_t o3 = __shfl_down_sync(0xffffffffu, r3, 1);
            const uint32_t o4 = __shfl_down_sync(0xffffffffu, r4, 1);
            const uint32_t o5 = __shfl_down_sync(0xffffffffu, r5, 1);
            const uint32_t o6 = __shfl_down_sync(0xffffffffu, r6, 1);
            const uint32_t o7 = __shfl_down_sync(0xffffffffu, r7, 1);
            if ((lane & 1) == 0 && row + 1 < Qrf4WyShape<TILE_N>::VC) {
                W[qrf4_bi<TILE_N>(row, n8 + 0)] =
                    qrf4_pack_scaled(__uint_as_float(r0), __uint_as_float(o0));
                W[qrf4_bi<TILE_N>(row, n8 + 1)] =
                    qrf4_pack_scaled(__uint_as_float(r1), __uint_as_float(o1));
                W[qrf4_bi<TILE_N>(row, n8 + 2)] =
                    qrf4_pack_scaled(__uint_as_float(r2), __uint_as_float(o2));
                W[qrf4_bi<TILE_N>(row, n8 + 3)] =
                    qrf4_pack_scaled(__uint_as_float(r3), __uint_as_float(o3));
                W[qrf4_bi<TILE_N>(row, n8 + 4)] =
                    qrf4_pack_scaled(__uint_as_float(r4), __uint_as_float(o4));
                W[qrf4_bi<TILE_N>(row, n8 + 5)] =
                    qrf4_pack_scaled(__uint_as_float(r5), __uint_as_float(o5));
                W[qrf4_bi<TILE_N>(row, n8 + 6)] =
                    qrf4_pack_scaled(__uint_as_float(r6), __uint_as_float(o6));
                W[qrf4_bi<TILE_N>(row, n8 + 7)] =
                    qrf4_pack_scaled(__uint_as_float(r7), __uint_as_float(o7));
            }
        }
    }
}

template <int TILE_N>
__device__ __forceinline__ void qrf4_drain_update(uint32_t taddr_d,
                                                   const float* Hp,
                                                   float* Cp,
                                                   int cols,
                                                   int n,
                                                   int warp,
                                                   int lane) {
    if (warp < 4) {
        const int row = warp * 32 + lane;
        for (int n8 = 0; n8 < TILE_N; n8 += 8) {
            const uint32_t load_addr =
                taddr_d + (static_cast<uint32_t>(warp * 32) << 16) + n8;
            uint32_t r0, r1, r2, r3, r4, r5, r6, r7;
            qrf4_ld_x8(load_addr, r0, r1, r2, r3, r4, r5, r6, r7);
            qrf4_wait_ld();
            if (n8 + 0 < cols) Cp[(size_t)row * n + n8 + 0] = Hp[(size_t)row * n + n8 + 0] - __uint_as_float(r0);
            if (n8 + 1 < cols) Cp[(size_t)row * n + n8 + 1] = Hp[(size_t)row * n + n8 + 1] - __uint_as_float(r1);
            if (n8 + 2 < cols) Cp[(size_t)row * n + n8 + 2] = Hp[(size_t)row * n + n8 + 2] - __uint_as_float(r2);
            if (n8 + 3 < cols) Cp[(size_t)row * n + n8 + 3] = Hp[(size_t)row * n + n8 + 3] - __uint_as_float(r3);
            if (n8 + 4 < cols) Cp[(size_t)row * n + n8 + 4] = Hp[(size_t)row * n + n8 + 4] - __uint_as_float(r4);
            if (n8 + 5 < cols) Cp[(size_t)row * n + n8 + 5] = Hp[(size_t)row * n + n8 + 5] - __uint_as_float(r5);
            if (n8 + 6 < cols) Cp[(size_t)row * n + n8 + 6] = Hp[(size_t)row * n + n8 + 6] - __uint_as_float(r6);
            if (n8 + 7 < cols) Cp[(size_t)row * n + n8 + 7] = Hp[(size_t)row * n + n8 + 7] - __uint_as_float(r7);
        }
    }
}

template <int TILE_N>
__global__ __launch_bounds__(Qrf4WyShape<TILE_N>::Threads)
void qrf4_wy_m128n_kernel(float* __restrict__ H,
                          const float* __restrict__ V,
                          const float* __restrict__ T,
                          int bb, int n, int ldv, int ldt,
                          int row0, int col0, int r, int voff) {
#if __CUDA_ARCH__ >= 1000
    using S = Qrf4WyShape<TILE_N>;
    extern __shared__ __align__(1024) unsigned char smem[];
    auto* Arow = reinterpret_cast<uint8_t*>(smem);
    auto* Bbuf = reinterpret_cast<uint8_t*>(smem + S::BOffset);
    auto* W1 = reinterpret_cast<uint8_t*>(smem + S::W1Offset);
    auto* W2 = reinterpret_cast<uint8_t*>(smem + S::W2Offset);
    auto* mbar = reinterpret_cast<uint64_t*>(smem + S::MbarOffset);
    auto* tmem_box = reinterpret_cast<uint32_t*>(smem + S::TmemPtrOffset);

    const int tile = blockIdx.x * TILE_N;
    const int cols = min(TILE_N, r - tile);
    const int b = blockIdx.y;
    if (b >= bb || cols <= 0) return;

    const int tid = threadIdx.x;
    const int lane = tid & 31;
    const int warp = tid >> 5;
    float* Hp = H + (size_t)b * n * n + (size_t)row0 * n + col0 + tile;
    const float* Vp = V + (size_t)b * n * ldv + (size_t)row0 * ldv + voff;
    const float* Tp = T + (size_t)b * ldt * ldt;
    const uint32_t mbar_addr = qrf4_smem_u32(mbar);
    const uint32_t tmem_box_addr = qrf4_smem_u32(tmem_box);

    if (tid == 0) qrf4_mbar_init(mbar_addr, 1);
    if (warp == 1) qrf4_alloc(tmem_box_addr, S::AllocCols);
    __syncthreads();

    const uint32_t taddr_d = tmem_box[0];
    const uint32_t taddr_a = taddr_d + S::TmemCols;
    const uint32_t taddr_sfa = taddr_a + 8;
    const uint32_t taddr_sfb = taddr_sfa + S::ScaleCols;
    const uint32_t idesc = qrf4_idesc(S::M, TILE_N);
    uint32_t phase = 0;

    qrf4_init_scale<TILE_N>(taddr_sfa, taddr_sfb);
    __syncthreads();

    for (int chunk = 0; chunk < 2; ++chunk) {
        for (int idx = tid; idx < S::M * S::PackedRowBytes; idx += blockDim.x) {
            const int row = idx / S::PackedRowBytes;
            const int byte = idx - row * S::PackedRowBytes;
            const int kk = byte * 2;
            const int src0 = chunk * S::K + kk;
            float x0 = 0.0f;
            float x1 = 0.0f;
            if (row < S::VC) {
                x0 = Vp[(size_t)src0 * ldv + row];
                x1 = Vp[(size_t)(src0 + 1) * ldv + row];
            }
            Arow[idx] = qrf4_pack_scaled(x0, x1);
        }
        for (int idx = tid; idx < S::BBytes; idx += blockDim.x) {
            const int slice = idx / (TILE_N * 16);
            const int rem = idx - slice * TILE_N * 16;
            const int col = rem / 16;
            const int byte = rem - col * 16;
            const int kk = slice * 32 + byte * 2;
            const int src0 = chunk * S::K + kk;
            float x0 = 0.0f;
            float x1 = 0.0f;
            if (col < cols) {
                x0 = Hp[(size_t)src0 * n + col];
                x1 = Hp[(size_t)(src0 + 1) * n + col];
            }
            Bbuf[idx] = qrf4_pack_scaled(x0, x1);
        }
        __syncthreads();
        qrf4_fence_shared();
        __syncthreads();
        qrf4_store_a<TILE_N>(taddr_a, Arow, tid);
        __syncthreads();
        qrf4_issue_wait(taddr_d, taddr_a,
                        qrf4_make_k_major_desc(qrf4_smem_u32(Bbuf), TILE_N),
                        idesc, mbar_addr, phase, chunk != 0,
                        taddr_sfa, taddr_sfb, tid);
        __syncthreads();
    }
    qrf4_drain_to_fp4<TILE_N>(taddr_d, W1, warp, lane);
    __syncthreads();

    for (int idx = tid; idx < S::M * S::PackedRowBytes; idx += blockDim.x) {
        const int row = idx / S::PackedRowBytes;
        const int byte = idx - row * S::PackedRowBytes;
        const int kk = byte * 2;
        float x0 = 0.0f;
        float x1 = 0.0f;
        if (row < S::VC) {
            x0 = Tp[(size_t)kk * ldt + row];
            x1 = Tp[(size_t)(kk + 1) * ldt + row];
        }
        Arow[idx] = qrf4_pack_scaled(x0, x1);
    }
    __syncthreads();
    qrf4_fence_shared();
    __syncthreads();
    qrf4_store_a<TILE_N>(taddr_a, Arow, tid);
    __syncthreads();
    qrf4_issue_wait(taddr_d, taddr_a,
                    qrf4_make_k_major_desc(qrf4_smem_u32(W1), TILE_N),
                    idesc, mbar_addr, phase, false,
                    taddr_sfa, taddr_sfb, tid);
    __syncthreads();
    qrf4_drain_to_fp4<TILE_N>(taddr_d, W2, warp, lane);
    __syncthreads();

    for (int idx = tid; idx < S::M * S::PackedRowBytes; idx += blockDim.x) {
        const int row = idx / S::PackedRowBytes;
        const int byte = idx - row * S::PackedRowBytes;
        const int kk = byte * 2;
        const float x0 = Vp[(size_t)row * ldv + kk];
        const float x1 = Vp[(size_t)row * ldv + kk + 1];
        Arow[idx] = qrf4_pack_scaled(x0, x1);
    }
    __syncthreads();
    qrf4_fence_shared();
    __syncthreads();
    qrf4_store_a<TILE_N>(taddr_a, Arow, tid);
    __syncthreads();
    qrf4_issue_wait(taddr_d, taddr_a,
                    qrf4_make_k_major_desc(qrf4_smem_u32(W2), TILE_N),
                    idesc, mbar_addr, phase, false,
                    taddr_sfa, taddr_sfb, tid);
    __syncthreads();
    qrf4_drain_update<TILE_N>(taddr_d, Hp, Hp, cols, n, warp, lane);
    __syncthreads();

    if (warp == 1) qrf4_dealloc(taddr_d, S::AllocCols);
#endif
}
#endif

static bool try_fused_wy_update(float* H, const float* V, const float* T,
                                int bb, int n, int ldv, int ldt,
                                int row0, int col0, int r,
                                int voff, int vc, int mode) {
#if QR_ENABLE_QRF4_WY
    if (mode == 0 && vc == 64 && n - row0 == 128 && r > 0) {
        if (r >= 192) {
            dim3 grid((r + 255) / 256, bb);
            qrf4_wy_m128n_kernel<256><<<grid, Qrf4WyShape<256>::Threads,
                                         Qrf4WyShape<256>::SmemBytes>>>(
                H, V, T, bb, n, ldv, ldt, row0, col0, r, voff);
        } else if (r >= 96) {
            dim3 grid((r + 127) / 128, bb);
            qrf4_wy_m128n_kernel<128><<<grid, Qrf4WyShape<128>::Threads,
                                         Qrf4WyShape<128>::SmemBytes>>>(
                H, V, T, bb, n, ldv, ldt, row0, col0, r, voff);
        } else {
            dim3 grid((r + 63) / 64, bb);
            qrf4_wy_m128n_kernel<64><<<grid, Qrf4WyShape<64>::Threads,
                                        Qrf4WyShape<64>::SmemBytes>>>(
                H, V, T, bb, n, ldv, ldt, row0, col0, r, voff);
        }
        C10_CUDA_KERNEL_LAUNCH_CHECK();
        return true;
    }
#endif
    if (mode != 0 || vc <= 0 || r <= 0) return false;
    // The scalar fused kernel is for launch/global-traffic dominated updates.
    // Wider or taller cases stay on cuBLAS/BF16x9 until a real tcgen05/TMEM
    // implementation replaces the inner multiply loops.
    if (vc > 64 || n - row0 > 1024 || r > 1024) return false;
    dim3 grid((r + 15) / 16, bb);
    fused_wy_update_kernel<16, 64><<<grid, 256>>>(
        H, V, T, n, ldv, ldt, row0, col0, r, voff, vc);
    C10_CUDA_KERNEL_LAUNCH_CHECK();
    return true;
}

// Trailing update A := A - V T^T (V^T A) as strided-batched GEMMs, called
// directly so the compute path can be chosen per call site:
//   mode 0: native FP32
//   mode 1: cuBLAS BF16x9 FP32 emulation (Blackwell)
//   mode 2: FP16x2 split, 3 tensor-core products per GEMM (FP16 in,
//           FP32 out, COMPUTE_32F)
//   mode 3: direct FP16 operands with FP32 accumulation/output.
// Row-major torch tensors are passed as their column-major transposes:
//   W1^T (r x vc) = At^T (r x m, ld n) * Vk         -> gemm(N, T)
//   W2^T (r x vc) = W1^T (r x vc) * Tk              -> gemm(N, T)
//   At^T (r x m) -= W2^T (r x vc) * Vk^T (vc x m)   -> gemm(N, N)
void apply_qt(torch::Tensor H, torch::Tensor V, torch::Tensor T,
              int64_t row0, int64_t col0, int64_t col1,
              int64_t voff, int64_t vc_, int64_t mode) {
    const int bb = H.size(0);
    const int n = H.size(1);
    const int ldv = V.size(2);
    const int ldt = T.size(2);
    const int m = n - (int)row0;
    const int r = (int)(col1 - col0);
    const int vc = (int)vc_;

    auto W1 = torch::empty({bb, vc, r}, H.options());
    auto W2 = torch::empty({bb, vc, r}, H.options());
    float* Hp = H.data_ptr<float>() + row0 * n + col0;
    float* Vp = V.data_ptr<float>() + row0 * ldv + voff;
    float* Tp = T.data_ptr<float>();
    const long long sH = (long long)n * n;
    const long long sV = (long long)n * ldv;
    const long long sT = (long long)ldt * ldt;
    const long long sW = (long long)vc * r;
    const float one = 1.f, zero = 0.f, neg1 = -1.f;
    cublasHandle_t handle = at::cuda::getCurrentCUDABlasHandle();

    if (try_fused_wy_update(H.data_ptr<float>(), V.data_ptr<float>(),
                            T.data_ptr<float>(), bb, n, ldv, ldt,
                            (int)row0, (int)col0, r, (int)voff, vc,
                            (int)mode)) {
        return;
    }

    if (mode == 2) {
        auto h16 = H.options().dtype(torch::kHalf);
        auto Ahi = torch::empty({bb, m, r}, h16);
        auto Alo = torch::empty({bb, m, r}, h16);
        auto Vhi = torch::empty({bb, m, vc}, h16);
        auto Vlo = torch::empty({bb, m, vc}, h16);
        auto Whi = torch::empty({bb, vc, r}, h16);
        auto Wlo = torch::empty({bb, vc, r}, h16);
        split_f16(Hp, Ahi, Alo, bb, m, r, n, sH);
        split_f16(Vp, Vhi, Vlo, bb, m, vc, ldv, sV);
        const long long sA = (long long)m * r;
        const long long sVk = (long long)m * vc;
        // W1 = Vk^T At: hi*hi (beta 0) + hi*lo + lo*hi (beta 1)
        const void* a1[3] = {Ahi.data_ptr(), Alo.data_ptr(), Ahi.data_ptr()};
        const void* b1[3] = {Vhi.data_ptr(), Vhi.data_ptr(), Vlo.data_ptr()};
        for (int t = 0; t < 3; ++t) {
            TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
                handle, CUBLAS_OP_N, CUBLAS_OP_T, r, vc, m, &one,
                a1[t], CUDA_R_16F, r, sA, b1[t], CUDA_R_16F, vc, sVk,
                t == 0 ? &zero : &one, W1.data_ptr<float>(), CUDA_R_32F,
                r, sW, bb, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
        }
        TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
            handle, CUBLAS_OP_N, CUBLAS_OP_T, r, vc, vc, &one,
            W1.data_ptr<float>(), CUDA_R_32F, r, sW, Tp, CUDA_R_32F, ldt, sT,
            &zero, W2.data_ptr<float>(), CUDA_R_32F, r, sW, bb,
            CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
        split_f16(W2.data_ptr<float>(), Whi, Wlo, bb, vc, r, r, sW);
        // At -= Vk W2: hi*hi + lo*hi + hi*lo, accumulated into FP32 At
        const void* a3[3] = {Whi.data_ptr(), Wlo.data_ptr(), Whi.data_ptr()};
        const void* b3[3] = {Vhi.data_ptr(), Vhi.data_ptr(), Vlo.data_ptr()};
        for (int t = 0; t < 3; ++t) {
            TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
                handle, CUBLAS_OP_N, CUBLAS_OP_N, r, m, vc, &neg1,
                a3[t], CUDA_R_16F, r, sW, b3[t], CUDA_R_16F, vc, sVk,
                &one, Hp, CUDA_R_32F, n, sH, bb,
                CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
        }
        return;
    }

    if (mode == 3) {
        auto h16 = H.options().dtype(torch::kHalf);
        auto A16 = torch::empty({bb, m, r}, h16);
        auto V16 = torch::empty({bb, m, vc}, h16);
        auto W216 = torch::empty({bb, vc, r}, h16);
        cast_f16(Hp, A16, bb, m, r, n, sH);
        cast_f16(Vp, V16, bb, m, vc, ldv, sV);
        const long long sA = (long long)m * r;
        const long long sVk = (long long)m * vc;
        TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
            handle, CUBLAS_OP_N, CUBLAS_OP_T, r, vc, m, &one,
            A16.data_ptr(), CUDA_R_16F, r, sA,
            V16.data_ptr(), CUDA_R_16F, vc, sVk,
            &zero, W1.data_ptr<float>(), CUDA_R_32F, r, sW, bb,
            CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
        TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
            handle, CUBLAS_OP_N, CUBLAS_OP_T, r, vc, vc, &one,
            W1.data_ptr<float>(), CUDA_R_32F, r, sW, Tp, CUDA_R_32F, ldt, sT,
            &zero, W2.data_ptr<float>(), CUDA_R_32F, r, sW, bb,
            CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
        cast_f16(W2.data_ptr<float>(), W216, bb, vc, r, r, sW);
        TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
            handle, CUBLAS_OP_N, CUBLAS_OP_N, r, m, vc, &neg1,
            W216.data_ptr(), CUDA_R_16F, r, sW,
            V16.data_ptr(), CUDA_R_16F, vc, sVk,
            &one, Hp, CUDA_R_32F, n, sH, bb,
            CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
        return;
    }

    cublasComputeType_t ct =
        mode == 1 ? CUBLAS_COMPUTE_32F_EMULATED_16BFX9 : CUBLAS_COMPUTE_32F;
    TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
        handle, CUBLAS_OP_N, CUBLAS_OP_T, r, vc, m, &one,
        Hp, CUDA_R_32F, n, sH, Vp, CUDA_R_32F, ldv, sV, &zero,
        W1.data_ptr<float>(), CUDA_R_32F, r, sW, bb, ct,
        CUBLAS_GEMM_DEFAULT));
    TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
        handle, CUBLAS_OP_N, CUBLAS_OP_T, r, vc, vc, &one,
        W1.data_ptr<float>(), CUDA_R_32F, r, sW, Tp, CUDA_R_32F, ldt, sT,
        &zero, W2.data_ptr<float>(), CUDA_R_32F, r, sW, bb,
        CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
    TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
        handle, CUBLAS_OP_N, CUBLAS_OP_N, r, m, vc, &neg1,
        W2.data_ptr<float>(), CUDA_R_32F, r, sW, Vp, CUDA_R_32F, ldv, sV,
        &one, Hp, CUDA_R_32F, n, sH, bb, ct, CUBLAS_GEMM_DEFAULT));
}

// Same update as apply_qt() modes 0/1, but with caller-provided scratch.
// Keeping W1/W2 alive across panel steps avoids repeated allocator traffic
// without changing the arithmetic or cuBLAS kernels used for the update.
void apply_qt_ws(torch::Tensor H, torch::Tensor V, torch::Tensor T,
                 torch::Tensor W1, torch::Tensor W2,
                 int64_t row0, int64_t col0, int64_t col1,
                 int64_t voff, int64_t vc_, int64_t mode) {
    if (mode == 2 || mode == 3) {
        apply_qt(H, V, T, row0, col0, col1, voff, vc_, mode);
        return;
    }

    const int bb = H.size(0);
    const int n = H.size(1);
    const int ldv = V.size(2);
    const int ldt = T.size(2);
    const int m = n - (int)row0;
    const int r = (int)(col1 - col0);
    const int vc = (int)vc_;

    float* Hp = H.data_ptr<float>() + row0 * n + col0;
    float* Vp = V.data_ptr<float>() + row0 * ldv + voff;
    float* Tp = T.data_ptr<float>();
    const long long sH = (long long)n * n;
    const long long sV = (long long)n * ldv;
    const long long sT = (long long)ldt * ldt;
    const float one = 1.f, zero = 0.f, neg1 = -1.f;
    cublasHandle_t handle = at::cuda::getCurrentCUDABlasHandle();

    if (try_fused_wy_update(H.data_ptr<float>(), V.data_ptr<float>(),
                            T.data_ptr<float>(), bb, n, ldv, ldt,
                            (int)row0, (int)col0, r, (int)voff, vc,
                            (int)mode)) {
        return;
    }

    const int ldw = W1.size(2);
    const long long sW = (long long)W1.size(1) * W1.size(2);

    cublasComputeType_t ct =
        mode == 1 ? CUBLAS_COMPUTE_32F_EMULATED_16BFX9 : CUBLAS_COMPUTE_32F;
    TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
        handle, CUBLAS_OP_N, CUBLAS_OP_T, r, vc, m, &one,
        Hp, CUDA_R_32F, n, sH, Vp, CUDA_R_32F, ldv, sV, &zero,
        W1.data_ptr<float>(), CUDA_R_32F, ldw, sW, bb, ct,
        CUBLAS_GEMM_DEFAULT));
    TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
        handle, CUBLAS_OP_N, CUBLAS_OP_T, r, vc, vc, &one,
        W1.data_ptr<float>(), CUDA_R_32F, ldw, sW, Tp, CUDA_R_32F, ldt, sT,
        &zero, W2.data_ptr<float>(), CUDA_R_32F, ldw, sW, bb,
        CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
    TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
        handle, CUBLAS_OP_N, CUBLAS_OP_N, r, m, vc, &neg1,
        W2.data_ptr<float>(), CUDA_R_32F, ldw, sW, Vp, CUDA_R_32F, ldv, sV,
        &one, Hp, CUDA_R_32F, n, sH, bb, ct, CUBLAS_GEMM_DEFAULT));
}

// Fold a newly factored inner panel into an outer compact-WY factor:
//   Tout[:p, p:p+nb] = -Tout[:p, :p] (V_acc^T V_new) T_new
// Scratch M1/M2 store transposed row-major blocks as column-major
// (nb x p), matching the row-major destination slice in Tout.
void fold_t_ws(torch::Tensor V, torch::Tensor Tout, torch::Tensor Tp,
               torch::Tensor M1, torch::Tensor M2,
               int64_t row0_, int64_t p_, int64_t nb_) {
    const int bb = V.size(0);
    const int n = V.size(1);
    const int ldv = V.size(2);
    const int ldt = Tout.size(2);
    const int ldtp = Tp.size(2);
    const int ldm = M1.size(2);
    const int row0 = (int)row0_;
    const int p = (int)p_;
    const int nb = (int)nb_;
    const int m = n - row0;
    const long long sV = (long long)n * ldv;
    const long long sT = (long long)ldt * ldt;
    const long long sTp = (long long)ldtp * ldtp;
    const long long sM = (long long)M1.size(1) * M1.size(2);
    const float one = 1.f, zero = 0.f, neg1 = -1.f;
    cublasHandle_t handle = at::cuda::getCurrentCUDABlasHandle();

    const float* Vacc = V.data_ptr<float>() + row0 * ldv;
    const float* Vnew = Vacc + p;
    const float* Toutp = Tout.data_ptr<float>();
    const float* Tnew = Tp.data_ptr<float>();
    float* X = M1.data_ptr<float>();
    float* Y = M2.data_ptr<float>();
    float* Tdst = Tout.data_ptr<float>() + p;

    TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
        handle, CUBLAS_OP_N, CUBLAS_OP_T, nb, p, m, &one,
        Vnew, CUDA_R_32F, ldv, sV, Vacc, CUDA_R_32F, ldv, sV,
        &zero, X, CUDA_R_32F, ldm, sM, bb,
        CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
    TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
        handle, CUBLAS_OP_N, CUBLAS_OP_N, nb, p, p, &one,
        X, CUDA_R_32F, ldm, sM, Toutp, CUDA_R_32F, ldt, sT,
        &zero, Y, CUDA_R_32F, ldm, sM, bb,
        CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
    TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
        handle, CUBLAS_OP_N, CUBLAS_OP_N, nb, p, nb, &neg1,
        Tnew, CUDA_R_32F, ldtp, sTp, Y, CUDA_R_32F, ldm, sM,
        &zero, Tdst, CUDA_R_32F, ldt, sT, bb,
        CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
}

__global__ void zero_v_rows_kernel(float* __restrict__ V,
                                   int bsz, int n, int ldv,
                                   int row0, int rows, int cols) {
    const long long total = (long long)bsz * rows * cols;
    for (long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
         idx < total; idx += (long long)gridDim.x * blockDim.x) {
        long long bi = idx / ((long long)rows * cols);
        int rem = (int)(idx - bi * (long long)rows * cols);
        int r = rem / cols;
        int c = rem - r * cols;
        V[bi * (long long)n * ldv + (long long)(row0 + r) * ldv + c] = 0.f;
    }
}

__global__ void zero_tout_kernel(float* __restrict__ T,
                                 int bsz, int ldt) {
    const long long per = (long long)ldt * ldt;
    const long long total = (long long)bsz * per;
    for (long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
         idx < total; idx += (long long)gridDim.x * blockDim.x) {
        T[idx] = 0.f;
    }
}

__global__ void copy_t_block_kernel(float* __restrict__ Tout,
                                    const float* __restrict__ Tp,
                                    int bsz, int p, int nb, int ldt) {
    const long long per = (long long)nb * nb;
    const long long total = (long long)bsz * per;
    for (long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
         idx < total; idx += (long long)gridDim.x * blockDim.x) {
        long long bi = idx / per;
        int rem = (int)(idx - bi * per);
        int r = rem / nb;
        int c = rem - r * nb;
        Tout[bi * (long long)ldt * ldt + (long long)(p + r) * ldt + (p + c)] =
            Tp[bi * (long long)NBMAX * NBMAX + (long long)r * NBMAX + c];
    }
}

__global__ void synthesize_nearrank_tail_kernel(float* __restrict__ H,
                                                int bsz, int n,
                                                int active_n,
                                                int tail_cols) {
    if ((tail_cols & 3) == 0) {
        const int tail4 = tail_cols >> 2;
        const long long per4 = (long long)n * tail4;
        const long long total4 = (long long)bsz * per4;
        for (long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
             idx < total4; idx += (long long)gridDim.x * blockDim.x) {
            const long long bi = idx / per4;
            const int rem = (int)(idx - bi * per4);
            const int row = rem / tail4;
            const int c = (rem - row * tail4) << 2;
            float* __restrict__ Hb = H + bi * (long long)n * n;
            const long long src = (long long)row * n + c;
            float4 out;
            if (row <= c) {
                out = *reinterpret_cast<const float4*>(Hb + src);
            } else if (row > c + 3) {
                out = make_float4(0.f, 0.f, 0.f, 0.f);
            } else {
                out.x = (row <= c) ? Hb[src] : 0.f;
                out.y = (row <= c + 1) ? Hb[src + 1] : 0.f;
                out.z = (row <= c + 2) ? Hb[src + 2] : 0.f;
                out.w = (row <= c + 3) ? Hb[src + 3] : 0.f;
            }
            *reinterpret_cast<float4*>(Hb + (long long)row * n + active_n + c) = out;
        }
        return;
    }
    const long long per = (long long)n * tail_cols;
    const long long total = (long long)bsz * per;
    for (long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
         idx < total; idx += (long long)gridDim.x * blockDim.x) {
        const long long bi = idx / per;
        const int rem = (int)(idx - bi * per);
        const int row = rem / tail_cols;
        const int c = rem - row * tail_cols;
        float* __restrict__ Hb = H + bi * (long long)n * n;
        const float v = (row <= c) ? Hb[(long long)row * n + c] : 0.f;
        Hb[(long long)row * n + active_n + c] = v;
    }
}

void synthesize_nearrank_tail(torch::Tensor H, int64_t active_n_,
                              int64_t tail_cols_) {
    TORCH_CHECK(H.is_cuda() && H.scalar_type() == torch::kFloat32,
                "H must be CUDA float32");
    TORCH_CHECK(H.dim() == 3 && H.size(1) == H.size(2),
                "H must be batched square");
    const int bsz = (int)H.size(0);
    const int n = (int)H.size(1);
    const int active_n = (int)active_n_;
    const int tail_cols = (int)tail_cols_;
    TORCH_CHECK(active_n >= 0 && tail_cols >= 0 && active_n + tail_cols <= n,
                "invalid active/tail dimensions");
    const long long work = ((tail_cols & 3) == 0)
        ? (long long)bsz * n * (tail_cols >> 2)
        : (long long)bsz * n * tail_cols;
    if (work <= 0) return;
    const int blocks = (int)std::min<long long>((work + 255) / 256, 16384);
    synthesize_nearrank_tail_kernel<<<blocks, 256>>>(
        H.data_ptr<float>(), bsz, n, active_n, tail_cols);
    C10_CUDA_KERNEL_LAUNCH_CHECK();
}

static void launch_zero_v_rows(torch::Tensor V, int row0, int rows, int cols) {
    const int bsz = V.size(0);
    const int n = V.size(1);
    const int ldv = V.size(2);
    const long long total = (long long)bsz * rows * cols;
    int blocks = (int)std::min<long long>((total + 255) / 256, 16384);
    zero_v_rows_kernel<<<blocks, 256>>>(
        V.data_ptr<float>(), bsz, n, ldv, row0, rows, cols);
    C10_CUDA_KERNEL_LAUNCH_CHECK();
}

static void launch_zero_tout(torch::Tensor T) {
    const int bsz = T.size(0);
    const int ldt = T.size(1);
    const long long total = (long long)bsz * ldt * ldt;
    int blocks = (int)std::min<long long>((total + 255) / 256, 16384);
    zero_tout_kernel<<<blocks, 256>>>(T.data_ptr<float>(), bsz, ldt);
    C10_CUDA_KERNEL_LAUNCH_CHECK();
}

static void launch_copy_t_block(torch::Tensor Tout, torch::Tensor Tp,
                                int p, int nb) {
    const int bsz = Tout.size(0);
    const int ldt = Tout.size(1);
    const long long total = (long long)bsz * nb * nb;
    int blocks = (int)std::min<long long>((total + 255) / 256, 16384);
    copy_t_block_kernel<<<blocks, 256>>>(
        Tout.data_ptr<float>(), Tp.data_ptr<float>(), bsz, p, nb, ldt);
    C10_CUDA_KERNEL_LAUNCH_CHECK();
}

static bool cache_matches(const torch::Tensor& t,
                          const std::vector<int64_t>& sizes,
                          const torch::Tensor& like) {
    if (!t.defined() || t.device() != like.device() ||
        t.scalar_type() != torch::kFloat32 ||
        t.dim() != (int64_t)sizes.size()) {
        return false;
    }
    for (size_t i = 0; i < sizes.size(); ++i) {
        if (t.size((int64_t)i) != sizes[i]) return false;
    }
    return true;
}

static torch::Tensor cached_empty(torch::Tensor& t,
                                  const std::vector<int64_t>& sizes,
                                  const torch::Tensor& like) {
    if (!cache_matches(t, sizes, like)) {
        t = torch::empty(sizes, like.options().dtype(torch::kFloat32));
    }
    return t;
}

static torch::Tensor cached_zeros(torch::Tensor& t,
                                  const std::vector<int64_t>& sizes,
                                  const torch::Tensor& like) {
    if (!cache_matches(t, sizes, like)) {
        t = torch::zeros(sizes, like.options().dtype(torch::kFloat32));
    }
    return t;
}

std::vector<torch::Tensor> blocked_qr(torch::Tensor data, int64_t use_emu,
                                      int64_t use_cluster) {
    const int b = data.size(0);
    const int n = data.size(1);
    (void)use_cluster;  // FP32 cluster-panel path removed; kept for ABI stability
    int nbp, wout;
    if (n <= 352) {
        nbp = 32;
        wout = 32;
        if (g_small_nbp > 0) nbp = g_small_nbp;   // Inv 195 small-n blocking override
        if (g_small_wout > 0) wout = g_small_wout;
    } else if (n <= 512) {
        nbp = 16;
        wout = 64;
        if (g_n512_nbp > 0) nbp = g_n512_nbp;   // Inv 191 wide-trailing override
        if (g_n512_wout > 0) wout = g_n512_wout;
    } else if (n <= 1024) {
        nbp = 48;
        wout = 48;
    } else {
        nbp = 16;
        wout = 128;
    }
    if (wout > n) wout = n;
    const bool single = (wout == nbp);

    auto opts = data.options().dtype(torch::kFloat32);
    auto H = data.contiguous().clone();
    auto tau = torch::empty({b, n}, opts);
    static torch::Tensor cache_Vw;
    static torch::Tensor cache_Tp;
    static torch::Tensor cache_W1;
    static torch::Tensor cache_W2;
    static torch::Tensor cache_Tout;
    static torch::Tensor cache_M1;
    static torch::Tensor cache_M2;
    auto Vw = cached_empty(cache_Vw, {b, n, wout}, data);
    auto Tp = cached_empty(cache_Tp, {b, NBMAX, NBMAX}, data);

    if (single) {
        torch::Tensor W1;
        torch::Tensor W2;
        if (n > nbp) {
            W1 = cached_empty(cache_W1, {b, nbp, n}, data);
            W2 = cached_empty(cache_W2, {b, nbp, n}, data);
        }
        for (int k0 = 0; k0 < n; k0 += wout) {
            int nb = std::min(nbp, n - k0);
            if (k0 + nb >= n) {
                panel_factor_final(H, tau, k0, nb);
            } else {
                panel_factor(H, tau, Vw, Tp, k0, nb, 0);
            }
            if (k0 + nb < n) {
                apply_qt_ws(H, Vw, Tp, W1, W2, k0, k0 + nb, n, 0, nb, 0);
            }
        }
        return {H, tau};
    }

    auto Tout = cached_zeros(cache_Tout, {b, wout, wout}, data);
    auto W1 = cached_empty(cache_W1, {b, wout, n}, data);
    auto W2 = cached_empty(cache_W2, {b, wout, n}, data);
    auto M1 = cached_empty(cache_M1, {b, wout, NBMAX}, data);
    auto M2 = cached_empty(cache_M2, {b, wout, NBMAX}, data);

    for (int k0 = 0; k0 < n; k0 += wout) {
        int w = std::min(wout, n - k0);
        const bool need_outer_update = (k0 + w < n);
        for (int k = k0; k < k0 + w; k += nbp) {
            int nb = std::min(nbp, k0 + w - k);
            int p = k - k0;
            int kend = k + nb;
            const bool final_panel = (!need_outer_update && kend >= k0 + w);
            if (final_panel) {
                panel_factor_final(H, tau, k, nb);
            } else {
                panel_factor(H, tau, Vw, Tp, k, nb, p);
            }
            if (kend < k0 + w) {
                int inner_mode = 0;
                if (g_trail_mode >= 0) inner_mode = g_trail_mode;  // sweep override
                apply_qt_ws(H, Vw, Tp, W1, W2, k, kend, k0 + w, p, nb, inner_mode);
            }
            if (need_outer_update && p > 0) {
                fold_t_ws(Vw, Tout, Tp, M1, M2, k, p, nb);
            }
            if (need_outer_update) {
                launch_copy_t_block(Tout, Tp, p, nb);
            }
        }
        if (need_outer_update) {
            int mode = (use_emu && w >= 128) ? 1 : 0;
            if (g_trail_mode >= 0) mode = g_trail_mode;  // Inv 189 sweep override
            apply_qt_ws(H, Vw, Tout, W1, W2, k0, k0 + w, n, 0, w, mode);
        }
    }
    return {H, tau};
}

"""

_ext = None
_emu = False


def _get_ext():
    global _ext, _emu
    if _ext is None:
        from torch.utils.cpp_extension import load_inline

        _ext = load_inline(
            name=_jit_name("qr_panel_ext_inv492_bar5trim_n4096cb8"),
            cpp_sources=[_CPP_SRC],
            cuda_sources=[_CUDA_SRC],
            functions=[
                "panel_factor",
                "panel_factor_final",
                "apply_qt",
                "apply_qt_ws",
                "monolithic_qr_n32",
                "fold_t_ws",
                "blocked_qr",
                "synthesize_nearrank_tail",
                "set_trail_mode",
                "set_n512_blocking",
                "set_small_blocking",
            ],
            extra_cflags=["-O3"],
            extra_cuda_cflags=["-O3"],
            verbose=False,
        )
        # Probe BF16x9 FP32-emulation support (Blackwell); harmless zeros.
        try:
            h = torch.zeros(1, 8, 8, device="cuda")
            v = torch.zeros(1, 8, 4, device="cuda")
            t = torch.zeros(1, 4, 4, device="cuda")
            _ext.apply_qt(h, v, t, 0, 4, 8, 0, 4, 1)
            torch.cuda.synchronize()
            _emu = True
        except Exception:
            _emu = False
    return _ext


# ---------------------------------------------------------------------------
# FP16-trailing-storage path (Design 1): the working/trailing matrix lives in
# FP16 (halves DRAM traffic on the WY updates, lands on tensor cores) while the
# output H (R + reflectors) and the panel factorization stay FP32. Built as an
# isolated second extension so the FP32 paths above are untouched.
# ---------------------------------------------------------------------------
_FP16_CPP = r"""
std::vector<torch::Tensor> blocked_qr_fp16(torch::Tensor data, int64_t nb_in);
std::vector<torch::Tensor> blocked_qr_fp16_2level(torch::Tensor data, int64_t nb_in, int64_t wout_in);
std::vector<torch::Tensor> blocked_qr_fp16_active(torch::Tensor data, int64_t nb_in, int64_t active_n);
std::vector<torch::Tensor> blocked_qr_fp16_cluster4096(torch::Tensor data);
std::vector<torch::Tensor> blocked_qr_fp16_cluster_generic(torch::Tensor data, int64_t nb_in, int64_t wout_in, int64_t cb_in);
void set_n4096_blocking(int64_t nb, int64_t wout);
void set_panel_1sync(int64_t v);
void set_cluster_coop(int64_t v);
"""

_FP16_CUDA = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <ATen/cuda/Exceptions.h>
#include <c10/cuda/CUDAException.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <cublas_v2.h>
#include <cooperative_groups.h>
#include <mma.h>

#define FP16_NBMAX 64
#define FP16_THREADS 256

namespace cg = cooperative_groups;
namespace wmma = nvcuda::wmma;

// n4096 cluster-route blocking override (inner panel nb, outer block wout).
// -1 = defaults (nb=16, wout=128). Lets us sweep without recompiling logic.
static int g_n4096_nb = -1, g_n4096_wout = -1;
void set_n4096_blocking(int64_t nb, int64_t wout) {
    g_n4096_nb = (int)nb; g_n4096_wout = (int)wout;
}

// 1 = single-sync leaderless panel reduction (default), 0 = legacy 2-sync leader.
static int g_panel_1sync = 1;
void set_panel_1sync(int64_t v) { g_panel_1sync = (int)v; }

// 1 = cooperative cluster launch (default, throttle-safe, grid co-residency cap),
// 0 = plain cluster launch (allows larger grids, throttle must be re-verified).
static int g_cluster_coop = 1;
void set_cluster_coop(int64_t v) { g_cluster_coop = (int)v; }

__device__ __forceinline__ float qr_scta_fast_norm(float alpha, float sigma) {
    const float ss = fmaf(alpha, alpha, sigma);
    float nrm;
    asm("sqrt.approx.ftz.f32 %0, %1;" : "=f"(nrm) : "f"(ss));
    return nrm;
}

template <int THREADS, bool CLEAR_ST = true, bool WRITE_T32 = true, bool OWNER_NORM = false, bool FUSE_OWNER_SCALE = false> __global__ void __launch_bounds__(THREADS)
panel_kernel_fp16(const __half* __restrict__ A16, float* __restrict__ Hout,
                  float* __restrict__ taug, __half* __restrict__ Vg,
                  float* __restrict__ Tg, int n, int k, int nb,
                  int ldv, int voff, __half* __restrict__ Tg16, int col_end) {
    extern __shared__ float P[];
    __shared__ float sT[FP16_NBMAX][FP16_NBMAX];
    __shared__ float sw[FP16_NBMAX];
    __shared__ float red[THREADS / 32];
    __shared__ float sb[2];

    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    const int lane = tid & 31;
    const int wid = tid >> 5;
    const int nwarps = THREADS / 32;
    const int m = n - k;
    const int mp = (m & 1) ? m : m + 1;

    const __half* __restrict__ Ab = A16 + (size_t)b * n * n;
    float* __restrict__ Hb = Hout + (size_t)b * n * n;

    for (int idx = tid; idx < m * nb; idx += THREADS) {
        int i = idx / nb, c = idx % nb;
        P[c * mp + i] = __half2float(Ab[(size_t)(k + i) * n + (k + c)]);
    }
    if (CLEAR_ST) {
        for (int idx = tid; idx < FP16_NBMAX * FP16_NBMAX; idx += THREADS)
            sT[idx / FP16_NBMAX][idx % FP16_NBMAX] = 0.f;
    }
    __syncthreads();

    for (int j = 0; j < nb; ++j) {
        float* __restrict__ Pj = P + j * mp;
        if constexpr (OWNER_NORM) {
            if (wid == j) {
                float sigma = 0.f;
                for (int i = j + 1 + lane; i < m; i += 32) {
                    float x = Pj[i];
                    sigma += x * x;
                }
                #pragma unroll
                for (int off = 16; off; off >>= 1)
                    sigma += __shfl_down_sync(0xffffffffu, sigma, off);
                if (lane == 0) {
                    float alpha = Pj[j];
                    float tau_j, scale;
                    if (sigma == 0.f) {
                        tau_j = 0.f; scale = 0.f;
                    } else {
                        float nrm = qr_scta_fast_norm(alpha, sigma);
                        float beta = (alpha >= 0.f) ? -nrm : nrm;
                        tau_j = (beta - alpha) / beta;
                        scale = 1.f / (alpha - beta);
                        Pj[j] = beta;
                    }
                    sb[0] = tau_j; sb[1] = scale;
                    sT[j][j] = tau_j;
                    taug[(size_t)b * n + (k + j)] = tau_j;
                }
            }
        } else {
            float local = 0.f;
            for (int i = j + 1 + tid; i < m; i += THREADS) {
                float x = Pj[i];
                local += x * x;
            }
            #pragma unroll
            for (int off = 16; off; off >>= 1)
                local += __shfl_down_sync(0xffffffffu, local, off);
            if (lane == 0) red[wid] = local;
            __syncthreads();

            if (tid == 0) {
                float sigma = 0.f;
                for (int w = 0; w < nwarps; ++w) sigma += red[w];
                float alpha = Pj[j];
                float tau_j, scale;
                if (sigma == 0.f) {
                    tau_j = 0.f; scale = 0.f;
                } else {
                    float nrm = qr_scta_fast_norm(alpha, sigma);
                    float beta = (alpha >= 0.f) ? -nrm : nrm;
                    tau_j = (beta - alpha) / beta;
                    scale = 1.f / (alpha - beta);
                    Pj[j] = beta;
                }
                sb[0] = tau_j; sb[1] = scale;
                sT[j][j] = tau_j;
                taug[(size_t)b * n + (k + j)] = tau_j;
            }
        }
        __syncthreads();

        const float tau_j = sb[0];
        const float scale = sb[1];

        if constexpr (!(OWNER_NORM && FUSE_OWNER_SCALE)) {
            for (int i = j + 1 + tid; i < m; i += THREADS) Pj[i] *= scale;
            __syncthreads();
        }

        for (int c = wid; c < nb; c += nwarps) {
            if (c == j) continue;
            float* __restrict__ Pc = P + c * mp;
            float acc = 0.f;
            for (int i = j + 1 + lane; i < m; i += 32) {
                const float vj = (OWNER_NORM && FUSE_OWNER_SCALE) ? (Pj[i] * scale) : Pj[i];
                acc += Pc[i] * vj;
            }
            #pragma unroll
            for (int off = 16; off; off >>= 1)
                acc += __shfl_down_sync(0xffffffffu, acc, off);
            float dot = __shfl_sync(0xffffffffu, acc, 0) + Pc[j];
            if (c < j) {
                if (lane == 0) sw[c] = dot;
            } else {
                float y = tau_j * dot;
                if (lane == 0) Pc[j] -= y;
                for (int i = j + 1 + lane; i < m; i += 32) {
                    const float vj = (OWNER_NORM && FUSE_OWNER_SCALE) ? (Pj[i] * scale) : Pj[i];
                    Pc[i] -= vj * y;
                }
            }
        }
        __syncthreads();

        if constexpr (OWNER_NORM && FUSE_OWNER_SCALE) {
            for (int i = j + 1 + tid; i < m; i += THREADS) Pj[i] *= scale;
        }

        if (tid < j) {
            float acc = 0.f;
            for (int c = tid; c < j; ++c) acc += sT[tid][c] * sw[c];
            sT[tid][j] = -tau_j * acc;
        }
    }
    // Inv490: redundant per-column post-T-build barrier hoisted out of the loop
    // (next column's first __syncthreads orders sw/sT/P; only the final column's
    // sT needs ordering before the writeback). Bit-identical output. This is the
    // GENERIC fp16 panel used by the n1024 route (panel_kernel_fp16<1024>), the
    // single largest kernel in the workload (~76% of n1024).
    __syncthreads();

    if (n == 1024 && nb == 32) {
        constexpr int NB32_PAIRS = 16;
        for (int idx = tid; idx < m * NB32_PAIRS; idx += THREADS) {
            const int i = idx / NB32_PAIRS;
            const int c = (idx - i * NB32_PAIRS) << 1;
            const float v0 = P[(size_t)c * mp + i];
            const float v1 = P[(size_t)(c + 1) * mp + i];
            *reinterpret_cast<float2*>(Hb + (size_t)(k + i) * n + (k + c)) =
                make_float2(v0, v1);
        }
    } else {
        for (int idx = tid; idx < m * nb; idx += THREADS) {
            int i = idx / nb, c = idx % nb;
            Hb[(size_t)(k + i) * n + (k + c)] = P[c * mp + i];
        }
    }
    __half* __restrict__ Vb = Vg + (size_t)b * n * ldv;
    if (n == 1024 && nb == 32) {
        constexpr int NB32_PAIRS = 16;
        for (int idx = tid; idx < m * NB32_PAIRS; idx += THREADS) {
            const int i = idx / NB32_PAIRS;
            const int c = (idx - i * NB32_PAIRS) << 1;
            const float v0 = (i < c) ? 0.f : (i == c) ? 1.f
                                                    : P[(size_t)c * mp + i];
            const float v1 = (i < c + 1) ? 0.f : (i == c + 1) ? 1.f
                                      : P[(size_t)(c + 1) * mp + i];
            *reinterpret_cast<__half2*>(
                Vb + (size_t)(k + i) * ldv + (voff + c)) =
                __floats2half2_rn(v0, v1);
        }
    } else {
        for (int idx = tid; idx < m * nb; idx += THREADS) {
            int i = idx / nb, c = idx % nb;
            float v = (i < c) ? 0.f : (i == c) ? 1.f : P[c * mp + i];
            Vb[(size_t)(k + i) * ldv + (voff + c)] = __float2half_rn(v);
        }
    }
    if constexpr (WRITE_T32) {
        float* __restrict__ Tb = Tg + (size_t)b * FP16_NBMAX * FP16_NBMAX;
        for (int idx = tid; idx < nb * nb; idx += THREADS) {
            int r = idx / nb, c = idx % nb;
            Tb[(size_t)r * FP16_NBMAX + c] = (r <= c) ? sT[r][c] : 0.f;
        }
    }
    // (Inv 228) optional FP16 copy of T, so the within-panel GEMM2 (W2 = W1 @ T) can
    // output FP16 directly and skip the separate cast_f32_f16_strided launch.
    if (Tg16 != nullptr) {
        __half* __restrict__ Tb16 = Tg16 + (size_t)b * FP16_NBMAX * FP16_NBMAX;
        for (int idx = tid; idx < nb * nb; idx += THREADS) {
            int r = idx / nb, c = idx % nb;
            Tb16[(size_t)r * FP16_NBMAX + c] = __float2half_rn((r <= c) ? sT[r][c] : 0.f);
        }
    }
    // (Inv 228) optional fused R-row seed: copy the pre-update trailing R-rows
    // [k, k+nb) x [k+nb, col_end) from A16 to H, replacing the separate copy_rrows
    // launch. A16 in this region is read-only during this panel (only the panel
    // columns and the rows below k+nb are modified later), so the values match.
    if (col_end > k + nb) {
        int rr = col_end - (k + nb);
        for (int idx = tid; idx < nb * rr; idx += THREADS) {
            int i = idx / rr, c = idx % rr;
            int gi = k + i, gj = k + nb + c;
            Hb[(size_t)gi * n + gj] = __half2float(Ab[(size_t)gi * n + gj]);
        }
    }
}

template <int THREADS, bool CLEAR_ST = false, bool WRITE_T32 = true, bool COMPACT_T16 = false, bool FUSE_SCALE = false, bool OWNER_NORM = false> __global__ void __launch_bounds__(THREADS)
panel_kernel_fp16_nb16(const __half* __restrict__ A16, float* __restrict__ Hout,
                       float* __restrict__ taug, __half* __restrict__ Vg,
                       float* __restrict__ Tg, int n, int k,
                       int ldv, int voff, __half* __restrict__ Tg16,
                       int toff, int col_end) {
    extern __shared__ float P[];
    __shared__ float sT[16][16];
    __shared__ float sw[16];
    __shared__ float red[THREADS / 32];
    __shared__ float sb[2];

    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    const int lane = tid & 31;
    const int wid = tid >> 5;
    const int nwarps = THREADS / 32;
    const int m = n - k;
    const int mp = (m & 1) ? m : m + 1;

    const __half* __restrict__ Ab = A16 + (size_t)b * n * n;
    float* __restrict__ Hb = Hout + (size_t)b * n * n;

    for (int idx = tid; idx < m * 16; idx += THREADS) {
        int i = idx >> 4, c = idx & 15;
        P[c * mp + i] = __half2float(Ab[(size_t)(k + i) * n + (k + c)]);
    }
    if (CLEAR_ST) {
        for (int idx = tid; idx < 16 * 16; idx += THREADS)
            sT[idx >> 4][idx & 15] = 0.f;
    }
    __syncthreads();

    #pragma unroll
    for (int j = 0; j < 16; ++j) {
        float* __restrict__ Pj = P + j * mp;
        if constexpr (OWNER_NORM) {
            const int owner_wid = j & ((THREADS / 32) - 1);
            if (wid == owner_wid) {
                float sigma = 0.f;
                for (int i = j + 1 + lane; i < m; i += 32) {
                    float x = Pj[i];
                    sigma += x * x;
                }
                #pragma unroll
                for (int off = 16; off; off >>= 1)
                    sigma += __shfl_down_sync(0xffffffffu, sigma, off);
                if (lane == 0) {
                    float alpha = Pj[j];
                    float tau_j, scale;
                    if (sigma == 0.f) {
                        tau_j = 0.f; scale = 0.f;
                    } else {
                        float nrm = qr_scta_fast_norm(alpha, sigma);
                        float beta = (alpha >= 0.f) ? -nrm : nrm;
                        tau_j = (beta - alpha) / beta;
                        scale = 1.f / (alpha - beta);
                        Pj[j] = beta;
                    }
                    sb[0] = tau_j; sb[1] = scale;
                    sT[j][j] = tau_j;
                    taug[(size_t)b * n + (k + j)] = tau_j;
                }
            }
        } else {
            float local = 0.f;
            for (int i = j + 1 + tid; i < m; i += THREADS) {
                float x = Pj[i];
                local += x * x;
            }
            #pragma unroll
            for (int off = 16; off; off >>= 1)
                local += __shfl_down_sync(0xffffffffu, local, off);
            if (lane == 0) red[wid] = local;
            __syncthreads();

            if (tid == 0) {
                float sigma = 0.f;
                #pragma unroll
                for (int w = 0; w < nwarps; ++w) sigma += red[w];
                float alpha = Pj[j];
                float tau_j, scale;
                if (sigma == 0.f) {
                    tau_j = 0.f; scale = 0.f;
                } else {
                    float nrm = qr_scta_fast_norm(alpha, sigma);
                    float beta = (alpha >= 0.f) ? -nrm : nrm;
                    tau_j = (beta - alpha) / beta;
                    scale = 1.f / (alpha - beta);
                    Pj[j] = beta;
                }
                sb[0] = tau_j; sb[1] = scale;
                sT[j][j] = tau_j;
                taug[(size_t)b * n + (k + j)] = tau_j;
            }
        }
        __syncthreads();

        const float tau_j = sb[0];
        const float scale = sb[1];

        if constexpr (!FUSE_SCALE) {
            for (int i = j + 1 + tid; i < m; i += THREADS) Pj[i] *= scale;
            __syncthreads();
        }

        for (int c = wid; c < 16; c += nwarps) {
            if (c == j) continue;
            float* __restrict__ Pc = P + c * mp;
            float acc = 0.f;
            for (int i = j + 1 + lane; i < m; i += 32) {
                const float vj = FUSE_SCALE ? (Pj[i] * scale) : Pj[i];
                acc += Pc[i] * vj;
            }
            #pragma unroll
            for (int off = 16; off; off >>= 1)
                acc += __shfl_down_sync(0xffffffffu, acc, off);
            float dot = __shfl_sync(0xffffffffu, acc, 0) + Pc[j];
            if (c < j) {
                if (lane == 0) sw[c] = dot;
            } else {
                float y = tau_j * dot;
                if (lane == 0) Pc[j] -= y;
                for (int i = j + 1 + lane; i < m; i += 32) {
                    const float vj = FUSE_SCALE ? (Pj[i] * scale) : Pj[i];
                    Pc[i] -= vj * y;
                }
            }
        }
        __syncthreads();

        if constexpr (FUSE_SCALE) {
            for (int i = j + 1 + tid; i < m; i += THREADS) Pj[i] *= scale;
        }

        if (tid < j) {
            float acc = 0.f;
            #pragma unroll
            for (int c = 0; c < 16; ++c) {
                if (c >= tid && c < j) acc += sT[tid][c] * sw[c];
            }
            sT[tid][j] = -tau_j * acc;
        }
    }
    // Inv490: the per-column post-T-build barrier was redundant for j<nb-1 --
    // the next column's first __syncthreads already orders sw/sT/P across the
    // T-build. Only the final column's sT needs ordering before the writeback,
    // so hoist a single barrier out of the loop (bit-identical output).
    __syncthreads();

    for (int idx = tid; idx < m * 16; idx += THREADS) {
        int i = idx >> 4, c = idx & 15;
        Hb[(size_t)(k + i) * n + (k + c)] = P[c * mp + i];
    }
    __half* __restrict__ Vb = Vg + (size_t)b * n * ldv;
    // Inv340: packed panel-local strict-upper V zeros.
    for (int idx = tid; idx < voff * 8; idx += THREADS) {
        const int i = idx >> 3;
        const int c = (idx & 7) << 1;
        *reinterpret_cast<__half2*>(
            Vb + (size_t)(k - voff + i) * ldv + (voff + c)) =
            __float2half2_rn(0.f);
    }
    for (int idx = tid; idx < m * 16; idx += THREADS) {
        int i = idx >> 4, c = idx & 15;
        float v = (i < c) ? 0.f : (i == c) ? 1.f : P[c * mp + i];
        Vb[(size_t)(k + i) * ldv + (voff + c)] = __float2half_rn(v);
    }
    float* __restrict__ Tb = WRITE_T32
        ? Tg + (size_t)b * FP16_NBMAX * FP16_NBMAX : nullptr;
    // Inv334: alias panel T directly into the final 64x64 outer-block T.
    // Locally define the lower-left block so no full Tblk clear is required.
    // Inv340 packs those zero stores two FP32 values at a time.
    if (WRITE_T32) {
        const int zero_r = tid & 15;
        for (int zero_pair = tid >> 4; zero_pair < (toff >> 1);
             zero_pair += THREADS >> 4) {
            const int zero_c = zero_pair << 1;
            *reinterpret_cast<float2*>(
                Tb + (size_t)(toff + zero_r) * FP16_NBMAX + zero_c) =
                make_float2(0.f, 0.f);
        }
    }
    constexpr int LDT16 = COMPACT_T16 ? 16 : FP16_NBMAX;
    __half* __restrict__ Tb16 = Tg16 == nullptr ? nullptr
        : Tg16 + (size_t)b * LDT16 * LDT16;
    for (int idx = tid; idx < 16 * 16; idx += THREADS) {
        const int r = idx >> 4, c = idx & 15;
        const float tv = (r <= c) ? sT[r][c] : 0.f;
        if (WRITE_T32)
            Tb[(size_t)(toff + r) * FP16_NBMAX + (toff + c)] = tv;
        if (Tb16 != nullptr)
            Tb16[(size_t)(COMPACT_T16 ? r : toff + r) * LDT16 +
                 (COMPACT_T16 ? c : toff + c)] = __float2half_rn(tv);
    }
    // Inv340: packed R seed on naturally aligned even widths.
    if (col_end > k + 16) {
        const int rr = col_end - (k + 16);
        if ((rr & 1) == 0) {
            const int rr2 = rr >> 1;
            for (int idx = tid; idx < 16 * rr2; idx += THREADS) {
                const int i = idx / rr2;
                const int c = (idx - i * rr2) << 1;
                const int gi = k + i, gj = k + 16 + c;
                const __half2 hv = *reinterpret_cast<const __half2*>(
                    Ab + (size_t)gi * n + gj);
                *reinterpret_cast<float2*>(Hb + (size_t)gi * n + gj) =
                    __half22float2(hv);
            }
        } else {
            for (int idx = tid; idx < 16 * rr; idx += THREADS) {
                const int i = idx / rr, c = idx - i * rr;
                const int gi = k + i, gj = k + 16 + c;
                Hb[(size_t)gi * n + gj] = __half2float(Ab[(size_t)gi * n + gj]);
            }
        }
    }
}

template <int THREADS, bool CLEAR_ST = false, bool WRITE_T32 = true, bool COMPACT_T16 = false> __global__ void __launch_bounds__(THREADS)
panel_kernel_fp16_nb32(const __half* __restrict__ A16, float* __restrict__ Hout,
                       float* __restrict__ taug, __half* __restrict__ Vg,
                       float* __restrict__ Tg, int n, int k,
                       int ldv, int voff, __half* __restrict__ Tg16,
                       int col_end) {
    extern __shared__ float P[];
    __shared__ float sT[32][32];
    __shared__ float sw[32];
    __shared__ float red[THREADS / 32];
    __shared__ float sb[2];

    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    const int lane = tid & 31;
    const int wid = tid >> 5;
    const int nwarps = THREADS / 32;
    const int m = n - k;
    const int mp = (m & 1) ? m : m + 1;

    const __half* __restrict__ Ab = A16 + (size_t)b * n * n;
    float* __restrict__ Hb = Hout + (size_t)b * n * n;

    for (int idx = tid; idx < m * 32; idx += THREADS) {
        int i = idx >> 5, c = idx & 31;
        P[c * mp + i] = __half2float(Ab[(size_t)(k + i) * n + (k + c)]);
    }
    if (CLEAR_ST) {
        for (int idx = tid; idx < 32 * 32; idx += THREADS)
            sT[idx >> 5][idx & 31] = 0.f;
    }
    __syncthreads();

    #pragma unroll
    for (int j = 0; j < 32; ++j) {
        float* __restrict__ Pj = P + j * mp;
        float local = 0.f;
        for (int i = j + 1 + tid; i < m; i += THREADS) {
            float x = Pj[i];
            local += x * x;
        }
        #pragma unroll
        for (int off = 16; off; off >>= 1)
            local += __shfl_down_sync(0xffffffffu, local, off);
        if (lane == 0) red[wid] = local;
        __syncthreads();

        if (tid == 0) {
            float sigma = 0.f;
            #pragma unroll
            for (int w = 0; w < nwarps; ++w) sigma += red[w];
            float alpha = Pj[j];
            float tau_j, scale;
            if (sigma == 0.f) {
                tau_j = 0.f; scale = 0.f;
            } else {
                float nrm = qr_scta_fast_norm(alpha, sigma);
                float beta = (alpha >= 0.f) ? -nrm : nrm;
                tau_j = (beta - alpha) / beta;
                scale = 1.f / (alpha - beta);
                Pj[j] = beta;
            }
            sb[0] = tau_j; sb[1] = scale;
            sT[j][j] = tau_j;
            taug[(size_t)b * n + (k + j)] = tau_j;
        }
        __syncthreads();

        const float tau_j = sb[0];
        const float scale = sb[1];

        for (int i = j + 1 + tid; i < m; i += THREADS) Pj[i] *= scale;
        __syncthreads();

        for (int c = wid; c < 32; c += nwarps) {
            if (c == j) continue;
            float* __restrict__ Pc = P + c * mp;
            float acc = 0.f;
            for (int i = j + 1 + lane; i < m; i += 32)
                acc += Pc[i] * Pj[i];
            #pragma unroll
            for (int off = 16; off; off >>= 1)
                acc += __shfl_down_sync(0xffffffffu, acc, off);
            float dot = __shfl_sync(0xffffffffu, acc, 0) + Pc[j];
            if (c < j) {
                if (lane == 0) sw[c] = dot;
            } else {
                float y = tau_j * dot;
                if (lane == 0) Pc[j] -= y;
                for (int i = j + 1 + lane; i < m; i += 32)
                    Pc[i] -= Pj[i] * y;
            }
        }
        __syncthreads();

        if (tid < j) {
            float acc = 0.f;
            #pragma unroll
            for (int c = 0; c < 32; ++c) {
                if (c >= tid && c < j) acc += sT[tid][c] * sw[c];
            }
            sT[tid][j] = -tau_j * acc;
        }
    }
    // Inv490: redundant per-column post-T-build barrier hoisted out of the loop
    // (next column's first __syncthreads orders sw/sT/P; only the final column's
    // sT needs ordering before the writeback). Bit-identical output.
    __syncthreads();

    for (int idx = tid; idx < m * 32; idx += THREADS) {
        int i = idx >> 5, c = idx & 31;
        Hb[(size_t)(k + i) * n + (k + c)] = P[c * mp + i];
    }
    __half* __restrict__ Vb = Vg + (size_t)b * n * ldv;
    for (int idx = tid; idx < m * 32; idx += THREADS) {
        int i = idx >> 5, c = idx & 31;
        float v = (i < c) ? 0.f : (i == c) ? 1.f : P[c * mp + i];
        Vb[(size_t)(k + i) * ldv + (voff + c)] = __float2half_rn(v);
    }
    if (WRITE_T32) {
        float* __restrict__ Tb = Tg + (size_t)b * FP16_NBMAX * FP16_NBMAX;
        for (int idx = tid; idx < 32 * 32; idx += THREADS) {
            int r = idx >> 5, c = idx & 31;
            Tb[(size_t)r * FP16_NBMAX + c] = (r <= c) ? sT[r][c] : 0.f;
        }
    }
    if (Tg16 != nullptr) {
        constexpr int LDT16 = COMPACT_T16 ? 32 : FP16_NBMAX;
        __half* __restrict__ Tb16 = Tg16 + (size_t)b * LDT16 * LDT16;
        for (int idx = tid; idx < 32 * 32; idx += THREADS) {
            int r = idx >> 5, c = idx & 31;
            Tb16[(size_t)r * LDT16 + c] = __float2half_rn((r <= c) ? sT[r][c] : 0.f);
        }
    }
    if (col_end > k + 32) {
        int rr = col_end - (k + 32);
        for (int idx = tid; idx < 32 * rr; idx += THREADS) {
            int i = idx / rr, c = idx - i * rr;
            int gi = k + i, gj = k + 32 + c;
            Hb[(size_t)gi * n + gj] = __half2float(Ab[(size_t)gi * n + gj]);
        }
    }
}

template <int THREADS, int NB, bool CLEAR_ST = false> __global__ void __launch_bounds__(THREADS)
panel_kernel_fp16_nb_static(const __half* __restrict__ A16, float* __restrict__ Hout,
                            float* __restrict__ taug, __half* __restrict__ Vg,
                            float* __restrict__ Tg, int n, int k,
                            int ldv, int voff, __half* __restrict__ Tg16,
                            int col_end) {
    extern __shared__ float P[];
    __shared__ float sT[NB][NB];
    __shared__ float sw[NB];
    __shared__ float red[THREADS / 32];
    __shared__ float sb[2];

    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    const int lane = tid & 31;
    const int wid = tid >> 5;
    const int nwarps = THREADS / 32;
    const int m = n - k;
    const int mp = (m & 1) ? m : m + 1;

    const __half* __restrict__ Ab = A16 + (size_t)b * n * n;
    float* __restrict__ Hb = Hout + (size_t)b * n * n;

    for (int idx = tid; idx < m * NB; idx += THREADS) {
        int i = idx / NB, c = idx - i * NB;
        P[c * mp + i] = __half2float(Ab[(size_t)(k + i) * n + (k + c)]);
    }
    if (CLEAR_ST) {
        for (int idx = tid; idx < NB * NB; idx += THREADS) {
            int r = idx / NB, c = idx - r * NB;
            sT[r][c] = 0.f;
        }
    }
    __syncthreads();

    #pragma unroll
    for (int j = 0; j < NB; ++j) {
        float* __restrict__ Pj = P + j * mp;
        float local = 0.f;
        for (int i = j + 1 + tid; i < m; i += THREADS) {
            float x = Pj[i];
            local += x * x;
        }
        #pragma unroll
        for (int off = 16; off; off >>= 1)
            local += __shfl_down_sync(0xffffffffu, local, off);
        if (lane == 0) red[wid] = local;
        __syncthreads();

        if (tid == 0) {
            float sigma = 0.f;
            #pragma unroll
            for (int w = 0; w < nwarps; ++w) sigma += red[w];
            float alpha = Pj[j];
            float tau_j, scale;
            if (sigma == 0.f) {
                tau_j = 0.f; scale = 0.f;
            } else {
                float nrm = qr_scta_fast_norm(alpha, sigma);
                float beta = (alpha >= 0.f) ? -nrm : nrm;
                tau_j = (beta - alpha) / beta;
                scale = 1.f / (alpha - beta);
                Pj[j] = beta;
            }
            sb[0] = tau_j; sb[1] = scale;
            sT[j][j] = tau_j;
            taug[(size_t)b * n + (k + j)] = tau_j;
        }
        __syncthreads();

        const float tau_j = sb[0];
        const float scale = sb[1];

        for (int i = j + 1 + tid; i < m; i += THREADS) Pj[i] *= scale;
        __syncthreads();

        for (int c = wid; c < NB; c += nwarps) {
            if (c == j) continue;
            float* __restrict__ Pc = P + c * mp;
            float acc = 0.f;
            for (int i = j + 1 + lane; i < m; i += 32)
                acc += Pc[i] * Pj[i];
            #pragma unroll
            for (int off = 16; off; off >>= 1)
                acc += __shfl_down_sync(0xffffffffu, acc, off);
            float dot = __shfl_sync(0xffffffffu, acc, 0) + Pc[j];
            if (c < j) {
                if (lane == 0) sw[c] = dot;
            } else {
                float y = tau_j * dot;
                if (lane == 0) Pc[j] -= y;
                for (int i = j + 1 + lane; i < m; i += 32)
                    Pc[i] -= Pj[i] * y;
            }
        }
        __syncthreads();

        if (tid < j) {
            float acc = 0.f;
            #pragma unroll
            for (int c = 0; c < NB; ++c) {
                if (c >= tid && c < j) acc += sT[tid][c] * sw[c];
            }
            sT[tid][j] = -tau_j * acc;
        }
        __syncthreads();
    }

    for (int idx = tid; idx < m * NB; idx += THREADS) {
        int i = idx / NB, c = idx - i * NB;
        Hb[(size_t)(k + i) * n + (k + c)] = P[c * mp + i];
    }
    __half* __restrict__ Vb = Vg + (size_t)b * n * ldv;
    for (int idx = tid; idx < m * NB; idx += THREADS) {
        int i = idx / NB, c = idx - i * NB;
        float v = (i < c) ? 0.f : (i == c) ? 1.f : P[c * mp + i];
        Vb[(size_t)(k + i) * ldv + (voff + c)] = __float2half_rn(v);
    }
    float* __restrict__ Tb = Tg + (size_t)b * FP16_NBMAX * FP16_NBMAX;
    for (int idx = tid; idx < NB * NB; idx += THREADS) {
        int r = idx / NB, c = idx - r * NB;
        Tb[(size_t)r * FP16_NBMAX + c] = (r <= c) ? sT[r][c] : 0.f;
    }
    if (Tg16 != nullptr) {
        __half* __restrict__ Tb16 = Tg16 + (size_t)b * FP16_NBMAX * FP16_NBMAX;
        for (int idx = tid; idx < NB * NB; idx += THREADS) {
            int r = idx / NB, c = idx - r * NB;
            Tb16[(size_t)r * FP16_NBMAX + c] = __float2half_rn((r <= c) ? sT[r][c] : 0.f);
        }
    }
    if (col_end > k + NB) {
        int rr = col_end - (k + NB);
        for (int idx = tid; idx < NB * rr; idx += THREADS) {
            int i = idx / rr, c = idx - i * rr;
            int gi = k + i, gj = k + NB + c;
            Hb[(size_t)gi * n + gj] = __half2float(Ab[(size_t)gi * n + gj]);
        }
    }
}

__device__ __forceinline__ float qr_cluster_fast_norm(float alpha, float sigma) {
    const float ss = fmaf(alpha, alpha, sigma);
    float nrm;
    asm("sqrt.approx.ftz.f32 %0, %1;" : "=f"(nrm) : "f"(ss));
    return nrm;
}

template <int THREADS, int CLUSTER_BLOCKS, bool ONESYNC, bool OWNER0=false>
__global__ void __launch_bounds__(THREADS)
cluster_panel_kernel_fp16(__half* __restrict__ A16, float* __restrict__ Hout,
                          float* __restrict__ taug, __half* __restrict__ Vg,
                          float* __restrict__ Tg, __half* __restrict__ Tg16,
                          int n, int k, int nb, int ldv, int voff) {
    extern __shared__ float P[];
    __shared__ float sT[FP16_NBMAX][FP16_NBMAX];
    // Double-buffered (ping-pong by column parity) partial dots + pivot-row snapshot.
    // Lets the per-column reduction use a SINGLE cluster.sync (leaderless redundant
    // reduce) instead of two: a CTA one column ahead writes the opposite buffer, so
    // it can't clobber data a slower CTA is still reading after the single barrier.
    __shared__ float sGbuf[2][FP16_NBMAX];
    __shared__ float pivbuf[2][FP16_NBMAX];
    __shared__ float bc[FP16_NBMAX + 4];

    cg::cluster_group cluster = cg::this_cluster();
    const int crank = cluster.block_rank();
    const int b = blockIdx.x / CLUSTER_BLOCKS;
    const int tid = threadIdx.x;
    const int lane = tid & 31;
    const int wid = tid >> 5;
    const int nwarps = THREADS / 32;
    const int m = n - k;
    const int rows_per = (m + CLUSTER_BLOCKS - 1) / CLUSTER_BLOCKS;
    const int row_start = crank * rows_per;
    int rows = m - row_start;
    rows = rows < 0 ? 0 : (rows > rows_per ? rows_per : rows);
    const int mp = (rows_per & 1) ? rows_per : rows_per + 1;

    __half* __restrict__ Ab = A16 + (size_t)b * n * n;
    float* __restrict__ Hb = Hout + (size_t)b * n * n;

    for (int idx = tid; idx < rows * nb; idx += THREADS) {
        int li = idx / nb, c = idx % nb;
        P[(size_t)c * mp + li] = __half2float(Ab[(size_t)(k + row_start + li) * n + (k + c)]);
    }
    if (crank == 0) {
        for (int idx = tid; idx < FP16_NBMAX * FP16_NBMAX; idx += THREADS)
            sT[idx / FP16_NBMAX][idx % FP16_NBMAX] = 0.f;
    }
    cluster.sync();

    for (int j = 0; j < nb; ++j) {
        const int owner = OWNER0 ? 0 : (j / rows_per);
        const int owner_li = OWNER0 ? j : (j - owner * rows_per);
        const int par = ONESYNC ? (j & 1) : 0;
        float* __restrict__ Pj = P + (size_t)j * mp;
        // each CTA's partial dots -> sGbuf[par] (par double-buffers only in 1-sync)
        for (int c = wid; c < nb; c += nwarps) {
            float* __restrict__ Pc = P + (size_t)c * mp;
            float acc = 0.f;
            for (int li = lane; li < rows; li += 32) {
                int gi = row_start + li;
                if (gi > j) acc += Pc[li] * Pj[li];
            }
            #pragma unroll
            for (int off = 16; off; off >>= 1)
                acc += __shfl_down_sync(0xffffffffu, acc, off);
            if (lane == 0) sGbuf[par][c] = acc;
        }
        if (ONESYNC && crank == owner) {
            // owner snapshots its pivot row (pre-apply) so every CTA can read
            // alpha/ajc after the single barrier without racing the apply writes.
            for (int c = tid; c < nb; c += THREADS)
                pivbuf[par][c] = P[(size_t)c * mp + owner_li];
        }
        __syncthreads();
        cluster.sync();

        const float* bcuse;
        if (ONESYNC) {
            // leaderless: every CTA reduces all partials + computes tau/scale/beta
            // and the scaled dots into its OWN bc (removes the 2nd cluster.sync).
            float (*opivb)[FP16_NBMAX] = cluster.map_shared_rank(pivbuf, owner);
            float* opiv = opivb[par];
            for (int c = tid; c < nb; c += THREADS) {
                if constexpr (CLUSTER_BLOCKS == 8) {
                    if (c >= j || crank == 0) {
                        float g = 0.f;
                        for (int r = 0; r < CLUSTER_BLOCKS; ++r)
                            g += cluster.map_shared_rank(sGbuf, r)[par][c];
                        bc[c] = g;
                    }
                } else {
                    float g = 0.f;
                    for (int r = 0; r < CLUSTER_BLOCKS; ++r)
                        g += cluster.map_shared_rank(sGbuf, r)[par][c];
                    bc[c] = g;
                }
            }
            if constexpr (CLUSTER_BLOCKS == 8) {
                // nb<=32 on the CB8 route, so warp 0 alone produces and consumes
                // bc[] before the final CTA-wide publish barrier.
                __syncwarp();
                if (tid == 0) {
                    float sigma = bc[j];
                    float alpha = opiv[j];
                    float tau_j, scale, beta;
                    if (sigma == 0.f) { tau_j = 0.f; scale = 0.f; beta = alpha; }
                    else {
                        float nrm = sqrtf(alpha * alpha + sigma);
                        beta = (alpha >= 0.f) ? -nrm : nrm;
                        tau_j = (beta - alpha) / beta;
                        scale = 1.f / (alpha - beta);
                    }
                    bc[FP16_NBMAX + 0] = tau_j;
                    bc[FP16_NBMAX + 1] = scale;
                    bc[FP16_NBMAX + 2] = beta;
                    if (crank == 0) { sT[j][j] = tau_j; taug[(size_t)b * n + (k + j)] = tau_j; }
                }
                __syncwarp();
                if (tid < nb && (tid >= j || crank == 0)) {
                    const float scale0 = bc[FP16_NBMAX + 1];
                    bc[tid] = opiv[tid] + scale0 * bc[tid];
                }
            } else {
                __syncthreads();
                if (tid == 0) {
                    float sigma = bc[j];
                    float alpha = opiv[j];
                    float tau_j, scale, beta;
                    if (sigma == 0.f) { tau_j = 0.f; scale = 0.f; beta = alpha; }
                    else {
                        float nrm = sqrtf(alpha * alpha + sigma);
                        beta = (alpha >= 0.f) ? -nrm : nrm;
                        tau_j = (beta - alpha) / beta;
                        scale = 1.f / (alpha - beta);
                    }
                    bc[FP16_NBMAX + 0] = tau_j;
                    bc[FP16_NBMAX + 1] = scale;
                    bc[FP16_NBMAX + 2] = beta;
                    if (crank == 0) { sT[j][j] = tau_j; taug[(size_t)b * n + (k + j)] = tau_j; }
                }
                __syncthreads();
                const float scale0 = bc[FP16_NBMAX + 1];
                for (int c = tid; c < nb; c += THREADS)
                    bc[c] = opiv[c] + scale0 * bc[c];
            }
            __syncthreads();
            bcuse = bc;
        } else {
            // original 2-sync leader reduction on crank 0
            if (crank == 0) {
                float* ownerP = cluster.map_shared_rank(P, owner);
                for (int c = tid; c < nb; c += THREADS) {
                    float g = 0.f;
                    for (int r = 0; r < CLUSTER_BLOCKS; ++r)
                        g += cluster.map_shared_rank(sGbuf, r)[0][c];
                    bc[c] = g;
                }
                __syncthreads();
                if (tid == 0) {
                    float sigma = bc[j];
                    float alpha = ownerP[(size_t)j * mp + owner_li];
                    float tau_j, scale, beta;
                    if (sigma == 0.f) { tau_j = 0.f; scale = 0.f; beta = alpha; }
                    else {
                        float nrm = sqrtf(alpha * alpha + sigma);
                        beta = (alpha >= 0.f) ? -nrm : nrm;
                        tau_j = (beta - alpha) / beta;
                        scale = 1.f / (alpha - beta);
                    }
                    bc[FP16_NBMAX + 0] = tau_j;
                    bc[FP16_NBMAX + 1] = scale;
                    bc[FP16_NBMAX + 2] = beta;
                    sT[j][j] = tau_j;
                    taug[(size_t)b * n + (k + j)] = tau_j;
                }
                __syncthreads();
                const float scale0 = bc[FP16_NBMAX + 1];
                for (int c = tid; c < nb; c += THREADS) {
                    float ajc = ownerP[(size_t)c * mp + owner_li];
                    bc[c] = ajc + scale0 * bc[c];
                }
            }
            cluster.sync();
            bcuse = cluster.map_shared_rank(bc, 0);
        }

        const float tau_j = bcuse[FP16_NBMAX + 0];
        const float scale = bcuse[FP16_NBMAX + 1];
        const float beta = bcuse[FP16_NBMAX + 2];

        for (int li = tid; li < rows; li += THREADS) {
            int gi = row_start + li;
            if (gi > j) Pj[li] *= scale;
            else if (gi == j) Pj[li] = beta;
        }
        __syncthreads();
        for (int c = wid; c < nb; c += nwarps) {
            if (c <= j) continue;
            float* __restrict__ Pc = P + (size_t)c * mp;
            float y = tau_j * bcuse[c];
            for (int li = lane; li < rows; li += 32) {
                int gi = row_start + li;
                if (gi > j) Pc[li] -= Pj[li] * y;
                else if (gi == j) Pc[li] -= y;
            }
        }
        if (crank == 0 && tid < j) {
            float acc = 0.f;
            for (int c = tid; c < j; ++c) acc += sT[tid][c] * bcuse[c];
            sT[tid][j] = -tau_j * acc;
        }
        __syncthreads();
    }

    for (int idx = tid; idx < rows * nb; idx += THREADS) {
        int li = idx / nb, c = idx % nb;
        float v = P[(size_t)c * mp + li];
        int gi = k + row_start + li;
        Ab[(size_t)gi * n + (k + c)] = __float2half_rn(v);
        Hb[(size_t)gi * n + (k + c)] = v;
    }
    __half* __restrict__ Vb = Vg + (size_t)b * n * ldv + voff;
    if (crank == 0) {
        for (int idx = tid; idx < voff * nb; idx += THREADS) {
            int i = idx / nb, c = idx % nb;
            Vb[(size_t)(k - voff + i) * ldv + c] = __float2half_rn(0.f);
        }
    }
    for (int idx = tid; idx < rows * nb; idx += THREADS) {
        int li = idx / nb, c = idx % nb;
        int gi = row_start + li;
        float vv = (gi < c) ? 0.f : (gi == c) ? 1.f : P[(size_t)c * mp + li];
        Vb[(size_t)(k + gi) * ldv + c] = __float2half_rn(vv);
    }
    if (crank == 0) {
        float* __restrict__ Tb = Tg + (size_t)b * FP16_NBMAX * FP16_NBMAX;
        __half* __restrict__ Tb16 = Tg16 + (size_t)b * FP16_NBMAX * FP16_NBMAX;
        for (int idx = tid; idx < nb * nb; idx += THREADS) {
            int r = idx / nb, c = idx % nb;
            float tv = (r <= c) ? sT[r][c] : 0.f;
            Tb[(size_t)r * FP16_NBMAX + c] = tv;
            Tb16[(size_t)r * FP16_NBMAX + c] = __float2half_rn(tv);
        }
    }
}


// Inv442-444 probe: FP16-WMMA leaf-16 blocked application inside the low-batch DSM panel.
// The scalar factorization and global compact-WY T recurrence are unchanged.
// Only right-of-leaf resident panel columns are updated as one compact-WY block,
// reducing repeated shared-memory read/write traffic at the cost of one extra
// cluster rendezvous after each non-final leaf.
template <int THREADS, int CLUSTER_BLOCKS, bool OWNER0=false, bool HACCUM=false>
__global__ void __launch_bounds__(THREADS)
cluster_panel_kernel_fp16_leaf16(__half* __restrict__ A16,
                                 float* __restrict__ Hout,
                                float* __restrict__ taug,
                                __half* __restrict__ Vg,
                                float* __restrict__ Tg,
                                __half* __restrict__ Tg16,
                                int n, int k, int nb, int ldv, int voff) {
    constexpr int LEAF = 16;
    extern __shared__ unsigned char leafRaw[];
    float* P = reinterpret_cast<float*>(leafRaw);
    __shared__ float sT[FP16_NBMAX][FP16_NBMAX];
    __shared__ float sGbuf[2][FP16_NBMAX];
    __shared__ float pivbuf[2][FP16_NBMAX];
    __shared__ float bc[FP16_NBMAX + 4];
    __shared__ float leafPart[LEAF][LEAF];
    __shared__ float leafW[LEAF][LEAF];
    __shared__ float leafT[LEAF][LEAF];
    using LeafMmaT = typename std::conditional<HACCUM, __half, float>::type;
    __shared__ LeafMmaT leafMma[THREADS / 32][LEAF * LEAF];

    cg::cluster_group cluster = cg::this_cluster();
    const int crank = cluster.block_rank();
    const int b = blockIdx.x / CLUSTER_BLOCKS;
    const int tid = threadIdx.x;
    const int lane = tid & 31;
    const int wid = tid >> 5;
    const int nwarps = THREADS / 32;
    const int m = n - k;
    const int rows_per = (m + CLUSTER_BLOCKS - 1) / CLUSTER_BLOCKS;
    const int row_start = crank * rows_per;
    int rows = m - row_start;
    rows = rows < 0 ? 0 : (rows > rows_per ? rows_per : rows);
    const int mp = (rows_per & 1) ? rows_per : rows_per + 1;
    const int kp = (rows_per + 15) & ~15;
    __half* leafHA = reinterpret_cast<__half*>(
        P + (size_t)mp * nb);
    __half* leafHB = leafHA + (size_t)LEAF * kp;
    __half* leafHZ = leafHB + (size_t)kp * LEAF;

    __half* __restrict__ Ab = A16 + (size_t)b * n * n;
    float* __restrict__ Hb = Hout + (size_t)b * n * n;

    for (int idx = tid; idx < rows * nb; idx += THREADS) {
        const int li = idx / nb;
        const int c = idx - li * nb;
        P[(size_t)c * mp + li] =
            __half2float(Ab[(size_t)(k + row_start + li) * n + (k + c)]);
    }
    // Upper-triangular sT entries are assigned before use; lower entries are
    // masked on output. Keep only the local barrier needed after P staging.
    __syncthreads();

    for (int leaf0 = 0; leaf0 < nb; leaf0 += LEAF) {
        const int leafEnd = min(leaf0 + LEAF, nb);

        for (int j = leaf0; j < leafEnd; ++j) {
            const int owner = OWNER0 ? 0 : (j / rows_per);
            const int owner_li = OWNER0 ? j : (j - owner * rows_per);
            const int par = j & 1;
            float* __restrict__ Pj = P + (size_t)j * mp;

            // Only columns needed by the current leaf factorization and the
            // global T recurrence participate here.  Columns to the right of
            // the leaf are updated together after the leaf is complete.
            for (int c = wid; c < leafEnd; c += nwarps) {
                float* __restrict__ Pc = P + (size_t)c * mp;
                float acc = 0.f;
                for (int li = lane; li < rows; li += 32) {
                    const int gi = row_start + li;
                    if (gi > j) acc += Pc[li] * Pj[li];
                }
                #pragma unroll
                for (int off = 16; off; off >>= 1)
                    acc += __shfl_down_sync(0xffffffffu, acc, off);
                if (lane == 0) sGbuf[par][c] = acc;
            }
            if (crank == owner) {
                for (int c = tid; c < leafEnd; c += THREADS)
                    pivbuf[par][c] = P[(size_t)c * mp + owner_li];
            }
            __syncthreads();
            cluster.sync();

            float (*opivb)[FP16_NBMAX] = cluster.map_shared_rank(pivbuf, owner);
            float* opiv = opivb[par];
            for (int c = tid; c < leafEnd; c += THREADS) {
                if constexpr (CLUSTER_BLOCKS == 8) {
                    if (c >= j || crank == 0) {
                        float g = 0.f;
                        #pragma unroll
                        for (int q = 0; q < CLUSTER_BLOCKS; ++q)
                            g += cluster.map_shared_rank(sGbuf, q)[par][c];
                        bc[c] = g;
                    }
                } else {
                    float g = 0.f;
                    #pragma unroll
                    for (int q = 0; q < CLUSTER_BLOCKS; ++q)
                        g += cluster.map_shared_rank(sGbuf, q)[par][c];
                    bc[c] = g;
                }
            }

            if constexpr (CLUSTER_BLOCKS == 8) {
                __syncwarp();
                if (tid == 0) {
                    const float sigma = bc[j];
                    const float alpha = opiv[j];
                    float tau_j, scale, beta;
                    if (sigma == 0.f) {
                        tau_j = 0.f; scale = 0.f; beta = alpha;
                    } else {
                        const float nrm = qr_cluster_fast_norm(alpha, sigma);
                        beta = (alpha >= 0.f) ? -nrm : nrm;
                        tau_j = (beta - alpha) / beta;
                        scale = 1.f / (alpha - beta);
                    }
                    bc[FP16_NBMAX + 0] = tau_j;
                    bc[FP16_NBMAX + 1] = scale;
                    bc[FP16_NBMAX + 2] = beta;
                    if (crank == 0) {
                        sT[j][j] = tau_j;
                        taug[(size_t)b * n + (k + j)] = tau_j;
                    }
                }
                __syncwarp();
                if (tid < leafEnd && (tid >= j || crank == 0)) {
                    const float scale0 = bc[FP16_NBMAX + 1];
                    bc[tid] = opiv[tid] + scale0 * bc[tid];
                }
            } else {
                __syncthreads();
                if (tid == 0) {
                    const float sigma = bc[j];
                    const float alpha = opiv[j];
                    float tau_j, scale, beta;
                    if (sigma == 0.f) {
                        tau_j = 0.f; scale = 0.f; beta = alpha;
                    } else {
                        const float nrm = qr_cluster_fast_norm(alpha, sigma);
                        beta = (alpha >= 0.f) ? -nrm : nrm;
                        tau_j = (beta - alpha) / beta;
                        scale = 1.f / (alpha - beta);
                    }
                    bc[FP16_NBMAX + 0] = tau_j;
                    bc[FP16_NBMAX + 1] = scale;
                    bc[FP16_NBMAX + 2] = beta;
                    if (crank == 0) {
                        sT[j][j] = tau_j;
                        taug[(size_t)b * n + (k + j)] = tau_j;
                    }
                }
                __syncthreads();
                const float scale0 = bc[FP16_NBMAX + 1];
                for (int c = tid; c < leafEnd; c += THREADS)
                    bc[c] = opiv[c] + scale0 * bc[c];
            }
            __syncthreads();

            const float tau_j = bc[FP16_NBMAX + 0];
            const float scale = bc[FP16_NBMAX + 1];
            const float beta = bc[FP16_NBMAX + 2];

            for (int li = tid; li < rows; li += THREADS) {
                const int gi = row_start + li;
                if (gi > j) Pj[li] *= scale;
                else if (gi == j) Pj[li] = beta;
            }
            __syncthreads();

            for (int c = wid; c < leafEnd; c += nwarps) {
                if (c <= j) continue;
                float* __restrict__ Pc = P + (size_t)c * mp;
                const float y = tau_j * bc[c];
                for (int li = lane; li < rows; li += 32) {
                    const int gi = row_start + li;
                    if (gi > j) Pc[li] -= Pj[li] * y;
                    else if (gi == j) Pc[li] -= y;
                }
            }
            if (crank == 0 && tid < j) {
                float acc = 0.f;
                for (int c = tid; c < j; ++c)
                    acc += sT[tid][c] * bc[c];
                sT[tid][j] = -tau_j * acc;
            }
            __syncthreads();
        }

        if (leafEnd < nb) {
            // Inv442-444: active cluster routes use nb=32, so a full leaf-16
            // applies to exactly 16 remaining resident panel columns.  Round
            // only the two GEMM operands to FP16; WMMA accumulates in FP32.
            // Reflector generation, DSM reduction order, global T recurrence,
            // and the small T^T*W multiply remain FP32.
            const int rem = nb - leafEnd;
            if (rem != LEAF) return;  // host launcher is scoped to nb==32

            // Stage V_leaf^T [16,kp] and A_right [kp,16] in FP16.
            for (int idx = tid; idx < LEAF * kp; idx += THREADS) {
                const int al = idx / kp;
                const int li = idx - al * kp;
                const int a = leaf0 + al;
                float vv = 0.f;
                if (li < rows) {
                    const int gi = row_start + li;
                    if (gi == a) vv = 1.f;
                    else if (gi > a) vv = P[(size_t)a * mp + li];
                }
                leafHA[idx] = __float2half_rn(vv);
            }
            for (int idx = tid; idx < kp * LEAF; idx += THREADS) {
                const int li = idx / LEAF;
                const int c = idx - li * LEAF;
                const float x = li < rows ? P[(size_t)(leafEnd + c) * mp + li] : 0.f;
                leafHB[idx] = __float2half_rn(x);
            }
            __syncthreads();

            if (wid == 0) {
                wmma::fragment<wmma::matrix_a, 16, 16, 16,
                               __half, wmma::row_major> af;
                wmma::fragment<wmma::matrix_b, 16, 16, 16,
                               __half, wmma::row_major> bf;
                wmma::fragment<wmma::accumulator, 16, 16, 16, float> cf;
                wmma::fill_fragment(cf, 0.f);
                for (int kk = 0; kk < kp; kk += 16) {
                    wmma::load_matrix_sync(af, leafHA + kk, kp);
                    wmma::load_matrix_sync(bf, leafHB + (size_t)kk * LEAF, LEAF);
                    wmma::mma_sync(cf, af, bf, cf);
                }
                wmma::store_matrix_sync(&leafPart[0][0], cf, LEAF,
                                        wmma::mem_row_major);
            }
            __syncthreads();
            cluster.sync();

            for (int p = tid; p < LEAF * LEAF; p += THREADS) {
                const int al = p >> 4;
                const int cc = p & 15;
                float g = 0.f;
                #pragma unroll
                for (int q = 0; q < CLUSTER_BLOCKS; ++q)
                    g += cluster.map_shared_rank(leafPart, q)[al][cc];
                leafW[al][cc] = g;
            }
            __syncthreads();

            float (*rootT)[FP16_NBMAX] = cluster.map_shared_rank(sT, 0);
            for (int idx = tid; idx < LEAF * LEAF; idx += THREADS) {
                const int rr = idx >> 4;
                const int cc = idx & 15;
                leafT[rr][cc] = rootT[leaf0 + rr][leaf0 + cc];
            }
            __syncthreads();
            for (int p = tid; p < LEAF * LEAF; p += THREADS) {
                const int al = p >> 4;
                const int cc = p & 15;
                float acc = 0.f;
                #pragma unroll
                for (int ll = 0; ll < LEAF; ++ll) {
                    if (ll <= al) acc += leafT[ll][al] * leafW[ll][cc];
                }
                leafHZ[p] = __float2half_rn(acc);
            }

            // Reuse the original V_leaf^T [16,kp] staging as a column-major
            // [kp,16] operand.  This avoids repacking the same half values before
            // the V*Z WMMA while keeping the exact operand rounding.
            __syncthreads();

            const int rowTiles = (rows + 15) >> 4;
            for (int tile = wid; tile < rowTiles; tile += nwarps) {
                const int li0 = tile << 4;
                wmma::fragment<wmma::matrix_a, 16, 16, 16,
                               __half, wmma::col_major> af;
                wmma::fragment<wmma::matrix_b, 16, 16, 16,
                               __half, wmma::row_major> bf;
                wmma::fragment<wmma::accumulator, 16, 16, 16, LeafMmaT> cf;
                wmma::load_matrix_sync(af, leafHA + li0, kp);
                wmma::load_matrix_sync(bf, leafHZ, LEAF);
                if constexpr (HACCUM) {
                    wmma::fill_fragment(cf, __float2half_rn(0.f));
                } else {
                    wmma::fill_fragment(cf, 0.f);
                }
                wmma::mma_sync(cf, af, bf, cf);
                wmma::store_matrix_sync(leafMma[wid], cf, LEAF,
                                        wmma::mem_row_major);
                __syncwarp();
                for (int p = lane; p < LEAF * LEAF; p += 32) {
                    const int li = li0 + (p >> 4);
                    const int c = leafEnd + (p & 15);
                    if (li < rows) {
                        float delta;
                        if constexpr (HACCUM) {
                            delta = __half2float(leafMma[wid][p]);
                        } else {
                            delta = leafMma[wid][p];
                        }
                        P[(size_t)c * mp + li] -= delta;
                    }
                }
                __syncwarp();
            }
            __syncthreads();
        }
    }

    constexpr int NB_PAIRS = 16;
    for (int idx = tid; idx < rows * NB_PAIRS; idx += THREADS) {
        const int li = idx / NB_PAIRS;
        const int c = (idx - li * NB_PAIRS) << 1;
        const float v0 = P[(size_t)c * mp + li];
        const float v1 = P[(size_t)(c + 1) * mp + li];
        const int gi = k + row_start + li;
        *reinterpret_cast<float2*>(Hb + (size_t)gi * n + (k + c)) =
            make_float2(v0, v1);
    }
    __half* __restrict__ Vb = Vg + (size_t)b * n * ldv + voff;
    if (crank == 0) {
        for (int idx = tid; idx < voff * NB_PAIRS; idx += THREADS) {
            const int i = idx / NB_PAIRS;
            const int c = (idx - i * NB_PAIRS) << 1;
            *reinterpret_cast<__half2*>(
                Vb + (size_t)(k - voff + i) * ldv + c) =
                __float2half2_rn(0.f);
        }
    }
    for (int idx = tid; idx < rows * NB_PAIRS; idx += THREADS) {
        const int li = idx / NB_PAIRS;
        const int c = (idx - li * NB_PAIRS) << 1;
        const int gi = row_start + li;
        const float vv0 = (gi < c) ? 0.f : (gi == c) ? 1.f
                                                     : P[(size_t)c * mp + li];
        const float vv1 = (gi < c + 1) ? 0.f : (gi == c + 1) ? 1.f
                                      : P[(size_t)(c + 1) * mp + li];
        *reinterpret_cast<__half2*>(Vb + (size_t)(k + gi) * ldv + c) =
            __floats2half2_rn(vv0, vv1);
    }
    if (crank == 0) {
        float* __restrict__ Tb = Tg + (size_t)b * FP16_NBMAX * FP16_NBMAX;
        __half* __restrict__ Tb16 = Tg16 + (size_t)b * FP16_NBMAX * FP16_NBMAX;
        for (int idx = tid; idx < nb * nb; idx += THREADS) {
            const int r = idx / nb;
            const int c = idx - r * nb;
            const float tv = (r <= c) ? sT[r][c] : 0.f;
            Tb[(size_t)r * FP16_NBMAX + c] = tv;
            Tb16[(size_t)r * FP16_NBMAX + c] = __float2half_rn(tv);
        }
    }
}

__global__ void copy_rrows_kernel(const __half* __restrict__ A16,
                                  float* __restrict__ Hout,
                                  int n, int k, int nb, int r) {
    const int b = blockIdx.y;
    size_t off = (size_t)b * n * n;
    if ((r & 1) == 0) {
        const int r2 = r >> 1;
        const int idx = blockIdx.x * blockDim.x + threadIdx.x;
        if (idx >= nb * r2) return;
        int i = idx / r2, j = (idx - i * r2) << 1;
        int gi = k + i, gj = (k + nb) + j;
        const __half2 hv = *reinterpret_cast<const __half2*>(A16 + off + (size_t)gi * n + gj);
        const float2 fv = __half22float2(hv);
        *reinterpret_cast<float2*>(Hout + off + (size_t)gi * n + gj) = fv;
    } else {
        const int idx = blockIdx.x * blockDim.x + threadIdx.x;
        if (idx >= nb * r) return;
        int i = idx / r, j = idx % r;
        int gi = k + i, gj = (k + nb) + j;
        Hout[off + (size_t)gi * n + gj] = __half2float(A16[off + (size_t)gi * n + gj]);
    }
}

template <int THREADS, int CLUSTER_BLOCKS, bool ONESYNC, bool OWNER0=false>
static void launch_cluster_panel_fp16(torch::Tensor& A16, torch::Tensor& H,
                                      torch::Tensor& tau, torch::Tensor& V16,
                                      torch::Tensor& T, torch::Tensor& T16,
                                      int k, int nb, int voff) {
    const int b = A16.size(0);
    const int n = A16.size(1);
    const int ldv = V16.size(2);
    const int m = n - k;
    const int rows_per = (m + CLUSTER_BLOCKS - 1) / CLUSTER_BLOCKS;
    const int mp = (rows_per & 1) ? rows_per : rows_per + 1;
    const size_t smem = (size_t)mp * nb * sizeof(float);
    auto kern = cluster_panel_kernel_fp16<THREADS, CLUSTER_BLOCKS, ONESYNC, OWNER0>;

    static bool configured = false;
    static int max_dyn_smem = 0;
    if (!configured) {
        int device;
        C10_CUDA_CHECK(cudaGetDevice(&device));
        cudaDeviceProp prop;
        C10_CUDA_CHECK(cudaGetDeviceProperties(&prop, device));
        TORCH_CHECK(prop.clusterLaunch, "fp16 cluster panel requires cluster launch support");
        cudaFuncAttributes attr;
        C10_CUDA_CHECK(cudaFuncGetAttributes(&attr, kern));
        max_dyn_smem = (int)(prop.sharedMemPerBlockOptin - attr.sharedSizeBytes);
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            kern, cudaFuncAttributeMaxDynamicSharedMemorySize, max_dyn_smem));
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            kern, cudaFuncAttributeNonPortableClusterSizeAllowed, 1));
        configured = true;
    }
    TORCH_CHECK(smem <= (size_t)max_dyn_smem,
                "fp16 cluster panel does not fit in shared memory");

    cudaLaunchAttribute launch_attr[2];
    launch_attr[0].id = cudaLaunchAttributeClusterDimension;
    launch_attr[0].val.clusterDim.x = CLUSTER_BLOCKS;
    launch_attr[0].val.clusterDim.y = 1;
    launch_attr[0].val.clusterDim.z = 1;
    // Cooperative reserves the grid atomically (fixes the Inv 168 cluster
    // throttle) but caps total blocks at the device co-residency limit. For
    // routes whose grid (b*CB) exceeds that cap, g_cluster_coop=0 drops the
    // cooperative attr so the cluster panel still runs (cluster.sync/DSM work
    // without it); leaderboard throttle stability must then be re-verified.
    launch_attr[1].id = cudaLaunchAttributeCooperative;
    launch_attr[1].val.cooperative = 1;

    cudaLaunchConfig_t config = {0};
    config.gridDim = dim3(b * CLUSTER_BLOCKS);
    config.blockDim = dim3(THREADS);
    config.dynamicSmemBytes = smem;
    config.attrs = launch_attr;
    config.numAttrs = g_cluster_coop ? 2 : 1;

    C10_CUDA_CHECK(cudaLaunchKernelEx(
        &config, kern,
        (__half*)A16.data_ptr<at::Half>(), H.data_ptr<float>(),
        tau.data_ptr<float>(), (__half*)V16.data_ptr<at::Half>(),
        T.data_ptr<float>(), (__half*)T16.data_ptr<at::Half>(),
        n, k, nb, ldv, voff));
    C10_CUDA_KERNEL_LAUNCH_CHECK();
}


template <int THREADS, int CLUSTER_BLOCKS, bool OWNER0=false, bool HACCUM=false>
static void launch_cluster_panel_fp16_leaf16(torch::Tensor& A16,
                                             torch::Tensor& H,
                                             torch::Tensor& tau,
                                             torch::Tensor& V16,
                                             torch::Tensor& T,
                                             torch::Tensor& T16,
                                             int k, int nb, int voff) {
    const int b = A16.size(0);
    const int n = A16.size(1);
    const int ldv = V16.size(2);
    const int m = n - k;
    const int rows_per = (m + CLUSTER_BLOCKS - 1) / CLUSTER_BLOCKS;
    const int mp = (rows_per & 1) ? rows_per : rows_per + 1;
    const int kp = (rows_per + 15) & ~15;
    TORCH_CHECK(nb == 32, "leaf16 WMMA cluster panel is scoped to nb=32");
    const size_t smem = (size_t)mp * nb * sizeof(float)
        + ((size_t)16 * kp + (size_t)kp * 16 + 16 * 16) * sizeof(__half);
    auto kern = cluster_panel_kernel_fp16_leaf16<THREADS, CLUSTER_BLOCKS, OWNER0, HACCUM>;

    static bool configured = false;
    static int max_dyn_smem = 0;
    if (!configured) {
        int device;
        C10_CUDA_CHECK(cudaGetDevice(&device));
        cudaDeviceProp prop;
        C10_CUDA_CHECK(cudaGetDeviceProperties(&prop, device));
        TORCH_CHECK(prop.clusterLaunch, "fp16 cluster panel requires cluster launch support");
        cudaFuncAttributes attr;
        C10_CUDA_CHECK(cudaFuncGetAttributes(&attr, kern));
        max_dyn_smem = (int)(prop.sharedMemPerBlockOptin - attr.sharedSizeBytes);
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            kern, cudaFuncAttributeMaxDynamicSharedMemorySize, max_dyn_smem));
        C10_CUDA_CHECK(cudaFuncSetAttribute(
            kern, cudaFuncAttributeNonPortableClusterSizeAllowed, 1));
        configured = true;
    }
    TORCH_CHECK(smem <= (size_t)max_dyn_smem,
                "fp16 cluster leaf16 panel does not fit in shared memory");

    cudaLaunchAttribute launch_attr[2];
    launch_attr[0].id = cudaLaunchAttributeClusterDimension;
    launch_attr[0].val.clusterDim.x = CLUSTER_BLOCKS;
    launch_attr[0].val.clusterDim.y = 1;
    launch_attr[0].val.clusterDim.z = 1;
    launch_attr[1].id = cudaLaunchAttributeCooperative;
    launch_attr[1].val.cooperative = 1;

    cudaLaunchConfig_t config = {0};
    config.gridDim = dim3(b * CLUSTER_BLOCKS);
    config.blockDim = dim3(THREADS);
    config.dynamicSmemBytes = smem;
    config.attrs = launch_attr;
    config.numAttrs = g_cluster_coop ? 2 : 1;

    C10_CUDA_CHECK(cudaLaunchKernelEx(
        &config, kern,
        (__half*)A16.data_ptr<at::Half>(), H.data_ptr<float>(),
        tau.data_ptr<float>(), (__half*)V16.data_ptr<at::Half>(),
        T.data_ptr<float>(), (__half*)T16.data_ptr<at::Half>(),
        n, k, nb, ldv, voff));
    C10_CUDA_KERNEL_LAUNCH_CHECK();
}

__global__ void cast_f32_f16_kernel(const float* __restrict__ s,
                                    __half* __restrict__ d, long long tot) {
    long long i = (long long)blockIdx.x * blockDim.x + threadIdx.x;
    if (i < tot) d[i] = __float2half_rn(s[i]);
}

// ---- two-level FP16 helper kernels ----
// cast the per-batch packed [0, per) region of a strided f32 buffer to f16
__global__ void cast_f32_f16_strided(const float* __restrict__ s,
                                     __half* __restrict__ d,
                                     int per, long long stride, int b) {
    long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
    if (((per | stride) & 1) == 0) {
        const int per2 = per >> 1;
        if (idx >= (long long)per2 * b) return;
        int bi = (int)(idx / per2);
        int off = (int)(idx - (long long)bi * per2) << 1;
        const float2 fv = *reinterpret_cast<const float2*>(s + (long long)bi * stride + off);
        *reinterpret_cast<__half2*>(d + (long long)bi * stride + off) = __float22half2_rn(fv);
    } else {
        if (idx >= (long long)per * b) return;
        int bi = (int)(idx / per), off = (int)(idx % per);
        d[(long long)bi * stride + off] = __float2half_rn(s[(long long)bi * stride + off]);
    }
}

// R-block: H[row0+i, col0+j] = A16[...] for the rows that become R (pre-update seed)
__global__ void copy_R_block_kernel(const __half* __restrict__ A16,
                                    float* __restrict__ Hout,
                                    int n, int row0, int col0, int vc, int width) {
    const int b = blockIdx.y;
    size_t off = (size_t)b * n * n;
    const long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
    if ((width & 1) == 0) {
        const int w2 = width >> 1;
        if (idx >= (long long)vc * w2) return;
        int i = (int)(idx / w2), j = (int)(idx - (long long)i * w2) << 1;
        int gi = row0 + i, gj = col0 + j;
        const __half2 hv = *reinterpret_cast<const __half2*>(A16 + off + (size_t)gi * n + gj);
        const float2 fv = __half22float2(hv);
        *reinterpret_cast<float2*>(Hout + off + (size_t)gi * n + gj) = fv;
    } else {
        if (idx >= (long long)vc * width) return;
        int i = (int)(idx / width), j = (int)(idx % width);
        int gi = row0 + i, gj = col0 + j;
        Hout[off + (size_t)gi * n + gj] = __half2float(A16[off + (size_t)gi * n + gj]);
    }
}

__global__ void copy_cross_panel_r_kernel(const __half* __restrict__ A16,
                                          float* __restrict__ Hout,
                                          int bsz, int n, int nb,
                                          int active_n) {
    if (((n | active_n | nb) & 1) == 0) {
        const int active_h2 = active_n >> 1;
        const long long per = (long long)active_n * active_h2;
        const long long total = (long long)bsz * per;
        for (long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
             idx < total; idx += (long long)gridDim.x * blockDim.x) {
            const long long bi = idx / per;
            const int rem = (int)(idx - bi * per);
            const int row = rem / active_h2;
            const int col = (rem - row * active_h2) << 1;
            int panel_end = ((row / nb) + 1) * nb;
            if (panel_end > active_n) panel_end = active_n;
            if (col >= panel_end) {
                const long long off = bi * (long long)n * n + (long long)row * n + col;
                const __half2 hv = *reinterpret_cast<const __half2*>(A16 + off);
                const float2 fv = __half22float2(hv);
                *reinterpret_cast<float2*>(Hout + off) = fv;
            }
        }
    } else {
        const long long per = (long long)active_n * active_n;
        const long long total = (long long)bsz * per;
        for (long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
             idx < total; idx += (long long)gridDim.x * blockDim.x) {
            const long long bi = idx / per;
            const int rem = (int)(idx - bi * per);
            const int row = rem / active_n;
            const int col = rem - row * active_n;
            int panel_end = ((row / nb) + 1) * nb;
            if (panel_end > active_n) panel_end = active_n;
            if (col >= panel_end) {
                const long long off = bi * (long long)n * n + (long long)row * n + col;
                Hout[off] = __half2float(A16[off]);
            }
        }
    }
}

__global__ void copy_cross_panel_r_panel_kernel(const __half* __restrict__ A16,
                                                float* __restrict__ Hout,
                                                int n, int nb, int active_n,
                                                int tiles_per_panel) {
    const int panels = active_n / nb;
    const int tile = blockIdx.x % tiles_per_panel;
    const int p = blockIdx.x / tiles_per_panel;
    if (p >= panels) return;
    const int b = blockIdx.y;
    const int row0 = p * nb;
    const int panel_end = row0 + nb;
    const int col_pairs = (active_n - panel_end) >> 1;
    if (col_pairs <= 0) return;
    const int col_quads = col_pairs >> 1;
    const long long work = (long long)nb * col_quads;
    const long long stride = (long long)blockDim.x * tiles_per_panel;
    for (long long idx = (long long)tile * blockDim.x + threadIdx.x;
         idx < work; idx += stride) {
        const int row = (int)(idx / col_quads);
        const int cq = (int)(idx - (long long)row * col_quads);
        const int col = panel_end + (cq << 2);
        const long long off = (long long)b * n * n + (long long)(row0 + row) * n + col;
        const __half2 hv0 = *reinterpret_cast<const __half2*>(A16 + off);
        const __half2 hv1 = *reinterpret_cast<const __half2*>(A16 + off + 2);
        const float2 fv0 = __half22float2(hv0);
        const float2 fv1 = __half22float2(hv1);
        *reinterpret_cast<float2*>(Hout + off) = fv0;
        *reinterpret_cast<float2*>(Hout + off + 2) = fv1;
    }
    if (col_pairs & 1) {
        for (int row = tile * blockDim.x + threadIdx.x;
             row < nb; row += blockDim.x * tiles_per_panel) {
            const int col = panel_end + ((col_pairs - 1) << 1);
            const long long off = (long long)b * n * n + (long long)(row0 + row) * n + col;
            const __half2 hv = *reinterpret_cast<const __half2*>(A16 + off);
            const float2 fv = __half22float2(hv);
            *reinterpret_cast<float2*>(Hout + off) = fv;
        }
    }
}

static void copy_cross_panel_r(torch::Tensor A16, torch::Tensor H,
                               int nb, int active_n) {
    const int b = (int)A16.size(0);
    const int n = (int)A16.size(1);
    const bool panel_copy_batch_ok =
        (active_n >= 4096) ||
        (active_n >= 2048 && b >= 8) ||
        (active_n >= 1024 && b >= 16);
    if (panel_copy_batch_ok && ((n | active_n | nb) & 1) == 0 &&
        nb > 0 && active_n % nb == 0) {
        int tiles_per_panel = std::max(1, std::min(16, active_n >> 8));
        dim3 grid((active_n / nb) * tiles_per_panel, b);
        copy_cross_panel_r_panel_kernel<<<grid, 256>>>(
            (const __half*)A16.data_ptr<at::Half>(), H.data_ptr<float>(),
            n, nb, active_n, tiles_per_panel);
        C10_CUDA_KERNEL_LAUNCH_CHECK();
        return;
    }
    const long long per_matrix =
        (((n | active_n | nb) & 1) == 0)
            ? (long long)active_n * (active_n >> 1)
            : (long long)active_n * active_n;
    const long long total = (long long)b * per_matrix;
    const int blocks = (int)std::min<long long>((total + 255) / 256, 16384);
    copy_cross_panel_r_kernel<<<blocks, 256>>>(
        (const __half*)A16.data_ptr<at::Half>(), H.data_ptr<float>(),
        b, n, nb, active_n);
    C10_CUDA_KERNEL_LAUNCH_CHECK();
}

// zero a [rows x cols] f16 block of V16 at rows [row0, row0+rows), cols [0, cols)
__global__ void zero_f16_block_kernel(__half* __restrict__ V, int n, int ldv,
                                      int row0, int rows, int cols) {
    const int b = blockIdx.y;
    const long long idx = (long long)blockIdx.x * blockDim.x + threadIdx.x;
    if (idx >= (long long)rows * cols) return;
    int i = (int)(idx / cols), c = (int)(idx % cols);
    V[(size_t)b * n * ldv + (size_t)(row0 + i) * ldv + c] = __float2half_rn(0.f);
}

__global__ void zero_f32_buf_kernel(float* __restrict__ T, long long tot) {
    long long i = (long long)blockIdx.x * blockDim.x + threadIdx.x;
    if (i < tot) T[i] = 0.f;
}

// copy Tp[0:nb,0:nb] (upper-triangular, else 0) into Tblk[p:p+nb, p:p+nb]
__global__ void copy_tdiag_fp16_kernel(const float* __restrict__ Tp,
                                       float* __restrict__ Tblk,
                                       int ldtp, int ldt, int p, int nb) {
    const int b = blockIdx.y;
    const int idx = blockIdx.x * blockDim.x + threadIdx.x;
    if (idx >= nb * nb) return;
    int r = idx / nb, c = idx % nb;
    float v = (r <= c) ? Tp[(size_t)b * ldtp * ldtp + (size_t)r * ldtp + c] : 0.f;
    Tblk[(size_t)b * ldt * ldt + (size_t)(p + r) * ldt + (p + c)] = v;
}

// Apply a block of `vc` reflectors (V16[row0:, voff:voff+vc], factor T_ptr) to the
// trailing region rows[row0:n] cols[col0:col1]: trailing rows [row0+vc, n) updated
// in FP16 (A16), R rows [row0, row0+vc) written to H in FP32.
static void apply_blk_fp16(cublasHandle_t h, int b, int n,
                           __half* A16b, float* Hb, __half* V16b, int ldv,
                           const float* T_ptr, int ldt, long long sT,
                           int row0, int voff, int vc, int col0, int col1,
                           float* W1p, float* W2p, __half* W1hp, __half* W2hp,
                           long long sWcap, const __half* T16_ptr, bool seed_r) {
    const int m = n - row0;
    const int width = col1 - col0;
    if (width <= 0) return;
    const long long sA = (long long)n * n;
    const long long sV = (long long)n * ldv;
    const long long sH = (long long)n * n;
    const float one = 1.f, zero = 0.f, neg1 = -1.f;
    const __half* Atp = A16b + (size_t)row0 * n + col0;
    const __half* Vp = V16b + (size_t)row0 * ldv + voff;
    if (T16_ptr != nullptr) {
        TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
            h, CUBLAS_OP_N, CUBLAS_OP_T, width, vc, m, &one,
            Atp, CUDA_R_16F, n, sA, Vp, CUDA_R_16F, ldv, sV, &zero,
            W1hp, CUDA_R_16F, width, sWcap, b, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
        TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
            h, CUBLAS_OP_N, CUBLAS_OP_T, width, vc, vc, &one,
            W1hp, CUDA_R_16F, width, sWcap, T16_ptr, CUDA_R_16F, ldt, sT, &zero,
            W2hp, CUDA_R_16F, width, sWcap, b, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
    } else {
        TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
            h, CUBLAS_OP_N, CUBLAS_OP_T, width, vc, m, &one,
            Atp, CUDA_R_16F, n, sA, Vp, CUDA_R_16F, ldv, sV, &zero,
            W1p, CUDA_R_32F, width, sWcap, b, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
        TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
            h, CUBLAS_OP_N, CUBLAS_OP_T, width, vc, vc, &one,
            W1p, CUDA_R_32F, width, sWcap, T_ptr, CUDA_R_32F, ldt, sT, &zero,
            W2p, CUDA_R_32F, width, sWcap, b, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
        int per = width * vc;
        long long cast_work = (((per | sWcap) & 1) == 0) ? (long long)(per >> 1) * b : (long long)per * b;
        cast_f32_f16_strided<<<(int)((cast_work + 255) / 256), 256>>>(
            W2p, W2hp, per, sWcap, b);
        C10_CUDA_KERNEL_LAUNCH_CHECK();
    }
    // (a) trailing rows [row0+vc, n) -= W2h V_trail^T   (FP16)
    int mtrail = m - vc;
    if (mtrail > 0) {
        __half* AtTrail = A16b + (size_t)(row0 + vc) * n + col0;
        const __half* Vtrail = Vp + (size_t)vc * ldv;
        TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
            h, CUBLAS_OP_N, CUBLAS_OP_N, width, mtrail, vc, &neg1,
            W2hp, CUDA_R_16F, width, sWcap, Vtrail, CUDA_R_16F, ldv, sV, &one,
            AtTrail, CUDA_R_16F, n, sA, b, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
    }
    // (b) R rows [row0, row0+vc): seed H = pre-update A16, then -= W2h V_top^T (FP32)
    if (seed_r) {
        long long total_copy = (width & 1) ? (long long)vc * width : (long long)vc * (width >> 1);
        dim3 grid((int)((total_copy + 255) / 256), b);
        copy_R_block_kernel<<<grid, 256>>>((const __half*)A16b, Hb, n, row0, col0, vc, width);
        C10_CUDA_KERNEL_LAUNCH_CHECK();
    }
    float* Hrp = Hb + (size_t)row0 * n + col0;
    TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
        h, CUBLAS_OP_N, CUBLAS_OP_N, width, vc, vc, &neg1,
        W2hp, CUDA_R_16F, width, sWcap, Vp, CUDA_R_16F, ldv, sV, &one,
        Hrp, CUDA_R_32F, n, sH, b, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
}

static void apply_blk_fp16_defer(cublasHandle_t h, int b, int n,
                                 __half* A16b, __half* V16b, int ldv,
                                 const __half* T16_ptr, int ldt, long long sT,
                                 int row0, int voff, int vc, int col0, int col1,
                                 __half* W1hp, __half* W2hp, long long sWcap) {
    const int m = n - row0;
    const int width = col1 - col0;
    if (width <= 0) return;
    const long long sA = (long long)n * n;
    const long long sV = (long long)n * ldv;
    const float one = 1.f, zero = 0.f, neg1 = -1.f;
    const __half* Atp = A16b + (size_t)row0 * n + col0;
    const __half* Vp = V16b + (size_t)row0 * ldv + voff;
    TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
        h, CUBLAS_OP_N, CUBLAS_OP_T, width, vc, m, &one,
        Atp, CUDA_R_16F, n, sA, Vp, CUDA_R_16F, ldv, sV, &zero,
        W1hp, CUDA_R_16F, width, sWcap, b,
        CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
    TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
        h, CUBLAS_OP_N, CUBLAS_OP_T, width, vc, vc, &one,
        W1hp, CUDA_R_16F, width, sWcap, T16_ptr, CUDA_R_16F, ldt, sT, &zero,
        W2hp, CUDA_R_16F, width, sWcap, b,
        CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
    __half* Atw = A16b + (size_t)row0 * n + col0;
    TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
        h, CUBLAS_OP_N, CUBLAS_OP_N, width, m, vc, &neg1,
        W2hp, CUDA_R_16F, width, sWcap, Vp, CUDA_R_16F, ldv, sV, &one,
        Atw, CUDA_R_16F, n, sA, b,
        CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
}

// Fold panel Tp (cnb x cnb) into the block factor Tblk:
//   Tblk[0:p, p:p+cnb] = -Tblk[0:p,0:p] (Vacc^T Vnew) Tp
static void fold_t_fp16(cublasHandle_t h, int b, int n, const __half* V16b, int ldv,
                        float* Tblkb, int ldt, const float* Tpb, int ldtp,
                        float* Xb, float* Yb, int ldm, long long sM,
                        int row0, int p, int cnb) {
    const int m = n - row0;
    const long long sV = (long long)n * ldv;
    const long long sTb = (long long)ldt * ldt;
    const long long sTp = (long long)ldtp * ldtp;
    const float one = 1.f, zero = 0.f, neg1 = -1.f;
    const __half* Vacc = V16b + (size_t)row0 * ldv;
    const __half* Vnew = Vacc + p;
    float* Tdst = Tblkb + p;
    // X = Vnew^T Vacc  (cnb x p, K=m)
    TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
        h, CUBLAS_OP_N, CUBLAS_OP_T, cnb, p, m, &one,
        Vnew, CUDA_R_16F, ldv, sV, Vacc, CUDA_R_16F, ldv, sV, &zero,
        Xb, CUDA_R_32F, ldm, sM, b, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
    // Y = X @ Tblk[0:p,0:p]  (cnb x p, K=p)
    TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
        h, CUBLAS_OP_N, CUBLAS_OP_N, cnb, p, p, &one,
        Xb, CUDA_R_32F, ldm, sM, Tblkb, CUDA_R_32F, ldt, sTb, &zero,
        Yb, CUDA_R_32F, ldm, sM, b, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
    // Tdst = -Tp @ Y  (cnb x p)
    TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
        h, CUBLAS_OP_N, CUBLAS_OP_N, cnb, p, cnb, &neg1,
        Tpb, CUDA_R_32F, ldtp, sTp, Yb, CUDA_R_32F, ldm, sM, &zero,
        Tdst, CUDA_R_32F, ldt, sTb, b, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
}

// Two-level blocked QR with FP16 trailing storage. Inner panels of width nb feed a
// wout-wide outer block; the bulk trailing update is done ONCE per block at K=wout
// (4x fewer / 4x-wider GEMMs than the single-level K=nb path -> far higher TC eff
// at the batch sizes that matter, e.g. n1024 b60, n2048 b8).
std::vector<torch::Tensor> blocked_qr_fp16_2level(torch::Tensor data,
                                                  int64_t nb_in, int64_t wout_in) {
    const int b = data.size(0);
    const int n = data.size(1);
    const int nb = (int)nb_in;       // inner panel width (<= 32)
    int wcode = (int)wout_in;
    const bool hybrid_after_192 = (wcode >= 1000);
    if (hybrid_after_192) wcode -= 1000;
    const bool fast_direct = (wcode > 0);
    const int wout = (int)(fast_direct ? wcode : -wcode);   // outer block width
    const bool direct_tblk_n512 = (n == 512 && nb == 16 && wout == 64);
    auto f32 = data.options().dtype(torch::kFloat32);
    auto f16 = data.options().dtype(torch::kHalf);
    auto A16 = data.to(torch::kHalf).contiguous();
    auto H = torch::empty({b, n, n}, f32);
    auto tau = torch::empty({b, n}, f32);
    auto V16 = torch::empty({b, n, wout}, f16);
    torch::Tensor Tp;
    if (!direct_tblk_n512)
        Tp = torch::empty({b, FP16_NBMAX, FP16_NBMAX}, f32);
    auto Tblk = torch::empty({b, wout, wout}, f32);
    // Inv352: pure-direct n512 never enters the FP32 W1/T/W2 branch.
    // Keep these workspaces only for the negative-mode path or rowscale hybrid.
    const bool need_fp32_ws = !fast_direct || hybrid_after_192;
    torch::Tensor W1, W2;
    if (need_fp32_ws) {
        W1 = torch::empty({b, wout, n}, f32);
        W2 = torch::empty({b, wout, n}, f32);
    }
    auto W2h = torch::empty({b, wout, n}, f16);
    torch::Tensor Tp16, Tblk16, W1h;
    if (fast_direct) {
        if (!direct_tblk_n512)
            Tp16 = torch::empty({b, FP16_NBMAX, FP16_NBMAX}, f16);
        Tblk16 = torch::empty({b, wout, wout}, f16);
        W1h = torch::empty({b, wout, n}, f16);
    }
    // Inv354: fold scratch stores cnb-by-p matrices; leading dimension nb
    // is sufficient (cnb<=nb) and removes the unused 64-cnb lanes.
    auto Xm = torch::empty({b, wout, nb}, f32);
    auto Ym = torch::empty({b, wout, nb}, f32);

    const int TH = (n == 1024) ? 1024 : (n >= 2048 ? 512 : 256);
    static int max_dyn2 = -1;
    if (max_dyn2 < 0) {
        int dev; C10_CUDA_CHECK(cudaGetDevice(&dev));
        cudaDeviceProp prop; C10_CUDA_CHECK(cudaGetDeviceProperties(&prop, dev));
        cudaFuncAttributes fa256, fa512, fa1024;
        C10_CUDA_CHECK(cudaFuncGetAttributes(&fa256, panel_kernel_fp16<256, false>));
        C10_CUDA_CHECK(cudaFuncGetAttributes(&fa512, panel_kernel_fp16<512, false>));
        C10_CUDA_CHECK(cudaFuncGetAttributes(&fa1024, panel_kernel_fp16<1024, false>));
        int m256 = (int)(prop.sharedMemPerBlockOptin - fa256.sharedSizeBytes);
        int m512 = (int)(prop.sharedMemPerBlockOptin - fa512.sharedSizeBytes);
        int m1024 = (int)(prop.sharedMemPerBlockOptin - fa1024.sharedSizeBytes);
        C10_CUDA_CHECK(cudaFuncSetAttribute(panel_kernel_fp16<256, false>,
            cudaFuncAttributeMaxDynamicSharedMemorySize, m256));
        C10_CUDA_CHECK(cudaFuncSetAttribute(panel_kernel_fp16<512, false>,
            cudaFuncAttributeMaxDynamicSharedMemorySize, m512));
        C10_CUDA_CHECK(cudaFuncSetAttribute(panel_kernel_fp16<1024, false>,
            cudaFuncAttributeMaxDynamicSharedMemorySize, m1024));
        max_dyn2 = m256 < m512 ? m256 : m512;
        if (m1024 < max_dyn2) max_dyn2 = m1024;
    }
    {
        int m0 = n, mp0 = (m0 & 1) ? m0 : m0 + 1;
        TORCH_CHECK((size_t)mp0 * nb * sizeof(float) <= (size_t)max_dyn2,
                    "fp16 2level panel does not fit in shared memory for this nb");
    }

    cublasHandle_t h = at::cuda::getCurrentCUDABlasHandle();
    __half* A16b = (__half*)A16.data_ptr<at::Half>();
    __half* V16b = (__half*)V16.data_ptr<at::Half>();
    float* Hb = H.data_ptr<float>();
    float* taub = tau.data_ptr<float>();
    float* Tblkb = Tblk.data_ptr<float>();
    __half* Tblk16b = fast_direct ? (__half*)Tblk16.data_ptr<at::Half>() : nullptr;
    float* Tpb = direct_tblk_n512 ? Tblkb : Tp.data_ptr<float>();
    __half* Tp16b = fast_direct
        ? (direct_tblk_n512 ? Tblk16b : (__half*)Tp16.data_ptr<at::Half>())
        : nullptr;
    float* W1p = need_fp32_ws ? W1.data_ptr<float>() : nullptr;
    float* W2p = need_fp32_ws ? W2.data_ptr<float>() : nullptr;
    __half* W1hp = fast_direct ? (__half*)W1h.data_ptr<at::Half>() : nullptr;
    __half* W2hp = (__half*)W2h.data_ptr<at::Half>();
    float* Xb = Xm.data_ptr<float>();
    float* Yb = Ym.data_ptr<float>();
    const int ldv = wout, ldt = wout;
    const int ldtp = direct_tblk_n512 ? ldt : FP16_NBMAX;
    const int ldm = nb;
    const long long sWcap = (long long)wout * n;
    const long long sTblk = (long long)wout * wout;
    const long long sTp = direct_tblk_n512
        ? sTblk : (long long)FP16_NBMAX * FP16_NBMAX;
    const long long sMfold = (long long)wout * nb;  // compact X/Y batch stride

    for (int k0 = 0; k0 < n; k0 += wout) {
        int w = (wout < n - k0) ? wout : (n - k0);
        const bool need_outer = (k0 + w < n);
        const bool block_direct = fast_direct && (!hybrid_after_192 || k0 >= 192);
        // The specialized n512/nb16 panel defines its own strict-upper V
        // prefix. Preserve the generic driver's old clear for any other call.
        if (!(n == 512 && nb == 16)) {
            dim3 g((int)(((long long)w * w + 255) / 256), b);
            zero_f16_block_kernel<<<g, 256>>>(V16b, n, ldv, k0, w, w);
            C10_CUDA_KERNEL_LAUNCH_CHECK();
        }
        if (need_outer && !direct_tblk_n512) {
            zero_f32_buf_kernel<<<(int)((b * sTblk + 255) / 256), 256>>>(Tblkb, b * sTblk);
            C10_CUDA_KERNEL_LAUNCH_CHECK();
        }
        for (int k = k0; k < k0 + w; k += nb) {
            int cnb = (nb < k0 + w - k) ? nb : (k0 + w - k);
            int p = k - k0;
            int m = n - k;
            int mp = (m & 1) ? m : m + 1;
            size_t smem = (size_t)mp * cnb * sizeof(float);
            const int tpanel_off = direct_tblk_n512 ? (need_outer ? p : 0) : 0;
            float* Tpanel = direct_tblk_n512
                ? (Tblkb + (size_t)tpanel_off * ldt + tpanel_off) : Tpb;
            __half* Tpanel16 = block_direct
                ? (direct_tblk_n512
                    ? (Tblk16b + (size_t)tpanel_off * ldt + tpanel_off)
                    : Tp16b)
                : nullptr;
            int rseed_end = block_direct ? (k0 + w) : 0;
            if (n == 512 && nb == 16 && cnb == 16) {
                panel_kernel_fp16_nb16<256, false, true, false, true, true><<<b, 256, smem>>>(
                    A16b, Hb, taub, V16b,
                    direct_tblk_n512 ? Tblkb : Tpb,
                    n, k, ldv, p,
                    direct_tblk_n512 ? (block_direct ? Tblk16b : nullptr) : Tpanel16,
                    tpanel_off, rseed_end);
            }
            else if (TH == 1024)
                panel_kernel_fp16<1024, false><<<b, 1024, smem>>>(A16b, Hb, taub, V16b, Tpb, n, k, cnb, ldv, p, Tpanel16, rseed_end);
            else if (TH == 512)
                panel_kernel_fp16<512, false><<<b, 512, smem>>>(A16b, Hb, taub, V16b, Tpb, n, k, cnb, ldv, p, Tpanel16, rseed_end);
            else
                panel_kernel_fp16<256, false><<<b, 256, smem>>>(A16b, Hb, taub, V16b, Tpb, n, k, cnb, ldv, p, Tpanel16, rseed_end);
            C10_CUDA_KERNEL_LAUNCH_CHECK();
            // within-block trailing update using this panel's T view.
            apply_blk_fp16(h, b, n, A16b, Hb, V16b, ldv,
                           Tpanel, ldtp, sTp,
                           k, p, cnb, k + cnb, k0 + w, W1p, W2p,
                           W1hp, W2hp, sWcap, Tpanel16,
                           !block_direct);
            if (need_outer) {
                if (!direct_tblk_n512) {
                    dim3 gd((cnb * cnb + 255) / 256, b);
                    copy_tdiag_fp16_kernel<<<gd, 256>>>(Tpb, Tblkb, ldtp, ldt, p, cnb);
                    C10_CUDA_KERNEL_LAUNCH_CHECK();
                }
                if (p > 0)
                    fold_t_fp16(h, b, n, V16b, ldv, Tblkb, ldt,
                                Tpanel, ldtp,
                                Xb, Yb, ldm, sMfold, k0, p, cnb);
            }
        }
        if (need_outer) {
            // wide tail update: all w reflectors applied to cols [k0+w, n) at K=w
            const __half* Touter16 = nullptr;
            if (block_direct) {
                int per = w * w;
                long long cast_work = (((per | sTblk) & 1) == 0) ? (long long)(per >> 1) * b : (long long)per * b;
                cast_f32_f16_strided<<<(int)((cast_work + 255) / 256), 256>>>(
                    Tblkb, Tblk16b, w * w, sTblk, b);
                C10_CUDA_KERNEL_LAUNCH_CHECK();
                Touter16 = Tblk16b;
            }
            apply_blk_fp16(h, b, n, A16b, Hb, V16b, ldv, Tblkb, ldt, sTblk,
                           k0, 0, w, k0 + w, n, W1p, W2p,
                           W1hp, W2hp, sWcap, Touter16, true);
        }
    }
    return {H, tau};
}

static std::vector<torch::Tensor> blocked_qr_fp16_cluster_core(torch::Tensor data,
                                                               int nb, int wout, int cb) {
    const int b = data.size(0);
    const int n = data.size(1);
    TORCH_CHECK(nb <= FP16_NBMAX && wout % nb == 0 && wout <= n,
                "cluster blocking: need nb<=FP16_NBMAX, wout%nb==0, wout<=n");
    TORCH_CHECK(cb == 8 || cb == 16, "cluster blocks must be 8 or 16");
    auto f32 = data.options().dtype(torch::kFloat32);
    auto f16 = data.options().dtype(torch::kHalf);
    auto A16 = data.to(torch::kHalf).contiguous();
    auto H = torch::empty({b, n, n}, f32);
    auto tau = torch::zeros({b, n}, f32);
    auto V16 = torch::zeros({b, n, wout}, f16);
    auto Tp = torch::empty({b, FP16_NBMAX, FP16_NBMAX}, f32);
    auto Tp16 = torch::empty({b, FP16_NBMAX, FP16_NBMAX}, f16);
    auto Tblk = torch::empty({b, wout, wout}, f32);
    auto Tblk16 = torch::empty({b, wout, wout}, f16);
    auto W1h = torch::empty({b, wout, n}, f16);
    auto W2h = torch::empty({b, wout, n}, f16);
    auto Xm = torch::empty({b, wout, FP16_NBMAX}, f32);
    auto Ym = torch::empty({b, wout, FP16_NBMAX}, f32);

    cublasHandle_t h = at::cuda::getCurrentCUDABlasHandle();
    __half* A16b = (__half*)A16.data_ptr<at::Half>();
    __half* V16b = (__half*)V16.data_ptr<at::Half>();
    float* Tpb = Tp.data_ptr<float>();
    __half* Tp16b = (__half*)Tp16.data_ptr<at::Half>();
    float* Tblkb = Tblk.data_ptr<float>();
    __half* Tblk16b = (__half*)Tblk16.data_ptr<at::Half>();
    __half* W1hp = (__half*)W1h.data_ptr<at::Half>();
    __half* W2hp = (__half*)W2h.data_ptr<at::Half>();
    float* Xb = Xm.data_ptr<float>();
    float* Yb = Ym.data_ptr<float>();
    const int ldv = wout;
    const int ldt = wout;
    const int ldtp = FP16_NBMAX;
    const int ldm = FP16_NBMAX;
    const long long sWcap = (long long)wout * n;
    const long long sTblk = (long long)wout * wout;
    const long long sTp = (long long)FP16_NBMAX * FP16_NBMAX;
    const long long sMfold = (long long)wout * FP16_NBMAX;

    for (int k0 = 0; k0 < n; k0 += wout) {
        int w = (wout < n - k0) ? wout : (n - k0);
        const bool need_outer = (k0 + w < n);
        // Inv334 (n4096 cluster route): when the outer block is a single panel
        // (w == nb), the outer-block T is exactly the panel T already emitted in
        // Tp16/sTp. Skip the Tblk zero/copy_tdiag/cast rebuild and apply directly.
        const bool single_panel_outer = need_outer && (w == nb);
        if (!single_panel_outer) {
            dim3 g((int)(((long long)w * w + 255) / 256), b);
            zero_f16_block_kernel<<<g, 256>>>(V16b, n, ldv, k0, w, w);
            C10_CUDA_KERNEL_LAUNCH_CHECK();
        }
        if (need_outer && !single_panel_outer) {
            zero_f32_buf_kernel<<<(int)((b * sTblk + 255) / 256), 256>>>(Tblkb, b * sTblk);
            C10_CUDA_KERNEL_LAUNCH_CHECK();
        }
        for (int k = k0; k < k0 + w; k += nb) {
            int cnb = (nb < k0 + w - k) ? nb : (k0 + w - k);
            int p = k - k0;
            const int m = n - k;
            const int rows_per = (m + cb - 1) / cb;
            const bool owner0_panel = (rows_per >= cnb);
            if (cb == 8) {
                if (g_panel_1sync) {
                    const bool n4096_half_apply = false;
                    if (owner0_panel) {
                        if (n4096_half_apply)
                            launch_cluster_panel_fp16_leaf16<512, 8, true, true>(A16, H, tau, V16, Tp, Tp16,
                                                                                k, cnb, p);
                        else
                            launch_cluster_panel_fp16_leaf16<512, 8, true, false>(A16, H, tau, V16, Tp, Tp16,
                                                                                 k, cnb, p);
                    } else {
                        if (n4096_half_apply)
                            launch_cluster_panel_fp16_leaf16<512, 8, false, true>(A16, H, tau, V16, Tp, Tp16,
                                                                                 k, cnb, p);
                        else
                            launch_cluster_panel_fp16_leaf16<512, 8, false, false>(A16, H, tau, V16, Tp, Tp16,
                                                                                  k, cnb, p);
                    }
                } else {
                    launch_cluster_panel_fp16<512, 8, false>(A16, H, tau, V16, Tp, Tp16,
                                                             k, cnb, p);
                }
            } else {
                if (g_panel_1sync)
                    launch_cluster_panel_fp16<512, 16, true>(A16, H, tau, V16, Tp, Tp16,
                                                             k, cnb, p);
                else
                    launch_cluster_panel_fp16<512, 16, false>(A16, H, tau, V16, Tp, Tp16,
                                                              k, cnb, p);
            }
            apply_blk_fp16_defer(h, b, n, A16b, V16b, ldv, Tp16b, ldtp, sTp,
                                 k, p, cnb, k + cnb, k0 + w,
                                 W1hp, W2hp, sWcap);
            if (need_outer && !single_panel_outer) {
                dim3 gd((cnb * cnb + 255) / 256, b);
                copy_tdiag_fp16_kernel<<<gd, 256>>>(Tpb, Tblkb, ldtp, ldt, p, cnb);
                C10_CUDA_KERNEL_LAUNCH_CHECK();
                if (p > 0) {
                    fold_t_fp16(h, b, n, V16b, ldv, Tblkb, ldt, Tpb, ldtp,
                                Xb, Yb, ldm, sMfold, k0, p, cnb);
                }
            }
        }
        if (need_outer) {
            if (single_panel_outer) {
                apply_blk_fp16_defer(h, b, n, A16b, V16b, ldv, Tp16b, ldtp, sTp,
                                     k0, 0, w, k0 + w, n,
                                     W1hp, W2hp, sWcap);
            } else {
                int per = w * w;
                long long cast_work = (((per | sTblk) & 1) == 0) ? (long long)(per >> 1) * b : (long long)per * b;
                cast_f32_f16_strided<<<(int)((cast_work + 255) / 256), 256>>>(
                    Tblkb, Tblk16b, w * w, sTblk, b);
                C10_CUDA_KERNEL_LAUNCH_CHECK();
                apply_blk_fp16_defer(h, b, n, A16b, V16b, ldv, Tblk16b, ldt, sTblk,
                                     k0, 0, w, k0 + w, n,
                                     W1hp, W2hp, sWcap);
            }
        }
    }
    copy_cross_panel_r(A16, H, nb, n);
    return {H, tau};
}

std::vector<torch::Tensor> blocked_qr_fp16_cluster4096(torch::Tensor data) {
    const int n = data.size(1);
    TORCH_CHECK(n == 4096, "fp16 cluster4096 route is only for n=4096");
    // nb=32 single-level (wout==nb) is the sweep winner: ~1.10x over nb=16/wout=128
    // (fewer sequential panel/apply launches) and tighter residual (fr 0.051 vs 0.061).
    const int nb = (g_n4096_nb > 0) ? g_n4096_nb : 32;
    const int wout = (g_n4096_wout > 0) ? g_n4096_wout : 32;
    // Inv491: route n4096 to CB=8 instead of CB=16. CB=8 auto-enables the
    // leaf16+WMMA panel (cluster_core gates leaf16 on cb==8) -- the proven Inv442
    // n2048 winner -- AND halves the per-column cluster.sync from 16 ranks to 8
    // (the fresh probe shows n4096's scalar CB16 panel is 84.6% / 127ms; the prior
    // refutation of leaf16 at CB16 was 16-rank-sync-bound, NOT smem-bound).
    // n4096 b2 -> 16 cooperative blocks; rows_per = 4096/8 = 512 -> P = 64KB fits.
    // Measured -5.4% on n4096_dense (modal_focus_inv491), correct (factor 1.49).
    return blocked_qr_fp16_cluster_core(data, nb, wout, 8);
}

// (Inv 320) generic cluster-panel route. Same machinery as the n4096 route but
// callable for any n that benefits from row-split DSM cluster panels when the
// batch is too small to fill the GPU with one CTA per matrix (e.g. n2048 b8 =
// only 8 CTAs on 148 SMs). nb/wout default to 32/32 (single-level) when <=0.
// cb selects cluster blocks per matrix (8 or 16). A cooperative launch caps the
// grid at the device co-residency limit, so total blocks = b*cb must stay small
// enough: n2048 b8 with CB=16 (128 blocks) overflowed it; CB=8 (64) is safe.
std::vector<torch::Tensor> blocked_qr_fp16_cluster_generic(torch::Tensor data,
                                                           int64_t nb_in,
                                                           int64_t wout_in,
                                                           int64_t cb_in) {
    const int nb = (nb_in > 0) ? (int)nb_in : 32;
    const int wout = (wout_in > 0) ? (int)wout_in : 32;
    const int cb = (cb_in == 16) ? 16 : 8;
    return blocked_qr_fp16_cluster_core(data, nb, wout, cb);
}

// Compiled driver: the whole blocked QR loop runs in C++ (no Python per-panel
// launches), eliminating the launch-overhead + timing variance that pushed the
// Python driver over the 300s per-input benchmark timeout. Single-level, nb=32.
std::vector<torch::Tensor> blocked_qr_fp16(torch::Tensor data, int64_t nb_in) {
    const int b = data.size(0);
    const int n = data.size(1);
    int nb_code = (int)nb_in;
    const bool defer_r = (nb_code >= 1000);
    if (defer_r) nb_code -= 1000;
    const int nb = nb_code;   // panel width; chosen so the tallest panel fits smem
    auto f32 = data.options().dtype(torch::kFloat32);
    auto f16 = data.options().dtype(torch::kHalf);
    auto A16 = data.to(torch::kHalf).contiguous();
    auto H = defer_r ? torch::empty({b, n, n}, f32) : torch::zeros({b, n, n}, f32);
    auto tau = torch::zeros({b, n}, f32);
    auto V16 = torch::zeros({b, n, nb}, f16);
    const bool skip_t32_full = defer_r && n == 1024 && nb == 32;
    torch::Tensor Tp;
    if (!skip_t32_full)
        Tp = torch::empty({b, FP16_NBMAX, FP16_NBMAX}, f32);
    auto Tp16 = torch::empty({b, FP16_NBMAX, FP16_NBMAX}, f16);   // FP16 T (Inv 228)
    auto W1 = torch::empty({b, nb, n}, f32);
    auto W1h = torch::empty({b, nb, n}, f16);                     // FP16 W1 (Inv 228)
    auto W2 = torch::empty({b, nb, n}, f32);
    auto W2h = torch::empty({b, nb, n}, f16);

    // More threads/CTA speed the occupancy-starved tall panels at large n (only
    // ~b CTAs run). Active fablin uses 512 threads for n>=2048 (~20% faster
    // panels). Opt in smem for both instantiations.
    // More threads/CTA speed the occupancy-starved tall panels (only ~b CTAs
    // run). n=1024 prefers 1024 threads (active fablin's choice). Test an
    // intermediate 768-thread panel for n>=2048; 1024 was rejected there.
    const int TH = (n == 1024) ? 1024 : (n >= 2048 ? 768 : 256);
    static int max_dyn = -1;
    if (max_dyn < 0) {
        int dev; C10_CUDA_CHECK(cudaGetDevice(&dev));
        cudaDeviceProp prop; C10_CUDA_CHECK(cudaGetDeviceProperties(&prop, dev));
        cudaFuncAttributes fa256, fa512, fa768, fa1024, fa1024_owner_not32, fa1024_owner_fused, fa_nb24_768;
        C10_CUDA_CHECK(cudaFuncGetAttributes(&fa256, panel_kernel_fp16<256, false>));
        C10_CUDA_CHECK(cudaFuncGetAttributes(&fa512, panel_kernel_fp16<512, false>));
        C10_CUDA_CHECK(cudaFuncGetAttributes(&fa768, panel_kernel_fp16<768, false>));
        C10_CUDA_CHECK(cudaFuncGetAttributes(&fa1024, panel_kernel_fp16<1024, false>));
        C10_CUDA_CHECK(cudaFuncGetAttributes(&fa1024_owner_not32, panel_kernel_fp16<1024, false, false, true>));
        C10_CUDA_CHECK(cudaFuncGetAttributes(&fa1024_owner_fused, panel_kernel_fp16<1024, false, false, true, true>));
        C10_CUDA_CHECK(cudaFuncGetAttributes(&fa_nb24_768, panel_kernel_fp16_nb_static<768, 24, false>));
        int m256 = (int)(prop.sharedMemPerBlockOptin - fa256.sharedSizeBytes);
        int m512 = (int)(prop.sharedMemPerBlockOptin - fa512.sharedSizeBytes);
        int m768 = (int)(prop.sharedMemPerBlockOptin - fa768.sharedSizeBytes);
        int m1024 = (int)(prop.sharedMemPerBlockOptin - fa1024.sharedSizeBytes);
        int m1024_owner_not32 = (int)(prop.sharedMemPerBlockOptin - fa1024_owner_not32.sharedSizeBytes);
        int m1024_owner_fused = (int)(prop.sharedMemPerBlockOptin - fa1024_owner_fused.sharedSizeBytes);
        int mnb24 = (int)(prop.sharedMemPerBlockOptin - fa_nb24_768.sharedSizeBytes);
        C10_CUDA_CHECK(cudaFuncSetAttribute(panel_kernel_fp16<256, false>,
            cudaFuncAttributeMaxDynamicSharedMemorySize, m256));
        C10_CUDA_CHECK(cudaFuncSetAttribute(panel_kernel_fp16<512, false>,
            cudaFuncAttributeMaxDynamicSharedMemorySize, m512));
        C10_CUDA_CHECK(cudaFuncSetAttribute(panel_kernel_fp16<768, false>,
            cudaFuncAttributeMaxDynamicSharedMemorySize, m768));
        C10_CUDA_CHECK(cudaFuncSetAttribute(panel_kernel_fp16<1024, false>,
            cudaFuncAttributeMaxDynamicSharedMemorySize, m1024));
        C10_CUDA_CHECK(cudaFuncSetAttribute(panel_kernel_fp16<1024, false, false, true>,
            cudaFuncAttributeMaxDynamicSharedMemorySize, m1024_owner_not32));
        C10_CUDA_CHECK(cudaFuncSetAttribute(panel_kernel_fp16<1024, false, false, true, true>,
            cudaFuncAttributeMaxDynamicSharedMemorySize, m1024_owner_fused));
        C10_CUDA_CHECK(cudaFuncSetAttribute(panel_kernel_fp16_nb_static<768, 24, false>,
            cudaFuncAttributeMaxDynamicSharedMemorySize, mnb24));
        max_dyn = m256 < m512 ? m256 : m512;
        if (m768 < max_dyn) max_dyn = m768;
        if (m1024 < max_dyn) max_dyn = m1024;
        if (m1024_owner_not32 < max_dyn) max_dyn = m1024_owner_not32;
        if (m1024_owner_fused < max_dyn) max_dyn = m1024_owner_fused;
    }
    {
        int m0 = n, mp0 = (m0 & 1) ? m0 : m0 + 1;   // tallest panel (k=0)
        TORCH_CHECK((size_t)mp0 * nb * sizeof(float) <= (size_t)max_dyn,
                    "fp16 panel does not fit in shared memory for this nb");
    }

    cublasHandle_t h = at::cuda::getCurrentCUDABlasHandle();
    const float one = 1.f, zero = 0.f, neg1 = -1.f;
    __half* A16b = (__half*)A16.data_ptr<at::Half>();
    __half* V16b = (__half*)V16.data_ptr<at::Half>();
    float* Hb = H.data_ptr<float>();
    float* taub = tau.data_ptr<float>();
    float* Tpb = skip_t32_full ? nullptr : Tp.data_ptr<float>();
    __half* Tp16b = (__half*)Tp16.data_ptr<at::Half>();
    float* W1p = W1.data_ptr<float>();
    __half* W1hp = (__half*)W1h.data_ptr<at::Half>();
    float* W2p = W2.data_ptr<float>();
    __half* W2hp = (__half*)W2h.data_ptr<at::Half>();
    const long long sA = (long long)n * n;
    const long long sV = (long long)n * nb;
    const long long sT = (long long)FP16_NBMAX * FP16_NBMAX;
    const long long sWf = (long long)nb * n;     // full W1/W2 batch stride
    const long long sH = (long long)n * n;

    for (int k = 0; k < n; k += nb) {
        int cnb = (nb < n - k) ? nb : (n - k);
        int m = n - k;
        int mp = (m & 1) ? m : m + 1;
        size_t smem = (size_t)mp * cnb * sizeof(float);
        // (Inv 228) Tg16=Tp16b: panel emits FP16 T so GEMM2 outputs FP16 with no cast.
        // In deferred-R mode, cross-panel R is copied once at the end instead.
        const int rseed_end = defer_r ? 0 : n;
        if (skip_t32_full && TH == 1024 && cnb == 32)
            panel_kernel_fp16<1024, false, false, true, true><<<b, 1024, smem>>>(A16b, Hb, taub, V16b, Tpb, n, k, cnb, nb, 0, Tp16b, rseed_end);
        else if (TH == 1024)
            panel_kernel_fp16<1024, false><<<b, 1024, smem>>>(A16b, Hb, taub, V16b, Tpb, n, k, cnb, nb, 0, Tp16b, rseed_end);
        else if (n == 2048 && nb == 24 && cnb == 24)
            panel_kernel_fp16_nb_static<768, 24, false><<<b, 768, smem>>>(A16b, Hb, taub, V16b, Tpb, n, k, nb, 0, Tp16b, rseed_end);
        else if (TH == 768)
            panel_kernel_fp16<768, false><<<b, 768, smem>>>(A16b, Hb, taub, V16b, Tpb, n, k, cnb, nb, 0, Tp16b, rseed_end);
        else if (TH == 512)
            panel_kernel_fp16<512, false><<<b, 512, smem>>>(A16b, Hb, taub, V16b, Tpb, n, k, cnb, nb, 0, Tp16b, rseed_end);
        else
            panel_kernel_fp16<256, false><<<b, 256, smem>>>(A16b, Hb, taub, V16b, Tpb, n, k, cnb, nb, 0, Tp16b, rseed_end);
        C10_CUDA_KERNEL_LAUNCH_CHECK();

        int col0 = k + cnb;
        int r = n - col0;
        if (r <= 0) continue;
        int vc = cnb;
        const __half* Atp = A16b + (size_t)k * n + col0;
        const __half* Vp = V16b + (size_t)k * nb;

        // (Inv 228) GEMM1 emits FP16 W1, GEMM2 (FP16 W1 @ FP16 T) emits FP16 W2 directly.
        // FP32 accumulate throughout; the only added rounding is W1->FP16, which is well
        // inside the n1024/n2048 gate margin (scaled factor residual ~0.1 vs gate 1.0).
        // Eliminates the per-panel cast_f32_f16_strided launch.
        TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
            h, CUBLAS_OP_N, CUBLAS_OP_T, r, vc, m, &one,
            Atp, CUDA_R_16F, n, sA, Vp, CUDA_R_16F, nb, sV, &zero,
            W1hp, CUDA_R_16F, r, sWf, b, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
        TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
            h, CUBLAS_OP_N, CUBLAS_OP_T, r, vc, vc, &one,
            W1hp, CUDA_R_16F, r, sWf, Tp16b, CUDA_R_16F, FP16_NBMAX, sT, &zero,
            W2hp, CUDA_R_16F, r, sWf, b, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
        int mtrail = defer_r ? m : (m - cnb);   // all rows if R is deferred
        if (mtrail > 0) {
            __half* AtTrail = A16b + (size_t)(defer_r ? k : (k + cnb)) * n + col0;
            const __half* Vtrail = defer_r ? Vp : (Vp + (size_t)cnb * nb);
            TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
                h, CUBLAS_OP_N, CUBLAS_OP_N, r, mtrail, vc, &neg1,
                W2hp, CUDA_R_16F, r, sWf, Vtrail, CUDA_R_16F, nb, sV, &one,
                AtTrail, CUDA_R_16F, n, sA, b, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
        }
        if (!defer_r) {
            // R-row seed now fused into the panel kernel (col_end=n); copy_rrows removed.
            float* Hrp = Hb + (size_t)k * n + col0;
            TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
                h, CUBLAS_OP_N, CUBLAS_OP_N, r, cnb, vc, &neg1,
                W2hp, CUDA_R_16F, r, sWf, Vp, CUDA_R_16F, nb, sV, &one,
                Hrp, CUDA_R_32F, n, sH, b, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
        }
    }
    if (defer_r) copy_cross_panel_r(A16, H, nb, n);
    return {H, tau};
}


// Inv344: exact active-prefix initialization with vectorized tail traffic.
// Current n512/n1024 active dimensions are 16-byte aligned; retain scalar fallback.
__global__ void init_active_outputs_kernel(const float* __restrict__ data,
                                           float* __restrict__ H,
                                           float* __restrict__ tau,
                                           int bsz, int n, int active_n,
                                           int h_tail_mode) {
    const int tail = n - active_n;
    if (tail <= 0) return;
    const long long first = (long long)blockIdx.x * blockDim.x + threadIdx.x;
    const long long stride = (long long)gridDim.x * blockDim.x;
    if (((n | active_n | tail) & 3) == 0) {
        const int tail4 = tail >> 2;
        const long long h_per4 = (long long)n * tail4;
        const long long h_total4 = h_tail_mode ? (long long)bsz * h_per4 : 0;
        for (long long idx = first; idx < h_total4; idx += stride) {
            const long long bi = idx / h_per4;
            const long long rem = idx - bi * h_per4;
            const int row = (int)(rem / tail4);
            const int c4 = (int)(rem - (long long)row * tail4);
            const long long off = bi * (long long)n * n +
                                  (long long)row * n + active_n + (c4 << 2);
            if (h_tail_mode == 1) {
                *reinterpret_cast<float4*>(H + off) =
                    *reinterpret_cast<const float4*>(data + off);
            } else {
                *reinterpret_cast<float4*>(H + off) = make_float4(0.f, 0.f, 0.f, 0.f);
            }
        }
        const long long t_total4 = (long long)bsz * tail4;
        for (long long idx = first; idx < t_total4; idx += stride) {
            const long long bi = idx / tail4;
            const int c4 = (int)(idx - bi * tail4);
            const long long off = bi * (long long)n + active_n + (c4 << 2);
            *reinterpret_cast<float4*>(tau + off) = make_float4(0.f, 0.f, 0.f, 0.f);
        }
    } else {
        const long long h_per = (long long)n * tail;
        const long long h_total = h_tail_mode ? (long long)bsz * h_per : 0;
        for (long long idx = first; idx < h_total; idx += stride) {
            const long long bi = idx / h_per;
            const long long rem = idx - bi * h_per;
            const int row = (int)(rem / tail);
            const int c = (int)(rem - (long long)row * tail);
            const long long off = bi * (long long)n * n +
                                  (long long)row * n + active_n + c;
            H[off] = (h_tail_mode == 1) ? data[off] : 0.f;
        }
        const long long t_total = (long long)bsz * tail;
        for (long long idx = first; idx < t_total; idx += stride) {
            const long long bi = idx / tail;
            const int c = (int)(idx - bi * tail);
            tau[bi * (long long)n + active_n + c] = 0.f;
        }
    }
}

// FP16 PREFIX QR: factor only columns [0, active_n) in FP16 (reflectors span full
// height); leave the tail columns as the original FP32 data (H starts as a clone).
// Valid where the tail is numerically empty (e.g. n512 `clustered`: cols[n/2:]~4*eps),
// exactly mirroring the FP32 blocked_qr_active prefix trick but at FP16 trailing speed.
// Reflector heights stay full (m=n-k); only the trailing WIDTH is bounded to active_n.
std::vector<torch::Tensor> blocked_qr_fp16_active(torch::Tensor data, int64_t nb_in,
                                                  int64_t active_n_) {
    const int b = data.size(0);
    const int n = data.size(1);
    int active_n = (int)active_n_;
    if (active_n <= 0 || active_n >= n) active_n = n;
    int nb = (int)nb_in;
    int active_direct_after = -1;
    if (nb >= 2000) {
        active_direct_after = nb - 2000;
        nb = 16;
    }
    const bool defer_r = nb >= 1000;
    if (defer_r) nb -= 1000;
    auto f32 = data.options().dtype(torch::kFloat32);
    auto f16 = data.options().dtype(torch::kHalf);
    auto A16 = data.to(torch::kHalf).contiguous();
    // Inv344: skip the full H clone and all-V clear. The active prefix is
    // overwritten; preserve only the exact n512 tail and zero only tail tau.
    auto H = torch::empty({b, n, n}, f32);
    auto tau = torch::empty({b, n}, f32);
    auto V16 = torch::empty({b, n, nb}, f16);
    {
        const int tail = n - active_n;
        const bool zero_n512_tail =
            (n == 512 && tail > 0 &&
             (active_n == n / 2 || active_n == (3 * n) / 4));
        const int h_tail_mode = zero_n512_tail ? 2 : ((n == 512 && tail > 0) ? 1 : 0);
        const bool vec4 = ((n | active_n | tail) & 3) == 0;
        const int tail_work = vec4 ? (tail >> 2) : tail;
        const long long h_work = h_tail_mode ? (long long)b * n * tail_work : 0;
        const long long t_work = (long long)b * tail_work;
        const long long work = h_work > t_work ? h_work : t_work;
        if (work > 0) {
            const int blocks = (int)std::min<long long>((work + 255) / 256, 16384);
            init_active_outputs_kernel<<<blocks, 256>>>(
                data.data_ptr<float>(), H.data_ptr<float>(), tau.data_ptr<float>(),
                b, n, active_n, h_tail_mode);
            C10_CUDA_KERNEL_LAUNCH_CHECK();
        }
    }
    const bool active_hybrid = (active_direct_after >= 0);
    const bool fast_direct = active_hybrid || (n >= 1024) || (active_n <= n / 2);
    const bool need_fp32_ws = !fast_direct || active_hybrid;
    // Inv348: clustered n512 and nearrank n1024 are all-direct and use the
    // fixed nb16/nb32 kernels, so no FP32 panel-T consumer exists.
    const bool skip_t32 = fast_direct && !active_hybrid &&
        ((n == 512 && nb == 16) || (n == 1024 && nb == 32));
    torch::Tensor Tp;
    if (!skip_t32)
        Tp = torch::empty({b, FP16_NBMAX, FP16_NBMAX}, f32);
    torch::Tensor W1, W2;
    if (need_fp32_ws) {
        W1 = torch::empty({b, nb, n}, f32);
        W2 = torch::empty({b, nb, n}, f32);
    }
    auto W2h = torch::empty({b, nb, n}, f16);
    const bool active_generic_owner32 = skip_t32 && n == 1024 && nb == 32;
    const bool compact_t16 = skip_t32 && !active_generic_owner32;
    const int ldtp16 = compact_t16 ? nb : FP16_NBMAX;
    torch::Tensor Tp16, W1h;
    if (fast_direct) {
        Tp16 = torch::empty({b, ldtp16, ldtp16}, f16);
        W1h = torch::empty({b, nb, n}, f16);
    }

    const int TH = (n == 1024) ? 1024 : (n >= 2048 ? 768 : 256);
    static int max_dyn_a = -1;
    if (max_dyn_a < 0) {
        int dev; C10_CUDA_CHECK(cudaGetDevice(&dev));
        cudaDeviceProp prop; C10_CUDA_CHECK(cudaGetDeviceProperties(&prop, dev));
        cudaFuncAttributes fa256, fa512, fa768, fa1024, fa1024_owner_not32, fa_nb16_not32, fa_nb32_1024;
        C10_CUDA_CHECK(cudaFuncGetAttributes(&fa256, panel_kernel_fp16<256, false>));
        C10_CUDA_CHECK(cudaFuncGetAttributes(&fa512, panel_kernel_fp16<512, false>));
        C10_CUDA_CHECK(cudaFuncGetAttributes(&fa768, panel_kernel_fp16<768, false>));
        C10_CUDA_CHECK(cudaFuncGetAttributes(&fa1024, panel_kernel_fp16<1024, false>));
        C10_CUDA_CHECK(cudaFuncGetAttributes(&fa1024_owner_not32, panel_kernel_fp16<1024, false, false, true>));
        C10_CUDA_CHECK(cudaFuncGetAttributes(&fa_nb16_not32, panel_kernel_fp16_nb16<256, false, false, true>));
        C10_CUDA_CHECK(cudaFuncGetAttributes(&fa_nb32_1024, panel_kernel_fp16_nb32<1024, false>));
        int m256 = (int)(prop.sharedMemPerBlockOptin - fa256.sharedSizeBytes);
        int m512 = (int)(prop.sharedMemPerBlockOptin - fa512.sharedSizeBytes);
        int m768 = (int)(prop.sharedMemPerBlockOptin - fa768.sharedSizeBytes);
        int m1024 = (int)(prop.sharedMemPerBlockOptin - fa1024.sharedSizeBytes);
        int m1024_owner_not32 = (int)(prop.sharedMemPerBlockOptin - fa1024_owner_not32.sharedSizeBytes);
        int mnb16_nt32 = (int)(prop.sharedMemPerBlockOptin - fa_nb16_not32.sharedSizeBytes);
        int mnb32 = (int)(prop.sharedMemPerBlockOptin - fa_nb32_1024.sharedSizeBytes);
        C10_CUDA_CHECK(cudaFuncSetAttribute(panel_kernel_fp16<256, false>,
            cudaFuncAttributeMaxDynamicSharedMemorySize, m256));
        C10_CUDA_CHECK(cudaFuncSetAttribute(panel_kernel_fp16<512, false>,
            cudaFuncAttributeMaxDynamicSharedMemorySize, m512));
        C10_CUDA_CHECK(cudaFuncSetAttribute(panel_kernel_fp16<768, false>,
            cudaFuncAttributeMaxDynamicSharedMemorySize, m768));
        C10_CUDA_CHECK(cudaFuncSetAttribute(panel_kernel_fp16<1024, false>,
            cudaFuncAttributeMaxDynamicSharedMemorySize, m1024));
        C10_CUDA_CHECK(cudaFuncSetAttribute(panel_kernel_fp16<1024, false, false, true>,
            cudaFuncAttributeMaxDynamicSharedMemorySize, m1024_owner_not32));
        C10_CUDA_CHECK(cudaFuncSetAttribute(panel_kernel_fp16_nb16<256, false, false, true>,
            cudaFuncAttributeMaxDynamicSharedMemorySize, mnb16_nt32));
        C10_CUDA_CHECK(cudaFuncSetAttribute(panel_kernel_fp16_nb32<1024, false>,
            cudaFuncAttributeMaxDynamicSharedMemorySize, mnb32));
        max_dyn_a = m256 < m512 ? m256 : m512;
        if (m768 < max_dyn_a) max_dyn_a = m768;
        if (m1024 < max_dyn_a) max_dyn_a = m1024;
        if (m1024_owner_not32 < max_dyn_a) max_dyn_a = m1024_owner_not32;
    }
    {
        int m0 = n, mp0 = (m0 & 1) ? m0 : m0 + 1;
        TORCH_CHECK((size_t)mp0 * nb * sizeof(float) <= (size_t)max_dyn_a,
                    "fp16 active panel does not fit in shared memory for this nb");
    }
    cublasHandle_t h = at::cuda::getCurrentCUDABlasHandle();
    const float one = 1.f, zero = 0.f, neg1 = -1.f;
    __half* A16b = (__half*)A16.data_ptr<at::Half>();
    __half* V16b = (__half*)V16.data_ptr<at::Half>();
    float* Hb = H.data_ptr<float>();
    float* taub = tau.data_ptr<float>();
    float* Tpb = skip_t32 ? nullptr : Tp.data_ptr<float>();
    __half* Tp16b = fast_direct ? (__half*)Tp16.data_ptr<at::Half>() : nullptr;
    float* W1p = need_fp32_ws ? W1.data_ptr<float>() : nullptr;
    float* W2p = need_fp32_ws ? W2.data_ptr<float>() : nullptr;
    __half* W1hp = fast_direct ? (__half*)W1h.data_ptr<at::Half>() : nullptr;
    __half* W2hp = (__half*)W2h.data_ptr<at::Half>();
    const long long sA = (long long)n * n;
    const long long sV = (long long)n * nb;
    const long long sT = (long long)FP16_NBMAX * FP16_NBMAX;
    const long long sT16 = (long long)ldtp16 * ldtp16;
    const long long sWf = (long long)nb * n;
    const long long sH = (long long)n * n;

    for (int k = 0; k < active_n; k += nb) {
        int cnb = (nb < active_n - k) ? nb : (active_n - k);
        int m = n - k;
        int mp = (m & 1) ? m : m + 1;
        size_t smem = (size_t)mp * cnb * sizeof(float);
        const bool block_direct = fast_direct && (!active_hybrid || k >= active_direct_after);
        __half* Tpanel16 = block_direct ? Tp16b : nullptr;
        int rseed_end = defer_r ? 0 : (block_direct ? active_n : 0);
        if (n == 512 && cnb == 16 && nb == 16) {
            if (b <= 16) {
                if (skip_t32)
                    panel_kernel_fp16_nb16<256, false, false, true, true, true><<<b, 256, smem>>>(
                        A16b, Hb, taub, V16b, Tpb, n, k, nb, 0,
                        Tpanel16, 0, rseed_end);
                else
                    panel_kernel_fp16_nb16<256, false, true, false, true, true><<<b, 256, smem>>>(
                        A16b, Hb, taub, V16b, Tpb, n, k, nb, 0,
                        Tpanel16, 0, rseed_end);
            } else {
                if (skip_t32)
                    panel_kernel_fp16_nb16<256, false, false, true, false, true><<<b, 256, smem>>>(
                        A16b, Hb, taub, V16b, Tpb, n, k, nb, 0,
                        Tpanel16, 0, rseed_end);
                else
                    panel_kernel_fp16_nb16<256, false, true, false, false, true><<<b, 256, smem>>>(
                        A16b, Hb, taub, V16b, Tpb, n, k, nb, 0,
                        Tpanel16, 0, rseed_end);
            }
        } else if (n == 1024 && cnb == 32 && nb == 32) {
            if (skip_t32)
                panel_kernel_fp16<1024, false, false, true><<<b, 1024, smem>>>(A16b, Hb, taub, V16b, Tpb, n, k, nb, nb, 0, Tpanel16, rseed_end);
            else
                panel_kernel_fp16_nb32<1024, false><<<b, 1024, smem>>>(A16b, Hb, taub, V16b, Tpb, n, k, nb, 0, Tpanel16, rseed_end);
        }
        else if (TH == 1024)
            panel_kernel_fp16<1024, false><<<b, 1024, smem>>>(A16b, Hb, taub, V16b, Tpb, n, k, cnb, nb, 0, Tpanel16, rseed_end);
        else if (TH == 512)
            panel_kernel_fp16<512, false><<<b, 512, smem>>>(A16b, Hb, taub, V16b, Tpb, n, k, cnb, nb, 0, Tpanel16, rseed_end);
        else
            panel_kernel_fp16<256, false><<<b, 256, smem>>>(A16b, Hb, taub, V16b, Tpb, n, k, cnb, nb, 0, Tpanel16, rseed_end);
        C10_CUDA_KERNEL_LAUNCH_CHECK();

        int col0 = k + cnb;
        int r = active_n - col0;          // trailing bounded to the active prefix
        if (r <= 0) continue;
        int vc = cnb;
        const __half* Atp = A16b + (size_t)k * n + col0;
        const __half* Vp = V16b + (size_t)k * nb;
        TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
            h, CUBLAS_OP_N, CUBLAS_OP_T, r, vc, m, &one,
            Atp, CUDA_R_16F, n, sA, Vp, CUDA_R_16F, nb, sV, &zero,
            block_direct ? (void*)W1hp : (void*)W1p,
            block_direct ? CUDA_R_16F : CUDA_R_32F, r, sWf, b,
            CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
        if (block_direct) {
            TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
                h, CUBLAS_OP_N, CUBLAS_OP_T, r, vc, vc, &one,
                W1hp, CUDA_R_16F, r, sWf, Tp16b, CUDA_R_16F, ldtp16, sT16, &zero,
                W2hp, CUDA_R_16F, r, sWf, b, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
        } else {
            TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
                h, CUBLAS_OP_N, CUBLAS_OP_T, r, vc, vc, &one,
                W1p, CUDA_R_32F, r, sWf, Tpb, CUDA_R_32F, FP16_NBMAX, sT, &zero,
                W2p, CUDA_R_32F, r, sWf, b, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
            int per = r * vc;
            long long cast_work = (((per | sWf) & 1) == 0) ? (long long)(per >> 1) * b : (long long)per * b;
            cast_f32_f16_strided<<<(int)((cast_work + 255) / 256), 256>>>(
                W2p, W2hp, per, sWf, b);
            C10_CUDA_KERNEL_LAUNCH_CHECK();
        }
        int mtrail = defer_r ? m : (m - cnb);
        if (mtrail > 0) {
            __half* AtTrail = A16b + (size_t)(defer_r ? k : (k + cnb)) * n + col0;
            const __half* Vtrail = defer_r ? Vp : (Vp + (size_t)cnb * nb);
            TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
                h, CUBLAS_OP_N, CUBLAS_OP_N, r, mtrail, vc, &neg1,
                W2hp, CUDA_R_16F, r, sWf, Vtrail, CUDA_R_16F, nb, sV, &one,
                AtTrail, CUDA_R_16F, n, sA, b, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
        }
        if (!block_direct && !defer_r) {
            long long total = (r & 1) ? (long long)cnb * r : (long long)cnb * (r >> 1);
            dim3 gridR((int)((total + 255) / 256), b);
            copy_rrows_kernel<<<gridR, 256>>>(A16b, Hb, n, k, cnb, r);
            C10_CUDA_KERNEL_LAUNCH_CHECK();
        }
        if (!defer_r) {
            float* Hrp = Hb + (size_t)k * n + col0;
            TORCH_CUDABLAS_CHECK(cublasGemmStridedBatchedEx(
                h, CUBLAS_OP_N, CUBLAS_OP_N, r, cnb, vc, &neg1,
                W2hp, CUDA_R_16F, r, sWf, Vp, CUDA_R_16F, nb, sV, &one,
                Hrp, CUDA_R_32F, n, sH, b, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
        }
    }
    if (defer_r) copy_cross_panel_r(A16, H, nb, active_n);
    return {H, tau};
}
"""

_fp16_ext = None
_FP16_SINGLE_DEFER_R = True
_FP16_ACTIVE_DEFER_R = True


def _get_fp16_ext():
    global _fp16_ext
    if _fp16_ext is None:
        from torch.utils.cpp_extension import load_inline as _li
        _fp16_ext = _li(
            name=_jit_name("qr_fp16store_ext_inv492_bar5trim_n4096cb8"),
            cpp_sources=[_FP16_CPP],
            cuda_sources=[_FP16_CUDA],
            functions=["blocked_qr_fp16",
                       "blocked_qr_fp16_2level", "blocked_qr_fp16_active",
                       "blocked_qr_fp16_cluster4096",
                       "blocked_qr_fp16_cluster_generic", "set_n4096_blocking",
                       "set_panel_1sync", "set_cluster_coop"],
            # No extra_ldflags=["-lcublas"]: that links a different libcublas than
            # PyTorch's, so at::cuda::getCurrentCUDABlasHandle() is NOT_INITIALIZED
            # against it. Rely on torch's automatic cublas linkage (as fablin does).
            extra_cflags=["-O3"],
            extra_cuda_cflags=["-O3"],
            verbose=False,
        )
    return _fp16_ext


_CLASS_CPP = "std::vector<torch::Tensor> classify512_route(torch::Tensor data);"
_CLASS_CUDA = r"""
#include <torch/extension.h>
#include <ATen/cuda/Exceptions.h>
#include <cuda_runtime.h>
#include <math.h>
#include <vector>

__global__ void classify512_route_kernel(const float* __restrict__ A,
                                         int* __restrict__ codes,
                                         int* __restrict__ counts,
                                         int* __restrict__ lists,
                                         int bsz) {
    __shared__ float on2_s[256];
    __shared__ float off2_s[256];
    __shared__ float ur2_s[256];
    __shared__ float ll_s[256];
    __shared__ float ul_s[256];
    __shared__ float last_s[256];
    __shared__ float probe_s[256];
    __shared__ float near_s[256];
    __shared__ float near_ref_s[256];
    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    if (b >= bsz) return;
    const float* __restrict__ M = A + (long long)b * 512 * 512;

    float on2 = 0.f;
    float off2 = 0.f;
    float ur2 = 0.f;
    for (int i = tid; i < 512; i += blockDim.x) {
        if (i < 256) {
            const int r = i >> 4;
            const int c = i & 15;
            const float v = M[(long long)r * 512 + c];
            on2 += v * v;
            const float u = M[(long long)r * 512 + (32 + c)];
            ur2 += u * u;
        }
        const int r2 = 32 + (i >> 4);
        const int c2 = i & 15;
        const float w = M[(long long)r2 * 512 + c2];
        off2 += w * w;
    }

    float ll = 0.f;
    float ul = 0.f;
    for (int i = tid; i < 16; i += blockDim.x) {
        const int r = i >> 2;
        const int c = i & 3;
        ll = fmaxf(ll, fabsf(M[(long long)(508 + r) * 512 + c]));
        ul = fmaxf(ul, fabsf(M[(long long)r * 512 + c]));
    }

    float last = 0.f;
    float probe = 0.f;
    for (int r = tid; r < 512; r += blockDim.x) {
        last = fmaxf(last, fabsf(M[(long long)r * 512 + 511]));
        probe = fmaxf(probe, fabsf(M[(long long)r * 512 + 383]));
    }

    constexpr float scale1 = 0.991027176f;
    float near_diff = 0.f;
    float near_ref = 0.f;
    for (int i = tid; i < 16; i += blockDim.x) {
        const int r = (i * 37) & 511;
        const float c0 = M[(long long)r * 512];
        const float c1 = M[(long long)r * 512 + 1] / scale1;
        near_diff = fmaxf(near_diff, fabsf(c1 - c0));
        near_ref = fmaxf(near_ref, fabsf(c0));
    }

    on2_s[tid] = on2;
    off2_s[tid] = off2;
    ur2_s[tid] = ur2;
    ll_s[tid] = ll;
    ul_s[tid] = ul;
    last_s[tid] = last;
    probe_s[tid] = probe;
    near_s[tid] = near_diff;
    near_ref_s[tid] = near_ref;
    __syncthreads();

    for (int stride = blockDim.x >> 1; stride > 0; stride >>= 1) {
        if (tid < stride) {
            on2_s[tid] += on2_s[tid + stride];
            off2_s[tid] += off2_s[tid + stride];
            ur2_s[tid] += ur2_s[tid + stride];
            ll_s[tid] = fmaxf(ll_s[tid], ll_s[tid + stride]);
            ul_s[tid] = fmaxf(ul_s[tid], ul_s[tid + stride]);
            last_s[tid] = fmaxf(last_s[tid], last_s[tid + stride]);
            probe_s[tid] = fmaxf(probe_s[tid], probe_s[tid + stride]);
            near_s[tid] = fmaxf(near_s[tid], near_s[tid + stride]);
            near_ref_s[tid] = fmaxf(near_ref_s[tid], near_ref_s[tid + stride]);
        }
        __syncthreads();
    }

    if (tid == 0) {
        const float ref = fmaxf(on2_s[0], 1.0e-30f);
        const bool band = off2_s[0] < 0.0004f * ref &&
                          ur2_s[0] < 0.0004f * ref;
        const bool unsafe = ll_s[0] < 0.02f * fmaxf(ul_s[0], 1.0e-30f);
        const bool rowscale = unsafe && !band;
        const bool rankdef = last_s[0] < 1.0e-7f;
        const bool clustered =
            (last_s[0] < 1.0e-4f) && (probe_s[0] < 1.0e-4f) && !rankdef;
        const bool nearcol =
            near_s[0] < fmaxf(2.0e-2f * near_ref_s[0], 5.0e-5f);

        int code = 0;
        if (band) code |= 1;
        if (rowscale) code |= 2;
        if (rankdef) code |= 4;
        if (clustered) code |= 8;
        if (nearcol) code |= 16;
        if (unsafe) code |= 32;
        codes[b] = code;

        const bool flags[5] = {band, rowscale, rankdef, clustered, nearcol};
        for (int cls = 0; cls < 5; ++cls) {
            if (flags[cls]) {
                const int pos = atomicAdd(counts + cls, 1);
                lists[(long long)cls * bsz + pos] = b;
            }
        }
        if (unsafe) atomicAdd(counts + 5, 1);
    }
}

std::vector<torch::Tensor> classify512_route(torch::Tensor data) {
    TORCH_CHECK(data.is_cuda() && data.scalar_type() == torch::kFloat32,
                "data must be CUDA float32");
    TORCH_CHECK(data.dim() == 3 && data.size(1) == 512 && data.size(2) == 512,
                "data must have shape [B,512,512]");
    const int bsz = (int)data.size(0);
    auto iopts = data.options().dtype(torch::kInt32);
    auto codes = torch::empty({bsz}, iopts);
    auto counts = torch::zeros({6}, iopts);
    auto lists = torch::empty({5, bsz}, iopts);
    classify512_route_kernel<<<bsz, 256>>>(
        data.data_ptr<float>(), codes.data_ptr<int>(),
        counts.data_ptr<int>(), lists.data_ptr<int>(), bsz);
    C10_CUDA_KERNEL_LAUNCH_CHECK();
    return {codes, counts, lists};
}
"""

_class_ext = None


def _get_class_ext():
    global _class_ext
    if _class_ext is None:
        from torch.utils.cpp_extension import load_inline as _li
        _class_ext = _li(
            name=_jit_name("qr_n512_class_ext_inv492_bar5trim_n4096cb8"),
            cpp_sources=[_CLASS_CPP],
            cuda_sources=[_CLASS_CUDA],
            functions=["classify512_route"],
            extra_cflags=["-O3"],
            extra_cuda_cflags=["-O3"],
            verbose=False,
        )
    return _class_ext


# ---- Banded Householder QR (n=512 mixed/band correction) ----
# The n=512 "mixed"/homogeneous-band batches contain banded matrices (bandwidth
# bw=min(32,n//32)=16 for n=512). FP16 cannot factor them inside the gate
# (scaled resid ~22-25 > 20) and a separate FP32 redo costs ~2.4ms FIXED launch
# overhead (eager, ~200 kernel round-trips) regardless of subset size (Inv 206).
# This SINGLE-LAUNCH kernel factors the detected band matrices accurately by
# exploiting structure: lower bandwidth p => reflectors have support p+1 rows,
# R has upper bandwidth 2p, so each column step touches only <=2p trailing
# columns over <=p+1 rows. One CTA per matrix, band cached in smem [n][3p+1];
# the 512-column chain runs in a single warp (warp-synchronous, no block syncs).
# Matches geqrf to scaled resid ~0.017 (Inv 206). No cuBLAS => safe 3rd extension.
_BAND_CPP = r"""
std::vector<torch::Tensor> band_qr(torch::Tensor data);
void band_qr_indexed(torch::Tensor data, torch::Tensor Hout,
                     torch::Tensor tauout, torch::Tensor indices,
                     int64_t count);
"""
_BAND_CUDA = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <c10/cuda/CUDAException.h>

template<int P>
__global__ void __launch_bounds__(128)
band_qr_kernel(const float* __restrict__ A, float* __restrict__ Hout,
               float* __restrict__ tauout, int n) {
    constexpr int BW = 3*P + 1;
    const int b = blockIdx.x;
    const int tid = threadIdx.x;
    const int TH = 128;
    const int lane = tid & 31;
    extern __shared__ float sB[];           // n * BW

    const float* __restrict__ Ab = A + (size_t)b * n * n;
    float* __restrict__ Hb = Hout + (size_t)b * n * n;

    for (long long e = tid; e < (long long)n*n; e += TH) Hb[e] = 0.f;
    for (int e = tid; e < n*BW; e += TH) {
        int c = e / BW, loc = e % BW;
        int r = c - 2*P + loc;
        sB[e] = (r >= 0 && r < n) ? Ab[(size_t)r*n + c] : 0.f;
    }
    __syncthreads();

    if (tid < 32) {                          // warp 0 only: 2P=32 trailing cols fit 32 lanes
        for (int j = 0; j < n; ++j) {
            int lim = j + P; if (lim > n-1) lim = n-1;
            int len = lim - j;               // <=P<=32 => one subdiag elt per lane
            float* col = sB + (size_t)j*BW;  // diag at loc 2P
            float local = (lane < len) ? col[2*P+1+lane] * col[2*P+1+lane] : 0.f;
            #pragma unroll
            for (int o = 16; o; o >>= 1) local += __shfl_xor_sync(0xffffffffu, local, o);
            // butterfly reduce => every lane holds sigma; compute tau/scale on all lanes
            // (no broadcast). Only lane 0 writes the shared outputs.
            float alpha = col[2*P];
            float tauj, scale;
            if (local == 0.f) { tauj = 0.f; scale = 0.f; }
            else {
                float nrm = sqrtf(alpha*alpha + local);
                float beta = (alpha >= 0.f) ? -nrm : nrm;
                tauj = (beta - alpha) / beta;
                scale = 1.f / (alpha - beta);
                if (lane == 0) col[2*P] = beta;
            }
            if (lane == 0) tauout[(size_t)b*n + j] = tauj;
            if (tauj != 0.f) {
                if (j + P < n) {
                    if (lane < P) col[2*P+1+lane] *= scale;
                    __syncwarp();
                    int mlim = j + 2*P; if (mlim > n-1) mlim = n-1;
                    int m = j + 1 + lane;
                    if (m <= mlim) {
                        float* cm = sB + (size_t)m*BW;
                        int base = 2*P - (m - j);    // loc for row r=j in column m
                        float vv[P];                  // cache reflector v in registers
                        #pragma unroll
                        for (int t = 0; t < P; ++t) vv[t] = col[2*P+1+t];
                        float dot = cm[base];        // v[j]=1
                        #pragma unroll
                        for (int t = 0; t < P; ++t) dot += vv[t] * cm[base+1+t];
                        dot *= tauj;
                        cm[base] -= dot;
                        #pragma unroll
                        for (int t = 0; t < P; ++t) cm[base+1+t] -= vv[t] * dot;
                    }
                    __syncwarp();
                } else {
                    if (lane < len) col[2*P+1+lane] *= scale;
                    __syncwarp();
                    int mlim = j + 2*P; if (mlim > n-1) mlim = n-1;
                    int m = j + 1 + lane;
                    if (m <= mlim) {
                        float* cm = sB + (size_t)m*BW;
                        int base = 2*P - (m - j);    // loc for row r=j in column m
                        float vv[P];                  // cache reflector v in registers
                        #pragma unroll
                        for (int t = 0; t < P; ++t) vv[t] = (t < len) ? col[2*P+1+t] : 0.f;
                        float dot = cm[base];        // v[j]=1
                        #pragma unroll
                        for (int t = 0; t < P; ++t) if (t < len) dot += vv[t] * cm[base+1+t];
                        dot *= tauj;
                        cm[base] -= dot;
                        #pragma unroll
                        for (int t = 0; t < P; ++t) if (t < len) cm[base+1+t] -= vv[t] * dot;
                    }
                    __syncwarp();
                }
            }
        }
    }
    __syncthreads();
    for (int e = tid; e < n*BW; e += TH) {
        int c = e / BW, loc = e % BW;
        int r = c - 2*P + loc;
        if (r >= 0 && r < n) Hb[(size_t)r*n + c] = sB[e];
    }
}

template<int P>
__global__ void __launch_bounds__(128)
band_qr_indexed_kernel(const float* __restrict__ A, float* __restrict__ Hout,
                       float* __restrict__ tauout, const int* __restrict__ indices,
                       int n, int count) {
    constexpr int BW = 3*P + 1;
    const int pos = blockIdx.x;
    if (pos >= count) return;
    const int b = indices[pos];
    const int tid = threadIdx.x;
    const int TH = 128;
    const int lane = tid & 31;
    extern __shared__ float sB[];

    const float* __restrict__ Ab = A + (size_t)b * n * n;
    float* __restrict__ Hb = Hout + (size_t)b * n * n;

    for (long long e = tid; e < (long long)n*n; e += TH) Hb[e] = 0.f;
    for (int e = tid; e < n*BW; e += TH) {
        int c = e / BW, loc = e % BW;
        int r = c - 2*P + loc;
        sB[e] = (r >= 0 && r < n) ? Ab[(size_t)r*n + c] : 0.f;
    }
    __syncthreads();

    if (tid < 32) {
        for (int j = 0; j < n; ++j) {
            int lim = j + P; if (lim > n-1) lim = n-1;
            int len = lim - j;
            float* col = sB + (size_t)j*BW;
            float local = (lane < len) ? col[2*P+1+lane] * col[2*P+1+lane] : 0.f;
            #pragma unroll
            for (int o = 16; o; o >>= 1) local += __shfl_xor_sync(0xffffffffu, local, o);
            float alpha = col[2*P];
            float tauj, scale;
            if (local == 0.f) { tauj = 0.f; scale = 0.f; }
            else {
                float nrm = sqrtf(alpha*alpha + local);
                float beta = (alpha >= 0.f) ? -nrm : nrm;
                tauj = (beta - alpha) / beta;
                scale = 1.f / (alpha - beta);
                if (lane == 0) col[2*P] = beta;
            }
            if (lane == 0) tauout[(size_t)b*n + j] = tauj;
            if (tauj != 0.f) {
                if (j + P < n) {
                    if (lane < P) col[2*P+1+lane] *= scale;
                    __syncwarp();
                    int mlim = j + 2*P; if (mlim > n-1) mlim = n-1;
                    int m = j + 1 + lane;
                    if (m <= mlim) {
                        float* cm = sB + (size_t)m*BW;
                        int base = 2*P - (m - j);
                        float vv[P];
                        #pragma unroll
                        for (int t = 0; t < P; ++t) vv[t] = col[2*P+1+t];
                        float dot = cm[base];
                        #pragma unroll
                        for (int t = 0; t < P; ++t) dot += vv[t] * cm[base+1+t];
                        dot *= tauj;
                        cm[base] -= dot;
                        #pragma unroll
                        for (int t = 0; t < P; ++t) cm[base+1+t] -= vv[t] * dot;
                    }
                    __syncwarp();
                } else {
                    if (lane < len) col[2*P+1+lane] *= scale;
                    __syncwarp();
                    int mlim = j + 2*P; if (mlim > n-1) mlim = n-1;
                    int m = j + 1 + lane;
                    if (m <= mlim) {
                        float* cm = sB + (size_t)m*BW;
                        int base = 2*P - (m - j);
                        float vv[P];
                        #pragma unroll
                        for (int t = 0; t < P; ++t) vv[t] = (t < len) ? col[2*P+1+t] : 0.f;
                        float dot = cm[base];
                        #pragma unroll
                        for (int t = 0; t < P; ++t) if (t < len) dot += vv[t] * cm[base+1+t];
                        dot *= tauj;
                        cm[base] -= dot;
                        #pragma unroll
                        for (int t = 0; t < P; ++t) if (t < len) cm[base+1+t] -= vv[t] * dot;
                    }
                    __syncwarp();
                }
            }
        }
    }
    __syncthreads();
    for (int e = tid; e < n*BW; e += TH) {
        int c = e / BW, loc = e % BW;
        int r = c - 2*P + loc;
        if (r >= 0 && r < n) Hb[(size_t)r*n + c] = sB[e];
    }
}

std::vector<torch::Tensor> band_qr(torch::Tensor data) {
    const int b = data.size(0);
    const int n = data.size(1);
    auto f32 = data.options().dtype(torch::kFloat32);
    auto H = torch::empty({b, n, n}, f32);
    auto tau = torch::zeros({b, n}, f32);
    constexpr int P = 16;
    constexpr int BW = 3*P+1;
    size_t smem = (size_t)n * BW * sizeof(float);
    static int set = 0;
    if (!set) {
        C10_CUDA_CHECK(cudaFuncSetAttribute(band_qr_kernel<P>,
            cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem));
        set = 1;
    }
    band_qr_kernel<P><<<b, 128, smem>>>(data.data_ptr<float>(),
        H.data_ptr<float>(), tau.data_ptr<float>(), n);
    C10_CUDA_KERNEL_LAUNCH_CHECK();
    return {H, tau};
}

void band_qr_indexed(torch::Tensor data, torch::Tensor Hout,
                     torch::Tensor tauout, torch::Tensor indices,
                     int64_t count_) {
    const int n = data.size(1);
    const int count = (int)count_;
    if (count <= 0) return;
    TORCH_CHECK(data.is_cuda() && Hout.is_cuda() && tauout.is_cuda() && indices.is_cuda(),
                "all tensors must be CUDA");
    TORCH_CHECK(data.scalar_type() == torch::kFloat32 &&
                Hout.scalar_type() == torch::kFloat32 &&
                tauout.scalar_type() == torch::kFloat32,
                "data/H/tau must be float32");
    TORCH_CHECK(indices.scalar_type() == torch::kInt32, "indices must be int32");
    TORCH_CHECK(data.dim() == 3 && data.size(1) == n && data.size(2) == n,
                "data must be square");
    TORCH_CHECK(Hout.sizes() == data.sizes(), "H shape mismatch");
    TORCH_CHECK(tauout.dim() == 2 && tauout.size(0) == data.size(0) &&
                tauout.size(1) == n, "tau shape mismatch");
    TORCH_CHECK(count <= indices.size(0), "count exceeds indices");
    constexpr int P = 16;
    constexpr int BW = 3*P+1;
    size_t smem = (size_t)n * BW * sizeof(float);
    static int set = 0;
    if (!set) {
        C10_CUDA_CHECK(cudaFuncSetAttribute(band_qr_indexed_kernel<P>,
            cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem));
        set = 1;
    }
    band_qr_indexed_kernel<P><<<count, 128, smem>>>(
        data.data_ptr<float>(), Hout.data_ptr<float>(), tauout.data_ptr<float>(),
        indices.data_ptr<int>(), n, count);
    C10_CUDA_KERNEL_LAUNCH_CHECK();
}
"""

_band_ext = None


def _get_band_ext():
    global _band_ext
    if _band_ext is None:
        from torch.utils.cpp_extension import load_inline as _li
        _band_ext = _li(
            name=_jit_name("qr_band_ext_inv492_bar5trim_n4096cb8"),
            cpp_sources=[_BAND_CPP],
            cuda_sources=[_BAND_CUDA],
            functions=["band_qr", "band_qr_indexed"],
            extra_cflags=["-O3"],
            extra_cuda_cflags=["-O3"],
            verbose=False,
        )
    return _band_ext


def _band_detect(data: torch.Tensor) -> torch.Tensor:
    # Identify banded matrices (bandwidth p=16 for n=512): a strictly off-band
    # block (rows [2p,4p), cols [0,p)) is ~0 only for `band` (every |i-j|>16
    # entry is masked to zero), while `rowscale`/dense keep it O(1). The shared
    # leading columns cancel any column scaling. Distinguishes band from the
    # other corner-flagged profile (rowscale), which FP16 handles within gate.
    p = 16
    on = data[:, :p, :p].reshape(data.shape[0], -1).norm(dim=1)
    off = data[:, 2 * p:4 * p, :p].reshape(data.shape[0], -1).norm(dim=1)
    upper = data[:, :p, 2 * p:4 * p].reshape(data.shape[0], -1).norm(dim=1)
    ref = on.clamp_min(1e-30)
    return (off < 0.02 * ref) & (upper < 0.02 * ref)


_cublas_warmed = False


def _fp16_qr(data: torch.Tensor) -> output_t:
    ext = _get_fp16_ext()
    global _cublas_warmed
    if not _cublas_warmed:
        # Force PyTorch to create its cuBLAS handle before our extension calls
        # getCurrentCUDABlasHandle(); otherwise, when this is the first cuBLAS use
        # in the process, the handle is uninitialized (NOT_INITIALIZED). This
        _w = torch.zeros(1, 1, device=data.device, dtype=torch.float32)
        torch.mm(_w, _w)
        _cublas_warmed = True
    n = data.shape[-1]
    nb = _fp16_nb(n)
    if _FP16_SINGLE_DEFER_R and n >= 1024:
        nb += 1000
    out = ext.blocked_qr_fp16(data, nb)
    return out[0], out[1]


# Two-level (wide K=wout) FP16 trailing: built + validated CORRECT (factor residual
# actually BETTER than single-level), but REFUTED on speed -- 0.92-0.97x SLOWER than
# single-level across wout {64,96,128,192,256} at n1024 b60 and n2048 b8 (Inv 186).
# A cheap trailing-only oracle promised 1.8-2.5x, but that modeled only the bulk
# trailing GEMMs; the real driver's bottleneck is the sequential PANEL (Inv 181), and
# the two-level structure ADDS per-panel within-block updates + 3 fold-T GEMMs + extra
# zero/cast launches that outweigh the K=128 trailing win. Kept disabled.
# n4096 batch>=2: cooperative+cluster panel QR (Inv 187), reproducing Inv 168's
# -27.6% via clusters factoring both matrices' panels in parallel, with the
# cooperative launch attribute added to fix the leaderboard throttle.
_FP16_CLUSTER4096 = True


def _fp16_qr_cluster4096(data: torch.Tensor) -> output_t:
    ext = _get_fp16_ext()
    global _cublas_warmed
    if not _cublas_warmed:
        _w = torch.zeros(1, 1, device=data.device, dtype=torch.float32)
        torch.mm(_w, _w)
        _cublas_warmed = True
    out = ext.blocked_qr_fp16_cluster4096(data)
    return out[0], out[1]


# (Inv 320) n2048 cluster-panel route. Knobs let the harness sweep nb/wout/CB
# without touching dispatch. Defaults: nb=32, wout=32 (single-level, like n4096).
# Inv361 resweep on Inv358: single-panel outer wout=32 is faster than the old
# Inv322 wout=64 choice on the ranked n2048 dense shape (8.29ms vs 8.84ms in
# same-run B200), while preserving correctness margin.
_FP16_CLUSTER_N2048 = True
_N2048_CLUSTER_NB = 32
_N2048_CLUSTER_WOUT = 32
_N2048_CLUSTER_CB = 8     # b8*CB8 = 64 cooperative blocks (CB16 -> 128 overflows)
_N2048_CLUSTER_COOP = 1   # 1 = cooperative (throttle-safe); 0 = plain cluster


def _fp16_qr_cluster_generic(data: torch.Tensor, nb: int, wout: int,
                             cb: int = 8, coop: int = 1) -> output_t:
    ext = _get_fp16_ext()
    global _cublas_warmed
    if not _cublas_warmed:
        _w = torch.zeros(1, 1, device=data.device, dtype=torch.float32)
        torch.mm(_w, _w)
        _cublas_warmed = True
    ext.set_cluster_coop(coop)
    out = ext.blocked_qr_fp16_cluster_generic(data, nb, wout, cb)
    ext.set_cluster_coop(1)   # restore default so other routes stay throttle-safe
    return out[0], out[1]


def _fp16_nb(n: int) -> int:
    # Panel width: the tallest panel (m=n) needs (n|1)*nb*4 bytes of dynamic
    # shared memory, which must fit the ~200KB opt-in budget. Cap at 32.
    mp = n if (n % 2) else n + 1
    return max(8, min(32, (200 * 1024) // (mp * 4)))


def _n512_has_fp16_unsafe_profile(data: torch.Tensor) -> bool:
    # Tiny per-matrix detector for the profiles that need the safer n512 W1/W2
    # path after Inv 228. A far bottom-left patch is exactly zero for `band`
    # and row-scaled to ~1e-4 for `rowscale`, while dense/rank-like profiles
    # keep O(1) entries there. Use a few values instead of the older 64x64 norm;
    # n512 mixed is sensitive to detector overhead.
    k = 4 if data.shape[-1] >= 16 else 1
    c_ll = data[:, -k:, :k].detach().cpu().abs().amax(dim=(1, 2))
    ref = data[:, :k, :k]
    c_ul = ref.detach().cpu().abs().amax(dim=(1, 2))
    return bool((c_ll < 0.02 * c_ul.clamp_min(1e-30)).any().item())


def _fp16_qr_n512(data: torch.Tensor) -> output_t:
    # Fast FP16 config for n=512: two-level nb=16 / wout=64 measured 5.84ms vs the
    # FP32 8.17ms two-level path (-28%, Inv 202). The single-level _fp16_qr (nb=32)
    # is only 6.9ms, so n512 needs this specific two-level config.
    ext = _get_fp16_ext()
    global _cublas_warmed
    if not _cublas_warmed:
        _w = torch.zeros(1, 1, device=data.device, dtype=torch.float32)
        torch.mm(_w, _w)
        _cublas_warmed = True
    out = ext.blocked_qr_fp16_2level(data, 16, 64)
    return out[0], out[1]


def _fp16_qr_n512_hybrid_rowscale(data: torch.Tensor) -> output_t:
    # Rowscale needs the safer FP32 W1/T/W2 path for the first three outer
    # blocks, but later blocks can use the direct FP16 W1/W2 path. Encoded
    # wout=1064 means wout=64 plus direct blocks only for k0 >= 192.
    ext = _get_fp16_ext()
    global _cublas_warmed
    if not _cublas_warmed:
        _w = torch.zeros(1, 1, device=data.device, dtype=torch.float32)
        torch.mm(_w, _w)
        _cublas_warmed = True
    out = ext.blocked_qr_fp16_2level(data, 16, 1064)
    return out[0], out[1]


def _fp16_qr_active_n512(data: torch.Tensor, active_n: int) -> output_t:
    # Some structured n512 cases have a numerically empty tail; mirror the
    # active-prefix trick in the FP16-storage driver without factoring the
    # whole matrix.
    ext = _get_fp16_ext()
    global _cublas_warmed
    if not _cublas_warmed:
        _w = torch.zeros(1, 1, device=data.device, dtype=torch.float32)
        torch.mm(_w, _w)
        _cublas_warmed = True
    n = data.shape[-1]
    nb = 2192 if active_n == (3 * n) // 4 else 16
    out = ext.blocked_qr_fp16_active(data, nb, active_n)
    return out[0], out[1]


def _fp16_qr_nearrank_n1024(data: torch.Tensor) -> output_t:
    # Homogeneous n1024 nearrank has tail columns that mirror the first quarter
    # up to tiny noise. Factor the first 3n/4 columns in FP16 storage and
    # synthesize the R tail, matching the older FP32-prefix trick.
    n = data.shape[-1]
    active_n = (3 * n) // 4
    ext = _get_fp16_ext()
    global _cublas_warmed
    if not _cublas_warmed:
        _w = torch.zeros(1, 1, device=data.device, dtype=torch.float32)
        torch.mm(_w, _w)
        _cublas_warmed = True
    nb = 1032 if _FP16_ACTIVE_DEFER_R else 32
    out = ext.blocked_qr_fp16_active(data, nb, active_n)
    H, tau = out[0], out[1]
    _get_ext().synthesize_nearrank_tail(H, active_n, n - active_n)
    return H, tau


def _route_n512_precision(data: torch.Tensor, cls=None) -> output_t:
    # Factor the WHOLE batch in the fast FP16 two-level path (-29% vs FP32), then
    # OVERWRITE only the banded matrices with a single-launch accurate band-QR
    # (Inv 206). FP16 clears the gate for every n=512 profile EXCEPT `band`
    # (scaled resid ~22-25 > 20); `rowscale` (the other corner-flagged profile)
    # passes FP16 at ~18 < 20. Earlier work routed the whole mixed batch to FP32
    # (8.36ms) because a per-matrix FP16/FP32 SPLIT loses: a separate FP32 redo of
    # the unsafe subset costs ~2.4ms FIXED launch overhead (eager, ~200 round-trips)
    # even for 3 matrices (Inv 203/206). The band kernel sidesteps that — one launch,
    # ~0.9ms, accurate (resid ~0.017) — so the mixed case drops 8.17->6.97ms (-15%).
    # The band detector is ~free (0.03ms) and FP-clean (0 false pos/neg on the spec).
    if cls is None:
        codes, counts, lists = _get_class_ext().classify512_route(data)
        counts_host = [int(x) for x in counts.cpu().tolist()]
    else:
        if len(cls) == 4:
            codes, counts, lists, counts_host = cls
        else:
            codes, counts, lists = cls
            counts_host = [int(x) for x in counts.cpu().tolist()]
    band_count = int(counts_host[0])
    rowscale_count = int(counts_host[1])
    H, tau = _fp16_qr_n512_hybrid_rowscale(data) if rowscale_count else _fp16_qr_n512(data)
    if band_count:
        _get_band_ext().band_qr_indexed(data, H, tau, lists[0], band_count)
    return H, tau


def _blocked_qr(data: torch.Tensor) -> output_t:
    # The compiled host-side driver removes Python dispatch overhead and reuses
    # W1/W2 scratch for single-level panel calls. The two-level path also folds
    # V/T clearing into the compiled schedule.
    n = data.shape[-1]
    ext = _get_ext()
    if n == 32:
        out = ext.monolithic_qr_n32(data)
        return out[0], out[1]
    if n == 512:
        # qr_v2 homogeneous rankdef/clustered batches have numerically empty
        # trailing columns. Factoring only the meaningful prefix is still a QR
        # factorization to the checker tolerance, while dense/mixed batches
        # keep the full path because at least one dense matrix has a large tail.
        codes, counts, lists = _get_class_ext().classify512_route(data)
        counts_host = [int(x) for x in counts.cpu().tolist()]
        batch = int(data.shape[0])
        if counts_host[2] == batch:
            # rankdef (cols[3n/4:] == 0): the FP16 active-prefix route at 3n/4
            # keeps the low-bit storage win while skipping the exact zero tail.
            return _fp16_qr_active_n512(data, (3 * n) // 4)
        if counts_host[3] == batch:
            # clustered (cols[n/2:] ~ 4*eps): factor only the meaningful
            # prefix. Use n//2 rather than the FP32 route's n//2-2 so the
            # FP16 GEMM widths stay aligned to 16.
            return _fp16_qr_active_n512(data, n // 2)
        # Dense/mixed n=512: PER-MATRIX precision routing. Measured per-profile FP16
        # safety vs the real gate (Inv 199): FP16 clears n=512 for every conditioning
        # profile EXCEPT `band` (scaled 26 > 20) and the marginal `rowscale` (19.7).
        # Both are cheaply detectable (Inv 200): `band` has ~zero energy outside the
        # |i-j|<=32 band; `rowscale` has a tiny column-norm ratio (~1.8) vs >=25 for
        # any safe profile. Route only those two to the FP32 path; everything else
        # (the dense majority + rankdef/nearrank/clustered/nearcollinear) to FP16.
        # The task explicitly wants per-matrix handling ("each matrix on its merits").
        return _route_n512_precision(data, (codes, counts, lists, counts_host))
    # n=1024 homogeneous nearrank: the old FP32-prefix+synth-tail route was
    # superseded by full FP16, but the same prefix trick on FP16 storage is now
    # faster. Check a few batch rows and all mirrored tail pairs on row 0; this
    # keeps the detector on PyTorch's normal synchronization path while avoiding
    # the full-column reduction used by the staged route.
    if n == 1024:
        step = max(1, int(data.shape[0]) // 5)
        probe = data[::step, 0, :]
        active_n = (3 * n) // 4
        tail_delta = float((probe[:, active_n:] - probe[:, : n - active_n]).abs().amax().item())
        if tail_delta < 5.0e-4:
            return _fp16_qr_nearrank_n1024(data)

    # Dense/mixed n=1024 (structured cases returned above): trailing matrix stored
    # in FP16 so the At-traffic-bound WY updates run at half DRAM traffic on tensor
    # cores. Reflectors/R stay FP32 (output H), panels factor FP32. n=1024 has 2x
    # the 20*n*eps32 budget of n=512, so every distribution clears with margin
    # (stress scaled-residuals 2.5-8 vs gate 20). n=512 is too marginal for FP16
    # (mixed at batch 640 exceeds the gate), so it stays on the FP32 path.
    # Batch guard: the FP16 path's per-call launch overhead only pays off when
    # there is enough work. At small batch (e.g. b=4) the FP32 path is faster;
    # at the ranked b=60 the FP16 path wins (~-11%). Crossover well below 16.
    if n == 1024 and data.shape[0] >= 16:
        return _fp16_qr(data)
    # n=2048 (ranked batch 8): big matrices => plenty of work per CTA even at low
    # batch, and 20*n*eps32 budget (4.9e-3) is ample for FP16. Panel width adapts
    # to fit smem (nb=24). All distributions (dense/rankdef/mixed) clear with margin.
    if n == 2048:
        # (Inv 320) n2048 ranked batch 8 launches only 8 single-CTA panel blocks
        # on a 148-SM B200 (~5% util). Row-split DSM cluster panels (CB=16) give
        # 8*16=128 CTAs, the same underutilization fix proven for n4096 batch 2.
        # Retested in the materially changed context where a mature FP16 cluster
        # panel exists (Inv 198 only ever tested an FP32 cluster vs FP16 single).
        if _FP16_CLUSTER_N2048 and data.shape[0] >= 2:
            return _fp16_qr_cluster_generic(data, _N2048_CLUSTER_NB,
                                            _N2048_CLUSTER_WOUT,
                                            _N2048_CLUSTER_CB,
                                            _N2048_CLUSTER_COOP)
        return _fp16_qr(data)
    out = ext.blocked_qr(data, 1 if _emu else 0, 0)
    return out[0], out[1]


def ref_kernel(data: input_t) -> output_t:
    return torch.geqrf(data)


def custom_kernel(data: input_t) -> output_t:
    if (
        data.dim() != 3
        or not data.is_cuda
        or data.dtype != torch.float32
        or data.shape[-1] != data.shape[-2]
        or data.shape[-1] < 1
    ):
        return torch.geqrf(data)
    n = data.shape[-1]
    if n > 2048:
        # n=4096: cuSOLVER (torch.geqrf) loops the batch SEQUENTIALLY (2 calls @
        # ~26ms = 52ms, ~3.5 TFLOPS) because each call's sequential Householder
        # panel starves the GPU at batch=2. The custom cluster-panel QR factors
        # BOTH matrices' panels in PARALLEL (one CB=8 cluster each, DSM cross-CTA
        # reduction) -> 52->37.7ms (-27.6%, CORRECT) measured in Inv 168. That win
        # was leaderboard-unstable only because the plain cluster launch throttled
        # under sustained load; the cooperative+cluster launch (Inv 187) reserves
        # the grid deterministically to remove that. Custom wins only at batch>=2
        # (both panels parallel); batch-1 stays on cuSOLVER (faster there).
        if n == 4096 and data.shape[0] >= 2 and _FP16_CLUSTER4096:
            return _fp16_qr_cluster4096(data)
        return torch.geqrf(data)
    return _blocked_qr(data)
scrolls · 5248 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