submission 489481
jackkhuu · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 219 lines, June 9 Researcher Reciprocity License v1.0.
prefixsum_py_H100_claude-opus-4.5_ka_submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-prefixsum-v2-489481?include=source"interfacepython
Compatibility
measured onNVIDIA H100
declared hardwareNVIDIA H100
architecturessm_90
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:5676d180d757d26c30a335cdcccb0b60873c014a4670443a4fc3aeaf94d6f1fd
license declaredunknown
license concludedunknown
authorsjackkhuu
imported2026-08-15
Kernel source
prefixsum_py_H100_claude-opus-4.5_ka_submission.py219 lines
import triton
import triton.language as tl
import torch
@triton.jit
def _prefix_sum_phase1_kernel(
input_ptr,
output_ptr,
block_sums_ptr,
n_elements,
BLOCK_SIZE: tl.constexpr,
):
"""
Phase 1: Compute local prefix sums within each block and store block totals.
Each block computes inclusive prefix sum for its elements.
"""
pid = tl.program_id(0)
block_start = pid * BLOCK_SIZE
# Load block elements
offsets = block_start + tl.arange(0, BLOCK_SIZE)
mask = offsets < n_elements
# Load input data
x = tl.load(input_ptr + offsets, mask=mask, other=0.0)
# Compute inclusive prefix sum using Hillis-Steele style algorithm
# This is an O(n log n) work algorithm but O(log n) depth
result = x
# Iterative doubling for prefix sum
# For BLOCK_SIZE elements, we need log2(BLOCK_SIZE) iterations
stride = 1
while stride < BLOCK_SIZE:
# Shift and add
shifted = tl.shift_left(result, stride)
result = result + shifted
stride = stride * 2
# Store local prefix sums
tl.store(output_ptr + offsets, result, mask=mask)
# Store the total sum of this block (last valid element)
# We need to find the last valid element in this block
last_idx = tl.minimum(block_start + BLOCK_SIZE - 1, n_elements - 1)
if block_start <= last_idx:
# The block sum is the prefix sum at the last valid position
last_offset = last_idx - block_start
block_sum = tl.sum(x, axis=0) # Total sum of elements in this block
tl.store(block_sums_ptr + pid, block_sum)
@triton.jit
def _inclusive_scan_kernel(
input_ptr,
output_ptr,
n_elements,
BLOCK_SIZE: tl.constexpr,
):
"""
Single-pass inclusive prefix sum for small arrays that fit in one block.
Uses sequential accumulation for correctness.
"""
pid = tl.program_id(0)
# For inclusive scan, we process sequentially within the block
# Each thread handles one element and accumulates from previous
offsets = tl.arange(0, BLOCK_SIZE)
mask = offsets < n_elements
# Load all elements
x = tl.load(input_ptr + offsets, mask=mask, other=0.0)
# Inclusive scan using cumsum pattern
# Kogge-Stone style parallel prefix sum
result = x
# Log2(BLOCK_SIZE) iterations
offset = 1
while offset < BLOCK_SIZE:
# For each element i, add element i-offset if it exists
shifted_result = tl.zeros((BLOCK_SIZE,), dtype=tl.float32)
shifted_mask = offsets >= offset
src_offsets = offsets - offset
src_mask = (src_offsets >= 0) & mask
# Manual shift: load from offset position
shifted_vals = tl.where(shifted_mask,
tl.load(input_ptr + src_offsets, mask=src_mask & shifted_mask, other=0.0),
0.0)
# Actually we need to work with result, not input
# This requires a different approach
offset = offset * 2
tl.store(output_ptr + offsets, result, mask=mask)
@triton.jit
def _sequential_scan_kernel(
input_ptr,
output_ptr,
n_elements,
BLOCK_SIZE: tl.constexpr,
):
"""
Compute inclusive prefix sum using parallel Blelloch-style scan.
"""
pid = tl.program_id(0)
block_start = pid * BLOCK_SIZE
offsets = block_start + tl.arange(0, BLOCK_SIZE)
mask = offsets < n_elements
# Load input
x = tl.load(input_ptr + offsets, mask=mask, other=0.0)
# Inclusive scan via cumsum
result = tl.cumsum(x, axis=0)
tl.store(output_ptr + offsets, result, mask=mask)
@triton.jit
def _add_block_prefix_kernel(
output_ptr,
block_prefix_ptr,
n_elements,
BLOCK_SIZE: tl.constexpr,
):
"""
Add the prefix sum of previous blocks to each block's elements.
"""
pid = tl.program_id(0)
if pid == 0:
return # First block doesn't need adjustment
block_start = pid * BLOCK_SIZE
offsets = block_start + tl.arange(0, BLOCK_SIZE)
mask = offsets < n_elements
# Load the prefix sum to add (sum of all previous blocks)
prefix = tl.load(block_prefix_ptr + pid - 1)
# Load current values and add prefix
current = tl.load(output_ptr + offsets, mask=mask)
result = current + prefix
tl.store(output_ptr + offsets, result, mask=mask)
def kernel_function(data: torch.Tensor, output: torch.Tensor) -> torch.Tensor:
"""
Computes inclusive prefix sum (cumulative sum) of input data.
This is a fused implementation that computes the complete prefix sum
using a multi-phase approach for larger arrays:
1. Phase 1: Compute local prefix sums within each block
2. Phase 2: Compute prefix sums of block totals
3. Phase 3: Add block prefixes to get final result
For small arrays, a single-pass approach is used.
"""
n_elements = data.numel()
# Choose block size - power of 2
BLOCK_SIZE = 1024
# Number of blocks needed
n_blocks = triton.cdiv(n_elements, BLOCK_SIZE)
if n_blocks == 1:
# Single block - direct scan
grid = (1,)
_sequential_scan_kernel[grid](data, output, n_elements, BLOCK_SIZE)
else:
# Multi-block approach
# Phase 1: Local scans and block sums
block_sums = torch.empty(n_blocks, device=data.device, dtype=torch.float32)
grid = (n_blocks,)
_sequential_scan_kernel[grid](data, output, n_elements, BLOCK_SIZE)
# Compute block sums (sum of each block)
# We need to extract the last element of each block's prefix sum
# and compute their differences to get block totals, then prefix sum those
# Actually simpler: compute block sums directly
for i in range(n_blocks):
start = i * BLOCK_SIZE
end = min(start + BLOCK_SIZE, n_elements)
block_sums[i] = data[start:end].sum()
# Compute prefix sum of block sums
block_prefix = torch.cumsum(block_sums, dim=0)
# Phase 3: Add block prefixes
_add_block_prefix_kernel[grid](output, block_prefix, n_elements, BLOCK_SIZE)
return output
import inspect
def custom_kernel(input):
sig = inspect.signature(kernel_function)
num_params = len(sig.parameters)
if len(input) == num_params:
return kernel_function(*input)
return kernel_function(input)
# Ensure deterministic cuBLAS.
import os
if os.environ.get("CUBLAS_WORKSPACE_CONFIG", "") not in (":4096:8", ":16:8"):
os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8"
scrolls · 219 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