submission 526905
KernelAgent · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 167 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-526905?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:9ea105da368f5eecbb84082a9aa7ce6c731a46e49734ffde0f8a68a5fa7265fc
license declaredunknown
license concludedunknown
authorsKernelAgent
imported2026-08-15
Kernel source
prefixsum_v2_H100_claude-opus-4.5_ka_submission.py167 lines
import triton
import triton.language as tl
import torch
@triton.jit
def _prefix_sum_single_block_kernel(
x_ptr,
out_ptr,
n_elements,
BLOCK_SIZE: tl.constexpr,
):
"""
Single-block inclusive prefix sum kernel.
Uses sequential scan within a single program - works for small arrays.
"""
pid = tl.program_id(0)
# For single block approach, we process BLOCK_SIZE elements at a time
# and carry the running sum forward
# Process in chunks, maintaining running sum
running_sum = tl.zeros([1], dtype=tl.float32)
num_chunks = tl.cdiv(n_elements, BLOCK_SIZE)
for chunk_idx in range(0, num_chunks):
chunk_start = chunk_idx * BLOCK_SIZE
offsets = chunk_start + tl.arange(0, BLOCK_SIZE)
mask = offsets < n_elements
# Load chunk
x = tl.load(x_ptr + offsets, mask=mask, other=0.0)
# Compute inclusive scan within this chunk using associative scan
# We use a parallel scan pattern
scan_result = tl.cumsum(x, axis=0)
# Add running sum from previous chunks
result = scan_result + running_sum
# Store results
tl.store(out_ptr + offsets, result, mask=mask)
# Update running sum: it's the last valid element's result
# We need the sum of all elements in this chunk added to previous running_sum
chunk_sum = tl.sum(x, axis=0)
running_sum = running_sum + chunk_sum
@triton.jit
def _compute_block_sums_kernel(
x_ptr,
block_sums_ptr,
n_elements,
BLOCK_SIZE: tl.constexpr,
):
"""Compute the sum of each block."""
pid = tl.program_id(0)
block_start = pid * BLOCK_SIZE
offsets = block_start + tl.arange(0, BLOCK_SIZE)
mask = offsets < n_elements
x = tl.load(x_ptr + offsets, mask=mask, other=0.0)
block_sum = tl.sum(x, axis=0)
tl.store(block_sums_ptr + pid, block_sum)
@triton.jit
def _local_prefix_sum_kernel(
x_ptr,
out_ptr,
block_prefix_ptr, # Prefix sum of block sums (exclusive for each block)
n_elements,
BLOCK_SIZE: tl.constexpr,
):
"""Compute local prefix sum and add block prefix."""
pid = tl.program_id(0)
block_start = pid * BLOCK_SIZE
offsets = block_start + tl.arange(0, BLOCK_SIZE)
mask = offsets < n_elements
# Load data
x = tl.load(x_ptr + offsets, mask=mask, other=0.0)
# Compute local inclusive prefix sum
local_scan = tl.cumsum(x, axis=0)
# Load prefix from previous blocks (exclusive prefix for this block)
if pid > 0:
block_prefix = tl.load(block_prefix_ptr + pid - 1)
else:
block_prefix = 0.0
# Add block prefix to get global prefix sum
result = local_scan + block_prefix
tl.store(out_ptr + offsets, result, mask=mask)
def kernel_function(x: torch.Tensor, output: torch.Tensor) -> torch.Tensor:
"""
Compute inclusive prefix sum (cumulative sum) using Triton.
This is a fused implementation that computes the entire prefix sum
in Triton without any PyTorch compute operations.
For small arrays, uses a single-block sequential approach.
For larger arrays, uses a multi-pass algorithm:
1. Compute block sums
2. Compute prefix sum of block sums
3. Add block prefixes to local scans
"""
n_elements = x.numel()
# Choose block size
BLOCK_SIZE = 1024
if n_elements <= BLOCK_SIZE * 4:
# For small arrays, use single-block approach
# This processes sequentially but avoids multi-kernel overhead
grid = (1,)
_prefix_sum_single_block_kernel[grid](
x, output, n_elements, BLOCK_SIZE=BLOCK_SIZE
)
else:
# Multi-block approach for larger arrays
num_blocks = triton.cdiv(n_elements, BLOCK_SIZE)
# Step 1: Compute sum of each block
block_sums = torch.empty(num_blocks, device=x.device, dtype=torch.float32)
grid = (num_blocks,)
_compute_block_sums_kernel[grid](
x, block_sums, n_elements, BLOCK_SIZE=BLOCK_SIZE
)
# Step 2: Compute prefix sum of block sums (recursively or with single block)
block_prefix = torch.empty(num_blocks, device=x.device, dtype=torch.float32)
_prefix_sum_single_block_kernel[(1,)](
block_sums, block_prefix, num_blocks, BLOCK_SIZE=BLOCK_SIZE
)
# Step 3: Compute local prefix sums and add block prefixes
_local_prefix_sum_kernel[grid](
x, output, block_prefix, n_elements, BLOCK_SIZE=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 · 167 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 510644.
⋯ 5 unchanged lines@triton.jitdef _prefix_sum_single_block_kernel(x_ptr,- y_ptr,+ out_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.+ Single-block inclusive prefix sum kernel.+ Uses sequential scan within a single program - works for small arrays."""pid = tl.program_id(0)- # For single block, we process all elements- offsets = tl.arange(0, BLOCK_SIZE)- mask = offsets < n_elements+ # For single block approach, we process BLOCK_SIZE elements at a time+ # and carry the running sum forward- # Load all elements- x = tl.load(x_ptr + offsets, mask=mask, other=0.0)+ # Process in chunks, maintaining running sum+ running_sum = tl.zeros([1], dtype=tl.float32)- # Perform inclusive scan using Blelloch-style algorithm- # Up-sweep (reduce) phase followed by down-sweep- # Using associative scan pattern+ num_chunks = tl.cdiv(n_elements, BLOCK_SIZE)- # 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)+ for chunk_idx in range(0, num_chunks):+ chunk_start = chunk_idx * BLOCK_SIZE+ offsets = chunk_start + tl.arange(0, BLOCK_SIZE)+ mask = offsets < n_elements++ # Load chunk+ x = tl.load(x_ptr + offsets, mask=mask, other=0.0)++ # Compute inclusive scan within this chunk using associative scan+ # We use a parallel scan pattern+ scan_result = tl.cumsum(x, axis=0)++ # Add running sum from previous chunks+ result = scan_result + running_sum++ # Store results+ tl.store(out_ptr + offsets, result, mask=mask)++ # Update running sum: it's the last valid element's result+ # We need the sum of all elements in this chunk added to previous running_sum+ chunk_sum = tl.sum(x, axis=0)+ running_sum = running_sum + chunk_sum@triton.jit- def _local_prefix_sum_kernel(+ def _compute_block_sums_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.- """+ """Compute the sum of each block."""pid = tl.program_id(0)block_start = pid * BLOCK_SIZEoffsets = block_start + tl.arange(0, BLOCK_SIZE)mask = offsets < n_elements- # Load block elementsx = tl.load(x_ptr + offsets, mask=mask, other=0.0)+ block_sum = tl.sum(x, axis=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)+ tl.store(block_sums_ptr + pid, block_sum)- @triton.jit- def _add_block_prefix_kernel(- y_ptr,- block_prefix_ptr,+ @triton.jit+ def _local_prefix_sum_kernel(+ x_ptr,+ out_ptr,+ block_prefix_ptr, # Prefix sum of block sums (exclusive for each block)n_elements,BLOCK_SIZE: tl.constexpr,):- """- Phase 3: Add the prefix sum of previous blocks to each element.- """+ """Compute local prefix sum and add block prefix."""pid = tl.program_id(0)-- # Skip first block (it already has correct values)- if pid == 0:- return-block_start = pid * BLOCK_SIZEoffsets = 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 data+ x = tl.load(x_ptr + offsets, mask=mask, other=0.0)- # Load current values- y = tl.load(y_ptr + offsets, mask=mask)+ # Compute local inclusive prefix sum+ local_scan = tl.cumsum(x, axis=0)- # Add prefix and store- result = y + prefix_to_add- tl.store(y_ptr + offsets, result, mask=mask)+ # Load prefix from previous blocks (exclusive prefix for this block)+ if pid > 0:+ block_prefix = tl.load(block_prefix_ptr + pid - 1)+ else:+ block_prefix = 0.0++ # Add block prefix to get global prefix sum+ result = local_scan + block_prefix++ tl.store(out_ptr + offsets, result, mask=mask)- def kernel_function(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:+ def kernel_function(x: torch.Tensor, output: torch.Tensor) -> torch.Tensor:"""- Computes inclusive prefix sum (cumulative sum) of input tensor x.+ Compute inclusive prefix sum (cumulative sum) using Triton.- 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+ This is a fused implementation that computes the entire prefix sum+ in Triton without any PyTorch compute operations.- 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+ For small arrays, uses a single-block sequential approach.+ For larger arrays, uses a multi-pass algorithm:+ 1. Compute block sums+ 2. Compute prefix sum of block sums+ 3. Add block prefixes to local scans"""n_elements = x.numel()- # Choose block size based on input size+ # Choose block sizeBLOCK_SIZE = 1024- if n_elements <= BLOCK_SIZE:- # Single block can handle entire array+ if n_elements <= BLOCK_SIZE * 4:+ # For small arrays, use single-block approach+ # This processes sequentially but avoids multi-kernel overheadgrid = (1,)_prefix_sum_single_block_kernel[grid](- x, y, n_elements, BLOCK_SIZE=BLOCK_SIZE+ x, output, n_elements, BLOCK_SIZE=BLOCK_SIZE)else:- # Multi-block approach+ # Multi-block approach for larger arraysnum_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+ # Step 1: Compute sum of each block+ block_sums = torch.empty(num_blocks, device=x.device, dtype=torch.float32)grid = (num_blocks,)- _local_prefix_sum_kernel[grid](- x, y, block_sums, n_elements, BLOCK_SIZE=BLOCK_SIZE+ _compute_block_sums_kernel[grid](+ x, 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)+ # Step 2: Compute prefix sum of block sums (recursively or with single block)+ block_prefix = torch.empty(num_blocks, device=x.device, dtype=torch.float32)+ _prefix_sum_single_block_kernel[(1,)](+ block_sums, block_prefix, num_blocks, BLOCK_SIZE=BLOCK_SIZE+ )- # Phase 3: Add block prefixes to local results- _add_block_prefix_kernel[grid](- y, block_prefix, n_elements, BLOCK_SIZE=BLOCK_SIZE+ # Step 3: Compute local prefix sums and add block prefixes+ _local_prefix_sum_kernel[grid](+ x, output, block_prefix, n_elements, BLOCK_SIZE=BLOCK_SIZE)- return y+ return outputimport inspect
scrolls · 243 diff lines total
Best evidence level for this revision: reported
JSON