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