DeepSelect
DeepSelect 是 DeepSeek 稀疏注意力(DSA)中所用 TopK 内核及采样器的高性能实现(DSA 用于 DeepSeek V3.2、DeepSeek V4 和 DeepSeek V4.1 模型)。与原生 torch.topk 相比,它可实现 2 ~ 20 倍加速。
新闻
支持的场景
TopK 工作负载差异很大,最快的算法与实现高度依赖于输入 dtype、batch_size、vocab_size 和 topk。本仓库仅关注以下场景:
Lightning Indexer 场景
该场景涵盖:
- 输入 dtype:
torch.bfloat16 batch_size:$1 \sim +\infty$(大、小批量均已优化)vocab_size:$1 \sim +\infty$(大、小词表均已优化)topk:较小(必须 $\le 4096$;不支持更大的值)
建议:
- 除非输出需要按索引或按值排序,否则请禁用
sorted_index;启用任一项都会损失性能。 - 不需要值时,设置
return_value=False。这将跳过值的输出,速度更快。
采样场景
该场景涵盖:
- 输入 dtype:
torch.float32 batch_size:$1 \sim +\infty$vocab_size:约 128Ktopk:较小(必须 $\le 4096$;不支持更大的值)
性能
使用 tests/test.py 中的基准测试进行测量
(python3 tests/test.py --perf-only),该测试报告在同一输入下相对于 torch.topk 的比值。指标为有效内存带宽:TopK 不进行浮点运算,因此 FLOP 速率在此没有意义。
Lightning Indexer 场景
bfloat16,topk = 512,每个批量大小一个子图,共享 0 - 7 TB/s 坐标轴。

采样场景
float32,vocab_size = 129280,topk = 512。

安装
git clone https://github.com/deepseek-ai/DeepSelect.git
cd DeepSelect
git submodule update --init --recursive
pip install -v .
用法
import torch
import deep_select
# input: (batch_size, vocab_size), torch.bfloat16 or torch.float32.
# Its row stride must be a multiple of `deep_select.get_stride_requirement()[0]` bytes, and its last dimension must be contiguous.
batch_size, vocab_size, topk = 4, 204800, 1024
x = torch.randn(batch_size, vocab_size, dtype=torch.bfloat16, device="cuda")
values, indices = deep_select.topk(
x,
topk,
sorted_index=True, # return each row's indices in ascending order
indices_type=torch.int32, # torch.int32 or torch.int64
return_value=True, # False skips the value output (~10% faster)
)
# values: (batch_size, topk) of x.dtype
# indices: (batch_size, topk) of indices_type
输入张量(x)的行步长必须按 deep_select.get_stride_requirement()[0] 字节对齐。对于未对齐的输入,需要进行填充。
两个输出均由调用方分配,其步长按 deep_select.get_stride_requirement()[1] 字节对齐(因此可能是非连续的)。传入 output_idx= 可将索引写入你自有的缓冲区,该缓冲区必须满足相同的步长要求。
完整函数签名请参见 deep_select/interface.py。
变长行
end 用于设置每行的上界(不含)。短于 topk 的行将使用
value_oob_fill_value / idx_oob_fill_value 进行填充:
batch_size, vocab_size = 2, 129280 # 129280 is a multiple of 256, so float32 is fine
x = torch.randn(batch_size, vocab_size, dtype=torch.float32, device="cuda")
end = torch.tensor([129280, 100000], dtype=torch.int32, device="cuda") # (batch_size,)
values, indices = deep_select.topk(x, 1000, end=end, sorted=True,
indices_type=torch.int64)
NaN 处理
NaN 检查始终开启。在默认的 abort_when_nan_found=True 下,内核会调用 trap() 并中止。长度 <= topk 的行不会进行 NaN 检查。
引用
@misc{deepselect2026,
title={DeepSelect: High-Performance TopK Kernels for DeepSeek Sparse Attention and Sampling},
author={Yi Qian and Shengyu Liu and Yichen Li},
year={2026},
publisher = {GitHub},
howpublished = {\url{https://github.com/deepseek-ai/DeepSelect}},
}