CODA: 把 Transformer Block 重写成 GEMM-Epilogue 程序 · 干货攻略

  • 链接:https://x.com/tri_dao/status/2057640492020469845
  • 分类:x-tips
  • 来源:X @tri_dao
  • 作者:Jay
  • 更新:2026-10-08

这是什么

CODA 是一套 GPU kernel 抽象框架,核心主张是:Transformer 里那些看似必须单独启动 kernel 的内存密集型算子(normalization、activation、residual update、reduction),可以通过数学重参数化,全部塞进 GEMM 的 epilogue 阶段执行——在 GEMM 输出 tile 还在芯片上时就把这些操作做完,不用写回 HBM 再读出来。

论文:CODA: Rewriting Transformer Blocks as GEMM-Epilogue Programs,Han Guo(Princeton/Together AI)、Jack Zhang、Arjun Menon、Driss Guessous、Vijay Thakkar(Meta)、Yoon Kim(MIT)、Tri Dao(Princeton/Together AI),arXiv:2605.19269,2026 年 5 月。

代码仓库:open-lm-engine/coda-kernels(GitHub,256 ⭐),基于 CUTLASS CuTeDSL 构建,目标硬件为 NVIDIA Hopper(H100)。


为什么值得关注

谁在关注这个问题

Tri Dao(FlashAttention 作者)亲自转发说明这套方向在他看来是重要的工程问题,而论文团队同时有 Princeton 系统组 + Meta 硬件团队 + Together AI 的背景——这是业内做 LLM training 系统效率最顶尖的组合之一。

解决什么问题

现代 LLM 训练的计算分布已经不是"只要优化 GEMM 就够了"。当矩阵乘法被 FP8/FP4 Tensor Core 加速到几乎瞬间完成时,数据搬运(memory-bound op)成了新的瓶颈。以 LLaMA-3-style 1B 模型在单卡 H100 上训练为例(非 GEMM 算子占据约 30-40% 的 end-to-end 时间),这个问题随低精度格式普及只会越来越严重。

现有的编程模型让这个问题很难解决:PyTorch 的 operator 序列天然产生 materialization boundary,框架层看到的 operator 边界恰好是 GEMM 和 memory-bound op 分离的地方——融合机会被抹掉了。生产系统通常靠手写 backward pass 或完全绕过框架来解决,但门槛极高。

CODA 想要填上"框架生产力"和"手工 kernel 效率"之间的 gap。


核验过程

官方来源

  1. arXiv abstract(arxiv.org/abs/2605.19269):确认论文标题、作者阵容、主要贡献(GEMM-plus-epilogue 重参数化、五类 epilogue 原语、forward+backward 覆盖)、代码仓库地址。抽象描述与 X 帖内容一致。
  2. GitHub README(github.com/open-lm-engine/coda-kernels):确认仓库当前维护状态、最新版本 v0.3.1(2026-10-06),确认安装方式(pip install git+https://github.com/open-lm-engine/coda-kernels.git@v0.3.1),确认依赖 quack fork(v0.6.5+fork.1),确认已实现的 functional API 包括 linear_swiglu、linear_sigmoid、linear_cross_entropy、linear_qknorm_rope。
  3. Han Guo X 帖(x.com/HanGuo97/status/2057588533595136229):作者本人确认 CODA 的结构化约束(从优化的 GEMM 模板出发组合少量 fast primitives)恰好给了 LLM 足够但不过多的结构,使 LLM 写出的 CuTeDSL kernel 能达到高性能。此说法与论文 abstract 中"both human- and LLM-authored CODA kernels achieve high performance"一致。

交叉验证

通过 YouTube 视频(Tri Dao 转发,2026 年 5 月发布)时间戳交叉验证以下说法: - "backward pass 比标准 PyTorch 快 1.7x":视频 16:05 处明确提到 CODA 在 backward pass 实现 up to 1.7x speedup。核验结论:论文有明确数据支撑,与 X 帖原意相符。 - "LLM 写的 kernel 可以接近 SoL":论文原文(未直接读取 PDF,但多方摘要一致)描述为"both human- and LLM-authored CODA kernels achieve high performance",Tri Dao X 帖表述为"接近 SoL"。核验结论:论文确实报告了 LLM-authored kernel 的 high performance,但"接近 SoL"的具体百分比数字(x% of roofline)未在本文档中核验到,此说法为原帖主张,未经逐字核验。


上手步骤

安装

# 依赖 quack fork,安装时会替换 PyPI 上的 quack-kernels
pip install git+https://github.com/open-lm-engine/coda-kernels.git@v0.3.1

# 或从源码
git clone https://github.com/open-lm-engine/coda-kernels.git
cd coda-kernels
pip install -e .

⚠️ 注意:coda-kernels 依赖的 quack 版本(v0.6.5+fork.1)与 PyPI 上公开的 quack-kernels 0.6.5 不同,安装时会覆盖 PyPI 版本。如果同时使用其他依赖 quack 的项目,需要注意版本隔离。

使用 Functional API

import torch
from coda.kernels.functional import linear_swiglu, linear_qknorm_rope

# SwiGLU activation fused into GEMM epilogue
# hidden_dim=4096, intermediate_dim=11008(LLaMA-3 style)
h = torch.randn(32, 128, 4096, device='cuda', dtype=torch.bfloat16)
w1 = torch.randn(11008, 4096, device='cuda', dtype=torch.bfloat16)
w2 = torch.randn(4096, 11008, device='cuda', dtype=torch.bfloat16)
out = linear_swiglu(h, w1, w2)  # 无需 separate SiLU + multiply kernel

# QK-Norm + RoPE fused
q = torch.randn(32, 8, 128, 64, device='cuda', dtype=torch.bfloat16)
k = torch.randn(32, 8, 128, 64, device='cuda', dtype=torch.bfloat16)
q_out, k_out = linear_qknorm_rope(q, k)  # QK norm + rotary 一次做掉

用 Benchmark 验证效果

cd coda-kernels
pip install -e ".[benchmark]"
# benchmarks/ 目录下有 block-level benchmark 脚本
# hidden_sizes: {2048, 4096, 8192} 对应约 1B / 7B / 70B 模型规模

benchmark 测量 GEMM + epilogue kernel sequence 的 end-to-end 延迟,包含 auxiliary reduction 和 glue operation,与 isolate kernel benchmark 不同(后者不包含这些 overhead)。


坑与适用边界

⚠️ 适用场景有限

CODA 主要解决 training 场景的 memory-bound overhead,对 inference 的适用性需要单独评估(inference 的 memory-bound pattern 与 training 不同)。

⚠️ 依赖 CUTLASS CuTeDSL,不支持所有 GPU

目前目标硬件为 NVIDIA Hopper(H100)。对 Ampere、Ada 等旧架构的支持未在文档中明确说明。AMD GPU 完全不适用。

⚠️ 安装复杂

依赖 quack fork,与 PyPI 版本不兼容。如果项目已经在用其他版本的 quack-kernels,安装 CODA 会造成版本冲突。建议在 virtual environment 或 docker container 中单独测试。

⚠️ API 粒度是 functional kernel,不是框架级透明替换

目前提供的 API 是 linear_swiglu 这种粒度的 kernel 函数,不是能直接丢进 PyTorch model 替换掉 nn.Linear + activation 的东西。集成到训练框架需要额外的工程工作。

⚠️ LLM 写 kernel 的质量依赖结构约束

论文强调 LLM 能写出好 kernel 是因为 epilogue 的结构约束恰好限制了 LLM 产生错误的 grid-wide sync barrier 等问题,但这不代表 LLM 能在没有 CODA 框架的情况下写出等价性能的手写 CUDA。


一句话结论

CODA 用数学重参数化的思路,把 Transformer 里散落的 memory-bound op 全部融合进 GEMM epilogue——backward pass 实测最高 1.7x 加速,LLM-authored kernel 也能达到 near-hardware-efficiency,是值得关注的 LLM training 系统效率方向;上手门槛(CUDA/Hopper、quack 依赖)较高,适合有底层优化需求的团队深入评估。