Skip to content
KernelIndex
Search⌘K

submission 150154

angel · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 169 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-sort-v2-150154?include=source"
interfacepython
Compatibility
measured onNVIDIA A100
declared hardwareNVIDIA A100
architecturessm_80
dtypesfp32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
Sortsuite of 5 cases
NVIDIA A100
15.1ms
#22 of 28
2025-12-12

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e9880e1ac56b7b6d2318e48c1966fc1f7f3647e072187174b748611b511c1717
license declaredunknown
license concludedunknown
authorsangel
imported2026-08-15

Kernel source

submission.py169 lines
import torch
import triton
import triton.language as tl


@triton.jit
def bitonic_step_kernel(
    data_ptr,
    n_elements,
    stage: tl.constexpr,
    step: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    """
    One step of bitonic sort.

    In bitonic sort, we have nested loops:
    - Outer loop (stage): builds sequences of size 2, 4, 8, ..., n
    - Inner loop (step): merges bitonic sequences with decreasing stride (jump in memory to move to next element)

    Each kernel call handles one (stage, step) combination.
    """
    pid = tl.program_id(0)
    idx = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)

    mask = idx < n_elements

    # Calculate partner index using XOR
    partner_idx = idx ^ step

    # Only process if BOTH indices are valid and we're the lower index
    both_valid = mask & (partner_idx < n_elements)
    should_process = both_valid & (idx < partner_idx)

    # Load our value and partner's value (only if both are valid)
    our_val = tl.load(data_ptr + idx, mask=should_process, other=0.0)
    partner_val = tl.load(data_ptr + partner_idx, mask=should_process, other=0.0)

    # Determine sort direction for this sequence
    # Sequences alternate between ascending and descending based on stage
    direction_bit = (idx >> stage) & 1
    ascending = direction_bit == 0

    # Determine if we need to swap
    need_swap = (ascending & (our_val > partner_val)) | (~ascending & (our_val < partner_val))

    # If we need to swap, exchange values
    new_our_val = tl.where(need_swap, partner_val, our_val)
    new_partner_val = tl.where(need_swap, our_val, partner_val)

    # Store back (only the lower index writes both values)
    tl.store(data_ptr + idx, new_our_val, mask=should_process)
    tl.store(data_ptr + partner_idx, new_partner_val, mask=should_process)


@triton.jit
def bitonic_final_step_kernel(
    data_ptr,
    n_elements,
    step: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
):
    """
    Final merge steps to ensure everything is sorted ascending.
    This is the last stage where we always sort ascending.
    """
    pid = tl.program_id(0)
    idx = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)

    mask = idx < n_elements
    partner_idx = idx ^ step

    # Only process if BOTH indices are valid and we're the lower index
    both_valid = mask & (partner_idx < n_elements)
    should_process = both_valid & (idx < partner_idx)

    # Load values (only if both are valid)
    our_val = tl.load(data_ptr + idx, mask=should_process, other=0.0)
    partner_val = tl.load(data_ptr + partner_idx, mask=should_process, other=0.0)

    # Always ascending for final merge
    need_swap = our_val > partner_val

    new_our_val = tl.where(need_swap, partner_val, our_val)
    new_partner_val = tl.where(need_swap, our_val, partner_val)

    # Store back
    tl.store(data_ptr + idx, new_our_val, mask=should_process)
    tl.store(data_ptr + partner_idx, new_partner_val, mask=should_process)


def bitonic_sort_triton(data, BLOCK_SIZE=1024):
    """
    Bitonic sort implementation using multiple kernel launches.

    Args:
        data: Tensor to sort in-place
        BLOCK_SIZE: Number of threads per block
    """
    n = data.numel()

    # Find next power of 2 >= n for bitonic sort
    n_padded = 1
    while n_padded < n:
        n_padded *= 2

    grid = lambda meta: (triton.cdiv(n, meta['BLOCK_SIZE']),)

    # Bitonic sort: Build sequences of size 2, 4, 8, ..., n_padded
    # For each size, we alternate direction and merge
    size = 2
    stage_num = 1
    while size <= n_padded:
        # For this size, do all merge steps with decreasing stride
        step = size >> 1
        while step > 0:
            # Use regular kernel for all stages except the very last one
            # The last stage (size == n_padded) must sort everything ascending
            if size == n_padded:
                # Final stage: always ascending
                bitonic_final_step_kernel[grid](
                    data,
                    n,
                    step=step,
                    BLOCK_SIZE=BLOCK_SIZE,
                )
            else:
                # Earlier stages: alternating directions
                bitonic_step_kernel[grid](
                    data,
                    n,
                    stage=stage_num,
                    step=step,
                    BLOCK_SIZE=BLOCK_SIZE,
                )
            step >>= 1

        size <<= 1
        stage_num += 1


def custom_kernel(data):
    """
    Entry point for the custom sorting kernel.

    Args:
        data: Tuple of (input_tensor, output_tensor)

    Returns:
        output_tensor: Sorted version of input_tensor
    """
    input_tensor, output_tensor = data
    n = input_tensor.numel()

    # Bitonic sort only works correctly for power-of-2 sizes
    # For non-power-of-2, fall back to PyTorch
    is_power_of_2 = (n & (n - 1)) == 0

    if is_power_of_2 and n <= 65536:
        # Copy input to output
        output_tensor.copy_(input_tensor)
        # Use bitonic sort for power-of-2 sizes
        bitonic_sort_triton(output_tensor, BLOCK_SIZE=1024)
    else:
        # Fall back to PyTorch's sort for non-power-of-2 or large arrays
        output_tensor[...] = torch.sort(input_tensor)[0]

    return output_tensor
scrolls · 169 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