Skip to content
KernelIndex
Search⌘K

submission 555497

ronitmathur_37198 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-gated-deltanet-chunk-fwd-h-555497?include=source"
interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVIDIA B200
12.2µs
#7 of 28
2026-03-15

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e6a4422226945f7cdd05fdfc497f98e5a375fc9332cca1d5d331053f28293098
license declaredunknown
license concludedunknown
authorsronitmathur_37198
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

autotune@helion.kernel(static_shapes=True, dot_precision="ieee", config=config, autotune_effort="none")
num-warps = 16…', '', '', 'first'], loop_orders=[[0, 1, 2]], num_stages=1, num_warps=16, pid_type='flat', range_flattens=[None, None], range_multi_buffers=[None, None], range_num_stages=[0, 0], r…
persistent-kernel…, num_sm_multiplier=1, num_stages=1, num_warps=8, pid_type='persistent_blocked', range_flattens=[None, None], range_multi_buffers=[None, True], range_num_stages=[0, 0], range_unrol…
stages = 1…olicies=['', '', '', '', 'first'], loop_orders=[[0, 1, 2]], num_stages=1, num_warps=16, pid_type='flat', range_flattens=[None, None], range_multi_buffers=[None, None], range_num_st…
warp-specialization…range_num_stages=[0, 0], range_unroll_factors=[0, 0], range_warp_specializes=[None, None], static_ranges=[False]),…

Kernel source

submission.py182 lines
from task import input_t, output_t

import torch
import helion
import math
import helion.language as hl

# ---------------------------------------------------------------------------
# Per-shape configs: map (B, T, H, K, V) to optimized helion.Config objects.
#
# With static_shapes=True and block_size=[1, 1, block_v], Helion collapses
# the size-1 B and H tile dimensions, leaving exactly ONE tunable block_size
# slot — for the V dimension. So block_sizes must be a single-element list.
#
# block_sizes=[V]  -> process full V dimension per thread block (recommended)
# block_sizes=[32] -> split V into 32-wide tiles (more parallelism, smaller matmul)
#
# Autotune locally for each shape, then paste the best config here.
# ---------------------------------------------------------------------------
SHAPE_CONFIGS: dict[tuple, helion.Config] = {
    # ---- test shapes (correctness gate) ------------------------------------
    # (B, T, H, K, V)
    (1,   64, 2,  64,  64): helion.Config(block_sizes=[8], indexing=['pointer', 'pointer', 'pointer', 'pointer', 'pointer', 'pointer', 'pointer'], l2_groupings=[1], load_eviction_policies=['', '', '', '', 'first'], loop_orders=[[0, 1, 2]], num_stages=1, num_warps=16, pid_type='flat', range_flattens=[None, None], range_multi_buffers=[None, None], range_num_stages=[0, 0], range_unroll_factors=[0, 0], range_warp_specializes=[None, None], static_ranges=[False]),
    (2,  128, 4,  64,  64): helion.Config(block_sizes=[4], indexing=['pointer', 'pointer', 'pointer', 'pointer', 'pointer', 'pointer', 'pointer'], l2_groupings=[2], load_eviction_policies=['', '', '', '', ''], loop_orders=[[2, 1, 0]], num_stages=1, num_warps=8, pid_type='flat', range_flattens=[None, None], range_multi_buffers=[None, None], range_num_stages=[0, 0], range_unroll_factors=[0, 0], range_warp_specializes=[None, None], static_ranges=[False]),
    (1,  256, 4,  64, 128): helion.Config(block_sizes=[4], indexing=['pointer', 'pointer', 'pointer', 'tensor_descriptor', 'pointer', 'pointer', 'pointer'], l2_groupings=[4], load_eviction_policies=['', '', '', '', 'last'], loop_orders=[[2, 1, 0]], num_stages=1, num_warps=4, pid_type='flat', range_flattens=[None, None], range_multi_buffers=[None, None], range_num_stages=[0, 0], range_unroll_factors=[0, 0], range_warp_specializes=[None, None], static_ranges=[False]),

    # ---- benchmark shapes (performance) ------------------------------------
    (1,   64, 1,  64,  64): helion.Config(block_sizes=[8], indexing=['pointer', 'pointer', 'pointer', 'pointer', 'pointer', 'pointer', 'pointer'], l2_groupings=[1], load_eviction_policies=['', '', '', '', ''], loop_orders=[[0, 2, 1]], num_stages=1, num_warps=16, pid_type='flat', range_flattens=[None, None], range_multi_buffers=[None, None], range_num_stages=[0, 0], range_unroll_factors=[0, 0], range_warp_specializes=[None, None], static_ranges=[False]),
    (2,  512, 3,  64,  64): helion.Config(block_sizes=[4], indexing=['pointer', 'tensor_descriptor', 'pointer', 'pointer', 'tensor_descriptor', 'pointer', 'pointer'], l2_groupings=[1], load_eviction_policies=['', '', '', '', 'first'], loop_orders=[[0, 1, 2]], num_stages=1, num_warps=4, pid_type='flat', range_flattens=[None, None], range_multi_buffers=[None, None], range_num_stages=[0, 0], range_unroll_factors=[0, 0], range_warp_specializes=[None, None], static_ranges=[True]),
    (2, 1024, 3,  64,  64): helion.Config(block_sizes=[4], indexing=['pointer', 'pointer', 'pointer', 'pointer', 'pointer', 'pointer', 'pointer'], l2_groupings=[1], load_eviction_policies=['', '', '', 'first', 'last'], loop_orders=[[1, 0, 2]], num_stages=1, num_warps=4, pid_type='flat', range_flattens=[None, None], range_multi_buffers=[None, True], range_num_stages=[0, 0], range_unroll_factors=[0, 0], range_warp_specializes=[None, None]),
    (3, 1024, 4, 100, 100): helion.Config(block_sizes=[16], indexing=['pointer', 'pointer', 'pointer', 'pointer', 'pointer', 'pointer', 'pointer'], l2_groupings=[1], load_eviction_policies=['', '', '', '', ''], loop_orders=[[0, 1, 2]], num_stages=1, num_warps=8, pid_type='flat', range_flattens=[None, None], range_multi_buffers=[None, False], range_num_stages=[0, 1], range_unroll_factors=[0, 0], range_warp_specializes=[None, None]),
    (4, 1024, 4, 128, 128): helion.Config(block_sizes=[16], indexing=['pointer', 'tensor_descriptor', 'pointer', 'pointer', 'pointer', 'pointer', 'pointer'], l2_groupings=[1], load_eviction_policies=['', '', '', '', ''], loop_orders=[[0, 1, 2]], num_stages=1, num_warps=16, pid_type='flat', range_flattens=[None, None], range_multi_buffers=[None, None], range_num_stages=[0, 0], range_unroll_factors=[0, 0], range_warp_specializes=[None, False]),
    (2, 1536, 4, 128, 128): helion.Config(block_sizes=[8], indexing=['pointer', 'tensor_descriptor', 'tensor_descriptor', 'pointer', 'pointer', 'pointer', 'pointer'], l2_groupings=[1], load_eviction_policies=['', '', '', '', ''], loop_orders=[[0, 1, 2]], num_sm_multiplier=1, num_stages=1, num_warps=8, pid_type='persistent_blocked', range_flattens=[None, None], range_multi_buffers=[None, True], range_num_stages=[0, 0], range_unroll_factors=[0, 2], range_warp_specializes=[None, None]),
    (4, 2048, 8,  64,  64): helion.Config(block_sizes=[16], indexing=['pointer', 'pointer', 'pointer', 'pointer', 'pointer', 'tensor_descriptor', 'pointer'], l2_groupings=[1], load_eviction_policies=['', 'last', 'first', 'first', ''], loop_orders=[[0, 1, 2]], num_stages=1, num_warps=8, pid_type='flat', range_flattens=[None, None], range_multi_buffers=[None, None], range_num_stages=[0, 0], range_unroll_factors=[0, 0], range_warp_specializes=[None, None]),
}

# ---------------------------------------------------------------------------
# Kernel factory
# ---------------------------------------------------------------------------

def _make_kernel(config: helion.Config):
    @helion.kernel(static_shapes=True, dot_precision="ieee", config=config, autotune_effort="none")
    def kernel(
        k: torch.Tensor,  # [B, T, H, K]
        w: torch.Tensor,  # [B, T, H, K]
        u: torch.Tensor,  # [B, T, H, V]
        g: torch.Tensor,  # [B, T, H]   cumulative gate (log-domain, chunk-local cumsum)
    ) -> tuple[torch.Tensor, torch.Tensor]:
        """
        Chunk-wise forward recurrence for Gated DeltaNet (arXiv:2412.06464).

        For each (b, h) pair, processes NT = T // C chunks of size C = 64:
          state = zeros(K, V)
          for c in range(NT):
              h_out[b, c, h]  = state                         # store before update
              proj            = w[c] @ state                  # [C, K] @ [K, V] -> [C, V]
              v_new           = u[c] - proj                   # corrected values
              g_last          = g[c, -1]                      # scalar gate for chunk end
              gate            = exp(g_last - g[c])            # [C] per-token decay
              v_gated         = v_new * gate[:, None]         # [C, V]
              state           = state * exp(g_last)           # decay existing state
              state           = state + k[c]^T @ v_gated      # [K, C] @ [C, V] -> [K, V]
          v_out[b, :, h] = v_new across all chunks
        """
        B, T, H, K = k.shape
        V = u.shape[-1]
        C = 64  # chunk size (T must be multiple of 64)

        # Specialize K and V so hl.dot gets constexpr tile shapes for tl.dot.
        K = hl.specialize(K)
        V = hl.specialize(V)

        NT = T // C

        # Output tensors allocated on host side (outside the tile loop).
        h_out = torch.empty(B, NT, H, K, V, dtype=k.dtype, device=k.device)
        v_out = torch.empty_like(u)

        # ----------------------------------------------------------------
        # Register V-dimension block size so the autotuner can explore it.
        # Block sizes must be powers of two — round V up if needed.
        # V=64 -> 64, V=100 -> 128, V=128 -> 128.
        # ----------------------------------------------------------------
        V_pow2 = 2 ** math.ceil(math.log2(max(V, 1)))
        block_v = hl.register_block_size(V_pow2)

        # ----------------------------------------------------------------
        # Outer tile: parallelise over (B, H, V-columns).
        # block_size=[1, 1, block_v]:
        #   - B and H tiled at 1 → Helion collapses them, leaving ONE
        #     tunable block_size slot (for V) in the compiled config.
        #   - block_v columns of V per GPU thread block.
        # The sequential chunk loop inside is NOT tiled — it carries the
        # recurrence dependency and must execute serially per (b, h).
        # ----------------------------------------------------------------
        for tile_b, tile_h, tile_v in hl.tile(
            [B, H, V], block_size=[1, 1, block_v]
        ):
            i_b = tile_b.begin   # scalar batch index
            i_h = tile_h.begin   # scalar head  index

            # Accumulator: hidden state for this (b, h) pair, V-column slice.
            # Shape: [K, block_v], held in registers across all chunk iterations.
            state = hl.zeros([K, tile_v], dtype=torch.float32)

            # ----------------------------------------------------------------
            # Sequential recurrence over NT chunks.
            # block_size=C is fixed (chunk boundary), not autotuned.
            # ----------------------------------------------------------------
            for t_chunk in hl.tile(T, block_size=C):
                chunk_idx = t_chunk.begin // C   # integer chunk index (0 … NT-1)

                # --- Step 1: store state BEFORE update ----------------------
                # h_out[b, chunk_idx, h, :, v_cols] = state  (cast to output dtype)
                h_out[i_b, chunk_idx, i_h, :, tile_v] = state.to(k.dtype)

                # --- Step 2: load chunk inputs ------------------------------
                # k_tile : [C, K]  — keys for this chunk
                # w_tile : [C, K]  — WY-transformed keys
                # u_tile : [C, block_v] — WY-transformed values (V-slice)
                # g_chunk: [C]    — cumulative gate for this chunk
                k_tile  = k[i_b, t_chunk, i_h, :]                         # [C, K]
                w_tile  = w[i_b, t_chunk, i_h, :]                         # [C, K]
                u_tile  = u[i_b, t_chunk, i_h, tile_v].to(torch.float32)  # [C, block_v]
                g_chunk = g[i_b, t_chunk, i_h].to(torch.float32)          # [C]

                # g_last: scalar gate at the final timestep of this chunk.
                # Uses the same min() pattern as the official gdn_fwd_h example
                # to safely derive the last-timestep index from the tile begin.
                t_last = min(t_chunk.begin + C, T) - 1
                g_last = g[i_b, t_last, i_h].to(torch.float32)

                # --- Step 3: compute proj = w_tile @ state ------------------
                # [C, K] @ [K, block_v] -> [C, block_v]
                # No duplicate matmul — single hl.dot call.
                proj = hl.dot(w_tile, state.to(k.dtype), out_dtype=torch.float32)

                # --- Step 4: compute v_new = u - proj -----------------------
                v_new = u_tile - proj   # [C, block_v]

                # --- Step 5: gate v_new -------------------------------------
                # gate[t] = exp(g_last - g[t])  for t in chunk   -> [C]
                gate = torch.exp(g_last - g_chunk)       # [C]
                v_gated = v_new * gate[:, None]          # [C, block_v]

                # --- Step 6: write v_new to output --------------------------
                v_out[i_b, t_chunk, i_h, tile_v] = v_new.to(u.dtype)

                # --- Step 7: decay state and update -------------------------
                # state = state * exp(g_last) + k^T @ v_gated
                # Fused via hl.dot(acc=state * decay):
                #   hl.dot lowers to tl.dot which accumulates into acc,
                #   so the scale-add is fused into the GEMM epilogue.
                decay = torch.exp(g_last)   # scalar
                state = hl.dot(
                    k_tile.T,                    # [K, C]
                    v_gated.to(k.dtype),         # [C, block_v]
                    acc=state * decay,           # fused: decay existing state
                    out_dtype=torch.float32,
                )

        return h_out, v_out

    return kernel


# ---------------------------------------------------------------------------
# Pre-compile one kernel variant per shape.
# ---------------------------------------------------------------------------
_KERNELS = {shape: _make_kernel(cfg) for shape, cfg in SHAPE_CONFIGS.items()}


# ---------------------------------------------------------------------------
# Entry point called by the evaluation harness.
# ---------------------------------------------------------------------------
def custom_kernel(data: input_t) -> output_t:
    k, w, u, g = data
    B, T, H, K = k.shape
    V = u.shape[-1]
    kernel = _KERNELS[(B, T, H, K, V)]
    return kernel(k, w, u, g)
scrolls · 182 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