Skip to content
KernelIndex
Search⌘K

submission 920149

shikhar · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-920149?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
1.12ms
#132 of 337
2026-07-28

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:25a960b876d2e14e2bdc3b9b371432f881c174c27484e3e501b016a183af21cf
license declaredunknown
license concludedunknown
authorsshikhar
imported2026-08-26

Techniques

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

cluster__cluster_dims__(2, 1, 1)
mbarrier"mbarrier.init.shared::cta.b64 [%0], %1;"
shared-memoryextern __shared__ float shared[];
tma"cp.async.bulk.shared::cluster.shared::cta."
vector-width = float4const float4* input_vectors =

Kernel source

submission.py1673 lines
from __future__ import annotations

import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline


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

torch::Tensor cholesky_small_cuda(const torch::Tensor& input);
void factor_owned(at::Tensor A, at::Tensor L, at::Tensor Lh, at::Tensor Ll,
                  at::Tensor sneg, at::Tensor spos,
                  at::Tensor invD, at::Tensor invDh, at::Tensor invDl,
                  at::Tensor Ph, at::Tensor Pl, at::Tensor PF,
                  at::Tensor ws, int64_t leaf);
"""


CUDA_SRC = r"""
#include <ATen/ATen.h>
#include <ATen/cuda/CUDAContextLight.h>
#include <c10/cuda/CUDAGuard.h>
#include <c10/cuda/CUDAException.h>
#include <cublasLt.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <torch/extension.h>

