RadixMLP——用于因果 Transformer 的批内去重
机器学习
2026-01-22 v1 分布式、并行与集群计算
摘要
因果 transformer 模型的批推断工作负载经常处理共享公共前缀的序列,如系统提示、few-shot 示例或共享查询。标准推理引擎将每个序列视为独立单元,导致对共享前缀中每个副本冗余地重新计算相同的 MLP 激活。我们引入 RadixMLP 技术,利用 MLP、LayerNorm、线性投影和嵌入的位点依赖性,消除这种冗余。RadixMLP 在批次上动态映射到前缀字典树,将共享分段压缩到紧凑表示中进行位点计算,仅在注意力边界处将结果重新分发。RadixMLP 无状态,能在单次前向传播内完成。在 MS~MARCO v1.1 上采用 Qwen3 模型(0.6B 至 8B 参数)的端到端服务基准测试中,RadixMLP 在实际的重排序工作负载中实现 1.44-1.59 倍的加速,在带有更长共享前缀的合成基准中可达最高 5 倍的加速。我们的代码已公开于 https://github.com/michaelfeil/radix-mlp。
引用
@article{arxiv.2601.15013,
title = {RadixMLP -- Intra-batch Deduplication for Causal Transformers},
author = {Michael Feil and Julius Lipp},
journal= {arXiv preprint arXiv:2601.15013},
year = {2026}
}