Skip to content
KernelIndex
Search⌘K

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
NVIDIA B200
4.09ms
#331 of 337
2026-07-28

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-memoryextern __shared__ float s_A[];
tile-k = 8template <int BM = 64, int BN = 64, int BK = 8>
tile-m = 64template <int BM = 64, int BN = 64, int BK = 8>
tile-n = 64template <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