namespace {

__device__ __forceinline__ void mbar_init(int address, int count) {
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
  asm volatile(
      "mbarrier.init.shared::cta.b64 [%0], %1;"
      :: "r"(address), "r"(count));
#endif
}

__device__ __forceinline__ void mbar_arrive(int address) {
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
  asm volatile(
      "mbarrier.arrive.release.cta.shared::cluster.b64 _, [%0];"
      :: "r"(address)
      : "memory");
#endif
}

__device__ inline void mbar_wait(int address) {
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
  constexpr int ticks = 0x989680;
  asm volatile(
      "{\n\t"
      ".reg .pred ready;\n\t"
      "chol_mbar_wait_loop:\n\t"
      "mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 "
      "ready, [%0], 0, %1;\n\t"
      "@!ready bra.uni chol_mbar_wait_loop;\n\t"
      "}"
      :: "r"(address), "r"(ticks));
#endif
}

__device__ __forceinline__ void mbar_expect_tx(int address, int bytes) {
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
  asm volatile(
      "mbarrier.arrive.expect_tx.relaxed.cluster.shared::cluster.b64 "
      "_, [%0], %1;"
      :: "r"(address), "r"(bytes)
      : "memory");
#endif
}

__device__ __forceinline__ void tma_shared_to_remote_shared(
    int destination,
    int source,
    int bytes,
    int mbar) {
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
  asm volatile(
      "cp.async.bulk.shared::cluster.shared::cta."
      "mbarrier::complete_tx::bytes [%0], [%1], %2, [%3];"
      :: "r"(destination), "r"(source), "r"(bytes), "r"(mbar));
#endif
}

// One warp owns one complete 32x32 matrix. Lane i keeps lower-triangular row i
// in registers. At pivot k, lane k publishes the diagonal and lane j publishes
// L[j,k] with shuffles; there is no shared memory or CTA-wide synchronization.
__global__ void __launch_bounds__(256)
warp_cholesky_32_kernel(
    const float* __restrict__ input,
    float* __restrict__ output,
    int batch) {
  constexpr int N = 32;
  constexpr int kWarpsPerBlock = 8;

  const int lane = threadIdx.x & 31;
  const int warp = threadIdx.x >> 5;
  const int matrix = blockIdx.x * kWarpsPerBlock + warp;
  if (matrix >= batch) return;

  const int64_t matrix_offset = static_cast<int64_t>(matrix) * N * N;
  const float* matrix_input = input + matrix_offset;
  float* matrix_output = output + matrix_offset;

  float row[N];
  const float4* input_vectors =
      reinterpret_cast<const float4*>(matrix_input + lane * N);
  #pragma unroll
  for (int vector_index = 0; vector_index < N / 4; ++vector_index) {
    const float4 values = input_vectors[vector_index];
    const int j = vector_index * 4;
    row[j + 0] = lane >= j + 0 ? values.x : 0.0f;
    row[j + 1] = lane >= j + 1 ? values.y : 0.0f;
    row[j + 2] = lane >= j + 2 ? values.z : 0.0f;
    row[j + 3] = lane >= j + 3 ? values.w : 0.0f;
  }

  constexpr unsigned kFullMask = 0xffffffffu;
  #pragma unroll
  for (int k = 0; k < N; ++k) {
    const float pivot = __shfl_sync(kFullMask, row[k], k);
    const float reciprocal = rsqrtf(pivot);
    const float diagonal = pivot * reciprocal;
    if (lane == k) {
      row[k] = diagonal;
    } else if (lane > k) {
      row[k] *= reciprocal;
    }

    const float row_factor = row[k];
    #pragma unroll
    for (int j = k + 1; j < N; ++j) {
      const float pivot_factor =
          __shfl_sync(kFullMask, row_factor, j);
      if (lane >= j) {
        row[j] = fmaf(-row_factor, pivot_factor, row[j]);
      }
    }
  }

  float4* output_vectors =
      reinterpret_cast<float4*>(matrix_output + lane * N);
  #pragma unroll
  for (int vector_index = 0; vector_index < N / 4; ++vector_index) {
    const int j = vector_index * 4;
    float4 values;
    values.x = lane >= j + 0 ? row[j + 0] : 0.0f;
    values.y = lane >= j + 1 ? row[j + 1] : 0.0f;
    values.z = lane >= j + 2 ? row[j + 2] : 0.0f;
    values.w = lane >= j + 3 ? row[j + 3] : 0.0f;
    output_vectors[vector_index] = values;
  }
}

// A warp owns one 64x64 matrix as two rows per lane. The first row only needs
// its first 32 lower-triangular entries; the second needs all 64, for 96 live
// matrix values per lane. The two 32-column phases make every shuffle source
// compile-time constant and keep the arrays scalarized into registers.
__global__ void __launch_bounds__(256)
warp_cholesky_64_kernel(
    const float* __restrict__ input,
    float* __restrict__ output,
    int batch) {
  constexpr int N = 64;
  constexpr int kWarpsPerBlock = 8;

  const int lane = threadIdx.x & 31;
  const int warp = threadIdx.x >> 5;
  const int matrix = blockIdx.x * kWarpsPerBlock + warp;
  if (matrix >= batch) return;

  const int64_t matrix_offset = static_cast<int64_t>(matrix) * N * N;
  const float* matrix_input = input + matrix_offset;
  float* matrix_output = output + matrix_offset;

  float first[32];
  float second[64];
  const float4* first_input =
      reinterpret_cast<const float4*>(matrix_input + lane * N);
  const float4* second_input =
      reinterpret_cast<const float4*>(
          matrix_input + (lane + 32) * N);

  #pragma unroll
  for (int vector_index = 0; vector_index < 8; ++vector_index) {
    const float4 values = first_input[vector_index];
    const int j = vector_index * 4;
    first[j + 0] = lane >= j + 0 ? values.x : 0.0f;
    first[j + 1] = lane >= j + 1 ? values.y : 0.0f;
    first[j + 2] = lane >= j + 2 ? values.z : 0.0f;
    first[j + 3] = lane >= j + 3 ? values.w : 0.0f;
  }
  #pragma unroll
  for (int vector_index = 0; vector_index < 16; ++vector_index) {
    const float4 values = second_input[vector_index];
    const int j = vector_index * 4;
    second[j + 0] = lane + 32 >= j + 0 ? values.x : 0.0f;
    second[j + 1] = lane + 32 >= j + 1 ? values.y : 0.0f;
    second[j + 2] = lane + 32 >= j + 2 ? values.z : 0.0f;
    second[j + 3] = lane + 32 >= j + 3 ? values.w : 0.0f;
  }

  constexpr unsigned kFullMask = 0xffffffffu;
  #pragma unroll
  for (int k = 0; k < 32; ++k) {
    const float pivot = __shfl_sync(kFullMask, first[k], k);
    const float reciprocal = rsqrtf(pivot);
    const float diagonal = pivot * reciprocal;
    if (lane == k) {
      first[k] = diagonal;
    } else if (lane > k) {
      first[k] *= reciprocal;
    }
    second[k] *= reciprocal;

    const float first_factor = first[k];
    const float second_factor = second[k];
    #pragma unroll
    for (int j = k + 1; j < 32; ++j) {
      const float pivot_factor =
          __shfl_sync(kFullMask, first_factor, j);
      if (lane >= j) {
        first[j] = fmaf(-first_factor, pivot_factor, first[j]);
      }
      second[j] = fmaf(-second_factor, pivot_factor, second[j]);
    }
    #pragma unroll
    for (int j = 32; j < 64; ++j) {
      const float pivot_factor =
          __shfl_sync(kFullMask, second_factor, j - 32);
      if (lane + 32 >= j) {
        second[j] = fmaf(-second_factor, pivot_factor, second[j]);
      }
    }
  }

  #pragma unroll
  for (int k = 32; k < 64; ++k) {
    const int pivot_lane = k - 32;
    const float pivot =
        __shfl_sync(kFullMask, second[k], pivot_lane);
    const float reciprocal = rsqrtf(pivot);
    const float diagonal = pivot * reciprocal;
    if (lane == pivot_lane) {
      second[k] = diagonal;
    } else if (lane > pivot_lane) {
      second[k] *= reciprocal;
    }

    const float row_factor = second[k];
    #pragma unroll
    for (int j = k + 1; j < 64; ++j) {
      const float pivot_factor =
          __shfl_sync(kFullMask, row_factor, j - 32);
      if (lane + 32 >= j) {
        second[j] = fmaf(-row_factor, pivot_factor, second[j]);
      }
    }
  }

  float4* first_output =
      reinterpret_cast<float4*>(matrix_output + lane * N);
  float4* second_output =
      reinterpret_cast<float4*>(
          matrix_output + (lane + 32) * N);
  #pragma unroll
  for (int vector_index = 0; vector_index < 16; ++vector_index) {
    const int j = vector_index * 4;
    float4 first_values;
    first_values.x = j + 0 < 32 && lane >= j + 0
        ? first[j + 0] : 0.0f;
    first_values.y = j + 1 < 32 && lane >= j + 1
        ? first[j + 1] : 0.0f;
    first_values.z = j + 2 < 32 && lane >= j + 2
        ? first[j + 2] : 0.0f;
    first_values.w = j + 3 < 32 && lane >= j + 3
        ? first[j + 3] : 0.0f;
    first_output[vector_index] = first_values;

    float4 second_values;
    second_values.x = lane + 32 >= j + 0 ? second[j + 0] : 0.0f;
    second_values.y = lane + 32 >= j + 1 ? second[j + 1] : 0.0f;
    second_values.z = lane + 32 >= j + 2 ? second[j + 2] : 0.0f;
    second_values.w = lane + 32 >= j + 3 ? second[j + 3] : 0.0f;
    second_output[vector_index] = second_values;
  }
}

// One warp owns eight complete columns. Lanes stripe rows and keep the active
// matrix in registers. A completed Cholesky column is broadcast through a
// padded shared vector and applied to every later resident column.
template <int N>
__global__ void register_cholesky_kernel(
    const float* __restrict__ input,
    float* __restrict__ output) {
  constexpr int kColsPerWarp = 8;
  constexpr int kWarps = N / kColsPerWarp;
  constexpr int kThreads = kWarps * 32;
  constexpr int kRowItems = N / 32;
  constexpr int kLd = N + 1;

  static_assert(N == 32 || N == 64 || N == 128);

  extern __shared__ float shared[];
  float* tile = shared;
  float* column = tile + N * kLd;

  const int tid = threadIdx.x;
  const int warp = tid >> 5;
  const int lane = tid & 31;
  const int64_t matrix_offset = static_cast<int64_t>(blockIdx.x) * N * N;
  const float* matrix_input = input + matrix_offset;
  float* matrix_output = output + matrix_offset;

  for (int index = tid; index < N * N; index += kThreads) {
    const int row = index / N;
    const int col = index - row * N;
    tile[row * kLd + col] = matrix_input[index];
  }
  __syncthreads();

  float columns[kRowItems][kColsPerWarp];
  #pragma unroll
  for (int item = 0; item < kRowItems; ++item) {
    const int row = item * 32 + lane;
    #pragma unroll
    for (int p = 0; p < kColsPerWarp; ++p) {
      const int col = warp * kColsPerWarp + p;
      columns[item][p] = row >= col ? tile[row * kLd + col] : 0.0f;
    }
  }
  __syncthreads();

  #pragma unroll 1
  for (int k = 0; k < N; ++k) {
    const int owner = k / kColsPerWarp;
    const int local_col = k & (kColsPerWarp - 1);

    if (warp == owner) {
      float pivot = 0.0f;
      #pragma unroll
      for (int item = 0; item < kRowItems; ++item) {
        if (item * 32 + lane == k) {
          pivot = columns[item][local_col];
        }
      }
      pivot = __shfl_sync(0xffffffffu, pivot, k & 31);

      const float diagonal = sqrtf(pivot);
      const float reciprocal = 1.0f / diagonal;

      #pragma unroll
      for (int item = 0; item < kRowItems; ++item) {
        const int row = item * 32 + lane;
        float value = columns[item][local_col];
        if (row == k) {
          value = diagonal;
          columns[item][local_col] = value;
          column[row] = value;
        } else if (row > k) {
          value *= reciprocal;
          columns[item][local_col] = value;
          column[row] = value;
        }
      }
    }
    __syncthreads();

    #pragma unroll
    for (int p = 0; p < kColsPerWarp; ++p) {
      const int col = warp * kColsPerWarp + p;
      if (col > k) {
        const float col_factor = column[col];
        #pragma unroll
        for (int item = 0; item < kRowItems; ++item) {
          const int row = item * 32 + lane;
          if (row >= col) {
            columns[item][p] =
                fmaf(-column[row], col_factor, columns[item][p]);
          }
        }
      }
    }
    __syncthreads();
  }

  #pragma unroll
  for (int item = 0; item < kRowItems; ++item) {
    const int row = item * 32 + lane;
    #pragma unroll
    for (int p = 0; p < kColsPerWarp; ++p) {
      const int col = warp * kColsPerWarp + p;
      tile[row * kLd + col] = row >= col ? columns[item][p] : 0.0f;
    }
  }
  __syncthreads();

  for (int index = tid; index < N * N; index += kThreads) {
    const int row = index / N;
    const int col = index - row * N;
    matrix_output[index] = tile[row * kLd + col];
  }
}

// GAU-style warp wavefront. There are no per-column CTA barriers: warp w keeps
// its eight columns resident, consumes the 8*w earlier columns, and then
// publishes its own columns through one 32-arrival mbarrier per column.
template <int N>
__global__ void wavefront_cholesky_kernel(
    const float* __restrict__ input,
    float* __restrict__ output) {
  constexpr int kColsPerWarp = 8;
  constexpr int kRowItems = N / 32;

  static_assert(N == 64 || N == 128);

  extern __shared__ float shared[];
  float* reflectors = shared;  // [column, row]
  const int mbar_base =
      __cvta_generic_to_shared(reflectors + N * N);

  const int tid = threadIdx.x;
  const int warp = tid >> 5;
  const int lane = tid & 31;
  const int column_base = warp * kColsPerWarp;
  const int64_t matrix_offset =
      static_cast<int64_t>(blockIdx.x) * N * N;
  const float* matrix_input = input + matrix_offset;
  float* matrix_output = output + matrix_offset;

  if (warp == 0 && lane == 0) {
    #pragma unroll 1
    for (int k = 0; k < N; ++k) {
      mbar_init(mbar_base + k * 8, 32);
    }
  }
  __syncthreads();

  float columns[kRowItems][kColsPerWarp];
  #pragma unroll
  for (int item = 0; item < kRowItems; ++item) {
    const int row = item * 32 + lane;
    const float* source = matrix_input + row * N + column_base;
    const float4 lo = reinterpret_cast<const float4*>(source)[0];
    const float4 hi = reinterpret_cast<const float4*>(source)[1];
    columns[item][0] = row >= column_base + 0 ? lo.x : 0.0f;
    columns[item][1] = row >= column_base + 1 ? lo.y : 0.0f;
    columns[item][2] = row >= column_base + 2 ? lo.z : 0.0f;
    columns[item][3] = row >= column_base + 3 ? lo.w : 0.0f;
    columns[item][4] = row >= column_base + 4 ? hi.x : 0.0f;
    columns[item][5] = row >= column_base + 5 ? hi.y : 0.0f;
    columns[item][6] = row >= column_base + 6 ? hi.z : 0.0f;
    columns[item][7] = row >= column_base + 7 ? hi.w : 0.0f;
  }

  #pragma unroll 1
  for (int earlier_warp = 0; earlier_warp < warp; ++earlier_warp) {
    #pragma unroll
    for (int p = 0; p < kColsPerWarp; ++p) {
      const int k = earlier_warp * kColsPerWarp + p;
      mbar_wait(mbar_base + k * 8);

      float factors[kColsPerWarp];
      #pragma unroll
      for (int q = 0; q < kColsPerWarp; ++q) {
        factors[q] = reflectors[k * N + column_base + q];
      }
      #pragma unroll
      for (int item = 0; item < kRowItems; ++item) {
        const int row = item * 32 + lane;
        const float row_factor = reflectors[k * N + row];
        #pragma unroll
        for (int q = 0; q < kColsPerWarp; ++q) {
          columns[item][q] =
              fmaf(-row_factor, factors[q], columns[item][q]);
        }
      }
    }
  }

  #pragma unroll
  for (int p = 0; p < kColsPerWarp; ++p) {
    const int k = column_base + p;
    float pivot = 0.0f;
    #pragma unroll
    for (int item = 0; item < kRowItems; ++item) {
      if (item * 32 + lane == k) {
        pivot = columns[item][p];
      }
    }
    pivot = __shfl_sync(0xffffffffu, pivot, k & 31);
    const float reciprocal = rsqrtf(pivot);
    const float diagonal = pivot * reciprocal;

    float completed[kRowItems];
    #pragma unroll
    for (int item = 0; item < kRowItems; ++item) {
      const int row = item * 32 + lane;
      float value = 0.0f;
      if (row == k) {
        value = diagonal;
      } else if (row > k) {
        value = columns[item][p] * reciprocal;
      }
      completed[item] = value;
      columns[item][p] = value;
      reflectors[k * N + row] = value;
    }
    mbar_arrive(mbar_base + k * 8);

    #pragma unroll
    for (int q = p + 1; q < kColsPerWarp; ++q) {
      const int trailing_col = column_base + q;
      float col_factor = 0.0f;
      #pragma unroll
      for (int item = 0; item < kRowItems; ++item) {
        if (item * 32 + lane == trailing_col) {
          col_factor = completed[item];
        }
      }
      col_factor =
          __shfl_sync(0xffffffffu, col_factor, trailing_col & 31);
      #pragma unroll
      for (int item = 0; item < kRowItems; ++item) {
        columns[item][q] =
            fmaf(-completed[item], col_factor, columns[item][q]);
      }
    }
  }

  #pragma unroll
  for (int item = 0; item < kRowItems; ++item) {
    const int row = item * 32 + lane;
    float* destination = matrix_output + row * N + column_base;
    const float4 lo = {
        columns[item][0], columns[item][1],
        columns[item][2], columns[item][3]};
    const float4 hi = {
        columns[item][4], columns[item][5],
        columns[item][6], columns[item][7]};
    reinterpret_cast<float4*>(destination)[0] = lo;
    reinterpret_cast<float4*>(destination)[1] = hi;
  }
}

// The N=256 benchmark has batch=64: exactly 64 two-SM clusters / 128 CTAs.
__global__
__cluster_dims__(2, 1, 1)
__launch_bounds__(512, 1)
void cluster_cholesky_256_kernel(
    const float* __restrict__ input,
    float* __restrict__ output) {
  constexpr int N = 256;
  constexpr int kColsPerWarp = 8;
  constexpr int kRowItems = 8;
  constexpr int kLocalCols = 128;

  extern __shared__ float shared[];
  float* reflectors = shared;  // [local column, row]
  const int reflector_address =
      __cvta_generic_to_shared(reflectors);
  const int mbar_base =
      reflector_address + kLocalCols * N * sizeof(float);
  const int remote_reflector_address =
      reflector_address | 0x01000000;

  const int tid = threadIdx.x;
  const int warp = tid >> 5;
  const int lane = tid & 31;
  const int rank = blockIdx.x & 1;
  const int batch = blockIdx.x >> 1;
  const int column_base = rank * kLocalCols + warp * kColsPerWarp;
  const int64_t matrix_offset =
      static_cast<int64_t>(batch) * N * N;
  const float* matrix_input = input + matrix_offset;
  float* matrix_output = output + matrix_offset;

  if (warp == 0 && lane == 0) {
    #pragma unroll 1
    for (int k = 0; k < N; ++k) {
      mbar_init(mbar_base + k * 8, 1);
    }
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
    asm volatile("fence.mbarrier_init.release.cluster;");
#endif
  }
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
  asm volatile("barrier.cluster.arrive.relaxed.aligned;");
  asm volatile("barrier.cluster.wait.acquire.aligned;");
#else
  __syncthreads();
#endif

  float columns[kRowItems][kColsPerWarp];
  #pragma unroll
  for (int item = 0; item < kRowItems; ++item) {
    const int row = item * 32 + lane;
    const float* source = matrix_input + row * N + column_base;
    const float4 lo = reinterpret_cast<const float4*>(source)[0];
    const float4 hi = reinterpret_cast<const float4*>(source)[1];
    columns[item][0] = row >= column_base + 0 ? lo.x : 0.0f;
    columns[item][1] = row >= column_base + 1 ? lo.y : 0.0f;
    columns[item][2] = row >= column_base + 2 ? lo.z : 0.0f;
    columns[item][3] = row >= column_base + 3 ? lo.w : 0.0f;
    columns[item][4] = row >= column_base + 4 ? hi.x : 0.0f;
    columns[item][5] = row >= column_base + 5 ? hi.y : 0.0f;
    columns[item][6] = row >= column_base + 6 ? hi.z : 0.0f;
    columns[item][7] = row >= column_base + 7 ? hi.w : 0.0f;
  }

  if (rank == 1) {
    #pragma unroll 1
    for (int k = 0; k < kLocalCols; ++k) {
      mbar_wait(mbar_base + k * 8);
      float factors[kColsPerWarp];
      #pragma unroll
      for (int q = 0; q < kColsPerWarp; ++q) {
        factors[q] = reflectors[k * N + column_base + q];
      }
      #pragma unroll
      for (int item = 0; item < kRowItems; ++item) {
        const int row = item * 32 + lane;
        const float row_factor = reflectors[k * N + row];
        #pragma unroll
        for (int q = 0; q < kColsPerWarp; ++q) {
          columns[item][q] =
              fmaf(-row_factor, factors[q], columns[item][q]);
        }
      }
    }
  }
  __syncthreads();

  #pragma unroll 1
  for (int earlier_warp = 0; earlier_warp < warp; ++earlier_warp) {
    #pragma unroll
    for (int p = 0; p < kColsPerWarp; ++p) {
      const int global_k =
          rank * kLocalCols + earlier_warp * kColsPerWarp + p;
      const int local_k = earlier_warp * kColsPerWarp + p;
      mbar_wait(mbar_base + global_k * 8);
      float factors[kColsPerWarp];
      #pragma unroll
      for (int q = 0; q < kColsPerWarp; ++q) {
        factors[q] = reflectors[local_k * N + column_base + q];
      }
      #pragma unroll
      for (int item = 0; item < kRowItems; ++item) {
        const int row = item * 32 + lane;
        const float row_factor = reflectors[local_k * N + row];
        #pragma unroll
        for (int q = 0; q < kColsPerWarp; ++q) {
          columns[item][q] =
              fmaf(-row_factor, factors[q], columns[item][q]);
        }
      }
    }
  }

  #pragma unroll
  for (int p = 0; p < kColsPerWarp; ++p) {
    const int global_k = column_base + p;
    const int local_k = warp * kColsPerWarp + p;

    float pivot = 0.0f;
    #pragma unroll
    for (int item = 0; item < kRowItems; ++item) {
      if (item * 32 + lane == global_k) {
        pivot = columns[item][p];
      }
    }
    pivot = __shfl_sync(0xffffffffu, pivot, global_k & 31);
    const float reciprocal = rsqrtf(pivot);
    const float diagonal = pivot * reciprocal;

    float completed[kRowItems];
    #pragma unroll
    for (int item = 0; item < kRowItems; ++item) {
      const int row = item * 32 + lane;
      float value = 0.0f;
      if (row == global_k) {
        value = diagonal;
      } else if (row > global_k) {
        value = columns[item][p] * reciprocal;
      }
      completed[item] = value;
      columns[item][p] = value;
      reflectors[local_k * N + row] = value;
    }

    __syncwarp();
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900
    asm volatile("fence.proxy.async.shared::cta;");
#endif
    if (lane == 0) {
      mbar_arrive(mbar_base + global_k * 8);
      if (rank == 0) {
        const int remote_mbar =
            (mbar_base + global_k * 8) | 0x01000000;
        mbar_expect_tx(remote_mbar, N * sizeof(float));
        tma_shared_to_remote_shared(
            remote_reflector_address + global_k * N * sizeof(float),
            reflector_address + local_k * N * sizeof(float),
            N * sizeof(float),
            remote_mbar);
      }
    }

    #pragma unroll
    for (int q = p + 1; q < kColsPerWarp; ++q) {
      const int trailing_col = column_base + q;
      float col_factor = 0.0f;
      #pragma unroll
      for (int item = 0; item < kRowItems; ++item) {
        if (item * 32 + lane == trailing_col) {
          col_factor = completed[item];
        }
      }
      col_factor =
          __shfl_sync(0xffffffffu, col_factor, trailing_col & 31);
      #pragma unroll
      for (int item = 0; item < kRowItems; ++item) {
        columns[item][q] =
            fmaf(-completed[item], col_factor, columns[item][q]);
      }
    }
  }

  #pragma unroll
  for (int item = 0; item < kRowItems; ++item) {
    const int row = item * 32 + lane;
    float* destination = matrix_output + row * N + column_base;
    const float4 lo = {
        columns[item][0], columns[item][1],
        columns[item][2], columns[item][3]};
    const float4 hi = {
        columns[item][4], columns[item][5],
        columns[item][6], columns[item][7]};
    reinterpret_cast<float4*>(destination)[0] = lo;
    reinterpret_cast<float4*>(destination)[1] = hi;
  }
}

// ---- Blocked/recursive leaf: 128x128 diagonal factor + explicit inverse ----
//
// Layout inside one CTA (256 threads, one matrix): the 128x128 tile is
// processed in four 32-wide steps. Each step factors its 32x32 corner and
// its inverse entirely inside one warp (register-resident, shuffle
// reductions), then solves and updates the rest of the tile element-parallel
// across all 256 threads. The full 128x128 lower-triangular inverse is then
// assembled from the 32x32 blocks, and emitted in fp32 plus fp16 hi/lo halves
// for the compensated tensor-core panel GEMMs.
__global__ void __launch_bounds__(256)
potrf_inv_kernel(float* __restrict__ L,
                 float* __restrict__ invD,
                 __half* __restrict__ invDh,
                 __half* __restrict__ invDl,
                 long n, long mat_elems, long inv_mat_elems, int k,
                 bool emit_half_inverse) {
  constexpr int NB = 128;
  constexpr int LD = 129;
  constexpr int THREADS = 256;
  extern __shared__ float smem[];
  float* S = smem;             // NB * LD
  float* X = smem + NB * LD;   // NB * LD  (inverse blocks)
  float* W = X + NB * LD;      // 32 * 33 scratch
  const long b = blockIdx.x;
  float* Db = L + b * mat_elems + (long)k * n + k;
  const long inv_off = b * inv_mat_elems + (long)(k / NB) * NB * NB;
  const int tid = threadIdx.x;
  const int warp = tid >> 5;
  const int lane = tid & 31;

  for (int idx = tid; idx < NB * NB; idx += THREADS) {
    const int i = idx / NB;
    const int j = idx - i * NB;
    S[i * LD + j] = Db[(long)i * n + j];
  }
  __syncthreads();

  #pragma unroll 1
  for (int kb = 0; kb < 4; ++kb) {
    const int c0 = kb * 32;

    // (a) warp 0: factor the 32x32 corner and its inverse in registers.
    if (warp == 0) {
      float r0[32];
      #pragma unroll
      for (int j = 0; j < 32; ++j) {
        r0[j] = S[(c0 + lane) * LD + c0 + j];
      }
      #pragma unroll
      for (int kk = 0; kk < 32; ++kk) {
        const float piv = __shfl_sync(0xffffffffu, r0[kk], kk);
        const float rc = rsqrtf(piv);
        const float d = piv * rc;
        if (lane == kk) {
          r0[kk] = d;
        } else if (lane > kk) {
          r0[kk] *= rc;
        }
        const float lrk = r0[kk];
        #pragma unroll
        for (int j = 0; j < 32; ++j) {
          if (j > kk) {
            const float ljk = __shfl_sync(0xffffffffu, r0[kk], j);
            if (lane >= j) {
              r0[j] = fmaf(-lrk, ljk, r0[j]);
            }
          }
        }
      }
      #pragma unroll
      for (int j = 0; j < 32; ++j) {
        S[(c0 + lane) * LD + c0 + j] = r0[j];
      }
      // Inverse of the corner: lane c solves L y = e_c. y[m] is zero above
      // the column, which drops the guard from the inner reduction.
      float y[32];
      #pragma unroll
      for (int r = 0; r < 32; ++r) {
        float acc = (r == lane) ? 1.0f : 0.0f;
        #pragma unroll
        for (int m = 0; m < 32; ++m) {
          if (m < r) {
            const float lrm = __shfl_sync(0xffffffffu, r0[m], r);
            acc = fmaf(-lrm, y[m], acc);
          }
        }
        const float drr = __shfl_sync(0xffffffffu, r0[r], r);
        y[r] = (r >= lane) ? acc / drr : 0.0f;
      }
      #pragma unroll
      for (int r = 0; r < 32; ++r) {
        X[(c0 + r) * LD + c0 + lane] = y[r];
      }
    }
    __syncthreads();

    const int rows = NB - c0 - 32;
    if (rows > 0) {
      // (b) panel solve below the corner: rows (c0+32..127) times inv^T,
      // element-parallel with a register stash so reads complete before
      // writes. Static trip count keeps the stash out of local memory.
      float out[12];
      #pragma unroll
      for (int t = 0; t < 12; ++t) {
        const int idx = tid + t * THREADS;
        if (idx < rows * 32) {
          const int i = c0 + 32 + idx / 32;
          const int j = idx & 31;
          float acc = 0.0f;
          #pragma unroll
          for (int m = 0; m < 32; ++m) {
            acc = fmaf(S[i * LD + c0 + m], X[(c0 + j) * LD + c0 + m], acc);
          }
          out[t] = acc;
        }
      }
      __syncthreads();
      #pragma unroll
      for (int t = 0; t < 12; ++t) {
        const int idx = tid + t * THREADS;
        if (idx < rows * 32) {
          const int i = c0 + 32 + idx / 32;
          const int j = idx & 31;
          S[i * LD + c0 + j] = out[t];
        }
      }
      __syncthreads();

      // (c) rank-32 trailing update; writes land strictly right of the
      // panel columns, so in-place is safe.
      #pragma unroll 1
      for (int idx = tid; idx < rows * rows; idx += THREADS) {
        const int i = c0 + 32 + idx / rows;
        const int j = c0 + 32 + (idx - (idx / rows) * rows);
        if (j <= i) {
          float acc = S[i * LD + j];
          #pragma unroll
          for (int m = 0; m < 32; ++m) {
            acc = fmaf(-S[i * LD + c0 + m], S[j * LD + c0 + m], acc);
          }
          S[i * LD + j] = acc;
        }
      }
      __syncthreads();
    }
  }

  // Assemble off-diagonal inverse blocks: X_ic = -X_ii (sum_m L_im X_mc).
  #pragma unroll 1
  for (int d = 1; d < 4; ++d) {
    #pragma unroll 1
    for (int c = 0; c + d < 4; ++c) {
      const int i = c + d;
      #pragma unroll 1
      for (int idx = tid; idx < 32 * 32; idx += THREADS) {
        const int r = idx >> 5;
        const int cc = idx & 31;
        float acc = 0.0f;
        for (int mb = c; mb < i; ++mb) {
          #pragma unroll
          for (int m = 0; m < 32; ++m) {
            acc = fmaf(S[(i * 32 + r) * LD + mb * 32 + m],
                       X[(mb * 32 + m) * LD + c * 32 + cc], acc);
          }
        }
        W[r * 33 + cc] = acc;
      }
      __syncthreads();
      #pragma unroll 1
      for (int idx = tid; idx < 32 * 32; idx += THREADS) {
        const int r = idx >> 5;
        const int cc = idx & 31;
        float acc = 0.0f;
        #pragma unroll
        for (int m = 0; m < 32; ++m) {
          acc = fmaf(X[(i * 32 + r) * LD + i * 32 + m], W[m * 33 + cc], acc);
        }
        X[(i * 32 + r) * LD + c * 32 + cc] = -acc;
      }
      __syncthreads();
    }
  }

  for (int idx = tid; idx < NB * NB; idx += THREADS) {
    const int i = idx / NB;
    const int j = idx - i * NB;
    const bool low = (j <= i);
    Db[(long)i * n + j] = low ? S[i * LD + j] : 0.0f;
    const float x = low ? X[i * LD + j] : 0.0f;
    const __half h = __float2half(x);
    invD[inv_off + idx] = x;
    if (emit_half_inverse) {
      invDh[inv_off + idx] = h;
      invDl[inv_off + idx] = __float2half(x - __half2float(h));
    }
  }
}

// ---- Equilibration / epilogue / operand splitting ----

__global__ void equil_diag_kernel(const float* __restrict__ A,
                                  float* __restrict__ sneg,
                                  float* __restrict__ spos,
                                  long n, long total) {
  const long idx = blockIdx.x * (long)blockDim.x + threadIdx.x;
  if (idx >= total) return;
  const long b = idx / n;
  const long i = idx - b * n;
  const float v = A[b * n * n + i * n + i];
  int e = 0;
  if (isfinite(v) && v > 0.0f) {
    frexpf(v, &e);
    e = e >> 1;
  }
  sneg[idx] = exp2f((float)(-e));
  spos[idx] = exp2f((float)(e));
}

__global__ void equil_apply_kernel(const float* __restrict__ A,
                                   float* __restrict__ L,
                                   const float* __restrict__ sneg,
                                   long n, long total) {
  const long idx = blockIdx.x * (long)blockDim.x + threadIdx.x;
  if (idx >= total) return;
  const long nn = n * n;
  const long b = idx / nn;
  const long within = idx - b * nn;
  const long i = within / n;
  const long j = within - i * n;
  L[idx] = A[idx] * sneg[b * n + i] * sneg[b * n + j];
}

__global__ void unscale_zero_upper_kernel(float* __restrict__ L,
                                          const float* __restrict__ spos,
                                          long n, long total) {
  const long idx = blockIdx.x * (long)blockDim.x + threadIdx.x;
  if (idx >= total) return;
  const long nn = n * n;
  const long b = idx / nn;
  const long within = idx - b * nn;
  const long i = within / n;
  const long j = within - i * n;
  L[idx] = (j <= i) ? L[idx] * spos[b * n + i] : 0.0f;
}

__global__ void split_pair_kernel(const float* __restrict__ src,
                                  __half* __restrict__ dst_h,
                                  __half* __restrict__ dst_l,
                                  long s0, long s1, long s2,
                                  long a0, long a1, long a2,
                                  long h0, long h1, long h2,
                                  long l0, long l1, long l2) {
  const long total = s0 * s1 * s2;
  const long idx = blockIdx.x * (long)blockDim.x + threadIdx.x;
  if (idx >= total) return;
  const long b = idx / (s1 * s2);
  const long rem = idx - b * s1 * s2;
  const long i = rem / s2;
  const long j = rem - i * s2;
  const float v = src[b * a0 + i * a1 + j * a2];
  const __half h = __float2half(v);
  dst_h[b * h0 + i * h1 + j * h2] = h;
  dst_l[b * l0 + i * l1 + j * l2] = __float2half(v - __half2float(h));
}

__global__ void copy_float_kernel(const float* __restrict__ src,
                                  float* __restrict__ dst,
                                  long s0, long s1, long s2,
                                  long a0, long a1, long a2,
                                  long d0, long d1, long d2) {
  const long total = s0 * s1 * s2;
  const long idx = blockIdx.x * (long)blockDim.x + threadIdx.x;
  if (idx >= total) return;
  const long b = idx / (s1 * s2);
  const long rem = idx - b * s1 * s2;
  const long i = rem / s2;
  const long j = rem - i * s2;
  dst[b * d0 + i * d1 + j * d2] =
      src[b * a0 + i * a1 + j * a2];
}

constexpr int kPotrfInvShared = (2 * 128 * 129 + 32 * 33) * sizeof(float);

// ---- cuBLASLt fp16-operand / fp32-accumulate matmul, workspace-backed ----

cublasLtMatrixLayout_t make_lt_layout(
    const at::Tensor& t,
    cudaDataType_t dtype) {
  TORCH_CHECK(t.dim() == 3);

  const int batch = static_cast<int>(t.size(0));
  const int64_t rows = t.size(1);
  const int64_t cols = t.size(2);

  cublasLtOrder_t order;
  int64_t ld;

  if (t.stride(2) == 1) {
    order = CUBLASLT_ORDER_ROW;
    ld = t.stride(1);
  } else if (t.stride(1) == 1) {
    order = CUBLASLT_ORDER_COL;
    ld = t.stride(2);
  } else {
    TORCH_CHECK(false, "tensor must be row-major or column-major");
  }

  cublasLtMatrixLayout_t layout = nullptr;
  auto status = cublasLtMatrixLayoutCreate(&layout, dtype, rows, cols, ld);
  TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, "layout create failed: ", status);

