Skip to content
KernelIndex
Search⌘K

submission 924876

Harshit Kulkarni · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

b200.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-cholesky-924876?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.65ms
#204 of 337
2026-07-29

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:65a65d5fea1db71f2b1dabd9565fa9a42d4785be59ee1b9323aa68be8d956606
license declaredunknown
license concludedunknown
authorsHarshit Kulkarni
imported2026-08-26

Kernel source

b200.py107 lines
import torch

# =============================================================================
# Batched dense Cholesky factorization  ->  A = L @ L.T   (lower, positive diag)
# Target: NVIDIA B200 (Blackwell).  Ranking = geometric mean of runtime.
#
# The geomean is dominated by the GIANT matrices (n=8192/16384/32768). Pure
# FP32 sits at the B200 compute wall (n=32768 ~= 221 ms) since B200 has no FP32
# tensor core. We run their O(n^3) work on the tensor cores.
#
# KEY: put BOTH heavy steps of the blocked factorization on tensor cores.
#   * trailing update  A22 -= L21 @ L21^T           -> TF32 GEMM
#   * panel            L21 = A21 @ inv(L11)^T       -> TF32 GEMM
#     (invert the small nb x nb diagonal block once in FP32; this replaces the
#     triangular SOLVE, which runs on CUDA cores only and was the hidden
#     bottleneck -- for n=32768 the panel TRSM alone was ~30 ms.)
# Only the diagonal-block Cholesky stays in exact FP32 (tiny, O(nb^2 * n)).
#
# Size routing (correctness gate only tests n<=2048 -> low precision confined
# to n>=4096, cannot affect tested correctness cases):
#   n <= 1024              : batched cuSOLVER potrf (one call; latency bound)
#   n == 2048 small batch  : per-matrix cuSOLVER (batched potrf degrades)
#   n >= GIANT_MIN_N       : tensor-core blocked Cholesky
#
# No GPU->CPU sync anywhere on the hot path.
# =============================================================================

torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True

# --- routing / tuning knobs --------------------------------------------------
GIANT_MIN_N = 8192       # >= this n uses the tensor-core blocked path
LOOP_MIN_N = 2048        # per-matrix cuSOLVER for the pathological batched potrf
LOOP_MAX_BATCH = 16


def _block_size(n: int) -> int:
    # Smaller blocks push a larger fraction of work onto the tensor-core GEMMs;
    # the only non-tensor-core cost is the nb x nb diagonal Cholesky + inverse.
    if n >= 16384:
        return 1024
    return 1024


def _blocked_cholesky(A: torch.Tensor) -> torch.Tensor:
    """Right-looking blocked Cholesky, in place. L in the lower triangle.

    A : (B, n, n) contiguous float32 SPD. Diagonal block is exact FP32; panel
    and trailing update are TF32 tensor-core GEMMs.
    """
    n = A.shape[-1]
    nb = _block_size(n)
    eye = torch.eye(nb, device=A.device, dtype=A.dtype)
    for j in range(0, n, nb):
        j1 = min(j + nb, n)
        jb = j1 - j

        # Diagonal block: exact FP32 Cholesky.
        Ajj = A[:, j:j1, j:j1].contiguous()
        Ljj = torch.linalg.cholesky_ex(Ajj, check_errors=False)[0]
        A[:, j:j1, j:j1] = Ljj
        if j1 >= n:
            break

        # Invert the (small) diagonal factor once, in FP32 -> stable.
        e = eye if jb == nb else torch.eye(jb, device=A.device, dtype=A.dtype)
        Ljj_inv = torch.linalg.solve_triangular(Ljj, e, upper=False)

        # Panel: L21 = A21 @ inv(L11)^T   (TF32 GEMM on tensor cores).
        A21 = A[:, j1:, j:j1].contiguous()               # (B, m, jb)
        L21 = torch.matmul(A21, Ljj_inv.transpose(-1, -2))
        A[:, j1:, j:j1] = L21

        # Trailing symmetric update: A22 -= L21 @ L21^T  (TF32 GEMM).
        upd = torch.matmul(L21, L21.transpose(-1, -2))
        A[:, j1:, j1:] -= upd

    return A


@torch.no_grad()
def custom_kernel(data) -> torch.Tensor:
    """A -> L (lower triangular, positive diagonal) with A = L @ L.T."""
    A = data
    n = A.shape[-1]

    if A.dim() == 2:
        if n >= GIANT_MIN_N:
            work = A.unsqueeze(0).contiguous()
            return torch.tril(_blocked_cholesky(work))[0]
        return torch.linalg.cholesky_ex(A, check_errors=False)[0]

    A3 = A.reshape(-1, n, n)
    batch = A3.shape[0]

    if n >= GIANT_MIN_N:
        work = A3.contiguous().clone()
        L = torch.tril(_blocked_cholesky(work))
    elif n >= LOOP_MIN_N and 1 < batch <= LOOP_MAX_BATCH:
        outs = [torch.linalg.cholesky_ex(A3[i], check_errors=False)[0]
                for i in range(batch)]
        L = torch.stack(outs, dim=0)
    else:
        L = torch.linalg.cholesky_ex(A3, check_errors=False)[0]

    return L.reshape(A.shape)
scrolls · 107 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