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
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