  status = cublasLtMatrixLayoutSetAttribute(
      layout, CUBLASLT_MATRIX_LAYOUT_ORDER, &order, sizeof(order));
  TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, "set order failed: ", status);

  status = cublasLtMatrixLayoutSetAttribute(
      layout, CUBLASLT_MATRIX_LAYOUT_BATCH_COUNT, &batch, sizeof(batch));
  TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, "set batch failed: ", status);

  const int64_t batch_stride = t.stride(0);
  status = cublasLtMatrixLayoutSetAttribute(
      layout,
      CUBLASLT_MATRIX_LAYOUT_STRIDED_BATCH_OFFSET,
      &batch_stride,
      sizeof(batch_stride));
  TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, "set batch off failed: ", status);

  return layout;
}

void destroy_lt_layouts(std::initializer_list<cublasLtMatrixLayout_t> layouts) {
  for (auto layout : layouts) {
    if (layout) cublasLtMatrixLayoutDestroy(layout);
  }
}

void lt_gemm(const at::Tensor& C, const at::Tensor& A, const at::Tensor& B,
             float beta, float alpha, void* ws, size_t ws_bytes) {
  cublasLtHandle_t handle = at::cuda::getCurrentCUDABlasLtHandle();

  cublasLtMatmulDesc_t op = nullptr;
  auto status = cublasLtMatmulDescCreate(&op, CUBLAS_COMPUTE_32F, CUDA_R_32F);
  TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, "desc create failed: ", status);

  auto a_layout = make_lt_layout(A, CUDA_R_16F);
  auto b_layout = make_lt_layout(B, CUDA_R_16F);
  auto c_layout = make_lt_layout(C, CUDA_R_32F);

  status = cublasLtMatmul(
      handle, op,
      &alpha,
      A.data_ptr(), a_layout,
      B.data_ptr(), b_layout,
      &beta,
      C.data_ptr<float>(), c_layout,
      C.data_ptr<float>(), c_layout,
      nullptr, ws, ws_bytes, 0);

  destroy_lt_layouts({c_layout, b_layout, a_layout});
  if (op) cublasLtMatmulDescDestroy(op);

  TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, "cublasLtMatmul failed: ", status);
}

void tf32_gemm(const at::Tensor& C, const at::Tensor& A, const at::Tensor& B,
               float beta, float alpha, void* ws, size_t ws_bytes) {
  cublasLtHandle_t handle = at::cuda::getCurrentCUDABlasLtHandle();

  cublasLtMatmulDesc_t op = nullptr;
  auto status = cublasLtMatmulDescCreate(
      &op, CUBLAS_COMPUTE_32F_FAST_TF32, CUDA_R_32F);
  TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, "desc create failed: ", status);

  auto a_layout = make_lt_layout(A, CUDA_R_32F);
  auto b_layout = make_lt_layout(B, CUDA_R_32F);
  auto c_layout = make_lt_layout(C, CUDA_R_32F);

  status = cublasLtMatmul(
      handle, op,
      &alpha,
      A.data_ptr<float>(), a_layout,
      B.data_ptr<float>(), b_layout,
      &beta,
      C.data_ptr<float>(), c_layout,
      C.data_ptr<float>(), c_layout,
      nullptr, ws, ws_bytes, 0);

  destroy_lt_layouts({c_layout, b_layout, a_layout});
  if (op) cublasLtMatmulDescDestroy(op);

  TORCH_CHECK(status == CUBLAS_STATUS_SUCCESS, "cublasLtMatmul failed: ", status);
}

