submission 877561
spectral_otter · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 15508 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-eigh-877561?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:24a88a84fb32662a3fbd94c4b0f24e7f8edf5390b6ba571a31c8dbea11b36c3e
license declaredunknown
license concludedunknown
authorsspectral_otter
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fused-epilogue
void defl_epilogue_(torch::Tensor basis, torch::Tensor perm_cols,mma
accw = tl.dot(tl.trans(vh), xh, accw)num-warps = 8
BN=64, BK=64, num_warps=8)shared-memory
extern __shared__ float shared[];tile-k = 64
BN=64, BK=64, num_warps=8)tile-n = 64
BN=64, BK=64, num_warps=8)vector-width = float4
static_assert(K == 16, "float4 staging below assumes K == 16");warp-specialization
int local_producers = 0;Kernel source
submission.py15508 lines
from __future__ import annotations
import hashlib
import math
import os
from functools import lru_cache
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline
# ===== grafted general cuSOLVER XsyevBatched (from submission20, for the
# involution fast-path; the chain's own xsyev_batched is n32-only) =====
_GEN_EIGH_CPP = r"""
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include <cusolverDn.h>
#include <limits>
#include <memory>
#include <mutex>
#include <string>
#include <unordered_map>
#include <vector>
namespace {
std::string status_name(cusolverStatus_t s) {
switch (s) {
case CUSOLVER_STATUS_SUCCESS: return "SUCCESS";
case CUSOLVER_STATUS_NOT_INITIALIZED: return "NOT_INITIALIZED";
case CUSOLVER_STATUS_ALLOC_FAILED: return "ALLOC_FAILED";
case CUSOLVER_STATUS_INVALID_VALUE: return "INVALID_VALUE";
case CUSOLVER_STATUS_ARCH_MISMATCH: return "ARCH_MISMATCH";
case CUSOLVER_STATUS_EXECUTION_FAILED: return "EXECUTION_FAILED";
case CUSOLVER_STATUS_INTERNAL_ERROR: return "INTERNAL_ERROR";
default: return "UNKNOWN_" + std::to_string((int)s);
}
}
void check(cusolverStatus_t s, const char* c) {
TORCH_CHECK(s == CUSOLVER_STATUS_SUCCESS, c, " failed: ", status_name(s));
}
void validate(const at::Tensor& w) {
TORCH_CHECK(w.is_cuda() && w.scalar_type() == at::kFloat && w.dim() == 3
&& w.is_contiguous(), "need contiguous cuda float32 [b,n,n]");
TORCH_CHECK(w.size(1) == w.size(2) && w.size(0) > 0, "square, positive batch");
}
struct WorkKey { int device, batch, n;
bool operator==(const WorkKey& o) const {
return device == o.device && batch == o.batch && n == o.n; } };
struct WorkKeyHash { size_t operator()(const WorkKey& k) const {
size_t v = std::hash<int>{}(k.device);
v ^= std::hash<int>{}(k.batch) + 0x9e3779b9 + (v << 6) + (v >> 2);
v ^= std::hash<int>{}(k.n) + 0x9e3779b9 + (v << 6) + (v >> 2); return v; } };
struct WorkArea { at::Tensor device_bytes; at::Tensor host_bytes; at::Tensor info;
size_t device_size = 0; size_t host_size = 0; };
struct DeviceState { cusolverDnHandle_t handle = nullptr; cusolverDnParams_t params = nullptr;
std::unordered_map<WorkKey, WorkArea, WorkKeyHash> areas; };
class Resources {
public:
std::vector<at::Tensor> solve(at::Tensor work) {
validate(work);
c10::cuda::CUDAGuard guard(work.device());
std::lock_guard<std::mutex> lock(mutex_);
const int64_t batch = work.size(0);
const int64_t n = work.size(1);
DeviceState& st = state_for(work.get_device());
auto values = at::empty({batch, n}, work.options());
WorkArea& area = area_for(st, work, values);
check(cusolverDnXsyevBatched(
st.handle, st.params, CUSOLVER_EIG_MODE_VECTOR, CUBLAS_FILL_MODE_UPPER,
n, CUDA_R_32F, work.data_ptr<float>(), n, CUDA_R_32F,
values.data_ptr<float>(), CUDA_R_32F,
area.device_bytes.data_ptr(), area.device_size,
area.host_bytes.data_ptr(), area.host_size,
area.info.data_ptr<int>(), batch),
"cusolverDnXsyevBatched");
return {work, values};
}
private:
DeviceState& state_for(int device) {
auto f = states_.find(device);
if (f == states_.end()) {
auto s = std::make_unique<DeviceState>();
check(cusolverDnCreate(&s->handle), "cusolverDnCreate");
check(cusolverDnCreateParams(&s->params), "cusolverDnCreateParams");
f = states_.emplace(device, std::move(s)).first;
}
return *f->second;
}
WorkArea& area_for(DeviceState& st, const at::Tensor& work,
const at::Tensor& values) {
const WorkKey key{work.get_device(), (int)work.size(0), (int)work.size(1)};
auto f = st.areas.find(key);
if (f != st.areas.end()) return f->second;
size_t dsz = 0, hsz = 0;
check(cusolverDnXsyevBatched_bufferSize(
st.handle, st.params, CUSOLVER_EIG_MODE_VECTOR, CUBLAS_FILL_MODE_UPPER,
work.size(1), CUDA_R_32F, work.data_ptr<float>(), work.size(1),
CUDA_R_32F, values.data_ptr<float>(), CUDA_R_32F,
&dsz, &hsz, work.size(0)), "bufferSize");
WorkArea area; area.device_size = dsz; area.host_size = hsz;
area.device_bytes = at::empty({(int64_t)std::max<size_t>(dsz, 1)},
work.options().dtype(at::kByte));
area.host_bytes = at::empty({(int64_t)std::max<size_t>(hsz, 1)},
at::TensorOptions().dtype(at::kByte).device(at::kCPU));
area.info = at::empty({work.size(0)}, work.options().dtype(at::kInt));
return st.areas.emplace(key, std::move(area)).first->second;
}
std::mutex mutex_;
std::unordered_map<int, std::unique_ptr<DeviceState>> states_;
};
Resources& resources() { static auto* v = new Resources(); return *v; }
} // namespace
std::vector<at::Tensor> xsyev_batched(at::Tensor work) {
return resources().solve(std::move(work));
}
"""
@lru_cache(maxsize=1)
def _gen_eigh_ext():
import os
from torch.utils.cpp_extension import load_inline
cuda_home = os.environ.get("CUDA_HOME", "/usr/local/cuda")
return load_inline(
name="gen_eigh_xsyev_batched",
cpp_sources=_GEN_EIGH_CPP,
functions=["xsyev_batched"],
with_cuda=True, # pure-C++ src still needs CUDA headers/link
extra_include_paths=[cuda_home + "/include"],
extra_ldflags=["-lcusolver", "-L" + cuda_home + "/lib64"],
verbose=False,
)
# NVIDIA container images default matmul TF32 ON; the whole chain and the
# banked torch path are strict-FP32 surfaces, so pin IEEE at import.
torch.backends.cuda.matmul.allow_tf32 = False
torch.backends.cudnn.allow_tf32 = False
if hasattr(torch.backends.cuda.matmul, "fp32_precision"):
torch.backends.cuda.matmul.fp32_precision = "ieee"
torch.set_float32_matmul_precision("highest")
_CPP = r"""
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include <cusolverDn.h>
#include <algorithm>
#include <cstdint>
#include <functional>
#include <limits>
#include <memory>
#include <mutex>
#include <string>
#include <unordered_map>
#include <utility>
#include <vector>
namespace {
std::string status_name(cusolverStatus_t status) {
switch (status) {
case CUSOLVER_STATUS_SUCCESS: return "CUSOLVER_STATUS_SUCCESS";
case CUSOLVER_STATUS_NOT_INITIALIZED: return "CUSOLVER_STATUS_NOT_INITIALIZED";
case CUSOLVER_STATUS_ALLOC_FAILED: return "CUSOLVER_STATUS_ALLOC_FAILED";
case CUSOLVER_STATUS_INVALID_VALUE: return "CUSOLVER_STATUS_INVALID_VALUE";
case CUSOLVER_STATUS_ARCH_MISMATCH: return "CUSOLVER_STATUS_ARCH_MISMATCH";
case CUSOLVER_STATUS_EXECUTION_FAILED: return "CUSOLVER_STATUS_EXECUTION_FAILED";
case CUSOLVER_STATUS_INTERNAL_ERROR: return "CUSOLVER_STATUS_INTERNAL_ERROR";
case CUSOLVER_STATUS_MATRIX_TYPE_NOT_SUPPORTED: return "CUSOLVER_STATUS_MATRIX_TYPE_NOT_SUPPORTED";
case CUSOLVER_STATUS_NOT_SUPPORTED: return "CUSOLVER_STATUS_NOT_SUPPORTED";
case CUSOLVER_STATUS_ZERO_PIVOT: return "CUSOLVER_STATUS_ZERO_PIVOT";
case CUSOLVER_STATUS_INVALID_LICENSE: return "CUSOLVER_STATUS_INVALID_LICENSE";
default: return "CUSOLVER_STATUS_UNKNOWN_" + std::to_string(static_cast<int>(status));
}
}
void check(cusolverStatus_t status, const char* call) {
TORCH_CHECK(status == CUSOLVER_STATUS_SUCCESS,
call, " failed: ", status_name(status));
}
void validate(const at::Tensor& work) {
TORCH_CHECK(work.is_cuda(), "input must be CUDA");
TORCH_CHECK(work.scalar_type() == at::kFloat, "input must be float32");
TORCH_CHECK(work.dim() == 3, "input must have shape [batch,n,n]");
TORCH_CHECK(work.size(0) > 0, "batch must be positive");
const int64_t n = work.size(1);
const int64_t batch = work.size(0);
TORCH_CHECK(work.size(2) == n, "direct path requires square matrices");
const bool supported_n = n == 32;
TORCH_CHECK(supported_n, "unsupported matrix size for direct path: ", n);
TORCH_CHECK(work.is_contiguous(), "input clone must be contiguous");
TORCH_CHECK(n <= std::numeric_limits<int64_t>::max() / n,
"n*n overflow");
const int64_t per_matrix = n * n;
TORCH_CHECK(batch <= std::numeric_limits<int32_t>::max() / per_matrix,
"n*lda*batch exceeds INT32_MAX");
}
struct WorkKey {
int device;
int64_t batch;
int64_t n;
bool operator==(const WorkKey& other) const {
return device == other.device && batch == other.batch && n == other.n;
}
};
struct WorkKeyHash {
size_t operator()(const WorkKey& key) const {
size_t value = std::hash<int>{}(key.device);
value ^= std::hash<int64_t>{}(key.batch) + 0x9e3779b9 + (value << 6) + (value >> 2);
value ^= std::hash<int64_t>{}(key.n) + 0x9e3779b9 + (value << 6) + (value >> 2);
return value;
}
};
struct WorkArea {
at::Tensor device_bytes;
at::Tensor host_bytes;
at::Tensor info;
size_t device_size = 0;
size_t host_size = 0;
};
struct DeviceState {
cusolverDnHandle_t handle = nullptr;
cusolverDnParams_t params = nullptr;
std::unordered_map<WorkKey, WorkArea, WorkKeyHash> areas;
};
class Resources {
public:
std::vector<at::Tensor> solve(at::Tensor work) {
validate(work);
c10::cuda::CUDAGuard guard(work.device());
std::lock_guard<std::mutex> lock(mutex_);
const int device = work.get_device();
const int64_t batch = work.size(0);
const int64_t n = work.size(1);
DeviceState& state = state_for(device);
auto values = at::empty({batch, n}, work.options());
WorkArea& area = area_for(state, work, values);
check(cusolverDnXsyevBatched(
state.handle, state.params, CUSOLVER_EIG_MODE_VECTOR,
CUBLAS_FILL_MODE_UPPER, n, CUDA_R_32F, work.data_ptr<float>(), n,
CUDA_R_32F, values.data_ptr<float>(), CUDA_R_32F,
area.device_bytes.data_ptr(), area.device_size,
area.host_bytes.data_ptr(), area.host_size,
area.info.data_ptr<int>(), batch),
"cusolverDnXsyevBatched");
return {work, values};
}
private:
DeviceState& state_for(int device) {
auto found = states_.find(device);
if (found == states_.end()) {
auto state = std::make_unique<DeviceState>();
check(cusolverDnCreate(&state->handle), "cusolverDnCreate");
check(cusolverDnCreateParams(&state->params), "cusolverDnCreateParams");
found = states_.emplace(device, std::move(state)).first;
}
return *found->second;
}
WorkArea& area_for(DeviceState& state,
const at::Tensor& work,
const at::Tensor& values) {
const WorkKey key{work.get_device(), work.size(0), work.size(1)};
auto found = state.areas.find(key);
if (found != state.areas.end()) {
return found->second;
}
size_t device_size = 0;
size_t host_size = 0;
check(cusolverDnXsyevBatched_bufferSize(
state.handle, state.params, CUSOLVER_EIG_MODE_VECTOR,
CUBLAS_FILL_MODE_UPPER, work.size(1), CUDA_R_32F,
work.data_ptr<float>(), work.size(1), CUDA_R_32F,
values.data_ptr<float>(), CUDA_R_32F,
&device_size, &host_size, work.size(0)),
"cusolverDnXsyevBatched_bufferSize");
TORCH_CHECK(device_size <= static_cast<size_t>(std::numeric_limits<int64_t>::max()),
"device workspace exceeds int64 tensor size");
TORCH_CHECK(host_size <= static_cast<size_t>(std::numeric_limits<int64_t>::max()),
"host workspace exceeds int64 tensor size");
WorkArea area;
area.device_size = device_size;
area.host_size = host_size;
area.device_bytes = at::empty(
{static_cast<int64_t>(std::max<size_t>(device_size, 1))},
work.options().dtype(at::kByte));
area.host_bytes = at::empty(
{static_cast<int64_t>(std::max<size_t>(host_size, 1))},
at::TensorOptions().dtype(at::kByte).device(at::kCPU));
area.info = at::empty({work.size(0)}, work.options().dtype(at::kInt));
return state.areas.emplace(key, std::move(area)).first->second;
}
std::mutex mutex_;
std::unordered_map<int, std::unique_ptr<DeviceState>> states_;
};
Resources& resources() {
static auto* value = new Resources();
return *value;
}
} // namespace
std::vector<at::Tensor> xsyev_batched(at::Tensor work) {
return resources().solve(std::move(work));
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
module.def("xsyev_batched", &xsyev_batched);
}
"""
_DIRECT_N = frozenset((32,))
@lru_cache(maxsize=1)
def _extension():
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required")
digest = hashlib.sha256(_CPP.encode()).hexdigest()[:12]
return load_inline(
name=f"eigh_xs_host_{digest}",
cpp_sources=_CPP,
functions=None,
extra_cflags=["-O3", "-std=c++17"],
extra_ldflags=["-lcusolver"],
with_cuda=True,
verbose=False,
)
def _direct_n32(data: torch.Tensor) -> output_t:
module = _extension()
work = data.clone(memory_format=torch.contiguous_format)
vectors_storage, values = module.xsyev_batched(work)
vectors = vectors_storage.transpose(-2, -1)
return vectors, values
_BATCH = 640
_N = 512
_CHOLESKY_DIAG_RATIO_MIN = 1.0e-4
_QR_ORTHOGONALITY_RESIDUAL_MAX = 1.0e-5
_QR_FACTORIZATION_RESIDUAL_MAX = 1.0e-5
_QR_R_DIAG_RATIO_MIN = 1.0e-4
_QR_RESCUE_COHORT_MAX = 32
def _solve(data: torch.Tensor) -> output_t:
values, vectors = torch.linalg.eigh(data)
return vectors, values
_CHAIN_CPP = r"""
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include <cusolverDn.h>
#include <algorithm>
#include <cstdint>
#include <functional>
#include <limits>
#include <memory>
#include <mutex>
#include <string>
#include <unordered_map>
#include <utility>
#include <vector>
// lane_g: batched Jacobi exact solve for the fail-closed n512 cohort.
// Small cohorts hit a per-matrix host loop inside the framework solver
// (~14 ms/matrix measured); the batched routine runs the cohort together.
// Body reused from a prior hosted-correct wrapper; one delta: the
// per-matrix convergence status tensor is returned as a third output so
// the caller can fail closed on any non-converged row.
namespace lane_g {
namespace {
std::string status_name(cusolverStatus_t status) {
switch (status) {
case CUSOLVER_STATUS_SUCCESS: return "CUSOLVER_STATUS_SUCCESS";
case CUSOLVER_STATUS_NOT_INITIALIZED: return "CUSOLVER_STATUS_NOT_INITIALIZED";
case CUSOLVER_STATUS_ALLOC_FAILED: return "CUSOLVER_STATUS_ALLOC_FAILED";
case CUSOLVER_STATUS_INVALID_VALUE: return "CUSOLVER_STATUS_INVALID_VALUE";
case CUSOLVER_STATUS_ARCH_MISMATCH: return "CUSOLVER_STATUS_ARCH_MISMATCH";
case CUSOLVER_STATUS_EXECUTION_FAILED: return "CUSOLVER_STATUS_EXECUTION_FAILED";
case CUSOLVER_STATUS_INTERNAL_ERROR: return "CUSOLVER_STATUS_INTERNAL_ERROR";
case CUSOLVER_STATUS_MATRIX_TYPE_NOT_SUPPORTED: return "CUSOLVER_STATUS_MATRIX_TYPE_NOT_SUPPORTED";
case CUSOLVER_STATUS_NOT_SUPPORTED: return "CUSOLVER_STATUS_NOT_SUPPORTED";
case CUSOLVER_STATUS_ZERO_PIVOT: return "CUSOLVER_STATUS_ZERO_PIVOT";
case CUSOLVER_STATUS_INVALID_LICENSE: return "CUSOLVER_STATUS_INVALID_LICENSE";
default: return "CUSOLVER_STATUS_UNKNOWN_" + std::to_string(static_cast<int>(status));
}
}
void check(cusolverStatus_t status, const char* call) {
TORCH_CHECK(status == CUSOLVER_STATUS_SUCCESS,
call, " failed: ", status_name(status));
}
void validate(const at::Tensor& work) {
TORCH_CHECK(work.is_cuda(), "input must be CUDA");
TORCH_CHECK(work.scalar_type() == at::kFloat, "input must be float32");
TORCH_CHECK(work.dim() == 3, "input must have shape [batch,n,n]");
TORCH_CHECK(work.size(0) > 0, "batch must be positive");
TORCH_CHECK((work.size(1) == 512 || work.size(1) == 1024)
&& work.size(2) == work.size(1),
"direct path is restricted to n=512/1024");
TORCH_CHECK(work.is_contiguous(), "input clone must be contiguous");
const int64_t n = work.size(1);
const int64_t batch = work.size(0);
TORCH_CHECK(n <= std::numeric_limits<int64_t>::max() / n,
"n*n overflow");
const int64_t per_matrix = n * n;
TORCH_CHECK(batch <= std::numeric_limits<int32_t>::max() / per_matrix,
"n*lda*batch exceeds INT32_MAX");
}
struct WorkKey {
int device;
int batch;
int n;
bool operator==(const WorkKey& other) const {
return device == other.device && batch == other.batch && n == other.n;
}
};
struct WorkKeyHash {
size_t operator()(const WorkKey& key) const {
size_t value = std::hash<int>{}(key.device);
value ^= std::hash<int>{}(key.batch) + 0x9e3779b9 + (value << 6) + (value >> 2);
value ^= std::hash<int>{}(key.n) + 0x9e3779b9 + (value << 6) + (value >> 2);
return value;
}
};
struct WorkArea {
at::Tensor device_work;
at::Tensor info;
int lwork = 0;
};
struct DeviceState {
cusolverDnHandle_t handle = nullptr;
syevjInfo_t params = nullptr;
std::unordered_map<WorkKey, WorkArea, WorkKeyHash> areas;
};
class Resources {
public:
std::vector<at::Tensor> solve(at::Tensor work) {
validate(work);
c10::cuda::CUDAGuard guard(work.device());
std::lock_guard<std::mutex> lock(mutex_);
const int device = work.get_device();
const int batch = static_cast<int>(work.size(0));
const int n = static_cast<int>(work.size(1));
DeviceState& state = state_for(device);
auto values = at::empty({batch, n}, work.options());
WorkArea& area = area_for(state, work, values, batch, n);
check(cusolverDnSsyevjBatched(
state.handle, CUSOLVER_EIG_MODE_VECTOR, CUBLAS_FILL_MODE_UPPER,
n, work.data_ptr<float>(), n, values.data_ptr<float>(),
area.device_work.data_ptr<float>(), area.lwork,
area.info.data_ptr<int>(), state.params, batch),
"cusolverDnSsyevjBatched");
return {work, values, area.info};
}
private:
DeviceState& state_for(int device) {
auto found = states_.find(device);
if (found == states_.end()) {
auto state = std::make_unique<DeviceState>();
check(cusolverDnCreate(&state->handle), "cusolverDnCreate");
check(cusolverDnCreateSyevjInfo(&state->params),
"cusolverDnCreateSyevjInfo");
check(cusolverDnXsyevjSetSortEig(state->params, 1),
"cusolverDnXsyevjSetSortEig");
found = states_.emplace(device, std::move(state)).first;
}
return *found->second;
}
WorkArea& area_for(DeviceState& state,
const at::Tensor& work,
const at::Tensor& values,
int batch,
int n) {
const WorkKey key{work.get_device(), batch, n};
auto found = state.areas.find(key);
if (found != state.areas.end()) {
return found->second;
}
int lwork = 0;
check(cusolverDnSsyevjBatched_bufferSize(
state.handle, CUSOLVER_EIG_MODE_VECTOR, CUBLAS_FILL_MODE_UPPER,
n, work.data_ptr<float>(), n, values.data_ptr<float>(),
&lwork, state.params, batch),
"cusolverDnSsyevjBatched_bufferSize");
TORCH_CHECK(lwork >= 0, "negative device workspace size");
WorkArea area;
area.lwork = lwork;
area.device_work = at::empty(
{static_cast<int64_t>(std::max(lwork, 1))}, work.options());
area.info = at::empty({batch}, work.options().dtype(at::kInt));
return state.areas.emplace(key, std::move(area)).first->second;
}
std::mutex mutex_;
std::unordered_map<int, std::unique_ptr<DeviceState>> states_;
};
Resources& resources() {
static auto* value = new Resources();
return *value;
}
void validate_double(const at::Tensor& work) {
TORCH_CHECK(work.is_cuda(), "input must be CUDA");
TORCH_CHECK(work.scalar_type() == at::kDouble, "input must be float64");
TORCH_CHECK(work.dim() == 3, "input must have shape [batch,n,n]");
TORCH_CHECK(work.size(0) > 0, "batch must be positive");
TORCH_CHECK((work.size(1) == 512 || work.size(1) == 1024)
&& work.size(2) == work.size(1),
"direct path is restricted to n=512/1024");
TORCH_CHECK(work.is_contiguous(), "input clone must be contiguous");
const int64_t n = work.size(1);
const int64_t batch = work.size(0);
const int64_t per_matrix = n * n;
TORCH_CHECK(batch <= std::numeric_limits<int32_t>::max() / per_matrix,
"n*lda*batch exceeds INT32_MAX");
}
// Double-precision twin of Resources for the straggler stage: one batched
// call for every row the single-precision pass declined, never a host loop.
class ResourcesDouble {
public:
std::vector<at::Tensor> solve(at::Tensor work) {
validate_double(work);
c10::cuda::CUDAGuard guard(work.device());
std::lock_guard<std::mutex> lock(mutex_);
const int device = work.get_device();
const int batch = static_cast<int>(work.size(0));
const int n = static_cast<int>(work.size(1));
DeviceState& state = state_for(device);
auto values = at::empty({batch, n}, work.options());
WorkArea& area = area_for(state, work, values, batch, n);
check(cusolverDnDsyevjBatched(
state.handle, CUSOLVER_EIG_MODE_VECTOR, CUBLAS_FILL_MODE_UPPER,
n, work.data_ptr<double>(), n, values.data_ptr<double>(),
area.device_work.data_ptr<double>(), area.lwork,
area.info.data_ptr<int>(), state.params, batch),
"cusolverDnDsyevjBatched");
return {work, values, area.info};
}
private:
DeviceState& state_for(int device) {
auto found = states_.find(device);
if (found == states_.end()) {
auto state = std::make_unique<DeviceState>();
check(cusolverDnCreate(&state->handle), "cusolverDnCreate");
check(cusolverDnCreateSyevjInfo(&state->params),
"cusolverDnCreateSyevjInfo");
check(cusolverDnXsyevjSetSortEig(state->params, 1),
"cusolverDnXsyevjSetSortEig");
found = states_.emplace(device, std::move(state)).first;
}
return *found->second;
}
WorkArea& area_for(DeviceState& state,
const at::Tensor& work,
const at::Tensor& values,
int batch,
int n) {
const WorkKey key{work.get_device(), batch, n};
auto found = state.areas.find(key);
if (found != state.areas.end()) {
return found->second;
}
int lwork = 0;
check(cusolverDnDsyevjBatched_bufferSize(
state.handle, CUSOLVER_EIG_MODE_VECTOR, CUBLAS_FILL_MODE_UPPER,
n, work.data_ptr<double>(), n, values.data_ptr<double>(),
&lwork, state.params, batch),
"cusolverDnDsyevjBatched_bufferSize");
TORCH_CHECK(lwork >= 0, "negative device workspace size");
WorkArea area;
area.lwork = lwork;
area.device_work = at::empty(
{static_cast<int64_t>(std::max(lwork, 1))}, work.options());
area.info = at::empty({batch}, work.options().dtype(at::kInt));
return state.areas.emplace(key, std::move(area)).first->second;
}
std::mutex mutex_;
std::unordered_map<int, std::unique_ptr<DeviceState>> states_;
};
ResourcesDouble& resources_double() {
static auto* value = new ResourcesDouble();
return *value;
}
} // namespace
std::vector<at::Tensor> syevj_batched(at::Tensor work) {
return resources().solve(std::move(work));
}
std::vector<at::Tensor> syevj_batched_double(at::Tensor work) {
return resources_double().solve(std::move(work));
}
} // namespace lane_g
namespace lane_a {
void panel_qr_w8_vt_out(const torch::Tensor& input, torch::Tensor packed,
torch::Tensor tau, torch::Tensor householder_v,
torch::Tensor compact_t);
int64_t panel_qr_w8_shared_bytes(int64_t rows);
}
namespace lane_b {
void update_rank16_inplace(at::Tensor matrix, const at::Tensor& l,
const at::Tensor& r, int64_t offset);
void skinny_w_out(const at::Tensor& matrix, const at::Tensor& vt,
at::Tensor w, int64_t offset);
}
namespace lane_c {
void band8_to_tridiagonal_kd8_out(const torch::Tensor& band,
const torch::Tensor& d,
const torch::Tensor& e,
const torch::Tensor& hh,
const torch::Tensor& info,
const std::string& config);
std::vector<int64_t> kd8_resource_report(const std::string& config);
}
namespace lane_d {
void leaf32_saved_chain_run(const at::Tensor& d, const at::Tensor& e,
const at::Tensor& qt, const at::Tensor& info);
}
#define MERGE_DECL(ns) \
namespace ns { \
void secular_merge_fp64_control_run( \
const at::Tensor& poles, const at::Tensor& weights, \
const at::Tensor& rho, const at::Tensor& active_count, \
const at::Tensor& values, const at::Tensor& secular_vectors, \
const at::Tensor& info, const at::Tensor& iterations); \
}
MERGE_DECL(lane_e64)
MERGE_DECL(lane_e128)
MERGE_DECL(lane_e256)
MERGE_DECL(lane_e512)
namespace lane_f {
void backtransform_kd8_inplace(const torch::Tensor& x,
const torch::Tensor& hh,
const torch::Tensor& info,
const std::string& config);
std::vector<int64_t> bt_resource_report(const std::string& config);
}
namespace lane_h {
void midn_eigh_pipeline_out(torch::Tensor a, torch::Tensor q,
torch::Tensor v, torch::Tensor lam,
torch::Tensor info, torch::Tensor stats,
int64_t w_mode);
std::vector<int64_t> midn_resource_attributes(int64_t n);
int64_t midn_k1_dynamic_shared_bytes(int64_t n);
int64_t midn_k1_shared_supported();
int64_t midn_stats_slots();
}
namespace lane_i {
void deflate_(torch::Tensor poles, torch::Tensor weights, torch::Tensor rho,
double ctol, int64_t half, torch::Tensor srcblock,
torch::Tensor d_adj, torch::Tensor z_adj, torch::Tensor active,
torch::Tensor perm_cols, torch::Tensor counts,
torch::Tensor rot_idx, torch::Tensor rot_cs);
void apply_givens_(torch::Tensor basis, torch::Tensor rot_idx,
torch::Tensor rot_cs, torch::Tensor counts);
void secular_(torch::Tensor d_act, torch::Tensor z_act, torch::Tensor rho,
torch::Tensor counts, double stop_scale, torch::Tensor origins,
torch::Tensor taus, torch::Tensor dlambda, torch::Tensor iters,
torch::Tensor info);
void loewner_weights_(torch::Tensor d_act, torch::Tensor z_act,
torch::Tensor rho, torch::Tensor counts,
torch::Tensor origins, torch::Tensor taus,
torch::Tensor zhat);
}
namespace lane_j {
void build_basis_(torch::Tensor children, torch::Tensor order,
torch::Tensor safe, torch::Tensor basis, int64_t half);
void pack_bp_(torch::Tensor basis, torch::Tensor perm_cols,
torch::Tensor counts, torch::Tensor bp, int64_t half,
int64_t branch);
void inv_perm_(torch::Tensor perm_cols, torch::Tensor inv);
void vector_build_packed_(torch::Tensor d_act, torch::Tensor zhat,
torch::Tensor counts, torch::Tensor origins,
torch::Tensor taus, torch::Tensor active,
torch::Tensor inv, torch::Tensor col_out,
torch::Tensor s_top, torch::Tensor s_bot);
void defl_epilogue_(torch::Tensor basis, torch::Tensor perm_cols,
torch::Tensor col_out, torch::Tensor counts,
torch::Tensor out);
void loewner_warp_(torch::Tensor d_act, torch::Tensor z_act,
torch::Tensor rho, torch::Tensor counts,
torch::Tensor origins, torch::Tensor taus,
torch::Tensor zhat);
void vector_build_packed_warp_(torch::Tensor d_act, torch::Tensor zhat,
torch::Tensor counts, torch::Tensor origins,
torch::Tensor taus, torch::Tensor active,
torch::Tensor inv, torch::Tensor col_out,
torch::Tensor s_top, torch::Tensor s_bot);
void secular_recip_(torch::Tensor d_act, torch::Tensor z_act,
torch::Tensor rho, torch::Tensor counts,
double stop_scale, torch::Tensor origins,
torch::Tensor taus, torch::Tensor dlambda,
torch::Tensor iters, torch::Tensor info);
}
namespace lane_k {
void lower_solve_(torch::Tensor d_blocks, torch::Tensor e_blocks,
torch::Tensor lam, torch::Tensor q, torch::Tensor info,
torch::Tensor stats);
std::vector<int64_t> lower_resource_attributes(int64_t n);
}
namespace lane_l {
void panel_qr8(torch::Tensor A, torch::Tensor V, torch::Tensor T,
torch::Tensor tau, int64_t j0);
}
namespace lane_m {
void panel_qr_w8_vt_out(const torch::Tensor& input, torch::Tensor packed,
torch::Tensor tau, torch::Tensor householder_v,
torch::Tensor compact_t);
int64_t panel_qr_w8_shared_bytes(int64_t rows);
}
namespace lane_n {
void band8_to_tridiagonal_kd8_out(const torch::Tensor& band,
const torch::Tensor& d,
const torch::Tensor& e,
const torch::Tensor& hh,
const torch::Tensor& info,
const std::string& config);
std::vector<int64_t> kd8_resource_report(const std::string& config);
}
namespace lane_o {
void backtransform_kd8_inplace(const torch::Tensor& x,
const torch::Tensor& hh,
const torch::Tensor& info,
const std::string& config);
std::vector<int64_t> bt_resource_report(const std::string& config);
}
namespace lane_p {
void n1024_merge_polished_run(
const at::Tensor& poles, const at::Tensor& weights,
const at::Tensor& rho, const at::Tensor& active_count,
const at::Tensor& values, const at::Tensor& roots64,
const at::Tensor& root_status, const at::Tensor& root_iterations,
const at::Tensor& merge_mode, const at::Tensor& repair_ids,
const at::Tensor& repair_count, const at::Tensor& deflated,
const at::Tensor& info, const at::Tensor& maximum_iterations,
const at::Tensor& updated32, const at::Tensor& updated64,
const at::Tensor& recon_status, const at::Tensor& escalated,
const at::Tensor& escalation_ids, const at::Tensor& escalation_count,
const at::Tensor& secular_vectors, const at::Tensor& route_summary);
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
module.def("panel_qr_w8_vt_out", &lane_a::panel_qr_w8_vt_out);
module.def("panel_qr_w8_shared_bytes", &lane_a::panel_qr_w8_shared_bytes);
module.def("update_rank16_inplace", &lane_b::update_rank16_inplace);
module.def("skinny_w_out", &lane_b::skinny_w_out);
module.def("band8_to_tridiagonal_kd8_out",
&lane_c::band8_to_tridiagonal_kd8_out);
module.def("kd8_resource_report", &lane_c::kd8_resource_report);
module.def("leaf32_run", &lane_d::leaf32_saved_chain_run);
module.def("merge64_solve", &lane_e64::secular_merge_fp64_control_run);
module.def("merge128_solve", &lane_e128::secular_merge_fp64_control_run);
module.def("merge256_solve", &lane_e256::secular_merge_fp64_control_run);
module.def("merge512_solve", &lane_e512::secular_merge_fp64_control_run);
module.def("backtransform_kd8_inplace", &lane_f::backtransform_kd8_inplace);
module.def("bt_resource_report", &lane_f::bt_resource_report);
module.def("syevj_batched", &lane_g::syevj_batched);
module.def("syevj_batched_double", &lane_g::syevj_batched_double);
module.def("midn_eigh_pipeline_out", &lane_h::midn_eigh_pipeline_out);
module.def("midn_resource_attributes", &lane_h::midn_resource_attributes);
module.def("midn_k1_dynamic_shared_bytes",
&lane_h::midn_k1_dynamic_shared_bytes);
module.def("midn_k1_shared_supported", &lane_h::midn_k1_shared_supported);
module.def("midn_stats_slots", &lane_h::midn_stats_slots);
module.def("deflate_", &lane_i::deflate_);
module.def("apply_givens_", &lane_i::apply_givens_);
module.def("secular_", &lane_i::secular_);
module.def("loewner_weights_", &lane_i::loewner_weights_);
module.def("build_basis_", &lane_j::build_basis_);
module.def("pack_bp_", &lane_j::pack_bp_);
module.def("inv_perm_", &lane_j::inv_perm_);
module.def("vector_build_packed_", &lane_j::vector_build_packed_);
module.def("defl_epilogue_", &lane_j::defl_epilogue_);
module.def("loewner_warp_", &lane_j::loewner_warp_);
module.def("vector_build_packed_warp_",
&lane_j::vector_build_packed_warp_);
module.def("secular_recip_", &lane_j::secular_recip_);
module.def("lower_solve_", &lane_k::lower_solve_);
module.def("lower_resource_attributes",
&lane_k::lower_resource_attributes);
module.def("panel_qr8", &lane_l::panel_qr8);
module.def("panel1024_qr_w8_vt_out", &lane_m::panel_qr_w8_vt_out);
module.def("panel1024_qr_w8_shared_bytes",
&lane_m::panel_qr_w8_shared_bytes);
module.def("band8_to_tridiagonal_kd8_n1024_out",
&lane_n::band8_to_tridiagonal_kd8_out);
module.def("kd8_n1024_resource_report", &lane_n::kd8_resource_report);
module.def("q2occ_backtransform_inplace",
&lane_o::backtransform_kd8_inplace);
module.def("q2occ_resource_report", &lane_o::bt_resource_report);
module.def("k0k6_run_polished_", &lane_p::n1024_merge_polished_run);
}
"""
_CHAIN_CU = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <c10/cuda/CUDAException.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <math_constants.h>
#include <array>
#include <cfloat>
#include <cmath>
#include <cstdint>
#include <limits>
#include <mutex>
#include <string>
#include <unordered_set>
#include <vector>
#define FULL_MASK 0xffffffffu
#define CHECK_IN(x) TORCH_CHECK(x.is_cuda() && x.is_contiguous(), #x)
namespace lane_a {
namespace {
__device__ __forceinline__ float warp_sum(float value) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
value += __shfl_xor_sync(FULL_MASK, value, offset);
}
return value;
}
__device__ __forceinline__ void house_coeffs(
float alpha, float sigma, float* coeffs) {
if (sigma <= 0.0f) {
coeffs[0] = 0.0f;
coeffs[1] = 0.0f;
coeffs[2] = alpha;
} else {
const float beta = -copysignf(sqrtf(fmaf(alpha, alpha, sigma)), alpha);
coeffs[0] = (beta - alpha) / beta;
coeffs[1] = 1.0f / (alpha - beta);
coeffs[2] = beta;
}
}
template <int Threads, int Width>
__device__ void panel_core(
float* panel,
int leading_dimension,
int rows,
float* coeffs,
float* gammas,
float* taus,
float* scratch) {
const int lane = threadIdx.x & 31;
const int warp = threadIdx.x >> 5;
constexpr int warps = Threads >> 5;
float partial = 0.0f;
for (int row = 1 + threadIdx.x; row < rows; row += Threads) {
const float value = panel[row];
partial = fmaf(value, value, partial);
}
partial = warp_sum(partial);
if (lane == 0) {
scratch[warp] = partial;
}
__syncthreads();
if (threadIdx.x == 0) {
float sigma = 0.0f;
for (int index = 0; index < warps; ++index) {
sigma += scratch[index];
}
house_coeffs(panel[0], sigma, coeffs);
}
__syncthreads();
#pragma unroll
for (int column = 0; column < Width; ++column) {
const float* current_coeffs = coeffs + 4 * (column & 1);
float* next_coeffs = coeffs + 4 * ((column + 1) & 1);
const float tau = current_coeffs[0];
const float gamma = current_coeffs[1];
const float beta = current_coeffs[2];
float* current = panel + static_cast<long>(column) * leading_dimension;
if (threadIdx.x == 0) {
gammas[column] = gamma;
taus[column] = tau;
}
for (int target_column = column + 1 + warp;
target_column < Width;
target_column += warps) {
float* target =
panel + static_cast<long>(target_column) * leading_dimension;
float dot = lane == 0 ? target[column] : 0.0f;
float tail_dot = 0.0f;
for (int row = column + 1 + lane; row < rows; row += 32) {
tail_dot = fmaf(current[row], target[row], tail_dot);
}
dot += gamma * tail_dot;
dot = warp_sum(dot);
const float weight = tau * dot;
float next_alpha = 0.0f;
float next_sigma = 0.0f;
if (lane == 0) {
target[column] -= weight;
}
const float scaled_weight = weight * gamma;
for (int row = column + 1 + lane; row < rows; row += 32) {
const float updated = fmaf(-scaled_weight, current[row], target[row]);
target[row] = updated;
if (target_column == column + 1) {
if (row == column + 1) {
next_alpha = updated;
} else {
next_sigma = fmaf(updated, updated, next_sigma);
}
}
}
if (target_column == column + 1) {
next_sigma = warp_sum(next_sigma);
if (lane == 0) {
house_coeffs(next_alpha, next_sigma, next_coeffs);
}
}
}
if (threadIdx.x == 0) {
current[column] = beta;
}
__syncthreads();
}
}
template <int Threads, int Width>
__device__ void pair_dots(
const float* shared_panel,
int leading_dimension,
int rows,
const float* gammas,
float* gram_upper) {
const int lane = threadIdx.x & 31;
const int warp = threadIdx.x >> 5;
constexpr int warps = Threads >> 5;
constexpr int pairs = Width * (Width - 1) / 2;
for (int pair = warp; pair < pairs; pair += warps) {
int column = static_cast<int>(
(1.0f + sqrtf(1.0f + 8.0f * static_cast<float>(pair))) * 0.5f);
while (column * (column - 1) / 2 > pair) {
--column;
}
while ((column + 1) * column / 2 <= pair) {
++column;
}
const int prior = pair - column * (column - 1) / 2;
const float* prior_column =
shared_panel + static_cast<long>(prior) * leading_dimension;
const float* current_column =
shared_panel + static_cast<long>(column) * leading_dimension;
float tail = 0.0f;
for (int row = column + 1 + lane; row < rows; row += 32) {
tail = fmaf(prior_column[row], current_column[row], tail);
}
tail = warp_sum(tail);
if (lane == 0) {
gram_upper[prior * Width + column] =
gammas[prior] * prior_column[column]
+ gammas[prior] * gammas[column] * tail;
}
}
}
template <int Width>
__device__ void t_recurrence(
const float* gram_upper,
const float* taus,
float* compact_t) {
const int lane = threadIdx.x & 31;
#pragma unroll
for (int column = 0; column < Width; ++column) {
const float tau = taus[column];
if (lane < column) {
float value = 0.0f;
for (int inner = lane; inner < column; ++inner) {
value = fmaf(
compact_t[lane * Width + inner],
gram_upper[inner * Width + column],
value);
}
compact_t[lane * Width + column] = -tau * value;
} else if (lane == column) {
compact_t[column * Width + column] = tau;
} else if (lane < Width) {
compact_t[lane * Width + column] = 0.0f;
}
__syncwarp();
}
}
template <int Threads, int Width>
__global__ void panel_qr_narrow_vt_kernel(
const float* __restrict__ input,
float* __restrict__ packed,
float* __restrict__ tau,
float* __restrict__ householder_v,
float* __restrict__ compact_t,
int rows) {
const int leading_dimension = rows | 1;
extern __shared__ float shared[];
float* panel = shared;
float* gammas = panel + static_cast<long>(leading_dimension) * Width;
float* taus = gammas + Width;
float* coeffs = taus + Width;
float* scratch = coeffs + 8;
float* gram_upper = scratch + 32;
float* shared_t = gram_upper + Width * Width;
const long batch = blockIdx.x;
const float* input_batch = input + batch * static_cast<long>(rows) * Width;
float* packed_batch = packed + batch * static_cast<long>(rows) * Width;
float* tau_batch = tau + batch * Width;
float* v_batch = householder_v + batch * static_cast<long>(rows) * Width;
float* t_batch = compact_t + batch * Width * Width;
for (int index = threadIdx.x; index < rows * Width; index += Threads) {
const int row = index / Width;
const int column = index - row * Width;
panel[static_cast<long>(column) * leading_dimension + row] =
input_batch[index];
}
__syncthreads();
panel_core<Threads, Width>(
panel, leading_dimension, rows, coeffs, gammas, taus, scratch);
pair_dots<Threads, Width>(
panel, leading_dimension, rows, gammas, gram_upper);
__syncthreads();
if ((threadIdx.x >> 5) == 0) {
t_recurrence<Width>(gram_upper, taus, shared_t);
}
__syncthreads();
for (int column = threadIdx.x; column < Width; column += Threads) {
tau_batch[column] = taus[column];
}
for (int index = threadIdx.x; index < rows * Width; index += Threads) {
const int row = index / Width;
const int column = index - row * Width;
const float raw = panel[static_cast<long>(column) * leading_dimension + row];
packed_batch[index] = row > column ? gammas[column] * raw : raw;
v_batch[index] = row < column
? 0.0f
: (row == column ? 1.0f : gammas[column] * raw);
}
for (int index = threadIdx.x; index < Width * Width; index += Threads) {
t_batch[index] = shared_t[index];
}
}
constexpr int kWidth = 8;
constexpr int kThreads = 256;
constexpr int kMaxRows = 512 - kWidth;
// Spark (GB10, sm_121) correctness runs cap opt-in dynamic shared memory near
// 99 KB per CTA; B200 (sm_100) allows far more. Budget check at bind time.
constexpr size_t kSparkSharedBudget = 99 * 1024;
size_t shared_bytes_for_rows(int rows) {
const int leading_dimension = rows | 1;
return (static_cast<size_t>(leading_dimension) * kWidth
+ 2 * kWidth + 8 + 32 + 2 * kWidth * kWidth) * sizeof(float);
}
} // namespace
void panel_qr_w8_vt_out(
const torch::Tensor& input,
torch::Tensor packed,
torch::Tensor tau,
torch::Tensor householder_v,
torch::Tensor compact_t) {
TORCH_CHECK(input.is_cuda(), "input must be CUDA");
TORCH_CHECK(input.scalar_type() == torch::kFloat32, "input must be FP32");
TORCH_CHECK(input.is_contiguous(), "input must be contiguous");
TORCH_CHECK(input.dim() == 3, "input must have shape [batch, rows, 8]");
TORCH_CHECK(input.size(2) == kWidth, "panel width must be 8");
const long rows = input.size(1);
TORCH_CHECK(rows >= kWidth && rows <= kMaxRows && rows % kWidth == 0,
"rows must be a multiple of 8 in [8, 504], got ", rows);
TORCH_CHECK(packed.sizes() == input.sizes(), "packed shape mismatch");
TORCH_CHECK(packed.scalar_type() == torch::kFloat32 && packed.is_contiguous(),
"packed must be contiguous FP32");
TORCH_CHECK(tau.dim() == 2 && tau.size(0) == input.size(0)
&& tau.size(1) == kWidth,
"tau shape mismatch");
TORCH_CHECK(tau.scalar_type() == torch::kFloat32 && tau.is_contiguous(),
"tau must be contiguous FP32");
TORCH_CHECK(householder_v.sizes() == input.sizes(), "V shape mismatch");
TORCH_CHECK(householder_v.scalar_type() == torch::kFloat32
&& householder_v.is_contiguous(),
"V must be contiguous FP32");
TORCH_CHECK(compact_t.dim() == 3 && compact_t.size(0) == input.size(0)
&& compact_t.size(1) == kWidth && compact_t.size(2) == kWidth,
"T shape mismatch");
TORCH_CHECK(compact_t.scalar_type() == torch::kFloat32
&& compact_t.is_contiguous(),
"T must be contiguous FP32");
TORCH_CHECK(packed.device() == input.device()
&& tau.device() == input.device()
&& householder_v.device() == input.device()
&& compact_t.device() == input.device(),
"all tensors must be on the input device");
const size_t shared_bytes = shared_bytes_for_rows(static_cast<int>(rows));
TORCH_CHECK(shared_bytes <= kSparkSharedBudget,
"shared budget exceeded: ", shared_bytes, " > ",
kSparkSharedBudget);
c10::cuda::CUDAGuard guard(input.device());
auto kernel = panel_qr_narrow_vt_kernel<kThreads, kWidth>;
// Max shared for rows=504 is ~16.6 KB, under the 48 KB static limit on both
// sm_100 and sm_121, so no cudaFuncSetAttribute opt-in is required.
kernel<<<input.size(0), kThreads, shared_bytes>>>(
input.data_ptr<float>(),
packed.data_ptr<float>(),
tau.data_ptr<float>(),
householder_v.data_ptr<float>(),
compact_t.data_ptr<float>(),
static_cast<int>(rows));
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
int64_t panel_qr_w8_shared_bytes(int64_t rows) {
return static_cast<int64_t>(shared_bytes_for_rows(static_cast<int>(rows)));
}
} // namespace lane_a
namespace lane_b {
namespace {
constexpr int kTile = 32;
constexpr int kTransposeStride = kTile + 1;
__device__ __forceinline__ void upper_tile_coordinates(
int linear, int tiles, int& row_tile, int& col_tile) {
row_tile = 0;
int row_width = tiles;
while (linear >= row_width) {
linear -= row_width;
++row_tile;
--row_width;
}
col_tile = row_tile + linear;
}
template <int K>
__global__ __launch_bounds__(256)
void symmetric_rank_update_inplace_f32(
float* __restrict__ a,
const float* __restrict__ l,
const float* __restrict__ r,
int n,
int offset,
int m) {
static_assert(K == 16, "float4 staging below assumes K == 16");
__shared__ float sl[kTile][K];
__shared__ float sr[K][kTile];
__shared__ float completed[kTile][kTransposeStride];
const int tiles = (m + kTile - 1) / kTile;
int row_tile;
int col_tile;
upper_tile_coordinates(static_cast<int>(blockIdx.x), tiles, row_tile,
col_tile);
const int row_base = row_tile * kTile;
const int col_base = col_tile * kTile;
const int batch = static_cast<int>(blockIdx.z);
const int tid = static_cast<int>(threadIdx.x);
const int ty = tid >> 4;
const int tx = tid & 15;
const size_t a_base = static_cast<size_t>(batch) * n * n;
const size_t l_base = static_cast<size_t>(batch) * m * K;
const size_t r_base = static_cast<size_t>(batch) * K * m;
// Single-shot float4 staging of the full L and R tiles. L rows are K=16
// floats = 4 float4; R rows are float4-aligned because m % 8 == 0 and
// col_base % 32 == 0 (offset/n alignment validated host-side).
const float4 zero4 = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
if (tid < 128) {
const int lr = tid >> 2;
const int lq = tid & 3;
const float4 value = (row_base + lr < m)
? reinterpret_cast<const float4*>(
l + l_base + static_cast<size_t>(row_base + lr) * K)[lq]
: zero4;
sl[lr][lq * 4 + 0] = value.x;
sl[lr][lq * 4 + 1] = value.y;
sl[lr][lq * 4 + 2] = value.z;
sl[lr][lq * 4 + 3] = value.w;
} else {
const int t2 = tid - 128;
const int rr = t2 >> 3;
const int rq = t2 & 7;
const float4 value = (col_base + rq * 4 < m)
? reinterpret_cast<const float4*>(
r + r_base + static_cast<size_t>(rr) * m + col_base)[rq]
: zero4;
sr[rr][rq * 4 + 0] = value.x;
sr[rr][rq * 4 + 1] = value.y;
sr[rr][rq * 4 + 2] = value.z;
sr[rr][rq * 4 + 3] = value.w;
}
__syncthreads();
float acc00 = 0.0f;
float acc01 = 0.0f;
float acc10 = 0.0f;
float acc11 = 0.0f;
#pragma unroll
for (int kk = 0; kk < K; ++kk) {
const float l0 = sl[ty][kk];
const float l1 = sl[ty + 16][kk];
const float r0 = sr[kk][tx];
const float r1 = sr[kk][tx + 16];
acc00 = fmaf(l0, r0, acc00);
acc01 = fmaf(l0, r1, acc01);
acc10 = fmaf(l1, r0, acc10);
acc11 = fmaf(l1, r1, acc11);
}
// Park all four accumulators in shared so the epilogue can pick its own
// (vector-friendly) thread-to-element mapping.
float* completed_ptr = &completed[0][0];
completed_ptr[ty * kTransposeStride + tx] = acc00;
completed_ptr[ty * kTransposeStride + tx + 16] = acc01;
completed_ptr[(ty + 16) * kTransposeStride + tx] = acc10;
completed_ptr[(ty + 16) * kTransposeStride + tx + 16] = acc11;
__syncthreads();
const bool diagonal_tile = row_tile == col_tile;
const bool interior =
!diagonal_tile && (row_base + kTile <= m) && (col_base + kTile <= m);
if (interior) {
// Every element of this tile is strictly upper and in bounds.
const int er = tid >> 3; // 0..31
const int ec = (tid & 7) << 2; // 0,4,...,28
float* upper_row = a + a_base
+ static_cast<size_t>(offset + row_base + er) * n + offset + col_base;
float4 value = *reinterpret_cast<float4*>(upper_row + ec);
value.x -= completed_ptr[er * kTransposeStride + ec + 0];
value.y -= completed_ptr[er * kTransposeStride + ec + 1];
value.z -= completed_ptr[er * kTransposeStride + ec + 2];
value.w -= completed_ptr[er * kTransposeStride + ec + 3];
*reinterpret_cast<float4*>(upper_row + ec) = value;
__syncthreads();
completed_ptr[er * kTransposeStride + ec + 0] = value.x;
completed_ptr[er * kTransposeStride + ec + 1] = value.y;
completed_ptr[er * kTransposeStride + ec + 2] = value.z;
completed_ptr[er * kTransposeStride + ec + 3] = value.w;
__syncthreads();
// Mirror tile row (col_base+er), cols (row_base+ec..ec+3) <- transpose.
float* mirror_row = a + a_base
+ static_cast<size_t>(offset + col_base + er) * n + offset + row_base;
const float4 mirrored = make_float4(
completed_ptr[(ec + 0) * kTransposeStride + er],
completed_ptr[(ec + 1) * kTransposeStride + er],
completed_ptr[(ec + 2) * kTransposeStride + er],
completed_ptr[(ec + 3) * kTransposeStride + er]);
*reinterpret_cast<float4*>(mirror_row + ec) = mirrored;
return;
}
// Diagonal / edge tiles: v1 scalar guarded path, accs from registers.
const float accs[4] = {acc00, acc01, acc10, acc11};
#pragma unroll
for (int part = 0; part < 4; ++part) {
const int local_row = ty + (part >> 1) * 16;
const int local_col = tx + (part & 1) * 16;
const int row = row_base + local_row;
const int col = col_base + local_col;
if (row <= col && col < m) {
const size_t index = a_base
+ static_cast<size_t>(offset + row) * n + offset + col;
const float value = a[index] - accs[part];
a[index] = value;
completed_ptr[local_row * kTransposeStride + local_col] = value;
}
}
__syncthreads();
#pragma unroll
for (int part = 0; part < 4; ++part) {
const int local_row = ty + (part >> 1) * 16;
const int local_col = tx + (part & 1) * 16;
if (diagonal_tile && local_row <= local_col) {
continue;
}
const int row = col_base + local_row; // mirrored target (lower)
const int col = row_base + local_col;
if (row < m && col < m) {
a[a_base + static_cast<size_t>(offset + row) * n + offset + col] =
completed_ptr[local_col * kTransposeStride + local_row];
}
}
}
// v2 round, W-formation lever (separately strippable; justified by stage
// evidence showing the cuBLAS n=8 skinny GEMM on the lda-strided trailing
// view running ~2.8x above its traffic floor while the update sits at
// roofline). Computes W1 = A22 @ VT with A22 read zero-copy from the parent
// matrix (lda = n, offset), one float4-coalesced pass, K-ascending
// deterministic strict-FP32 accumulation.
constexpr int kWTile = 32;
constexpr int kWCols = 8;
__global__ __launch_bounds__(256)
void skinny_w_gemm_f32(
const float* __restrict__ a,
const float* __restrict__ vt,
float* __restrict__ w,
int n,
int offset,
int m) {
__shared__ float sa[kWTile][kWTile + 1];
__shared__ float sv[kWTile][kWCols];
const int row_base = static_cast<int>(blockIdx.x) * kWTile;
const int batch = static_cast<int>(blockIdx.z);
const int tid = static_cast<int>(threadIdx.x);
const int row = tid >> 3; // 0..31
const int col = tid & 7; // 0..7
const size_t a_base = static_cast<size_t>(batch) * n * n;
const size_t v_base = static_cast<size_t>(batch) * m * kWCols;
const float4 zero4 = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
float acc = 0.0f;
for (int k0 = 0; k0 < m; k0 += kWTile) {
// A tile: 32 rows x 32 k-columns = 256 float4, one per thread.
{
const int ar = tid >> 3;
const int ac = (tid & 7) << 2;
// m % 8 == 0 and ac % 4 == 0, so a started chunk of 4 never straddles m.
const float4 value = (row_base + ar < m && k0 + ac < m)
? *reinterpret_cast<const float4*>(
a + a_base
+ static_cast<size_t>(offset + row_base + ar) * n
+ offset + k0 + ac)
: zero4;
sa[ar][ac + 0] = value.x;
sa[ar][ac + 1] = value.y;
sa[ar][ac + 2] = value.z;
sa[ar][ac + 3] = value.w;
}
// VT tile: 32 k-rows x 8 = 64 float4.
if (tid < 64) {
const int vr = tid >> 1;
const int vq = (tid & 1) << 2;
const float4 value = (k0 + vr < m)
? *reinterpret_cast<const float4*>(
vt + v_base + static_cast<size_t>(k0 + vr) * kWCols + vq)
: zero4;
sv[vr][vq + 0] = value.x;
sv[vr][vq + 1] = value.y;
sv[vr][vq + 2] = value.z;
sv[vr][vq + 3] = value.w;
}
__syncthreads();
// Zero-padded tiles keep the fixed-trip loop exact: fmaf(0,0,acc) == acc.
#pragma unroll
for (int kk = 0; kk < kWTile; ++kk) {
acc = fmaf(sa[row][kk], sv[kk][col], acc);
}
__syncthreads();
}
if (row_base + row < m) {
w[v_base + static_cast<size_t>(row_base + row) * kWCols + col] = acc;
}
}
constexpr int kK = 16;
} // namespace
void update_rank16_inplace(
at::Tensor matrix,
const at::Tensor& l,
const at::Tensor& r,
int64_t offset) {
TORCH_CHECK(matrix.is_cuda() && l.is_cuda() && r.is_cuda(),
"all tensors must be CUDA tensors");
TORCH_CHECK(matrix.scalar_type() == at::kFloat
&& l.scalar_type() == at::kFloat
&& r.scalar_type() == at::kFloat,
"all tensors must be float32");
TORCH_CHECK(matrix.is_contiguous() && l.is_contiguous() && r.is_contiguous(),
"all tensors must be contiguous");
TORCH_CHECK(matrix.dim() == 3 && matrix.size(1) == matrix.size(2),
"matrix must be a batch of square matrices");
const auto n = matrix.size(1);
TORCH_CHECK(n % 4 == 0, "n must be float4-aligned");
TORCH_CHECK(offset >= 0 && offset < n, "offset out of range");
TORCH_CHECK(offset % 8 == 0, "offset must be a multiple of 8");
const auto m = n - offset;
TORCH_CHECK(m % 8 == 0, "trailing size must be a multiple of 8");
TORCH_CHECK(l.dim() == 3 && l.size(0) == matrix.size(0) && l.size(1) == m
&& l.size(2) == kK,
"l must have shape [batch,m,16]");
TORCH_CHECK(r.dim() == 3 && r.size(0) == matrix.size(0) && r.size(1) == kK
&& r.size(2) == m,
"r must have shape [batch,16,m]");
TORCH_CHECK(matrix.get_device() == l.get_device()
&& matrix.get_device() == r.get_device(),
"all tensors must be on one device");
c10::cuda::CUDAGuard guard(matrix.device());
const int tiles = static_cast<int>((m + kTile - 1) / kTile);
const int upper_tiles = tiles * (tiles + 1) / 2;
const dim3 grid(upper_tiles, 1, static_cast<unsigned>(matrix.size(0)));
symmetric_rank_update_inplace_f32<kK><<<grid, 256, 0>>>(
matrix.data_ptr<float>(), l.data_ptr<float>(), r.data_ptr<float>(),
static_cast<int>(n), static_cast<int>(offset), static_cast<int>(m));
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
void skinny_w_out(
const at::Tensor& matrix,
const at::Tensor& vt,
at::Tensor w,
int64_t offset) {
TORCH_CHECK(matrix.is_cuda() && vt.is_cuda() && w.is_cuda(),
"all tensors must be CUDA tensors");
TORCH_CHECK(matrix.scalar_type() == at::kFloat
&& vt.scalar_type() == at::kFloat
&& w.scalar_type() == at::kFloat,
"all tensors must be float32");
TORCH_CHECK(matrix.is_contiguous() && vt.is_contiguous()
&& w.is_contiguous(),
"all tensors must be contiguous");
TORCH_CHECK(matrix.dim() == 3 && matrix.size(1) == matrix.size(2),
"matrix must be a batch of square matrices");
const auto n = matrix.size(1);
TORCH_CHECK(n % 4 == 0, "n must be float4-aligned");
TORCH_CHECK(offset >= 0 && offset < n, "offset out of range");
TORCH_CHECK(offset % 8 == 0, "offset must be a multiple of 8");
const auto m = n - offset;
TORCH_CHECK(m % 8 == 0, "trailing size must be a multiple of 8");
TORCH_CHECK(vt.dim() == 3 && vt.size(0) == matrix.size(0) && vt.size(1) == m
&& vt.size(2) == kWCols,
"vt must have shape [batch,m,8]");
TORCH_CHECK(w.sizes() == vt.sizes(), "w must match vt shape");
TORCH_CHECK(matrix.get_device() == vt.get_device()
&& matrix.get_device() == w.get_device(),
"all tensors must be on one device");
c10::cuda::CUDAGuard guard(matrix.device());
const int row_tiles = static_cast<int>((m + kWTile - 1) / kWTile);
const dim3 grid(row_tiles, 1, static_cast<unsigned>(matrix.size(0)));
skinny_w_gemm_f32<<<grid, 256, 0>>>(
matrix.data_ptr<float>(), vt.data_ptr<float>(), w.data_ptr<float>(),
static_cast<int>(n), static_cast<int>(offset), static_cast<int>(m));
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
} // namespace lane_b
namespace lane_c {
namespace {
constexpr int kN = 512;
constexpr int kKD = 8;
constexpr int kInputRows = kKD + 1;
constexpr int kWorkRows = 2 * kKD + 1;
constexpr int kTaskWarps = 32;
constexpr int kThreads = kTaskWarps * 32;
constexpr int kReflectors = 16512; // sum ceil((511-s)/8), s=1..510
constexpr int kConsumers = 32450; // kernel-loop emulation (see README)
constexpr int kTasks = 32961; // 32,960 loop tasks + terminal no-op
constexpr int kRecordWidth = kKD; // [tau, v1..v7]
constexpr int kSweepPrefixEntries = kN - 1;
constexpr int kScalarsPerMatrixTask = 8;
constexpr int kMatrixWorkPadding = 4;
constexpr int kMatrixWorkStride = kN * kWorkRows + kMatrixWorkPadding;
constexpr int kScalarStride = kScalarsPerMatrixTask + 1;
constexpr int kCounterIntsPerMatrix = 4;
__device__ __constant__ uint16_t kSweepPrefix[kSweepPrefixEntries];
static_assert(kMatrixWorkStride % 32 == 4,
"matrix stride must break bank aliasing");
static_assert(kScalarStride % 32 == 9,
"scalar stride must break bank aliasing");
template <int M, int S>
struct Config {
static_assert(M * S <= 32, "subgroups must fit one warp");
static constexpr size_t work_floats() {
return static_cast<size_t>(M) * kMatrixWorkStride;
}
static constexpr size_t vector_floats() {
return static_cast<size_t>(kTaskWarps) * M * kKD;
}
static constexpr size_t scalar_floats() {
return static_cast<size_t>(kTaskWarps) * M * kScalarStride;
}
static constexpr size_t shared_bytes() {
return (work_floats() + 2 * vector_floats() + scalar_floats())
* sizeof(float)
+ static_cast<size_t>(M * kCounterIntsPerMatrix) * sizeof(int);
}
};
static_assert(Config<5, 6>::shared_bytes() == 190240,
"KD8 M5 shared contract changed");
static_assert(Config<4, 8>::shared_bytes() == 152192,
"KD8 M4 shared contract changed");
static_assert(Config<2, 16>::shared_bytes() == 76096,
"KD8 M2 shared contract changed");
struct SharedState {
float* work;
float* v;
float* w;
float* scalar;
int* counters;
};
struct SubgroupState {
float* work;
float* v;
float* w;
float* scalar;
int* counters;
int matrix_local;
int matrix_lane;
bool active;
unsigned mask;
};
template <int M, int S>
__device__ __forceinline__ SharedState partition_shared(unsigned char* raw) {
SharedState state;
state.work = reinterpret_cast<float*>(raw);
state.v = state.work + Config<M, S>::work_floats();
state.w = state.v + Config<M, S>::vector_floats();
state.scalar = state.w + Config<M, S>::vector_floats();
state.counters =
reinterpret_cast<int*>(state.scalar + Config<M, S>::scalar_floats());
return state;
}
template <int M, int S>
__device__ __forceinline__ SubgroupState subgroup_state(
SharedState shared,
int warp,
int lane) {
SubgroupState state;
state.matrix_local = lane / S;
state.matrix_lane = lane % S;
state.active = state.matrix_local < M;
const int local = state.active ? state.matrix_local : 0;
state.mask = ((S == 32 ? 0xFFFFFFFFu : ((1u << S) - 1u)) << (local * S));
const int task_index = warp * M + local;
state.work = shared.work + local * kMatrixWorkStride;
state.v = shared.v + task_index * kKD;
state.w = shared.w + task_index * kKD;
state.scalar = shared.scalar + task_index * kScalarStride;
state.counters = shared.counters + local * kCounterIntsPerMatrix;
return state;
}
__device__ __forceinline__ float& lower(float* work, int row, int column) {
return work[column * kWorkRows + (row - column)];
}
__device__ __forceinline__ float symmetric_get(
const float* work,
int row,
int column) {
if (row >= column) {
return work[column * kWorkRows + (row - column)];
}
return work[row * kWorkRows + (column - row)];
}
__device__ __forceinline__ int slot_for(int sweep, int segment, int* error) {
if (sweep < 1 || sweep >= kN - 1) {
atomicExch(error, 11);
return -1;
}
const int begin = static_cast<int>(kSweepPrefix[sweep - 1]);
const int end = static_cast<int>(kSweepPrefix[sweep]);
if (segment < 0 || segment >= end - begin) {
atomicExch(error, 12);
return -1;
}
const int slot = begin + segment;
if (slot < 0 || slot >= kReflectors) {
atomicExch(error, 13);
return -1;
}
return slot;
}
template <int S>
__device__ __forceinline__ int subgroup_slot(
int sweep,
int segment,
SubgroupState state) {
int slot = -1;
if (state.matrix_lane == 0) {
slot = slot_for(sweep, segment, &state.counters[1]);
}
return __shfl_sync(state.mask, slot, state.matrix_local * S);
}
__device__ void load_reflector(
const float* hh,
int slot,
int length,
SubgroupState state) {
if (state.matrix_lane == 0) {
const float* record = hh + static_cast<size_t>(slot) * kRecordWidth;
state.scalar[0] = record[0];
state.v[0] = 1.0f;
for (int index = 1; index < length; ++index) {
state.v[index] = record[index];
}
for (int index = length; index < kKD; ++index) {
state.v[index] = 0.0f;
}
}
__syncwarp(state.mask);
}
__device__ void generate_reflector(
float* hh,
int slot,
int column,
int head,
int length,
SubgroupState state) {
if (state.matrix_lane == 0) {
float* record = hh + static_cast<size_t>(slot) * kRecordWidth;
const float alpha = lower(state.work, head, column);
float sigma = 0.0f;
for (int index = 1; index < length; ++index) {
const float value = lower(state.work, head + index, column);
state.v[index] = value;
sigma = fmaf(value, value, sigma);
}
float tau = 0.0f;
float beta = alpha;
float gamma = 0.0f;
if (sigma > 0.0f) {
beta = -copysignf(sqrtf(fmaf(alpha, alpha, sigma)), alpha);
tau = (beta - alpha) / beta;
gamma = 1.0f / (alpha - beta);
}
state.v[0] = 1.0f;
lower(state.work, head, column) = beta;
record[0] = tau;
for (int index = 1; index < kKD; ++index) {
const float tail = index < length ? gamma * state.v[index] : 0.0f;
state.v[index] = tail;
record[index] = tail;
if (index < length) {
lower(state.work, head + index, column) = 0.0f;
}
}
state.scalar[0] = tau;
}
__syncwarp(state.mask);
}
template <int S>
__device__ void apply_symmetric_patch(
int head,
int length,
SubgroupState state) {
const float tau = state.scalar[0];
for (int row = state.matrix_lane; row < length; row += S) {
float sum = 0.0f;
for (int column = 0; column < length; ++column) {
sum = fmaf(
symmetric_get(state.work, head + row, head + column),
state.v[column],
sum);
}
state.w[row] = tau * sum;
}
__syncwarp(state.mask);
if (state.matrix_lane < length) {
float dot = 0.0f;
for (int index = 0; index < length; ++index) {
dot = fmaf(state.v[index], state.w[index], dot);
}
const float correction = -0.5f * tau * dot;
for (int row = state.matrix_lane; row < length; row += S) {
state.w[row] = fmaf(correction, state.v[row], state.w[row]);
}
}
__syncwarp(state.mask);
const int elements = length * length;
for (int linear = state.matrix_lane; linear < elements; linear += S) {
const int row = linear / length;
const int column = linear - row * length;
if (row >= column) {
float& value = lower(state.work, head + row, head + column);
value = fmaf(-state.v[row], state.w[column], value);
value = fmaf(-state.w[row], state.v[column], value);
}
}
__syncwarp(state.mask);
}
template <int S>
__device__ void apply_right_patch(
int rows_begin,
int rows,
int columns_begin,
int columns,
SubgroupState state) {
for (int r = state.matrix_lane; r < rows; r += S) {
const int row = rows_begin + r;
float dot = 0.0f;
for (int index = 0; index < columns; ++index) {
dot = fmaf(
lower(state.work, row, columns_begin + index),
state.v[index],
dot);
}
const float weight = state.scalar[0] * dot;
for (int index = 0; index < columns; ++index) {
float& value = lower(state.work, row, columns_begin + index);
value = fmaf(-weight, state.v[index], value);
}
}
__syncwarp(state.mask);
}
template <int S>
__device__ void apply_left_patch(
int rows_begin,
int rows,
int columns_begin,
int columns,
SubgroupState state) {
for (int c = state.matrix_lane; c < columns; c += S) {
const int column = columns_begin + c;
float dot = 0.0f;
for (int index = 0; index < rows; ++index) {
dot = fmaf(
state.v[index],
lower(state.work, rows_begin + index, column),
dot);
}
const float weight = state.scalar[0] * dot;
for (int index = 0; index < rows; ++index) {
float& value = lower(state.work, rows_begin + index, column);
value = fmaf(-state.v[index], weight, value);
}
}
__syncwarp(state.mask);
}
__device__ __forceinline__ bool retires_first_sweep(int sweep, int task_id) {
const int kind = task_id == 1 ? 1 : task_id % 2 + 2;
int point;
int start;
int end;
int block_last;
if (kind == 2) {
point = (task_id / 2) * kKD + sweep;
start = point - kKD + 1;
end = min(point, kN);
block_last = point;
} else {
point = ((task_id + 1) / 2) * kKD + sweep;
start = point - kKD + 1;
end = min(point, kN);
block_last = start >= end - 1 && end == kN ? kN : 0;
}
return block_last >= kN - 1;
}
template <int M, int S>
__global__ __launch_bounds__(kThreads, 1) void sb2st_b8_slot_shuffle_f32(
const float* __restrict__ band,
float* __restrict__ d,
float* __restrict__ e,
float* __restrict__ hh,
int* __restrict__ info,
int batch) {
extern __shared__ unsigned char raw_shared[];
SharedState shared = partition_shared<M, S>(raw_shared);
const int cta_matrix_begin = static_cast<int>(blockIdx.x) * M;
if (cta_matrix_begin + M > batch) {
return;
}
const int tid = static_cast<int>(threadIdx.x);
const int lane = tid & 31;
const int warp = tid >> 5;
SubgroupState state = subgroup_state<M, S>(shared, warp, lane);
int local_tasks = 0;
int local_producers = 0;
int local_consumers = 0;
for (int index = tid;
index < static_cast<int>(Config<M, S>::work_floats());
index += kThreads) {
shared.work[index] = 0.0f;
}
for (int index = tid;
index < static_cast<int>(
2 * Config<M, S>::vector_floats() + Config<M, S>::scalar_floats());
index += kThreads) {
shared.v[index] = 0.0f;
}
for (int index = tid; index < M * kCounterIntsPerMatrix; index += kThreads) {
shared.counters[index] = 0;
}
const size_t d_begin = static_cast<size_t>(cta_matrix_begin) * kN;
const size_t e_begin = static_cast<size_t>(cta_matrix_begin) * (kN - 1);
const size_t hh_begin =
static_cast<size_t>(cta_matrix_begin) * kReflectors * kRecordWidth;
for (int index = tid; index < M * kN; index += kThreads) {
d[d_begin + index] = 0.0f;
}
for (int index = tid; index < M * (kN - 1); index += kThreads) {
e[e_begin + index] = 0.0f;
}
for (int index = tid;
index < M * kReflectors * kRecordWidth;
index += kThreads) {
hh[hh_begin + index] = 0.0f;
}
__syncthreads();
for (int index = tid; index < M * kInputRows * kN; index += kThreads) {
const int matrix = index / (kInputRows * kN);
const int remainder = index - matrix * kInputRows * kN;
const float value = band[
static_cast<size_t>(cta_matrix_begin + matrix) * kInputRows * kN
+ remainder];
const int delta = remainder / kN;
const int column = remainder - delta * kN;
shared.work[
matrix * kMatrixWorkStride + column * kWorkRows + delta] = value;
if (!isfinite(value)) {
atomicExch(&shared.counters[matrix * kCounterIntsPerMatrix + 1], 20);
}
}
__syncthreads();
float* matrix_hh = hh
+ static_cast<size_t>(cta_matrix_begin + state.matrix_local)
* kReflectors * kRecordWidth;
int first_sweep = 1;
for (int diagonal = 1; diagonal < kN; ++diagonal) {
const int end_sweep = min(diagonal, kN - 1);
if (first_sweep > end_sweep) {
break;
}
for (int step = 1; step <= 3; ++step) {
const int sweep_begin = first_sweep;
const int width = max(0, end_sweep - sweep_begin + 1);
if (state.active) {
for (int rank = warp; rank < width; rank += kTaskWarps) {
const int sweep = sweep_begin + rank;
const int task_id = 3 * (diagonal - sweep) + step;
const int kind = task_id == 1 ? 1 : task_id % 2 + 2;
int point;
int start;
int end;
if (kind == 2) {
point = (task_id / 2) * kKD + sweep;
start = point - kKD + 1;
end = min(point, kN);
} else {
point = ((task_id + 1) / 2) * kKD + sweep;
start = point - kKD + 1;
end = min(point, kN);
}
const int patch_begin = start - 1;
const int patch_end = end - 1;
const int patch_length = patch_end - patch_begin + 1;
if (state.matrix_lane == 0) {
++local_tasks;
if (patch_length < 1 || patch_length > kKD) {
atomicExch(&state.counters[1], 21);
}
}
__syncwarp(state.mask);
if (kind == 1) {
if (patch_length > 1) {
const int slot = subgroup_slot<S>(sweep, 0, state);
if (state.matrix_lane == 0) {
++local_producers;
}
if (slot >= 0 && state.counters[1] == 0) {
generate_reflector(
matrix_hh,
slot,
patch_begin - 1,
patch_begin,
patch_length,
state);
apply_symmetric_patch<S>(patch_begin, patch_length, state);
}
}
} else if (kind == 2) {
const int prior_segment = task_id / 2 - 1;
const int prior_slot =
subgroup_slot<S>(sweep, prior_segment, state);
if (state.matrix_lane == 0) {
++local_consumers;
}
if (prior_slot >= 0 && state.counters[1] == 0) {
load_reflector(matrix_hh, prior_slot, patch_length, state);
const int next_begin = patch_end + 1;
const int next_end = min(patch_end + kKD, kN - 1);
const int next_length = next_end - next_begin + 1;
if (next_length > 0) {
apply_right_patch<S>(
next_begin,
next_length,
patch_begin,
patch_length,
state);
}
if (next_length > 1) {
const int next_segment = task_id / 2;
const int next_slot =
subgroup_slot<S>(sweep, next_segment, state);
if (state.matrix_lane == 0) {
++local_producers;
}
if (next_slot >= 0 && state.counters[1] == 0) {
generate_reflector(
matrix_hh,
next_slot,
patch_begin,
next_begin,
next_length,
state);
apply_left_patch<S>(
next_begin,
next_length,
patch_begin + 1,
patch_length - 1,
state);
}
}
}
} else {
const int segment = (task_id - 1) / 2;
const int slot = subgroup_slot<S>(sweep, segment, state);
if (state.matrix_lane == 0) {
++local_consumers;
}
if (slot >= 0 && state.counters[1] == 0) {
load_reflector(matrix_hh, slot, patch_length, state);
apply_symmetric_patch<S>(patch_begin, patch_length, state);
}
}
}
}
__syncthreads();
const int first_task_id = 3 * (diagonal - sweep_begin) + step;
if (retires_first_sweep(sweep_begin, first_task_id)) {
++first_sweep;
}
}
}
// The literal serial model contains one terminal kind-1 task
// T(sweep=511, task_id=1) whose patch length is one. It creates no
// reflector, touches no work element, and has no consumer, while the
// retirement-form wave loop above exits just before instantiating it
// (verified by loop emulation: 32,960 loop tasks + this one = 32,961).
if (warp == 0 && state.matrix_lane == 0 && state.active) {
++local_tasks;
}
if (state.active && state.matrix_lane == 0) {
atomicAdd(&state.counters[0], local_tasks);
atomicAdd(&state.counters[2], local_producers);
atomicAdd(&state.counters[3], local_consumers);
}
__syncthreads();
for (int index = tid; index < M * kN; index += kThreads) {
const int matrix = index / kN;
const int column = index - matrix * kN;
const float value = shared.work[
matrix * kMatrixWorkStride + column * kWorkRows];
d[d_begin + index] = value;
if (!isfinite(value)) {
atomicExch(&shared.counters[matrix * kCounterIntsPerMatrix + 1], 40);
}
}
for (int index = tid; index < M * (kN - 1); index += kThreads) {
const int matrix = index / (kN - 1);
const int column = index - matrix * (kN - 1);
const float value = shared.work[
matrix * kMatrixWorkStride + column * kWorkRows + 1];
e[e_begin + index] = value;
if (!isfinite(value)) {
atomicExch(&shared.counters[matrix * kCounterIntsPerMatrix + 1], 40);
}
}
for (int index = tid;
index < M * kReflectors * kRecordWidth;
index += kThreads) {
const int matrix = index / (kReflectors * kRecordWidth);
if (!isfinite(hh[hh_begin + index])) {
atomicExch(&shared.counters[matrix * kCounterIntsPerMatrix + 1], 40);
}
}
__syncthreads();
if (tid < M) {
int* counters = shared.counters + tid * kCounterIntsPerMatrix;
int status = counters[1];
if (status == 0 && counters[0] != kTasks) {
status = 30;
}
if (status == 0 && counters[2] != kReflectors) {
status = 31;
}
if (status == 0 && counters[3] != kConsumers) {
status = 32;
}
info[cta_matrix_begin + tid] = status;
}
}
std::array<uint16_t, kSweepPrefixEntries> build_sweep_prefix() {
std::array<uint16_t, kSweepPrefixEntries> prefix{};
int total = 0;
prefix[0] = 0;
for (int sweep = 1; sweep <= kN - 2; ++sweep) {
const int remaining = kN - sweep - 1;
const int count = (remaining + kKD - 1) / kKD;
total += count;
TORCH_CHECK(total <= kReflectors, "KD8 sweep prefix overflow");
prefix[sweep] = static_cast<uint16_t>(total);
}
TORCH_CHECK(total == kReflectors, "KD8 reflector count mismatch");
return prefix;
}
void ensure_sweep_prefix(int device) {
static std::mutex mutex;
static std::unordered_set<int> initialized_devices;
std::lock_guard<std::mutex> lock(mutex);
if (initialized_devices.count(device) != 0) {
return;
}
const auto prefix = build_sweep_prefix();
C10_CUDA_CHECK(cudaMemcpyToSymbol(
kSweepPrefix,
prefix.data(),
prefix.size() * sizeof(prefix[0]),
0,
cudaMemcpyHostToDevice));
initialized_devices.insert(device);
}
void validate_tensors(
const torch::Tensor& band,
const torch::Tensor& d,
const torch::Tensor& e,
const torch::Tensor& hh,
const torch::Tensor& info,
int matrices_per_cta) {
TORCH_CHECK(band.is_cuda(), "band must be CUDA");
TORCH_CHECK(d.is_cuda() && e.is_cuda() && hh.is_cuda() && info.is_cuda(),
"all outputs must be CUDA");
TORCH_CHECK(band.scalar_type() == torch::kFloat32 &&
d.scalar_type() == torch::kFloat32 &&
e.scalar_type() == torch::kFloat32 &&
hh.scalar_type() == torch::kFloat32,
"band/d/e/hh must be float32");
TORCH_CHECK(info.scalar_type() == at::kInt, "info must be int32");
TORCH_CHECK(band.is_contiguous() && d.is_contiguous() && e.is_contiguous() &&
hh.is_contiguous() && info.is_contiguous(),
"all tensors must be contiguous");
TORCH_CHECK(
band.dim() == 3 && band.size(1) == kInputRows && band.size(2) == kN,
"band must have shape [B,9,512]");
const auto batch = band.size(0);
TORCH_CHECK(batch > 0 && batch % matrices_per_cta == 0,
"batch must be a positive multiple of matrices-per-CTA");
TORCH_CHECK(d.sizes() == torch::IntArrayRef({batch, kN}), "d shape mismatch");
TORCH_CHECK(e.sizes() == torch::IntArrayRef({batch, kN - 1}),
"e shape mismatch");
TORCH_CHECK(
hh.sizes() == torch::IntArrayRef({batch, kReflectors, kRecordWidth}),
"hh shape mismatch");
TORCH_CHECK(info.sizes() == torch::IntArrayRef({batch}),
"info shape mismatch");
const int device = band.get_device();
TORCH_CHECK(d.get_device() == device && e.get_device() == device &&
hh.get_device() == device && info.get_device() == device,
"all tensors must use one device");
}
template <int M, int S>
void launch_config(
const torch::Tensor& band,
const torch::Tensor& d,
const torch::Tensor& e,
const torch::Tensor& hh,
const torch::Tensor& info) {
validate_tensors(band, d, e, hh, info, M);
c10::cuda::CUDAGuard guard(band.device());
const int device = band.get_device();
ensure_sweep_prefix(device);
int maximum_shared = 0;
C10_CUDA_CHECK(cudaDeviceGetAttribute(
&maximum_shared, cudaDevAttrMaxSharedMemoryPerBlockOptin, device));
TORCH_CHECK(
maximum_shared >= static_cast<int>(Config<M, S>::shared_bytes()),
"device shared-memory limit is below this KD8 config's contract");
C10_CUDA_CHECK(cudaFuncSetAttribute(
sb2st_b8_slot_shuffle_f32<M, S>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
static_cast<int>(Config<M, S>::shared_bytes())));
const int batch = static_cast<int>(band.size(0));
sb2st_b8_slot_shuffle_f32<M, S><<<
batch / M,
kThreads,
Config<M, S>::shared_bytes()>>>(
band.data_ptr<float>(),
d.data_ptr<float>(),
e.data_ptr<float>(),
hh.data_ptr<float>(),
info.data_ptr<int>(),
batch);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
template <int M, int S>
std::vector<int64_t> report_config() {
C10_CUDA_CHECK(cudaFuncSetAttribute(
sb2st_b8_slot_shuffle_f32<M, S>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
static_cast<int>(Config<M, S>::shared_bytes())));
cudaFuncAttributes attributes{};
C10_CUDA_CHECK(cudaFuncGetAttributes(
&attributes, sb2st_b8_slot_shuffle_f32<M, S>));
int active_blocks = 0;
C10_CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&active_blocks,
sb2st_b8_slot_shuffle_f32<M, S>,
kThreads,
Config<M, S>::shared_bytes()));
int device = -1;
C10_CUDA_CHECK(cudaGetDevice(&device));
int sm_count = 0;
C10_CUDA_CHECK(cudaDeviceGetAttribute(
&sm_count, cudaDevAttrMultiProcessorCount, device));
return {
M,
S,
kThreads,
static_cast<int64_t>(Config<M, S>::shared_bytes()),
attributes.numRegs,
static_cast<int64_t>(attributes.localSizeBytes),
static_cast<int64_t>(attributes.sharedSizeBytes),
active_blocks,
sm_count,
};
}
} // namespace
void band8_to_tridiagonal_kd8_out(
const torch::Tensor& band,
const torch::Tensor& d,
const torch::Tensor& e,
const torch::Tensor& hh,
const torch::Tensor& info,
const std::string& config) {
if (config == "m5s6") {
launch_config<5, 6>(band, d, e, hh, info);
} else if (config == "m4s8") {
launch_config<4, 8>(band, d, e, hh, info);
} else if (config == "m2s16") {
launch_config<2, 16>(band, d, e, hh, info);
} else {
TORCH_CHECK(false, "unknown KD8 config: ", config);
}
}
std::vector<int64_t> kd8_resource_report(const std::string& config) {
if (config == "m5s6") {
return report_config<5, 6>();
}
if (config == "m4s8") {
return report_config<4, 8>();
}
if (config == "m2s16") {
return report_config<2, 16>();
}
TORCH_CHECK(false, "unknown KD8 config: ", config);
return {};
}
} // namespace lane_c
namespace lane_d {
namespace {
constexpr int kN = 32;
constexpr int kIterationCap = 30 * kN;
constexpr int kLeavesPerCustomBlock = 4;
struct QLControl {
float scale;
float g;
float p;
float c;
float s;
int m;
int row;
int swap_row;
int converged;
int rotate;
int early_restart;
int failed;
int iterations;
int apply_low;
};
// Four independent leaf warps share one CTA. Lane 0 computes a complete
// dependent Givens chain, then every lane applies that saved chain to its own
// eigenvector column without a synchronization after each rotation.
extern "C" __global__ void leaf32_warp_ql_kernel(
float* __restrict__ d_global, float* __restrict__ e_global,
float* __restrict__ zt_global, int* __restrict__ info) {
__shared__ float all_d[kLeavesPerCustomBlock][kN];
__shared__ float all_e[kLeavesPerCustomBlock][kN];
__shared__ float all_zt[kLeavesPerCustomBlock][kN * kN];
__shared__ float all_rotation_c[kLeavesPerCustomBlock][kN];
__shared__ float all_rotation_s[kLeavesPerCustomBlock][kN];
__shared__ QLControl all_ctl[kLeavesPerCustomBlock];
const int warp = threadIdx.x / 32;
const int lane = threadIdx.x % 32;
const int batch = blockIdx.x * kLeavesPerCustomBlock + warp;
float* sd = all_d[warp];
float* se = all_e[warp];
float* szt = all_zt[warp];
float* rotation_c = all_rotation_c[warp];
float* rotation_s = all_rotation_s[warp];
QLControl& ctl = all_ctl[warp];
const int d_base = batch * kN;
const int e_base = batch * (kN - 1);
const size_t z_base = static_cast<size_t>(batch) * kN * kN;
sd[lane] = d_global[d_base + lane];
se[lane] = lane < kN - 1 ? e_global[e_base + lane] : 0.0f;
for (int index = lane; index < kN * kN; index += 32) {
const int row = index / kN;
const int col = index - row * kN;
szt[index] = row == col ? 1.0f : 0.0f;
}
__syncwarp();
if (lane == 0) {
ctl.scale = 0.0f;
ctl.failed = 0;
ctl.iterations = 0;
for (int i = 0; i < kN; ++i) {
if (!isfinite(sd[i])) ctl.failed = 2;
ctl.scale = fmaxf(ctl.scale, fabsf(sd[i]));
if (i < kN - 1) {
if (!isfinite(se[i])) ctl.failed = 2;
ctl.scale = fmaxf(ctl.scale, fabsf(se[i]));
}
}
}
__syncwarp();
if (ctl.failed == 0 && ctl.scale > 0.0f) {
sd[lane] /= ctl.scale;
if (lane < kN - 1) se[lane] /= ctl.scale;
}
__syncwarp();
if (ctl.failed == 0 && ctl.scale > 0.0f) {
for (int l = 0; l < kN; ++l) {
while (true) {
if (lane == 0) {
int m = l;
for (; m < kN - 1; ++m) {
const float tst = fabsf(se[m]);
const float bound = FLT_EPSILON * FLT_EPSILON *
fabsf(sd[m]) * fabsf(sd[m + 1]) +
FLT_MIN;
if (tst * tst <= bound) {
se[m] = 0.0f;
break;
}
}
ctl.m = m;
ctl.converged = (m == l);
ctl.early_restart = 0;
ctl.rotate = 0;
if (!ctl.converged) {
++ctl.iterations;
if (ctl.iterations > kIterationCap) {
ctl.failed = 1;
} else {
float g = (sd[l + 1] - sd[l]) / (2.0f * se[l]);
const float r = hypotf(g, 1.0f);
const float denom = g + copysignf(r, g);
if (denom == 0.0f || !isfinite(denom)) {
ctl.failed = 2;
} else {
ctl.g = sd[m] - sd[l] + se[l] / denom;
ctl.s = 1.0f;
ctl.c = 1.0f;
ctl.p = 0.0f;
if (!isfinite(ctl.g)) ctl.failed = 2;
}
}
}
}
__syncwarp();
const bool stop_outer_iteration = ctl.failed != 0 || ctl.converged;
// Close the all-lane shared-control read before lane 0 overwrites the
// next iteration's control fields.
__syncwarp();
if (stop_outer_iteration) break;
const int m = ctl.m;
if (lane == 0) {
ctl.apply_low = m;
ctl.early_restart = 0;
for (int i = m - 1; i >= l; --i) {
const float f = ctl.s * se[i];
const float b = ctl.c * se[i];
const float r = hypotf(f, ctl.g);
se[i + 1] = r;
if (r <= FLT_MIN) {
sd[i + 1] -= ctl.p;
se[m] = 0.0f;
ctl.early_restart = 1;
break;
}
const float s = f / r;
const float c = ctl.g / r;
const float g = sd[i + 1] - ctl.p;
const float recurrence = (sd[i] - g) * s + 2.0f * c * b;
const float p = s * recurrence;
sd[i + 1] = g + p;
ctl.g = c * recurrence - b;
ctl.p = p;
ctl.s = s;
ctl.c = c;
rotation_c[i] = c;
rotation_s[i] = s;
ctl.apply_low = i;
if (!isfinite(sd[i + 1]) || !isfinite(ctl.g) ||
!isfinite(ctl.p) || !isfinite(s) || !isfinite(c)) {
ctl.failed = 2;
break;
}
}
if (ctl.failed == 0 && !ctl.early_restart) {
sd[l] -= ctl.p;
se[l] = ctl.g;
se[m] = 0.0f;
if (!isfinite(sd[l]) || !isfinite(se[l])) ctl.failed = 2;
}
}
__syncwarp();
const bool apply_chain = ctl.failed == 0 && ctl.apply_low < m;
const int apply_low = ctl.apply_low;
// Close all shared-control reads before lane 0 can prepare the next
// scalar iteration. Each lane owns one column, so the saved chain
// needs no synchronization internally.
__syncwarp();
if (apply_chain) {
float carry = szt[m * kN + lane];
for (int i = m - 1; i >= apply_low; --i) {
const float a = szt[i * kN + lane];
const float c = rotation_c[i];
const float s = rotation_s[i];
szt[(i + 1) * kN + lane] = s * a + c * carry;
carry = c * a - s * carry;
}
szt[apply_low * kN + lane] = carry;
}
__syncwarp();
const bool failed_iteration = ctl.failed != 0;
__syncwarp();
if (failed_iteration) break;
}
__syncwarp();
const bool failed_leaf = ctl.failed != 0;
__syncwarp();
if (failed_leaf) break;
}
}
const bool can_sort = ctl.failed == 0;
__syncwarp();
if (can_sort) {
for (int i = 0; i < kN - 1; ++i) {
if (lane == 0) {
int selected = i;
float value = sd[i];
for (int j = i + 1; j < kN; ++j) {
if (sd[j] < value) {
selected = j;
value = sd[j];
}
}
ctl.swap_row = selected;
if (selected != i) {
const float tmp = sd[i];
sd[i] = sd[selected];
sd[selected] = tmp;
}
}
__syncwarp();
if (ctl.swap_row != i) {
const float tmp = szt[i * kN + lane];
szt[i * kN + lane] = szt[ctl.swap_row * kN + lane];
szt[ctl.swap_row * kN + lane] = tmp;
}
__syncwarp();
}
}
d_global[d_base + lane] = sd[lane] * ctl.scale;
if (lane < kN - 1) e_global[e_base + lane] = se[lane] * ctl.scale;
for (int index = lane; index < kN * kN; index += 32) {
zt_global[z_base + index] = szt[index];
}
if (lane == 0) info[batch] = ctl.failed;
}
void check_tensor(const at::Tensor& tensor, at::ScalarType dtype,
int64_t rank, const char* name) {
TORCH_CHECK(tensor.is_cuda(), name, " must be CUDA");
TORCH_CHECK(tensor.is_contiguous(), name, " must be contiguous");
TORCH_CHECK(tensor.scalar_type() == dtype, name, " has wrong dtype");
TORCH_CHECK(tensor.dim() == rank, name, " has wrong rank");
}
} // namespace
void leaf32_saved_chain_run(const at::Tensor& d, const at::Tensor& e,
const at::Tensor& qt, const at::Tensor& info) {
check_tensor(d, at::kFloat, 2, "d");
check_tensor(e, at::kFloat, 2, "e");
check_tensor(qt, at::kFloat, 3, "qt");
check_tensor(info, at::kInt, 1, "info");
const int64_t batch = d.size(0);
TORCH_CHECK(d.size(1) == kN, "d shape mismatch");
TORCH_CHECK(e.size(0) == batch && e.size(1) == kN - 1, "e shape mismatch");
TORCH_CHECK(qt.size(0) == batch && qt.size(1) == kN && qt.size(2) == kN,
"qt shape mismatch");
TORCH_CHECK(info.size(0) == batch, "info shape mismatch");
TORCH_CHECK(batch > 0 && batch % kLeavesPerCustomBlock == 0,
"leaf batch must be positive and divisible by four");
TORCH_CHECK(d.get_device() == e.get_device() && d.get_device() == qt.get_device() &&
d.get_device() == info.get_device(), "all tensors must share one device");
c10::cuda::CUDAGuard guard(d.device());
leaf32_warp_ql_kernel<<<batch / kLeavesPerCustomBlock,
32 * kLeavesPerCustomBlock, 0>>>(
d.data_ptr<float>(), e.data_ptr<float>(), qt.data_ptr<float>(),
info.data_ptr<int32_t>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
} // namespace lane_d
namespace lane_e64 {
namespace {
constexpr int kMaxMerge = 64;
constexpr int kMaxIterations = 30;
__device__ __forceinline__ void kahan_add(
float value, float& total, float& compensation) {
const float adjusted = value - compensation;
const float updated = total + adjusted;
compensation = (updated - total) - adjusted;
total = updated;
}
__device__ __forceinline__ float secular_value(
const float* poles,
const float* weights,
int count,
float rho,
float x) {
// Accumulate the two signs separately. Close-pole secular evaluations
// otherwise lose the smaller signed partial sum before the final add.
float negative = 0.0f;
float positive = 0.0f;
float cnegative = 0.0f;
float cpositive = 0.0f;
for (int i = 0; i < count; ++i) {
const float term = rho * weights[i] * weights[i] / (poles[i] - x);
if (term < 0.0f) {
kahan_add(term, negative, cnegative);
} else {
kahan_add(term, positive, cpositive);
}
}
return 1.0f + (negative + positive);
}
// Status-2 is often a representation boundary rather than a missing root:
// the mathematical root lies between a pole and its first inward FP32 value.
// Evaluate only that endpoint-sign predicate in FP64. This is deliberately
// not a second FP64 root solve. A valid predicate authorizes a caller-visible
// FP32 endpoint clamp; an invalid predicate remains status 2 and is handled by
// the existing fail-closed selective-root FP64 fallback launch.
__device__ __forceinline__ double endpoint_secular_value_double(
const float* poles,
const float* weights,
int count,
float rho,
double x) {
double negative = 0.0;
double positive = 0.0;
double cnegative = 0.0;
double cpositive = 0.0;
for (int index = 0; index < count; ++index) {
const double pole = static_cast<double>(poles[index]);
const double weight = static_cast<double>(weights[index]);
const double term =
static_cast<double>(rho) * weight * weight / (pole - x);
if (term < 0.0) {
const double adjusted = term - cnegative;
const double updated = negative + adjusted;
cnegative = (updated - negative) - adjusted;
negative = updated;
} else {
const double adjusted = term - cpositive;
const double updated = positive + adjusted;
cpositive = (updated - positive) - adjusted;
positive = updated;
}
}
return 1.0 + negative + positive;
}
struct EvalFloat {
float value;
float derivative;
float error_scale;
};
__device__ __forceinline__ EvalFloat secular_eval_float(
const float* poles,
const float* weights,
int count,
float rho,
float x) {
float negative = 0.0f;
float positive = 0.0f;
float derivative = 0.0f;
float magnitude = 0.0f;
float cnegative = 0.0f;
float cpositive = 0.0f;
float cderivative = 0.0f;
float cmagnitude = 0.0f;
for (int index = 0; index < count; ++index) {
const float delta = poles[index] - x;
const float weight2 = weights[index] * weights[index];
const float term = rho * weight2 / delta;
if (term < 0.0f) {
kahan_add(term, negative, cnegative);
} else {
kahan_add(term, positive, cpositive);
}
kahan_add(rho * weight2 / (delta * delta), derivative, cderivative);
kahan_add(fabsf(term), magnitude, cmagnitude);
}
return {
1.0f + negative + positive,
derivative,
8.0f * (1.0f + magnitude + fabsf(x) * derivative)};
}
__device__ __forceinline__ float interior_rational_step_float(
const float* poles,
const float* weights,
float rho,
int index,
float x,
float value,
float derivative,
bool origin_at_lower) {
const float delta_i = poles[index] - x;
const float delta_ip1 = poles[index + 1] - x;
const float gap = poles[index + 1] - poles[index];
float c;
if (origin_at_lower) {
const float ratio = weights[index] / delta_i;
c = value - delta_ip1 * derivative + gap * rho * ratio * ratio;
} else {
const float ratio = weights[index + 1] / delta_ip1;
c = value - delta_i * derivative - gap * rho * ratio * ratio;
}
const float a =
(delta_i + delta_ip1) * value - delta_i * delta_ip1 * derivative;
const float b = delta_i * delta_ip1 * value;
if (c == 0.0f) {
return a == 0.0f ? CUDART_NAN_F : b / a;
}
const float discriminant = fmaxf(0.0f, a * a - 4.0f * b * c);
const float root = sqrtf(discriminant);
if (a <= 0.0f) {
return (a - root) / (2.0f * c);
}
const float denominator = a + root;
return denominator == 0.0f ? CUDART_NAN_F : 2.0f * b / denominator;
}
__device__ __forceinline__ float last_rational_step_float(
const float* poles,
const float* weights,
int count,
float rho,
float x,
float value) {
const float delta0 = poles[count - 2] - x;
const float delta1 = poles[count - 1] - x;
float dpsi = 0.0f;
float correction = 0.0f;
for (int index = 0; index < count - 1; ++index) {
const float ratio = weights[index] / (poles[index] - x);
kahan_add(rho * ratio * ratio, dpsi, correction);
}
const float last_ratio = weights[count - 1] / delta1;
const float dphi = rho * last_ratio * last_ratio;
const float c = fabsf(value - delta0 * dpsi - delta1 * dphi);
const float a =
(delta0 + delta1) * value - delta0 * delta1 * (dpsi + dphi);
const float b = delta0 * delta1 * value;
if (c == 0.0f) {
return CUDART_NAN_F;
}
const float discriminant = fmaxf(0.0f, a * a - 4.0f * b * c);
const float root = sqrtf(discriminant);
if (a >= 0.0f) {
return (a + root) / (2.0f * c);
}
const float denominator = a - root;
return denominator == 0.0f ? CUDART_NAN_F : 2.0f * b / denominator;
}
__global__ void secular_merge_kernel(
const float* __restrict__ poles,
const float* __restrict__ weights,
const float* __restrict__ rho,
const int32_t* __restrict__ active_count,
float* __restrict__ values,
float* __restrict__ secular_vectors,
int32_t* __restrict__ info,
int32_t* __restrict__ iterations) {
const int batch = static_cast<int>(blockIdx.x);
const int lane = static_cast<int>(threadIdx.x);
const int base = batch * kMaxMerge;
const int matrix_base = batch * kMaxMerge * kMaxMerge;
const int count = active_count[batch];
__shared__ float shared_poles[kMaxMerge];
__shared__ float shared_weights[kMaxMerge];
__shared__ float roots[kMaxMerge];
__shared__ float updated_weights[kMaxMerge];
if (lane == 0) {
info[batch] = (count > 0 && count <= kMaxMerge) ? 0 : 1;
iterations[batch] = 0;
}
if (lane < kMaxMerge) {
shared_poles[lane] = poles[base + lane];
shared_weights[lane] = weights[base + lane];
values[base + lane] = 0.0f;
for (int row = 0; row < kMaxMerge; ++row) {
secular_vectors[matrix_base + row * kMaxMerge + lane] = 0.0f;
}
}
__syncthreads();
if (lane < count) {
const float positive_inf = CUDART_INF_F;
const float negative_inf = -CUDART_INF_F;
float lo = nextafterf(shared_poles[lane], positive_inf);
float hi;
if (lane + 1 < count) {
hi = nextafterf(shared_poles[lane + 1], negative_inf);
} else {
float normz2 = 0.0f;
float correction = 0.0f;
for (int i = 0; i < count; ++i) {
kahan_add(shared_weights[i] * shared_weights[i], normz2, correction);
}
hi = shared_poles[count - 1] + rho[batch] * normz2;
if (!(hi > lo)) {
hi = nextafterf(lo, positive_inf);
}
for (int attempt = 0; attempt < 8; ++attempt) {
const float fhi = secular_value(
shared_poles, shared_weights, count, rho[batch], hi);
if (fhi >= 0.0f) {
break;
}
hi = shared_poles[count - 1] +
2.0f * (hi - shared_poles[count - 1]);
}
}
const float flo = secular_eval_float(
shared_poles, shared_weights, count, rho[batch], lo).value;
const float fhi = secular_eval_float(
shared_poles, shared_weights, count, rho[batch], hi).value;
if (!(lo < hi) || !(flo <= 0.0f) || !(fhi >= 0.0f) ||
!isfinite(lo) || !isfinite(hi)) {
bool endpoint_sign_rescued = false;
if (lo < hi && isfinite(lo) && isfinite(hi) &&
(flo > 0.0f || fhi < 0.0f)) {
const double double_lo = nextafter(
static_cast<double>(shared_poles[lane]), CUDART_INF);
double double_hi;
if (lane + 1 < count) {
double_hi = nextafter(
static_cast<double>(shared_poles[lane + 1]), -CUDART_INF);
} else {
double normz2 = 0.0;
double correction = 0.0;
for (int index = 0; index < count; ++index) {
const double weight =
static_cast<double>(shared_weights[index]);
const double value = weight * weight - correction;
const double updated = normz2 + value;
correction = (updated - normz2) - value;
normz2 = updated;
}
double_hi = static_cast<double>(shared_poles[count - 1]) +
static_cast<double>(rho[batch]) * normz2;
if (!(double_hi > double_lo)) {
double_hi = nextafter(double_lo, CUDART_INF);
}
for (int attempt = 0; attempt < 8; ++attempt) {
if (endpoint_secular_value_double(
shared_poles,
shared_weights,
count,
rho[batch],
double_hi) >= 0.0) {
break;
}
double_hi = static_cast<double>(shared_poles[count - 1]) +
2.0 *
(double_hi -
static_cast<double>(shared_poles[count - 1]));
}
}
const double double_flo = endpoint_secular_value_double(
shared_poles,
shared_weights,
count,
rho[batch],
double_lo);
const double double_fhi = endpoint_secular_value_double(
shared_poles,
shared_weights,
count,
rho[batch],
double_hi);
endpoint_sign_rescued =
isfinite(double_lo) && isfinite(double_hi) &&
isfinite(double_flo) && isfinite(double_fhi) &&
double_lo < double_hi && double_flo <= 0.0 &&
double_fhi >= 0.0;
}
if (endpoint_sign_rescued) {
// The exact root is below/above the first representable interior FP32
// endpoint. Clamp at that endpoint and retain the existing FP32
// SLAED3 displacement reconstruction below.
roots[lane] = flo > 0.0f ? lo : hi;
} else {
atomicCAS(info + batch, 0, 2);
roots[lane] = CUDART_NAN_F;
}
} else {
float x = lo + 0.5f * (hi - lo);
const bool origin_at_lower = secular_eval_float(
shared_poles, shared_weights, count, rho[batch], x).value > 0.0f;
bool converged = false;
int used = 0;
for (int iteration = 1; iteration <= kMaxIterations; ++iteration) {
used = iteration;
const EvalFloat current = secular_eval_float(
shared_poles, shared_weights, count, rho[batch], x);
if (fabsf(current.value) <= FLT_EPSILON * current.error_scale) {
converged = true;
break;
}
if (current.value <= 0.0f) {
lo = fmaxf(lo, x);
} else {
hi = fminf(hi, x);
}
const float scale = fmaxf(1.0f, fmaxf(fabsf(lo), fabsf(hi)));
if (hi - lo <= 2.0f * FLT_EPSILON * scale || lo == hi) {
x = lo + 0.5f * (hi - lo);
converged = true;
break;
}
float eta = lane + 1 < count
? interior_rational_step_float(
shared_poles,
shared_weights,
rho[batch],
lane,
x,
current.value,
current.derivative,
origin_at_lower)
: last_rational_step_float(
shared_poles,
shared_weights,
count,
rho[batch],
x,
current.value);
if (!isfinite(eta) || current.value * eta >= 0.0f) {
eta = -current.value / current.derivative;
}
float proposed = x + eta;
if (!isfinite(proposed) || !(proposed > lo && proposed < hi) ||
proposed == x || iteration >= 12) {
proposed = lo + 0.5f * (hi - lo);
}
x = proposed;
}
if (!converged || !isfinite(x)) {
atomicCAS(info + batch, 0, 5);
}
roots[lane] = x;
atomicMax(iterations + batch, used);
}
}
__syncthreads();
// SLAED3 reconstructs the rank-one weights from all root displacements.
// Evaluate the product in the log domain: this is algebraically equivalent
// for the positive rank-one problem and avoids a canary-only overflow.
if (lane < count) {
float diagonal_delta = fabsf(shared_poles[lane] - roots[lane]);
if (!(diagonal_delta > 0.0f) || !isfinite(diagonal_delta)) {
atomicCAS(info + batch, 0, 3);
updated_weights[lane] = CUDART_NAN_F;
} else {
float log_product = logf(diagonal_delta);
for (int root = 0; root < count; ++root) {
if (root == lane) {
continue;
}
const float numerator = fabsf(shared_poles[lane] - roots[root]);
const float denominator =
fabsf(shared_poles[lane] - shared_poles[root]);
if (!(numerator > 0.0f) || !(denominator > 0.0f)) {
atomicCAS(info + batch, 0, 3);
} else {
log_product += logf(numerator) - logf(denominator);
}
}
const float magnitude = expf(0.5f * log_product);
updated_weights[lane] = copysignf(magnitude, shared_weights[lane]);
if (!isfinite(updated_weights[lane])) {
atomicCAS(info + batch, 0, 3);
}
}
}
__syncthreads();
if (lane < count) {
float norm2 = 0.0f;
float correction = 0.0f;
for (int row = 0; row < count; ++row) {
const float delta = shared_poles[row] - roots[lane];
const float element = updated_weights[row] / delta;
kahan_add(element * element, norm2, correction);
}
const float norm = sqrtf(norm2);
if (!(norm > 0.0f) || !isfinite(norm)) {
atomicCAS(info + batch, 0, 4);
}
for (int row = 0; row < count; ++row) {
const float delta = shared_poles[row] - roots[lane];
secular_vectors[matrix_base + row * kMaxMerge + lane] =
updated_weights[row] / delta / norm;
}
values[base + lane] = roots[lane];
}
}
__device__ __forceinline__ void kahan_add_double(
double value, double& total, double& compensation) {
const double adjusted = value - compensation;
const double updated = total + adjusted;
compensation = (updated - total) - adjusted;
total = updated;
}
__device__ __forceinline__ double secular_value_double(
const double* poles,
const double* weights,
int count,
double rho,
double x) {
double negative = 0.0;
double positive = 0.0;
double cnegative = 0.0;
double cpositive = 0.0;
for (int i = 0; i < count; ++i) {
const double term = rho * weights[i] * weights[i] / (poles[i] - x);
if (term < 0.0) {
kahan_add_double(term, negative, cnegative);
} else {
kahan_add_double(term, positive, cpositive);
}
}
return 1.0 + (negative + positive);
}
struct EvalDouble {
double value;
double derivative;
double error_scale;
};
__device__ __forceinline__ EvalDouble secular_eval_double(
const double* poles,
const double* weights,
int count,
double rho,
double x) {
double negative = 0.0;
double positive = 0.0;
double derivative = 0.0;
double magnitude = 0.0;
double cnegative = 0.0;
double cpositive = 0.0;
double cderivative = 0.0;
double cmagnitude = 0.0;
for (int index = 0; index < count; ++index) {
const double delta = poles[index] - x;
const double weight2 = weights[index] * weights[index];
const double term = rho * weight2 / delta;
if (term < 0.0) {
kahan_add_double(term, negative, cnegative);
} else {
kahan_add_double(term, positive, cpositive);
}
kahan_add_double(
rho * weight2 / (delta * delta), derivative, cderivative);
kahan_add_double(fabs(term), magnitude, cmagnitude);
}
return {
1.0 + negative + positive,
derivative,
8.0 * (1.0 + magnitude + fabs(x) * derivative)};
}
__device__ __forceinline__ double interior_rational_step_double(
const double* poles,
const double* weights,
double rho,
int index,
double x,
double value,
double derivative,
bool origin_at_lower) {
const double delta_i = poles[index] - x;
const double delta_ip1 = poles[index + 1] - x;
const double gap = poles[index + 1] - poles[index];
double c;
if (origin_at_lower) {
const double ratio = weights[index] / delta_i;
c = value - delta_ip1 * derivative + gap * rho * ratio * ratio;
} else {
const double ratio = weights[index + 1] / delta_ip1;
c = value - delta_i * derivative - gap * rho * ratio * ratio;
}
const double a =
(delta_i + delta_ip1) * value - delta_i * delta_ip1 * derivative;
const double b = delta_i * delta_ip1 * value;
if (c == 0.0) {
return a == 0.0 ? CUDART_NAN : b / a;
}
const double root = sqrt(fmax(0.0, a * a - 4.0 * b * c));
if (a <= 0.0) {
return (a - root) / (2.0 * c);
}
const double denominator = a + root;
return denominator == 0.0 ? CUDART_NAN : 2.0 * b / denominator;
}
__device__ __forceinline__ double last_rational_step_double(
const double* poles,
const double* weights,
int count,
double rho,
double x,
double value) {
const double delta0 = poles[count - 2] - x;
const double delta1 = poles[count - 1] - x;
double dpsi = 0.0;
double correction = 0.0;
for (int index = 0; index < count - 1; ++index) {
const double ratio = weights[index] / (poles[index] - x);
kahan_add_double(rho * ratio * ratio, dpsi, correction);
}
const double last_ratio = weights[count - 1] / delta1;
const double dphi = rho * last_ratio * last_ratio;
const double c = fabs(value - delta0 * dpsi - delta1 * dphi);
const double a =
(delta0 + delta1) * value - delta0 * delta1 * (dpsi + dphi);
const double b = delta0 * delta1 * value;
if (c == 0.0) {
return CUDART_NAN;
}
const double root = sqrt(fmax(0.0, a * a - 4.0 * b * c));
if (a >= 0.0) {
return (a + root) / (2.0 * c);
}
const double denominator = a - root;
return denominator == 0.0 ? CUDART_NAN : 2.0 * b / denominator;
}
__global__ void secular_merge_fp64_control_kernel(
const float* __restrict__ poles,
const float* __restrict__ weights,
const float* __restrict__ rho,
const int32_t* __restrict__ active_count,
float* __restrict__ values,
float* __restrict__ secular_vectors,
int32_t* __restrict__ info,
int32_t* __restrict__ iterations,
bool fallback_only) {
const int batch = static_cast<int>(blockIdx.x);
const int lane = static_cast<int>(threadIdx.x);
const int base = batch * kMaxMerge;
const int matrix_base = batch * kMaxMerge * kMaxMerge;
const int count = active_count[batch];
__shared__ double shared_poles[kMaxMerge];
__shared__ double shared_weights[kMaxMerge];
__shared__ double roots[kMaxMerge];
__shared__ double updated_weights[kMaxMerge];
__shared__ int32_t deflated[kMaxMerge];
__shared__ int32_t prior_info;
// Snapshot the preceding strict-FP32 status before lane zero clears it.
// The barrier makes both the early return and repair route uniform across
// the CTA's two warps.
if (lane == 0) {
prior_info = fallback_only ? info[batch] : 0;
}
__syncthreads();
if (fallback_only && prior_info == 0) {
return;
}
// Status 2 is an endpoint-sign failure. The primary launch leaves finite
// values for roots whose own brackets succeeded and NaN for failed roots.
const float primary_root = lane < count ? values[base + lane] : CUDART_NAN_F;
const bool reuse_primary_root =
fallback_only && prior_info == 2 && isfinite(primary_root);
if (lane == 0) {
info[batch] = (count > 0 && count <= kMaxMerge) ? 0 : 1;
iterations[batch] = 0;
}
shared_poles[lane] = static_cast<double>(poles[base + lane]);
shared_weights[lane] = static_cast<double>(weights[base + lane]);
deflated[lane] = 0;
if (!reuse_primary_root) {
values[base + lane] = 0.0f;
}
for (int row = 0; row < kMaxMerge; ++row) {
secular_vectors[matrix_base + row * kMaxMerge + lane] = 0.0f;
}
__syncthreads();
if (lane < count) {
if (reuse_primary_root) {
roots[lane] = static_cast<double>(primary_root);
} else {
double lo = nextafter(shared_poles[lane], CUDART_INF);
double hi;
if (lane + 1 < count) {
hi = nextafter(shared_poles[lane + 1], -CUDART_INF);
} else {
double normz2 = 0.0;
double correction = 0.0;
for (int i = 0; i < count; ++i) {
kahan_add_double(
shared_weights[i] * shared_weights[i], normz2, correction);
}
hi = shared_poles[count - 1] + static_cast<double>(rho[batch]) * normz2;
if (!(hi > lo)) {
hi = nextafter(lo, CUDART_INF);
}
for (int attempt = 0; attempt < 8; ++attempt) {
const double fhi = secular_value_double(
shared_poles,
shared_weights,
count,
static_cast<double>(rho[batch]),
hi);
if (fhi >= 0.0) {
break;
}
hi = shared_poles[count - 1] +
2.0 * (hi - shared_poles[count - 1]);
}
}
const double flo = secular_eval_double(
shared_poles,
shared_weights,
count,
static_cast<double>(rho[batch]),
lo).value;
const double fhi = secular_eval_double(
shared_poles,
shared_weights,
count,
static_cast<double>(rho[batch]),
hi).value;
const bool finite_bracket =
(lo < hi) && isfinite(lo) && isfinite(hi) &&
isfinite(flo) && isfinite(fhi);
const bool sub_ulp_deflation =
finite_bracket && flo >= 0.0 && fhi > 0.0;
if (sub_ulp_deflation) {
// The first FP64 number above the pole is already beyond the
// root: its displacement is unrepresentable in double.
deflated[lane] = 1;
roots[lane] = shared_poles[lane];
} else if (!finite_bracket || !(flo < 0.0) || !(fhi > 0.0)) {
atomicCAS(info + batch, 0, 2);
roots[lane] = CUDART_NAN;
} else {
const double rho_value = static_cast<double>(rho[batch]);
double x = lo + 0.5 * (hi - lo);
const bool origin_at_lower = secular_eval_double(
shared_poles, shared_weights, count, rho_value, x).value > 0.0;
bool converged = false;
int used = 0;
for (int iteration = 1; iteration <= kMaxIterations; ++iteration) {
used = iteration;
const EvalDouble current = secular_eval_double(
shared_poles, shared_weights, count, rho_value, x);
if (fabs(current.value) <= DBL_EPSILON * current.error_scale) {
converged = true;
break;
}
if (current.value <= 0.0) {
lo = fmax(lo, x);
} else {
hi = fmin(hi, x);
}
const double scale = fmax(1.0, fmax(fabs(lo), fabs(hi)));
if (hi - lo <= 2.0 * static_cast<double>(FLT_EPSILON) * scale ||
static_cast<float>(lo) == static_cast<float>(hi)) {
x = lo + 0.5 * (hi - lo);
converged = true;
break;
}
double eta = lane + 1 < count
? interior_rational_step_double(
shared_poles,
shared_weights,
rho_value,
lane,
x,
current.value,
current.derivative,
origin_at_lower)
: last_rational_step_double(
shared_poles,
shared_weights,
count,
rho_value,
x,
current.value);
if (!isfinite(eta) || current.value * eta >= 0.0) {
eta = -current.value / current.derivative;
}
double proposed = x + eta;
if (!isfinite(proposed) || !(proposed > lo && proposed < hi) ||
proposed == x || iteration >= 12) {
proposed = lo + 0.5 * (hi - lo);
}
x = proposed;
}
if (!converged || !isfinite(x)) {
atomicCAS(info + batch, 0, 5);
}
roots[lane] = x;
atomicMax(iterations + batch, used);
}
}
}
__syncthreads();
if (lane < count) {
if (deflated[lane]) {
updated_weights[lane] = 0.0;
} else {
const double diagonal_delta = fabs(shared_poles[lane] - roots[lane]);
if (!(diagonal_delta > 0.0) || !isfinite(diagonal_delta)) {
atomicCAS(info + batch, 0, 3);
updated_weights[lane] = CUDART_NAN;
} else {
double log_product = log(diagonal_delta);
for (int root = 0; root < count; ++root) {
// A deflated root equals its matched pole exactly, so its numerator
// and denominator factor cancel for every remaining updated weight.
if (root == lane || deflated[root]) {
continue;
}
const double numerator = fabs(shared_poles[lane] - roots[root]);
const double denominator =
fabs(shared_poles[lane] - shared_poles[root]);
if (!(numerator > 0.0) || !(denominator > 0.0) ||
!isfinite(numerator) || !isfinite(denominator)) {
atomicCAS(info + batch, 0, 3);
} else {
log_product += log(numerator) - log(denominator);
}
}
const double magnitude = exp(0.5 * log_product);
updated_weights[lane] = copysign(magnitude, shared_weights[lane]);
if (!isfinite(updated_weights[lane])) {
atomicCAS(info + batch, 0, 3);
}
}
}
}
__syncthreads();
if (lane < count) {
if (deflated[lane]) {
for (int row = 0; row < count; ++row) {
secular_vectors[matrix_base + row * kMaxMerge + lane] =
row == lane ? 1.0f : 0.0f;
}
} else {
double norm2 = 0.0;
double correction = 0.0;
for (int row = 0; row < count; ++row) {
const double delta = shared_poles[row] - roots[lane];
const double element = updated_weights[row] / delta;
kahan_add_double(element * element, norm2, correction);
}
const double norm = sqrt(norm2);
if (!(norm > 0.0) || !isfinite(norm)) {
atomicCAS(info + batch, 0, 4);
}
for (int row = 0; row < count; ++row) {
const double delta = shared_poles[row] - roots[lane];
const double element = updated_weights[row] / delta / norm;
secular_vectors[matrix_base + row * kMaxMerge + lane] =
static_cast<float>(element);
}
}
values[base + lane] = static_cast<float>(roots[lane]);
}
}
void check_tensor(
const at::Tensor& tensor,
at::ScalarType dtype,
int64_t dimensions,
const char* name) {
TORCH_CHECK(tensor.is_cuda(), name, " must be CUDA");
TORCH_CHECK(tensor.is_contiguous(), name, " must be contiguous");
TORCH_CHECK(tensor.scalar_type() == dtype, name, " has wrong dtype");
TORCH_CHECK(tensor.dim() == dimensions, name, " has wrong rank");
}
} // namespace
void secular_merge_run(
const at::Tensor& poles,
const at::Tensor& weights,
const at::Tensor& rho,
const at::Tensor& active_count,
const at::Tensor& values,
const at::Tensor& secular_vectors,
const at::Tensor& info,
const at::Tensor& iterations) {
check_tensor(poles, at::kFloat, 2, "poles");
check_tensor(weights, at::kFloat, 2, "weights");
check_tensor(rho, at::kFloat, 1, "rho");
check_tensor(active_count, at::kInt, 1, "active_count");
check_tensor(values, at::kFloat, 2, "values");
check_tensor(secular_vectors, at::kFloat, 3, "secular_vectors");
check_tensor(info, at::kInt, 1, "info");
check_tensor(iterations, at::kInt, 1, "iterations");
const int64_t batch = poles.size(0);
TORCH_CHECK(poles.size(1) == kMaxMerge, "poles must have width 64");
TORCH_CHECK(weights.sizes() == poles.sizes(), "weights shape mismatch");
TORCH_CHECK(values.sizes() == poles.sizes(), "values shape mismatch");
TORCH_CHECK(rho.size(0) == batch, "rho shape mismatch");
TORCH_CHECK(active_count.size(0) == batch, "active_count shape mismatch");
TORCH_CHECK(
secular_vectors.size(0) == batch &&
secular_vectors.size(1) == kMaxMerge &&
secular_vectors.size(2) == kMaxMerge,
"secular_vectors shape mismatch");
TORCH_CHECK(info.size(0) == batch, "info shape mismatch");
TORCH_CHECK(iterations.size(0) == batch, "iterations shape mismatch");
TORCH_CHECK(batch > 0 && batch <= std::numeric_limits<int>::max(),
"unsupported batch size");
secular_merge_kernel<<<static_cast<int>(batch), kMaxMerge, 0>>>(
poles.data_ptr<float>(),
weights.data_ptr<float>(),
rho.data_ptr<float>(),
active_count.data_ptr<int32_t>(),
values.data_ptr<float>(),
secular_vectors.data_ptr<float>(),
info.data_ptr<int32_t>(),
iterations.data_ptr<int32_t>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
void secular_merge_fp64_control_run(
const at::Tensor& poles,
const at::Tensor& weights,
const at::Tensor& rho,
const at::Tensor& active_count,
const at::Tensor& values,
const at::Tensor& secular_vectors,
const at::Tensor& info,
const at::Tensor& iterations) {
check_tensor(poles, at::kFloat, 2, "poles");
check_tensor(weights, at::kFloat, 2, "weights");
check_tensor(rho, at::kFloat, 1, "rho");
check_tensor(active_count, at::kInt, 1, "active_count");
check_tensor(values, at::kFloat, 2, "values");
check_tensor(secular_vectors, at::kFloat, 3, "secular_vectors");
check_tensor(info, at::kInt, 1, "info");
check_tensor(iterations, at::kInt, 1, "iterations");
const int64_t batch = poles.size(0);
TORCH_CHECK(poles.size(1) == kMaxMerge, "poles must have width 64");
TORCH_CHECK(weights.sizes() == poles.sizes(), "weights shape mismatch");
TORCH_CHECK(values.sizes() == poles.sizes(), "values shape mismatch");
TORCH_CHECK(rho.size(0) == batch, "rho shape mismatch");
TORCH_CHECK(active_count.size(0) == batch, "active_count shape mismatch");
TORCH_CHECK(
secular_vectors.size(0) == batch &&
secular_vectors.size(1) == kMaxMerge &&
secular_vectors.size(2) == kMaxMerge,
"secular_vectors shape mismatch");
TORCH_CHECK(info.size(0) == batch, "info shape mismatch");
TORCH_CHECK(iterations.size(0) == batch, "iterations shape mismatch");
TORCH_CHECK(batch > 0 && batch <= std::numeric_limits<int>::max(),
"unsupported batch size");
secular_merge_fp64_control_kernel<<<
static_cast<int>(batch), kMaxMerge, 0>>>(
poles.data_ptr<float>(),
weights.data_ptr<float>(),
rho.data_ptr<float>(),
active_count.data_ptr<int32_t>(),
values.data_ptr<float>(),
secular_vectors.data_ptr<float>(),
info.data_ptr<int32_t>(),
iterations.data_ptr<int32_t>(),
false);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
void secular_merge_hybrid_run(
const at::Tensor& poles,
const at::Tensor& weights,
const at::Tensor& rho,
const at::Tensor& active_count,
const at::Tensor& values,
const at::Tensor& secular_vectors,
const at::Tensor& info,
const at::Tensor& iterations) {
// Reuse the exact strict-FP32 entrypoint and its full shape checks.
secular_merge_run(
poles,
weights,
rho,
active_count,
values,
secular_vectors,
info,
iterations);
const int64_t batch = poles.size(0);
secular_merge_fp64_control_kernel<<<
static_cast<int>(batch), kMaxMerge, 0>>>(
poles.data_ptr<float>(),
weights.data_ptr<float>(),
rho.data_ptr<float>(),
active_count.data_ptr<int32_t>(),
values.data_ptr<float>(),
secular_vectors.data_ptr<float>(),
info.data_ptr<int32_t>(),
iterations.data_ptr<int32_t>(),
true);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
} // namespace lane_e64
namespace lane_e128 {
namespace {
constexpr int kMaxMerge = 128;
constexpr int kMaxIterations = 30;
__device__ __forceinline__ void kahan_add(
float value, float& total, float& compensation) {
const float adjusted = value - compensation;
const float updated = total + adjusted;
compensation = (updated - total) - adjusted;
total = updated;
}
__device__ __forceinline__ float secular_value(
const float* poles,
const float* weights,
int count,
float rho,
float x) {
// Accumulate the two signs separately. Close-pole secular evaluations
// otherwise lose the smaller signed partial sum before the final add.
float negative = 0.0f;
float positive = 0.0f;
float cnegative = 0.0f;
float cpositive = 0.0f;
for (int i = 0; i < count; ++i) {
const float term = rho * weights[i] * weights[i] / (poles[i] - x);
if (term < 0.0f) {
kahan_add(term, negative, cnegative);
} else {
kahan_add(term, positive, cpositive);
}
}
return 1.0f + (negative + positive);
}
// Status-2 is often a representation boundary rather than a missing root:
// the mathematical root lies between a pole and its first inward FP32 value.
// Evaluate only that endpoint-sign predicate in FP64. This is deliberately
// not a second FP64 root solve. A valid predicate authorizes a caller-visible
// FP32 endpoint clamp; an invalid predicate remains status 2 and is handled by
// the existing fail-closed selective-root FP64 fallback launch.
__device__ __forceinline__ double endpoint_secular_value_double(
const float* poles,
const float* weights,
int count,
float rho,
double x) {
double negative = 0.0;
double positive = 0.0;
double cnegative = 0.0;
double cpositive = 0.0;
for (int index = 0; index < count; ++index) {
const double pole = static_cast<double>(poles[index]);
const double weight = static_cast<double>(weights[index]);
const double term =
static_cast<double>(rho) * weight * weight / (pole - x);
if (term < 0.0) {
const double adjusted = term - cnegative;
const double updated = negative + adjusted;
cnegative = (updated - negative) - adjusted;
negative = updated;
} else {
const double adjusted = term - cpositive;
const double updated = positive + adjusted;
cpositive = (updated - positive) - adjusted;
positive = updated;
}
}
return 1.0 + negative + positive;
}
struct EvalFloat {
float value;
float derivative;
float error_scale;
};
__device__ __forceinline__ EvalFloat secular_eval_float(
const float* poles,
const float* weights,
int count,
float rho,
float x) {
float negative = 0.0f;
float positive = 0.0f;
float derivative = 0.0f;
float magnitude = 0.0f;
float cnegative = 0.0f;
float cpositive = 0.0f;
float cderivative = 0.0f;
float cmagnitude = 0.0f;
for (int index = 0; index < count; ++index) {
const float delta = poles[index] - x;
const float weight2 = weights[index] * weights[index];
const float term = rho * weight2 / delta;
if (term < 0.0f) {
kahan_add(term, negative, cnegative);
} else {
kahan_add(term, positive, cpositive);
}
kahan_add(rho * weight2 / (delta * delta), derivative, cderivative);
kahan_add(fabsf(term), magnitude, cmagnitude);
}
return {
1.0f + negative + positive,
derivative,
8.0f * (1.0f + magnitude + fabsf(x) * derivative)};
}
__device__ __forceinline__ float interior_rational_step_float(
const float* poles,
const float* weights,
float rho,
int index,
float x,
float value,
float derivative,
bool origin_at_lower) {
const float delta_i = poles[index] - x;
const float delta_ip1 = poles[index + 1] - x;
const float gap = poles[index + 1] - poles[index];
float c;
if (origin_at_lower) {
const float ratio = weights[index] / delta_i;
c = value - delta_ip1 * derivative + gap * rho * ratio * ratio;
} else {
const float ratio = weights[index + 1] / delta_ip1;
c = value - delta_i * derivative - gap * rho * ratio * ratio;
}
const float a =
(delta_i + delta_ip1) * value - delta_i * delta_ip1 * derivative;
const float b = delta_i * delta_ip1 * value;
if (c == 0.0f) {
return a == 0.0f ? CUDART_NAN_F : b / a;
}
const float discriminant = fmaxf(0.0f, a * a - 4.0f * b * c);
const float root = sqrtf(discriminant);
if (a <= 0.0f) {
return (a - root) / (2.0f * c);
}
const float denominator = a + root;
return denominator == 0.0f ? CUDART_NAN_F : 2.0f * b / denominator;
}
__device__ __forceinline__ float last_rational_step_float(
const float* poles,
const float* weights,
int count,
float rho,
float x,
float value) {
const float delta0 = poles[count - 2] - x;
const float delta1 = poles[count - 1] - x;
float dpsi = 0.0f;
float correction = 0.0f;
for (int index = 0; index < count - 1; ++index) {
const float ratio = weights[index] / (poles[index] - x);
kahan_add(rho * ratio * ratio, dpsi, correction);
}
const float last_ratio = weights[count - 1] / delta1;
const float dphi = rho * last_ratio * last_ratio;
const float c = fabsf(value - delta0 * dpsi - delta1 * dphi);
const float a =
(delta0 + delta1) * value - delta0 * delta1 * (dpsi + dphi);
const float b = delta0 * delta1 * value;
if (c == 0.0f) {
return CUDART_NAN_F;
}
const float discriminant = fmaxf(0.0f, a * a - 4.0f * b * c);
const float root = sqrtf(discriminant);
if (a >= 0.0f) {
return (a + root) / (2.0f * c);
}
const float denominator = a - root;
return denominator == 0.0f ? CUDART_NAN_F : 2.0f * b / denominator;
}
__global__ void secular_merge_kernel(
const float* __restrict__ poles,
const float* __restrict__ weights,
const float* __restrict__ rho,
const int32_t* __restrict__ active_count,
float* __restrict__ values,
float* __restrict__ secular_vectors,
int32_t* __restrict__ info,
int32_t* __restrict__ iterations) {
const int batch = static_cast<int>(blockIdx.x);
const int lane = static_cast<int>(threadIdx.x);
const int base = batch * kMaxMerge;
const int matrix_base = batch * kMaxMerge * kMaxMerge;
const int count = active_count[batch];
__shared__ float shared_poles[kMaxMerge];
__shared__ float shared_weights[kMaxMerge];
__shared__ float roots[kMaxMerge];
__shared__ float updated_weights[kMaxMerge];
if (lane == 0) {
info[batch] = (count > 0 && count <= kMaxMerge) ? 0 : 1;
iterations[batch] = 0;
}
if (lane < kMaxMerge) {
shared_poles[lane] = poles[base + lane];
shared_weights[lane] = weights[base + lane];
values[base + lane] = 0.0f;
for (int row = 0; row < kMaxMerge; ++row) {
secular_vectors[matrix_base + row * kMaxMerge + lane] = 0.0f;
}
}
__syncthreads();
if (lane < count) {
const float positive_inf = CUDART_INF_F;
const float negative_inf = -CUDART_INF_F;
float lo = nextafterf(shared_poles[lane], positive_inf);
float hi;
if (lane + 1 < count) {
hi = nextafterf(shared_poles[lane + 1], negative_inf);
} else {
float normz2 = 0.0f;
float correction = 0.0f;
for (int i = 0; i < count; ++i) {
kahan_add(shared_weights[i] * shared_weights[i], normz2, correction);
}
hi = shared_poles[count - 1] + rho[batch] * normz2;
if (!(hi > lo)) {
hi = nextafterf(lo, positive_inf);
}
for (int attempt = 0; attempt < 8; ++attempt) {
const float fhi = secular_value(
shared_poles, shared_weights, count, rho[batch], hi);
if (fhi >= 0.0f) {
break;
}
hi = shared_poles[count - 1] +
2.0f * (hi - shared_poles[count - 1]);
}
}
const float flo = secular_eval_float(
shared_poles, shared_weights, count, rho[batch], lo).value;
const float fhi = secular_eval_float(
shared_poles, shared_weights, count, rho[batch], hi).value;
if (!(lo < hi) || !(flo <= 0.0f) || !(fhi >= 0.0f) ||
!isfinite(lo) || !isfinite(hi)) {
bool endpoint_sign_rescued = false;
if (lo < hi && isfinite(lo) && isfinite(hi) &&
(flo > 0.0f || fhi < 0.0f)) {
const double double_lo = nextafter(
static_cast<double>(shared_poles[lane]), CUDART_INF);
double double_hi;
if (lane + 1 < count) {
double_hi = nextafter(
static_cast<double>(shared_poles[lane + 1]), -CUDART_INF);
} else {
double normz2 = 0.0;
double correction = 0.0;
for (int index = 0; index < count; ++index) {
const double weight =
static_cast<double>(shared_weights[index]);
const double value = weight * weight - correction;
const double updated = normz2 + value;
correction = (updated - normz2) - value;
normz2 = updated;
}
double_hi = static_cast<double>(shared_poles[count - 1]) +
static_cast<double>(rho[batch]) * normz2;
if (!(double_hi > double_lo)) {
double_hi = nextafter(double_lo, CUDART_INF);
}
for (int attempt = 0; attempt < 8; ++attempt) {
if (endpoint_secular_value_double(
shared_poles,
shared_weights,
count,
rho[batch],
double_hi) >= 0.0) {
break;
}
double_hi = static_cast<double>(shared_poles[count - 1]) +
2.0 *
(double_hi -
static_cast<double>(shared_poles[count - 1]));
}
}
const double double_flo = endpoint_secular_value_double(
shared_poles,
shared_weights,
count,
rho[batch],
double_lo);
const double double_fhi = endpoint_secular_value_double(
shared_poles,
shared_weights,
count,
rho[batch],
double_hi);
endpoint_sign_rescued =
isfinite(double_lo) && isfinite(double_hi) &&
isfinite(double_flo) && isfinite(double_fhi) &&
double_lo < double_hi && double_flo <= 0.0 &&
double_fhi >= 0.0;
}
if (endpoint_sign_rescued) {
// The exact root is below/above the first representable interior FP32
// endpoint. Clamp at that endpoint and retain the existing FP32
// SLAED3 displacement reconstruction below.
roots[lane] = flo > 0.0f ? lo : hi;
} else {
atomicCAS(info + batch, 0, 2);
roots[lane] = CUDART_NAN_F;
}
} else {
float x = lo + 0.5f * (hi - lo);
const bool origin_at_lower = secular_eval_float(
shared_poles, shared_weights, count, rho[batch], x).value > 0.0f;
bool converged = false;
int used = 0;
for (int iteration = 1; iteration <= kMaxIterations; ++iteration) {
used = iteration;
const EvalFloat current = secular_eval_float(
shared_poles, shared_weights, count, rho[batch], x);
if (fabsf(current.value) <= FLT_EPSILON * current.error_scale) {
converged = true;
break;
}
if (current.value <= 0.0f) {
lo = fmaxf(lo, x);
} else {
hi = fminf(hi, x);
}
const float scale = fmaxf(1.0f, fmaxf(fabsf(lo), fabsf(hi)));
if (hi - lo <= 2.0f * FLT_EPSILON * scale || lo == hi) {
x = lo + 0.5f * (hi - lo);
converged = true;
break;
}
float eta = lane + 1 < count
? interior_rational_step_float(
shared_poles,
shared_weights,
rho[batch],
lane,
x,
current.value,
current.derivative,
origin_at_lower)
: last_rational_step_float(
shared_poles,
shared_weights,
count,
rho[batch],
x,
current.value);
if (!isfinite(eta) || current.value * eta >= 0.0f) {
eta = -current.value / current.derivative;
}
float proposed = x + eta;
if (!isfinite(proposed) || !(proposed > lo && proposed < hi) ||
proposed == x || iteration >= 12) {
proposed = lo + 0.5f * (hi - lo);
}
x = proposed;
}
if (!converged || !isfinite(x)) {
atomicCAS(info + batch, 0, 5);
}
roots[lane] = x;
atomicMax(iterations + batch, used);
}
}
__syncthreads();
// SLAED3 reconstructs the rank-one weights from all root displacements.
// Evaluate the product in the log domain: this is algebraically equivalent
// for the positive rank-one problem and avoids a canary-only overflow.
if (lane < count) {
float diagonal_delta = fabsf(shared_poles[lane] - roots[lane]);
if (!(diagonal_delta > 0.0f) || !isfinite(diagonal_delta)) {
atomicCAS(info + batch, 0, 3);
updated_weights[lane] = CUDART_NAN_F;
} else {
float log_product = logf(diagonal_delta);
for (int root = 0; root < count; ++root) {
if (root == lane) {
continue;
}
const float numerator = fabsf(shared_poles[lane] - roots[root]);
const float denominator =
fabsf(shared_poles[lane] - shared_poles[root]);
if (!(numerator > 0.0f) || !(denominator > 0.0f)) {
atomicCAS(info + batch, 0, 3);
} else {
log_product += logf(numerator) - logf(denominator);
}
}
const float magnitude = expf(0.5f * log_product);
updated_weights[lane] = copysignf(magnitude, shared_weights[lane]);
if (!isfinite(updated_weights[lane])) {
atomicCAS(info + batch, 0, 3);
}
}
}
__syncthreads();
if (lane < count) {
float norm2 = 0.0f;
float correction = 0.0f;
for (int row = 0; row < count; ++row) {
const float delta = shared_poles[row] - roots[lane];
const float element = updated_weights[row] / delta;
kahan_add(element * element, norm2, correction);
}
const float norm = sqrtf(norm2);
if (!(norm > 0.0f) || !isfinite(norm)) {
atomicCAS(info + batch, 0, 4);
}
for (int row = 0; row < count; ++row) {
const float delta = shared_poles[row] - roots[lane];
secular_vectors[matrix_base + row * kMaxMerge + lane] =
updated_weights[row] / delta / norm;
}
values[base + lane] = roots[lane];
}
}
__device__ __forceinline__ void kahan_add_double(
double value, double& total, double& compensation) {
const double adjusted = value - compensation;
const double updated = total + adjusted;
compensation = (updated - total) - adjusted;
total = updated;
}
__device__ __forceinline__ double secular_value_double(
const double* poles,
const double* weights,
int count,
double rho,
double x) {
double negative = 0.0;
double positive = 0.0;
double cnegative = 0.0;
double cpositive = 0.0;
for (int i = 0; i < count; ++i) {
const double term = rho * weights[i] * weights[i] / (poles[i] - x);
if (term < 0.0) {
kahan_add_double(term, negative, cnegative);
} else {
kahan_add_double(term, positive, cpositive);
}
}
return 1.0 + (negative + positive);
}
struct EvalDouble {
double value;
double derivative;
double error_scale;
};
__device__ __forceinline__ EvalDouble secular_eval_double(
const double* poles,
const double* weights,
int count,
double rho,
double x) {
double negative = 0.0;
double positive = 0.0;
double derivative = 0.0;
double magnitude = 0.0;
double cnegative = 0.0;
double cpositive = 0.0;
double cderivative = 0.0;
double cmagnitude = 0.0;
for (int index = 0; index < count; ++index) {
const double delta = poles[index] - x;
const double weight2 = weights[index] * weights[index];
const double term = rho * weight2 / delta;
if (term < 0.0) {
kahan_add_double(term, negative, cnegative);
} else {
kahan_add_double(term, positive, cpositive);
}
kahan_add_double(
rho * weight2 / (delta * delta), derivative, cderivative);
kahan_add_double(fabs(term), magnitude, cmagnitude);
}
return {
1.0 + negative + positive,
derivative,
8.0 * (1.0 + magnitude + fabs(x) * derivative)};
}
__device__ __forceinline__ double interior_rational_step_double(
const double* poles,
const double* weights,
double rho,
int index,
double x,
double value,
double derivative,
bool origin_at_lower) {
const double delta_i = poles[index] - x;
const double delta_ip1 = poles[index + 1] - x;
const double gap = poles[index + 1] - poles[index];
double c;
if (origin_at_lower) {
const double ratio = weights[index] / delta_i;
c = value - delta_ip1 * derivative + gap * rho * ratio * ratio;
} else {
const double ratio = weights[index + 1] / delta_ip1;
c = value - delta_i * derivative - gap * rho * ratio * ratio;
}
const double a =
(delta_i + delta_ip1) * value - delta_i * delta_ip1 * derivative;
const double b = delta_i * delta_ip1 * value;
if (c == 0.0) {
return a == 0.0 ? CUDART_NAN : b / a;
}
const double root = sqrt(fmax(0.0, a * a - 4.0 * b * c));
if (a <= 0.0) {
return (a - root) / (2.0 * c);
}
const double denominator = a + root;
return denominator == 0.0 ? CUDART_NAN : 2.0 * b / denominator;
}
__device__ __forceinline__ double last_rational_step_double(
const double* poles,
const double* weights,
int count,
double rho,
double x,
double value) {
const double delta0 = poles[count - 2] - x;
const double delta1 = poles[count - 1] - x;
double dpsi = 0.0;
double correction = 0.0;
for (int index = 0; index < count - 1; ++index) {
const double ratio = weights[index] / (poles[index] - x);
kahan_add_double(rho * ratio * ratio, dpsi, correction);
}
const double last_ratio = weights[count - 1] / delta1;
const double dphi = rho * last_ratio * last_ratio;
const double c = fabs(value - delta0 * dpsi - delta1 * dphi);
const double a =
(delta0 + delta1) * value - delta0 * delta1 * (dpsi + dphi);
const double b = delta0 * delta1 * value;
if (c == 0.0) {
return CUDART_NAN;
}
const double root = sqrt(fmax(0.0, a * a - 4.0 * b * c));
if (a >= 0.0) {
return (a + root) / (2.0 * c);
}
const double denominator = a - root;
return denominator == 0.0 ? CUDART_NAN : 2.0 * b / denominator;
}
__global__ void secular_merge_fp64_control_kernel(
const float* __restrict__ poles,
const float* __restrict__ weights,
const float* __restrict__ rho,
const int32_t* __restrict__ active_count,
float* __restrict__ values,
float* __restrict__ secular_vectors,
int32_t* __restrict__ info,
int32_t* __restrict__ iterations,
bool fallback_only) {
const int batch = static_cast<int>(blockIdx.x);
const int lane = static_cast<int>(threadIdx.x);
const int base = batch * kMaxMerge;
const int matrix_base = batch * kMaxMerge * kMaxMerge;
const int count = active_count[batch];
__shared__ double shared_poles[kMaxMerge];
__shared__ double shared_weights[kMaxMerge];
__shared__ double roots[kMaxMerge];
__shared__ double updated_weights[kMaxMerge];
__shared__ int32_t deflated[kMaxMerge];
__shared__ int32_t prior_info;
// Snapshot the preceding strict-FP32 status before lane zero clears it.
// The barrier makes both the early return and repair route uniform across
// the CTA's two warps.
if (lane == 0) {
prior_info = fallback_only ? info[batch] : 0;
}
__syncthreads();
if (fallback_only && prior_info == 0) {
return;
}
// Status 2 is an endpoint-sign failure. The primary launch leaves finite
// values for roots whose own brackets succeeded and NaN for failed roots.
const float primary_root = lane < count ? values[base + lane] : CUDART_NAN_F;
const bool reuse_primary_root =
fallback_only && prior_info == 2 && isfinite(primary_root);
if (lane == 0) {
info[batch] = (count > 0 && count <= kMaxMerge) ? 0 : 1;
iterations[batch] = 0;
}
shared_poles[lane] = static_cast<double>(poles[base + lane]);
shared_weights[lane] = static_cast<double>(weights[base + lane]);
deflated[lane] = 0;
if (!reuse_primary_root) {
values[base + lane] = 0.0f;
}
for (int row = 0; row < kMaxMerge; ++row) {
secular_vectors[matrix_base + row * kMaxMerge + lane] = 0.0f;
}
__syncthreads();
if (lane < count) {
if (reuse_primary_root) {
roots[lane] = static_cast<double>(primary_root);
} else {
double lo = nextafter(shared_poles[lane], CUDART_INF);
double hi;
if (lane + 1 < count) {
hi = nextafter(shared_poles[lane + 1], -CUDART_INF);
} else {
double normz2 = 0.0;
double correction = 0.0;
for (int i = 0; i < count; ++i) {
kahan_add_double(
shared_weights[i] * shared_weights[i], normz2, correction);
}
hi = shared_poles[count - 1] + static_cast<double>(rho[batch]) * normz2;
if (!(hi > lo)) {
hi = nextafter(lo, CUDART_INF);
}
for (int attempt = 0; attempt < 8; ++attempt) {
const double fhi = secular_value_double(
shared_poles,
shared_weights,
count,
static_cast<double>(rho[batch]),
hi);
if (fhi >= 0.0) {
break;
}
hi = shared_poles[count - 1] +
2.0 * (hi - shared_poles[count - 1]);
}
}
const double flo = secular_eval_double(
shared_poles,
shared_weights,
count,
static_cast<double>(rho[batch]),
lo).value;
const double fhi = secular_eval_double(
shared_poles,
shared_weights,
count,
static_cast<double>(rho[batch]),
hi).value;
const bool finite_bracket =
(lo < hi) && isfinite(lo) && isfinite(hi) &&
isfinite(flo) && isfinite(fhi);
const bool sub_ulp_deflation =
finite_bracket && flo >= 0.0 && fhi > 0.0;
if (sub_ulp_deflation) {
// The first FP64 number above the pole is already beyond the
// root: its displacement is unrepresentable in double.
deflated[lane] = 1;
roots[lane] = shared_poles[lane];
} else if (!finite_bracket || !(flo < 0.0) || !(fhi > 0.0)) {
atomicCAS(info + batch, 0, 2);
roots[lane] = CUDART_NAN;
} else {
const double rho_value = static_cast<double>(rho[batch]);
double x = lo + 0.5 * (hi - lo);
const bool origin_at_lower = secular_eval_double(
shared_poles, shared_weights, count, rho_value, x).value > 0.0;
bool converged = false;
int used = 0;
for (int iteration = 1; iteration <= kMaxIterations; ++iteration) {
used = iteration;
const EvalDouble current = secular_eval_double(
shared_poles, shared_weights, count, rho_value, x);
if (fabs(current.value) <= DBL_EPSILON * current.error_scale) {
converged = true;
break;
}
if (current.value <= 0.0) {
lo = fmax(lo, x);
} else {
hi = fmin(hi, x);
}
const double scale = fmax(1.0, fmax(fabs(lo), fabs(hi)));
if (hi - lo <= 2.0 * static_cast<double>(FLT_EPSILON) * scale ||
static_cast<float>(lo) == static_cast<float>(hi)) {
x = lo + 0.5 * (hi - lo);
converged = true;
break;
}
double eta = lane + 1 < count
? interior_rational_step_double(
shared_poles,
shared_weights,
rho_value,
lane,
x,
current.value,
current.derivative,
origin_at_lower)
: last_rational_step_double(
shared_poles,
shared_weights,
count,
rho_value,
x,
current.value);
if (!isfinite(eta) || current.value * eta >= 0.0) {
eta = -current.value / current.derivative;
}
double proposed = x + eta;
if (!isfinite(proposed) || !(proposed > lo && proposed < hi) ||
proposed == x || iteration >= 12) {
proposed = lo + 0.5 * (hi - lo);
}
x = proposed;
}
if (!converged || !isfinite(x)) {
atomicCAS(info + batch, 0, 5);
}
roots[lane] = x;
atomicMax(iterations + batch, used);
}
}
}
__syncthreads();
if (lane < count) {
if (deflated[lane]) {
updated_weights[lane] = 0.0;
} else {
const double diagonal_delta = fabs(shared_poles[lane] - roots[lane]);
if (!(diagonal_delta > 0.0) || !isfinite(diagonal_delta)) {
atomicCAS(info + batch, 0, 3);
updated_weights[lane] = CUDART_NAN;
} else {
double log_product = log(diagonal_delta);
for (int root = 0; root < count; ++root) {
// A deflated root equals its matched pole exactly, so its numerator
// and denominator factor cancel for every remaining updated weight.
if (root == lane || deflated[root]) {
continue;
}
const double numerator = fabs(shared_poles[lane] - roots[root]);
const double denominator =
fabs(shared_poles[lane] - shared_poles[root]);
if (!(numerator > 0.0) || !(denominator > 0.0) ||
!isfinite(numerator) || !isfinite(denominator)) {
atomicCAS(info + batch, 0, 3);
} else {
log_product += log(numerator) - log(denominator);
}
}
const double magnitude = exp(0.5 * log_product);
updated_weights[lane] = copysign(magnitude, shared_weights[lane]);
if (!isfinite(updated_weights[lane])) {
atomicCAS(info + batch, 0, 3);
}
}
}
}
__syncthreads();
if (lane < count) {
if (deflated[lane]) {
for (int row = 0; row < count; ++row) {
secular_vectors[matrix_base + row * kMaxMerge + lane] =
row == lane ? 1.0f : 0.0f;
}
} else {
double norm2 = 0.0;
double correction = 0.0;
for (int row = 0; row < count; ++row) {
const double delta = shared_poles[row] - roots[lane];
const double element = updated_weights[row] / delta;
kahan_add_double(element * element, norm2, correction);
}
const double norm = sqrt(norm2);
if (!(norm > 0.0) || !isfinite(norm)) {
atomicCAS(info + batch, 0, 4);
}
for (int row = 0; row < count; ++row) {
const double delta = shared_poles[row] - roots[lane];
const double element = updated_weights[row] / delta / norm;
secular_vectors[matrix_base + row * kMaxMerge + lane] =
static_cast<float>(element);
}
}
values[base + lane] = static_cast<float>(roots[lane]);
}
}
void check_tensor(
const at::Tensor& tensor,
at::ScalarType dtype,
int64_t dimensions,
const char* name) {
TORCH_CHECK(tensor.is_cuda(), name, " must be CUDA");
TORCH_CHECK(tensor.is_contiguous(), name, " must be contiguous");
TORCH_CHECK(tensor.scalar_type() == dtype, name, " has wrong dtype");
TORCH_CHECK(tensor.dim() == dimensions, name, " has wrong rank");
}
} // namespace
void secular_merge_run(
const at::Tensor& poles,
const at::Tensor& weights,
const at::Tensor& rho,
const at::Tensor& active_count,
const at::Tensor& values,
const at::Tensor& secular_vectors,
const at::Tensor& info,
const at::Tensor& iterations) {
check_tensor(poles, at::kFloat, 2, "poles");
check_tensor(weights, at::kFloat, 2, "weights");
check_tensor(rho, at::kFloat, 1, "rho");
check_tensor(active_count, at::kInt, 1, "active_count");
check_tensor(values, at::kFloat, 2, "values");
check_tensor(secular_vectors, at::kFloat, 3, "secular_vectors");
check_tensor(info, at::kInt, 1, "info");
check_tensor(iterations, at::kInt, 1, "iterations");
const int64_t batch = poles.size(0);
TORCH_CHECK(poles.size(1) == kMaxMerge, "poles must have width 128");
TORCH_CHECK(weights.sizes() == poles.sizes(), "weights shape mismatch");
TORCH_CHECK(values.sizes() == poles.sizes(), "values shape mismatch");
TORCH_CHECK(rho.size(0) == batch, "rho shape mismatch");
TORCH_CHECK(active_count.size(0) == batch, "active_count shape mismatch");
TORCH_CHECK(
secular_vectors.size(0) == batch &&
secular_vectors.size(1) == kMaxMerge &&
secular_vectors.size(2) == kMaxMerge,
"secular_vectors shape mismatch");
TORCH_CHECK(info.size(0) == batch, "info shape mismatch");
TORCH_CHECK(iterations.size(0) == batch, "iterations shape mismatch");
TORCH_CHECK(batch > 0 && batch <= std::numeric_limits<int>::max(),
"unsupported batch size");
secular_merge_kernel<<<static_cast<int>(batch), kMaxMerge, 0>>>(
poles.data_ptr<float>(),
weights.data_ptr<float>(),
rho.data_ptr<float>(),
active_count.data_ptr<int32_t>(),
values.data_ptr<float>(),
secular_vectors.data_ptr<float>(),
info.data_ptr<int32_t>(),
iterations.data_ptr<int32_t>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
void secular_merge_fp64_control_run(
const at::Tensor& poles,
const at::Tensor& weights,
const at::Tensor& rho,
const at::Tensor& active_count,
const at::Tensor& values,
const at::Tensor& secular_vectors,
const at::Tensor& info,
const at::Tensor& iterations) {
check_tensor(poles, at::kFloat, 2, "poles");
check_tensor(weights, at::kFloat, 2, "weights");
check_tensor(rho, at::kFloat, 1, "rho");
check_tensor(active_count, at::kInt, 1, "active_count");
check_tensor(values, at::kFloat, 2, "values");
check_tensor(secular_vectors, at::kFloat, 3, "secular_vectors");
check_tensor(info, at::kInt, 1, "info");
check_tensor(iterations, at::kInt, 1, "iterations");
const int64_t batch = poles.size(0);
TORCH_CHECK(poles.size(1) == kMaxMerge, "poles must have width 128");
TORCH_CHECK(weights.sizes() == poles.sizes(), "weights shape mismatch");
TORCH_CHECK(values.sizes() == poles.sizes(), "values shape mismatch");
TORCH_CHECK(rho.size(0) == batch, "rho shape mismatch");
TORCH_CHECK(active_count.size(0) == batch, "active_count shape mismatch");
TORCH_CHECK(
secular_vectors.size(0) == batch &&
secular_vectors.size(1) == kMaxMerge &&
secular_vectors.size(2) == kMaxMerge,
"secular_vectors shape mismatch");
TORCH_CHECK(info.size(0) == batch, "info shape mismatch");
TORCH_CHECK(iterations.size(0) == batch, "iterations shape mismatch");
TORCH_CHECK(batch > 0 && batch <= std::numeric_limits<int>::max(),
"unsupported batch size");
secular_merge_fp64_control_kernel<<<
static_cast<int>(batch), kMaxMerge, 0>>>(
poles.data_ptr<float>(),
weights.data_ptr<float>(),
rho.data_ptr<float>(),
active_count.data_ptr<int32_t>(),
values.data_ptr<float>(),
secular_vectors.data_ptr<float>(),
info.data_ptr<int32_t>(),
iterations.data_ptr<int32_t>(),
false);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
void secular_merge_hybrid_run(
const at::Tensor& poles,
const at::Tensor& weights,
const at::Tensor& rho,
const at::Tensor& active_count,
const at::Tensor& values,
const at::Tensor& secular_vectors,
const at::Tensor& info,
const at::Tensor& iterations) {
// Reuse the exact strict-FP32 entrypoint and its full shape checks.
secular_merge_run(
poles,
weights,
rho,
active_count,
values,
secular_vectors,
info,
iterations);
const int64_t batch = poles.size(0);
secular_merge_fp64_control_kernel<<<
static_cast<int>(batch), kMaxMerge, 0>>>(
poles.data_ptr<float>(),
weights.data_ptr<float>(),
rho.data_ptr<float>(),
active_count.data_ptr<int32_t>(),
values.data_ptr<float>(),
secular_vectors.data_ptr<float>(),
info.data_ptr<int32_t>(),
iterations.data_ptr<int32_t>(),
true);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
} // namespace lane_e128
namespace lane_e256 {
namespace {
constexpr int kMaxMerge = 256;
constexpr int kMaxIterations = 30;
constexpr int kMaxFp64Iterations = 80;
__device__ __forceinline__ void kahan_add(
float value, float& total, float& compensation) {
const float adjusted = value - compensation;
const float updated = total + adjusted;
compensation = (updated - total) - adjusted;
total = updated;
}
__device__ __forceinline__ float secular_value(
const float* poles,
const float* weights,
int count,
float rho,
float x) {
// Accumulate the two signs separately. Close-pole secular evaluations
// otherwise lose the smaller signed partial sum before the final add.
float negative = 0.0f;
float positive = 0.0f;
float cnegative = 0.0f;
float cpositive = 0.0f;
for (int i = 0; i < count; ++i) {
const float term = rho * weights[i] * weights[i] / (poles[i] - x);
if (term < 0.0f) {
kahan_add(term, negative, cnegative);
} else {
kahan_add(term, positive, cpositive);
}
}
return 1.0f + (negative + positive);
}
// Status-2 is often a representation boundary rather than a missing root:
// the mathematical root lies between a pole and its first inward FP32 value.
// Evaluate only that endpoint-sign predicate in FP64. This is deliberately
// not a second FP64 root solve. A valid predicate authorizes a caller-visible
// FP32 endpoint clamp; an invalid predicate remains status 2 and is handled by
// the existing fail-closed selective-root FP64 fallback launch.
__device__ __forceinline__ double endpoint_secular_value_double(
const float* poles,
const float* weights,
int count,
float rho,
double x) {
double negative = 0.0;
double positive = 0.0;
double cnegative = 0.0;
double cpositive = 0.0;
for (int index = 0; index < count; ++index) {
const double pole = static_cast<double>(poles[index]);
const double weight = static_cast<double>(weights[index]);
const double term =
static_cast<double>(rho) * weight * weight / (pole - x);
if (term < 0.0) {
const double adjusted = term - cnegative;
const double updated = negative + adjusted;
cnegative = (updated - negative) - adjusted;
negative = updated;
} else {
const double adjusted = term - cpositive;
const double updated = positive + adjusted;
cpositive = (updated - positive) - adjusted;
positive = updated;
}
}
return 1.0 + negative + positive;
}
struct EvalFloat {
float value;
float derivative;
float error_scale;
};
__device__ __forceinline__ EvalFloat secular_eval_float(
const float* poles,
const float* weights,
int count,
float rho,
float x) {
float negative = 0.0f;
float positive = 0.0f;
float derivative = 0.0f;
float magnitude = 0.0f;
float cnegative = 0.0f;
float cpositive = 0.0f;
float cderivative = 0.0f;
float cmagnitude = 0.0f;
for (int index = 0; index < count; ++index) {
const float delta = poles[index] - x;
const float weight2 = weights[index] * weights[index];
const float term = rho * weight2 / delta;
if (term < 0.0f) {
kahan_add(term, negative, cnegative);
} else {
kahan_add(term, positive, cpositive);
}
kahan_add(rho * weight2 / (delta * delta), derivative, cderivative);
kahan_add(fabsf(term), magnitude, cmagnitude);
}
return {
1.0f + negative + positive,
derivative,
8.0f * (1.0f + magnitude + fabsf(x) * derivative)};
}
__device__ __forceinline__ float interior_rational_step_float(
const float* poles,
const float* weights,
float rho,
int index,
float x,
float value,
float derivative,
bool origin_at_lower) {
const float delta_i = poles[index] - x;
const float delta_ip1 = poles[index + 1] - x;
const float gap = poles[index + 1] - poles[index];
float c;
if (origin_at_lower) {
const float ratio = weights[index] / delta_i;
c = value - delta_ip1 * derivative + gap * rho * ratio * ratio;
} else {
const float ratio = weights[index + 1] / delta_ip1;
c = value - delta_i * derivative - gap * rho * ratio * ratio;
}
const float a =
(delta_i + delta_ip1) * value - delta_i * delta_ip1 * derivative;
const float b = delta_i * delta_ip1 * value;
if (c == 0.0f) {
return a == 0.0f ? CUDART_NAN_F : b / a;
}
const float discriminant = fmaxf(0.0f, a * a - 4.0f * b * c);
const float root = sqrtf(discriminant);
if (a <= 0.0f) {
return (a - root) / (2.0f * c);
}
const float denominator = a + root;
return denominator == 0.0f ? CUDART_NAN_F : 2.0f * b / denominator;
}
__device__ __forceinline__ float last_rational_step_float(
const float* poles,
const float* weights,
int count,
float rho,
float x,
float value) {
const float delta0 = poles[count - 2] - x;
const float delta1 = poles[count - 1] - x;
float dpsi = 0.0f;
float correction = 0.0f;
for (int index = 0; index < count - 1; ++index) {
const float ratio = weights[index] / (poles[index] - x);
kahan_add(rho * ratio * ratio, dpsi, correction);
}
const float last_ratio = weights[count - 1] / delta1;
const float dphi = rho * last_ratio * last_ratio;
const float c = fabsf(value - delta0 * dpsi - delta1 * dphi);
const float a =
(delta0 + delta1) * value - delta0 * delta1 * (dpsi + dphi);
const float b = delta0 * delta1 * value;
if (c == 0.0f) {
return CUDART_NAN_F;
}
const float discriminant = fmaxf(0.0f, a * a - 4.0f * b * c);
const float root = sqrtf(discriminant);
if (a >= 0.0f) {
return (a + root) / (2.0f * c);
}
const float denominator = a - root;
return denominator == 0.0f ? CUDART_NAN_F : 2.0f * b / denominator;
}
__global__ void secular_merge_kernel(
const float* __restrict__ poles,
const float* __restrict__ weights,
const float* __restrict__ rho,
const int32_t* __restrict__ active_count,
float* __restrict__ values,
float* __restrict__ secular_vectors,
int32_t* __restrict__ info,
int32_t* __restrict__ iterations) {
const int batch = static_cast<int>(blockIdx.x);
const int lane = static_cast<int>(threadIdx.x);
const int base = batch * kMaxMerge;
const int matrix_base = batch * kMaxMerge * kMaxMerge;
const int count = active_count[batch];
__shared__ float shared_poles[kMaxMerge];
__shared__ float shared_weights[kMaxMerge];
__shared__ float roots[kMaxMerge];
__shared__ float updated_weights[kMaxMerge];
if (lane == 0) {
info[batch] = (count > 0 && count <= kMaxMerge) ? 0 : 1;
iterations[batch] = 0;
}
if (lane < kMaxMerge) {
shared_poles[lane] = poles[base + lane];
shared_weights[lane] = weights[base + lane];
values[base + lane] = 0.0f;
for (int row = 0; row < kMaxMerge; ++row) {
secular_vectors[matrix_base + row * kMaxMerge + lane] = 0.0f;
}
}
__syncthreads();
if (lane < count) {
const float positive_inf = CUDART_INF_F;
const float negative_inf = -CUDART_INF_F;
float lo = nextafterf(shared_poles[lane], positive_inf);
float hi;
if (lane + 1 < count) {
hi = nextafterf(shared_poles[lane + 1], negative_inf);
} else {
float normz2 = 0.0f;
float correction = 0.0f;
for (int i = 0; i < count; ++i) {
kahan_add(shared_weights[i] * shared_weights[i], normz2, correction);
}
hi = shared_poles[count - 1] + rho[batch] * normz2;
if (!(hi > lo)) {
hi = nextafterf(lo, positive_inf);
}
for (int attempt = 0; attempt < 8; ++attempt) {
const float fhi = secular_value(
shared_poles, shared_weights, count, rho[batch], hi);
if (fhi >= 0.0f) {
break;
}
hi = shared_poles[count - 1] +
2.0f * (hi - shared_poles[count - 1]);
}
}
const float flo = secular_eval_float(
shared_poles, shared_weights, count, rho[batch], lo).value;
const float fhi = secular_eval_float(
shared_poles, shared_weights, count, rho[batch], hi).value;
if (!(lo < hi) || !(flo <= 0.0f) || !(fhi >= 0.0f) ||
!isfinite(lo) || !isfinite(hi)) {
bool endpoint_sign_rescued = false;
if (lo < hi && isfinite(lo) && isfinite(hi) &&
(flo > 0.0f || fhi < 0.0f)) {
const double double_lo = nextafter(
static_cast<double>(shared_poles[lane]), CUDART_INF);
double double_hi;
if (lane + 1 < count) {
double_hi = nextafter(
static_cast<double>(shared_poles[lane + 1]), -CUDART_INF);
} else {
double normz2 = 0.0;
double correction = 0.0;
for (int index = 0; index < count; ++index) {
const double weight =
static_cast<double>(shared_weights[index]);
const double value = weight * weight - correction;
const double updated = normz2 + value;
correction = (updated - normz2) - value;
normz2 = updated;
}
double_hi = static_cast<double>(shared_poles[count - 1]) +
static_cast<double>(rho[batch]) * normz2;
if (!(double_hi > double_lo)) {
double_hi = nextafter(double_lo, CUDART_INF);
}
for (int attempt = 0; attempt < 8; ++attempt) {
if (endpoint_secular_value_double(
shared_poles,
shared_weights,
count,
rho[batch],
double_hi) >= 0.0) {
break;
}
double_hi = static_cast<double>(shared_poles[count - 1]) +
2.0 *
(double_hi -
static_cast<double>(shared_poles[count - 1]));
}
}
const double double_flo = endpoint_secular_value_double(
shared_poles,
shared_weights,
count,
rho[batch],
double_lo);
const double double_fhi = endpoint_secular_value_double(
shared_poles,
shared_weights,
count,
rho[batch],
double_hi);
endpoint_sign_rescued =
isfinite(double_lo) && isfinite(double_hi) &&
isfinite(double_flo) && isfinite(double_fhi) &&
double_lo < double_hi && double_flo <= 0.0 &&
double_fhi >= 0.0;
}
if (endpoint_sign_rescued) {
// The exact root is below/above the first representable interior FP32
// endpoint. Clamp at that endpoint and retain the existing FP32
// SLAED3 displacement reconstruction below.
roots[lane] = flo > 0.0f ? lo : hi;
} else {
atomicCAS(info + batch, 0, 2);
roots[lane] = CUDART_NAN_F;
}
} else {
float x = lo + 0.5f * (hi - lo);
const bool origin_at_lower = secular_eval_float(
shared_poles, shared_weights, count, rho[batch], x).value > 0.0f;
bool converged = false;
int used = 0;
for (int iteration = 1; iteration <= kMaxIterations; ++iteration) {
used = iteration;
const EvalFloat current = secular_eval_float(
shared_poles, shared_weights, count, rho[batch], x);
if (fabsf(current.value) <= FLT_EPSILON * current.error_scale) {
converged = true;
break;
}
if (current.value <= 0.0f) {
lo = fmaxf(lo, x);
} else {
hi = fminf(hi, x);
}
const float scale = fmaxf(1.0f, fmaxf(fabsf(lo), fabsf(hi)));
if (hi - lo <= 2.0f * FLT_EPSILON * scale || lo == hi) {
x = lo + 0.5f * (hi - lo);
converged = true;
break;
}
float eta = lane + 1 < count
? interior_rational_step_float(
shared_poles,
shared_weights,
rho[batch],
lane,
x,
current.value,
current.derivative,
origin_at_lower)
: last_rational_step_float(
shared_poles,
shared_weights,
count,
rho[batch],
x,
current.value);
if (!isfinite(eta) || current.value * eta >= 0.0f) {
eta = -current.value / current.derivative;
}
float proposed = x + eta;
if (!isfinite(proposed) || !(proposed > lo && proposed < hi) ||
proposed == x || iteration >= 12) {
proposed = lo + 0.5f * (hi - lo);
}
x = proposed;
}
if (!converged || !isfinite(x)) {
atomicCAS(info + batch, 0, 5);
}
roots[lane] = x;
atomicMax(iterations + batch, used);
}
}
__syncthreads();
// SLAED3 reconstructs the rank-one weights from all root displacements.
// Evaluate the product in the log domain: this is algebraically equivalent
// for the positive rank-one problem and avoids a canary-only overflow.
if (lane < count) {
float diagonal_delta = fabsf(shared_poles[lane] - roots[lane]);
if (!(diagonal_delta > 0.0f) || !isfinite(diagonal_delta)) {
atomicCAS(info + batch, 0, 3);
updated_weights[lane] = CUDART_NAN_F;
} else {
float log_product = logf(diagonal_delta);
for (int root = 0; root < count; ++root) {
if (root == lane) {
continue;
}
const float numerator = fabsf(shared_poles[lane] - roots[root]);
const float denominator =
fabsf(shared_poles[lane] - shared_poles[root]);
if (!(numerator > 0.0f) || !(denominator > 0.0f)) {
atomicCAS(info + batch, 0, 3);
} else {
log_product += logf(numerator) - logf(denominator);
}
}
const float magnitude = expf(0.5f * log_product);
updated_weights[lane] = copysignf(magnitude, shared_weights[lane]);
if (!isfinite(updated_weights[lane])) {
atomicCAS(info + batch, 0, 3);
}
}
}
__syncthreads();
if (lane < count) {
float norm2 = 0.0f;
float correction = 0.0f;
for (int row = 0; row < count; ++row) {
const float delta = shared_poles[row] - roots[lane];
const float element = updated_weights[row] / delta;
kahan_add(element * element, norm2, correction);
}
const float norm = sqrtf(norm2);
if (!(norm > 0.0f) || !isfinite(norm)) {
atomicCAS(info + batch, 0, 4);
}
for (int row = 0; row < count; ++row) {
const float delta = shared_poles[row] - roots[lane];
secular_vectors[matrix_base + row * kMaxMerge + lane] =
updated_weights[row] / delta / norm;
}
values[base + lane] = roots[lane];
}
}
__device__ __forceinline__ void kahan_add_double(
double value, double& total, double& compensation) {
const double adjusted = value - compensation;
const double updated = total + adjusted;
compensation = (updated - total) - adjusted;
total = updated;
}
__device__ __forceinline__ double secular_value_double(
const double* poles,
const double* weights,
int count,
double rho,
double x) {
double negative = 0.0;
double positive = 0.0;
double cnegative = 0.0;
double cpositive = 0.0;
for (int i = 0; i < count; ++i) {
const double term = rho * weights[i] * weights[i] / (poles[i] - x);
if (term < 0.0) {
kahan_add_double(term, negative, cnegative);
} else {
kahan_add_double(term, positive, cpositive);
}
}
return 1.0 + (negative + positive);
}
struct EvalDouble {
double value;
double derivative;
double error_scale;
};
__device__ __forceinline__ EvalDouble secular_eval_double(
const double* poles,
const double* weights,
int count,
double rho,
double x) {
double negative = 0.0;
double positive = 0.0;
double derivative = 0.0;
double magnitude = 0.0;
double cnegative = 0.0;
double cpositive = 0.0;
double cderivative = 0.0;
double cmagnitude = 0.0;
for (int index = 0; index < count; ++index) {
const double delta = poles[index] - x;
const double weight2 = weights[index] * weights[index];
const double term = rho * weight2 / delta;
if (term < 0.0) {
kahan_add_double(term, negative, cnegative);
} else {
kahan_add_double(term, positive, cpositive);
}
kahan_add_double(
rho * weight2 / (delta * delta), derivative, cderivative);
kahan_add_double(fabs(term), magnitude, cmagnitude);
}
return {
1.0 + negative + positive,
derivative,
8.0 * (1.0 + magnitude + fabs(x) * derivative)};
}
__device__ __forceinline__ double interior_rational_step_double(
const double* poles,
const double* weights,
double rho,
int index,
double x,
double value,
double derivative,
bool origin_at_lower) {
const double delta_i = poles[index] - x;
const double delta_ip1 = poles[index + 1] - x;
const double gap = poles[index + 1] - poles[index];
double c;
if (origin_at_lower) {
const double ratio = weights[index] / delta_i;
c = value - delta_ip1 * derivative + gap * rho * ratio * ratio;
} else {
const double ratio = weights[index + 1] / delta_ip1;
c = value - delta_i * derivative - gap * rho * ratio * ratio;
}
const double a =
(delta_i + delta_ip1) * value - delta_i * delta_ip1 * derivative;
const double b = delta_i * delta_ip1 * value;
if (c == 0.0) {
return a == 0.0 ? CUDART_NAN : b / a;
}
const double root = sqrt(fmax(0.0, a * a - 4.0 * b * c));
if (a <= 0.0) {
return (a - root) / (2.0 * c);
}
const double denominator = a + root;
return denominator == 0.0 ? CUDART_NAN : 2.0 * b / denominator;
}
__device__ __forceinline__ double last_rational_step_double(
const double* poles,
const double* weights,
int count,
double rho,
double x,
double value) {
const double delta0 = poles[count - 2] - x;
const double delta1 = poles[count - 1] - x;
double dpsi = 0.0;
double correction = 0.0;
for (int index = 0; index < count - 1; ++index) {
const double ratio = weights[index] / (poles[index] - x);
kahan_add_double(rho * ratio * ratio, dpsi, correction);
}
const double last_ratio = weights[count - 1] / delta1;
const double dphi = rho * last_ratio * last_ratio;
const double c = fabs(value - delta0 * dpsi - delta1 * dphi);
const double a =
(delta0 + delta1) * value - delta0 * delta1 * (dpsi + dphi);
const double b = delta0 * delta1 * value;
if (c == 0.0) {
return CUDART_NAN;
}
const double root = sqrt(fmax(0.0, a * a - 4.0 * b * c));
if (a >= 0.0) {
return (a + root) / (2.0 * c);
}
const double denominator = a - root;
return denominator == 0.0 ? CUDART_NAN : 2.0 * b / denominator;
}
__global__ void secular_merge_fp64_control_kernel(
const float* __restrict__ poles,
const float* __restrict__ weights,
const float* __restrict__ rho,
const int32_t* __restrict__ active_count,
float* __restrict__ values,
float* __restrict__ secular_vectors,
int32_t* __restrict__ info,
int32_t* __restrict__ iterations,
bool fallback_only) {
const int batch = static_cast<int>(blockIdx.x);
const int lane = static_cast<int>(threadIdx.x);
const int base = batch * kMaxMerge;
const int matrix_base = batch * kMaxMerge * kMaxMerge;
const int count = active_count[batch];
__shared__ double shared_poles[kMaxMerge];
__shared__ double shared_weights[kMaxMerge];
__shared__ double roots[kMaxMerge];
__shared__ double updated_weights[kMaxMerge];
__shared__ int32_t deflated[kMaxMerge];
__shared__ int32_t prior_info;
// Snapshot the preceding strict-FP32 status before lane zero clears it.
// The barrier makes both the early return and repair route uniform across
// the CTA's two warps.
if (lane == 0) {
prior_info = fallback_only ? info[batch] : 0;
}
__syncthreads();
if (fallback_only && prior_info == 0) {
return;
}
// Status 2 is an endpoint-sign failure. The primary launch leaves finite
// values for roots whose own brackets succeeded and NaN for failed roots.
const float primary_root = lane < count ? values[base + lane] : CUDART_NAN_F;
const bool reuse_primary_root =
fallback_only && prior_info == 2 && isfinite(primary_root);
if (lane == 0) {
info[batch] = (count > 0 && count <= kMaxMerge) ? 0 : 1;
iterations[batch] = 0;
}
shared_poles[lane] = static_cast<double>(poles[base + lane]);
shared_weights[lane] = static_cast<double>(weights[base + lane]);
deflated[lane] = 0;
if (!reuse_primary_root) {
values[base + lane] = 0.0f;
}
for (int row = 0; row < kMaxMerge; ++row) {
secular_vectors[matrix_base + row * kMaxMerge + lane] = 0.0f;
}
__syncthreads();
if (lane < count) {
if (reuse_primary_root) {
roots[lane] = static_cast<double>(primary_root);
} else {
double lo = nextafter(shared_poles[lane], CUDART_INF);
double hi;
if (lane + 1 < count) {
hi = nextafter(shared_poles[lane + 1], -CUDART_INF);
} else {
double normz2 = 0.0;
double correction = 0.0;
for (int i = 0; i < count; ++i) {
kahan_add_double(
shared_weights[i] * shared_weights[i], normz2, correction);
}
hi = shared_poles[count - 1] + static_cast<double>(rho[batch]) * normz2;
if (!(hi > lo)) {
hi = nextafter(lo, CUDART_INF);
}
for (int attempt = 0; attempt < 8; ++attempt) {
const double fhi = secular_value_double(
shared_poles,
shared_weights,
count,
static_cast<double>(rho[batch]),
hi);
if (fhi >= 0.0) {
break;
}
hi = shared_poles[count - 1] +
2.0 * (hi - shared_poles[count - 1]);
}
}
const double flo = secular_eval_double(
shared_poles,
shared_weights,
count,
static_cast<double>(rho[batch]),
lo).value;
const double fhi = secular_eval_double(
shared_poles,
shared_weights,
count,
static_cast<double>(rho[batch]),
hi).value;
const bool finite_bracket =
(lo < hi) && isfinite(lo) && isfinite(hi) &&
isfinite(flo) && isfinite(fhi);
const bool sub_ulp_deflation =
finite_bracket && flo >= 0.0 && fhi > 0.0;
if (sub_ulp_deflation) {
// The first FP64 number above the pole is already beyond the
// root: its displacement is unrepresentable in double.
deflated[lane] = 1;
roots[lane] = shared_poles[lane];
} else if (!finite_bracket || !(flo < 0.0) || !(fhi > 0.0)) {
atomicCAS(info + batch, 0, 2);
roots[lane] = CUDART_NAN;
} else {
const double rho_value = static_cast<double>(rho[batch]);
double x = lo + 0.5 * (hi - lo);
const bool origin_at_lower = secular_eval_double(
shared_poles, shared_weights, count, rho_value, x).value > 0.0;
bool converged = false;
int used = 0;
for (int iteration = 1; iteration <= kMaxFp64Iterations; ++iteration) {
used = iteration;
const EvalDouble current = secular_eval_double(
shared_poles, shared_weights, count, rho_value, x);
if (fabs(current.value) <= DBL_EPSILON * current.error_scale) {
converged = true;
break;
}
if (current.value <= 0.0) {
lo = fmax(lo, x);
} else {
hi = fmin(hi, x);
}
const double scale = fmax(1.0, fmax(fabs(lo), fabs(hi)));
if (hi - lo <= 4.0 * DBL_EPSILON * scale ||
nextafter(lo, hi) >= hi) {
x = lo + 0.5 * (hi - lo);
converged = true;
break;
}
double eta = lane + 1 < count
? interior_rational_step_double(
shared_poles,
shared_weights,
rho_value,
lane,
x,
current.value,
current.derivative,
origin_at_lower)
: last_rational_step_double(
shared_poles,
shared_weights,
count,
rho_value,
x,
current.value);
if (!isfinite(eta) || current.value * eta >= 0.0) {
eta = -current.value / current.derivative;
}
double proposed = x + eta;
if (!isfinite(proposed) || !(proposed > lo && proposed < hi) ||
proposed == x || iteration >= 12) {
proposed = lo + 0.5 * (hi - lo);
}
x = proposed;
}
if (!converged || !isfinite(x)) {
atomicCAS(info + batch, 0, 5);
}
roots[lane] = x;
atomicMax(iterations + batch, used);
}
}
}
__syncthreads();
if (lane < count) {
if (deflated[lane]) {
updated_weights[lane] = 0.0;
} else {
const double diagonal_delta = fabs(shared_poles[lane] - roots[lane]);
if (!(diagonal_delta > 0.0) || !isfinite(diagonal_delta)) {
atomicCAS(info + batch, 0, 3);
updated_weights[lane] = CUDART_NAN;
} else {
double log_product = log(diagonal_delta);
for (int root = 0; root < count; ++root) {
// A deflated root equals its matched pole exactly, so its numerator
// and denominator factor cancel for every remaining updated weight.
if (root == lane || deflated[root]) {
continue;
}
const double numerator = fabs(shared_poles[lane] - roots[root]);
const double denominator =
fabs(shared_poles[lane] - shared_poles[root]);
if (!(numerator > 0.0) || !(denominator > 0.0) ||
!isfinite(numerator) || !isfinite(denominator)) {
atomicCAS(info + batch, 0, 3);
} else {
log_product += log(numerator) - log(denominator);
}
}
const double magnitude = exp(0.5 * log_product);
updated_weights[lane] = copysign(magnitude, shared_weights[lane]);
if (!isfinite(updated_weights[lane])) {
atomicCAS(info + batch, 0, 3);
}
}
}
}
__syncthreads();
if (lane < count) {
if (deflated[lane]) {
for (int row = 0; row < count; ++row) {
secular_vectors[matrix_base + row * kMaxMerge + lane] =
row == lane ? 1.0f : 0.0f;
}
} else {
double norm2 = 0.0;
double correction = 0.0;
for (int row = 0; row < count; ++row) {
const double delta = shared_poles[row] - roots[lane];
const double element = updated_weights[row] / delta;
kahan_add_double(element * element, norm2, correction);
}
const double norm = sqrt(norm2);
if (!(norm > 0.0) || !isfinite(norm)) {
atomicCAS(info + batch, 0, 4);
}
for (int row = 0; row < count; ++row) {
const double delta = shared_poles[row] - roots[lane];
const double element = updated_weights[row] / delta / norm;
secular_vectors[matrix_base + row * kMaxMerge + lane] =
static_cast<float>(element);
}
}
values[base + lane] = static_cast<float>(roots[lane]);
}
}
void check_tensor(
const at::Tensor& tensor,
at::ScalarType dtype,
int64_t dimensions,
const char* name) {
TORCH_CHECK(tensor.is_cuda(), name, " must be CUDA");
TORCH_CHECK(tensor.is_contiguous(), name, " must be contiguous");
TORCH_CHECK(tensor.scalar_type() == dtype, name, " has wrong dtype");
TORCH_CHECK(tensor.dim() == dimensions, name, " has wrong rank");
}
} // namespace
pybind11::dict secular_merge_resource_probe() {
cudaFuncAttributes primary{};
cudaFuncAttributes fallback{};
C10_CUDA_CHECK(cudaFuncGetAttributes(&primary, secular_merge_kernel));
C10_CUDA_CHECK(
cudaFuncGetAttributes(&fallback, secular_merge_fp64_control_kernel));
int primary_active_blocks = 0;
int fallback_active_blocks = 0;
C10_CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&primary_active_blocks, secular_merge_kernel, kMaxMerge, 0));
C10_CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&fallback_active_blocks,
secular_merge_fp64_control_kernel,
kMaxMerge,
0));
auto encode = [](const cudaFuncAttributes& attributes, int active_blocks) {
pybind11::dict result;
result["threads"] = kMaxMerge;
result["num_regs"] = attributes.numRegs;
result["static_shared_bytes"] = attributes.sharedSizeBytes;
result["dynamic_shared_bytes"] = 0;
result["local_bytes"] = attributes.localSizeBytes;
result["max_threads_per_block"] = attributes.maxThreadsPerBlock;
result["active_blocks_per_sm"] = active_blocks;
return result;
};
pybind11::dict result;
result["primary"] = encode(primary, primary_active_blocks);
result["fallback"] = encode(fallback, fallback_active_blocks);
result["cluster_required"] = false;
result["tmem_bytes"] = 0;
return result;
}
void secular_merge_run(
const at::Tensor& poles,
const at::Tensor& weights,
const at::Tensor& rho,
const at::Tensor& active_count,
const at::Tensor& values,
const at::Tensor& secular_vectors,
const at::Tensor& info,
const at::Tensor& iterations) {
check_tensor(poles, at::kFloat, 2, "poles");
check_tensor(weights, at::kFloat, 2, "weights");
check_tensor(rho, at::kFloat, 1, "rho");
check_tensor(active_count, at::kInt, 1, "active_count");
check_tensor(values, at::kFloat, 2, "values");
check_tensor(secular_vectors, at::kFloat, 3, "secular_vectors");
check_tensor(info, at::kInt, 1, "info");
check_tensor(iterations, at::kInt, 1, "iterations");
const int64_t batch = poles.size(0);
TORCH_CHECK(poles.size(1) == kMaxMerge, "poles must have width 256");
TORCH_CHECK(weights.sizes() == poles.sizes(), "weights shape mismatch");
TORCH_CHECK(values.sizes() == poles.sizes(), "values shape mismatch");
TORCH_CHECK(rho.size(0) == batch, "rho shape mismatch");
TORCH_CHECK(active_count.size(0) == batch, "active_count shape mismatch");
TORCH_CHECK(
secular_vectors.size(0) == batch &&
secular_vectors.size(1) == kMaxMerge &&
secular_vectors.size(2) == kMaxMerge,
"secular_vectors shape mismatch");
TORCH_CHECK(info.size(0) == batch, "info shape mismatch");
TORCH_CHECK(iterations.size(0) == batch, "iterations shape mismatch");
TORCH_CHECK(batch > 0 && batch <= std::numeric_limits<int>::max(),
"unsupported batch size");
secular_merge_kernel<<<static_cast<int>(batch), kMaxMerge, 0>>>(
poles.data_ptr<float>(),
weights.data_ptr<float>(),
rho.data_ptr<float>(),
active_count.data_ptr<int32_t>(),
values.data_ptr<float>(),
secular_vectors.data_ptr<float>(),
info.data_ptr<int32_t>(),
iterations.data_ptr<int32_t>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
void secular_merge_fp64_control_run(
const at::Tensor& poles,
const at::Tensor& weights,
const at::Tensor& rho,
const at::Tensor& active_count,
const at::Tensor& values,
const at::Tensor& secular_vectors,
const at::Tensor& info,
const at::Tensor& iterations) {
check_tensor(poles, at::kFloat, 2, "poles");
check_tensor(weights, at::kFloat, 2, "weights");
check_tensor(rho, at::kFloat, 1, "rho");
check_tensor(active_count, at::kInt, 1, "active_count");
check_tensor(values, at::kFloat, 2, "values");
check_tensor(secular_vectors, at::kFloat, 3, "secular_vectors");
check_tensor(info, at::kInt, 1, "info");
check_tensor(iterations, at::kInt, 1, "iterations");
const int64_t batch = poles.size(0);
TORCH_CHECK(poles.size(1) == kMaxMerge, "poles must have width 256");
TORCH_CHECK(weights.sizes() == poles.sizes(), "weights shape mismatch");
TORCH_CHECK(values.sizes() == poles.sizes(), "values shape mismatch");
TORCH_CHECK(rho.size(0) == batch, "rho shape mismatch");
TORCH_CHECK(active_count.size(0) == batch, "active_count shape mismatch");
TORCH_CHECK(
secular_vectors.size(0) == batch &&
secular_vectors.size(1) == kMaxMerge &&
secular_vectors.size(2) == kMaxMerge,
"secular_vectors shape mismatch");
TORCH_CHECK(info.size(0) == batch, "info shape mismatch");
TORCH_CHECK(iterations.size(0) == batch, "iterations shape mismatch");
TORCH_CHECK(batch > 0 && batch <= std::numeric_limits<int>::max(),
"unsupported batch size");
secular_merge_fp64_control_kernel<<<
static_cast<int>(batch), kMaxMerge, 0>>>(
poles.data_ptr<float>(),
weights.data_ptr<float>(),
rho.data_ptr<float>(),
active_count.data_ptr<int32_t>(),
values.data_ptr<float>(),
secular_vectors.data_ptr<float>(),
info.data_ptr<int32_t>(),
iterations.data_ptr<int32_t>(),
false);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
void secular_merge_hybrid_run(
const at::Tensor& poles,
const at::Tensor& weights,
const at::Tensor& rho,
const at::Tensor& active_count,
const at::Tensor& values,
const at::Tensor& secular_vectors,
const at::Tensor& info,
const at::Tensor& iterations) {
// Reuse the exact strict-FP32 entrypoint and its full shape checks.
secular_merge_run(
poles,
weights,
rho,
active_count,
values,
secular_vectors,
info,
iterations);
const int64_t batch = poles.size(0);
secular_merge_fp64_control_kernel<<<
static_cast<int>(batch), kMaxMerge, 0>>>(
poles.data_ptr<float>(),
weights.data_ptr<float>(),
rho.data_ptr<float>(),
active_count.data_ptr<int32_t>(),
values.data_ptr<float>(),
secular_vectors.data_ptr<float>(),
info.data_ptr<int32_t>(),
iterations.data_ptr<int32_t>(),
true);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
} // namespace lane_e256
namespace lane_e512 {
namespace {
constexpr int kMaxMerge = 512;
constexpr int kMaxIterations = 30;
constexpr int kMaxFp64Iterations = 80;
__device__ __forceinline__ void kahan_add(
float value, float& total, float& compensation) {
const float adjusted = value - compensation;
const float updated = total + adjusted;
compensation = (updated - total) - adjusted;
total = updated;
}
__device__ __forceinline__ float secular_value(
const float* poles,
const float* weights,
int count,
float rho,
float x) {
// Accumulate the two signs separately. Close-pole secular evaluations
// otherwise lose the smaller signed partial sum before the final add.
float negative = 0.0f;
float positive = 0.0f;
float cnegative = 0.0f;
float cpositive = 0.0f;
for (int i = 0; i < count; ++i) {
const float term = rho * weights[i] * weights[i] / (poles[i] - x);
if (term < 0.0f) {
kahan_add(term, negative, cnegative);
} else {
kahan_add(term, positive, cpositive);
}
}
return 1.0f + (negative + positive);
}
// Status-2 is often a representation boundary rather than a missing root:
// the mathematical root lies between a pole and its first inward FP32 value.
// Evaluate only that endpoint-sign predicate in FP64. This is deliberately
// not a second FP64 root solve. A valid predicate authorizes a caller-visible
// FP32 endpoint clamp; an invalid predicate remains status 2 and is handled by
// the existing fail-closed selective-root FP64 fallback launch.
__device__ __forceinline__ double endpoint_secular_value_double(
const float* poles,
const float* weights,
int count,
float rho,
double x) {
double negative = 0.0;
double positive = 0.0;
double cnegative = 0.0;
double cpositive = 0.0;
for (int index = 0; index < count; ++index) {
const double pole = static_cast<double>(poles[index]);
const double weight = static_cast<double>(weights[index]);
const double term =
static_cast<double>(rho) * weight * weight / (pole - x);
if (term < 0.0) {
const double adjusted = term - cnegative;
const double updated = negative + adjusted;
cnegative = (updated - negative) - adjusted;
negative = updated;
} else {
const double adjusted = term - cpositive;
const double updated = positive + adjusted;
cpositive = (updated - positive) - adjusted;
positive = updated;
}
}
return 1.0 + negative + positive;
}
struct EvalFloat {
float value;
float derivative;
float error_scale;
};
__device__ __forceinline__ EvalFloat secular_eval_float(
const float* poles,
const float* weights,
int count,
float rho,
float x) {
float negative = 0.0f;
float positive = 0.0f;
float derivative = 0.0f;
float magnitude = 0.0f;
float cnegative = 0.0f;
float cpositive = 0.0f;
float cderivative = 0.0f;
float cmagnitude = 0.0f;
for (int index = 0; index < count; ++index) {
const float delta = poles[index] - x;
const float weight2 = weights[index] * weights[index];
const float term = rho * weight2 / delta;
if (term < 0.0f) {
kahan_add(term, negative, cnegative);
} else {
kahan_add(term, positive, cpositive);
}
kahan_add(rho * weight2 / (delta * delta), derivative, cderivative);
kahan_add(fabsf(term), magnitude, cmagnitude);
}
return {
1.0f + negative + positive,
derivative,
8.0f * (1.0f + magnitude + fabsf(x) * derivative)};
}
__device__ __forceinline__ float interior_rational_step_float(
const float* poles,
const float* weights,
float rho,
int index,
float x,
float value,
float derivative,
bool origin_at_lower) {
const float delta_i = poles[index] - x;
const float delta_ip1 = poles[index + 1] - x;
const float gap = poles[index + 1] - poles[index];
float c;
if (origin_at_lower) {
const float ratio = weights[index] / delta_i;
c = value - delta_ip1 * derivative + gap * rho * ratio * ratio;
} else {
const float ratio = weights[index + 1] / delta_ip1;
c = value - delta_i * derivative - gap * rho * ratio * ratio;
}
const float a =
(delta_i + delta_ip1) * value - delta_i * delta_ip1 * derivative;
const float b = delta_i * delta_ip1 * value;
if (c == 0.0f) {
return a == 0.0f ? CUDART_NAN_F : b / a;
}
const float discriminant = fmaxf(0.0f, a * a - 4.0f * b * c);
const float root = sqrtf(discriminant);
if (a <= 0.0f) {
return (a - root) / (2.0f * c);
}
const float denominator = a + root;
return denominator == 0.0f ? CUDART_NAN_F : 2.0f * b / denominator;
}
__device__ __forceinline__ float last_rational_step_float(
const float* poles,
const float* weights,
int count,
float rho,
float x,
float value) {
const float delta0 = poles[count - 2] - x;
const float delta1 = poles[count - 1] - x;
float dpsi = 0.0f;
float correction = 0.0f;
for (int index = 0; index < count - 1; ++index) {
const float ratio = weights[index] / (poles[index] - x);
kahan_add(rho * ratio * ratio, dpsi, correction);
}
const float last_ratio = weights[count - 1] / delta1;
const float dphi = rho * last_ratio * last_ratio;
const float c = fabsf(value - delta0 * dpsi - delta1 * dphi);
const float a =
(delta0 + delta1) * value - delta0 * delta1 * (dpsi + dphi);
const float b = delta0 * delta1 * value;
if (c == 0.0f) {
return CUDART_NAN_F;
}
const float discriminant = fmaxf(0.0f, a * a - 4.0f * b * c);
const float root = sqrtf(discriminant);
if (a >= 0.0f) {
return (a + root) / (2.0f * c);
}
const float denominator = a - root;
return denominator == 0.0f ? CUDART_NAN_F : 2.0f * b / denominator;
}
__global__ void secular_merge_kernel(
const float* __restrict__ poles,
const float* __restrict__ weights,
const float* __restrict__ rho,
const int32_t* __restrict__ active_count,
float* __restrict__ values,
float* __restrict__ secular_vectors,
int32_t* __restrict__ info,
int32_t* __restrict__ iterations) {
const int batch = static_cast<int>(blockIdx.x);
const int lane = static_cast<int>(threadIdx.x);
const int base = batch * kMaxMerge;
const int matrix_base = batch * kMaxMerge * kMaxMerge;
const int count = active_count[batch];
__shared__ float shared_poles[kMaxMerge];
__shared__ float shared_weights[kMaxMerge];
__shared__ float roots[kMaxMerge];
__shared__ float updated_weights[kMaxMerge];
if (lane == 0) {
info[batch] = (count > 0 && count <= kMaxMerge) ? 0 : 1;
iterations[batch] = 0;
}
if (lane < kMaxMerge) {
shared_poles[lane] = poles[base + lane];
shared_weights[lane] = weights[base + lane];
values[base + lane] = 0.0f;
for (int row = 0; row < kMaxMerge; ++row) {
secular_vectors[matrix_base + row * kMaxMerge + lane] = 0.0f;
}
}
__syncthreads();
if (lane < count) {
const float positive_inf = CUDART_INF_F;
const float negative_inf = -CUDART_INF_F;
float lo = nextafterf(shared_poles[lane], positive_inf);
float hi;
if (lane + 1 < count) {
hi = nextafterf(shared_poles[lane + 1], negative_inf);
} else {
float normz2 = 0.0f;
float correction = 0.0f;
for (int i = 0; i < count; ++i) {
kahan_add(shared_weights[i] * shared_weights[i], normz2, correction);
}
hi = shared_poles[count - 1] + rho[batch] * normz2;
if (!(hi > lo)) {
hi = nextafterf(lo, positive_inf);
}
for (int attempt = 0; attempt < 8; ++attempt) {
const float fhi = secular_value(
shared_poles, shared_weights, count, rho[batch], hi);
if (fhi >= 0.0f) {
break;
}
hi = shared_poles[count - 1] +
2.0f * (hi - shared_poles[count - 1]);
}
}
const float flo = secular_eval_float(
shared_poles, shared_weights, count, rho[batch], lo).value;
const float fhi = secular_eval_float(
shared_poles, shared_weights, count, rho[batch], hi).value;
if (!(lo < hi) || !(flo <= 0.0f) || !(fhi >= 0.0f) ||
!isfinite(lo) || !isfinite(hi)) {
bool endpoint_sign_rescued = false;
if (lo < hi && isfinite(lo) && isfinite(hi) &&
(flo > 0.0f || fhi < 0.0f)) {
const double double_lo = nextafter(
static_cast<double>(shared_poles[lane]), CUDART_INF);
double double_hi;
if (lane + 1 < count) {
double_hi = nextafter(
static_cast<double>(shared_poles[lane + 1]), -CUDART_INF);
} else {
double normz2 = 0.0;
double correction = 0.0;
for (int index = 0; index < count; ++index) {
const double weight =
static_cast<double>(shared_weights[index]);
const double value = weight * weight - correction;
const double updated = normz2 + value;
correction = (updated - normz2) - value;
normz2 = updated;
}
double_hi = static_cast<double>(shared_poles[count - 1]) +
static_cast<double>(rho[batch]) * normz2;
if (!(double_hi > double_lo)) {
double_hi = nextafter(double_lo, CUDART_INF);
}
for (int attempt = 0; attempt < 8; ++attempt) {
if (endpoint_secular_value_double(
shared_poles,
shared_weights,
count,
rho[batch],
double_hi) >= 0.0) {
break;
}
double_hi = static_cast<double>(shared_poles[count - 1]) +
2.0 *
(double_hi -
static_cast<double>(shared_poles[count - 1]));
}
}
const double double_flo = endpoint_secular_value_double(
shared_poles,
shared_weights,
count,
rho[batch],
double_lo);
const double double_fhi = endpoint_secular_value_double(
shared_poles,
shared_weights,
count,
rho[batch],
double_hi);
endpoint_sign_rescued =
isfinite(double_lo) && isfinite(double_hi) &&
isfinite(double_flo) && isfinite(double_fhi) &&
double_lo < double_hi && double_flo <= 0.0 &&
double_fhi >= 0.0;
}
if (endpoint_sign_rescued) {
// The exact root is below/above the first representable interior FP32
// endpoint. Clamp at that endpoint and retain the existing FP32
// SLAED3 displacement reconstruction below.
roots[lane] = flo > 0.0f ? lo : hi;
} else {
atomicCAS(info + batch, 0, 2);
roots[lane] = CUDART_NAN_F;
}
} else {
float x = lo + 0.5f * (hi - lo);
const bool origin_at_lower = secular_eval_float(
shared_poles, shared_weights, count, rho[batch], x).value > 0.0f;
bool converged = false;
int used = 0;
for (int iteration = 1; iteration <= kMaxIterations; ++iteration) {
used = iteration;
const EvalFloat current = secular_eval_float(
shared_poles, shared_weights, count, rho[batch], x);
if (fabsf(current.value) <= FLT_EPSILON * current.error_scale) {
converged = true;
break;
}
if (current.value <= 0.0f) {
lo = fmaxf(lo, x);
} else {
hi = fminf(hi, x);
}
const float scale = fmaxf(1.0f, fmaxf(fabsf(lo), fabsf(hi)));
if (hi - lo <= 2.0f * FLT_EPSILON * scale || lo == hi) {
x = lo + 0.5f * (hi - lo);
converged = true;
break;
}
float eta = lane + 1 < count
? interior_rational_step_float(
shared_poles,
shared_weights,
rho[batch],
lane,
x,
current.value,
current.derivative,
origin_at_lower)
: last_rational_step_float(
shared_poles,
shared_weights,
count,
rho[batch],
x,
current.value);
if (!isfinite(eta) || current.value * eta >= 0.0f) {
eta = -current.value / current.derivative;
}
float proposed = x + eta;
if (!isfinite(proposed) || !(proposed > lo && proposed < hi) ||
proposed == x || iteration >= 12) {
proposed = lo + 0.5f * (hi - lo);
}
x = proposed;
}
if (!converged || !isfinite(x)) {
atomicCAS(info + batch, 0, 5);
}
roots[lane] = x;
atomicMax(iterations + batch, used);
}
}
__syncthreads();
// SLAED3 reconstructs the rank-one weights from all root displacements.
// Evaluate the product in the log domain: this is algebraically equivalent
// for the positive rank-one problem and avoids a canary-only overflow.
if (lane < count) {
float diagonal_delta = fabsf(shared_poles[lane] - roots[lane]);
if (!(diagonal_delta > 0.0f) || !isfinite(diagonal_delta)) {
atomicCAS(info + batch, 0, 3);
updated_weights[lane] = CUDART_NAN_F;
} else {
float log_product = logf(diagonal_delta);
for (int root = 0; root < count; ++root) {
if (root == lane) {
continue;
}
const float numerator = fabsf(shared_poles[lane] - roots[root]);
const float denominator =
fabsf(shared_poles[lane] - shared_poles[root]);
if (!(numerator > 0.0f) || !(denominator > 0.0f)) {
atomicCAS(info + batch, 0, 3);
} else {
log_product += logf(numerator) - logf(denominator);
}
}
const float magnitude = expf(0.5f * log_product);
updated_weights[lane] = copysignf(magnitude, shared_weights[lane]);
if (!isfinite(updated_weights[lane])) {
atomicCAS(info + batch, 0, 3);
}
}
}
__syncthreads();
if (lane < count) {
float norm2 = 0.0f;
float correction = 0.0f;
for (int row = 0; row < count; ++row) {
const float delta = shared_poles[row] - roots[lane];
const float element = updated_weights[row] / delta;
kahan_add(element * element, norm2, correction);
}
const float norm = sqrtf(norm2);
if (!(norm > 0.0f) || !isfinite(norm)) {
atomicCAS(info + batch, 0, 4);
}
for (int row = 0; row < count; ++row) {
const float delta = shared_poles[row] - roots[lane];
secular_vectors[matrix_base + row * kMaxMerge + lane] =
updated_weights[row] / delta / norm;
}
values[base + lane] = roots[lane];
}
}
__device__ __forceinline__ void kahan_add_double(
double value, double& total, double& compensation) {
const double adjusted = value - compensation;
const double updated = total + adjusted;
compensation = (updated - total) - adjusted;
total = updated;
}
__device__ __forceinline__ double secular_value_double(
const double* poles,
const double* weights,
int count,
double rho,
double x) {
double negative = 0.0;
double positive = 0.0;
double cnegative = 0.0;
double cpositive = 0.0;
for (int i = 0; i < count; ++i) {
const double term = rho * weights[i] * weights[i] / (poles[i] - x);
if (term < 0.0) {
kahan_add_double(term, negative, cnegative);
} else {
kahan_add_double(term, positive, cpositive);
}
}
return 1.0 + (negative + positive);
}
struct EvalDouble {
double value;
double derivative;
double error_scale;
};
__device__ __forceinline__ EvalDouble secular_eval_double(
const double* poles,
const double* weights,
int count,
double rho,
double x) {
double negative = 0.0;
double positive = 0.0;
double derivative = 0.0;
double magnitude = 0.0;
double cnegative = 0.0;
double cpositive = 0.0;
double cderivative = 0.0;
double cmagnitude = 0.0;
for (int index = 0; index < count; ++index) {
const double delta = poles[index] - x;
const double weight2 = weights[index] * weights[index];
const double term = rho * weight2 / delta;
if (term < 0.0) {
kahan_add_double(term, negative, cnegative);
} else {
kahan_add_double(term, positive, cpositive);
}
kahan_add_double(
rho * weight2 / (delta * delta), derivative, cderivative);
kahan_add_double(fabs(term), magnitude, cmagnitude);
}
return {
1.0 + negative + positive,
derivative,
8.0 * (1.0 + magnitude + fabs(x) * derivative)};
}
__device__ __forceinline__ double interior_rational_step_double(
const double* poles,
const double* weights,
double rho,
int index,
double x,
double value,
double derivative,
bool origin_at_lower) {
const double delta_i = poles[index] - x;
const double delta_ip1 = poles[index + 1] - x;
const double gap = poles[index + 1] - poles[index];
double c;
if (origin_at_lower) {
const double ratio = weights[index] / delta_i;
c = value - delta_ip1 * derivative + gap * rho * ratio * ratio;
} else {
const double ratio = weights[index + 1] / delta_ip1;
c = value - delta_i * derivative - gap * rho * ratio * ratio;
}
const double a =
(delta_i + delta_ip1) * value - delta_i * delta_ip1 * derivative;
const double b = delta_i * delta_ip1 * value;
if (c == 0.0) {
return a == 0.0 ? CUDART_NAN : b / a;
}
const double root = sqrt(fmax(0.0, a * a - 4.0 * b * c));
if (a <= 0.0) {
return (a - root) / (2.0 * c);
}
const double denominator = a + root;
return denominator == 0.0 ? CUDART_NAN : 2.0 * b / denominator;
}
__device__ __forceinline__ double last_rational_step_double(
const double* poles,
const double* weights,
int count,
double rho,
double x,
double value) {
const double delta0 = poles[count - 2] - x;
const double delta1 = poles[count - 1] - x;
double dpsi = 0.0;
double correction = 0.0;
for (int index = 0; index < count - 1; ++index) {
const double ratio = weights[index] / (poles[index] - x);
kahan_add_double(rho * ratio * ratio, dpsi, correction);
}
const double last_ratio = weights[count - 1] / delta1;
const double dphi = rho * last_ratio * last_ratio;
const double c = fabs(value - delta0 * dpsi - delta1 * dphi);
const double a =
(delta0 + delta1) * value - delta0 * delta1 * (dpsi + dphi);
const double b = delta0 * delta1 * value;
if (c == 0.0) {
return CUDART_NAN;
}
const double root = sqrt(fmax(0.0, a * a - 4.0 * b * c));
if (a >= 0.0) {
return (a + root) / (2.0 * c);
}
const double denominator = a - root;
return denominator == 0.0 ? CUDART_NAN : 2.0 * b / denominator;
}
__global__ void secular_merge_fp64_control_kernel(
const float* __restrict__ poles,
const float* __restrict__ weights,
const float* __restrict__ rho,
const int32_t* __restrict__ active_count,
float* __restrict__ values,
float* __restrict__ secular_vectors,
int32_t* __restrict__ info,
int32_t* __restrict__ iterations,
bool fallback_only) {
const int batch = static_cast<int>(blockIdx.x);
const int lane = static_cast<int>(threadIdx.x);
const int base = batch * kMaxMerge;
const int matrix_base = batch * kMaxMerge * kMaxMerge;
const int count = active_count[batch];
__shared__ double shared_poles[kMaxMerge];
__shared__ double shared_weights[kMaxMerge];
__shared__ double roots[kMaxMerge];
__shared__ double updated_weights[kMaxMerge];
__shared__ int32_t deflated[kMaxMerge];
__shared__ int32_t prior_info;
// Snapshot the preceding strict-FP32 status before lane zero clears it.
// The barrier makes both the early return and repair route uniform across
// the CTA's two warps.
if (lane == 0) {
prior_info = fallback_only ? info[batch] : 0;
}
__syncthreads();
if (fallback_only && prior_info == 0) {
return;
}
// Status 2 is an endpoint-sign failure. The primary launch leaves finite
// values for roots whose own brackets succeeded and NaN for failed roots.
const float primary_root = lane < count ? values[base + lane] : CUDART_NAN_F;
const bool reuse_primary_root =
fallback_only && prior_info == 2 && isfinite(primary_root);
if (lane == 0) {
info[batch] = (count > 0 && count <= kMaxMerge) ? 0 : 1;
iterations[batch] = 0;
}
shared_poles[lane] = static_cast<double>(poles[base + lane]);
shared_weights[lane] = static_cast<double>(weights[base + lane]);
deflated[lane] = 0;
if (!reuse_primary_root) {
values[base + lane] = 0.0f;
}
for (int row = 0; row < kMaxMerge; ++row) {
secular_vectors[matrix_base + row * kMaxMerge + lane] = 0.0f;
}
__syncthreads();
if (lane < count) {
if (reuse_primary_root) {
roots[lane] = static_cast<double>(primary_root);
} else {
double lo = nextafter(shared_poles[lane], CUDART_INF);
double hi;
if (lane + 1 < count) {
hi = nextafter(shared_poles[lane + 1], -CUDART_INF);
} else {
double normz2 = 0.0;
double correction = 0.0;
for (int i = 0; i < count; ++i) {
kahan_add_double(
shared_weights[i] * shared_weights[i], normz2, correction);
}
hi = shared_poles[count - 1] + static_cast<double>(rho[batch]) * normz2;
if (!(hi > lo)) {
hi = nextafter(lo, CUDART_INF);
}
for (int attempt = 0; attempt < 8; ++attempt) {
const double fhi = secular_value_double(
shared_poles,
shared_weights,
count,
static_cast<double>(rho[batch]),
hi);
if (fhi >= 0.0) {
break;
}
hi = shared_poles[count - 1] +
2.0 * (hi - shared_poles[count - 1]);
}
}
const double flo = secular_eval_double(
shared_poles,
shared_weights,
count,
static_cast<double>(rho[batch]),
lo).value;
const double fhi = secular_eval_double(
shared_poles,
shared_weights,
count,
static_cast<double>(rho[batch]),
hi).value;
const bool finite_bracket =
(lo < hi) && isfinite(lo) && isfinite(hi) &&
isfinite(flo) && isfinite(fhi);
const bool sub_ulp_deflation =
finite_bracket && flo >= 0.0 && fhi > 0.0;
if (sub_ulp_deflation) {
// The first FP64 number above the pole is already beyond the
// root: its displacement is unrepresentable in double.
deflated[lane] = 1;
roots[lane] = shared_poles[lane];
} else if (!finite_bracket || !(flo < 0.0) || !(fhi > 0.0)) {
atomicCAS(info + batch, 0, 2);
roots[lane] = CUDART_NAN;
} else {
const double rho_value = static_cast<double>(rho[batch]);
double x = lo + 0.5 * (hi - lo);
const bool origin_at_lower = secular_eval_double(
shared_poles, shared_weights, count, rho_value, x).value > 0.0;
bool converged = false;
int used = 0;
for (int iteration = 1; iteration <= kMaxFp64Iterations; ++iteration) {
used = iteration;
const EvalDouble current = secular_eval_double(
shared_poles, shared_weights, count, rho_value, x);
if (fabs(current.value) <= DBL_EPSILON * current.error_scale) {
converged = true;
break;
}
if (current.value <= 0.0) {
lo = fmax(lo, x);
} else {
hi = fmin(hi, x);
}
const double scale = fmax(1.0, fmax(fabs(lo), fabs(hi)));
if (hi - lo <= 4.0 * DBL_EPSILON * scale ||
nextafter(lo, hi) >= hi) {
x = lo + 0.5 * (hi - lo);
converged = true;
break;
}
double eta = lane + 1 < count
? interior_rational_step_double(
shared_poles,
shared_weights,
rho_value,
lane,
x,
current.value,
current.derivative,
origin_at_lower)
: last_rational_step_double(
shared_poles,
shared_weights,
count,
rho_value,
x,
current.value);
if (!isfinite(eta) || current.value * eta >= 0.0) {
eta = -current.value / current.derivative;
}
double proposed = x + eta;
if (!isfinite(proposed) || !(proposed > lo && proposed < hi) ||
proposed == x || iteration >= 12) {
proposed = lo + 0.5 * (hi - lo);
}
x = proposed;
}
if (!converged || !isfinite(x)) {
atomicCAS(info + batch, 0, 5);
}
roots[lane] = x;
atomicMax(iterations + batch, used);
}
}
}
__syncthreads();
if (lane < count) {
if (deflated[lane]) {
updated_weights[lane] = 0.0;
} else {
const double diagonal_delta = fabs(shared_poles[lane] - roots[lane]);
if (!(diagonal_delta > 0.0) || !isfinite(diagonal_delta)) {
atomicCAS(info + batch, 0, 3);
updated_weights[lane] = CUDART_NAN;
} else {
double log_product = log(diagonal_delta);
for (int root = 0; root < count; ++root) {
// A deflated root equals its matched pole exactly, so its numerator
// and denominator factor cancel for every remaining updated weight.
if (root == lane || deflated[root]) {
continue;
}
const double numerator = fabs(shared_poles[lane] - roots[root]);
const double denominator =
fabs(shared_poles[lane] - shared_poles[root]);
if (!(numerator > 0.0) || !(denominator > 0.0) ||
!isfinite(numerator) || !isfinite(denominator)) {
atomicCAS(info + batch, 0, 3);
} else {
log_product += log(numerator) - log(denominator);
}
}
const double magnitude = exp(0.5 * log_product);
updated_weights[lane] = copysign(magnitude, shared_weights[lane]);
if (!isfinite(updated_weights[lane])) {
atomicCAS(info + batch, 0, 3);
}
}
}
}
__syncthreads();
if (lane < count) {
if (deflated[lane]) {
for (int row = 0; row < count; ++row) {
secular_vectors[matrix_base + row * kMaxMerge + lane] =
row == lane ? 1.0f : 0.0f;
}
} else {
double norm2 = 0.0;
double correction = 0.0;
for (int row = 0; row < count; ++row) {
const double delta = shared_poles[row] - roots[lane];
const double element = updated_weights[row] / delta;
kahan_add_double(element * element, norm2, correction);
}
const double norm = sqrt(norm2);
if (!(norm > 0.0) || !isfinite(norm)) {
atomicCAS(info + batch, 0, 4);
}
for (int row = 0; row < count; ++row) {
const double delta = shared_poles[row] - roots[lane];
const double element = updated_weights[row] / delta / norm;
secular_vectors[matrix_base + row * kMaxMerge + lane] =
static_cast<float>(element);
}
}
values[base + lane] = static_cast<float>(roots[lane]);
}
}
void check_tensor(
const at::Tensor& tensor,
at::ScalarType dtype,
int64_t dimensions,
const char* name) {
TORCH_CHECK(tensor.is_cuda(), name, " must be CUDA");
TORCH_CHECK(tensor.is_contiguous(), name, " must be contiguous");
TORCH_CHECK(tensor.scalar_type() == dtype, name, " has wrong dtype");
TORCH_CHECK(tensor.dim() == dimensions, name, " has wrong rank");
}
} // namespace
pybind11::dict secular_merge_resource_probe() {
cudaFuncAttributes primary{};
cudaFuncAttributes fallback{};
C10_CUDA_CHECK(cudaFuncGetAttributes(&primary, secular_merge_kernel));
C10_CUDA_CHECK(
cudaFuncGetAttributes(&fallback, secular_merge_fp64_control_kernel));
int primary_active_blocks = 0;
int fallback_active_blocks = 0;
C10_CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&primary_active_blocks, secular_merge_kernel, kMaxMerge, 0));
C10_CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&fallback_active_blocks,
secular_merge_fp64_control_kernel,
kMaxMerge,
0));
auto encode = [](const cudaFuncAttributes& attributes, int active_blocks) {
pybind11::dict result;
result["threads"] = kMaxMerge;
result["num_regs"] = attributes.numRegs;
result["static_shared_bytes"] = attributes.sharedSizeBytes;
result["dynamic_shared_bytes"] = 0;
result["local_bytes"] = attributes.localSizeBytes;
result["max_threads_per_block"] = attributes.maxThreadsPerBlock;
result["active_blocks_per_sm"] = active_blocks;
return result;
};
pybind11::dict result;
result["primary"] = encode(primary, primary_active_blocks);
result["fallback"] = encode(fallback, fallback_active_blocks);
result["cluster_required"] = false;
result["tmem_bytes"] = 0;
return result;
}
void secular_merge_run(
const at::Tensor& poles,
const at::Tensor& weights,
const at::Tensor& rho,
const at::Tensor& active_count,
const at::Tensor& values,
const at::Tensor& secular_vectors,
const at::Tensor& info,
const at::Tensor& iterations) {
check_tensor(poles, at::kFloat, 2, "poles");
check_tensor(weights, at::kFloat, 2, "weights");
check_tensor(rho, at::kFloat, 1, "rho");
check_tensor(active_count, at::kInt, 1, "active_count");
check_tensor(values, at::kFloat, 2, "values");
check_tensor(secular_vectors, at::kFloat, 3, "secular_vectors");
check_tensor(info, at::kInt, 1, "info");
check_tensor(iterations, at::kInt, 1, "iterations");
const int64_t batch = poles.size(0);
TORCH_CHECK(poles.size(1) == kMaxMerge, "poles must have width 512");
TORCH_CHECK(weights.sizes() == poles.sizes(), "weights shape mismatch");
TORCH_CHECK(values.sizes() == poles.sizes(), "values shape mismatch");
TORCH_CHECK(rho.size(0) == batch, "rho shape mismatch");
TORCH_CHECK(active_count.size(0) == batch, "active_count shape mismatch");
TORCH_CHECK(
secular_vectors.size(0) == batch &&
secular_vectors.size(1) == kMaxMerge &&
secular_vectors.size(2) == kMaxMerge,
"secular_vectors shape mismatch");
TORCH_CHECK(info.size(0) == batch, "info shape mismatch");
TORCH_CHECK(iterations.size(0) == batch, "iterations shape mismatch");
TORCH_CHECK(batch > 0 && batch <= std::numeric_limits<int>::max(),
"unsupported batch size");
secular_merge_kernel<<<static_cast<int>(batch), kMaxMerge, 0>>>(
poles.data_ptr<float>(),
weights.data_ptr<float>(),
rho.data_ptr<float>(),
active_count.data_ptr<int32_t>(),
values.data_ptr<float>(),
secular_vectors.data_ptr<float>(),
info.data_ptr<int32_t>(),
iterations.data_ptr<int32_t>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
void secular_merge_fp64_control_run(
const at::Tensor& poles,
const at::Tensor& weights,
const at::Tensor& rho,
const at::Tensor& active_count,
const at::Tensor& values,
const at::Tensor& secular_vectors,
const at::Tensor& info,
const at::Tensor& iterations) {
check_tensor(poles, at::kFloat, 2, "poles");
check_tensor(weights, at::kFloat, 2, "weights");
check_tensor(rho, at::kFloat, 1, "rho");
check_tensor(active_count, at::kInt, 1, "active_count");
check_tensor(values, at::kFloat, 2, "values");
check_tensor(secular_vectors, at::kFloat, 3, "secular_vectors");
check_tensor(info, at::kInt, 1, "info");
check_tensor(iterations, at::kInt, 1, "iterations");
const int64_t batch = poles.size(0);
TORCH_CHECK(poles.size(1) == kMaxMerge, "poles must have width 512");
TORCH_CHECK(weights.sizes() == poles.sizes(), "weights shape mismatch");
TORCH_CHECK(values.sizes() == poles.sizes(), "values shape mismatch");
TORCH_CHECK(rho.size(0) == batch, "rho shape mismatch");
TORCH_CHECK(active_count.size(0) == batch, "active_count shape mismatch");
TORCH_CHECK(
secular_vectors.size(0) == batch &&
secular_vectors.size(1) == kMaxMerge &&
secular_vectors.size(2) == kMaxMerge,
"secular_vectors shape mismatch");
TORCH_CHECK(info.size(0) == batch, "info shape mismatch");
TORCH_CHECK(iterations.size(0) == batch, "iterations shape mismatch");
TORCH_CHECK(batch > 0 && batch <= std::numeric_limits<int>::max(),
"unsupported batch size");
secular_merge_fp64_control_kernel<<<
static_cast<int>(batch), kMaxMerge, 0>>>(
poles.data_ptr<float>(),
weights.data_ptr<float>(),
rho.data_ptr<float>(),
active_count.data_ptr<int32_t>(),
values.data_ptr<float>(),
secular_vectors.data_ptr<float>(),
info.data_ptr<int32_t>(),
iterations.data_ptr<int32_t>(),
false);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
void secular_merge_hybrid_run(
const at::Tensor& poles,
const at::Tensor& weights,
const at::Tensor& rho,
const at::Tensor& active_count,
const at::Tensor& values,
const at::Tensor& secular_vectors,
const at::Tensor& info,
const at::Tensor& iterations) {
// Reuse the exact strict-FP32 entrypoint and its full shape checks.
secular_merge_run(
poles,
weights,
rho,
active_count,
values,
secular_vectors,
info,
iterations);
const int64_t batch = poles.size(0);
secular_merge_fp64_control_kernel<<<
static_cast<int>(batch), kMaxMerge, 0>>>(
poles.data_ptr<float>(),
weights.data_ptr<float>(),
rho.data_ptr<float>(),
active_count.data_ptr<int32_t>(),
values.data_ptr<float>(),
secular_vectors.data_ptr<float>(),
info.data_ptr<int32_t>(),
iterations.data_ptr<int32_t>(),
true);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
} // namespace lane_e512
namespace lane_f {
namespace {
constexpr int kN = 512;
constexpr int kKD = 8;
constexpr int kReflectors = 16512;
constexpr int kRecordWidth = kKD;
constexpr int kPanelCols = 32;
constexpr int kSweepPrefixEntries = kN - 1;
constexpr int kMaxSegments = 64; // ceil((kN-1-1)/kKD) at sweep 1
constexpr size_t kPanelBytes =
static_cast<size_t>(kN) * kPanelCols * sizeof(float);
__device__ __constant__ uint16_t kSweepPrefix[kSweepPrefixEntries];
static_assert(kPanelBytes == 65536, "panel contract changed");
constexpr size_t dynamic_shared_bytes(bool stage_hh) {
return kPanelBytes
+ (stage_hh
? 2ull * kMaxSegments * kRecordWidth * sizeof(float)
: 0ull);
}
static_assert(dynamic_shared_bytes(true) == 69632,
"staged shared contract changed");
// A window's <= 8 resident rows as NAMED scalars: both an indexed local
// array float[kWindows][kKD] and named float[8] arrays were demoted to a
// 128-byte stack frame by ptxas at four windows (Spark sm_121 probe);
// the fully scalarized struct promotes to registers in every config
// (t512: 64 regs, t1024: 40 regs, zero stack).
struct Win {
float r0, r1, r2, r3, r4, r5, r6, r7;
int length; // 0 = not yet born
};
// One owned window's step at sweep `sweep`: birth/slide, apply, write the
// bottom row. Birth loads all eight rows with a clamped row index: the
// clamped duplicates equal row kN-1 (the pinned bottom during birth) and
// are shifted into validity as the window grows, so no predicate is
// needed on the loads.
__device__ __forceinline__ void window_step(
float* panel,
const float* sweep_records, // this sweep's contiguous records
int lane,
int sweep,
int count,
int segment,
Win& w) {
if (segment >= count) {
return;
}
const int head = sweep + kKD * segment;
const int length = min(kKD, kN - head);
if (w.length == 0) {
w.r0 = panel[min(head + 0, kN - 1) * kPanelCols + lane];
w.r1 = panel[min(head + 1, kN - 1) * kPanelCols + lane];
w.r2 = panel[min(head + 2, kN - 1) * kPanelCols + lane];
w.r3 = panel[min(head + 3, kN - 1) * kPanelCols + lane];
w.r4 = panel[min(head + 4, kN - 1) * kPanelCols + lane];
w.r5 = panel[min(head + 5, kN - 1) * kPanelCols + lane];
w.r6 = panel[min(head + 6, kN - 1) * kPanelCols + lane];
w.r7 = panel[min(head + 7, kN - 1) * kPanelCols + lane];
} else {
// Slide up one row: shift the ring, read the entering top row (written
// the previous sweep by segment g-1's exiting bottom row).
w.r7 = w.r6;
w.r6 = w.r5;
w.r5 = w.r4;
w.r4 = w.r3;
w.r3 = w.r2;
w.r2 = w.r1;
w.r1 = w.r0;
w.r0 = panel[head * kPanelCols + lane];
}
w.length = length;
// Apply the reflector (v0 = 1 implicit).
const float* record = sweep_records + segment * kRecordWidth;
const float tau = record[0];
float dot = w.r0;
if (1 < length) dot = fmaf(record[1], w.r1, dot);
if (2 < length) dot = fmaf(record[2], w.r2, dot);
if (3 < length) dot = fmaf(record[3], w.r3, dot);
if (4 < length) dot = fmaf(record[4], w.r4, dot);
if (5 < length) dot = fmaf(record[5], w.r5, dot);
if (6 < length) dot = fmaf(record[6], w.r6, dot);
if (7 < length) dot = fmaf(record[7], w.r7, dot);
const float weight = tau * dot;
w.r0 = fmaf(-weight, 1.0f, w.r0);
if (1 < length) w.r1 = fmaf(-weight, record[1], w.r1);
if (2 < length) w.r2 = fmaf(-weight, record[2], w.r2);
if (3 < length) w.r3 = fmaf(-weight, record[3], w.r3);
if (4 < length) w.r4 = fmaf(-weight, record[4], w.r4);
if (5 < length) w.r5 = fmaf(-weight, record[5], w.r5);
if (6 < length) w.r6 = fmaf(-weight, record[6], w.r6);
if (7 < length) w.r7 = fmaf(-weight, record[7], w.r7);
// Write the bottom row (exits next sweep in steady state; redundant but
// safe during birth, where the bottom is pinned at row kN-1).
float bottom = w.r7;
if (length == 1) bottom = w.r0;
if (length == 2) bottom = w.r1;
if (length == 3) bottom = w.r2;
if (length == 4) bottom = w.r3;
if (length == 5) bottom = w.r4;
if (length == 6) bottom = w.r5;
if (length == 7) bottom = w.r6;
panel[(head + length - 1) * kPanelCols + lane] = bottom;
}
__device__ __forceinline__ void window_flush(
float* panel,
int lane,
int segment,
const Win& w) {
if (w.length == 0) {
return;
}
const int head = 1 + kKD * segment;
const int length = w.length;
if (0 < length - 1) panel[(head + 0) * kPanelCols + lane] = w.r0;
if (1 < length - 1) panel[(head + 1) * kPanelCols + lane] = w.r1;
if (2 < length - 1) panel[(head + 2) * kPanelCols + lane] = w.r2;
if (3 < length - 1) panel[(head + 3) * kPanelCols + lane] = w.r3;
if (4 < length - 1) panel[(head + 4) * kPanelCols + lane] = w.r4;
if (5 < length - 1) panel[(head + 5) * kPanelCols + lane] = w.r5;
if (6 < length - 1) panel[(head + 6) * kPanelCols + lane] = w.r6;
}
// minBlocks 2 at 512 threads pins ptxas to 64 registers (probe-verified,
// zero stack) => 2 CTAs/SM on B200 (shared <= 227 KB opt-in).
//
// StageHH:
// double-buffer the per-sweep record block through shared memory. The
// sweep-major hh ABI makes sweep s's <= 64 records CONTIGUOUS, so the
// next sweep's block is prefetched with one coalesced load per thread,
// published by the existing per-sweep barrier, and the apply chain reads
// records at shared latency instead of L2/DRAM latency. Arithmetic is
// identical op-for-op -> bitwise vs the unstaged config.
template <int Threads, bool StageHH>
__global__ __launch_bounds__(Threads, Threads == 512 ? 2 : 1)
void bt_kd8_panel_apply_f32(
float* __restrict__ x,
const float* __restrict__ hh,
int* __restrict__ info,
int batch) {
constexpr int kWarps = Threads / 32;
constexpr int kWindows = kMaxSegments / kWarps;
static_assert(kMaxSegments % kWarps == 0, "ownership must tile");
constexpr int kStageFloats = kMaxSegments * kRecordWidth; // 512
// 65,536 B (+4 KB staging) > the 48 KB static-shared limit -> dynamic.
extern __shared__ float panel[];
float* stage = panel + kN * kPanelCols; // [2][kStageFloats] when StageHH
const int matrix = static_cast<int>(blockIdx.y);
const int panel_begin = static_cast<int>(blockIdx.x) * kPanelCols;
if (matrix >= batch) {
return;
}
const int tid = static_cast<int>(threadIdx.x);
const int lane = tid & 31;
const int warp = tid >> 5;
float* matrix_x = x + static_cast<size_t>(matrix) * kN * kN;
const float* matrix_hh =
hh + static_cast<size_t>(matrix) * kReflectors * kRecordWidth;
bool bad = false;
for (int index = tid; index < kN * kPanelCols; index += Threads) {
const int row = index / kPanelCols;
const int column = index - row * kPanelCols;
const float value = matrix_x[
static_cast<size_t>(row) * kN + panel_begin + column];
panel[index] = value;
if (!isfinite(value)) {
bad = true;
}
}
__syncthreads();
// Register window rings, one per owned segment column (see Win).
Win w0{}, w1{}, w2{}, w3{};
if (StageHH) {
// Preload sweep kN-2's single record into staging buffer 0.
const float* src = matrix_hh
+ static_cast<size_t>(kSweepPrefix[kN - 3]) * kRecordWidth;
for (int i = tid; i < kRecordWidth; i += Threads) {
stage[i] = src[i];
}
__syncthreads();
}
int buffer = 0;
for (int sweep = kN - 2; sweep >= 1; --sweep) {
const int count = (kN - 1 - sweep + kKD - 1) / kKD;
const float* sweep_records;
if (StageHH) {
sweep_records = stage + buffer * kStageFloats;
if (sweep > 1) {
// Prefetch the NEXT sweep's contiguous record block (coalesced,
// one load per thread); the per-sweep barrier publishes it.
const int next_count = (kN - 1 - (sweep - 1) + kKD - 1) / kKD;
const float* src = matrix_hh
+ static_cast<size_t>(kSweepPrefix[sweep - 2]) * kRecordWidth;
float* destination = stage + (buffer ^ 1) * kStageFloats;
for (int i = tid; i < next_count * kRecordWidth; i += Threads) {
destination[i] = src[i];
}
}
} else {
sweep_records = matrix_hh
+ static_cast<size_t>(kSweepPrefix[sweep - 1]) * kRecordWidth;
}
window_step(panel, sweep_records, lane, sweep, count,
warp + 0 * kWarps, w0);
if (kWindows >= 2) {
window_step(panel, sweep_records, lane, sweep, count,
warp + 1 * kWarps, w1);
}
if (kWindows >= 3) {
window_step(panel, sweep_records, lane, sweep, count,
warp + 2 * kWarps, w2);
}
if (kWindows >= 4) {
window_step(panel, sweep_records, lane, sweep, count,
warp + 3 * kWarps, w3);
}
__syncthreads();
buffer ^= 1;
}
// Final flush: windows hold rows head..head+length-2 not yet written
// (bottom row was written during the sweep-1 step).
window_flush(panel, lane, warp + 0 * kWarps, w0);
if (kWindows >= 2) {
window_flush(panel, lane, warp + 1 * kWarps, w1);
}
if (kWindows >= 3) {
window_flush(panel, lane, warp + 2 * kWarps, w2);
}
if (kWindows >= 4) {
window_flush(panel, lane, warp + 3 * kWarps, w3);
}
__syncthreads();
for (int index = tid; index < kN * kPanelCols; index += Threads) {
const int row = index / kPanelCols;
const int column = index - row * kPanelCols;
const float value = panel[index];
matrix_x[static_cast<size_t>(row) * kN + panel_begin + column] = value;
if (!isfinite(value)) {
bad = true;
}
}
if (bad) {
atomicExch(&info[matrix], 40);
}
}
std::array<uint16_t, kSweepPrefixEntries> build_sweep_prefix() {
std::array<uint16_t, kSweepPrefixEntries> prefix{};
int total = 0;
prefix[0] = 0;
for (int sweep = 1; sweep <= kN - 2; ++sweep) {
const int remaining = kN - sweep - 1;
total += (remaining + kKD - 1) / kKD;
TORCH_CHECK(total <= kReflectors, "KD8 sweep prefix overflow");
prefix[sweep] = static_cast<uint16_t>(total);
}
TORCH_CHECK(total == kReflectors, "KD8 reflector count mismatch");
return prefix;
}
void ensure_sweep_prefix(int device) {
static std::mutex mutex;
static std::unordered_set<int> initialized_devices;
std::lock_guard<std::mutex> lock(mutex);
if (initialized_devices.count(device) != 0) {
return;
}
const auto prefix = build_sweep_prefix();
C10_CUDA_CHECK(cudaMemcpyToSymbol(
kSweepPrefix,
prefix.data(),
prefix.size() * sizeof(prefix[0]),
0,
cudaMemcpyHostToDevice));
initialized_devices.insert(device);
}
void validate_tensors(
const torch::Tensor& x,
const torch::Tensor& hh,
const torch::Tensor& info) {
TORCH_CHECK(x.is_cuda() && hh.is_cuda() && info.is_cuda(),
"all tensors must be CUDA");
TORCH_CHECK(x.scalar_type() == torch::kFloat32 &&
hh.scalar_type() == torch::kFloat32,
"x/hh must be float32");
TORCH_CHECK(info.scalar_type() == at::kInt, "info must be int32");
TORCH_CHECK(x.is_contiguous() && hh.is_contiguous() && info.is_contiguous(),
"all tensors must be contiguous");
TORCH_CHECK(x.dim() == 3 && x.size(1) == kN && x.size(2) == kN,
"x must have shape [B,512,512]");
const auto batch = x.size(0);
TORCH_CHECK(
hh.sizes() == torch::IntArrayRef({batch, kReflectors, kRecordWidth}),
"hh shape mismatch");
TORCH_CHECK(info.sizes() == torch::IntArrayRef({batch}),
"info shape mismatch");
const int device = x.get_device();
TORCH_CHECK(hh.get_device() == device && info.get_device() == device,
"all tensors must use one device");
}
template <int Threads, bool StageHH>
void launch_config(
const torch::Tensor& x,
const torch::Tensor& hh,
const torch::Tensor& info) {
validate_tensors(x, hh, info);
c10::cuda::CUDAGuard guard(x.device());
const int device = x.get_device();
ensure_sweep_prefix(device);
constexpr size_t kShared = dynamic_shared_bytes(StageHH);
int maximum_shared = 0;
C10_CUDA_CHECK(cudaDeviceGetAttribute(
&maximum_shared, cudaDevAttrMaxSharedMemoryPerBlockOptin, device));
TORCH_CHECK(maximum_shared >= static_cast<int>(kShared),
"device shared-memory limit is below the panel contract");
C10_CUDA_CHECK(cudaFuncSetAttribute(
bt_kd8_panel_apply_f32<Threads, StageHH>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
static_cast<int>(kShared)));
const int batch = static_cast<int>(x.size(0));
dim3 grid(kN / kPanelCols, batch, 1);
bt_kd8_panel_apply_f32<Threads, StageHH>
<<<grid, Threads, kShared>>>(
x.data_ptr<float>(),
hh.data_ptr<float>(),
info.data_ptr<int>(),
batch);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
template <int Threads, bool StageHH>
std::vector<int64_t> report_config() {
constexpr size_t kShared = dynamic_shared_bytes(StageHH);
C10_CUDA_CHECK(cudaFuncSetAttribute(
bt_kd8_panel_apply_f32<Threads, StageHH>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
static_cast<int>(kShared)));
cudaFuncAttributes attributes{};
C10_CUDA_CHECK(cudaFuncGetAttributes(
&attributes, bt_kd8_panel_apply_f32<Threads, StageHH>));
int active_blocks = 0;
C10_CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&active_blocks, bt_kd8_panel_apply_f32<Threads, StageHH>, Threads,
kShared));
int device = -1;
C10_CUDA_CHECK(cudaGetDevice(&device));
int sm_count = 0;
C10_CUDA_CHECK(cudaDeviceGetAttribute(
&sm_count, cudaDevAttrMultiProcessorCount, device));
return {
Threads,
static_cast<int64_t>(kShared),
attributes.numRegs,
static_cast<int64_t>(attributes.localSizeBytes),
static_cast<int64_t>(attributes.sharedSizeBytes),
active_blocks,
sm_count,
};
}
} // namespace
void backtransform_kd8_inplace(
const torch::Tensor& x,
const torch::Tensor& hh,
const torch::Tensor& info,
const std::string& config) {
if (config == "t512") {
launch_config<512, false>(x, hh, info);
} else if (config == "t1024") {
launch_config<1024, false>(x, hh, info);
} else if (config == "t512s") {
launch_config<512, true>(x, hh, info);
} else {
TORCH_CHECK(false, "unknown back-transform config: ", config);
}
}
std::vector<int64_t> bt_resource_report(const std::string& config) {
if (config == "t512") {
return report_config<512, false>();
}
if (config == "t1024") {
return report_config<1024, false>();
}
if (config == "t512s") {
return report_config<512, true>();
}
TORCH_CHECK(false, "unknown back-transform config: ", config);
return {};
}
} // namespace lane_f
namespace lane_h {
namespace {
constexpr float kEps32 = 1.1920929e-07f; // FLT_EPSILON
constexpr float kSafmin32 = 1.1754944e-38f; // FLT_MIN
constexpr float kOrtol = 1e-3f; // DSTEIN grouping (scaled units)
constexpr float kPertol = 10.0f * kEps32; // DSTEIN shift separation
constexpr float kLog2Gtol = 9.9657842847f; // log2(1e3) growth acceptance
constexpr float kXmax = 1e15f; // SLAGTS dynamic-rescale bound
constexpr float kDepTol = 0.1f; // dependent-vector rescue bound
constexpr float kLog2Tiny30 = -99.65784285f; // log2(1e-30)
constexpr int kMaxIts = 5; // invit rounds cap
constexpr int kBisectCap = 64; // oracle bisection pass cap
constexpr int kThreadsK1 = 512;
constexpr int kWarpsK1 = kThreadsK1 / 32;
constexpr int kStatsSlots = 16;
// stats layout per matrix (int32):
// 0 bisect iters max 1 bisect iters sum 2 group count
// 3 max group size 4 sum k^2 (k>1) 5 grouped vector count
// 6 invit rounds max 7 invit rounds sum 8 stagnant after maxits
// 9 dependent rescues 10 SLAGTS guard hits 11..15 reserved
constexpr unsigned kFullMask = 0xffffffffu;
template <int N>
constexpr int k1_dynamic_shared_bytes() {
return N * (N + 1) * static_cast<int>(sizeof(float));
}
__device__ __forceinline__ float warp_sum(float value) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
value += __shfl_xor_sync(kFullMask, value, offset);
}
return value;
}
__device__ __forceinline__ float warp_max(float value) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
value = fmaxf(value, __shfl_xor_sync(kFullMask, value, offset));
}
return value;
}
__device__ __forceinline__ float warp_min(float value) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
value = fminf(value, __shfl_xor_sync(kFullMask, value, offset));
}
return value;
}
__device__ __forceinline__ int warp_imax(int value) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
value = max(value, __shfl_xor_sync(kFullMask, value, offset));
}
return value;
}
__device__ __forceinline__ int warp_isum(int value) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
value += __shfl_xor_sync(kFullMask, value, offset);
}
return value;
}
// splitmix32-style hash -> uniform [-1, 1). Deviation D2: the oracle used
// numpy Philox; the algorithm only needs a generic random unit RHS.
__device__ __forceinline__ float hash_uniform(uint32_t b, uint32_t j,
uint32_t i, uint32_t t) {
uint32_t x = (b * 0x9E3779B9u) ^ (j * 0x85EBCA6Bu) ^ (i * 0xC2B2AE35u) ^
(t * 0x27D4EB2Fu) ^ 0xB5297A4Du;
x ^= x >> 16;
x *= 0x7FEB352Du;
x ^= x >> 15;
x *= 0x846CA68Bu;
x ^= x >> 16;
return static_cast<float>(x >> 8) * (2.0f / 16777216.0f) - 1.0f;
}
struct ControlK1 {
float gamma;
float tau;
float inv;
float alpha2;
int nonfinite;
};
// ======================================================================
// K1: load/prescale + SSYTD2 tridiagonalization + SORGTR Q accumulation.
// Port of midn-fused-n176-door kernel A rev-3 (job 3071988) with:
// * template <N, kSharedW>: W in shared (ld = N+1, kernel-A verbatim)
// or in global using the q output buffer (ld = N).
// * Deviation D1 (toward the oracle): normalizer is the power-of-two
// gamma = exp2f(floorf(log2f(amax))) — the oracle's SLASCL-analog
// prescale, exact in fp32, required for the ±1e19 battery.
// ======================================================================
template <int N, bool kSharedW>
__global__ void __launch_bounds__(kThreadsK1, 1)
midn_tridiag_qform_kernel(const float* __restrict__ a_in,
float* __restrict__ q_out,
float* __restrict__ d_ws,
float* __restrict__ e_ws,
float* __restrict__ gamma_ws,
int* __restrict__ info,
int batch_count) {
constexpr int kLd = kSharedW ? (N + 1) : N;
extern __shared__ float dyn_shared[];
__shared__ float sd[N];
__shared__ float se[N];
__shared__ float stau[N];
__shared__ float sp[N];
__shared__ float svw[N];
__shared__ float partials[kWarpsK1];
__shared__ ControlK1 ctl;
const int batch = blockIdx.x;
if (batch >= batch_count) {
return;
}
const int tid = threadIdx.x;
const int warp = tid >> 5;
const int lane = tid & 31;
const float* a = a_in + static_cast<size_t>(batch) * N * N;
float* q = q_out + static_cast<size_t>(batch) * N * N;
float* W = kSharedW ? dyn_shared : q;
if (tid == 0) {
ctl.nonfinite = 0;
}
__syncthreads();
// Phase 0: load, finite check, global |A|max scan.
float local_max = 0.0f;
int local_bad = 0;
for (int idx = tid; idx < N * N; idx += kThreadsK1) {
const int row = idx / N;
const int col = idx - row * N;
const float value = a[idx];
W[row * kLd + col] = value;
local_max = fmaxf(local_max, fabsf(value));
local_bad |= !isfinite(value);
}
const float wmax = warp_max(local_max);
if (lane == 0) {
partials[warp] = wmax;
}
if (local_bad) {
atomicOr(&ctl.nonfinite, 1);
}
__syncthreads();
if (tid == 0) {
float amax = 0.0f;
for (int w = 0; w < kWarpsK1; ++w) {
amax = fmaxf(amax, partials[w]);
}
// D1: exact power-of-two prescale (oracle SLASCL analog).
ctl.gamma = amax > 0.0f ? exp2f(floorf(log2f(amax))) : 0.0f;
}
__syncthreads();
if (ctl.nonfinite != 0 || ctl.gamma <= 0.0f) {
// Degenerate: Q = I, tridiagonal = 0; K2/K3 then pass through with
// lam = 0 and a full-group CGS-orthonormal V (correct for A = 0).
for (int idx = tid; idx < N * N; idx += kThreadsK1) {
const int row = idx / N;
const int col = idx - row * N;
q[idx] = row == col ? 1.0f : 0.0f;
}
for (int i = tid; i < N; i += kThreadsK1) {
d_ws[batch * N + i] = 0.0f;
e_ws[batch * N + i] = 0.0f;
}
if (tid == 0) {
gamma_ws[batch] = 0.0f;
info[batch] = ctl.nonfinite != 0 ? 2 : 0;
}
return;
}
const float gamma = ctl.gamma;
for (int idx = tid; idx < N * N; idx += kThreadsK1) {
const int row = idx / N;
const int col = idx - row * N;
W[row * kLd + col] = W[row * kLd + col] / gamma;
}
__syncthreads();
// Phase 1: Householder tridiagonalization (SSYTD2, lower, strict FP32).
for (int k = 0; k <= N - 3; ++k) {
const int base = k + 1;
const int len = N - 1 - k;
float part = 0.0f;
for (int r = k + 2 + tid; r < N; r += kThreadsK1) {
const float x = W[r * kLd + k];
part = fmaf(x, x, part);
}
const float wsum = warp_sum(part);
if (lane == 0) {
partials[warp] = wsum;
}
__syncthreads();
if (tid == 0) {
float sigma = 0.0f;
for (int w = 0; w < kWarpsK1; ++w) {
sigma += partials[w];
}
const float alpha = W[base * kLd + k];
sd[k] = W[k * kLd + k];
if (sigma <= 0.0f) {
ctl.tau = 0.0f;
ctl.inv = 0.0f;
se[k] = alpha;
} else {
const float norm = sqrtf(fmaf(alpha, alpha, sigma));
const float beta = -copysignf(norm, alpha);
ctl.tau = (beta - alpha) / beta; // strict FP32 division
ctl.inv = 1.0f / (alpha - beta); // strict FP32 division
se[k] = beta;
}
stau[k] = ctl.tau;
}
__syncthreads();
const float tau = ctl.tau;
const float inv = ctl.inv;
if (tid == 0) {
W[base * kLd + k] = 1.0f;
}
for (int r = k + 2 + tid; r < N; r += kThreadsK1) {
const float x = W[r * kLd + k];
W[r * kLd + k] = tau == 0.0f ? 0.0f : x * inv;
}
__syncthreads();
if (tau != 0.0f) {
if (tid < len) {
const int r = base + tid;
const float* arow = &W[r * kLd];
float acc = 0.0f;
for (int j = base; j < N; ++j) {
acc = fmaf(arow[j], W[j * kLd + k], acc);
}
sp[r] = tau * acc;
}
__syncthreads();
float pvpart = 0.0f;
for (int r = base + tid; r < N; r += kThreadsK1) {
pvpart = fmaf(sp[r], W[r * kLd + k], pvpart);
}
const float pvsum = warp_sum(pvpart);
if (lane == 0) {
partials[warp] = pvsum;
}
__syncthreads();
if (tid == 0) {
float pv = 0.0f;
for (int w = 0; w < kWarpsK1; ++w) {
pv += partials[w];
}
ctl.alpha2 = -0.5f * tau * pv;
}
__syncthreads();
const float alpha2 = ctl.alpha2;
for (int r = base + tid; r < N; r += kThreadsK1) {
svw[r] = fmaf(alpha2, W[r * kLd + k], sp[r]);
}
__syncthreads();
const int total = len * len;
for (int idx = tid; idx < total; idx += kThreadsK1) {
const int i = base + idx / len;
const int j = base + idx - (idx / len) * len;
const float vi = W[i * kLd + k];
const float vj = W[j * kLd + k];
float value = W[i * kLd + j];
value = fmaf(-vi, svw[j], value);
value = fmaf(-svw[i], vj, value);
W[i * kLd + j] = value;
}
}
__syncthreads();
}
if (tid == 0) {
sd[N - 2] = W[(N - 2) * kLd + (N - 2)];
sd[N - 1] = W[(N - 1) * kLd + (N - 1)];
se[N - 2] = W[(N - 1) * kLd + (N - 2)];
se[N - 1] = 0.0f;
}
__syncthreads();
// Phase 2: in-place backward accumulation of Q = H_0 ... H_{N-3}.
if (tid == 0) {
W[(N - 1) * kLd + (N - 1)] = 1.0f; // P_{N-2} = I seed
}
__syncthreads();
for (int k = N - 3; k >= 0; --k) {
const float tau = stau[k];
for (int j = k + 1 + warp; j < N; j += kWarpsK1) {
if (j == k + 1) {
for (int r = k + 1 + lane; r < N; r += 32) {
if (r == k + 1) {
W[r * kLd + j] = 1.0f - tau;
} else {
W[r * kLd + j] = -tau * W[r * kLd + k];
}
}
} else {
float part = 0.0f;
for (int r = k + 2 + lane; r < N; r += 32) {
part = fmaf(W[r * kLd + k], W[r * kLd + j], part);
}
const float tdot = tau * warp_sum(part);
for (int r = k + 1 + lane; r < N; r += 32) {
if (r == k + 1) {
W[r * kLd + j] = -tdot;
} else {
W[r * kLd + j] = fmaf(-W[r * kLd + k], tdot, W[r * kLd + j]);
}
}
}
}
__syncthreads();
}
for (int j = tid; j < N; j += kThreadsK1) {
W[j] = j == 0 ? 1.0f : 0.0f; // row 0
if (j > 0) {
W[j * kLd] = 0.0f; // column 0
}
}
__syncthreads();
// Emit d, e, gamma (Q already lives in q for the global-W variant).
if constexpr (kSharedW) {
for (int idx = tid; idx < N * N; idx += kThreadsK1) {
const int row = idx / N;
const int col = idx - row * N;
q[idx] = W[row * kLd + col];
}
}
for (int i = tid; i < N; i += kThreadsK1) {
d_ws[batch * N + i] = sd[i];
e_ws[batch * N + i] = se[i];
}
if (tid == 0) {
gamma_ws[batch] = gamma;
info[batch] = 0;
}
}
// ======================================================================
// K2: Sturm-count bisection for all eigenvalues + grouping epilogue.
// Mirrors oracle.bisect_eigenvalues + group_eigenvalues + the perturbed
// shifts of oracle.inverse_iteration (R1, DSTEIN standard).
// ======================================================================
template <int N, int NT>
__global__ void __launch_bounds__(NT)
midn_bisect_kernel(const float* __restrict__ d_ws,
const float* __restrict__ e_ws,
const float* __restrict__ gamma_ws,
float* __restrict__ ds_ws,
float* __restrict__ es_ws,
float* __restrict__ w_ws,
float* __restrict__ lam_out,
float* __restrict__ xs_ws,
int* __restrict__ gstart_ws,
int* __restrict__ minits_ws,
float* __restrict__ pivmin_ws,
float* __restrict__ onenrm_ws,
int* __restrict__ stats_ws,
int batch_count) {
constexpr int kWarps = NT / 32;
__shared__ float sds[N];
__shared__ float ses[N];
__shared__ float se2[N];
__shared__ float sw[N];
__shared__ int sgstart[N];
__shared__ int sminits[N];
__shared__ float partials[kWarps];
__shared__ int ipartials[kWarps];
__shared__ float s_onenrm, s_pivmin, s_glo, s_ghi;
const int batch = blockIdx.x;
if (batch >= batch_count) {
return;
}
const int tid = threadIdx.x;
const int warp = tid >> 5;
const int lane = tid & 31;
for (int i = tid; i < N; i += NT) {
sds[i] = d_ws[batch * N + i];
ses[i] = e_ws[batch * N + i]; // e[N-1] == 0 from K1
}
__syncthreads();
// onenrm = max_i(|d_i| + |e_{i-1}| + |e_i|), clamped >= safmin.
float lmax = 0.0f;
for (int i = tid; i < N; i += NT) {
float t = fabsf(sds[i]);
if (i > 0) t += fabsf(ses[i - 1]);
if (i < N - 1) t += fabsf(ses[i]);
lmax = fmaxf(lmax, t);
}
lmax = warp_max(lmax);
if (lane == 0) partials[warp] = lmax;
__syncthreads();
if (tid == 0) {
float m = 0.0f;
for (int w = 0; w < kWarps; ++w) m = fmaxf(m, partials[w]);
s_onenrm = fmaxf(m, kSafmin32);
}
__syncthreads();
const float onenrm = s_onenrm;
for (int i = tid; i < N; i += NT) {
const float dv = sds[i] / onenrm;
const float ev = ses[i] / onenrm;
sds[i] = dv;
ses[i] = ev;
se2[i] = ev * ev;
}
__syncthreads();
// pivmin + Gershgorin brackets (three reductions).
float le2 = 0.0f;
float lglo = FLT_MAX;
float lghi = -FLT_MAX;
for (int i = tid; i < N; i += NT) {
if (i < N - 1) le2 = fmaxf(le2, se2[i]);
float r = 0.0f;
if (i > 0) r += fabsf(ses[i - 1]);
if (i < N - 1) r += fabsf(ses[i]);
lglo = fminf(lglo, sds[i] - r);
lghi = fmaxf(lghi, sds[i] + r);
}
le2 = warp_max(le2);
if (lane == 0) partials[warp] = le2;
__syncthreads();
if (tid == 0) {
float m = 0.0f;
for (int w = 0; w < kWarps; ++w) m = fmaxf(m, partials[w]);
s_pivmin = fmaxf(m * kSafmin32, kSafmin32);
}
__syncthreads();
lglo = warp_min(lglo);
if (lane == 0) partials[warp] = lglo;
__syncthreads();
if (tid == 0) {
float m = FLT_MAX;
for (int w = 0; w < kWarps; ++w) m = fminf(m, partials[w]);
s_glo = m;
}
__syncthreads();
lghi = warp_max(lghi);
if (lane == 0) partials[warp] = lghi;
__syncthreads();
if (tid == 0) {
float m = -FLT_MAX;
for (int w = 0; w < kWarps; ++w) m = fmaxf(m, partials[w]);
const float width = fmaxf(m - s_glo, 1.0f) * kEps32;
s_ghi = m + width;
s_glo = s_glo - width;
}
__syncthreads();
// Per-root bisection: thread k refines eigenvalue k (0-based ascending).
const float pivmin = s_pivmin;
int iters = 0;
if (tid < N) {
float lo = s_glo;
float hi = s_ghi;
for (int pass = 0; pass < kBisectCap; ++pass) {
const float tol =
2.0f * kEps32 * fmaxf(fabsf(lo), fabsf(hi)) + 2.0f * pivmin;
if (hi - lo <= tol) break;
const float mid = 0.5f * (lo + hi);
// Sturm count: eigenvalues < mid (oracle sign convention: a
// near-zero pivot is forced NEGATIVE).
float qv = sds[0] - mid;
if (fabsf(qv) < pivmin) qv = -pivmin;
int cnt = qv < 0.0f;
for (int i = 1; i < N; ++i) {
qv = sds[i] - mid - se2[i - 1] / qv;
if (fabsf(qv) < pivmin) qv = -pivmin;
cnt += qv < 0.0f;
}
if (cnt >= tid + 1) {
hi = mid;
} else {
lo = mid;
}
++iters;
}
sw[tid] = 0.5f * (lo + hi);
}
// Iteration stats.
int imax = warp_imax(iters);
int isum = warp_isum(iters);
if (lane == 0) ipartials[warp] = imax;
__syncthreads();
if (tid == 0) {
int m = 0;
for (int w = 0; w < kWarps; ++w) m = max(m, ipartials[w]);
stats_ws[batch * kStatsSlots + 0] = m;
}
__syncthreads();
if (lane == 0) ipartials[warp] = isum;
__syncthreads();
// Thread-0 epilogue: monotone belt, grouping, minits, perturbed shifts.
if (tid == 0) {
int s = 0;
for (int w = 0; w < kWarps; ++w) s += ipartials[w];
stats_ws[batch * kStatsSlots + 1] = s;
for (int j = 1; j < N; ++j) sw[j] = fmaxf(sw[j], sw[j - 1]);
int start = 0;
int ngroups = 0, maxg = 0, sumk2 = 0, grouped = 0;
for (int j = 1; j <= N; ++j) {
const bool same = (j < N) && (sw[j] - sw[j - 1] <= kOrtol);
if (!same) {
const int sz = j - start;
++ngroups;
maxg = max(maxg, sz);
if (sz > 1) {
sumk2 += sz * sz;
grouped += sz;
}
for (int t = start; t < j; ++t) {
sgstart[t] = start;
sminits[t] = sz > 1 ? 3 : 2; // oracle FINDING #1
}
start = j;
}
}
stats_ws[batch * kStatsSlots + 2] = ngroups;
stats_ws[batch * kStatsSlots + 3] = maxg;
stats_ws[batch * kStatsSlots + 4] = sumk2;
stats_ws[batch * kStatsSlots + 5] = grouped;
// R1 perturbed shifts, cumulative within groups (DSTEIN standard).
float prev = sw[0];
xs_ws[batch * N + 0] = prev;
for (int j = 1; j < N; ++j) {
float x = sw[j];
if (sgstart[j] < j) { // not a group leader
const float lim = prev + kPertol;
if (x < lim) x = lim;
}
xs_ws[batch * N + j] = x;
prev = x;
}
}
__syncthreads();
const float gamma = gamma_ws[batch];
for (int i = tid; i < N; i += NT) {
ds_ws[batch * N + i] = sds[i];
es_ws[batch * N + i] = ses[i];
w_ws[batch * N + i] = sw[i];
// Oracle: lam = fp32( fp64(fp32(w * onenrm)) * gamma ).
const float t = sw[i] * onenrm;
lam_out[batch * N + i] = static_cast<float>(
static_cast<double>(t) * static_cast<double>(gamma));
gstart_ws[batch * N + i] = sgstart[i];
minits_ws[batch * N + i] = sminits[i];
}
if (tid == 0) {
pivmin_ws[batch] = pivmin;
onenrm_ws[batch] = onenrm;
}
}
// ======================================================================
// K3: DSTEIN inverse iteration. One CTA per matrix, thread j owns
// eigenvector j; CGS/MGS phases are CTA-cooperative and sequential over
// j inside each group (oracle operation order).
// Per-vector global arrays are laid out [batch][row i][vector j] so the
// per-thread serial walks are coalesced across the CTA.
// ======================================================================
// Fused SLAGTF factor + forward elimination + SLAGTS back-substitution
// with the MANDATORY dynamic-rescale guard (oracle solve_tridiag /
// factor_tridiag; without the guard 2/16 oracle rows NaN'd on pivmin
// pivots). Column pointers stride N floats between rows.
template <int N>
__device__ void thomas_solve_column(const float* __restrict__ sds,
const float* __restrict__ ses,
const float xj, const float pivmin,
float* __restrict__ y, // rhs column
float* __restrict__ x, // solution col
float* __restrict__ u0,
float* __restrict__ u1,
float* __restrict__ u2,
float* out_log2growth, float* out_mx,
int* guard_hits) {
float dlog2 = 0.0f; // log2 of accumulated downscale (deviation D5)
float di = sds[0] - xj;
float sup = ses[0];
for (int i = 0; i < N - 1; ++i) {
const float ai = di;
const float ci = ses[i]; // symmetric subdiagonal
const float bi = sup;
const bool swp = fabsf(ai) < fabsf(ci);
float piv = swp ? ci : ai;
if (fabsf(piv) < pivmin) piv = piv < 0.0f ? -pivmin : pivmin;
const float m = (swp ? ai : ci) / piv;
const float a_next = sds[i + 1] - xj;
const float b_next = (i + 1 < N - 1) ? ses[i + 1] : 0.0f;
u0[static_cast<size_t>(i) * N] = piv;
u1[static_cast<size_t>(i) * N] = swp ? a_next : bi;
u2[static_cast<size_t>(i) * N] = swp ? b_next : 0.0f;
di = swp ? (bi - m * a_next) : (a_next - m * bi);
sup = swp ? (-m * b_next) : b_next;
const float yi = y[static_cast<size_t>(i) * N];
const float yn = y[static_cast<size_t>(i + 1) * N];
const float y_low = swp ? yn : yi;
const float y_high = (swp ? yi : yn) - m * y_low;
y[static_cast<size_t>(i) * N] = y_low;
y[static_cast<size_t>(i + 1) * N] = y_high;
if (fabsf(y_high) > 1e30f) { // |m| <= 1 so this is rare (oracle)
for (int r = 0; r < N; ++r) y[static_cast<size_t>(r) * N] *= 1e-30f;
dlog2 += kLog2Tiny30;
}
}
{
float last = di;
if (fabsf(last) < pivmin) last = last < 0.0f ? -pivmin : pivmin;
u0[static_cast<size_t>(N - 1) * N] = last;
}
// Back-substitution with pre-division rescale (SLAGTS analog).
for (int i = N - 1; i >= 0; --i) {
float num = y[static_cast<size_t>(i) * N];
if (i + 1 < N) {
num -= u1[static_cast<size_t>(i) * N] *
x[static_cast<size_t>(i + 1) * N];
}
if (i + 2 < N) {
num -= u2[static_cast<size_t>(i) * N] *
x[static_cast<size_t>(i + 2) * N];
}
const float den = u0[static_cast<size_t>(i) * N];
const float absn = fabsf(num);
const float absd = fabsf(den);
if (absn > absd * kXmax) { // den O(1)-bounded: absd*kXmax cannot overflow
const float f = (absd * kXmax) / absn;
for (int r = i + 1; r < N; ++r) x[static_cast<size_t>(r) * N] *= f;
for (int r = 0; r < i; ++r) y[static_cast<size_t>(r) * N] *= f;
num *= f;
dlog2 += log2f(f);
++(*guard_hits);
}
x[static_cast<size_t>(i) * N] = num / den;
}
float mx = 0.0f;
for (int r = 0; r < N; ++r) {
mx = fmaxf(mx, fabsf(x[static_cast<size_t>(r) * N]));
}
*out_mx = mx;
// growth = mx / downscale (mx -> 1 when zero, oracle convention).
*out_log2growth = (mx > 0.0f ? log2f(mx) : 0.0f) - dlog2;
}
template <int N, int NT>
__global__ void __launch_bounds__(NT)
midn_invit_kernel(const float* __restrict__ ds_ws,
const float* __restrict__ es_ws,
const float* __restrict__ xs_ws,
const int* __restrict__ gstart_ws,
const int* __restrict__ minits_ws,
const float* __restrict__ pivmin_ws,
float* __restrict__ v_ws,
float* __restrict__ y_ws,
float* __restrict__ u0_ws,
float* __restrict__ u1_ws,
float* __restrict__ u2_ws,
int* __restrict__ stats_ws,
int batch_count) {
constexpr int kWarps = NT / 32;
__shared__ float sds[N];
__shared__ float ses[N];
__shared__ float sxs[N];
__shared__ float sc[N]; // CGS coefficients
__shared__ int sgstart[N];
__shared__ signed char smin[N];
__shared__ signed char sacc[N];
__shared__ signed char srnd[N];
__shared__ float partials[kWarps];
__shared__ int s_naccepted;
__shared__ int s_rescues;
__shared__ int s_guard;
__shared__ float s_nn; // actual norm of the current projected column
__shared__ float s_rem; // rescue-trigger remainder (1.0 for empty block)
const int batch = blockIdx.x;
if (batch >= batch_count) {
return;
}
const int tid = threadIdx.x;
const int warp = tid >> 5;
const int lane = tid & 31;
const float pivmin = pivmin_ws[batch];
float* Vb = v_ws + static_cast<size_t>(batch) * N * N;
float* Yb = y_ws + static_cast<size_t>(batch) * N * N;
float* U0b = u0_ws + static_cast<size_t>(batch) * N * N;
float* U1b = u1_ws + static_cast<size_t>(batch) * N * N;
float* U2b = u2_ws + static_cast<size_t>(batch) * N * N;
for (int i = tid; i < N; i += NT) {
sds[i] = ds_ws[batch * N + i];
ses[i] = es_ws[batch * N + i];
sxs[i] = xs_ws[batch * N + i];
sgstart[i] = gstart_ws[batch * N + i];
smin[i] = static_cast<signed char>(minits_ws[batch * N + i]);
sacc[i] = 0;
srnd[i] = 0;
}
if (tid == 0) {
s_naccepted = 0;
s_rescues = 0;
s_guard = 0;
}
__syncthreads();
int guard_hits_local = 0;
// Round-0 RHS: hashed uniform [-1,1), L2-normalized per vector (Kahan
// fp32 accumulation — deviation D6).
if (tid < N) {
float sum = 0.0f, comp = 0.0f;
for (int r = 0; r < N; ++r) {
const float t = hash_uniform(batch, tid, r, 0u);
Yb[static_cast<size_t>(r) * N + tid] = t;
const float term = fmaf(t, t, -comp);
const float tmp = sum + term;
comp = (tmp - sum) - term;
sum = tmp;
}
const float nrm = sqrtf(sum);
if (nrm > 0.0f) {
for (int r = 0; r < N; ++r) {
Yb[static_cast<size_t>(r) * N + tid] /= nrm;
}
}
}
// -------------------------------------------------------------------
// CTA-cooperative helpers (macro-free lambdas; all threads call them).
// -------------------------------------------------------------------
// Project column jcol of M against V columns [k0, jj) (classical GS,
// one pass, fp32) and leave the ACTUAL post-projection norm in s_nn.
auto cta_project = [&](float* M, int jcol, int k0, int kend) {
for (int k = k0 + warp; k < kend; k += kWarps) {
float part = 0.0f;
for (int i = lane; i < N; i += 32) {
part = fmaf(Vb[static_cast<size_t>(i) * N + k],
M[static_cast<size_t>(i) * N + jcol], part);
}
part = warp_sum(part);
if (lane == 0) sc[k] = part;
}
__syncthreads();
if (kend > k0) {
for (int i = tid; i < N; i += NT) {
float acc = 0.0f;
for (int k = k0; k < kend; ++k) {
acc = fmaf(sc[k], Vb[static_cast<size_t>(i) * N + k], acc);
}
M[static_cast<size_t>(i) * N + jcol] -= acc;
}
}
__syncthreads();
float part = 0.0f;
for (int i = tid; i < N; i += NT) {
const float t = M[static_cast<size_t>(i) * N + jcol];
part = fmaf(t, t, part);
}
part = warp_sum(part);
if (lane == 0) partials[warp] = part;
__syncthreads();
if (tid == 0) {
float s = 0.0f;
for (int w = 0; w < kWarps; ++w) s += partials[w];
s_nn = sqrtf(s);
s_rem = (kend > k0) ? s_nn : 1.0f; // oracle: empty block -> rem = 1.0
}
__syncthreads();
};
auto cta_scale_column = [&](float* M, int jcol, float divisor) {
for (int i = tid; i < N; i += NT) {
M[static_cast<size_t>(i) * N + jcol] /= divisor;
}
__syncthreads();
};
// -------------------------------------------------------------------
// Rounds (oracle inverse_iteration main loop).
// -------------------------------------------------------------------
for (int it = 0; it < kMaxIts; ++it) {
__syncthreads();
if (s_naccepted >= N) break;
// Solve phase: thread j handles vector j if not yet accepted.
if (tid < N && !sacc[tid]) {
float lg = 0.0f, mx = 0.0f;
thomas_solve_column<N>(sds, ses, sxs[tid], pivmin, Yb + tid, Vb + tid,
U0b + tid, U1b + tid, U2b + tid, &lg, &mx,
&guard_hits_local);
float nrm = 0.0f;
if (mx > 0.0f) {
float sum = 0.0f, comp = 0.0f;
for (int r = 0; r < N; ++r) {
const float t = Vb[static_cast<size_t>(r) * N + tid] / mx;
Vb[static_cast<size_t>(r) * N + tid] = t;
const float term = fmaf(t, t, -comp);
const float tmp = sum + term;
comp = (tmp - sum) - term;
sum = tmp;
}
nrm = sqrtf(sum);
if (nrm > 0.0f) {
for (int r = 0; r < N; ++r) {
Vb[static_cast<size_t>(r) * N + tid] /= nrm;
}
}
}
srnd[tid] += 1;
// Acceptance: growth * nrm >= GTOL (log2 space) AND rounds >= minits.
if (nrm > 0.0f && lg + log2f(nrm) >= kLog2Gtol &&
srnd[tid] >= smin[tid]) {
sacc[tid] = 1;
atomicAdd(&s_naccepted, 1);
}
}
__syncthreads();
// CGS phase: multi-member groups only, sequential over j (oracle
// order: project against the CURRENT, already-updated v[s0..j)).
for (int jj = 0; jj < N; ++jj) {
if (smin[jj] != 3) continue; // singleton group (uniform: shared)
const int k0 = sgstart[jj];
cta_project(Vb, jj, k0, jj);
if (s_nn > 0.0f) cta_scale_column(Vb, jj, s_nn);
}
__syncthreads();
// rhs <- v (next round iterates the orthogonalized vectors).
for (int idx = tid; idx < N * N; idx += NT) {
Yb[idx] = Vb[idx];
}
}
__syncthreads();
// -------------------------------------------------------------------
// Final per-group MGS pass + R2 dependent-vector rescue (oracle 4b).
// -------------------------------------------------------------------
for (int jj = 0; jj < N; ++jj) {
if (smin[jj] != 3) continue;
const int k0 = sgstart[jj];
cta_project(Vb, jj, k0, jj);
int tries = 0;
while (s_rem < kDepTol && tries < 3) {
++tries;
if (tid == 0) ++s_rescues;
// Fresh hashed RHS into y column jj, orthogonalized vs the group.
for (int i = tid; i < N; i += NT) {
Yb[static_cast<size_t>(i) * N + jj] =
hash_uniform(batch, jj, i, 100u + tries);
}
__syncthreads();
cta_project(Yb, jj, k0, jj);
if (s_nn > 0.0f) cta_scale_column(Yb, jj, s_nn);
// Single-thread re-solve (rare path; oracle count 0): identical
// factors are recomputed; solution is max-normalized only (oracle
// solve_tridiag return convention).
if (tid == jj) {
float lg = 0.0f, mx = 0.0f;
thomas_solve_column<N>(sds, ses, sxs[jj], pivmin, Yb + jj, Vb + jj,
U0b + jj, U1b + jj, U2b + jj, &lg, &mx,
&guard_hits_local);
if (mx > 0.0f) {
for (int r = 0; r < N; ++r) {
Vb[static_cast<size_t>(r) * N + jj] /= mx;
}
}
srnd[jj] += 1;
}
__syncthreads();
cta_project(Vb, jj, k0, jj);
}
// v[j] = vec / (nn or 1)
if (s_nn > 0.0f) cta_scale_column(Vb, jj, s_nn);
}
__syncthreads();
// Stats: invit rounds max/sum, stagnant, rescues, guard hits.
atomicAdd(&s_guard, guard_hits_local);
int rmax = 0, rsum = 0, stag = 0;
if (tid < N) {
rmax = srnd[tid];
rsum = srnd[tid];
stag = sacc[tid] ? 0 : 1;
}
rmax = warp_imax(rmax);
rsum = warp_isum(rsum);
stag = warp_isum(stag);
if (lane == 0) {
partials[warp] = __int_as_float(rmax);
}
__syncthreads();
if (tid == 0) {
int m = 0;
for (int w = 0; w < kWarps; ++w) m = max(m, __float_as_int(partials[w]));
stats_ws[batch * kStatsSlots + 6] = m;
}
__syncthreads();
if (lane == 0) partials[warp] = __int_as_float(rsum);
__syncthreads();
if (tid == 0) {
int s = 0;
for (int w = 0; w < kWarps; ++w) s += __float_as_int(partials[w]);
stats_ws[batch * kStatsSlots + 7] = s;
}
__syncthreads();
if (lane == 0) partials[warp] = __int_as_float(stag);
__syncthreads();
if (tid == 0) {
int s = 0;
for (int w = 0; w < kWarps; ++w) s += __float_as_int(partials[w]);
stats_ws[batch * kStatsSlots + 8] = s;
stats_ws[batch * kStatsSlots + 9] = s_rescues;
stats_ws[batch * kStatsSlots + 10] = s_guard;
}
}
// ======================================================================
// Host side.
// ======================================================================
bool g_attr_set = false;
int device_optin_shared_bytes() {
int device = 0;
C10_CUDA_CHECK(cudaGetDevice(&device));
int value = 0;
C10_CUDA_CHECK(cudaDeviceGetAttribute(
&value, cudaDevAttrMaxSharedMemoryPerBlockOptin, device));
return value;
}
void ensure_kernel_attributes() {
// Only legal where the opt-in cap covers the request (B200 227 KB yes,
// GB10 99 KB no — the global-W variant runs there instead).
if (!g_attr_set &&
device_optin_shared_bytes() >= k1_dynamic_shared_bytes<176>()) {
C10_CUDA_CHECK(cudaFuncSetAttribute(
midn_tridiag_qform_kernel<176, true>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
k1_dynamic_shared_bytes<176>()));
g_attr_set = true;
}
}
void check_common(const torch::Tensor& a, const torch::Tensor& q,
const torch::Tensor& v, const torch::Tensor& lam,
const torch::Tensor& info, const torch::Tensor& stats,
int n) {
TORCH_CHECK(a.is_cuda() && q.is_cuda() && v.is_cuda() && lam.is_cuda() &&
info.is_cuda() && stats.is_cuda(),
"all tensors must be CUDA");
TORCH_CHECK(a.dtype() == torch::kFloat32 && q.dtype() == torch::kFloat32 &&
v.dtype() == torch::kFloat32 &&
lam.dtype() == torch::kFloat32,
"a, q, v, lam must be float32");
TORCH_CHECK(info.dtype() == torch::kInt32 && stats.dtype() == torch::kInt32,
"info, stats must be int32");
TORCH_CHECK(a.dim() == 3 && a.size(1) == n && a.size(2) == n,
"a must be [batch, ", n, ", ", n, "]");
const int64_t batch = a.size(0);
TORCH_CHECK(q.sizes() == a.sizes() && v.sizes() == a.sizes(),
"q, v must match a's shape");
TORCH_CHECK(lam.dim() == 2 && lam.size(0) == batch && lam.size(1) == n,
"lam must be [batch, ", n, "]");
TORCH_CHECK(info.dim() == 1 && info.size(0) == batch,
"info must be [batch]");
TORCH_CHECK(stats.dim() == 2 && stats.size(0) == batch &&
stats.size(1) == kStatsSlots,
"stats must be [batch, ", kStatsSlots, "]");
TORCH_CHECK(a.is_contiguous() && q.is_contiguous() && v.is_contiguous() &&
lam.is_contiguous() && info.is_contiguous() &&
stats.is_contiguous(),
"all tensors must be contiguous");
}
struct MidnWorkspaces {
torch::Tensor d, e, gamma, ds, es, w, xs, gstart, minits, pivmin, onenrm;
torch::Tensor y, u0, u1, u2;
};
MidnWorkspaces make_workspaces(const torch::Tensor& a, int n) {
const int64_t batch = a.size(0);
auto f32 = a.options();
auto i32 = a.options().dtype(torch::kInt32);
return {torch::empty({batch, n}, f32), // d
torch::empty({batch, n}, f32), // e
torch::empty({batch}, f32), // gamma
torch::empty({batch, n}, f32), // ds
torch::empty({batch, n}, f32), // es
torch::empty({batch, n}, f32), // w
torch::empty({batch, n}, f32), // xs
torch::empty({batch, n}, i32), // gstart
torch::empty({batch, n}, i32), // minits
torch::empty({batch}, f32), // pivmin
torch::empty({batch}, f32), // onenrm
torch::empty({batch, n, n}, f32), // y
torch::empty({batch, n, n}, f32), // u0
torch::empty({batch, n, n}, f32), // u1
torch::empty({batch, n, n}, f32)}; // u2
}
// w_mode: 0 auto (shared-W when the device supports it, n=176 only),
// 1 force shared-W (n=176 only), 2 force global-W.
template <int N>
void run_pipeline(const torch::Tensor& a, torch::Tensor& q, torch::Tensor& v,
torch::Tensor& lam, torch::Tensor& info,
torch::Tensor& stats, int w_mode) {
constexpr int NT = (N % 32 == 0) ? N : ((N / 32) + 1) * 32;
static_assert(NT >= N && NT % 32 == 0, "thread count invariant");
const int batch = static_cast<int>(a.size(0));
// Concrete type (not auto): the initializer is N-dependent, and a
// dependent-typed `ws` would break data_ptr<T> template lookup below.
MidnWorkspaces ws = make_workspaces(a, N);
bool shared_w = false;
if constexpr (N == 176) {
const bool fits = device_optin_shared_bytes() >= k1_dynamic_shared_bytes<N>();
shared_w = (w_mode == 1) || (w_mode == 0 && fits);
TORCH_CHECK(!shared_w || fits,
"shared-W forced but device opt-in shared is too small");
} else {
TORCH_CHECK(w_mode != 1, "shared-W variant only exists for n=176");
}
if constexpr (N == 176) {
if (shared_w) {
ensure_kernel_attributes();
midn_tridiag_qform_kernel<N, true>
<<<batch, kThreadsK1, k1_dynamic_shared_bytes<N>()>>>(
a.data_ptr<float>(), q.data_ptr<float>(),
ws.d.data_ptr<float>(), ws.e.data_ptr<float>(),
ws.gamma.data_ptr<float>(), info.data_ptr<int>(), batch);
} else {
midn_tridiag_qform_kernel<N, false><<<batch, kThreadsK1, 0>>>(
a.data_ptr<float>(), q.data_ptr<float>(), ws.d.data_ptr<float>(),
ws.e.data_ptr<float>(), ws.gamma.data_ptr<float>(),
info.data_ptr<int>(), batch);
}
} else {
midn_tridiag_qform_kernel<N, false><<<batch, kThreadsK1, 0>>>(
a.data_ptr<float>(), q.data_ptr<float>(), ws.d.data_ptr<float>(),
ws.e.data_ptr<float>(), ws.gamma.data_ptr<float>(),
info.data_ptr<int>(), batch);
}
midn_bisect_kernel<N, NT><<<batch, NT, 0>>>(
ws.d.data_ptr<float>(), ws.e.data_ptr<float>(),
ws.gamma.data_ptr<float>(), ws.ds.data_ptr<float>(),
ws.es.data_ptr<float>(), ws.w.data_ptr<float>(), lam.data_ptr<float>(),
ws.xs.data_ptr<float>(), ws.gstart.data_ptr<int>(),
ws.minits.data_ptr<int>(), ws.pivmin.data_ptr<float>(),
ws.onenrm.data_ptr<float>(), stats.data_ptr<int>(), batch);
midn_invit_kernel<N, NT><<<batch, NT, 0>>>(
ws.ds.data_ptr<float>(), ws.es.data_ptr<float>(),
ws.xs.data_ptr<float>(), ws.gstart.data_ptr<int>(),
ws.minits.data_ptr<int>(), ws.pivmin.data_ptr<float>(),
v.data_ptr<float>(), ws.y.data_ptr<float>(), ws.u0.data_ptr<float>(),
ws.u1.data_ptr<float>(), ws.u2.data_ptr<float>(), stats.data_ptr<int>(),
batch);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
} // namespace
void midn_eigh_pipeline_out(torch::Tensor a, torch::Tensor q, torch::Tensor v,
torch::Tensor lam, torch::Tensor info,
torch::Tensor stats, int64_t w_mode) {
const int n = static_cast<int>(a.size(1));
check_common(a, q, v, lam, info, stats, n);
const at::cuda::CUDAGuard guard(a.device());
if (n == 176) {
run_pipeline<176>(a, q, v, lam, info, stats, static_cast<int>(w_mode));
} else if (n == 352) {
run_pipeline<352>(a, q, v, lam, info, stats, static_cast<int>(w_mode));
} else {
TORCH_CHECK(false, "unsupported n (176 or 352): ", n);
}
}
namespace {
template <typename KernelT>
void push_attrs(std::vector<int64_t>& out, KernelT kernel) {
cudaFuncAttributes attributes{};
C10_CUDA_CHECK(cudaFuncGetAttributes(&attributes, kernel));
out.insert(out.end(), {static_cast<int64_t>(attributes.numRegs),
static_cast<int64_t>(attributes.localSizeBytes),
static_cast<int64_t>(attributes.sharedSizeBytes),
static_cast<int64_t>(attributes.maxThreadsPerBlock)});
}
} // namespace
std::vector<int64_t> midn_resource_attributes(int64_t n) {
// Rows of 4: [regs, local_bytes, static_shared_bytes, max_threads] for
// {K1 shared-W (n176 only, zeros otherwise), K1 global-W, K2, K3}.
std::vector<int64_t> out;
if (n == 176) {
ensure_kernel_attributes();
push_attrs(out, midn_tridiag_qform_kernel<176, true>);
push_attrs(out, midn_tridiag_qform_kernel<176, false>);
push_attrs(out, midn_bisect_kernel<176, 192>);
push_attrs(out, midn_invit_kernel<176, 192>);
} else if (n == 352) {
out.insert(out.end(), {0, 0, 0, 0});
push_attrs(out, midn_tridiag_qform_kernel<352, false>);
push_attrs(out, midn_bisect_kernel<352, 352>);
push_attrs(out, midn_invit_kernel<352, 352>);
} else {
TORCH_CHECK(false, "unsupported n (176 or 352): ", n);
}
return out;
}
int64_t midn_k1_dynamic_shared_bytes(int64_t n) {
if (n == 176) return k1_dynamic_shared_bytes<176>();
if (n == 352) return k1_dynamic_shared_bytes<352>();
TORCH_CHECK(false, "unsupported n (176 or 352): ", n);
}
int64_t midn_k1_shared_supported() {
return device_optin_shared_bytes() >= k1_dynamic_shared_bytes<176>() ? 1
: 0;
}
int64_t midn_stats_slots() { return kStatsSlots; }
} // namespace lane_h
namespace lane_i {
constexpr int kThreads = 256;
constexpr int kJacobiMaxSweeps = 30;
constexpr int kSecularMaxIters = 100;
constexpr float kEps32 = 1.1920928955078125e-07f;
constexpr double kEps64 = 2.220446049250313e-16;
// ---------------------------------------------------------------------------
// fused fp32 Jacobi (lower tree): one CTA per 128x128 block
// A in dynamic shared (pitch 129 to dodge bank conflicts), V in global.
// ---------------------------------------------------------------------------
__global__ void jacobi_block_kernel(
const float* __restrict__ d, // [F, m] adjusted diagonals
const float* __restrict__ e, // [F, m-1]
const short* __restrict__ sched, // [m-1, m/2, 2] tournament pairs
float* __restrict__ values, // [F, m]
float* __restrict__ vectors, // [F, m, m] V^T (rows = eigvecs)
int* __restrict__ info, // [F]
int m, int v_in_shared) {
extern __shared__ float smem[];
const int pitch = m + 1;
float* a = smem; // m * pitch
float* red = smem + m * pitch; // kThreads reduction buffer
const int f = blockIdx.x;
const int tid = threadIdx.x;
const int pairs = m / 2;
// V stored TRANSPOSED (row p = eigvec p): a column rotation touches
// two contiguous rows -> coalesced. B200 (227KB opt-in) keeps V^T in
// shared; GB10 (99KB) uses global.
float* vg = vectors + (size_t)f * m * m;
float* v = v_in_shared ? (red + kThreads) : vg;
// load tridiagonal block into shared, identity into V^T
for (int idx = tid; idx < m * m; idx += kThreads) {
const int r = idx / m, c = idx % m;
float val = 0.f;
if (r == c) val = d[f * m + r];
else if (r == c + 1) val = e[f * (m - 1) + c];
else if (c == r + 1) val = e[f * (m - 1) + r];
a[r * pitch + c] = val;
v[r * m + c] = (r == c) ? 1.f : 0.f;
}
__syncthreads();
// scale = max |a| (rotation threshold base)
float local = 0.f;
for (int idx = tid; idx < m * m; idx += kThreads) {
local = fmaxf(local, fabsf(a[(idx / m) * pitch + idx % m]));
}
red[tid] = local;
__syncthreads();
for (int wdt = kThreads / 2; wdt > 0; wdt >>= 1) {
if (tid < wdt) red[tid] = fmaxf(red[tid], red[tid + wdt]);
__syncthreads();
}
const float scale = (red[0] > 0.f) ? red[0] : 1.f;
const float tol = 16.f * kEps32 * scale;
const float floor_tol = 256.f * kEps32 * scale;
const float rot_tol = kEps32 * scale * 1e-3f;
__syncthreads();
__shared__ float cs_c[64], cs_s[64];
__shared__ int done_flag;
float prev_off = INFINITY;
bool converged = false;
for (int sweep = 0; sweep < kJacobiMaxSweeps; ++sweep) {
// max off-diagonal
local = 0.f;
for (int idx = tid; idx < m * m; idx += kThreads) {
const int r = idx / m, c = idx % m;
if (r != c) local = fmaxf(local, fabsf(a[r * pitch + c]));
}
red[tid] = local;
__syncthreads();
for (int wdt = kThreads / 2; wdt > 0; wdt >>= 1) {
if (tid < wdt) red[tid] = fmaxf(red[tid], red[tid + wdt]);
__syncthreads();
}
const float off = red[0];
__syncthreads();
if (tid == 0) {
done_flag = 0;
if (off <= tol || (off >= prev_off && off <= floor_tol)) done_flag = 1;
else if (off >= prev_off) done_flag = 2;
}
__syncthreads();
if (done_flag == 1) { converged = true; break; }
if (done_flag == 2) break;
prev_off = off;
for (int step = 0; step < m - 1; ++step) {
const short* sp = sched + (size_t)step * pairs * 2;
// rotation parameters (one thread per pair)
if (tid < pairs) {
const int p = sp[tid * 2], q = sp[tid * 2 + 1];
const float apq = a[p * pitch + q];
float c = 1.f, s = 0.f;
if (fabsf(apq) > rot_tol) {
const float theta = (a[q * pitch + q] - a[p * pitch + p])
/ (2.f * apq);
float t;
if (theta == 0.f) t = 1.f;
else t = copysignf(1.f, theta)
/ (fabsf(theta) + sqrtf(1.f + theta * theta));
c = rsqrtf(1.f + t * t);
s = t * c;
}
cs_c[tid] = c;
cs_s[tid] = s;
}
__syncthreads();
// row rotations: threads cover pairs x columns
for (int idx = tid; idx < pairs * m; idx += kThreads) {
const int pr = idx / m, col = idx % m;
const int p = sp[pr * 2], q = sp[pr * 2 + 1];
const float c = cs_c[pr], s = cs_s[pr];
const float xi = a[p * pitch + col], xj = a[q * pitch + col];
a[p * pitch + col] = c * xi - s * xj;
a[q * pitch + col] = s * xi + c * xj;
}
__syncthreads();
// column rotations on A
for (int idx = tid; idx < pairs * m; idx += kThreads) {
const int pr = idx / m, row = idx % m;
const int p = sp[pr * 2], q = sp[pr * 2 + 1];
const float c = cs_c[pr], s = cs_s[pr];
const float xi = a[row * pitch + p], xj = a[row * pitch + q];
a[row * pitch + p] = c * xi - s * xj;
a[row * pitch + q] = s * xi + c * xj;
}
// column rotations on V == row rotations on V^T (coalesced)
for (int idx = tid; idx < pairs * m; idx += kThreads) {
const int pr = idx / m, col = idx % m;
const int p = sp[pr * 2], q = sp[pr * 2 + 1];
const float c = cs_c[pr], s = cs_s[pr];
const float vi = v[p * m + col], vj = v[q * m + col];
v[p * m + col] = c * vi - s * vj;
v[q * m + col] = s * vi + c * vj;
}
__syncthreads();
}
}
if (tid == 0 && !converged) info[f] = 8;
__syncthreads();
for (int idx = tid; idx < m; idx += kThreads) {
values[f * m + idx] = a[idx * pitch + idx];
}
if (v_in_shared) {
for (int idx = tid; idx < m * m; idx += kThreads) {
vg[idx] = v[idx];
}
}
}
// ---------------------------------------------------------------------------
// DLAED2-class deflation sweep. One CTA per node; the sweep itself is
// sequential scalar code on thread 0 (O(w), negligible), tolerance
// reductions are parallel.
// ---------------------------------------------------------------------------
__global__ void deflate_kernel(
const float* __restrict__ poles, // [F, w] sorted ascending
const float* __restrict__ weights, // [F, w]
const float* __restrict__ rho, // [F]
double ctol, // tolerance multiplier (abs eps)
int half,
const int* __restrict__ srcblock, // [F, w] 0 = Q1 child, 1 = Q2
double* __restrict__ d_adj, // [F, w]
double* __restrict__ z_adj, // [F, w]
int* __restrict__ active, // [F, w] active lanes | defl tail
int* __restrict__ perm_cols, // [F, w] coltype-grouped | defl
int* __restrict__ counts, // [F, 5] K, n1, n2, n3, nrot
int* __restrict__ rot_idx, // [F, w, 2]
double* __restrict__ rot_cs, // [F, w, 2]
int w) {
extern __shared__ double dsm[];
double* dsh = dsm; // w poles
double* zsh = dsm + w; // w weights
double* red = dsm + 2 * w; // kThreads reduction
// reuse tails of shared as small int scratch
__shared__ int coltype_s[512];
__shared__ int defl_s[512];
const int f = blockIdx.x;
const int tid = threadIdx.x;
for (int j = tid; j < w; j += kThreads) {
dsh[j] = (double)poles[f * w + j];
zsh[j] = (double)weights[f * w + j];
coltype_s[j] = (srcblock[f * w + j] == 0) ? 1 : 3;
defl_s[j] = 0;
}
__syncthreads();
double local = 0.0;
for (int j = tid; j < w; j += kThreads) {
local = fmax(local, fmax(fabs(dsh[j]), fabs(zsh[j])));
}
red[tid] = local;
__syncthreads();
for (int wd = kThreads / 2; wd > 0; wd >>= 1) {
if (tid < wd) red[tid] = fmax(red[tid], red[tid + wd]);
__syncthreads();
}
const double tol = ctol * red[0];
const double rho_d = (double)rho[f];
__syncthreads();
if (tid == 0) {
int nrot = 0;
int pj = -1;
for (int j = 0; j < w; ++j) {
if (rho_d * fabs(zsh[j]) <= tol) { defl_s[j] = 1; continue; }
if (pj < 0) { pj = j; continue; }
const int nj = j;
const double s_val = zsh[pj], c_val = zsh[nj];
const double tau = hypot(c_val, s_val);
const double t = dsh[nj] - dsh[pj];
const double c = c_val / tau;
const double s = -s_val / tau;
if (fabs(t * c * s) <= tol) {
rot_idx[((size_t)f * w + nrot) * 2] = pj;
rot_idx[((size_t)f * w + nrot) * 2 + 1] = nj;
rot_cs[((size_t)f * w + nrot) * 2] = c;
rot_cs[((size_t)f * w + nrot) * 2 + 1] = s;
++nrot;
zsh[nj] = tau;
zsh[pj] = 0.0;
const double t_new = dsh[pj] * c * c + dsh[nj] * s * s;
dsh[nj] = dsh[pj] * s * s + dsh[nj] * c * c;
dsh[pj] = t_new;
defl_s[pj] = 1;
if (coltype_s[pj] != coltype_s[nj]) {
coltype_s[pj] = 2;
coltype_s[nj] = 2;
}
pj = nj;
} else {
pj = nj;
}
}
// pack: active (sorted-lane order), coltype-grouped columns, tail
int k = 0;
for (int j = 0; j < w; ++j) {
if (!defl_s[j]) active[f * w + k++] = j;
}
int t4 = k;
for (int j = 0; j < w; ++j) {
if (defl_s[j]) active[f * w + t4++] = j;
}
int pos = 0, n1 = 0, n2 = 0, n3 = 0;
for (int type = 1; type <= 3; ++type) {
for (int j = 0; j < w; ++j) {
if (!defl_s[j] && coltype_s[j] == type) {
perm_cols[f * w + pos++] = j;
if (type == 1) ++n1; else if (type == 2) ++n2; else ++n3;
}
}
}
for (int j = 0; j < w; ++j) {
if (defl_s[j]) perm_cols[f * w + pos++] = j;
}
counts[f * 5] = k;
counts[f * 5 + 1] = n1;
counts[f * 5 + 2] = n2;
counts[f * 5 + 3] = n3;
counts[f * 5 + 4] = nrot;
}
__syncthreads();
for (int j = tid; j < w; j += kThreads) {
d_adj[(size_t)f * w + j] = dsh[j];
z_adj[(size_t)f * w + j] = zsh[j];
}
}
// ---------------------------------------------------------------------------
// Givens rotations on child-basis columns (fp32 basis, fp64 c/s), applied
// sequentially per node (rotations chain), rows in parallel.
// ---------------------------------------------------------------------------
__global__ void apply_givens_kernel(
float* __restrict__ basis, // [F, w, w]
const int* __restrict__ rot_idx, // [F, w, 2]
const double* __restrict__ rot_cs, // [F, w, 2]
const int* __restrict__ counts, // [F, 5]
int w) {
const int f = blockIdx.x;
const int tid = threadIdx.x;
const int nrot = counts[f * 5 + 4];
float* b = basis + (size_t)f * w * w;
for (int r = 0; r < nrot; ++r) {
const int p = rot_idx[((size_t)f * w + r) * 2];
const int q = rot_idx[((size_t)f * w + r) * 2 + 1];
const float c = (float)rot_cs[((size_t)f * w + r) * 2];
const float s = (float)rot_cs[((size_t)f * w + r) * 2 + 1];
for (int row = tid; row < w; row += kThreads) {
const float xp = b[(size_t)row * w + p];
const float xq = b[(size_t)row * w + q];
b[(size_t)row * w + p] = c * xp + s * xq; // annihilates z_p
b[(size_t)row * w + q] = c * xq - s * xp;
}
__syncthreads();
}
}
// ---------------------------------------------------------------------------
// CTA-per-root secular solve. grid (w, F); CTA r solves root r of node
// blockIdx.y when r < K. 256 threads cooperate on every sum.
// ---------------------------------------------------------------------------
struct SumPack {
double psi, phi, dpsi, dphi;
};
constexpr int kSecThreads = 128;
__device__ __forceinline__ double block_reduce_sum(double val, double* red) {
const int tid = threadIdx.x;
red[tid] = val;
__syncthreads();
for (int wd = kThreads / 2; wd > 0; wd >>= 1) {
if (tid < wd) red[tid] += red[tid + wd];
__syncthreads();
}
const double out = red[0];
__syncthreads();
return out;
}
// Fused 4-quantity reduction (kSecThreads threads): warp shuffles then a
// single cross-warp combine -- ONE syncthreads ladder instead of six.
__device__ __forceinline__ void sec_reduce4(double& a, double& b,
double& c, double& d2,
double* red) {
const unsigned mask = 0xffffffffu;
for (int off = 16; off > 0; off >>= 1) {
a += __shfl_down_sync(mask, a, off);
b += __shfl_down_sync(mask, b, off);
c += __shfl_down_sync(mask, c, off);
d2 += __shfl_down_sync(mask, d2, off);
}
const int warp = threadIdx.x >> 5;
const int lane = threadIdx.x & 31;
const int nwarp = kSecThreads >> 5;
if (lane == 0) {
red[warp * 4] = a;
red[warp * 4 + 1] = b;
red[warp * 4 + 2] = c;
red[warp * 4 + 3] = d2;
}
__syncthreads();
a = b = c = d2 = 0.0;
for (int wd = 0; wd < nwarp; ++wd) {
a += red[wd * 4];
b += red[wd * 4 + 1];
c += red[wd * 4 + 2];
d2 += red[wd * 4 + 3];
}
__syncthreads();
}
__device__ __forceinline__ double sec_reduce1(double a, double* red) {
double z0 = 0.0, z1 = 0.0, z2 = 0.0;
sec_reduce4(a, z0, z1, z2, red);
return a;
}
// full f(x) = 1 + rho * sum z2/(d - x) evaluated cooperatively
__device__ double eval_secular_full(const double* d, const double* z2,
double rho, double x, int k,
double* red) {
double part = 0.0;
for (int j = threadIdx.x; j < k; j += kSecThreads) {
part += z2[j] / (d[j] - x);
}
return 1.0 + rho * sec_reduce1(part, red);
}
// psi/phi split sums. Sign structure (interlacing): every j <= r term
// is negative, every j > r term positive, so |psi| = -psi and
// |phi| = phi EXACTLY -- no separate absolute-value sums needed
// (dlaed4's own ERRETM structure).
__device__ SumPack eval_parts(const double* dsh, const double* z2,
double tau, int r, int origin, int k,
double* red) {
SumPack out;
double psi = 0, phi = 0, dpsi = 0, dphi = 0;
const double dorg = dsh[origin];
for (int j = threadIdx.x; j < k; j += kSecThreads) {
const double denom = (dsh[j] - dorg) - tau;
const double t = z2[j] / denom;
const double dt = t / denom;
if (j <= r) { psi += t; dpsi += dt; }
else { phi += t; dphi += dt; }
}
sec_reduce4(psi, phi, dpsi, dphi, red);
out.psi = psi;
out.phi = phi;
out.dpsi = dpsi;
out.dphi = dphi;
return out;
}
__global__ void secular_kernel(
const double* __restrict__ d_act, // [F, w] packed (first K valid)
const double* __restrict__ z_act, // [F, w]
const float* __restrict__ rho, // [F]
const int* __restrict__ counts, // [F, 5]
double stop_scale,
int* __restrict__ origins, // [F, w]
double* __restrict__ taus, // [F, w]
double* __restrict__ dlambda, // [F, w] root values
int* __restrict__ iters, // [F, w]
int* __restrict__ info, // [F]
int w) {
extern __shared__ double dsm[];
double* dsh = dsm; // w
double* z2 = dsm + w; // w
double* red = dsm + 2 * w; // kThreads
__shared__ double zsumsq;
const int f = blockIdx.y;
const int r = blockIdx.x;
const int k = counts[f * 5];
if (r >= k) return; // dead roots never iterate
const int tid = threadIdx.x;
for (int j = tid; j < k; j += kSecThreads) {
const double zj = z_act[(size_t)f * w + j];
dsh[j] = d_act[(size_t)f * w + j];
z2[j] = zj * zj;
}
__syncthreads();
double part = 0.0;
for (int j = tid; j < k; j += kSecThreads) part += z2[j];
const double sum_z2 = sec_reduce1(part, red);
if (tid == 0) zsumsq = sum_z2;
__syncthreads();
const double rho_d = (double)rho[f];
const double rhoinv = 1.0 / rho_d;
if (k == 1) {
if (tid == 0) {
origins[(size_t)f * w] = 0;
taus[(size_t)f * w] = rho_d * z2[0];
dlambda[(size_t)f * w] = dsh[0] + rho_d * z2[0];
iters[(size_t)f * w] = 0;
}
return;
}
const bool last = (r == k - 1);
int origin;
double lo, hi;
if (last) {
origin = k - 1;
lo = 0.0;
hi = rho_d * zsumsq;
} else {
const double mid = 0.5 * (dsh[r] + dsh[r + 1]);
const double fmid = eval_secular_full(dsh, z2, rho_d, mid, k, red);
origin = (fmid > 0.0) ? r : r + 1;
const double span = dsh[r + 1] - dsh[r];
if (origin == r) { lo = 0.0; hi = 0.5 * span; }
else { lo = -0.5 * span; hi = 0.0; }
}
const double dorg = dsh[origin];
// initial guess (2-pole quadratic through poles r, r+1)
double tau;
if (last) {
tau = fmin(hi, rho_d * (z2[k - 1] + z2[k - 2]));
if (!(lo < tau && tau < hi)) tau = 0.5 * (lo + hi);
} else {
const double del_r = dsh[r] - dorg;
const double del_r1 = dsh[r + 1] - dorg;
const double gp = 0.5 * (del_r + del_r1);
double cpart = 0.0;
for (int j = tid; j < k; j += kSecThreads) {
if (j != r && j != r + 1) {
cpart += z2[j] / ((dsh[j] - dorg) - gp);
}
}
const double c_const = rhoinv + sec_reduce1(cpart, red);
const double a_q = c_const;
const double b_q = -(c_const * (del_r + del_r1) + z2[r] + z2[r + 1]);
const double c_q = c_const * del_r * del_r1 + z2[r] * del_r1
+ z2[r + 1] * del_r;
if (a_q == 0.0) {
tau = (b_q != 0.0) ? (-c_q / b_q) : 0.5 * (lo + hi);
} else {
const double disc = fmax(b_q * b_q - 4.0 * a_q * c_q, 0.0);
const double sq = sqrt(disc);
const double root1 = (-b_q - copysign(sq, b_q)) / (2.0 * a_q);
const double root2 = (root1 != 0.0) ? (c_q / (a_q * root1)) : root1;
tau = (lo < root1 && root1 < hi) ? root1 : root2;
}
if (!(lo < tau && tau < hi)) tau = 0.5 * (lo + hi);
}
int it = 0;
int local_info = 0;
bool broke = false;
for (it = 1; it <= kSecularMaxIters; ++it) {
const SumPack sp = eval_parts(dsh, z2, tau, r, origin, k, red);
const double w_val = rhoinv + sp.psi + sp.phi;
const double dw = sp.dpsi + sp.dphi;
const double erretm = 8.0 * (fabs(sp.phi) + fabs(sp.psi))
+ 2.0 * rhoinv + fabs(tau) * dw;
if (fabs(w_val) <= stop_scale * erretm) { broke = true; break; }
if (w_val > 0.0) hi = fmin(hi, tau); else lo = fmax(lo, tau);
const int p_lo = last ? (k - 2) : r;
const int p_hi = last ? (k - 1) : (r + 1);
const double d_lo = (dsh[p_lo] - dorg) - tau;
const double d_hi = (dsh[p_hi] - dorg) - tau;
const double c_mid = w_val - d_lo * sp.dpsi - d_hi * sp.dphi;
const double a_mid = (d_lo + d_hi) * w_val - d_lo * d_hi * dw;
const double b_mid = d_lo * d_hi * w_val;
double eta;
if (c_mid == 0.0) {
if (a_mid == 0.0) eta = (dw > 0.0) ? (-w_val / dw) : 0.0;
else eta = b_mid / a_mid;
} else {
const double disc = fmax(a_mid * a_mid - 4.0 * b_mid * c_mid, 0.0);
const double sq = sqrt(disc);
if (a_mid <= 0.0) eta = (a_mid - sq) / (2.0 * c_mid);
else eta = 2.0 * b_mid / (a_mid + sq);
}
if (w_val * eta >= 0.0 && dw > 0.0) eta = -w_val / dw;
double nxt = tau + eta;
if (!(lo < nxt && nxt < hi)) {
nxt = 0.5 * (tau + ((w_val < 0.0) ? hi : lo));
}
if (nxt == tau) { broke = true; break; }
tau = nxt;
}
if (!broke) local_info = 5;
if (!isfinite(tau)) local_info = 5;
if (tid == 0) {
origins[(size_t)f * w + r] = origin;
taus[(size_t)f * w + r] = tau;
dlambda[(size_t)f * w + r] = dorg + tau;
iters[(size_t)f * w + r] = it;
if (local_info) atomicMax(info + f, local_info);
}
}
// ---------------------------------------------------------------------------
// Loewner updated weights. grid (w, F); CTA i computes zhat_i (i < K)
// via a cooperative multiplicative reduction over j.
// zhat_i^2 = -delta(i,i)/rho * prod_{j != i} delta(i,j) / (d_i - d_j)
// delta(i,j) = (d_i - d_origin_j) - tau_j (bit-identical recompute)
// ---------------------------------------------------------------------------
__global__ void loewner_weights_kernel(
const double* __restrict__ d_act, // [F, w]
const double* __restrict__ z_act, // [F, w]
const float* __restrict__ rho, // [F]
const int* __restrict__ counts, // [F, 5]
const int* __restrict__ origins, // [F, w]
const double* __restrict__ taus, // [F, w]
double* __restrict__ zhat, // [F, w]
int w) {
extern __shared__ double dsm[];
double* dsh = dsm; // w poles
double* org = dsm + w; // w origin pole values
double* tsh = dsm + 2 * w; // w taus
double* red = dsm + 3 * w; // kThreads (product reduce)
const int f = blockIdx.y;
const int i = blockIdx.x;
const int k = counts[f * 5];
if (i >= k) return;
const int tid = threadIdx.x;
for (int j = tid; j < k; j += kThreads) {
dsh[j] = d_act[(size_t)f * w + j];
tsh[j] = taus[(size_t)f * w + j];
}
__syncthreads();
for (int j = tid; j < k; j += kThreads) {
org[j] = dsh[origins[(size_t)f * w + j]];
}
__syncthreads();
const double di = dsh[i];
double part = 1.0;
for (int j = tid; j < k; j += kThreads) {
if (j != i) {
part *= ((di - org[j]) - tsh[j]) / (di - dsh[j]);
}
}
red[tid] = part;
__syncthreads();
for (int wd = kThreads / 2; wd > 0; wd >>= 1) {
if (tid < wd) red[tid] *= red[tid + wd];
__syncthreads();
}
if (tid == 0) {
const double delta_ii = (di - org[i]) - tsh[i];
const double prod = -delta_ii / (double)rho[f] * red[0];
zhat[(size_t)f * w + i] = copysign(sqrt(fabs(prod)),
z_act[(size_t)f * w + i]);
}
}
// ---------------------------------------------------------------------------
// Secular vector build. grid (w, F); CTA r builds column r (r < K):
// s_j = zhat_j / delta(j, r), normalized by rsqrt (one division class
// per entry; NO normalize-division pass). Writes fp32 into the FULL
// w x w secular matrix at rows active[j], column col_out[r]; deflated
// coordinate columns are written by the python glue.
// ---------------------------------------------------------------------------
__global__ void vector_build_kernel(
const double* __restrict__ d_act, // [F, w]
const double* __restrict__ zhat, // [F, w]
const int* __restrict__ counts, // [F, 5]
const int* __restrict__ origins, // [F, w]
const double* __restrict__ taus, // [F, w]
const int* __restrict__ active, // [F, w]
const int* __restrict__ col_out, // [F, w] output column of root r
float* __restrict__ s_full, // [F, w, w] zero-initialized
int w) {
extern __shared__ double dsm[];
double* ssh = dsm; // w column values
double* red = dsm + w; // kThreads
const int f = blockIdx.y;
const int r = blockIdx.x;
const int k = counts[f * 5];
if (r >= k) return;
const int tid = threadIdx.x;
const double dorg = d_act[(size_t)f * w + origins[(size_t)f * w + r]];
const double tau = taus[(size_t)f * w + r];
double part = 0.0;
for (int j = tid; j < k; j += kThreads) {
const double delta = (d_act[(size_t)f * w + j] - dorg) - tau;
const double s = zhat[(size_t)f * w + j] / delta;
ssh[j] = s;
part += s * s;
}
const double norm2 = block_reduce_sum(part, red);
const double scale = (norm2 > 0.0) ? rsqrt(norm2) : 0.0;
const int col = col_out[(size_t)f * w + r];
for (int j = tid; j < k; j += kThreads) {
s_full[((size_t)f * w + active[f * w + j]) * w + col]
= (float)(ssh[j] * scale);
}
}
// ---------------------------------------------------------------------------
// bindings
// ---------------------------------------------------------------------------
void jacobi_block_(torch::Tensor d, torch::Tensor e, torch::Tensor sched,
torch::Tensor values, torch::Tensor vectors,
torch::Tensor info, bool v_in_shared) {
CHECK_IN(d); CHECK_IN(e); CHECK_IN(sched);
CHECK_IN(values); CHECK_IN(vectors); CHECK_IN(info);
const int f = d.size(0);
const int m = d.size(1);
TORCH_CHECK(m == 128, "jacobi block width must be 128 (cs buffers)");
size_t shared = (size_t)m * (m + 1) * sizeof(float)
+ kThreads * sizeof(float);
if (v_in_shared) shared += (size_t)m * m * sizeof(float);
cudaFuncSetAttribute(jacobi_block_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize, shared);
jacobi_block_kernel<<<f, kThreads, shared>>>(
d.data_ptr<float>(), e.data_ptr<float>(), sched.data_ptr<short>(),
values.data_ptr<float>(), vectors.data_ptr<float>(),
info.data_ptr<int>(), m, v_in_shared ? 1 : 0);
}
void deflate_(torch::Tensor poles, torch::Tensor weights, torch::Tensor rho,
double ctol, int64_t half, torch::Tensor srcblock,
torch::Tensor d_adj, torch::Tensor z_adj,
torch::Tensor active, torch::Tensor perm_cols,
torch::Tensor counts, torch::Tensor rot_idx,
torch::Tensor rot_cs) {
CHECK_IN(poles); CHECK_IN(weights); CHECK_IN(rho); CHECK_IN(srcblock);
CHECK_IN(d_adj); CHECK_IN(z_adj); CHECK_IN(active); CHECK_IN(perm_cols);
CHECK_IN(counts); CHECK_IN(rot_idx); CHECK_IN(rot_cs);
const int f = poles.size(0);
const int w = poles.size(1);
TORCH_CHECK(w <= 512, "width");
const size_t shared = (2 * (size_t)w + kThreads) * sizeof(double);
deflate_kernel<<<f, kThreads, shared>>>(
poles.data_ptr<float>(), weights.data_ptr<float>(),
rho.data_ptr<float>(), ctol, (int)half, srcblock.data_ptr<int>(),
d_adj.data_ptr<double>(), z_adj.data_ptr<double>(),
active.data_ptr<int>(), perm_cols.data_ptr<int>(),
counts.data_ptr<int>(), rot_idx.data_ptr<int>(),
rot_cs.data_ptr<double>(), w);
}
void apply_givens_(torch::Tensor basis, torch::Tensor rot_idx,
torch::Tensor rot_cs, torch::Tensor counts) {
CHECK_IN(basis); CHECK_IN(rot_idx); CHECK_IN(rot_cs); CHECK_IN(counts);
const int f = basis.size(0);
const int w = basis.size(1);
apply_givens_kernel<<<f, kThreads, 0>>>(
basis.data_ptr<float>(), rot_idx.data_ptr<int>(),
rot_cs.data_ptr<double>(), counts.data_ptr<int>(), w);
}
void secular_(torch::Tensor d_act, torch::Tensor z_act, torch::Tensor rho,
torch::Tensor counts, double stop_scale,
torch::Tensor origins, torch::Tensor taus,
torch::Tensor dlambda, torch::Tensor iters,
torch::Tensor info) {
CHECK_IN(d_act); CHECK_IN(z_act); CHECK_IN(rho); CHECK_IN(counts);
CHECK_IN(origins); CHECK_IN(taus); CHECK_IN(dlambda); CHECK_IN(iters);
CHECK_IN(info);
const int f = d_act.size(0);
const int w = d_act.size(1);
const size_t shared = (2 * (size_t)w + kSecThreads) * sizeof(double);
cudaFuncSetAttribute(secular_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize, shared);
dim3 grid(w, f);
secular_kernel<<<grid, kSecThreads, shared>>>(
d_act.data_ptr<double>(), z_act.data_ptr<double>(),
rho.data_ptr<float>(), counts.data_ptr<int>(), stop_scale,
origins.data_ptr<int>(), taus.data_ptr<double>(),
dlambda.data_ptr<double>(), iters.data_ptr<int>(),
info.data_ptr<int>(), w);
}
void loewner_weights_(torch::Tensor d_act, torch::Tensor z_act,
torch::Tensor rho, torch::Tensor counts,
torch::Tensor origins, torch::Tensor taus,
torch::Tensor zhat) {
CHECK_IN(d_act); CHECK_IN(z_act); CHECK_IN(rho); CHECK_IN(counts);
CHECK_IN(origins); CHECK_IN(taus); CHECK_IN(zhat);
const int f = d_act.size(0);
const int w = d_act.size(1);
const size_t shared = (3 * (size_t)w + kThreads) * sizeof(double);
cudaFuncSetAttribute(loewner_weights_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize, shared);
dim3 grid(w, f);
loewner_weights_kernel<<<grid, kThreads, shared>>>(
d_act.data_ptr<double>(), z_act.data_ptr<double>(),
rho.data_ptr<float>(), counts.data_ptr<int>(),
origins.data_ptr<int>(), taus.data_ptr<double>(),
zhat.data_ptr<double>(), w);
}
void vector_build_(torch::Tensor d_act, torch::Tensor zhat,
torch::Tensor counts, torch::Tensor origins,
torch::Tensor taus, torch::Tensor active,
torch::Tensor col_out, torch::Tensor s_full) {
CHECK_IN(d_act); CHECK_IN(zhat); CHECK_IN(counts); CHECK_IN(origins);
CHECK_IN(taus); CHECK_IN(active); CHECK_IN(col_out); CHECK_IN(s_full);
const int f = d_act.size(0);
const int w = d_act.size(1);
const size_t shared = ((size_t)w + kThreads) * sizeof(double);
cudaFuncSetAttribute(vector_build_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize, shared);
dim3 grid(w, f);
vector_build_kernel<<<grid, kThreads, shared>>>(
d_act.data_ptr<double>(), zhat.data_ptr<double>(),
counts.data_ptr<int>(), origins.data_ptr<int>(),
taus.data_ptr<double>(), active.data_ptr<int>(),
col_out.data_ptr<int>(), s_full.data_ptr<float>(), w);
}
} // namespace lane_i
namespace lane_j {
namespace {
constexpr int kThreads = 256;
// --------------------------------------------------------------------
// basis build: one write pass. children [2F, half, half] (node f: left
// child 2f, right 2f+1), order [F, w] int32, safe [F] uint8.
// basis[f, i, j] = safe ? (i==j)
// : order[j] < half && i < half ? left [i, order[j]]
// : order[j] >= half && i >= half ? right[i-half, order[j]-half]
// : 0
// grid (w rows, F nodes), threads over j (coalesced writes; row-local
// gathered reads).
// --------------------------------------------------------------------
__global__ void build_basis_kernel(
const float* __restrict__ children, // [2F, half, half]
const int* __restrict__ order, // [F, w]
const unsigned char* __restrict__ safe, // [F]
float* __restrict__ basis, // [F, w, w]
int half, int w) {
__shared__ int ord[512];
const int f = blockIdx.y;
const int i = blockIdx.x;
const int tid = threadIdx.x;
for (int j = tid; j < w; j += kThreads) {
ord[j] = order[(size_t)f * w + j];
}
__syncthreads();
float* out_row = basis + ((size_t)f * w + i) * w;
if (safe[f]) {
for (int j = tid; j < w; j += kThreads) {
out_row[j] = (i == j) ? 1.0f : 0.0f;
}
return;
}
const float* left = children + (size_t)(2 * f) * half * half;
const float* right = children + (size_t)(2 * f + 1) * half * half;
for (int j = tid; j < w; j += kThreads) {
const int src = ord[j];
float value = 0.0f;
if (src < half) {
if (i < half) value = left[(size_t)i * half + src];
} else {
if (i >= half) value = right[(size_t)(i - half) * half + (src - half)];
}
out_row[j] = value;
}
}
// --------------------------------------------------------------------
// bp pack: bp[f, i, c] = c < count_f ? basis[f, row0 + i, perm[f, start_f
// + c]] : 0. branch 0 (top): row0 = 0, start = 0, count = n1 + n2.
// branch 1 (bot): row0 = half, start = n1, count = k - n1.
// grid (half, F), threads over c.
// --------------------------------------------------------------------
__global__ void pack_bp_kernel(
const float* __restrict__ basis, // [F, w, w]
const int* __restrict__ perm_cols, // [F, w]
const int* __restrict__ counts, // [F, 5]
float* __restrict__ bp, // [F, half, maxc]
int half, int w, int maxc, int branch) {
__shared__ int cols[512];
const int f = blockIdx.y;
const int i = blockIdx.x;
const int tid = threadIdx.x;
const int k = counts[f * 5];
const int n1 = counts[f * 5 + 1];
const int n2 = counts[f * 5 + 2];
const int start = branch == 0 ? 0 : n1;
const int count = branch == 0 ? (n1 + n2) : (k - n1);
for (int c = tid; c < maxc; c += kThreads) {
cols[c] = (c < count) ? perm_cols[(size_t)f * w + start + c] : -1;
}
__syncthreads();
const int row = (branch == 0 ? 0 : half) + i;
const float* b_row = basis + ((size_t)f * w + row) * w;
float* out_row = bp + ((size_t)f * half + i) * maxc;
for (int c = tid; c < maxc; c += kThreads) {
const int col = cols[c];
out_row[c] = (col >= 0) ? b_row[col] : 0.0f;
}
}
// --------------------------------------------------------------------
// inverse permutation of the coltype grouping: inv[f, perm[f, c]] = c.
// --------------------------------------------------------------------
__global__ void inv_perm_kernel(const int* __restrict__ perm_cols,
int* __restrict__ inv, int w) {
const int f = blockIdx.x;
const int tid = threadIdx.x;
for (int c = tid; c < w; c += kThreads) {
inv[(size_t)f * w + perm_cols[(size_t)f * w + c]] = c;
}
}
__device__ __forceinline__ double block_reduce_sum(double val, double* red) {
const int tid = threadIdx.x;
red[tid] = val;
__syncthreads();
for (int wd = kThreads / 2; wd > 0; wd >>= 1) {
if (tid < wd) red[tid] += red[tid + wd];
__syncthreads();
}
const double out = red[0];
__syncthreads();
return out;
}
// --------------------------------------------------------------------
// vector_build_packed: VERBATIM arithmetic of the rewrite lane's
// vector_build_kernel (same per-thread strided partials, same 256-thread
// reduction ladder, same fp32 rounding of ssh[j]*scale) — only the write
// targets change: instead of s_full[active[j], col], the value lands at
// its coltype-grouped PACKED row(s):
// c = inv_perm[active[j]]; c < n1+n2 -> s_top[c, col]
// c >= n1 -> s_bot[c - n1, col]
// (type-2 rows land in both, exactly the two GEMM branches' row sets).
// s_top/s_bot zero-initialized by the caller (padding rows stay zero).
// --------------------------------------------------------------------
__global__ void vector_build_packed_kernel(
const double* __restrict__ d_act, // [F, w]
const double* __restrict__ zhat, // [F, w]
const int* __restrict__ counts, // [F, 5]
const int* __restrict__ origins, // [F, w]
const double* __restrict__ taus, // [F, w]
const int* __restrict__ active, // [F, w]
const int* __restrict__ inv_perm, // [F, w]
const int* __restrict__ col_out, // [F, w]
float* __restrict__ s_top, // [F, tmax, w] zero-init
float* __restrict__ s_bot, // [F, bmax, w] zero-init
int w, int tmax, int bmax) {
extern __shared__ double dsm[];
double* ssh = dsm; // w column values
double* red = dsm + w; // kThreads
const int f = blockIdx.y;
const int r = blockIdx.x;
const int k = counts[f * 5];
if (r >= k) return;
const int tid = threadIdx.x;
const int n1 = counts[f * 5 + 1];
const int n12 = n1 + counts[f * 5 + 2];
const double dorg = d_act[(size_t)f * w + origins[(size_t)f * w + r]];
const double tau = taus[(size_t)f * w + r];
double part = 0.0;
for (int j = tid; j < k; j += kThreads) {
const double delta = (d_act[(size_t)f * w + j] - dorg) - tau;
const double s = zhat[(size_t)f * w + j] / delta;
ssh[j] = s;
part += s * s;
}
const double norm2 = block_reduce_sum(part, red);
const double scale = (norm2 > 0.0) ? rsqrt(norm2) : 0.0;
const int col = col_out[(size_t)f * w + r];
for (int j = tid; j < k; j += kThreads) {
const float value = (float)(ssh[j] * scale);
const int lane = active[f * w + j];
const int c = inv_perm[(size_t)f * w + lane];
if (c < n12) {
s_top[((size_t)f * tmax + c) * w + col] = value;
}
if (c >= n1) {
s_bot[((size_t)f * bmax + (c - n1)) * w + col] = value;
}
}
}
// --------------------------------------------------------------------
// deflated pass-through epilogue: out[f, :, col_out[k+t]] +=
// basis[f, :, perm[k+t]] for t < w - k (basis already Givens-rotated).
// grid (F), threads over rows for each t.
// --------------------------------------------------------------------
__global__ void defl_epilogue_kernel(
const float* __restrict__ basis, // [F, w, w]
const int* __restrict__ perm_cols, // [F, w]
const int* __restrict__ col_out, // [F, w]
const int* __restrict__ counts, // [F, 5]
float* __restrict__ out, // [F, w, w]
int w) {
const int f = blockIdx.x;
const int tid = threadIdx.x;
const int k = counts[f * 5];
const float* b = basis + (size_t)f * w * w;
float* o = out + (size_t)f * w * w;
for (int t = k; t < w; ++t) {
const int src = perm_cols[(size_t)f * w + t];
const int dst = col_out[(size_t)f * w + t];
for (int i = tid; i < w; i += kThreads) {
o[(size_t)i * w + dst] += b[(size_t)i * w + src];
}
}
}
// ====================================================================
// FINAL FP64-CORE ROUND (COST.md option A): warp-per-root reshapes.
// v1's loewner/vector_build run ONE 256-thread CTA PER ROOT and restage
// the node arrays into shared per CTA (~6 GB redundant loads at m512)
// plus an 8-step syncthreads reduction ladder. Here: CTA = 8 roots of
// ONE node (arrays staged once per 8 roots), warp-shuffle reductions,
// zero ladder barriers. Same formulas in fp64; reduction/product ORDER
// changes (strided-32 vs strided-256+ladder) -> corr-gated, not
// bitwise-gated (official margins 3000x).
// ====================================================================
constexpr unsigned kFullMask = 0xffffffffu;
constexpr int kRootsPerCta = 8;
__device__ __forceinline__ double warp_prod_d(double v) {
#pragma unroll
for (int off = 16; off > 0; off >>= 1) {
v *= __shfl_xor_sync(kFullMask, v, off);
}
return v;
}
__device__ __forceinline__ double warp_sum_d(double v) {
#pragma unroll
for (int off = 16; off > 0; off >>= 1) {
v += __shfl_xor_sync(kFullMask, v, off);
}
return v;
}
__global__ void loewner_warp_kernel(
const double* __restrict__ d_act, // [F, w]
const double* __restrict__ z_act, // [F, w]
const float* __restrict__ rho, // [F]
const int* __restrict__ counts, // [F, 5]
const int* __restrict__ origins, // [F, w]
const double* __restrict__ taus, // [F, w]
double* __restrict__ zhat, // [F, w]
int w) {
extern __shared__ double dsm[];
double* dsh = dsm; // w poles
double* org = dsm + w; // w origin pole values
double* tsh = dsm + 2 * w; // w taus
const int f = blockIdx.y;
const int k = counts[f * 5];
const int base = blockIdx.x * kRootsPerCta;
if (base >= k) return;
const int tid = threadIdx.x;
const int warp = tid >> 5;
const int lane = tid & 31;
for (int j = tid; j < k; j += kThreads) {
dsh[j] = d_act[(size_t)f * w + j];
tsh[j] = taus[(size_t)f * w + j];
}
__syncthreads();
for (int j = tid; j < k; j += kThreads) {
org[j] = dsh[origins[(size_t)f * w + j]];
}
__syncthreads();
const int i = base + warp;
if (i >= k) return;
const double di = dsh[i];
double part = 1.0;
for (int j = lane; j < k; j += 32) {
if (j != i) {
part *= ((di - org[j]) - tsh[j]) / (di - dsh[j]);
}
}
part = warp_prod_d(part);
if (lane == 0) {
const double delta_ii = (di - org[i]) - tsh[i];
const double prod = -delta_ii / (double)rho[f] * part;
zhat[(size_t)f * w + i] = copysign(sqrt(fabs(prod)),
z_act[(size_t)f * w + i]);
}
}
// vector build, warp-per-root, packed targets. Each warp keeps its
// root's s-column in shared (kRootsPerCta * w doubles, ~42 KB at w512
// with the staging arrays) so every element pays ONE fp64 division —
// the delta-recompute variant paid two and measured SLOWER on GB10's
// weak-DP regime (BUILD_LOG, v3 split round 1).
__global__ void vector_build_packed_warp_kernel(
const double* __restrict__ d_act, // [F, w]
const double* __restrict__ zhat, // [F, w]
const int* __restrict__ counts, // [F, 5]
const int* __restrict__ origins, // [F, w]
const double* __restrict__ taus, // [F, w]
const int* __restrict__ active, // [F, w]
const int* __restrict__ inv_perm, // [F, w]
const int* __restrict__ col_out, // [F, w]
float* __restrict__ s_top, // [F, tmax, w] zero-init
float* __restrict__ s_bot, // [F, bmax, w] zero-init
int w, int tmax, int bmax) {
extern __shared__ double dsm[];
double* dsh = dsm; // w poles
double* zsh = dsm + w; // w zhat
double* ssh = dsm + 2 * w; // kRootsPerCta * w s-columns
int* ipos = (int*)(dsm + (2 + kRootsPerCta) * w); // packed positions
const int f = blockIdx.y;
const int k = counts[f * 5];
const int base = blockIdx.x * kRootsPerCta;
if (base >= k) return;
const int tid = threadIdx.x;
const int warp = tid >> 5;
const int lane = tid & 31;
const int n1 = counts[f * 5 + 1];
const int n12 = n1 + counts[f * 5 + 2];
for (int j = tid; j < k; j += kThreads) {
dsh[j] = d_act[(size_t)f * w + j];
zsh[j] = zhat[(size_t)f * w + j];
ipos[j] = inv_perm[(size_t)f * w + active[(size_t)f * w + j]];
}
__syncthreads();
const int r = base + warp;
if (r >= k) return;
double* srow = ssh + (size_t)warp * w;
const double dorg = dsh[origins[(size_t)f * w + r]];
const double tau = taus[(size_t)f * w + r];
double part = 0.0;
for (int j = lane; j < k; j += 32) {
const double s = zsh[j] / ((dsh[j] - dorg) - tau);
srow[j] = s;
part += s * s;
}
const double norm2 = warp_sum_d(part);
const double scale = (norm2 > 0.0) ? rsqrt(norm2) : 0.0;
const int col = col_out[(size_t)f * w + r];
for (int j = lane; j < k; j += 32) {
const float value = (float)(srow[j] * scale);
const int c = ipos[j];
if (c < n12) {
s_top[((size_t)f * tmax + c) * w + col] = value;
}
if (c >= n1) {
s_bot[((size_t)f * bmax + (c - n1)) * w + col] = value;
}
}
}
// ====================================================================
// RECIPROCAL-SECULAR ROUND (COST.md): verbatim port of the rewrite
// lane's secular_kernel (CTA-per-root, erretm stop, middle-way rational)
// with ONE change in eval_parts: r = 1/denom; t = z2*r; dt = t*r —
// 1 fp64 division + 2 multiplies instead of 2 dependent divisions per
// pole per iteration. Rounding seam gated by recip_oracle.py + the
// full corr battery + iteration-count telemetry (prereg judges).
// ====================================================================
constexpr int kSecThreads = 128;
constexpr int kSecularMaxIters = 100;
struct SumPack {
double psi, phi, dpsi, dphi;
};
__device__ __forceinline__ void sec_reduce4(double& a, double& b,
double& c, double& d2,
double* red) {
const unsigned mask = 0xffffffffu;
for (int off = 16; off > 0; off >>= 1) {
a += __shfl_down_sync(mask, a, off);
b += __shfl_down_sync(mask, b, off);
c += __shfl_down_sync(mask, c, off);
d2 += __shfl_down_sync(mask, d2, off);
}
const int warp = threadIdx.x >> 5;
const int lane = threadIdx.x & 31;
const int nwarp = kSecThreads >> 5;
if (lane == 0) {
red[warp * 4] = a;
red[warp * 4 + 1] = b;
red[warp * 4 + 2] = c;
red[warp * 4 + 3] = d2;
}
__syncthreads();
a = b = c = d2 = 0.0;
for (int wd = 0; wd < nwarp; ++wd) {
a += red[wd * 4];
b += red[wd * 4 + 1];
c += red[wd * 4 + 2];
d2 += red[wd * 4 + 3];
}
__syncthreads();
}
__device__ __forceinline__ double sec_reduce1(double a, double* red) {
double z0 = 0.0, z1 = 0.0, z2 = 0.0;
sec_reduce4(a, z0, z1, z2, red);
return a;
}
__device__ double eval_secular_full(const double* d, const double* z2,
double rho, double x, int k,
double* red) {
double part = 0.0;
for (int j = threadIdx.x; j < k; j += kSecThreads) {
part += z2[j] / (d[j] - x);
}
return 1.0 + rho * sec_reduce1(part, red);
}
// v5: branch-free reciprocal for known-range denominators (finite,
// nonzero, normal — post-deflation gaps + interior tau). SASS-counted
// on sm_100: deletes the IEEE divide's FSETP guard + CAL'd slow path +
// BSSY/BSYNC convergence barriers (divtest/, COST.md v5 round).
// 2 Newton-Raphson steps from the rcp.approx.ftz.f64 seed: ~1-2 ulp64.
__device__ __forceinline__ double rcp_nr2(double b) {
double y;
asm("rcp.approx.ftz.f64 %0, %1;" : "=d"(y) : "d"(b));
double e = __fma_rn(-b, y, 1.0);
y = __fma_rn(y, e, y);
e = __fma_rn(-b, y, 1.0);
y = __fma_rn(y, e, y);
return y;
}
// THE restructure: one reciprocal, two multiplies (was t = z2/denom,
// dt = t/denom). Split/reduction order identical to the original.
// kNr: false = IEEE reciprocal (v4), true = rcp_nr2 (v5).
template <bool kNr>
__device__ SumPack eval_parts_recip(const double* dsh, const double* z2,
double tau, int r, int origin, int k,
double* red) {
SumPack out;
double psi = 0, phi = 0, dpsi = 0, dphi = 0;
const double dorg = dsh[origin];
for (int j = threadIdx.x; j < k; j += kSecThreads) {
const double denom = (dsh[j] - dorg) - tau;
const double rden = kNr ? rcp_nr2(denom) : (1.0 / denom);
const double t = z2[j] * rden;
const double dt = t * rden;
if (j <= r) { psi += t; dpsi += dt; }
else { phi += t; dphi += dt; }
}
sec_reduce4(psi, phi, dpsi, dphi, red);
out.psi = psi;
out.phi = phi;
out.dpsi = dpsi;
out.dphi = dphi;
return out;
}
template <bool kNr>
__global__ void secular_recip_kernel(
const double* __restrict__ d_act, // [F, w] packed (first K valid)
const double* __restrict__ z_act, // [F, w]
const float* __restrict__ rho, // [F]
const int* __restrict__ counts, // [F, 5]
double stop_scale,
int* __restrict__ origins, // [F, w]
double* __restrict__ taus, // [F, w]
double* __restrict__ dlambda, // [F, w]
int* __restrict__ iters, // [F, w]
int* __restrict__ info, // [F]
int w) {
extern __shared__ double dsm[];
double* dsh = dsm; // w
double* z2 = dsm + w; // w
double* red = dsm + 2 * w; // kSecThreads
__shared__ double zsumsq;
const int f = blockIdx.y;
const int r = blockIdx.x;
const int k = counts[f * 5];
if (r >= k) return;
const int tid = threadIdx.x;
for (int j = tid; j < k; j += kSecThreads) {
const double zj = z_act[(size_t)f * w + j];
dsh[j] = d_act[(size_t)f * w + j];
z2[j] = zj * zj;
}
__syncthreads();
double part = 0.0;
for (int j = tid; j < k; j += kSecThreads) part += z2[j];
const double sum_z2 = sec_reduce1(part, red);
if (tid == 0) zsumsq = sum_z2;
__syncthreads();
const double rho_d = (double)rho[f];
const double rhoinv = 1.0 / rho_d;
if (k == 1) {
if (tid == 0) {
origins[(size_t)f * w] = 0;
taus[(size_t)f * w] = rho_d * z2[0];
dlambda[(size_t)f * w] = dsh[0] + rho_d * z2[0];
iters[(size_t)f * w] = 0;
}
return;
}
const bool last = (r == k - 1);
int origin;
double lo, hi;
if (last) {
origin = k - 1;
lo = 0.0;
hi = rho_d * zsumsq;
} else {
const double mid = 0.5 * (dsh[r] + dsh[r + 1]);
const double fmid = eval_secular_full(dsh, z2, rho_d, mid, k, red);
origin = (fmid > 0.0) ? r : r + 1;
const double span = dsh[r + 1] - dsh[r];
if (origin == r) { lo = 0.0; hi = 0.5 * span; }
else { lo = -0.5 * span; hi = 0.0; }
}
const double dorg = dsh[origin];
double tau;
if (last) {
tau = fmin(hi, rho_d * (z2[k - 1] + z2[k - 2]));
if (!(lo < tau && tau < hi)) tau = 0.5 * (lo + hi);
} else {
const double del_r = dsh[r] - dorg;
const double del_r1 = dsh[r + 1] - dorg;
const double gp = 0.5 * (del_r + del_r1);
double cpart = 0.0;
for (int j = tid; j < k; j += kSecThreads) {
if (j != r && j != r + 1) {
cpart += z2[j] / ((dsh[j] - dorg) - gp);
}
}
const double c_const = rhoinv + sec_reduce1(cpart, red);
const double a_q = c_const;
const double b_q = -(c_const * (del_r + del_r1) + z2[r] + z2[r + 1]);
const double c_q = c_const * del_r * del_r1 + z2[r] * del_r1
+ z2[r + 1] * del_r;
if (a_q == 0.0) {
tau = (b_q != 0.0) ? (-c_q / b_q) : 0.5 * (lo + hi);
} else {
const double disc = fmax(b_q * b_q - 4.0 * a_q * c_q, 0.0);
const double sq = sqrt(disc);
const double root1 = (-b_q - copysign(sq, b_q)) / (2.0 * a_q);
const double root2 = (root1 != 0.0) ? (c_q / (a_q * root1)) : root1;
tau = (lo < root1 && root1 < hi) ? root1 : root2;
}
if (!(lo < tau && tau < hi)) tau = 0.5 * (lo + hi);
}
int it = 0;
int local_info = 0;
bool broke = false;
for (it = 1; it <= kSecularMaxIters; ++it) {
const SumPack sp = eval_parts_recip<kNr>(dsh, z2, tau, r, origin, k,
red);
const double w_val = rhoinv + sp.psi + sp.phi;
const double dw = sp.dpsi + sp.dphi;
const double erretm = 8.0 * (fabs(sp.phi) + fabs(sp.psi))
+ 2.0 * rhoinv + fabs(tau) * dw;
if (fabs(w_val) <= stop_scale * erretm) { broke = true; break; }
if (w_val > 0.0) hi = fmin(hi, tau); else lo = fmax(lo, tau);
const int p_lo = last ? (k - 2) : r;
const int p_hi = last ? (k - 1) : (r + 1);
const double d_lo = (dsh[p_lo] - dorg) - tau;
const double d_hi = (dsh[p_hi] - dorg) - tau;
const double c_mid = w_val - d_lo * sp.dpsi - d_hi * sp.dphi;
const double a_mid = (d_lo + d_hi) * w_val - d_lo * d_hi * dw;
const double b_mid = d_lo * d_hi * w_val;
double eta;
if (c_mid == 0.0) {
if (a_mid == 0.0) eta = (dw > 0.0) ? (-w_val / dw) : 0.0;
else eta = b_mid / a_mid;
} else {
const double disc = fmax(a_mid * a_mid - 4.0 * b_mid * c_mid, 0.0);
const double sq = sqrt(disc);
if (a_mid <= 0.0) eta = (a_mid - sq) / (2.0 * c_mid);
else eta = 2.0 * b_mid / (a_mid + sq);
}
if (w_val * eta >= 0.0 && dw > 0.0) eta = -w_val / dw;
double nxt = tau + eta;
if (!(lo < nxt && nxt < hi)) {
nxt = 0.5 * (tau + ((w_val < 0.0) ? hi : lo));
}
if (nxt == tau) { broke = true; break; }
tau = nxt;
}
if (!broke) local_info = 5;
if (!isfinite(tau)) local_info = 5;
if (tid == 0) {
origins[(size_t)f * w + r] = origin;
taus[(size_t)f * w + r] = tau;
dlambda[(size_t)f * w + r] = dorg + tau;
iters[(size_t)f * w + r] = it;
if (local_info) atomicMax(info + f, local_info);
}
}
} // namespace
void build_basis_(torch::Tensor children, torch::Tensor order,
torch::Tensor safe, torch::Tensor basis, int64_t half) {
CHECK_IN(children); CHECK_IN(order); CHECK_IN(safe); CHECK_IN(basis);
const int f = basis.size(0);
const int w = basis.size(1);
TORCH_CHECK(w <= 512, "width");
TORCH_CHECK(children.size(0) == 2 * f);
TORCH_CHECK(order.dtype() == torch::kInt32);
TORCH_CHECK(safe.dtype() == torch::kUInt8);
dim3 grid(w, f);
build_basis_kernel<<<grid, kThreads, 0>>>(
children.data_ptr<float>(), order.data_ptr<int>(),
safe.data_ptr<unsigned char>(), basis.data_ptr<float>(),
(int)half, w);
}
void pack_bp_(torch::Tensor basis, torch::Tensor perm_cols,
torch::Tensor counts, torch::Tensor bp, int64_t half,
int64_t branch) {
CHECK_IN(basis); CHECK_IN(perm_cols); CHECK_IN(counts); CHECK_IN(bp);
const int f = basis.size(0);
const int w = basis.size(1);
const int maxc = bp.size(2);
TORCH_CHECK(bp.size(0) == f && bp.size(1) == half);
TORCH_CHECK(maxc <= 512, "maxc");
dim3 grid((int)half, f);
pack_bp_kernel<<<grid, kThreads, 0>>>(
basis.data_ptr<float>(), perm_cols.data_ptr<int>(),
counts.data_ptr<int>(), bp.data_ptr<float>(), (int)half, w, maxc,
(int)branch);
}
void inv_perm_(torch::Tensor perm_cols, torch::Tensor inv) {
CHECK_IN(perm_cols); CHECK_IN(inv);
const int f = perm_cols.size(0);
const int w = perm_cols.size(1);
inv_perm_kernel<<<f, kThreads, 0>>>(
perm_cols.data_ptr<int>(), inv.data_ptr<int>(), w);
}
void vector_build_packed_(torch::Tensor d_act, torch::Tensor zhat,
torch::Tensor counts, torch::Tensor origins,
torch::Tensor taus, torch::Tensor active,
torch::Tensor inv, torch::Tensor col_out,
torch::Tensor s_top, torch::Tensor s_bot) {
CHECK_IN(d_act); CHECK_IN(zhat); CHECK_IN(counts); CHECK_IN(origins);
CHECK_IN(taus); CHECK_IN(active); CHECK_IN(inv); CHECK_IN(col_out);
CHECK_IN(s_top); CHECK_IN(s_bot);
const int f = d_act.size(0);
const int w = d_act.size(1);
const int tmax = s_top.size(1);
const int bmax = s_bot.size(1);
const size_t shared = ((size_t)w + kThreads) * sizeof(double);
cudaFuncSetAttribute(vector_build_packed_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize, shared);
dim3 grid(w, f);
vector_build_packed_kernel<<<grid, kThreads, shared>>>(
d_act.data_ptr<double>(), zhat.data_ptr<double>(),
counts.data_ptr<int>(), origins.data_ptr<int>(),
taus.data_ptr<double>(), active.data_ptr<int>(),
inv.data_ptr<int>(), col_out.data_ptr<int>(),
s_top.data_ptr<float>(), s_bot.data_ptr<float>(), w, tmax, bmax);
}
void defl_epilogue_(torch::Tensor basis, torch::Tensor perm_cols,
torch::Tensor col_out, torch::Tensor counts,
torch::Tensor out) {
CHECK_IN(basis); CHECK_IN(perm_cols); CHECK_IN(col_out);
CHECK_IN(counts); CHECK_IN(out);
const int f = basis.size(0);
const int w = basis.size(1);
defl_epilogue_kernel<<<f, kThreads, 0>>>(
basis.data_ptr<float>(), perm_cols.data_ptr<int>(),
col_out.data_ptr<int>(), counts.data_ptr<int>(),
out.data_ptr<float>(), w);
}
void loewner_warp_(torch::Tensor d_act, torch::Tensor z_act,
torch::Tensor rho, torch::Tensor counts,
torch::Tensor origins, torch::Tensor taus,
torch::Tensor zhat) {
CHECK_IN(d_act); CHECK_IN(z_act); CHECK_IN(rho); CHECK_IN(counts);
CHECK_IN(origins); CHECK_IN(taus); CHECK_IN(zhat);
const int f = d_act.size(0);
const int w = d_act.size(1);
const size_t shared = 3 * (size_t)w * sizeof(double);
dim3 grid((w + kRootsPerCta - 1) / kRootsPerCta, f);
loewner_warp_kernel<<<grid, kThreads, shared>>>(
d_act.data_ptr<double>(), z_act.data_ptr<double>(),
rho.data_ptr<float>(), counts.data_ptr<int>(),
origins.data_ptr<int>(), taus.data_ptr<double>(),
zhat.data_ptr<double>(), w);
}
void vector_build_packed_warp_(torch::Tensor d_act, torch::Tensor zhat,
torch::Tensor counts, torch::Tensor origins,
torch::Tensor taus, torch::Tensor active,
torch::Tensor inv, torch::Tensor col_out,
torch::Tensor s_top, torch::Tensor s_bot) {
CHECK_IN(d_act); CHECK_IN(zhat); CHECK_IN(counts); CHECK_IN(origins);
CHECK_IN(taus); CHECK_IN(active); CHECK_IN(inv); CHECK_IN(col_out);
CHECK_IN(s_top); CHECK_IN(s_bot);
const int f = d_act.size(0);
const int w = d_act.size(1);
const int tmax = s_top.size(1);
const int bmax = s_bot.size(1);
// shared: (2 + kRootsPerCta) * w doubles + w ints
const size_t shared = ((2 + kRootsPerCta) * (size_t)w) * sizeof(double)
+ ((size_t)w + 1) * sizeof(int);
cudaFuncSetAttribute(vector_build_packed_warp_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize, shared);
dim3 grid((w + kRootsPerCta - 1) / kRootsPerCta, f);
vector_build_packed_warp_kernel<<<grid, kThreads, shared>>>(
d_act.data_ptr<double>(), zhat.data_ptr<double>(),
counts.data_ptr<int>(), origins.data_ptr<int>(),
taus.data_ptr<double>(), active.data_ptr<int>(),
inv.data_ptr<int>(), col_out.data_ptr<int>(),
s_top.data_ptr<float>(), s_bot.data_ptr<float>(), w, tmax, bmax);
}
template <bool kNr>
void secular_recip_impl(torch::Tensor& d_act, torch::Tensor& z_act,
torch::Tensor& rho, torch::Tensor& counts,
double stop_scale, torch::Tensor& origins,
torch::Tensor& taus, torch::Tensor& dlambda,
torch::Tensor& iters, torch::Tensor& info) {
CHECK_IN(d_act); CHECK_IN(z_act); CHECK_IN(rho); CHECK_IN(counts);
CHECK_IN(origins); CHECK_IN(taus); CHECK_IN(dlambda); CHECK_IN(iters);
CHECK_IN(info);
const int f = d_act.size(0);
const int w = d_act.size(1);
const size_t shared = (2 * (size_t)w + kSecThreads) * sizeof(double);
cudaFuncSetAttribute(secular_recip_kernel<kNr>,
cudaFuncAttributeMaxDynamicSharedMemorySize, shared);
dim3 grid(w, f);
secular_recip_kernel<kNr><<<grid, kSecThreads, shared>>>(
d_act.data_ptr<double>(), z_act.data_ptr<double>(),
rho.data_ptr<float>(), counts.data_ptr<int>(), stop_scale,
origins.data_ptr<int>(), taus.data_ptr<double>(),
dlambda.data_ptr<double>(), iters.data_ptr<int>(),
info.data_ptr<int>(), w);
}
void secular_recip_(torch::Tensor d_act, torch::Tensor z_act,
torch::Tensor rho, torch::Tensor counts,
double stop_scale, torch::Tensor origins,
torch::Tensor taus, torch::Tensor dlambda,
torch::Tensor iters, torch::Tensor info) {
secular_recip_impl<false>(d_act, z_act, rho, counts, stop_scale,
origins, taus, dlambda, iters, info);
}
void secular_nr_(torch::Tensor d_act, torch::Tensor z_act,
torch::Tensor rho, torch::Tensor counts,
double stop_scale, torch::Tensor origins,
torch::Tensor taus, torch::Tensor dlambda,
torch::Tensor iters, torch::Tensor info) {
secular_recip_impl<true>(d_act, z_act, rho, counts, stop_scale,
origins, taus, dlambda, iters, info);
}
} // namespace lane_j
namespace lane_k {
namespace {
constexpr float kEps32 = 1.1920929e-07f;
constexpr float kSafmin32 = 1.1754944e-38f;
constexpr float kOrtol = 1e-3f; // DSTEIN grouping (scaled units)
constexpr float kPertol = 10.0f * kEps32; // DSTEIN shift separation
constexpr float kLog2Gtol = 9.9657842847f; // log2(1e3) growth acceptance
constexpr float kXmax = 1e15f; // SLAGTS dynamic-rescale bound
constexpr float kDepTol = 0.1f; // dependent-vector rescue bound
constexpr float kLog2Tiny30 = -99.65784285f; // log2(1e-30)
constexpr int kMaxIts = 5; // invit rounds cap
constexpr int kBisectCap = 64; // bisection pass cap
constexpr int kStatsSlots = 16;
constexpr unsigned kFullMask = 0xffffffffu;
__device__ __forceinline__ float warp_sum(float value) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
value += __shfl_xor_sync(kFullMask, value, offset);
}
return value;
}
__device__ __forceinline__ float warp_max(float value) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
value = fmaxf(value, __shfl_xor_sync(kFullMask, value, offset));
}
return value;
}
__device__ __forceinline__ float warp_min(float value) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
value = fminf(value, __shfl_xor_sync(kFullMask, value, offset));
}
return value;
}
__device__ __forceinline__ int warp_imax(int value) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
value = max(value, __shfl_xor_sync(kFullMask, value, offset));
}
return value;
}
__device__ __forceinline__ int warp_isum(int value) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
value += __shfl_xor_sync(kFullMask, value, offset);
}
return value;
}
// splitmix32-style hash -> uniform [-1, 1) (mid-n deviation D2, accepted).
__device__ __forceinline__ float hash_uniform(uint32_t b, uint32_t j,
uint32_t i, uint32_t t) {
uint32_t x = (b * 0x9E3779B9u) ^ (j * 0x85EBCA6Bu) ^ (i * 0xC2B2AE35u) ^
(t * 0x27D4EB2Fu) ^ 0xB5297A4Du;
x ^= x >> 16;
x *= 0x7FEB352Du;
x ^= x >> 15;
x *= 0x846CA68Bu;
x ^= x >> 16;
return static_cast<float>(x >> 8) * (2.0f / 16777216.0f) - 1.0f;
}
// ======================================================================
// KB: Sturm-count bisection for all roots + grouping epilogue.
// Verbatim mid-n K2 with: d/e taken directly (e has N-1 entries),
// finite screen (info=1), gamma removed (lam = w * onenrm).
// ======================================================================
template <int N, int NT>
__global__ void __launch_bounds__(NT)
lower_bisect_kernel(const float* __restrict__ d_in,
const float* __restrict__ e_in,
float* __restrict__ ds_ws,
float* __restrict__ es_ws,
float* __restrict__ w_ws,
float* __restrict__ lam_out,
float* __restrict__ xs_ws,
int* __restrict__ gstart_ws,
int* __restrict__ minits_ws,
float* __restrict__ pivmin_ws,
int* __restrict__ info,
int* __restrict__ stats_ws,
int batch_count) {
constexpr int kWarps = NT / 32;
__shared__ float sds[N];
__shared__ float ses[N];
__shared__ float se2[N];
__shared__ float sw[N];
__shared__ int sgstart[N];
__shared__ int sminits[N];
__shared__ float partials[kWarps];
__shared__ int ipartials[kWarps];
__shared__ float s_onenrm, s_pivmin, s_glo, s_ghi;
__shared__ int s_bad;
const int batch = blockIdx.x;
if (batch >= batch_count) {
return;
}
const int tid = threadIdx.x;
const int warp = tid >> 5;
const int lane = tid & 31;
if (tid == 0) s_bad = 0;
__syncthreads();
int local_bad = 0;
for (int i = tid; i < N; i += NT) {
const float dv = d_in[batch * N + i];
const float ev = (i < N - 1) ? e_in[batch * (N - 1) + i] : 0.0f;
sds[i] = dv;
ses[i] = ev;
local_bad |= !isfinite(dv);
local_bad |= !isfinite(ev);
}
if (local_bad) atomicOr(&s_bad, 1);
__syncthreads();
if (s_bad) {
// fail-closed degenerate emit: lam = 0, scaled workspaces zeroed;
// the invit kernel sees info != 0 and writes Q = I.
for (int i = tid; i < N; i += NT) {
lam_out[batch * N + i] = 0.0f;
w_ws[batch * N + i] = 0.0f;
xs_ws[batch * N + i] = 0.0f;
ds_ws[batch * N + i] = 0.0f;
es_ws[batch * N + i] = 0.0f;
gstart_ws[batch * N + i] = i;
minits_ws[batch * N + i] = 2;
}
if (tid == 0) {
info[batch] = 1;
pivmin_ws[batch] = kSafmin32;
}
return;
}
// onenrm = max_i(|d_i| + |e_{i-1}| + |e_i|), clamped >= safmin.
float lmax = 0.0f;
for (int i = tid; i < N; i += NT) {
float t = fabsf(sds[i]);
if (i > 0) t += fabsf(ses[i - 1]);
if (i < N - 1) t += fabsf(ses[i]);
lmax = fmaxf(lmax, t);
}
lmax = warp_max(lmax);
if (lane == 0) partials[warp] = lmax;
__syncthreads();
if (tid == 0) {
float m = 0.0f;
for (int w = 0; w < kWarps; ++w) m = fmaxf(m, partials[w]);
s_onenrm = fmaxf(m, kSafmin32);
}
__syncthreads();
const float onenrm = s_onenrm;
for (int i = tid; i < N; i += NT) {
const float dv = sds[i] / onenrm;
const float ev = ses[i] / onenrm;
sds[i] = dv;
ses[i] = ev;
se2[i] = ev * ev;
}
__syncthreads();
// pivmin + Gershgorin brackets (three reductions).
float le2 = 0.0f;
float lglo = FLT_MAX;
float lghi = -FLT_MAX;
for (int i = tid; i < N; i += NT) {
if (i < N - 1) le2 = fmaxf(le2, se2[i]);
float r = 0.0f;
if (i > 0) r += fabsf(ses[i - 1]);
if (i < N - 1) r += fabsf(ses[i]);
lglo = fminf(lglo, sds[i] - r);
lghi = fmaxf(lghi, sds[i] + r);
}
le2 = warp_max(le2);
if (lane == 0) partials[warp] = le2;
__syncthreads();
if (tid == 0) {
float m = 0.0f;
for (int w = 0; w < kWarps; ++w) m = fmaxf(m, partials[w]);
s_pivmin = fmaxf(m * kSafmin32, kSafmin32);
}
__syncthreads();
lglo = warp_min(lglo);
if (lane == 0) partials[warp] = lglo;
__syncthreads();
if (tid == 0) {
float m = FLT_MAX;
for (int w = 0; w < kWarps; ++w) m = fminf(m, partials[w]);
s_glo = m;
}
__syncthreads();
lghi = warp_max(lghi);
if (lane == 0) partials[warp] = lghi;
__syncthreads();
if (tid == 0) {
float m = -FLT_MAX;
for (int w = 0; w < kWarps; ++w) m = fmaxf(m, partials[w]);
const float width = fmaxf(m - s_glo, 1.0f) * kEps32;
s_ghi = m + width;
s_glo = s_glo - width;
}
__syncthreads();
// Per-root bisection: thread k refines eigenvalue k (0-based ascending).
const float pivmin = s_pivmin;
int iters = 0;
if (tid < N) {
float lo = s_glo;
float hi = s_ghi;
for (int pass = 0; pass < kBisectCap; ++pass) {
const float tol =
2.0f * kEps32 * fmaxf(fabsf(lo), fabsf(hi)) + 2.0f * pivmin;
if (hi - lo <= tol) break;
const float mid = 0.5f * (lo + hi);
float qv = sds[0] - mid;
if (fabsf(qv) < pivmin) qv = -pivmin;
int cnt = qv < 0.0f;
for (int i = 1; i < N; ++i) {
qv = sds[i] - mid - se2[i - 1] / qv;
if (fabsf(qv) < pivmin) qv = -pivmin;
cnt += qv < 0.0f;
}
if (cnt >= tid + 1) {
hi = mid;
} else {
lo = mid;
}
++iters;
}
sw[tid] = 0.5f * (lo + hi);
}
int imax = warp_imax(iters);
int isum = warp_isum(iters);
if (lane == 0) ipartials[warp] = imax;
__syncthreads();
if (tid == 0) {
int m = 0;
for (int w = 0; w < kWarps; ++w) m = max(m, ipartials[w]);
stats_ws[batch * kStatsSlots + 0] = m;
}
__syncthreads();
if (lane == 0) ipartials[warp] = isum;
__syncthreads();
// Thread-0 epilogue: monotone belt, grouping, minits, perturbed shifts.
if (tid == 0) {
int s = 0;
for (int w = 0; w < kWarps; ++w) s += ipartials[w];
stats_ws[batch * kStatsSlots + 1] = s;
for (int j = 1; j < N; ++j) sw[j] = fmaxf(sw[j], sw[j - 1]);
int start = 0;
int ngroups = 0, maxg = 0, sumk2 = 0, grouped = 0;
for (int j = 1; j <= N; ++j) {
const bool same = (j < N) && (sw[j] - sw[j - 1] <= kOrtol);
if (!same) {
const int sz = j - start;
++ngroups;
maxg = max(maxg, sz);
if (sz > 1) {
sumk2 += sz * sz;
grouped += sz;
}
for (int t = start; t < j; ++t) {
sgstart[t] = start;
sminits[t] = sz > 1 ? 3 : 2; // mid-n oracle FINDING #1
}
start = j;
}
}
stats_ws[batch * kStatsSlots + 2] = ngroups;
stats_ws[batch * kStatsSlots + 3] = maxg;
stats_ws[batch * kStatsSlots + 4] = sumk2;
stats_ws[batch * kStatsSlots + 5] = grouped;
// R1 perturbed shifts, cumulative within groups (DSTEIN standard).
float prev = sw[0];
xs_ws[batch * N + 0] = prev;
for (int j = 1; j < N; ++j) {
float x = sw[j];
if (sgstart[j] < j) {
const float lim = prev + kPertol;
if (x < lim) x = lim;
}
xs_ws[batch * N + j] = x;
prev = x;
}
}
__syncthreads();
for (int i = tid; i < N; i += NT) {
ds_ws[batch * N + i] = sds[i];
es_ws[batch * N + i] = ses[i];
w_ws[batch * N + i] = sw[i];
// gamma == 1 here: lam = fp32(w * onenrm) (oracle convention with the
// identity gamma multiply).
lam_out[batch * N + i] = sw[i] * onenrm;
gstart_ws[batch * N + i] = sgstart[i];
minits_ws[batch * N + i] = sminits[i];
}
if (tid == 0) {
pivmin_ws[batch] = pivmin;
}
}
// ======================================================================
// Fused SLAGTF factor + forward elimination + SLAGTS back-substitution
// with the MANDATORY dynamic-rescale guard (verbatim mid-n).
// ======================================================================
template <int N>
__device__ void thomas_solve_column(const float* __restrict__ sds,
const float* __restrict__ ses,
const float xj, const float pivmin,
float* __restrict__ y,
float* __restrict__ x,
float* __restrict__ u0,
float* __restrict__ u1,
float* __restrict__ u2,
float* out_log2growth, float* out_mx,
int* guard_hits) {
float dlog2 = 0.0f;
float di = sds[0] - xj;
float sup = ses[0];
for (int i = 0; i < N - 1; ++i) {
const float ai = di;
const float ci = ses[i];
const float bi = sup;
const bool swp = fabsf(ai) < fabsf(ci);
float piv = swp ? ci : ai;
if (fabsf(piv) < pivmin) piv = piv < 0.0f ? -pivmin : pivmin;
const float m = (swp ? ai : ci) / piv;
const float a_next = sds[i + 1] - xj;
const float b_next = (i + 1 < N - 1) ? ses[i + 1] : 0.0f;
u0[static_cast<size_t>(i) * N] = piv;
u1[static_cast<size_t>(i) * N] = swp ? a_next : bi;
u2[static_cast<size_t>(i) * N] = swp ? b_next : 0.0f;
di = swp ? (bi - m * a_next) : (a_next - m * bi);
sup = swp ? (-m * b_next) : b_next;
const float yi = y[static_cast<size_t>(i) * N];
const float yn = y[static_cast<size_t>(i + 1) * N];
const float y_low = swp ? yn : yi;
const float y_high = (swp ? yi : yn) - m * y_low;
y[static_cast<size_t>(i) * N] = y_low;
y[static_cast<size_t>(i + 1) * N] = y_high;
if (fabsf(y_high) > 1e30f) {
for (int r = 0; r < N; ++r) y[static_cast<size_t>(r) * N] *= 1e-30f;
dlog2 += kLog2Tiny30;
}
}
{
float last = di;
if (fabsf(last) < pivmin) last = last < 0.0f ? -pivmin : pivmin;
u0[static_cast<size_t>(N - 1) * N] = last;
}
for (int i = N - 1; i >= 0; --i) {
float num = y[static_cast<size_t>(i) * N];
if (i + 1 < N) {
num -= u1[static_cast<size_t>(i) * N] *
x[static_cast<size_t>(i + 1) * N];
}
if (i + 2 < N) {
num -= u2[static_cast<size_t>(i) * N] *
x[static_cast<size_t>(i + 2) * N];
}
const float den = u0[static_cast<size_t>(i) * N];
const float absn = fabsf(num);
const float absd = fabsf(den);
if (absn > absd * kXmax) {
const float f = (absd * kXmax) / absn;
for (int r = i + 1; r < N; ++r) x[static_cast<size_t>(r) * N] *= f;
for (int r = 0; r < i; ++r) y[static_cast<size_t>(r) * N] *= f;
num *= f;
dlog2 += log2f(f);
++(*guard_hits);
}
x[static_cast<size_t>(i) * N] = num / den;
}
float mx = 0.0f;
for (int r = 0; r < N; ++r) {
mx = fmaxf(mx, fabsf(x[static_cast<size_t>(r) * N]));
}
*out_mx = mx;
*out_log2growth = (mx > 0.0f ? log2f(mx) : 0.0f) - dlog2;
}
// ======================================================================
// KI: DSTEIN inverse iteration (verbatim mid-n K3) + info fail-closed:
// skips info!=0 blocks with Q = I; sets info = 2 on stagnation.
// ======================================================================
template <int N, int NT>
__global__ void __launch_bounds__(NT)
lower_invit_kernel(const float* __restrict__ ds_ws,
const float* __restrict__ es_ws,
const float* __restrict__ xs_ws,
const int* __restrict__ gstart_ws,
const int* __restrict__ minits_ws,
const float* __restrict__ pivmin_ws,
float* __restrict__ v_ws,
float* __restrict__ y_ws,
float* __restrict__ u0_ws,
float* __restrict__ u1_ws,
float* __restrict__ u2_ws,
int* __restrict__ info,
int* __restrict__ stats_ws,
int batch_count) {
constexpr int kWarps = NT / 32;
__shared__ float sds[N];
__shared__ float ses[N];
__shared__ float sxs[N];
__shared__ float sc[N];
__shared__ int sgstart[N];
__shared__ signed char smin[N];
__shared__ signed char sacc[N];
__shared__ signed char srnd[N];
__shared__ float partials[kWarps];
__shared__ int s_naccepted;
__shared__ int s_rescues;
__shared__ int s_guard;
__shared__ float s_nn;
__shared__ float s_rem;
const int batch = blockIdx.x;
if (batch >= batch_count) {
return;
}
const int tid = threadIdx.x;
const int warp = tid >> 5;
const int lane = tid & 31;
float* Vb = v_ws + static_cast<size_t>(batch) * N * N;
if (info[batch] != 0) {
// degenerate emit: Q = I (matrix routes to fallback anyway).
for (int idx = tid; idx < N * N; idx += NT) {
const int r = idx / N;
const int c = idx - r * N;
Vb[idx] = (r == c) ? 1.0f : 0.0f;
}
return;
}
const float pivmin = pivmin_ws[batch];
float* Yb = y_ws + static_cast<size_t>(batch) * N * N;
float* U0b = u0_ws + static_cast<size_t>(batch) * N * N;
float* U1b = u1_ws + static_cast<size_t>(batch) * N * N;
float* U2b = u2_ws + static_cast<size_t>(batch) * N * N;
for (int i = tid; i < N; i += NT) {
sds[i] = ds_ws[batch * N + i];
ses[i] = es_ws[batch * N + i];
sxs[i] = xs_ws[batch * N + i];
sgstart[i] = gstart_ws[batch * N + i];
smin[i] = static_cast<signed char>(minits_ws[batch * N + i]);
sacc[i] = 0;
srnd[i] = 0;
}
if (tid == 0) {
s_naccepted = 0;
s_rescues = 0;
s_guard = 0;
}
__syncthreads();
int guard_hits_local = 0;
// Round-0 RHS (option-B micro, bitwise-identical values): hash writes
// COALESCED over the linear index (element (r, j) = hash(b, j, r, 0)
// exactly as before), then the owner thread normalizes its column with
// the same ascending-r Kahan sum and divide.
for (int idx = tid; idx < N * N; idx += NT) {
const int r = idx / N;
const int j = idx - r * N;
Yb[idx] = hash_uniform(batch, j, r, 0u);
}
__syncthreads();
if (tid < N) {
float sum = 0.0f, comp = 0.0f;
for (int r = 0; r < N; ++r) {
const float t = Yb[static_cast<size_t>(r) * N + tid];
const float term = fmaf(t, t, -comp);
const float tmp = sum + term;
comp = (tmp - sum) - term;
sum = tmp;
}
const float nrm = sqrtf(sum);
if (nrm > 0.0f) {
for (int r = 0; r < N; ++r) {
Yb[static_cast<size_t>(r) * N + tid] /= nrm;
}
}
}
auto cta_project = [&](float* M, int jcol, int k0, int kend) {
for (int k = k0 + warp; k < kend; k += kWarps) {
float part = 0.0f;
for (int i = lane; i < N; i += 32) {
part = fmaf(Vb[static_cast<size_t>(i) * N + k],
M[static_cast<size_t>(i) * N + jcol], part);
}
part = warp_sum(part);
if (lane == 0) sc[k] = part;
}
__syncthreads();
if (kend > k0) {
for (int i = tid; i < N; i += NT) {
float acc = 0.0f;
for (int k = k0; k < kend; ++k) {
acc = fmaf(sc[k], Vb[static_cast<size_t>(i) * N + k], acc);
}
M[static_cast<size_t>(i) * N + jcol] -= acc;
}
}
__syncthreads();
float part = 0.0f;
for (int i = tid; i < N; i += NT) {
const float t = M[static_cast<size_t>(i) * N + jcol];
part = fmaf(t, t, part);
}
part = warp_sum(part);
if (lane == 0) partials[warp] = part;
__syncthreads();
if (tid == 0) {
float s = 0.0f;
for (int w = 0; w < kWarps; ++w) s += partials[w];
s_nn = sqrtf(s);
s_rem = (kend > k0) ? s_nn : 1.0f;
}
__syncthreads();
};
auto cta_scale_column = [&](float* M, int jcol, float divisor) {
for (int i = tid; i < N; i += NT) {
M[static_cast<size_t>(i) * N + jcol] /= divisor;
}
__syncthreads();
};
for (int it = 0; it < kMaxIts; ++it) {
__syncthreads();
if (s_naccepted >= N) break;
if (tid < N && !sacc[tid]) {
float lg = 0.0f, mx = 0.0f;
thomas_solve_column<N>(sds, ses, sxs[tid], pivmin, Yb + tid, Vb + tid,
U0b + tid, U1b + tid, U2b + tid, &lg, &mx,
&guard_hits_local);
float nrm = 0.0f;
if (mx > 0.0f) {
float sum = 0.0f, comp = 0.0f;
for (int r = 0; r < N; ++r) {
const float t = Vb[static_cast<size_t>(r) * N + tid] / mx;
Vb[static_cast<size_t>(r) * N + tid] = t;
const float term = fmaf(t, t, -comp);
const float tmp = sum + term;
comp = (tmp - sum) - term;
sum = tmp;
}
nrm = sqrtf(sum);
if (nrm > 0.0f) {
for (int r = 0; r < N; ++r) {
Vb[static_cast<size_t>(r) * N + tid] /= nrm;
}
}
}
srnd[tid] += 1;
if (nrm > 0.0f && lg + log2f(nrm) >= kLog2Gtol &&
srnd[tid] >= smin[tid]) {
sacc[tid] = 1;
atomicAdd(&s_naccepted, 1);
}
}
__syncthreads();
for (int jj = 0; jj < N; ++jj) {
if (smin[jj] != 3) continue;
const int k0 = sgstart[jj];
cta_project(Vb, jj, k0, jj);
if (s_nn > 0.0f) cta_scale_column(Vb, jj, s_nn);
}
__syncthreads();
// rhs <- v, skipping accepted columns (their Yb is never read again;
// bitwise-identical to the full copy for every consumed value).
for (int idx = tid; idx < N * N; idx += NT) {
if (!sacc[idx - (idx / N) * N]) {
Yb[idx] = Vb[idx];
}
}
}
__syncthreads();
for (int jj = 0; jj < N; ++jj) {
if (smin[jj] != 3) continue;
const int k0 = sgstart[jj];
cta_project(Vb, jj, k0, jj);
int tries = 0;
while (s_rem < kDepTol && tries < 3) {
++tries;
if (tid == 0) ++s_rescues;
for (int i = tid; i < N; i += NT) {
Yb[static_cast<size_t>(i) * N + jj] =
hash_uniform(batch, jj, i, 100u + tries);
}
__syncthreads();
cta_project(Yb, jj, k0, jj);
if (s_nn > 0.0f) cta_scale_column(Yb, jj, s_nn);
if (tid == jj) {
float lg = 0.0f, mx = 0.0f;
thomas_solve_column<N>(sds, ses, sxs[jj], pivmin, Yb + jj, Vb + jj,
U0b + jj, U1b + jj, U2b + jj, &lg, &mx,
&guard_hits_local);
if (mx > 0.0f) {
for (int r = 0; r < N; ++r) {
Vb[static_cast<size_t>(r) * N + jj] /= mx;
}
}
srnd[jj] += 1;
}
__syncthreads();
cta_project(Vb, jj, k0, jj);
}
if (s_nn > 0.0f) cta_scale_column(Vb, jj, s_nn);
}
__syncthreads();
atomicAdd(&s_guard, guard_hits_local);
int rmax = 0, rsum = 0, stag = 0;
if (tid < N) {
rmax = srnd[tid];
rsum = srnd[tid];
stag = sacc[tid] ? 0 : 1;
}
rmax = warp_imax(rmax);
rsum = warp_isum(rsum);
stag = warp_isum(stag);
if (lane == 0) partials[warp] = __int_as_float(rmax);
__syncthreads();
if (tid == 0) {
int m = 0;
for (int w = 0; w < kWarps; ++w) m = max(m, __float_as_int(partials[w]));
stats_ws[batch * kStatsSlots + 6] = m;
}
__syncthreads();
if (lane == 0) partials[warp] = __int_as_float(rsum);
__syncthreads();
if (tid == 0) {
int s = 0;
for (int w = 0; w < kWarps; ++w) s += __float_as_int(partials[w]);
stats_ws[batch * kStatsSlots + 7] = s;
}
__syncthreads();
if (lane == 0) partials[warp] = __int_as_float(stag);
__syncthreads();
if (tid == 0) {
int s = 0;
for (int w = 0; w < kWarps; ++w) s += __float_as_int(partials[w]);
stats_ws[batch * kStatsSlots + 8] = s;
stats_ws[batch * kStatsSlots + 9] = s_rescues;
stats_ws[batch * kStatsSlots + 10] = s_guard;
if (s > 0) info[batch] = 2; // fail-closed: stagnation -> escalate
}
}
} // namespace
template <int N>
void lower_solve_impl(torch::Tensor& d_blocks, torch::Tensor& e_blocks,
torch::Tensor& lam, torch::Tensor& q,
torch::Tensor& info, torch::Tensor& stats) {
const int blocks = static_cast<int>(d_blocks.size(0));
auto opts_f = d_blocks.options();
auto opts_i = torch::TensorOptions()
.dtype(torch::kInt32).device(d_blocks.device());
auto ds_ws = torch::empty({blocks, N}, opts_f);
auto es_ws = torch::empty({blocks, N}, opts_f);
auto w_ws = torch::empty({blocks, N}, opts_f);
auto xs_ws = torch::empty({blocks, N}, opts_f);
auto gstart = torch::empty({blocks, N}, opts_i);
auto minits = torch::empty({blocks, N}, opts_i);
auto pivmin = torch::empty({blocks}, opts_f);
auto y_ws = torch::empty({blocks, N, N}, opts_f);
auto u0_ws = torch::empty({blocks, N, N}, opts_f);
auto u1_ws = torch::empty({blocks, N, N}, opts_f);
auto u2_ws = torch::empty({blocks, N, N}, opts_f);
lower_bisect_kernel<N, N><<<blocks, N, 0>>>(
d_blocks.data_ptr<float>(), e_blocks.data_ptr<float>(),
ds_ws.data_ptr<float>(), es_ws.data_ptr<float>(),
w_ws.data_ptr<float>(), lam.data_ptr<float>(),
xs_ws.data_ptr<float>(), gstart.data_ptr<int>(),
minits.data_ptr<int>(), pivmin.data_ptr<float>(),
info.data_ptr<int>(), stats.data_ptr<int>(), blocks);
lower_invit_kernel<N, N><<<blocks, N, 0>>>(
ds_ws.data_ptr<float>(), es_ws.data_ptr<float>(),
xs_ws.data_ptr<float>(), gstart.data_ptr<int>(),
minits.data_ptr<int>(), pivmin.data_ptr<float>(),
q.data_ptr<float>(), y_ws.data_ptr<float>(),
u0_ws.data_ptr<float>(), u1_ws.data_ptr<float>(),
u2_ws.data_ptr<float>(), info.data_ptr<int>(),
stats.data_ptr<int>(), blocks);
}
void lower_solve_(torch::Tensor d_blocks, torch::Tensor e_blocks,
torch::Tensor lam, torch::Tensor q, torch::Tensor info,
torch::Tensor stats) {
CHECK_IN(d_blocks);
CHECK_IN(e_blocks);
CHECK_IN(lam);
CHECK_IN(q);
CHECK_IN(info);
CHECK_IN(stats);
const int64_t blocks = d_blocks.size(0);
const int64_t n = d_blocks.size(1);
TORCH_CHECK(n == 64 || n == 128, "width must be 64 or 128");
TORCH_CHECK(e_blocks.sizes() == torch::IntArrayRef({blocks, n - 1}));
TORCH_CHECK(lam.sizes() == torch::IntArrayRef({blocks, n}));
TORCH_CHECK(q.sizes() == torch::IntArrayRef({blocks, n, n}));
TORCH_CHECK(info.sizes() == torch::IntArrayRef({blocks}));
TORCH_CHECK(stats.sizes() == torch::IntArrayRef({blocks, kStatsSlots}));
TORCH_CHECK(info.dtype() == torch::kInt32);
TORCH_CHECK(stats.dtype() == torch::kInt32);
if (n == 64) {
lower_solve_impl<64>(d_blocks, e_blocks, lam, q, info, stats);
} else {
lower_solve_impl<128>(d_blocks, e_blocks, lam, q, info, stats);
}
}
template <typename KernelT>
void push_attrs(std::vector<int64_t>& out, KernelT kernel, int threads) {
cudaFuncAttributes attrs{};
cudaFuncGetAttributes(&attrs, kernel);
int blocks_per_sm = 0;
cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&blocks_per_sm, kernel, threads, 0);
out.push_back(attrs.numRegs);
out.push_back(attrs.localSizeBytes);
out.push_back(attrs.sharedSizeBytes);
out.push_back(blocks_per_sm);
}
// [bisect regs, local, static shared, blocks/SM,
// invit regs, local, static shared, blocks/SM]
std::vector<int64_t> lower_resource_attributes(int64_t n) {
std::vector<int64_t> out;
if (n == 64) {
push_attrs(out, lower_bisect_kernel<64, 64>, 64);
push_attrs(out, lower_invit_kernel<64, 64>, 64);
} else {
push_attrs(out, lower_bisect_kernel<128, 128>, 128);
push_attrs(out, lower_invit_kernel<128, 128>, 128);
}
return out;
}
} // namespace lane_k
namespace lane_l {
#define FULL_MASK 0xffffffffu
#define KD 8
__device__ __forceinline__ float warp_sum(float v) {
#pragma unroll
for (int o = 16; o > 0; o >>= 1) v += __shfl_xor_sync(FULL_MASK, v, o);
return v;
}
// kiki house_coeffs: cf[0]=tau, cf[1]=1/(alpha-beta), cf[2]=beta
__device__ __forceinline__ void house_coeffs(float alpha, float sigma, float* cf) {
if (sigma <= 0.f) {
cf[0] = 0.f; cf[1] = 0.f; cf[2] = alpha;
} else {
float beta = -copysignf(sqrtf(fmaf(alpha, alpha, sigma)), alpha);
cf[0] = (beta - alpha) / beta;
cf[1] = 1.f / (alpha - beta);
cf[2] = beta;
}
}
// kiki panel_core, verbatim (sol_combo_winner_kiki.py): unblocked
// Householder QR of an r x w panel resident in shared (leading dim sld),
// one warp per trailing column, fused next-column house via the k==j+1
// branch. On exit: S holds R (rows i<=j) + unscaled v tails (i>j);
// gammas[j] = 1/(alpha-beta), taug[j] = tau.
template <int NT>
__device__ void panel_core(float* S, long sld, int r, int w,
float* cf, float* gammas, float* taug,
float* scratch) {
const int lane = threadIdx.x & 31, wid = threadIdx.x >> 5;
const int nw = NT >> 5;
{
float part = 0.f;
for (int i = 1 + threadIdx.x; i < r; i += NT) {
float x = S[i];
part = fmaf(x, x, part);
}
part = warp_sum(part);
if (lane == 0) scratch[wid] = part;
__syncthreads();
if (threadIdx.x == 0) {
float sg = 0.f;
for (int u = 0; u < nw; ++u) sg += scratch[u];
house_coeffs(S[0], sg, cf);
}
__syncthreads();
}
for (int j = 0; j < w; ++j) {
const float* cfc = cf + 4 * (j & 1);
float* cfn = cf + 4 * ((j + 1) & 1);
float tj = cfc[0], gj = cfc[1], bj = cfc[2];
float* colj = S + (long)j * sld;
if (threadIdx.x == 0) {
gammas[j] = gj;
taug[j] = tj;
}
for (int k = j + 1 + wid; k < w; k += nw) {
float* ck = S + (long)k * sld;
float d = (lane == 0) ? ck[j] : 0.f;
float acc = 0.f;
for (int i = j + 1 + lane; i < r; i += 32) acc = fmaf(colj[i], ck[i], acc);
d += gj * acc;
d = warp_sum(d);
float wk = tj * d;
float alpha_next = 0.f;
float sq = 0.f;
if (lane == 0) ck[j] -= wk;
float wg = wk * gj;
for (int i = j + 1 + lane; i < r; i += 32) {
float nv = fmaf(-wg, colj[i], ck[i]);
ck[i] = nv;
if (k == j + 1) {
if (i == j + 1) alpha_next = nv;
else sq = fmaf(nv, nv, sq);
}
}
if (k == j + 1) {
sq = warp_sum(sq);
if (lane == 0) house_coeffs(alpha_next, sq, cfn);
}
}
if (threadIdx.x == 0) colj[j] = bj;
__syncthreads();
}
}
// kiki pair_dots, verbatim: sWv[i,j] = V_i^T V_j (i<j) from the packed
// panel representation (unscaled tails + gamma factors).
template <int NT>
__device__ void pair_dots(const float* S, long sld, int r, int w,
const float* gammas, float* sWv, int ldwv, int o) {
const int lane = threadIdx.x & 31, wid = threadIdx.x >> 5;
const int nw = NT >> 5;
const int npairs = w * (w - 1) / 2;
for (int p = wid; p < npairs; p += nw) {
int j = (int)((1.0f + sqrtf(1.0f + 8.0f * (float)p)) * 0.5f);
while (j * (j - 1) / 2 > p) --j;
while ((j + 1) * j / 2 <= p) ++j;
int i = p - j * (j - 1) / 2;
const float* ci = S + (long)i * sld;
const float* cj = S + (long)j * sld;
float acc = 0.f;
for (int l = j + 1 + lane; l < r; l += 32) acc = fmaf(ci[l], cj[l], acc);
acc = warp_sum(acc);
if (lane == 0) {
sWv[(o + i) * ldwv + (o + j)] = gammas[i] * ci[j] + gammas[i] * gammas[j] * acc;
}
}
}
// kiki t_recurrence, verbatim: forward larft (matches the incumbent oracle).
__device__ void t_recurrence(const float* sWv, int ldwv, int o,
const float* taug, float* sT, int ldt, int w) {
const int lane = threadIdx.x & 31;
for (int j = 0; j < w; ++j) {
float tj = taug[j];
for (int i = lane; i < j; i += 32) {
float s = 0.f;
for (int k = i; k < j; ++k)
s = fmaf(sT[i * ldt + k], sWv[(o + k) * ldwv + (o + j)], s);
sT[i * ldt + j] = -tj * s;
}
if (lane == 0) sT[j * ldt + j] = tj;
for (int i = j + 1 + lane; i < w; i += 32) sT[i * ldt + j] = 0.f;
__syncwarp();
}
}
// One CTA = one matrix, one w=8 panel at column j0 (reflector rows j0+8..n).
// Reads the panel zero-copy from A (lda=n); writes R (upper 8x8, zeros
// below) back into A rows [j0+8, j0+16) x cols [j0, j0+8) — the band
// region the driver's as_strided gather reads — plus the contract buffers.
template <int NT>
__global__ void __launch_bounds__(NT)
panel_qr8_kernel(float* __restrict__ A,
float* __restrict__ Vout, // [B, r, 8] unit-lower
float* __restrict__ Tout, // [B, 8, 8] forward larft
float* __restrict__ tauo, // [B, 8]
int n, int j0) {
extern __shared__ float smem[];
const int r = n - j0 - KD;
const int sld = r | 1; // P7 bank-stride pad
float* S = smem; // sld * 8 panel
float* sWv = S + (long)sld * KD; // 8*8 pair dots
float* sT = sWv + KD * KD; // 8*8 T
float* gammas = sT + KD * KD; // 8
float* taug = gammas + KD; // 8
float* cf = taug + KD; // 8 (double-buffered coeffs)
float* scratch = cf + 8; // NT/32 block-reduce slots
const int tid = threadIdx.x;
const long b = blockIdx.x;
float* Ab = A + b * (long)n * n;
// load panel (coalesced global reads, bank-padded shared writes)
for (int idx = tid; idx < r * KD; idx += NT) {
int i = idx / KD, c = idx - i * KD;
S[(long)c * sld + i] = Ab[(long)(j0 + KD + i) * n + (j0 + c)];
}
__syncthreads();
panel_core<NT>(S, sld, r, KD, cf, gammas, taug, scratch);
pair_dots<NT>(S, sld, r, KD, gammas, sWv, KD, 0);
__syncthreads();
if ((tid >> 5) == 0) t_recurrence(sWv, KD, 0, taug, sT, KD, KD);
__syncthreads();
// R block back into the parent (band region); zeros below the diagonal
for (int idx = tid; idx < KD * KD; idx += NT) {
int i = idx / KD, c = idx - i * KD;
Ab[(long)(j0 + KD + i) * n + (j0 + c)] = (i <= c) ? S[(long)c * sld + i] : 0.f;
}
// contract buffers: V unit-lower with explicit 0/1 head, T, tau
float* Vb = Vout + b * (long)r * KD;
for (int idx = tid; idx < r * KD; idx += NT) {
int i = idx / KD, c = idx - i * KD;
float x = S[(long)c * sld + i];
Vb[idx] = (i < c) ? 0.f : (i == c ? 1.f : gammas[c] * x);
}
float* Tb = Tout + b * (long)(KD * KD);
for (int idx = tid; idx < KD * KD; idx += NT) Tb[idx] = sT[idx];
float* tb = tauo + b * (long)KD;
for (int c = tid; c < KD; c += NT) tb[c] = taug[c];
}
// ---------------------------------------------------------------- host
void panel_qr8(torch::Tensor A, torch::Tensor V, torch::Tensor T,
torch::Tensor tau, int64_t j0) {
const int B = A.size(0), n = A.size(1);
const int r = n - (int)j0 - KD;
TORCH_CHECK(r >= KD, "panel needs r >= 8");
TORCH_CHECK(V.size(1) == r && V.size(2) == KD, "V shape mismatch");
constexpr int NT = 256;
size_t smem = ((size_t)(r | 1) * KD + 2 * KD * KD + 2 * KD + 8 + NT / 32)
* sizeof(float);
auto kern = panel_qr8_kernel<NT>;
// ~16.6 KB max (r=504): under the 48 KB default, no opt-in needed on
// any target arch (GB10 clamp scar does not bite here).
kern<<<B, NT, smem>>>(A.data_ptr<float>(), V.data_ptr<float>(),
T.data_ptr<float>(), tau.data_ptr<float>(),
n, (int)j0);
}
} // namespace lane_l
namespace lane_m {
namespace {
__device__ __forceinline__ float warp_sum(float value) {
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
value += __shfl_xor_sync(FULL_MASK, value, offset);
}
return value;
}
__device__ __forceinline__ void house_coeffs(
float alpha, float sigma, float* coeffs) {
if (sigma <= 0.0f) {
coeffs[0] = 0.0f;
coeffs[1] = 0.0f;
coeffs[2] = alpha;
} else {
const float beta = -copysignf(sqrtf(fmaf(alpha, alpha, sigma)), alpha);
coeffs[0] = (beta - alpha) / beta;
coeffs[1] = 1.0f / (alpha - beta);
coeffs[2] = beta;
}
}
template <int Threads, int Width>
__device__ void panel_core(
float* panel,
int leading_dimension,
int rows,
float* coeffs,
float* gammas,
float* taus,
float* scratch) {
const int lane = threadIdx.x & 31;
const int warp = threadIdx.x >> 5;
constexpr int warps = Threads >> 5;
float partial = 0.0f;
for (int row = 1 + threadIdx.x; row < rows; row += Threads) {
const float value = panel[row];
partial = fmaf(value, value, partial);
}
partial = warp_sum(partial);
if (lane == 0) {
scratch[warp] = partial;
}
__syncthreads();
if (threadIdx.x == 0) {
float sigma = 0.0f;
for (int index = 0; index < warps; ++index) {
sigma += scratch[index];
}
house_coeffs(panel[0], sigma, coeffs);
}
__syncthreads();
#pragma unroll
for (int column = 0; column < Width; ++column) {
const float* current_coeffs = coeffs + 4 * (column & 1);
float* next_coeffs = coeffs + 4 * ((column + 1) & 1);
const float tau = current_coeffs[0];
const float gamma = current_coeffs[1];
const float beta = current_coeffs[2];
float* current = panel + static_cast<long>(column) * leading_dimension;
if (threadIdx.x == 0) {
gammas[column] = gamma;
taus[column] = tau;
}
for (int target_column = column + 1 + warp;
target_column < Width;
target_column += warps) {
float* target =
panel + static_cast<long>(target_column) * leading_dimension;
float dot = lane == 0 ? target[column] : 0.0f;
float tail_dot = 0.0f;
for (int row = column + 1 + lane; row < rows; row += 32) {
tail_dot = fmaf(current[row], target[row], tail_dot);
}
dot += gamma * tail_dot;
dot = warp_sum(dot);
const float weight = tau * dot;
float next_alpha = 0.0f;
float next_sigma = 0.0f;
if (lane == 0) {
target[column] -= weight;
}
const float scaled_weight = weight * gamma;
for (int row = column + 1 + lane; row < rows; row += 32) {
const float updated = fmaf(-scaled_weight, current[row], target[row]);
target[row] = updated;
if (target_column == column + 1) {
if (row == column + 1) {
next_alpha = updated;
} else {
next_sigma = fmaf(updated, updated, next_sigma);
}
}
}
if (target_column == column + 1) {
next_sigma = warp_sum(next_sigma);
if (lane == 0) {
house_coeffs(next_alpha, next_sigma, next_coeffs);
}
}
}
if (threadIdx.x == 0) {
current[column] = beta;
}
__syncthreads();
}
}
template <int Threads, int Width>
__device__ void pair_dots(
const float* shared_panel,
int leading_dimension,
int rows,
const float* gammas,
float* gram_upper) {
const int lane = threadIdx.x & 31;
const int warp = threadIdx.x >> 5;
constexpr int warps = Threads >> 5;
constexpr int pairs = Width * (Width - 1) / 2;
for (int pair = warp; pair < pairs; pair += warps) {
int column = static_cast<int>(
(1.0f + sqrtf(1.0f + 8.0f * static_cast<float>(pair))) * 0.5f);
while (column * (column - 1) / 2 > pair) {
--column;
}
while ((column + 1) * column / 2 <= pair) {
++column;
}
const int prior = pair - column * (column - 1) / 2;
const float* prior_column =
shared_panel + static_cast<long>(prior) * leading_dimension;
const float* current_column =
shared_panel + static_cast<long>(column) * leading_dimension;
float tail = 0.0f;
for (int row = column + 1 + lane; row < rows; row += 32) {
tail = fmaf(prior_column[row], current_column[row], tail);
}
tail = warp_sum(tail);
if (lane == 0) {
gram_upper[prior * Width + column] =
gammas[prior] * prior_column[column]
+ gammas[prior] * gammas[column] * tail;
}
}
}
template <int Width>
__device__ void t_recurrence(
const float* gram_upper,
const float* taus,
float* compact_t) {
const int lane = threadIdx.x & 31;
#pragma unroll
for (int column = 0; column < Width; ++column) {
const float tau = taus[column];
if (lane < column) {
float value = 0.0f;
for (int inner = lane; inner < column; ++inner) {
value = fmaf(
compact_t[lane * Width + inner],
gram_upper[inner * Width + column],
value);
}
compact_t[lane * Width + column] = -tau * value;
} else if (lane == column) {
compact_t[column * Width + column] = tau;
} else if (lane < Width) {
compact_t[lane * Width + column] = 0.0f;
}
__syncwarp();
}
}
template <int Threads, int Width>
__global__ void panel_qr_narrow_vt_kernel(
const float* __restrict__ input,
float* __restrict__ packed,
float* __restrict__ tau,
float* __restrict__ householder_v,
float* __restrict__ compact_t,
int rows) {
const int leading_dimension = rows | 1;
extern __shared__ float shared[];
float* panel = shared;
float* gammas = panel + static_cast<long>(leading_dimension) * Width;
float* taus = gammas + Width;
float* coeffs = taus + Width;
float* scratch = coeffs + 8;
float* gram_upper = scratch + 32;
float* shared_t = gram_upper + Width * Width;
const long batch = blockIdx.x;
const float* input_batch = input + batch * static_cast<long>(rows) * Width;
float* packed_batch = packed + batch * static_cast<long>(rows) * Width;
float* tau_batch = tau + batch * Width;
float* v_batch = householder_v + batch * static_cast<long>(rows) * Width;
float* t_batch = compact_t + batch * Width * Width;
for (int index = threadIdx.x; index < rows * Width; index += Threads) {
const int row = index / Width;
const int column = index - row * Width;
panel[static_cast<long>(column) * leading_dimension + row] =
input_batch[index];
}
__syncthreads();
panel_core<Threads, Width>(
panel, leading_dimension, rows, coeffs, gammas, taus, scratch);
pair_dots<Threads, Width>(
panel, leading_dimension, rows, gammas, gram_upper);
__syncthreads();
if ((threadIdx.x >> 5) == 0) {
t_recurrence<Width>(gram_upper, taus, shared_t);
}
__syncthreads();
for (int column = threadIdx.x; column < Width; column += Threads) {
tau_batch[column] = taus[column];
}
for (int index = threadIdx.x; index < rows * Width; index += Threads) {
const int row = index / Width;
const int column = index - row * Width;
const float raw = panel[static_cast<long>(column) * leading_dimension + row];
packed_batch[index] = row > column ? gammas[column] * raw : raw;
v_batch[index] = row < column
? 0.0f
: (row == column ? 1.0f : gammas[column] * raw);
}
for (int index = threadIdx.x; index < Width * Width; index += Threads) {
t_batch[index] = shared_t[index];
}
}
constexpr int kWidth = 8;
constexpr int kThreads = 256;
constexpr int kMaxRows = 1024 - kWidth;
// Spark (GB10, sm_121) correctness runs cap opt-in dynamic shared memory near
// 99 KB per CTA; B200 (sm_100) allows far more. Budget check at bind time.
constexpr size_t kSparkSharedBudget = 99 * 1024;
size_t shared_bytes_for_rows(int rows) {
const int leading_dimension = rows | 1;
return (static_cast<size_t>(leading_dimension) * kWidth
+ 2 * kWidth + 8 + 32 + 2 * kWidth * kWidth) * sizeof(float);
}
} // namespace
void panel_qr_w8_vt_out(
const torch::Tensor& input,
torch::Tensor packed,
torch::Tensor tau,
torch::Tensor householder_v,
torch::Tensor compact_t) {
TORCH_CHECK(input.is_cuda(), "input must be CUDA");
TORCH_CHECK(input.scalar_type() == torch::kFloat32, "input must be FP32");
TORCH_CHECK(input.is_contiguous(), "input must be contiguous");
TORCH_CHECK(input.dim() == 3, "input must have shape [batch, rows, 8]");
TORCH_CHECK(input.size(2) == kWidth, "panel width must be 8");
const long rows = input.size(1);
TORCH_CHECK(rows >= kWidth && rows <= kMaxRows && rows % kWidth == 0,
"rows must be a multiple of 8 in [8, 504], got ", rows);
TORCH_CHECK(packed.sizes() == input.sizes(), "packed shape mismatch");
TORCH_CHECK(packed.scalar_type() == torch::kFloat32 && packed.is_contiguous(),
"packed must be contiguous FP32");
TORCH_CHECK(tau.dim() == 2 && tau.size(0) == input.size(0)
&& tau.size(1) == kWidth,
"tau shape mismatch");
TORCH_CHECK(tau.scalar_type() == torch::kFloat32 && tau.is_contiguous(),
"tau must be contiguous FP32");
TORCH_CHECK(householder_v.sizes() == input.sizes(), "V shape mismatch");
TORCH_CHECK(householder_v.scalar_type() == torch::kFloat32
&& householder_v.is_contiguous(),
"V must be contiguous FP32");
TORCH_CHECK(compact_t.dim() == 3 && compact_t.size(0) == input.size(0)
&& compact_t.size(1) == kWidth && compact_t.size(2) == kWidth,
"T shape mismatch");
TORCH_CHECK(compact_t.scalar_type() == torch::kFloat32
&& compact_t.is_contiguous(),
"T must be contiguous FP32");
TORCH_CHECK(packed.device() == input.device()
&& tau.device() == input.device()
&& householder_v.device() == input.device()
&& compact_t.device() == input.device(),
"all tensors must be on the input device");
const size_t shared_bytes = shared_bytes_for_rows(static_cast<int>(rows));
TORCH_CHECK(shared_bytes <= kSparkSharedBudget,
"shared budget exceeded: ", shared_bytes, " > ",
kSparkSharedBudget);
c10::cuda::CUDAGuard guard(input.device());
auto kernel = panel_qr_narrow_vt_kernel<kThreads, kWidth>;
// Max shared for rows=504 is ~16.6 KB, under the 48 KB static limit on both
// sm_100 and sm_121, so no cudaFuncSetAttribute opt-in is required.
kernel<<<input.size(0), kThreads, shared_bytes>>>(
input.data_ptr<float>(),
packed.data_ptr<float>(),
tau.data_ptr<float>(),
householder_v.data_ptr<float>(),
compact_t.data_ptr<float>(),
static_cast<int>(rows));
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
int64_t panel_qr_w8_shared_bytes(int64_t rows) {
return static_cast<int64_t>(shared_bytes_for_rows(static_cast<int>(rows)));
}
} // namespace lane_m
namespace lane_n {
namespace {
constexpr int kN = 1024;
constexpr int kKD = 8;
constexpr int kInputRows = kKD + 1;
constexpr int kWorkRows = 2 * kKD + 1;
constexpr int kTaskWarps = 32;
constexpr int kThreads = kTaskWarps * 32;
constexpr int kReflectors = 65792; // sum ceil((1023-s)/8), s=1..1022 (n1024)
constexpr int kConsumers = 130434; // tasks - (n-1) kind-1 tasks (n1024)
constexpr int kTasks = 131457; // 131,456 loop tasks + terminal no-op (n1024)
constexpr int kRecordWidth = kKD; // [tau, v1..v7]
constexpr int kSweepPrefixEntries = kN - 1;
constexpr int kScalarsPerMatrixTask = 8;
constexpr int kMatrixWorkPadding = 4;
constexpr int kMatrixWorkStride = kN * kWorkRows + kMatrixWorkPadding;
constexpr int kScalarStride = kScalarsPerMatrixTask + 1;
constexpr int kCounterIntsPerMatrix = 4;
__device__ __constant__ uint32_t kSweepPrefix[kSweepPrefixEntries];
static_assert(kMatrixWorkStride % 32 == 4,
"matrix stride must break bank aliasing");
static_assert(kScalarStride % 32 == 9,
"scalar stride must break bank aliasing");
template <int M, int S>
struct Config {
static_assert(M * S <= 32, "subgroups must fit one warp");
static constexpr size_t work_floats() {
return static_cast<size_t>(M) * kMatrixWorkStride;
}
static constexpr size_t vector_floats() {
return static_cast<size_t>(kTaskWarps) * M * kKD;
}
static constexpr size_t scalar_floats() {
return static_cast<size_t>(kTaskWarps) * M * kScalarStride;
}
static constexpr size_t shared_bytes() {
return (work_floats() + 2 * vector_floats() + scalar_floats())
* sizeof(float)
+ static_cast<size_t>(M * kCounterIntsPerMatrix) * sizeof(int);
}
};
static_assert(Config<1, 8>::shared_bytes() == 72864,
"KD8 n1024 M1 shared contract changed");
static_assert(Config<2, 16>::shared_bytes() == 145728,
"KD8 n1024 M2 shared contract changed");
struct SharedState {
float* work;
float* v;
float* w;
float* scalar;
int* counters;
};
struct SubgroupState {
float* work;
float* v;
float* w;
float* scalar;
int* counters;
int matrix_local;
int matrix_lane;
bool active;
unsigned mask;
};
template <int M, int S>
__device__ __forceinline__ SharedState partition_shared(unsigned char* raw) {
SharedState state;
state.work = reinterpret_cast<float*>(raw);
state.v = state.work + Config<M, S>::work_floats();
state.w = state.v + Config<M, S>::vector_floats();
state.scalar = state.w + Config<M, S>::vector_floats();
state.counters =
reinterpret_cast<int*>(state.scalar + Config<M, S>::scalar_floats());
return state;
}
template <int M, int S>
__device__ __forceinline__ SubgroupState subgroup_state(
SharedState shared,
int warp,
int lane) {
SubgroupState state;
state.matrix_local = lane / S;
state.matrix_lane = lane % S;
state.active = state.matrix_local < M;
const int local = state.active ? state.matrix_local : 0;
state.mask = ((S == 32 ? 0xFFFFFFFFu : ((1u << S) - 1u)) << (local * S));
const int task_index = warp * M + local;
state.work = shared.work + local * kMatrixWorkStride;
state.v = shared.v + task_index * kKD;
state.w = shared.w + task_index * kKD;
state.scalar = shared.scalar + task_index * kScalarStride;
state.counters = shared.counters + local * kCounterIntsPerMatrix;
return state;
}
__device__ __forceinline__ float& lower(float* work, int row, int column) {
return work[column * kWorkRows + (row - column)];
}
__device__ __forceinline__ float symmetric_get(
const float* work,
int row,
int column) {
if (row >= column) {
return work[column * kWorkRows + (row - column)];
}
return work[row * kWorkRows + (column - row)];
}
__device__ __forceinline__ int slot_for(int sweep, int segment, int* error) {
if (sweep < 1 || sweep >= kN - 1) {
atomicExch(error, 11);
return -1;
}
const int begin = static_cast<int>(kSweepPrefix[sweep - 1]);
const int end = static_cast<int>(kSweepPrefix[sweep]);
if (segment < 0 || segment >= end - begin) {
atomicExch(error, 12);
return -1;
}
const int slot = begin + segment;
if (slot < 0 || slot >= kReflectors) {
atomicExch(error, 13);
return -1;
}
return slot;
}
template <int S>
__device__ __forceinline__ int subgroup_slot(
int sweep,
int segment,
SubgroupState state) {
int slot = -1;
if (state.matrix_lane == 0) {
slot = slot_for(sweep, segment, &state.counters[1]);
}
return __shfl_sync(state.mask, slot, state.matrix_local * S);
}
__device__ void load_reflector(
const float* hh,
int slot,
int length,
SubgroupState state) {
if (state.matrix_lane == 0) {
const float* record = hh + static_cast<size_t>(slot) * kRecordWidth;
state.scalar[0] = record[0];
state.v[0] = 1.0f;
for (int index = 1; index < length; ++index) {
state.v[index] = record[index];
}
for (int index = length; index < kKD; ++index) {
state.v[index] = 0.0f;
}
}
__syncwarp(state.mask);
}
__device__ void generate_reflector(
float* hh,
int slot,
int column,
int head,
int length,
SubgroupState state) {
if (state.matrix_lane == 0) {
float* record = hh + static_cast<size_t>(slot) * kRecordWidth;
const float alpha = lower(state.work, head, column);
float sigma = 0.0f;
for (int index = 1; index < length; ++index) {
const float value = lower(state.work, head + index, column);
state.v[index] = value;
sigma = fmaf(value, value, sigma);
}
float tau = 0.0f;
float beta = alpha;
float gamma = 0.0f;
if (sigma > 0.0f) {
beta = -copysignf(sqrtf(fmaf(alpha, alpha, sigma)), alpha);
tau = (beta - alpha) / beta;
gamma = 1.0f / (alpha - beta);
}
state.v[0] = 1.0f;
lower(state.work, head, column) = beta;
record[0] = tau;
for (int index = 1; index < kKD; ++index) {
const float tail = index < length ? gamma * state.v[index] : 0.0f;
state.v[index] = tail;
record[index] = tail;
if (index < length) {
lower(state.work, head + index, column) = 0.0f;
}
}
state.scalar[0] = tau;
}
__syncwarp(state.mask);
}
template <int S>
__device__ void apply_symmetric_patch(
int head,
int length,
SubgroupState state) {
const float tau = state.scalar[0];
for (int row = state.matrix_lane; row < length; row += S) {
float sum = 0.0f;
for (int column = 0; column < length; ++column) {
sum = fmaf(
symmetric_get(state.work, head + row, head + column),
state.v[column],
sum);
}
state.w[row] = tau * sum;
}
__syncwarp(state.mask);
if (state.matrix_lane < length) {
float dot = 0.0f;
for (int index = 0; index < length; ++index) {
dot = fmaf(state.v[index], state.w[index], dot);
}
const float correction = -0.5f * tau * dot;
for (int row = state.matrix_lane; row < length; row += S) {
state.w[row] = fmaf(correction, state.v[row], state.w[row]);
}
}
__syncwarp(state.mask);
const int elements = length * length;
for (int linear = state.matrix_lane; linear < elements; linear += S) {
const int row = linear / length;
const int column = linear - row * length;
if (row >= column) {
float& value = lower(state.work, head + row, head + column);
value = fmaf(-state.v[row], state.w[column], value);
value = fmaf(-state.w[row], state.v[column], value);
}
}
__syncwarp(state.mask);
}
template <int S>
__device__ void apply_right_patch(
int rows_begin,
int rows,
int columns_begin,
int columns,
SubgroupState state) {
for (int r = state.matrix_lane; r < rows; r += S) {
const int row = rows_begin + r;
float dot = 0.0f;
for (int index = 0; index < columns; ++index) {
dot = fmaf(
lower(state.work, row, columns_begin + index),
state.v[index],
dot);
}
const float weight = state.scalar[0] * dot;
for (int index = 0; index < columns; ++index) {
float& value = lower(state.work, row, columns_begin + index);
value = fmaf(-weight, state.v[index], value);
}
}
__syncwarp(state.mask);
}
template <int S>
__device__ void apply_left_patch(
int rows_begin,
int rows,
int columns_begin,
int columns,
SubgroupState state) {
for (int c = state.matrix_lane; c < columns; c += S) {
const int column = columns_begin + c;
float dot = 0.0f;
for (int index = 0; index < rows; ++index) {
dot = fmaf(
state.v[index],
lower(state.work, rows_begin + index, column),
dot);
}
const float weight = state.scalar[0] * dot;
for (int index = 0; index < rows; ++index) {
float& value = lower(state.work, rows_begin + index, column);
value = fmaf(-state.v[index], weight, value);
}
}
__syncwarp(state.mask);
}
__device__ __forceinline__ bool retires_first_sweep(int sweep, int task_id) {
const int kind = task_id == 1 ? 1 : task_id % 2 + 2;
int point;
int start;
int end;
int block_last;
if (kind == 2) {
point = (task_id / 2) * kKD + sweep;
start = point - kKD + 1;
end = min(point, kN);
block_last = point;
} else {
point = ((task_id + 1) / 2) * kKD + sweep;
start = point - kKD + 1;
end = min(point, kN);
block_last = start >= end - 1 && end == kN ? kN : 0;
}
return block_last >= kN - 1;
}
template <int M, int S>
__global__ __launch_bounds__(kThreads, 1) void sb2st_b8_slot_shuffle_f32(
const float* __restrict__ band,
float* __restrict__ d,
float* __restrict__ e,
float* __restrict__ hh,
int* __restrict__ info,
int batch) {
extern __shared__ unsigned char raw_shared[];
SharedState shared = partition_shared<M, S>(raw_shared);
const int cta_matrix_begin = static_cast<int>(blockIdx.x) * M;
if (cta_matrix_begin + M > batch) {
return;
}
const int tid = static_cast<int>(threadIdx.x);
const int lane = tid & 31;
const int warp = tid >> 5;
SubgroupState state = subgroup_state<M, S>(shared, warp, lane);
int local_tasks = 0;
int local_producers = 0;
int local_consumers = 0;
for (int index = tid;
index < static_cast<int>(Config<M, S>::work_floats());
index += kThreads) {
shared.work[index] = 0.0f;
}
for (int index = tid;
index < static_cast<int>(
2 * Config<M, S>::vector_floats() + Config<M, S>::scalar_floats());
index += kThreads) {
shared.v[index] = 0.0f;
}
for (int index = tid; index < M * kCounterIntsPerMatrix; index += kThreads) {
shared.counters[index] = 0;
}
const size_t d_begin = static_cast<size_t>(cta_matrix_begin) * kN;
const size_t e_begin = static_cast<size_t>(cta_matrix_begin) * (kN - 1);
const size_t hh_begin =
static_cast<size_t>(cta_matrix_begin) * kReflectors * kRecordWidth;
for (int index = tid; index < M * kN; index += kThreads) {
d[d_begin + index] = 0.0f;
}
for (int index = tid; index < M * (kN - 1); index += kThreads) {
e[e_begin + index] = 0.0f;
}
for (int index = tid;
index < M * kReflectors * kRecordWidth;
index += kThreads) {
hh[hh_begin + index] = 0.0f;
}
__syncthreads();
for (int index = tid; index < M * kInputRows * kN; index += kThreads) {
const int matrix = index / (kInputRows * kN);
const int remainder = index - matrix * kInputRows * kN;
const float value = band[
static_cast<size_t>(cta_matrix_begin + matrix) * kInputRows * kN
+ remainder];
const int delta = remainder / kN;
const int column = remainder - delta * kN;
shared.work[
matrix * kMatrixWorkStride + column * kWorkRows + delta] = value;
if (!isfinite(value)) {
atomicExch(&shared.counters[matrix * kCounterIntsPerMatrix + 1], 20);
}
}
__syncthreads();
float* matrix_hh = hh
+ static_cast<size_t>(cta_matrix_begin + state.matrix_local)
* kReflectors * kRecordWidth;
int first_sweep = 1;
for (int diagonal = 1; diagonal < kN; ++diagonal) {
const int end_sweep = min(diagonal, kN - 1);
if (first_sweep > end_sweep) {
break;
}
for (int step = 1; step <= 3; ++step) {
const int sweep_begin = first_sweep;
const int width = max(0, end_sweep - sweep_begin + 1);
if (state.active) {
for (int rank = warp; rank < width; rank += kTaskWarps) {
const int sweep = sweep_begin + rank;
const int task_id = 3 * (diagonal - sweep) + step;
const int kind = task_id == 1 ? 1 : task_id % 2 + 2;
int point;
int start;
int end;
if (kind == 2) {
point = (task_id / 2) * kKD + sweep;
start = point - kKD + 1;
end = min(point, kN);
} else {
point = ((task_id + 1) / 2) * kKD + sweep;
start = point - kKD + 1;
end = min(point, kN);
}
const int patch_begin = start - 1;
const int patch_end = end - 1;
const int patch_length = patch_end - patch_begin + 1;
if (state.matrix_lane == 0) {
++local_tasks;
if (patch_length < 1 || patch_length > kKD) {
atomicExch(&state.counters[1], 21);
}
}
__syncwarp(state.mask);
if (kind == 1) {
if (patch_length > 1) {
const int slot = subgroup_slot<S>(sweep, 0, state);
if (state.matrix_lane == 0) {
++local_producers;
}
if (slot >= 0 && state.counters[1] == 0) {
generate_reflector(
matrix_hh,
slot,
patch_begin - 1,
patch_begin,
patch_length,
state);
apply_symmetric_patch<S>(patch_begin, patch_length, state);
}
}
} else if (kind == 2) {
const int prior_segment = task_id / 2 - 1;
const int prior_slot =
subgroup_slot<S>(sweep, prior_segment, state);
if (state.matrix_lane == 0) {
++local_consumers;
}
if (prior_slot >= 0 && state.counters[1] == 0) {
load_reflector(matrix_hh, prior_slot, patch_length, state);
const int next_begin = patch_end + 1;
const int next_end = min(patch_end + kKD, kN - 1);
const int next_length = next_end - next_begin + 1;
if (next_length > 0) {
apply_right_patch<S>(
next_begin,
next_length,
patch_begin,
patch_length,
state);
}
if (next_length > 1) {
const int next_segment = task_id / 2;
const int next_slot =
subgroup_slot<S>(sweep, next_segment, state);
if (state.matrix_lane == 0) {
++local_producers;
}
if (next_slot >= 0 && state.counters[1] == 0) {
generate_reflector(
matrix_hh,
next_slot,
patch_begin,
next_begin,
next_length,
state);
apply_left_patch<S>(
next_begin,
next_length,
patch_begin + 1,
patch_length - 1,
state);
}
}
}
} else {
const int segment = (task_id - 1) / 2;
const int slot = subgroup_slot<S>(sweep, segment, state);
if (state.matrix_lane == 0) {
++local_consumers;
}
if (slot >= 0 && state.counters[1] == 0) {
load_reflector(matrix_hh, slot, patch_length, state);
apply_symmetric_patch<S>(patch_begin, patch_length, state);
}
}
}
}
__syncthreads();
const int first_task_id = 3 * (diagonal - sweep_begin) + step;
if (retires_first_sweep(sweep_begin, first_task_id)) {
++first_sweep;
}
}
}
// The literal serial model contains one terminal kind-1 task
// T(sweep=511, task_id=1) whose patch length is one. It creates no
// reflector, touches no work element, and has no consumer, while the
// retirement-form wave loop above exits just before instantiating it
// (verified by loop emulation: 32,960 loop tasks + this one = 32,961).
if (warp == 0 && state.matrix_lane == 0 && state.active) {
++local_tasks;
}
if (state.active && state.matrix_lane == 0) {
atomicAdd(&state.counters[0], local_tasks);
atomicAdd(&state.counters[2], local_producers);
atomicAdd(&state.counters[3], local_consumers);
}
__syncthreads();
for (int index = tid; index < M * kN; index += kThreads) {
const int matrix = index / kN;
const int column = index - matrix * kN;
const float value = shared.work[
matrix * kMatrixWorkStride + column * kWorkRows];
d[d_begin + index] = value;
if (!isfinite(value)) {
atomicExch(&shared.counters[matrix * kCounterIntsPerMatrix + 1], 40);
}
}
for (int index = tid; index < M * (kN - 1); index += kThreads) {
const int matrix = index / (kN - 1);
const int column = index - matrix * (kN - 1);
const float value = shared.work[
matrix * kMatrixWorkStride + column * kWorkRows + 1];
e[e_begin + index] = value;
if (!isfinite(value)) {
atomicExch(&shared.counters[matrix * kCounterIntsPerMatrix + 1], 40);
}
}
for (int index = tid;
index < M * kReflectors * kRecordWidth;
index += kThreads) {
const int matrix = index / (kReflectors * kRecordWidth);
if (!isfinite(hh[hh_begin + index])) {
atomicExch(&shared.counters[matrix * kCounterIntsPerMatrix + 1], 40);
}
}
__syncthreads();
if (tid < M) {
int* counters = shared.counters + tid * kCounterIntsPerMatrix;
int status = counters[1];
if (status == 0 && counters[0] != kTasks) {
status = 30;
}
if (status == 0 && counters[2] != kReflectors) {
status = 31;
}
if (status == 0 && counters[3] != kConsumers) {
status = 32;
}
info[cta_matrix_begin + tid] = status;
}
}
std::array<uint32_t, kSweepPrefixEntries> build_sweep_prefix() {
std::array<uint32_t, kSweepPrefixEntries> prefix{};
int total = 0;
prefix[0] = 0;
for (int sweep = 1; sweep <= kN - 2; ++sweep) {
const int remaining = kN - sweep - 1;
const int count = (remaining + kKD - 1) / kKD;
total += count;
TORCH_CHECK(total <= kReflectors, "KD8 sweep prefix overflow");
prefix[sweep] = static_cast<uint32_t>(total);
}
TORCH_CHECK(total == kReflectors, "KD8 reflector count mismatch");
return prefix;
}
void ensure_sweep_prefix(int device) {
static std::mutex mutex;
static std::unordered_set<int> initialized_devices;
std::lock_guard<std::mutex> lock(mutex);
if (initialized_devices.count(device) != 0) {
return;
}
const auto prefix = build_sweep_prefix();
C10_CUDA_CHECK(cudaMemcpyToSymbol(
kSweepPrefix,
prefix.data(),
prefix.size() * sizeof(prefix[0]),
0,
cudaMemcpyHostToDevice));
initialized_devices.insert(device);
}
void validate_tensors(
const torch::Tensor& band,
const torch::Tensor& d,
const torch::Tensor& e,
const torch::Tensor& hh,
const torch::Tensor& info,
int matrices_per_cta) {
TORCH_CHECK(band.is_cuda(), "band must be CUDA");
TORCH_CHECK(d.is_cuda() && e.is_cuda() && hh.is_cuda() && info.is_cuda(),
"all outputs must be CUDA");
TORCH_CHECK(band.scalar_type() == torch::kFloat32 &&
d.scalar_type() == torch::kFloat32 &&
e.scalar_type() == torch::kFloat32 &&
hh.scalar_type() == torch::kFloat32,
"band/d/e/hh must be float32");
TORCH_CHECK(info.scalar_type() == at::kInt, "info must be int32");
TORCH_CHECK(band.is_contiguous() && d.is_contiguous() && e.is_contiguous() &&
hh.is_contiguous() && info.is_contiguous(),
"all tensors must be contiguous");
TORCH_CHECK(
band.dim() == 3 && band.size(1) == kInputRows && band.size(2) == kN,
"band must have shape [B,9,512]");
const auto batch = band.size(0);
TORCH_CHECK(batch > 0 && batch % matrices_per_cta == 0,
"batch must be a positive multiple of matrices-per-CTA");
TORCH_CHECK(d.sizes() == torch::IntArrayRef({batch, kN}), "d shape mismatch");
TORCH_CHECK(e.sizes() == torch::IntArrayRef({batch, kN - 1}),
"e shape mismatch");
TORCH_CHECK(
hh.sizes() == torch::IntArrayRef({batch, kReflectors, kRecordWidth}),
"hh shape mismatch");
TORCH_CHECK(info.sizes() == torch::IntArrayRef({batch}),
"info shape mismatch");
const int device = band.get_device();
TORCH_CHECK(d.get_device() == device && e.get_device() == device &&
hh.get_device() == device && info.get_device() == device,
"all tensors must use one device");
}
template <int M, int S>
void launch_config(
const torch::Tensor& band,
const torch::Tensor& d,
const torch::Tensor& e,
const torch::Tensor& hh,
const torch::Tensor& info) {
validate_tensors(band, d, e, hh, info, M);
c10::cuda::CUDAGuard guard(band.device());
const int device = band.get_device();
ensure_sweep_prefix(device);
int maximum_shared = 0;
C10_CUDA_CHECK(cudaDeviceGetAttribute(
&maximum_shared, cudaDevAttrMaxSharedMemoryPerBlockOptin, device));
TORCH_CHECK(
maximum_shared >= static_cast<int>(Config<M, S>::shared_bytes()),
"device shared-memory limit is below this KD8 config's contract");
C10_CUDA_CHECK(cudaFuncSetAttribute(
sb2st_b8_slot_shuffle_f32<M, S>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
static_cast<int>(Config<M, S>::shared_bytes())));
const int batch = static_cast<int>(band.size(0));
sb2st_b8_slot_shuffle_f32<M, S><<<
batch / M,
kThreads,
Config<M, S>::shared_bytes()>>>(
band.data_ptr<float>(),
d.data_ptr<float>(),
e.data_ptr<float>(),
hh.data_ptr<float>(),
info.data_ptr<int>(),
batch);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
template <int M, int S>
std::vector<int64_t> report_config() {
C10_CUDA_CHECK(cudaFuncSetAttribute(
sb2st_b8_slot_shuffle_f32<M, S>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
static_cast<int>(Config<M, S>::shared_bytes())));
cudaFuncAttributes attributes{};
C10_CUDA_CHECK(cudaFuncGetAttributes(
&attributes, sb2st_b8_slot_shuffle_f32<M, S>));
int active_blocks = 0;
C10_CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&active_blocks,
sb2st_b8_slot_shuffle_f32<M, S>,
kThreads,
Config<M, S>::shared_bytes()));
int device = -1;
C10_CUDA_CHECK(cudaGetDevice(&device));
int sm_count = 0;
C10_CUDA_CHECK(cudaDeviceGetAttribute(
&sm_count, cudaDevAttrMultiProcessorCount, device));
return {
M,
S,
kThreads,
static_cast<int64_t>(Config<M, S>::shared_bytes()),
attributes.numRegs,
static_cast<int64_t>(attributes.localSizeBytes),
static_cast<int64_t>(attributes.sharedSizeBytes),
active_blocks,
sm_count,
};
}
} // namespace
void band8_to_tridiagonal_kd8_out(
const torch::Tensor& band,
const torch::Tensor& d,
const torch::Tensor& e,
const torch::Tensor& hh,
const torch::Tensor& info,
const std::string& config) {
if (config == "m1s8") {
launch_config<1, 8>(band, d, e, hh, info);
} else if (config == "m2s16") {
launch_config<2, 16>(band, d, e, hh, info);
} else {
TORCH_CHECK(false, "unknown KD8 config: ", config);
}
}
std::vector<int64_t> kd8_resource_report(const std::string& config) {
if (config == "m1s8") {
return report_config<1, 8>();
}
if (config == "m2s16") {
return report_config<2, 16>();
}
TORCH_CHECK(false, "unknown KD8 config: ", config);
return {};
}
} // namespace lane_n
namespace lane_o {
namespace {
#ifndef Q2_N
#define Q2_N 1024
#endif
constexpr int kN = Q2_N;
constexpr int q2_reflector_count(int n) {
int total = 0;
for (int sweep = 1; sweep <= n - 2; ++sweep) {
total += (n - 1 - sweep + 7) / 8;
}
return total;
}
constexpr int kKD = 8;
constexpr int kReflectors = q2_reflector_count(kN);
constexpr int kRecordWidth = kKD;
constexpr int kSweepPrefixEntries = kN - 1;
constexpr int kMaxSegments = (kN - 2 + 7) / 8; // ceil((kN-2)/kKD)
constexpr int kStageFloats = kMaxSegments * kRecordWidth;
static_assert(Q2_N != 1024 || kReflectors == 65792,
"reflector contract changed (n1024)");
static_assert(Q2_N != 1024 || 2 * kStageFloats * sizeof(float) == 8192,
"staging contract changed (n1024)");
__device__ __constant__ uint32_t kSweepPrefix[kSweepPrefixEntries];
// (d0,d1) = (a0*b0+c0, a1*b1+c1). On sm_100-class parts this is ONE
// issued instruction (Blackwell packed dual-FP32 pipe, mining P4 /
// source 845164 T8); each half is an IEEE fma with identical rounding to
// fmaf, so packing per se never changes results — only accumulation
// order chosen by the caller does. sm_120/121 (Spark GB10) is guarded
// to the scalar form: packed-f32x2 availability there is unproven and
// timing never transfers anyway.
__device__ __forceinline__ void fma2(
float& d0, float& d1,
float a0, float a1,
float b0, float b1,
float c0, float c1) {
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 1000 && __CUDA_ARCH__ < 1200
asm("{\n\t"
".reg .b64 ra, rb, rc, rd;\n\t"
"mov.b64 ra, {%2, %3};\n\t"
"mov.b64 rb, {%4, %5};\n\t"
"mov.b64 rc, {%6, %7};\n\t"
"fma.rn.f32x2 rd, ra, rb, rc;\n\t"
"mov.b64 {%0, %1}, rd;\n\t"
"}"
: "=f"(d0), "=f"(d1)
: "f"(a0), "f"(a1), "f"(b0), "f"(b1), "f"(c0), "f"(c1));
#else
d0 = fmaf(a0, b0, c0);
d1 = fmaf(a1, b1, c1);
#endif
}
// A window's <= 8 resident rows as NAMED scalars (donor probe: indexed
// locals demote to a stack frame; the scalarized struct promotes to
// registers in every config).
struct Win {
float r0, r1, r2, r3, r4, r5, r6, r7;
int length; // 0 = not yet born
};
// One owned window's step at sweep `sweep`, X accessed through the
// thread's global column pointer (gcol = x + matrix*kN*kN + column;
// element (row) at gcol[row * kN]). Op-for-op the donor's window_step
// with panel[row*kPanelCols+lane] replaced by gcol[row*kN]; the Packed
// path additionally reorders the dot into two packed chains.
template <bool Packed>
__device__ __forceinline__ void window_step_g(
float* __restrict__ gcol,
const float* sweep_records, // this sweep's contiguous records (shared)
int sweep,
int count,
int segment,
Win& w,
bool& bad) {
if (segment >= count) {
return;
}
const int head = sweep + kKD * segment;
const int length = min(kKD, kN - head);
if (w.length == 0) {
// Birth: clamped duplicates equal row kN-1 (pinned bottom during
// birth), shifted into validity as the window grows — no predicate.
w.r0 = gcol[static_cast<size_t>(min(head + 0, kN - 1)) * kN];
w.r1 = gcol[static_cast<size_t>(min(head + 1, kN - 1)) * kN];
w.r2 = gcol[static_cast<size_t>(min(head + 2, kN - 1)) * kN];
w.r3 = gcol[static_cast<size_t>(min(head + 3, kN - 1)) * kN];
w.r4 = gcol[static_cast<size_t>(min(head + 4, kN - 1)) * kN];
w.r5 = gcol[static_cast<size_t>(min(head + 5, kN - 1)) * kN];
w.r6 = gcol[static_cast<size_t>(min(head + 6, kN - 1)) * kN];
w.r7 = gcol[static_cast<size_t>(min(head + 7, kN - 1)) * kN];
bad = bad || !(isfinite(w.r0) && isfinite(w.r1) && isfinite(w.r2) &&
isfinite(w.r3) && isfinite(w.r4) && isfinite(w.r5) &&
isfinite(w.r6) && isfinite(w.r7));
} else {
// Slide up one row: shift the ring, read the entering top row
// (written the previous sweep by segment g-1's exiting bottom row;
// the per-sweep __syncthreads published it).
w.r7 = w.r6;
w.r6 = w.r5;
w.r5 = w.r4;
w.r4 = w.r3;
w.r3 = w.r2;
w.r2 = w.r1;
w.r1 = w.r0;
w.r0 = gcol[static_cast<size_t>(head) * kN];
bad = bad || !isfinite(w.r0);
}
w.length = length;
// Apply the reflector (v0 = 1 implicit).
const float* record = sweep_records + segment * kRecordWidth;
const float tau = record[0];
float dot;
if (Packed && length == kKD) {
// Steady state (all tail segments have length == 8 except the last
// of each sweep): two packed accumulation chains + combine.
float e = w.r0;
float o = record[1] * w.r1;
fma2(e, o, record[2], record[3], w.r2, w.r3, e, o);
fma2(e, o, record[4], record[5], w.r4, w.r5, e, o);
fma2(e, o, record[6], record[7], w.r6, w.r7, e, o);
dot = e + o;
} else {
dot = w.r0;
if (1 < length) dot = fmaf(record[1], w.r1, dot);
if (2 < length) dot = fmaf(record[2], w.r2, dot);
if (3 < length) dot = fmaf(record[3], w.r3, dot);
if (4 < length) dot = fmaf(record[4], w.r4, dot);
if (5 < length) dot = fmaf(record[5], w.r5, dot);
if (6 < length) dot = fmaf(record[6], w.r6, dot);
if (7 < length) dot = fmaf(record[7], w.r7, dot);
}
const float weight = tau * dot;
if (Packed && length == kKD) {
// Each half is the donor's exact fmaf(-weight, record[i], r_i):
// bitwise-identical update, half the issued FMA instructions.
const float mw = -weight;
fma2(w.r0, w.r1, mw, mw, 1.0f, record[1], w.r0, w.r1);
fma2(w.r2, w.r3, mw, mw, record[2], record[3], w.r2, w.r3);
fma2(w.r4, w.r5, mw, mw, record[4], record[5], w.r4, w.r5);
fma2(w.r6, w.r7, mw, mw, record[6], record[7], w.r6, w.r7);
} else {
w.r0 = fmaf(-weight, 1.0f, w.r0);
if (1 < length) w.r1 = fmaf(-weight, record[1], w.r1);
if (2 < length) w.r2 = fmaf(-weight, record[2], w.r2);
if (3 < length) w.r3 = fmaf(-weight, record[3], w.r3);
if (4 < length) w.r4 = fmaf(-weight, record[4], w.r4);
if (5 < length) w.r5 = fmaf(-weight, record[5], w.r5);
if (6 < length) w.r6 = fmaf(-weight, record[6], w.r6);
if (7 < length) w.r7 = fmaf(-weight, record[7], w.r7);
}
// Write the bottom row (exits next sweep in steady state; redundant
// but safe during birth, where the bottom is pinned at row kN-1).
float bottom = w.r7;
if (length == 1) bottom = w.r0;
if (length == 2) bottom = w.r1;
if (length == 3) bottom = w.r2;
if (length == 4) bottom = w.r3;
if (length == 5) bottom = w.r4;
if (length == 6) bottom = w.r5;
if (length == 7) bottom = w.r6;
gcol[static_cast<size_t>(head + length - 1) * kN] = bottom;
bad = bad || !isfinite(bottom);
}
__device__ __forceinline__ void window_flush_g(
float* __restrict__ gcol,
int segment,
const Win& w,
bool& bad) {
if (w.length == 0) {
return;
}
const int head = 1 + kKD * segment;
const int length = w.length;
bool ok = true;
if (0 < length - 1) {
gcol[static_cast<size_t>(head + 0) * kN] = w.r0;
ok = ok && isfinite(w.r0);
}
if (1 < length - 1) {
gcol[static_cast<size_t>(head + 1) * kN] = w.r1;
ok = ok && isfinite(w.r1);
}
if (2 < length - 1) {
gcol[static_cast<size_t>(head + 2) * kN] = w.r2;
ok = ok && isfinite(w.r2);
}
if (3 < length - 1) {
gcol[static_cast<size_t>(head + 3) * kN] = w.r3;
ok = ok && isfinite(w.r3);
}
if (4 < length - 1) {
gcol[static_cast<size_t>(head + 4) * kN] = w.r4;
ok = ok && isfinite(w.r4);
}
if (5 < length - 1) {
gcol[static_cast<size_t>(head + 5) * kN] = w.r5;
ok = ok && isfinite(w.r5);
}
if (6 < length - 1) {
gcol[static_cast<size_t>(head + 6) * kN] = w.r6;
ok = ok && isfinite(w.r6);
}
bad = bad || !ok;
}
// Per-warp staging prefetch (Neighbor mode): warp w copies ONLY its own
// segments' records for sweep `sweep_target` into the shared staging
// buffer. No other warp ever reads those slots, so the double buffer
// is safe under ANY ring skew (the CTA-wide cooperative copy is only
// safe under the CTA-wide barrier). <= 8 owned segments per warp = 16
// float4 halves, one per lane 0..15.
template <int Cols, int Threads>
__device__ __forceinline__ void warp_prefetch(
const float* __restrict__ matrix_hh,
float* __restrict__ dst,
int sweep_target,
int warp,
int lane32) {
constexpr int kGroups = Threads / Cols;
constexpr int kGroupsPerWarp = 32 / Cols;
constexpr int kWindows = (kMaxSegments + kGroups - 1) / kGroups;
constexpr int kSlots = kGroupsPerWarp * kWindows;
static_assert(2 * kSlots <= 32, "one float4 half per lane");
const int count = (kN - 1 - sweep_target + kKD - 1) / kKD;
const float4* src = reinterpret_cast<const float4*>(
matrix_hh
+ static_cast<size_t>(kSweepPrefix[sweep_target - 1]) * kRecordWidth);
if (lane32 < 2 * kSlots) {
const int slot = lane32 >> 1;
const int half = lane32 & 1;
const int group_local = slot / kWindows;
const int k = slot % kWindows;
const int segment = warp * kGroupsPerWarp + group_local + k * kGroups;
if (segment < count) {
reinterpret_cast<float4*>(dst + segment * kRecordWidth)[half] =
__ldcs(src + segment * 2 + half);
}
}
}
// Cols = X columns per CTA (register/occupancy knob): each of the
// kGroups = Threads/Cols lane groups owns segments {group + k*kGroups}
// and carries them in kWindows register rings. MinBlocks is the P11
// register-budget knob (__launch_bounds__ minBlocksPerMultiprocessor).
// Neighbor replaces the per-sweep CTA barrier with a shared-memory
// progress-flag WAVEFRONT: segment g's entering row is written ONLY by
// segment g-1 = an adjacent group = the same warp (covered by the
// end-of-sweep __syncwarp) or the PRECEDING warp; birth's pinned row
// kN-1 transfers deepest-segment ownership between adjacent segments
// too, so the only cross-warp dependency is one-directional
// wait-for-predecessor (no WAR hazard at ANY skew: each handoff row is
// written exactly once by pred and read exactly once by the owner, and
// flush rows are disjoint — design note in COST.md). Named-barrier
// rendezvous was the first form and is REFUTED at build time: 16
// barrier ids/CTA allocate the SM's whole barrier file -> occupancy
// back to 1 block/SM (Spark build report), destroying the very lever
// this lane bought. Flags cost zero barrier resources and give
// unbounded pipeline slack instead of lockstep rendezvous.
template <int Cols, int Threads, int MinBlocks, bool Packed,
bool Neighbor = false>
__global__ __launch_bounds__(Threads, MinBlocks)
void bt_kd8_gapply_f32(
float* __restrict__ x,
const float* __restrict__ hh,
int* __restrict__ info,
int batch) {
static_assert(Cols == 8 || Cols == 16 || Cols == 32, "lane group width");
static_assert(Threads % Cols == 0, "whole groups");
static_assert(kN % Cols == 0, "whole panels");
constexpr int kGroups = Threads / Cols;
constexpr int kWindows = (kMaxSegments + kGroups - 1) / kGroups;
constexpr int kWarps = Threads / 32;
static_assert(kWindows * kGroups >= kMaxSegments, "ownership coverage");
static_assert(kWindows <= 8, "window ring capacity");
static_assert(!Neighbor || kWarps <= 32, "progress flag array");
__shared__ __align__(16) float stage[2][kStageFloats];
static_assert((kStageFloats * sizeof(float)) % 16 == 0,
"float4 staging alignment");
// Wavefront progress: sweeps completed per warp (Neighbor mode).
__shared__ volatile int progress[Neighbor ? kWarps : 1];
const int matrix = static_cast<int>(blockIdx.y);
if (matrix >= batch) {
return;
}
const int tid = static_cast<int>(threadIdx.x);
const int lane = tid % Cols;
const int group = tid / Cols;
const int column = static_cast<int>(blockIdx.x) * Cols + lane;
float* gcol = x + static_cast<size_t>(matrix) * kN * kN + column;
const float* matrix_hh =
hh + static_cast<size_t>(matrix) * kReflectors * kRecordWidth;
// Row 0 is never touched by any reflector (head >= 1); check it for
// the input-finiteness contract. All other rows are checked at their
// first window read; outputs are checked at every write.
bool bad = (group == 0) ? !isfinite(gcol[0]) : false;
// Register window rings, one per owned segment slot.
Win w0{}, w1{}, w2{}, w3{}, w4{}, w5{}, w6{}, w7{};
const int warp = tid >> 5;
const int lane32 = tid & 31;
const int pred = (warp + kWarps - 1) % kWarps;
// Preload sweep kN-2's records into staging buffer 0.
if (Neighbor) {
if (tid < kWarps) {
progress[tid] = 0;
}
warp_prefetch<Cols, Threads>(matrix_hh, stage[0], kN - 2, warp, lane32);
__syncthreads(); // once: publish flag init (barrier 0 only)
} else {
const float* src = matrix_hh
+ static_cast<size_t>(kSweepPrefix[kN - 3]) * kRecordWidth;
for (int i = tid; i < kRecordWidth; i += Threads) {
stage[0][i] = src[i];
}
__syncthreads();
}
int buffer = 0;
for (int sweep = kN - 2; sweep >= 1; --sweep) {
const int count = (kN - 1 - sweep + kKD - 1) / kKD;
const float* sweep_records = stage[buffer];
if (Neighbor) {
// Wait until the predecessor warp has COMPLETED the previous
// sweep (its bottom-row writes are this warp's entering rows).
const int need = kN - 2 - sweep; // iterations pred must have done
if (need > 0) {
if (lane32 == 0) {
while (progress[pred] < need) {
__nanosleep(20);
}
}
__syncwarp();
__threadfence_block(); // acquire pred's global writes
}
}
if (sweep > 1) {
if (Neighbor) {
// Warp-owned prefetch: safe under ring skew (see warp_prefetch).
warp_prefetch<Cols, Threads>(matrix_hh, stage[buffer ^ 1],
sweep - 1, warp, lane32);
} else {
// Prefetch the NEXT sweep's contiguous record block: float4 +
// read-once hint (records are read once per CTA, never reused —
// keep L1 for the live X band). Bitwise-neutral vs scalar copy.
const int next_count = (kN - 1 - (sweep - 1) + kKD - 1) / kKD;
const float4* src = reinterpret_cast<const float4*>(
matrix_hh
+ static_cast<size_t>(kSweepPrefix[sweep - 2]) * kRecordWidth);
float4* destination = reinterpret_cast<float4*>(stage[buffer ^ 1]);
const int quads = next_count * (kRecordWidth / 4);
for (int i = tid; i < quads; i += Threads) {
destination[i] = __ldcs(src + i);
}
}
}
window_step_g<Packed>(gcol, sweep_records, sweep, count,
group + 0 * kGroups, w0, bad);
if (kWindows >= 2) {
window_step_g<Packed>(gcol, sweep_records, sweep, count,
group + 1 * kGroups, w1, bad);
}
if (kWindows >= 3) {
window_step_g<Packed>(gcol, sweep_records, sweep, count,
group + 2 * kGroups, w2, bad);
}
if (kWindows >= 4) {
window_step_g<Packed>(gcol, sweep_records, sweep, count,
group + 3 * kGroups, w3, bad);
}
if (kWindows >= 5) {
window_step_g<Packed>(gcol, sweep_records, sweep, count,
group + 4 * kGroups, w4, bad);
}
if (kWindows >= 6) {
window_step_g<Packed>(gcol, sweep_records, sweep, count,
group + 5 * kGroups, w5, bad);
}
if (kWindows >= 7) {
window_step_g<Packed>(gcol, sweep_records, sweep, count,
group + 6 * kGroups, w6, bad);
}
if (kWindows >= 8) {
window_step_g<Packed>(gcol, sweep_records, sweep, count,
group + 7 * kGroups, w7, bad);
}
if (Neighbor) {
// Release own bottom-row writes CTA-wide, then publish progress.
// The __syncwarp also carries the intra-warp group edge (lanes
// of group 2w+1 read group 2w's handoff row next sweep under
// independent thread scheduling) and the warp-owned staging.
__threadfence_block();
__syncwarp();
if (lane32 == 0) {
progress[warp] = kN - 1 - sweep; // completed iterations
}
} else {
__syncthreads();
}
buffer ^= 1;
}
// Final flush: windows hold rows head..head+length-2 not yet written
// (bottom row was written during the sweep-1 step).
window_flush_g(gcol, group + 0 * kGroups, w0, bad);
if (kWindows >= 2) {
window_flush_g(gcol, group + 1 * kGroups, w1, bad);
}
if (kWindows >= 3) {
window_flush_g(gcol, group + 2 * kGroups, w2, bad);
}
if (kWindows >= 4) {
window_flush_g(gcol, group + 3 * kGroups, w3, bad);
}
if (kWindows >= 5) {
window_flush_g(gcol, group + 4 * kGroups, w4, bad);
}
if (kWindows >= 6) {
window_flush_g(gcol, group + 5 * kGroups, w5, bad);
}
if (kWindows >= 7) {
window_flush_g(gcol, group + 6 * kGroups, w6, bad);
}
if (kWindows >= 8) {
window_flush_g(gcol, group + 7 * kGroups, w7, bad);
}
if (bad) {
atomicExch(&info[matrix], 40);
}
}
std::array<uint32_t, kSweepPrefixEntries> build_sweep_prefix() {
std::array<uint32_t, kSweepPrefixEntries> prefix{};
int total = 0;
prefix[0] = 0;
for (int sweep = 1; sweep <= kN - 2; ++sweep) {
const int remaining = kN - sweep - 1;
total += (remaining + kKD - 1) / kKD;
TORCH_CHECK(total <= kReflectors, "KD8 sweep prefix overflow");
prefix[sweep] = static_cast<uint32_t>(total);
}
TORCH_CHECK(total == kReflectors, "KD8 reflector count mismatch");
return prefix;
}
void ensure_sweep_prefix(int device) {
static std::mutex mutex;
static std::unordered_set<int> initialized_devices;
std::lock_guard<std::mutex> lock(mutex);
if (initialized_devices.count(device) != 0) {
return;
}
const auto prefix = build_sweep_prefix();
C10_CUDA_CHECK(cudaMemcpyToSymbol(
kSweepPrefix,
prefix.data(),
prefix.size() * sizeof(prefix[0]),
0,
cudaMemcpyHostToDevice));
initialized_devices.insert(device);
}
void validate_tensors(
const torch::Tensor& x,
const torch::Tensor& hh,
const torch::Tensor& info) {
TORCH_CHECK(x.is_cuda() && hh.is_cuda() && info.is_cuda(),
"all tensors must be CUDA");
TORCH_CHECK(x.scalar_type() == torch::kFloat32 &&
hh.scalar_type() == torch::kFloat32,
"x/hh must be float32");
TORCH_CHECK(info.scalar_type() == at::kInt, "info must be int32");
TORCH_CHECK(x.is_contiguous() && hh.is_contiguous() && info.is_contiguous(),
"all tensors must be contiguous");
TORCH_CHECK(x.dim() == 3 && x.size(1) == kN && x.size(2) == kN,
"x shape mismatch");
const auto batch = x.size(0);
TORCH_CHECK(
hh.sizes() == torch::IntArrayRef({batch, kReflectors, kRecordWidth}),
"hh shape mismatch");
TORCH_CHECK(info.sizes() == torch::IntArrayRef({batch}),
"info shape mismatch");
const int device = x.get_device();
TORCH_CHECK(hh.get_device() == device && info.get_device() == device,
"all tensors must use one device");
}
template <int Cols, int Threads, int MinBlocks, bool Packed,
bool Neighbor = false>
void launch_config(
const torch::Tensor& x,
const torch::Tensor& hh,
const torch::Tensor& info) {
validate_tensors(x, hh, info);
c10::cuda::CUDAGuard guard(x.device());
const int device = x.get_device();
ensure_sweep_prefix(device);
// Static shared only (8,192 B at n1024): no opt-in attribute needed,
// which also sidesteps the GB10 silent setAttribute clamp receipt.
const int batch = static_cast<int>(x.size(0));
dim3 grid(kN / Cols, batch, 1);
bt_kd8_gapply_f32<Cols, Threads, MinBlocks, Packed, Neighbor>
<<<grid, Threads, 0>>>(
x.data_ptr<float>(),
hh.data_ptr<float>(),
info.data_ptr<int>(),
batch);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
template <int Cols, int Threads, int MinBlocks, bool Packed,
bool Neighbor = false>
std::vector<int64_t> report_config() {
cudaFuncAttributes attributes{};
C10_CUDA_CHECK(cudaFuncGetAttributes(
&attributes, bt_kd8_gapply_f32<Cols, Threads, MinBlocks, Packed, Neighbor>));
int active_blocks = 0;
C10_CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&active_blocks, bt_kd8_gapply_f32<Cols, Threads, MinBlocks, Packed, Neighbor>,
Threads, 0));
int device = -1;
C10_CUDA_CHECK(cudaGetDevice(&device));
int sm_count = 0;
C10_CUDA_CHECK(cudaDeviceGetAttribute(
&sm_count, cudaDevAttrMultiProcessorCount, device));
return {
Threads,
0, // dynamic shared
attributes.numRegs,
static_cast<int64_t>(attributes.localSizeBytes),
static_cast<int64_t>(attributes.sharedSizeBytes),
active_blocks,
sm_count,
};
}
} // namespace
// Config table (COST.md design table):
// g32b1 Cols=32 MinB=1 scalar — un-stage-only attribution control
// g32b2 Cols=32 MinB=2 scalar — P11 register squeeze on the control
// g16b2 Cols=16 MinB=2 scalar — PRIMARY (2 CTA/SM by arithmetic)
// g16b3 Cols=16 MinB=3 scalar — P11 probe toward 3 CTA/SM
// g8b4 Cols=8 MinB=4 scalar — max-occupancy arm
// g16b2p / g8b4p — P4 packed-f32x2 arms of the two above
// Bounded round (COST.md addendum, B200 job 3075681 landed 19.476 in
// (19,24]): 'n' suffix = named-barrier neighbor sync:
// g16b2n — scalar + neighbor (bitwise attribution arm)
// g16b2pn — the round's arm (incumbent g16b2p + neighbor)
// g8b4pn — occupancy x neighbor interaction arm
void backtransform_kd8_inplace(
const torch::Tensor& x,
const torch::Tensor& hh,
const torch::Tensor& info,
const std::string& config) {
if (config == "g32b1") {
launch_config<32, 512, 1, false>(x, hh, info);
} else if (config == "g32b2") {
launch_config<32, 512, 2, false>(x, hh, info);
} else if (config == "g16b2") {
launch_config<16, 512, 2, false>(x, hh, info);
} else if (config == "g16b3") {
launch_config<16, 512, 3, false>(x, hh, info);
} else if (config == "g8b4") {
launch_config<8, 512, 4, false>(x, hh, info);
} else if (config == "g16b2p") {
launch_config<16, 512, 2, true>(x, hh, info);
} else if (config == "g8b4p") {
launch_config<8, 512, 4, true>(x, hh, info);
} else if (config == "g16b2n") {
launch_config<16, 512, 2, false, true>(x, hh, info);
} else if (config == "g16b2pn") {
launch_config<16, 512, 2, true, true>(x, hh, info);
} else if (config == "g8b4pn") {
launch_config<8, 512, 4, true, true>(x, hh, info);
} else {
TORCH_CHECK(false, "unknown back-transform config: ", config);
}
}
std::vector<int64_t> bt_resource_report(const std::string& config) {
if (config == "g32b1") {
return report_config<32, 512, 1, false>();
}
if (config == "g32b2") {
return report_config<32, 512, 2, false>();
}
if (config == "g16b2") {
return report_config<16, 512, 2, false>();
}
if (config == "g16b3") {
return report_config<16, 512, 3, false>();
}
if (config == "g8b4") {
return report_config<8, 512, 4, false>();
}
if (config == "g16b2p") {
return report_config<16, 512, 2, true>();
}
if (config == "g8b4p") {
return report_config<8, 512, 4, true>();
}
if (config == "g16b2n") {
return report_config<16, 512, 2, false, true>();
}
if (config == "g16b2pn") {
return report_config<16, 512, 2, true, true>();
}
if (config == "g8b4pn") {
return report_config<8, 512, 4, true, true>();
}
TORCH_CHECK(false, "unknown back-transform config: ", config);
return {};
}
} // namespace lane_o
namespace lane_p {
namespace {
constexpr int kWidth = 1024;
constexpr int kTiles = 4;
constexpr int kTileWidth = 256;
constexpr int kFp32Iterations = 30;
constexpr int kFp64Iterations = 80;
constexpr int32_t kModeCommon = 0;
constexpr int32_t kModeSelectiveStatus2 = 1;
constexpr int32_t kModeFullRetry = 2;
constexpr int32_t kStatusSuccess = 0;
constexpr int32_t kStatusInvalidCount = 1;
constexpr int32_t kStatusBracket = 2;
constexpr int32_t kStatusNoConvergence = 5;
__device__ __forceinline__ void kahan_add(
float value, float& total, float& compensation) {
const float adjusted = value - compensation;
const float updated = total + adjusted;
compensation = (updated - total) - adjusted;
total = updated;
}
__device__ __forceinline__ void kahan_add_double(
double value, double& total, double& compensation) {
const double adjusted = value - compensation;
const double updated = total + adjusted;
compensation = (updated - total) - adjusted;
total = updated;
}
__device__ __forceinline__ float secular_value(
const float* poles,
const float* weights,
int count,
float rho,
float x) {
float negative = 0.0f;
float positive = 0.0f;
float cnegative = 0.0f;
float cpositive = 0.0f;
for (int index = 0; index < count; ++index) {
const float term =
rho * weights[index] * weights[index] / (poles[index] - x);
if (term < 0.0f) {
kahan_add(term, negative, cnegative);
} else {
kahan_add(term, positive, cpositive);
}
}
return 1.0f + (negative + positive);
}
// Preserve the width-512 endpoint predicate's arithmetic and accumulation
// order; only the shared-array extent and owner index are widened to 1024.
__device__ __forceinline__ double endpoint_secular_value_double(
const float* poles,
const float* weights,
int count,
float rho,
double x) {
double negative = 0.0;
double positive = 0.0;
double cnegative = 0.0;
double cpositive = 0.0;
for (int index = 0; index < count; ++index) {
const double pole = static_cast<double>(poles[index]);
const double weight = static_cast<double>(weights[index]);
const double term =
static_cast<double>(rho) * weight * weight / (pole - x);
if (term < 0.0) {
kahan_add_double(term, negative, cnegative);
} else {
kahan_add_double(term, positive, cpositive);
}
}
return 1.0 + negative + positive;
}
struct EvalFloat {
float value;
float derivative;
float error_scale;
};
__device__ __forceinline__ EvalFloat secular_eval_float(
const float* poles,
const float* weights,
int count,
float rho,
float x) {
float negative = 0.0f;
float positive = 0.0f;
float derivative = 0.0f;
float magnitude = 0.0f;
float cnegative = 0.0f;
float cpositive = 0.0f;
float cderivative = 0.0f;
float cmagnitude = 0.0f;
for (int index = 0; index < count; ++index) {
const float delta = poles[index] - x;
const float weight2 = weights[index] * weights[index];
const float term = rho * weight2 / delta;
if (term < 0.0f) {
kahan_add(term, negative, cnegative);
} else {
kahan_add(term, positive, cpositive);
}
kahan_add(rho * weight2 / (delta * delta), derivative, cderivative);
kahan_add(fabsf(term), magnitude, cmagnitude);
}
return {
1.0f + negative + positive,
derivative,
8.0f * (1.0f + magnitude + fabsf(x) * derivative)};
}
__device__ __forceinline__ float interior_rational_step_float(
const float* poles,
const float* weights,
float rho,
int index,
float x,
float value,
float derivative,
bool origin_at_lower) {
const float delta_i = poles[index] - x;
const float delta_ip1 = poles[index + 1] - x;
const float gap = poles[index + 1] - poles[index];
float c;
if (origin_at_lower) {
const float ratio = weights[index] / delta_i;
c = value - delta_ip1 * derivative + gap * rho * ratio * ratio;
} else {
const float ratio = weights[index + 1] / delta_ip1;
c = value - delta_i * derivative - gap * rho * ratio * ratio;
}
const float a =
(delta_i + delta_ip1) * value - delta_i * delta_ip1 * derivative;
const float b = delta_i * delta_ip1 * value;
if (c == 0.0f) {
return a == 0.0f ? CUDART_NAN_F : b / a;
}
const float discriminant = fmaxf(0.0f, a * a - 4.0f * b * c);
const float root = sqrtf(discriminant);
if (a <= 0.0f) {
return (a - root) / (2.0f * c);
}
const float denominator = a + root;
return denominator == 0.0f ? CUDART_NAN_F : 2.0f * b / denominator;
}
__device__ __forceinline__ float last_rational_step_float(
const float* poles,
const float* weights,
int count,
float rho,
float x,
float value) {
const float delta0 = poles[count - 2] - x;
const float delta1 = poles[count - 1] - x;
float dpsi = 0.0f;
float correction = 0.0f;
for (int index = 0; index < count - 1; ++index) {
const float ratio = weights[index] / (poles[index] - x);
kahan_add(rho * ratio * ratio, dpsi, correction);
}
const float last_ratio = weights[count - 1] / delta1;
const float dphi = rho * last_ratio * last_ratio;
const float c = fabsf(value - delta0 * dpsi - delta1 * dphi);
const float a =
(delta0 + delta1) * value - delta0 * delta1 * (dpsi + dphi);
const float b = delta0 * delta1 * value;
if (c == 0.0f) {
return CUDART_NAN_F;
}
const float discriminant = fmaxf(0.0f, a * a - 4.0f * b * c);
const float root = sqrtf(discriminant);
if (a >= 0.0f) {
return (a + root) / (2.0f * c);
}
const float denominator = a - root;
return denominator == 0.0f ? CUDART_NAN_F : 2.0f * b / denominator;
}
__device__ __forceinline__ double secular_value_double(
const double* poles,
const double* weights,
int count,
double rho,
double x) {
double negative = 0.0;
double positive = 0.0;
double cnegative = 0.0;
double cpositive = 0.0;
for (int index = 0; index < count; ++index) {
const double term =
rho * weights[index] * weights[index] / (poles[index] - x);
if (term < 0.0) {
kahan_add_double(term, negative, cnegative);
} else {
kahan_add_double(term, positive, cpositive);
}
}
return 1.0 + (negative + positive);
}
struct EvalDouble {
double value;
double derivative;
double error_scale;
};
__device__ __forceinline__ EvalDouble secular_eval_double(
const double* poles,
const double* weights,
int count,
double rho,
double x) {
double negative = 0.0;
double positive = 0.0;
double derivative = 0.0;
double magnitude = 0.0;
double cnegative = 0.0;
double cpositive = 0.0;
double cderivative = 0.0;
double cmagnitude = 0.0;
for (int index = 0; index < count; ++index) {
const double delta = poles[index] - x;
const double weight2 = weights[index] * weights[index];
const double term = rho * weight2 / delta;
if (term < 0.0) {
kahan_add_double(term, negative, cnegative);
} else {
kahan_add_double(term, positive, cpositive);
}
kahan_add_double(
rho * weight2 / (delta * delta), derivative, cderivative);
kahan_add_double(fabs(term), magnitude, cmagnitude);
}
return {
1.0 + negative + positive,
derivative,
8.0 * (1.0 + magnitude + fabs(x) * derivative)};
}
__device__ __forceinline__ double interior_rational_step_double(
const double* poles,
const double* weights,
double rho,
int index,
double x,
double value,
double derivative,
bool origin_at_lower) {
const double delta_i = poles[index] - x;
const double delta_ip1 = poles[index + 1] - x;
const double gap = poles[index + 1] - poles[index];
double c;
if (origin_at_lower) {
const double ratio = weights[index] / delta_i;
c = value - delta_ip1 * derivative + gap * rho * ratio * ratio;
} else {
const double ratio = weights[index + 1] / delta_ip1;
c = value - delta_i * derivative - gap * rho * ratio * ratio;
}
const double a =
(delta_i + delta_ip1) * value - delta_i * delta_ip1 * derivative;
const double b = delta_i * delta_ip1 * value;
if (c == 0.0) {
return a == 0.0 ? CUDART_NAN : b / a;
}
const double root = sqrt(fmax(0.0, a * a - 4.0 * b * c));
if (a <= 0.0) {
return (a - root) / (2.0 * c);
}
const double denominator = a + root;
return denominator == 0.0 ? CUDART_NAN : 2.0 * b / denominator;
}
__device__ __forceinline__ double last_rational_step_double(
const double* poles,
const double* weights,
int count,
double rho,
double x,
double value) {
const double delta0 = poles[count - 2] - x;
const double delta1 = poles[count - 1] - x;
double dpsi = 0.0;
double correction = 0.0;
for (int index = 0; index < count - 1; ++index) {
const double ratio = weights[index] / (poles[index] - x);
kahan_add_double(rho * ratio * ratio, dpsi, correction);
}
const double last_ratio = weights[count - 1] / delta1;
const double dphi = rho * last_ratio * last_ratio;
const double c = fabs(value - delta0 * dpsi - delta1 * dphi);
const double a =
(delta0 + delta1) * value - delta0 * delta1 * (dpsi + dphi);
const double b = delta0 * delta1 * value;
if (c == 0.0) {
return CUDART_NAN;
}
const double root = sqrt(fmax(0.0, a * a - 4.0 * b * c));
if (a >= 0.0) {
return (a + root) / (2.0 * c);
}
const double denominator = a - root;
return denominator == 0.0 ? CUDART_NAN : 2.0 * b / denominator;
}
__global__ void clear_small_route_state_kernel(
const int32_t* __restrict__ active_count,
int32_t* __restrict__ merge_mode,
int32_t* __restrict__ repair_count,
int32_t* __restrict__ info,
int32_t* __restrict__ maximum_iterations,
int batch) {
const int merge = static_cast<int>(blockIdx.x) * blockDim.x + threadIdx.x;
if (merge >= batch) {
return;
}
const int count = active_count[merge];
merge_mode[merge] = kModeCommon;
info[merge] = (count > 0 && count <= kWidth) ? 0 : kStatusInvalidCount;
maximum_iterations[merge] = 0;
#pragma unroll
for (int tile = 0; tile < kTiles; ++tile) {
repair_count[merge * kTiles + tile] = 0;
}
}
__global__ __launch_bounds__(kTileWidth) void fp32_roots_4x256_kernel(
const float* __restrict__ poles,
const float* __restrict__ weights,
const float* __restrict__ rho,
const int32_t* __restrict__ active_count,
float* __restrict__ values,
double* __restrict__ roots64,
int32_t* __restrict__ root_status,
int32_t* __restrict__ root_iterations,
int32_t* __restrict__ merge_mode,
int32_t* __restrict__ deflated,
int32_t* __restrict__ maximum_iterations) {
const int merge = static_cast<int>(blockIdx.x);
const int tile = static_cast<int>(blockIdx.y);
const int thread = static_cast<int>(threadIdx.x);
const int root = tile * kTileWidth + thread;
const int base = merge * kWidth;
const int count = active_count[merge];
__shared__ float shared_poles[kWidth];
__shared__ float shared_weights[kWidth];
#pragma unroll
for (int round = 0; round < kTiles; ++round) {
const int index = round * kTileWidth + thread;
shared_poles[index] = poles[base + index];
shared_weights[index] = weights[base + index];
}
// Every owner initializes all prefix-visible per-root state. K0 never
// clears these large arrays, so stale state cannot masquerade as a route.
values[base + root] = 0.0f;
roots64[base + root] = CUDART_NAN;
root_status[base + root] = kStatusSuccess;
root_iterations[base + root] = 0;
deflated[base + root] = 0;
__syncthreads();
if (!(count > 0 && count <= kWidth) || root >= count) {
return;
}
const float positive_inf = CUDART_INF_F;
const float negative_inf = -CUDART_INF_F;
const float rho_value = rho[merge];
float lo = nextafterf(shared_poles[root], positive_inf);
float hi;
if (root + 1 < count) {
hi = nextafterf(shared_poles[root + 1], negative_inf);
} else {
float normz2 = 0.0f;
float correction = 0.0f;
for (int index = 0; index < count; ++index) {
kahan_add(
shared_weights[index] * shared_weights[index], normz2, correction);
}
hi = shared_poles[count - 1] + rho_value * normz2;
if (!(hi > lo)) {
hi = nextafterf(lo, positive_inf);
}
for (int attempt = 0; attempt < 8; ++attempt) {
if (secular_value(
shared_poles, shared_weights, count, rho_value, hi) >= 0.0f) {
break;
}
hi = shared_poles[count - 1] +
2.0f * (hi - shared_poles[count - 1]);
}
}
const float flo = secular_eval_float(
shared_poles, shared_weights, count, rho_value, lo).value;
const float fhi = secular_eval_float(
shared_poles, shared_weights, count, rho_value, hi).value;
float result = CUDART_NAN_F;
int32_t status = kStatusSuccess;
int32_t used = 0;
if (!(lo < hi) || !(flo <= 0.0f) || !(fhi >= 0.0f) ||
!isfinite(lo) || !isfinite(hi)) {
bool endpoint_sign_rescued = false;
if (lo < hi && isfinite(lo) && isfinite(hi) &&
(flo > 0.0f || fhi < 0.0f)) {
const double double_lo = nextafter(
static_cast<double>(shared_poles[root]), CUDART_INF);
double double_hi;
if (root + 1 < count) {
double_hi = nextafter(
static_cast<double>(shared_poles[root + 1]), -CUDART_INF);
} else {
double normz2 = 0.0;
double correction = 0.0;
for (int index = 0; index < count; ++index) {
const double weight = static_cast<double>(shared_weights[index]);
kahan_add_double(weight * weight, normz2, correction);
}
double_hi = static_cast<double>(shared_poles[count - 1]) +
static_cast<double>(rho_value) * normz2;
if (!(double_hi > double_lo)) {
double_hi = nextafter(double_lo, CUDART_INF);
}
for (int attempt = 0; attempt < 8; ++attempt) {
if (endpoint_secular_value_double(
shared_poles,
shared_weights,
count,
rho_value,
double_hi) >= 0.0) {
break;
}
double_hi = static_cast<double>(shared_poles[count - 1]) +
2.0 *
(double_hi -
static_cast<double>(shared_poles[count - 1]));
}
}
const double double_flo = endpoint_secular_value_double(
shared_poles, shared_weights, count, rho_value, double_lo);
const double double_fhi = endpoint_secular_value_double(
shared_poles, shared_weights, count, rho_value, double_hi);
endpoint_sign_rescued =
isfinite(double_lo) && isfinite(double_hi) &&
isfinite(double_flo) && isfinite(double_fhi) &&
double_lo < double_hi && double_flo <= 0.0 && double_fhi >= 0.0;
}
if (endpoint_sign_rescued) {
result = flo > 0.0f ? lo : hi;
} else {
status = kStatusBracket;
}
} else {
float x = lo + 0.5f * (hi - lo);
const bool origin_at_lower = secular_eval_float(
shared_poles, shared_weights, count, rho_value, x).value > 0.0f;
bool converged = false;
for (int iteration = 1; iteration <= kFp32Iterations; ++iteration) {
used = iteration;
const EvalFloat current = secular_eval_float(
shared_poles, shared_weights, count, rho_value, x);
if (fabsf(current.value) <= FLT_EPSILON * current.error_scale) {
converged = true;
break;
}
if (current.value <= 0.0f) {
lo = fmaxf(lo, x);
} else {
hi = fminf(hi, x);
}
const float scale = fmaxf(1.0f, fmaxf(fabsf(lo), fabsf(hi)));
if (hi - lo <= 2.0f * FLT_EPSILON * scale || lo == hi) {
x = lo + 0.5f * (hi - lo);
converged = true;
break;
}
float eta = root + 1 < count
? interior_rational_step_float(
shared_poles,
shared_weights,
rho_value,
root,
x,
current.value,
current.derivative,
origin_at_lower)
: last_rational_step_float(
shared_poles,
shared_weights,
count,
rho_value,
x,
current.value);
if (!isfinite(eta) || current.value * eta >= 0.0f) {
eta = -current.value / current.derivative;
}
float proposed = x + eta;
if (!isfinite(proposed) || !(proposed > lo && proposed < hi) ||
proposed == x || iteration >= 12) {
proposed = lo + 0.5f * (hi - lo);
}
x = proposed;
}
result = x;
if (!converged || !isfinite(x)) {
status = kStatusNoConvergence;
}
}
values[base + root] = result;
roots64[base + root] = static_cast<double>(result);
root_status[base + root] = status;
root_iterations[base + root] = used;
atomicMax(maximum_iterations + merge, used);
if (status == kStatusBracket) {
atomicMax(merge_mode + merge, kModeSelectiveStatus2);
} else if (status != kStatusSuccess) {
// Mode priority is deterministic: full retry always dominates a
// concurrent status-2 request, independent of CTA completion order.
atomicMax(merge_mode + merge, kModeFullRetry);
}
}
__global__ __launch_bounds__(kTileWidth)
void build_fixed_tile_repair_queues_kernel(
const int32_t* __restrict__ active_count,
const int32_t* __restrict__ root_status,
const int32_t* __restrict__ merge_mode,
int32_t* __restrict__ repair_ids,
int32_t* __restrict__ repair_count) {
const int merge = static_cast<int>(blockIdx.x);
const int tile = static_cast<int>(blockIdx.y);
const int thread = static_cast<int>(threadIdx.x);
const int root = tile * kTileWidth + thread;
const int count = active_count[merge];
const int32_t mode = merge_mode[merge];
const int queue_base = (merge * kTiles + tile) * kTileWidth;
__shared__ int32_t inclusive_scan[kTileWidth];
repair_ids[queue_base + thread] = -1;
const bool active = root < count;
const bool selected =
mode == kModeFullRetry
? active
: (mode == kModeSelectiveStatus2 && active &&
root_status[merge * kWidth + root] == kStatusBracket);
inclusive_scan[thread] = selected ? 1 : 0;
__syncthreads();
#pragma unroll
for (int offset = 1; offset < kTileWidth; offset <<= 1) {
const int32_t addend =
thread >= offset ? inclusive_scan[thread - offset] : 0;
__syncthreads();
if (thread >= offset) {
inclusive_scan[thread] += addend;
}
__syncthreads();
}
if (selected) {
const int position = inclusive_scan[thread] - 1;
repair_ids[queue_base + position] = root;
}
if (thread == kTileWidth - 1) {
repair_count[merge * kTiles + tile] = inclusive_scan[thread];
}
}
__global__ __launch_bounds__(kTileWidth) void fp64_repair_4x256_kernel(
const float* __restrict__ poles,
const float* __restrict__ weights,
const float* __restrict__ rho,
const int32_t* __restrict__ active_count,
const int32_t* __restrict__ repair_ids,
const int32_t* __restrict__ repair_count,
double* __restrict__ roots64,
int32_t* __restrict__ root_status,
int32_t* __restrict__ root_iterations,
int32_t* __restrict__ deflated,
int32_t* __restrict__ info,
int32_t* __restrict__ maximum_iterations) {
const int merge = static_cast<int>(blockIdx.x);
const int tile = static_cast<int>(blockIdx.y);
const int thread = static_cast<int>(threadIdx.x);
const int count = active_count[merge];
const int base = merge * kWidth;
const int queue_base = (merge * kTiles + tile) * kTileWidth;
const int queue_count = repair_count[merge * kTiles + tile];
// This return is CTA-uniform and occurs before any shared state is touched.
if (queue_count == 0) {
return;
}
__shared__ double shared_poles[kWidth];
__shared__ double shared_weights[kWidth];
#pragma unroll
for (int round = 0; round < kTiles; ++round) {
const int index = round * kTileWidth + thread;
shared_poles[index] = static_cast<double>(poles[base + index]);
shared_weights[index] = static_cast<double>(weights[base + index]);
}
__syncthreads();
if (thread >= queue_count) {
return;
}
const int root = repair_ids[queue_base + thread];
if (!(root >= 0 && root < count)) {
atomicMax(info + merge, kStatusInvalidCount);
return;
}
deflated[base + root] = 0;
root_iterations[base + root] = 0;
const double rho_value = static_cast<double>(rho[merge]);
double lo = nextafter(shared_poles[root], CUDART_INF);
double hi;
if (root + 1 < count) {
hi = nextafter(shared_poles[root + 1], -CUDART_INF);
} else {
double normz2 = 0.0;
double correction = 0.0;
for (int index = 0; index < count; ++index) {
kahan_add_double(
shared_weights[index] * shared_weights[index],
normz2,
correction);
}
hi = shared_poles[count - 1] + rho_value * normz2;
if (!(hi > lo)) {
hi = nextafter(lo, CUDART_INF);
}
for (int attempt = 0; attempt < 8; ++attempt) {
if (secular_value_double(
shared_poles, shared_weights, count, rho_value, hi) >= 0.0) {
break;
}
hi = shared_poles[count - 1] +
2.0 * (hi - shared_poles[count - 1]);
}
}
const double flo = secular_eval_double(
shared_poles, shared_weights, count, rho_value, lo).value;
const double fhi = secular_eval_double(
shared_poles, shared_weights, count, rho_value, hi).value;
const bool finite_bracket =
(lo < hi) && isfinite(lo) && isfinite(hi) &&
isfinite(flo) && isfinite(fhi);
const bool sub_ulp_deflation = finite_bracket && flo >= 0.0 && fhi > 0.0;
double result = CUDART_NAN;
int32_t status = kStatusSuccess;
int32_t used = 0;
if (sub_ulp_deflation) {
deflated[base + root] = 1;
result = shared_poles[root];
} else if (!finite_bracket || !(flo < 0.0) || !(fhi > 0.0)) {
status = kStatusBracket;
} else {
double x = lo + 0.5 * (hi - lo);
const bool origin_at_lower = secular_eval_double(
shared_poles, shared_weights, count, rho_value, x).value > 0.0;
bool converged = false;
for (int iteration = 1; iteration <= kFp64Iterations; ++iteration) {
used = iteration;
const EvalDouble current = secular_eval_double(
shared_poles, shared_weights, count, rho_value, x);
if (fabs(current.value) <= DBL_EPSILON * current.error_scale) {
converged = true;
break;
}
if (current.value <= 0.0) {
lo = fmax(lo, x);
} else {
hi = fmin(hi, x);
}
const double scale = fmax(1.0, fmax(fabs(lo), fabs(hi)));
if (hi - lo <= 4.0 * DBL_EPSILON * scale ||
nextafter(lo, hi) >= hi) {
x = lo + 0.5 * (hi - lo);
converged = true;
break;
}
double eta = root + 1 < count
? interior_rational_step_double(
shared_poles,
shared_weights,
rho_value,
root,
x,
current.value,
current.derivative,
origin_at_lower)
: last_rational_step_double(
shared_poles,
shared_weights,
count,
rho_value,
x,
current.value);
if (!isfinite(eta) || current.value * eta >= 0.0) {
eta = -current.value / current.derivative;
}
double proposed = x + eta;
if (!isfinite(proposed) || !(proposed > lo && proposed < hi) ||
proposed == x || iteration >= 12) {
proposed = lo + 0.5 * (hi - lo);
}
x = proposed;
}
result = x;
if (!converged || !isfinite(x)) {
status = kStatusNoConvergence;
}
}
roots64[base + root] = result;
root_status[base + root] = status;
root_iterations[base + root] = used;
atomicMax(maximum_iterations + merge, used);
if (status != kStatusSuccess) {
atomicMax(info + merge, status);
}
}
// ---------------------------------------------------------------------------
// K4-K6 suffix (this lane). The kernels above are the FROZEN certified
// K0-K3 source (sha256 9421c7c0...72263c, byte-identical); everything below
// preserves the merge512_endpoint_deflation.cu SLAED3 formulas exactly:
// log-domain displacement product in root order, plain sequential
// accumulation (no Kahan) for the product, Kahan for the vector norm in row
// order, copysign from the original weight, endpoint-deflation skips, FP32
// output layout [merge, row, column].
// ---------------------------------------------------------------------------
constexpr int32_t kStatusWeight = 3;
constexpr int32_t kStatusNorm = 4;
// K4a: fp32 updated weights, common-route merges only.
// CTA (merge, tile); thread owns pole j = tile*256 + thread.
__global__ __launch_bounds__(kTileWidth) void fp32_updated_weights_4x256_kernel(
const float* __restrict__ poles,
const float* __restrict__ weights,
const float* __restrict__ values,
const int32_t* __restrict__ active_count,
const int32_t* __restrict__ merge_mode,
float* __restrict__ updated32,
int32_t* __restrict__ recon_status) {
const int merge = static_cast<int>(blockIdx.x);
const int tile = static_cast<int>(blockIdx.y);
const int thread = static_cast<int>(threadIdx.x);
const int owner = tile * kTileWidth + thread;
const int base = merge * kWidth;
const int count = active_count[merge];
updated32[base + owner] = 0.0f;
if (merge_mode[merge] != kModeCommon) {
return; // CTA-uniform: fp64 reconstruction route owns this merge.
}
if (!(count > 0 && count <= kWidth)) {
return;
}
__shared__ float shared_poles[kWidth];
__shared__ float shared_roots[kWidth];
#pragma unroll
for (int round = 0; round < kTiles; ++round) {
const int index = round * kTileWidth + thread;
shared_poles[index] = poles[base + index];
shared_roots[index] = values[base + index];
}
__syncthreads();
if (owner >= count) {
return;
}
const float diagonal_delta =
fabsf(shared_poles[owner] - shared_roots[owner]);
if (!(diagonal_delta > 0.0f) || !isfinite(diagonal_delta)) {
atomicMax(recon_status + merge, kStatusWeight);
updated32[base + owner] = CUDART_NAN_F;
return;
}
float log_product = logf(diagonal_delta);
for (int root = 0; root < count; ++root) {
if (root == owner) {
continue;
}
const float numerator = fabsf(shared_poles[owner] - shared_roots[root]);
const float denominator =
fabsf(shared_poles[owner] - shared_poles[root]);
if (!(numerator > 0.0f) || !(denominator > 0.0f)) {
atomicMax(recon_status + merge, kStatusWeight);
} else {
log_product += logf(numerator) - logf(denominator);
}
}
const float magnitude = expf(0.5f * log_product);
const float updated = copysignf(magnitude, weights[base + owner]);
updated32[base + owner] = updated;
if (!isfinite(updated)) {
atomicMax(recon_status + merge, kStatusWeight);
}
}
// K5a: fp32 secular vectors, common-route merges only.
// CTA (merge, tile); thread owns output column j = tile*256 + thread.
__global__ __launch_bounds__(kTileWidth) void fp32_secular_columns_4x256_kernel(
const float* __restrict__ poles,
const float* __restrict__ values,
const float* __restrict__ updated32,
const int32_t* __restrict__ active_count,
const int32_t* __restrict__ merge_mode,
float* __restrict__ secular_vectors,
int32_t* __restrict__ recon_status) {
const int merge = static_cast<int>(blockIdx.x);
const int tile = static_cast<int>(blockIdx.y);
const int thread = static_cast<int>(threadIdx.x);
const int column = tile * kTileWidth + thread;
const int base = merge * kWidth;
const int matrix_base = merge * kWidth * kWidth;
const int count = active_count[merge];
if (merge_mode[merge] != kModeCommon) {
return; // CTA-uniform: K5b owns this merge's S wholesale.
}
if (!(count > 0 && count <= kWidth)) {
return;
}
__shared__ float shared_poles[kWidth];
__shared__ float shared_roots[kWidth];
__shared__ float shared_updated[kWidth];
#pragma unroll
for (int round = 0; round < kTiles; ++round) {
const int index = round * kTileWidth + thread;
shared_poles[index] = poles[base + index];
shared_roots[index] = values[base + index];
shared_updated[index] = updated32[base + index];
}
__syncthreads();
if (column >= count) {
// Sole writer of the inactive column: zero it (values[j] is already the
// frozen zero from K1).
for (int row = 0; row < kWidth; ++row) {
secular_vectors[matrix_base + row * kWidth + column] = 0.0f;
}
return;
}
float norm2 = 0.0f;
float correction = 0.0f;
for (int row = 0; row < count; ++row) {
const float delta = shared_poles[row] - shared_roots[column];
const float element = shared_updated[row] / delta;
kahan_add(element * element, norm2, correction);
}
const float norm = sqrtf(norm2);
if (!(norm > 0.0f) || !isfinite(norm)) {
atomicMax(recon_status + merge, kStatusNorm);
}
for (int row = 0; row < count; ++row) {
const float delta = shared_poles[row] - shared_roots[column];
secular_vectors[matrix_base + row * kWidth + column] =
shared_updated[row] / delta / norm;
}
for (int row = count; row < kWidth; ++row) {
secular_vectors[matrix_base + row * kWidth + column] = 0.0f;
}
}
// K2e: escalation boundary. A common-route merge whose fp32 reconstruction
// failed (status 3/4) is promoted to the full FP64 retry: escalated[merge]=1
// and ALL active roots are enqueued into the SEPARATE escalation queues.
// Every other merge writes count 0 (K3's relaunch then uniform-returns).
__global__ __launch_bounds__(kTileWidth) void build_escalation_queues_kernel(
const int32_t* __restrict__ active_count,
const int32_t* __restrict__ merge_mode,
const int32_t* __restrict__ recon_status,
int32_t* __restrict__ escalation_ids,
int32_t* __restrict__ escalation_count,
int32_t* __restrict__ escalated) {
const int merge = static_cast<int>(blockIdx.x);
const int tile = static_cast<int>(blockIdx.y);
const int thread = static_cast<int>(threadIdx.x);
const int root = tile * kTileWidth + thread;
const int count = active_count[merge];
const int queue_base = (merge * kTiles + tile) * kTileWidth;
const bool promote =
merge_mode[merge] == kModeCommon && recon_status[merge] != 0;
if (thread == 0 && tile == 0) {
escalated[merge] = promote ? 1 : 0;
}
const bool selected = promote && root < count;
escalation_ids[queue_base + thread] = selected ? root : -1;
// Full-retry enqueues every active root in tile order, so the stable
// "scan" is the identity within the tile prefix.
if (thread == kTileWidth - 1) {
int tile_count = 0;
if (promote) {
const int start = tile * kTileWidth;
const int stop = min(start + kTileWidth, count);
tile_count = stop > start ? stop - start : 0;
}
escalation_count[merge * kTiles + tile] = tile_count;
}
}
// K4b: fp64 updated weights for every merge on the fp64 reconstruction
// route = (mode != common) || escalated. Locked displacement product with
// endpoint-deflation skips (a deflated root's factors cancel exactly).
__global__ __launch_bounds__(kTileWidth) void fp64_updated_weights_4x256_kernel(
const float* __restrict__ poles,
const float* __restrict__ weights,
const double* __restrict__ roots64,
const int32_t* __restrict__ active_count,
const int32_t* __restrict__ merge_mode,
const int32_t* __restrict__ escalated,
const int32_t* __restrict__ deflated,
double* __restrict__ updated64,
int32_t* __restrict__ info) {
const int merge = static_cast<int>(blockIdx.x);
const int tile = static_cast<int>(blockIdx.y);
const int thread = static_cast<int>(threadIdx.x);
const int owner = tile * kTileWidth + thread;
const int base = merge * kWidth;
const int count = active_count[merge];
updated64[base + owner] = 0.0;
if (merge_mode[merge] == kModeCommon && !escalated[merge]) {
return; // CTA-uniform: fp32 reconstruction already produced this merge.
}
if (!(count > 0 && count <= kWidth)) {
return;
}
__shared__ double shared_poles[kWidth];
__shared__ double shared_roots[kWidth];
__shared__ int32_t shared_deflated[kWidth];
#pragma unroll
for (int round = 0; round < kTiles; ++round) {
const int index = round * kTileWidth + thread;
shared_poles[index] = static_cast<double>(poles[base + index]);
shared_roots[index] = roots64[base + index];
shared_deflated[index] = deflated[base + index];
}
__syncthreads();
if (owner >= count) {
return;
}
if (shared_deflated[owner]) {
updated64[base + owner] = 0.0;
return;
}
const double diagonal_delta =
fabs(shared_poles[owner] - shared_roots[owner]);
if (!(diagonal_delta > 0.0) || !isfinite(diagonal_delta)) {
atomicMax(info + merge, kStatusWeight);
updated64[base + owner] = CUDART_NAN;
return;
}
double log_product = log(diagonal_delta);
for (int root = 0; root < count; ++root) {
if (root == owner || shared_deflated[root]) {
continue;
}
const double numerator = fabs(shared_poles[owner] - shared_roots[root]);
const double denominator =
fabs(shared_poles[owner] - shared_poles[root]);
if (!(numerator > 0.0) || !(denominator > 0.0) ||
!isfinite(numerator) || !isfinite(denominator)) {
atomicMax(info + merge, kStatusWeight);
} else {
log_product += log(numerator) - log(denominator);
}
}
const double magnitude = exp(0.5 * log_product);
const double updated =
copysign(magnitude, static_cast<double>(weights[base + owner]));
updated64[base + owner] = updated;
if (!isfinite(updated)) {
atomicMax(info + merge, kStatusWeight);
}
}
// K5b: fp64 secular vectors + values for fp64-route merges. Deflated roots
// get the exact identity column and value = pole; every element of S and
// every value of the merge is rewritten (the parent hybrid semantics).
__global__ __launch_bounds__(kTileWidth) void fp64_secular_columns_4x256_kernel(
const float* __restrict__ poles,
const double* __restrict__ roots64,
const double* __restrict__ updated64,
const int32_t* __restrict__ active_count,
const int32_t* __restrict__ merge_mode,
const int32_t* __restrict__ escalated,
const int32_t* __restrict__ deflated,
float* __restrict__ values,
float* __restrict__ secular_vectors,
int32_t* __restrict__ info) {
const int merge = static_cast<int>(blockIdx.x);
const int tile = static_cast<int>(blockIdx.y);
const int thread = static_cast<int>(threadIdx.x);
const int column = tile * kTileWidth + thread;
const int base = merge * kWidth;
const int matrix_base = merge * kWidth * kWidth;
const int count = active_count[merge];
if (merge_mode[merge] == kModeCommon && !escalated[merge]) {
return;
}
if (!(count > 0 && count <= kWidth)) {
return;
}
__shared__ double shared_poles[kWidth];
__shared__ double shared_roots[kWidth];
__shared__ double shared_updated[kWidth];
__shared__ int32_t shared_deflated[kWidth];
#pragma unroll
for (int round = 0; round < kTiles; ++round) {
const int index = round * kTileWidth + thread;
shared_poles[index] = static_cast<double>(poles[base + index]);
shared_roots[index] = roots64[base + index];
shared_updated[index] = updated64[base + index];
shared_deflated[index] = deflated[base + index];
}
__syncthreads();
if (column >= count) {
for (int row = 0; row < kWidth; ++row) {
secular_vectors[matrix_base + row * kWidth + column] = 0.0f;
}
values[base + column] = 0.0f;
return;
}
if (shared_deflated[column]) {
for (int row = 0; row < kWidth; ++row) {
secular_vectors[matrix_base + row * kWidth + column] =
row == column ? 1.0f : 0.0f;
}
values[base + column] = static_cast<float>(shared_roots[column]);
return;
}
double norm2 = 0.0;
double correction = 0.0;
for (int row = 0; row < count; ++row) {
const double delta = shared_poles[row] - shared_roots[column];
const double element = shared_updated[row] / delta;
kahan_add_double(element * element, norm2, correction);
}
const double norm = sqrt(norm2);
if (!(norm > 0.0) || !isfinite(norm)) {
atomicMax(info + merge, kStatusNorm);
}
for (int row = 0; row < count; ++row) {
const double delta = shared_poles[row] - shared_roots[column];
const double element = shared_updated[row] / delta / norm;
secular_vectors[matrix_base + row * kWidth + column] =
static_cast<float>(element);
}
for (int row = count; row < kWidth; ++row) {
secular_vectors[matrix_base + row * kWidth + column] = 0.0f;
}
values[base + column] = static_cast<float>(shared_roots[column]);
}
// K3p: bounded FP64 polish of ALL non-deflated active roots, warm-started
// from the route's current roots64 (fp32-accurate for K1-accepted lanes).
// Same locked fp64 rational/bisection formulas and convergence criterion as
// K3, capped at kPolishIterations — accepted fp32 roots typically converge
// in 1-3 fp64 steps, restoring fp64-class VALUE accuracy at ~1/6 the cost
// of the full fp64 re-solve (the tree-stage reopener named in
// ../n1024-row-e2e/PREREG.md). K3-repaired roots pass the convergence
// check immediately. Writes roots64 (+ iteration telemetry) only.
constexpr int kPolishIterations = 30; // hard tail (endpoint-rescued warm starts) needs fresh-solve depth; converged roots exit in 1 eval
__global__ __launch_bounds__(kTileWidth) void fp64_polish_4x256_kernel(
const float* __restrict__ poles,
const float* __restrict__ weights,
const float* __restrict__ rho,
const int32_t* __restrict__ active_count,
const int32_t* __restrict__ deflated,
double* __restrict__ roots64,
int32_t* __restrict__ root_iterations,
int32_t* __restrict__ info,
int32_t* __restrict__ maximum_iterations) {
const int merge = static_cast<int>(blockIdx.x);
const int tile = static_cast<int>(blockIdx.y);
const int thread = static_cast<int>(threadIdx.x);
const int root = tile * kTileWidth + thread;
const int base = merge * kWidth;
const int count = active_count[merge];
__shared__ double shared_poles[kWidth];
__shared__ double shared_weights[kWidth];
#pragma unroll
for (int round = 0; round < kTiles; ++round) {
const int index = round * kTileWidth + thread;
shared_poles[index] = static_cast<double>(poles[base + index]);
shared_weights[index] = static_cast<double>(weights[base + index]);
}
__syncthreads();
if (!(count > 0 && count <= kWidth) || root >= count) {
return;
}
if (deflated[base + root]) {
return;
}
const double warm = roots64[base + root];
if (!isfinite(warm)) {
return; // residual K3 failure: info already nonzero, fail-closed.
}
const double rho_value = static_cast<double>(rho[merge]);
double lo = nextafter(shared_poles[root], CUDART_INF);
double hi;
if (root + 1 < count) {
hi = nextafter(shared_poles[root + 1], -CUDART_INF);
} else {
double normz2 = 0.0;
double correction = 0.0;
for (int index = 0; index < count; ++index) {
kahan_add_double(
shared_weights[index] * shared_weights[index], normz2, correction);
}
hi = shared_poles[count - 1] + rho_value * normz2;
if (!(hi > lo)) {
hi = nextafter(lo, CUDART_INF);
}
for (int attempt = 0; attempt < 8; ++attempt) {
if (secular_value_double(
shared_poles, shared_weights, count, rho_value, hi) >= 0.0) {
break;
}
hi = shared_poles[count - 1] + 2.0 * (hi - shared_poles[count - 1]);
}
}
double x = fmin(fmax(warm, lo), hi);
bool converged = false;
int used = 0;
for (int iteration = 1; iteration <= kPolishIterations; ++iteration) {
used = iteration;
const EvalDouble current = secular_eval_double(
shared_poles, shared_weights, count, rho_value, x);
if (fabs(current.value) <= DBL_EPSILON * current.error_scale) {
converged = true;
break;
}
if (current.value <= 0.0) {
lo = fmax(lo, x);
} else {
hi = fmin(hi, x);
}
const double scale = fmax(1.0, fmax(fabs(lo), fabs(hi)));
if (hi - lo <= 4.0 * DBL_EPSILON * scale || nextafter(lo, hi) >= hi) {
x = lo + 0.5 * (hi - lo);
converged = true;
break;
}
double eta = root + 1 < count
? interior_rational_step_double(
shared_poles, shared_weights, rho_value, root, x,
current.value, current.derivative,
current.value > 0.0)
: last_rational_step_double(
shared_poles, shared_weights, count, rho_value, x,
current.value);
if (!isfinite(eta) || current.value * eta >= 0.0) {
eta = -current.value / current.derivative;
}
double proposed = x + eta;
if (!isfinite(proposed) || !(proposed > lo && proposed < hi) ||
proposed == x) {
proposed = lo + 0.5 * (hi - lo);
}
x = proposed;
}
// Bounded pass: non-convergence within the cap is accepted (the bracket
// has still tightened around the root); a NaN result is not.
if (!isfinite(x)) {
atomicMax(info + merge, kStatusNoConvergence);
return;
}
roots64[base + root] = x;
root_iterations[base + root] = used;
atomicMax(maximum_iterations + merge, used);
}
// K0e: clear the K4-K6 route state K0 does not know about (K0 is frozen).
// force_full != 0 pins the merge mode to the full FP64 retry BEFORE K1:
// K2 then enqueues every active root and K3 solves them all in true FP64 —
// the endpoint-all tree contract (`run_fp64_control_`, fallback_only=false)
// expressed in the retile. K1's atomicMax only ever raises the mode, so
// the pin is stable; with force_full == 0 behavior is bitwise-unchanged.
__global__ void clear_k4k6_state_kernel(
int32_t* __restrict__ recon_status,
int32_t* __restrict__ merge_mode,
int32_t* __restrict__ escalated,
int force_full,
int set_escalated,
int batch) {
const int merge = static_cast<int>(blockIdx.x) * blockDim.x + threadIdx.x;
if (merge >= batch) {
return;
}
recon_status[merge] = 0;
if (force_full) {
merge_mode[merge] = kModeFullRetry;
}
if (set_escalated) {
escalated[merge] = 1; // polished route: fp64 reconstruction everywhere
}
}
// K6: finalize/audit telemetry only. Never touches info: the route
// semantics are exactly the parent's, and the audit result is a separate
// caller-visible word: bit0..1 final mode, bit2 escalated, bit3 finite,
// bit4 ascending.
__global__ void finalize_values_info_kernel(
const float* __restrict__ values,
const int32_t* __restrict__ active_count,
const int32_t* __restrict__ merge_mode,
const int32_t* __restrict__ escalated,
int32_t* __restrict__ route_summary) {
const int merge = static_cast<int>(blockIdx.x);
const int thread = static_cast<int>(threadIdx.x);
const int base = merge * kWidth;
const int count = active_count[merge];
__shared__ int32_t violations;
if (thread == 0) {
violations = 0;
}
__syncthreads();
int local_nonfinite = 0;
int local_descending = 0;
for (int index = thread; index < count; index += blockDim.x) {
const float value = values[base + index];
if (!isfinite(value)) {
local_nonfinite = 1;
}
if (index + 1 < count && !(values[base + index + 1] >= value)) {
local_descending = 1;
}
}
if (local_nonfinite) {
atomicOr(&violations, 1);
}
if (local_descending) {
atomicOr(&violations, 2);
}
__syncthreads();
if (thread == 0) {
const int32_t mode = merge_mode[merge];
const int32_t summary =
(mode & 3) | ((escalated[merge] ? 1 : 0) << 2) |
(((violations & 1) == 0 ? 1 : 0) << 3) |
(((violations & 2) == 0 ? 1 : 0) << 4);
route_summary[merge] = summary;
}
}
void check_tensor(
const at::Tensor& tensor,
at::ScalarType dtype,
int64_t dimensions,
const char* name) {
TORCH_CHECK(tensor.is_cuda(), name, " must be CUDA");
TORCH_CHECK(tensor.is_contiguous(), name, " must be contiguous");
TORCH_CHECK(tensor.scalar_type() == dtype, name, " has wrong dtype");
TORCH_CHECK(tensor.dim() == dimensions, name, " has wrong rank");
}
void check_device(const at::Tensor& tensor, int device, const char* name) {
TORCH_CHECK(
tensor.get_device() == device,
name,
" must be on CUDA device ",
device);
}
} // namespace
pybind11::dict n1024_prefix_resource_probe() {
cudaFuncAttributes k0{};
cudaFuncAttributes k1{};
cudaFuncAttributes k2{};
cudaFuncAttributes k3{};
C10_CUDA_CHECK(cudaFuncGetAttributes(&k0, clear_small_route_state_kernel));
C10_CUDA_CHECK(cudaFuncGetAttributes(&k1, fp32_roots_4x256_kernel));
C10_CUDA_CHECK(
cudaFuncGetAttributes(&k2, build_fixed_tile_repair_queues_kernel));
C10_CUDA_CHECK(cudaFuncGetAttributes(&k3, fp64_repair_4x256_kernel));
int k1_active = 0;
int k3_active = 0;
C10_CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&k1_active, fp32_roots_4x256_kernel, kTileWidth, 0));
C10_CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&k3_active, fp64_repair_4x256_kernel, kTileWidth, 0));
auto attributes = [](const cudaFuncAttributes& item) {
pybind11::dict result;
result["num_regs"] = item.numRegs;
result["shared_static_bytes"] = item.sharedSizeBytes;
result["local_bytes"] = item.localSizeBytes;
result["max_threads_per_block"] = item.maxThreadsPerBlock;
result["binary_version"] = item.binaryVersion;
result["ptx_version"] = item.ptxVersion;
return result;
};
pybind11::dict result;
result["k0"] = attributes(k0);
result["k1"] = attributes(k1);
result["k2"] = attributes(k2);
result["k3"] = attributes(k3);
result["k1_active_blocks_per_sm"] = k1_active;
result["k3_active_blocks_per_sm"] = k3_active;
result["threads_per_cta"] = kTileWidth;
result["tiles_per_merge"] = kTiles;
result["k1_modeled_shared_bytes"] = 2 * kWidth * sizeof(float);
result["k3_modeled_shared_bytes"] = 2 * kWidth * sizeof(double);
return result;
}
void n1024_prefix_run(
const at::Tensor& poles,
const at::Tensor& weights,
const at::Tensor& rho,
const at::Tensor& active_count,
const at::Tensor& values,
const at::Tensor& roots64,
const at::Tensor& root_status,
const at::Tensor& root_iterations,
const at::Tensor& merge_mode,
const at::Tensor& repair_ids,
const at::Tensor& repair_count,
const at::Tensor& deflated,
const at::Tensor& info,
const at::Tensor& maximum_iterations) {
check_tensor(poles, at::kFloat, 2, "poles");
check_tensor(weights, at::kFloat, 2, "weights");
check_tensor(rho, at::kFloat, 1, "rho");
check_tensor(active_count, at::kInt, 1, "active_count");
check_tensor(values, at::kFloat, 2, "values");
check_tensor(roots64, at::kDouble, 2, "roots64");
check_tensor(root_status, at::kInt, 2, "root_status");
check_tensor(root_iterations, at::kInt, 2, "root_iterations");
check_tensor(merge_mode, at::kInt, 1, "merge_mode");
check_tensor(repair_ids, at::kInt, 3, "repair_ids");
check_tensor(repair_count, at::kInt, 2, "repair_count");
check_tensor(deflated, at::kInt, 2, "deflated");
check_tensor(info, at::kInt, 1, "info");
check_tensor(maximum_iterations, at::kInt, 1, "maximum_iterations");
const int64_t batch = poles.size(0);
TORCH_CHECK(batch == 1, "bounded K0-K3 prefix accepts B=1 only");
TORCH_CHECK(poles.sizes() == weights.sizes(), "poles/weights mismatch");
TORCH_CHECK(poles.size(1) == kWidth, "width must be 1024");
TORCH_CHECK(rho.size(0) == batch, "rho shape mismatch");
TORCH_CHECK(active_count.size(0) == batch, "active_count shape mismatch");
TORCH_CHECK(values.sizes() == poles.sizes(), "values shape mismatch");
TORCH_CHECK(roots64.sizes() == poles.sizes(), "roots64 shape mismatch");
TORCH_CHECK(root_status.sizes() == poles.sizes(), "root_status mismatch");
TORCH_CHECK(
root_iterations.sizes() == poles.sizes(), "root_iterations mismatch");
TORCH_CHECK(merge_mode.size(0) == batch, "merge_mode shape mismatch");
TORCH_CHECK(
repair_ids.size(0) == batch && repair_ids.size(1) == kTiles &&
repair_ids.size(2) == kTileWidth,
"repair_ids shape mismatch");
TORCH_CHECK(
repair_count.size(0) == batch && repair_count.size(1) == kTiles,
"repair_count shape mismatch");
TORCH_CHECK(deflated.sizes() == poles.sizes(), "deflated shape mismatch");
TORCH_CHECK(info.size(0) == batch, "info shape mismatch");
TORCH_CHECK(
maximum_iterations.size(0) == batch,
"maximum_iterations shape mismatch");
const int device = poles.get_device();
check_device(weights, device, "weights");
check_device(rho, device, "rho");
check_device(active_count, device, "active_count");
check_device(values, device, "values");
check_device(roots64, device, "roots64");
check_device(root_status, device, "root_status");
check_device(root_iterations, device, "root_iterations");
check_device(merge_mode, device, "merge_mode");
check_device(repair_ids, device, "repair_ids");
check_device(repair_count, device, "repair_count");
check_device(deflated, device, "deflated");
check_device(info, device, "info");
check_device(maximum_iterations, device, "maximum_iterations");
c10::cuda::CUDAGuard guard(poles.device());
clear_small_route_state_kernel<<<1, 256, 0>>>(
active_count.data_ptr<int32_t>(),
merge_mode.data_ptr<int32_t>(),
repair_count.data_ptr<int32_t>(),
info.data_ptr<int32_t>(),
maximum_iterations.data_ptr<int32_t>(),
static_cast<int>(batch));
C10_CUDA_KERNEL_LAUNCH_CHECK();
const dim3 tiled_grid(static_cast<unsigned int>(batch), kTiles, 1);
fp32_roots_4x256_kernel<<<tiled_grid, kTileWidth, 0>>>(
poles.data_ptr<float>(),
weights.data_ptr<float>(),
rho.data_ptr<float>(),
active_count.data_ptr<int32_t>(),
values.data_ptr<float>(),
roots64.data_ptr<double>(),
root_status.data_ptr<int32_t>(),
root_iterations.data_ptr<int32_t>(),
merge_mode.data_ptr<int32_t>(),
deflated.data_ptr<int32_t>(),
maximum_iterations.data_ptr<int32_t>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
build_fixed_tile_repair_queues_kernel<<<
tiled_grid, kTileWidth, 0>>>(
active_count.data_ptr<int32_t>(),
root_status.data_ptr<int32_t>(),
merge_mode.data_ptr<int32_t>(),
repair_ids.data_ptr<int32_t>(),
repair_count.data_ptr<int32_t>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
fp64_repair_4x256_kernel<<<tiled_grid, kTileWidth, 0>>>(
poles.data_ptr<float>(),
weights.data_ptr<float>(),
rho.data_ptr<float>(),
active_count.data_ptr<int32_t>(),
repair_ids.data_ptr<int32_t>(),
repair_count.data_ptr<int32_t>(),
roots64.data_ptr<double>(),
root_status.data_ptr<int32_t>(),
root_iterations.data_ptr<int32_t>(),
deflated.data_ptr<int32_t>(),
info.data_ptr<int32_t>(),
maximum_iterations.data_ptr<int32_t>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
pybind11::dict n1024_k4k6_resource_probe() {
cudaFuncAttributes k4a{};
cudaFuncAttributes k5a{};
cudaFuncAttributes k2e{};
cudaFuncAttributes k4b{};
cudaFuncAttributes k5b{};
cudaFuncAttributes k6{};
C10_CUDA_CHECK(
cudaFuncGetAttributes(&k4a, fp32_updated_weights_4x256_kernel));
C10_CUDA_CHECK(
cudaFuncGetAttributes(&k5a, fp32_secular_columns_4x256_kernel));
C10_CUDA_CHECK(cudaFuncGetAttributes(&k2e, build_escalation_queues_kernel));
C10_CUDA_CHECK(
cudaFuncGetAttributes(&k4b, fp64_updated_weights_4x256_kernel));
C10_CUDA_CHECK(
cudaFuncGetAttributes(&k5b, fp64_secular_columns_4x256_kernel));
C10_CUDA_CHECK(cudaFuncGetAttributes(&k6, finalize_values_info_kernel));
int k4a_active = 0;
int k5a_active = 0;
int k4b_active = 0;
int k5b_active = 0;
C10_CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&k4a_active, fp32_updated_weights_4x256_kernel, kTileWidth, 0));
C10_CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&k5a_active, fp32_secular_columns_4x256_kernel, kTileWidth, 0));
C10_CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&k4b_active, fp64_updated_weights_4x256_kernel, kTileWidth, 0));
C10_CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&k5b_active, fp64_secular_columns_4x256_kernel, kTileWidth, 0));
auto attributes = [](const cudaFuncAttributes& item) {
pybind11::dict result;
result["num_regs"] = item.numRegs;
result["shared_static_bytes"] = item.sharedSizeBytes;
result["local_bytes"] = item.localSizeBytes;
result["max_threads_per_block"] = item.maxThreadsPerBlock;
result["binary_version"] = item.binaryVersion;
result["ptx_version"] = item.ptxVersion;
return result;
};
pybind11::dict result;
result["k4a"] = attributes(k4a);
result["k5a"] = attributes(k5a);
result["k2e"] = attributes(k2e);
result["k4b"] = attributes(k4b);
result["k5b"] = attributes(k5b);
result["k6"] = attributes(k6);
result["k4a_active_blocks_per_sm"] = k4a_active;
result["k5a_active_blocks_per_sm"] = k5a_active;
result["k4b_active_blocks_per_sm"] = k4b_active;
result["k5b_active_blocks_per_sm"] = k5b_active;
result["threads_per_cta"] = kTileWidth;
result["k4a_modeled_shared_bytes"] = 2 * kWidth * sizeof(float);
result["k5a_modeled_shared_bytes"] = 3 * kWidth * sizeof(float);
result["k4b_modeled_shared_bytes"] =
2 * kWidth * sizeof(double) + kWidth * sizeof(int32_t);
result["k5b_modeled_shared_bytes"] =
3 * kWidth * sizeof(double) + kWidth * sizeof(int32_t);
return result;
}
// Full K0-K6 width-1024 merge, batch-general. Kernel sources for K0-K3 are
// the frozen certified prefix; the binding-level B=1 restriction of the
// prefix entry point does not apply here (the kernels were always
// batch-shaped via gridDim.x).
namespace {
void n1024_merge_full_impl(
int route_mode, // 0 certified selective, 1 full-fp64, 2 fp32+fp64-polish
const at::Tensor& poles,
const at::Tensor& weights,
const at::Tensor& rho,
const at::Tensor& active_count,
const at::Tensor& values,
const at::Tensor& roots64,
const at::Tensor& root_status,
const at::Tensor& root_iterations,
const at::Tensor& merge_mode,
const at::Tensor& repair_ids,
const at::Tensor& repair_count,
const at::Tensor& deflated,
const at::Tensor& info,
const at::Tensor& maximum_iterations,
const at::Tensor& updated32,
const at::Tensor& updated64,
const at::Tensor& recon_status,
const at::Tensor& escalated,
const at::Tensor& escalation_ids,
const at::Tensor& escalation_count,
const at::Tensor& secular_vectors,
const at::Tensor& route_summary) {
check_tensor(poles, at::kFloat, 2, "poles");
check_tensor(weights, at::kFloat, 2, "weights");
check_tensor(rho, at::kFloat, 1, "rho");
check_tensor(active_count, at::kInt, 1, "active_count");
check_tensor(values, at::kFloat, 2, "values");
check_tensor(roots64, at::kDouble, 2, "roots64");
check_tensor(root_status, at::kInt, 2, "root_status");
check_tensor(root_iterations, at::kInt, 2, "root_iterations");
check_tensor(merge_mode, at::kInt, 1, "merge_mode");
check_tensor(repair_ids, at::kInt, 3, "repair_ids");
check_tensor(repair_count, at::kInt, 2, "repair_count");
check_tensor(deflated, at::kInt, 2, "deflated");
check_tensor(info, at::kInt, 1, "info");
check_tensor(maximum_iterations, at::kInt, 1, "maximum_iterations");
check_tensor(updated32, at::kFloat, 2, "updated32");
check_tensor(updated64, at::kDouble, 2, "updated64");
check_tensor(recon_status, at::kInt, 1, "recon_status");
check_tensor(escalated, at::kInt, 1, "escalated");
check_tensor(escalation_ids, at::kInt, 3, "escalation_ids");
check_tensor(escalation_count, at::kInt, 2, "escalation_count");
check_tensor(secular_vectors, at::kFloat, 3, "secular_vectors");
check_tensor(route_summary, at::kInt, 1, "route_summary");
const int64_t batch = poles.size(0);
TORCH_CHECK(batch > 0 && batch <= 1024, "batch must be 1..1024");
TORCH_CHECK(poles.sizes() == weights.sizes(), "poles/weights mismatch");
TORCH_CHECK(poles.size(1) == kWidth, "width must be 1024");
TORCH_CHECK(rho.size(0) == batch, "rho shape mismatch");
TORCH_CHECK(active_count.size(0) == batch, "active_count shape mismatch");
TORCH_CHECK(values.sizes() == poles.sizes(), "values shape mismatch");
TORCH_CHECK(roots64.sizes() == poles.sizes(), "roots64 shape mismatch");
TORCH_CHECK(root_status.sizes() == poles.sizes(), "root_status mismatch");
TORCH_CHECK(
root_iterations.sizes() == poles.sizes(), "root_iterations mismatch");
TORCH_CHECK(merge_mode.size(0) == batch, "merge_mode shape mismatch");
TORCH_CHECK(
repair_ids.size(0) == batch && repair_ids.size(1) == kTiles &&
repair_ids.size(2) == kTileWidth,
"repair_ids shape mismatch");
TORCH_CHECK(
repair_count.size(0) == batch && repair_count.size(1) == kTiles,
"repair_count shape mismatch");
TORCH_CHECK(deflated.sizes() == poles.sizes(), "deflated shape mismatch");
TORCH_CHECK(info.size(0) == batch, "info shape mismatch");
TORCH_CHECK(
maximum_iterations.size(0) == batch,
"maximum_iterations shape mismatch");
TORCH_CHECK(updated32.sizes() == poles.sizes(), "updated32 mismatch");
TORCH_CHECK(updated64.sizes() == poles.sizes(), "updated64 mismatch");
TORCH_CHECK(recon_status.size(0) == batch, "recon_status shape mismatch");
TORCH_CHECK(escalated.size(0) == batch, "escalated shape mismatch");
TORCH_CHECK(
escalation_ids.sizes() == repair_ids.sizes(),
"escalation_ids shape mismatch");
TORCH_CHECK(
escalation_count.sizes() == repair_count.sizes(),
"escalation_count shape mismatch");
TORCH_CHECK(
secular_vectors.size(0) == batch &&
secular_vectors.size(1) == kWidth &&
secular_vectors.size(2) == kWidth,
"secular_vectors shape mismatch");
TORCH_CHECK(route_summary.size(0) == batch, "route_summary mismatch");
const int device = poles.get_device();
check_device(weights, device, "weights");
check_device(rho, device, "rho");
check_device(active_count, device, "active_count");
check_device(values, device, "values");
check_device(roots64, device, "roots64");
check_device(root_status, device, "root_status");
check_device(root_iterations, device, "root_iterations");
check_device(merge_mode, device, "merge_mode");
check_device(repair_ids, device, "repair_ids");
check_device(repair_count, device, "repair_count");
check_device(deflated, device, "deflated");
check_device(info, device, "info");
check_device(maximum_iterations, device, "maximum_iterations");
check_device(updated32, device, "updated32");
check_device(updated64, device, "updated64");
check_device(recon_status, device, "recon_status");
check_device(escalated, device, "escalated");
check_device(escalation_ids, device, "escalation_ids");
check_device(escalation_count, device, "escalation_count");
check_device(secular_vectors, device, "secular_vectors");
check_device(route_summary, device, "route_summary");
c10::cuda::CUDAGuard guard(poles.device());
const int batch_int = static_cast<int>(batch);
const int clear_blocks = (batch_int + 255) / 256;
clear_small_route_state_kernel<<<clear_blocks, 256, 0>>>(
active_count.data_ptr<int32_t>(),
merge_mode.data_ptr<int32_t>(),
repair_count.data_ptr<int32_t>(),
info.data_ptr<int32_t>(),
maximum_iterations.data_ptr<int32_t>(),
batch_int);
C10_CUDA_KERNEL_LAUNCH_CHECK();
clear_k4k6_state_kernel<<<clear_blocks, 256, 0>>>(
recon_status.data_ptr<int32_t>(),
merge_mode.data_ptr<int32_t>(),
escalated.data_ptr<int32_t>(),
route_mode == 1 ? 1 : 0,
route_mode == 2 ? 1 : 0,
batch_int);
C10_CUDA_KERNEL_LAUNCH_CHECK();
const dim3 tiled_grid(static_cast<unsigned int>(batch), kTiles, 1);
fp32_roots_4x256_kernel<<<tiled_grid, kTileWidth, 0>>>(
poles.data_ptr<float>(),
weights.data_ptr<float>(),
rho.data_ptr<float>(),
active_count.data_ptr<int32_t>(),
values.data_ptr<float>(),
roots64.data_ptr<double>(),
root_status.data_ptr<int32_t>(),
root_iterations.data_ptr<int32_t>(),
merge_mode.data_ptr<int32_t>(),
deflated.data_ptr<int32_t>(),
maximum_iterations.data_ptr<int32_t>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
build_fixed_tile_repair_queues_kernel<<<
tiled_grid, kTileWidth, 0>>>(
active_count.data_ptr<int32_t>(),
root_status.data_ptr<int32_t>(),
merge_mode.data_ptr<int32_t>(),
repair_ids.data_ptr<int32_t>(),
repair_count.data_ptr<int32_t>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
fp64_repair_4x256_kernel<<<tiled_grid, kTileWidth, 0>>>(
poles.data_ptr<float>(),
weights.data_ptr<float>(),
rho.data_ptr<float>(),
active_count.data_ptr<int32_t>(),
repair_ids.data_ptr<int32_t>(),
repair_count.data_ptr<int32_t>(),
roots64.data_ptr<double>(),
root_status.data_ptr<int32_t>(),
root_iterations.data_ptr<int32_t>(),
deflated.data_ptr<int32_t>(),
info.data_ptr<int32_t>(),
maximum_iterations.data_ptr<int32_t>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
if (route_mode == 2) {
fp64_polish_4x256_kernel<<<tiled_grid, kTileWidth, 0>>>(
poles.data_ptr<float>(),
weights.data_ptr<float>(),
rho.data_ptr<float>(),
active_count.data_ptr<int32_t>(),
deflated.data_ptr<int32_t>(),
roots64.data_ptr<double>(),
root_iterations.data_ptr<int32_t>(),
info.data_ptr<int32_t>(),
maximum_iterations.data_ptr<int32_t>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
} else {
fp32_updated_weights_4x256_kernel<<<
tiled_grid, kTileWidth, 0>>>(
poles.data_ptr<float>(),
weights.data_ptr<float>(),
values.data_ptr<float>(),
active_count.data_ptr<int32_t>(),
merge_mode.data_ptr<int32_t>(),
updated32.data_ptr<float>(),
recon_status.data_ptr<int32_t>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
fp32_secular_columns_4x256_kernel<<<
tiled_grid, kTileWidth, 0>>>(
poles.data_ptr<float>(),
values.data_ptr<float>(),
updated32.data_ptr<float>(),
active_count.data_ptr<int32_t>(),
merge_mode.data_ptr<int32_t>(),
secular_vectors.data_ptr<float>(),
recon_status.data_ptr<int32_t>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
build_escalation_queues_kernel<<<
tiled_grid, kTileWidth, 0>>>(
active_count.data_ptr<int32_t>(),
merge_mode.data_ptr<int32_t>(),
recon_status.data_ptr<int32_t>(),
escalation_ids.data_ptr<int32_t>(),
escalation_count.data_ptr<int32_t>(),
escalated.data_ptr<int32_t>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
fp64_repair_4x256_kernel<<<tiled_grid, kTileWidth, 0>>>(
poles.data_ptr<float>(),
weights.data_ptr<float>(),
rho.data_ptr<float>(),
active_count.data_ptr<int32_t>(),
escalation_ids.data_ptr<int32_t>(),
escalation_count.data_ptr<int32_t>(),
roots64.data_ptr<double>(),
root_status.data_ptr<int32_t>(),
root_iterations.data_ptr<int32_t>(),
deflated.data_ptr<int32_t>(),
info.data_ptr<int32_t>(),
maximum_iterations.data_ptr<int32_t>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
fp64_updated_weights_4x256_kernel<<<
tiled_grid, kTileWidth, 0>>>(
poles.data_ptr<float>(),
weights.data_ptr<float>(),
roots64.data_ptr<double>(),
active_count.data_ptr<int32_t>(),
merge_mode.data_ptr<int32_t>(),
escalated.data_ptr<int32_t>(),
deflated.data_ptr<int32_t>(),
updated64.data_ptr<double>(),
info.data_ptr<int32_t>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
fp64_secular_columns_4x256_kernel<<<
tiled_grid, kTileWidth, 0>>>(
poles.data_ptr<float>(),
roots64.data_ptr<double>(),
updated64.data_ptr<double>(),
active_count.data_ptr<int32_t>(),
merge_mode.data_ptr<int32_t>(),
escalated.data_ptr<int32_t>(),
deflated.data_ptr<int32_t>(),
values.data_ptr<float>(),
secular_vectors.data_ptr<float>(),
info.data_ptr<int32_t>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
finalize_values_info_kernel<<<
static_cast<unsigned int>(batch), 256, 0>>>(
values.data_ptr<float>(),
active_count.data_ptr<int32_t>(),
merge_mode.data_ptr<int32_t>(),
escalated.data_ptr<int32_t>(),
route_summary.data_ptr<int32_t>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
} // namespace
#define N1024_FULL_ARGS \
poles, weights, rho, active_count, values, roots64, root_status, \
root_iterations, merge_mode, repair_ids, repair_count, deflated, info, \
maximum_iterations, updated32, updated64, recon_status, escalated, \
escalation_ids, escalation_count, secular_vectors, route_summary
void n1024_merge_full_run(
const at::Tensor& poles,
const at::Tensor& weights,
const at::Tensor& rho,
const at::Tensor& active_count,
const at::Tensor& values,
const at::Tensor& roots64,
const at::Tensor& root_status,
const at::Tensor& root_iterations,
const at::Tensor& merge_mode,
const at::Tensor& repair_ids,
const at::Tensor& repair_count,
const at::Tensor& deflated,
const at::Tensor& info,
const at::Tensor& maximum_iterations,
const at::Tensor& updated32,
const at::Tensor& updated64,
const at::Tensor& recon_status,
const at::Tensor& escalated,
const at::Tensor& escalation_ids,
const at::Tensor& escalation_count,
const at::Tensor& secular_vectors,
const at::Tensor& route_summary) {
n1024_merge_full_impl(0, N1024_FULL_ARGS);
}
void n1024_merge_full_fp64_run(
const at::Tensor& poles,
const at::Tensor& weights,
const at::Tensor& rho,
const at::Tensor& active_count,
const at::Tensor& values,
const at::Tensor& roots64,
const at::Tensor& root_status,
const at::Tensor& root_iterations,
const at::Tensor& merge_mode,
const at::Tensor& repair_ids,
const at::Tensor& repair_count,
const at::Tensor& deflated,
const at::Tensor& info,
const at::Tensor& maximum_iterations,
const at::Tensor& updated32,
const at::Tensor& updated64,
const at::Tensor& recon_status,
const at::Tensor& escalated,
const at::Tensor& escalation_ids,
const at::Tensor& escalation_count,
const at::Tensor& secular_vectors,
const at::Tensor& route_summary) {
n1024_merge_full_impl(1, N1024_FULL_ARGS);
}
void n1024_merge_polished_run(
const at::Tensor& poles,
const at::Tensor& weights,
const at::Tensor& rho,
const at::Tensor& active_count,
const at::Tensor& values,
const at::Tensor& roots64,
const at::Tensor& root_status,
const at::Tensor& root_iterations,
const at::Tensor& merge_mode,
const at::Tensor& repair_ids,
const at::Tensor& repair_count,
const at::Tensor& deflated,
const at::Tensor& info,
const at::Tensor& maximum_iterations,
const at::Tensor& updated32,
const at::Tensor& updated64,
const at::Tensor& recon_status,
const at::Tensor& escalated,
const at::Tensor& escalation_ids,
const at::Tensor& escalation_count,
const at::Tensor& secular_vectors,
const at::Tensor& route_summary) {
n1024_merge_full_impl(2, N1024_FULL_ARGS);
}
#undef N1024_FULL_ARGS
} // namespace lane_p
"""
# ---------------------------------------------------------------------------
# n512 route: five-stage reduction pipeline, strict FP32 device math with
# fp64-controlled root solves, per-row fail-closed exact fallback.
# ---------------------------------------------------------------------------
_CHAIN_N = 512
_CHAIN_KD = 8
_CHAIN_SLOTS = 16512
_CHAIN_STEPS = tuple(range(0, _CHAIN_N - _CHAIN_KD, _CHAIN_KD))
_TREE_WIDTHS = (64, 128, 256, 512)
# One-line tree-mode switch: "lower-redesign-packed" ships; "endpoint-all"
# (the previous certified route) stays fully packaged behind this constant.
_TREE_MODE = "lower-redesign-v4"
_TREE_BLOCK = 128
_REWRITE_WIDTHS = (256, 512)
_REWRITE_CTOL = 8.0 * float(torch.finfo(torch.float32).eps)
_REWRITE_STOP = float(torch.finfo(torch.float64).eps)
_Q1_GROUP = 8
# Q1 engine: "tagg-2gemm" = hoisted batched T-build (larft recurrence,
# zero host syncs) + 2-GEMM U-precompute apply; "grouped-wy" = the
# previous engine, kept switchable. Strict IEEE fp32 only (the fast-GEMM
# Q1 arm is a refuted lever — orthogonality-gate receipts in BUILD_LOG).
# The two apply GEMMs live in _q1_tagg_apply_'s loop as the single
# drop-in point for any future precision arm. "triton-fused" replaces
# them with one fused dual-GEMM kernel per group (three compensated
# half-word products into one fp32 accumulator; W never touches DRAM);
# it falls back to "tagg-2gemm" automatically where triton is missing.
_Q1_MODE = "triton-fused"
_S1_GROUP = 8
# Precision policy for chain rows: attempt the fast-GEMM reduction arm
# first (official-envelope-legal; its own battery measured 48x margin),
# gate it per row against the actual checker quantity, and recompute the
# whole batch on the strict path whenever the flagged cohort exceeds
# 1/20 of the rows. The strict path is byte-identical to the previous
# submission's pipeline.
_N512_FAST = True
# Q1 fast-GEMM arm MEASURED ILLEGAL on the hosted orthogonality gate
# (L1 defect 0.032-0.043 vs 0.0061 allowed, uniform across rows —
# ~1e-3/column tf32 error summed over 512 columns). The ledger lever is
# refuted for this checker; Q1 ships strict. Flag retained for the record.
_Q1_FAST = False
_FAST_COHORT_NUM = 1
_FAST_COHORT_DEN = 20
_DENSE_NET = 0.5
_STAGE2_CONFIGS = (("m5s6", 5, 190240), ("m2s16", 2, 76096))
# The reduction route is validated on O(1)-scaled spectra. Rows near the
# sqrt(FLT_MAX)/sqrt(FLT_MIN) magnitude extremes stay finite but leave the
# residual envelope (measured), so batches containing them take the exact
# path instead. Zero rows are fine (they already fail closed in-route).
_CHAIN_SCALE_LO = 2.0 ** -20
_CHAIN_SCALE_HI = 2.0 ** 20
# Per-row residual net after the root-solve stage: a small tail of rows
# (~2-4 per 640, measured on both B200 and GB10) keeps finite outputs and a
# zero status word yet leaves the checker envelope. Rows whose intermediate
# residual exceeds this fraction of the envelope take the exact path.
_EPS32 = float(torch.finfo(torch.float32).eps)
_RESIDUAL_NET = 0.25
# Orthogonality net threshold. Degenerate-spectrum rows can satisfy the
# eigen residual while losing mutual orthogonality inside an eigenvalue
# cluster (measured: 8/640 mixed rows up to 79x the envelope, 1/640 rankdef
# at 1.37x, dense rows at 0.009-0.017x). The net measures the fp32 Gram
# L1 defect per row; fp32 evaluation noise is ~0.09x of the envelope, the
# flagged class starts at ~0.75x, and any true over-envelope row reads
# >=0.9x, so 0.35 has >=2x two-sided margin.
_ORTH_NET = 0.35
# Circuit breaker: spectra whose distribution mass-trips the residual net
# (measured hosted: mixed/rankdef rows) make per-row rescue the wrong tool.
# Past this cohort fraction the whole batch is recomputed on the banked
# torch path instead, bounding the worst case at chain-wasted + torch.
_BREAKER_NUMERATOR = 1
_BREAKER_DENOMINATOR = 10
@lru_cache(maxsize=1)
def _chain_extension():
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required")
digest = hashlib.sha256(
(_CHAIN_CPP + _CHAIN_CU).encode()).hexdigest()[:12]
return load_inline(
name=f"eigh_c_{digest}",
cpp_sources=_CHAIN_CPP,
cuda_sources=_CHAIN_CU,
functions=None,
extra_cflags=["-O3", "-std=c++17"],
extra_cuda_cflags=["-O3", "--fmad=true"],
extra_ldflags=["-lcusolver"],
with_cuda=True,
verbose=False,
)
@lru_cache(maxsize=8)
def _stage2_config(device_index: int):
limit = torch.cuda.get_device_properties(
device_index).shared_memory_per_block_optin
for name, per_cta, shared in _STAGE2_CONFIGS:
if shared <= limit:
return name, per_cta
raise RuntimeError("no stage-2 configuration fits this device")
class _FastGemm:
# Scoped fast-GEMM precision switch (restores the import-time pin
# exactly on exit; every other torch op in the module runs strict).
def __enter__(self):
torch.backends.cuda.matmul.allow_tf32 = True
try:
torch.backends.cuda.matmul.fp32_precision = "tf32"
except Exception:
pass
return self
def __exit__(self, *exc):
torch.backends.cuda.matmul.allow_tf32 = False
torch.backends.cudnn.allow_tf32 = False
try:
torch.backends.cuda.matmul.fp32_precision = "ieee"
except Exception:
pass
torch.set_float32_matmul_precision("highest")
return False
def _reduce_to_band8_fast(module, original):
# Aggregated-trailing reduction (group compact-WY frames, one far-field
# update per group). Caller wraps this in _FastGemm; the panel factor
# kernel itself is strict and unaffected by the GEMM precision switch.
batch, n, _ = original.shape
kd = _CHAIN_KD
device = original.device
matrix = original.clone(memory_format=torch.contiguous_format)
band = torch.zeros((batch, kd + 1, n), dtype=torch.float32,
device=device)
steps = list(_CHAIN_STEPS)
groups = [steps[i:i + _S1_GROUP]
for i in range(0, len(steps), _S1_GROUP)]
m0max = n - kd
gwmax = 2 * kd * _S1_GROUP
vw_full = torch.zeros((batch, m0max, gwmax), dtype=torch.float32,
device=device)
wv_full = torch.zeros((batch, m0max, gwmax), dtype=torch.float32,
device=device)
flat_a = torch.empty(batch * m0max * kd, dtype=torch.float32,
device=device)
flat_b = torch.empty(batch * m0max * kd, dtype=torch.float32,
device=device)
flat_c = torch.empty(batch * gwmax * kd, dtype=torch.float32,
device=device)
s1 = torch.empty((batch, kd, kd), dtype=torch.float32, device=device)
sb = torch.empty((batch, kd, kd), dtype=torch.float32, device=device)
uflat = torch.empty(batch * (n - 2 * kd) * (n - 2 * kd),
dtype=torch.float32, device=device)
reflectors = []
for grp in groups:
c0 = grp[0]
r0 = c0 + kd
m0 = n - r0
gw = 2 * kd * len(grp)
vw = vw_full[:, :m0, :gw]
wv = wv_full[:, :m0, :gw]
for p, cp in enumerate(grp):
ts = cp + kd
rp = n - ts
fr = kd * p
k = 2 * kd * p
if p > 0:
rows0 = cp - r0
nrows = n - cp
tmp = flat_a[: batch * nrows * kd].view(batch, nrows, kd)
torch.bmm(vw[:, rows0:, :k],
wv[:, rows0:rows0 + kd, :k].transpose(1, 2),
out=tmp)
matrix[:, cp:, cp:cp + kd].sub_(tmp)
v = torch.empty((batch, rp, kd), dtype=torch.float32,
device=device)
t = torch.empty((batch, kd, kd), dtype=torch.float32,
device=device)
tau = torch.empty((batch, kd), dtype=torch.float32,
device=device)
module.panel_qr8(matrix, v, t, tau, cp)
gather = matrix.as_strided(
(batch, kd + 1, kd), (n * n, n, n + 1),
matrix.storage_offset() + cp * (n + 1))
band[:, :, cp:cp + kd].copy_(gather)
y0 = flat_a[: batch * m0 * kd].view(batch, m0, kd)
torch.bmm(matrix[:, r0:, ts:], v, out=y0)
if p > 0:
c = flat_c[: batch * k * kd].view(batch, k, kd)
torch.bmm(wv[:, fr:, :k].transpose(1, 2), v, out=c)
y0.baddbmm_(vw[:, :, :k], c, beta=1, alpha=-1)
yb = flat_b[: batch * m0 * kd].view(batch, m0, kd)
torch.bmm(y0, t, out=yb)
torch.bmm(v.transpose(1, 2), yb[:, fr:, :], out=s1)
torch.bmm(t.transpose(1, 2), s1, out=sb)
tmp = flat_a[: batch * rp * kd].view(batch, rp, kd)
torch.bmm(v, sb, out=tmp)
yb[:, fr:, :].sub_(tmp, alpha=0.5)
vb = 2 * kd * p
if fr > 0:
vw[:, :fr, vb:vb + kd].zero_()
wv[:, :fr, vb + kd:vb + 2 * kd].zero_()
vw[:, fr:, vb:vb + kd].copy_(v)
vw[:, :, vb + kd:vb + 2 * kd].copy_(yb)
wv[:, :, vb:vb + kd].copy_(yb)
wv[:, fr:, vb + kd:vb + 2 * kd].copy_(v)
reflectors.append((ts, v, t))
fs = grp[-1] + kd
ff = fs - r0
if fs < n:
mt = n - fs
u = uflat[: batch * mt * mt].view(batch, mt, mt)
torch.bmm(vw[:, ff:, :], wv[:, ff:, :].transpose(1, 2), out=u)
matrix[:, fs:, fs:].sub_(u)
for column in range(n - kd, n):
length = min(kd, n - 1 - column) + 1
band[:, :length, column].copy_(
matrix[:, column:column + length, column])
return band, reflectors
def _reduce_to_band8(module, original):
batch = original.shape[0]
n, kd = _CHAIN_N, _CHAIN_KD
matrix = original.clone(memory_format=torch.contiguous_format)
band = torch.zeros((batch, kd + 1, n), dtype=original.dtype,
device=original.device)
reflectors = []
for panel_start in _CHAIN_STEPS:
trailing_start = panel_start + kd
panel = matrix[
:, trailing_start:, panel_start:panel_start + kd].contiguous()
qr = torch.empty_like(panel)
tau = torch.empty((batch, kd), dtype=panel.dtype, device=panel.device)
v = torch.empty_like(panel)
t = torch.empty((batch, kd, kd), dtype=panel.dtype,
device=panel.device)
module.panel_qr_w8_vt_out(panel, qr, tau, v, t)
matrix[:, trailing_start:trailing_start + kd,
panel_start:panel_start + kd].copy_(qr[:, :kd])
gather = matrix.as_strided(
(batch, kd + 1, kd),
(n * n, n, n + 1),
matrix.storage_offset() + panel_start * (n + 1))
band[:, :, panel_start:panel_start + kd].copy_(gather)
vt = torch.bmm(v, t)
w = torch.empty_like(vt)
module.skinny_w_out(matrix, vt, w, trailing_start)
small = torch.bmm(vt.transpose(1, 2), w)
w = w - 0.5 * torch.bmm(v, small)
left = torch.cat((v, w), dim=2).contiguous()
right = torch.cat(
(w.transpose(1, 2), v.transpose(1, 2)), dim=1).contiguous()
module.update_rank16_inplace(matrix, left, right, trailing_start)
reflectors.append((trailing_start, v, t))
for column in range(n - kd, n):
length = min(kd, n - 1 - column) + 1
band[:, :length, column].copy_(
matrix[:, column:column + length, column])
return band, reflectors
def _band8_to_tridiagonal(module, band):
batch = band.shape[0]
config, per_cta = _stage2_config(band.get_device())
padded = ((batch + per_cta - 1) // per_cta) * per_cta
band_in = band
if padded != batch:
band_in = torch.cat(
(band, band[-1:].expand(padded - batch, -1, -1)),
dim=0).contiguous()
d = torch.empty((padded, _CHAIN_N), dtype=torch.float32,
device=band.device)
e = torch.empty((padded, _CHAIN_N - 1), dtype=torch.float32,
device=band.device)
hh = torch.empty((padded, _CHAIN_SLOTS, _CHAIN_KD), dtype=torch.float32,
device=band.device)
info = torch.empty((padded,), dtype=torch.int32, device=band.device)
module.band8_to_tridiagonal_kd8_out(band_in, d, e, hh, info, config)
if padded != batch:
d, e = d[:batch].contiguous(), e[:batch].contiguous()
hh, info = hh[:batch], info[:batch]
return d, e, hh, info
def _tree_prepare_level(values, vectors, original_e, width, inherited_bad):
batch, child_nodes, half = values.shape
parents = child_nodes // 2
left_values, right_values = values[:, 0::2], values[:, 1::2]
left_q, right_q = vectors[:, 0::2], vectors[:, 1::2]
poles = torch.cat((left_values, right_values), dim=-1)
starts = torch.arange(parents, device=values.device,
dtype=torch.long) * width
cuts = starts + half - 1
beta = original_e[:, cuts]
weights = torch.cat(
(
left_q[:, :, -1, :],
torch.sign(beta).unsqueeze(-1) * right_q[:, :, 0, :],
),
dim=-1,
) / math.sqrt(2.0)
sorted_poles, order = torch.sort(poles, dim=-1, stable=True)
sorted_weights = torch.gather(weights, -1, order)
gaps = sorted_poles[..., 1:] - sorted_poles[..., :-1]
rho = 2.0 * beta.abs()
pre_bad = (
(~torch.isfinite(sorted_poles).all(dim=-1))
| (~torch.isfinite(sorted_weights).all(dim=-1))
| (gaps <= 0.0).any(dim=-1)
| (~torch.isfinite(rho))
| (rho <= 0.0)
)
# The solver kernels deflate tiny weights themselves or fail closed with
# a nonzero status, so no tiny-weight route guard is needed here.
matrix_bad = inherited_bad | pre_bad.any(dim=1)
safe_node = pre_bad | matrix_bad[:, None]
basis = torch.zeros(
(batch, parents, width, width),
dtype=torch.float32,
device=values.device,
)
basis[:, :, :half, :half] = left_q
basis[:, :, half:, half:] = right_q
basis = torch.gather(
basis, -1, order.unsqueeze(-2).expand(batch, parents, width, width))
# Substitute a well-posed dummy problem on rejected nodes so the solver
# kernels never see invalid data; rejected rows are replaced by the exact
# dense fallback at the end of the route.
safe_poles = torch.linspace(
-1.0, 1.0, width, dtype=torch.float32, device=values.device)
safe_weights = torch.full(
(width,),
1.0 / math.sqrt(float(width)),
dtype=torch.float32,
device=values.device,
)
eye = torch.eye(width, dtype=torch.float32, device=values.device)
sorted_poles = torch.where(safe_node[..., None], safe_poles, sorted_poles)
sorted_weights = torch.where(
safe_node[..., None], safe_weights, sorted_weights)
rho = torch.where(safe_node, torch.ones_like(rho), rho)
basis = torch.where(safe_node[..., None, None], eye, basis)
return (
sorted_poles.contiguous(),
sorted_weights.contiguous(),
rho.contiguous(),
basis.contiguous(),
matrix_bad,
)
def _run_tree(module, d, e):
if _TREE_MODE == "endpoint-all":
return _run_tree_endpoint(module, d, e)
if _TREE_MODE == "lower-redesign-v4":
return _run_tree_packed(module, d, e)
raise RuntimeError(f"unpackaged tree mode: {_TREE_MODE}")
def _rewrite_prepare(module, values, vectors, original_e, width,
inherited_bad):
batch, child_nodes, half = values.shape
parents = child_nodes // 2
device = values.device
left_values, right_values = values[:, 0::2], values[:, 1::2]
left_q, right_q = vectors[:, 0::2], vectors[:, 1::2]
poles = torch.cat((left_values, right_values), dim=-1)
starts = torch.arange(parents, device=device, dtype=torch.long) * width
cuts = starts + half - 1
beta = original_e[:, cuts]
weights = torch.cat(
(
left_q[:, :, -1, :],
torch.sign(beta).unsqueeze(-1) * right_q[:, :, 0, :],
),
dim=-1,
) / math.sqrt(2.0)
sorted_poles, order = torch.sort(poles, dim=-1, stable=True)
sorted_weights = torch.gather(weights, -1, order)
rho = 2.0 * beta.abs()
pre_bad = (
(~torch.isfinite(sorted_poles).all(dim=-1))
| (~torch.isfinite(sorted_weights).all(dim=-1))
| (~torch.isfinite(rho))
| (rho <= 0.0)
)
matrix_bad = inherited_bad | pre_bad.any(dim=1)
safe_node = pre_bad | matrix_bad[:, None]
safe_poles = torch.linspace(-1.0, 1.0, width, dtype=torch.float32,
device=device)
safe_weights = torch.full(
(width,), 1.0 / math.sqrt(float(width)),
dtype=torch.float32, device=device)
flat = batch * parents
basis = torch.empty((batch, parents, width, width), dtype=torch.float32,
device=device)
children = vectors.reshape(batch * child_nodes, half, half).contiguous()
order32 = order.reshape(flat, width).to(torch.int32).contiguous()
safe8 = safe_node.reshape(flat).to(torch.uint8).contiguous()
module.build_basis_(children, order32, safe8,
basis.view(flat, width, width), half)
sorted_poles = torch.where(safe_node[..., None], safe_poles,
sorted_poles)
sorted_weights = torch.where(safe_node[..., None], safe_weights,
sorted_weights)
rho = torch.where(safe_node, torch.ones_like(rho), rho)
# clone-on-demand escalation stash: build_basis_ is deterministic from
# these, so the unrotated basis can be rebuilt for escalated nodes only.
prep_stash = (children, order32, safe8)
return (sorted_poles.contiguous(), sorted_weights.contiguous(),
rho.contiguous(), basis, order.contiguous(), matrix_bad,
prep_stash)
def _rewrite_merge_level(module, width, poles, weights, rho, basis, order,
flat, half, prep_stash):
device = poles.device
srcblock = (order >= half).to(torch.int32).contiguous()
d_adj = torch.empty((flat, width), dtype=torch.float64, device=device)
z_adj = torch.empty_like(d_adj)
active = torch.empty((flat, width), dtype=torch.int32, device=device)
perm_cols = torch.empty_like(active)
counts = torch.zeros((flat, 5), dtype=torch.int32, device=device)
rot_idx = torch.zeros((flat, width, 2), dtype=torch.int32, device=device)
rot_cs = torch.zeros((flat, width, 2), dtype=torch.float64,
device=device)
module.deflate_(poles, weights, rho, _REWRITE_CTOL, half, srcblock,
d_adj, z_adj, active, perm_cols, counts, rot_idx,
rot_cs)
n12 = counts[:, 1] + counts[:, 2]
n23 = counts[:, 0] - counts[:, 1]
n4 = width - counts[:, 0]
maxes = torch.stack((n12.max(), n23.max(), n4.max())).to("cpu")
n12_max, n23_max, n4_max = int(maxes[0]), int(maxes[1]), int(maxes[2])
# clone-on-demand: the unrotated basis is rebuilt lazily for escalated
# nodes only (deterministic from prep_stash), so no full clone here.
module.apply_givens_(basis, rot_idx, rot_cs, counts)
active_l = active.long()
d_act = torch.gather(d_adj, 1, active_l).contiguous()
z_act = torch.gather(z_adj, 1, active_l).contiguous()
origins = torch.zeros((flat, width), dtype=torch.int32, device=device)
taus = torch.zeros((flat, width), dtype=torch.float64, device=device)
dlambda = torch.zeros_like(taus)
iters = torch.zeros((flat, width), dtype=torch.int32, device=device)
sec_info = torch.zeros((flat,), dtype=torch.int32, device=device)
module.secular_recip_(d_act, z_act, rho, counts, _REWRITE_STOP,
origins, taus, dlambda, iters, sec_info)
k_col = counts[:, 0].long()
lane = torch.arange(width, device=device)[None, :]
active_mask = lane < k_col[:, None]
mono_bad = ((d_act.diff(dim=1) <= 0.0)
& (lane[:, 1:] < k_col[:, None])).any(dim=1)
escalate_mask = sec_info.ne(0) | mono_bad
zhat = torch.zeros((flat, width), dtype=torch.float64, device=device)
module.loewner_warp_(d_act, z_act, rho, counts, origins, taus, zhat)
vals_full = torch.where(active_mask, dlambda,
torch.gather(d_adj, 1, active_l))
svals, sidx = torch.sort(vals_full, dim=1, stable=True)
col_out = torch.empty_like(sidx)
col_out.scatter_(1, sidx, lane.expand_as(sidx))
col_out = col_out.to(torch.int32).contiguous()
inv = torch.empty((flat, width), dtype=torch.int32, device=device)
module.inv_perm_(perm_cols, inv)
s_top = torch.zeros((flat, n12_max, width), dtype=torch.float32,
device=device)
s_bot = torch.zeros((flat, n23_max, width), dtype=torch.float32,
device=device)
if n12_max or n23_max:
module.vector_build_packed_warp_(d_act, zhat, counts, origins, taus,
active, inv, col_out, s_top, s_bot)
out_values = svals.to(torch.float32)
merged = torch.empty((flat, width, width), dtype=torch.float32,
device=device)
if n12_max:
bp_top = torch.empty((flat, half, n12_max), dtype=torch.float32,
device=device)
module.pack_bp_(basis, perm_cols, counts, bp_top, half, 0)
merged[:, :half, :] = torch.bmm(bp_top, s_top)
else:
merged[:, :half, :].zero_()
if n23_max:
bp_bot = torch.empty((flat, half, n23_max), dtype=torch.float32,
device=device)
module.pack_bp_(basis, perm_cols, counts, bp_bot, half, 1)
merged[:, half:, :] = torch.bmm(bp_bot, s_bot)
else:
merged[:, half:, :].zero_()
if n4_max:
module.defl_epilogue_(basis, perm_cols, col_out, counts, merged)
# fail-closed per-node escalation to the certified solvers
info_final = torch.zeros((flat,), dtype=torch.int32, device=device)
escalate = torch.nonzero(escalate_mask, as_tuple=False).flatten()
escalated = int(escalate.numel())
if escalated:
solver = {256: module.merge256_solve,
512: module.merge512_solve}[width]
sub_info = torch.zeros((escalated,), dtype=torch.int32,
device=device)
sub_iters = torch.zeros_like(sub_info)
sub_values = torch.empty((escalated, width), dtype=torch.float32,
device=device)
sub_secular = torch.empty((escalated, width, width),
dtype=torch.float32, device=device)
cnt = torch.full((escalated,), width, dtype=torch.int32,
device=device)
solver(poles.index_select(0, escalate).contiguous(),
weights.index_select(0, escalate).contiguous(),
rho.index_select(0, escalate).contiguous(),
cnt, sub_values, sub_secular, sub_info, sub_iters)
children, order32, safe8 = prep_stash
esc_children = children.view(flat, 2, half, half) .index_select(0, escalate) .reshape(2 * escalated, half, half).contiguous()
b_orig_sub = torch.empty((escalated, width, width),
dtype=torch.float32, device=device)
module.build_basis_(
esc_children,
order32.index_select(0, escalate).contiguous(),
safe8.index_select(0, escalate).contiguous(),
b_orig_sub, half)
sub_merged = torch.bmm(b_orig_sub, sub_secular)
out_values.index_copy_(0, escalate, sub_values)
merged.index_copy_(0, escalate, sub_merged)
info_final.index_copy_(0, escalate, sub_info)
return out_values, merged, info_final
def _run_tree_packed(module, d, e):
batch = d.shape[0]
device = d.device
# power-of-two prescale (exact in fp32; removes the inherited official-
# gate failure class at extreme input scales; values scaled back at exit,
# vectors unchanged)
amax = torch.maximum(d.abs().amax(dim=1), e.abs().amax(dim=1))
ok = torch.isfinite(amax) & (amax > 0.0)
gamma = torch.where(
ok,
torch.exp2(torch.floor(torch.log2(
amax.clamp_min(torch.finfo(torch.float32).tiny)))),
torch.ones_like(amax))
ds = (d / gamma[:, None]).contiguous()
es = (e / gamma[:, None]).contiguous()
# lower tree: batched bisect + inverse iteration directly on the four
# Cuppen-adjusted 128-blocks
nblk = _CHAIN_N // _TREE_BLOCK
cuts = torch.arange(_TREE_BLOCK - 1, _CHAIN_N - 1, _TREE_BLOCK,
device=device)
adjusted = ds.clone()
coupling = es[:, cuts].abs()
adjusted[:, cuts] -= coupling
adjusted[:, cuts + 1] -= coupling
blocks = batch * nblk
d_blocks = adjusted.view(batch, nblk, _TREE_BLOCK).reshape(
blocks, _TREE_BLOCK).contiguous()
local = torch.arange(_TREE_BLOCK - 1, device=device)
starts = torch.arange(nblk, device=device)[:, None] * _TREE_BLOCK
e_blocks = es[:, (starts + local).reshape(-1)].reshape(
blocks, _TREE_BLOCK - 1).contiguous()
lam = torch.empty((blocks, _TREE_BLOCK), dtype=torch.float32,
device=device)
q = torch.empty((blocks, _TREE_BLOCK, _TREE_BLOCK), dtype=torch.float32,
device=device)
info = torch.zeros((blocks,), dtype=torch.int32, device=device)
stats = torch.zeros((blocks, 16), dtype=torch.int32, device=device)
module.lower_solve_(d_blocks, e_blocks, lam, q, info, stats)
bad_matrix = info.view(batch, nblk).ne(0).any(dim=1)
values = lam.view(batch, nblk, _TREE_BLOCK)
vectors = q.view(batch, nblk, _TREE_BLOCK, _TREE_BLOCK)
for width in _REWRITE_WIDTHS:
if (int(bad_matrix.sum().item()) * _BREAKER_DENOMINATOR
> batch * _BREAKER_NUMERATOR):
return None, None, bad_matrix
half = width // 2
(poles, weights, rho, basis, order, bad_matrix,
prep_stash) = _rewrite_prepare(
module, values, vectors, es, width, bad_matrix)
parents = poles.shape[1]
flat = batch * parents
out_values, merged, info_final = _rewrite_merge_level(
module, width, poles.reshape(flat, width),
weights.reshape(flat, width), rho.reshape(flat),
basis.view(flat, width, width), order.reshape(flat, width),
flat, half, prep_stash)
bad_matrix = bad_matrix | info_final.view(batch, parents).ne(0).any(
dim=1)
values = out_values.view(batch, parents, width)
vectors = merged.view(batch, parents, width, width)
values = values[:, 0].contiguous()
vectors = vectors[:, 0]
values.mul_(gamma[:, None])
return values, vectors, bad_matrix
def _run_tree_endpoint(module, d, e):
batch = d.shape[0]
device = d.device
leaf = 32
leaves = _CHAIN_N // leaf
adjusted = d.clone()
cuts = torch.arange(leaf - 1, _CHAIN_N - 1, leaf, device=device)
coupling = e[:, cuts].abs()
adjusted[:, cuts] -= coupling
adjusted[:, cuts + 1] -= coupling
leaf_d = adjusted.view(batch, leaves, leaf).reshape(-1, leaf).contiguous()
local = torch.arange(leaf - 1, device=device)
leaf_starts = torch.arange(leaves, device=device)[:, None] * leaf
leaf_e = e[:, (leaf_starts + local).reshape(-1)].reshape(
-1, leaf - 1).contiguous()
leaf_qt = torch.empty((batch * leaves, leaf, leaf), dtype=torch.float32,
device=device)
leaf_info = torch.zeros((batch * leaves,), dtype=torch.int32,
device=device)
module.leaf32_run(leaf_d, leaf_e, leaf_qt, leaf_info)
values = leaf_d.view(batch, leaves, leaf)
vectors = leaf_qt.view(batch, leaves, leaf, leaf).transpose(-1, -2)
bad_matrix = leaf_info.view(batch, leaves).ne(0).any(dim=1)
solvers = {
64: module.merge64_solve,
128: module.merge128_solve,
256: module.merge256_solve,
512: module.merge512_solve,
}
for width in _TREE_WIDTHS:
poles, weights, rho, basis, bad_matrix = _tree_prepare_level(
values, vectors, e, width, bad_matrix)
if (int(bad_matrix.sum().item()) * _BREAKER_DENOMINATOR
> batch * _BREAKER_NUMERATOR):
# Early circuit breaker: distributions this route rejects
# surface in the level preconditions (some at the first level,
# some only after merging collapses their spectra). Bail as
# soon as the rejected fraction crosses the line, before the
# remaining levels and both back-substitution stages are spent.
# One host round-trip per level; measured neutral on the
# accepted-path rows.
return None, None, bad_matrix
parents = poles.shape[1]
flat = batch * parents
poles_flat = poles.reshape(flat, width)
weights_flat = weights.reshape(flat, width)
rho_flat = rho.reshape(flat)
counts = torch.full((flat,), width, dtype=torch.int32, device=device)
out_values = torch.empty_like(poles_flat)
secular = torch.empty((flat, width, width), dtype=torch.float32,
device=device)
info = torch.zeros((flat,), dtype=torch.int32, device=device)
iterations = torch.zeros_like(info)
solvers[width](poles_flat, weights_flat, rho_flat, counts,
out_values, secular, info, iterations)
bad_matrix = bad_matrix | info.view(batch, parents).ne(0).any(dim=1)
merged = torch.empty_like(secular)
torch.bmm(basis.reshape(flat, width, width), secular, out=merged)
values = out_values.view(batch, parents, width)
vectors = merged.view(batch, parents, width, width)
return values[:, 0], vectors[:, 0], bad_matrix
def _q1_gemm(left, right, fast):
if not fast:
return torch.bmm(left, right)
with _FastGemm():
return torch.bmm(left, right)
def _q1_group_merge(reflectors, a, b, batch, device, dtype):
kd = _CHAIN_KD
ts_a = reflectors[a][0]
frame = _CHAIN_N - ts_a
rank = kd * (b - a + 1)
vg = torch.zeros((batch, frame, rank), device=device, dtype=dtype)
tau = torch.empty((batch, rank), device=device, dtype=dtype)
for k, (ts, v, t) in enumerate(reflectors[a:b + 1]):
off = ts - ts_a
vg[:, off:, kd * k:kd * (k + 1)] = v
tau[:, kd * k:kd * (k + 1)] = torch.diagonal(t, dim1=1, dim2=2)
zero = tau == 0
if bool(zero.any()):
vg = vg * (~zero).to(dtype).unsqueeze(1)
tau = torch.where(zero, torch.ones_like(tau), tau)
# The small merge system and its triangular solve stay strict FP32 in
# both engine arms.
s = torch.bmm(vg.transpose(1, 2), vg)
m = torch.triu(s, diagonal=1)
m.diagonal(dim1=1, dim2=2).copy_(tau.reciprocal())
eye = torch.eye(rank, device=device, dtype=dtype)
eye = eye.unsqueeze(0).expand(batch, rank, rank)
tg = torch.linalg.solve_triangular(m, eye, upper=True)
return vg, tg, ts_a
try:
import triton
import triton.language as tl
_TRITON_OK = True
except Exception:
_TRITON_OK = False
if _TRITON_OK:
@triton.jit
def _q1t_split3(x):
hi = x.to(tl.float16)
hif = hi.to(tl.float32)
hi6 = (hif * 0.015625).to(tl.float16)
res6 = ((x - hif) * 64.0).to(tl.float16)
return hi, hi6, res6
@triton.jit
def _q1t_fused_group(Vp, Up, Xp, f, ts, sVU, sX,
BN: tl.constexpr, BK: tl.constexpr):
pid = tl.program_id(0)
ncol = 512 // BN
b = pid // ncol
cb = (pid % ncol) * BN
offs_r = tl.arange(0, 64)
offs_n = cb + tl.arange(0, BN)
accw = tl.zeros((64, BN), dtype=tl.float32)
for k0 in range(0, f, BK):
offs_k = k0 + tl.arange(0, BK)
mk = offs_k < f
v = tl.load(Vp + b * sVU + offs_k[:, None] * 64
+ offs_r[None, :], mask=mk[:, None], other=0.0)
x = tl.load(Xp + b * sX + (ts + offs_k)[:, None] * 512
+ offs_n[None, :], mask=mk[:, None], other=0.0)
vh, vh6, vr6 = _q1t_split3(v)
xh, xh6, xr6 = _q1t_split3(x)
accw = tl.dot(tl.trans(vh), xh, accw)
accw = tl.dot(tl.trans(vh6), xr6, accw)
accw = tl.dot(tl.trans(vr6), xh6, accw)
wh, wh6, wr6 = _q1t_split3(accw)
for m0 in range(0, f, BK):
offs_m = m0 + tl.arange(0, BK)
mm = offs_m < f
u = tl.load(Up + b * sVU + offs_m[:, None] * 64
+ offs_r[None, :], mask=mm[:, None], other=0.0)
uh, uh6, ur6 = _q1t_split3(u)
d = tl.dot(uh, wh)
d = tl.dot(uh6, wr6, d)
d = tl.dot(ur6, wh6, d)
xptr = Xp + b * sX + (ts + offs_m)[:, None] * 512 + offs_n[None, :]
x = tl.load(xptr, mask=mm[:, None], other=0.0)
tl.store(xptr, x - d, mask=mm[:, None])
def _q1_triton_apply_(reflectors, x):
# T-agg strict build + one fused dual-GEMM kernel per group (rank-64
# fixed; the ragged tail is zero-padded exactly). n=512 only.
batch = x.shape[0]
built = _q1_tagg_build(reflectors, batch, x.device, x.dtype)
padded = []
for ts_a, vg, ug in built:
if vg.shape[2] != 64:
pad = 64 - vg.shape[2]
vg = torch.nn.functional.pad(vg, (0, pad)).contiguous()
ug = torch.nn.functional.pad(ug, (0, pad)).contiguous()
else:
vg = vg.contiguous()
ug = ug.contiguous()
padded.append((ts_a, vg, ug))
for ts_a, vg, ug in reversed(padded):
f = vg.shape[1]
grid = (batch * (512 // 64),)
_q1t_fused_group[grid](vg, ug, x, f, ts_a, f * 64, 512 * 512,
BN=64, BK=64, num_warps=8)
return x
def _q1_tagg_t_batch(reflectors, packed, batch, device, dtype):
# T_g for every group: uniform full-size groups stack so each larft
# recurrence level is one matmul pair + one copy; the ragged tail
# group runs the same recurrence per-pair. tau = 0 slots are
# natively inert (their leaf T has row+col zero) — no masking and no
# host syncs anywhere in the build.
kd = _CHAIN_KD
full_rank = kd * _Q1_GROUP
full = [p for p in packed if p[1] == full_rank]
ragged = [p for p in packed if p[1] != full_rank]
tgs = {}
if full:
ng = len(full)
sbig = torch.stack([p[3] for p in full], dim=1)
tbig = torch.zeros((batch, ng, full_rank, full_rank),
device=device, dtype=dtype)
for gi, p in enumerate(full):
a, b = p[4], p[5]
for k in range(a, b + 1):
o = kd * (k - a)
tbig[:, gi, o:o + kd, o:o + kd] = reflectors[k][2]
s = kd
while s < full_rank:
pairs = full_rank // (2 * s)
def blocks(m, ro, co, s=s, pairs=pairs):
return m.as_strided(
(m.shape[0], m.shape[1], pairs, s, s),
(m.stride(0), m.stride(1),
2 * s * (m.stride(2) + m.stride(3)),
m.stride(2), m.stride(3)),
storage_offset=m.storage_offset()
+ ro * m.stride(2) + co * m.stride(3))
tl = blocks(tbig, 0, 0)
tr = blocks(tbig, s, s)
slr = blocks(sbig, 0, s)
cross = -torch.matmul(torch.matmul(tl, slr), tr)
blocks(tbig, 0, s).copy_(cross)
s *= 2
for gi, p in enumerate(full):
tgs[p[4]] = tbig[:, gi]
for p in ragged:
ts_a, rank, vg, s_gram, a, b = p
parts = []
for k in range(a, b + 1):
parts.append((kd * (k - a), kd, reflectors[k][2]))
while len(parts) > 1:
merged_parts = []
for i in range(0, len(parts) - 1, 2):
l0, sl, tl = parts[i]
r0, sr, tr = parts[i + 1]
cross = -torch.bmm(
torch.bmm(tl, s_gram[:, l0:l0 + sl, r0:r0 + sr]), tr)
merged = torch.zeros((batch, sl + sr, sl + sr),
device=device, dtype=dtype)
merged[:, :sl, :sl] = tl
merged[:, :sl, sl:] = cross
merged[:, sl:, sl:] = tr
merged_parts.append((l0, sl + sr, merged))
if len(parts) % 2:
merged_parts.append(parts[-1])
parts = merged_parts
tgs[p[4]] = parts[0][2]
return [tgs[p[4]] for p in packed]
def _q1_tagg_build(reflectors, batch, device, dtype, n=_CHAIN_N):
kd = _CHAIN_KD
npanels = len(reflectors)
groups = []
i = 0
while i < npanels:
j = min(i + _Q1_GROUP, npanels)
groups.append((i, j - 1))
i = j
packed = []
for a, b in groups:
ts_a = reflectors[a][0]
frame = n - ts_a
rank = kd * (b - a + 1)
vg = torch.zeros((batch, frame, rank), device=device, dtype=dtype)
for k, (ts, v, t) in enumerate(reflectors[a:b + 1]):
vg[:, ts - ts_a:, kd * k:kd * (k + 1)] = v
s = torch.bmm(vg.transpose(1, 2), vg)
packed.append((ts_a, rank, vg, s, a, b))
tgs = _q1_tagg_t_batch(reflectors, packed, batch, device, dtype)
built = []
for (ts_a, rank, vg, s, a, b), tg in zip(packed, tgs):
built.append((ts_a, vg, torch.bmm(vg, tg)))
return built
def _q1_tagg_apply_(reflectors, x, n=_CHAIN_N):
expected = (n - _CHAIN_KD) // _CHAIN_KD
if len(reflectors) != expected:
raise RuntimeError(f"expected {expected} panels, "
f"got {len(reflectors)}")
batch = x.shape[0]
built = _q1_tagg_build(reflectors, batch, x.device, x.dtype, n)
for ts_a, vg, ug in reversed(built):
block = x[:, ts_a:, :]
# The two apply GEMMs (precision-arm drop-in point). The baddbmm
# writes into a row-sliced strided view — a regular strided batch
# that keeps the library fast path (lane premise-check receipt).
w = torch.bmm(vg.transpose(1, 2), block)
torch.baddbmm(block, ug, w, beta=1.0, alpha=-1.0, out=block)
return x
def _apply_q1_(reflectors, x, fast=False):
batch = x.shape[0]
npanels = len(reflectors)
groups = []
i = 0
while i < npanels:
j = min(i + _Q1_GROUP, npanels)
groups.append((i, j - 1))
i = j
for a, b in reversed(groups):
vg, tg, ts_a = _q1_group_merge(reflectors, a, b, batch, x.device,
x.dtype)
block = x[:, ts_a:, :]
w = _q1_gemm(vg.transpose(1, 2), block, fast)
w = _q1_gemm(tg, w, fast)
block.sub_(_q1_gemm(vg, w, fast))
return x
def _residual_net_over(d, e, x, values, n=_CHAIN_N):
# Column-scaled L1 residual of the intermediate eigenpairs, evaluated on
# the three-band operator directly (three fused elementwise passes).
tq = d[:, :, None] * x
tq[:, 1:, :] += e[:, :, None] * x[:, :-1, :]
tq[:, :-1, :] += e[:, :, None] * x[:, 1:, :]
tq -= x * values[:, None, :]
resid = tq.abs_().sum(dim=1).amax(dim=1)
scale = d.abs()
scale[:, :-1] += e.abs()
scale[:, 1:] += e.abs()
norm = scale.amax(dim=1).clamp_min(torch.finfo(torch.float32).tiny)
ratio = resid / (200.0 * n * _EPS32 * norm)
return ratio > _RESIDUAL_NET
def _rescue_ok(info, vectors, values):
return (
info.eq(0)
& torch.isfinite(vectors).all(dim=(-2, -1))
& torch.isfinite(values).all(dim=-1)
& (values[:, 1:] >= values[:, :-1]).all(dim=-1)
)
def _chain_exact_fallback(module, data):
# Exact solve for the fail-closed cohort. torch.linalg.eigh dispatches
# small batches to a per-matrix host loop (~14 ms per n512 matrix,
# measured on B200), so the rescue NEVER loops: one batched fp32 Jacobi
# call on the whole cohort, then ONE batched fp64 Jacobi call on any
# rows the fp32 pass declined (non-converged / non-finite / unsorted).
# torch.linalg.eigh remains only as a terminal per-call safety for rows
# the fp64 pass also declines (counter-instrumented; expected never).
work = data.clone(memory_format=torch.contiguous_format)
storage, values, info = module.syevj_batched(work)
vectors = storage.transpose(-2, -1)
good = _rescue_ok(info, vectors, values)
if bool(good.all().item()):
return vectors, values
rows = torch.nonzero(~good, as_tuple=False).flatten()
work64 = data[rows].to(torch.float64).contiguous()
storage64, values64, info64 = module.syevj_batched_double(work64)
row_vectors = storage64.transpose(-2, -1).to(torch.float32)
row_values = values64.to(torch.float32)
good64 = _rescue_ok(info64, row_vectors, row_values)
vectors = vectors.contiguous()
vectors[rows] = row_vectors
values[rows] = row_values
if not bool(good64.all().item()):
last = rows[~good64]
last_vectors, last_values = _solve(data[last])
vectors[last] = last_vectors
values[last] = last_values
return vectors, values
_MIDN_N = 176
_MIDN_STATS = 16
def _midn_n176(data: torch.Tensor) -> output_t:
# Three-kernel direct route for the n176 rows: reduction front end,
# per-root refinement, per-vector solve, then one batched GEMM.
module = _chain_extension()
work = data.contiguous()
batch = work.shape[0]
q = torch.empty((batch, _MIDN_N, _MIDN_N), dtype=torch.float32,
device=work.device)
v = torch.empty_like(q)
lam = torch.empty((batch, _MIDN_N), dtype=torch.float32,
device=work.device)
info = torch.empty((batch,), dtype=torch.int32, device=work.device)
stats = torch.zeros((batch, _MIDN_STATS), dtype=torch.int32,
device=work.device)
module.midn_eigh_pipeline_out(work, q, v, lam, info, stats, 0)
x = torch.bmm(q, v)
# Residual insurance at the dense level (one small batched GEMM). The
# lane's own correctness margins were 30-43x, but unflagged accuracy
# tails on real generator families cost this package a hosted round
# once already; ~0.1 ms buys the same fail-closed guarantee here.
resid = torch.bmm(work, x) - x * lam[:, None, :]
resid = resid.abs().sum(dim=1).amax(dim=1)
scale = work.abs().sum(dim=1).amax(dim=1).clamp_min(
torch.finfo(torch.float32).tiny)
ratio = resid / (200.0 * _MIDN_N * _EPS32 * scale)
bad = (
info.ne(0)
| ~torch.isfinite(x).all(dim=(-2, -1))
| ~torch.isfinite(lam).all(dim=-1)
| ~(lam[:, 1:] >= lam[:, :-1]).all(dim=-1)
| (ratio > _RESIDUAL_NET)
)
if bool(bad.any().item()):
rows = torch.nonzero(bad, as_tuple=False).flatten()
row_vectors, row_values = _solve(work[rows])
x[rows] = row_vectors
lam[rows] = row_values
return x, lam
def _orth_net_over(x, n=_CHAIN_N):
gram = torch.bmm(x.transpose(1, 2), x)
gram.diagonal(dim1=1, dim2=2).sub_(1.0)
ratio = gram.abs().sum(dim=1).amax(dim=1) / (100.0 * n * _EPS32)
return ratio > _ORTH_NET
def _chain_scale_ok(data: torch.Tensor) -> bool:
amax = data.abs().amax(dim=(-2, -1))
ok = (torch.isfinite(amax)
& (amax < _CHAIN_SCALE_HI)
& ((amax == 0.0) | (amax > _CHAIN_SCALE_LO)))
return bool(ok.all().item())
_INVOLUTION_TOL = 0.1
def _chain_dispatch(work: torch.Tensor) -> output_t:
if _N512_FAST:
out = _chain_pipeline(work, fast=True)
if out is not None:
return out
# The fast arm declined this batch; recompute everything strict.
return _chain_pipeline(work, fast=False)
def _chain_n512(data: torch.Tensor) -> output_t:
work = data.contiguous()
# Pre-detect involution matrices (A^2 ~= I => clustered +-1 spectrum). The
# D&C tree's secular solve brackets fail on near-degenerate spectra and route to
# the slow Jacobi fallback (~226ms); sending these straight to the fast cuSOLVER
# XsyevBatched (D&C) is far cheaper. This is a STRUCTURAL property of the input
# (the matrix squares to the identity), NOT a seed or value fingerprint.
# Cheap O(n^2) necessary-condition pre-gate before the O(n^3) A^2 GEMM: an
# involution (A^2 = I, A symmetric) has unit-norm rows -- the diagonal of
# A^2 is all ones, i.e. sum_j A_ij^2 = 1 for every row i. That row
# sum-of-squares is a reduction, not a GEMM. Non-involution spectra (mixed)
# fail it and skip the expensive A^2 entirely, so the detector costs ~0 on
# rows it cannot help (fixes the +13ms mixed regression from f1). The
# diagonal check is a strict subset of the full A^2~=I test, so it never
# rejects a real involution -- the clustered fast-path is preserved.
row_sq = (work * work).sum(dim=-1)
inv_cand = (row_sq - 1.0).abs().amax(dim=-1) < _INVOLUTION_TOL
if not bool(inv_cand.any().item()):
return _chain_dispatch(work)
cand_idx = torch.nonzero(inv_cand, as_tuple=False).flatten()
work_c = work[cand_idx].contiguous()
a2 = torch.bmm(work_c, work_c)
eye = torch.eye(work.shape[-1], device=work.device,
dtype=work.dtype).unsqueeze(0)
is_inv_c = (a2 - eye).abs().amax(dim=-1).amax(dim=-1) < _INVOLUTION_TOL
is_involution = torch.zeros((work.shape[0],), dtype=torch.bool,
device=work.device)
is_involution[cand_idx] = is_inv_c
n_inv = int(is_involution.sum().item())
# Fraction gate: only split when involutions DOMINATE the batch (clustered
# ~100%), where the chain's secular solve falls to the ~226ms Jacobi and the
# cuSOLVER D&C is far cheaper. When they are a sparse minority (the mixed row
# is ~6%), the chain already handles them at full speed, so splitting off a
# partial batch + a separate cuSOLVER call + scatter/gather only adds cost --
# this is the actual +13ms mixed regression from f1 (the detector bmm itself
# is only ~1ms). Fall through to the full-batch chain unless the gate clears.
if n_inv * 2 < work.shape[0]:
return _chain_dispatch(work)
n = work.shape[-1]
out_vecs = work.new_empty((work.shape[0], n, n))
out_vals = work.new_empty((work.shape[0], n))
inv_idx = torch.nonzero(is_involution, as_tuple=False).flatten()
gvecs, gvals = _gen_eigh_ext().xsyev_batched(
work[inv_idx].contiguous().clone())
out_vecs[inv_idx] = gvecs.transpose(-1, -2).contiguous()
out_vals[inv_idx] = gvals
non_idx = torch.nonzero(~is_involution, as_tuple=False).flatten()
if int(non_idx.numel()):
cv, cvals = _chain_dispatch(work[non_idx].contiguous())
out_vecs[non_idx] = cv
out_vals[non_idx] = cvals
return out_vecs, out_vals
def _chain_pipeline(work, fast):
module = _chain_extension()
batch = work.shape[0]
if fast:
with _FastGemm():
band, reflectors = _reduce_to_band8_fast(module, work)
else:
band, reflectors = _reduce_to_band8(module, work)
d, e, hh, stage2_info = _band8_to_tridiagonal(module, band)
values, x, bad = _run_tree(module, d, e)
if values is None:
# Early tree breaker. Fast arm: hand the batch back for a strict
# recompute; strict arm: whole-batch exact path as before.
return None if fast else _solve(work)
x = x.contiguous()
values = values.contiguous()
bad = bad | stage2_info.ne(0)
bad = bad | _residual_net_over(d, e, x, values)
bad = bad | _orth_net_over(x)
if not fast and (int(bad.sum().item()) * _BREAKER_DENOMINATOR
> batch * _BREAKER_NUMERATOR):
# Circuit breaker BEFORE the two back-substitution stages: past
# this cohort fraction per-row rescue is the wrong tool.
return _solve(work)
apply_info = torch.zeros((batch,), dtype=torch.int32, device=work.device)
module.backtransform_kd8_inplace(x, hh, apply_info, "t512s")
if _Q1_MODE == "triton-fused" and _TRITON_OK:
_q1_triton_apply_(reflectors, x)
elif _Q1_MODE in ("triton-fused", "tagg-2gemm"):
_q1_tagg_apply_(reflectors, x)
else:
_apply_q1_(reflectors, x, fast and _Q1_FAST)
bad = bad | apply_info.ne(0)
finite = (torch.isfinite(x).all(dim=(-2, -1))
& torch.isfinite(values).all(dim=-1))
bad = bad | ~finite
if fast:
# Dense-level eigen residual net: the intermediate nets cannot see
# the fast-GEMM reduction error, so the fast arm is additionally
# gated per row against the actual checker quantity.
resid = torch.bmm(work, x) - x * values[:, None, :]
resid = resid.abs().sum(dim=1).amax(dim=1)
scale = work.abs().sum(dim=1).amax(dim=1).clamp_min(
torch.finfo(torch.float32).tiny)
bad = bad | (resid / (200.0 * _CHAIN_N * _EPS32 * scale)
> _DENSE_NET)
if (int(bad.sum().item()) * _FAST_COHORT_DEN
> batch * _FAST_COHORT_NUM):
return None
if bool(bad.any().item()):
rows = torch.nonzero(bad, as_tuple=False).flatten()
fallback_vectors, fallback_values = _chain_exact_fallback(
module, work[rows])
x[rows] = fallback_vectors
values[rows] = fallback_values
return x, values
_N1024 = 1024
_N1024_SLOTS = sum((_N1024 - s - 1 + _CHAIN_KD - 1) // _CHAIN_KD
for s in range(1, _N1024 - 1))
_N1024_TREE_WIDTHS = (64, 128, 256, 512, 1024)
# The binding requirement is the m1s8 reduction kernel's dynamic shared
# contract (the shipping ring-buffer back-transform config needs only 8 KB
# static shared); devices below the opt-in keep the plain torch route.
_N1024_KD8_SHARED = 72864
_K0K6_STATE = {}
@lru_cache(maxsize=8)
def _n1024_supported(device_index: int) -> bool:
limit = torch.cuda.get_device_properties(
device_index).shared_memory_per_block_optin
return bool(limit >= _N1024_KD8_SHARED)
def _reduce_to_band8_n1024(module, original):
batch, n = original.shape[0], _N1024
kd = _CHAIN_KD
matrix = original.clone(memory_format=torch.contiguous_format)
band = torch.zeros((batch, kd + 1, n), dtype=original.dtype,
device=original.device)
reflectors = []
for panel_start in range(0, n - kd, kd):
trailing_start = panel_start + kd
panel = matrix[
:, trailing_start:, panel_start:panel_start + kd].contiguous()
qr = torch.empty_like(panel)
tau = torch.empty((batch, kd), dtype=panel.dtype, device=panel.device)
v = torch.empty_like(panel)
t = torch.empty((batch, kd, kd), dtype=panel.dtype,
device=panel.device)
module.panel1024_qr_w8_vt_out(panel, qr, tau, v, t)
matrix[:, trailing_start:trailing_start + kd,
panel_start:panel_start + kd].copy_(qr[:, :kd])
gather = matrix.as_strided(
(batch, kd + 1, kd), (n * n, n, n + 1),
matrix.storage_offset() + panel_start * (n + 1))
band[:, :, panel_start:panel_start + kd].copy_(gather)
vt = torch.bmm(v, t)
w = torch.empty_like(vt)
module.skinny_w_out(matrix, vt, w, trailing_start)
small = torch.bmm(vt.transpose(1, 2), w)
w = w - 0.5 * torch.bmm(v, small)
left = torch.cat((v, w), dim=2).contiguous()
right = torch.cat(
(w.transpose(1, 2), v.transpose(1, 2)), dim=1).contiguous()
module.update_rank16_inplace(matrix, left, right, trailing_start)
reflectors.append((trailing_start, v, t))
for column in range(n - kd, n):
length = min(kd, n - 1 - column) + 1
band[:, :length, column].copy_(
matrix[:, column:column + length, column])
return band, reflectors
def _band8_to_tridiagonal_n1024(module, band):
batch = band.shape[0]
n = _N1024
d = torch.empty((batch, n), dtype=torch.float32, device=band.device)
e = torch.empty((batch, n - 1), dtype=torch.float32, device=band.device)
hh = torch.empty((batch, _N1024_SLOTS, _CHAIN_KD), dtype=torch.float32,
device=band.device)
info = torch.zeros((batch,), dtype=torch.int32, device=band.device)
module.band8_to_tridiagonal_kd8_n1024_out(band, d, e, hh, info, "m1s8")
return d, e, hh, info
def _k0k6_solve(module, poles, weights, rho, counts, out_values, secular,
info, iterations):
# Adapter around the batch-general four-tile merge (same call shape as
# the certified control solvers; nonzero info == node failed closed).
batch, device = poles.shape[0], poles.device
key = (batch, device.index)
state = _K0K6_STATE.get(key)
if state is None:
f32 = dict(dtype=torch.float32, device=device)
f64 = dict(dtype=torch.float64, device=device)
i32 = dict(dtype=torch.int32, device=device)
w, tiles, tw = 1024, 4, 256
state = {
"roots64": torch.zeros(batch, w, **f64),
"root_status": torch.zeros(batch, w, **i32),
"root_iterations": torch.zeros(batch, w, **i32),
"merge_mode": torch.zeros(batch, **i32),
"repair_ids": torch.zeros(batch, tiles, tw, **i32),
"repair_count": torch.zeros(batch, tiles, **i32),
"deflated": torch.zeros(batch, w, **i32),
"updated32": torch.zeros(batch, w, **f32),
"updated64": torch.zeros(batch, w, **f64),
"recon_status": torch.zeros(batch, **i32),
"escalated": torch.zeros(batch, **i32),
"escalation_ids": torch.zeros(batch, tiles, tw, **i32),
"escalation_count": torch.zeros(batch, tiles, **i32),
"route_summary": torch.zeros(batch, **i32),
}
_K0K6_STATE[key] = state
module.k0k6_run_polished_(
poles, weights, rho, counts, out_values,
state["roots64"], state["root_status"], state["root_iterations"],
state["merge_mode"], state["repair_ids"], state["repair_count"],
state["deflated"], info, iterations,
state["updated32"], state["updated64"], state["recon_status"],
state["escalated"], state["escalation_ids"],
state["escalation_count"], secular, state["route_summary"])
def _run_tree_n1024(module, d, e):
batch = d.shape[0]
device = d.device
n = _N1024
leaf = 32
leaves = n // leaf
adjusted = d.clone()
cuts = torch.arange(leaf - 1, n - 1, leaf, device=device)
coupling = e[:, cuts].abs()
adjusted[:, cuts] -= coupling
adjusted[:, cuts + 1] -= coupling
leaf_d = adjusted.view(batch, leaves, leaf).reshape(-1, leaf).contiguous()
local = torch.arange(leaf - 1, device=device)
leaf_starts = torch.arange(leaves, device=device)[:, None] * leaf
leaf_e = e[:, (leaf_starts + local).reshape(-1)].reshape(
-1, leaf - 1).contiguous()
leaf_qt = torch.empty((batch * leaves, leaf, leaf), dtype=torch.float32,
device=device)
leaf_info = torch.zeros((batch * leaves,), dtype=torch.int32,
device=device)
module.leaf32_run(leaf_d, leaf_e, leaf_qt, leaf_info)
values = leaf_d.view(batch, leaves, leaf)
vectors = leaf_qt.view(batch, leaves, leaf, leaf).transpose(-1, -2)
bad_matrix = leaf_info.view(batch, leaves).ne(0).any(dim=1)
solvers = {
64: module.merge64_solve,
128: module.merge128_solve,
256: module.merge256_solve,
512: module.merge512_solve,
}
for width in _N1024_TREE_WIDTHS:
poles, weights, rho, basis, bad_matrix = _tree_prepare_level(
values, vectors, e, width, bad_matrix)
if (int(bad_matrix.sum().item()) * _BREAKER_DENOMINATOR
> batch * _BREAKER_NUMERATOR):
return None, None, bad_matrix
parents = poles.shape[1]
flat = batch * parents
poles_flat = poles.reshape(flat, width)
weights_flat = weights.reshape(flat, width)
rho_flat = rho.reshape(flat)
counts = torch.full((flat,), width, dtype=torch.int32, device=device)
out_values = torch.empty_like(poles_flat)
secular = torch.empty((flat, width, width), dtype=torch.float32,
device=device)
info = torch.zeros((flat,), dtype=torch.int32, device=device)
iterations = torch.zeros_like(info)
if width == 1024:
_k0k6_solve(module, poles_flat, weights_flat, rho_flat, counts,
out_values, secular, info, iterations)
else:
solvers[width](poles_flat, weights_flat, rho_flat, counts,
out_values, secular, info, iterations)
bad_matrix = bad_matrix | info.view(batch, parents).ne(0).any(dim=1)
merged = torch.empty_like(secular)
torch.bmm(basis.reshape(flat, width, width), secular, out=merged)
values = out_values.view(batch, parents, width)
vectors = merged.view(batch, parents, width, width)
return values[:, 0], vectors[:, 0], bad_matrix
# One-line route switch, OFF per hosted benchmark 870908: the n1024
# quartet regressed on the real board (mixed 131.48 / nearrank 131.73 /
# geometric 196.01 vs torch ~88-92; dense 90.93 a wash) — the breaker path
# loses hosted and even the clean-flow rows do not pay. All four n1024
# rows take the plain torch route; the machinery stays packaged for a
# future dense-only routing round.
_N1024_ROUTE = False
def _chain_n1024(data: torch.Tensor) -> output_t:
# Strict pipeline only for this row (the fast reduction arm is NOT
# integrated here: it priced at a loss on this row's net recomputes).
module = _chain_extension()
work = data.contiguous()
batch = work.shape[0]
band, reflectors = _reduce_to_band8_n1024(module, work)
d, e, hh, stage2_info = _band8_to_tridiagonal_n1024(module, band)
values, x, bad = _run_tree_n1024(module, d, e)
if values is None:
return _solve(work)
x = x.contiguous()
values = values.contiguous()
bad = bad | stage2_info.ne(0)
bad = bad | _residual_net_over(d, e, x, values, _N1024)
bad = bad | _orth_net_over(x, _N1024)
if (int(bad.sum().item()) * _BREAKER_DENOMINATOR
> batch * _BREAKER_NUMERATOR):
return _solve(work)
apply_info = torch.zeros((batch,), dtype=torch.int32, device=work.device)
module.q2occ_backtransform_inplace(x, hh, apply_info, "g16b2p")
_q1_tagg_apply_(reflectors, x, _N1024)
bad = bad | apply_info.ne(0)
finite = (torch.isfinite(x).all(dim=(-2, -1))
& torch.isfinite(values).all(dim=-1))
bad = bad | ~finite
if bool(bad.any().item()):
rows = torch.nonzero(bad, as_tuple=False).flatten()
fallback_vectors, fallback_values = _chain_exact_fallback(
module, work[rows])
x[rows] = fallback_vectors
values[rows] = fallback_values
return x, values
def custom_kernel(data: input_t) -> output_t:
if (
data.dim() == 3
and data.shape[-2] == data.shape[-1]
and data.shape[-1] in _DIRECT_N
):
return _direct_n32(data)
if (
data.dim() == 3
and data.shape[-2] == data.shape[-1] == _MIDN_N
and data.dtype == torch.float32
and data.is_cuda
and _chain_scale_ok(data)
):
return _midn_n176(data)
if (
data.dim() == 3
and data.shape[-2] == data.shape[-1] == _CHAIN_N
and data.dtype == torch.float32
and data.is_cuda
and _chain_scale_ok(data)
):
return _chain_n512(data)
if (
data.dim() == 3
and data.shape[-2] == data.shape[-1] == _N1024
and data.dtype == torch.float32
and data.is_cuda
and _N1024_ROUTE
and _n1024_supported(data.get_device())
and _chain_scale_ok(data)
):
return _chain_n1024(data)
return _solve(data)
scrolls · 15508 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