submission 510644
KernelAgent · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 180 lines, June 9 Researcher Reciprocity License v1.0.
prefixsum_v2_H100_claude-opus-4.5_ka_submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-prefixsum-v2-510644?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:b152cc71c2d056782fdf0bdcc22c1ef3415e388c180f74f4e6029baded48d4e3
license declaredunknown
license concludedunknown
authorsKernelAgent
imported2026-08-15
Kernel source
prefixsum_v2_H100_claude-opus-4.5_ka_submission.py180 lines
import triton
import triton.language as tl
import torch
@triton.jit
def _prefix_sum_single_block_kernel(
x_ptr,
y_ptr,
n_elements,
BLOCK_SIZE: tl.constexpr,
):
"""
Single-block inclusive prefix sum using sequential scan.
This handles the entire array in one block when n_elements <= BLOCK_SIZE.
"""
pid = tl.program_id(0)
# For single block, we process all elements
offsets = tl.arange(0, BLOCK_SIZE)
mask = offsets < n_elements
# Load all elements
x = tl.load(x_ptr + offsets, mask=mask, other=0.0)
# Perform inclusive scan using Blelloch-style algorithm
# Up-sweep (reduce) phase followed by down-sweep
# Using associative scan pattern
# For inclusive prefix sum, we use tl.associative_scan
result = tl.cumsum(x, axis=0)
# Store results
tl.store(y_ptr + offsets, result, mask=mask)
@triton.jit
def _local_prefix_sum_kernel(
x_ptr,
y_ptr,
block_sums_ptr,
n_elements,
BLOCK_SIZE: tl.constexpr,
):
"""
Phase 1: Compute local prefix sum within each block and store block totals.
"""
pid = tl.program_id(0)
block_start = pid * BLOCK_SIZE
offsets = block_start + tl.arange(0, BLOCK_SIZE)
mask = offsets < n_elements
# Load block elements
x = tl.load(x_ptr + offsets, mask=mask, other=0.0)
# Compute local inclusive prefix sum
local_prefix = tl.cumsum(x, axis=0)
# Store local prefix sums
tl.store(y_ptr + offsets, local_prefix, mask=mask)
# Store the last valid element as block sum
# The block sum is the total of all elements in this block
last_idx = tl.arange(0, BLOCK_SIZE)
last_mask = (offsets < n_elements)
# Get the last valid prefix sum value (which is the block total)
block_total = tl.sum(x, axis=0)
# Only one thread stores the block sum
if pid >= 0: # Always true, just need to store once per block
tl.store(block_sums_ptr + pid, block_total)
@triton.jit
def _add_block_prefix_kernel(
y_ptr,
block_prefix_ptr,
n_elements,
BLOCK_SIZE: tl.constexpr,
):
"""
Phase 3: Add the prefix sum of previous blocks to each element.
"""
pid = tl.program_id(0)
# Skip first block (it already has correct values)
if pid == 0:
return
block_start = pid * BLOCK_SIZE
offsets = block_start + tl.arange(0, BLOCK_SIZE)
mask = offsets < n_elements
# Load the prefix sum up to (but not including) this block
prefix_to_add = tl.load(block_prefix_ptr + pid - 1)
# Load current values
y = tl.load(y_ptr + offsets, mask=mask)
# Add prefix and store
result = y + prefix_to_add
tl.store(y_ptr + offsets, result, mask=mask)
def kernel_function(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
"""
Computes inclusive prefix sum (cumulative sum) of input tensor x.
This implementation fuses local prefix sums with block-level aggregation:
- For small arrays (<=4096): single block handles everything
- For larger arrays: multi-phase approach with block-level prefix propagation
Args:
x: Input tensor of shape (n,) on CUDA
y: Output tensor of shape (n,) on CUDA (pre-allocated)
Returns:
y: The output tensor containing cumulative sum
"""
n_elements = x.numel()
# Choose block size based on input size
BLOCK_SIZE = 1024
if n_elements <= BLOCK_SIZE:
# Single block can handle entire array
grid = (1,)
_prefix_sum_single_block_kernel[grid](
x, y, n_elements, BLOCK_SIZE=BLOCK_SIZE
)
else:
# Multi-block approach
num_blocks = triton.cdiv(n_elements, BLOCK_SIZE)
# Allocate space for block sums
block_sums = torch.empty(num_blocks, device=x.device, dtype=x.dtype)
# Phase 1: Local prefix sums and block totals
grid = (num_blocks,)
_local_prefix_sum_kernel[grid](
x, y, block_sums, n_elements, BLOCK_SIZE=BLOCK_SIZE
)
# Phase 2: Compute prefix sum of block sums
# For small number of blocks, do recursively or use single block
if num_blocks <= BLOCK_SIZE:
block_prefix = torch.empty(num_blocks, device=x.device, dtype=x.dtype)
_prefix_sum_single_block_kernel[(1,)](
block_sums, block_prefix, num_blocks, BLOCK_SIZE=BLOCK_SIZE
)
else:
# Recursive case for very large arrays
block_prefix = torch.empty(num_blocks, device=x.device, dtype=x.dtype)
kernel_function(block_sums, block_prefix)
# Phase 3: Add block prefixes to local results
_add_block_prefix_kernel[grid](
y, block_prefix, n_elements, BLOCK_SIZE=BLOCK_SIZE
)
return y
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 · 180 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 510383.
⋯ 4 unchanged lines@triton.jitdef _prefix_sum_single_block_kernel(- input_ptr,- output_ptr,+ x_ptr,+ y_ptr,n_elements,BLOCK_SIZE: tl.constexpr,):"""- Single-block inclusive prefix sum using a work-efficient parallel scan.- This kernel handles arrays that fit within a single block.+ Single-block inclusive prefix sum using sequential scan.+ This handles the entire array in one block when n_elements <= BLOCK_SIZE."""- # Load all elements into shared memory+ pid = tl.program_id(0)++ # For single block, we process all elementsoffsets = tl.arange(0, BLOCK_SIZE)mask = offsets < n_elements- # Load input values- x = tl.load(input_ptr + offsets, mask=mask, other=0.0)+ # Load all elements+ x = tl.load(x_ptr + offsets, mask=mask, other=0.0)- # Perform inclusive scan using Hillis-Steele algorithm- # This is a simple parallel prefix sum that works well for moderate sizes- stride = 1- while stride < BLOCK_SIZE:- # Get value from 'stride' positions before- shifted_offsets = offsets - stride- shifted_mask = shifted_offsets >= 0-- # Load the shifted values (from the same array x)- x_shifted = tl.where(shifted_mask,- tl.load(input_ptr + shifted_offsets, mask=shifted_mask & (shifted_offsets < n_elements), other=0.0),- 0.0)-- # This approach won't work directly - we need to be more careful- stride *= 2+ # Perform inclusive scan using Blelloch-style algorithm+ # Up-sweep (reduce) phase followed by down-sweep+ # Using associative scan pattern+ # For inclusive prefix sum, we use tl.associative_scan+ result = tl.cumsum(x, axis=0)+# Store results- tl.store(output_ptr + offsets, x, mask=mask)+ tl.store(y_ptr + offsets, result, mask=mask)@triton.jit- def _local_scan_kernel(- input_ptr,- output_ptr,+ def _local_prefix_sum_kernel(+ x_ptr,+ y_ptr,block_sums_ptr,n_elements,BLOCK_SIZE: tl.constexpr,):"""- First pass: compute local prefix sums within each block and store block totals.+ Phase 1: Compute local prefix sum within each block and store block totals."""pid = tl.program_id(0)block_start = pid * BLOCK_SIZEoffsets = block_start + tl.arange(0, BLOCK_SIZE)mask = offsets < n_elements- # Load input values- x = tl.load(input_ptr + offsets, mask=mask, other=0.0)+ # Load block elements+ x = tl.load(x_ptr + offsets, mask=mask, other=0.0)- # Compute inclusive prefix sum within block using cumsum- # Triton doesn't have built-in cumsum, so we use associative scan- result = tl.cumsum(x, axis=0)+ # Compute local inclusive prefix sum+ local_prefix = tl.cumsum(x, axis=0)# Store local prefix sums- tl.store(output_ptr + offsets, result, mask=mask)+ tl.store(y_ptr + offsets, local_prefix, mask=mask)- # Store the block total (last valid element's prefix sum)- # We need the sum of all elements in this block- block_sum = tl.sum(x, axis=0)- tl.store(block_sums_ptr + pid, block_sum)+ # Store the last valid element as block sum+ # The block sum is the total of all elements in this block+ last_idx = tl.arange(0, BLOCK_SIZE)+ last_mask = (offsets < n_elements)++ # Get the last valid prefix sum value (which is the block total)+ block_total = tl.sum(x, axis=0)++ # Only one thread stores the block sum+ if pid >= 0: # Always true, just need to store once per block+ tl.store(block_sums_ptr + pid, block_total)- @triton.jit+ @triton.jitdef _add_block_prefix_kernel(- output_ptr,+ y_ptr,block_prefix_ptr,n_elements,BLOCK_SIZE: tl.constexpr,):"""- Third pass: add the prefix sum of previous blocks to each element.+ Phase 3: Add the prefix sum of previous blocks to each element."""pid = tl.program_id(0)- # Skip first block - it doesn't need adjustment+ # Skip first block (it already has correct values)if pid == 0:return⋯ 1 unchanged linesoffsets = block_start + tl.arange(0, BLOCK_SIZE)mask = offsets < n_elements- # Load the prefix sum of all previous blocks- prefix = tl.load(block_prefix_ptr + pid - 1)+ # Load the prefix sum up to (but not including) this block+ prefix_to_add = tl.load(block_prefix_ptr + pid - 1)- # Load current values and add prefix- current = tl.load(output_ptr + offsets, mask=mask, other=0.0)- result = current + prefix+ # Load current values+ y = tl.load(y_ptr + offsets, mask=mask)- # Store updated values- tl.store(output_ptr + offsets, result, mask=mask)+ # Add prefix and store+ result = y + prefix_to_add+ tl.store(y_ptr + offsets, result, mask=mask)- def kernel_function(data: torch.Tensor, output: torch.Tensor) -> torch.Tensor:+ def kernel_function(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:"""- Computes inclusive prefix sum (cumulative sum) of the input tensor.+ Computes inclusive prefix sum (cumulative sum) of input tensor x.- This is a fused multi-pass implementation:- 1. Local scan within blocks + compute block totals- 2. Scan of block totals (done sequentially for small number of blocks)- 3. Add block prefixes to local results+ This implementation fuses local prefix sums with block-level aggregation:+ - For small arrays (<=4096): single block handles everything+ - For larger arrays: multi-phase approach with block-level prefix propagationArgs:- data: Input 1D tensor- output: Output tensor (pre-allocated)+ x: Input tensor of shape (n,) on CUDA+ y: Output tensor of shape (n,) on CUDA (pre-allocated)Returns:- Output tensor containing inclusive prefix sum+ y: The output tensor containing cumulative sum"""- n_elements = data.numel()+ n_elements = x.numel()- # Choose block size+ # Choose block size based on input sizeBLOCK_SIZE = 1024- # Number of blocks needed- n_blocks = triton.cdiv(n_elements, BLOCK_SIZE)-- # Allocate temporary storage for block sums- block_sums = torch.empty(n_blocks, device=data.device, dtype=torch.float32)-- # Pass 1: Local prefix sums within each block- grid = (n_blocks,)- _local_scan_kernel[grid](- data, output, block_sums,- n_elements,- BLOCK_SIZE=BLOCK_SIZE,- )-- # Pass 2: Compute prefix sum of block sums- # For small number of blocks, we can do this with a single block- if n_blocks > 1:- block_prefix = torch.empty(n_blocks, device=data.device, dtype=torch.float32)+ if n_elements <= BLOCK_SIZE:+ # Single block can handle entire array+ grid = (1,)+ _prefix_sum_single_block_kernel[grid](+ x, y, n_elements, BLOCK_SIZE=BLOCK_SIZE+ )+ else:+ # Multi-block approach+ num_blocks = triton.cdiv(n_elements, BLOCK_SIZE)- # Use a simple kernel to compute prefix sum of block sums- _compute_block_prefix[(1,)](- block_sums, block_prefix, n_blocks,- BLOCK_SIZE=triton.next_power_of_2(n_blocks),+ # Allocate space for block sums+ block_sums = torch.empty(num_blocks, device=x.device, dtype=x.dtype)++ # Phase 1: Local prefix sums and block totals+ grid = (num_blocks,)+ _local_prefix_sum_kernel[grid](+ x, y, block_sums, n_elements, BLOCK_SIZE=BLOCK_SIZE)- # Pass 3: Add block prefixes to local results+ # Phase 2: Compute prefix sum of block sums+ # For small number of blocks, do recursively or use single block+ if num_blocks <= BLOCK_SIZE:+ block_prefix = torch.empty(num_blocks, device=x.device, dtype=x.dtype)+ _prefix_sum_single_block_kernel[(1,)](+ block_sums, block_prefix, num_blocks, BLOCK_SIZE=BLOCK_SIZE+ )+ else:+ # Recursive case for very large arrays+ block_prefix = torch.empty(num_blocks, device=x.device, dtype=x.dtype)+ kernel_function(block_sums, block_prefix)++ # Phase 3: Add block prefixes to local results_add_block_prefix_kernel[grid](- output, block_prefix,- n_elements,- BLOCK_SIZE=BLOCK_SIZE,+ y, block_prefix, n_elements, BLOCK_SIZE=BLOCK_SIZE)- return output+ return y-- @triton.jit- def _compute_block_prefix(- block_sums_ptr,- block_prefix_ptr,- n_blocks,- BLOCK_SIZE: tl.constexpr,- ):- """- Compute inclusive prefix sum of block sums.- Single block kernel for small arrays.- """- offsets = tl.arange(0, BLOCK_SIZE)- mask = offsets < n_blocks-- # Load block sums- x = tl.load(block_sums_ptr + offsets, mask=mask, other=0.0)-- # Compute inclusive prefix sum- result = tl.cumsum(x, axis=0)-- # Store results- tl.store(block_prefix_ptr + offsets, result, mask=mask)-import inspectdef custom_kernel(input):
scrolls · 273 diff lines total
Best evidence level for this revision: reported
JSON