void comp_mm(const at::Tensor& C, const at::Tensor& Ah, const at::Tensor& Al,
             const at::Tensor& Bh, const at::Tensor& Bl,
             float beta, float alpha, void* ws, size_t ws_bytes) {
  // C = beta*C + alpha*(A @ B) with fp16 hi/lo compensated operands and
  // fp32 accumulation; the lo*lo term is below fp32 roundoff and dropped.
  lt_gemm(C, Ah, Bh, beta, alpha, ws, ws_bytes);
  lt_gemm(C, Ah, Bl, 1.0f, alpha, ws, ws_bytes);
  lt_gemm(C, Al, Bh, 1.0f, alpha, ws, ws_bytes);
}

// ---- device-kernel launch helpers ----

void launch_warp_32(
    const at::Tensor& input,
    at::Tensor& output) {
  constexpr int kWarpsPerBlock = 8;
  constexpr int kThreads = kWarpsPerBlock * 32;
  const int batch = static_cast<int>(input.size(0));
  const int blocks = (batch + kWarpsPerBlock - 1) / kWarpsPerBlock;
  warp_cholesky_32_kernel<<<blocks, kThreads>>>(
      input.data_ptr<float>(), output.data_ptr<float>(), batch);
}

void launch_warp_64(
    const at::Tensor& input,
    at::Tensor& output) {
  constexpr int kWarpsPerBlock = 8;
  constexpr int kThreads = kWarpsPerBlock * 32;
  const int batch = static_cast<int>(input.size(0));
  const int blocks = (batch + kWarpsPerBlock - 1) / kWarpsPerBlock;
  warp_cholesky_64_kernel<<<blocks, kThreads>>>(
      input.data_ptr<float>(), output.data_ptr<float>(), batch);
}

