LogoCookLLM文档
LogoCookLLM文档
首页CookLLM

原理精讲

词元化
Tokenization 基础BPE 算法详解GPT 系列 TokenizerBPE 训练工程化
模型架构
Transformer LM
从 token ids 到 logitsEmbedding 与 LM Head
Attention 机制
Self-Attention 到 GQAAttention Sink
位置编码
位置编码基础RoPE 数学推导RoPE 代码实现长度外推
GPU 编程基础
GPU 架构基础张量布局Triton 入门:向量加法
FlashAttention
Flash Attention 原理详解从朴素实现到 Auto-TuningBlock Pointer 与多维支持Causal Masking 优化Grouped Query Attention反向传播实现
分布式训练
数据并行ZeRO 优化器全分片数据并行张量并行流水线并行多维混合并行
推理优化
KV CacheContinuous BatchingPagedAttention

动手训练

概述
预训练
预训练数据Tokenizer 训练模型架构数据流水线训练循环监控与验证
X (Twitter)
系统工程FlashAttention

Causal Masking 优化

会员专享

为自回归模型实现因果注意力机制,通过跳过上三角计算实现 ~2x 加速。

在权益中心获取代码

在前三章中,我们实现了一个完整的 Flash Attention Kernel,支持任意长度的序列和多维并行。但我们还没有针对一个非常重要的应用场景进行优化:自回归生成(Autoregressive Generation)。

这正是 GPT、LLaMA 等 Decoder-Only 模型的核心计算模式。通过引入 Causal Masking,我们可以进一步获得约 2倍的加速。

Causal Attention 快速回顾

Causal Attention 是自回归语言模型的核心机制:在预测第 iii 个 token 时,只能看到位置 0 到 i-1 的信息,不能看到未来的 token。

在注意力矩阵中,这通过将上三角部分置为 −∞-\infty−∞ 实现:

         K₀  K₁  K₂  K₃
    Q₀ [ ·  -∞  -∞  -∞ ]  ← Q₀ 只能看 K₀
    Q₁ [ ·   ·  -∞  -∞ ]  ← Q₁ 只能看 K₀,K₁
    Q₂ [ ·   ·   ·  -∞ ]  ← Q₂ 只能看 K₀,K₁,K₂
    Q₃ [ ·   ·   ·   · ]  ← Q₃ 能看到所有

登录以继续阅读

这是一篇付费内容,请登录您的账户以访问完整内容。

如果您对 Causal Attention 的概念和原理还不熟悉,建议先阅读 Attention 机制详解 中的详细介绍。

本章重点关注如何在 Flash Attention 中高效实现 Causal Masking 以获得性能提升。

Block Pointer 与多维支持

从单序列扩展到 Batch/Head 并行,并使用 Block Pointer 简化指针管理。

Grouped Query Attention

实现 GQA/MQA 支持,让多个 Query Head 共享 KV,优化 KV Cache 内存占用。

目录

Causal Attention 快速回顾
Causal Masking 的性能机会
计算量减半的数学原理
可视化:跳过的计算区域
代码实现解析
修改点1: 粗粒度跳过 — 循环边界优化
修改点2: 细粒度 Mask — Block内部的边界处理
两层 Mask 的必要性
性能对比与验证
数值正确性验证
性能对比
不同序列长度的加速比
实现技巧总结
Triton 编译期常量的妙用
总结