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
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.
cluster
cluster.sync();fused-epilogue
void launch_trail_epilogue(float* A, __half* Ah, const float* P,mbarrier
asm volatile("mbarrier.init.shared::cta.b64 [%0], 1;" :: "r"(smem_u32(p)));mma
tmp = tl.dot(v1t, t)num-warps = 4
a, vv, pairs, active, _N512, _TB512, _P2_512, prec, num_warps=4)shared-memory
__shared__ float sIJ[TES_TS][TES_TS + 1];tma
asm volatile("cp.async.bulk.shared::cluster.global.mbarrier::complete_tx::bytes [%0], [%1], %2, [%3];"vector-width = float4
float4 a4 = *reinterpret_cast<const float4*>(&A[aidx]);warp-specialization
static constexpr int F_ROT = 4 * 16 * 2; // per-producer float2 rotsKernel 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