template <int N>
void launch_register(
    const at::Tensor& input,
    at::Tensor& output) {
  constexpr int threads = (N / 8) * 32;
  constexpr int shared_bytes = (N * (N + 1) + N) * sizeof(float);
  if constexpr (shared_bytes > 48 * 1024) {
    C10_CUDA_CHECK(cudaFuncSetAttribute(
        register_cholesky_kernel<N>,
        cudaFuncAttributeMaxDynamicSharedMemorySize,
        shared_bytes));
  }
  register_cholesky_kernel<N><<<input.size(0), threads, shared_bytes>>>(
      input.data_ptr<float>(), output.data_ptr<float>());
}

template <int N>
void launch_wavefront(
    const at::Tensor& input,
    at::Tensor& output) {
  constexpr int threads = (N / 8) * 32;
  constexpr int shared_bytes = N * N * sizeof(float) + N * 8;
  if constexpr (shared_bytes > 48 * 1024) {
    C10_CUDA_CHECK(cudaFuncSetAttribute(
        wavefront_cholesky_kernel<N>,
        cudaFuncAttributeMaxDynamicSharedMemorySize,
        shared_bytes));
  }
  wavefront_cholesky_kernel<N><<<input.size(0), threads, shared_bytes>>>(
      input.data_ptr<float>(), output.data_ptr<float>());
}

void launch_cluster_256(
    const at::Tensor& input,
    at::Tensor& output) {
  constexpr int shared_bytes =
      256 * 128 * sizeof(float) + 256 * 8;
  C10_CUDA_CHECK(cudaFuncSetAttribute(
      cluster_cholesky_256_kernel,
      cudaFuncAttributeMaxDynamicSharedMemorySize,
      shared_bytes));
  cluster_cholesky_256_kernel<<<
      input.size(0) * 2, 512, shared_bytes>>>(
      input.data_ptr<float>(), output.data_ptr<float>());
}

void launch_equil_diag(const at::Tensor& A, const at::Tensor& sneg,
                       const at::Tensor& spos) {
  const long n = A.size(-1);
  const long total = A.size(0) * n;
  const long blocks = (total + 255) / 256;
  equil_diag_kernel<<<(unsigned int)blocks, 256>>>(
      A.data_ptr<float>(), sneg.data_ptr<float>(), spos.data_ptr<float>(),
      n, total);
  C10_CUDA_KERNEL_LAUNCH_CHECK();
}

void launch_equil_apply(const at::Tensor& A, const at::Tensor& L,
                        const at::Tensor& sneg) {
  const long n = A.size(-1);
  const long total = A.numel();
  const long blocks = (total + 255) / 256;
  equil_apply_kernel<<<(unsigned int)blocks, 256>>>(
      A.data_ptr<float>(), L.data_ptr<float>(), sneg.data_ptr<float>(),
      n, total);
  C10_CUDA_KERNEL_LAUNCH_CHECK();
}

void launch_unscale_zero_upper(const at::Tensor& L, const at::Tensor& spos) {
  const long n = L.size(-1);
  const long total = L.numel();
  const long blocks = (total + 255) / 256;
  unscale_zero_upper_kernel<<<(unsigned int)blocks, 256>>>(
      L.data_ptr<float>(), spos.data_ptr<float>(), n, total);
  C10_CUDA_KERNEL_LAUNCH_CHECK();
}

void launch_potrf_inv(const at::Tensor& L, const at::Tensor& invD,
                      const at::Tensor& invDh, const at::Tensor& invDl,
                      int k, bool emit_half_inverse) {
  const int batch = static_cast<int>(L.size(0));
  const long n = L.size(1);
  static bool armed = false;
  if (!armed) {
    C10_CUDA_CHECK(cudaFuncSetAttribute(
        potrf_inv_kernel,
        cudaFuncAttributeMaxDynamicSharedMemorySize,
        kPotrfInvShared));
    armed = true;
  }
  potrf_inv_kernel<<<batch, 256, kPotrfInvShared>>>(
      L.data_ptr<float>(), invD.data_ptr<float>(),
      reinterpret_cast<__half*>(invDh.data_ptr<at::Half>()),
      reinterpret_cast<__half*>(invDl.data_ptr<at::Half>()),
      n, n * n, invD.stride(0), k, emit_half_inverse);
  C10_CUDA_KERNEL_LAUNCH_CHECK();
}

void launch_split_pair(const at::Tensor& src, const at::Tensor& dst_h,
                       const at::Tensor& dst_l) {
  const long total = src.numel();
  const long blocks = (total + 255) / 256;
  split_pair_kernel<<<(unsigned int)blocks, 256>>>(
      src.data_ptr<float>(),
      reinterpret_cast<__half*>(dst_h.data_ptr<at::Half>()),
      reinterpret_cast<__half*>(dst_l.data_ptr<at::Half>()),
      src.size(0), src.size(1), src.size(2),
      src.stride(0), src.stride(1), src.stride(2),
      dst_h.stride(0), dst_h.stride(1), dst_h.stride(2),
      dst_l.stride(0), dst_l.stride(1), dst_l.stride(2));
  C10_CUDA_KERNEL_LAUNCH_CHECK();
}

void launch_copy_float(const at::Tensor& src, const at::Tensor& dst) {
  const long total = src.numel();
  const long blocks = (total + 255) / 256;
  copy_float_kernel<<<(unsigned int)blocks, 256>>>(
      src.data_ptr<float>(), dst.data_ptr<float>(),
      src.size(0), src.size(1), src.size(2),
      src.stride(0), src.stride(1), src.stride(2),
      dst.stride(0), dst.stride(1), dst.stride(2));
  C10_CUDA_KERNEL_LAUNCH_CHECK();
}

