submission 877439
drillyb · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 4630 lines, June 9 Researcher Reciprocity License v1.0.
new_sub.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-eigh-877439?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:c842a95ee0214131c4c418b21de187f5aa15f44b0bf5de4f3bb0701f9484a776
license declaredunknown
license concludedunknown
authorsdrillyb
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
cluster
class ClusterShape_ = Shape<_1, _1, _1>,fused-epilogue
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<persistent-kernel
void persistent_panel_small_kernel(shared-memory
double x, double y, double* smem_x, double* smem_y) {Kernel source
new_sub.py4630 lines
#!POPCORN leaderboard eigh
#!POPCORN gpu B200
from __future__ import annotations
import os
from pathlib import Path
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
CUDA_SOURCE = r"""
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include <c10/cuda/CUDAException.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <cooperative_groups.h>
#include <math_constants.h>
#include <cutlass/cutlass.h>
#include <cutlass/numeric_types.h>
#include <cute/tensor.hpp>
#include <cutlass/gemm/dispatch_policy.hpp>
#include <cutlass/gemm/collective/collective_builder.hpp>
#include <cutlass/epilogue/collective/collective_builder.hpp>
#include <cutlass/gemm/device/gemm_universal_adapter.h>
#include <cutlass/gemm/kernel/gemm_universal.hpp>
#include <algorithm>
#include <cmath>
#include <cfloat>
#include <cstring>
#include <cstdint>
#include <limits>
#include <tuple>
#include <type_traits>
#include <vector>
namespace b200_eigh {
using namespace cute;
namespace cg = cooperative_groups;
constexpr int kPanel = 32;
constexpr int kLeaf = 32;
constexpr int kThreads = 256;
constexpr int kLeafJacobiThreads = 256;
constexpr int kDenseJacobiThreads = 512;
constexpr int kPanel512Threads = 512;
constexpr int kPanel512Chunk = 16;
constexpr int kPanel512PrefixCols = 16;
constexpr unsigned kFullMask = 0xffffffffu;
__host__ __device__ static inline int ceil_div(int x, int y) { return (x + y - 1) / y; }
__host__ __device__ static inline int round_up(int x, int y) { return ceil_div(x, y) * y; }
static inline int next_pow2(int x) {
int p = 1;
while (p < x) p <<= 1;
return p;
}
inline void cutlass_check(cutlass::Status status, const char* where) {
TORCH_CHECK(status == cutlass::Status::kSuccess,
where, ": CUTLASS status ", static_cast<int>(status));
}
template<class LayoutA, class LayoutB,
class ElementA_ = float, class ElementB_ = float,
int AlignmentAB = 4, int AlignmentC_ = 4,
class ClusterShape_ = Shape<_1, _1, _1>,
class TileShape_ = Shape<_128, _128, _32>>
struct Sm100Gemm {
using ElementA = ElementA_;
using ElementB = ElementB_;
using ElementC = float;
using ElementAccumulator = float;
using ArchTag = cutlass::arch::Sm100;
using OperatorClass = cutlass::arch::OpClassTensorOp;
static constexpr int AlignmentA = AlignmentAB;
static constexpr int AlignmentB = AlignmentAB;
static constexpr int AlignmentC = AlignmentC_;
using TileShape = TileShape_;
using ClusterShape = ClusterShape_;
using LayoutC = cutlass::layout::ColumnMajor;
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
ArchTag, OperatorClass,
TileShape, ClusterShape,
cutlass::epilogue::collective::EpilogueTileAuto,
ElementAccumulator, ElementAccumulator,
ElementC, LayoutC, AlignmentC,
ElementC, LayoutC, AlignmentC,
cutlass::epilogue::collective::EpilogueScheduleAuto>::CollectiveOp;
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
ArchTag, OperatorClass,
ElementA, LayoutA, AlignmentA,
ElementB, LayoutB, AlignmentB,
ElementAccumulator,
TileShape, ClusterShape,
cutlass::gemm::collective::StageCountAutoCarveout<
static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
cutlass::gemm::collective::KernelScheduleAuto>::CollectiveOp;
using Kernel = cutlass::gemm::kernel::GemmUniversal<
Shape<int, int, int, int>, CollectiveMainloop, CollectiveEpilogue, void>;
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<Kernel>;
using StrideA = typename Kernel::StrideA;
using StrideB = typename Kernel::StrideB;
using StrideC = typename Kernel::StrideC;
using StrideD = typename Kernel::StrideD;
static StrideA stride_a(int64_t ld, int64_t batch_stride) {
if constexpr (std::is_same_v<LayoutA, cutlass::layout::ColumnMajor>) {
return StrideA{cute::Int<1>{}, ld, batch_stride};
} else {
return StrideA{ld, cute::Int<1>{}, batch_stride};
}
}
static StrideB stride_b(int64_t ld, int64_t batch_stride) {
if constexpr (std::is_same_v<LayoutB, cutlass::layout::ColumnMajor>) {
return StrideB{ld, cute::Int<1>{}, batch_stride};
} else {
return StrideB{cute::Int<1>{}, ld, batch_stride};
}
}
static StrideC stride_c(int64_t ld, int64_t batch_stride) {
return StrideC{cute::Int<1>{}, ld, batch_stride};
}
static StrideD stride_d(int64_t ld, int64_t batch_stride) {
return StrideD{cute::Int<1>{}, ld, batch_stride};
}
static typename Gemm::Arguments make_args(
int m, int n, int k, int batch,
const ElementA* A, int64_t lda, int64_t batch_a,
const ElementB* B, int64_t ldb, int64_t batch_b,
const float* C, int64_t ldc, int64_t batch_c,
float* D, int64_t ldd, int64_t batch_d,
float alpha, float beta) {
return typename Gemm::Arguments{
cutlass::gemm::GemmUniversalMode::kGemm,
{m, n, k, batch},
{A, stride_a(lda, batch_a), B, stride_b(ldb, batch_b)},
{{alpha, beta}, C, stride_c(ldc, batch_c),
D, stride_d(ldd, batch_d)}};
}
static size_t workspace_size(
int m, int n, int k, int batch,
const ElementA* A, int64_t lda, int64_t batch_a,
const ElementB* B, int64_t ldb, int64_t batch_b,
const float* C, int64_t ldc, int64_t batch_c,
float* D, int64_t ldd, int64_t batch_d) {
auto args = make_args(m, n, k, batch,
A, lda, batch_a, B, ldb, batch_b,
C, ldc, batch_c, D, ldd, batch_d,
1.0f, 0.0f);
return Gemm::get_workspace_size(args);
}
static bool can_run(
int m, int n, int k, int batch,
const ElementA* A, int64_t lda, int64_t batch_a,
const ElementB* B, int64_t ldb, int64_t batch_b,
const float* C, int64_t ldc, int64_t batch_c,
float* D, int64_t ldd, int64_t batch_d,
float alpha = 1.0f, float beta = 0.0f) {
auto args = make_args(m, n, k, batch,
A, lda, batch_a, B, ldb, batch_b,
C, ldc, batch_c, D, ldd, batch_d,
alpha, beta);
Gemm gemm;
return gemm.can_implement(args) == cutlass::Status::kSuccess;
}
static void run(
int m, int n, int k, int batch,
const ElementA* A, int64_t lda, int64_t batch_a,
const ElementB* B, int64_t ldb, int64_t batch_b,
const float* C, int64_t ldc, int64_t batch_c,
float* D, int64_t ldd, int64_t batch_d,
float alpha, float beta,
void* workspace, size_t workspace_bytes) {
auto args = make_args(m, n, k, batch,
A, lda, batch_a, B, ldb, batch_b,
C, ldc, batch_c, D, ldd, batch_d,
alpha, beta);
size_t need = Gemm::get_workspace_size(args);
TORCH_CHECK(need <= workspace_bytes,
"insufficient CUTLASS workspace: need ", need,
", have ", workspace_bytes);
Gemm gemm;
auto implement_status = gemm.can_implement(args);
TORCH_CHECK(implement_status == cutlass::Status::kSuccess,
"can_implement failed for m=", m, ", n=", n,
", k=", k, ", batch=", batch,
", alignment=", AlignmentAB, ": ",
cutlassGetStatusString(implement_status), " (",
static_cast<int>(implement_status), ")");
cutlass_check(gemm.initialize(args, workspace), "initialize");
cutlass_check(gemm.run(), "run");
}
};
using GemmCR = Sm100Gemm<cutlass::layout::ColumnMajor,
cutlass::layout::RowMajor>;
using GemmCC = Sm100Gemm<cutlass::layout::ColumnMajor,
cutlass::layout::ColumnMajor>;
using FloatGemmRC = Sm100Gemm<cutlass::layout::RowMajor,
cutlass::layout::ColumnMajor,
float, float, 4>;
using HalfGemmRC = Sm100Gemm<cutlass::layout::RowMajor,
cutlass::layout::ColumnMajor,
cutlass::half_t, cutlass::half_t, 8>;
using HalfGemmCC = Sm100Gemm<cutlass::layout::ColumnMajor,
cutlass::layout::ColumnMajor,
cutlass::half_t, cutlass::half_t, 8>;
using GemmCRTmaCluster2 = Sm100Gemm<
cutlass::layout::ColumnMajor, cutlass::layout::RowMajor,
float, float, 4, 4, Shape<_2, _1, _1>>;
using HalfGemmCCTmaCluster2 = Sm100Gemm<
cutlass::layout::ColumnMajor, cutlass::layout::ColumnMajor,
cutlass::half_t, cutlass::half_t, 8, 4, Shape<_2, _1, _1>>;
__device__ __forceinline__ double warp_sum_double(double x) {
for (int d = 16; d > 0; d >>= 1) x += __shfl_down_sync(kFullMask, x, d);
return x;
}
__device__ __forceinline__ float warp_sum_float(float x) {
for (int d = 16; d > 0; d >>= 1) x += __shfl_down_sync(kFullMask, x, d);
return x;
}
__device__ __forceinline__ double warp_max_double(double x) {
for (int d = 16; d > 0; d >>= 1)
x = fmax(x, __shfl_down_sync(kFullMask, x, d));
return x;
}
template<int Threads>
__device__ __forceinline__ double block_sum_double(double x, double* smem) {
constexpr int Warps = (Threads + 31) / 32;
int lane = threadIdx.x & 31;
int warp = threadIdx.x >> 5;
x = warp_sum_double(x);
if (lane == 0) smem[warp] = x;
__syncthreads();
double y = 0.0;
if (warp == 0) {
y = lane < Warps ? smem[lane] : 0.0;
y = warp_sum_double(y);
if (lane == 0) smem[0] = y;
}
__syncthreads();
return smem[0];
}
template<int Threads>
__device__ __forceinline__ double2 block_sum_double_pair(
double x, double y, double* smem_x, double* smem_y) {
constexpr int Warps = (Threads + 31) / 32;
int lane = threadIdx.x & 31;
int warp = threadIdx.x >> 5;
x = warp_sum_double(x);
y = warp_sum_double(y);
if (lane == 0) {
smem_x[warp] = x;
smem_y[warp] = y;
}
__syncthreads();
if (warp == 0) {
x = lane < Warps ? smem_x[lane] : 0.0;
y = lane < Warps ? smem_y[lane] : 0.0;
x = warp_sum_double(x);
y = warp_sum_double(y);
if (lane == 0) {
smem_x[0] = x;
smem_y[0] = y;
}
}
__syncthreads();
return make_double2(smem_x[0], smem_y[0]);
}
template<int Threads>
__device__ __forceinline__ double block_max_double(double x, double* smem) {
constexpr int Warps = (Threads + 31) / 32;
int lane = threadIdx.x & 31;
int warp = threadIdx.x >> 5;
x = warp_max_double(x);
if (lane == 0) smem[warp] = x;
__syncthreads();
double y = 0.0;
if (warp == 0) {
y = lane < Warps ? smem[lane] : 0.0;
y = warp_max_double(y);
if (lane == 0) smem[0] = y;
}
__syncthreads();
return smem[0];
}
__global__ void copy_symmetrize_to_colmajor_kernel(
const float* __restrict__ input,
float* __restrict__ A,
int n, int ld) {
int b = blockIdx.z;
int i = blockIdx.y * blockDim.y + threadIdx.y;
int j = blockIdx.x * blockDim.x + threadIdx.x;
if (i >= n || j >= n) return;
const float* Ib = input + static_cast<int64_t>(b) * n * n;
float* Ab = A + static_cast<int64_t>(b) * ld * n;
float x = 0.5f * (Ib[i * n + j] + Ib[j * n + i]);
Ab[i + static_cast<int64_t>(j) * ld] = x;
}
__global__ void copy_colmajor_q_to_rowmajor_kernel(
const float* __restrict__ Qin,
float* __restrict__ Qout,
int n, int ld) {
int b = blockIdx.z;
int i = blockIdx.y * blockDim.y + threadIdx.y;
int j = blockIdx.x * blockDim.x + threadIdx.x;
if (i >= n || j >= n) return;
const float* Qb = Qin + static_cast<int64_t>(b) * ld * n;
float* Ob = Qout + static_cast<int64_t>(b) * n * n;
Ob[i * n + j] = Qb[i + static_cast<int64_t>(j) * ld];
}
__global__ void normalize_matrix_pow2_kernel(
float* __restrict__ A,
double* __restrict__ scale_back,
int n, int ld) {
int b = blockIdx.x;
int tid = threadIdx.x;
float* Ab = A + static_cast<int64_t>(b) * ld * n;
__shared__ double red[8];
__shared__ double sh_forward;
double local_max = 0.0;
for (int64_t t = tid; t < static_cast<int64_t>(n) * n;
t += blockDim.x) {
int i = static_cast<int>(t % n);
int j = static_cast<int>(t / n);
local_max = fmax(local_max,
fabs(static_cast<double>(
Ab[i + static_cast<int64_t>(j) * ld])));
}
double mx = block_max_double<kThreads>(local_max, red);
if (tid == 0) {
double forward = 1.0;
double backward = 1.0;
if (mx > 0.0 && isfinite(mx)) {
int exponent = 0;
frexp(mx, &exponent);
forward = ldexp(1.0, -exponent);
backward = ldexp(1.0, exponent);
}
sh_forward = forward;
scale_back[b] = backward;
}
__syncthreads();
for (int64_t t = tid; t < static_cast<int64_t>(n) * n;
t += blockDim.x) {
int i = static_cast<int>(t % n);
int j = static_cast<int>(t / n);
int64_t idx = i + static_cast<int64_t>(j) * ld;
Ab[idx] = static_cast<float>(
static_cast<double>(Ab[idx]) * sh_forward);
}
}
__global__ void rescale_eigenvalues_kernel(
float* __restrict__ values,
const double* __restrict__ scale_back,
int n) {
int b = blockIdx.y;
int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i >= n) return;
values[static_cast<int64_t>(b) * n + i] =
static_cast<float>(
static_cast<double>(values[static_cast<int64_t>(b) * n + i]) *
scale_back[b]);
}
__global__ void symmetrize_submatrix_kernel(
float* __restrict__ A, int n, int ld, int begin) {
int b = blockIdx.z;
int i = begin + blockIdx.y * blockDim.y + threadIdx.y;
int j = begin + blockIdx.x * blockDim.x + threadIdx.x;
if (i >= n || j >= n || i >= j) return;
float* Ab = A + static_cast<int64_t>(b) * ld * n;
float s = 0.5f * (Ab[i + static_cast<int64_t>(j) * ld] +
Ab[j + static_cast<int64_t>(i) * ld]);
Ab[i + static_cast<int64_t>(j) * ld] = s;
Ab[j + static_cast<int64_t>(i) * ld] = s;
}
__global__ void trailing_rank2_update_kernel(
float* __restrict__ A,
const float* __restrict__ V,
const float* __restrict__ W,
int n, int ld, int panel_b, int begin, int bcols) {
int b = blockIdx.z;
int i = begin + blockIdx.y * blockDim.y + threadIdx.y;
int j = begin + blockIdx.x * blockDim.x + threadIdx.x;
if (i >= n || j >= n) return;
float* Ab = A + static_cast<int64_t>(b) * ld * n;
const float* Vb = V + static_cast<int64_t>(b) * ld * panel_b;
const float* Wb = W + static_cast<int64_t>(b) * ld * panel_b;
float update = 0.0f;
for (int s = 0; s < bcols; ++s) {
update += Vb[i + static_cast<int64_t>(s) * ld] *
Wb[j + static_cast<int64_t>(s) * ld];
update += Wb[i + static_cast<int64_t>(s) * ld] *
Vb[j + static_cast<int64_t>(s) * ld];
}
Ab[i + static_cast<int64_t>(j) * ld] -= update;
}
__global__ void pack_rank2_factors_kernel(
const float* __restrict__ V,
const float* __restrict__ W,
float* __restrict__ U2,
float* __restrict__ Z2,
int ld, int panel_b, int bcols) {
int b = blockIdx.y;
int64_t t = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
int64_t total = static_cast<int64_t>(ld) * (2 * bcols);
if (t >= total) return;
int r = static_cast<int>(t % ld);
int c = static_cast<int>(t / ld);
const float* Vb = V + static_cast<int64_t>(b) * ld * panel_b;
const float* Wb = W + static_cast<int64_t>(b) * ld * panel_b;
float* Ub = U2 + static_cast<int64_t>(b) * ld * (2 * panel_b);
float* Zb = Z2 + static_cast<int64_t>(b) * ld * (2 * panel_b);
if (c < bcols) {
Ub[r + static_cast<int64_t>(c) * ld] =
Vb[r + static_cast<int64_t>(c) * ld];
Zb[r + static_cast<int64_t>(c) * ld] =
Wb[r + static_cast<int64_t>(c) * ld];
} else {
int q = c - bcols;
Ub[r + static_cast<int64_t>(c) * ld] =
Wb[r + static_cast<int64_t>(q) * ld];
Zb[r + static_cast<int64_t>(c) * ld] =
Vb[r + static_cast<int64_t>(q) * ld];
}
}
__global__ void detect_diagonal_kernel(
const float* __restrict__ A, int n, int ld,
int* __restrict__ flags) {
int b = blockIdx.x;
const float* Ab = A + static_cast<int64_t>(b) * ld * n;
__shared__ int any;
if (threadIdx.x == 0) any = 0;
__syncthreads();
for (int64_t t = threadIdx.x; t < static_cast<int64_t>(n) * n;
t += blockDim.x) {
int i = static_cast<int>(t % n);
int j = static_cast<int>(t / n);
if (i != j && Ab[i + static_cast<int64_t>(j) * ld] != 0.0f)
atomicExch(&any, 1);
}
__syncthreads();
if (threadIdx.x == 0) flags[b] = any == 0 ? 1 : 0;
}
__global__ void fill_diagonal_solution_kernel(
const float* __restrict__ A, int n, int ld,
const int* __restrict__ flags,
float* __restrict__ evals,
float* __restrict__ Q) {
int b = blockIdx.z;
if (!flags[b]) return;
const float* Ab = A + static_cast<int64_t>(b) * ld * n;
float* Lb = evals + static_cast<int64_t>(b) * n;
float* Qb = Q + static_cast<int64_t>(b) * ld * n;
for (int t = blockIdx.x * blockDim.x + threadIdx.x;
t < n * n; t += blockDim.x * gridDim.x) {
int i = t % n;
int j = t / n;
Qb[i + static_cast<int64_t>(j) * ld] = (i == j) ? 1.0f : 0.0f;
}
for (int i = blockIdx.x * blockDim.x + threadIdx.x;
i < n; i += blockDim.x * gridDim.x) {
Lb[i] = Ab[i + static_cast<int64_t>(i) * ld];
}
}
__global__ __launch_bounds__(kThreads, 1)
void persistent_panel_small_kernel(
float* __restrict__ A,
float* __restrict__ T,
float* __restrict__ tau,
float* __restrict__ diag,
float* __restrict__ offdiag,
float* __restrict__ V,
float* __restrict__ W,
float* __restrict__ U2,
float* __restrict__ Z2,
int n, int ld, int panel_b, int num_panels,
int panel_id, int k, int bcols) {
int b = blockIdx.x;
int tid = threadIdx.x;
float* Ab = A + static_cast<int64_t>(b) * ld * n;
float* taub = tau + static_cast<int64_t>(b) * n;
float* db = diag + static_cast<int64_t>(b) * n;
float* eb = offdiag + static_cast<int64_t>(b) * n;
float* Vb = V + static_cast<int64_t>(b) * ld * panel_b;
float* Wb = W + static_cast<int64_t>(b) * ld * panel_b;
float* Ub = U2 + static_cast<int64_t>(b) * ld * (2 * panel_b);
float* Zb = Z2 + static_cast<int64_t>(b) * ld * (2 * panel_b);
float* Tb = T + (static_cast<int64_t>(b) * num_panels + panel_id) *
panel_b * panel_b;
extern __shared__ unsigned char dynamic_smem[];
float* sV = reinterpret_cast<float*>(dynamic_smem);
float* sW = sV + static_cast<int64_t>(ld) * panel_b;
float* sx = sW + static_cast<int64_t>(ld) * panel_b;
float* sy = sx + ld;
float* sT = sy + ld;
float* scv = sT + panel_b * panel_b;
float* scw = scv + panel_b;
double* red = reinterpret_cast<double*>(scw + panel_b);
double* red_pair = red + 8;
__shared__ double sh_max;
__shared__ double sh_tau;
__shared__ double sh_beta;
__shared__ double sh_inv;
__shared__ double sh_alpha;
int panel_storage = ld * panel_b;
for (int t = tid; t < 2 * panel_storage; t += blockDim.x)
sV[t] = 0.0f; // sV and sW are contiguous.
for (int t = tid; t < panel_b * panel_b; t += blockDim.x)
sT[t] = 0.0f;
__syncthreads();
for (int s = 0; s < bcols; ++s) {
int j = k + s;
int start = j + 1;
if (tid == 0) {
double d = Ab[j + static_cast<int64_t>(j) * ld];
for (int p = 0; p < s; ++p) {
d -= 2.0 * static_cast<double>(sV[j + static_cast<int64_t>(p) * ld]) *
static_cast<double>(sW[j + static_cast<int64_t>(p) * ld]);
}
db[j] = static_cast<float>(d);
}
for (int r = start + tid; r < n; r += blockDim.x) {
double value = Ab[r + static_cast<int64_t>(j) * ld];
for (int p = 0; p < s; ++p) {
value -= static_cast<double>(sV[r + static_cast<int64_t>(p) * ld]) *
sW[j + static_cast<int64_t>(p) * ld];
value -= static_cast<double>(sW[r + static_cast<int64_t>(p) * ld]) *
sV[j + static_cast<int64_t>(p) * ld];
}
sx[r] = static_cast<float>(value);
}
__syncthreads();
double local_max = 0.0;
for (int r = start + tid; r < n; r += blockDim.x)
local_max = fmax(local_max, fabs(static_cast<double>(sx[r])));
double mx = block_max_double<kThreads>(local_max, red);
if (tid == 0) sh_max = mx;
__syncthreads();
double total = 0.0;
double tail = 0.0;
if (sh_max > 0.0) {
for (int r = start + tid; r < n; r += blockDim.x) {
double q = static_cast<double>(sx[r]) / sh_max;
total += q * q;
if (r > start) tail += q * q;
}
}
double2 norm_sums = block_sum_double_pair<kThreads>(
total, tail, red, red_pair);
total = norm_sums.x;
tail = norm_sums.y;
if (tid == 0) {
double x0 = sx[start];
if (sh_max == 0.0 || tail == 0.0) {
sh_beta = x0;
sh_tau = 0.0;
sh_inv = 0.0;
} else {
double norm = sh_max * sqrt(total);
double beta = -copysign(norm, x0);
double denom = x0 - beta;
sh_beta = beta;
sh_tau = (beta - x0) / beta;
sh_inv = 1.0 / denom;
}
taub[j] = static_cast<float>(sh_tau);
eb[j] = static_cast<float>(sh_beta);
Ab[start + static_cast<int64_t>(j) * ld] =
static_cast<float>(sh_beta);
}
__syncthreads();
for (int r = tid; r < ld; r += blockDim.x) {
float vv = 0.0f;
if (r == start) {
vv = 1.0f;
} else if (r > start && r < n && sh_tau != 0.0) {
vv = static_cast<float>(static_cast<double>(sx[r]) * sh_inv);
}
sV[r + static_cast<int64_t>(s) * ld] = vv;
if (r > start && r < n)
Ab[r + static_cast<int64_t>(j) * ld] = vv;
}
__syncthreads();
int lane = tid & 31;
int warp = tid >> 5;
for (int p = warp; p < s; p += 8) {
double av = 0.0;
double aw = 0.0;
for (int r = start + lane; r < n; r += 32) {
double vr = sV[r + static_cast<int64_t>(s) * ld];
av += static_cast<double>(sV[r + static_cast<int64_t>(p) * ld]) * vr;
aw += static_cast<double>(sW[r + static_cast<int64_t>(p) * ld]) * vr;
}
av = warp_sum_double(av);
aw = warp_sum_double(aw);
if (lane == 0) {
scv[p] = static_cast<float>(av);
scw[p] = static_cast<float>(aw);
}
}
__syncthreads();
for (int r = start + tid; r < n; r += blockDim.x) {
double acc = 0.0;
for (int c = start; c < n; ++c) {
acc += static_cast<double>(Ab[r + static_cast<int64_t>(c) * ld]) *
sV[c + static_cast<int64_t>(s) * ld];
}
double value = static_cast<double>(static_cast<float>(acc));
for (int p = 0; p < s; ++p) {
value -= static_cast<double>(sV[r + static_cast<int64_t>(p) * ld]) *
scw[p];
value -= static_cast<double>(sW[r + static_cast<int64_t>(p) * ld]) *
scv[p];
}
sy[r] = static_cast<float>(value);
}
__syncthreads();
double dot = 0.0;
for (int r = start + tid; r < n; r += blockDim.x) {
dot += static_cast<double>(sV[r + static_cast<int64_t>(s) * ld]) *
sy[r];
}
dot = block_sum_double<kThreads>(dot, red);
if (tid == 0) {
sh_alpha = -0.5 * sh_tau * sh_tau * dot;
float tf = static_cast<float>(sh_tau);
if (tf == 0.0f) {
for (int r = 0; r <= s; ++r)
sT[r + static_cast<int64_t>(s) * panel_b] = 0.0f;
} else {
float tmp[kPanel];
for (int i = 0; i < s; ++i) tmp[i] = -tf * scv[i];
for (int r = 0; r < s; ++r) {
double acc = 0.0;
for (int q = r; q < s; ++q) {
acc += static_cast<double>(
sT[r + static_cast<int64_t>(q) * panel_b]) * tmp[q];
}
sT[r + static_cast<int64_t>(s) * panel_b] =
static_cast<float>(acc);
}
sT[s + static_cast<int64_t>(s) * panel_b] = tf;
}
}
__syncthreads();
for (int r = tid; r < ld; r += blockDim.x) {
float w = 0.0f;
if (r >= start && r < n) {
w = static_cast<float>(
sh_tau * static_cast<double>(sy[r]) +
sh_alpha * sV[r + static_cast<int64_t>(s) * ld]);
}
sW[r + static_cast<int64_t>(s) * ld] = w;
}
__syncthreads();
}
int r0 = k + bcols;
int m = n - r0;
if (r0 == n - 1 && tid == 0) {
int last = n - 1;
double d = Ab[last + static_cast<int64_t>(last) * ld];
for (int p = 0; p < bcols; ++p) {
d -= 2.0 * static_cast<double>(sV[last + static_cast<int64_t>(p) * ld]) *
static_cast<double>(sW[last + static_cast<int64_t>(p) * ld]);
}
db[last] = static_cast<float>(d);
eb[last] = 0.0f;
}
for (int t = tid; t < panel_b * panel_b; t += blockDim.x)
Tb[t] = sT[t];
int factor_elems = ld * bcols;
if (r0 < n - 1 && (m & 3) == 0) {
for (int t = tid; t < factor_elems; t += blockDim.x) {
int r = t % ld;
int p = t / ld;
float v = sV[r + static_cast<int64_t>(p) * ld];
float w = sW[r + static_cast<int64_t>(p) * ld];
Ub[r + static_cast<int64_t>(p) * ld] = v;
Zb[r + static_cast<int64_t>(p) * ld] = w;
Ub[r + static_cast<int64_t>(bcols + p) * ld] = w;
Zb[r + static_cast<int64_t>(bcols + p) * ld] = v;
}
} else if (r0 < n - 1) {
for (int t = tid; t < factor_elems; t += blockDim.x) {
int r = t % ld;
int p = t / ld;
Vb[r + static_cast<int64_t>(p) * ld] =
sV[r + static_cast<int64_t>(p) * ld];
Wb[r + static_cast<int64_t>(p) * ld] =
sW[r + static_cast<int64_t>(p) * ld];
}
}
}
template<bool CachePrefix16>
__global__ __launch_bounds__(kPanel512Threads)
void panel_chunk8_512_kernel(
float* __restrict__ A,
float* __restrict__ T,
float* __restrict__ tau,
float* __restrict__ diag,
float* __restrict__ offdiag,
float* __restrict__ V,
float* __restrict__ W,
int n, int ld, int panel_b, int num_panels,
int panel_id, int k, int s_begin, int chunk_cols) {
int b = blockIdx.x;
int tid = threadIdx.x;
int r = tid;
int lane = tid & 31;
int warp = tid >> 5;
float* Ab = A + static_cast<int64_t>(b) * ld * n;
float* taub = tau + static_cast<int64_t>(b) * n;
float* db = diag + static_cast<int64_t>(b) * n;
float* eb = offdiag + static_cast<int64_t>(b) * n;
float* Vb = V + static_cast<int64_t>(b) * ld * panel_b;
float* Wb = W + static_cast<int64_t>(b) * ld * panel_b;
float* Tb = T + (static_cast<int64_t>(b) * num_panels + panel_id) *
panel_b * panel_b;
extern __shared__ unsigned char panel512_prefix_raw[];
float* sVprefix = reinterpret_cast<float*>(panel512_prefix_raw);
float* sWprefix = sVprefix + 512 * kPanel512PrefixCols;
__shared__ float sx[512];
__shared__ float sy[512];
__shared__ float scv[kPanel];
__shared__ float scw[kPanel];
__shared__ float sjv[kPanel];
__shared__ float sjw[kPanel];
__shared__ float sT[kPanel * kPanel];
__shared__ double red[16];
__shared__ double red_pair[16];
__shared__ double sh_max;
__shared__ double sh_tau;
__shared__ double sh_beta;
__shared__ double sh_inv;
__shared__ double sh_alpha;
for (int idx = tid; idx < panel_b * panel_b; idx += blockDim.x) {
int col = idx / panel_b;
sT[idx] = (col < s_begin) ? Tb[idx] : 0.0f;
}
if constexpr (CachePrefix16) {
if (r >= k + 1 && r < n) {
#pragma unroll
for (int p = 0; p < kPanel512PrefixCols; ++p) {
sVprefix[r + p * 512] =
Vb[r + static_cast<int64_t>(p) * ld];
sWprefix[r + p * 512] =
Wb[r + static_cast<int64_t>(p) * ld];
}
}
}
__syncthreads();
for (int local_s = 0; local_s < chunk_cols; ++local_s) {
int s = s_begin + local_s;
int j = k + s;
int start = j + 1;
if (tid < s) {
if constexpr (CachePrefix16) {
if (tid < kPanel512PrefixCols) {
sjv[tid] = sVprefix[j + tid * 512];
sjw[tid] = sWprefix[j + tid * 512];
} else {
sjv[tid] = Vb[j + static_cast<int64_t>(tid) * ld];
sjw[tid] = Wb[j + static_cast<int64_t>(tid) * ld];
}
} else {
sjv[tid] = Vb[j + static_cast<int64_t>(tid) * ld];
sjw[tid] = Wb[j + static_cast<int64_t>(tid) * ld];
}
}
__syncthreads();
if (tid == 0) {
double d = Ab[j + static_cast<int64_t>(j) * ld];
for (int p = 0; p < s; ++p) {
d -= 2.0 * static_cast<double>(sjv[p]) *
static_cast<double>(sjw[p]);
}
db[j] = static_cast<float>(d);
}
if (r >= start && r < n) {
double value = Ab[r + static_cast<int64_t>(j) * ld];
if constexpr (CachePrefix16) {
#pragma unroll
for (int p = 0; p < kPanel512PrefixCols; ++p) {
value -= static_cast<double>(sVprefix[r + p * 512]) * sjw[p];
value -= static_cast<double>(sWprefix[r + p * 512]) * sjv[p];
}
for (int p = kPanel512PrefixCols; p < s; ++p) {
value -= static_cast<double>(
Vb[r + static_cast<int64_t>(p) * ld]) * sjw[p];
value -= static_cast<double>(
Wb[r + static_cast<int64_t>(p) * ld]) * sjv[p];
}
} else {
for (int p = 0; p < s; ++p) {
value -= static_cast<double>(
Vb[r + static_cast<int64_t>(p) * ld]) * sjw[p];
value -= static_cast<double>(
Wb[r + static_cast<int64_t>(p) * ld]) * sjv[p];
}
}
sx[r] = static_cast<float>(value);
}
__syncthreads();
double local_max = (r >= start && r < n)
? fabs(static_cast<double>(sx[r])) : 0.0;
double mx = block_max_double<kPanel512Threads>(local_max, red);
if (tid == 0) sh_max = mx;
__syncthreads();
double q2 = 0.0;
double tail2 = 0.0;
if (r >= start && r < n && sh_max > 0.0) {
double q = static_cast<double>(sx[r]) / sh_max;
q2 = q * q;
if (r > start) tail2 = q2;
}
double2 norm_sums = block_sum_double_pair<kPanel512Threads>(
q2, tail2, red, red_pair);
double total = norm_sums.x;
double tail = norm_sums.y;
if (tid == 0) {
double x0 = sx[start];
if (sh_max == 0.0 || tail == 0.0) {
sh_beta = x0;
sh_tau = 0.0;
sh_inv = 0.0;
} else {
double norm = sh_max * sqrt(total);
double beta = -copysign(norm, x0);
sh_beta = beta;
sh_tau = (beta - x0) / beta;
sh_inv = 1.0 / (x0 - beta);
}
taub[j] = static_cast<float>(sh_tau);
eb[j] = static_cast<float>(sh_beta);
Ab[start + static_cast<int64_t>(j) * ld] =
static_cast<float>(sh_beta);
}
__syncthreads();
float vv = 0.0f;
if (r == start) vv = 1.0f;
else if (r > start && r < n && sh_tau != 0.0)
vv = static_cast<float>(static_cast<double>(sx[r]) * sh_inv);
sx[r] = vv;
Vb[r + static_cast<int64_t>(s) * ld] = vv;
if (r > start && r < n)
Ab[r + static_cast<int64_t>(j) * ld] = vv;
__syncthreads();
for (int p = warp; p < s; p += 16) {
double av = 0.0;
double aw = 0.0;
for (int rr = start + lane; rr < n; rr += 32) {
double vr = sx[rr];
if constexpr (CachePrefix16) {
if (p < kPanel512PrefixCols) {
av += static_cast<double>(sVprefix[rr + p * 512]) * vr;
aw += static_cast<double>(sWprefix[rr + p * 512]) * vr;
} else {
av += static_cast<double>(
Vb[rr + static_cast<int64_t>(p) * ld]) * vr;
aw += static_cast<double>(
Wb[rr + static_cast<int64_t>(p) * ld]) * vr;
}
} else {
av += static_cast<double>(
Vb[rr + static_cast<int64_t>(p) * ld]) * vr;
aw += static_cast<double>(
Wb[rr + static_cast<int64_t>(p) * ld]) * vr;
}
}
av = warp_sum_double(av);
aw = warp_sum_double(aw);
if (lane == 0) {
scv[p] = static_cast<float>(av);
scw[p] = static_cast<float>(aw);
}
}
__syncthreads();
if (r >= start && r < n) {
double acc = 0.0;
for (int c = start; c < n; ++c) {
acc += static_cast<double>(Ab[r + static_cast<int64_t>(c) * ld]) *
sx[c];
}
double value = static_cast<double>(static_cast<float>(acc));
if constexpr (CachePrefix16) {
#pragma unroll
for (int p = 0; p < kPanel512PrefixCols; ++p) {
value -= static_cast<double>(sVprefix[r + p * 512]) * scw[p];
value -= static_cast<double>(sWprefix[r + p * 512]) * scv[p];
}
for (int p = kPanel512PrefixCols; p < s; ++p) {
value -= static_cast<double>(
Vb[r + static_cast<int64_t>(p) * ld]) * scw[p];
value -= static_cast<double>(
Wb[r + static_cast<int64_t>(p) * ld]) * scv[p];
}
} else {
for (int p = 0; p < s; ++p) {
value -= static_cast<double>(
Vb[r + static_cast<int64_t>(p) * ld]) * scw[p];
value -= static_cast<double>(
Wb[r + static_cast<int64_t>(p) * ld]) * scv[p];
}
}
sy[r] = static_cast<float>(value);
}
__syncthreads();
double dot = 0.0;
if (r >= start && r < n) {
dot = static_cast<double>(sx[r]) * sy[r];
}
dot = block_sum_double<kPanel512Threads>(dot, red);
if (tid == 0)
sh_alpha = -0.5 * sh_tau * sh_tau * dot;
__syncthreads();
if (warp == 0) {
float tf = static_cast<float>(sh_tau);
if (tf == 0.0f) {
if (lane <= s)
sT[lane + static_cast<int64_t>(s) * panel_b] = 0.0f;
} else {
if (lane < s) {
double acc = 0.0;
for (int q = lane; q < s; ++q) {
float tmpq = -tf * scv[q];
acc += static_cast<double>(
sT[lane + static_cast<int64_t>(q) * panel_b]) *
tmpq;
}
sT[lane + static_cast<int64_t>(s) * panel_b] =
static_cast<float>(acc);
}
if (lane == s)
sT[s + static_cast<int64_t>(s) * panel_b] = tf;
}
}
__syncthreads();
float w = 0.0f;
if (r >= start && r < n) {
w = static_cast<float>(
sh_tau * static_cast<double>(sy[r]) +
sh_alpha * static_cast<double>(sx[r]));
}
Wb[r + static_cast<int64_t>(s) * ld] = w;
__syncthreads();
}
int chunk_t_elems = panel_b * chunk_cols;
for (int idx = tid; idx < chunk_t_elems; idx += blockDim.x) {
int local_col = idx / panel_b;
int row = idx - local_col * panel_b;
int col = s_begin + local_col;
Tb[row + static_cast<int64_t>(col) * panel_b] =
sT[row + static_cast<int64_t>(col) * panel_b];
}
}
template<int Capacity>
__global__ void __cluster_dims__(2, 1, 1)
persistent_panel_cluster2_kernel(
float* __restrict__ A,
float* __restrict__ T,
float* __restrict__ tau,
float* __restrict__ diag,
float* __restrict__ offdiag,
float* __restrict__ V,
float* __restrict__ W,
float* __restrict__ U2,
float* __restrict__ Z2,
int n, int ld, int panel_b, int num_panels,
int panel_id, int k, int bcols) {
static_assert(Capacity == 512 || Capacity == 1024,
"cluster panel is specialized for 512/1024 capacity tiers");
cg::cluster_group cluster = cg::this_cluster();
int rank = static_cast<int>(cluster.block_rank());
int b = static_cast<int>(blockIdx.x) >> 1;
int tid = threadIdx.x;
int rows_per_cta = ceil_div(ld, 2);
int row_begin = rank * rows_per_cta;
int row_end = min(ld, row_begin + rows_per_cta);
float* Ab = A + static_cast<int64_t>(b) * ld * n;
float* taub = tau + static_cast<int64_t>(b) * n;
float* db = diag + static_cast<int64_t>(b) * n;
float* eb = offdiag + static_cast<int64_t>(b) * n;
float* Vb = V + static_cast<int64_t>(b) * ld * panel_b;
float* Wb = W + static_cast<int64_t>(b) * ld * panel_b;
float* Ub = U2 + static_cast<int64_t>(b) * ld * (2 * panel_b);
float* Zb = Z2 + static_cast<int64_t>(b) * ld * (2 * panel_b);
float* Tb = T + (static_cast<int64_t>(b) * num_panels + panel_id) *
panel_b * panel_b;
extern __shared__ unsigned char dynamic_smem[];
float* sV = reinterpret_cast<float*>(dynamic_smem);
float* sW = sV + static_cast<int64_t>(rows_per_cta) * panel_b;
float* sx = sW + static_cast<int64_t>(rows_per_cta) * panel_b;
float* sy = sx + rows_per_cta;
float* svfull = sy + rows_per_cta;
float* sT = svfull + ld;
float* scv = sT + panel_b * panel_b;
float* scw = scv + panel_b;
float* sjv = scw + panel_b;
float* sjw = sjv + panel_b;
double* red = reinterpret_cast<double*>(sjw + panel_b);
double* red_pair = red + 8;
double* state = red_pair + 8;
constexpr int ST_PART_MAX = 0;
constexpr int ST_PART_TOTAL = 1;
constexpr int ST_PART_TAIL = 2;
constexpr int ST_PART_DOT = 3;
constexpr int ST_MAX = 4;
constexpr int ST_TOTAL = 5;
constexpr int ST_TAIL = 6;
constexpr int ST_TAU = 7;
constexpr int ST_BETA = 8;
constexpr int ST_INV = 9;
constexpr int ST_ALPHA = 10;
int local_panel_elems = rows_per_cta * panel_b;
for (int t = tid; t < 2 * local_panel_elems; t += blockDim.x)
sV[t] = 0.0f; // sV and sW are contiguous.
for (int t = tid; t < panel_b * panel_b; t += blockDim.x)
sT[t] = 0.0f;
for (int t = tid; t < ld; t += blockDim.x)
svfull[t] = 0.0f;
cluster.sync();
for (int s = 0; s < bcols; ++s) {
int j = k + s;
int start = j + 1;
if (rank == 0) {
int owner = min(j / rows_per_cta, 1);
int local_j = j - owner * rows_per_cta;
float* owner_v = cluster.map_shared_rank(sV, owner);
float* owner_w = cluster.map_shared_rank(sW, owner);
for (int p = tid; p < s; p += blockDim.x) {
sjv[p] = owner_v[local_j + static_cast<int64_t>(p) * rows_per_cta];
sjw[p] = owner_w[local_j + static_cast<int64_t>(p) * rows_per_cta];
}
}
cluster.sync();
float* leader_jv = cluster.map_shared_rank(sjv, 0);
float* leader_jw = cluster.map_shared_rank(sjw, 0);
if (rank == 0 && tid == 0) {
double d = Ab[j + static_cast<int64_t>(j) * ld];
for (int p = 0; p < s; ++p) {
d -= 2.0 * static_cast<double>(leader_jv[p]) *
static_cast<double>(leader_jw[p]);
}
db[j] = static_cast<float>(d);
}
int active_begin = max(start, row_begin);
for (int r = active_begin + tid; r < min(n, row_end);
r += blockDim.x) {
int lr = r - row_begin;
double value = Ab[r + static_cast<int64_t>(j) * ld];
for (int p = 0; p < s; ++p) {
value -= static_cast<double>(
sV[lr + static_cast<int64_t>(p) * rows_per_cta]) *
leader_jw[p];
value -= static_cast<double>(
sW[lr + static_cast<int64_t>(p) * rows_per_cta]) *
leader_jv[p];
}
sx[lr] = static_cast<float>(value);
}
double local_max = 0.0;
for (int r = active_begin + tid; r < min(n, row_end);
r += blockDim.x) {
local_max = fmax(local_max,
fabs(static_cast<double>(sx[r - row_begin])));
}
local_max = block_max_double<kThreads>(local_max, red);
if (tid == 0) state[ST_PART_MAX] = local_max;
cluster.sync();
if (rank == 0 && tid == 0) {
double* state1 = cluster.map_shared_rank(state, 1);
state[ST_MAX] = fmax(state[ST_PART_MAX], state1[ST_PART_MAX]);
}
cluster.sync();
double* leader_state = cluster.map_shared_rank(state, 0);
double global_max = leader_state[ST_MAX];
double local_total = 0.0;
double local_tail = 0.0;
if (global_max > 0.0) {
for (int r = active_begin + tid; r < min(n, row_end);
r += blockDim.x) {
double q = static_cast<double>(sx[r - row_begin]) / global_max;
local_total += q * q;
if (r > start) local_tail += q * q;
}
}
double2 norm_sums = block_sum_double_pair<kThreads>(
local_total, local_tail, red, red_pair);
local_total = norm_sums.x;
local_tail = norm_sums.y;
if (tid == 0) {
state[ST_PART_TOTAL] = local_total;
state[ST_PART_TAIL] = local_tail;
}
cluster.sync();
if (rank == 0 && tid == 0) {
double* state1 = cluster.map_shared_rank(state, 1);
double total = state[ST_PART_TOTAL] + state1[ST_PART_TOTAL];
double tail = state[ST_PART_TAIL] + state1[ST_PART_TAIL];
state[ST_TOTAL] = total;
state[ST_TAIL] = tail;
int owner = min(start / rows_per_cta, 1);
int local_start = start - owner * rows_per_cta;
float* owner_x = cluster.map_shared_rank(sx, owner);
double x0 = owner_x[local_start];
double beta = x0;
double tauv = 0.0;
double inv = 0.0;
if (global_max != 0.0 && tail != 0.0) {
double norm = global_max * sqrt(total);
beta = -copysign(norm, x0);
tauv = (beta - x0) / beta;
inv = 1.0 / (x0 - beta);
}
state[ST_BETA] = beta;
state[ST_TAU] = tauv;
state[ST_INV] = inv;
taub[j] = static_cast<float>(tauv);
eb[j] = static_cast<float>(beta);
Ab[start + static_cast<int64_t>(j) * ld] = static_cast<float>(beta);
}
cluster.sync();
double tauv = leader_state[ST_TAU];
double inv = leader_state[ST_INV];
for (int lr = tid; lr < rows_per_cta; lr += blockDim.x) {
int r = row_begin + lr;
float vv = 0.0f;
if (r == start) {
vv = 1.0f;
} else if (r > start && r < n && tauv != 0.0) {
vv = static_cast<float>(static_cast<double>(sx[lr]) * inv);
}
sV[lr + static_cast<int64_t>(s) * rows_per_cta] = vv;
if (r > start && r < n)
Ab[r + static_cast<int64_t>(j) * ld] = vv;
}
cluster.sync();
for (int r = tid; r < ld; r += blockDim.x) {
int owner = min(r / rows_per_cta, 1);
int lr = r - owner * rows_per_cta;
float* owner_v = cluster.map_shared_rank(sV, owner);
svfull[r] = owner_v[lr + static_cast<int64_t>(s) * rows_per_cta];
}
__syncthreads();
int lane = tid & 31;
int warp = tid >> 5;
for (int p = warp; p < s; p += 8) {
double av = 0.0;
double aw = 0.0;
for (int r = active_begin + lane; r < min(n, row_end); r += 32) {
int lr = r - row_begin;
double vr = sV[lr + static_cast<int64_t>(s) * rows_per_cta];
av += static_cast<double>(
sV[lr + static_cast<int64_t>(p) * rows_per_cta]) * vr;
aw += static_cast<double>(
sW[lr + static_cast<int64_t>(p) * rows_per_cta]) * vr;
}
av = warp_sum_double(av);
aw = warp_sum_double(aw);
if (lane == 0) {
scv[p] = static_cast<float>(av);
scw[p] = static_cast<float>(aw);
}
}
cluster.sync();
if (rank == 0) {
float* cv1 = cluster.map_shared_rank(scv, 1);
float* cw1 = cluster.map_shared_rank(scw, 1);
for (int p = tid; p < s; p += blockDim.x) {
scv[p] += cv1[p];
scw[p] += cw1[p];
}
}
cluster.sync();
float* leader_cv = cluster.map_shared_rank(scv, 0);
float* leader_cw = cluster.map_shared_rank(scw, 0);
for (int r = active_begin + tid; r < min(n, row_end);
r += blockDim.x) {
int lr = r - row_begin;
double acc = 0.0;
for (int c = start; c < n; ++c) {
acc += static_cast<double>(Ab[r + static_cast<int64_t>(c) * ld]) *
svfull[c];
}
double value = static_cast<double>(static_cast<float>(acc));
for (int p = 0; p < s; ++p) {
value -= static_cast<double>(
sV[lr + static_cast<int64_t>(p) * rows_per_cta]) *
leader_cw[p];
value -= static_cast<double>(
sW[lr + static_cast<int64_t>(p) * rows_per_cta]) *
leader_cv[p];
}
sy[lr] = static_cast<float>(value);
}
double local_dot = 0.0;
for (int r = active_begin + tid; r < min(n, row_end);
r += blockDim.x) {
int lr = r - row_begin;
local_dot += static_cast<double>(
sV[lr + static_cast<int64_t>(s) * rows_per_cta]) *
sy[lr];
}
local_dot = block_sum_double<kThreads>(local_dot, red);
if (tid == 0) state[ST_PART_DOT] = local_dot;
cluster.sync();
if (rank == 0 && tid == 0) {
double* state1 = cluster.map_shared_rank(state, 1);
double dot = state[ST_PART_DOT] + state1[ST_PART_DOT];
state[ST_ALPHA] = -0.5 * tauv * tauv * dot;
float tf = static_cast<float>(tauv);
if (tf == 0.0f) {
for (int r = 0; r <= s; ++r)
sT[r + static_cast<int64_t>(s) * panel_b] = 0.0f;
} else {
float tmp[kPanel];
for (int i = 0; i < s; ++i) tmp[i] = -tf * leader_cv[i];
for (int r = 0; r < s; ++r) {
double acc = 0.0;
for (int q = r; q < s; ++q) {
acc += static_cast<double>(
sT[r + static_cast<int64_t>(q) * panel_b]) * tmp[q];
}
sT[r + static_cast<int64_t>(s) * panel_b] =
static_cast<float>(acc);
}
sT[s + static_cast<int64_t>(s) * panel_b] = tf;
}
}
cluster.sync();
double alpha = leader_state[ST_ALPHA];
for (int lr = tid; lr < rows_per_cta; lr += blockDim.x) {
int r = row_begin + lr;
float w = 0.0f;
if (r >= start && r < n) {
w = static_cast<float>(
tauv * static_cast<double>(sy[lr]) +
alpha * sV[lr + static_cast<int64_t>(s) * rows_per_cta]);
}
sW[lr + static_cast<int64_t>(s) * rows_per_cta] = w;
}
cluster.sync();
}
int r0 = k + bcols;
int m = n - r0;
if (r0 == n - 1 && rank == 0 && tid == 0) {
int last = n - 1;
int owner = min(last / rows_per_cta, 1);
int local_last = last - owner * rows_per_cta;
float* owner_v = cluster.map_shared_rank(sV, owner);
float* owner_w = cluster.map_shared_rank(sW, owner);
double d = Ab[last + static_cast<int64_t>(last) * ld];
for (int p = 0; p < bcols; ++p) {
d -= 2.0 * static_cast<double>(
owner_v[local_last + static_cast<int64_t>(p) * rows_per_cta]) *
static_cast<double>(
owner_w[local_last + static_cast<int64_t>(p) * rows_per_cta]);
}
db[last] = static_cast<float>(d);
eb[last] = 0.0f;
}
cluster.sync();
if (rank == 0) {
for (int t = tid; t < panel_b * panel_b; t += blockDim.x)
Tb[t] = sT[t];
}
if (r0 < n - 1 && (m & 3) == 0) {
for (int p = 0; p < bcols; ++p) {
for (int lr = tid; lr < rows_per_cta; lr += blockDim.x) {
int r = row_begin + lr;
if (r >= ld) continue;
float v = sV[lr + static_cast<int64_t>(p) * rows_per_cta];
float w = sW[lr + static_cast<int64_t>(p) * rows_per_cta];
Ub[r + static_cast<int64_t>(p) * ld] = v;
Zb[r + static_cast<int64_t>(p) * ld] = w;
Ub[r + static_cast<int64_t>(bcols + p) * ld] = w;
Zb[r + static_cast<int64_t>(bcols + p) * ld] = v;
}
}
} else if (r0 < n - 1) {
for (int p = 0; p < bcols; ++p) {
for (int lr = tid; lr < rows_per_cta; lr += blockDim.x) {
int r = row_begin + lr;
if (r >= ld) continue;
Vb[r + static_cast<int64_t>(p) * ld] =
sV[lr + static_cast<int64_t>(p) * rows_per_cta];
Wb[r + static_cast<int64_t>(p) * ld] =
sW[lr + static_cast<int64_t>(p) * rows_per_cta];
}
}
}
}
template<int Capacity, int ClusterCtas>
__device__ __forceinline__ void persistent_panel_large_body(
unsigned char* __restrict__ dynamic_smem,
float* __restrict__ A,
float* __restrict__ T,
float* __restrict__ tau,
float* __restrict__ diag,
float* __restrict__ offdiag,
float* __restrict__ V,
float* __restrict__ W,
float* __restrict__ U2,
float* __restrict__ Z2,
int n, int ld, int panel_b, int num_panels,
int panel_id, int k, int bcols) {
static_assert(Capacity == 1024 || Capacity == 2048,
"large panel capacity must be 1024 or 2048");
static_assert(ClusterCtas == 4 || ClusterCtas == 8,
"large panel cluster must use four or eight CTAs");
static_assert(Capacity / ClusterCtas == 256,
"large panel owns 256 rows per CTA");
cg::cluster_group cluster = cg::this_cluster();
int rank = static_cast<int>(cluster.block_rank());
int b = static_cast<int>(blockIdx.x) / ClusterCtas;
int tid = threadIdx.x;
constexpr int RowsPerCta = Capacity / ClusterCtas;
int row_begin = rank * RowsPerCta;
int row_end = min(ld, row_begin + RowsPerCta);
float* Ab = A + static_cast<int64_t>(b) * ld * n;
float* taub = tau + static_cast<int64_t>(b) * n;
float* db = diag + static_cast<int64_t>(b) * n;
float* eb = offdiag + static_cast<int64_t>(b) * n;
float* Vb = V + static_cast<int64_t>(b) * ld * panel_b;
float* Wb = W + static_cast<int64_t>(b) * ld * panel_b;
float* Ub = U2 + static_cast<int64_t>(b) * ld * (2 * panel_b);
float* Zb = Z2 + static_cast<int64_t>(b) * ld * (2 * panel_b);
float* Tb = T + (static_cast<int64_t>(b) * num_panels + panel_id) *
panel_b * panel_b;
float* sV = reinterpret_cast<float*>(dynamic_smem);
float* sW = sV + static_cast<int64_t>(RowsPerCta) * panel_b;
float* sx = sW + static_cast<int64_t>(RowsPerCta) * panel_b;
float* sy = sx + RowsPerCta;
float* svfull = sy + RowsPerCta;
float* sT = svfull + Capacity;
float* scv = sT + panel_b * panel_b;
float* scw = scv + panel_b;
float* sjv = scw + panel_b;
float* sjw = sjv + panel_b;
double* red = reinterpret_cast<double*>(sjw + panel_b);
double* red_pair = red + 8;
double* state = red_pair + 8;
constexpr int ST_PART_MAX = 0;
constexpr int ST_PART_TOTAL = 1;
constexpr int ST_PART_TAIL = 2;
constexpr int ST_PART_DOT = 3;
constexpr int ST_MAX = 4;
constexpr int ST_TOTAL = 5;
constexpr int ST_TAIL = 6;
constexpr int ST_TAU = 7;
constexpr int ST_BETA = 8;
constexpr int ST_INV = 9;
constexpr int ST_ALPHA = 10;
int local_panel_elems = RowsPerCta * panel_b;
for (int t = tid; t < 2 * local_panel_elems; t += blockDim.x)
sV[t] = 0.0f;
if (rank == 0) {
for (int t = tid; t < panel_b * panel_b; t += blockDim.x)
sT[t] = 0.0f;
}
for (int t = tid; t < Capacity; t += blockDim.x)
svfull[t] = 0.0f;
cluster.sync();
for (int s = 0; s < bcols; ++s) {
int j = k + s;
int start = j + 1;
if (rank == 0) {
int owner = min(j / RowsPerCta, ClusterCtas - 1);
int local_j = j - owner * RowsPerCta;
float* owner_v = cluster.map_shared_rank(sV, owner);
float* owner_w = cluster.map_shared_rank(sW, owner);
for (int p = tid; p < s; p += blockDim.x) {
sjv[p] = owner_v[local_j + static_cast<int64_t>(p) * RowsPerCta];
sjw[p] = owner_w[local_j + static_cast<int64_t>(p) * RowsPerCta];
}
}
cluster.sync();
float* leader_jv = cluster.map_shared_rank(sjv, 0);
float* leader_jw = cluster.map_shared_rank(sjw, 0);
if (rank == 0 && tid == 0) {
double d = Ab[j + static_cast<int64_t>(j) * ld];
for (int p = 0; p < s; ++p) {
d -= 2.0 * static_cast<double>(leader_jv[p]) *
static_cast<double>(leader_jw[p]);
}
db[j] = static_cast<float>(d);
}
int active_begin = max(start, row_begin);
int active_end = min(n, row_end);
for (int gr = active_begin + tid; gr < active_end; gr += blockDim.x) {
int lr = gr - row_begin;
double value = Ab[gr + static_cast<int64_t>(j) * ld];
for (int p = 0; p < s; ++p) {
value -= static_cast<double>(
sV[lr + static_cast<int64_t>(p) * RowsPerCta]) *
leader_jw[p];
value -= static_cast<double>(
sW[lr + static_cast<int64_t>(p) * RowsPerCta]) *
leader_jv[p];
}
sx[lr] = static_cast<float>(value);
}
double local_max = 0.0;
for (int gr = active_begin + tid; gr < active_end; gr += blockDim.x) {
local_max = fmax(local_max,
fabs(static_cast<double>(sx[gr - row_begin])));
}
local_max = block_max_double<kThreads>(local_max, red);
if (tid == 0) state[ST_PART_MAX] = local_max;
cluster.sync();
if (rank == 0 && tid == 0) {
double global_max = 0.0;
for (int rr = 0; rr < ClusterCtas; ++rr) {
double* remote_state = cluster.map_shared_rank(state, rr);
global_max = fmax(global_max, remote_state[ST_PART_MAX]);
}
state[ST_MAX] = global_max;
}
cluster.sync();
double* leader_state = cluster.map_shared_rank(state, 0);
double global_max = leader_state[ST_MAX];
double local_total = 0.0;
double local_tail = 0.0;
if (global_max > 0.0) {
for (int gr = active_begin + tid; gr < active_end; gr += blockDim.x) {
double q = static_cast<double>(sx[gr - row_begin]) / global_max;
local_total += q * q;
if (gr > start) local_tail += q * q;
}
}
double2 norm_sums = block_sum_double_pair<kThreads>(
local_total, local_tail, red, red_pair);
if (tid == 0) {
state[ST_PART_TOTAL] = norm_sums.x;
state[ST_PART_TAIL] = norm_sums.y;
}
cluster.sync();
if (rank == 0 && tid == 0) {
double total = 0.0;
double tail = 0.0;
for (int rr = 0; rr < ClusterCtas; ++rr) {
double* remote_state = cluster.map_shared_rank(state, rr);
total += remote_state[ST_PART_TOTAL];
tail += remote_state[ST_PART_TAIL];
}
state[ST_TOTAL] = total;
state[ST_TAIL] = tail;
int owner = min(start / RowsPerCta, ClusterCtas - 1);
int local_start = start - owner * RowsPerCta;
float* owner_x = cluster.map_shared_rank(sx, owner);
double x0 = owner_x[local_start];
double beta = x0;
double tauv = 0.0;
double inv = 0.0;
if (global_max != 0.0 && tail != 0.0) {
double norm = global_max * sqrt(total);
beta = -copysign(norm, x0);
tauv = (beta - x0) / beta;
inv = 1.0 / (x0 - beta);
}
state[ST_BETA] = beta;
state[ST_TAU] = tauv;
state[ST_INV] = inv;
taub[j] = static_cast<float>(tauv);
eb[j] = static_cast<float>(beta);
Ab[start + static_cast<int64_t>(j) * ld] = static_cast<float>(beta);
}
cluster.sync();
double tauv = leader_state[ST_TAU];
double inv = leader_state[ST_INV];
for (int lr = tid; lr < RowsPerCta; lr += blockDim.x) {
int gr = row_begin + lr;
float vv = 0.0f;
if (gr == start) {
vv = 1.0f;
} else if (gr > start && gr < n && tauv != 0.0) {
vv = static_cast<float>(static_cast<double>(sx[lr]) * inv);
}
sV[lr + static_cast<int64_t>(s) * RowsPerCta] = vv;
if (gr > start && gr < n)
Ab[gr + static_cast<int64_t>(j) * ld] = vv;
}
cluster.sync();
for (int gc = tid; gc < ld; gc += blockDim.x) {
int owner = min(gc / RowsPerCta, ClusterCtas - 1);
int lr = gc - owner * RowsPerCta;
float* owner_v = cluster.map_shared_rank(sV, owner);
svfull[gc] = owner_v[lr + static_cast<int64_t>(s) * RowsPerCta];
}
__syncthreads();
int lane = tid & 31;
int warp = tid >> 5;
for (int p = warp; p < s; p += 8) {
double av = 0.0;
double aw = 0.0;
for (int gr = active_begin + lane; gr < active_end; gr += 32) {
int lr = gr - row_begin;
double vr = sV[lr + static_cast<int64_t>(s) * RowsPerCta];
av += static_cast<double>(
sV[lr + static_cast<int64_t>(p) * RowsPerCta]) * vr;
aw += static_cast<double>(
sW[lr + static_cast<int64_t>(p) * RowsPerCta]) * vr;
}
av = warp_sum_double(av);
aw = warp_sum_double(aw);
if (lane == 0) {
scv[p] = static_cast<float>(av);
scw[p] = static_cast<float>(aw);
}
}
cluster.sync();
if (rank == 0) {
for (int p = tid; p < s; p += blockDim.x) {
double av = 0.0;
double aw = 0.0;
for (int rr = 0; rr < ClusterCtas; ++rr) {
float* remote_cv = cluster.map_shared_rank(scv, rr);
float* remote_cw = cluster.map_shared_rank(scw, rr);
av += static_cast<double>(remote_cv[p]);
aw += static_cast<double>(remote_cw[p]);
}
scv[p] = static_cast<float>(av);
scw[p] = static_cast<float>(aw);
}
}
cluster.sync();
float* leader_cv = cluster.map_shared_rank(scv, 0);
float* leader_cw = cluster.map_shared_rank(scw, 0);
for (int gr = active_begin + tid; gr < active_end; gr += blockDim.x) {
int lr = gr - row_begin;
double acc = 0.0;
for (int c = start; c < n; ++c) {
acc += static_cast<double>(Ab[gr + static_cast<int64_t>(c) * ld]) *
svfull[c];
}
double value = static_cast<double>(static_cast<float>(acc));
for (int p = 0; p < s; ++p) {
value -= static_cast<double>(
sV[lr + static_cast<int64_t>(p) * RowsPerCta]) *
leader_cw[p];
value -= static_cast<double>(
sW[lr + static_cast<int64_t>(p) * RowsPerCta]) *
leader_cv[p];
}
sy[lr] = static_cast<float>(value);
}
double local_dot = 0.0;
for (int gr = active_begin + tid; gr < active_end; gr += blockDim.x) {
int lr = gr - row_begin;
local_dot += static_cast<double>(
sV[lr + static_cast<int64_t>(s) * RowsPerCta]) *
sy[lr];
}
local_dot = block_sum_double<kThreads>(local_dot, red);
if (tid == 0) state[ST_PART_DOT] = local_dot;
cluster.sync();
if (rank == 0 && tid == 0) {
double dot = 0.0;
for (int rr = 0; rr < ClusterCtas; ++rr) {
double* remote_state = cluster.map_shared_rank(state, rr);
dot += remote_state[ST_PART_DOT];
}
state[ST_ALPHA] = -0.5 * tauv * tauv * dot;
float tf = static_cast<float>(tauv);
if (tf == 0.0f) {
for (int rr = 0; rr <= s; ++rr)
sT[rr + static_cast<int64_t>(s) * panel_b] = 0.0f;
} else {
float tmp[kPanel];
for (int i = 0; i < s; ++i) tmp[i] = -tf * leader_cv[i];
for (int rr = 0; rr < s; ++rr) {
double acc = 0.0;
for (int q = rr; q < s; ++q) {
acc += static_cast<double>(
sT[rr + static_cast<int64_t>(q) * panel_b]) * tmp[q];
}
sT[rr + static_cast<int64_t>(s) * panel_b] =
static_cast<float>(acc);
}
sT[s + static_cast<int64_t>(s) * panel_b] = tf;
}
}
cluster.sync();
double alpha = leader_state[ST_ALPHA];
for (int lr = tid; lr < RowsPerCta; lr += blockDim.x) {
int gr = row_begin + lr;
float w = 0.0f;
if (gr >= start && gr < n) {
w = static_cast<float>(
tauv * static_cast<double>(sy[lr]) +
alpha * sV[lr + static_cast<int64_t>(s) * RowsPerCta]);
}
sW[lr + static_cast<int64_t>(s) * RowsPerCta] = w;
}
cluster.sync();
}
int r0 = k + bcols;
int m = n - r0;
if (r0 == n - 1 && rank == 0 && tid == 0) {
int last = n - 1;
int owner = min(last / RowsPerCta, ClusterCtas - 1);
int local_last = last - owner * RowsPerCta;
float* owner_v = cluster.map_shared_rank(sV, owner);
float* owner_w = cluster.map_shared_rank(sW, owner);
double d = Ab[last + static_cast<int64_t>(last) * ld];
for (int p = 0; p < bcols; ++p) {
d -= 2.0 * static_cast<double>(
owner_v[local_last + static_cast<int64_t>(p) * RowsPerCta]) *
static_cast<double>(
owner_w[local_last + static_cast<int64_t>(p) * RowsPerCta]);
}
db[last] = static_cast<float>(d);
eb[last] = 0.0f;
}
cluster.sync();
if (rank == 0) {
for (int t = tid; t < panel_b * panel_b; t += blockDim.x)
Tb[t] = sT[t];
}
if (r0 < n - 1 && (m & 3) == 0) {
for (int p = 0; p < bcols; ++p) {
for (int lr = tid; lr < RowsPerCta; lr += blockDim.x) {
int gr = row_begin + lr;
if (gr >= ld) continue;
float v = sV[lr + static_cast<int64_t>(p) * RowsPerCta];
float w = sW[lr + static_cast<int64_t>(p) * RowsPerCta];
Ub[gr + static_cast<int64_t>(p) * ld] = v;
Zb[gr + static_cast<int64_t>(p) * ld] = w;
Ub[gr + static_cast<int64_t>(bcols + p) * ld] = w;
Zb[gr + static_cast<int64_t>(bcols + p) * ld] = v;
}
}
} else if (r0 < n - 1) {
for (int p = 0; p < bcols; ++p) {
for (int lr = tid; lr < RowsPerCta; lr += blockDim.x) {
int gr = row_begin + lr;
if (gr >= ld) continue;
Vb[gr + static_cast<int64_t>(p) * ld] =
sV[lr + static_cast<int64_t>(p) * RowsPerCta];
Wb[gr + static_cast<int64_t>(p) * ld] =
sW[lr + static_cast<int64_t>(p) * RowsPerCta];
}
}
}
cluster.sync();
}
__global__ void __cluster_dims__(4, 1, 1)
persistent_panel_large1024_kernel(
float* __restrict__ A,
float* __restrict__ T,
float* __restrict__ tau,
float* __restrict__ diag,
float* __restrict__ offdiag,
float* __restrict__ V,
float* __restrict__ W,
float* __restrict__ U2,
float* __restrict__ Z2,
int n, int ld, int panel_b, int num_panels,
int panel_id, int k, int bcols) {
extern __shared__ unsigned char dynamic_smem[];
persistent_panel_large_body<1024, 4>(
dynamic_smem, A, T, tau, diag, offdiag, V, W, U2, Z2,
n, ld, panel_b, num_panels, panel_id, k, bcols);
}
__global__ void __cluster_dims__(8, 1, 1)
persistent_panel_large2048_kernel(
float* __restrict__ A,
float* __restrict__ T,
float* __restrict__ tau,
float* __restrict__ diag,
float* __restrict__ offdiag,
float* __restrict__ V,
float* __restrict__ W,
float* __restrict__ U2,
float* __restrict__ Z2,
int n, int ld, int panel_b, int num_panels,
int panel_id, int k, int bcols) {
extern __shared__ unsigned char dynamic_smem[];
persistent_panel_large_body<2048, 8>(
dynamic_smem, A, T, tau, diag, offdiag, V, W, U2, Z2,
n, ld, panel_b, num_panels, panel_id, k, bcols);
}
__global__ void panel_correct_column_kernel(
const float* __restrict__ A,
const float* __restrict__ V,
const float* __restrict__ W,
float* __restrict__ x,
float* __restrict__ diag,
int n, int ld, int panel_b, int j, int s) {
int b = blockIdx.y;
const float* Ab = A + static_cast<int64_t>(b) * ld * n;
const float* Vb = V + static_cast<int64_t>(b) * ld * panel_b;
const float* Wb = W + static_cast<int64_t>(b) * ld * panel_b;
if (blockIdx.x == 0 && threadIdx.x == 0) {
double d = Ab[j + static_cast<int64_t>(j) * ld];
for (int p = 0; p < s; ++p)
d -= 2.0 * static_cast<double>(Vb[j + static_cast<int64_t>(p) * ld]) *
static_cast<double>(Wb[j + static_cast<int64_t>(p) * ld]);
diag[static_cast<int64_t>(b) * n + j] = static_cast<float>(d);
}
int r = j + 1 + blockIdx.x * blockDim.x + threadIdx.x;
if (r >= n) return;
float* xb = x + static_cast<int64_t>(b) * ld;
double value = Ab[r + static_cast<int64_t>(j) * ld];
for (int p = 0; p < s; ++p) {
value -= static_cast<double>(Vb[r + static_cast<int64_t>(p) * ld]) *
Wb[j + static_cast<int64_t>(p) * ld];
value -= static_cast<double>(Wb[r + static_cast<int64_t>(p) * ld]) *
Vb[j + static_cast<int64_t>(p) * ld];
}
xb[r] = static_cast<float>(value);
}
__global__ void panel_householder_kernel(
float* __restrict__ A,
const float* __restrict__ x,
float* __restrict__ V,
float* __restrict__ tau,
float* __restrict__ offdiag,
int n, int ld, int panel_b, int j, int s) {
int b = blockIdx.x;
int tid = threadIdx.x;
int start = j + 1;
const float* xb = x + static_cast<int64_t>(b) * ld;
float* Ab = A + static_cast<int64_t>(b) * ld * n;
float* Vb = V + static_cast<int64_t>(b) * ld * panel_b;
__shared__ double red[8];
__shared__ double red_pair[8];
__shared__ double sh_max, sh_norm, sh_tau, sh_beta, sh_inv;
double local_max = 0.0;
for (int r = start + tid; r < n; r += blockDim.x)
local_max = fmax(local_max, fabs(static_cast<double>(xb[r])));
double mx = block_max_double<kThreads>(local_max, red);
if (tid == 0) sh_max = mx;
__syncthreads();
double total = 0.0;
double tail = 0.0;
if (sh_max > 0.0) {
for (int r = start + tid; r < n; r += blockDim.x) {
double q = static_cast<double>(xb[r]) / sh_max;
total += q * q;
if (r > start) tail += q * q;
}
}
double2 norm_sums = block_sum_double_pair<kThreads>(
total, tail, red, red_pair);
total = norm_sums.x;
tail = norm_sums.y;
if (tid == 0) {
double x0 = xb[start];
if (sh_max == 0.0 || tail == 0.0) {
sh_beta = x0;
sh_tau = 0.0;
sh_inv = 0.0;
sh_norm = fabs(x0);
} else {
double norm = sh_max * sqrt(total);
double beta = -copysign(norm, x0);
double denom = x0 - beta;
sh_beta = beta;
sh_tau = (beta - x0) / beta;
sh_inv = 1.0 / denom;
sh_norm = norm;
}
tau[static_cast<int64_t>(b) * n + j] = static_cast<float>(sh_tau);
offdiag[static_cast<int64_t>(b) * n + j] = static_cast<float>(sh_beta);
Ab[start + static_cast<int64_t>(j) * ld] = static_cast<float>(sh_beta);
}
__syncthreads();
for (int r = tid; r < n; r += blockDim.x) {
float vv = 0.0f;
if (r == start) vv = 1.0f;
else if (r > start && sh_tau != 0.0)
vv = static_cast<float>(static_cast<double>(xb[r]) * sh_inv);
Vb[r + static_cast<int64_t>(s) * ld] = vv;
if (r > start)
Ab[r + static_cast<int64_t>(j) * ld] = vv;
}
}
__global__ void panel_matvec_kernel(
const float* __restrict__ A,
const float* __restrict__ V,
const float* __restrict__ W,
const float* __restrict__ coeff_v,
const float* __restrict__ coeff_w,
float* __restrict__ y,
int n, int ld, int panel_b, int start, int s) {
int b = blockIdx.y;
int r = start + blockIdx.x * blockDim.x + threadIdx.x;
if (r >= n) return;
const float* Ab = A + static_cast<int64_t>(b) * ld * n;
const float* Vb = V + static_cast<int64_t>(b) * ld * panel_b;
const float* Wb = W + static_cast<int64_t>(b) * ld * panel_b;
const float* cv = coeff_v + static_cast<int64_t>(b) * panel_b;
const float* cw = coeff_w + static_cast<int64_t>(b) * panel_b;
float* yb = y + static_cast<int64_t>(b) * ld;
double acc = 0.0;
for (int c = start; c < n; ++c)
acc += static_cast<double>(Ab[r + static_cast<int64_t>(c) * ld]) *
Vb[c + static_cast<int64_t>(s) * ld];
double value = static_cast<double>(static_cast<float>(acc));
for (int p = 0; p < s; ++p) {
value -= static_cast<double>(Vb[r + static_cast<int64_t>(p) * ld]) * cw[p];
value -= static_cast<double>(Wb[r + static_cast<int64_t>(p) * ld]) * cv[p];
}
yb[r] = static_cast<float>(value);
}
__global__ void panel_coefficients_kernel(
const float* __restrict__ V,
const float* __restrict__ W,
float* __restrict__ coeff_v,
float* __restrict__ coeff_w,
int n, int ld, int panel_b, int start, int s) {
int b = blockIdx.x;
int tid = threadIdx.x;
const float* Vb = V + static_cast<int64_t>(b) * ld * panel_b;
const float* Wb = W + static_cast<int64_t>(b) * ld * panel_b;
float* cv = coeff_v + static_cast<int64_t>(b) * panel_b;
float* cw = coeff_w + static_cast<int64_t>(b) * panel_b;
__shared__ double red[8];
for (int p = 0; p < s; ++p) {
double av = 0.0, aw = 0.0;
for (int r = start + tid; r < n; r += blockDim.x) {
double vr = Vb[r + static_cast<int64_t>(s) * ld];
av += static_cast<double>(Vb[r + static_cast<int64_t>(p) * ld]) * vr;
aw += static_cast<double>(Wb[r + static_cast<int64_t>(p) * ld]) * vr;
}
av = block_sum_double<kThreads>(av, red);
aw = block_sum_double<kThreads>(aw, red);
if (tid == 0) {
cv[p] = static_cast<float>(av);
cw[p] = static_cast<float>(aw);
}
__syncthreads();
}
}
__global__ void panel_coefficients_matvec_fused_kernel(
const float* __restrict__ A,
const float* __restrict__ V,
const float* __restrict__ W,
float* __restrict__ coeff_v,
float* __restrict__ coeff_w,
float* __restrict__ y,
int n, int ld, int panel_b, int start, int s) {
int b = blockIdx.y;
int tid = threadIdx.x;
const float* Ab = A + static_cast<int64_t>(b) * ld * n;
const float* Vb = V + static_cast<int64_t>(b) * ld * panel_b;
const float* Wb = W + static_cast<int64_t>(b) * ld * panel_b;
float* cv = coeff_v + static_cast<int64_t>(b) * panel_b;
float* cw = coeff_w + static_cast<int64_t>(b) * panel_b;
float* yb = y + static_cast<int64_t>(b) * ld;
__shared__ float scv[kPanel];
__shared__ float scw[kPanel];
int lane = tid & 31;
int warp = tid >> 5;
for (int p = warp; p < s; p += 8) {
double av = 0.0;
double aw = 0.0;
for (int r = start + lane; r < n; r += 32) {
double vr = Vb[r + static_cast<int64_t>(s) * ld];
av += static_cast<double>(Vb[r + static_cast<int64_t>(p) * ld]) * vr;
aw += static_cast<double>(Wb[r + static_cast<int64_t>(p) * ld]) * vr;
}
av = warp_sum_double(av);
aw = warp_sum_double(aw);
if (lane == 0) {
scv[p] = static_cast<float>(av);
scw[p] = static_cast<float>(aw);
if (blockIdx.x == 0) {
cv[p] = scv[p];
cw[p] = scw[p];
}
}
}
__syncthreads();
int r = start + blockIdx.x * blockDim.x + tid;
if (r >= n) return;
double acc = 0.0;
for (int c = start; c < n; ++c)
acc += static_cast<double>(Ab[r + static_cast<int64_t>(c) * ld]) *
Vb[c + static_cast<int64_t>(s) * ld];
double value = static_cast<double>(static_cast<float>(acc));
for (int p = 0; p < s; ++p) {
value -= static_cast<double>(Vb[r + static_cast<int64_t>(p) * ld]) * scw[p];
value -= static_cast<double>(Wb[r + static_cast<int64_t>(p) * ld]) * scv[p];
}
yb[r] = static_cast<float>(value);
}
__global__ void panel_finalize_w_kernel(
const float* __restrict__ V,
const float* __restrict__ y,
const float* __restrict__ tau,
const float* __restrict__ coeff_v,
float* __restrict__ W,
float* __restrict__ T,
int n, int ld, int panel_b, int num_panels,
int panel_id, int j, int s) {
int b = blockIdx.x;
int tid = threadIdx.x;
int start = j + 1;
const float* Vb = V + static_cast<int64_t>(b) * ld * panel_b;
const float* yb = y + static_cast<int64_t>(b) * ld;
const float* cv = coeff_v + static_cast<int64_t>(b) * panel_b;
float* Wb = W + static_cast<int64_t>(b) * ld * panel_b;
float* Tb = T + (static_cast<int64_t>(b) * num_panels + panel_id) *
panel_b * panel_b;
double t = tau[static_cast<int64_t>(b) * n + j];
__shared__ double red[8];
__shared__ double alpha;
double dot = 0.0;
for (int r = start + tid; r < n; r += blockDim.x)
dot += static_cast<double>(Vb[r + static_cast<int64_t>(s) * ld]) * yb[r];
dot = block_sum_double<kThreads>(dot, red);
if (tid == 0) {
alpha = -0.5 * t * t * dot;
float tf = static_cast<float>(t);
if (tf == 0.0f) {
for (int r = 0; r <= s; ++r)
Tb[r + static_cast<int64_t>(s) * panel_b] = 0.0f;
} else {
float tmp[kPanel];
for (int i = 0; i < s; ++i) tmp[i] = -tf * cv[i];
for (int r = 0; r < s; ++r) {
double acc = 0.0;
for (int q = r; q < s; ++q)
acc += static_cast<double>(
Tb[r + static_cast<int64_t>(q) * panel_b]) * tmp[q];
Tb[r + static_cast<int64_t>(s) * panel_b] = static_cast<float>(acc);
}
Tb[s + static_cast<int64_t>(s) * panel_b] = tf;
}
}
__syncthreads();
for (int r = tid; r < n; r += blockDim.x) {
float w = 0.0f;
if (r >= start)
w = static_cast<float>(t * static_cast<double>(yb[r]) +
alpha * Vb[r + static_cast<int64_t>(s) * ld]);
Wb[r + static_cast<int64_t>(s) * ld] = w;
}
}
__global__ void panel_build_t_kernel(
const float* __restrict__ tau,
const float* __restrict__ coeff_v,
float* __restrict__ T,
int n, int panel_b, int num_panels, int panel_id, int j, int s) {
int b = blockIdx.x;
if (threadIdx.x != 0) return;
const float* cv = coeff_v + static_cast<int64_t>(b) * panel_b;
float* Tb = T + (static_cast<int64_t>(b) * num_panels + panel_id) *
panel_b * panel_b;
float t = tau[static_cast<int64_t>(b) * n + j];
if (t == 0.0f) {
for (int r = 0; r <= s; ++r) Tb[r + static_cast<int64_t>(s) * panel_b] = 0.0f;
return;
}
float tmp[kPanel];
for (int i = 0; i < s; ++i) tmp[i] = -t * cv[i];
for (int r = 0; r < s; ++r) {
double acc = 0.0;
for (int q = r; q < s; ++q)
acc += static_cast<double>(Tb[r + static_cast<int64_t>(q) * panel_b]) * tmp[q];
Tb[r + static_cast<int64_t>(s) * panel_b] = static_cast<float>(acc);
}
Tb[s + static_cast<int64_t>(s) * panel_b] = t;
}
__global__ void final_diagonal_kernel(
const float* __restrict__ A,
const float* __restrict__ V,
const float* __restrict__ W,
float* __restrict__ diag,
int n, int ld, int panel_b, int last, int bcols) {
int b = blockIdx.x;
if (threadIdx.x != 0) return;
const float* Ab = A + static_cast<int64_t>(b) * ld * n;
const float* Vb = V + static_cast<int64_t>(b) * ld * panel_b;
const float* Wb = W + static_cast<int64_t>(b) * ld * panel_b;
double d = Ab[last + static_cast<int64_t>(last) * ld];
for (int p = 0; p < bcols; ++p)
d -= 2.0 * static_cast<double>(Vb[last + static_cast<int64_t>(p) * ld]) *
Wb[last + static_cast<int64_t>(p) * ld];
diag[static_cast<int64_t>(b) * n + last] = static_cast<float>(d);
}
__device__ __forceinline__ int jacobi_order32(int step, int pos) {
if (pos == 0) return 0;
int r = (pos - 1) - step;
if (r < 0) r += 31;
return r + 1;
}
__device__ __forceinline__ void jacobi_update_packed32(
const float* __restrict__ cur,
float* __restrict__ nxt,
const int* __restrict__ mate,
const int* __restrict__ pair_id,
const float* __restrict__ cs,
const float* __restrict__ ss,
unsigned packed) {
constexpr int LD = 33;
int i = static_cast<int>(packed >> 5);
int j = static_cast<int>(packed & 31u);
int mi = mate[i];
int mj = mate[j];
int pi = pair_id[i];
int pj = pair_id[j];
float ci = cs[pi];
float cj = cs[pj];
float si = ss[pi];
float sj = ss[pj];
if (i < mi) si = -si;
if (j < mj) sj = -sj;
float a00 = cur[i * LD + j];
float a01 = cur[i * LD + mj];
float a10 = cur[mi * LD + j];
float a11 = cur[mi * LD + mj];
float value = ci * cj * a00 + ci * sj * a01 +
si * cj * a10 + si * sj * a11;
nxt[i * LD + j] = value;
nxt[j * LD + i] = value;
}
__global__ __launch_bounds__(kLeafJacobiThreads)
void dc_leaf_jacobi_kernel(
const float* __restrict__ diag,
const float* __restrict__ offdiag,
const int* __restrict__ leaf_meta, // [leaf][start,len]
int num_leaves,
int n,
double* __restrict__ vals,
float* __restrict__ vecs) {
constexpr int BS = 32;
constexpr int LD = 33;
constexpr int Pairs = 16;
constexpr int Steps = 31;
constexpr int Sweeps = 7;
constexpr float JacobiEps = 1.0e-20f;
int leaf = blockIdx.x;
int b = blockIdx.y;
int tid = threadIdx.x;
if (leaf >= num_leaves) return;
int start = leaf_meta[2 * leaf + 0];
int len = leaf_meta[2 * leaf + 1];
const float* db = diag + static_cast<int64_t>(b) * n;
const float* eb = offdiag + static_cast<int64_t>(b) * n;
double* vb = vals + static_cast<int64_t>(b) * n;
float* Qout = vecs + static_cast<int64_t>(b) * n * n;
__shared__ float A0[BS * LD];
__shared__ float A1[BS * LD];
__shared__ float Q[BS * LD];
__shared__ float cs[Pairs], ss[Pairs];
__shared__ int mate[BS], pair_id[BS], perm[BS];
__shared__ float evals[BS];
__shared__ float bound;
__shared__ uint16_t upper_coord[BS * (BS + 1) / 2];
if (tid < BS) {
int row = tid;
int base = row * (2 * BS - row + 1) / 2;
for (int col = row; col < BS; ++col)
upper_coord[base + col - row] =
static_cast<uint16_t>((row << 5) | col);
}
if (tid == 0) {
float bd = 0.0f;
for (int i = 0; i < len; ++i) {
float radius = 0.0f;
if (i > 0) radius += fabsf(eb[start + i - 1]);
if (i + 1 < len) radius += fabsf(eb[start + i]);
bd = fmaxf(bd, fabsf(db[start + i]) + radius);
}
bound = fmaxf(1.0f, bd);
}
__syncthreads();
for (int t = tid; t < BS * BS; t += blockDim.x) {
int r = t / BS;
int c = t - r * BS;
float a = 0.0f;
if (r < len && c < len) {
if (r == c) {
int g = start + r;
a = db[g];
if (r == 0 && start > 0) a -= fabsf(eb[start - 1]);
if (r == len - 1 && start + len < n)
a -= fabsf(eb[start + len - 1]);
} else if (r + 1 == c) {
a = eb[start + r];
} else if (c + 1 == r) {
a = eb[start + c];
}
} else if (r == c) {
a = bound * (4.0f + static_cast<float>(r - len));
}
A0[r * LD + c] = a;
A1[r * LD + c] = 0.0f;
Q[r * LD + c] = (r == c) ? 1.0f : 0.0f;
}
if (tid < BS) {
mate[tid] = 0;
pair_id[tid] = 0;
perm[tid] = tid;
evals[tid] = 0.0f;
}
__syncthreads();
const unsigned upper0 = upper_coord[tid];
const unsigned upper1 = upper_coord[tid + kLeafJacobiThreads];
const unsigned upper2 =
tid < 16 ? upper_coord[tid + 2 * kLeafJacobiThreads] : 0u;
float* cur = A0;
float* nxt = A1;
for (int sweep = 0; sweep < Sweeps; ++sweep) {
for (int step = 0; step < Steps; ++step) {
if (tid < Pairs) {
int a = jacobi_order32(step, tid);
int bb = jacobi_order32(step, BS - 1 - tid);
int p = min(a, bb), q = max(a, bb);
float app = cur[p * LD + p];
float aqq = cur[q * LD + q];
float apq = cur[p * LD + q];
float c = 1.0f, s = 0.0f;
if (fabsf(apq) > JacobiEps) {
float tauj = (aqq - app) / (2.0f * apq);
float denom = fabsf(tauj) + sqrtf(fmaf(tauj, tauj, 1.0f));
float tt = (tauj >= 0.0f ? 1.0f : -1.0f) / denom;
c = rsqrtf(fmaf(tt, tt, 1.0f));
s = tt * c;
}
cs[tid] = c;
ss[tid] = s;
mate[p] = q; mate[q] = p;
pair_id[p] = tid; pair_id[q] = tid;
}
__syncthreads();
jacobi_update_packed32(
cur, nxt, mate, pair_id, cs, ss, upper0);
jacobi_update_packed32(
cur, nxt, mate, pair_id, cs, ss, upper1);
if (tid < 16)
jacobi_update_packed32(
cur, nxt, mate, pair_id, cs, ss, upper2);
for (int t = tid; t < BS * Pairs; t += blockDim.x) {
int r = t / Pairs, pair = t - r * Pairs;
int a = jacobi_order32(step, pair);
int bb = jacobi_order32(step, BS - 1 - pair);
int p = min(a, bb), q = max(a, bb);
float c = cs[pair], s = ss[pair];
float qp = Q[r * LD + p], qq = Q[r * LD + q];
Q[r * LD + p] = c * qp - s * qq;
Q[r * LD + q] = s * qp + c * qq;
}
__syncthreads();
float* tmp = cur; cur = nxt; nxt = tmp;
}
}
if (tid == 0) {
for (int i = 0; i < BS; ++i) {
evals[i] = cur[i * LD + i];
perm[i] = i;
}
for (int i = 0; i < BS - 1; ++i) {
int best = i;
for (int j = i + 1; j < BS; ++j)
if (evals[j] < evals[best]) best = j;
if (best != i) {
float tv = evals[i]; evals[i] = evals[best]; evals[best] = tv;
int tp = perm[i]; perm[i] = perm[best]; perm[best] = tp;
}
}
for (int i = 0; i < len; ++i) vb[start + i] = evals[i];
}
__syncthreads();
for (int t = tid; t < len * len; t += blockDim.x) {
int r = t % len;
int c = t / len;
int src = perm[c];
Qout[(start + r) + static_cast<int64_t>(start + c) * n] = Q[r * LD + src];
}
}
__global__ void dc_gather_sort_kernel(
const double* __restrict__ vals_in,
const float* __restrict__ vecs_in,
const float* __restrict__ offdiag,
const int* __restrict__ nodes,
int num_nodes, int n, int sort_width,
double* __restrict__ d_sorted,
double* __restrict__ z_sorted,
int* __restrict__ perm,
double* __restrict__ rho_out) {
int node = blockIdx.x;
int b = blockIdx.y;
int tid = threadIdx.x;
int start = nodes[3 * node + 0];
int L = nodes[3 * node + 1];
int R = nodes[3 * node + 2];
int M = L + R;
const double* vb = vals_in + static_cast<int64_t>(b) * n;
const float* Qb = vecs_in + static_cast<int64_t>(b) * n * n;
const float* eb = offdiag + static_cast<int64_t>(b) * n;
double* dout = d_sorted + static_cast<int64_t>(b) * n + start;
double* zout = z_sorted + static_cast<int64_t>(b) * n + start;
int* pout = perm + static_cast<int64_t>(b) * n + start;
extern __shared__ unsigned char raw[];
double* sd = reinterpret_cast<double*>(raw);
double* sz = sd + sort_width;
int* sp = reinterpret_cast<int*>(sz + sort_width);
double beta = static_cast<double>(eb[start + L - 1]);
double sign = beta >= 0.0 ? 1.0 : -1.0;
constexpr double inv_sqrt2 = 0.7071067811865475244008443621048490;
for (int i = tid; i < sort_width; i += blockDim.x) {
if (i < M) {
sd[i] = vb[start + i];
if (i < L) {
sz[i] = static_cast<double>(
Qb[(start + L - 1) + static_cast<int64_t>(start + i) * n]) *
inv_sqrt2;
} else {
int q = i - L;
sz[i] = sign * static_cast<double>(
Qb[(start + L) + static_cast<int64_t>(start + L + q) * n]) *
inv_sqrt2;
}
sp[i] = i;
} else {
sd[i] = CUDART_INF;
sz[i] = 0.0;
sp[i] = i;
}
}
if (tid == 0)
rho_out[static_cast<int64_t>(b) * num_nodes + node] = 2.0 * fabs(beta);
__syncthreads();
for (int k = 2; k <= sort_width; k <<= 1) {
for (int j = k >> 1; j > 0; j >>= 1) {
for (int i = tid; i < sort_width; i += blockDim.x) {
int ix = i ^ j;
if (ix > i) {
bool up = ((i & k) == 0);
if ((sd[i] > sd[ix]) == up) {
double td = sd[i]; sd[i] = sd[ix]; sd[ix] = td;
double tz = sz[i]; sz[i] = sz[ix]; sz[ix] = tz;
int tp = sp[i]; sp[i] = sp[ix]; sp[ix] = tp;
}
}
}
__syncthreads();
}
}
for (int i = tid; i < M; i += blockDim.x) {
dout[i] = sd[i];
zout[i] = sz[i];
pout[i] = sp[i];
}
}
__global__ void dc_permute_basis_kernel(
const float* __restrict__ Qin,
const int* __restrict__ perm,
const int* __restrict__ nodes,
int num_nodes, int n, int max_m,
float* __restrict__ Qsorted) {
int b = blockIdx.z;
int row_tiles = ceil_div(max_m, 16);
int node = blockIdx.y / row_tiles;
int rt = blockIdx.y - node * row_tiles;
if (node >= num_nodes) return;
int start = nodes[3 * node + 0];
int L = nodes[3 * node + 1];
int R = nodes[3 * node + 2];
int M = L + R;
int r = rt * 16 + threadIdx.y;
int c = blockIdx.x * 16 + threadIdx.x;
if (r >= M || c >= M) return;
int src = perm[static_cast<int64_t>(b) * n + start + c];
const float* Qb = Qin + static_cast<int64_t>(b) * n * n;
float* Sb = Qsorted + static_cast<int64_t>(b) * n * n;
bool same_child = (r < L) == (src < L);
Sb[(start + r) + static_cast<int64_t>(start + c) * n] =
same_child
? Qb[(start + r) + static_cast<int64_t>(start + src) * n]
: 0.0f;
}
__global__ void dc_deflate_scan_kernel(
double* __restrict__ d_sorted,
double* __restrict__ z_sorted,
const double* __restrict__ rho,
const int* __restrict__ nodes,
int num_nodes, int n,
double* __restrict__ d_active,
double* __restrict__ z_active,
int* __restrict__ active_idx,
int* __restrict__ defl_idx,
int* __restrict__ active_count,
int* __restrict__ defl_count,
int* __restrict__ rot_count,
int* __restrict__ rot_i,
int* __restrict__ rot_j,
double* __restrict__ rot_c,
double* __restrict__ rot_s) {
int node = blockIdx.x;
int b = blockIdx.y;
if (threadIdx.x != 0) return;
int start = nodes[3 * node + 0];
int M = nodes[3 * node + 1] + nodes[3 * node + 2];
int64_t base = static_cast<int64_t>(b) * n + start;
int64_t slot = static_cast<int64_t>(b) * num_nodes + node;
double* d = d_sorted + base;
double* z = z_sorted + base;
double rrho = rho[slot];
double maxd = 0.0, maxz = 0.0;
for (int i = 0; i < M; ++i) {
maxd = fmax(maxd, fabs(d[i]));
maxz = fmax(maxz, fabs(z[i]));
}
double tol = 8.0 * static_cast<double>(FLT_EPSILON) *
fmax(maxd, maxz);
tol = fmax(tol, DBL_MIN);
int K = 0, D = 0, NR = 0;
int j = 0;
while (j < M && rrho * fabs(z[j]) <= tol) {
defl_idx[base + D++] = j++;
}
if (j < M) {
int pj = j++;
while (j < M) {
int nj = j++;
if (rrho * fabs(z[nj]) <= tol) {
defl_idx[base + D++] = nj;
continue;
}
double tau = hypot(z[nj], z[pj]);
double c = z[nj] / tau;
double s = -z[pj] / tau;
double gap = d[nj] - d[pj];
if (fabs(gap * c * s) <= tol) {
rot_i[base + NR] = pj;
rot_j[base + NR] = nj;
rot_c[base + NR] = c;
rot_s[base + NR] = s;
++NR;
z[nj] = tau;
z[pj] = 0.0;
double dp = d[pj], dn = d[nj];
d[pj] = dp * c * c + dn * s * s;
d[nj] = dp * s * s + dn * c * c;
defl_idx[base + D++] = pj;
pj = nj;
} else {
active_idx[base + K] = pj;
d_active[base + K] = d[pj];
z_active[base + K] = z[pj];
++K;
pj = nj;
}
}
active_idx[base + K] = pj;
d_active[base + K] = d[pj];
z_active[base + K] = z[pj];
++K;
}
active_count[slot] = K;
defl_count[slot] = D;
rot_count[slot] = NR;
}
__global__ void dc_apply_rotations_kernel(
float* __restrict__ Qsorted,
const int* __restrict__ nodes,
int num_nodes, int n,
const int* __restrict__ rot_count,
const int* __restrict__ rot_i,
const int* __restrict__ rot_j,
const double* __restrict__ rot_c,
const double* __restrict__ rot_s) {
int node = blockIdx.x;
int b = blockIdx.y;
int tid = threadIdx.x;
int start = nodes[3 * node + 0];
int M = nodes[3 * node + 1] + nodes[3 * node + 2];
int64_t base = static_cast<int64_t>(b) * n + start;
int64_t slot = static_cast<int64_t>(b) * num_nodes + node;
float* Qb = Qsorted + static_cast<int64_t>(b) * n * n;
int nr = rot_count[slot];
for (int q = 0; q < nr; ++q) {
int i = rot_i[base + q];
int j = rot_j[base + q];
float c = static_cast<float>(rot_c[base + q]);
float s = static_cast<float>(rot_s[base + q]);
for (int r = tid; r < M; r += blockDim.x) {
int64_t xi = (start + r) + static_cast<int64_t>(start + i) * n;
int64_t xj = (start + r) + static_cast<int64_t>(start + j) * n;
float x = Qb[xi], y = Qb[xj];
Qb[xi] = c * x + s * y;
Qb[xj] = c * y - s * x;
}
__syncthreads();
}
}
__global__ void dc_compact_basis_kernel(
const float* __restrict__ Qsorted,
const int* __restrict__ nodes,
int num_nodes, int n, int max_m,
const int* __restrict__ active_count,
const int* __restrict__ active_idx,
const int* __restrict__ defl_idx,
float* __restrict__ Qbasis) {
int b = blockIdx.z;
int row_tiles = ceil_div(max_m, 16);
int node = blockIdx.y / row_tiles;
int rt = blockIdx.y - node * row_tiles;
if (node >= num_nodes) return;
int start = nodes[3 * node + 0];
int M = nodes[3 * node + 1] + nodes[3 * node + 2];
int64_t slot = static_cast<int64_t>(b) * num_nodes + node;
int K = active_count[slot];
int r = rt * 16 + threadIdx.y;
int c = blockIdx.x * 16 + threadIdx.x;
if (r >= M || c >= M) return;
int64_t base = static_cast<int64_t>(b) * n + start;
int src = (c < K) ? active_idx[base + c] : defl_idx[base + (c - K)];
const float* Sb = Qsorted + static_cast<int64_t>(b) * n * n;
float* Bb = Qbasis + static_cast<int64_t>(b) * n * n;
Bb[(start + r) + static_cast<int64_t>(start + c) * n] =
Sb[(start + r) + static_cast<int64_t>(start + src) * n];
}
__device__ __forceinline__ double warp_secular_value(
double x, const double* d, const double* z, int K, double rho) {
int lane = threadIdx.x & 31;
double part = 0.0;
double scale = 1.0;
for (int i = lane; i < K; i += 32) {
double den = d[i] - x;
double floor_den = 16.0 * DBL_EPSILON *
fmax(scale, fmax(fabs(d[i]), fabs(x)));
if (fabs(den) < floor_den)
den = copysign(floor_den, den == 0.0 ? 1.0 : den);
part += z[i] * z[i] / den;
}
part = warp_sum_double(part);
return 1.0 + rho * __shfl_sync(kFullMask, part, 0);
}
constexpr int kSecularWarps = 4;
__device__ __forceinline__ void warp_secular_value_derivative(
double x, const double* d, const double* z, int K, double rho,
double& f, double& fp) {
int lane = threadIdx.x & 31;
double part = 0.0;
double deriv = 0.0;
for (int i = lane; i < K; i += 32) {
double den = d[i] - x;
double floor_den = 16.0 * DBL_EPSILON *
fmax(1.0, fmax(fabs(d[i]), fabs(x)));
if (fabs(den) < floor_den)
den = copysign(floor_den, den == 0.0 ? 1.0 : den);
double zi2 = z[i] * z[i];
double inv = 1.0 / den;
part += zi2 * inv;
deriv += zi2 * inv * inv;
}
part = warp_sum_double(part);
deriv = warp_sum_double(deriv);
f = 1.0 + rho * __shfl_sync(kFullMask, part, 0);
fp = rho * __shfl_sync(kFullMask, deriv, 0);
}
__global__ void dc_secular_roots_fast512_kernel(
const double* __restrict__ d_active,
const double* __restrict__ z_active,
const double* __restrict__ rho,
const int* __restrict__ active_count,
const int* __restrict__ nodes,
int num_nodes, int n,
double* __restrict__ roots) {
int warp = threadIdx.x >> 5;
int lane = threadIdx.x & 31;
int j = blockIdx.x * kSecularWarps + warp;
int node = blockIdx.y;
int b = blockIdx.z;
int start = nodes[3 * node + 0];
int64_t slot = static_cast<int64_t>(b) * num_nodes + node;
int K = active_count[slot];
if (j >= K) return;
int64_t base = static_cast<int64_t>(b) * n + start;
const double* d = d_active + base;
const double* z = z_active + base;
double rrho = rho[slot];
if (K == 1) {
if (lane == 0) roots[base] = d[0] + rrho * z[0] * z[0];
return;
}
double local_max = 0.0;
double local_z2 = 0.0;
for (int i = lane; i < K; i += 32) {
local_max = fmax(local_max, fabs(d[i]));
local_z2 += z[i] * z[i];
}
local_max = warp_max_double(local_max);
local_z2 = warp_sum_double(local_z2);
double maxd = __shfl_sync(kFullMask, local_max, 0);
double z2sum = __shfl_sync(kFullMask, local_z2, 0);
double scale = 1.0 + maxd + rrho * z2sum;
double lo = 0.0, hi = 0.0;
if (lane == 0) {
if (j < K - 1) {
double gap = d[j + 1] - d[j];
double eps = fmin(0.25 * gap, 32.0 * DBL_EPSILON * scale);
eps = fmax(eps, DBL_MIN * scale);
lo = d[j] + eps;
hi = d[j + 1] - eps;
if (!(hi > lo)) {
lo = nextafter(d[j], CUDART_INF);
hi = nextafter(d[j + 1], -CUDART_INF);
}
} else {
double eps = fmax(32.0 * DBL_EPSILON * scale, DBL_MIN * scale);
lo = d[K - 1] + eps;
hi = d[K - 1] + fmax(scale, rrho * z2sum + scale);
}
}
lo = __shfl_sync(kFullMask, lo, 0);
hi = __shfl_sync(kFullMask, hi, 0);
if (j == K - 1) {
for (int grow = 0; grow < 32; ++grow) {
double fhi = warp_secular_value(hi, d, z, K, rrho);
if (fhi > 0.0 && isfinite(fhi)) break;
if (lane == 0) hi = d[K - 1] + 2.0 * (hi - d[K - 1]);
hi = __shfl_sync(kFullMask, hi, 0);
}
}
double x = 0.5 * (lo + hi);
for (int it = 0; it < 32; ++it) {
double fx, dfx;
warp_secular_value_derivative(x, d, z, K, rrho, fx, dfx);
if (lane == 0) {
if (isfinite(fx)) {
if (fx <= 0.0) lo = x;
else hi = x;
} else if (fx < 0.0) {
lo = x;
} else {
hi = x;
}
double midpoint = 0.5 * (lo + hi);
double candidate = midpoint;
if (isfinite(fx) && isfinite(dfx) && dfx > 0.0) {
double newton = x - fx / dfx;
double guard = 0.015625 * (hi - lo);
if (newton > lo + guard && newton < hi - guard)
candidate = newton;
}
x = candidate;
}
lo = __shfl_sync(kFullMask, lo, 0);
hi = __shfl_sync(kFullMask, hi, 0);
x = __shfl_sync(kFullMask, x, 0);
double width = hi - lo;
double sc = fmax(DBL_MIN, fmax(fabs(lo), fabs(hi)));
if (width <= 8.0 * DBL_EPSILON * sc || x == lo || x == hi) break;
}
if (lane == 0) roots[base + j] = x;
}
__global__ void dc_secular_roots_kernel(
const double* __restrict__ d_active,
const double* __restrict__ z_active,
const double* __restrict__ rho,
const int* __restrict__ active_count,
const int* __restrict__ nodes,
int num_nodes, int n,
double* __restrict__ roots) {
int warp = threadIdx.x >> 5;
int lane = threadIdx.x & 31;
int j = blockIdx.x * kSecularWarps + warp;
int node = blockIdx.y;
int b = blockIdx.z;
int start = nodes[3 * node + 0];
int64_t slot = static_cast<int64_t>(b) * num_nodes + node;
int K = active_count[slot];
if (j >= K) return;
int64_t base = static_cast<int64_t>(b) * n + start;
const double* d = d_active + base;
const double* z = z_active + base;
double rrho = rho[slot];
if (K == 1) {
if (lane == 0) roots[base] = d[0] + rrho * z[0] * z[0];
return;
}
double local_max = 0.0;
double local_z2 = 0.0;
for (int i = lane; i < K; i += 32) {
local_max = fmax(local_max, fabs(d[i]));
local_z2 += z[i] * z[i];
}
local_max = warp_max_double(local_max);
local_z2 = warp_sum_double(local_z2);
double maxd = __shfl_sync(kFullMask, local_max, 0);
double z2sum = __shfl_sync(kFullMask, local_z2, 0);
double scale = 1.0 + maxd + rrho * z2sum;
double lo = 0.0;
double hi = 0.0;
if (lane == 0) {
if (j < K - 1) {
double gap = d[j + 1] - d[j];
double eps = fmin(0.25 * gap, 32.0 * DBL_EPSILON * scale);
eps = fmax(eps, DBL_MIN * scale);
lo = d[j] + eps;
hi = d[j + 1] - eps;
if (!(hi > lo)) {
lo = nextafter(d[j], CUDART_INF);
hi = nextafter(d[j + 1], -CUDART_INF);
}
} else {
double eps = fmax(32.0 * DBL_EPSILON * scale, DBL_MIN * scale);
lo = d[K - 1] + eps;
hi = d[K - 1] + fmax(scale, rrho * z2sum + scale);
}
}
lo = __shfl_sync(kFullMask, lo, 0);
hi = __shfl_sync(kFullMask, hi, 0);
if (j == K - 1) {
for (int grow = 0; grow < 32; ++grow) {
double fhi = warp_secular_value(hi, d, z, K, rrho);
if (fhi > 0.0 && isfinite(fhi)) break;
if (lane == 0) hi = d[K - 1] + 2.0 * (hi - d[K - 1]);
hi = __shfl_sync(kFullMask, hi, 0);
}
}
for (int it = 0; it < 64; ++it) {
double mid = 0.5 * (lo + hi);
double fm = warp_secular_value(mid, d, z, K, rrho);
if (lane == 0) {
if (!isfinite(fm) || fm <= 0.0) lo = mid;
else hi = mid;
}
lo = __shfl_sync(kFullMask, lo, 0);
hi = __shfl_sync(kFullMask, hi, 0);
double width = hi - lo;
double midpoint = 0.5 * (lo + hi);
double local_scale = fmax(DBL_MIN, fmax(fabs(lo), fabs(hi)));
if (midpoint == lo || midpoint == hi ||
width <= 8.0 * DBL_EPSILON * local_scale) break;
}
if (lane == 0) roots[base + j] = 0.5 * (lo + hi);
}
__global__ void dc_log_zhat_kernel(
const double* __restrict__ d_active,
const double* __restrict__ z_active,
const double* __restrict__ roots,
const int* __restrict__ active_count,
const int* __restrict__ nodes,
int num_nodes, int n,
double* __restrict__ zhat) {
int i = blockIdx.x * blockDim.x + threadIdx.x;
int node = blockIdx.y;
int b = blockIdx.z;
int start = nodes[3 * node + 0];
int64_t slot = static_cast<int64_t>(b) * num_nodes + node;
int K = active_count[slot];
if (i >= K) return;
int64_t base = static_cast<int64_t>(b) * n + start;
const double* d = d_active + base;
const double* lam = roots + base;
double di = d[i];
double logabs = 0.0;
for (int j = 0; j < K; ++j) {
double q = di - lam[j];
double floor_q = 16.0 * DBL_EPSILON *
fmax(1.0, fmax(fabs(di), fabs(lam[j])));
logabs += log(fmax(fabs(q), floor_q));
}
for (int j = 0; j < K; ++j) {
if (j == i) continue;
double q = di - d[j];
double floor_q = 16.0 * DBL_EPSILON *
fmax(1.0, fmax(fabs(di), fabs(d[j])));
logabs -= log(fmax(fabs(q), floor_q));
}
double mag = exp(fmin(700.0, fmax(-700.0, 0.5 * logabs)));
zhat[base + i] = copysign(mag, z_active[base + i]);
}
__global__ void dc_scaled_zhat_fast512_kernel(
const double* __restrict__ d_active,
const double* __restrict__ z_active,
const double* __restrict__ roots,
const int* __restrict__ active_count,
const int* __restrict__ nodes,
int num_nodes, int n,
double* __restrict__ zhat) {
int i = blockIdx.x * blockDim.x + threadIdx.x;
int node = blockIdx.y;
int b = blockIdx.z;
int start = nodes[3 * node + 0];
int64_t slot = static_cast<int64_t>(b) * num_nodes + node;
int K = active_count[slot];
if (i >= K) return;
int64_t base = static_cast<int64_t>(b) * n + start;
const double* d = d_active + base;
const double* lam = roots + base;
double di = d[i];
double mant = 1.0;
int exp2 = 0;
bool fallback = false;
for (int j = 0; j < K; ++j) {
double num = fabs(di - lam[j]);
double floor_num = 16.0 * DBL_EPSILON *
fmax(1.0, fmax(fabs(di), fabs(lam[j])));
num = fmax(num, floor_num);
double factor = num;
if (j != i) {
double den = fabs(di - d[j]);
double floor_den = 16.0 * DBL_EPSILON *
fmax(1.0, fmax(fabs(di), fabs(d[j])));
factor /= fmax(den, floor_den);
}
if (!(factor > 0.0) || !isfinite(factor)) {
fallback = true;
break;
}
mant *= factor;
int e = 0;
mant = frexp(mant, &e);
exp2 += e;
if (!isfinite(mant) || exp2 > 2000 || exp2 < -2000) {
fallback = true;
break;
}
}
double mag = 0.0;
if (!fallback) {
if ((exp2 % 2) != 0) {
mant *= 2.0;
--exp2;
}
mag = ldexp(sqrt(mant), exp2 / 2);
fallback = !isfinite(mag);
}
if (fallback) {
double logabs = 0.0;
for (int j = 0; j < K; ++j) {
double q = di - lam[j];
double floor_q = 16.0 * DBL_EPSILON *
fmax(1.0, fmax(fabs(di), fabs(lam[j])));
logabs += log(fmax(fabs(q), floor_q));
}
for (int j = 0; j < K; ++j) {
if (j == i) continue;
double q = di - d[j];
double floor_q = 16.0 * DBL_EPSILON *
fmax(1.0, fmax(fabs(di), fabs(d[j])));
logabs -= log(fmax(fabs(q), floor_q));
}
mag = exp(fmin(700.0, fmax(-700.0, 0.5 * logabs)));
}
zhat[base + i] = copysign(mag, z_active[base + i]);
}
constexpr int kBuildUWarps = 8;
__global__ void dc_build_u_kernel(
const double* __restrict__ d_active,
const double* __restrict__ roots,
const double* __restrict__ zhat,
const int* __restrict__ active_count,
const int* __restrict__ nodes,
int num_nodes, int n,
float* __restrict__ U) {
int warp = threadIdx.x >> 5;
int lane = threadIdx.x & 31;
int j = blockIdx.x * kBuildUWarps + warp;
int node = blockIdx.y;
int b = blockIdx.z;
int start = nodes[3 * node + 0];
int64_t slot = static_cast<int64_t>(b) * num_nodes + node;
int K = active_count[slot];
if (j >= K) return;
int64_t base = static_cast<int64_t>(b) * n + start;
const double* d = d_active + base;
double lambda = roots[base + j];
double norm2 = 0.0;
for (int i = lane; i < K; i += 32) {
double den = d[i] - lambda;
double floor_den = 16.0 * DBL_EPSILON *
fmax(1.0, fmax(fabs(d[i]), fabs(lambda)));
if (fabs(den) < floor_den)
den = copysign(floor_den, den == 0.0 ? 1.0 : den);
double x = zhat[base + i] / den;
norm2 += x * x;
}
norm2 = warp_sum_double(norm2);
double invnorm = 0.0;
if (lane == 0) invnorm = norm2 > 0.0 ? 1.0 / sqrt(norm2) : 0.0;
invnorm = __shfl_sync(kFullMask, invnorm, 0);
float* Ub = U + static_cast<int64_t>(b) * n * n;
for (int i = lane; i < K; i += 32) {
double den = d[i] - lambda;
double floor_den = 16.0 * DBL_EPSILON *
fmax(1.0, fmax(fabs(d[i]), fabs(lambda)));
if (fabs(den) < floor_den)
den = copysign(floor_den, den == 0.0 ? 1.0 : den);
Ub[(start + i) + static_cast<int64_t>(start + j) * n] =
static_cast<float>((zhat[base + i] / den) * invnorm);
}
}
__global__ void dc_assemble_values_kernel(
const double* __restrict__ roots,
const double* __restrict__ d_sorted,
const int* __restrict__ defl_idx,
const int* __restrict__ active_count,
const int* __restrict__ defl_count,
const int* __restrict__ nodes,
int num_nodes, int n,
double* __restrict__ merged_values) {
int node = blockIdx.x;
int b = blockIdx.y;
int tid = threadIdx.x;
int start = nodes[3 * node + 0];
int M = nodes[3 * node + 1] + nodes[3 * node + 2];
int64_t slot = static_cast<int64_t>(b) * num_nodes + node;
int K = active_count[slot];
int D = defl_count[slot];
int64_t base = static_cast<int64_t>(b) * n + start;
for (int i = tid; i < K; i += blockDim.x)
merged_values[base + i] = roots[base + i];
for (int i = tid; i < D; i += blockDim.x) {
int src = defl_idx[base + i];
merged_values[base + K + i] = d_sorted[base + src];
}
for (int i = K + D + tid; i < M; i += blockDim.x)
merged_values[base + i] = CUDART_INF;
}
__device__ __forceinline__ float round_to_tf32_software(float x) {
unsigned bits = __float_as_uint(x);
unsigned exponent = bits & 0x7f800000u;
if (exponent == 0x7f800000u) return x;
unsigned lsb = (bits >> 13) & 1u;
bits += 0x00000fffu + lsb;
bits &= 0xffffe000u;
return __uint_as_float(bits);
}
__global__ void dc_pack_compensated_tf32_fused_basis_kernel(
const float* __restrict__ Qsorted,
const float* __restrict__ U,
const int* __restrict__ active_count,
const int* __restrict__ active_idx,
const int* __restrict__ defl_idx,
const int* __restrict__ nodes,
int num_nodes, int batch, int n, int M,
float* __restrict__ Apack,
float* __restrict__ Bpack) {
int64_t total = static_cast<int64_t>(batch) * num_nodes * M * M;
int64_t linear = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
if (linear >= total) return;
int64_t matrix_elems = static_cast<int64_t>(M) * M;
int merge_matrix = static_cast<int>(linear / matrix_elems);
int local = static_cast<int>(linear - static_cast<int64_t>(merge_matrix) * matrix_elems);
int r = local % M;
int c = local / M;
int b = merge_matrix / num_nodes;
int node = merge_matrix - b * num_nodes;
int start = nodes[3 * node + 0];
int64_t slot = static_cast<int64_t>(b) * num_nodes + node;
int K = active_count[slot];
int64_t base = static_cast<int64_t>(b) * n + start;
int src = (c < K) ? active_idx[base + c]
: defl_idx[base + (c - K)];
const float* Qb = Qsorted + static_cast<int64_t>(b) * n * n;
const float* Ub = U + static_cast<int64_t>(b) * n * n;
float a = Qb[(start + r) + static_cast<int64_t>(start + src) * n];
float transform = 0.0f;
if (c < K) {
if (r < K)
transform = Ub[(start + r) + static_cast<int64_t>(start + c) * n];
} else {
transform = (r == c) ? 1.0f : 0.0f;
}
constexpr float Scale = 4096.0f;
constexpr float Inv64 = 1.0f / 64.0f;
float ahi = round_to_tf32_software(a);
float alo = (a - ahi) * Scale;
float bhi = round_to_tf32_software(transform);
float blo = (transform - bhi) * Scale;
int64_t a_stride = static_cast<int64_t>(M) * (3 * M);
int64_t b_stride = static_cast<int64_t>(3 * M) * M;
float* Am = Apack + static_cast<int64_t>(merge_matrix) * a_stride;
float* Bm = Bpack + static_cast<int64_t>(merge_matrix) * b_stride;
Am[r + static_cast<int64_t>(c) * M] = ahi;
Am[r + static_cast<int64_t>(c + M) * M] = ahi * Inv64;
Am[r + static_cast<int64_t>(c + 2 * M) * M] = alo * Inv64;
Bm[r + static_cast<int64_t>(c) * (3 * M)] = bhi;
Bm[(r + M) + static_cast<int64_t>(c) * (3 * M)] = blo * Inv64;
Bm[(r + 2 * M) + static_cast<int64_t>(c) * (3 * M)] = bhi * Inv64;
}
template<bool UseDouble>
__global__ void dc_active_matmul_kernel(
const float* __restrict__ Qbasis,
const float* __restrict__ U,
const int* __restrict__ active_count,
const int* __restrict__ nodes,
int num_nodes, int n, int max_m,
float* __restrict__ Qtmp) {
constexpr int T = 16;
__shared__ float As[T][T + 1];
__shared__ float Bs[T][T + 1];
int b = blockIdx.z;
int row_tiles = ceil_div(max_m, T);
int node = blockIdx.y / row_tiles;
int rt = blockIdx.y - node * row_tiles;
if (node >= num_nodes) return;
int start = nodes[3 * node + 0];
int M = nodes[3 * node + 1] + nodes[3 * node + 2];
int K = active_count[static_cast<int64_t>(b) * num_nodes + node];
int r = rt * T + threadIdx.y;
int c = blockIdx.x * T + threadIdx.x;
const float* Ab = Qbasis + static_cast<int64_t>(b) * n * n;
const float* Bb = U + static_cast<int64_t>(b) * n * n;
float* Cb = Qtmp + static_cast<int64_t>(b) * n * n;
double accd = 0.0;
float acc = 0.0f, corr = 0.0f;
for (int k0 = 0; k0 < K; k0 += T) {
int ka = k0 + threadIdx.x;
int kb = k0 + threadIdx.y;
As[threadIdx.y][threadIdx.x] =
(r < M && ka < K) ? Ab[(start + r) + static_cast<int64_t>(start + ka) * n] : 0.0f;
Bs[threadIdx.y][threadIdx.x] =
(kb < K && c < K) ? Bb[(start + kb) + static_cast<int64_t>(start + c) * n] : 0.0f;
__syncthreads();
#pragma unroll
for (int q = 0; q < T; ++q) {
if constexpr (UseDouble) {
accd += static_cast<double>(As[threadIdx.y][q]) * Bs[q][threadIdx.x];
} else {
float prod = As[threadIdx.y][q] * Bs[q][threadIdx.x];
float y = prod - corr;
float t = acc + y;
corr = (t - acc) - y;
acc = t;
}
}
__syncthreads();
}
if (r < M && c < K)
Cb[(start + r) + static_cast<int64_t>(start + c) * n] =
UseDouble ? static_cast<float>(accd) : acc;
}
__global__ void dc_copy_deflated_kernel(
const float* __restrict__ Qbasis,
const int* __restrict__ active_count,
const int* __restrict__ nodes,
int num_nodes, int n, int max_m,
float* __restrict__ Qtmp) {
int b = blockIdx.z;
int row_tiles = ceil_div(max_m, 16);
int node = blockIdx.y / row_tiles;
int rt = blockIdx.y - node * row_tiles;
if (node >= num_nodes) return;
int start = nodes[3 * node + 0];
int M = nodes[3 * node + 1] + nodes[3 * node + 2];
int K = active_count[static_cast<int64_t>(b) * num_nodes + node];
int r = rt * 16 + threadIdx.y;
int c = blockIdx.x * 16 + threadIdx.x;
if (r >= M || c < K || c >= M) return;
const float* Sb = Qbasis + static_cast<int64_t>(b) * n * n;
float* Tb = Qtmp + static_cast<int64_t>(b) * n * n;
Tb[(start + r) + static_cast<int64_t>(start + c) * n] =
Sb[(start + r) + static_cast<int64_t>(start + c) * n];
}
__global__ void dc_sort_merged_values_kernel(
const double* __restrict__ merged_values,
const int* __restrict__ nodes,
int num_nodes, int n, int sort_width,
double* __restrict__ vals_out,
int* __restrict__ final_perm) {
int node = blockIdx.x;
int b = blockIdx.y;
int tid = threadIdx.x;
int start = nodes[3 * node + 0];
int M = nodes[3 * node + 1] + nodes[3 * node + 2];
int64_t base = static_cast<int64_t>(b) * n + start;
extern __shared__ unsigned char raw[];
double* sv = reinterpret_cast<double*>(raw);
int* sp = reinterpret_cast<int*>(sv + sort_width);
for (int i = tid; i < sort_width; i += blockDim.x) {
sv[i] = i < M ? merged_values[base + i] : CUDART_INF;
sp[i] = i;
}
__syncthreads();
for (int k = 2; k <= sort_width; k <<= 1) {
for (int j = k >> 1; j > 0; j >>= 1) {
for (int i = tid; i < sort_width; i += blockDim.x) {
int ix = i ^ j;
if (ix > i) {
bool up = ((i & k) == 0);
if ((sv[i] > sv[ix]) == up) {
double tv = sv[i]; sv[i] = sv[ix]; sv[ix] = tv;
int tp = sp[i]; sp[i] = sp[ix]; sp[ix] = tp;
}
}
}
__syncthreads();
}
}
for (int i = tid; i < M; i += blockDim.x) {
vals_out[base + i] = sv[i];
final_perm[base + i] = sp[i];
}
}
__global__ void dc_permute_final_kernel(
const float* __restrict__ Qtmp,
const int* __restrict__ final_perm,
const int* __restrict__ nodes,
int num_nodes, int n, int max_m,
float* __restrict__ Qout) {
int b = blockIdx.z;
int row_tiles = ceil_div(max_m, 16);
int node = blockIdx.y / row_tiles;
int rt = blockIdx.y - node * row_tiles;
if (node >= num_nodes) return;
int start = nodes[3 * node + 0];
int M = nodes[3 * node + 1] + nodes[3 * node + 2];
int r = rt * 16 + threadIdx.y;
int c = blockIdx.x * 16 + threadIdx.x;
if (r >= M || c >= M) return;
int src = final_perm[static_cast<int64_t>(b) * n + start + c];
const float* Tb = Qtmp + static_cast<int64_t>(b) * n * n;
float* Ob = Qout + static_cast<int64_t>(b) * n * n;
Ob[(start + r) + static_cast<int64_t>(start + c) * n] =
Tb[(start + r) + static_cast<int64_t>(start + src) * n];
}
__global__ void copy_dc_q_to_padded_kernel(
const float* __restrict__ Qin,
float* __restrict__ Qout,
int n, int ld) {
int b = blockIdx.z;
int i = blockIdx.y * blockDim.y + threadIdx.y;
int j = blockIdx.x * blockDim.x + threadIdx.x;
if (i >= n || j >= n) return;
Qout[static_cast<int64_t>(b) * ld * n + i + static_cast<int64_t>(j) * ld] =
Qin[static_cast<int64_t>(b) * n * n + i + static_cast<int64_t>(j) * n];
}
__global__ void copy_double_values_to_float_kernel(
const double* __restrict__ in,
float* __restrict__ out,
int total) {
int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i < total) out[i] = static_cast<float>(in[i]);
}
__device__ __forceinline__ void split_reflector_fp16(
float x, cutlass::half_t& high, cutlass::half_t& low) {
high = static_cast<cutlass::half_t>(x);
low = static_cast<cutlass::half_t>(
(x - static_cast<float>(high)) * 2048.0f);
}
__global__ void materialize_pack_panel_v_kernel(
const float* __restrict__ A,
cutlass::half_t* __restrict__ Vrc,
cutlass::half_t* __restrict__ Vcc,
int n, int ld, int panel_b, int k, int bcols,
int row0, int m, int mld) {
constexpr float kInv64 = 1.0f / 64.0f;
int b = blockIdx.y;
int64_t t = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
int64_t total = static_cast<int64_t>(mld) * panel_b;
if (t >= total) return;
int rr = static_cast<int>(t % mld);
int s = static_cast<int>(t / mld);
int r = row0 + rr;
float value = 0.0f;
if (s < bcols && rr < m) {
int j = k + s;
if (r == j + 1) value = 1.0f;
else if (r > j + 1 && r < n) {
const float* Ab = A + static_cast<int64_t>(b) * ld * n;
value = Ab[r + static_cast<int64_t>(j) * ld];
}
}
cutlass::half_t hi, lo;
split_reflector_fp16(value, hi, lo);
int64_t rc = static_cast<int64_t>(b) * panel_b * (3 * mld) +
static_cast<int64_t>(s) * (3 * mld) + rr;
Vrc[rc] = hi;
Vrc[rc + mld] = static_cast<cutlass::half_t>(
static_cast<float>(hi) * kInv64);
Vrc[rc + 2 * mld] = static_cast<cutlass::half_t>(
static_cast<float>(lo) * kInv64);
int64_t cc = static_cast<int64_t>(b) * mld * (3 * panel_b) + rr +
static_cast<int64_t>(s) * mld;
Vcc[cc] = hi;
Vcc[cc + static_cast<int64_t>(panel_b) * mld] =
static_cast<cutlass::half_t>(static_cast<float>(hi) * kInv64);
Vcc[cc + static_cast<int64_t>(2 * panel_b) * mld] =
static_cast<cutlass::half_t>(static_cast<float>(lo) * kInv64);
}
__global__ void pack_rhs_fp16_kernel(
const float* __restrict__ input,
cutlass::half_t* __restrict__ packed,
int batch, int n, int ld, int row0, int m, int mld) {
constexpr float kInv32 = 1.0f / 32.0f;
int64_t i = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
int64_t per_batch = static_cast<int64_t>(mld) * n;
int64_t total = static_cast<int64_t>(batch) * per_batch;
if (i >= total) return;
int b = static_cast<int>(i / per_batch);
int64_t local = i - static_cast<int64_t>(b) * per_batch;
int r = static_cast<int>(local % mld);
int c = static_cast<int>(local / mld);
float value = 0.0f;
if (r < m) {
const float* Qb = input + static_cast<int64_t>(b) * ld * n;
value = Qb[(row0 + r) + static_cast<int64_t>(c) * ld];
}
cutlass::half_t hi, lo;
split_reflector_fp16(value, hi, lo);
int64_t base = static_cast<int64_t>(b) * (3 * mld) * n +
static_cast<int64_t>(c) * (3 * mld) + r;
packed[base] = hi;
packed[base + mld] = static_cast<cutlass::half_t>(
static_cast<float>(lo) * kInv32);
packed[base + 2 * mld] = static_cast<cutlass::half_t>(
static_cast<float>(hi) * kInv32);
}
__global__ void small_t_times_y_pack_kernel(
const float* __restrict__ T,
const float* __restrict__ Y,
cutlass::half_t* __restrict__ Y2pack,
int n, int panel_b, int num_panels, int panel_id, int bcols) {
constexpr float kInv32 = 1.0f / 32.0f;
int b = blockIdx.y;
int t = blockIdx.x * blockDim.x + threadIdx.x;
int total = panel_b * n;
if (t >= total) return;
int r = t % panel_b;
int c = t / panel_b;
const float* Tb = T + (static_cast<int64_t>(b) * num_panels + panel_id) *
panel_b * panel_b;
const float* Yb = Y + static_cast<int64_t>(b) * panel_b * n;
double acc = 0.0;
if (r < bcols) {
for (int s = r; s < bcols; ++s)
acc += static_cast<double>(Tb[r + static_cast<int64_t>(s) * panel_b]) *
Yb[s + static_cast<int64_t>(c) * panel_b];
}
cutlass::half_t hi, lo;
split_reflector_fp16(static_cast<float>(acc), hi, lo);
int64_t base = static_cast<int64_t>(b) * (3 * panel_b) * n +
static_cast<int64_t>(c) * (3 * panel_b) + r;
Y2pack[base] = hi;
Y2pack[base + panel_b] = static_cast<cutlass::half_t>(
static_cast<float>(lo) * kInv32);
Y2pack[base + 2 * panel_b] = static_cast<cutlass::half_t>(
static_cast<float>(hi) * kInv32);
}
__global__ void materialize_mixed_panel_v_kernel(
const float* __restrict__ A,
float* __restrict__ Vt32,
cutlass::half_t* __restrict__ V16,
int n, int ld, int panel_b, int k, int bcols,
int row0, int m) {
int b = blockIdx.y;
int64_t t = static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x;
int64_t total = static_cast<int64_t>(m) * panel_b;
if (t >= total) return;
int rr = static_cast<int>(t % m);
int s = static_cast<int>(t / m);
int r = row0 + rr;
float value = 0.0f;
if (s < bcols) {
int j = k + s;
if (r == j + 1) value = 1.0f;
else if (r > j + 1 && r < n) {
const float* Ab = A + static_cast<int64_t>(b) * ld * n;
value = Ab[r + static_cast<int64_t>(j) * ld];
}
}
Vt32[static_cast<int64_t>(b) * panel_b * ld +
static_cast<int64_t>(s) * ld + rr] = value;
V16[static_cast<int64_t>(b) * ld * panel_b +
rr + static_cast<int64_t>(s) * ld] =
static_cast<cutlass::half_t>(value);
}
__global__ void cast_panel_fp32_to_fp16_kernel(
const float* __restrict__ input,
cutlass::half_t* __restrict__ output,
int total) {
int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i < total) output[i] = static_cast<cutlass::half_t>(input[i]);
}
__global__ void build_newton_polar_factor_kernel(
float* __restrict__ gram_factor, int n, int ld) {
int b = blockIdx.z;
int i = blockIdx.y * blockDim.y + threadIdx.y;
int j = blockIdx.x * blockDim.x + threadIdx.x;
if (i >= n || j >= n || i > j) return;
float* Gb = gram_factor + static_cast<int64_t>(b) * ld * n;
float gij = Gb[i + static_cast<int64_t>(j) * ld];
float gji = Gb[j + static_cast<int64_t>(i) * ld];
float gsym = 0.5f * (gij + gji);
float value = (i == j ? 1.5f : 0.0f) - 0.5f * gsym;
Gb[i + static_cast<int64_t>(j) * ld] = value;
if (i != j) Gb[j + static_cast<int64_t>(i) * ld] = value;
}
__global__ void normalize_columns_kernel(float* Q, int n, int ld) {
int b = blockIdx.y;
int col = blockIdx.x;
int tid = threadIdx.x;
float* Qb = Q + static_cast<int64_t>(b) * ld * n;
__shared__ double red[8];
__shared__ double inv;
double norm2 = 0.0;
for (int r = tid; r < n; r += blockDim.x) {
double q = Qb[r + static_cast<int64_t>(col) * ld];
norm2 += q * q;
}
norm2 = block_sum_double<kThreads>(norm2, red);
if (tid == 0) inv = norm2 > 0.0 ? 1.0 / sqrt(norm2) : 0.0;
__syncthreads();
for (int r = tid; r < n; r += blockDim.x)
Qb[r + static_cast<int64_t>(col) * ld] =
static_cast<float>(Qb[r + static_cast<int64_t>(col) * ld] * inv);
}
torch::Tensor make_cuda_int_tensor(
const std::vector<int>& values,
const c10::Device& device) {
auto cpu = torch::empty({static_cast<int64_t>(values.size())},
torch::TensorOptions().dtype(torch::kInt32).device(torch::kCPU));
std::memcpy(cpu.data_ptr<int>(), values.data(), values.size() * sizeof(int));
return cpu.to(torch::TensorOptions().device(device).dtype(torch::kInt32),
/*non_blocking=*/false, /*copy=*/true);
}
void blocked_tridiagonalize(
torch::Tensor A,
torch::Tensor T,
torch::Tensor tau,
torch::Tensor diag,
torch::Tensor offdiag,
torch::Tensor V,
torch::Tensor W,
torch::Tensor x,
torch::Tensor y,
torch::Tensor coeff_v,
torch::Tensor coeff_w,
torch::Tensor U2,
torch::Tensor Z2,
torch::Tensor cutlass_workspace,
int batch, int n, int ld) {
int num_panels = ceil_div(n - 1, kPanel);
float* Ap = A.data_ptr<float>();
float* Tp = T.data_ptr<float>();
float* taup = tau.data_ptr<float>();
float* dp = diag.data_ptr<float>();
float* ep = offdiag.data_ptr<float>();
float* Vp = V.data_ptr<float>();
float* Wp = W.data_ptr<float>();
float* U2p = U2.data_ptr<float>();
float* Z2p = Z2.data_ptr<float>();
float* xp = x.data_ptr<float>();
float* yp = y.data_ptr<float>();
float* cvp = coeff_v.data_ptr<float>();
float* cwp = coeff_w.data_ptr<float>();
void* ws = cutlass_workspace.data_ptr();
size_t ws_bytes = cutlass_workspace.numel();
const bool use_persistent_small = (n <= 384);
const bool use_chunked_512 = (n == 512);
const bool use_large_1024 = (n > 512 && n <= 1024);
const bool use_large_2048 = (n > 1024 && n <= 2048);
const bool use_large_panel = use_large_1024 || use_large_2048;
const bool use_any_persistent = use_persistent_small || use_large_panel;
size_t persistent_smem_bytes = 0;
size_t large_smem_bytes = 0;
if (use_persistent_small) {
persistent_smem_bytes =
static_cast<size_t>(2) * ld * kPanel * sizeof(float) +
static_cast<size_t>(2) * ld * sizeof(float) +
static_cast<size_t>(kPanel) * kPanel * sizeof(float) +
static_cast<size_t>(2) * kPanel * sizeof(float) +
static_cast<size_t>(16) * sizeof(double);
static const bool persistent_smem_configured = []() {
constexpr int max_ld = 384;
constexpr int max_bytes =
2 * max_ld * kPanel * static_cast<int>(sizeof(float)) +
2 * max_ld * static_cast<int>(sizeof(float)) +
kPanel * kPanel * static_cast<int>(sizeof(float)) +
2 * kPanel * static_cast<int>(sizeof(float)) +
16 * static_cast<int>(sizeof(double));
C10_CUDA_CHECK(cudaFuncSetAttribute(
persistent_panel_small_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
max_bytes));
return true;
}();
(void)persistent_smem_configured;
}
if (use_large_panel) {
constexpr int rows_per_cta = 256;
int capacity = use_large_1024 ? 1024 : 2048;
large_smem_bytes =
static_cast<size_t>(2) * rows_per_cta * kPanel * sizeof(float) +
static_cast<size_t>(2) * rows_per_cta * sizeof(float) +
static_cast<size_t>(capacity) * sizeof(float) +
static_cast<size_t>(kPanel) * kPanel * sizeof(float) +
static_cast<size_t>(4) * kPanel * sizeof(float) +
static_cast<size_t>(8 + 8 + 16) * sizeof(double);
if (use_large_1024) {
static const bool large1024_smem_configured = []() {
constexpr int rows = 256;
constexpr int capacity = 1024;
constexpr int max_bytes =
2 * rows * kPanel * static_cast<int>(sizeof(float)) +
2 * rows * static_cast<int>(sizeof(float)) +
capacity * static_cast<int>(sizeof(float)) +
kPanel * kPanel * static_cast<int>(sizeof(float)) +
4 * kPanel * static_cast<int>(sizeof(float)) +
(8 + 8 + 16) * static_cast<int>(sizeof(double));
C10_CUDA_CHECK(cudaFuncSetAttribute(
persistent_panel_large1024_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
max_bytes));
return true;
}();
(void)large1024_smem_configured;
} else {
static const bool large2048_smem_configured = []() {
constexpr int rows = 256;
constexpr int capacity = 2048;
constexpr int max_bytes =
2 * rows * kPanel * static_cast<int>(sizeof(float)) +
2 * rows * static_cast<int>(sizeof(float)) +
capacity * static_cast<int>(sizeof(float)) +
kPanel * kPanel * static_cast<int>(sizeof(float)) +
4 * kPanel * static_cast<int>(sizeof(float)) +
(8 + 8 + 16) * static_cast<int>(sizeof(double));
C10_CUDA_CHECK(cudaFuncSetAttribute(
persistent_panel_large2048_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
max_bytes));
return true;
}();
(void)large2048_smem_configured;
}
}
if (use_chunked_512) {
static const bool panel512_prefix_smem_configured = []() {
constexpr int prefix_bytes =
2 * 512 * kPanel512PrefixCols * static_cast<int>(sizeof(float));
C10_CUDA_CHECK(cudaFuncSetAttribute(
panel_chunk8_512_kernel<true>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
prefix_bytes));
return true;
}();
(void)panel512_prefix_smem_configured;
}
if (!use_any_persistent) {
C10_CUDA_CHECK(cudaMemset(Tp, 0, T.numel() * sizeof(float)));
C10_CUDA_CHECK(cudaMemset(ep, 0, offdiag.numel() * sizeof(float)));
}
for (int panel_id = 0, k = 0; k < n - 1; k += kPanel, ++panel_id) {
int bcols = std::min(kPanel, n - k - 1);
if (use_persistent_small) {
persistent_panel_small_kernel<<<batch, kThreads, persistent_smem_bytes>>>(
Ap, Tp, taup, dp, ep, Vp, Wp, U2p, Z2p,
n, ld, kPanel, num_panels, panel_id, k, bcols);
} else if (use_large_1024) {
persistent_panel_large1024_kernel<<<
batch * 4, kThreads, large_smem_bytes>>>(
Ap, Tp, taup, dp, ep, Vp, Wp, U2p, Z2p,
n, ld, kPanel, num_panels, panel_id, k, bcols);
} else if (use_large_2048) {
persistent_panel_large2048_kernel<<<
batch * 8, kThreads, large_smem_bytes>>>(
Ap, Tp, taup, dp, ep, Vp, Wp, U2p, Z2p,
n, ld, kPanel, num_panels, panel_id, k, bcols);
} else if (use_chunked_512) {
constexpr int prefix_bytes =
2 * 512 * kPanel512PrefixCols * static_cast<int>(sizeof(float));
for (int s0 = 0; s0 < bcols; s0 += kPanel512Chunk) {
int chunk_cols = std::min(kPanel512Chunk, bcols - s0);
if (s0 == kPanel512PrefixCols) {
panel_chunk8_512_kernel<true><<<
batch, kPanel512Threads, prefix_bytes>>>(
Ap, Tp, taup, dp, ep, Vp, Wp,
n, ld, kPanel, num_panels, panel_id, k, s0, chunk_cols);
} else {
panel_chunk8_512_kernel<false><<<batch, kPanel512Threads>>>(
Ap, Tp, taup, dp, ep, Vp, Wp,
n, ld, kPanel, num_panels, panel_id, k, s0, chunk_cols);
}
}
} else {
for (int s = 0; s < bcols; ++s) {
int j = k + s;
int start = j + 1;
int rows = n - start;
dim3 grid_col(ceil_div(rows, kThreads), batch);
panel_correct_column_kernel<<<grid_col, kThreads>>>(
Ap, Vp, Wp, xp, dp, n, ld, kPanel, j, s);
panel_householder_kernel<<<batch, kThreads>>>(
Ap, xp, Vp, taup, ep, n, ld, kPanel, j, s);
int row_ctas = ceil_div(rows, kThreads);
bool use_redundant_fused_coefficients = row_ctas <= 8;
if (use_redundant_fused_coefficients) {
panel_coefficients_matvec_fused_kernel<<<grid_col, kThreads>>>(
Ap, Vp, Wp, cvp, cwp, yp, n, ld, kPanel, start, s);
} else {
if (s > 0) {
panel_coefficients_kernel<<<batch, kThreads>>>(
Vp, Wp, cvp, cwp, n, ld, kPanel, start, s);
}
panel_matvec_kernel<<<grid_col, kThreads>>>(
Ap, Vp, Wp, cvp, cwp, yp, n, ld, kPanel, start, s);
}
panel_finalize_w_kernel<<<batch, kThreads>>>(
Vp, yp, taup, cvp, Wp, Tp,
n, ld, kPanel, num_panels, panel_id, j, s);
}
}
int r0 = k + bcols;
if (!use_any_persistent && r0 == n - 1) {
final_diagonal_kernel<<<batch, 1>>>(
Ap, Vp, Wp, dp, n, ld, kPanel, n - 1, bcols);
}
int m = n - r0;
if (r0 < n - 1) {
float* A22 = Ap + r0 + static_cast<int64_t>(r0) * ld;
if ((m & 3) == 0) {
if (!use_any_persistent) {
dim3 pack_grid(ceil_div(ld * (2 * bcols), kThreads), batch);
pack_rank2_factors_kernel<<<pack_grid, kThreads>>>(
Vp, Wp, U2p, Z2p, ld, kPanel, bcols);
}
const float* Utail = U2p + r0;
const float* Ztail = Z2p + r0;
bool use_tma_cluster =
n >= 2048 && m >= 256 &&
GemmCRTmaCluster2::can_run(
m, m, 2 * bcols, batch,
Utail, ld, static_cast<int64_t>(ld) * (2 * kPanel),
Ztail, ld, static_cast<int64_t>(ld) * (2 * kPanel),
A22, ld, static_cast<int64_t>(ld) * n,
A22, ld, static_cast<int64_t>(ld) * n,
-1.0f, 1.0f);
if (use_tma_cluster) {
GemmCRTmaCluster2::run(
m, m, 2 * bcols, batch,
Utail, ld, static_cast<int64_t>(ld) * (2 * kPanel),
Ztail, ld, static_cast<int64_t>(ld) * (2 * kPanel),
A22, ld, static_cast<int64_t>(ld) * n,
A22, ld, static_cast<int64_t>(ld) * n,
-1.0f, 1.0f, ws, ws_bytes);
} else {
GemmCR::run(
m, m, 2 * bcols, batch,
Utail, ld, static_cast<int64_t>(ld) * (2 * kPanel),
Ztail, ld, static_cast<int64_t>(ld) * (2 * kPanel),
A22, ld, static_cast<int64_t>(ld) * n,
A22, ld, static_cast<int64_t>(ld) * n,
-1.0f, 1.0f, ws, ws_bytes);
}
} else {
dim3 block(16, 16);
dim3 grid(ceil_div(m, 16), ceil_div(m, 16), batch);
trailing_rank2_update_kernel<<<grid, block>>>(
Ap, Vp, Wp, n, ld, kPanel, r0, bcols);
}
dim3 block(16, 16);
dim3 grid(ceil_div(m, 16), ceil_div(m, 16), batch);
symmetrize_submatrix_kernel<<<grid, block>>>(Ap, n, ld, r0);
}
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
}
std::tuple<torch::Tensor, torch::Tensor> tridiagonal_divide_conquer(
torch::Tensor diag,
torch::Tensor offdiag,
int batch, int n) {
auto fopts = diag.options();
auto dopts = diag.options().dtype(torch::kFloat64);
auto iopts = diag.options().dtype(torch::kInt32);
int requested_leaves = ceil_div(n, kLeaf);
int num_leaves = next_pow2(requested_leaves);
std::vector<int> leaf_values;
leaf_values.reserve(2 * num_leaves);
std::vector<std::pair<int, int>> segments;
segments.reserve(num_leaves);
int q = n / num_leaves;
int rem = n % num_leaves;
int start = 0;
for (int i = 0; i < num_leaves; ++i) {
int len = q + (i < rem ? 1 : 0);
TORCH_CHECK(len > 0 && len <= kLeaf, "invalid D&C leaf partition");
leaf_values.push_back(start);
leaf_values.push_back(len);
segments.push_back({start, len});
start += len;
}
auto leaf_meta = make_cuda_int_tensor(leaf_values, diag.device());
auto vals_a = torch::empty({batch, n}, dopts);
auto vals_b = torch::empty({batch, n}, dopts);
auto vec_a = torch::empty({batch, n, n}, fopts);
auto vec_b = torch::empty({batch, n, n}, fopts);
auto qsorted = torch::empty({batch, n, n}, fopts);
auto uwork = torch::empty({batch, n, n}, fopts);
auto qtmp = torch::empty({batch, n, n}, fopts);
torch::Tensor dc_apack;
torch::Tensor dc_bpack;
torch::Tensor dc_gemm_workspace;
size_t dc_gemm_workspace_bytes = 0;
if (n == 512 || n == 1024) {
std::vector<int> tensor_merge_sizes =
n == 512 ? std::vector<int>{256, 512}
: std::vector<int>{256};
int64_t max_pack_elems_per_batch = 1;
for (int M : tensor_merge_sizes) {
int nodes_at_level = n / M;
int64_t level_pack_elems =
static_cast<int64_t>(nodes_at_level) * M * (3 * M);
max_pack_elems_per_batch =
std::max(max_pack_elems_per_batch, level_pack_elems);
}
dc_apack = torch::empty(
{batch, max_pack_elems_per_batch}, fopts);
dc_bpack = torch::empty(
{batch, max_pack_elems_per_batch}, fopts);
size_t max_ws = 1;
for (int M : tensor_merge_sizes) {
int nodes_at_level = n / M;
int64_t a_stride = static_cast<int64_t>(M) * (3 * M);
int64_t b_stride = static_cast<int64_t>(3 * M) * M;
size_t need = GemmCC::workspace_size(
M, M, 3 * M, batch,
dc_apack.data_ptr<float>(), M,
static_cast<int64_t>(nodes_at_level) * a_stride,
dc_bpack.data_ptr<float>(), 3 * M,
static_cast<int64_t>(nodes_at_level) * b_stride,
qtmp.data_ptr<float>(), n,
static_cast<int64_t>(n) * n,
qtmp.data_ptr<float>(), n,
static_cast<int64_t>(n) * n);
max_ws = std::max(max_ws, need);
}
dc_gemm_workspace_bytes = max_ws;
dc_gemm_workspace = torch::empty(
{static_cast<int64_t>(max_ws)}, fopts.dtype(torch::kUInt8));
}
auto d_sorted = torch::empty({batch, n}, dopts);
auto z_sorted = torch::empty({batch, n}, dopts);
auto d_active = torch::empty({batch, n}, dopts);
auto z_active = torch::empty({batch, n}, dopts);
auto roots = torch::empty({batch, n}, dopts);
auto zhat = torch::empty({batch, n}, dopts);
auto merged_values = torch::empty({batch, n}, dopts);
int max_nodes = num_leaves;
auto rho = torch::empty({batch, max_nodes}, dopts);
auto active_count = torch::empty({batch, max_nodes}, iopts);
auto defl_count = torch::empty({batch, max_nodes}, iopts);
auto rot_count = torch::empty({batch, max_nodes}, iopts);
auto perm = torch::empty({batch, n}, iopts);
auto active_idx = torch::empty({batch, n}, iopts);
auto defl_idx = torch::empty({batch, n}, iopts);
auto rot_i = torch::empty({batch, n}, iopts);
auto rot_j = torch::empty({batch, n}, iopts);
auto rot_c = torch::empty({batch, n}, dopts);
auto rot_s = torch::empty({batch, n}, dopts);
dim3 leaf_grid(num_leaves, batch);
dc_leaf_jacobi_kernel<<<leaf_grid, kLeafJacobiThreads>>>(
diag.data_ptr<float>(), offdiag.data_ptr<float>(),
leaf_meta.data_ptr<int>(), num_leaves, n,
vals_a.data_ptr<double>(), vec_a.data_ptr<float>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
torch::Tensor vals_in = vals_a, vals_out = vals_b;
torch::Tensor vec_in = vec_a, vec_out = vec_b;
std::vector<torch::Tensor> metadata_keepalive;
metadata_keepalive.push_back(leaf_meta);
while (segments.size() > 1) {
int num_nodes = static_cast<int>(segments.size() / 2);
std::vector<int> node_values;
node_values.reserve(3 * num_nodes);
std::vector<std::pair<int, int>> next_segments;
next_segments.reserve(num_nodes);
int max_m = 0;
for (int i = 0; i < num_nodes; ++i) {
auto left = segments[2 * i];
auto right = segments[2 * i + 1];
TORCH_CHECK(left.first + left.second == right.first,
"noncontiguous D&C node");
node_values.push_back(left.first);
node_values.push_back(left.second);
node_values.push_back(right.second);
int m = left.second + right.second;
max_m = std::max(max_m, m);
next_segments.push_back({left.first, m});
}
auto nodes = make_cuda_int_tensor(node_values, diag.device());
metadata_keepalive.push_back(nodes);
int sort_width = next_pow2(max_m);
size_t gather_smem = static_cast<size_t>(sort_width) *
(2 * sizeof(double) + sizeof(int));
size_t sort_smem = static_cast<size_t>(sort_width) *
(sizeof(double) + sizeof(int));
C10_CUDA_CHECK(cudaFuncSetAttribute(
dc_gather_sort_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
static_cast<int>(gather_smem)));
C10_CUDA_CHECK(cudaFuncSetAttribute(
dc_sort_merged_values_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
static_cast<int>(sort_smem)));
dc_gather_sort_kernel<<<dim3(num_nodes, batch), kThreads,
gather_smem>>>(
vals_in.data_ptr<double>(), vec_in.data_ptr<float>(),
offdiag.data_ptr<float>(), nodes.data_ptr<int>(),
num_nodes, n, sort_width,
d_sorted.data_ptr<double>(), z_sorted.data_ptr<double>(),
perm.data_ptr<int>(), rho.data_ptr<double>());
dim3 block2(16, 16);
int row_tiles = ceil_div(max_m, 16);
dim3 grid2(ceil_div(max_m, 16), num_nodes * row_tiles, batch);
dc_permute_basis_kernel<<<grid2, block2>>>(
vec_in.data_ptr<float>(), perm.data_ptr<int>(), nodes.data_ptr<int>(),
num_nodes, n, max_m, qsorted.data_ptr<float>());
dc_deflate_scan_kernel<<<dim3(num_nodes, batch), 1>>>(
d_sorted.data_ptr<double>(), z_sorted.data_ptr<double>(),
rho.data_ptr<double>(), nodes.data_ptr<int>(), num_nodes, n,
d_active.data_ptr<double>(), z_active.data_ptr<double>(),
active_idx.data_ptr<int>(), defl_idx.data_ptr<int>(),
active_count.data_ptr<int>(), defl_count.data_ptr<int>(),
rot_count.data_ptr<int>(), rot_i.data_ptr<int>(), rot_j.data_ptr<int>(),
rot_c.data_ptr<double>(), rot_s.data_ptr<double>());
dc_apply_rotations_kernel<<<dim3(num_nodes, batch), kThreads>>>(
qsorted.data_ptr<float>(), nodes.data_ptr<int>(), num_nodes, n,
rot_count.data_ptr<int>(), rot_i.data_ptr<int>(), rot_j.data_ptr<int>(),
rot_c.data_ptr<double>(), rot_s.data_ptr<double>());
bool use_tensor_merge =
(n == 512 && max_m >= 256) ||
(n == 1024 && max_m == 256);
if (!use_tensor_merge) {
dc_compact_basis_kernel<<<grid2, block2>>>(
qsorted.data_ptr<float>(), nodes.data_ptr<int>(), num_nodes, n, max_m,
active_count.data_ptr<int>(), active_idx.data_ptr<int>(),
defl_idx.data_ptr<int>(), vec_out.data_ptr<float>());
}
if (n == 512) {
dc_secular_roots_fast512_kernel<<<
dim3(ceil_div(max_m, kSecularWarps), num_nodes, batch),
32 * kSecularWarps>>>(
d_active.data_ptr<double>(), z_active.data_ptr<double>(),
rho.data_ptr<double>(), active_count.data_ptr<int>(),
nodes.data_ptr<int>(), num_nodes, n, roots.data_ptr<double>());
dc_scaled_zhat_fast512_kernel<<<
dim3(ceil_div(max_m, 128), num_nodes, batch), 128>>>(
d_active.data_ptr<double>(), z_active.data_ptr<double>(),
roots.data_ptr<double>(), active_count.data_ptr<int>(),
nodes.data_ptr<int>(), num_nodes, n, zhat.data_ptr<double>());
} else {
dc_secular_roots_kernel<<<
dim3(ceil_div(max_m, kSecularWarps), num_nodes, batch),
32 * kSecularWarps>>>(
d_active.data_ptr<double>(), z_active.data_ptr<double>(),
rho.data_ptr<double>(), active_count.data_ptr<int>(),
nodes.data_ptr<int>(), num_nodes, n, roots.data_ptr<double>());
dc_log_zhat_kernel<<<
dim3(ceil_div(max_m, 128), num_nodes, batch), 128>>>(
d_active.data_ptr<double>(), z_active.data_ptr<double>(),
roots.data_ptr<double>(), active_count.data_ptr<int>(),
nodes.data_ptr<int>(), num_nodes, n, zhat.data_ptr<double>());
}
dc_build_u_kernel<<<
dim3(ceil_div(max_m, kBuildUWarps), num_nodes, batch),
32 * kBuildUWarps>>>(
d_active.data_ptr<double>(), roots.data_ptr<double>(),
zhat.data_ptr<double>(), active_count.data_ptr<int>(),
nodes.data_ptr<int>(), num_nodes, n, uwork.data_ptr<float>());
dc_assemble_values_kernel<<<dim3(num_nodes, batch), kThreads>>>(
roots.data_ptr<double>(), d_sorted.data_ptr<double>(),
defl_idx.data_ptr<int>(), active_count.data_ptr<int>(),
defl_count.data_ptr<int>(), nodes.data_ptr<int>(), num_nodes, n,
merged_values.data_ptr<double>());
if (use_tensor_merge) {
int merge_batch = batch * num_nodes;
int64_t merge_elems = static_cast<int64_t>(merge_batch) * max_m * max_m;
dc_pack_compensated_tf32_fused_basis_kernel<<<
ceil_div(static_cast<int>(merge_elems), kThreads), kThreads>>>(
qsorted.data_ptr<float>(), uwork.data_ptr<float>(),
active_count.data_ptr<int>(), active_idx.data_ptr<int>(),
defl_idx.data_ptr<int>(), nodes.data_ptr<int>(),
num_nodes, batch, n, max_m,
dc_apack.data_ptr<float>(), dc_bpack.data_ptr<float>());
int64_t a_stride = static_cast<int64_t>(max_m) * (3 * max_m);
int64_t b_stride = static_cast<int64_t>(3 * max_m) * max_m;
for (int node = 0; node < num_nodes; ++node) {
int node_start = node_values[3 * node + 0];
const float* A_node = dc_apack.data_ptr<float>() +
static_cast<int64_t>(node) * a_stride;
const float* B_node = dc_bpack.data_ptr<float>() +
static_cast<int64_t>(node) * b_stride;
float* Q_node = qtmp.data_ptr<float>() + node_start +
static_cast<int64_t>(node_start) * n;
GemmCC::run(
max_m, max_m, 3 * max_m, batch,
A_node, max_m,
static_cast<int64_t>(num_nodes) * a_stride,
B_node, 3 * max_m,
static_cast<int64_t>(num_nodes) * b_stride,
Q_node, n, static_cast<int64_t>(n) * n,
Q_node, n, static_cast<int64_t>(n) * n,
1.0f, 0.0f,
dc_gemm_workspace.data_ptr(), dc_gemm_workspace_bytes);
}
} else {
if (max_m <= 128) {
dc_active_matmul_kernel<true><<<grid2, block2>>>(
vec_out.data_ptr<float>(), uwork.data_ptr<float>(),
active_count.data_ptr<int>(), nodes.data_ptr<int>(),
num_nodes, n, max_m, qtmp.data_ptr<float>());
} else {
dc_active_matmul_kernel<false><<<grid2, block2>>>(
vec_out.data_ptr<float>(), uwork.data_ptr<float>(),
active_count.data_ptr<int>(), nodes.data_ptr<int>(),
num_nodes, n, max_m, qtmp.data_ptr<float>());
}
dc_copy_deflated_kernel<<<grid2, block2>>>(
vec_out.data_ptr<float>(), active_count.data_ptr<int>(),
nodes.data_ptr<int>(), num_nodes, n, max_m, qtmp.data_ptr<float>());
}
dc_sort_merged_values_kernel<<<dim3(num_nodes, batch), kThreads,
sort_smem>>>(
merged_values.data_ptr<double>(), nodes.data_ptr<int>(),
num_nodes, n, sort_width, vals_out.data_ptr<double>(),
perm.data_ptr<int>());
dc_permute_final_kernel<<<grid2, block2>>>(
qtmp.data_ptr<float>(), perm.data_ptr<int>(), nodes.data_ptr<int>(),
num_nodes, n, max_m, vec_out.data_ptr<float>());
C10_CUDA_KERNEL_LAUNCH_CHECK();
std::swap(vals_in, vals_out);
std::swap(vec_in, vec_out);
segments.swap(next_segments);
}
return {vals_in, vec_in};
}
void blocked_backtransform_mixed512(
torch::Tensor A,
torch::Tensor T,
torch::Tensor Q,
torch::Tensor Y,
torch::Tensor Z,
torch::Tensor Vt32,
torch::Tensor V16,
torch::Tensor Z16,
torch::Tensor polar_factor,
torch::Tensor cutlass_workspace,
int batch, int n, int ld) {
TORCH_CHECK(n == 512, "mixed compact-WY path is specialized for n=512");
int num_panels = ceil_div(n - 1, kPanel);
float* Ap = A.data_ptr<float>();
float* Tp = T.data_ptr<float>();
float* Qp = Q.data_ptr<float>();
float* Yp = Y.data_ptr<float>();
float* Zp = Z.data_ptr<float>();
float* Vt32p = Vt32.data_ptr<float>();
auto* V16p = reinterpret_cast<cutlass::half_t*>(V16.data_ptr());
auto* Z16p = reinterpret_cast<cutlass::half_t*>(Z16.data_ptr());
float* Sp = polar_factor.data_ptr<float>();
void* ws = cutlass_workspace.data_ptr();
size_t ws_bytes = cutlass_workspace.numel();
constexpr int64_t panel_elems = static_cast<int64_t>(kPanel) * kPanel;
constexpr int64_t yn_elems = static_cast<int64_t>(kPanel) * 512;
for (int panel_id = num_panels - 1; panel_id >= 0; --panel_id) {
int k = panel_id * kPanel;
int bcols = std::min(kPanel, n - k - 1);
int row0 = k;
int m = n - row0;
dim3 vgrid(ceil_div(m * kPanel, kThreads), batch);
materialize_mixed_panel_v_kernel<<<vgrid, kThreads>>>(
Ap, Vt32p, V16p, n, ld, kPanel, k, bcols, row0, m);
float* Qpanel = Qp + row0;
FloatGemmRC::run(
kPanel, n, m, batch,
Vt32p, ld, static_cast<int64_t>(kPanel) * ld,
Qpanel, ld, static_cast<int64_t>(ld) * n,
Yp, kPanel, static_cast<int64_t>(kPanel) * n,
Yp, kPanel, static_cast<int64_t>(kPanel) * n,
1.0f, 0.0f, ws, ws_bytes);
const float* Tpanel = Tp + static_cast<int64_t>(panel_id) * panel_elems;
GemmCC::run(
kPanel, n, kPanel, batch,
Tpanel, kPanel, static_cast<int64_t>(num_panels) * panel_elems,
Yp, kPanel, static_cast<int64_t>(kPanel) * n,
Zp, kPanel, static_cast<int64_t>(kPanel) * n,
Zp, kPanel, static_cast<int64_t>(kPanel) * n,
1.0f, 0.0f, ws, ws_bytes);
int z_total = batch * static_cast<int>(yn_elems);
cast_panel_fp32_to_fp16_kernel<<<ceil_div(z_total, kThreads), kThreads>>>(
Zp, Z16p, z_total);
HalfGemmCC::run(
m, n, kPanel, batch,
V16p, ld, static_cast<int64_t>(ld) * kPanel,
Z16p, kPanel, static_cast<int64_t>(kPanel) * n,
Qpanel, ld, static_cast<int64_t>(ld) * n,
Qpanel, ld, static_cast<int64_t>(ld) * n,
-1.0f, 1.0f, ws, ws_bytes);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
FloatGemmRC::run(
n, n, n, batch,
Qp, ld, static_cast<int64_t>(ld) * n,
Qp, ld, static_cast<int64_t>(ld) * n,
Sp, ld, static_cast<int64_t>(ld) * n,
Sp, ld, static_cast<int64_t>(ld) * n,
1.0f, 0.0f, ws, ws_bytes);
dim3 polar_block(16, 16);
dim3 polar_grid(ceil_div(n, 16), ceil_div(n, 16), batch);
build_newton_polar_factor_kernel<<<polar_grid, polar_block>>>(Sp, n, ld);
GemmCC::run(
n, n, n, batch,
Qp, ld, static_cast<int64_t>(ld) * n,
Sp, ld, static_cast<int64_t>(ld) * n,
Ap, ld, static_cast<int64_t>(ld) * n,
Ap, ld, static_cast<int64_t>(ld) * n,
1.0f, 0.0f, ws, ws_bytes);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
void blocked_backtransform(
torch::Tensor A,
torch::Tensor T,
torch::Tensor Q,
torch::Tensor Y,
torch::Tensor Vrc,
torch::Tensor Vcc,
torch::Tensor Qpack,
torch::Tensor Y2pack,
torch::Tensor cutlass_workspace,
int batch, int n, int ld) {
int num_panels = ceil_div(n - 1, kPanel);
float* Ap = A.data_ptr<float>();
float* Tp = T.data_ptr<float>();
float* Qp = Q.data_ptr<float>();
float* Yp = Y.data_ptr<float>();
auto* Vrcp = reinterpret_cast<cutlass::half_t*>(Vrc.data_ptr());
auto* Vccp = reinterpret_cast<cutlass::half_t*>(Vcc.data_ptr());
auto* Qpackp = reinterpret_cast<cutlass::half_t*>(Qpack.data_ptr());
auto* Y2packp = reinterpret_cast<cutlass::half_t*>(Y2pack.data_ptr());
void* ws = cutlass_workspace.data_ptr();
size_t ws_bytes = cutlass_workspace.numel();
for (int panel_id = num_panels - 1; panel_id >= 0; --panel_id) {
int k = panel_id * kPanel;
int bcols = std::min(kPanel, n - k - 1);
int row0 = k;
int m = n - row0;
int m_gemm = round_up(m, 16);
int mld = round_up(m_gemm, 8);
dim3 vgrid(ceil_div(mld * kPanel, kThreads), batch);
materialize_pack_panel_v_kernel<<<vgrid, kThreads>>>(
Ap, Vrcp, Vccp, n, ld, kPanel, k, bcols, row0, m, mld);
int64_t q_total = static_cast<int64_t>(batch) * mld * n;
pack_rhs_fp16_kernel<<<
static_cast<int>((q_total + kThreads - 1) / kThreads),
kThreads>>>(
Qp, Qpackp, batch, n, ld, row0, m, mld);
HalfGemmRC::run(
kPanel, n, 3 * mld, batch,
Vrcp, 3 * mld, static_cast<int64_t>(3 * mld) * kPanel,
Qpackp, 3 * mld, static_cast<int64_t>(3 * mld) * n,
Yp, kPanel, static_cast<int64_t>(kPanel) * n,
Yp, kPanel, static_cast<int64_t>(kPanel) * n,
1.0f, 0.0f, ws, ws_bytes);
int panel_total = kPanel * n;
small_t_times_y_pack_kernel<<<
dim3(ceil_div(panel_total, kThreads), batch), kThreads>>>(
Tp, Yp, Y2packp, n, kPanel, num_panels, panel_id, bcols);
float* Qpanel = Qp + row0;
bool use_tma_cluster =
n >= 2048 && m_gemm >= 256 &&
HalfGemmCCTmaCluster2::can_run(
m_gemm, n, 3 * kPanel, batch,
Vccp, mld, static_cast<int64_t>(mld) * (3 * kPanel),
Y2packp, 3 * kPanel, static_cast<int64_t>(3 * kPanel) * n,
Qpanel, ld, static_cast<int64_t>(ld) * n,
Qpanel, ld, static_cast<int64_t>(ld) * n,
-1.0f, 1.0f);
if (use_tma_cluster) {
HalfGemmCCTmaCluster2::run(
m_gemm, n, 3 * kPanel, batch,
Vccp, mld, static_cast<int64_t>(mld) * (3 * kPanel),
Y2packp, 3 * kPanel, static_cast<int64_t>(3 * kPanel) * n,
Qpanel, ld, static_cast<int64_t>(ld) * n,
Qpanel, ld, static_cast<int64_t>(ld) * n,
-1.0f, 1.0f, ws, ws_bytes);
} else {
HalfGemmCC::run(
m_gemm, n, 3 * kPanel, batch,
Vccp, mld, static_cast<int64_t>(mld) * (3 * kPanel),
Y2packp, 3 * kPanel, static_cast<int64_t>(3 * kPanel) * n,
Qpanel, ld, static_cast<int64_t>(ld) * n,
Qpanel, ld, static_cast<int64_t>(ld) * n,
-1.0f, 1.0f, ws, ws_bytes);
}
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
normalize_columns_kernel<<<dim3(n, batch), kThreads>>>(
Qp, n, ld);
C10_CUDA_KERNEL_LAUNCH_CHECK();
}
template<int FixedN = 0>
__global__ __launch_bounds__(kDenseJacobiThreads)
void dense_jacobi32_kernel(
const float* __restrict__ input,
int runtime_n,
float* __restrict__ evals_out,
float* __restrict__ Qout) {
constexpr int BS = 32;
constexpr int LD = 33;
constexpr int Pairs = 16;
constexpr int Steps = 31;
constexpr int Sweeps = FixedN == 32 ? 7 : 24;
constexpr float JacobiEps = 1.0e-20f;
int n = FixedN == 0 ? runtime_n : FixedN;
int b = blockIdx.x;
int tid = threadIdx.x;
const float* Ab = input + static_cast<int64_t>(b) * n * n;
float* Lb = evals_out + static_cast<int64_t>(b) * n;
float* Ob = Qout + static_cast<int64_t>(b) * n * n;
__shared__ float A0[BS * LD], A1[BS * LD], Q[BS * LD];
__shared__ float cs[Pairs], ss[Pairs], evals[BS], bound;
__shared__ int mate[BS], pair_id[BS], perm[BS];
__shared__ uint16_t upper_coord[BS * (BS + 1) / 2];
if (tid < BS) {
int row = tid;
int base = row * (2 * BS - row + 1) / 2;
for (int col = row; col < BS; ++col)
upper_coord[base + col - row] =
static_cast<uint16_t>((row << 5) | col);
}
if (tid == 0) {
float mx = 0.0f;
if (n < BS) {
for (int j = 0; j < n; ++j)
for (int i = 0; i < n; ++i)
mx = fmaxf(mx, fabsf(0.5f * (Ab[i * n + j] + Ab[j * n + i])));
}
bound = n < BS ? fmaxf(1.0f, mx * n) : 0.0f;
}
__syncthreads();
for (int t = tid; t < BS * BS; t += blockDim.x) {
int r = t / BS, c = t - r * BS;
float a = 0.0f;
if (r < n && c < n) a = 0.5f * (Ab[r * n + c] + Ab[c * n + r]);
else if (r == c) a = bound * (4.0f + r - n);
A0[r * LD + c] = a;
A1[r * LD + c] = 0.0f;
Q[r * LD + c] = r == c ? 1.0f : 0.0f;
}
if (tid < BS) { mate[tid]=0; pair_id[tid]=0; perm[tid]=tid; }
__syncthreads();
const unsigned upper0 = upper_coord[tid];
const unsigned upper1 = tid < 16 ? upper_coord[tid + kDenseJacobiThreads] : 0u;
float* cur=A0; float* nxt=A1;
for (int sweep=0; sweep<Sweeps; ++sweep) {
for (int step=0; step<Steps; ++step) {
if (tid<Pairs) {
int aa=jacobi_order32(step,tid);
int bb=jacobi_order32(step,BS-1-tid);
int p=min(aa,bb), q=max(aa,bb);
float app=cur[p*LD+p], aqq=cur[q*LD+q], apq=cur[p*LD+q];
float c=1.0f,s=0.0f;
if constexpr (FixedN == 32) {
if (fabsf(apq)>JacobiEps) {
float tau=(aqq-app)/(2.0f*apq);
float denom=fabsf(tau)+sqrtf(fmaf(tau,tau,1.0f));
float tt=(tau>=0.0f?1.0f:-1.0f)/denom;
c=rsqrtf(fmaf(tt,tt,1.0f)); s=tt*c;
}
} else {
float scale=fmaxf(1.0f,fmaxf(fabsf(app),fabsf(aqq)));
if (fabsf(apq)>8.0f*FLT_EPSILON*scale) {
double tj=(static_cast<double>(aqq)-app)/(2.0*apq);
double tt=copysign(1.0/(fabs(tj)+hypot(tj,1.0)),tj);
double cc=1.0/sqrt(1.0+tt*tt);
c=static_cast<float>(cc); s=static_cast<float>(tt*cc);
}
}
cs[tid]=c; ss[tid]=s; mate[p]=q; mate[q]=p;
pair_id[p]=tid; pair_id[q]=tid;
}
__syncthreads();
jacobi_update_packed32(
cur, nxt, mate, pair_id, cs, ss, upper0);
if (tid < 16)
jacobi_update_packed32(
cur, nxt, mate, pair_id, cs, ss, upper1);
for (int t=tid;t<BS*Pairs;t+=blockDim.x) {
int r=t/Pairs,pair=t-r*Pairs;
int aa=jacobi_order32(step,pair),bb=jacobi_order32(step,BS-1-pair);
int p=min(aa,bb),q=max(aa,bb);
float c=cs[pair],s=ss[pair],x=Q[r*LD+p],y=Q[r*LD+q];
Q[r*LD+p]=c*x-s*y; Q[r*LD+q]=s*x+c*y;
}
__syncthreads(); float* tmp=cur;cur=nxt;nxt=tmp;
}
}
if(tid==0){
for(int i=0;i<BS;++i){evals[i]=cur[i*LD+i];perm[i]=i;}
for(int i=0;i<BS-1;++i){int best=i;for(int j=i+1;j<BS;++j)
if(evals[j]<evals[best])best=j;
if(best!=i){float tv=evals[i];evals[i]=evals[best];evals[best]=tv;
int tp=perm[i];perm[i]=perm[best];perm[best]=tp;}}
for(int i=0;i<n;++i)Lb[i]=evals[i];
}
__syncthreads();
for(int t=tid;t<n*n;t+=blockDim.x){int r=t%n,c=t/n;
Ob[r*n+c]=Q[r*LD+perm[c]];}
}
std::tuple<torch::Tensor, torch::Tensor> eigh_cuda(torch::Tensor input) {
TORCH_CHECK(input.is_cuda(), "input must be CUDA");
TORCH_CHECK(input.scalar_type() == torch::kFloat32, "input must be float32");
TORCH_CHECK(input.dim() == 3, "input must have shape [batch,n,n]");
TORCH_CHECK(input.size(1) == input.size(2), "input matrices must be square");
TORCH_CHECK(input.size(1) >= 1 && input.size(1) <= 4096,
"supported range is 1 <= n <= 4096");
c10::cuda::CUDAGuard guard(input.device());
input = input.contiguous();
int batch = static_cast<int>(input.size(0));
int n = static_cast<int>(input.size(1));
int ld = round_up(n, 32);
auto fopts = input.options();
auto L = torch::empty({batch, n}, fopts);
if (n <= 32) {
auto Qout = torch::empty({batch, n, n}, fopts);
if (n == 32) {
dense_jacobi32_kernel<32><<<batch, kDenseJacobiThreads>>>(
input.data_ptr<float>(), n, L.data_ptr<float>(), Qout.data_ptr<float>());
} else {
dense_jacobi32_kernel<0><<<batch, kDenseJacobiThreads>>>(
input.data_ptr<float>(), n, L.data_ptr<float>(), Qout.data_ptr<float>());
}
C10_CUDA_KERNEL_LAUNCH_CHECK();
return {Qout, L};
}
auto A = torch::empty({batch, ld, n}, fopts);
C10_CUDA_CHECK(cudaDeviceSynchronize());
dim3 copy_block(16, 16);
dim3 copy_grid(ceil_div(n, 16), ceil_div(n, 16), batch);
copy_symmetrize_to_colmajor_kernel<<<copy_grid, copy_block>>>(
input.data_ptr<float>(), A.data_ptr<float>(), n, ld);
auto Qinternal = torch::empty({batch, ld, n}, fopts);
torch::Tensor matrix_scale_back;
if (n == 1024) {
matrix_scale_back = torch::empty(
{batch}, fopts.dtype(torch::kFloat64));
normalize_matrix_pow2_kernel<<<batch, kThreads>>>(
A.data_ptr<float>(), matrix_scale_back.data_ptr<double>(), n, ld);
}
int num_panels = ceil_div(n - 1, kPanel);
auto T = torch::empty({batch, num_panels, kPanel, kPanel}, fopts);
auto tau = torch::empty({batch, n}, fopts);
auto diag = torch::empty({batch, n}, fopts);
auto offdiag = torch::empty({batch, n}, fopts);
auto V = torch::empty({batch, ld, kPanel}, fopts);
auto W = torch::empty({batch, ld, kPanel}, fopts);
auto x = torch::empty({batch, ld}, fopts);
auto y = torch::empty({batch, ld}, fopts);
auto coeff_v = torch::empty({batch, kPanel}, fopts);
auto coeff_w = torch::empty({batch, kPanel}, fopts);
auto U2 = torch::empty({batch, ld, 2 * kPanel}, fopts);
auto Z2 = torch::empty({batch, ld, 2 * kPanel}, fopts);
auto Y = torch::empty({batch, kPanel, n}, fopts);
auto hopts = fopts.dtype(torch::kFloat16);
torch::Tensor Vrc, Vcc, Qpack, Y2pack;
torch::Tensor Vt32_mix, V16_mix, Z32_mix, Z16_mix, PolarS_mix;
if (n == 512) {
Vt32_mix = torch::empty({batch, kPanel, ld}, fopts);
V16_mix = torch::empty({batch, ld, kPanel}, hopts);
Z32_mix = torch::empty({batch, kPanel, n}, fopts);
Z16_mix = torch::empty({batch, kPanel, n}, hopts);
PolarS_mix = torch::empty({batch, ld, n}, fopts);
} else {
Vrc = torch::empty({batch, kPanel, 3 * ld}, hopts);
Vcc = torch::empty({batch, ld, 3 * kPanel}, hopts);
Qpack = torch::empty({batch, 3 * ld, n}, hopts);
Y2pack = torch::empty({batch, 3 * kPanel, n}, hopts);
}
int bmax = std::min(kPanel, n - 1);
int mtrail = std::max(1, n - bmax);
float* Ap = A.data_ptr<float>();
float* Vp = V.data_ptr<float>();
float* Wp = W.data_ptr<float>();
float* U2p = U2.data_ptr<float>();
float* Z2p = Z2.data_ptr<float>();
float* Qp = Qinternal.data_ptr<float>();
float* Yp = Y.data_ptr<float>();
size_t ws0 = GemmCR::workspace_size(
mtrail, mtrail, 2 * bmax, batch,
U2p + bmax, ld, static_cast<int64_t>(ld) * (2 * kPanel),
Z2p + bmax, ld, static_cast<int64_t>(ld) * (2 * kPanel),
Ap + bmax + static_cast<int64_t>(bmax) * ld,
ld, static_cast<int64_t>(ld) * n,
Ap + bmax + static_cast<int64_t>(bmax) * ld,
ld, static_cast<int64_t>(ld) * n);
int bt_m = std::max(1, n);
int bt_m_gemm = round_up(bt_m, 16);
int bt_mld = round_up(bt_m_gemm, 8);
size_t ws1 = 0, ws2 = 0, ws4 = 0, ws5 = 0, ws6 = 0;
if (n == 512) {
float* Vt32p = Vt32_mix.data_ptr<float>();
auto* V16p = reinterpret_cast<cutlass::half_t*>(V16_mix.data_ptr());
float* Z32p = Z32_mix.data_ptr<float>();
auto* Z16p = reinterpret_cast<cutlass::half_t*>(Z16_mix.data_ptr());
ws1 = FloatGemmRC::workspace_size(
kPanel, n, n, batch,
Vt32p, ld, static_cast<int64_t>(kPanel) * ld,
Qp, ld, static_cast<int64_t>(ld) * n,
Yp, kPanel, static_cast<int64_t>(kPanel) * n,
Yp, kPanel, static_cast<int64_t>(kPanel) * n);
ws2 = GemmCC::workspace_size(
kPanel, n, kPanel, batch,
T.data_ptr<float>(), kPanel,
static_cast<int64_t>(num_panels) * kPanel * kPanel,
Yp, kPanel, static_cast<int64_t>(kPanel) * n,
Z32p, kPanel, static_cast<int64_t>(kPanel) * n,
Z32p, kPanel, static_cast<int64_t>(kPanel) * n);
ws4 = HalfGemmCC::workspace_size(
n, n, kPanel, batch,
V16p, ld, static_cast<int64_t>(ld) * kPanel,
Z16p, kPanel, static_cast<int64_t>(kPanel) * n,
Qp, ld, static_cast<int64_t>(ld) * n,
Qp, ld, static_cast<int64_t>(ld) * n);
float* Sp = PolarS_mix.data_ptr<float>();
ws5 = FloatGemmRC::workspace_size(
n, n, n, batch,
Qp, ld, static_cast<int64_t>(ld) * n,
Qp, ld, static_cast<int64_t>(ld) * n,
Sp, ld, static_cast<int64_t>(ld) * n,
Sp, ld, static_cast<int64_t>(ld) * n);
ws6 = GemmCC::workspace_size(
n, n, n, batch,
Qp, ld, static_cast<int64_t>(ld) * n,
Sp, ld, static_cast<int64_t>(ld) * n,
Ap, ld, static_cast<int64_t>(ld) * n,
Ap, ld, static_cast<int64_t>(ld) * n);
} else {
auto* Vrcp = reinterpret_cast<cutlass::half_t*>(Vrc.data_ptr());
auto* Vccp = reinterpret_cast<cutlass::half_t*>(Vcc.data_ptr());
auto* Qpackp = reinterpret_cast<cutlass::half_t*>(Qpack.data_ptr());
auto* Y2packp = reinterpret_cast<cutlass::half_t*>(Y2pack.data_ptr());
ws1 = HalfGemmRC::workspace_size(
kPanel, n, 3 * bt_mld, batch,
Vrcp, 3 * bt_mld, static_cast<int64_t>(3 * bt_mld) * kPanel,
Qpackp, 3 * bt_mld, static_cast<int64_t>(3 * bt_mld) * n,
Yp, kPanel, static_cast<int64_t>(kPanel) * n,
Yp, kPanel, static_cast<int64_t>(kPanel) * n);
ws2 = HalfGemmCC::workspace_size(
bt_m_gemm, n, 3 * kPanel, batch,
Vccp, bt_mld, static_cast<int64_t>(bt_mld) * (3 * kPanel),
Y2packp, 3 * kPanel, static_cast<int64_t>(3 * kPanel) * n,
Qp, ld, static_cast<int64_t>(ld) * n,
Qp, ld, static_cast<int64_t>(ld) * n);
ws4 = HalfGemmCCTmaCluster2::workspace_size(
bt_m_gemm, n, 3 * kPanel, batch,
Vccp, bt_mld, static_cast<int64_t>(bt_mld) * (3 * kPanel),
Y2packp, 3 * kPanel, static_cast<int64_t>(3 * kPanel) * n,
Qp, ld, static_cast<int64_t>(ld) * n,
Qp, ld, static_cast<int64_t>(ld) * n);
}
size_t ws3 = GemmCRTmaCluster2::workspace_size(
mtrail, mtrail, 2 * bmax, batch,
U2p + bmax, ld, static_cast<int64_t>(ld) * (2 * kPanel),
Z2p + bmax, ld, static_cast<int64_t>(ld) * (2 * kPanel),
Ap + bmax + static_cast<int64_t>(bmax) * ld,
ld, static_cast<int64_t>(ld) * n,
Ap + bmax + static_cast<int64_t>(bmax) * ld,
ld, static_cast<int64_t>(ld) * n);
size_t ws_bytes = std::max({size_t{1}, ws0, ws1, ws2, ws3, ws4, ws5, ws6});
auto cutlass_workspace = torch::empty(
{static_cast<int64_t>(ws_bytes)}, fopts.dtype(torch::kUInt8));
blocked_tridiagonalize(
A, T, tau, diag, offdiag, V, W, x, y, coeff_v, coeff_w, U2, Z2,
cutlass_workspace, batch, n, ld);
auto dc_result = tridiagonal_divide_conquer(diag, offdiag, batch, n);
auto vals_double = std::get<0>(dc_result);
auto Qdc = std::get<1>(dc_result);
copy_dc_q_to_padded_kernel<<<copy_grid, copy_block>>>(
Qdc.data_ptr<float>(), Qinternal.data_ptr<float>(), n, ld);
if (n == 512) {
blocked_backtransform_mixed512(
A, T, Qinternal, Y, Z32_mix,
Vt32_mix, V16_mix, Z16_mix, PolarS_mix, cutlass_workspace,
batch, n, ld);
Qinternal = A;
} else {
blocked_backtransform(
A, T, Qinternal, Y,
Vrc, Vcc, Qpack, Y2pack, cutlass_workspace,
batch, n, ld);
}
copy_double_values_to_float_kernel<<<
ceil_div(batch * n, kThreads), kThreads>>>(
vals_double.data_ptr<double>(), L.data_ptr<float>(), batch * n);
if (n == 1024) {
dim3 value_grid(ceil_div(n, kThreads), batch);
rescale_eigenvalues_kernel<<<value_grid, kThreads>>>(
L.data_ptr<float>(), matrix_scale_back.data_ptr<double>(), n);
}
C10_CUDA_KERNEL_LAUNCH_CHECK();
C10_CUDA_CHECK(cudaDeviceSynchronize());
auto Qview = Qinternal.as_strided(
{static_cast<int64_t>(batch), static_cast<int64_t>(n), static_cast<int64_t>(n)},
{static_cast<int64_t>(ld) * n, static_cast<int64_t>(1), static_cast<int64_t>(ld)});
return {Qview, L};
}
} // namespace b200_eigh
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("eigh", &b200_eigh::eigh_cuda,
"Blocked Compact-WY symmetric eigensolver for SM100");
}
"""
_B200_EXT = None
def _resolve_cutlass_include_paths() -> list[str]:
candidates: list[Path] = []
for env_key in ("CUTLASS_PATH", "CUTLASS_ROOT"):
env_val = os.environ.get(env_key)
if env_val:
candidates.append(Path(env_val))
here = Path(__file__).resolve().parent
candidates.extend([
here / "cutlass",
here.parent / "cutlass",
Path("/workspace/cutlass"),
Path("/opt/cutlass"),
Path("/usr/local/cutlass"),
Path.home() / "cutlass",
Path.home() / "CUTLASS",
])
include_paths: list[str] = []
seen: set[str] = set()
for root in candidates:
for inc in (root / "include", root):
marker = inc / "cutlass"
if marker.is_dir():
resolved = str(inc.resolve())
if resolved not in seen:
seen.add(resolved)
include_paths.append(resolved)
if not include_paths:
raise RuntimeError(
"CUTLASS headers not found. Set CUTLASS_PATH to a CUTLASS checkout "
"that contains include/cutlass."
)
return include_paths
def _load_eigh_extension():
global _B200_EXT
if _B200_EXT is None:
_B200_EXT = load_inline(
name="b200_eigh_inline_v2_31_pow2_scale_1024",
cpp_sources="",
cuda_sources=CUDA_SOURCE,
functions=None,
with_cuda=True,
extra_cflags=["-O3", "-std=c++20"],
extra_cuda_cflags=[
"-O3",
"-std=c++20",
"-lineinfo",
"-gencode=arch=compute_100a,code=sm_100a",
"--expt-relaxed-constexpr",
"--expt-extended-lambda",
],
extra_include_paths=_resolve_cutlass_include_paths(),
verbose=False,
)
return _B200_EXT
def custom_kernel(data: input_t) -> output_t:
if (not data.is_cuda) or data.dtype != torch.float32 or data.ndim != 3 or data.shape[-1] != data.shape[-2]:
values, vectors = torch.linalg.eigh(data)
return vectors, values
ext = _load_eigh_extension()
q, l = ext.eigh(data.contiguous())
return q, lscrolls · 4630 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