claude-opus-4-1-20250805 / tritonc569cd
claude-opus-4-1-20250805_triton_c569cd · claude-opus-4-1-20250805 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 498 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-20250805-triton-c569cd?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, fp8_e4m3, int32
Benchmark evidence
No published measurement for this revision.
No evidence · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:a5dd063cf5eb7853d45f6ed001d7fd12cacb973c42b06e44787a5c948a9af436
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fused-epilogue
s_with_bias = tl.zeros((256,), dtype=tl.float32)Kernel source
main.py498 lines
import torch
import triton
import triton.language as tl
import math
@triton.jit
def moe_fp8_routing_kernel(
# Routing inputs
routing_logits_ptr, routing_bias_ptr,
# Routing outputs
topk_idx_ptr, weights_ptr,
# Dimensions
seq_len, num_experts,
routed_scaling_factor,
# Strides
stride_rl_t, stride_rl_e,
stride_topk_t, stride_topk_k,
stride_w_t, stride_w_e,
# Block sizes
TOP_K: tl.constexpr,
N_GROUP: tl.constexpr,
TOPK_GROUP: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
"""First pass: compute routing (topk selection and weights)"""
pid_t = tl.program_id(axis=0)
# Constants
GROUP_SIZE: tl.constexpr = 32 # 256 / 8
# Process a single token
token_idx = pid_t
if token_idx >= seq_len:
return
# Load all routing logits and bias for this token
logits_base = routing_logits_ptr + token_idx * stride_rl_t
# Process in blocks to compute sigmoid and add bias
s_vals = tl.zeros((256,), dtype=tl.float32)
s_with_bias = tl.zeros((256,), dtype=tl.float32)
for e_block in range(0, num_experts, BLOCK_SIZE):
e_offs = e_block + tl.arange(0, BLOCK_SIZE)
mask = e_offs < num_experts
logits = tl.load(logits_base + e_offs * stride_rl_e, mask=mask, other=0.0)
bias = tl.load(routing_bias_ptr + e_offs, mask=mask, other=0.0).to(tl.float32)
# Compute sigmoid
s = tl.sigmoid(logits)
s_wb = s + bias
# Store in arrays
for i in range(BLOCK_SIZE):
if e_block + i < num_experts:
idx = e_block + i
val_s = tl.sum(tl.where(tl.arange(0, BLOCK_SIZE) == i, s, 0.0))
val_swb = tl.sum(tl.where(tl.arange(0, BLOCK_SIZE) == i, s_wb, 0.0))
s_vals = tl.where(tl.arange(0, 256) == idx, val_s, s_vals)
s_with_bias = tl.where(tl.arange(0, 256) == idx, val_swb, s_with_bias)
# Compute group scores (top-2 sum per group)
group_scores = tl.zeros((N_GROUP,), dtype=tl.float32)
for g in range(N_GROUP):
g_start = g * GROUP_SIZE
# Find top-2 in this group
max1_val = -1e10
max2_val = -1e10
for i in range(GROUP_SIZE):
idx = g_start + i
val = tl.sum(tl.where(tl.arange(0, 256) == idx, s_with_bias, 0.0))
# Update top-2
is_new_max1 = val > max1_val
is_new_max2 = (val > max2_val) & (~is_new_max1)
# Shift values
max2_val = tl.where(is_new_max1, max1_val, tl.where(is_new_max2, val, max2_val))
max1_val = tl.where(is_new_max1, val, max1_val)
score = max1_val + max2_val
group_scores = tl.where(tl.arange(0, N_GROUP) == g, score, group_scores)
# Select top TOPK_GROUP groups using insertion sort without break
selected_groups = tl.zeros((TOPK_GROUP,), dtype=tl.int32)
selected_scores = tl.full((TOPK_GROUP,), -1e10, dtype=tl.float32)
for g in range(N_GROUP):
g_score = tl.sum(tl.where(tl.arange(0, N_GROUP) == g, group_scores, 0.0))
# Find position to insert - use a flag to track if inserted
insert_pos = TOPK_GROUP # Default to end (won't insert)
for pos in range(TOPK_GROUP):
curr_score = tl.sum(tl.where(tl.arange(0, TOPK_GROUP) == pos, selected_scores, 0.0))
# Only set insert_pos for the first position where g_score > curr_score
should_insert = (g_score > curr_score) & (insert_pos == TOPK_GROUP)
insert_pos = tl.where(should_insert, pos, insert_pos)
# Perform insertion if we found a valid position
for pos in range(TOPK_GROUP):
is_insert_pos = (pos == insert_pos)
# Shift elements after insert_pos
for k in range(TOPK_GROUP - 1, 0, -1):
should_shift = (k > insert_pos) & (insert_pos < TOPK_GROUP)
prev_group = tl.sum(tl.where(tl.arange(0, TOPK_GROUP) == k-1, selected_groups, 0))
prev_score = tl.sum(tl.where(tl.arange(0, TOPK_GROUP) == k-1, selected_scores, -1e10))
selected_groups = tl.where((tl.arange(0, TOPK_GROUP) == k) & should_shift,
prev_group, selected_groups)
selected_scores = tl.where((tl.arange(0, TOPK_GROUP) == k) & should_shift,
prev_score, selected_scores)
# Insert at position
selected_groups = tl.where((tl.arange(0, TOPK_GROUP) == insert_pos) & is_insert_pos,
g, selected_groups)
selected_scores = tl.where((tl.arange(0, TOPK_GROUP) == insert_pos) & is_insert_pos,
g_score, selected_scores)
# Find top-k experts from selected groups
topk_experts = tl.full((TOP_K,), -1, dtype=tl.int32)
topk_s = tl.zeros((TOP_K,), dtype=tl.float32)
topk_scores = tl.full((TOP_K,), -1e10, dtype=tl.float32)
for g_idx in range(TOPK_GROUP):
g = tl.sum(tl.where(tl.arange(0, TOPK_GROUP) == g_idx, selected_groups, 0))
g_start = g * GROUP_SIZE
# Process experts in this group
for i in range(GROUP_SIZE):
expert_id = g_start + i
val_swb = tl.sum(tl.where(tl.arange(0, 256) == expert_id, s_with_bias, 0.0))
val_s = tl.sum(tl.where(tl.arange(0, 256) == expert_id, s_vals, 0.0))
# Find minimum in current top-k
min_score = 1e10
min_pos = 0
for k in range(TOP_K):
curr = tl.sum(tl.where(tl.arange(0, TOP_K) == k, topk_scores, 1e10))
is_min = curr < min_score
min_score = tl.where(is_min, curr, min_score)
min_pos = tl.where(is_min, k, min_pos)
# Replace if better
should_replace = val_swb > min_score
topk_experts = tl.where((tl.arange(0, TOP_K) == min_pos) & should_replace,
expert_id, topk_experts)
topk_s = tl.where((tl.arange(0, TOP_K) == min_pos) & should_replace,
val_s, topk_s)
topk_scores = tl.where((tl.arange(0, TOP_K) == min_pos) & should_replace,
val_swb, topk_scores)
# Store top-k indices
topk_base = topk_idx_ptr + token_idx * stride_topk_t
tl.store(topk_base + tl.arange(0, TOP_K) * stride_topk_k, topk_experts)
# Compute normalized weights
weight_sum = tl.sum(topk_s) + 1e-20
norm_factor = routed_scaling_factor / weight_sum
# Initialize all weights to zero
weights_base = weights_ptr + token_idx * stride_w_t
for e_block in range(0, num_experts, BLOCK_SIZE):
e_offs = e_block + tl.arange(0, BLOCK_SIZE)
mask = e_offs < num_experts
tl.store(weights_base + e_offs * stride_w_e,
tl.zeros((BLOCK_SIZE,), dtype=tl.float32), mask=mask)
# Set weights for selected experts
for k in range(TOP_K):
expert_id = tl.sum(tl.where(tl.arange(0, TOP_K) == k, topk_experts, -1))
weight_val = tl.sum(tl.where(tl.arange(0, TOP_K) == k, topk_s, 0.0)) * norm_factor
valid = expert_id >= 0
if valid:
tl.store(weights_base + expert_id * stride_w_e, weight_val)
@triton.jit
def moe_fp8_compute_kernel(
# Inputs
hidden_states_ptr, hidden_states_scale_ptr,
gemm1_weights_ptr, gemm1_weights_scale_ptr,
gemm2_weights_ptr, gemm2_weights_scale_ptr,
# Routing
topk_idx_ptr, weights_ptr,
# Output
output_ptr,
# Dimensions
seq_len, num_local_experts,
hidden_size, intermediate_size,
local_expert_offset,
# Strides - hidden states
stride_hs_t, stride_hs_h,
stride_hss_b, stride_hss_t,
# Strides - gemm1
stride_g1_e, stride_g1_o, stride_g1_h,
stride_g1s_e, stride_g1s_ob, stride_g1s_hb,
# Strides - gemm2
stride_g2_e, stride_g2_h, stride_g2_i,
stride_g2s_e, stride_g2s_hb, stride_g2s_ib,
# Strides - routing and output
stride_topk_t, stride_topk_k,
stride_w_t, stride_w_e,
stride_out_t, stride_out_h,
# Block configuration
BLOCK_T: tl.constexpr,
BLOCK_H: tl.constexpr,
TOP_K: tl.constexpr,
):
"""Compute kernel for MoE with FP8 weights - optimized for B200"""
pid = tl.program_id(axis=0)
# 2D grid: [seq_len/BLOCK_T, hidden_size/BLOCK_H]
num_t_blocks = tl.cdiv(seq_len, BLOCK_T)
num_h_blocks = tl.cdiv(hidden_size, BLOCK_H)
t_block_idx = pid // num_h_blocks
h_block_idx = pid % num_h_blocks
if t_block_idx >= num_t_blocks:
return
# Token and hidden dimension ranges
t_start = t_block_idx * BLOCK_T
t_offs = t_start + tl.arange(0, BLOCK_T)
t_mask = t_offs < seq_len
h_start = h_block_idx * BLOCK_H
h_offs = h_start + tl.arange(0, BLOCK_H)
h_mask = h_offs < hidden_size
# Initialize output accumulator for token block
output_acc = tl.zeros((BLOCK_T, BLOCK_H), dtype=tl.float32)
# Process each token in the block
for t_idx in range(BLOCK_T):
token_idx = t_start + t_idx
if token_idx >= seq_len:
continue
# Load and dequantize hidden states for this token
hs_fp8 = tl.load(
hidden_states_ptr + token_idx * stride_hs_t + h_offs * stride_hs_h,
mask=h_mask, other=0.0
).to(tl.float32)
# Load scale for this block
h_scale_idx = h_start // 128
hs_scale = tl.load(
hidden_states_scale_ptr + h_scale_idx * stride_hss_b + token_idx * stride_hss_t
).to(tl.float32)
hs_dequant = hs_fp8 * hs_scale
# Accumulator for this token
token_output = tl.zeros((BLOCK_H,), dtype=tl.float32)
# Process each selected expert
for k in range(TOP_K):
# Load expert index and weight
global_expert_id = tl.load(topk_idx_ptr + token_idx * stride_topk_t + k * stride_topk_k)
local_expert_id = global_expert_id - local_expert_offset
# Check if this is a local expert
is_local = (local_expert_id >= 0) & (local_expert_id < num_local_experts)
if is_local:
weight = tl.load(weights_ptr + token_idx * stride_w_t + global_expert_id * stride_w_e).to(tl.float32)
# Skip if weight is too small
if weight > 1e-10:
# GEMM1: Compute gate and up projections
gate_acc = tl.zeros((intermediate_size,), dtype=tl.float32)
up_acc = tl.zeros((intermediate_size,), dtype=tl.float32)
# Process in tiles for GEMM1
for i_idx in range(intermediate_size):
# Gate projection
gate_val = 0.0
up_val = 0.0
for h_tile in range(0, BLOCK_H, 32):
h_tile_offs = h_tile + tl.arange(0, 32)
h_tile_mask = (h_tile_offs < BLOCK_H) & h_mask[h_tile:h_tile+32]
# Load weight values for gate
w1_gate = tl.load(
gemm1_weights_ptr + local_expert_id * stride_g1_e +
i_idx * stride_g1_o + (h_start + h_tile_offs) * stride_g1_h,
mask=h_tile_mask, other=0.0
).to(tl.float32)
# Load weight values for up
w1_up = tl.load(
gemm1_weights_ptr + local_expert_id * stride_g1_e +
(intermediate_size + i_idx) * stride_g1_o + (h_start + h_tile_offs) * stride_g1_h,
mask=h_tile_mask, other=0.0
).to(tl.float32)
# Load scales
h_scale_idx_w = (h_start + h_tile) // 128
gate_scale = tl.load(
gemm1_weights_scale_ptr + local_expert_id * stride_g1s_e +
(i_idx // 128) * stride_g1s_ob + h_scale_idx_w * stride_g1s_hb
).to(tl.float32)
up_scale = tl.load(
gemm1_weights_scale_ptr + local_expert_id * stride_g1s_e +
((intermediate_size + i_idx) // 128) * stride_g1s_ob + h_scale_idx_w * stride_g1s_hb
).to(tl.float32)
# Get hidden states tile
hs_tile = tl.where(h_tile_mask, hs_dequant[h_tile:h_tile+32], 0.0)
# Accumulate
gate_val += tl.sum(w1_gate * gate_scale * hs_tile)
up_val += tl.sum(w1_up * up_scale * hs_tile)
gate_acc = tl.where(tl.arange(0, intermediate_size) == i_idx, gate_val, gate_acc)
up_acc = tl.where(tl.arange(0, intermediate_size) == i_idx, up_val, up_acc)
# Apply SwiGLU activation
gate_silu = gate_acc * tl.sigmoid(gate_acc)
intermediate = gate_silu * up_acc
# GEMM2: Down projection
for h_idx in range(BLOCK_H):
if h_start + h_idx < hidden_size:
out_val = 0.0
for i_tile in range(0, intermediate_size, 32):
i_tile_offs = i_tile + tl.arange(0, 32)
i_tile_mask = i_tile_offs < intermediate_size
# Load weight tile
w2_tile = tl.load(
gemm2_weights_ptr + local_expert_id * stride_g2_e +
(h_start + h_idx) * stride_g2_h + i_tile_offs * stride_g2_i,
mask=i_tile_mask, other=0.0
).to(tl.float32)
# Load scale
w2_scale = tl.load(
gemm2_weights_scale_ptr + local_expert_id * stride_g2s_e +
((h_start + h_idx) // 128) * stride_g2s_hb + (i_tile // 128) * stride_g2s_ib
).to(tl.float32)
# Get intermediate values
inter_tile = tl.where(i_tile_mask, intermediate[i_tile:i_tile+32], 0.0)
# Accumulate
out_val += tl.sum(w2_tile * w2_scale * inter_tile)
# Accumulate weighted output
token_output = tl.where(tl.arange(0, BLOCK_H) == h_idx,
token_output[h_idx] + out_val * weight,
token_output)
# Store token output in block accumulator
for h_idx in range(BLOCK_H):
val = tl.sum(tl.where(tl.arange(0, BLOCK_H) == h_idx, token_output, 0.0))
output_acc = tl.where((tl.arange(0, BLOCK_T)[:, None] == t_idx) &
(tl.arange(0, BLOCK_H)[None, :] == h_idx),
val, output_acc)
# Store final output block
out_ptr = output_ptr + t_offs[:, None] * stride_out_t + h_offs[None, :] * stride_out_h
tl.store(out_ptr, output_acc.to(tl.bfloat16), mask=t_mask[:, None] & h_mask[None, :])
def run(
routing_logits,
routing_bias,
hidden_states,
hidden_states_scale,
gemm1_weights,
gemm1_weights_scale,
gemm2_weights,
gemm2_weights_scale,
local_expert_offset,
routed_scaling_factor,
):
"""Main entry point for the MoE FP8 kernel"""
# Check CUDA availability
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available but this kernel requires GPU")
# Device management
device = None
tensors = {
'routing_logits': routing_logits,
'routing_bias': routing_bias,
'hidden_states': hidden_states,
'hidden_states_scale': hidden_states_scale,
'gemm1_weights': gemm1_weights,
'gemm1_weights_scale': gemm1_weights_scale,
'gemm2_weights': gemm2_weights,
'gemm2_weights_scale': gemm2_weights_scale
}
# Track original devices and move to GPU if needed
original_devices = {}
gpu_tensors = {}
for name, tensor in tensors.items():
if tensor is not None:
original_devices[name] = tensor.device
if tensor.device.type != 'cuda':
gpu_tensors[name] = tensor.cuda()
else:
gpu_tensors[name] = tensor
if device is None:
device = tensor.device
if device is None:
device = torch.device('cuda:0')
# Ensure tensors are contiguous
for name in gpu_tensors:
if not gpu_tensors[name].is_contiguous():
gpu_tensors[name] = gpu_tensors[name].contiguous()
# Get dimensions
seq_len = gpu_tensors['routing_logits'].shape[0]
num_experts = gpu_tensors['routing_logits'].shape[1]
num_local_experts = gpu_tensors['gemm1_weights'].shape[0]
hidden_size = 7168
intermediate_size = 2048
# Routing constants
TOP_K = 8
N_GROUP = 8
TOPK_GROUP = 4
BLOCK_T = 1 # Tokens per block
BLOCK_H = 128 # Block size for hidden dimension
BLOCK_SIZE = 32 # Block size for expert processing
# Allocate outputs
output = torch.zeros((seq_len, hidden_size), dtype=torch.bfloat16, device=device)
topk_idx = torch.zeros((seq_len, TOP_K), dtype=torch.int32, device=device)
weights = torch.zeros((seq_len, num_experts), dtype=torch.float32, device=device)
# Launch routing kernel
grid_routing = (seq_len,)
moe_fp8_routing_kernel[grid_routing](
gpu_tensors['routing_logits'], gpu_tensors['routing_bias'],
topk_idx, weights,
seq_len, num_experts,
routed_scaling_factor,
gpu_tensors['routing_logits'].stride(0), gpu_tensors['routing_logits'].stride(1),
topk_idx.stride(0), topk_idx.stride(1),
weights.stride(0), weights.stride(1),
TOP_K, N_GROUP, TOPK_GROUP, BLOCK_SIZE,
)
# Launch compute kernel
num_t_blocks = triton.cdiv(seq_len, BLOCK_T)
num_h_blocks = triton.cdiv(hidden_size, BLOCK_H)
grid_compute = (num_t_blocks * num_h_blocks,)
moe_fp8_compute_kernel[grid_compute](
gpu_tensors['hidden_states'], gpu_tensors['hidden_states_scale'],
gpu_tensors['gemm1_weights'], gpu_tensors['gemm1_weights_scale'],
gpu_tensors['gemm2_weights'], gpu_tensors['gemm2_weights_scale'],
topk_idx, weights,
output,
seq_len, num_local_experts,
hidden_size, intermediate_size,
local_expert_offset,
# Hidden states strides
gpu_tensors['hidden_states'].stride(0), gpu_tensors['hidden_states'].stride(1),
gpu_tensors['hidden_states_scale'].stride(0), gpu_tensors['hidden_states_scale'].stride(1),
# GEMM1 strides
gpu_tensors['gemm1_weights'].stride(0), gpu_tensors['gemm1_weights'].stride(1),
gpu_tensors['gemm1_weights'].stride(2),
gpu_tensors['gemm1_weights_scale'].stride(0), gpu_tensors['gemm1_weights_scale'].stride(1),
gpu_tensors['gemm1_weights_scale'].stride(2),
# GEMM2 strides
gpu_tensors['gemm2_weights'].stride(0), gpu_tensors['gemm2_weights'].stride(1),
gpu_tensors['gemm2_weights'].stride(2),
gpu_tensors['gemm2_weights_scale'].stride(0), gpu_tensors['gemm2_weights_scale'].stride(1),
gpu_tensors['gemm2_weights_scale'].stride(2),
# Routing and output strides
topk_idx.stride(0), topk_idx.stride(1),
weights.stride(0), weights.stride(1),
output.stride(0), output.stride(1),
BLOCK_T, BLOCK_H, TOP_K,
)
# Move output back to original device if needed
if 'hidden_states' in original_devices and original_devices['hidden_states'].type != 'cuda':
output = output.to(original_devices['hidden_states'])
return outputscrolls · 498 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON