Skip to content
KernelIndex
Search⌘K

submission 858389

ozamatash · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-eigh-858389?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
17.5ms
#28 of 286
2026-07-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:ba915d913b41d44e403034eda71e141c16355424a92cc2b2e4eb0bf21009f8e5
license declaredunknown
license concludedunknown
authorsozamatash
imported2026-08-26

Techniques

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

clustercluster.sync();
fused-epiloguevoid launch_trail_epilogue(float* A, __half* Ah, const float* P,
mbarrierasm volatile("mbarrier.init.shared::cta.b64 [%0], 1;" :: "r"(smem_u32(p)));
mmatmp = tl.dot(v1t, t)
num-warps = 4a, vv, pairs, active, _N512, _TB512, _P2_512, prec, num_warps=4)
shared-memory__shared__ float sIJ[TES_TS][TES_TS + 1];
tmaasm volatile("cp.async.bulk.shared::cluster.global.mbarrier::complete_tx::bytes [%0], [%1], %2, [%3];"
vector-width = float4float4 a4 = *reinterpret_cast<const float4*>(&A[aidx]);
warp-specializationstatic constexpr int F_ROT = 4 * 16 * 2; // per-producer float2 rots

Kernel source

submission.py6785 lines
#!POPCORN leaderboard eigh
#!POPCORN gpu B200

# GENERATED FILE — submission.py is produced by make_submission.py, which
# inlines csrc/bindings.cpp and csrc/kernels.cu into this template. Edit
# those sources (and submission_template.py) instead, then regenerate:
#   python make_submission.py

from pathlib import Path

import torch
import triton
import triton.language as tl
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline

CPP_SRC = r"""
#include <ATen/ATen.h>
#include <ATen/cuda/CUDAContextLight.h>
#include <cuda_fp16.h>
#include <cublasLt.h>
#include <cstdint>
#include <torch/library.h>
#include <optional>
#include <vector>
#include <unordered_map>

void launch_jacobi_smem(
    const float* input,
    float* q,
    float* l,
    int batch,
    int n,
    float tol2,
    int max_sweeps,
    int* sweep_count);

void launch_steqr_tri(
    const double* d,
    const double* e,
    float* q,
    double* l,
    int P,
    int n,
    int use_f32,
    int do_sort);

void launch_jacobi_cluster(
    const float* input,
    float* q,
    float* l,
    int batch,
    int n,
    float tol2,
    int max_sweeps,
    int* sweep_count);

void launch_seated_pivot_solve(
    const void* a,
    bool a_is_half,
    void* v,
    const int* pairs,
    const int* active,
    int batch,
    int n,
    int max_sweeps,
    float tol2);

void launch_block_jacobi(
    const float* input,
    float* q,
    float* l,
    float* a_work,
    float* qt_work,
    int batch,
    int n,
    float tol2,
    int max_sweeps,
    int* sweep_count);

namespace dc {
void launch_dc_prep(
    const double* Dl, const double* Drr, const float* zL, const float* zR,
    const double* rho_in, double* dc_out, double* zc_out, int* k_out,
    double* rho_out, int* invorder, int* perm_out, double* hv_out,
    double* hbeta_out, int* segstart, int* nseg_out, int P, int h, int m,
    double defl_zk, double defl_gk);
void launch_merge_build(
    const double* dc_in, const double* zc_in, const int* k_in,
    const double* rho_in, const int* invorder, const int* perm,
    const double* hv_in, const double* hbeta_in, const int* segstart,
    const int* nseg_in, float* Vchild, double* Lam_out,
    int P, int m, int newt, int nf32, double res_tol, double step_tol);
void launch_merge_build_multi(
    const double* dc_in, const double* zc_in, const int* k_in,
    const double* rho_in, const int* invorder, const int* perm,
    const double* hv_in, const double* hbeta_in, const int* segstart,
    const int* nseg_in, float* Vchild, double* Lam_out,
    int P, int m, int newt, int nf32, int G);
}

namespace os1 {
void launch_latrd(const float* A, const __half* Ah, float* Vout, float* Wout,
                  float* dout, float* eout, int b, int n, int p0);
void launch_latrd_tail(const float* A, float* Vout, float* dout, float* eout,
                       int b, int n, int p0, int m0);
void launch_trec(const float* G, float* T, int b, int nb);
}

namespace os1cl {
void launch_latrd_cluster(const float* A, const __half* Ah, float* V, float* W, float* d, float* e,
                          int b, int n, int p0, int C, int nb);
}

void launch_split16(const float* in, __half* hi, __half* lo,
                    long total, int Y, int XY, long bs, long rs, long cs);

namespace {

// ---------------------------------------------------------------------------
// cublasLt descriptor cache.  The D&C fp16x3 back-transform (and the tf32/fp16
// baddbmm helpers) issue a storm of tiny batched GEMMs whose (dtype, shape,
// stride, batch) tuples repeat across the 3 error-corrected passes, across the
// top/bot products of a merge, across levels, and across benchmark calls.
// Creating+destroying 4 cublasLtMatrixLayout_t + 1 cublasLtMatmulDesc_t per
// GEMM was measured as host-side dead time between the tiny kernels.  We cache
// the immutable layouts/descriptors (they carry no data pointer, only shape),
// keyed by their defining attributes, and never destroy them (leaked at
// process exit — fine for a single-process solver, single host thread).
// ---------------------------------------------------------------------------
struct LtLayoutKey {
  int dtype;
  int order;
  int64_t rows, cols, ld, batch, batch_stride;
  bool operator==(const LtLayoutKey& o) const {
    return dtype == o.dtype && order == o.order && rows == o.rows &&
           cols == o.cols && ld == o.ld && batch == o.batch &&
           batch_stride == o.batch_stride;
  }
};
struct LtLayoutKeyHash {
  size_t operator()(const LtLayoutKey& k) const {
    size_t h = 1469598103934665603ULL;
    auto mix = [&](int64_t v) {
      h ^= static_cast<size_t>(v);
      h *= 1099511628211ULL;
    };
    mix(k.dtype);
    mix(k.order);
    mix(k.rows);
    mix(k.cols);
    mix(k.ld);
    mix(k.batch);
    mix(k.batch_stride);
    return h;
  }
};

cublasLtMatrixLayout_t make_lt_layout(
    const at::Tensor& t,
    cudaDataType_t dtype) {
  TORCH_CHECK(t.dim() == 3);

  const int batch = static_cast<int>(t.size(0));
  const int64_t rows = t.size(1);
  const int64_t cols = t.size(2);

  cublasLtOrder_t order;
  int64_t ld;

  if (t.stride(2) == 1) {
    order = CUBLASLT_ORDER_ROW;
    ld = t.stride(1);
  } else if (t.stride(1) == 1) {
    order = CUBLASLT_ORDER_COL;
    ld = t.stride(2);
  } else {
    TORCH_CHECK(false, "tensor must be row-major or column-major, strides=", t.strides());
  }

  const int64_t batch_stride = t.stride(0);
  LtLayoutKey key{static_cast<int>(dtype), static_cast<int>(order),
                  rows, cols, ld, static_cast<int64_t>(batch), batch_stride};
  static std::unordered_map<LtLayoutKey, cublasLtMatrixLayout_t, LtLayoutKeyHash> cache;
  auto it = cache.find(key);
  if (it != cache.end()) return it->second;

  cublasLtMatrixLayout_t layout = nullptr;
  auto status = cublasLtMatrixLayoutCreate(&layout, dtype, rows, cols, ld);
  TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, "layout create failed: ", status);

  status = cublasLtMatrixLayoutSetAttribute(
      layout, CUBLASLT_MATRIX_LAYOUT_ORDER, &order, sizeof(order));
  TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, "set order failed: ", status);

  status = cublasLtMatrixLayoutSetAttribute(
      layout, CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT, &batch, sizeof(batch));
  TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, "set batch count failed: ", status);

  status = cublasLtMatrixLayoutSetAttribute(
      layout,
      CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET,
      &batch_stride,
      sizeof(batch_stride));
  TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, "set batch stride failed: ", status);

  cache.emplace(key, layout);
  return layout;
}

cublasLtMatmulDesc_t get_matmul_desc(cublasComputeType_t compute_type) {
  static std::unordered_map<int, cublasLtMatmulDesc_t> cache;
  const int key = static_cast<int>(compute_type);
  auto it = cache.find(key);
  if (it != cache.end()) return it->second;
  cublasLtMatmulDesc_t op = nullptr;
  auto status = cublasLtMatmulDescCreate(&op, compute_type, CUDA_R_32F);
  TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, "matmul desc create failed: ", status);
  cache.emplace(key, op);
  return op;
}

void lt_baddbmm_out(
    const at::Tensor& input,
    const at::Tensor& left,
    const at::Tensor& right,
    at::Tensor& output,
    cudaDataType_t ab_dtype,
    cublasComputeType_t compute_type,
    float alpha,
    float beta) {
  TORCH_CHECK(input.dim() == 3);
  TORCH_CHECK(left.dim() == 3);
  TORCH_CHECK(right.dim() == 3);
  TORCH_CHECK(output.dim() == 3);

  TORCH_CHECK(left.size(0) == right.size(0));
  TORCH_CHECK(left.size(0) == input.size(0));
  TORCH_CHECK(left.size(0) == output.size(0));
  TORCH_CHECK(left.size(2) == right.size(1));

  TORCH_CHECK(input.size(1) == left.size(1));
  TORCH_CHECK(input.size(2) == right.size(2));
  TORCH_CHECK(output.size(1) == input.size(1));
  TORCH_CHECK(output.size(2) == input.size(2));

  TORCH_CHECK(input.dtype() == at::kFloat);
  TORCH_CHECK(output.dtype() == at::kFloat);

  cublasLtHandle_t handle = at::cuda::getCurrentCUDABlasLtHandle();

  cublasLtMatmulDesc_t op = get_matmul_desc(compute_type);

  auto a_layout = make_lt_layout(left, ab_dtype);
  auto b_layout = make_lt_layout(right, ab_dtype);
  auto c_layout = make_lt_layout(input, CUDA_R_32F);
  auto d_layout = make_lt_layout(output, CUDA_R_32F);

  auto status = cublasLtMatmul(
      handle,
      op,
      &alpha,
      left.data_ptr(),
      a_layout,
      right.data_ptr(),
      b_layout,
      &beta,
      input.data_ptr<float>(),
      c_layout,
      output.data_ptr<float>(),
      d_layout,
      nullptr,
      nullptr,
      0,
      0);

  TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, "cublasLtMatmul failed: ", status);
}

} // namespace

void tf32_baddbmm_out(
    const at::Tensor& input,
    const at::Tensor& left,
    const at::Tensor& right,
    at::Tensor& output,
    double beta,
    double alpha) {
  TORCH_CHECK(left.dtype() == at::kFloat);
  TORCH_CHECK(right.dtype() == at::kFloat);

  lt_baddbmm_out(
      input,
      left,
      right,
      output,
      CUDA_R_32F,
      CUBLAS_COMPUTE_32F_FAST_TF32,
      static_cast<float>(alpha),
      static_cast<float>(beta));
}

void fp16_baddbmm_out(
    const at::Tensor& input,
    const at::Tensor& left,
    const at::Tensor& right,
    at::Tensor& output,
    double beta,
    double alpha) {
  TORCH_CHECK(left.dtype() == at::kHalf);
  TORCH_CHECK(right.dtype() == at::kHalf);

  lt_baddbmm_out(
      input,
      left,
      right,
      output,
      CUDA_R_16F,
      CUBLAS_COMPUTE_32F,
      static_cast<float>(alpha),
      static_cast<float>(beta));
}

void split16(
    const at::Tensor& x,
    at::Tensor& hi,
    at::Tensor& lo) {
  TORCH_CHECK(x.dim() == 3);
  TORCH_CHECK(x.dtype() == at::kFloat && x.is_cuda());
  TORCH_CHECK(hi.dtype() == at::kHalf && hi.is_contiguous());
  TORCH_CHECK(lo.dtype() == at::kHalf && lo.is_contiguous());
  const int P = static_cast<int>(x.size(0));
  const int X = static_cast<int>(x.size(1));
  const int Y = static_cast<int>(x.size(2));
  TORCH_CHECK(hi.size(0) == P && hi.size(1) == X && hi.size(2) == Y);
  TORCH_CHECK(lo.size(0) == P && lo.size(1) == X && lo.size(2) == Y);
  const long total = static_cast<long>(P) * X * Y;
  launch_split16(
      x.data_ptr<float>(),
      reinterpret_cast<__half*>(hi.data_ptr<at::Half>()),
      reinterpret_cast<__half*>(lo.data_ptr<at::Half>()),
      total, Y, X * Y,
      x.stride(0), x.stride(1), x.stride(2));
}

void launch_trail_epilogue(float* A, __half* Ah, const float* P,
                           int b, int n, int off, int tm);

// Fused trailing-update epilogue: A[:, off:, off:] -= P; Ah[:, off:, off:] = half.
// A/Ah are the FULL (b,n,n) contiguous matrices; P is (b,tm,tm) contiguous with
// tm = n - off. Replaces the torch subtract + writeback + half-cast chain.
void trail_epilogue(
    at::Tensor& A,
    at::Tensor& Ah,
    const at::Tensor& P,
    int64_t off) {
  TORCH_CHECK(A.dim() == 3 && A.dtype() == at::kFloat && A.is_cuda() && A.is_contiguous());
  TORCH_CHECK(Ah.dtype() == at::kHalf && Ah.is_contiguous());
  TORCH_CHECK(P.dtype() == at::kFloat && P.is_contiguous());
  const int b = static_cast<int>(A.size(0));
  const int n = static_cast<int>(A.size(1));
  const int tm = static_cast<int>(P.size(1));
  TORCH_CHECK(P.size(0) == b && P.size(2) == tm);
  TORCH_CHECK((int)off + tm == n);
  launch_trail_epilogue(
      A.data_ptr<float>(),
      reinterpret_cast<__half*>(Ah.data_ptr<at::Half>()),
      P.data_ptr<float>(),
      b, n, (int)off, tm);
}

void launch_trail_epilogue_sym(float* A, __half* Ah, const float* Q,
                               int b, int n, int off, int tm);

// Symmetric fused trailing-update epilogue: A[:, off:, off:] -= (Q + Q^T);
// Ah[:, off:, off:] = half. Q = V W^T is the single-side (K=nb, no cat) rank-2b
// product; the kernel adds Q^T on the fly (shared-memory tile transpose).
void trail_epilogue_sym(
    at::Tensor& A,
    at::Tensor& Ah,
    const at::Tensor& Q,
    int64_t off) {
  TORCH_CHECK(A.dim() == 3 && A.dtype() == at::kFloat && A.is_cuda() && A.is_contiguous());
  TORCH_CHECK(Ah.dtype() == at::kHalf && Ah.is_contiguous());
  TORCH_CHECK(Q.dtype() == at::kFloat && Q.is_contiguous());
  const int b = static_cast<int>(A.size(0));
  const int n = static_cast<int>(A.size(1));
  const int tm = static_cast<int>(Q.size(1));
  TORCH_CHECK(Q.size(0) == b && Q.size(2) == tm);
  TORCH_CHECK((int)off + tm == n);
  launch_trail_epilogue_sym(
      A.data_ptr<float>(),
      reinterpret_cast<__half*>(Ah.data_ptr<at::Half>()),
      Q.data_ptr<float>(),
      b, n, (int)off, tm);
}

void launch_prescale(const float* data, float* A, __half* Ah, float* scale,
                     int b, int n);

// Fused amax-prescale: scale = amax(|data|) per matrix (clamped 1e-30);
// A = data/scale; Ah = A.half(). Replaces abs-temp + amax + divide + half-cast.
void prescale(
    const at::Tensor& data,
    at::Tensor& A,
    at::Tensor& Ah,
    at::Tensor& scale) {
  TORCH_CHECK(data.dim() == 3 && data.dtype() == at::kFloat && data.is_cuda() && data.is_contiguous());
  TORCH_CHECK(A.dtype() == at::kFloat && A.is_contiguous());
  TORCH_CHECK(Ah.dtype() == at::kHalf && Ah.is_contiguous());
  TORCH_CHECK(scale.dtype() == at::kFloat && scale.is_contiguous());
  const int b = static_cast<int>(data.size(0));
  const int n = static_cast<int>(data.size(1));
  launch_prescale(
      data.data_ptr<float>(),
      A.data_ptr<float>(),
      reinterpret_cast<__half*>(Ah.data_ptr<at::Half>()),
      scale.data_ptr<float>(),
      b, n);
}

void jacobi_smem(
    const at::Tensor& input,
    at::Tensor& q,
    at::Tensor& l,
    double tol2,
    int64_t max_sweeps,
    std::optional<at::Tensor> sweep_count) {
  TORCH_CHECK(input.dim() == 3);
  const int batch = static_cast<int>(input.size(0));
  const int n = static_cast<int>(input.size(1));
  TORCH_CHECK(input.size(2) == n);
  TORCH_CHECK(input.is_cuda());
  TORCH_CHECK(input.dtype() == at::kFloat);
  TORCH_CHECK(input.is_contiguous());
  TORCH_CHECK(q.sizes() == input.sizes());
  TORCH_CHECK(q.is_contiguous());
  TORCH_CHECK(l.size(0) == batch && l.size(1) == n);
  TORCH_CHECK(l.is_contiguous());

  launch_jacobi_smem(
      input.data_ptr<float>(),
      q.data_ptr<float>(),
      l.data_ptr<float>(),
      batch,
      n,
      static_cast<float>(tol2),
      static_cast<int>(max_sweeps),
      sweep_count.has_value() ? sweep_count->data_ptr<int>() : nullptr);
}

// Batched symmetric-tridiagonal leaf eigensolver (implicit-shift QL).
// d (P,N) fp64 diagonal, e (P,N-1) fp64 subdiagonal ->
// q (P,N,N) fp32 eigenvectors, l (P,N) fp64 eigenvalues ascending.
void steqr_tri(
    const at::Tensor& d,
    const at::Tensor& e,
    at::Tensor& q,
    at::Tensor& l,
    int64_t use_f32,
    int64_t do_sort) {
  TORCH_CHECK(d.dim() == 2);
  const int P = static_cast<int>(d.size(0));
  const int n = static_cast<int>(d.size(1));
  TORCH_CHECK(e.size(0) == P && e.size(1) == n - 1);
  TORCH_CHECK(d.is_cuda() && d.dtype() == at::kDouble && d.is_contiguous());
  TORCH_CHECK(e.is_cuda() && e.dtype() == at::kDouble && e.is_contiguous());
  TORCH_CHECK(q.dim() == 3 && q.size(0) == P && q.size(1) == n && q.size(2) == n);
  TORCH_CHECK(q.dtype() == at::kFloat && q.is_contiguous());
  TORCH_CHECK(l.size(0) == P && l.size(1) == n);
  TORCH_CHECK(l.dtype() == at::kDouble && l.is_contiguous());
  launch_steqr_tri(
      d.data_ptr<double>(),
      e.data_ptr<double>(),
      q.data_ptr<float>(),
      l.data_ptr<double>(),
      P,
      n,
      static_cast<int>(use_f32),
      static_cast<int>(do_sort));
}

void jacobi_cluster(
    const at::Tensor& input,
    at::Tensor& q,
    at::Tensor& l,
    double tol2,
    int64_t max_sweeps,
    std::optional<at::Tensor> sweep_count) {
  TORCH_CHECK(input.dim() == 3);
  const int batch = static_cast<int>(input.size(0));
  const int n = static_cast<int>(input.size(1));
  TORCH_CHECK(input.size(2) == n);
  TORCH_CHECK(input.is_cuda());
  TORCH_CHECK(input.dtype() == at::kFloat);
  TORCH_CHECK(input.is_contiguous());
  TORCH_CHECK(q.sizes() == input.sizes());
  TORCH_CHECK(q.is_contiguous());
  TORCH_CHECK(l.size(0) == batch && l.size(1) == n);
  TORCH_CHECK(l.is_contiguous());

  launch_jacobi_cluster(
      input.data_ptr<float>(),
      q.data_ptr<float>(),
      l.data_ptr<float>(),
      batch,
      n,
      static_cast<float>(tol2),
      static_cast<int>(max_sweeps),
      sweep_count.has_value() ? sweep_count->data_ptr<int>() : nullptr);
}

void seated_pivot_solve(
    const at::Tensor& a,
    at::Tensor& v,
    const at::Tensor& pairs,
    const at::Tensor& active,
    int64_t max_sweeps,
    double tol2) {
  TORCH_CHECK(a.dim() == 3);
  const int batch = static_cast<int>(a.size(0));
  const int n = static_cast<int>(a.size(1));
  TORCH_CHECK(a.size(2) == n);
  TORCH_CHECK(a.is_cuda() && a.is_contiguous());
  const bool a_half = a.dtype() == at::kHalf;
  TORCH_CHECK(a_half || a.dtype() == at::kFloat);
  TORCH_CHECK(v.dtype() == at::kHalf && v.is_contiguous());
  TORCH_CHECK(v.size(2) == 64 && v.size(3) == 64);
  TORCH_CHECK(pairs.dtype() == at::kInt && pairs.is_contiguous());
  TORCH_CHECK(pairs.size(0) == batch);
  TORCH_CHECK(active.dtype() == at::kInt && active.is_contiguous());

  launch_seated_pivot_solve(
      a.data_ptr(),
      a_half,
      v.data_ptr(),
      pairs.data_ptr<int>(),
      active.data_ptr<int>(),
      batch,
      n,
      static_cast<int>(max_sweeps),
      static_cast<float>(tol2));
}

void block_jacobi(
    const at::Tensor& input,
    at::Tensor& q,
    at::Tensor& l,
    at::Tensor& a_work,
    at::Tensor& qt_work,
    double tol2,
    int64_t max_sweeps,
    std::optional<at::Tensor> sweep_count) {
  TORCH_CHECK(input.dim() == 3);
  const int batch = static_cast<int>(input.size(0));
  const int n = static_cast<int>(input.size(1));
  TORCH_CHECK(input.size(2) == n);
  TORCH_CHECK(input.is_cuda() && input.dtype() == at::kFloat && input.is_contiguous());
  TORCH_CHECK(q.sizes() == input.sizes() && q.is_contiguous());
  TORCH_CHECK(l.size(0) == batch && l.size(1) == n && l.is_contiguous());
  TORCH_CHECK(a_work.sizes() == input.sizes() && a_work.is_contiguous());
  TORCH_CHECK(qt_work.sizes() == input.sizes() && qt_work.is_contiguous());

  launch_block_jacobi(
      input.data_ptr<float>(),
      q.data_ptr<float>(),
      l.data_ptr<float>(),
      a_work.data_ptr<float>(),
      qt_work.data_ptr<float>(),
      batch,
      n,
      static_cast<float>(tol2),
      static_cast<int>(max_sweeps),
      sweep_count.has_value() ? sweep_count->data_ptr<int>() : nullptr);
}

void dc_prep(
    const at::Tensor& Dl, const at::Tensor& Drr, const at::Tensor& zL,
    const at::Tensor& zR, const at::Tensor& rho,
    at::Tensor& dc, at::Tensor& zc, at::Tensor& k, at::Tensor& rho_s,
    at::Tensor& invorder, at::Tensor& perm, at::Tensor& hv, at::Tensor& hbeta,
    at::Tensor& segstart, at::Tensor& nseg, double defl_zk, double defl_gk) {
  const int P = static_cast<int>(Dl.size(0));
  const int h = static_cast<int>(Dl.size(1));
  const int m = static_cast<int>(dc.size(1));
  TORCH_CHECK(Dl.dtype() == at::kDouble && Dl.is_contiguous());
  TORCH_CHECK(Drr.dtype() == at::kDouble && Drr.is_contiguous());
  TORCH_CHECK(zL.dtype() == at::kFloat && zL.is_contiguous());
  TORCH_CHECK(zR.dtype() == at::kFloat && zR.is_contiguous());
  TORCH_CHECK(rho.dtype() == at::kDouble && rho.is_contiguous());
  dc::launch_dc_prep(
      Dl.data_ptr<double>(), Drr.data_ptr<double>(), zL.data_ptr<float>(),
      zR.data_ptr<float>(), rho.data_ptr<double>(),
      dc.data_ptr<double>(), zc.data_ptr<double>(), k.data_ptr<int>(),
      rho_s.data_ptr<double>(), invorder.data_ptr<int>(), perm.data_ptr<int>(),
      hv.data_ptr<double>(), hbeta.data_ptr<double>(), segstart.data_ptr<int>(),
      nseg.data_ptr<int>(), P, h, m, defl_zk, defl_gk);
}

void merge_build(
    const at::Tensor& dc, const at::Tensor& zc, const at::Tensor& k,
    const at::Tensor& rho, const at::Tensor& invorder, const at::Tensor& perm,
    const at::Tensor& hv, const at::Tensor& hbeta, const at::Tensor& segstart,
    const at::Tensor& nseg, at::Tensor& Vchild, at::Tensor& Lam,
    int64_t newt, int64_t nf32, double res_tol, double step_tol) {
  const int P = static_cast<int>(dc.size(0));
  const int m = static_cast<int>(dc.size(1));
  TORCH_CHECK(dc.dtype() == at::kDouble && dc.is_contiguous());
  TORCH_CHECK(zc.dtype() == at::kDouble && zc.is_contiguous());
  TORCH_CHECK(k.dtype() == at::kInt && k.is_contiguous());
  TORCH_CHECK(rho.dtype() == at::kDouble && rho.is_contiguous());
  TORCH_CHECK(invorder.dtype() == at::kInt && invorder.is_contiguous());
  TORCH_CHECK(perm.dtype() == at::kInt && perm.is_contiguous());
  TORCH_CHECK(hv.dtype() == at::kDouble && hv.is_contiguous());
  TORCH_CHECK(hbeta.dtype() == at::kDouble && hbeta.is_contiguous());
  TORCH_CHECK(segstart.dtype() == at::kInt && segstart.is_contiguous());
  TORCH_CHECK(nseg.dtype() == at::kInt && nseg.is_contiguous());
  TORCH_CHECK(Vchild.dtype() == at::kFloat && Vchild.is_contiguous());
  TORCH_CHECK(Lam.dtype() == at::kDouble && Lam.is_contiguous());
  dc::launch_merge_build(
      dc.data_ptr<double>(), zc.data_ptr<double>(), k.data_ptr<int>(),
      rho.data_ptr<double>(), invorder.data_ptr<int>(), perm.data_ptr<int>(),
      hv.data_ptr<double>(), hbeta.data_ptr<double>(), segstart.data_ptr<int>(),
      nseg.data_ptr<int>(), Vchild.data_ptr<float>(), Lam.data_ptr<double>(),
      P, m, static_cast<int>(newt), static_cast<int>(nf32), res_tol, step_tol);
}

void merge_build_multi(
    const at::Tensor& dc, const at::Tensor& zc, const at::Tensor& k,
    const at::Tensor& rho, const at::Tensor& invorder, const at::Tensor& perm,
    const at::Tensor& hv, const at::Tensor& hbeta, const at::Tensor& segstart,
    const at::Tensor& nseg, at::Tensor& Vchild, at::Tensor& Lam,
    int64_t newt, int64_t nf32, int64_t G) {
  const int P = static_cast<int>(dc.size(0));
  const int m = static_cast<int>(dc.size(1));
  dc::launch_merge_build_multi(
      dc.data_ptr<double>(), zc.data_ptr<double>(), k.data_ptr<int>(),
      rho.data_ptr<double>(), invorder.data_ptr<int>(), perm.data_ptr<int>(),
      hv.data_ptr<double>(), hbeta.data_ptr<double>(), segstart.data_ptr<int>(),
      nseg.data_ptr<int>(), Vchild.data_ptr<float>(), Lam.data_ptr<double>(),
      P, m, static_cast<int>(newt), static_cast<int>(nf32), static_cast<int>(G));
}

// One-stage tridiagonalization panel kernel + compact-WY T builder.
void latrd(
    const at::Tensor& A, const at::Tensor& Ah, at::Tensor& V, at::Tensor& W,
    at::Tensor& d, at::Tensor& e, int64_t p0) {
  const int b = static_cast<int>(A.size(0));
  const int n = static_cast<int>(A.size(1));
  TORCH_CHECK(A.dtype() == at::kFloat && A.is_contiguous());
  TORCH_CHECK(Ah.dtype() == at::kHalf && Ah.is_contiguous());
  TORCH_CHECK(V.dtype() == at::kFloat && W.dtype() == at::kFloat);
  TORCH_CHECK(d.dtype() == at::kFloat && e.dtype() == at::kFloat);
  os1::launch_latrd(
      A.data_ptr<float>(), reinterpret_cast<const __half*>(Ah.data_ptr<at::Half>()),
      V.data_ptr<float>(), W.data_ptr<float>(),
      d.data_ptr<float>(), e.data_ptr<float>(), b, n, static_cast<int>(p0));
}

// Tail-collapse: finish the last M0-column trailing block in ONE SMEM-resident
// launch. Vfull (b,M0,M0) trapezoidal reflectors + d,e (b,M0).
void latrd_tail(
    const at::Tensor& A, at::Tensor& Vfull, at::Tensor& d, at::Tensor& e,
    int64_t p0) {
  const int b = static_cast<int>(A.size(0));
  const int n = static_cast<int>(A.size(1));
  const int m0 = static_cast<int>(Vfull.size(1));
  TORCH_CHECK(A.dtype() == at::kFloat && A.is_contiguous());
  TORCH_CHECK(Vfull.dtype() == at::kFloat && Vfull.is_contiguous());
  TORCH_CHECK(d.dtype() == at::kFloat && e.dtype() == at::kFloat);
  os1::launch_latrd_tail(
      A.data_ptr<float>(), Vfull.data_ptr<float>(),
      d.data_ptr<float>(), e.data_ptr<float>(), b, n, static_cast<int>(p0), m0);
}

void latrd_cluster(
    const at::Tensor& A, const at::Tensor& Ah, at::Tensor& V, at::Tensor& W,
    at::Tensor& d, at::Tensor& e, int64_t p0, int64_t C) {
  const int b = static_cast<int>(A.size(0));
  const int n = static_cast<int>(A.size(1));
  TORCH_CHECK(A.dtype() == at::kFloat && A.is_contiguous());
  TORCH_CHECK(Ah.dtype() == at::kHalf && Ah.is_contiguous());
  TORCH_CHECK(V.dtype() == at::kFloat && W.dtype() == at::kFloat);
  TORCH_CHECK(d.dtype() == at::kFloat && e.dtype() == at::kFloat);
  const int nb = static_cast<int>(V.size(2));
  os1cl::launch_latrd_cluster(
      A.data_ptr<float>(), reinterpret_cast<const __half*>(Ah.data_ptr<at::Half>()),
      V.data_ptr<float>(), W.data_ptr<float>(),
      d.data_ptr<float>(), e.data_ptr<float>(), b, n,
      static_cast<int>(p0), static_cast<int>(C), nb);
}

void trec(const at::Tensor& G, at::Tensor& T) {
  const int b = static_cast<int>(G.size(0));
  TORCH_CHECK(G.dtype() == at::kFloat && G.is_contiguous());
  TORCH_CHECK(T.dtype() == at::kFloat && T.is_contiguous());
  const int nb = static_cast<int>(G.size(1));
  os1::launch_trec(G.data_ptr<float>(), T.data_ptr<float>(), b, nb);
}

TORCH_LIBRARY(eigh_ops, m) {
  m.def("jacobi_smem(Tensor input, Tensor(a!) q, Tensor(b!) l, float tol2=1e-10, int max_sweeps=30, Tensor(c!)? sweep_count=None) -> ()");
  m.impl("jacobi_smem", &jacobi_smem);
  m.def("steqr_tri(Tensor d, Tensor e, Tensor(a!) q, Tensor(b!) l, int use_f32=0, int do_sort=1) -> ()");
  m.impl("steqr_tri", &steqr_tri);
  m.def("jacobi_cluster(Tensor input, Tensor(a!) q, Tensor(b!) l, float tol2=1e-10, int max_sweeps=30, Tensor(c!)? sweep_count=None) -> ()");
  m.impl("jacobi_cluster", &jacobi_cluster);
  m.def("block_jacobi(Tensor input, Tensor(a!) q, Tensor(b!) l, Tensor(c!) a_work, Tensor(d!) qt_work, float tol2=1e-7, int max_sweeps=16, Tensor(e!)? sweep_count=None) -> ()");
  m.impl("block_jacobi", &block_jacobi);
  m.def("seated_pivot_solve(Tensor a, Tensor(a!) v, Tensor pairs, Tensor active, int max_sweeps, float tol2) -> ()");
  m.impl("seated_pivot_solve", &seated_pivot_solve);
  m.def("tf32_baddbmm_out(Tensor input, Tensor left, Tensor right, Tensor(a!) output, float beta=1.0, float alpha=1.0) -> ()");
  m.impl("tf32_baddbmm_out", &tf32_baddbmm_out);
  m.def("fp16_baddbmm_out(Tensor input, Tensor left, Tensor right, Tensor(a!) output, float beta=1.0, float alpha=1.0) -> ()");
  m.impl("fp16_baddbmm_out", &fp16_baddbmm_out);
  m.def("split16(Tensor x, Tensor(a!) hi, Tensor(b!) lo) -> ()");
  m.impl("split16", &split16);
  m.def("trail_epilogue(Tensor(a!) A, Tensor(b!) Ah, Tensor P, int off) -> ()");
  m.impl("trail_epilogue", &trail_epilogue);
  m.def("trail_epilogue_sym(Tensor(a!) A, Tensor(b!) Ah, Tensor Q, int off) -> ()");
  m.impl("trail_epilogue_sym", &trail_epilogue_sym);
  m.def("prescale(Tensor data, Tensor(a!) A, Tensor(b!) Ah, Tensor(c!) scale) -> ()");
  m.impl("prescale", &prescale);
  m.def("dc_prep(Tensor Dl, Tensor Drr, Tensor zL, Tensor zR, Tensor rho, Tensor(a!) dc, Tensor(b!) zc, Tensor(c!) k, Tensor(d!) rho_s, Tensor(e!) invorder, Tensor(f!) perm, Tensor(g!) hv, Tensor(h!) hbeta, Tensor(i!) segstart, Tensor(j!) nseg, float defl_zk=1.0, float defl_gk=1.0) -> ()");
  m.impl("dc_prep", &dc_prep);
  m.def("merge_build(Tensor dc, Tensor zc, Tensor k, Tensor rho, Tensor invorder, Tensor perm, Tensor hv, Tensor hbeta, Tensor segstart, Tensor nseg, Tensor(a!) Vchild, Tensor(b!) Lam, int newt, int nf32, float res_tol, float step_tol) -> ()");
  m.impl("merge_build", &merge_build);
  m.def("merge_build_multi(Tensor dc, Tensor zc, Tensor k, Tensor rho, Tensor invorder, Tensor perm, Tensor hv, Tensor hbeta, Tensor segstart, Tensor nseg, Tensor(a!) Vchild, Tensor(b!) Lam, int newt, int nf32, int G) -> ()");
  m.impl("merge_build_multi", &merge_build_multi);
  m.def("latrd(Tensor A, Tensor Ah, Tensor(a!) V, Tensor(b!) W, Tensor(c!) d, Tensor(e!) ee, int p0) -> ()");
  m.impl("latrd", &latrd);
  m.def("latrd_tail(Tensor A, Tensor(a!) Vfull, Tensor(b!) d, Tensor(c!) ee, int p0) -> ()");
  m.impl("latrd_tail", &latrd_tail);
  m.def("latrd_cluster(Tensor A, Tensor Ah, Tensor(a!) V, Tensor(b!) W, Tensor(c!) d, Tensor(e!) ee, int p0, int C) -> ()");
  m.impl("latrd_cluster", &latrd_cluster);
  m.def("trec(Tensor G, Tensor(a!) T) -> ()");
  m.impl("trec", &trec);
}
"""

CUDA_SRC = r"""
#include <cooperative_groups.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <cuda.h>            // driver API: CUtensorMap + enums (2D-TMA tensor-map symv)
#include <cudaTypedefs.h>    // PFN_cuTensorMapEncodeTiled_v12000
#include <math_constants.h>
#include <cstdint>
#include <cstdlib>
#include <cstdio>

namespace cg = cooperative_groups;

__device__ __host__
constexpr int cdiv(int a, int b) { return (a + b - 1) / b; }

constexpr unsigned FULL_MASK = 0xffffffffu;

// Evict-first fp16 global load (ld.global.cs cache hint): the trailing-matrix symv
// reads the whole m*m fp16 shadow once per column and never reuses a line before it
// is evicted (512KB working set >> L1), so caching it only thrashes L1. The cs hint
// tells the LSU not to keep it resident.
__device__ __forceinline__
float ld_half_cs(const __half* p){
  unsigned short h;
  asm volatile("ld.global.cs.u16 %0, [%1];" : "=h"(h) : "l"(p));
  return __half2float(*reinterpret_cast<const __half*>(&h));
}

__device__ __forceinline__
float warp_sum(float value, int size = 32) {
  #pragma unroll
  for (int offset = size / 2; offset > 0; offset >>= 1)
    value += __shfl_xor_sync(FULL_MASK, value, offset);
  return value;
}

// Fused double-single fp16 split: read one fp32 element, emit (hi, lo) fp16 in a
// single pass with contiguous coalesced writes.  Bit-identical to the torch glue
//   hi = x.to(fp16);  lo = (x - hi.float()).to(fp16)
// (round-to-nearest-even in __float2half_rn, exact __half2float), but replaces
// ~4 memory-bound elementwise kernels + their temporaries with one (1 read + 2
// half writes instead of ~18N read / 12N write bytes).  The input may be a
// strided sub-block view (e.g. Vchild[:, :h, :]) so no .contiguous() copy is
// needed; outputs hi/lo are contiguous (P,X,Y).
__global__ void split16_kernel(
    const float* __restrict__ in,
    __half* __restrict__ hi,
    __half* __restrict__ lo,
    long total, int Y, int XY,
    long bs, long rs, long cs) {
  for (long t = blockIdx.x * (long)blockDim.x + threadIdx.x; t < total;
       t += (long)gridDim.x * blockDim.x) {
    int p = (int)(t / XY);
    int rem = (int)(t - (long)p * XY);
    int i = rem / Y;
    int j = rem - i * Y;
    long idx = (long)p * bs + (long)i * rs + (long)j * cs;
    float x = in[idx];
    __half h = __float2half_rn(x);
    hi[t] = h;
    lo[t] = __float2half_rn(x - __half2float(h));
  }
}

void launch_split16(const float* in, __half* hi, __half* lo,
                    long total, int Y, int XY, long bs, long rs, long cs) {
  if (total <= 0) return;
  const int threads = 256;
  long nb = (total + threads - 1) / threads;
  if (nb > 65535) nb = 65535;
  split16_kernel<<<(int)nb, threads>>>(in, hi, lo, total, Y, XY, bs, rs, cs);
}

// Fused trailing-update epilogue for the one-stage tridiagonalization.  Given
// the rank-2b GEMM result P = VW @ WV^T (contiguous (b, tm, tm)) and the trailing
// sub-block of A at offset `off`, computes IN PLACE
//     A[b, off+r, off+c] -= P[b, r, c]
//     Ah[b, off+r, off+c] = (half) A[...]
// in a SINGLE pass, replacing the torch chain `res = sub - P; A[slice]=res;
// Ah[slice]=res.half()` (3 big memory-bound elementwise kernels + a full-size
// temp) with one (read A-slice + read P, write A-slice + write Ah-slice).
// Bit-identical: fp32 subtract + __float2half_rn (== torch .half() RN).
__global__ void trail_epilogue_kernel(
    float* __restrict__ A, __half* __restrict__ Ah,
    const float* __restrict__ P,
    int n, int off, int tm, long total) {
  const long nn = (long)n * n;
  const long tmtm = (long)tm * tm;
  for (long t = blockIdx.x * (long)blockDim.x + threadIdx.x; t < total;
       t += (long)gridDim.x * blockDim.x) {
    int bidx = (int)(t / tmtm);
    long rem = t - (long)bidx * tmtm;
    int r = (int)(rem / tm);
    int c = (int)(rem - (long)r * tm);
    long aidx = (long)bidx * nn + (long)(off + r) * n + (off + c);
    float v = A[aidx] - P[t];
    A[aidx] = v;
    Ah[aidx] = __float2half_rn(v);
  }
}

// float4-vectorized variant: 4 contiguous columns per thread.  Valid when
// tm % 4 == 0 (full 32-wide panels), which also guarantees 16-byte alignment
// of every float4 access (off is a multiple of the panel width, n = 512/1024/
// 2048 are multiples of 4).  Coalesces the strided A/Ah row writes into 128-bit
// transactions and halves the fp16 store count via __half2.
__global__ void trail_epilogue_vec4_kernel(
    float* __restrict__ A, __half* __restrict__ Ah,
    const float* __restrict__ P,
    int n, int off, int tm, long total4) {
  const long nn = (long)n * n;
  const long tmtm = (long)tm * tm;
  const int tm4 = tm >> 2;
  for (long q = blockIdx.x * (long)blockDim.x + threadIdx.x; q < total4;
       q += (long)gridDim.x * blockDim.x) {
    int bidx = (int)(q / ((long)tm * tm4));
    long rem = q - (long)bidx * tm * tm4;
    int r = (int)(rem / tm4);
    int cg = (int)(rem - (long)r * tm4);
    int c = cg << 2;
    long t = (long)bidx * tmtm + (long)r * tm + c;
    long aidx = (long)bidx * nn + (long)(off + r) * n + (off + c);
    float4 a4 = *reinterpret_cast<const float4*>(&A[aidx]);
    float4 p4 = *reinterpret_cast<const float4*>(&P[t]);
    a4.x -= p4.x; a4.y -= p4.y; a4.z -= p4.z; a4.w -= p4.w;
    *reinterpret_cast<float4*>(&A[aidx]) = a4;
    __half2 h01 = __floats2half2_rn(a4.x, a4.y);
    __half2 h23 = __floats2half2_rn(a4.z, a4.w);
    *reinterpret_cast<__half2*>(&Ah[aidx])     = h01;
    *reinterpret_cast<__half2*>(&Ah[aidx + 2]) = h23;
  }
}

void launch_trail_epilogue(float* A, __half* Ah, const float* P,
                           int b, int n, int off, int tm) {
  const long total = (long)b * tm * tm;
  if (total <= 0) return;
  const int threads = 256;
  if ((tm & 3) == 0) {
    const long total4 = total >> 2;
    long nb = (total4 + threads - 1) / threads;
    if (nb > 65535) nb = 65535;
    trail_epilogue_vec4_kernel<<<(int)nb, threads>>>(A, Ah, P, n, off, tm, total4);
  } else {
    long nb = (total + threads - 1) / threads;
    if (nb > 65535) nb = 65535;
    trail_epilogue_kernel<<<(int)nb, threads>>>(A, Ah, P, n, off, tm, total);
  }
}

// Symmetric fused trailing-update epilogue.  The rank-2b trailing update
//     A -= V W^T + W V^T
// is symmetric, so with the single (K=nb, no cat) product Q = V W^T the update
// is A -= (Q + Q^T).  This kernel forms (Q + Q^T), subtracts it from the
// trailing A block, and writes the fp16 shadow -- reading Q exactly ONCE
// (coalesced) via a shared-memory tile transpose, replacing the two torch.cat
// copies (the K=64 [V|W]/[W|V] operands) AND halving the trailing GEMM (one
// K=nb product instead of one K=2nb product).  Grid: one block per (batch,
// upper-triangular tile-pair); each block updates the (I,J) block and its
// mirror (J,I) so every A/Ah element is touched once.
// 32x32 tiles, float4-vectorized global I/O (4 columns per thread) so the
// A/Ah reads/writes and the Q loads coalesce into 128-bit transactions exactly
// like the non-symmetric vec4 epilogue.  block = 8 col-groups x 32 rows (256
// threads).  Requires tm % 4 == 0 (always true: off is a multiple of the panel
// width and n is a multiple of 4).  The +Q^T term is a scalar strided read out
// of the on-chip Q_JI tile (shared memory; a few-way bank conflict there is far
// cheaper than an uncoalesced global transpose read).
#define TES_TS 32
#define TES_CG 8
__global__ void trail_epilogue_sym_kernel(
    float* __restrict__ A, __half* __restrict__ Ah,
    const float* __restrict__ Q,
    int n, int off, int tm, int nt) {
  // +1 padded column stride: the transposed reads sJI[cb][ty] / sIJ[ty][cb]
  // hit distinct banks ((4*tx+ty)%32 all distinct across the warp) instead of
  // the 8-way conflict of a 32-stride tile. Stride 33 breaks float4-store
  // alignment, so the SMEM writes are done as 4 scalar stores (themselves
  // conflict-free at stride 33). Byte-identical math.
  __shared__ float sIJ[TES_TS][TES_TS + 1];
  __shared__ float sJI[TES_TS][TES_TS + 1];
  const int bidx = blockIdx.y;
  int pair = blockIdx.x;
  int I = 0;
  while (pair >= nt - I) { pair -= (nt - I); ++I; }
  const int J = I + pair;
  const int ty = threadIdx.y;      // row within tile (0..31)
  const int tx = threadIdx.x;      // col-group (0..7)
  const int cb = tx << 2;          // base column within tile
  const long tmtm = (long)tm * tm;
  const long nn = (long)n * n;
  const long qbase = (long)bidx * tmtm;
  const long abase = (long)bidx * nn;
  const int gi = I * TES_TS;
  const int gj = J * TES_TS;
  // load Q_IJ (row gi+ty, col gj+cb..) and Q_JI (row gj+ty, col gi+cb..)
  float4 z4 = make_float4(0.f, 0.f, 0.f, 0.f);
  {
    float4 q = z4;
    if (gi + ty < tm && gj + cb < tm)
      q = *reinterpret_cast<const float4*>(&Q[qbase + (long)(gi + ty) * tm + (gj + cb)]);
    sIJ[ty][cb] = q.x; sIJ[ty][cb+1] = q.y; sIJ[ty][cb+2] = q.z; sIJ[ty][cb+3] = q.w;
    q = z4;
    if (gj + ty < tm && gi + cb < tm)
      q = *reinterpret_cast<const float4*>(&Q[qbase + (long)(gj + ty) * tm + (gi + cb)]);
    sJI[ty][cb] = q.x; sJI[ty][cb+1] = q.y; sJI[ty][cb+2] = q.z; sJI[ty][cb+3] = q.w;
  }
  __syncthreads();
  // block (I,J): p = Q_IJ[ty][cb+i] + Q^T = Q_JI[cb+i][ty]
  if (gi + ty < tm && gj + cb < tm) {
    long aidx = abase + (long)(off + gi + ty) * n + (off + gj + cb);
    float4 a = *reinterpret_cast<const float4*>(&A[aidx]);
    a.x -= sIJ[ty][cb]   + sJI[cb][ty];
    a.y -= sIJ[ty][cb+1] + sJI[cb+1][ty];
    a.z -= sIJ[ty][cb+2] + sJI[cb+2][ty];
    a.w -= sIJ[ty][cb+3] + sJI[cb+3][ty];
    *reinterpret_cast<float4*>(&A[aidx]) = a;
    *reinterpret_cast<__half2*>(&Ah[aidx])     = __floats2half2_rn(a.x, a.y);
    *reinterpret_cast<__half2*>(&Ah[aidx + 2]) = __floats2half2_rn(a.z, a.w);
  }
  // mirror block (J,I) when off-diagonal
  if (I != J && gj + ty < tm && gi + cb < tm) {
    long aidx = abase + (long)(off + gj + ty) * n + (off + gi + cb);
    float4 a = *reinterpret_cast<const float4*>(&A[aidx]);
    a.x -= sJI[ty][cb]   + sIJ[cb][ty];
    a.y -= sJI[ty][cb+1] + sIJ[cb+1][ty];
    a.z -= sJI[ty][cb+2] + sIJ[cb+2][ty];
    a.w -= sJI[ty][cb+3] + sIJ[cb+3][ty];
    *reinterpret_cast<float4*>(&A[aidx]) = a;
    *reinterpret_cast<__half2*>(&Ah[aidx])     = __floats2half2_rn(a.x, a.y);
    *reinterpret_cast<__half2*>(&Ah[aidx + 2]) = __floats2half2_rn(a.z, a.w);
  }
}

// Scalar fallback for the tiny tail panels where tm is not a multiple of 4
// (e.g. the last n512 panel has tm=1): one thread per element, no float4.
__global__ void trail_epilogue_sym_scalar_kernel(
    float* __restrict__ A, __half* __restrict__ Ah,
    const float* __restrict__ Q, int n, int off, int tm, long total) {
  const long tmtm = (long)tm * tm;
  const long nn = (long)n * n;
  for (long t = blockIdx.x * (long)blockDim.x + threadIdx.x; t < total;
       t += (long)gridDim.x * blockDim.x) {
    int bidx = (int)(t / tmtm);
    long rem = t - (long)bidx * tmtm;
    int r = (int)(rem / tm);
    int c = (int)(rem - (long)r * tm);
    long qb = (long)bidx * tmtm;
    float p = Q[qb + (long)r * tm + c] + Q[qb + (long)c * tm + r];
    long aidx = (long)bidx * nn + (long)(off + r) * n + (off + c);
    float v = A[aidx] - p;
    A[aidx] = v;
    Ah[aidx] = __float2half_rn(v);
  }
}

void launch_trail_epilogue_sym(float* A, __half* Ah, const float* Q,
                               int b, int n, int off, int tm) {
  const long total = (long)b * tm * tm;
  if (total <= 0) return;
  if (tm % 4 == 0) {
    const int nt = (tm + TES_TS - 1) / TES_TS;
    const int npair = nt * (nt + 1) / 2;
    dim3 grid(npair, b);
    dim3 block(TES_CG, TES_TS);
    trail_epilogue_sym_kernel<<<grid, block>>>(A, Ah, Q, n, off, tm, nt);
  } else {
    const int threads = 256;
    long nb = (total + threads - 1) / threads;
    if (nb > 65535) nb = 65535;
    trail_epilogue_sym_scalar_kernel<<<(int)nb, threads>>>(A, Ah, Q, n, off, tm, total);
  }
}

// Fused amax-prescale.  Replaces the torch chain
//   scale = data.abs().amax(dim=(1,2)).clamp_min(1e-30)   # abs temp + reduce
//   A = data / scale                                       # divide
//   Ah = A.half()                                          # fp16 shadow cast
// (7 memory passes incl. a full (b,n,n) abs temporary) with two custom passes:
//   (1) a 2D-grid float4 reduction (CHUNKS blocks/matrix, atomicMax) -> scale[b];
//   (2) a float4 elementwise pass writes A = data*inv and Ah = A.half().
// prescale_reduce_kernel: 2D grid (blockIdx.y = matrix, blockIdx.x = chunk).
// The prior one-block-per-matrix layout serialized a grid-stride amax over the
// whole n*n matrix in a single block (8 blocks at n2048 b8 => 2.1 ms, ~64 GB/s,
// SM-starved). Here CHUNKS blocks cooperatively sweep each matrix with float4
// loads and atomicMax (int-reinterpret is monotone for non-negative floats) into
// a zero-initialized scale[]. Fills the machine and quarters the load count.
__global__ void prescale_reduce_kernel(
    const float* __restrict__ data, float* __restrict__ scale, long nn) {
  const int bidx = blockIdx.y;
  const long base = (long)bidx * nn;
  const long nn4 = nn >> 2;                 // n is a multiple of 4 (512/1024/2048)
  const float4* __restrict__ d4 =
      reinterpret_cast<const float4*>(data + base);
  float m = 0.0f;
  const long stride = (long)gridDim.x * blockDim.x;
  for (long t = (long)blockIdx.x * blockDim.x + threadIdx.x; t < nn4; t += stride) {
    float4 v = d4[t];
    m = fmaxf(m, fmaxf(fmaxf(fabsf(v.x), fabsf(v.y)),
                       fmaxf(fabsf(v.z), fabsf(v.w))));
  }
  // warp then block reduction
  #pragma unroll
  for (int o = 16; o > 0; o >>= 1) m = fmaxf(m, __shfl_xor_sync(FULL_MASK, m, o));
  __shared__ float sm[32];
  int lane = threadIdx.x & 31, wid = threadIdx.x >> 5;
  if (lane == 0) sm[wid] = m;
  __syncthreads();
  if (wid == 0) {
    int nw = (blockDim.x + 31) >> 5;
    m = (lane < nw) ? sm[lane] : 0.0f;
    #pragma unroll
    for (int o = 16; o > 0; o >>= 1) m = fmaxf(m, __shfl_xor_sync(FULL_MASK, m, o));
    // int-reinterpret atomicMax: valid because m >= 0 (fabsf) and scale is
    // zero-initialized -> IEEE non-negative floats order as their int bits.
    if (lane == 0) atomicMax((int*)&scale[bidx], __float_as_int(m));
  }
}

// prescale_apply_kernel: A = data * (1/scale[b]); Ah = A.half().  float4 I/O.
__global__ void prescale_apply_kernel(
    const float* __restrict__ data, float* __restrict__ A,
    __half* __restrict__ Ah, const float* __restrict__ scale,
    long nn, long total4) {
  const long nn4 = nn >> 2;
  for (long q = blockIdx.x * (long)blockDim.x + threadIdx.x; q < total4;
       q += (long)gridDim.x * blockDim.x) {
    int bidx = (int)(q / nn4);
    float inv = 1.0f / fmaxf(scale[bidx], 1e-30f);
    long idx = q << 2;
    float4 d = *reinterpret_cast<const float4*>(&data[idx]);
    d.x *= inv; d.y *= inv; d.z *= inv; d.w *= inv;
    *reinterpret_cast<float4*>(&A[idx]) = d;
    *reinterpret_cast<__half2*>(&Ah[idx])     = __floats2half2_rn(d.x, d.y);
    *reinterpret_cast<__half2*>(&Ah[idx + 2]) = __floats2half2_rn(d.z, d.w);
  }
}

void launch_prescale(const float* data, float* A, __half* Ah, float* scale,
                     int b, int n) {
  const long nn = (long)n * n;
  if (b <= 0 || nn <= 0) return;
  int rthreads = 256;
  // atomicMax reduction needs scale zero-initialized. Cooperative CHUNKS blocks
  // per matrix (target ~64 float4 loads/thread) fill the SMs even at tiny batch.
  cudaMemsetAsync(scale, 0, (size_t)b * sizeof(float));
  const long nn4 = nn >> 2;
  int chunks = (int)((nn4 + (long)rthreads * 64 - 1) / ((long)rthreads * 64));
  if (chunks < 1) chunks = 1;
  if (chunks > 512) chunks = 512;
  dim3 rgrid(chunks, b);
  prescale_reduce_kernel<<<rgrid, rthreads>>>(data, scale, nn);
  const long total4 = (long)b * nn >> 2;
  const int threads = 256;
  long nb = (total4 + threads - 1) / threads;
  if (nb > 65535) nb = 65535;
  prescale_apply_kernel<<<(int)nb, threads>>>(data, A, Ah, scale, nn, total4);
}

// Round-robin tournament pairing (circle method): N players, N-1 rounds of
// N/2 disjoint pairs that together cover every (p, q) combination exactly
// once per sweep. Player 0 is fixed; player x in [1, N-1] sits at position
// (x - 1 + round) mod (N - 1).
template <int N>
__device__ __forceinline__
void round_robin_pair(int round, int k, int& p, int& q) {
  constexpr int M = N - 1;
  auto player_at = [round](int pos) {
    int x = pos - round;
    x %= M;
    if (x < 0) x += M;
    return 1 + x;
  };
  int a, b;
  if (k == 0) {
    a = 0;
    b = player_at(0);
  } else {
    a = player_at(k);
    b = player_at(M - k);
  }
  p = min(a, b);
  q = max(a, b);
}

constexpr int next_pow2(int v) {
  int p = 1;
  while (p < v) p <<= 1;
  return p;
}

// One matrix per CTA, A and Q resident in shared memory with a padded row
// stride (N+1) so column accesses are bank-conflict free. Two-sided parallel
// Jacobi with the round-robin ordering. Each round:
//   phase R (one thread per pair): generate the rotation from the 2x2 diagonal
//     block and immediately apply it to that block (annihilating A[p][q]).
//   phase U: fused 2x2-block update — for every unordered pair-of-pairs
//     (k1 < k2) rotate the 2x2 block from both sides in registers and write it
//     plus its mirror; independently rotate Q's column pairs.
// Eigenvalues are sorted ascending in-kernel (bitonic) and Q's columns are
// permuted to match during writeout, so no host-side launches are needed.
template <int N, int THREADS>
__global__
__launch_bounds__(THREADS, 1)
void jacobi_smem_kernel(
    const float* __restrict__ input,
    float* __restrict__ q_out,
    float* __restrict__ l_out,
    float tolerance2,
    int max_sweeps,
    int* __restrict__ sweep_count) {
  static_assert(N % 2 == 0);
  constexpr int S = N + 1;                        // padded row stride
  constexpr int PAIRS = N / 2;
  constexpr int ROUNDS = N - 1;
  constexpr int NOFF = PAIRS * (PAIRS - 1) / 2;   // pair-of-pairs blocks
  constexpr int WARPS = THREADS / 32;
  static_assert(PAIRS <= 256, "pair table uses uint8");
  const int tid = threadIdx.x;
  const int batch = blockIdx.x;

  input += static_cast<size_t>(batch) * N * N;
  q_out += static_cast<size_t>(batch) * N * N;
  l_out += static_cast<size_t>(batch) * N;

  extern __shared__ float smem[];
  float* const A = smem;                          // N * S
  float* const Q = A + N * S;                     // N * S
  float* const rot_c = Q + N * S;                 // PAIRS
  float* const rot_s = rot_c + PAIRS;             // PAIRS
  int* const rot_p = reinterpret_cast<int*>(rot_s + PAIRS);   // PAIRS
  int* const rot_q = rot_p + PAIRS;               // PAIRS
  uint8_t* const tab1 = reinterpret_cast<uint8_t*>(rot_q + PAIRS);  // NOFF
  uint8_t* const tab2 = tab1 + NOFF;              // NOFF
  __shared__ float partials[WARPS];
  __shared__ float fro2_shared;
  __shared__ int done;

  // Pair-of-pairs lookup table (linear triangular index -> (k1, k2)).
  for (int t = tid; t < NOFF; t += THREADS) {
    int k1 = 0;
    int rem = t;
    while (rem >= PAIRS - 1 - k1) {
      rem -= PAIRS - 1 - k1;
      ++k1;
    }
    tab1[t] = static_cast<uint8_t>(k1);
    tab2[t] = static_cast<uint8_t>(k1 + 1 + rem);
  }

  // Load and symmetrize (input is symmetric up to FP32 roundoff); Q = I.
  float fro2_local = 0.0f;
  for (int i = tid; i < N * N; i += THREADS) {
    const int row = i / N;
    const int col = i % N;
    const float value = 0.5f * (input[i] + input[col * N + row]);
    A[row * S + col] = value;
    Q[row * S + col] = row == col ? 1.0f : 0.0f;
    fro2_local += value * value;
  }
  fro2_local = warp_sum(fro2_local);
  if ((tid & 31) == 0) partials[tid / 32] = fro2_local;
  __syncthreads();
  if (tid == 0) {
    float total = 0.0f;
    for (int w = 0; w < WARPS; ++w) total += partials[w];
    fro2_shared = total;
    done = total == 0.0f;
  }
  __syncthreads();

  int sweeps_run = 0;
  for (int sweep = 0; sweep < max_sweeps && !done; ++sweep) {
    ++sweeps_run;
    for (int round = 0; round < ROUNDS; ++round) {
      // Phase R: rotation generation + own 2x2 diagonal block update.
      if (tid < PAIRS) {
        int p, q;
        round_robin_pair<N>(round, tid, p, q);
        const float app = A[p * S + p];
        const float aqq = A[q * S + q];
        const float apq = A[p * S + q];
        float c = 1.0f;
        float s = 0.0f;
        float t = 0.0f;
        if (fabsf(apq) > 0.0f) {
          const float theta = 0.5f * (aqq - app) / apq;
          t = copysignf(
              1.0f / (fabsf(theta) + sqrtf(1.0f + theta * theta)), theta);
          c = rsqrtf(1.0f + t * t);
          s = t * c;
        }
        rot_c[tid] = c;
        rot_s[tid] = s;
        rot_p[tid] = p;
        rot_q[tid] = q;
        A[p * S + p] = app - t * apq;
        A[q * S + q] = aqq + t * apq;
        A[p * S + q] = 0.0f;
        A[q * S + p] = 0.0f;
      }
      __syncthreads();

      // Phase U(a): fused two-sided update of off blocks, mirror written.
      for (int t = tid; t < NOFF; t += THREADS) {
        const int k1 = tab1[t];
        const int k2 = tab2[t];
        const int p1 = rot_p[k1];
        const int q1 = rot_q[k1];
        const int p2 = rot_p[k2];
        const int q2 = rot_q[k2];
        const float c1 = rot_c[k1];
        const float s1 = rot_s[k1];
        const float c2 = rot_c[k2];
        const float s2 = rot_s[k2];
        const float xpp = A[p1 * S + p2];
        const float xpq = A[p1 * S + q2];
        const float xqp = A[q1 * S + p2];
        const float xqq = A[q1 * S + q2];
        // Left rotation (rows p1, q1), then right rotation (cols p2, q2).
        const float ypp = c1 * xpp - s1 * xqp;
        const float ypq = c1 * xpq - s1 * xqq;
        const float yqp = s1 * xpp + c1 * xqp;
        const float yqq = s1 * xpq + c1 * xqq;
        const float zpp = c2 * ypp - s2 * ypq;
        const float zpq = s2 * ypp + c2 * ypq;
        const float zqp = c2 * yqp - s2 * yqq;
        const float zqq = s2 * yqp + c2 * yqq;
        A[p1 * S + p2] = zpp;
        A[p1 * S + q2] = zpq;
        A[q1 * S + p2] = zqp;
        A[q1 * S + q2] = zqq;
        A[p2 * S + p1] = zpp;
        A[q2 * S + p1] = zpq;
        A[p2 * S + q1] = zqp;
        A[q2 * S + q1] = zqq;
      }

      // Phase U(b): rotate Q's column pairs (row-parallel, conflict-free via
      // the padded stride).
      for (int t = tid; t < PAIRS * N; t += THREADS) {
        const int k = t / N;
        const int i = t % N;
        const int p = rot_p[k];
        const int q = rot_q[k];
        const float c = rot_c[k];
        const float s = rot_s[k];
        const float qp = Q[i * S + p];
        const float qq = Q[i * S + q];
        Q[i * S + p] = c * qp - s * qq;
        Q[i * S + q] = s * qp + c * qq;
      }
      __syncthreads();
    }

    // Sweep-end convergence check on the off-diagonal Frobenius norm.
    float off2_local = 0.0f;
    for (int i = tid; i < N * N; i += THREADS) {
      const int row = i / N;
      const int col = i % N;
      if (row != col) {
        const float value = A[row * S + col];
        off2_local += value * value;
      }
    }
    off2_local = warp_sum(off2_local);
    if ((tid & 31) == 0) partials[tid / 32] = off2_local;
    __syncthreads();
    if (tid == 0) {
      float total = 0.0f;
      for (int w = 0; w < WARPS; ++w) total += partials[w];
      done = total <= tolerance2 * fro2_shared;
    }
    __syncthreads();
  }

  if (sweep_count != nullptr && tid == 0) sweep_count[batch] = sweeps_run;

  // Ascending sort of the diagonal (bitonic network on NP2 slots) and Q
  // column permutation, entirely in-kernel.
  constexpr int NP2 = next_pow2(N);
  float* const sort_key = reinterpret_cast<float*>(
      (reinterpret_cast<uintptr_t>(tab2 + NOFF) + 15) & ~uintptr_t(15));
  int* const sort_idx = reinterpret_cast<int*>(sort_key + NP2);
  for (int i = tid; i < NP2; i += THREADS) {
    sort_key[i] = i < N ? A[i * S + i] : CUDART_INF_F;
    sort_idx[i] = i;
  }
  __syncthreads();
  for (int k = 2; k <= NP2; k <<= 1) {
    for (int j = k >> 1; j > 0; j >>= 1) {
      for (int i = tid; i < NP2; i += THREADS) {
        const int ixj = i ^ j;
        if (ixj > i) {
          const float a_val = sort_key[i];
          const float b_val = sort_key[ixj];
          const bool ascending = (i & k) == 0;
          if (ascending ? (a_val > b_val) : (a_val < b_val)) {
            sort_key[i] = b_val;
            sort_key[ixj] = a_val;
            const int tmp = sort_idx[i];
            sort_idx[i] = sort_idx[ixj];
            sort_idx[ixj] = tmp;
          }
        }
      }
      __syncthreads();
    }
  }

  if (tid < N) l_out[tid] = sort_key[tid];
  for (int i = tid; i < N * N; i += THREADS) {
    const int row = i / N;
    const int col = i % N;
    q_out[i] = Q[row * S + sort_idx[col]];
  }
}

// ---------------------------------------------------------------------------
// Batched TRIDIAGONAL leaf eigensolver (implicit-shift QL, EISPACK tql2 /
// Numerical-Recipes tqli). The D&C leaf blocks are SYMMETRIC TRIDIAGONAL, so
// running a dense Jacobi (jacobi_smem) wastes O(N^2) SMEM traffic on the zeros.
// This solver keeps the tridiagonal structure throughout (Givens bulge-chase),
// touching only O(N) diag/offdiag per step plus the O(N) eigenvector column
// rotation -> ~N/const less work and near-zero SMEM (everything in registers /
// per-thread local, warp-uniform indexed => coalesced L1).
//
// Layout: ONE WARP owns one matrix (N == 32 lanes). lane r holds ROW r of the
// eigenvector matrix Z in a per-thread local array zrow[N] (init = identity).
// The scalar QL chase (d,e updates, shift, deflation) is computed REDUNDANTLY
// and identically on all 32 lanes (uniform control flow -> no divergence, no
// broadcast); each lane applies the resulting Givens (c,s) to its own Z row.
// d,e are held in fp64 per-thread local (dd[N], ee[N]) for eigenvalue accuracy;
// Z stays fp32 (orthogonality is exact under c^2+s^2=1). Output L is fp64.
template <typename CT>
__device__ __forceinline__ CT steqr_pythag(CT a, CT b) {
  CT absa = fabs(a), absb = fabs(b);
  if (absa > absb) {
    CT r = absb / absa;
    return absa * sqrt(CT(1) + r * r);
  }
  if (absb == CT(0)) return CT(0);
  CT r = absa / absb;
  return absb * sqrt(CT(1) + r * r);
}
template <typename CT> __device__ __forceinline__ CT steqr_eps();
template <> __device__ __forceinline__ double steqr_eps<double>() { return 2.220446049250313e-16; }
template <> __device__ __forceinline__ float  steqr_eps<float>()  { return 1.1920929e-07f; }

// One WARP owns one matrix (N == 32 lanes, lane r = row r). The eigenvector
// row lives in a per-thread LOCAL array zrow[N] (NOT shared) -> no shared cap on
// occupancy; the QL chase publishes each Givens column index bi UNIFORMLY across
// the warp (via __shfl), so zrow[bi]/zrow[bi+1] are warp-uniform-indexed local
// accesses -> the hardware coalesces them into single L1 transactions. The
// scalar tridiagonal (dsh,esh) + sort perm live in a tiny per-warp SHARED slice
// (only lane 0 mutates dsh/esh). lane 0 runs the fp64 chase ONCE and broadcasts
// (bi,bc,bs); the other 31 lanes apply it to their Z row. Occupancy is now
// limited only by registers (shared is ~640 B/warp), which is what a per-matrix
// serial dependency chain needs to hide its latency across many warps.
template <int N, int WARPS, typename CT, bool SORT>
__global__
__launch_bounds__(WARPS * 32, 6)
void steqr_tri_kernel(
    const double* __restrict__ d_in,   // (P, N)   diagonal
    const double* __restrict__ e_in,   // (P, N-1) subdiagonal
    float* __restrict__ q_out,         // (P, N, N) eigenvectors (row-major)
    double* __restrict__ l_out,        // (P, N) eigenvalues (ascending if SORT)
    int P) {
  static_assert(N <= 32, "one warp per matrix requires N <= 32");
  const int warp = threadIdx.x >> 5;
  const int lane = threadIdx.x & 31;
  const int matrix = blockIdx.x * WARPS + warp;

  extern __shared__ char steqr_smem[];
  CT* const dbase = reinterpret_cast<CT*>(steqr_smem);
  int* const pbase = reinterpret_cast<int*>(dbase + WARPS * 2 * N);
  CT* dsh = dbase + warp * (2 * N);               // [N]
  CT* esh = dsh + N;                              // [N]
  int* perm = pbase + warp * N;                   // [N]

  if (matrix >= P) return;

  const CT EPS = steqr_eps<CT>();
  const int nm1 = N - 1;

  // Load tridiagonal (lane 0). Z row in local; identity init.
  const double* dptr = d_in + static_cast<size_t>(matrix) * N;
  const double* eptr = e_in + static_cast<size_t>(matrix) * nm1;
  if (lane == 0) {
    for (int i = 0; i < N; ++i) {
      dsh[i] = (CT)dptr[i];
      esh[i] = (i < nm1) ? (CT)eptr[i] : CT(0);
    }
  }
  float zrow[N];
  for (int c = 0; c < N; ++c) zrow[c] = (c == lane) ? 1.0f : 0.0f;
  __syncwarp();

  // Implicit-shift QL (tql2), driven by lane 0.
  for (int l = 0; l < N; ++l) {
    int iter = 0;
    do {
      // lane 0 locates a negligible subdiagonal at/below l.
      int m = l;
      if (lane == 0) {
        for (m = l; m <= nm1 - 1; ++m) {
          CT dda = fabs(dsh[m]) + fabs(dsh[m + 1]);
          if (fabs(esh[m]) <= EPS * dda) break;
        }
      }
      m = __shfl_sync(FULL_MASK, m, 0);
      if (m == l) break;          // eigenvalue l converged
      if (++iter > 50) break;     // safety

      // Inner QL bulge-chase. lane 0 holds the scalar state (g,s,c,p,i) and
      // emits one Givens (bi,bc,bs) per step; all lanes apply it to their Z row.
      CT g = CT(0), s = CT(1), c = CT(1), p = CT(0);
      int i = m - 1;
      if (lane == 0) {
        g = (dsh[l + 1] - dsh[l]) / (CT(2) * esh[l]);
        CT r = steqr_pythag<CT>(g, CT(1));
        g = dsh[m] - dsh[l] + esh[l] / (g + copysign(r, g));
      }
      while (true) {
        int bi = -1, active = 0;
        float bc = 1.0f, bs = 0.0f;
        if (lane == 0) {
          if (i >= l) {
            CT f = s * esh[i];
            CT b = c * esh[i];
            CT r = steqr_pythag<CT>(f, g);
            esh[i + 1] = r;
            if (r == CT(0)) {
              dsh[i + 1] -= p;
              esh[m] = CT(0);
              active = 0;          // sweep aborts; do-while will restart
            } else {
              s = f / r;
              c = g / r;
              g = dsh[i + 1] - p;
              r = (dsh[i] - g) * s + CT(2) * c * b;
              p = s * r;
              dsh[i + 1] = g + p;
              g = c * r - b;
              bi = i; bc = (float)c; bs = (float)s;
              active = 1;
            }
          } else {
            dsh[l] -= p;           // normal sweep completion
            esh[l] = g;
            esh[m] = CT(0);
            active = 0;
          }
        }
        active = __shfl_sync(FULL_MASK, active, 0);
        if (!active) break;
        bi = __shfl_sync(FULL_MASK, bi, 0);
        bc = __shfl_sync(FULL_MASK, bc, 0);
        bs = __shfl_sync(FULL_MASK, bs, 0);
        float zi = zrow[bi];       // bi warp-uniform -> coalesced local
        float zj = zrow[bi + 1];
        zrow[bi] = bc * zi - bs * zj;
        zrow[bi + 1] = bs * zi + bc * zj;
        if (lane == 0) --i;
      }
    } while (true);
  }

  // Optionally ascending selection-sort the N eigenvalues (lane 0) -> perm.
  // When the caller's D&C merge re-sorts (dc_prep), SORT=false skips this serial
  // O(N^2) sort on lane 0's critical path and emits natural (column) order.
  if (SORT) {
    if (lane == 0) {
      for (int i = 0; i < N; ++i) perm[i] = i;
      for (int i = 0; i < N - 1; ++i) {
        int best = i;
        CT bv = dsh[perm[i]];
        for (int j = i + 1; j < N; ++j) {
          CT v = dsh[perm[j]];
          if (v < bv) { bv = v; best = j; }
        }
        if (best != i) { int t = perm[i]; perm[i] = perm[best]; perm[best] = t; }
      }
    }
    __syncwarp();
    if (lane < N) {
      float* q_row = q_out + (static_cast<size_t>(matrix) * N + lane) * N;
      for (int c = 0; c < N; ++c) q_row[c] = zrow[perm[c]];
      l_out[static_cast<size_t>(matrix) * N + lane] = (double)dsh[perm[lane]];
    }
  } else {
    if (lane < N) {
      float* q_row = q_out + (static_cast<size_t>(matrix) * N + lane) * N;
      for (int c = 0; c < N; ++c) q_row[c] = zrow[c];
      l_out[static_cast<size_t>(matrix) * N + lane] = (double)dsh[lane];
    }
  }
}

// ---------------------------------------------------------------------------
// Warp Jacobi: ONE WARP owns one matrix (small n). Same fused round structure
// as jacobi_smem_kernel but fully warp-synchronous: no __syncthreads on the
// 31-round x ~6-sweep critical path, several matrices per CTA, one launch for
// the whole batch. A and Q live in a per-warp SMEM slice (padded stride).
// ---------------------------------------------------------------------------
template <int N>
__device__ __forceinline__ int tri_addr(int i, int j) {
  // Packed row-major upper triangle, i <= j.
  return i * N - (i * (i - 1)) / 2 + (j - i);
}

template <int N>
__device__ __forceinline__ int tri_addr_sym(int a, int b) {
  const int i = min(a, b);
  const int j = max(a, b);
  return tri_addr<N>(i, j);
}

template <int N, int THREADS>
__global__
__launch_bounds__(THREADS, 1)
void jacobi_cluster_kernel(
    const float* __restrict__ input,
    float* __restrict__ q_out,
    float* __restrict__ l_out,
    float tolerance2,
    int max_sweeps,
    int* __restrict__ sweep_count) {
  static_assert(N % 4 == 0);
  constexpr int S = N + 1;                       // padded stride for Q bands
  constexpr int PAIRS = N / 2;
  constexpr int ROUNDS = N - 1;
  constexpr int NOFF = PAIRS * (PAIRS - 1) / 2;
  constexpr int TRI = N * (N + 1) / 2;
  constexpr int HALF = N / 2;                    // rows per Q-CTA
  constexpr int NP2 = next_pow2(N);
  constexpr int WARPS = THREADS / 32;
  static_assert(PAIRS <= 256, "pair table uses uint8");

  cg::cluster_group cluster = cg::this_cluster();
  const int rank = cluster.block_rank();
  const int matrix = blockIdx.x / 3;
  const int tid = threadIdx.x;

  input += static_cast<size_t>(matrix) * N * N;
  q_out += static_cast<size_t>(matrix) * N * N;
  l_out += static_cast<size_t>(matrix) * N;

  // --- shared memory layout: identical comm prefix in every CTA ---
  extern __shared__ float4 smem4[];
  float4* const rot_buf = smem4;                 // [2][PAIRS]
  int* const comm_int = reinterpret_cast<int*>(rot_buf + 2 * PAIRS);
  int* const done_flag = comm_int;               // [4] (padded)
  int* const perm = comm_int + 4;                // [N]
  float* const tail = reinterpret_cast<float*>(perm + N);
  // rank 0 tail: A[N*S] (full, padded stride), sort_key[NP2], sort_idx[NP2],
  // tabs. Full storage costs 2x the SMEM of a packed triangle but makes every
  // fused-update access bank-conflict free (loads run along rows, mirror
  // stores run down columns with odd stride).
  float* const Afull = tail;
  float* const sort_key = Afull + N * S;
  int* const sort_idx = reinterpret_cast<int*>(sort_key + NP2);
  uint8_t* const tab1 = reinterpret_cast<uint8_t*>(sort_idx + NP2);
  uint8_t* const tab2 = tab1 + NOFF;
  // rank 1/2 tail: Qband[HALF * S]
  float* const Qband = tail;

  __shared__ float partials[WARPS];
  __shared__ float fro2_shared;
  __shared__ float tau2_shared;  // threshold-Jacobi skip level (see below)

  // --- init ---
  if (rank == 0) {
    // Pair-of-pairs lookup table.
    for (int t = tid; t < NOFF; t += THREADS) {
      int k1 = 0;
      int rem = t;
      while (rem >= PAIRS - 1 - k1) {
        rem -= PAIRS - 1 - k1;
        ++k1;
      }
      tab1[t] = static_cast<uint8_t>(k1);
      tab2[t] = static_cast<uint8_t>(k1 + 1 + rem);
    }
    // Load + symmetrize; Frobenius norm^2 (each off pair counted twice).
    float fro2_local = 0.0f;
    for (int t = tid; t < N * N; t += THREADS) {
      const int i = t / N;
      const int j = t % N;
      const float value = 0.5f * (input[i * N + j] + input[j * N + i]);
      Afull[i * S + j] = value;
      fro2_local += value * value;
    }
    fro2_local = warp_sum(fro2_local);
    if ((tid & 31) == 0) partials[tid / 32] = fro2_local;
    __syncthreads();
    if (tid == 0) {
      float total = 0.0f;
      for (int w = 0; w < WARPS; ++w) total += partials[w];
      fro2_shared = total;
      // Threshold-Jacobi: rotations with apq^2 below tau2 are skipped as
      // exact identities. If ALL off elements sat at tau, off^2 would be
      // tol2*fro2/4, still a 4x margin inside the convergence tolerance.
      tau2_shared = 0.25f * tolerance2 * total / float(N * N);
      const int is_done = total == 0.0f;
      done_flag[0] = is_done;
      *static_cast<int*>(cluster.map_shared_rank(done_flag, 1)) = is_done;
      *static_cast<int*>(cluster.map_shared_rank(done_flag, 2)) = is_done;
    }
    __syncthreads();
    // Rotations for global round 0 into buffer 0.
    if (tid < PAIRS) {
      int p, q;
      round_robin_pair<N>(0, tid, p, q);
      const float app = Afull[p * S + p];
      const float aqq = Afull[q * S + q];
      const float apq = Afull[p * S + q];
      float c = 1.0f;
      float s = 0.0f;
      if (fabsf(apq) > 0.0f) {
        const float theta = 0.5f * (aqq - app) / apq;
        const float t = copysignf(
            1.0f / (fabsf(theta) + sqrtf(1.0f + theta * theta)), theta);
        c = rsqrtf(1.0f + t * t);
        s = t * c;
        Afull[p * S + p] = app - t * apq;
        Afull[q * S + q] = aqq + t * apq;
        Afull[p * S + q] = 0.0f;
        Afull[q * S + p] = 0.0f;
      }
      const float4 rot = make_float4(c, s, __int_as_float(p), __int_as_float(q));
      rot_buf[tid] = rot;
      static_cast<float4*>(cluster.map_shared_rank(rot_buf, 1))[tid] = rot;
      static_cast<float4*>(cluster.map_shared_rank(rot_buf, 2))[tid] = rot;
    }
  } else {
    // Q band = identity slice.
    const int band = (rank - 1) * HALF;
    for (int t = tid; t < HALF * N; t += THREADS) {
      const int i = t / N;
      const int j = t % N;
      Qband[i * S + j] = (band + i) == j ? 1.0f : 0.0f;
    }
  }
  cluster.sync();

  // --- sweeps ---
  int sweeps_run = 0;
  bool finished = done_flag[0] != 0;
  unsigned long long dbg_update = 0, dbg_gen = 0, dbg_sync = 0;
  const bool dbg = sweep_count != nullptr && tid == 0;
  for (int sweep = 0; sweep < max_sweeps && !finished; ++sweep) {
    ++sweeps_run;
    for (int round = 0; round < ROUNDS; ++round) {
      const int g = sweep * ROUNDS + round;
      const float4* const rot_cur = rot_buf + (g & 1) * PAIRS;
      const bool sweep_end = round == ROUNDS - 1;
      unsigned long long c0 = dbg ? clock64() : 0, c1 = 0, c2 = 0;

      if (rank == 0) {
        // Fused two-sided update of all off blocks. All of a thread's loads
        // are issued into registers BEFORE any store: distinct blocks touch
        // disjoint elements, so the load and store phases can't alias, but
        // the compiler cannot prove that for interleaved SMEM ld/st — the
        // split is what buys memory-level parallelism.
        constexpr int UA_BATCH = cdiv(NOFF, THREADS);
        float xs[UA_BATCH][4];
        float cs[UA_BATCH][4];
        int idx[UA_BATCH][4];
        #pragma unroll
        for (int it = 0; it < UA_BATCH; ++it) {
          const int t = tid + it * THREADS;
          const bool valid = t < NOFF;
          const int tt = valid ? t : 0;
          const float4 r1 = rot_cur[tab1[tt]];
          const float4 r2 = rot_cur[tab2[tt]];
          const int p1 = __float_as_int(r1.z);
          const int q1 = __float_as_int(r1.w);
          const int p2 = __float_as_int(r2.z);
          const int q2 = __float_as_int(r2.w);
          idx[it][0] = valid ? p1 : -1;
          idx[it][1] = q1;
          idx[it][2] = p2;
          idx[it][3] = q2;
          cs[it][0] = r1.x;
          cs[it][1] = r1.y;
          cs[it][2] = r2.x;
          cs[it][3] = r2.y;
          xs[it][0] = Afull[p1 * S + p2];
          xs[it][1] = Afull[p1 * S + q2];
          xs[it][2] = Afull[q1 * S + p2];
          xs[it][3] = Afull[q1 * S + q2];
        }
        #pragma unroll
        for (int it = 0; it < UA_BATCH; ++it) {
          const int p1 = idx[it][0];
          if (p1 < 0) continue;
          const int q1 = idx[it][1];
          const int p2 = idx[it][2];
          const int q2 = idx[it][3];
          const float c1 = cs[it][0];
          const float s1 = cs[it][1];
          const float c2 = cs[it][2];
          const float s2 = cs[it][3];
          const float ypp = c1 * xs[it][0] - s1 * xs[it][2];
          const float ypq = c1 * xs[it][1] - s1 * xs[it][3];
          const float yqp = s1 * xs[it][0] + c1 * xs[it][2];
          const float yqq = s1 * xs[it][1] + c1 * xs[it][3];
          const float zpp = c2 * ypp - s2 * ypq;
          const float zpq = s2 * ypp + c2 * ypq;
          const float zqp = c2 * yqp - s2 * yqq;
          const float zqq = s2 * yqp + c2 * yqq;
          Afull[p1 * S + p2] = zpp;
          Afull[p1 * S + q2] = zpq;
          Afull[q1 * S + p2] = zqp;
          Afull[q1 * S + q2] = zqq;
          Afull[p2 * S + p1] = zpp;
          Afull[q2 * S + p1] = zpq;
          Afull[p2 * S + q1] = zqp;
          Afull[q2 * S + q1] = zqq;
        }
        __syncthreads();
        if (dbg) c1 = clock64();

        if (sweep_end) {
          // Convergence check: accumulate strictly-off-diagonal squares only
          // (summing everything and subtracting the diagonal cancels
          // catastrophically in fp32 once nearly converged).
          float off2_local = 0.0f;
          for (int t = tid; t < N * N; t += THREADS) {
            const int i = t / N;
            const int j = t % N;
            if (i < j) {
              const float value = Afull[i * S + j];
              off2_local += value * value;
            }
          }
          off2_local = warp_sum(off2_local);
          if ((tid & 31) == 0) partials[tid / 32] = off2_local;
          __syncthreads();
          if (tid == 0) {
            float total = 0.0f;
            for (int w = 0; w < WARPS; ++w) total += partials[w];
            // 2x off-diagonal weight (packed stores each pair once).
            const int is_done = 2.0f * total <= tolerance2 * fro2_shared;
            done_flag[0] = is_done;
            *static_cast<int*>(cluster.map_shared_rank(done_flag, 1)) = is_done;
            *static_cast<int*>(cluster.map_shared_rank(done_flag, 2)) = is_done;
          }
          __syncthreads();
        }

        // Rotations for round g+1 into the other buffer.
        if (tid < PAIRS) {
          const int next_round = sweep_end ? 0 : round + 1;
          int p, q;
          round_robin_pair<N>(next_round, tid, p, q);
          const float app = Afull[p * S + p];
          const float aqq = Afull[q * S + q];
          const float apq = Afull[p * S + q];
          float c = 1.0f;
          float s = 0.0f;
          if (fabsf(apq) > 0.0f) {
            const float theta = 0.5f * (aqq - app) / apq;
            const float t = copysignf(
                1.0f / (fabsf(theta) + sqrtf(1.0f + theta * theta)), theta);
            c = rsqrtf(1.0f + t * t);
            s = t * c;
            Afull[p * S + p] = app - t * apq;
            Afull[q * S + q] = aqq + t * apq;
            Afull[p * S + q] = 0.0f;
            Afull[q * S + p] = 0.0f;
          }
          const float4 rot = make_float4(c, s, __int_as_float(p), __int_as_float(q));
          float4* const dst = rot_buf + ((g + 1) & 1) * PAIRS;
          dst[tid] = rot;
          static_cast<float4*>(cluster.map_shared_rank(dst, 1))[tid] = rot;
          static_cast<float4*>(cluster.map_shared_rank(dst, 2))[tid] = rot;
        }
      } else {
        // Q-CTA: rotate the band's column pairs. Same load/store phase split
        // as the A update, in chunks to bound register pressure.
        constexpr int UQ_TOTAL = PAIRS * HALF;
        constexpr int UQ_BATCH = 8;
        constexpr int UQ_CHUNKS = cdiv(cdiv(UQ_TOTAL, THREADS), UQ_BATCH);
        #pragma unroll
        for (int chunk = 0; chunk < UQ_CHUNKS; ++chunk) {
          float qv[UQ_BATCH][2];
          float qcs[UQ_BATCH][2];
          int qidx[UQ_BATCH][2];
          #pragma unroll
          for (int it = 0; it < UQ_BATCH; ++it) {
            const int t = tid + (chunk * UQ_BATCH + it) * THREADS;
            const bool valid = t < UQ_TOTAL;
            const int tt = valid ? t : 0;
            const int k = tt / HALF;
            const int i = tt - k * HALF;
            const float4 r = rot_cur[k];
            const int p = __float_as_int(r.z);
            const int q = __float_as_int(r.w);
            qidx[it][0] = valid ? i * S + p : -1;
            qidx[it][1] = i * S + q;
            qcs[it][0] = r.x;
            qcs[it][1] = r.y;
            qv[it][0] = Qband[i * S + p];
            qv[it][1] = Qband[i * S + q];
          }
          #pragma unroll
          for (int it = 0; it < UQ_BATCH; ++it) {
            if (qidx[it][0] < 0) continue;
            Qband[qidx[it][0]] = qcs[it][0] * qv[it][0] - qcs[it][1] * qv[it][1];
            Qband[qidx[it][1]] = qcs[it][1] * qv[it][0] + qcs[it][0] * qv[it][1];
          }
        }
      }
      if (dbg) c2 = clock64();
      cluster.sync();
      if (dbg) {
        const unsigned long long c3 = clock64();
        dbg_update += (rank == 0 ? c1 : c2) - c0;
        dbg_gen += rank == 0 ? c2 - c1 : 0;
        dbg_sync += c3 - c2;
      }
      finished = done_flag[0] != 0;
      if (finished) break;
    }
  }
  // Debug phase breakdown for matrix 0 (units: cycles/1024), stored past the
  // per-matrix sweep counts.
  if (dbg && matrix == 0) {
    const int base = gridDim.x / 3 + rank * 3;
    sweep_count[base + 0] = static_cast<int>(dbg_update >> 10);
    sweep_count[base + 1] = static_cast<int>(dbg_gen >> 10);
    sweep_count[base + 2] = static_cast<int>(dbg_sync >> 10);
  }

  // --- sort + writeout ---
  if (rank == 0) {
    if (sweep_count != nullptr && tid == 0) sweep_count[matrix] = sweeps_run;
    for (int i = tid; i < NP2; i += THREADS) {
      sort_key[i] = i < N ? Afull[i * S + i] : CUDART_INF_F;
      sort_idx[i] = i;
    }
    __syncthreads();
    for (int k = 2; k <= NP2; k <<= 1) {
      for (int j = k >> 1; j > 0; j >>= 1) {
        for (int i = tid; i < NP2; i += THREADS) {
          const int ixj = i ^ j;
          if (ixj > i) {
            const float a_val = sort_key[i];
            const float b_val = sort_key[ixj];
            const bool ascending = (i & k) == 0;
            if (ascending ? (a_val > b_val) : (a_val < b_val)) {
              sort_key[i] = b_val;
              sort_key[ixj] = a_val;
              const int tmp = sort_idx[i];
              sort_idx[i] = sort_idx[ixj];
              sort_idx[ixj] = tmp;
            }
          }
        }
        __syncthreads();
      }
    }
    if (tid < N) {
      l_out[tid] = sort_key[tid];
      const int idx = sort_idx[tid];
      perm[tid] = idx;
      static_cast<int*>(cluster.map_shared_rank(perm, 1))[tid] = idx;
      static_cast<int*>(cluster.map_shared_rank(perm, 2))[tid] = idx;
    }
  }
  cluster.sync();

  if (rank != 0) {
    const int band = (rank - 1) * HALF;
    for (int t = tid; t < HALF * N; t += THREADS) {
      const int i = t / N;
      const int j = t % N;
      q_out[(band + i) * N + j] = Qband[i * S + perm[j]];
    }
  }
}

void launch_jacobi_cluster(
    const float* input,
    float* q,
    float* l,
    int batch,
    int n,
    float tol2,
    int max_sweeps,
    int* sweep_count) {
#define LAUNCH_JACOBI_CLUSTER(N, THREADS) \
  if (n == N) { \
    constexpr int kPairs = N / 2; \
    constexpr int kNoff = kPairs * (kPairs - 1) / 2; \
    constexpr int kNp2 = next_pow2(N); \
    constexpr int kTri = N * (N + 1) / 2; \
    constexpr size_t kComm = \
        2 * kPairs * sizeof(float4) + (4 + N) * sizeof(int); \
    constexpr size_t kTail0 = \
        (N * (N + 1) + kNp2) * sizeof(float) + kNp2 * sizeof(int) + 2 * kNoff; \
    constexpr size_t kTailQ = (N / 2) * (N + 1) * sizeof(float); \
    constexpr size_t kBytes = \
        kComm + (kTail0 > kTailQ ? kTail0 : kTailQ) + 256; \
    auto* const kfn = jacobi_cluster_kernel<N, THREADS>; \
    static bool configured = [kfn] { \
      cudaFuncSetAttribute( \
          kfn, cudaFuncAttributeMaxDynamicSharedMemorySize, kBytes); \
      return true; \
    }(); \
    (void)configured; \
    cudaLaunchConfig_t config = {}; \
    config.gridDim = dim3(3 * batch); \
    config.blockDim = dim3(THREADS); \
    config.dynamicSmemBytes = kBytes; \
    cudaLaunchAttribute attrs[1]; \
    attrs[0].id = cudaLaunchAttributeClusterDimension; \
    attrs[0].val.clusterDim.x = 3; \
    attrs[0].val.clusterDim.y = 1; \
    attrs[0].val.clusterDim.z = 1; \
    config.attrs = attrs; \
    config.numAttrs = 1; \
    cudaLaunchKernelEx(&config, kfn, input, q, l, tol2, max_sweeps, sweep_count); \
    return; \
  }

  LAUNCH_JACOBI_CLUSTER(176, 512)

#undef LAUNCH_JACOBI_CLUSTER
}

// ===========================================================================
// GEMM-structured block Jacobi (n = 352): one matrix per 3-CTA cluster.
//
// A is a 22x22 grid of 16x16 blocks. Each round pairs the 22 block indices
// into 11 disjoint pairs (round-robin tournament, 21 rounds/sweep); the
// 32x32 pivot of each pair is diagonalized by a warp-scope element-Jacobi
// mini-solve (ONE cyclic sweep = 31 warp-synchronous rounds on padded SMEM;
// measured in simulation: 8 global sweeps to off2 <= 1e-7 * fro2, gate
// margins 12-16x), and the resulting 32x32 rotations V_k are applied to A
// (pair-tile x pair-tile, two-sided) and to Q^T (pair-row panels,
// one-sided) as register-tiled FP32 warp GEMMs on GMEM-resident data (the
// 40-matrix working set is ~40 MB: L2-resident; full SMEM residency is
// impossible at this shape — measured/settled in NOTES.md).
//
// Pipelining (the mini-eig serial chain must hide behind the panel work):
// during round g the 4 producer warps of each CTA solve the pivots of round
// g+1 while the 12 consumer warps apply round g's rotations. A round-(g+1)
// pivot (p', q') depends only on (a) the D-tiles written by round-g-1's
// producers and (b) ONE round-g cross tile — which the producer computes
// itself before solving. Verified exhaustively for NB=22 over all 21
// rounds (incl. the sweep wraparound): the 11 producer->cross-tile mappings
// of a round are always DISTINCT and a next pair never lies inside one
// current pair, so each priority tile is read+written by exactly one warp
// (program order safe) and skipped by consumers. That makes each round need
// exactly ONE cluster.sync and no cross-warp handoff. D-tile writes are
// deferred at sweep boundaries so the convergence check and the early exit
// see a consistent matrix state.
//
// Precision: everything FP32. TF32 tile GEMMs were simulated and KILLED:
// A-update rounding is permanent w.r.t. the original A (eigen residual
// 2.2x over gate even with FP32 final sweeps).
// ===========================================================================

namespace bj {

constexpr int BS = 16;   // block size
constexpr int TD = 32;   // pair-tile dim = 2*BS
constexpr int PST = TD + 1;  // padded pivot/VT stride (33)
constexpr int SST = TD + 4;  // scratch stride (36; 16B-aligned rows)

template <int N> struct Cfg {
  static constexpr int NB = N / BS;                      // 22 block indices
  static constexpr int NPAIR = NB / 2;                   // 11 pairs/round
  static constexpr int NROUND = NB - 1;                  // 21 rounds/sweep
  static constexpr int NCROSS = NPAIR * (NPAIR - 1) / 2; // 55 cross tiles
  static constexpr int NQCHUNK = N / TD;                 // 11 Q col chunks
  static constexpr int NTASK = NCROSS + NPAIR * NQCHUNK; // 176 tasks
  static constexpr int NP2 = next_pow2(N);
  static constexpr int NWARP = 16;                       // 512 threads
  static constexpr int NCONS = NWARP - 4;                // consumer warps/CTA
  // dynamic SMEM (floats): V double buffer, 4 pivot + 4 V^T workspaces,
  // per-warp GEMM scratch, sort keys, static mini-eig schedule tables.
  static constexpr int F_VBUF = 2 * NPAIR * TD * TD;
  static constexpr int F_PWS = 4 * TD * PST;
  static constexpr int F_VTWS = 4 * TD * SST;   // fl4-aligned rows
  static constexpr int F_SCR = NWARP * TD * SST;
  static constexpr int F_ROT = 4 * 16 * 2;      // per-producer float2 rots
  static constexpr int F_KEY = NP2;
  static constexpr int I_IDX = NP2;
  static constexpr int I_PERM = N;
  static constexpr int I_BLK = 31 * 120;        // packed p1q1p2q2 per round
  static constexpr int I_TABPQ = 31 * 16;       // packed p,q per round/pair
  static constexpr int I_MISC = 8 * NPAIR + 2 * NB + 64;  // tables + flags
  static constexpr size_t BYTES =
      (F_VBUF + F_PWS + F_VTWS + F_SCR + F_ROT + F_KEY) * sizeof(float) +
      (I_IDX + I_PERM + I_BLK + I_TABPQ + I_MISC) * sizeof(int) +
      NCROSS + 2 * NCROSS /*crossIJ*/ + 240 /*minieig tabs*/ + 256;
};

__device__ __forceinline__ int srow(int2 pr, int k) {
  // Row/col of strip element k (0..31) for block pair pr (two 16-strips).
  return k < BS ? pr.x * BS + k : pr.y * BS + (k - BS);
}

__device__ __forceinline__ void fma8x4(float4 acc[8], float4 x0, float4 x1,
                                       float4 y) {
  const float xs[8] = {x0.x, x0.y, x0.z, x0.w, x1.x, x1.y, x1.z, x1.w};
  #pragma unroll
  for (int i = 0; i < 8; ++i) {
    acc[i].x += xs[i] * y.x;
    acc[i].y += xs[i] * y.y;
    acc[i].z += xs[i] * y.z;
    acc[i].w += xs[i] * y.w;
  }
}

// Two-sided pair-tile update: tile (rows = strips of pi, cols = strips of
// pj) gets T' = Vi^T T Vj; both T' and its mirror T'^T are written to the
// full-storage A. Lane owns an 8x4 fragment; the intermediate goes through
// the per-warp scratch (stride 36, conflict-free rows).
template <int N>
__device__ void a_tile_task(float* __restrict__ Aw, const float* __restrict__ Vi,
                            const float* __restrict__ Vj, int2 pi, int2 pj,
                            float* __restrict__ sc) {
  const int lane = threadIdx.x & 31;
  const int a0 = (lane >> 3) * 8;
  const int b0 = (lane & 7) * 4;
  const int ca = srow(pj, a0);   // global col of tile col a0 (8 contiguous)
  // Stage T into scratch with 8 batched global loads per lane (full MLP;
  // per-k loads inside the GEMM loop serialize on L2 latency instead).
  {
    const int seg = (lane & 7) * 4;
    const int cg = srow(pj, seg);
    float4 tv[8];
    #pragma unroll
    for (int it = 0; it < 8; ++it) {
      const int k = it * 4 + (lane >> 3);
      tv[it] = *reinterpret_cast<const float4*>(
          Aw + static_cast<size_t>(srow(pi, k)) * N + cg);
    }
    __syncwarp();
    #pragma unroll
    for (int it = 0; it < 8; ++it)
      *reinterpret_cast<float4*>(sc + (it * 4 + (lane >> 3)) * SST + seg) =
          tv[it];
  }
  __syncwarp();
  float4 acc[8];
  #pragma unroll
  for (int i = 0; i < 8; ++i) acc[i] = make_float4(0.f, 0.f, 0.f, 0.f);
  // GEMM1: MT[c][r] = sum_k T[k][c] * Vi[k][r]  (MT = (Vi^T T)^T)
  #pragma unroll
  for (int k = 0; k < TD; ++k) {
    const float* xr = sc + k * SST + a0;
    const float4 x0 = *reinterpret_cast<const float4*>(xr);
    const float4 x1 = *reinterpret_cast<const float4*>(xr + 4);
    const float4 y = *reinterpret_cast<const float4*>(Vi + k * TD + b0);
    fma8x4(acc, x0, x1, y);
  }
  __syncwarp();
  #pragma unroll
  for (int i = 0; i < 8; ++i)
    *reinterpret_cast<float4*>(sc + (a0 + i) * SST + b0) = acc[i];
  __syncwarp();
  // GEMM2: TT[c'][r] = sum_k Vj[k][c'] * MT[k][r]  (TT = T'^T)
  #pragma unroll
  for (int i = 0; i < 8; ++i) acc[i] = make_float4(0.f, 0.f, 0.f, 0.f);
  #pragma unroll
  for (int k = 0; k < TD; ++k) {
    const float* xr = Vj + k * TD + a0;
    const float4 x0 = *reinterpret_cast<const float4*>(xr);
    const float4 x1 = *reinterpret_cast<const float4*>(xr + 4);
    const float4 y = *reinterpret_cast<const float4*>(sc + k * SST + b0);
    fma8x4(acc, x0, x1, y);
  }
  // mirror tile (pj, pi): element TT[c'][r] lands at row srow(pj,c'),
  // col srow(pi,r) — direct float4 stores.
  const int cb = srow(pi, b0);
  #pragma unroll
  for (int i = 0; i < 8; ++i)
    *reinterpret_cast<float4*>(
        Aw + static_cast<size_t>(srow(pj, a0 + i)) * N + cb) = acc[i];
  __syncwarp();
  #pragma unroll
  for (int i = 0; i < 8; ++i)
    *reinterpret_cast<float4*>(sc + (a0 + i) * SST + b0) = acc[i];
  __syncwarp();
  // upper tile (pi, pj): out[r][c'] = TT[c'][r] = sc[c'][r] (transpose read).
  const int cu = srow(pj, b0);
  #pragma unroll
  for (int i = 0; i < 8; ++i) {
    const int r = a0 + i;
    float4 v;
    v.x = sc[(b0 + 0) * SST + r];
    v.y = sc[(b0 + 1) * SST + r];
    v.z = sc[(b0 + 2) * SST + r];
    v.w = sc[(b0 + 3) * SST + r];
    *reinterpret_cast<float4*>(
        Aw + static_cast<size_t>(srow(pi, r)) * N + cu) = v;
  }
}

// One-sided Q^T panel update: rows = strips of pair pk, cols = chunk cc.
// QT[R] <- Vk^T QT[R]:  C[r][c] = sum_k Vk[k][r] * QT[R(k)][c].
template <int N>
__device__ void q_tile_task(float* __restrict__ QTw, const float* __restrict__ Vk,
                            int2 pk, int cc, float* __restrict__ sc) {
  const int lane = threadIdx.x & 31;
  const int a0 = (lane >> 3) * 8;   // r
  const int b0 = (lane & 7) * 4;    // c
  const int col = cc * TD + b0;
  // Stage the Q^T panel into scratch (batched global loads, full MLP).
  {
    const int seg = (lane & 7) * 4;
    float4 tv[8];
    #pragma unroll
    for (int it = 0; it < 8; ++it) {
      const int k = it * 4 + (lane >> 3);
      tv[it] = *reinterpret_cast<const float4*>(
          QTw + static_cast<size_t>(srow(pk, k)) * N + cc * TD + seg);
    }
    __syncwarp();
    #pragma unroll
    for (int it = 0; it < 8; ++it)
      *reinterpret_cast<float4*>(sc + (it * 4 + (lane >> 3)) * SST + seg) =
          tv[it];
  }
  __syncwarp();
  float4 acc[8];
  #pragma unroll
  for (int i = 0; i < 8; ++i) acc[i] = make_float4(0.f, 0.f, 0.f, 0.f);
  #pragma unroll
  for (int k = 0; k < TD; ++k) {
    const float* xr = Vk + k * TD + a0;
    const float4 x0 = *reinterpret_cast<const float4*>(xr);
    const float4 x1 = *reinterpret_cast<const float4*>(xr + 4);
    const float4 y = *reinterpret_cast<const float4*>(sc + k * SST + b0);
    fma8x4(acc, x0, x1, y);
  }
  #pragma unroll
  for (int i = 0; i < 8; ++i)
    *reinterpret_cast<float4*>(
        QTw + static_cast<size_t>(srow(pk, a0 + i)) * N + col) = acc[i];
  __syncwarp();
}

// Warp-scope element Jacobi on the 32x32 pivot P (padded stride 33): one
// cyclic sweep (31 rounds), accumulating V^T (float4 rows, stride 36).
// The pair schedule is STATIC, so the round's block addresses come from
// precomputed SMEM tables (sBlk: packed p1|q1|p2|q2 per block, sTabPQ:
// packed p|q per pair) — block loads issue before the rotation-generation
// divide/sqrt chain completes, and no shuffles sit on the critical path.
// Rotations pass through a per-producer float2 buffer.
__device__ void mini_eig(float* __restrict__ P, float* __restrict__ VT,
                         float2* __restrict__ rot,
                         const uint8_t* __restrict__ tab1,
                         const uint8_t* __restrict__ tab2,
                         const unsigned* __restrict__ blk,
                         const unsigned* __restrict__ tabpq) {
  const int lane = threadIdx.x & 31;
  #pragma unroll
  for (int r = 0; r < TD; ++r) VT[r * SST + lane] = r == lane ? 1.f : 0.f;
  __syncwarp();
  for (int rnd = 0; rnd < TD - 1; ++rnd) {
    // Block-update load phase: addresses are static (independent of the
    // rotations), so these SMEM loads overlap the rotgen latency chain.
    unsigned pk[4];
    float xs[4][4];
    int k1v[4], k2v[4];
    #pragma unroll
    for (int it = 0; it < 4; ++it) {
      const int t = lane + it * 32;
      const bool valid = t < 120;
      const int tt = valid ? t : 0;
      pk[it] = blk[rnd * 120 + tt];
      k1v[it] = valid ? tab1[tt] : -1;
      k2v[it] = tab2[tt];
      const int p1 = pk[it] & 255, q1 = (pk[it] >> 8) & 255;
      const int p2 = (pk[it] >> 16) & 255, q2 = pk[it] >> 24;
      xs[it][0] = P[p1 * PST + p2];
      xs[it][1] = P[p1 * PST + q2];
      xs[it][2] = P[q1 * PST + p2];
      xs[it][3] = P[q1 * PST + q2];
    }
    // Rotation generation (lanes 0..15) — overlapped with the loads above.
    if (lane < 16) {
      const unsigned pq = tabpq[rnd * 16 + lane];
      const int p = pq & 255, q = pq >> 8;
      const float app = P[p * PST + p];
      const float aqq = P[q * PST + q];
      const float apq = P[p * PST + q];
      float c = 1.f, s = 0.f;
      if (fabsf(apq) > 0.f) {
        const float theta = 0.5f * (aqq - app) / apq;
        const float t = copysignf(
            1.f / (fabsf(theta) + sqrtf(1.f + theta * theta)), theta);
        c = rsqrtf(1.f + t * t);
        s = t * c;
        P[p * PST + p] = app - t * apq;
        P[q * PST + q] = aqq + t * apq;
        P[p * PST + q] = 0.f;
        P[q * PST + p] = 0.f;
      }
      rot[lane] = make_float2(c, s);
    }
    __syncwarp();
    #pragma unroll
    for (int it = 0; it < 4; ++it) {
      if (k1v[it] < 0) continue;
      const float2 r1 = rot[k1v[it]];
      const float2 r2 = rot[k2v[it]];
      const int p1 = pk[it] & 255, q1 = (pk[it] >> 8) & 255;
      const int p2 = (pk[it] >> 16) & 255, q2 = pk[it] >> 24;
      const float ypp = r1.x * xs[it][0] - r1.y * xs[it][2];
      const float ypq = r1.x * xs[it][1] - r1.y * xs[it][3];
      const float yqp = r1.y * xs[it][0] + r1.x * xs[it][2];
      const float yqq = r1.y * xs[it][1] + r1.x * xs[it][3];
      const float zpp = r2.x * ypp - r2.y * ypq;
      const float zpq = r2.y * ypp + r2.x * ypq;
      const float zqp = r2.x * yqp - r2.y * yqq;
      const float zqq = r2.y * yqp + r2.x * yqq;
      P[p1 * PST + p2] = zpp;
      P[p1 * PST + q2] = zpq;
      P[q1 * PST + p2] = zqp;
      P[q1 * PST + q2] = zqq;
      P[p2 * PST + p1] = zpp;
      P[q2 * PST + p1] = zpq;
      P[p2 * PST + q1] = zqp;
      P[q2 * PST + q1] = zqq;
    }
    // V^T row mix, float4 columns: pair k = it*4 + (lane>>3), cols (lane&7)*4.
    // Two batches of 2 keep peak register pressure down (128-reg cap).
    const int vk = lane >> 3;
    const int vj = (lane & 7) * 4;
    #pragma unroll
    for (int half = 0; half < 2; ++half) {
      float4 vv[2][2];
      float2 vr[2];
      int vp[2], vq[2];
      #pragma unroll
      for (int it = 0; it < 2; ++it) {
        const int k = (half * 2 + it) * 4 + vk;
        const unsigned pq = tabpq[rnd * 16 + k];
        vr[it] = rot[k];
        vp[it] = pq & 255;
        vq[it] = pq >> 8;
        vv[it][0] = *reinterpret_cast<const float4*>(VT + vp[it] * SST + vj);
        vv[it][1] = *reinterpret_cast<const float4*>(VT + vq[it] * SST + vj);
      }
      #pragma unroll
      for (int it = 0; it < 2; ++it) {
        const float c = vr[it].x, s = vr[it].y;
        float4 a = vv[it][0], b = vv[it][1];
        float4 na, nb;
        na.x = c * a.x - s * b.x;
        na.y = c * a.y - s * b.y;
        na.z = c * a.z - s * b.z;
        na.w = c * a.w - s * b.w;
        nb.x = s * a.x + c * b.x;
        nb.y = s * a.y + c * b.y;
        nb.z = s * a.z + c * b.z;
        nb.w = s * a.w + c * b.w;
        *reinterpret_cast<float4*>(VT + vp[it] * SST + vj) = na;
        *reinterpret_cast<float4*>(VT + vq[it] * SST + vj) = nb;
      }
    }
    __syncwarp();
  }
}

}  // namespace bj

template <int N, int THREADS>
__global__ __launch_bounds__(THREADS, 1)
void block_jacobi_kernel(
    const float* __restrict__ input,
    float* __restrict__ q_out,
    float* __restrict__ l_out,
    float* __restrict__ a_work,
    float* __restrict__ qt_work,
    float tolerance2,
    int max_sweeps,
    int* __restrict__ sweep_count) {
  using C = bj::Cfg<N>;
  using bj::srow;
  constexpr int TD = bj::TD;
  constexpr int PST = bj::PST;
  constexpr int SST = bj::SST;
  constexpr int NPAIR = C::NPAIR;
  constexpr int NROUND = C::NROUND;
  constexpr int NCROSS = C::NCROSS;
  constexpr int NTASK = C::NTASK;
  constexpr int NP2 = C::NP2;
  constexpr int NWARP = C::NWARP;
  static_assert(THREADS == 32 * C::NWARP);

  cg::cluster_group cluster = cg::this_cluster();
  const int rank = cluster.block_rank();
  const int matrix = blockIdx.x / 3;
  const int tid = threadIdx.x;
  const int warp = tid >> 5;
  const int lane = tid & 31;

  input += static_cast<size_t>(matrix) * N * N;
  q_out += static_cast<size_t>(matrix) * N * N;
  l_out += static_cast<size_t>(matrix) * N;
  float* const Aw = a_work + static_cast<size_t>(matrix) * N * N;
  float* const QTw = qt_work + static_cast<size_t>(matrix) * N * N;

  // --- dynamic SMEM carve (identical layout in every CTA) ---
  extern __shared__ float smem[];
  float* const vbuf = smem;                       // [2][NPAIR][TD*TD]
  float* const pws = vbuf + C::F_VBUF;            // [4][TD*PST]
  float* const vtws = pws + C::F_PWS;             // [4][TD*SST]
  float* const scr = vtws + C::F_VTWS;            // [NWARP][TD*SST]
  float* const rotws = scr + C::F_SCR;            // [4][16] float2
  float* const sKey = rotws + C::F_ROT;           // [NP2]
  int* const sIdx = reinterpret_cast<int*>(sKey + C::F_KEY);  // [NP2]
  int* const sPerm = sIdx + C::I_IDX;             // [N]
  unsigned* const sBlk = reinterpret_cast<unsigned*>(sPerm + C::I_PERM);
  unsigned* const sTabPQ = sBlk + C::I_BLK;       // [31*16]
  int* const sMisc = reinterpret_cast<int*>(sTabPQ + C::I_TABPQ);
  int2* const sPairsCur = reinterpret_cast<int2*>(sMisc);          // [NPAIR]
  int2* const sPairsNxt = sPairsCur + NPAIR;                       // [NPAIR]
  int2* const sProd = sPairsNxt + NPAIR;   // producer k' -> (k1, k2)
  int* const sPairOf = reinterpret_cast<int*>(sProd + NPAIR);      // [NB]
  int* const sFlags = sPairOf + C::NB + (C::NB & 1);
  // sFlags[0] = done, sFlags[1..3] partial slots base as float
  float* const sFro2 = reinterpret_cast<float*>(sFlags + 4);
  float* const sPartial = sFro2 + 4;                               // [NWARP]
  uint8_t* const sSkip = reinterpret_cast<uint8_t*>(sPartial + NWARP);  // [NCROSS]
  uint8_t* const sCrossI = sSkip + ((NCROSS + 3) & ~3);            // [NCROSS]
  uint8_t* const sCrossJ = sCrossI + NCROSS;                       // [NCROSS]
  uint8_t* const sTab1 = sCrossJ + NCROSS + 2;                     // [120]
  uint8_t* const sTab2 = sTab1 + 120;                              // [120]

  // --- one-time tables ---
  // mini-eig triangular decode (16 inner pairs -> 120 pair-of-pairs)
  for (int t = tid; t < 120; t += THREADS) {
    int k1 = 0, rem = t;
    while (rem >= 15 - k1) {
      rem -= 15 - k1;
      ++k1;
    }
    sTab1[t] = static_cast<uint8_t>(k1);
    sTab2[t] = static_cast<uint8_t>(k1 + 1 + rem);
  }
  // cross-tile triangular decode (NPAIR pairs -> NCROSS tiles)
  for (int t = tid; t < NCROSS; t += THREADS) {
    int k1 = 0, rem = t;
    while (rem >= NPAIR - 1 - k1) {
      rem -= NPAIR - 1 - k1;
      ++k1;
    }
    sCrossI[t] = static_cast<uint8_t>(k1);
    sCrossJ[t] = static_cast<uint8_t>(k1 + 1 + rem);
  }
  // static mini-eig schedules: per inner round, packed pair table and
  // packed per-block (p1,q1,p2,q2) addresses
  for (int t = tid; t < 31 * 16; t += THREADS) {
    int p, q;
    round_robin_pair<TD>(t / 16, t % 16, p, q);
    sTabPQ[t] = static_cast<unsigned>(p | (q << 8));
  }
  for (int t = tid; t < 31 * 120; t += THREADS) {
    const int rnd = t / 120;
    int k1 = 0, rem = t % 120;
    while (rem >= 15 - k1) {
      rem -= 15 - k1;
      ++k1;
    }
    const int k2 = k1 + 1 + rem;
    int p1, q1, p2, q2;
    round_robin_pair<TD>(rnd, k1, p1, q1);
    round_robin_pair<TD>(rnd, k2, p2, q2);
    sBlk[t] = static_cast<unsigned>(p1 | (q1 << 8) | (p2 << 16) | (q2 << 24));
  }

  // --- init: QT = I; A_work = symmetrized input; fro2 ---
  for (int t = rank * THREADS + tid; t < N * N; t += 3 * THREADS) {
    QTw[t] = (t / N) == (t % N) ? 1.f : 0.f;
  }
  float fro2_local = 0.f;
  {
    // natural 32x32 tile pairs (I <= J), staged transpose through scratch
    constexpr int NT = N / TD;                 // 11
    constexpr int NUP = NT * (NT + 1) / 2;     // 66
    float* const sc = scr + warp * TD * SST;
    const int a0 = (lane >> 3) * 8;
    const int b0 = (lane & 7) * 4;
    for (int t = rank * NWARP + warp; t < NUP; t += 3 * NWARP) {
      int I = 0, rem = t;
      while (rem >= NT - I) {
        rem -= NT - I;
        ++I;
      }
      const int J = I + rem;
      // stage tile (J, I) rows into scratch
      #pragma unroll
      for (int it = 0; it < 8; ++it) {
        const int r = it * 4 + (lane >> 3);
        const int seg = (lane & 7) * 4;
        *reinterpret_cast<float4*>(sc + r * SST + seg) =
            *reinterpret_cast<const float4*>(
                input + static_cast<size_t>(J * TD + r) * N + I * TD + seg);
      }
      __syncwarp();
      float4 frag[8];
      #pragma unroll
      for (int i = 0; i < 8; ++i) {
        const int r = a0 + i;
        const float4 x = *reinterpret_cast<const float4*>(
            input + static_cast<size_t>(I * TD + r) * N + J * TD + b0);
        float4 v;
        v.x = 0.5f * (x.x + sc[(b0 + 0) * SST + r]);
        v.y = 0.5f * (x.y + sc[(b0 + 1) * SST + r]);
        v.z = 0.5f * (x.z + sc[(b0 + 2) * SST + r]);
        v.w = 0.5f * (x.w + sc[(b0 + 3) * SST + r]);
        frag[i] = v;
        const float w = I == J ? 1.f : 2.f;
        fro2_local += w * (v.x * v.x + v.y * v.y + v.z * v.z + v.w * v.w);
        *reinterpret_cast<float4*>(
            Aw + static_cast<size_t>(I * TD + r) * N + J * TD + b0) = v;
      }
      __syncwarp();
      if (I != J) {
        #pragma unroll
        for (int i = 0; i < 8; ++i)
          *reinterpret_cast<float4*>(sc + (a0 + i) * SST + b0) = frag[i];
        __syncwarp();
        #pragma unroll
        for (int i = 0; i < 8; ++i) {
          const int r = a0 + i;
          float4 v;
          v.x = sc[(b0 + 0) * SST + r];
          v.y = sc[(b0 + 1) * SST + r];
          v.z = sc[(b0 + 2) * SST + r];
          v.w = sc[(b0 + 3) * SST + r];
          *reinterpret_cast<float4*>(
              Aw + static_cast<size_t>(J * TD + r) * N + I * TD + b0) = v;
        }
        __syncwarp();
      }
    }
  }
  fro2_local = warp_sum(fro2_local);
  if (lane == 0) sPartial[warp] = fro2_local;
  __syncthreads();
  if (tid == 0) {
    float total = 0.f;
    for (int w = 0; w < NWARP; ++w) total += sPartial[w];
    sFro2[1] = total;   // CTA partial
  }
  cluster.sync();
  if (rank == 0 && tid == 0) {
    float total = sFro2[1];
    for (int r = 1; r < 3; ++r)
      total += static_cast<float*>(cluster.map_shared_rank(sFro2, r))[1];
    for (int r = 0; r < 3; ++r) {
      float* remote = static_cast<float*>(cluster.map_shared_rank(sFro2, r));
      remote[0] = total;
      reinterpret_cast<int*>(
          static_cast<void*>(cluster.map_shared_rank(sFlags, r)))[0] =
          total == 0.f ? 1 : 0;
    }
  }
  cluster.sync();

  const float fro2 = sFro2[0];
  bool done = sFlags[0] != 0;
  // Threshold-Jacobi pivot skip: a pivot whose 32x32 off-diagonal mass is
  // negligible relative to the per-sweep convergence budget (tolerance2*fro2)
  // is already ~diagonal, so its rotation is ~identity. Skip the 31-round
  // mini_eig (leave V as the identity it is initialised to) -- the producer is
  // the wall (consumers idle on cluster.sync waiting for it), so cutting
  // producer mini_eig work shortens the round even though the consumers still
  // apply identity. The skipped pivot's tiny off-diagonals stay in A and are
  // counted by the honest end-of-sweep off2 convergence check (no false
  // convergence). SKIP_FRAC is a small fraction of the total budget per pivot,
  // so total skipped mass stays far under tolerance2*fro2 (no accuracy/sweep
  // regression). SKIP_FRAC<=0 disables (byte-identical).
  constexpr float SKIP_FRAC = 3.0e-2f;
  const float tau2_skip = SKIP_FRAC * tolerance2 * fro2;
  // Producer warps 0-3 spread over the four SM sub-partitions (measured:
  // co-locating all four on SMSP0 stretches the solve 65k->78k cycles — the
  // chains hurt each other more than consumer competition hurts them).
  // After producing, they join the shared task pool.
  const bool is_prod = warp < 4;
  const int slot = is_prod ? warp : 0;
  const int kprime = is_prod ? warp * 3 + rank : NPAIR;  // producer pair id
  float* const Pws = pws + slot * TD * PST;
  float* const VTws = vtws + slot * TD * SST;
  float2* const Rot = reinterpret_cast<float2*>(rotws) + slot * 16;
  float* const sc = scr + warp * TD * SST;

  // batched pivot gather / D write (8 float4 global ops per lane, full MLP)
  auto gather_pivot = [&](int2 pv) {
    const int seg = (lane & 7) * 4;
    float4 gv[8];
    #pragma unroll
    for (int it = 0; it < 8; ++it) {
      const int r = it * 4 + (lane >> 3);
      gv[it] = *reinterpret_cast<const float4*>(
          Aw + static_cast<size_t>(srow(pv, r)) * N + srow(pv, seg));
    }
    #pragma unroll
    for (int it = 0; it < 8; ++it) {
      const int r = it * 4 + (lane >> 3);
      Pws[r * PST + seg + 0] = gv[it].x;
      Pws[r * PST + seg + 1] = gv[it].y;
      Pws[r * PST + seg + 2] = gv[it].z;
      Pws[r * PST + seg + 3] = gv[it].w;
    }
    __syncwarp();
  };
  auto write_d = [&](int2 pv) {
    const int seg = (lane & 7) * 4;
    #pragma unroll
    for (int it = 0; it < 8; ++it) {
      const int r = it * 4 + (lane >> 3);
      float4 v;
      v.x = Pws[r * PST + seg + 0];
      v.y = Pws[r * PST + seg + 1];
      v.z = Pws[r * PST + seg + 2];
      v.w = Pws[r * PST + seg + 3];
      *reinterpret_cast<float4*>(
          Aw + static_cast<size_t>(srow(pv, r)) * N + srow(pv, seg)) = v;
    }
  };
  auto broadcast_v = [&](float* slot) {
    float* const dst0 = static_cast<float*>(cluster.map_shared_rank(slot, 0));
    float* const dst1 = static_cast<float*>(cluster.map_shared_rank(slot, 1));
    float* const dst2 = static_cast<float*>(cluster.map_shared_rank(slot, 2));
    #pragma unroll
    for (int k = 0; k < TD; ++k) {
      const float v = VTws[lane * SST + k];   // V[k][lane]
      dst0[k * TD + lane] = v;
      dst1[k * TD + lane] = v;
      dst2[k * TD + lane] = v;
    }
  };
  // Off-diagonal mass of the gathered 32x32 pivot (both triangles, matching the
  // end-of-sweep convergence metric). Lane l owns row l; warp-reduce to all.
  auto pivot_off2 = [&]() -> float {
    float s = 0.f;
    #pragma unroll
    for (int j = 0; j < TD; ++j) {
      if (j != lane) { const float v = Pws[lane * PST + j]; s += v * v; }
    }
    #pragma unroll
    for (int o = 16; o > 0; o >>= 1) s += __shfl_xor_sync(FULL_MASK, s, o);
    return s;
  };
  // mini_eig unless the pivot is already ~diagonal (threshold skip); on skip,
  // set V = identity so broadcast_v/consumers apply a no-op rotation.
  auto mini_eig_or_skip = [&]() {
    if (SKIP_FRAC > 0.f && pivot_off2() <= tau2_skip) {
      #pragma unroll
      for (int r = 0; r < TD; ++r) VTws[r * SST + lane] = (r == lane) ? 1.f : 0.f;
      __syncwarp();
    } else {
      bj::mini_eig(Pws, VTws, Rot, sTab1, sTab2, sBlk, sTabPQ);
    }
  };

  // --- prime: solve round-0 pivots into vbuf[0], write D(0) ---
  if (!done) {
    if (tid < NPAIR) {
      int p, q;
      round_robin_pair<C::NB>(0, tid, p, q);
      sPairsNxt[tid] = make_int2(p, q);
    }
    __syncthreads();
    if (kprime < NPAIR) {
      const int2 pv = sPairsNxt[kprime];
      gather_pivot(pv);
      mini_eig_or_skip();
      broadcast_v(vbuf + kprime * TD * TD);
      write_d(pv);
    }
  }
  cluster.sync();

  // --- sweeps ---
  int g = 0;
  int sweep = 0;
  // Phase-cycle instrumentation (matrix 0 only, when sweep_count given):
  // [0] producer tile GEMM+gather, [1] producer mini-eig+bcast+D,
  // [2] consumer task loop, [3] wait at the round cluster.sync.
  unsigned long long dbg[4] = {0, 0, 0, 0};
  const bool dbg_on = sweep_count != nullptr && matrix == 0 && lane == 0 &&
                      (warp == 0 || warp == 4);
  while (!done && sweep < max_sweeps) {
    const int rnd = g % NROUND;
    const bool sweep_end = rnd == NROUND - 1;
    // schedule build (identical in every CTA)
    if (tid < NPAIR) {
      int p, q;
      round_robin_pair<C::NB>(rnd, tid, p, q);
      sPairsCur[tid] = make_int2(p, q);
      round_robin_pair<C::NB>((rnd + 1) % NROUND, tid, p, q);
      sPairsNxt[tid] = make_int2(p, q);
    }
    for (int t = tid; t < NCROSS; t += THREADS) sSkip[t] = 0;
    if (tid == 0) sFlags[1] = 0;   // shared task counter
    __syncthreads();
    if (tid < NPAIR) {
      sPairOf[sPairsCur[tid].x] = tid;
      sPairOf[sPairsCur[tid].y] = tid;
    }
    __syncthreads();
    if (tid < NPAIR) {
      const int k1 = sPairOf[sPairsNxt[tid].x];
      const int k2 = sPairOf[sPairsNxt[tid].y];
      const int lo = min(k1, k2), hi = max(k1, k2);
      sProd[tid] = make_int2(lo, hi);
      // linear triangular index of cross tile (lo, hi)
      const int tri = lo * (2 * NPAIR - lo - 1) / 2 + (hi - lo - 1);
      sSkip[tri] = 1;
    }
    __syncthreads();

    float* const vcur = vbuf + (g & 1) * NPAIR * TD * TD;
    float* const vnxt = vbuf + ((g + 1) & 1) * NPAIR * TD * TD;
    const unsigned long long dc0 = dbg_on ? clock64() : 0;
    unsigned long long dc1 = dc0;

    unsigned long long dc1b = dc0;
    if (kprime < NPAIR) {
      // producer: own cross tile, gather, solve, broadcast, (deferred) D
      const int k1 = sProd[kprime].x;
      const int k2 = sProd[kprime].y;
      bj::a_tile_task<N>(Aw, vcur + k1 * TD * TD, vcur + k2 * TD * TD,
                         sPairsCur[k1], sPairsCur[k2], sc);
      __threadfence_block();
      __syncwarp();
      const int2 pv = sPairsNxt[kprime];
      gather_pivot(pv);
      if (dbg_on) dc1 = clock64();
      mini_eig_or_skip();
      broadcast_v(vnxt + kprime * TD * TD);
      if (!sweep_end) write_d(pv);
      if (dbg_on) dc1b = clock64();
    }
    // shared task pool over cross tiles + Q panels: CTA rank r pulls tasks
    // r, r+3, r+6, ... — consumers immediately, producers once done solving
    for (;;) {
      int i = 0;
      if (lane == 0) i = atomicAdd(sFlags + 1, 1);
      i = __shfl_sync(FULL_MASK, i, 0);
      const int t = rank + 3 * i;
      if (t >= NTASK) break;
      if (t < NCROSS) {
        if (sSkip[t]) continue;
        const int ki = sCrossI[t];
        const int kj = sCrossJ[t];
        bj::a_tile_task<N>(Aw, vcur + ki * TD * TD, vcur + kj * TD * TD,
                           sPairsCur[ki], sPairsCur[kj], sc);
      } else {
        const int tq = t - NCROSS;
        const int kq = tq % NPAIR;
        const int cc = tq / NPAIR;
        bj::q_tile_task<N>(QTw, vcur + kq * TD * TD, sPairsCur[kq], cc, sc);
      }
    }
    unsigned long long dc2 = 0;
    if (dbg_on) {
      dc2 = clock64();
      if (warp == 0) {
        dbg[0] += dc1 - dc0;
        dbg[1] += dc1b - dc1;
        dbg[2] += dc2 - dc1b;
      } else {
        dbg[2] += dc2 - dc0;
      }
    }
    cluster.sync();
    if (dbg_on) dbg[3] += clock64() - dc2;

    if (sweep_end) {
      // convergence check on the strictly-off-diagonal mass (state = end of
      // round g: D(g+1) writes are deferred until after this check)
      float off2_local = 0.f;
      for (int t = rank * THREADS + tid; t < N * N; t += 3 * THREADS) {
        const int r = t / N;
        const int c = t % N;
        if (r != c) {
          const float v = Aw[t];
          off2_local += v * v;
        }
      }
      off2_local = warp_sum(off2_local);
      if (lane == 0) sPartial[warp] = off2_local;
      __syncthreads();
      if (tid == 0) {
        float total = 0.f;
        for (int w = 0; w < NWARP; ++w) total += sPartial[w];
        sFro2[1] = total;
      }
      cluster.sync();
      if (rank == 0 && tid == 0) {
        float total = sFro2[1];
        for (int r = 1; r < 3; ++r)
          total +=
              static_cast<float*>(cluster.map_shared_rank(sFro2, r))[1];
        const int is_done = total <= tolerance2 * fro2 ? 1 : 0;
        for (int r = 0; r < 3; ++r)
          reinterpret_cast<int*>(
              static_cast<void*>(cluster.map_shared_rank(sFlags, r)))[0] =
              is_done;
      }
      ++sweep;
      cluster.sync();
      done = sFlags[0] != 0 || sweep >= max_sweeps;
      if (!done && kprime < NPAIR) {
        // deferred D(g+1) writes, now that the check has read the matrix
        write_d(sPairsNxt[kprime]);
      }
      cluster.sync();
    }
    ++g;
  }

  if (sweep_count != nullptr && rank == 0 && tid == 0)
    sweep_count[matrix] = sweep;
  if (dbg_on) {
    // layout: [batch + rank*8 + slot] (units: cycles >> 10)
    const int base = gridDim.x / 3 + rank * 8 + (warp == 0 ? 0 : 4);
    for (int i = 0; i < 4; ++i)
      sweep_count[base + i] = static_cast<int>(dbg[i] >> 10);
  }

  // --- ascending sort of diag(A) + permuted transposed Q writeout ---
  if (rank == 0) {
    for (int i = tid; i < NP2; i += THREADS) {
      sKey[i] = i < N ? Aw[static_cast<size_t>(i) * N + i] : CUDART_INF_F;
      sIdx[i] = i;
    }
    __syncthreads();
    for (int k = 2; k <= NP2; k <<= 1) {
      for (int j = k >> 1; j > 0; j >>= 1) {
        for (int i = tid; i < NP2; i += THREADS) {
          const int ixj = i ^ j;
          if (ixj > i) {
            const float a_val = sKey[i];
            const float b_val = sKey[ixj];
            const bool ascending = (i & k) == 0;
            if (ascending ? (a_val > b_val) : (a_val < b_val)) {
              sKey[i] = b_val;
              sKey[ixj] = a_val;
              const int tmp = sIdx[i];
              sIdx[i] = sIdx[ixj];
              sIdx[ixj] = tmp;
            }
          }
        }
        __syncthreads();
      }
    }
    if (tid < N) {
      l_out[tid] = sKey[tid];
      const int src = sIdx[tid];
      sPerm[tid] = src;
      static_cast<int*>(cluster.map_shared_rank(sPerm, 1))[tid] = src;
      static_cast<int*>(cluster.map_shared_rank(sPerm, 2))[tid] = src;
    }
  }
  cluster.sync();

  // q_out[r][c] = QT[perm[c]][r]: 32x32 tile transpose through scratch
  {
    constexpr int NT = N / TD;   // 11
    for (int t = rank; t < NT * NT; t += 3) {
      const int w = (t / 3) % NWARP;
      if (w != warp) continue;
      const int rt = t / NT;
      const int ct = t % NT;
      #pragma unroll
      for (int it = 0; it < 8; ++it) {
        const int cl = it * 4 + (lane >> 3);
        const int seg = (lane & 7) * 4;
        *reinterpret_cast<float4*>(sc + cl * SST + seg) =
            *reinterpret_cast<const float4*>(
                QTw + static_cast<size_t>(sPerm[ct * TD + cl]) * N + rt * TD +
                seg);
      }
      __syncwarp();
      #pragma unroll
      for (int it = 0; it < 8; ++it) {
        const int r = it * 4 + (lane >> 3);
        const int c0 = (lane & 7) * 4;
        float4 v;
        v.x = sc[(c0 + 0) * SST + r];
        v.y = sc[(c0 + 1) * SST + r];
        v.z = sc[(c0 + 2) * SST + r];
        v.w = sc[(c0 + 3) * SST + r];
        *reinterpret_cast<float4*>(
            q_out + static_cast<size_t>(rt * TD + r) * N + ct * TD + c0) = v;
      }
      __syncwarp();
    }
  }
}

void launch_block_jacobi(
    const float* input,
    float* q,
    float* l,
    float* a_work,
    float* qt_work,
    int batch,
    int n,
    float tol2,
    int max_sweeps,
    int* sweep_count) {
#define LAUNCH_BLOCK_JACOBI(N, THREADS) \
  if (n == N) { \
    constexpr size_t kBytes = bj::Cfg<N>::BYTES; \
    auto* const kfn = block_jacobi_kernel<N, THREADS>; \
    static bool configured = [kfn] { \
      cudaFuncSetAttribute( \
          kfn, cudaFuncAttributeMaxDynamicSharedMemorySize, kBytes); \
      return true; \
    }(); \
    (void)configured; \
    cudaLaunchConfig_t config = {}; \
    config.gridDim = dim3(3 * batch); \
    config.blockDim = dim3(THREADS); \
    config.dynamicSmemBytes = kBytes; \
    cudaLaunchAttribute attrs[1]; \
    attrs[0].id = cudaLaunchAttributeClusterDimension; \
    attrs[0].val.clusterDim.x = 3; \
    attrs[0].val.clusterDim.y = 1; \
    attrs[0].val.clusterDim.z = 1; \
    config.attrs = attrs; \
    config.numAttrs = 1; \
    cudaLaunchKernelEx(&config, kfn, input, q, l, a_work, qt_work, tol2, \
                       max_sweeps, sweep_count); \
    return; \
  }

  LAUNCH_BLOCK_JACOBI(352, 512)
  LAUNCH_BLOCK_JACOBI(192, 512)

#undef LAUNCH_BLOCK_JACOBI
}

__device__ __forceinline__ float load_f32(const float* p) { return *p; }
__device__ __forceinline__ float load_f32(const __half* p) { return __half2float(*p); }
__device__ __forceinline__ void store_f32(float* p, float v) { *p = v; }
__device__ __forceinline__ void store_f32(__half* p, float v) { *p = __float2half(v); }

// ---------------------------------------------------------------------------
// Seated warp pivot solver (TB = 64, fp16 shared memory, one warp per pivot).
//
// The 64 pivot columns sit in "seats"; the round-robin tournament is realized
// by a FIXED reseating permutation SIGMA applied every inner round, so the
// rotation pairs are ALWAYS the adjacent seats (2m, 2m+1) — i.e. both halves
// of one __half2 word. Each lane owns one seat-pair: it pulls rows 2m and
// 2m+1 into registers, mixes them (row phase), applies all 32 in-word column
// rotations (coefficients shuffled across the warp), scatters the columns to
// their next-round seats (compile-time SIGMA => pure register moves), and
// stores to the reseated rows. V accumulates the same column operations with
// unpermuted rows. Per round: ~256 half2 SMEM ops + ALU vs ~1650 scalar ops
// for the naive kernel.
// ---------------------------------------------------------------------------

__device__ constexpr int SEAT64_S2P[64] = {
    0, 1, 2, 63, 3, 62, 4, 61, 5, 60, 6, 59, 7, 58, 8, 57, 9, 56, 10, 55, 11,
    54, 12, 53, 13, 52, 14, 51, 15, 50, 16, 49, 17, 48, 18, 47, 19, 46, 20, 45,
    21, 44, 22, 43, 23, 42, 24, 41, 25, 40, 26, 39, 27, 38, 28, 37, 29, 36, 30,
    35, 31, 34, 32, 33};
__device__ constexpr int SEAT64_SIGMA[64] = {
    0, 2, 4, 1, 6, 3, 8, 5, 10, 7, 12, 9, 14, 11, 16, 13, 18, 15, 20, 17, 22,
    19, 24, 21, 26, 23, 28, 25, 30, 27, 32, 29, 34, 31, 36, 33, 38, 35, 40, 37,
    42, 39, 44, 41, 46, 43, 48, 45, 50, 47, 52, 49, 54, 51, 56, 53, 58, 55, 60,
    57, 62, 59, 63, 61};

__device__ __forceinline__ __half2 rot_word_f32(__half2 w, float c, float s) {
  // (x, y) -> (c*x - s*y, s*x + c*y), computed in fp32 (fp16 c/s arithmetic
  // accumulates rotation non-orthogonality ~30x faster than storage rounding)
  const float x = __low2float(w);
  const float y = __high2float(w);
  return __floats2half2_rn(fmaf(c, x, -s * y), fmaf(s, x, c * y));
}

__device__ __forceinline__ void rowmix_f32(
    __half2& p, __half2& q, float c, float s) {
  const float px = __low2float(p), py = __high2float(p);
  const float qx = __low2float(q), qy = __high2float(q);
  p = __floats2half2_rn(fmaf(c, px, -s * qx), fmaf(c, py, -s * qy));
  q = __floats2half2_rn(fmaf(s, px, c * qx), fmaf(s, py, c * qy));
}

template <int NGLOB, int P, int WARPS, typename TIn>
__global__ __launch_bounds__(WARPS * 32, 1)
void seated_pivot_solve_kernel(
    const TIn* __restrict__ A,
    __half* __restrict__ v_out,
    const int* __restrict__ pairs,  // (batch, P/2, 2) block indices
    const int* __restrict__ active,  // (batch,) per-matrix convergence flag
    int max_sweeps,
    float tol2) {
  constexpr int TB = 64;
  constexpr int BLK = TB / 2;
  constexpr int ROUNDS = TB - 1;
  constexpr int WORDS = TB / 2;   // half2 words per row
  constexpr int LDW = WORDS + 1;  // padded row stride in words
  constexpr int P2 = P / 2;
  const int warp = threadIdx.x / 32;
  const int lane = threadIdx.x % 32;
  const int pair = blockIdx.x * WARPS + warp;
  const int batch = blockIdx.y;
  const unsigned mask = 0xffffffffu;
  if (active[batch] == 0) return;

  extern __shared__ unsigned char raw_smem[];
  __half2* smem = reinterpret_cast<__half2*>(raw_smem);
  __half2* As = smem + (size_t)warp * 2 * TB * LDW;
  __half2* Vs = As + TB * LDW;
  __shared__ float sort_d[WARPS][TB];
  __shared__ int sort_i[WARPS][TB];
  // fp32 diagonal tracked exactly: reading the diagonal from fp16 storage
  // injects ~5e-4*|d| angle errors that re-seed the off-diagonal at
  // ~5e-4*||A|| per visit (the convergence floor of a pure-fp16 solver).
  __shared__ float diag_f[WARPS][TB];

  const int bp = pairs[(batch * P2 + pair) * 2];
  const int bq = pairs[(batch * P2 + pair) * 2 + 1];
  const TIn* Ab = A + static_cast<size_t>(batch) * NGLOB * NGLOB;

  // Load pivot in seat order (seat s holds pivot-local index SEAT64_S2P[s]).
  float fro2_local = 0.0f;
  float off2_local = 0.0f;
  for (int si = 0; si < TB; ++si) {
    const int pi = SEAT64_S2P[si];
    const int gi = pi < BLK ? bp * BLK + pi : bq * BLK + (pi - BLK);
    const int sj0 = 2 * lane;
    const int pj0 = SEAT64_S2P[sj0];
    const int pj1 = SEAT64_S2P[sj0 + 1];
    const int gj0 = pj0 < BLK ? bp * BLK + pj0 : bq * BLK + (pj0 - BLK);
    const int gj1 = pj1 < BLK ? bp * BLK + pj1 : bq * BLK + (pj1 - BLK);
    const float v0 = 0.5f * (
        load_f32(Ab + (size_t)gi * NGLOB + gj0) +
        load_f32(Ab + (size_t)gj0 * NGLOB + gi));
    const float v1 = 0.5f * (
        load_f32(Ab + (size_t)gi * NGLOB + gj1) +
        load_f32(Ab + (size_t)gj1 * NGLOB + gi));
    As[si * LDW + lane] = __floats2half2_rn(v0, v1);
    if (si == sj0) diag_f[warp][si] = v0;
    if (si == sj0 + 1) diag_f[warp][si] = v1;
    // V rows are pivot-local (never reseated): V[i][s] = (i == S2P[s]).
    Vs[si * LDW + lane] = __floats2half2_rn(
        si == pj0 ? 1.0f : 0.0f, si == pj1 ? 1.0f : 0.0f);
    fro2_local += v0 * v0 + v1 * v1;
    if (si != sj0) off2_local += v0 * v0;
    if (si != sj0 + 1) off2_local += v1 * v1;
  }
  const float fro2 = warp_sum(fro2_local);
  const float off2_in = warp_sum(off2_local);
  bool done = fro2 == 0.0f || off2_in <= tol2 * fro2;
  __syncwarp();

  for (int sweep = 0; sweep < max_sweeps && !done; ++sweep) {
    for (int round = 0; round < ROUNDS; ++round) {
      // --- rotation for seat pair (2*lane, 2*lane+1), computed in fp32 ---
      const int rp = 2 * lane;
      const int rq = rp + 1;
      const __half2 wpp = As[rp * LDW + lane];  // (app, apq)
      const __half2 wqq = As[rq * LDW + lane];  // (aqp, aqq)
      const float app = diag_f[warp][rp];
      const float aqq = diag_f[warp][rq];
      const float apq = 0.5f * (__high2float(wpp) + __low2float(wqq));
      float c = 1.0f;
      float s = 0.0f;
      float t = 0.0f;
      if (fabsf(apq) > 0.0f) {
        const float theta = 0.5f * (aqq - app) / apq;
        t = copysignf(
            1.0f / (fabsf(theta) + sqrtf(1.0f + theta * theta)), theta);
        c = rsqrtf(1.0f + t * t);
        s = t * c;
      }
      const float dp_new = app - t * apq;
      const float dq_new = aqq + t * apq;


      // --- load rows rp, rq ---
      __half2 rowp[WORDS];
      __half2 rowq[WORDS];
      #pragma unroll
      for (int w = 0; w < WORDS; ++w) {
        rowp[w] = As[rp * LDW + w];
        rowq[w] = As[rq * LDW + w];
      }
      __syncwarp();

      // --- row phase in registers (fp32 math) ---
      #pragma unroll
      for (int w = 0; w < WORDS; ++w) {
        rowmix_f32(rowp[w], rowq[w], c, s);
      }

      // --- column phase: in-word rotations, coefficients from lane k ---
      #pragma unroll
      for (int k = 0; k < WORDS; ++k) {
        const float ck = __shfl_sync(mask, c, k);
        const float sk = __shfl_sync(mask, s, k);
        rowp[k] = rot_word_f32(rowp[k], ck, sk);
        rowq[k] = rot_word_f32(rowq[k], ck, sk);
      }

      // --- column reseat (compile-time scatter) + row reseat on store ---
      __half outp[TB];
      __half outq[TB];
      #pragma unroll
      for (int k = 0; k < WORDS; ++k) {
        outp[SEAT64_SIGMA[2 * k]] = __low2half(rowp[k]);
        outp[SEAT64_SIGMA[2 * k + 1]] = __high2half(rowp[k]);
        outq[SEAT64_SIGMA[2 * k]] = __low2half(rowq[k]);
        outq[SEAT64_SIGMA[2 * k + 1]] = __high2half(rowq[k]);
      }
      const int dp = SEAT64_SIGMA[rp];
      const int dq = SEAT64_SIGMA[rq];
      #pragma unroll
      for (int w = 0; w < WORDS; ++w) {
        As[dp * LDW + w] = __halves2half2(outp[2 * w], outp[2 * w + 1]);
        As[dq * LDW + w] = __halves2half2(outq[2 * w], outq[2 * w + 1]);
      }
      diag_f[warp][dp] = dp_new;
      diag_f[warp][dq] = dq_new;

      // --- V: same column ops, rows stay in place ---
      #pragma unroll
      for (int w = 0; w < WORDS; ++w) {
        rowp[w] = Vs[rp * LDW + w];
        rowq[w] = Vs[rq * LDW + w];
      }
      #pragma unroll
      for (int k = 0; k < WORDS; ++k) {
        const float ck = __shfl_sync(mask, c, k);
        const float sk = __shfl_sync(mask, s, k);
        rowp[k] = rot_word_f32(rowp[k], ck, sk);
        rowq[k] = rot_word_f32(rowq[k], ck, sk);
      }
      #pragma unroll
      for (int k = 0; k < WORDS; ++k) {
        outp[SEAT64_SIGMA[2 * k]] = __low2half(rowp[k]);
        outp[SEAT64_SIGMA[2 * k + 1]] = __high2half(rowp[k]);
        outq[SEAT64_SIGMA[2 * k]] = __low2half(rowq[k]);
        outq[SEAT64_SIGMA[2 * k + 1]] = __high2half(rowq[k]);
      }
      #pragma unroll
      for (int w = 0; w < WORDS; ++w) {
        Vs[rp * LDW + w] = __halves2half2(outp[2 * w], outp[2 * w + 1]);
        Vs[rq * LDW + w] = __halves2half2(outq[2 * w], outq[2 * w + 1]);
      }
      __syncwarp();
    }

    if (sweep + 1 < max_sweeps) {
      float o2 = 0.0f;
      for (int i = 0; i < TB; ++i) {
        const __half2 w = As[i * LDW + lane];
        const float x = __low2float(w);
        const float y = __high2float(w);
        if (i != 2 * lane) o2 += x * x;
        if (i != 2 * lane + 1) o2 += y * y;
      }
      done = warp_sum(o2) <= tol2 * fro2;
      __syncwarp();
    }
  }

  // Sort diagonal (seat space) ascending; permute V columns to match.
  {
    sort_d[warp][2 * lane] = diag_f[warp][2 * lane];
    sort_d[warp][2 * lane + 1] = diag_f[warp][2 * lane + 1];
    sort_i[warp][2 * lane] = 2 * lane;
    sort_i[warp][2 * lane + 1] = 2 * lane + 1;
  }
  __syncwarp();
  #pragma unroll
  for (int k = 2; k <= 64; k <<= 1) {
    #pragma unroll
    for (int j = k >> 1; j > 0; j >>= 1) {
      #pragma unroll
      for (int h = 0; h < 2; ++h) {
        const int i = lane + h * 32;
        const int ixj = i ^ j;
        if (ixj > i) {
          const bool up = (i & k) == 0;
          const float di = sort_d[warp][i];
          const float dj = sort_d[warp][ixj];
          if ((di > dj) == up) {
            sort_d[warp][i] = dj;
            sort_d[warp][ixj] = di;
            const int ti = sort_i[warp][i];
            sort_i[warp][i] = sort_i[warp][ixj];
            sort_i[warp][ixj] = ti;
          }
        }
      }
      __syncwarp();
    }
  }

  // Emit V: rows = pivot-local basis (unpermuted), cols sorted by eigenvalue.
  __half* Vb = v_out + (static_cast<size_t>(batch) * P2 + pair) * TB * TB;
  for (int i = 0; i < TB; ++i) {
    #pragma unroll
    for (int h = 0; h < 2; ++h) {
      const int j = lane + h * 32;
      const int sj = sort_i[warp][j];
      const __half2 w = Vs[i * LDW + (sj >> 1)];
      Vb[i * TB + j] = (sj & 1) ? __high2half(w) : __low2half(w);
    }
  }
}

template <int NGLOB, int P, int WARPS, typename TIn>
void run_seated_pivot_solve(
    const void* a,
    void* v,
    const int* pairs,
    const int* active,
    int batch,
    int max_sweeps,
    float tol2) {
  constexpr int TB = 64;
  constexpr int LDW = TB / 2 + 1;
  constexpr size_t SMEM = (size_t)WARPS * 2 * TB * LDW * sizeof(__half2);
  constexpr int P2 = P / 2;
  static_assert(P2 % WARPS == 0);
  auto* kernel = seated_pivot_solve_kernel<NGLOB, P, WARPS, TIn>;
  static bool configured = false;
  if (!configured) {
    cudaFuncSetAttribute(
        kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM);
    configured = true;
  }
  dim3 grid(P2 / WARPS, batch);
  kernel<<<grid, WARPS * 32, SMEM>>>(
      static_cast<const TIn*>(a), static_cast<__half*>(v), pairs, active,
      max_sweeps, tol2);
}

void launch_seated_pivot_solve(
    const void* a,
    bool a_is_half,
    void* v,
    const int* pairs,
    const int* active,
    int batch,
    int n,
    int max_sweeps,
    float tol2) {
  if (n == 512) {
    if (a_is_half) {
      run_seated_pivot_solve<512, 16, 4, __half>(
          a, v, pairs, active, batch, max_sweeps, tol2);
    } else {
      run_seated_pivot_solve<512, 16, 4, float>(
          a, v, pairs, active, batch, max_sweeps, tol2);
    }
    return;
  }
}

void launch_jacobi_smem(
    const float* input,
    float* q,
    float* l,
    int batch,
    int n,
    float tol2,
    int max_sweeps,
    int* sweep_count) {
#define LAUNCH_JACOBI_SMEM(N, THREADS) \
  if (n == N) { \
    constexpr int kPairs = N / 2; \
    constexpr int kNoff = kPairs * (kPairs - 1) / 2; \
    constexpr int kNp2 = next_pow2(N); \
    constexpr size_t kBytes = \
        (2 * N * (N + 1) + 4 * kPairs + 2 * kNp2) * sizeof(float) + \
        2 * kNoff + 256; \
    static bool configured = [] { \
      cudaFuncSetAttribute( \
          jacobi_smem_kernel<N, THREADS>, \
          cudaFuncAttributeMaxDynamicSharedMemorySize, kBytes); \
      return true; \
    }(); \
    (void)configured; \
    jacobi_smem_kernel<N, THREADS> \
        <<<batch, THREADS, kBytes>>>(input, q, l, tol2, max_sweeps, sweep_count); \
    return; \
  }

  LAUNCH_JACOBI_SMEM(16, 64)
  LAUNCH_JACOBI_SMEM(32, 128)
  LAUNCH_JACOBI_SMEM(88, 256)

#undef LAUNCH_JACOBI_SMEM
}

void launch_steqr_tri(
    const double* d,
    const double* e,
    float* q,
    double* l,
    int P,
    int n,
    int use_f32,
    int do_sort) {
#define LAUNCH_STEQR(N, WARPS, CT, SORT) \
  if (n == N) { \
    constexpr size_t kBytes = \
        WARPS * (2 * N * sizeof(CT) + N * sizeof(int)); \
    const int grid = cdiv(P, WARPS); \
    steqr_tri_kernel<N, WARPS, CT, SORT><<<grid, WARPS * 32, kBytes>>>(d, e, q, l, P); \
    return; \
  }
  if (use_f32 && do_sort) {
    LAUNCH_STEQR(32, 8, float, true)
    LAUNCH_STEQR(22, 8, float, true)
    LAUNCH_STEQR(16, 8, float, true)
  } else if (use_f32) {
    LAUNCH_STEQR(32, 8, float, false)
    LAUNCH_STEQR(22, 8, float, false)
    LAUNCH_STEQR(16, 8, float, false)
  } else if (do_sort) {
    LAUNCH_STEQR(32, 8, double, true)
    LAUNCH_STEQR(22, 8, double, true)
    LAUNCH_STEQR(16, 8, double, true)
  } else {
    LAUNCH_STEQR(32, 8, double, false)
    LAUNCH_STEQR(22, 8, double, false)
    LAUNCH_STEQR(16, 8, double, false)
  }
#undef LAUNCH_STEQR
}

// ===========================================================================
// Batched tridiagonal divide&conquer MERGE kernel (dc:: namespace).
// One CTA per merge subproblem. Torch does the O(P*m) deflation bookkeeping
// (sort, clusters, Householder, compaction); this kernel does the O(m^2)
// secular solve (fp64 scalar), Gu-Eisenstat zhat (log-space), eigenvector V
// build with normalization + deflated unit columns + Householder G, writing
// Vchild (child-pole-row order, compacted-root-col order) and Lam.
// See dev/dc_proto.py (_merge) for the validated reference math.
// ===========================================================================
namespace dc {

// IEEE round-to-nearest divide (unaffected by -use_fast_math's rcp.approx);
// the fp32 secular iteration oscillates at the approx-divide noise floor
// otherwise (coordinator lever 2).
__device__ __forceinline__ float ieee_div(float a, float b) { return __fdiv_rn(a, b); }
__device__ __forceinline__ double ieee_div(double a, double b) { return a / b; }
// IEEE round-to-nearest reciprocal (explicit intrinsics defeat -use_fast_math's
// imprecise-division substitution; fp64 has no rcp.approx HW so __drcp_rn is exact).
__device__ __forceinline__ float ieee_rcp(float b) { return __frcp_rn(b); }
__device__ __forceinline__ double ieee_rcp(double b) { return __drcp_rn(b); }
// Mixed-precision reciprocal SEED for the secular polish: fp32 rcp.approx (one
// cheap MUFU) + ONE fp64 Newton step -> ~2^-48 accurate 1/b. Used only where the
// value is a SEED that a subsequent fp64 Newton correction (the zz - r*den
// residual, formed in full fp64) refines to full 2^-52 accuracy regardless. This
// replaces __drcp_rn (fp32 seed + ~2 internal fp64 Newton) with one Newton, so
// per pole we drop ~2 fp64 FMAs off the hottest loop on this 2:1 fp64 box.
__device__ __forceinline__ float refined_rcp(float b) { return __frcp_rn(b); }
__device__ __forceinline__ double refined_rcp(double b) {
  double x = (double)__frcp_rn((float)b);
  return fma(fma(-x, b, 1.0), x, x);
}

// One middle-way (Li) rational-interpolation secular step for root i, generic
// in precision T. Reads poles dc/zc (T), returns new eta; updates lo/hi bracket
// and outputs the secular value g at the *current* eta (for |g| convergence).
template <typename T, bool ILP = false>
__device__ __forceinline__ T mw_step(
    const T* __restrict__ dc, const T* __restrict__ zc, int k, int iU,
    T di, T rho, T gL, T gU, T et, T& lo, T& hi, bool pos, T& g_out, T& erretm_out) {
  T psi, psip, phi, phip;
  if constexpr (ILP) {
    // ILP path (cluster/low-wave kernel, has register room): the plain loop
    // carries a per-j branch (j<iU?psi:phi) and ONE serial fp64 add-chain per
    // side -> add-latency-bound, stalls at No-Eligible 32% (ncu). iU splits the
    // poles CONTIGUOUSLY (psi=[0,iU), phi=[iU,k)), so hoist the branch into two
    // straight-line loops, each unrolled x2 with independent accumulators: the
    // add-chain latency halves and the two chains + independent fp64 reciprocals
    // fill the issue stalls. Reassociating the fp64 sum is a <1ulp perturbation,
    // far inside the 8e-16*erretm convergence tol. NOT used by the single-CTA
    // (high-wave n512) kernel: there the 4 extra fp64 accs spill at its 48-reg
    // launch_bounds and occupancy (not latency) is the wall.
    T psi0 = 0, psi1 = 0, psip0 = 0, psip1 = 0;
    int j = 0;
    for (; j + 1 < iU; j += 2) {
      T den0 = (dc[j] - di) - et, den1 = (dc[j + 1] - di) - et;
      T inv0 = ieee_rcp(den0), inv1 = ieee_rcp(den1);
      T r0 = (zc[j] * zc[j]) * inv0, r1 = (zc[j + 1] * zc[j + 1]) * inv1;
      psi0 += r0; psi1 += r1; psip0 += r0 * inv0; psip1 += r1 * inv1;
    }
    for (; j < iU; ++j) {
      T den = (dc[j] - di) - et; T inv = ieee_rcp(den);
      T r = (zc[j] * zc[j]) * inv; psi0 += r; psip0 += r * inv;
    }
    T phi0 = 0, phi1 = 0, phip0 = 0, phip1 = 0;
    j = iU;
    for (; j + 1 < k; j += 2) {
      T den0 = (dc[j] - di) - et, den1 = (dc[j + 1] - di) - et;
      T inv0 = ieee_rcp(den0), inv1 = ieee_rcp(den1);
      T r0 = (zc[j] * zc[j]) * inv0, r1 = (zc[j + 1] * zc[j + 1]) * inv1;
      phi0 += r0; phi1 += r1; phip0 += r0 * inv0; phip1 += r1 * inv1;
    }
    for (; j < k; ++j) {
      T den = (dc[j] - di) - et; T inv = ieee_rcp(den);
      T r = (zc[j] * zc[j]) * inv; phi0 += r; phip0 += r * inv;
    }
    psi = (psi0 + psi1) * rho; psip = (psip0 + psip1) * rho;
    phi = (phi0 + phi1) * rho; phip = (phip0 + phip1) * rho;
  } else {
    // Original serial loop (single-CTA n512 kernel: register-tight, occupancy-bound).
    T psi_ = 0, psip_ = 0, phi_ = 0, phip_ = 0;
    for (int j = 0; j < k; ++j) {
      T den = (dc[j] - di) - et;
      T zz = zc[j] * zc[j];
      T inv = ieee_rcp(den);
      T r = zz * inv;
      T rp = r * inv;
      if (j < iU) { psi_ += r; psip_ += rp; }
      else        { phi_ += r; phip_ += rp; }
    }
    psi = psi_ * rho; psip = psip_ * rho; phi = phi_ * rho; phip = phip_ * rho;
  }
  T g = (T)1 + psi + phi;
  g_out = g;
  // roundoff error estimate for g: sum of the magnitudes that (nearly) cancel.
  erretm_out = fabs(psi) + fabs(phi) + (T)1;
  bool up = pos ? (g < (T)0) : (g > (T)0);
  if (up) lo = et; else hi = et;
  T dL = gL - et, dU = gU - et;
  T A = psip * dL * dL, Bc = phip * dU * dU;
  T K = (T)1 + (psi - psip * dL) + (phi - phip * dU);
  T a1 = -(K * (gL + gU) + A + Bc);
  T a0 = K * gL * gU + A * gU + Bc * gL;
  T disc = a1 * a1 - (T)4 * K * a0;
  if (K != (T)0 && disc >= (T)0) {
    T sq = sqrt(disc);
    T q = (T)(-0.5) * (a1 + ((a1 >= (T)0) ? sq : -sq));
    T r1 = q / K, r2 = a0 / q;
    bool i1 = (r1 > lo && r1 < hi), i2 = (r2 > lo && r2 < hi);
    if (i1 && !i2) return r1;
    if (i2 && !i1) return r2;
    if (i1 && i2)  return (fabs(r1 - et) < fabs(r2 - et)) ? r1 : r2;
  }
  return (T)0.5 * (lo + hi);
}

// Neumaier compensated add: sum += x with a running error term e.
__device__ __forceinline__ void neu(float& s, float& e, float x) {
  float t = s + x;
  e += (fabsf(s) >= fabsf(x)) ? ((s - t) + x) : ((x - t) + s);
  s = t;
}

// Compensated fp32 middle-way step. The secular sums psi/phi cancel by
// definition at the root, so plain fp32 evaluation stalls at a false floor ~
// eps32*sum|terms| (huge on wide spectra) -> wrong roots. Neumaier-compensated
// accumulation (~2^-45 effective) resolves f to gap-relative precision using
// only fp32 units, so fp32 Newton converges like fp64 (coordinator diagnosis).
template <bool ILP = false>
__device__ __forceinline__ float mw_step_comp(
    const float* __restrict__ dc, const float* __restrict__ zc, int k, int iU,
    float di, float rho, float gL, float gU, float et, float& lo, float& hi,
    bool pos, float& g_out, float& erretm_out) {
  float psi, psie, psip, phi, phie, phip;
  if constexpr (ILP) {
    // ILP path (cluster/low-wave kernel): branchless two-loop split (iU splits
    // the poles contiguously) with TWO Neumaier accumulators per side -> hoists
    // the per-j branch and breaks the compensated-add serial chain into two
    // independent chains, filling No-Eligible issue stalls. Combining the two
    // compensated sums preserves the ~2^-45 effective precision the cancelling
    // secular sum needs. one correctly-rounded fp32 reciprocal reused for
    // r=z^2/den and rp=r/den (pole-dominated derivative, plain fp32 ample).
    float psiA = 0, psieA = 0, psiB = 0, psieB = 0, psip0 = 0, psip1 = 0;
    int j = 0;
    for (; j + 1 < iU; j += 2) {
      float d0 = (dc[j] - di) - et, d1 = (dc[j + 1] - di) - et;
      float i0 = __frcp_rn(d0), i1 = __frcp_rn(d1);
      float r0 = (zc[j] * zc[j]) * i0, r1 = (zc[j + 1] * zc[j + 1]) * i1;
      neu(psiA, psieA, r0); neu(psiB, psieB, r1);
      psip0 += r0 * i0; psip1 += r1 * i1;
    }
    for (; j < iU; ++j) {
      float d = (dc[j] - di) - et; float i0 = __frcp_rn(d);
      float r = (zc[j] * zc[j]) * i0; neu(psiA, psieA, r); psip0 += r * i0;
    }
    neu(psiA, psieA, psiB); psieA += psieB;
    psi = psiA; psie = psieA; psip = psip0 + psip1;
    float phiA = 0, phieA = 0, phiB = 0, phieB = 0, phip0 = 0, phip1 = 0;
    j = iU;
    for (; j + 1 < k; j += 2) {
      float d0 = (dc[j] - di) - et, d1 = (dc[j + 1] - di) - et;
      float i0 = __frcp_rn(d0), i1 = __frcp_rn(d1);
      float r0 = (zc[j] * zc[j]) * i0, r1 = (zc[j + 1] * zc[j + 1]) * i1;
      neu(phiA, phieA, r0); neu(phiB, phieB, r1);
      phip0 += r0 * i0; phip1 += r1 * i1;
    }
    for (; j < k; ++j) {
      float d = (dc[j] - di) - et; float i0 = __frcp_rn(d);
      float r = (zc[j] * zc[j]) * i0; neu(phiA, phieA, r); phip0 += r * i0;
    }
    neu(phiA, phieA, phiB); phieA += phieB;
    phi = phiA; phie = phieA; phip = phip0 + phip1;
  } else {
    // Original single serial loop (single-CTA n512 kernel, register-tight).
    float ps = 0, pse = 0, pp = 0, ph = 0, phe = 0, phpp = 0;
    for (int j = 0; j < k; ++j) {
      float den = (dc[j] - di) - et;
      float zz = zc[j] * zc[j];
      float inv = __frcp_rn(den);
      float r = zz * inv;
      float rp = r * inv;
      if (j < iU) { neu(ps, pse, r); pp += rp; }
      else        { neu(ph, phe, r); phpp += rp; }
    }
    psi = ps; psie = pse; psip = pp; phi = ph; phie = phe; phip = phpp;
  }
  psi = (psi + psie) * rho; psip = psip * rho;
  phi = (phi + phie) * rho; phip = phip * rho;
  float g = 1.0f + psi + phi; g_out = g;
  erretm_out = fabsf(psi) + fabsf(phi) + 1.0f;
  bool up = pos ? (g < 0.0f) : (g > 0.0f);
  if (up) lo = et; else hi = et;
  float dL = gL - et, dU = gU - et;
  float A = psip * dL * dL, Bc = phip * dU * dU;
  float K = 1.0f + (psi - psip * dL) + (phi - phip * dU);
  float a1 = -(K * (gL + gU) + A + Bc);
  float a0 = K * gL * gU + A * gU + Bc * gL;
  float disc = a1 * a1 - 4.0f * K * a0;
  if (K != 0.0f && disc >= 0.0f) {
    float sq = sqrtf(disc);
    float q = -0.5f * (a1 + ((a1 >= 0.0f) ? sq : -sq));
    float r1 = __fdiv_rn(q, K), r2 = __fdiv_rn(a0, q);
    bool i1 = (r1 > lo && r1 < hi), i2 = (r2 > lo && r2 < hi);
    if (i1 && !i2) return r1;
    if (i2 && !i1) return r2;
    if (i1 && i2)  return (fabsf(r1 - et) < fabsf(r2 - et)) ? r1 : r2;
  }
  return 0.5f * (lo + hi);
}

// ============ two-float (df32) arithmetic on the FP32 pipe =================
// On this box FP32:FP64 = 64:1, so the fp64 secular polish is the whole wall.
// The denominator (poles - eta) is the only precision-critical quantity (z is
// fp32-adequate per the validated precision plan); deflation guarantees min pole
// gap ~1e-8*scale while df32 (~2^-46) resolves ~1e-14*scale => 6 digits of
// headroom, and 2^-46 eta storage gives eigen residual ~0.1 (gate 200). We do
// the O(k) inner reduction in df32 and the O(1) quadratic solve in native fp64.
struct df2 { float h, l; };
__device__ __forceinline__ df2 two_sum(float a, float b) {
  float s = a + b, bb = s - a;
  return {s, (a - (s - bb)) + (b - bb)};
}
__device__ __forceinline__ df2 quick_two_sum(float a, float b) {
  float s = a + b;
  return {s, b - (s - a)};
}
__device__ __forceinline__ df2 df_add(df2 a, df2 b) {
  df2 s = two_sum(a.h, b.h);
  df2 t = two_sum(a.l, b.l);
  s.l += t.h;
  s = quick_two_sum(s.h, s.l);
  s.l += t.l;
  return quick_two_sum(s.h, s.l);
}
__device__ __forceinline__ df2 df_sub(df2 a, df2 b) {
  return df_add(a, df2{-b.h, -b.l});
}
__device__ __forceinline__ df2 df_mul(df2 a, df2 b) {
  float p = a.h * b.h;
  float e = fmaf(a.h, b.h, -p);
  e = fmaf(a.h, b.l, e);
  e = fmaf(a.l, b.h, e);
  return quick_two_sum(p, e);
}
// df2 * fp32 scalar (b exact in fp32)
__device__ __forceinline__ df2 df_mul_f(df2 a, float b) {
  float p = a.h * b;
  float e = fmaf(a.h, b, -p);
  e = fmaf(a.l, b, e);
  return quick_two_sum(p, e);
}
// reciprocal via one df Newton step from a fp32 seed (~2^-46)
__device__ __forceinline__ df2 df_recip(df2 b) {
  float x = __frcp_rn(b.h);
  df2 bx = df_mul_f(b, x);              // b*x ~ 1
  df2 r  = df_sub(df2{2.0f, 0.0f}, bx); // 2 - b*x
  return df_mul_f(r, x);               // x*(2 - b*x)
}

// df32 middle-way secular step for root i. Poles as df32 (dch/dcl), z as fp32.
// et/lo/hi/gL/gU stay fp64 (bracket bookkeeping + O(1) solve). Returns new et.
__device__ __forceinline__ double mw_step_df(
    const float* __restrict__ dch, const float* __restrict__ dcl,
    const float* __restrict__ zc, int k, int iU,
    float dih, float dil, double rho, double gL, double gU, double et,
    double& lo, double& hi, bool pos, double& g_out) {
  // split et into df32
  float eth = (float)et;
  float etl = (float)(et - (double)eth);
  df2 di = df2{dih, dil};
  df2 etd = df2{eth, etl};
  df2 psi{0, 0}, psip{0, 0}, phi{0, 0}, phip{0, 0};
  for (int j = 0; j < k; ++j) {
    df2 den = df_sub(df_sub(df2{dch[j], dcl[j]}, di), etd);  // (dc[j]-di)-et
    df2 inv = df_recip(den);
    float zz = zc[j] * zc[j];
    df2 r  = df_mul_f(inv, zz);   // z^2/den
    df2 rp = df_mul(r, inv);      // z^2/den^2
    if (j < iU) { psi = df_add(psi, r); psip = df_add(psip, rp); }
    else        { phi = df_add(phi, r); phip = df_add(phip, rp); }
  }
  double psid = ((double)psi.h + (double)psi.l) * rho;
  double psipd = ((double)psip.h + (double)psip.l) * rho;
  double phid = ((double)phi.h + (double)phi.l) * rho;
  double phipd = ((double)phip.h + (double)phip.l) * rho;
  double g = 1.0 + psid + phid; g_out = g;
  bool up = pos ? (g < 0.0) : (g > 0.0);
  if (up) lo = et; else hi = et;
  double dL = gL - et, dU = gU - et;
  double A = psipd * dL * dL, Bc = phipd * dU * dU;
  double K = 1.0 + (psid - psipd * dL) + (phid - phipd * dU);
  double a1 = -(K * (gL + gU) + A + Bc);
  double a0 = K * gL * gU + A * gU + Bc * gL;
  double disc = a1 * a1 - 4.0 * K * a0;
  if (K != 0.0 && disc >= 0.0) {
    double sq = sqrt(disc);
    double q = -0.5 * (a1 + ((a1 >= 0.0) ? sq : -sq));
    double r1 = q / K, r2 = a0 / q;
    bool i1 = (r1 > lo && r1 < hi), i2 = (r2 > lo && r2 < hi);
    if (i1 && !i2) return r1;
    if (i2 && !i1) return r2;
    if (i1 && i2)  return (fabs(r1 - et) < fabs(r2 - et)) ? r1 : r2;
  }
  return 0.5 * (lo + hi);
}

// ===========================================================================
// dc_prep: fused per-level D&C bookkeeping (one CTA per merge subproblem).
// Replaces the ~130-launch torch glue chain in the _dc_merge python (cat, sort,
// gather, gap-cluster, segmented reductions, Householder deflation, active
// compaction, invorder/perm/segstart build) with ONE kernel. Emits exactly the
// merge_build_kernel input contract. Per-column arithmetic mirrors the python
// (bit-identical for distinct-eigenvalue inputs; gauge-equivalent on ties).
//   d_sec[i] = i<h ? Dl[p,i] : Drr[p,i-h]         (fp64)
//   z[i]     = i<h ? zL[p,i] : zR[p,i-h]          (fp32 widened to fp64)
// Sorts d_sec ascending (bitonic, carrying child index) -> ds/perm; z_s=z[perm].
// ===========================================================================
__device__ __forceinline__ void dcp_scan_incl(int* a, int* tmp, int m, int tid, int T) {
  for (int off = 1; off < m; off <<= 1) {
    for (int i = tid; i < m; i += T) tmp[i] = a[i] + ((i >= off) ? a[i - off] : 0);
    __syncthreads();
    for (int i = tid; i < m; i += T) a[i] = tmp[i];
    __syncthreads();
  }
}

// PAD=false: native power-of-2 merge size m (the 12/13 dc:: shapes n512/n1024/
// n2048, m in {64,128,256,512,1024,2048}). Byte-identical to the original
// single-`m` kernel (S folds to m, the pad branch compiles out).
// PAD=true: NON-power-of-2 m (native-352 D&C, m in {44,88,176,352}). The bitonic
// sort REQUIRES a power-of-2 length (i^j indexes up to next_pow2(m)-1, else OOB
// SMEM read + wrong order). Fix WITHOUT touching the problem size: size the SMEM
// sort arrays to mp = next_pow2(m), fill the pad slots [m,mp) with +inf keys so
// they sort strictly above every real pole (prescaled eigenvalues <= n << 1e300),
// then run all later bookkeeping over the real [0,m) as before. Only ds/pm
// are read at padded indices; everything else stays within [0,m).
template <bool PAD>
__global__ __launch_bounds__(256, 1) void dc_prep_kernel(
    const double* __restrict__ Dl,   // (P,h)
    const double* __restrict__ Drr,  // (P,h)
    const float*  __restrict__ zL,   // (P,h)  Ql[:,h-1,:]
    const float*  __restrict__ zR,   // (P,h)  Qrr[:,0,:]
    const double* __restrict__ rho_in, // (P,)
    double* __restrict__ dc_out,     // (P,m) compacted poles (jittered)
    double* __restrict__ zc_out,     // (P,m) compacted z (post-householder)
    int*    __restrict__ k_out,      // (P,)
    double* __restrict__ rho_out,    // (P,) de-zeroed rho
    int*    __restrict__ invorder,   // (P,m) sorted pos -> compacted idx
    int*    __restrict__ perm_out,   // (P,m) sorted pos -> child idx
    double* __restrict__ hv_out,     // (P,m) sorted-order householder v
    double* __restrict__ hbeta_out,  // (P,m) sorted-order per-pole beta
    int*    __restrict__ segstart,   // (P,m) cluster start (sorted idx), pad m
    int*    __restrict__ nseg_out,   // (P,)
    int h, int m, int mp, double zk_mul, double gk_mul) {
  const int p = blockIdx.x;
  const int tid = threadIdx.x;
  const int T = blockDim.x;
  const double EPS = 2.220446049250313e-16;
  const double GAP_K = 4.5e7;
  // SMEM stride: pow2-padded length for the bitonic sort when PAD, else m.
  const int S = PAD ? mp : m;

  extern __shared__ double sdp[];
  double* ds = sdp;            // S
  double* zs = ds + S;         // S
  double* hv = zs + S;         // S
  double* a0 = hv + S;         // S  (cluster accumulator)
  double* a1 = a0 + S;         // S  (cluster accumulator)
  int* pm    = (int*)(a1 + S); // S  perm
  int* cid   = pm + S;         // S
  int* nc    = cid + S;        // S  new_clu 0/1 (later reused as active)
  int* ss    = nc + S;         // S  segstart-by-cluster
  int* tmp   = ss + S;         // S  scan temp
  int* cpos  = tmp + S;        // S  compacted position (invorder)
  __shared__ double scale_sh, znorm_sh, rho_sh;
  __shared__ int nseg_sh, k_sh;
  __shared__ double red[48];

  const size_t hb = (size_t)p * h;
  for (int i = tid; i < S; i += T) {
    if (!PAD || i < m) {
      ds[i] = (i < h) ? Dl[hb + i] : Drr[hb + (i - h)];
      zs[i] = (i < h) ? (double)zL[hb + i] : (double)zR[hb + (i - h)];
    } else {
      ds[i] = 1.0e300;   // +inf sort key -> pad slots land in [m,mp)
    }
    pm[i] = i;
  }
  if (tid == 0) rho_sh = rho_in[p];
  __syncthreads();

  // ---- bitonic sort ds ascending carrying pm (over the pow2 length S) ----
  for (int kk = 2; kk <= S; kk <<= 1) {
    for (int j = kk >> 1; j > 0; j >>= 1) {
      for (int i = tid; i < S; i += T) {
        int ixj = i ^ j;
        if (ixj > i) {
          bool up = ((i & kk) == 0);
          if ((ds[i] > ds[ixj]) == up) {
            double td = ds[i]; ds[i] = ds[ixj]; ds[ixj] = td;
            int tp = pm[i]; pm[i] = pm[ixj]; pm[ixj] = tp;
          }
        }
      }
      __syncthreads();
    }
  }
  // z_s = z[perm]  (gather via a1 scratch then copy back)
  for (int i = tid; i < m; i += T) a1[i] = zs[pm[i]];
  __syncthreads();
  for (int i = tid; i < m; i += T) zs[i] = a1[i];
  __syncthreads();

  // ---- scale = max|ds|, znorm = sqrt(sum zs^2) ----
  double smax = 0.0, ssum = 0.0;
  for (int i = tid; i < m; i += T) { double a = fabs(ds[i]); if (a > smax) smax = a; ssum += zs[i]*zs[i]; }
  for (int off = 16; off > 0; off >>= 1) { smax = fmax(smax, __shfl_down_sync(FULL_MASK, smax, off)); ssum += __shfl_down_sync(FULL_MASK, ssum, off); }
  if ((tid & 31) == 0) { red[tid >> 5] = smax; red[(tid >> 5) + 24] = ssum; }
  __syncthreads();
  if (tid == 0) {
    double mx = 0.0, sm = 0.0; int nw = (T + 31) >> 5;
    for (int w = 0; w < nw; ++w) { mx = fmax(mx, red[w]); sm += red[w + 24]; }
    scale_sh = fmax(mx, 1e-30);
    znorm_sh = fmax(sqrt(sm), 1e-300);
  }
  __syncthreads();
  const double scale = scale_sh;
  const double znorm = znorm_sh;
  const double gtol = GAP_K * EPS * scale * gk_mul;

  // ---- new_clu + cluster id (scan) ----
  for (int i = tid; i < m; i += T) {
    int flag = (i == 0) ? 1 : ((ds[i] - ds[i-1] > gtol) ? 1 : 0);
    nc[i] = flag;
    cid[i] = flag;
  }
  __syncthreads();
  dcp_scan_incl(cid, tmp, m, tid, T);
  if (tid == 0) nseg_sh = cid[m-1];
  __syncthreads();
  const int nseg = nseg_sh;
  for (int i = tid; i < m; i += T) cid[i] -= 1;
  __syncthreads();

  // ---- segstart (by cluster) + csize/sum-z^2 (atomics) ----
  for (int c = tid; c < m; c += T) { ss[c] = m; a0[c] = 0.0; a1[c] = 0.0; }
  __syncthreads();
  for (int i = tid; i < m; i += T) {
    if (nc[i]) ss[cid[i]] = i;
    atomicAdd(&a0[cid[i]], 1.0);
    atomicAdd(&a1[cid[i]], zs[i]*zs[i]);
  }
  __syncthreads();

  // ---- householder v within multi-clusters ----
  for (int i = tid; i < m; i += T) {
    int c = cid[i];
    if (a0[c] > 1.5) {
      if (nc[i]) {
        double segnorm = sqrt(fmax(a1[c], 0.0));
        double zrep = zs[ss[c]];
        double rsgn = (zrep >= 0.0) ? 1.0 : -1.0;
        hv[i] = zs[i] + rsgn * segnorm;
      } else {
        hv[i] = zs[i];
      }
    } else {
      hv[i] = 0.0;
    }
  }
  __syncthreads();
  // ---- segvv = sum hv^2, vz = sum hv*z ----
  for (int c = tid; c < m; c += T) { a0[c] = 0.0; a1[c] = 0.0; }
  __syncthreads();
  for (int i = tid; i < m; i += T) {
    atomicAdd(&a0[cid[i]], hv[i]*hv[i]);
    atomicAdd(&a1[cid[i]], hv[i]*zs[i]);
  }
  __syncthreads();
  const size_t mb = (size_t)p * m;
  for (int i = tid; i < m; i += T) {
    int c = cid[i];
    double segvv = a0[c];
    double hbeta_c = (segvv > 0.0) ? (2.0 / segvv) : 0.0;
    zs[i] = zs[i] - hv[i] * hbeta_c * a1[c];
    hbeta_out[mb + i] = hbeta_c;
    hv_out[mb + i] = hv[i];
  }
  __syncthreads();

  // ---- active mask + stable compaction (active first, sorted order) ----
  const double rho = rho_sh;
  const bool decoupled = fabs(rho) < 1e-300;
  const double ztol = 8.0 * m * EPS * znorm * zk_mul;
  for (int i = tid; i < m; i += T) {
    int act = (!decoupled && fabs(zs[i]) > ztol) ? 1 : 0;
    nc[i] = act;
    cpos[i] = act;
  }
  __syncthreads();
  dcp_scan_incl(cpos, tmp, m, tid, T);
  if (tid == 0) k_sh = cpos[m-1];
  __syncthreads();
  const int k = k_sh;
  for (int i = tid; i < m; i += T) {
    int inc_act = cpos[i];
    int excl_act = inc_act - nc[i];
    cpos[i] = nc[i] ? excl_act : (k + (i - excl_act));
  }
  __syncthreads();
  const double jit = 8.0 * EPS * scale;
  for (int i = tid; i < m; i += T) {
    int cp = cpos[i];
    double d = ds[i];
    if (nc[i]) d += jit * (double)cp;
    dc_out[mb + cp] = d;
    zc_out[mb + cp] = zs[i];
    invorder[mb + i] = cp;
    perm_out[mb + i] = pm[i];
    segstart[mb + i] = ss[i];
  }
  if (tid == 0) {
    k_out[p] = k;
    nseg_out[p] = nseg;
    rho_out[p] = decoupled ? 1.0 : rho;
  }
}

__global__ __launch_bounds__(256, 5) void merge_build_kernel(
    const double* __restrict__ dc_in,    // (P,m) compacted poles (jittered)
    const double* __restrict__ zc_in,    // (P,m) compacted z (signed)
    const int*    __restrict__ k_in,     // (P,) active count
    const double* __restrict__ rho_in,   // (P,) rho_s (de-zeroed)
    const int*    __restrict__ invorder, // (P,m) sorted pos -> compacted idx
    const int*    __restrict__ perm,     // (P,m) sorted pos -> child idx
    const double* __restrict__ hv_in,    // (P,m) sorted-order householder v
    const double* __restrict__ hbeta_in, // (P,m) sorted-order per-pole beta
    const int*    __restrict__ segstart, // (P,m) cluster start (sorted idx), pad m
    const int*    __restrict__ nseg_in,  // (P,) number of clusters
    float*  __restrict__ Vchild,         // (P,m,m) OUT eigenvectors
    double* __restrict__ Lam_out,        // (P,m) OUT eigenvalues (compacted)
    int m, int NEWT, int NF32, double RES_TOL, double STEP_TOL) {
  const int p = blockIdx.x;
  const int tid = threadIdx.x;
  const int T = blockDim.x;

  extern __shared__ double sdm[];
  double* dc    = sdm;            // m
  double* zc    = dc + m;         // m
  double* eta   = zc + m;         // m
  double* zhat  = eta + m;        // m
  double* csc   = zhat + m;       // m (column scale)
  double* hv    = csc + m;        // m
  double* hbeta = hv + m;         // m
  float*  dcf   = (float*)(hbeta + m); // m (fp32 poles hi, for fp32-first newton)
  float*  zcf   = dcf + m;             // m (fp32 z)
  int*    iord  = (int*)(zcf + m);     // m  (invorder)
  int*    prm   = iord + m;            // m  (perm)
  int*    sst   = prm + m;             // m  (segstart)
  __shared__ int k_sh;
  __shared__ double rho_sh;
  __shared__ int nseg_sh;
  __shared__ double far_sh;

  const size_t base = (size_t)p * m;
  for (int i = tid; i < m; i += T) {
    dc[i]    = dc_in[base + i];
    zc[i]    = zc_in[base + i];
    dcf[i]   = (float)dc[i];
    zcf[i]   = (float)zc[i];
    hv[i]    = hv_in[base + i];
    hbeta[i] = hbeta_in[base + i];
    iord[i]  = invorder[base + i];
    prm[i]   = perm[base + i];
    sst[i]   = segstart[base + i];
    eta[i]   = 0.0;
    zhat[i]  = 0.0;
    csc[i]   = 1.0;
  }
  if (tid == 0) { k_sh = k_in[p]; rho_sh = rho_in[p]; nseg_sh = nseg_in[p]; }
  __syncthreads();
  const int k = k_sh;
  const double rho = rho_sh;
  const bool pos = rho > 0.0;

  // ---- far bound = rho * sum_{j<k} zc[j]^2 ----
  double part = 0.0;
  for (int j = tid; j < k; j += T) part += zc[j] * zc[j];
  // block reduce
  __shared__ double red[32];
  for (int off = 16; off > 0; off >>= 1) part += __shfl_down_sync(0xffffffffu, part, off);
  if ((tid & 31) == 0) red[tid >> 5] = part;
  __syncthreads();
  if (tid == 0) {
    double s = 0.0; int nw = (T + 31) >> 5;
    for (int w = 0; w < nw; ++w) s += red[w];
    far_sh = rho * s;
  }
  __syncthreads();
  const double far = far_sh;

  // ---- secular solve: one root i per thread (strided), fp64 safeguarded
  //      Newton with in-bracket clamp + bisection fallback. ----
  const float rhof = (float)rho;
  const float farf = (float)far;
  for (int i = tid; i < k; i += T) {
    const int iU = pos ? (i + 1) : i;   // first index of phi side
    // ---- fp32-first phase (fast on this box) ----
    const float dif = dcf[i];
    float gLf, gUf;
    if (pos) { gLf = 0.0f; gUf = (i == k - 1) ? farf : (dcf[i + 1] - dif); }
    else     { gLf = (i == 0) ? farf : -(dif - dcf[i - 1]); gUf = 0.0f; }
    const float gapwf = gUf - gLf;
    float lof = gLf, hif = gUf, etf = 0.5f * (gLf + gUf);
    #pragma unroll 1
    for (int it = 0; it < NF32; ++it) {
      float gf, eef;
      float tn = mw_step_comp(dcf, zcf, k, iU, dif, rhof, gLf, gUf, etf, lof, hif, pos, gf, eef);
      // residual at the compensated-fp32 roundoff floor (~2^-24*erretm from the
      // single-precision divides): keep etf as the fp64 seed. Cuts the fp32
      // straggler warp-max (clustered/edge roots that used to burn the NF32 cap).
      if (fabsf(gf) <= 2e-7f * eef) break;
      float step = fabsf(tn - etf);
      etf = tn;
      if (step <= 1e-6f * gapwf) break;
    }
    // ---- fp64 polish from the fp32 seed ----
    const double di = dc[i];
    double lo, hi;
    if (pos) { lo = 0.0; hi = (i == k - 1) ? far : (dc[i + 1] - di); }
    else     { lo = (i == 0) ? far : -(di - dc[i - 1]); hi = 0.0; }
    const double gL = lo, gU = hi;
    const double gapw = gU - gL;
    double et = (double)etf;
    if (!(et > lo && et < hi)) et = 0.5 * (lo + hi);
    #pragma unroll 1
    for (int it = 0; it < NEWT; ++it) {
      double gg, ee;
      // Native fp64 polish (2:1 box). Converge on |g| <= RES_TOL*erretm OR
      // |step| <= STEP_TOL*gapwidth. The prior 8e-16/1e-9 chased the fp64 roundoff
      // floor, but tight-gap "straggler" roots reach fp64 eigenvalue/eigenvector
      // accuracy far earlier and then spun to the ~17-iter SIMT warp-max for no
      // gain; RES_TOL/STEP_TOL retire them (n512/n1024 relax to 1e-12/1e-7 -> clustered
      // -1.3ms, all gates byte-identical; n2048 keeps 8e-16/1e-9 to not perturb its
      // later multi-CTA merges). KEEP et on the residual exit (taking the step
      // could let the ill-conditioned near-root quadratic overshoot the wide bracket).
      double tnew = mw_step<double>(dc, zc, k, iU, di, rho, gL, gU, et, lo, hi, pos, gg, ee);
      if (fabs(gg) <= RES_TOL * ee) break;
      double step = fabs(tnew - et);
      et = tnew;
      if (step <= STEP_TOL * gapw) break;
    }
    eta[i] = et;
    Lam_out[base + i] = di + et;
  }
  // deflated eigenvalues
  for (int i = k + tid; i < m; i += T) Lam_out[base + i] = dc[i];
  __syncthreads();

  // ---- zhat (Gu-Eisenstat) in log-space, one pole j per thread ----
  const double inv_sqrt_rho = rsqrt(fabs(rho));
  for (int j = tid; j < k; j += T) {
    const double dj = dc[j];
    const double etj = eta[j];
    // clamp each ratio away from 0 before log (a root landing on a pole makes
    // one (lam_kk - d_j) factor exactly 0 -> log(0) = -inf; proto clamps to
    // 1e-300 so zhat stays a tiny-but-nonzero value, not exactly 0).
    // fp32 hardware logf per term (MUFU, fast) with fp64 accumulation: per-term
    // rel err ~1e-7 in log domain -> ~m*1e-7 rel in zhat after exp, far inside
    // the orthogonality gate (coordinator lever 3). tiny-clamp guards a root
    // landing on a pole (one factor exactly 0 -> log(0)=-inf).
    double logsum = (double)logf(fmaxf(fabsf((float)etj), 1e-30f));
    for (int kk = 0; kk < k; ++kk) {
      if (kk == j) continue;
      double pdkj = dc[kk] - dj;
      // form the numerator difference in fp64 FIRST (roots near a pole give
      // ratio -> 0 on clustered spectra; the cancellation must survive), then
      // multiply by the fp64 reciprocal (cheaper than a divide; the product is
      // cast to fp32 for logf anyway so 1 rounding is ample).
      double ratio = (pdkj + eta[kk]) * ieee_rcp(pdkj);
      logsum += (double)logf(fmaxf(fabsf((float)ratio), 1e-30f));
    }
    double zmag = exp(0.5 * logsum) * inv_sqrt_rho;
    double zsgn = (zc[j] >= 0.0) ? 1.0 : -1.0;
    zhat[j] = zmag * zsgn;
  }
  __syncthreads();

  // ---- column norms for active columns (scale-aware, overflow-safe: a root
  //      landing on a pole gives 1/denom ~ 1e300; naive sum-of-squares would
  //      overflow to inf and corrupt the unit-vector collapse) ----
  for (int i = tid; i < k; i += T) {
    const double di = dc[i];
    const double eti = eta[i];
    // Overflow-scale is any upper bound on max_j |zhat_j/d_j| (csc is
    // scale-invariant: ||raw||^2 = vmax^2 * s regardless of vmax). Use
    // max|zhat| / min|d| -> two reduction passes WITHOUT a per-term reciprocal
    // (only one rcp for the whole column), halving csc's fp64 reciprocals.
    double zmax = 0.0, dmin = 1e300;
    for (int j = 0; j < k; ++j) {
      double d = (dc[j] - di) - eti;
      double ad = fabs(d);
      if (ad < 1e-300) ad = 1e-300;
      if (ad < dmin) dmin = ad;
      double az = fabs(zhat[j]);
      if (az > zmax) zmax = az;
    }
    double vmax = zmax * ieee_rcp(dmin);
    double s = 0.0;
    if (vmax > 0.0) {
      const double inv_vmax = ieee_rcp(vmax);  // scalar reciprocal once, not per-term
      for (int j = 0; j < k; ++j) {
        double d = (dc[j] - di) - eti;
        if (fabs(d) < 1e-300) d = (d < 0.0) ? -1e-300 : 1e-300;
        double inv = ieee_rcp(d);
        double r = zhat[j] * inv;
        r = fma(fma(-r, d, zhat[j]), inv, r);   // accurate zhat/d, fp64 rcp not divide
        double v = r * inv_vmax;
        s += v * v;
      }
    }
    csc[i] = (vmax > 0.0 && s > 0.0) ? (1.0 / (vmax * sqrt(s))) : 1.0;
  }
  __syncthreads();

  // ---- raw V build (no G): Vchild[perm[t], i] = raw(t,i), coalesced over i ----
  const size_t vbase = (size_t)p * m * m;
  for (int idx = tid; idx < m * m; idx += T) {
    const int t = idx / m;    // sorted pole position
    const int i = idx % m;    // compacted root (column)
    const int j = iord[t];    // compacted pole index of sorted pos t
    double val;
    if (i < k) {
      if (j < k) {
        double d = (dc[j] - dc[i]) - eta[i];
        if (fabs(d) < 1e-300) d = (d < 0.0) ? -1e-300 : 1e-300;
        double inv = ieee_rcp(d);
        double r = zhat[j] * inv;
        r = fma(fma(-r, d, zhat[j]), inv, r);   // accurate zhat/d, fp64 rcp not divide
        val = r * csc[i];
      } else {
        val = 0.0;
      }
    } else {
      val = (j == i) ? 1.0 : 0.0;
    }
    Vchild[vbase + (size_t)prm[t] * m + i] = (float)val;
  }
  __syncthreads();

  // ---- Householder G correction over multi-clusters (sorted basis, applied
  //      in child-row layout via perm) ----
  const int nseg = nseg_sh;
  for (int c = 0; c < nseg; ++c) {
    const int t0 = sst[c];
    const int t1 = (c + 1 < nseg) ? sst[c + 1] : m;
    if (t1 - t0 < 2) continue;
    for (int i = tid; i < m; i += T) {
      double w = 0.0;
      for (int t = t0; t < t1; ++t)
        w += hv[t] * (double)Vchild[vbase + (size_t)prm[t] * m + i];
      for (int t = t0; t < t1; ++t) {
        size_t off = vbase + (size_t)prm[t] * m + i;
        Vchild[off] = (float)((double)Vchild[off] - hbeta[t] * hv[t] * w);
      }
    }
    __syncthreads();
  }
}

void launch_dc_prep(
    const double* Dl, const double* Drr, const float* zL, const float* zR,
    const double* rho_in, double* dc_out, double* zc_out, int* k_out,
    double* rho_out, int* invorder, int* perm_out, double* hv_out,
    double* hbeta_out, int* segstart, int* nseg_out, int P, int h, int m,
    double defl_zk, double defl_gk) {
  // Bitonic sort needs a power-of-2 length. For the native pow2-m shapes mp==m
  // and the PAD=false kernel is byte-identical to the original. For native-352
  // (m in {44,88,176,352}) the SMEM sort arrays are padded to mp = next_pow2(m)
  // (pad slots keyed +inf) so the sort is legal; later work stays over m.
  // Deflation-tolerance multipliers: caller passes the shipped value; the
  // DEFL_ZK/DEFL_GK env vars (dev-only) override for sweeps.
  double zk_mul = defl_zk, gk_mul = defl_gk;
  { const char* zs = getenv("DEFL_ZK"); if (zs) zk_mul = atof(zs);
    const char* gs = getenv("DEFL_GK"); if (gs) gk_mul = atof(gs); }
  int mp = 1;
  while (mp < m) mp <<= 1;
  const bool pad = (mp != m);
  const int S = pad ? mp : m;
  int T = (m <= 256) ? 128 : 256;
  size_t smem = (size_t)(5 * S) * sizeof(double) + (size_t)(6 * S) * sizeof(int);
  if (pad) {
    static int cfgp = -1;
    if (cfgp < (int)smem) {
      cudaFuncSetAttribute(dc_prep_kernel<true>,
          cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
      cfgp = (int)smem;
    }
    dc_prep_kernel<true><<<P, T, smem>>>(
        Dl, Drr, zL, zR, rho_in, dc_out, zc_out, k_out, rho_out,
        invorder, perm_out, hv_out, hbeta_out, segstart, nseg_out, h, m, mp, zk_mul, gk_mul);
  } else {
    static int cfg = -1;
    if (cfg < (int)smem) {
      cudaFuncSetAttribute(dc_prep_kernel<false>,
          cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
      cfg = (int)smem;
    }
    dc_prep_kernel<false><<<P, T, smem>>>(
        Dl, Drr, zL, zR, rho_in, dc_out, zc_out, k_out, rho_out,
        invorder, perm_out, hv_out, hbeta_out, segstart, nseg_out, h, m, mp, zk_mul, gk_mul);
  }
}

void launch_merge_build(
    const double* dc_in, const double* zc_in, const int* k_in,
    const double* rho_in, const int* invorder, const int* perm,
    const double* hv_in, const double* hbeta_in, const int* segstart,
    const int* nseg_in, float* Vchild, double* Lam_out,
    int P, int m, int newt, int nf32, double res_tol, double step_tol) {
  // merge_build_kernel is __launch_bounds__(256,5), so the block is always 256
  // threads (a per-level "more threads when the batch under-fills" idea was tried
  // and is moot under that cap).
  constexpr int T = 256;
  size_t smem = (size_t)(7 * m) * sizeof(double) + (size_t)(2 * m) * sizeof(float) + (size_t)(3 * m) * sizeof(int);
  static int cfg = -1;
  if (cfg < (int)smem) {
    cudaFuncSetAttribute(merge_build_kernel,
        cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
    cfg = (int)smem;
  }
  merge_build_kernel<<<P, T, smem>>>(
      dc_in, zc_in, k_in, rho_in, invorder, perm, hv_in, hbeta_in,
      segstart, nseg_in, Vchild, Lam_out, m, newt, nf32, res_tol, step_tol);
}

// ===========================================================================
// MULTI-CTA (thread-block cluster) merge: G CTAs cooperate on one merge.
// The top-of-tree merges (m=512..2048) run with P = B*newK CTAs << 148 SMs
// (n1024 b60 top merge: 60 CTAs -> 0.41 waves; n2048 b8: 8 CTAs). One CTA per
// merge leaves most SMs idle AND each CTA serialises the O(m^2) secular + zhat
// + V-build. Splitting a merge across G CTAs fills the machine and shortens the
// per-CTA critical path G-fold. Partition: secular over roots, zhat over poles,
// csc/V-build/G-correction over COLUMNS. eta and zhat are exchanged across the
// cluster (cg::map_shared_rank) with a cluster.sync between phases. The
// per-column arithmetic is byte-identical to merge_build_kernel, so precision is
// preserved; only loop bounds + the two gathers change.
template<int G>
__global__ __launch_bounds__(1024, 1)
void merge_build_cluster_kernel(
    const double* __restrict__ dc_in, const double* __restrict__ zc_in,
    const int* __restrict__ k_in, const double* __restrict__ rho_in,
    const int* __restrict__ invorder, const int* __restrict__ perm,
    const double* __restrict__ hv_in, const double* __restrict__ hbeta_in,
    const int* __restrict__ segstart, const int* __restrict__ nseg_in,
    float* __restrict__ Vchild, double* __restrict__ Lam_out,
    int m, int NEWT, int NF32) {
  cg::cluster_group cluster = cg::this_cluster();
  const int rank = cluster.block_rank();
  const int p = blockIdx.x / G;
  const int tid = threadIdx.x;
  const int T = blockDim.x;

  extern __shared__ double sdm[];
  double* dc    = sdm;            // m
  double* zc    = dc + m;         // m
  double* eta   = zc + m;         // m
  double* zhat  = eta + m;        // m
  double* csc   = zhat + m;       // m
  double* hv    = csc + m;        // m
  double* hbeta = hv + m;         // m
  float*  dcf   = (float*)(hbeta + m); // m
  float*  zcf   = dcf + m;             // m
  int*    iord  = (int*)(zcf + m);     // m
  int*    prm   = iord + m;            // m
  int*    sst   = prm + m;             // m
  __shared__ int k_sh;
  __shared__ double rho_sh;
  __shared__ int nseg_sh;
  __shared__ double far_sh;

  const size_t base = (size_t)p * m;
  for (int i = tid; i < m; i += T) {
    dc[i]=dc_in[base+i]; zc[i]=zc_in[base+i];
    dcf[i]=(float)dc[i]; zcf[i]=(float)zc[i];
    hv[i]=hv_in[base+i]; hbeta[i]=hbeta_in[base+i];
    iord[i]=invorder[base+i]; prm[i]=perm[base+i]; sst[i]=segstart[base+i];
    eta[i]=0.0; zhat[i]=0.0; csc[i]=1.0;
  }
  if (tid==0){ k_sh=k_in[p]; rho_sh=rho_in[p]; nseg_sh=nseg_in[p]; }
  __syncthreads();
  const int k = k_sh;
  const double rho = rho_sh;
  const bool pos = rho > 0.0;

  const int ks = (k + G - 1) / G;            // roots/poles per rank
  const int cs = (m + G - 1) / G;            // columns per rank
  const int rlo = rank*ks, rhi = min((rank+1)*ks, k);
  const int clo = rank*cs, chi = min((rank+1)*cs, m);

  // ---- far bound (redundant per CTA) ----
  double part = 0.0;
  for (int j = tid; j < k; j += T) part += zc[j]*zc[j];
  __shared__ double red[32];
  for (int off=16; off>0; off>>=1) part += __shfl_down_sync(0xffffffffu, part, off);
  if ((tid&31)==0) red[tid>>5]=part;
  __syncthreads();
  if (tid==0){ double s=0.0; int nw=(T+31)>>5; for(int w=0;w<nw;++w) s+=red[w]; far_sh=rho*s; }
  __syncthreads();
  const double far = far_sh;

  // ---- secular solve over ROOT slice [rlo,rhi) ----
  const float rhof=(float)rho, farf=(float)far;
  for (int i = rlo + tid; i < rhi; i += T) {
    const int iU = pos ? (i+1) : i;
    const float dif = dcf[i];
    float gLf, gUf;
    if (pos){ gLf=0.0f; gUf=(i==k-1)?farf:(dcf[i+1]-dif); }
    else    { gLf=(i==0)?farf:-(dif-dcf[i-1]); gUf=0.0f; }
    const float gapwf=gUf-gLf;
    float lof=gLf, hif=gUf, etf=0.5f*(gLf+gUf);
    #pragma unroll 1
    for (int it=0; it<NF32; ++it){
      float gf, eef; float tn=mw_step_comp<true>(dcf,zcf,k,iU,dif,rhof,gLf,gUf,etf,lof,hif,pos,gf,eef);
      if (fabsf(gf)<=2e-7f*eef) break;                 // keep etf (residual at fp32 floor)
      if (fabsf(tn-etf)<=1e-6f*gapwf){ etf=tn; break; }
      etf=tn;
    }
    const double di=dc[i]; double lo,hi;
    if (pos){ lo=0.0; hi=(i==k-1)?far:(dc[i+1]-di); }
    else    { lo=(i==0)?far:-(di-dc[i-1]); hi=0.0; }
    const double gL=lo, gU=hi, gapw=gU-gL;
    double et=(double)etf; if(!(et>lo&&et<hi)) et=0.5*(lo+hi);
    #pragma unroll 1
    for (int it=0; it<NEWT; ++it){
      // native fp64 polish (2:1 box): residual-floor convergence.
      double gg, ee; double tnew=mw_step<double,true>(dc,zc,k,iU,di,rho,gL,gU,et,lo,hi,pos,gg,ee);
      if (fabs(gg)<=8e-16*ee) break;            // keep et (residual at floor)
      if (fabs(tnew-et)<=1e-9*gapw){ et=tnew; break; }
      et=tnew;
    }
    eta[i]=et; Lam_out[base+i]=di+et;
  }
  // deflated eigenvalues over this rank's COLUMN slice (i>=k)
  for (int i = max(clo,k) + tid; i < chi; i += T) Lam_out[base+i]=dc[i];
  // NOTE: eta must be visible cluster-wide before zhat -> cluster.sync + gather.
  cluster.sync();
  for (int i = tid; i < k; i += T) {
    int owner = i / ks;
    if (owner != rank) eta[i] = ((const double*)cluster.map_shared_rank(eta, owner))[i];
  }
  __syncthreads();

  // ---- zhat (Gu-Eisenstat) over POLE slice [rlo,rhi) ----
  const double inv_sqrt_rho = rsqrt(fabs(rho));
  for (int j = rlo + tid; j < rhi; j += T) {
    const double dj=dc[j], etj=eta[j];
    double logsum=(double)logf(fmaxf(fabsf((float)etj),1e-30f));
    for (int kk=0; kk<k; ++kk){
      if (kk==j) continue;
      double pdkj=dc[kk]-dj;
      double ratio=(pdkj+eta[kk])*ieee_rcp(pdkj);
      logsum += (double)logf(fmaxf(fabsf((float)ratio),1e-30f));
    }
    double zmag=exp(0.5*logsum)*inv_sqrt_rho;
    double zsgn=(zc[j]>=0.0)?1.0:-1.0;
    zhat[j]=zmag*zsgn;
  }
  cluster.sync();
  for (int j = tid; j < k; j += T) {
    int owner = j / ks;
    if (owner != rank) zhat[j] = ((const double*)cluster.map_shared_rank(zhat, owner))[j];
  }
  __syncthreads();

  // ---- csc over this rank's COLUMN slice (i<k) ----
  for (int i = max(clo,0) + tid; i < min(chi,k); i += T) {
    const double di=dc[i], eti=eta[i];
    double zmax=0.0, dmin=1e300;
    for (int j=0;j<k;++j){
      double d=(dc[j]-di)-eti; double ad=fabs(d); if(ad<1e-300) ad=1e-300;
      if(ad<dmin) dmin=ad; double az=fabs(zhat[j]); if(az>zmax) zmax=az;
    }
    double vmax=zmax*ieee_rcp(dmin); double s=0.0;
    if (vmax>0.0){
      const double inv_vmax=ieee_rcp(vmax);
      for(int j=0;j<k;++j){
        double d=(dc[j]-di)-eti; if(fabs(d)<1e-300) d=(d<0.0)?-1e-300:1e-300;
        double inv=ieee_rcp(d); double r=zhat[j]*inv; r=fma(fma(-r,d,zhat[j]),inv,r);
        double v=r*inv_vmax; s+=v*v;
      }
    }
    csc[i]=(vmax>0.0&&s>0.0)?(1.0/(vmax*sqrt(s))):1.0;
  }
  __syncthreads();

  // ---- V build over this rank's COLUMN slice: coalesced over i within slab ----
  const size_t vbase=(size_t)p*m*m;
  const int ncol = chi - clo;
  for (int idx=tid; idx < m*ncol; idx += T) {
    const int t = idx / ncol;
    const int i = clo + idx % ncol;
    const int j = iord[t];
    double val;
    if (i<k){
      if (j<k){
        double d=(dc[j]-dc[i])-eta[i]; if(fabs(d)<1e-300) d=(d<0.0)?-1e-300:1e-300;
        double inv=ieee_rcp(d); double r=zhat[j]*inv; r=fma(fma(-r,d,zhat[j]),inv,r);
        val=r*csc[i];
      } else val=0.0;
    } else val=(j==i)?1.0:0.0;
    Vchild[vbase+(size_t)prm[t]*m+i]=(float)val;
  }
  __syncthreads();

  // ---- Householder G correction over multi-clusters, this rank's columns ----
  const int nseg=nseg_sh;
  for (int c=0;c<nseg;++c){
    const int t0=sst[c]; const int t1=(c+1<nseg)?sst[c+1]:m;
    if (t1-t0<2) continue;
    for (int i=clo+tid; i<chi; i+=T){
      double w=0.0;
      for(int t=t0;t<t1;++t) w += hv[t]*(double)Vchild[vbase+(size_t)prm[t]*m+i];
      for(int t=t0;t<t1;++t){ size_t off=vbase+(size_t)prm[t]*m+i;
        Vchild[off]=(float)((double)Vchild[off]-hbeta[t]*hv[t]*w); }
    }
    __syncthreads();
  }
  cluster.sync();  // exit guard: no CTA may leave while a peer still maps its SMEM
}

void launch_merge_build_multi(
    const double* dc_in, const double* zc_in, const int* k_in,
    const double* rho_in, const int* invorder, const int* perm,
    const double* hv_in, const double* hbeta_in, const int* segstart,
    const int* nseg_in, float* Vchild, double* Lam_out,
    int P, int m, int newt, int nf32, int G) {
  constexpr int T = 512;  // merge_build_cluster_kernel is __launch_bounds__(1024,1)
  size_t smem = (size_t)(7*m)*sizeof(double)+(size_t)(2*m)*sizeof(float)+(size_t)(3*m)*sizeof(int);
  auto launch=[&](auto kfn){
    cudaFuncSetAttribute(kfn, cudaFuncAttributeMaxDynamicSharedMemorySize,(int)smem);
    cudaLaunchConfig_t cfg={};
    cfg.gridDim=dim3(G*P); cfg.blockDim=dim3(T); cfg.dynamicSmemBytes=smem;
    cudaLaunchAttribute attr[1];
    attr[0].id=cudaLaunchAttributeClusterDimension;
    attr[0].val.clusterDim.x=G; attr[0].val.clusterDim.y=1; attr[0].val.clusterDim.z=1;
    cfg.attrs=attr; cfg.numAttrs=1;
    cudaLaunchKernelEx(&cfg, kfn, dc_in, zc_in, k_in, rho_in, invorder, perm,
                       hv_in, hbeta_in, segstart, nseg_in, Vchild, Lam_out, m, newt, nf32);
  };
  if (G==2) launch(merge_build_cluster_kernel<2>);
  else if (G==3) launch(merge_build_cluster_kernel<3>);
  else if (G==4) launch(merge_build_cluster_kernel<4>);
  else if (G==6) launch(merge_build_cluster_kernel<6>);
  else if (G==8) launch(merge_build_cluster_kernel<8>);
}

}  // namespace dc

// ===========================================================================
// One-stage tridiagonalization (blocked sytrd / LAPACK slatrd) panel kernel.
// One CTA per matrix. Reduces NB=32 columns of the trailing block At =
// A[p0:, p0:] (A is b x n x n, symmetric, full storage). The per-column symv
// p = At @ v is the BLAS-2 wall; everything else (column correction, reflector
// gen, p correction, w) is on-chip. Accumulates V,W (m x NB) in SMEM (fp16
// for the correction dots -> higher occupancy) and writes fp32 V,W out for the
// deferred rank-2b trailing update and the WY back-transform.
// NOTE: reuses the file-scope warp_sum(float,int=32) and FULL_MASK.
// ===========================================================================
namespace os1 {

template<int NB, int THREADS, int MAXV>
__global__ __launch_bounds__(THREADS,3)
void latrd_kernel(const float* __restrict__ A, const __half* __restrict__ Ah,
                  float* __restrict__ Vout, float* __restrict__ Wout,
                  float* __restrict__ dout, float* __restrict__ eout,
                  int n, int p0){
  const int b = blockIdx.x;
  const int tid = threadIdx.x;
  const int lane = tid & 31;
  const int warp = tid >> 5;
  constexpr int WARPS = THREADS/32;
  const int m = n - p0;
  constexpr int SS = NB + 1;                       // padded SMEM row stride (odd -> no bank conflict)
  const float* At = A + (size_t)b*n*n + (size_t)p0*n + p0;   // At[r][c] = At[r*n + c]
  // fp16 shadow of the trailing matrix, used ONLY for the O(m^2) symv (halves the
  // dominant DRAM traffic). Inputs are amax-normalized to O(1) -> fp16 (11-bit
  // mantissa) is in-range and ~8x more accurate than bf16. Column-correction still
  // reads the fp32 A (row i) so d/e stay fp32-accurate.
  const __half* Aht = Ah + (size_t)b*n*n + (size_t)p0*n + p0;
  Vout += (size_t)b*m*NB;  Wout += (size_t)b*m*NB;
  dout += (size_t)b*NB;    eout += (size_t)b*NB;

  extern __shared__ char smemc[];
  __half* Vs = (__half*)smemc;                    // [m*SS]  fp16 reflectors (padded stride)
  __half* Ws = Vs + (size_t)m*SS;                 // [m*SS]  fp16 W (padded stride)
  float* col = (float*)(Ws + (size_t)m*SS);       // [m]  working column / p
  float* vcur = col + m;                          // [m]  current reflector, fp32 contiguous
  float* red = vcur + m;                          // [WARPS]
  float* Vtv = red + WARPS;                       // [NB]
  float* Wtv = Vtv + NB;                          // [NB]

  // zero V,W panel (fp16)
  for(int idx=tid; idx<m*SS; idx+=THREADS){ Vs[idx]=__float2half(0.f); Ws[idx]=__float2half(0.f); }
  __syncthreads();

  for(int i=0;i<NB;++i){
    // ---- 1. column correction: col[r] = At[r][i] - sum_{k<i} Vs[r][k]*Ws[i][k] + Ws[r][k]*Vs[i][k]
    // At is symmetric (rank-2b update preserves symmetry) -> read ROW i (coalesced)
    // instead of COLUMN i (stride-n, uncoalesced): At[i][r] == At[r][i].
    for(int r=tid; r<m; r+=THREADS){
      float v = At[(size_t)i*n + r];
      #pragma unroll 4
      for(int k=0;k<i;++k){
        v -= __half2float(Vs[r*SS+k])*__half2float(Ws[i*SS+k])
           + __half2float(Ws[r*SS+k])*__half2float(Vs[i*SS+k]);
      }
      col[r] = v;
    }
    __syncthreads();
    if(tid==0) dout[i] = col[i];
    int Llen = m - i - 1;                 // subcolumn length below diagonal
    if(Llen <= 0){ __syncthreads(); continue; }
    float alpha = col[i+1];
    if(Llen == 1){
      if(tid==0) eout[i] = alpha;
      __syncthreads();
      continue;
    }
    // ---- 2. reflector: tail = sum_{r>=i+2} col[r]^2
    float part=0.f;
    for(int r=i+2+tid; r<m; r+=THREADS){ float x=col[r]; part+=x*x; }
    part = warp_sum(part);
    if(lane==0) red[warp]=part;
    __syncthreads();
    float tail=0.f;
    #pragma unroll
    for(int w=0;w<WARPS;++w) tail += red[w];
    float norm = sqrtf(alpha*alpha + tail);
    bool has = norm > 0.f;
    float beta = (alpha>=0.f)? -norm : norm;
    float tau  = has ? (beta-alpha)/beta : 0.f;
    float invd = has ? 1.f/(alpha-beta) : 0.f;
    if(tid==0){ eout[i] = has? beta : alpha; }
    // build v: v[i+1]=1, v[r]=col[r]*invd for r>=i+2 (into vcur fp32, Vs fp16, Vout fp32)
    if(tid==0){ vcur[i+1]=1.f; Vs[(i+1)*SS+i]=__float2half(1.f); Vout[(i+1)*NB+i]=1.f; }
    for(int r=i+2+tid; r<m; r+=THREADS){
      float vv = col[r]*invd; vcur[r]=vv; Vs[r*SS+i]=__float2half(vv); Vout[r*NB+i]=vv;
    }
    __syncthreads();
    // ---- 3. symv p = At @ v (fp32 vcur), warp-per-row; preload v into registers ONCE.
    // Process 2 rows/warp-iter with interleaved loads -> 2x memory-level parallelism
    // (2*MAXV independent loads in flight before the warp reduction) to hide DRAM
    // latency on this latency-bound kernel. Numerically identical.
    float vreg[MAXV];
    #pragma unroll
    for(int t=0;t<MAXV;++t){ int c=i+1+lane+32*t; vreg[t] = (c<m)? vcur[c] : 0.f; }
    int r0 = i+1+warp;
    constexpr int RU = 4;
    for(; r0+(RU-1)*WARPS<m; r0+=RU*WARPS){
      const __half* Arp[RU];
      #pragma unroll
      for(int u=0;u<RU;++u) Arp[u] = Aht + (size_t)(r0+u*WARPS)*n;
      float dd[RU];
      #pragma unroll
      for(int u=0;u<RU;++u) dd[u]=0.f;
      #pragma unroll
      for(int t=0;t<MAXV;++t){ int c=i+1+lane+32*t; if(c<m){ float vt=vreg[t];
        #pragma unroll
        for(int u=0;u<RU;++u) dd[u]+=ld_half_cs(Arp[u]+c)*vt; } }
      #pragma unroll
      for(int u=0;u<RU;++u){ dd[u]=warp_sum(dd[u]); if(lane==0) col[r0+u*WARPS]=dd[u]; }
    }
    for(; r0<m; r0+=WARPS){
      const __half* Ar = Aht + (size_t)r0*n;
      float d=0.f;
      #pragma unroll
      for(int t=0;t<MAXV;++t){ int c=i+1+lane+32*t; if(c<m) d += ld_half_cs(Ar+c)*vreg[t]; }
      d = warp_sum(d);
      if(lane==0) col[r0] = d;             // reuse col[] as p
    }
    __syncthreads();
    // ---- 4. correct p: Vtv[k]=sum_r Vs[r][k]*v[r], Wtv[k]=sum_r Ws[r][k]*v[r]  (k<i)
    for(int k=warp; k<i; k+=WARPS){
      float sv=0.f, sw=0.f;
      for(int r=i+1+lane; r<m; r+=32){ float vr=vcur[r]; sv+=__half2float(Vs[r*SS+k])*vr; sw+=__half2float(Ws[r*SS+k])*vr; }
      sv=warp_sum(sv); sw=warp_sum(sw);
      if(lane==0){ Vtv[k]=sv; Wtv[k]=sw; }
    }
    __syncthreads();
    // p[r] -= sum_k Ws[r][k]*Vtv[k] + Vs[r][k]*Wtv[k]  ; then p*=tau
    for(int r=i+1+tid; r<m; r+=THREADS){
      float pr = col[r];
      #pragma unroll 4
      for(int k=0;k<i;++k) pr -= __half2float(Ws[r*SS+k])*Vtv[k] + __half2float(Vs[r*SS+k])*Wtv[k];
      col[r] = pr*tau;
    }
    __syncthreads();
    // ---- 5. wtv = p^T v ; w = p - (tau/2 * wtv) v
    float pv=0.f;
    for(int r=i+1+tid; r<m; r+=THREADS){ pv += col[r]*vcur[r]; }
    pv = warp_sum(pv);
    if(lane==0) red[warp]=pv;
    __syncthreads();
    float wtv=0.f;
    #pragma unroll
    for(int w=0;w<WARPS;++w) wtv += red[w];
    float coef = 0.5f*tau*wtv;
    for(int r=i+1+tid; r<m; r+=THREADS){
      float wv = col[r] - coef*vcur[r];
      Ws[r*SS+i]=__float2half(wv); Wout[r*NB+i]=wv;
    }
    __syncthreads();
  }
}

// ===========================================================================
// TAIL COLLAPSE: once the trailing block m <= M0 fits in SMEM, finish the
// ENTIRE remaining tridiagonalization in ONE launch (one CTA/matrix), fully
// SMEM-resident. Replaces the last few blocked latrd panels + their per-panel
// trailing GEMM + epilogue + glue launches with a single unblocked LAPACK
// ssytd2 (in-place rank-2 update in SMEM, fp32). Emits the reflector matrix
// Vfull (b,M0,M0) trapezoidal [column i holds the i-th Householder vector,
// unit at row i+1, zero above] and d/e for columns p0..p0+M0-1, in the exact
// layout _aggregate_refl expects (the host slices Vfull into M0/32 panel views).
// d/e stay fp32 (eigenvalue/cluster accuracy) since the update is fp32 SMEM.
// 4 __syncthreads per column: (tail-reduce, v-ready, xtv-reduce, post-update).
// The symv + xtv are fused (lane0 accumulates x^T v while writing x), and the
// w = x - (tau/2 x^Tv) v vector is folded directly into the rank-2 update
// (A -= v x^T + x v^T - 2*coef v v^T) so no separate w pass/sync is needed.
// ===========================================================================
template<int M0, int THREADS, int OCC>
__global__ __launch_bounds__(THREADS, OCC)
void latrd_tail_kernel(const float* __restrict__ A,
                       float* __restrict__ Vout,   // (b, M0, M0)
                       float* __restrict__ dout,   // (b, M0)
                       float* __restrict__ eout,   // (b, M0)
                       int n, int p0){
  const int bb  = blockIdx.x;
  const int tid = threadIdx.x;
  const int lane = tid & 31;
  const int warp = tid >> 5;
  constexpr int WARPS = THREADS/32;
  constexpr int SS = M0 + 1;                    // padded SMEM stride (no bank conflict)
  const float* At = A + (size_t)bb*n*n + (size_t)p0*n + p0;
  Vout += (size_t)bb*M0*M0;
  dout += (size_t)bb*M0;
  eout += (size_t)bb*M0;

  extern __shared__ char smemc[];
  float* As  = (float*)smemc;                   // [M0*SS]  symmetric block, updated in place
  float* vv  = As + (size_t)M0*SS;              // [M0]     current reflector
  float* pw  = vv + M0;                         // [M0]     x = tau*A*v
  float* red = pw + M0;                         // [WARPS]

  for(int idx=tid; idx<M0*M0; idx+=THREADS){
    int r = idx / M0, c = idx - r*M0;
    As[r*SS + c] = At[(size_t)r*n + c];
    Vout[idx] = 0.f;
  }
  __syncthreads();

  for(int i=0;i<M0-1;++i){
    if(tid==0) dout[i] = As[i*SS+i];
    float alpha = As[(i+1)*SS + i];
    // ---- reflector tail = sum_{r>=i+2} As[r][i]^2 (block reduction)
    float part=0.f;
    for(int r=i+2+tid; r<M0; r+=THREADS){ float x=As[r*SS+i]; part+=x*x; }
    part = warp_sum(part);
    if(lane==0) red[warp]=part;
    __syncthreads();
    float tail=0.f;
    #pragma unroll
    for(int w=0;w<WARPS;++w) tail += red[w];
    float norm = sqrtf(alpha*alpha + tail);
    bool has = norm > 0.f;
    float beta = (alpha>=0.f)? -norm : norm;
    float tau  = has ? (beta-alpha)/beta : 0.f;
    float invd = has ? 1.f/(alpha-beta) : 0.f;
    if(tid==0) eout[i] = has ? beta : alpha;
    // ---- build v: v[i+1]=1, v[r]=As[r][i]*invd (r>=i+2)
    if(tid==0){ vv[i+1]=1.f; Vout[(i+1)*M0+i]=1.f; }
    for(int r=i+2+tid; r<M0; r+=THREADS){ float x=As[r*SS+i]*invd; vv[r]=x; Vout[r*M0+i]=x; }
    __syncthreads();
    // ---- symv x = tau*A*v (rows i+1..), warp-per-row; fuse xtv = x^T v
    float xtv_local = 0.f;
    for(int r=i+1+warp; r<M0; r+=WARPS){
      float acc=0.f;
      for(int c=i+1+lane; c<M0; c+=32) acc += As[r*SS+c]*vv[c];
      acc = warp_sum(acc);
      float xr = acc*tau;
      if(lane==0){ pw[r]=xr; xtv_local += xr*vv[r]; }
    }
    xtv_local = warp_sum(xtv_local);
    if(lane==0) red[warp]=xtv_local;
    __syncthreads();
    float xtv=0.f;
    #pragma unroll
    for(int w=0;w<WARPS;++w) xtv += red[w];
    float coef = 0.5f*tau*xtv;
    // ---- rank-2 update A -= v x^T + x v^T - 2*coef v v^T (rows,cols i+1..)
    int base=i+1, span=M0-base;
    for(int idx=tid; idx<span*span; idx+=THREADS){
      int rr=idx/span, cc=idx-rr*span;
      int r=base+rr, c=base+cc;
      As[r*SS+c] -= vv[r]*pw[c] + pw[r]*vv[c] - 2.f*coef*vv[r]*vv[c];
    }
    __syncthreads();
  }
  if(tid==0) dout[M0-1] = As[(size_t)(M0-1)*SS+(M0-1)];
}

void launch_latrd_tail(const float* A, float* Vout, float* dout, float* eout,
                       int b, int n, int p0, int m0){
  auto go = [&](auto kern, int threads){
    size_t smem = ((size_t)m0*(m0+1) + 2*m0 + threads/32)*sizeof(float);
    cudaFuncSetAttribute(kern, cudaFuncAttributeMaxDynamicSharedMemorySize, smem);
    kern<<<b, threads, smem>>>(A, Vout, dout, eout, n, p0);
  };
  // Small M0: latency/occupancy bound -> favour MANY CTAs/SM (fewer threads/CTA,
  // fp32 SMEM block caps residency). Larger M0: more threads to shorten the
  // per-column serial passes.
  if(m0==128)      go(latrd_tail_kernel<128,256,3>, 256);
  else if(m0==160) go(latrd_tail_kernel<160,256,2>, 256);
  else if(m0==192) go(latrd_tail_kernel<192,256,1>, 192);
  else if(m0==96)  go(latrd_tail_kernel<96,128,8>, 128);
  else if(m0==64)  go(latrd_tail_kernel<64,128,12>, 128);
}

// Compact-WY T (nb x nb upper) from the Gram G = V^T V (b,nb,nb), which the
// caller computes with a tensor-core bmm. One warp per matrix: tau_j =
// 2/G[j][j]; T[j][j]=tau_j; T[0:j][j] = -tau_j T[0:j][0:j] G[0:j][j].
// (T is orthogonality-critical -> Gram must be fp32.)
template<int NB>
__global__ void trec_kernel(const float* __restrict__ Gin, float* __restrict__ Tout){
  const int b = blockIdx.x, lane = threadIdx.x;   // one warp
  Gin += (size_t)b*NB*NB; Tout += (size_t)b*NB*NB;
  __shared__ float G[NB*NB];
  __shared__ float Tm[NB*NB];
  for(int idx=lane; idx<NB*NB; idx+=32) G[idx]=Gin[idx];
  __syncwarp();
  float tauL = (lane<NB && G[lane*NB+lane]>0.f) ? 2.f/G[lane*NB+lane] : 0.f;
  for(int j=0;j<NB;++j){
    float tj = __shfl_sync(FULL_MASK, tauL, j);
    int a = lane;
    if(a<j){ float z=0.f; for(int k=a;k<j;++k) z += Tm[a*NB+k]*G[k*NB+j]; Tm[a*NB+j] = -tj*z; }
    if(a==j) Tm[j*NB+j]=tj;
    __syncwarp();
  }
  for(int idx=lane; idx<NB*NB; idx+=32){ int a=idx/NB,bb=idx-a*NB; Tout[idx]=(a<=bb)?Tm[idx]:0.f; }
}

void launch_latrd(const float* A, const __half* Ah, float* Vout, float* Wout,
                  float* dout, float* eout, int b, int n, int p0){
  int m = n - p0;
  constexpr int NB=32, THREADS=256;
  size_t smem = (size_t)2*m*(NB+1)*sizeof(__half) + ((size_t)2*m + THREADS/32 + 2*NB)*sizeof(float);
  int maxv = (m + 31) / 32;
  auto k = (maxv<=16)? latrd_kernel<NB,THREADS,16> : latrd_kernel<NB,THREADS,32>;
  cudaFuncSetAttribute(k, cudaFuncAttributeMaxDynamicSharedMemorySize, smem);
  k<<<b, THREADS, smem>>>(A, Ah, Vout, Wout, dout, eout, n, p0);
}

void launch_trec(const float* G, float* T, int b, int nb){
  if(nb==16) trec_kernel<16><<<b,32>>>(G, T);
  else       trec_kernel<32><<<b,32>>>(G, T);
}

}  // namespace os1

// ===========================================================================
// Thread-block CLUSTER latrd for one-stage tridiagonalization (n1024/n2048).
// C CTAs per matrix; symv rows split across the cluster, V/W replicated.
// Same 4 symv wins as os1::latrd (padded SMEM stride, symmetric coalesced
// column load, RU=4 row-unrolled symv). Fills SMs when batch << SM count.
// ===========================================================================

// TMA (cp.async.bulk) + mbarrier helpers for the warp-private prefetch symv.
__device__ __forceinline__ unsigned smem_u32(const void* p){ return (unsigned)__cvta_generic_to_shared(p); }
__device__ __forceinline__ void mbar_init1(void* p){
  asm volatile("mbarrier.init.shared::cta.b64 [%0], 1;" :: "r"(smem_u32(p)));
}
__device__ __forceinline__ void mbar_init_fence(){ asm volatile("fence.mbarrier_init.release.cluster;"); }
__device__ __forceinline__ void mbar_expect(void* p, int bytes){
  asm volatile("mbarrier.arrive.expect_tx.relaxed.cluster.shared::cluster.b64 _, [%0], %1;"
               :: "r"(smem_u32(p)), "r"(bytes) : "memory");
}
__device__ __forceinline__ void mbar_wait(void* p, int phase){
  asm volatile("{\n\t.reg .pred P;\n\tLwt_cltma:\n\t"
    "mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P, [%0], %1, 0x989680;\n\t"
    "@!P bra Lwt_cltma;\n\t}" :: "r"(smem_u32(p)), "r"(phase));
}
__device__ __forceinline__ void tma_g2s(void* dst, const void* src, int bytes, void* mbar){
  asm volatile("cp.async.bulk.shared::cluster.global.mbarrier::complete_tx::bytes [%0], [%1], %2, [%3];"
               :: "r"(smem_u32(dst)), "l"(src), "r"(bytes), "r"(smem_u32(mbar)) : "memory");
}
// 2D tensor-map bulk copy: one arrival brings a TR-row x TC-col contiguous tile
// (amortizes the mbarrier/issue over TR*TC elements). coords = {col, row}
// (innermost first), matching cuTensorMapEncodeTiled's globalDim[0]=cols layout.
__device__ __forceinline__ void tma_2d(void* dst, const CUtensorMap* tmap,
                                       int col, int row, void* mbar){
  asm volatile(
    "cp.async.bulk.tensor.2d.shared::cluster.global.mbarrier::complete_tx::bytes"
    " [%0], [%1, {%2, %3}], [%4];"
    :: "r"(smem_u32(dst)), "l"(tmap), "r"(col), "r"(row), "r"(smem_u32(mbar)) : "memory");
}

namespace os1cl {
template<int NB, int THREADS, int MAXV, int C>
__global__ __launch_bounds__(THREADS,1)
void latrd_cluster_kernel(const float* __restrict__ A, const __half* __restrict__ Ah,
                  float* __restrict__ Vout, float* __restrict__ Wout,
                  float* __restrict__ dout, float* __restrict__ eout,
                  int n, int p0){
  cg::cluster_group cluster = cg::this_cluster();
  const int rank = cluster.block_rank();
  const int b = blockIdx.x / C;
  const int tid = threadIdx.x;
  const int lane = tid & 31;
  const int warp = tid >> 5;
  constexpr int WARPS = THREADS/32;
  const int m = n - p0;
  constexpr int SS = NB + 1;                       // padded SMEM row stride (odd -> no 16-way bank conflict)
  const float* At = A + (size_t)b*n*n + (size_t)p0*n + p0;
  // fp16 shadow of the trailing matrix, read ONLY by the O(m^2) symv sweep. ncu on
  // the m=1024 first panel shows this kernel is memory-LATENCY bound (8 warps, IPC
  // 0.40) with an L2 hit rate of only 18.5% — the 60-matrix x 4MB fp32 working set
  // thrashes L2. Halving the symv bytes shrinks the working set (better L2 residency
  // => lower avg load latency) AND halves DRAM traffic. NORMAL cached loads (NOT the
  // os1 ld.global.cs evict-first, which would defeat the L2-reuse this path needs).
  // Column-correction still reads the fp32 A (row i) so d/e stay fp32-accurate.
  const __half* Aht = Ah + (size_t)b*n*n + (size_t)p0*n + p0;
  Vout += (size_t)b*m*NB;  Wout += (size_t)b*m*NB;
  dout += (size_t)b*NB;    eout += (size_t)b*NB;
  const int slab = (m + C - 1)/C;
  const int rlo = rank*slab, rhi = min((rank+1)*slab, m);

  extern __shared__ char smemc[];
  // cross-CTA region (identical offset in every CTA -> DSMEM addressable)
  float* xw   = (float*)smemc;                    // [m]  w-column exchange
  float* xred = xw + m;                            // [4]  scalar reductions per rank
  // private
  __half* Vs = (__half*)(xred + 4);               // [m*SS]
  __half* Ws = Vs + (size_t)m*SS;                 // [m*SS]
  float* col = (float*)(Ws + (size_t)m*SS);       // [m]  working column / v-source
  float* vcur = col + m;                          // [m]  current reflector, fp32
  float* pslab = vcur + m;                        // [m]  symv output (only slab filled)
  float* red = pslab + m;                         // [WARPS]
  float* Vtv = red + WARPS;                       // [NB]
  float* Wtv = Vtv + NB;                          // [NB]

  for(int idx=tid; idx<m*SS; idx+=THREADS){ Vs[idx]=__float2half(0.f); Ws[idx]=__float2half(0.f); }
  __syncthreads();

  for(int i=0;i<NB;++i){
    // 1. column correction (full, redundant per CTA). At is symmetric -> read ROW i
    // (coalesced) not COLUMN i (stride-n, uncoalesced): At[i][r] == At[r][i].
    for(int r=tid; r<m; r+=THREADS){
      float v = At[(size_t)i*n + r];
      #pragma unroll 4
      for(int k=0;k<i;++k){
        v -= __half2float(Vs[r*SS+k])*__half2float(Ws[i*SS+k])
           + __half2float(Ws[r*SS+k])*__half2float(Vs[i*SS+k]);
      }
      col[r] = v;
    }
    __syncthreads();
    if(tid==0 && rank==0) dout[i] = col[i];
    int Llen = m - i - 1;
    if(Llen <= 0){ __syncthreads(); continue; }
    float alpha = col[i+1];
    if(Llen == 1){
      if(tid==0 && rank==0) eout[i] = alpha;
      __syncthreads();
      continue;
    }
    // 2. reflector (full, redundant)
    float part=0.f;
    for(int r=i+2+tid; r<m; r+=THREADS){ float x=col[r]; part+=x*x; }
    part = warp_sum(part);
    if(lane==0) red[warp]=part;
    __syncthreads();
    float tail=0.f;
    #pragma unroll
    for(int w=0;w<WARPS;++w) tail += red[w];
    float norm = sqrtf(alpha*alpha + tail);
    bool has = norm > 0.f;
    float beta = (alpha>=0.f)? -norm : norm;
    float tau  = has ? (beta-alpha)/beta : 0.f;
    float invd = has ? 1.f/(alpha-beta) : 0.f;
    if(tid==0 && rank==0){ eout[i] = has? beta : alpha; }
    if(tid==0){ vcur[i+1]=1.f; Vs[(i+1)*SS+i]=__float2half(1.f); }
    if(tid==0 && rank==0){ Vout[(i+1)*NB+i]=1.f; }
    for(int r=i+2+tid; r<m; r+=THREADS){
      float vv = col[r]*invd; vcur[r]=vv; Vs[r*SS+i]=__float2half(vv);
      if(r>=rlo && r<rhi) Vout[r*NB+i]=vv;   // split global write
    }
    // owner of row i+1 writes its Vout slab entry
    if(tid==0 && (i+1)>=rlo && (i+1)<rhi) Vout[(i+1)*NB+i]=1.f;
    __syncthreads();
    // 3. symv p = At @ v  (SPLIT rows across cluster). RU rows/warp-iter with
    // interleaved loads -> RU*MAXV independent global loads in flight before the
    // warp reduction, hiding DRAM latency on this request-starved kernel.
    float vreg[MAXV];
    #pragma unroll
    for(int t=0;t<MAXV;++t){ int c=i+1+lane+32*t; vreg[t] = (c<m)? vcur[c] : 0.f; }
    int rstart = (rlo > i+1)? rlo : i+1;
    constexpr int RU = 4;
    int r0 = rstart+warp;
    for(; r0+(RU-1)*WARPS<rhi; r0+=RU*WARPS){
      const __half* Arp[RU];
      #pragma unroll
      for(int u=0;u<RU;++u) Arp[u] = Aht + (size_t)(r0+u*WARPS)*n;
      float dd[RU];
      #pragma unroll
      for(int u=0;u<RU;++u) dd[u]=0.f;
      #pragma unroll
      for(int t=0;t<MAXV;++t){ int c=i+1+lane+32*t; if(c<m){ float vt=vreg[t];
        #pragma unroll
        for(int u=0;u<RU;++u) dd[u]+=__half2float(Arp[u][c])*vt; } }
      #pragma unroll
      for(int u=0;u<RU;++u){ dd[u]=warp_sum(dd[u]); if(lane==0) pslab[r0+u*WARPS]=dd[u]; }
    }
    for(; r0<rhi; r0+=WARPS){
      const __half* Ar = Aht + (size_t)r0*n;
      float d=0.f;
      #pragma unroll
      for(int t=0;t<MAXV;++t){ int c=i+1+lane+32*t; if(c<m) d += __half2float(Ar[c])*vreg[t]; }
      d = warp_sum(d);
      if(lane==0) pslab[r0] = d;
    }
    __syncthreads();
    // 4. correct p (Vtv,Wtv full redundant); apply to slab rows only
    for(int k=warp; k<i; k+=WARPS){
      float sv=0.f, sw=0.f;
      for(int r=i+1+lane; r<m; r+=32){ float vr=vcur[r]; sv+=__half2float(Vs[r*SS+k])*vr; sw+=__half2float(Ws[r*SS+k])*vr; }
      sv=warp_sum(sv); sw=warp_sum(sw);
      if(lane==0){ Vtv[k]=sv; Wtv[k]=sw; }
    }
    __syncthreads();
    for(int r=rstart+tid; r<rhi; r+=THREADS){
      float pr = pslab[r];
      #pragma unroll 4
      for(int k=0;k<i;++k) pr -= __half2float(Ws[r*SS+k])*Vtv[k] + __half2float(Vs[r*SS+k])*Wtv[k];
      pslab[r] = pr*tau;
    }
    __syncthreads();
    // 5. wtv = p^T v  (cross-CTA reduction over split p)
    float pv=0.f;
    for(int r=rstart+tid; r<rhi; r+=THREADS){ pv += pslab[r]*vcur[r]; }
    pv = warp_sum(pv);
    if(lane==0) red[warp]=pv;
    __syncthreads();
    if(tid==0){ float s=0.f; for(int w=0;w<WARPS;++w) s+=red[w]; xred[0]=s; }
    cluster.sync();
    float wtv=0.f;
    #pragma unroll
    for(int rr=0;rr<C;++rr) wtv += ((float*)cluster.map_shared_rank(xred, rr))[0];
    float coef = 0.5f*tau*wtv;
    // 6. w = p - coef*v for slab; write to xw slab; exchange; store Ws(full)+Wout(slab)
    for(int r=rstart+tid; r<rhi; r+=THREADS){
      float wv = pslab[r] - coef*vcur[r];
      xw[r] = wv;
      Wout[r*NB+i] = wv;                 // split global write
    }
    // rows < rstart in this slab: their w is 0 (support rows only); ensure Ws cleared -> already 0 from init, but overwrite path below gathers all
    cluster.sync();
    // gather full w-column into replicated Ws[:,i]
    for(int rr=0;rr<C;++rr){
      const float* xwr = (const float*)cluster.map_shared_rank(xw, rr);
      int a = rr*slab, bnd = min((rr+1)*slab, m);
      int as = (a > i+1)? a : i+1;       // that rank filled [max(a,i+1), bnd)
      for(int r=as+tid; r<bnd; r+=THREADS){ Ws[r*SS+i]=__float2half(xwr[r]); }
    }
    // MUST cluster.sync after remote reads: else a fast CTA overwrites xw (next
    // column) or EXITS (last column) while a slow CTA still maps its SMEM -> fault.
    cluster.sync();
  }
}

// TMA-pipelined variant of the cluster symv (n1024 regime). The cluster latrd
// is at 1 CTA/SM (SMEM-capped) with SPARE SMEM and is latency-bound (IPC ~0.40,
// 8 warps): a warp-private cp.async.bulk prefetch pipeline over each rank's row
// slab hides the global-load latency for FREE (no occupancy cost). Phase-3 is
// the only difference vs latrd_cluster_kernel; all other phases are verbatim so
// outputs are BIT-IDENTICAL. DD = pipeline depth. TMA buffers are PREPENDED so
// the cross-CTA xw/xred offsets stay identical for DSMEM map_shared_rank.
template<int NB, int THREADS, int MAXV, int C, int DD>
__global__ __launch_bounds__(THREADS,1)
void latrd_cluster_tma_kernel(const float* __restrict__ A, const __half* __restrict__ Ah,
                  float* __restrict__ Vout, float* __restrict__ Wout,
                  float* __restrict__ dout, float* __restrict__ eout,
                  int n, int p0){
  cg::cluster_group cluster = cg::this_cluster();
  const int rank = cluster.block_rank();
  const int b = blockIdx.x / C;
  const int tid = threadIdx.x, lane = tid & 31, warp = tid >> 5;
  constexpr int WARPS = THREADS/32;
  const int m = n - p0;
  constexpr int SS = NB + 1;
  const float* At = A + (size_t)b*n*n + (size_t)p0*n + p0;
  const __half* Aht = Ah + (size_t)b*n*n + (size_t)p0*n + p0;
  Vout += (size_t)b*m*NB;  Wout += (size_t)b*m*NB;
  dout += (size_t)b*NB;    eout += (size_t)b*NB;
  const int slab = (m + C - 1)/C;
  const int rlo = rank*slab, rhi = min((rank+1)*slab, m);

  extern __shared__ char smemc[];
  __half* wbuf = (__half*)smemc;                        // [WARPS*DD*m] fp16 (16B-aligned)
  unsigned long long* mbar = (unsigned long long*)(wbuf + (size_t)WARPS*DD*m);  // [WARPS*DD]
  float* xw   = (float*)(mbar + WARPS*DD);              // [m]
  float* xred = xw + m;                                 // [4]
  __half* Vs = (__half*)(xred + 4);                     // [m*SS]
  __half* Ws = Vs + (size_t)m*SS;
  float* col = (float*)(Ws + (size_t)m*SS);
  float* vcur = col + m;
  float* pslab = vcur + m;
  float* red = pslab + m;
  float* Vtv = red + WARPS;
  float* Wtv = Vtv + NB;

  for(int idx=tid; idx<m*SS; idx+=THREADS){ Vs[idx]=__float2half(0.f); Ws[idx]=__float2half(0.f); }
  for(int idx=tid; idx<WARPS*DD; idx+=THREADS) mbar_init1((void*)&mbar[idx]);
  mbar_init_fence();
  __syncthreads();
  int wphase[DD];
  #pragma unroll
  for(int dd=0; dd<DD; ++dd) wphase[dd]=0;

  for(int i=0;i<NB;++i){
    for(int r=tid; r<m; r+=THREADS){
      // 2 INDEPENDENT accumulators (v0,v1): the correction sum over k is a serial
      // FADD chain with SMEM operands. cicc-13 reassociates + batches the SMEM loads
      // (fast-math); cicc-12.9 (grader) keeps it serial -> short_scoreboard-stalled.
      // Splitting makes the ILP toolchain-independent. (NOTES PORTABILITY hardening.)
      float v0 = At[(size_t)i*n + r], v1 = 0.f;
      int k=0;
      #pragma unroll 4
      for(; k+1<i; k+=2){
        v0 -= __half2float(Vs[r*SS+k  ])*__half2float(Ws[i*SS+k  ])
            + __half2float(Ws[r*SS+k  ])*__half2float(Vs[i*SS+k  ]);
        v1 -= __half2float(Vs[r*SS+k+1])*__half2float(Ws[i*SS+k+1])
            + __half2float(Ws[r*SS+k+1])*__half2float(Vs[i*SS+k+1]);
      }
      if(k<i) v0 -= __half2float(Vs[r*SS+k])*__half2float(Ws[i*SS+k])
                  + __half2float(Ws[r*SS+k])*__half2float(Vs[i*SS+k]);
      col[r] = v0+v1;
    }
    __syncthreads();
    if(tid==0 && rank==0) dout[i] = col[i];
    int Llen = m - i - 1;
    if(Llen <= 0){ __syncthreads(); continue; }
    float alpha = col[i+1];
    if(Llen == 1){ if(tid==0 && rank==0) eout[i] = alpha; __syncthreads(); continue; }
    float part=0.f;
    for(int r=i+2+tid; r<m; r+=THREADS){ float x=col[r]; part+=x*x; }
    part = warp_sum(part);
    if(lane==0) red[warp]=part;
    __syncthreads();
    float tail=0.f;
    #pragma unroll
    for(int w=0;w<WARPS;++w) tail += red[w];
    float norm = sqrtf(alpha*alpha + tail);
    bool has = norm > 0.f;
    float beta = (alpha>=0.f)? -norm : norm;
    float tau  = has ? (beta-alpha)/beta : 0.f;
    float invd = has ? 1.f/(alpha-beta) : 0.f;
    if(tid==0 && rank==0){ eout[i] = has? beta : alpha; }
    if(tid==0){ vcur[i+1]=1.f; Vs[(i+1)*SS+i]=__float2half(1.f); }
    if(tid==0 && rank==0){ Vout[(i+1)*NB+i]=1.f; }
    for(int r=i+2+tid; r<m; r+=THREADS){
      float vv = col[r]*invd; vcur[r]=vv; Vs[r*SS+i]=__float2half(vv);
      if(r>=rlo && r<rhi) Vout[r*NB+i]=vv;
    }
    if(tid==0 && (i+1)>=rlo && (i+1)<rhi) Vout[(i+1)*NB+i]=1.f;
    __syncthreads();
    // 3. symv p = At@v over this rank's slab rows, TMA warp-private pipeline.
    float vreg[MAXV];
    #pragma unroll
    for(int t=0;t<MAXV;++t){ int c=i+1+lane+32*t; vreg[t] = (c<m)? vcur[c] : 0.f; }
    const int rstart = (rlo > i+1)? rlo : i+1;
    const int ca = (i+1) & ~7;
    const int segbytes = (m - ca) * 2;
    __half* const wb = wbuf + (size_t)warp*DD*m;
    unsigned long long* const wmb = mbar + warp*DD;
    const int r0w = rstart + warp;
    const int nk = (r0w < rhi) ? ((rhi - r0w + WARPS - 1)/WARPS) : 0;
    auto prefetch = [&](int k, int buf){
      int rr = r0w + k*WARPS;
      if(lane==0){
        mbar_expect((void*)&wmb[buf], segbytes);
        tma_g2s(wb + (size_t)buf*m + ca, Aht + (size_t)rr*n + ca, segbytes, (void*)&wmb[buf]);
      }
    };
    #pragma unroll
    for(int dd=0; dd<DD; ++dd) if(dd<nk) prefetch(dd, dd);
    for(int k=0;k<nk;++k){
      int buf = k % DD;
      mbar_wait((void*)&wmb[buf], wphase[buf]); wphase[buf]^=1;
      int rr = r0w + k*WARPS;
      const __half* trow = wb + (size_t)buf*m;
      // 4 INDEPENDENT accumulators: force ILP=4 in the dot-product structurally so
      // BOTH front-ends emit it. A single `d += ...` chain over MAXV is a length-MAXV
      // serial FADD dependency; -use_fast_math lets cicc-13 reassociate it into
      // parallel partial sums, but cicc-12.9 (grader) keeps it serial (latency-bound,
      // -40% on the leaderboard). Splitting the accumulator makes the ILP toolchain-
      // independent. MAXV in {16,32} (both %4==0). (see NOTES PORTABILITY hardening.)
      float d0=0.f,d1=0.f,d2=0.f,d3=0.f;
      #pragma unroll
      for(int t2=0;t2<MAXV;t2+=4){ int c0=i+1+lane+32*t2;
        if(c0    <m) d0 += __half2float(trow[c0    ])*vreg[t2  ];
        if(c0+32 <m) d1 += __half2float(trow[c0+32 ])*vreg[t2+1];
        if(c0+64 <m) d2 += __half2float(trow[c0+64 ])*vreg[t2+2];
        if(c0+96 <m) d3 += __half2float(trow[c0+96 ])*vreg[t2+3]; }
      float d = (d0+d1)+(d2+d3);
      d = warp_sum(d);
      if(lane==0) pslab[rr] = d;
      int kn = k + DD;
      if(kn < nk) prefetch(kn, buf);
    }
    __syncthreads();
    for(int k=warp; k<i; k+=WARPS){
      float sv=0.f, sw=0.f;
      for(int r=i+1+lane; r<m; r+=32){ float vr=vcur[r]; sv+=__half2float(Vs[r*SS+k])*vr; sw+=__half2float(Ws[r*SS+k])*vr; }
      sv=warp_sum(sv); sw=warp_sum(sw);
      if(lane==0){ Vtv[k]=sv; Wtv[k]=sw; }
    }
    __syncthreads();
    for(int r=rstart+tid; r<rhi; r+=THREADS){
      // 2 INDEPENDENT accumulators: same serial-chain / short-scoreboard fix as
      // phase 1 (see comment there). Toolchain-independent ILP for the k-sum.
      float p0=pslab[r], p1=0.f;
      int k=0;
      #pragma unroll 4
      for(; k+1<i; k+=2){
        p0 -= __half2float(Ws[r*SS+k  ])*Vtv[k  ] + __half2float(Vs[r*SS+k  ])*Wtv[k  ];
        p1 -= __half2float(Ws[r*SS+k+1])*Vtv[k+1] + __half2float(Vs[r*SS+k+1])*Wtv[k+1];
      }
      if(k<i) p0 -= __half2float(Ws[r*SS+k])*Vtv[k] + __half2float(Vs[r*SS+k])*Wtv[k];
      pslab[r] = (p0+p1)*tau;
    }
    __syncthreads();
    float pv=0.f;
    for(int r=rstart+tid; r<rhi; r+=THREADS){ pv += pslab[r]*vcur[r]; }
    pv = warp_sum(pv);
    if(lane==0) red[warp]=pv;
    __syncthreads();
    if(tid==0){ float s=0.f; for(int w=0;w<WARPS;++w) s+=red[w]; xred[0]=s; }
    cluster.sync();
    float wtv=0.f;
    #pragma unroll
    for(int rr=0;rr<C;++rr) wtv += ((float*)cluster.map_shared_rank(xred, rr))[0];
    float coef = 0.5f*tau*wtv;
    for(int r=rstart+tid; r<rhi; r+=THREADS){
      float wv = pslab[r] - coef*vcur[r];
      xw[r] = wv;
      Wout[r*NB+i] = wv;
    }
    cluster.sync();
    for(int rr=0;rr<C;++rr){
      const float* xwr = (const float*)cluster.map_shared_rank(xw, rr);
      int a = rr*slab, bnd = min((rr+1)*slab, m);
      int as = (a > i+1)? a : i+1;
      for(int r=as+tid; r<bnd; r+=THREADS){ Ws[r*SS+i]=__float2half(xwr[r]); }
    }
    cluster.sync();
  }
}

// 2D-TENSOR-MAP cluster symv (n2048 regime). The chunked-TMA kill (1D, CK=512)
// hid the L2 latency but paid 16x the mbarrier/issue overhead — each warp issued
// its own 1-row chunk (1024 issues/column, ~512 FMAs/arrival, below the ~1-2K
// amortization floor). Here ONE thread (tid==0) issues a CTA-COLLECTIVE 2D tile
// spanning TR contiguous rows x TC cols via cp.async.bulk.tensor.2d: one arrival
// delivers TR*TC elements (= TR*TC FMAs/arrival, e.g. 32*256=8192, 16x over the
// floor) with 16x fewer issues + 16x fewer mbarrier objects. The whole CTA reads
// the shared tile; a per-tile __syncthreads gates buffer reuse. SMEM for the tile
// pipeline is DD*TR*TC*2 (shared across warps, NOT x16 as full-row would be), so
// it FITS beside the 172KB nb16 base (48KB at TR=32,TC=256,DD=3). Phase-3 only
// differs vs latrd_cluster_kernel; other phases verbatim -> BIT-IDENTICAL (same
// fp16 bits, per-lane column order preserved, masked cols contribute exact +0).
// NRW = TR/WARPS rows owned per warp per super-block (requires TR % WARPS == 0).
template<int NB, int THREADS, int C, int DD, int TR, int TC>
__global__ __launch_bounds__(THREADS,1)
void latrd_cluster_2dtma_kernel(const float* __restrict__ A, const __half* __restrict__ Ah,
                  float* __restrict__ Vout, float* __restrict__ Wout,
                  float* __restrict__ dout, float* __restrict__ eout,
                  int n, int p0, const __grid_constant__ CUtensorMap tmap){
  cg::cluster_group cluster = cg::this_cluster();
  const int rank = cluster.block_rank();
  const int b = blockIdx.x / C;
  const int tid = threadIdx.x, lane = tid & 31, warp = tid >> 5;
  constexpr int WARPS = THREADS/32;
  constexpr int NRW = TR/WARPS;
  const int m = n - p0;
  constexpr int SS = NB + 1;
  const float* At = A + (size_t)b*n*n + (size_t)p0*n + p0;
  Vout += (size_t)b*m*NB;  Wout += (size_t)b*m*NB;
  dout += (size_t)b*NB;    eout += (size_t)b*NB;
  const int slab = (m + C - 1)/C;
  const int rlo = rank*slab, rhi = min((rank+1)*slab, m);

  extern __shared__ char smemc[];
  __half* wbuf = (__half*)smemc;                        // [DD*TR*TC] fp16 (128B-aligned tile pipeline)
  unsigned long long* mbar = (unsigned long long*)(wbuf + (size_t)DD*TR*TC);  // full[DD]
  float* xw   = (float*)(mbar + DD);                    // [m]
  float* xred = xw + m;                                 // [4]
  __half* Vs = (__half*)(xred + 4);                     // [m*SS]
  __half* Ws = Vs + (size_t)m*SS;
  float* col = (float*)(Ws + (size_t)m*SS);
  float* vcur = col + m;
  float* pslab = vcur + m;
  float* red = pslab + m;
  float* Vtv = red + WARPS;
  float* Wtv = Vtv + NB;

  for(int idx=tid; idx<m*SS; idx+=THREADS){ Vs[idx]=__float2half(0.f); Ws[idx]=__float2half(0.f); }
  for(int idx=tid; idx<DD; idx+=THREADS) mbar_init1((void*)&mbar[idx]);
  mbar_init_fence();
  __syncthreads();
  int wphase[DD];
  #pragma unroll
  for(int dd=0; dd<DD; ++dd) wphase[dd]=0;

  for(int i=0;i<NB;++i){
    for(int r=tid; r<m; r+=THREADS){
      float v = At[(size_t)i*n + r];
      #pragma unroll 4
      for(int k=0;k<i;++k){
        v -= __half2float(Vs[r*SS+k])*__half2float(Ws[i*SS+k])
           + __half2float(Ws[r*SS+k])*__half2float(Vs[i*SS+k]);
      }
      col[r] = v;
    }
    __syncthreads();
    if(tid==0 && rank==0) dout[i] = col[i];
    int Llen = m - i - 1;
    if(Llen <= 0){ __syncthreads(); continue; }
    float alpha = col[i+1];
    if(Llen == 1){ if(tid==0 && rank==0) eout[i] = alpha; __syncthreads(); continue; }
    float part=0.f;
    for(int r=i+2+tid; r<m; r+=THREADS){ float x=col[r]; part+=x*x; }
    part = warp_sum(part);
    if(lane==0) red[warp]=part;
    __syncthreads();
    float tail=0.f;
    #pragma unroll
    for(int w=0;w<WARPS;++w) tail += red[w];
    float norm = sqrtf(alpha*alpha + tail);
    bool has = norm > 0.f;
    float beta = (alpha>=0.f)? -norm : norm;
    float tau  = has ? (beta-alpha)/beta : 0.f;
    float invd = has ? 1.f/(alpha-beta) : 0.f;
    if(tid==0 && rank==0){ eout[i] = has? beta : alpha; }
    if(tid==0){ vcur[i+1]=1.f; Vs[(i+1)*SS+i]=__float2half(1.f); }
    if(tid==0 && rank==0){ Vout[(i+1)*NB+i]=1.f; }
    for(int r=i+2+tid; r<m; r+=THREADS){
      float vv = col[r]*invd; vcur[r]=vv; Vs[r*SS+i]=__float2half(vv);
      if(r>=rlo && r<rhi) Vout[r*NB+i]=vv;
    }
    if(tid==0 && (i+1)>=rlo && (i+1)<rhi) Vout[(i+1)*NB+i]=1.f;
    __syncthreads();
    // 3. symv p = At@v over this rank's slab rows, 2D-TMA CTA-collective pipeline.
    const int rstart = (rlo > i+1)? rlo : i+1;
    const int ca = (i+1) & ~7;                           // 16B-aligned trailing col start
    const int ncol = m - ca;
    const int nchunks = (ncol + TC - 1)/TC;
    const int nrows = rhi - rstart;
    if(nrows > 0 && nchunks > 0){
      const int nsb = (nrows + TR - 1)/TR;
      const int TOT = nsb * nchunks;
      const int rowbase_abs = b*n + p0 + rstart;         // abs tensor row of local row 0
      const int colbase_abs = p0 + ca;                   // abs tensor col of chunk 0
      auto issue = [&](int j, int buf){
        if(tid==0){
          int sb = j / nchunks, ch = j - sb*nchunks;
          mbar_expect((void*)&mbar[buf], TR*TC*2);
          tma_2d(wbuf + (size_t)buf*TR*TC, &tmap,
                 colbase_abs + ch*TC, rowbase_abs + sb*TR, (void*)&mbar[buf]);
        }
      };
      #pragma unroll
      for(int dd=0; dd<DD; ++dd) if(dd<TOT) issue(dd, dd);
      float acc[NRW];
      #pragma unroll
      for(int q=0;q<NRW;++q) acc[q]=0.f;
      for(int j=0;j<TOT;++j){
        int buf = j%DD;
        mbar_wait((void*)&mbar[buf], wphase[buf]); wphase[buf]^=1;
        int sb = j/nchunks, ch = j - sb*nchunks;
        int colbase = ca + ch*TC;
        float vchunk[TC/32];
        #pragma unroll
        for(int t=0;t<TC/32;++t){ int gc=colbase+lane+32*t; vchunk[t]=(gc>=i+1 && gc<m)? vcur[gc]:0.f; }
        const __half* tile = wbuf + (size_t)buf*TR*TC;
        #pragma unroll
        for(int q=0;q<NRW;++q){
          const __half* trow = tile + (size_t)(warp + q*WARPS)*TC;
          float dloc=0.f;
          #pragma unroll
          for(int t=0;t<TC/32;++t){ int cl=lane+32*t; dloc += __half2float(trow[cl])*vchunk[t]; }
          acc[q] += dloc;
        }
        if(ch == nchunks-1){
          #pragma unroll
          for(int q=0;q<NRW;++q){
            int gr = rstart + sb*TR + warp + q*WARPS;
            float dsum = warp_sum(acc[q]);
            if(lane==0 && gr < rhi) pslab[gr] = dsum;
            acc[q]=0.f;
          }
        }
        __syncthreads();                                 // all warps done reading buf -> safe to reissue
        int jn = j+DD; if(jn<TOT) issue(jn, buf);
      }
    }
    __syncthreads();
    for(int k=warp; k<i; k+=WARPS){
      float sv=0.f, sw=0.f;
      for(int r=i+1+lane; r<m; r+=32){ float vr=vcur[r]; sv+=__half2float(Vs[r*SS+k])*vr; sw+=__half2float(Ws[r*SS+k])*vr; }
      sv=warp_sum(sv); sw=warp_sum(sw);
      if(lane==0){ Vtv[k]=sv; Wtv[k]=sw; }
    }
    __syncthreads();
    for(int r=rstart+tid; r<rhi; r+=THREADS){
      float pr = pslab[r];
      #pragma unroll 4
      for(int k=0;k<i;++k) pr -= __half2float(Ws[r*SS+k])*Vtv[k] + __half2float(Vs[r*SS+k])*Wtv[k];
      pslab[r] = pr*tau;
    }
    __syncthreads();
    float pv=0.f;
    for(int r=rstart+tid; r<rhi; r+=THREADS){ pv += pslab[r]*vcur[r]; }
    pv = warp_sum(pv);
    if(lane==0) red[warp]=pv;
    __syncthreads();
    if(tid==0){ float s=0.f; for(int w=0;w<WARPS;++w) s+=red[w]; xred[0]=s; }
    cluster.sync();
    float wtv=0.f;
    #pragma unroll
    for(int rr=0;rr<C;++rr) wtv += ((float*)cluster.map_shared_rank(xred, rr))[0];
    float coef = 0.5f*tau*wtv;
    for(int r=rstart+tid; r<rhi; r+=THREADS){
      float wv = pslab[r] - coef*vcur[r];
      xw[r] = wv;
      Wout[r*NB+i] = wv;
    }
    cluster.sync();
    for(int rr=0;rr<C;++rr){
      const float* xwr = (const float*)cluster.map_shared_rank(xw, rr);
      int a = rr*slab, bnd = min((rr+1)*slab, m);
      int as = (a > i+1)? a : i+1;
      for(int r=as+tid; r<bnd; r+=THREADS){ Ws[r*SS+i]=__float2half(xwr[r]); }
    }
    cluster.sync();
  }
}

void launch_latrd_cluster(const float* A, const __half* Ah, float* V, float* W, float* d, float* e,
                          int b, int n, int p0, int C, int nb){
  int m = n - p0;
  constexpr int THREADS=256;
  int maxv = (m + 31) / 32;
  constexpr int WARPS = THREADS/32;
  constexpr int TMA_D = 3;                          // pipeline depth (D=3 == D=4, use less SMEM)
  int blkthreads = THREADS;                         // blockDim (scalar path may override -> more MLP)
  auto launch_cfg = [&](auto kfn, size_t smem){
    cudaFuncSetAttribute(kfn, cudaFuncAttributeMaxDynamicSharedMemorySize, smem);
    if(C > 8) cudaFuncSetAttribute(kfn, cudaFuncAttributeNonPortableClusterSizeAllowed, 1);
    cudaLaunchConfig_t cfg = {};
    cfg.gridDim = dim3(C*b); cfg.blockDim = dim3(blkthreads); cfg.dynamicSmemBytes = smem;
    cudaLaunchAttribute attr[1];
    attr[0].id = cudaLaunchAttributeClusterDimension;
    attr[0].val.clusterDim.x = C; attr[0].val.clusterDim.y=1; attr[0].val.clusterDim.z=1;
    cfg.attrs = attr; cfg.numAttrs = 1;
    cudaLaunchKernelEx(&cfg, kfn, A, Ah, V, W, d, e, n, p0);
  };
  auto smem_base = [&](int NBv, int warps){
    return (size_t)m*sizeof(float) + 4*sizeof(float)
         + (size_t)2*m*(NBv+1)*sizeof(__half)
         + ((size_t)3*m + warps + 2*NBv)*sizeof(float);
  };
  size_t tma_extra = (size_t)WARPS*TMA_D*m*sizeof(__half) + (size_t)WARPS*TMA_D*sizeof(unsigned long long);
  // Scalar cluster symv is memory-LATENCY bound (IPC ~0.40, 8 warps at 256 threads,
  // 1 CTA/SM SMEM-capped). Widening blockDim to 512 (16 warps) doubles the in-flight
  // loads for FREE (no extra SMEM, still 1 CTA/SM) -> ~2x MLP on the request-starved
  // symv. Used for the n2048 nb=16 large-m path where the O(m^2) symv dominates.
  constexpr int SCALTHR = 512;
  constexpr int SCALW = SCALTHR/32;
  auto do_launch = [&](auto kfn, int NBv){ blkthreads = SCALTHR; launch_cfg(kfn, smem_base(NBv, SCALW)); };
  // WIN (default ON): TMA@512 (16 warps) + DD=2 full-row buffers BEATS scalar@512
  // on n1024 by -4.3% (locked b60: dense 45.55->43.56, mixed 46.46->44.49, nearrank
  // 45.65->43.68, every round). The prior "TMA@512 doesn't fit / scalar supersedes"
  // verdict was a DD=3 artifact: DD=3 full-row = 16*3*m*2B = 96KB overflows, but
  // DD=2 = 64KB FITS (152KB base -> ~213KB < 231KB cap). ncu confirms the mechanism:
  // the async prefetch converts the L2-load-latency wall (long_scoreboard 30354->6108,
  // -80%) into SMEM reads, IPC 1.44->2.32. BIT-IDENTICAL to scalar (same fp16 bits,
  // only the fetch path changes) -> gates unchanged. n2048 (nb=16) falls through
  // to the scalar path below (full-row TMA buffers don't fit at m=2048).
  constexpr int TMA_THR = 512;
  constexpr int TMA_W = TMA_THR/32;
  constexpr int TMA_DD = 2;
  size_t tma_extra512 = (size_t)TMA_W*TMA_DD*m*sizeof(__half) + (size_t)TMA_W*TMA_DD*sizeof(unsigned long long);
  bool tma_ok = (nb==32) && (maxv<=32) && (C<=4) && (smem_base(32, TMA_W) + tma_extra512 <= 231000);
  if(tma_ok){
    size_t sm = smem_base(32, TMA_W) + tma_extra512;
    blkthreads = TMA_THR;
    if(maxv<=16){
      if(C==2) launch_cfg(latrd_cluster_tma_kernel<32,TMA_THR,16,2,TMA_DD>, sm);
      else if(C==3) launch_cfg(latrd_cluster_tma_kernel<32,TMA_THR,16,3,TMA_DD>, sm);
      else launch_cfg(latrd_cluster_tma_kernel<32,TMA_THR,16,4,TMA_DD>, sm);
    } else {
      if(C==2) launch_cfg(latrd_cluster_tma_kernel<32,TMA_THR,32,2,TMA_DD>, sm);
      else if(C==3) launch_cfg(latrd_cluster_tma_kernel<32,TMA_THR,32,3,TMA_DD>, sm);
      else launch_cfg(latrd_cluster_tma_kernel<32,TMA_THR,32,4,TMA_DD>, sm);
    }
    return;
  }
  // 2D-TENSOR-MAP path for the n2048 nb=16 large-m symv (C==8). The scalar cluster
  // kernel is register-saturated at m=2048 (vreg[64] eats the file, 112 regs @512thr)
  // AND cicc-12.9 serializes its symv accumulator -> the grader (12.9) runs it at
  // 80.6ms vs 56.5 on cu13 (a codegen ILP loss the source-hardening can't reach here
  // without spilling). The 2D-TMA moves the fp16 row loads to a hardware async tile
  // pipeline (long_scoreboard -87%): codegen-robust (the async copy is an explicit
  // hw instruction, not cicc-scheduled MLP), so it should NOT carry the 12.9 penalty.
  // Toolchain-conditional by default (grader-12.9 uses 2D-TMA; our cu13 keeps the
  // faster scalar) — EIGH_USE_2DTMA / EIGH_NO_2DTMA override for A/B on either stack.
  {
    bool want_2dtma;
  #if defined(__CUDACC_VER_MAJOR__) && (__CUDACC_VER_MAJOR__ < 13)
    want_2dtma = (getenv("EIGH_NO_2DTMA") == nullptr);   // grader (12.9): default ON
  #else
    want_2dtma = (getenv("EIGH_USE_2DTMA") != nullptr);  // dev (cu13): default OFF (scalar faster)
  #endif
    if(want_2dtma && nb==16 && C==8 && maxv<=64){
      constexpr int T_THR = 512, T_W = T_THR/32;
      int t_tr=48, t_tc=256, t_dd=2;
      if(const char* cfg = getenv("EIGH_2DTMA")){ sscanf(cfg, "%d,%d,%d", &t_tr, &t_tc, &t_dd); }
      static PFN_cuTensorMapEncodeTiled_v12000 encode = nullptr;
      if(!encode) cudaGetDriverEntryPoint("cuTensorMapEncodeTiled", (void**)&encode, cudaEnableDefault, nullptr);
      CUtensorMap tmap;
      uint64_t gdim[2] = { (uint64_t)n, (uint64_t)b*(uint64_t)n };
      uint64_t gstr[1] = { (uint64_t)n*sizeof(__half) };
      uint32_t estr[2] = { 1, 1 };
      auto launch2d = [&](auto kfn, int TRv, int TCv, int DDv){
        uint32_t bdim[2] = { (uint32_t)TCv, (uint32_t)TRv };
        encode(&tmap, CU_TENSOR_MAP_DATA_TYPE_FLOAT16, 2, (void*)Ah, gdim, gstr, bdim, estr,
               CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_NONE,
               CU_TENSOR_MAP_L2_PROMOTION_L2_128B, CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
        size_t tile_extra = (size_t)DDv*TRv*TCv*sizeof(__half) + (size_t)DDv*sizeof(unsigned long long);
        size_t sm = smem_base(16, T_W) + tile_extra;
        cudaFuncSetAttribute(kfn, cudaFuncAttributeMaxDynamicSharedMemorySize, sm);
        cudaLaunchConfig_t cf = {};
        cf.gridDim = dim3(C*b); cf.blockDim = dim3(T_THR); cf.dynamicSmemBytes = sm;
        cudaLaunchAttribute at[1];
        at[0].id = cudaLaunchAttributeClusterDimension;
        at[0].val.clusterDim.x = C; at[0].val.clusterDim.y=1; at[0].val.clusterDim.z=1;
        cf.attrs = at; cf.numAttrs = 1;
        cudaLaunchKernelEx(&cf, kfn, A, Ah, V, W, d, e, n, p0, tmap);
      };
      #define L2D(TRv,TCv,DDv) launch2d(latrd_cluster_2dtma_kernel<16,T_THR,8,DDv,TRv,TCv>, TRv, TCv, DDv)
      #define TRY(TRv,TCv,DDv) if(t_tr==TRv&&t_tc==TCv&&t_dd==DDv){ L2D(TRv,TCv,DDv); return; }
      TRY(48,256,2) TRY(32,256,2) TRY(64,256,2) TRY(32,256,3) TRY(48,256,3)
      TRY(80,256,2) TRY(48,128,2) TRY(64,128,3) TRY(96,256,2) TRY(32,512,2)
      #undef L2D
      #undef TRY
      // unlisted geometry -> fall through to scalar
    }
  }
  // NB=16 path (large m, e.g. n2048): halves the fp16 panel SMEM so the first
  // panel (m up to 2048) fits under the 228 KB cap, and adds MAXV=64 so the symv
  // covers reflectors longer than 1024 columns without silently dropping any.
  #define DISP16(CC) do{ if(maxv<=16) do_launch(latrd_cluster_kernel<16,SCALTHR,16,CC>,16); \
                         else if(maxv<=32) do_launch(latrd_cluster_kernel<16,SCALTHR,32,CC>,16); \
                         else do_launch(latrd_cluster_kernel<16,SCALTHR,64,CC>,16); }while(0)
  #define DISP32(CC) do{ if(maxv<=16) do_launch(latrd_cluster_kernel<32,SCALTHR,16,CC>,32); \
                         else do_launch(latrd_cluster_kernel<32,SCALTHR,32,CC>,32); }while(0)
  #define DISP(CC) do{ if(nb==16) DISP16(CC); else DISP32(CC); }while(0)
  if(C==2) DISP(2); else if(C==3) DISP(3); else if(C==4) DISP(4);
  else if(C==5) DISP(5); else if(C==6) DISP(6); else if(C==7) DISP(7); else if(C==8) DISP(8);
  else if(C==10) DISP16(10); else if(C==12) DISP16(12); else if(C==16) DISP16(16);
  #undef DISP
  #undef DISP16
  #undef DISP32
}

}  // namespace os1cl
"""

BUILD_DIR = Path(__file__).resolve().parent / ".build"
BUILD_DIR.mkdir(exist_ok=True)
CU13_ROOT = Path(torch.__file__).resolve().parent.parent / "nvidia" / "cu13"
CU13_LIB = CU13_ROOT / "lib"

# Dev-box only (env set by gpu_run.sh, never on the grader): the canonical local
# build uses the grader's CUDA-12.9 nvcc with cu13 torch — the grader's own
# pairing — so skip torch's major-version guard for that combination.
import os as _os_tc
if _os_tc.environ.get("EIGH_CU129") == "1":
    import torch.utils.cpp_extension as _cpp_ext_mod
    _cpp_ext_mod._check_cuda_version = lambda *a, **k: None

load_inline(
    "eigh_ext",
    cpp_sources=CPP_SRC,
    cuda_sources=CUDA_SRC,
    is_python_module=False,
    no_implicit_headers=True,
    extra_include_paths=[str(CU13_ROOT / "include")],
    extra_cflags=["-O3", "-std=c++17"],
    extra_cuda_cflags=["-O3", "-std=c++17", "-lineinfo", "-use_fast_math"],
    extra_ldflags=[
        f"-L{CU13_LIB}",
        f"-Wl,-rpath,{CU13_LIB}",
        "-l:libcublas.so.13",
        "-l:libcublasLt.so.13",
    ],
    build_directory=str(BUILD_DIR),
)

# Per-shape solver for the small shapes (n -> Jacobi variant). n=176 is routed
# separately (padded block-Jacobi) in custom_kernel.
SMALL_STRATEGY = {
    32: "jacobi",
    352: "block",
}

# Convergence tolerance^2 per shape (off-diagonal Frobenius^2 relative to
# ||A||_F^2). Tuned against the gates with >10x margin over multiple seeds.
JACOBI_TOL2 = {
    32: 1e-9,
    # 352: quadratic convergence jumps ~1.3e-6 -> ~1.2e-8 between sweeps 7 and
    # 8; 3e-6 stops at 7 sweeps with measured gate margins >= 5x (multi-seed).
    352: 2e-6,
}


def _jacobi_smem(data: torch.Tensor, tol2: float = 1e-10) -> output_t:
    batch, n, _ = data.shape
    q = torch.empty_like(data)
    l = data.new_empty(batch, n)
    torch.ops.eigh_ops.jacobi_smem(data, q, l, tol2)
    return q, l


def _block_jacobi(data: torch.Tensor, tol2: float = 1e-7) -> output_t:
    batch, n, _ = data.shape
    q = torch.empty_like(data)
    l = data.new_empty(batch, n)
    a_work = torch.empty_like(data)
    qt_work = torch.empty_like(data)
    torch.ops.eigh_ops.block_jacobi(data, q, l, a_work, qt_work, tol2)
    return q, l


def _block_jacobi_pad(data: torch.Tensor, npad: int, tol2: float = 1e-7) -> output_t:
    """Embed an n x n symmetric problem into npad x npad (npad divisible by 32,
    npad/16 even) so it runs through the A-parallel block-Jacobi cluster kernel.

    The pad block (rows/cols n..npad) is left at ZERO (diagonal and off), which is
    an EXACT invariant subspace under two-sided Jacobi: A[real, pad] starts at 0
    and every rotation keeps it 0 (c*0 - s*0), so the pad eigenvectors stay inside
    span(e_n..e_npad) and the real eigenvectors stay inside span(e_0..e_n) with
    exactly zero mass on the pad rows. The pad also contributes 0 to ||A||_F, so
    the convergence criterion is identical to the un-padded problem. Real columns
    are recovered by their eigenvector support in the first n rows.
    """
    batch, n, _ = data.shape
    ap = data.new_zeros(batch, npad, npad)
    ap[:, :n, :n] = data
    q = torch.empty_like(ap)
    l = ap.new_empty(batch, npad)
    a_work = torch.empty_like(ap)
    qt_work = torch.empty_like(ap)
    torch.ops.eigh_ops.block_jacobi(ap, q, l, a_work, qt_work, tol2)
    # Select the n columns whose mass lives in the real rows (0..n), keeping
    # ascending-eigenvalue order (l is already sorted; stable argsort preserves it).
    mass_real = (q[:, :n, :] ** 2).sum(dim=1)          # (batch, npad)
    order = torch.argsort(mass_real, dim=1, descending=True, stable=True)
    sel = order[:, :n]                                  # (batch, n) real col indices
    sel, _ = torch.sort(sel, dim=1)                     # restore ascending eigenvalue order
    l_real = torch.gather(l, 1, sel)                    # (batch, n)
    q_real = torch.gather(q[:, :n, :], 2, sel.unsqueeze(1).expand(-1, n, -1))
    return q_real.contiguous(), l_real.contiguous()


# ---------------------------------------------------------------------------
# n=512: blocked two-sided Jacobi with dynamic (greedy) block-pair ordering.
#
# A (prescaled to unit max-magnitude, fp16) is partitioned into a 16x16 grid
# of 32x32 blocks. Each round: (1) per-matrix block Frobenius norms -> greedy
# maximum-weight matching picks 8 disjoint block pairs (dynamic ordering,
# ~1.5-2.3x fewer rounds than round-robin); (2) a seated warp-per-pivot CUDA
# kernel solves the 8 batched 64x64 block-pivot eigenproblems on-chip
# (eigenvalues sorted ascending -> van Kempen sorting effect); (3) Triton
# tensor-core kernels apply V two-sided to A (tile-fused T <- V1^T T V2, in
# place) and to the accumulated Q, all fp16. The loop stops on a norm-based
# convergence criterion (max over batch). Low-precision drift is then repaired
# by Newton-Schulz re-orthonormalization of Q (TF32 step + bf16x9 step), an
# accurate bf16x9 refresh A_r = (Q^T A Q)/amax against the ORIGINAL input, and
# a short accurate phase (fp16 pivot solves + per-V fp32 NS-orth + tf32x3
# updates). Finally L = diag(A_r)*amax, argsort ascending, gather Q columns.
# ---------------------------------------------------------------------------

_N512 = 512
_TB512 = 64        # pivot size (2 blocks)
_BLK512 = _TB512 // 2
_P512 = _N512 // _BLK512         # column blocks (16)
_P2_512 = _P512 // 2             # pairs per round (8)

_LO_INNER512 = 2
_LO_STOP512 = 3.0e-3             # rel off-diagonal Frobenius, per matrix
_LO_MAX_ROUNDS512 = 96
_HI_INNER512 = 2
_HI_STOP512 = 4.0e-4
_HI_MAX_ROUNDS512 = 32
_CHECK_EVERY512 = 4


@triton.jit
def _block_norms(a_ptr, out_ptr, active_ptr,
                 N: tl.constexpr, BLK: tl.constexpr, P: tl.constexpr):
    """out[b, br, bc] = ||A[b, br-block, bc-block]||_F^2 (fp32)."""
    bc = tl.program_id(0)
    br = tl.program_id(1)
    b = tl.program_id(2)
    if tl.load(active_ptr + b) == 0:
        return
    offs = tl.arange(0, BLK)
    rows = br * BLK + offs
    cols = bc * BLK + offs
    a = tl.load(a_ptr + b.to(tl.int64) * N * N + rows[:, None] * N + cols[None, :])
    v = a.to(tl.float32)
    total = tl.sum(tl.sum(v * v, axis=1), axis=0)
    tl.store(out_ptr + (b * P + br) * P + bc, total)


@triton.jit
def _greedy_pairs(norms_ptr, pairs_ptr, stats_ptr, active_ptr, stop2,
                  P: tl.constexpr, P2: tl.constexpr):
    """Greedy max-weight matching on the off-diagonal block-norm matrix.

    Emits per-matrix (off2_sum, fro2_sum) and freezes converged matrices:
    active[b] <- off2 > stop2 * fro2 (all later kernels skip frozen b)."""
    b = tl.program_id(0)
    if tl.load(active_ptr + b) == 0:
        return
    offs = tl.arange(0, P)
    w = tl.load(norms_ptr + (b * P + offs[:, None]) * P + offs[None, :])
    fro2 = tl.sum(tl.sum(w, axis=1), axis=0)
    diag = offs[:, None] == offs[None, :]
    off2 = tl.sum(tl.sum(tl.where(diag, 0.0, w), axis=1), axis=0)
    tl.store(stats_ptr + b * 2, off2)
    tl.store(stats_ptr + b * 2 + 1, fro2)
    if off2 <= stop2 * fro2:
        tl.store(active_ptr + b, 0)
        return
    avail = tl.full((P,), 1, tl.int32)
    for k in tl.static_range(P2):
        # +1.0 keeps every eligible entry above the -1 sentinel even for an
        # all-zero norm matrix, so the pick is always a valid disjoint pair.
        ok = (avail[:, None] > 0) & (avail[None, :] > 0) & (~diag)
        wm = tl.where(ok, w + 1.0, -1.0)
        # two-stage argmax (tl.reshape+argmax mis-indexes at 16x16)
        rowmax = tl.max(wm, axis=1)
        i = tl.argmax(rowmax, axis=0).to(tl.int32)
        rowvals = tl.max(tl.where(offs[:, None] == i, wm, -2.0), axis=0)
        j = tl.argmax(rowvals, axis=0).to(tl.int32)
        lo = tl.minimum(i, j)
        hi = tl.maximum(i, j)
        tl.store(pairs_ptr + (b * P2 + k) * 2, lo)
        tl.store(pairs_ptr + (b * P2 + k) * 2 + 1, hi)
        avail = tl.where((offs == i) | (offs == j), 0, avail)


@triton.jit
def _jacobi_tile_update(
    a_ptr, v_ptr, pairs_ptr, active_ptr,
    N: tl.constexpr, TB: tl.constexpr, P2: tl.constexpr, PREC: tl.constexpr,
):
    """In-place two-sided block update: tile(k1,k2) <- V1^T @ tile @ V2.

    One block owns block-row k1 and walks all P2 column-pivots k2. Each output
    tile depends only on its own input (V is block-diagonal), so the P2 tiles
    are independent -> the loop exposes P2-way ILP that hides the global-load
    latency (this update was latency-bound at ~20% occupancy) while V1 is loaded
    once and reused across the whole row-band."""
    k1 = tl.program_id(0)
    b = tl.program_id(1)
    BLK: tl.constexpr = TB // 2
    if tl.load(active_ptr + b) == 0:
        return

    pbase = pairs_ptr + b * P2 * 2
    bp1 = tl.load(pbase + k1 * 2)
    bq1 = tl.load(pbase + k1 * 2 + 1)

    offs = tl.arange(0, TB)
    rows = tl.where(offs < BLK, bp1 * BLK + offs, bq1 * BLK + offs - BLK)

    a_base = a_ptr + b.to(tl.int64) * N * N
    v_base = v_ptr + (b.to(tl.int64) * P2) * TB * TB
    v1 = tl.load(v_base + k1 * TB * TB + offs[:, None] * TB + offs[None, :])
    dt = a_ptr.dtype.element_ty

    if PREC == "fp16x3":
        v1t = tl.trans(v1)
        v1h = v1t.to(tl.float16)
        v1l = (v1t - v1h.to(tl.float32)).to(tl.float16)
    elif PREC == "fp16":
        v1t = tl.trans(v1)
    else:
        v1t = tl.trans(v1).to(tl.float32)

    for k2 in range(P2):
        bp2 = tl.load(pbase + k2 * 2)
        bq2 = tl.load(pbase + k2 * 2 + 1)
        cols = tl.where(offs < BLK, bp2 * BLK + offs, bq2 * BLK + offs - BLK)
        t_ptrs = a_base + rows[:, None] * N + cols[None, :]
        t = tl.load(t_ptrs)
        v2 = tl.load(v_base + k2 * TB * TB + offs[:, None] * TB + offs[None, :])
        if PREC == "fp16":
            tmp = tl.dot(v1t, t)
            out = tl.dot(tmp.to(dt), v2)
        elif PREC == "fp16x3":
            th = t.to(tl.float16)
            tl_ = (t - th.to(tl.float32)).to(tl.float16)
            tmp = tl.dot(v1h, th) + tl.dot(v1h, tl_) + tl.dot(v1l, th)
            tmph = tmp.to(tl.float16)
            tmpl = (tmp - tmph.to(tl.float32)).to(tl.float16)
            v2h = v2.to(tl.float16)
            v2l = (v2 - v2h.to(tl.float32)).to(tl.float16)
            out = tl.dot(tmph, v2h) + tl.dot(tmph, v2l) + tl.dot(tmpl, v2h)
        else:
            tmp = tl.dot(v1t, t.to(tl.float32), input_precision=PREC)
            out = tl.dot(tmp, v2.to(tl.float32), input_precision=PREC)
        tl.store(t_ptrs, out.to(dt))


@triton.jit
def _jacobi_q_update(
    q_ptr, v_ptr, pairs_ptr, active_ptr,
    N: tl.constexpr, TB: tl.constexpr, P2: tl.constexpr, PREC: tl.constexpr,
    RBM: tl.constexpr,
):
    """In-place eigenvector update: Q[:, cols(k2)] <- Q[:, cols(k2)] @ V2.

    Each block processes a TALL RBM-row strip (RBM >> TB) against one V2 tile.
    The bigger m dimension gives the tiny 64-wide GEMM enough independent work
    to hide global-load latency (the update was latency-bound at ~20% occupancy
    with a 64-row tile) while V2 is loaded once and reused across all RBM rows.
    (A k2-loop variant with disjoint column blocks measured slower — 183 vs
    180 ms — the scattered per-k2 column access outweighs the extra ILP.)"""
    k2 = tl.program_id(0)
    rt = tl.program_id(1)
    b = tl.program_id(2)
    BLK: tl.constexpr = TB // 2
    if tl.load(active_ptr + b) == 0:
        return

    pbase = pairs_ptr + b * P2 * 2
    bp2 = tl.load(pbase + k2 * 2)
    bq2 = tl.load(pbase + k2 * 2 + 1)

    coffs = tl.arange(0, TB)
    roffs = tl.arange(0, RBM)
    rows = rt * RBM + roffs
    cols = tl.where(coffs < BLK, bp2 * BLK + coffs, bq2 * BLK + coffs - BLK)

    q_base = q_ptr + b.to(tl.int64) * N * N
    q_ptrs = q_base + rows[:, None] * N + cols[None, :]
    q = tl.load(q_ptrs)

    v_base = v_ptr + (b.to(tl.int64) * P2) * TB * TB
    v2 = tl.load(v_base + k2 * TB * TB + coffs[:, None] * TB + coffs[None, :])

    dt = q_ptr.dtype.element_ty
    if PREC == "fp16":
        out = tl.dot(q, v2)
    elif PREC == "fp16x3":
        qh = q.to(tl.float16)
        ql = (q - qh.to(tl.float32)).to(tl.float16)
        v2h = v2.to(tl.float16)
        v2l = (v2 - v2h.to(tl.float32)).to(tl.float16)
        out = tl.dot(qh, v2h) + tl.dot(qh, v2l) + tl.dot(ql, v2h)
    else:
        out = tl.dot(q.to(tl.float32), v2.to(tl.float32), input_precision=PREC)
    tl.store(q_ptrs, out.to(dt))


@triton.jit
def _ns_orth_v(v16_ptr, v32_ptr, active_ptr, P2: tl.constexpr, TB: tl.constexpr):
    """Two fp32 Newton-Schulz steps per block rotation:
    V <- V (1.5 I - 0.5 V^T V), twice, in fp32/tf32x3 (the fp16 solver's V
    carries ~1e-2 L1 non-orthogonality; two quadratic steps reach ~1e-6)."""
    j = tl.program_id(0)
    if tl.load(active_ptr + j // P2) == 0:
        return
    offs = tl.arange(0, TB)
    v = tl.load(v16_ptr + j.to(tl.int64) * TB * TB +
                offs[:, None] * TB + offs[None, :]).to(tl.float32)
    eye = (offs[:, None] == offs[None, :]).to(tl.float32)
    g = tl.dot(tl.trans(v), v, input_precision="tf32x3")
    w = 1.5 * eye - 0.5 * g
    v = tl.dot(v, w, input_precision="tf32x3")
    g = tl.dot(tl.trans(v), v, input_precision="tf32x3")
    w = 1.5 * eye - 0.5 * g
    out = tl.dot(v, w, input_precision="tf32x3")
    tl.store(v32_ptr + j.to(tl.int64) * TB * TB +
             offs[:, None] * TB + offs[None, :], out)


def _jacobi_round(a, q, v, pairs, norms, stats, active, stop2, batch, inner,
                  tol2, prec, vq=None):
    """One dynamic-ordering block-Jacobi round (norms -> greedy [+ freeze] ->
    solve -> updates). Frozen (converged) matrices are skipped by every
    kernel. `vq` (if given) is the fp32 V buffer for accurate updates."""
    _block_norms[(_P512, _P512, batch)](a, norms, active, _N512, _BLK512, _P512)
    _greedy_pairs[(batch,)](norms, pairs, stats, active, stop2, _P512, _P2_512)
    torch.ops.eigh_ops.seated_pivot_solve(a, v, pairs, active, inner, tol2)
    vv = v
    if vq is not None:
        _ns_orth_v[(batch * _P2_512,)](v, vq, active, _P2_512, _TB512)
        vv = vq
    _jacobi_tile_update[(_P2_512, batch)](
        a, vv, pairs, active, _N512, _TB512, _P2_512, prec, num_warps=4)
    nw = 4 if vq is None else 2
    _RBM512 = 128
    _jacobi_q_update[(_P2_512, _N512 // _RBM512, batch)](
        q, vv, pairs, active, _N512, _TB512, _P2_512, prec, _RBM512, num_warps=nw)


def _fp16x3_bmm_out(out, lh, ll, rh, rl):
    """out = (lh+ll) @ (rh+rl) to ~2^-22 relative accuracy on fp16 tensor
    cores (drop the ll@rl term): the QR winner's fp16x3 trick."""
    torch.ops.eigh_ops.fp16_baddbmm_out(out, lh, rh, out, 0.0, 1.0)
    torch.ops.eigh_ops.fp16_baddbmm_out(out, lh, rl, out, 1.0, 1.0)
    torch.ops.eigh_ops.fp16_baddbmm_out(out, ll, rh, out, 1.0, 1.0)


def _split16(x):
    # Fused double-single fp16 split in ONE custom kernel (1 fp32 read -> hi/lo
    # fp16 writes), bit-identical to `hi = x.to(fp16); lo = (x - hi.float())
    # .to(fp16)` but ~4x less memory traffic and 1 launch instead of ~4. Accepts
    # a strided sub-block view for x (no .contiguous() copy needed).
    P, X, Y = x.shape
    hi = torch.empty(P, X, Y, device=x.device, dtype=torch.float16)
    lo = torch.empty(P, X, Y, device=x.device, dtype=torch.float16)
    torch.ops.eigh_ops.split16(x, hi, lo)
    return hi, lo


def _eigh512(data: input_t) -> output_t:
    batch = data.shape[0]
    n = _N512
    dev = data.device

    amax = data.abs().amax(dim=(1, 2), keepdim=True).clamp_min_(1e-30)
    inv = 1.0 / amax

    a_scaled = data * inv  # fp32, range-normalized; refresh target
    a16 = a_scaled.to(torch.float16)
    q16 = torch.eye(n, device=dev, dtype=torch.float16).expand(
        batch, n, n).contiguous()
    v16 = torch.empty(batch, _P2_512, _TB512, _TB512,
                      device=dev, dtype=torch.float16)
    norms = torch.empty(batch, _P512, _P512, device=dev, dtype=torch.float32)
    pairs = torch.empty(batch, _P2_512, 2, device=dev, dtype=torch.int32)
    stats = torch.zeros(batch, 2, device=dev, dtype=torch.float32)
    active = torch.ones(batch, device=dev, dtype=torch.int32)

    lo_stop2 = _LO_STOP512 * _LO_STOP512
    rnd = 0
    prev_rel2 = float("inf")
    while rnd < _LO_MAX_ROUNDS512:
        _jacobi_round(a16, q16, v16, pairs, norms, stats, active, lo_stop2,
                      batch, _LO_INNER512, 1e-7, "fp16")
        rnd += 1
        if rnd % _CHECK_EVERY512 == 0:
            if not active.any().item():
                break
            # plateau break: fp16 storage floors some structures above the
            # stop threshold; further rounds are wasted (the accurate phase
            # finishes the job). Norm-based, structure-agnostic.
            rel2 = (stats[:, 0] / stats[:, 1].clamp_min(1e-30)).amax().item()
            if rnd >= 16 and rel2 > 0.72 * prev_rel2:
                break
            prev_rel2 = rel2

    # Newton-Schulz re-orthonormalization: Q <- Q (1.5 I - 0.5 Q^T Q), twice.
    # Step 1 in TF32 (fp16-phase drift -> TF32 noise floor ~1e-3 L1); step 2
    # must beat TF32 input rounding (the accurate phase's block rotations
    # smear L1 non-orthogonality up to ~8x), so it runs fp16x3.
    q = q16.float()
    del q16
    old_tf32 = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = True
    try:
        w = torch.matmul(q.transpose(1, 2), q)
        w.mul_(-0.5)
        w.diagonal(dim1=1, dim2=2).add_(1.5)
        q = torch.matmul(q, w)
    finally:
        torch.backends.cuda.matmul.allow_tf32 = old_tf32
    qh, ql = _split16(q)
    w2 = w  # reuse
    _fp16x3_bmm_out(w2, qh.transpose(1, 2), ql.transpose(1, 2), qh, ql)
    w2.mul_(-0.5)
    w2.diagonal(dim1=1, dim2=2).add_(1.5)
    wh, wl = _split16(w2)
    q = torch.empty_like(w2)
    _fp16x3_bmm_out(q, qh, ql, wh, wl)

    # Accurate refresh vs the (scaled) ORIGINAL input: A_r = Q^T A_s Q (fp16x3)
    qh, ql = _split16(q)
    ah, al = _split16(a_scaled)
    del a_scaled
    t = w2  # reuse
    _fp16x3_bmm_out(t, ah, al, qh, ql)
    del ah, al
    th, tl_ = _split16(t)
    ar = t  # reuse (splits captured)
    _fp16x3_bmm_out(ar, qh.transpose(1, 2), ql.transpose(1, 2), th, tl_)
    del th, tl_, qh, ql
    ar = 0.5 * (ar + ar.transpose(1, 2))

    # Accurate phase: fp16 seated pivot solves + per-V fp32 NS-orth + tf32x3
    # tensor-core updates on the fp32 A_r / Q.
    v32 = torch.empty(batch, _P2_512, _TB512, _TB512,
                      device=dev, dtype=torch.float32)
    hi_stop2 = _HI_STOP512 * _HI_STOP512
    active.fill_(1)
    rnd = 0
    prev_rel2 = float("inf")
    while rnd < _HI_MAX_ROUNDS512:
        _jacobi_round(ar, q, v16, pairs, norms, stats, active, hi_stop2,
                      batch, _HI_INNER512, 1e-8, "fp16x3", vq=v32)
        rnd += 1
        if rnd % 2 == 0:
            if not active.any().item():
                break
            rel2 = (stats[:, 0] / stats[:, 1].clamp_min(1e-30)).amax().item()
            if rnd >= 6 and rel2 > 0.85 * prev_rel2:
                break
            prev_rel2 = rel2
    del v32, v16

    l = ar.diagonal(dim1=1, dim2=2) * amax.view(batch, 1)
    order = l.argsort(dim=1)
    l = l.gather(1, order)
    q = q.gather(2, order.unsqueeze(1).expand_as(q))
    return q.contiguous(), l.contiguous()




# ============================================================================
# n=512 COMPOSED pipeline:  one-stage CUDA tridiagonalization (blocked slatrd)
#   -> df32 divide-and-conquer tridiagonal eigensolver  -> WY back-transform.
#
# This replaces the two-sided block-Jacobi `_eigh512` at the scored n512 b640
# family. Reduce (latrd panel kernel, one CTA/matrix) produces the tridiagonal
# (d,e) plus the Householder reflectors (V,T) per 32-col panel; the D&C solve
# gives the tridiagonal eigenpairs (W,L) in double-single (df32) precision on
# fp16 tensor cores; the WY back-transform applies the reflectors to W.
# All GEMMs run FP32 (tf32 OFF) except the fp16x3 error-corrected products
# inside the D&C merges (those are precision-independent of the tf32 flag).
# See NOTES.md: composed 135.9 ms vs Jacobi 181 ms at n512 b640, gates PASS.
# ============================================================================

_DC_NEWT = 18
_DC_NF32 = 16

# fp64 secular-root polish convergence tolerances for the single-CTA merge_build.
# res: |g| <= RESTOL*erretm (residual-floor exit); step: |Δeta| <= STEPTOL*gapwidth.
# BASELINE (8e-16/1e-9) targets the fp64 roundoff floor — the default, kept for n2048
# and the small shapes. RELAX (1e-12/1e-7) retires "straggler" secular roots: on
# clustered/lapack_even spectra a few tight-gap roots kept iterating to ~17 (SIMT
# warp-max, so the whole warp waited) with ZERO accuracy gain (max-over-b640
# eigen/recon/orth BYTE-IDENTICAL 8e-16..1e-11 over 9 families x 3 seeds;
# dev/tol_robust.py). RELAX cuts them ~4-5 iters sooner -> n512 clustered -1.3ms,
# lapack_even -0.4, dense/mixed/rankdef -0.3..-0.7; n1024 mixed/nearrank/lapack_geo
# -0.3..-0.7 -- all gates unchanged. Applied to n512 + n1024 only: on n2048 the
# relaxed eta perturbs the later multi-CTA merges (+1.1ms), so n2048 keeps
# baseline. Held-out margin: 2 extra orders vs the proven-flat 1e-11/1e-6.
# Env-overridable for dev sweeps.
import os as _os
_DC_RESTOL = 8e-16
_DC_STEPTOL = 1e-9
_DC_RESTOL_RELAX = float(_os.environ.get("DC_RESTOL", "1e-12"))
_DC_STEPTOL_RELAX = float(_os.environ.get("DC_STEPTOL", "1e-7"))

# Divide-and-conquer z-deflation tolerance multiplier. The LAPACK slaed2-style
# deflation threshold (ztol = 8*m*eps*znorm) is machine-eps grade; the eigh
# gate is percent-level (rtol ~ 200*n*eps ~ 1e-2) and the dominant error floor
# is the fp16/tf32 tridiagonalisation+compose (~1e-5 scaled). Raising the
# threshold by 1e8 gives a deflation backward error ~ 1e8*eps ~ 2e-8, still
# ~500x below that floor, so residuals are BYTE-IDENTICAL to the eps-tol run
# while more poles deflate (each deflated pole skips an O(k) secular Newton
# solve + its rank-1 eigenvector column -> the secular phase is O(k^2), so the
# saving is quadratic). Measured: n512/n1024 merge_build ~-25%, full-solve
# n512 -2.2 ms / n1024 -1.8 ms geomean, all gates unchanged.
_DEFL_ZK = 1.0e8
# n2048 has a far looser gate margin (all residuals 40-160x below the caps), so
# it tolerates far more aggressive z-deflation than n512/n1024. At the SM-starved
# b8 batch the fp64 secular merge_build is ~11% of the solve; pushing the
# deflation tolerance to 1e9 deflates the extra fp16-noise poles -> merge_build
# shrinks (quadratic in the active-pole fraction) for -1.7 ms full-solve. 1e9
# keeps ~14x eigen margin (dense 14/200, mixed 8.3/200) and stays one order below
# the failure knee (1e10 -> dense eigen 251, FAILS). n512/n1024 keep 1e8 (their
# mixed cases sit at 58% of gate already at 1e9 -> not safe there; NOTES).
_DEFL_ZK_2048 = 1.0e9


def _dc_merge_G(P):
    """CTAs-per-subproblem G for the multi-CTA (thread-block-cluster) merge.
    Targets grid = G*P ~ 148 (fill the machine) while snapping to the
    instantiated cluster sizes {2,3,4,6,8}. G is capped at 8: measured on the
    n2048 top level (P=8), pushing to a Blackwell non-portable G=16 cluster
    (grid 64 -> 128) REGRESSED 5.42 -> 7.52 ms — cluster.sync + DSMEM remote
    map_shared_rank gather cost across 16 CTAs outweighs the shorter O(k^2/G)
    per-rank path, so 8-wide clusters are the accessible optimum."""
    return min(8, max(2, 148 // P))


def _dc_fp16x3(A, B, out=None):
    """(A@B) via 3-pass error-corrected fp16 tensor cores -> ~fp32 accuracy.
    If `out` is given (may be a strided sub-block view of a preallocated buffer),
    the product is written there in place — lets the caller skip a later cat."""
    P, x, y = A.shape
    z = B.shape[2]
    Ah, Al = _split16(A)
    Bh, Bl = _split16(B)
    if out is None:
        out = torch.empty(P, x, z, device=A.device, dtype=torch.float32)
    torch.ops.eigh_ops.fp16_baddbmm_out(out, Ah, Bh, out, 0.0, 1.0)
    torch.ops.eigh_ops.fp16_baddbmm_out(out, Ah, Bl, out, 1.0, 1.0)
    torch.ops.eigh_ops.fp16_baddbmm_out(out, Al, Bh, out, 1.0, 1.0)
    return out


def _tf32_trunc(x):
    """Truncate an fp32 tensor to tf32 precision (zero the low 13 mantissa bits)
    without changing dtype -> the value is exactly representable in tf32, so a
    subsequent tf32 tensor-core GEMM does not re-round it."""
    return (x.view(torch.int32) & -8192).view(torch.float32)


def _dc_fp16x1(A, B, out):
    """(A@B) via a SINGLE fp16 (fp32-accumulate) tensor-core GEMM. Both operands
    are rounded to fp16 (~2^-11 relative); the compose orthogonality error this
    incurs is repaired later by the terminal Newton-Schulz purify on the final
    Q = backT(compose(...)) (see _os_back_transform_fp16ns). Cheapest compose:
    1 TC pass vs fp16x3's 3 passes + hi/lo splits, or tf32x2's 2 passes. Since
    the terminal purify already re-orthogonalizes the FINAL Q, the per-merge-level
    exactness the older schedules paid is redundant -- measured (dev derisk,
    scored b640/b60): fp16x1-everywhere raises n512 mixed eigen only 77->81 (gate
    200) and leaves orth/recon unchanged, while cutting the compose ~1.2-1.8 ms
    at n512. `out` may be a strided sub-block view of a preallocated buffer."""
    torch.ops.eigh_ops.fp16_baddbmm_out(out, A.half(), B.half(), out, 0.0, 1.0)
    return out


def _dc_tf32x2(A, B, out):
    """(A@B) via a 2-pass error-corrected tf32 tensor-core GEMM: split A into a
    tf32-exact high part + fp32 residual low part and run BOTH through tf32 GEMMs
    (out = A_hi@B + A_lo@B), so A is captured to ~fp32 while B carries a single
    tf32 rounding. Removes A's tf32 rounding error vs a plain tf32 GEMM (~halves
    the compose orthogonality error) at ~2 tf32 passes -- still far cheaper than
    fp16x3, whose cost is dominated by the hi/lo split of BOTH operands, not the
    3 GEMMs. Caller sets allow_tf32=True. `out` may be a strided sub-block view."""
    Ahi = _tf32_trunc(A)
    Alo = A - Ahi
    torch.bmm(Ahi, B, out=out)
    torch.baddbmm(out, Alo, B, beta=1.0, alpha=1.0, out=out)
    return out


_LEAF_N = 32  # D&C leaf block size


def _dc_leaf_solve(d_leaf, e_leaf, leaf):
    """Solve the (P,leaf) SYMMETRIC-TRIDIAGONAL leaf blocks with a warp-per-matrix
    implicit-shift QL (steqr_tri) directly on d,e (no densification). Returns D
    (P,leaf) fp64 ascending, Q (P,leaf,leaf) fp32. The fp32 chase (arg 5) and the
    skipped in-kernel sort (arg 6, the merge re-sorts) are the tuned defaults."""
    P = d_leaf.shape[0]
    dev = d_leaf.device
    d = d_leaf.double().contiguous()
    e = e_leaf.double().contiguous()
    Q = torch.empty(P, leaf, leaf, device=dev, dtype=torch.float32)
    L = torch.empty(P, leaf, device=dev, dtype=torch.float64)
    torch.ops.eigh_ops.steqr_tri(d, e, Q, L, 1, 0)
    return L, Q


def _dc_prep_fused(Ql, Qrr, Dl, Drr, rho, defl_zk=1.0, defl_gk=1.0):
    """Fused per-level D&C bookkeeping in ONE custom kernel (one CTA per merge
    subproblem): sort + gap-cluster + segmented reductions + Householder
    deflation + active compaction + invorder/perm/segstart build. Replaces a
    ~130-launch torch glue chain. Returns the merge_build input tuple
    (dc, zc, k, rho_s, invorder, perm, hv, hbeta, segstart, nseg)."""
    P, h, _ = Ql.shape
    m = 2 * h
    dev = Ql.device
    zL = Ql[:, h - 1, :].contiguous()
    zR = Qrr[:, 0, :].contiguous()
    Dlc = Dl.contiguous(); Drrc = Drr.contiguous()
    rhoc = rho.contiguous()
    dc = torch.empty(P, m, dtype=torch.float64, device=dev)
    zc = torch.empty(P, m, dtype=torch.float64, device=dev)
    k = torch.empty(P, dtype=torch.int32, device=dev)
    rho_s = torch.empty(P, dtype=torch.float64, device=dev)
    invorder = torch.empty(P, m, dtype=torch.int32, device=dev)
    perm = torch.empty(P, m, dtype=torch.int32, device=dev)
    hv = torch.empty(P, m, dtype=torch.float64, device=dev)
    hbeta = torch.empty(P, m, dtype=torch.float64, device=dev)
    segstart = torch.empty(P, m, dtype=torch.int32, device=dev)
    nseg = torch.empty(P, dtype=torch.int32, device=dev)
    torch.ops.eigh_ops.dc_prep(Dlc, Drrc, zL, zR, rhoc, dc, zc, k, rho_s,
                               invorder, perm, hv, hbeta, segstart, nseg,
                               defl_zk, defl_gk)
    return dc, zc, k, rho_s, invorder, perm, hv, hbeta, segstart, nseg


def _dc_merge(Ql, Qrr, Dl, Drr, rho, compose_tf32_hmin=1 << 30, compose_tf32x2_hmin=1 << 30,
              defl_zk=1.0, defl_gk=1.0, compose_fp16x1_hmax=0,
              res_tol=_DC_RESTOL, step_tol=_DC_STEPTOL):
    """One batched divide-and-conquer merge. Ql,Qrr (P,h,h) fp32; Dl,Drr (P,h)
    fp64. Returns Qnew (P,m,m) fp32, Lam (P,m) fp64."""
    P, h, _ = Ql.shape
    m = 2 * h
    dev = Ql.device
    (dc, zc, k, rho_s, invorder, perm, hv, hbeta,
     segstart, nseg) = _dc_prep_fused(Ql, Qrr, Dl, Drr, rho, defl_zk, defl_gk)
    Vchild = torch.empty(P, m, m, device=dev, dtype=torch.float32)
    Lam = torch.empty(P, m, device=dev, dtype=torch.float64)
    args = (dc, zc, k, rho_s, invorder, perm, hv, hbeta, segstart, nseg,
            Vchild, Lam, _DC_NEWT, _DC_NF32)
    # Multi-CTA (thread-block cluster) merge when the batch under-fills the
    # machine: G CTAs cooperate on each merge (secular over roots, zhat over
    # poles, V-build over columns), filling idle SMs and shortening the O(m^2)
    # critical path. P>=148 keeps the single-CTA path -> n512 b640 (every level
    # P>=640) is UNCHANGED. Per-column arithmetic is byte-identical.
    if P < 148:
        G = _dc_merge_G(P)
        torch.ops.eigh_ops.merge_build_multi(*args, G)
    else:
        torch.ops.eigh_ops.merge_build(*args, res_tol, step_tol)
    # Write the two child-block products straight into a preallocated Qnew
    # (top -> rows [:h], bot -> rows [h:]) via strided out slices, skipping
    # the torch.cat([top,bot]) copy (a full (P,m,m) read+write per level).
    # compose_tf32: at the SM-STARVED small-batch shapes (n1024 P=60..960,
    # n2048 P<=32) a single tf32 tensor-core bmm is 6-8x faster than fp16x3
    # (3 fp16 passes + hi/lo splits) -- fp16x3's per-matrix-small/batch-large
    # advantage inverts to a big LOSS when the matrix is large and the batch
    # tiny. The child eigenvectors Ql/Qrr and the merge V are both orthonormal,
    # so the tf32 product (fp32 accumulation, tf32 operand rounding ~2^-11)
    # keeps the eigenvectors well inside the n1024/n2048 orth+eigen gates
    # (validated at scored batches). n512 (small matrices, large batch) stays
    # fp16x3 where it wins and the gates are tighter.
    Qnew = torch.empty(P, m, m, device=dev, dtype=torch.float32)
    if h <= compose_fp16x1_hmax:
        # Single-pass fp16 compose at this merge level; orthogonality repaired by
        # the terminal NS purify (see _os_back_transform_fp16ns). Cheapest compose.
        _dc_fp16x1(Ql, Vchild[:, :h, :], Qnew[:, :h, :])
        _dc_fp16x1(Qrr, Vchild[:, h:, :], Qnew[:, h:, :])
    elif h >= compose_tf32_hmin:
        old = torch.backends.cuda.matmul.allow_tf32
        torch.backends.cuda.matmul.allow_tf32 = True
        torch.bmm(Ql, Vchild[:, :h, :], out=Qnew[:, :h, :])
        torch.bmm(Qrr, Vchild[:, h:, :], out=Qnew[:, h:, :])
        torch.backends.cuda.matmul.allow_tf32 = old
    elif h >= compose_tf32x2_hmin:
        # tf32x2 splits Ql via .view(int32), which needs a C-contiguous operand;
        # Ql/Qrr arrive as reshape views (strided batch) -> materialize here.
        Ql = Ql.contiguous(); Qrr = Qrr.contiguous()
        old = torch.backends.cuda.matmul.allow_tf32
        torch.backends.cuda.matmul.allow_tf32 = True
        _dc_tf32x2(Ql, Vchild[:, :h, :], Qnew[:, :h, :])
        _dc_tf32x2(Qrr, Vchild[:, h:, :], Qnew[:, h:, :])
        torch.backends.cuda.matmul.allow_tf32 = old
    else:
        _dc_fp16x3(Ql, Vchild[:, :h, :], out=Qnew[:, :h, :])
        _dc_fp16x3(Qrr, Vchild[:, h:, :], out=Qnew[:, h:, :])
    return Qnew, Lam


def _dc_eigh(d, e, leaf=32, compose_tf32_hmin=1 << 30, compose_tf32x2_hmin=1 << 30,
            defl_zk=1.0, defl_gk=1.0, compose_fp16x1_hmax=0,
            res_tol=_DC_RESTOL, step_tol=_DC_STEPTOL):
    """Batched divide-and-conquer symmetric tridiagonal eigensolver.
    d (B,n), e (B,n-1) -> (W (B,n,n) fp32 eigenvectors, L (B,n) fp32 ascending)."""
    B, n = d.shape
    dev = d.device
    d = d.double(); e = e.double()
    assert n % leaf == 0
    K = n // leaf
    d_mod = d.clone()
    bnds = torch.arange(leaf, n, leaf, device=dev)
    d_mod[:, bnds - 1] -= e[:, bnds - 1]
    d_mod[:, bnds] -= e[:, bnds - 1]
    d_leaf = d_mod.reshape(B * K, leaf)
    # e_leaf[b,k,:] = e[b, k*leaf : k*leaf+leaf-1] (the leaf-1 intra-block
    # couplings, dropping the boundary coupling e[k*leaf+leaf-1]). Vectorized:
    # pad e to K*leaf, view as (B,K,leaf), drop the last column of each block --
    # replaces the K-iteration Python slice-copy loop with two ops.
    ep = torch.nn.functional.pad(e, (0, 1))              # (B, K*leaf)
    e_leaf = ep.reshape(B, K, leaf)[:, :, :leaf - 1].reshape(B * K, leaf - 1).contiguous()
    D, Q = _dc_leaf_solve(d_leaf, e_leaf, leaf)

    h = leaf; curK = K
    while curK > 1:
        newK = curK // 2
        P = B * newK
        m = 2 * h
        Dr = D.reshape(B, curK, h)
        Qr = Q.reshape(B, curK, h, h)
        # Reshape-only VIEW (batch stride 2*h*h): the even/odd child slices are a
        # valid view, so skip the .contiguous() copy. The fp16x3 compose reads Ql
        # only through split16 (arbitrary-stride aware) + strided-batch bmm, and
        # dc_prep re-materializes just the boundary row (zL/zR). Only the tf32x2
        # compose (which uses .view() for the hi/lo split) needs a contiguous Ql,
        # so it re-contiguates locally. Drops the per-level (P,h,h) reshape copies.
        Ql = Qr[:, 0::2].reshape(P, h, h)
        Qrr = Qr[:, 1::2].reshape(P, h, h)
        Dl = Dr[:, 0::2].reshape(P, h)
        Drr = Dr[:, 1::2].reshape(P, h)
        mlocal = torch.arange(newK, device=dev)
        split_idx = mlocal * m + h - 1
        rho = e[:, split_idx].reshape(P)
        Q, D = _dc_merge(Ql, Qrr, Dl, Drr, rho, compose_tf32_hmin=compose_tf32_hmin, compose_tf32x2_hmin=compose_tf32x2_hmin,
                         defl_zk=defl_zk, defl_gk=defl_gk, compose_fp16x1_hmax=compose_fp16x1_hmax,
                         res_tol=res_tol, step_tol=step_tol)
        h = m; curK = newK

    L, si = torch.sort(D.reshape(B, n), dim=1)
    W = Q.reshape(B, n, n)
    W = torch.gather(W, 2, si.unsqueeze(1).expand(-1, n, -1))
    return W.to(torch.float32), L.to(torch.float32)


def _os_buildT(V):
    """Compact-WY T (b,nb,nb upper) from the Gram G = V^T V. Gram must be FP32
    (orthogonality-critical); caller must have tf32 disabled."""
    b, m, nb = V.shape
    G = torch.bmm(V.transpose(1, 2), V)                # (b,nb,nb) Gram (fp32)
    T = torch.zeros(b, nb, nb, device=V.device, dtype=torch.float32)
    torch.ops.eigh_ops.trec(G.contiguous(), T)
    return T


def _prescale(data):
    """Fused amax-prescale: scale = amax(|data|) per matrix (clamped), A =
    data/scale, Ah = A.half() (fp16 shadow) -- one reduction + one float4 pass,
    replacing the abs-temp + amax + divide + half-cast torch chain."""
    b, n, _ = data.shape
    data = data.contiguous()
    A = torch.empty_like(data)
    Ah = torch.empty(b, n, n, device=data.device, dtype=torch.float16)
    scale = torch.empty(b, device=data.device, dtype=torch.float32)
    torch.ops.eigh_ops.prescale(data, A, Ah, scale)
    return A, Ah, scale


def _os_sytrd(A, nb=32, cluster=0, tf32_trailing=False, nb_big=0, Ah=None, skip_T=False, agg_buildT=False, tail_m0=0):
    """Dense -> tridiagonal via CUDA latrd panels + torch rank-2b trailing
    update. Returns (d (b,n), e (b,n-1), refl [(p0,V,T)...]).
    cluster>=2 -> thread-block-cluster latrd (C CTAs/matrix, symv rows split):
    fills the machine when batch << SM count (n1024 b60, n2048 b8). The panel
    reflector build stays FP32 (orthogonality-critical); tf32_trailing only
    affects the BLAS-3 rank-2b trailing update, which is gate-safe.
    skip_T=True: the caller's back-transform rebuilds each compound-WY T from
    the assembled compound Gram (closed-form Schreiber-Van Loan), so the per-
    panel T is later DISCARDED by the aggregate -- skip building it entirely
    (drop the deferred batched Gram+trec launches). refl then carries
    (p0, V, None). Only valid when every reflector group is closed-form
    aggregated (n2048).
    nb_big>0 (n2048): once the trailing block m <= 1024 the fp16 WY panel fits
    SMEM at the wider nb_big, so widen the panel there — this HALVES the number
    of panels over the second half of the reduction, cutting the per-panel
    trailing-GEMM / buildT launch overhead (which dominates at nb=16, b8).
    The symv work is nb-independent, so the wide panels don't cost the reduce.
    (m<=1024 keeps MAXV=32 valid for the nb_big cluster kernel.)"""
    b, n, _ = A.shape
    A = A.contiguous()
    # fp16 shadow of the trailing matrix: read only by the latrd symv (the O(m^2)
    # DRAM wall) to halve its traffic. Inputs are amax-normalized to O(1) so fp16 is
    # in-range and ~8x more accurate than bf16; d/e stay fp32 (column-correction
    # reads the fp32 A). Updated on the same trailing slice each panel. When the
    # caller already produced the fp16 shadow (fused prescale) it is passed in.
    if Ah is None:
        Ah = A.half()
    d = torch.zeros(b, n, device=A.device, dtype=torch.float32)
    e = torch.zeros(b, n - 1, device=A.device, dtype=torch.float32)
    refl = []
    old_tf32 = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = tf32_trailing
    # Cluster (n1024/n2048) path: DEFER the compact-WY T build. Per panel the
    # T-recurrence (`trec`, 1 warp/matrix) and the nb x nb Gram are tiny but
    # LATENCY-bound, so ~32-96 separate launches dominate (measured: trec 1.67ms
    # n2048 / 0.82ms n1024 as a launch-storm). Default-queue serialization means
    # (default-queue serialized) these can't overlap latrd, so instead collect
    # them and build ALL T's in ONE
    # batched trec launch (Grams stacked over panels by nb) after the loop -> the
    # 96 trec launches collapse to ~2 (per nb bucket). Byte-identical T. n512
    # (cluster==0) keeps the per-panel path (its 16 launches are already cheap).
    defer_T = cluster >= 2
    pend = []                                  # (p0, V, G, nb) when deferring
    p0 = 0
    # Tail collapse (n512): once the trailing block m <= tail_m0 fits in SMEM,
    # a SINGLE kernel finishes the remaining tridiagonalization (unblocked
    # ssytd2, fp32 SMEM), emitting the reflectors + d/e in the panel layout.
    # This deletes the last tail_m0/nb latrd launches + their per-panel trailing
    # GEMM / epilogue / glue (all barrier/launch-dominated on the tiny blocks).
    stop = (n - tail_m0) if tail_m0 else (n - 1)
    while p0 < stop:
        m = n - p0
        cur_nb = nb_big if (nb_big and m <= 1024) else nb
        cw = min(cur_nb, m - 1)
        # DEAD-WORK ELIMINATION: only V's zero-init is live. latrd writes V's
        # column i at rows {i+1 (=1), >=i+2 (=vv)} and leaves rows <=i for the
        # implicit-identity upper triangle -> V MUST be zeroed.
        # W zero-init is DEAD only for the SINGLE-BLOCK latrd (n512, cluster==0):
        # it writes W's column i at every row >=i+1 in one contiguous sweep, so the
        # trailing Q=Vt@Wt^T reader (the [cw:,:cw] slice, rows>=cw>=i+1) sees only
        # written rows. But the CLUSTER latrd (cluster>=2, n1024/n2048) SPLITS the
        # symv rows across C CTAs and writes Wout only on each rank's
        # [max(rlo,i+1), rhi) slab: the cross-rank "support rows" (kernel comment
        # "their w is 0") are deliberately LEFT to the host zero-init. They fall in
        # the reader's [cw:,:cw] slice, so with torch.empty they read RECYCLED,
        # possibly non-finite, allocator memory -> NaN in the trailing update ->
        # NaN d/e -> NaN L and Q. This only surfaces on the 2nd+ call on a reused
        # input (the caching allocator hands back dirty blocks), i.e. the grader's
        # benchmark loop rechecks every iteration while local_benchmark checks only
        # the first (fresh, ~zeroed) call -> silent on cu13 local, fails the grader.
        # So W MUST be zeroed whenever cluster>=2. n512 keeps empty (a b640 W memset
        # per panel is NOT free, and its single-block W read region is complete).
        # dp/ep ARE fully written in their read ranges on both paths (dout[i] every
        # i; eout[i] for i<=m-2, read range i<min(nb,m-1)) -> empty() safe.
        V = torch.zeros(b, m, cur_nb, device=A.device, dtype=torch.float32)
        W = (torch.zeros if cluster >= 2 else torch.empty)(b, m, cur_nb, device=A.device, dtype=torch.float32)
        dp = torch.empty(b, cur_nb, device=A.device, dtype=torch.float32)
        ep = torch.empty(b, cur_nb, device=A.device, dtype=torch.float32)
        if cluster >= 2:
            torch.ops.eigh_ops.latrd_cluster(A, Ah, V, W, dp, ep, p0, cluster)
        else:
            torch.ops.eigh_ops.latrd(A, Ah, V, W, dp, ep, p0)
        d[:, p0:p0 + cw] = dp[:, :cw]
        ecnt = min(cw, (n - 1) - p0)
        e[:, p0:p0 + ecnt] = ep[:, :ecnt]
        if p0 + cw < n:
            Vt = V[:, cw:, :cw]; Wt = W[:, cw:, :cw]
            # rank-2b update A -= V W^T + W V^T is SYMMETRIC, so form only the
            # single-side product Q = Vt @ Wt^T (ONE K=nb GEMM, no cat) and let
            # the epilogue add Q^T on the fly (shared-memory tile transpose):
            # A[slice] -= (Q + Q^T); Ah[slice] = A.half(). This drops the two
            # torch.cat([V|W]/[W|V]) copies AND halves the trailing GEMM (K=nb
            # vs K=2nb) at identical epilogue traffic. FP32 (n512) / tf32 (flag).
            Q = torch.bmm(Vt, Wt.transpose(1, 2)).contiguous()
            torch.ops.eigh_ops.trail_epilogue_sym(A, Ah, Q, p0 + cw)
        # buildT/back-transform use the full 32-col panel; unused trailing
        # columns are zero -> tau=0 -> identity reflector (harmless). Gram stays
        # FP32 (orthogonality-critical) regardless of the trailing tf32 flag.
        if skip_T or agg_buildT:
            # skip_T (n2048): closed-form aggregate rebuilds T from V.
            # agg_buildT (n512): the aggregate ALSO forms the compound Gram
            # M = V_group^T V_group whose 32x32 diagonal blocks are EXACTLY the
            # per-panel Grams V_j^T V_j -- so the per-panel buildT (Gram + trec)
            # is 100% redundant with M. Defer T to the aggregate, which recovers
            # each panel's compact-WY T from M's diagonal via ONE batched trec
            # (byte-identical), dropping the 16 redundant Gram bmms + 16 trec
            # launches + 16 T zero-fills.
            refl.append((p0, V, None))
        elif defer_T:
            pend.append((p0, V, cur_nb))       # T built in one batched pass below
        else:
            prev = torch.backends.cuda.matmul.allow_tf32
            torch.backends.cuda.matmul.allow_tf32 = False
            refl.append((p0, V, _os_buildT(V)))
            torch.backends.cuda.matmul.allow_tf32 = prev
        p0 += cw
    if tail_m0:
        # p0 == n - tail_m0; the trailing block A[p0:,p0:] (tail_m0 x tail_m0)
        # has been produced by the last panel's trailing update. Reduce it whole
        # in one SMEM-resident launch. Vfull (b,tail_m0,tail_m0) trapezoidal.
        m0 = tail_m0
        Vfull = torch.zeros(b, m0, m0, device=A.device, dtype=torch.float32)
        dp2 = torch.zeros(b, m0, device=A.device, dtype=torch.float32)
        ep2 = torch.zeros(b, m0, device=A.device, dtype=torch.float32)
        torch.ops.eigh_ops.latrd_tail(A, Vfull, dp2, ep2, p0)
        d[:, p0:p0 + m0] = dp2
        e[:, p0:p0 + m0 - 1] = ep2[:, :m0 - 1]
        # slice Vfull into nb-wide panel views (p0 = base + nb*k) -> refl entries
        # in the exact (p0, V, None) form the agg_buildT aggregate consumes.
        for k in range(0, m0, nb):
            refl.append((p0 + k, Vfull[:, k:, k:k + nb], None))
        torch.backends.cuda.matmul.allow_tf32 = old_tf32
        return d, e, refl
    if defer_T and not skip_T:
        # ONE batched Gram + ONE batched trec per nb bucket. The per-panel Gram
        # (V^T V, nb x nb out) and trec (1 warp/matrix) are both TINY but
        # LATENCY-bound, so ~32-96 launches each dominate (measured n2048:
        # gram 2.27ms + trec 1.67ms as launch-storms). Stack all panels of a
        # given nb (zero-padded to the bucket's max panel height -> the pad rows
        # contribute 0 to V^T V) into one bmm, and stack the resulting Grams into
        # one trec launch. Byte-identical T; FP32 Gram (orthogonality-critical).
        from collections import defaultdict as _dd
        buck = _dd(list)
        for i, (_pp, _V, _nb) in enumerate(pend):
            buck[_nb].append(i)
        prev = torch.backends.cuda.matmul.allow_tf32
        torch.backends.cuda.matmul.allow_tf32 = False
        Ts = [None] * len(pend)
        for _nb, idxs in buck.items():
            ng = len(idxs)
            mmax = max(pend[i][1].shape[1] for i in idxs)
            Vpad = torch.zeros(ng * b, mmax, _nb, device=A.device, dtype=torch.float32)
            for k, i in enumerate(idxs):
                mv = pend[i][1].shape[1]
                Vpad[k * b:(k + 1) * b, :mv, :] = pend[i][1]
            Gst = torch.bmm(Vpad.transpose(1, 2), Vpad).contiguous()   # (ng*b,nb,nb)
            Tst = torch.zeros(ng * b, _nb, _nb, device=A.device, dtype=torch.float32)
            torch.ops.eigh_ops.trec(Gst, Tst)
            for k, i in enumerate(idxs):
                Ts[i] = Tst[k * b:(k + 1) * b]
        torch.backends.cuda.matmul.allow_tf32 = prev
        for i, (pp, V, _nb) in enumerate(pend):
            refl.append((pp, V, Ts[i]))
    d[:, n - 1] = A[:, n - 1, n - 1]
    torch.backends.cuda.matmul.allow_tf32 = old_tf32
    return d, e, refl


def _os_back_transform(Z, refl):
    """Apply the stored WY reflectors to the tridiagonal eigenvectors Z (in
    place). FP32 (caller must have tf32 disabled)."""
    for (r0, V, T) in reversed(refl):
        Zt = Z[:, r0:, :]
        VtZ = torch.bmm(V.transpose(1, 2), Zt)
        # in-place subtract: 1 memory pass (read Zt + P, write Zt) instead of
        # `Zt - P` (new temp) followed by a `Z[slice] = ...` writeback copy.
        Zt.sub_(torch.bmm(V, torch.bmm(T, VtZ)))
    return Z


def _os_back_transform_x3(Z, refl):
    """WY back-transform with the two big BLAS-3 products in 3-pass error-
    corrected FP16 (tensor cores). Blackwell has NO native FP32 tensor core, so
    the plain-FP32 back-transform (_os_back_transform) runs on the SIMT/CUDA
    cores; fp16x3 gives ~FP32 accuracy (err ~1e-7, gate-validated identical
    scaled residual across all families) at tensor-core throughput. The small
    nb x nb x k T-product stays FP32 (negligible). In-place on Z."""
    fp16 = torch.ops.eigh_ops.fp16_baddbmm_out
    for (r0, V, T) in reversed(refl):
        Zt = Z[:, r0:, :]
        VtZ = _dc_fp16x3(V.transpose(1, 2), Zt)   # (b,nb,k); split16 takes strided
        TVtZ = torch.bmm(T, VtZ)                   # (b,nb,k) fp32
        Vh, Vl = _split16(V)
        Bh, Bl = _split16(TVtZ)
        # Zt <- Zt - V @ TVtZ, subtract FUSED into the 3 fp16 TC passes (beta=1,
        # alpha=-1) writing in place into the strided Z slice -> no separate
        # elementwise subtract, no intermediate (killed 16 big binary kernels).
        fp16(Zt, Vh, Bh, Zt, 1.0, -1.0)
        fp16(Zt, Vh, Bl, Zt, 1.0, -1.0)
        fp16(Zt, Vl, Bh, Zt, 1.0, -1.0)
    return Z


def _os_back_transform_tf32x2(Z, refl):
    """WY back-transform with the two big BLAS-3 products (and the small T @ VtZ)
    run as 3-term error-corrected tf32 tensor-core GEMMs -> ~fp32 accuracy on the
    higher-efficiency tf32 tensor cores. Blackwell has NO fp32 tensor core, so
    the plain-fp32 apply (_os_back_transform) is SIMT; fp16x3 is the other TC
    route. This captures BOTH operands of every product to ~fp32 by the hi/lo
    split (hi=_tf32_trunc, lo=residual) and summing the three cross terms
    hi*hi + hi*lo + lo*hi (the lo*lo term ~2^-22 is dropped) in tf32.

    WHY the full both-operand + T split is REQUIRED (measured, NOTES tf32x2-backT):
    the reflector back-transform is a CHAIN of orthogonal applies I - V T V^T, so
    -- unlike the D&C compose's single product where a one-operand tf32x2 split
    passes -- rounding ANY of the three (the carried orthonormal Z, the reflector
    V, or T which must stay consistent with V) to a single tf32 pass blows the
    tight orth gate: measured n512 b640 orth 370-684/100 for split-V / split-Z /
    T-single-tf32, vs 51 for the full split (== fp32). So the one-operand trick
    canNOT rescue this apply; only the full ~fp32 tf32 route is gate-legal.

    Shipped at n1024 ONLY: at the fat NB=256 aggregated blocks (big per-matrix,
    launch-bound b60) the tf32 tensor cores edge out fp16x3 by ~0.3 ms at
    identical accuracy. At n512 (small blocks, saturated b640) the extra tf32
    passes LOSE +1.5 ms to fp32 SIMT; at n2048 b8 it is FLAT vs fp16x3."""
    old = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = True
    try:
        for (r0, V, T) in reversed(refl):
            Zt = Z[:, r0:, :]
            Vhi = _tf32_trunc(V); Vlo = V - Vhi
            VhiT = Vhi.transpose(1, 2); VloT = Vlo.transpose(1, 2)
            Zc = Zt.contiguous()
            Zhi = _tf32_trunc(Zc); Zlo = Zc - Zhi
            # GEMM1: VtZ = V^T @ Zt   (both operands captured to ~fp32)
            VtZ = torch.bmm(VhiT, Zhi)
            torch.baddbmm(VtZ, VhiT, Zlo, beta=1.0, alpha=1.0, out=VtZ)
            torch.baddbmm(VtZ, VloT, Zhi, beta=1.0, alpha=1.0, out=VtZ)
            # T-product: TVtZ = T @ VtZ  (T defines the operator -> also split)
            Thi = _tf32_trunc(T); Tlo = T - Thi
            Whi = _tf32_trunc(VtZ); Wlo = VtZ - Whi
            TVtZ = torch.bmm(Thi, Whi)
            torch.baddbmm(TVtZ, Thi, Wlo, beta=1.0, alpha=1.0, out=TVtZ)
            torch.baddbmm(TVtZ, Tlo, Whi, beta=1.0, alpha=1.0, out=TVtZ)
            # GEMM2: Zt <- Zt - V @ TVtZ   (subtract fused into the 3 tf32 passes)
            Bhi = _tf32_trunc(TVtZ); Blo = TVtZ - Bhi
            torch.baddbmm(Zt, Vhi, Bhi, beta=1.0, alpha=-1.0, out=Zt)
            torch.baddbmm(Zt, Vhi, Blo, beta=1.0, alpha=-1.0, out=Zt)
            torch.baddbmm(Zt, Vlo, Bhi, beta=1.0, alpha=-1.0, out=Zt)
    finally:
        torch.backends.cuda.matmul.allow_tf32 = old
    return Z


def _os_back_transform_fp16ns(Z, refl):
    """Plain fp16 (fp32-accumulate) tensor-core WY back-transform + ONE
    Newton-Schulz orthogonality purify.

    Motivation (NOTES fp16-first scout idea #1+#5): the fp32 SIMT WY apply
    (_os_back_transform) and the fp16x3 / tf32-full-split routes all pay
    ~fp32-accurate arithmetic (3-9 GEMM passes / SIMT cores) to keep Q
    orthonormal through the reflector chain. But the eigh gates are loose
    (orth 100*n*eps): a SINGLE-pass fp16 apply (1 GEMM/product on the tensor
    cores, ~2x the fp32-SIMT throughput) plus ONE Newton-Schulz half-step
    Q <- Q(1.5 I - 0.5 Q^T Q) reaches the gate with wide margin, at a fraction
    of the arithmetic. Measured (dev/derisk_fp16_backt.py, real solver state,
    scored batches, L kept from D&C unchanged -- diag(Q^T A0 Q) recovery is
    both costlier and no more accurate here):
      shape/worst-family  fp16-apply orth  ->  +NS(fp16) orth   eigen  recon
      n512  / mixed             766         ->      17.2         77.1   69.6
      n1024 / mixed             499         ->       8.9         42.8   34.6
      n2048 / mixed             334         ->       4.5          7.5    5.1
    (gates orth<100 eigen<200 recon<400 -- all pass; fp16 purify is enough,
    tf32/fp16x3 purify only tighten orth further at more cost.)

    The tf32x2-backT kill (NOTES ~355) DIED trying to reach fp32 accuracy with
    more tf32 GEMMs and still lost n512; this route accepts fp16 accuracy and
    REPAIRS orthogonality with the purify step that kill never tried."""
    b = Z.shape[0]
    fp16 = torch.ops.eigh_ops.fp16_baddbmm_out
    for (r0, V, T) in reversed(refl):
        Zt = Z[:, r0:, :]
        nb = V.shape[2]
        k = Zt.shape[2]
        Vh = V.half()
        VtZ = torch.empty(b, nb, k, device=Z.device, dtype=torch.float32)
        fp16(VtZ, Vh.transpose(1, 2), Zt.half(), VtZ, 0.0, 1.0)   # V^T @ Zt
        TVtZ = torch.empty(b, nb, k, device=Z.device, dtype=torch.float32)
        fp16(TVtZ, T.half(), VtZ.half(), TVtZ, 0.0, 1.0)          # T @ VtZ
        fp16(Zt, Vh, TVtZ.half(), Zt, 1.0, -1.0)                  # Zt -= V @ TVtZ
    return _ns_purify_fp16(Z)


def _ns_purify_fp16(Q):
    """One Newton-Schulz orthogonality half-step Q <- Q(1.5 I - 0.5 Q^T Q) with
    fp16 (fp32-accumulate) tensor-core GEMMs. Purifies the orthogonality of a
    low-precision Q back to the eigh orth gate (quadratic convergence for
    ||Q||_2 < sqrt(3), which a fp16-applied orthonormal WY product satisfies by
    a wide margin). Two batched GEMMs; L (from D&C) stays unchanged."""
    b, n, _ = Q.shape
    fp16 = torch.ops.eigh_ops.fp16_baddbmm_out
    Qh = Q.half()
    G = torch.empty(b, n, n, device=Q.device, dtype=torch.float32)
    fp16(G, Qh.transpose(1, 2), Qh, G, 0.0, 1.0)     # G = Q^T Q
    G.mul_(-0.5)                                       # M = 1.5 I - 0.5 G
    G.diagonal(dim1=1, dim2=2).add_(1.5)
    Qout = torch.empty_like(Q)
    fp16(Qout, Qh, G.half(), Qout, 0.0, 1.0)          # Q @ M
    return Qout


def _aggregate_refl(refl, group, closed=False, build_T_from_M=False):
    """Compound `group` consecutive nb panel WY reflectors into ONE wider
    compact-WY block (V (b,m,G*nb), T (b,G*nb,G*nb) block-upper-triangular), so
    the back-transform runs as FEWER, FATTER GEMMs. Reflectors are aligned to
    the outermost group row r0 (later panels zero-padded on top).

    Two ways to build the block-WY T from the assembled reflectors V:
      * recurrence (default): the LAPACK LARFT block recurrence
        T_agg = [[Ta, -Ta (Va^T Vb) Tb],[0, Tb]]; a chain of tiny bmms. Batched
        across all same-shape groups so the launch count drops ~num_groups-fold.
        Byte-identical to the original per-group loop.
      * closed form (`closed=True`, SM-starved n1024/n2048 launch-bound backT):
        the compact-WY factor is the Schreiber-Van Loan triangular inverse
        T = (0.5*diag(M) + strictuu(M))^-1 with M = V^T V (since M_ii=||v_i||^2=
        2/tau_i so 1/tau_i = M_ii/2). The ENTIRE recurrence collapses to ONE
        batched upper-triangular solve -> ~30 tiny bmms become 1 launch (~1.1 ms
        host saved at n2048 b8). Null reflectors (M_ii<=0) get R_ii=1 (their V
        column is zero, so the T entry is inert). Reproduces the recurrence T to
        fp32 roundoff (rel ~2e-7); the invariant-based gates keep >70x margin."""
    p = len(refl)
    slots = []            # ('single', tuple) | ('multi', idx)
    metas = []            # (r0, V, M) [closed] or (r0, V, M, T0, widths, starts) [recur]
    from collections import defaultdict as _dd
    buckets = _dd(list)   # signature -> [meta index]
    i = 0
    while i < p:
        grp = refl[i:i + group]
        i += group
        if len(grp) == 1:
            slots.append(('single', grp[0]))
            continue
        r0 = grp[0][0]
        V0 = grp[0][1]
        b, m, _ = V0.shape
        widths = [g[1].shape[2] for g in grp]
        NB = sum(widths)
        V = torch.zeros(b, m, NB, device=V0.device, dtype=torch.float32)
        starts = []
        cc = 0
        for (rj, Vj, Tj) in grp:
            ro = rj - r0
            w = Vj.shape[2]
            V[:, ro:, cc:cc + w] = Vj
            starts.append(cc)
            cc += w
        # cross-block Gram (kept FP32: feeds the orthogonality-relevant T blocks;
        # an fp16x3 TC Gram was a measured wash).
        M = torch.bmm(V.transpose(1, 2), V)              # (b, NB, NB) fp32
        idx = len(metas)
        if closed:
            metas.append((r0, V, M))
            buckets[NB].append(idx)
        else:
            if build_T_from_M:
                # T0 recovered from M's diagonal blocks in ONE batched trec
                # after the loop (the per-panel Grams == M's diagonal blocks).
                T0 = None
            else:
                T0 = torch.zeros(b, NB, NB, device=V0.device, dtype=torch.float32)
                for j, (rj, Vj, Tj) in enumerate(grp):
                    s = starts[j]; w = widths[j]
                    T0[:, s:s + w, s:s + w] = Tj
            metas.append([r0, V, M, T0, tuple(widths), tuple(starts)])
            buckets[(tuple(widths), tuple(starts))].append(idx)
        slots.append(('multi', idx))

    if build_T_from_M and not closed and metas:
        # Recover every panel's compact-WY T from the already-formed compound
        # Grams M (diagonal 32x32 blocks == V_j^T V_j) with ONE batched trec,
        # replacing the 16 per-panel Gram bmms + 16 trec launches in _os_buildT.
        # Grouped by block width w so trec's fixed-nb kernel applies. Blocks
        # stacked over (meta, panel) -> single launch per width.
        wbuck = _dd(list)                     # w -> [(meta_idx, s)]
        for mi, meta in enumerate(metas):
            _, _, Mm, T0m, widths, starts = meta
            if T0m is not None:
                continue
            b_ = Mm.shape[0]; NB_ = Mm.shape[1]
            meta[3] = torch.zeros(b_, NB_, NB_, device=Mm.device, dtype=torch.float32)
            for s, w in zip(starts, widths):
                wbuck[w].append((mi, s))
        for w, entries in wbuck.items():
            blocks = [metas[mi][2][:, s:s + w, s:s + w] for (mi, s) in entries]
            Gcat = torch.cat(blocks, 0).contiguous()      # (len*b, w, w)
            Tcat = torch.empty_like(Gcat)
            torch.ops.eigh_ops.trec(Gcat, Tcat)
            bb = blocks[0].shape[0]
            for e, (mi, s) in enumerate(entries):
                metas[mi][3][:, s:s + w, s:s + w] = Tcat[e * bb:(e + 1) * bb]

    Tres = [None] * len(metas)
    if closed:
        # closed-form T: one batched triangular solve per NB bucket
        for NB, idxs in buckets.items():
            ng = len(idxs)
            b = metas[idxs[0]][2].shape[0]
            Mst = torch.stack([metas[j][2] for j in idxs], 0).reshape(ng * b, NB, NB)
            diagv = torch.diagonal(Mst, dim1=1, dim2=2)          # ||v_i||^2
            invtau = torch.where(diagv > 0, 0.5 * diagv,
                                 torch.ones_like(diagv))          # 1/tau (null->1)
            R = torch.triu(Mst, 1) + torch.diag_embed(invtau)     # upper-tri
            eye = torch.eye(NB, device=Mst.device,
                            dtype=torch.float32).expand(ng * b, NB, NB)
            Tst = torch.linalg.solve_triangular(R, eye, upper=True)
            Tst = Tst.reshape(ng, b, NB, NB)
            for gi, j in enumerate(idxs):
                Tres[j] = Tst[gi]
    else:
        # batched block-WY recurrence per (widths,starts) signature bucket
        for sig, idxs in buckets.items():
            widths, starts = sig
            nblk = len(widths)
            ng = len(idxs)
            b = metas[idxs[0]][3].shape[0]
            NB = metas[idxs[0]][3].shape[1]
            Mst = torch.stack([metas[j][2] for j in idxs], 0).reshape(ng * b, NB, NB)
            Tst = torch.stack([metas[j][3] for j in idxs], 0).reshape(ng * b, NB, NB)
            for k in range(1, nblk):
                s_k = starts[k]; w_k = widths[k]
                Tk = Tst[:, s_k:s_k + w_k, s_k:s_k + w_k]
                cross = Mst[:, 0:s_k, s_k:s_k + w_k]
                Tul = Tst[:, 0:s_k, 0:s_k]
                Tst[:, 0:s_k, s_k:s_k + w_k] = -torch.bmm(torch.bmm(Tul, cross), Tk)
            Tst = Tst.reshape(ng, b, NB, NB)
            for gi, j in enumerate(idxs):
                Tres[j] = Tst[gi]

    out = []
    for kind, payload in slots:
        if kind == 'single':
            out.append(payload)
        else:
            out.append((metas[payload][0], metas[payload][1], Tres[payload]))
    return out


# Panel-aggregation groups for the WY back-transform: compound this many
# consecutive nb=32 tridiag panels into one wider compact-WY block, so the apply
# runs as fewer/fatter GEMMs. n512 (saturated b640) stays FP32 SIMT (fp16x3 is a
# wash there); n1024/n2048 (SM-starved, launch-bound backT) win with fp16x3 TC on
# the fattened blocks. Tuned by full-kernel A/B (see NOTES.md).
_AGG512 = 4     # n512 backT panel aggregation group
_TAIL512 = 64   # n512 tail-collapse block (last _TAIL512/nb panels -> 1 SMEM kernel)
_AGG1024 = 8    # n1024 backT panel aggregation group
_AGG2048 = 16   # n2048 backT panel aggregation group


def _eigh512_composed(data: input_t, nb: int = 32, compose_tf32_hmin: int = 1 << 30, compose_tf32x2_hmin: int = 1 << 30, defl_zk: float = 1.0, compose_fp16x1_hmax: int = 0) -> output_t:
    b, n, _ = data.shape
    old_tf32 = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = False
    try:
        A, Ah, scale = _prescale(data)
        # agg_buildT: per-panel buildT is redundant with the aggregate's compound
        # Gram (its 32x32 diagonal blocks ARE the per-panel Grams); recover every
        # panel T from M's diagonal in one batched trec (glue agent, -0.6 ms).
        d, e, refl = _os_sytrd(A, nb, tf32_trailing=True, Ah=Ah, agg_buildT=True, tail_m0=_TAIL512)
        W, L = _dc_eigh(d.float(), e.float(), leaf=_LEAF_N, compose_tf32_hmin=compose_tf32_hmin, compose_tf32x2_hmin=compose_tf32x2_hmin, defl_zk=defl_zk, compose_fp16x1_hmax=compose_fp16x1_hmax, res_tol=_DC_RESTOL_RELAX, step_tol=_DC_STEPTOL_RELAX)
        refl = _aggregate_refl(refl, _AGG512, build_T_from_M=True)
        # fp16 tensor-core WY apply + one Newton-Schulz orthogonality purify
        # (NOTES fp16-first idea #1+#5): single-pass fp16 (vs fp32 SIMT / fp16x3)
        # then Q<-Q(1.5I-0.5 Q^TQ) repairs orth. Replaces the fp32 SIMT apply
        # (~6.8 ms). L from D&C kept unchanged (diag(Q^TA0Q) recovery no better).
        Q = _os_back_transform_fp16ns(W, refl)
        L = L * scale.view(b, 1)
    finally:
        torch.backends.cuda.matmul.allow_tf32 = old_tf32
    return Q.contiguous(), L.float().contiguous()


def _eigh1024_composed(data: input_t, nb: int = 32, cluster: int = 2, nb_big: int = 0,
                       compose_tf32_hmin: int = 1 << 30, compose_tf32x2_hmin: int = 1 << 30,
                       back_transform=_os_back_transform_x3, defl_zk: float = 1.0,
                       compose_fp16x1_hmax: int = 0) -> output_t:
    """n1024 composed pipeline: thread-block-cluster one-stage tridiagonal
    reduction (symv rows split across C=2 CTAs/matrix -> fills the SM-starved
    b60 batch) + multi-CTA divide-and-conquer tridiagonal solve + fp16x3 WY
    back-transform. Reduce ~50 ms, D&C ~32 ms, backT ~10 ms => ~90 ms vs
    torch.linalg.eigh 122 ms (1.35x). tf32 only on the BLAS-3 trailing update."""
    b, n, _ = data.shape
    old_tf32 = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = False
    try:
        A, Ah, scale = _prescale(data)
        # n2048 aggregates the reflectors closed-form (Schreiber-Van Loan), which
        # rebuilds each compound-WY T from the assembled Gram and discards the
        # per-panel T -> skip building it in sytrd (drops the deferred Gram+trec).
        d, e, refl = _os_sytrd(A, nb, cluster=cluster, tf32_trailing=True,
                               nb_big=nb_big, Ah=Ah, skip_T=(n == 2048))
        # n2048's single-CTA merge is only its m=64 first level; relaxing its secular
        # tolerance perturbs eta enough to change the later multi-CTA merges'
        # cost (+1.1 ms measured), so n2048 keeps the tight baseline 8e-16/1e-9.
        # n1024's lower levels all benefit from relaxed straggler retirement.
        _rt, _st = (_DC_RESTOL, _DC_STEPTOL) if n == 2048 else (_DC_RESTOL_RELAX, _DC_STEPTOL_RELAX)
        W, L = _dc_eigh(d.float(), e.float(), leaf=32, compose_tf32_hmin=compose_tf32_hmin, compose_tf32x2_hmin=compose_tf32x2_hmin, defl_zk=defl_zk, compose_fp16x1_hmax=compose_fp16x1_hmax, res_tol=_rt, step_tol=_st)
        agg = _AGG2048 if n == 2048 else _AGG1024
        # closed-form (Schreiber-Van Loan triangular-inverse) block-WY build:
        # ONE batched triangular solve vs the launch-bound bmm recurrence. Net
        # win only at n2048 b8 (deeply launch-bound); at n1024 b60 the (ng*b)=240
        # -batch triangular solve is slower than the batched recurrence, so n1024
        # keeps the recurrence (measured +0.2 ms closed).
        refl = _aggregate_refl(refl, agg, closed=(n == 2048))
        # Shape-specific WY back-transform, chosen by the caller (see call sites):
        # n1024 -> _os_back_transform_tf32x2 (full tf32-split apply, fp16x3 accuracy
        # at higher tf32-TC efficiency on its fat NB=256 blocks); n2048 -> the
        # default _os_back_transform_x3 (fp16x3; tf32-full-split is flat at b8).
        Q = back_transform(W, refl)
        L = L * scale.view(b, 1)
    finally:
        torch.backends.cuda.matmul.allow_tf32 = old_tf32
    return Q.contiguous(), L.float().contiguous()


def _eigh352_composed(data: input_t) -> output_t:
    """n352 b40: composed cluster one-stage tridiag -> D&C -> FP32 WY back-
    transform. 6.94 ms vs 7.72 ms block-Jacobi (-10%) at 13.6x eigen margin
    (vs Jacobi's 2.5x).

    Why the tridiagonal is PADDED to 512: the dc:: D&C merge is specialized for
    merge sizes m % 64 == 0 (it crashes for m in {44,88,176,352}); the native
    352 tridiagonal has no leaf*2^k factorisation with m%64==0. So after the
    352->tridiagonal reduction, the (d,e) tridiagonal is embedded into a 512
    problem (leaf=32, K=16, every merge m in {64,128,256,512}) by appending 160
    ISOLATED sentinel eigenvalues on the pad diagonal (value 1e3+i, coupling
    e=0). Sentinels exceed any real eigenvalue (|A|max=1 after prescale ->
    |eig| <= ||A||_1 <= 352 << 1e3) and are decoupled, so the D&C returns them
    as an exact invariant block: the real 352 eigenpairs are precisely the 352
    smallest -> the first 352 columns after the ascending sort, with exactly
    zero eigenvector mass on the pad rows. The 352-sized reduction reflectors
    then back-transform the real 352x352 tridiagonal-eigenvector block.

    cluster=3 splits each matrix's symv reduction across 3 CTAs (40*3 = 120
    CTAs ~ fills 148 SMs in one wave; cluster=4 -> 160 CTAs spills to a 2nd wave
    and regresses). Plain FP32 back-transform (fp16x3 measured slower here) and
    no reflector aggregation (agg=1: 11 nb=32 panels, the plain loop wins).

    NATIVE-352 D&C (no 512-pad): 352 = 22 * 16, so leaf=22 (K=16, a power-of-2
    leaf count) builds a balanced binary merge tree with uniform per-level merge
    sizes m in {44,88,176,352}. steqr_tri solves the 22-wide leaves (N<=32 OK).
    The dc:: merge kernels are already m-generic; the ONLY power-of-2 assumption
    was the bitonic sort in dc_prep, now padded to next_pow2(m) inside SMEM
    (pad slots keyed +inf, later work stays over the real m). This does the
    exact 352-sized D&C instead of the 512-pad's ~2x work."""
    b, n, _ = data.shape
    old_tf32 = torch.backends.cuda.matmul.allow_tf32
    torch.backends.cuda.matmul.allow_tf32 = False
    try:
        A, Ah, scale = _prescale(data)
        d, e, refl = _os_sytrd(A, 32, cluster=3, tf32_trailing=True, Ah=Ah)
        d = d.float(); e = e.float()
        W, L = _dc_eigh(d, e, leaf=22)     # native 352 = 22*16 (K=16 balanced)
        Q = _os_back_transform(W, refl)    # FP32 SIMT WY back-transform
        L = L * scale.view(b, 1)
    finally:
        torch.backends.cuda.matmul.allow_tf32 = old_tf32
    return Q.contiguous(), L.float().contiguous()


def custom_kernel(data: input_t) -> output_t:
    batch, n, _ = data.shape

    if n == 1024:
        # Composed cluster one-stage reduce + multi-CTA D&C beats torch at the
        # SM-starved b60 batch (60 matrices << 148 SMs). Shape-specialized.
        # C=2 (120 CTAs) is the reduction-grid sweet spot: it is fastest on BOTH
        # the local CUDA-13 toolchain AND the grader's CUDA-12.9 toolchain (the
        # kernelbot Modal image builds submissions with nvcc/cicc 12.9). The
        # earlier C=4 was a stale optimum from an intermediate code state and
        # regressed BOTH stacks once the fp16-shadow-symv / dead-zero-fill /
        # fp16x1-compose changes landed: measured n1024-geomean C=4 -> C=2 is
        # 36.0 -> 29.3 ms on CUDA-13 and 56.3 -> 48.8 ms on CUDA-12.9. C is a
        # pure launch-grid knob (byte-identical numerics), so this is a free win
        # and a portability guard (cicc-12.9 gives the cluster kernel much less
        # ILP than cicc-13, so a finer C multiplies its per-CTA overhead). See
        # NOTES "PORTABILITY".
        return _eigh1024_composed(data, cluster=2, compose_fp16x1_hmax=1 << 30,
                                  back_transform=_os_back_transform_fp16ns,
                                  defl_zk=_DEFL_ZK)

    if n == 2048:
        # Same composed one-stage pipeline at n2048 b8: nb=16 (the fp16 panel
        # SMEM fits the 228 KB cap only at nb<=16 for m=2048) + cluster C=8 (8
        # matrices x 8 CTAs = 64 CTAs; C>8 regresses on non-portable GPC co-
        # residency). Reduce ~86 ms, D&C ~30 ms, backT ~22 ms => ~138 ms vs
        # torch.linalg.eigh ~188 ms (1.36x). Shape-specialized.
        return _eigh1024_composed(data, nb=16, cluster=8, nb_big=32,
                                  back_transform=_os_back_transform_fp16ns,
                                  defl_zk=_DEFL_ZK_2048)

    if n == _N512:
        # The composed one-stage tridiag reduction launches one CTA per matrix,
        # so its win depends on the batch saturating the machine (~148 SMs): at
        # the scored b640 it beats the block-Jacobi ~1.3x, but at a small batch
        # the 16-CTA reduction SM-starves and the intra-matrix-parallel Jacobi
        # is faster (measured b16: Jacobi ~23-27 ms vs composed ~31 ms). Route
        # on batch (a shape property, not input data); both paths are correct.
        if batch >= 128:
            return _eigh512_composed(data, compose_fp16x1_hmax=1 << 30, defl_zk=_DEFL_ZK)
        return _eigh512(data)

    if n == 176:
        # Route n176 through the A-parallel block-Jacobi cluster kernel (padded
        # 176->192, NB=12). Element-Jacobi serialises the entire off-block A
        # update on rank0 (the barrier-bound long pole); block-Jacobi distributes
        # it across all consumer warps, so even at 8 sweeps (needed for accuracy)
        # it beats the 7-sweep element-Jacobi. tol2=2e-7 -> 8 sweeps, eigen ~52/200
        # (3.9x margin, > the n352 block path's shipped 2.5x). 3.67 ms -> ~2.9 ms.
        return _block_jacobi_pad(data, 192, 2e-7)

    if n == 352:
        # Composed cluster tridiag -> D&C(512-padded) -> FP32 WY back-transform.
        # 7.72 ms block-Jacobi -> 6.94 ms (-10%) at 13.6x eigen margin. The b40
        # batch (120 CTAs at cluster=3) fills the machine; the block-Jacobi's
        # 32x32 mini-eig producer chain is the wall it replaces. See NOTES.
        return _eigh352_composed(data)

    strategy = SMALL_STRATEGY.get(n)
    if strategy == "jacobi":
        return _jacobi_smem(data, JACOBI_TOL2[n])
    if strategy == "block":
        return _block_jacobi(data, JACOBI_TOL2[n])
    values, vectors = torch.linalg.eigh(data)
    return vectors, values
scrolls · 6785 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