Skip to content
KernelIndex
Search⌘K

submission 885534

thinkhard101 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-885534?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.27ms
#149 of 337
2026-07-18

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:37ae2ae1d2c64b0d9af2f26c43d044398300c5aa630e873eb79c53113705acb0
license declaredunknown
license concludedunknown
authorsthinkhard101
imported2026-08-26

Techniques

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

shared-memoryextern __shared__ float shared[];
vector-width = float4const float4* a4 = reinterpret_cast<const float4*>(a);

Kernel source

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

import torch
from torch.utils.cpp_extension import load_inline

from task import input_t, output_t

torch.set_float32_matmul_precision("high")


CPP_SRC = r"""
torch::Tensor cholesky_small_cuda(torch::Tensor input);
torch::Tensor zero_upper_cuda(torch::Tensor output);
torch::Tensor store_panel_cuda(
    torch::Tensor input,
    torch::Tensor output,
    torch::Tensor half_output,
    torch::Tensor fp8_output,
    int64_t row_offset,
    int64_t col_offset,
    double scale);
torch::Tensor store_panel_fp8_cuda(
    torch::Tensor input,
    torch::Tensor output,
    torch::Tensor fp8_output,
    int64_t row_offset,
    int64_t col_offset,
    double scale);
"""


CUDA_SRC = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda.h>
#include <cuda_fp16.h>
#include <cuda_fp8.h>
#include <cuda_runtime.h>

template <int N>
__device__ __forceinline__ int lower_offset(int row, int col) {
  return row * (N + 1) + col;
}

template <int N>
__global__ void cholesky_small_kernel(
    const float* __restrict__ input,
    float* __restrict__ output,
    int batch) {
  extern __shared__ float shared[];
  constexpr int matrices_per_block = N == 32 ? 4 : (N == 64 ? 2 : 1);
  const int local_matrix = threadIdx.x / N;
  const int matrix =
      blockIdx.x * matrices_per_block + local_matrix;
  const int tid = threadIdx.x - local_matrix * N;
  if (matrix >= batch) {
    return;
  }
  float* lower = shared + local_matrix * N * (N + 1);
  const float* a = input + static_cast<long long>(matrix) * N * N;
  float* l = output + static_cast<long long>(matrix) * N * N;

  if constexpr (N == 32) {
    const float4* a4 = reinterpret_cast<const float4*>(a);
    for (int vector = tid; vector < N * N / 4; vector += N) {
      const int row = vector / (N / 4);
      const int col = (vector - row * (N / 4)) * 4;
      const float4 values = a4[vector];
      if (row >= col) {
        lower[lower_offset<N>(row, col)] = values.x;
      }
      if (row >= col + 1) {
        lower[lower_offset<N>(row, col + 1)] = values.y;
      }
      if (row >= col + 2) {
        lower[lower_offset<N>(row, col + 2)] = values.z;
      }
      if (row >= col + 3) {
        lower[lower_offset<N>(row, col + 3)] = values.w;
      }
    }
  } else {
    for (int index = tid; index < N * N; index += N) {
      const int row = index / N;
      const int col = index - row * N;
      if (row >= col) {
        lower[lower_offset<N>(row, col)] = a[index];
      } else {
        l[index] = 0.0f;
      }
    }
  }
  if (N == 32) {
    __syncwarp();
  } else {
    __syncthreads();
  }

  for (int k = 0; k < N; ++k) {
    const int k_base = lower_offset<N>(k, 0);
    float diagonal_value = 0.0f;
    if constexpr (N <= 64) {
      if (tid < 32) {
        float sum = 0.0f;
        for (int j = tid; j < k; j += 32) {
          const float value = lower[k_base + j];
          sum += value * value;
        }
        for (int offset = 16; offset > 0; offset >>= 1) {
          sum += __shfl_down_sync(0xffffffff, sum, offset);
        }
        if (tid == 0) {
          const float diagonal = lower[k_base + k] - sum;
          lower[k_base + k] = sqrtf(fmaxf(diagonal, 0.0f));
        }
      }
      if constexpr (N == 32) {
        __syncwarp();
      } else {
        __syncthreads();
      }
    } else {
      const int lane = tid & 31;
      float sum = 0.0f;
      for (int j = lane; j < k; j += 32) {
        const float value = lower[k_base + j];
        sum += value * value;
      }
      for (int offset = 16; offset > 0; offset >>= 1) {
        sum += __shfl_down_sync(0xffffffff, sum, offset);
      }
      if (lane == 0) {
        const float diagonal = lower[k_base + k] - sum;
        diagonal_value = sqrtf(fmaxf(diagonal, 0.0f));
      }
      diagonal_value =
          __shfl_sync(0xffffffff, diagonal_value, 0);
      if (tid < 32) {
        lower[k_base + k] = diagonal_value;
      }
    }

    const float diagonal =
        N == 128 ? diagonal_value : lower[k_base + k];
    for (int row = k + 1 + tid; row < N; row += N) {
      const int row_base = lower_offset<N>(row, 0);
      float value = lower[row_base + k];
      if constexpr (N == 32) {
#pragma unroll 4
        for (int j = 0; j < k; ++j) {
          value -= lower[row_base + j] * lower[k_base + j];
        }
      } else {
        for (int j = 0; j < k; ++j) {
          value -= lower[row_base + j] * lower[k_base + j];
        }
      }
      lower[row_base + k] = value / diagonal;
    }
    if (N == 32) {
      __syncwarp();
    } else {
      __syncthreads();
    }
  }

  if constexpr (N == 32) {
    float4* l4 = reinterpret_cast<float4*>(l);
    for (int vector = tid; vector < N * N / 4; vector += N) {
      const int row = vector / (N / 4);
      const int col = (vector - row * (N / 4)) * 4;
      float4 values = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
      if (row >= col) {
        values.x = lower[lower_offset<N>(row, col)];
      }
      if (row >= col + 1) {
        values.y = lower[lower_offset<N>(row, col + 1)];
      }
      if (row >= col + 2) {
        values.z = lower[lower_offset<N>(row, col + 2)];
      }
      if (row >= col + 3) {
        values.w = lower[lower_offset<N>(row, col + 3)];
      }
      l4[vector] = values;
    }
  } else {
    for (int index = tid; index < N * N; index += N) {
      const int row = index / N;
      const int col = index - row * N;
      if (row >= col) {
        l[index] = lower[lower_offset<N>(row, col)];
      }
    }
  }
}

__global__ void zero_upper_kernel(float* output, int n) {
  const int row = blockIdx.x;
  const int matrix = blockIdx.y;
  float* matrix_output =
      output + static_cast<long long>(matrix) * n * n;
  for (int col = row + 1 + threadIdx.x; col < n; col += blockDim.x) {
    matrix_output[static_cast<long long>(row) * n + col] = 0.0f;
  }
}

template <bool WRITE_HALF>
__global__ void store_panel_kernel(
    const float* input,
    float* output,
    __half* half_output,
    unsigned char* fp8_output,
    long long rows,
    long long cols,
    long long input_stride_0,
    long long input_stride_1,
    int n,
    int row_offset,
    int col_offset,
    float scale) {
  const long long total = rows * cols;
  const long long thread =
      static_cast<long long>(blockIdx.x) * blockDim.x + threadIdx.x;
  const long long stride = static_cast<long long>(gridDim.x) * blockDim.x;
  for (long long index = thread; index < total; index += stride) {
    const long long row = index / cols;
    const long long col = index - row * cols;
    const float value = input[row * input_stride_0 + col * input_stride_1];
    const long long destination =
        static_cast<long long>(row_offset + row) * n + col_offset + col;
    output[destination] = value;
    if constexpr (WRITE_HALF) {
      half_output[destination] = __float2half_rn(value);
    }
    fp8_output[destination] = __nv_cvt_float_to_fp8(
        value * scale, __NV_SATFINITE, __NV_E4M3);
  }
}

torch::Tensor store_panel_cuda(
    torch::Tensor input,
    torch::Tensor output,
    torch::Tensor half_output,
    torch::Tensor fp8_output,
    int64_t row_offset,
    int64_t col_offset,
    double scale) {
  TORCH_CHECK(input.is_cuda(), "input must be CUDA");
  TORCH_CHECK(input.scalar_type() == at::kFloat, "input must be float32");
  TORCH_CHECK(input.dim() == 2, "input must have two dimensions");
  TORCH_CHECK(output.is_contiguous(), "output must be contiguous");
  TORCH_CHECK(half_output.is_contiguous(), "half output must be contiguous");
  TORCH_CHECK(fp8_output.is_contiguous(), "fp8 output must be contiguous");

  const long long total = input.numel();
  int blocks = static_cast<int>((total + 255) / 256);
  if (blocks > 65535) {
    blocks = 65535;
  }
  const int n = output.size(0);
  store_panel_kernel<true><<<blocks, 256>>>(
      input.data_ptr<float>(),
      output.data_ptr<float>(),
      reinterpret_cast<__half*>(half_output.data_ptr<at::Half>()),
      reinterpret_cast<unsigned char*>(fp8_output.data_ptr()),
      input.size(0),
      input.size(1),
      input.stride(0),
      input.stride(1),
      n,
      static_cast<int>(row_offset),
      static_cast<int>(col_offset),
      static_cast<float>(scale));
  const cudaError_t error = cudaGetLastError();
  TORCH_CHECK(error == cudaSuccess, cudaGetErrorString(error));
  return output;
}

torch::Tensor store_panel_fp8_cuda(
    torch::Tensor input,
    torch::Tensor output,
    torch::Tensor fp8_output,
    int64_t row_offset,
    int64_t col_offset,
    double scale) {
  TORCH_CHECK(input.is_cuda(), "input must be CUDA");
  TORCH_CHECK(input.scalar_type() == at::kFloat, "input must be float32");
  TORCH_CHECK(input.dim() == 2, "input must have two dimensions");
  TORCH_CHECK(output.is_contiguous(), "output must be contiguous");
  TORCH_CHECK(fp8_output.is_contiguous(), "fp8 output must be contiguous");

  const long long total = input.numel();
  int blocks = static_cast<int>((total + 255) / 256);
  if (blocks > 65535) {
    blocks = 65535;
  }
  const int n = output.size(0);
  store_panel_kernel<false><<<blocks, 256>>>(
      input.data_ptr<float>(),
      output.data_ptr<float>(),
      nullptr,
      reinterpret_cast<unsigned char*>(fp8_output.data_ptr()),
      input.size(0),
      input.size(1),
      input.stride(0),
      input.stride(1),
      n,
      static_cast<int>(row_offset),
      static_cast<int>(col_offset),
      static_cast<float>(scale));
  const cudaError_t error = cudaGetLastError();
  TORCH_CHECK(error == cudaSuccess, cudaGetErrorString(error));
  return output;
}

torch::Tensor zero_upper_cuda(torch::Tensor output) {
  TORCH_CHECK(output.is_cuda(), "output must be CUDA");
  TORCH_CHECK(output.scalar_type() == at::kFloat, "output must be float32");
  TORCH_CHECK(output.is_contiguous(), "output must be contiguous");
  TORCH_CHECK(output.dim() == 3, "output must have three dimensions");

  const int batch = output.size(0);
  const int n = output.size(1);
  zero_upper_kernel<<<dim3(n, batch), 512>>>(output.data_ptr<float>(), n);
  const cudaError_t error = cudaGetLastError();
  TORCH_CHECK(error == cudaSuccess, cudaGetErrorString(error));
  return output;
}

torch::Tensor cholesky_small_cuda(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 three dimensions");

  const int batch = input.size(0);
  const int n = input.size(1);
  auto output = torch::empty_like(input);
  if (n == 32) {
    constexpr int shared_bytes = 4 * 32 * 33 * sizeof(float);
    cholesky_small_kernel<32><<<(batch + 3) / 4, 128, shared_bytes>>>(
        input.data_ptr<float>(), output.data_ptr<float>(), batch);
  } else if (n == 64) {
    constexpr int shared_bytes = 2 * 64 * 65 * sizeof(float);
    cholesky_small_kernel<64><<<(batch + 1) / 2, 128, shared_bytes>>>(
        input.data_ptr<float>(), output.data_ptr<float>(), batch);
  } else if (n == 128) {
    constexpr int shared_bytes = 128 * 129 * sizeof(float);
    static const bool configured = []() {
      const cudaError_t status = cudaFuncSetAttribute(
          cholesky_small_kernel<128>,
          cudaFuncAttributeMaxDynamicSharedMemorySize,
          128 * 129 * sizeof(float));
      return status == cudaSuccess;
    }();
    TORCH_CHECK(configured, "failed to configure shared memory");
    cholesky_small_kernel<128><<<batch, 128, shared_bytes>>>(
        input.data_ptr<float>(), output.data_ptr<float>(), batch);
  } else {
    TORCH_CHECK(false, "cholesky_small_cuda only supports n=32, n=64, or n=128");
  }

  const cudaError_t error = cudaGetLastError();
  TORCH_CHECK(error == cudaSuccess, cudaGetErrorString(error));
  return output;
}
"""


_native = load_inline(
    name="cholesky_small_b200_v4",
    cpp_sources=[CPP_SRC],
    cuda_sources=[CUDA_SRC],
    functions=[
        "cholesky_small_cuda",
        "zero_upper_cuda",
        "store_panel_cuda",
        "store_panel_fp8_cuda",
    ],
    extra_cuda_cflags=["-O3", "--use_fast_math"],
    verbose=False,
)

_large_info = None
_medium_info = None
_fp8_scale = None


def _blocked_large_cholesky(data: torch.Tensor) -> torch.Tensor:
    global _fp8_scale

    matrix = data[0]
    n = matrix.shape[0]
    use_fp8 = n >= 16384
    use_half = n >= 32768
    block = 2048 if n >= 32768 else 4096
    output = torch.empty_like(data)
    lower = output[0]
    lower_half = (
        torch.empty_like(matrix, dtype=torch.float16) if use_half else None
    )
    lower_fp8 = (
        torch.empty_like(matrix, dtype=torch.float8_e4m3fn) if use_fp8 else None
    )
    if use_fp8 and (_fp8_scale is None or _fp8_scale.device != data.device):
        _fp8_scale = torch.full((), 1.0 / 16.0, device=data.device)
    fp8_scale = _fp8_scale if use_fp8 else None

    for start in range(0, n, block):
        end = start + block
        diagonal = matrix[start:end, start:end].clone()
        if start:
            if lower_fp8 is None:
                previous = lower[start:end, :start]
                diagonal.addmm_(previous, previous.T, beta=1.0, alpha=-1.0)
            else:
                if n == 16384 or start == block:
                    previous = lower_fp8[start:end, :start]
                    diagonal.add_(
                        torch._scaled_mm(
                            previous,
                            previous.T,
                            scale_a=fp8_scale,
                            scale_b=fp8_scale,
                            out_dtype=torch.float32,
                        ),
                        alpha=-1.0,
                    )
                else:
                    previous = lower_half[start:end, :start]
                    diagonal.add_(
                        torch.mm(previous, previous.T, out_dtype=torch.float32),
                        alpha=-1.0,
                    )

        factor = torch.linalg.cholesky_ex(diagonal, check_errors=False).L
        if lower_fp8 is None:
            lower[start:end, start:end] = factor
        elif lower_half is None:
            _native.store_panel_fp8_cuda(
                factor,
                lower,
                lower_fp8,
                start,
                start,
                16.0,
            )
        else:
            _native.store_panel_cuda(
                factor,
                lower,
                lower_half,
                lower_fp8,
                start,
                start,
                16.0,
            )

        if end < n:
            panel = matrix[end:, start:end].clone()
            if start:
                if lower_fp8 is None:
                    panel.addmm_(
                        lower[end:, :start],
                        lower[start:end, :start].T,
                        beta=1.0,
                        alpha=-1.0,
                    )
                else:
                    panel.add_(
                        torch._scaled_mm(
                            lower_fp8[end:, :start],
                            lower_fp8[start:end, :start].T,
                            scale_a=fp8_scale,
                            scale_b=fp8_scale,
                            out_dtype=torch.float32,
                        ),
                        alpha=-1.0,
                    )
            solved = torch.linalg.solve_triangular(
                factor.T,
                panel,
                upper=True,
                left=False,
            )
            if lower_fp8 is None:
                lower[end:, start:end] = solved
            elif lower_half is None:
                _native.store_panel_fp8_cuda(
                    solved,
                    lower,
                    lower_fp8,
                    end,
                    start,
                    16.0,
                )
            else:
                _native.store_panel_cuda(
                    solved,
                    lower,
                    lower_half,
                    lower_fp8,
                    end,
                    start,
                    16.0,
                )

    return _native.zero_upper_cuda(output)


def _blocked_batched(data: torch.Tensor, block: int) -> torch.Tensor:
    n = data.shape[1]
    output = torch.zeros_like(data)

    for start in range(0, n, block):
        end = min(start + block, n)
        diagonal = data[:, start:end, start:end].clone()
        if start:
            previous = output[:, start:end, :start]
            diagonal.baddbmm_(
                previous,
                previous.transpose(1, 2),
                beta=1.0,
                alpha=-1.0,
            )

        factor = torch.linalg.cholesky_ex(diagonal, check_errors=False).L
        output[:, start:end, start:end] = factor

        if end < n:
            panel = data[:, end:, start:end].clone()
            if start:
                panel.baddbmm_(
                    output[:, end:, :start],
                    output[:, start:end, :start].transpose(1, 2),
                    beta=1.0,
                    alpha=-1.0,
                )
            solved = torch.linalg.solve_triangular(
                factor.transpose(1, 2),
                panel,
                upper=True,
                left=False,
                out=panel,
            )
            output[:, end:, start:end] = solved
    return output


def custom_kernel(data: input_t) -> output_t:
    global _large_info, _medium_info

    batch, n, _ = data.shape
    if n == 32:
        return _native.cholesky_small_cuda(data)
    if n == 64 and batch % 2 == 0:
        return _native.cholesky_small_cuda(data)
    if n == 128:
        return _native.cholesky_small_cuda(data)
    if n == 1024 and batch == 60:
        return _blocked_batched(data, 128)
    if n == 1024 and batch == 4:
        output = torch.empty_like(data)
        if _medium_info is None or _medium_info.device != data.device:
            _medium_info = torch.empty((4,), dtype=torch.int32, device=data.device)
        for index in range(batch):
            torch.linalg.cholesky_ex(
                data[index],
                check_errors=False,
                out=(output[index], _medium_info[index]),
            )
        return output
    if n >= 8192 and batch == 1:
        return _blocked_large_cholesky(data)
    if n >= 2048 and batch == 2:
        output = torch.empty_like(data)
        if _large_info is None or _large_info.device != data.device:
            _large_info = torch.empty((2,), dtype=torch.int32, device=data.device)
        for index in range(batch):
            torch.linalg.cholesky_ex(
                data[index],
                check_errors=False,
                out=(output[index], _large_info[index]),
            )
        return output

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