首页
应用
关于
Search
1
Pytorch DDP
2,482 阅读
2
Pytorch 常见问题
1,515 阅读
3
视频时序切分
1,341 阅读
4
中文场景下的CLIP图文预训练
1,044 阅读
5
Semi-Supervised + Noisy Label
1,028 阅读
AIGC
AIGC Daily Papers
AIGC Fundamentals
其他
职场经验复盘
多模态理解
购房/投资
阅读
广告
Segmentation
LeetCode
Pytorch
Python
Shell
C++
常用链接
Search
标签搜索
AIGC
人工智能
论文速读
ai
视频生成
DiT
对齐
蒸馏
扩散模型
attention
transformer
视频编辑
图像生成
diffusion
基础知识
稀疏注意力
多模态
文生图
NVIDIA
llm
Jefxiong
累计撰写
203
篇文章
累计收到
8
条评论
首页
应用
栏目
AIGC
AIGC Daily Papers
AIGC Fundamentals
其他
职场经验复盘
多模态理解
购房/投资
阅读
广告
Segmentation
LeetCode
Pytorch
Python
Shell
C++
常用链接
页面
关于
搜索到
203
篇与
人工智能炼丹君
的结果
2026-10-05
AIGC 每日速读|2026-10-05|NeurIPS省七成token保孔洞SILSA
今日 AIGC 论文速览 今日共 10 篇 · 3D生成 1 篇 · 持久世界建模 2 篇 · 视频表征与3D控制 3 篇 · 个性化图像生成 1 篇 · 音频生成 2 篇 · 视频生成评测 1 篇 重点论文标题列表 SILSA(伊利诺伊大学):省七成token保住孔洞 Oneira(莫纳什大学):新物体能互动且留下后果 World Observer(KAIST):镜头外的物体也继续演化 SemanTok(Stability AI):小模型先学会视频语义 Tacit-TTS(杜比实验室):免转写克隆语音快十倍 今日论文速览 1. SILSA:省七成token保住孔洞 SILSA: Sliding-Window Slice Latents for Topology-Preserving High-Resolution 3D Generation | 伊利诺伊大学厄巴纳—香槟分校 | arXiv:2610.02201 关键词:图像生成3D, 切片潜空间, Rectified Flow, 拓扑约束 前序问题:高分辨率 3D 生成常先预测活跃体素,再补局部几何;连续曲面被拆成大量局部 token,细杆、孔洞和远距离连接容易断裂。增加体素数能补细节,却同时推高序列长度与训练显存。问题不只是模型画得像不像,还包括生成的网格是否保留原有结构。 本文贡献:SILSA 沿 x、y、z 三轴使用重叠的滑动窗口切片,每个 token 概括一段局部深度窗口。Slice VAE 将带朝向的表面采样编码成切片潜变量,再通过稀疏体积解码器重建网格;Volumetric Anchor Lattice 给三组切片共享的 3D 工作空间。切片级持久同调与相邻切片 Betti 转变损失约束连通分量和孔洞,图像条件 Rectified Flow 则一次生成整组潜变量。 论文 Figure 2:SILSA 总览。沿 x、y、z 轴以重叠切片 latent 表示形状。拓扑感知 SliceVAE 将表面编码为固定的多轴切片 latent,再通过稀疏体积上采样解码为高分辨率网格。图像条件 Rectified Flow Transformer 生成 latent,并以 Volumetric Anchor Lattice 作为共享 3D 记忆协调不同轴。VAE 训练中的切片级持久同调与 Betti 转变损失监督连通分量、孔洞和拓扑一致性。 实验效果:表 3 中固定 384 个 token,对比最紧凑基线 Dora 的 1280 个减少 70.0%;batch size 4、单张 A100 下训练显存 8.7GB,对比 Dora 的 14.6GB 减少 40.4%,端到端推理 0.34 秒/形状,对比 0.82 秒减少约 58.5%。表 1 的 PSNR 为 32.74,SparseFlex 为 30.12;表 2 的 Betti-Err 为 1.582,对比 1.743。arXiv comments 明确标注 Accepted at NeurIPS 2026。 论文 Figure 3:自然场景输入图像的 image-to-3D 生成结果。图中展示 SILSA 从不同输入得到的 3D 形状;这是定性展示,不能仅凭样例推出分布外形状可靠性。 批判点评:这些数字来自不同对照:70% token、显存和延迟以 Dora 为参照,PSNR 与 Betti-Err 的强基线是 SparseFlex,不能拼成对同一模型的全面领先。切片拓扑监督改善结构,并不提供任意形状的拓扑保证;作者也承认分布外结构及信息不足的输入视角会降低保真。值得关注的是表示方式,而不是把单物体结果直接外推到完整场景。 2. Oneira:新物体能互动且留下后果 Oneira: From Open-Ended Generation to Open-World Interaction in Video World Models | 莫纳什大学;大连理工大学;香港中文大学(深圳);牛津大学;布里斯托大学 | arXiv:2610.01614 关键词:交互式视频生成, 显式世界状态, 状态持久性, 条件视频 前序问题:能生成新房间,不代表能持续与新房间里的东西互动。镜头发现一个新苹果,模型要把它变成可操作对象;把苹果拿走之后,再回头也应看到空桌面。纯视频续写往往把交互当作局部视觉效果,缺少记录事件后果并供后续片段查询的状态载体。 本文贡献:Oneira 用显式世界状态表把视频生成与交互闭环连接。coding agent 读取当前观察和目标,定位实体、规划动作并把结果写回表;探索中发现的新物体也会注册。渲染引擎沿相机轨迹把更新后的表转成粗条件视频,视频生成器结合首帧、可选记忆帧和分块 caption 补齐外观与动作;末帧和更新后的状态进入下一段。它的贡献落在生成可互动且状态连续的视频环境。 论文 Figure 2:Oneira 单片段流程。给定输入图像与目标,coding agent A 读取世界状态表并规划动作 P_k。引擎 E 将每个动作写入表,把更新结果 T_k 沿相机路径 C_k 渲染成粗条件视频 V_k(第 3.2 节)。交互表现为物体框的删除或重新着色;被移动物体的框暂时移除,静止后从生成帧重新注册,新露出的物体也如此。视频模型 G 结合首帧、可选记忆帧、V_k 与分块 caption c_k 生成片段(第 3.3 节),末帧与更新后的表启动下一段。球体仅为示意,条件与生成帧来自同一个 Oneira 片段。 实验效果:表 1 的 InteractionBench 包含 100 个交互案例和 50 条记忆链。Oneira 的交互成功率为 0.78、目标正确率 0.82,LingBot-World-V2 分别为 0.49、0.39;记忆部分 State 通过率 0.80,MiniMax-H3 Ref2VA 为 0.68。图 3 展示三个片段组成的 30 秒链,比较执行交互后转回原处的状态。 论文 Figure 3:三个片段组成的 30 秒链的定性比较。各列上方标出当时的动作,橙色单元格表示 caption 携带交互的 chunk,插图显示 Oneira 的条件视频。红色标签标注 baseline 帧中的失败。Oneira 在指定 chunk 控制两次交互,并在相机转回时维持改变后的状态。 批判点评:状态表使“后果被记住”更明确,却没有让底层视频模型成为物理仿真器。作者列出的动作空间仍是离散的,连续控制和精细接触并未解决;长时域视觉一致性仍受生成器能力限制。100 个交互案例与 50 条记忆链是作者的专用评测,不能等同于开放世界任意操作均可靠。 3. World Observer:镜头外的物体也继续演化 World Observer: Joint Actor-Observer Generation for Persistent World Modeling | 韩国科学技术院(KAIST)人工智能学院 | arXiv:2610.02162 关键词:视频世界模型, 全景观察者, 多视角生成, 镜头外动态 前序问题:视频世界模型通常只生成行动者当前看见的画面。汽车驶出视野之后,没有可见证据约束它继续前进;再次入镜时可能丢失、停住或变成别的车。仅保存旧图像能帮助记住外观,却不一定能维护镜头外持续发生的动态。 本文贡献:World Observer 把“行动”和“观察”拆成两个同步生成流:透视 actor 负责当前视角,全景 observer 持续观察选定区域,二者由共享 DiT 联合生成。共同全景源的 warping 建立几何对应;Observer Sink 从高分辨率初始全景取四个透视参考,帮助物体重返视野时恢复细节。观察者还能放在不同位置,或接收单独条件来控制画面外的事件。 论文 Figure 4:模型总览。(a)World Observer 使用共享 DiT 联合生成 actor 与全景 observer 两条流。(b)Observer Sink 将高分辨率初始全景转换为四个透视参考,以保留精细外观信息。 实验效果:表 1 的真实/合成 OOV 测试中,单观察者 2B 模型的 FVD 为 291.7/196.2,LingBot-World 为 420.0/333.3;OOV-D_gt 为 0.426/0.531,对方为 0.337/0.340。多观察者仅在合成集测试,FVD 为 147.0,但 OOV-D_gt 为 0.528,略低于单观察者 0.531。因此视角增加改善部分指标,并非所有动态指标同步上涨。 论文 Figure 7:与其他模型的定性比较。红框标出目标离开视野前与重新进入视野后的物体。已有模型常丢失物体、冻结其状态或不一致地重建它;World Observer 在视野外间隔后更好地保持物体身份与状态演化。 批判点评:方法目前需要全景条件,作者将从普通透视图补全全景列为后续方向;真实世界多观察者条件稀缺,多观察者优势只在合成集有证据。有限的观察者预算也无法覆盖整个世界。它更像把算力分配到感兴趣区域,而不是获得无限范围的世界记忆。 4. SemanTok:小模型先学会视频语义 SemanTok: Predictable Semantic Tokens for Efficient Autoregressive Video Generation | Stability AI;卡尔斯鲁厄理工学院 | arXiv:2610.00686 关键词:视频tokenizer, 自回归生成, DINO语义监督, 可变长度token 前序问题:自回归视频生成需要先确定场景是什么,再补纹理细节。可变长度 tokenizer 虽能从粗到细输出,已有 decoder-REPA 对齐却可能利用带噪输入满足损失,未必迫使短 token 前缀本身承载语义;小 AR 模型因此仍可能生成类别不对或随时间漂移的画面。 本文贡献:SemanTok 保留 VideoFlexTok 的可变长度训练和 nested dropout,在编码端加入冻结 DINO 特征,并用轻量 Dense DINO 与 Class DINO 读出头,要求每个保留下来的 token 前缀单独预测教师语义。冻结 tokenizer 后,AR 模型按时间优先、由粗到细的顺序预测 token,再由扩散 decoder 还原视频。语义监督直接落在短前缀,而不只落在 decoder 的中间状态。 论文 Figure 2:方法总览。(1)tokenizer 训练:encoder、FSQ 与扩散 decoder 联合训练,nested dropout 使 decoder 能从任意 token 前缀重建视频。SemanTok 保留 VideoFlexTok 的训练方案(橙色),再加入冻结教师的语义监督:教师特征输入 encoder,每个保留的前缀都学习预测这些特征(紫色;细节见图 3)。(2)AR 训练与生成:冻结 tokenizer,AR 模型按时间优先、粗到细的顺序预测 token,步骤(1)的 decoder 将任意生成前缀渲染为视频。 实验效果:摘要与实验分析报告:201M 的 SemanTok AR 模型可匹配或超过约 3.4 倍规模的 VideoFlexTok AR 模型;附录表 4 中相应规模阶梯为 201M 和 679M。正文还报告在相同 AR 模型下,k=16 时每 token 预测约少用 32% 的 bits。图 7 比较同一 Kinetics-600 弹吉他类别的 rollout,201M 的 SemanTok 随时间退化更慢。 论文 Figure 7:SemanTok 在两种 AR 模型规模下保持优势。图中是同一 Kinetics-600 弹吉他类别的 class-to-video rollout。在 201M 模型下,SemanTok 随时间缓慢退化,而 VideoFlexTok 只在 t=1 时清晰。项目页提供视频。 批判点评:“小模型追上大模型”限定在作者的数据、token 预算与指标配置;参数量减少不等于端到端生成耗时同比减少。正文专门指出每帧仅 1 个 token 时不能兼顾全部目标,附录重建表也显示部分 PSNR/SSIM 逊于基线。语义更准和逐像素重建更准是两种目标,部署时仍需按预算权衡。 5. Tacit-TTS:免转写克隆语音快十倍 Tacit-TTS: From Autoregressive Decoding to Masked Prediction for Efficient Transcript-Free Voice Cloning | 杜比实验室 | arXiv:2609.38658 关键词:语音克隆, Masked LM, 非自回归生成, ReFlow蒸馏 前序问题:自回归 TTS 逐个预测语义 token,零样本音色克隆效果好但延迟高。非自回归系统更快,却常要求参考音频的转写;参考是陌生语言、婴儿咿呀或无意义发音时,ASR 生成的错误文本反而污染条件。目标是在不要求参考转写的同时降低生成耗时。 本文贡献:Tacit-TTS 从 IndexTTS2 蒸馏,将 text-to-semantic 阶段替换为 masked 非自回归生成,同时保留预训练的说话人和情绪条件。免训练长度估计结合参考音频语速与目标文本音节数,确定输出长度;连续潜表示和 ReFlow 蒸馏则加速下游 semantic-to-mel 的 flow-matching 渲染器,最后由 vocoder 合成波形。 论文 Figure 2:Tacit-TTS 流程总览。输入为参考音频与目标文本,encoder 提取条件表征;text-to-semantic 模型通过 masked-LM 生成语义 token,semantic-to-mel 模型将其转成 mel 频谱,最后由 vocoder 合成波形。 实验效果:单张 NVIDIA A100 上,作者对超过 5 秒的语音报告相对 IndexTTS2 超过 10 倍加速。图 3 的端到端测试包含参考条件计算:4 秒输出约 9.0 倍、31 秒附近最高约 14.9 倍、125 秒为 11.8 倍。表 1 中 LibriSpeech speaker similarity 为 0.875,教师为 0.870;但 WER 为 5.54,教师为 3.115。参考语言另测试了 8 种语言及非词汇声音。 论文 Figure 3:IndexTTS2 与 Tacit-TTS 的端到端生成时间(包含参考条件计算)随输出时长的变化。使用同一参考片段,以固定段落重复延长输出。虚线表示实时速度;Tacit-TTS 曲线上方数字是相对教师的加速比,插图放大 Tacit-TTS 曲线。4 秒、31 秒附近和 125 秒的示例加速比约为 9.0、14.9 和 11.8 倍,不能当作全数据平均。 批判点评:更快没有消除质量损失:普通中文测试上的音色相似度和内容准确度仍落后于教师。图 3 使用固定段落重复拉长输出,最高 14.9 倍不能写成所有语音平均加速;长序列的全注意力成本也开始侵蚀收益。四个常规数据集的其他基线质量数字来自教师论文,而跨语言与非词汇测试由作者重跑,比较范围应分开理解。 6. 4Director:完整网格让物体转身不丢脸 4Director: Controlling Video World Models with Rigid 3D Geometry | Stability AI;伊利诺伊大学厄巴纳—香槟分校 | arXiv:2610.02160 关键词:可控视频生成, 4D场景, 刚体网格, Motion Adapter 前序问题:2D 框只能指定物体在屏幕上的位置,不能完整表达深度和转向;3D 点或 blob 也缺少完整表面。镜头与物体同时移动时,模型容易每帧重新猜未观测几何,导致物体转了一半、遮挡出错或身份变化。专业创作需要把相机和物体轨迹放进统一的 3D 坐标系。 本文贡献:4Director 从输入图为每个物体一次重建 canonical mesh,每帧只施加一个指定刚体变换。网格、背景点云与相机共享坐标系,用户编排轨迹后渲染成深度视频;可训练 Motion Adapter 将几何条件注入预训练视频模型,后者补充外观、光照与非刚体动态。团队构建含 20,774 个片段的 RealCOD-Rigid,并以身份判断门控 mask IoU,形成 IG-IoU 指标。 论文 Figure 3:4Director 总览。左侧:每个物体只重建一次为完整 canonical mesh,每帧用一次刚体变换移动;网格、背景点云与相机共享坐标系,用户在该坐标系指定物体与相机轨迹,并将场景渲染成深度视频(第 3.1 节)。右侧:可训练的 Motion Adapter 分支将刚体渲染注入预训练视频生成器,由后者补充外观、光照与非刚体动态(第 3.2 节);训练使用 RealCOD-Rigid(第 4 节)。 实验效果:表 1 中 4Director 的 FVD 为 370.4,SymphoMotion 为 405.4、VerseCrafter 为 484.7;IG-IoU 为 60.4,对应 52.9 和 54.8,camera RotErr 为 3.65°。图 5 对比共同相机与物体控制,显示其更完整地跟随指定转向;图 6 另展示物体离开视野后重返的定性案例。 论文 Figure 5:相机与物体共同控制的定性比较。每个片段顶行展示输入图像和规定的轨迹,其下为四个 baseline、4Director 以及源视频按时间对齐的帧。箭头表示物体朝向,源视频中为绿色、生成视频中为红色。该组案例里只有 4Director 完成源视频的转向;baseline 保留原朝向、只转了一部分、模糊或丢失主体。 批判点评:完整网格减少逐帧几何猜测,却会把错误重建固定进控制条件。论文失败案例中,倒立舞者被当作一个刚体平移,生成结果一直保持初始姿态而没有跳舞,说明刚体轮廓可能压制应有的非刚体动作。IG-IoU 还依赖视觉语言裁判判断身份;作者的 VBench 优势多数在 2 个点以内,应避免把图中单例写成稳定保证。 7. GenCine:相机与物体在3D里一起编排 Generative Cinematographer: Composing Camera and Object Motion in 3D | 约翰斯霍普金斯大学 | arXiv:2610.02180 关键词:视频运动控制, 3D运动handle, XYZ引导图, LoRA 前序问题:同一条 2D 拖拽轨迹可能对应完全不同的 3D 运动;相机和物体一起动时,屏幕坐标更难区分“物体真的转身”和“镜头绕过去”。另一方面,整物体刚体控制过硬,人体、动物或机械臂的不同部分需要分别编排,而不必为每类对象建立物理模型。 本文贡献:GenCine 把单张图提升为可编辑的 3D 场景 scaffold,艺术家在共同世界坐标系里设置相机路径,并以局部 3D handle 控制前景不同部分。背景 XYZ、前景 XYZ 和恒定前景身份颜色图分别表达静态位置、动态位置和跨帧对应;有效性 mask 标记缺失条件。轻量 guidance branch 与 Wan 模型的 LoRA 学习这些控制,从真实视频恢复轨迹并结合合成数据训练。 论文 Figure 3:GenCine 的 3D 控制流程。首先从 RGB 图估计深度与动态 mask,再提升为供动作编排的 3D 场景。艺术家通过 Blender 等 3D 界面,在共同世界坐标系指定相机和物体运动。控制转为背景 XYZ、前景 XYZ 和前景身份图,分别指示静态场景位置、移动 handle 的位置及身份。有效性 mask 区分可用条件与缺失值,动态 mask 标出初始前景。冻结的 Wan VAE 将图编码给视频模型;训练配对从视频运动恢复的条件,推理时同样的图表达艺术家编排的控制。 实验效果:表 1 的 Camera Shooting 上,GenCine 的 TransErr 为 1.00,Wan-Move 为 1.02;受控前景 R-LPIPS 为 0.18,对方为 0.19。Object Interactions 上整体 LPIPS 为 0.35,Wan-Move 为 0.37,但整体 PSNR 为 16.54,略低于 Go-w-Track 的 16.59。图 5 给出艺术家编排运动的定性对比,体现 3D 意图与 2D 投影控制的区别。 论文 Figure 5:艺术家编排运动的定性比较。自上而下为 Living room、Bus 和 Robotic arm;每组按行展示 Wan-Move、VerseCrafter、Go-with-the-Track 和 GenCine。第一列为 3D 控制(GenCine)或共享的 2D 投影(baseline),随后四帧采样各控制区间。实线与虚线分别表示已执行和剩余目标路径,圆点表示当前目标,半透明 mask 表示投影目标区域。 批判点评:主要价值是可表达的控制接口,数值优势相当局部,不支持“全指标碾压”的结论。多个 handle 只是分段刚体近似,不能可靠表示复杂形变;单图深度估计出错、遮挡和超出归一化范围的运动都会影响控制。它与 4Director 各有取舍:局部操作更灵活,但没有完整网格带来的全表面约束。 8. PEARL:从用户历史推理个性化画面 Personalized Image Generation with Reasoning and Reflection | 范德堡大学;Adobe;马里兰大学帕克分校;佐治亚大学 | arXiv:2610.00737 关键词:个性化图像生成, 用户历史, 推理反思, DPO 前序问题:个性化图像生成通常依赖挑选好的参考图片,而真实用户偏好分散在评论、帖子、图片和 metadata 中。把某件商品放到什么环境、某个主题应采用什么风格,需要综合历史线索;检索一张最像的图片并不能完整表达生活方式与审美。 本文贡献:PEARL 将多模态 planner、reflector 与冻结图像生成器串成 reason-reflect 循环:先从历史形成场景计划并渲染,再根据实际图像与历史的不匹配修订计划。训练先蒸馏 silver personalization trajectory,再采用交替策略 DPO,用与用户历史相关的检索奖励优化两种策略。配套 benchmark 区分电商商品的 Personalized Scene Generation 和社交媒体的 Personalized Creative Generation。 论文 Figure 2:Pearl 的训练流程。先用 Silver Personalization Trajectory 做 warm start,再通过 Render-in-the-Loop 优化两种策略。图左的教师产生推理与计划供 planner 和 reflector 蒸馏;图右把两种策略与冻结渲染器串接,按奖励给候选排序、构建偏好对并做交替 DPO 更新。 实验效果:表 1 的 Amazon 场景生成中,H@5 为 0.2297,PMG 为 0.2143;MLLM Overall 为 3.92,对方为 3.48。表 2 的 Instagram 创作生成中,跨类别 R@1 为 0.876,Pigeon 为 0.748,MLLM Overall 为 3.541,对方为 3.334。图 4 将两种任务的指标归一化后汇总,比较完整 PEARL 与移除 reflection 的版本。 论文 Figure 4:消融:Pearl 与 Pearl -Reflection 的比较。该可视化对两种任务的指标进行归一化和汇总,展示保留 reflection 前后的差异;坐标是汇总后的相对值,不是原始 benchmark 分数。 批判点评:用户“喜欢”不等于模型能识别出用户:检索与 MLLM 分数只是偏好对齐的代理指标。作者明确指出电商场景没有每个用户—商品组合的自然真实目标图,只能间接评估;策略训练也只做了一轮交替 DPO。Instagram 的 LPIPS 和 MS-SSIM 并非最优,因此不能把摘要中的平均个性化提升写成所有画质指标普遍提升。 9. PLACE:让双耳声像跟随多模态条件 PLACE: Positional Latent Adaptation via Conditioned Embeddings for Binaural Audio Generation | 杜比实验室;佐治亚理工学院 | arXiv:2610.00630 关键词:双耳音频生成, 空间条件, 低秩latent适配, ILD与ITD 前序问题:音频与画面语义一致,还不一定听得出声音来自左边还是右边。现有空间声音生成器的条件接口常较固定,难以同时接收文本、视频和可选音频。想让耳机中的声像随场景变化,需要条件中的空间信息真正改变生成 latent,而不是只描述声源类别。 本文贡献:PLACE 基于 AudioX,增加 Perception Encoder Core 视频特征,将文本与视频表征对齐以提取空间线索,再用依赖条件的低秩变换调整生成 latent。适配器通过解码音频上的耳间声级差 ILD 和时间差 ITD 损失训练;五个预训练 encoder 与 SAO decoder 冻结,AudioX 通过 LoRA 微调。系统支持不同模态条件组合,而非只处理视频输入。 论文 Figure 1:PLACE 框架总览,蓝色突出视频特征路径,橙色突出文本路径。(a)AudioX 生成 latent,空间条件在冻结 SAO decoder 之前驱动其变换;五个预训练 encoder 与 decoder 全部冻结,AudioX 用 LoRA 微调。(b)投影后的 PE Core 特征在 MAF 模块之前补充 CLIP 与 Synchformer 特征。(c)文本—视频 cross-attention 与联合 self-attention 产生空间特征 H_c-sp。 实验效果:表 1 的 FAIR-Play split 1 上,ITD 误差 0.1077ms、ILD 误差 2.0495dB,ViSAGe 分别为 0.1136ms、2.1306dB;split 2 的 ITD 却较差。表 3 的分布外 T2A 听测平均分为 55.5,SpatialSonic 为 44.8;域内 V2A 为 65.3,对方 ViSAGe 为 44.9。原文只写 Submitted to ICASSP 2027,不能称为会议接收。 论文 Table 3:主观听测表。评分范围为 1–100,越高越好;SC 为语义一致性,SI 为空间感,SP 为空间一致性,Sy 为同步性(T2A 不评此项),AR 为音频真实感,Avg 为所评维度的均值。GT 是真实音频参考,不参与最佳模型加粗;粗体表示每种条件各列最优模型。表中区分域内与分布外 T2A、域内与 Gemini 视频条件 V2A;n 统计评分数。PLACE 在域内 T2A 均分 40.5,分布外 T2A 为 55.5,域内 V2A 为 65.3,Gemini V2A 为 58.4。 批判点评:收益取决于模态与分布。域内 T2A 听测均分仅 40.5,低于 SpatialSonic 的 44.4 和 AudioX 的 41.1;BEWO-1M SS-set 上 FSAD、ITD、ILD 也逊于 SpatialSonic,优势集中在原文定义的 SpatialCLAP 差值指标。听测表中的 n 是评分数,不能当作独立听众人数。空间一致性进步不等于全面音质领先。 10. VTR-Bench:视频逼真了字却还常写错 VTR-Bench: A Systematic Benchmark for Evaluating Visual Text Rendering in Video Generation | 香港城市大学;香港科技大学(广州);香港科技大学;西湖大学;电子科技大学 | arXiv:2610.01499 关键词:视频生成评测, 视觉文字渲染, WER, 关键帧引导 前序问题:广告、科学演示和界面视频里的文字常承担核心信息。模型能把灯光与动作做得漂亮,招牌仍可能拼错、缺字或小字不可读;只评审美与物理合理性无法发现这些问题。需要把文本正确性与画面任务完成度分开衡量。 本文贡献:VTR-Bench 用 300 个 prompt 覆盖五类场景,对具体文字载体分别转写并计算 WER,同时用每条 prompt 的 20 问 checklist 检查场景与运动要求。作者另外提出关键帧引导生成流程,由 Director agent 根据视觉反馈反复调整图像、首帧及运动计划,再生成视频。主贡献是生成模型评测与改进,符合 AIGC 范围,而不是通用视频理解。 论文 Figure 2:VTR-Bench 总览。左侧:专家设计五种场景类别,引导场景 seed 生成与人工过滤;保留的 seed 用于构造 prompt,再通过参考图生成和 VLM 审查迭代完善,最后人工复核。右侧:生成视频在两个维度分别评估,Video Score 以 20 问 checklist 衡量场景与视频要求完成度,WER 比较针对具体文字载体的 VLM 转写与参考文字,衡量视觉文字保真。 实验效果:主表 1 比较 11 个视频模型,Wan3.0 的整体 WER 最低,为 0.250,Video Score 为 0.849;MiniMax H3 为 0.447/0.756,Seedance2.5 为 0.641/0.791。表 2 中 Qwen3.8-27B 对文字块 WER 的 Pearson 为 0.9541,视频级为 0.9842,支持其与人工转写的相关性。官方代码:VTR-Bench。 论文 Figure 1:即使视频视觉上可信,现有视频模型仍难正确渲染文字。上方是 Seedance2.5 生成的视频;Video Score 是 20 问 checklist 中满足要求的比例,WER 是相对参考文字的词错误率。左下展示视频中的四种文字渲染问题,右下展示 11 个模型在 VTR-Bench 上的整体 WER。 批判点评:WER 是词级编辑距离,0.250 不等于“25% 的视频失败”,也不能推出每个视频都错四分之一。不同模型的分辨率与生成管线不同,排行榜还受评审模型影响;相关性高不意味着每条转写都正确。关键帧 agent 的迭代成本也应单独核算,不能把更高预算的改进当作单次生成能力。 趋势观察 从像素记忆走向状态载体 Oneira 显式记录交互后果,World Observer 持续生成选定区域的视角;两条路线都让镜头外状态获得约束,但尚未证明无限时域或全面物理可靠性。 控制接口先解决几何歧义 4Director 用完整刚体网格,GenCine 用局部运动 handle;它们共同把相机和物体放进世界坐标系,分别在完整几何与形变灵活性之间取舍。 短序列必须携带可用信息 SILSA 的切片潜变量与 SemanTok 的短语义前缀都减少局部 token 负担;这是不同任务中的表示设计趋势,不能把各论文加速数字横向相乘。 评测拆开能力与代价 视频里的字、音频的空间感、用户历史匹配和端到端耗时都需要独立指标。各论文测试集、评分口径和硬件不同,本期不做跨论文统一排行榜。 今日讨论 SILSA 用切片拓扑监督改善单物体孔洞与连通性;面对输入视角不充分或分布外形状,你会优先增加多视角条件,还是设计可验证的拓扑约束? 人工智能炼丹君 整理 | 2026-10-05
2026年10月05日
0 阅读
0 评论
0 点赞
2026-10-05
AIGC 基本功|扩散模型的跨步缓存复用-DiffCache
AIGC 基本功|扩散模型的跨步缓存复用-DiffCache 同样走完采样器的每一步,能不能少算几次昂贵的去噪网络?跨步缓存把这个问题拆成两件事:复用什么,以及什么时候必须刷新。本文用残差缓存的最小实验解释机制,再对照 DeepCache 和 TeaCache 的真实实现。实验是人工构造的数值系统,不是视频模型的质量或速度复现。 本文承接 DDIM 与高阶采样器、DiT 和 性能建模与 Profiling。默认读者知道采样器会反复调用网络,不要求先读缓存论文。论文版本与官方代码核对日期为 2026-10-05。 01. 为什么需要它 一次视频生成中,潜变量在逐步变化,文本条件通常固定,网络却要一遍遍穿过相同的模块。假如两次调用的深层特征几乎一样,重新计算这部分就可能浪费时间。但“几乎一样”只是观测,不能推出“任意一次都可以省”。运动切换、条件变化和采样末期的细节修整,都可能让陈旧特征的误差显现。 这里最具体的失败场景是固定间隔缓存:先完整算一次,之后连续复用,直到计数器要求刷新。如果变化突然出现在间隔中间,策略不会因为内容变了而提前刷新。反过来,如果某段变化很慢,固定间隔也会做多余的完整计算。问题在于日历式排期没有感知当前输入,周期短会损失收益,周期长会增加近似误差。 本文实验保留全部 40 次状态更新。精确版本执行 40 次完整残差计算;自适应阈值为 0.16 时,只执行 11 次,复用 29 次。终点相对误差为 0.008222。它证明了在这个特定数值系统中确实可以减少计算次数,同时也证明输出发生了变化。它不证明视频质量达标,更不证明真实 GPU 上能得到相同的提速。 缓存适合回答“这个步骤里,哪些计算可以暂时沿用”。减少采样步数则改变采样器的离散路径。两种方法可以组合,但组合后的误差要重新评估:单独用缓存通过了一组样例,不意味着换成少步采样器后仍然通过。 02. 最小可用理解 第一,把网络拆成每步仍要计算的部分,以及准备复用的昂贵部分。第二,保存上次完整计算得到的特征或残差,用便宜的变化指标决定是否刷新。第三,即使命中缓存,当前输入仍然要参与输出,采样器仍然继续更新状态。 以残差形式为例,设当前嵌入是 $h_i$,完整模块输出为 $G(h_i,t_i,c)$。缓存保存的是 $R_i=G(h_i,t_i,c)-h_i$。如果上次刷新发生在 $r$,当前近似模块输出是 $h_i+R_r$。符号 $i$ 表示第几次网络调用,$t_i$ 是这次调用的时间条件,$c$ 是固定的文本等条件;$r$ 是缓存的生成时刻,通常小于当前 $i$。 这三句话的重点是“当前输入加旧残差”。如果直接把整个旧输出拿回来,连当前输入的变化也被抹掉了。残差复用仍然有误差,只是保留了输入的直通更新。DeepCache 的 U-Net 高层特征缓存并不等于这个残差写法,不能把所有缓存方法都称为同一个算法。 03. 数学推导 先定义缓存误差发生在哪里 设 $E$ 是便宜的输入嵌入,$G$ 是昂贵模块,$H$ 是输出头。完整调用写为: $$h_i=E(x_i,t_i,c),\quad R_i=G(h_i,t_i,c)-h_i,\quad y_i=H(h_i+R_i,t_i,c)$$ $x_i$ 是采样器当前的潜变量;$h_i$ 是网络内部的嵌入,不必与 $x_i$ 形状相同;$R_i$ 是昂贵模块对嵌入的修正;$y_i$ 是送回采样器的网络预测。这里的预测可以按模型定义代表噪声或速度,缓存逻辑本身不替模型选择参数化。等式只是代数分解,不要求模型专门训练一个名叫“残差缓存”的层。 命中缓存后,输入嵌入和输出头仍按当前条件计算: $$\widehat{y}_i=H(h_i+R_r,t_i,c),\quad r<i$$ 帽子表示近似计算。比较精确与近似时,先固定当前同一个 $h_i$:误差来自 $R_r$ 替代了 $R_i$,而不是来自状态已经分叉。若输出头在所讨论的局部区域满足 Lipschitz 条件,常数记作 $L_H$,那么: $$\|\widehat{y}_i-y_i\|\leq L_H\|R_r-R_i\|$$ Lipschitz 条件的意思是输入变化不能被这个局部映射无限放大;$L_H$ 是放大上限。本文没有测出真实模型的 $L_H$,这个式子是带假设的分析工具,不是线上质量保证。它告诉我们,缓存内部误差小仍然需要考虑输出头的敏感性。 为什么使用便宜的变化代理 直接计算 $\|R_i-R_r\|$ 能判断缓存是否陈旧,但得到 $R_i$ 就已经做了昂贵计算。为了节省计算,需要先观察一个便宜的量 $m_i$。它可以由当前嵌入与时间条件构造。相邻调用的相对 L1 变化写为: $$d_i=\frac{\operatorname{mean}|m_i-m_{i-1}|}{\max(\operatorname{mean}|m_{i-1}|,\varepsilon)}$$ $\operatorname{mean}$ 对张量所有元素取平均,绝对值逐元素计算,$\varepsilon$ 是防止除零的小正数。两次张量形状一致时,均值之比等于 L1 范数之比。这个值没有单位,反映变化相对之前幅度有多大。分母很小时,它也可能变得敏感,不能把这种数值问题误判成内容剧烈变化。 本文教学实现取 $\varepsilon=10^{-12}$。该保护和后面非负累计属于教学实现的防御性处理;不能据此声称官方代码逐字采用相同逻辑。 最直觉的想法是:如果每一步变化都很小,就一直复用。但缓存对应的是刷新时刻 $r$,不是上一时刻。利用三角不等式: $$\|m_i-m_r\|\leq\sum_{j=r+1}^{i}\|m_j-m_{j-1}\|$$ 这里 $j$ 遍历刷新之后的调用。即使单步变化很小,连续多步也会积累成明显差异,因此策略应保留“自上次刷新以来的变化预算”。这个式子约束绝对变化;把每项分别除以不同的幅度后,不能直接称相对变化之和为严格误差界。 实际决策可以采用累计代理: $$A_i=A_{i-1}+g(d_i),\quad \text{refresh if } A_i\geq\delta$$ $A_i$ 是累计预算,$g$ 将输入变化映射到预估的输出变化,$\delta$ 是刷新阈值。完整计算后把预算归零并保存新残差。TeaCache 使用多项式重标定来改善这种代理;拟合关系是经验估计,不会把统计相关性变成对所有样本成立的数学上界。TeaCache 原文第 3.2–3.3 节 给出了指标与累计决策。 本文取 $g(d)=d$,以便只观察状态机。阈值 0.16 因而只属于本实验的代理尺度,不能迁移成真实模型的推荐阈值。换掉嵌入幅度、时间调制或采样路径,同一个阈值对应的刷新频率就会改变。 一次局部近似如何影响后续轨迹 现在让两条轨迹分别使用精确和近似输出。为便于推导,考虑显式 Euler 更新: $$x_{i+1}=x_i+\Delta s_i v(x_i,s_i),\quad \widehat{x }_{i+1}=\widehat{x }_i+\Delta s_i\widehat{v}(\widehat{x }_i,s_i)$$ $s_i$ 是积分坐标,$\Delta s_i>0$ 是步长,$v$ 是精确向量场,$\widehat{v}$ 使用缓存近似。令 $e_i=\|\widehat{x }_i-x_i\|$ 表示轨迹差异,令 $\eta_i=\|\widehat{v}(\widehat{x }_i,s_i)-v(\widehat{x }_i,s_i)\|$ 表示在近似轨迹当前状态上的局部缓存误差。 先在更新差值中加减 $v(\widehat{x }_i,s_i)$,把“缓存误差”和“输入状态不同”分开。再使用三角不等式,并假设 $v$ 对状态的 Lipschitz 常数为 $L$,得到: $$e_{i+1}\leq(1+\Delta s_iL)e_i+\Delta s_i\eta_i$$ 若初始状态相同,则 $e_0=0$。重复代入可得: $$e_N\leq\sum_{i=0}^{N-1}\Delta s_i\eta_i\prod_{j=i+1}^{N-1}(1+\Delta s_jL)$$ 空乘积取 1。这个展开说明早期误差会通过后续更新传播;最后一次强制刷新只能去掉最后一次局部缓存误差,不能恢复此前已经走偏的轨迹。真实系统可能局部收缩,实际误差小于这个上界;高阶求解器还涉及历史预测,不能把 Euler 的式子直接冒充所有采样器的误差定理。 为什么命中率不能直接当加速比 令 $N$ 是总调用数,$K$ 是完整计算数,$C_f$ 是完整调用成本,$C_h$ 是缓存命中调用成本,$C_o$ 是其他固定开销。则一个简单的成本模型是: $$S=\frac{NC_f+C_o}{KC_f+(N-K)C_h+C_o}$$ 命中调用仍要嵌入、判定、读缓存、运行输出头和更新采样器,$C_h$ 不会自动等于零。缓存张量跨设备读取甚至可能很贵。端到端还包括文本编码、VAE 解码和输出处理,跳过 Transformer 并不会同时省掉这些部分。 本文把 $C_f=1$、$C_h=0.08$、$C_o=0$ 当成教学假设。按实跑统计的 $K$ 代入可以算成本比,但没有用这些假设伪装 GPU 计时。部署时需要测出真实成本,再决定缓存是否值得加入。 04. 代码实现 完整脚本是文末的 cache_demo.py,依赖 NumPy 与 matplotlib,无需模型权重。它构造一个四层非线性残差函数,用固定随机种子产生输入和矩阵。输入形状是 $(B,L,D)=(1,4,16)$:一个批次、四个位置、十六个通道。张量使用 float64,目的是让读者稳定复核误差,不模拟混合精度推理。 脚本里的 cheap_input 对应 $E$,residual 对应 $R$,proxy 对应 $m$,budget 对应 $A$,threshold 对应 $\delta$。输出头取恒等映射。每次用当前 $h$ 加缓存残差得到预测,然后执行状态更新。时间标签从 1 递减到 0,而状态更新的步长固定为 $1/N$;这只是合成系统的定义,没有冒充 DDIM 或真实扩散调度器。 决策的核心代码如下;变量准备、精确对照、绘图及断言都在完整附录中: budget += max(change, 0.0) boundary = i in (0, STEPS - 1) calculate = boundary or budget >= threshold if calculate: saved_residual = residual(h, t) budget = 0.0 output = h + saved_residual 首次调用没有缓存,所以必须完整算。最后一次强制刷新属于本实现的策略。previous_proxy 每一步都更新,saved_residual 只有刷新时更新:把这两件事混淆,就会把“相邻变化累计”写成另一种算法。simulate 每次创建新的状态,避免后一条请求接着用前一条请求的残差。 为了画误差曲线,脚本还在每个步骤额外执行一次精确函数作为离线诊断。这些额外执行没有算进策略的 full_calls;因此实际运行这个带诊断的脚本并不是加速测速。计数回答的是策略本来会执行多少次昂贵计算,图回答的是当前近似轨迹上的局部输出误差有多大。 真实运行环境为 Python 3.11.9、NumPy 2.4.3。基线终点的元素范围为 [-0.791508, 0.614856]。输出如下: shape=(1, 4, 16), steps=40, seed=7 baseline_range=[-0.791508, 0.614856] name full_calls hits endpoint_relative_error assumed_speed_ratio exact 40 0 0.000000 1.000 adaptive_0.08 18 22 0.004939 2.024 adaptive_0.16 11 29 0.008222 3.003 adaptive_0.32 7 33 0.026326 4.149 uniform_4 11 29 0.008672 3.003 All assertions passed; figure=figures/cache_schedule.png 终点相对误差的定义是近似终点与精确终点之差的 L2 范数,除以精确终点的 L2 范数。因此 0.008222 是这个定义下约 0.8222% 的差异,不是 VBench 下降,也不是图像中有 0.8222% 的像素出错。最后一列来自前节假设的成本模型,没有真实延迟单位。 图上半部分每个方块表示一次完整残差计算,空缺表示复用;下半部分是离线诊断得到的局部输出 L2 误差。请先观察刷新是否集中在同一段,再看复用之间误差如何变化。曲线没有用感知质量指标,纵轴不能解读成画面可接受程度。配图由附录脚本生成,数据与上述输出来自同一次实验。 阈值 0.16 与固定间隔 4 都执行了 11 次完整计算,但终点误差略有差异。这是“同样预算也会因刷新位置不同而得到不同结果”的一个实例。它只有一个种子和一个人工系统,不能证明自适应策略在所有输入上优于固定间隔,更不能作为统计显著性结论。 05. 工业级实现对照 DeepCache 原文第 3.3 节 利用 U-Net 的分支结构:当前浅层特征继续更新,较深分支的高层特征从缓存取回并拼接。它缓存的是特定网络位置的特征,并非本文抽象的整段残差。官方 horseee/DeepCache 的 deepcache.py 提供 forward 包装与状态复位;对照时应先找缓存边界,再判断到底省掉了哪些模块。 TeaCache 的 HunyuanVideo 示例 中,teacache_forward 用时间调制后的输入计算变化代理,累计重标定后的变化,并保存 Transformer 主体的输入输出残差。首尾调用完整计算,命中后把旧残差加到当前图像嵌入上;输出层仍然运行。这里对应的是 2026-10-05 实际取回的源码,后续上游重构需要重新核对。 教学版与上述示例的差异很明确:教学版没有文本分支、真实 patchify、注意力或权重,也没有拟合官方多项式;它只验证预算与残差状态机。真实代码还可能把张量变化转成 CPU 标量,这种同步的代价不能从“指标计算量很小”中自动排除。需要把判定开销也纳入测量。 工程上,缓存状态应属于一次请求,而不是一个可被并发请求共享的随意全局变量。prompt、负向条件、guidance、分辨率、帧数、模型权重或调度器变化时,复用旧状态没有语义依据。只检查张量形状远远不够;两个请求可以形状相同、内容完全不同。失败中断后重新开始生成,也要建立新的缓存生命周期。 若模型采用双分支 CFG,必须考虑条件与无条件分支的缓存分别有效。按 $y=y_u+w(y_c-y_u)$ 写,$y_u$ 与 $y_c$ 分别是无条件和条件预测,$w$ 是 guidance 系数。预测误差满足: $$\|\Delta y\|\leq |w|\|\Delta y_c\|+|1-w|\|\Delta y_u\|$$ 因此较大的 guidance 可能放大分支近似误差。这个式子是代数上的工程分析,不声称所有视频模型都执行双分支 CFG;有些采用蒸馏后的 guidance 条件。接入前要先确认实际预测路径,再决定状态如何隔离。 还要区分采样步和网络调用。调度器可能在一个名义步里求值多次,或重访相同的时间标签。缓存的计数器应跟随真实调用语义。如果只拿界面上显示的步数当数组长度,很容易在错误的位置刷新,或者提前把状态清零。应记录每次实际调用的时间标签、是否刷新和缓存来源,再与调度器对齐。 06. 代价与边界 省的是昂贵模块重复计算,增加的是缓存内存、生命周期管理、判定开销和近似误差。保存一个形状为 $(B,L,D)$ 的残差张量,最低存储量为 $BLDb$ 字节,其中 $b$ 是每个元素的字节数;克隆的输入、多个分支及其他缓存还会额外占空间。长视频中 token 数很大,省算力并不自动省显存。 阈值没有通用刻度。代理的归一化方式、拟合数据、模型层、时间条件与采样器都影响其意义。在某模型上有效的多项式可能在另一个模型或调度器上失配。应在目标分布上先记录输入代理与真实残差变化的对应关系,再评估误判,不能只沿用别人仓库里的常数。 端点刷新是一种保守策略,不是一种证明。即使最后一次完整计算,网络读到的潜变量也已经经过缓存轨迹;它不能把积累误差自动洗掉。判断质量时要看整段生成结果,尤其是动作连续性、物体身份、细节与提示词对应关系,而不是只看最后一个调用的缓存误差归零。 已有很少采样步的蒸馏模型可能没有多少安全的复用空间。短轨迹上每次预测都更重要,进一步减少完整网络调用可能迅速损害质量。动态改变条件的交互式生成、频繁切换控制信号的任务,也不应默认共享同一段缓存。先建立可靠基线,才能判断具体组合是否值得部署。 评测时固定权重、prompt、种子、分辨率、帧数、采样器和 guidance,分别测完整网络计算时间与端到端延迟。预热、设备同步、重复测量和尾延迟都要说明;不要把首次编译的基线与已预热的缓存版放在一起比较。质量样例应包含快速运动、镜头切换、文字、遮挡和多主体交互等容易暴露近似问题的内容。 本文给出的误差曲线只能帮助理解缓存的数值行为。没有真实模型权重,没有样本集,也没有测量硬件延迟。因此本文不提供生产阈值,不给出“无损”的结论,也不声称这些 toy 误差能预测真实视频的主观质量。 还要分别检查代理误判的两种代价。代理高估变化会增加完整计算,主要损失效率;代理低估变化会继续使用陈旧残差,主要增加质量风险。这两种错误不能只用一个平均相关系数概括。即使整体相关性很高,少数快速变化的步骤被漏掉,也可能成为最明显的失败样例。应保留逐步记录,把异常视频定位到具体刷新区间,再判断是代理不敏感、阈值过大还是缓存位置选错。 为避免测试集上的阈值被过度优化,可以先用一组提示词选择阈值,再在另一组提示词上固定参数评估。报告平均质量以外,还要记录最差样例以及变化最大的场景。缓存有时能带来某项评分的小幅上升,但近似轨迹改变本身并不保证改进;应通过更多独立样本判断这是随机波动还是稳定效果。不能挑出一段更好看的视频,就把整条路线称作质量提升技术。 缓存决策在不同请求中可能产生不同的完整计算次数,所以服务端延迟也会随内容变化。只报告平均命中率,会遮蔽慢请求的行为;应同时关注完整计算次数的分布和尾延迟。如果业务要求固定延迟,可以限制最长复用间隔或最低刷新次数,但这些额外约束需要重新纳入成本账本与质量验证。 07. 经典论文脉络 DDIM:Denoising Diffusion Implicit Models,2010.02502。它提供理解不同采样路径的基础;跨步缓存要先说明是否保留原有离散调用,再比较误差,不能把省模块与改步长混为一谈。 DeepCache,2312.00858。它把 U-Net 高层特征的跨步复用转成免训练加速方案,主要启发是选择合适的结构边界,而不是对所有中间张量一概缓存。 Pyramid Attention Broadcast,2408.12588。它把视频生成中的注意力复用纳入设计,提示我们不同模块可以有不同刷新节奏。本文只把它作为技术脉络,不复现其性能结果。 TeaCache,2411.19108。它用时间调制输入的变化代理和重标定来选择刷新时刻,把问题从固定周期推进到输入感知的缓存决策。 论文数值也要带条件。DeepCache 的摘要报告 SD v1.5 上 2.3 倍加速与 CLIP Score 下降 0.05;这是作者特定设置下的结果,不能泛化成所有扩散模型都同样受益。原始摘要 支持这组数字。 TeaCache v2 的表 1 中,OpenSora-Plan、65 帧、512×512 的基线延迟为 99.65 秒,slow 配置为 22.62 秒,对应 4.41 倍;VBench 从 80.39% 到 80.32%,差值是 0.07 个百分点。fast 配置在同表报告更高加速,但质量也有不同取舍,所以本文引用 slow 配置,避免把不同配置的最优速度与最优质量拼成一个结果。原文表 1 可直接复核。 这些工作属于相关路线,没有一个统一的阈值或通用缓存张量。先检查缓存边界,再看刷新规则,最后检查采样路径是否改变,是比较它们时更稳定的顺序。方法名称相似、都声称免训练,不代表工程实现可以直接互换。 08. 常见误解 “跨步缓存与 KV cache 是同一种精确复用。” 自回归 KV cache 通常复用固定前缀的计算结果,而扩散过程中潜变量和时间条件在变化。这里复用的是可能陈旧的近似特征,必须单独评估误差。是否等价还取决于模型路径,不能仅凭有个 cache 字段就判断。 “相邻特征相似,所以可以一直复用。” 相邻相似不等于相隔很多步仍相似,更不等于输出质量不变。累计预算关注缓存年龄,刷新机制关注何时恢复完整计算。二者缺一不可。 “阈值 0.16 表示允许 16% 的质量损失。” 阈值只是累计代理的比较值。它不是终点误差、FID、VBench 或主观质量的直接比例。本实验阈值与官方多项式的阈值甚至不在相同刻度上。 “缓存命中后跳过了采样器更新。” 本文所有版本都更新状态 40 次。省掉的是一部分网络计算,当前输入与输出头仍然可以变化。如果代码命中后直接跳过调度器,它执行的是另一种算法。 “最后一次完整计算能补回前面所有误差。” 最后一次局部预测变准确,不等于输入轨迹回到基线。误差传播推导和实验都把这两件事区分开。 “调用次数少三倍,就是端到端快三倍。” 命中步骤、其他固定模块、同步与内存操作仍有成本。本文最后一列只是显式假设的成本比,实际脚本还带额外离线诊断,因此不能拿它的运行时间当加速报告。 09. 动手验证 先把附录完整代码保存为 cache_demo.py,准备 Python 3.10 以上环境,执行: python -m pip install numpy matplotlib python cache_demo.py 脚本按照所在目录生成实验结果文件和 figures/cache_schedule.png。数值应与第 04 节一致;换 NumPy 或底层数值库时,浮点尾数可能略有不同。比对完整计算次数、刷新时刻和误差到所列精度,比要求截图或图片压缩字节完全一致更合适。 首先把阈值设成 0。每个步骤都会刷新,完整计算次数应等于 40,终点应与精确基线逐元素一致。附录已经用断言验证这个条件;如果失败,要先修缓存状态机,不要调阈值掩盖问题。 其次比较 0.08、0.16 和 0.32。在本次固定实验里,完整计算次数逐步减少,终点误差逐步增加。这是一个具体样本的趋势,不是所有系统的单调性定理:刷新位置改变、误差相互抵消、轨迹分叉,都可能让某个中间阈值的结果反而更接近基线。 再比较 adaptive_0.16 与 uniform_4。两者的成本模型使用相同的 11 次完整计算,但时刻并不相同。观察图中误差峰值,不只盯终点一个数。如果想研究预算分配,固定计算次数,再调整刷新位置,比仅仅换一个阈值更容易把原因看清。 最后修改种子或残差函数中时间变化的频率,重新运行整套对照,并保留新的输出。这一操作没有本文预先承诺的性能结果;它用来观察原先代理是否仍能判断残差变化。若输入代理很平缓、真实残差突然变化,就能构造“指标漏报”的反例。把反例加入测试,比只展示一条顺利的曲线更有价值。 实际接入模型之前,还应验证两次独立请求是否有相同基线行为、错误中断是否清理状态、同形状不同条件是否互相污染。本文实验用每次 simulate 独立初始化来验证复位;真实服务还需要针对并发与重入的检查。 10. 延伸阅读 先从 DDIM 与高阶采样器 理解调用路径,再用 DiT 找嵌入、主体和输出头的边界,最后通过 性能建模与 Profiling 实测哪些部分值得省。 可以对照 KV cache 理解缓存有效性,也可继续读 量化 比较另一种减少推理成本的方法。知识树中的少步蒸馏节点仍为 planned;它是下一条相关学习路线,尚无已发布链接。跨步缓存、量化与蒸馏组合时,应重新建立质量和延迟基线。 附录:完整代码 09 节用到的脚本全文如下(cache_demo.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 cache_demo.py """Deterministic residual-cache toy; Python 3.10+, pip install numpy matplotlib. Run: python cache_demo.py No trained diffusion weights; cost ratios below are accounting assumptions. """ from pathlib import Path import json import platform import numpy as np import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt STEPS = 40 SHAPE = (1, 4, 16) SEED = 7 rng = np.random.default_rng(SEED) X0 = rng.normal(size=SHAPE).astype(np.float64) W = rng.normal(size=(16, 16)) / 8.0 V = rng.normal(size=(16, 16)) / 8.0 def cheap_input(x, t): return x + 0.05 * np.sin(2 * np.pi * t) def proxy(h, t): return h * (1 + 0.3 * np.sin(2 * np.pi * t)) + 0.1 * np.cos(2 * np.pi * t) def residual(h, t): z = h for _ in range(4): z = np.tanh(z @ W) return 0.15 * h + 0.4 * (z @ V) + 0.08 * np.sin(6 * np.pi * t) def simulate(threshold=None, period=None): x = X0.copy() previous_proxy = None saved_residual = None budget = 0.0 refreshed, local_errors = [], [] for i in range(STEPS): t = 1.0 - i / (STEPS - 1) h = cheap_input(x, t) m = proxy(h, t) change = 0.0 if previous_proxy is None else float( np.mean(np.abs(m - previous_proxy)) / max(np.mean(np.abs(previous_proxy)), 1e-12)) budget += max(change, 0.0) boundary = i in (0, STEPS - 1) if threshold is None and period is None: calculate = True elif period is not None: calculate = boundary or i % period == 0 else: calculate = boundary or budget >= threshold if calculate: saved_residual = residual(h, t) budget = 0.0 output = h + saved_residual # Offline diagnostic only; extra evaluation is excluded from work accounting. exact = h + residual(h, t) local_errors.append(float(np.linalg.norm(output - exact))) refreshed.append(bool(calculate)) previous_proxy = m.copy() x = x - output / STEPS return {"x": x, "refreshed": refreshed, "local_errors": local_errors} def experiment(): baseline = simulate() rows = [] results = [] settings = [("exact", None, None), ("adaptive_0.08", 0.08, None), ("adaptive_0.16", 0.16, None), ("adaptive_0.32", 0.32, None), ("uniform_4", None, 4)] for name, threshold, period in settings: result = simulate(threshold, period) calls = sum(result["refreshed"]) relative_error = float(np.linalg.norm(result["x"] - baseline["x"]) / np.linalg.norm(baseline["x"])) # Full call = 1 unit; cache hit = 0.08 unit, solely for illustration. assumed_work = calls + (STEPS - calls) * 0.08 rows.append({"name": name, "full_calls": calls, "hits": STEPS - calls, "relative_endpoint_error": relative_error, "assumed_speed_ratio": STEPS / assumed_work, "refresh_indices": [i for i, flag in enumerate(result["refreshed"]) if flag]}) results.append(result) zero = simulate(threshold=0.0) assert np.array_equal(zero["x"], baseline["x"]) assert sum(zero["refreshed"]) == STEPS assert all(r["refreshed"][0] and r["refreshed"][-1] for r in results) assert np.array_equal(simulate(threshold=0.16)["x"], results[2]["x"]) assert all(np.isfinite(r["x"]).all() for r in results) return rows, results, baseline def make_figure(rows, results, destination): fig, axes = plt.subplots(2, 1, figsize=(11, 7), constrained_layout=True) colors = ["#2563eb", "#ea580c", "#16a34a", "#dc2626", "#9333ea"] for y, (row, result) in enumerate(zip(rows, results)): idx = np.flatnonzero(result["refreshed"]) axes[0].scatter(idx, np.full(len(idx), y), marker="s", s=38, color=colors[y]) axes[0].set_yticks(range(len(rows)), [r["name"] for r in rows]) axes[0].set_xlim(-1, STEPS) axes[0].set_xlabel("Sampling call index (0-based)") axes[0].set_title("Squares = full residual computation; gaps = cache reuse") for y, (row, result) in enumerate(zip(rows[1:], results[1:]), start=1): axes[1].plot(result["local_errors"], label=row["name"], linewidth=1.7, color=colors[y]) axes[1].set_xlabel("Sampling call index (0-based)") axes[1].set_ylabel("Local output L2 error") axes[1].set_title("Offline diagnostics on each cached trajectory; no quality metric") axes[1].legend(ncol=2, fontsize=9) axes[1].grid(alpha=0.25) fig.suptitle("Residual caching: refresh decisions and approximation error", fontsize=14) fig.savefig(destination, dpi=150) plt.close(fig) if __name__ == "__main__": rows, results, baseline = experiment() folder = Path(__file__).resolve().parent.parent (folder / "figures").mkdir(parents=True, exist_ok=True) make_figure(rows, results, folder / "figures" / "cache_schedule.png") record = {"seed": SEED, "steps": STEPS, "shape": list(SHAPE), "python": platform.python_version(), "numpy": np.__version__, "baseline_min": float(baseline["x"].min()), "baseline_max": float(baseline["x"].max()), "results": rows, "assertions": "zero-threshold exactness, endpoints, request reset, finite outputs passed"} (folder / "experiment_results.json").write_text( json.dumps(record, ensure_ascii=False, indent=2) + "\n", encoding="utf-8") print(f"shape={SHAPE}, steps={STEPS}, seed={SEED}") print(f"baseline_range=[{record['baseline_min']:.6f}, {record['baseline_max']:.6f}]") print("name full_calls hits endpoint_relative_error assumed_speed_ratio") for row in rows: print(f"{row['name']} {row['full_calls']} {row['hits']} " f"{row['relative_endpoint_error']:.6f} {row['assumed_speed_ratio']:.3f}") print("All assertions passed; figure=figures/cache_schedule.png")
2026年10月05日
4 阅读
0 评论
0 点赞
2026-10-04
AIGC 每日推送目录
关注公众号,不错过每一篇 每天精选 AIGC 论文速读、深度解读与周末专题,把「值得读」的那几篇送到你面前。扫码或微信搜一搜 人工智能炼丹君,第一时间收到更新。 微信搜一搜 · 人工智能炼丹君 下方为博客归档目录:点击日期即可进入对应推送全文。 当前共收录 120 期,按时间倒序排列。 2026 年 10 月 2026-10-05|NeurIPS省七成token保孔洞SILSA 2026-10-04|商汤小模型绕圈反超6.5倍大模型,Looped-DiT 2026-10-02|NVIDIA砍掉视觉编码器,PixelUMM图像视频同模 2026-10-01|一次蒸馏插遍54个模型,LongLive-Plug 2026-10-01|深度解读|字节Seed×UCSD VSA2|砍掉95%注意力计算 720p端到端快4.62倍 2026 年 9 月 2026-09-30|港科大先写谱再唱,YuE2 逼近 Suno 2026-09-29|阿里双Agent语音涨10.4分,Qwen-Audio 2026-09-28|阿里 397B 只改提示词就涨 50 分,WanPE 2026-09-25|快手开源修图登顶,GEdit仍输GPT 2026-09-24|阿里Qwen3.8-Omni发布,视频推理输Gemini 2026-09-23|NVIDIA像素扩散FID 1.46,仍输latent 2026-09-22|阶跃星辰全双工98.9登顶,清华W4A4逼近FP16 2026-09-21|伯克利线性注意力让H3视频快14.5倍-VDN 2026-09-18|高通揪出长视频KV错配-Recency Forcing 2026-09-17|Meta一个模型既写歌又改歌-SongCraft 2026-09-16|22帧视频377毫秒-LynnReal-Omni 2026-09-15|阿里五路音频控制-DiffSynth-Music 2026-09-14|SenseNova 8B原生视觉直出4K-SN-U15 2026-09-11|阿里13模型不会专业剪辑-CutCraft 2026-09-10|90%稀疏视频DiT加速2.63倍-RoLA 2026-09-09|阿里两步口型配音跑到7.13FPS-TBDub 2026-09-08|蚂蚁从零训6B开源生图-LLaDA-Image 2026-09-07|腾讯3B活跃参数逼近万亿代理-WeAgent 2026-09-06|长视频世界模型的记忆与状态概览 2026-09-04|港中深5秒训练撑起一小时世界推演-SolarWM 2026-09-03|Runway让界面直接生成而不再写代码-Solaris 2026-09-02|阿里7B原生音视频生成硬刚33B大模型-DreamX 2026-09-01|西湖大学让写代码的智能体当世界大脑CWM 2026 年 8 月 2026-08-31|复旦GameWAM同时生成画面和键鼠动作 2026-08-29|字节斯坦福让视频记住走远的人RingForcing 2026-08-28|京东JoyAI-Echo-1.5给长视频装上跨镜头记忆 2026-08-27|浙大清华19B音视频4步2.5秒出片TurboT2VA 2026-08-26|京东开源世界模型边走边生720p视听EchoWM 2026-08-25|Meta让15B DiT优化器时间砍半Muon-DiT 2026-08-24|快手可灵证明奖励劫持源于流形漂移ThermoDPO 2026-08-23|3DGS工程化转折:传得动跑得起管得住 2026-08-22|阿里6B小模型编辑分反超20B的Swift-Image 2026-08-21|中科大揪出蒸馏后数字人变木头人DynaForcing 2026-08-20|LumaAI首证扩散缩放定律要10倍数据(Abra) 2026-08-19|阿里像素空间T2I反超潜空间Z-Image-Pixel 2026-08-18|港科大数字人比LTX-2快33倍还能连说一分钟Omni… 2026-08-17|AlayaLab世界模型连跑两小时画面不塌Evoke 2026-08-16|NVIDIA领衔 On-Policy 蒸馏让学生反超教师 2026-08-14|阿里14B数字人实时跑三分钟LiveAnimate 2026-08-14|深度解读|NVIDIA Sol-Engine|Agent自动调栈视频推理快2.77倍 2026-08-13|Adobe视频世界模型首次外推物理定律LDR 2026-08-12|阿里DUET两步视频质量多样性双赢 2026-08-11|NVIDIA免训练治好世界模型失忆WorldTrace 2026-08-10|阿里Wan-Animate-2角色动画冲进实时24FPS 2026-08-07|Meta实测视觉懒惰生成数据只留5%反而更强 2026-08-07|深度解读|Sand.ai MAGI-2|114B视频MoE把token切12份挑专家 2026-08-06|京东22B音画同修让百年老片复活OmniVR 2026-08-03|深度解读|MiniMax H3|33B统一音视频与2K链路拆解 2026 年 7 月 2026-07-31|Adobe混合DiT算效提升7.3倍Chimera 2026-07-30|Adobe世界模型16FPS实时生成分钟视频 2026-07-29|Apple端侧Siri语音16倍实时仅21MB 2026-07-28|Meta单H200实时视频22.84帧MsForcing 2026-07-27|NVIDIA单卡视频生成提速120倍SANA2.0 2026-07-25|深度解读|清华SLA2|97%稀疏对比DeepSeek DSA 2026-07-23|阿里高德5090单卡实时世界模型ABot-World-0 2026-07-22|华为40万美元逼近闭源Boogu-Image 2026-07-21|腾讯混元4步视频胜50步强化扩散MeanFlowNFT 2026-07-16|字节Seed读回提示当奖励SpectraReward 2026-07-14|DeepMind断言视频生成即通用视觉底座 2026-07-10|高通MobileWan让5B视频扩散上手机16FPS 2026-07-03|北航ETH免训练让FLUX画图提速25倍MrFlow 2026-07-02|快手可灵MemLearner可学习记忆视频世界模型 2026 年 6 月 2026-06-29|清华LiveEdit实时流式视频编辑12.66FPS 2026-06-22|近两周精选-MaineCoon实时音视频世界模型 2026-06-02|微软实时流式数字人视频比肩大模型 2026-06-01|英伟达SANA单卡24FPS实时流式视频编辑 2026 年 5 月 2026-05-29|生数科技minWM开源实时交互视频世界模型 2026-05-28|北大OSP-Next视频生成跨硬件加速 2026-05-27|美团LongCat-Avatar 1.5开源逼近闭源数字人 2026-05-26|百度ERNIE-Image开源8B DiT追平闭源 2026-05-25|字节Bernini让MLLM规划DiT渲染视频 2026-05-21|智能编辑成统一模型通用任务Uni-Edit 2026-05-20|视频生成补物理常识NEWTON 2026-05-19|长视频生成FP4训推全栈LongLive-2.0 2026-05-18|14B视频对齐单步训练Flash-GRPO 2026-05-17|实时自回归视频生成加速 2026-05-15|实时视频2步出帧Causal Forcing++ 2026-05-14|视频扩散从少步到任意步AnyFlow 2026-05-13|INSET图像即词汇开启统一视觉生成新范式 2026-05-12|Forcing-KV 视频扩散2.82倍加速突破实时 2026-05-11|Cola DLM 扩散语言模型挑战自回归范式 2026-05-09|视频编辑最新进展 2026-05-08|JoyAI统一模型唤醒空间智能,通义D-OPSD破解少步… 2026-05-06|DiT-MoE统一多模态模型25B仅激活3B,运动感知缓… 2026-05-05|纯视觉流统一生成颠覆文本管线,1D-Token端到端FI… 2026-05-04|20260504|4步打败40步!AdvDMD蒸馏加速刷新SD3.5 2026-05-02|13.7-18.6x注意力加速,稀疏注意力撕开视频Di… 2026-05-01|20秒训练7x加速DiT,SAMG零开销解锁空间自适应… 2026 年 4 月 2026-04-30|V-GRPO让RL对齐提速3倍,64token暴力生图 2026-04-28|20FPS实时数字人Hallo-Live 2026-04-27|多目标Pareto后训练ParetoSlider 2026-04-26|统一多模态生成大爆发:Omni五模态SOTA碾压Qwen 2026-04-25|视频编辑评测方法全景:从传统指标到 Reward Model 的范式跃迁 2026-04-24|Wan-Image 2026-04-23|淘宝试穿上线-Google城市视频-GRPO优化扩散 2026-04-22|武大MemWN按需记忆撑起长视频一致性 2026-04-21|Qwen3.5-Omni全模态215项SOTA 2026-04-10|重新审视可控扩散训练目标——直接x₀监督实现2倍加速 2026-04-09|一图多改不再崩-MIRAGE并行编辑 2026-04-08|分数步蒸馏新范式1.x-Distill 2026-04-07|SC-DMD蒸馏2-4步高质量视频生成Salt 2026-04-06|VOID因果推理视频编辑|DynaVid CVPR2026|SteerFlow免训练编辑 2026-04-05|Mistral Voxtral TTS胜ElevenLabs 2026-04-04|视频生成前沿|统一框架|长视频|物理一致性 2026-04-03|Dynin-Omni|OmniVoice 2026-04-02|MacTok 64-token SOTA 2026 年 3 月 2026-03-30|BiFM|Wan-Weaver|PackForcing|Voxtral TTS 2026-03-29|GIDE|ScaleEdit-12M|Calibri 2026-03-28|视觉生成后训练与偏好优化 2026-03-27|ScrollScape|OmniWeaving 2026-03-25|上交ScaleEdit-12M+CVPR跨时间步自校准 2026-03-24|CubiD高维离散扩散|扩散通用加速|FoleyDirector V2A 2026-03-23|MOSS-TTS|ColourCrafter|Q-Drift 2026-03-22|视频生成与编辑前沿进展 2026-03-21|偏好对齐|RL后训练|SOLACE|CRAFT|CRD|VIGOR 最后更新:2026-10-05 12:39
2026年10月04日
28 阅读
0 评论
4 点赞
2026-10-04
AIGC 基本功知识树
每日论文速读解决的是广度——今天世界上发生了什么。但读论文有个前提:你得先看得懂。 这里是另一条线:把视觉生成的底层零件一个个拆开讲透。每篇都给数学推导 + 能跑的代码 + 经典论文出处,并标注前置知识,你可以顺着依赖链一路读下来。 当前规划 38 个知识点,已发布 33 篇。 图例:● 已发布 · ◍ 已发布待更新 · ◐ 正在写 · ○ 计划中 数学与优化基础(2/2) 概率、变分推断、随机微分方程——读懂扩散模型公式的最小前置集合。 ● 变分下界与重参数化(ELBO) · 难度:入门前置 不理解 ELBO 就没法理解 VAE 的 KL 项为什么要加权,也没法理解扩散模型的训练目标从哪来。 ● 扩散过程的前向与反向推导(SDE) · 难度:入门前置 · 前置:ELBO 把 DDPM 的离散公式和 SDE 的连续视角对上,后面所有采样器的差异都能一句话解释。 注意力与核心零件(5/7) 注意力、位置编码、稀疏激活与循环记忆——视频生成里最吃显存、最难外推也最影响长程建模的部件。 ● 自注意力机制的计算与显存账本(MHA) · 难度:核心必修 先把 O(N²) 的常数项算清楚,才能判断后面各种稀疏/线性方案到底省在哪一项。 ● 旋转位置编码 RoPE 的原理与实现(RoPE) · 难度:核心必修 · 前置:MHA 许多现代视频 DiT 使用 RoPE;理解旋转和相对位置性质,才能区分坐标可计算与生成质量可外推。 ● FlashAttention 为什么不需要存下注意力矩阵(FlashAttn) · 难度:进阶 · 前置:MHA online softmax 这一个技巧撑起了整个长序列时代,值得把递推式一步步推一遍。 ● 视频 DiT 里的 3D RoPE 与分辨率外推(3D-RoPE) · 难度:进阶 · 前置:RoPE 换分辨率或帧数时,3D RoPE 的坐标、频率与轴分配是需排查的因素之一,还要检查训练分布、VAE 与注意力实现。 ● 混合专家 MoE:稀疏激活怎么省算力(MoE) · 难度:工程实战 · 前置:MHA 统一多模态模型开始普遍用 MoE 扛参数量,但路由不均衡带来的训练不稳定很少被讲清楚。 ○ 线性、循环与记忆架构(LinearMem) · 难度:前沿 · 前置:MHA 长序列架构正在从保存完整注意力矩阵转向可更新状态和深度循环;把复杂度、记忆容量与并行性放进同一框架,才能判断它们何时真能替代 Softmax Attention。 ○ 视频生成里的稀疏注意力(SparseAttn) · 难度:前沿 · 前置:FlashAttn、3D-RoPE 视频长序列的全注意力成本很高;稀疏化可以减少计算,但免训练部署的质量取舍需和精确 IO 优化、缓存等方案比较。 生成范式(7/9) 从 DDPM 到流匹配,以及把 50 步压到 4 步的蒸馏路线。 ● DDPM 训练目标与采样流程(DDPM) · 难度:核心必修 · 前置:SDE DDPM 是理解扩散模型的重要起点;先读懂其 loss 与采样循环,再比较其他生成路径和训练目标。 ● 分类器无关引导 CFG 的代价与调法(CFG) · 难度:核心必修 · 前置:DDPM CFG 让每步算两遍,是推理成本里最容易被忽视的 2×,也是蒸馏首先要干掉的对象。 ● 从 DDIM 到高阶采样器(DDIM) · 难度:进阶 · 前置:DDPM 采样器换一个、缓存策略就得重调——这是推理加速最常见的踩坑点。 ● 流匹配与 Rectified Flow(FlowMatching) · 难度:进阶 · 前置:SDE、DDPM 流匹配被许多近期生成模型采用;它与 DDPM 的概率路径、监督目标及采样方式值得系统比较。 ● 潜空间扩散与 Stable Diffusion 架构(LDM) · 难度:进阶 · 前置:DDPM、VAE 把扩散搬进潜空间这一步,直接决定了今天视觉生成的算力可行性。 ● DiT:用 Transformer 替掉 UNet(DiT) · 难度:进阶 · 前置:LDM、MHA adaLN-Zero 这个小设计是 DiT 能稳定训起来的关键,值得逐行对照公式看。 ● 自回归视频生成与 Forcing 范式(Forcing) · 难度:工程实战 · 前置:DDPM、DiT 双向扩散没法流式出帧,Forcing 系列是把视频生成变成可交互的关键一步,也是当下最活跃的范式。 ○ 少步蒸馏:从 50 步到 4 步(Distill) · 难度:前沿 · 前置:DDIM、CFG 蒸馏是过去两年推理加速收益最大的一条线,也是最容易把质量搞崩的一条。 ○ 世界模型:从视频生成到可交互环境(WorldModel) · 难度:前沿 · 前置:Forcing 视频生成正在从「出片」转向「可交互环境」,这是范式转变而不是又一个 SOTA 分数。 表征与压缩(4/4) VAE / Tokenizer——视频生成的「地基」,决定了上限和伪影形态。 ● VAE 结构与训练目标(VAE) · 难度:核心必修 · 前置:ELBO 潜空间扩散依赖编码器输出的统计尺度;像素空间扩散不使用 VAE,不能将 scaling_factor 推广到所有扩散模型。 ● 离散化表征:VQ-VAE 与 VQGAN(VQGAN) · 难度:进阶 · 前置:VAE 离散自回归视觉模型依赖 tokenizer,码本利用率是重要问题;连续潜变量自回归路线不适用这一前提。 ● 视频 VAE 的时空压缩结构(VideoVAE) · 难度:工程实战 · 前置:VAE 时间维压缩比和因果卷积的实现方式,直接决定了长视频能不能逐块解码而不接缝。 ● 视频 VAE 的常见 loss 组合(VAELoss) · 难度:工程实战 · 前置:VideoVAE L1 + KL + LPIPS + GAN 四项的权重配比是玄学重灾区,把每项在管什么讲透很有价值。 对齐与强化学习(4/4) 从 PPO 到 GRPO,以及怎么把 RL 用到扩散和视频生成上。 ● 策略梯度与 PPO 基础(PPO) · 难度:核心必修 PPO 是在线策略优化的重要基础;掌握优势估计和裁剪后,再区分 GRPO 与直接偏好优化等不同路线。 ● 从 DPO 到 GRPO:去掉价值网络(GRPO) · 难度:进阶 · 前置:PPO GRPO 用组内相对奖励省去 critic;净成本还取决于组采样、参考模型、奖励评估和显存,不能统一声称降低一个量级。 ● 把 RL 用到扩散模型上(DiffusionRL) · 难度:前沿 · 前置:GRPO、DDIM 把多步去噪当成一条 MDP 轨迹,是理解视频 RL 各种做法的统一视角。 ● 视频生成中的强化学习与奖励模型(VideoRL) · 难度:前沿 · 前置:DiffusionRL、FlowMatching 视频的奖励要同时管画质、运动和指令遵循,reward hacking 在这里表现得最明显。 分布式训练(4/4) 参数、数据、序列三个维度怎么切,以及切完之后通信量变成多少。 ● 数据并行与 ZeRO 显存切分(ZeRO) · 难度:进阶 先把「参数 + 梯度 + 优化器状态」的显存账算清楚,才知道该切哪一部分。 ● 混合精度与数值稳定性(AMP) · 难度:进阶 BF16、FP16 与 FP8 的范围、精度和缩放机制不同,排查异常时要区分前向溢出、梯度下溢与舍入误差。 ● 张量并行与流水线并行(TP-PP) · 难度:工程实战 · 前置:ZeRO TP 的通信量和 PP 的气泡率都能手算,算完就知道并行度该怎么配。 ● 序列并行与 Ring Attention(SP) · 难度:前沿 · 前置:TP-PP、FlashAttn 超长视频序列可能需要序列或上下文并行;是否值得使用取决于单卡容量、通信和计算重叠。 推理加速与部署(5/5) 先用性能模型定位算力、带宽和显存瓶颈,再用缓存、稀疏、量化与算子融合把生成时间从分钟压到秒。 ● 性能建模与 Profiling:算力、带宽与显存账本(Roofline) · 难度:进阶 · 前置:AMP 不先判断算力、带宽还是显存容量在卡住系统,量化、融合、缓存很容易做成负优化;这篇提供所有推理优化节点共同的测量方法。 ● KV Cache 与自回归视频生成(KVCache) · 难度:进阶 · 前置:MHA、Roofline 自回归视频模型正在把 LLM 的这套缓存工程整体搬过来,值得先打好底。 ● 算子融合与 CUDA Graph(Fusion) · 难度:工程实战 · 前置:MHA、Roofline 小算子多的模型往往卡在访存和 launch 开销上,融合的收益比换算法更确定。 ● 量化:从 INT8 到 FP4(Quant) · 难度:工程实战 · 前置:AMP、Roofline 激活里的离群值是量化掉点的主因,SmoothQuant 的迁移技巧值得逐步推演一遍。 ● 扩散模型的跨步缓存复用(DiffCache) · 难度:前沿 · 前置:DDIM、DiT、Roofline 相邻去噪步的特征高度相似,这是免训练加速里性价比最高的一类手段。 评测与指标(2/3) FID/CLIP/VBench 各自测的是什么,以及它们什么时候会骗人。 ● FID / CLIP Score 到底测了什么(FID) · 难度:核心必修 FID 对样本量和预处理极其敏感,不同论文的数字经常根本不可比。 ○ 视频生成评测:VBench 与人工验收(VBench) · 难度:进阶 · 前置:FID 加速类工作最爱报 VBench 总分不变,但分维度看往往能发现明显退化。 ● 音频质量评测:MOS、PESQ 与 FAD 各测什么(AudioEval) · 难度:工程实战 音视频生成里音频质量几乎全靠几个数字说话,但这些数字测的根本不是一回事:SI-SDR 只管波形对齐、PESQ 只为通信语音设计、FAD 看的是分布而不是单条样本的保真度。选错指标会得到自欺欺人的结论——比如用波形指标验收 codec,数字很好听但高频毛刺全在。 接下来会写 线性、循环与记忆架构(LinearMem) —— 长序列架构正在从保存完整注意力矩阵转向可更新状态和深度循环;把复杂度、记忆容量与并行性放进同一框架,才能判断它们何时真能替代 Softmax Attention。 视频生成里的稀疏注意力(SparseAttn) —— 视频长序列的全注意力成本很高;稀疏化可以减少计算,但免训练部署的质量取舍需和精确 IO 优化、缓存等方案比较。 少步蒸馏:从 50 步到 4 步(Distill) —— 蒸馏是过去两年推理加速收益最大的一条线,也是最容易把质量搞崩的一条。 最后更新:2026-10-05 12:40
2026年10月04日
11 阅读
0 评论
3 点赞
2026-10-04
AIGC 每日速读|2026-10-04|商汤小模型绕圈反超6.5倍大模型,Looped-DiT
今日 AIGC 论文速览 今日共 10 篇 · 高效文生图架构 3 篇 · 统一多模态 1 篇 · 长视频与音视频生成 2 篇 · 生成推理加速 2 篇 · 视频强化学习与图像评测 2 篇 重点论文标题列表 Looped-DiT(商汤):260M循环4圈赢6.5倍大模型 Multimodal Flow(华科):图文全连续流建模,仅用150B词元 NEPA-DiT(密歇根大学):预测嵌入当条件,FID 1.32 MosaiChunk(康奈尔大学):只训13.6M路由,长视频回访不失忆 Soundwich(西蒙弗雷泽大学):冻结LTX-2.5免训练拆分音轨 今日论文速览 1. Looped-DiT:260M循环4圈赢6.5倍大模型 Looped Diffusion Transformer | 商汤科技;清华大学;南洋理工大学 | arXiv:2609.40305 关键词:文生图, 循环计算, 参数共享, 深度监督, 潜在推理 前序问题:文生图模型提质通常只有两条路:加参数或加去噪步数。语言模型里已验证的「循环复用同一组层」能在不增参数的前提下加深计算,但直接搬进 MMDiT 并不稳定:推理圈数超过训练值后性能先涨后跌,中间表示里可线性解码的空间位置信息逐圈流失。 本文贡献:Looped-DiT 把网络切成前段 6 块、循环段 5 块、后段 6 块,每个去噪步内让循环段共享权重重复 N 次(训练取 4 圈)。两项稳定手段:深度监督,每圈输出都经共享后段解码并计入 flow loss;自调制注意力,用门控或无参的 Exclusive Self-Attention 抑制注意力对局部信息的过度改写。 实验效果:260M 的 B/16 在 GenEval、DPG、PRISM、CoRe、Spatial、TIIF 六项均值 71.5,六项全部高于 1.7B 的 InternVL-U(均值 69.0),也高于 6.6B 的 Uni-CoT(66.8),单图推理算力约为 InternVL-U 的 1/4.9。参数对齐对照里,B/32 均值从 55.2 提到 59.1;同等推理预算下,加圈比加去噪步更划算。 批判点评:「参数不变」不等于「算力不变」:循环版每步推理 267 GFLOPs,是同参数基线的 1.83 倍,训练 1246 GFLOPs 更是 2.83 倍。按推理算力对齐时,加宽版 MiniT2I 已到 58.6,与 Looped-DiT 的 59.1 只差 0.5。分项看,循环擅长约束求解,需要常识推断的 CoRe/Generalization 只 +7.5,远不如文本 CoT 的 +22.2。 2. Multimodal Flow:图文全连续流建模,仅用150B词元 Multimodal Flow: Unified Flow Modeling of Language and Vision in Embedding Spaces | 华中科技大学;北京交通大学;地平线 | arXiv:2609.40362 关键词:统一多模态, 连续流匹配, 文生图, 视觉理解, 多模态预训练 前序问题:统一理解生成模型多走两条路:文本和图像都离散化,图像要过量化瓶颈;或文本离散、图像连续,两种目标与采样流程各写一套。全连续建模能让两种模态共用一个生成过程,但在多模态预训练规模上很少被系统验证。 本文贡献:MF-1 用冻结的 T5-small 与 SigLIP2 把文本块(每 8 个词元一块)和图像编码成连续表示,按顺序组织成 hyperchunk,交给一个 chunk-causal 的 flow 主干学习单一速度场;注意力投影共享、FFN 按模态分开。训练时并行预测多个目标块,推理时逐块生成,再由预训练图像解码器与单独训练的文本解码器还原。 实验效果:1.6B 版本只用 150B 预训练词元,GenEval 0.82、DPG 83.44,理解侧 POPE 86.1、MMBench 67.2。数据、优化与参数完全对齐时,全连续版 GenEval 0.713、MMBench 46.7,高于 Transfusion 式混合(0.669 / 33.7)与 Chameleon 式全离散(0.374 / 38.4)。 批判点评:生成强、理解弱:MMBench 67.2 落后同量级的 Janus-Pro(75.5)与 JanusFlow(74.9),VQAv2 72.6 也低于 JanusFlow 的 79.8(对方有预训练 LLM 初始化);DPG 83.44 还没追上纯生成的 SD3 Medium(84.08)。scaling 曲线上 1.6B 预训练阶段 GenEval 只到约 0.32,表中 0.82 是再经 5B 词元微调后的结果;「全连续」也依赖冻结编码器与外部解码器,图像输入仅 224 分辨率。 3. NEPA-DiT:预测嵌入当条件,FID 1.32 Embedding Prediction Helps Image Generation | 密歇根大学;卡内基梅隆大学 | arXiv:2610.02203 关键词:类别条件生成, 嵌入预测, 扩散 Transformer, 表征对齐, ImageNet 前序问题:DiT 里类别或文本只嵌入一次,每个去噪步复用同一个条件,「这个条件对眼前这张噪声图意味着什么」全留给生成器自己推断。NEPA(下一嵌入预测自回归)已能在连续嵌入上做自回归预测,作者想知道:能否用预测出的干净图嵌入替代固定条件。 本文贡献:把生成写成「条件 → 噪声图 → 干净图」的嵌入序列,干净图的 patch 嵌入正是条件与噪声图之后的「下一组嵌入」。作者训练 NEPA 模型用多嵌入预测(MEP)一次性预测全部干净嵌入;生成时 DiT 以这些预测为条件,并在每个去噪步按当前噪声状态重新计算,条件随采样进程自适应。 实验效果:ImageNet 256×256 上,NEPA-DiT-XL 结合 REPA,SDE-250 采样 FID 1.32、IS 311.0,ODE-96 采样 FID 1.57;训练为 NEPA 240 + 生成器 80 epoch,约为 REPA(800 epoch)训练算力的三分之一。九组规模交叉实验中,生成器或 NEPA 模型变大,FID 都单调下降。 批判点评:训练省了,推理贵了:多挂一个 711M 的 NEPA 网络,总参数约 1.39B,是 SiT-XL/2+REPA(675M)的两倍;同为 SDE-250,单图 160 TFLOPs,比对方的 91 TFLOPs 多 76%。换成更省的 ODE-96,FID 1.57 反而不如 REPA 原版的 1.42。若放开 tokenizer,表中 SFD、RAEv2 已做到 1.06。 4. MosaiChunk:只训13.6M路由,长视频回访不失忆 MosaiChunk: Compositing Spatio-Temporal Memory for Autoregressive Video Generation | 康奈尔大学;加州大学伯克利分校;哈佛大学;Impossible Inc. | arXiv:2610.02153 关键词:自回归视频生成, 长时记忆, KV 缓存, 记忆路由, 场景回访 前序问题:自回归长视频受限于上下文窗口:物体或场景移出窗口后细节丢失,镜头转回来时往往「重新编」一个。整块检索历史 chunk 能补记忆,但固定显存预算下塞不了几块;滑动窗口则干脆遗忘。 本文贡献:作者先验证冻结的视频生成器能直接消费非连续拼接的历史 KV 并还原对应内容。MosaiChunk 据此把每个历史 chunk 聚成若干 section 并编码描述子,由轻量路由器全局挑出 top-N,在固定活跃显存预算内拼成一块跨时空的 KV 马赛克;生成器全程冻结,只自蒸馏训练路由器。同时发布 RememBench 回访基准,含 T2V 与 I2V 两个子集。 实验效果:在 MiniMax-H3 改造的自回归版 H3-AR(T2V)上,两块远程记忆预算下回访 CLIP 从滑窗基线 0.751 升到 0.936,LPIPS 从 0.652 降到 0.500,优于整块检索 MoC(0.841);LingBot-World-Infinity(I2V)180° 旋转回访 CLIP 0.793→0.863。路由器只有 13.65M 参数。 批判点评:自建基准涨幅大,公开基准涨幅小:WBench 的 Subject 一致性仅 88.10→88.38,两块预算下 Segment 一致性反从 97.47 掉到 94.94;T2V 的 Drift 也略升(0.050→0.055)。论文还写明评测用的是早期 checkpoint:T2V 计划训 4000 步、实取第 416 步,I2V 计划 4768 步、实取第 500 步,未解释为何不用训完的权重。 5. Soundwich:冻结LTX-2.5免训练拆分音轨 Soundwich: Video Generation with Layered and Controllable Audio | 西蒙弗雷泽大学;塞浦路斯大学;CYENS 卓越中心;特拉维夫大学 | arXiv:2610.00691 关键词:音视频联合生成, 音轨分离, 免训练, 时间控制, 可编辑音频 前序问题:联合音视频模型已能生成带同步声音的视频,但音频是一条混好的总轨:谁在何时说话、配乐和环境声都没法单独改。实际制作中,对白、音乐、音效、环境声都是分轨编辑的。 本文贡献:Soundwich 不做训练,直接改造冻结的 LTX-2.5:先把音频轨迹扩成 N 条可编辑分轨加一条内部场景轨,用缓存的「出声 / 静音」特征(timing carrier)在指定时间窗内回放或压制;再用 gather–broadcast 注意力让各分轨共享全局声学语境,并用实体掩码把每条分轨的跨模态注意力路由到画面中对应的发声者。 实验效果:15 个多声源场景上,时间窗成功率 93.1%、窗外误发声 1.7%、声源归属 90.0%,窗内 WER 0.060;同场景原版 LTX-2.5 窗成功率仅 20.4%,MiniMax H3 为 27.0%,Gemini Omni 为 85.0%。去掉 timing carrier,窗成功率从 93.3% 跌到 15.6%。 批判点评:评测规模很小:主表只有 15 个场景、每场景 3 个固定种子,场景广播模块的人评只收回 15 份问卷。真正接近的对手是闭源 Gemini Omni,WER 0.072 对 0.060、窗成功率 85.0% 对 93.1%,且它只有 14 个场景可计 WER。声源归属 90.0% 比朴素批量扩展的 82.2% 只高 7.8 个点。 6. DMAD:判别器替掉辅助分数网络,4步出片 DMAD: Distribution Matching as Adversarial Distillation for Fast Visual Generation | 得州农工大学;字节跳动智能创作 | arXiv:2610.02188 关键词:少步蒸馏, 分布匹配, 对抗蒸馏, 视频生成, 音视频生成 前序问题:DMD 系列少步蒸馏要从「目标分数 − 学生分数」的差里取梯度,因此必须额外维护一个持续拟合学生分布的辅助扩散模型,显存与算力开销都大,在 14B、33B 级视频模型上尤其吃紧。 本文贡献:DMAD 把分布匹配改写成分类:共享主干上挂两个判别头,分别区分「真实数据 vs 学生」与「教师样本 vs 学生」,用 logit 直接学习对数密度比,学生只用作用在 logit 上的线性损失训练,无需拟合辅助分数。作者证明判别器最优时该损失恰好还原 DMD 的梯度,并用真实头在真实与教师样本间的 logit 差,自适应调节各噪声级上教师监督的权重。 实验效果:ImageNet-64 一步 FID 1.04(加投影判别器),SDXL 四步 COCO-10K FID 14.47,Wan2.1-14B 四步 VBench 85.15,均优于所比少步方法与多步教师。在 MiniMax-H3-33B 联合音视频生成上,四步学生人评总体偏好 79.1% 胜 DMD2、84.6% 胜 rCM;每次生成器更新耗时从 818.1s 降到 233.3s。 批判点评:人评高分主要来自画质与音质(对 DMD2 胜率 84.2% / 87.9%),提示词对齐只有 54.3% 对 45.7%,接近打平。AVGen-Bench 上音画同步 AV 分 49.1,低于 DMD2 的 56.7;唇形同步 30.6,离教师的 39.0 也远。SDXL 一步时 Patch FID 32.42,反而比 DMD2 的 26.98 差,细节保真不如整体 FID 好看。 7. FlashForward:复用在途KV,4卡流水线生成长视频 In-Flight KV Cache with Clean Anchors for Faster Autoregressive Video Diffusion | Meta;南洋理工大学 | arXiv:2609.32540 关键词:自回归视频扩散, KV 缓存, 流水线并行, 长视频生成, 推理加速 前序问题:少步自回归视频扩散按 chunk 逐段生成。为记住已生成内容,Self-Forcing、HiAR 等方法要额外跑「只更新缓存、不推进输出」的前向来重建干净或低噪 KV。可每次去噪前向本来就算出了当前 chunk 的 KV,这部分被白白浪费。 本文贡献:FlashForward 直接复用每个去噪阶段算出的「在途」KV:当前 chunk 完成一个阶段后,该阶段的缓存即可供下一个 chunk 使用,于是每张 GPU 负责一个去噪阶段,不同 chunk 在卡间流水并行。为弥补噪声历史带来的外观与运动漂移,再由 planner 预先生成稀疏的干净锚点 latent,从前后两侧约束 renderer 的轨迹。 实验效果:4 卡下,16 FPS、20 秒以上视频比 HiAR 快 1.16–1.69 倍、比 Self-Forcing 快 1.42–2.92 倍(1.3B 与 14B,480p / 720p)。1.3B 480p 生成 65 秒视频去噪耗时 21.3s(HiAR 24.7s、Self-Forcing 65.1s);VBench 总分 0.838,65 秒 VBench-Long 0.8395,高于 Self-Forcing 的 0.7806。 批判点评:提速依赖多卡流水:单卡 480p 下它比 Self-Forcing 还慢(20 秒视频 17.77s 对 17.44s,65 秒 59.64s 对 57.89s)。生成 5 秒短片时 4 卡耗时 2.21s,慢于 HiAR 的 2.06s,更慢于双向模型的 1.13s。14B 720p 峰值显存 112.8 GiB/卡,比 Self-Forcing 多约 20 GiB;20 秒视频语义分 0.748 也略低于 HiAR 的 0.757。 8. TVRL:用VLM梯度给视频token分功劳 Token-Level Video Reinforcement Learning | 美国东北大学 | arXiv:2610.01973 关键词:视频生成, 强化学习, GRPO, 信用分配, VLM 奖励 前序问题:视频生成的 RL 后训练通常给整段视频一个标量奖励。可视频的缺陷是局部的:有的 token 已满足提示词,有的才需要修。标量奖励定位不了错误,优化时既会扰动本来正确的区域,又对真正出错的区域用力不足。 本文贡献:TVRL 让奖励自己给出 token 级功劳:冻结 VLM 对提示词拆出的若干是非题给出答案似然,平均后作为视频级奖励;同一 VLM 对视频输入的梯度幅值,则标出哪些生成 token 最影响该得分。在 GRPO 中,组相对优势决定更新方向与强度,经 3×3 平滑、按问题条件化的 credit 图在裁剪策略比内重加权各去噪步的转移对数概率。 实验效果:以 HunyuanVideo-1.5 为底座、Qwen3.5-9B 作奖励,VBench-2.0 Overall 从 54.09 升到 57.69(+3.60),同设置 GRPO 仅 54.54。换 SAGE、Flow、Dance 三种 SDE 采样器,比匹配 GRPO 高 2.68–3.15;换 VideoAlign 等四种奖励模型,高 1.33–3.15。 批判点评:增益并不均匀:常识维度上 TVRL 常不如 GRPO(Qwen3.5-9B 奖励下 64.31 对 64.88,VideoScore2 下 64.60 对 64.89),VideoScore2 下 Human 维度也从 91.52 降到 90.79。评论家规模很关键:换成 Qwen3.5-0.8B 时 Overall 仅 51.85,比未训练的底座 54.09 还低。所有结果都只训 100 步(64 张 A100),单步耗时比 GRPO 多 18%。 9. PixelDense:4个冻结老师分两路对齐像素扩散 PixelDense: Dense Prediction as Representation Alignment for Pixel Diffusion | 弗吉尼亚大学;密歇根州立大学;Adobe;Arcade AI | arXiv:2610.00483 关键词:像素扩散, 表征对齐, 稠密预测, 文生图, 几何先验 前序问题:REPA 通过对齐预训练编码器特征来加速 DiT 训练,但对齐目标几乎都是 DINOv2、CLIP 这类语义编码器。已有分析指出起作用的是空间结构而非全局语义,那么专门预测结构的稠密预测模型(分割、深度)为何没人拿来当对齐老师? 本文贡献:在像素空间扩散上,作者先验证 SAM2、Depth Anything v2、Metric3D v2 单独加入都优于只用 DINOv2,但四个老师直接相加反而不如最好的单个几何老师,语义与几何梯度在同一投影上互相争抢。PixelDense 因此让 DINOv2+SAM2 走语义投影流、两个深度模型走几何投影流,再加权重空间正交惩罚让两流处在不相交子空间;老师全部冻结、推理时丢弃。 实验效果:在 PixelGen-XXL(512×512)上,GenEval Overall 0.7927→0.8093,DPG 78.7→78.9,HPS 0.280→0.282;同一配方不调参迁移到 DeCo,GenEval 0.8620→0.8690。部分加噪重建中,τ=0.5 时 COCO 全景 PQ 23.23→31.43;PIE-Bench 上 SDEdit 背景 PSNR 最多高 2.2 dB。论文已接收 NeurIPS 2026。 批判点评:几何探针涨得多,生成指标涨得少:GenEval 只 +0.017、DPG +0.2、HPS +0.002,DeCo 上 +0.007,计数项 58.75 仍低于 PixelGen 原版的 59。消融里单独的几何流(DA2+M3D)只有 0.7960,还不如 DINOv2+DA2 单老师的 0.8069。主实验只是在已发布权重上用 2 张 H200 微调 1 万步,从头训练仅给出「1.23 倍更快达到基线峰值」一项。 10. VIEScore2:16×16网格同时打分并圈出缺陷 VIEScore2: Unified Image Evaluation with Spatially Grounded Explanations | 滑铁卢大学;NVIDIA;台湾大学(中国台湾);南洋理工大学 | arXiv:2610.00994 关键词:图像评测, 缺陷定位, 生成与编辑, GRPO, VLM 评估器 前序问题:现有合成图评估器大多只给一个标量分数,不说明依据在图上哪里;能定位缺陷的模型又通常不打分、也不支持编辑任务的条件图输入。生成与编辑的评测和奖励信号,缺一个「既能打分也能指位置」的统一接口。 本文贡献:VIEScore2 把图像表示为 N×N(默认 16×16)文本网格,单次前向同时输出感知质量、语义一致性分数与缺陷格子位置,支持可选条件图,覆盖文生图、编辑、参考图生成等任务。它基于 Qwen3-VL-8B 在 38K 条混合监督上先 SFT,再以 Dice 重叠、分数准确度与格式奖励做 GRPO,最后用无参数解析器把网格结果转成可读解释。 实验效果:主测试集整体分数 SRCC 0.601,高于同输入下最强的零样本通用 VLM Gemini-3-Flash(0.491)与 GPT-5.6-sol(0.437);联合评测网格 IoU 0.324,Qwen3-VL-8B 原版仅 0.070。缺陷定位在 6 个基准中 3 个第一、5 个进前三。 批判点评:整体领先主要来自训练同源数据:RichHF(0.692)与 EvalMuse(0.803)拉高均值,编辑类子集反被零样本 API 压过,TIE 0.429 对 Gemini-3-Flash 的 0.686,MRIE 0.417 对 GPT-5.6-sol 的 0.684。定位上,轻量的 SegFormer-b0 在主测试集 F1 拿到 0.493,与它的 0.506 相差无几;GRPO 之后 PQ 的 SRCC 还从 0.580 降到 0.558。 趋势观察 「加算力不加参数」成了今天的共同母题 Looped-DiT 让同一组块在每个去噪步里绕 4 圈,NEPA-DiT 每步重算一次预测嵌入当条件,两者都用额外推理算力换质量而不是加参数;但细看算力表,前者每步 1.83 倍 GFLOPs,后者单图多 76% TFLOPs,「小模型赢大模型」的账要连推理成本一起算。另一头,DMAD 与 FlashForward 继续在蒸馏和 KV 复用上压缩视频生成的步数与等待时间,TVRL 与 VIEScore2 则把奖励和评测从一个标量推进到 token 与网格级定位。 人工智能炼丹君 整理 | 2026-10-04 更多 AIGC 论文解读,关注微信公众号「人工智能炼丹君」 每日更新 · 论文精选 · 深度解读 · 技术脉络 微信搜索 人工智能炼丹君 或扫描下方二维码关注
2026年10月04日
2 阅读
0 评论
0 点赞
2026-10-04
AIGC 基本功|量化:从 INT8 到 FP4-Quant
量化:从 INT8 到 FP4 所属方向:推理加速与部署 | 难度:工程实战 | 前置知识:混合精度与数值稳定性、性能建模与 Profiling 关键词:INT8、W8A8、INT4、W4A16、SmoothQuant、AWQ、NVFP4、校准、离群值、分组缩放 01. 为什么需要它 设想一个视频生成服务:DiT 的线性层权重很大,单次生成又要反复经过这些层。团队把 BF16 权重直接换成 4 bit 文件,磁盘占用马上下降;上线后却发现小 batch 时并没有按 4 倍提速,字幕边缘还出现抖动。问题不在于「4 bit 失效」,而是把存储位宽、矩阵乘法位宽、缩放元数据、反量化成本和生成质量混成了同一个数字。若运行时先把 4 bit 权重解回 BF16 再调用普通 GEMM,压缩带来的只是存储或部分读带宽收益;若没有匹配硬件和内核,也不会凭空得到 FP4 算力。 再看一个可复现的小实验:附录的 quant_demo.py 生成 64 维输入,其中第 0 个通道有显著离群值。普通 W8A8 在留出的 64 条输入上的输出均方误差是 0.222085;做一次数学上严格等价的 SmoothQuant 式通道变换后,误差变成 0.001384。这两个数来自本机实际运行,不是论文模型的评测分数。它说明一个关键事实:量化前的浮点函数完全相同,量化后的误差却能相差很大。要读懂 INT8、AWQ 或 FP4 的论文,先要分清误差落在什么张量、哪个通道,以及内核真正执行了什么。 本文用一条线串起这些问题:先推整数映射与输出误差,再拆 SmoothQuant 如何搬走激活离群值、AWQ 如何保护重要权重,最后解释 NVFP4 为什么需要按块缩放。所有示例都在 CPU 上跑,不依赖模型权重;它们验证机制,不声称复现生产吞吐或画质。 02. 最小可用理解 第一句:量化把连续值映到少数离散码字;位数越少,步长或裁剪误差通常越大,实际误差还取决于缩放粒度和数据分布。第二句:SmoothQuant 面向 W8A8,把难量化的激活通道缩小、相应权重通道放大,使原始浮点线性层保持等价;AWQ 面向权重低比特,用激活统计找重要通道,并搜索权重缩放以减小输出误差。第三句:INT4 与 FP4 只是数字格式,真正的速度由硬件支持、打包布局、反量化位置、batch 和带宽瓶颈共同决定。 这里统一写线性层为 $Y=XW$:$X$ 是形状 $[T,K]$ 的输入激活,$W$ 是 $[K,N]$ 的权重,$Y$ 是 $[T,N]$ 的输出;$T$ 可理解为当前批的 token 数,$K$ 是输入通道,$N$ 是输出通道。真实视频 DiT 的 $T$ 可能包含时间与空间 patch,且随分辨率、帧数、去噪步变化。下文的校准统计只能代表采样过的条件分布,不能自动外推到所有视频场景。 因此,阅读任何“压到 4 bit 后提速”的结论时,先追问四件事:原始精度是什么、量化了权重还是激活、scale 按什么粒度共享、测试时是否调用了对应硬件的低比特内核。若论文只给模型文件大小与单项质量分数,就还不足以判断生产服务的成本收益。 03. 数学推导 3.1 从实数到整数:误差在哪里出现 给一组实数选择步长 $\Delta>0$ 与零点 $z$。一般的仿射量化写成一行: $$q=\operatorname{clip}\bigl(\operatorname{round}(x/\Delta)+z,\ q_{\min},q_{\max}\bigr),\qquad \hat x=\Delta(q-z)$$ $q$ 是存储的整数码字,$\hat x$ 是反量化后参与近似计算的值;$q_{\min},q_{\max}$ 是格式允许的最小、最大整数。对本文 W8A8 实验的对称 INT8,取 $z=0$、范围 $[-127,127]$、$\Delta=\max|x|/127$。没有裁剪且落在舍入区间内时,单个值的绝对误差不超过 $\Delta/2$;超出校准范围被截断后,这个上界不再成立。全张量共用一个步长时,一个大离群值会把 $\Delta$ 拉大,使大多数小值在相邻码字间跳得更粗。按通道或按组各给一个步长能局部化这种损失,但要保存更多缩放因子,并且内核未必支持同样快的计算路径。 线性层两边都有量化误差时,令 $\hat X=X+E_X$、$\hat W=W+E_W$。直接展开,而不是笼统说「精度下降」: $$\hat Y-Y=(X+E_X)(W+E_W)-XW=E_XW+XE_W+E_XE_W$$ 第一项是激活误差被权重放大,第二项是权重误差被输入激活放大,第三项是二者相乘。W8A8 三项都有;W4A16 的激活通常仍用高精度,主要关心第二项。即使某个权重元素误差很小,只要它对应的输入通道经常很大,也可能显著改变输出。这解释了为什么不能只用权重自身的 MSE 判断生成模型是否安全。 为了把“一个离群值拖累整组”算到具体数字,假设 127 级对称量化器要同时覆盖一个大小 20 的值与许多大小约 0.1 的值。若整个张量共用尺度,步长约为 $20/127=0.1575$;一个 0.1 会被舍入到 1 个码字,反量化约 0.1575,误差约 0.0575。假如该张量的最大幅度只有 1,步长便约为 $1/127=0.00787$,0.1 反量化约 0.1024,误差约 0.0024。两种情形都只用 8 bit,差别来自谁和谁共享 scale。实际矩阵乘还要把这些单点误差乘上权重并求和,因此不能只拿这个标量例子推最终 MSE,但它解释了为什么先观察通道直方图比直接改位宽更有用。 校准也有两层口径。静态量化先用代表性输入估计范围,推理时直接复用 scale;如果后来出现更大的输入,就有裁剪风险。动态量化可按当前 token 或当前批重新估计激活范围,减轻分布漂移,却要在运行时付出求最大值、计算 scale 和可能的同步成本。权重通常固定,离线按输出通道或按组量化较容易;激活每次都变,统计粒度过细会增加内核复杂度。本文故意采用校准集估计的静态全张量激活 scale,让离群问题足够清晰;它不是宣称这种设置在所有线上服务中最佳。 3.2 SmoothQuant:把离群值迁移到更好量化的一侧 取一个所有元素都为正的通道缩放向量 $s\in\mathbb R^K$,令 $D=\operatorname{diag}(s)$。在量化前做: $$X^{\prime}=XD^{-1},\qquad W^{\prime}=DW,\qquad X^{\prime}W^{\prime}=XD^{-1}DW=XW$$ 等号最后一步只用到 $D^{-1}D=I$,所以浮点层完全等价。第 $j$ 个激活通道除以 $s_j$,权重的第 $j$ 个输入通道乘以 $s_j$。如果异常大的激活集中在少数通道,选较大的 $s_j$ 就能把它们压下去;代价是相应权重幅度变大。SmoothQuant 论文的核心判断是:在其研究的 LLM 线性层中,权重一侧通常比激活一侧容易承受这种量化难度。它并非数学定理,具体模型仍要测。 记校准输入第 $j$ 通道的绝对最大值为 $a_j$,权重第 $j$ 个输入通道的绝对最大值为 $b_j$,SmoothQuant 的一种尺度选择是: $$s_j=\frac{a_j^{\alpha}}{b_j^{1-\alpha}},\qquad 0\le\alpha\le1$$ $a_j$ 和 $b_j$ 都要设正下界以免全零通道除零。$\alpha$ 控制把多少难度从激活侧搬往权重侧;本文实验固定 $\alpha=0.5$,不是所有模型的默认最优值。生产实现还要把 $1/s_j$ 融入前一层归一化的参数,避免推理时额外插一个逐元素除法。等价变换只保证未量化的 $XW$ 不变,不保证量化输出相同,这正是需要校准和误差测量的原因。 图 1:纵轴为绝对最大值的对数刻度,横轴为输入通道。橙色是原始值,蓝色是等价变换后。第 0 通道的激活离群值向权重侧迁移;这是一组程序构造的数据,不是任何论文模型的实测激活分布。 3.3 AWQ:为何要看激活,而不只看权重 若只量化权重,输出误差近似为 $XE_W$。对留出的输入矩阵平方求和,可写成: $$\|XE_W\|_F^2=\operatorname{tr}\bigl(E_W^{\mathsf T}X^{\mathsf T}XE_W\bigr)$$ $\|\cdot\|_F$ 是把所有元素平方后求和再开根号,$X^{\mathsf T}X$ 编码各输入通道的能量与相关性。若暂时忽略通道间相关性,式子近似为各通道「输入能量 × 对应权重误差能量」之和: $$\|XE_W\|_F^2\approx\sum_{j=1}^{K}\|X_{:j}\|_2^2\|E_{W,j:}\|_2^2$$ 所以「权重数值小」不等于「量化它不重要」:如果 $X_{:j}$ 经常大,该通道的权重误差就会被放大。AWQ 论文据此用激活统计找显著权重通道,采用等价缩放和校准误差搜索,而不是把少数通道改成混合精度来增加内核复杂度。其论文报告保护约 1% 显著权重即可明显降低量化误差;这属于论文实验结果,不能直接当成任何 DiT 的固定比例。 我们的教学版对每个输出通道的权重按 32 个输入元素成组做非对称 INT4 伪量化:组内用最小值和最大值求 $\Delta=(\max-\min)/15$,零点把实数零映到 $[0,15]$ 的整数区间。再用激活平均幅度构造 $s_j$,扫 20 个候选指数,在校准输入上找最小输出 MSE。这样抓住了「激活统计 + 缩放搜索 + 分组权重量化」的骨架;它没有 AWQ 的完整模型层搜索、真实 W4A16 内核或端到端精度验证,因此代码称为 awq_toy。 3.4 FP4 不等于把 INT4 改个名字 INT4 通常按整数码字加组缩放与零点解释;FP4 的码字本身有符号、指数和尾数。NVIDIA 的 NVFP4 文档把单个数据码字定义为 E2M1:1 个符号位、2 个指数位、1 个尾数位,再乘以每 16 个元素一组的 FP8 E4M3 局部缩放和一个 FP32 全局缩放。因此真实值不是一个孤立的 4 bit 码字,而是: $$x_{\mathrm{recon}}=x_{\mathrm{E2M1}}\,s_{\mathrm{block}}\,s_{\mathrm{global}}$$ 只数存储位数,忽略对齐和打包时,长度为 $L$ 的张量平均每元素约用 $4+8/16+32/L$ bit:大张量趋近 4.5 bit/元素,相对 16 bit 的 BF16 理想压缩比约 $16/4.5=3.56$,而不是整齐的 4 倍。实际实现还受 padding、布局和额外元数据影响。我们的 e2m1_toy 用最近邻 E2M1 可表示值和精确的浮点块缩放说明概念,没有模拟 FP8 缩放舍入、全局缩放、硬件矩阵乘法或 NVFP4 的完整训练配方,不能用它的误差预测真实 NVFP4 模型质量。 分组大小还有一笔容易被省略的元数据账。假设某种 INT4 实现对每 32 个权重同时保存一个 16 bit scale 和一个 16 bit zero point,理想平均是 $4+(16+16)/32=5$ bit/权重,相对 BF16 的理论压缩比是 $16/5=3.2$。若每组 128 个权重,在同一假设下是 $4.25$ bit/权重、约 $3.76$ 倍。这里的数字只是指定元数据布局后的算术示例;AWQ 不同内核可能把零点、scale 以其他精度或打包方式存放,padding 也会改变实际占用。组变小往往改善局部拟合,却可能增加元数据带宽和内核约束。因此选组大小时,至少同时报告量化后的实际显存、输出质量和目标设备延迟,而不能只报告“4 bit”。 还要区分权重量化与激活量化的计算链。W4A16 常见路径保留较高精度激活,只压缩权重;W8A8 则要求两侧都进入 8 bit 路径,整数累加再按 scale 还原。前者通常能让权重读取变轻,却不自动让激活流量变小;后者更可能用 INT8 矩阵乘硬件,但激活离群问题也更直接。训练中的 FP4、推理中的 FP4、权重专用 INT4 同样不能只按“都是四位”合并对比,它们量化统计、可接受误差和内核目标都不同。 3.5 再把性能账接回来 设原权重数为 $KN$。BF16 仅权重约 $2KN$ 字节,裸 INT8 约 $KN$ 字节,裸 4 bit 约 $KN/2$ 字节;后两者还要加各自的缩放、零点及对齐开销。这是容量账。速度账要写成性能建模那篇的下界: $$T\ge\max(F/P,\ D/B)$$ $F$ 是内核运算量,$P$ 是该硬件上该精度路径的有效算力,$D$ 是实际搬运字节,$B$ 是有效带宽。若小 batch 解码反复读大权重,压缩 $D$ 可能有效;若 prefill 的矩阵乘已受算力限制,就要有真正的 INT8 或 FP4 计算内核才可能受益。若先把低比特权重展开回 BF16,反量化、临时缓冲和 kernel launch 也要入账。低比特文件大小、显存占用、单次延迟和吞吐量是四个不同的观测量,必须分别报告。 04. 代码实现 完整可运行的 quant_demo.py 和 make_figures.py 放在文末附录。只需 Python、NumPy、Pillow;没有模型权重、CUDA 或外部数据下载。固定随机种子 7,前 192 条输入作校准,后 64 条作评估,避免直接拿选择缩放的样本当成绩。张量形状是 $X_{\text{cal}}\in\mathbb R^{192\times64}$、$X_{\text{eval}}\in\mathbb R^{64\times64}$、$W\in\mathbb R^{64\times32}$。第 0 激活通道被人为放大,相应权重通道被缩小;这是为了把问题放到显微镜下,不代表真实模型普遍这样分布。 最核心的四行与 3.2 节逐项对应:s 是公式的 $s$,x_s 是 $XD^{-1}$,w_s 是 $DW$。先用未量化的矩阵乘核对代数等价,再分别走 INT8 路径: s = smooth_scales(x_cal, w, alpha=0.5) x_s, w_s = x_eval / s, w * s[:, None] assert np.allclose(x_eval @ w, x_s @ w_s) after = int8_matmul(x_cal / s, x_s, w_s) int8_matmul 把评估输入量化成整数、权重量化成整数,真正用 int32 矩阵乘累加,然后按激活和每输出通道的权重步长反量化。它没有调用真实 GPU INT8 Tensor Core,所以这里只能谈数值误差,不能谈运行速度。AWQ 教学路径也故意返回反量化浮点权重以便检查输出;若要测 W4A16 吞吐,必须换成真实打包内核。 本次在项目 Python 3.11 + NumPy 上执行 python outputs/fundamentals_files/quantization/code/quant_demo.py 的原始输出如下,MSE 均针对同一个留出集的浮点 $XW$;relative 是 MSE 除以该输出的均方值: X_cal (192, 64) X_eval (64, 64) W (64, 32) Smooth alpha=0.50, s[0]=16.2463, median(s)=0.9942 equivalent transform max_abs_error=3.553e-15 int8_before MSE=0.222085 relative=0.015861 int8_after MSE=0.001384 relative=0.000099 int4_plain MSE=0.492117 relative=0.035146 int4_awq MSE=0.089772 relative=0.006411 fp4_toy MSE=0.335760 relative=0.023979 AWQ toy alpha=0.95 calibration_MSE=0.091065 第一个可检验的结论是等价变换在浮点下只剩 $3.553\times10^{-15}$ 的舍入差。第二个是这组构造数据里 SmoothQuant 式迁移降低了 W8A8 的输出误差。第三个是 AWQ 教学搜索把 INT4 权重量化的留出集误差从 0.492117 降到 0.089772;选择指数 0.95 只是这批构造样本上的选择。fp4_toy 的 0.335760 不能与前三者直接排模型名次:它的缩放编码与计算路径不同,而且没有真实 FP8 scale 舍入。图 1 由同一组数组生成,因此数值与图是可追溯的一套实验。 05. 工业级实现对照 SmoothQuant 的工业细节在「把除法折进上一层」。 作者公开实现的 smooth_ln_fcs 先从校准激活和多个相邻线性层的权重求通道尺度,再对 LayerNorm 的 weight、bias 除以尺度,并对后续线性层的相应输入权重乘以尺度。多个 Q/K/V 投影可能共享上一层归一化,不能各自随意选一套不一致的缩放。本文代码直接显式写 x/s,便于看清代数关系;真正部署需确认融合位置、归一化类型与残差分支,不能机械把所有线性层都改一遍。链接以该仓库所示 commit 为准。 AWQ 的工程路径比教学版多了布局与搜索。 作者 pseudo_quantize_tensor 接受位宽、零点与组大小,把权重按末维分组,计算每组 min/max、scale、zero,再舍入裁剪;auto_scale_block 用校准输入取模块原输出、扫 20 个尺度候选并比较输出 MSE。本文的 int4_groups 与 awq_toy 只复现这两个可解释步骤。生产还要选择组大小、权重打包顺序、激活精度、融合反量化的 GEMM,以及对 tokenizer、任务和模态有代表性的校准数据。伪量化权重占用浮点内存,不能拿它冒充真实 4 bit 显存占用。 NVFP4 则是格式与硬件共同定义的方案。 NVIDIA Transformer Engine 文档 明确写出 E2M1 码字、16 元素 FP8 局部缩放和 FP32 全局缩放,也讨论训练时缩放与舍入细节。这与 AWQ 的非对称 INT4 零点方案不同。若在 Blackwell 以外的设备上模拟 E2M1,再调用普通 BF16 GEMM,只是在研究量化误差;没有证据证明获得 NVFP4 Tensor Core 吞吐。对视频 DiT,还应分别统计文本投影、注意力、MLP、VAE 与不同去噪时刻的敏感度,而不是用一个 LLM 基准替代画质验收。 5.1 把这套办法移到视频 DiT 时怎么验 先在模型推理图中列出每个线性层的输入形状、权重字节、调用次数与实测耗时,按去噪步、分辨率和 batch 分桶。一个只调用一次的小投影层即使压到 4 bit,也不如在每一步反复运行的大 MLP 值得优先优化。随后对候选层采样真实生成输入:文本条件、无条件分支、不同帧长和空间分辨率、去噪早中晚步都应覆盖。统计每个输入通道的最大值、分位数及其随条件变化的范围,再决定是对称、非对称、按通道还是按组,而不是从 LLM 的一套 scale 直接复制过来。 然后做逐层替换实验:固定随机种子,只量化一组层,记录该层输出误差、整网中间特征差、生成结果与延迟。若数值误差集中在少数层,可保留这些层为 BF16,其他层继续压缩;但必须把混合精度带来的格式转换也计入耗时。对视频尤其要检查相邻帧的纹理闪烁和文字稳定性,因为单帧指标可能掩盖时间不一致。最后把“模型文件大小、峰值显存、单次生成时延、吞吐、质量”五列并排记录,并列出硬件型号、内核版本和量化格式。这样才能判断收益来自低位宽计算、权重少读、还是单纯把模型装进了原先放不下的卡。 本文的 CPU 脚本只覆盖上述流程的数值诊断第一步,不覆盖真实 DiT 校准、低比特 kernel 或视频验收。这一边界并非小字备注:若读者要据此选择线上格式,必须在实际模型和实际设备上补完后续测量。 06. 代价与边界 精度代价首先来自动态范围:静态校准没见过的 prompt、分辨率或高运动视频,可能把激活推过已选范围,产生裁剪而非普通舍入误差。多步扩散还会把每一步的小偏差沿采样轨迹累积。本文的线性层输出 MSE 只能当局部诊断;最终还需在固定种子和足够多提示词上比较画面结构、文本可读性、时间一致性与人评,不能把 MSE 下降直接翻译成 VBench 上升。 系统代价是尺度与布局。按张量一个 scale 便宜却易被离群值支配;按通道、按组或按 16 元素块能更细地适应分布,但多了元数据、打包、反量化与特定内核限制。对容量,要记 weights + scales + zeros + padding + workspace;对延迟,要测 prefill、单步 decode、小 batch 和大 batch。生产图中若有不支持低比特的算子导致来回格式转换,转换成本可能抵掉 GEMM 收益。 方法边界也不同。SmoothQuant 的迁移假设权重能承受放大,且前后层可安全折叠;AWQ 搜索的是给定校准集、给定位宽和组大小的局部误差,不能保证分布外质量;FP4 的动态范围和非均匀码字不等同于 INT4,不能把一个格式的超参数照搬到另一个。没有兼容内核时,先考虑稳定的 BF16/FP16 路径或更保守的 INT8。对已经被算力、launch 或跨卡通信限制的工作负载,权重文件变小也可能几乎不缩短时间。 一个容易漏掉的验收维度是误差的结构。同样的总体 MSE,误差若集中在少数文本 token、眼睛边缘或连续帧的同一位置,主观影响可能远大于均匀噪声。因此可把误差按层、通道、去噪时刻和内容类型拆开画分布,并检查最大误差与高分位数,而非只看全局平均。对生成视频,建议把原始模型与量化模型用相同 prompt、相同随机种子逐对比较;除整体指标外,抽查运动边界、物体身份、OCR 文本、暗部纹理和高频闪烁。若只是量化器的局部输出误差变小,却没有带来终端质量改善,也应如实报告“局部数值改善,端到端效果未证实”。这种分层验收也能告诉工程师应该先恢复哪一层的精度,而不必把整网退回 BF16。 07. 经典论文脉络 Jacob 等,Quantization and Training of Neural Networks for Efficient Integer-Arithmetic-Only Inference,CVPR 2018 / arXiv:1712.05877:把 scale、zero point 与整数推理串成可部署的量化路径,是理解本文 3.1 节映射的起点;它的主要实验对象是移动端视觉网络。 Xiao 等,SmoothQuant: Accurate and Efficient Post-Training Quantization for Large Language Models,arXiv:2211.10438:指出 LLM 激活离群值阻碍 W8A8,通过等价通道变换把难度迁到权重侧,目标是同时量化权重与激活。 Frantar 等,GPTQ: Accurate Post-Training Quantization for Generative Pre-trained Transformers,arXiv:2210.17323:走权重低比特的另一条路,用近似二阶信息做一次性量化与误差补偿;它说明 INT4 不是只能靠逐元素四舍五入。 Lin 等,AWQ: Activation-aware Weight Quantization for LLM Compression and Acceleration,arXiv:2306.00978:把激活统计用于识别重要权重,并通过等价缩放保护输出;与 SmoothQuant 都用通道缩放,但目标分别是权重低比特和 W8A8,不能混为一篇算法。 从这四篇再看 NVFP4 格式文档,会发现研究问题已经从「如何挑量化码字」延伸到「硬件支持哪种码字、块尺度和矩阵布局」。本文关于 NVFP4 的字段与块大小以该官方文档为准;本文没有把它列作四篇论文之一。 08. 常见误解 误解一:INT4 一定比 INT8 快两倍。 4 bit 只说明每个裸权重码字更短。若内核先展开、如果 GEMM 已受算力限制、若尺度加载很重,延迟可能没有相应收益。必须在目标 GPU 上测端到端而非只量文件大小。 误解二:SmoothQuant 把模型函数改好了,所以精度一定升。 它在浮点下是等价的;被改善或恶化的是后续量化误差。不同层、不同 $\alpha$、不同校准集都可能改变结论。本文的 160 倍左右 MSE 改善来自特意制造的离群通道,不是通用倍率。 误解三:AWQ 只看权重最大值。 3.3 节说明输出误差被输入放大。官方搜索用校准激活和模块输出误差,不是简单挑最大的权重元素。本文虽然也使用权重分组 min/max,但那是码字映射,不是重要性判据。 误解四:FP4 就是 INT4 加一个不同的 scale。 E2M1 的码字间隔非均匀,NVFP4 还有 16 元素局部 FP8 scale 与全局 scale。即使两个方案都写「4 bit」,误差形态、元数据和硬件路径也不同。 误解五:单层 MSE 合格,视频质量便合格。 视频生成的时间一致性、字符与细节可能对少数层或特定步数敏感;线性层 MSE 只帮助定位,应继续做固定种子的视频样例、人评和任务指标核验。 09. 动手验证 把附录两段代码分别保存为 quant_demo.py 和 make_figures.py,放在同一目录,安装 numpy 与 Pillow 后运行: python quant_demo.py python make_figures.py 图会写到上一级 figures/smooth_migration.png。先确认 equivalent transform max_abs_error 接近零,再把 case() 里的 x[:, 0] *= 20.0 改为 *= 1.0 重跑:离群值消失后,普通 INT8 与 SmoothQuant 的误差差距应明显缩小,但具体数值以你的实跑结果为准。第二个实验把 awq_toy 中的候选指数固定为 0,观察 int4_awq 是否退化到 int4_plain 附近;这对应“不使用激活引导缩放”的基线。第三个实验把分组大小从 32 改成 16 或 64,注意 int4_groups 要同步修改调用处,比较输出 MSE 与理论元数据开销:更细的组通常更能适应局部分布,但真实内核性能仍需另测。 10. 延伸阅读 若不清楚 BF16、FP16 与 FP8 的表示范围,先读混合精度与数值稳定性;若想判断量化后为何没提速,回看性能建模与 Profiling的带宽、算力与容量账;若关心自回归生成的权重带宽之外还有哪些显存项,接着读KV Cache 与自回归视频生成。下一步面对实际视频 DiT,应把本文的输出误差实验移到真实层、真实校准集和多个去噪时刻,再用质量指标与人评决定哪些层保留高精度。 附录:完整代码 09 节用到的脚本全文如下(quant_demo.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 quant_demo.py """Minimal CPU quantization experiment. Requires numpy; run: python quant_demo.py.""" import numpy as np def case(): rng = np.random.default_rng(7) x = rng.normal(0, 0.6, (256, 64)) x[:, 0] *= 20.0 w = rng.normal(0, 0.8, (64, 32)) w[0, :] *= 0.08 return x[:192], x[192:], w def mse(y, y_hat): return float(np.mean((y - y_hat) ** 2)) def smooth_scales(x_cal, w, alpha=0.5): a = np.maximum(np.max(np.abs(x_cal), axis=0), 1e-8) b = np.maximum(np.max(np.abs(w), axis=1), 1e-8) return a**alpha / b**(1.0 - alpha) def int8_matmul(x_cal, x_eval, w): # Static per-tensor activation scale, per-output-channel weight scales. sx = max(float(np.max(np.abs(x_cal))) / 127.0, 1e-8) sw = np.maximum(np.max(np.abs(w), axis=0) / 127.0, 1e-8) qx = np.clip(np.rint(x_eval / sx), -127, 127).astype(np.int32) qw = np.clip(np.rint(w / sw), -127, 127).astype(np.int32) return (qx @ qw).astype(np.float64) * sx * sw def int4_groups(w, group=32): # W is [input, output]; each output row is split across input groups. out_dim = w.shape[1] rows = w.T.reshape(-1, group) lo = rows.min(axis=1, keepdims=True) hi = rows.max(axis=1, keepdims=True) scale = np.maximum((hi - lo) / 15.0, 1e-8) zero = np.clip(np.rint(-lo / scale), 0, 15) q = np.clip(np.rint(rows / scale) + zero, 0, 15) return ((q - zero) * scale).reshape(out_dim, -1).T def awq_toy(x_cal, w): # Search a channel scale using calibration output error, as in AWQ's idea. importance = np.maximum(np.mean(np.abs(x_cal), axis=0), 1e-8) target = x_cal @ w trials = [] for alpha in np.linspace(0.0, 0.95, 20): s = importance**alpha s /= np.sqrt(s.max() * s.min()) w_hat = int4_groups(w * s[:, None]) / s[:, None] trials.append((mse(target, x_cal @ w_hat), float(alpha), w_hat)) return min(trials, key=lambda row: row[0]) def e2m1_toy(w, block=16): # Nearest E2M1 value with exact float64 block scale; NOT NVFP4 encoding. levels = np.array([0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0]) out_dim = w.shape[1] rows = w.T.reshape(-1, block) scale = np.maximum(np.max(np.abs(rows), axis=1, keepdims=True) / 6.0, 1e-8) normalized = np.abs(rows) / scale idx = np.abs(normalized[..., None] - levels).argmin(axis=-1) return (np.sign(rows) * levels[idx] * scale).reshape(out_dim, -1).T def results(): x_cal, x_eval, w = case() y = x_eval @ w s = smooth_scales(x_cal, w) x_s, w_s = x_eval / s, w * s[:, None] before = int8_matmul(x_cal, x_eval, w) after = int8_matmul(x_cal / s, x_s, w_s) awq_cal_error, awq_alpha, w_awq = awq_toy(x_cal, w) w_int4 = int4_groups(w) w_fp4 = e2m1_toy(w) return { "x_cal": x_cal, "x_eval": x_eval, "w": w, "s": s, "x_s": x_s, "w_s": w_s, "exact_error": float(np.max(np.abs(y - x_s @ w_s))), "int8_before": mse(y, before), "int8_after": mse(y, after), "int4_plain": mse(y, x_eval @ w_int4), "int4_awq": mse(y, x_eval @ w_awq), "awq_alpha": awq_alpha, "awq_cal_error": awq_cal_error, "fp4_toy": mse(y, x_eval @ w_fp4), "signal": float(np.mean(y**2)), } def main(): r = results() print("X_cal", r["x_cal"].shape, "X_eval", r["x_eval"].shape, "W", r["w"].shape) print("Smooth alpha=0.50, s[0]=%.4f, median(s)=%.4f" % (r["s"][0], np.median(r["s"]))) print("equivalent transform max_abs_error=%.3e" % r["exact_error"]) for name in ("int8_before", "int8_after", "int4_plain", "int4_awq", "fp4_toy"): print("%s MSE=%.6f relative=%.6f" % (name, r[name], r[name] / r["signal"])) print("AWQ toy alpha=%.2f calibration_MSE=%.6f" % (r["awq_alpha"], r["awq_cal_error"])) if __name__ == "__main__": main() make_figures.py """Draw a two-panel SmoothQuant diagram. Requires numpy and Pillow.""" from pathlib import Path import numpy as np from PIL import Image, ImageDraw, ImageFont from quant_demo import case, smooth_scales def font(size): path = Path("C:/Windows/Fonts/arial.ttf") return ImageFont.truetype(str(path), size) if path.exists() else ImageFont.load_default() def panel(draw, left, title, before, after): top, width, height = 95, 570, 280 lo = min(float(before.min()), float(after.min())) hi = max(float(before.max()), float(after.max())) low, high = np.floor(np.log10(lo)), np.ceil(np.log10(hi)) draw.rectangle((left, top, left + width, top + height), outline="#9ca3af", width=2) for power in range(int(low), int(high) + 1): value = 10.0**power y = top + height - (power - low) / (high - low) * height draw.line((left, y, left + width, y), fill="#e5e7eb", width=2) draw.text((left - 62, y - 12), f"{value:g}", fill="#4b5563", font=font(21)) draw.text((left, 42), title, fill="#111827", font=font(28)) for values, color in ((before, "#e76f51"), (after, "#2563eb")): points = [] for i, val in enumerate(values): x = left + i / (len(values) - 1) * width y = top + height - (np.log10(val) - low) / (high - low) * height points.append((float(x), float(y))) draw.line(points, fill=color, width=4) draw.ellipse((points[0][0] - 6, points[0][1] - 6, points[0][0] + 6, points[0][1] + 6), fill=color) draw.text((left, top + height + 14), "channel 0", fill="#4b5563", font=font(19)) draw.text((left + width - 100, top + height + 14), "channel 63", fill="#4b5563", font=font(19)) def main(): x_cal, _, w = case() s = smooth_scales(x_cal, w) before_x = np.max(np.abs(x_cal), axis=0) after_x = np.max(np.abs(x_cal / s), axis=0) before_w = np.max(np.abs(w), axis=1) after_w = np.max(np.abs(w * s[:, None]), axis=1) canvas = Image.new("RGB", (1400, 510), "#ffffff") draw = ImageDraw.Draw(canvas) panel(draw, 110, "Activation channel maximum", before_x, after_x) panel(draw, 800, "Weight input-channel maximum", before_w, after_w) draw.line((860, 455, 910, 455), fill="#e76f51", width=5) draw.text((920, 441), "before", fill="#111827", font=font(22)) draw.line((1040, 455, 1090, 455), fill="#2563eb", width=5) draw.text((1100, 441), "after smoothing", fill="#111827", font=font(22)) out = Path(__file__).resolve().parent.parent / "figures" / "smooth_migration.png" out.parent.mkdir(parents=True, exist_ok=True) canvas.save(out) print("saved", out, "bytes", out.stat().st_size) if __name__ == "__main__": main()
2026年10月04日
2 阅读
0 评论
0 点赞
2026-10-03
AIGC 基本功|混合专家 MoE:稀疏激活怎么省算力-MoE
混合专家 MoE:稀疏激活怎么省算力 所属方向:注意力与位置编码 | 难度:进阶 | 前置知识:自注意力机制的计算与显存账本 关键词:MoE、稀疏激活、专家路由、top-k、负载均衡、容量溢出 本文的算例在 CPU 上用 PyTorch 实际运行;它验证路由与容量机制,并不代表 GPU 集群吞吐。参数量账本只计算专家前馈层,明确排除注意力、优化器状态与通信。两篇锚点论文的标题与编号分别核对了 Shazeer 等人,1701.06538 和 Fedus 等人,2101.03961。 01. 为什么需要它 设想你已经把 Transformer 的注意力优化得很好,但还想让模型记住更多视觉对象、动作模式和语言知识。直接把每层前馈网络加宽,参数量和每个 token 的前向计算都会一起涨;训练显存、推理算力也跟着涨。MoE 的出发点是:保留多套前馈参数,让每个 token 只调用其中少数几套。总容量可以长得快,单 token 实际执行的专家计算长得慢。 先看一笔能复算的账。取隐藏宽度 $d=4096$、SwiGLU 中间宽度 $h=14336$。一个专家有 gate、up、down 三块矩阵,忽略偏置后共有 $3dh=176,160,768$ 个参数。放八个专家,总计 $1,409,286,144$ 个专家参数;每个 token 选两个,只触及 $352,321,536$ 个专家参数。在“同样八个专家都计算”的假想稠密基线下,专家矩阵乘的主项是四分之一。这里的“四倍”只属于这一层的专家部分:注意力、路由、分发、合并、跨设备通信一个都没算进去。八套 bf16 专家权重仍占约 2.625 GiB,不会因为每次只用两套就自动缩成四分之一。 真正会让方案失效的是路由。假如 16 个 token 中有 10 个奔向同一专家,其他专家分别只有 3、2、1 个,最忙的卡可能决定整步耗时。若像 Switch Transformer 那样给每个专家固定容量 4,六个 token 会溢出;把容量提到 6,仍有四个溢出。附录的 CPU 实验确实打印了 6/16 与 4/16。容量继续提到 10 才没有溢出,却要为大量空槽留空间。稀疏激活解决了“每个 token 算太多专家”的问题,不能自动解决“专家分工是否均匀”的问题。 图中上半部分故意把总 token 数保持为 16:均衡分配是 [4,4,4,4],坍缩分配是 [16,0,0,0];下半部分用 [10,3,2,1] 计算容量为 4、6、10 时的溢出。它画的是调度代价,不是模型质量曲线。 02. 最小可用理解 第一句:MoE 通常替换 Transformer block 的前馈子层,注意力仍先让 token 互相交换信息,然后每个 token 在前馈阶段自己选专家;Mixtral 原论文明确采用每个 token 选两个 SwiGLU 专家。第二句:路由器先给所有专家打分,再只执行 top-k;参数总数随专家数增加,但单 token 专家计算主要随 $k$ 增加。第三句:专家偏科会造成容量溢出、尾部等待和训练不稳,因此必须把专家负载、溢出率与通信量和任务损失一起看。 它不是“每个专家懂一种人类可命名的技能”的硬分工。专家是网络权重,路由是在每一层、对每个 token 重新做的决定。同一个视频片段里的不同 patch,或同一个句子的不同 token,都可能走不同专家;解释一个专家“负责什么”需要额外分析,不能凭编号想象。 03. 数学推导 3.1 从一层普通前馈网络出发 先把一个 token 的隐藏向量记作 $u$,维度为 $d$。SwiGLU 前馈层可写成: $$F(u)=W_{\mathrm{down}}\bigl(\operatorname{SiLU}(W_{\mathrm{gate}}u)\odot W_{\mathrm{up}}u\bigr)$$ $W_{\mathrm{gate}}$ 和 $W_{\mathrm{up}}$ 各把 $d$ 维映射到 $h$ 维,$W_{\mathrm{down}}$ 再映射回 $d$ 维;$\odot$ 是逐元素乘法。三块矩阵分别有 $dh$、$dh$、$hd$ 个权重,所以主参数量是 $3dh$。一次前向的矩阵乘主项也近似与 $3dh$ 成正比;SiLU、逐元素乘、读写权重及内核发射是额外成本。 现在复制出 $N$ 套这样的前馈层,记作 $F_1,\ldots,F_N$。如果每个 token 都执行全部 $N$ 套,参数量和前馈计算都乘 $N$,这只是昂贵的稠密集成。MoE 的关键是:路由仍观察全部 $N$ 个候选,但只执行其中 $k$ 个专家,通常 $k\ll N$。 3.2 路由概率与 top-k 合并 路由矩阵 $W_{\mathrm{route}}$ 的形状是 $N\times d$。对 token $u_t$,先得分,再做 softmax: $$z_{t,i}=(W_{\mathrm{route}}u_t)_i,\qquad p_{t,i}=\frac{\exp z_{t,i}}{\sum_{j=1}^{N}\exp z_{t,j}}$$ $t$ 是 token 序号,$i$ 是专家序号;$z_{t,i}$ 是尚未归一化的偏好,$p_{t,i}$ 是归一化概率。选出概率最大的 $k$ 个专家,形成集合 $S_t$。以 Mixtral 式 top-2 归一化为例,选中专家的合并权重是: $$a_{t,i}=\frac{p_{t,i}}{\sum_{j\in S_t}p_{t,j}},\quad i\in S_t;\qquad y_t=\sum_{i\in S_t}a_{t,i}F_i(u_t)$$ 分母只对已选中的专家求和,故其权重之和是 1。没有选中的专家既不运行前馈网络,也不贡献输出。路由计算本身的矩阵乘大致是 $Nd$,专家计算主项是 $k(3dh)$,所以在 $h$ 很大而 $k$ 很小时路由通常比专家矩阵乘小;实际耗时仍会受到 token 重排和通信影响。top-k 索引是离散的;反向传播在当前选中集合内可以沿连续权重求导,但不能把“换成另一个专家”的离散跳变当成普通连续导数。 这里有个容易漏掉的细节:若 $k=1$ 还把唯一选中的概率除以自己,合并权重恒为 1,主任务损失就无法经这个权重训练路由器。Switch 的 top-1 路由保留了所选概率作为乘数,并另加负载均衡损失;不能把上面的 top-2 归一化公式机械套到所有 top-1 实现上。实现差别要看原论文与代码,而不是只看“top-1”三个字。 3.3 “省算力”到底在比较什么 专家总参数约为 $P_{\mathrm{all}}=N(3dh)$;单 token 运行的专家参数约为 $P_{\mathrm{active}}=k(3dh)$,因此两者的比值是 $N/k$。这个推导只说明同一 MoE 层内全部专家与激活专家的差异。它没有证明 MoE 相对“参数更少、但充分训练的稠密模型”一定更快或更准;也没有证明端到端延迟会按 $N/k$ 缩短。公平比较要明确横轴是总参数、激活参数、每 token FLOPs、训练吞吐、墙钟时间中的哪一个。 以本文数值为例,$N=8$、$k=2$,所以 $N/k=4$。如果改成 $k=4$,专家计算主项会翻倍而总参数不变;如果把 $N$ 从 8 扩到 16 且保持 $k=2$,激活的专家计算主项近似不变,但权重驻留、路由维度与分布式通信压力会增加。这才是条件计算的交易:用更多存储与更复杂的调度,换更多参数容量与较少的激活计算。 3.4 为什么会需要负载均衡损失 取一批 $T$ 个 token,先讨论 Switch 的 top-1。令 $f_i$ 为实际派给专家 $i$ 的 token 比例,令 $P_i$ 为路由器分给该专家的平均概率。原论文使用下面的辅助项: $$f_i=\frac{1}{T}\sum_{t=1}^{T}\mathbf{1}[\operatorname{argmax}_j p_{t,j}=i],\qquad P_i=\frac{1}{T}\sum_{t=1}^{T}p_{t,i}$$ $$L_{\mathrm{aux}}=\alpha N\sum_{i=1}^{N}f_iP_i$$ $\mathbf{1}$ 是指示函数;$\alpha$ 控制辅助项相对主任务损失的强度。$f_i$ 是由离散 argmax 得到的实际负载,作为本批次统计量不走梯度;$P_i$ 连续可导,给路由器提供调整信号。若四个专家都恰好分到四分之一 token 且平均概率也是四分之一,则不乘 $\alpha$ 的项是 $4\times4\times(1/4)(1/4)=1$。若 16 个 token 全选第 0 个专家,且它的平均概率是 0.7112,同一项约为 $4\times0.7112=2.8449$。附录代码打印的就是这两个值。 不要把辅助项误认为“越小一定越好”的独立目标。它和语言、图像或视频的主任务损失一起优化;给得太重会强迫本来有意义的专家分化变得机械均匀,给得太轻会让路由坍缩。Shazeer 等人的早期方案分别约束 importance 与 load;Switch 把它简化成一个点积项。Switch 论文第 2.2 节报告了 $\alpha=10^{-2}$ 的实验设置,但这个数不是跨任务的通用常数。 3.5 容量与 token 溢出 路由概率决定“想去哪里”,硬件要决定“哪里放得下”。Switch 的 top-1 固定容量在概念上是: $$C=\left\lceil c\frac{T}{N}\right\rceil$$ $C$ 是每个专家本批次最多接收的 token 数,$c$ 是容量因子。若 $T=16$、$N=4$、$c=1$,每个专家只有四个槽;[10,3,2,1] 的第一位会溢出六个。把 $c$ 设成 1.5,容量变六、溢出四个;到 2.5 容量才变十、溢出为零。增加容量会减少丢弃,但静态张量中的空槽也增加。Switch 原论文说明:溢出的 token 跳过该专家层的计算,经残差路径传给下一层;这不是把 token 从整个网络删除。不同生产实现可以使用不同的溢出策略,不能从论文机制推断所有框架都会丢 token。 这一定义专门对应 top-1。top-2 每个 token 会产生两个派发名额,负载与容量的账要按派发次数重新算;直接拿 $T/N$ 给 top-2 算容量会少估。文章后面的代码将 top-2 正确性实验与 top-1 容量实验明确分开。 3.6 再检查三条可验证的守恒关系 把一批输入写成张量形状能看出程序为何必须“先拆再合”。设 batch 有 $B$ 条样本,每条 $S$ 个 token,压平后 $T=BS$;输入是 [T,d],路由 logits 是 [T,N],top-k 索引和权重都是 [T,k]。按专家重新排列后,第 $i$ 位专家只拿到自己的 $L_i$ 行,结果是 [L_i,d];再按原 token 索引累加回 [T,d]。只要每个 token 恰好被派发 $k$ 次,就必有 $\sum_i L_i=kT$。附录的 expert_loads=[8,7,8,9] 总和为 32,恰好等于 $16\times2$。这是发现漏派发或重复派发的第一道检查。 第二条关系来自合并权重。对同一个 token,Mixtral 式归一化后 $\sum_{i\in S_t}a_{t,i}=1$;若实测不等于 1,可能把全部专家的 softmax 概率直接拿来加权,却忘了对选中集合重新归一化。第三条关系是输出维度不变:专家前馈网络虽可扩到中间宽度 $h$,最终都要回到 $d$ 维,才能与 Transformer block 的残差相加。三条关系分别守住派发次数、权重尺度和残差形状,比直接观察最终 loss 更容易定位代码错误。 3.7 辅助项的梯度从哪里来 一次反向传播中,把当前 batch 已统计出的 $f_i$ 当作常量。softmax 的导数是 $\partial p_{t,i}/\partial z_{t,j}=p_{t,i}(\mathbf{1}[i=j]-p_{t,j})$:提高专家 $j$ 的 logit,会抬高它自己的概率,同时压低其他专家的概率。代入上一小节的 $L_{\mathrm{aux}}$,得到: $$\frac{\partial L_{\mathrm{aux}}}{\partial z_{t,j}}=\frac{\alpha N}{T}p_{t,j}\left(f_j-\sum_{i=1}^{N}f_ip_{t,i}\right)$$ 若专家 $j$ 的实际负载 $f_j$ 高于这个 token 所见的加权平均负载,梯度为正,梯度下降倾向于压低它的 logit。低负载专家可能得到相反推力。这只是一次梯度步的局部解释:$f_i$ 由 argmax 决定,会在选择边界突然跳变;主任务梯度、容量限制和数据分布也同时影响下一次路由。因此“有可导辅助项”不等于“必然均衡”,更不等于“均衡后质量一定最好”。 3.8 溢出与空槽是两本账 若第 $i$ 位专家收到 $L_i$ 个 token、每位静态容量都是 $C$,则溢出数为 $\sum_i\max(L_i-C,0)$,空槽数为 $\sum_i\max(C-L_i,0)$。用前面的 [10,3,2,1] 代入:$C=4$ 时溢出 6、空槽 6;$C=6$ 时溢出 4、空槽 12;$C=10$ 时溢出 0、空槽 24。容量越大,溢出变少,却给静态张量留下更多空位。空槽是否真的造成同等比例的计算浪费,还要看内核能否跳过填充;它至少会影响形状、存储或通信安排。 在多模态场景,整体平均负载还可能掩盖局部高峰。文本、图像 patch 与连续视频帧混在一批时,整批的四位专家统计可能接近均匀,但某一类 token 在某一层仍高度集中;按帧解码时的瞬时负载又可能不同于整段视频的平均。一个实用的诊断表应按层、模态和时间片记录最大 $L_i$、平均 $L_i$、溢出数与通信时间。这里是从路由与队列机制推出的监控建议,并非本文已经做过的多模态训练实测。 04. 代码实现 完整代码在文末 moe_lab.py,只依赖 PyTorch,CPU 上直接执行:python moe_lab.py。它先生成形状为 [16,4] 的 token 矩阵,以一个线性路由器选择四个专家中的两个。每个专家是一层很小的 Linear → ReLU → Linear,用于验证路由逻辑;前面的大模型参数账本仍按真实 SwiGLU 的三矩阵结构计算,不能把玩具专家误当成生产模型。 核心的稀疏计算可以缩到下面几行。top_idx 的形状是 [token,k],token_id 和 slot_id 一起定位“这个专家接收了哪个 token、占该 token 的第几个席位”;index_add_ 把两路专家结果加回对应 token。 probs = F.softmax(logits.float(), dim=-1) top_prob, top_idx = probs.topk(k, dim=-1) top_weight = top_prob / top_prob.sum(dim=-1, keepdim=True) result = torch.zeros_like(x) for expert_id, expert in enumerate(experts): token_id, slot_id = torch.where(top_idx == expert_id) if token_id.numel() == 0: continue value = expert(x[token_id]) * top_weight[token_id, slot_id, None] result.index_add_(0, token_id, value) print(result.shape, top_idx.shape) 固定随机种子为 7 后,脚本实际输出 input=(16, 4) output=(16, 4) top_idx=(16, 2),四位专家各接到 [8,7,8,9] 次派发,总和为 32,正好是 $16\times2$。它另用“全部专家都算、最后把没选中的输出乘零”的稠密参考结果做正确性对照,最大绝对误差在打印精度内为 0.000000000。参考实现会浪费计算,但适合检验稀疏实现的索引和加权有没有写错。 接着脚本不训练模型,而是构造两组受控路由 logits:一组让四个专家每人接四个 token,另一组让全部 token 偏向同一个专家。这样能把“负载变了”与“训练也变了”分开。打印结果为:均衡时 $f=P=[0.25,0.25,0.25,0.25]$,未乘 $\alpha$ 的辅助项是 1.0000;坍缩时 $f=[1,0,0,0]$、$P=[0.7112,0.0963,0.0963,0.0963]$,辅助项是 2.8449。容量实验又输出 6/16、4/16、0/16 三档溢出。读者可以独立复跑,不需要下载权重或数据集。 这组数字只证明公式、分发与计数代码互相吻合,不证明训练时加辅助项就一定达到均衡。真实训练还要监控各层的负载直方图、丢弃率、主任务损失、路由熵及设备间 all-to-all 的时间;单看一张 token 直方图不足以判断模型质量。 05. 工业级实现对照 以 2026-10-03 检查的 Hugging Face modeling_mixtral.py 为准,MixtralTopKRouter 先把隐藏状态展平,用 F.linear 得到 [token, expert] logits;它在 float32 中做 softmax,再 top-k,并对选中权重重新归一化。MixtralSparseMoeBlock 负责接路由输出与专家集合,最终还原 batch、sequence、hidden 三维。知识树里的 code_refs 指向这个文件的 MixtralSparseMoeBlock,路径与符号已经实际打开核对。 同一文件中的 MixtralExperts 不为每个 token 单独启动一次专家网络,而是先找出哪些专家被选中,再按专家收集 token,用 index_add_ 汇总。它把多位专家的矩阵存成带专家维的三维权重,并通过 @use_experts_implementation 接入不同执行实现。我们的小脚本保持“每个专家一个模块”,方便看公式;生产代码则必须考虑连续内存布局、分组矩阵乘、编译路径和设备利用率。文件在主分支上会变,正文只核对了本日可见的结构与函数名,没有声称这些细节永久不变。 Mixtral 的 top-2 与 Switch 的 top-1 还不能混成一个实现。前者在路由器里对两个已选概率归一化;后者的论文重点是单专家派发、静态容量、溢出路径与辅助负载损失。Hugging Face 这份 Mixtral 文件的前向代码不是 Switch 论文的 TPU 分布式训练系统,也不能从它有 index_add_ 就断言大集群没有 all-to-all。若专家分布在多卡上,token 必须先到相应设备、计算后再汇总;代价由网络拓扑、批量大小、token 倾斜和实现决定。 还有一个工业层面的数字边界:本文的 3dh 只含 SwiGLU 的三个权重矩阵。真实 Transformer block 还有 Q、K、V、输出投影、归一化、嵌入及可能的共享专家。训练时优化器状态和梯度常比 bf16 权重本身占更多空间;推理时即使每个 token 只用两位专家,服务系统仍需让其他专家随时可访问,或者承担按需调入的时延。因此“激活参数少”不等于“部署内存少”。 06. 代价与边界 收益首先是参数容量对激活计算的比值。 对 $N$ 位同宽专家取 $k$ 位,专家主计算的理想比值是 $N/k$;训练可在近似固定的每 token 专家 FLOPs 下试更多参数。它是否换来更好的任务质量,仍要做同预算、同数据、同训练时长的实验。早期论文报告过优于稠密对照的结果,但不能把某篇任务的收益搬给任意视觉生成模型。 第一笔代价是权重与状态。 参数量按 $N$ 增长,bf16 权重、梯度、优化器状态与检查点大小都要算。只看前向激活的 $k$ 位专家会严重低估训练资源。专家可以并行放在不同设备上,但又引出网络通信。 第二笔代价是路由偏斜。 每个专家的处理时间近似由收到的 token 数决定,整步会等最慢专家。容量限制可以把形状固定,却在溢出时改变有效计算;容量放大又会浪费填充。高平均负载不够,还要检查最忙专家、各层差异、长尾 token 和各设备之间的均衡。附录中的 [10,3,2,1] 正是最小反例。 第三笔代价是通信与小批量。 当专家跨设备切分,派发和合并通常需要集体通信;batch 很小或自回归逐 token 解码时,每位专家分到的 token 少,矩阵乘可能跑不满,固定通信开销更显眼。用单机 CPU 玩具实验不能估算这部分,也不能用“理论 FLOPs 四分之一”宣称端到端四倍加速。 第四笔代价是路由训练本身。 top-k 的硬选择会让专家早期获得的样本不均,冷门专家因训练较少又更难被选中,形成反馈。负载项、噪声、容量因子与专家初始化都会影响稳定性,但也可能伤害有意义的专门化。需要同时报告主任务指标与路由指标,而不是把“每个专家等量接单”当终点。 什么时候不该急着用?模型还小、普通前馈层已足够,或者服务以低 batch、低延迟为主且跨卡通信昂贵时,先用 性能建模与 Profiling 找到实际瓶颈。如果数据与训练预算不足以让多位专家学出差异,多出来的参数只会变成管理负担。这是工程判断,不是 MoE 论文对所有应用的普适否定。 07. 经典论文脉络 Shazeer 等,2017,Outrageously Large Neural Networks:把稀疏门控专家用于大容量模型,明确讨论专家偏科的自强化现象,并分别引入 importance 与 load 的均衡约束。它是理解“为什么光有 top-k 不够”的起点。 GShard,2020:把专家路由与自动分片放进大规模 Transformer 训练,提醒我们专家计算之外还有设备布局与通信问题。 Switch Transformers,2021/2022:用 top-1 简化派发,给出容量因子、溢出处理和可微的 $f\cdot P$ 负载项;本文第 3.4 与 3.5 节主要沿它推导。 ST-MoE,2022:系统研究稀疏专家训练的稳定性和迁移,说明“能扩参数”之后仍要解决训练动态。 Mixtral of Experts,2024:在每个相关层采用 top-2 SwiGLU 专家,是对照当前公开推理实现的具体案例;本文并不把它的结果视为所有视频或多模态 MoE 的结果。 这五篇连起来的主线是:先提出条件计算,再解决大规模分布式派发,接着简化路由与容量控制,最后面对稳定性和实际模型实现。论文里的速度、质量数字都有各自的设备与数据条件;这篇文章只借它们确认机制,不拼接成一个不存在的统一基准。 08. 常见误解 误解一:“八专家选二就是整个模型快四倍。” 四倍只来自专家矩阵乘的理想主项比值。注意力、路由、all-to-all、token 重排及等待最慢专家都在分母里。实验报告应分别列专家 FLOPs 与端到端时延。 误解二:“只激活两位专家,显存里只需放两套权重。” 路由在运行时根据每个 token 的状态变化,未选中的六位专家下一 token 可能被选到。权重需要驻留、分片或按需加载;哪一种都要付存储或调入成本。 误解三:“softmax 概率均匀就说明负载均衡。” 硬 top-1 派发由 argmax 决定,平均概率 $P_i$ 与实际负载 $f_i$ 是两个量。大量 token 的首选专家仍可能相同,所以要同时观察两者。Switch 的辅助项特意把它们相乘。 误解四:“溢出意味着 token 从模型中消失。” 在 Switch 论文描述的残差结构里,溢出 token 跳过这一专家层,后续层仍能接到它。它确实失去本层的专家变换,但不等于整条序列被删除。 误解五:“专家会自动对应数学、代码、图像等人类类别。” 路由学的是降低训练目标的内部划分,可能依赖位置、频率或更难解释的因素。必须用受控输入、路由统计和质量消融来验证专门化;给专家起名字只是叙事。 09. 动手验证 先执行文末两份完整脚本:python moe_lab.py 复现数值,python make_figures.py 生成配图。脚本固定随机种子,CPU 不需要模型权重。练习时一次只改一个条件,并把 主计算、路由负载、溢出率 分开记录。 把 k=2 改为 k=1。top_idx 应从 (16,2) 变成 (16,1),总派发次数从 32 变成 16。注意玩具实现会对唯一概率做归一化,权重恒为 1;这正好验证第 3.2 节提醒的训练梯度陷阱,因此它只能用来做前向路由演示,不能直接拿来训练 Switch。 保持 T=16,N=4,把不均匀分配 [10,3,2,1] 改成 [4,4,4,4]。容量因子 1.0 时,溢出应从 6/16 变 0/16。这是不改网络结构、只改路由就改变有效计算的最小实验。 把容量因子从 1.0、1.5、2.5 依次试过。原分配下脚本应给出容量 4、6、10 和溢出 6、4、0;同时算空槽总数:四位专家的静态槽分别是 16、24、40。容量增大并非免费。 将账本里的 $N=8,k=2$ 改成 $N=16,k=2$。专家总权重翻倍,单 token 激活专家参数不变,理想 $N/k$ 从 4 变 8。再想一想:若专家分散在更多设备上,为什么实际延迟未必更短? 若要进一步实验训练效果,可以在小数据集上加入任务损失与 $L_{\mathrm{aux}}$,逐档扫描 $\alpha$,同时画验证损失和各专家负载。预期是过小的 $\alpha$ 可能坍缩、过大的 $\alpha$ 可能牺牲任务目标;具体拐点依赖数据和实现,本文没有训练实验,因此不提供虚构数值。 10. 延伸阅读 先复习知识树的 自注意力机制:MoE 一般替换前馈子层,不替读者省去注意力的计算与显存账。然后读 性能建模与 Profiling,把“专家 FLOPs 少了”落到实际延迟账本;跨卡部署时再接 数据并行与 ZeRO,理解参数和状态如何分布。知识树中规划中的“量化:从 INT8 到 FP4”会继续回答:专家权重多、驻留贵时,能否用更低精度换内存与带宽,以及质量会付出什么代价。 把这篇用在自己的模型上时,建议先固定总训练 token、硬件与优化器配置,记录稠密前馈基线的质量和耗时;再逐步加入多专家、top-k 路由、容量限制和均衡项。每加一项都单独保存路由直方图、每层最忙专家、溢出比例及端到端吞吐。这样最后即使质量没有提升,也能判断问题出在专家容量不足、调度不均,还是通信把理论上的稀疏收益吃掉了。只报告总参数和激活参数两个数字,很难让读者复现或比较。另外,比较实验要记录随机种子和路由器初始化。路由决策依赖输入分布,同一个模型换一批数据,专家负载也可能改变;只截取一次顺利的运行,无法说明服务长期稳定。 本文最重要的读法是把三根轴分开:总参数量决定可用容量,激活专家数决定每 token 专家计算,实际路由负载决定系统是否跑得顺。 只有把它们放到同一张账本里,MoE 的收益和代价才不会被一个“四倍”掩盖。 附录:完整代码 09 节用到的脚本全文如下(moe_lab.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 moe_lab.py """A small, reproducible MoE routing lab. Requires PyTorch, runs on CPU.""" from __future__ import annotations import math import torch from torch import nn from torch.nn import functional as F def sparse_moe(x: torch.Tensor, logits: torch.Tensor, experts: nn.ModuleList, k: int): """Route each token to k experts and merge the selected outputs.""" probs = F.softmax(logits.float(), dim=-1) top_prob, top_idx = probs.topk(k, dim=-1) top_weight = top_prob / top_prob.sum(dim=-1, keepdim=True) result = torch.zeros_like(x) for expert_id, expert in enumerate(experts): token_id, slot_id = torch.where(top_idx == expert_id) if token_id.numel() == 0: continue value = expert(x[token_id]) * top_weight[token_id, slot_id, None] result.index_add_(0, token_id, value) return result, probs, top_idx, top_weight def switch_aux_loss(probs: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """Switch's top-1 balancing term, without the tunable alpha coefficient.""" token_count, expert_count = probs.shape chosen = probs.argmax(dim=-1) fractions = torch.bincount(chosen, minlength=expert_count).float() / token_count mean_prob = probs.mean(dim=0) return expert_count * (fractions * mean_prob).sum(), fractions, mean_prob def capacity_drops(assignments: torch.Tensor, expert_count: int, factor: float): """Count top-1 overflow with a fixed expert capacity.""" capacity = math.ceil(assignments.numel() * factor / expert_count) loads = torch.bincount(assignments, minlength=expert_count) dropped = torch.clamp(loads - capacity, min=0).sum().item() return capacity, loads.tolist(), int(dropped) def main(): torch.manual_seed(7) token_count, dim, hidden, expert_count, k = 16, 4, 8, 4, 2 x = torch.randn(token_count, dim) gate = nn.Linear(dim, expert_count, bias=False) experts = nn.ModuleList( [nn.Sequential(nn.Linear(dim, hidden), nn.ReLU(), nn.Linear(hidden, dim)) for _ in range(expert_count)] ) out, probs, chosen, weights = sparse_moe(x, gate(x), experts, k) # A dense reference evaluates every expert. It is only a correctness oracle. dense_ref = torch.zeros_like(x) for expert_id, expert in enumerate(experts): contribution = torch.zeros(token_count, 1) for slot in range(k): contribution += (chosen[:, slot] == expert_id)[:, None] * weights[:, slot, None] dense_ref += contribution * expert(x) max_error = (out - dense_ref).abs().max().item() print(f"input={tuple(x.shape)} output={tuple(out.shape)} top_idx={tuple(chosen.shape)}") print(f"expert_loads={torch.bincount(chosen.flatten(), minlength=expert_count).tolist()}") print(f"sparse_dense_max_error={max_error:.9f}") # Controlled router logits: the difference is caused by routing, not training. balanced_logits = torch.zeros(token_count, expert_count) balanced_logits[torch.arange(token_count), torch.arange(token_count) % expert_count] = 2.0 collapsed_logits = torch.zeros_like(balanced_logits) collapsed_logits[:, 0] = 2.0 for name, logits in (("balanced", balanced_logits), ("collapsed", collapsed_logits)): loss, fractions, mean_prob = switch_aux_loss(F.softmax(logits, dim=-1)) print(f"{name}: f={fractions.tolist()} P={[round(v, 4) for v in mean_prob.tolist()]} " f"N_sum_fP={loss.item():.4f}") uneven = torch.tensor([0] * 10 + [1] * 3 + [2] * 2 + [3]) even = torch.arange(token_count) % expert_count for name, assignments in (("even", even), ("uneven", uneven)): for factor in (1.0, 1.5, 2.5): capacity, loads, dropped = capacity_drops(assignments, expert_count, factor) print(f"{name} factor={factor:.1f}: capacity={capacity} " f"loads={loads} dropped={dropped}/{token_count}") width, expansion, model_experts, active = 4096, 14336, 8, 2 params_one = 3 * width * expansion # SwiGLU: gate, up, down matrices. params_all = model_experts * params_one params_active = active * params_one print(f"SwiGLU ledger: one={params_one:,} all={params_all:,} " f"active_per_token={params_active:,} all_over_active={params_all / params_active:.1f}x") print(f"expert_weights_bf16={params_all * 2 / 2**30:.3f} GiB; " "this excludes attention, optimizer states, activations and communication") if __name__ == "__main__": main() make_figures.py """Draw the MoE routing/capacity experiment. Requires Pillow only.""" from pathlib import Path from PIL import Image, ImageDraw, ImageFont ROOT = Path(__file__).resolve().parents[1] OUT = ROOT / "figures" OUT.mkdir(exist_ok=True) def font(size): for path in ("C:/Windows/Fonts/arial.ttf", "/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf"): if Path(path).is_file(): return ImageFont.truetype(path, size) return ImageFont.load_default() im = Image.new("RGB", (1200, 540), "#f8fafc") d = ImageDraw.Draw(im) title = font(31) text = font(21) small = font(17) d.text((48, 30), "MoE routing: same token count, different costs", font=title, fill="#0f172a") d.text((50, 100), "Top-1 token load across 4 experts (T = 16)", font=text, fill="#334155") for row, (name, counts, color) in enumerate(( ("Balanced", [4, 4, 4, 4], "#0ea5e9"), ("Collapsed", [16, 0, 0, 0], "#f97316"), )): y = 160 + row * 130 d.text((50, y + 23), name, font=text, fill="#0f172a") for i, value in enumerate(counts): x = 205 + i * 220 d.rectangle((x, y + 20, x + 175, y + 55), fill="#e2e8f0") if value: d.rectangle((x, y + 20, x + 175 * value // 16, y + 55), fill=color) d.text((x + 52, y + 65), f"E{i}: {value}", font=small, fill="#334155") d.line((48, 423, 1150, 423), fill="#cbd5e1", width=2) d.text((50, 443), "Uneven load [10, 3, 2, 1]: drops at capacity 4 / 6 / 10", font=text, fill="#334155") for i, (cap, drops) in enumerate(((4, 6), (6, 4), (10, 0))): x = 650 + i * 170 d.rounded_rectangle((x, 438, x + 145, 490), radius=8, fill="#dbeafe") d.text((x + 12, 453), f"C={cap}: {drops} drop", font=small, fill="#1e3a8a") path = OUT / "routing_capacity.png" im.save(path) print(path)
2026年10月03日
2 阅读
0 评论
0 点赞
2026-10-02
AIGC 基本功|算子融合与 CUDA Graph-Fusion
算子融合与 CUDA Graph 所属方向:推理加速 | 难度:进阶 | 前置知识:性能建模与 Profiling(performance_profiling)、混合精度(mixed_precision)、自注意力机制(attention_basics) 关键词:算子融合、kernel fusion、CUDA Graph、发射开销、访存瓶颈、torch.compile、Inductor 关于本文的数字:作者手里这台机器是 Apple M1 Pro,没有 NVIDIA GPU,也没有安装 torch。所以凡是标「实测」的数字,都出自文末附录里那三个能在纯 numpy 上跑通的脚本;凡是 CUDA/HBM 相关的数字,都来自公开资料并显式标注出处,我没有拿 CPU 的数字去外推 GPU。第六节 6.5 专门交代了这条边界。 01. 为什么需要它 先摆三个数,都是本篇附录里真跑出来的。 第一个数:一步 decode 大约要发射 1093 个 kernel。 这不是实测,是按 Llama-3-8B 的公开结构一层层手数出来的:每层 34 个 kernel(输入 RMSNorm 拆 4 个、q/k/v 投影各 1 个、两组 RoPE 各 3 个、KV 写入 2 个、QK 转置乘 1 个、缩放加掩码 2 个、softmax 3 个、乘 V 1 个、o 投影 1 个、残差 1 个、后置 RMSNorm 4 个、gate/up 各 1 个、SiLU 1 个、逐元素乘 1 个、down 投影 1 个、残差 1 个),乘 32 层再加输出头约 5 个,得到 1093。公开资料给的量级是 300~1300,这个数落在里面,可以互为印证。 第二个数:这 1093 个 kernel 里,绝大多数只干几微秒的活。 每次 kernel 发射在 CPU 侧要花 1~5 微秒——NVIDIA 开发者论坛上 njuffa 给的口径是「空 kernel 约 5 微秒」,《CUDA Handbook》实测 NULL launch 约 4.9 微秒(老机器)、GeForce RTX 3060 上约 1.2 微秒,A100 上常见引用是 2~3 微秒。假设发射 3 微秒、执行 2 微秒,附录 launch_model.py 的时间线推演给出的结果是:eager 模式总时长 2102 微秒,其中 GPU 空转 702 微秒,也就是 33% 的时间卡在等 CPU 把下一个 kernel 发过来;换成一次提交的 CUDA Graph,总时长降到 1753 微秒,加速 1.20 倍,GPU 基本不空转了。 第三个数:8 个串起来的逐元素算子,其中 93.8% 的时间花在「发射」上而不是「算」。 这是 fusion_lab.py 在 512 个元素的张量上实测的。同样一条 8 段链,在 400 万元素的张量上反过来:加速的 100% 来自少搬字节,发射部分只占 0%。同一个优化在两个极端上赚的是完全不同的钱,这件事后面会用一整套账把它算清楚。 还有一个更刺眼的数在注意力上。按 $B=8$、$H=16$、$S=2048$、$d=4096$、bf16 算:输入激活 $[B,S,d]$ 是 128.0 MiB,而 score 矩阵 $[B,H,S,S]$ 是 1.000 GiB,正好是输入的 8 倍。放大倍数是 $H \cdot S / d$——$S$ 翻一倍,它翻一倍;$S=8192$ 时是 32 倍。softmax 如果不融合,max / 减 / exp / 求和 / 除五步各读写一遍 score,就是 8.000 GiB 的访存;融合后只要 2.000 GiB。 你以为模型卡在算力上,其实它经常卡在「把中间结果写出去再读回来」和「让 CPU 一次又一次地通知 GPU 开工」这两件事上。 02. 最小可用理解 三句话: 融合(fusion):一串相邻算子的中间结果,本来每个算子都要写回显存、下一个再读回来。融合就是把它们放进同一个 kernel,中间结果留在寄存器或共享内存里,只在最开头读一次输入、最末尾写一次输出。它不产生新的数学,只是让已经算出来的东西少走几趟路。 CUDA Graph:那串 kernel 的发射动作(函数名、参数、依赖关系、显存地址)可以录下来,之后一次 cudaGraphLaunch 提交整张图。它一个字节都不省,省的是 CPU 侧那 1093 次提交。 两者的收益都能提前算出来,也都有限。融合的上限是字节比 $K$($K$ 段链最多省 $K$ 倍访存);CUDA Graph 的上限是「发射时间 / 执行时间」,而且只在单 kernel 执行时间接近或小于发射时间时才为正。 一句总结:融合省字节,Graph 省发射。把两者混为一谈,是这一块最常见的错误。 03. 数学推导 3.1 先回到那条判据:roofline 前置文章 performance_profiling 已经建过这条判据,这里只取结论。一个算子的耗时下界由两件事里更慢的那个决定: $$T \ge \max\left(\frac{F}{P_{\text{peak}}},\ \frac{B}{B_{\text{off}}}\right)$$ 其中 $F$ 是浮点运算次数(FLOPs),$P_{\text{peak}}$ 是峰值算力(FLOP/s),$B$ 是需要搬动的字节数(含读和写),$B_{\text{off}}$ 是片外内存带宽(字节/秒)。两者的比值就是算术强度: $$I = \frac{F}{B}$$ $I$ 高的算子卡在算力上,$I$ 低的算子卡在带宽上。逐元素算子的 $I$ 是常数:算一个元素做 1 次操作、读 4 字节写 4 字节,$I = 1/8$。在 fp32 下算力峰值除以带宽的量级是每字节几十次操作,所以逐元素算子永远在带宽那一侧。 这就是融合的着力点:它不改变 $F$,只把 $B$ 变小。 3.2 字节账本:K 段链为什么最多省 K 倍 考虑 $K$ 个串起来的逐元素算子,作用在一个 $N$ 元素的张量上($b$ 为每元素字节数)。 不融合:每个算子单独成一个 kernel,各自读一遍输入、写一遍输出。第 $k$ 个 kernel 的访存是 $2Nb$ 字节(因为输入和输出都是 $N$ 个元素),$K$ 个加起来: $$B_{\text{unfused}} = 2KNb$$ 融合:一个 kernel 从头做到尾。读输入 $Nb$、写输出 $Nb$,中间那 $K-1$ 步全在片上: $$B_{\text{fused}} = 2Nb$$ 两式相除,字节比正好是 $K$: $$\frac{B_{\text{unfused}}}{B_{\text{fused}}} = K$$ 本文附录 fusion_lab.py 的 [A] 节把这个账算在了三处:8 段链 244.1 MiB 降到 30.5 MiB(8 倍);RMSNorm 从 4 遍降到 1 遍,122.1 MiB 降到 30.5 MiB(4 倍);attention softmax 从 8.000 GiB 降到 2.000 GiB(4 倍),另外还有一份被物化的 mask 值 1.000 GiB,融合后直接消失。 注意字节比 $K$ 只是「访存减少到 1/K」,不是「时间减少到 1/K」。 时间还取决于中间结果到底停在多快的存储上。设片上带宽为 $B_{\text{on}}$,定义片内片外带宽比: $$\rho = \frac{B_{\text{on}}}{B_{\text{off}}}$$ 那么融合后的时间是「片外读一次写一次」加上「$K-1$ 步在片上走」: $$T_{\text{fused}} = \frac{2Nb}{B_{\text{off}}} + \frac{2(K-1)Nb}{B_{\text{on}}}$$ 不融合的时间是: $$T_{\text{unfused}} = \frac{2KNb}{B_{\text{off}}}$$ 两者相除,$N$ 和 $b$ 全部约掉: $$S = \frac{T_{\text{unfused}}}{T_{\text{fused}}} = \frac{K}{1 + (K-1)/\rho}$$ 这个式子值得盯着看三秒。 它的两个极限都很干净: $\rho \to 1$(片上不比片外快):$S \to 1$,融合白干。 $\rho \to \infty$(片上无限快):$S \to K$,也就是字节比。 所以「融合能加速几倍」这个问题,答案既不是 $K$,也不是玄学,而是 $K$ 和 $\rho$ 共同决定的。记住这一点,第六节会看到它把一台机器上的实测结果解释得干干净净。 3.3 单次调用的固定成本,和它的临界规模 一次算子调用的耗时,在规模足够小时和数据量无关。写成仿射形式: $$t(n) = a + b_{\text{el}} \cdot n$$ $a$ 是固定开销(派发、参数检查、缓冲建立),$b_{\text{el}}$ 是每元素的边际成本。开销和数据各占一半的那个规模是: $$n^{*} = \frac{a}{b_{\text{el}}}$$ fusion_lab.py 的 [D] 节实测这台机器:$a = 0.430$ 微秒/次,$b_{\text{el}} = 55.47$ 微秒/百万元素,$n^{*} = 7746$ 个元素,也就是 30.3 KiB。张量小于 30 KiB 时,大部分时间不是在算数据。 如果把这条链的 $2K$ 次调用全加起来,固定开销的占比是: $$\text{share} = \frac{2K a}{2K(a + b_{\text{el}} n)} = \frac{1}{1 + n/n^{*}}$$ 实测($K=8$,16 次调用):$n=512$ 时 93.8%,$n=4096$ 时 65.4%,$n=65536$ 时 10.6%,$n=1048576$ 时 0.7%。张量一小,你花在「组织计算」上的钱就超过「做计算」的钱。 这就是 GPU 上 kernel launch 开销在 CPU 上的同构物,也是为什么 GPU 上会出现同样性质的墙。 3.4 发射与执行的时间线:CUDA Graph 到底省什么 考虑单流、异步提交的 $N$ 个 kernel。CPU 发射第 $i$ 个要花 $t_{\text{launch}}$,发完就可以去发下一个,不等 GPU。GPU 执行第 $i$ 个要花 $t_{\text{exec}}$,它有两个前提:前一个跑完了,并且这一个已经发出去了。所以 GPU 开工时刻是: $$\text{start}_i = \max\left(\text{end}_{i-1},\ (i+1) \cdot t_{\text{launch}}\right)$$ 两个条件谁慢,谁决定节奏。分两种情形: $t_{\text{launch}} \le t_{\text{exec}}$:CPU 能跑在 GPU 前面,GPU 一个接一个不停,总时长约 $N \cdot t_{\text{exec}} + t_{\text{launch}}$。 $t_{\text{launch}} > t_{\text{exec}}$:CPU 喂不上,GPU 每跑完一个就得等 $t_{\text{launch}} - t_{\text{exec}}$,总时长约 $N \cdot t_{\text{launch}}$。GPU 空转的比例是 $1 - t_{\text{exec}}/t_{\text{launch}}$。 CUDA Graph 把 $N$ 次提交压成 1 次,但 GPU 前端仍然要逐个节点过一遍(记 $t_{\text{disp}}$): $$T_{\text{graph}} = t_{\text{launch}} + N \cdot (t_{\text{exec}} + t_{\text{disp}})$$ 于是加速比是: $$S_{\text{graph}} = \frac{\max(N t_{\text{launch}},\ N t_{\text{exec}})}{t_{\text{launch}} + N(t_{\text{exec}} + t_{\text{disp}})} \approx \frac{\max(t_{\text{launch}}, t_{\text{exec}})}{t_{\text{exec}} + t_{\text{disp}}}$$ 这个式子给出了盈亏平衡点:只有当 $t_{\text{exec}} \lesssim t_{\text{launch}}$ 时 $S_{\text{graph}} > 1$。 launch_model.py 的 [F2] 节用 $N=700$、$t_{\text{launch}}=3$ 微秒、$t_{\text{disp}}=0.5$ 微秒扫了一遍:$t_{\text{exec}}=0.5$ 微秒时 2.99 倍,1 微秒时 2.00 倍,2 微秒时 1.20 倍,到 3 微秒正好跌破 1(0.86 倍),40 微秒时 0.99 倍——大 kernel 上 CUDA Graph 是纯负担。 3.5 因果掩码:白算的那一半 因果注意力里,第 $i$ 个 query 只能看见前 $i$ 个 key。如果老老实实算满 $S \times S$ 的 score 矩阵,被掩掉的部分是纯浪费。精确的保留比例是: $$\text{kept} = \frac{S(S+1)/2}{S^2} = \frac{S+1}{2S}$$ $S=2048$ 时是 50.0%——一半的 score 元素算完就扔。逐元素地跳过在硬件上不现实,实际做法是按块跳:块边长 $B_r$ 时共有 $n_b = S/B_r$ 个 query 块,第 $i$ 个 query 块只算前 $i$ 个 key 块,算下来是 $n_b(n_b+1)/2$ 个块: $$\text{kept}_{\text{block}} = \frac{n_b(n_b+1)/2}{n_b^2} = \frac{n_b+1}{2 n_b} = \frac{1}{2} + \frac{1}{2 n_b}$$ 实测($S=2048$):块边长 64 时 51.6%,128 时 53.1%,256 时 56.2%。误差随 $n_b$ 按 $1/(2n_b)$ 衰减,所以块越小越接近理论上界——但块越小,片上数据被切得越碎,调度开销越大。这是块大小必须折中的原因,不是调参玄学。 这张图要看什么:左图是 score 矩阵相对输入激活的放大倍数随 $S$ 的变化,灰虚线是「若按 $S$ 线性增长」的参考——实际曲线比线性还陡($H \cdot S/d$ 是 $S$ 的一次式,但在对数轴上叠加了 $H/d$ 的常数放大,$S=8192$ 时已到 32 倍)。右图是块级跳过的实际工作量占比,绿虚线是理论上界 50.0%:块边长 256 时要多算 6.2 个百分点,块边长 64 时只多算 1.6 个——这就是「块越小越准、代价越碎」的量化版本。 04. 代码实现 三个脚本,全部纯 numpy,/usr/local/bin/python3 直接跑: fusion_lab.py:[A] 字节账本、[B] 带宽与工作集、[C] 三种融合形态、[D] 固定开销、[E] attention 尾部 launch_model.py:[F1] 时间线、[F2] 扫描、[F3] kernel 计数、[F4] 收益分解 make_figures.py:把上面两个脚本落盘的 json 画成 5 张图 4.1 三种融合形态 我把「融合」拆成三种可测量的形态,这是本篇的核心实验设计: K = 8 # 8 段链,每段 h <- A*h + B A_S, B_S = 1.01, 0.01 def _v1_unfused(x, buf, K=K): """V1 不融合:K 段,每段 2 次算子调用,每段结果都写回内存。""" np.multiply(x, A_S, out=buf) np.add(buf, B_S, out=buf) for _ in range(K - 1): np.multiply(buf, A_S, out=buf) np.add(buf, B_S, out=buf) return buf def _v3_closed(x, out, K=K): """V3 代数合并:K 段仿射合成一个仿射,只剩 2 次调用。""" Ak = A_S ** K Bk = B_S * (Ak - 1.0) / (A_S - 1.0) np.multiply(x, Ak, out=out) np.add(out, Bk, out=out) return out 这三种形态的物理含义不同,必须分清: V1 不融合:每个算子一个 kernel,中间结果每次都落回主存。 V2 分块融合(_v2_tiled):一块读进来,K 段都在这一块上算完,只写回一次。tile 就是片上缓冲。这是 GPU 融合 kernel 在 CPU 上能做到的最好近似——numpy 做不到寄存器级融合,每个 ufunc 调用仍然要把结果写回 tile。 V3 代数合并:把 $K$ 段仿射在数学上合成一段,中间结果根本不存在。它对应的是融合给编译器创造的二阶机会:串起来看得见全貌之后,可以化简。 实测结果(fusion_lab.py ALL 的 [C] 节): 张量元素数 V1 不融合 V3 代数合并 V1/V3 相对误差 512 7.75 μs 1.12 μs 6.89x 1.35e-07 4096 10.58 μs 1.46 μs 7.26x 1.97e-07 65536 84.63 μs 11.25 μs 7.52x 5.12e-07 262144 310.83 μs 41.62 μs 7.47x 5.80e-07 1048576 1240.04 μs 176.54 μs 7.02x 5.13e-07 4000000 5205.08 μs 792.83 μs 6.57x 6.16e-07 V1/V3 稳定在 6.57~7.52 倍,围着 $K=8$ 上下浮动——正如 3.2 节的推导,字节比就是上限。最后一列的相对误差约 5e-7,正好是 fp32 的舍入量级:代数合并改变了运算顺序,所以数值不会逐位相同,但量级完全在许可范围内。 同一节里 V2 的结果才是本篇最值得记住的一个数: tile(元素) 耗时 相对 V1 16384 5.835 ms 0.89x 65536 5.626 ms 0.93x 262144 5.290 ms 0.98x 1048576 5.361 ms 0.97x V2 一点都没变快,甚至略慢。 而 3.2 节的模型早就预测到了——代入 [B] 节测出的 $\rho = 1.124$: $$S = \frac{8}{1 + 7/1.124} = 1.11$$ 模型说 1.11 倍,实测 0.98 倍。同一量级,方向一致。为什么这台机器上 $\rho$ 这么小?因为单线程 numpy 逐元素循环的吞吐上限只有 91.5 GB/s,而主存带宽是 81.4 GB/s——两个数字离得太近,片上片下几乎没有差价可赚。 这张图要看什么:左图是两种形态的耗时随规模的变化,两条虚线是各自的发射地板——$V1$ 是 $2K \cdot a = 6.9$ μs,$V3$ 是 $2a = 0.86$ μs。两条实线在小规模处几乎是水平的(贴着各自的地板走),到大规模才分开,这就是「小规模赚发射、大规模赚字节」的直接证据。灰色菱形是 V2 在最大规模上的结果,它几乎落在 V1 那条线上——$V2$ 省了字节但没省调用,所以拿不到 V3 的收益。右图把这件事翻译成加速比:橙线(V1/V3)贴着绿色上限 $K=8$ 走,而 V2 的实测点落在灰虚线(模型预测 1.11x)附近,离 8 差着一个数量级。 4.2 把带宽层级测出来 $\rho$ 不是查来的,是测来的: for nbytes in [32 << 10, 128 << 10, 512 << 10, 2 << 20, 8 << 20, 32 << 20, 128 << 20]: n = nbytes // 4 x = rng.standard_normal(n).astype(np.float32) y = np.empty_like(x) bench(lambda: np.add(x, 1.0, out=y), reps=3) # 预热 t = bench(lambda: np.add(x, 1.0, out=y), reps=11) bw = 2 * nbytes / t / 1e9 # 1 读 1 写 实测:32 KiB 62.9 GB/s、128 KiB 79.6、512 KiB 91.5、2 MiB 95.3、8 MiB 94.7、32 MiB 79.5、128 MiB 83.3。取 ≤8 MiB 的中位数 91.5 GB/s 作 $B_{\text{on}}$,≥32 MiB 的中位数 81.4 GB/s 作 $B_{\text{off}}$,得 $\rho = 1.124$。 注意这两个数字不是硬件规格里的峰值带宽,而是「单线程 numpy 逐元素循环能跑出来的吞吐」。它们的差距不代表 L2 和主存的差距,只代表在这条代码路径上「留片上」值多少钱。这个诚实的限定很重要——换一台有真 GPU 的机器,$\rho$ 会是另一个数,$K/(1+(K-1)/\rho)$ 会给出完全不同的答案。 4.3 固定开销 reps = 20000 if n <= 4096 else 500 t0 = time.perf_counter() for _ in range(reps): np.add(a, 1.0, out=b) t = (time.perf_counter() - t0) / reps 实测:$n=1$ 时 0.417 μs,$n=8$ 时 0.419,$n=64$ 时 0.425,$n=512$ 时 0.477,$n=4096$ 时 0.675,$n=16384$ 时 1.333,$n=65536$ 时 5.966,$n=262144$ 时 22.388,$n=1048576$ 时 89.331。 前四个点几乎是一条水平线——从 1 个元素到 512 个元素,元素数涨了 512 倍,耗时只从 0.417 涨到 0.477。用 $n \le 16384$ 的六个点做最小二乘,得 $a = 0.430$ μs、$b_{\text{el}}$ 对应 55.47 μs/百万元素。 4.4 图:先看两张 这张图要看什么:三组 pattern 的字节账本,横轴是对数刻度。橙红是不融合、绿是融合后,条右边的倍数就是字节比。注意最上面那根——attention softmax 的 8192.0 MiB 比最下面那根逐元素链的 244.1 MiB 大了三十多倍,真正需要融合的不是那串小算子,而是注意力里那张被反复读写的 score 矩阵。 这张图要看什么:左图是单次调用耗时随张量规模的变化,红色虚线是固定开销地板,紫色竖线是「开销和数据各占一半」的临界规模(7746 个元素);右图是 8 段链里固定开销占总时间的比例。曲线在最左边几乎是水平的——那一段里,你增加 500 倍的数据量,耗时只涨 14%。 4.5 分解:两种收益各占多少 launch_model.py 的 [F4] 节把本机实测的加速拆成两部分: 张量元素数 V1 实测 V1 模型 V3 实测 V3 模型 发射贡献 512 7.75 μs 7.33 μs 1.12 μs 0.92 μs 94% 4096 10.58 μs 10.51 μs 1.46 μs 1.31 μs 65% 65536 84.63 μs 65.04 μs 11.25 μs 8.13 μs 11% 262144 310.83 μs 239.53 μs 41.62 μs 29.94 μs 3% 1048576 1240.04 μs 937.51 μs 176.54 μs 117.19 μs 1% 4000000 5205.08 μs 3556.95 μs 792.83 μs 444.62 μs 0% $n=512$(2.0 KiB)时,加速的 94% 来自少发射,只有 6% 来自少搬字节;$n=4000000$(15.3 MiB)时反过来,100% 来自少搬字节。 模型在大规模那一端明显低估(5205 实测 vs 3557 模型)——因为 $b_{\text{el}}$ 是用缓存内的点拟合的,外推到 15 MiB 后,边际成本实际上比拟合值高。这是线性模型的能力边界,写在这里免得读者拿它当精确预测器。 这张图要看什么:左图是时间线甘特图,上排是 eager。蓝条是 CPU 发射,绿条是 GPU 执行,斜纹是 GPU 空转——空转的宽度正好等于发射和执行的时间差。下排是 CUDA Graph,一次发射之后 GPU 一路不停。右图是加速比随单 kernel 执行时间的变化,红线是发射成本 3 μs:红线左边 Graph 赚,红线右边 Graph 亏。 05. 工业级实现对照 5.1 torch.compile / Inductor:融合是调度器的决策 PyTorch 的主入口在 torch/_inductor/compile_fx.py 的 compile_fx(当前实现在该文件第 3122 行),它把 FX 图接给 Inductor,后者负责切分、调度、生成 Triton kernel: https://github.com/pytorch/pytorch/blob/main/torch/_inductor/compile_fx.py 真正做融合决策的是调度器。torch/_inductor/scheduler.py 里: Scheduler.fuse_nodes(nodes)(约 7048 行)是融合主循环; can_fuse_vertical / can_fuse_horizontal / can_fuse_reduction_epilogue 决定两个节点能不能合; score_fusion_memory(node1, node2, count_bytes=...)(约 11278 行)给一次融合打分——它的第一个形参就叫 count_bytes,因为融合的分数就是省下的字节数。 融合决定返回 FusionResult(约 130 行),里面既可以是布尔值,也可以是一个待求值的 callable_fn,让代价模型并行算。 这就是本篇 3.2 节的账本在生产代码里的样子。 区别在于:编译器要在不能融合的时候正确地放弃。不能融合的情形包括跨块归纳(reduction 的中间结果必须先写完)、原地写与别名(写坏了输入后续还要用)、随机数算子(融合会改变随机流)、以及动态 shape 下无法预先确定 tile 大小。 5.2 一个必须在 05 节点名的差异:epilogue fusion 最小实现里的融合是「一串逐元素算子合成一个 kernel」。生产实现里最有价值的融合形态是 epilogue fusion:把接在 GEMM 后面的 bias、激活、缩放并进 GEMM 的 kernel 里,让 GEMM 的输出直接以最终形态写出,不经过显存。 用 3.2 节的账本算一下就明白为什么值钱:一个 $[M,N]$ 的 fp32 输出,不融合时要写 GEMM 输出($4MN$ 字节写)、bias 加法读+写($8MN$)、激活读+写($8MN$),合计 $20MN$;融合后 GEMM 一趟写 $4MN$,加上读输入的开销,量级降到 $1/5$。表达式没变一行,访存少了五分之四。 5.3 CUDA Graph 在 PyTorch 里的落地 torch/cuda/graphs.py 里的 CUDAGraph(约 289 行)暴露出和 CUDA 一一对应的四个动作:capture_begin / capture_end / instantiate / replay: https://github.com/pytorch/pytorch/blob/main/torch/cuda/graphs.py 三个工程细节值得单独说: graph_pool_handle()(约 96 行)必须存在。graph 在录制时把显存地址写死在节点里,回放时不会重新分配。所以你必须让图里所有张量都来自一个固定的内存池。这不是 API 的怪癖,是「录制-回放」这个机制的必然代价。 上游警告 instantiate 要在第一次 replay 之前显式调用,否则第一次回放的延迟会变高(图要在那时才编译)。这个问题在实时推理里就是一次 P99 抖动。 torch/_inductor/cudagraph_trees.py 里是 CUDAGraphNode(约 982 行)和 TreeManagerContainer(约 220 行),还有一个 CUDAWarmupNode。为什么是「树」而不是「一张图」:训练时反向图的形状依赖前向的实际形状,一批数据一个形状;而且每次回放都会产生新的输出张量,需要知道哪些显存可以复用。Inductor 用 mode="reduce-overhead" 打开这套机制。 https://github.com/pytorch/pytorch/blob/main/torch/_inductor/cudagraph_trees.py 5.4 差异从哪来:为什么生产实现不能只有最小实现 差异 来源 要判断「能不能融合」,甚至要算清代价 融合不是永远划算,见第六节 6.1 与 6.3 要处理动态 shape tile 大小不能预知,必须切分或退化成不融合 要维护固定的显存池 CUDA Graph 把地址写死了 要为每个形状/profile 各录一份图 图是静态的;变长输入需要分段或重录 要处理随机流、原地写、别名 融合会改变这些语义 06. 代价与边界 6.1 收益的上限是字节比,不是魔法 3.2 节的 $S = K/(1+(K-1)/\rho)$ 是硬约束。在写本文这台机器上 $\rho = 1.124$,理论上限只有 1.11 倍,实测 V2 是 0.98 倍——收益在哪台机器上有、有多大,完全由 $\rho$ 决定。所以当有人告诉你「融合能快 6 倍」,第一个该问的问题不是「怎么融合」,而是「他的 $\rho$ 是多少、$K$ 是多少」。 顺带说,V3(代数合并)之所以能跑到 6.57~7.52 倍而不受 $\rho$ 限制,是因为它根本没产生中间结果——它走的是 $S \to K$ 那条极限路径。但代数合并只在可化简的模式上成立(本链是仿射复合,可以闭式求解),一般算子是合不掉的。它展示的是融合的第二层价值:让编译器看见全貌,从而有机会化简。 6.2 融合会改变数值 实测相对误差约 5e-7(fp32 舍入量级)。这不是 bug,是代数重排的必然结果。任何声称「融合不改变数值」的说法都需要限定条件。在需要严格复现的训练里,这会让「同一份代码换个后端跑出不同 loss」——通常无害,但你必须知道它从哪来。 6.3 融合省不了算错的东西 这是本篇最反直觉的一条,也是附录 [E] 节专门测的。softmax 五步在 $[8,1024,1024]$ 的 score 矩阵上:不融合 18.804 ms,按行分块融合 23.870 ms,加速比 0.79x——融合反而慢了 21%。 算一下就知道为什么。五步共读写 8 份 score 矩阵 = 256 MB,按 81.4 GB/s 的访存下界是 3.298 ms。实测 18.804 ms,是访存下界的 5.70 倍。也就是说这个 kernel 的时间几乎全在 exp 的计算上,不在访存上。 分块省了字节,但一个 exp 都没少,反而因为切碎了向量化还赔了一点。 所以:融合只对「本来就卡在访存或发射上」的算子有效。 一个计算受限的算子,融合不动它。这也是 FlashAttention 要在融合之外再做「不实例化 S×S 矩阵」的原因——它同时省了访存和一部分计算。 6.4 CUDA Graph 的硬约束 静态 shape:图录下来的形状是死的。变长输入要么分段(prefill 一段、decode 一段各录一份),要么放弃。 静态地址:所有张量必须来自固定内存池,中间不能有新的 cudaMalloc。 不能有 host 同步和 D2H 拷贝:录制期间不允许任何把控制权交回 CPU 的操作。 首次录制和实例化有成本,且出错时栈是「图里的第 N 个节点」,比直接调试难得多。 大 kernel 上纯亏:3.4 节的模型里,$t_{\text{exec}} = 20$ μs 时加速比 0.98 倍。你付出了录制、显存和调试的成本,换来 2% 的倒退。 一句话:CUDA Graph 是给「小 kernel 洪水」用的药,不是通用加速开关。 一个理想的判断顺序是:先用 profiler 确认 GPU 有空转;再确认单个 kernel 确实小于发射成本;然后才考虑上 Graph。而如果那些 kernel 本来就该被融合掉,融合是更根本的解法——融合让 kernel 数从 1093 降到 355,Graph 只是把这 1093 次提交合并成 1 次;两者正交,可以叠加。 6.5 证据的范围 本机实测:全部来自附录三个脚本在这台 Apple M1 Pro 上的真实运行输出([A]~[F4])。这台机器没有 NVIDIA GPU、没有 torch,所以没有任何一个 CUDA 数字是实测的。 公开资料(非实测):kernel launch 的 1~5 μs 量级、一步 decode 的 300~1300 个 kernel、HBM 与片上带宽的量级差。出处见 launch_model.py 文件头。 模型推演:launch_model.py 的 [F1]/[F2]/[F3]/[F4] 是在上述输入上做的算术,脚本可跑、输入可改。请把这些当作「给定这些输入会得出什么」,而不是「GPU 上就是这么快」。 不确定的部分:[F3] 的 1093 是按架构手数的,真实值取决于编译器怎么 lowering,本文没有在真实 GPU 上验证过。Inductor 的融合策略也在持续变化,5.1 节引用的行号与函数名以 2026-10 时的 main 分支为准。 07. 经典论文脉络 四篇,按「融合是怎么从手工技巧变成系统能力,又怎么被反过来审视」串起来。 TVM: An Automated End-to-End Optimizing Compiler for Deep Learning(arXiv:1802.04799,2018)。第一次把「算符融合」从工程师的手工活变成编译器的自动决策:给定计算图,搜索融合方案和调度模板。本篇 3.2 节那个「融合省多少字节」的账,在这篇里是被当作搜索目标函数的一部分来算的。 The Deep Learning Compiler: A Comprehensive Survey(arXiv:2002.03794,2020,本文知识树的锚点)。给出一张完整坐标系:图级优化(融合、常量折叠、CSE)、算子级优化、内存分配、后端代码生成。它把融合明确归到「图级优化」——融合对单个算子的数学一无所知,它只改写算子之间的边界。这句话是理解本篇全部内容的前提。 FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness(arXiv:2205.14135,2022)。把融合推到极致:不只把 softmax 融进注意力,还根本不实例化 $S \times S$ 的 score 矩阵,用在线 softmax 一边遍历一边归一化。这正是本篇 [A] 节那 1.000 GiB 中间产物的解法,而且它额外做到了 6.3 节说的那件事——省的不只是访存,还有被掩码白算的一半。 Operator Fusion in XLA: Analysis and Evaluation(arXiv:2301.13062,2023)。反过来审视这件事:作者去读 XLA 的融合 pass 源码,实测在 Cartpole 上不同融合策略的效果,最好的实现拿到 10.56 倍。这篇的价值在于它把「融合是好事」这个默认假设变成了可度量的对象——和本篇的 V2 实测 0.98 倍是同一个姿势:先问一句「到底快了多少」,再决定要不要相信。 至于 CUDA Graph,它不是论文,是 CUDA 10 起提供的一项运行时机制(录制-实例化-回放),说明见 CUDA C++ Programming Guide 的 CUDA Graphs 一节。把它和融合并列讨论,是因为它经常被当成融合的替代品——而它其实解决另一个问题。 08. 常见误解 「融合越大越好」。 融得越狠,寄存器压力和共享内存占用越大,occupancy 掉下来,反而更慢。Inductor 的 can_fuse_* 那一组函数之所以是「判断」而不是「尽量合」,就是因为融合有代价。而本文 [E] 节的实测给了一个更直接的极端例子:分块融合的 softmax 比不融合慢 21%。 「CUDA Graph 是通用加速开关」。 它一个字节都不省,也不减少 kernel 数。3.4 节的公式说得很清楚:$t_{\text{exec}} \ge t_{\text{launch}}$ 时它只会让你付 $t_{\text{disp}}$ 的额外成本。先测 GPU 有没有空转,再决定要不要上。 「融合不改变数值」。 代数重排会改变舍入顺序。本文 V3 实测相对误差约 5e-7。无害,但要知道它存在。 「把中间结果留在片上就一定快」。 收益取决于 $\rho$。本文这台机器 $\rho = 1.124$,V2 实测 0.98 倍——留片上的动作做对了,收益是零。反面同样成立:GPU 上 $\rho$ 大得多,同一个动作就值几倍。 「融合省的是计算量」。 它不改变 FLOPs,只改变访存和发射。softmax 的 exp 一个都没少([E] 节:实测是访存下界的 5.70 倍,说明瓶颈在计算)。 09. 动手验证 都是纯 numpy,不需要 GPU。 实验一:拿到你自己机器的地板。 跑 python fusion_lab.py D,读出 $a$ 和 $n^{*}$。在 M1 Pro 上预期 $a \approx 0.43$ μs、$n^{*} \approx 7746$ 元素(30.3 KiB)。换台机器这两个数会变,但「小张量上耗时和元素数无关」这段平台一定会出现。 实验二:把 K 调大。 把 fusion_lab.py 里的 K_CHAIN 从 8 改成 16,重跑 ALL。预期两件事:V1/V3 的加速比向 16 靠拢(而不是翻倍到 16 以上——字节比就是上限),以及 V1 的发射地板从 $2\times8\times0.43 \approx 6.9$ μs 涨到约 13.8 μs,曲线左端整体抬高。 实验三:看 $\rho$ 怎么被改坏。 把 [B] 节的 np.add(x, 1.0, out=y) 换成跨大步长的切片访问(例如每隔 32 个元素取一个),重测。预期 $B_{\text{on}}$ 明显下降、$\rho$ 趋近 1,随后 V2 的收益会进一步向 1 靠拢。这能让你亲眼看到「$\rho$ 小 ⇒ 融合白干」这条因果链。 实验四:自己数一遍 kernel。 跑 python launch_model.py,在 [F3] 的 per_layer 列表里按你自己熟悉的模型改,看总数落在哪。预期 Llama-3-8B 的 1093 落在公开资料给的 300~1300 区间内。 10. 延伸阅读 performance_profiling|性能建模与 Profiling:本篇 3.1 节的 roofline 判据来自那里,$\rho$ 和算术强度这两个量在那里第一次被建立起来。 attention_basics|自注意力机制:理解 3.5 节因果掩码和 [A] 节 score 矩阵形状的前提。 flash_attention|FlashAttention:本篇 [A] 节那 1.000 GiB 中间产物的正式解法,融合走到极致的样子。 mixed_precision|混合精度:数据类型决定 $b$(每元素字节数),直接进 3.2 节的账本。 kv_cache|KV Cache 与自回归视频生成:decode 场景为什么 kernel 又多又小,那篇给的 cache 账本是本篇 6.4 节「分段 capture」动机的来源。 附录:完整代码 09 节用到的脚本全文如下(launch_model.py、fusion_lab.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 launch_model.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """ launch_model.py —— 发射开销、CUDA Graph,以及它们和融合的分工 融合省的是字节,CUDA Graph 省的是 CPU 侧的发射次数,一个字节都不省。 这个脚本把两件事分开算: [F1] 离散事件时间线:eager 逐个发射 vs CUDA Graph 一次发射 [F2] 扫描:单个 kernel 执行多久时,CUDA Graph 才赚? [F3] 按 Llama-3-8B 架构手数一步 decode 有多少个 kernel [F4] 分解:本机实测的 K 段链加速里,多少来自「少发射」,多少来自「少搬字节」 公开量级的出处(都不是本机实测,本机没有 NVIDIA GPU): * kernel launch 的 CPU 侧成本:NVIDIA 开发者论坛 njuffa 给的是"空 kernel 约 5 us";CUDA Handbook 实测 NULL launch 约 4.9 us(老机器)、RTX 3060 约 1.2 us;A100 上常见引用是 2~3 us。所以本文取 1~5 us 这个区间, 并且强调**比值比绝对值重要**。 * 一个 decode step 的 kernel 数:公开资料给的量级是 300~1300。 [F3] 按架构手数出来的数落在这个区间里,可互为印证。 运行: python launch_model.py """ from __future__ import annotations import json import os import numpy as np HERE = os.path.dirname(os.path.abspath(__file__)) # ══════════════════════════════════════════════════════════════ # [F1] 离散事件时间线 # ══════════════════════════════════════════════════════════════ def simulate_eager(n_kernel, t_launch, t_exec): """eager:CPU 逐个 cudaLaunchKernel,GPU 逐个执行,单流异步。 CPU 发第 i 个要花 t_launch,发完就不用管了,可以继续发下一个; GPU 要等「前一个跑完」且「这个已经被发出去」才能开始。 两个条件里慢的那个决定 GPU 什么时候开工,差出来的就是 GPU 空转。 """ gpu_free = 0.0 idle = 0.0 for i in range(n_kernel): cpu_done = (i + 1) * t_launch # CPU 发完第 i 个的时刻 start = max(gpu_free, cpu_done) # GPU 能开工的时刻 idle += start - gpu_free gpu_free = start + t_exec return {"total": gpu_free, "idle": idle, "idle_frac": idle / gpu_free if gpu_free > 0 else 0.0} def simulate_graph(n_kernel, t_launch, t_exec, t_dispatch): """CUDA Graph:一次 cudaGraphLaunch 提交整张图。 GPU 仍然要逐个节点过一遍前端,所以每个 kernel 还留一个 t_dispatch 的 GPU 侧派发成本——只是不再需要 CPU 每次都插手。 """ total = t_launch + n_kernel * (t_exec + t_dispatch) return {"total": total, "idle": t_launch, "idle_frac": t_launch / total if total > 0 else 0.0} def section_F1(): print("\n" + "=" * 72) print("[F1] 时间线:eager 逐个发射 vs CUDA Graph 一次发射") print("=" * 72) N, tl, te, td = 700, 3.0, 2.0, 0.5 e = simulate_eager(N, tl, te) g = simulate_graph(N, tl, te, td) print(f"\n 设定:N={N} 个 kernel,发射 {tl} us,执行 {te} us," f"graph 内派发 {td} us") print(f" eager 总时长 {e['total']:8.1f} us,GPU 空转 {e['idle']:7.1f} us" f"({e['idle_frac']*100:.0f}%)") print(f" CUDA Graph 总时长 {g['total']:8.1f} us,GPU 空转 {g['idle']:7.1f} us" f"({g['idle_frac']*100:.0f}%)") print(f" 加速比 = {e['total']/g['total']:.2f}x") # 换一个大 kernel 场景:执行时间远大于发射 e2 = simulate_eager(N, tl, 20.0) g2 = simulate_graph(N, tl, 20.0, td) print(f"\n 换成大 kernel(执行 20 us):") print(f" eager {e2['total']:8.1f} us,空转 {e2['idle_frac']*100:.0f}%") print(f" CUDA Graph {g2['total']:8.1f} us,空转 {g2['idle_frac']*100:.0f}%") print(f" 加速比 = {e2['total']/g2['total']:.2f}x -> 几乎没用") print(f"\n 结论:CUDA Graph 只在「单个 kernel 的执行时间接近或小于发射时间」") print(f" 时才赚钱。大 kernel 上它是纯负担(录制、显存、调试成本)。") return {"N": N, "t_launch": tl, "t_exec": te, "t_dispatch": td, "eager": e, "graph": g, "speedup": e["total"] / g["total"], "big": {"eager": e2, "graph": g2, "speedup": e2["total"] / g2["total"]}} # ══════════════════════════════════════════════════════════════ # [F2] 扫描:什么时候值得上 CUDA Graph # ══════════════════════════════════════════════════════════════ def section_F2(): print("\n" + "=" * 72) print("[F2] 扫描:单 kernel 执行时间 t_exec 对加速比的影响") print("=" * 72) N, tl, td = 700, 3.0, 0.5 print(f"\n N={N}, 发射 {tl} us, graph 内派发 {td} us") print(f"\n {'t_exec(us)':>11} {'eager(us)':>11} {'graph(us)':>11} " f"{'加速':>7} {'eager空转':>9}") rows = [] for te in [0.5, 1.0, 2.0, 3.0, 5.0, 8.0, 12.0, 20.0, 40.0]: e = simulate_eager(N, tl, te) g = simulate_graph(N, tl, te, td) rows.append({"t_exec": te, "eager_us": e["total"], "graph_us": g["total"], "speedup": e["total"] / g["total"], "idle_frac": e["idle_frac"]}) print(f" {te:>11.1f} {e['total']:>11.1f} {g['total']:>11.1f} " f"{e['total']/g['total']:6.2f}x {e['idle_frac']*100:8.0f}%") # 盈亏平衡点:eager 总时长 == graph 总时长 print(f"\n 盈亏平衡:eager 靠 CPU 逐个发射,graph 每个 kernel 多付 {td} us 派发。") print(f" 当 t_exec > t_launch 时 eager 已经不让 GPU 空转,graph 只是白付 {td}。" f"") print(f" 本例盈亏点在 t_exec ≈ t_launch = {tl} us 附近。") return {"N": N, "t_launch": tl, "t_dispatch": td, "rows": rows} # ══════════════════════════════════════════════════════════════ # [F3] 一步 decode 有多少个 kernel(按架构手数) # ══════════════════════════════════════════════════════════════ # Llama-3-8B 的公开配置 LLAMA3_8B = dict(n_layer=32, d_model=4096, n_head=32, n_kv_head=8, head_dim=128, d_ffn=14336, vocab=128256) def section_F3(): print("\n" + "=" * 72) print("[F3] 一步 decode 有多少个 kernel(按 Llama-3-8B 架构手数)") print("=" * 72) print("\n 这是**按架构推导**,不是实测。真实数字取决于编译器怎么 lowering,") print(" 融合后能少一半以上。公开资料给的量级是 300~1300。") cfg = LLAMA3_8B # 每层:括号里是不融合时这个模块会拆成几个 kernel per_layer = [ ("输入 RMSNorm", 4), # 平方 / 归约 / 乘 / 缩放 ("q_proj", 1), ("k_proj", 1), ("v_proj", 1), ("RoPE on q", 3), # cos/sin 表取 + 旋转 + 拼接 ("RoPE on k", 3), ("KV cache 写入", 2), # k、v 各一次 scatter ("QK^T", 1), ("scale + mask", 2), ("softmax", 3), # max / sub+exp / sum+div ("@V", 1), ("o_proj", 1), ("残差加", 1), ("后注意力 RMSNorm", 4), ("gate_proj", 1), ("up_proj", 1), ("SiLU", 1), ("逐元素乘", 1), ("down_proj", 1), ("残差加", 1), ] total_per_layer = sum(k for _, k in per_layer) n_layer = cfg["n_layer"] total = total_per_layer * n_layer + 5 # + 输出层 norm / lm_head / 采样等 print(f"\n 每层 {total_per_layer} 个 kernel:") for name, k in per_layer: print(f" {name:<22} {k}") print(f"\n x {n_layer} 层 = {total_per_layer*n_layer},加输出头约 5 个") print(f" 合计 ≈ {total} 个 kernel / token") print(f" (公开资料给的区间 300~1300,这个数落在里面)") # 融合后:把每层里能合的合掉 fused_per_layer = [ ("RMSNorm 融合", 1), ("QKV 一次 GEMM", 1), ("RoPE 融合", 1), ("KV cache 写入", 1), ("注意力融合(Flash)", 1), ("o_proj", 1), ("残差 + RMSNorm 融合", 1), ("gate/up 一次 GEMM", 1), ("SiLU + 乘 融合", 1), ("down_proj", 1), ("残差加", 1), ] fused_pl = sum(k for _, k in fused_per_layer) fused_total = fused_pl * n_layer + 3 print(f"\n 融合后每层 {fused_pl} 个 -> 合计 ≈ {fused_total} 个/token") print(f" kernel 数减少 {total/fused_total:.1f}x") print(f"\n 注意:融合减少的是 kernel 数,CUDA Graph 不减少 kernel 数,") print(f" 它只是把 {total} 次 CPU 发射合并成 1 次。两者正交,可以同时用。") return {"cfg": cfg, "per_layer": per_layer, "total_per_layer": total_per_layer, "total": total, "fused_per_layer": fused_per_layer, "fused_total": fused_total, "reduce": total / fused_total} # ══════════════════════════════════════════════════════════════ # [F4] 分解:本机实测的加速里,发射和字节各占多少 # ══════════════════════════════════════════════════════════════ def section_F4(fus): print("\n" + "=" * 72) print("[F4] 分解:本机 K 段链的加速里,多少来自少发射、多少来自少搬字节") print("=" * 72) D = fus["D"] C = fus["C"] a = D["a_us"] # 每次调用的固定开销(实测) b = D["b_us_per_elem"] # 每元素的边际成本(实测) K = C["K"] print(f"\n 实测固定开销 a = {a:.3f} us/次,边际 b = {b*1e6:.2f} us/百万元素") print(f"\n {'n':>9} {'V1 实测':>10} {'V1 模型':>10} {'V3 实测':>10} " f"{'V3 模型':>10} {'发射贡献':>9}") rows = [] for c in C["curves"]: n = c["n"] # V1: 2K 次调用,每次搬 n 个元素;V3: 2 次调用 t1_model = 2 * K * (a + b * n) t3_model = 2 * (a + b * n) # 反事实:只把调用次数从 2K 降到 2,数据量不变(纯发射收益) t_launch_only = 2 * K * a + 2 * K * b * n - (2 * a + 2 * K * b * n) total_gain = t1_model - t3_model frac = t_launch_only / total_gain if total_gain > 0 else 1.0 rows.append({"n": n, "t1_model": t1_model, "t3_model": t3_model, "launch_frac": frac}) print(f" {n:>9} {c['t1_us']:9.2f} us {t1_model:9.2f} us " f"{c['t3_us']:9.2f} us {t3_model:9.2f} us {frac*100:8.0f}%") small = rows[0] big = rows[-1] print(f"\n n={small['n']}({small['n']*4/1024:.1f} KiB):" f"模型说加速的 {small['launch_frac']*100:.0f}% 来自少发射," f"只有 {(1-small['launch_frac'])*100:.0f}% 来自少搬字节。") print(f" n={big['n']}({big['n']*4/1024/1024:.1f} MiB):" f"反过来,{big['launch_frac']*100:.0f}% 来自少发射," f"{(1-big['launch_frac'])*100:.0f}% 来自少搬字节。") print(f"\n 这就是本文最想说的一句话:") print(f" **小张量上,融合赚的几乎全是发射次数;大张量上,赚的才是字节。**") print(f" CUDA Graph 只做前一件,而且做得更彻底(一次发射整张图)。") return {"a_us": a, "b_us_per_elem": b, "rows": rows} # ══════════════════════════════════════════════════════════════ def main(): res = {} res["F1"] = section_F1() res["F2"] = section_F2() res["F3"] = section_F3() p = os.path.join(HERE, "_fusion_results.json") if os.path.exists(p): with open(p, encoding="utf-8") as f: fus = json.load(f) if "C" in fus and "D" in fus: res["F4"] = section_F4(fus) else: print("\n[F4] 跳过:先跑 fusion_lab.py ALL 生成 _fusion_results.json") p = os.path.join(HERE, "_launch_results.json") with open(p, "w", encoding="utf-8") as f: json.dump(res, f, ensure_ascii=False, indent=2) print(f"\n结果已写入 {p}") if __name__ == "__main__": main() fusion_lab.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """ fusion_lab.py —— 算子融合的「字节账本 + 实测校验」 写这篇文章的机器是一台 Apple M1 Pro,**没有 NVIDIA GPU,也没有 torch**, 所以这里不假装能测 CUDA kernel。能测的是两件在这台机器上真实存在的事: 1. 内存有层级:数据留在片上和落回主存,代价差一个可测的倍数。 融合做的事就是「少落回几次」,这个倍数能测,收益也能算。 2. 每次调用算子都有一个与数据规模无关的固定成本(函数派发、缓冲检查、 循环建立)。它能测出来,是 GPU 上 kernel launch 开销在 CPU 上的同构物。 GPU 上的具体数字(HBM 带宽、launch 微秒数)本文一律引公开资料并明确标注, 不拿这台机器外推。反过来,凡是标「实测」的数字,都出自本文件。 五组实验: [A] 字节账本 —— 几类常见 pattern 融合前后各搬多少字节(纯算术,精确) [B] 带宽–工作集 —— 实测片上/片外带宽比 rho [C] 融合三形态 —— 不融合 / 分块(留片上) / 代数合并,扫规模看各自值多少 [D] 固定开销 —— 拟合出每次调用的固定成本,算它占总耗时的比例 [E] attention尾 —— score 矩阵账本、本机 softmax 卡在哪 运行: python fusion_lab.py ALL """ from __future__ import annotations import json import os import time import numpy as np HERE = os.path.dirname(os.path.abspath(__file__)) RNG_SEED = 20261002 MIB = float(1024 ** 2) GIB = float(1024 ** 3) # 逐元素链的形态:每段 h <- A*h + B,共 K 段 K_CHAIN = 8 A_S, B_S = 1.01, 0.01 def bench(fn, reps=7): """取多次里的最小值:最小值受系统抖动影响最小。""" ts = [] for _ in range(reps): t0 = time.perf_counter() fn() ts.append(time.perf_counter() - t0) return min(ts) def passes_to_bytes(n_elem, dtype_bytes, n_passes): """1 个 pass = 把整个张量读一遍再写一遍 = 2 * 字节数。""" return 2.0 * n_elem * dtype_bytes * n_passes # ══════════════════════════════════════════════════════════════ # [A] 字节账本:融合前后各搬多少字节 # ══════════════════════════════════════════════════════════════ def section_A(): print("\n" + "=" * 72) print("[A] 字节账本:融合前后各搬多少字节(纯算术,精确)") print("=" * 72) out = {} n = 4_000_000 # 一个中等激活张量,fp32 b = 4 # ── A1. K 段逐元素链 ────────────────────────────────── K = K_CHAIN unf = passes_to_bytes(n, b, K) fus = passes_to_bytes(n, b, 1) print(f"\n A1 {K} 段逐元素链({n/1e6:.0f}M 元素,fp32)") print(f" 不融合 {K} 个 kernel : {unf/MIB:8.1f} MiB") print(f" 融合成 1 个 kernel : {fus/MIB:8.1f} MiB -> 省 {unf/fus:.0f}x") out["chain"] = {"K": K, "unfused_MiB": unf / MIB, "fused_MiB": fus / MIB, "ratio": unf / fus} # ── A2. RMSNorm ─────────────────────────────────────── # 不融合:x*x(读写) / reduce mean(读) / x*r(读写) / *g(读写) unf = passes_to_bytes(n, b, 4) fus = passes_to_bytes(n, b, 1) print(f"\n A2 RMSNorm({n/1e6:.0f}M 元素)") print(f" 不融合 4 遍 : {unf/MIB:8.1f} MiB") print(f" 融合 1 遍 : {fus/MIB:8.1f} MiB -> 省 {unf/fus:.0f}x") out["rmsnorm"] = {"unfused_MiB": unf / MIB, "fused_MiB": fus / MIB, "ratio": unf / fus} # ── A3. attention 的 score 矩阵 ─────────────────────── Bs, H, S, D = 8, 16, 2048, 4096 dt = 2 # bf16 act_bytes = Bs * S * D * dt score_bytes = Bs * H * S * S * dt half = 1 + 2 + 2 + 1 + 2 # softmax 五步各读写几份 score unf = half * score_bytes fus = 2 * score_bytes print(f"\n A3 attention score(B={Bs}, H={H}, S={S}, d={D}, bf16)") print(f" 输入激活 [B,S,d] : {act_bytes/MIB:8.1f} MiB") print(f" score 矩阵 [B,H,S,S] : {score_bytes/GIB:8.3f} GiB" f" = 输入的 {score_bytes/act_bytes:.0f}x") print(f" softmax 不融合 {half} 份读写 : {unf/GIB:8.3f} GiB") print(f" softmax 融合 : {fus/GIB:8.3f} GiB -> 省 {unf/fus:.1f}x") print(f" 额外物化一份 mask : +{score_bytes/GIB:.3f} GiB(融合后消失)") print(f"\n 放大倍数 = H*S/d = {H}*{S}/{D} = {H*S/D:.0f}x") print(f" S 翻一倍 -> 放大倍数翻一倍(score 按 S 的平方长,激活按 S 长)") for s in (1024, 2048, 4096, 8192): print(f" S={s:>5}: score/激活 = {H*s/D:5.1f}x") out["softmax"] = {"B": Bs, "H": H, "S": S, "d": D, "act_MiB": act_bytes / MIB, "score_GiB": score_bytes / GIB, "blowup": score_bytes / act_bytes, "unfused_GiB": unf / GIB, "fused_GiB": fus / GIB, "ratio": unf / fus, "mask_GiB": score_bytes / GIB, "blowup_curve": [{"S": s, "x": H * s / D} for s in (1024, 2048, 4096, 8192)]} # ── A4. 因果掩码:块级跳过能省多少 ──────────────────── total = S * S kept_exact = S * (S + 1) // 2 rows = [] print(f"\n A4 因果掩码 S={S}:块级跳过 vs 精确跳过") for br in (64, 128, 256): nb = S // br blocks_done = nb * (nb + 1) // 2 # query 块 i 只算 j<=i 的 key 块 rows.append({"br": br, "nb": nb, "frac": blocks_done / (nb * nb)}) print(f" 块边长 {br:>4}: 算 {blocks_done:>4}/{nb*nb:<4} 块" f" = {blocks_done/(nb*nb)*100:5.1f}% 工作量") print(f" 理论上界(逐元素精确): {kept_exact/total*100:5.1f}%") print(f" 完全不跳 : {100.0:5.1f}% -> 白算 " f"{100-kept_exact/total*100:.1f}%") out["causal"] = {"S": S, "exact_frac": kept_exact / total, "blocks": rows} return out # ══════════════════════════════════════════════════════════════ # [B] 带宽 vs 工作集:实测片上/片外带宽比 # ══════════════════════════════════════════════════════════════ def section_B(): print("\n" + "=" * 72) print("[B] 带宽 vs 工作集(np.add(x, 1, out=y):1 读 1 写)") print("=" * 72) rng = np.random.default_rng(RNG_SEED) print(f"\n {'工作集':>12} {'耗时':>10} {'带宽':>10} 归属") rows = [] for nbytes in [32 << 10, 128 << 10, 512 << 10, 2 << 20, 8 << 20, 32 << 20, 128 << 20]: n = nbytes // 4 x = rng.standard_normal(n).astype(np.float32) y = np.empty_like(x) bench(lambda: np.add(x, 1.0, out=y), reps=3) t = bench(lambda: np.add(x, 1.0, out=y), reps=11) bw = 2 * nbytes / t / 1e9 tag = "onchip" if nbytes <= (8 << 20) else "dram" rows.append({"KiB": nbytes >> 10, "ms": t * 1e3, "GBs": bw, "tag": tag}) print(f" {nbytes>>10:>10} KiB {t*1e3:9.3f} ms {bw:9.1f} GB/s {tag}") del x, y b_cache = float(np.median([r["GBs"] for r in rows if r["tag"] == "onchip"])) b_dram = float(np.median([r["GBs"] for r in rows if r["tag"] == "dram"])) rho = b_cache / b_dram print(f"\n 片上带宽中位数 B_cache = {b_cache:.1f} GB/s") print(f" 主存带宽中位数 B_dram = {b_dram:.1f} GB/s") print(f" 比值 rho = {rho:.3f} <- 决定融合能兑现多少的关键参数") print(f"\n 注意:这里的 B_cache 不是缓存的峰值带宽,而是单线程 numpy 逐元素") print(f" 循环的吞吐上限。它只比主存快 {rho:.2f} 倍,所以在这台机器上") print(f" 「把中间结果留在片上」这件事本身几乎不值钱(见 [C] 的 V2)。") return {"rows": rows, "B_cache_GBs": b_cache, "B_dram_GBs": b_dram, "rho": rho} # ══════════════════════════════════════════════════════════════ # [C] 融合三形态 × 规模扫描 # ══════════════════════════════════════════════════════════════ def _v1_unfused(x, buf, K=K_CHAIN): """V1 不融合:K 段,每段 2 次算子调用,每段结果都写回内存。""" np.multiply(x, A_S, out=buf) np.add(buf, B_S, out=buf) for _ in range(K - 1): np.multiply(buf, A_S, out=buf) np.add(buf, B_S, out=buf) return buf def _v2_tiled(x, out, tile, K=K_CHAIN): """V2 分块融合:一块读进来,K 段都在片上算完,只写回一次。 这是 GPU 融合 kernel 的 CPU 近似:tile 就是寄存器/共享内存那块片上缓冲。 numpy 做不到寄存器级融合(每个 ufunc 仍要写回 tile),所以 V2 是本机能 做到的最好近似。 """ t = np.empty(tile, dtype=np.float32) for s in range(0, x.size, tile): m = min(tile, x.size - s) v = t[:m] np.multiply(x[s:s + m], A_S, out=v) np.add(v, B_S, out=v) for _ in range(K - 1): np.multiply(v, A_S, out=v) np.add(v, B_S, out=v) out[s:s + m] = v return out def _v3_closed(x, out, K=K_CHAIN): """V3 代数合并:K 段仿射合成一个仿射,只剩 2 次调用。 h_K = A^K * h_0 + B*(A^K - 1)/(A - 1) 中间结果根本不存在,连「留在片上」都不需要。 """ Ak = A_S ** K Bk = B_S * (Ak - 1.0) / (A_S - 1.0) np.multiply(x, Ak, out=out) np.add(out, Bk, out=out) return out def section_C(rho): print("\n" + "=" * 72) print(f"[C] 融合三形态 × 规模扫描(K={K_CHAIN} 段 h <- {A_S}*h + {B_S})") print("=" * 72) rng = np.random.default_rng(RNG_SEED) sizes = [512, 4096, 65536, 262144, 1 << 20, 4_000_000] curves = [] print(f"\n {'n':>9} {'V1 不融合':>11} {'V3 代数合并':>11} " f"{'V1/V3':>7} {'相对误差':>10}") for n in sizes: x = rng.standard_normal(n).astype(np.float32) buf, out = np.empty_like(x), np.empty_like(x) reps = 5000 if n <= 65536 else 11 bench(lambda: _v1_unfused(x, buf), reps=3) t1 = bench(lambda: _v1_unfused(x, buf), reps=reps) bench(lambda: _v3_closed(x, out), reps=3) t3 = bench(lambda: _v3_closed(x, out), reps=reps) rel = float(np.max(np.abs(buf - out)) / np.max(np.abs(buf))) curves.append({"n": n, "t1_us": t1 * 1e6, "t3_us": t3 * 1e6, "speedup": t1 / t3, "rel_err": rel}) print(f" {n:>9} {t1*1e6:9.2f} us {t3*1e6:9.2f} us " f"{t1/t3:6.2f}x {rel:10.2e}") # V2 只在最大规模上跑:它慢,且只有这里才谈得上「主存 vs 片上」 print(f"\n V2 分块融合(n=4,000,000,扫 tile):") x = rng.standard_normal(4_000_000).astype(np.float32) buf, out = np.empty_like(x), np.empty_like(x) bench(lambda: _v1_unfused(x, buf), reps=3) t1_big = bench(lambda: _v1_unfused(x, buf), reps=11) print(f" V1 不融合基准: {t1_big*1e3:.3f} ms") v2 = [] for tile in [16384, 65536, 262144, 1 << 20]: bench(lambda: _v2_tiled(x, out, tile), reps=3) t = bench(lambda: _v2_tiled(x, out, tile), reps=11) v2.append({"tile": tile, "ms": t * 1e3, "speedup": t1_big / t}) print(f" tile={tile:>8}: {t*1e3:8.3f} ms {t1_big/t:5.2f}x") best2 = max(v2, key=lambda r: r["speedup"]) K = K_CHAIN pred = K / (1 + (K - 1) / rho) print(f"\n V2 最佳 {best2['speedup']:.2f}x(tile={best2['tile']})") print(f" 模型预测 K/(1+(K-1)/rho) = {K}/(1+{K-1}/{rho:.3f}) = {pred:.2f}x") print(f" -> 模型说 V2 几乎没收益,实测确实几乎没有。两者一致。") print(f"\n V3 的收益不受 rho 限制:中间结果根本不存在,") print(f" 直接就是字节比 K = {K}x(实测 " f"{min(c['speedup'] for c in curves):.2f}~" f"{max(c['speedup'] for c in curves):.2f}x,在 K 附近浮动)。") return {"K": K, "rho": rho, "curves": curves, "v2": v2, "v2_best": best2, "v2_pred": pred, "t1_big_ms": t1_big * 1e3} # ══════════════════════════════════════════════════════════════ # [D] 固定开销:每次调用的地板成本 # ══════════════════════════════════════════════════════════════ def section_D(): print("\n" + "=" * 72) print("[D] 单次调用的固定开销(np.add(x, 1, out=y))") print("=" * 72) rng = np.random.default_rng(RNG_SEED) sizes = [1, 8, 64, 512, 4096, 16384, 65536, 262144, 1 << 20] print(f"\n {'n':>9} {'单次耗时':>11}") rows = [] for n in sizes: a = rng.standard_normal(n).astype(np.float32) b = np.empty_like(a) reps = 20000 if n <= 4096 else 500 t0 = time.perf_counter() for _ in range(reps): np.add(a, 1.0, out=b) t = (time.perf_counter() - t0) / reps rows.append({"n": n, "us": t * 1e6}) print(f" {n:>9} {t*1e6:9.3f} us") # 只用小端拟合:大端会被缓存/主存的拐点污染 ns = np.array([r["n"] for r in rows if r["n"] <= 16384], dtype=float) ts = np.array([r["us"] for r in rows if r["n"] <= 16384], dtype=float) A = np.stack([np.ones_like(ns), ns], axis=1) coef, *_ = np.linalg.lstsq(A, ts, rcond=None) a_us, slope = float(coef[0]), float(coef[1]) print(f"\n 拟合 t = a + b*n(n <= 16384):") print(f" 固定开销 a = {a_us:.3f} us / 次") print(f" 边际成本 b = {slope*1e6:.3f} us / 百万元素") n_half = a_us / slope if slope > 0 else float("inf") print(f" 开销与数据各占一半的临界规模 n* = a/b = {n_half:.0f} 元素" f"({n_half*4/1024:.1f} KiB)") print(f"\n {K_CHAIN} 段链({2*K_CHAIN} 次调用)里固定开销的占比:") share = [] for r in rows: n = r["n"] fixed = 2 * K_CHAIN * a_us tc = 2 * K_CHAIN * (a_us + slope * n) frac = fixed / tc if tc > 0 else 1.0 share.append({"n": n, "fixed_frac": frac}) print(f" n={n:>9}: 固定部分占 {frac*100:5.1f}%") return {"rows": rows, "a_us": a_us, "b_us_per_elem": slope, "n_half": n_half, "share": share, "K": K_CHAIN} # ══════════════════════════════════════════════════════════════ # [E] attention 尾部:本机 softmax 卡在哪 # ══════════════════════════════════════════════════════════════ def _softmax_unfused(x, out): m = np.max(x, axis=-1, keepdims=True) # 读 np.subtract(x, m, out=out) # 读 x + 写 out np.exp(out, out=out) # 读 + 写 s = np.sum(out, axis=-1, keepdims=True) # 读 np.divide(out, s, out=out) # 读 + 写 return out def _softmax_tiled(x, out, rows): for i in range(0, x.shape[1], rows): sl = slice(i, i + rows) blk, dst = x[:, sl, :], out[:, sl, :] m = np.max(blk, axis=-1, keepdims=True) np.subtract(blk, m, out=dst) np.exp(dst, out=dst) s = np.sum(dst, axis=-1, keepdims=True) np.divide(dst, s, out=dst) return out def section_E(b_dram): print("\n" + "=" * 72) print("[E] attention 尾部:本机 softmax 到底卡在哪") print("=" * 72) rng = np.random.default_rng(RNG_SEED) H, S = 8, 1024 x = rng.standard_normal((H, S, S)).astype(np.float32) nbytes = x.nbytes print(f"\n score 矩阵 [H={H}, S={S}, S={S}] = {nbytes/MIB:.1f} MB") out = np.empty_like(x) bench(lambda: _softmax_unfused(x, out), reps=3) t_unf = bench(lambda: _softmax_unfused(x, out), reps=11) ref = out.copy() o2 = np.empty_like(x) bench(lambda: _softmax_tiled(x, o2, 64), reps=3) t_til = bench(lambda: _softmax_tiled(x, o2, 64), reps=11) diff = float(np.max(np.abs(ref - o2))) traffic = 8 * nbytes # [A3] 里的 8 份读写 roof = traffic / (b_dram * 1e9) * 1e3 print(f"\n 不融合(5 个 kernel,共 {traffic/MIB:.0f} MB 读写): {t_unf*1e3:8.3f} ms") print(f" 分块融合(rows=64) : {t_til*1e3:8.3f} ms" f" speedup={t_unf/t_til:.2f}x diff={diff:.1e}") print(f"\n 纯访存下界 = {traffic/MIB:.0f} MB / {b_dram:.0f} GB/s = {roof:.3f} ms") print(f" 实测 / 下界 = {t_unf*1e3/roof:.2f}x") print(f" -> 远高于 1:这里卡的是 exp 的计算本身,不是访存。") print(f" 分块省了字节但一个 exp 都没少,所以反而更慢。") print(f" 这正是「融合省不了算错的东西」的直接证据。") return {"H": H, "S": S, "score_MB": nbytes / MIB, "t_unfused_ms": t_unf * 1e3, "t_tiled_ms": t_til * 1e3, "speedup": t_unf / t_til, "diff": diff, "traffic_MB": traffic / MIB, "roof_ms": roof, "over_roof": t_unf * 1e3 / roof} # ══════════════════════════════════════════════════════════════ def main(): import sys which = sys.argv[1] if len(sys.argv) > 1 else "ALL" res = {} if which in ("ALL", "A"): res["A"] = section_A() if which in ("ALL", "B"): res["B"] = section_B() if which in ("ALL", "C"): rho = res.get("B", {}).get("rho") if rho is None: res["B"] = section_B() rho = res["B"]["rho"] res["C"] = section_C(rho) if which in ("ALL", "D"): res["D"] = section_D() if which in ("ALL", "E"): bd = res.get("B", {}).get("B_dram_GBs") if bd is None: res["B"] = section_B() bd = res["B"]["B_dram_GBs"] res["E"] = section_E(bd) p = os.path.join(HERE, "_fusion_results.json") with open(p, "w", encoding="utf-8") as f: json.dump(res, f, ensure_ascii=False, indent=2) print(f"\n结果已写入 {p}") if __name__ == "__main__": main() make_figures.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """ make_figures.py —— 画本文的 5 张配图。 数据全部读已经跑完的实验(_fusion_results.json / _launch_results.json), 不在这里重新算,避免图上的数字和正文漂移。 运行: python make_figures.py """ from __future__ import annotations import json import os import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np HERE = os.path.dirname(os.path.abspath(__file__)) FIGDIR = os.path.join(HERE, "..", "figures") try: os.makedirs(FIGDIR, exist_ok=True) except FileExistsError: pass # 配色(正文写「这张图要看什么」时按这几个名字描述) C_MAIN = "#1f4e79" # 深蓝:主曲线 C_ALT = "#c1440e" # 橙红:对照曲线 C_GREEN = "#2e7d32" # 绿:第三组 / 融合后 C_PURPLE = "#7b1fa2" # 紫:标注线 C_GRAY = "#8a8a8a" # 灰:参考线 C_LIGHT = "#bcd7ee" # 浅蓝:填充 C_RED = "#b71c1c" # 红:越界 / 地板 plt.rcParams["font.sans-serif"] = ["PingFang SC", "Arial Unicode MS", "DejaVu Sans"] plt.rcParams["axes.unicode_minus"] = False plt.rcParams["figure.dpi"] = 130 plt.rcParams["savefig.dpi"] = 130 MIB = float(1024 ** 2) def _load(name): p = os.path.join(HERE, name) if not os.path.exists(p): raise SystemExit(f"缺少 {name},先跑 fusion_lab.py ALL / launch_model.py") with open(p, encoding="utf-8") as f: return json.load(f) # ══════════════════════════════════════════════════════════════ # 图 1:字节账本 # ══════════════════════════════════════════════════════════════ def fig_traffic(fus): A = fus["A"] items = [ ("8 段逐元素链\n(4M 元素 fp32)", A["chain"]["unfused_MiB"], A["chain"]["fused_MiB"]), ("RMSNorm\n(4M 元素 fp32)", A["rmsnorm"]["unfused_MiB"], A["rmsnorm"]["fused_MiB"]), ("attention softmax\n(B=8,H=16,S=2048,bf16)", A["softmax"]["unfused_GiB"] * 1024, A["softmax"]["fused_GiB"] * 1024), ] labels = [t[0] for t in items] unf = [t[1] for t in items] fs = [t[2] for t in items] fig, ax = plt.subplots(figsize=(9.2, 4.4)) y = np.arange(len(items)) h = 0.34 ax.barh(y + h / 2, unf, height=h, color=C_ALT, label="不融合") ax.barh(y - h / 2, fs, height=h, color=C_GREEN, label="融合后") for i, (u, f) in enumerate(zip(unf, fs)): ax.text(u * 1.15, i + h / 2, f"{u:.1f} MiB", va="center", fontsize=9, color=C_ALT) ax.text(f * 1.15, i - h / 2, f"{f:.1f} MiB", va="center", fontsize=9, color=C_GREEN) ax.annotate(f"{u/f:.0f}x", xy=(u, i), xytext=(u * 1.15, i + 0.42), fontsize=9, color=C_MAIN, fontweight="bold") ax.set_yticks(y) ax.set_yticklabels(labels, fontsize=9) ax.set_xscale("log") ax.set_xlim(8, 40000) ax.set_xlabel("搬动的字节数(对数刻度)", fontsize=10) ax.set_title("图 1:融合前后各搬多少字节", fontsize=12, color=C_MAIN) ax.legend(loc="lower right", fontsize=9) ax.grid(axis="x", alpha=0.3, linestyle=":") ax.set_axisbelow(True) fig.tight_layout() fig.savefig(os.path.join(FIGDIR, "fig_traffic.png")) plt.close(fig) # ══════════════════════════════════════════════════════════════ # 图 2:score 矩阵的放大倍数 + 因果块级跳过 # ══════════════════════════════════════════════════════════════ def fig_blowup(fus): A = fus["A"] sm = A["softmax"] ca = A["causal"] fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(10.4, 4.2)) # 左:放大倍数随 S 增长 ss = [d["S"] for d in sm["blowup_curve"]] xx = [d["x"] for d in sm["blowup_curve"]] ax1.plot(ss, xx, "o-", color=C_MAIN, linewidth=2, markersize=7) ax1.plot(ss, ss, "--", color=C_GRAY, linewidth=1.2, label="线性参考(若按 S 增长)") ax1.set_xscale("log", base=2) ax1.set_yscale("log", base=2) ax1.set_xticks(ss) ax1.set_xticklabels([str(s) for s in ss], fontsize=9) ax1.set_yticks(xx) ax1.set_yticklabels([f"{v:.0f}x" for v in xx], fontsize=9) ax1.minorticks_off() ax1.set_xlabel("序列长度 S", fontsize=10) ax1.set_ylabel("score 矩阵 / 输入激活", fontsize=10) ax1.set_title(f"中间产物被放大 H*S/d = {sm['H']}*S/{sm['d']} 倍", fontsize=11, color=C_MAIN) ax1.grid(alpha=0.3, linestyle=":") ax1.legend(fontsize=8, loc="upper left") ax1.set_axisbelow(True) for s, v in zip(ss, xx): ax1.annotate(f"{v:.0f}x", (s, v), textcoords="offset points", xytext=(6, -12), fontsize=8, color=C_MAIN) # 右:因果块级跳过 brs = [str(b["br"]) for b in ca["blocks"]] fr = [b["frac"] * 100 for b in ca["blocks"]] bars = ax2.bar(brs, fr, color=C_LIGHT, edgecolor=C_MAIN, width=0.55) ax2.axhline(ca["exact_frac"] * 100, color=C_GREEN, linestyle="--", linewidth=1.6, label=f"理论上界 {ca['exact_frac']*100:.1f}%") ax2.axhline(100, color=C_GRAY, linestyle=":", linewidth=1.2, label="不跳过 100%") for b, v in zip(bars, fr): ax2.text(b.get_x() + b.get_width() / 2, v + 1.2, f"{v:.1f}%", ha="center", fontsize=9, color=C_MAIN) ax2.set_ylim(0, 118) ax2.set_xlabel("块边长", fontsize=10) ax2.set_ylabel("实际算的工作量占比 (%)", fontsize=10) ax2.set_title(f"因果掩码块级跳过(S={ca['S']})", fontsize=11, color=C_MAIN) ax2.legend(fontsize=8, loc="upper left") ax2.grid(axis="y", alpha=0.3, linestyle=":") ax2.set_axisbelow(True) fig.suptitle("图 2:attention 的两本账——中间产物多大、白算多少", fontsize=12, color=C_MAIN, y=1.0) fig.tight_layout() fig.savefig(os.path.join(FIGDIR, "fig_blowup.png")) plt.close(fig) # ══════════════════════════════════════════════════════════════ # 图 3:融合三形态 × 规模扫描 # ══════════════════════════════════════════════════════════════ def fig_three_forms(fus): C = fus["C"] D = fus["D"] K = C["K"] a = D["a_us"] cur = C["curves"] ns = [c["n"] for c in cur] t1 = [c["t1_us"] for c in cur] t3 = [c["t3_us"] for c in cur] sp = [c["speedup"] for c in cur] fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(10.6, 4.3)) # 左:耗时 vs 规模 ax1.loglog(ns, t1, "o-", color=C_ALT, linewidth=2, markersize=6, label="V1 不融合(2K 次调用)") ax1.loglog(ns, t3, "s-", color=C_GREEN, linewidth=2, markersize=6, label="V3 代数合并(2 次调用)") floor = 2 * K * a ax1.axhline(floor, color=C_RED, linestyle="--", linewidth=1.4, label=f"V1 的发射地板 2K*a = {floor:.1f} μs") ax1.axhline(2 * a, color=C_PURPLE, linestyle=":", linewidth=1.4, label=f"V3 的发射地板 2*a = {2*a:.1f} μs") # V2 在最大规模上的结果 v2b = C["v2_best"] ax1.plot([ns[-1]], [v2b["ms"] * 1e3], "D", color=C_GRAY, markersize=9, label=f"V2 分块融合 {v2b['ms']*1e3:.0f} μs") ax1.set_xlabel("张量元素数", fontsize=10) ax1.set_ylabel("单次耗时 (μs)", fontsize=10) ax1.set_title(f"三种形态的耗时(K={K})", fontsize=11, color=C_MAIN) ax1.legend(fontsize=8, loc="upper left") ax1.grid(alpha=0.3, linestyle=":") ax1.set_axisbelow(True) # 右:加速比 ax2.semilogx(ns, sp, "o-", color=C_MAIN, linewidth=2, markersize=6, label="V1 / V3(代数合并)") ax2.axhline(K, color=C_GREEN, linestyle="--", linewidth=1.6, label=f"字节比上限 K = {K}x") ax2.axhline(C["v2_pred"], color=C_GRAY, linestyle=":", linewidth=1.6, label=f"V2 模型预测 {C['v2_pred']:.2f}x") ax2.plot([ns[-1]], [v2b["speedup"]], "D", color=C_ALT, markersize=9, label=f"V2 实测 {v2b['speedup']:.2f}x") ax2.set_ylim(0, K * 1.25) ax2.set_xlabel("张量元素数", fontsize=10) ax2.set_ylabel("加速比", fontsize=10) ax2.set_title("各自兑现了多少", fontsize=11, color=C_MAIN) ax2.legend(fontsize=8, loc="lower left") ax2.grid(alpha=0.3, linestyle=":") ax2.set_axisbelow(True) fig.suptitle("图 3:融合三形态 × 规模扫描(本机实测,M1 Pro / numpy)", fontsize=12, color=C_MAIN, y=1.0) fig.tight_layout() fig.savefig(os.path.join(FIGDIR, "fig_three_forms.png")) plt.close(fig) # ══════════════════════════════════════════════════════════════ # 图 4:单次调用的固定开销 # ══════════════════════════════════════════════════════════════ def fig_launch(fus): D = fus["D"] rows = D["rows"] ns = np.array([r["n"] for r in rows], dtype=float) us = np.array([r["us"] for r in rows], dtype=float) a, b = D["a_us"], D["b_us_per_elem"] n_half = D["n_half"] fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(10.6, 4.2)) # 左:延迟 vs 规模,含地板 ax1.loglog(ns, us, "o", color=C_MAIN, markersize=7, label="实测") grid = np.logspace(0, np.log10(ns.max()), 60) ax1.loglog(grid, a + b * grid, "-", color=C_ALT, linewidth=1.8, label=f"拟合 a + b*n(a={a:.2f} μs)") ax1.axhline(a, color=C_RED, linestyle="--", linewidth=1.6, label=f"固定开销地板 a = {a:.2f} μs") ax1.axvline(n_half, color=C_PURPLE, linestyle=":", linewidth=1.6, label=f"各占一半 n* = {n_half:.0f} 元素") ax1.set_xlabel("张量元素数", fontsize=10) ax1.set_ylabel("单次调用耗时 (μs)", fontsize=10) ax1.set_title("小到一定程度,耗时就和数据无关了", fontsize=11, color=C_MAIN) ax1.legend(fontsize=8, loc="upper left") ax1.grid(alpha=0.3, linestyle=":") ax1.set_axisbelow(True) # 右:一条 K 段链里固定开销的占比 share = D["share"] sn = np.array([s["n"] for s in share], dtype=float) sf = np.array([s["fixed_frac"] for s in share]) * 100 ax2.semilogx(sn, sf, "o-", color=C_MAIN, linewidth=2, markersize=6) ax2.axhline(50, color=C_GRAY, linestyle=":", linewidth=1.3) ax2.fill_between(sn, 0, sf, color=C_LIGHT, alpha=0.55) for x, v in zip(sn, sf): if v > 55 or v < 12: ax2.annotate(f"{v:.0f}%", (x, v), textcoords="offset points", xytext=(0, 8), fontsize=8, color=C_MAIN, ha="center") ax2.set_ylim(0, 108) ax2.set_xlabel("张量元素数", fontsize=10) ax2.set_ylabel("固定开销占总耗时 (%)", fontsize=10) ax2.set_title(f"{D['K']} 段链({2*D['K']} 次调用)里,发射占多少", fontsize=11, color=C_MAIN) ax2.grid(alpha=0.3, linestyle=":") ax2.set_axisbelow(True) fig.suptitle("图 4:每次调用都有一个和数据规模无关的地板(本机实测)", fontsize=12, color=C_MAIN, y=1.0) fig.tight_layout() fig.savefig(os.path.join(FIGDIR, "fig_launch.png")) plt.close(fig) # ══════════════════════════════════════════════════════════════ # 图 5:CUDA Graph 的时间线 # ══════════════════════════════════════════════════════════════ def fig_graph(lnc): F1 = lnc["F1"] F2 = lnc["F2"] tl, te, td = F1["t_launch"], F1["t_exec"], F1["t_dispatch"] fig, (ax1, ax2) = plt.subplots( 1, 2, figsize=(11.2, 4.3), gridspec_kw={"width_ratios": [1.25, 1]}) # ── 左:甘特图,只画前 8 个 kernel ── n_show = 8 # eager:CPU 逐个发射,GPU 要等「前一个跑完」且「这个已发出」才能开工。 # 时刻逐事件推算,和图里的条形严格同源。 gpu_free, segs_cpu, segs_gpu, segs_idle = 0.0, [], [], [] for i in range(n_show): cpu_done = (i + 1) * tl start = max(gpu_free, cpu_done) if start > gpu_free: segs_idle.append((gpu_free, start - gpu_free)) segs_gpu.append((start, te)) segs_cpu.append((i * tl, tl)) gpu_free = start + te ax1.broken_barh(segs_cpu, (3.4, 0.8), facecolors=C_LIGHT, edgecolor=C_MAIN, linewidth=0.6) ax1.broken_barh(segs_gpu, (2.2, 0.8), facecolors=C_GREEN, edgecolor=C_GREEN, linewidth=0.6) ax1.broken_barh(segs_idle, (2.2, 0.8), facecolors="#f0c9c9", edgecolor=C_RED, linewidth=0.6, hatch="//") # graph:一次发射,之后背靠背 g_cpu = [(0, tl)] g_gpu = [(tl, n_show * (te + td))] ax1.broken_barh(g_cpu, (1.0, 0.8), facecolors=C_LIGHT, edgecolor=C_MAIN, linewidth=0.6) ax1.broken_barh(g_gpu, (-0.2, 0.8), facecolors=C_GREEN, edgecolor=C_GREEN, linewidth=0.6) ax1.set_yticks([3.8, 2.6, 1.4, 0.2]) ax1.set_yticklabels(["CPU 发射", "GPU 执行", "CPU 发射", "GPU 执行"], fontsize=9) ax1.set_xlabel("时间 (μs)", fontsize=10) ax1.set_xlim(-0.5, n_show * tl + te + 3) ax1.set_ylim(-0.6, 5.4) ax1.set_title(f"eager(上)vs CUDA Graph(下) " f"发射 {tl} μs / 执行 {te} μs", fontsize=10.5, color=C_MAIN) ax1.grid(axis="x", alpha=0.3, linestyle=":") ax1.set_axisbelow(True) ax1.text(0.4, 4.95, "eager:每次发射 GPU 都要等(斜纹 = 空转)", fontsize=8.5, color=C_ALT, va="center") ax1.text(0.4, 0.78, "graph:只发射一次,之后 GPU 背靠背(无空转)", fontsize=8.5, color=C_GREEN, va="center") # ── 右:加速比 vs 单 kernel 执行时间 ── rws = F2["rows"] tes = [r["t_exec"] for r in rws] sps = [r["speedup"] for r in rws] ax2.plot(tes, sps, "o-", color=C_MAIN, linewidth=2, markersize=6) ax2.axhline(1.0, color=C_GRAY, linestyle=":", linewidth=1.3, label="盈亏线 1.0x") ax2.axvline(F2["t_launch"], color=C_RED, linestyle="--", linewidth=1.5, label=f"发射成本 {F2['t_launch']} μs") ax2.fill_between(tes, 0, 1.0, where=np.array(sps) < 1.0, color="#f0c9c9", alpha=0.6, label="graph 反而更慢") ax2.set_xscale("log") ax2.set_xlabel("单个 kernel 的执行时间 (μs)", fontsize=10) ax2.set_ylabel("CUDA Graph 加速比", fontsize=10) ax2.set_title(f"N={F2['N']} 个 kernel", fontsize=11, color=C_MAIN) ax2.legend(fontsize=8, loc="lower right") ax2.grid(alpha=0.3, linestyle=":") ax2.set_axisbelow(True) fig.suptitle("图 5:CUDA Graph 省的是发射,不是字节(模型推演)", fontsize=12, color=C_MAIN, y=1.0) fig.tight_layout() fig.savefig(os.path.join(FIGDIR, "fig_graph.png")) plt.close(fig) # ══════════════════════════════════════════════════════════════ def main(): fus = _load("_fusion_results.json") lnc = _load("_launch_results.json") fig_traffic(fus) fig_blowup(fus) fig_three_forms(fus) fig_launch(fus) fig_graph(lnc) print("已生成 5 张图:") for f in sorted(os.listdir(FIGDIR)): if f.endswith(".png"): p = os.path.join(FIGDIR, f) print(f" {f} {os.path.getsize(p)/1024:.0f} KiB") if __name__ == "__main__": main()
2026年10月02日
1 阅读
0 评论
0 点赞
2026-10-02
AIGC 基本功|离散化表征:VQ-VAE 与 VQGAN-VQGAN
离散化表征:VQ-VAE 与 VQGAN 到底在解决什么问题 所属方向:表征 | 难度:进阶 | 前置知识:变分下界与重参数化(ELBO)、VAE 结构与训练目标(知道「重参数化」和「KL 项怎么来的」就够了) 关键词:VQ-VAE、VQGAN、码本、straight-through、码本坍缩、感知损失 01. 为什么需要它 先算一笔账,这笔账是 VQGAN 那篇论文(Taming Transformers, 2012.09841)的全部动机。把一张 256×256 的 RGB 图像当成序列直接做自回归:每个像素三个通道各算一个 token,序列长度是 256 × 256 × 3 = 196,608。自回归 Transformer 的注意力矩阵是序列长度的平方,也就是 3.87 × 10¹⁰ 个元素——一张图都喂不进去,更别说训练了。 我按这个口径把不同下采样率的账都算了一遍(token_budget.py,本机实跑): 下采样率 f 256² 图像的 token 数 注意力矩阵规模 相对像素级 像素级 RGB 196,608 3.87 × 10¹⁰ 1 f=4 4,096 1.68 × 10⁷ 4.3 × 10⁻⁴ f=8 1,024 1.05 × 10⁶ 2.7 × 10⁻⁵ f=16 256 6.55 × 10⁴ 1.7 × 10⁻⁶ 这张图要看的是两件事:左图里三条 tokenizer 曲线与像素级虚线之间的纵向鸿沟——f=16 时 256² 图像只有 256 个 token,是像素级序列的 1/768;右图是同一个事实在注意力开销上的投影,柱子从 3.87 × 10¹⁰ 掉到 6.55 × 10⁴,省了约 59 万倍。分辨率越高鸿沟越陡(1024² 图像在 f=16 下是 4,096 个 token),视频更极端:5 秒 24fps 共 121 帧,时间下采样 4 倍、空间 f=16 时只要 7,680 个 token。没有这个压缩,自回归视频生成连第一步都迈不出去。 但「短」只是必要条件,自回归还要求「离散」。语言模型之所以好训,是因为下一个 token 的预测是 K 分类交叉熵——一个良性的、方差可控的目标。连续潜变量上做自回归就得给每个位置配一个连续密度(混合密度网络那一路),训练病态且难以和大模型基建兼容。把潜变量离散成「码本里的编号」之后,图像生成就变成了「图像版的语言建模」:DALL-E 用 8192 个码字的 dVAE 加 32×32=1024 个 token 的自回归,VQGAN 用 16×16=256 个 token 的自回归,后来的 MaskGIT、VAR 走的都是这条路。 不过离散化是有代价的,而且真正容易被忽视的代价不在「码本有多大」,在「码本用得怎么样」。我在 8×8 小块的玩具上训了一个 K=512 的 VQ 自编码器(collapse_lab.py,实跑):512 个码字里只有 17 个被用到,96.7% 的码一次都没被选中;码本标称 9 bit,实际只用出去 log2(14.50) = 3.86 bit。这篇就讲三件事:量化怎么写进损失函数、码本怎么训练才不会死、以及为什么重建损失必须从 MSE 换成感知损失加对抗损失。 02. 最小可用理解 三句话: 机制:编码器把图像压成连续向量序列 $z_e$;每个 $z_e$ 在码本(K 个可学习的向量)里找最近邻,换成那个码字得到 $z_q$;解码器只从 $z_q$ 重建图像。量化器是唯一的信息瓶颈——解码器看不到任何量化误差之外的信息。 训练:三个损失各管一段。重建损失管「编码器+解码器」 jointly,但 argmin 不可导,梯度靠 straight-through(把 $z_q$ 的梯度原样抄给 $z_e$);码本本身要么用字典损失往编码器输出上拉,要么用 EMA 直接滑向被选中样本的均值;commitment 项(权重 $\beta$)把编码器往码字上拉,防止两头越走越远。 代价:量化误差是硬地板,而且高维码本的容量收益极差(失真只能按 $K^{-2/d}$ 衰减);码本会坍缩;MSE 重建的最优解是条件均值,必然糊——所以 VQGAN 在重建侧换成了 LPIPS 加 PatchGAN。 03. 数学推导 3.1 从 VAE 到 VQ-VAE:KL 项去哪了 VQ-VAE(Neural Discrete Representation Learning, 1711.00937)的名字里有 VAE,推导起点也确实是 ELBO: $$\log p(x) \ge \mathbb{E}_{q(z \mid x)} \big[ \log p(x \mid z) \big] - \mathrm{KL}\big( q(z \mid x) \,\Vert\, p(z) \big)$$ 每一项的含义:$q(z \mid x)$ 是编码器给出的「后验」,$p(z)$ 是我们先验地相信 latent 该有的分布,$p(x \mid z)$ 是解码器。VAE 里这三样都是连续分布,KL 项把后验往先验上压,重参数化让采样可导。VQ-VAE 把这三样全换了: $q(z \mid x)$ 不再是分布,而是确定性的:$z$ 就是被选中的那个码字 $e_k$,配合 one-hot 指示变量,相当于 $q(z = e_k \mid x) = 1$,对其他码字取 0; $p(z)$ 取 K 个码字上的均匀分布; $p(x \mid z)$ 是解码器(高斯均值或离散化的像素分布),第一项就是重建损失。 把这个确定性后验和均匀先验代进 KL:$q$ 在 $e_k$ 处为 1、其余为 0,求和只剩被选中那一项: $$\mathrm{KL}\big( q \,\Vert\, p \big) = \sum_{j=1}^{K} q_j \log \frac{q_j}{p_j} = 1 \cdot \log \frac{1}{1/K} = \log K$$ $\log K$ 是一个常数,对梯度没有任何贡献。VQ-VAE 的 KL 项就此消失——没有 KL、没有重参数化、没有「均值方差都被压向先验」的正则,ELBO 退化成「重建项 + 两个逐样本的 L2 距离项」。名字里的 V 是历史包袱,这也是后面 08 节第一条误解的来源。 3.2 argmin 不可导,梯度要靠「装傻」 量化操作本身是: $$z_q = e_k, \quad k = \arg\min_{j} \, \Vert z_e - e_j \Vert_2^2$$ 符号含义:$z_e \in \mathbb{R}^d$ 是编码器输出(e 指 encoder),$e_j \in \mathbb{R}^d$ 是第 j 个码字,$z_q$ 是替换后的向量(q 指 quantized)。问题出在 $\arg\min$:它是分段常数函数。训练中把 $z_e$ 微动一点点,只要不跨过两个码字的垂直平分面,选中的 $k$ 根本不变,$z_q$ 不变,重建损失也不变。 这不是「梯度小」,是梯度恒等于零。我在训好的 VQ-AE 上用有限差分实测过(vq_core.py 的 [D] 段,512 个样本,沿随机单位方向扰动 $z_e$): 扰动步长 eps 有限差分 $\partial L/\partial z_e$ argmin 保持不变的样本比例 10⁻¹ −7.6 × 10⁻⁶ 0.9941 10⁻² 0.0(精确为零) 1.0000 10⁻³ 0.0(精确为零) 1.0000 10⁻⁴ 0.0(精确为零) 1.0000 eps=10⁻¹ 那一行有 0.59% 的样本跨过了平分面,所以差分不为零——这也正是 argmin「分段常数、边界处跳变」的直接展示。而在平分面之间的整片区域里,重建损失对编码器没有任何梯度。VQ-VAE 的解法是 straight-through:假装量化是恒等映射,把解码器对 $z_q$ 的梯度原封不动地抄给 $z_e$: $$\frac{\partial L}{\partial z_e} \mathrel{:=} \frac{\partial L}{\partial z_q}$$ 要强调的是:这不是真实梯度的估计,是替代品。真实梯度是 0,直通梯度实测平均范数 1.49 × 10⁻⁴(同一组权重),它携带的信息是「如果量化不存在,往哪边调编码器能让重建更好」。它能工作的原因是:编码器的真正职责不是让 $z_e$ 落在哪个精确位置,而是让「选出来的码字」是对的——直通梯度恰好只优化这件事。 实现上只有一行(taming-transformers 的写法,见 05 节):z_q = z + (z_q - z).detach()。前向时 $z_q$ 是真的码字,反向时梯度绕过 $(z_q - z)$ 这个常量直接流向 $z$。 3.3 码本的两条更新路线 直通梯度有个致命遗漏:码本自己拿不到任何来自重建损失的梯度。$z_q$ 是查表查出来的,对 $e_j$ 的导数被查表操作挡住了;直通又把全部梯度引向 $z_e$。如果什么都不加,码本会永远停在初始化的位置。我在玩具上实测过这条「什么都不加」的路线(vq_core.py [E] 段,K=64,800 步):码本位移精确为 0.000000,重建 MSE 0.185734,是正常训练的 3 倍多。 所以码本必须有独立的更新机制,VQ-VAE 给了第一条路——字典损失。把编码器输出当成常数(stop-gradient,记作 sg),把码字往它身上拉: $$L_{\text{codebook}} = \big\Vert \mathrm{sg}[z_e] - e_k \big\Vert_2^2, \qquad \frac{\partial L_{\text{codebook}}}{\partial e_k} = 2\,(e_k - z_e)$$ 第二条路是 EMA(VQ-VAE-2 之后的主流):不用梯度,直接对「每个码字被选中样本的均值」做指数滑动。记第 t 步里码字 j 被选中了 $n_j$ 次、被选中样本之和为 $s_j$,则 $$c_j \leftarrow \gamma c_j + (1-\gamma)\, n_j, \quad m_j \leftarrow \gamma m_j + (1-\gamma)\, s_j, \quad e_j \leftarrow \frac{m_j}{c_j + \epsilon_{\text{smooth}}}$$ $\gamma$ 是滑动系数(taming 里 decay=0.99),$c_j$ 是每个码字的滑动计数,$m_j$ 是滑动累加的样本和,$\epsilon_{\text{smooth}}$ 用来防止某个几乎没人用的码字除以接近零的数(taming 的平滑是 $(c_j + \epsilon) \cdot n / (n + K\epsilon)$,其中 $n$ 是全部计数之和)。EMA 的本质是把码字变成「最近被它编码过的那些向量的滑动平均」,没有学习率要调,这也是它取代字典损失的原因。 3.4 commitment 项:另一头的绳子 现在把绳子接上另一头。EMA 把码字往编码器输出上拉,但没有任何东西把编码器输出往码字上拉——编码器完全可以漂走,让量化误差 $\Vert z_e - e_k \Vert^2$ 失控。commitment 项就是拴住编码器的那根绳: $$L_{\text{commit}} = \beta \,\big\Vert z_e - \mathrm{sg}[e_k] \big\Vert_2^2$$ 注意方向和字典损失正好相反:字典损失动了 $e_k$($z_e$ 被 stop-gradient 冻住),commitment 动了 $z_e$($e_k$ 被冻住)。两项合起来,VQ-VAE 的完整训练目标是: $$L = \underbrace{\log p(x \mid z_q)}_{\text{重建,经直通传给编码器}} + \underbrace{\big\Vert \mathrm{sg}[z_e] - e_k \big\Vert_2^2}_{\text{码本项或 EMA}} + \underbrace{\beta \big\Vert z_e - \mathrm{sg}[e_k] \big\Vert_2^2}_{\text{commitment}}$$ $\beta$ 不是一个可以随手抄的超参。在同一组玩具权重上,我把直通梯度和 commitment 项的梯度范数都量了一下(vq_core.py [D] 段):直通梯度平均范数 1.49 × 10⁻⁴,commitment 梯度平均范数 6.40 × 10⁻²,差 429 倍。$\beta$ 扫描的实测结果(collapse_lab.py [A] 段,K=64,1200 步): $\beta$ 0 0.05 0.25 1.0 4.0 重建 MSE 0.0673 0.0484 0.0600 0.0387 0.0385 活跃码数 14 14 16 17 18 在这个玩具上 $\beta=1.0$ 反而比论文默认的 0.25 好 35%。原因就藏在那个 429 倍里:当重建梯度相对太弱时,加大 $\beta$ 相当于在帮编码器「站稳」在码字附近,量化误差随之下降。$\beta$ 的最优值和你的重建损失量纲绑死——这就是为什么换损失(比如 VQGAN 换成 L1 + LPIPS)之后不能照抄别人的 $\beta$。 3.5 VQGAN 补上的两块 VQ-VAE 的重建损失是逐像素的,这有一个数学上无解的毛病(06 节用实验展开):最优解是条件均值,纹理会被平均掉。VQGAN 把重建侧换成三件套: $$L_{\text{VQGAN}} = \underbrace{\big\Vert x - \hat x \big\Vert_1 + \lambda_{\text{lpips}} L_{\text{LPIPS}}(x, \hat x)}_{\text{感知重建}} + \underbrace{\lambda_{\text{GAN}} \big( -\log D(\hat x) \big)}_{\text{对抗}} + \underbrace{\lambda_{\text{cb}} L_{\text{codebook}} + \beta L_{\text{commit}}}_{\text{量化}}$$ $\hat x$ 是解码器输出,$D$ 是 PatchGAN 判别器,LPIPS 是在 VGG16 五层特征上算距离再加一层学出来的 1×1 卷积(lpips.py 里五个 NetLinLayer,实读源码确认)。对抗项的权重不是手调的,而是自适应的:对解码器最后一层分别求重建损失和对抗损失的梯度,取范数比 $\lambda_{\text{GAN}} \leftarrow \Vert \nabla L_{\text{rec}} \Vert / \Vert \nabla L_{\text{GAN}} \Vert$,让两边的梯度量级匹配——GAN 一开判别器就抢梯度主导权,这是压住它的办法。 3.6 怎么量「码本用得怎么样」 三个从粗到细的指标,别混用: 活跃码数:训练中至少被选中过一次的码字数除以 K。最直观,但它是二值的——一个只被用过 3 次的码字和用过 3 万次的算得一样。 perplexity:把码字的使用频率 $p_j$ 看成一个分布,取 $\mathrm{ppl} = \exp\big( -\sum_{j} p_j \log p_j \big)$。完全均匀使用时等于 K,退化到只用一个码时等于 1。它把长尾压成一个数,是训练日志里最常盯的那个量。 有效 bit:$\log_2 \mathrm{ppl}$,可以直接和标称 bit $\log_2 K$ 比。本文玩具里 K=512 的标称 9 bit 只传出 3.86 bit,57% 的编码容量是白付的。 两个容易踩的坑。第一,perplexity 是分布层面的量:它下降只说明使用分布变尖了,既不告诉你死码落在哪,也不等价于重建质量——要看死码分布得直接画使用次数的直方图(图 4 干的就是这件事)。第二,taming 是在当前 batch 上算 avg_probs 的,batch 越小噪声越大;小 batch 上看到 ppl 上下抖动不等于坍缩,别急着改 decay 或加重启。 04. 代码实现 最小实现不需要网络,一个「线性编码器 + VQ + 线性解码器」就能把所有机制跑出来。数据是我合成的 8×8 小块:16 个低频类心,类内加连续变化和白噪声(make_patch_data,4096 个样本,种子 0,全部结果可复现)。核心的量化与直通就这几行(完整脚本在附录): def quantize(z_e, codebook): d2 = ((z_e[:, None, :] - codebook[None, :, :]) ** 2).sum(-1) # [N, K] idx = d2.argmin(axis=1) # 最近邻 return codebook[idx], idx, d2[np.arange(len(idx)), idx] # 训练循环里(z_q 是查表结果,e_k = z_q): x_hat = z_q @ params["Wd"] + params["bd"] dxh = 2.0 * (x_hat - x) / (len(x) * d_in) # 对 x_hat 的梯度 dz_q = dxh @ params["Wd"].T # 解码器传给 z_q 的梯度 dz_e = dz_q + beta * 2.0 * (z_e - z_q) / (len(x) * d_lat) # 直通 + commitment # 码本走 EMA(onehot 统计被选中的次数与样本和,见附录 vq_core.py) 第一件事:误差到底由哪几块组成。 用同一组训好的权重,把「走不走量化」作为开关(vq_core.py [C] 段,K=64,d=8): 通路 重建 MSE/像素 连续自编码器基线(无 VQ,单独训练到收敛) 0.011082 同一组 VQ-AE 权重,绕开量化(直接用 $z_e$ 过解码器) 0.023360 同一组 VQ-AE 权重,正常走量化 0.060345 两个结论都值得停下来想。第一,量化把误差从 0.0234 推到 0.0603,多出来的 0.0370(+158%)全是量化的账。第二,VQ-AE 的连续通路 0.0234 比独立训练的连续基线 0.0111 差了一倍——量化还会反过来把编码器带偏:commitment 在拉编码器,编码器为迁就码字牺牲了一部分子空间的质量。这部分「隐性代价」在只报一个重建指标时完全看不见。 第二件事:码本容量 K 的收益被码本维度 d 卡死。 固定一个训好的连续编码器,对它的潜变量做 k-means(k-means 就是「给定 K 个码字的最优最近邻量化器」的近似),在留出集上测失真随 K 的变化: 码本维度 d 理论斜率 −2/d 实测斜率(K≥16 段) K 从 8 加到 512,失真降多少 2 −1.00 −0.87 66.5 倍 4 −0.50 −0.40 33.8 倍 8 −0.25 −0.28 13.9 倍 16 −0.125 −0.31 12.9 倍 这张图要看的是实测线(实线)和理论斜率(点线)的贴合程度:d=2、4、8 都贴得不错(量化理论的 Zador 渐近:最优失真 $\propto K^{-2/d}$),d=16 在大 K 端因为每个码字只剩 8 个训练样本而偏离。直白地说:d=8 时把码本从 8 加到 512(64 倍),量化失真只降到 1/13.9——这就是为什么后面 LFQ、FSQ、残差量化都要在「维度」上做文章,而不是无脑加码本。 第三件事:码本坍缩与两副解药。 K=512、随机初始化、EMA 更新,训练 1200 步: 配置 重建 MSE 活跃码数 perplexity 有效 bit(log2 ppl) 随机初始化 0.049067 17 / 512 14.50 3.86 k-means 初始化 0.019695 325 / 512 217.26 7.76 死码重采样重启 0.018368 500 / 512 426.71 8.74 这张图要看的是三条线的起点和走势:红线(随机初始化)从第 200 步起就钉死在 20 以下——坍缩发生在训练极早期,一旦码字没被选中过,它就再也没有机会被选中(EMA 的计数是零,均值是零除零);蓝线(k-means 初始化)起点就是 261;绿线(死码重采样,把长期没人用的码字随机替换成真实的编码器输出)稳定在 470 上下。 这张图要看的是横轴超过 20 之后红线的缺席:随机初始化的码本只有 17 个码被用过,剩下的在 log 轴上根本画不出来;绿线的使用占比分布虽然仍是长尾(最热的码占 0.13%),但整条尾巴被抬起来了。两个数字合起来读:坍缩让 9 bit 的码本只传出 3.86 bit,而这是在重建 MSE 上实打实付了 2.7 倍代价换来的(0.0491 对 0.0184)。 顺带一提,如果不用 EMA 而用字典损失的梯度更新码本,学习率极其敏感(码本是普通 SGD,不走 Adam):在同一个玩具上把码本学习率从 0.02 加到 10,重建 MSE 从 0.2418 降到 0.0538,活跃码从 8 涨到 14——始终追不上 EMA。这就是主流实现全部转向 EMA 的实证原因。 05. 工业级实现对照 上面是最小实现,生产代码在 taming-transformers(VQGAN 官方仓库,本文对照的是 taming/modules/vqvae/quantize.py 与 taming/modules/losses/vqperceptual.py,见 GitHub 源码,以下均以 2026-10 时的 master 为准)。逐个对照: 距离计算用展开式。 最小实现里我直接算 (z[:, None] - cb[None]) ** 2,会物化一个 N×K×d 的中间张量;taming 用 $\Vert z - e \Vert^2 = \Vert z \Vert^2 + \Vert e \Vert^2 - 2 z \cdot e$ 展开,只算 N×K 的矩阵乘: d = torch.sum(z_flattened ** 2, dim=1, keepdim=True) \ + torch.sum(self.embedding.weight**2, dim=1) \ - 2 * torch.einsum('bd,dn->bn', z_flattened, ...) 图像 tokenizer 的 N 是 batch × H/16 × W/16,K 上万,这个展开是省显存的关键。 直通就一行,和推导完全一致:z_q = z + (z_q - z).detach()。 taming 保留了一个带 bug 的版本。 VectorQuantizer2 有个 legacy 开关,legacy=True(默认)时损失是 loss = mean((z_q.detach()-z)**2) + beta * mean((z_q - z.detach())**2)——$\beta$ 被加在了码本项上,而不是 commitment 项,与论文公式相反。源码注释直说这是历史 bug,为了兼容旧 checkpoint 保留默认。读老代码、对老权重时要注意这个坑。 EMA 码本的权重是 requires_grad=False 的。EmbeddingEMA 里 weight、cluster_size、embed_avg 全是关闭梯度的参数,更新完全靠滑动平均——和我 3.3 节的推导一致。taming 还顺手算了 perplexity:exp(-sum(avg_probs * log(avg_probs))),这正是我用来诊断坍缩的量。 decay 不是「越接近 1 越稳」,它是一个时间常数。 衰减系数 $\gamma$ 决定码字「记得多久以前的样本」:$\gamma=0.99$ 时一个历史样本的贡献半衰期是 $\log 0.5 / \log 0.99 \approx 69$ 步,458 步后只剩 1%;$\gamma=0.999$ 半衰期拉到 693 步,$\gamma=0.95$ 只有 13.5 步。这个数要和 batch 里的 token 数对着调:VQGAN f=16 时一张 256² 图给 256 个 token,batch=8 一共 2048 个 token,摊到 K=16384 的码本上平均每步每个码字只被选中 0.125 次。也就是说大码本下 EMA 的计数极其稀疏,decay 太小会让码字在两次命中之间就被洗回零——这是「大码本更容易坍缩」的一条工程解释,也是 LFQ/FSQ 从结构上绕开它的动机。 损失在 losses/vqperceptual.py:L1 + LPIPS + hinge GAN。 三个值得抄的细节:一是 adopt_weight,判别器从第 disc_start 步才介入(先让重建学好,再上对抗);二是 calculate_adaptive_weight,取 $\Vert \nabla_{\text{last}} L_{\text{rec}} \Vert / \Vert \nabla_{\text{last}} L_{\text{GAN}} \Vert$ 并 clamp 到 10⁴;三是判别器是 NLayerPatchGAN,按 patch 判真假,这样高分辨率下判别器参数量不随分辨率爆炸。 官方数字(仓库 README 的重建 FID 表): 模型 f 码本 K 重建 rFID DALL-E dVAE(Gumbel) 8 8192 33.88 VQGAN ImageNet 16 1024 10.54 VQGAN ImageNet 16 16384 7.41 VQGAN OpenImages 8 256 1.49 VQGAN OpenImages 8 16384 1.14 两行读法:同为 f=16,K 从 1024 加到 16384,rFID 从 10.54 降到 7.41——码本容量确实有用,但注意这是在 d=256 的码本维度上(vq_model.py 里 quant_conv 把通道投影到 embed_dim=256),按 $K^{-2/d}$ 的规律这个收益已经非常温和。另一个对照更惊人:VQGAN f=8 K=256 的 rFID(1.49)比 DALL-E dVAE f=8 K=8192(33.88)好了 22 倍——感知损失加对抗带来的提升,比码本大 32 倍带来的提升大得多。这就是 VQGAN 论文标题里 "taming" 的真正含义。 06. 代价与边界 代价一:量化误差是硬地板,而且维度惩罚很重。 04 节的表已经给了 d=8 的数字:K=512 时量化误差仍把重建误差推高 80.5%(相对连续通路 0.011378)。想压低这块,有两条数学上已知的路:降有效维度(FSQ 干的事:把码本从「K 个 d 维向量」换成「每维只有少量取值」,维度语义变了但利用率为 100%),或者分层量化(残差 VQ、VAR 用的 RQ-VAE:一层量化不完的残差给下一层)。蛮力加 K 是最差的一条路。 代价二:码本坍缩几乎必然发生,解药都有副作用。 实测里随机初始化有 96.7% 死码;k-means 初始化要额外跑一次 k-means 且只在训练初期有用;死码重采样最有效,但它等价于「用随机重启换利用率」——被重启的码字携带的信息丢了,而且重启阈值又是一个新超参。工业界还有第三条路:直接改量化方式让坍缩在结构上不可能发生(LFQ 把每个维度独立二值化/多值化,FSQ 同理),这超出了本文范围,07 节给出处。 代价三:感知重建会「编」细节。 这是 3.5 节埋的伏笔,用一个能算清楚的玩具展开(perceptual_lab.py)。构造:一个 latent 对应两种等概率的纹理 $x = m \pm p$(m 是低频内容,p 是棋盘纹理,纹理占 83.5% 的梯度能量)。在候选输出 $m + c \cdot p$ 上比较两种损失: 候选输出 MSE PSNR 梯度能量保留 到最近模态的距离 $c=0$(MSE 最优,条件均值) 0.250000 5.30 dB 18.6% 0.250 $c=-0.79$(特征空间最优) 0.404056 3.22 dB 70.7% 0.012 $c=1$(选一个清晰模态) 0.500000 2.29 dB 97.5% 0.000 这张图要看三处:上排四张 8×8 小图里,MSE 最优那格的棋盘纹理消失了(两种极性平均成平色),特征最优那格纹理回来了;左下柱状图说明这个纹理消失在 PSNR 上是加分的(5.30 dB 最高);右下的曲线说明只要特征里给高频一点权重($\alpha > 0.74$,实跑二分定位),最优解就会离开模糊均值往清晰模态走。三条事实合起来的结论是:PSNR 和「看起来真」在这类问题上方向相反——VQGAN 之后没人用 PSNR 报告 tokenizer 的重建质量,rFID 成了标配,原因就在这。 但这笔账的另一面是:对抗训练出来的「细节」不保证是真的。判别器只关心「像不像真图」,不关心「是不是这张图」,所以感知+GAN 的重建会补出数据集里常见的纹理——做压缩、做医学影像这类需要像素保真的任务时,这条路要慎走(GAN 训练不稳的代价也真实存在,05 节的 disc_start 和自适应权重都是为此付的工程税)。 边界:什么时候不该用。 下游不是自回归/离散先验时,离散化没有收益只有损失——Stable Diffusion 的第一-stage 就是连续 VAE(latent_diffusion 那篇讲过),因为扩散模型要的是连续潜空间上的 score,不是离散 token。同样,「码本利用率」这一整章在连续 VAE 里没有对应物。选型的判据一句话:下游要对离散序列做自回归或掩码预测,才需要 VQ tokenizer。 最后坦诚标注边界:本文所有数字来自 8×8 合成小块上的线性自编码器玩具(numpy 实跑,可复现),规模和真实 VQGAN(f=16、d=256、K=16384、百万级图像)差几个数量级;机制层面(直通、EMA、坍缩、感知损失的偏好)我认为可以直接外推,但具体数字(比如 $\beta=1.0$ 更好)不能外推,它依赖我的损失量纲。 07. 经典论文脉络 VQ-VAE(1711.00937, Neural Discrete Representation Learning):第一次把离散 latent 做成端到端可训——straight-through 加字典损失/EMA 的组合沿用至今。 VQ-VAE-2(1906.00446):层级化(全局加局部两级码本),并正式用 EMA 替代字典梯度,perplexity 作为利用率指标从这里普及。 VQGAN(2012.09841, Taming Transformers):感知损失加 PatchGAN 把重建质量拉到 rFID 个位数,第一次让「Transformer 学图像 token」在算力和效果上同时成立。 DALL-E(2102.12092):zero-shot 文生图,dVAE 用 Gumbel-softmax 变分训练(8192 码本、f=8),证明了离散 token 路线在多模态上的可扩展性。 ViT-VQGAN(2110.04627):把卷积 tokenizer 换成 ViT 结构,码本效率(利用率)被单独拿出来分析。 MaskGIT(2202.04200):放弃自回归,改用掩码并行解码——依赖的仍是 VQGAN 的离散 token,说明离散化红利不止自回归一条路。 LFQ / FSQ(2310.05737 MagViT-2 / 2309.15505):从结构上消灭码本坍缩——LFQ 把每维独立量化到固定格点,FSQ 直接用少量取值的整数网格,码本利用率都能到 100%。 VAR(2404.08560):残差 VQ(粗到细多级量化)加「下一尺度预测」,把自回归视觉生成推到与扩散相当的区间,是 2024 年后 tokenizer 论文的必引坐标。 08. 常见误解 误解一:「VQ-VAE 是 VAE 的一种」。 3.1 节推过:后验是确定性的 one-hot,先验是均匀分布,KL 精确等于 $\log K$,是常数,对训练没有任何贡献。没有变分、没有重参数化、没有 KL 正则——ELBO 在这里只是叙事起点,不是训练目标。 误解二:「码本是梯度下降学出来的」。 码本拿不到重建损失的梯度(查表挡住了,直通又把梯度全引向编码器),它的更新要么靠字典损失这一项、要么靠 EMA。taming 的 EMA 实现里码本权重干脆是 requires_grad=False 的。 误解三:「straight-through 是梯度的无偏估计」。 实测(04 节表):真实有限差分精确为 0,直通梯度非零。它不是对真实梯度的估计,是「假装量化不存在」的替代品;真正把它扶正的是 commitment 项,否则编码器会漂走。 误解四:「码本越大越好」。 $K^{-2/d}$ 的维度惩罚加上坍缩风险,让大码本的收益远低于直觉——K=1024 到 16384 在 d=256 上只把 rFID 从 10.54 拉到 7.41(05 节表)。利用率不到 100% 时,标称 bit 和有效 bit 的差距更离谱(本文玩具里 9 bit 只传出 3.86 bit)。 误解五:「$\beta=0.25$ 是默认值,照抄就行」。 $\beta$ 控制的是 commitment 梯度和重建直通梯度的量级比,本文玩具里两者天然差 429 倍,$\beta$ 扫描显示 1.0 反而最好。换了重建损失(MSE 换 L1+LPIPS)量纲就变,$\beta$ 必须重调。 09. 动手验证 三个实验都可以在附录代码里一键复现(numpy only,不需要 GPU): 亲手确认 argmin 的梯度是零:跑 python vq_core.py,看 [D] 段——把 $z_e$ 沿随机方向扰动 10⁻² 到 10⁻⁴,有限差分精确为 0,argmin 保持不变的样本比例是 1.0000;对比直通梯度平均范数 1.49 × 10⁻⁴。 亲手制造并修好码本坍缩:跑 python collapse_lab.py,看 [B][C] 段——K=512 随机初始化最终只有 17 个活码(96.7% 死码)、重建 MSE 0.0491;换死码重采样后 500 个活码、MSE 0.0184。你也可以把 restart_dead=True 关掉再跑一遍,确认结果回到 17。 亲手验证「MSE 必然糊」:跑 python perceptual_lab.py,把文件顶部的 AMP(纹理幅度)从 0.5 改成 0.1 再跑——你会发现梯度能量保留率对纹理幅度极其敏感,而 MSE 最优解的纹理保留永远是 0(条件均值把 ±纹理精确抵消)。 亲手确认「加码本不如降维度」:跑 python codebook_size_law.py,读 [A] 段输出——d=16 时 K 从 8 加到 512 只把失真降到 1/12.9,d=2 时同样 64 倍的预算能降到 1/66.5。再把 make_patch_data 的 n_cluster 从 16 改成 4(类更少、潜变量更集中)重跑,斜率会明显变陡:失真下降的速度由数据分布本身决定,码本参数只是顺着它走。 10. 延伸阅读 前置:变分下界与重参数化(ELBO)、VAE 结构与训练目标——本文 3.1 节是 ELBO 在离散情形下的特例。 同方向:视频 VAE 的时空压缩结构(f_t 的时间维压缩,01 节视频账本用到了)、视频 VAE 的常见 loss 组合(GAN/LPIPS 组合的连续版)。 相关:潜空间扩散与 Stable Diffusion 架构(06 节「什么时候不用 VQ」的那条连续路线)、自回归视频生成(离散 token 的下游)、FID / CLIP Score 到底测了什么(05 节 rFID 的定义与陷阱)。 附录:完整代码 09 节用到的脚本全文如下(token_budget.py、collapse_lab.py、vq_core.py、perceptual_lab.py、codebook_size_law.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 token_budget.py """token_budget.py —— 离散化到底换来多少「序列长度」上的便宜。 VQGAN 那篇论文的动机只有一句话:Transformer 的注意力和序列长度是平方关系, 在像素上做自回归根本不可能,所以必须先把图像压成一串短的离散 token。 本脚本把这笔账算成具体的数: [A] 不同分辨率 / 下采样率下的 token 数与注意力规模 [B] 码本大小 K 决定每个 token 多少 bit,折算成每张图多少 bpp [C] 视频的 token 数(时空压缩一起算) [D] 码本坍缩要付的比特代价:perplexity 才是真实容量 运行:/usr/local/bin/python3 token_budget.py """ from __future__ import annotations import numpy as np def n_tokens(h, w, f, t=1, f_t=1): return (t // f_t) * (h // f) * (w // f) def main(): print("[A] 图像:边长 H 与下采样率 f 决定 token 数(注意力按 n^2 涨)") print(" H f token 数 注意力矩阵 n^2 相对像素级") for H in (256, 512, 1024): base = H * H * 3 # 像素级(RGB 逐通道) for f in (4, 8, 16): n = n_tokens(H, H, f) print(" %-5d %-4d %-12d %-18.3e %.4g" % ( H, f, n, float(n) ** 2, float(n * n) / float(base * base))) print(" (像素级 RGB 序列 %d,n^2 = %.3e)" % (base, float(base) ** 2)) print() print("[B] 码本容量 K 决定每个 token 的 bit 数(256x256 图像)") print(" K bit/token f=8: bit/图 bpp f=16: bit/图 bpp") for K in (256, 512, 1024, 4096, 16384, 262144): bit = np.log2(K) row = [K, bit] for f in (8, 16): n = n_tokens(256, 256, f) total = n * bit row += [total, total / (256 * 256)] print(" %-8d %-10.2f %-12.0f %-7.3f %-12.0f %.3f" % tuple(row)) print(" bpp = bit per pixel。作为参照,JPEG 在中等质量下大约 0.5~1 bpp。") print() print("[C] 视频:时间维也要压(5 秒 24fps = 121 帧,256x256)") print(" f_t f token 数 上下文长度对比(相对 121x256x256x3 像素)") pixel = 121 * 256 * 256 * 3 for f_t in (1, 4, 8): for f in (8, 16): n = n_tokens(256, 256, f, t=120, f_t=f_t) print(" %-5d %-4d %-14d %.5g" % ( f_t, f, n, float(n) / pixel)) print() print("[D] 码本坍缩的比特代价:真实容量看 perplexity,不是看 K") print(" K perplexity 标称 bit 有效 bit 浪费") for K, ppl in [(512, 14.50), (512, 217.26), (512, 426.71), (16384, 14.50), (16384, 1000.0), (16384, 16384.0)]: nominal = np.log2(K) eff = np.log2(ppl) print(" %-6d %-11.2f %-10.2f %-10.2f %.1f%%" % ( K, ppl, nominal, eff, 100 * (nominal - eff) / nominal)) print(" 前两行来自 collapse_lab.py 的实测:K=512 随机初始化只有 17 个码活着,") print(" 9 bit 的码本只传出 3.86 bit;死码重启后回到 8.74 bit。") print() print("[E] 一句话总结") f16 = n_tokens(256, 256, 16) print(" 256x256 图像在 f=16 下是 %d 个 token,是像素级 RGB 序列的 1/%.0f;" % (f16, (256 * 256 * 3) / f16)) print(" 注意力规模从 %.2e 降到 %.2e,省了 %.0f 倍——这才是必须先做 tokenizer 的原因。" % (float(256 * 256 * 3) ** 2, float(f16) ** 2, float(256 * 256 * 3) ** 2 / float(f16) ** 2)) if __name__ == "__main__": main() collapse_lab.py """collapse_lab.py —— 码本坍缩(codebook collapse)是怎么发生的,能救回来多少。 码本坍缩指的是:K 个码字里只有一小撮被用到,剩下的从头到尾一次都没被选中 (死码)。花了一整个 K×d 的码本,只买到 log2(perplexity) 比特的表达力。 [A] commitment 权重 beta 怎么影响坍缩程度 [B] 大码本在训练过程中怎么一步步坍缩(活跃码数曲线) [C] 两种常用解药有多大用:k-means 初始化码本 / 死码重采样重启 运行:/usr/local/bin/python3 collapse_lab.py """ from __future__ import annotations import argparse import numpy as np from vq_core import kmeans, make_patch_data, perplexity, quantize, train_ae BIG_K = 512 def stats(X, params, codebook): """给定训练好的权重与码本,算重建误差与使用分布统计量。""" z_e = X @ params["We"] z_q, idx, _ = quantize(z_e, codebook) cnt = np.bincount(idx, minlength=len(codebook)).astype(np.float64) rec = float((((z_q @ params["Wd"] + params["bd"]) - X) ** 2).mean()) frac = np.sort(cnt / cnt.sum())[::-1] return { "rec": rec, "alive": int((cnt > 0).sum()), "ppl": perplexity(cnt), "top1": float(frac[0]), "bottom_half": float(frac[len(frac) // 2:].sum()), "counts": cnt, } def run_vq(X, K=BIG_K, beta=0.25, steps=1200, seed=0, mode="ema", init_cb=None, restart=False, log_every=100): """跑一次 VQ-AE 训练,返回 (stats, hist)。""" p, c, h, _ = train_ae(X, 8, mode, K=K, beta=beta, steps=steps, seed=seed, restart_dead=restart, init_cb=init_cb, log_every=log_every) return stats(X, p, c), h def main(): ap = argparse.ArgumentParser() ap.add_argument("--steps", type=int, default=1200) ap.add_argument("--K", type=int, default=BIG_K) ap.add_argument("--seed", type=int, default=0) args = ap.parse_args() X, _, _, _ = make_patch_data(n=4096, seed=args.seed) print("数据 X%s,训练 %d 步" % (X.shape, args.steps)) print() # ---------------- [A] beta 扫描 ---------------- print("[A] commitment 权重 beta 的影响(K=64,码本 EMA,%d 步)" % args.steps) print(" beta 重建MSE 活跃码 perplexity") for beta in (0.0, 0.05, 0.25, 1.0, 4.0): s, _ = run_vq(X, K=64, beta=beta, steps=args.steps, seed=args.seed) print(" %-9.2f %.6f %4d %.2f" % ( beta, s["rec"], s["alive"], s["ppl"])) print(" -> beta=0 时编码器完全不被拉向码本,量化误差最大;beta 加大能压住") print(" 误差,但也把编码器往码本上拽,两头都不免费(本玩具上 1.0 最好)。") print() # ---------------- [B] 大码本的坍缩过程 ---------------- print("[B] 大码本 K=%d 的坍缩过程(beta=0.25,EMA,随机初始化)" % args.K) s0, h0 = run_vq(X, K=args.K, steps=args.steps, seed=args.seed, log_every=200) print(" 步数 重建MSE 活跃码 perplexity") for i in range(len(h0["step"])): print(" %5d %.6f %4d %.2f" % ( h0["step"][i], h0["recon"][i], h0["alive"][i], h0["ppl"][i])) print(" 最终:%d 个码字里只有 %d 个活着;perplexity %.2f 对应 %.2f bit," "而码本容量是 %.2f bit" % ( args.K, s0["alive"], s0["ppl"], np.log2(max(s0["ppl"], 1e-12)), np.log2(args.K))) print(" 最热的 1 个码占 %.3f 的使用量,最冷的一半码一共只占 %.4f" % (s0["top1"], s0["bottom_half"])) print() # ---------------- [C] 解药 ---------------- print("[C] 两种解药(同为 K=%d,%d 步)" % (args.K, args.steps)) p_init, _, _, _ = train_ae(X, 8, "continuous", steps=800, seed=args.seed) init_cb = kmeans(X @ p_init["We"], args.K, seed=args.seed, iters=20) s1, h1 = run_vq(X, K=args.K, steps=args.steps, seed=args.seed, init_cb=init_cb, log_every=200) s2, h2 = run_vq(X, K=args.K, steps=args.steps, seed=args.seed, restart=True, log_every=200) print(" 配置 重建MSE 活跃码 perplexity 有效bit") for tag, s in [("随机初始化", s0), ("k-means 初始化", s1), ("死码重采样重启", s2)]: print(" %-16s %.6f %4d/%4d %6.2f %.2f" % ( tag, s["rec"], s["alive"], args.K, s["ppl"], np.log2(max(s["ppl"], 1e-12)))) print() print(" 注:本玩具的真实类心只有 16 个,活跃码数的上界本来就远小于 K,") print(" 所以这里看的是「死码能不能被救活」,不是表达力真的翻了多少倍。") print() print("[D] 活跃码数随训练步数的变化(供配图)") print(" 随机初始化 :", list(zip(h0["step"], h0["alive"]))) print(" k-means 初始化:", list(zip(h1["step"], h1["alive"]))) print(" 死码重采样 :", list(zip(h2["step"], h2["alive"]))) if __name__ == "__main__": main() vq_core.py """vq_core.py —— 向量量化器(VQ)的最小实现,外加一个真跑得起来的训练实验。 无 torch 依赖,纯 numpy。本机解释器:/usr/local/bin/python3(3.10.5)。 运行: /usr/local/bin/python3 vq_core.py /usr/local/bin/python3 vq_core.py --steps 1200 --K 64 --beta 0.25 输出分五段: [A] 连续自编码器基线(无量化)——量化误差的下界 [B] VQ-AE 训练结果:重建误差 / 量化误差 / 码本使用率 / perplexity [C] 误差分解:总误差 = 子空间残差 + 量化误差(与 [A] 对拍) [D] straight-through 的梯度核验:真实有限差分 vs 直通梯度 [E] 码本更新方式对照(不更新 / 梯度 / EMA / EMA+死码重启) """ from __future__ import annotations import argparse import os import numpy as np HERE = os.path.dirname(os.path.abspath(__file__)) # ------------------------------------------------------------------ 数据 def cos_basis(patch: int = 8, n_freq: int = 3): """8x8 小块的低频余弦基,共 n_freq^2 = 9 个,已按行归一化。""" t = np.arange(patch) + 0.5 one_d = [np.ones(patch)] for k in range(1, n_freq): one_d.append(np.cos(np.pi * k * t / patch)) B = np.stack([np.outer(a, b).ravel() for a in one_d for b in one_d]) return B / np.linalg.norm(B, axis=1, keepdims=True) def make_patch_data(n=4096, patch=8, n_cluster=16, seed=0, center_scale=2.0, intra=0.35, noise=0.05): """合成一批 8x8 小块:16 个类心 + 类内低频连续变化 + 白噪声。 返回 X[N, 64]、类标、类心、基。类间方差远大于类内,所以码本有机会学到 「块类别」这种离散结构;白噪声部分不可压缩,构成误差地板。 """ rng = np.random.default_rng(seed) B = cos_basis(patch) # [9, 64] dim_b = B.shape[0] centers = center_scale * (rng.normal(size=(n_cluster, dim_b)) @ B) labels = rng.integers(0, n_cluster, size=n) X = centers[labels].copy() X += intra * (rng.normal(size=(n, dim_b)) @ B) # 类内连续变化 X += noise * rng.normal(size=(n, patch * patch)) # 不可压缩噪声 return X, labels, centers, B # ------------------------------------------------------------------ 量化 def quantize(z_e, codebook): """最近邻量化。z_e [N, d],codebook [K, d]。 返回 z_q[N, d]、索引 idx[N]、量化误差(每样本 d 维平方和)。 """ d2 = ((z_e[:, None, :] - codebook[None, :, :]) ** 2).sum(-1) # [N, K] idx = d2.argmin(axis=1) return codebook[idx], idx, d2[np.arange(len(idx)), idx] def kmeans(data, k, seed=0, iters=25): """Lloyd 迭代 + kmeans++ 初始化。空簇用「当前最差点」补齐。""" rng = np.random.default_rng(seed) n, d = data.shape k = min(k, n) centers = np.empty((k, d), dtype=data.dtype) centers[0] = data[rng.integers(n)] closest = ((data - centers[0]) ** 2).sum(1) for j in range(1, k): tot = closest.sum() if tot <= 0: centers[j] = data[rng.integers(n)] else: centers[j] = data[rng.choice(n, p=closest / tot)] closest = np.minimum(closest, ((data - centers[j]) ** 2).sum(1)) for _ in range(iters): _, assign, dist = quantize(data, centers) dist = dist.copy() for j in range(k): mask = assign == j if mask.any(): centers[j] = data[mask].mean(0) else: # 空簇:拿当前最差的点填 j_worst = int(np.argmax(dist)) centers[j] = data[j_worst] dist[j_worst] = -1.0 return centers def perplexity(counts): """码本 perplexity = exp(使用分布的熵),上界是码本大小 K。""" p = np.asarray(counts, dtype=np.float64) p = p / p.sum() nz = p[p > 0] return float(np.exp(-(nz * np.log(nz)).sum())) # ------------------------------------------------------------------ 训练 class Adam: """够用的 Adam,只处理一组 numpy 参数。""" def __init__(self, params, lr=0.02, b1=0.9, b2=0.999, eps=1e-8): self.p = params self.m = {k: np.zeros_like(v) for k, v in params.items()} self.v = {k: np.zeros_like(v) for k, v in params.items()} self.lr, self.b1, self.b2, self.eps = lr, b1, b2, eps self.t = 0 def step(self, grads): self.t += 1 for k, g in grads.items(): self.m[k] = self.b1 * self.m[k] + (1 - self.b1) * g self.v[k] = self.b2 * self.v[k] + (1 - self.b2) * (g * g) mhat = self.m[k] / (1 - self.b1 ** self.t) vhat = self.v[k] / (1 - self.b2 ** self.t) self.p[k] -= self.lr * mhat / (np.sqrt(vhat) + self.eps) def make_params(d_in, d_lat, seed=0): rng = np.random.default_rng(seed) return { "We": rng.normal(scale=0.1, size=(d_in, d_lat)), "Wd": rng.normal(scale=0.1, size=(d_lat, d_in)), "bd": np.zeros(d_in), } def evaluate(X, params, codebook, eval_idx): """在固定评估集上算重建 MSE、量化误差、活跃码数、perplexity。""" x = X[eval_idx] z_e = x @ params["We"] d_lat = params["We"].shape[1] if codebook is None: z_q, qerr, alive, ppl, counts = z_e, 0.0, 0, 0.0, None else: z_q, idx, qerr = quantize(z_e, codebook) counts = np.bincount(idx, minlength=len(codebook)).astype(np.float64) alive = int((counts > 0).sum()) ppl = perplexity(counts) x_hat = z_q @ params["Wd"] + params["bd"] return { "recon": float(((x_hat - x) ** 2).mean()), "quant": float(np.mean(qerr)) / d_lat if codebook is not None else 0.0, "alive": alive, "ppl": ppl, "counts": counts, } def train_ae(X, d_lat, codebook_mode, K=64, beta=0.25, steps=800, batch=512, lr=0.02, seed=0, decay=0.99, restart_dead=False, log_every=100, cb_lr=10.0, init_cb=None): """训练「线性编码器 + VQ + 线性解码器」。 codebook_mode: "continuous" —— 无量化,连续自编码器基线 "none" —— 码本完全不更新(只有 straight-through) "grad" —— 码本用 ||sg[z_e] - e||^2 的梯度更新 "ema" —— 码本用指数滑动平均更新(cluster_size + embed_avg) """ rng = np.random.default_rng(seed) eval_rng = np.random.default_rng(1000 + seed) eval_idx = eval_rng.choice(len(X), size=min(2048, len(X)), replace=False) n, d_in = X.shape params = make_params(d_in, d_lat, seed=seed) opt = Adam(params, lr=lr) has_vq = codebook_mode != "continuous" if has_vq: # 码本初始化:小方差(呼应 taming 里 uniform(-1/n_e, 1/n_e) 的量级) codebook = (rng.normal(scale=1.0 / K, size=(K, d_lat)) if init_cb is None else np.asarray(init_cb).copy()) init_codebook = codebook.copy() cluster_size = np.ones(K, dtype=np.float64) # EMA 用 embed_avg = codebook.copy() # EMA 用 else: codebook = init_codebook = None cluster_size = embed_avg = None hist = {"step": [], "recon": [], "quant": [], "alive": [], "ppl": []} for step in range(1, steps + 1): b = rng.choice(n, size=min(batch, n), replace=False) x = X[b] z_e = x @ params["We"] # [B, d] if has_vq: z_q, idx, _ = quantize(z_e, codebook) e_k = z_q x_hat = z_q @ params["Wd"] + params["bd"] # 反传:解码器对 z_q 的梯度,原封不动地当作对 z_e 的梯度 dxh = 2.0 * (x_hat - x) / (len(x) * d_in) # [B, d_in] dz_q = dxh @ params["Wd"].T # [B, d] dz_e = dz_q + beta * 2.0 * (z_e - e_k) / (len(x) * d_lat) opt.step({"We": x.T @ dz_e, "Wd": z_q.T @ dxh, "bd": dxh.sum(0)}) if codebook_mode == "grad": # d/de ||sg[z_e] - e||^2 = 2 (e - z_e)。注意码本是普通 SGD, # 不走 Adam,所以它的有效步长要单独调(见 [E] 的 cb_lr 扫描)。 g = 2.0 * (e_k - z_e) / (len(x) * d_lat) np.add.at(codebook, idx, -cb_lr * g) elif codebook_mode == "ema": onehot = np.zeros((len(x), K)) onehot[np.arange(len(x)), idx] = 1.0 cnt = onehot.sum(0) cluster_size = decay * cluster_size + (1 - decay) * cnt embed_avg = decay * embed_avg + (1 - decay) * (onehot.T @ z_e) tot = cluster_size.sum() smoothed = (cluster_size + 1e-5) / (tot + K * 1e-5) * tot codebook = embed_avg / smoothed[:, None] if restart_dead: dead = np.where(cluster_size < 1.0)[0] for j in dead: # 死码重采样为真实的编码器输出 codebook[j] = z_e[rng.integers(len(z_e))] cluster_size[j] = 1.0 embed_avg[j] = codebook[j] else: x_hat = z_e @ params["Wd"] + params["bd"] dxh = 2.0 * (x_hat - x) / (len(x) * d_in) opt.step({"We": x.T @ (dxh @ params["Wd"].T), "Wd": z_e.T @ dxh, "bd": dxh.sum(0)}) if step % log_every == 0 or step == steps: m = evaluate(X, params, codebook, eval_idx) hist["step"].append(step) for k in ("recon", "quant", "alive", "ppl"): hist[k].append(m[k]) shift = None if has_vq: shift = float(np.abs(codebook - init_codebook).mean()) return params, codebook, hist, shift # ------------------------------------------------------------------ 主流程 def main(): ap = argparse.ArgumentParser() ap.add_argument("--steps", type=int, default=800) ap.add_argument("--K", type=int, default=64) ap.add_argument("--d", type=int, default=8) ap.add_argument("--beta", type=float, default=0.25) ap.add_argument("--seed", type=int, default=0) args = ap.parse_args() X, labels, centers, B = make_patch_data(n=4096, seed=args.seed) n, d_in = X.shape print("数据: X%s 样本范数均值 %.4f" % (X.shape, np.linalg.norm(X, axis=1).mean())) print("码本 K=%d, 潜维度 d=%d, beta=%.2f" % (args.K, args.d, args.beta)) print() # ---------------- [A] 连续自编码器基线 ---------------- pc, _, hc, _ = train_ae(X, args.d, "continuous", steps=args.steps, seed=args.seed) recon_c = hc["recon"][-1] print("[A] 连续自编码器(无量化,最优线性重建)") print(" 重建 MSE/像素 = %.6f" % recon_c) print() # ---------------- [B] VQ-AE ---------------- pv, cb, hv, _ = train_ae(X, args.d, "ema", K=args.K, beta=args.beta, steps=args.steps, seed=args.seed) print("[B] VQ-AE(码本 EMA 更新, decay=0.99)训练曲线") print(" 步数 重建MSE 量化误差/d 活跃码 perplexity") for i in range(len(hv["step"])): print(" %5d %.6f %.6f %5d %.2f" % ( hv["step"][i], hv["recon"][i], hv["quant"][i], hv["alive"][i], hv["ppl"][i])) counts = np.bincount(quantize(X @ pv["We"], cb)[1], minlength=args.K).astype(np.float64) print(" 最终活跃码 %d / %d,perplexity %.2f(上界 %d)" % ( (counts > 0).sum(), args.K, perplexity(counts), args.K)) print() # ---------------- [C] 误差分解 ---------------- z_e = X @ pv["We"] z_q, _, qerr = quantize(z_e, cb) rec_vq = (((z_q @ pv["Wd"] + pv["bd"]) - X) ** 2).mean() rec_cont = (((z_e @ pv["Wd"] + pv["bd"]) - X) ** 2).mean() print("[C] 误差分解(同一组 VQ-AE 权重,只换「走不走量化」)") print(" 连续通路重建 MSE = %.6f (子空间残差,与 [A] 同量级)" % rec_cont) print(" 量化后重建 MSE = %.6f" % rec_vq) print(" 差值(量化引入) = %.6f (%.1f%%)" % ( rec_vq - rec_cont, 100 * (rec_vq - rec_cont) / rec_cont)) print(" [A] 连续基线 = %.6f" % recon_c) print(" 量化误差 E||z_e-e||^2/d = %.6f" % (qerr.mean() / args.d)) print() # ---------------- [D] straight-through 梯度核验 ---------------- # 用 [B] 训好的 VQ-AE,在真实工作点上做有限差分。 rng = np.random.default_rng(7) x0 = X[:512] z_e0 = x0 @ pv["We"] z_q0, idx0, _ = quantize(z_e0, cb) xh0 = z_q0 @ pv["Wd"] + pv["bd"] L0 = ((xh0 - x0) ** 2).mean() direction = rng.normal(size=args.d) direction /= np.linalg.norm(direction) print("[D] straight-through 梯度核验(512 样本,取 [B] 训好的权重)") print(" 有限差分:把 z_e 沿随机单位方向微扰 eps,看 L 怎么变") for eps in (1e-1, 1e-2, 1e-3, 1e-4): z_p = z_e0 + eps * direction z_qp, idxp, _ = quantize(z_p, cb) Lp = ((z_qp @ pv["Wd"] + pv["bd"] - x0) ** 2).mean() same = float((idxp == idx0).mean()) print(" eps=%.0e : dL/dz_e = %+.10f (argmin 保持不变的样本 %.4f)" % ( eps, (Lp - L0) / eps, same)) dxh = 2.0 * (xh0 - x0) / (len(x0) * d_in) dz_q = dxh @ pv["Wd"].T st_grad = float((dz_q * direction).mean()) st_norm = float(np.linalg.norm(dz_q, axis=1).mean()) commit = 2.0 * args.beta * (z_e0 - z_q0) / args.d commit_norm = float(np.linalg.norm(commit, axis=1).mean()) print(" 直通梯度 dL/dz_q 投影到同一方向 = %+.10f" % st_grad) print(" 直通梯度平均范数 ||dL/dz_q|| = %.6e" % st_norm) print(" commitment 项平均范数 = %.6e" % commit_norm) print(" -> 真实梯度恒为 0(argmin 是分段常数),直通梯度非零;") print(" 它是「假装量化是恒等映射」的替代品,不是真实梯度的估计。") print() # ---------------- [E] 码本更新方式对照 ---------------- print("[E] 码本更新方式对照(同为 %d 步,K=%d,beta=%.2f)" % ( args.steps, args.K, args.beta)) print(" 模式 重建MSE 活跃码 perplexity 码本位移") for mode, rst, cbl in [("none", False, 10.0), ("grad", False, 10.0), ("ema", False, 10.0), ("ema", True, 10.0)]: p, c, h, shift = train_ae(X, args.d, mode, K=args.K, beta=args.beta, steps=args.steps, seed=args.seed, restart_dead=rst, cb_lr=cbl) cnt = np.bincount(quantize(X @ p["We"], c)[1], minlength=args.K).astype(np.float64) tag = mode + ("+restart" if rst else "") print(" %-18s %.6f %5d %.2f %.6f" % ( tag, h["recon"][-1], int((cnt > 0).sum()), perplexity(cnt), shift)) print() print(" 附:grad 模式对码本步长极敏感(码本是普通 SGD,不走 Adam)") for cbl in (0.02, 0.1, 1.0, 10.0): p, c, h, _ = train_ae(X, args.d, "grad", K=args.K, beta=args.beta, steps=args.steps, seed=args.seed, cb_lr=cbl) cnt = np.bincount(quantize(X @ p["We"], c)[1], minlength=args.K).astype(np.float64) print(" cb_lr=%-6.2f 重建MSE %.6f 活跃码 %3d perplexity %.2f" % ( cbl, h["recon"][-1], int((cnt > 0).sum()), perplexity(cnt))) print() if __name__ == "__main__": main() perceptual_lab.py """perceptual_lab.py —— 为什么 VQGAN 不能只用 MSE:一个能算清的玩具。 VQ-VAE 用 MSE(或像素空间的似然)训练解码器。MSE 的最优解是条件均值, 而条件均值会把「同一 latent 对应多种合理细节」平均掉 —— 这就是重建发糊的 数学根源,不是玄学。VQGAN 的解法是换掉重建损失:改成特征空间距离(LPIPS) 外加一个对抗项。 本脚本用一个双模态玩具把这个机制算清楚: [A] 构造:同一个 latent 对应两种等概率的纹理(+p 和 -p) [B] MSE 最优 = 条件均值,纹理被平均掉;高频(梯度)能量只剩多少 [C] 换成非线性特征空间后,最优解跳到清晰模态;求阈值 alpha 并与解析式对拍 [D] 三种候选输出的三项指标对比:PSNR / 梯度能量保留 / 到最近模态的距离 注意:这里的特征映射 phi(x) = [x, alpha * |grad x|] 是我手工造的非线性特征, 用来演示「非线性」这一步为什么关键;真实的 LPIPS 用的是 VGG16 五层特征加 一层学习的 1x1 卷积,机制相同但权重是学出来的。 运行:/usr/local/bin/python3 perceptual_lab.py """ from __future__ import annotations import numpy as np PATCH = 8 AMP = 0.5 # 纹理幅度 SEED = 0 def gradient_magnitude(img): """前向差分后取绝对值:|grad x| 的展平向量(水平 56 个 + 垂直 56 个)。""" gx = np.diff(img, axis=1) gy = np.diff(img, axis=0) return np.concatenate([np.abs(gx).ravel(), np.abs(gy).ravel()]) def phi(img, alpha): """非线性特征映射:图像本身 + alpha 乘梯度幅值。""" return np.concatenate([img.ravel(), alpha * gradient_magnitude(img)]) def make_toy(): """低频内容 m + 等概率的 ±棋盘纹理 p。""" rng = np.random.default_rng(SEED) u = np.arange(PATCH) + 0.5 m = np.outer(np.cos(np.pi * u / PATCH), np.cos(np.pi * u / PATCH)) * 1.5 m += 0.3 * rng.normal(size=(PATCH, PATCH)) # 让内容不那么对称 ii, jj = np.meshgrid(np.arange(PATCH), np.arange(PATCH), indexing="ij") p = AMP * ((-1.0) ** (ii + jj)) return m, p def evaluate_candidate(m, p, c, alpha): """候选输出 x_hat = m + c * p,返回三项指标。""" x_hat = m + c * p x_plus, x_minus = m + p, m - p # MSE(对两种真值取期望) mse = 0.5 * (((x_hat - x_plus) ** 2).mean() + ((x_hat - x_minus) ** 2).mean()) # 特征空间距离(对两种真值取期望) f_hat = phi(x_hat, alpha) feat = 0.5 * (((f_hat - phi(x_plus, alpha)) ** 2).mean() + ((f_hat - phi(x_minus, alpha)) ** 2).mean()) # 梯度能量保留(相对真值的期望梯度能量) g_hat = np.concatenate([np.diff(x_hat, axis=1).ravel(), np.diff(x_hat, axis=0).ravel()]) g_true = 0.5 * (np.concatenate([np.diff(x_plus, axis=1).ravel(), np.diff(x_plus, axis=0).ravel()]) ** 2).sum() g_true += 0.5 * (np.concatenate([np.diff(x_minus, axis=1).ravel(), np.diff(x_minus, axis=0).ravel()]) ** 2).sum() grad_keep = float((g_hat ** 2).sum() / g_true) # 到最近模态的距离(越小说明落在数据流形上) to_mode = float(min(((x_hat - x_plus) ** 2).mean(), ((x_hat - x_minus) ** 2).mean())) return {"c": c, "mse": float(mse), "feat": float(feat), "grad_keep": grad_keep, "to_mode": to_mode} def best_c(m, p, alpha, grid=None): if grid is None: grid = np.linspace(-1.5, 1.5, 601) losses = np.array([evaluate_candidate(m, p, c, alpha)["feat"] for c in grid]) return float(grid[int(np.argmin(losses))]), grid, losses def main(): m, p = make_toy() signal_power = float(((0.5 * ((m + p) ** 2).mean() + 0.5 * ((m - p) ** 2).mean()))) print("[A] 玩具构造:8x8 patch,x = m + s * p,s = +1 / -1 各 0.5") print(" 低频内容 m: 范数 %.4f,梯度能量 %.4f" % ( np.linalg.norm(m), (np.concatenate([np.diff(m, axis=1).ravel(), np.diff(m, axis=0).ravel()]) ** 2).sum())) print(" 棋盘纹理 p: 范数 %.4f,梯度能量 %.4f(纹理占了 %.1f%% 的梯度能量)" % ( np.linalg.norm(p), (np.concatenate([np.diff(p, axis=1).ravel(), np.diff(p, axis=0).ravel()]) ** 2).sum(), 100 * (np.concatenate([np.diff(p, axis=1).ravel(), np.diff(p, axis=0).ravel()]) ** 2).sum() / (np.concatenate([np.diff(m + p, axis=1).ravel(), np.diff(m + p, axis=0).ravel()]) ** 2).sum())) print() print("[B] MSE 最优 = 条件均值(c=0,纹理被平均掉)") c_star, grid, losses = best_c(m, p, 0.0) r0 = evaluate_candidate(m, p, 0.0, 0.0) psnr0 = 10 * np.log10(signal_power / r0["mse"]) print(" MSE 最优的 c* = %.3f(理论值 0)" % c_star) print(" 重建 MSE = %.6f,PSNR = %.2f dB" % (r0["mse"], psnr0)) print(" 梯度能量保留 = %.4f(条件均值必然丢高频:E||grad x||^2 = " "||grad E[x]||^2 + E||grad(x - E[x])||^2)" % r0["grad_keep"]) print() print("[C] 换成特征空间后,最优解跳到清晰模态") print(" alpha 最优 c* 该点的 MSE PSNR(dB) 梯度能量保留") for alpha in (0.0, 0.2, 0.3, 0.378, 0.4, 0.5, 1.0, 2.0): c_a, _, _ = best_c(m, p, alpha) r = evaluate_candidate(m, p, c_a, alpha) psnr = 10 * np.log10(signal_power / r["mse"]) print(" %-8.3f %-10.3f %-13.6f %-10.2f %.4f" % ( alpha, c_a, r["mse"], psnr, r["grad_keep"])) # 解析阈值:alpha^2 > ||p||^2 / ||grad p||^2 gp = np.concatenate([np.diff(p, axis=1).ravel(), np.diff(p, axis=0).ravel()]) thresh = float(np.sqrt((p ** 2).sum() / (gp ** 2).sum())) print(" 粗略解析估计 alpha* = sqrt(||p||^2 / ||grad p||^2) = %.4f" "(假设纹理梯度远大于内容梯度)" % thresh) # 精确的「两个候选」比较:L(c) = (c^2+1)||p||^2 + alpha^2 * G(c), # 于是「清晰模态 L(1)」优于「模糊均值 L(0)」的条件是 alpha^2 > ||p||^2/(G(0)-G(1)) g0 = 0.5 * ((gradient_magnitude(m) - gradient_magnitude(m + p)) ** 2).sum() g0 += 0.5 * ((gradient_magnitude(m) - gradient_magnitude(m - p)) ** 2).sum() g1 = 0.5 * ((gradient_magnitude(m + p) - gradient_magnitude(m - p)) ** 2).sum() exact = float(np.sqrt((p ** 2).sum() / (g0 - g1))) print(" 精确阈值:G(0)=%.4f, G(1)=%.4f,L(1)<L(0) 要求 alpha > %.4f" % (g0, g1, exact)) print(" -> alpha 超过 %.3f 之后,「输出一个清晰模态」在特征损失上严格优于" "「输出模糊均值」;" % exact) print(" alpha 继续加大,最优 c 沿坐标轴继续往 ±1 移(不是跳变," "因为 |grad(m+c*p)| 关于 c 连续)。") print() print("[D] 四种候选输出的指标对比(PSNR 用信号能量 %.4f 作基准)" % signal_power) print(" 候选 MSE PSNR(dB) 梯度保留 到最近模态距离") c_feat, _, _ = best_c(m, p, 1.0) for tag, c in [("MSE 最优 (c=0)", 0.0), ("折中 (c=0.5)", 0.5), ("特征最优 (c=%.3f)" % c_feat, c_feat), ("清晰模态 (c=1)", 1.0)]: r = evaluate_candidate(m, p, c, 1.0) psnr = 10 * np.log10(signal_power / r["mse"]) print(" %-20s %.6f %-10.2f %-9.4f %.6f" % ( tag, r["mse"], psnr, r["grad_keep"], r["to_mode"])) print() r_mean = evaluate_candidate(m, p, 0.0, 1.0) r_sharp = evaluate_candidate(m, p, 1.0, 1.0) r_feat = evaluate_candidate(m, p, c_feat, 1.0) print(" PSNR 的绝对值很小,是因为这个玩具里纹理完全无法从 latent 预测,") print(" 要看的是相对差:") print(" · 特征最优解 (c=%.3f) 的 MSE 是模糊均值的 %.2f 倍(PSNR 低 %.2f dB);" % (c_feat, r_feat["mse"] / r_mean["mse"], -10 * np.log10(r_mean["mse"] / r_feat["mse"]))) print(" · 但它保留了 %.0f%% 的梯度能量,模糊均值只保留 %.0f%%;" % (100 * r_feat["grad_keep"], 100 * r_mean["grad_keep"])) print(" · 完全选一个模态 (c=1) 时 MSE 是均值的 %.1f 倍(PSNR 低 %.2f dB)," "梯度能量保留 %.0f%%,且到最近模态距离为 0 —— 它落在数据流形上。" % (r_sharp["mse"] / r_mean["mse"], -10 * np.log10(r_mean["mse"] / r_sharp["mse"]), 100 * r_sharp["grad_keep"])) print(" LPIPS/FID 站在后者一边,PSNR 站在前者一边 —— 这就是 VQGAN 之后") print(" 没人再用 PSNR 报告 tokenizer 重建质量的原因。") if __name__ == "__main__": main() codebook_size_law.py """codebook_size_law.py —— 码本容量 K 的收益到底有多大? 核心问题:把码本从 512 加到 16384,重建能好多少?直觉是「容量越大越好」, 但高维量化的经典结论(Zador 定理的渐近形式)说:最优量化失真随码本大小 只能按 K^(-2/d) 衰减,d 是码本向量维度。d=256 时指数只有 -1/128,也就是 K 翻一倍、失真只降 0.5%。 本脚本的实测方式:先训一个连续自编码器把潜分布固定住,再对潜变量做 k-means (k-means 就是「给定 K 个码字的最优最近邻量化器」的近似),在**留出集**上 测失真(避免用训练集测失真造成的过拟合假象)。 [A] 不同 d 下,量化失真 vs K 的双对数斜率,与 -2/d 对拍 [B] 量化误差什么时候降到「子空间残差」之下(继续加 K 的收益拐点) 运行:/usr/local/bin/python3 codebook_size_law.py """ from __future__ import annotations import argparse import numpy as np from vq_core import kmeans, make_patch_data, quantize, train_ae def distortion_curve(z_fit, z_test, k_list, seed=0, iters=20): """在 z_fit 上拟合 k-means,在 z_test 上测失真(留出集估计)。""" out = [] for k in k_list: c = kmeans(z_fit, k, seed=seed, iters=iters) _, _, d2 = quantize(z_test, c) out.append(float(d2.mean()) / z_fit.shape[1]) return np.array(out) def fit_slope(k_list, dist): """双对数线性拟合:log D = a * log K + b,返回斜率 a。""" lk = np.log(np.asarray(k_list, dtype=np.float64)) ld = np.log(np.asarray(dist, dtype=np.float64)) a, b = np.polyfit(lk, ld, 1) return float(a), float(b) def main(): ap = argparse.ArgumentParser() ap.add_argument("--steps", type=int, default=1500, help="连续自编码器训练步数") ap.add_argument("--seed", type=int, default=0) args = ap.parse_args() X, _, _, _ = make_patch_data(n=8192, seed=args.seed) d_in = X.shape[1] # 前一半拟合码本,后一半当留出集测失真 X_fit, X_test = X[:4096], X[4096:] k_list = [2, 4, 8, 16, 32, 64, 128, 256, 512] print("数据 X%s(前 4096 拟合码本,后 4096 留出评估)" % (X.shape,)) print() print("[A] 量化失真 D(K) = E||z - q(z)||^2 / d 的双对数斜率") print(" d 理论 -2/d 实测斜率(K>=16) D(8) D(512) 衰减倍数") results = {} for d in (2, 4, 8, 16): params, _, hist, _ = train_ae(X_fit, d, "continuous", steps=args.steps, seed=args.seed) z_fit = X_fit @ params["We"] z_test = X_test @ params["We"] dist = distortion_curve(z_fit, z_test, k_list, seed=args.seed) slope, _ = fit_slope(k_list[3:], dist[3:]) # 只拟合 K>=16 的渐近段 results[d] = {"dist": dist, "slope": slope, "params": params, "z_fit": z_fit, "z_test": z_test} print(" %-4d %-11.4f %-17.4f %.4e %.4e %.1fx" % ( d, -2.0 / d, slope, dist[2], dist[-1], dist[2] / dist[-1])) print() print(" K 从 8 加到 512(64 倍):d=2 失真降 %.1f 倍,d=16 只降 %.1f 倍。" % (results[2]["dist"][2] / results[2]["dist"][-1], results[16]["dist"][2] / results[16]["dist"][-1])) print(" d=2/4 的斜率与 -2/d 吻合;d>=8 时 K=512 已经逼近每个码字 8 个样本,") print(" 斜率被有限样本抬高(同样的码本在训练集上测会更陡),真值应更接近理论。") print() # ---------------- [B] 拐点:量化误差 vs 子空间残差 ---------------- params = results[8]["params"] z_fit, z_test = results[8]["z_fit"], results[8]["z_test"] rec_cont = (((z_test @ params["Wd"] + params["bd"]) - X_test) ** 2).mean() print("[B] 码本大到什么程度,量化误差才降到子空间残差之下(d=8,留出集)") print(" 连续通路重建 MSE(子空间残差) = %.6f" % rec_cont) print(" K 量化失真/d 量化后重建MSE 相对残差涨幅") for k in k_list: c = kmeans(z_fit, k, seed=args.seed, iters=20) z_q, _, d2 = quantize(z_test, c) rec = (((z_q @ params["Wd"] + params["bd"]) - X_test) ** 2).mean() print(" %-7d %.4e %.6f %+.1f%%" % ( k, d2.mean() / 8, rec, 100 * (rec - rec_cont) / rec_cont)) print() if __name__ == "__main__": main() make_figures.py """make_figures.py —— 画正文用到的五张示意图。 所有数字都现场重算(不读缓存),来源是同目录下的实验脚本: token_budget.py -> 图 1(序列长度与注意力规模) codebook_size_law.py-> 图 2(码本容量 K 的收益曲线) collapse_lab.py -> 图 3、图 4(坍缩过程与使用分布) perceptual_lab.py -> 图 5(MSE 最优 vs 特征空间最优) 运行:/usr/local/bin/python3 make_figures.py 输出:../figures/*.png """ from __future__ import annotations import os import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np import codebook_size_law as CSL import collapse_lab as CL import perceptual_lab as PL import token_budget as TB from vq_core import kmeans, make_patch_data, quantize, train_ae HERE = os.path.dirname(os.path.abspath(__file__)) FIGDIR = os.path.join(HERE, os.pardir, "figures") os.makedirs(FIGDIR, exist_ok=True) plt.rcParams["font.sans-serif"] = ["PingFang SC", "Heiti TC", "Arial Unicode MS"] plt.rcParams["axes.unicode_minus"] = False C_A = "#2E6DB4" # 蓝:主曲线 / 基线 C_B = "#C0504D" # 红:第二种配置 C_C = "#4FA96B" # 绿:第三种配置 / 好的一方 C_D = "#E08A2E" # 橙:强调 C_GREY = "#8C8C8C" def _save(fig, name): path = os.path.join(FIGDIR, name) fig.savefig(path, dpi=140, bbox_inches="tight", facecolor="white") plt.close(fig) print(" -> %s" % path) # ---------------------------------------------------------------- 图 1 def fig_token_budget(): sizes = np.array([128, 256, 512, 1024, 2048], dtype=float) fig, axes = plt.subplots(1, 2, figsize=(11.5, 4.3)) ax = axes[0] for f, col, mk in [(4, C_D, "o"), (8, C_A, "s"), (16, C_C, "^")]: n = (sizes / f) ** 2 ax.plot(sizes, n, color=col, marker=mk, lw=1.8, label=r"tokenizer $f$=%d" % f) ax.plot(sizes, sizes ** 2 * 3, color=C_GREY, marker="d", lw=1.8, ls="--", label="像素级 RGB") ax.set_xscale("log", base=2) ax.set_yscale("log") ax.set_xlabel("边长 H (像素)") ax.set_ylabel("序列长度(token 数)") ax.set_title("(a) 压缩率决定序列长度") ax.grid(alpha=0.3, which="both") ax.legend(fontsize=9) ax.annotate("H=256, f=16\n只有 256 个 token", xy=(256, 256), xytext=(300, 60), fontsize=9, color=C_C, arrowprops=dict(arrowstyle="->", color=C_C, lw=1.2)) ax = axes[1] labels = ["像素级\nRGB", "f=4", "f=8", "f=16"] vals = [(256 * 256 * 3) ** 2, ((256 / 4) ** 2) ** 2, ((256 / 8) ** 2) ** 2, ((256 / 16) ** 2) ** 2] colors = [C_GREY, C_D, C_A, C_C] bars = ax.bar(labels, vals, color=colors) ax.set_yscale("log") ax.set_ylabel(r"注意力矩阵规模 $n^2$") ax.set_title("(b) 256x256 图像的自注意力开销") ax.grid(alpha=0.3, axis="y") for b, v in zip(bars, vals): ax.text(b.get_x() + b.get_width() / 2, v * 1.6, "%.1e" % v, ha="center", fontsize=9) ax.annotate("省 5.9e5 倍", xy=(3, vals[3]), xytext=(2.55, 1e8), fontsize=9, color=C_C, arrowprops=dict(arrowstyle="->", color=C_C, lw=1.2)) fig.suptitle("图 1:为什么必须先做 tokenizer(数字来自 token_budget.py)", fontsize=11) fig.tight_layout() _save(fig, "token_budget.png") # ---------------------------------------------------------------- 图 2 def fig_codebook_law(): X, _, _, _ = make_patch_data(n=8192, seed=0) X_fit, X_test = X[:4096], X[4096:] k_list = [2, 4, 8, 16, 32, 64, 128, 256, 512] fig, ax = plt.subplots(figsize=(7.6, 5.2)) color_map = {2: C_C, 4: C_A, 8: C_D, 16: C_B} slopes = {} for d in (2, 4, 8, 16): params, _, _, _ = train_ae(X_fit, d, "continuous", steps=1500, seed=0) z_fit = X_fit @ params["We"] z_test = X_test @ params["We"] dist = CSL.distortion_curve(z_fit, z_test, k_list, seed=0) slope, _ = CSL.fit_slope(k_list[3:], dist[3:]) slopes[d] = slope ax.plot(k_list, dist, color=color_map[d], marker="o", lw=1.8, label=r"$d$=%d 实测斜率 %.2f" % (d, slope)) # 理论斜率 -2/d,锚定在 K=16 处 anchor = dist[3] theo = anchor * (np.array(k_list[3:], dtype=float) / 16.0) ** (-2.0 / d) ax.plot(k_list[3:], theo, color=color_map[d], lw=1.0, ls=":", alpha=0.75) ax.set_xscale("log", base=2) ax.set_yscale("log") ax.set_xlabel(r"码本大小 $K$") ax.set_ylabel(r"量化失真 $E\Vert z-q(z)\Vert^2/d$(留出集)") ax.set_title("图 2:码本容量 K 的收益被码本维度 d 卡死\n" "实线=实测,点线=理论斜率 -2/d(数字来自 codebook_size_law.py)", fontsize=11) ax.grid(alpha=0.3, which="both") ax.legend(fontsize=9) _save(fig, "codebook_size_law.png") return slopes # ---------------------------------------------------------------- 图 3、4 def fig_collapse(): X, _, _, _ = make_patch_data(n=4096, seed=0) K = CL.BIG_K s_rand, h_rand = CL.run_vq(X, K=K, steps=1200, seed=0, log_every=100) p_init, _, _, _ = train_ae(X, 8, "continuous", steps=800, seed=0) init_cb = kmeans(X @ p_init["We"], K, seed=0, iters=20) s_km, h_km = CL.run_vq(X, K=K, steps=1200, seed=0, init_cb=init_cb, log_every=100) s_rs, h_rs = CL.run_vq(X, K=K, steps=1200, seed=0, restart=True, log_every=100) fig, ax = plt.subplots(figsize=(7.6, 5.0)) for h, col, mk, tag in [(h_rand, C_B, "o", "随机初始化"), (h_km, C_A, "s", "k-means 初始化"), (h_rs, C_C, "^", "死码重采样重启")]: ax.plot(h["step"], h["alive"], color=col, marker=mk, lw=1.8, label=tag) ax.axhline(K, color=C_GREY, ls="--", lw=1.2) ax.text(max(h_rand["step"]) * 0.55, K * 1.08, "码本容量 K=%d(全活)" % K, color=C_GREY, fontsize=9) ax.set_xlabel("训练步数") ax.set_ylabel("活跃码数(至少被选中过一次)") ax.set_title("图 3:码本坍缩过程\n" "K=%d 的码本,随机初始化下最终只有 %d 个码活着(%.1f%%)" % (K, s_rand["alive"], 100.0 * s_rand["alive"] / K), fontsize=11) ax.grid(alpha=0.3) ax.legend(fontsize=9) ax.set_ylim(0, K * 1.25) _save(fig, "collapse_curve.png") fig, ax = plt.subplots(figsize=(7.6, 4.8)) for s, col, tag in [(s_rand, C_B, "随机初始化"), (s_rs, C_C, "死码重采样重启")]: frac = np.sort(s["counts"] / s["counts"].sum())[::-1] ax.plot(np.arange(1, len(frac) + 1), frac, color=col, lw=1.8, label="%s:%d 个活码,perplexity %.1f" % (tag, s["alive"], s["ppl"])) ax.set_yscale("log") ax.set_xlabel("码字按使用频次从高到低排序") ax.set_ylabel("使用占比") ax.set_title("图 4:使用分布的长尾与死码\n" "随机初始化有 %.1f%% 的码一次都没被用到" % (100.0 * (1 - s_rand["alive"] / K)), fontsize=11) ax.grid(alpha=0.3, which="both") ax.legend(fontsize=9) _save(fig, "usage_hist.png") return s_rand, s_km, s_rs # ---------------------------------------------------------------- 图 5 def fig_perceptual(): m, p = PL.make_toy() signal_power = float(0.5 * ((m + p) ** 2).mean() + 0.5 * ((m - p) ** 2).mean()) c_feat, _, _ = PL.best_c(m, p, 1.0) cands = [("真值样本 s=+1", m + p), ("真值样本 s=-1", m - p), ("MSE 最优 c=0", m), ("特征最优 c=%.2f" % c_feat, m + c_feat * p)] fig, axes = plt.subplots(2, 4, figsize=(12.0, 5.6), gridspec_kw={"height_ratios": [1.0, 0.92]}) vmin = min(a.min() for _, a in cands) vmax = max(a.max() for _, a in cands) for ax, (tag, img) in zip(axes[0], cands): im = ax.imshow(img, cmap="gray", vmin=vmin, vmax=vmax, interpolation="nearest") ax.set_title(tag, fontsize=10) ax.set_xticks([]) ax.set_yticks([]) axes[0][0].set_ylabel("8x8 patch", fontsize=9) rows = [("MSE 最优 c=0", 0.0), ("折中 c=0.5", 0.5), ("特征最优 c=%.2f" % c_feat, c_feat), ("清晰模态 c=1", 1.0)] psnrs, keeps = [], [] for tag, c in rows: r = PL.evaluate_candidate(m, p, c, 1.0) psnrs.append(10 * np.log10(signal_power / r["mse"])) keeps.append(100 * r["grad_keep"]) xpos = np.arange(len(rows)) ax = axes[1][0] bs = ax.bar(xpos, psnrs, color=[C_A, C_GREY, C_D, C_C]) ax.set_xticks(xpos) ax.set_xticklabels(["MSE\n最优", "折中", "特征\n最优", "清晰\n模态"], fontsize=8) ax.set_ylabel("PSNR (dB)") ax.set_title("(e) 像素误差:模糊均值最好", fontsize=10) ax.grid(alpha=0.3, axis="y") ax.set_ylim(0, max(psnrs) * 1.3) for b, v in zip(bs, psnrs): ax.text(b.get_x() + b.get_width() / 2, v + 0.08, "%.2f" % v, ha="center", fontsize=8) ax = axes[1][1] bs = ax.bar(xpos, keeps, color=[C_A, C_GREY, C_D, C_C]) ax.set_xticks(xpos) ax.set_xticklabels(["MSE\n最优", "折中", "特征\n最优", "清晰\n模态"], fontsize=8) ax.set_ylabel("梯度能量保留 (%)") ax.set_title("(f) 细节保留:清晰模态最好", fontsize=10) ax.grid(alpha=0.3, axis="y") ax.set_ylim(0, 118) for b, v in zip(bs, keeps): ax.text(b.get_x() + b.get_width() / 2, v + 1.5, "%.0f%%" % v, ha="center", fontsize=8) # 右侧两格合并成一条 alpha 扫描曲线(先把占位的两个空轴关掉) axes[1][2].axis("off") axes[1][3].axis("off") ax = plt.subplot2grid((2, 4), (1, 2), colspan=2) alphas = np.linspace(0.0, 2.0, 41) cs = [PL.best_c(m, p, a)[0] for a in alphas] ax.plot(alphas, np.abs(cs), color=C_A, lw=1.8) ax.axhline(1.0, color=C_GREY, ls="--", lw=1.0) ax.axhline(0.0, color=C_GREY, ls="--", lw=1.0) ax.set_xlabel(r"特征里高频项的权重 $\alpha$") ax.set_ylabel("最优输出的纹理系数 |c|") ax.set_title(r"(g) $\alpha$ 越大,最优解越靠近清晰模态", fontsize=10) ax.grid(alpha=0.3) ax.set_ylim(-0.05, 1.15) ax.set_xticks([0.0, 0.5, 1.0, 1.5, 2.0]) fig.suptitle("图 5:MSE 最优必然糊,特征空间最优不糊" "(数字来自 perceptual_lab.py)", fontsize=11) fig.tight_layout() _save(fig, "perceptual_tradeoff.png") def main(): print("画图:所有数字现场重算,不读缓存") print("[1/5] token_budget.png") fig_token_budget() print("[2/5] codebook_size_law.png") slopes = fig_codebook_law() print(" 实测斜率:", {k: round(v, 3) for k, v in slopes.items()}) print("[3/5] collapse_curve.png + usage_hist.png") s_rand, s_km, s_rs = fig_collapse() print(" 随机初始化: alive=%d ppl=%.2f rec=%.6f" % ( s_rand["alive"], s_rand["ppl"], s_rand["rec"])) print(" k-means 初始化: alive=%d ppl=%.2f rec=%.6f" % ( s_km["alive"], s_km["ppl"], s_km["rec"])) print(" 死码重采样: alive=%d ppl=%.2f rec=%.6f" % ( s_rs["alive"], s_rs["ppl"], s_rs["rec"])) print("[5/5] perceptual_tradeoff.png") fig_perceptual() print("完成,图在 %s" % os.path.abspath(FIGDIR)) if __name__ == "__main__": main()
2026年10月02日
0 阅读
0 评论
0 点赞
2026-10-02
AIGC 每日速读|2026-10-02|NVIDIA砍掉视觉编码器,PixelUMM图像视频同模
今日 AIGC 论文速览 今日共 10 篇 · 统一理解生成与三维细节 2 篇 · 推理式与可控视频生成 2 篇 · 音视频联合生成 1 篇 · 视频推理加速 4 篇 · 视频文字编辑评测 1 篇 重点论文标题列表 PixelUMM(NVIDIA):像素直进图像视频同模型 BTC3D(香港城市大学(东莞)):分块条件补3D细节 ThinkV2V(港中大):先推理再改视频 LIFT(加州大学圣迭戈分校):只画终帧也能控未来 StereoBind(西电):声源移动立体声跟着走 今日论文速览 1. PixelUMM:像素直进图像视频同模型 PixelUMM: Encoder-Free Unified Image and Video Understanding and Generation | NVIDIA;滑铁卢大学 | arXiv:2609.38597 关键词:统一多模态, 像素空间, 图像生成, 视频生成, 视觉理解 前序问题:统一理解与生成模型通常给同一张图准备两套表示:ViT 特征服务理解,VAE latent 服务生成。这不只让视觉上下文变长,也把成熟的视觉语言预训练与生成接口绑在两条管线上。图像还能直接切 patch,视频却要同时兼容理解侧的采样帧与生成侧的连续时空块,统一接口更难。 本文贡献:PixelUMM 取消预训练视觉编码器和 VAE,把图像变成空间 patch、视频变成时空 tubelet,都只用一层线性投影接入 8B Mixture-of-Transformers。理解走自回归文本预测,生成在原始像素上做 flow matching;共享 attention 让干净条件与带噪目标直接交互。论文还系统比较 patch/tubelet 大小、线性头与卷积解码头,发现卷积头能压掉格子伪影。 实验效果:不使用视觉 encoder/tokenizer 时,PixelUMM 在图像理解 MMMU 得 41.67、视频理解 MVBench 得 70.53;图像生成 GenEval 原始 prompt 为 0.77,用 BAGEL-long 重写后为 0.83,DPG-Bench 为 85.74。视频生成 VBench 总分 83.24,处于统一模型可比区间,但不是各榜第一。 批判点评:“统一”仍有边界:理解与生成使用两套 Transformer expert,只共享 attention,并非一组参数同时做两件事。原始像素 token 让更小 patch 和 tubelet 虽然降低 loss,却直接抬高序列长度与训练成本。作者也提醒各模型训练数据不同,当前表格不能证明 encoder-free 架构本身更优;MMMU 41.67 与主流 8B VLM 仍有明显差距。 2. BTC3D:分块条件补3D细节 BTC3D: Blended Tile Conditioning for Detail-Enhancing Image-to-3D Generation | 香港城市大学(东莞);SB Intuitions;巴塞罗那自治大学 | arXiv:2609.39709 关键词:图像生成3D, 免训练, 分块条件, 细节保真, 流匹配 前序问题:扩散式 image-to-3D 已经能出像样模型,可一旦输入图细节丰富,纹理就糊成一团。主流做法把整张图压成一个全局条件特征,空间信息在这一步被压掉,作者称之为 detail attenuation;而想补细节通常要重训或微调大模型,对 3D 管线代价过高。 本文贡献:BTC3D 是推理时即插即用的免训练框架。它先验证 image-to-3D 模型存在「图像特征可加性」:把局部图块的特征相加,仍能得到语义一致且保留局部细节的 3D 输出。据此把输入图切成若干 tile,各自提取局部条件并按前景权重聚合成 blended tile embedding;再配一个 dynamic conditioning 调度,在扩散后期低噪声阶段逐步加大 tile 条件的权重,早期仍交给全局条件主导结构。 实验效果:在 3D-Arena 与 Toys4K 上,把 BTC3D 接到 TRELLIS、TRELLIS.2、Hunyuan3D-v2.1 上都提升了纹理与视觉保真,且保持全局结构一致。20 名参与者、400 次两两偏好测试中,接入 BTC3D 的 TRELLIS.2 拿到 67.25% 选票(269 票),基线只有 32.75%(131 票),纹理密集的样本优势最明显。 批判点评:收益主要靠人眼偏好而非几何或纹理指标说话,67.25% 也是相对自家基线的增幅,缺少跨方法横评;论文同时承认纹理稀疏的样本上两者偏好接近持平,说明方法吃的是「有没有细节可补」。另外 tile 切分与融合系数 α、β 都是人工超参,动态调度只按扩散时间步推进,不看当前生成内容。 3. ThinkV2V:先推理再改视频 ThinkV2V: Unleashing the Reasoning Capability of MLLMs for Instruction-Guided Video Editing | 香港中文大学;字节跳动;浙江大学;俄亥俄州立大学 | arXiv:2609.38541 关键词:视频编辑, 多模态大模型, 显式推理, DiT, 指令跟随 前序问题:现有指令视频编辑器多把 MLLM 当语义编码器,直接把 prompt 压成条件。碰到“把会被雨淋湿的东西收起来”这类不直接点名目标的指令,模型必须先理解因果与对象关系,再决定改哪里;只抓关键词往往画面能改,却改错对象或违背真实意图。 本文贡献:ThinkV2V 让 MLLM 先显式分析源视频和指令,再通过 learnable-query connector 把推理结果送给多条件 DiT。训练采用从直接编辑到推理密集样本的 progressive curriculum;推理时生成多份候选编辑 prompt,迭代反思并选择最可靠的一份。团队同时整理 ThinkV2V-150K 与专门考隐式意图、因果编辑的 ThinkV2V-Bench。 实验效果:5B 编辑器在 ThinkV2V-Bench 上由 Seed-1.6-VL 评测的 Overall 为 2.72,高于 OpenVE-Edit 的 2.09 和 14B DITTO 的 2.02;Gemini-2.5-Pro 评测下为 2.61,也高于 ICVE 的 2.52。模型在 32 卡上分三阶段训练,最终 720p 推理,以小于多条 10B–14B 基线的规模拿到最佳综合分。 批判点评:核心结论依赖 Seed-1.6-VL 与 Gemini-2.5-Pro 充当裁判,复杂编辑的自动分数仍可能偏向更会“讲意图”的输出。Inference-Time Thinking Scaling 需要多候选生成、反思和筛选,论文没有把新增延迟与算力单独量化。150K 数据由既有模型与筛选流程构建,推理能力有多少来自 MLLM、多少来自数据课程,还缺更严格的因果拆分。 4. LIFT:只画终帧也能控未来 LIFT: Layout-In-Future Video Generation under Large Viewpoint Change via On-Policy Self-Distillation | 加州大学圣迭戈分校;弗吉尼亚大学;Meta;Amazon;Lambda | arXiv:2609.38146 关键词:视频生成, 相机控制, 未来布局, on-policy self-distillation, 图生视频 前序问题:相机大幅移动时,首帧之外会露出全新区域。只给相机轨迹能控制“镜头往哪走”,却不能指定新区域里出现什么;逐帧框轨迹虽然能控布局,标注和交互成本又太高。真正实用的接口应允许用户只在未来关键帧画几个框,同时让中间过程自然长出来。 本文贡献:LIFT 把最后一帧的 bounding box 与局部 prompt 编成稀疏 layout latent,与首帧和噪声视频一起送入 DiT,相机 token 另行注入。训练时用拥有逐帧 dense layout 的冻结教师,在学生自己访问的 ODE 状态上做 on-policy self-distillation,把 dense 轨迹知识蒸馏给只看终帧布局的学生;同一个学生交替学习有布局与仅相机两种模式。 实验效果:LIFT 1.3B 在终帧布局控制上 mIoU 0.51,高于需要逐帧框的 MagicMotion 0.41;同时相机 RotErr 2.97、TransErr 0.59,视频 FVD 99.35。OPSD 只用 500×16=8000 次样本更新,直接或 dense-to-sparse SFT 都要 4000×32=12.8 万次,训练样本量缩到 1/16。 批判点评:布局接口仍只是 2D 框加局部文字,没有深度、朝向、遮挡和实例关系;相机一转,两个框在 3D 中究竟谁挡谁并未显式建模。所谓 1/16 训练量比较的是作者设定的更新预算,并不等于端到端数据与教师成本也缩到 1/16;OPSD 还依赖一个先用 dense layout 训练好的教师。 5. StereoBind:声源移动立体声跟着走 Here the World in Stereo: Learning Dynamic Spatial Correspondence for Immersive Joint Video-Audio Generation | 西安电子科技大学;vivo 蓝图影像实验室 | arXiv:2609.38748 关键词:音视频联合生成, 立体声, 空间音频, 运动轨迹, 跨模态对应 前序问题:联合音视频模型已经能做到“画面里有什么就听到什么”和大致同步,却很少管声音来自哪里。AR/VR 中声源从左往右移动时,左右声道也应连续变化;只用文本或整段视频条件生成声音,往往知道是汽车声,却无法让声像跟随汽车轨迹。 本文贡献:StereoBind 把可见声源的 motion track 作为共同坐标系:Visual Motion Binding 绑定视觉运动与音频 token,Spatial Track Encoder 提供绝对位置,Residual Track RoPE 编码相对位移。团队用真实视频定位跟踪加合成立体声、再配合可控合成场景,构建 StereoWorld-29K,并用 StereoWorldBench 单独评测动态空间对应。 实验效果:在 StereoWorldBench 上,StereoBind 的 ILD-W 为 2.754、SELD-Acc 0.757、SMR-Err 14.86、AST-Ang 0.465、AST-Cal 0.314,五项空间指标都优于表中 Ovi、MiniMax-H3、LTX-2.5/2.3 及两条 video-to-spatial-audio 基线;视觉 Subject Consistency 0.955、Motion Smoothness 0.996,整体画面质量没有明显牺牲。 批判点评:当前条件是一条主要声源轨迹,训练和评测强调水平声像移动;多声源竞争、声源与听者同时运动、遮挡和完整 3D 声场仍被列为未来工作。空间指标全面领先不代表总体音质全面领先,Subject Consistency 0.955 低于 MiniMax-H3 的 0.968,Audio PQ 6.86 也只是略高于 6.81。 6. ReCaVSR:1080p单卡跑21.2FPS ReCaVSR: One-Step Streaming Diffusion Video Super-Resolution with Recycled Latents and Learned Cache Routing | 中国科学技术大学 | arXiv:2609.37831 关键词:视频超分辨率, 流式扩散, KV Cache, 一步生成, 实时推理 前序问题:扩散式视频超分画质好,却很难满足直播级延迟。现有流式方案要么多步去噪,要么每层都保留完整历史 KV;但视频超分当前帧已有低清观测,上一块超分 latent 也带着局部时序信息,继续为所有层保存同样长的历史既占显存又拖慢注意力。 本文贡献:ReCaVSR 基于 Wan2.2 做一步流式超分,把上一块生成的 SR latent 直接回收为下一块条件,再为每个 DiT 层学习不同的历史缓存范围,训练完导出固定路由。生成器用 sequential self-rollout 对齐真实因果推理;Multi-Scope Query 判别器同时看全局、空间窗口和时间 tube;LR-conditioned FlashDecoder 在解码时再次利用低清观测。 实验效果:单张 A100-80GB 输出 1080×1920 时达到 21.20 FPS、峰值显存 15.16GB,相对 FlashVSR Tiny 快 2.72 倍、少用 38.0% 显存,首个完整 RGB 输出的模型时间为 0.982 秒。LongVSR60 上 MUSIQ 57.89、CLIP-IQA 0.4932、DOVER 0.6075,后三项都高于列出的流式基线。 批判点评:它并非所有重建指标都第一:REDS30 的 PSNR 21.67 明显低于 RealViformer 23.32,说明生成式纹理仍会牺牲逐像素保真。缓存路由在训练后固定,遇到运动强度突变或硬切镜不能自适应;论文的长视频测试还是单段连续镜头,跨场景时旧 latent 和 KV 会不会污染新镜头尚未验证。 7. PARK:稀疏注意力快1.76倍 PARK: Accurate Block Retrieval for Sparse Attention in Video Diffusion Transformers | 华南理工大学 | arXiv:2609.38978 关键词:视频生成, 稀疏注意力, DiT, block retrieval, 免训练加速 前序问题:视频 DiT 的稀疏注意力通常先把 query 与 key 分块求均值,再按块检索。这个近似有两重错位:Softmax 前把多个 query 平均会抹掉各自偏好的 key;在原始 key 空间做欧氏聚类,也不保证同一簇的 key 在当前 query 下拥有相似 QK 分数。检索一旦选错块,要么掉画质,要么为了兜底算太多。 本文贡献:PARK 保留块内每一个原始 query,分别对 key block 做归一化后再平均分布,避免 query-centroid 代替整块偏好;key 侧则从当前 query 推导度量,先变换 key 再做 K-means,让欧氏距离对应 QK-score 误差。方法完全免训练,并把 PAMA 检索写成融合 GPU kernel,把准确率收益留住而不让检索开销反噬。 实验效果:在 Wan2.1-1.3B 上,PARK 将单次测试从 745 秒降到 423 秒,端到端 1.76 倍加速,同时 PSNR 27.52、SSIM 0.908、LPIPS 0.177,优于同表其他稀疏基线。跨 Wan2.2-14B、Wan2.1-14B、Wan2.1-1.3B 和 HunyuanVideo 四个模型,相对 dense 的加速范围为 1.53–1.76 倍。 批判点评:收益建立在四个视频扩散模型和固定检索配置上,作者明确没有验证 LLM 或其他多模态任务。主表中 Wan2.2 的 Aesthetic Quality 65.29% 仍低于 dense 的 65.56%,Wan2.1-1.3B 的 Background Consistency 96.87% 也低于 dense 的 96.99%;它更像高质量近似,而不是无损替换。 8. UnStep:免训练把视频推到50FPS UnStep: Training-Free Acceleration of Causal Video Diffusion with Fewer Steps Than Distillation | Meta 超级智能实验室 | arXiv:2609.32518 关键词:视频生成, 因果扩散, 免训练加速, KV Cache, 推理优化 前序问题:四步因果蒸馏已经比 50 步双向教师快,但 Self Forcing 在 H100 上仍只有 17 FPS。步数继续砍到一两步会掉质量,历史 KV 随视频变长继续膨胀;同时当 DiT 已经很少步,RoPE、cache 索引、attention 调度和 VAE 解码这些过去被遮住的系统开销反而成了瓶颈。 本文贡献:UnStep 是不改权重的推理 wrapper:后续块从四步减到两步,只保留 sink 加最近窗口;把已生成 latent 重新加到近干净噪声,再复用原有 clean-cache pass 做一次修正,并对 attention 的 V/O 投影做截断 SVD 抵消少步误差。系统侧再融合 RoPE、缓存系数、固定长度 attention,并把 VAE 改成 FP16 与 channels-last 卷积。 实验效果:同一 Self Forcing checkpoint 在 H100 上从 17.0 FPS 提到 49.8 FPS,VBench Total 反而从 84.31 到 84.56;相同 wrapper 套在 Causal Forcing 和 LongLive 1.0 上也都约 49.8 FPS。若推理时改成一步,H100 达 59.6 FPS;GB200 最快报告 77 FPS,全程没有重训。 批判点评:所有超参数只在 Self Forcing、5 秒 T2V、单张 H100 这一套设置上调过,迁移结果证明“能用”,不代表每个 checkpoint、GPU、任务和长度都最优。50 FPS 的成绩同时叠加算法减步与大量 H100 专用 kernel 优化,不能全部归功于新的生成算法;而训练自由的约束也限制了更激进稀疏或量化后的质量恢复。 9. DeCoPrune:剪85%缓存仍保上下文 DeCoPrune: Efficient KV-Cache Pruning for Autoregressive Video Diffusion via Denoising Consistency | 南洋理工大学;清华大学;普林斯顿大学 | arXiv:2609.39096 关键词:自回归视频, KV Cache, 训练自由剪枝, 长上下文, 一致性 前序问题:自回归视频扩散把每段历史都塞进 KV Cache,时间越长,显存和 attention 成本越大。固定滑窗只看远近,attention 权重或相似度剪枝又不直接回答“这个历史块是否提供了不可替代的信息”,所以容易把稍后 prompt 才会重新提到的物体细节一起删掉。 本文贡献:DeCoPrune 用同一条去噪轨迹里的中间 probe 与最终预测做差:没有历史支持时难以稳定去噪、step-to-final discrepancy 较大的 token 被判为更需要上下文,因此保留高差异 token,再物理压紧历史 KV。论文同时构建 CMBench,用一分钟上下文中的物体取放、场景恢复等续写任务专测“前文记没记住”,并提供按 attention head 类型分组的 HS 版本。 实验效果:主设置剪掉 85.43% 历史 KV,续写 FPS 从 FullKV 的 1.568 提到 6.489,快 4.14 倍;CMBench DINO 从 0.6803 只降到 0.6701。head-specialized 版本在 86.19% 剪枝率下 DINO 为 0.6783,几乎追平 FullKV,且 VBench 的图像质量项没有出现系统性崩塌。 批判点评:“剪掉 85%”并不等于显存有硬上限,作者明确承认剩余缓存仍随视频长度线性增长。Denoising consistency 是 future-agnostic 的:某个细节此刻很好去噪、因此被删掉,但未来指令恰好再次询问它时仍可能失忆。主干又集中在 LingBot World v2,其他开源模型的分钟级 FullKV 本身就不稳定,方法与 backbone 上限难完全拆开。 10. ViTeX-Bench:387段视频检验动态改字 ViTeX-Bench: Benchmarking High-Fidelity Video Scene Text Editing | 得州农工大学 | arXiv:2609.40356 关键词:视频编辑, 场景文字, 评测基准, 时序一致性, 编辑局部性 前序问题:视频里的招牌、球衣或屏幕改字,要同时满足字符串正确、跨帧不漂、背景不被改。通用视频编辑指标可能把“画面很稳但根本没改字”的结果评得很高,静态 OCR 又看不见闪烁与字符漂移,现有公开数据也缺少真实视频的成对编辑参考。 本文贡献:ViTeX-Bench 收集 387 段 720p、120 帧、24 FPS 真实视频,其中 230 段提供人工复核的成对训练编辑,157 段冻结测试。评测拆成文字正确性、视觉/时序质量、编辑局部性三轴共 13 个指标,并用 OCR 校准、人评相关性与 Pareto 前沿解释取舍;配套的 ViTeX-Edit-14B 用随运动对齐的 glyph-video 条件微调 VACE。 实验效果:八条基线中,ViTeX-Edit-14B 在 video-native 编辑器里 CharAcc 最高,为 0.688,text-crop Warp 1.53 也最低;相较 VideoPainter 的 CharAcc 0.619 提升 0.069。但逐帧 FLUX-Text 的 CharAcc 仍达 0.737,代价是 Warpc 13.01、字符跨帧严重漂。人评与 OCR 方法排名的 Spearman 相关系数为 0.95。 批判点评:这个基准把“漂亮、正确、稳定、局部”拆开是优点,也暴露了参考模型并非全面第一:ViTeX-Edit 的 SeqAcc 0.341 低于 FLUX-Text 0.528。数据以清晰、局部、拉丁文字为主;非拉丁切片里除 AnyText2 外其余编辑器 CharAcc 都低于 0.19,严重运动、曲面文字、密集排版的覆盖仍不足。 趋势观察 加速的战场从「少走几步」挪到了「每层只记该记的」 ReCaVSR 给每个 DiT 层学一个缓存跨度,PARK 修正稀疏注意力的块检索误差,UnStep 把两步采样、限窗 KV 与 kernel 优化叠进同一个免训练 wrapper,DeCoPrune 则用去噪一致性挑出真正依赖历史的 token,四篇都指向同一件事:视频模型已进入少步时代,下一轮效率竞争发生在缓存、检索与运行时细节。另一头,PixelUMM 去掉视觉编码器让像素直进统一模型,说明「统一」正在从拼模块转向砍模块。 人工智能炼丹君 整理 | 2026-10-02 更多 AIGC 论文解读,关注微信公众号「人工智能炼丹君」 每日更新 · 论文精选 · 深度解读 · 技术脉络 微信搜索 人工智能炼丹君 或扫描下方二维码关注
2026年10月02日
1 阅读
0 评论
1 点赞
2026-10-02
AIGC 基本功|KV Cache 与自回归视频生成-KVCache
KV Cache 与自回归视频生成 所属方向:推理加速 | 难度:进阶 | 前置知识:自注意力机制的计算与显存账本(attention_basics)、性能建模与 Profiling(performance_profiling) 关键词:KV Cache、显存带宽、自回归生成、因果注意力、缓存驱逐、PagedAttention 01. 为什么需要它 先给三个数,都是本篇附录按明确假设计算的账本;它们不是实际部署峰值显存。 第一个数:一个 token 的 KV cache 是 128 KiB。 按 Meta Llama 模型配置 中 Llama-3-8B 的结构算(32 层、8 个 KV 头、每头 128 维、BF16 两个字节):每个 token 要缓存 $2 \times 32 \times 8 \times 128 \times 2 = 131072$ 字节,正好 128 KiB。所以一条 8192 token 的序列,KV cache 恰好是 1.00 GiB。这个整数不是巧合,是配置凑出来的——记住它,后面所有账都从它出发。 第二个数:batch=32 时,KV cache 是权重的 2.14 倍。 权重 14.96 GiB,KV cache $32 \times 1.00 = 32.00$ GiB,合计 46.96 GiB。batch 加到 64,合计 78.96 GiB——已经接近本文假设的 80 GiB 总预算,其中约 81% 是 KV cache;还未计入激活、临时空间与框架开销。实际 A100 型号标称 80 GB,可用字节应以设备查询为准,不能把理想 80 GiB 预算当作部署保证。你以为你在被权重压垮,其实在被缓存压垮。 第三个数:一组视频假设会产生 52.7 GiB 的 KV。 假设沿用上面的 Llama 结构、缓存全部历史、不做额外 latent patch 化:5 秒 24fps、时间维 4 倍压缩、空间 8×8 压缩,是 30 帧 × 14400 token = 432000 token,按同样的 128 KiB/token 算就是 52.73 GiB——还没算权重,一条视频就吃掉大半张卡。10 秒 105.47 GiB,30 秒 316.41 GiB。这是展示量级的假想配置,不是真实视频模型的统一用量;VAE、patch、窗口与模型层数都能改变它。 玩具 decoder 的全量重算与缓存增量解码还展示了一点:计算量减少不等于同倍数加速。修正了密集注意力 FLOPs 计数、输出头计数和最后一次无用 decode 后,当前代码解析比为 105.61 倍,本次 NumPy 墙钟为 10.29 倍。Python 调度、临时数组、重复 KV 头和小矩阵效率都会影响耗时,单靠这个比值不能认定差额全来自内存带宽。 02. 最小可用理解 三句话: 机制:因果注意力里,每层位置 $j$ 的 key 和 value 由固定前缀 $1\ldots j$ 的隐藏状态决定,跟它后面来了什么 token 无关。所以生成第 $t$ 个 token 时,前 $t-1$ 个位置的 K、V 和上一步完全一样——把它们留在显存里,每步只算新 token 自己的那一个 query、一对 K/V。用空间换时间,空间就是 KV cache。 成本:省的算力是真实的(每步从 $O(S^2)$ 降到 $O(S)$,整段生成从 $O(S^3)$ 降到 $O(S^2)$),长上下文下注意力常受访存约束。理想融合的单 query 注意力有 $I=2g/p$,但实际瓶颈还取决于 batch、kernel、缓存命中、并行与硬件,decode 仍有投影和 FFN 运算。同时缓存自己按 $2 L n_{\text{kv}} d p$ 字节每 token 线性膨胀,长上下文和视频场景下反过来成了显存的主宰。 效果与代价:文本场景它是推理加速的第一功臣;视频自回归场景 token 数大两个数量级,于是问题从「要不要缓存」变成「怎么让缓存装得下」——分页管理、GQA、缓存量化、块级因果,全是被这个量级逼出来的。 03. 数学推导 3.1 因果注意力里,什么是死的 自注意力一步的计算是 $$\mathrm{Attn}(Q, K, V) = \mathrm{softmax}\left( \frac{Q K^{\top}}{\sqrt{d}} + M \right) V$$ 其中 $M$ 是因果掩码,$M_{ij} = 0$($j \le i$)或 $-\infty$($j > i$)。关键在 $Q$、$K$、$V$ 是怎么来的:第 $i$ 个位置的 query、key、value 是 $$q_i = W_q x_i, \quad k_i = W_k x_i, \quad v_i = W_v x_i$$ 这里的 $x_i$ 是该层输入隐藏状态,不是只含第 $i$ 个 token 的原始 embedding。在多层 decoder 中,它已经聚合了前面位置的信息。正确的论证是逐层归纳:固定权重、位置、条件与推理随机性后,因果掩码保证历史位置不能看未来;追加 token 不会改变旧位置的隐藏状态,所以其 K/V 可复用。此处 $t$ 表示正在处理的输入位置,$j<t$ 的缓存已存在,$q_t,k_t,v_t$ 是本步新计算的量。改变前缀、RoPE 位置、条件、模型权重或噪声等级,都可能使旧缓存失效。 所以增量解码的正确姿势是: $$o_t = \sum_{j=1}^{t} \mathrm{softmax}_j \left( \frac{q_t k_j^{\top}}{\sqrt{d}} \right) v_j$$ 注意这个式子里只有 $q_t$、$k_t$、$v_t$ 是新的,$k_j$、$v_j$($j < t$)全部从缓存里读。softmax 的分母也只在这一行上归一化——不需要重算别的行,因为别的行的输出早就有了,而且以后也不会变。 这里有一个值得停一下的对比:训练时我们并行算所有位置,$Q K^{\top}$ 是一个 $S \times S$ 的矩阵;decode 时一次只有一个 query,$Q K^{\top}$ 退化成一个 $1 \times t$ 的向量。同一个算子,在两个阶段里形状完全不同,这让长 prefill 通常更容易利用矩阵计算资源,而单 token decode 的注意力通常更容易受带宽或并行度限制,仍需实测确认。 顺带回答一个常见疑问:为什么缓存的是 K 和 V,而不是 Q?因为它们的生命周期不同。$q_t$ 在这一步算完、和缓存做完内积之后就没用了——下一个 token 不会再来问它;而 $k_j$、$v_j$ 是「将来所有 query 都要来查一遍」的公共数据,未来第 $t+1$、$t+2$ 步的注意力都要用。缓存的对象必须是「写入后不再变、且会被反复读」的东西,这正好是 3.4 节那个算术强度问题的另一半来源:省下的是重算 K/V 的算力,付出的是每步把这块只读数据整个搬一遍的带宽。 3.2 省了多少算力 先算不用缓存的账。每一步要对长度为 $t$ 的前缀做一次完整因果前向,注意力部分是 $O(t^2)$;生成 $S$ 个 token 总共是 $$\sum_{t=1}^{S} c \cdot t^2 \approx \frac{c \, S^3}{3}$$ 用缓存之后,每步只算一个新 query 对 $t$ 个缓存条的注意力,是 $O(t)$;总共 $$\sum_{t=1}^{S} 2c \cdot t \approx c \, S^2$$ 按上面这套只计有效因果三角的常数约定,比值趋于 $S/3$,不是 $2S/3$。若全量实现先计算完整 $t\times t$ 分数再掩码(本文 NumPy 就是这种实现),它做了约两倍的注意力乘加,比值才趋于 $2S/3$。这两个系数对应不同实现,不能混在一起。 投影与 FFN 的账不同:全量路线每步重算前缀各 token,累计为 $O(S^2D^2)$;缓存路线累计为 $O(SD^2)$,$D$ 为隐藏宽度。注意力累计阶数则分别是 $O(S^3D)$ 和 $O(S^2D)$。有长度 $P$ 的 prompt 时应从 $P$ 开始求和,生成 $G$ 个输出只需一次 prefill 加 $G-1$ 次 decode,因为 prefill 已给出首个预测。附录的解析计数按实际密集 NumPy 矩阵乘与每步单个输出头计数,省略 norm、softmax 等逐元素操作。 3.3 缓存自己要多大:显存公式 每生成一个 token,要在缓存里留下这一层的 $k_t$ 和 $v_t$。数一数字节数: $$B_{\text{tok}} = \underbrace{2}_{K,V} \times \; L \times n_{\text{kv}} \times d \times p$$ $L$ 是层数,$n_{\text{kv}}$ 是 KV 头数(GQA 下小于 query 头数 $n_q$),$d$ 是每头维度,$p$ 是 dtype 字节数(BF16 是 2)。这个公式里没有 S——每个 token 的缓存占用与上下文长度无关,缓存总量才随 $S$ 线性增长: $$B_{\text{kv}} = B_{\text{tok}} \times S \times B_{\text{batch}}$$ 代 Llama-3-8B($L=32$,$n_{\text{kv}}=8$,$d=128$,$p=2$):$B_{\text{tok}} = 131072$ 字节。8192 token 一条序列 = 1.00 GiB;batch=32 就是 32.00 GiB,是 14.96 GiB 权重的 2.14 倍。如果换成 MHA($n_{\text{kv}} = 32$),每个 token 变成 512 KiB,batch=32 时 128 GiB——GQA 在这里把 KV 显存降低 4 倍,也会减少 K/V 投影计算;query 头上的注意力乘加不同比例下降。 这张图要看什么:左图是 batch 从 1 扫到 128 时显存账本的构成(seq=8192),深蓝是权重、橙红是 KV cache、浅蓝是激活和余量,红色虚线是 80 GiB——注意从 batch=32 开始橙红就盖过深蓝,计入示意的 4 GiB 余量后 batch=64 已超过 80 GiB 预算、batch=128 直接越界,增长主要来自 KV cache。右图固定 batch 看 KV cache 随序列长度的增长(对数-对数坐标),三条线都是斜率 1 的直线(线性增长),紫色虚线(MHA)比橙色(GQA batch=8)高一截;61 GiB 灰线表示扣除假设权重和余量后的缓存预算,曲线与它的交点给出该简化预算下的长度上限。 3.4 算术强度:为什么上下文再长,decode 也快不起来 这是全篇最要紧的一节。上一篇(性能建模)说过,一个操作的算术强度 $I$ = 算力 / 访存字节数,它和硬件的山脊点(ridge point)比一比,就知道这个操作是算力受限还是带宽受限。 先只看理想融合的单 query 注意力,假设每份 KV 从所分析的存储层读一次,并在一组 query 头之间共享,忽略 Q/O 和中间量:对每个 query 头,$QK^{\top}$ 是 $1 \times d$ 乘 $d \times S$,$2 S d$ FLOP;再加权和 $AV$ 同样 $2 S d$ FLOP。$n_q$ 个头合计 $4 n_q d S$ FLOP。访存呢?要把 K 和 V 的缓存全部读一遍:$2 n_{\text{kv}} d S p$ 字节。两者一除: $$I_{\text{dec}} = \frac{4 \, n_q \, d \, S}{2 \, n_{\text{kv}} \, d \, S \, p} = \frac{2 g}{p}, \quad g = \frac{n_q}{n_{\text{kv}}}$$ $S$ 在这个理想模型中约掉了。 算术强度不随 $S$ 变化,但实际效率可能随上下文而变:短序列并行不足,长序列跨越缓存容量,kernel 的分块和归约成本也会变化。而且注意力之外仍有读权重、投影与 FFN,不能用一个注意力公式概括完整 decode。 代数字(BF16,$p=2$,$I = g$):MHA($g=1$)的 $I = 1$ FLOP/byte;Llama-3 的 GQA($g=4$)是 4;MQA($g=32$)是 32。而 A100 的山脊点是 $312\ \text{TFLOP/s} \div 2.04\ \text{TB/s} = 153$ FLOP/byte。这些理想强度低于该 BF16 山脊点,提示带宽约束;实际 kernel 还可能受并行度与延迟约束。 下图把本机 FP32 微基准(1201.61 GFLOP/s、71.89 GB/s)与 A100 80GB SXM 官方规格 的 BF16 理论峰值分别作为屋顶。CPU 与 GPU 使用不同 dtype,因此同一 $g$ 的 CPU 理想强度为 GPU 的一半。 图中各点是把解析强度放到屋顶上计算出的上界,不是测得的 attention 性能。本机屋顶来自大矩阵和流式数组微基准,不保证小算子能达到。改变 GQA、精度、融合、序列并行或多 query 批处理都可能改善性能;MLA 有自己的压缩结构,不能简单当成增大 $g$。 长 prefill 一次处理多个 query,更容易复用权重和 KV,因此通常有更高强度。附录使用 $2NS$ 近似投影计算量($N$ 为参数量),并假定每层激活搬运为 10*S*D*p,得到示意强度 3973.9 FLOP/byte(BF16)。这是粗略模型,不是逐算子的显存流量测量;输入很短、注意力未融合、张量并行通信或低效 kernel 都可能改变实际瓶颈。不能仅凭这个数断言 prefill 必定吃满 GPU。 这个视角还能直接写出 decode 一步的耗时下限。若权重与全部 KV 每步都需从 HBM 读取,字节数除带宽给出一个下界;还须同时满足 FLOPs/算力下界: $$T_{\text{step}} \ge \frac{W + B_{\text{batch}} \, B_{\text{tok}} \, S}{\mathrm{BW}}$$ $W$ 是权重大小(每步都要读一遍,batch 摊薄),$B_{\text{batch}} B_{\text{tok}} S$ 是全部序列的缓存(batch 摊不了)。代 Llama-3-8B、seq=8192、A100 的 2.04 TB/s:batch=1 时 $(14.96 + 1.00)$ GiB 除以带宽得 8.4 ms,也就是单流吞吐的上限约 119 token/s——仅是单设备、该精度、该访存假设下的上界;权重量化、共享前缀、投机解码、多设备等会改变假设。第 6.1 节进一步说明 batch 对这个模型的影响。 04. 代码实现 核心是三段:完整前向(prefill)、单步(decode)、以及承载它们的缓存。下面是 kv_cache_lab.py 的主干,完整版在文末附录,纯 numpy 可跑。 def attn_full(Q, K, V): """完整因果注意力。Q:[S,Hq,Dh] K/V:[S,Hkv,Dh] -> [S,Hq,Dh]""" S, Hq, Dh = Q.shape g = Hq // K.shape[1] Kk = np.repeat(K, g, axis=1) # GQA:把 kv 头复制 g 份对齐 q 头 Vv = np.repeat(V, g, axis=1) logits = np.einsum("qhd,khd->hqk", Q, Kk) / np.sqrt(Dh) mask = np.triu(np.ones((S, S), dtype=bool), 1) logits = np.where(mask[None, :, :], -np.inf, logits) logits = logits - logits.max(axis=-1, keepdims=True) p = np.exp(logits) p = p / p.sum(axis=-1, keepdims=True) return np.einsum("hqk,khd->qhd", p, Vv) def attn_step(q, K, V): """decode 一步:一个 query 对长度 S 的缓存。q:[Hq,Dh] K/V:[S,Hkv,Dh] -> [Hq,Dh]""" Hq, Dh = q.shape Kk = np.repeat(K, Hq // K.shape[1], axis=1) Vv = np.repeat(V, Hq // K.shape[1], axis=1) logits = np.einsum("hd,khd->hk", q, Kk) / np.sqrt(Dh) # [Hq,S],只有一行 logits = logits - logits.max(axis=-1, keepdims=True) p = np.exp(logits) p = p / p.sum(axis=-1, keepdims=True) return np.einsum("hk,khd->hd", p, Vv) def block_step(x_new, w, cache, pos, cfg=CFG): """decode 一步:只算新 token,并把它的 K/V 写进缓存的 pos 位置。""" D = x_new.shape[1] hq, hkv, dh = cfg["n_q_head"], cfg["n_kv_head"], cfg["d_head"] h = rmsnorm(x_new) q = (h @ w["wq"]).reshape(hq, dh) k = (h @ w["wk"]).reshape(hkv, dh) v = (h @ w["wv"]).reshape(hkv, dh) cache["k"][pos] = k # 写缓存:只有这一步是新算的 cache["v"][pos] = v o = attn_step(q, cache["k"][:pos + 1], cache["v"][:pos + 1]).reshape(1, hq * dh) y = x_new + o @ w["wo"] y = y + np.maximum(y @ w["w1"], 0.0) @ w["w2"] return y 变量名和 03 节的符号一一对应:Q/K/V 是 $Q$、$K$、$V$,Hq/Hkv 是 $n_q$、$n_{\text{kv}}$,g 就是 GQA 的分组数 $g$,pos 是当前长度 $t$。缓存预分配成最大长度(prompt + gen),KV 更新按下标原地写入;注意力仍创建临时数组,并用 np.repeat 实际复制 KV 头,因此它没有实现理论分析所假定的完美 GQA 访存复用。 修订后的代码在本机运行得到(时间随机器、BLAS 和负载变化): 生成 token 一致数:192 / 192 最后 logits 最大绝对误差:7.344e-06 相对误差:4.028e-07 无缓存矩阵乘 FLOPs:1.3005e+11 缓存矩阵乘 FLOPs:1.2314e+09 解析比值:105.61x 无缓存 / 缓存墙钟:6.986s / 0.679s 本次加速比:10.29x 全量与增量在精确算术下等价,FP32 累加顺序会带来小误差。真实模型若候选 logits 很接近,舍入也可能改变 argmax,所以不要求所有模型都逐 token 位级一致。附录固定种子玩具例子对序列与误差都加了断言。本例省略位置编码,用于验证缓存数据流;RoPE 和视频块缓存的正确性还需要另行验证。 解析 FLOPs 与墙钟的差额只说明实现效率不同。若要确认带宽瓶颈,还需测实际带宽、kernel 时间、线程设置和缓存命中,不能由加速比倒推出原因。 05. 工业级实现对照 HF transformers:动态缓存与静态缓存是不同路径。 cache_utils.py 的 DynamicLayer.update(以 2026-10 的实现为准)是这么追加一个 token 的: self.keys = torch.cat([self.keys, key_states], dim=-2) self.values = torch.cat([self.values, value_states], dim=-2) torch.cat 每一步都分配一块新内存、把旧缓存整个抄过去。缓存越大这一步越贵,到长上下文时光是拷贝就吃掉不少带宽——而带宽恰恰是 decode 最缺的东西。它换来的是简单:形状任意增长、随便 crop、随便回滚,对 batch=1 的研究代码完全够用。我第一版玩具实现也是这么写的,后来才意识到「预分配 + 下标写入」差在哪:一个把带宽花在拷贝上,一个把带宽花在读缓存上,前者是纯浪费。 同一份 cache_utils.py 还提供 StaticLayer,预分配后原地更新;不能把动态 torch.cat 描述为 transformers 唯一方式。动态拼接与按最大长度预留是两种不同策略,前者可能复制和碎片化,后者有未用容量。 vLLM:把缓存当成页来管。 当服务请求长度不确定时,按最大长度预留可能浪费容量。vLLM(PagedAttention,arXiv:2309.06180)的解法是把操作系统管内存的那一套搬过来:缓存切成固定大小的块(block),每条请求维护一张块表(block table),逻辑上连续、物理上散落。vLLM v0.16.0 的 FlashAttention 后端文档 中可检查 FlashAttentionImpl.forward、key_cache, value_cache = kv_cache.unbind(0) 与 block_table 的使用。缓存张量布局会随版本和后端变化,不能把某个 2 * head_size 排列写成 vLLM 的统一规定。稳定的设计要点是逻辑块映射到物理块,kernel 按块表寻址,新增 K/V 按 slot 映射写入;prefill/decode 通过长度等元数据区分。 分页到底值多少?附录 paged_alloc.py 在同一个 61 GiB 预算下做了模拟(Llama-3-8B,块 16): --- 负载:长度均匀 512~4096 --- 策略 并发请求 占用 真实用到 利用率 连续预留 122 61.00 GiB 34.52 GiB 56.59% 分页(块16) 214 60.85 GiB 60.67 GiB 99.70% 并发提升 : 1.75x --- 负载:重尾:八成 256~1024 --- 连续预留 122 61.00 GiB 18.37 GiB 30.11% 分页(块16) 405 60.79 GiB 60.42 GiB 99.40% 并发提升 : 3.32x 浪费的两半也拆开了:连续预留平均每条请求浪费 223.26 MiB(预留 4096、平均只用 2310),分页只有 0.93 MiB(最后一个块没填满),239 倍。分页确实减少预留浪费和碎片,未改变有效 KV 每 token 的字节数。这里是静态容量模拟:1.75~3.32 倍是可容纳请求数之比,不是测得的吞吐倍数;尚未模拟到达、释放、动态增长和调度。 这张图要看什么:左图是分页的时空图,上面各行是每条请求的逻辑视图(块连续),下面一行是物理块池(按申请顺序排列,颜色表示属于谁)——灰色细线从逻辑块指向它真正的物理块,能看到同一条请求的块在物理上是散的;块里的灰底数字是「已用/容量」,只有每条请求的最后一块没填满。右图是两种策略在同一预算下的并发数,灰色(连续预留)在重尾负载下利用率只剩 30%,蓝色(分页)两种负载都在 99% 以上。 还有一个 decode 特有的并行技巧。 prefill 可以把 $S^2$ 的注意力摊到很多 SM 上,decode 只有一个 query、$S$ 个 key——并行度天然不足(FlashDecoding 的出发点)。做法是把序列维切成几段,各段独立算局部 softmax 再合并(online softmax 的分治),用更多并行度换带宽利用率。这与 3.4 节的结论一致:decode 的问题是算术强度低,切序列不改变 $I$,但能把空闲的算力单元动员起来去搬字节。 生产部署还要考虑块表处理、CPU 调度、编译与内核启动开销。分页增加了寻址和管理成本,是否获益取决于请求负载,不能只看有效容量。 06. 代价与边界 6.1 batch 摊权重,但独立请求的 KV 随 batch 增长 在第 3.4 节的带宽模型里,一批请求共享一次权重读取,而各自拥有不同 KV。若上下文相同,batch 增大时总 KV 字节数线性增长;单步吞吐上界是 $B_{\text{batch}}/T_{\text{step}}$,不会无限线性增加。共享前缀、不同长度调度、投机解码和张量并行会改变这张账,需重新列出复用假设。 6.2 分页不改变有效 KV 大小,但减少浪费 块大小越小,最后一块的空槽通常越少,块表却越长。若长度模块大小的余数均匀,块大小 $b$ 的平均空槽是 $(b-1)/2$,不是无条件精确的 $b/2$。本例平均长度约 2310,16/32 token 块的空槽比例约 0.32%/0.67%;真实工作负载与 kernel 对齐要求应共同决定块大小。分页能让原本浪费的显存参与服务,也支持一些共享场景;batch=1 同样可能减少预留容量,但未必带来明显延迟收益。 6.3 视频缓存:因果掩码只是必要条件之一 一些自回归视频系统按帧或块推进,块内双向、块间因果;另一些按离散 token 生成或采用不同的窗口。下图仅展示一种块因果结构: 图中右上角为空表示不能读未来块,对角块为满表示当前块内双向。允许读取历史,不等于历史 K/V 在所有去噪步都不变。 当前块的噪声状态随去噪变化,其 K/V 通常必须重算;历史块只有在输入、位置、时间/噪声条件、模型权重均固定且架构允许时才能精确缓存。若历史也被重新加噪或全局时间条件改变,需按模型缓存策略重建,不能直接套文本 decoder 的永久缓存假设。 本例假设 720p、空间压缩 8×8、时间压缩 4×、每 latent 位置一个 token、全历史保留,并套用 Llama-3-8B 的缓存结构;5/10/30 秒分别为 52.73/105.47/316.41 GiB。额外 2×2 patch 化会让空间 token 数约为四分之一;因果 VAE 的首帧约定、边界取整、滑动窗口、层数与头数还会继续改变数字。一帧的 1.76 GiB 是这组假设的计算结果,不是视频模型的普遍常数。 语义块和分配块不必相同。 一帧可包含很多物理缓存页,原有 token 分页机制仍可复用;需要调整的是掩码、批处理和缓存生命周期。不能仅凭帧内双向就断言块表必须整帧分配,或元数据一定压垮 CPU。 位置必须保留逻辑含义。 3D RoPE 需要时间/高/宽坐标,不能只用缓存里的 token 条数推断它们。滑窗驱逐后,存储长度尤其不等于全局逻辑位置,需单独维护 position_ids 或绝对偏移。丢缓存并不删除已输出的视频帧,只会改变未来生成能读到的上下文,可能影响长时一致性。 量化、驱逐、压缩各有条件。 理想地把 KV 每元素字节从 2 降到 1,载荷减半,但总存储还包括量化 scale、元数据和工作区。FP8 不同格式有不同动态范围,离群值与精度误差要验证,不能只按「减半」判断能部署。MLA 是训练时设计的潜在注意力结构,不是给任意既有模型套一个无损压缩器;滑窗或驱逐会改变可见历史。 6.4 证据的范围 本篇证实了玩具文本 decoder 的缓存等价性,计算了指定配置的显存账,并做了静态分页容量模拟。Roofline 给出假设下的上界,未测 A100 kernel,未实现完整视频去噪缓存,也没有证明任何真实视频模型必须采用某个驱逐策略。CPU FP32 微基准、GPU BF16 理论峰值与 NumPy 教学实现的访存行为要分开理解。 07. 经典论文脉络 Attention Is All You Need(arXiv:1706.03762)——decoder 的因果结构允许缓存固定历史,缓存是利用因果结构减少重复计算的常见优化,并非数学正确性所必需。 Fast Transformer Decoding: One Write-Head is All You Need(arXiv:1911.02150,MQA)——系统讨论了 decode 的瓶颈是带宽而非算力,并把所有 query 头共享一对 KV 头,把缓存压到 $1/g$。贡献是把「算术强度」这个视角带进了推理优化。 GQA: Training Generalized Multi-Query Transformer Models from Multi-head Checkpoints(arXiv:2305.13245)——MQA 掉点太狠,这篇用「分组共享 + 上游检查点升级」折中:8 个 KV 头保住大部分质量,缓存仍压到 1/4。具体分组数与是否采用 GQA 随模型规格而异。 FlashAttention(arXiv:2205.14135)与后续的 FlashDecoding——前者说明注意力可以不把 $S \times S$ 矩阵写回显存(本系列已写过);后者把 decode 的序列维切开并行,专治「一个 query、一长串 key」的并行度不足。它们改善不同形状的 IO 与并行效率,但不保证始终达到理论带宽。 Efficient Memory Management for Large Language Model Serving with PagedAttention(arXiv:2309.06180,vLLM)——把虚拟内存的分页思想搬进 KV cache,解决「输出长度未知导致的预留浪费与外部碎片」,本篇 05 节用自设负载模拟预留浪费,不是论文 benchmark 的直接复现。这是推理服务从「单条请求优化」走向「系统优化」的分水岭。 两条缓存压缩路线的起点:StreamingLLM(arXiv:2309.17453)发现「开头几个 token + 最近窗口」就能稳定外推,给出了驱逐策略的最简形式;DeepSeek-V2 的 MLA(arXiv:2405.04434)则把 K/V 联合投影到低秩隐空间再缓存,压缩比远超 GQA。前者改变可见上下文,后者是在训练中学习的注意力参数化,都直接对应 6.3 节视频场景里那道「丢什么、怎么丢」的选择题。 08. 常见误解 误解 1:「算力少 100 倍就一定快 100 倍。」 墙钟还受 kernel、调度、带宽和分配影响,加速比需实测。上下文变长使缓存路线自身更慢,但相对全量重算的加速比可能反而增大,不能说必然恶化。 误解 2:「所有 decode 都是纯访存。」 低 batch、长上下文的注意力常受带宽限制,但整体 decode 还有投影、FFN、通信与 CPU 开销。先用 profiler 定位。 误解 3:「GQA 只影响质量。」 它减少 KV 容量、KV 投影与理想读取字节,但 query 头数不变;代价要结合训练质量和实际 kernel 评估。 误解 4:「分页不省显存,容量增益直接等于吞吐增益。」 分页减少的是预留与碎片浪费,有效 KV 本身不变。能放更多请求并不保证同倍数 tokens/s。 误解 5:「缓存长度就是新 token 的位置。」 仅在简单连续全缓存情形下成立。驱逐、packing、padding 或 3D 坐标都需要额外位置元数据,不能从当前存储长度猜。 误解 6:「因果视频掩码就保证跨去噪步复用 KV。」 掩码约束依赖方向,缓存还要求被缓存的隐藏状态不变;当前噪声块及发生条件变化的历史必须重新计算。 09. 动手验证 数值脚本依赖 numpy,配图另需 matplotlib,均不需要 torch: python kv_cache_lab.py ALL # 约 8 秒:等价性 + 算力账 + 显存账 + Roofline python paged_alloc.py ALL # 约 1 秒:分页 vs 连续预分配的并发与浪费 预期结果(实跑输出,可以直接对): kv_cache_lab.py 的 [A1] 必须是 192 / 192 完全一致,[A2] 的相对误差在 $10^{-7}$ 量级。如果你的 [A2] 是 $10^{-2}$ 量级,九成是参照解喂多了 token([A] 段注释里那个 off-by-one,我第一版就踩了)。 [A3] 修订后的矩阵乘计算量比约 105.61;[A4] 时间不设固定范围,它取决于机器和运行条件。 [B1] 每个 token 131072 字节、8192 token 正好 1.0000 GiB;[B4] 三条视频账 52.73 / 105.47 / 316.41 GiB。 paged_alloc.py 的 [A]:均匀负载 122 → 214(1.75x),重尾负载 122 → 405(3.32x);[B] 两类浪费之比约 239 倍。 想自己碰一下边界?改 kv_cache_lab.py 顶部的 CFG["n_kv_head"](从 2 改到 8,变成 MHA),再跑一遍:缓存载荷变 4 倍,但投影结构、临时复制与两条路线的耗时也会变化,实测加速比不保证单调。NumPy 的 repeat 会削弱理想 GQA 访存收益,因此这个实验不能直接验证 GPU 带宽模型。 10. 延伸阅读 性能建模与 Profiling:算力、带宽与显存账本(本系列已发布)——本篇 3.4 节的算术强度和山脊点在那里有完整推导和实跑标定方法;读那篇再看本篇的 Roofline 图会非常顺。 自注意力机制的计算与显存账本(本系列已发布)——训练态注意力的 $O(S^2)$ 账本,本篇是它在推理态的续集。 FlashAttention 为什么不需要存下注意力矩阵(本系列已发布)——07 节第 4 条的展开,online softmax 的分治细节在那里。 自回归视频生成与 Forcing 范式(本系列已发布)——6.3 节的因果结构(帧内双向、帧间因果)是从那套范式来的;读完范式再看本篇的缓存账,能对上号。 扩散模型的跨步缓存复用(本系列规划中)——同样是「缓存」,扩散模型的跨去噪步近似复用可能针对中间特征或注意力状态,与本文的精确历史缓存要区分,动机相同、机制完全不同,适合对照着读。 附录:完整代码 09 节用到的脚本全文如下(kv_cache_lab.py、paged_alloc.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 kv_cache_lab.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """KV Cache 实验室。 三个问题,全部用可复现的实跑回答,不靠记忆里的结论: [A] 增量解码(用缓存)和「每步把整个前缀重算一遍」,得到的到底是不是同一个东西? 以及真实加速比是多少、算力量省了多少倍。 [B] 缓存自己要吃掉多少显存?按 Llama-3-8B 的真实配置算一遍, 再算一遍自回归视频的 token 数,看哪个先爆。 [C] decode 一步的算术强度为什么和上下文长度无关?把它放到 Roofline 上看。 纯 numpy,不需要 torch / scipy。 python kv_cache_lab.py ALL # 全部,约 30 秒 python kv_cache_lab.py A # 只跑等价性与加速比 python kv_cache_lab.py B # 只跑显存账本 python kv_cache_lab.py C # 只跑算术强度与 Roofline """ from __future__ import annotations import argparse import json import math import os import time import numpy as np HERE = os.path.dirname(os.path.abspath(__file__)) # 玩具模型的配置:刻意做成 GQA(8 个 q 头共享 2 个 kv 头,g = 4), # 它是常见的一种结构;MHA 是 g = 1 的特例,其他模型也可能采用 MQA/MLA。 CFG = dict( n_layer=4, n_q_head=8, n_kv_head=2, d_head=32, ffn_mult=4, prompt=16, gen=192, vocab=32, dtype_bytes=4, # 玩具模型用 fp32 跑,记账就按 4 字节 ) def d_model(cfg=CFG): return cfg["n_q_head"] * cfg["d_head"] # ══════════════════════════════════════════════════════════════ # 0. 一个能跑的小 Transformer decoder # ══════════════════════════════════════════════════════════════ def make_weights(cfg=CFG, seed=0): rng = np.random.default_rng(seed) D = d_model(cfg) Dk = cfg["n_kv_head"] * cfg["d_head"] F = cfg["ffn_mult"] * D s = 1.0 / math.sqrt(D) blocks = [] for _ in range(cfg["n_layer"]): blocks.append(dict( wq=rng.normal(0, s, (D, D)).astype(np.float32), wk=rng.normal(0, s, (D, Dk)).astype(np.float32), wv=rng.normal(0, s, (D, Dk)).astype(np.float32), wo=rng.normal(0, s, (D, D)).astype(np.float32), w1=rng.normal(0, s, (D, F)).astype(np.float32), w2=rng.normal(0, s, (F, D)).astype(np.float32), )) E = rng.normal(0, s, (cfg["vocab"], D)).astype(np.float32) # token embedding head = rng.normal(0, s, (D, cfg["vocab"])).astype(np.float32) return blocks, E, head def rmsnorm(x, eps=1e-8): return x / np.sqrt(np.mean(x * x, axis=-1, keepdims=True) + eps) def attn_full(Q, K, V): """完整因果注意力。Q:[S,Hq,Dh] K/V:[S,Hkv,Dh] -> [S,Hq,Dh]""" S, Hq, Dh = Q.shape g = Hq // K.shape[1] Kk = np.repeat(K, g, axis=1) Vv = np.repeat(V, g, axis=1) logits = np.einsum("qhd,khd->hqk", Q, Kk) / np.sqrt(Dh) # [Hq,S,S] mask = np.triu(np.ones((S, S), dtype=bool), 1) logits = np.where(mask[None, :, :], -np.inf, logits) logits = logits - logits.max(axis=-1, keepdims=True) p = np.exp(logits) p = p / p.sum(axis=-1, keepdims=True) return np.einsum("hqk,khd->qhd", p, Vv) def attn_step(q, K, V): """单步注意力:一个 query 对长度 S 的缓存。q:[Hq,Dh] K/V:[S,Hkv,Dh] -> [Hq,Dh]""" Hq, Dh = q.shape g = Hq // K.shape[1] Kk = np.repeat(K, g, axis=1) Vv = np.repeat(V, g, axis=1) logits = np.einsum("hd,khd->hk", q, Kk) / np.sqrt(Dh) # [Hq,S] logits = logits - logits.max(axis=-1, keepdims=True) p = np.exp(logits) p = p / p.sum(axis=-1, keepdims=True) return np.einsum("hk,khd->hd", p, Vv) def block_forward(x, w, cfg=CFG, cache=None, base=0): """对一个前缀做完整因果前向。x:[S,D] -> [S,D]。 cache 不为 None 时,顺手把这一层的 K/V 写进 cache[base:base+S]。""" S, D = x.shape hq, hkv, dh = cfg["n_q_head"], cfg["n_kv_head"], cfg["d_head"] h = rmsnorm(x) Q = (h @ w["wq"]).reshape(S, hq, dh) K = (h @ w["wk"]).reshape(S, hkv, dh) V = (h @ w["wv"]).reshape(S, hkv, dh) if cache is not None: cache["k"][base:base + S] = K cache["v"][base:base + S] = V o = attn_full(Q, K, V).reshape(S, D) x = x + o @ w["wo"] x = x + np.maximum(x @ w["w1"], 0.0) @ w["w2"] return x def block_step(x_new, w, cache, pos, cfg=CFG): """decode 一步:只算新 token。x_new:[1,D] -> [1,D],并把它的 K/V 写进 pos。""" D = x_new.shape[1] hq, hkv, dh = cfg["n_q_head"], cfg["n_kv_head"], cfg["d_head"] h = rmsnorm(x_new) q = (h @ w["wq"]).reshape(hq, dh) k = (h @ w["wk"]).reshape(hkv, dh) v = (h @ w["wv"]).reshape(hkv, dh) cache["k"][pos] = k cache["v"][pos] = v o = attn_step(q, cache["k"][:pos + 1], cache["v"][:pos + 1]).reshape(1, hq * dh) y = x_new + o @ w["wo"] y = y + np.maximum(y @ w["w1"], 0.0) @ w["w2"] return y def new_cache(cfg=CFG, total=None): total = total or (cfg["prompt"] + cfg["gen"]) return [dict(k=np.zeros((total, cfg["n_kv_head"], cfg["d_head"]), np.float32), v=np.zeros((total, cfg["n_kv_head"], cfg["d_head"]), np.float32)) for _ in range(cfg["n_layer"])] # ══════════════════════════════════════════════════════════════ # [A] 等价性与加速比 # ══════════════════════════════════════════════════════════════ def gen_naive(blocks, E, head, prompt_tokens, cfg=CFG): """不用缓存:每一步把整个前缀重算一遍,只取最后一个位置的输出。""" X = E[prompt_tokens].copy() toks, last_logits = [], None for _ in range(cfg["gen"]): h = X for w in blocks: h = block_forward(h, w, cfg) logits = h[-1] @ head tok = int(np.argmax(logits)) toks.append(tok) last_logits = logits X = np.vstack([X, E[tok][None, :]]) return toks, last_logits def gen_cached(blocks, E, head, prompt_tokens, cfg=CFG): """用 KV cache:prefill 一次,之后每步只算一个新 token。""" total = cfg["prompt"] + cfg["gen"] caches = new_cache(cfg, total) X = E[prompt_tokens].copy() h = X for i, w in enumerate(blocks): h = block_forward(h, w, cfg, cache=caches[i], base=0) pos = len(prompt_tokens) - 1 h_last = h[-1:] toks, last_logits = [], None for t in range(cfg["gen"]): logits = h_last[0] @ head tok = int(np.argmax(logits)) toks.append(tok) last_logits = logits if t == cfg["gen"] - 1: break # 已得到最后一个预测,不再做一次未使用的 decode x_new = E[tok][None, :] pos += 1 h_new = x_new for i, w in enumerate(blocks): h_new = block_step(h_new, w, caches[i], pos, cfg) h_last = h_new return toks, last_logits, caches def forward_full(blocks, E, head, tokens, cfg=CFG): """一次性对整段 token 做完整前向(参照解)。""" h = E[tokens] for w in blocks: h = block_forward(h, w, cfg) return h[-1] @ head def flops_full(S, cfg=CFG): """一次长度 S 的完整因果前向的 FLOPs(只算矩阵乘,2*macs)。""" D = d_model(cfg) Dk = cfg["n_kv_head"] * cfg["d_head"] F = cfg["ffn_mult"] * D per_layer = (2 * (D * D) # wq + 2 * (D * Dk) # wk + 2 * (D * Dk) # wv + 2 * (D * D) # wo + 2 * (D * F) # w1 + 2 * (F * D)) # w2 # 实际代码只对最后一个位置计算输出头,计 2*D*vocab FLOP。 proj = cfg["n_layer"] * S * per_layer + 2 * D * cfg["vocab"] # NumPy 实现先做完整密集矩阵乘再掩码,不能按跳过上三角计数。 attn = cfg["n_layer"] * 4 * cfg["n_q_head"] * cfg["d_head"] * S * S return proj + attn def flops_step(S, cfg=CFG): """decode 一步(上下文长度 S)的 FLOPs。""" D = d_model(cfg) Dk = cfg["n_kv_head"] * cfg["d_head"] F = cfg["ffn_mult"] * D per_layer = (2 * (D * D) + 2 * (D * Dk) + 2 * (D * Dk) + 2 * (D * D) + 2 * (D * F) + 2 * (F * D)) proj = cfg["n_layer"] * per_layer + 2 * D * cfg["vocab"] attn = cfg["n_layer"] * 2 * (2 * cfg["n_q_head"] * cfg["d_head"] * S) return proj + attn def section_A(cfg=CFG): print("=" * 72) print("[A] 增量解码 vs 每步全量重算") print("=" * 72) blocks, E, head = make_weights(cfg) rng = np.random.default_rng(7) prompt_tokens = list(rng.integers(0, cfg["vocab"], cfg["prompt"])) S_end = cfg["prompt"] + cfg["gen"] # ── A1 缓存路线 ── t0 = time.perf_counter() toks_c, logits_c, caches = gen_cached(blocks, E, head, prompt_tokens, cfg) t_cached = time.perf_counter() - t0 # ── A2 无缓存路线 ── t0 = time.perf_counter() toks_n, logits_n = gen_naive(blocks, E, head, prompt_tokens, cfg) t_naive = time.perf_counter() - t0 # ── A3 参照解:一次性完整前向 ── # 对齐位置很容易错:最后一步的 logits 是在「倒数第二个 token」上算出来的 # (它用来预测最后一个 token),所以参照解只能喂到 toks_c[:-1]。 # 多喂一个 token,比的就是下一步的 logits 了。 all_tokens = list(prompt_tokens) + toks_c[:-1] logits_ref = forward_full(blocks, E, head, all_tokens, cfg) same = sum(1 for a, b in zip(toks_c, toks_n) if a == b) dif = float(np.max(np.abs(logits_c - logits_ref))) scale = float(np.max(np.abs(logits_ref))) assert same == cfg["gen"], "token sequences differ" assert dif <= 1e-5 * max(scale, 1.0), "cached/full logits mismatch" print("\n[A1] 两条路线生成的 token 序列") print(f" 序列长度 : {cfg['gen']}") print(f" 完全一致的 token : {same} / {cfg['gen']}") print("\n[A2] 缓存路线最后一步的 logits vs 一次性完整前向(参照解)") print(f" max |Δ| : {dif:.3e}") print(f" 参照解的量级 : {scale:.6f}") print(f" 相对误差 : {dif / scale:.3e}") # ── A4 算力账 ── f_naive = sum(flops_full(cfg["prompt"] + t, cfg) for t in range(cfg["gen"])) f_cached = (flops_full(cfg["prompt"], cfg) + sum(flops_step(cfg["prompt"] + t, cfg) for t in range(1, cfg["gen"]))) print("\n[A3] 算力账(解析计数,单位 FLOP)") print(f" 无缓存总算力 : {f_naive:.4e}") print(f" 有缓存总算力 : {f_cached:.4e}") print(f" 算力节省倍数 : {f_naive / f_cached:.2f}x") print("\n[A4] 真实墙钟(同一台机器,各跑一次)") print(f" 无缓存 : {t_naive:.3f} s") print(f" 有缓存 : {t_cached:.3f} s") print(f" 实测加速比 : {t_naive / t_cached:.2f}x") print(f" (算力省了 {f_naive / f_cached:.1f}x,墙钟只快 {t_naive / t_cached:.1f}x —— " f"差额还含 Python、分配、KV 复制和小矩阵效率,不能仅归因于带宽)") return dict( gen=cfg["gen"], same_tokens=same, max_diff=dif, ref_scale=scale, flops_naive=f_naive, flops_cached=f_cached, flops_ratio=f_naive / f_cached, t_naive=t_naive, t_cached=t_cached, speedup=t_naive / t_cached, seq_end=S_end, ) # ══════════════════════════════════════════════════════════════ # [B] 显存账本 # ══════════════════════════════════════════════════════════════ # Llama-3-8B 的公开配置 LLAMA3_8B = dict(name="Llama-3-8B", n_layer=32, n_q_head=32, n_kv_head=8, d_head=128, params=8.03e9, dtype_bytes=2) def kv_bytes_per_token(cfg, n_kv_head=None): """每个 token 的 KV cache 字节数(所有层)。""" nkv = cfg["n_kv_head"] if n_kv_head is None else n_kv_head return 2 * cfg["n_layer"] * nkv * cfg["d_head"] * cfg["dtype_bytes"] def kv_bytes_total(cfg, S, batch, n_kv_head=None): return kv_bytes_per_token(cfg, n_kv_head) * S * batch def section_B(): print("\n" + "=" * 72) print("[B] KV cache 自己吃掉多少显存") print("=" * 72) c = LLAMA3_8B bpt = kv_bytes_per_token(c) print("\n[B1] 每个 token 的 KV cache(Llama-3-8B,BF16)") print(f" 2 (K,V) x {c['n_layer']} 层 x {c['n_kv_head']} kv头 x {c['d_head']} 维 x 2 字节") print(f" = {bpt} 字节/token = {bpt / 1024:.0f} KiB/token") print(f" 8192 token 一条序列 = {bpt * 8192 / (1024**3):.4f} GiB") w_bytes = c["params"] * c["dtype_bytes"] print(f"\n[B2] 和权重比一比(权重 {w_bytes / (1024**3):.2f} GiB)") print(f" {'batch':>6} {'seq':>6} {'KV cache':>12} {'KV/权重':>9} {'合计':>10}") rows = [] for batch, S in [(1, 8192), (8, 8192), (16, 8192), (32, 8192), (64, 8192), (32, 32768)]: kb = kv_bytes_total(c, S, batch) rows.append(dict(batch=batch, seq=S, kv=kb, total=kb + w_bytes, ratio=kb / w_bytes)) print(f" {batch:>6} {S:>6} {kb / (1024**3):>9.2f} GiB " f"{kb / w_bytes:>8.2f}x {(kb + w_bytes) / (1024**3):>7.2f} GiB") # MHA 对照 bpt_mha = kv_bytes_per_token(c, n_kv_head=c["n_q_head"]) print(f"\n[B3] 如果换成 MHA(32 个 kv 头而不是 8 个)") print(f" {bpt_mha} 字节/token = {bpt_mha / 1024:.0f} KiB/token" f" (GQA 的 {bpt_mha / bpt:.0f} 倍)") print(f" batch=32 / seq=8192 时:" f"{kv_bytes_total(c, 8192, 32, c['n_q_head']) / (1024**3):.2f} GiB" f" vs GQA {kv_bytes_total(c, 8192, 32) / (1024**3):.2f} GiB") # ── 视频 ── print("\n[B4] 自回归视频:token 数先把你压垮") vcfg = dict(c, dtype_bytes=2) cases = [] for name, frames, tf, h, w, fps, sec in [ ("5s 720p", 120, 4, 1280, 720, 24, 5), ("10s 720p", 240, 4, 1280, 720, 24, 10), ("30s 720p", 720, 4, 1280, 720, 24, 30), ]: lat_frames = frames // tf tok_per_frame = (h // 8) * (w // 8) # 空间 8x8 压缩 ntok = lat_frames * tok_per_frame kb = ntok * bpt cases.append(dict(name=name, lat_frames=lat_frames, tok_per_frame=tok_per_frame, ntok=ntok, kv=kb)) print(f" {name:>8}: 潜在帧 {lat_frames:>3} x 每帧 {tok_per_frame:>5} token" f" = {ntok:>8} token -> KV cache {kb / (1024**3):>8.2f} GiB(单条视频)") return dict(bytes_per_token=bpt, bytes_per_token_mha=bpt_mha, weight_bytes=w_bytes, rows=rows, video=cases) # ══════════════════════════════════════════════════════════════ # [C] 算术强度与 Roofline # ══════════════════════════════════════════════════════════════ def probe_peak(n=3072, n_bytes=20_000_000, repeat=5): """本机标定:峰值算力(大矩阵乘)与峰值带宽(大数组流式读写)。 带宽不能用 x.sum() 测:标量归约是延迟受限的,实测只有 23 GB/s, 而同一块内存的流式读写能到 75 GB/s。差 3 倍,用错了整个 Roofline 就歪了。 这里取几种流式算子里最快的一个。 """ rng = np.random.default_rng(0) a = rng.normal(size=(n, n)).astype(np.float32) b = rng.normal(size=(n, n)).astype(np.float32) best = float("inf") for _ in range(repeat): t0 = time.perf_counter() a @ b best = min(best, time.perf_counter() - t0) peak_flops = 2.0 * n ** 3 / best x = rng.normal(size=n_bytes).astype(np.float32) y = np.empty_like(x) best = float("inf") for fn in (lambda: np.copyto(y, x), lambda: np.add(x, 1.0, out=y)): for _ in range(4): t0 = time.perf_counter() fn() best = min(best, time.perf_counter() - t0) peak_bw = (2 * x.nbytes) / best # 读一份 + 写一份 return peak_flops, peak_bw def intensity_decode(g, dtype_bytes=2): """decode 一步「注意力部分」的算术强度 I = 2g/p,与上下文长度无关。""" return 2.0 * g / dtype_bytes def section_C(): print("\n" + "=" * 72) print("[C] 算术强度:为什么上下文再长,decode 也改善不了") print("=" * 72) print("\n[C1] decode 一步的注意力部分(上下文长度 S,GQA 分组数 g,dtype 字节 p)") print(" 算力 = 4 * n_q * d * S (QK^T 与 AV 各 2*n_q*d*S)") print(" 访存 = 2 * n_kv * d * S * p (K、V 各一份)") print(" I = 4 n_q d S / (2 n_kv d S p) = 2 (n_q/n_kv) / p = 2g/p") print(" —— 理想融合注意力中 S 被约去;不代表实际 kernel 效率随 S 不变。") print() print(f" {'g (n_q/n_kv)':>12} {'I = 2g/p (BF16)':>16}") rows_g = [] for g in [1, 2, 4, 8, 16, 32]: I = intensity_decode(g, 2) rows_g.append(dict(g=g, I=I)) print(f" {g:>12} {I:>16.1f}") # ── prefill 对照 ── c = LLAMA3_8B D = c["n_q_head"] * c["d_head"] print("\n[C2] prefill 的算术强度(Llama-3-8B,S=8192)") S = 8192 attn_flops = c["n_layer"] * 2 * (2 * c["n_q_head"] * c["d_head"] * S * (S + 1) / 2) proj_flops = 2 * c["params"] * S # 粗略 2NS 估计,不是逐层实测 FLOPs w_bytes = c["params"] * c["dtype_bytes"] act_bytes = c["n_layer"] * 10 * S * D * c["dtype_bytes"] I_pre = (attn_flops + proj_flops) / (w_bytes + act_bytes) print(f" 注意力算力 : {attn_flops / 1e12:.2f} TFLOP") print(f" 投影层算力 : {proj_flops / 1e12:.2f} TFLOP <-- dominates") print(f" 访存(权重+激活): {(w_bytes + act_bytes) / (1024**3):.2f} GiB") print(f" I_prefill : {I_pre:.1f} FLOP/byte") # ── Roofline ── peak_flops, peak_bw = probe_peak() ridge = peak_flops / peak_bw print("\n[C3] 本机标定(numpy / CPU,实跑)") print(f" 峰值算力 : {peak_flops / 1e9:.2f} GFLOP/s") print(f" 峰值带宽 : {peak_bw / 1e9:.2f} GB/s") print(f" 山脊点 : {ridge:.2f} FLOP/byte") # A100-80GB 公开规格 a100 = dict(flops=312e12, bw=2.039e12) print(f"\n A100-80GB(公开规格,非实测):{a100['flops'] / 1e12:.0f} TFLOP/s BF16 / " f"{a100['bw'] / 1e12:.2f} TB/s -> 山脊点 {a100['flops'] / a100['bw']:.1f} FLOP/byte") print("\n[C4] 理想注意力 Roofline:CPU 按 FP32 (I 为 BF16 的一半),A100 按 BF16") print(f" {'工作负载':>18} {'I':>10} {'本机判定':>10} {'A100 判定':>10}") pts = [] for g in [1, 4, 8]: I = intensity_decode(g, 2) pts.append(dict(name=f"decode g={g}", I=I, I_local=I / 2, local="带宽" if I / 2 < ridge else "算力", a100="带宽" if I < a100["flops"] / a100["bw"] else "算力")) print(f" {f'decode g={g}':>18} {I:>10.1f} {pts[-1]['local']:>10} {pts[-1]['a100']:>10}") pts.append(dict(name="prefill S=8192", I=I_pre, I_local=I_pre / 2, local="带宽" if I_pre / 2 < ridge else "算力", a100="带宽" if I_pre < a100["flops"] / a100["bw"] else "算力")) print(f" {'prefill S=8192':>18} {I_pre:>10.1f} {pts[-1]['local']:>10} {pts[-1]['a100']:>10}") return dict(rows_g=rows_g, I_prefill=I_pre, I_prefill_local=I_pre / 2, peak_flops=peak_flops, peak_bw=peak_bw, ridge=ridge, a100_flops=a100["flops"], a100_bw=a100["bw"], a100_ridge=a100["flops"] / a100["bw"], points=pts) def main(): ap = argparse.ArgumentParser() ap.add_argument("which", nargs="?", default="ALL") args = ap.parse_args() w = args.which.upper() out = {} if w in ("ALL", "A"): out["A"] = section_A() if w in ("ALL", "B"): out["B"] = section_B() if w in ("ALL", "C"): out["C"] = section_C() if w == "ALL": with open(os.path.join(HERE, "_kv_results.json"), "w", encoding="utf-8") as f: json.dump(out, f, ensure_ascii=False, indent=1) print("\n结果已写入 _kv_results.json(画图脚本读它,避免图上的数字和正文漂移)") if __name__ == "__main__": main() paged_alloc.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """PagedAttention 的分页分配,到底省在哪。 KV cache 有一个和别的张量都不一样的地方:**它是在请求进行中一点点长出来的**。 你事先不知道这条请求最终会有多长,所以要么按最大长度预留(浪费), 要么让它能非连续地增长(分页)。这个脚本把两种做法放在同一个显存预算下对比。 [A] 同一个显存预算,两种分配策略各能同时服务多少条请求 [B] 浪费来自哪里:预留浪费 vs 块内碎片 [C] 块大小怎么选:越小越省,但不是越小越好 python paged_alloc.py ALL """ from __future__ import annotations import argparse import json import math import os import numpy as np HERE = os.path.dirname(os.path.abspath(__file__)) # Llama-3-8B / BF16 在假设 80 GiB 总预算下的账(不是设备可用显存实测)(与 kv_cache_lab.py [B] 同源) N_LAYER, N_KV, D_HEAD, DTYPE = 32, 8, 128, 2 BYTES_PER_TOKEN = 2 * N_LAYER * N_KV * D_HEAD * DTYPE # 131072 = 128 KiB GPU_BYTES = 80 * (1024 ** 3) WEIGHT_BYTES = 8.03e9 * DTYPE OTHER_BYTES = 4 * (1024 ** 3) # 激活 / 框架 / 碎片余量 MAX_LEN = 4096 def budget_tokens(): return int((GPU_BYTES - WEIGHT_BYTES - OTHER_BYTES) // BYTES_PER_TOKEN) def workload(kind, n=4000, seed=11): """两种典型负载:长度均匀的(批处理)和重尾的(线上对话)。""" rng = np.random.default_rng(seed) if kind == "uniform": lens = rng.integers(512, MAX_LEN + 1, n) else: # heavy:八成短请求、两成长请求 short = rng.integers(256, 1025, int(n * 0.8)) long_ = rng.integers(3072, MAX_LEN + 1, n - int(n * 0.8)) lens = np.concatenate([short, long_]) rng.shuffle(lens) return lens def admit(lens, bytes_each, budget): """贪心接纳:按请求顺序一直加,直到预算装不下。返回接纳条数。""" used, k = 0, 0 for b in bytes_each: if used + b > budget: break used += b k += 1 return k, used def section_A(): print("=" * 72) print("[A] 同一块显存,两种分配策略能同时服务多少条请求") print("=" * 72) cap = budget_tokens() budget = cap * BYTES_PER_TOKEN print(f"\n显存预算:80 GiB - 权重 {WEIGHT_BYTES / (1024**3):.2f} GiB" f" - 其他 {OTHER_BYTES / (1024**3):.0f} GiB = {budget / (1024**3):.2f} GiB") print(f"折合 token 容量 : {cap} tokens({BYTES_PER_TOKEN} B/token)") out = {} for kind, label in [("uniform", "长度均匀 512~4096"), ("heavy", "重尾:八成 256~1024")]: print(f"\n--- 负载:{label} ---") lens = workload(kind) b_cont = np.full(len(lens), MAX_LEN * BYTES_PER_TOKEN) # 按最大长度预留 b_paged = (np.ceil(lens / 16) * 16) * BYTES_PER_TOKEN # 分页,块 16 k_c, u_c = admit(lens, b_cont, budget) k_p, u_p = admit(lens, b_paged, budget) used_tok_c = lens[:k_c].sum() used_tok_p = lens[:k_p].sum() util_c = used_tok_c * BYTES_PER_TOKEN / u_c util_p = used_tok_p * BYTES_PER_TOKEN / u_p print(f" {'策略':>12} {'并发请求':>9} {'占用':>10} {'真实用到':>10} {'利用率':>8}") print(f" {'连续预留':>12} {k_c:>9} {u_c / (1024**3):>7.2f} GiB " f"{used_tok_c * BYTES_PER_TOKEN / (1024**3):>7.2f} GiB {util_c:>7.2%}") print(f" {'分页(块16)':>12} {k_p:>9} {u_p / (1024**3):>7.2f} GiB " f"{used_tok_p * BYTES_PER_TOKEN / (1024**3):>7.2f} GiB {util_p:>7.2%}") print(f" 并发提升 : {k_p / k_c:.2f}x") out[kind] = dict(k_cont=int(k_c), k_paged=int(k_p), util_cont=float(util_c), util_paged=float(util_p), gain=float(k_p / k_c), bytes_cont=float(u_c), bytes_paged=float(u_p), real_cont=float(used_tok_c * BYTES_PER_TOKEN), real_paged=float(used_tok_p * BYTES_PER_TOKEN)) out["cap_tokens"] = cap out["budget_bytes"] = float(budget) return out def section_B(): print("\n" + "=" * 72) print("[B] 浪费的两半:预留浪费 和 块内碎片") print("=" * 72) lens = workload("uniform") bs = 16 reserve_waste = (MAX_LEN - lens).mean() * BYTES_PER_TOKEN frag = ((np.ceil(lens / bs) * bs) - lens).mean() * BYTES_PER_TOKEN print(f"\n每条请求的平均长度 : {lens.mean():.1f} tokens") print(f"连续预留的浪费/条 : {reserve_waste / (1024**2):.2f} MiB" f"(预留 {MAX_LEN},平均只用 {lens.mean():.0f})") print(f"分页的块内碎片/条 : {frag / (1024**2):.2f} MiB" f"(只有最后一个块没填满,余数均匀时约 {(bs - 1) / 2} 个空槽)") print(f"两者之比 : {reserve_waste / frag:.1f}x") print("\n分页把「不知道会多长」这个不确定性,从「按最坏情况预留」" "换成了「最多浪费一个块」——这是整个设计的关键一跳。") return dict(reserve_waste=float(reserve_waste), frag=float(frag), ratio=float(reserve_waste / frag), mean_len=float(lens.mean())) def section_C(): print("\n" + "=" * 72) print("[C] 块大小怎么选") print("=" * 72) lens = workload("uniform") mean_len = lens.mean() print(f"\n平均请求长度 {mean_len:.1f} tokens。块越大,最后一块的空槽越多;" f"块越小,块表越长、kernel 里要 Gather 的次数越多。") print(f"\n {'块大小':>6} {'碎片/条':>10} {'理论利用率':>10} {'块表条目/请求':>14}") rows = [] for bs in [1, 4, 8, 16, 32, 64, 128, 256]: frag = ((np.ceil(lens / bs) * bs) - lens).mean() util = mean_len / (mean_len + frag) nblk = float(np.ceil(lens / bs).mean()) rows.append(dict(bs=bs, frag=float(frag), util=float(util), nblk=nblk)) print(f" {bs:>6} {frag:>8.1f} tk {util:>9.3%} {nblk:>14.1f}") print("\n长度模 bs 的余数均匀时,块内空槽期望为 (bs-1)/2;" "真实碎片取决于长度分布。16 或 32 是常见候选,需实测:" "本例平均长度约 2310,16/32 块约有 0.32%/0.67% 空槽;元数据和 kernel 开销另计。") return dict(rows=rows, mean_len=float(mean_len)) def main(): ap = argparse.ArgumentParser() ap.add_argument("which", nargs="?", default="ALL") w = ap.parse_args().which.upper() out = {} if w in ("ALL", "A"): out["A"] = section_A() if w in ("ALL", "B"): out["B"] = section_B() if w in ("ALL", "C"): out["C"] = section_C() if w == "ALL": # 画图用的示意数据:3 条请求怎么被切成块 lens = [37, 21, 45] bs = 16 seqs = [] for i, L in enumerate(lens): seqs.append(dict(idx=i, length=L, n_blocks=int(math.ceil(L / bs)), last_used=L % bs or bs)) out["demo"] = dict(block_size=bs, seqs=seqs) with open(os.path.join(HERE, "_paged_results.json"), "w", encoding="utf-8") as f: json.dump(out, f, ensure_ascii=False, indent=1) print("\n结果已写入 _paged_results.json") if __name__ == "__main__": main() make_figures.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """ make_figures.py —— 画本文的 4 张配图。 数据全部来自已经跑完的实验(_kv_results.json / _paged_results.json), 不在这里重新算,避免图上的数字和正文漂移。 运行: python make_figures.py """ from __future__ import annotations import json import os import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import matplotlib.patches as mpatches import numpy as np HERE = os.path.dirname(os.path.abspath(__file__)) FIGDIR = os.path.join(HERE, "..", "figures") try: os.makedirs(FIGDIR, exist_ok=True) except FileExistsError: pass # 配色(正文里写「这张图要看什么」时按这六个名字来描述) C_MAIN = "#1f4e79" # 深蓝:主曲线 C_ALT = "#c1440e" # 橙红:对照曲线 C_GREEN = "#2e7d32" # 绿:第三组 C_PURPLE = "#7b1fa2" # 紫:标注线 C_GRAY = "#8a8a8a" # 灰:参考线 C_LIGHT = "#bcd7ee" # 浅蓝:填充 C_RED = "#b71c1c" # 红:越界线 plt.rcParams["font.sans-serif"] = ["PingFang SC", "Arial Unicode MS", "DejaVu Sans"] plt.rcParams["axes.unicode_minus"] = False plt.rcParams["figure.dpi"] = 130 plt.rcParams["savefig.dpi"] = 130 def _load(name): p = os.path.join(HERE, name) if not os.path.exists(p): raise SystemExit(f"缺少 {name},先跑 kv_cache_lab.py ALL / paged_alloc.py ALL") with open(p, encoding="utf-8") as f: return json.load(f) GIB = float(1024 ** 3) # ══════════════════════════════════════════════════════════════ # 图 1:显存账本 # ══════════════════════════════════════════════════════════════ def fig_mem_ledger(res): B = res["B"] w = B["weight_bytes"] bpt, bpt_mha = B["bytes_per_token"], B["bytes_per_token_mha"] fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12.6, 4.9)) # ── 左:batch 扫描下的显存构成 ── batches = [1, 8, 16, 32, 64, 128] seq = 8192 kv = [bpt * seq * b for b in batches] others = 4 * GIB xs = np.arange(len(batches)) ax1.bar(xs, [w / GIB] * len(batches), color=C_MAIN, label="模型权重(14.96 GiB)") ax1.bar(xs, [k / GIB for k in kv], bottom=[w / GIB] * len(batches), color=C_ALT, label="KV cache") ax1.bar(xs, [others / GIB] * len(batches), bottom=[(w + k) / GIB for k in kv], color=C_LIGHT, label="激活 / 框架 / 余量") ax1.axhline(80, color=C_RED, ls="--", lw=1.6) ax1.text(0.05, 81.4, "假设总预算 80 GiB", color=C_RED, fontsize=9) ax1.set_xticks(xs) ax1.set_xticklabels([str(b) for b in batches]) ax1.set_xlabel("batch size") ax1.set_ylabel("显存(GiB)") ax1.set_title("显存账本:batch 一大,权重就不再是主角", fontsize=11) ax1.legend(fontsize=8, loc="upper left") ax1.set_ylim(0, 200) # ── 右:序列长度扫描,GQA vs MHA ── seqs = np.array([512, 1024, 2048, 4096, 8192, 16384, 32768, 65536, 131072]) for b in (1, 8): ax2.plot(seqs, bpt * seqs * b / GIB, "-o", ms=3.5, color=C_MAIN if b == 1 else C_ALT, label=f"GQA batch={b}") ax2.plot(seqs, bpt_mha * seqs * 8 / GIB, "--s", ms=3.5, color=C_PURPLE, label="MHA batch=8(kv 头 32 个)") ax2.axhline(61, color=C_GRAY, ls=":", lw=1.4) ax2.text(seqs[0], 66, "KV 预算约 61 GiB", color=C_GRAY, fontsize=8.5) ax2.axhline(80, color=C_RED, ls="--", lw=1.4) ax2.text(seqs[0] * 1.6, 88, "80 GiB 显存上限", color=C_RED, fontsize=8.5) ax2.set_xscale("log", base=2) ax2.set_yscale("log") ax2.set_xlabel("每条序列的长度(token)") ax2.set_ylabel("KV cache(GiB,对数刻度)") ax2.set_title("KV cache 随长度线性增长,随 batch 线性增长", fontsize=11) ax2.legend(fontsize=8) ax2.grid(alpha=0.25, which="both") fig.tight_layout() p = os.path.join(FIGDIR, "fig_mem_ledger.png") fig.savefig(p, bbox_inches="tight") plt.close(fig) print(" ", os.path.basename(p)) # ══════════════════════════════════════════════════════════════ # 图 2:Roofline # ══════════════════════════════════════════════════════════════ def fig_roofline(res): C = res["C"] pf, pb = C["peak_flops"], C["peak_bw"] ridge = pf / pb af, ab = C["a100_flops"], C["a100_bw"] aridge = af / ab pts = [(1, "decode g=1(MHA)", C_RED), (4, "decode g=4(GQA)", C_GREEN), (8, "decode g=8", C_PURPLE), (C["I_prefill"], "prefill S=8192", C_MAIN)] fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12.6, 5.2)) I = np.logspace(-1, 5, 400) # ── 左:本机(实跑标定,GFLOP/s)── ax1.plot(I, np.minimum(pf, I * pb) / 1e9, color=C_MAIN, lw=2.0) ax1.fill_betweenx([1e-2, pf / 1e9 * 2], 1e-1, ridge, color=C_LIGHT, alpha=0.3) ax1.axvline(ridge, color=C_MAIN, ls=":", lw=1.2) ax1.text(ridge * 1.12, 3.0, f"山脊点 {ridge:.0f}", color=C_MAIN, fontsize=8.5, rotation=90, va="bottom") for i, name, col in pts: i = i / 2 # 本机峰值来自 FP32,字节数是 BF16 的两倍 y = min(pf, i * pb) / 1e9 ax1.plot([i], [y], "o", ms=8, color=col, zorder=5) off = (-108, -20) if i >= 100 else (7, -13 if i < 100 else -3) ax1.annotate(name, (i, y), textcoords="offset points", xytext=off, fontsize=8.5, color=col) ax1.text(0.13, pf / 1e9 * 1.35, "带宽受限区", color=C_MAIN, fontsize=9) ax1.set_xscale("log") ax1.set_yscale("log") ax1.set_xlim(0.1, 1e4) ax1.set_ylim(1.0, pf / 1e9 * 2.2) ax1.set_xlabel("算术强度 I(FLOP / byte)") ax1.set_ylabel("理论上界(GFLOP/s)") ax1.set_title(f"本机 FP32 微基准屋顶:{pf / 1e9:.0f} GFLOP/s / {pb / 1e9:.0f} GB/s", fontsize=10.5) ax1.grid(alpha=0.25, which="both") # ── 右:A100(公开规格,TFLOP/s)── ax2.plot(I, np.minimum(af, I * ab) / 1e12, color=C_ALT, lw=2.0) ax2.fill_betweenx([1e-2, af / 1e12 * 2], 1e-1, aridge, color="#f6d9c9", alpha=0.45) ax2.axvline(aridge, color=C_ALT, ls=":", lw=1.2) ax2.text(aridge * 1.12, 0.9, f"山脊点 {aridge:.0f}", color=C_ALT, fontsize=8.5, rotation=90, va="bottom") for i, name, col in pts: y = min(af, i * ab) / 1e12 ax2.plot([i], [y], "o", ms=8, color=col, zorder=5) off = (-108, -20) if i >= 100 else (7, -13 if i < 100 else -3) ax2.annotate(name, (i, y), textcoords="offset points", xytext=off, fontsize=8.5, color=col) ax2.text(0.13, af / 1e12 * 1.35, "带宽受限区", color=C_ALT, fontsize=9) ax2.set_xscale("log") ax2.set_yscale("log") ax2.set_xlim(0.1, 1e5) ax2.set_ylim(0.3, af / 1e12 * 2.2) ax2.set_xlabel("算术强度 I(FLOP / byte)") ax2.set_ylabel("理论上界(TFLOP/s)") ax2.set_title(f"A100-80GB(公开规格):{af / 1e12:.0f} TFLOP/s / {ab / 1e12:.2f} TB/s", fontsize=10.5) ax2.grid(alpha=0.25, which="both") fig.suptitle("理想融合注意力的 Roofline 上界(点不是实测 kernel 性能)", fontsize=12, y=1.02) fig.tight_layout() p = os.path.join(FIGDIR, "fig_roofline.png") fig.savefig(p, bbox_inches="tight") plt.close(fig) print(" ", os.path.basename(p)) # ══════════════════════════════════════════════════════════════ # 图 3:PagedAttention 的分页布局 + 并发对比 # ══════════════════════════════════════════════════════════════ def fig_paged(res): demo = res.get("demo", {}) bs = demo.get("block_size", 16) lens = [s["length"] for s in demo.get("seqs", [])] or [37, 21, 45, 29] colors = [C_MAIN, C_ALT, C_GREEN, C_PURPLE, C_GRAY] n_seq = len(lens) # 模拟「边生成边分配」:每条请求轮流出 1 个 token,块不够了才申请新的 pos = [0] * n_seq blocks = [[] for _ in range(n_seq)] # 逻辑块 -> 物理块号 phys_owner, phys_fill, phys_cap = [], [], [] while any(pos[i] < lens[i] for i in range(n_seq)): for i in range(n_seq): if pos[i] >= lens[i]: continue if pos[i] % bs == 0: # 需要一个新的物理块 phys_owner.append(i) phys_fill.append(0) phys_cap.append(bs) blocks[i].append(len(phys_owner) - 1) phys_fill[blocks[i][-1]] += 1 pos[i] += 1 n_phys = len(phys_owner) + 3 # 末尾留 3 个空块 fig = plt.figure(figsize=(12.6, 5.0)) gs = fig.add_gridspec(1, 2, width_ratios=[1.25, 1.0]) # ── 左:逻辑视图 -> 物理块池 ── axL = fig.add_subplot(gs[0, 0]) axL.set_xlim(-0.6, n_phys + 0.6) axL.set_ylim(-0.5, n_seq + 2.6) axL.axis("off") axL.set_title("逻辑块 → 物理块:请求可以非连续地长", fontsize=11, loc="left") bw_l, bh = 0.82, 0.62 # 逻辑视图:每条请求一行,块连续排列 for i in range(n_seq): y = n_seq - i + 1.35 axL.text(-0.5, y + bh / 2, f"请求 {i + 1}", fontsize=8.5, ha="right", va="center") for j, p in enumerate(blocks[i]): x = j * (bw_l + 0.06) axL.add_patch(mpatches.Rectangle( (x, y), bw_l, bh, facecolor=colors[i], alpha=0.30, edgecolor=colors[i], lw=1.2)) used = phys_fill[p] axL.add_patch(mpatches.Rectangle( (x, y), bw_l * used / bs, bh, facecolor=colors[i], alpha=0.85)) if used < bs: axL.text(x + bw_l * used / bs / 2, y + bh / 2, f"{used}/{bs}", fontsize=6.4, ha="center", va="center", color="white") # 物理块池:一行,按申请顺序排列 y_p = 0.15 axL.text(-0.5, y_p + bh / 2, "物理块池", fontsize=8.5, ha="right", va="center") for p in range(n_phys): x = p * (bw_l + 0.06) if p < len(phys_owner): c = colors[phys_owner[p]] axL.add_patch(mpatches.Rectangle( (x, y_p), bw_l, bh, facecolor=c, alpha=0.30, edgecolor=c, lw=1.2)) axL.add_patch(mpatches.Rectangle( (x, y_p), bw_l * phys_fill[p] / bs, bh, facecolor=c, alpha=0.85)) axL.text(x + bw_l / 2, y_p - 0.28, str(p), fontsize=6.4, ha="center", color=C_GRAY) else: axL.add_patch(mpatches.Rectangle( (x, y_p), bw_l, bh, facecolor="white", edgecolor=C_GRAY, lw=1.0, hatch="//")) axL.text(x + bw_l / 2, y_p - 0.28, str(p), fontsize=6.4, ha="center", color=C_GRAY) axL.text(n_phys * (bw_l + 0.06) + 0.1, y_p + bh / 2, "空闲", fontsize=8, va="center", color=C_GRAY) # 几条连线:从逻辑块指到它真正的物理块 for i in range(n_seq): for j, p in enumerate(blocks[i]): x0 = j * (bw_l + 0.06) + bw_l / 2 y0 = n_seq - i + 1.35 x1 = p * (bw_l + 0.06) + bw_l / 2 axL.annotate("", xy=(x1, y_p + bh), xytext=(x0, y0), arrowprops=dict(arrowstyle="-", color=C_GRAY, lw=0.7, alpha=0.55)) axL.text(0, n_seq + 2.2, f"块大小 {bs}:只有每条请求的最后一块没填满(灰底数字是「已用/容量」)", fontsize=8.5, color=C_GRAY) # ── 右:并发对比 ── axR = fig.add_subplot(gs[0, 1]) A = res["A"] labels = ["长度均匀\n512~4096", "重尾\n八成短请求"] kc = [A["uniform"]["k_cont"], A["heavy"]["k_cont"]] kp = [A["uniform"]["k_paged"], A["heavy"]["k_paged"]] uc = [A["uniform"]["util_cont"], A["heavy"]["util_cont"]] up = [A["uniform"]["util_paged"], A["heavy"]["util_paged"]] xs = np.arange(2) wbar = 0.36 b1 = axR.bar(xs - wbar / 2, kc, wbar, color=C_GRAY, label="连续预分配(按 4096 预留)") b2 = axR.bar(xs + wbar / 2, kp, wbar, color=C_MAIN, label="分页(块 16)") for i, (a, b) in enumerate(zip(kc, kp)): axR.text(i - wbar / 2, a + 6, f"{a}\n利用率 {uc[i]:.1%}", ha="center", fontsize=8) axR.text(i + wbar / 2, b + 6, f"{b}\n利用率 {up[i]:.1%}", ha="center", fontsize=8, color=C_MAIN) axR.text(i, max(a, b) + 58, f"{b / a:.2f}x", ha="center", fontsize=10, color=C_RED, fontweight="bold") axR.set_xticks(xs) axR.set_xticklabels(labels, fontsize=9) axR.set_ylabel("同一 61 GiB 预算下的并发请求数") axR.set_title("分页减少预留浪费,提高静态可容纳请求数", fontsize=11) axR.legend(fontsize=8) axR.set_ylim(0, max(kp) * 1.42) axR.grid(alpha=0.2, axis="y") fig.tight_layout() p = os.path.join(FIGDIR, "fig_paged.png") fig.savefig(p, bbox_inches="tight") plt.close(fig) print(" ", os.path.basename(p)) # ══════════════════════════════════════════════════════════════ # 图 4:自回归视频的掩码结构 + 缓存规模 # ══════════════════════════════════════════════════════════════ def fig_video(res): B = res["B"] fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12.6, 4.9)) # ── 左:块内双向 + 块间因果的掩码 ── n_frame, tok_per_frame = 5, 12 n = n_frame * tok_per_frame M = np.zeros((n, n)) for a in range(n_frame): for b in range(n_frame): if b <= a: # 只能看当前帧和之前的帧 M[a * tok_per_frame:(a + 1) * tok_per_frame, b * tok_per_frame:(b + 1) * tok_per_frame] = 1.0 ax1.imshow(M, cmap="Blues", vmin=0, vmax=1.4, interpolation="nearest") for f in range(n_frame + 1): ax1.axhline(f * tok_per_frame - 0.5, color=C_ALT, lw=1.0) ax1.axvline(f * tok_per_frame - 0.5, color=C_ALT, lw=1.0) ax1.set_xlabel("key / value 位置(时间从前到后)") ax1.set_ylabel("query 位置") ax1.set_title("块内双向、块间因果", fontsize=11) ax1.set_xticks([f * tok_per_frame + tok_per_frame / 2 for f in range(n_frame)]) ax1.set_xticklabels([f"帧{f + 1}" for f in range(n_frame)], fontsize=8) ax1.set_yticks([f * tok_per_frame + tok_per_frame / 2 for f in range(n_frame)]) ax1.set_yticklabels([f"帧{f + 1}" for f in range(n_frame)], fontsize=8) ax1.text(1, n - 3, "对角块 = 块内双向示例\n(扩散当前块仍须重算)", fontsize=8.5, color=C_ALT, bbox=dict(fc="white", ec=C_ALT, alpha=0.85)) ax1.text(n * 0.34, n * 0.12, "下三角块 = 能看到过去帧\n(仅固定且兼容的历史可缓存)", fontsize=8.5, color=C_MAIN, bbox=dict(fc="white", ec=C_MAIN, alpha=0.85)) # ── 右:单条视频的 KV cache 规模 ── cases = B["video"] names = [c["name"] for c in cases] vals = [c["kv"] / GIB for c in cases] toks = [c["ntok"] for c in cases] bars = ax2.bar(names, vals, color=[C_GREEN, C_ALT, C_RED], width=0.55) ax2.axhline(80, color=C_RED, ls="--", lw=1.6) ax2.axhline(61, color=C_GRAY, ls=":", lw=1.4) ax2.text(2.52, 83, "假设总预算 80 GiB", color=C_RED, fontsize=8.5, ha="right", bbox=dict(fc="white", ec="none", alpha=0.9)) ax2.text(2.52, 64, "扣除权重与余量后 KV 约 61 GiB", color=C_GRAY, fontsize=8.5, ha="right", bbox=dict(fc="white", ec="none", alpha=0.9)) for b, v, t in zip(bars, vals, toks): ax2.text(b.get_x() + b.get_width() / 2, v + 10, f"{v:.1f} GiB\n{t // 1000}k tokens", ha="center", fontsize=8.5) ax2.set_ylabel("单条视频的 KV cache(GiB)") ax2.set_title("720p 假设账本:无额外 patch 化,缓存全部历史", fontsize=11) ax2.set_ylim(0, max(vals) * 1.30) ax2.grid(alpha=0.2, axis="y") fig.tight_layout() p = os.path.join(FIGDIR, "fig_video_mask.png") fig.savefig(p, bbox_inches="tight") plt.close(fig) print(" ", os.path.basename(p)) def main(): kv = _load("_kv_results.json") pg = _load("_paged_results.json") print("画图:") fig_mem_ledger(kv) fig_roofline(kv) fig_paged(pg) fig_video(kv) if __name__ == "__main__": main()
2026年10月02日
3 阅读
0 评论
0 点赞
2026-10-02
AIGC 基本功|FID / CLIP Score 到底测了什么-FID
FID / CLIP Score 到底测了什么 所属方向:评测 | 难度:进阶 | 前置知识:无(本篇自洽,只需要你见过「协方差矩阵」和「余弦相似度」这两个词) 关键词:FID、Inception Score、CLIP Score、2-Wasserstein 距离、样本量偏差、指标失效 01. 为什么需要它 先给一个我自己跑出来的数字,它比任何论述都更能说明问题。 我从同一个分布里抽两组样本:真实组和生成组来自同一个人工构造的 2048 维多元高斯:协方差特征值按 $1/k$ 衰减,迹归一化为 2048。总体的高斯距离为 0,但两组有限样本的估计值不必为 0。每组 10000 个样本时,实测 FID = 71.20;每组加到 50000 个样本,实测 FID = 14.21。两组的总体分布相同,但具体抽样不同。这里改变样本量并重复抽样取平均;不是对真实 Inception 图像特征的测量。 这不是实现写错了,而是 FID 的固有性质:它是一个自带正偏差的估计量。偏差来自哪里?Inception-v3 的 pool3 特征是 2048 维,一个 2048×2048 的协方差矩阵有 $2048 \times 2049 / 2 \approx 2.10 \times 10^{6}$ 个自由参数,全靠有限样本去填。有限样本让均值与协方差发生波动,经非线性距离公式后产生估计偏差。同分布基线的期望非负,但不能把自由参数数量直接当作最低样本量。我做过分解(附录 fid_bias_lab.py 的 [A2] 段):在 $d=512$、$n=40000$ 时,把均值换成真值只能消掉 1.4% 的偏差,把协方差换成真值能消掉 98.6%——在这组谱和尺度下,偏差主要来自协方差估计。 这些数字展示了样本量可以显著改变同分布基线,但不能把 71 或 57 分推广为所有真实模型的偏差。特征整体乘 $c$,距离就乘 $c^2$;谱结构、参考集是否固定、两个分布的差异也会影响偏差。因此跨论文比较前必须核对特征器、数据集和样本量等协议。Chong 与 Forsyth 的研究 还说明偏差依赖生成器,同样的样本量并不能自动消除排序偏差。 第二个坑在另一头。FID 只比较分布,不比较单张图。例如,若只是把同一组图像与 prompt 重新错配,图像集合不变,FID 就完全不变;但把所有橘猫换成黑猫可能改变图像分布,不能保证 FID 不变。因此文生图评测常配合 CLIP 相似度或其他条件一致性指标——它不需要真实图片做参考(reference-free),能逐样本给分。但 CLIP Score 有自己的洞:它先给每个图文对打分,常见的数据集均值会丢失分布信息,好样本和坏样本可以互相平均掉。在合成共享空间里,可以构造汇总分数近似相同、逐样本分布却不同的两组输出。第 6.3 节列出四组例子;标准差比最大约 17.55 倍对应 $p=0.9$,并非 $p=0.3$ 那一行。分数一样,产品体验完全不同。 所以这篇要讲清楚的是三件事:FID 的闭式解是怎么推出来的、它的偏差有多大且怎么补救、以及 FID 和 CLIP Score 各自测的是哪一半。 02. 最小可用理解 三句话: 机制:把真实图和生成图都过一遍 Inception-v3,取 pool3 层的 2048 维特征;对两组特征各拟合一个多元高斯 $\mathcal{N}(\mu_{\text{real}}, \Sigma_{\text{real}})$ 和 $\mathcal{N}(\mu_{\text{gen}}, \Sigma_{\text{gen}})$;然后算这两个高斯之间的 Fréchet 距离(也就是 2-Wasserstein 距离的平方)。全部闭式,几行代码。 成本:只需要一阶矩和二阶矩,不需要知道分布的形状,也不需要先训一个判别器。代价是它只看得见前两阶矩,有限样本协方差可以计算,但估计误差会进入 FID;大样本区间常用 $1/n$ 展开描述偏差,系数依赖具体分布。 效果与代价:FID 对改变前两阶矩的分布变化敏感,也可能漏掉矩匹配的模式变化;CLIPScore 提供逐图文对相似度,但不等于综合质量。本文合成例子说明两个目标可能冲突,不能据此断言真实 FID 与 CLIPScore 永远相反。 03. 数学推导 3.1 为什么不能逐图打分 生成模型的评测有一个结构性困难:无条件生成或自由文生图评测通常没有逐图配对的真值。你拿不到「这张生成图对应的真值图」,因此 MSE、PSNR、SSIM 等有参考指标不能直接用于这种非配对比较——它们要求两张图逐像素对齐。 一个自然的替代是「只给生成图打分」,这就是 Inception Score(IS)的思路:把生成图送进 Inception-v3,看分类分布 $p(y \mid x)$ 是不是既尖锐又有多样性。但它有一个致命缺陷:它根本不看真实图片。记忆训练集也可能得到高 IS,因为 IS 不检查是否抄袭,也不比较目标数据分布。高 IS 还要求预测类别清晰且类别边缘分布有多样性,并非真实图片自动「满分」。FID 同样不直接检验记忆训练集。 FID 的出发点就是要补上这一半:把真实分布也纳入比较,比较两个分布之间的距离,而不是给单张图打分。 3.2 为什么是高斯 图像在 Inception 特征空间(2048 维)里的分布形状未知,而且在这个维度上你没法可靠地估计它的形状。FID 选择用一阶矩和二阶矩近似描述,估计质量仍依赖样本量。 那么问题变成:在只知道均值和协方差的条件下,应该假设什么分布?答案是高斯——它是给定前两阶矩时熵最大的分布,也就是「在已知信息下最不作额外假设」的那个选择。这个选择的物理含义是:FID 只承诺比较前两阶矩,不承诺比较形状。后面 3.5 节会看到,这个妥协是有代价的。 3.3 Fréchet 距离的闭式解 设 $X \sim \mathcal{N}(\mu_1, \Sigma_1)$、$Y \sim \mathcal{N}(\mu_2, \Sigma_2)$。2-Wasserstein 距离的平方定义为所有耦合(joint distribution)中传输代价的最小值: $$W_2^2 = \min_{\text{coupling}} \mathbb{E} \big[ \| X - Y \|^2 \big]$$ 先看任意一个耦合的代价是多少。设交叉协方差 $C = \mathrm{Cov}(X, Y)$,把 $\|X - Y\|^2$ 展开成三项并分别取期望:第一项 $\mathbb{E}\|X\|^2 = \|\mu_1\|^2 + \mathrm{Tr}(\Sigma_1)$(因为 $\mathrm{Tr}(\Sigma_1)$ 就是 $X$ 各维方差之和);第二项同理;第三项交叉项 $\mathbb{E}[X^\top Y] = \mu_1^\top \mu_2 + \mathrm{Tr}(C)$。三项合并,$\|\mu_1\|^2 + \|\mu_2\|^2 - 2\mu_1^\top\mu_2$ 正好凑成 $\|\mu_1 - \mu_2\|^2$,于是 $$\mathbb{E} \big[ \| X - Y \|^2 \big] = \| \mu_1 - \mu_2 \|^2 + \mathrm{Tr}(\Sigma_1) + \mathrm{Tr}(\Sigma_2) - 2\,\mathrm{Tr}(C)$$ 前三项由边缘分布决定,动不了。所以要让传输代价最小,等价于让 $\mathrm{Tr}(C)$ 最大。约束是联合协方差矩阵必须半正定: $$\begin{pmatrix} \Sigma_1 & C \\ C^\top & \Sigma_2 \end{pmatrix} \succeq 0$$ 这个半正定约束下的迹最大化有闭式解(协方差补全问题的标准结果)。$\Sigma_1$ 可逆时最优解为 $$C^\star = \Sigma_1^{1/2} \big( \Sigma_1^{1/2} \Sigma_2 \Sigma_1^{1/2} \big)^{1/2} \Sigma_1^{-1/2}$$ 代回去,利用迹的循环性质把外面的 $\Sigma_1^{1/2}$ 和 $\Sigma_1^{-1/2}$ 抵消掉,得到 $$\mathrm{Tr}(C^\star) = \mathrm{Tr} \Big( \big( \Sigma_1^{1/2} \Sigma_2 \Sigma_1^{1/2} \big)^{1/2} \Big) = \mathrm{Tr} \Big( \big( \Sigma_1 \Sigma_2 \big)^{1/2} \Big)$$ 最后一个等号值得停一下,因为它是实现环节最容易写错的地方。$\Sigma_1 \Sigma_2$ 两个对称矩阵的乘积一般不是对称矩阵,不能直接交给对称特征值求解器。当 $\Sigma_1$ 正定时,$\Sigma_1 \Sigma_2$ 与 $\Sigma_1^{1/2} \Sigma_2 \Sigma_1^{1/2}$ 相似: $$\Sigma_1^{1/2} \big( \Sigma_1^{1/2} \Sigma_2 \Sigma_1^{1/2} \big) \Sigma_1^{-1/2} = \Sigma_1 \Sigma_2$$ 而后者 $\Sigma_1^{1/2} \Sigma_2 \Sigma_1^{1/2}$ 是对称半正定的(对任意 $v$ 有 $v^\top \Sigma_1^{1/2} \Sigma_2 \Sigma_1^{1/2} v = (\Sigma_1^{1/2} v)^\top \Sigma_2 (\Sigma_1^{1/2} v) \ge 0$)。相似矩阵特征值相同,而主平方根与相似变换可交换,所以两者的平方根迹相等: $$\mathrm{Tr} \big( (\Sigma_1 \Sigma_2)^{1/2} \big) = \sum_i \sqrt{\lambda_i}, \quad \lambda_i = \text{eig} \big( \Sigma_1^{1/2} \Sigma_2 \Sigma_1^{1/2} \big)$$ 奇异协方差可由连续性取极限,实际计算仍使用半正定夹心矩阵而无需显式求逆。 把所有项拼起来,就是 FID 的完整定义: $$\mathrm{FID} = \| \mu_{\text{real}} - \mu_{\text{gen}} \|^2 + \mathrm{Tr}(\Sigma_{\text{real}}) + \mathrm{Tr}(\Sigma_{\text{gen}}) - 2\,\mathrm{Tr} \Big( \big( \Sigma_{\text{real}} \Sigma_{\text{gen}} \big)^{1/2} \Big)$$ 三项的物理含义分别是:均值项管「两组图的平均特征偏了多远」(对应内容/风格的整体偏移),两个迹项管「各自铺开得多宽」(对应多样性),交叉项管「两者铺开的形状有多重合」。注意如果两组分布只是整体缩放 $c$ 倍,那么 $\mu$ 变 $c$ 倍、$\Sigma$ 变 $c^2$ 倍,四项一起变 $c^2$ 倍——FID 不是尺度不变的,这一点后面会用到。 3.4 对称求解器必须使用对称半正定输入 上面那个「最后一个等号」在实现里就是一道坎。看看三种写法差多少(附录 fid_core.py 的 [2][3] 段,随机生成的对称正定 $\Sigma_1, \Sigma_2$): $d$ 对称化路线(正确) eigh(Σ1Σ2)(错误) 通用特征值参考 8 11.5957639872 10.8132126304 11.5957639872 32 41.7985007127 40.2403835447 41.7985007127 128 173.3801162800 165.4519506579 173.3801162800 本表里错误写法偏小,且绝对误差随所选维度增加;这不是任意矩阵都成立的单调律。原因很具体:np.linalg.eigh 是专供对称矩阵的求解器,它只读矩阵的上三角(或下三角)并假设输入对称。按 NumPy 官方文档,默认 UPLO="L" 只读下三角并按其镜像解释上三角,并不是计算 $(A+A^\top)/2$。它与半正定夹心矩阵不是一回事。 误差会原样传进 FID。同一对 128 维特征($n=4096$),正确路线算出 FID = 45.584468,错误路线算出 50.484178,差 +4.899711。这个量级足以让你以为模型退化了。 3.5 FID 只看前两阶矩,所以有结构性盲区 3.2 节的高斯假设现在来收账了。构造一个极端例子: 真实分布 $P = \frac{1}{2}\mathcal{N}(+m, I) + \frac{1}{2}\mathcal{N}(-m, I)$——两个分离的模式; 生成分布 $Q = \mathcal{N}(0, I + m m^\top)$——把两个模式糊成一团的单个高斯。 两者的均值都是 0,协方差都是 $I + m m^\top$。前两阶矩完全一样,所以 FID 的真值严格等于 0,不管两个模式离多远。 但人(或者一个简单的分类器)一眼就能看出区别。实测(附录 fid_bias_lab.py 的 [D] 段,$d=32$、$n=40000$;记 $a = \|m\|$,横轴是两个模式中心的间距 $2a$): 模式间距 $2a$ FID(总体真值) FID(经验估计) 贝叶斯最优 AUC 1-NN 双样本检验准确率 二次特征 ridge AUC 2 1.42e-14 0.0156 0.5365 0.4956 0.5022 4 2.84e-14 0.0162 0.6601 0.5056 0.5003 8 −2.84e-14 0.0163 0.8106 0.6204 0.4972 12 0.0 0.0165 0.8739 0.7010 0.4977 16 8.53e-14 0.0239 0.9033 0.7518 0.5042 FID 从头到尾是 0(那 0.015~0.024 是前面说的有限样本估计误差,不是信号),而贝叶斯最优判别器的 AUC 已经到了 0.9033,1-NN 双样本检验准确率到了 0.7518(0.5 表示完全无法区分)。两个模式明明越离越远,FID 一动不动。 更值得玩味的是最后两列:二次特征上的 ridge 分类器 AUC 也一直是 0.50。这不是巧合——这里的平方损失 ridge 在类平衡、矩匹配时缺少均值层面的监督信号。不能推广为所有二次判别器都只看前两阶矩:例如对 $x^2$ 设阈值也可以利用平方值分布的尾部差异。为了确认这一点我做了对照实验:固定模式间距,改成把 $Q$ 的协方差整体放大 $s$ 倍(破坏矩匹配),于是 FID 和二次 ridge AUC 一起抬头: 协方差放大倍数 $s$ FID 二次特征 ridge AUC 1.0 8.53e-14 0.5021 1.05 0.0585 0.5022 1.2 0.8745 0.5198 1.5 4.8490 0.5972 2.0 16.4710 0.7806 本实验中,矩匹配使总体 FID 为零,二次特征 ridge 的测试 AUC 接近随机水平 0.5;改变协方差后两者均发生变化。 1-NN 那种基于局部密度的判别器则不吃这一套——它看的是密度本身的形状,不是矩。 这张图要看什么:左图是 $P$(蓝)和 $Q$(橙)在二维上的真实散点,连同它们各自的 1σ/2σ 椭圆——注意两个椭圆几乎完全重合,这就是「前两阶矩完全一样」的几何含义:二阶统计量把两个分离的团和一个糊在一起的团画成了同一个椭圆。右图是同一个实验扫过模式间距的结果,蓝线(FID,左轴对数刻度)从头到尾贴在 $10^{-14}$ 量级纹丝不动,橙线(贝叶斯最优 AUC)和绿线(1-NN 准确率,右轴)一路爬到 0.90 / 0.75。两条线之间的那片空白,就是 FID 用高斯假设换来的盲区。 04. 代码实现 核心只有三件事:对称矩阵的平方根、$\mathrm{Tr}((\Sigma_1\Sigma_2)^{1/2})$ 的对称化路线、以及协方差怎么估。下面这段是 fid_core.py 的主干(完整版见附录)。 import numpy as np def sqrtm_sym(C: np.ndarray, eps: float = 1e-6) -> np.ndarray: r"""对称半正定矩阵的平方根。 对 C = V diag(w) V^T,有 C^{1/2} = V diag(sqrt(w)) V^T。 eps 是相对谱尺度的负特征值容差,不是给正特征值设置下限。 明显非半正定输入报错;容差内负值截到 0,保留真正的零特征值。 """ C = np.asarray(C, dtype=np.float64) if C.ndim != 2 or C.shape[0] != C.shape[1] or not np.isfinite(C).all(): raise ValueError("C must be a finite square matrix") scale = max(np.linalg.norm(C, ord=np.inf), np.finfo(float).tiny) if not np.allclose(C, C.T, rtol=0.0, atol=eps * scale): raise ValueError("C must be symmetric") w, V = np.linalg.eigh((C + C.T) / 2) if w.min() < -eps * max(np.abs(w).max(), np.finfo(float).tiny): raise ValueError("C must be positive semidefinite") w = np.sqrt(np.clip(w, 0.0, None)) return (V * w) @ V.T def trace_sqrt_product(sigma1: np.ndarray, sigma2: np.ndarray, eps: float = 1e-6) -> float: r"""计算 Tr((Σ1 Σ2)^{1/2}),走对称化路线。 Σ1 正定时,Σ1 Σ2 与 Σ1^{1/2} Σ2 Σ1^{1/2} 相似: Σ1^{1/2} (Σ1^{1/2} Σ2 Σ1^{1/2}) Σ1^{-1/2} = Σ1 Σ2 二者特征值相同,而后者是**对称半正定**的,可以安全用 eigh。 主平方根与相似变换可交换,所以迹也相同: Tr((Σ1 Σ2)^{1/2}) = Σ_i sqrt(λ_i) """ s1 = sqrtm_sym(sigma1, eps) M = s1 @ sigma2 @ s1 M = 0.5 * (M + M.T) # 强制对称,压掉浮点不对称 w = np.linalg.eigvalsh(M) return float(np.sqrt(np.clip(w, 0.0, None)).sum()) 这里先展示平方根与交叉项;完整均值、协方差与距离实现见附录。 四个符号和 3.3 节的推导一一对应:mu1/mu2 是 $\mu_1/\mu_2$,sigma1/sigma2 是 $\Sigma_1/\Sigma_2$,trace_sqrt_product 就是 $\mathrm{Tr}((\Sigma_1\Sigma_2)^{1/2})$,eps 是判定负特征值是否属于舍入误差的相对容差,不是把所有小特征值抬高的正则项。 跑 python fid_core.py 的自检输出(这些数字全部是实跑结果): [1] 恒等性:FID(P, P) 必须为 0 FID(P,P) = -4.263e-14 [2] 对称化路线 vs 错误写法 vs 一般特征值参考实现 d sym(正确) naive(错误) eigvals(参考) 8 11.5957639872 10.8132126304 11.5957639872 32 41.7985007127 40.2403835447 41.7985007127 128 173.3801162800 165.4519506579 173.3801162800 [3] 这个差异会传进 FID:同一对特征,两种写法差多少 FID(sym) = 45.584468 FID(naive) = 50.484178 (差值 +4.899711) [4] 尺度不是不变的:特征整体乘 c,FID 变 c^2 倍 c=0.5 FID= 11.396117 期望 c^2*base= 11.396117 c=2.0 FID= 182.337871 期望 c^2*base= 182.337871 c=4.0 FID= 729.351482 期望 c^2*base= 729.351482 [5] 有偏 vs 无偏协方差(n 越小差得越多) n unbiased biased 差值 256 43.932626 43.765931 -0.166696 1024 10.959338 10.948996 -0.010341 8192 1.518005 1.517825 -0.000180 逐条对一下这几个数为什么要看: [1] 按绝对值和尺度相关容差检查恒等性;0.0、极小正值和极小负值都可能正确。明显超出容差才需要排查。附录还验证奇异、小尺度协方差与非 PSD 输入。 [4] 验证 3.3 节末尾那个推论:$c=2$ 时 $4 \times 45.584468 = 182.337871$,完全吻合。这意味着任何改变特征尺度的预处理都会改变 FID,而且不是线性地改。 [5] 本例两组样本数相同,用 $1/n$ 会把双方协方差一起缩小,因此表中的有偏协方差版本距离略小。协方差无偏不代表 FID 无偏;不同样本量时不能直接套这个单调结论。$n=8192$ 时差 0.00018 可以忽略,$n=256$ 时差 0.167 就不能忽略了。工业实现和手写实现在这上面分叉,跨实现比较时要注意。 05. 工业级实现对照 生产里大家用的是 mseitzer/pytorch-fid(截至 2026-10 的实现为准),核心函数在 src/pytorch_fid/fid_score.py#calculate_frechet_distance。和上面的最小实现有五处差异,每一处都有原因: 1. 矩阵平方根用的是 scipy 而不是 eigh。 官方写的是 covmean, _ = linalg.sqrtm(sigma1.dot(sigma2), disp=False) if not np.isfinite(covmean).all(): offset = np.eye(sigma1.shape[0]) * eps covmean = linalg.sqrtm((sigma1 + offset).dot(sigma2 + offset)) if np.iscomplexobj(covmean): if not np.allclose(np.diagonal(covmean).imag, 0, atol=1e-3): raise ValueError("Imaginary component {}".format(np.max(np.abs(covmean.imag)))) covmean = covmean.real scipy.linalg.sqrtm 是通用矩阵的主平方根求解器,不假设输入对称,所以直接喂 $\Sigma_1\Sigma_2$ 是对的——这跟我的对称化路线在数学上等价,只是浮点路径不同。代价是结果可能带极小的虚部(数值误差),所以有了那段 .real 兜底:先看对角线虚部是不是都在 $10^{-3}$ 以内,超了就报错而不是默默取实部。我的最小实现用 eigvalsh 绕开了复数问题,代价是必须自己先把矩阵对称化。真正的坑是有人为了去掉 scipy 依赖把它换成 np.linalg.eigh——那就是 3.4 节那个 +4.9 的错误。 2. 数值容差与正则化不同。 本文只截掉容差内的负特征值并保留零值,以便奇异或小尺度协方差仍满足恒等性;不能无条件把所有特征值抬到 eps。pytorch-fid 在 sqrtm 失败时给协方差加 eps*I 重试,这是改变问题的正则化兜底,应记录是否触发。 3. 协方差用 np.cov(act, rowvar=False),默认无偏。 也就是 04 节 [5] 那张表的第一列。这一点在 $n$ 小的时候会造成跨实现的系统差异。 4. 特征提取有一整套约定。 src/pytorch_fid/inception.py 里的 InceptionV3 取的是第 3 个 block(最终平均池化之后)的 2048 维;resize_input=True 会先把输入双线性缩放到 299×299;use_fid_inception=True 用的是与 TensorFlow 版对齐的 FID 专用权重,而不是 torchvision 默认的 ImageNet 权重。命令行还有 --dims,可以选 64 / 192 / 768 / 2048——换了 dims 就是换了另一个指标,数字不可比,这是跨论文比较时最常被忽略的一项。 5. 保留数值诊断。 该实现返回原始结果,不自动截零。生产实现也可以先验证误差在容差内再截零,但不能不加检查地掩盖大负值;正好为 0 并不说明出错。 还有一个官方不管、但你必须自己管的事:预处理的一致性。Parmar 等人在 arXiv:2104.11222 里指出,resize 是否抗锯齿、图片是否被 JPEG 量化,都会显著改变 FID 的数值——大到足以改变两篇论文的排序。所以在比对任何两个 FID 之前,先确认两边的 Inception 权重、dims、resize 方式、量化流程、样本量、协方差是否有偏,这六项是不是一致。六项里任何一项对不上,两个 FID 就不是同一个东西。 06. 代价与边界 6.1 偏差有多大:随 $1/n$ 衰减,随维度放大 把 01 节那个实验做全(附录 fid_bias_lab.py 的 [A] 段,$P$ 和 $Q$ 是同一个分布,所以真值是 0): 每组样本量 $n$ FID($d=512$) FID($d=2048$) 50 431.54 2131.57 1000 53.85 640.65 10000 5.39 71.20 50000 1.09 14.21 两个观察: 双对数坐标下这些点近乎落在一条直线上,斜率接近 $-1$,也就是偏差按 $\sim 1/n$ 衰减。在本实验的大样本区间,样本量翻倍时平均偏差约减半。 维度从 512 涨到 2048(4 倍),$n=10000$ 处的偏差从 5.40 涨到 71.14(13 倍)。把 $d$ 从 64 扫到 2048 做拟合([B] 段,固定 $n=10000$),实测指数约为 1.814: $d$ 64 128 256 512 1024 2048 FID($P$,$P$) 0.1326 0.4517 1.5584 5.3967 19.5333 71.1371 这里约 $d^{1.8}$ 的拟合只适用于所选的 $1/k$ 协方差谱、迹归一化与扫描范围。真实 Inception 的谱和尺度不同,不能把这个指数当作通用样本量定律。 这张图要看什么:左图是 $P=Q$(真值 0)时 FID 随样本量的变化,双对数刻度,蓝线 $d=512$、橙线 $d=2048$,灰色虚线是斜率 $-1$ 的参考——两条实测线几乎与它平行,这就是「偏差按 $1/n$ 衰减」的直接证据;注意橙线在 $n=50000$ 时还有 14.21,远没有收敛到 0。右图固定 $n=10000$ 扫维度,同样是对数刻度,拟合斜率约 1.814,意味着维度翻一倍偏差涨约 3.5 倍;最右端采用了与 Inception pool3 相同的维度,但使用的是合成高斯特征,并非实际图像嵌入。 6.2 能不能把偏差外推掉 在偏差近似服从 $1/n$ 展开的样本区间,可以尝试外推;需检验拟合稳定性。Chong & Forsyth(arXiv:1911.07023)的做法是:用 $n = N, N/2, N/4, N/8$ 四个点算四个 FID,对 $1/n$ 做线性拟合 $\mathrm{FID}(n) \approx F_{\infty} + \beta / n$,截距 $F_{\infty}$ 就是外推到无穷样本量的估计。 实测([C] 段,$d=512$): $P = Q$(真值 0):$n=40000$ 估 1.3473、$n=5000$ 估 10.8717,外推得 $F_{\infty} = -0.0121$。直接报 $n=40000$ 的数字(1.35)比外推差得多。 $P \neq Q$(真值 0.3635):$n=40000$ 直接估 1.7430,误差 +1.3795;外推得 0.4845,误差 +0.1210。 外推把误差压掉了约 11 倍。代价是要多算三次特征统计量(不过统计量可以复用——从大样本里抽子集就行,不用重新过一遍 Inception)。 报告时同时给参考/生成样本数、原始 FID 和完整协议;使用外推还要报告拟合点、重复采样与不确定性。负截距反映估计误差,不是负的总体距离;外推不保证每个有限样本实验都更准。 6.3 FID 和 CLIP Score 在给不同的东西打高分 CLIP Score 的定义(Hessel 等,arXiv:2104.08718)比 FID 简单得多: $$\mathrm{CLIP\text{-}S} = w \cdot \max \big( \cos(f_{\text{img}}, f_{\text{txt}}), 0 \big), \quad w = 2.5$$ 其中 $f_{\text{img}}$ 和 $f_{\text{txt}}$ 是 CLIP 的图像/文本编码,$w=2.5$ 只是把数值放大到好读的量级。原始 CLIPScore 为图像描述评价提出,reference-free 指不需要人工参考描述;它仍需要待评图像和文本。迁移到文生图时不需要配对真值图。这里按原论文 $w=2.5$,其他实现也会用 100 等缩放,必须注明模型和约定。 下面的人工共享空间反例中,两者最优点不同(附录 clip_alignment_lab.py 的 [A] 段:一个结构化的 CLIP 替身,共享表示空间 $d_s=64$、$K=24$ 个概念,真实数据的多样性固定为 $\sigma_{\text{real}}=0.5$,$n=20000$): 生成多样性 $\sigma_g$ 0.05 0.2 0.4 0.5 0.6 0.8 1.0 FID 10.97 5.24 0.63 0.029 0.65 5.65 15.61 CLIP Score 2.323 1.323 0.744 0.607 0.515 0.399 0.335 FID 在 $\sigma_g = 0.5$(真实数据的多样性)处取最小 0.029,呈 U 形;CLIP Score 单调递减,在 $\sigma_g = 0.05$(几乎退化成确定性输出)处取最大 2.323。该构造下,相似度奖励靠近概念中心,分布距离奖励匹配设定的方差;不能解释为所有 CLIP 模型都偏好确定性。 顺带一提,真实数据自己的 CLIP Score 是 0.6082,而 $\sigma_g=0.5$ 那个「FID 最优」的模型是 0.6073——跟真实数据几乎一样。也就是说在这个实验里,FID 最优点才对应「和真实数据一致」,CLIP Score 的最优点对应的是「模式收敛」。 而且 CLIP Score 的均值性质会掩盖分布。构造两个模型([B] 段):$M_1$ 以概率 $p$ 输出完美匹配的图、以 $1-p$ 输出纯噪声;$M_2$ 每张图都中等匹配(二分调 $\sigma$ 让它的 CLIP Score 与 $M_1$ 相同): $p$ $M_1$ CLIP Score $M_1$ 逐样本标准差 $M_1$ 好图占比 $M_1$ FID $M_2$ $\sigma$ $M_2$ CLIP Score $M_2$ 逐样本标准差 $M_2$ FID 0.3 0.8330 0.4695 0.2980 9.920 0.353 0.8332 0.1078 1.323 0.5 1.3039 0.5078 0.4966 10.191 0.204 1.3016 0.0847 5.112 0.7 1.7896 0.4627 0.7012 10.677 0.122 1.7895 0.0536 8.079 0.9 2.2681 0.2992 0.9023 11.584 0.058 2.2688 0.0171 10.663 两组 CLIP Score 的最大差距只有 0.0023(按构造应该同分),但 $M_1$ 的逐样本相似度标准差是 $M_2$ 的 17.55 倍($p=0.9$ 时 0.2992 对 0.0171)。同一个分数,一个是「九成图完美、一成完全不沾边」,另一个是相似度更集中的输出($p=0.9$ 时均值也很高,不能称为勉强沾边)。汇总均值看不见这个区别,但逐样本 CLIPScore 的直方图或分位数可以显示它。 这张图要看什么:左图是同一个 $\sigma_g$ 扫描下两个指标的走向,蓝线(FID,左轴,越低越好)呈 U 形、在 $\sigma_g=0.5$ 处触底,橙线(CLIP Score,右轴,越高越好)单调下降、在最左端 $\sigma_g=0.05$ 处封顶——两条线的最优点不同,说明这组构造下两个目标存在冲突;0.5 并不是横轴右端,也不能据此断言现实模型的普遍走势;灰色竖虚线标出的是真实数据的多样性 $\sigma_{\text{real}}=0.5$。右图是 $p=0.5$ 那一行两个模型的逐样本相似度分布:橙色的 $M_1$ 是明显的双峰(一半堆在接近 1 的位置、一半堆在 0 附近),蓝色的 $M_2$ 是一根集中在 0.5 附近的单峰,图中直方图直接来自附录实际实验数据,橙/蓝虚线分别是原始余弦均值。CLIPScore 先对每个相似度截零,分数接近不保证原始余弦均值相同;汇总分数仍可能掩盖两种很不一样的样本分布。 6.4 什么时候不该用 FID 样本量小的时候:估计偏差可能淹没模型差异,具体量级取决于数据与特征器。先做同分布基线和重复抽样,再判断是否增样本或尝试外推;没有通用的「10k 以下偏差大于 70」门槛。 要评价单张图的时候:FID 根本没有「单张图」这个概念。需要逐样本打分就上 CLIP Score / ImageReward,但要区分逐样本评分与数据集均值,要配着直方图或者分位数一起看。 要区分「质量」和「覆盖」的时候:FID 把两者压成一个数。一个只生成 10 张高质量图的模型和一个生成 10000 张中等质量图的模型,FID 可能很接近,但产品含义完全不同。这种情况应该用 Improved Precision & Recall(arXiv:1904.06991)拆成两个数。 两个分布形状不同但矩相同的时候:3.5 节的实验——FID 严格为 0,而 1-NN 双样本检验有 0.75 的判别率。 要跨论文比较的时候:除非核实了 05 节末尾那六项完全一致,否则请把数字当成「同一篇论文内部的相对量」,不要当成绝对值。 07. 经典论文脉络 Inception Score(arXiv:1606.03498,Improved Techniques for Training GANs)——第一个被广泛采用的自动指标:用 Inception 的分类分布衡量「单图是否清晰可辨」+「整体是否有多样性」。贡献是让 GAN 评测摆脱了人工打分;根本缺陷是完全不看真实分布,所以无法检测记忆训练集,也无法检测 mode collapse 的另一种形式。 FID(arXiv:1706.08500,TTUR 那篇)——把真实分布拉进比较,用 Inception pool3 特征上的 2-Wasserstein 距离平方评价分布差异。贡献是「两个分布之间的距离」这个范式,在论文实验中展示了对若干退化的敏感性,后来成为常用指标;这种表现不保证覆盖所有退化。 Improved Precision and Recall Metric for Assessing Generative Models(2019)——指出单个标量无法同时表达「生成质量」和「分布覆盖」,拆成 P(生成样本落在真实流形内的比例)和 R(真实样本能被生成覆盖的比例)。贡献是提供了 FID 缺失的那个维度:FID 相同的一对模型,可以在 P/R 平面上处于完全不同、甚至此消彼长的位置。 Effectively Unbiased FID(arXiv:1911.07023)——把 FID 当成一个统计估计量来审视,指出它是有偏的、偏差随 $1/n$ 衰减,并给出用多个样本量外推到 $F_{\infty}$ 的方法。6.2 节那组数字就是照它的做法复现的。它解释了有限样本偏差为何会影响模型排序。 CLIPScore(arXiv:2104.08718)——把评测从「分布 vs 分布」拉回「图 vs 文」,提出不需要参考图的图文对齐指标。贡献是让文生图有了逐样本的自动化对齐分数;局限是相似度不能覆盖全部质量维度,汇总均值也会丢失样本分布;与 FID 的关系依赖实际模型。 补充一条横向的:Borji 的 Pros and Cons of GAN Evaluation Measures(arXiv:1802.03446)系统比较了十几种指标的失效模式,结论是没有任何单一指标能在所有场景下胜出——这也是本篇反复强调「两个指标一起看、连同它们的盲区一起看」的依据。 08. 常见误解 误解 1:「FID 越低,生成的图越好。」 FID 只比较两个分布的前两阶矩,不比较单张图。实测证据:把两个模式越拉越远,FID 一动不动(3.5 节,FID 恒为 0,1-NN 判别率 0.7518)。反过来的方向也成立——FID 很低但每张图都文不对题,是完全可能的。 误解 2:「FID = 0 说明两个分布一样。」 3.5 节整节都在反驳这一点。前两阶矩匹配就够了,形状随便怎么不同都行。高斯假设换来的就是这个。 误解 3:「两篇论文的 FID 可以直接比大小。」 我自己踩过最狠的一个。至少六项要对齐:Inception 权重、特征维度 dims、resize 方式与是否抗锯齿、量化流程、样本量、协方差是否有偏。实测的敏感度:$d=2048$ 时样本量从 10k 到 50k,本文特定人工高斯模型的估计距离从 71.20 变 14.21(差 57);特征整体乘 $c$ 倍,FID 乘 $c^2$ 倍($c=2$ 时 45.58 → 182.34)。这些数字说明 FID 更像一个有单位的物理量,不是一个无量纲分数。 误解 4:「FID 与 CLIP Score 必然反向变化。」 两者测的内容不同,可能同好同坏,也可能冲突。本文合成实验给出冲突的一个例子,不是普遍定律。CFG 的变化也不能精确等同于给合成嵌入加某个固定高斯噪声。 误解 5:「平均 CLIP Score 高,说明每张图都对。」 数据集均值会掩盖尾部。实测两个 CLIP Score 差 0.0023 的模型,逐样本相似度标准差差 17.55 倍(0.2992 对 0.0171)。要看单张图的质量分布,得看直方图或者低分位数,不能只看均值。 误解 6:「FID 非负,所以所有负输出都直接夹为 0。」 应先确认输入和平方根实现,再检查负值是否在尺度相关容差内。小负数可以记录后截零,明显负数必须报错;恰好输出 0 本身既不能证明正确,也不能证明有错。 09. 动手验证 三个数值脚本在文末附录,依赖 numpy;配图另需 matplotlib,不使用真实 Inception/CLIP 权重: python fid_core.py # 约 1 秒:FID 最小实现 + 5 组自检 python fid_bias_lab.py ALL # 约 1 分钟:样本量偏差 / 维度 / 外推 / 矩盲区 python clip_alignment_lab.py ALL # 约 10 秒:FID 与 CLIP Score 的相反最优 预期结果(这些数字是实跑输出,可以直接对): fid_core.py 会按尺度相关容差断言恒等性,包括奇异和小尺度 PSD 输入。数值可以恰好为 0;不要要求固定尾数或符号。明显超出容差时再查矩阵平方根与输入。 fid_bias_lab.py 的 [A] 段,$d=2048$ 那一列:n=10000 应约 71.20、n=50000 应约 14.21($P=Q$,真值 0)。[D] 段最后一行的贝叶斯 AUC 应约 0.9033、1-NN 准确率应约 0.7518,而总体 FID 的数值实现应在零附近的浮点容差内。 clip_alignment_lab.py 的 [A] 段,最优 $\sigma_g$:FID 是 0.5、CLIP Score 是 0.05。[B] 段最后打印的「CLIP Score 最大差距」应约 0.0023、「标准差比值」应约 17.55 倍。 想自己造一个「FID 失效」的例子?改 fid_bias_lab.py 的 [D] 段里那个 a(模式间距的一半)就行:把 $a$ 从 1 扫到 8(中心间距从 2 到 16),总体 FID 保持 0;本例 1-NN 准确率从约 0.50 增至 0.75。 10. 延伸阅读 音频质量评测:MOS、PESQ 与 FAD 各测什么(本系列已发布)——里面的 FAD(Fréchet Audio Distance)就是 FID 换了个特征提取器,高斯矩估计的风险同样存在,但偏差系数与特征器、尺度和样本相关性有关。 视频生成评测:VBench 与人工验收(本系列规划中)——VBench 是多维度视频评测框架,不是简单把 FID 升维;FVD 才是采用视频特征的相关高斯距离指标。两者都不能直接套本文合成模型的 $d^{1.8}$ 指数。 DDPM 训练目标与采样流程、Classifier-Free Guidance(本系列已发布)——CFG 会影响条件一致性与多样性,但不是 6.3 节 $\sigma_g$ 的严格等价参数。实际趋势需要在具体模型上测量。 附录:完整代码 09 节用到的脚本全文如下(fid_bias_lab.py、fid_core.py、clip_alignment_lab.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 fid_bias_lab.py """ fid_bias_lab.py —— FID 的样本量偏差实验。 回答三个问题: [A] 两个分布**完全一样**时,FID 是多少?(答:不是 0,而且 n 越小越离谱) [B] 这个偏差随维度 d、样本量 n 怎么变? [C] 能不能把它外推掉?(Chong & Forsyth, arXiv:1911.07023 的做法) [D] FID 只看前两阶矩,那「前两阶矩完全一样、分布完全不同」能骗过去吗? 运行: python fid_bias_lab.py # 全跑 python fid_bias_lab.py A # 只跑 A 段 """ from __future__ import annotations import json import os import sys import numpy as np from fid_core import fid_from_features, frechet_distance, covariance HERE = os.path.dirname(os.path.abspath(__file__)) OUT_JSON = os.path.join(HERE, "_fid_bias_results.json") # ────────────────────────────────────────────────────────────── # 造一个明确指定谱与尺度的合成协方差:特征值按 1/k 衰减,迹归一到 d # ────────────────────────────────────────────────────────────── def spectrum_cov(d: int, trace: float | None = None, alpha: float = 1.0) -> np.ndarray: r"""对角协方差,特征值 λ_k ∝ k^{-alpha},归一化到 Tr(Σ)=trace。 这是人为选择的谱衰减模型,并非对 Inception pool3 的实测拟合: 少数几个方向撑着大部分方差,长尾方向方差很小。 trace 默认取 d,也就是「每个维度平均方差为 1」。 """ k = np.arange(1, d + 1, dtype=np.float64) lam = k.astype(np.float64) ** (-alpha) if trace is None: trace = float(d) lam = lam * (trace / lam.sum()) return np.diag(lam) def sample_gaussian(mean: np.ndarray, cov_diag: np.ndarray, n: int, rng: np.random.Generator) -> np.ndarray: """从对角协方差的高斯里采样(对角阵直接按列缩放,不用 Cholesky)。""" d = mean.shape[0] lam = np.diag(cov_diag) if cov_diag.ndim == 2 else cov_diag z = rng.standard_normal((n, d)) return mean[None, :] + z * np.sqrt(lam)[None, :] # ────────────────────────────────────────────────────────────── # [A] P == Q 时的 FID # ────────────────────────────────────────────────────────────── def section_a(d_list=(512, 2048), reps=3, verbose=True): print("=" * 74) print("[A] 两个分布完全一样时,FID 不是 0") print(" 真实分布 = 生成分布,理论上 FID = 0。实测:") print("=" * 74) n_list = [50, 100, 200, 500, 1000, 2000, 5000, 10000, 20000, 50000] out = {} for d in d_list: rng = np.random.default_rng(20261001 + d) cov = spectrum_cov(d) mu = np.zeros(d) rows = [] for n in n_list: vals = [] for r in range(reps): x1 = sample_gaussian(mu, cov, n, rng) x2 = sample_gaussian(mu, cov, n, rng) vals.append(fid_from_features(x1, x2)) m = float(np.mean(vals)) s = float(np.std(vals)) rows.append((n, m, s)) if verbose: print(f" d={d:>5} n={n:>6} FID = {m:>10.4f} (std {s:.4f})") out[d] = rows if verbose: print() return {"n_list": n_list, "by_d": {str(k): v for k, v in out.items()}} # ────────────────────────────────────────────────────────────── # [A2] 偏差到底来自均值还是协方差 # ────────────────────────────────────────────────────────────── def section_a2(d=512, n=40000, reps=3, verbose=True): print("=" * 74) print("[A2] 偏差来自哪里:均值还是协方差?(P == Q,真值 0)") print("=" * 74) rng = np.random.default_rng(556677) cov = spectrum_cov(d) mu = np.zeros(d) full, mean_known, cov_known = [], [], [] for _ in range(reps): x1 = sample_gaussian(mu, cov, n, rng) x2 = sample_gaussian(mu, cov, n, rng) m1, s1 = x1.mean(0), covariance(x1) m2, s2 = x2.mean(0), covariance(x2) full.append(frechet_distance(m1, s1, m2, s2)) # 假设均值已知(用真实 mu=0),只估协方差 mean_known.append(frechet_distance(mu, s1, mu, s2)) # 假设协方差已知(用真实 cov),只估均值 cov_known.append(frechet_distance(m1, cov, m2, cov)) f, mk, ck = float(np.mean(full)), float(np.mean(mean_known)), float(np.mean(cov_known)) theory_mean = 2.0 * np.trace(cov) / n if verbose: print(f" d={d} n={n}") print(f" 两个都估(正常做法) FID = {f:>10.4f}") print(f" 均值已知、只估协方差 FID = {mk:>10.4f} " f"占全部偏差的 {mk / f * 100:5.1f}%") print(f" 协方差已知、只估均值 FID = {ck:>10.4f} " f"占全部偏差的 {ck / f * 100:5.1f}%") print(f" 理论值 2*Tr(Sigma)/n = {theory_mean:.6f}(对照上一行)") print() print(" -> 偏差几乎全部来自**协方差估计**。均值那一项理论上就是") print(" 2*Tr(Sigma)/n,小到可以忽略;麻烦的是 d x d 个协方差元素。") print() return {"d": d, "n": n, "full": f, "mean_known": mk, "cov_known": ck, "theory_mean": float(theory_mean)} # ────────────────────────────────────────────────────────────── # [B] 偏差随 d 与 n 的缩放 # ────────────────────────────────────────────────────────────── def section_b(verbose=True): print("=" * 74) print("[B] 偏差随维度 d 怎么长(固定 n=10000)") print("=" * 74) n = 10000 rows = [] for d in (64, 128, 256, 512, 1024, 2048): rng = np.random.default_rng(777000 + d) cov = spectrum_cov(d) mu = np.zeros(d) vals = [] for r in range(3): x1 = sample_gaussian(mu, cov, n, rng) x2 = sample_gaussian(mu, cov, n, rng) vals.append(fid_from_features(x1, x2)) m = float(np.mean(vals)) rows.append((d, m, d / n)) if verbose: print(f" d={d:>5} n={n} d/n={d / n:>7.4f} FID = {m:>10.4f}") print() return {"n": n, "rows": rows} # ────────────────────────────────────────────────────────────── # [C] 外推:FID(n) ≈ F_inf + beta / n # ────────────────────────────────────────────────────────────── def section_c(verbose=True): print("=" * 74) print("[C] 能不能把偏差外推掉?(arXiv:1911.07023 的做法)") print(" 用 n = N, N/2, N/4, N/8 四个点,对 1/n 做线性拟合,截距即 F_inf") print("=" * 74) d = 512 rng = np.random.default_rng(31337) cov = spectrum_cov(d) # C1: P == Q,真值 0 N = 40000 sizes = [N, N // 2, N // 4, N // 8] est = [] for n in sizes: v = [] for r in range(3): x1 = sample_gaussian(np.zeros(d), cov, n, rng) x2 = sample_gaussian(np.zeros(d), cov, n, rng) v.append(fid_from_features(x1, x2)) est.append(float(np.mean(v))) inv_n = np.array([1.0 / n for n in sizes]) beta, a0 = np.polyfit(inv_n, np.array(est), 1) if verbose: for n, e in zip(sizes, est): print(f" P==Q n={n:>6} FID = {e:>9.4f}") print(f" 外推 F_inf = {a0:>9.4f} (真值 0.0000, 斜率 beta={beta:.2f})") c1 = {"sizes": sizes, "est": est, "extrap": float(a0), "slope": float(beta)} # C2: P != Q,真值可以直接从矩算出来 print() shift = np.zeros(d) shift[0] = 0.5 # 只在第 0 维上挪一点 mu_q = shift cov_q = cov * 1.03 # 协方差整体放大 3% true_fid = frechet_distance(np.zeros(d), cov, mu_q, cov_q) est2 = [] for n in sizes: v = [] for r in range(3): x1 = sample_gaussian(np.zeros(d), cov, n, rng) x2 = sample_gaussian(mu_q, cov_q, n, rng) v.append(fid_from_features(x1, x2)) est2.append(float(np.mean(v))) beta2, a02 = np.polyfit(inv_n, np.array(est2), 1) if verbose: for n, e in zip(sizes, est2): print(f" P!=Q n={n:>6} FID = {e:>9.4f}") print(f" 真值 FID = {true_fid:>9.4f}") print(f" 外推 F_inf = {a02:>9.4f} (斜率 beta={beta2:.2f})") print(f" 直接用 n={N} 的估计误差 = {est2[0] - true_fid:+.4f}") print(f" 外推后的误差 = {a02 - true_fid:+.4f}") print() return {"c1": c1, "c2": {"sizes": sizes, "est": est2, "true": float(true_fid), "extrap": float(a02), "slope": float(beta2)}} # ────────────────────────────────────────────────────────────── # [D] 前两阶矩一样、分布完全不同 # ────────────────────────────────────────────────────────────── def _auc(scores_pos: np.ndarray, scores_neg: np.ndarray) -> float: """Mann-Whitney U 形式的 AUC:P(score_pos > score_neg)。""" a = np.sort(scores_pos) b = np.sort(scores_neg) # 对每个 b,统计有多少 a 严格大于它 cnt = a.size - np.searchsorted(a, b, side="right") return float(cnt.sum() / (a.size * b.size)) def _quad_features(X: np.ndarray) -> np.ndarray: """二次特征展开:[x_i, x_i x_j (i<=j)]。""" n, d = X.shape iu = np.triu_indices(d) quad = X[:, iu[0]] * X[:, iu[1]] return np.concatenate([X, quad], axis=1) def _gauss_logpdf(X: np.ndarray, mean: np.ndarray, cov: np.ndarray) -> np.ndarray: """对角/一般协方差下的高斯 log 密度。""" d = X.shape[1] Xc = X - mean[None, :] if cov.ndim == 2 and cov.shape[0] == cov.shape[1]: lam, V = np.linalg.eigh(cov) lam = np.clip(lam, 1e-12, None) proj = Xc @ V quad = (proj ** 2 / lam[None, :]).sum(axis=1) logdet = np.log(lam).sum() else: lam = np.asarray(cov).ravel() quad = (Xc ** 2 / lam[None, :]).sum(axis=1) logdet = np.log(lam).sum() return -0.5 * (quad + logdet + d * np.log(2 * np.pi)) def _nn_two_sample(Xa: np.ndarray, Xb: np.ndarray, m: int = 2500) -> float: """1-近邻两样本检验的准确率。 把两组样本混在一起,对每个点找它的最近邻,看这个邻居是不是同组的。 P == Q 时这个比例趋近 0.5(纯随机),分布有差别时会明显大于 0.5。 """ A, B = Xa[:m], Xb[:m] P = np.concatenate([A, B], axis=0) D = ((P[:, None, :] - P[None, :, :]) ** 2).sum(-1) np.fill_diagonal(D, np.inf) idx = np.argmin(D, axis=1) lab = np.concatenate([np.zeros(m), np.ones(m)]) return float((lab[idx] == lab).mean()) def _ridge_auc(Xp: np.ndarray, Xq: np.ndarray, feat, ntr: int, lam: float, rng: np.random.Generator) -> float: """在给定特征映射上训一个 ridge 二分类器,返回测试集 AUC。""" Fp, Fq = feat(Xp), feat(Xq) n = Xp.shape[0] Xtr = np.concatenate([Fp[:ntr], Fq[:ntr]], axis=0) ytr = np.concatenate([np.ones(ntr), -np.ones(ntr)]) Xte = np.concatenate([Fp[ntr:], Fq[ntr:]], axis=0) yte = np.concatenate([np.ones(n - ntr), -np.ones(n - ntr)]) sd = Xtr.std(axis=0) sd[sd < 1e-12] = 1.0 Xtr, Xte = Xtr / sd, Xte / sd Phi = Xtr.T @ Xtr w = np.linalg.solve(Phi + lam * np.eye(Phi.shape[0]), Xtr.T @ ytr) sc = Xte @ w return _auc(sc[yte > 0], sc[yte < 0]) def section_d(verbose=True): print("=" * 74) print("[D] 前两阶矩完全一样、分布完全不同 —— FID 看得见吗") print(" 真实 P = 0.5*N(+m, I) + 0.5*N(-m, I) (两个分离的模式)") print(" 生成 Q = N(0, I + m m^T) (一个把两个模式糊在一起的团)") print(" 两者均值都是 0、协方差都是 I + m m^T => FID 真值严格等于 0") print("=" * 74) d = 32 n = 40000 rng = np.random.default_rng(24680) u = rng.standard_normal(d) u /= np.linalg.norm(u) if verbose: print(f" {'模式间距 2|m|':>12} {'FID(总体)':>12} {'FID(经验)':>10} " f"{'最优AUC':>9} {'1NN':>7} {'二次AUC':>8} {'线性AUC':>8}") rows = [] for a in (1, 2, 3, 4, 6, 8): m = a * u sign = rng.integers(0, 2, size=n) * 2 - 1 Xp = rng.standard_normal((n, d)) + sign[:, None] * m[None, :] cov_q = np.eye(d) + np.outer(m, m) Xq = rng.multivariate_normal(np.zeros(d), cov_q, size=n) # 总体 FID:直接用矩算,理论值 0 cov_p = np.eye(d) + np.outer(m, m) fid_pop = frechet_distance(np.zeros(d), cov_p, np.zeros(d), cov_q) # 经验 FID fid_emp = frechet_distance(Xp.mean(0), covariance(Xp), Xq.mean(0), covariance(Xq)) # 最优判别(log 密度比) inv_q = np.linalg.inv(cov_q) def logp_mix(X): return np.logaddexp(-0.5 * ((X - m) ** 2).sum(1), -0.5 * ((X + m) ** 2).sum(1)) def logq(X): return -0.5 * (X @ inv_q * X).sum(1) auc_bayes = _auc(logp_mix(Xp) - logq(Xp), logp_mix(Xq) - logq(Xq)) # 1-NN 两样本检验 acc_nn = _nn_two_sample(Xp, Xq) # 二次特征 / 线性特征的 ridge 分类器 auc_quad = _ridge_auc(Xp, Xq, _quad_features, 8000, 1.0, rng) auc_lin = _ridge_auc(Xp, Xq, lambda X: X, 8000, 1.0, rng) rows.append((2 * a, float(fid_pop), float(fid_emp), auc_bayes, acc_nn, auc_quad, auc_lin)) if verbose: print(f" {2 * a:>12} {fid_pop:>12.2e} {fid_emp:>10.4f} " f"{auc_bayes:>9.4f} {acc_nn:>7.4f} {auc_quad:>8.4f} {auc_lin:>8.4f}") # 对照组:把 Q 的协方差整体放大 s 倍(破坏矩匹配),FID 与二次判别一起醒过来 print() print(" 对照组(2|m|=8):把 Q 的协方差整体放大 s 倍,破坏矩匹配") print(f" {'s':>6} {'FID':>10} {'二次特征 ridge AUC':>20}") m = 8 * u sign = rng.integers(0, 2, size=n) * 2 - 1 Xp = rng.standard_normal((n, d)) + sign[:, None] * m[None, :] cov_p = np.eye(d) + np.outer(m, m) ctrl = [] for s in (1.0, 1.05, 1.2, 1.5, 2.0): cov_q_s = cov_p * s Xq_s = rng.multivariate_normal(np.zeros(d), cov_q_s, size=n) fid_s = frechet_distance(np.zeros(d), cov_p, np.zeros(d), cov_q_s) auc_s = _ridge_auc(Xp, Xq_s, _quad_features, 8000, 1.0, rng) ctrl.append((s, float(fid_s), auc_s)) print(f" {s:>6} {fid_s:>10.4f} {auc_s:>20.4f}") print(" -> 本例矩匹配时总体 FID 为 0,平方损失 ridge 的 AUC 近 0.5。") print(" 这不代表任意二次分类器都无法区分两组分布。") print() return {"d": d, "n": n, "rows": rows, "control": ctrl} # ────────────────────────────────────────────────────────────── def main(): which = sys.argv[1].upper() if len(sys.argv) > 1 else "ALL" res = {} if which in ("ALL", "A"): res["A"] = section_a() if which in ("ALL", "A2"): res["A2"] = section_a2() if which in ("ALL", "B"): res["B"] = section_b() if which in ("ALL", "C"): res["C"] = section_c() if which in ("ALL", "D"): res["D"] = section_d() if which == "ALL": with open(OUT_JSON, "w", encoding="utf-8") as f: json.dump(res, f, ensure_ascii=False, indent=1) print(f"结果已写入 {OUT_JSON}") if __name__ == "__main__": main() fid_core.py """ fid_core.py —— FID(Fréchet Inception Distance)的最小可用实现。 本机没有 torch,也没有 scipy(scipy.linalg.sqrtm 是官方实现的核心依赖), 所以这里的矩阵平方根全部用 numpy 的对称特征分解手算。 好处是每一步都看得见,也正好能把「手写实现最容易踩的那个坑」暴露出来。 运行: python fid_core.py """ from __future__ import annotations import numpy as np # ────────────────────────────────────────────────────────────── # 1. 对称 PSD 矩阵的平方根 # ────────────────────────────────────────────────────────────── def sqrtm_sym(C: np.ndarray, eps: float = 1e-6) -> np.ndarray: r"""对称半正定矩阵的平方根。 对 C = V diag(w) V^T,有 C^{1/2} = V diag(sqrt(w)) V^T。 eps 是相对谱尺度的负特征值容差,不是给正特征值设置下限。 明显非半正定输入报错;容差内负值截到 0,保留真正的零特征值。 """ C = np.asarray(C, dtype=np.float64) if C.ndim != 2 or C.shape[0] != C.shape[1] or not np.isfinite(C).all(): raise ValueError("C must be a finite square matrix") scale = max(np.linalg.norm(C, ord=np.inf), np.finfo(float).tiny) if not np.allclose(C, C.T, rtol=0.0, atol=eps * scale): raise ValueError("C must be symmetric") w, V = np.linalg.eigh((C + C.T) / 2) if w.min() < -eps * max(np.abs(w).max(), np.finfo(float).tiny): raise ValueError("C must be positive semidefinite") w = np.sqrt(np.clip(w, 0.0, None)) return (V * w) @ V.T def sqrtm_naive(A: np.ndarray, eps: float = 1e-6) -> np.ndarray: """「把 A 直接当对称矩阵开方」。 np.linalg.eigh 只读矩阵的上/下三角并**假设输入对称**, 默认 UPLO="L",以 A 的下三角及其镜像构造对称矩阵, 并不等于 (A + A^T)/2。 很多手写 FID 就是这么写的,而 A = Σ1 Σ2 恰恰不是对称矩阵。 """ w, V = np.linalg.eigh(A) w = np.sqrt(np.clip(w, eps, None)) return (V * w) @ V.T # ────────────────────────────────────────────────────────────── # 2. Tr((Σ1 Σ2)^{1/2}) # ────────────────────────────────────────────────────────────── def trace_sqrt_product(sigma1: np.ndarray, sigma2: np.ndarray, eps: float = 1e-6) -> float: r"""计算 Tr((Σ1 Σ2)^{1/2}),走对称化路线。 Σ1 正定时,Σ1 Σ2 与 Σ1^{1/2} Σ2 Σ1^{1/2} 相似: Σ1^{1/2} (Σ1^{1/2} Σ2 Σ1^{1/2}) Σ1^{-1/2} = Σ1 Σ2 二者特征值相同,而后者是**对称半正定**的,可以安全用 eigh。 主平方根与相似变换可交换,所以迹也相同: Tr((Σ1 Σ2)^{1/2}) = Σ_i sqrt(λ_i) """ s1 = sqrtm_sym(sigma1, eps) M = s1 @ sigma2 @ s1 M = 0.5 * (M + M.T) # 强制对称,压掉浮点不对称 w = np.linalg.eigvalsh(M) return float(np.sqrt(np.clip(w, 0.0, None)).sum()) def trace_sqrt_product_naive(sigma1: np.ndarray, sigma2: np.ndarray, eps: float = 1e-6) -> float: """对照用的错误写法:直接对 Σ1 Σ2 调 eigh。""" return float(np.trace(sqrtm_naive(sigma1 @ sigma2, eps))) def trace_sqrt_product_ref(sigma1: np.ndarray, sigma2: np.ndarray) -> float: """参考实现:用一般矩阵的特征值求解器 eigvals(不假设对称)。""" ev = np.linalg.eigvals(sigma1 @ sigma2) return float(np.sqrt(np.clip(ev.real, 0.0, None)).sum()) # ────────────────────────────────────────────────────────────── # 3. FID 本体 # ────────────────────────────────────────────────────────────── def frechet_distance(mu1: np.ndarray, sigma1: np.ndarray, mu2: np.ndarray, sigma2: np.ndarray, eps: float = 1e-6, mode: str = "sym") -> float: r"""两个多元高斯之间的 Fréchet 距离(= 2-Wasserstein 距离的平方)。 FID = ||μ1 - μ2||^2 + Tr(Σ1) + Tr(Σ2) - 2 Tr((Σ1 Σ2)^{1/2}) mode="sym" 用对称化路线(正确) mode="naive" 用 eigh(Σ1 Σ2)(错误,保留它只为对照) """ diff = mu1 - mu2 if mode == "sym": tr = trace_sqrt_product(sigma1, sigma2, eps) elif mode == "naive": tr = trace_sqrt_product_naive(sigma1, sigma2, eps) else: raise ValueError(f"unknown mode: {mode}") # 注意:这里**不做** max(val, 0)。浮点误差确实会让 FID(P,P) 变成 -1e-9 量级, # 但把负数夹成 0 会把真正的实现 bug(比如开方写错)一起藏掉—— # 我自己第一版就是因为夹了 0,测试全绿而结果是错的。 return float(diff @ diff + np.trace(sigma1) + np.trace(sigma2) - 2.0 * tr) def covariance(X: np.ndarray, unbiased: bool = True) -> np.ndarray: r"""样本协方差。unbiased=True 用 1/(n-1)(np.cov 默认), False 用 1/n(高斯最大似然估计的常见约定)。""" X = np.asarray(X, dtype=np.float64) if X.ndim != 2 or not np.isfinite(X).all(): raise ValueError("X must be a finite [n, d] array") n = X.shape[0] if n < (2 if unbiased else 1): raise ValueError("not enough samples for covariance") Xc = X - X.mean(axis=0, keepdims=True) denom = (n - 1) if unbiased else n return (Xc.T @ Xc) / denom def fid_from_features(X1: np.ndarray, X2: np.ndarray, unbiased: bool = True, eps: float = 1e-6, mode: str = "sym") -> float: """直接从两组特征算 FID。X1: 真实 [n1, d],X2: 生成 [n2, d]。""" mu1, mu2 = X1.mean(axis=0), X2.mean(axis=0) sig1 = covariance(X1, unbiased=unbiased) sig2 = covariance(X2, unbiased=unbiased) return frechet_distance(mu1, sig1, mu2, sig2, eps=eps, mode=mode) # ────────────────────────────────────────────────────────────── # 4. 自检 # ────────────────────────────────────────────────────────────── def _rand_psd(d: int, rng: np.random.Generator, k: int | None = None) -> np.ndarray: """随机对称正定矩阵:A A^T/k + 0.5 I;加单位阵后满秩。""" k = k or d A = rng.standard_normal((d, k)) return A @ A.T / k + 0.5 * np.eye(d) def self_test() -> None: rng = np.random.default_rng(20261001) print("=" * 68) print("[1] 恒等性:FID(P, P) 必须为 0") d = 32 mu = rng.standard_normal(d) sig = _rand_psd(d, rng) identity = frechet_distance(mu, sig, mu, sig) assert abs(identity) < 1e-10 * np.trace(sig) print(f" FID(P,P) = {identity:.3e} (按容差判断,不要求固定符号或尾数)") for scale in (1.0, 1e-12): singular = np.diag([scale, 0.0, 2 * scale]) z = np.zeros(3) got = frechet_distance(z, singular, z, singular) assert abs(got) < 1e-10 * np.trace(singular) try: sqrtm_sym(np.diag([1.0, -0.1])) except ValueError: pass else: raise AssertionError("non-PSD input was accepted") print(" 奇异/小尺度 PSD 恒等性、非 PSD 拒绝:通过") print() print("[2] 对称化路线 vs 错误写法 vs 一般特征值参考实现") print(f" {'d':>6} {'sym(正确)':>16} {'naive(错误)':>16} {'eigvals(参考)':>16}") for d in (8, 32, 128): s1 = _rand_psd(d, rng) s2 = _rand_psd(d, rng) a = trace_sqrt_product(s1, s2) b = trace_sqrt_product_naive(s1, s2) c = trace_sqrt_product_ref(s1, s2) assert np.isclose(a, c, rtol=1e-9) print(f" {d:>6} {a:>16.10f} {b:>16.10f} {c:>16.10f}") print() print("[3] 这个差异会传进 FID:同一对特征,两种写法差多少") d = 128 n = 4096 mu_a = rng.standard_normal(d) * 0.3 sa = _rand_psd(d, rng) sb = sa + 0.05 * np.eye(d) xa = rng.multivariate_normal(mu_a, sa, size=n) xb = rng.multivariate_normal(-mu_a, sb, size=n) fa = fid_from_features(xa, xb, mode="sym") fb = fid_from_features(xa, xb, mode="naive") print(f" FID(sym) = {fa:.6f}") print(f" FID(naive) = {fb:.6f} (差值 {fb - fa:+.6f})") print() print("[4] 尺度不是不变的:特征整体乘 c,FID 变 c^2 倍") base = fid_from_features(xa, xb, mode="sym") for c in (0.5, 2.0, 4.0): got = fid_from_features(xa * c, xb * c, mode="sym") assert np.isclose(got, c * c * base, rtol=1e-9) print(f" c={c:<4} FID={got:>12.6f} 期望 c^2*base={c * c * base:>12.6f}") print() print("[5] 有偏 vs 无偏协方差(n 越小差得越多)") d = 128 truth_a = _rand_psd(d, rng) truth_b = truth_a + 0.08 * np.eye(d) print(f" {'n':>7} {'unbiased':>13} {'biased':>13} {'差值':>12}") for n in (256, 1024, 8192): pa = rng.multivariate_normal(np.zeros(d), truth_a, size=n) pb = rng.multivariate_normal(np.zeros(d), truth_b, size=n) fu = fid_from_features(pa, pb, unbiased=True) fb2 = fid_from_features(pa, pb, unbiased=False) print(f" {n:>7} {fu:>13.6f} {fb2:>13.6f} {fb2 - fu:>+12.6f}") print() print("=" * 68) if __name__ == "__main__": self_test() clip_alignment_lab.py """ clip_alignment_lab.py —— CLIP Score 与 FID 到底在给谁打高分。 本机没有 torch,跑不了真的 CLIP。这里搭的是一个**结构替身**: - 两个编码器把图像和文本投到同一个共享空间(真 CLIP 就是这么干的) - 打分用余弦相似度,并且照抄 CLIPScore 论文的定义 CLIP-S = w * max(cos(image, text), 0),w = 2.5(arXiv:2104.08718) - 相似度用归一化向量,FID 用未归一化的原始特征(实际评测也是这么用的) 替身复现不了真 CLIP 的具体数值,但复现了它的**结构**。 下面两个结论都只依赖结构,不依赖具体权重: [A] CLIP Score 随生成多样性单调下降,FID 是 U 形 —— 两者的最优解不在一个地方 [B] CLIP Score 只用一个均值,好坏样本可以互相平均掉 运行: python clip_alignment_lab.py """ from __future__ import annotations import json import os import sys import numpy as np from fid_core import fid_from_features HERE = os.path.dirname(os.path.abspath(__file__)) OUT_JSON = os.path.join(HERE, "_clip_results.json") W = 2.5 # CLIPScore 论文的缩放系数 # ────────────────────────────────────────────────────────────── # 共享空间与两个(替身)编码器 # ────────────────────────────────────────────────────────────── def build_concepts(K: int, ds: int, rng: np.random.Generator) -> np.ndarray: """K 个语义概念的类心,单位范数。ds >> K 时它们近似两两正交。""" C = rng.standard_normal((K, ds)) C /= np.linalg.norm(C, axis=1, keepdims=True) return C def encode_image(C: np.ndarray, idx: np.ndarray, sigma: float, rng: np.random.Generator) -> np.ndarray: """「图像编码器」:类心 + 各向同性噪声。sigma 就是生成多样性。""" Z = rng.standard_normal((idx.shape[0], C.shape[1])) return C[idx] + sigma * Z def encode_text(C: np.ndarray, idx: np.ndarray) -> np.ndarray: """「文本编码器」:prompt 直接就是类心本身。""" return C[idx] def clip_score(images: np.ndarray, texts: np.ndarray, w: float = W) -> dict: """CLIP-S = w * max(cos(image, text), 0),逐样本取均值。""" a = images / np.linalg.norm(images, axis=1, keepdims=True) b = texts / np.linalg.norm(texts, axis=1, keepdims=True) cos = (a * b).sum(axis=1) return { "score": float(w * np.maximum(cos, 0.0).mean()), "cos_mean": float(cos.mean()), "cos_std": float(cos.std()), "cos_frac_gt_half": float((cos > 0.5).mean()), "cos": cos, } # ────────────────────────────────────────────────────────────── # [A] 多样性扫描:FID 与 CLIP Score 的最优解不在一起 # ────────────────────────────────────────────────────────────── def section_a(K=24, ds=64, n=20000, verbose=True): print("=" * 76) print("[A] 生成多样性 sigma_g 扫描:FID 与 CLIP Score 分别给谁打高分") print(f" 真实数据 sigma_real = 0.5,K={K} 个概念,共享空间维度 ds={ds}") print("=" * 76) rng = np.random.default_rng(90210) C = build_concepts(K, ds, rng) idx = rng.integers(0, K, size=n) real_img = encode_image(C, idx, 0.5, rng) real_txt = encode_text(C, idx) ref = clip_score(real_img, real_txt) if verbose: print(f" 真实数据自己的 CLIP Score = {ref['score']:.4f} " f"(cos 均值 {ref['cos_mean']:.4f})") print() print(f" {'sigma_g':>8} {'FID':>10} {'CLIPScore':>11} " f"{'cos均值':>9} {'cos标准差':>10}") rows = [] for sg in (0.05, 0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.8, 1.0): idx2 = rng.integers(0, K, size=n) gen_img = encode_image(C, idx2, sg, rng) gen_txt = encode_text(C, idx2) fid = fid_from_features(real_img, gen_img) cs = clip_score(gen_img, gen_txt) rows.append((sg, float(fid), cs["score"], cs["cos_mean"], cs["cos_std"])) if verbose: mark = " <- 真实值" if abs(sg - 0.5) < 1e-9 else "" print(f" {sg:>8} {fid:>10.4f} {cs['score']:>11.4f} " f"{cs['cos_mean']:>9.4f} {cs['cos_std']:>10.4f}{mark}") fids = [r[1] for r in rows] scores = [r[2] for r in rows] best_fid = rows[int(np.argmin(fids))][0] best_cs = rows[int(np.argmax(scores))][0] if verbose: print() print(f" FID 最小时 sigma_g = {best_fid}") print(f" CLIP Score 最大时 sigma_g = {best_cs}") print(" -> 在本合成实验中,FID 偏好的方差与 CLIPScore 不同;") print(" 不能据此把真实模型中的多样性与文本对齐视为必然冲突。") print() return {"K": K, "ds": ds, "n": n, "rows": rows, "real_score": ref["score"], "best_fid_sigma": best_fid, "best_clip_sigma": best_cs} # ────────────────────────────────────────────────────────────── # [B] 同一个均值,完全不同的现实 # ────────────────────────────────────────────────────────────── def _score_for_sigma(C, idx, sigma) -> float: """用固定种子探一次,避免二分过程本身消耗主随机流。""" probe = np.random.default_rng(4242) img = encode_image(C, idx, sigma, probe) return clip_score(img, encode_text(C, idx))["score"] def section_b(K=24, ds=64, n=20000, verbose=True): print("=" * 76) print("[B] 汇总 CLIP Score 是均值:好样本和坏样本可以互相平均掉") print(" 模型 M1:p 的概率输出完美匹配,1-p 的概率输出纯噪声") print(" 模型 M2:各样本相似度较集中(调 sigma 让截断后的 CLIPScore 与 M1 相同)") print(" 两者的 CLIP Score 一样,现实完全不一样。") print("=" * 76) rng = np.random.default_rng(1357) C = build_concepts(K, ds, rng) idx = rng.integers(0, K, size=n) real_img = encode_image(C, idx, 0.5, rng) if verbose: print(f" {'p':>6} {'M1 CLIP':>9} {'M1 cos标准差':>13} {'M1 好图占比':>12} " f"{'M1 FID':>10} | {'M2 sigma':>9} {'M2 CLIP':>9} {'M2 cos标准差':>13} " f"{'M2 好图占比':>12} {'M2 FID':>10}") rows = [] for p in (0.3, 0.5, 0.7, 0.9): # M1: 混合 good = rng.random(n) < p img1 = np.where(good[:, None], C[idx], 0.0) # 完美命中类心 noise = rng.standard_normal((n, ds)) noise /= np.linalg.norm(noise, axis=1, keepdims=True) img1 = img1 + np.where(good[:, None], 0.0, noise) # 否则是随机方向 cs1 = clip_score(img1, encode_text(C, idx)) fid1 = fid_from_features(real_img, img1) # M2: 二分法找 sigma,使 CLIP Score 与 M1 对齐 # 注意要对齐的是 score(含 max(cos, 0) 截断),不是裸的 cos 均值 target = cs1["score"] lo, hi = 1e-3, 50.0 for _ in range(60): mid = 0.5 * (lo + hi) if _score_for_sigma(C, idx, mid) > target: lo = mid else: hi = mid sg2 = 0.5 * (lo + hi) img2 = encode_image(C, idx, sg2, rng) cs2 = clip_score(img2, encode_text(C, idx)) fid2 = fid_from_features(real_img, img2) rows.append({"p": p, "m1_clip": cs1["score"], "m1_std": cs1["cos_std"], "m1_good": cs1["cos_frac_gt_half"], "m1_fid": float(fid1), "m2_sigma": float(sg2), "m2_clip": cs2["score"], "m2_std": cs2["cos_std"], "m2_good": cs2["cos_frac_gt_half"], "m2_fid": float(fid2)}) bins = np.linspace(-1.0, 1.0, 51) rows[-1].update(hist_bins=bins.tolist(), hist1=np.histogram(np.clip(cs1["cos"], -1, 1), bins)[0].tolist(), hist2=np.histogram(np.clip(cs2["cos"], -1, 1), bins)[0].tolist(), m1_cos_mean=cs1["cos_mean"], m2_cos_mean=cs2["cos_mean"]) if verbose: print(f" {p:>6} {cs1['score']:>9.4f} {cs1['cos_std']:>13.4f} " f"{cs1['cos_frac_gt_half']:>12.4f} {fid1:>10.3f} | " f"{sg2:>9.3f} {cs2['score']:>9.4f} {cs2['cos_std']:>13.4f} " f"{cs2['cos_frac_gt_half']:>12.4f} {fid2:>10.3f}") if verbose: print() d_clip = max(abs(r["m1_clip"] - r["m2_clip"]) for r in rows) d_std = max(r["m1_std"] / max(r["m2_std"], 1e-9) for r in rows) print(f" CLIP Score 最大差距 = {d_clip:.4f} (按构造两者应当同分)") print(f" 逐样本相似度标准差的比值最大 = {d_std:.2f} 倍") print(" -> 同一个 CLIP Score 背后,可以是「p 的图完美、其余完全不沾边」,") print(" 也可以是「各样本有相近的相似度」。汇总均值看不见这个区别,逐样本分数分布可以。") print() return {"K": K, "ds": ds, "n": n, "rows": rows} def main(): which = sys.argv[1].upper() if len(sys.argv) > 1 else "ALL" res = {} if which in ("ALL", "A"): res["A"] = section_a() if which in ("ALL", "B"): res["B"] = section_b() if which == "ALL": with open(OUT_JSON, "w", encoding="utf-8") as f: json.dump(res, f, ensure_ascii=False, indent=1) print(f"结果已写入 {OUT_JSON}") if __name__ == "__main__": main() make_figures.py """ make_figures.py —— 画本文的配图。 数据来源都是已经跑完的实验(_fid_bias_results.json / _clip_results.json), 不在这里重新算,避免图上的数字和正文漂移。 运行: python make_figures.py """ from __future__ import annotations import json import os import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np HERE = os.path.dirname(os.path.abspath(__file__)) FIGDIR = os.path.join(HERE, "..", "figures") os.makedirs(FIGDIR, exist_ok=True) # 配色(正文里写「这张图要看什么」时按这六个名字来描述) C_MAIN = "#1f4e79" # 深蓝:主曲线 C_ALT = "#c1440e" # 橙红:对照曲线 C_GREEN = "#2e7d32" # 绿:第三组 C_PURPLE = "#7b1fa2" # 紫:标注线 C_GRAY = "#8a8a8a" # 灰:参考线 C_LIGHT = "#bcd7ee" # 浅蓝:填充 plt.rcParams["font.sans-serif"] = ["PingFang SC", "Arial Unicode MS", "DejaVu Sans"] plt.rcParams["axes.unicode_minus"] = False plt.rcParams["figure.dpi"] = 130 plt.rcParams["savefig.dpi"] = 130 def _load(name): p = os.path.join(HERE, name) if not os.path.exists(p): return None with open(p, encoding="utf-8") as f: return json.load(f) # ────────────────────────────────────────────────────────────── # 图 1:FID 的样本量偏差 # ────────────────────────────────────────────────────────────── def fig_bias(res): rows512 = res["A"]["by_d"]["512"] rows2048 = res["A"]["by_d"]["2048"] rows_b = res["B"]["rows"] fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12.4, 4.7)) # ── 左:偏差 vs n ── n5 = [r[0] for r in rows512] f5 = [r[1] for r in rows512] n2 = [r[0] for r in rows2048] f2 = [r[1] for r in rows2048] ax1.loglog(n5, f5, "o-", color=C_MAIN, lw=2, ms=5, label=r"$d=512$") ax1.loglog(n2, f2, "s-", color=C_ALT, lw=2, ms=5, label=r"$d=2048$") # 参考斜率 1/n ref_n = np.array([n2[3], n2[-1]], dtype=float) ref_y = f2[-1] * (ref_n / n2[-1]) ** (-1.0) ax1.loglog(ref_n, ref_y, "--", color=C_GRAY, lw=1.6, label=r"$\mathrm{slope}=-1$") ax1.axvline(2048, color=C_PURPLE, ls=":", lw=1.6) ax1.annotate(r"$n=d=2048$", xy=(2048, 1.0), xytext=(2600, 1.6), color=C_PURPLE, fontsize=10) ax1.annotate(r"$n=50000$ 时仍有 $14.21$", xy=(n2[-1], f2[-1]), xytext=(9000, 30), color=C_ALT, fontsize=10.5, arrowprops=dict(arrowstyle="->", color=C_ALT, lw=1.2)) ax1.set_xlabel("样本量 $n$(两组各 $n$ 张)", fontsize=11) ax1.set_ylabel(r"$\mathrm{FID}$(真值 $0$)", fontsize=11) ax1.set_title("偏差随样本量衰减:$1/n$", fontsize=12, pad=8) ax1.legend(loc="upper right", fontsize=10, framealpha=0.95) ax1.grid(True, which="both", alpha=0.25) # ── 右:偏差 vs d ── dd = [r[0] for r in rows_b] bb = [r[1] for r in rows_b] ax2.loglog(dd, bb, "o-", color=C_MAIN, lw=2, ms=6) lo, hi = np.log(dd[0]), np.log(dd[-1]) slope = (np.log(bb[-1]) - np.log(bb[0])) / (hi - lo) ax2.annotate(r"$\mathrm{slope}\approx %.2f$" % slope, xy=(dd[2], bb[2]), xytext=(90, 20), fontsize=11.5, color=C_PURPLE, arrowprops=dict(arrowstyle="->", color=C_PURPLE, lw=1.2)) ax2.scatter([2048], [bb[-1]], s=90, facecolors="none", edgecolors=C_ALT, lw=2, zorder=5) ax2.annotate(r"$d=2048$ 时 $71.14$", xy=(2048, bb[-1]), xytext=(600, 40), color=C_ALT, fontsize=10.5, arrowprops=dict(arrowstyle="->", color=C_ALT, lw=1.2)) ax2.set_xlabel("特征维度 $d$", fontsize=11) ax2.set_ylabel(r"$\mathrm{FID}$(真值 $0$)", fontsize=11) ax2.set_title(r"固定 $n=10000$,偏差随维度暴涨", fontsize=12, pad=8) ax2.grid(True, which="both", alpha=0.25) fig.suptitle("图 1:人工高斯同分布,有限样本估计仍有偏差", fontsize=13.5, y=1.0) fig.tight_layout() out = os.path.join(FIGDIR, "fig_fid_bias.png") fig.savefig(out, bbox_inches="tight", facecolor="white") plt.close(fig) print("wrote", out, f"slope={slope:.3f}") # ────────────────────────────────────────────────────────────── # 图 2:前两阶矩一样、分布完全不同 # ────────────────────────────────────────────────────────────── def fig_moment_blind(res): rows = res["D"]["rows"] sep = [r[0] for r in rows] fpop = [abs(r[1]) for r in rows] femp = [r[2] for r in rows] aucb = [r[3] for r in rows] nn = [r[4] for r in rows] fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12.4, 4.7)) # ── 左:二维示意 ── rng = np.random.default_rng(20261001) a = 4.2 # 二维示意里把模式拉开一点,让「两个团」一眼可见 n = 1400 sgn = rng.integers(0, 2, size=n) * 2 - 1 P = rng.standard_normal((n, 2)) + np.stack([sgn * a, np.zeros(n)], axis=1) covq = np.eye(2) + np.array([[a * a, 0.0], [0.0, 0.0]]) Q = rng.multivariate_normal(np.zeros(2), covq, size=n) ax1.scatter(P[:, 0], P[:, 1], s=9, alpha=0.5, color=C_MAIN, label="真实分布 $P$(两个模式)") ax1.scatter(Q[:, 0], Q[:, 1], s=9, alpha=0.5, color=C_ALT, label="生成分布 $Q$(糊成一团)") # 画 Q 的 1 个标准差椭圆 w, V = np.linalg.eigh(covq) ang = np.degrees(np.arctan2(V[1, -1], V[0, -1])) from matplotlib.patches import Ellipse for k, col in ((1, C_ALT), (2, C_ALT)): e = Ellipse((0, 0), 2 * k * np.sqrt(w[0]), 2 * k * np.sqrt(w[1]), angle=ang, fill=False, ls="--", lw=1.4, edgecolor=col, alpha=0.75) ax1.add_patch(e) ax1.set_xlim(-9, 9) ax1.set_ylim(-4.2, 4.2) ax1.set_aspect("equal", adjustable="box") ax1.set_xlabel(r"$x_1$", fontsize=11) ax1.set_ylabel(r"$x_2$", fontsize=11) ax1.set_title(r"$\mathrm{FID}=0$,但一眼就能看出不是一回事", fontsize=12, pad=8) ax1.legend(loc="upper left", fontsize=10, framealpha=0.95) ax1.grid(True, alpha=0.25) # ── 右:FID 与可区分度 ── ax2.plot(sep, femp, "o-", color=C_MAIN, lw=2.2, ms=6, label=r"$\mathrm{FID}$(左边刻度)") ax2.set_yscale("log") ax2.set_ylim(1e-3, 1e1) ax2.axhline(0.5, color=C_GRAY, ls=":", lw=1.2) ax2.set_xlabel("模式间距 $2|m|$", fontsize=11) ax2.set_ylabel(r"$\mathrm{FID}$(对数刻度,真值严格为 $0$)", fontsize=11, color=C_MAIN) ax2.tick_params(axis="y", labelcolor=C_MAIN) ax3 = ax2.twinx() ax3.plot(sep, aucb, "s--", color=C_ALT, lw=2.2, ms=6, label=r"$\mathrm{AUC}$(最优判别)") ax3.plot(sep, nn, "^--", color=C_GREEN, lw=2.2, ms=6, label=r"$1$-$\mathrm{NN}$ 两样本准确率") ax3.axhline(0.5, color=C_GRAY, ls="-", lw=1.0) ax3.set_ylim(0.45, 1.0) ax3.set_ylabel(r"$\mathrm{AUC}$ / $1$-$\mathrm{NN}$(右边刻度)", fontsize=11) ax3.text(sep[-1], 0.47, r"$0.5=$ 随机猜测", color=C_GRAY, fontsize=9.5, ha="right") h1, l1 = ax2.get_legend_handles_labels() h2, l2 = ax3.get_legend_handles_labels() ax2.legend(h1 + h2, l1 + l2, loc="center left", fontsize=9.5, framealpha=0.95) ax2.set_title("模式越分开,FID 越是一动不动,判别器越看得清", fontsize=12, pad=8) fig.suptitle("图 2:FID 只匹配前两阶矩,形状对不对它不管", fontsize=13.5, y=1.0) fig.tight_layout() out = os.path.join(FIGDIR, "fig_moment_blind.png") fig.savefig(out, bbox_inches="tight", facecolor="white") plt.close(fig) print("wrote", out) # ────────────────────────────────────────────────────────────── # 图 3:FID 与 CLIP Score 的最优解不在一起 # ────────────────────────────────────────────────────────────── def fig_clip_vs_fid(cres): rows = cres["A"]["rows"] sg = [r[0] for r in rows] fid = [r[1] for r in rows] cs = [r[2] for r in rows] fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12.4, 4.7)) # ── 左:sigma 扫描 ── ax1.plot(sg, fid, "o-", color=C_MAIN, lw=2.4, ms=6) ax1.set_xlabel(r"生成多样性 $\sigma_g$", fontsize=11) ax1.set_ylabel("FID(越低越好)", fontsize=11, color=C_MAIN) ax1.tick_params(axis="y", labelcolor=C_MAIN) ax1.set_ylim(-1.0, 19.5) imin = int(np.argmin(fid)) ax1.scatter([sg[imin]], [fid[imin]], s=170, facecolors="none", edgecolors=C_MAIN, lw=2.2, zorder=5) ax1.annotate(r"FID 最小,$\sigma_g=%.2f$" % sg[imin], xy=(sg[imin], fid[imin]), xytext=(0.58, 9.6), color=C_MAIN, fontsize=10.5, arrowprops=dict(arrowstyle="->", color=C_MAIN, lw=1.2)) ax3 = ax1.twinx() ax3.plot(sg, cs, "s--", color=C_ALT, lw=2.4, ms=6) ax3.set_ylabel("CLIP Score(越高越好)", fontsize=11, color=C_ALT) ax3.tick_params(axis="y", labelcolor=C_ALT) ax3.set_ylim(0.15, 2.85) imax = int(np.argmax(cs)) ax3.scatter([sg[imax]], [cs[imax]], s=170, facecolors="none", edgecolors=C_ALT, lw=2.2, zorder=5) ax3.annotate(r"CLIP Score 最大,$\sigma_g=%.2f$" % sg[imax], xy=(sg[imax], cs[imax]), xytext=(0.21, 1.85), color=C_ALT, fontsize=10.5, arrowprops=dict(arrowstyle="->", color=C_ALT, lw=1.2)) ax1.axvline(0.5, color=C_GRAY, ls=":", lw=1.4) ax1.text(0.31, 2.4, r"$\sigma_{\mathrm{real}}=0.5$", color=C_GRAY, fontsize=10) ax1.set_title("合成共享空间:两个目标的最优点不同", fontsize=12, pad=8) ax1.grid(True, alpha=0.22) # ── 右:同一个 CLIP Score 的两种现实 ── brows = cres["B"]["rows"] target = [r for r in brows if abs(r["p"] - 0.5) < 1e-9][0] # 直接读取实验的真实 cos 直方图;不重新捏造近似分布。 bins = np.asarray(target["hist_bins"]) ax2.stairs(target["hist1"], bins, fill=True, alpha=0.72, color=C_ALT, label=r"$M_1$:半数精确对齐,半数随机方向") ax2.stairs(target["hist2"], bins, fill=True, alpha=0.72, color=C_MAIN, label=r"$M_2$:相似度较集中") for key, color in [("m1_cos_mean", C_ALT), ("m2_cos_mean", C_MAIN)]: ax2.axvline(target[key], color=color, ls="--", lw=1.3) ax2.text(0.03, 0.74, "虚线为各自原始 cos 均值\n分数含截断,等分不等于 cos 均值相同", transform=ax2.transAxes, color="#444444", fontsize=8.5) ax2.set_xlabel(r"单张样本的相似度 $\cos(f_{\mathrm{img}}, f_{\mathrm{txt}})$", fontsize=11) ax2.set_ylabel(r"$\mathrm{count}$", fontsize=11) ax2.set_title(r"$M_1$ 标准差 $%.2f$,$M_2$ 标准差 $%.2f$,近似同分" % (target["m1_std"], target["m2_std"]), fontsize=12, pad=8) ax2.legend(loc="upper center", fontsize=9.5, framealpha=0.95) ax2.grid(True, alpha=0.22) fig.suptitle("图 3:合成替身实验,不是真实 CLIP 或 Inception 测评", fontsize=13.5, y=1.0) fig.tight_layout() out = os.path.join(FIGDIR, "fig_clip_vs_fid.png") fig.savefig(out, bbox_inches="tight", facecolor="white") plt.close(fig) print("wrote", out) # ────────────────────────────────────────────────────────────── def main(): res = _load("_fid_bias_results.json") cres = _load("_clip_results.json") made = [] if res: fig_bias(res) fig_moment_blind(res) made += ["fig_fid_bias.png", "fig_moment_blind.png"] if cres: fig_clip_vs_fid(cres) made.append("fig_clip_vs_fid.png") print("figures:", made) if __name__ == "__main__": main()
2026年10月02日
2 阅读
0 评论
0 点赞
1
2
...
17
粤ICP备2021042327号