block-transformer
Block Transformer 通过层级化全局到局部建模,将 LLM 推理速度提升 10-20
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
Block Transformer 通过层级化全局到局部建模,将 LLM 推理速度提升 10-20
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
想象一下,你在读一本没有目录和章节的百科全书,要回答任何问题都必须从第一页翻到最后一页——这就是传统大语言模型(Transformer)在生成文本时面临的困境。每生成一个新 token,都要回头看一遍所有已生成的 token,随着序列越来越长,这个「回头看」的计算量呈平方级增长,成为 LLM 推理的主要瓶颈。
Block Transformer 就是来解决这个问题的。它由 KAIST AI、LG AI Research 和 Google DeepMind 的研究团队提出,在 2024 年的 NeurIPS 会议上发表,核心思想是:把文本分成块(Block),先用粗粒度全局注意力把握整体结构,再在每个块内用细粒度局部注意力解码具体 token——就像先读懂一本书的目录和章节标题,再精读每个段落。
Block Transformer 的架构设计非常精妙,包含三个核心模块:
1. Embedder(嵌入器):负责把原始 token 序列切分成块,再嵌入到高维向量空间。它支持多种嵌入策略——RoBERTa 嵌入(用预训练语言模型做 embedding)、T5 嵌入,以及轻量的 Lookup 嵌入(直接查表)。每个块内的 token 数(block_length)是一个关键超参数,决定了全局注意力的压缩比。
2. Block Decoder(块解码器):这是整个架构的核心创新点。它只处理块级别的表示(而非逐 token),通过自注意力机制学习块与块之间的全局依赖关系。Block Decoder 基于 GPT-Neo 或 GPT-NeoX 架构实现,在较低的 transformer 层捕获句子级别甚至段落级别的语义关联。由于处理的 token 数量大幅减少(从 N 个 token 变成 N/block_length 个块),注意力计算复杂度从 O(N²) 降低到 O((N/B)²),其中 B 是块大小。
3. Token Decoder(Token 解码器):接收 Block Decoder 的输出作为前缀(block embeddings),在此基础上精细化解码每个块内的具体 token。它也是基于 GPT-NeoX 实现的,利用 FlashAttention 加速计算。由于每个块只需要解码 block_length 个 token,解码路径被大大缩短。
三层模块协同工作的流程是:输入文本先被切分成块并嵌入,然后 Block Decoder 在块级别建立全局依赖,最后 Token Decoder 在块内精细化解码。每个块的前缀嵌入(block embeddings)来自 Block Decoder 的隐藏状态,实现了全局信息和局部解码的有效结合。
论文的实验结果令人印象深刻。在 The Pile 数据集上,Block Transformer 与同等困惑度(perplexity)的 vanilla Transformer 相比,推理吞吐量提升了 10-20 倍。这一提升来自于两个层面:一是 Block Decoder 减少了需要处理的 token 数量(全局注意力压缩);二是 Token Decoder 解码路径更短(只需解码块内 token)。
在实际部署场景中,这种加速意义重大。以批量推理为例,在 batch_size=128 的设置下,Block Transformer 可以在单卡 A100 上高效运行,而 vanilla transformer 在同样硬件上由于 KV-cache 过大,batch size 受限严重。作者提供的 inference_demo.py 脚本可以直接复现这一对比实验。演示视频见:https://youtu.be/9k9n0RkPBCI
项目代码工程化程度较高,使用了多个工业级训练工具:
训练入口脚本有两个:pretrain_block_transformer.py(Block Transformer 训练)和 pretrain_vanilla_transformer.py(baseline 对比训练)。预训练数据集采用 EleutherAI 的 Pythia(The Pile 去重版),需要额外下载并预处理成 memory-mapped numpy 格式。
适用场景:
局限性:
安装依赖(建议 conda 环境):
pip install torch datasets transformers==4.39.3 accelerate==0.33.0 hydra-core wandb deepspeed flash-attn
# 编译 flash-attn(耗时约 10 分钟)
pip install packaging ninja && pip install flash-attn --no-build-isolation
python setup.py develop # 启用绝对导入
推理演示:
# 从 Google Drive 下载 checkpoint,解压到 ./results/block_main_b4_1.2b/
CUDA_VISIBLE_DEVICES=0 python inference_demo.py --model=block_main_b4_1.2b --batch_size=128
零样本评估:
CUDA_VISIBLE_DEVICES=0 python eval_zero_shot_task.py \
--config-name=eval_multiple_ckpt configs.block=["block_main_b4_5"] batch_size=64
Block Transformer 代表了 LLM 架构优化的一个重要方向——通过层级化的全局到局部建模,在保持模型能力的同时显著降低推理计算量。10-20 倍的吞吐量提升对于需要部署 LLM 服务的团队来说非常有价值。虽然目前只是一个研究代码库(无 Docker、无推理服务框架),但其核心思想已经被工业界关注。配合 HuggingFace Transformers 生态使用,适合有深度学习工程能力的团队进行二次开发和部署。