Skip to content
KernelIndex
Search⌘K

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
NVIDIA B200
42.4ms
#86 of 286
2026-07-14

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.

clusterclass ClusterShape_ = Shape<_1, _1, _1>,
fused-epilogueusing CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
persistent-kernelvoid persistent_panel_small_kernel(
shared-memorydouble 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, l
scrolls · 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