Skip to content
KernelIndex
Search⌘K

submission 899876

leanyoshi · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-899876?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
985.7µs
#115 of 337
2026-07-23

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:1ff203f9f5b46d764251648bd8c993a0fa413644f195b8726c6140cc9291bb40
license declaredunknown
license concludedunknown
authorsleanyoshi
imported2026-08-26

Techniques

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

fp8const __nv_fp8_e4m3* primary,
mmanvcuda::wmma::fragment<nvcuda::wmma::accumulator, 16, 16, 16, float>
shared-memory__shared__ float tile[kMatrixSize][kSharedStride];
vector-width = float4const float4 packed = *reinterpret_cast<const float4*>(

Kernel source

submission.py3999 lines
#!POPCORN leaderboard cholesky
#!POPCORN gpu B200

import torch
import triton
import triton.language as tl
from torch.utils.cpp_extension import load_inline

from task import input_t, output_t


_CHOLESKY64_CPP = r"""
#include <torch/extension.h>

torch::Tensor cholesky64_block32_hacc4safe_cuda(torch::Tensor input);
torch::Tensor cholesky32_warp_cuda(torch::Tensor input);
torch::Tensor cholesky128_wmma_schur_cuda(torch::Tensor input);
torch::Tensor cholesky128_wmma_schur_exact_cuda(torch::Tensor input);
torch::Tensor cholesky128_block32_warp16_lowerio_anybatch_cuda(
    torch::Tensor input);
torch::Tensor cholesky256_batch64_register_panel64_cuda(torch::Tensor input);
torch::Tensor cholesky256_combined_batched_cuda(torch::Tensor input);
torch::Tensor cholesky512_batch16_register_panel64_cuda(torch::Tensor input);
torch::Tensor cholesky512_combined_batched_cuda(torch::Tensor input);
torch::Tensor cholesky512_highbatch_blocked_cuda(torch::Tensor input);
torch::Tensor cholesky1024_batch4_blocked_cuda(torch::Tensor input);
torch::Tensor cholesky1024_batch60_blocked_cuda(torch::Tensor input);
torch::Tensor cholesky2048_batch2_blocked_cuda(torch::Tensor input);
torch::Tensor cholesky2048_batch8_blocked_cuda(torch::Tensor input);
torch::Tensor cholesky4096_batch2_truebatched_fast16bf_cuda(
    torch::Tensor input);
torch::Tensor cholesky_large_explicit_bf16_cuda(torch::Tensor input);
torch::Tensor cholesky_large_combined_bf16x9_cuda(torch::Tensor input);
"""


_CHOLESKY64_CUDA = r"""
#include <torch/extension.h>

#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <c10/cuda/CUDAGuard.h>

#include <cuda.h>
#include <cuda_bf16.h>
#include <cuda_fp8.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <cublasLt.h>
#include <cublas_v2.h>
#include <cusolverDn.h>
#include <algorithm>
#include <cstdint>
#include <limits>
#include <memory>
#include <mma.h>
#include <mutex>
#include <vector>

namespace {

constexpr int kMatrixSize = 64;
constexpr int kSharedStride = 65;
constexpr int kMatrixElements = kMatrixSize * kMatrixSize;
constexpr int kBlockSize = 32;
constexpr int kWarpsPerBlock = 4;
constexpr int kThreads = 32 * kWarpsPerBlock;
constexpr int kSchurHalfTerms = 14;
constexpr int kSchurHalfStride = 16;

// Accumulate a product of two binary16 inputs directly into binary32.  On
// sm_100 this maps to the scalar mixed-precision FMA without explicit unpack
// instructions or binary16 rounding of the product.
__device__ __forceinline__ float n64_fhfma(
    unsigned short a, unsigned short b, float accumulator) {
  float result;
  asm("{.reg .b16 ha, hb;\n\t"
      "mov.b16 ha, %1;\n\t"
      "mov.b16 hb, %2;\n\t"
      "fma.rn.f32.f16 %0, ha, hb, %3;}\n"
      : "=f"(result)
      : "h"(a), "h"(b), "f"(accumulator));
  return result;
}

__device__ __forceinline__ float n64_fhfma2_acc(
    __half2 a, __half2 b, float accumulator) {
  const unsigned int a_bits = *reinterpret_cast<const unsigned int*>(&a);
  const unsigned int b_bits = *reinterpret_cast<const unsigned int*>(&b);
  accumulator = n64_fhfma(
      static_cast<unsigned short>(a_bits),
      static_cast<unsigned short>(b_bits), accumulator);
  return n64_fhfma(
      static_cast<unsigned short>(a_bits >> 16),
      static_cast<unsigned short>(b_bits >> 16), accumulator);
}

constexpr int kMatrixSize32 = 32;
constexpr int kMatrixElements32 = kMatrixSize32 * kMatrixSize32;
constexpr int kBlockSize32 = 16;
constexpr int kIoStride32 = 33;
constexpr int kIoTileElements32 = kMatrixSize32 * kIoStride32;
constexpr int kWarpsPerBlock32 = 4;
constexpr int kThreads32 = 32 * kWarpsPerBlock32;
constexpr int kSharedBytes32 =
    kWarpsPerBlock32 * kIoTileElements32 * sizeof(float);
static_assert(kSharedBytes32 == 16896,
              "n=32 register-stage I/O shared-memory size changed");

constexpr int kMatrixSize128 = 128;
constexpr int kSharedStride128 = 129;
// WMMA accumulator loads of float require a leading dimension divisible by
// four.  Padding by four also keeps every 16-column tile origin aligned.
constexpr int kWmmaSharedStride128 = 132;
constexpr int kMatrixElements128 = kMatrixSize128 * kMatrixSize128;
constexpr int kWarpsPerBlock128 = 16;
constexpr int kThreads128 = 32 * kWarpsPerBlock128;
constexpr int kFloatSharedBytes128 =
    kMatrixSize128 * kSharedStride128 * static_cast<int>(sizeof(float));
constexpr int kWmmaFloatSharedBytes128 =
    kMatrixSize128 * kWmmaSharedStride128 * static_cast<int>(sizeof(float));
constexpr int kHalfPanelRows128 = kMatrixSize128 - kBlockSize;
constexpr int kHalfPanelStride128 = 40;
constexpr int kHalfPanelBytes128 =
    kHalfPanelRows128 * kHalfPanelStride128 * static_cast<int>(sizeof(__half));
constexpr int kHalfPanelsBytes128 = 2 * kHalfPanelBytes128;
constexpr int kWmmaSharedBytes128 =
    kWmmaFloatSharedBytes128 + kHalfPanelsBytes128;

template <int... Indices>
struct N64IndexSequence {};

template <int Count, int... Indices>
struct N64MakeIndexSequence
    : N64MakeIndexSequence<Count - 1, Count - 1, Indices...> {};

template <int... Indices>
struct N64MakeIndexSequence<0, Indices...> {
  using type = N64IndexSequence<Indices...>;
};

constexpr int kNestedBlockSize64 = 16;
using N64AllIndices =
    typename N64MakeIndexSequence<kNestedBlockSize64>::type;

struct N64RegisterRow {
  float v0;
  float v1;
  float v2;
  float v3;
  float v4;
  float v5;
  float v6;
  float v7;
  float v8;
  float v9;
  float v10;
  float v11;
  float v12;
  float v13;
  float v14;
  float v15;
};

template <int Index>
struct N64RegisterAccess;

#define N64_REGISTER_ACCESS(Index)                                      \
  template <>                                                          \
  struct N64RegisterAccess<Index> {                                    \
    __device__ __forceinline__ static float get(                       \
        const N64RegisterRow& row) {                                   \
      return row.v##Index;                                             \
    }                                                                  \
    __device__ __forceinline__ static void set(                        \
        N64RegisterRow& row, float value) {                            \
      row.v##Index = value;                                            \
    }                                                                  \
  };

N64_REGISTER_ACCESS(0)
N64_REGISTER_ACCESS(1)
N64_REGISTER_ACCESS(2)
N64_REGISTER_ACCESS(3)
N64_REGISTER_ACCESS(4)
N64_REGISTER_ACCESS(5)
N64_REGISTER_ACCESS(6)
N64_REGISTER_ACCESS(7)
N64_REGISTER_ACCESS(8)
N64_REGISTER_ACCESS(9)
N64_REGISTER_ACCESS(10)
N64_REGISTER_ACCESS(11)
N64_REGISTER_ACCESS(12)
N64_REGISTER_ACCESS(13)
N64_REGISTER_ACCESS(14)
N64_REGISTER_ACCESS(15)

#undef N64_REGISTER_ACCESS

template <int Stride, int... Indices>
__device__ __forceinline__ void n64_load_register_row(
    N64RegisterRow& values,
    float (*tile)[Stride],
    int row,
    int col_base,
    N64IndexSequence<Indices...>) {
  int unused[] = {
      0, (N64RegisterAccess<Indices>::set(
              values, tile[row][col_base + Indices]),
          0)...};
  (void)unused;
}

template <int Stride, int... Indices>
__device__ __forceinline__ void n64_store_register_row(
    const N64RegisterRow& values,
    float (*tile)[Stride],
    int row,
    int col_base,
    N64IndexSequence<Indices...>) {
  int unused[] = {
      0, ((tile[row][col_base + Indices] =
               N64RegisterAccess<Indices>::get(values)),
          0)...};
  (void)unused;
}

template <int Pivot, int... Previous>
__device__ __forceinline__ void n64_factor_step(
    N64RegisterRow& values,
    int lane,
    N64IndexSequence<Previous...>) {
  float diagonal = 0.0f;
  if (lane == Pivot) {
    diagonal = N64RegisterAccess<Pivot>::get(values);
    int diagonal_terms[] = {
        0, ((diagonal = fmaf(
                 -N64RegisterAccess<Previous>::get(values),
                 N64RegisterAccess<Previous>::get(values), diagonal)),
            0)...};
    (void)diagonal_terms;
    diagonal = sqrtf(fmaxf(diagonal, 0.0f));
  }
  diagonal = __shfl_sync(0xffffffffu, diagonal, Pivot);

  float value = N64RegisterAccess<Pivot>::get(values);
  int update_terms[] = {
      0, ((value = fmaf(
               -N64RegisterAccess<Previous>::get(values),
               __shfl_sync(0xffffffffu,
                           N64RegisterAccess<Previous>::get(values), Pivot),
               value)),
          0)...};
  (void)update_terms;
  if (lane == Pivot) {
    N64RegisterAccess<Pivot>::set(values, diagonal);
  } else if (lane > Pivot && lane < kNestedBlockSize64) {
    N64RegisterAccess<Pivot>::set(values, value / diagonal);
  }
}

template <int... Pivots>
__device__ __forceinline__ void n64_factor_register_block(
    N64RegisterRow& values,
    int lane,
    N64IndexSequence<Pivots...>) {
  int unused[] = {
      0, (n64_factor_step<Pivots>(
              values, lane,
              typename N64MakeIndexSequence<Pivots>::type{}),
          0)...};
  (void)unused;
}

template <int Pivot, int... Previous>
__device__ __forceinline__ void n64_nested_panel_step(
    N64RegisterRow& values,
    int lane,
    N64IndexSequence<Previous...>) {
  float value = N64RegisterAccess<Pivot>::get(values);
  int update_terms[] = {
      0, ((value = fmaf(
               -N64RegisterAccess<Previous>::get(values),
               __shfl_sync(0xffffffffu,
                           N64RegisterAccess<Previous>::get(values), Pivot),
               value)),
          0)...};
  (void)update_terms;
  const float diagonal = __shfl_sync(
      0xffffffffu, N64RegisterAccess<Pivot>::get(values), Pivot);
  if (lane >= kNestedBlockSize64) {
    N64RegisterAccess<Pivot>::set(values, value / diagonal);
  }
}

template <int... Pivots>
__device__ __forceinline__ void n64_solve_nested_panel(
    N64RegisterRow& values,
    int lane,
    N64IndexSequence<Pivots...>) {
  int unused[] = {
      0, (n64_nested_panel_step<Pivots>(
              values, lane,
              typename N64MakeIndexSequence<Pivots>::type{}),
          0)...};
  (void)unused;
}

template <int... Terms>
__device__ __forceinline__ float n64_nested_panel_dot(
    const N64RegisterRow& values,
    int row,
    int col,
    float value,
    N64IndexSequence<Terms...>) {
  int unused[] = {
      0, ((value = fmaf(
               -__shfl_sync(0xffffffffu,
                            N64RegisterAccess<Terms>::get(values),
                            kNestedBlockSize64 + row),
               __shfl_sync(0xffffffffu,
                           N64RegisterAccess<Terms>::get(values),
                           kNestedBlockSize64 + col),
               value)),
          0)...};
  (void)unused;
  return value;
}

template <int Stride>
__device__ __forceinline__ void n64_factor_nested32(
    float (*tile)[Stride],
    int lane,
    int base) {
  N64RegisterRow values;
  if (lane < kNestedBlockSize64) {
    n64_load_register_row(values, tile, base + lane, base, N64AllIndices{});
  } else {
    values = {};
  }

  n64_factor_register_block(values, lane, N64AllIndices{});

  if (lane < kNestedBlockSize64) {
    n64_store_register_row(values, tile, base + lane, base, N64AllIndices{});
  } else {
    n64_load_register_row(values, tile, base + lane, base, N64AllIndices{});
  }

  n64_solve_nested_panel(values, lane, N64AllIndices{});

  if (lane >= kNestedBlockSize64) {
    n64_store_register_row(values, tile, base + lane, base, N64AllIndices{});
  }

#pragma unroll
  for (int linear = lane;
       linear < kNestedBlockSize64 * kNestedBlockSize64;
       linear += 32) {
    const int row = linear >> 4;
    const int col = linear & 15;
    float value = tile[base + kNestedBlockSize64 + row]
                      [base + kNestedBlockSize64 + col];
    value = n64_nested_panel_dot(
        values, row, col, value, N64AllIndices{});
    if (row >= col) {
      tile[base + kNestedBlockSize64 + row]
          [base + kNestedBlockSize64 + col] = value;
    }
  }
  __syncwarp();

  if (lane < kNestedBlockSize64) {
    n64_load_register_row(
        values, tile, base + kNestedBlockSize64 + lane,
        base + kNestedBlockSize64, N64AllIndices{});
  }

  n64_factor_register_block(values, lane, N64AllIndices{});

  if (lane < kNestedBlockSize64) {
    n64_store_register_row(
        values, tile, base + kNestedBlockSize64 + lane,
        base + kNestedBlockSize64, N64AllIndices{});
  }
}

template <int Pivot, int Stride, int... Previous>
__device__ __forceinline__ void n64_outer_panel_first_step(
    N64RegisterRow& values,
    float (*tile)[Stride],
    int factor_base,
    N64IndexSequence<Previous...>) {
  float value = N64RegisterAccess<Pivot>::get(values);
  int update_terms[] = {
      0, ((value = fmaf(
               -N64RegisterAccess<Previous>::get(values),
               tile[factor_base + Pivot][factor_base + Previous], value)),
          0)...};
  (void)update_terms;
  N64RegisterAccess<Pivot>::set(
      values, value / tile[factor_base + Pivot][factor_base + Pivot]);
}

template <int Stride, int... Pivots>
__device__ __forceinline__ void n64_solve_outer_panel_first(
    N64RegisterRow& values,
    float (*tile)[Stride],
    int factor_base,
    N64IndexSequence<Pivots...>) {
  int unused[] = {
      0, (n64_outer_panel_first_step<Pivots>(
              values, tile, factor_base,
              typename N64MakeIndexSequence<Pivots>::type{}),
          0)...};
  (void)unused;
}

template <int Pivot, int Stride, int... History, int... Previous>
__device__ __forceinline__ void n64_outer_panel_second_step(
    N64RegisterRow& values,
    float (*tile)[Stride],
    int row,
    int factor_base,
    N64IndexSequence<History...>,
    N64IndexSequence<Previous...>) {
  float value = N64RegisterAccess<Pivot>::get(values);
  int history_terms[] = {
      0, ((value = fmaf(
               -tile[row][factor_base + History],
               tile[factor_base + kNestedBlockSize64 + Pivot]
                   [factor_base + History],
               value)),
          0)...};
  (void)history_terms;
  int update_terms[] = {
      0, ((value = fmaf(
               -N64RegisterAccess<Previous>::get(values),
               tile[factor_base + kNestedBlockSize64 + Pivot]
                   [factor_base + kNestedBlockSize64 + Previous],
               value)),
          0)...};
  (void)update_terms;
  N64RegisterAccess<Pivot>::set(
      values,
      value /
          tile[factor_base + kNestedBlockSize64 + Pivot]
              [factor_base + kNestedBlockSize64 + Pivot]);
}

template <int Stride, int... Pivots>
__device__ __forceinline__ void n64_solve_outer_panel_second(
    N64RegisterRow& values,
    float (*tile)[Stride],
    int row,
    int factor_base,
    N64IndexSequence<Pivots...>) {
  int unused[] = {
      0, (n64_outer_panel_second_step<Pivots>(
              values, tile, row, factor_base, N64AllIndices{},
              typename N64MakeIndexSequence<Pivots>::type{}),
          0)...};
  (void)unused;
}

// Factor a 64x64 matrix as two 32x32 diagonal blocks.  FP32 shared memory is
// authoritative throughout.  A compact, aligned FP16 mirror holds selected
// finalized L10 entries for grouped half2 work in the rank-32 Schur update.
__global__ __launch_bounds__(kThreads) void cholesky64_block32_hacc4safe_kernel(
    const float* __restrict__ input,
    float* __restrict__ output) {
  const int matrix_offset = static_cast<int>(blockIdx.x) * kMatrixElements;
  const int tid = static_cast<int>(threadIdx.x);
  const int lane = tid & 31;
  const int warp = tid >> 5;

  __shared__ float tile[kMatrixSize][kSharedStride];
  __shared__ float schur_diagonal[kBlockSize];
  __shared__ __align__(16) __half
      schur_half[kBlockSize][kSchurHalfStride];

#pragma unroll
  for (int vector = tid; vector < kMatrixElements / 4;
       vector += kThreads) {
    const int load_row = vector >> 4;
    const int col = (vector & 15) << 2;
    const float4 packed = *reinterpret_cast<const float4*>(
        input + matrix_offset + load_row * kMatrixSize + col);
    tile[load_row][col + 0] = load_row >= col + 0 ? packed.x : 0.0f;
    tile[load_row][col + 1] = load_row >= col + 1 ? packed.y : 0.0f;
    tile[load_row][col + 2] = load_row >= col + 2 ? packed.z : 0.0f;
    tile[load_row][col + 3] = load_row >= col + 3 ? packed.w : 0.0f;
  }
  __syncthreads();

  // Phase 1: warp zero factors A00 as two bounded 16-column register stages.
  // Every lane carries at most one explicit 16-float row, while shuffles move
  // finalized recurrence values between row owners.
  if (warp == 0) {
    n64_factor_nested32(tile, lane, 0);
  }
  __syncthreads();

  // Phase 2: two warps each solve sixteen L10 rows.  A lane owns one bounded
  // 16-float row at a time; the second stage consumes the first stage from
  // shared memory and preserves the original left-to-right arithmetic order.
  if (warp < 2 && lane < kNestedBlockSize64) {
    const int row = kBlockSize + warp * kNestedBlockSize64 + lane;
    N64RegisterRow values;
    n64_load_register_row(values, tile, row, 0, N64AllIndices{});
    n64_solve_outer_panel_first(values, tile, 0, N64AllIndices{});
    n64_store_register_row(values, tile, row, 0, N64AllIndices{});

    n64_load_register_row(
        values, tile, row, kNestedBlockSize64, N64AllIndices{});
    n64_solve_outer_panel_second(
        values, tile, row, 0, N64AllIndices{});
    n64_store_register_row(
        values, tile, row, kNestedBlockSize64, N64AllIndices{});
  }
  __syncthreads();

  if (warp == 0) {
    const int diagonal = kBlockSize + lane;
    schur_diagonal[lane] = tile[diagonal][diagonal];
  }

  // Compact odd columns 5..31 from the finalized panel.  Seven adjacent half2
  // loads cover fourteen products; the 32-byte row stride keeps every load
  // naturally aligned and avoids per-product half-to-float conversion.
#pragma unroll
  for (int linear = tid; linear < kBlockSize * kSchurHalfTerms;
       linear += kThreads) {
    const int local_row = linear / kSchurHalfTerms;
    const int slot = linear - local_row * kSchurHalfTerms;
    schur_half[local_row][slot] =
        __float2half_rn(tile[kBlockSize + local_row][5 + 2 * slot]);
  }
  __syncthreads();

  // Phase 3: all threads update the lower triangle of A11 in shared memory.
#pragma unroll
  for (int linear = tid; linear < kBlockSize * kBlockSize;
       linear += kThreads) {
    const int local_row = linear >> 5;
    const int local_col = linear & 31;
    if (local_row >= local_col) {
      const int row = kBlockSize + local_row;
      const int col = kBlockSize + local_col;
      float value = tile[row][col];

      // Keep eighteen products in authoritative FP32: all columns 0..4 and
      // the thirteen even columns 6..30.  The remaining fourteen odd-column
      // products use binary16 inputs with four independent FP32 FMA chains.
#pragma unroll 4
      for (int k = 0; k < 5; ++k) {
        value = fmaf(-tile[row][k], tile[col][k], value);
      }
#pragma unroll
      for (int k = 6; k < kBlockSize; k += 2) {
        value = fmaf(-tile[row][k], tile[col][k], value);
      }

      const __half2 row0 = *reinterpret_cast<const __half2*>(
          &schur_half[local_row][0]);
      const __half2 row1 = *reinterpret_cast<const __half2*>(
          &schur_half[local_row][2]);
      const __half2 row2 = *reinterpret_cast<const __half2*>(
          &schur_half[local_row][4]);
      const __half2 row3 = *reinterpret_cast<const __half2*>(
          &schur_half[local_row][6]);
      const __half2 col0 = *reinterpret_cast<const __half2*>(
          &schur_half[local_col][0]);
      const __half2 col1 = *reinterpret_cast<const __half2*>(
          &schur_half[local_col][2]);
      const __half2 col2 = *reinterpret_cast<const __half2*>(
          &schur_half[local_col][4]);
      const __half2 col3 = *reinterpret_cast<const __half2*>(
          &schur_half[local_col][6]);
      float half_dot0 = n64_fhfma2_acc(row0, col0, 0.0f);
      half_dot0 = n64_fhfma2_acc(row1, col1, half_dot0);
      float half_dot1 = n64_fhfma2_acc(row2, col2, 0.0f);
      half_dot1 = n64_fhfma2_acc(row3, col3, half_dot1);

      const __half2 row4 = *reinterpret_cast<const __half2*>(
          &schur_half[local_row][8]);
      const __half2 row5 = *reinterpret_cast<const __half2*>(
          &schur_half[local_row][10]);
      const __half2 row6 = *reinterpret_cast<const __half2*>(
          &schur_half[local_row][12]);
      const __half2 col4 = *reinterpret_cast<const __half2*>(
          &schur_half[local_col][8]);
      const __half2 col5 = *reinterpret_cast<const __half2*>(
          &schur_half[local_col][10]);
      const __half2 col6 = *reinterpret_cast<const __half2*>(
          &schur_half[local_col][12]);
      float half_dot2 = n64_fhfma2_acc(row4, col4, 0.0f);
      half_dot2 = n64_fhfma2_acc(row5, col5, half_dot2);
      const float half_dot3 = n64_fhfma2_acc(row6, col6, 0.0f);
      value -= (half_dot0 + half_dot1) + (half_dot2 + half_dot3);
      tile[row][col] = value;
    }
  }
  __syncthreads();

  // Restore all Schur diagonal entries from their original values with exact
  // FP32 rank-32 dots.  Two warps each own sixteen diagonals, and adjacent
  // lanes split one diagonal into two independent sixteen-term chains.
  if (warp < 2) {
    const int pair = lane >> 1;
    const int half = lane & 1;
    const int local_diagonal = warp * 16 + pair;
    const int row = kBlockSize + local_diagonal;
    float value = half == 0 ? schur_diagonal[local_diagonal] : 0.0f;
#pragma unroll 4
    for (int k = half * 16; k < half * 16 + 16; ++k) {
      const float panel_value = tile[row][k];
      value = fmaf(-panel_value, panel_value, value);
    }
    value += __shfl_xor_sync(0xffffffffu, value, 1);
    if (half == 0) {
      tile[row][row] = value;
    }
  }
  __syncthreads();

  // Phase 4: warp zero factors the repaired A11 with the same two bounded
  // 16-column register stages used for A00.
  if (warp == 0) {
    n64_factor_nested32(tile, lane, kBlockSize);
  }
  __syncthreads();

#pragma unroll
  for (int vector = tid; vector < kMatrixElements / 4;
       vector += kThreads) {
    const int store_row = vector >> 4;
    const int col = (vector & 15) << 2;
    float4 packed;
    packed.x = store_row >= col + 0 ? tile[store_row][col + 0] : 0.0f;
    packed.y = store_row >= col + 1 ? tile[store_row][col + 1] : 0.0f;
    packed.z = store_row >= col + 2 ? tile[store_row][col + 2] : 0.0f;
    packed.w = store_row >= col + 3 ? tile[store_row][col + 3] : 0.0f;
    *reinterpret_cast<float4*>(
        output + matrix_offset + store_row * kMatrixSize + col) = packed;
  }
}

template <int... Indices>
struct N32IndexSequence {};

template <int Count, int... Indices>
struct N32MakeIndexSequence
    : N32MakeIndexSequence<Count - 1, Count - 1, Indices...> {};

template <int... Indices>
struct N32MakeIndexSequence<0, Indices...> {
  using type = N32IndexSequence<Indices...>;
};

using N32AllIndices = typename N32MakeIndexSequence<kBlockSize32>::type;

struct N32RegisterRow {
  float v0;
  float v1;
  float v2;
  float v3;
  float v4;
  float v5;
  float v6;
  float v7;
  float v8;
  float v9;
  float v10;
  float v11;
  float v12;
  float v13;
  float v14;
  float v15;
};

template <int Index>
struct N32RegisterAccess;

#define N32_REGISTER_ACCESS(Index)                                      \
  template <>                                                          \
  struct N32RegisterAccess<Index> {                                    \
    __device__ __forceinline__ static float get(                       \
        const N32RegisterRow& row) {                                   \
      return row.v##Index;                                             \
    }                                                                  \
    __device__ __forceinline__ static void set(                        \
        N32RegisterRow& row, float value) {                            \
      row.v##Index = value;                                            \
    }                                                                  \
  };

N32_REGISTER_ACCESS(0)
N32_REGISTER_ACCESS(1)
N32_REGISTER_ACCESS(2)
N32_REGISTER_ACCESS(3)
N32_REGISTER_ACCESS(4)
N32_REGISTER_ACCESS(5)
N32_REGISTER_ACCESS(6)
N32_REGISTER_ACCESS(7)
N32_REGISTER_ACCESS(8)
N32_REGISTER_ACCESS(9)
N32_REGISTER_ACCESS(10)
N32_REGISTER_ACCESS(11)
N32_REGISTER_ACCESS(12)
N32_REGISTER_ACCESS(13)
N32_REGISTER_ACCESS(14)
N32_REGISTER_ACCESS(15)

#undef N32_REGISTER_ACCESS

template <int... Indices>
__device__ __forceinline__ void n32_load_register_row(
    N32RegisterRow& values,
    float (*io)[kIoStride32],
    int row,
    int col_base,
    N32IndexSequence<Indices...>) {
  int unused[] = {
      0, (N32RegisterAccess<Indices>::set(
              values, io[row][col_base + Indices]),
          0)...};
  (void)unused;
}

template <int... Indices>
__device__ __forceinline__ void n32_store_register_row(
    const N32RegisterRow& values,
    float (*io)[kIoStride32],
    int row,
    int col_base,
    N32IndexSequence<Indices...>) {
  int unused[] = {
      0, ((io[row][col_base + Indices] =
               N32RegisterAccess<Indices>::get(values)),
          0)...};
  (void)unused;
}

template <int Pivot, int... Previous>
__device__ __forceinline__ void n32_factor_step(
    N32RegisterRow& values,
    int lane,
    N32IndexSequence<Previous...>) {
  float diagonal = 0.0f;
  if (lane == Pivot) {
    diagonal = N32RegisterAccess<Pivot>::get(values);
    int diagonal_terms[] = {
        0, ((diagonal = fmaf(
                 -N32RegisterAccess<Previous>::get(values),
                 N32RegisterAccess<Previous>::get(values), diagonal)),
            0)...};
    (void)diagonal_terms;
    diagonal = sqrtf(fmaxf(diagonal, 0.0f));
  }
  diagonal = __shfl_sync(0xffffffffu, diagonal, Pivot);

  float value = N32RegisterAccess<Pivot>::get(values);
  int update_terms[] = {
      0, ((value = fmaf(
               -N32RegisterAccess<Previous>::get(values),
               __shfl_sync(0xffffffffu,
                           N32RegisterAccess<Previous>::get(values), Pivot),
               value)),
          0)...};
  (void)update_terms;
  if (lane == Pivot) {
    N32RegisterAccess<Pivot>::set(values, diagonal);
  } else if (lane > Pivot && lane < kBlockSize32) {
    N32RegisterAccess<Pivot>::set(values, value / diagonal);
  }
}

template <int... Pivots>
__device__ __forceinline__ void n32_factor_register_block(
    N32RegisterRow& values,
    int lane,
    N32IndexSequence<Pivots...>) {
  int unused[] = {
      0, (n32_factor_step<Pivots>(
              values, lane,
              typename N32MakeIndexSequence<Pivots>::type{}),
          0)...};
  (void)unused;
}

template <int Pivot, int... Previous>
__device__ __forceinline__ void n32_panel_step(
    N32RegisterRow& values,
    int lane,
    N32IndexSequence<Previous...>) {
  float value = N32RegisterAccess<Pivot>::get(values);
  int update_terms[] = {
      0, ((value = fmaf(
               -N32RegisterAccess<Previous>::get(values),
               __shfl_sync(0xffffffffu,
                           N32RegisterAccess<Previous>::get(values), Pivot),
               value)),
          0)...};
  (void)update_terms;
  const float diagonal = __shfl_sync(
      0xffffffffu, N32RegisterAccess<Pivot>::get(values), Pivot);
  if (lane >= kBlockSize32) {
    N32RegisterAccess<Pivot>::set(values, value / diagonal);
  }
}

template <int... Pivots>
__device__ __forceinline__ void n32_solve_register_panel(
    N32RegisterRow& values,
    int lane,
    N32IndexSequence<Pivots...>) {
  int unused[] = {
      0, (n32_panel_step<Pivots>(
              values, lane,
              typename N32MakeIndexSequence<Pivots>::type{}),
          0)...};
  (void)unused;
}

template <int... Terms>
__device__ __forceinline__ float n32_panel_dot(
    const N32RegisterRow& values,
    int row,
    int col,
    float value,
    N32IndexSequence<Terms...>) {
  int unused[] = {
      0, ((value = fmaf(
               -__shfl_sync(0xffffffffu,
                            N32RegisterAccess<Terms>::get(values),
                            kBlockSize32 + row),
               __shfl_sync(0xffffffffu,
                           N32RegisterAccess<Terms>::get(values),
                           kBlockSize32 + col),
               value)),
          0)...};
  (void)unused;
  return value;
}

// Each warp factors one matrix as two exact 16-column register stages.  Lanes
// 0..15 first own D00 rows while lanes 16..31 take over for L10, then those
// roles reverse for D11.  Thus every lane carries at most one 16-float row.
// Shared memory is only a padded whole-matrix I/O tile; all recurrence values
// move between lanes through full-warp shuffles.
__global__ __launch_bounds__(kThreads32) void
cholesky32_lower_blocktiles_kernel(
    const float* __restrict__ input,
    float* __restrict__ output,
    int batch) {
  const int lane = static_cast<int>(threadIdx.x) & 31;
  const int warp = static_cast<int>(threadIdx.x) >> 5;
  const int matrix = static_cast<int>(blockIdx.x) * kWarpsPerBlock32 + warp;
  if (matrix >= batch) {
    return;
  }
  const int matrix_offset = matrix * kMatrixElements32;

  __shared__ float io_tiles[kWarpsPerBlock32][kMatrixSize32][kIoStride32];
  float (*io)[kIoStride32] = io_tiles[warp];

  // Coalesced whole-matrix load.  Upper entries become the final explicit
  // zeros immediately and are never touched by the factorization.
#pragma unroll
  for (int linear = lane; linear < kMatrixElements32;
       linear += 32) {
    const int row = linear >> 5;
    const int col = linear & 31;
    io[row][col] = row >= col ? input[matrix_offset + linear] : 0.0f;
  }
  __syncwarp();

  N32RegisterRow values;
  if (lane < kBlockSize32) {
    n32_load_register_row(values, io, lane, 0, N32AllIndices{});
  } else {
    values = {};
  }

  // Stage 1: exact FP32 Cholesky of D00.  Every lane executes every shuffle;
  // only lanes 0..15 retain the row recurrence in their private registers.
  n32_factor_register_block(values, lane, N32AllIndices{});

  if (lane < kBlockSize32) {
    n32_store_register_row(values, io, lane, 0, N32AllIndices{});
  } else {
    n32_load_register_row(values, io, lane, 0, N32AllIndices{});
  }

  // Lanes 16..31 now own L10 rows.  Lanes 0..15 keep D00 resident so the
  // triangular solve also consumes recurrence data only through shuffles.
  n32_solve_register_panel(values, lane, N32AllIndices{});

  if (lane >= kBlockSize32) {
    n32_store_register_row(values, io, lane, 0, N32AllIndices{});
  }

  // Form D11 with its 256 entries striped across the full warp.  The solved
  // panel remains in lanes 16..31; every lane executes every shuffle, and the
  // completed lower entries return to the sole shared I/O tile.
#pragma unroll
  for (int linear = lane; linear < kBlockSize32 * kBlockSize32;
       linear += 32) {
    const int row = linear >> 4;
    const int col = linear & 15;
    float value = io[kBlockSize32 + row][kBlockSize32 + col];
    value = n32_panel_dot(values, row, col, value, N32AllIndices{});
    if (row >= col) {
      io[kBlockSize32 + row][kBlockSize32 + col] = value;
    }
  }
  __syncwarp();

  if (lane < kBlockSize32) {
    n32_load_register_row(
        values, io, kBlockSize32 + lane, kBlockSize32, N32AllIndices{});
  }

  // Stage 2: exact FP32 Cholesky of the register-resident D11 tile.
  n32_factor_register_block(values, lane, N32AllIndices{});

  if (lane < kBlockSize32) {
    n32_store_register_row(
        values, io, kBlockSize32 + lane, kBlockSize32, N32AllIndices{});
  }
  __syncwarp();

  // Coalesced whole-matrix store includes all explicitly zeroed upper entries.
#pragma unroll
  for (int linear = lane; linear < kMatrixElements32;
       linear += 32) {
    const int row = linear >> 5;
    const int col = linear & 31;
    output[matrix_offset + linear] = io[row][col];
  }
}

// Four 32x32 blocked stages factor one 128x128 matrix per CTA.  FP32 remains
// authoritative while a padded FP16 panel mirror feeds rank-32 WMMA updates.
// Full-FP32 fallback for non-public n=128 batch sizes. Four 32x32 blocked
// stages factor one matrix per CTA while transferring only the lower triangle.
__global__ __launch_bounds__(kThreads128) void
cholesky128_block32_warp16_lowerio_anybatch_kernel(
    const float* __restrict__ input,
    float* __restrict__ output) {
  const int matrix_offset = static_cast<int>(blockIdx.x) * kMatrixElements128;
  const int tid = static_cast<int>(threadIdx.x);
  const int lane = tid & 31;
  const int warp = tid >> 5;

  extern __shared__ float fp32_shared_values[];
  float (*tile)[kSharedStride128] =
      reinterpret_cast<float (*)[kSharedStride128]>(fp32_shared_values);

#pragma unroll
  for (int row = warp; row < kMatrixSize128; row += kWarpsPerBlock128) {
#pragma unroll
    for (int col = lane; col <= row; col += 32) {
      tile[row][col] = input[matrix_offset + row * kMatrixSize128 + col];
    }
  }
  __syncthreads();

#pragma unroll
  for (int block = 0; block < 4; ++block) {
    const int base = block * kBlockSize;

    if (warp == 0) {
      const int row = base + lane;
#pragma unroll 1
      for (int k = 0; k < kBlockSize; ++k) {
        const int col = base + k;
        float diagonal = 0.0f;
        if (lane == k) {
          diagonal = tile[row][col];
#pragma unroll 4
          for (int j = 0; j < k; ++j) {
            const float value = tile[row][base + j];
            diagonal = fmaf(-value, value, diagonal);
          }
          diagonal = sqrtf(fmaxf(diagonal, 0.0f));
        }
        diagonal = __shfl_sync(0xffffffffu, diagonal, k);

        if (lane == k) {
          tile[row][col] = diagonal;
        } else if (lane > k) {
          float value = tile[row][col];
#pragma unroll 4
          for (int j = 0; j < k; ++j) {
            value = fmaf(-tile[row][base + j], tile[col][base + j], value);
          }
          tile[row][col] = value / diagonal;
        }
        __syncwarp();
      }
    }
    __syncthreads();

    const int remaining_blocks = 3 - block;
    if (warp < remaining_blocks) {
      const int row = base + kBlockSize * (warp + 1) + lane;
#pragma unroll 1
      for (int k = 0; k < kBlockSize; ++k) {
        const int col = base + k;
        float value = tile[row][col];
#pragma unroll 4
        for (int j = 0; j < k; ++j) {
          value = fmaf(-tile[row][base + j], tile[col][base + j], value);
        }
        tile[row][col] = value / tile[col][col];
      }
    }
    __syncthreads();

    const int trailing_base = base + kBlockSize;
    const int trailing_size = kMatrixSize128 - trailing_base;
    const int trailing_elements = trailing_size * trailing_size;
    for (int linear = tid; linear < trailing_elements;
         linear += kThreads128) {
      const int local_row = linear / trailing_size;
      const int local_col = linear - local_row * trailing_size;
      if (local_row >= local_col) {
        const int row = trailing_base + local_row;
        const int col = trailing_base + local_col;
        float value = tile[row][col];
#pragma unroll
        for (int k = 0; k < kBlockSize; ++k) {
          value = fmaf(-tile[row][base + k], tile[col][base + k], value);
        }
        tile[row][col] = value;
      }
    }
    __syncthreads();
  }

#pragma unroll
  for (int row = warp; row < kMatrixSize128; row += kWarpsPerBlock128) {
#pragma unroll
    for (int col = lane; col <= row; col += 32) {
      output[matrix_offset + row * kMatrixSize128 + col] = tile[row][col];
    }
#pragma unroll
    for (int col = row + 1 + lane; col < kMatrixSize128; col += 32) {
      output[matrix_offset + row * kMatrixSize128 + col] = 0.0f;
    }
  }
}

template <bool Compensated>
__global__ __launch_bounds__(kThreads128) void cholesky128_wmma_schur_kernel(
    const float* __restrict__ input,
    float* __restrict__ output) {
  const int matrix_offset = static_cast<int>(blockIdx.x) * kMatrixElements128;
  const int tid = static_cast<int>(threadIdx.x);
  const int lane = tid & 31;
  const int warp = tid >> 5;

  extern __shared__ __align__(32) unsigned char shared_values[];
  float (*tile)[kWmmaSharedStride128] =
      reinterpret_cast<float (*)[kWmmaSharedStride128]>(shared_values);
  __half (*factor_half_hi)[kHalfPanelStride128] =
      reinterpret_cast<__half (*)[kHalfPanelStride128]>(
          shared_values + kWmmaFloatSharedBytes128);
  __half (*factor_half_lo)[kHalfPanelStride128] =
      reinterpret_cast<__half (*)[kHalfPanelStride128]>(
          shared_values + kWmmaFloatSharedBytes128 + kHalfPanelBytes128);

  // Direct accumulator loads consume complete 16x16 diagonal tiles, so every
  // shared element must be initialized. The input contract is symmetric.
#pragma unroll
  for (int row = warp; row < kMatrixSize128; row += kWarpsPerBlock128) {
#pragma unroll
    for (int col = lane;
         col < (Compensated ? kMatrixSize128 : row + 1); col += 32) {
      tile[row][col] =
          input[matrix_offset + row * kMatrixSize128 + col];
    }
  }
  __syncthreads();

#pragma unroll
  for (int block = 0; block < 4; ++block) {
    const int base = block * kBlockSize;

    // The exact public path factors each 32x32 diagonal tile as two bounded
    // 16-column register stages.  The compensated fallback retains its
    // established shared-memory recurrence.
    if (warp == 0) {
      if constexpr (!Compensated) {
        n64_factor_nested32(tile, lane, base);
      } else {
        const int row = base + lane;
#pragma unroll 1
        for (int k = 0; k < kBlockSize; ++k) {
          const int col = base + k;
          float diagonal = 0.0f;
          if (lane == k) {
            diagonal = tile[row][col];
#pragma unroll 4
            for (int j = 0; j < k; ++j) {
              const float value = tile[row][base + j];
              diagonal = fmaf(-value, value, diagonal);
            }
            diagonal = sqrtf(fmaxf(diagonal, 0.0f));
          }
          diagonal = __shfl_sync(0xffffffffu, diagonal, k);

          if (lane == k) {
            tile[row][col] = diagonal;
          } else if (lane > k) {
            float value = tile[row][col];
#pragma unroll 4
            for (int j = 0; j < k; ++j) {
              value = fmaf(-tile[row][base + j],
                           tile[col][base + j], value);
            }
            tile[row][col] = value / diagonal;
          }
          __syncwarp();
        }
      }
    }
    __syncthreads();

    // One warp owns each remaining 32-row panel block.  The exact path keeps
    // each half-row in registers and performs the two triangular solves with
    // fully expanded dependencies; no inter-stage polling is required.
    const int remaining_blocks = 3 - block;
    if (warp < remaining_blocks) {
      const int row = base + kBlockSize * (warp + 1) + lane;
      if constexpr (!Compensated) {
        N64RegisterRow values;
        n64_load_register_row(values, tile, row, base, N64AllIndices{});
        n64_solve_outer_panel_first(
            values, tile, base, N64AllIndices{});
        n64_store_register_row(values, tile, row, base, N64AllIndices{});

        n64_load_register_row(
            values, tile, row, base + kNestedBlockSize64, N64AllIndices{});
        n64_solve_outer_panel_second(
            values, tile, row, base, N64AllIndices{});
        n64_store_register_row(
            values, tile, row, base + kNestedBlockSize64, N64AllIndices{});
      } else {
#pragma unroll 1
        for (int k = 0; k < kBlockSize; ++k) {
          const int col = base + k;
          float value = tile[row][col];
#pragma unroll 4
          for (int j = 0; j < k; ++j) {
            value = fmaf(-tile[row][base + j],
                         tile[col][base + j], value);
          }
          tile[row][col] = value / tile[col][col];
        }
      }
    }
    __syncthreads();

    const int trailing_base = base + kBlockSize;
    const int trailing_size = kMatrixSize128 - trailing_base;
    if (trailing_size == 0) {
      continue;
    }

    // Split each finalized panel value into a rounded half and a rounded-half
    // residual. The 40-half leading dimension satisfies WMMA alignment.
    const int panel_elements = trailing_size * kBlockSize;
    for (int linear = tid; linear < panel_elements; linear += kThreads128) {
      const int local_row = linear >> 5;
      const int k = linear & 31;
      const int row = trailing_base + local_row;
      const float value = tile[row][base + k];
      const __half high = __float2half_rn(value);
      factor_half_hi[local_row][k] = high;
      if constexpr (Compensated) {
        factor_half_lo[local_row][k] =
            __float2half_rn(value - __half2float(high));
      }
    }
    __syncthreads();

    // Enumerate lower 16x16 trailing tiles exactly once.  Each warp loads its
    // current FP32 tile directly into an accumulator, negates the FP16 A
    // operand, and computes C + (-A) * B before storing directly back to the
    // authoritative tile.  Full diagonal tiles are updated; their upper
    // entries are never consumed by the lower-triangular factorization.
    const int trailing_tiles = trailing_size >> 4;
    const int task_count = trailing_tiles * (trailing_tiles + 1) / 2;
    const int task_rounds = (task_count + kWarpsPerBlock128 - 1) /
                            kWarpsPerBlock128;

    for (int round = 0; round < task_rounds; ++round) {
      const int triangular_task = round * kWarpsPerBlock128 + warp;
      if (triangular_task < task_count) {
        int tile_row = 0;
        int tile_col = triangular_task;
#pragma unroll
        for (int candidate_row = 1; candidate_row < 8; ++candidate_row) {
          if (tile_col >= candidate_row) {
            tile_col -= candidate_row;
            tile_row = candidate_row;
          }
        }

        const int row = trailing_base + (tile_row << 4);
        const int col = trailing_base + (tile_col << 4);
        nvcuda::wmma::fragment<nvcuda::wmma::accumulator, 16, 16, 16, float>
            updated;
        nvcuda::wmma::load_matrix_sync(
            updated, &tile[row][col], kWmmaSharedStride128,
            nvcuda::wmma::mem_row_major);

#pragma unroll
        for (int k = 0; k < kBlockSize; k += 16) {
          nvcuda::wmma::fragment<nvcuda::wmma::matrix_a, 16, 16, 16, __half,
                                 nvcuda::wmma::row_major>
              panel_a;
          nvcuda::wmma::fragment<nvcuda::wmma::matrix_b, 16, 16, 16, __half,
                                 nvcuda::wmma::col_major>
              panel_b;
          nvcuda::wmma::load_matrix_sync(
              panel_a, &factor_half_hi[tile_row << 4][k],
              kHalfPanelStride128);
#pragma unroll
          for (int element = 0; element < panel_a.num_elements; ++element) {
            panel_a.x[element] = __hneg(panel_a.x[element]);
          }

          nvcuda::wmma::load_matrix_sync(
              panel_b, &factor_half_hi[tile_col << 4][k],
              kHalfPanelStride128);
          nvcuda::wmma::mma_sync(updated, panel_a, panel_b, updated);
          if constexpr (Compensated) {
            nvcuda::wmma::load_matrix_sync(
                panel_b, &factor_half_lo[tile_col << 4][k],
                kHalfPanelStride128);
            nvcuda::wmma::mma_sync(updated, panel_a, panel_b, updated);

            nvcuda::wmma::load_matrix_sync(
                panel_a, &factor_half_lo[tile_row << 4][k],
                kHalfPanelStride128);
#pragma unroll
            for (int element = 0; element < panel_a.num_elements; ++element) {
              panel_a.x[element] = __hneg(panel_a.x[element]);
            }

            nvcuda::wmma::mma_sync(updated, panel_a, panel_b, updated);
            nvcuda::wmma::load_matrix_sync(
                panel_b, &factor_half_hi[tile_col << 4][k],
                kHalfPanelStride128);
            nvcuda::wmma::mma_sync(updated, panel_a, panel_b, updated);
          }
        }

        nvcuda::wmma::store_matrix_sync(
            &tile[row][col], updated, kWmmaSharedStride128,
            nvcuda::wmma::mem_row_major);
      }
    }
    __syncthreads();
  }

  // Read back only the factor.  Upper entries bypass shared memory and are
  // written directly as zero, preserving the exact dense-output contract.
#pragma unroll
  for (int row = warp; row < kMatrixSize128; row += kWarpsPerBlock128) {
#pragma unroll
    for (int col = lane; col <= row; col += 32) {
      output[matrix_offset + row * kMatrixSize128 + col] = tile[row][col];
    }
#pragma unroll
    for (int col = row + 1 + lane; col < kMatrixSize128; col += 32) {
      output[matrix_offset + row * kMatrixSize128 + col] = 0.0f;
    }
  }
}

}  // namespace

torch::Tensor cholesky64_block32_hacc4safe_cuda(torch::Tensor input) {
  TORCH_CHECK(input.is_cuda(),
              "cholesky64_block32_hacc4safe_cuda expects a CUDA tensor");
  TORCH_CHECK(input.scalar_type() == torch::kFloat32,
              "cholesky64_block32_hacc4safe_cuda expects float32 input");
  TORCH_CHECK(input.dim() == 3 && input.size(0) > 0 &&
                  input.size(1) == kMatrixSize &&
                  input.size(2) == kMatrixSize,
              "cholesky64_block32_hacc4safe_cuda expects a nonempty "
              "[batch, 64, 64] input");
  TORCH_CHECK(input.is_contiguous(),
              "cholesky64_block32_hacc4safe_cuda expects contiguous input");

  const auto batch = input.size(0);
  auto output = torch::empty_like(input);
  cholesky64_block32_hacc4safe_kernel<<<batch, kThreads>>>(
      input.data_ptr<float>(), output.data_ptr<float>());
  C10_CUDA_KERNEL_LAUNCH_CHECK();
  return output;
}

torch::Tensor cholesky32_warp_cuda(torch::Tensor input) {
  TORCH_CHECK(input.is_cuda(), "cholesky32_warp_cuda expects a CUDA tensor");
  TORCH_CHECK(input.scalar_type() == torch::kFloat32,
              "cholesky32_warp_cuda expects float32 input");
  TORCH_CHECK(input.dim() == 3 && input.size(0) > 0 &&
                  input.size(1) == kMatrixSize32 &&
                  input.size(2) == kMatrixSize32,
              "cholesky32_warp_cuda expects a nonempty [batch, 32, 32] input");
  TORCH_CHECK(input.is_contiguous(),
              "cholesky32_warp_cuda expects contiguous input");

  const auto batch = input.size(0);
  TORCH_CHECK(batch <= std::numeric_limits<int>::max(),
              "cholesky32_warp_cuda batch is too large");
  const auto blocks =
      (batch + kWarpsPerBlock32 - 1) / kWarpsPerBlock32;
  auto output = torch::empty_like(input);
  cholesky32_lower_blocktiles_kernel<<<blocks, kThreads32>>>(
      input.data_ptr<float>(), output.data_ptr<float>(),
      static_cast<int>(batch));
  C10_CUDA_KERNEL_LAUNCH_CHECK();
  return output;
}

torch::Tensor cholesky128_wmma_schur_cuda(torch::Tensor input) {
  TORCH_CHECK(input.is_cuda(),
              "cholesky128_wmma_schur_cuda expects a CUDA tensor");
  TORCH_CHECK(input.scalar_type() == torch::kFloat32,
              "cholesky128_wmma_schur_cuda expects float32 input");
  TORCH_CHECK(input.dim() == 3 && input.size(0) > 0 &&
                  input.size(1) == kMatrixSize128 &&
                  input.size(2) == kMatrixSize128,
              "cholesky128_wmma_schur_cuda expects a nonempty "
              "[batch, 128, 128] input");
  TORCH_CHECK(input.is_contiguous(),
              "cholesky128_wmma_schur_cuda expects contiguous input");

  static const cudaError_t attribute_status = cudaFuncSetAttribute(
      cholesky128_wmma_schur_kernel<true>,
      cudaFuncAttributeMaxDynamicSharedMemorySize, kWmmaSharedBytes128);
  TORCH_CHECK(attribute_status == cudaSuccess,
              "cholesky128 shared-memory opt-in failed: ",
              cudaGetErrorString(attribute_status));

  const auto batch = input.size(0);
  auto output = torch::empty_like(input);
  cholesky128_wmma_schur_kernel<true>
      <<<batch, kThreads128, kWmmaSharedBytes128>>>(
      input.data_ptr<float>(), output.data_ptr<float>());
  C10_CUDA_KERNEL_LAUNCH_CHECK();
  return output;
}

torch::Tensor cholesky128_wmma_schur_exact_cuda(torch::Tensor input) {
  TORCH_CHECK(input.is_cuda(),
              "cholesky128_wmma_schur_exact_cuda expects a CUDA tensor");
  TORCH_CHECK(input.scalar_type() == torch::kFloat32,
              "cholesky128_wmma_schur_exact_cuda expects float32 input");
  TORCH_CHECK(input.dim() == 3 && input.size(0) == 256 &&
                  input.size(1) == kMatrixSize128 &&
                  input.size(2) == kMatrixSize128,
              "cholesky128_wmma_schur_exact_cuda expects "
              "[256, 128, 128] input");
  TORCH_CHECK(input.is_contiguous(),
              "cholesky128_wmma_schur_exact_cuda expects contiguous input");

  constexpr int shared_bytes =
      kWmmaFloatSharedBytes128 + kHalfPanelBytes128;
  static const cudaError_t attribute_status = cudaFuncSetAttribute(
      cholesky128_wmma_schur_kernel<false>,
      cudaFuncAttributeMaxDynamicSharedMemorySize, shared_bytes);
  TORCH_CHECK(attribute_status == cudaSuccess,
              "cholesky128 exact shared-memory opt-in failed: ",
              cudaGetErrorString(attribute_status));

  auto output = torch::empty_like(input);
  cholesky128_wmma_schur_kernel<false>
      <<<256, kThreads128, shared_bytes>>>(
          input.data_ptr<float>(), output.data_ptr<float>());
  C10_CUDA_KERNEL_LAUNCH_CHECK();
  return output;
}

torch::Tensor cholesky128_block32_warp16_lowerio_anybatch_cuda(
    torch::Tensor input) {
  TORCH_CHECK(input.is_cuda(),
              "cholesky128_block32_warp16_lowerio_anybatch_cuda expects a "
              "CUDA tensor");
  TORCH_CHECK(input.scalar_type() == torch::kFloat32,
              "cholesky128_block32_warp16_lowerio_anybatch_cuda expects "
              "float32 input");
  TORCH_CHECK(input.dim() == 3 && input.size(0) > 0 &&
                  input.size(1) == kMatrixSize128 &&
                  input.size(2) == kMatrixSize128,
              "cholesky128_block32_warp16_lowerio_anybatch_cuda expects a "
              "nonempty [batch, 128, 128] input");
  TORCH_CHECK(input.is_contiguous(),
              "cholesky128_block32_warp16_lowerio_anybatch_cuda expects "
              "contiguous input");

  static const cudaError_t attribute_status = cudaFuncSetAttribute(
      cholesky128_block32_warp16_lowerio_anybatch_kernel,
      cudaFuncAttributeMaxDynamicSharedMemorySize, kFloatSharedBytes128);
  TORCH_CHECK(attribute_status == cudaSuccess,
              "cholesky128 FP32 shared-memory opt-in failed: ",
              cudaGetErrorString(attribute_status));

  const auto batch = input.size(0);
  auto output = torch::empty_like(input);
  cholesky128_block32_warp16_lowerio_anybatch_kernel
      <<<batch, kThreads128, kFloatSharedBytes128>>>(
          input.data_ptr<float>(), output.data_ptr<float>());
  C10_CUDA_KERNEL_LAUNCH_CHECK();
  return output;
}

namespace {

constexpr int kCombinedSize256 = 256;
constexpr int kCombinedSize512 = 512;

inline void check_combined_solver(cusolverStatus_t status,
                                  const char* operation) {
  TORCH_CHECK(status == CUSOLVER_STATUS_SUCCESS, operation,
              " failed with cuSOLVER status ", static_cast<int>(status));
}

inline void check_combined_blas(cublasStatus_t status,
                                const char* operation) {
  TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, operation,
              " failed with cuBLAS status ", static_cast<int>(status));
}

void lt_packed_bf16_block_column(
    const __nv_bfloat16* packed,
    const __nv_bfloat16* packed_j,
    const float* source_j,
    float* destination_j,
    int rows,
    int cols,
    int inner,
    int output_ld,
    const float* alpha,
    const float* beta) {
  cublasLtMatmulDescOpaque_t operation_storage{};
  cublasLtMatmulDesc_t operation = &operation_storage;
  check_combined_blas(
      cublasLtMatmulDescInit(
          operation, CUBLAS_COMPUTE_32F, CUDA_R_32F),
      "cublasLtMatmulDescInit BF16 block-column update");
  const cublasOperation_t transpose = CUBLAS_OP_T;
  const cublasOperation_t identity = CUBLAS_OP_N;
  check_combined_blas(
      cublasLtMatmulDescSetAttribute(
          operation, CUBLASLT_MATMUL_DESC_TRANSA,
          &transpose, sizeof(transpose)),
      "cublasLt set transpose-A BF16 block-column update");
  check_combined_blas(
      cublasLtMatmulDescSetAttribute(
          operation, CUBLASLT_MATMUL_DESC_TRANSB,
          &identity, sizeof(identity)),
      "cublasLt set identity-B BF16 block-column update");

  cublasLtMatrixLayoutOpaque_t a_storage{};
  cublasLtMatrixLayoutOpaque_t b_storage{};
  cublasLtMatrixLayoutOpaque_t c_storage{};
  cublasLtMatrixLayoutOpaque_t d_storage{};
  cublasLtMatrixLayout_t a_layout = &a_storage;
  cublasLtMatrixLayout_t b_layout = &b_storage;
  cublasLtMatrixLayout_t c_layout = &c_storage;
  cublasLtMatrixLayout_t d_layout = &d_storage;
  check_combined_blas(
      cublasLtMatrixLayoutInit(a_layout, CUDA_R_16BF, inner, rows, inner),
      "cublasLt A-layout init BF16 block-column update");
  check_combined_blas(
      cublasLtMatrixLayoutInit(b_layout, CUDA_R_16BF, inner, cols, inner),
      "cublasLt B-layout init BF16 block-column update");
  check_combined_blas(
      cublasLtMatrixLayoutInit(
          c_layout, CUDA_R_32F, rows, cols, output_ld),
      "cublasLt C-layout init BF16 block-column update");
  check_combined_blas(
      cublasLtMatrixLayoutInit(
          d_layout, CUDA_R_32F, rows, cols, output_ld),
      "cublasLt D-layout init BF16 block-column update");

  cublasLtMatmulPreferenceOpaque_t preference_storage{};
  cublasLtMatmulPreference_t preference = &preference_storage;
  check_combined_blas(
      cublasLtMatmulPreferenceInit(preference),
      "cublasLt preference init BF16 block-column update");
  const size_t scratch_bytes = 0;
  check_combined_blas(
      cublasLtMatmulPreferenceSetAttribute(
          preference, CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES,
          &scratch_bytes, sizeof(scratch_bytes)),
      "cublasLt scratch limit BF16 block-column update");

  cublasLtMatmulHeuristicResult_t heuristic{};
  int returned = 0;
  cublasLtHandle_t handle = at::cuda::getCurrentCUDABlasLtHandle();
  check_combined_blas(
      cublasLtMatmulAlgoGetHeuristic(
          handle, operation, a_layout, b_layout,
          c_layout, d_layout, preference, 1,
          &heuristic, &returned),
      "cublasLt heuristic BF16 block-column update");
  TORCH_CHECK(
      returned == 1 && heuristic.state == CUBLAS_STATUS_SUCCESS &&
          heuristic.workspaceSize == 0,
      "cublasLt found no scratch-free BF16 block-column algorithm");

  check_combined_blas(
      cublasLtMatmul(
          handle, operation, alpha, packed, a_layout,
          packed_j, b_layout, beta, source_j, c_layout,
          destination_j, d_layout, &heuristic.algo, nullptr, 0, nullptr),
      "cublasLtMatmul BF16 block-column update");
}

struct Mxfp8LtCacheKey {
  int device;
  int rows;
  int cols;
  int inner;
  int output_ld;
  cudaDataType_t a_type;
  cudaDataType_t b_type;
  cudaDataType_t c_type;
  cudaDataType_t d_type;
  cublasComputeType_t compute_type;
  cudaDataType_t scale_type;
  cublasOperation_t transa;
  cublasOperation_t transb;
  cublasLtMatmulMatrixScale_t a_scale_mode;
  cublasLtMatmulMatrixScale_t b_scale_mode;
  int8_t fast_accumulation;
  int a_alignment;
  int b_alignment;
  int c_alignment;
  int d_alignment;
  int a_scale_alignment;
  int b_scale_alignment;
  int workspace_alignment;
  int alpha_alignment;
  int beta_alignment;
  size_t scratch_bytes;

  bool operator==(const Mxfp8LtCacheKey& other) const {
    return device == other.device && rows == other.rows &&
           cols == other.cols && inner == other.inner &&
           output_ld == other.output_ld && a_type == other.a_type &&
           b_type == other.b_type && c_type == other.c_type &&
           d_type == other.d_type && compute_type == other.compute_type &&
           scale_type == other.scale_type && transa == other.transa &&
           transb == other.transb && a_scale_mode == other.a_scale_mode &&
           b_scale_mode == other.b_scale_mode &&
           fast_accumulation == other.fast_accumulation &&
           a_alignment == other.a_alignment &&
           b_alignment == other.b_alignment &&
           c_alignment == other.c_alignment &&
           d_alignment == other.d_alignment &&
           a_scale_alignment == other.a_scale_alignment &&
           b_scale_alignment == other.b_scale_alignment &&
           workspace_alignment == other.workspace_alignment &&
           alpha_alignment == other.alpha_alignment &&
           beta_alignment == other.beta_alignment &&
           scratch_bytes == other.scratch_bytes;
  }
};

struct Mxfp8LtCacheEntry {
  Mxfp8LtCacheKey key;
  cublasLtMatrixLayoutOpaque_t a_storage{};
  cublasLtMatrixLayoutOpaque_t b_storage{};
  cublasLtMatrixLayoutOpaque_t c_storage{};
  cublasLtMatrixLayoutOpaque_t d_storage{};
  cublasLtMatmulPreferenceOpaque_t preference_storage{};
  cublasLtMatmulHeuristicResult_t heuristic{};

  cublasLtMatrixLayout_t a_layout() { return &a_storage; }
  cublasLtMatrixLayout_t b_layout() { return &b_storage; }
  cublasLtMatrixLayout_t c_layout() { return &c_storage; }
  cublasLtMatrixLayout_t d_layout() { return &d_storage; }
  cublasLtMatmulPreference_t preference() { return &preference_storage; }
};

constexpr size_t kMxfp8LtCacheCapacity = 64;

int capped_pointer_alignment(const void* pointer) {
  const uintptr_t address = reinterpret_cast<uintptr_t>(pointer);
  int alignment = 1;
  while (alignment < 256 && (address & (alignment * 2 - 1)) == 0) {
    alignment *= 2;
  }
  return alignment;
}

Mxfp8LtCacheEntry* find_mxfp8_lt_cache_entry(
    std::vector<std::unique_ptr<Mxfp8LtCacheEntry>>& cache,
    const Mxfp8LtCacheKey& key) {
  for (auto& entry : cache) {
    if (entry->key == key) {
      return entry.get();
    }
  }
  return nullptr;
}

Mxfp8LtCacheEntry* create_mxfp8_lt_cache_entry(
    std::vector<std::unique_ptr<Mxfp8LtCacheEntry>>& cache,
    const Mxfp8LtCacheKey& key,
    cublasLtHandle_t handle,
    cublasLtMatmulDesc_t operation) {
  auto entry = std::make_unique<Mxfp8LtCacheEntry>();
  entry->key = key;
  check_combined_blas(
      cublasLtMatrixLayoutInit(
          entry->a_layout(), key.a_type, key.inner, key.rows, key.inner),
      "cublasLt cached A-layout init MXFP8 block-column update");
  check_combined_blas(
      cublasLtMatrixLayoutInit(
          entry->b_layout(), key.b_type, key.inner, key.cols, key.inner),
      "cublasLt cached B-layout init MXFP8 block-column update");
  check_combined_blas(
      cublasLtMatrixLayoutInit(
          entry->c_layout(), key.c_type, key.rows, key.cols, key.output_ld),
      "cublasLt cached C-layout init MXFP8 block-column update");
  check_combined_blas(
      cublasLtMatrixLayoutInit(
          entry->d_layout(), key.d_type, key.rows, key.cols, key.output_ld),
      "cublasLt cached D-layout init MXFP8 block-column update");
  check_combined_blas(
      cublasLtMatmulPreferenceInit(entry->preference()),
      "cublasLt cached preference init MXFP8 block-column update");
  check_combined_blas(
      cublasLtMatmulPreferenceSetAttribute(
          entry->preference(), CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES,
          &key.scratch_bytes, sizeof(key.scratch_bytes)),
      "cublasLt cached scratch limit MXFP8 block-column update");

  int returned = 0;
  check_combined_blas(
      cublasLtMatmulAlgoGetHeuristic(
          handle, operation, entry->a_layout(), entry->b_layout(),
          entry->c_layout(), entry->d_layout(), entry->preference(), 1,
          &entry->heuristic, &returned),
      "cublasLt cached heuristic MXFP8 block-column update");
  TORCH_CHECK(
      returned == 1 && entry->heuristic.state == CUBLAS_STATUS_SUCCESS &&
          entry->heuristic.workspaceSize <= key.scratch_bytes,
      "cublasLt found no cached bounded-workspace MXFP8 algorithm");

  if (cache.size() == kMxfp8LtCacheCapacity) {
    cache.erase(cache.begin());
  }
  Mxfp8LtCacheEntry* const result = entry.get();
  cache.push_back(std::move(entry));
  return result;
}

void lt_packed_mxfp8_block_column(
    const __nv_fp8_e4m3* primary,
    const __nv_fp8_e4m3* primary_j,
    const __nv_fp8_e8m0* primary_scale,
    const __nv_fp8_e8m0* primary_scale_j,
    const float* source_j,
    float* destination_j,
    int rows,
    int cols,
    int inner,
    int output_ld,
    int device,
    std::vector<std::unique_ptr<Mxfp8LtCacheEntry>>& cache,
    void* scratch,
    size_t scratch_bytes,
    const float* alpha,
    const float* beta) {
  cublasLtMatmulDescOpaque_t operation_storage{};
  cublasLtMatmulDesc_t operation = &operation_storage;
  check_combined_blas(
      cublasLtMatmulDescInit(
          operation, CUBLAS_COMPUTE_32F, CUDA_R_32F),
      "cublasLtMatmulDescInit MXFP8 block-column update");
  const cublasOperation_t transpose = CUBLAS_OP_T;
  const cublasOperation_t identity = CUBLAS_OP_N;
  check_combined_blas(
      cublasLtMatmulDescSetAttribute(
          operation, CUBLASLT_MATMUL_DESC_TRANSA,
          &transpose, sizeof(transpose)),
      "cublasLt set transpose-A MXFP8 block-column update");
  check_combined_blas(
      cublasLtMatmulDescSetAttribute(
          operation, CUBLASLT_MATMUL_DESC_TRANSB,
          &identity, sizeof(identity)),
      "cublasLt set identity-B MXFP8 block-column update");
  const cublasLtMatmulMatrixScale_t block_scale =
      CUBLASLT_MATMUL_MATRIX_SCALE_VEC32_UE8M0;
  check_combined_blas(
      cublasLtMatmulDescSetAttribute(
          operation, CUBLASLT_MATMUL_DESC_A_SCALE_MODE,
          &block_scale, sizeof(block_scale)),
      "cublasLt set A block scale mode MXFP8 block-column update");
  check_combined_blas(
      cublasLtMatmulDescSetAttribute(
          operation, CUBLASLT_MATMUL_DESC_B_SCALE_MODE,
          &block_scale, sizeof(block_scale)),
      "cublasLt set B block scale mode MXFP8 block-column update");
  const int8_t standard_accumulation = 0;
  check_combined_blas(
      cublasLtMatmulDescSetAttribute(
          operation, CUBLASLT_MATMUL_DESC_FAST_ACCUM,
          &standard_accumulation, sizeof(standard_accumulation)),
      "cublasLt set FP32 accumulation MXFP8 block-column update");

  const __nv_fp8_e8m0* const a_scale_pointer = primary_scale;
  const __nv_fp8_e8m0* const b_scale_pointer = primary_scale_j;
  check_combined_blas(
      cublasLtMatmulDescSetAttribute(
          operation, CUBLASLT_MATMUL_DESC_A_SCALE_POINTER,
          &a_scale_pointer, sizeof(a_scale_pointer)),
      "cublasLt set A scale pointer MXFP8 block-column update");
  check_combined_blas(
      cublasLtMatmulDescSetAttribute(
          operation, CUBLASLT_MATMUL_DESC_B_SCALE_POINTER,
          &b_scale_pointer, sizeof(b_scale_pointer)),
      "cublasLt set B scale pointer MXFP8 block-column update");

  cublasLtHandle_t handle = at::cuda::getCurrentCUDABlasLtHandle();
  const Mxfp8LtCacheKey cache_key{
      device,
      rows,
      cols,
      inner,
      output_ld,
      CUDA_R_8F_E4M3,
      CUDA_R_8F_E4M3,
      CUDA_R_32F,
      CUDA_R_32F,
      CUBLAS_COMPUTE_32F,
      CUDA_R_32F,
      transpose,
      identity,
      block_scale,
      block_scale,
      standard_accumulation,
      capped_pointer_alignment(primary),
      capped_pointer_alignment(primary_j),
      capped_pointer_alignment(source_j),
      capped_pointer_alignment(destination_j),
      capped_pointer_alignment(primary_scale),
      capped_pointer_alignment(primary_scale_j),
      capped_pointer_alignment(scratch),
      capped_pointer_alignment(alpha),
      capped_pointer_alignment(beta),
      scratch_bytes};
  Mxfp8LtCacheEntry* cache_entry =
      find_mxfp8_lt_cache_entry(cache, cache_key);
  if (cache_entry == nullptr) {
    cache_entry = create_mxfp8_lt_cache_entry(
        cache, cache_key, handle, operation);
  }
  cublasLtMatrixLayout_t a_layout = cache_entry->a_layout();
  cublasLtMatrixLayout_t b_layout = cache_entry->b_layout();
  cublasLtMatrixLayout_t c_layout = cache_entry->c_layout();
  cublasLtMatrixLayout_t d_layout = cache_entry->d_layout();
  const cublasLtMatmulAlgo_t* const algorithm =
      &cache_entry->heuristic.algo;

  check_combined_blas(
      cublasLtMatmul(
          handle, operation, alpha, primary, a_layout,
          primary_j, b_layout, beta, source_j, c_layout,
          destination_j, d_layout, algorithm, scratch,
          scratch_bytes, nullptr),
      "cublasLtMatmul MXFP8 primary block-column update");
}

struct CombinedSolverState {
  cusolverDnHandle_t batched_handle = nullptr;
  cusolverDnHandle_t large_handle = nullptr;
  torch::Tensor pointers;
  torch::Tensor batched_info;
  torch::Tensor workspace;
  torch::Tensor large_info;
  int pointer_capacity = 0;
  int workspace_elements = 0;
  int device = -1;
  std::mutex mutex;

  ~CombinedSolverState() {
    if (batched_handle != nullptr) {
      cusolverDnDestroy(batched_handle);
    }
    if (large_handle != nullptr) {
      cusolverDnDestroy(large_handle);
    }
  }

  void prepare_device(const torch::Tensor& input) {
    const int input_device = input.get_device();
    if (device == input_device && batched_handle != nullptr &&
        large_handle != nullptr) {
      return;
    }

    if (batched_handle != nullptr) {
      cusolverDnDestroy(batched_handle);
      batched_handle = nullptr;
    }
    if (large_handle != nullptr) {
      cusolverDnDestroy(large_handle);
      large_handle = nullptr;
    }

    check_combined_solver(cusolverDnCreate(&batched_handle),
                          "cusolverDnCreate batched");
    check_combined_solver(cusolverDnCreate(&large_handle),
                          "cusolverDnCreate large");
#if CUSOLVER_VERSION >= 12000
    check_combined_solver(
        cusolverDnSetMathMode(
            large_handle, CUSOLVER_FP32_EMULATED_BF16X9_MATH),
        "cusolverDnSetMathMode");
    check_combined_solver(
        cusolverDnSetEmulationStrategy(
            large_handle, CUDA_EMULATION_STRATEGY_EAGER),
        "cusolverDnSetEmulationStrategy");
#endif

    pointers = torch::Tensor();
    batched_info = torch::Tensor();
    workspace = torch::Tensor();
    large_info = torch::Tensor();
    pointer_capacity = 0;
    workspace_elements = 0;
    device = input_device;
  }

  void prepare_batched(const torch::Tensor& input, int batch) {
    prepare_device(input);
    if (pointer_capacity < batch) {
      pointers = torch::empty(
          {batch}, input.options().dtype(torch::kInt64));
      batched_info = torch::empty(
          {batch}, input.options().dtype(torch::kInt32));
      pointer_capacity = batch;
    }
  }
};

CombinedSolverState& combined_solver_state() {
  static CombinedSolverState state;
  return state;
}

struct ExplicitLowPrecisionState {
  cusolverDnHandle_t solver_handle = nullptr;
  cublasHandle_t blas_handle = nullptr;
  torch::Tensor workspace;
  torch::Tensor info;
  torch::Tensor panel_bf16;
  torch::Tensor mxfp8_primary_scale;
  torch::Tensor mxfp8_diagonal_correction;
  torch::Tensor mxfp8_scratch;
  std::vector<std::unique_ptr<Mxfp8LtCacheEntry>> mxfp8_lt_cache;
  std::vector<torch::Tensor> retired_buffers;
  int workspace_elements = 0;
  long long panel_capacity = 0;
  int device = -1;
  std::mutex mutex;

  ~ExplicitLowPrecisionState() {
    if (solver_handle != nullptr) {
      cusolverDnDestroy(solver_handle);
    }
    if (blas_handle != nullptr) {
      cublasDestroy(blas_handle);
    }
  }

  void prepare(const torch::Tensor& input) {
    const int input_device = input.get_device();
    if (device == input_device && solver_handle != nullptr &&
        blas_handle != nullptr) {
      return;
    }
    if (solver_handle != nullptr) {
      cusolverDnDestroy(solver_handle);
    }
    if (blas_handle != nullptr) {
      cublasDestroy(blas_handle);
    }
    check_combined_solver(cusolverDnCreate(&solver_handle),
                          "cusolverDnCreate explicit BF16");
    check_combined_blas(cublasCreate(&blas_handle),
                        "cublasCreate explicit BF16");
    workspace = torch::Tensor();
    info = torch::Tensor();
    panel_bf16 = torch::Tensor();
    mxfp8_primary_scale = torch::Tensor();
    mxfp8_diagonal_correction = torch::Tensor();
    mxfp8_scratch = torch::Tensor();
    mxfp8_lt_cache.clear();
    retired_buffers.clear();
    workspace_elements = 0;
    panel_capacity = 0;
    device = input_device;
  }
};

ExplicitLowPrecisionState& explicit_low_precision_state() {
  static ExplicitLowPrecisionState state;
  return state;
}

struct TrueBatchedFast16BFState {
  cusolverDnHandle_t solver = nullptr;
  cublasHandle_t blas = nullptr;
  torch::Tensor pointers;
  torch::Tensor info;
  std::vector<torch::Tensor> retired_buffers;
  int pointer_capacity = 0;
  int device = -1;
  std::mutex mutex;

  ~TrueBatchedFast16BFState() {
    if (solver != nullptr) {
      cusolverDnDestroy(solver);
    }
    if (blas != nullptr) {
      cublasDestroy(blas);
    }
  }

  void prepare(const torch::Tensor& input, int batch) {
    const int input_device = input.get_device();
    if (device == -1) {
      check_combined_solver(cusolverDnCreate(&solver),
                            "cusolverDnCreate true-batched fast16bf");
      check_combined_blas(cublasCreate(&blas),
                          "cublasCreate true-batched fast16bf");
      check_combined_blas(
          cublasSetPointerMode(blas, CUBLAS_POINTER_MODE_HOST),
          "cublasSetPointerMode true-batched fast16bf");
      device = input_device;
    } else {
      TORCH_CHECK(device == input_device && solver != nullptr &&
                      blas != nullptr,
                  "true-batched fast16bf state cannot migrate between CUDA "
                  "devices");
    }

    if (pointer_capacity < batch) {
      if (pointers.defined()) {
        retired_buffers.push_back(pointers);
      }
      if (info.defined()) {
        retired_buffers.push_back(info);
      }
      pointers = torch::empty(
          {2, batch}, input.options().dtype(torch::kInt64));
      info = torch::empty(
          {batch}, input.options().dtype(torch::kInt32));
      pointer_capacity = batch;
    }
  }
};

TrueBatchedFast16BFState& true_batched_fast16bf_state() {
  static TrueBatchedFast16BFState state;
  return state;
}

template <int MatrixSize>
__global__ void initialize_combined_batched_triangle(
    const float* __restrict__ input, float* __restrict__ output,
    float** pointers, int batch, long long total) {
  constexpr long long matrix_elements =
      static_cast<long long>(MatrixSize) * MatrixSize;
  long long linear =
      static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
  const long long stride =
      static_cast<long long>(gridDim.x) * blockDim.x;

  for (long long matrix = linear; matrix < batch; matrix += stride) {
    pointers[matrix] = output + matrix * matrix_elements;
  }

  for (; linear < total; linear += stride) {
    const int within = static_cast<int>(linear % matrix_elements);
    const int row = within / MatrixSize;
    const int col = within - row * MatrixSize;
    if (col <= row) {
      output[linear] = input[linear];
    }
  }
}

template <int MatrixSize>
__global__ void fill_highbatch_panel_pointers(
    float* base, float** diagonal_pointers, float** panel_pointers,
    int diagonal_offset, int panel_offset, int batch) {
  const int matrix =
      static_cast<int>(blockIdx.x) * blockDim.x + threadIdx.x;
  if (matrix < batch) {
    constexpr long long matrix_elements =
        static_cast<long long>(MatrixSize) * MatrixSize;
    float* const matrix_base =
        base + static_cast<long long>(matrix) * matrix_elements;
    diagonal_pointers[matrix] = matrix_base + diagonal_offset;
    panel_pointers[matrix] = matrix_base + panel_offset;
  }
}

template <int MatrixSize>
__global__ void copy_highbatch_lower_zero_upper(
    const float* __restrict__ input, float* __restrict__ output,
    long long total) {
  constexpr long long matrix_elements =
      static_cast<long long>(MatrixSize) * MatrixSize;
  long long linear =
      static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
  const long long stride =
      static_cast<long long>(gridDim.x) * blockDim.x;
  for (; linear < total; linear += stride) {
    const int within = static_cast<int>(linear % matrix_elements);
    const int row = within / MatrixSize;
    const int col = within - row * MatrixSize;
    output[linear] = col <= row ? input[linear] : 0.0f;
  }
}

template <int MatrixSize>
__global__ void copy_highbatch_lower_only(
    const float* __restrict__ input, float* __restrict__ output,
    long long total) {
  constexpr long long matrix_elements =
      static_cast<long long>(MatrixSize) * MatrixSize;
  long long linear =
      static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
  const long long stride =
      static_cast<long long>(gridDim.x) * blockDim.x;
  for (; linear < total; linear += stride) {
    const int within = static_cast<int>(linear % matrix_elements);
    const int row = within / MatrixSize;
    const int col = within - row * MatrixSize;
    if (col <= row) {
      output[linear] = input[linear];
    }
  }
}

template <int MatrixSize>
__global__ void copy_highbatch_lower_only_float4(
    const float* __restrict__ input, float* __restrict__ output,
    long long vector_total) {
  static_assert((MatrixSize & 3) == 0,
                "float4 triangular copy requires four-column rows");
  constexpr int vectors_per_row = MatrixSize / 4;
  constexpr long long vectors_per_matrix =
      static_cast<long long>(MatrixSize) * vectors_per_row;
  const auto* input4 = reinterpret_cast<const float4*>(input);
  auto* output4 = reinterpret_cast<float4*>(output);
  long long vector_linear =
      static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
  const long long stride =
      static_cast<long long>(gridDim.x) * blockDim.x;
  for (; vector_linear < vector_total; vector_linear += stride) {
    const int within =
        static_cast<int>(vector_linear % vectors_per_matrix);
    const int row = within / vectors_per_row;
    const int col = (within - row * vectors_per_row) * 4;
    if (col + 3 <= row) {
      output4[vector_linear] = input4[vector_linear];
    } else if (col <= row) {
      const long long scalar = vector_linear * 4;
#pragma unroll
      for (int lane = 0; lane < 4; ++lane) {
        if (col + lane <= row) {
          output[scalar + lane] = input[scalar + lane];
        }
      }
    }
  }
}

template <int MatrixSize>
__global__ void clear_combined_batched_row_major_upper(float* output,
                                                       long long total) {
  constexpr long long matrix_elements =
      static_cast<long long>(MatrixSize) * MatrixSize;
  long long linear =
      static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
  const long long stride =
      static_cast<long long>(gridDim.x) * blockDim.x;
  for (; linear < total; linear += stride) {
    const int within = static_cast<int>(linear % matrix_elements);
    const int row = within / MatrixSize;
    const int col = within - row * MatrixSize;
    if (col > row) {
      output[linear] = 0.0f;
    }
  }
}

template <int MatrixSize>
__global__ void clear_highbatch_upper_float4(float* output,
                                             long long vector_total) {
  static_assert((MatrixSize & 3) == 0,
                "float4 triangular clear requires four-column rows");
  constexpr int vectors_per_row = MatrixSize / 4;
  constexpr long long vectors_per_matrix =
      static_cast<long long>(MatrixSize) * vectors_per_row;
  auto* output4 = reinterpret_cast<float4*>(output);
  long long vector_linear =
      static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
  const long long stride =
      static_cast<long long>(gridDim.x) * blockDim.x;
  for (; vector_linear < vector_total; vector_linear += stride) {
    const int within =
        static_cast<int>(vector_linear % vectors_per_matrix);
    const int row = within / vectors_per_row;
    const int col = (within - row * vectors_per_row) * 4;
    if (col > row) {
      output4[vector_linear] = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
    } else if (col + 3 > row) {
      const long long scalar = vector_linear * 4;
#pragma unroll
      for (int lane = 0; lane < 4; ++lane) {
        if (col + lane > row) {
          output[scalar + lane] = 0.0f;
        }
      }
    }
  }
}

__global__ void clear_combined_large_upper(float* output, int n) {
  const long long total = static_cast<long long>(n) * n;
  long long linear =
      static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
  const long long stride =
      static_cast<long long>(gridDim.x) * blockDim.x;
  for (; linear < total; linear += stride) {
    const int row = static_cast<int>(linear / n);
    const int col =
        static_cast<int>(linear - static_cast<long long>(row) * n);
    if (col > row) {
      output[linear] = 0.0f;
    }
  }
}

__global__ void clear_explicit_batch_upper(float* output, int n,
                                           long long total) {
  const long long matrix_elements = static_cast<long long>(n) * n;
  long long linear =
      static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
  const long long stride =
      static_cast<long long>(gridDim.x) * blockDim.x;
  for (; linear < total; linear += stride) {
    const long long within = linear % matrix_elements;
    const int row = static_cast<int>(within / n);
    const int col = static_cast<int>(
        within - static_cast<long long>(row) * n);
    if (col > row) {
      output[linear] = 0.0f;
    }
  }
}

__global__ void clear_explicit_large_upper_vectorized(float* output, int n) {
  const int row = static_cast<int>(blockIdx.x);
  const int upper_begin = row + 1;
  const int vector_begin = (upper_begin + 3) & ~3;
  const int prefix_end = vector_begin < n ? vector_begin : n;
  const long long row_offset = static_cast<long long>(row) * n;
  float* const output_row = output + row_offset;

  for (int col = upper_begin + static_cast<int>(threadIdx.x);
       col < prefix_end; col += static_cast<int>(blockDim.x)) {
    output_row[col] = 0.0f;
  }

  const int vector_count = (n - vector_begin) >> 2;
  auto* const output_vectors =
      reinterpret_cast<float4*>(output_row + vector_begin);
  for (int vector = static_cast<int>(threadIdx.x);
       vector < vector_count;
       vector += static_cast<int>(blockDim.x)) {
    output_vectors[vector] = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
  }

  const int tail_begin = vector_begin + (vector_count << 2);
  for (int col = tail_begin + static_cast<int>(threadIdx.x); col < n;
       col += static_cast<int>(blockDim.x)) {
    output_row[col] = 0.0f;
  }
}

__global__ void initialize_explicit_large_first_panel_lower(
    const float* input, float* output, int n, int initialized_cols) {
  const int row = static_cast<int>(blockIdx.x);
  const long long row_offset = static_cast<long long>(row) * n;
  for (int col = static_cast<int>(threadIdx.x) * 4;
       col < initialized_cols;
       col += static_cast<int>(blockDim.x) * 4) {
    const long long first = row_offset + col;
    if (col + 3 < initialized_cols && col + 3 <= row) {
      reinterpret_cast<float4*>(output + first)[0] =
          reinterpret_cast<const float4*>(input + first)[0];
    } else {
#pragma unroll
      for (int lane = 0; lane < 4; ++lane) {
        const int scalar_col = col + lane;
        if (scalar_col < initialized_cols && scalar_col <= row) {
          output[row_offset + scalar_col] = input[row_offset + scalar_col];
        }
      }
    }
  }
}

__global__ void pack_panel_bf16(const float* panel, __nv_bfloat16* packed,
                                int rows, int cols, int source_ld) {
  const long long total = static_cast<long long>(rows) * cols;
  long long linear =
      static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
  const long long stride =
      static_cast<long long>(gridDim.x) * blockDim.x;
  for (; linear < total; linear += stride) {
    const int row = static_cast<int>(linear % rows);
    const int col = static_cast<int>(linear / rows);
    packed[linear] = __float2bfloat16_rn(
        panel[row + static_cast<long long>(col) * source_ld]);
  }
}

__global__ void pack_outer_panel_bf16(
    const float* panel, __nv_bfloat16* packed, int rows, int cols,
    int source_ld, int packed_ld, int packed_row, int packed_col) {
  const long long total = static_cast<long long>(rows) * cols;
  long long linear =
      static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
  const long long stride =
      static_cast<long long>(gridDim.x) * blockDim.x;
  for (; linear < total; linear += stride) {
    const int row = static_cast<int>(linear % rows);
    const int col = static_cast<int>(linear / rows);
    packed[packed_row + row +
           static_cast<long long>(packed_col + col) * packed_ld] =
        __float2bfloat16_rn(
            panel[row + static_cast<long long>(col) * source_ld]);
  }
}

union alignas(16) PackedOuterBf16x8 {
  struct {
    __nv_bfloat162 pair0;
    __nv_bfloat162 pair1;
    __nv_bfloat162 pair2;
    __nv_bfloat162 pair3;
  } pairs;
  uint4 vector;
};

static_assert(sizeof(PackedOuterBf16x8) == 16,
              "eight packed BF16 values must occupy 16 bytes");

__global__ void pack_outer_panel_bf16_vector8(
    const float* panel, __nv_bfloat16* packed, int rows, int cols,
    int source_ld, int packed_ld, int packed_row, int packed_col) {
  const int row_groups = rows / 8;
  const long long total = static_cast<long long>(row_groups) * cols;
  long long linear =
      static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
  const long long stride =
      static_cast<long long>(gridDim.x) * blockDim.x;
  for (; linear < total; linear += stride) {
    const int row_group = static_cast<int>(linear % row_groups);
    const int col = static_cast<int>(linear / row_groups);
    const int row = row_group * 8;
    const float* const source =
        panel + row + static_cast<long long>(col) * source_ld;
    const float4 first = reinterpret_cast<const float4*>(source)[0];
    const float4 second = reinterpret_cast<const float4*>(source)[1];
    PackedOuterBf16x8 converted;
    converted.pairs.pair0 = __floats2bfloat162_rn(first.x, first.y);
    converted.pairs.pair1 = __floats2bfloat162_rn(first.z, first.w);
    converted.pairs.pair2 = __floats2bfloat162_rn(second.x, second.y);
    converted.pairs.pair3 = __floats2bfloat162_rn(second.z, second.w);
    auto* const destination = reinterpret_cast<uint4*>(
        packed + packed_row + row +
        static_cast<long long>(packed_col + col) * packed_ld);
    destination[0] = converted.vector;
  }
}

__device__ __forceinline__ int mxfp8_e8m0_exponent(float maximum) {
  const unsigned bits = __float_as_uint(maximum);
  const int exponent = static_cast<int>((bits >> 23) & 255U);
  const unsigned mantissa = bits & 0x7FFFFFU;
  const int rounded = exponent + (mantissa > 0x600000U) - 8;
  return max(1, min(253, rounded));
}

__device__ __forceinline__ long long mxfp8_scale_offset(
    int group, int column, int inner) {
  constexpr int scale_block_elements = 512;
  const int scaled_rows = (inner + 31) / 32;
  const int scaled_ld = (scaled_rows + 3) & ~3;
  const long long block =
      static_cast<long long>(group / 4) +
      static_cast<long long>(column / 128) * (scaled_ld / 4);
  return block * scale_block_elements +
         static_cast<long long>(column & 31) * 16 +
         static_cast<long long>((column / 32) & 3) * 4 +
         (group & 3);
}

inline long long mxfp8_scale_column_base(int inner, int column) {
  const int scaled_rows = (inner + 31) / 32;
  const int scaled_ld = (scaled_rows + 3) & ~3;
  return static_cast<long long>(column / 128) *
         (scaled_ld / 4) * 512;
}

__global__ void clear_mxfp8_diagonal_correction(
    float* correction, int count) {
  for (int index = static_cast<int>(blockIdx.x) * blockDim.x +
                       static_cast<int>(threadIdx.x);
       index < count;
       index += static_cast<int>(gridDim.x) * blockDim.x) {
    correction[index] = 0.0f;
  }
}

__global__ void pack_outer_far_mxfp8_block32(
    const float* panel,
    __nv_fp8_e4m3* primary,
    __nv_fp8_e8m0* primary_scale,
    float* diagonal_correction,
    int inner,
    int columns,
    int source_ld) {
  const int column = static_cast<int>(blockIdx.x);
  if (column >= columns) {
    return;
  }
  const int lane = static_cast<int>(threadIdx.x) & 31;
  const int warp = static_cast<int>(threadIdx.x) >> 5;
  const int warps = static_cast<int>(blockDim.x) >> 5;
  const int groups = inner / 32;
  for (int group = warp; group < groups; group += warps) {
    const int row = group * 32 + lane;
    const long long packed_offset =
        row + static_cast<long long>(column) * inner;
    const float value =
        panel[row + static_cast<long long>(column) * source_ld];
    float maximum = fabsf(value);
#pragma unroll
    for (int delta = 16; delta > 0; delta >>= 1) {
      maximum = fmaxf(
          maximum, __shfl_down_sync(0xFFFFFFFFU, maximum, delta));
    }
    maximum = __shfl_sync(0xFFFFFFFFU, maximum, 0);
    const int primary_exponent = mxfp8_e8m0_exponent(maximum);
    const float primary_value_scale =
        __uint_as_float(static_cast<unsigned>(primary_exponent) << 23);
    const float primary_inverse = __uint_as_float(
        static_cast<unsigned>(254 - primary_exponent) << 23);
    const __nv_fp8_e4m3 primary_value(value * primary_inverse);
    primary[packed_offset] = primary_value;
    const float primary_reconstruction =
        static_cast<float>(primary_value) * primary_value_scale;

    if (lane == 0) {
      const long long scale_offset =
          mxfp8_scale_offset(group, column, inner);
      __nv_fp8_e8m0 primary_scale_value;
      primary_scale_value.__x =
          static_cast<__nv_fp8_storage_t>(primary_exponent);
      primary_scale[scale_offset] = primary_scale_value;
    }

    float correction =
        value * value - primary_reconstruction * primary_reconstruction;
#pragma unroll
    for (int delta = 16; delta > 0; delta >>= 1) {
      correction +=
          __shfl_down_sync(0xFFFFFFFFU, correction, delta);
    }
    if (lane == 0) {
      atomicAdd(diagonal_correction + column, correction);
    }
  }
}

__global__ void add_mxfp8_diagonal_correction(
    float* diagonal, const float* correction, int count, int ld) {
  for (int index = static_cast<int>(blockIdx.x) * blockDim.x +
                       static_cast<int>(threadIdx.x);
       index < count;
       index += static_cast<int>(gridDim.x) * blockDim.x) {
    diagonal[static_cast<long long>(index) * (ld + 1)] +=
        correction[index];
  }
}

__global__ void copy_combined_large_lower(const float* input, float* output,
                                          int n) {
  const int row = static_cast<int>(blockIdx.x);
  const long long row_offset = static_cast<long long>(row) * n;
  for (int col = static_cast<int>(threadIdx.x) * 4; col <= row;
       col += static_cast<int>(blockDim.x) * 4) {
    const long long first = row_offset + col;
    if ((n & 3) == 0 && col + 3 <= row) {
      reinterpret_cast<float4*>(output + first)[0] =
          reinterpret_cast<const float4*>(input + first)[0];
    } else {
#pragma unroll
      for (int lane = 0; lane < 4; ++lane) {
        const int scalar_col = col + lane;
        if (scalar_col <= row) {
          output[row_offset + scalar_col] = input[row_offset + scalar_col];
        }
      }
    }
  }
}

template <int MatrixSize>
torch::Tensor cholesky_combined_batched_impl(torch::Tensor input) {
  c10::cuda::CUDAGuard device_guard(input.device());
  const int batch = static_cast<int>(input.size(0));
  auto output = torch::empty_like(input);
  auto& state = combined_solver_state();
  std::lock_guard<std::mutex> lock(state.mutex);
  state.prepare_batched(input, batch);

  constexpr int threads = 256;
  auto pointer_data = reinterpret_cast<float**>(
      state.pointers.data_ptr<int64_t>());
  constexpr long long matrix_elements =
      static_cast<long long>(MatrixSize) * MatrixSize;
  const long long total = static_cast<long long>(batch) * matrix_elements;
  const int blocks = static_cast<int>(
      std::min<long long>(4096, (total + threads - 1) / threads));
  initialize_combined_batched_triangle<MatrixSize><<<blocks, threads>>>(
      input.data_ptr<float>(), output.data_ptr<float>(), pointer_data, batch,
      total);
  C10_CUDA_KERNEL_LAUNCH_CHECK();

  check_combined_solver(
      cusolverDnSpotrfBatched(
          state.batched_handle, CUBLAS_FILL_MODE_UPPER, MatrixSize,
          pointer_data, MatrixSize, state.batched_info.data_ptr<int>(), batch),
      "cusolverDnSpotrfBatched");

  clear_combined_batched_row_major_upper<MatrixSize><<<blocks, threads>>>(
      output.data_ptr<float>(), total);
  C10_CUDA_KERNEL_LAUNCH_CHECK();
  return output;
}

void check_combined_batched_input(const torch::Tensor& input, int n,
                                  const char* operation) {
  TORCH_CHECK(input.is_cuda(), operation, " expects a CUDA tensor");
  TORCH_CHECK(input.scalar_type() == torch::kFloat32,
              operation, " expects float32 input");
  TORCH_CHECK(input.dim() == 3 && input.size(0) > 0 &&
                  input.size(1) == n && input.size(2) == n,
              operation, " received an invalid shape");
  TORCH_CHECK(input.is_contiguous(), operation, " expects contiguous input");
  TORCH_CHECK(input.size(0) <= std::numeric_limits<int>::max(),
              operation, " batch is too large");
  static_assert(sizeof(void*) == sizeof(int64_t),
                "pointer storage requires 64-bit pointers");
}

struct TrueBatchedBlockedState {
  cusolverDnHandle_t solver = nullptr;
  cublasHandle_t blas = nullptr;
  torch::Tensor pointers;
  torch::Tensor info;
  int pointer_capacity = 0;
  int device = -1;
  std::mutex mutex;

  ~TrueBatchedBlockedState() {
    if (solver != nullptr) {
      cusolverDnDestroy(solver);
    }
    if (blas != nullptr) {
      cublasDestroy(blas);
    }
  }

  void prepare(const torch::Tensor& input, int batch) {
    const int input_device = input.get_device();
    if (device == -1) {
      check_combined_solver(cusolverDnCreate(&solver),
                            "cusolverDnCreate true-batched blocked");
      check_combined_blas(cublasCreate(&blas),
                          "cublasCreate true-batched blocked");
      check_combined_blas(
          cublasSetPointerMode(blas, CUBLAS_POINTER_MODE_HOST),
          "cublasSetPointerMode true-batched blocked");
      device = input_device;
    } else {
      TORCH_CHECK(device == input_device && solver != nullptr &&
                      blas != nullptr,
                  "true-batched blocked state cannot migrate between CUDA "
                  "devices");
    }

    if (pointer_capacity < batch) {
      pointers = torch::empty(
          {2, batch}, input.options().dtype(torch::kInt64));
      info = torch::empty({batch}, input.options().dtype(torch::kInt32));
      pointer_capacity = batch;
    }
  }
};

TrueBatchedBlockedState& true_batched_blocked_state() {
  static TrueBatchedBlockedState state;
  return state;
}

constexpr int kRegisterPanel64 = 64;
constexpr int kRegisterPanelStride64 = 65;
constexpr int kRegisterPanelThreads64 = 128;
constexpr int kRegisterPanelSharedBytes64 =
    kRegisterPanel64 * kRegisterPanelStride64 * sizeof(float);
static_assert(kRegisterPanelSharedBytes64 < 17 * 1024,
              "register panel shared storage exceeds its bound");

constexpr int kRegisterDiagonal128 = 128;
constexpr int kRegisterDiagonalRows128 = 128;
constexpr int kRegisterDiagonalStride128 = 65;
constexpr int kRegisterDiagonalThreads128 = 128;
constexpr int kRegisterDiagonalSharedBytes128 =
    kRegisterDiagonalRows128 * kRegisterDiagonalStride128 * sizeof(float);
static_assert(kRegisterDiagonalSharedBytes128 < 33 * 1024,
              "register diagonal shared storage exceeds its bound");

template <int Base>
__device__ __forceinline__ void factor_register_panel64_stage(
    float (*tile)[kRegisterPanelStride64], int tid, int lane, int warp) {
  if (warp == 0) {
    N64RegisterRow values;
    if (lane < kNestedBlockSize64) {
      n64_load_register_row(values, tile, Base + lane, Base, N64AllIndices{});
    } else {
      values = {};
    }
    n64_factor_register_block(values, lane, N64AllIndices{});
    if (lane < kNestedBlockSize64) {
      n64_store_register_row(values, tile, Base + lane, Base,
                             N64AllIndices{});
    }
  }
  __syncthreads();

  constexpr int next = Base + kNestedBlockSize64;
  for (int row = next + tid; row < kRegisterPanel64;
       row += kRegisterPanelThreads64) {
    N64RegisterRow values;
    n64_load_register_row(values, tile, row, Base, N64AllIndices{});
    n64_solve_outer_panel_first(values, tile, Base, N64AllIndices{});
    n64_store_register_row(values, tile, row, Base, N64AllIndices{});
  }
  __syncthreads();

  constexpr int remaining = kRegisterPanel64 - next;
  if constexpr (remaining > 0) {
    for (int linear = tid; linear < remaining * remaining;
         linear += kRegisterPanelThreads64) {
      const int row = next + linear / remaining;
      const int col = next + linear % remaining;
      if (row >= col) {
        float value = tile[row][col];
#pragma unroll
        for (int k = Base; k < next; ++k) {
          value = fmaf(-tile[row][k], tile[col][k], value);
        }
        tile[row][col] = value;
      }
    }
  }
  __syncthreads();
}

template <int Base>
__device__ __forceinline__ void factor_register_diagonal128_first_stage(
    float (*tile)[kRegisterDiagonalStride128], int tid, int lane, int warp) {
  if (warp == 0) {
    N64RegisterRow values;
    if (lane < kNestedBlockSize64) {
      n64_load_register_row(values, tile, Base + lane, Base, N64AllIndices{});
    } else {
      values = {};
    }
    n64_factor_register_block(values, lane, N64AllIndices{});
    if (lane < kNestedBlockSize64) {
      n64_store_register_row(values, tile, Base + lane, Base,
                             N64AllIndices{});
    }
  }
  __syncthreads();

  constexpr int next = Base + kNestedBlockSize64;
  for (int row = next + tid; row < kRegisterDiagonalRows128;
       row += kRegisterDiagonalThreads128) {
    N64RegisterRow values;
    n64_load_register_row(values, tile, row, Base, N64AllIndices{});
    n64_solve_outer_panel_first(values, tile, Base, N64AllIndices{});
    n64_store_register_row(values, tile, row, Base, N64AllIndices{});
  }
  __syncthreads();

  constexpr int remaining_columns = kRegisterPanel64 - next;
  constexpr int remaining_rows = kRegisterDiagonalRows128 - next;
  if constexpr (remaining_columns > 0) {
    for (int linear = tid;
         linear < remaining_rows * remaining_columns;
         linear += kRegisterDiagonalThreads128) {
      const int row = next + linear / remaining_columns;
      const int col = next + linear % remaining_columns;
      if (row >= col) {
        float value = tile[row][col];
#pragma unroll
        for (int k = Base; k < next; ++k) {
          value = fmaf(-tile[row][k], tile[col][k], value);
        }
        tile[row][col] = value;
      }
    }
  }
  __syncthreads();
}

template <int MatrixSize>
__global__ __launch_bounds__(kRegisterDiagonalThreads128)
void factor_register_diagonal128_kernel(
    float* matrices, float** diagonal_pointers, float** panel_pointers,
    int offset) {
  static_assert(MatrixSize == 512 || MatrixSize == 1024 ||
                    MatrixSize == 2048,
                "register diagonal kernel supports the high-batch paths");
  constexpr long long matrix_elements =
      static_cast<long long>(MatrixSize) * MatrixSize;
  const int matrix = static_cast<int>(blockIdx.x);
  const int tid = static_cast<int>(threadIdx.x);
  const int lane = tid & 31;
  const int warp = tid >> 5;
  float* const matrix_base =
      matrices + static_cast<long long>(matrix) * matrix_elements;

  __shared__ float tile[kRegisterDiagonalRows128]
                       [kRegisterDiagonalStride128];
  for (int linear = tid;
       linear < kRegisterDiagonalRows128 * kRegisterPanel64;
       linear += kRegisterDiagonalThreads128) {
    const int row = linear >> 6;
    const int col = linear & 63;
    tile[row][col] =
        row >= col
            ? matrix_base[(offset + row) * MatrixSize + offset + col]
            : 0.0f;
  }
  __syncthreads();

  factor_register_diagonal128_first_stage<0>(tile, tid, lane, warp);
  factor_register_diagonal128_first_stage<16>(tile, tid, lane, warp);
  factor_register_diagonal128_first_stage<32>(tile, tid, lane, warp);
  factor_register_diagonal128_first_stage<48>(tile, tid, lane, warp);

  for (int linear = tid;
       linear < kRegisterDiagonalRows128 * kRegisterPanel64;
       linear += kRegisterDiagonalThreads128) {
    const int row = linear >> 6;
    const int col = linear & 63;
    if (row >= col) {
      matrix_base[(offset + row) * MatrixSize + offset + col] =
          tile[row][col];
    }
  }

  for (int linear = tid;
       linear < kRegisterPanel64 * kRegisterPanel64;
       linear += kRegisterDiagonalThreads128) {
    const int row = linear >> 6;
    const int col = linear & 63;
    float value = 0.0f;
    if (row >= col) {
      value = matrix_base[(offset + kRegisterPanel64 + row) * MatrixSize +
                          offset + kRegisterPanel64 + col];
#pragma unroll
      for (int k = 0; k < kRegisterPanel64; ++k) {
        value = fmaf(-tile[kRegisterPanel64 + row][k],
                     tile[kRegisterPanel64 + col][k], value);
      }
    }
    tile[row][col] = value;
  }
  __syncthreads();

  factor_register_panel64_stage<0>(tile, tid, lane, warp);
  factor_register_panel64_stage<16>(tile, tid, lane, warp);
  factor_register_panel64_stage<32>(tile, tid, lane, warp);
  factor_register_panel64_stage<48>(tile, tid, lane, warp);

  for (int linear = tid;
       linear < kRegisterPanel64 * kRegisterPanel64;
       linear += kRegisterDiagonalThreads128) {
    const int row = linear >> 6;
    const int col = linear & 63;
    if (row >= col) {
      matrix_base[(offset + kRegisterPanel64 + row) * MatrixSize +
                  offset + kRegisterPanel64 + col] = tile[row][col];
    }
  }

  if (tid == 0) {
    float* const diagonal = matrix_base + offset * (MatrixSize + 1);
    diagonal_pointers[matrix] = diagonal;
    panel_pointers[matrix] =
        offset + kRegisterDiagonal128 < MatrixSize
            ? matrix_base + (offset + kRegisterDiagonal128) * MatrixSize +
                  offset
            : diagonal;
  }
}

template <int MatrixSize>
__global__ __launch_bounds__(kRegisterPanelThreads64)
void factor_register_panel64_kernel(
    float* matrices, float** diagonal_pointers, float** panel_pointers,
    int offset) {
  constexpr long long matrix_elements =
      static_cast<long long>(MatrixSize) * MatrixSize;
  const int matrix = static_cast<int>(blockIdx.x);
  const int tid = static_cast<int>(threadIdx.x);
  const int lane = tid & 31;
  const int warp = tid >> 5;
  float* const matrix_base =
      matrices + static_cast<long long>(matrix) * matrix_elements;

  __shared__ float tile[kRegisterPanel64][kRegisterPanelStride64];
  for (int linear = tid;
       linear < kRegisterPanel64 * kRegisterPanel64;
       linear += kRegisterPanelThreads64) {
    const int row = linear >> 6;
    const int col = linear & 63;
    tile[row][col] = row >= col
                         ? matrix_base[(offset + row) * MatrixSize +
                                       offset + col]
                         : 0.0f;
  }
  __syncthreads();

  factor_register_panel64_stage<0>(tile, tid, lane, warp);
  factor_register_panel64_stage<16>(tile, tid, lane, warp);
  factor_register_panel64_stage<32>(tile, tid, lane, warp);
  factor_register_panel64_stage<48>(tile, tid, lane, warp);

  for (int linear = tid;
       linear < kRegisterPanel64 * kRegisterPanel64;
       linear += kRegisterPanelThreads64) {
    const int row = linear >> 6;
    const int col = linear & 63;
    if (row >= col) {
      matrix_base[(offset + row) * MatrixSize + offset + col] =
          tile[row][col];
    }
  }

  if (tid == 0) {
    float* const diagonal =
        matrix_base + offset * (MatrixSize + 1);
    diagonal_pointers[matrix] = diagonal;
    panel_pointers[matrix] =
        offset + kRegisterPanel64 < MatrixSize
            ? matrix_base + (offset + kRegisterPanel64) * MatrixSize + offset
            : diagonal;
  }
}

torch::Tensor cholesky256_batch64_register_panel64_impl(torch::Tensor input) {
  constexpr int matrix_size = 256;
  constexpr int batch = 64;
  constexpr long long matrix_elements =
      static_cast<long long>(matrix_size) * matrix_size;
  c10::cuda::CUDAGuard device_guard(input.device());
  auto output = input.clone();
  auto& state = true_batched_blocked_state();
  std::lock_guard<std::mutex> lock(state.mutex);
  state.prepare(input, batch);

  auto diagonal_pointers = reinterpret_cast<float**>(
      state.pointers.data_ptr<int64_t>());
  auto panel_pointers = diagonal_pointers + batch;
  const float one = 1.0f;
  const float minus_one = -1.0f;
  float* const base = output.data_ptr<float>();

  for (int offset = 0; offset < matrix_size; offset += kRegisterPanel64) {
    const int trailing = matrix_size - offset - kRegisterPanel64;
    factor_register_panel64_kernel<matrix_size>
        <<<batch, kRegisterPanelThreads64>>>(
        base, diagonal_pointers, panel_pointers, offset);
    C10_CUDA_KERNEL_LAUNCH_CHECK();

    if (trailing == 0) {
      continue;
    }

    check_combined_blas(
        cublasStrsmBatched(
            state.blas, CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_UPPER,
            CUBLAS_OP_T, CUBLAS_DIAG_NON_UNIT, kRegisterPanel64, trailing,
            &one,
            reinterpret_cast<const float* const*>(diagonal_pointers),
            matrix_size, panel_pointers, matrix_size, batch),
        "cublasStrsmBatched register panel");

    float* const panel =
        base + offset + (offset + kRegisterPanel64) * matrix_size;
    float* const trailing_matrix =
        base + (offset + kRegisterPanel64) * (matrix_size + 1);
    check_combined_blas(
        cublasSgemmStridedBatched(
            state.blas, CUBLAS_OP_T, CUBLAS_OP_N, trailing, trailing,
            kRegisterPanel64, &minus_one, panel, matrix_size, matrix_elements,
            panel, matrix_size, matrix_elements, &one, trailing_matrix,
            matrix_size, matrix_elements, batch),
        "cublasSgemmStridedBatched register trailing update");
  }

  constexpr int threads = 256;
  constexpr long long total = batch * matrix_elements;
  constexpr int blocks = 4096;
  clear_combined_batched_row_major_upper<matrix_size>
      <<<blocks, threads>>>(base, total);
  C10_CUDA_KERNEL_LAUNCH_CHECK();
  return output;
}

torch::Tensor cholesky512_batch16_register_panel64_impl(torch::Tensor input) {
  constexpr int matrix_size = 512;
  constexpr int batch = 16;
  constexpr long long matrix_elements =
      static_cast<long long>(matrix_size) * matrix_size;
  c10::cuda::CUDAGuard device_guard(input.device());
  auto output = input.clone();
  auto& state = true_batched_blocked_state();
  std::lock_guard<std::mutex> lock(state.mutex);
  state.prepare(input, batch);

  auto diagonal_pointers = reinterpret_cast<float**>(
      state.pointers.data_ptr<int64_t>());
  auto panel_pointers = diagonal_pointers + state.pointer_capacity;
  const float one = 1.0f;
  const float minus_one = -1.0f;
  float* const base = output.data_ptr<float>();

  for (int offset = 0; offset < matrix_size; offset += kRegisterPanel64) {
    const int trailing = matrix_size - offset - kRegisterPanel64;
    factor_register_panel64_kernel<matrix_size>
        <<<batch, kRegisterPanelThreads64>>>(
            base, diagonal_pointers, panel_pointers, offset);
    C10_CUDA_KERNEL_LAUNCH_CHECK();

    if (trailing == 0) {
      continue;
    }

    check_combined_blas(
        cublasStrsmBatched(
            state.blas, CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_UPPER,
            CUBLAS_OP_T, CUBLAS_DIAG_NON_UNIT, kRegisterPanel64, trailing,
            &one,
            reinterpret_cast<const float* const*>(diagonal_pointers),
            matrix_size, panel_pointers, matrix_size, batch),
        "cublasStrsmBatched n512 batch16 register panel");

    float* const panel =
        base + offset + (offset + kRegisterPanel64) * matrix_size;
    float* const trailing_matrix =
        base + (offset + kRegisterPanel64) * (matrix_size + 1);
    check_combined_blas(
        cublasSgemmStridedBatched(
            state.blas, CUBLAS_OP_T, CUBLAS_OP_N, trailing, trailing,
            kRegisterPanel64, &minus_one, panel, matrix_size, matrix_elements,
            panel, matrix_size, matrix_elements, &one, trailing_matrix,
            matrix_size, matrix_elements, batch),
        "cublasSgemmStridedBatched n512 batch16 register trailing update");
  }

  constexpr int threads = 256;
  constexpr long long total = batch * matrix_elements;
  constexpr int blocks = 4096;
  clear_combined_batched_row_major_upper<matrix_size>
      <<<blocks, threads>>>(base, total);
  C10_CUDA_KERNEL_LAUNCH_CHECK();
  return output;
}

template <int MatrixSize, int Batch, int PanelSize, int ColumnTile = 0>
torch::Tensor cholesky_true_batched_blocked_impl(
    torch::Tensor input, const char* operation) {
  static_assert(MatrixSize % PanelSize == 0,
                "panel size must divide matrix size");
  static_assert(ColumnTile == 0 ||
                    (ColumnTile > 0 && ColumnTile <= MatrixSize),
                "column tile must be disabled or fit the matrix");
  constexpr long long matrix_elements =
      static_cast<long long>(MatrixSize) * MatrixSize;
  check_combined_batched_input(input, MatrixSize, operation);
  TORCH_CHECK(input.size(0) == Batch, operation, " received an invalid batch");

  c10::cuda::CUDAGuard device_guard(input.device());
  auto output = torch::empty_like(input);
  auto& state = true_batched_blocked_state();
  std::lock_guard<std::mutex> lock(state.mutex);
  state.prepare(input, Batch);

  constexpr int threads = 256;
  constexpr long long total = static_cast<long long>(Batch) * matrix_elements;
  const int matrix_blocks = static_cast<int>(
      std::min<long long>(4096, (total + threads - 1) / threads));
  auto diagonal_pointers = reinterpret_cast<float**>(
      state.pointers.data_ptr<int64_t>());
  auto panel_pointers = diagonal_pointers + Batch;
  if constexpr ((MatrixSize == 512 && Batch == 640) ||
                (MatrixSize == 1024 && Batch == 60)) {
    copy_highbatch_lower_only<MatrixSize>
        <<<matrix_blocks, threads>>>(
            input.data_ptr<float>(), output.data_ptr<float>(), total);
  } else {
    copy_highbatch_lower_zero_upper<MatrixSize>
        <<<matrix_blocks, threads>>>(
            input.data_ptr<float>(), output.data_ptr<float>(), total);
  }
  C10_CUDA_KERNEL_LAUNCH_CHECK();

  constexpr int pointer_blocks = (Batch + threads - 1) / threads;
  const float one = 1.0f;
  const float minus_one = -1.0f;
  float* const base = output.data_ptr<float>();

  for (int offset = 0; offset < MatrixSize; offset += PanelSize) {
    const int trailing = MatrixSize - offset - PanelSize;
    const int diagonal_offset = offset + offset * MatrixSize;
    const int panel_offset =
        trailing > 0 ? offset + (offset + PanelSize) * MatrixSize
                     : diagonal_offset;
    fill_highbatch_panel_pointers<MatrixSize>
        <<<pointer_blocks, threads>>>(
            base, diagonal_pointers, panel_pointers, diagonal_offset,
            panel_offset, Batch);
    C10_CUDA_KERNEL_LAUNCH_CHECK();

    check_combined_solver(
        cusolverDnSpotrfBatched(
            state.solver, CUBLAS_FILL_MODE_UPPER, PanelSize,
            diagonal_pointers, MatrixSize, state.info.data_ptr<int>(), Batch),
        "cusolverDnSpotrfBatched true-batched blocked diagonal");

    if (trailing == 0) {
      continue;
    }

    check_combined_blas(
        cublasStrsmBatched(
            state.blas, CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_UPPER,
            CUBLAS_OP_T, CUBLAS_DIAG_NON_UNIT, PanelSize, trailing, &one,
            reinterpret_cast<const float* const*>(diagonal_pointers),
            MatrixSize, panel_pointers, MatrixSize, Batch),
        "cublasStrsmBatched true-batched blocked panel");

    float* const panel = base + panel_offset;
    float* const trailing_matrix =
        base + (offset + PanelSize) * (MatrixSize + 1);
    if constexpr (ColumnTile > 0) {
      for (int column = 0; column < trailing; column += ColumnTile) {
        const int width = std::min(ColumnTile, trailing - column);
        const int rows = trailing - column;
        float* const panel_column =
            panel + static_cast<long long>(column) * MatrixSize;
        float* const trailing_column =
            trailing_matrix + static_cast<long long>(column) *
                                  (MatrixSize + 1);
        check_combined_blas(
            cublasGemmStridedBatchedEx(
                state.blas, CUBLAS_OP_T, CUBLAS_OP_N, width, rows,
                PanelSize, &minus_one, panel_column, CUDA_R_32F, MatrixSize,
                matrix_elements, panel_column, CUDA_R_32F, MatrixSize,
                matrix_elements, &one, trailing_column, CUDA_R_32F,
                MatrixSize, matrix_elements, Batch,
                CUBLAS_COMPUTE_32F_FAST_TF32, CUBLAS_GEMM_DEFAULT),
            "cublasGemmStridedBatchedEx true-batched blocked block-column");
      }
    } else {
      check_combined_blas(
          cublasGemmStridedBatchedEx(
              state.blas, CUBLAS_OP_T, CUBLAS_OP_N, trailing, trailing,
              PanelSize, &minus_one, panel, CUDA_R_32F, MatrixSize,
              matrix_elements, panel, CUDA_R_32F, MatrixSize, matrix_elements,
              &one, trailing_matrix, CUDA_R_32F, MatrixSize, matrix_elements,
              Batch, CUBLAS_COMPUTE_32F_FAST_TF32, CUBLAS_GEMM_DEFAULT),
          "cublasGemmStridedBatchedEx true-batched blocked trailing");
    }
  }

  clear_combined_batched_row_major_upper<MatrixSize>
      <<<matrix_blocks, threads>>>(base, total);
  C10_CUDA_KERNEL_LAUNCH_CHECK();
  return output;
}

template <int MatrixSize, int Batch, int InnerPanel, int OuterPanel,
          bool FastBf16, bool CloneInput = false>
torch::Tensor cholesky_true_batched_two_level_impl(
    torch::Tensor input, const char* operation) {
  static_assert(InnerPanel > 0 && OuterPanel >= InnerPanel,
                "invalid two-level panel sizes");
  constexpr long long matrix_elements =
      static_cast<long long>(MatrixSize) * MatrixSize;
  check_combined_batched_input(input, MatrixSize, operation);
  TORCH_CHECK(input.size(0) == Batch, operation, " received an invalid batch");

  c10::cuda::CUDAGuard device_guard(input.device());
  torch::Tensor output;
  if constexpr (CloneInput) {
    output = input.clone();
  } else {
    output = torch::empty_like(input);
  }
  auto& state = true_batched_blocked_state();
  std::lock_guard<std::mutex> lock(state.mutex);
  state.prepare(input, Batch);

  constexpr int threads = 256;
  constexpr long long total = static_cast<long long>(Batch) * matrix_elements;
  const int matrix_blocks = static_cast<int>(
      std::min<long long>(4096, (total + threads - 1) / threads));
  if constexpr (!CloneInput) {
    copy_highbatch_lower_zero_upper<MatrixSize>
        <<<matrix_blocks, threads>>>(
            input.data_ptr<float>(), output.data_ptr<float>(), total);
    C10_CUDA_KERNEL_LAUNCH_CHECK();
  }

  auto diagonal_pointers = reinterpret_cast<float**>(
      state.pointers.data_ptr<int64_t>());
  auto panel_pointers = diagonal_pointers + Batch;
  constexpr int pointer_blocks = (Batch + threads - 1) / threads;
  const float one = 1.0f;
  const float minus_one = -1.0f;
  float* const base = output.data_ptr<float>();
  constexpr cublasComputeType_t update_compute =
      FastBf16 ? CUBLAS_COMPUTE_32F_FAST_16BF
                : CUBLAS_COMPUTE_32F_FAST_TF32;

  for (int outer = 0; outer < MatrixSize; outer += OuterPanel) {
    const int outer_end = std::min(MatrixSize, outer + OuterPanel);
    const int outer_width = outer_end - outer;

    for (int offset = outer; offset < outer_end; offset += InnerPanel) {
      const int jb = std::min(InnerPanel, outer_end - offset);
      const int trailing = MatrixSize - offset - jb;
      const int inner_remaining = outer_end - offset - jb;
      const int diagonal_offset = offset + offset * MatrixSize;
      const int panel_offset =
          trailing > 0 ? offset + (offset + jb) * MatrixSize
                       : diagonal_offset;
      if constexpr ((((MatrixSize == 2048 || MatrixSize == 4096) &&
                       Batch == 2) ||
                      (MatrixSize == 2048 && Batch == 8)) &&
                     InnerPanel == kRegisterPanel64) {
        factor_register_panel64_kernel<MatrixSize>
            <<<Batch, kRegisterPanelThreads64>>>(
                base, diagonal_pointers, panel_pointers, offset);
        C10_CUDA_KERNEL_LAUNCH_CHECK();
      } else if constexpr (MatrixSize == 2048 && Batch == 8 &&
                           InnerPanel == kRegisterDiagonal128) {
        factor_register_diagonal128_kernel<MatrixSize>
            <<<Batch, kRegisterDiagonalThreads128>>>(
                base, diagonal_pointers, panel_pointers, offset);
        C10_CUDA_KERNEL_LAUNCH_CHECK();
      } else {
        fill_highbatch_panel_pointers<MatrixSize>
            <<<pointer_blocks, threads>>>(
                base, diagonal_pointers, panel_pointers, diagonal_offset,
                panel_offset, Batch);
        C10_CUDA_KERNEL_LAUNCH_CHECK();

        check_combined_solver(
            cusolverDnSpotrfBatched(
                state.solver, CUBLAS_FILL_MODE_UPPER, jb,
                diagonal_pointers, MatrixSize, state.info.data_ptr<int>(),
                Batch),
            "cusolverDnSpotrfBatched two-level diagonal");
      }
      if (trailing == 0) {
        continue;
      }

      check_combined_blas(
          cublasStrsmBatched(
              state.blas, CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_UPPER,
              CUBLAS_OP_T, CUBLAS_DIAG_NON_UNIT, jb, trailing, &one,
              reinterpret_cast<const float* const*>(diagonal_pointers),
              MatrixSize, panel_pointers, MatrixSize, Batch),
          "cublasStrsmBatched two-level panel");

      if (inner_remaining > 0) {
        float* const panel = base + panel_offset;
        float* const next_strip =
            base + (offset + jb) * (MatrixSize + 1);
        check_combined_blas(
            cublasGemmStridedBatchedEx(
                state.blas, CUBLAS_OP_T, CUBLAS_OP_N, inner_remaining,
                trailing, jb, &minus_one, panel, CUDA_R_32F, MatrixSize,
                matrix_elements, panel, CUDA_R_32F, MatrixSize,
                matrix_elements, &one, next_strip, CUDA_R_32F, MatrixSize,
                matrix_elements, Batch, update_compute,
                CUBLAS_GEMM_DEFAULT),
            "cublasGemmStridedBatchedEx two-level panel strip");
      }
    }

    const int far = MatrixSize - outer_end;
    if (far > 0) {
      float* const outer_to_far =
          base + outer + static_cast<long long>(outer_end) * MatrixSize;
      float* const far_diagonal =
          base + static_cast<long long>(outer_end) * (MatrixSize + 1);
      check_combined_blas(
          cublasGemmStridedBatchedEx(
              state.blas, CUBLAS_OP_T, CUBLAS_OP_N, far, far, outer_width,
              &minus_one, outer_to_far, CUDA_R_32F, MatrixSize,
              matrix_elements, outer_to_far, CUDA_R_32F, MatrixSize,
              matrix_elements, &one, far_diagonal, CUDA_R_32F, MatrixSize,
              matrix_elements, Batch, update_compute, CUBLAS_GEMM_DEFAULT),
          "cublasGemmStridedBatchedEx two-level wide trailing");
    }
  }

  clear_combined_batched_row_major_upper<MatrixSize>
      <<<matrix_blocks, threads>>>(base, total);
  C10_CUDA_KERNEL_LAUNCH_CHECK();
  return output;
}

template <int MatrixSize, int Batch, int InnerPanel, int OuterPanel>
torch::Tensor cholesky_highbatch_two_level_impl(
    torch::Tensor input, const char* operation) {
  static_assert(InnerPanel > 0 && OuterPanel >= InnerPanel,
                "invalid two-level panel sizes");
  static_assert(OuterPanel % InnerPanel == 0,
                "outer panel must contain whole inner panels");
  static_assert(MatrixSize % OuterPanel == 0,
                "outer panel must divide matrix size");
  constexpr long long matrix_elements =
      static_cast<long long>(MatrixSize) * MatrixSize;
  check_combined_batched_input(input, MatrixSize, operation);
  TORCH_CHECK(input.size(0) == Batch, operation, " received an invalid batch");

  c10::cuda::CUDAGuard device_guard(input.device());
  auto output = torch::empty_like(input);
  auto& state = true_batched_blocked_state();
  std::lock_guard<std::mutex> lock(state.mutex);
  state.prepare(input, Batch);

  constexpr int threads = 256;
  constexpr long long total = static_cast<long long>(Batch) * matrix_elements;
  const int matrix_blocks = static_cast<int>(
      std::min<long long>(4096, (total + threads - 1) / threads));
  constexpr long long vector_total = total / 4;
  const int vector_blocks = static_cast<int>(
      std::min<long long>(4096, (vector_total + threads - 1) / threads));
  if constexpr (MatrixSize == 1024 && (Batch == 4 || Batch == 60)) {
    copy_highbatch_lower_only_float4<MatrixSize>
        <<<vector_blocks, threads>>>(input.data_ptr<float>(),
                                    output.data_ptr<float>(), vector_total);
  } else {
    copy_highbatch_lower_only<MatrixSize>
        <<<matrix_blocks, threads>>>(
            input.data_ptr<float>(), output.data_ptr<float>(), total);
  }
  C10_CUDA_KERNEL_LAUNCH_CHECK();

  auto diagonal_pointers = reinterpret_cast<float**>(
      state.pointers.data_ptr<int64_t>());
  auto panel_pointers = diagonal_pointers + Batch;
  constexpr int pointer_blocks = (Batch + threads - 1) / threads;
  const float one = 1.0f;
  const float minus_one = -1.0f;
  float* const base = output.data_ptr<float>();

  for (int outer = 0; outer < MatrixSize; outer += OuterPanel) {
    constexpr int outer_width = OuterPanel;
    const int outer_end = outer + OuterPanel;

    for (int offset = outer; offset < outer_end; offset += InnerPanel) {
      constexpr int jb = InnerPanel;
      const int trailing = MatrixSize - offset - jb;
      const int inner_remaining = outer_end - offset - jb;
      const int diagonal_offset = offset + offset * MatrixSize;
      const int panel_offset =
          trailing > 0 ? offset + (offset + jb) * MatrixSize
                       : diagonal_offset;
      if constexpr (MatrixSize == 1024 &&
                    (Batch == 4 || Batch == 60) &&
                    InnerPanel == kRegisterPanel64) {
        factor_register_panel64_kernel<MatrixSize>
            <<<Batch, kRegisterPanelThreads64>>>(
                base, diagonal_pointers, panel_pointers, offset);
        C10_CUDA_KERNEL_LAUNCH_CHECK();
      } else if constexpr (MatrixSize == 512 && Batch == 640 &&
                           InnerPanel == kRegisterDiagonal128) {
        factor_register_diagonal128_kernel<MatrixSize>
            <<<Batch, kRegisterDiagonalThreads128>>>(
                base, diagonal_pointers, panel_pointers, offset);
        C10_CUDA_KERNEL_LAUNCH_CHECK();
      } else {
        fill_highbatch_panel_pointers<MatrixSize>
            <<<pointer_blocks, threads>>>(
                base, diagonal_pointers, panel_pointers, diagonal_offset,
                panel_offset, Batch);
        C10_CUDA_KERNEL_LAUNCH_CHECK();

        check_combined_solver(
            cusolverDnSpotrfBatched(
                state.solver, CUBLAS_FILL_MODE_UPPER, jb,
                diagonal_pointers, MatrixSize, state.info.data_ptr<int>(),
                Batch),
            "cusolverDnSpotrfBatched high-batch two-level diagonal");
      }
      if (trailing == 0) {
        continue;
      }

      check_combined_blas(
          cublasStrsmBatched(
              state.blas, CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_UPPER,
              CUBLAS_OP_T, CUBLAS_DIAG_NON_UNIT, jb, trailing, &one,
              reinterpret_cast<const float* const*>(diagonal_pointers),
              MatrixSize, panel_pointers, MatrixSize, Batch),
          "cublasStrsmBatched high-batch two-level panel");

      if (inner_remaining > 0) {
        float* const panel = base + panel_offset;
        float* const next_strip =
            base + (offset + jb) * (MatrixSize + 1);
        check_combined_blas(
            cublasGemmStridedBatchedEx(
                state.blas, CUBLAS_OP_T, CUBLAS_OP_N, inner_remaining,
                trailing, jb, &minus_one, panel, CUDA_R_32F, MatrixSize,
                matrix_elements, panel, CUDA_R_32F, MatrixSize,
                matrix_elements, &one, next_strip, CUDA_R_32F, MatrixSize,
                matrix_elements, Batch, CUBLAS_COMPUTE_32F_FAST_TF32,
                CUBLAS_GEMM_DEFAULT),
            "cublasGemmStridedBatchedEx high-batch two-level panel strip");
      }
    }

    const int far = MatrixSize - outer_end;
    if (far > 0) {
      float* const outer_to_far =
          base + outer + static_cast<long long>(outer_end) * MatrixSize;
      float* const far_diagonal =
          base + static_cast<long long>(outer_end) * (MatrixSize + 1);
      check_combined_blas(
          cublasGemmStridedBatchedEx(
              state.blas, CUBLAS_OP_T, CUBLAS_OP_N, far, far, outer_width,
              &minus_one, outer_to_far, CUDA_R_32F, MatrixSize,
              matrix_elements, outer_to_far, CUDA_R_32F, MatrixSize,
              matrix_elements, &one, far_diagonal, CUDA_R_32F, MatrixSize,
              matrix_elements, Batch, CUBLAS_COMPUTE_32F_FAST_TF32,
              CUBLAS_GEMM_DEFAULT),
          "cublasGemmStridedBatchedEx high-batch two-level wide trailing");
    }
  }

  if constexpr (MatrixSize == 1024 && (Batch == 4 || Batch == 60)) {
    clear_highbatch_upper_float4<MatrixSize>
        <<<vector_blocks, threads>>>(base, vector_total);
  } else {
    clear_combined_batched_row_major_upper<MatrixSize>
        <<<matrix_blocks, threads>>>(base, total);
  }
  C10_CUDA_KERNEL_LAUNCH_CHECK();
  return output;
}

}  // namespace

torch::Tensor cholesky256_batch64_register_panel64_cuda(torch::Tensor input) {
  check_combined_batched_input(
      input, kCombinedSize256,
      "cholesky256_batch64_register_panel64_cuda");
  TORCH_CHECK(input.size(0) == 64,
              "cholesky256_batch64_register_panel64_cuda expects batch 64");
  return cholesky256_batch64_register_panel64_impl(input);
}

torch::Tensor cholesky256_combined_batched_cuda(torch::Tensor input) {
  check_combined_batched_input(
      input, kCombinedSize256, "cholesky256_combined_batched_cuda");
  return cholesky_combined_batched_impl<kCombinedSize256>(input);
}

torch::Tensor cholesky512_batch16_register_panel64_cuda(torch::Tensor input) {
  check_combined_batched_input(
      input, kCombinedSize512,
      "cholesky512_batch16_register_panel64_cuda");
  TORCH_CHECK(input.size(0) == 16,
              "cholesky512_batch16_register_panel64_cuda expects batch 16");
  return cholesky512_batch16_register_panel64_impl(input);
}

torch::Tensor cholesky512_combined_batched_cuda(torch::Tensor input) {
  check_combined_batched_input(
      input, kCombinedSize512, "cholesky512_combined_batched_cuda");
  return cholesky_combined_batched_impl<kCombinedSize512>(input);
}

torch::Tensor cholesky512_highbatch_blocked_cuda(torch::Tensor input) {
  return cholesky_highbatch_two_level_impl<512, 640, 128, 256>(
      input, "cholesky512_highbatch_blocked_cuda");
}

torch::Tensor cholesky1024_batch4_blocked_cuda(torch::Tensor input) {
  return cholesky_highbatch_two_level_impl<1024, 4, 64, 128>(
      input, "cholesky1024_batch4_blocked_cuda");
}

torch::Tensor cholesky1024_batch60_blocked_cuda(torch::Tensor input) {
  return cholesky_highbatch_two_level_impl<1024, 60, 64, 256>(
      input, "cholesky1024_batch60_blocked_cuda");
}

torch::Tensor cholesky2048_batch2_blocked_cuda(torch::Tensor input) {
  return cholesky_true_batched_two_level_impl<2048, 2, 64, 256, true>(
      input, "cholesky2048_batch2_blocked_cuda");
}

torch::Tensor cholesky2048_batch8_blocked_cuda(torch::Tensor input) {
  return cholesky_true_batched_two_level_impl<2048, 8, 64, 512, false>(
      input, "cholesky2048_batch8_blocked_cuda");
}

torch::Tensor cholesky4096_batch2_truebatched_fast16bf_cuda(
    torch::Tensor input) {
  return cholesky_true_batched_two_level_impl<4096, 2, 64, 512, true,
                                               true>(
      input, "cholesky4096_batch2_truebatched_fast16bf_cuda");
}

torch::Tensor cholesky_large_explicit_bf16_cuda(torch::Tensor input) {
  TORCH_CHECK(input.is_cuda(),
              "cholesky_large_explicit_bf16_cuda expects a CUDA tensor");
  TORCH_CHECK(input.scalar_type() == torch::kFloat32,
              "cholesky_large_explicit_bf16_cuda expects float32 input");
  TORCH_CHECK(
      input.dim() == 3 && input.size(1) == input.size(2) &&
          ((input.size(0) == 2 && input.size(1) == 4096) ||
           (input.size(0) == 1 &&
            (input.size(1) == 16384 || input.size(1) == 32768))),
      "cholesky_large_explicit_bf16_cuda expects batch/n in "
      "{(2,4096),(1,16384),(1,32768)}");
  TORCH_CHECK(input.is_contiguous(),
              "cholesky_large_explicit_bf16_cuda expects contiguous input");

  c10::cuda::CUDAGuard device_guard(input.device());
  const int device = input.get_device();
  const int batch = static_cast<int>(input.size(0));
  const int n = static_cast<int>(input.size(1));
  const long long matrix_elements = static_cast<long long>(n) * n;
  constexpr int standard_nb = 512;
  constexpr int champion_inner_nb = 128;
  constexpr int wide_inner_nb = 384;
  constexpr int panel_correction_nb = 96;
  constexpr int wide_prefix_end = 23040;
  constexpr int outer_nb = 1536;
  constexpr int trailing_tile = 3072;
  const bool lower_only_initialization = n >= 16384;
  const bool regime_schedule = n == 32768;
  const int factor_nb = regime_schedule
                            ? wide_inner_nb
                            : (lower_only_initialization ? champion_inner_nb
                                                         : standard_nb);
  auto output = lower_only_initialization ? torch::empty_like(input)
                                          : input.clone();
  float* const base = output.data_ptr<float>();
  if (lower_only_initialization) {
    constexpr int initialization_threads = 256;
    initialize_explicit_large_first_panel_lower
        <<<n, initialization_threads>>>(
            input.data_ptr<float>(), base, n, outer_nb);
    C10_CUDA_KERNEL_LAUNCH_CHECK();
  }
  auto& state = explicit_low_precision_state();
  std::lock_guard<std::mutex> lock(state.mutex);
  state.prepare(input);

  int required_elements = 0;
  check_combined_solver(
      cusolverDnSpotrf_bufferSize(
          state.solver_handle, CUBLAS_FILL_MODE_UPPER, factor_nb, base, n,
          &required_elements),
      "cusolverDnSpotrf_bufferSize explicit BF16");
  if (state.workspace_elements < required_elements) {
    if (state.workspace.defined()) {
      state.retired_buffers.push_back(state.workspace);
    }
    state.workspace = torch::empty(
        {required_elements}, input.options().dtype(torch::kFloat32));
    state.workspace_elements = required_elements;
  }
  if (!state.info.defined()) {
    state.info = torch::empty({1}, input.options().dtype(torch::kInt32));
  }
  const int packed_rows = lower_only_initialization ? outer_nb : standard_nb;
  const long long required_panel_capacity =
      static_cast<long long>(packed_rows) * n;
  if (state.panel_capacity < required_panel_capacity) {
    if (state.panel_bf16.defined()) {
      state.retired_buffers.push_back(state.panel_bf16);
    }
    state.panel_bf16 = torch::empty(
        {required_panel_capacity}, input.options().dtype(torch::kBFloat16));
    state.panel_capacity = required_panel_capacity;
  }
  auto* const packed = reinterpret_cast<__nv_bfloat16*>(
      state.panel_bf16.data_ptr<at::BFloat16>());
  constexpr long long mxfp8_scratch_bytes = 32LL * 1024 * 1024;
  if (regime_schedule) {
    const long long scale_elements =
        static_cast<long long>((outer_nb + 127) / 128) *
        ((n + 127) / 128) * 512;
    if (!state.mxfp8_primary_scale.defined() ||
        state.mxfp8_primary_scale.numel() < scale_elements) {
      if (state.mxfp8_primary_scale.defined()) {
        state.retired_buffers.push_back(state.mxfp8_primary_scale);
      }
      state.mxfp8_primary_scale = torch::empty(
          {scale_elements}, input.options().dtype(torch::kUInt8));
    }
    if (!state.mxfp8_diagonal_correction.defined() ||
        state.mxfp8_diagonal_correction.numel() < n) {
      if (state.mxfp8_diagonal_correction.defined()) {
        state.retired_buffers.push_back(state.mxfp8_diagonal_correction);
      }
      state.mxfp8_diagonal_correction = torch::empty(
          {n}, input.options().dtype(torch::kFloat32));
    }
    if (!state.mxfp8_scratch.defined() ||
        state.mxfp8_scratch.numel() < mxfp8_scratch_bytes) {
      if (state.mxfp8_scratch.defined()) {
        state.retired_buffers.push_back(state.mxfp8_scratch);
      }
      state.mxfp8_scratch = torch::empty(
          {mxfp8_scratch_bytes}, input.options().dtype(torch::kUInt8));
    }
  }

  const float one = 1.0f;
  const float minus_one = -1.0f;
  for (int matrix = 0; matrix < batch; ++matrix) {
    const float* const input_matrix_base =
        input.data_ptr<float>() +
        static_cast<long long>(matrix) * matrix_elements;
    float* const matrix_base =
        base + static_cast<long long>(matrix) * matrix_elements;
    if (lower_only_initialization) {
      constexpr int threads = 256;
      for (int outer = 0; outer < n; outer += outer_nb) {
        const int outer_end = std::min(n, outer + outer_nb);
        const int outer_width = outer_end - outer;
        const bool wide_prefix =
            regime_schedule && outer < wide_prefix_end;
        const int inner_nb =
            wide_prefix ? wide_inner_nb : champion_inner_nb;

        for (int k = outer; k < outer_end; k += inner_nb) {
          const int jb = std::min(inner_nb, outer_end - k);
          const int r = n - k - jb;
          const int inner_remaining = outer_end - k - jb;
          float* const diagonal =
              matrix_base + static_cast<long long>(k) * (n + 1);
          check_combined_solver(
              cusolverDnSpotrf(
                  state.solver_handle, CUBLAS_FILL_MODE_UPPER, jb, diagonal,
                  n, state.workspace.data_ptr<float>(),
                  state.workspace_elements, state.info.data_ptr<int>()),
              "cusolverDnSpotrf two-level explicit BF16");
          if (r == 0) {
            continue;
          }

          const int packed_row = k - outer;
          const int packed_col = k + jb - outer;
          if (wide_prefix) {
            for (int q = 0; q < jb; q += panel_correction_nb) {
              const int qb = std::min(panel_correction_nb, jb - q);
              float* const correction_diagonal =
                  matrix_base + (k + q) +
                  static_cast<long long>(k + q) * n;
              float* const correction_panel =
                  matrix_base + (k + q) +
                  static_cast<long long>(k + jb) * n;
              check_combined_blas(
                  cublasStrsm(
                      state.blas_handle, CUBLAS_SIDE_LEFT,
                      CUBLAS_FILL_MODE_UPPER, CUBLAS_OP_T,
                      CUBLAS_DIAG_NON_UNIT, qb, r, &one,
                      correction_diagonal, n, correction_panel, n),
                  "cublasStrsm FP32 diagonal correction for BF16 panel");

              const int source_col = k + q + qb;
              const int packed_source_col = source_col - outer;
              const int packed_source_row = k + q - outer;
              const int source_cols = n - source_col;
              const float* const source =
                  matrix_base + (k + q) +
                  static_cast<long long>(source_col) * n;
              const long long packed_offset =
                  packed_source_row +
                  static_cast<long long>(packed_source_col) * outer_width;
              const bool vector_pack =
                  (qb & 7) == 0 && (n & 3) == 0 &&
                  (outer_width & 7) == 0 &&
                  (reinterpret_cast<std::uintptr_t>(source) & 15) == 0 &&
                  (reinterpret_cast<std::uintptr_t>(packed + packed_offset) &
                   15) == 0;
              if (vector_pack) {
                const long long source_groups =
                    static_cast<long long>(qb / 8) * source_cols;
                const int blocks = static_cast<int>(std::min<long long>(
                    4096, (source_groups + threads - 1) / threads));
                pack_outer_panel_bf16_vector8<<<blocks, threads>>>(
                    source, packed, qb, source_cols, n, outer_width,
                    packed_source_row, packed_source_col);
              } else {
                const long long source_elements =
                    static_cast<long long>(qb) * source_cols;
                const int blocks = static_cast<int>(std::min<long long>(
                    4096, (source_elements + threads - 1) / threads));
                pack_outer_panel_bf16<<<blocks, threads>>>(
                    source, packed, qb, source_cols, n, outer_width,
                    packed_source_row, packed_source_col);
              }
              C10_CUDA_KERNEL_LAUNCH_CHECK();

              const int unsolved_rows = jb - q - qb;
              if (unsolved_rows > 0) {
                const __nv_bfloat16* const packed_diagonal_row =
                    packed + packed_source_row +
                    static_cast<long long>(packed_source_col) * outer_width;
                const __nv_bfloat16* const packed_panel_row =
                    packed + packed_source_row +
                    static_cast<long long>(packed_col) * outer_width;
                float* const unsolved_panel =
                    matrix_base + (k + q + qb) +
                    static_cast<long long>(k + jb) * n;
                check_combined_blas(
                    cublasGemmEx(
                        state.blas_handle, CUBLAS_OP_T, CUBLAS_OP_N,
                        unsolved_rows, r, qb, &minus_one,
                        packed_diagonal_row, CUDA_R_16BF, outer_width,
                        packed_panel_row, CUDA_R_16BF, outer_width, &one,
                        unsolved_panel, CUDA_R_32F, n, CUBLAS_COMPUTE_32F,
                        CUBLAS_GEMM_DEFAULT_TENSOR_OP),
                    "cublasGemmEx BF16 tensor panel solve update");
              }
            }
          } else {
            float* const panel =
                matrix_base + k + static_cast<long long>(k + jb) * n;
            check_combined_blas(
                cublasStrsm(
                    state.blas_handle, CUBLAS_SIDE_LEFT,
                    CUBLAS_FILL_MODE_UPPER, CUBLAS_OP_T,
                    CUBLAS_DIAG_NON_UNIT, jb, r, &one, diagonal, n, panel,
                    n),
                "cublasStrsm two-level explicit BF16 panel");

            const long long packed_offset =
                packed_row +
                static_cast<long long>(packed_col) * outer_width;
            const bool vector_pack =
                (jb & 7) == 0 && (n & 3) == 0 &&
                (outer_width & 7) == 0 &&
                (reinterpret_cast<std::uintptr_t>(panel) & 15) == 0 &&
                (reinterpret_cast<std::uintptr_t>(packed + packed_offset) &
                 15) == 0;
            if (vector_pack) {
              const long long panel_groups =
                  static_cast<long long>(jb / 8) * r;
              const int blocks = static_cast<int>(std::min<long long>(
                  4096, (panel_groups + threads - 1) / threads));
              pack_outer_panel_bf16_vector8<<<blocks, threads>>>(
                  panel, packed, jb, r, n, outer_width, packed_row,
                  packed_col);
            } else {
              const long long panel_elements =
                  static_cast<long long>(jb) * r;
              const int blocks = static_cast<int>(std::min<long long>(
                  4096, (panel_elements + threads - 1) / threads));
              pack_outer_panel_bf16<<<blocks, threads>>>(
                  panel, packed, jb, r, n, outer_width, packed_row,
                  packed_col);
            }
            C10_CUDA_KERNEL_LAUNCH_CHECK();
          }

          if (inner_remaining > 0) {
            const __nv_bfloat16* const packed_panel =
                packed + packed_row +
                static_cast<long long>(packed_col) * outer_width;
            float* const next_strip =
                matrix_base + static_cast<long long>(k + jb) * (n + 1);
            check_combined_blas(
                cublasGemmEx(
                    state.blas_handle, CUBLAS_OP_T, CUBLAS_OP_N,
                    inner_remaining, r, jb, &minus_one, packed_panel,
                    CUDA_R_16BF, outer_width, packed_panel, CUDA_R_16BF,
                    outer_width, &one, next_strip, CUDA_R_32F, n,
                    CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP),
                "cublasGemmEx two-level BF16 unfinished outer strip");
          }
        }

        const int far = n - outer_end;
        if (far > 0) {
          const __nv_bfloat16* const packed_far =
              packed + static_cast<long long>(outer_width) * outer_width;
          float* const far_diagonal =
              matrix_base + static_cast<long long>(outer_end) * (n + 1);
          const float* const input_far_diagonal =
              input_matrix_base +
              static_cast<long long>(outer_end) * (n + 1);
          if (regime_schedule) {
            auto* const primary = reinterpret_cast<__nv_fp8_e4m3*>(packed);
            auto* const primary_scale =
                reinterpret_cast<__nv_fp8_e8m0*>(
                    state.mxfp8_primary_scale.data_ptr<uint8_t>());
            float* const diagonal_correction =
                state.mxfp8_diagonal_correction.data_ptr<float>();
            const int clear_blocks =
                std::min(4096, (far + threads - 1) / threads);
            clear_mxfp8_diagonal_correction<<<clear_blocks, threads>>>(
                diagonal_correction, far);
            const float* const finalized_far_panel =
                matrix_base + outer +
                static_cast<long long>(outer_end) * n;
            pack_outer_far_mxfp8_block32<<<far, threads>>>(
                finalized_far_panel, primary, primary_scale,
                diagonal_correction, outer_width, far, n);
            C10_CUDA_KERNEL_LAUNCH_CHECK();

            for (int j = 0; j < far; j += trailing_tile) {
              const int nj = std::min(trailing_tile, far - j);
              const int m = j + nj;
              const long long packed_offset =
                  static_cast<long long>(j) * outer_width;
              const long long scale_offset =
                  mxfp8_scale_column_base(outer_width, j);
              float* const trailing_j =
                  far_diagonal + static_cast<long long>(j) * n;
              const float* const source_j =
                  outer == 0
                      ? input_far_diagonal + static_cast<long long>(j) * n
                      : trailing_j;
              lt_packed_mxfp8_block_column(
                  primary, primary + packed_offset,
                  primary_scale, primary_scale + scale_offset,
                  source_j, trailing_j, m, nj, outer_width, n,
                  device, state.mxfp8_lt_cache,
                  state.mxfp8_scratch.data_ptr<uint8_t>(),
                  mxfp8_scratch_bytes, &minus_one, &one);
            }
            add_mxfp8_diagonal_correction<<<clear_blocks, threads>>>(
                far_diagonal, diagonal_correction, far, n);
            C10_CUDA_KERNEL_LAUNCH_CHECK();
            continue;
          }
          for (int j = 0; j < far; j += trailing_tile) {
            const int nj = std::min(trailing_tile, far - j);
            const int m = j + nj;
            const __nv_bfloat16* const packed_j =
                packed_far + static_cast<long long>(j) * outer_width;
            float* const trailing_j =
                far_diagonal + static_cast<long long>(j) * n;
            if (outer == 0) {
              const float* const source_j =
                  input_far_diagonal + static_cast<long long>(j) * n;
              lt_packed_bf16_block_column(
                  packed_far, packed_j, source_j, trailing_j, m, nj,
                  outer_width, n, &minus_one, &one);
            } else {
              check_combined_blas(
                  cublasGemmEx(
                      state.blas_handle, CUBLAS_OP_T, CUBLAS_OP_N, m, nj,
                      outer_width, &minus_one, packed_far, CUDA_R_16BF,
                      outer_width, packed_j, CUDA_R_16BF, outer_width, &one,
                      trailing_j, CUDA_R_32F, n, CUBLAS_COMPUTE_32F,
                      CUBLAS_GEMM_DEFAULT_TENSOR_OP),
                  "cublasGemmEx two-level wide packed BF16 update");
            }
          }
        }
      }
    } else {
      for (int k = 0; k < n; k += standard_nb) {
        const int jb = std::min(standard_nb, n - k);
        const int r = n - k - jb;
        float* const diagonal =
            matrix_base + static_cast<long long>(k) * (n + 1);
        check_combined_solver(
            cusolverDnSpotrf(
                state.solver_handle, CUBLAS_FILL_MODE_UPPER, jb, diagonal, n,
                state.workspace.data_ptr<float>(), state.workspace_elements,
                state.info.data_ptr<int>()),
            "cusolverDnSpotrf explicit BF16");
        if (r == 0) {
          continue;
        }

        float* const panel =
            matrix_base + k + static_cast<long long>(k + jb) * n;
        check_combined_blas(
            cublasStrsm(
                state.blas_handle, CUBLAS_SIDE_LEFT, CUBLAS_FILL_MODE_UPPER,
                CUBLAS_OP_T, CUBLAS_DIAG_NON_UNIT, jb, r, &one, diagonal, n,
                panel, n),
            "cublasStrsm explicit BF16 panel");

        constexpr int threads = 256;
        const long long panel_elements = static_cast<long long>(jb) * r;
        const int blocks = static_cast<int>(
            std::min<long long>(4096,
                                (panel_elements + threads - 1) / threads));
        pack_panel_bf16<<<blocks, threads>>>(panel, packed, jb, r, n);
        C10_CUDA_KERNEL_LAUNCH_CHECK();

        float* const trailing =
            matrix_base + static_cast<long long>(k + jb) * (n + 1);
        for (int j = 0; j < r; j += trailing_tile) {
          const int nj = std::min(trailing_tile, r - j);
          const int m = j + nj;
          float* const trailing_j = trailing + static_cast<long long>(j) * n;
          const __nv_bfloat16* const packed_j =
              packed + static_cast<long long>(j) * jb;
          check_combined_blas(
              cublasGemmEx(
                  state.blas_handle, CUBLAS_OP_T, CUBLAS_OP_N, m, nj, jb,
                  &minus_one, packed, CUDA_R_16BF, jb, packed_j,
                  CUDA_R_16BF, jb, &one, trailing_j, CUDA_R_32F, n,
                  CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP),
              "cublasGemmEx packed BF16 block-column update");
        }
      }
    }
  }

  constexpr int threads = 256;
  const long long total = static_cast<long long>(batch) * matrix_elements;
  if (lower_only_initialization) {
    clear_explicit_large_upper_vectorized<<<n, threads>>>(base, n);
  } else {
    const int blocks = static_cast<int>(
        std::min<long long>(4096, (total + threads - 1) / threads));
    clear_explicit_batch_upper<<<blocks, threads>>>(base, n, total);
  }
  C10_CUDA_KERNEL_LAUNCH_CHECK();
  return output;
}

torch::Tensor cholesky_large_combined_bf16x9_cuda(torch::Tensor input) {
  TORCH_CHECK(input.is_cuda(),
              "cholesky_large_combined_bf16x9_cuda expects a CUDA tensor");
  TORCH_CHECK(input.scalar_type() == torch::kFloat32,
              "cholesky_large_combined_bf16x9_cuda expects float32 input");
  TORCH_CHECK(input.dim() == 3 && input.size(0) == 1 &&
                  input.size(1) == input.size(2) && input.size(1) >= 24576,
              "cholesky_large_combined_bf16x9_cuda expects [1, n, n] with "
              "n >= 24576");
  TORCH_CHECK(input.is_contiguous(),
              "cholesky_large_combined_bf16x9_cuda expects contiguous input");
  TORCH_CHECK(input.size(1) <= std::numeric_limits<int>::max(),
              "cholesky_large_combined_bf16x9_cuda matrix is too large");

  c10::cuda::CUDAGuard device_guard(input.device());
  const int n = static_cast<int>(input.size(1));
  auto output = torch::empty_like(input);
  constexpr int threads = 256;
  const long long total = static_cast<long long>(n) * n;
  copy_combined_large_lower<<<n, threads>>>(
      input.data_ptr<float>(), output.data_ptr<float>(), n);
  C10_CUDA_KERNEL_LAUNCH_CHECK();

  auto& state = combined_solver_state();
  std::lock_guard<std::mutex> lock(state.mutex);
  state.prepare_device(input);

  int required_elements = 0;
  check_combined_solver(
      cusolverDnSpotrf_bufferSize(
          state.large_handle, CUBLAS_FILL_MODE_UPPER, n,
          output.data_ptr<float>(), n, &required_elements),
      "cusolverDnSpotrf_bufferSize");

  if (state.workspace_elements < required_elements) {
    state.workspace = torch::empty(
        {required_elements}, input.options().dtype(torch::kFloat32));
    state.workspace_elements = required_elements;
  }
  if (!state.large_info.defined()) {
    state.large_info = torch::empty(
        {1}, input.options().dtype(torch::kInt32));
  }

  check_combined_solver(
      cusolverDnSpotrf(
          state.large_handle, CUBLAS_FILL_MODE_UPPER, n,
          output.data_ptr<float>(), n, state.workspace.data_ptr<float>(),
          required_elements, state.large_info.data_ptr<int>()),
      "cusolverDnSpotrf");

  const int blocks = static_cast<int>(
      std::min<long long>(4096, (total + threads - 1) / threads));
  clear_combined_large_upper<<<blocks, threads>>>(output.data_ptr<float>(), n);
  C10_CUDA_KERNEL_LAUNCH_CHECK();
  return output;
}
"""


_cholesky64_left_extension = load_inline(
    name="popcorn_cholesky_integration_medium_schedules_v3",
    cpp_sources=_CHOLESKY64_CPP,
    cuda_sources=_CHOLESKY64_CUDA,
    functions=[
        "cholesky64_block32_hacc4safe_cuda",
        "cholesky32_warp_cuda",
        "cholesky128_wmma_schur_cuda",
        "cholesky128_wmma_schur_exact_cuda",
        "cholesky128_block32_warp16_lowerio_anybatch_cuda",
        "cholesky256_batch64_register_panel64_cuda",
        "cholesky256_combined_batched_cuda",
        "cholesky512_batch16_register_panel64_cuda",
        "cholesky512_combined_batched_cuda",
        "cholesky512_highbatch_blocked_cuda",
        "cholesky1024_batch4_blocked_cuda",
        "cholesky1024_batch60_blocked_cuda",
        "cholesky2048_batch2_blocked_cuda",
        "cholesky2048_batch8_blocked_cuda",
        "cholesky4096_batch2_truebatched_fast16bf_cuda",
        "cholesky_large_explicit_bf16_cuda",
        "cholesky_large_combined_bf16x9_cuda",
    ],
    extra_cflags=["-O3"],
    extra_cuda_cflags=["-O3", "--use_fast_math"],
    extra_ldflags=["-lcusolver", "-lcublas", "-lcublasLt"],
    with_cuda=True,
    verbose=False,
)


@triton.jit
def _cholesky32_left_kernel(input_ptr, output_ptr, matrix_stride: tl.constexpr):
    matrix = tl.program_id(0)
    row_ids = tl.arange(0, 32)
    col_ids = tl.arange(0, 32)
    rows = row_ids[:, None]
    cols = col_ids[None, :]
    offsets = matrix * matrix_stride + rows * 32 + cols
    values = tl.where(rows >= cols, tl.load(input_ptr + offsets), 0.0)

    for k in range(32):
        row = tl.sum(tl.where(rows == k, values, 0.0), axis=0)
        diagonal = tl.sum(tl.where(col_ids == k, row, 0.0), axis=0)
        diagonal -= tl.sum(tl.where(col_ids < k, row * row, 0.0), axis=0)
        diagonal = tl.sqrt(tl.maximum(diagonal, 0.0))

        column = tl.sum(tl.where(cols == k, values, 0.0), axis=1)
        products = tl.where(cols < k, values * row[None, :], 0.0)
        column = (column - tl.sum(products, axis=1)) / diagonal
        values = tl.where((rows == k) & (cols == k), diagonal, values)
        values = tl.where((rows > k) & (cols == k), column[:, None], values)

    tl.store(output_ptr + offsets, values)


def custom_kernel(data: input_t) -> output_t:
    batch, n, _ = data.shape
    if batch == 2 and n == 4096:
        return (
            _cholesky64_left_extension
            .cholesky4096_batch2_truebatched_fast16bf_cuda(data)
        )

    if batch == 1 and n in (16384, 32768):
        return _cholesky64_left_extension.cholesky_large_explicit_bf16_cuda(
            data
        )

    if batch == 1 and n >= 24576:
        return _cholesky64_left_extension.cholesky_large_combined_bf16x9_cuda(
            data
        )

    if batch == 64 and n == 256:
        return (
            _cholesky64_left_extension
            .cholesky256_batch64_register_panel64_cuda(data)
        )

    if batch >= 128 and n == 256:
        return _cholesky64_left_extension.cholesky256_combined_batched_cuda(
            data
        )

    if batch == 640 and n == 512:
        return (
            _cholesky64_left_extension
            .cholesky512_highbatch_blocked_cuda(data)
        )

    if batch == 16 and n == 512:
        return (
            _cholesky64_left_extension
            .cholesky512_batch16_register_panel64_cuda(data)
        )

    if batch == 4 and n == 1024:
        return (
            _cholesky64_left_extension
            .cholesky1024_batch4_blocked_cuda(data)
        )

    if batch == 60 and n == 1024:
        return (
            _cholesky64_left_extension
            .cholesky1024_batch60_blocked_cuda(data)
        )

    if batch == 2 and n == 2048:
        return (
            _cholesky64_left_extension
            .cholesky2048_batch2_blocked_cuda(data)
        )

    if batch == 8 and n == 2048:
        return (
            _cholesky64_left_extension
            .cholesky2048_batch8_blocked_cuda(data)
        )

    if batch >= 256 and n == 512:
        return _cholesky64_left_extension.cholesky512_combined_batched_cuda(
            data
        )

    if batch > 0 and n == 32:
        return _cholesky64_left_extension.cholesky32_warp_cuda(data)

    if batch > 0 and n == 64:
        return _cholesky64_left_extension.cholesky64_block32_hacc4safe_cuda(data)

    if batch == 256 and n == 128:
        return _cholesky64_left_extension.cholesky128_wmma_schur_exact_cuda(
            data
        )

    if batch >= 16 and n == 128:
        return (
            _cholesky64_left_extension
            .cholesky128_block32_warp16_lowerio_anybatch_cuda(data)
        )

    return torch.linalg.cholesky_ex(data, check_errors=False).L
scrolls · 3999 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