Skip to content
KernelIndex
Search⌘K

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
Inclusive prefix sumsuite of 11 cases
NVIDIA H100
1.27ms
#17 of 23
2026-03-10

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.jit
def _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_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)
+ 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_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 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 size
BLOCK_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 overhead
grid = (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 arrays
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
+ # 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 output
import inspect
scrolls · 243 diff lines total

Best evidence level for this revision: reported

JSON