Skip to content
KernelIndex
Search⌘K

submission 888263

Ziron · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-888263?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.40ms
#166 of 337
2026-07-19

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:89fb951d14eec1b0c11fe303f6254efaa2035da084992623786a89478951570a
license declaredunknown
license concludedunknown
authorsZiron
imported2026-08-26

Techniques

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

shared-memory__shared__ float tile[32][33]; // padded to avoid bank conflicts

Kernel source

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

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

CUDA_SRC = """
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cusolverDn.h>

static cusolverDnHandle_t g_handle = nullptr;
static cusolverDnParams_t g_params = nullptr;
static void* g_workspace = nullptr;
static size_t g_workspace_size = 0;
static int* g_info = nullptr;

// Pre-allocated buffers for SpotrfBatched
static float** g_Aarray = nullptr;
static int* g_info_batch = nullptr;
static int g_Aarray_size = 0;
static int g_info_batch_size = 0;

// Pre-allocated A_col buffer
static float* g_A_col = nullptr;
static size_t g_A_col_size = 0;

void ensure_resources(int batch_size, size_t needed_A_col) {
    if (g_handle == nullptr) cusolverDnCreate(&g_handle);
    if (g_params == nullptr) cusolverDnCreateParams(&g_params);
    if (g_info == nullptr) cudaMalloc(&g_info, sizeof(int));
    
    if (batch_size > g_Aarray_size) {
        if (g_Aarray) cudaFree(g_Aarray);
        cudaMalloc(&g_Aarray, batch_size * sizeof(float*));
        g_Aarray_size = batch_size;
    }
    if (batch_size > g_info_batch_size) {
        if (g_info_batch) cudaFree(g_info_batch);
        cudaMalloc(&g_info_batch, batch_size * sizeof(int));
        g_info_batch_size = batch_size;
    }
    
    if (needed_A_col > g_A_col_size) {
        if (g_A_col) cudaFree(g_A_col);
        cudaMalloc(&g_A_col, needed_A_col * sizeof(float));
        g_A_col_size = needed_A_col;
    }
}

void ensure_workspace(size_t needed) {
    if (needed > g_workspace_size) {
        if (g_workspace) cudaFree(g_workspace);
        cudaMalloc(&g_workspace, needed);
        g_workspace_size = needed;
    }
}

// 1D thread mapping which is much better for coalesced writes
__global__ void zero_upper_triangle_1d_kernel(float* A, int N, int batch_size) {
    int idx = blockIdx.x * blockDim.x + threadIdx.x;
    int N2 = N * N;
    int total = batch_size * N2;
    if (idx < total) {
        int local_idx = idx % N2;
        int r = local_idx / N;
        int c = local_idx % N;
        if (c > r) {
            A[idx] = 0.0f;
        }
    }
}

// Coalesced transpose of in to out (no zeroing)
__global__ void transpose_coalesced(const float* __restrict__ in, float* __restrict__ out, int N) {
    __shared__ float tile[32][33]; // padded to avoid bank conflicts
    
    int tile_i = blockIdx.y * 32;
    int tile_j = blockIdx.x * 32;
    int b = blockIdx.z;
    
    int thread_x = threadIdx.x;
    int thread_y = threadIdx.y;
    
    #pragma unroll
    for (int k = 0; k < 32; k += 8) {
        int r = thread_y + k;
        int c = thread_x;
        
        int global_r = tile_i + r;
        int global_c = tile_j + c;
        
        if (global_r < N && global_c < N) {
            tile[r][c] = in[(long long)b * N * N + global_r * N + global_c];
        } else {
            tile[r][c] = 0.0f;
        }
    }
    __syncthreads();
    
    #pragma unroll
    for (int k = 0; k < 32; k += 8) {
        int r_trans = thread_y + k;
        int c_trans = thread_x;
        
        int global_row = tile_j + r_trans;
        int global_col = tile_i + c_trans;
        
        if (global_row < N && global_col < N) {
            out[(long long)b * N * N + global_row * N + global_col] = tile[c_trans][r_trans];
        }
    }
}

// Coalesced transpose and zero upper tril from in to out
__global__ void transpose_tril_coalesced(const float* __restrict__ in, float* __restrict__ out, int N) {
    __shared__ float tile[32][33];
    
    int tile_i = blockIdx.y * 32;
    int tile_j = blockIdx.x * 32;
    int b = blockIdx.z;
    
    int thread_x = threadIdx.x;
    int thread_y = threadIdx.y;
    
    #pragma unroll
    for (int k = 0; k < 32; k += 8) {
        int r = thread_y + k;
        int c = thread_x;
        
        int global_r = tile_i + r;
        int global_c = tile_j + c;
        
        if (global_r < N && global_c < N) {
            tile[r][c] = in[(long long)b * N * N + global_r * N + global_c];
        } else {
            tile[r][c] = 0.0f;
        }
    }
    __syncthreads();
    
    #pragma unroll
    for (int k = 0; k < 32; k += 8) {
        int r_trans = thread_y + k;
        int c_trans = thread_x;
        
        int global_row = tile_j + r_trans;
        int global_col = tile_i + c_trans;
        
        if (global_row < N && global_col < N) {
            if (global_col > global_row) {
                out[(long long)b * N * N + global_row * N + global_col] = 0.0f;
            } else {
                out[(long long)b * N * N + global_row * N + global_col] = tile[c_trans][r_trans];
            }
        }
    }
}

// Custom warp-shuffle kernel for N=32
// Processes 8 matrices per block (256 threads)
// Uses padded shared memory (33 floats width) to eliminate bank conflicts.
__global__ void batched_cholesky_32_kernel(const float* __restrict__ A, float* __restrict__ L, int batch_size) {
    __shared__ float s_mem[8 * 1056]; // 8.448 KB for 8 warps (padded)

    int warp_id = threadIdx.x / 32;
    int lane_id = threadIdx.x % 32;
    int batch_idx = blockIdx.x * 8 + warp_id;

    if (batch_idx < batch_size) {
        #pragma unroll
        for (int k = 0; k < 32; ++k) {
            s_mem[warp_id * 1056 + k * 33 + lane_id] = A[batch_idx * 1024 + k * 32 + lane_id];
        }
        __syncwarp();

        float row[32];
        #pragma unroll
        for (int col = 0; col < 32; ++col) {
            row[col] = s_mem[warp_id * 1056 + lane_id * 33 + col];
        }

        #pragma unroll
        for (int j = 0; j < 32; ++j) {
            if (lane_id == j) {
                row[j] = sqrtf(row[j]);
            }
            float diag = __shfl_sync(0xffffffff, row[j], j);
            if (lane_id > j) {
                row[j] /= diag;
            }
            #pragma unroll
            for (int k = j + 1; k < 32; ++k) {
                float lk_j = __shfl_sync(0xffffffff, row[j], k);
                if (lane_id >= k) {
                    row[k] -= row[j] * lk_j;
                }
            }
        }

        #pragma unroll
        for (int col = 0; col < 32; ++col) {
            if (col <= lane_id) {
                s_mem[warp_id * 1056 + lane_id * 33 + col] = row[col];
            } else {
                s_mem[warp_id * 1056 + lane_id * 33 + col] = 0.0f;
            }
        }
        __syncwarp();

        #pragma unroll
        for (int k = 0; k < 32; ++k) {
            L[batch_idx * 1024 + k * 32 + lane_id] = s_mem[warp_id * 1056 + k * 33 + lane_id];
        }
    }
}

// Register warp-shuffle kernel for N=64
// Processes 2 matrices per block (64 threads)
// Uses padded shared memory (65 floats width) to eliminate bank conflicts.
__global__ void batched_cholesky_64_kernel(const float* __restrict__ A, float* __restrict__ L, int batch_size) {
    __shared__ float s_mem[2 * 4160]; // 33.28 KB (padded)

    int warp_id = threadIdx.x / 32; // 0 or 1
    int lane_id = threadIdx.x % 32; // 0 to 31
    int batch_idx = blockIdx.x * 2 + warp_id;

    if (batch_idx < batch_size) {
        #pragma unroll
        for (int k = 0; k < 128; ++k) {
            int offset = k * 32 + lane_id;
            int r = offset / 64;
            int c = offset % 64;
            s_mem[warp_id * 4160 + r * 65 + c] = A[batch_idx * 4096 + offset];
        }
        __syncwarp();

        float r0[64];
        float r1[64];
        #pragma unroll
        for (int col = 0; col < 64; ++col) {
            r0[col] = s_mem[warp_id * 4160 + lane_id * 65 + col];
            r1[col] = s_mem[warp_id * 4160 + (lane_id + 32) * 65 + col];
        }

        #pragma unroll
        for (int j = 0; j < 64; ++j) {
            if (j < 32) {
                if (lane_id == j) {
                    r0[j] = sqrtf(r0[j]);
                }
            } else {
                if (lane_id == j - 32) {
                    r1[j] = sqrtf(r1[j]);
                }
            }
            
            float diag = __shfl_sync(0xffffffff, (j < 32) ? r0[j] : r1[j], (j < 32) ? j : (j - 32));
            
            if (j < 32) {
                if (lane_id > j) {
                    r0[j] /= diag;
                }
                r1[j] /= diag;
            } else {
                if (lane_id > j - 32) {
                    r1[j] /= diag;
                }
            }

            #pragma unroll
            for (int k = j + 1; k < 64; ++k) {
                float lk_j = __shfl_sync(0xffffffff, (k < 32) ? r0[j] : r1[j], (k < 32) ? k : (k - 32));
                if (k < 32) {
                    if (lane_id >= k) {
                        r0[k] -= r0[j] * lk_j;
                    }
                }
                if (lane_id + 32 >= k) {
                    r1[k] -= r1[j] * lk_j;
                }
            }
        }

        #pragma unroll
        for (int col = 0; col < 64; ++col) {
            s_mem[warp_id * 4160 + lane_id * 65 + col] = (col <= lane_id) ? r0[col] : 0.0f;
            s_mem[warp_id * 4160 + (lane_id + 32) * 65 + col] = (col <= lane_id + 32) ? r1[col] : 0.0f;
        }
        __syncwarp();

        #pragma unroll
        for (int k = 0; k < 128; ++k) {
            int offset = k * 32 + lane_id;
            int r = offset / 64;
            int c = offset % 64;
            L[batch_idx * 4096 + offset] = s_mem[warp_id * 4160 + r * 65 + c];
        }
    }
}

// Templated full-shared-memory blocked Cholesky solver
// Uses dynamic thread count (nthreads = blockDim.x) for best occupancy.
// For M=4 (N=128): use 512 threads for best SYRK parallelism.
template <int M>
__global__ void blocked_cholesky_shared_kernel(const float* __restrict__ A, float* __restrict__ L, int batch_size) {
    int batch_idx = blockIdx.x;
    if (batch_idx >= batch_size) return;

    extern __shared__ float s_mem[];

    int tid = threadIdx.x;
    int nthreads = blockDim.x;
    const int num_lower_blocks = M * (M + 1) / 2;
    const int total_elements = num_lower_blocks * 1024;
    const int n = M * 32;

    // Load lower triangle of A to s_mem
    for (int idx = tid; idx < total_elements; idx += nthreads) {
        int t_block = idx >> 10;
        int elem_idx = idx & 1023;
        int r = elem_idx >> 5;
        int c = elem_idx & 31;
        int i = 0, j = 0;
        int accum = 0;
        #pragma unroll
        for (int k = 0; k < M; ++k) {
            int next_accum = accum + (k + 1);
            if (t_block < next_accum) {
                i = k;
                j = t_block - accum;
                break;
            }
            accum = next_accum;
        }
        int block_offset = t_block * 1056;
        s_mem[block_offset + r * 33 + c] = A[batch_idx * n * n + (i * 32 + r) * n + (j * 32 + c)];
    }
    __syncthreads();

    // Blocked Cholesky
    #pragma unroll
    for (int J = 0; J < M; ++J) {
        // Step 1: Factor diagonal block with warp 0
        if (tid < 32) {
            int block_offset = (J * (J + 1) / 2 + J) * 1056;
            float row[32];
            #pragma unroll
            for (int col = 0; col < 32; ++col) row[col] = s_mem[block_offset + tid * 33 + col];
            #pragma unroll
            for (int j_inner = 0; j_inner < 32; ++j_inner) {
                if (tid == j_inner) row[j_inner] = sqrtf(row[j_inner]);
                float diag = __shfl_sync(0xffffffff, row[j_inner], j_inner);
                if (tid > j_inner) row[j_inner] /= diag;
                #pragma unroll
                for (int k = j_inner + 1; k < 32; ++k) {
                    float lk_j = __shfl_sync(0xffffffff, row[j_inner], k);
                    if (tid >= k) row[k] -= row[j_inner] * lk_j;
                }
            }
            #pragma unroll
            for (int col = 0; col < 32; ++col) s_mem[block_offset + tid * 33 + col] = (col <= tid) ? row[col] : 0.0f;
        }
        __syncthreads();

        // Step 2: Solve column panel (TRSM)
        const int num_col_blocks = M - J - 1;
        const int num_rows = num_col_blocks * 32;
        if (tid < num_rows) {
            int block_idx = tid >> 5;
            int r = tid & 31;
            int I = J + 1 + block_idx;
            int block_offset = (I * (I + 1) / 2 + J) * 1056;
            int diag_offset = (J * (J + 1) / 2 + J) * 1056;
            float row[32];
            #pragma unroll
            for (int col = 0; col < 32; ++col) row[col] = s_mem[block_offset + r * 33 + col];
            #pragma unroll
            for (int c = 0; c < 32; ++c) {
                float sum = 0.0f;
                #pragma unroll
                for (int k = 0; k < c; ++k) sum += row[k] * s_mem[diag_offset + c * 33 + k];
                row[c] = (row[c] - sum) / s_mem[diag_offset + c * 33 + c];
            }
            #pragma unroll
            for (int col = 0; col < 32; ++col) s_mem[block_offset + r * 33 + col] = row[col];
        }
        __syncthreads();

        // Step 3: Update trailing submatrix (SYRK) - uses all threads for max parallelism
        const int num_trailing_blocks = num_col_blocks * (num_col_blocks + 1) / 2;
        const int total_trailing_elements = num_trailing_blocks * 1024;
        for (int idx = tid; idx < total_trailing_elements; idx += nthreads) {
            int t_block = idx >> 10;
            int elem_idx = idx & 1023;
            int r = elem_idx >> 5;
            int c = elem_idx & 31;
            int I_block = 0, K_block = 0;
            int accum = 0;
            #pragma unroll
            for (int k = 0; k < M; ++k) {
                if (k >= num_col_blocks) break;
                int next_accum = accum + (num_col_blocks - k);
                if (t_block < next_accum) {
                    K_block = k;
                    I_block = k + (t_block - accum);
                    break;
                }
                accum = next_accum;
            }
            int I = J + 1 + I_block;
            int K = J + 1 + K_block;
            int block_offset = (I * (I + 1) / 2 + K) * 1056;
            int I_col_offset = (I * (I + 1) / 2 + J) * 1056;
            int K_col_offset = (K * (K + 1) / 2 + J) * 1056;
            float sum = 0.0f;
            #pragma unroll
            for (int k = 0; k < 32; ++k) sum += s_mem[I_col_offset + r * 33 + k] * s_mem[K_col_offset + c * 33 + k];
            s_mem[block_offset + r * 33 + c] -= sum;
        }
        __syncthreads();
    }

    // Store lower triangular L to output
    for (int idx = tid; idx < n * n; idx += nthreads) {
        int r = idx / n;
        int c = idx % n;
        if (c <= r) {
            int block_r = r >> 5;
            int block_c = c >> 5;
            int elem_r = r & 31;
            int elem_c = c & 31;
            int block_idx = block_r * (block_r + 1) / 2 + block_c;
            L[batch_idx * n * n + idx] = s_mem[block_idx * 1056 + elem_r * 33 + elem_c];
        } else {
            L[batch_idx * n * n + idx] = 0.0f;
        }
    }
}

// Fill pointer array for batched cuSOLVER
__global__ void fill_ptrs_kernel(float** ptrs, float* base, int stride, int count) {
    int idx = blockIdx.x * blockDim.x + threadIdx.x;
    if (idx < count) ptrs[idx] = base + (long long)idx * stride;
}

torch::Tensor batched_cholesky(torch::Tensor A) {
    int batch_size = A.size(0);
    int N = A.size(1);

    if (N == 32) {
        auto L = torch::empty_like(A);
        const int threads = 256;
        const int blocks = (batch_size + 7) / 8;
        batched_cholesky_32_kernel<<<blocks, threads>>>(A.data_ptr<float>(), L.data_ptr<float>(), batch_size);
        return L;
    } 
    if (N == 64) {
        auto L = torch::empty_like(A);
        const int threads = 64;
        const int blocks = (batch_size + 1) / 2;
        batched_cholesky_64_kernel<<<blocks, threads>>>(A.data_ptr<float>(), L.data_ptr<float>(), batch_size);
        return L;
    } 
    if (N == 128) {
        int shmem_bytes = 10 * 1056 * sizeof(float);
        auto L = torch::empty_like(A);
        cudaFuncSetAttribute(blocked_cholesky_shared_kernel<4>, cudaFuncAttributeMaxDynamicSharedMemorySize, shmem_bytes);
        blocked_cholesky_shared_kernel<4><<<batch_size, 512, shmem_bytes>>>(A.data_ptr<float>(), L.data_ptr<float>(), batch_size);
        return L;
    }
    if (N == 256) {
        // N=256 batch=64 is slightly faster with our custom UPPER SpotrfBatched + custom in-place zeroing.
        ensure_resources(batch_size, (size_t)batch_size * N * N);
        
        auto L = A.clone();
        float* base_ptr = L.data_ptr<float>();
        
        fill_ptrs_kernel<<<(batch_size + 255) / 256, 256>>>(g_Aarray, base_ptr, N * N, batch_size);
        
        cusolverDnSpotrfBatched(
            g_handle,
            CUBLAS_FILL_MODE_UPPER,
            N,
            g_Aarray,
            N,
            g_info_batch,
            batch_size
        );
        
        int total = batch_size * N * N;
        zero_upper_triangle_1d_kernel<<<(total + 255) / 256, 256>>>(base_ptr, N, batch_size);
        return L;
    }
    
    // Dispatch for N >= 512
    if (N < 2048) {
        // N=512, 1024
        // These sizes use cuSOLVER's batched potrf internally in ATen, which is already optimal.
        auto result = at::linalg_cholesky_ex(A, false, false);
        return std::get<0>(result);
    } else {
        // N >= 2048 (Ultimate coalesced transpose-based pipeline)
        size_t elements = (size_t)batch_size * N * N;
        ensure_resources(batch_size, elements);
        
        // Coalesced transpose of A to g_A_col
        dim3 block(32, 8);
        dim3 grid((N + 31)/32, (N + 31)/32, batch_size);
        transpose_coalesced<<<grid, block>>>(A.data_ptr<float>(), g_A_col, N);
        
        size_t workspaceInBytes = 0;
        size_t testWorkspaceInBytes = 0;
        cusolverDnXpotrf_bufferSize(
            g_handle, g_params, CUBLAS_FILL_MODE_LOWER, N,
            CUDA_R_32F, g_A_col, N, CUDA_R_32F,
            &workspaceInBytes, &testWorkspaceInBytes
        );
        ensure_workspace(workspaceInBytes);
        
        for (int i = 0; i < batch_size; i++) {
            cusolverDnXpotrf(
                g_handle, g_params, CUBLAS_FILL_MODE_LOWER, N,
                CUDA_R_32F, g_A_col + (long long)i * N * N, N,
                CUDA_R_32F, g_workspace, workspaceInBytes, nullptr, 0, g_info
            );
        }
        
        // Coalesced transpose + tril from g_A_col to output L
        auto L = torch::empty_like(A);
        transpose_tril_coalesced<<<grid, block>>>(g_A_col, L.data_ptr<float>(), N);
        
        return L;
    }
}
"""

CPP_SRC = """
torch::Tensor batched_cholesky(torch::Tensor A);
"""

# Compile the inline extension
module = load_inline(
    name='my_cholesky_submission',
    cpp_sources=[CPP_SRC],
    cuda_sources=[CUDA_SRC],
    functions=['batched_cholesky'],
    extra_cuda_cflags=['-O3', '--use_fast_math'],
    extra_ldflags=['-lcusolver'],
    verbose=False,
)

def custom_kernel(data: input_t) -> output_t:
    return module.batched_cholesky(data)
scrolls · 556 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