Skip to content
KernelIndex
Search⌘K

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
Inclusive prefix sumsuite of 11 cases
NVIDIA H100
3.52s
#23 of 23
2026-02-12

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