Skip to content
KernelIndex
Search⌘K

submission 889544

debadree25 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-889544?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.35ms
#159 of 337
2026-07-20

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:07edb4089ed2db81a7dda8397899d34130d1c7d97b7a7612c2a99a4f4c82a5e0
license declaredunknown
license concludedunknown
authorsdebadree25
imported2026-08-26

Techniques

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

shared-memorycholesky64_smem_persist(const float* __restrict__ A,
vector-width = float4float4 v = *reinterpret_cast<const float4*>(a + lane * 32 + t * 4);

Kernel source

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

"""Hybrid: n32/64/128 + n256 dense-panel + medium high-batch + large3 TF32.

- n=32/64/128: champ CUDA
- medium high-batch blocked batched (512@≥64, 1024@≥16, 2048@≥8)
- large TF32 selective (n≥16384 or n≥4096&batch≥2)
- else torch
"""

from __future__ import annotations

import torch
from torch.utils.cpp_extension import load_inline

from task import input_t, output_t

SMALL_CUDA_SRC = r"""

#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cmath>
#include <algorithm>
#include <stdexcept>
#include <string>

#define CHECK_CUDA(expr) do { \
  cudaError_t _e = (expr); \
  if (_e != cudaSuccess) { \
    throw std::runtime_error(std::string("CUDA error: ") + cudaGetErrorString(_e)); \
  } \
} while (0)

// -------------------- n=32: one warp / matrix, rows in registers --------------------
template <int MPB>
__global__ void __launch_bounds__(32 * MPB, 4)
cholesky32_warp_persist(const float* __restrict__ A,
                        float* __restrict__ L,
                        int batch) {
  const int lane = threadIdx.x;
  const int slot = threadIdx.y;
  const int stride = gridDim.x * MPB;

  for (int b = blockIdx.x * MPB + slot; b < batch; b += stride) {
    const float* a = A + (size_t)b * 1024u;
    float* out = L + (size_t)b * 1024u;

    float row[32];
#pragma unroll
    for (int t = 0; t < 8; ++t) {
      float4 v = *reinterpret_cast<const float4*>(a + lane * 32 + t * 4);
      row[t * 4 + 0] = v.x;
      row[t * 4 + 1] = v.y;
      row[t * 4 + 2] = v.z;
      row[t * 4 + 3] = v.w;
    }

#pragma unroll
    for (int k = 0; k < 32; ++k) {
      float sum = 0.f;
#pragma unroll
      for (int j = 0; j < 32; ++j) {
        if (j < k) {
          float rkj = __shfl_sync(0xffffffffu, row[j], k);
          sum = fmaf(row[j], rkj, sum);
        }
      }
      float aik = row[k] - sum;
      if (lane == k) {
        row[k] = sqrtf(fmaxf(aik, 0.f));
      }
      float diag = __shfl_sync(0xffffffffu, row[k], k);
      float inv = 1.f / diag;
      if (lane > k) {
        row[k] = aik * inv;
      } else if (lane < k) {
        row[k] = 0.f;
      }
    }

#pragma unroll
    for (int t = 0; t < 8; ++t) {
      float4 v;
      int j0 = t * 4;
      v.x = (lane >= (j0 + 0)) ? row[j0 + 0] : 0.f;
      v.y = (lane >= (j0 + 1)) ? row[j0 + 1] : 0.f;
      v.z = (lane >= (j0 + 2)) ? row[j0 + 2] : 0.f;
      v.w = (lane >= (j0 + 3)) ? row[j0 + 3] : 0.f;
      *reinterpret_cast<float4*>(out + lane * 32 + j0) = v;
    }
  }
}

torch::Tensor cholesky_n32(torch::Tensor A) {
  TORCH_CHECK(A.is_cuda(), "A must be CUDA");
  TORCH_CHECK(A.scalar_type() == at::kFloat, "A must be float32");
  TORCH_CHECK(A.dim() == 3 && A.size(1) == 32 && A.size(2) == 32, "expected (B,32,32)");

  auto Ac = A.contiguous();
  auto L = torch::empty_like(Ac);
  const int batch = static_cast<int>(Ac.size(0));
  if (batch == 0) return L;

  constexpr int mpb = 8;
  dim3 block(32, mpb);
  int need = (batch + mpb - 1) / mpb;
  int grid = need < 1024 ? need : 1024;
  if (grid < 1) grid = 1;

  cholesky32_warp_persist<mpb><<<grid, block>>>(
      Ac.data_ptr<float>(), L.data_ptr<float>(), batch);
  CHECK_CUDA(cudaGetLastError());
  return L;
}

// -------------------- n=64: left-looking, SMEM, full unroll, multi-matrix --------------------
// Each of 64 threads owns one matrix row in shared memory (bank-padded stride 65).
// Outer k and inner j fully unrolled so j<k is compile-time (n32-style quality).
// Dynamic SMEM: MPB * 64 * 65 floats (static 48KB cap is too small for MPB>=3).
template <int MPB>
__global__ void __launch_bounds__(64 * MPB, 2)
cholesky64_smem_persist(const float* __restrict__ A,
                        float* __restrict__ L,
                        int batch) {
  const int tid = threadIdx.x;   // 0..63 row owner
  const int slot = threadIdx.y;  // matrix slot in CTA
  const int stride = gridDim.x * MPB;

  extern __shared__ float smem[];
  float* tile = smem + slot * (64 * 65);

  for (int b = blockIdx.x * MPB + slot; b < batch; b += stride) {
    const float* a = A + (size_t)b * 4096u;
    float* out = L + (size_t)b * 4096u;

    // Cooperative float4 load of full 64x64 into padded rows.
#pragma unroll
    for (int t = 0; t < 16; ++t) {
      float4 v = *reinterpret_cast<const float4*>(a + tid * 64 + t * 4);
      int j0 = t * 4;
      tile[tid * 65 + j0 + 0] = v.x;
      tile[tid * 65 + j0 + 1] = v.y;
      tile[tid * 65 + j0 + 2] = v.z;
      tile[tid * 65 + j0 + 3] = v.w;
    }
    __syncthreads();

#pragma unroll
    for (int k = 0; k < 64; ++k) {
      float sum = 0.f;
#pragma unroll
      for (int j = 0; j < 64; ++j) {
        if (j < k) {
          sum = fmaf(tile[tid * 65 + j], tile[k * 65 + j], sum);
        }
      }
      float aik = tile[tid * 65 + k] - sum;

      if (tid == k) {
        tile[k * 65 + k] = sqrtf(fmaxf(aik, 0.f));
      }
      __syncthreads();

      float inv = 1.f / tile[k * 65 + k];
      if (tid > k) {
        tile[tid * 65 + k] = aik * inv;
      } else if (tid < k) {
        tile[tid * 65 + k] = 0.f;
      }
      __syncthreads();
    }

#pragma unroll
    for (int t = 0; t < 16; ++t) {
      float4 v;
      int j0 = t * 4;
      v.x = (tid >= (j0 + 0)) ? tile[tid * 65 + j0 + 0] : 0.f;
      v.y = (tid >= (j0 + 1)) ? tile[tid * 65 + j0 + 1] : 0.f;
      v.z = (tid >= (j0 + 2)) ? tile[tid * 65 + j0 + 2] : 0.f;
      v.w = (tid >= (j0 + 3)) ? tile[tid * 65 + j0 + 3] : 0.f;
      *reinterpret_cast<float4*>(out + tid * 64 + j0) = v;
    }
    if (b + stride < batch) {
      __syncthreads();
    }
  }
}

// Register-resident row variant: lower SMEM traffic (only diag broadcast).
// Higher register pressure; occupancy kept low via launch_bounds.
template <int MPB>
__global__ void __launch_bounds__(64 * MPB, 1)
cholesky64_reg_persist(const float* __restrict__ A,
                       float* __restrict__ L,
                       int batch) {
  const int tid = threadIdx.x;
  const int slot = threadIdx.y;
  const int stride = gridDim.x * MPB;

  __shared__ float sk[MPB][64];  // row-k lower + diag publish

  for (int b = blockIdx.x * MPB + slot; b < batch; b += stride) {
    const float* a = A + (size_t)b * 4096u;
    float* out = L + (size_t)b * 4096u;

    float row[64];
#pragma unroll
    for (int t = 0; t < 16; ++t) {
      float4 v = *reinterpret_cast<const float4*>(a + tid * 64 + t * 4);
      row[t * 4 + 0] = v.x;
      row[t * 4 + 1] = v.y;
      row[t * 4 + 2] = v.z;
      row[t * 4 + 3] = v.w;
    }

#pragma unroll
    for (int k = 0; k < 64; ++k) {
      // Publish L[k, 0:k] (and later diag at sk[k]).
      if (tid == k) {
#pragma unroll
        for (int j = 0; j < 64; ++j) {
          if (j < k) sk[slot][j] = row[j];
        }
      }
      __syncthreads();

      float sum = 0.f;
#pragma unroll
      for (int j = 0; j < 64; ++j) {
        if (j < k) {
          sum = fmaf(row[j], sk[slot][j], sum);
        }
      }
      float aik = row[k] - sum;

      if (tid == k) {
        float d = sqrtf(fmaxf(aik, 0.f));
        row[k] = d;
        sk[slot][k] = d;
      }
      __syncthreads();

      float inv = 1.f / sk[slot][k];
      if (tid > k) {
        row[k] = aik * inv;
      } else if (tid < k) {
        row[k] = 0.f;
      }
      __syncthreads();
    }

#pragma unroll
    for (int t = 0; t < 16; ++t) {
      float4 v;
      int j0 = t * 4;
      v.x = (tid >= (j0 + 0)) ? row[j0 + 0] : 0.f;
      v.y = (tid >= (j0 + 1)) ? row[j0 + 1] : 0.f;
      v.z = (tid >= (j0 + 2)) ? row[j0 + 2] : 0.f;
      v.w = (tid >= (j0 + 3)) ? row[j0 + 3] : 0.f;
      *reinterpret_cast<float4*>(out + tid * 64 + j0) = v;
    }
  }
}

// Right-looking register row + SMEM column broadcast, full unroll.
template <int MPB>
__global__ void __launch_bounds__(64 * MPB, 1)
cholesky64_right_persist(const float* __restrict__ A,
                         float* __restrict__ L,
                         int batch) {
  const int tid = threadIdx.x;
  const int slot = threadIdx.y;
  const int stride = gridDim.x * MPB;

  __shared__ float col[MPB][64];
  __shared__ float diag_s[MPB];

  for (int b = blockIdx.x * MPB + slot; b < batch; b += stride) {
    const float* a = A + (size_t)b * 4096u;
    float* out = L + (size_t)b * 4096u;

    float row[64];
#pragma unroll
    for (int t = 0; t < 16; ++t) {
      float4 v = *reinterpret_cast<const float4*>(a + tid * 64 + t * 4);
      row[t * 4 + 0] = v.x;
      row[t * 4 + 1] = v.y;
      row[t * 4 + 2] = v.z;
      row[t * 4 + 3] = v.w;
    }

#pragma unroll
    for (int k = 0; k < 64; ++k) {
      if (tid == k) {
        float d = sqrtf(fmaxf(row[k], 0.f));
        row[k] = d;
        diag_s[slot] = d;
      }
      __syncthreads();

      float inv = 1.f / diag_s[slot];
      float lik;
      if (tid > k) {
        lik = row[k] * inv;
        row[k] = lik;
      } else if (tid == k) {
        lik = diag_s[slot];
      } else {
        lik = 0.f;
        row[k] = 0.f;
      }
      col[slot][tid] = lik;
      __syncthreads();

      float lik_i = col[slot][tid];
#pragma unroll
      for (int j = 0; j < 64; ++j) {
        if (j > k && tid >= j) {
          row[j] = fmaf(-lik_i, col[slot][j], row[j]);
        }
      }
      __syncthreads();
    }

#pragma unroll
    for (int t = 0; t < 16; ++t) {
      float4 v;
      int j0 = t * 4;
      v.x = (tid >= (j0 + 0)) ? row[j0 + 0] : 0.f;
      v.y = (tid >= (j0 + 1)) ? row[j0 + 1] : 0.f;
      v.z = (tid >= (j0 + 2)) ? row[j0 + 2] : 0.f;
      v.w = (tid >= (j0 + 3)) ? row[j0 + 3] : 0.f;
      *reinterpret_cast<float4*>(out + tid * 64 + j0) = v;
    }
  }
}

// Variant selector: 0=smem left, 1=reg left, 2=right
#ifndef CHOLESKY64_VARIANT
#define CHOLESKY64_VARIANT 0
#endif

torch::Tensor cholesky_n64(torch::Tensor A) {
  TORCH_CHECK(A.is_cuda(), "A must be CUDA");
  TORCH_CHECK(A.scalar_type() == at::kFloat, "A must be float32");
  TORCH_CHECK(A.dim() == 3 && A.size(1) == 64 && A.size(2) == 64, "expected (B,64,64)");

  auto Ac = A.contiguous();
  auto L = torch::empty_like(Ac);
  const int batch = static_cast<int>(Ac.size(0));
  if (batch == 0) return L;

  constexpr int mpb = 2;  // Modal sweep: mpb2 ~37us > mpb4 ~43 > mpb8 ~50 @1024x64
  dim3 block(64, mpb);
  int need = (batch + mpb - 1) / mpb;
  int grid = need < 2048 ? need : 2048;
  if (grid < 1) grid = 1;

#if CHOLESKY64_VARIANT == 1
  cholesky64_reg_persist<mpb><<<grid, block>>>(
      Ac.data_ptr<float>(), L.data_ptr<float>(), batch);
#elif CHOLESKY64_VARIANT == 2
  cholesky64_right_persist<mpb><<<grid, block>>>(
      Ac.data_ptr<float>(), L.data_ptr<float>(), batch);
#else
  {
    // 4 * 64 * 65 * 4 = 66560 bytes dynamic SMEM (over static 48KB cap).
    const int shmem = mpb * 64 * 65 * (int)sizeof(float);
    CHECK_CUDA(cudaFuncSetAttribute(
        cholesky64_smem_persist<mpb>,
        cudaFuncAttributeMaxDynamicSharedMemorySize,
        shmem));
    cholesky64_smem_persist<mpb><<<grid, block, shmem>>>(
        Ac.data_ptr<float>(), L.data_ptr<float>(), batch);
  }
#endif
  CHECK_CUDA(cudaGetLastError());
  return L;
}


"""

SMALL_CPP_SRC = r"""
torch::Tensor cholesky_n32(torch::Tensor A);
torch::Tensor cholesky_n64(torch::Tensor A);

"""

_mod_small = load_inline(
    name="chol_s3264_fast_v1",  # same name as champ for compile cache hit
    cpp_sources=[SMALL_CPP_SRC],
    cuda_sources=[SMALL_CUDA_SRC],
    functions=["cholesky_n32", "cholesky_n64"],
    verbose=False,
    extra_cuda_cflags=["-O3", "--use_fast_math", "-DCHOLESKY64_VARIANT=1"],
)

N128_CUDA_SRC = r"""

#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cmath>
#include <algorithm>
#include <stdexcept>
#include <string>

#define CHECK_CUDA(expr) do { \
  cudaError_t _e = (expr); \
  if (_e != cudaSuccess) { \
    throw std::runtime_error(std::string("CUDA error: ") + cudaGetErrorString(_e)); \
  } \
} while (0)


// -------------------- n=128 research kernels (VARIANT 0..4) --------------------
// Failed baselines (do not regress to these alone):
//   full row[128]/thr ~665us; naive blocked BS32 ~258us; hier 2x64 ~310us.
// Target: Modal beat ~151us torch; Popcorn <=130us @256x128.
//
// VARIANT 0: phase-split halfreg left (known best ~113us Modal @ MPB=2)
// VARIANT 1: blocked BS multi-matrix SMEM SYRK
// VARIANT 2: interleaved even/odd halfreg (load-balanced dots)
// VARIANT 3: right-looking halfreg
// VARIANT 4: 256-thread blocked BS=16, 2 thr/row SYRK

#ifndef CHOLESKY128_BS
#define CHOLESKY128_BS 16
#endif
#ifndef CHOLESKY128_VARIANT
#define CHOLESKY128_VARIANT 0
#endif
#ifndef CHOLESKY128_MPB
#define CHOLESKY128_MPB 2
#endif

// ---- V0: phase-split halfreg left-looking (optimized) ----
// 2 threads/row hold col[64]. Phase0 k in [0,64): only half0 owns columns;
// no shfl (half1 partial always 0). Phase1 k in [64,128): both halves + xor.
template <int MPB>
__global__ void __launch_bounds__(256 * MPB, 1)
cholesky128_halfreg_phase(const float* __restrict__ A,
                          float* __restrict__ L,
                          int batch) {
  const int tid = threadIdx.x;   // 0..255
  const int slot = threadIdx.y;
  const int stride = gridDim.x * MPB;
  const int half = tid & 1;
  const int row = tid >> 1;
  const int col0 = half << 6;

  __shared__ float sk[MPB][128];

  for (int b = blockIdx.x * MPB + slot; b < batch; b += stride) {
    const float* a = A + (size_t)b * 16384u;
    float* out = L + (size_t)b * 16384u;

    float col[64];
#pragma unroll
    for (int t = 0; t < 16; ++t) {
      float4 v = *reinterpret_cast<const float4*>(a + row * 128 + col0 + t * 4);
      col[t * 4 + 0] = v.x;
      col[t * 4 + 1] = v.y;
      col[t * 4 + 2] = v.z;
      col[t * 4 + 3] = v.w;
    }

    // Phase 0: k in [0,64). half1 never contributes; skip shfl_xor.
#pragma unroll
    for (int k = 0; k < 64; ++k) {
      if (row == k && half == 0) {
#pragma unroll
        for (int j = 0; j < 64; ++j) {
          if (j < k) sk[slot][j] = col[j];
        }
      }
      __syncthreads();

      float sum = 0.f;
      if (half == 0) {
#pragma unroll
        for (int j = 0; j < 64; ++j) {
          if (j < k) sum = fmaf(col[j], sk[slot][j], sum);
        }
        float aik = col[k] - sum;
        if (row == k) {
          float d = sqrtf(fmaxf(aik, 0.f));
          col[k] = d;
          sk[slot][k] = d;
        } else {
          // stash aik in a register; need inv after sync — use sk write only for diag
          // For non-diag rows keep aik in col temporarily via reinterpret: store in unused
          // upper half slot area is free for phase0 — use sk[slot][64+row] is wrong size.
          // Keep classic path: write diag, sync, then scale using recomputed aik... 
          // Recompute is free enough with j<k already done; stash aik in col[k] temp:
          if (row > k) col[k] = aik; // temp: pre-scale residual
        }
      }
      __syncthreads();

      if (half == 0) {
        float inv = 1.f / sk[slot][k];
        if (row > k) col[k] = col[k] * inv;
        else if (row < k) col[k] = 0.f;
      }
      __syncthreads();
    }

    // Phase 1: k in [64,128)
#pragma unroll
    for (int kk = 0; kk < 64; ++kk) {
      const int k = 64 + kk;
      if (row == k) {
        if (half == 0) {
#pragma unroll
          for (int j = 0; j < 64; ++j) sk[slot][j] = col[j];
        } else {
#pragma unroll
          for (int j = 0; j < 64; ++j) {
            if (j < kk) sk[slot][64 + j] = col[j];
          }
        }
      }
      __syncthreads();

      float sum = 0.f;
      if (half == 0) {
#pragma unroll
        for (int j = 0; j < 64; ++j) {
          sum = fmaf(col[j], sk[slot][j], sum);
        }
      } else {
#pragma unroll
        for (int j = 0; j < 64; ++j) {
          if (j < kk) sum = fmaf(col[j], sk[slot][64 + j], sum);
        }
      }
      sum += __shfl_xor_sync(0xffffffffu, sum, 1);

      if (half == 1) {
        float aik = col[kk] - sum;
        if (row == k) {
          float d = sqrtf(fmaxf(aik, 0.f));
          col[kk] = d;
          sk[slot][k] = d;
        } else if (row > k) {
          col[kk] = aik; // temp residual
        }
      }
      __syncthreads();

      if (half == 1) {
        float inv = 1.f / sk[slot][k];
        if (row > k) col[kk] = col[kk] * inv;
        else if (row < k) col[kk] = 0.f;
      }
      __syncthreads();
    }

#pragma unroll
    for (int t = 0; t < 16; ++t) {
      float4 v;
      int j0 = col0 + t * 4;
      v.x = (row >= (j0 + 0)) ? col[t * 4 + 0] : 0.f;
      v.y = (row >= (j0 + 1)) ? col[t * 4 + 1] : 0.f;
      v.z = (row >= (j0 + 2)) ? col[t * 4 + 2] : 0.f;
      v.w = (row >= (j0 + 3)) ? col[t * 4 + 3] : 0.f;
      *reinterpret_cast<float4*>(out + row * 128 + j0) = v;
    }
  }
}

// ---- V2: interleaved even/odd column halfreg ----
// col[j] holds global column (2*j + half). Both threads contribute ~k/2 terms.
template <int MPB>
__global__ void __launch_bounds__(256 * MPB, 1)
cholesky128_halfreg_interleave(const float* __restrict__ A,
                               float* __restrict__ L,
                               int batch) {
  const int tid = threadIdx.x;
  const int slot = threadIdx.y;
  const int stride = gridDim.x * MPB;
  const int half = tid & 1;
  const int row = tid >> 1;

  __shared__ float sk[MPB][128];

  for (int b = blockIdx.x * MPB + slot; b < batch; b += stride) {
    const float* a = A + (size_t)b * 16384u;
    float* out = L + (size_t)b * 16384u;

    float col[64];
#pragma unroll
    for (int j = 0; j < 64; ++j) {
      col[j] = a[row * 128 + (j * 2 + half)];
    }

#pragma unroll
    for (int k = 0; k < 128; ++k) {
      // Publish L[k, 0:k] from row-k's two halves
      if (row == k) {
#pragma unroll
        for (int j = 0; j < 64; ++j) {
          int c = j * 2 + half;
          if (c < k) sk[slot][c] = col[j];
        }
      }
      __syncthreads();

      float sum = 0.f;
#pragma unroll
      for (int j = 0; j < 64; ++j) {
        int c = j * 2 + half;
        if (c < k) sum = fmaf(col[j], sk[slot][c], sum);
      }
      sum += __shfl_xor_sync(0xffffffffu, sum, 1);

      // Owner of column k updates A[row,k]
      if (half == (k & 1)) {
        int jk = k >> 1;
        float aik = col[jk] - sum;
        if (row == k) {
          float d = sqrtf(fmaxf(aik, 0.f));
          col[jk] = d;
          sk[slot][k] = d;
        } else if (row > k) {
          col[jk] = aik;
        }
      }
      __syncthreads();

      if (half == (k & 1)) {
        int jk = k >> 1;
        float inv = 1.f / sk[slot][k];
        if (row > k) col[jk] = col[jk] * inv;
        else if (row < k) col[jk] = 0.f;
      }
      __syncthreads();
    }

    // Store (scatter even/odd)
#pragma unroll
    for (int j = 0; j < 64; ++j) {
      int c = j * 2 + half;
      out[row * 128 + c] = (row >= c) ? col[j] : 0.f;
    }
  }
}

// ---- V3: right-looking halfreg ----
// Publish L[:,k] to SMEM, rank-1 update trailing columns in both halves.
template <int MPB>
__global__ void __launch_bounds__(256 * MPB, 1)
cholesky128_halfreg_right(const float* __restrict__ A,
                          float* __restrict__ L,
                          int batch) {
  const int tid = threadIdx.x;
  const int slot = threadIdx.y;
  const int stride = gridDim.x * MPB;
  const int half = tid & 1;
  const int row = tid >> 1;
  const int col0 = half << 6;

  __shared__ float colk[MPB][128];
  __shared__ float diag_s[MPB];

  for (int b = blockIdx.x * MPB + slot; b < batch; b += stride) {
    const float* a = A + (size_t)b * 16384u;
    float* out = L + (size_t)b * 16384u;

    float col[64];
#pragma unroll
    for (int t = 0; t < 16; ++t) {
      float4 v = *reinterpret_cast<const float4*>(a + row * 128 + col0 + t * 4);
      col[t * 4 + 0] = v.x;
      col[t * 4 + 1] = v.y;
      col[t * 4 + 2] = v.z;
      col[t * 4 + 3] = v.w;
    }

#pragma unroll
    for (int k = 0; k < 128; ++k) {
      // Diag from owner of column k
      const int kh = k >> 6;          // which half owns k
      const int kk = k & 63;          // local index
      if (half == kh && row == k) {
        float d = sqrtf(fmaxf(col[kk], 0.f));
        col[kk] = d;
        diag_s[slot] = d;
      }
      __syncthreads();

      float inv = 1.f / diag_s[slot];
      float lik = 0.f;
      if (half == kh) {
        if (row > k) {
          lik = col[kk] * inv;
          col[kk] = lik;
        } else if (row == k) {
          lik = diag_s[slot];
        } else {
          col[kk] = 0.f;
        }
      }
      // Broadcast full column k via SMEM (only half owner writes active rows)
      if (half == kh) {
        colk[slot][row] = (row >= k) ? ((row == k) ? diag_s[slot] : lik) : 0.f;
      }
      __syncthreads();

      float lik_i = colk[slot][row];
      // Rank-1: for all j > k owned by this half, update if row >= j
#pragma unroll
      for (int jloc = 0; jloc < 64; ++jloc) {
        int j = col0 + jloc;
        if (j > k && row >= j) {
          col[jloc] = fmaf(-lik_i, colk[slot][j], col[jloc]);
        }
      }
      __syncthreads();
    }

#pragma unroll
    for (int t = 0; t < 16; ++t) {
      float4 v;
      int j0 = col0 + t * 4;
      v.x = (row >= (j0 + 0)) ? col[t * 4 + 0] : 0.f;
      v.y = (row >= (j0 + 1)) ? col[t * 4 + 1] : 0.f;
      v.z = (row >= (j0 + 2)) ? col[t * 4 + 2] : 0.f;
      v.w = (row >= (j0 + 3)) ? col[t * 4 + 3] : 0.f;
      *reinterpret_cast<float4*>(out + row * 128 + j0) = v;
    }
  }
}

// ---- V1: blocked BS multi-matrix, 128 thr/matrix, SYRK panel-in-regs ----
// Right-looking blocked: potf2 panel (intra-panel left-looking), rank-BS SYRK.
// Modal best: BS=32 MPB=1 ~75.6us @256x128 (torch ~151). Prior naive BS32 ~258.
// NOTE: do NOT float4-cast tile[row*129+col] — 129-stride breaks 16B alignment.
template <int MPB, int BS>
__global__ void __launch_bounds__(128 * MPB, 2)
cholesky128_blocked_mpb(const float* __restrict__ A,
                        float* __restrict__ L,
                        int batch) {
  const int tid = threadIdx.x;
  const int slot = threadIdx.y;
  const int stride = gridDim.x * MPB;

  extern __shared__ float smem[];
  constexpr int TILE = 128 * 129;
  float* tile = smem + slot * TILE;

  for (int b = blockIdx.x * MPB + slot; b < batch; b += stride) {
    const float* a = A + (size_t)b * 16384u;
    float* out = L + (size_t)b * 16384u;

#pragma unroll
    for (int t = 0; t < 32; ++t) {
      float4 v = *reinterpret_cast<const float4*>(a + tid * 128 + t * 4);
      int j0 = t * 4;
      tile[tid * 129 + j0 + 0] = v.x;
      tile[tid * 129 + j0 + 1] = v.y;
      tile[tid * 129 + j0 + 2] = v.z;
      tile[tid * 129 + j0 + 3] = v.w;
    }
    __syncthreads();

#pragma unroll
    for (int kb = 0; kb < 128; kb += BS) {
#pragma unroll
      for (int kk = 0; kk < BS; ++kk) {
        const int k = kb + kk;
        float sum = 0.f;
#pragma unroll
        for (int j = 0; j < BS; ++j) {
          if (j < kk) {
            sum = fmaf(tile[tid * 129 + (kb + j)],
                       tile[k * 129 + (kb + j)], sum);
          }
        }
        float aik = tile[tid * 129 + k] - sum;
        if (tid == k) {
          tile[k * 129 + k] = sqrtf(fmaxf(aik, 0.f));
        }
        __syncthreads();
        float inv = 1.f / tile[k * 129 + k];
        if (tid > k) tile[tid * 129 + k] = aik * inv;
        else if (tid < k) tile[tid * 129 + k] = 0.f;
        __syncthreads();
      }

      if (kb + BS < 128) {
        float panel[BS];
#pragma unroll
        for (int t = 0; t < BS; ++t) panel[t] = tile[tid * 129 + kb + t];
        __syncthreads();

#pragma unroll 1
        for (int j = kb + BS; j < 128; ++j) {
          if (tid >= j) {
            float s = 0.f;
#pragma unroll
            for (int t = 0; t < BS; ++t) {
              s = fmaf(panel[t], tile[j * 129 + kb + t], s);
            }
            tile[tid * 129 + j] -= s;
          }
        }
        __syncthreads();
      }
    }

#pragma unroll
    for (int t = 0; t < 32; ++t) {
      float4 v;
      int j0 = t * 4;
      v.x = (tid >= (j0 + 0)) ? tile[tid * 129 + j0 + 0] : 0.f;
      v.y = (tid >= (j0 + 1)) ? tile[tid * 129 + j0 + 1] : 0.f;
      v.z = (tid >= (j0 + 2)) ? tile[tid * 129 + j0 + 2] : 0.f;
      v.w = (tid >= (j0 + 3)) ? tile[tid * 129 + j0 + 3] : 0.f;
      *reinterpret_cast<float4*>(out + tid * 128 + j0) = v;
    }
    if (b + stride < batch) __syncthreads();
  }
}

// ---- V4: 256-thread CTA, BS=16, 2 threads/row for vectorized SYRK ----
// thr = row*2+lane; lane 0/1 split panel BS/2 for SYRK partial dots + xor reduce.
// Panel factor still uses lane0 as row owner (lane1 assists SYRK only).
template <int MPB, int BS>
__global__ void __launch_bounds__(256 * MPB, 1)
cholesky128_blocked_vecsyrk(const float* __restrict__ A,
                            float* __restrict__ L,
                            int batch) {
  const int tid = threadIdx.x; // 0..255
  const int slot = threadIdx.y;
  const int stride = gridDim.x * MPB;
  const int lane = tid & 1;
  const int row = tid >> 1; // 0..127

  extern __shared__ float smem[];
  constexpr int TILE = 128 * 129;
  float* tile = smem + slot * TILE;

  for (int b = blockIdx.x * MPB + slot; b < batch; b += stride) {
    const float* a = A + (size_t)b * 16384u;
    float* out = L + (size_t)b * 16384u;

    // Cooperative load: 256 thr * 64 elems? each thr loads 64 floats of its row half
    {
      const int col0 = lane << 6;
#pragma unroll
      for (int t = 0; t < 16; ++t) {
        float4 v = *reinterpret_cast<const float4*>(a + row * 128 + col0 + t * 4);
        int j0 = col0 + t * 4;
        tile[row * 129 + j0 + 0] = v.x;
        tile[row * 129 + j0 + 1] = v.y;
        tile[row * 129 + j0 + 2] = v.z;
        tile[row * 129 + j0 + 3] = v.w;
      }
    }
    __syncthreads();

#pragma unroll 1
    for (int kb = 0; kb < 128; kb += BS) {
      // Panel potf2/trsm: only lane0 rows (one thr per row)
#pragma unroll
      for (int kk = 0; kk < BS; ++kk) {
        const int k = kb + kk;
        float sum = 0.f;
        if (lane == 0) {
#pragma unroll
          for (int j = 0; j < BS; ++j) {
            if (j < kk) {
              sum = fmaf(tile[row * 129 + (kb + j)],
                         tile[k * 129 + (kb + j)], sum);
            }
          }
          float aik = tile[row * 129 + k] - sum;
          if (row == k) {
            tile[k * 129 + k] = sqrtf(fmaxf(aik, 0.f));
          } else if (row > k) {
            // stash
            tile[row * 129 + k] = aik;
          }
        }
        __syncthreads();
        if (lane == 0) {
          float inv = 1.f / tile[k * 129 + k];
          if (row > k) tile[row * 129 + k] = tile[row * 129 + k] * inv;
          else if (row < k) tile[row * 129 + k] = 0.f;
        }
        __syncthreads();
      }

      if (kb + BS < 128) {
        // Each thr holds BS/2 panel elems for its row
        float panel_half[BS / 2];
        const int t0 = lane * (BS / 2);
#pragma unroll
        for (int t = 0; t < BS / 2; ++t) {
          panel_half[t] = tile[row * 129 + kb + t0 + t];
        }
        __syncthreads();

#pragma unroll 1
        for (int j = kb + BS; j < 128; ++j) {
          float s = 0.f;
          if (row >= j) {
#pragma unroll
            for (int t = 0; t < BS / 2; ++t) {
              s = fmaf(panel_half[t], tile[j * 129 + kb + t0 + t], s);
            }
          }
          s += __shfl_xor_sync(0xffffffffu, s, 1);
          if (lane == 0 && row >= j) {
            tile[row * 129 + j] -= s;
          }
        }
        __syncthreads();
      }
    }

    if (lane == 0) {
#pragma unroll
      for (int t = 0; t < 32; ++t) {
        float4 v;
        int j0 = t * 4;
        v.x = (row >= (j0 + 0)) ? tile[row * 129 + j0 + 0] : 0.f;
        v.y = (row >= (j0 + 1)) ? tile[row * 129 + j0 + 1] : 0.f;
        v.z = (row >= (j0 + 2)) ? tile[row * 129 + j0 + 2] : 0.f;
        v.w = (row >= (j0 + 3)) ? tile[row * 129 + j0 + 3] : 0.f;
        *reinterpret_cast<float4*>(out + row * 128 + j0) = v;
      }
    }
    if (b + stride < batch) __syncthreads();
  }
}


// ---- V5: pure right-looking rank-1, 128 thr, full SMEM tile ----
template <int MPB>
__global__ void __launch_bounds__(128 * MPB, 2)
cholesky128_right_rank1(const float* __restrict__ A,
                        float* __restrict__ L,
                        int batch) {
  const int tid = threadIdx.x;
  const int slot = threadIdx.y;
  const int stride = gridDim.x * MPB;
  extern __shared__ float smem[];
  constexpr int TILE = 128 * 129;
  float* tile = smem + slot * (TILE + 128);
  float* colk = tile + TILE;

  for (int b = blockIdx.x * MPB + slot; b < batch; b += stride) {
    const float* a = A + (size_t)b * 16384u;
    float* out = L + (size_t)b * 16384u;
#pragma unroll
    for (int t = 0; t < 32; ++t) {
      float4 v = *reinterpret_cast<const float4*>(a + tid * 128 + t * 4);
      int j0 = t * 4;
      tile[tid * 129 + j0 + 0] = v.x;
      tile[tid * 129 + j0 + 1] = v.y;
      tile[tid * 129 + j0 + 2] = v.z;
      tile[tid * 129 + j0 + 3] = v.w;
    }
    __syncthreads();

#pragma unroll
    for (int k = 0; k < 128; ++k) {
      if (tid == k) {
        tile[k * 129 + k] = sqrtf(fmaxf(tile[k * 129 + k], 0.f));
      }
      __syncthreads();
      float inv = 1.f / tile[k * 129 + k];
      float lik;
      if (tid > k) {
        lik = tile[tid * 129 + k] * inv;
        tile[tid * 129 + k] = lik;
      } else if (tid == k) {
        lik = tile[k * 129 + k];
      } else {
        lik = 0.f;
        tile[tid * 129 + k] = 0.f;
      }
      colk[tid] = lik;
      __syncthreads();
      float li = colk[tid];
#pragma unroll 1
      for (int j = k + 1; j < 128; ++j) {
        if (tid >= j) {
          tile[tid * 129 + j] = fmaf(-li, colk[j], tile[tid * 129 + j]);
        }
      }
      __syncthreads();
    }

#pragma unroll
    for (int t = 0; t < 32; ++t) {
      float4 v;
      int j0 = t * 4;
      v.x = (tid >= (j0 + 0)) ? tile[tid * 129 + j0 + 0] : 0.f;
      v.y = (tid >= (j0 + 1)) ? tile[tid * 129 + j0 + 1] : 0.f;
      v.z = (tid >= (j0 + 2)) ? tile[tid * 129 + j0 + 2] : 0.f;
      v.w = (tid >= (j0 + 3)) ? tile[tid * 129 + j0 + 3] : 0.f;
      *reinterpret_cast<float4*>(out + tid * 128 + j0) = v;
    }
    if (b + stride < batch) __syncthreads();
  }
}

torch::Tensor cholesky_n128(torch::Tensor A) {
  TORCH_CHECK(A.is_cuda(), "A must be CUDA");
  TORCH_CHECK(A.scalar_type() == at::kFloat, "A must be float32");
  TORCH_CHECK(A.dim() == 3 && A.size(1) == 128 && A.size(2) == 128, "expected (B,128,128)");

  auto Ac = A.contiguous();
  auto L = torch::empty_like(Ac);
  const int batch = static_cast<int>(Ac.size(0));
  if (batch == 0) return L;

  constexpr int mpb = CHOLESKY128_MPB;
  constexpr int bs = CHOLESKY128_BS;

#if CHOLESKY128_VARIANT == 1
  {
    dim3 block(128, mpb);
    int need = (batch + mpb - 1) / mpb;
    int grid = need < 2048 ? need : 2048;
    if (grid < 1) grid = 1;
    const int shmem = mpb * (128 * 129) * (int)sizeof(float);
    CHECK_CUDA(cudaFuncSetAttribute(
        cholesky128_blocked_mpb<mpb, bs>,
        cudaFuncAttributeMaxDynamicSharedMemorySize, shmem));
    cholesky128_blocked_mpb<mpb, bs><<<grid, block, shmem>>>(
        Ac.data_ptr<float>(), L.data_ptr<float>(), batch);
  }
#elif CHOLESKY128_VARIANT == 2
  {
    dim3 block(256, mpb);
    int need = (batch + mpb - 1) / mpb;
    int grid = need < 2048 ? need : 2048;
    if (grid < 1) grid = 1;
    cholesky128_halfreg_interleave<mpb><<<grid, block>>>(
        Ac.data_ptr<float>(), L.data_ptr<float>(), batch);
  }
#elif CHOLESKY128_VARIANT == 3
  {
    dim3 block(256, mpb);
    int need = (batch + mpb - 1) / mpb;
    int grid = need < 2048 ? need : 2048;
    if (grid < 1) grid = 1;
    cholesky128_halfreg_right<mpb><<<grid, block>>>(
        Ac.data_ptr<float>(), L.data_ptr<float>(), batch);
  }
#elif CHOLESKY128_VARIANT == 4
  {
    dim3 block(256, mpb);
    int need = (batch + mpb - 1) / mpb;
    int grid = need < 2048 ? need : 2048;
    if (grid < 1) grid = 1;
    const int shmem = mpb * (128 * 129) * (int)sizeof(float);
    CHECK_CUDA(cudaFuncSetAttribute(
        cholesky128_blocked_vecsyrk<mpb, bs>,
        cudaFuncAttributeMaxDynamicSharedMemorySize, shmem));
    cholesky128_blocked_vecsyrk<mpb, bs><<<grid, block, shmem>>>(
        Ac.data_ptr<float>(), L.data_ptr<float>(), batch);
  }
#elif CHOLESKY128_VARIANT == 5
  {
    dim3 block(128, mpb);
    int need = (batch + mpb - 1) / mpb;
    int grid = need < 2048 ? need : 2048;
    if (grid < 1) grid = 1;
    const int shmem = mpb * (128 * 129 + 128) * (int)sizeof(float);
    CHECK_CUDA(cudaFuncSetAttribute(
        cholesky128_right_rank1<mpb>,
        cudaFuncAttributeMaxDynamicSharedMemorySize, shmem));
    cholesky128_right_rank1<mpb><<<grid, block, shmem>>>(
        Ac.data_ptr<float>(), L.data_ptr<float>(), batch);
  }
#else
  {
    dim3 block(256, mpb);
    int need = (batch + mpb - 1) / mpb;
    int grid = need < 2048 ? need : 2048;
    if (grid < 1) grid = 1;
    cholesky128_halfreg_phase<mpb><<<grid, block>>>(
        Ac.data_ptr<float>(), L.data_ptr<float>(), batch);
  }
#endif
  CHECK_CUDA(cudaGetLastError());
  return L;
}
// ---- n=256 dense-panel SYRK (beam_mid V1) ----

#ifndef CHOLESKY256_BS
#define CHOLESKY256_BS 32
#endif

__device__ __forceinline__ int lt256(int i, int j) {
  return (i * (i + 1)) / 2 + j;
}

// ---- V1 STRUCTURAL: dense-panel SYRK + multi-j ILP (JU=4) ----
// After each panel potf, publish L[:, kb:kb+BS] into dense dens[row*(BS+1)+t]
// so SYRK is two dense vectors (far fewer bank conflicts than packed gather).
// Multi-j unrolling amortizes dens row loads for L[j,*].
template <int BS>
__global__ void __launch_bounds__(256, 1)
cholesky256_dense_syrk(const float* __restrict__ A,
                       float* __restrict__ L,
                       int batch) {
  const int tid = threadIdx.x;  // row
  extern __shared__ float smem[];
  constexpr int PACK = 256 * 257 / 2;
  constexpr int DENS_STRIDE = BS + 1;  // bank pad
  float* tile = smem;
  float* dens = smem + PACK;  // 256 * DENS_STRIDE

  for (int b = blockIdx.x; b < batch; b += gridDim.x) {
    const float* a = A + (size_t)b * 65536u;
    float* out = L + (size_t)b * 65536u;

#pragma unroll 1
    for (int t = 0; t < 64; ++t) {
      int j0 = t * 4;
      float4 v = *reinterpret_cast<const float4*>(a + tid * 256 + j0);
      if (j0 + 0 <= tid) tile[lt256(tid, j0 + 0)] = v.x;
      if (j0 + 1 <= tid) tile[lt256(tid, j0 + 1)] = v.y;
      if (j0 + 2 <= tid) tile[lt256(tid, j0 + 2)] = v.z;
      if (j0 + 3 <= tid) tile[lt256(tid, j0 + 3)] = v.w;
    }
    __syncthreads();

#pragma unroll 1
    for (int kb = 0; kb < 256; kb += BS) {
      // Panel potf2 (left-looking within panel)
#pragma unroll
      for (int kk = 0; kk < BS; ++kk) {
        const int k = kb + kk;
        float sum = 0.f;
        float aik = 0.f;
        if (tid >= k) {
#pragma unroll
          for (int j = 0; j < BS; ++j) {
            if (j < kk) {
              sum = fmaf(tile[lt256(tid, kb + j)],
                         tile[lt256(k, kb + j)], sum);
            }
          }
          aik = tile[lt256(tid, k)] - sum;
          if (tid == k) {
            tile[lt256(k, k)] = sqrtf(fmaxf(aik, 0.f));
          }
        }
        __syncthreads();
        if (tid > k) {
          tile[lt256(tid, k)] = aik * (1.f / tile[lt256(k, k)]);
        } else if (tid < k) {
          // keep upper-of-panel zero in packed (already unused)
        }
        __syncthreads();
      }

      // Publish dense panel
#pragma unroll
      for (int t = 0; t < BS; ++t) {
        dens[tid * DENS_STRIDE + t] =
            (tid >= (kb + t)) ? tile[lt256(tid, kb + t)] : 0.f;
      }
      __syncthreads();

      if (kb + BS < 256) {
        float my_panel[BS];
#pragma unroll
        for (int t = 0; t < BS; ++t) my_panel[t] = dens[tid * DENS_STRIDE + t];

        // Multi-j SYRK: update columns j in steps of 4
        const int j0 = kb + BS;
        const int nfull = ((256 - j0) / 4) * 4;
#pragma unroll 1
        for (int jj = 0; jj < nfull; jj += 4) {
          const int j = j0 + jj;
          float s0 = 0.f, s1 = 0.f, s2 = 0.f, s3 = 0.f;
          float p0[BS], p1[BS], p2[BS], p3[BS];
#pragma unroll
          for (int t = 0; t < BS; ++t) {
            p0[t] = dens[j * DENS_STRIDE + t];
            p1[t] = dens[(j + 1) * DENS_STRIDE + t];
            p2[t] = dens[(j + 2) * DENS_STRIDE + t];
            p3[t] = dens[(j + 3) * DENS_STRIDE + t];
          }
#pragma unroll
          for (int t = 0; t < BS; ++t) {
            float mp = my_panel[t];
            s0 = fmaf(mp, p0[t], s0);
            s1 = fmaf(mp, p1[t], s1);
            s2 = fmaf(mp, p2[t], s2);
            s3 = fmaf(mp, p3[t], s3);
          }
          if (tid >= j)     tile[lt256(tid, j)]     -= s0;
          if (tid >= j + 1) tile[lt256(tid, j + 1)] -= s1;
          if (tid >= j + 2) tile[lt256(tid, j + 2)] -= s2;
          if (tid >= j + 3) tile[lt256(tid, j + 3)] -= s3;
        }
#pragma unroll 1
        for (int j = j0 + nfull; j < 256; ++j) {
          float s = 0.f;
#pragma unroll
          for (int t = 0; t < BS; ++t) {
            s = fmaf(my_panel[t], dens[j * DENS_STRIDE + t], s);
          }
          if (tid >= j) tile[lt256(tid, j)] -= s;
        }
        __syncthreads();
      }
    }

#pragma unroll 1
    for (int t = 0; t < 64; ++t) {
      int j0 = t * 4;
      float4 v;
      v.x = (tid >= (j0 + 0)) ? tile[lt256(tid, j0 + 0)] : 0.f;
      v.y = (tid >= (j0 + 1)) ? tile[lt256(tid, j0 + 1)] : 0.f;
      v.z = (tid >= (j0 + 2)) ? tile[lt256(tid, j0 + 2)] : 0.f;
      v.w = (tid >= (j0 + 3)) ? tile[lt256(tid, j0 + 3)] : 0.f;
      *reinterpret_cast<float4*>(out + tid * 256 + j0) = v;
    }
  }
}


static void launch_n256(const float* A, float* L, int batch) {
  constexpr int bs = CHOLESKY256_BS;
  int grid = batch < 2048 ? batch : 2048;
  if (grid < 1) grid = 1;
  const int shmem = (256 * 257 / 2 + 256 * (bs + 1)) * (int)sizeof(float);
  CHECK_CUDA(cudaFuncSetAttribute(
      cholesky256_dense_syrk<bs>,
      cudaFuncAttributeMaxDynamicSharedMemorySize, shmem));
  cholesky256_dense_syrk<bs><<<grid, 256, shmem>>>(A, L, batch);
  CHECK_CUDA(cudaGetLastError());
}

torch::Tensor cholesky_n256(torch::Tensor A) {
  TORCH_CHECK(A.is_cuda(), "A must be CUDA");
  TORCH_CHECK(A.scalar_type() == at::kFloat, "A must be float32");
  TORCH_CHECK(A.dim() == 3 && A.size(1) == 256 && A.size(2) == 256,
              "expected (B,256,256)");
  auto Ac = A.contiguous();
  auto L = torch::empty_like(Ac);
  const int batch = static_cast<int>(Ac.size(0));
  if (batch == 0) return L;
  launch_n256(Ac.data_ptr<float>(), L.data_ptr<float>(), batch);
  return L;
}

"""

N128_CPP_SRC = r"""
torch::Tensor cholesky_n128(torch::Tensor A);
torch::Tensor cholesky_n256(torch::Tensor A);

"""

_mod_n128 = load_inline(
    name="chol_n128_n256_dense_v1",
    cpp_sources=[N128_CPP_SRC],
    cuda_sources=[N128_CUDA_SRC],
    functions=["cholesky_n128", "cholesky_n256"],
    verbose=False,
    extra_cuda_cflags=[
        "-O3",
        "-DCHOLESKY128_VARIANT=1",
        "-DCHOLESKY128_MPB=1",
        "-DCHOLESKY128_BS=32",
        "-DCHOLESKY256_BS=32",
    ],
)

# ---------------------------------------------------------------------------
# Medium high-batch: blocked right-looking with batched potf / TRSM / GEMM
# Row-major SPD as col-major FILL_UPPER → L in lower triangle (champ convention).
# ---------------------------------------------------------------------------
MEDIUM2_CUDA_SRC = r"""
#include <torch/extension.h>
#include <cublas_v2.h>
#include <cusolverDn.h>
#include <cuda_runtime.h>
#include <algorithm>
#include <mutex>
#include <stdexcept>
#include <string>

// large3-style TF32 trailing knobs
#ifndef LARGE3_GEMMEX_M_MIN
#define LARGE3_GEMMEX_M_MIN 4096
#endif
#ifndef LARGE3_REFINE_TAIL
#define LARGE3_REFINE_TAIL 0
#endif
#ifndef LARGE3_TF32_MATH
#define LARGE3_TF32_MATH 1
#endif

#define CHECK_CUDA(expr) do { \
  cudaError_t _err = (expr); \
  if (_err != cudaSuccess) { \
    throw std::runtime_error(std::string("CUDA error: ") + cudaGetErrorString(_err)); \
  } \
} while (0)

#define CHECK_CUBLAS(expr) do { \
  cublasStatus_t _st = (expr); \
  if (_st != CUBLAS_STATUS_SUCCESS) { \
    throw std::runtime_error(std::string("cuBLAS error ") + std::to_string((int)_st)); \
  } \
} while (0)

#define CHECK_CUSOLVER(expr) do { \
  cusolverStatus_t _st = (expr); \
  if (_st != CUSOLVER_STATUS_SUCCESS) { \
    throw std::runtime_error(std::string("cuSOLVER error ") + std::to_string((int)_st)); \
  } \
} while (0)

static cublasHandle_t g_blas = nullptr;
static cusolverDnHandle_t g_solver = nullptr;
static std::once_flag g_once;

// Persistent device pointer arrays + potf workspace.
static float** g_d_A = nullptr;
static float** g_d_B = nullptr;
static int g_ptr_cap = 0;
static float* g_work = nullptr;
static int g_work_cap = 0;
static int* g_info = nullptr;
static int g_info_cap = 0;

static void init_handles() {
  std::call_once(g_once, []() {
    CHECK_CUBLAS(cublasCreate(&g_blas));
    CHECK_CUBLAS(cublasSetMathMode(g_blas, CUBLAS_TF32_TENSOR_OP_MATH));
    CHECK_CUSOLVER(cusolverDnCreate(&g_solver));
  });
}

static void ensure_ptrs(int batch) {
  if (batch > g_ptr_cap) {
    if (g_d_A) CHECK_CUDA(cudaFree(g_d_A));
    if (g_d_B) CHECK_CUDA(cudaFree(g_d_B));
    CHECK_CUDA(cudaMalloc(reinterpret_cast<void**>(&g_d_A),
                          (size_t)batch * sizeof(float*)));
    CHECK_CUDA(cudaMalloc(reinterpret_cast<void**>(&g_d_B),
                          (size_t)batch * sizeof(float*)));
    g_ptr_cap = batch;
  }
}

static float* ensure_work(int lwork) {
  if (lwork > g_work_cap) {
    if (g_work) CHECK_CUDA(cudaFree(g_work));
    CHECK_CUDA(cudaMalloc(reinterpret_cast<void**>(&g_work),
                          (size_t)std::max(lwork, 1) * sizeof(float)));
    g_work_cap = std::max(lwork, 1);
  }
  return g_work;
}

static int* ensure_info(int ninfo) {
  if (ninfo > g_info_cap) {
    if (g_info) CHECK_CUDA(cudaFree(g_info));
    CHECK_CUDA(cudaMalloc(reinterpret_cast<void**>(&g_info),
                          (size_t)std::max(ninfo, 1) * sizeof(int)));
    g_info_cap = std::max(ninfo, 1);
  }
  return g_info;
}

__global__ void zero_upper_kernel(float* __restrict__ L, int n, int batch) {
  const int b = blockIdx.z;
  const int i = blockIdx.y * blockDim.y + threadIdx.y;
  const int j = blockIdx.x * blockDim.x + threadIdx.x;
  if (b >= batch || i >= n || j >= n) return;
  if (j > i) {
    L[(size_t)b * (size_t)n * (size_t)n + (size_t)i * (size_t)n + (size_t)j] = 0.f;
  }
}

// ptrs[b] = base + b * mat_stride + offset
__global__ void fill_offset_ptrs_kernel(
    float** __restrict__ ptrs,
    float* base,
    long long mat_stride,
    long long offset,
    int batch) {
  int b = blockIdx.x * blockDim.x + threadIdx.x;
  if (b < batch) {
    ptrs[b] = base + (long long)b * mat_stride + offset;
  }
}

static void fill_ptrs(float** d_ptrs, float* base, long long mat_stride,
                      long long offset, int batch) {
  int threads = 256;
  int blocks = (batch + threads - 1) / threads;
  if (blocks < 1) blocks = 1;
  fill_offset_ptrs_kernel<<<blocks, threads>>>(
      d_ptrs, base, mat_stride, offset, batch);
  CHECK_CUDA(cudaGetLastError());
}

// mode: 0 = SpotrfBatched on each diagonal tile; 1 = host loop potf reused workspace
torch::Tensor cholesky_blocked_batched(torch::Tensor A, int64_t nb_in, int64_t mode_in) {
  TORCH_CHECK(A.is_cuda() && A.scalar_type() == torch::kFloat32 && A.dim() == 3);
  TORCH_CHECK(A.size(1) == A.size(2));

  init_handles();

  auto L = A.contiguous().clone();
  const int batch = static_cast<int>(L.size(0));
  const int n = static_cast<int>(L.size(1));
  if (batch == 0) return L;

  int nb = static_cast<int>(nb_in);
  if (nb < 32) nb = 32;
  if (nb > n) nb = n;
  const int mode = static_cast<int>(mode_in);

  float* Lptr = L.data_ptr<float>();
  const long long mat_stride = (long long)n * (long long)n;

  ensure_ptrs(batch);
  int* d_info = ensure_info(std::max(batch, 1));
  // Clear info once; potf writes per call (batched fills all).
  CHECK_CUDA(cudaMemset(d_info, 0, (size_t)std::max(batch, 1) * sizeof(int)));

  int lwork = 0;
  CHECK_CUSOLVER(cusolverDnSpotrf_bufferSize(
      g_solver, CUBLAS_FILL_MODE_UPPER, nb, Lptr, n, &lwork));
  float* work = ensure_work(lwork);

  const float one = 1.f;
  const float minus_one = -1.f;

  for (int k = 0; k < n; k += nb) {
    const int kb = std::min(nb, n - k);
    const long long off_kk = (long long)k + (long long)k * (long long)n;

    // 1) Factor diagonal tiles Akk (kb x kb, lda = n)
    if (mode == 0) {
      fill_ptrs(g_d_A, Lptr, mat_stride, off_kk, batch);
      CHECK_CUSOLVER(cusolverDnSpotrfBatched(
          g_solver,
          CUBLAS_FILL_MODE_UPPER,
          kb,
          g_d_A,
          n,
          d_info,
          batch));
    } else {
      // Tight host loop, single workspace reused across batch items.
      for (int b = 0; b < batch; ++b) {
        float* Akk = Lptr + (long long)b * mat_stride + off_kk;
        CHECK_CUSOLVER(cusolverDnSpotrf(
            g_solver, CUBLAS_FILL_MODE_UPPER, kb, Akk, n,
            work, lwork, d_info));
      }
    }

    const int m = n - (k + kb);
    if (m <= 0) continue;

    const long long off_12 = (long long)k + (long long)(k + kb) * (long long)n;
    const long long off_22 = (long long)(k + kb) + (long long)(k + kb) * (long long)n;

    // 2) Panel TRSM: A12 := inv(Akk)^T * A12  (batched)
    // SIDE_LEFT, UPPER, OP_T: Akk^T X = A12, X is kb x m col-major
    fill_ptrs(g_d_A, Lptr, mat_stride, off_kk, batch);
    fill_ptrs(g_d_B, Lptr, mat_stride, off_12, batch);
    CHECK_CUBLAS(cublasStrsmBatched(
        g_blas,
        CUBLAS_SIDE_LEFT,
        CUBLAS_FILL_MODE_UPPER,
        CUBLAS_OP_T,
        CUBLAS_DIAG_NON_UNIT,
        kb,
        m,
        &one,
        (const float**)g_d_A,
        n,
        g_d_B,
        n,
        batch));

    // 3) Trailing GEMM: A22 -= A12^T * A12  (strided batched, TF32)
    float* A12 = Lptr + off_12;
    float* A22 = Lptr + off_22;
    CHECK_CUBLAS(cublasSgemmStridedBatched(
        g_blas,
        CUBLAS_OP_T,
        CUBLAS_OP_N,
        m,
        m,
        kb,
        &minus_one,
        A12,
        n,
        mat_stride,
        A12,
        n,
        mat_stride,
        &one,
        A22,
        n,
        mat_stride,
        batch));
  }

  dim3 block(16, 16);
  dim3 grid((n + 15) / 16, (n + 15) / 16, batch);
  zero_upper_kernel<<<grid, block>>>(Lptr, n, batch);
  CHECK_CUDA(cudaGetLastError());
  // No end-sync: leave work async like torch path.
  return L;
}

// Trailing A22 -= A12^T @ A12.
// Large m: GemmEx FAST_TF32. Small m / refine tail: Sgemm under handle math mode
// (TF32_TENSOR_OP when LARGE3_TF32_MATH=1 - same residual as champ/large2).
static inline void trailing_update(
    cublasHandle_t blas,
    int m,
    int kb,
    int gemmex_m_min,
    int force_fp32,
    const float* A12,
    float* A22,
    int ld) {
  const float one = 1.f;
  const float minus_one = -1.f;

  if (!force_fp32 && m >= gemmex_m_min) {
    CHECK_CUBLAS(cublasGemmEx(
        blas,
        CUBLAS_OP_T,
        CUBLAS_OP_N,
        m,
        m,
        kb,
        &minus_one,
        A12,
        CUDA_R_32F,
        ld,
        A12,
        CUDA_R_32F,
        ld,
        &one,
        A22,
        CUDA_R_32F,
        ld,
        CUBLAS_COMPUTE_32F_FAST_TF32,
        CUBLAS_GEMM_DEFAULT_TENSOR_OP));
  } else {
    CHECK_CUBLAS(cublasSgemm(
        blas,
        CUBLAS_OP_T,
        CUBLAS_OP_N,
        m,
        m,
        kb,
        &minus_one,
        A12,
        ld,
        A12,
        ld,
        &one,
        A22,
        ld));
  }
}

static void potf_blocked_one(
    float* A,
    int n,
    int nb,
    int gemmex_m_min,
    int refine_tail,
    float* work,
    int lwork,
    int* d_info) {
  const float one = 1.f;
  cublasHandle_t blas = g_blas;
  cusolverDnHandle_t solver = g_solver;

  for (int k = 0; k < n; k += nb) {
    const int kb = std::min(nb, n - k);
    float* Akk = A + (size_t)k + (size_t)k * (size_t)n;

    // Panel always FP32 cuSOLVER.
    CHECK_CUSOLVER(cusolverDnSpotrf(
        solver, CUBLAS_FILL_MODE_UPPER, kb, Akk, n, work, lwork, d_info));

    const int m = n - (k + kb);
    if (m <= 0) {
      continue;
    }

    float* A12 = A + (size_t)k + (size_t)(k + kb) * (size_t)n;
    float* A22 = A + (size_t)(k + kb) + (size_t)(k + kb) * (size_t)n;

    CHECK_CUBLAS(cublasStrsm(
        blas,
        CUBLAS_SIDE_LEFT,
        CUBLAS_FILL_MODE_UPPER,
        CUBLAS_OP_T,
        CUBLAS_DIAG_NON_UNIT,
        kb,
        m,
        &one,
        Akk,
        n,
        A12,
        n));

    // refine_tail: prefer Sgemm path for small remaining m (residual knob;
    // under TF32 math mode residual stays ~champ - true FP32 needs DEFAULT math).
    const int force_fp32 = (refine_tail > 0 && m <= refine_tail) ? 1 : 0;
    trailing_update(blas, m, kb, gemmex_m_min, force_fp32, A12, A22, n);
  }
}

torch::Tensor cholesky_tf32_blocked(
    torch::Tensor A,
    int64_t nb_in,
    int64_t gemmex_m_min_in,
    int64_t refine_tail_in) {
  TORCH_CHECK(A.is_cuda() && A.is_contiguous());
  TORCH_CHECK(A.scalar_type() == torch::kFloat32);
  TORCH_CHECK(A.dim() == 3 && A.size(1) == A.size(2));

  init_handles();

  const int batch = static_cast<int>(A.size(0));
  const int n = static_cast<int>(A.size(1));
  int nb = static_cast<int>(nb_in);
  if (nb < 32) nb = 32;
  if (nb > n) nb = n;
  int gemmex_m_min = static_cast<int>(gemmex_m_min_in);
  if (gemmex_m_min < 0) gemmex_m_min = LARGE3_GEMMEX_M_MIN;
  int refine_tail = static_cast<int>(refine_tail_in);
  if (refine_tail < 0) refine_tail = LARGE3_REFINE_TAIL;

  auto L = A.clone();
  float* Lptr = L.data_ptr<float>();
  const long long mat_stride = (long long)n * (long long)n;

  int lwork = 0;
  CHECK_CUSOLVER(cusolverDnSpotrf_bufferSize(
      g_solver, CUBLAS_FILL_MODE_UPPER, nb, Lptr, n, &lwork));
  auto work = torch::empty({std::max(lwork, 1)}, A.options());
  auto info = torch::zeros(
      {1}, torch::TensorOptions().dtype(torch::kInt32).device(A.device()));
  float* wptr = work.data_ptr<float>();
  int* iptr = info.data_ptr<int>();

  if (batch == 1) {
    potf_blocked_one(Lptr, n, nb, gemmex_m_min, refine_tail, wptr, lwork, iptr);
  } else {
    for (int b = 0; b < batch; ++b) {
      potf_blocked_one(
          Lptr + (long long)b * mat_stride,
          n,
          nb,
          gemmex_m_min,
          refine_tail,
          wptr,
          lwork,
          iptr);
    }
  }

  dim3 block(16, 16);
  dim3 grid((n + 15) / 16, (n + 15) / 16, batch);
  zero_upper_kernel<<<grid, block>>>(Lptr, n, batch);
  CHECK_CUDA(cudaGetLastError());
  // No host device-sync (large2 speed path); torch/popcorn sync externally.
  return L;
}


"""

MEDIUM2_CPP_SRC = r"""
torch::Tensor cholesky_blocked_batched(torch::Tensor A, int64_t nb_in, int64_t mode_in);
torch::Tensor cholesky_tf32_blocked(
    torch::Tensor A,
    int64_t nb_in,
    int64_t gemmex_m_min_in,
    int64_t refine_tail_in);
"""

# Single extension (medium batched + large TF32) — 3 load_inline total like champ.
_mod_lib = load_inline(
    name="chol_hybrid_m2l3_v3",
    cpp_sources=[MEDIUM2_CPP_SRC],
    cuda_sources=[MEDIUM2_CUDA_SRC],
    functions=["cholesky_blocked_batched", "cholesky_tf32_blocked"],
    # large3 gemmex + medium2 batched
    verbose=False,
    extra_cuda_cflags=[
        "-O3",
        "-DLARGE3_GEMMEX_M_MIN=4096",
        "-DLARGE3_REFINE_TAIL=0",
        "-DLARGE3_TF32_MATH=1",
    ],
    extra_ldflags=["-lcublas", "-lcusolver"],
)


def _choose_nb_large(n: int) -> int:
    # large3: 2048@32k was best public 88.7ms; 512@16k; 256@4k
    if n >= 32768:
        return 2048
    if n >= 16384:
        return 512
    if n >= 4096:
        return 256
    return 128


def _choose_gemmex_m_min(n: int) -> int:
    if n >= 4096:
        return 4096 if n >= 16384 else 2048
    return 10**9


def _choose_nb_medium(n: int) -> int:
    # medium3 Modal: nb=256@512 edges 128 (~1% @640×512); same for 1024/2048
    if n >= 512:
        return 256
    return 64


def _use_tf32_blocked(batch: int, n: int) -> bool:
    if n >= 16384:
        return True
    if n >= 4096 and batch >= 2:
        return True
    return False


def _use_medium_blocked(batch: int, n: int) -> bool:
    """High-batch only. 2×2048 stays torch (Popcorn torch faster than medium)."""
    if n == 512 and batch >= 64:
        return True
    if n == 1024 and batch >= 16:
        return True
    if n == 2048 and batch >= 8:
        return True
    return False


_MEDIUM_POTF_MODE = 0


def custom_kernel(data: input_t) -> output_t:
    A = data if data.is_contiguous() else data.contiguous()
    batch, n, n2 = A.shape
    if n != n2:
        raise ValueError("expected square matrices")
    if n == 32:
        return _mod_small.cholesky_n32(A)
    if n == 64:
        return _mod_small.cholesky_n64(A)
    if n == 128:
        return _mod_n128.cholesky_n128(A)
    if n == 256:
        return _mod_n128.cholesky_n256(A)
    b_i, n_i = int(batch), int(n)
    if _use_medium_blocked(b_i, n_i):
        return _mod_lib.cholesky_blocked_batched(
            A, _choose_nb_medium(n_i), _MEDIUM_POTF_MODE
        )
    if _use_tf32_blocked(b_i, n_i):
        return _mod_lib.cholesky_tf32_blocked(
            A, _choose_nb_large(n_i), _choose_gemmex_m_min(n_i), 0
        )
    return torch.linalg.cholesky_ex(A, check_errors=False).L
scrolls · 1792 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