Skip to content
KernelIndex
Search⌘K

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-epilogues_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 output
scrolls · 498 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON