submission 921444
Vadym Doroshenko · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 338 lines, June 9 Researcher Reciprocity License v1.0.
submission2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-921444?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
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:6cfc486e8489616e8bd79d1121e6a27dc15a00011156b2f4e930c891d5657cf0
license declaredunknown
license concludedunknown
authorsVadym Doroshenko
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
extern __shared__ float s_A[];tile-k = 8
template <int BM = 64, int BN = 64, int BK = 8>tile-m = 64
template <int BM = 64, int BN = 64, int BK = 8>tile-n = 64
template <int BM = 64, int BN = 64, int BK = 8>Kernel source
submission2.py338 lines
import torch
from torch.utils.cpp_extension import load_inline
# ---------------------------------------------------------
# 1. CUDA Kernel Source Code
# ---------------------------------------------------------
cuda_source = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <math.h>
// ============================================================================
// CASE 1: BATCHED UNBLOCKED KERNEL (For N <= 64)
// ============================================================================
__global__ void batched_unblocked_potrf_kernel(float* A, int N) {
extern __shared__ float s_A[];
int batch_idx = blockIdx.x;
int tx = threadIdx.x;
float* A_b = A + batch_idx * N * N;
int stride = N + 1; // Padded for bank conflicts
for (int r = 0; r < N; ++r) {
if (tx < N) s_A[r * stride + tx] = A_b[r * N + tx];
}
__syncthreads();
for (int j = 0; j < N; ++j) {
if (tx == 0) s_A[j * stride + j] = sqrtf(s_A[j * stride + j]);
__syncthreads();
if (tx > j && tx < N) s_A[tx * stride + j] /= s_A[j * stride + j];
__syncthreads();
if (tx > j && tx < N) {
float L_ij = s_A[tx * stride + j];
for (int k = j + 1; k <= tx; ++k) {
s_A[tx * stride + k] -= L_ij * s_A[k * stride + j];
}
}
__syncthreads();
}
for (int r = 0; r < N; ++r) {
if (r >= tx && tx < N) A_b[r * N + tx] = s_A[r * stride + tx];
}
}
// ============================================================================
// CASE 2: BATCHED BLOCKED KERNELS (For 64 < N <= 2048)
// ============================================================================
template <int BLOCK_SIZE>
__global__ void batched_potrf_kernel(float* A, int N, int k) {
__shared__ float s_A[BLOCK_SIZE][BLOCK_SIZE + 1];
int batch_idx = blockIdx.x;
float* A_b = A + batch_idx * N * N + k * N + k;
int tx = threadIdx.x;
for (int r = 0; r < BLOCK_SIZE; ++r) s_A[r][tx] = A_b[r * N + tx];
__syncthreads();
for (int j = 0; j < BLOCK_SIZE; ++j) {
if (tx == 0) s_A[j][j] = sqrtf(s_A[j][j]);
__syncthreads();
if (tx > j && tx < BLOCK_SIZE) s_A[tx][j] /= s_A[j][j];
__syncthreads();
if (tx > j && tx < BLOCK_SIZE) {
float L_ij = s_A[tx][j];
for (int step = j + 1; step <= tx; ++step) {
s_A[tx][step] -= L_ij * s_A[step][j];
}
}
__syncthreads();
}
for (int r = 0; r < BLOCK_SIZE; ++r) {
if (r >= tx) A_b[r * N + tx] = s_A[r][tx];
}
}
template <int BLOCK_SIZE>
__global__ void batched_invert_LT_kernel(const float* __restrict__ L, float* __restrict__ U_inv,
int N, int k, int ld_u) {
__shared__ float s_L[BLOCK_SIZE][BLOCK_SIZE + 1];
int batch_idx = blockIdx.x;
const float* L_b = L + batch_idx * N * N + k * N + k;
float* U_inv_b = U_inv + batch_idx * ld_u * ld_u;
int tx = threadIdx.x;
for (int r = 0; r < BLOCK_SIZE; ++r) {
s_L[r][tx] = (r >= tx) ? L_b[r * N + tx] : 0.0f;
}
__syncthreads();
for (int i = 0; i < BLOCK_SIZE; ++i) {
float L_ii = s_L[i][i];
float val = 0.0f;
if (tx < i) {
float sum = 0.0f;
for (int step = tx; step < i; ++step) {
sum += s_L[i][step] * s_L[step][tx];
}
val = -sum / L_ii;
} else if (tx == i) {
val = 1.0f / L_ii;
}
__syncthreads();
if (tx <= i) s_L[i][tx] = val;
__syncthreads();
}
for (int r = 0; r < BLOCK_SIZE; ++r) {
if (tx >= r) U_inv_b[r * ld_u + tx] = s_L[tx][r];
}
}
template <int BLOCK_ROWS = 8, int BLOCK_COLS = 64>
__global__ void batched_panel_update_kernel(float* __restrict__ A21,
const float* __restrict__ U_inv,
int M, int N, int k, int ld_u) {
__shared__ float s_A[BLOCK_ROWS][BLOCK_COLS];
int batch_idx = blockIdx.z;
float* A21_b = A21 + batch_idx * N * N + (k + BLOCK_COLS) * N + k;
const float* U_inv_b = U_inv + batch_idx * ld_u * ld_u;
int tx = threadIdx.x;
int ty = threadIdx.y;
int row = blockIdx.x * BLOCK_ROWS + ty;
if (row < M) s_A[ty][tx] = A21_b[row * N + tx];
__syncthreads();
if (row < M) {
float sum = 0.0f;
for (int step = 0; step <= tx; ++step) {
sum += s_A[ty][step] * U_inv_b[step * ld_u + tx];
}
A21_b[row * N + tx] = sum;
}
}
// ============================================================================
// REGISTER TILED SYRK KERNEL (4x4 Tiling per Thread)
// ============================================================================
// Block Dim: 16x16 threads (256 threads)
// Tile Output: 64x64
// Inner Loop Chunk (BK): 8
template <int BM = 64, int BN = 64, int BK = 8>
__global__ void batched_syrk_kernel(float* __restrict__ A22,
const float* __restrict__ L21,
int M, int b, int N, int k) {
// Early exit for strictly upper triangular tiles
if (blockIdx.x > blockIdx.y) return;
int batch_idx = blockIdx.z;
float* A22_b = A22 + batch_idx * N * N + (k + b) * N + (k + b);
const float* L21_b = L21 + batch_idx * N * N + (k + b) * N + k;
// Padded to prevent 4-way bank conflicts during transposed reads
__shared__ float s_A[BM][BK + 1];
__shared__ float s_B[BN][BK + 1];
int tx = threadIdx.x; // 0..15
int ty = threadIdx.y; // 0..15
int tid = ty * 16 + tx;
// Thread-local accumulation registers (4x4 grid per thread)
float c[4][4] = {{0.0f}};
float a_reg[4];
float b_reg[4];
// IO Mapping: 256 threads collaboratively load 64x8 chunks (512 elements)
// Each thread loads exactly 2 elements per inner loop iteration
int load_r = tid / BK; // 0..31
int load_c = tid % BK; // 0..7
for (int k_base = 0; k_base < b; k_base += BK) {
// 1. Load A (Row chunk of L21)
int g_rA1 = blockIdx.y * BM + load_r;
int g_rA2 = blockIdx.y * BM + load_r + 32;
int g_cA = k_base + load_c;
s_A[load_r][load_c] = (g_rA1 < M && g_cA < b) ? L21_b[g_rA1 * N + g_cA] : 0.0f;
s_A[load_r + 32][load_c] = (g_rA2 < M && g_cA < b) ? L21_b[g_rA2 * N + g_cA] : 0.0f;
// 2. Load B (Column chunk of L21 - effectively computing L21^T)
int g_rB1 = blockIdx.x * BN + load_r;
int g_rB2 = blockIdx.x * BN + load_r + 32;
int g_cB = k_base + load_c;
s_B[load_r][load_c] = (g_rB1 < M && g_cB < b) ? L21_b[g_rB1 * N + g_cB] : 0.0f;
s_B[load_r + 32][load_c] = (g_rB2 < M && g_cB < b) ? L21_b[g_rB2 * N + g_cB] : 0.0f;
__syncthreads();
// 3. Compute Outer Product strictly in registers
#pragma unroll
for (int kk = 0; kk < BK; ++kk) {
// Load 8 values from slow Shared Memory into fast Registers
#pragma unroll
for (int i = 0; i < 4; ++i) a_reg[i] = s_A[ty + i * 16][kk];
#pragma unroll
for (int j = 0; j < 4; ++j) b_reg[j] = s_B[tx + j * 16][kk];
// Execute 16 Fused Multiply-Adds (FMA) instantly
#pragma unroll
for (int i = 0; i < 4; ++i) {
#pragma unroll
for (int j = 0; j < 4; ++j) {
c[i][j] += a_reg[i] * b_reg[j];
}
}
}
__syncthreads();
}
// 4. Write exactly the lower-triangular portion of the 64x64 tile back to global memory
#pragma unroll
for (int i = 0; i < 4; ++i) {
#pragma unroll
for (int j = 0; j < 4; ++j) {
int global_r = blockIdx.y * BM + ty + i * 16;
int global_c = blockIdx.x * BN + tx + j * 16;
if (global_r < M && global_c < M && global_r >= global_c) {
A22_b[global_r * N + global_c] -= c[i][j];
}
}
}
}
__global__ void batched_zero_upper_kernel(float* A, int N) {
int batch_idx = blockIdx.z;
float* A_b = A + batch_idx * N * N;
int col = blockIdx.x * 32 + threadIdx.x;
int row = blockIdx.y * 32 + threadIdx.y;
if (blockIdx.x < blockIdx.y) return;
if (row < N && col < N && col > row) {
A_b[row * N + col] = 0.0f;
}
}
void cholesky_batched_forward(torch::Tensor A) {
TORCH_CHECK(A.is_cuda() && A.is_contiguous(), "Input must be CUDA & contiguous");
TORCH_CHECK(A.dim() == 3 && A.size(1) == A.size(2), "Input must be (B, N, N)");
int B = A.size(0);
int N = A.size(1);
float* d_A = A.data_ptr<float>();
const int BLOCK_SIZE = 64; // Ensures we fit in 48KB default shared memory
if (N <= 64) {
// CASE 1: Small Matrices (Unblocked)
int target_smem = N * (N + 1) * sizeof(float);
batched_unblocked_potrf_kernel<<<B, N, target_smem>>>(d_A, N);
}
else {
// CASE 2: Medium Matrices (Blocked with Register Tiling)
TORCH_CHECK(N % BLOCK_SIZE == 0, "N must be a multiple of 64");
float* d_U_workspace;
cudaMalloc(&d_U_workspace, B * BLOCK_SIZE * BLOCK_SIZE * sizeof(float));
for (int k = 0; k < N; k += BLOCK_SIZE) {
int trailing_size = N - k - BLOCK_SIZE;
batched_potrf_kernel<BLOCK_SIZE><<<B, BLOCK_SIZE>>>(d_A, N, k);
if (trailing_size > 0) {
batched_invert_LT_kernel<BLOCK_SIZE><<<B, BLOCK_SIZE>>>(
d_A, d_U_workspace, N, k, BLOCK_SIZE
);
// Panel Update Block Config
dim3 panel_block(BLOCK_SIZE, 8);
dim3 panel_grid((trailing_size + panel_block.y - 1) / panel_block.y, 1, B);
batched_panel_update_kernel<8, BLOCK_SIZE><<<panel_grid, panel_block>>>(
d_A, d_U_workspace, trailing_size, N, k, BLOCK_SIZE
);
// SYRK Grid Mapping
// Since each thread block processes a 64x64 output tile:
dim3 syrk_block(16, 16);
dim3 syrk_grid((trailing_size + 63) / 64, (trailing_size + 63) / 64, B);
batched_syrk_kernel<64, 64, 8><<<syrk_grid, syrk_block>>>(
d_A, d_A, trailing_size, BLOCK_SIZE, N, k
);
}
}
cudaFree(d_U_workspace);
}
dim3 zero_block(32, 32);
dim3 zero_grid((N + 31) / 32, (N + 31) / 32, B);
batched_zero_upper_kernel<<<zero_grid, zero_block>>>(d_A, N);
}
"""
# ---------------------------------------------------------
# 2. C++ Host API (Standard 48KB Limit Maintained)
# ---------------------------------------------------------
cpp_source = "void cholesky_batched_forward(torch::Tensor A);"
# Compile extension
batched_cholesky = load_inline(
name='batched_cholesky_v2',
cpp_sources=cpp_source,
cuda_sources=cuda_source,
functions=['cholesky_batched_forward'],
with_cuda=True,
extra_cflags=['-O3'],
extra_cuda_cflags=['-O3', '-lineinfo']
)
from task import input_t, output_t
def custom_kernel(data: input_t) -> output_t:
L = data.clone().contiguous().float().cuda()
batched_cholesky.cholesky_batched_forward(L)
return L
scrolls · 338 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