// ---- full blocked/recursive driver, host-sequenced in C++ ----

constexpr int64_t kPB = 128;
constexpr int64_t kSplitCut = 1024;

inline at::Tensor rgn(const at::Tensor& t, int64_t r0, int64_t nr,
                      int64_t c0, int64_t nc) {
  return t.narrow(1, r0, nr).narrow(2, c0, nc);
}

inline at::Tensor tr(const at::Tensor& t) {
  return t.transpose(1, 2);
}

struct Ctx {
  at::Tensor L, Lh, Ll, invD, invDh, invDl, Ph, Pl, PF;
  void* ws;
  size_t ws_bytes;
  bool fast_tf32;
};

void leaf_blocked(Ctx& cx, int64_t base, int64_t nreg) {
  for (int64_t kb = 0; kb < nreg / kPB; ++kb) {
    const int64_t k = base + kb * kPB;
    launch_potrf_inv(
        cx.L, cx.invD, cx.invDh, cx.invDl, (int)k, !cx.fast_tf32);
    const int64_t m = base + nreg - (k + kPB);
    if (m > 0) {
      const int64_t slot = k / kPB;
      auto P = rgn(cx.L, k + kPB, m, k, kPB);
      if (cx.fast_tf32) {
        auto F = cx.PF.narrow(1, 0, m);
        launch_copy_float(P, F);
        auto iv = tr(cx.invD.select(1, slot));
        tf32_gemm(P, F, iv, 0.0f, 1.0f, cx.ws, cx.ws_bytes);
        auto T = rgn(cx.L, k + kPB, m, k + kPB, m);
        tf32_gemm(T, P, tr(P), 1.0f, -1.0f, cx.ws, cx.ws_bytes);
      } else {
        auto Ph = cx.Ph.narrow(1, 0, m);
        auto Pl = cx.Pl.narrow(1, 0, m);
        launch_split_pair(P, Ph, Pl);
        auto ih = tr(cx.invDh.select(1, slot));
        auto il = tr(cx.invDl.select(1, slot));
        // Panel solve P <- P @ inv(D)^T lands in place: operands were saved
        // by the split, so overwriting the region is safe.
        comp_mm(P, Ph, Pl, ih, il, 0.0f, 1.0f,
                cx.ws, cx.ws_bytes);
        // The panel is final: publish its fp16 hi/lo shadow once.
        auto Sh = rgn(cx.Lh, k + kPB, m, k, kPB);
        auto Sl = rgn(cx.Ll, k + kPB, m, k, kPB);
        launch_split_pair(P, Sh, Sl);
        auto T = rgn(cx.L, k + kPB, m, k + kPB, m);
        comp_mm(T, Sh, Sl, tr(Sh), tr(Sl), 1.0f, -1.0f,
                cx.ws, cx.ws_bytes);
      }
    }
  }
}

void rec_trsm(Ctx& cx, int64_t rbase, int64_t m, int64_t cbase, int64_t w) {
  // Solve X @ L11^T = X in place for X = L[rbase:rbase+m, cbase:cbase+w].
  if (w == kPB) {
    const int64_t slot = cbase / kPB;
    auto X = rgn(cx.L, rbase, m, cbase, kPB);
    if (cx.fast_tf32) {
      auto F = cx.PF.narrow(1, 0, m);
      launch_copy_float(X, F);
      auto iv = tr(cx.invD.select(1, slot));
      tf32_gemm(X, F, iv, 0.0f, 1.0f, cx.ws, cx.ws_bytes);
    } else {
      auto Ph = cx.Ph.narrow(1, 0, m);
      auto Pl = cx.Pl.narrow(1, 0, m);
      launch_split_pair(X, Ph, Pl);
      auto ih = tr(cx.invDh.select(1, slot));
      auto il = tr(cx.invDl.select(1, slot));
      comp_mm(X, Ph, Pl, ih, il, 0.0f, 1.0f,
              cx.ws, cx.ws_bytes);
      auto Sh = rgn(cx.Lh, rbase, m, cbase, kPB);
      auto Sl = rgn(cx.Ll, rbase, m, cbase, kPB);
      launch_split_pair(X, Sh, Sl);
    }
    return;
  }
  const int64_t hw = (w / 2) / kPB * kPB;
  rec_trsm(cx, rbase, m, cbase, hw);
  auto Xb = rgn(cx.L, rbase, m, cbase + hw, w - hw);
  if (cx.fast_tf32) {
    auto Xa = rgn(cx.L, rbase, m, cbase, hw);
    auto Lb = rgn(cx.L, cbase + hw, w - hw, cbase, hw);
    tf32_gemm(Xb, Xa, tr(Lb), 1.0f, -1.0f, cx.ws, cx.ws_bytes);
  } else {
    auto Xah = rgn(cx.Lh, rbase, m, cbase, hw);
    auto Xal = rgn(cx.Ll, rbase, m, cbase, hw);
    auto Lbh = rgn(cx.Lh, cbase + hw, w - hw, cbase, hw);
    auto Lbl = rgn(cx.Ll, cbase + hw, w - hw, cbase, hw);
    comp_mm(Xb, Xah, Xal, tr(Lbh), tr(Lbl), 1.0f, -1.0f,
            cx.ws, cx.ws_bytes);
  }
  rec_trsm(cx, rbase, m, cbase + hw, w - hw);
}

void rec_syrk(Ctx& cx, int64_t base, int64_t t, int64_t pcbase, int64_t pw) {
  // T -= P @ P^T on the lower triangle; P is finalized factor, read from
  // the fp16 shadows. Cutoff blocks pay the full-square waste only.
  if (t <= kSplitCut) {
    auto T = rgn(cx.L, base, t, base, t);
    if (cx.fast_tf32) {
      auto P = rgn(cx.L, base, t, pcbase, pw);
      tf32_gemm(T, P, tr(P), 1.0f, -1.0f, cx.ws, cx.ws_bytes);
    } else {
      auto Ph = rgn(cx.Lh, base, t, pcbase, pw);
      auto Pl = rgn(cx.Ll, base, t, pcbase, pw);
      comp_mm(T, Ph, Pl, tr(Ph), tr(Pl), 1.0f, -1.0f,
              cx.ws, cx.ws_bytes);
    }
    return;
  }
  const int64_t ht = (t / 2) / kPB * kPB;
  rec_syrk(cx, base, ht, pcbase, pw);
  auto T21 = rgn(cx.L, base + ht, t - ht, base, ht);
  if (cx.fast_tf32) {
    auto P2 = rgn(cx.L, base + ht, t - ht, pcbase, pw);
    auto P1 = rgn(cx.L, base, ht, pcbase, pw);
    tf32_gemm(T21, P2, tr(P1), 1.0f, -1.0f, cx.ws, cx.ws_bytes);
  } else {
    auto P2h = rgn(cx.Lh, base + ht, t - ht, pcbase, pw);
    auto P2l = rgn(cx.Ll, base + ht, t - ht, pcbase, pw);
    auto P1h = rgn(cx.Lh, base, ht, pcbase, pw);
    auto P1l = rgn(cx.Ll, base, ht, pcbase, pw);
    comp_mm(T21, P2h, P2l, tr(P1h), tr(P1l), 1.0f, -1.0f,
            cx.ws, cx.ws_bytes);
  }
  rec_syrk(cx, base + ht, t - ht, pcbase, pw);
}

void rec_chol(Ctx& cx, int64_t base, int64_t n, int64_t leaf) {
  if (n <= leaf) {
    leaf_blocked(cx, base, n);
    return;
  }
  const int64_t h = (n / 2) / kPB * kPB;
  rec_chol(cx, base, h, leaf);
  rec_trsm(cx, base + h, n - h, base, h);
  rec_syrk(cx, base + h, n - h, base, h);
  rec_chol(cx, base + h, n - h, leaf);
}

}  // namespace

torch::Tensor cholesky_small_cuda(const torch::Tensor& input) {
  TORCH_CHECK(input.is_cuda(), "input must be CUDA");
  TORCH_CHECK(input.scalar_type() == at::kFloat, "input must be float32");
  TORCH_CHECK(input.is_contiguous(), "input must be contiguous");
  TORCH_CHECK(input.dim() == 3, "input must have shape [batch,n,n]");
  TORCH_CHECK(input.size(1) == input.size(2), "matrices must be square");

  c10::cuda::CUDAGuard guard(input.device());
  auto output = at::empty_like(input);
  const int n = static_cast<int>(input.size(1));

  switch (n) {
    case 32:
      launch_warp_32(input, output);
      break;
    case 64:
      launch_warp_64(input, output);
      break;
    case 128:
      launch_wavefront<128>(input, output);
      break;
    case 256:
      launch_cluster_256(input, output);
      break;
    default:
      TORCH_CHECK(false, "unsupported small Cholesky size: ", n);
  }
  C10_CUDA_KERNEL_LAUNCH_CHECK();
  return output;
}

void factor_owned(at::Tensor A, at::Tensor L, at::Tensor Lh, at::Tensor Ll,
                  at::Tensor sneg, at::Tensor spos,
                  at::Tensor invD, at::Tensor invDh, at::Tensor invDl,
                  at::Tensor Ph, at::Tensor Pl, at::Tensor PF,
                  at::Tensor ws, int64_t leaf) {
  TORCH_CHECK(A.is_contiguous() && L.is_contiguous());
  c10::cuda::CUDAGuard guard(A.device());
  const int64_t batch = A.size(0);
  const int64_t n = A.size(1);

  launch_equil_diag(A, sneg, spos);
  launch_equil_apply(A, L, sneg);

  Ctx cx;
  cx.invD = invD; cx.invDh = invDh; cx.invDl = invDl;
  cx.Ph = Ph; cx.Pl = Pl; cx.PF = PF;
  cx.ws = ws.data_ptr();
  cx.ws_bytes = (size_t)ws.numel();
  cx.fast_tf32 = n >= 16384;

  if (n <= leaf) {
    cx.L = L; cx.Lh = Lh; cx.Ll = Ll;
    leaf_blocked(cx, 0, n);
  } else {
    for (int64_t bi = 0; bi < batch; ++bi) {
      cx.L = L.narrow(0, bi, 1);
      cx.Lh = Lh.narrow(0, bi, 1);
      cx.Ll = Ll.narrow(0, bi, 1);
      rec_chol(cx, 0, n, leaf);
    }
  }

  launch_unscale_zero_upper(L, spos);
}
"""


def _cublaslt_build_flags():
    import glob
    import os

    roots = []
    try:
        import nvidia

        roots += [p for p in list(getattr(nvidia, "__path__", []) or []) if p]
    except ImportError:
        pass
    roots.append(os.path.join(os.path.dirname(os.path.dirname(torch.__file__)), "nvidia"))

    ldflags = []
    includes = []
    for root in roots:
        for lib_dir in sorted(glob.glob(os.path.join(root, "*", "lib"))):
            sos = glob.glob(os.path.join(lib_dir, "libcublasLt.so.*"))
            if sos:
                name = os.path.basename(sorted(sos, key=len)[0])
                ldflags += [f"-L{lib_dir}", f"-Wl,-rpath,{lib_dir}", f"-l:{name}"]
                # The toolkit's own headers match its nvcc; only fall back to
                # the pip-bundled headers when the toolkit has no cublasLt.h.
                inc = os.path.join(os.path.dirname(lib_dir), "include")
                cuda_home = os.environ.get("CUDA_HOME", "/usr/local/cuda")
                if not os.path.exists(os.path.join(cuda_home, "include", "cublasLt.h")):
                    if os.path.isdir(inc):
                        includes.append(inc)
                break
        if ldflags:
            break
    if not ldflags:
        ldflags = ["-L/usr/local/cuda/lib64", "-lcublasLt"]
    return ldflags, includes


_LDFLAGS, _INCLUDES = _cublaslt_build_flags()

_module = load_inline(
    name="cholesky_gau_v27_large_tf32",
    cpp_sources=CPP_SRC,
    cuda_sources=CUDA_SRC,
    functions=["cholesky_small_cuda", "factor_owned"],
    extra_cuda_cflags=["-O3"],
    extra_ldflags=_LDFLAGS,
    extra_include_paths=_INCLUDES,
    with_cuda=True,
    verbose=False,
)


_PB = 128        # leaf panel width; the diagonal pivot chain never leaves one CTA
_LEAF = 4096     # recursion cutoff for the single-matrix path
_WS_BYTES = 32 * 1024 * 1024   # cuBLASLt workspace, allocated fresh per call


def _run_owned(data):
    # Fully eager, allocation-per-call path; the whole panel loop, recursion,
    # and cuBLASLt sequencing run inside one C++ driver call. The input is
    # only ever read.
    batch, n, _ = data.shape
    device = data.device
    nb = n // _PB
    batched = n <= _LEAF
    inv_b = batch if batched else 1
    split_b = batch if batched else 1

    fast_tf32 = n >= 16384
    L = torch.empty_like(data)
    half_matrix_shape = (1, 1, 1) if fast_tf32 else (batch, n, n)
    Lh = torch.empty(half_matrix_shape, device=device, dtype=torch.float16)
    Ll = torch.empty(half_matrix_shape, device=device, dtype=torch.float16)
    sneg = torch.empty((batch, n), device=device, dtype=torch.float32)
    spos = torch.empty((batch, n), device=device, dtype=torch.float32)
    invD = torch.empty((inv_b, nb, _PB, _PB), device=device, dtype=torch.float32)
    half_inverse_shape = (
        (1, 1, 1, 1) if fast_tf32 else (inv_b, nb, _PB, _PB)
    )
    invDh = torch.empty(half_inverse_shape, device=device, dtype=torch.float16)
    invDl = torch.empty(half_inverse_shape, device=device, dtype=torch.float16)
    half_panel_shape = (
        (1, 1, 1) if fast_tf32 else (split_b, n - _PB, _PB)
    )
    Ph = torch.empty(half_panel_shape, device=device, dtype=torch.float16)
    Pl = torch.empty(half_panel_shape, device=device, dtype=torch.float16)
    float_panel_shape = (
        (split_b, n - _PB, _PB) if fast_tf32 else (1, 1, 1)
    )
    PF = torch.empty(float_panel_shape, device=device, dtype=torch.float32)
    ws = torch.empty(_WS_BYTES, device=device, dtype=torch.uint8)

    _module.factor_owned(
        data, L, Lh, Ll, sneg, spos,
        invD, invDh, invDl, Ph, Pl, PF, ws, _LEAF,
    )
    return L


def _per_matrix_reference(data):
    # cuSOLVER dispatches low-batch batched calls poorly; a per-matrix loop
    # was measured much faster (3.22 ms vs 11.1 ms at n=4096 batch=2).
    output = torch.empty_like(data)
    info = torch.empty((1,), dtype=torch.int32, device=data.device)
    for index in range(data.shape[0]):
        torch.linalg.cholesky_ex(
            data[index:index + 1],
            check_errors=False,
            out=(output[index:index + 1], info),
        )
    return output


def _reference(data):
    return torch.linalg.cholesky_ex(data, check_errors=False).L


def custom_kernel(data: input_t) -> output_t:
    n = data.shape[-1]

    if data.is_contiguous() and n in (32, 64, 128):
        return _module.cholesky_small_cuda(data)

    if data.is_contiguous() and n == 256 and torch.cuda.get_device_capability(data.device)[0] >= 10:
        return _module.cholesky_small_cuda(data)

    owned = (
        isinstance(data, torch.Tensor)
        and data.is_cuda
        and data.dtype == torch.float32
        and data.dim() == 3
        and data.shape[1] == data.shape[2]
        and data.is_contiguous()
    )

    # Per-row routing, each entry the measured winner on B200 among the owned
    # compensated-fp16 path, one batched reference call, and a per-matrix
    # reference loop (benchmark_ref.txt / benchmark_v7_outsplit.txt / v10).
    if owned:
        b = data.shape[0]
        if n == 512:
            return _run_owned(data) if b <= 64 else _reference(data)
        if n == 1024:
            return _run_owned(data)
        if n == 2048:
            return _per_matrix_reference(data) if b <= 2 else _run_owned(data)
        if n == 4096:
            if b == 1:
                return _reference(data)
            if b <= 4:
                return _per_matrix_reference(data)
            return _run_owned(data)
        if n == 8192:
            return _reference(data) if b == 1 else _per_matrix_reference(data)
        if n in (16384, 32768):
            return _run_owned(data)

    # Correctness floor for anything off the benchmark grid.
    return _reference(data)

scrolls · 1673 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