Skip to content
KernelIndex
Search⌘K

flashinfer / deepgemm / wrapper2ba145

flashinfer_deepgemm_wrapper_2ba145 · FlashInfer-Bench baselines · python · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

No package. Vendor the mirrored source: 63 lines, Apache-2.0, pinned at da91508.

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-flashinfer-deepgemm-wrapper-2ba145?include=source"
interfacepython
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
declared hardwareNVIDIA B200
architecturesunknown
dtypesfp32, fp8_e4m3, int32, int8

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:bd7e96756f4144107473b0907f7af3d2904490f0b8291618d8af0fbe5a26faa2
license declaredApache-2.0
license concludedApache-2.0
authorsbaseline
imported2026-08-16

Kernel source

main.py63 lines
import torch
import deep_gemm
import flashinfer


@torch.no_grad()
def run(q_index_fp8, k_index_cache_fp8, weights, seq_lens, block_table):
    """
    DeepSeek sparse attention top-K indexer using deep_gemm FP8 kernel + FlashInfer.
    
    Pipeline: deep_gemm.fp8_paged_mqa_logits -> flashinfer.top_k_page_table_transform
    """
    batch_size, num_index_heads, index_head_dim = q_index_fp8.shape
    num_pages, page_size, _, _ = k_index_cache_fp8.shape
    topk = 2048

    # Check constants
    assert num_index_heads == 64
    assert index_head_dim == 128
    assert page_size == 64

    device = q_index_fp8.device
    max_num_pages = block_table.shape[1]
    max_context_len = max_num_pages * page_size

    # deep_gemm expects q shape: [batch, next_n, heads, head_dim]
    q_index_fp8_4d = q_index_fp8.unsqueeze(1)  # [batch, 1, heads, head_dim]
    k_index_cache_uint8 = k_index_cache_fp8.view(torch.uint8)

    # Get schedule metadata for deep_gemm
    num_sms = torch.cuda.get_device_properties(device).multi_processor_count
    schedule_meta = deep_gemm.get_paged_mqa_logits_metadata(seq_lens, page_size, num_sms)

    # Compute FP8 attention scores using deep_gemm
    logits = deep_gemm.fp8_paged_mqa_logits(
        q_index_fp8_4d,
        k_index_cache_uint8,
        weights,
        seq_lens,
        block_table,
        schedule_meta,
        max_context_len,
        clean_logits=False,
    )

    # Build token-level page table for FlashInfer
    offsets = torch.arange(page_size, device=device, dtype=torch.int32)
    physical = block_table.unsqueeze(-1) * page_size + offsets  # [batch, max_num_pages, page_size]
    physical_flat = physical.reshape(batch_size, -1)  # [batch, max_num_pages * page_size]
    token_indices = torch.arange(max_num_pages * page_size, device=device)
    mask = token_indices.unsqueeze(0) < seq_lens.unsqueeze(1)
    token_page_table = torch.where(mask, physical_flat, torch.zeros_like(physical_flat))

    # Run FlashInfer top-k selection
    topk_indices = flashinfer.top_k_page_table_transform(
        input=logits.to(torch.float16), 
        src_page_table=token_page_table, 
        lengths=seq_lens, 
        k=topk
    )

    return (topk_indices,)
scrolls · 63 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON