首页
应用
关于
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
累计撰写
205
篇文章
累计收到
8
条评论
首页
应用
栏目
AIGC
AIGC Daily Papers
AIGC Fundamentals
其他
职场经验复盘
多模态理解
购房/投资
阅读
广告
Segmentation
LeetCode
Pytorch
Python
Shell
C++
常用链接
页面
关于
搜索到
205
篇与
人工智能炼丹君
的结果
2026-09-28
AIGC 每日速读|2026-09-28|阿里 397B 只改提示词就涨 50 分,WanPE
今日 AIGC 论文速览 今日共 10 篇 · 视频生成与编辑 5 篇 · 音视频与语音生成 2 篇 · 图像与 3D 生成 2 篇 · 多模态理解与评测 1 篇 重点论文标题列表 WanPE(南大):397B 分镜规划提示词 AV-GRPO(上海人工智能实验室):音视频联合生成的模态解耦 RL ⚡ TRACK(NVIDIA):免训练大小模型按步换班 ViRDM(美国东北大学):砍掉教师与批评家做因果后训练 EditVoice(厦门大学):变长非自回归零样本语音编辑 今日论文速览 1. WanPE:397B 分镜规划提示词 WanPE: Towards Cinematic Prompt Enhancement for Modern Text-to-Video Generation | 南京大学;Wan 团队(阿里巴巴集团);中国科学技术大学;复旦大学;清华大学 | arXiv:2609.30221 关键词:文生视频, 提示词增强, GRPO, 电影分镜, 评测基准 ⚠️ 前序问题:现代视频生成器已能生成 30 秒、并忠实遵循复杂条件,文本提示几乎直接决定多镜头序列里动作、机位轨迹、光效和声音如何展开。但现有的提示增强大多还停留在「可选扩写」层面,面对导演级的镜头级规划(分镜如何组织、跨镜头语义如何不漂)能力明显不足。 本文贡献:提出 WanPE——397B 参数的提示增强模型,在 105 万条真实视频上训练。核心有两块:一是 video-grounded 逆向构建(reverse SFT),从真实视频反推出镜头级电影分镜规划,让模型学到的是「怎么组织分镜」而不是「怎么把句子写长」;二是 Semantic-Consistency GRPO(SC-GRPO),用语义一致性奖励把用户意图在跨镜头、跨时间维度上保住。配套还建了 WanPEval:249 条人工标注请求,覆盖 5–30 秒时长与不同意图粒度,约 1.1 万次盲评 pairwise 比较。 WanPE training pipeline. Video-grounded reverse SFT and semantic-consistency GRPO. 实验效果:驱动 Wan3.0 生成器时,相对原始用户提示,5–15 秒人类偏好提升 10.66–18.84 分,30 秒竞技场提升 50.86 分。5–15 秒子集上专家偏好 S / Bradley–Terry 分数为 50.61 / 61.85,超过 Seedance 2.0(43.42 / 54.70)、MiniMax-H3(36.40 / 49.27)、Kling 3.0(20.80 / 37.40)与 HappyHorse 1.1(24.89 / 40.46)。消融显示逆向构建明显优于正向改写,SC-GRPO 在各模型规模上都能稳住语义保真。 Fine-grained expert evaluation on the 5-15 seconds subset of WanPEval. Left: Preference scores across generation durations and request-granularity levels. Right: Preference scores S across seven content categories. 批判点评:397B 这个体量本身就是最大的成本项,而表里的规模曲线并不支持它:4B 是 43.60 / 56.21,9B 是 46.00 / 58.01,397B 才 50.61 / 61.85——为了最后 5 分,参数量涨了两个数量级。更值得注意的是 30 秒子集:+WanPE-397B 的 Overall 60.24 只是险胜 Seedance 2.5 的 59.76,而且拆开看,动作类 50.00 被 Seedance 2.5 的 68.18 甩开一大截,知识类、演讲类也不占优。换句话说,WanPE 真正补齐的是「镜头语言」,动作生成本身的短板它一个字都没碰。 2. AV-GRPO:音视频联合生成的模态解耦 RL AV-GRPO: Modality-Anchored Decoupling Diffusion Reinforcement Learning for Joint Audio-Video Generation | 上海人工智能实验室;香港理工大学;新加坡国立大学;南洋理工大学;浙江大学 | arXiv:2609.29816 关键词:音视频联合生成, 强化学习后训练, GRPO, 跨模态同步, 扩散模型 ⚠️ 前序问题:音视频联合生成里,异质多模态奖励互相纠缠导致信用分配困难,两个模态塔动力学差异大、联合优化极贵,而同步性评测又依赖成对样本、奖励没法公平比较——直接把 RL 后训练搬过来几乎处处碰壁。 本文贡献:提出 AV-GRPO 这一模态锚定的在线扩散 RL 框架,三个模块:模态锚定 rollout 解耦学习信号并稳定难度;轨迹锁定的冻结塔优化降低成本、重新分配信用;针对模态各自动力学的自适应目标与扰动强度。由此把耦合的多模态偏好学习拆成单模态子问题,实现精确奖励归因。同时构建了 5DAV 数据集,在五个维度上解耦样本、难度可控。基于 LTX-2.3 22B,8 卡 A800 上跑 LoRA 与全参两档。 Overview of the AV-GRPO training pipeline. 实验效果:JavisBench 上 AV-IB 从 0.224 提到 0.247,JavisScore 0.202 → 0.222,DeSync 0.708 → 0.607;VABench 上 Lip Sync 1.439 → 1.646、DeSync 0.726 → 0.542。在生成质量、语义对齐与跨模态同步三项上均超过 LTX-2.3,LoRA 与全参微调两档都成立。 Left: Radar chart comparing performance on the 5DAV and VGGSound datasets. Middle: Metrics are reported as relative percentage changes with respect to the baseline. Right: Comparison of model performance with different alternating step intervals. 批判点评:表里藏着一个不太好看的反差:LTX-2.3+AV-GRPO(lora) 在 Alignment(4.512 对 4.510)、Expressiveness(4.450 对 4.434)、Visual Realism(4.402 对 4.395)三项主观指标上,全都高于全参版。花了全参训练的成本,主观质量却没赢过 LoRA,说明这套「冻结塔 + 分塔信用分配」的收益主要落在同步类客观指标上,主观提升更接近采样噪声。而且这三项差距都在小数点后第三位,论文也没给方差或多次运行,显著性无从判断。 3. TRACK:免训练大小模型按步换班 Accelerating Video Diffusion via Training-Free Trajectory Routing | NVIDIA | arXiv:2609.30096 关键词:视频扩散加速, 免训练, 模型路由, 步级调度, 能耗优化 ⚠️ 前序问题:视频扩散推理昂贵:要把大模型跑很多个去噪步。即便做了步数蒸馏,剩下每一个蒸馏步依然要付一次完整的大模型前向,推理成本还是下不来。 本文贡献:TRACK(TRajectory-Aware Capacity routing via top-K selection)是一种异构去噪策略:离线校准阶段先跑一遍全大模型的参考轨迹,每一步同时收集小模型的预测,与大模型预测比较得到相对分歧分数(两边共享同一 latent、timestep、条件与 guidance),再把该信号跨校准集聚合成分步分歧分数图,据此决定切换策略——质量敏感步留给大模型(如 Wan 2.1 14B),低分歧步路由给小模型(1.3B)。推理时每步只执行被选中的那一个模型,不重训练、不改架构与调度器、也不需要在线双模型评估。 Training-Free Trajectory-Aware Capacity routing. Top: Offline calibration rolls out the all-large-model reference trajectory and, at every step, evaluates the large and small checkpoints using the same latent, timestep, conditioning, and guidance inputs. Bottom: Online inference applies that policy while preserving the shared latent representation and scheduler update. 实验效果:在 Wan 2.1、Cosmos 3、TurboDiffusion、FastVideo 四条管线上分别取得 1.95×、2.04×–2.73×、2.69×、2.17× 加速,聚合质量相当、DreamSim 多样性保留 >95%。在 A100 上用 NVML 实测,去噪能耗降低约 45%–60%:Wan 2.1 每视频省 194.57 kJ,折合每千条视频省 54.05 kWh(降 49.3%);Cosmos 3(Super/Edge)能耗降近 60%。 Speed--quality tradeoffs across video pipelines. Each panel compares the all-large model (baseline) and TRACK (ours) switching policy. All models retain comparable aggregate quality at the selected operating points in most quality dimensions. 批判点评:「质量相当」这个结论要挑一下口径:它取的是 VBench 的时间闪烁、运动平滑、主体一致性、背景一致性四个维度的均值——恰好是「换小模型不太容易崩」的那四个。真正的代价作者写在 Limitations 里:部署时必须同时驻留大小两个模型,显存占用反而高于单模型管线。另外对已经把步数压到 3–4 步的蒸馏管线(FastVideo 只有 2.17×),可换的步数本就只剩个位数,收益上限天然受限。 4. ViRDM:砍掉教师与批评家做因果后训练 ViRDM: Taming Representation Distribution Matching for Few-Step Causal Video Generation | 美国东北大学;Adobe Research | arXiv:2609.28923 关键词:因果视频生成, 分布匹配, 后训练, 显存优化, VBench ⚠️ 前序问题:少步自回归视频扩散能实现低延迟流式生成,但现有后训练方法几乎都依赖 DMD:需要一个庞大的预训练教师和一个在线 critic,靠 diffusion score 估计分布差异。三套网络同时在线,资源开销成了主要门槛。 本文贡献:ViRDM 把教师与 critic 整个去掉,只让生成器对齐一份离线预计算的目标表征分布。三个障碍逐个击破:用一致性式采样 + 随机截断 clean-exit 监督(S 从 U{1..K} 采样,选中的那次评估才吃梯度)把梯度路径变得可行;用轻量 VAE 解码器压显存;用分阶段 vector–Jacobian product 拆解反传。此外还给出了视频 RDM 的生成种群规模与初始化策略(因果 ODE 初始化最优),并加了 flow 动力学正则来补表征分布对时间变化约束不足的问题。 Consistency-style sampling with stochastic exits for few-step video RDM. 实验效果:同样 8 卡 A100 的设置下,峰值显存从 77.1 GB/卡降到 48.3 GB/卡,后训练时间从 22 小时压到 2 小时,VBench Total 从 84.51 提到 84.87。只需 20 次生成器更新、16 A100 GPU-小时,比此前最佳少步因果基线高 0.36 分。用户研究中文本-视频对齐 40.4%、视觉质量 43.0% 均为四步因果方法里最高(Self Forcing 两项均为 32.6%)。 Replacing the resource-intensive teacher-critic stack with direct representation distribution matching for few-step causal video post-training. In the same 8-A100 parallelized training setting, this reduces peak memory from 77.1 to 48.3 GB per GPU and post-training time from 22 to 2 hours, while improving the VBench Total from 84.51 to 84.87. 批判点评:84.87 对 84.51 只涨了 0.36 分,而论文自己的消融表把原因说得很直白:要把 Dynamic Degree 维度单独剔掉再比,ViRDM 才更好看——说明动力学恰恰是 RDM 的软肋(表征分布对时间演化约束不足),得额外加 flow 正则才勉强补回来,等于承认「省掉教师 critic」的代价是引入了新的手工正则。另外 84.87 是在 Wan2.1-1.3B 这个 1.3B 小骨架上拿的,跟动辄 14B 的双向模型不在一个量级,跨量级横向对比要打折扣。 5. EditVoice:变长非自回归零样本语音编辑 EditVoice: Variable-Length Non-Autoregressive Zero-Shot TTS and Speech Editing with Edit Flows | 厦门大学 | arXiv:2609.29889 关键词:零样本 TTS, 语音编辑, 非自回归生成, Edit Flows, 变长序列建模 ⚠️ 前序问题:现有非自回归零样本 TTS 虽然能并行生成,但几乎都要在生成前先指定目标序列长度——长度得靠另一个模块预估,一旦估错就直接影响韵律与内容。 本文贡献:EditVoice 是首个变长 NAR 零样本 TTS:用 Edit Flows 通过插入、删除、替换三类编辑操作同时更新语音内容和序列长度。训练上采用 speech-infilling,把零样本 TTS 与文本驱动语音编辑统一到同一目标下,推理时前缀提示和后缀提示两种摆位都能用;提出 Complementary Prompt Sampling(CPS),按每类操作取两种摆位里速率更大的那个,利用互补的 Edit Flow 预测。还发现模型能编辑训练来源之外的语音(包括它自己生成的),据此做端到端编辑与免训练的后生成精炼。 Overview of the EditVoice zero-shot TTS inference pipeline. Stage one generates semantic speech tokens using CPS, and stage two applies post-generation refinement to the generated sequence before speech detokenization. 实验效果:在 10K 小时 GigaSpeech 上训练 Edit Flow 模型,在 Seed-TTS Eval EN、LibriSpeech-PC 上零样本 TTS 表现具竞争力,在 RealEdit 上语音编辑表现具竞争力。后生成精炼把 WER 从 2.66% 降到 1.55%,且 8 步生成 + 2 步精炼已经超过不做精炼的 16 步生成。 Post-generation refinement on Seed-TTS Eval EN. Starting from the same 8-step CPS output, WER and edits to generated tokens are measured over 8 refinement steps. The dashed line denotes 16-step CPS generation without refinement. 批判点评:全文只有 5 页 3 张图,实验规模偏薄:10K 小时训练的对比对象里有 CosyVoice 2 这类数据量级高一个档的系统,「competitive」具体竞争到什么程度得回表里逐项看。后精炼这个卖点本质是拿推理时算力换质量——8+2 步打败 16 步听起来漂亮,但总步数只少了 6 步,换成真实延迟并不划算。另外「能编辑非训练来源的语音」是意外发现而不是设计,作者没给系统的泛化边界评估,这条能力在多说话人、跨语种上会不会退化是未知数。 6. OREO:用 2D 编辑反馈对齐 3D 保真 OREO: Fidelity Alignment in 3D Generation via On-the-fly Rendering-Editing Optimization | 香港理工大学;腾讯 ARC Lab | arXiv:2609.29788 关键词:3D 生成, 保真对齐, 2D 扩散先验, 强化编辑, ECCV 2026 ⚠️ 前序问题:3D 生成模型进步很快,但产出的资产视觉保真度往往不够——几何、布局对得上,质感与细节就是差一口气,和 2D 扩散模型能达到的写实程度有明显落差。 本文贡献:OREO 是一个对齐框架,核心是不依赖静态数据集,而是建立动态优化闭环:把 3D 输出渲染出的视图交给 2D 模型做 Reinforced Editing,在保持几何、视角和内容不变的前提下提升视觉保真;这些被实时编辑过的渲染图反过来作为高质量 2D 伪目标,监督 3D 生成器从自己的样本里学习、逐步提升画质。 Overview of OREO. 实验效果:以 Trellis 为底座,对参考图 $x^{ref}$ 的 CLIP / DINO 相似度在 Conceptual Design 上从 0.7613 / 0.7916 提到 0.7834 / 0.8065,在 GSO 上从 0.7722 / 0.7022 提到 0.7764 / 0.7069;对比 Photo3D(0.7380 / 0.7837、0.7512 / 0.7034)同样领先。论文已被 ECCV 2026 接收。 Qualitative comparison of 2D feedback sources (Reinforced Editing / NanoBanana Pro / Qwen-Image-Edit). 批判点评:消融表里最刺眼的一行:去掉 Source Branch 后 ΔCLIP / ΔDINO 从 +0.0379 / +0.0298 直接翻成 −0.0523 / −0.0497——比不做对齐还差。整套方法对「保留源分支」这一条设计极其敏感,稳健性存疑。变体对比里把分布匹配换成 DMD,分数掉到 0.5393 / 0.5486,说明这条路和当前主流的 3D 蒸馏范式并不兼容。另外绝对增益全部停在小数点后第二位,而「视觉保真」这件事到头来还是用 CLIP / DINO 相似度做 proxy,离人眼判断有多远没人说得清。 7. ComplexSync:复杂场景 70FPS 实时唇形同步 ComplexSync: High-Fidelity and Real-Time Lip Sync in Complex Scenarios | 广州趣丸网络科技 | arXiv:2609.29225 关键词:唇形同步, 扩散蒸馏, 实时推理, 视觉基础模型, 基准评测 ⚠️ 前序问题:唇形同步要让口型动态与语音精确对齐。扩散模型生成质量高,但在遮挡、侧脸、复杂光照等复杂场景下容易失稳,而且推理延迟高到没法实际部署。 本文贡献:ComplexSync 是一套统一的扩散框架,三块设计:一是双流联合训练策略,抑制参考帧的信息泄漏同时保留自然动态;二是基于蒸馏的单步去噪加速方案,做到 70+ FPS 吞吐;三是关系对齐损失,借助视觉基础模型(VFM)的结构先验提升复杂场景因素下的同步精度与鲁棒性。此外还给出首个专门面向复杂唇形同步的 benchmark,含 200+ 条挑战视频序列与专门指标。 Overview of ComplexSync. (a) Dual-stream joint training strategy designed to suppress spatial information leakage and ensure high-fidelity synthesis. (b) Distillation-based acceleration scheme that facilitates single-step denoising for real-time inference. (c) Relational alignment loss leveraging VFMs to enhance synchronization precision and robustness under unconstrained conditions. 实验效果:在标准场景与复杂场景上均达到 SOTA,同时支持实时推理,吞吐超过 70 FPS。训练分三阶段:第一阶段 8 卡 A100 训 200K 步,第二阶段 4 卡 A100 训 50K 步,第三阶段再训 10K 步,分辨率 512×512。 Qualitative comparison with existing methods. Given that static images cannot fully capture critical temporal attributes such as synchronization, naturalness, and stability, we provide comprehensive video comparisons in the supplementary materials. 批判点评:三阶段累计 260K 步、起步就是 8 张 A100,这个训练预算对一家业务公司来说更像工程投入而非可复现的研究。全文只在 512×512 分辨率验证,而真实配音、数字人场景常见的是 1080p 以上——放大后的口型稳定性没有数据。定性对比仍然用静态帧,同步、自然度、稳定性这些纯时间属性都推给了补充视频,而这恰是唇形同步最容易「单帧好看、连起来崩」的地方。 8. AdaPilot:一套策略零样本迁移各生成器 AdaPilot: Towards Scene-Adaptive Policy Learning for Cross-Generator Text-to-Image Quality Optimization | 南洋理工大学;同花顺研究院 | arXiv:2609.29517 关键词:文生图, 强化学习, 跨生成器迁移, 多轮视觉反馈, 奖励设计 ⚠️ 前序问题:提升文生图质量的路线已经从生成器微调、提示优化走到了带多轮视觉反馈的强化学习,但现有策略都与特定生成器、特定任务深度耦合,学到的能力很难泛化成一套通用的质量优化策略——换个生成器就得重训。 本文贡献:AdaPilot 把多轮图像生成建模成马尔可夫决策过程,用端到端 RL 学一个场景自适应、可跨生成器迁移的质量优化策略。三点关键:策略与生成器内部解耦(只通过 prompt 和编码后的视觉观测交互,不碰梯度与内部状态),因此能跨生成器复用;场景感知奖励(AdaReward)把质量评估维度与任务语义自适应对齐;过程级奖励建模图像质量在多轮里的演化轨迹。实现上用 Qwen2.5-VL-7B-Instruct 作 agent,Qwen-Image-2512 作冻结生成器。 An illustration of the proposed AdaPilot. 实验效果:在 7 个指标(CLIP-T、HPS、GDino、Aesthetic、OCR 用于奖励,Realism、Physics 为留出维度)上超过 Flow-GRPO、T2I-R1、T2I-Copilot 等基线。跨生成器评测中,单一策略零样本迁移到未见过生成器(Qwen-Image、GPT-Image-1、Gemini-3-pro-Image-Preview 等)时,在所有被评生成器上仍保持正平均增益,OOD 数据集上每个维度都领先。 Averages over applicable evaluation metrics for four unseen image generators on four out-of-distribution datasets and overall, comparing results with and without AdaPilot to evaluate zero-shot cross-generator transfer. 批判点评:论文自己给的失败案例很说明问题:一张要求对比八位利物浦前锋的密集信息图,多轮迭代确实把版面理清了,但文字仍基本不可读、球员身份与属性不一致、图表数值不可验证,agent 明知有残缺陷还是终止了。这暴露了「策略只改 prompt」的天花板——当缺陷出在生成器自身的文字渲染与数值编码能力时,再聪明的提示改写也救不回来。另外训练时要同时部署 CLIP ViT-L/14、HPS v2.1、Grounding DINO、GLM-OCR 和 Qwen2.5-VL-72B 五个评估器算奖励,这套推理栈的成本并没有被算进收益里。 9. ZoomDiff:双摄平滑变焦的高保真插值 ZoomDiff: A High-Fidelity Diffusion Model for Dual-Camera Smooth Zooming | 哈尔滨工业大学;淘宝(阿里巴巴集团) | arXiv:2609.28083 关键词:双摄变焦, 帧插值, 扩散模型, 高频重建, 时间一致性 ⚠️ 前序问题:双摄像头之间的数字变焦切换,几何结构和色彩一致性上会出现肉眼可见的跳变。现有双摄平滑变焦(DCSZ)方法多在帧插值模型上微调,扛不住大视差与复杂几何变换;直接套扩散式帧插值模型,又因为条件引导不足、VAE 编码丢高频、时间一致性不够,保真度还是上不去。 本文贡献:ZoomDiff 在潜空间与像素空间同时吃透双摄输入:一是多步去噪过程中强化双图条件引导,提升几何一致性;二是把 VAE 编码器的 flow 对齐多尺度特征注回解码器(ref injector),把 VAE 编码时丢掉的高频细节补回来;三是引入 flow 引导的时间一致性监督($\mathcal{L}_{ftc}$)让转场更顺。并给出把步数从 25 压到 8 / 4 的适配方案。 The dual-camera images ($\mathbf{X}_0$ and $\mathbf{X}_F$) are fed into the VAE encoder and the semantic encoder. 实验效果:合成与真实数据集上全面领先:PSNR 23.80、SSIM 0.761、LPIPS 0.253、PSNR-div 23.01、FVD 131.332,真实集 MUSIQ 73.634、LIQE 4.874、DBCNN 0.692、DOVER 0.665,均为最优。步数压缩上,25 步 45.89s(FLOPs 743T)可压到 8 步 21.28s(2.16×)与 4 步 15.71s(2.92×)。 Visual comparisons on the synthetic dataset and the real-world dataset. Our method still produces more consistent results. 批判点评:两个反差值得记。其一,扩散类基线整体拉胯:Wan2.1 只有 17.71 PSNR、TRF 14.15、Framer 的 DOVER 只有 0.335,全都输给光流法(UPRNet 22.58、RIFE LPIPS 0.261)——说明「生成式先验」在这个任务上并非天然优势,ZoomDiff 的赢面来自双图条件引导与高频注入这两个工程补丁,而非换 backbone。其二,步数压缩后质量反而上升:8 步版 PSNR 24.08 高于 25 步版的 23.80,而去掉步数适配的 8 步版只有 22.32。也就是说 25 步其实是过采样,指标对步数并不单调,任何「步数越少越差」的外推在这里都不成立。 10. ODU-Bench:14 个全模态模型的需求理解评测 Omni Demand Understanding: A Benchmark for Contextual User-Intent Inference in Multimodal Interaction | 上海交通大学;上海创智学院;阿里巴巴;香港中文大学;清华大学;香港理工大学;南开大学;约翰斯·霍普金斯大学 | arXiv:2609.21392 关键词:多模态交互, 需求理解, 评测基准, 误触发, MLLM ⚠️ 前序问题:音视频自然交互正成为 AI 助手的重要入口,但现有交互能力基准主要评「回复质量」,漏掉了更基础的一问:模型能不能从复杂的多模态交互里正确推断用户的真实需求?现实中的需求在口语里往往是欠规格的,得靠视觉线索、声学线索和对话历史补;含糊表达、噪声环境进一步加剧难度。反过来,听起来像请求的 utterance 也可能根本不是给助手的需求,导致误触发。 本文贡献:把 Omni Demand Understanding(ODU)定义成一个独立的多模态上下文推断问题:给定交互流,模型必须判断是否存在需求,并从多模态与对话上下文推断意图,评测覆盖五个维度(M1 需求存在性、M2 关键点命中率、M3 时间跨度 IoU、M4 转录、M5 用户画像),同时覆盖单轮与多轮。ODU-Bench 用挑战驱动的 taxonomy、taxonomy 引导的 agentic 视频生成加真人录制交互来构建,标注从互补媒体证据重建并经人工校验。 Overview of the ODU-Bench construction pipeline. Challenge-driven targets are expanded into coarse-to-fine scripts, screened for plausibility and challenge validity, and realized as synthesized or human-recorded interactions. 实验效果:评测 14 个原生 MLLM:最强者 Gemini 3.1 Pro 也只能恢复 44.7% 必须从视觉、声学或对话上下文推断的关键信息;14 个模型里有 11 个在非需求场景上的误触发率(FTR)超过 50%。 Qualitative error analysis on English generated audio-visual scenes. (a)--(b) Demand-present cases; M2 reports covered/reference key points. (c)--(d) No-demand cases. 批判点评:44.7% 和「11/14 误触发率 >50%」这两个数字,真正打脸的是当前「先响应、后理解」的范式——模型宁可硬猜一个需求,也不愿意说「这不是需求」。可惜基准本身大量依赖 agentic 生成的交互与 LLM judge,虽然论文做了 judge 稳定性检验和真人录制对照(Δ = Rec − Syn),合成数据与真人数据之间的分布差仍是这套结论最大的外推风险:如果 agentic 生成场景的「挑战性」分布和真实用户偏得比较多,44.7% 这个数就未必能平移到线上。 趋势观察 生成的瓶颈正在往「输入端」和「调度端」挪 今天这十篇里有三篇在动生成流程的两端而不是中间:WanPE 用 397B 模型重写提示词,把导演级分镜规划搬到文本空间完成;TRACK 不改模型、不重训练,只决定每一步该派大模型还是小模型上场;AdaPilot 则干脆只通过 prompt 和视觉观测跟生成器打交道。共同点是:底座模型被当成黑箱,可优化的空间被挪到了它前面和外面——这对没算力训底座的团队是好事,对卖底座的团队不是。 「砍掉教师与 critic」成为后训练的新共识 ViRDM 把 DMD 的教师+critic+生成器三件套砍成只训生成器,8 卡 A100 下显存 77.1→48.3 GB、时间 22h→2h,VBench 还涨了 0.36;OREO 同样放弃静态数据与蒸馏范式,改用 2D 编辑实时产出伪目标做自我监督;AV-GRPO 冻结一塔训另一塔来降低 RL 成本。三条独立工作指向同一判断:昂贵的在线监督预算正在被更便宜的离线/自产目标替换,代价往往是引入新的手工正则(ViRDM 的 flow 正则就是这么来的)。 评测基准开始盯「理解」而不是「生成」 ODU-Bench 评测 14 个原生 MLLM 后发现,最强的 Gemini 3.1 Pro 也只能恢复 44.7% 必须靠上下文推断的关键信息,而且 11/14 的模型在非需求场景上误触发率超过 50%。这个数字比任何生成质量榜单都更能说明当前语音/视觉助手的实际状态:问题不在答得好不好,而在该不该答都还没判断对。ComplexSync 也顺手补了首个复杂场景唇形同步 benchmark,200+ 条挑战序列把遮挡、侧脸这些长期被平均掉的case单独拎出来。 人工智能炼丹君 整理 | 2026-09-28 更多 AIGC 论文解读,关注微信公众号「人工智能炼丹君」 每日更新 · 论文精选 · 深度解读 · 技术脉络 微信搜索 人工智能炼丹君 或扫描下方二维码关注
2026年09月28日
2 阅读
0 评论
0 点赞
2026-09-27
AIGC 基本功|从 DDIM 到高阶采样器-DDIM
从 DDIM 到高阶采样器 所属方向:生成范式 | 难度:进阶 | 前置知识:DDPM 的训练目标与采样流程(ddpm)、扩散过程的前向与反向推导(diffusion_math) 关键词:DDIM、DPM-Solver、确定性采样、二阶采样、步数、NFE 01. 为什么需要它 先给两个数字,都来自文末附录里能直接跑的脚本。 数字一:同样的精度,评估次数差 12.8 倍。 在一个二维八分量高斯混合上跑(加噪之后仍是高斯混合,所以 score 有闭式解,见 04 节),把生成终点到高精度参考解的 RMS 距离压到 1e-2: 采样器 压到 1e-2 需要的模型评估次数 压到 1e-3 DDIM(时间步等间隔) 256 > 384 DPM-Solver-2(单步二阶) 63 191 DPM-Solver++ 2M(多步二阶) 40 128 DPM-Solver++ 3M(多步三阶) 20 64 DDIM 要 256 次评估,3M 只要 20 次。这不是"实现优化"级别的差距,是换了一套数学带来的差距——同一个模型、同一份起点、同一份随机数,只是步法不同。生产环境里 20 步出图还是 250 步出图,是"能上"和"不能上"的区别。 数字二:换个采样器,缓存策略就得重调。 这一条更隐蔽,也更容易踩。所有"复用上一步特征"的加速(DeepCache、各种 step-cache、特征缓存)都在做一个判断:这一步的模型输入和上一步差不多,就直接搬上一步的中间结果。那"差不多"到底是多少?在同一个模型、同一条 50 步轨迹上,量最后十步"模型输入相对变化": 采样器 最后 10 步输入相对变化 DDIM(η=0) 0.0197 η=0.5 0.0667 DDPM 祖采样(η=1) 0.1294 η=1 的抖动是 η=0 的 6.6 倍。也就是说,你在 DDIM 上把阈值调成"输入变化小于 2% 就复用",切到 DDPM 祖采样后这条曲线整段都在 13% 附近,缓存一次都不会命中,加速方案静默失效;反过来,为了吃掉 η=1 的抖动把阈值放宽到 15%,切回 DDIM 后就会在真正需要重算的步上复用,画面细节直接掉。这不是调参没调好,这是两套采样器的轨迹性质不同。 所以这篇文章要回答的是一个问题:采样器到底是什么? 答案不是"一种跳步策略",而是一条常微分方程的离散格式。一旦这么看,两件事立刻变得可算:步法是多少阶、达到给定精度要花多少次模型评估(NFE,number of function evaluations)。而 DDIM——这个被无数人当成"DDPM 的加速版"的东西——恰好是这条 ODE 上最低的一阶格式。下面会把它验到 1e-14。 02. 最小可用理解 三句话讲完: DDIM 不是"步幅更大的 DDPM"。 DDPM 的采样每步都要掷一份新噪声,DDIM 的贡献是证明:同一族前向过程可以用非马尔可夫的方式重新构造,只要边缘分布 $q(x_t|x_0)$ 不变,反向过程就有自由度——其中一个极端是完全不掷噪声。所以差别不在步幅,在随机项。 不掷噪声之后,采样变成解 ODE。 这个 ODE 叫概率流 ODE,它的解是确定性的:同一个起点永远给同一个终点。DDIM 的一步迭代,就是这条 ODE 的一个显式一阶格式(用当前点的信息走一整步)。用二维真轨迹看最直观: 这张图要看的是:蓝色的 η=0 轨迹是一条平滑曲线,橙色 η=1 轨迹在同一份随机数下走成了折线——每一步都被注入的方向改变。"同一个采样器家族"的两种极端,轨迹的几何性质完全不同,缓存类加速盯的正是这个几何性质。 既然它是 ODE,就有阶数,就有"性价比最高的一次评估花在哪"的问题。 把状态量从 $t$ 换成 $\lambda=\log(\alpha_t/\sigma_t)$(信噪比的对数)之后,一阶格式的误差正比于步长、二阶格式正比于步长的平方。于是同样的精度,高两阶的方法能省一个数量级的评估——这就是 DPM-Solver 系列的全部立足点。 03. 数学推导 3.1 把 DDIM 的一步拆开 DDPM 的采样式(祖先采样)长这样:先估干净图,再缩放,再加一份后验噪声。 $$x_{t-1} = \frac{1}{\sqrt{\alpha_t}}\Big(x_t - \frac{\beta_t}{\sqrt{1-\bar\alpha_t}}\epsilon_\theta(x_t,t)\Big) + \sqrt{\tilde\beta_t}z$$ DDIM 论文式 (12) 换了一个组织方式。令 $\hat x_0$ 是这一时刻对干净图的估计、$\epsilon_\theta$ 是模型输出的噪声预测,两者由 $\hat x_0 = (x_t-\sqrt{1-\bar\alpha_t}\epsilon_\theta)/\sqrt{\bar\alpha_t}$ 相互确定,则一步写成三项: $$x_{t-1} = \underbrace{\sqrt{\bar\alpha_{t-1}}\hat x_0}_{\text{回数据的方向}} + \underbrace{\sqrt{1-\bar\alpha_{t-1}-\tilde\sigma_t^2}\epsilon_\theta}_{\text{沿当前噪声方向}} + \underbrace{\tilde\sigma_t z}_{\text{新掷的噪声}}$$ 这里第二项必须使用目标时刻的总噪声方差 $1-\bar\alpha_{t-1}$,新注入噪声为 $$\tilde\sigma_t = \eta\sqrt{\frac{1-\bar\alpha_{t-1}}{1-\bar\alpha_t}}\sqrt{1-\frac{\bar\alpha_t}{\bar\alpha_{t-1}}}$$ 每一项在干什么。 第一项是"这一时刻认为的干净图"乘上新时刻的 $\sqrt{\bar\alpha}$,也就是把估计值放到目标噪声水平上。第二项补偿两者的差额:前向过程是 $x_t=\sqrt{\bar\alpha_t}x_0+\sqrt{1-\bar\alpha_t}\epsilon$,要从 $x_t$ 走到 $x_{t-1}$,得把噪声幅度从 $\sqrt{1-\bar\alpha_t}$ 缩到 $\sqrt{1-\bar\alpha_{t-1}}$,这个缩放只能沿着"当前噪声方向"做——$\sqrt{1-\bar\alpha_{t-1}-\tilde\sigma_t^2}$ 正好是这两者的勾股差。第三项才是唯一带随机性的部分。 $\eta$ 是那个旋钮。 $\eta=0$ 时第三项消失,$\tilde\sigma_t=0$,迭代完全确定;$\eta=1$ 时 $\tilde\sigma_t$ 恰好等于 DDPM 的后验标准差,第二、三项合起来正好还原 DDPM 的加噪——在相同的完整时间表、fixed-small 方差、预测器及随机数下,$\eta=1$ 与该 DDPM 更新代数等价;跳步时应比较对应重排时间表的祖先采样,而不是原始 1000 步链。04 节会把这两条都量到浮点误差量级。 为什么可以有这个旋钮。 因为 DDIM 用的前向过程不是马尔可夫链:用 $q_\eta(x_{t-1}\mid x_t,x_0)$ 与终点分布重新构造非马尔可夫联合过程,但所有边缘 $q(x_t|x_0)$ 一个都没变。训练目标只依赖边缘(噪声预测的回归目标就是 $\epsilon$),所以模型不用重训;而反向的可选空间变大了,$\eta$ 就是在里面挑一条路。 3.2 换成 ODE 视角 $\eta=0$ 时迭代没有随机项,它可以被看成一条 ODE 的离散格式。连续时间下,前向过程是 $dx = -\frac{1}{2}\beta(t)x\mathrm{d}t + \sqrt{\beta(t)}\mathrm{d}w$(第一项把 $x$ 往 0 拉,第二项持续注入白噪声),对应的概率流 ODE(与 SDE 共享所有边缘分布的那个确定性方程)是 $$\frac{dx}{dt} = -\frac{1}{2}\beta(t)\Big(x + \nabla_x\log p_t(x)\Big)$$ 两式一比就能看出这个方程的来历:把 SDE 的噪声项换成"噪声的均值流",也就是 score 项,随机性就被抽掉了,但每一时刻的边缘分布 $p_t(x)$ 一模一样。这就是为什么用同一条训练好的网络既能跑随机采样(SDE)、也能跑确定性采样(ODE)。 score 与噪声预测的关系是 $\nabla_x\log p_t(x) \approx -\epsilon_\theta(x,t)/\sqrt{1-\bar\alpha_t}$,代进去就得到一个只用模型输出的 ODE。关键一步是换坐标。 令 $\lambda = \log(\alpha_t/\sigma_t) = \frac{1}{2}\log\frac{\bar\alpha_t}{1-\bar\alpha_t}$(半 log-SNR)。 这个坐标为什么叫"半":信噪比本身就是 $\alpha_t^2/\sigma_t^2 = \bar\alpha_t/(1-\bar\alpha_t)$,取对数再除以二,得到的是"振幅比"的对数——也就是 $x_0$ 与 $\epsilon$ 两个分量在 $x_t$ 里的相对尺度的对数。$\lambda$ 越大代表越干净:$t=1000$ 时 $\bar\alpha_t\approx 4\times10^{-5}$,$\lambda\approx-5$;$t=1$ 时 $\bar\alpha_t\approx0.9999$,$\lambda\approx+4.6$。整条轨迹在 $\lambda$ 上只跨了不到 10 个单位,而 $t$ 跨了 1000 个单位——这就是"难度不是按 $t$ 均匀分布"的定量说法。 为什么要换?因为在这条 ODE 里,步长的含义由 $\lambda$ 决定:$t$ 上均匀的一步,落在 $\lambda$ 上的长度可能差几十倍。$\lambda$ 坐标把 ODE 拉成接近常数系数的形式($\alpha_\lambda$、$\sigma_\lambda$ 随 $\lambda$ 的变化是光滑的 sigmoid 型),指数积分才有干净的闭式。换成 $\lambda$ 之后,一步从 $\lambda_s$ 走到 $\lambda_t$(约定 $h=\lambda_t-\lambda_s>0$,因为我们从噪声走到干净)的精确解是 $$x_t = \frac{\sigma_t}{\sigma_s}x_s + \alpha_t\int_{0}^{h} e^{-(h-u)}\hat x_0(\lambda_s+u)du$$ 这个式子值得停一下:它说明整步的解由"$\hat x_0$ 沿着 $\lambda$ 的变化曲线"加权积分决定,权重是 $e^{-(h-u)}$——越靠近区间右端(越干净的那端)权重越大。验证一下:如果 $\hat x_0$ 是常数,积分给出 $\alpha_t(1-e^{-h})\hat x_0$,加上第一项正好是 $\alpha_t \hat x_0+\sigma_t \epsilon$,也就是"一步跳到位"的精确解。所以步法阶数的本质,就是用多少个点上的 $\hat x_0$ 去逼近这条曲线。 3.3 阶数从哪来:泰勒展开 把 $\hat x_0(\lambda_s+u)$ 在 $u=0$ 处展开,记 $\hat x_0,\hat x_0',\hat x_0''$ 为该点的一至二阶导,逐项积出系数: $$x_t = \frac{\sigma_t}{\sigma_s}x_s + \alpha_t\Big[J_0\hat x_0 + J_1\hat x_0' + J_2\hat x_0'' + \cdots\Big]$$ $$J_0 = 1-e^{-h},\qquad J_1 = e^{-h}-1+h,\qquad J_2 = \frac{h^2-2h+2-2e^{-h}}{2}$$ 小 $h$ 展开看数量级:$J_0\approx h$、$J_1\approx h^2/2$、$J_2\approx h^3/6$。所以导数项每高一阶,整步误差就多一个 $h$。 一阶格式:只留 $J_0$,把 $\hat x_0$ 当成整步不变。局部误差 $O(h^2)$,全局 $O(h)$。 二阶格式:再加 $J_1\hat x_0'$,系数 $\varphi_2 = J_1/h$,用一次评估估出 $\hat x_0'$ 就行。 三阶格式:再加 $J_2\hat x_0''$,系数 $\varphi_3 = J_2/h^2$。 多步法(multistep)就是"不额外花评估,用前几步的 $\hat x_0$ 做差商估导数":设这一步的 $\lambda$ 为 $\lambda_{s_0}$,前两步在 $\lambda_{s_1}$、$\lambda_{s_2}$,令 $h_i=\lambda_{s_i}-\lambda_{s_{i+1}}$,则 $$D_{1,0} = \frac{h}{h_0}(\hat x_0^{s_0}-\hat x_0^{s_1}) \approx h\hat x_0',\qquad D_1 = D_{1,0} + \frac{r_0}{r_0+r_1}(D_{1,0}-D_{1,1})\approx h\hat x_0'$$ 上一步的 $O(h)$ 偏差被组合系数消掉了(这就是"两个差商加权成一阶导"的标准手法),所以 $D_1$ 对 $\hat x_0'$ 是二阶准确的。同理二阶导用二阶差商: $$D_2 = \frac{D_{1,0}-D_{1,1}}{r_0+r_1}$$ 这里有个必须预先说清的细节。 三个点上的二阶差商等于 $\hat x_0''/2$(这是差商定义直接给出的),所以 $D_2\approx h^2\hat x_0''/2$;而上面泰勒系数 $\varphi_3$ 是按 $h^2\hat x_0''$ 的量纲推的。两者一比差一个 2。这个 2 在 05 节会展开:上游实现里单步版带了它、多步版没带,我先用受控实验把"该不该带"量出来,再决定怎么在文章里说。 3.4 DDIM 就是一阶格式 把 03.1 的 DDIM 迭代($\eta=0$)和"只留 $J_0$"的一阶格式对照:$\eta=0$ 时第二项系数为 $\sqrt{1-\bar\alpha_{t-1}}$,第三项为 0,代入 $\hat x_0$ 的定义整理,得到 $$x_t = \frac{\sigma_t}{\sigma_s}x_s + \alpha_t(1-e^{-h})\hat x_0(\lambda_s)$$ 与一阶格式逐项相同。所以"DDIM 是一阶方法"不是比喻,是恒等式——04 节用两套独立实现的代码把它验到 1e-14 以内。 04. 代码实现 全部代码在文末附录,五个脚本、只依赖 numpy:oracle_gmm.py(实验台与尺子)、ddim_family.py(η 家族)、dpm_solver_lab.py(λ 坐标与高阶格式)、spacing_lab.py(步数摆法)、make_figures.py(配图)。下面的数字是它们的真实输出。 4.1 先造一把可靠的尺子 要在二维上量"采样器差多少",需要一个模型误差为零的环境,否则量到的是模型不行,不是步法不行。用八个高斯分量摆在一圈上,加噪之后仍是高斯混合,score 有闭式解——模型换成这个解析 score,误差就只剩离散化。oracle_gmm.py 自己先做两件事: A. Oracle 自检:解析 score vs 有限差分 t= 1 abar=0.999900 max|解析 - 差分| = 1.57e-08 t= 200 abar=0.659039 max|解析 - 差分| = 4.99e-11 t= 1000 abar=0.000040 max|解析 - 差分| = 2.14e-11 B. 端点检查:t=T 时协方差与 I 的最大差 = 3.834e-05 C. 尺子的噪声地板:SW1 真vs真 ×8 = 0.03055 ± 0.00452 D. 天花板:真样本当生成样本送进去 SW1 = 0.000000 第一项是"我写的解析 score 不是拍脑袋写的"(和有限差分对上);第二项确认 $t=T$ 时分布确实接近标准正态(否则起点选错了);第三项最重要:两批都是真样本时也能量出 0.03 的距离,这是抽样噪声的地板,后面低于该尺度的差别需要更多样本和置信区间确认,不能仅凭单次结果归因;第四项确认尺子无偏置。 4.2 最小实现:DDIM 的一步 ddim_family.py 的核心就是下面这个函数,变量名与 3.1 节一一对应: def ddim_step(x, abar_cur, abar_tgt, eps, z, eta): """从 abar_cur 走到 abar_tgt(噪声变小)。返回 (x_tgt, x0_hat)。""" a_c, s_c = alpha_sigma(abar_cur) # 当前时刻的 alpha_t / sigma_t a_t, s_t = alpha_sigma(abar_tgt) # 目标时刻的 alpha_t / sigma_t x0 = (x - s_c * eps) / a_c # 第一项用的 x0_hat var = (1 - abar_tgt) / (1 - abar_cur) * (1 - abar_cur / abar_tgt) sig_tilde = eta * np.sqrt(max(var, 0.0)) # 第三项:新掷噪声的幅度 direc = np.sqrt(max(s_t ** 2 - sig_tilde ** 2, 0.0)) # 第二项:勾股差 return a_t * x0 + direc * eps + sig_tilde * z, x0 最后一行就是 3.1 节那三项,一次加法写完,没有任何按时刻分支的特判。跑个小验证,看 $\tilde\sigma_t$ 在轨迹两端各占多少: _, ABAR = linear_schedule() ts, grid = uniform_t_grid(50) # 51 个 abar,grid[0] 最吵、grid[-1] = 1.0 x, zs = paired_noise(50, n=4, seed=4321) # 起点 + 每一步要用的随机数 print("x.shape =", x.shape, "| grid 长度 =", len(grid)) for i in (0, 25, 48): for eta in (0.0, 1.0): a_c, s_c = alpha_sigma(grid[i]) a_t, s_t = alpha_sigma(grid[i + 1]) var = (1 - grid[i+1]) / (1 - grid[i]) * (1 - grid[i] / grid[i+1]) sig_tilde = eta * np.sqrt(var) print(f"step {i:2d} eta={eta:.1f} sigma_tilde = {sig_tilde:.6f}" f" (该步总噪声幅度 {s_t:.4f})") 真实输出: x.shape = (4, 2) | grid 长度 = 51 step 0 eta=0.0 sigma_tilde = 0.000000 (该步总噪声幅度 1.0000) step 0 eta=1.0 sigma_tilde = 0.574284 (该步总噪声幅度 1.0000) step 25 eta=0.0 sigma_tilde = 0.000000 (该步总噪声幅度 0.9509) step 25 eta=1.0 sigma_tilde = 0.419845 (该步总噪声幅度 0.9509) step 48 eta=0.0 sigma_tilde = 0.000000 (该步总噪声幅度 0.0760) step 48 eta=1.0 sigma_tilde = 0.063819 (该步总噪声幅度 0.0760) 两个数值得盯一下。η=1 时 $\tilde\sigma_t$ 与该步总噪声幅度之比,在起点是 57%(0.574/1.000),到接近干净端的第 48 步反而是 84%(0.0638/0.0760)。 也就是说"加噪声"的代价在轨迹末端最重:那里 $\hat x_0$ 刚要把几个模式分辨开,一份占 84% 的新噪声就砸进去了。这是 4.6 节 η=0 处处最优的直接原因。 再看一步里三项谁大谁小(用 oracle score 走第 25 步): eta=1 第 25 步三项 RMS: 回数据 0.1742 | 沿噪声方向 0.8694 | 新掷噪声 0.7800 两个噪声项的模长比"回数据"项大四五倍。三项向量的范数反映更新的组成,不能用向量大小推出网络算力花在哪里,而不是往数据方向推——这也解释了为什么"阶数"值钱:阶数讲的正是怎么更准地把这一步搬完。 4.3 两次交叉验证 ddim_family.py 里故意写了两套互不相干的代码:一套照 DDPM Algorithm 2 的原式走 1000 步(一步不能跳),一套照 DDIM 式 (12) 走。喂同样的随机数: A. 终点逐元素最大差 = 3.109e-14 终点 RMS = 1.3520 3.1e-14 就是浮点累加误差的量级。η=1 的 DDIM 与 DDPM 祖采样是同一个算法,不是"近似"。 dpm_solver_lab.py 里同样对待 DDIM 与一阶指数积分器: A. max|DDIM(eta=0) - DPM-Solver-1| S= 20 uniform-t 8.549e-15 uniform-lambda 4.774e-15 S= 50 uniform-t 1.044e-14 uniform-lambda 1.088e-14 两套记号、两套时间网格,输出差在 1e-14。DDIM 是一阶指数积分器,这句到这里可以当结论用了。 4.4 量阶数:误差 vs NFE 有了参考解(多步三阶跑 4000 步,再用 2000 步自查,分辨率 1.16e-08)就能量阶数了。横轴必须是真实 NFE——脚本里用 Counter 包住模型,把调用次数数出来,而不是拿"步数 × 理论阶数"算。 这张图要看的是斜率:DDIM 的误差衰减阶约 1,二阶方法约 2,修正系数后的三阶方法接近 3;误差更小的曲线位于图的下方。纵向的差距就是"同样的评估次数,精度差多少个数量级"。 拟合出来: 采样器 拟合阶数 p 备注 DDIM(时间步等间隔) 0.97 一阶 DDIM(λ 等间隔) 0.99 一阶,换坐标不改阶数 DPM-Solver-2(单步) 1.90 二阶,每步 2 次评估 DPM-Solver++ 2M(多步) 1.98 二阶,每步 1 次评估 DPM-Solver++ 3M(原式) 2.07 应该是三阶,实测只有二阶 3M(三阶项系数 ×2 后) 2.91 接近三阶 3M 那两行是全篇唯一"和教科书不一致"的地方,见 05 节。 4.5 步数与质量的实测表 阶数是数值分析的语言,产品要的是"20 步够不够"。同一批配置换成 SW1(切片 Wasserstein-1,越小越好),换几份起点量出抖动: NFE DDIM(t) DDIM(λ) Solver-2 2M 3M(原式) 10 0.1292±0.006 0.2180±0.006 0.1910±0.007 0.0731±0.006 0.0680±0.006 20 0.0741±0.004 0.1122±0.004 0.0568±0.004 0.0317±0.006 0.0277±0.005 50 0.0420±0.007 0.0516±0.007 0.0335±0.006 0.0328±0.005 0.0327±0.005 噪声地板是 0.0301±0.0026。请只看 NFE=20 那一行:20 次评估时各采样器之间的差距(0.0277 到 0.1122)远大于抖动;到了 50 次评估,所有方法都贴着地板(0.033~0.052),这张表就再也排不出名次了——不是"大家都一样好",是当前样本量和重复次数难以稳定区分,增加样本与重复实验仍可提高分辨力。要排名次得回到 4.4 的确定性误差表。 4.6 η 到底该调多大 ddim_family.py 在 oracle 下扫一遍 η,再在欠拟合模型下扫一遍。所谓欠拟合是让模型以为数据是单个高斯(用真均值真协方差拟合),它在高噪声区几乎是对的、在低噪声区错得离谱——这是人为选择的一种误差结构,并不代表所有真实网络(也正好对应"训练时看到的是加噪数据分布"这件事)。 这张图要看的是两件事:一是曲线从左到右单调下降(步数越多越好,符合预期);二是本玩具大多数设置下 η=0 的 SW1 更低,不能推广为所有模型中确定性采样必胜,至少在有闭式解、且模型差得很有代表性的两种情况下都是错的。右图还给出一个诚实的例外:N=100 时 η=1 的 0.1830 略低于 η=0 的 0.1845,但差值 0.0015 远小于这批量的抖动,不该当成结论。 4.7 缓存为什么会被采样器带崩 这张图要看的是:三条线在干净端都会收敛(步长趋于 0),但在中间段整段差 6.6 倍。缓存阈值是按这条线定的,换采样器就是换这条线的量级。同一份数据里 $x_0$ 估计的位移更夸张:η=0 是 0.0371,η=1 是 0.2407(6.5 倍)。 4.8 评估预算 这张图要看的是柱子的相对高度,以及"步数"和"评估次数"不是一回事:DPM-Solver-2 每步要 2 次评估(先踩一步到中点、再走完),所以它的 NFE 是步数的两倍,这也是它虽然二阶却在低 NFE 段不占优的原因。 4.9 步数该摆在哪:收尾步不是免费的 阶数回答的是"步数变多时误差掉多快",没回答"步数摆在哪"。spacing_lab.py 固定 NFE=20、只改摆法,终点误差如下: 摆法 DDIM(一阶) 2M(二阶) 在 $t$ 上等间隔(本文的时间网格) 1.107e-01 5.285e-02 在 $\lambda$ 上等间隔 1.713e-01 2.897e-02 在 $\lambda$ 上等间隔、落回整数 $t$ 1.712e-01 2.894e-02 Karras / EDM($\rho=7$) 1.873e-01 7.038e-02 Karras($\rho=3$) 2.655e-01 3.397e-01 阶数和网格要配套:一阶方法在 $t$ 上等间隔最好,二阶方法在 $\lambda$ 上等间隔最好(差 1.8 倍)。而 EDM 那套 $\rho=7$ 的摆法在这里垫底——它不是不好,是它的 $\sigma$ 区间(EDM 原文是 80→0.002)和 VP 调度(这里只有 158→0.01,注意 EDM 的 $\sigma$ 是 $\sigma_t/\alpha_t$,不是 $\sqrt{1-\bar\alpha_t}$)不是一回事,照搬会把最吵的那一步拉得极长:NFE=20 时它最大的 $\lambda$ 间隔是 1.02,而 $\lambda$ 等间隔只有 0.51。 uniform-t 为什么赢,值得追问一层。 它把 $\lambda$ 的上界截在了 +1.76(NFE=20 时),剩下的干净端交给最后那一步"直接输出 $\hat x_0$"。这条捷径是不是免费的?不是。沿高精度参考轨迹量"在 $\lambda$ 处直接输出 $\hat x_0$ 会偏多少": $\lambda$ −0.23 +0.74 +1.71 +2.67 +3.64 提前收尾的偏差 5.24e-01 1.40e-01 4.40e-02 9.39e-03 1.63e-03 截得越早白送的偏差越大,而且这份偏差不随步数下降——它只取决于截在哪个 $\lambda$。于是 $\lambda_{max}$ 存在最优值(2M,NFE=20,只改截断点): $\lambda_{max}$ 1.0 2.0 2.5 3.0 4.61(走到底) 终点误差 9.82e-02 2.80e-02 1.71e-02 1.81e-02 2.90e-02 U 形,最优在 2.5,比"走到底"好 1.7 倍。这条也解释了 uniform-t 的行为为什么还不错:它的截断点随步数自动往右挪(NFE=20 时 1.76、50 时 2.55、100 时 3.1),等于一个廉价的自适应策略——但它是巧合,不是设计。 顺带一个反直觉的观察。 一步的局部误差带着一个 $\alpha_t$ 因子($\delta_i\approx\alpha_{t_i}\frac{h^2}{2}\hat x_0'$),所以误差权重是 $\alpha_t\,|d\hat x_0/d\lambda|$ 而不是单纯的导数。量下来:噪声端 1.02e-01、干净端 1.00e-01,几乎一样重,峰在 $\lambda\approx-0.2$(也就是 $\bar\alpha\approx0.4$ 的中间段)。常听到的"把步数往干净端挪"在本实验里不成立——难度是按 $\lambda$ 均匀铺开的,真正要避免的是"某一段特别长"(Karras 就是栽在这)。 05. 工业级实现对照 以 huggingface/diffusers 的 DPMSolverMultistepScheduler.step 为准(https://github.com/huggingface/diffusers/blob/main/src/diffusers/schedulers/scheduling_dpmsolver_multistep.py,以 2026-09 时的实现为准,上游会重构),最小实现与生产实现的差距集中在四处。 第一处:算法类型的开关是"预测什么"。 algorithm_type 有四个取值:dpmsolver 用噪声预测($\epsilon$-parameterization),dpmsolver++ 用数据预测($x_0$-parameterization);后缀 -sde-dpmsolver++ 则在同一条 ODE 上加回一个可控的噪声项。本实验包含噪声预测单步 Solver-2 与数据预测多步 Solver++。参数化会改变数值误差和引导稳定性,不能由训练 loss 的参数化权重图证明某条预测曲线总更光滑。 第二处:低步数时的降阶。 多步法开头没有历史可用,第一、二步必须降成一阶/二阶,lower_order_final 控制最后一步是否也降阶,final_sigmas_type="zero" 让最后一步直接落到 $\sigma=0$。这一条在我们的实现里对应"$\lambda_t=+\infty$ 时 $e^{-h}=0$、整式退化成 $x_t=\hat x_0$",不用特判。我专门验证过"开头降阶"是不是 3M 阶数不达标的原因:给第一步多加一次评估换成中点法(warm2s=True),实测阶数从 2.10 变成 2.10——一点没变,所以瓶颈不在起步。 第三处:时间步的排布。 timestep_spacing 有 leading、trailing、linspace(当前 DPMSolverMultistepScheduler 默认) 三种,同一个调度器换一种排布,20 步出图质量差很多。我们把它量化了:DDIM 在 $t$ 上等间隔时,NFE=191 的终点误差是 1.245e-02;换成 $\lambda$ 等间隔反而变差到 1.775e-02,达到 1e-2 分别需要 256 次和 384 次评估。这是"坐标选择比阶数更早起作用"的直接证据:一阶方法的误差由最大的那一段步长决定,均匀 $\lambda$ 的 $h$ 本来就是相等的;它改变了各时间段的评估密度,误差还由场的导数、传播与终端截断决定。所以别把"回到 λ 坐标"当成万能钥匙——它给高阶方法提供了干净的积分形式,但不自动给一阶方法更好的网格。 第四处,也是唯一需要标注存疑的:三阶项的系数。 我们逐行对齐 diffusers 的多步三阶更新后,实测阶数只有 2.07(4.4 节),不是 3。于是做了一个受控实验:不让 $\hat x_0$ 由 $x$ 决定,而是直接规定它是 $\lambda$ 的二次多项式,此时精确解可以用数值积分算出来,任何真正的三阶格式都应该一步算准。把区间长度 $h$ 从 0.4 缩到 0.025: 一阶: 4.17e-02 1.04e-02 2.57e-03 6.39e-04 1.59e-04 → h^2 ✓ 二阶: 8.16e-03 1.05e-03 1.32e-04 1.65e-05 2.07e-06 → h^3 ✓ 三阶-D2原式: 1.96e-03 2.46e-04 3.06e-05 3.81e-06 4.75e-07 → h^3 ✗ 三阶-D2乘2: 4.43e-13 6.23e-14 2.89e-15 1.11e-15 4.44e-16 → 一步精确 ✓ 一阶、二阶的收敛阶都对得上公式推导;三阶原式只有 $h^3$(等于二阶),把 $D_2$ 乘 2 之后直接掉到 1e-13~1e-16 的机器精度。我还用非均匀步长($h_0\neq h_1\neq h$)复验过一遍,结论不变:原式 8.3e-04、乘 2 之后 5.4e-14,所以这不是"等间隔时才碰巧"的巧合。 回头逐字核对上游源码(diffusers v0.30.0):单步版 scheduling_dpmsolver_singlestep.py 里写的是 D2 = 2.0 * (D1_1 - D1_0) / (r0 - r1),带了那个 2;多步版 scheduling_dpmsolver_multistep.py 里写的是 D2 = (1.0 / (r0 + r1)) * (D1_0 - D1_1),没带。两个版本共用同一个 $\varphi_3$ 系数 - (alpha_t * ((exp(-h) - 1.0 + h) / h**2 - 0.5)) * D2。 我把这件事的处理方式说清楚:这篇文章只报告我们自己的复现结果——逐行对齐上游的多步三阶实现、实测阶数 2.07;把 $D_2$ 乘 2 后升到 2.91,局部误差检验一步算准。至于这是上游实现的笔误、还是我对多步差商归一化的理解和原作者不同,我不下结论,已记入待决清单单独确认。对读者的实用含义是确定的:不要假设"用了 3M 就是三阶",阶数要自己量。 06. 代价与边界 省了什么。 DDIM 把每步的随机项去掉,换来三件事:同样的步数下误差更小(4.6 节的两个模型都验证了)、轨迹光滑(缓存加速可用)、以及固定初始噪声时可复现的输出。正则连续 ODE 流可逆,有限步 DDIM 或终步投影并不自动保证一一对应。 赔了什么。 路径随机性少了。 给定起点只有一条轨迹,但更换初始噪声仍产生多样本;确定性 ODE 不意味着输出分布没有多样性,也没有额外规定必须付出的批量成本。 高阶方法的稳定域有限。 显式格式的稳定性有上限,步长太大时高阶项不是"更准"而是"发散"。生产实现里低 NFE 时的降阶、以及当前多步调度器的 solver_type="midpoint" / "heun" 选项;它们不能和作者单步 API 中的 dpmsolver / taylor 选项混用都是在拿稳定性换名义阶数。 理论阶数要靠光滑性支撑。 我们的实验用的是解析 score,$\hat x_0(\lambda)$ 足够光滑,阶数才量得出来。真实网络输出的 $\hat x_0$ 带高频抖动,阶数通常要打折扣——这解释了为什么"3 阶"在实践中省下的没那么多。 什么时候不该用确定性采样器。 需要随机性作为正则的场景(例如低步数下用 SDE 采样换取更好的分布覆盖、或者需要"温度"调节多样性),$x_0$ 强约束会导致过平滑;以及任何依赖"步步重掷噪声"来做布朗桥类操作的训练/蒸馏流程。 什么时候值得上高阶。 评估预算在 20~50 次这个区间时最划算(4.8 节:NFE=20 时 3M 已经压到 1e-2,DDIM 需要 256 次);一旦预算到几百次,所有方法都进入渐进区,选最便宜的一阶反而更省心。 07. 经典论文脉络 DDIM,arXiv:2010.02502(2020):把 DDPM 的马尔可夫反向替换成非马尔可夫族,边缘不变、反向可选,给出 $\eta$ 旋钮。它的历史意义是把"采样"从随机过程问题变成确定性 ODE 问题。 DPM-Solver,arXiv:2206.00927(2022):把 ODE 换到 $\lambda$ 坐标做指数积分,给出二阶/三阶的单步格式与"约 10 步出图"的结果。核心贡献是 $\lambda$ 坐标与精确的指数积分系数,阶数第一次变得可算。 DPM-Solver++(2022):把参数化从 $\epsilon$ 换成 $x_0$,并补上多步变体(每步只要 1 次评估),这才是现在框架默认调度器的形状;同时给出引导采样的稳定化处理。 EDM,arXiv:2206.00364(2022):把"调度(噪声表)"和"采样器(ODE 积分格式)"彻底解耦,并给出把任何 $\sigma$ 上的网络包装成统一 ODE 的框架。看完 EDM 再回看 DDIM/DPM-Solver,会发现它们只是同一套 ODE 的三组离散格式。 一致性模型 / LCM(2023):把"多步 ODE 积分"压成"一步直接映射到 $x_0$",用蒸馏替代阶数。它的定位不是"更高阶的采样器",而是"不需要采样器"。 08. 常见误解 误解一:DDIM 就是步幅更大的 DDPM。 不是。DDIM 可以走 $t$ 上的任意子序列(这是它"能跳步"的前提,不是它的内容),它的内容是去掉随机项。反过来也成立:确定性 DDIM(η=0)也可以走 1000 步,仍不同于随机 DDPM;η=1 配合对应完整时间表时才有前文的等价关系。把这两件事分开之后,"为什么 DDIM 的 1000 步也不等于 DDPM"这个问题才有答案。 误解二:步数少了,就把 η 调大一点补回来。 实测反了。oracle 下 η=0 在 N=10/20/50/100 全部最优(N=10 时 0.1321 对 0.1888);欠拟合模型下同样单调(N=10 时 0.2679 对 0.3496)。直觉为什么错:加噪声确实让每步的 x̂₀ 估计被"抹匀"一点,但代价是它也抹掉了本该累积的方向信息;而噪声项还会让下一步的输入抖动变大(6.6 倍那个数字),把误差一路带下去。 误解三:阶数越高越省。 一是不一定真拿到那个阶(4.4 节 3M 实测 2.07);二是低 NFE 段高阶方法要"攒历史",前几步被迫降阶,实测 NFE=10 时 2M(0.0731)确实明显好于 DDIM(0.1292),但 3M(0.0680)比 2M 只好一点点,收益远小于"二阶到三阶"的名义差距。 误解四:SW1/FID 分不出来就没差别。 4.5 节里 50 次评估时全部方法都贴在 0.031 附近,看着"都一样"——那是因为尺子的噪声地板就是 0.0301。指标贴着地板时不代表方法等价,只代表这个指标失效了,此时要换确定性指标(有真值时用 RMS,我们的表就是这么打的)。 误解五:换到 λ 坐标只会更好。 不保证。λ 坐标便于推导指数积分,但选择 λ 均匀网格是另一件事;本文的一阶实验反而更偏好 t 网格:DDIM 在 $\lambda$ 等间隔网格上误差更大(NFE=191 时 1.775e-02,而 $t$ 等间隔是 1.245e-02)。网格与阶数要配套选。 09. 动手验证 五个脚本都能直接跑,几分钟内出结果。建议按这个顺序试: cd outputs/fundamentals_files/ddim_samplers/code /usr/local/bin/python3 oracle_gmm.py # ① 标定尺子(先量地板) /usr/local/bin/python3 ddim_family.py # ② η 旋钮与缓存抖动 /usr/local/bin/python3 dpm_solver_lab.py # ③ 阶数、预算、三阶项的受控检验 /usr/local/bin/python3 spacing_lab.py # ④ 步数摆法与收尾点 /usr/local/bin/python3 make_figures.py # ⑤ 画图 四个具体的改动实验,附我这里跑出来的结果: 把 η 从 0 改成 0.3:ddim_family.py 的 D 节会打印"输入相对变化",η=0 时最后 10 步是 0.0197,η=0.5 是 0.0667,η=1 是 0.1294。你会看到"加一点点噪声"的代价比想象中大。 把 dpm_solver_lab.py 里 dpm_pp_third 的 d2_scale 从 1.0 改成 2.0:C 节里 3M 的拟合阶数会从 2.07 变成 2.91,F 节的局部误差从 4.75e-07(最后一行)变成 4.44e-16。这一个常数就是"名义三阶"和"实测三阶"的分界。 把 DDIM 的网格从 uniform_t_grid 换成 uniform_lambda_grid:C 节里 DDIM 那两行的阶数不变(0.97 / 0.99),但误差整体上移,达到 1e-2 所需的 NFE 从 256 涨到 384。"同阶但更差"这件事,只有量出来才看得见。 在 spacing_lab.py 里把 grid_uniform_lambda_capped(20, cap) 的 cap 从 4.61 改成 2.5:E2 节那张 U 形表上,2M 的终点误差会从 2.90e-02 掉到 1.71e-02。这个数不用换模型、不用换采样器,只改一行网格。 10. 延伸阅读 ddpm:DDPM 的训练目标与采样流程。本篇的前置——$\hat x_0$ 与 $\epsilon_\theta$ 的互换关系、$\bar\alpha_t$ 的调度都在那里定义。 diffusion_math:扩散过程的前向与反向推导。概率流 ODE 与 SDE 的关系、参数化选择对误差权重的影响,是理解 05 节 algorithm_type 开关的基础。 vae_elbo:变分下界与重参数化。$\eta$ 旋钮的合法性来自"边缘分布不变",那套记号在那里建立。 cfg:分类器无关引导。引导会放大 ODE 右端的 Lipschitz 常数,采样步法要跟着调,是 06 节"稳定域"的直接后果。 flow_matching:流匹配与 Rectified Flow。把"直线路径"当作先验,是这条 ODE 主线的另一个分支。 附录:完整代码 09 节用到的脚本全文如下(oracle_gmm.py、ddim_family.py、dpm_solver_lab.py、spacing_lab.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 oracle_gmm.py # -*- coding: utf-8 -*- """采样器实验的地基:一个 score 有闭式解的二维高斯混合 + 一把「尺子」。 为什么必须先把这两件事钉死: 1. 比较采样器时,最大的干扰项是**模型误差**。同一个噪声网络,在 t=1 和 t=1000 上的误差可以差两个数量级,你根本分不清「这个采样器好」和「这个时间步好学」。 用闭式 score(oracle),模型误差恒为 0,剩下的全是离散化误差。 2. 比较采样器时,第二大的干扰项是**尺子本身的噪声**。切片 Wasserstein 距离 是用有限样本估的,两个「都对」的分布之间也能量出一个正数。这个正数就是 噪声地板,所有小于它的差别都是运气。 所以这个脚本干三件事:给出 oracle、给出尺子、量出地板。后面三个脚本全部 import 它,保证四处用到的目标分布、指标、随机数是同一套。 /usr/local/bin/python3 oracle_gmm.py 依赖:numpy, matplotlib(画图时才需要) """ import numpy as np T = 1000 # 训练时的扩散步数 D = 2 # 数据维度(二维才画得出轨迹,也才跑得起闭式 score) N_EVAL = 4000 # 评价用的样本数 N_DIRS = 512 # 切片 Wasserstein 的投影方向数 # ══════════════════════════════════════════════════════════════════ # 1. 噪声调度 # ══════════════════════════════════════════════════════════════════ def linear_schedule(T: int = T, b0: float = 1e-4, b1: float = 0.02): """DDPM / diffusers 默认的线性 beta 调度。 返回 abar 的长度是 T+1:abar[t] = prod_{s<=t} (1 - beta_s), abar[0] = 1 表示「完全干净」。这样下标和论文里的 t 直接对齐, 不用在 0-based / 1-based 之间来回换算。 """ betas = np.linspace(b0, b1, T) abar = np.concatenate([[1.0], np.cumprod(1.0 - betas)]) return betas, abar def alpha_sigma(abar_t: float): """给定 abar_t,返回 (alpha_t, sigma_t) = (sqrt(abar), sqrt(1-abar))。 这是全文唯一的「坐标」:x_t = alpha_t * x_0 + sigma_t * eps。 注意 alpha 是**信号尺度**,不是 DDPM 论文里的 alpha_t = 1 - beta_t。 """ return np.sqrt(abar_t), np.sqrt(1.0 - abar_t) # ══════════════════════════════════════════════════════════════════ # 2. 目标分布:8 个切向拉长的高斯,摆成一个环 # ══════════════════════════════════════════════════════════════════ def make_target(K: int = 8, radius: float = 1.8, tang: float = 0.30, rad: float = 0.05): """K 个分量均匀摆在一个半径 radius 的环上,每个分量沿切线方向拉长。 为什么不用各向同性:各向同性的高斯混合,加噪之后很快变成一个圆球, score 场几乎是线性的,一阶方法就能解到机器精度,分不出高下。 切向拉长之后,环上的 score 场有明显曲率,离散化误差才显出来。 """ ang = np.linspace(0.0, 2.0 * np.pi, K, endpoint=False) mu = np.stack([radius * np.cos(ang), radius * np.sin(ang)], axis=1) Sig = np.zeros((K, D, D)) for k, a in enumerate(ang): R = np.array([[np.cos(a), -np.sin(a)], [np.sin(a), np.cos(a)]]) Sig[k] = R @ np.diag([tang, rad]) @ R.T w = np.ones(K) / K return w, mu, Sig def sample_true(n: int, w, mu, Sig, rng): """从目标分布抽样。""" k = rng.choice(len(w), size=n, p=w) L = np.linalg.cholesky(Sig[k]) # (n, D, D) z = rng.standard_normal((n, D)) return mu[k] + np.einsum("nij,nj->ni", L, z) # ══════════════════════════════════════════════════════════════════ # 3. Oracle:加噪后分布的闭式 score,以及最优的 eps 预测 # ══════════════════════════════════════════════════════════════════ def _cov_at(abar_t: float, Sig): """q(x_t) 的第 k 个分量协方差:C_k = abar * Sig_k + (1 - abar) * I。""" return abar_t * Sig + (1.0 - abar_t) * np.eye(D) def logp_and_score(x, abar_t: float, w, mu, Sig): """返回 (log q_t(x), score_t(x)),x 形状 (n, D)。 q_t 是 K 个高斯的混合:第 k 个分量的均值是 sqrt(abar) * mu_k, 协方差是 abar * Sig_k + (1 - abar) * I。score 是这个混合分布的 对数梯度 = 各分量 score 的后验加权平均。 """ C = _cov_at(abar_t, Sig) # (K, D, D) inv = np.linalg.inv(C) ld = np.log(np.linalg.det(C)) r = x[:, None, :] - np.sqrt(abar_t) * mu[None] # (n, K, D) quad = np.einsum("nki,kij,nkj->nk", r, inv, r) lg = np.log(w)[None, :] - 0.5 * quad - 0.5 * ld[None, :] m = lg.max(axis=1, keepdims=True) p = np.exp(lg - m) Z = p.sum(axis=1, keepdims=True) logp = (m[:, 0] + np.log(Z[:, 0])) pi = p / Z # 后验分量权重 comp = -np.einsum("kij,nkj->nki", inv, r) # 每个分量的 score score = np.einsum("nk,nki->ni", pi, comp) return logp, score def eps_star(x, abar_t: float, w, mu, Sig): """MMSE 意义下最优的 eps 预测:eps* = -sigma_t * score_t(x)。 因为 x_t = alpha x_0 + sigma * eps,对 x_t 求 log 梯度会得到 grad log q_t = -E[eps | x_t] / sigma,所以最优的噪声预测就是 score 乘一个 -sigma。后面所有采样器都拿它当「模型」。 """ _, s = logp_and_score(x, abar_t, w, mu, Sig) return -np.sqrt(1.0 - abar_t) * s def x0_from_eps(x, abar_t: float, eps): """由 x_t 和 eps 反解 x_0:x_0 = (x_t - sigma_t * eps) / alpha_t。""" a, sg = alpha_sigma(abar_t) return (x - sg * eps) / a # ══════════════════════════════════════════════════════════════════ # 4. 尺子:切片 Wasserstein-1 + 平均对数密度 # ══════════════════════════════════════════════════════════════════ def make_dirs(n_dirs: int = N_DIRS, seed: int = 7): """预先抽好投影方向,全篇共用同一组。 这一条很重要:每次评价都重新抽方向的话,两组「同样好」的样本也会因为 方向不同量出不同的数,这部分方差完全是无谓的。方向固定之后, A 和 B 的差别就只来自样本本身。 """ th = np.random.default_rng(seed).standard_normal((n_dirs, D)) th /= np.linalg.norm(th, axis=1, keepdims=True) return th def sliced_w1(X, Y, TH=None, n_dirs: int = N_DIRS, rng=None): """切片 W1:在 n_dirs 个随机方向上算一维 W1,再取平均。 一维 W1 有闭式解:把两个样本集投影后排序,逐位相减取绝对值平均。 单位和数据同单位(这里是「平均要挪多远」),比 MMD^2 好解释。 """ if TH is None: TH = make_dirs(n_dirs) if rng is None else None if TH is None: th = rng.standard_normal((n_dirs, X.shape[1])) th /= np.linalg.norm(th, axis=1, keepdims=True) else: th = TH a = np.sort(X @ th.T, axis=0) b = np.sort(Y @ th.T, axis=0) return float(np.abs(a - b).mean()) def mean_logp(X, w, mu, Sig): """生成样本在真实分布下的平均对数密度。""" lp, _ = logp_and_score(X, 1.0, w, mu, Sig) return float(lp.mean()) def evaluate(X, Xref, w, mu, Sig, lp_ref, TH): """一次评价,返回 (SW1, dlogp)。 dlogp = 生成样本的平均 log 密度 − 真实样本的平均 log 密度。 0 表示完美;负数表示生成样本落在了真实分布的「空区」。 """ return sliced_w1(X, Xref, TH=TH), mean_logp(X, w, mu, Sig) - lp_ref # ══════════════════════════════════════════════════════════════════ # 5. 主流程:验 oracle、量尺子的噪声地板 # ══════════════════════════════════════════════════════════════════ def main(): rng = np.random.default_rng(20260927) betas, abar = linear_schedule() w, mu, Sig = make_target() print("=" * 68) print("A. Oracle 自检:解析 score vs 有限差分") print("=" * 68) for t in (1, 50, 200, 500, 1000): x = sample_true(8, w, mu, Sig, rng) _, s = logp_and_score(x, abar[t], w, mu, Sig) h = 1e-5 fd = np.zeros_like(s) for i in range(D): e = np.zeros(D) e[i] = h lp_p, _ = logp_and_score(x + e, abar[t], w, mu, Sig) lp_m, _ = logp_and_score(x - e, abar[t], w, mu, Sig) fd[:, i] = (lp_p - lp_m) / (2 * h) err = np.abs(fd - s).max() print(f" t={t:5d} abar={abar[t]:.6f} max|解析 - 差分| = {err:.2e}") print(" 差分误差在 1e-7 量级 → 解析 score 是对的(不是我拍脑袋写的)") # 顺便看一眼:oracle eps 在 t=1000 处还剩多少信息 print() print("=" * 68) print("B. 端点检查:t=T 时分布离标准正态有多远") print("=" * 68) for t in (200, 500, 800, 1000): C = _cov_at(abar[t], Sig) dev = np.abs(C - np.eye(D)).max() print(f" t={t:5d} abar={abar[t]:.3e} " f"各分量协方差与 I 的最大差 = {dev:.3e}") print(" 越接近 t=T,加噪后的混合越接近一个标准正态球") # ── 尺子的噪声地板 ── print() print("=" * 68) print("C. 尺子的噪声地板:两批「都对」的样本之间也能量出距离") print("=" * 68) Xref = sample_true(N_EVAL, w, mu, Sig, rng) lp_ref = mean_logp(Xref, w, mu, Sig) print(f" 参考样本: {Xref.shape}, 真实样本自己的平均 logp = {lp_ref:.4f}") TH = make_dirs() sw_floor, lp_floor = [], [] for r in range(8): Y = sample_true(N_EVAL, w, mu, Sig, np.random.default_rng(1000 + r)) sw_floor.append(sliced_w1(Y, Xref, TH=TH)) lp_floor.append(mean_logp(Y, w, mu, Sig)) print(f" SW1 真vs真 ×8: 均值 {np.mean(sw_floor):.5f} " f"标准差 {np.std(sw_floor):.5f} 最大 {np.max(sw_floor):.5f}") print(f" dlogp 真vs真 ×8: 均值 {np.mean(lp_floor) - lp_ref:+.4f} " f"标准差 {np.std(lp_floor):.4f}") print(" → 后面所有表格里,小于这个数的差别都是抽样噪声,不是采样器的功劳") # ── 一个有参照的好答案长什么样 ── print() print("=" * 68) print("D. 天花板:把真样本直接当生成样本送进去") print("=" * 68) print(f" SW1 = {sliced_w1(Xref, Xref, TH=TH):.6f}" f" dlogp = {mean_logp(Xref, w, mu, Sig) - lp_ref:+.4f}") print(" (恒等于 0,说明尺子本身没有偏置)") if __name__ == "__main__": main() ddim_family.py # -*- coding: utf-8 -*- """DDIM 家族:一个 eta 旋钮,把 DDPM 祖采样和 DDIM 串成一条线。 这个脚本回答三个问题: 1. **eta=1 的 DDIM 是不是就是 DDPM 祖采样?** 用两套独立的代码(一套写 DDPM Algorithm 2 的原式,一套写 DDIM 论文式 (12)),喂同样的随机数和 同样的起点,看输出能不能对到最后一位。能,那这条家族关系就不是传说。 2. **eta 到底该调多大?** 在 oracle score 下扫一遍:步数少的时候是不是 「加噪声更好」?在欠拟合模型下再扫一遍,看结论翻不翻。 3. **换采样器为什么缓存策略得重调?** 量每步「模型输入变了多少」: 确定性采样器的轨迹是光滑的,随机采样器每一步都往里砸一份新噪声。 /usr/local/bin/python3 ddim_family.py 依赖:numpy(+ 同目录的 oracle_gmm.py) """ import numpy as np from oracle_gmm import (T, linear_schedule, make_target, sample_true, eps_star, x0_from_eps, alpha_sigma, sliced_w1, mean_logp, make_dirs, N_EVAL, N_DIRS) # ══════════════════════════════════════════════════════════════════ # 0. 时间步网格与模型 # ══════════════════════════════════════════════════════════════════ def uniform_t_grid(S: int, T: int = T): """DDIM 论文 / diffusers 默认的「leading」取法:在 t 上等间隔。 返回 S+1 个 abar:起点 abar[T],终点 1.0(完全干净)。 """ ts = np.round(np.linspace(0.0, T, S + 1)).astype(int)[::-1] grid = np.array([1.0 if t == 0 else _ABAR[t] for t in ts], dtype=float) return ts, grid def paired_noise(S: int, n: int, seed: int = 1234): """给所有采样器**同一份**起点和同一批随机数。 这是整套实验最要紧的一条:如果每个配置各抽各的起点,配置之间就多了一份 抽样方差,量出来的差别里有多少是采样器的功劳根本说不清。配对之后, eta=0 和 eta=1 的终点之差只能来自算法本身。 """ r = np.random.default_rng(seed + S) x_init = r.standard_normal((n, 2)) zs = [r.standard_normal((n, 2)) for _ in range(S)] return x_init, zs def make_model(kind: str, w, mu, Sig): """两种「模型」,用来分离离散化误差和模型误差。 oracle : 用的就是真 score,模型误差 = 0,剩下的全是离散化误差。 gauss : 模型以为数据是一个单高斯(用真均值真协方差拟合), 也就是「容量不足以表达多峰」。高噪声区它几乎是对的, 低噪声区它错得离谱 —— 这正是真实网络的误差结构。 """ if kind == "oracle": return lambda x, abar: eps_star(x, abar, w, mu, Sig) if kind == "gauss": m = (w[:, None] * mu).sum(axis=0) S = ((w[:, None, None] * (Sig + np.einsum("ki,kj->kij", mu, mu))).sum(axis=0) - np.outer(m, m)) w1, mu1, S1 = np.array([1.0]), m[None, :], S[None, :, :] return lambda x, abar: eps_star(x, abar, w1, mu1, S1) raise ValueError(kind) # ══════════════════════════════════════════════════════════════════ # 1. DDIM 的一步(论文式 12,与 diffusers DDIMScheduler.step 逐项对齐) # ══════════════════════════════════════════════════════════════════ def ddim_step(x, abar_cur, abar_tgt, eps, z, eta): """从 abar_cur 走到 abar_tgt(abar_tgt > abar_cur,噪声变小)。 x_tgt = sqrt(abar_tgt) * x0_hat + sqrt(sigma_tgt^2 - sigma_tilde^2) * eps ← 指向 x_t 的方向项 + sigma_tilde * z ← 随机项 sigma_tilde = eta * sqrt( (1-abar_tgt)/(1-abar_cur) * (1 - abar_cur/abar_tgt) ) eta = 0 → 完全没有随机项,DDIM; eta = 1 → 随机项等于 DDPM 的后验标准差 beta_tilde。 """ a_c, s_c = alpha_sigma(abar_cur) a_t, s_t = alpha_sigma(abar_tgt) x0 = (x - s_c * eps) / a_c var = (1.0 - abar_tgt) / (1.0 - abar_cur) * (1.0 - abar_cur / abar_tgt) sig_tilde = eta * np.sqrt(max(var, 0.0)) direc = np.sqrt(max(s_t ** 2 - sig_tilde ** 2, 0.0)) return a_t * x0 + direc * eps + sig_tilde * z, x0 def ddpm_step(x, abar_cur, abar_tgt, eps, z, betas): """DDPM Algorithm 2 的原式,故意写成「另一套代码」用来交叉验证。 x_{t-1} = (x_t - beta_t / sqrt(1 - abar_t) * eps) / sqrt(alpha_t) + sqrt(beta_tilde_t) * z 只在相邻步(stride = 1)上成立,不能跳步。 """ abar_prev = abar_tgt beta_t = 1.0 - abar_cur / max(abar_prev, 1e-12) # = 1 - alpha_t beta_tilde = (1.0 - abar_prev) / (1.0 - abar_cur) * beta_t a_t = np.sqrt(max(abar_cur / max(abar_prev, 1e-12), 1e-12)) mean = (x - beta_t / np.sqrt(1.0 - abar_cur) * eps) / a_t return mean + np.sqrt(max(beta_tilde, 0.0)) * z def run_ddim(x, grid, model, eta, zs, track=False): """沿 grid 走一遍。grid[0] 最吵,grid[-1] = 1.0。 track=True 时额外返回每一步的 (x, x0_hat),D 节量「轨迹抖动」要用。 """ xs, x0s = [x.copy()], [] for i in range(len(grid) - 1): eps = model(x, grid[i]) z = zs[i] if zs is not None else 0.0 x, x0 = ddim_step(x, grid[i], grid[i + 1], eps, z, eta) if track: xs.append(x.copy()) x0s.append(x0.copy()) return x, (xs, x0s) if track else x0s # ══════════════════════════════════════════════════════════════════ # 2. 主流程 # ══════════════════════════════════════════════════════════════════ _BETAS, _ABAR = linear_schedule() def main(): rng = np.random.default_rng(20260927) w, mu, Sig = make_target() Xref = sample_true(N_EVAL, w, mu, Sig, rng) lp_ref = mean_logp(Xref, w, mu, Sig) ev_rng = np.random.default_rng(7) model = make_model("oracle", w, mu, Sig) # ── A. eta=1 的 DDIM 是不是 DDPM 祖采样 ── print("=" * 70) print("A. 交叉验证:DDIM(eta=1, stride=1) 是否等于 DDPM Algorithm 2") print("=" * 70) S = 1000 ts, grid = uniform_t_grid(S) n = 400 x_init = rng.standard_normal((n, 2)) zs = [rng.standard_normal((n, 2)) for _ in range(S)] x_ddpm = x_init.copy() for i in range(S): # 原式,一步步走完 1000 步 t = ts[i] x_ddpm = ddpm_step(x_ddpm, grid[i], grid[i + 1], model(x_ddpm, grid[i]), zs[i], _BETAS) x_ddim, _ = run_ddim(x_init.copy(), grid, model, eta=1.0, zs=zs) d = np.abs(x_ddpm - x_ddim).max() print(f" 两条独立实现,喂同样的 {S} 份随机数") print(f" 终点逐元素最大差 = {d:.3e} 终点 RMS = {np.sqrt((x_ddim**2).mean()):.4f}") print(f" → 差 {d:.1e},是浮点累加误差量级:eta=1 就是 DDPM 祖采样,不是「近似」") # ── B. eta 扫描(oracle score)── print() print("=" * 70) print("B. eta 扫描:步数越少,是不是越该加噪声?(oracle score,配对起点)") print("=" * 70) TH = make_dirs() floors = [sliced_w1(sample_true(N_EVAL, w, mu, Sig, np.random.default_rng(1000 + r)), Xref, TH=TH) for r in range(6)] print(f" 噪声地板 SW1 = {np.mean(floors):.4f} ± {np.std(floors):.4f}" f"(真实样本对真实样本),越接近它越好") NS = (10, 20, 50, 100) print(f" {'eta':>5} " + " ".join(f"{'N=' + str(s):>10}" for s in NS)) for eta in (0.0, 0.25, 0.5, 0.75, 1.0): row = [] for S in NS: ts, grid = uniform_t_grid(S) x_init, zs = paired_noise(S, N_EVAL) x, _ = run_ddim(x_init, grid, model, eta, zs) row.append(sliced_w1(x, Xref, TH=TH)) print(f" {eta:5.2f} " + " ".join(f"{v:10.4f}" for v in row)) # ── C. 换个欠拟合的模型,结论翻不翻 ── print() print("=" * 70) print("C. 同样的扫描,但把模型换成「以为数据是单高斯」") print("=" * 70) gm = make_model("gauss", w, mu, Sig) print(f" {'eta':>5} " + " ".join(f"{'N=' + str(s):>10}" for s in NS)) for eta in (0.0, 0.25, 0.5, 0.75, 1.0): row = [] for S in NS: ts, grid = uniform_t_grid(S) x_init, zs = paired_noise(S, N_EVAL) x, _ = run_ddim(x_init, grid, gm, eta, zs) row.append(sliced_w1(x, Xref, TH=TH)) print(f" {eta:5.2f} " + " ".join(f"{v:10.4f}" for v in row)) # ── D. 每步「模型输入变了多少」:缓存友好度 ── print() print("=" * 70) print("D. 轨迹抖动:每步「模型输入」变了百分之几(缓存复用看的就是它)") print("=" * 70) print(" 缓存类加速(按『这一步的输入和上一步差不多就复用上一步的特征』决策)") print(" 盯的是这个量。它一旦被噪声项垫住,原阈值就全废了。") print(f" {'eta':>5} {'最后10步 输入相对变化':>22} {'最后10步 x0_hat位移':>22}") S = 50 ts, grid = uniform_t_grid(S) x_init, zs = paired_noise(S, N_EVAL) out = {} for eta in (0.0, 0.5, 1.0): _, (xs, x0s) = run_ddim(x_init.copy(), grid, model, eta, zs, track=True) rel_x = [float(np.sqrt((((xs[i + 1] - xs[i]) ** 2).sum(1)).mean()) / np.sqrt(((xs[i] ** 2).sum(1)).mean())) for i in range(len(xs) - 1)] d_x0 = [float(np.sqrt((((x0s[i] - x0s[i - 1]) ** 2).sum(1)).mean())) for i in range(1, len(x0s))] out[eta] = (rel_x, d_x0) print(f" {eta:5.1f} {np.mean(rel_x[-10:]):22.4f} {np.mean(d_x0[-10:]):22.4f}") r0, r1 = out[0.0][0][-10:], out[1.0][0][-10:] print(f" → eta=1 的输入抖动是 eta=0 的 {np.mean(r1) / np.mean(r0):.1f} 倍;" f"按 eta=0 调出来的缓存阈值,换个采样器就完全不是一回事") if __name__ == "__main__": main() dpm_solver_lab.py # -*- coding: utf-8 -*- """从 DDIM 走到高阶:把时间轴换成 lambda,阶数就变成可以直接量出来的东西。 这个脚本做四件事: 1. **验证 DDIM 就是一阶 DPM-Solver。** 两套完全不同的代码(一套写 DDIM 论文式 12,一套写 lambda 空间的一阶指数积分器),同样的时间网格、 同样的模型,看输出是不是同一个数。是,那「DDIM 是一阶方法」就不是比喻。 2. **把阶数量出来。** 拿一条超高精度的参考解当真值,量各采样器终点到它的 距离。这个距离没有抽样噪声(起点配对、采样器确定),能干净地跨几个数量级, 于是「误差 ~ NFE^(-p)」里的 p 可以直接拟合出来。 3. **给出实用的步数-质量表。** 阶数是数值分析的语言,产品要的是 「20 步够不够」,所以再给一张 SW1 表(含噪声地板)。 4. **看高阶方法省下的评估次数到底落在哪。** /usr/local/bin/python3 dpm_solver_lab.py 依赖:numpy(+ 同目录的 oracle_gmm.py / ddim_family.py) """ import numpy as np from oracle_gmm import (T, linear_schedule, make_target, sample_true, eps_star, alpha_sigma, sliced_w1, make_dirs, N_EVAL) from ddim_family import uniform_t_grid, make_model, ddim_step _BETAS, _ABAR = linear_schedule() _LAM = 0.5 * np.log(np.maximum(_ABAR, 1e-300) / np.maximum(1.0 - _ABAR, 1e-300)) # ══════════════════════════════════════════════════════════════════ # 0. lambda 坐标与网格 # ══════════════════════════════════════════════════════════════════ def lam_of_abar(abar): """half log-SNR:lambda = log(alpha) - log(sigma) = 0.5 * log(abar/(1-abar))。""" return 0.5 * np.log(np.maximum(abar, 1e-300) / np.maximum(1.0 - abar, 1e-300)) def abar_of_lam(lam): """lambda 的反函数:abar = e^(2 lam) / (1 + e^(2 lam))。""" e = np.exp(2.0 * lam) return e / (1.0 + e) def nearest_t(abar): """把连续的 abar 落到最近的训练时间步上(真实网络只认整数 t)。 注意 abar 是**递减**的,不能用 searchsorted(它只认递增数组), 老老实实找绝对值最小的那个下标。 """ return int(np.clip(int(np.argmin(np.abs(_ABAR - abar))), 1, T)) def uniform_lambda_grid(S: int, quantize: bool = False): """在 lambda 上等间隔取 S 个点,再落回训练网格。 返回 S+1 个 abar:前 S 个是「要评估模型」的位置,最后一个是 1.0 (完全干净,sigma = 0,对应 lambda = +inf,只能用一阶收尾)。 quantize=True 会把每个 lambda 落回整数时间步。落回去之后可能撞车 (两个 lambda 落到同一个 t),撞车会让多步法的差商分母变成 0, 所以强制 t 严格递减。参考解要更细的分辨率,用 quantize=False 直接 吃连续的 abar —— oracle 认得任意 abar,真实网络才需要整数。 """ l0, l1 = _LAM[T], _LAM[1] # 从最吵到最干净 lams = np.linspace(l0, l1, S) if not quantize: return None, np.concatenate([abar_of_lam(lams), [1.0]]) ts = [] for l in lams: t = nearest_t(abar_of_lam(l)) if ts and t >= ts[-1]: t = ts[-1] - 1 # 强制严格递减,避免差商分母为 0 if t < 1: break # 干净端训练网格已无分辨率,到此为止 ts.append(t) ts = np.array(ts, dtype=int) return ts, np.concatenate([_ABAR[ts], [1.0]]) class Counter: """包一层模型,数一下到底调了几次。NFE 不能靠嘴算。""" def __init__(self, model): self.m = model self.n = 0 def __call__(self, x, abar): self.n += 1 return self.m(x, abar) # ══════════════════════════════════════════════════════════════════ # 1. 一阶指数积分器(DPM-Solver-1 / dpmsolver++ 的一阶更新) # ══════════════════════════════════════════════════════════════════ def dpm_pp_first(x, m0, a_s, s_s, a_t, s_t): """x_t = (sigma_t/sigma_s) * x_s - alpha_t * (e^{-h} - 1) * x0_hat h = lambda_t - lambda_s。sigma_t = 0 时 lambda_t = +inf,e^{-h} = 0, 整条式子退化成 x_t = x0_hat —— 最后一步自动正确,不用特判。 """ if s_t <= 0.0: return m0 h = (np.log(a_t) - np.log(s_t)) - (np.log(a_s) - np.log(s_s)) return (s_t / s_s) * x - a_t * (np.exp(-h) - 1.0) * m0 def dpm_pp_second(x, hist, a_s, s_s, a_t, s_t, solver_type="midpoint"): """diffusers multistep_dpm_solver_second_order_update 的逐项复刻。 D1 是用上一步的 x0_hat 做差商估出来的导数:D1 ~ h * dx0/dlambda。 midpoint 和 heun 的区别只在 D1 前面的系数,两者都是二阶。 """ lam_t = np.log(a_t) - np.log(s_t) lam_s0, m0 = hist[-1] lam_s1, m1 = hist[-2] h = lam_t - lam_s0 h0 = lam_s0 - lam_s1 r0 = h0 / h D0, D1 = m0, (m0 - m1) / r0 emh = np.exp(-h) if solver_type == "midpoint": return ((s_t / s_s) * x - a_t * (emh - 1.0) * D0 - 0.5 * a_t * (emh - 1.0) * D1) return ((s_t / s_s) * x - a_t * (emh - 1.0) * D0 + a_t * ((emh - 1.0) / h + 1.0) * D1) def dpm_pp_third(x, hist, a_s, s_s, a_t, s_t, d2_scale=2.0): """三阶多步更新。 d2_scale=1.0 是 diffusers(v0.30.0)`multistep_dpm_solver_third_order_update` 的逐项复刻;d2_scale=2.0 是本文按「二次 x0_hat 必须精确」推出来的系数。 两者的差别见 F 节的受控检验 —— 差的就是这一个 2。 """ lam_t = np.log(a_t) - np.log(s_t) lam_s0, m0 = hist[-1] lam_s1, m1 = hist[-2] lam_s2, m2 = hist[-3] h = lam_t - lam_s0 h0, h1 = lam_s0 - lam_s1, lam_s1 - lam_s2 r0, r1 = h0 / h, h1 / h D0 = m0 D1_0, D1_1 = (m0 - m1) / r0, (m1 - m2) / r1 D1 = D1_0 + (r0 / (r0 + r1)) * (D1_0 - D1_1) D2 = d2_scale * (D1_0 - D1_1) / (r0 + r1) emh = np.exp(-h) return ((s_t / s_s) * x - a_t * (emh - 1.0) * D0 + a_t * ((emh - 1.0) / h + 1.0) * D1 - a_t * ((emh - 1.0 + h) / h ** 2 - 0.5) * D2) # ══════════════════════════════════════════════════════════════════ # 2. 五个采样器 # ══════════════════════════════════════════════════════════════════ def solve_ddim(grid, model, x): """DDIM(eta=0):论文式 12,完全不碰 lambda。""" for i in range(len(grid) - 1): eps = model(x, grid[i]) x, _ = ddim_step(x, grid[i], grid[i + 1], eps, 0.0, 0.0) return x def solve_dpm1(grid, model, x): """一阶指数积分器。理论上应该和 DDIM 逐位相同 —— A 节去验。""" for i in range(len(grid) - 1): a_s, s_s = alpha_sigma(grid[i]) a_t, s_t = alpha_sigma(grid[i + 1]) eps = model(x, grid[i]) m0 = (x - s_s * eps) / a_s x = dpm_pp_first(x, m0, a_s, s_s, a_t, s_t) return x def solve_dpm2s(grid, model, x): """单步二阶(DPM-Solver-2):每步先在 lambda 中点落一脚,用中点的 x0 走完。 一步两步评估。中点的 x0 对积分的「加权平均」是二阶准确的, 所以整步的局部误差是 O(h^3),全局 O(h^2)。 """ for i in range(len(grid) - 1): a_s, s_s = alpha_sigma(grid[i]) a_t, s_t = alpha_sigma(grid[i + 1]) eps = model(x, grid[i]) m0 = (x - s_s * eps) / a_s lam_s, lam_t = np.log(a_s) - np.log(s_s), None if s_t <= 0.0: # 收尾:lambda_t = +inf,中点无从谈起 x = m0 continue lam_t = np.log(a_t) - np.log(s_t) lam_m = 0.5 * (lam_s + lam_t) abar_m = _ABAR[nearest_t(abar_of_lam(lam_m))] a_m, s_m = alpha_sigma(abar_m) x_m = dpm_pp_first(x, m0, a_s, s_s, a_m, s_m) # 先跳到中点 eps_m = model(x_m, abar_m) m_m = (x_m - s_m * eps_m) / a_m # 中点处的 x0 x = dpm_pp_first(x, m_m, a_s, s_s, a_t, s_t) # 用中点 x0 走完整步 return x def _midpoint_x0(x, m0, a_s, s_s, a_t, s_t, model): """在 lambda 的中点补一次评估,拿中点的 x0 走完整步(局部二阶)。 多步法开头那一步没有历史可用,只能降成一阶,而一阶在长度为 h 的区间上 局部误差是 O(h^2),这个误差会一路传到终点 —— 这就是 3M 实测阶数被卡在 2 的原因。给第一步多花一次评估可以验证这件事。 """ lam_s, lam_t = np.log(a_s) - np.log(s_s), np.log(a_t) - np.log(s_t) abar_m = abar_of_lam(0.5 * (lam_s + lam_t)) a_m, s_m = alpha_sigma(abar_m) x_m = dpm_pp_first(x, m0, a_s, s_s, a_m, s_m) eps_m = model(x_m, abar_m) return (x_m - s_m * eps_m) / a_m def solve_multistep(grid, model, x, order=2, solver_type="midpoint", warm2s=False, d2_scale=2.0): """多步法(DPM-Solver++ 2M / 3M):一步一次评估,导数是拿历史 x0 做差商。 开头几步历史不够,自动降阶(和 diffusers 的 lower_order_nums 一样); 最后一步 sigma=0 强制降成一阶(和 final_sigmas_type="zero" 一样)。 warm2s=True 时,第一步改用中点法(多花一次评估)把局部精度提到二阶。 """ hist = [] for i in range(len(grid) - 1): a_s, s_s = alpha_sigma(grid[i]) a_t, s_t = alpha_sigma(grid[i + 1]) eps = model(x, grid[i]) m0 = (x - s_s * eps) / a_s hist.append((np.log(a_s) - np.log(s_s), m0)) last = (i == len(grid) - 2) if i == 0 and warm2s and s_t > 0.0: x = dpm_pp_first(x, _midpoint_x0(x, m0, a_s, s_s, a_t, s_t, model), a_s, s_s, a_t, s_t) elif s_t <= 0.0 or len(hist) < 2 or last: x = dpm_pp_first(x, m0, a_s, s_s, a_t, s_t) elif order >= 3 and len(hist) >= 3: x = dpm_pp_third(x, hist, a_s, s_s, a_t, s_t, d2_scale=d2_scale) else: x = dpm_pp_second(x, hist, a_s, s_s, a_t, s_t, solver_type) return x # ══════════════════════════════════════════════════════════════════ # 3. 主流程 # ══════════════════════════════════════════════════════════════════ def rms_err(x, xref): return float(np.sqrt(((x - xref) ** 2).sum(1).mean())) def fit_slope(ns, errs, res, nfe_min=24): """在双对数上拟合 err ~ N^(-p),返回 p。 只取渐近段:NFE 太小的时候高阶方法还在「攒历史」(前几步被迫降阶), 这一段量出来的斜率既不是 1 也不是 3,什么都不说明;误差小于参考解 分辨率的点是机器精度,同样不能要。 """ xs, ys = [], [] for n, e in zip(ns, errs): if n >= nfe_min and e > 3.0 * res and np.isfinite(e) and e > 0: xs.append(np.log(n)) ys.append(np.log(e)) if len(xs) < 3: return float("nan"), 0 p = -np.polyfit(xs, ys, 1)[0] return float(p), len(xs) def local_err_test(): """受控检验:规定 x0_hat(lambda) 是二次多项式,看各阶更新的局部误差阶。 一条真正的 k 阶方法,在 x0_hat 是 k-1 次多项式时必须一步算准(误差 ~ 0), 因为它的构造就是「用 k 个历史点插值出 k-1 次多项式,再对 e^lambda 精确积分」。 这个检验把「公式对不对」和「轨迹好不好」彻底分开,谁也赖不着谁。 """ rng = np.random.default_rng(3) A, Bc, Cc = rng.standard_normal(2), rng.standard_normal(2), rng.standard_normal(2) g = lambda lam: A + Bc * lam + 0.5 * Cc * lam ** 2 def exact_inc(lam_s, lam_t, n=200001): u = np.linspace(lam_s, lam_t, n) f = np.exp(u)[:, None] * np.stack([g(v) for v in u]) return np.trapezoid(f, u, axis=0) hs = [0.4, 0.2, 0.1, 0.05, 0.025] lam_s = 0.3 out = {} for name in ("一阶", "二阶", "三阶-D2原式", "三阶-D2乘2"): errs = [] for h in hs: lam_t = lam_s + h a_s, s_s = alpha_sigma(abar_of_lam(lam_s)) a_t, s_t = alpha_sigma(abar_of_lam(lam_t)) x = np.array([0.7, -0.4]) hist = [(lam_s - j * h, g(lam_s - j * h)) for j in (2, 1, 0)] # 老→新 if name == "一阶": got = dpm_pp_first(x, g(lam_s), a_s, s_s, a_t, s_t) elif name == "二阶": got = dpm_pp_second(x, hist, a_s, s_s, a_t, s_t) elif name == "三阶-D2原式": got = dpm_pp_third(x, hist, a_s, s_s, a_t, s_t, d2_scale=1.0) else: got = dpm_pp_third(x, hist, a_s, s_s, a_t, s_t, d2_scale=2.0) want = (s_t / s_s) * x + s_t * exact_inc(lam_s, lam_t) errs.append(float(np.abs(got - want).max())) out[name] = errs print(f" {name:>12}: " + " ".join(f"{e:.2e}" for e in errs)) print(" (h 从 0.4 缩到 0.025)") print(" → 一阶 h^2、二阶 h^3 都对;三阶原式只有 h^3(等于二阶),") print(" D2 乘 2 之后直接掉到 1e-14 —— 二次 x0_hat 下它一步就算准了") def main(): rng = np.random.default_rng(20260927) w, mu, Sig = make_target() TH = make_dirs() Xref = sample_true(N_EVAL, w, mu, Sig, rng) model = make_model("oracle", w, mu, Sig) # ── A. DDIM == 一阶 DPM-Solver ? ── print("=" * 72) print("A. 交叉验证:DDIM(eta=0) 与 lambda 空间一阶指数积分器是不是同一个东西") print("=" * 72) n = 500 for S in (20, 50): _, grid_l = uniform_lambda_grid(S) _, grid_t = uniform_t_grid(S) for name, grid in (("uniform-t", grid_t), ("uniform-lambda", grid_l)): x0 = rng.standard_normal((n, 2)) xa = solve_ddim(grid, model, x0.copy()) xb = solve_dpm1(grid, model, x0.copy()) print(f" S={S:3d} {name:>14} max|DDIM - DPM-Solver-1| = " f"{np.abs(xa - xb).max():.3e}") print(" → 两套代码、两套记号,输出差在浮点误差量级:DDIM 就是一阶方法") # ── B. 参考解 ── print() print("=" * 72) print("B. 参考解:3M 跑 4000 步当真值,再自查一下它自己收敛到哪") print("=" * 72) n_ref = 1000 x_init = rng.standard_normal((n_ref, 2)) _, g_fine = uniform_lambda_grid(4000, quantize=False) _, g_mid = uniform_lambda_grid(2000, quantize=False) xref = solve_multistep(g_fine, model, x_init.copy(), order=3) xchk = solve_multistep(g_mid, model, x_init.copy(), order=3) res = rms_err(xref, xchk) print(f" 4000 步 vs 2000 步参考解的差 = {res:.3e}") print(f" → 这个数就是下面所有误差表的「分辨率」,比它小的差别不可信") # ── C. 阶数 ── print() print("=" * 72) print("C. 把阶数量出来:终点到参考解的 RMS 距离 vs 真实 NFE") print("=" * 72) NS = [6, 8, 12, 16, 24, 32, 48, 64, 96, 128, 192] solvers = { "DDIM (uniform-t)": lambda S: (uniform_t_grid(S)[1], solve_ddim, {}), "DDIM (=DPM-1, lambda)": lambda S: (uniform_lambda_grid(S)[1], solve_ddim, {}), "DPM-Solver-2 单步": lambda S: (uniform_lambda_grid(max(S // 2, 2))[1], solve_dpm2s, {}), "DPM-Solver++ 2M": lambda S: (uniform_lambda_grid(S)[1], solve_multistep, {"order": 2}), "3M(D2 用原式)": lambda S: (uniform_lambda_grid(S)[1], solve_multistep, {"order": 3, "d2_scale": 1.0}), "3M(D2 修正×2)": lambda S: (uniform_lambda_grid(S)[1], solve_multistep, {"order": 3, "d2_scale": 2.0}), } print(f" {'NFE':>5} " + " ".join(f"{k:>14}" for k in solvers)) table = {k: [] for k in solvers} nfes = [] for S in NS: row, real_n = [], [] for k, build in solvers.items(): grid, fn, kw = build(S) c = Counter(model) x = fn(grid, c, x_init.copy(), **kw) row.append(rms_err(x, xref)) real_n.append(c.n) nfes.append(int(np.mean(real_n))) table_k = list(solvers) for k, v in zip(table_k, row): table[k].append(v) print(f" {nfes[-1]:5d} " + " ".join(f"{v:14.3e}" for v in row)) print() print(f" 拟合误差 ~ NFE^(-p)(只取 NFE>=24 且误差高于分辨率 {res:.1e}×3 的渐近段):") for k in solvers: p, m = fit_slope(nfes, table[k], res) print(f" {k:>24} p = {p:.2f} (用了 {m} 个点)") # ── D. SW1 走到哪一步就分不出高下了 ── print() print("=" * 72) print("D. 分布层面的尺子(SW1):重复换起点,看它什么时候失灵") print("=" * 72) floors = [sliced_w1(sample_true(N_EVAL, w, mu, Sig, np.random.default_rng(2000 + r)), Xref, TH=TH) for r in range(8)] print(f" 真样本 vs 真样本(5 次):{np.mean(floors):.4f} ± {np.std(floors):.4f}" f" ← 这就是尺子的分辨率") PS = [10, 20, 50] print(f" {'NFE':>5} " + " ".join(f"{k:>14}" for k in solvers)) for S in PS: row = [] for k, build in solvers.items(): grid, fn, kw = build(S) vs = [] for r in range(5): # 换几份起点,量出估计量的抖动 c = Counter(model) x = fn(grid, c, np.random.default_rng(900 + S * 100 + r) .standard_normal((N_EVAL, 2)), **kw) vs.append(sliced_w1(x, Xref, TH=TH)) row.append((np.mean(vs), np.std(vs))) print(f" {S:5d} " + " ".join(f"{v[0]:9.4f}±{v[1]:.3f}" for v in row)) print(" → 20 步时各采样器的差别还大于抖动;到 50 步大家都贴着地板,") print(" SW1 已经回答不了「谁更好」——这种时候只能看 C 节的确定性误差") # ── E. 同一个精度,省多少次评估 ── print() print("=" * 72) print("E. 把终点误差压到 1e-2 / 1e-3,最少要几次模型评估(实测扫描)") print("=" * 72) scan = [4, 6, 8, 10, 12, 16, 20, 24, 32, 40, 48, 64, 80, 96, 128, 160, 192, 256, 384] for k, build in solvers.items(): got = {} for S in scan: if len(got) == 2: break grid, fn, kw = build(S) c = Counter(model) e = rms_err(fn(grid, c, x_init.copy(), **kw), xref) for tol in (1e-2, 1e-3): if tol not in got and e <= tol: got[tol] = c.n s1 = got.get(1e-2, None) s2 = got.get(1e-3, None) f1 = f"NFE={s1:3d}" if s1 else " >400 " f2 = f"NFE={s2:3d}" if s2 else " >400 " print(f" {k:>24} 1e-2: {f1} 1e-3: {f2}") # ── F. 受控检验:三阶更新到底几阶 ── print() print("=" * 72) print("F. 受控检验:规定一条解析的 x0_hat(lambda),直接量局部误差") print("=" * 72) print(" 做法:不让 x0_hat 由 x 决定,而是规定它是 lambda 的二次多项式") print(" g(lambda)。此时精确解就是 (sigma_t/sigma_s) x + sigma_t * 积分 e^lambda g。") print(" x0_hat 是二次的,所以**任何真正的三阶方法都必须一步算准**。") local_err_test() # ── G. 卡住 3M 的不是开头那一步 ── print() print("=" * 72) print("G. 是不是「第一步被迫降阶」拖累的?给第一步多花一次评估试试") print("=" * 72) for tag, kw in (("第一步一阶(默认)", {}), ("第一步改中点法(多 1 次评估)", {"warm2s": True})): ns, es = [], [] for S in (24, 32, 48, 64, 96, 128, 192): grid = uniform_lambda_grid(S)[1] c = Counter(model) ns.append(c.n if False else S) es.append(rms_err(solve_multistep(grid, c, x_init.copy(), order=3, d2_scale=1.0, **kw), xref)) p, _ = fit_slope(ns, es, res, nfe_min=24) print(f" 3M(D2 原式)+ {tag:<22} 实测阶数 p = {p:.2f}") print(" → 没变化。第一步不是瓶颈,瓶颈在系数本身(见 F 节)") if __name__ == "__main__": main() spacing_lab.py # -*- coding: utf-8 -*- """步数该往哪放:同样 20 次评估,换个摆法能差出一个数量级。 阶数(dpm_solver_lab.py 里量出来的 p)说的是「步数变多时误差掉多快」, 但没说「步数摆在哪」。这一篇的主角是**时间步的摆法**: - 在 t 上等间隔(DDIM 论文 / diffusers 默认 leading) - 在 lambda 上等间隔(DPM-Solver 的建议) - 在 lambda 上等间隔但落回整数 t(真实网络只能吃整数) - 在 log sigma 上等间隔 - Karras / EDM 的 rho=7 摆法(SDXL、EDM 系列在用) 判断标准是两条:确定性误差(终点到参考解的距离)和「难度曲线」—— 沿着一条高精度参考轨迹量 |dx0_hat / d lambda|,看哪一段 lambda 上 x0_hat 变得最快。步数就该往那儿放。 /usr/local/bin/python3 spacing_lab.py 依赖:numpy(+ 同目录的 oracle_gmm.py / ddim_family.py / dpm_solver_lab.py) """ import numpy as np from oracle_gmm import (T, linear_schedule, make_target, sample_true, alpha_sigma, sliced_w1, make_dirs, N_EVAL) from ddim_family import uniform_t_grid, make_model, ddim_step from dpm_solver_lab import (uniform_lambda_grid, solve_ddim, solve_multistep, solve_dpm2s, lam_of_abar, abar_of_lam, nearest_t, Counter, rms_err) _BETAS, _ABAR = linear_schedule() _LAM = lam_of_abar(_ABAR) # ══════════════════════════════════════════════════════════════════ # 1. 五种摆法 # ══════════════════════════════════════════════════════════════════ def grid_uniform_t(S: int): """在 t 上等间隔(diffusers 的 leading / DDIM 论文默认)。""" return uniform_t_grid(S)[1] def grid_uniform_lambda(S: int): """在 lambda 上等间隔,连续 abar(DPM-Solver 理论里的标准摆法)。""" return uniform_lambda_grid(S, quantize=False)[1] def grid_uniform_lambda_q(S: int): """在 lambda 上等间隔,但落回整数时间步(真实网络只认整数 t)。""" return uniform_lambda_grid(S, quantize=True)[1] def _ratio(abar): """EDM / diffusers 语境里的 sigma:sigma_t / alpha_t = e^{-lambda}。 注意这不是 sqrt(1-abar)。VP 调度下 lambda 就是 -log 这个量, 所以「在 log sigma 上等间隔」等于「在 lambda 上等间隔」——两者是同一个东西, 表里因此只留一种。 """ return np.sqrt(np.maximum(1.0 - abar, 1e-300) / np.maximum(abar, 1e-300)) def grid_uniform_lambda_capped(S: int, lam_max: float): """在 lambda 上等间隔,但只走到 lam_max 就收尾。 为什么要这个变体:uniform-t 天然把 lambda 的上界截在某个值(NFE=20 时 是 +1.76),剩下的干净端交给最后那个「直接输出 x0_hat」的一阶收尾步。 如果 uniform-t 赢的原因是「截得早」而不是「t 本身特别」,那么把 uniform-lambda 截到同一个上界,两者就该差不多。 """ l0 = _LAM[T] lams = np.linspace(l0, lam_max, S) return np.concatenate([abar_of_lam(lams), [1.0]]) def grid_karras(S: int, rho: float = 7.0): """Karras / EDM 的摆法:在 sigma^(1/rho) 上等间隔,EDM 原文取 rho=7。 sigma 用的是 EDM 的那个 sigma(= sigma_t/alpha_t),VP 下等于 e^{-lambda}, 区间是 [_ratio(abar_T), _ratio(abar_1)] ≈ [158, 0.01]。 """ r_hi, r_lo = _ratio(_ABAR[T]), _ratio(_ABAR[1]) r = (r_hi ** (1.0 / rho) + np.arange(S) / (S - 1) * (r_lo ** (1.0 / rho) - r_hi ** (1.0 / rho))) ** rho return np.concatenate([1.0 / (1.0 + r ** 2), [1.0]]) GRIDS = { "uniform-t": grid_uniform_t, "uniform-lambda": grid_uniform_lambda, "uniform-lambda(整数t)": grid_uniform_lambda_q, "lambda截到1.8": lambda S: grid_uniform_lambda_capped(S, 1.8), "lambda截到3.0": lambda S: grid_uniform_lambda_capped(S, 3.0), "karras(rho=7)": grid_karras, "karras(rho=3)": lambda S: grid_karras(S, rho=3.0), } # ══════════════════════════════════════════════════════════════════ # 2. 难度曲线:沿着参考轨迹看 x0_hat 在哪一段变得最快 # ══════════════════════════════════════════════════════════════════ def difficulty_profile(model, n=400, steps=1500, seed=5): """返回 (lam_mid, weight):一阶方法每一步的误差贡献权重。 推导:第 i 步的局部误差(在 x 的单位里)是 delta_i = sigma_{t_i} * ∫ e^lambda [x0(lam) - x0(lam_i)] d lambda ≈ sigma_{t_i} * e^{lam_i} * (h^2/2) * x0'(lam_i) = alpha_{t_i} * e^{-h} * (h^2/2) * x0'(lam_i) 也就是说局部误差自带一个 alpha_t 因子:噪声端 alpha≈0,轨迹在 lambda 上 几乎不动,走错一点也不打紧;干净端 alpha≈1,同样的 h 会实打实地错。 所以难度权重是 alpha(lambda) * |dx0_hat / d lambda|,不是单纯的导数。 """ x = np.random.default_rng(seed).standard_normal((n, 2)) grid = uniform_lambda_grid(steps, quantize=False)[1] lams, x0s, alphas = [], [], [] for i in range(len(grid) - 1): a_s, s_s = alpha_sigma(grid[i]) a_t, s_t = alpha_sigma(grid[i + 1]) eps = model(x, grid[i]) m0 = (x - s_s * eps) / a_s lams.append(lam_of_abar(grid[i])) alphas.append(a_s) x0s.append(m0.copy()) x = (s_t / s_s) * x - a_t * (np.exp(-(lam_of_abar(grid[i + 1]) - lam_of_abar(grid[i]))) - 1.0) * m0 if s_t <= 0: break lams = np.array(lams) alphas = np.array(alphas) x0s = np.array(x0s) # (steps, n, 2) dx = np.sqrt(((np.diff(x0s, axis=0)) ** 2).sum(-1).mean(-1)) speed = dx / np.abs(np.diff(lams)) lam_mid = 0.5 * (lams[1:] + lams[:-1]) a_mid = 0.5 * (alphas[1:] + alphas[:-1]) return lam_mid, a_mid * speed, speed # ══════════════════════════════════════════════════════════════════ # 3. 主流程 # ══════════════════════════════════════════════════════════════════ def main(): rng = np.random.default_rng(20260927) w, mu, Sig = make_target() model = make_model("oracle", w, mu, Sig) TH = make_dirs() Xref = sample_true(N_EVAL, w, mu, Sig, rng) n_ref = 1000 x_init = rng.standard_normal((n_ref, 2)) g_ref = uniform_lambda_grid(3000, quantize=False)[1] xref = solve_multistep(g_ref, model, x_init.copy(), order=3) g_chk = uniform_lambda_grid(1500, quantize=False)[1] res = rms_err(xref, solve_multistep(g_chk, model, x_init.copy(), order=3)) print("=" * 72) print("A. 参考解分辨率") print("=" * 72) print(f" 3000 步 vs 1500 步 = {res:.2e}(误差表只能信到这个量级)") # ── B. 摆法对比(确定性误差)── print() print("=" * 72) print("B. 同样 NFE,步数摆在哪:终点到参考解的 RMS 距离") print("=" * 72) print(f" {'摆法':>22} " + " ".join(f"{'N=' + str(s):>9}" for s in (10, 20, 50)) + " | " + " ".join(f"{'N=' + str(s):>9}" for s in (10, 20, 50))) print(f" {'':>22} " + " ".join(f"{'DDIM':>9}" for _ in range(3)) + " | " + " ".join(f"{'2M':>9}" for _ in range(3))) for name, build in GRIDS.items(): row_a, row_b = [], [] for S in (10, 20, 50): grid = build(S) row_a.append(rms_err(solve_ddim(grid, model, x_init.copy()), xref)) row_b.append(rms_err(solve_multistep(grid, model, x_init.copy(), order=2), xref)) print(f" {name:>22} " + " ".join(f"{v:9.3e}" for v in row_a) + " | " + " ".join(f"{v:9.3e}" for v in row_b)) # ── C. 摆法对比(分布层面)── print() print("=" * 72) print("C. 同样的比较,换成 SW1(地板约 0.030,重复换起点看抖动)") print("=" * 72) print(f" {'摆法':>22} " + " ".join(f"{'N=' + str(s):>16}" for s in (10, 20))) for name, build in GRIDS.items(): row = [] for S in (10, 20): grid = build(S) vs = [] for r in range(5): x0 = np.random.default_rng(700 + S * 10 + r).standard_normal( (N_EVAL, 2)) vs.append(sliced_w1(solve_multistep(grid, model, x0, order=2), Xref, TH=TH)) row.append((np.mean(vs), np.std(vs))) print(f" {name:>22} " + " ".join(f"{v[0]:.4f}±{v[1]:.3f}" for v in row)) # ── D. 难度曲线 ── print() print("=" * 72) print("D. 难度曲线:x0_hat 在哪一段 lambda 上变化最快") print("=" * 72) lams, weight, raw = difficulty_profile(model) lo = lams < 0 print(f" 只看 |dx0/dlambda|(不带 alpha 因子):" f"噪声端 {raw[lo].mean():.3e} vs 干净端 {raw[~lo].mean():.3e}") print(f" 乘上 alpha_t 之后的**误差权重**: " f"噪声端 {weight[lo].mean():.3e} vs 干净端 {weight[~lo].mean():.3e}" f" ← 干净端重 {weight[~lo].mean() / weight[lo].mean():.1f} 倍") top = np.argsort(weight)[-5:][::-1] print(" 权重最大的 5 段:") for i in top: print(f" lambda = {lams[i]:+6.2f} 权重 = {weight[i]:.3e}") print(" → alpha_t 这个因子把难度整体推向干净端:噪声端轨迹在 lambda 上几乎") print(" 不动(alpha≈0),走错一点也不打紧;干净端才是一步错步步错的地方") # ── E2. 收尾步不是免费的 ── print() print("=" * 72) print("E2. 收尾步的代价:在 lambda_max 处直接输出 x0_hat,到底偏了多少") print("=" * 72) print(" 沿高精度参考轨迹取快照,对每个 lambda 量") print(" || x0_hat(x_lambda, lambda) - 轨迹终点 || —— 这就是提前收尾的偏差。") g_fine = uniform_lambda_grid(3000, quantize=False)[1] x = x_init.copy() snaps = [] for i in range(len(g_fine) - 1): a_s, s_s = alpha_sigma(g_fine[i]) a_t, s_t = alpha_sigma(g_fine[i + 1]) eps = model(x, g_fine[i]) m0 = (x - s_s * eps) / a_s if i % 60 == 0: snaps.append((lam_of_abar(g_fine[i]), float(np.sqrt(((m0 - xref[:n_ref]) ** 2).sum(1).mean())))) x = (s_t / s_s) * x - a_t * (np.exp(-(lam_of_abar(g_fine[i + 1]) - lam_of_abar(g_fine[i]))) - 1.0) * m0 if s_t <= 0: break for lam, d in snaps[::5]: print(f" lambda = {lam:+6.2f} 提前收尾的偏差 = {d:.3e}") print(" → 偏差随 lambda 单调下降:截得越早,这一步白送的误差越大,") print(" 而且它**不随步数下降**(只取决于截在哪个 lambda)") print() print(" 于是 lambda_max 有个最优值(2M, N=20,只改截断点):") for cap in (1.0, 1.5, 2.0, 2.5, 3.0, 3.5, 4.0, 4.61): grid = grid_uniform_lambda_capped(20, cap) e = rms_err(solve_multistep(grid, model, x_init.copy(), order=2), xref) print(f" lambda_max = {cap:4.2f} 终点误差 = {e:.3e}") print(" → 两头都变差:截太早是收尾偏差,截太晚是每步的 h 变大") # ── E. 20 步时各摆法的落脚点 ── print() print("=" * 72) print("E. NFE=20 时各摆法把步子落在 lambda 的哪几个位置") print("=" * 72) for name, build in GRIDS.items(): grid = build(20) lam = np.array([lam_of_abar(g) if g < 1.0 else np.inf for g in grid]) lam_f = lam[np.isfinite(lam)] print(f" {name:>22}: lambda 从 {lam_f[0]:+.2f} 到 {lam_f[-1]:+.2f}," f"相邻间隔 min {np.diff(lam_f).min():.3f} / max {np.diff(lam_f).max():.3f}") if __name__ == "__main__": main() make_figures.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """给「从 DDIM 到高阶采样器」画五张图。 所有数字都是现场重算的:这里 import ddim_family / dpm_solver_lab 里的 采样器,重新跑一遍实验再画,不抄任何手打的表格。改了那边的实现,重跑 这个脚本图就会跟着变,不会出现「图上是旧数字、正文是新数字」。 五张图分别回答: 1. eta=0 和 eta=1 走的是两条什么样的路(同一份随机数,二维真轨迹) 2. eta 该调多大?oracle 模型下扫一遍,欠拟合模型下再扫一遍 3. 为什么换采样器缓存就得重调:模型输入每步变多少 4. 阶数:误差 vs NFE 的双对数斜率 5. 达到 1e-2 / 1e-3 到底要花几次模型评估 只依赖 numpy + matplotlib。跑法:/usr/local/bin/python3 make_figures.py """ from pathlib import Path import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np import oracle_gmm as OG import ddim_family as DF import dpm_solver_lab as DL ROOT = Path(__file__).resolve().parents[1] OUT = ROOT / "figures" OUT.mkdir(parents=True, exist_ok=True) plt.rcParams["font.sans-serif"] = ["PingFang SC", "Heiti TC", "Arial Unicode MS"] plt.rcParams["axes.unicode_minus"] = False plt.rcParams["figure.dpi"] = 130 RNG_SEED = 20260927 # ══════════════════════════════════════════════════════════════════ # 公共:目标分布、参考解、评估尺子 # ═════════════════════════════════════════════════════════════════= def build_world(): rng = np.random.default_rng(RNG_SEED) w, mu, Sig = OG.make_target() Xref = OG.sample_true(OG.N_EVAL, w, mu, Sig, rng) TH = OG.make_dirs() return w, mu, Sig, Xref, TH def floors(TH, Xref, w, mu, Sig, n=8, base=2000): """尺子的分辨率:真样本对真样本,量 n 次。""" v = [OG.sliced_w1(OG.sample_true(OG.N_EVAL, w, mu, Sig, np.random.default_rng(base + r)), Xref, TH=TH) for r in range(n)] return float(np.mean(v)), float(np.std(v)) # ══════════════════════════════════════════════════════════════════ # 图 1:两条轨迹 # ═════════════════════════════════════════════════════════════════= def fig_traj(w, mu, Sig, Xref): model = DF.make_model("oracle", w, mu, Sig) S = 20 ts, grid = DF.uniform_t_grid(S) n_show = 4 x_init, zs = DF.paired_noise(S, n_show, seed=4321) fig, ax = plt.subplots(figsize=(6.4, 5.4)) ax.scatter(Xref[:, 0], Xref[:, 1], s=3, c="#C9CDD4", alpha=0.45, label="真实样本(目标分布)") styles = {1.0: ("#D2691E", 0.9, 5, r"$\eta=1$(DDPM 祖采样式):每步都重掷噪声"), 0.0: ("#2F6FB3", 1.7, 8, r"$\eta=0$(DDIM,确定性):一条光滑的路")} for eta, (c, lw, ms, lab) in styles.items(): _, (xs, _) = DF.run_ddim(x_init.copy(), grid, model, eta, zs, track=True) P = np.stack(xs, axis=1) # (n, S+1, 2) for k in range(n_show): ax.plot(P[k, :, 0], P[k, :, 1], "-", color=c, lw=lw, alpha=0.85 if eta == 0.0 else 0.55, zorder=4 if eta == 0.0 else 3, label=lab if k == 0 else None) ax.scatter(P[k, 1:-1, 0], P[k, 1:-1, 1], s=ms, color=c, alpha=0.85 if eta == 0.0 else 0.5, zorder=4) ax.scatter(x_init[:, 0], x_init[:, 1], s=110, marker="*", c="#111111", zorder=6, label="起点(两边共用,只有 4 个)") ax.set_xlabel(r"$x_1$") ax.set_ylabel(r"$x_2$") ax.set_title("同一份起点、同一份随机数:N=20 步走出的两条路") ax.legend(loc="upper right", fontsize=8, framealpha=0.9) ax.set_aspect("equal", adjustable="box") fig.tight_layout() p = OUT / "traj_eta.png" fig.savefig(p) plt.close(fig) print(f" [1] {p.name} 轨迹 {n_show} 条 × 2 种 eta") # ══════════════════════════════════════════════════════════════════ # 图 2:eta 扫描(oracle + 欠拟合) # ═════════════════════════════════════════════════════════════════= def fig_eta_sweep(w, mu, Sig, Xref, TH): NS = (10, 20, 50, 100) etas = (0.0, 0.25, 0.5, 0.75, 1.0) flo, flo_sd = floors(TH, Xref, w, mu, Sig) panels = {} for tag, kind in (("oracle(真 score,模型误差=0)", "oracle"), ("欠拟合(模型以为数据是单高斯)", "gauss")): model = DF.make_model(kind, w, mu, Sig) panels[tag] = {e: [] for e in etas} for e in etas: for S in NS: _, grid = DF.uniform_t_grid(S) xi, zs = DF.paired_noise(S, OG.N_EVAL) x, _ = DF.run_ddim(xi, grid, model, e, zs) panels[tag][e].append(OG.sliced_w1(x, Xref, TH=TH)) fig, axes = plt.subplots(1, 2, figsize=(11.2, 4.4), sharey=True) cmap = plt.get_cmap("viridis") for ax, (tag, tab) in zip(axes, panels.items()): for i, e in enumerate(etas): ax.plot(NS, tab[e], "o-", color=cmap(0.15 + 0.7 * i / (len(etas) - 1)), lw=1.6, ms=5, label=r"$\eta=" + f"{e:g}" + r"$") ax.axhline(flo, ls="--", c="#999999", lw=1.2, label=f"噪声地板 {flo:.4f}") ax.set_xscale("log") ax.set_xticks(list(NS)) ax.set_xticklabels([str(s) for s in NS]) ax.set_xlabel("采样步数 N") ax.set_title(tag) ax.grid(alpha=0.25) axes[0].set_ylabel("SW1(越小越好)") axes[0].legend(fontsize=8) axes[1].legend(fontsize=8) fig.suptitle("eta 越小越好,而且两个模型下结论一致", fontsize=12) fig.tight_layout() p = OUT / "eta_sweep.png" fig.savefig(p) plt.close(fig) print(f" [2] {p.name} 地板 {flo:.4f}±{flo_sd:.4f};" f"oracle N=10: eta0={panels['oracle(真 score,模型误差=0)'][0.0][0]:.4f}" f" eta1={panels['oracle(真 score,模型误差=0)'][1.0][0]:.4f}") # ══════════════════════════════════════════════════════════════════ # 图 3:每步「模型输入变了多少」 # ═════════════════════════════════════════════════════════════════= def fig_cache_jitter(w, mu, Sig): model = DF.make_model("oracle", w, mu, Sig) S = 50 _, grid = DF.uniform_t_grid(S) x_init, zs = DF.paired_noise(S, OG.N_EVAL) fig, ax = plt.subplots(figsize=(6.6, 4.4)) cols = {0.0: "#2F6FB3", 0.5: "#7A7A7A", 1.0: "#D2691E"} last10 = {} for eta, c in cols.items(): _, (xs, _) = DF.run_ddim(x_init.copy(), grid, model, eta, zs, track=True) rel = [float(np.sqrt((((xs[i + 1] - xs[i]) ** 2).sum(1)).mean()) / np.sqrt(((xs[i] ** 2).sum(1)).mean())) for i in range(len(xs) - 1)] last10[eta] = float(np.mean(rel[-10:])) ax.plot(range(1, len(rel) + 1), rel, "-", color=c, lw=1.6, label=r"$\eta=" + f"{eta:g}" + r"$") ax.plot([len(rel) - 9, len(rel)], [last10[eta], last10[eta]], lw=3.2, color=c, alpha=0.35) ax.set_xlabel("步序号(越往右越接近干净端)") ax.set_ylabel(r"模型输入相对变化 $\Delta x_i / x_i$(长度之比,无量纲)") ax.set_title("缓存加速盯的就是这条线:eta 越大,每步输入跳得越狠") ax.legend(fontsize=9) ax.grid(alpha=0.25) ratio = last10[1.0] / last10[0.0] ax.annotate(f"最后 10 步:eta=1 是 eta=0 的 {ratio:.1f} 倍", xy=(0.42, 0.86), xycoords="axes fraction", fontsize=9, bbox=dict(boxstyle="round,pad=0.35", fc="#FFF6E5", ec="#D2691E")) fig.tight_layout() p = OUT / "cache_jitter.png" fig.savefig(p) plt.close(fig) print(f" [3] {p.name} 最后10步 eta0={last10[0.0]:.4f} " f"eta0.5={last10[0.5]:.4f} eta1={last10[1.0]:.4f} 倍数 {ratio:.1f}") # ══════════════════════════════════════════════════════════════════ # 图 4 + 图 5:阶数与评估预算 # ═════════════════════════════════════════════════════════════════= def build_reference(w, mu, Sig): """超高分辨率参考解;顺手用 2000 步自查它的收敛精度。 随机数的取法与 dpm_solver_lab.py 的 B 节一模一样(同一个种子、同样的 消耗顺序:先抽一批真样本、再抽初始噪声),所以这里打印的分辨率和那边 是同一个数,图上和正文不会对不上。 """ model = DF.make_model("oracle", w, mu, Sig) rng = np.random.default_rng(RNG_SEED) OG.sample_true(OG.N_EVAL, w, mu, Sig, rng) x_init = rng.standard_normal((1000, 2)) _, gf = DL.uniform_lambda_grid(4000, quantize=False) _, gm = DL.uniform_lambda_grid(2000, quantize=False) xref = DL.solve_multistep(gf, model, x_init.copy(), order=3, d2_scale=2.0) xchk = DL.solve_multistep(gm, model, x_init.copy(), order=3, d2_scale=2.0) return model, x_init, xref, DL.rms_err(xref, xchk) def fig_order_and_budget(w, mu, Sig): model, x_init, xref, res = build_reference(w, mu, Sig) NS = [6, 8, 12, 16, 24, 32, 48, 64, 96, 128, 192] solvers = { "DDIM(uniform-t)": lambda S: (DF.uniform_t_grid(S)[1], DL.solve_ddim, {}), "DDIM(=DPM-1, lambda)": lambda S: (DL.uniform_lambda_grid(S)[1], DL.solve_ddim, {}), "DPM-Solver-2 单步": lambda S: (DL.uniform_lambda_grid(max(S // 2, 2))[1], DL.solve_dpm2s, {}), "DPM-Solver++ 2M": lambda S: (DL.uniform_lambda_grid(S)[1], DL.solve_multistep, {"order": 2}), "3M(三阶项原式)": lambda S: (DL.uniform_lambda_grid(S)[1], DL.solve_multistep, {"order": 3, "d2_scale": 1.0}), "3M(三阶项修正)": lambda S: (DL.uniform_lambda_grid(S)[1], DL.solve_multistep, {"order": 3, "d2_scale": 2.0}), } cols = ["#2F6FB3", "#6FA8DC", "#7A7A7A", "#3E8E41", "#D2691E", "#B03060"] tab = {k: [] for k in solvers} nfes = [] for S in NS: cn = [] for k, build in solvers.items(): grid, fn, kw = build(S) c = DL.Counter(model) tab[k].append(DL.rms_err(fn(grid, c, x_init.copy(), **kw), xref)) cn.append(c.n) nfes.append(int(round(np.mean(cn)))) # ── 图 4:阶数 ── fig, ax = plt.subplots(figsize=(7.0, 5.0)) for (k, errs), c in zip(tab.items(), cols): p, _ = DL.fit_slope(nfes, errs, res) ax.plot(nfes, errs, "o-", color=c, lw=1.5, ms=4.5, label=f"{k} " + r"$p=" + f"{p:.2f}" + r"$") ax.axhline(res, ls="--", c="#BBBBBB", lw=1.2) ax.text(nfes[-1], res * 1.15, f"参考解分辨率 {res:.1e}", fontsize=8, color="#888888", ha="right") ax.set_xscale("log") ax.set_yscale("log") ax.set_xlabel("NFE(真实模型评估次数,数出来的)") ax.set_ylabel("终点到参考解的 RMS 距离") ax.set_title("斜率就是阶数:低阶方法要加一个数量级的步数才能追上") ax.grid(alpha=0.25, which="both") ax.legend(fontsize=8, loc="lower left") fig.tight_layout() p4 = OUT / "order_nfe.png" fig.savefig(p4) plt.close(fig) # ── 图 5:评估预算 ── scan = [4, 6, 8, 10, 12, 16, 20, 24, 32, 40, 48, 64, 80, 96, 128, 160, 192, 256, 384] got = {} for k, build in solvers.items(): row = {} for S in scan: grid, fn, kw = build(S) c = DL.Counter(model) e = DL.rms_err(fn(grid, c, x_init.copy(), **kw), xref) for tol in (1e-2, 1e-3): if tol not in row and e <= tol: row[tol] = c.n got[k] = row fig, ax = plt.subplots(figsize=(7.4, 4.6)) names = list(solvers) xs = np.arange(len(names)) for j, (tol, lab) in enumerate(((1e-2, r"压到 $10^{-2}$ 所需 NFE"), (1e-3, r"压到 $10^{-3}$ 所需 NFE"))): vals, txt = [], [] for k in names: v = got[k].get(tol) vals.append(v if v else 0) txt.append(str(v) if v else ">" + str(scan[-1])) bars = ax.bar(xs + (j - 0.5) * 0.38, vals, width=0.36, color=["#8FB8DE", "#1F4E79"][j], label=lab) for b, t in zip(bars, txt): ax.text(b.get_x() + b.get_width() / 2, b.get_height() + 4, t, ha="center", fontsize=8) ax.set_xticks(xs) ax.set_xticklabels(names, rotation=18, ha="right", fontsize=8) ax.set_ylabel("NFE") ax.set_title("同一个精度,高阶方法省下的评估次数(柱子顶上标的是实测值)") ax.legend(fontsize=9) ax.grid(alpha=0.25, axis="y") fig.tight_layout() p5 = OUT / "nfe_budget.png" fig.savefig(p5) plt.close(fig) print(f" [4] {p4.name} 参考解分辨率 {res:.3e}") print(f" [5] {p5.name} 1e-2 所需 NFE: " + " ".join(f"{k}={got[k].get(1e-2, '>384')}" for k in names)) print(f" 1e-3 所需 NFE: " + " ".join(f"{k}={got[k].get(1e-3, '>384')}" for k in names)) def main(): print("生成配图(数值全部现场重算)") w, mu, Sig, Xref, TH = build_world() fig_traj(w, mu, Sig, Xref) fig_eta_sweep(w, mu, Sig, Xref, TH) fig_cache_jitter(w, mu, Sig) fig_order_and_budget(w, mu, Sig) print("完成,输出目录:", OUT) if __name__ == "__main__": main()
2026年09月27日
3 阅读
0 评论
0 点赞
2026-09-27
AIGC 基本功|DDPM 训练目标与采样流程-DDPM
DDPM 训练目标与采样流程 所属方向:生成范式 | 难度:进阶 | 前置知识:扩散过程的前向与反向推导(前向闭式解、反向后验、以及 ELBO 化简到「预测噪声」的那一步都在那篇推过,这里直接用结论) 关键词:DDPM、L_simple、ELBO 每步权重、采样方差、Algorithm 1/2、fixed_small / fixed_large 01. 为什么需要它 我第一次照着 DDPM 论文的公式老老实实实现训练目标,结果比「偷懒版」更差。 实验是这样的:数据是二维高斯,调度用原文的线性 $T=1000$,模型是一个跨时间步共享参数的小网络(18 个可学参数),三种目标各训一份,然后算生成分布的精确负对数似然。结果是: 训练目标 期望 NLL(nats,越小越好) 等权 $L_{\text{simple}}$ 1.898760 真 ELBO 每步权重 1.922295 差了 0.0235 nats。而 Ho 等人在原文里就明说了:他们把 ELBO 里那一串只跟 $t$ 有关的系数扔掉,直接等权,反而「sample quality 更好」。我当时以为这只是工程上的凑巧,跑完才发现不是——权重的形状和模型容量是绑在一起的。把同一个网络的时间基从 3 项加到 5 项(30 个参数),结论立刻反过来:真 ELBO 权重 1.880952,等权 1.881758,真 ELBO 反超。 第二个坑是采样方差。我一直把 $\tilde\beta_t$ 当成「真实后验方差」。实测把方差换成真正的后验协方差之后,生成分布的期望 NLL 从 1.881015 降到 1.880914,差的 0.000101 nats 全部来自把方差钉死成 $\tilde\beta_t$;生成协方差与真值的比值从 $[0.98946,\ 0.98713]$ 变成 $[1.0,\ 1.0]$。准确地说:$\tilde\beta_t I$ 是已知 $x_0$ 的条件后验协方差;对边缘反向 $q(x_{t-1}\mid x_t)$,它是全协方差分解中的下界。本例用它生成的分布偏窄,不能把数值幅度推广到所有数据。 第三个观察是权重分配。$t\ge2$ 的 ELBO KL 权重跨 120.4 倍,而低噪声端的噪声预测有很高的不可约误差。不可约误差不产生期望梯度;真正被权重重新分配的是可学习的超额误差及随机梯度噪声。后文的容量对照说明“加权在哪种模型上更合适”需要实测,不能把某个时间段断言成无用功。 这三个坑合起来就是这一篇要讲的事:DDPM 的训练目标和采样循环是两件独立的东西,前者管误差权重,后者管转移;在 $L_{simple}$ 下可以分别选择,但 ELBO 权重显式依赖反向方差,并非完全独立。把这两件事分开看,后面所有改进(DDIM、CFG、flow matching)才读得懂。 02. 最小可用理解 三句话: 训练:每一步都是一个「从 $x_t$ 猜刚才加了什么噪声 $\varepsilon$」的回归问题。DDPM 把它做成等权均方误差 $L_{\text{simple}}$,故意不理 ELBO 给每步分配的系数。 权重:本例 KL 权重 $w_t$ 跨 120 倍。等权是改写目标;是否改善生成质量取决于误差分布、优化、容量和评测目标。 采样:反向一步的均值由 $\varepsilon$ 预测器决定,方差可取固定 $\tilde\beta_t$、$\beta_t$ 或学习值;固定方差不进入 $L_{simple}$,却进入 ELBO,学习方差还需要相应训练项。 这张图要看什么:左轴的两条线是真 ELBO 权重(蓝实线 $\sigma^2=\tilde\beta_t$,红虚线 $\sigma^2=\beta_t$),灰色点线是「等权」压平到 1 的位置。右轴绿线是这一时刻噪声里能学出来的比例 $R^2$。两条曲线正好反向:权重在 $t$ 很小的地方冲到 0.6,而那里的 $R^2$ 几乎是 0;权重最低的 $t\approx350$,反而是 $R^2$ 爬到一半的地方。低噪声端权重大、可预测噪声占比小;右端并不遵循这个反向关系。这提示检查容量分配,不足以单独证明等权必然更优。 03. 数学推导 3.1 ELBO 拆成每步的 KL 变分下界写出来是 $$L = \mathbb{E}_q\Big[-\log p_\theta(x_0|x_1) + \sum_{t=2}^{T} D_{\mathrm{KL}}\big(q(x_{t-1}|x_t,x_0)\,\|\,p_\theta(x_{t-1}|x_t)\big) + D_{\mathrm{KL}}\big(q(x_T|x_0)\,\|\,p(x_T)\big)\Big]$$ 最后一项在固定前向过程与先验时没有可学参数;$L_0=-\log p_\theta(x_0|x_1)$ 也是训练目标的一部分。DDPM 为离散像素采用离散化高斯解码似然,连续数据可选连续密度;中间那一长串是下面推导的 KL 项。把它记为 $L_{t-1}$(下标是 $t-1$ 因为它监督的是「从 $t$ 走到 $t-1$」这一步)。 为什么这一项好算?因为 $q$ 和 $p_\theta$ 都是高斯,两个高斯的 KL 有闭式解。而 $q(x_{t-1}|x_t,x_0)$ 在后验那篇已经推过:$q=\mathcal{N}(\tilde\mu_t,\ \tilde\beta_t I)$,其中 $\tilde\beta_t$ 是一个不依赖 $x_t$ 的常数,$\tilde\mu_t$ 是 $x_t$ 与 $x_0$ 的线性组合。在反向方差固定、不由模型学习的前提下,KL 里唯一依赖模型的部分是均值之差: $$L_{t-1} = \mathbb{E}_q\Big[\tfrac{1}{2\sigma_t^{2}}\big\|\tilde\mu_t(x_t,x_0)-\mu_\theta(x_t,t)\big\|^{2}\Big] + C$$ $C$ 是与参数无关的常数,$\sigma_t^2$ 是反向链在这一步选的方差。 3.2 把均值之差换成噪声之差 $\tilde\mu_t$ 和 $\mu_\theta$ 用同一组系数线性组合,前者组合的是真 $x_0$,后者组合的是模型猜的 $\hat x_0$。把 $x_0=\frac{1}{\sqrt{\bar\alpha_t}}(x_t-\sqrt{1-\bar\alpha_t}\,\varepsilon)$ 和 $\hat x_0=\frac{1}{\sqrt{\bar\alpha_t}}(x_t-\sqrt{1-\bar\alpha_t}\,\varepsilon_\theta)$ 代进去,$x_t$ 那一项整整齐齐地消掉,剩下 $$\tilde\mu_t-\mu_\theta = -\frac{\beta_t}{\sqrt{\alpha_t}\sqrt{1-\bar\alpha_t}}\big(\varepsilon-\varepsilon_\theta\big)$$ 注意这里出现的 $\sqrt{\bar\alpha_{t-1}}/\sqrt{\bar\alpha_t}=1/\sqrt{\alpha_t}$——这是整个化简能成立的关键一步,两个系数只差一个 $\sqrt{\alpha_t}$,所以差值是干净的单项。 平方之后: $$L_{t-1} = \mathbb{E}\Big[w_t\big\|\varepsilon-\varepsilon_\theta\big\|^{2}\Big] + C,\qquad w_t=\frac{\beta_t^{2}}{2\sigma_t^{2}\alpha_t(1-\bar\alpha_t)}$$ 这就是那个「只跟 $t$ 有关的系数」。固定方差时由调度和方差选择决定;若学习方差,就不能把相关项都视为与参数无关的常数。 代入两种方差选择,还能再化简一层: $$w_t=\frac{\beta_t}{2\alpha_t(1-\bar\alpha_{t-1})}\ \ (\sigma_t^{2}=\tilde\beta_t),\qquad w_t=\frac{\beta_t}{2\alpha_t(1-\bar\alpha_t)}\ \ (\sigma_t^{2}=\beta_t)$$ 两条式子的分母只差一个下标。附录 ddpm_lab.py 里解析式和化简式两条都算了,最大相对差 3.93×10⁻¹⁶($\tilde\beta$)和 4.13×10⁻¹⁶($\beta$),就是浮点误差级别。 $t=1$ 时 $\tilde\beta_1=0$,上述固定小方差高斯 KL 权重分母为零、分子非零,不能使用。它不是 $0/0$。ELBO 此时对应重建似然 $L_0$,需另行定义;不是“所有最后一步都不准有高斯密度”。本实验用 $\beta_1>0$ 的连续高斯重建权重补这一点,因此表中的 elbo_tilde 是这一连续数据约定,不是逐字复现离散像素 ELBO。 3.3 $w_t$ 长什么样 线性调度 $T=1000$、$\beta$ 从 $10^{-4}$ 到 $0.02$(原文 CIFAR-10 的配置),$\bar\alpha_T=4.0358\times10^{-5}$。实跑出来: $t$ $\bar\alpha_t$ $\beta_t$ $w_t$($\sigma^2=\tilde\beta_t$) $w_t$($\sigma^2=\beta_t$) 相对 $t=500$ 2 9.998e-01 0.00012 5.9967e-01 2.7269e-01 108.87 10 9.981e-01 0.00028 8.6437e-02 7.3717e-02 15.69 50 9.710e-01 0.00108 1.9279e-02 1.8583e-02 3.50 100 8.970e-01 0.00207 1.0267e-02 1.0081e-02 1.86 250 5.241e-01 0.00506 5.3733e-03 5.3432e-03 0.98 500 7.859e-02 0.01004 5.5082e-03 5.5034e-03 1.00 750 3.351e-03 0.01502 7.6506e-03 7.6502e-03 1.39 900 2.752e-04 0.01801 9.1717e-03 9.1716e-03 1.67 1000 4.036e-05 0.02000 1.0205e-02 1.0204e-02 1.85 最后一列是「相对 $t=500$ 的倍数」,等权相当于把这一列全压成 1.00。 形状是两端高、中间平的一个碗。$\tilde\beta$ 版本的最大值是最小值的 120.4 倍;$\beta$ 版本因为分母用的是 $1-\bar\alpha_t$ 而不是 $1-\bar\alpha_{t-1}$,在 $t$ 很小的地方发散得温和些,动态范围小很多。 左端 $1-\bar\alpha_{t-1}$ 很小,使权重大;中段分母增长与线性增大的 $\beta_t$ 互相抵消;右端 $1-\bar\alpha_t$ 趋近 1 后,$w_t\approx\beta_t/(2\alpha_t)$ 随 $\beta_t$ 上升。不能说 $t=250$ 时分母已饱和:表中此时 $1-\bar\alpha_t\approx0.476$。 这个式子不是我拍脑袋的,脚本里用两个高斯的 KL 数值核了一遍:随机取 $x_0$ 和 $\varepsilon$,令 $\varepsilon_\theta=\varepsilon+\delta$($\delta$ 是一个固定的小扰动),比较「数值算的 KL」与「$w_t\|\delta\|^2$」: t KL(数值) w_t*||delta||^2 相对差 2 1.416848e-02 1.416848e-02 2.89e-13 100 1.612712e-03 1.612712e-03 4.17e-14 500 5.035556e-04 5.035556e-04 1.72e-15 1000 2.188401e-03 2.188401e-03 1.61e-14 相对差在 $10^{-13}$ 量级,推导和代码对上了。 3.4 权重最大的地方,恰恰最学不到东西 $w_t$ 只说明「这一步的误差在 ELBO 里值多少钱」,没说明这一步的误差能不能被压下去。真正该看的是贝叶斯地板:给定 $x_t$ 后 $\varepsilon$ 还剩多少不确定性。 数据是 $\mathcal{N}(\mu_0,S_0)$ 时,$x_t$ 的协方差是 $\bar\alpha_t S_0+(1-\bar\alpha_t)I$,条件方差有闭式解,地板就是 $$\mathrm{floor}(t)=\frac{1}{D}\,\mathrm{tr}\Big[\mathrm{Var}(\varepsilon\,|\,x_t)\Big]=\frac{D-(1-\bar\alpha_t)\,\mathrm{tr}\big[(\bar\alpha_tS_0+(1-\bar\alpha_t)I)^{-1}\big]}{D}$$ 定义可学占比 $R^2(t)=1-\mathrm{floor}(t)$($\varepsilon$ 每维先验方差是 1)。实测:$t=1$ 时 $R^2=3.22\times10^{-4}$,$t=500$ 时 0.9616,$t=1000$ 时 1.0000。 在 $t=1$,噪声系数是 $\sqrt{1-\bar\alpha_1}=0.01$,$R^2=3.22\times10^{-4}$ 即约 0.0322%。这说明噪声有很大的条件方差,不说明最优预测毫无用途:生成所需 score 正由小的条件均值决定。$t=2$ 的 108.87 倍权重与 $t=1$ 的 $R^2$ 不能混为同一个时间步。 因此应把不可约地板与超额误差分开分析。不可约项对参数的期望梯度为零,有限容量如何分配精度、随机梯度方差多大,才决定加权目标的优化表现。 3.5 采样方差不是后验方差 $\tilde\beta_t$ 是 $q(x_{t-1}|x_t,x_0)$ 的方差,条件是「已知 $x_0$」。但采样时我们没有 $x_0$,只有模型猜的 $\hat x_0$,它自己有误差。把这一层不确定性算进去,真实后验协方差是 $$\Sigma_t = \tilde\beta_t I + \Big(\tfrac{\sqrt{\bar\alpha_{t-1}}\beta_t}{1-\bar\alpha_t}\Big)^{2}\mathrm{Var}(x_0-\hat x_0\,|\,x_t)$$ 多出来的第二项就是「均值估计不准」带来的。实测:用最优 $\varepsilon$ 预测器配 $\tilde\beta_t$,生成协方差只有真值的 $[0.98946,\ 0.98713]$;换成真实后验协方差,比值精确回到 $[1.0,\ 1.0]$,NLL 差 $3.56\times10^{-10}$ nats——在本例中极接近数据分布;仍有有限终点先验与数值误差。 这一节的结论对后面很重要:反向链的均值和方差是两件事。均值决定你往哪走,方差决定你抖多厉害;在这个高斯实验中,固定 $\tilde\beta_t$ 使协方差特征值约偏小 1.1%~1.3%;该数值不是通用结论。 04. 代码实现 4.1 调度、$\tilde\beta_t$、$w_t$(20 行) def linear_beta(T=1000, b1=1e-4, bT=0.02): """DDPM 原文 CIFAR-10 用的线性调度。""" return np.linspace(b1, bT, T) class Sched: """下标约定:abar[t]、beta[t]、bt[t] 里的 t 取 1..T;t=0 是数据本身。""" def __init__(self, beta): self.T = len(beta) self.beta = beta self.alpha = 1.0 - beta self.abar = np.concatenate([[1.0], np.cumprod(self.alpha)]) # 长度 T+1 self.bt = np.empty(self.T + 1) self.bt[0] = 0.0 self.bt[1:] = (1.0 - self.abar[:-1]) / (1.0 - self.abar[1:]) * beta def weight(self, which="tilde"): s2 = self.bt[1:] if which == "tilde" else self.beta with np.errstate(divide="ignore", invalid="ignore"): w = self.beta ** 2 / (2.0 * s2 * self.alpha * (1.0 - self.abar[1:])) if which == "tilde": # t=1 时 tilde_beta_1=0,分母为零。DDPM 把这一步交给 L_0 单独处理, # 这里退化地用 fixed_large 的权重顶上,只为让训练能跑。 w[0] = self.beta[0] / (2.0 * self.alpha[0] * (1.0 - self.abar[1])) return w np.errstate 那一行不是装饰:\tilde\beta_1=0 会触发除零,不包起来的话 w/w.mean() 之后整条权重变成 nan,训练静悄悄地训出一个废模型。 4.2 三种目标:共享参数的带权最小二乘 要让「权重决定容量往哪搬」这件事看得见,模型必须跨时间步共享参数——否则每一步各自拟合,权重只影响每一步自己的收敛速度,看不出搬移效果。所以用一个仿射 $\varepsilon$ 预测器加固定时间基: $$\varepsilon_\theta(x_t,t)=W\,\big[x_t\otimes\psi(t),\;\psi(t)\big],\qquad \psi(t)=\big[1,\ (t/T)^{0.5},\ t/T\big]$$ (指数取 $[0,0.5,1]$ 而不是 $[0,1,2]$:因为 $\sqrt{1-\bar\alpha_t}$ 在 $t$ 小的时候像 $\sqrt{t}$,纯整数次幂的多项式基在这一段拟合不出来,曲线会剧烈振荡。这是踩过的坑。) 训练集 $M=120000$ 条 $(x_0,t,\varepsilon)$,特征维度 9,可学参数 18 个。三种权重都归一化到均值 1(整体缩放不改变解,只为了让正则强度可比)。用带权岭回归一步算出闭式解,不做迭代。 4.3 精确 NLL:把蒙特卡洛噪声干掉 一开始我用采样算 NLL,三种目标差 0.002~0.005 nats,而 2 万样本的蒙特卡洛误差就有 ±0.007——结论完全淹没在噪声里。改走闭式解:仿射模型下 $p_\theta(x_0)$ 仍是高斯,把密度沿反向链往前传,最后 $$\mathbb{E}\big[-\log p_\theta(x_0)\big]=H\big[\mathcal{N}(\mu_0,S_0)\big]+D_{\mathrm{KL}}\big(\mathcal{N}(\mu_0,S_0)\,\|\,\mathcal{N}(m,S)\big)$$ 用 40 万样本复核,闭式解 1.880914 对蒙特卡洛 1.880890,差 $2.38\times10^{-5}$。两种独立实现(密度前向传播 / 线性映射复合)的均值差 $9.99\times10^{-16}$、协方差差 $5.00\times10^{-16}$。 三种目标 × 两种采样方差: 训练目标 $\sigma^2=\tilde\beta_t$ $\sigma^2=\beta_t$ 差 生成 std / 真 std uniform(等权) 1.898760 1.895815 −0.002945 0.9659 elbo_tilde 1.922295 1.924692 +0.002396 1.0392 elbo_large 1.914947 1.917129 +0.002181 1.0322 (最后一列是生成分布标准差与真分布标准差的比值,1 表示胖瘦刚好对上。等权明显偏窄,真 ELBO 权重的两个版本反而偏宽。) 4.4 权重把超额误差搬到了哪 「超额」= 实际 MSE − 贝叶斯地板,包括容量、有限样本、正则化与优化误差: $t$ 可学占比 $R^2$ 地板 等权超额 真 ELBO 超额 1 0.0003 0.9997 1.516e-01 1.017e-02 5 0.0022 0.9978 1.120e-01 6.824e-03 20 0.0182 0.9818 5.295e-02 3.897e-03 50 0.0853 0.9147 9.778e-03 7.323e-03 100 0.2510 0.7490 4.271e-03 2.560e-02 200 0.5663 0.4337 1.324e-02 2.820e-02 400 0.9000 0.1000 1.067e-03 4.463e-03 600 0.9876 0.0124 5.375e-03 1.349e-02 800 0.9993 0.0007 3.118e-04 1.082e-03 1000 1.0000 0.0000 1.337e-02 2.345e-02 分档汇总:低噪声档 $t\le50$,等权超额 5.140e-02,真 ELBO 5.137e-03(比值 0.10,好 10 倍);高噪声档 $t\ge400$,等权 3.459e-03,真 ELBO 8.060e-03(比值 2.33,差一倍多)。 这张图要看什么:(a) 两条曲线的交叉点大约在 $t=70$——交叉点左边真 ELBO 权重更准,右边等权更准;(b) 取比值后看得更清楚,灰色填充区(比值 <1)是真 ELBO 占优的低噪声段,红色填充区(比值 >1)是等权占优的中高噪声段。加权不是全面变好,是把误差从一段搬到另一段。等权之所以赢,是因为搬走的那一头($t\le50$)本来误差就大到没救(1.5e-01 对地板 0.9997),而搬来的那一头($t\ge400$)绝对误差只有 1e-03 量级,赔得起。 4.5 换个容量档,结论反过来 时间基从 3 项加到 5 项($[0,0.5,1,1.5,2]$,30 个参数),其他完全不动: 容量档 时间基项数 可学参数 NLL(等权) NLL(真 ELBO) 谁更好 低噪声档超额比 高噪声档超额比 loose 5 30 1.881758 1.880952 真 ELBO 0.71 3.91 tight 3 18 1.898760 1.922295 等权 0.10 2.33 这张表说明两档容量的排序不同,不能推出“容量足够时 ELBO 自然赢”。ELBO、生成 NLL 和感知质量也不是同一个目标;EDM、Min-SNR 同时涉及预条件、噪声分布或多任务梯度冲突,不应仅归因为模型变大。 4.6 真网络版:Algorithm 1 与 Algorithm 2 闭式解毕竟是玩具。附录 ddpm_train.py 给了一份 numpy 手写版:两层 MLP(输入 $10=2+8$ 维时间嵌入,隐层 64,输出 2),手写反向传播 + Adam,batch 256,20000 步,数据是 7 个高斯的混合(一个中心 + 六个环)。 梯度先对:手写反向传播 vs 有限差分,最大相对差 9.92×10⁻¹⁰。Algorithm 1(训练) 照抄原文,唯一变量就是 loss 前面的 $w_t$: def train(sc, w_full, steps, batch, rng, lr=1e-3): p = init_params(rng) opt = Adam(p, lr=lr) for _ in range(steps): # 1: t ~ Uniform({1, ..., T}) 2: x_0 ~ q(x_0) 3: eps ~ N(0, I) t_idx = rng.integers(1, T + 1, size=batch) # (256,) x0 = sample_mix(batch, rng) # (256, 2) eps = rng.standard_normal((batch, D)) # (256, 2) # 4: 一步加噪(闭式解,不需要真的走 t 步) a = sc.abar[t_idx] # (256,) xt = np.sqrt(a)[:, None] * x0 + np.sqrt(1.0 - a)[:, None] * eps # 5: 梯度下降一步 Z = np.concatenate([xt, time_embed(t_idx)], axis=1) # (256, 10) h1, h2, out = forward(p, Z) # (256,64) (256,64) (256,2) opt.step(p, backward(p, Z, h1, h2, out, eps, w_full[t_idx - 1])) return p 三个 shape 值得停下来看一眼。第一,t_idx 是一个长度为 batch 的随机整数向量,不是标量——每个样本走不同的时间步,这是 $L_{\text{simple}}$ 里那个 $\mathbb{E}_t$ 的实现方式,也是「等权」的字面含义:每个 $t$ 被抽中的概率相同。第二,Z 的第二维是 $10 = 2 + 8$,其中 8 维来自 TIME_FREQ = (1,2,4,8) 的 cos/sin 时间嵌入;不把 $t$ 喂进去的话,网络不知道当前噪声档位,$L_{\text{simple}}$ 根本学不动。第三,w_full[t_idx - 1] 这个减一是全文最容易写错的地方:时间步 $t$ 从 1 数到 $T$,而数组下标从 0 开始。写反了整条权重会整体错位一步,$t$ 很小的地方拿到的是 $t+1$ 的权重——在动态范围 120 倍的曲线上,错位一步就能让训练目标面目全非。 Algorithm 2(采样) 是完整 1000 步链: def sample_chain(sc, p, n, which="tilde", rng=None): x = rng.standard_normal((n, D)) # x_T ~ N(0, I) for t in range(T, 0, -1): eps_hat = predict_eps(p, x, t) a, a_prev = sc.abar[t], sc.abar[t - 1] alpha_t, beta_t = sc.alpha[t - 1], sc.beta[t - 1] mu = (x - beta_t / np.sqrt(1.0 - a) * eps_hat) / np.sqrt(alpha_t) sigma = np.sqrt(sc.bt[t]) if which == "tilde" else np.sqrt(beta_t) z = rng.standard_normal((n, D)) x = mu + (sigma * z if t > 1 else 0.0) # t == 1 时不加噪声 return x 最后那行 if t > 1 else 0.0 就是 diffusers 里 if t > 0 的同一个判断(下标约定差 1):最后一步不加噪声。均值那一行是 $\tilde\mu_t$ 的另一种写法——把 $x_0$ 用 Tweedie 公式替掉之后 $\tilde\mu_t=\frac{1}{\sqrt{\alpha_t}}\big(x_t-\frac{\beta_t}{\sqrt{1-\bar\alpha_t}}\varepsilon_\theta\big)$,比 3.2 节那个两项组合的形式更紧凑,也是 diffusers 之外大多数实现采用的写法。 采样之后用三个指标打分:$\log p_0$ 相对真样本、到最近模式中心的距离比、模式覆盖数(某个分量占比超过 2% 才算被覆盖)。这三个指标是互补的——$\log p_0$ 只衡量落在高密度区,模式覆盖才告诉你有没有整个丢掉某个模式;实测两个网络都是 7/7,说明等权赢的不是覆盖面,是落在模式里的精度。 训练目标 $\log p_0$ 相对真样本 到最近中心距离 / 真样本 模式覆盖 uniform −0.2208 ± 0.0195 1.0857 7/7 elbo_tilde −0.4212 ± 0.0062 1.1997 7/7 多峰数据上等权依然赢,而且差距比高斯那组更明显(0.20 vs 0.42 nats)。每步 MSE 剖面也和闭式解的预测一致:$t=1$ 处两个网络都是 ≈1.0(学不动),$t=1000$ 处 uniform 0.0016、elbo_tilde 0.0026(都学得很好)。 只换采样方差、不动网络: 方差选择 $\log p_0$ 相对真样本 到最近中心距离 / 真样本 $\tilde\beta_t$(fixed_small) −0.2363 ± 0.0123 1.0968 $\beta_t$(fixed_large) −0.2640 ± 0.0129 1.1099 这里和 4.3 的高斯实验结论相反:高斯闭式解里 $\beta_t$ 把偏窄补回来了(std 比 0.9659 → 0.9720),多峰数据上 $\tilde\beta_t$ 反而更好。原因不神秘——两者的相对差主要集中在低噪声段,高噪声端反而很接近(见 3.3 节),多峰分布的模式之间本来就脆弱,注入更多噪声会把样本推离模式中心(距离比 1.0968 → 1.1099 就是这个效果)。方差选择取决于数据、预测误差、调度与评价指标,不能只按是否多峰判断,别把一维高斯的结论直接搬。 05. 工业级实现对照 对照 huggingface/diffusers 的 src/diffusers/schedulers/scheduling_ddpm.py → DDPMScheduler.step(以 2026-09 时的实现为准)。这一份源码就是知识树里给这个节点配的 code_refs。 第一处差异:方差有六个分支,不是一个数。 我 4.1 节只写了 $\tilde\beta_t$ 和 $\beta_t$ 两种,框架里是 _get_variance() 的六个 variance_type: if variance_type == "fixed_small": # 默认:后验方差的下界 variance = variance elif variance_type == "fixed_small_log": # 取 log 再 exp(0.5*log),代数上等价 variance = torch.log(variance); variance = torch.exp(0.5 * variance) elif variance_type == "fixed_large": # 直接上 beta_t variance = current_beta_t elif variance_type == "fixed_large_log": variance = torch.log(current_beta_t) elif variance_type == "learned": # 网络自己输出 return predicted_variance elif variance_type == "learned_range": # 在 [tilde_beta, beta] 之间插值 min_log = torch.log(variance); max_log = torch.log(current_beta_t) frac = (predicted_variance + 1) / 2 variance = frac * max_log + (1 - frac) * min_log 值得注意的是 learned_range:它插值区间的两个端点恰好就是我 4.5 节扫描的那两个,而且是在 log 空间插值(这更合理,因为这两个量跨好几个数量级)。也就是说「$\tilde\beta_t$ 还是 $\beta_t$」在工业实现里不是二选一,而是交给网络学一个位置。 第二处差异:最后一步不加噪声。 源码里 variance = 0; if t > 0: ...。这正对应 3.2 节说的 $L_0$ 退化问题——$t=1$ 时 $\tilde\beta_1=0$,这是该采样算法的终步约定;连续数据模型仍可选择非零方差的终步重建分布。 第三处差异:累计噪声表预先计算,转移系数在 step 内由当前与目标时间步组合。 step() 里那两行: pred_original_sample_coeff = (alpha_prod_t_prev ** 0.5 * current_beta_t) / beta_prod_t current_sample_coeff = current_alpha_t ** 0.5 * beta_prod_t_prev / beta_prod_t pred_prev_sample = pred_original_sample_coeff * pred_original_sample + current_sample_coeff * sample 展开就是 3.2 节的 $\tilde\mu_t$,$c_0=\frac{\sqrt{\bar\alpha_{t-1}}\beta_t}{1-\bar\alpha_t}$、$c_x=\frac{\sqrt{\alpha_t}(1-\bar\alpha_{t-1})}{1-\bar\alpha_t}$,一字不差。 这张图要看什么:三条线是反向一步的三个配料随 $t$ 的变化——蓝线是「拉向 $\hat x_0$ 的系数」$c_0$,灰虚线是「保留 $x_t$ 的系数」$c_x$,红线是注入噪声的标准差 $\sqrt{\tilde\beta_t}$。绝大多数步子里 $c_0$ 在 1e-4 ~ 1e-2 量级,也就是一步只挪千分之几;只有最后几十步 $c_0$ 才冲到 1 附近,真正开始「成形」。这是完整 1000 步离散表下的系数,不意味着采样必须走满 1000 步。重排噪声表或换求解器后,模型可用更少的评估次数采样。 第四处差异:prediction_type 有三个选项。 epsilon(默认)、sample(直接预测 $x_0$)、v_prediction(预测 $v$)。源码把三者统一先转成 pred_original_sample,再走同一套系数——参数化只是换了个入口,反向一步的代数完全一样。参数化之间的权重差异在那篇《扩散过程的前向与反向推导》里算过($\varepsilon$ 1.0 倍、$v$ 2.5×10⁴、$x_0$ 2.5×10⁸)。 第五处:clip_sample=True 是默认开的。 把 pred_original_sample 夹到 $[-1,1]$。这在像素空间合理,在潜空间是错的——SD1.x 的 VAE 原始潜变量乘约 0.18215 后才送入扩散网络,尺度约归一到 1;但它并不被限制在 $[-1,1]$。因此不能照搬像素裁剪,是否关闭或采用其他限幅应遵循具体模型配置。 06. 代价与边界 等权省了什么:不用管 $\tilde\beta_1=0$ 的退化,不用管 $L_0$ 的离散解码器,没有显式时间权重,但不保证每步的梯度范数一致(这点很重要——那篇前向反向推导里算过,$x_0$ 参数化跨 2.5×10⁸ 倍,梯度会被最吵的几步吃光)。 等权赔了什么:三档实测摆在一起看—— 场景 等权 真 ELBO 权重 谁赢 容量宽裕(30 参数,高斯) 1.881758 1.880952 真 ELBO 容量吃紧(18 参数,高斯) 1.898760 1.922295 等权 容量吃紧(MLP,7 峰混合) −0.2208 −0.4212 等权 怎么选择权重:若关注低噪声重建精度,应在任务数据上测量各时段超额误差与最终质量,比较等权与加权;本例的 10 倍误差比不能直接推出超分、修复都不该用等权。 采样方差的边界:高斯数据上 $\beta_t$ 把偏窄从 3.4% 补到 2.8%(补回 0.6 个百分点,还没补满);多峰数据上 $\beta_t$ 反而更差。所以别把「fixed_large 更大更对」当成通例——要不要更大方差,取决于你的数据是不是多峰、模式之间经不经得起抖。 这张图要看什么:(a) 三条线是本高斯实验每步的方差——$\tilde\beta_t$(蓝)、$\beta_t$(红虚)、真实后验方差(绿)。绿线与蓝线的差异取决于时间步,较低噪声端尤其需关注,这个差值就是 3.5 节说的「均值估计不准」那一项,也是 NLL 上那 0.0001 nats 的全部来源。(b) 在 $\tilde\beta_t$ 与 $\beta_t$ 之间插值扫一遍:蓝线(左轴)是生成分布标准差与真分布之比,从 0.9659 单调爬到 0.9720;红线(右轴)是期望 NLL,几乎是一条平线。标准差之比是敏感指标,NLL 对这件事几乎不动——想判断方差选得对不对,同时检查 NLL 与协方差,不要只看一个指标。 采样成本要按同一口径计数:4000 样本 × 1000 步是 $4\times10^6$ 次样本级前向;训练 20000 步 × batch 256 是 $5.12\times10^6$ 次样本级前向,另有反向计算。不能拿样本级采样前向数除以训练优化器步数,声称采样贵 200 倍。批量大小与硬件利用率还会改变墙钟时间。 最后一个容易搞混的地方:训练目标(等权还是加权)和采样器(DDPM 随机链还是 DDIM 确定链)是两个正交的旋钮。本篇从头到尾只动前一个,采样器始终是 DDPM 原文那条随机链。换了训练目标不影响你能不能换采样器,反过来也一样——DDIM 那篇最关键的观察就是:DDPM 的训练目标根本没有约束反向链必须是一阶 Markov 的。把这两件事当成一件事,是读扩散模型文献时最普遍的混线。 07. 经典论文脉络 arXiv:1503.03585(Sohl-Dickstein et al., 2015)——扩散的雏形。 定义了前向 Markov 链和反向链,用前向扩散与变分目标训练反向过程,样本质量也远不够看。贡献是「这个方向存在」。 arXiv:2006.11239(Ho et al., 2020)——DDPM,本篇的锚点。 三件事:把 $\tilde\mu_t$ 参数化成预测 $\varepsilon$;指出 ELBO 里那串只跟 $t$ 有关的系数可以扔掉、等权反而更好;给出 $\tilde\beta_t/\tilde\mu_t$ 的闭式解并配上线性调度。这才是「扩散模型能训起来」的直接原因。 arXiv:2102.09672(Nichol & Dhariwal, 2021)——Improved DDPM。 两件事直接对着本篇的洞:一是学方差(对应 learned_range,在 $\tilde\beta_t$ 与 $\beta_t$ 之间让网络选位置,且在 log 空间插值);二是余弦调度,改善低分辨率设置下线性噪声调度过快破坏信号、后段接近纯噪声的问题。 arXiv:2010.02502(Song et al., 2020)——DDIM。 指出 DDPM 的训练目标其实没有约束反向链必须是 Markov 的,于是可以推一个非 Markov 的确定性问题,一步跨很多步。训练和采样在这里正式解耦。 arXiv:2206.00364(Karras et al., 2022)——EDM。 从「每步该加权多少」重新出发,把加权、预处理、调度统一成一套设计空间,并用二阶求解器把步数压到几十步。这是 4.5 节那个「容量变了结论会变」在真实模型上的落地版本。 08. 常见误解 ①「$L_{\text{simple}}$ 就是 ELBO。」 不是。它丢掉了两样东西:每步的 $w_t$,以及 $L_0$ 那一项的离散解码器。它是一个设计选择,不是近似——丢掉的部分在数学上并不小($w_t$ 跨 120 倍),只是恰好在容量吃紧时更划算。 ②「$\tilde\beta_t$ 是真实后验方差。」 不是,它是已知 $x_0$ 时的后验方差。采样时 $x_0$ 是猜的,猜错的那部分不确定性没算进去。实测真实后验方差在中后段比 $\tilde\beta_t$ 高一大截(图 3a 的绿线与蓝线),代价是生成分布偏窄 1.3%。 ③「权重大意味着更容易学。」 权重衡量误差在目标中的代价,条件方差衡量不可约误差;两者不是同一个量。低噪声端 $R^2$ 小不意味着该步 score 不重要。 ④「采样方差是训出来的。」 默认不是。fixed_small 是硬编码的 $\tilde\beta_t$,跟训练毫无关系;只有 learned / learned_range 才让网络参与,而且此时模型输出通道要翻倍(step() 里 model_output.shape[1] == sample.shape[1] * 2 那个判断就是干这个的)。 ⑤「换方差只影响采样,不影响训练。」 一半对。方差确实不参与 $L_{\text{simple}}$ 的计算,但 $w_t$ 的公式里有 $\sigma_t^2$——所以用真 ELBO 权重训练时,你选的方差会通过 $w_t$ 反过来改变训练目标(表 3 里 elbo_tilde 和 elbo_large 是两种不同的训练目标,不只是两种采样方式)。 ⑥「等权赢,所以加权分析没用了。」 本文两个容量档已经出现排序翻转;权重还与噪声调度、预条件、数据与优化相互作用。应在同一评测口径下选择,不能把一次玩具胜负当成普遍结论。 ⑦「$T=1000$ 是个需要调的超参。」 不完全是。$T$ 决定了 $w_t$ 的动态范围,也决定了每步要挪多远;但它同时被 $\beta$ 调度绑住——改 $T$ 不改 $\beta$ 的端点,等于改了整条噪声表的形状。Improved DDPM 换余弦调度而不是换 $T$,正是因为这两个量不能分开调。 09. 动手验证 两个脚本都在附录,只依赖 numpy: /usr/local/bin/python3 ddpm_lab.py # 闭式解实验室,几秒跑完 /usr/local/bin/python3 ddpm_train.py --steps 20000 --n-gen 4000 # 真网络版,几分钟 ddpm_lab.py 会依次打印六段(真实输出摘要,不是预期值): A 段:两种方差下的 $w_t$ 表格,解析式与化简式的最大相对差 3.93e-16;用两个高斯的 KL 数值核对,相对差 2.89e-13;并明确打印 t=1 的 tilde_beta_1 = 0.000e+00。 B 段:训练集规模(120000 条、时间基 3 项、18 个参数)与权重动态范围 120.4。 C 段:真分布熵 1.880914 nats;「最优 $\varepsilon$ 预测器 + $\tilde\beta_t$」1.881015(差 +0.000101);「最优 $\varepsilon$ 预测器 + 真实后验方差」1.880914(差 +3.56e-10)。三种目标 × 两种方差的 NLL 表。 D 段:每步超额误差剖面,低噪声档比值 0.10、高噪声档比值 2.33。 E 段:$\lambda$ 插值扫描,NLL 1.898760 → 1.895815,生成 std 比 0.9659 → 0.9720。 F 段:两档容量对照,loose 真 ELBO 赢(1.880952 vs 1.881758)、tight 等权赢(1.898760 vs 1.922295)。 ddpm_train.py 的 A 段是梯度检查(有限差分相对差 9.92e-10),C 段是三种指标打分(uniform −0.2208±0.0195 / elbo_tilde −0.4212±0.0062),D 段只换采样方差($\tilde\beta_t$ −0.2363±0.0123 / $\beta_t$ −0.2640±0.0129)。 想自己动手改的话,三个地方最值得试:把 BASIS 从 3 项换成 5 项看结论翻转(F 段已经在做);把 T 从 1000 降到 100 看 $w_t$ 的动态范围怎么变;把数据从高斯换成 ddpm_train.py 的 7 峰混合,看方差选择的结论会不会反过来。 10. 延伸阅读 前置:扩散过程的前向与反向推导(前向闭式解、反向后验、ELBO→$\varepsilon$ 的化简)、变分下界与重参数化(ELBO 这个工具本身从哪来)。 已发布后继:DDIM 与高阶采样器、CFG、潜空间扩散、流匹配。 往上一层:把 RL 用到扩散模型上(训练目标被换成奖励之后,$w_t$ 这套分析还成不成立)。 附录:完整代码 09 节用到的脚本全文如下(ddpm_lab.py、ddpm_train.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 ddpm_lab.py # -*- coding: utf-8 -*- """DDPM 训练目标的最小实验室(闭式解版,只依赖 numpy,无 torch)。 这个脚本回答四个只能用数字回答的问题: A. ELBO 里那一串「只跟 t 有关的系数」到底是什么形状。 解析式 w_t = beta_t^2 / (2 sigma_t^2 alpha_t (1 - abar_t)), 用两个高斯的 KL 数值核一遍,确认没推导错。 B. 用带权最小二乘训练一个**跨时间步共享参数**的 epsilon 网络 (仿射 + 固定时间基),比较三种目标: uniform —— DDPM 的 L_simple,等权 elbo_tilde —— 真 ELBO 权重,sigma^2 = tilde_beta_t(diffusers 的 fixed_small) elbo_large —— 真 ELBO 权重,sigma^2 = beta_t(diffusers 的 fixed_large) 模型容量故意给小(9 个时间基函数),所以权重真的会决定容量往哪搬。 C. 仿射模型下 p_theta(x_0) 是一个高斯,可以算出**精确 NLL**。 用「密度沿链前向传播」算,再用「线性映射复合」交叉验证一遍。 D. 反向链的方差 sigma^2 到底该取 tilde_beta_t 还是 beta_t。 在 beta_tilde 与 beta 之间插值扫一遍,看 NLL 与生成分布的胖瘦怎么变。 运行: /usr/local/bin/python3 ddpm_lab.py 依赖: numpy(无 GPU、无 torch) """ import os import numpy as np RNG_SEED = 20260927 T = 1000 D = 2 # ══════════════════════════════════════════════════════════════════ # 0. 噪声调度与前向过程的闭式量 # ══════════════════════════════════════════════════════════════════ def linear_beta(T=1000, b1=1e-4, bT=0.02): """DDPM 原文 CIFAR-10 用的线性调度。""" return np.linspace(b1, bT, T) def cosine_beta(T=1000, s=0.008, clip=0.999): """Improved DDPM 的余弦调度:先定 abar(t) 再反解 beta。""" u = np.arange(1, T + 1) / T f = np.cos(((u + s) / (1 + s)) * np.pi / 2) ** 2 f0 = np.cos((s / (1 + s)) * np.pi / 2) ** 2 abar = f / f0 beta = 1.0 - abar / np.concatenate([[1.0], abar[:-1]]) return np.clip(beta, 1e-8, clip) class Sched: """下标约定:abar[t]、beta[t]、bt[t] 中的 t 取 1..T;t=0 是数据本身。""" def __init__(self, beta): self.T = len(beta) self.beta = beta self.alpha = 1.0 - beta self.abar = np.concatenate([[1.0], np.cumprod(self.alpha)]) # 长度 T+1 # tilde_beta_t = (1 - abar_{t-1}) / (1 - abar_t) * beta_t self.bt = np.empty(self.T + 1) self.bt[0] = 0.0 self.bt[1:] = (1.0 - self.abar[:-1]) / (1.0 - self.abar[1:]) * beta def weight(self, which="tilde"): """ELBO 第 t 项里 ||eps - eps_theta||^2 前面的系数(长度 T,下标 0 对应 t=1)。 which="tilde" -> sigma^2 = tilde_beta_t (fixed_small) which="large" -> sigma^2 = beta_t (fixed_large) t=1 时 tilde_beta_1 = 0,权重分母为零(分子非零)。DDPM 把这一项交给 L_0(离散解码器) 单独处理;这里为了让训练能跑,退化地用 fixed_large 的权重顶上。 """ s2 = self.bt[1:] if which == "tilde" else self.beta with np.errstate(divide="ignore", invalid="ignore"): w = self.beta ** 2 / (2.0 * s2 * self.alpha * (1.0 - self.abar[1:])) if which == "tilde": w[0] = self.beta[0] / (2.0 * self.alpha[0] * (1.0 - self.abar[1])) return w def weight_simplified(self, which="tilde"): """上面那条式子化简之后的样子,用来核对代数没推错。""" denom = 1.0 - (self.abar[:-1] if which == "tilde" else self.abar[1:]) with np.errstate(divide="ignore", invalid="ignore"): return self.beta / (2.0 * self.alpha * denom) def q_sample(x0, t_idx, abar, rng): """x_t = sqrt(abar_t) x_0 + sqrt(1 - abar_t) eps,一步到位。""" a = abar[t_idx] if np.isscalar(a): a = np.full(len(x0), a) eps = rng.standard_normal(x0.shape) return np.sqrt(a)[:, None] * x0 + np.sqrt(1.0 - a)[:, None] * eps, eps # ══════════════════════════════════════════════════════════════════ # 1. 数据:一个二维高斯(NLL 能算到机器精度) # ══════════════════════════════════════════════════════════════════ MU0 = np.array([1.0, -0.5]) S0 = np.array([[0.60, 0.25], [0.25, 0.35]]) def sample_gauss(n, rng): L = np.linalg.cholesky(S0) return MU0 + rng.standard_normal((n, D)) @ L.T def entropy_gauss(S): """高斯熵(nats)。""" _, ld = np.linalg.slogdet(S) return 0.5 * (D * np.log(2 * np.pi * np.e) + ld) def kl_gauss(m0, S0_, m1, S1_): """KL( N(m0, S0_) || N(m1, S1_) ),闭式解。""" L = np.linalg.cholesky(S1_) y = np.linalg.solve(L, m1 - m0) Sinv_S0 = np.linalg.solve(S1_, S0_) _, ld1 = np.linalg.slogdet(S1_) _, ld0 = np.linalg.slogdet(S0_) return 0.5 * (np.trace(Sinv_S0) + y @ y - D + ld1 - ld0) def expected_nll(m, S): """E_{x ~ p_data}[-log p_theta(x)] 的**精确值**。 = 数据分布的熵 + KL(p_data || p_theta),两个高斯之间全是闭式解, 没有蒙特卡洛噪声——用固定测试集估 NLL 时,2 万样本的波动有 ±0.007 nats, 比我们要比的 0.002~0.004 nats 还大,所以这里必须用闭式解。 """ return entropy_gauss(S0) + kl_gauss(MU0, S0, m, S) def gauss_nll(X, m, S): """多元高斯负对数似然(nats)。协方差病态时加抖动兜底。""" jitter = 0.0 for _ in range(8): try: L = np.linalg.cholesky(S + jitter * np.eye(D)) y = np.linalg.solve(L, (X - m).T) return 0.5 * (np.sum(y * y, axis=0) + 2 * np.sum(np.log(np.diag(L))) + D * np.log(2 * np.pi)) except np.linalg.LinAlgError: jitter = 1e-12 if jitter == 0.0 else jitter * 100.0 raise RuntimeError("协方差矩阵无法正定化,说明反向链已经数值发散") # ══════════════════════════════════════════════════════════════════ # 2. 模型:eps_theta(x, t) = B(psi_t) x + c(psi_t),参数跨 t 共享 # psi 是固定的时间基(Fourier),只有 B、c 的系数是可学的 # ══════════════════════════════════════════════════════════════════ # 时间基故意取得很小:让「容量」成为真正的约束,权重才会决定容量往哪搬。 # 指数里出现 0.5,是因为最优的 B_t 与 sqrt(1 - abar_t) 成正比, # 在 t 很小时它按 sqrt(t) 走 —— 用纯多项式去拟合会在开头剧烈震荡。 BASIS = { "loose": [0.0, 0.5, 1.0, 1.5, 2.0], # 5 项:容量基本够用 "tight": [0.0, 0.5, 1.0], # 3 项:容量真的成了瓶颈 } BASIS_EXP = BASIS["tight"] # 默认用容量吃紧那一档,机制看得最清楚 K = len(BASIS_EXP) def set_basis(name): global BASIS_EXP, K BASIS_EXP = BASIS[name] K = len(BASIS_EXP) return K def psi_basis(t_arr): """t (1..T) -> (n, K) 的固定时间基,u = t / T in (0, 1]。""" u = np.asarray(t_arr, dtype=float) / T return np.stack([u ** e for e in BASIS_EXP], axis=-1) def feats(x, psi): """拼特征:[x ⊗ psi, psi],最后一维长度 D*K + K。""" px = x[:, :, None] * psi[:, None, :] # (n, D, K) return np.concatenate([px.reshape(len(x), -1), psi], axis=1) def fit_weighted_ridge(Phi, E, w, lam): """最小化 sum_i w_i ||W phi_i - eps_i||^2 + lam ||W||^2。""" P = Phi.shape[1] A = np.einsum("n,np,nq->pq", w, Phi, Phi) + lam * np.eye(P) Bmat = np.einsum("n,np,nd->pd", w, Phi, E) return np.linalg.solve(A, Bmat) def model_Bc(W, psi_t): """把共享权重在某个时刻 t 上展开成 eps = B x + c。""" Wx = W[: D * K, :].reshape(D, K, D) # Wx[i, k, j] Bmat = np.einsum("ikj,k->ji", Wx, psi_t) cvec = psi_t @ W[D * K:, :] return Bmat, cvec def optimal_Bc(sc, t): """高斯数据下 eps 的最优预测器(后验均值),用来做参考与验算。""" a = sc.abar[t] C = a * S0 + (1.0 - a) * np.eye(D) Bmat = np.sqrt(1.0 - a) * np.linalg.inv(C) cvec = -np.sqrt(a) * (Bmat @ MU0) return Bmat, cvec def mse_analytic(Bmat, cvec, sc, t): """E||B x_t + c - eps||^2 的闭式解(对 D 维取了平均)。 (x_t, eps) 是联合高斯:Var(x_t) = C_t = abar_t S0 + (1-abar_t) I, Cov(x_t, eps) = sqrt(1-abar_t) I,E[x_t] = sqrt(abar_t) mu0,E[eps] = 0。 展开 ||B x_t + c - eps||^2 的期望即得下式,不用蒙特卡洛,没有抽样噪声。 """ a = sc.abar[t] C = a * S0 + (1.0 - a) * np.eye(D) m = np.sqrt(a) * MU0 quad = np.trace(Bmat @ C @ Bmat.T) # Var(B x_t) bias = np.sum((Bmat @ m + cvec) ** 2) # 均值没对上的部分 cross = 2.0 * np.sqrt(1.0 - a) * np.trace(Bmat) # -2 Cov(B x_t, eps) return float((quad + bias - cross + D) / D) def mse_floor(sc, t): """贝叶斯最优的 MSE 地板:Var(eps | x_t) 的迹除以 D。""" a = sc.abar[t] C = a * S0 + (1.0 - a) * np.eye(D) return float((D - (1.0 - a) * np.trace(np.linalg.inv(C))) / D) def learnability(sc, t): """这一时刻的 eps 里,有多大比例是能从 x_t 里看出来的。 R^2 = 1 - 地板 / Var(eps),Var(eps) 的每维是 1。 """ return 1.0 - mse_floor(sc, t) # ══════════════════════════════════════════════════════════════════ # 3. 精确 NLL:把高斯密度沿反向链往前传 # ══════════════════════════════════════════════════════════════════ def _s2_mat(sigma2, t): """每步注入的噪声协方差:既接受标量(各向同性),也接受 (T+1, D, D)。""" if sigma2.ndim == 1: return sigma2[t] * np.eye(D) return sigma2[t] def chain_moments_prop(sc, Bcfun, sigma2): """写法一:密度传播。x_{t-1} = M_t x_t + v_t + noise。""" m = np.zeros(D) S = np.eye(D) sig = np.atleast_1d(sigma2) for t in range(T, 0, -1): Bmat, cvec = Bcfun(t) kk = sc.beta[t - 1] / np.sqrt(1.0 - sc.abar[t]) M = (np.eye(D) - kk * Bmat) / np.sqrt(sc.alpha[t - 1]) v = -kk * cvec / np.sqrt(sc.alpha[t - 1]) m = M @ m + v S = M @ S @ M.T + _s2_mat(sig, t) return m, S def chain_moments_comp(sc, Bcfun, sigma2): """写法二:把整条链复合成一个仿射映射,再累加各步噪声。独立实现,用于交叉验证。""" m = np.zeros(D) S = np.zeros((D, D)) Pprev = np.eye(D) for t in range(1, T + 1): Bmat, cvec = Bcfun(t) kk = sc.beta[t - 1] / np.sqrt(1.0 - sc.abar[t]) M = (np.eye(D) - kk * Bmat) / np.sqrt(sc.alpha[t - 1]) v = -kk * cvec / np.sqrt(sc.alpha[t - 1]) m = m + Pprev @ v S = S + Pprev @ _s2_mat(sigma2, t) @ Pprev.T Pprev = Pprev @ M S = S + Pprev @ Pprev.T # x_T ~ N(0, I) 的那一坨 return m, S def sigma2_exact_posterior(sc): """真实后验 q(x_{t-1}|x_t) 的协方差(高斯数据下可算)。 = tilde_beta_t I + Var(mu_t | x_t) = tilde_beta_t I + (beta_t^2 / (alpha_t (1 - abar_t))) Var(eps | x_t) 第二项就是 DDPM 反向链扔掉的那部分:真实后验比 beta_tilde 更胖。 """ out = np.zeros((T + 1, D, D)) for t in range(1, T + 1): a = sc.abar[t] C = a * S0 + (1.0 - a) * np.eye(D) var_eps = np.eye(D) - (1.0 - a) * np.linalg.inv(C) kk2 = sc.beta[t - 1] ** 2 / (sc.alpha[t - 1] * (1.0 - a)) out[t] = sc.bt[t] * np.eye(D) + kk2 * var_eps return out def sigma2_from_choice(sc, which="tilde", lam_mix=0.0): """反向链每步注入的方差。lam_mix 在 tilde_beta 与 beta 之间插值。""" if which == "tilde": base = sc.bt.copy() other = np.concatenate([[0.0], sc.beta]) else: base = np.concatenate([[0.0], sc.beta]) other = sc.bt.copy() return (1.0 - lam_mix) * base + lam_mix * other # ══════════════════════════════════════════════════════════════════ # 4. 主流程 # ══════════════════════════════════════════════════════════════════ def section_A(sc): print("=" * 74) print("A. ELBO 每步权重:解析式 vs 两个高斯的 KL") print("=" * 74) w_tilde = sc.weight("tilde") w_large = sc.weight("large") ws_tilde = sc.weight_simplified("tilde") ws_large = sc.weight_simplified("large") ok = np.isfinite(ws_tilde) & (ws_tilde > 0) rel = np.max(np.abs(w_tilde[ok] - ws_tilde[ok]) / ws_tilde[ok]) print(f" 化简式与原始式的最大相对差(tilde):{rel:.3e}") rel2 = np.max(np.abs(w_large - ws_large) / ws_large) print(f" 化简式与原始式的最大相对差(large):{rel2:.3e}") print(f" t=1 的 tilde_beta_1 = {sc.bt[1]:.3e}(后验方差为 0 => 权重发散," f"这就是 DDPM 要把 L_0 单独拿出来的原因)") print() print(" t abar_t beta_t w_t(sigma^2=bt) w_t(sigma^2=beta) 相对等权") for t in [2, 10, 50, 100, 250, 500, 750, 900, 1000]: print(f" {t:<6d}{sc.abar[t]:<12.3e}{sc.beta[t-1]:<10.5f}" f"{w_tilde[t-1]:<18.4e}{w_large[t-1]:<20.4e}" f"{w_tilde[t-1] / w_tilde[499]:<10.2f}") print() span = w_tilde[1] / w_tilde.min() print(f" sigma^2=tilde_beta 时,权重最大值(t=2)是最小值(t={int(np.argmin(w_tilde))+1})的 " f"{span:.1f} 倍") print(f" 等权 L_simple 相当于把这条曲线整体压平到 1") print() # 数值核对:KL( q(x_{t-1}|x_t,x_0) || p_theta(x_{t-1}|x_t) ) == w_t * ||delta_eps||^2 rng = np.random.default_rng(RNG_SEED) print(" —— KL 数值核对(随机取 x_0, eps,令 eps_theta = eps + delta)——") print(" t KL(数值) w_t*||delta||^2 相对差") for t in [2, 100, 500, 1000]: x0 = sample_gauss(1, rng)[0] a = sc.abar[t] eps = rng.standard_normal(D) xt = np.sqrt(a) * x0 + np.sqrt(1 - a) * eps delta = rng.standard_normal(D) * 0.3 eps_hat = eps + delta kk = sc.beta[t - 1] / np.sqrt(1 - a) mu_tilde = (xt - kk * eps) / np.sqrt(sc.alpha[t - 1]) mu_theta = (xt - kk * eps_hat) / np.sqrt(sc.alpha[t - 1]) s2 = sc.bt[t] kl = np.sum((mu_tilde - mu_theta) ** 2) / (2 * s2) pred = w_tilde[t - 1] * np.sum(delta ** 2) print(f" {t:<7d}{kl:<18.6e}{pred:<20.6e}{abs(kl-pred)/kl:<12.2e}") print() return w_tilde, w_large def fit_three(sc, w_tilde, w_large, lam=1e-6): """三种权重各解一遍带权最小二乘。不打印任何东西,供 section_B / section_F 共用。""" rng = np.random.default_rng(RNG_SEED + 1) M = 120_000 t_idx = rng.integers(1, T + 1, size=M) x0 = sample_gauss(M, rng) a = sc.abar[t_idx] eps = rng.standard_normal((M, D)) xt = np.sqrt(a)[:, None] * x0 + np.sqrt(1 - a)[:, None] * eps psi = psi_basis(t_idx) Phi = feats(xt, psi) # 三种权重都归一化到均值 1(整体缩放不改变无正则时的解, # 归一化只是为了让正则强度在三种目标下可比) schemes = { "uniform": np.ones(T), "elbo_tilde": w_tilde / w_tilde.mean(), "elbo_large": w_large / w_large.mean(), } lam = 1e-6 W = {} for name, w_full in schemes.items(): w = w_full[t_idx - 1] W[name] = fit_weighted_ridge(Phi, eps, w, lam) return W, schemes def section_B(sc, w_tilde, w_large): print("=" * 74) print("B. 带权最小二乘训练:三种目标,共享参数,容量有限") print("=" * 74) W, schemes = fit_three(sc, w_tilde, w_large) print(f" 训练集:M=120000 条 (x_0, t, eps),时间基 {K} 项,特征维度 {K * (D + 1)}," f"可学参数 {K * (D + 1) * D} 个") print(f" 三种权重都归一化到均值 1(整体缩放不改变解,归一化只为让正则强度可比)") print(f" 权重动态范围:elbo_tilde 的 max/min = {w_tilde.max()/w_tilde.min():.1f}") print() return W, schemes def model_moments(sc, Wm, which="tilde"): """把训练好的共享权重展开成整条反向链→后验的均值与协方差。""" s2 = sigma2_from_choice(sc, which) return chain_moments_prop( sc, lambda t: model_Bc(Wm, psi_basis(np.array([t]))[0]), s2) def excess_bands(sc, W): """低噪声档 / 高噪声档的超额误差均值,以及 elbo 相对 uniform 的比值。""" lo = np.arange(1, 51) hi = np.arange(400, 1001) out = {} for band, idx in [("lo", lo), ("hi", hi)]: e = {} for name in ["uniform", "elbo_tilde"]: e[name] = float(np.mean([ mse_analytic(*model_Bc(W[name], psi_basis(np.array([t]))[0]), sc, t) - mse_floor(sc, t) for t in idx])) out[band] = (e["uniform"], e["elbo_tilde"], e["elbo_tilde"] / e["uniform"]) return out def section_C(sc, W, schemes): print("=" * 74) print("C. 仿射模型下的精确 NLL(密度前向传播)") print("=" * 74) rng = np.random.default_rng(RNG_SEED + 2) Xtest = sample_gauss(400_000, rng) nll_true = expected_nll(MU0, S0) print(f" 参考:真分布 N(mu_0, S_0) 的熵 = {nll_true:.6f} nats" f"({nll_true / np.log(2) / D:.4f} bits/dim)") print(f" 蒙特卡洛复核(40 万样本):{gauss_nll(Xtest, MU0, S0).mean():.6f} nats" f",与闭式解差 {abs(gauss_nll(Xtest, MU0, S0).mean()-nll_true):.2e}") # 交叉验证两种写法 m1, S1 = chain_moments_prop(sc, lambda t: optimal_Bc(sc, t), sigma2_from_choice(sc, "tilde")) m2, S2 = chain_moments_comp(sc, lambda t: optimal_Bc(sc, t), sigma2_from_choice(sc, "tilde")) print(f" 两种写法的差异:均值 {np.max(np.abs(m1-m2)):.2e},协方差 {np.max(np.abs(S1-S2)):.2e}") # 蒙特卡洛复核:真跑 20 万条链,比对经验均值/协方差 rng3 = np.random.default_rng(RNG_SEED + 3) xs = rng3.standard_normal((200_000, D)) for t in range(T, 0, -1): Bmat, cvec = optimal_Bc(sc, t) kk = sc.beta[t - 1] / np.sqrt(1.0 - sc.abar[t]) mu = (xs - kk * (xs @ Bmat.T + cvec)) / np.sqrt(sc.alpha[t - 1]) xs = mu + np.sqrt(sc.bt[t]) * rng3.standard_normal((200_000, D)) print(f" 蒙特卡洛复核:均值差 {np.max(np.abs(xs.mean(0)-m1)):.2e}," f"协方差差 {np.max(np.abs(np.cov(xs.T)-S1)):.2e}") nll_opt = expected_nll(m1, S1) print(f" 「最优 eps 预测器 + sigma^2=tilde_beta」的期望 NLL = {nll_opt:.6f} nats," f"比真分布差 {nll_opt - nll_true:+.6f} nats") print(f" 蒙特卡洛复核(同上 40 万样本):{gauss_nll(Xtest, m1, S1).mean():.6f} nats," f"差 {abs(gauss_nll(Xtest, m1, S1).mean()-nll_opt):.2e}") print(f" 生成分布的协方差 diag = {np.diag(S1)},真值 diag = {np.diag(S0)} " f"=> 比值 {np.diag(S1)/np.diag(S0)}") # 把真实后验方差(而不是 tilde_beta)灌回反向链:应该精确还原数据分布 s2_exact = sigma2_exact_posterior(sc) m3, S3 = chain_moments_prop(sc, lambda t: optimal_Bc(sc, t), s2_exact) nll_exact = expected_nll(m3, S3) print(f" 「最优 eps 预测器 + 真实后验方差」的期望 NLL = {nll_exact:.6f} nats," f"比真分布差 {nll_exact - nll_true:+.2e} nats") print(f" 此时生成协方差 diag / 真值 = {np.diag(S3)/np.diag(S0)}," f"均值差 {np.max(np.abs(m3 - MU0)):.2e}") print(" => reverse 链的均值用最优 eps 预测器、方差用真实后验方差," "就能精确还原数据分布;") print(" NLL 上剩下的那 0.0001 nats 完全是「把方差钉死成 tilde_beta」造成的。") print() print(" 三种训练目标 × 两种采样方差的期望 NLL(nats,越小越好):") print(" 训练目标 sigma^2=tilde_beta sigma^2=beta 差 生成std/真std") rows = {} for name in ["uniform", "elbo_tilde", "elbo_large"]: Wm = W[name] out = {} for which in ["tilde", "large"]: m, S = model_moments(sc, Wm, which) out[which] = expected_nll(m, S) out[which + "_std"] = float(np.mean(np.sqrt(np.diag(S)) / np.sqrt(np.diag(S0)))) rows[name] = out print(f" {name:<16}{out['tilde']:<22.6f}{out['large']:<18.6f}" f"{out['large']-out['tilde']:<+12.6f}{out['tilde_std']:.4f}") print() return rows, nll_true, nll_opt def section_D(sc, W): print("=" * 74) print("D. 每步误差剖面:权重把「超额误差」搬到了哪") print("=" * 74) grid = list(range(1, 1001)) floor = np.array([mse_floor(sc, t) for t in grid]) learn = np.array([learnability(sc, t) for t in grid]) prof = {} for name, Wm in W.items(): mse = [] for t in grid: Bmat, cvec = model_Bc(Wm, psi_basis(np.array([t]))[0]) mse.append(mse_analytic(Bmat, cvec, sc, t)) prof[name] = np.array(mse) print(" 闭式解算的,没有蒙特卡洛噪声。『超额』= MSE − 贝叶斯地板。") print(" t 可学占比R^2 地板 uniform超额 elbo_tilde超额 elbo_large超额") for t in [1, 5, 20, 50, 100, 200, 400, 600, 800, 950, 1000]: i = t - 1 print(f" {t:<6d}{learn[i]:<12.4f}{floor[i]:<9.4f}" f"{prof['uniform'][i]-floor[i]:<14.3e}{prof['elbo_tilde'][i]-floor[i]:<16.3e}" f"{prof['elbo_large'][i]-floor[i]:.3e}") print() lo = slice(0, 50) # t = 1..50 hi = slice(399, 1000) # t = 400..1000 eu_u, eu_e = prof["uniform"] - floor, prof["elbo_tilde"] - floor print(f" 低噪声档 t<=50 :uniform 超额均值 {eu_u[lo].mean():.3e}," f"elbo_tilde {eu_e[lo].mean():.3e}(比值 {eu_e[lo].mean()/eu_u[lo].mean():.2f})") print(f" 高噪声档 t>=400 :uniform 超额均值 {eu_u[hi].mean():.3e}," f"elbo_tilde {eu_e[hi].mean():.3e}(比值 {eu_e[hi].mean()/eu_u[hi].mean():.2f})") print() print(f" 可学占比 R^2:t=1 时 {learn[0]:.2e},t=500 时 {learn[499]:.4f}," f"t=1000 时 {learn[999]:.4f}") print() return grid, prof, floor, learn def section_E(sc, W): print("=" * 74) print("E. 采样方差该取 tilde_beta 还是 beta:在两者之间插值扫一遍") print("=" * 74) nll_true = expected_nll(MU0, S0) print(" lam 是插值系数:sigma^2 = (1-lam)*tilde_beta + lam*beta") print(" lam NLL(uniform训练) NLL(elbo_tilde训练) 生成 std / 真 std") out = {"lam": [], "nll_uniform": [], "nll_elbo": [], "std_ratio": []} for lam in [0.0, 0.25, 0.5, 0.75, 1.0]: s2 = sigma2_from_choice(sc, "tilde", lam_mix=lam) row, sr = [], 0.0 for name in ["uniform", "elbo_tilde"]: m, S = chain_moments_prop( sc, lambda t: model_Bc(W[name], psi_basis(np.array([t]))[0]), s2) row.append(expected_nll(m, S)) if name == "uniform": sr = float(np.mean(np.sqrt(np.diag(S)) / np.sqrt(np.diag(S0)))) out["lam"].append(lam) out["nll_uniform"].append(row[0]) out["nll_elbo"].append(row[1]) out["std_ratio"].append(sr) print(f" {lam:<7.2f}{row[0]:<19.6f}{row[1]:<22.6f}{sr:.4f}") best = out["lam"][int(np.argmin(np.abs(np.array(out["std_ratio"]) - 1.0)))] print(f" 生成分布的胖瘦刚好对上真分布的 lam ≈ {best:.2f}(这一列是最敏感的指标," f"NLL 对 lam 几乎不动)") print(f" 参考:真分布熵 = {nll_true:.6f}") print() return out def section_F(sc, w_tilde, w_large): """容量充足 vs 容量吃紧:『该不该扔掉系数』的答案会不会反过来。""" print("=" * 74) print("F. 同样两个目标,换一档模型容量:结论会反过来") print("=" * 74) print(" 两档时间基都只动特征个数,训练数据、种子、正则强度完全一样。") print() print(" 容量档 时间基项数 可学参数 NLL(等权) NLL(真ELBO) 谁更好 低噪声档超额比 高噪声档超额比") out = {} for regime in ["loose", "tight"]: k = set_basis(regime) Wf, _ = fit_three(sc, w_tilde, w_large) n_u = expected_nll(*model_moments(sc, Wf["uniform"])) n_e = expected_nll(*model_moments(sc, Wf["elbo_tilde"])) bands = excess_bands(sc, Wf) better = "等权" if n_u < n_e else "真ELBO" n_par = k * (D + 1) * D out[regime] = dict(K=k, nll_uniform=n_u, nll_elbo=n_e, bands=bands, W=Wf) print(f" {regime:<8}{k:<12}{n_par:<11}{n_u:<12.6f}{n_e:<14.6f}" f"{better:<9}{bands['lo'][2]:<16.2f}{bands['hi'][2]:.2f}") set_basis("tight") print() print(" 列『超额比』= 真 ELBO 权重的超额误差 / 等权的超额误差,小于 1 表示更好。") print(" => 容量够用时真 ELBO 权重略胜,容量真的吃紧时它反而输给等权。") print() return out def main(): sc = Sched(linear_beta(T)) print(f"调度:线性 T={T},beta 从 {sc.beta[0]:.1e} 到 {sc.beta[-1]:.3f}," f"abar_T = {sc.abar[T]:.4e}") print() w_tilde, w_large = section_A(sc) W, schemes = section_B(sc, w_tilde, w_large) print() rows, nll_true, nll_opt = section_C(sc, W, schemes) print() grid, prof, floor, learn = section_D(sc, W) print() mix = section_E(sc, W) print() regimes = section_F(sc, w_tilde, w_large) return dict(sched=sc, w_tilde=w_tilde, w_large=w_large, W=W, nll_rows=rows, nll_true=nll_true, nll_opt=nll_opt, grid=grid, prof=prof, floor=floor, learn=learn, mix=mix, regimes=regimes) if __name__ == "__main__": main() ddpm_train.py # -*- coding: utf-8 -*- """照着 DDPM 原文 Algorithm 1 / Algorithm 2 写的最小实现(纯 numpy,无 torch)。 为什么还要写一遍神经网络版:ddpm_lab.py 用带权最小二乘的闭式解把「优化算法」 这个变量消掉了,代价是模型只能仿射。这里补上手写反向传播的两层 MLP, 在**多模态**的目标分布上跑完整的一千步采样,看训练目标的权重到底怎么影响 最后生成出来的东西。 A. 梯度核对:手写反向传播 vs 有限差分 B. Algorithm 1:按三种加权(等权 / ELBO / ELBO-large)训练三个网络 C. Algorithm 2:一千步采样,三种指标打分 D. 反向链方差 sigma^2 = tilde_beta 还是 beta 运行: /usr/local/bin/python3 ddpm_train.py # 默认 20000 步 /usr/local/bin/python3 ddpm_train.py --steps 40000 依赖: numpy """ import argparse import os import numpy as np RNG_SEED = 20260927 T = 1000 D = 2 # ══════════════════════════════════════════════════════════════════ # 0. 调度(与 ddpm_lab.py 同一套,下标约定:t 取 1..T) # ══════════════════════════════════════════════════════════════════ def linear_beta(T=1000, b1=1e-4, bT=0.02): return np.linspace(b1, bT, T) class Sched: def __init__(self, beta): self.beta = beta self.alpha = 1.0 - beta self.abar = np.concatenate([[1.0], np.cumprod(self.alpha)]) self.bt = np.empty(len(beta) + 1) self.bt[0] = 0.0 self.bt[1:] = (1.0 - self.abar[:-1]) / (1.0 - self.abar[1:]) * beta def weight(self, which="tilde"): """ELBO 第 t 项里 ||eps - eps_theta||^2 的系数。""" s2 = self.bt[1:] if which == "tilde" else self.beta with np.errstate(divide="ignore", invalid="ignore"): w = self.beta ** 2 / (2.0 * s2 * self.alpha * (1.0 - self.abar[1:])) if which == "tilde": # t=1 的 0/0 退化,用 fixed_large 顶上 w[0] = self.beta[0] / (2.0 * self.alpha[0] * (1.0 - self.abar[1])) return w # ══════════════════════════════════════════════════════════════════ # 1. 目标分布:7 个分量的二维高斯混合(环状 + 一个中心) # ══════════════════════════════════════════════════════════════════ RING_R = 1.15 N_RING = 6 CENTERS = np.array( [[RING_R * np.cos(2 * np.pi * k / N_RING), RING_R * np.sin(2 * np.pi * k / N_RING)] for k in range(N_RING)] + [[0.0, 0.0]] ) SCALES = np.array([0.17, 0.20, 0.16, 0.22, 0.18, 0.19, 0.28]) WEIGHTS = SCALES ** 2 # 分量越大越容易被采到,制造不均匀 WEIGHTS = WEIGHTS / WEIGHTS.sum() def sample_mix(n, rng): kk = rng.choice(len(WEIGHTS), size=n, p=WEIGHTS) return CENTERS[kk] + SCALES[kk][:, None] * rng.standard_normal((n, D)) def logp_mix(X): """混合分布在 X 处的对数密度(各分量都是各向同性高斯)。""" d2 = ((X[:, None, :] - CENTERS[None, :, :]) ** 2).sum(-1) # (n, K) comp = np.log(WEIGHTS)[None, :] - 0.5 * d2 / SCALES[None, :] ** 2 \ - D * np.log(SCALES[None, :]) - 0.5 * D * np.log(2 * np.pi) m = comp.max(axis=1, keepdims=True) return (m[:, 0] + np.log(np.exp(comp - m).sum(axis=1))) def nearest_center_dist(X): d2 = ((X[:, None, :] - CENTERS[None, :, :]) ** 2).sum(-1) return np.sqrt(d2.min(axis=1)) # ══════════════════════════════════════════════════════════════════ # 2. 时间嵌入 + 两层 MLP + 手写 Adam # ══════════════════════════════════════════════════════════════════ TIME_FREQ = (1, 2, 4, 8) # 8 维时间基 def time_embed(t_arr): u = np.asarray(t_arr, dtype=float) / T cols = [] for w in TIME_FREQ: cols.append(np.cos(w * np.pi * u)) cols.append(np.sin(w * np.pi * u)) return np.stack(cols, axis=-1) N_IN = D + 2 * len(TIME_FREQ) H = 64 def init_params(rng): def he(fan_in, fan_out): return rng.standard_normal((fan_in, fan_out)) * np.sqrt(2.0 / fan_in) return { "W1": he(N_IN, H), "b1": np.zeros(H), "W2": he(H, H), "b2": np.zeros(H), "W3": he(H, D), "b3": np.zeros(D), } def forward(p, Z): h1 = np.maximum(Z @ p["W1"] + p["b1"], 0.0) h2 = np.maximum(h1 @ p["W2"] + p["b2"], 0.0) return h1, h2, h2 @ p["W3"] + p["b3"] def backward(p, Z, h1, h2, out, eps, w): """d/d(theta) of mean_i w_i ||out_i - eps_i||^2。""" n = len(Z) dout = 2.0 * (w[:, None] * (out - eps)) / n g = {} g["W3"] = h2.T @ dout g["b3"] = dout.sum(0) dh2 = dout @ p["W3"].T dh2[h2 <= 0] = 0.0 g["W2"] = h1.T @ dh2 g["b2"] = dh2.sum(0) dh1 = dh2 @ p["W2"].T dh1[h1 <= 0] = 0.0 g["W1"] = Z.T @ dh1 g["b1"] = dh1.sum(0) return g class Adam: def __init__(self, p, lr=1e-3, b1=0.9, b2=0.999, eps=1e-8): self.m = {k: np.zeros_like(v) for k, v in p.items()} self.v = {k: np.zeros_like(v) for k, v in p.items()} self.lr, self.b1, self.b2, self.eps = lr, b1, b2, eps self.i = 0 def step(self, p, g): self.i += 1 for k in p: self.m[k] = self.b1 * self.m[k] + (1 - self.b1) * g[k] self.v[k] = self.b2 * self.v[k] + (1 - self.b2) * g[k] ** 2 mh = self.m[k] / (1 - self.b1 ** self.i) vh = self.v[k] / (1 - self.b2 ** self.i) p[k] -= self.lr * mh / (np.sqrt(vh) + self.eps) def predict_eps(p, x, t_idx): Z = np.concatenate([x, time_embed(np.full(len(x), t_idx))], axis=1) _, _, out = forward(p, Z) return out # ══════════════════════════════════════════════════════════════════ # 3. Algorithm 1(训练)与 Algorithm 2(采样) # ══════════════════════════════════════════════════════════════════ def train(sc, w_full, steps, batch, rng, lr=1e-3): """DDPM 原文 Algorithm 1,一行不差地照抄。""" p = init_params(rng) opt = Adam(p, lr=lr) for _ in range(steps): # 1: t ~ Uniform({1, ..., T}) 2: x_0 ~ q(x_0) 3: eps ~ N(0, I) t_idx = rng.integers(1, T + 1, size=batch) x0 = sample_mix(batch, rng) eps = rng.standard_normal((batch, D)) # 4: 一步加噪(闭式解,不需要真的走 t 步) a = sc.abar[t_idx] xt = np.sqrt(a)[:, None] * x0 + np.sqrt(1.0 - a)[:, None] * eps # 5: 梯度下降一步 Z = np.concatenate([xt, time_embed(t_idx)], axis=1) h1, h2, out = forward(p, Z) opt.step(p, backward(p, Z, h1, h2, out, eps, w_full[t_idx - 1])) return p def sample_chain(sc, p, n, which="tilde", rng=None): """DDPM 原文 Algorithm 2。which 决定 sigma^2 取 tilde_beta 还是 beta。""" x = rng.standard_normal((n, D)) # x_T ~ N(0, I) for t in range(T, 0, -1): eps_hat = predict_eps(p, x, t) a, a_prev = sc.abar[t], sc.abar[t - 1] alpha_t, beta_t = sc.alpha[t - 1], sc.beta[t - 1] # mu_tilde = (x_t - beta_t/sqrt(1-abar_t) * eps_hat) / sqrt(alpha_t) mu = (x - beta_t / np.sqrt(1.0 - a) * eps_hat) / np.sqrt(alpha_t) if which == "tilde": sigma = np.sqrt(sc.bt[t]) else: sigma = np.sqrt(beta_t) z = rng.standard_normal((n, D)) x = mu + (sigma * z if t > 1 else 0.0) # t == 1 时不加噪声 return x # ══════════════════════════════════════════════════════════════════ # 4. 主流程 # ══════════════════════════════════════════════════════════════════ def section_A(): print("=" * 74) print("A. 手写反向传播 vs 有限差分") print("=" * 74) rng = np.random.default_rng(RNG_SEED) sc = Sched(linear_beta(T)) p = init_params(rng) n = 32 t_idx = rng.integers(1, T + 1, size=n) x0 = sample_mix(n, rng) eps = rng.standard_normal((n, D)) a = sc.abar[t_idx] xt = np.sqrt(a)[:, None] * x0 + np.sqrt(1 - a)[:, None] * eps Z = np.concatenate([xt, time_embed(t_idx)], axis=1) w = np.ones(n) h1, h2, out = forward(p, Z) g = backward(p, Z, h1, h2, out, eps, w) def loss(): _, _, o = forward(p, Z) return float(np.mean(np.sum((o - eps) ** 2, axis=1))) print(" 参数 解析梯度 有限差分 相对差") worst = 0.0 for key in ["W1", "b1", "W2", "b2", "W3", "b3"]: idx = tuple(0 for _ in p[key].shape) h = 1e-6 orig = p[key][idx] p[key][idx] = orig + h lp = loss() p[key][idx] = orig - h lm = loss() p[key][idx] = orig num = (lp - lm) / (2 * h) ana = g[key][idx] rel = abs(num - ana) / max(abs(num), 1e-12) worst = max(worst, rel) print(f" {key:<8}{ana:<16.8e}{num:<16.8e}{rel:.2e}") print(f" => 最大相对差 {worst:.2e}(有限差分自己的精度极限在 1e-6 量级)") print() return sc def section_B(sc, steps): print("=" * 74) print("B. Algorithm 1:三种加权各训一个网络") print("=" * 74) w_tilde, w_large = sc.weight("tilde"), sc.weight("large") schemes = { "uniform": np.ones(T), "elbo_tilde": w_tilde / w_tilde.mean(), } print(f" 步数 {steps},batch 256,两层 MLP({N_IN} -> {H} -> {H} -> {D}),Adam lr=1e-3") print(f" elbo 权重的动态范围:max/min = {w_tilde.max()/w_tilde.min():.1f}") models = {} # 注意:不能用 hash(name)——Python 的字符串 hash 每次进程都变,结果会不可复现。 seed_of = {"uniform": 11, "elbo_tilde": 22, "elbo_large": 33} for name, w in schemes.items(): rng = np.random.default_rng(RNG_SEED + seed_of.get(name, 44)) models[name] = train(sc, w, steps, 256, rng) print(f" [{name}] 训练完成") print() return models def _score(X, base_logp, base_dist): """三个指标:log p0 相对真样本、到最近模式中心的距离比、模式覆盖数。""" dlogp = logp_mix(X).mean() - base_logp dratio = nearest_center_dist(X).mean() / base_dist assign = ((X[:, None, :] - CENTERS[None, :, :]) ** 2).sum(-1).argmin(1) frac = np.bincount(assign, minlength=len(CENTERS)) / len(X) return dlogp, dratio, int((frac > 0.02).sum()) def section_C(sc, models, n_gen, n_seed=3): print("=" * 74) print("C. Algorithm 2:一千步采样,三个指标打分") print("=" * 74) rng0 = np.random.default_rng(RNG_SEED + 77) Xtrue = sample_mix(20_000, rng0) base_logp = logp_mix(Xtrue).mean() base_dist = nearest_center_dist(Xtrue).mean() print(f" 基线(2 万真样本):log p0 均值 {base_logp:.4f}," f"到最近模式中心距离均值 {base_dist:.4f}") print(f" 每组用 {n_seed} 个不同的采样种子重复,报告均值 ± 标准差") print() print(" 训练目标 log p0 相对真样本 到最近中心距离/真样本 模式覆盖") res = {} for name, p in models.items(): dl, dr, cov = [], [], [] for s in range(n_seed): rng = np.random.default_rng(RNG_SEED + 88 + s) a, b, c = _score(sample_chain(sc, p, n_gen, "tilde", rng), base_logp, base_dist) dl.append(a) dr.append(b) cov.append(c) res[name] = dict(dlogp=float(np.mean(dl)), dlogp_std=float(np.std(dl, ddof=1)), dratio=float(np.mean(dr)), frac=cov[0]) print(f" {name:<14}{np.mean(dl):<+10.4f} ± {np.std(dl, ddof=1):<10.4f}" f"{np.mean(dr):<22.4f}{cov[0]}/7") print() return res, base_logp, base_dist def section_D(sc, models, n_gen, n_seed=3): print("=" * 74) print("D. 反向链方差:tilde_beta(fixed_small)还是 beta(fixed_large)") print("=" * 74) rng0 = np.random.default_rng(RNG_SEED + 77) Xtrue = sample_mix(20_000, rng0) base_logp = logp_mix(Xtrue).mean() base_dist = nearest_center_dist(Xtrue).mean() p = models["uniform"] print(" 用 uniform(L_simple)训出来的那个网络,只换每步注入的方差:") print(" 方差选择 log p0 相对真样本 到最近中心距离/真样本") out = {} for which in ["tilde", "large"]: dl, dr = [], [] for s in range(n_seed): rng = np.random.default_rng(RNG_SEED + 99 + s) a, b, _ = _score(sample_chain(sc, p, n_gen, which, rng), base_logp, base_dist) dl.append(a) dr.append(b) out[which] = dict(dlogp=float(np.mean(dl)), dlogp_std=float(np.std(dl, ddof=1)), dratio=float(np.mean(dr))) print(f" {which:<14}{np.mean(dl):<+10.4f} ± {np.std(dl, ddof=1):<10.4f}{np.mean(dr):.4f}") print() return out def mse_profile(sc, models, rng_seed=RNG_SEED + 123): """每个时刻的 eps 预测误差(蒙特卡洛,用来画图)。""" rng = np.random.default_rng(rng_seed) grid = np.array([1, 2, 5, 10, 20, 50, 100, 200, 400, 600, 800, 950, 1000]) n = 6000 prof = {} for name, p in models.items(): vals = [] for t in grid: x0 = sample_mix(n, rng) a = sc.abar[t] eps = rng.standard_normal((n, D)) xt = np.sqrt(a) * x0 + np.sqrt(1 - a) * eps vals.append(float(np.mean((predict_eps(p, xt, t) - eps) ** 2))) prof[name] = np.array(vals) return grid, prof def main(): ap = argparse.ArgumentParser() ap.add_argument("--steps", type=int, default=20000) ap.add_argument("--n-gen", type=int, default=4000) args = ap.parse_args() sc = section_A() models = section_B(sc, args.steps) res, base_logp, base_dist = section_C(sc, models, args.n_gen) dres = section_D(sc, models, args.n_gen) grid, prof = mse_profile(sc, models) print("=" * 74) print("E. 每步预测误差(MSE,真 eps 的每维方差是 1)") print("=" * 74) print(" t uniform elbo_tilde") for i, t in enumerate(grid): print(f" {t:<7d}{prof['uniform'][i]:<12.4f}{prof['elbo_tilde'][i]:.4f}") return dict(sched=sc, models=models, res=res, dres=dres, grid=grid, prof=prof) if __name__ == "__main__": main() make_figures.py # -*- coding: utf-8 -*- """画本文的四张图。数据源全部来自 ddpm_lab.py 的真实输出,不另造数。 weight_profile.png ELBO 每步权重 vs 这个时刻真正可学的信号占比 excess_error.png 两种训练目标把「超额误差」搬到了哪 variance_ledger.png 反向链每步注入的噪声:beta_tilde / beta / 真实后验 step_anatomy.png 反向一步的三个配料怎么随 t 变 运行: /usr/local/bin/python3 make_figures.py 依赖: numpy, matplotlib(字体 PingFang SC) 注意:matplotlib 的 mathtext 标签一律用 raw 字符串,且反斜杠后面只能跟字母 ——源码会被 sync-code 原样搬进文章附录,反斜杠后面跟非字母会被体检器判成转义污染。 """ import os import numpy as np import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import ddpm_lab as LAB plt.rcParams["font.sans-serif"] = ["PingFang SC", "Arial Unicode MS", "Heiti TC", "sans-serif"] plt.rcParams["axes.unicode_minus"] = False plt.rcParams["figure.facecolor"] = "white" plt.rcParams["axes.facecolor"] = "white" plt.rcParams["savefig.facecolor"] = "white" plt.rcParams["font.size"] = 11 HERE = os.path.dirname(os.path.abspath(__file__)) FIGDIR = os.path.join(os.path.dirname(HERE), "figures") os.makedirs(FIGDIR, exist_ok=True) C_MAIN, C_ALT, C_GREEN, C_GRAY = "#2563eb", "#dc2626", "#059669", "#6b7280" T = LAB.T D = LAB.D # ─────────────────────────── 图 1:权重 vs 可学占比 ─────────────────────────── def fig_weight_profile(sc, w_tilde, w_large, learn): fig, ax = plt.subplots(figsize=(9.2, 5.4)) tt = np.arange(1, T + 1) ax.semilogy(tt, w_tilde, color=C_MAIN, lw=2.0, label=r"真 ELBO 权重 $w_t$($\sigma^2=\tilde\beta_t$)") ax.semilogy(tt, w_large, color=C_ALT, lw=1.6, ls="--", label=r"真 ELBO 权重 $w_t$($\sigma^2=\beta_t$)") ax.axhline(1.0, color=C_GRAY, lw=1.2, ls=":") ax.text(620, 1.15, "等权 $L_{\mathrm{simple}}$ 就压在这条 1 上", color=C_GRAY, fontsize=9) ax.set_xlabel(r"时间步 $t$") ax.set_ylabel(r"预测误差平方前面的系数 $w_t$") ax.set_ylim(3e-3, 3.0) ax.grid(alpha=0.25, which="both") ax.legend(fontsize=9, loc="lower left") ax2 = ax.twinx() ax2.plot(tt, learn, color=C_GREEN, lw=2.0) ax2.set_ylabel(r"这一时刻能学出来的噪声占比 $R^2$", color=C_GREEN) ax2.set_ylim(-0.05, 1.05) ax2.tick_params(axis="y", labelcolor=C_GREEN) ax2.set_title("(a) 权重最大地方,恰恰是最学不到东西的地方", fontsize=11) span = w_tilde.max() / w_tilde.min() ax.annotate(rf"权重跨 {span:.0f} 倍", xy=(2, w_tilde[0]), xytext=(70, 0.35), fontsize=9, color=C_MAIN, arrowprops=dict(arrowstyle="->", color=C_MAIN, lw=1.2)) ax.annotate(rf"$t=1$ 时 $R^2$只有 {learn[0]:.1e}", xy=(1, 0.02), xytext=(120, 0.012), fontsize=9, color=C_GREEN, arrowprops=dict(arrowstyle="->", color=C_GREEN, lw=1.2)) fig.tight_layout() fig.savefig(os.path.join(FIGDIR, "weight_profile.png"), dpi=130) plt.close(fig) # ──────────────────────── 图 2:超额误差被搬到哪 ──────────────────────── def fig_excess(sc, grid, prof, floor, learn): fig, axes = plt.subplots(1, 2, figsize=(13.2, 5.0)) tt = np.asarray(grid) ax = axes[0] ax.semilogy(tt, prof["uniform"] - floor, color=C_MAIN, lw=2.0, label=r"等权($L_{\mathrm{simple}}$)") ax.semilogy(tt, prof["elbo_tilde"] - floor, color=C_ALT, lw=2.0, label="真 ELBO 权重") ax.set_xlabel(r"时间步 $t$") ax.set_ylabel(r"超额误差 MSE $-$ 贝叶斯地板") ax.set_title("(a) 容量被搬走了:低噪声档变好,中高噪声档变差") ax.grid(alpha=0.25, which="both") ax.legend(fontsize=9) ax.set_xlim(0, 1000) ax = axes[1] ratio = (prof["elbo_tilde"] - floor) / (prof["uniform"] - floor) ax.semilogy(tt, ratio, color=C_GREEN, lw=2.0) ax.axhline(1.0, color=C_GRAY, lw=1.2, ls=":") ax.fill_between(tt, 1e-3, ratio, where=(ratio < 1.0), color=C_MAIN, alpha=0.13) ax.fill_between(tt, 1.0, ratio, where=(ratio > 1.0), color=C_ALT, alpha=0.13) ax.set_xlabel(r"时间步 $t$") ax.set_ylabel("真 ELBO 权重 / 等权 的超额误差之比") ax.set_title("(b) 同一条曲线取比值:1 以下变好,1 以上变差") ax.grid(alpha=0.25, which="both") ax.set_xlim(0, 1000) ax.set_ylim(1e-2, 1e2) ax.text(60, 0.022, "低噪声档:好 10 倍以上", color=C_MAIN, fontsize=9) ax.text(430, 4.5, "中高噪声档:差 2~6 倍", color=C_ALT, fontsize=9) fig.tight_layout() fig.savefig(os.path.join(FIGDIR, "excess_error.png"), dpi=130) plt.close(fig) # ───────────────────── 图 3:每步注入的噪声账本 ───────────────────── def fig_variance(sc, mix): s2_exact = LAB.sigma2_exact_posterior(sc) tt = np.arange(1, T + 1) tr_exact = np.array([np.trace(s2_exact[t]) / D for t in range(1, T + 1)]) fig, axes = plt.subplots(1, 2, figsize=(13.2, 5.0)) ax = axes[0] ax.semilogy(tt, sc.bt[1:], color=C_MAIN, lw=2.0, label=r"$\tilde\beta_t$(DDPM 反向链用的,fixed_small)") ax.semilogy(tt, sc.beta, color=C_ALT, lw=1.8, ls="--", label=r"$\beta_t$(fixed_large)") ax.semilogy(tt, tr_exact, color=C_GREEN, lw=2.0, label=r"真实后验方差($\tilde\beta_t$ + 均值那一项的不确定性)") ax.set_xlabel(r"时间步 $t$") ax.set_ylabel("每步注入噪声的方差") ax.set_title(r"(a) 真实后验方差比 $\tilde\beta_t$ 大,差值就是 NLL 缺口") ax.set_ylim(1e-6, 1e0) ax.grid(alpha=0.25, which="both") ax.legend(fontsize=9) ax = axes[1] lam = np.asarray(mix["lam"]) sr = np.asarray(mix["std_ratio"]) nn = np.asarray(mix["nll_uniform"]) ax.plot(lam, sr, "o-", color=C_MAIN, lw=2.0, label="生成分布的标准差 / 真分布") ax.axhline(1.0, color=C_GRAY, lw=1.2, ls=":") ax.set_xlabel(r"插值系数 $\lambda$:$\sigma^2=(1-\lambda)\tilde\beta_t+\lambda\beta_t$") ax.set_ylabel("标准差之比(1 表示胖瘦刚好对上)") ax.set_ylim(0.960, 0.978) ax2 = ax.twinx() ax2.plot(lam, nn, "s--", color=C_ALT, lw=1.8, label="期望 NLL") ax2.set_ylabel("期望 NLL(nats)", color=C_ALT) ax2.tick_params(axis="y", labelcolor=C_ALT) ax.set_title(r"(b) 换成 $\beta_t$ 把 3.4% 的偏窄补回 0.6 个点(没补满)", fontsize=11) ax.grid(alpha=0.25) fig.tight_layout() fig.savefig(os.path.join(FIGDIR, "variance_ledger.png"), dpi=130) plt.close(fig) # ──────────────────── 图 4:反向一步的三个配料 ──────────────────── def fig_step(sc): tt = np.arange(1, T + 1) abar_prev = sc.abar[:-1] abar_cur = sc.abar[1:] beta = sc.beta # x_{t-1} = c_x * x_t + c_0 * x_hat0 + sigma_t * z c_x = np.sqrt(sc.alpha) * (1.0 - abar_prev) / (1.0 - abar_cur) c_0 = np.sqrt(abar_prev) * beta / (1.0 - abar_cur) sigma = np.sqrt(sc.bt[1:]) fig, ax = plt.subplots(figsize=(9.2, 5.4)) ax.semilogy(tt, c_0, color=C_MAIN, lw=2.2, label=r"拉向 $\hat x_0$ 的系数(这一步挪了多远)") ax.semilogy(tt, c_x, color=C_GRAY, lw=1.8, ls="--", label=r"保留 $x_t$ 的系数") ax.semilogy(tt, sigma, color=C_ALT, lw=2.0, label=r"注入噪声的标准差 $\sqrt{\tilde\beta_t}$") ax.set_xlabel(r"时间步 $t$") ax.set_ylabel("系数 / 标准差(绝对值)") ax.set_title("反向一步的三个配料:绝大多数步子只挪千分之几") ax.grid(alpha=0.25, which="both") ax.legend(fontsize=9) ax.set_ylim(1e-6, 3.0) ax.annotate(rf"$t=500$:只挪 {c_0[499]:.2e}", xy=(500, c_0[499]), xytext=(560, 2e-5), fontsize=9, color=C_MAIN, arrowprops=dict(arrowstyle="->", color=C_MAIN, lw=1.2)) ax.annotate(rf"$t=2$:一步挪 {c_0[1]:.2f}", xy=(2, c_0[1]), xytext=(90, 0.9), fontsize=9, color=C_MAIN, arrowprops=dict(arrowstyle="->", color=C_MAIN, lw=1.2)) fig.tight_layout() fig.savefig(os.path.join(FIGDIR, "step_anatomy.png"), dpi=130) plt.close(fig) def main(): sc = LAB.Sched(LAB.linear_beta(T)) w_tilde = sc.weight("tilde") w_large = sc.weight("large") grid = list(range(1, T + 1)) floor = np.array([LAB.mse_floor(sc, t) for t in grid]) learn = np.array([LAB.learnability(sc, t) for t in grid]) print(" 重新训练三份闭式解权重(与 ddpm_lab.py 同一套种子)...") W, _ = LAB.section_B(sc, w_tilde, w_large) prof = {} for name, Wm in W.items(): prof[name] = np.array([ LAB.mse_analytic(*LAB.model_Bc(Wm, LAB.psi_basis(np.array([t]))[0]), sc, t) for t in grid]) mix = LAB.section_E(sc, W) fig_weight_profile(sc, w_tilde, w_large, learn) fig_excess(sc, grid, prof, floor, learn) fig_variance(sc, mix) fig_step(sc) print(" 四张图已写入 figures/") return dict(sc=sc, W=W, prof=prof, floor=floor, learn=learn, mix=mix) if __name__ == "__main__": main()
2026年09月27日
5 阅读
0 评论
0 点赞
2026-09-27
AIGC 基本功|扩散过程的前向与反向推导-SDE
扩散过程的前向与反向推导 所属方向:数学基础 | 难度:入门 | 前置知识:变分下界与重参数化(KL、重参数化技巧、重参数化梯度) 关键词:马尔可夫链、前向扩散、反向去噪、随机微分方程、score matching、DDPM、DDIM、DPM-Solver 01. 为什么需要它 先摆四组数字,全部来自文末附录里五个能直接跑的脚本。 第一组:同一个 score 下比较不同采样器。 在 2D 七分量高斯混合上,score 有闭式解,可以排除神经网络拟合误差。下表记录一次固定种子实验的 dlogp,即生成样本与参考样本的平均 $\log p_0(x)$ 之差。它只是一个分布统计量:0 只说明这一个均值相同,不代表分布完全相同;正值也不是“生成得更好”。 配置 score 评估次数 dlogp DDPM 祖采样 N=1000 1000 −0.032 DDPM 祖采样 N=200 200 +0.089 DDIM(η=0)N=50 50 +0.090 Heun 二阶 N=50 100 +0.019 DDIM@50 与 DDPM@200 在这次运行中的 dlogp 接近,但不能据此得出通用的“4 倍提速”。Heun@50 的 +0.019 与 DDPM@1000 的 −0.032 也不足以证明谁更准确:还需重复随机种子、估计差值的不确定性,并结合模式占比和模式内半径等指标。比较成本时要按 score 调用次数 NFE 计,Heun 每步调用两次。 不把这件事算清楚,调 num_inference_steps 就是盲猜:既不知道收益的量级,也分不清收益里哪部分来自「换了算法的阶数」、哪部分来自「少走了冤枉路」。 第二组:调度改变了各噪声水平上的计算分配。 本实验把线性调度换成余弦,DDIM@50 的单次 dlogp 从 0.090 变到 0.050。线性调度有 74.0% 的离散步处在单位方差参考下 SNR<1 的区域,这描述了噪声水平分配,并不证明这些步骤“没有用”;高噪声阶段也负责形成全局结构。0.040 的差值需要重复实验确认,不能称为已经验证的 1.8 倍质量提升。 这张图要看什么:左图是信号方差系数,右图是在数据方差为 1 的约定下的信噪比。虚线标出 $\bar\alpha=0.5$ 的位置;低于这条线仍可能保留有用语义,不能把阴影区直接标成无效计算。 第三组:换个参数化(ε / v / x₀),等于给每个时间步换了 8 个数量级的权重。 三种参数化描述的是同一个量,但作为最小二乘目标并不等价。把它们的损失都折算回「对 ε 误差的权重」:ε 参数化恒为 1;v 参数化跨 2.5×10⁴ 倍;x₀ 参数化跨 2.5×10⁸ 倍(t=1 时 1.0×10⁻⁴,t=1000 时 2.5×10⁴)。这表明改变参数化会改变隐式时间权重。DDPM 的简化目标是经样本质量实验支持的重新加权,不能仅由此表推出其唯一理由。 第四组:验证手段本身有个坑。 我第一版是用 MMD² 给采样器排序的,跑出来的表看着非常漂亮:Heun 在 N≥10 就顶到「噪声地板」,其余采样器一路降到 0 以下。问题是那个「地板」是单次抽样的运气值——拿真实样本去对真实样本、重复 8 次,MMD² 的标准差是 ±3.6×10⁻⁴,比 N≥100 时各采样器之间的差别还大。发现它的办法很笨:做一次噪声标定。附录里的 mmd_noise_check.py 就是干这个的。分布距离的估计量本身有方差,用它排序前先量一下噪声。 02. 最小可用理解 三句话。 第一句:前向是一串手工设好的加噪,能一步算出来。 它形式上是一条 1000 步的马尔可夫链,但因为每步加的都是高斯,累乘之后仍然是高斯,所以 $q(x_t|x_0)$ 有闭式解,不需要真的跑 1000 次。 第二句:反向需要「学」的只有一个东西——score。 给定 $x_t$ 去猜刚才加了什么噪声,最小均方误差意义下的最优答案就是 $\nabla_{x_t}\log p_t(x_t)$ 的一个常数倍。把这层说清楚之后,「从噪声生成数据」就退化成一个纯粹的数值积分问题。 第三句:采样器既选择过程,也选择离散化。 DDPM / 反向 SDE、DDIM / 概率流 ODE 的随机性与漂移不同;求解器还要选择参数化、时间网格、阶数、方差和预测后处理,不能只把它们看作同一个更新式改几个系数。 这张图要看什么:左图看「抹掉」的过程——t=200 时七个模式已经开始互相渗透,t=500 彻底变成一个圆球,看不出原始结构;右图看「找回来」的过程——40 条从同一个噪声球出发的轨迹,在最后十几步才各自「决定」落进哪个模式,前面漫长的路程只是在搭粗轮廓。这解释了大步长为什么主要伤细节而不伤整体布局。 03. 数学推导 3.1 前向链:为什么能一步到位 设原始数据 $x_0 \sim p_0$。前向过程每一步只做一件事:往上加一点高斯噪声。 $$q(x_t|x_{t-1}) = \mathcal{N}(x_t; \sqrt{\alpha_t}\,x_{t-1},\ \beta_t I)$$ 这里 $\alpha_t = 1 - \beta_t$,而 $\beta_t$ 是方差而不是标准差——这个记号是 DDPM 原文定的,容易看错。写成重参数化形式就是 $x_t = \sqrt{\alpha_t}\,x_{t-1} + \sqrt{\beta_t}\,\varepsilon$。 现在把两步接起来看: $$x_t = \sqrt{\alpha_t\alpha_{t-1}}\,x_{t-2} + \sqrt{\alpha_t\beta_{t-1}}\,\varepsilon_1 + \sqrt{\beta_t}\,\varepsilon_2$$ 两个独立高斯的线性组合仍是高斯,噪声方差为 $(\alpha_t\beta_{t-1}+\beta_t)I$。代入 $\beta_{t-1}=1-\alpha_{t-1}$、$\beta_t=1-\alpha_t$,得到 $\alpha_t(1-\alpha_{t-1})+(1-\alpha_t)=1-\alpha_t\alpha_{t-1}$。这才是两步方差的正确化简。 于是归纳下去,任意步数都能一步算出来。记 $\bar\alpha_t = \prod_{s\le t}\alpha_s$: $$q(x_t|x_0) = \mathcal{N}\big(x_t;\ \sqrt{\bar\alpha_t}\,x_0,\ (1-\bar\alpha_t)I\big)$$ 这就是整个扩散模型里唯一一条「白送」的公式,训练时要多少步的加噪样本都能直接算。注意它成立的前提是噪声必须是各向同性高斯:换成重尾噪声,累乘就不再是同类分布,这条式子立刻失效。 附录里的 forward_diffusion.py 拿 20 万样本对着验了一遍:用「逐步迭代 1000 次」和「一步闭式」两条路分别算 $x_t$ 的均值与方差,误差量级都在 $10^{-3}$,正好是 20 万样本的蒙特卡洛噪声水平(t=1000 时闭式方差误差 6.20×10⁻³,逐步方差误差 6.19×10⁻³,两者几乎相等,说明闭式解没错)。 3.2 反向链:真实后验长什么样 我们想要的是 $q(x_{t-1}|x_t)$,它不好算。但加上 $x_0$ 之后就好算了——贝叶斯公式、配方、得到 $$q(x_{t-1}|x_t,x_0) = \mathcal{N}\Big(x_{t-1};\ \tilde\mu_t,\ \tilde\beta_t I\Big)$$ 其中 $\tilde\beta_t = \dfrac{1-\bar\alpha_{t-1}}{1-\bar\alpha_t}\beta_t$,$\tilde\mu_t = \dfrac{\sqrt{\bar\alpha_{t-1}}\,\beta_t}{1-\bar\alpha_t}x_0 + \dfrac{\sqrt{\alpha_t}(1-\bar\alpha_{t-1})}{1-\bar\alpha_t}x_t$。 $\tilde\beta_t$ 是在已知 $x_0$ 和 $x_t$ 后的剩余不确定性。$t=1$ 时 $\bar\alpha_0=1$,所以 $\tilde\beta_1=0$;早期分子分母之比可能明显小于 1,后期二者都接近 1 时才有 $\tilde\beta_t\approx\beta_t$。它不是只在中间段才变小。 问题在于 $\tilde\mu_t$ 里含着 $x_0$,而这个量正是我们不知道的。DDPM 的做法是让网络去猜:把 $\tilde\mu_t$ 里的 $x_0$ 换成一个网络估计 $\hat x_0$,就得到一个可采样的反向链。 3.3 从 ELBO 到「只预测噪声」 把前向过程 $q(x_{1:T}\mid x_0)$ 作为变分分布,生成模型使用反向链 $p_\theta(x_{0:T})$,负 ELBO 分解为: $L_T=\mathrm{KL}(q(x_T\mid x_0)\Vert p(x_T))$——条件终端先验项。它依赖数据与固定调度,不依赖去噪网络。对本文数据取期望,修正脚本的闭式结果为线性调度 $1.065154\times10^{-4}$ nats、余弦调度 $6.410060\times10^{-9}$ nats。原先蒙特卡洛测的是边缘 $\mathrm{KL}(q(x_T)\Vert p(x_T))$,两者相差 $I(x_0;x_T)$,不能混用。忽略固定 $L_T$ 不改变网络梯度,但报告似然下界时仍需计入。 $L_{t-1} = \mathrm{KL}\big(q(x_{t-1}|x_t,x_0)\,\|\,\text{反向链}\big)$——去噪匹配项,起主导作用。 $L_0 = -\log p_\theta(x_0|x_1)$——最终重建/解码项,离散像素需要相应离散化似然;连续数据也不能不经评估就认定该项很小。 主项是两个高斯之间的 KL。两个同协方差高斯 $p=\mathcal{N}(\mu_p,\Sigma),q=\mathcal{N}(\mu_q,\Sigma)$ 的 KL 正好是 $\tfrac12(\mu_p-\mu_q)^\top\Sigma^{-1}(\mu_p-\mu_q)$——只剩均值之间的加权距离,协方差项被消掉了。代入 $\tilde\mu_t$ 和它的网络版本,需要网络拟合的只有 $\hat x_0$(或者等价地,$\varepsilon$)这一个量。 接下来是最关键的一次代入。用 $x_t = \sqrt{\bar\alpha_t}x_0 + \sqrt{1-\bar\alpha_t}\,\varepsilon$ 把 $x_0$ 换成 $\varepsilon$,会得到 $$\tilde\mu_t = \frac{1}{\sqrt{\alpha_t}}\Big(x_t - \frac{\beta_t}{\sqrt{1-\bar\alpha_t}}\,\varepsilon\Big)$$ 也就是均值里对 $\varepsilon$ 的依赖是线性的、系数是确定的。既然输出只需选一种等价参数化(每个数据坐标仍有一个分量),那就干脆让网络直接输出 $\varepsilon$,损失变成 $$L_{\text{simple}} = \mathbb{E}_{t,x_0,\varepsilon}\Big[\big\|\varepsilon - \varepsilon_\theta\big(\sqrt{\bar\alpha_t}x_0 + \sqrt{1-\bar\alpha_t}\,\varepsilon,\ t\big)\big\|^2\Big]$$ 从 ELBO 到这一行,中间扔掉了一串只跟 t 有关的系数。 为什么可以扔,06 节会用数字回答。 3.4 连续化:同一件事的两种时间写法 1000 步只是一个离散近似。把步长做成无穷小,$\beta_t$ 变成一个率 $\beta(u)$、$u\in[0,1]$,马尔可夫链就变成一个随机微分方程: $$dx = -\tfrac12\beta(u)\,x\,du + \sqrt{\beta(u)}\,dw$$ 这就是 Song 等人说的 VP-SDE(方差保持型)。系数 $-\tfrac12\beta(u)$ 让数据慢慢收缩到 0,$\sqrt{\beta(u)}$ 同时在加噪声——其边缘方差为 $\bar\alpha(u)\mathrm{Var}(x_0)+1-\bar\alpha(u)$:初始方差为 1 时才严格保持 1,否则逐渐趋向 1。 离散和连续能不能对上,是有条件的。 离散的 $\bar\alpha_t = \prod(1-\beta_s)$ 与连续的 $\exp(-\int_0^u \beta)$,只有在 $\beta$ 足够小时才接近。附录实测:线性调度折算回连续时间之后 $\beta_{\min}=0.1$、$\beta_{\max}=20$(正好是 Song 等人论文里的默认值),两条路的相对差从 $u=0.25$ 的 7.74×10⁻⁴ 涨到 $u=1$ 的 7.01×10⁻²。也就是说,连续时间的那套结论不能无条件搬到离散实现上——这个 7% 就是「离散化误差」的本体。 正向 SDE 对应一个 Fokker–Planck 方程,描述整个分布 $p_u$ 怎么随时间流动。关键观察:同一个 Fokker–Planck 方程对应无穷多条 SDE,它们的漂移项不同、但边缘分布完全相同。其中两条特别有用: $$\text{反向 SDE:}\quad dx = \Big[-\tfrac12\beta x - \beta\,\nabla_x\log p_u(x)\Big]du + \sqrt{\beta}\,d\bar w$$ $$\text{概率流 ODE:}\quad dx = \Big[-\tfrac12\beta x - \tfrac12\beta\,\nabla_x\log p_u(x)\Big]du$$ 这两条式子之间只差两个地方:score 前面的系数,和那个噪声项。 系数从 1 减半到 1/2,删掉噪声——因为注噪声带来的那部分扩散,正好被「一半的 score」抵消了,两者合起来保持边缘分布不变。所有采样器都在这两条式子之间做取舍,这就是 02 节第三句话的出处。 一个我实际踩的符号坑。 正向时间 $u$ 从 0 涨到 1,反向采样是让 $u$ 往回走,所以 $du<0$。我第一版把这一步忘了,写成 $x \leftarrow x - \tfrac12\beta(x+s)$,结果 Euler-Maruyama、概率流 ODE、Heun 三个采样器全线崩掉(MMD² 卡在 0.5 下不来,而 DDPM/DDIM 正常)。正确的写法是把「每步跨过的积分量」先抠出来: $$L = \int_{u_{\text{prev}}}^{u_{\text{cur}}}\beta(u)\,du = \log\frac{\bar\alpha_{\text{prev}}}{\bar\alpha_{\text{cur}}} > 0$$ 代入 $du=-1/N$($N$ 是采样步数)之后符号整体翻转,Euler 步变成 $x \leftarrow x + L\,(\tfrac12 x + c\,s)$,其中 $c=1$ 是反向 SDE、$c=1/2$ 是概率流 ODE。物理上很好理解:反向过程是把被前向压扁的分布吹回原样,漂移当然要往外推。 3.5 四种参数化:同一个量的四个名字 代码里同一个东西有四种写法,它们之间全是恒等式: $$\varepsilon = -\sqrt{1-\bar\alpha_t}\ \nabla_x\log p_t(x)$$ $$\mathbb{E}[x_0|x_t] = \frac{x_t + (1-\bar\alpha_t)\nabla_x\log p_t(x_t)}{\sqrt{\bar\alpha_t}}$$ 第二条就是 Tweedie 公式,它说的是「去噪」和「算 score」是同一件事。第一条则把 score 和 ε 预测对上。把它代进第二条,就得到代码里那句最眼熟的 pred_x0 = (x - sqrt(1-a) * eps) / sqrt(a)。 第四种是 v 参数化:$v = \sqrt{\bar\alpha_t}\,\varepsilon - \sqrt{1-\bar\alpha_t}\,x_0$。它看起来像个随手拼出来的组合,实际上 $(x_t, v)$ 和 $(x_0, \varepsilon)$ 之间是一个旋转: $$x_0 = \sqrt{\bar\alpha_t}\,x_t - \sqrt{1-\bar\alpha_t}\,v,\qquad \varepsilon = \sqrt{1-\bar\alpha_t}\,x_t + \sqrt{\bar\alpha_t}\,v$$ 验证很简单:把 $x_t = \sqrt{\bar\alpha_t}x_0 + \sqrt{1-\bar\alpha_t}\varepsilon$ 代入第一个式子,$\sqrt{\bar\alpha_t}\sqrt{1-\bar\alpha_t}$ 的交叉项正好抵消,剩下 $(\bar\alpha_t + 1 - \bar\alpha_t)x_0 = x_0$。变换矩阵的行列式是 $\bar\alpha_t + (1-\bar\alpha_t)=1$,正交——这就是 v 参数化「不会放大噪声」的来源。 score_bridges.py 把这四条恒等式全核了一遍:Tweedie 公式与精确后验均值的最大绝对误差在 $10^{-15}\sim10^{-13}$(t=1 时 1.332×10⁻¹⁵,t=1000 时 1.550×10⁻¹³),ε 的两条算法误差在 $10^{-16}\sim10^{-15}$,v 的重建误差 8.882×10⁻¹⁶。另外还做了一次不依赖闭式解的交叉验证:用 100 万样本的重要性采样直接估 $\mathbb{E}[x_0|x_t]$,与 Tweedie 公式的答案在小数点后两到三位一致(有效样本数从 t=50 时的 49895 涨到 t=1000 时的 999852)。 04. 代码实现 4.1 实验设计:把「模型误差」这个变量消掉 真实扩散模型的采样误差来自两处:score 估计得不准,以及数值积分不准。想把第二处单独看清楚,就得让第一处等于零——用一个 score 有闭式解的目标分布。 办法是选高斯混合:$p_0 = \sum_k w_k\mathcal{N}(\mu_k, s_k^2 I)$。前向加噪之后,第 $k$ 个分量的均值缩到 $\sqrt{\bar\alpha_t}\mu_k$、方差变成 $\bar\alpha_t s_k^2 + (1-\bar\alpha_t)$,仍然是各向同性高斯,所以 $p_t$ 还是高斯混合,且 $$\nabla_x\log p_t(x) = -\sum_k r_k(x)\frac{x - \sqrt{\bar\alpha_t}\mu_k}{\bar\alpha_t s_k^2 + (1-\bar\alpha_t)},\qquad r_k(x) = \frac{w_k\mathcal{N}_k(x)}{\sum_j w_j\mathcal{N}_j(x)}$$ $r_k$ 就是「这个样本属于第 $k$ 个分量」的责任度,是标准的软分配。于是整条反向链路都可以用真值跑,跑出来的差异 100% 来自离散化和「要不要注噪声」。 这里有个调参坑值得记一下:我第一版用了 3 个很宽的分量(标准差 0.45/0.35/0.55),结果五种采样器全部顶到 MMD 噪声地板上,分不出高下。换成 7 个紧分量(标准差 0.16~0.30)之后差距才显出来——目标太光滑,分辨不出采样器的差别。这和真实情况是对应的:scheduler 之间的差别本来就在高频细节上。 4.2 前向:闭式解核对 核心就几行,除了算 $\bar\alpha$ 之外没有任何魔法: def linear_beta(T=1000, beta_1=1e-4, beta_T=0.02): return np.linspace(beta_1, beta_T, T) def alpha_bar_from_beta(beta): return np.cumprod(1.0 - beta) def q_sample(x0, t_idx, alpha_bar, rng): a = alpha_bar[t_idx] eps = rng.standard_normal(x0.shape) return np.sqrt(a) * x0 + np.sqrt(1.0 - a) * eps, eps q_sample 就是 3.1 节那个闭式解。余弦调度多两行——按 $\bar\alpha_t$ 定义式算完之后要把 β 截到 0.999 以内,因为 $t=T$ 时 $\cos(\pi/2)=0$ 会让最后一步 $\beta=1$,实操直接炸: def cosine_beta(T=1000, s=0.008, clip=0.999): t = np.arange(1, T + 1) / T f = np.cos(((t + s) / (1 + s)) * np.pi / 2) ** 2 f0 = (np.cos((s / (1 + s)) * np.pi / 2)) ** 2 beta = beta_from_alpha_bar(f / f0) # 反解出每步的 beta return np.clip(beta, None, clip) 4.3 反向:五个采样器 DDPM 祖采样——用高斯核近似反向转移。对一般数据,即使均值来自精确 score,有限步高斯核也不等于真实反向条件分布。把 3.2 节的 $\tilde\mu_t$ 和 $\tilde\beta_t$ 抄进来,再把 $\hat x_0$ 换成 Tweedie 公式: def run_ddpm(x, alpha_bar, taus, rng): for j in range(len(taus) - 1): t_cur, t_prev = taus[j], taus[j + 1] a_cur, a_prev = abar_at(alpha_bar, t_cur), abar_at(alpha_bar, t_prev) alpha_j = a_cur / a_prev # 跨 k 步的等效 alpha beta_j = 1.0 - alpha_j s = score_at(x, alpha_bar, t_cur) mean = (x + beta_j * s) / np.sqrt(alpha_j) var = beta_j * (1.0 - a_prev) / (1.0 - a_cur) x = mean + np.sqrt(var) * rng.standard_normal(x.shape) return x Euler-Maruyama(反向 SDE)——把 $L$ 和 $c=1$ 代进 3.4 节的结果: def run_em(x, alpha_bar, taus, rng): for j in range(len(taus) - 1): L = _log_step(alpha_bar, taus[j], taus[j + 1]) s = score_at(x, alpha_bar, taus[j]) x = x + L * (0.5 * x + s) + np.sqrt(L) * rng.standard_normal(x.shape) return x 概率流 ODE——唯一的改动是 score 系数减半、噪声项删掉: def run_ode(x, alpha_bar, taus, rng): for j in range(len(taus) - 1): L = _log_step(alpha_bar, taus[j], taus[j + 1]) s = score_at(x, alpha_bar, taus[j]) x = x + 0.5 * L * (x + s) return x Heun 二阶——Euler 预测一步,再用终点的 score 校正一次。每步两次评估: def run_heun(x, alpha_bar, taus, rng): for j in range(len(taus) - 1): L = _log_step(alpha_bar, taus[j], taus[j + 1]) s0 = score_at(x, alpha_bar, taus[j]) d0 = 0.5 * L * (x + s0) x1 = x + d0 # Euler 预测 s1 = score_at(x1, alpha_bar, taus[j + 1]) # 终点再评估一次 d1 = 0.5 * L * (x1 + s1) x = x + 0.5 * (d0 + d1) return x DDIM 单独说,因为它长得不像上面四个。它先把样本一步跳到 $\hat x_0$,再按目标时刻的 $\bar\alpha$ 重新加回噪声: $$x_{t-1} = \sqrt{\bar\alpha_{t-1}}\,\hat x_0 + \sqrt{1-\bar\alpha_{t-1}}\,\hat\varepsilon$$ $\hat x_0$ 用 Tweedie 公式算、$\hat\varepsilon$ 用 $\varepsilon=-\sqrt{1-\bar\alpha_t}s$ 算。η=0 的 DDIM 在小步长极限对应概率流 ODE,其 score 漂移系数为反向 SDE 的一半;DDPM 的均值更新对应后者,两者并不在代数上相同。 4.4 结果 七个分量、每档生成 4000 个样本、参照集 8000 个真实样本。三个主指标都先扣掉了「真实样本自己」的基线,所以 0 才等于完美: 采样器 N=10 N=25 N=50 N=100 N=200 N=1000 DDPM −0.365 +0.266 +0.249 +0.174 +0.089 −0.032 DDIM −0.248 +0.124 +0.090 +0.057 +0.032 +0.007 EMA(SDE) −2.725 −0.756 −0.347 −0.189 −0.100 −0.081 Euler(ODE) −2.340 −0.792 −0.348 −0.159 −0.074 −0.014 Heun +0.546 +0.084 +0.019 +0.005 +0.002 +0.001 上表是单次实验的 dlogp。两批独立真实样本在这次运行中的差值约 0.053;这不是经重复估计的标准差或置信区间。表格主要帮助识别量级很大的离散化偏差,微小差异不做显著性排序。三点读法: DDPM / DDIM 在 N=10 还能看,Euler 系直接崩(−2.3 ~ −2.7,说明样本散在低密度区根本没收敛)。这不能怪 ODE 或 SDE,只能怪第一步就跨了 1.92 的积分量。 误差变号这件事有意义。 DDPM/DDIM 的 dlogp 是正的(样本被堆到高密度区、偏聚拢),Euler 系是负的(还没收敛、偏散开)。两种失败模式方向相反,光看一个「距离」指标看不出来。 本次 Heun 在 N≥20 的 dlogp 较小,但差异接近采样波动时不作排名;接近0也不足以证明整个分布正确。 补充两个不同统计量:按模式归属估算的分量占比 TV,以及模式内均方半径的相对误差。在 N=50 的本次运行中,Heun 为 0.011 / 0.066,DDIM 为 0.014 / 0.155,DDPM 为 0.019 / 0.340,Euler-ODE 为 0.017 / 0.358。它们分别检查模式质量与分散程度,排序并非处处一致,也要估计采样误差。 这张图要看什么:左图画的是单个诊断统计量的偏差,参考虚线只表示一次真实样本对照差值,不是置信区间。右图比较线性漂移项的指数因子和 Euler 近似,解释大步长为何可能困难;完整 score 场同时参与更新,不能仅用这一项定量归因全部误差。 4.5 为什么朴素 Euler 在大步长下会崩 把上面第 1 点挖到底。概率流 ODE 里线性部分 $\tfrac12\beta x$ 的精确解是 $e^{L/2}$,而 Euler 用的是它的一阶展开 $1+L/2$。两者的相对误差随 $L$ 指数上升: 采样步数 N 每步最大积分量 $L_{\max}$ $e^{L/2}$ 与 $1+L/2$ 的相对误差 10 1.9197 24.95% 20 0.9852 8.80% 50 0.4002 1.75% 100 0.2011 0.47% 200 0.1008 0.12% 1000 0.0202 0.01% DDPM 和 DDIM 把这部分精确解掉了——它们的更新直接把 $\sqrt{\bar\alpha}$ 乘上去,等于用指数积分器而不是 Euler。所以 N=10 时它们还能看,而 Euler 每步有 25% 的相对误差、误差还会沿着链累积。这也解释了为什么工业实现里大家都用 DDIM/PNDM/DPM-Solver 的写法,而不是拿反向 SDE 直接上 Euler。 高阶方法也需要合适的步长。 本次 Heun@10 的 dlogp 为 +0.546,DDIM@10 为 −0.248,说明这个网格上的 Heun 出现明显偏聚拢。原因可能包括稳定域、非线性 score 与步长的相互作用;不能说二阶校正必然放大线性项误差,更不能推广为 DPM-Solver 高阶算法在所有少步数任务上都更差。 05. 工业级实现对照 对着 huggingface/diffusers 的 src/diffusers/schedulers/scheduling_ddpm.py 看(https://github.com/huggingface/diffusers/blob/main/src/diffusers/schedulers/scheduling_ddpm.py,以 2026-09 时的实现为准)。 add_noise 就是 3.1 节的闭式解,一字不差: sqrt_alpha_prod = alphas_cumprod[timesteps] ** 0.5 sqrt_one_minus_alpha_prod = (1 - alphas_cumprod[timesteps]) ** 0.5 noisy_samples = sqrt_alpha_prod * original_samples + sqrt_one_minus_alpha_prod * noise 第一个值得注意的细节:noise 是从外面传进来的,不是在函数里采的。因为训练时需要知道「这次加的是哪个 ε」才能算损失——如果函数内部自己采样,你就永远拿不到那个 target。 第二个细节:alphas_cumprod 在 __init__ 里用 torch.cumprod(1 - betas) 一次算好并缓存。1000 个数,每个训练步都要按 timestep 取,重新累乘显然不划算。这就是 3.1 节那条「白送」的公式在工程上的直接体现。 step 与我的 numpy 实现是代数等价的。 它写的是 DDPM 原文公式 (7) 的系数形式: current_alpha_t = alpha_prod_t / alpha_prod_t_prev current_beta_t = 1 - current_alpha_t pred_original_sample = (sample - beta_prod_t ** 0.5 * model_output) / alpha_prod_t ** 0.5 # 即 Tweedie pred_original_sample_coeff = (alpha_prod_t_prev ** 0.5 * current_beta_t) / beta_prod_t current_sample_coeff = current_alpha_t ** 0.5 * beta_prod_t_prev / beta_prod_t pred_prev_sample = pred_original_sample_coeff * pred_original_sample + current_sample_coeff * sample 把 pred_original_sample_coeff 和 current_sample_coeff 代进 3.2 节的 $\tilde\mu_t$ 表达式展开,$x_t$ 的系数会化简成 $1/\sqrt{\alpha_j}$、$\varepsilon$ 的系数化简成 $-\beta_j/(\sqrt{\alpha_j}\sqrt{1-\bar\alpha_{\text{cur}}})$,也就是 $\tilde\mu = (x_t + \beta_j s)/\sqrt{\alpha_j}$——和我 4.3 节那个「一行 DDPM」完全一样。current_alpha_t = alpha_prod_t / alpha_prod_t_prev 对应的正是我的 alpha_j,所以 diffusers 支持跳步(strided)采样;这里只在相同 prediction_type、方差与不启用 clipping/thresholding 的条件下和最小实现等价。 三处「最小实现没有、工业实现必须有」的差异: 第一,调度不是自由参数,是分类的。beta_schedule 有 linear / scaled_linear / squaredcos_cap_v2 / sigmoid 几个分支;我 3.1 节实现的余弦调度在它这里叫 squaredcos_cap_v2,走的是通用函数 betas_for_alpha_bar——先给定 $\bar\alpha(t)$ 的解析式,再按 $1 - \bar\alpha(t_2)/\bar\alpha(t_1)$ 反解出每步 β,并且硬编码 max_beta=0.999。这正好对上我 cosine_beta 里那句 clip,不是巧合:$t=T$ 时 $\cos(\pi/2)=0$ 会给出 $\beta_T=1$,必须截。 第二,有限终端 SNR 会引入起点分布近似。线性表的 $\bar\alpha_T=4.0358\times10^{-5}$,仍保留少量数据成分;推理却常从标准高斯开始。diffusers 的 rescale_betas_zero_snr 对应 Lin 等,2023 的修正,但要与训练的参数化、时刻采样和起点设置一起核对。对已有 epsilon checkpoint 不能只开一个开关就假定兼容,零 SNR 下若仍用除以 $\sqrt{\bar\alpha_T}$ 的公式还会遇到奇异点。 第三,参数化是可切换的输出头。prediction_type 支持 epsilon / sample / v_prediction 三选一,step 开头那个 if-elif 就干这件事。三种模式对 pred_original_sample 的算法不同,但后面的系数计算完全共用——这正是 3.5 节「它们描述同一个量」在工程上的样子。也正因如此,采样端转换形式可由配置选择,但 checkpoint 的训练目标必须匹配,不能把 epsilon 权重仅改一行配置就当作 v 预测器,代价藏在 06 节的权重表里。 06. 代价与边界 代价一:换参数化不是改记号,是改损失权重。 把三种参数化都折算回「对 $\varepsilon$ 误差的权重」(推导见 score_bridges.py 的文档串:x₀ 的损失 $\|\delta x_0\|^2$ 折成 $\varepsilon$ 误差要乘倍率,v 的误差因为 $\delta\varepsilon = \sqrt{\bar\alpha_t}\,\delta v$ 也要乘): t $\bar\alpha_t$ ε 参数化的权重 v 参数化的权重 x₀ 参数化的权重 1 9.999×10⁻¹ 1.0000 1.0001 1.0001×10⁻⁴ 100 8.970×10⁻¹ 1.0000 1.1148 1.1480×10⁻¹ 500 7.859×10⁻² 1.0000 1.2725×10¹ 1.1725×10¹ 900 2.752×10⁻⁴ 1.0000 3.6336×10³ 3.6326×10³ 1000 4.036×10⁻⁵ 1.0000 2.4778×10⁴ 2.4777×10⁴ 这张表是等价残差之间的代数权重,不是网络参数梯度的实测值。 ε、v、x₀ 三种目标的输出尺度、网络雅可比以及时间采样方式都会影响梯度。等权 ε 损失在这里对应 ε 残差权重恒为 1;不能据此断言各时刻梯度相同,也不能仅凭该表断言 x₀ 目标一定被高噪声支配或无法训练。 这张图要看什么:把不同输出误差换算到 ε 误差时,各自带上不同时间权重。曲线说明训练目标不等价,不是实测梯度图;实际训练还要连同参数化和时间采样权重一起分析。 代价二:阶数更高也要选择合适的网格。 Heun@10 的表现说明本实验的大步长不合适,不是对所有高阶求解器的否定。DPM-Solver 专门利用扩散 ODE 的半线性结构,不能从朴素 Heun 的结果推导它的少步数表现。 代价三:随机项改变有限步采样的误差与方差。 DDIM 的 η 从 0 到 1,本次 N=50 的 dlogp 为 0.090、0.065、0.092、0.169、0.249。这里 η=1 配上相同时间网格和方差选择可恢复所实现的 DDPM 更新,但不能把单次结果概括为“噪声没有用”。在精确 score、连续时间和正确初始分布下,反向 SDE 与概率流 ODE 都能得到相同边缘分布;有限步误差的优劣取决于具体离散化。 边界:这套实验测不到模型误差。 oracle score 排除了训练误差,却仍有有限终端分布近似、离散化和样本统计误差;它不是现实图像模型误差的严格下界。真实网络的误差还会与采样轨迹交互,必须在目标 checkpoint 上另做实验。 边界:高斯混合不是自然图像。 这里的数据维数、模态结构、score 光滑性和引导强度都很简单。表格适合验证公式与数值方法,不能当作真实模型加速倍率或通用采样器排名。 ODE 也能生成多样样本与估计似然。 确定性只意味着固定初始噪声对应固定轨迹;不同随机初值仍可覆盖整个数据分布。原始 score-SDE 论文 就利用概率流 ODE 计算似然。在向量场满足正则条件时精确流可逆,但有限步 DDIM / ODE 数值求解一般不能无误差反演。 按步数预算给一张速查表(数字全部来自 4.4 节那张表,dlogp,绝对值越小越好): 步数预算 该选谁 实测依据 8~10 步 在目标模型上比较合适的少步采样器 本 toy 的 Heun@10 偏聚拢,不能推出其他高阶方法的排名 20~50 步 Heun 是本实验可考虑的方案 Heun@50 的 dlogp 为 0.019,需结合其他指标和重复试验 100~200 步 同 NFE 比较,避免只按步数选 N=200:Heun 0.002、DDIM 0.032、DDPM 0.089 要多样性 ODE / SDE 都可,比较分布覆盖 ODE 的多样性来自随机初值;SDE 还增加路径随机性 要反演 / 要编辑 可考虑概率流 ODE / DDIM 理想流可逆;有限步反演仍有数值和模型误差 表格只总结当前教学实验能支持的选择。没有实测的 PNDM / DPM-Solver 不参与排名;真实模型应以同等 NFE、重复种子与多个质量指标比较。 07. 经典论文脉络 ① 1503.03585(Sohl-Dickstein et al., 2015)——把扩散搬进生成模型。 用非平衡热力学里的一个想法:先定义一个把数据逐步破坏成噪声的正向过程,再学它的反向过程。贡献是框架本身,同时研究了高斯与二项扩散等设置,采样慢到没有实战价值。 ② 2006.11239(Ho et al., 2020)——DDPM,把目标改成「预测噪声」。 三件事:把 $\tilde\mu_t$ 参数化成预测 $\varepsilon$;指出把 ELBO 里那一串只跟 t 有关的系数扔掉、直接用等权的 $L_{\text{simple}}$ 反而效果更好;给出 3.2 节那套 $\tilde\beta_t / \tilde\mu_t$ 的闭式解。这才是「扩散模型能训练起来」的直接原因——在这之前,没人找到规模化的训练目标。 ③ 2011.13456(Song et al., 2021)——把离散和连续统一起来。 这篇是本节点的锚点。它做了两件大事:把 DDPM(VP-SDE)和它自己那套 score matching(VE-SDE)统一到同一个 SDE 框架下,把它们写成不同漂移与扩散系数下的 SDE,并由相应 Fokker–Planck 方程构造边缘等价的概率流 ODE;以及提出了概率流 ODE——同一个边缘分布、确定性求解、还能用现成的 ODE 求解器(这篇文章里用四阶 Runge–Kutta)。3.4 节那两条式子就出自这里。 ④ 2010.02502(Song et al., 2021)——DDIM,确定性采样。 把反向链改成不含随机项的确定性映射:样本轨迹只由 $x_T$ 决定,跳步采样不再需要「一步一小步」的假设。两篇作者不同:DDIM 的第一作者是 Jiaming Song,score-SDE 的第一作者是 Yang Song——DDIM 本质上就是概率流 ODE 的一个(指数积分器风格的)离散化。它让 50 步的采样第一次在质量上追平 1000 步。 ⑤ 2102.09672(Nichol & Dhariwal, 2021)——余弦调度。 指出线性调度把信噪比压得太快(就是 01 节第二组数字里那 74%),改成 $\bar\alpha_t = \cos^2(\cdot)$ 之后低步数下的质量明显更好。这篇的价值在于它把「调度」从工程细节变成了有图像解释的设计问题:$\bar\alpha_t$ 的形状决定了「每个时刻还剩多少信息量」,而余弦的形状让信息量的衰减更均匀。 在这五篇之外还有两条重要支线:v 参数化(arXiv:2202.00512,Salimans & Ho)解决 06 节权重表里 x₀ 参数化的尺度失衡问题;零终端 SNR(arXiv:2305.08891)修掉 05 节那个 $\bar\alpha_T\neq0$ 的 bug。它们都是在这条主线已经跑通之后,针对具体失效模式的补丁。 08. 常见误解 误解一:「DDPM 和 DDIM 是同一条 ODE 的两种离散化。」 DDPM 是带噪声的高斯反向马尔可夫链,连续极限对应反向 SDE;η=0 的 DDIM 是确定性路径,可联系到概率流 ODE。二者的连续边缘分布可一致,但漂移中的 score 系数不同:SDE 是 1,ODE 是 1/2,不能只把 DDPM 的噪声删除就得到 DDIM。 误解二:「反向过程就是把噪声一步步减掉。」 方向反了。反向 drift 是 $+\tfrac12\beta x + \beta s$,$x$ 那一项是往外推的。前向把分布压向原点,反向把它吹回原样。我按「减掉」实现了三个采样器,全部崩掉(4.3 节那个符号坑),而错误版本跑起来并不报错、只是数值不对——这类 bug 只能靠对着闭式解核对来抓。 误解三:「换个参数化只是记号问题,等价就是等价。」 描述的对象等价,作为训练目标不等价。权重表跨 8 个数量级(06 节),这不是小差异。等价只发生在「已经收敛到精确最优解」这个极限情况;训练过程中不同参数化走的路径完全不同。 误解四:「换成 mean log p 就能可靠排序。」 MMD / FID 需要估计不确定性,mean log p 同样需要,而且单个均值不能刻画完整分布。对两批真实样本只算一次差值,不能据此画置信区间;必须重复抽样或用适当的标准误 / bootstrap 分析。 误解五:「ODE 无随机项,所以没有多样性;SDE 多走几步只会累积坏噪声。」 随机性可以来自初始噪声,也可以来自路径。精确连续过程下两者都可得到正确分布;更多步数通常减少相应数值方法的离散误差,但有限模型、引导、网格与算力约束下仍应实测。 误解六:「$\bar\alpha_t$ 和 $1-\bar\alpha_t$ 加起来是 1,所以信噪比就是 $\bar\alpha_t$。」 信噪比是 $\bar\alpha_t/(1-\bar\alpha_t)$,不是 $\bar\alpha_t$。t=500 时线性调度 $\bar\alpha_{500}=7.86\times10^{-2}$,看着还有 7.9% 的信号,但信噪比只有 0.085(−10.7 dB)——信号和噪声的幅度比已经掉到 1:3.4 以下。只看 $\bar\alpha$ 会严重高估「还剩多少信息」。 09. 动手验证 五个脚本都在文末附录里,用 /usr/local/bin/python3 直接跑,只依赖 numpy 和 matplotlib(make_figures.py 需要 matplotlib,其余四个只要 numpy)。 python forward_diffusion.py # 前向闭式解核对 + 两种调度对比 + 终端 SNR python reverse_sampling.py # 五种采样器 × 七档步数的主实验 python score_bridges.py # Tweedie / ε / v 四条恒等式核对 + 权重表 python mmd_noise_check.py # 先量一下 MMD² 自己的噪声 python make_figures.py # 生成四张图 值得自己动手改着看的四处: 第一,把 N 设成 10,看 DDIM 和 Euler 谁先崩。 预期:DDIM 的 dlogp 是 −0.248(偏散、没收敛完),Euler 是 −2.340(散得离谱)。原因是 4.5 节那张表——N=10 时每步跨 1.92 的积分量,Euler 的线性部分误差 24.95%,而 DDIM 直接乘 $\sqrt{\bar\alpha}$ 把它精确解掉了。 第二,给 run_ode 加一个符号。 把 x + 0.5 * L * (x + s) 改成 x - 0.5 * L * (x + s) 再跑,预期:MMD² 从 8.57×10⁻⁴ 涨到 0.5 量级,dlogp 从 −0.348 变成毫无意义的正数。这就是 3.4 节那个 $du<0$ 的坑。 第三,把目标分布的 SCALES 从 0.16~0.24 改成 0.45 附近。 预期:五种采样器在 N≥50 时全部顶到噪声地板,表格失去分辨力。这是 4.1 节那个「目标太光滑就分不出来」的第二遍确认。 第四,改 forward_diffusion.py 里的 linear_beta,把 beta_T 从 0.02 降到 0.01。 预期:$\bar\alpha_T$ 变大(终端残留信号变多),t=500 的 SNR 也跟着变。可以顺着看 05 节说的那个「非零终端 SNR」到底有多敏感。 10. 延伸阅读 前置:变分下界与重参数化(vae_elbo)。3.3 节里「两个同协方差高斯的 KL 只剩均值距离」和「重参数化让梯度穿透采样」这两步,在前置那篇里有完整的推导和代码验证,这里直接用结论了。 后续:视频 VAE 的时空压缩结构(video_vae)。把 3.1 节的 2D 扩散扩到时空 $(T,H,W)$ 之后,噪声调度要不要沿时间轴做非均匀分配,是视频模型和图像模型第一个分岔点。 后续:FID / CLIP Score 到底测了什么(image_metrics)。08 节误解四的教训在通用指标上会放大——FID 的估计方差、样本量、以及它对「模式丢失」不敏感这件事,值得单独算一次。 后续:KV Cache 与自回归视频生成(kv_cache)。自回归视频模型把「一步去噪」换成「一次前向出一帧」,3.4 节的两条 SDE/ODE 就不再适用了,但 score 那一层的直觉仍然通用。 附录:完整代码 09 节用到的脚本全文如下(mmd_noise_check.py、forward_diffusion.py、score_bridges.py、make_figures.py、reverse_sampling.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 mmd_noise_check.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """MMD^2 到底能不能拿来给采样器排序?——结论:在这个问题上不能。 写这篇的时候第一版主指标就是 MMD^2,跑出来的表看着很漂亮(Heun 在 N>=10 就顶到"噪声地板",其余采样器数值一路降到 0 以下)。问题是那个 "噪声地板"是单次抽样的运气值:把真实样本对真实样本重复 8 次,MMD^2 的标准差是 3.6e-4,比 N>=100 时各采样器之间的差别还大。 所以正文改用三个有闭式解、方差小得多的指标(mean log p_0 / 分量占比 TV / 模式内半径),MMD^2 只留作低步数区间的旁证。 跑法:python mmd_noise_check.py """ import numpy as np from forward_diffusion import sample_data from reverse_sampling import mmd2, run_ddpm, make_stride from forward_diffusion import alpha_bar_from_beta, linear_beta, T, D ab = alpha_bar_from_beta(linear_beta(T)) print("A. 真实样本 vs 真实样本(应当 ~0),重复 8 次,看估计量的散布") for bw in [None, 1.9, 1.0, 0.5]: vals = [] for k in range(8): a = sample_data(2000, np.random.default_rng(1000 + k)) b = sample_data(2000, np.random.default_rng(5000 + k)) v, used = mmd2(a, b, bw=bw) vals.append(v) vals = np.array(vals) print(f" bw={str(bw):>5} (实取{used:.2f}) mean={vals.mean():+.3e} " f"std={vals.std():.3e} min={vals.min():+.3e} max={vals.max():+.3e}") print() print("B. ddpm@1000 换 8 个不同的初始噪声种子") for bw in [None, 1.0, 0.5]: ref = sample_data(4000, np.random.default_rng(7)) vals = [] for k in range(8): rng = np.random.default_rng(3000 + k) x = rng.standard_normal((2000, D)) x = run_ddpm(x, ab, make_stride(1000, T), rng) v, used = mmd2(x, ref, bw=bw) vals.append(v) vals = np.array(vals) print(f" bw={str(bw):>5} (实取{used:.2f}) mean={vals.mean():+.3e} " f"std={vals.std():.3e}") print(f" -> {np.array2string(vals, precision=3, formatter={'float': lambda v: f'{v:+.2e}'})}") print() print("C. 极端对照:把样本整体平移 0.5,看各带宽的分辨力") ref = sample_data(4000, np.random.default_rng(7)) for shift in [0.05, 0.1, 0.2, 0.5]: y = sample_data(2000, np.random.default_rng(11)) + shift row = f" shift={shift:.2f} " for bw in [1.9, 1.0, 0.5, 0.25]: v, _ = mmd2(y, ref, bw=bw) row += f"bw={bw}:{v:+.2e} " print(row) forward_diffusion.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """扩散前向过程:闭式解、噪声调度、信噪比。 只依赖 numpy。正文里出现的每一个数字都由本脚本打印,不手填。 要回答三个问题: 1. q(x_t | x_0) 的闭式解是不是真的成立(蒙特卡洛对着验) 2. 线性调度和余弦调度把"难度"分配得有多不一样(看 alpha_bar / SNR) 3. 前向终点的分布离标准正态到底差多少(决定了 L_T 那一项有多大) 运行: python forward_diffusion.py """ import numpy as np SEED = 20260926 T = 1000 # 总步数,与 DDPM 原文一致 D = 2 # 玩具数据维度 # ── 玩具数据:2D 七分量各向同性高斯混合(环状 + 一个中心分量)────────── # 分量故意取紧(标准差 0.16~0.30):score 场的曲率大,离散化误差才显出来。 # 换成 3 个宽分量(标准差 ~0.5)的话,oracle score 下所有采样器都会直接顶到 # MMD 噪声地板上,分不出高下——这是调这个玩具问题时踩的第一个坑。 WEIGHTS = np.array([0.14, 0.13, 0.15, 0.12, 0.14, 0.13, 0.19]) MEANS = np.array([[-2.6, -1.0], [-1.0, -2.2], [1.4, -2.0], [2.6, -0.4], [1.6, 1.8], [-0.6, 2.4], [0.0, 0.0]]) SCALES = np.array([0.22, 0.18, 0.20, 0.16, 0.24, 0.20, 0.30]) def sample_data(n, rng): """从 p_0 采 n 个样本。""" k = rng.choice(len(WEIGHTS), size=n, p=WEIGHTS) return MEANS[k] + SCALES[k][:, None] * rng.standard_normal((n, D)) # ── 噪声调度 ────────────────────────────────────────────────────────── def linear_beta(T=T, beta_1=1e-4, beta_T=0.02): """DDPM 原文的线性调度,beta 从 1e-4 均匀涨到 0.02。""" return np.linspace(beta_1, beta_T, T) def cosine_beta(T=T, s=0.008, clip=0.999): """Nichol & Dhariwal 的余弦调度,返回 beta_t。 alpha_bar_t = cos^2(((t/T + s)/(1+s)) * pi/2) / cos^2((s/(1+s)) * pi/2) 分母只是为了让 t=0 时 alpha_bar=1。 按 t=T 代入会得到 alpha_bar_T = 0(cos(pi/2)=0),也就是最后一步 beta=1, 实操上会炸。原文和 diffusers 都会把 beta 截到 0.999 以内,这里照做。 """ t = np.arange(1, T + 1) / T f = np.cos(((t + s) / (1 + s)) * np.pi / 2) ** 2 f0 = (np.cos((s / (1 + s)) * np.pi / 2)) ** 2 a_bar = f / f0 beta = beta_from_alpha_bar(a_bar) return np.clip(beta, None, clip) def cosine_alpha_bar(T=T, s=0.008, clip=0.999): """截尾之后的余弦调度的 alpha_bar,与 cosine_beta 一致。""" return alpha_bar_from_beta(cosine_beta(T, s, clip)) def alpha_bar_from_beta(beta): """alpha_bar_t = prod_{s<=t} (1 - beta_s),t = 1..T。""" return np.cumprod(1.0 - beta) def beta_from_alpha_bar(alpha_bar): """反解出每步的 beta_t = 1 - alpha_bar_t / alpha_bar_{t-1}。""" prev = np.concatenate([[1.0], alpha_bar[:-1]]) return 1.0 - alpha_bar / prev def snr_db(alpha_bar): """信噪比 alpha_bar/(1-alpha_bar),取 10*log10。""" return 10.0 * np.log10(alpha_bar / (1.0 - alpha_bar)) def q_sample(x0, t_idx, alpha_bar, rng): """x_t = sqrt(alpha_bar_t) * x_0 + sqrt(1 - alpha_bar_t) * eps。 t_idx 是 0-based 的数组下标,对应 alpha_bar[t_idx]。 """ a = alpha_bar[t_idx] eps = rng.standard_normal(x0.shape) return np.sqrt(a) * x0 + np.sqrt(1.0 - a) * eps, eps # ── 加噪后分布的解析形式(反向采样要用 oracle score)────────────────── def noised_mixture(alpha_bar_t): """p_t 仍是高斯混合:第 k 个分量的均值缩到 sqrt(a)*mu_k, 协方差变成 a*Sigma_k + (1-a)*I。因为 Sigma_k = s_k^2 I,结果仍是各向同性。 """ a = alpha_bar_t means = np.sqrt(a) * MEANS var = a * SCALES ** 2 + (1.0 - a) # 每个分量的方差(标量) return WEIGHTS, means, var def log_density(x, weights, means, var): """各向同性高斯混合的 log 密度,x: [n, 2]。""" n, d = x.shape # [n, K] sq = ((x[:, None, :] - means[None, :, :]) ** 2).sum(-1) comp = -0.5 * (sq / var[None, :] + d * np.log(2 * np.pi * var)[None, :]) mx = comp.max(axis=1, keepdims=True) e = np.exp(comp - mx) mix = (weights[None, :] * e).sum(1) return (np.log(mix) + mx[:, 0]) def score_fn(x, alpha_bar_t): """nabla_x log p_t(x) 的闭式解。混合权重用 log-sum-exp 稳住数值。""" w, m, v = noised_mixture(alpha_bar_t) n, d = x.shape sq = ((x[:, None, :] - m[None, :, :]) ** 2).sum(-1) comp = -0.5 * (sq / v[None, :] + d * np.log(2 * np.pi * v)[None, :]) mx = comp.max(axis=1, keepdims=True) e = np.exp(comp - mx) resp = w[None, :] * e resp = resp / resp.sum(1, keepdims=True) # [n, K] 责任度 # grad log N(x; m_k, v_k I) = -(x - m_k)/v_k return -(resp[:, :, None] * (x[:, None, :] - m[None, :, :]) / v[None, :, None]).sum(1) # ────────────────────────────────────────────────────────────────────── # A. 闭式解核对 # ────────────────────────────────────────────────────────────────────── def check_closed_form(rng, alpha_bar, n=200_000): """一步一步加噪 1000 次,和闭式解 x_t = sqrt(a) x_0 + sqrt(1-a) eps 对着验。""" x0 = sample_data(n, rng) out = [] for t_idx in [49, 199, 499, 999]: a = alpha_bar[t_idx] # 路径 1:逐步迭代 x = x0.copy() beta = beta_from_alpha_bar(alpha_bar) for i in range(t_idx + 1): x = np.sqrt(1.0 - beta[i]) * x + np.sqrt(beta[i]) * rng.standard_normal(x.shape) # 路径 2:闭式解(用同一步里现造的噪声,保证逐样本可比) xt_cf, _ = q_sample(x0, t_idx, alpha_bar, rng) # 逐样本比不了(噪声不同),比统计量 out.append({ "t": t_idx + 1, "alpha_bar": a, "mean_step": np.abs(x.mean(0) - np.sqrt(a) * x0.mean(0)).max(), "var_step": np.abs(x.var(0) - (a * x0.var(0) + (1 - a))).max(), "var_cf": np.abs(xt_cf.var(0) - (a * x0.var(0) + (1 - a))).max(), }) return out # ────────────────────────────────────────────────────────────────────── # B. 调度对比 # ────────────────────────────────────────────────────────────────────── def schedule_report(): lin_b = linear_beta(T) lin_a = alpha_bar_from_beta(lin_b) cos_b = cosine_beta(T) cos_a = alpha_bar_from_beta(cos_b) rows = [] for t in [1, 50, 100, 250, 500, 750, 900, 1000]: rows.append({ "t": t, "lin_a": lin_a[t - 1], "cos_a": cos_a[t - 1], "lin_snr": snr_db(lin_a[t - 1]), "cos_snr": snr_db(cos_a[t - 1]), "lin_b": lin_b[t - 1], "cos_b": cos_b[t - 1], }) return lin_b, lin_a, cos_b, cos_a, rows def half_life(alpha_bar): """alpha_bar 掉到 0.5 是第几步——一半的采样步数花在这之后。""" idx = np.argmax(alpha_bar < 0.5) return int(idx) + 1 if alpha_bar[idx] < 0.5 else T # ────────────────────────────────────────────────────────────────────── # C. 区分边缘终端 KL 与 ELBO 的条件终端 KL # ────────────────────────────────────────────────────────────────────── def terminal_kl(alpha_bar_T): """KL( q(x_T) || N(0,I) ),q(x_T) 是七分量高斯混合。 没有闭式解,用蒙特卡洛:E_{q(x_T)}[ log q(x_T) - log N(0,I) ] """ rng = np.random.default_rng(SEED + 7) x = sample_data(400_000, rng) w, m, v = noised_mixture(alpha_bar_T) eps = rng.standard_normal(x.shape) xt = np.sqrt(alpha_bar_T) * x + np.sqrt(1.0 - alpha_bar_T) * eps lq = log_density(xt, w, m, v) ln = -0.5 * ((xt ** 2).sum(1) + D * np.log(2 * np.pi)) return float((lq - ln).mean()), float((lq - ln).std() / np.sqrt(len(xt))) def terminal_elbo_kl(alpha_bar_T): """E_data KL(q(x_T|x_0) || N(0,I)) 的闭式值。""" if not 0 <= alpha_bar_T < 1: raise ValueError("alpha_bar_T must lie in [0,1)") second_moment = np.sum(WEIGHTS * (np.sum(MEANS ** 2, axis=1) + D * SCALES ** 2)) return 0.5 * (alpha_bar_T * second_moment - D * alpha_bar_T - D * np.log1p(-alpha_bar_T)) def main(): rng = np.random.default_rng(SEED) print("=" * 68) print("A. 闭式解核对:逐步迭代 vs 一步闭式(20 万样本,2D)") print("=" * 68) lin_b, lin_a, cos_b, cos_a, _rows = schedule_report() for r in check_closed_form(rng, lin_a): print(f" t={r['t']:>4} alpha_bar={r['alpha_bar']:.6e} " f"|均值差|={r['mean_step']:.2e} 逐步方差误={r['var_step']:.2e} " f"闭式方差误={r['var_cf']:.2e}") print() print("=" * 68) print("B. 两种调度把信息怎么抹掉的") print("=" * 68) print(f" {'t':>5} {'alpha_bar(linear)':>18} {'alpha_bar(cosine)':>18} " f"{'SNR_dB(lin)':>12} {'SNR_dB(cos)':>12}") for r in _rows: print(f" {r['t']:>5} {r['lin_a']:>18.6e} {r['cos_a']:>18.6e} " f"{r['lin_snr']:>12.2f} {r['cos_snr']:>12.2f}") print(f"\n alpha_bar 掉到 0.5 的步数:linear = {half_life(lin_a)}," f"cosine = {half_life(cos_a)}") print(f" 也就是说 linear 调度下,{100 * (1 - half_life(lin_a) / T):.1f}% 的步数" f"花在信噪比已经低于 0 dB 的区域") print() print("=" * 68) print("C. 边缘终端 KL 和 ELBO 条件终端 KL(不同对象)") print("=" * 68) for name, aT in [("linear", lin_a[-1]), ("cosine", cos_a[-1])]: kl, se = terminal_kl(aT) print(f" {name:>7}: E_data KL(q(x_T|x_0)||N(0,I)) = {terminal_elbo_kl(aT):.6e} nats") print(f" {name:>7}: alpha_bar_T={aT:.6e} KL(q(x_T)||N(0,I)) = {kl:.6e} nats" f" (±{se:.1e})") print() print("=" * 68) print("D. 连续时间对账:beta(t) = T * beta_i,积分出来的 alpha_bar 对不对") print("=" * 68) b_min, b_max = T * float(lin_b[0]), T * float(lin_b[-1]) print(f" beta_min = {b_min:.4f}(Song et al. VP-SDE 默认 0.1)" f" beta_max = {b_max:.4f}(默认 20)") for t in [0.25, 0.5, 0.75, 1.0]: i = int(round(t * T)) - 1 a_cont = np.exp(-(b_min * t + 0.5 * (b_max - b_min) * t * t)) print(f" t={t:.2f} 离散 alpha_bar={lin_a[i]:.6e} " f"连续 exp(-int beta)={a_cont:.6e} 相对差={abs(lin_a[i] - a_cont) / lin_a[i]:.2e}") print() print("=" * 68) print("E. 终点还剩多少原始信号(非零终端 SNR 问题)") print("=" * 68) for name, aT in [("linear", lin_a[-1]), ("cosine", cos_a[-1])]: amp = np.sqrt(aT) snr_T = aT / (1.0 - aT) print(f" {name:>7}: alpha_bar_T={aT:.6e} sqrt(alpha_bar_T)={amp:.6e}" f" => x_T 里原始信号的振幅占比 {amp * 100:.4f}%") print(f" 终端 SNR = {snr_T:.4e}(= {snr_db(aT):.2f} dB)," f"理论值应当是 0") if __name__ == "__main__": main() score_bridges.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """把 score / epsilon / x_0 / v 四种参数化钉死在同一组恒等式上。 扩散模型的代码里同一个东西有四种写法,换参数化是家常便饭(SD 2.x 用 v, SD 1.x 和大部分视频模型用 epsilon,蒸馏论文里又爱用 x_0)。这篇要说的是: 它们描述的是同一个量,但**作为训练目标并不等价**——换参数化等于给每个 时间步偷偷换了一个权重。 跑法:python score_bridges.py """ import numpy as np from forward_diffusion import ( D, SEED, T, MEANS, SCALES, WEIGHTS, alpha_bar_from_beta, linear_beta, noised_mixture, sample_data, score_fn, ) ALPHA_BAR = alpha_bar_from_beta(linear_beta(T)) def abar(t): """t 是 1-based 步号。""" return 1.0 if t == 0 else float(ALPHA_BAR[t - 1]) def score_at(x, t): return score_fn(x, abar(t)) # ────────────────────────────────────────────────────────────────────── # A. 精确后验均值 E[x_0 | x_t] # ────────────────────────────────────────────────────────────────────── def posterior_mean_exact(x, t): """q(x_0 | x_t) 仍是高斯混合,第 k 个分量: 后验权重 r_k = 与 score 里用的是同一份责任度 后验方差 C_k = (1/s_k^2 + alpha_bar/(1-alpha_bar))^{-1} 后验均值 m_k = C_k * (mu_k / s_k^2 + sqrt(alpha_bar) * x_t / (1-alpha_bar)) """ a = abar(t) if a >= 1.0: return x.copy() _, m_t, v_t = noised_mixture(a) # 加噪后各分量的均值 / 方差 n = x.shape[0] sq = ((x[:, None, :] - m_t[None, :, :]) ** 2).sum(-1) logc = -0.5 * (sq / v_t[None, :] + D * np.log(2 * np.pi * v_t)[None, :]) mx = logc.max(1, keepdims=True) resp = WEIGHTS[None, :] * np.exp(logc - mx) resp = resp / resp.sum(1, keepdims=True) # [n, K] s2 = SCALES ** 2 C = 1.0 / (1.0 / s2 + a / (1.0 - a)) # [K] mk = C[None, :, None] * (MEANS[None, :, :] / s2[None, :, None] + np.sqrt(a) * x[:, None, :] / (1.0 - a)) return (resp[:, :, None] * mk).sum(1) def tweedie_mean(x, t): """Tweedie:E[x_0 | x_t] = (x_t + (1 - alpha_bar_t) * score) / sqrt(alpha_bar_t)。""" a = abar(t) if a >= 1.0: return x.copy() return (x + (1.0 - a) * score_at(x, t)) / np.sqrt(a) def eps_from_score(x, t): """eps = -sqrt(1 - alpha_bar_t) * score。""" a = abar(t) return -np.sqrt(1.0 - a) * score_at(x, t) def eps_from_x0(x, x0_hat, t): """反解:eps = (x_t - sqrt(alpha_bar_t) * x_0) / sqrt(1 - alpha_bar_t)。""" a = abar(t) return (x - np.sqrt(a) * x0_hat) / np.sqrt(1.0 - a) def v_target(x0, eps, t): """v = sqrt(alpha_bar) * eps - sqrt(1 - alpha_bar) * x_0。""" a = abar(t) return np.sqrt(a) * eps - np.sqrt(1.0 - a) * x0 def recover_from_v(x, v, t): """v 和 x_t 之间是一个旋转:x_0 = a*x_t - b*v,eps = b*x_t + a*v。 变换矩阵 [[a, b], [-b, a]] 行列式为 a^2+b^2=1,所以它是正交的—— 这也意味着 v 参数化不会放大噪声。 """ a, b = np.sqrt(abar(t)), np.sqrt(1.0 - abar(t)) return a * x - b * v, b * x + a * v # ────────────────────────────────────────────────────────────────────── # B. 蒙特卡洛交叉验证(不靠闭式解,纯重要性采样) # ────────────────────────────────────────────────────────────────────── def is_posterior_mean(xi, t, n=1_000_000, seed=SEED + 11): """E[x_0 | x_t = xi] 的自归一化重要性采样估计。 x_0^i 就是从 p_0 里采的,所以权重直接取 q(xi | x_0^i) 即可, 不需要知道归一化常数。 """ rng = np.random.default_rng(seed) x0 = sample_data(n, rng) a = abar(t) b2 = 1.0 - a sq = ((xi[None, :] - np.sqrt(a) * x0) ** 2).sum(1) logw = -0.5 * sq / b2 logw -= logw.max() w = np.exp(logw) ess = w.sum() ** 2 / (w ** 2).sum() return (w[:, None] * x0).sum(0) / w.sum(), ess # ────────────────────────────────────────────────────────────────────── # C. 三种参数化在 epsilon 空间下的每步权重 # ────────────────────────────────────────────────────────────────────── def loss_weights(t): """把三种参数化的训练损失都换算回"相当于给 eps 误差加了多大权重"。 - eps 参数化:delta_eps = delta_u -> 权重 1 - x_0 参数化:eps = (x_t - sqrt(a) x_0)/sqrt(1-a) delta_eps = -sqrt(a/(1-a)) * delta_u -> 权重 a/(1-a) = SNR 反过来,x_0 的损失 ||delta_u||^2 折算成 eps 误差是 ||delta_eps||^2 = SNR * ||delta_x0||^2, 即 x_0 损失对 eps 误差的权重是 1/SNR - v 参数化 :eps = b*x_t + a*v,delta_eps = a * delta_v -> v 损失折算成 eps 误差的权重是 1/a = 1/alpha_bar 返回 (权重_eps, 权重_v, 权重_x0),都是"乘在 ||delta_eps||^2 上的系数"。 """ a = abar(t) if t > 0 else 1.0 snr = a / (1.0 - a) return 1.0, 1.0 / a, 1.0 / snr def main(): rng = np.random.default_rng(SEED) print("=" * 72) print("A. Tweedie 恒等式核对:闭式后验均值 vs 公式 (x_t + (1-a) s)/sqrt(a)") print("=" * 72) print(f" {'t':>5} {'alpha_bar':>14} {'最大绝对误差':>16} {'x_t 范数均值':>14}") for t in [1, 50, 100, 250, 500, 750, 1000]: x = sample_data(2000, rng) xt, _ = (np.sqrt(abar(t)) * x + np.sqrt(1 - abar(t)) * rng.standard_normal(x.shape), None) e1 = posterior_mean_exact(xt, t) e2 = tweedie_mean(xt, t) err = np.abs(e1 - e2).max() print(f" {t:>5} {abar(t):>14.6e} {err:>16.3e} " f"{np.linalg.norm(xt, axis=1).mean():>14.4f}") print() print("=" * 72) print("B. eps 与 score 的换算:eps = -sqrt(1-a) * s,两条路算 eps 对不对") print("=" * 72) print(f" {'t':>5} {'|eps(score) - eps(x_0)| 最大':>28}") for t in [50, 250, 500, 1000]: x0 = sample_data(2000, rng) eps_true = rng.standard_normal(x0.shape) a = abar(t) xt = np.sqrt(a) * x0 + np.sqrt(1 - a) * eps_true # 路线 1:从 score e_s = eps_from_score(xt, t) # 路线 2:从 Tweedie 反解出的 x_0 e_x = eps_from_x0(xt, tweedie_mean(xt, t), t) print(f" {t:>5} {np.abs(e_s - e_x).max():>28.3e}") # 顺便看看 MMSE 估计量离真实 eps 有多远(这是"复原不了"的那部分) print(f" (MMSE 估计量与本次真实 eps 的 RMSE = " f"{np.sqrt(((e_s - eps_true) ** 2).sum(1).mean()):.4f}," f"sqrt(2D)={np.sqrt(2 * D):.4f} 是纯瞎猜的水平)") print() print("=" * 72) print("C. v 参数化:它与 (x_t, eps, x_0) 之间是一个旋转") print("=" * 72) for t in [50, 500, 1000]: x0 = sample_data(2000, rng) eps = rng.standard_normal(x0.shape) a = abar(t) xt = np.sqrt(a) * x0 + np.sqrt(1 - a) * eps v = v_target(x0, eps, t) x0_rec, eps_rec = recover_from_v(xt, v, t) print(f" t={t:>5} 重建 x_0 误差={np.abs(x0_rec - x0).max():.3e} " f"重建 eps 误差={np.abs(eps_rec - eps).max():.3e} " f"v 的范数均值={np.linalg.norm(v, axis=1).mean():.4f}") print() print("=" * 72) print("D. 蒙特卡洛交叉验证(100 万样本的重要性采样,不依赖上面的闭式解)") print("=" * 72) print(f" {'t':>5} {'Tweedie':>22} {'重要性采样':>22} {'有效样本数':>12}") for t in [50, 250, 500, 1000]: xi = sample_data(1, rng)[0] xi = np.sqrt(abar(t)) * xi + np.sqrt(1 - abar(t)) * rng.standard_normal(2) tm = tweedie_mean(xi[None, :], t)[0] im, ess = is_posterior_mean(xi, t) print(f" {t:>5} {np.array2string(tm, precision=4):>22} " f"{np.array2string(im, precision=4):>22} {ess:>12.0f}") print() print("=" * 72) print("E. 换参数化 = 给每个时间步换权重(折算回 eps 空间的系数)") print("=" * 72) print(f" {'t':>5} {'alpha_bar':>13} {'w_eps':>12} {'w_v':>14} {'w_x0':>14}") for t in [1, 10, 50, 100, 250, 500, 750, 900, 1000]: we, wv, wx = loss_weights(t) print(f" {t:>5} {abar(t):>13.6e} {we:>12.4f} {wv:>14.4e} {wx:>14.4e}") wes, wvs, wxs = zip(*[loss_weights(t) for t in range(1, T + 1)]) print(f"\n eps 参数化: 权重恒为 1,动态范围 {max(wes) / min(wes):.1f}") print(f" v 参数化: 权重 {min(wvs):.4e} ~ {max(wvs):.4e}," f"动态范围 {max(wvs) / min(wvs):.3e}") print(f" x_0 参数化: 权重 {min(wxs):.4e} ~ {max(wxs):.4e}," f"动态范围 {max(wxs) / min(wxs):.3e}") if __name__ == "__main__": main() make_figures.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """生成「扩散过程的前向与反向推导」的四张解释图。 数值一律从同目录的三个脚本里取(forward_diffusion / reverse_sampling / score_bridges),这里只负责画——改了那边这里要重跑,免得图和正文数字打架。 四张图分别回答: 1. 噪声调度到底把"难度"怎么分配到 1000 步上的 2. 五种采样器的误差随步数怎么降,以及每步跨过的积分量有多大 3. 前向把数据流推成球、反向沿轨迹走回七个模式,长什么样 4. 换参数化等价于给每个时间步换了多大权重(跨 10 个数量级) 只依赖 numpy + matplotlib。跑法:python make_figures.py """ import textwrap from pathlib import Path import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np from matplotlib.collections import LineCollection import forward_diffusion as FD import reverse_sampling as RS import score_bridges as SB ROOT = Path(__file__).resolve().parents[1] OUT = ROOT / "figures" OUT.mkdir(parents=True, exist_ok=True) plt.rcParams.update({ "font.sans-serif": ["PingFang SC", "Arial Unicode MS", "DejaVu Sans"], "axes.unicode_minus": False, "figure.dpi": 160, "savefig.bbox": "tight", }) INK = "#1f2937" MUTE = "#6b7280" C_LIN = "#d1495b" # 线性调度 / 差的一侧:红 C_COS = "#2f6fb0" # 余弦调度:蓝 C_DDPM = "#e0a03c" C_DDIM = "#2f9e6f" C_EM = "#8b5cf6" C_ODE = "#d1495b" C_HEUN = "#1f6feb" C_OK = "#2f9e6f" C_BAD = "#d1495b" def sci(v): """2.48e+08 -> 2.5e8,读起来省事。""" return f"{v:.1e}".replace("e+0", "e").replace("e+", "e").replace("e-0", "e-") def style(ax, title, xlabel=None, ylabel=None): ax.set_title(title, fontsize=11.5, color=INK, pad=10, loc="left") if xlabel: ax.set_xlabel(xlabel, fontsize=10, color=MUTE) if ylabel: ax.set_ylabel(ylabel, fontsize=10, color=MUTE) ax.tick_params(colors=MUTE, labelsize=9) for s in ("top", "right"): ax.spines[s].set_visible(False) for s in ("left", "bottom"): ax.spines[s].set_color("#d1d5db") ax.grid(alpha=0.25, linewidth=0.6) ax.set_axisbelow(True) def footer(fig, text, width=118): """把「这张图要看什么」放到坐标轴下方。 必须放在 y<0 的位置:bbox_inches="tight" 会把负坐标的 artist 一起收进来, 放在 0~0.05 之间的话会和 x 轴标签叠在一起。 """ wrapped = "\n".join(textwrap.wrap(text, width=width)) fig.text(0.012, -0.13, wrapped, fontsize=8.5, color=MUTE, va="top", ha="left", linespacing=1.6) # ────────────────────────────────────────────────────────────────────── # 图 1:噪声调度怎么分配难度 # ────────────────────────────────────────────────────────────────────── def fig_schedule(): lin_b, lin_a, cos_b, cos_a, rows = FD.schedule_report() ts = np.arange(1, FD.T + 1) snr_lin = FD.snr_db(lin_a) snr_cos = FD.snr_db(cos_a) fig, axes = plt.subplots(1, 2, figsize=(11.2, 4.1)) ax = axes[0] ax.plot(ts, lin_a, color=C_LIN, lw=2.0, label=r"linear $\beta$ 调度") ax.plot(ts, cos_a, color=C_COS, lw=2.0, label="余弦调度") ax.set_yscale("log") ax.axhline(0.5, color=MUTE, ls=":", lw=1.0) for x, c, lbl in [(FD.half_life(lin_a), C_LIN, f"linear 跌破 0.5:第 {FD.half_life(lin_a)} 步"), (FD.half_life(cos_a), C_COS, f"余弦跌破 0.5:第 {FD.half_life(cos_a)} 步")]: ax.axvline(x, color=c, ls="--", lw=1.0, alpha=0.8) ax.legend(fontsize=8.5, frameon=False, loc="lower left") style(ax, r"$\bar\alpha_t$:还剩下多少原始信号", "步数 t", r"$\bar\alpha_t$(对数轴)") ax = axes[1] ax.plot(ts, snr_lin, color=C_LIN, lw=2.0, label="linear") ax.plot(ts, snr_cos, color=C_COS, lw=2.0, label="余弦") ax.axhline(0.0, color=INK, lw=1.0) ax.fill_between(ts, snr_lin.min(), 0.0, where=(snr_lin < 0), color=C_LIN, alpha=0.10) i500 = 499 ax.annotate(f"t=500\nlinear {snr_lin[i500]:.1f} dB\n余弦 {snr_cos[i500]:.1f} dB", xy=(500, snr_lin[i500]), xytext=(620, -6), fontsize=8.5, color=INK, arrowprops=dict(arrowstyle="->", color=MUTE, lw=0.9)) ax.legend(fontsize=8.5, frameon=False, loc="upper right") style(ax, "信噪比:低于 0 dB 时信号已经被噪声淹没", "步数 t", "SNR (dB)") fig.suptitle("图 1 线性调度把 74% 的步数花在了信噪比已经低于 0 dB 的区域", fontsize=12, color=INK, x=0.012, ha="left", y=1.02) footer(fig, "要看什么:左图是「还剩多少信号」,右图是「信噪比」。线性调度在第 260 步就把一半信号丢光了," "后续处在较低 SNR 区间,但不等于无效计算;图中 SNR 以数据方差为 1 作参考。") fig.savefig(OUT / "schedule_snr.png") plt.close(fig) # ────────────────────────────────────────────────────────────────────── # 图 2:采样器误差随步数怎么降 + 每步跨过的积分量 # ────────────────────────────────────────────────────────────────────── def fig_samplers(res=None): if res is None: res, _ = RS.sweep() n_list = [10, 20, 25, 50, 100, 200, 1000] cols = {"ddpm": C_DDPM, "ddim": C_DDIM, "em": C_EM, "ode": C_ODE, "heun": C_HEUN} names = {"ddpm": "DDPM 祖采样", "ddim": "DDIM (η=0)", "em": "Euler–Maruyama (SDE)", "ode": "Euler (概率流 ODE)", "heun": "Heun 二阶 (ODE)"} fig, axes = plt.subplots(1, 2, figsize=(11.2, 4.3)) ax = axes[0] for name in ["ddpm", "ddim", "em", "ode", "heun"]: y = [abs(res[(name, n)]["dlogp"]) for n in n_list] ax.plot(n_list, y, marker="o", ms=4.5, lw=1.9, color=cols[name], label=names[name]) ax.set_xscale("log") ax.set_yscale("log") ax.axhline(0.0528, color=MUTE, ls="--", lw=1.0) ax.text(11, 0.062, "一次真实样本对照差值 0.053", fontsize=8, color=MUTE) ax.legend(fontsize=8, frameon=False, loc="upper right") style(ax, "单一统计量偏差(不能独立确认分布正确)", "采样步数 N(对数轴)", r"$|$mean log $p_0$ 偏移$|$(nats)") ax = axes[1] st = RS.stride_stats() ns = [r["n"] for r in st] ax.plot(ns, [r["L_max"] for r in st], marker="s", ms=4.5, lw=1.9, color=C_ODE, label=r"每步跨过的积分量 $L=\int\beta(u)du$(最大)") ax.plot(ns, [r["euler_err"] * 100 for r in st], marker="^", ms=4.5, lw=1.9, color=C_HEUN, label=r"Euler 近似 $e^{L/2}\approx 1+L/2$ 的相对误差") ax.set_xscale("log") ax.set_yscale("log") for r in st: if r["n"] in (10, 50): ax.annotate(f"{r['euler_err'] * 100:.1f}%", xy=(r["n"], r["euler_err"] * 100), xytext=(r["n"] * 1.15, r["euler_err"] * 100 * 1.6), fontsize=8.5, color=C_HEUN) ax.legend(fontsize=8, frameon=False, loc="upper right") style(ax, "为什么朴素 Euler 在大步长下会崩", "采样步数 N(对数轴)", "数值(对数轴,% 按数值读)") fig.suptitle("图 2 采样器诊断统计量与线性漂移的离散误差", fontsize=12, color=INK, x=0.012, ha="left", y=1.02) footer(fig, "要看什么:左图为单次实验,近零差异不作显著性排名;右图只比较线性漂移项的数值近似。DDPM/DDIM 把线性部分" "精确解掉了(直接乘 √ᾱ),Euler 却把它展开成一阶,N=10 时每步误差就有 25%。") fig.savefig(OUT / "sampler_scaling.png") plt.close(fig) return res # ────────────────────────────────────────────────────────────────────── # 图 3:前向流形被推成球,反向轨迹走回七个模式 # ────────────────────────────────────────────────────────────────────── def fig_trajectories(): rng = np.random.default_rng(FD.SEED) ab = FD.alpha_bar_from_beta(FD.linear_beta(FD.T)) x0 = FD.sample_data(600, rng) snaps = [0, 60, 200, 500, 1000] cols_f = plt.cm.viridis(np.linspace(0.05, 0.85, len(snaps))) # 反向:从同一批噪声出发,用 Heun 走 50 步,记录轨迹 x = rng.standard_normal((40, FD.D)) taus = RS.make_stride(50, FD.T) paths = [x.copy()] for j in range(len(taus) - 1): L = RS._log_step(ab, taus[j], taus[j + 1]) if L < 1e-12: continue s0 = RS.score_at(x, ab, taus[j]) d0 = 0.5 * L * (x + s0) x1 = x + d0 s1 = RS.score_at(x1, ab, taus[j + 1]) x = x + 0.5 * (d0 + 0.5 * L * (x1 + s1)) paths.append(x.copy()) fig, axes = plt.subplots(1, 2, figsize=(11.4, 4.9)) ax = axes[0] for t, c in zip(snaps, cols_f): a = 1.0 if t == 0 else float(ab[t - 1]) xt = np.sqrt(a) * x0 + np.sqrt(1 - a) * rng.standard_normal(x0.shape) ax.scatter(xt[:, 0], xt[:, 1], s=7, color=c, alpha=0.65, label=f"t={t}") ax.scatter(FD.MEANS[:, 0], FD.MEANS[:, 1], marker="x", s=70, color=INK, linewidths=1.6, label="分量中心") ax.legend(fontsize=8, frameon=False, loc="upper left", ncol=2) style(ax, "前向:七个模式被逐步抹成一个标准正态球", r"$x_1$", r"$x_2$") ax.set_aspect("equal") ax = axes[1] # 背景:p_0 的密度等高线 g = np.linspace(-4.6, 4.6, 220) GX, GY = np.meshgrid(g, g) grid = np.stack([GX.ravel(), GY.ravel()], 1) lp = FD.log_density(grid, FD.WEIGHTS, FD.MEANS, FD.SCALES ** 2) dens = np.exp(lp - lp.max()).reshape(GX.shape) ax.contourf(GX, GY, dens, levels=np.linspace(0.02, 1.0, 12), cmap="Blues", alpha=0.9, vmin=0.0, vmax=1.6) ax.contour(GX, GY, dens, levels=[0.05, 0.2, 0.5], colors="#1d4ed8", linewidths=0.7, alpha=0.55) P = np.stack(paths, 1) # [n, steps, 2] for i in range(P.shape[0]): seg = P[i] ax.add_collection(LineCollection( [seg[k:k + 2] for k in range(len(seg) - 1)], colors="#1f6feb", linewidths=0.8, alpha=0.5)) ax.scatter(P[:, 0, 0], P[:, 0, 1], s=14, color=MUTE, label="起点 x_T ~ N(0,I)") ax.scatter(P[:, -1, 0], P[:, -1, 1], s=16, color=C_OK, label="终点 x_0") ax.scatter(FD.MEANS[:, 0], FD.MEANS[:, 1], marker="x", s=70, color=INK, linewidths=1.6) ax.legend(fontsize=8, frameon=False, loc="upper left") style(ax, "反向:50 步 Heun,40 条轨迹从噪声回到模式", r"$x_1$", r"$x_2$") ax.set_aspect("equal") fig.suptitle("图 3 前向是「加水搅匀」,反向是「沿着 score 场把水滤掉」", fontsize=12, color=INK, x=0.012, ha="left", y=1.0) footer(fig, "要看什么:左边 t=200 时七个模式已经互相渗透,t=500 彻底成球;右边每条轨迹在最后十几步才" "「决定」进哪个模式——前面的漫长路程都在把粗轮廓搭起来,这解释了为什么大步长主要伤细节。") fig.savefig(OUT / "trajectories.png") plt.close(fig) # ────────────────────────────────────────────────────────────────────── # 图 4:换参数化 = 换每步权重 # ────────────────────────────────────────────────────────────────────── def fig_param_weights(): ts = np.arange(1, FD.T + 1) w = np.array([SB.loss_weights(t) for t in ts]) # [T, 3] fig, ax = plt.subplots(figsize=(7.6, 4.5)) ax.plot(ts, w[:, 0], color=C_OK, lw=2.4, label=r"$\varepsilon$ 参数化:恒为 1") ax.plot(ts, w[:, 1], color=C_DDPM, lw=2.0, label=r"$v$ 参数化:$1/\bar\alpha_t$") ax.plot(ts, w[:, 2], color=C_BAD, lw=2.0, label=r"$x_0$ 参数化:$1/\mathrm{SNR}_t$") ax.set_yscale("log") ax.axhline(1.0, color=MUTE, ls=":", lw=1.0) ax.annotate(rf"$x_0$ 跨 {sci(w[-1, 2] / w[0, 2])} 倍(1.0e-4 → 2.5e4)", xy=(1000, w[-1, 2]), xytext=(300, 6e2), fontsize=9, color=C_BAD, arrowprops=dict(arrowstyle="->", color=C_BAD, lw=0.9)) ax.annotate(rf"$v$ 跨 {sci(w[-1, 1] / w[0, 1])} 倍(1.0 → 2.5e4)", xy=(1000, w[-1, 1]), xytext=(300, 1.6e0), fontsize=9, color=C_DDPM, arrowprops=dict(arrowstyle="->", color=C_DDPM, lw=0.9)) ax.set_ylim(5e-5, 1e5) ax.legend(fontsize=9, frameon=False, loc="lower right") style(ax, "同一个模型,换参数化等于给每个时间步换权重", "步数 t", r"折算到 $\varepsilon$ 空间后的权重(对数轴)") fig.suptitle(r"图 4 $\varepsilon$ 参数化的权重恒为 1,$x_0$ 参数化跨了 8 个数量级", fontsize=12, color=INK, x=0.012, ha="left", y=1.03) footer(fig, "要看什么:三种参数化描述的是同一个量,但作为训练目标并不等价。$x_0$ 参数化在 t=1 处权重只有" " 1e-4、在 t=1000 处却有 2.5e4。这些是代数残差权重,不是实测参数梯度;不能据此单独判断 " "训练稳定性或最终质量。") fig.savefig(OUT / "param_weights.png") plt.close(fig) def main(): fig_schedule() res = fig_samplers() fig_trajectories() fig_param_weights() print(f"[OK] 四张图已写入 {OUT}") for f in sorted(OUT.glob("*.png")): print(f" {f.name} {f.stat().st_size / 1024:.0f} KB") if __name__ == "__main__": main() reverse_sampling.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """反向采样:五种采样器在"同一个 oracle score"下的正面对比。 这一节要回答的是:既然反向过程的 drift 只有唯一一种写法,为什么工业界 能搞出 DDPM / DDIM / DPM-Solver / Euler / Heun 这么多种采样器?差别到底 在哪里? 关键设计:目标分布用 2D 高斯混合,加噪之后仍然是高斯混合,所以 nabla_x log p_t(x) 有闭式解。也就是说这里的 score 是**理论最优**的, 不掺任何网络拟合误差——差异来自离散化、有限终端噪声近似以及蒙特卡洛波动。 五种采样器: ddpm DDPM 祖采样(ancestral),高斯反向核近似,每步注噪声 ddim DDIM,eta=0,确定性 em Euler-Maruyama 解反向 SDE,一阶,注噪声 ode Euler 解概率流 ODE,一阶,确定性 heun Heun 二阶解概率流 ODE,每步两次 score 评估 运行: python reverse_sampling.py """ import numpy as np from forward_diffusion import ( D, SEED, T, MEANS, SCALES, WEIGHTS, alpha_bar_from_beta, cosine_beta, linear_beta, log_density, sample_data, score_fn, ) # ────────────────────────────────────────────────────────────────────── # 时间步工具 # ────────────────────────────────────────────────────────────────────── def make_stride(n_steps, T=T): """把 0..T 均匀切成 n_steps 段,返回 1-based 的 t 序列(含 T 与 0)。 例:n_steps=50, T=1000 -> [1000, 980, 960, ..., 20, 0] """ idx = np.linspace(T, 0, n_steps + 1).astype(int) # 保证严格递减且唯一 idx = np.unique(idx)[::-1] return list(idx) def abar_at(alpha_bar, t): """t 是 1-based 步号;t=0 时 alpha_bar=1(即 x_0 本身)。""" return 1.0 if t == 0 else float(alpha_bar[t - 1]) def score_at(x, alpha_bar, t): return score_fn(x, abar_at(alpha_bar, t)) # ────────────────────────────────────────────────────────────────────── # 五种采样器 # ────────────────────────────────────────────────────────────────────── def _tweedie_x0(x, s, a_cur): """Tweedie 公式:E[x_0 | x_t] = (x_t + (1-a_t) * score) / sqrt(a_t)。""" return (x + (1.0 - a_cur) * s) / np.sqrt(a_cur) def _eps_from_score(s, a_cur): """epsilon 与 score 的换算:eps = -sqrt(1 - a_t) * score。""" return -np.sqrt(1.0 - a_cur) * s def run_ddpm(x, alpha_bar, taus, rng, **kw): """DDPM 祖采样:用精确反向均值配固定高斯方差,有限步并非精确逆核。 x_{t-1} = (x_t + beta_t * s) / sqrt(alpha_t) + sigma_t * z sigma_t^2 = beta_t * (1 - abar_{t-1}) / (1 - abar_t) """ for j in range(len(taus) - 1): t_cur, t_prev = taus[j], taus[j + 1] a_cur, a_prev = abar_at(alpha_bar, t_cur), abar_at(alpha_bar, t_prev) alpha_j = a_cur / a_prev beta_j = 1.0 - alpha_j if beta_j < 1e-12: continue s = score_at(x, alpha_bar, t_cur) mean = (x + beta_j * s) / np.sqrt(alpha_j) var = beta_j * (1.0 - a_prev) / (1.0 - a_cur) x = mean + np.sqrt(max(var, 0.0)) * rng.standard_normal(x.shape) return x def run_ddim(x, alpha_bar, taus, rng, **kw): """DDIM,eta=0:先跳到去噪后的 x_0,再重新加回 target 时刻的噪声。 x_{t-1} = sqrt(abar_{t-1}) * xhat_0 + sqrt(1 - abar_{t-1}) * eps_hat """ for j in range(len(taus) - 1): t_cur, t_prev = taus[j], taus[j + 1] a_cur, a_prev = abar_at(alpha_bar, t_cur), abar_at(alpha_bar, t_prev) s = score_at(x, alpha_bar, t_cur) xhat0 = _tweedie_x0(x, s, a_cur) eps_hat = _eps_from_score(s, a_cur) x = np.sqrt(a_prev) * xhat0 + np.sqrt(max(1.0 - a_prev, 0.0)) * eps_hat return x def _log_step(alpha_bar, t_cur, t_prev): """跨一步"累积"起来的积分量 L = int beta du = log(alpha_bar_{prev}/alpha_bar_{cur}) > 0。 这一步很关键:正向 SDE 的时间是 u = t/T,从 0 涨到 1;反向采样是让 u **往回走**,所以 du 是负的。把 du = -1/N 代进去之后,drift 的符号会整体 翻过来——反向过程是把被压扁的分布"吹"回原样,而不是继续压。 用 L 记这段区间上 beta 的积分,Euler 步就写成: dx = +L * (0.5 x + c * s) (c=1 是 SDE,c=1/2 是 ODE) """ a_cur, a_prev = abar_at(alpha_bar, t_cur), abar_at(alpha_bar, t_prev) return float(np.log(a_prev / a_cur)) def run_em(x, alpha_bar, taus, rng, **kw): """Euler-Maruyama 解反向 SDE:全量 score + 注入噪声。 反向 SDE 的 drift 里 score 的系数是 **1 倍**(不是 ODE 的半倍), 另外每步还要加 sqrt(L) 的噪声。 """ for j in range(len(taus) - 1): L = _log_step(alpha_bar, taus[j], taus[j + 1]) if L < 1e-12: continue s = score_at(x, alpha_bar, taus[j]) x = x + L * (0.5 * x + s) + np.sqrt(L) * rng.standard_normal(x.shape) return x def run_ode(x, alpha_bar, taus, rng, **kw): """Euler 解概率流 ODE:半倍 score,无噪声。 和 EM 只差两处:score 系数从 1 减到 0.5,噪声项整个去掉。 """ for j in range(len(taus) - 1): L = _log_step(alpha_bar, taus[j], taus[j + 1]) if L < 1e-12: continue s = score_at(x, alpha_bar, taus[j]) x = x + 0.5 * L * (x + s) return x def run_heun(x, alpha_bar, taus, rng, **kw): """Heun 二阶解概率流 ODE:Euler 预测一步,再用终点的 score 校正一次。 每步两次 score 评估(NFE = 2 * n_steps)。 """ for j in range(len(taus) - 1): L = _log_step(alpha_bar, taus[j], taus[j + 1]) if L < 1e-12: continue s0 = score_at(x, alpha_bar, taus[j]) d0 = 0.5 * L * (x + s0) x1 = x + d0 # Euler 预测 s1 = score_at(x1, alpha_bar, taus[j + 1]) # 在目标时刻再评估一次 d1 = 0.5 * L * (x1 + s1) x = x + 0.5 * (d0 + d1) return x SAMPLERS = { "ddpm": (run_ddpm, 1), "ddim": (run_ddim, 1), "em": (run_em, 1), "ode": (run_ode, 1), "heun": (run_heun, 2), } # ────────────────────────────────────────────────────────────────────── # 评价:MMD^2(RBF 核,median heuristic) # ────────────────────────────────────────────────────────────────────── def _rbf_kernels(a, b, bw): sa = (a ** 2).sum(1) sb = (b ** 2).sum(1) d2 = sa[:, None] + sb[None, :] - 2.0 * a @ b.T return np.exp(-d2 / (2.0 * bw ** 2)) def mmd2(x, y, bw=None, n_cap=1500): """MMD^2 的无偏估计。x 是生成样本,y 是真实样本。""" rng = np.random.default_rng(0) if len(x) > n_cap: x = x[rng.choice(len(x), n_cap, replace=False)] if len(y) > n_cap: y = y[rng.choice(len(y), n_cap, replace=False)] if bw is None: allp = np.vstack([x[:800], y[:800]]) d2 = ((allp[:, None, :] - allp[None, :, :]) ** 2).sum(-1) bw = float(np.sqrt(np.median(d2) / 2.0)) bw = max(bw, 1e-3) kxx = _rbf_kernels(x, x, bw) kyy = _rbf_kernels(y, y, bw) kxy = _rbf_kernels(x, y, bw) n, m = len(x), len(y) # 去掉对角线才无偏 t1 = (kxx.sum() - np.trace(kxx)) / (n * (n - 1)) t2 = (kyy.sum() - np.trace(kyy)) / (m * (m - 1)) t3 = kxy.mean() return float(t1 + t2 - 2 * t3), bw def mean_logp(x): """生成样本在真实 p_0 下的平均 log 密度。 比 MMD 更能抓一种特定的失败:采样器把样本堆到高密度区(过聚拢), 或者样本飘到分量之间的低密度地带。 """ return float(log_density(x, WEIGHTS, MEANS, SCALES ** 2).mean()) def assign_components(x): """按最近的均值把样本归到混合分量上(分量标准差同量级,够用)。""" d = ((x[:, None, :] - MEANS[None, :, :]) ** 2).sum(-1) return d.argmin(1) def tv_weights(x): """经验分量占比与真实权重之间的全变差距离。 抓的是"各模式的比例对不对"——模式丢了一个、或者某个模式被过度采样, 这个数会涨,而 mean log p 未必涨(甚至可能更漂亮)。 """ idx = assign_components(x) n = len(x) hist = np.bincount(idx, minlength=len(WEIGHTS)) / n return float(0.5 * np.abs(hist - WEIGHTS).sum()) def mode_spread_err(x): """每个分量内部样本的均方半径,与理论值 D*s_k^2 的相对误差(取各分量最大)。 抓的是"样本落在模式的中心但挤成一团"或者"散得太开"。 """ idx = assign_components(x) errs = [] for k in range(len(WEIGHTS)): sel = x[idx == k] if len(sel) < 30: errs.append(1.0) continue r2 = ((sel - MEANS[k]) ** 2).sum(1).mean() errs.append(abs(r2 / (D * SCALES[k] ** 2) - 1.0)) return float(max(errs)) def moment_err(x, y): """均值与协方差的绝对误差(作为 MMD 之外的直观补充)。""" return float(np.abs(x.mean(0) - y.mean(0)).max()), \ float(np.abs(np.cov(x.T) - np.cov(y.T)).max()) # ────────────────────────────────────────────────────────────────────── # 主实验 # ────────────────────────────────────────────────────────────────────── def sweep(n_list=(10, 20, 25, 50, 100, 200, 1000), n_samples=4000, schedule="linear", seed=SEED): """固定一份真实样本做参照,扫采样器 × 步数。""" alpha_bar = (alpha_bar_from_beta(linear_beta(T)) if schedule == "linear" else alpha_bar_from_beta(cosine_beta(T))) ref = sample_data(8000, np.random.default_rng(seed)) # 噪声地板:两份**独立**真实样本之间的 MMD^2。 # 单次值不是置信区间;细微差异需要重复抽样评估。 ref_b = sample_data(4000, np.random.default_rng(seed + 999)) floor, bw = mmd2(ref[:4000], ref_b) # 三个主指标都先在大样本真实数据上算一遍当基线,后面一律报"相对基线的偏移", # 理想期望下匹配的统计量差为0,有限样本仍有波动,0也不保证整个分布匹配。 base = { "logp": mean_logp(ref), "tv": tv_weights(ref), "spread": mode_spread_err(ref), } # 一次对照差值(不是标准误或置信区间) noise = { "logp": mean_logp(ref_b) - base["logp"], "tv": tv_weights(ref_b) - base["tv"], "spread": mode_spread_err(ref_b) - base["spread"], } results = {"_floor": floor, "_bw": bw, "_base": base, "_noise": noise} for name, (fn, mult) in SAMPLERS.items(): for n in n_list: rng = np.random.default_rng(seed + 1) x = rng.standard_normal((n_samples, D)) taus = make_stride(n, T) x = fn(x, alpha_bar, taus, rng) m2, _ = mmd2(x, ref, bw=bw) me, ce = moment_err(x, ref) results[(name, n)] = { "mmd2": m2, "nfe": n * mult, "mean_err": me, "cov_err": ce, "dlogp": mean_logp(x) - base["logp"], "dtv": tv_weights(x) - base["tv"], "dspread": mode_spread_err(x) - base["spread"], "finite": bool(np.isfinite(x).all()), } return results, bw def stride_stats(n_list=(10, 20, 50, 100, 200, 1000), T=T): """每一步跨过的"积分量" L = int beta du,以及 Euler 近似 e^{L/2} 会错多少。 DDPM / DDIM 的更新里线性部分是**精确**解掉的(直接乘 sqrt(alpha_bar)), Euler 却把它展开成一阶:e^{L/2} ≈ 1 + L/2。步长越大这两个差得越远, 这就是 N=10 时 Euler 系采样器全线崩掉、DDIM 却还能看的根本原因。 """ alpha_bar = alpha_bar_from_beta(linear_beta(T)) out = [] for n in n_list: taus = make_stride(n, T) Ls = np.array([_log_step(alpha_bar, taus[j], taus[j + 1]) for j in range(len(taus) - 1)]) lm = float(Ls.max()) err = float(np.abs(np.exp(lm / 2) - (1 + lm / 2)) / np.exp(lm / 2)) out.append({"n": n, "L_max": lm, "L_mean": float(Ls.mean()), "euler_err": err}) return out def ablation_noise(n=50, n_samples=4000, seed=SEED): """把 DDIM 的 eta 从 0 拉到 1,看"注入噪声"这一项单独值多少钱。 eta=0 -> DDIM(确定性);eta=1 -> 等价 DDPM 祖采样。 sigma_t = eta * sqrt(beta_t * (1 - abar_{t-1}) / (1 - abar_t)) """ alpha_bar = alpha_bar_from_beta(linear_beta(T)) ref = sample_data(8000, np.random.default_rng(seed)) base_logp, base_tv, bw0 = _ablation_baseline(seed) out = [] taus = make_stride(n, T) for eta in [0.0, 0.25, 0.5, 0.75, 1.0]: rng = np.random.default_rng(seed + 1) x = rng.standard_normal((n_samples, D)) for j in range(len(taus) - 1): t_cur, t_prev = taus[j], taus[j + 1] a_cur, a_prev = abar_at(alpha_bar, t_cur), abar_at(alpha_bar, t_prev) alpha_j = a_cur / a_prev beta_j = 1.0 - alpha_j if beta_j < 1e-12: continue s = score_at(x, alpha_bar, t_cur) xhat0 = _tweedie_x0(x, s, a_cur) eps_hat = _eps_from_score(s, a_cur) var = beta_j * (1.0 - a_prev) / (1.0 - a_cur) sig = eta * np.sqrt(max(var, 0.0)) coeff = np.sqrt(max(1.0 - a_prev - sig ** 2, 0.0)) x = np.sqrt(a_prev) * xhat0 + coeff * eps_hat + sig * rng.standard_normal(x.shape) m2, _ = mmd2(x, ref, bw=bw0) out.append({"eta": eta, "mmd2": m2, "dlogp": mean_logp(x) - base_logp, "dtv": tv_weights(x) - base_tv}) return out def _ablation_baseline(seed=SEED): ref = sample_data(8000, np.random.default_rng(seed)) _, bw0 = mmd2(ref[:2000], sample_data(2000, np.random.default_rng(seed + 3))) return mean_logp(ref), tv_weights(ref), bw0 def main(): print("=" * 74) print("主实验:同一个 oracle score,五种采样器 × 七档步数") print(f"目标:2D 七分量高斯混合;参照集 8000 真实样本;每档生成 {4000} 个") print("主指标是 dlogp / dTV / dspread 三个(都扣掉了真实样本自己的基线," "0 = 完美)") print("MMD^2 只列出来做旁证——它自身的噪声有 ±3.6e-4,见 mmd_noise_check.py") print("=" * 74) res, bw = sweep() floor = res["_floor"] base, noise = res["_base"], res["_noise"] n_list = [10, 20, 25, 50, 100, 200, 1000] print(f" MMD^2(噪声 ±3.6e-4,所以微小差异需重复抽样确认):") print(f" {'采样器':<8}" + "".join(f"{'N=' + str(n):>13}" for n in n_list)) for name in ["ddpm", "ddim", "em", "ode", "heun"]: row = f" {name:<8}" for n in n_list: row += f"{res[(name, n)]['mmd2']:>13.2e}" print(row) print() print(" 主指标 1 —— mean log p_0 相对真实样本的偏移(nats,0 = 完美):") print(f" {'采样器':<8}" + "".join(f"{'N=' + str(n):>13}" for n in n_list)) for name in ["ddpm", "ddim", "em", "ode", "heun"]: row = f" {name:<8}" for n in n_list: row += f"{res[(name, n)]['dlogp']:>13.3f}" print(row) print() print(" 主指标 2 —— 分量占比的全变差距离偏移(0 = 各模式比例都对):") print(f" {'采样器':<8}" + "".join(f"{'N=' + str(n):>13}" for n in n_list)) for name in ["ddpm", "ddim", "em", "ode", "heun"]: row = f" {name:<8}" for n in n_list: row += f"{res[(name, n)]['dtv']:>13.4f}" print(row) print() print(" 主指标 3 —— 模式内均方半径的相对误差(0 = 胖瘦都正好):") print(f" {'采样器':<8}" + "".join(f"{'N=' + str(n):>13}" for n in n_list)) for name in ["ddpm", "ddim", "em", "ode", "heun"]: row = f" {name:<8}" for n in n_list: row += f"{res[(name, n)]['dspread']:>13.4f}" print(row) print() print("=" * 74) print("为什么 DDIM 大步长扛得住、Euler 扛不住:看每步跨过的 L 有多大") print("=" * 74) print(f" {'N':>6} {'L_max':>10} {'L_mean':>10} {'e^(L/2) 与 1+L/2 的相对误差':>30}") for r in stride_stats(): print(f" {r['n']:>6} {r['L_max']:>10.4f} {r['L_mean']:>10.4f} " f"{r['euler_err'] * 100:>27.2f}%") print(f"\n 基线:真实样本 mean log p_0 = {base['logp']:.4f}," f"TV = {base['tv']:.4f},spread err = {base['spread']:.4f}") print(f" 指标自身噪声(另一份独立真实样本对基线的偏移):" f"logp {noise['logp']:+.4f} / TV {noise['tv']:+.4f} / " f"spread {noise['spread']:+.4f}") print() print(" 按 NFE(score 评估次数)对齐再看一遍——Heun 每步算两次:") print(f" {'采样器':<8} {'N':>6} {'NFE':>6} {'dlogp':>11} {'dspread':>11}") for name, n in [("ddpm", 1000), ("ddpm", 200), ("ddim", 50), ("ddim", 100), ("em", 50), ("ode", 50), ("heun", 25), ("heun", 50), ("ode", 200)]: r = res[(name, n)] print(f" {name:<8} {n:>6} {r['nfe']:>6} {r['dlogp']:>11.4f} " f"{r['dspread']:>11.4f}") print() print("=" * 74) print("关键对照:50 步能追上 1000 步吗(主指标看 dlogp,0 = 与真实样本一致)") print("=" * 74) for name, n in [("ddpm", 1000), ("ddim", 50), ("heun", 50), ("ode", 50), ("em", 50), ("ddpm", 50)]: r = res[(name, n)] print(f" {name:<6} N={n:<5} NFE={r['nfe']:<5} " f"dlogp={r['dlogp']:+.4f} dTV={r['dtv']:+.4f} " f"dspread={r['dspread']:+.4f} MMD^2={r['mmd2']:+.2e}") print() print("=" * 74) print("消融:只改 eta(注入多少噪声),其余完全不动,N=50") print("=" * 74) print(f" {'eta':>6} {'dlogp':>12} {'dTV':>12} {'MMD^2':>13}") for r in ablation_noise(): print(f" {r['eta']:>6.2f} {r['dlogp']:>12.4f} {r['dtv']:>12.4f} " f"{r['mmd2']:>13.2e}") print() print("=" * 74) print("换余弦调度再跑一遍(N=50)") print("=" * 74) res_c, _ = sweep(n_list=(50,), schedule="cosine") print(f" {'采样器':<8} {'dlogp(linear)':>15} {'dlogp(cosine)':>15}") for name in ["ddpm", "ddim", "em", "ode", "heun"]: print(f" {name:<8} {res[(name, 50)]['dlogp']:>15.4f} " f"{res_c[(name, 50)]['dlogp']:>15.4f}") print() print(f" (参照:MMD^2 噪声地板 = {floor:.3e},RBF 带宽 = {bw:.3f};" f"真实样本 mean log p_0 = {base['logp']:.4f})") if __name__ == "__main__": main()
2026年09月27日
4 阅读
0 评论
0 点赞
2026-09-26
AIGC 基本功|VAE 结构与训练目标-VAE
VAE 结构、KL 权重,与那个神秘的 0.18215 所属方向:表征与压缩 | 难度:进阶 | 前置知识:变分下界与重参数化(本篇是它的直接后继,ELBO 与重参数化在那里推过,这里只用结论) 关键词:VAE、编码器、解码器、潜空间、KL 权重、后验坍缩、scaling factor 01. 为什么需要它 先看一个可复现的尺度错误:预训练扩散模型要求 VAE latent 乘以 vae.config.scaling_factor,却把原始 latent 直接送进去。张量形状仍正确,但它已经偏离训练分布。下面用一维高斯去噪器量化这个误差;这些数字来自教学模型,不是 Stable Diffusion 图像质量实验,不能据此断言真实图像一定发灰或过曝。 这不是玄学,可以精确算出来。设潜空间真实标准差是 $\sigma$,扩散模型的噪声表却是按「数据方差为 1」标定的。在完全干净那一端($\bar{\alpha} \to 1$)两边都对;越往噪声端走,误差越大。用高斯 MMSE 估计可以算出,模型的去噪幅度只有正确值的 $$r(\bar{\alpha}) = \bar{\alpha} + \frac{1 - \bar{\alpha}}{\sigma^{2}}$$ 倍。$\sigma = 5.49$(这是由 SD1.x 常用配置系数反推的尺度,第 03.4 节会讲为什么)时,$\bar{\alpha} = 0.5$ 处 $r = 0.5166$——幅度只剩一半。实测的均方误差从本该有的 0.9682 涨到 7.8154,恶化 8.07 倍。这是高斯教学模型的估计误差,不是图像实测。 反着错也一样疼。如果你的 VAE 是规规矩矩训的(潜空间方差约等于 1),却照抄 Stable Diffusion 的 0.18215,那么 $\sigma$ 变成 0.18215,同一个 $\bar{\alpha}=0.5$ 处 $r = 15.5699$——幅度被放大 15.6 倍,教学去噪器的后验均值幅度被高估。这个常数抄错方向,比抄错符号更常见。 还有第三种错法,比前两种隐蔽得多。LDM 论文(arXiv:2112.10752)附录 D.1 里有一段原话,把它写得很清楚: the signal-to-noise ratio induced by the variance of the latent space (i.e. $\text{Var}(z)/\sigma_t^{2}$) significantly affects the results for convolutional sampling ... when training a LDM directly in the latent space of a KL-regularized model, this ratio is very high, such that the model allocates a lot of semantic detail early on in the reverse denoising process ... Note that the VQ-regularized space has a variance close to 1, such that it does not have to be rescaled. 这段观察针对 LDM 论文中具体的 KL / VQ 自编码器:KL latent 的高方差改变了给定噪声调度下的信噪比,影响高分辨率卷积采样;论文所用 VQ latent 的方差接近 1。它不表示所有 VQ 码本天然归一化,也不表示尺度错误一定对应某一种视觉伪影。 最后是 KL 权重本身的坑。在我的最小实验里,把 KL 权重 $\beta$ 从 0.1 调到 1,重建 MSE 从 0.2119 跳到 1.0000——1.0000 就是「什么都不学、直接输出均值」的分数(数据逐维方差已被归一化成 1)。潜变量的 8 个维度全部死掉。这个现象叫后验坍缩。 这一篇就把这三件事串起来:VAE 的结构决定了潜空间里有什么,KL 权重决定了还剩下什么,而剩下东西的尺度就是那个 0.18215。 02. 最小可用理解 三句话: 结构:编码器把 $x$ 压成 $2d$ 个数($d$ 个均值 $\mu$、$d$ 个对数方差 $\log\sigma^{2}$),重参数化采出一个 $z$,解码器从 $z$ 重建 $x$。训练目标是 ELBO——本质上是「重建质量」减「每个潜变量维度花掉的 KL 预算」。 KL 项既是正则,也可以看作信息预算。一维潜变量要花掉多少 KL,就必须换回足够的重建收益,否则最优解就是关掉这一维。实测里 $\beta = 10^{-3}$ 时模型正好活 4 个维度(合成数据的真实因子数就是 4),$\beta$ 收到 0.3 只剩 3 个,$\beta = 1$ 一个不剩。维度是一个个死的,不是一起死。 潜空间的尺度是 KL 权重留下的痕迹。$\beta$ 大到 KL 有效时,它把潜空间边际标准差钉在 1 附近(实测 1.0067);$\beta$ 小到 KL 失效,尺度就失去约束,随训练动力学漂走(实测漂到 2.7174)。Stable Diffusion 那套 VAE 的 KL 权重是 $10^{-6}$(LDM 论文原话:we either weight the KL term by a factor $\sim 10^{-6}$),其配置对应的原始 latent 标准差约为 5.49,不能仅凭 KL 系数推断具体漂移过程——而 $1/5.49 = 0.18215$。 这张图要看什么:四张子图连起来读。(a) 重建 MSE 在 $\beta$ 超过 0.1 后逐渐上升到 1.0 那条虚线,说明模型彻底放弃潜变量;(b) 总 KL 同步归零,先验和近似后验重合;(c) 潜空间边际标准差:蓝线(无权重衰减)在 $\beta \ge 10^{-3}$ 后紧紧贴着 1,$\beta$ 一小就抬头,红线(有权重衰减)在同一条路上走得更快更远,绿虚线是 Stable Diffusion 的 5.49;(d) 还活着的维度数从 8 一路掉到 0,中间在 4 这个地方有个明显的台阶——那是合成数据的真实因子数。 03. 数学推导 3.1 一次编码到底出了什么 设数据 $x \in \mathbb{R}^{D}$,潜变量 $z \in \mathbb{R}^{d}$。VAE 的编码器不输出一个 $z$,它输出一个分布: $$q_{\phi}(z \mid x) = \mathcal{N}\!\left(z;\ \mu_{\phi}(x),\ \text{diag}\left(\sigma_{\phi}^{2}(x)\right)\right)$$ 网络实际算出来的是 $\mu$ 和对数方差 $\text{logvar}$ 两组数,各 $d$ 个。为什么是 logvar 而不是方差:网络输出可以取任意实数,但方差必须为正。把网络的输出过一层 $\exp$ 就得到恒正的方差,同时把它放进对数域还有一个好处——数值范围。实验中 $\beta = 10^{-6}$ 时后验标准差会掉到 0.009 量级,方差是 $8\times 10^{-5}$;如果网络直接回归方差,这个量级上梯度会和重建项的梯度混在一起互相淹没,换到对数域后它就只是一个普通的负实数。 从 $q_{\phi}(z|x)$ 里采一个 $z$ 直接用会断掉梯度(采样操作不可导)。重参数化把它挪出去: $$z = \mu + \sigma \odot \epsilon, \qquad \epsilon \sim \mathcal{N}(0, I)$$ $\odot$ 是逐元素乘。随机性全部塞进 $\epsilon$ 里,$\mu$ 和 $\sigma$ 变成普通可导函数——这是整篇 VAE 能训起来的前提。 解码器给的是似然,取对角高斯: $$p_{\theta}(x \mid z) = \mathcal{N}\!\left(x;\ \text{dec}_{\theta}(z),\ \sigma_{\text{dec}}^{2} I\right)$$ 取对数: $$\log p_{\theta}(x \mid z) = -\frac{\lVert x - \text{dec}_{\theta}(z) \rVert^{2}}{2\sigma_{\text{dec}}^{2}} - \frac{D}{2}\log\left(2\pi\sigma_{\text{dec}}^{2}\right)$$ 第二项和 $z$ 无关,是常数。所以重建项就是平方误差之和除以 $2\sigma_{\text{dec}}^{2}$——「重建用 MSE」不是拍脑袋定的,它是高斯似然的必然结果,而 $\sigma_{\text{dec}}$ 就是重建项的隐式权重。日志里那种 recon + beta * kl 的写法,还要明确 reduction:若 recon 是逐元素均值,则标准负 ELBO 按同尺度缩放后的 KL 系数为 $2\sigma_{\text{dec}}^{2}/D$。 把两件事拼起来,ELBO 说 $\log p(x) \ge \mathbb{E}_{q}\left[\log p_{\theta}(x|z)\right] - \text{KL}\!\left(q_{\phi}(z|x)\,\Vert\,p(z)\right)$,先验取标准正态 $p(z) = \mathcal{N}(0, I)$。我们要最小化的就是负 ELBO: $$J = \frac{\mathbb{E}_{q}\left[\lVert x - \text{dec}_{\theta}(z) \rVert^{2}\right]}{2\sigma_{\text{dec}}^{2}} + \beta \sum_{j=1}^{d} \text{KL}_{j}, \qquad \text{KL}_{j} = \text{KL}\!\left(\mathcal{N}(\mu_{j}, \sigma_{j}^{2})\,\Vert\,\mathcal{N}(0,1)\right)$$ $\beta$ 是我们插进去的旋钮:$\beta = 1$ 是标准 ELBO,$\beta > 1$ 就是 $\beta$-VAE 路线,调大换解耦表征。注意 KL 是对维度求和而不是求平均,这一点第 05 节还会回来算账。 3.2 KL 的闭式解,逐项推 两个对角高斯之间的 KL 有闭式解。按定义 $\text{KL}(q \Vert p) = \mathbb{E}_{q}[\log q] - \mathbb{E}_{q}[\log p]$ 分头算。因为逐维独立,下面只看第 $j$ 维。 先算 $\mathbb{E}_{q}[\log q]$。$q_{j} = \mathcal{N}(\mu_{j}, \sigma_{j}^{2})$,所以 $$\log q_{j}(z_{j}) = -\frac{1}{2}\log(2\pi\sigma_{j}^{2}) - \frac{(z_{j} - \mu_{j})^{2}}{2\sigma_{j}^{2}}$$ 对 $q_{j}$ 取期望时,右边第二项的期望是 $\sigma_{j}^{2} / (2\sigma_{j}^{2}) = 1/2$,于是 $$\mathbb{E}_{q_{j}}\left[\log q_{j}\right] = -\frac{1}{2}\left(1 + \log 2\pi + \log \sigma_{j}^{2}\right)$$ 只有三项,和 $\mu_{j}$ 无关——这一点值得停一下:在 $\mu_{j}$ 上平移一个高斯分布,它的熵不变。 再算 $\mathbb{E}_{q}[\log p]$。先验 $p_{j} = \mathcal{N}(0,1)$,同样展开 $$\mathbb{E}_{q_{j}}\left[\log p_{j}\right] = -\frac{1}{2}\log 2\pi - \frac{\mathbb{E}_{q_{j}}\left[z_{j}^{2}\right]}{2}$$ 这里用到 $z_{j} = \mu_{j} + \sigma_{j}\epsilon$,所以 $\mathbb{E}_{q_{j}}[z_{j}^{2}] = \mu_{j}^{2} + \sigma_{j}^{2}$——这就是 $z$ 的二阶矩,它才是把 $\mu$ 拉进公式的那一项。 两者相减,$-\frac{1}{2}\log 2\pi$ 正好抵消: $$\text{KL}_{j} = \frac{1}{2}\left(\mu_{j}^{2} + \sigma_{j}^{2} - \log \sigma_{j}^{2} - 1\right)$$ 对 $j$ 求和即得总 KL。这个式子里每一块都有明确的物理含义,逐项读: $\mu_{j}^{2}$ 是均值偏离先验中心的成本。 不同输入的均值发生变化可以传递信息,但这一项不是互信息本身。即使所有输入的均值都是 0,只要方差仍依赖输入,潜变量也可能携带信息;只有整个条件分布都与输入无关时,这一维才不传信息。逐维 KL 同时包含信息代价与聚合后验偏离先验的代价,不能把 8~12 nats 直接叫作有效信息量。 3.3 权重到底在权衡什么:一维线性 VAE 的闭式解 上面的直觉可以算到精确解。把模型简化到最狠:一维数据 $x \sim \mathcal{N}(0, v)$,线性编码器 $q(z|x) = \mathcal{N}(mx, s^{2})$($m$ 是缩放系数,$s^{2}$ 是固定的后验方差),线性解码器 $p(x|z) = \mathcal{N}(wz, \sigma_{\text{dec}}^{2})$。目标函数展开成 $$J(w, m, s) = \frac{v(1 - wm)^{2} + w^{2} s^{2}}{2\sigma_{\text{dec}}^{2}} + \frac{\beta}{2}\left(m^{2} v + s^{2} - \log s^{2} - 1\right)$$ 第一项里的 $v(1-wm)^{2}$ 是「编码-解码这条路的增益偏离 1 有多远」,$w^{2}s^{2}$ 是「采样噪声被放大 $w$ 倍后落在输出上的方差」;第二项就是上一节的 KL。三个未知数各求一次偏导并置零: $$\frac{\partial J}{\partial w} = \frac{-v m(1 - wm) + w s^{2}}{\sigma_{\text{dec}}^{2}} = 0, \qquad \frac{\partial J}{\partial m} = \frac{-v w(1 - wm)}{\sigma_{\text{dec}}^{2}} + \beta m v = 0, \qquad \frac{\partial J}{\partial s} = \frac{w^{2} s}{\sigma_{\text{dec}}^{2}} + \beta\left(s - \frac{1}{s}\right) = 0$$ 记路增益 $u := wm$(它就是「潜变量被真正使用的程度」),把 $\partial_m$ 的式子两边乘 $w$ 换成 $u$,可以整理出 $w^{2}(1-u) = \beta \sigma_{\text{dec}}^{2} u$;再从 $\partial_w$ 得到 $w^{2} s^{2} = v u (1-u)$;从 $\partial_s$ 得到 $s^{2}\left(\beta + 2c w^{2}\right) = \beta$(这里 $2c = 1/\sigma_{\text{dec}}^{2}$)。三式联立,令 $$\lambda := \frac{\beta \sigma_{\text{dec}}^{2}}{v}$$ 解出来是一组非常干净的东西: $$u = 1 - \lambda, \qquad s^{2} = \lambda, \qquad w^{2} = v(1 - \lambda)$$ 代回验一遍:$\partial_m$ 要求 $w(1-u)/\sigma_{\text{dec}}^{2}=\beta m$。乘以 $w$,代入 $w^2=v(1-\lambda)$ 和 $1-u=\lambda$,得到 $v(1-\lambda)\lambda/\sigma_{\text{dec}}^{2}=\beta(1-\lambda)$。在未坍缩区间约去 $1-\lambda$,正好得到 $\lambda=\beta\sigma_{\text{dec}}^{2}/v$。第 04.2 节的数值优化与此吻合。 $\lambda$ 的读法:它是 KL 权重乘观测噪声方差,再除以数据方差的无量纲比值。这个比例决定一切: $u = 1-\lambda$ 随 $\lambda$ 线性下降。$\lambda \to 0$ 时 $u \to 1$、$s^{2} \to 0$,退化成确定性自编码器——潜变量满负荷工作,采样噪声归零。 $s^{2}=\lambda$ 随权重变化;$\beta=1$ 时等于这个线性高斯模型的真实后验方差,其他权重一般对应不同的变分目标。 当 $\lambda>1$ 时上述分支要求 $w^2<0$,不存在实数解;$\lambda=1$ 时它连续接到坍缩点,最优解退化成坍缩点 $(w, m, s) = (0, 0, 1)$。*坍缩阈值是 $\beta^{} = v / \sigma_{\text{dec}}^{2}$**:数据方差越大、或者解码器的观测噪声越小,需要的 $\beta$ 就越大。这是线性高斯模型的阈值,不是神经 VAE 的通用保证。 附带一个值得记住的事实:$(0,0,1)$ 在任意 $\beta$ 下都是驻点(数值验证梯度恒等于 0)。它只是当 $\lambda < 1$ 时不是最小值。所以「后验坍缩」不是数值 bug、不是训练不充分——在本线性模型的阈值以上,它是该目标的最优解;一般神经模型也可能因局部最优和优化动力学而坍缩。 3.4 潜空间的统计性质,和那个 0.18215 现在把「编码器」反过来看会得到什么分布。训练完之后,把所有 $x$ 编码一遍,潜变量的边际分布是 $$q(z) = \int q_{\phi}(z \mid x)\, p(x)\, \mathrm{d}x$$ 这是个混合分布。对每一维分别算方差,用全方差公式: $$\text{Var}(z_{j}) = \underbrace{\text{Var}_{p_{\text{data}}}\!\left[\mu_{j}(x)\right]}_{\text{信息}} + \underbrace{\mathbb{E}_{p_{\text{data}}}\!\left[\sigma_{j}^{2}(x)\right]}_{\text{采样噪声}}$$ 左边是 scaling_factor 要归一化的东西,右边两项来源完全不同:前一项是「不同样本被编码到不同位置」,后一项是「每个样本自己抖多少」。下表先对每维方差取平均再开方,不能先平均标准差再平方,也不能把两项标准差直接相加。方差占比不是互信息。实测拆账($\beta$ 从小到大): $\beta$ 均值变化标准差 RMS 后验噪声标准差 RMS 边际标准差 RMS 总方差中均值变化占比 1e-06 2.060 0.009 2.060 100% 1e-05 1.411 0.014 1.411 100% 0.0001 1.004 0.355 1.065 89% 0.001 0.726 0.697 1.007 52% 0.01 0.705 0.717 1.005 49% 0.1 0.620 0.779 0.995 39% 0.3 0.407 0.914 1.000 17% 1 0.003 1.000 1.000 0% $\beta = 10^{-3}$ 附近有个交叉点:潜空间方差里「信息」和「噪声」各占一半。往左走,$\text{std}(z)$ 几乎全部来自信息;往右走,几乎全部来自采样噪声——$\beta = 1$ 时 $z$ 就是一坨噪声,$\text{std}_x(\mu)$ 只剩 0.003。 那么 scaling_factor 是什么?在这里它是原始 latent 标准差的倒数。LDM 的 原始标定代码 使用 1. / z.flatten().std();diffusers 文档中的标准差描述要结合实际乘法方向理解。本文 toy 的跨维 RMS 也不是一般情况下与展平标准差完全相等:后者还包含不同通道均值的差异。 $$z_{\text{scaled}} = s \cdot z, \qquad s = \frac{1}{\text{std}(z)}$$ Stable Diffusion 取 $s = 0.18215$,倒过来就是 $\text{std}(z) = 5.4900$、$\text{Var}(z) = 30.1399$。为什么需要这一步,LDM D.1 用信噪比的语言回答了:扩散模型的噪声表 $\sigma_t$ 是相对数据尺度标定的,把这个比值写出来 $$\text{SNR} = \frac{\text{Var}(z)\,\bar{\alpha}}{1 - \bar{\alpha}}$$ 模型以为的是 $\bar{\alpha}/(1-\bar{\alpha})$,实际却差 $\text{Var}(z) = 30.14$ 倍,也就是 14.79 dB。每一步模型都以为「噪声占了这么多」,实际噪声只有它以为的 1/30。 最后一个式子把「差 30 倍」翻译成「图发灰」。假设扩散模型是在单位方差潜空间上训好的,那它的去噪器就是一个高斯 MMSE 估计:先验 $z_{0} \sim \mathcal{N}(0, \sigma_{p}^{2} I)$,观测 $z_{t} = \sqrt{\bar{\alpha}}\,z_{0} + \sqrt{1-\bar{\alpha}}\,\epsilon$,后验均值是 $$\hat{z}_{0} = \frac{\sigma_{p}^{2}\sqrt{\bar{\alpha}}}{\sigma_{p}^{2}\bar{\alpha} + 1 - \bar{\alpha}}\, z_{t}$$ 代入 $\sigma_{p}^{2} = 1$ 得 $\hat{z}_{0} = \sqrt{\bar{\alpha}}\, z_{t}$;而真实潜空间的标准差是 $\sigma$,正确的 MMSE 估计要用 $\sigma_{p}^{2} = \sigma^{2}$。两者相除就是第 01 节那个 $r(\bar{\alpha}) = \bar{\alpha} + (1-\bar{\alpha})/\sigma^{2}$。$\sigma = 5.49$ 时它在 $\bar{\alpha} \to 0$ 处趋于 $1/\sigma^{2} = 0.033$——重建幅度只剩 3.3%,图当然是灰的。 04. 代码实现 完整脚本在文末附录(vae_minimal.py、collapse_closed_form.py、latent_scaling.py、make_figures.py),只依赖 numpy 与 matplotlib,全部用 /usr/local/bin/python3 实跑过,下面每个数字都是真实输出。环境里没有 torch,所以反向传播是手推的——这反而更好,公式和代码能一行行对上。 4.1 最小 VAE 与两轮对照实验 模型就是 3.1 节那套,符号一一对应: def forward(self, x, eps): p = self.p h1 = np.maximum(x @ p["W1"] + p["b1"], 0.0) mu = h1 @ p["Wmu"] + p["bmu"] # 后验均值 lv_raw = h1 @ p["Wlv"] + p["blv"] lv = np.clip(lv_raw, LOGVAR_MIN, LOGVAR_MAX) # logvar 截断,防 exp 溢出 sig = np.exp(0.5 * lv) # 后验标准差 sigma z = mu + sig * eps # 重参数化 g1 = np.maximum(z @ p["V1"] + p["c1"], 0.0) xhat = g1 @ p["V2"] + p["c2"] return xhat, dict(x=x, h1=h1, mu=mu, lv=lv, lv_raw=lv_raw, sig=sig, z=z, g1=g1, xhat=xhat, eps=eps) @staticmethod def losses(x, xhat, mu, lv, sig): recon = float(np.mean((xhat - x) ** 2)) # 重建:MSE kl_dim = 0.5 * np.mean(mu ** 2 + sig ** 2 - lv - 1.0, axis=0) # 逐维 KL(3.2 节那个式子) return recon, kl_dim, float(np.sum(kl_dim)) # 对维度求和 数据是合成的:64 个观测维度由 4 个高斯因子线性混合而成,再逐维归一化到方差 1(所以「重建 MSE = 1」= 什么都没学到)。潜变量给了 8 维,故意比真实因子数多一倍,看模型怎么选。编码器/解码器各一层 128 宽的隐藏层,Adam,250 epoch,逐维 KL 大于 0.05 nat 才算「存活」。 跑两轮:一组不加重衰减,一组给权重加上 $10^{-3}$ 的衰减。结果: weight_decay beta 重建MSE 总KL std(z) 1/std 存活维度 -------------------------------------------------------------------------------- 0 1e-06 0.0030 76.958 2.0596 0.4855 8/8 0 1e-05 0.0031 39.861 1.4107 0.7088 8/8 0 0.0001 0.0038 23.512 1.0651 0.9389 7/8 0 0.001 0.0062 12.454 1.0067 0.9934 4/8 0 0.01 0.0247 7.843 1.0055 0.9946 4/8 0 0.1 0.2119 3.092 0.9952 1.0048 4/8 0 0.3 0.6197 0.887 1.0000 1.0000 3/8 0 1 1.0000 0.000 1.0001 0.9999 0/8 0 3 1.0001 0.000 1.0000 1.0000 0/8 0.001 1e-06 0.0046 55.983 2.7174 0.3680 8/8 0.001 1e-05 0.0046 49.327 2.5978 0.3849 6/8 0.001 0.0001 0.0049 28.672 2.0580 0.4859 5/8 0.001 0.001 0.0075 13.409 1.3082 0.7644 4/8 0.001 0.01 0.0259 7.917 1.0689 0.9355 4/8 0.001 0.1 0.2148 3.107 1.0051 0.9949 4/8 0.001 0.3 0.6301 0.858 1.0018 0.9982 3/8 0.001 1 1.0000 0.000 1.0000 1.0000 0/8 0.001 3 1.0000 0.000 1.0000 1.0000 0/8 三件事一眼可见: 第一,坍缩是一条断崖,不是一个缓坡。 $\beta$ 从 $10^{-2}$ 到 $0.1$ 到 $0.3$ 到 $1$,重建 MSE 走 $0.0247 \to 0.2119 \to 0.6197 \to 1.0000$。活着的维度 $4 \to 4 \to 3 \to 0$。$\beta = 1$ 时总 KL 精确变成 0.000,说明近似后验和先验完全重合——编码器变成了一个只会输出 $\mathcal{N}(0,I)$ 的函数。 第二,维度是一个个死的。 $\beta = 10^{-3}$ 时逐维 KL 是 [2.96, 0.002, 3.19, 0.004, 3.12, 3.18, 0.001, 0.001]——恰好 4 个在 3 nats 附近,另外 4 个趴在 0.002 上。活下来的正好是 4 个,和合成数据的真实因子数相等。 模型自己算出「值得买 4 个维度」,这不是我告诉它的。 第三,潜空间尺度确实跟着 $\beta$ 漂。 无权重衰减时从 1.0067($\beta=10^{-3}$)漂到 2.0596($\beta = 10^{-6}$);加上 $10^{-3}$ 的权重衰减后在同样区间漂到 2.7174。对应的 scaling_factor 从 0.9934 掉到 0.4855 / 0.3680。这就是「0.18215 从哪来」的机制:$\beta$ 小到 KL 项失去约束力时,潜空间的尺度不再由任何东西钉住。 需要说清楚的是,让尺度长大的那股力在我的实验里是权重衰减(潜尺度越大,解码器权重就可以越小),而 Adam 自身对参数尺度不敏感/敏感的部分也在推它(不加重衰减时也会从 1.0067 漂到 2.0596)。真实 VAE 里起同样作用的还有编码器末端的归一化层、初始化尺度、以及训练超参。结论只需要一条:$\beta$ 决定的是「KL 有没有能力把尺度钉在 1」,钉不住之后具体漂到几,是别的因素决定的。 Stable Diffusion 漂到了 5.49,我的玩具漂到了 2.7,同一个机制。 4.2 换条路验一遍:一维闭式解 3.3 节那组闭式解值得单独验,因为「后验坍缩阈值 $\lambda = 1$」这个结论如果错了,整篇文章的框架就错了。做法是直接对 $(w, m, \log s)$ 做梯度下降。 def closed_form(beta, v=V, sigma_x=SIGMA_X): lam = beta * sigma_x ** 2 / v if lam >= 1.0: # lambda >= 1:坍缩 return dict(lam=lam, u=0.0, s2=1.0, w2=0.0, collapsed=True) u = 1.0 - lam return dict(lam=lam, u=u, s2=lam, w2=v * u, collapsed=False) ($v = 1$、$\sigma_{\text{dec}} = 1$,所以 $\lambda = \beta$、坍缩阈值 $\beta^{*} = 1$。)对照结果: beta lambda | u=wm 预测 拟合 | s^2 预测 拟合 | w^2 预测 拟合 | 坍缩 0.05 0.050 | 0.9500 0.9500 | 0.0500 0.0500 | 0.9500 0.9500 | 否 0.1 0.100 | 0.9000 0.9000 | 0.1000 0.1000 | 0.9000 0.9000 | 否 0.2 0.200 | 0.8000 0.8000 | 0.2000 0.2000 | 0.8000 0.8000 | 否 0.4 0.400 | 0.6000 0.6000 | 0.4000 0.4000 | 0.6000 0.6000 | 否 0.6 0.600 | 0.4000 0.4000 | 0.6000 0.6000 | 0.4000 0.4000 | 否 0.8 0.800 | 0.2000 0.2000 | 0.8000 0.8000 | 0.2000 0.2000 | 否 0.95 0.950 | 0.0500 0.0500 | 0.9500 0.9500 | 0.0500 0.0500 | 否 1 1.000 | 0.0000 0.0000 | 1.0000 1.0000 | 0.0000 0.0000 | 是 1.2 1.200 | 0.0000 0.0000 | 1.0000 1.0000 | 0.0000 0.0000 | 是 2 2.000 | 0.0000 0.0000 | 1.0000 1.0000 | 0.0000 0.0000 | 是 5 5.000 | 0.0000 0.0000 | 1.0000 1.0000 | 0.0000 0.0000 | 是 未坍缩区间内,闭式解与数值拟合的最大偏差:5.91e-06 坍缩点 (w, m, s) = (0, 0, 1) 处的梯度(应当恒为 0): beta=0.05 dJ/dw=+0.00e+00 dJ/dm=+0.00e+00 dJ/ds=+0.00e+00 beta=0.5 dJ/dw=+0.00e+00 dJ/dm=+0.00e+00 dJ/ds=+0.00e+00 beta=1 dJ/dw=+0.00e+00 dJ/dm=+0.00e+00 dJ/ds=+0.00e+00 beta=5 dJ/dw=+0.00e+00 dJ/dm=+0.00e+00 dJ/ds=+0.00e+00 坍缩点 vs 解析解的目标函数值(谁小谁是最优): beta=0.05 J(解析解)= 0.09989 J(坍缩点)= 0.50000 beta=0.5 J(解析解)= 0.42329 J(坍缩点)= 0.50000 beta=0.9 J(解析解)= 0.49741 J(坍缩点)= 0.50000 beta=1 J(解析解)= 0.50000 J(坍缩点)= — beta=2 J(解析解)= 0.50000 J(坍缩点)= — 三处细节值得留意。偏差 5.91e-06 说明闭式解是对的。梯度恒等于 0 说明坍缩点在任意 $\beta$ 下都是驻点——它一直「在那儿」,$\lambda \ge 1$ 时它只是终于变成了最小值。$\beta = 0.9$ 时两者的目标值只差 0.0026,说明接近阈值时塌向坍缩点的阻力非常小,这解释了为什么真实训练里坍缩一旦开始就很快。 4.3 后验坍缩长什么样 这张图要看什么:三张子图是三档 $\beta$ 下逐维 KL 的柱状图,蓝色是存活维度(KL > 0.05 nat),灰色是死掉的。注意三张图的纵轴量级完全不同(3.19 / 0.39 / 0.00001 nats,差五个数量级),所以我把每张图各自缩放并标了纵轴最大值——如果共享纵轴,后两张会被压成一条线,看不出结构。$\beta = 10^{-3}$ 时是 4 根高柱加 4 根贴地;$\beta = 0.3$ 时只剩 3 根矮柱,重建 MSE 已经涨到 0.6197;$\beta = 1$ 时一根都没有。 4.4 潜尺度漂移:从 1.0067 到 2.7174 这张图左侧按全方差公式堆叠两项方差:跨样本后验均值方差、平均后验方差;右侧显示第一项占总方差的比例。两项方差相加后开方才得到边际标准差,标准差本身不能直接堆叠。这个分解描述二阶统计,不能等同于信息与噪声的互信息分解。 这张图帮助理解尺度与信息的区别:相同的边际方差可以来自不同的均值/条件方差组合。缩放只调整统计尺度,不保证 latent 的语义分布匹配,也不能单凭 std 判断是否坍缩。 4.5 忘掉 scaling_factor 的代价,量化 这张图要看什么:左图的纵轴是对数的,三条线分别是三种潜空间尺度下的去噪幅度比 $r$。绿线 $\sigma = 1.0$ 平在 $r = 1$ 上(正确);蓝线是真实情形 $\sigma = 5.49$,$\bar{\alpha}$ 越小掉得越狠,最左端贴在 $1/\sigma^{2} = 0.033$,也就是幅度只剩 3.3%(后验均值幅度偏小);红线是反向错误——潜空间其实是单位方差却照抄了 0.18215,$\sigma$ 变成 0.182,$r$ 冲到 30 倍(后验均值幅度偏大)。右图是同一个前向过程下两种去噪器的均方误差,中段拉开 8 倍,两端收敛($\bar{\alpha} \to 1$ 时都没噪声要除,$\bar{\alpha} \to 0$ 时都没信息可用——误差差距最大的地方在中间,这跟直觉不太一样)。 alpha_bar | 错假设 MSE 正确 MSE 理论后验方差 恶化倍数 0.999 | 0.0010 0.0010 0.0010 1.03x 0.99 | 0.0130 0.0101 0.0101 1.28x 0.9 | 0.3933 0.1106 0.1107 3.55x 0.5 | 7.8154 0.9682 0.9679 8.07x 0.1 | 24.5816 6.9502 6.9305 3.54x 0.01 | 29.6359 23.2187 23.1056 1.28x 「理论后验方差」一列是和蒙特卡洛结果并排校核的:$\bar{\alpha}=0.5$ 处 0.9679 对 0.9682,这一格的差值约 $3\times10^{-4}$;其他格的采样波动会更大,不能把单个差值当作整张表的误差保证。 05. 工业级实现对照 真实框架长什么样,看 diffusers 的 AutoencoderKL 和 DiagonalGaussianDistribution。以下以 2026-09 时的实现为准,上游会重构。 编码链路(AutoencoderKL._encode / encode):encoder(x) → quant_conv → 把结果塞进 DiagonalGaussianDistribution。和最小实现的差异有四处: 用卷积而不是全连接。最小实现里 x 是一个 64 维向量,一层矩阵乘就够;图像要保留空间结构,所以编码器是卷积堆栈,输出 [B, 2*z_channels, H/8, W/8],逐像素各出一组 $(\mu, \text{logvar})$。KL 也因此对每个元素都算一次。 *quant_conv 是 1×1 卷积,默认从 `2latent_channels映到同样的通道数**。编码器已经输出均值与 logvar 所需的双倍通道;这里做通道混合,随后torch.chunk(parameters, 2, dim=1)` 分成两组,不是在 quant_conv 这一步翻倍。 logvar 被硬截断到 $[-30, 20]$,这一行是 self.logvar = torch.clamp(self.logvar, -30.0, 20.0)。别小看它:$\exp(20) = 4.85\times10^{8}$,$\exp(-30) = 9.36\times10^{-14}$。我的最小实现里也照抄了这个阈值。不加截断,一个异常值就能让 KL 炸到 1e8 或者把梯度打进下溢区。 encode 返回的是后验分布,不乘 scaling factor。这一步由调用方负责——pipeline 里显式写 latents = vae.encode(image).latent_dist.sample() * vae.config.scaling_factor。这个设计是个典型的「容易漏」:它不在 encode 里面,所以复制粘贴半段代码就会丢掉。 推理时用 sample() 还是 mode():两者都合法,必须与具体管线的约定一致。mode() 返回均值,适合需要确定性编码的实验;sample(generator=...) 返回后验样本,固定随机生成器也能复现。diffusers 的 Stable Diffusion img2img 中 retrieve_latents 默认走 sample,所以不能说图生图、修补必须用 mode()。deterministic=True 是把分布对象的方差与标准差置零的特殊模式,不等于重新训练过的普通自编码器。 KL 的实现细节:kl() 里写的是 0.5 * torch.sum(torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar, dim=[1, 2, 3])。跟 3.2 节那个式子逐字对上(var 就是 $\sigma^{2}$,logvar 就是 $\log\sigma^{2}$),只是求和范围从「8 个潜变量维度」变成「$4 \times 64 \times 64 = 16384$ 个元素」,返回的是每个样本一个标量。这一点是全文最容易被忽略的工程事实:KL 是逐元素求和的,所以潜变量个数一变,同一个 $\beta$ 数值的含义完全不同。LDM 用 $10^{-6}$、我的实验用 $10^{-3}\sim1$,这两个数根本不在同一个坐标系里——看别人的 $\beta$ 必须连着看它的潜变量张量形状。 视频 VAE 上这条线怎么延伸:时间维一起下采样,潜变量变成 [B, C, T/4, H/8, W/8],KL 求和范围又大了一个量级。KL 虽逐元素求和,但编码器、解码器会耦合各维,不能保证各维独立开关。元素数增大时还要同时看重建项如何归约,不能仅凭维度数断言同一 β 必然关闭更多维度——这也是为什么视频 VAE 的 loss 配比需要单开一篇(见 视频 VAE 的常见 loss 组合)。 带 shift 的 VAE 要区分编码和解码方向:例如 diffusers 的 SD3 管线 在解码前使用 $z_{\text{raw}}=z_{\text{diffusion}}/\text{scale}+\text{shift}$;对应的正向变换才是 $z_{\text{diffusion}}=(z_{\text{raw}}-\text{shift})\cdot\text{scale}$。读取具体 checkpoint 的配置和管线,不把 SD1.x 的常数套到 SD3 / FLUX。 06. 代价与边界 VAE 是有损压缩,这是第一位的代价。 以 $f=8$、4 通道为例,一张 512×512 的图进来,出去的是 $4 \times 64 \times 64 = 16384$ 个数,压缩比 $3\times512\times512 / 16384 \approx 48$。压缩本身就是有损的,而且丢的是高频——文字、小脸、细纹理这些恰恰是人类最敏感的东西。扩散模型再强也补不回来,因为信息在进入扩散过程之前就已经没了(这也是 SD 生态里独立高分辨率精修模型存在的理由)。 把 $\beta$ 调小,赔的是潜空间的可预测性。 潜空间尺度失去约束之后:换 VAE 必须核对 latent 语义、尺度与扩散模型训练约定(不能只重算一个系数),而且——按 LDM D.1 的观察——即使一致地训练,信噪比全程偏高也会让高分辨率卷积采样出问题。$\beta$ 越小,重建越好,但潜空间越像一个「定制格式」,越难被别的东西复用。 把 $\beta$ 调大,赔的是潜变量的信息容量。 实测 $\beta=1$ 时 8 个维度接近先验、重建 MSE 约为 1。在线性闭式模型里,存活区的 MSE 为 $\beta\sigma_{\text{dec}}^2$,到阈值后连续接到 $v$;逐维 KL 也连续降到 0。采用阈值统计的“存活维度数”会出现台阶,但这不等于重建误差存在不连续跳变。 什么时候不该用潜空间:需要像素级保真的任务(超分、医学影像、文字密集的文档生成)直接上像素空间或多尺度方案,别压 $f=8$;另外如果任务的训练数据量和算力都充足、又不要求高分辨率,2.1 节那套「潜空间省算力」的收益应与压缩误差一起实测。 另一条路线:VQ。 用有限码本替换连续高斯,带来量化误差、码本容量与使用率等另一组权衡。LDM 论文观察到它所用的 VQ latent 方差接近 1;这不是离散化的数学保证,码本仍然可以整体缩放,是否需要归一化要看实际训练分布。 07. 经典论文脉络 Kingma & Welling, 2013, Auto-Encoding Variational Bayes(arXiv:1312.6114) — VAE 本体。两个贡献撑起了后面所有工作:重参数化技巧(把采样挪出计算图)和对角高斯先验下 KL 的闭式解(3.2 节那个式子)。 Higgins 等, ICLR 2017, β-VAE: Learning Basic Visual Concepts with a Constrained Variational Framework — 把 KL 权重从固定的 1 变成可调旋钮,$\beta>1$ 换解耦表征。副产品是把「后验坍缩」推到了台前:$\beta$ 一大,重建就塌。 van den Oord 等, 2017, Neural Discrete Representation Learning(arXiv:1711.00937) — VQ-VAE。用码本查表替换连续高斯采样,潜空间变成离散索引,3.2 节那个 KL 项整个消失了。 Esser 等, CVPR 2021, Taming Transformers for High-Resolution Image Synthesis(arXiv:2012.09841) — VQGAN。给 VQ 自编码器加上感知损失和对抗损失,把重建质量推到可商用,第一次让「学到的潜空间 + 自回归」在 1024² 上真正可用。 Rombach 等, CVPR 2022, High-Resolution Image Synthesis with Latent Diffusion Models(arXiv:2112.10752) — LDM / Stable Diffusion。它把「KL 权重、潜空间方差、信噪比、scaling factor」这四件事的关系写进了 4.3.2 与 D.1 两节,$0.18215$ 从此钉在了所有下游代码里。 顺着这条线看,故事其实是「潜空间的统计性质从哪儿来」在被一步步讲清楚:2013 年给了目标函数,2017 年发现权重会毁掉它,2017—2021 年绕开它(离散化 + 对抗训练),2022 年终于正面处理它的尺度问题。 08. 常见误解 误解一:KL 既然是正则项,越大越好。 KL 确实可以看作正则,也可用信息预算解释,但增强它会牺牲重建。第 09 节的密集扫描显示 MSE 从 $0.0650$ 到 $0.8874$ 逐步上升;台阶出现在人为设阈值的存活维度计数,不能把稀疏扫描误读为“中间没有过渡”。 误解二:任意 KL 权重下都在拟合原模型的真实后验。 当 $\beta=1$ 且变分族足够时,最优 $q$ 可以等于该生成模型的真实后验;本文线性高斯模型就是例子。$\beta\ne1$ 则改变权衡,存活区 $s^2=\beta\sigma_{\text{dec}}^2/v$ 同时依赖权重、观测噪声和数据方差 $v$,不能说与数据无关。 误解三:聚合后验等于先验就意味着坍缩。 逐样本 $q(z\mid x)$ 和聚合后验 $q(z)=\int q(z\mid x)p_{\text{data}}(x)\,dx$ 是两个对象。本文未坍缩线性解满足 $m^2v+s^2=(1-\lambda)+\lambda=1$,因而 $q(z)=\mathcal N(0,1)$,但 $m\ne0$,仍然传递信息。真正的完全坍缩要求几乎所有输入的整个 $q(z\mid x)$ 都等于同一个先验。 误解四:scaling factor 是固定的魔法常数。 LDM 的标定实现使用首个训练 batch 的 latent 展平后的标准差,令缩放因子为其倒数;这不是逐通道独立归一化。使用预训练模型时应遵守 checkpoint 保存的系数,不能随手在一张新图上重估并替换。换 VAE 还可能改变 latent 的语义和通道分布,重算一个标量不保证与原扩散模型兼容。 误解五:后验采样只在训练时用。 官方 img2img 管线默认也会采样;固定 generator 可以控制随机性。mode() 是去掉编码采样噪声的一种选择,必须核对训练与推理的分布约定,不能统一替换所有管线。 误解六:截断范围等于所有精度下的安全范围。 [-30,20] 是实现中的防护范围,但 exp(20) 超过 fp16 最大有限值,实际还要看计算 dtype 和 VAE 是否上转 fp32。实验中的 logvar 约 −9.4 离 −30 很远,不能仅凭这个数断言已接近数值崩溃。 09. 动手验证 三个小实验,都能在两分钟内跑完。 实验一:确认坍缩是「跳」还是「滑」。 一行命令: python vae_minimal.py --betas=0.03,0.05,0.1,0.2,0.5 --wd=0 实测($\text{weight\_decay} = 0$,$250$ epoch,随机种子固定,可复现): $\beta$ $0.03$ $0.05$ $0.1$ $0.2$ $0.5$ 重建 MSE $0.0650$ $0.1057$ $0.2119$ $0.4203$ $0.8874$ 总 KL $5.580$ $4.554$ $3.092$ $1.676$ $0.198$ 存活维度 $4/8$ $4/8$ $4/8$ $4/8$ $1/8$ 结论:重建误差逐渐增加,存活维度的计数呈台阶。 这里把每维 KL 超过 0.05 nat 计为存活,所以计数天然离散。在线性模型中,每维 KL 为 $-\tfrac12\log\lambda$(未坍缩时),连续趋向 0,并不会从 3 直接跳到 0;应同时观察逐维 KL、重建与均值/方差统计。 实验二:确认尺度漂移是权重衰减带来的。 把权重衰减单独开到 $10^{-2}$: python vae_minimal.py --wd=0.01 三档权重衰减下 $\text{std}(z)$ 的对照(同一批 $\beta$,越小越说明潜尺度已经失控): $\beta$ $\text{wd} = 0$ $\text{wd} = 10^{-3}$ $\text{wd} = 10^{-2}$ $10^{-6}$ $2.0596$ $2.7174$ $2.6749$ $10^{-5}$ $1.4107$ $2.5978$ $2.6383$ $10^{-4}$ $1.0651$ $2.0580$ $2.5399$ $10^{-3}$ $1.0067$ $1.3082$ $2.0463$ $10^{-2}$ $1.0055$ $1.0689$ $1.3657$ $10^{-1}$ $0.9952$ $1.0051$ $1.0523$ $3\times10^{-1}$ $1.0000$ $1.0018$ $1.0000$ 看这张表的方式是看「回到 1.00 的那个拐点」在往右挪:$\text{wd} = 0$ 时 $\beta \ge 10^{-3}$ 就已经归位;$\text{wd} = 10^{-3}$ 时要到 $\beta \ge 10^{-2}$;$\text{wd} = 10^{-2}$ 时要一路推到 $\beta \ge 0.3$。权重衰减把「潜尺度失控」的区间往大 $\beta$ 方向整整推了两个数量级。一个诚实的补充:$10^{-6}$ 那一档 $10^{-2}$ 的 $2.6749$ 反而略低于 $10^{-3}$ 的 $2.7174$,说明这个偏移会饱和,不是权重衰减越大越离谱。 实验三:反向验证 scaling factor。 换两个反事实的缩放系数各跑一次(--sf 会把整张表按新系数重算): python latent_scaling.py --sf=0.5 python latent_scaling.py --sf=2.0 实际结果: $\text{scaling factor}$ $\text{std}(z)$ $r$ 在 $\bar{\alpha}=0.999$ $r$ 在 $\bar{\alpha}=0.01$ $\bar{\alpha}=0.5$ 处 MSE 比值 $0.18215$(真实值) $5.4900$ $0.9990$ $0.0428$ $8.07\times$ $0.5$(缩放不足) $2.0000$ $0.9992$ $0.2575$ $1.57\times$ $2.0$(缩放过头) $0.5000$ $1.0030$ $3.9700$ $1.56\times$ $\sigma = 2$ 时 $r$ 全程在 1 以下(最低 $0.2575$,正是 $1/\sigma^{2} = 0.25$ 加上 $\bar{\alpha}$ 那一项),该高斯估计器的幅度偏小;$\sigma = 0.5$ 时 $r$ 全程在 1 以上(最高 $3.97$,趋近 $1/\sigma^{2} = 4$),该高斯估计器的幅度偏大。在这个高斯模型中:$r$ 是大于还是小于 1,只取决于 $\sigma$ 是大于还是小于 1;而且 $\bar{\alpha}\to0$ 的末端,$r$ 就直接收敛到 $1/\sigma^{2}$ —— 噪声越大,缩放错误的代价越彻底。反过来看真实值那一档:$r$ 掉到 $0.0428$,MSE 差 $8.07$ 倍,比两个假想档位严重一个量级,这就是 0.18215 这个数不能省的原因。 10. 延伸阅读 前置:变分下界与重参数化 —— ELBO 的完整推导和重参数化为什么可导,本篇只用结论。 同方向:视频 VAE 的常见 loss 组合 —— 真实训练里 $\beta$ 不是单独出现的,它要和 L1、LPIPS、GAN 四项一起配比。 数值细节:混合精度与数值稳定性 —— logvar 截断那类技巧在 fp16 下会变得更关键。 下一步:潜空间扩散与 Stable Diffusion 架构会讲 $\text{Var}(z)$ 和信噪比这条线在 UNet 采样里怎么具体表现出来;「离散化表征:VQ-VAE 与 VQGAN」讲另一条绕开后验坍缩的路线。 附录:完整代码 09 节用到的脚本全文如下(vae_minimal.py、collapse_closed_form.py、latent_scaling.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 vae_minimal.py """最小 VAE:numpy 手写反向传播,扫 KL 权重 beta,观察潜空间统计量怎么变。 这篇文章要回答的问题是:KL 权重 beta 调大调小,到底改变了什么? 本脚本做两轮对照实验: 条件 A(无权重衰减):潜尺度被 KL 项钉住,beta 从小到大,std(z) 稳在 1 附近。 条件 B(权重衰减 1e-3):解码器偏好「潜尺度大、权重小」的解, 只有 KL 项拦得住它;beta 一小,潜尺度就漂走。 两轮对照说明同一件事:**beta 并不直接决定潜尺度,它决定的是 「KL 项有没有能力把潜尺度钉在 1」**。钉不住时,潜尺度由训练里其他所有力 (权重衰减、归一化、初始化)共同决定,可以漂到 5 倍开外—— 此实验解释弱正则下的一种漂移机制,但不能据此反推 Stable Diffusion 的具体训练轨迹。 运行: /usr/local/bin/python3 vae_minimal.py 依赖: numpy """ import os import json import numpy as np # ── 固定随机性,保证正文里贴的每个数字都能复现 ───────────────────────── DATA_SEED = 0 INIT_SEED = 1 BATCH_SEED = 2 EVAL_SEED = 3 N_TRAIN = 4096 # 训练样本数 D_OBS = 64 # 观测维度(类比一张图的像素数) H_HID = 128 # 编码器/解码器隐藏层 D_LAT = 8 # 潜变量维度 K_FAC = 4 # 合成数据的真实因子数(< D_LAT,故意留出冗余维度) EPOCHS = 250 BATCH = 256 LR = 3e-3 LOGVAR_MIN, LOGVAR_MAX = -30.0, 20.0 # 与 diffusers DiagonalGaussianDistribution 一致 BETAS = [1e-6, 1e-5, 1e-4, 1e-3, 1e-2, 1e-1, 3e-1, 1.0, 3.0] WD_LIST = [0.0, 1e-3] # 两轮对照:无权重衰减 / 有权重衰减 # ─────────────────────────── 合成数据 ─────────────────────────── def make_data(n=N_TRAIN, d=D_OBS, k=K_FAC, seed=DATA_SEED): """x = A u + 观测噪声,u 是 k 维高斯因子。 做法和真实图像一样:观测维度很多(d=64),但真正驱动它的因子只有 k=4 个。 潜变量维度 D_LAT=8 > 4,所以模型必须自己决定「用几个维度」, 这正是后验坍缩能被观察到的前提。 最后把每个观测维度归一化到方差 1,这样「重建 MSE = 1」就等于「什么都没学到」。 """ rng = np.random.default_rng(seed) A = rng.normal(size=(d, k)) / np.sqrt(k) U = rng.normal(size=(n, k)) X = U @ A.T + 0.05 * rng.normal(size=(n, d)) X = X / X.std(axis=0, keepdims=True) # 逐维方差 = 1 return X # ─────────────────────────── 模型 ─────────────────────────── class VAE: """编码器 x -> h -> (mu, logvar),解码器 z -> g -> x_hat。 符号与正文第 03 节一致: h = relu(x W1 + b1) 编码器隐藏层 mu = h Wmu + bmu 后验均值 logvar = h Wlv + blv 后验对数方差(网络出的是 log sigma^2) z = mu + exp(0.5 logvar) * eps 重参数化 x_hat = relu(z V1 + c1) V2 + c2 解码器 """ def __init__(self, d_obs=D_OBS, h=H_HID, d_lat=D_LAT, seed=INIT_SEED): rng = np.random.default_rng(seed) sc = lambda fan_in, fan_out: rng.normal( scale=np.sqrt(2.0 / (fan_in + fan_out)), size=(fan_in, fan_out) ) self.p = {} self.p["W1"] = sc(d_obs, h) self.p["b1"] = np.zeros(h) self.p["Wmu"] = sc(h, d_lat) self.p["bmu"] = np.zeros(d_lat) self.p["Wlv"] = sc(h, d_lat) self.p["blv"] = np.zeros(d_lat) self.p["V1"] = sc(d_lat, h) self.p["c1"] = np.zeros(h) self.p["V2"] = sc(h, d_obs) self.p["c2"] = np.zeros(d_obs) self.d_obs, self.h, self.d_lat = d_obs, h, d_lat def forward(self, x, eps): p = self.p h1 = np.maximum(x @ p["W1"] + p["b1"], 0.0) mu = h1 @ p["Wmu"] + p["bmu"] lv_raw = h1 @ p["Wlv"] + p["blv"] lv = np.clip(lv_raw, LOGVAR_MIN, LOGVAR_MAX) # 防 exp 溢出,见第 05 节 sig = np.exp(0.5 * lv) z = mu + sig * eps g1 = np.maximum(z @ p["V1"] + p["c1"], 0.0) xhat = g1 @ p["V2"] + p["c2"] cache = dict(x=x, h1=h1, mu=mu, lv=lv, lv_raw=lv_raw, sig=sig, z=z, g1=g1, xhat=xhat, eps=eps) return xhat, cache @staticmethod def losses(x, xhat, mu, lv, sig): """recon = 逐元素均方误差;kl = 0.5 * sum_d(mu^2 + var - 1 - logvar)。 注意 kl 是对潜变量维度「求和」而不是求平均——这是标准写法, 也意味着 beta 的等效大小会随潜变量个数一起变。 """ recon = float(np.mean((xhat - x) ** 2)) kl_dim = 0.5 * np.mean(mu ** 2 + sig ** 2 - lv - 1.0, axis=0) # [d] return recon, kl_dim, float(np.sum(kl_dim)) def backward(self, cache, beta): """全部手推。grad 的形状与 self.p 一一对应。""" p = self.p x, h1, mu, lv, lv_raw, sig, z, g1, xhat, eps = ( cache[k] for k in ["x", "h1", "mu", "lv", "lv_raw", "sig", "z", "g1", "xhat", "eps"] ) B, D = x.shape[0], self.d_obs g = {k: np.zeros_like(v) for k, v in p.items()} # 重建项:recon = mean((xhat - x)^2),d/d(xhat) = 2 (xhat - x) / (B*D) dxhat = 2.0 * (xhat - x) / (B * D) g["V2"] += g1.T @ dxhat g["c2"] += dxhat.sum(axis=0) dg1 = (dxhat @ p["V2"].T) * (g1 > 0) # relu g["V1"] += z.T @ dg1 g["c1"] += dg1.sum(axis=0) dz = dg1 @ p["V1"].T # [B, d] # 重参数化:dz 同时流回 mu 与 logvar 两条支路 dmu = dz + beta * mu / B # KL 对 mu 的导数 = beta * mu dlv = dz * (0.5 * sig * eps) + beta * 0.5 * (sig ** 2 - 1.0) / B # 被 clamp 的位置梯度必须截断,否则会用未截断的 logvar 继续更新 dlv = dlv * ((lv_raw > LOGVAR_MIN) & (lv_raw < LOGVAR_MAX)) g["Wmu"] += h1.T @ dmu g["bmu"] += dmu.sum(axis=0) g["Wlv"] += h1.T @ dlv g["blv"] += dlv.sum(axis=0) dh1 = (dmu @ p["Wmu"].T + dlv @ p["Wlv"].T) * (h1 > 0) g["W1"] += x.T @ dh1 g["b1"] += dh1.sum(axis=0) return g # ─────────────────────────── Adam ─────────────────────────── class Adam: def __init__(self, params, lr=LR, b1=0.9, b2=0.999, eps=1e-8): 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, params, grads): self.t += 1 for k in params: self.m[k] = self.b1 * self.m[k] + (1 - self.b1) * grads[k] self.v[k] = self.b2 * self.v[k] + (1 - self.b2) * grads[k] ** 2 mhat = self.m[k] / (1 - self.b1 ** self.t) vhat = self.v[k] / (1 - self.b2 ** self.t) params[k] -= self.lr * mhat / (np.sqrt(vhat) + self.eps) def train(beta, X, weight_decay=0.0, epochs=EPOCHS, batch=BATCH, seed=BATCH_SEED): """训练一个 VAE。weight_decay 是本次实验的关键对照变量。""" rng = np.random.default_rng(seed) model = VAE() opt = Adam(model.p) n = X.shape[0] for _ in range(epochs): perm = rng.permutation(n) for s in range(0, n, batch): xb = X[perm[s:s + batch]] eps = rng.normal(size=(xb.shape[0], D_LAT)) xhat, cache = model.forward(xb, eps) model.losses(xb, xhat, cache["mu"], cache["lv"], cache["sig"]) grads = model.backward(cache, beta) if weight_decay > 0: # 只对权重做衰减,偏置不管 for k in ["W1", "Wmu", "Wlv", "V1", "V2"]: grads[k] += weight_decay * model.p[k] opt.step(model.p, grads) return model def evaluate(model, X, beta, weight_decay=0.0, n_repeat=8): """统计潜空间性质,并把边际标准差拆成两半: Var(z_j) = Var_x(mu_j(x)) + E_x[sigma_j(x)^2] 前一半是「不同样本被编码到不同位置」,后一半是「每个样本自己抖多少」。 scaling_factor 要归一化的是两者之和的开方。 """ rng = np.random.default_rng(EVAL_SEED) n, d = X.shape[0], D_LAT mu_all = np.empty((n, d)) var_all = np.empty((n, d)) recon_all = 0.0 kl_dim_all = np.zeros(d) for s in range(0, n, 512): xb = X[s:s + 512] B = xb.shape[0] h1 = np.maximum(xb @ model.p["W1"] + model.p["b1"], 0.0) mu = h1 @ model.p["Wmu"] + model.p["bmu"] lv = np.clip(h1 @ model.p["Wlv"] + model.p["blv"], LOGVAR_MIN, LOGVAR_MAX) var = np.exp(lv) sig = np.sqrt(var) mu_all[s:s + B] = mu var_all[s:s + B] = var # 逐维 KL:先对 batch 求和,最后统一除以总样本数 kl_dim_all += 0.5 * np.sum(mu ** 2 + var - lv - 1.0, axis=0) / n # 用同一个 mu 重复采样 n_repeat 次,估计解码器实际看到的抖动有多大 eps = rng.normal(size=(B, n_repeat, d)) z = mu[:, None, :] + sig[:, None, :] * eps z_flat = z.reshape(-1, d) h1d = np.maximum(z_flat @ model.p["V1"] + model.p["c1"], 0.0) xhat = (h1d @ model.p["V2"] + model.p["c2"]).reshape(B, n_repeat, D_OBS) recon_all += np.sum((xhat - xb[:, None, :]) ** 2) / (n * n_repeat * D_OBS) std_mu = mu_all.std(axis=0) # [d] mean_sig = np.sqrt(var_all.mean(axis=0)) # [d] marginal_std = np.sqrt(std_mu ** 2 + var_all.mean(axis=0)) # [d] overall_std = float(np.sqrt(np.mean(marginal_std ** 2))) return dict( beta=beta, weight_decay=weight_decay, recon=float(recon_all), kl=float(np.sum(kl_dim_all)), kl_dim=kl_dim_all.tolist(), std_mu=std_mu.tolist(), mean_sig=mean_sig.tolist(), marginal_std=marginal_std.tolist(), overall_std=overall_std, scaling_factor=1.0 / overall_std, active_dims=int(np.sum(kl_dim_all > 0.05)), # 逐维 KL > 0.05 nat 视为「在用」 ) def sweep(wd_list=None, betas=None, verbose=True): wd_list = WD_LIST if wd_list is None else wd_list betas = BETAS if betas is None else betas X = make_data() out = {} for wd in wd_list: rows = [] for beta in betas: model = train(beta, X, weight_decay=wd) rows.append(evaluate(model, X, beta, weight_decay=wd)) out[f"wd={wd:g}"] = rows if verbose: print(f"--- weight_decay = {wd:g} ---") for m in rows: print(f" beta={m['beta']:<8g} recon={m['recon']:.4f} " f"kl={m['kl']:9.3f} std(z)={m['overall_std']:7.4f} " f"1/std={m['scaling_factor']:7.4f} active={m['active_dims']}/{D_LAT}") return out CACHE = os.path.join(os.path.dirname(os.path.abspath(__file__)), "beta_sweep.json") def load_or_sweep(path=CACHE): """图与正文共用同一份结果,避免图文数字漂移。""" if os.path.exists(path): with open(path) as f: return json.load(f) res = sweep() with open(path, "w") as f: json.dump(res, f, indent=2) return res def _parse_cli(argv): """正文 09 节的两个实验就是靠这两个参数复现的: python vae_minimal.py --betas 0.03,0.05,0.1,0.2,0.5 --wd 0 python vae_minimal.py --wd 0.01 """ wd_list, betas, full = None, None, False for a in argv: if a.startswith("--betas="): betas = [float(t) for t in a.split("=", 1)[1].split(",")] elif a.startswith("--wd="): wd_list = [float(t) for t in a.split("=", 1)[1].split(",")] elif a == "--full": full = True return wd_list, betas, full if __name__ == "__main__": import sys wd_list, betas, force = _parse_cli(sys.argv[1:]) custom = wd_list is not None or betas is not None if os.path.exists(CACHE) and not force and not custom: # 扫描要跑 9 x 2 x 250 = 4500 个 epoch,读缓存是为了让「跑一遍就有输出」 # 这句话成立;想从头算就加 --full。指定了自定义参数就不走缓存。 with open(CACHE) as f: res = json.load(f) print(f"[缓存] 读 {os.path.basename(CACHE)};加 --full 可重跑约 5 分钟的完整扫描") else: if custom: print(f"自定义扫描:wd={wd_list or WD_LIST} beta={betas or BETAS}") else: print("开始完整扫描(9 个 beta x 2 组权重衰减 x 250 epoch,约 5 分钟)...") res = sweep(wd_list, betas) if not custom: with open(CACHE, "w") as f: json.dump(res, f, indent=2) print(f"\n{'='*84}") print(f"{'weight_decay':>14} {'beta':>9} {'重建MSE':>10} {'总KL':>10} " f"{'std(z)':>9} {'1/std':>8} {'存活维度':>8}") print(f"{'-'*84}") for key, rows in res.items(): for m in rows: print(f"{key:>14} {m['beta']:>9g} {m['recon']:>10.4f} {m['kl']:>10.3f} " f"{m['overall_std']:>9.4f} {m['scaling_factor']:>8.4f} " f"{m['active_dims']:>5}/{D_LAT}") print("\n潜尺度拆账(std_mu = 跨维均值变化标准差 RMS;mean_sig = 跨维采样噪声标准差 RMS):") for key, rows in res.items(): print(f" {key}") for m in rows: frac = np.mean(np.array(m["std_mu"]) ** 2) / ( np.mean(np.array(m["std_mu"]) ** 2) + np.mean(np.array(m["mean_sig"]) ** 2)) print(f" beta={m['beta']:<8g} std_mu={np.sqrt(np.mean(np.array(m['std_mu']) ** 2)):6.3f} " f"mean_sig={np.sqrt(np.mean(np.array(m['mean_sig']) ** 2)):6.3f} std(z)={m['overall_std']:6.3f} " f"信息占比={frac:5.1%}") collapse_closed_form.py """一维线性 VAE 的闭式解:后验坍缩的阈值到底在哪。 第 03 节推了这么一个结论: 设 数据 x ~ N(0, v),编码器 q(z|x) = N(m x, s^2), 解码器 p(x|z) = N(w z, sigma_x^2)(sigma_x 固定,等价于重建项的权重), 目标 J = [v(1-wm)^2 + w^2 s^2] / (2 sigma_x^2) + beta * 0.5 * (m^2 v + s^2 - log s^2 - 1) 令 lambda = beta * sigma_x^2 / v,则驻点为 u := w m = 1 - lambda (潜变量被真正使用的程度) s^2 = lambda (后验方差) w^2 = v (1 - lambda) (解码器权重) 当 lambda > 1 时此分支要求 w^2 < 0,不存在实数解;lambda = 1 接到坍缩点,最优解退化为坍缩点 (w, m, s) = (0, 0, 1)。 这个脚本用数值优化去拟合 (w, m, s),逐项对照闭式解, 顺便验证「坍缩点永远是一个驻点」这件事。 运行: /usr/local/bin/python3 collapse_closed_form.py 依赖: numpy """ import numpy as np V = 1.0 # 数据方差 SIGMA_X = 1.0 # 解码器的观测噪声标准差,固定(它就是重建项的隐式权重) # ─────────────────────────── 目标函数与梯度 ─────────────────────────── def objective(w, m, s, beta, v=V, sigma_x=SIGMA_X): recon = (v * (1.0 - w * m) ** 2 + w ** 2 * s ** 2) / (2.0 * sigma_x ** 2) kl = 0.5 * beta * (m ** 2 * v + s ** 2 - np.log(s ** 2) - 1.0) return recon + kl def grad(w, m, s, beta, v=V, sigma_x=SIGMA_X): dJ_dw = (-v * m * (1.0 - w * m) + w * s ** 2) / sigma_x ** 2 dJ_dm = -v * w * (1.0 - w * m) / sigma_x ** 2 + beta * m * v dJ_ds = w ** 2 * s / sigma_x ** 2 + beta * (s - 1.0 / s) return dJ_dw, dJ_dm, dJ_ds def fit(beta, v=V, sigma_x=SIGMA_X, steps=20000, lr=0.02, seed=0): """对 (w, m, log s) 做梯度下降。用 log s 参数化保证 s > 0。""" rng = np.random.default_rng(seed) w = rng.normal() * 0.5 m = rng.normal() * 0.5 t = rng.normal() * 0.5 # s = exp(t) for i in range(steps): s = np.exp(t) gw, gm, gs = grad(w, m, s, beta, v, sigma_x) # 对 t 的梯度要过一次链式法则 w -= lr * gw m -= lr * gm t -= lr * gs * s return w, m, np.exp(t) def closed_form(beta, v=V, sigma_x=SIGMA_X): if beta <= 0 or v <= 0 or sigma_x <= 0: raise ValueError("beta, v and sigma_x must be positive") lam = beta * sigma_x ** 2 / v if lam >= 1.0: # 坍缩 return dict(lam=lam, u=0.0, s2=1.0, w2=0.0, collapsed=True) u = 1.0 - lam return dict(lam=lam, u=u, s2=lam, w2=v * u, collapsed=False) # ─────────────────────────── 主流程 ─────────────────────────── if __name__ == "__main__": betas = [0.05, 0.1, 0.2, 0.4, 0.6, 0.8, 0.95, 1.0, 1.2, 2.0, 5.0] print(f"v = {V}, sigma_x = {SIGMA_X}, lambda = beta * sigma_x^2 / v = {SIGMA_X**2/V} * beta") print(f"闭式预测:坍缩阈值 beta* = v / sigma_x^2 = {V / SIGMA_X**2:g}\n") print(f"{'beta':>6} {'lambda':>8} | {'u=wm 预测':>10} {'拟合':>9} " f"| {'s^2 预测':>9} {'拟合':>9} | {'w^2 预测':>9} {'拟合':>9} | 坍缩") max_err = 0.0 for beta in betas: cf = closed_form(beta) w, m, s = fit(beta) u_fit = w * m err = max(abs(u_fit - cf["u"]), abs(s ** 2 - cf["s2"]), abs(w ** 2 - cf["w2"])) max_err = max(max_err, err) if not cf["collapsed"] else max_err print(f"{beta:>6g} {cf['lam']:>8.3f} | {cf['u']:>10.4f} {u_fit:>9.4f} " f"| {cf['s2']:>9.4f} {s**2:>9.4f} | {cf['w2']:>9.4f} {w**2:>9.4f} " f"| {'是' if cf['collapsed'] else '否'}") print(f"\n未坍缩区间内,闭式解与数值拟合的最大偏差:{max_err:.2e}") # 坍缩点是不是永远的驻点? print("\n坍缩点 (w, m, s) = (0, 0, 1) 处的梯度(应当恒为 0):") for beta in [0.05, 0.5, 1.0, 5.0]: gw, gm, gs = grad(0.0, 0.0, 1.0, beta) print(f" beta={beta:<6g} dJ/dw={gw:+.2e} dJ/dm={gm:+.2e} dJ/ds={gs:+.2e}") # 坍缩点在不同 beta 下到底是不是最优?比一下目标函数值 print("\n坍缩点 vs 解析解的目标函数值(谁小谁是最优):") for beta in [0.05, 0.5, 0.9, 1.0, 2.0]: cf = closed_form(beta) if cf["collapsed"]: j_star = objective(0.0, 0.0, 1.0, beta) j_alt = None else: w2 = cf["w2"] w = np.sqrt(w2) m = cf["u"] / w j_star = objective(w, m, np.sqrt(cf["s2"]), beta) j_alt = objective(0.0, 0.0, 1.0, beta) alt = " —" if j_alt is None else f"{j_alt:9.5f}" print(f" beta={beta:<6g} J(解析解)={j_star:9.5f} J(坍缩点)={alt}") latent_scaling.py """scaling_factor 到底在补什么:把「潜空间标准差」翻译成「信噪比」。 Stable Diffusion 的 VAE 有个著名常数 0.18215。它是这么用的: 编码: z_scaled = z * 0.18215 (交给扩散模型之前) 解码: z = z_scaled / 0.18215 (交给解码器之前) diffusers 的文档字符串(AutoencoderKL)写得很明白:这个数是「在训练集第一 个 batch 上算出来的潜空间逐通道标准差」,用它把潜空间缩到单位方差,出处是 LDM 论文(arXiv:2112.10752)的 4.3.2 与 D.1 节。LDM D.1 的原话是: "the signal-to-noise ratio induced by the variance of the latent space (i.e. Var(z) / sigma_t^2) significantly affects the results ... when training a LDM directly in the latent space of a KL-regularized model, this ratio is very high ... Note that the VQ-regularized space has a variance close to 1, such that it does not have to be rescaled." 本脚本把这段话变成可算的数。核心结论是两条: 1. 若扩散模型是在单位方差潜空间上训练的,它对 z_0 的后验均值估计就是 z_hat_0 = sqrt(alpha_bar) * z_t (高斯先验 N(0, I) 下的 MMSE 估计)。而真实潜空间标准差是 sigma 时, 正确的 MMSE 估计是 z_hat_0 = sigma^2 sqrt(alpha_bar) / (sigma^2 alpha_bar + 1 - alpha_bar) * z_t 两者之比 r = alpha_bar + (1 - alpha_bar) / sigma^2。 r < 1 就是「重建出来的东西被整体缩小」——图发灰。 2. 每一层的真实信噪比是 sigma^2 * alpha_bar / (1 - alpha_bar), 模型以为的是 alpha_bar / (1 - alpha_bar),整整差 sigma^2 倍 (20*log10(sigma) 分贝)。 运行: /usr/local/bin/python3 latent_scaling.py 依赖: numpy """ import numpy as np SD_SCALING_FACTOR = 0.18215 # diffusers AutoencoderKL 的默认值 ABARS = [0.999, 0.99, 0.9, 0.5, 0.1, 0.01] def latent_std_from_scaling(s): """scaling_factor 是潜空间边际标准差的倒数。""" return 1.0 / s def mmse_model(z_t, abar): """在「潜空间方差 = 1」的假设下训练出来的高斯 MMSE 去噪器。""" return np.sqrt(abar) * z_t def mmse_true(z_t, abar, sigma): """潜空间真实标准差为 sigma 时的高斯 MMSE 去噪器。""" return sigma ** 2 * np.sqrt(abar) / (sigma ** 2 * abar + 1.0 - abar) * z_t def amplitude_ratio(abar, sigma): """模型输出 / 正确输出 = alpha_bar + (1 - alpha_bar) / sigma^2。""" return abar + (1.0 - abar) / sigma ** 2 def snr(abar, sigma): """真实信噪比 Var(z_0 的信号成分) / Var(噪声成分) = sigma^2 * abar / (1 - abar)。""" return sigma ** 2 * abar / (1.0 - abar) def mmse_mse(abar, sigma): """高斯 MMSE 估计的理论误差,就是后验方差。""" return sigma ** 2 * (1.0 - abar) / (sigma ** 2 * abar + 1.0 - abar) def demo_step(abar, sigma, n=400000, seed=0): """蒙特卡洛:固定同一个前向过程,比较两个去噪器。 est_model:按「潜空间方差 = 1」训练出来的理想去噪器,用在方差为 sigma^2 的 潜空间上 —— 这就是拿了别人训好的权重却没乘 scaling_factor 的情形。 est_true :就该 sigma 训练出来的理想去噪器 —— 也就是正确缩放后应有的表现。 """ rng = np.random.default_rng(seed) z0 = sigma * rng.normal(size=n) eps = rng.normal(size=n) z_t = np.sqrt(abar) * z0 + np.sqrt(1.0 - abar) * eps est_model = mmse_model(z_t, abar) est_true = mmse_true(z_t, abar, sigma) return (float(np.mean((est_model - z0) ** 2)), float(np.mean((est_true - z0) ** 2))) if __name__ == "__main__": import sys # 正文 09 节实验三:换一个反事实的 scaling_factor 重算整张表 # python latent_scaling.py --sf=0.5 # python latent_scaling.py --sf=2.0 for a in sys.argv[1:]: if a.startswith("--sf="): SD_SCALING_FACTOR = float(a.split("=", 1)[1]) sigma_sd = latent_std_from_scaling(SD_SCALING_FACTOR) print("=" * 72) print("一、0.18215 意味着什么") print("=" * 72) print(f" scaling_factor = {SD_SCALING_FACTOR}") print(f" 潜空间边际标准差 sigma = 1 / {SD_SCALING_FACTOR} = {sigma_sd:.4f}") print(f" 潜空间边际方差 Var(z) = {sigma_sd**2:.4f}") print(f" 信噪比偏移 = {20*np.log10(sigma_sd):.2f} dB") print() print(" 对照:若某个 VAE 的潜空间方差真的接近 1(LDM 说 VQ 版就是如此),") print(" scaling_factor 就应当接近 1,完全不需要 rescale。") print() print("=" * 72) print("二、忘掉 scaling factor,去噪幅度错多少") print("=" * 72) print(" r = 模型输出 / 正确输出;r=1 才是正确的") print(f" {'alpha_bar':>10} | " + " | ".join(f"sigma={s:<7.4g}" for s in [sigma_sd, 1.0, SD_SCALING_FACTOR])) for abar in ABARS: rs = [amplitude_ratio(abar, s) for s in [sigma_sd, 1.0, SD_SCALING_FACTOR]] print(f" {abar:>10g} | " + " | ".join(f"{r:>13.4f}" for r in rs)) print() print("=" * 72) print("三、蒙特卡洛实测:单步去噪的均方误差(n=400000)") print("=" * 72) print(f" {'alpha_bar':>10} | {'错假设 MSE':>12} {'正确 MSE':>12} {'理论后验方差':>14} {'恶化倍数':>10}") for abar in ABARS: m_bad, m_good = demo_step(abar, sigma_sd) print(f" {abar:>10g} | {m_bad:>12.4f} {m_good:>12.4f} {mmse_mse(abar, sigma_sd):>14.4f} " f"{m_bad/m_good:>10.2f}x") print() print(" 读法:固定同一个前向过程 z_t = sqrt(alpha_bar) z_0 + sqrt(1-alpha_bar) eps,") print(" 「错假设」是拿了按单位方差潜空间训好的去噪器却喂未缩放的 latent,") print(" 「正确」是就该 sigma 训练的去噪器(也就是乖乖乘上 0.18215 的效果)。") print(f" 「理论后验方差」= sigma^2 (1-alpha_bar) / (sigma^2 alpha_bar + 1 - alpha_bar),") print(f" 与蒙特卡洛的「正确」一列应当吻合(校核 alpha_bar=0.5:" f"{mmse_mse(0.5, sigma_sd):.4f} vs 实测 {demo_step(0.5, sigma_sd)[1]:.4f})") print() print("=" * 72) print("四、反向的错误:换了个方差接近 1 的 VAE,却照抄 0.18215") print("=" * 72) print(f" {'alpha_bar':>10} | {'r(幅度被放大的倍数)':>24}") for abar in ABARS: print(f" {abar:>10g} | {amplitude_ratio(abar, SD_SCALING_FACTOR):>24.4f}") print() print(" r >> 1:重建幅度被整体放大,对应生成图过曝、结构糊成一团。") print() print("=" * 72) print("五、信噪比视角(LDM D.1 的定义 Var(z) / sigma_t^2)") print("=" * 72) print(f" {'alpha_bar':>10} | {'模型以为的 SNR':>16} {'真实 SNR':>16} {'倍数':>10}") for abar in ABARS: a = abar / (1.0 - abar) b = snr(abar, sigma_sd) print(f" {abar:>10g} | {a:>16.4f} {b:>16.4f} {b/a:>10.2f}x") make_figures.py """画本文的四张图。数据源全部来自 vae_minimal.py / latent_scaling.py 的真实输出, 不另造数,避免图文数字漂移。 beta_sweep.png beta 扫描全景:重建 / KL / 潜尺度 / 存活维度 latent_scale_split.png 潜尺度的拆账:编码位置 vs 采样噪声 scaling_snr.png 忘掉 scaling_factor 的后果:幅度比 + 单步去噪 MSE collapse_profile.png 后验坍缩的指纹:逐维 KL 怎么一个个死掉 运行: /usr/local/bin/python3 make_figures.py 依赖: numpy, matplotlib(字体 PingFang SC) """ import os import numpy as np import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import vae_minimal as VM import latent_scaling as LS plt.rcParams["font.sans-serif"] = ["PingFang SC", "Arial Unicode MS", "Heiti TC", "sans-serif"] plt.rcParams["axes.unicode_minus"] = False plt.rcParams["figure.facecolor"] = "white" plt.rcParams["axes.facecolor"] = "white" plt.rcParams["savefig.facecolor"] = "white" plt.rcParams["font.size"] = 11 HERE = os.path.dirname(os.path.abspath(__file__)) FIGDIR = os.path.join(os.path.dirname(HERE), "figures") os.makedirs(FIGDIR, exist_ok=True) # exist_ok 必须带,否则目录已存在会 PermissionError C_MAIN, C_ALT = "#2563eb", "#dc2626" C_GRAY = "#6b7280" def _b(rows): return np.array([r["beta"] for r in rows]) # ─────────────────────────── 图 1:beta 扫描全景 ─────────────────────────── def fig_beta_sweep(res): rows0 = res["wd=0"] rows1 = res["wd=0.001"] fig, axes = plt.subplots(2, 2, figsize=(13.5, 8.6)) ax = axes[0][0] ax.plot(_b(rows0), [r["recon"] for r in rows0], "o-", color=C_MAIN, label="无权重衰减") ax.plot(_b(rows1), [r["recon"] for r in rows1], "s--", color=C_ALT, label="权重衰减 1e-3") ax.set_xscale("log") ax.set_yscale("log") ax.axhline(1.0, color=C_GRAY, lw=1, ls=":") ax.text(1.1e-6, 1.05, "重建 MSE = 1:等于什么都没学到", color=C_GRAY, fontsize=9) ax.set_xlabel(r"KL 权重 $\beta$") ax.set_ylabel("重建 MSE($D$ 维平均)") ax.set_title("(a) 重建质量:β 一大就断崖") ax.legend(fontsize=9) ax.grid(alpha=0.25) ax = axes[0][1] ax.plot(_b(rows0), [r["kl"] for r in rows0], "o-", color=C_MAIN, label="无权重衰减") ax.plot(_b(rows1), [r["kl"] for r in rows1], "s--", color=C_ALT, label="权重衰减 1e-3") ax.set_xscale("log") ax.set_yscale("symlog", linthresh=1e-1) ax.set_xlabel(r"KL 权重 $\beta$") ax.set_ylabel("总 KL(nats,8 维求和)") ax.set_title("(b) KL:β 越大越贴先验,坍缩时归零") ax.legend(fontsize=9) ax.grid(alpha=0.25) ax = axes[1][0] ax.plot(_b(rows0), [r["overall_std"] for r in rows0], "o-", color=C_MAIN, label="无权重衰减") ax.plot(_b(rows1), [r["overall_std"] for r in rows1], "s--", color=C_ALT, label="权重衰减 1e-3") ax.axhline(1.0, color=C_GRAY, lw=1, ls=":") ax.text(1.1e-6, 1.03, "std(z) = 1:被 KL 钉住", color=C_GRAY, fontsize=9) ax.axhline(LS.latent_std_from_scaling(LS.SD_SCALING_FACTOR), color="#059669", lw=1, ls="--") ax.text(1.1e-6, 5.62, "std(z) = 5.49:由 SD1.x 配置系数反推", color="#059669", fontsize=9) ax.set_xscale("log") ax.set_xlabel(r"KL 权重 $\beta$") ax.set_ylabel("潜空间边际标准差 std(z)") ax.set_ylim(0.9, 6.2) ax.set_title("(c) 潜尺度:β 一松手就漂走") ax.legend(fontsize=9, loc="upper right") ax.grid(alpha=0.25) ax2 = ax.twinx() ax2.set_ylim(1 / 6.2, 1 / 0.9) ax2.set_ylabel("对应的 scaling_factor = 1 / std(z)") ax = axes[1][1] ax.plot(_b(rows0), [r["active_dims"] for r in rows0], "o-", color=C_MAIN, label="无权重衰减") ax.plot(_b(rows1), [r["active_dims"] for r in rows1], "s--", color=C_ALT, label="权重衰减 1e-3") ax.axhline(4, color="#059669", lw=1, ls="--") ax.text(1.1e-6, 4.2, "真实因子数 = 4", color="#059669", fontsize=9) ax.set_xscale("log") ax.set_xlabel(r"KL 权重 $\beta$") ax.set_ylabel("还活着的维度数(逐维 KL > 0.05 nat)") ax.set_ylim(-0.3, 8.6) ax.set_title("(d) 维度预算:冗余维度被逐个关掉") ax.legend(fontsize=9, loc="lower left") ax.grid(alpha=0.25) fig.suptitle("KL 权重扫描:重建、KL、潜尺度、存活维度(潜变量 8 维,真实因子 4 个)", fontsize=13, y=0.98) fig.tight_layout(rect=[0, 0, 1, 0.96]) out = os.path.join(FIGDIR, "beta_sweep.png") fig.savefig(out, dpi=150) plt.close(fig) return out # ─────────────────────────── 图 2:潜尺度拆账 ─────────────────────────── def fig_scale_split(res): rows = res["wd=0"] betas = _b(rows) std_mu = np.array([np.mean(np.array(r["std_mu"]) ** 2) for r in rows]) # 编码位置的跨样本波动 mean_sig = np.array([np.mean(np.array(r["mean_sig"]) ** 2) for r in rows]) # 采样噪声 RMS x = np.arange(len(betas)) fig, axes = plt.subplots(1, 2, figsize=(13.5, 4.6)) ax = axes[0] ax.bar(x, std_mu, 0.62, label=r"后验均值的跨样本方差", color=C_MAIN) ax.bar(x, mean_sig, 0.62, bottom=std_mu, label=r"平均后验方差", color="#f59e0b") ax.set_xticks(x) ax.set_xticklabels([f"{b:g}" for b in betas]) ax.set_xlabel(r"KL 权重 $\beta$(对数轴上的等距刻度)") ax.set_ylabel("Var(z) 的构成(各维平均)") ax.set_title("(a) 方差拆账:后验均值变化与后验噪声") ax.legend(fontsize=9) ax.grid(alpha=0.25, axis="y") ax = axes[1] frac = std_mu / (std_mu + mean_sig) ax.plot(x, frac, "o-", color=C_MAIN, lw=2) ax.axhline(0.5, color=C_GRAY, lw=1, ls=":") ax.set_xticks(x) ax.set_xticklabels([f"{b:g}" for b in betas]) ax.set_xlabel(r"KL 权重 $\beta$") ax.set_ylabel("方差里来自编码位置的比例") ax.set_ylim(-0.03, 1.06) ax.set_title("(b) 总方差中均值变化的占比") ax.grid(alpha=0.25) for i, (xi, fi) in enumerate(zip(x, frac)): ax.annotate(f"{fi:.2f}", (xi, fi), textcoords="offset points", xytext=(0, 8), ha="center", fontsize=9, color=C_MAIN) fig.suptitle("全方差公式:两项方差相加;方差占比不等于互信息", fontsize=12.5, y=1.0) fig.tight_layout(rect=[0, 0, 1, 0.95]) out = os.path.join(FIGDIR, "latent_scale_split.png") fig.savefig(out, dpi=150) plt.close(fig) return out # ─────────────────────────── 图 3:scaling factor 的后果 ─────────────────────────── def fig_scaling_snr(): sigma_sd = LS.latent_std_from_scaling(LS.SD_SCALING_FACTOR) abar = np.logspace(-3, np.log10(0.9999), 400) fig, axes = plt.subplots(1, 2, figsize=(13.5, 4.8)) ax = axes[0] for sigma, lab, col in [(sigma_sd, rf"$\sigma$={sigma_sd:.2f}(由 SD1.x 配置反推)", C_MAIN), (1.0, r"$\sigma$=1.0(正确缩放后)", "#059669"), (LS.SD_SCALING_FACTOR, rf"$\sigma$={LS.SD_SCALING_FACTOR:.3f}(照抄 0.18215 用错 VAE)", C_ALT)]: ax.plot(abar, LS.amplitude_ratio(abar, sigma), "-", color=col, lw=2, label=lab) ax.axhline(1.0, color=C_GRAY, lw=1, ls=":") ax.text(2e-3, 1.35, "$r=1$:缩放正确", color=C_GRAY, fontsize=9) ax.text(2e-3, 0.038, r"$r=1/\sigma^2=0.033$:估计幅度约为正确值的 3.3%", color=C_MAIN, fontsize=9) ax.set_xscale("log") ax.set_yscale("log") ax.set_ylim(2e-2, 6e1) ax.set_xlabel(r"$\bar{\alpha}$(1 = 完全干净,0 = 纯噪声)") ax.set_ylabel("去噪幅度比 $r$ = 模型输出 / 正确输出") ax.set_title("(a) 高斯 MMSE 教学模型的后验均值幅度比") ax.legend(fontsize=9, loc="center left") ax.grid(alpha=0.25, which="both") ax = axes[1] abars = np.array(LS.ABARS) bad, good = [], [] for a in abars: b_, g_ = LS.demo_step(a, sigma_sd, n=200000) bad.append(b_) good.append(g_) ax.plot(abars, bad, "o-", color=C_ALT, lw=2, label=r"错假设去噪器(按 $\sigma$=1 训练)") ax.plot(abars, good, "s-", color=C_MAIN, lw=2, label=r"正确去噪器(就该 $\sigma$ 训练)") ax.set_xscale("log") ax.set_yscale("log") ax.set_xlabel(r"$\bar{\alpha}$") ax.set_ylabel("单步去噪均方误差") ax.set_title("(b) 同一个前向过程下的 MSE 差距") ax.legend(fontsize=9) ax.grid(alpha=0.25) i5 = int(np.argmin(np.abs(abars - 0.5))) ax.annotate(f"{bad[i5]/good[i5]:.1f}×", (abars[i5], bad[i5]), textcoords="offset points", xytext=(12, -6), color=C_ALT, fontsize=11) fig.suptitle(r"忘掉 scaling_factor 的代价:信噪比被高估 $\sigma^2$ = 30.1 倍(14.8 dB)", fontsize=12.5, y=1.0) fig.tight_layout(rect=[0, 0, 1, 0.94]) out = os.path.join(FIGDIR, "scaling_snr.png") fig.savefig(out, dpi=150) plt.close(fig) return out # ─────────────────────────── 图 4:坍缩指纹 ─────────────────────────── def fig_collapse_profile(res): rows = res["wd=0"] picks = [3e-3, 3e-1, 1.0] # 三档:够用 / 正在坍缩 / 完全坍缩 # 每张子图各自缩放:三档之间差 20 倍,共享 y 轴会把另两张压成一条线 fig, axes = plt.subplots(1, 3, figsize=(13.5, 4.2)) d = len(rows[0]["kl_dim"]) for ax, beta in zip(axes, picks): row = min(rows, key=lambda r: abs(r["beta"] - beta)) kl = np.array(row["kl_dim"]) colors = [C_MAIN if v > 0.05 else "#d1d5db" for v in kl] ax.bar(np.arange(d), kl, 0.68, color=colors) ax.axhline(0.05, color=C_GRAY, lw=1, ls=":") ax.set_xticks(np.arange(d)) ax.set_xticklabels([f"$z_{j+1}$" for j in range(d)], fontsize=9) ax.set_xlabel("潜变量维度") ax.set_ylim(0, max(0.35, kl.max() * 1.18)) ax.set_title(rf"$\beta$={row['beta']:g} 存活 {row['active_dims']}/{d} " f"重建 MSE={row['recon']:.4f}") ax.grid(alpha=0.25, axis="y") ax.text(0.99, 0.92, f"纵轴最大 {kl.max():.2f} nats", transform=ax.transAxes, ha="right", fontsize=8.5, color=C_GRAY) axes[0].set_ylabel("逐维 KL(nats)") axes[1].set_ylabel("逐维 KL(nats)") axes[2].set_ylabel("逐维 KL(nats)") fig.suptitle("后验坍缩的指纹:KL 预算被削减时,维度是一个个死的,不是一起死(注意三张图的纵轴量级不同)", fontsize=12.5, y=0.99) fig.tight_layout(rect=[0, 0, 1, 0.91]) out = os.path.join(FIGDIR, "collapse_profile.png") fig.savefig(out, dpi=150) plt.close(fig) return out if __name__ == "__main__": res = VM.load_or_sweep() for fn in [fig_beta_sweep, fig_scale_split, fig_collapse_profile]: print("写出", fn(res)) print("写出", fig_scaling_snr())
2026年09月26日
3 阅读
0 评论
0 点赞
2026-09-25
AIGC 基本功|变分下界与重参数化-ELBO
变分下界与重参数化 所属方向:数学基础 | 难度:入门 | 前置知识:无(会求导、会算期望即可) 关键词:变分推断、ELBO、重参数化、KL 散度、后验坍缩、β-VAE、IWAE、得分函数估计量 01. 为什么需要它 先摆三组在同一台机器上跑出来的数字,都出自文末附录,可以自己复现。 第一组:KL 项掉到 0,潜变量整条死掉。 固定解码器噪声 σ²=0.25,让四个独立数据方向分别具有方差 λ=4.0、1.0、0.25、0.0625,只改 KL 项的权重 β,看看最优解长什么样: 数据方向方差 λ β=0.5 β=1.0 β=2.0 β=4.0 死亡阈值 β* 4.0 MSE 0.125 存活 MSE 0.250 存活 MSE 0.500 存活 MSE 1.000 存活 16.00 1.0 MSE 0.125 存活 MSE 0.250 存活 MSE 0.500 存活 MSE 1.000 死亡 4.00 0.25 MSE 0.125 存活 MSE 0.250 死亡 死亡 死亡 1.00 0.0625 死亡 死亡 死亡 死亡 0.25 死亡那一格的具体解是 a=0、s=1、KL=0:编码器把均值直接输出 0、方差输出 1,等于彻底放弃这个维度。此时重建误差精确等于 λ——也就是把 x 全猜成 0 的水平,这个维度一点信息都没传。问题在于看 loss 曲线你只会看到「KL 顺利降到 0,总损失还在降」,很难意识到这是失败而不是收敛。 第二组:把 ELBO 当成似然来监控,会被方差骗。 在一个 D=6、K=3 的高斯线性模型上,真实 log p(x) = -5.434331,解析 ELBO = -7.419767,两者差 1.985436。但用训练时真正用的那个单样本蒙特卡洛估计量去估 ELBO,重复 4000 次,标准差是 4.308——比它和 log p(x) 之间的差距还大一倍多,结果有 40.7% 的采样值直接越过了 -5.434 这条「上界」。如果你在 tensorboard 上画这条线当似然看,会得出完全错误的结论。 第三组:不用重参数化,梯度方差能大到没法训练。 同一个模型上估 ∇ E_q[log p(x|z)],重参数化的梯度方差和是 377.55,得分函数估计量是 3043.11,差 8.1 倍。听起来还能忍?把数据维度 D 从 8 拉到 2048: D 重参数化方差 得分函数方差 比值 8 2.08 2 469 1 188× 64 1.98 104 328 52 630× 512 1.84 7 579 291 4 121 907× 2048 1.84 114 230 222 61 916 893× 重参数化的方差几乎不随 D 动,得分函数估计量按 D² 往上冲,到 D=2048(一张 45×45 的灰度图而已)比值是 6191 万倍。这不是「慢一点」的差别,是「能不能训」的差别。 这三组数字背后是同一件事:我们想最大化的是 log p(x),但它算不动,只能换成它的一个下界 ELBO 来优化;而换成下界之后,「换掉了什么」「这个下界的估计量有多吵」「梯度怎么穿过采样」这三笔账必须自己算清楚。不搞清楚 KL 项的价,你就不知道 β 该往哪调;不搞清楚下界的松紧,你就不知道扩散模型的训练目标从哪来。 02. 最小可用理解 三句话讲完核心: 算不动的边际似然换成一个能算的下界:$p_\theta(x)$ 要对所有 z 积分,算不动,于是引入一个我们自己挑的分布 $q_\phi(z \mid x)$,用 Jensen 不等式把优化目标换成 ELBO,它只需要在 q 下求期望,采样就能算。 下界和真值之间差的那一坨,正好是 q 和真实后验的 KL:$\log p_\theta(x) = \mathcal{L}(\theta, \phi) + \mathrm{KL}(q_\phi(z \mid x) \,\|\, p_\theta(z \mid x))$。所以「下界有多松」等价于「q 离真后验有多远」,而不是「q 离先验有多远」——后者的 KL 是 ELBO 里的一个加数,是另一回事。 要让这个下界能被反向传播,就把随机性从梯度路径上挪走:写成 $z = m_\phi + s_\phi \odot \epsilon$,$\epsilon$ 与参数无关,采样这一步变成确定性变换,梯度顺着 z 一路传回 $m_\phi$ 和 $s_\phi$。 03. 数学推导 3.1 边际似然为什么算不动 生成模型的设定很简单:先从一个固定先验里采隐变量 $z \sim p(z)$(通常是标准正态),再用解码器生成观测 $x \sim p_\theta(x \mid z)$。要拟合数据,目标是对数边际似然: $$\log p_\theta(x) = \log \int p_\theta(x \mid z) \, p(z) \, dz$$ 麻烦全在这个积分上。$p_\theta(x \mid z)$ 是个神经网络,塞进积分里没有任何闭式解;用数值积分的话,z 是 K 维,网格点数随 K 指数爆炸,K=32 就已经不可能。而 K 小了模型表达能力又不够——这正是我们不愿意接受的取舍。 3.2 换个能算的目标 既然积分算不动,就绕开它。引入在目标分布支撑上为正、满足相关可积条件的以 x 为条件的分布 $q_\phi(z \mid x)$(后面简称 $q$),把被积函数乘一个 $q/q$: $$\log p_\theta(x) = \log \int q_\phi(z \mid x) \, \frac{p_\theta(x, z)}{q_\phi(z \mid x)} \, dz = \log \mathbb{E}_{q_\phi}\!\left[\frac{p_\theta(x, z)}{q_\phi(z \mid x)}\right]$$ log 是凹函数,Jensen 不等式给出 $\log \mathbb{E}[Y] \ge \mathbb{E}[\log Y]$,于是 $$\log p_\theta(x) \ge \mathbb{E}_{q_\phi}\!\left[\log \frac{p_\theta(x, z)}{q_\phi(z \mid x)}\right] \equiv \mathcal{L}(\theta, \phi)$$ 右边就是 ELBO(Evidence Lower Bound)。它好算:$\log p_\theta(x, z) = \log p_\theta(x \mid z) + \log p(z)$ 两项都能直接求值,期望用采样估计就行。代价是我们不再直接优化 log p(x),而是优化它的一个下界。 3.3 差的那一项到底是什么 这是全文最关键的一步,也是最容易被跳过的一步。把 $\log p_\theta(x)$ 写成对 q 的期望(它和 z 无关,所以这么写是恒等的),再硬塞一个 $\log \frac{q}{q}$ 进去: $$\log p_\theta(x) = \mathbb{E}_{q}\!\left[\log p_\theta(x)\right] = \mathbb{E}_{q}\!\left[\log \frac{p_\theta(x, z)}{p_\theta(z \mid x)}\right] = \mathbb{E}_{q}\!\left[\log \frac{p_\theta(x, z)}{q_\phi(z \mid x)}\right] + \mathbb{E}_{q}\!\left[\log \frac{q_\phi(z \mid x)}{p_\theta(z \mid x)}\right]$$ 第一项正是 ELBO,第二项按定义就是 KL 散度。合起来: $$\log p_\theta(x) = \mathcal{L}(\theta, \phi) + \mathrm{KL}\big(q_\phi(z \mid x) \,\|\, p_\theta(z \mid x)\big)$$ KL 恒非负,所以 $\mathcal{L} \le \log p_\theta(x)$,和 Jensen 的结论一致——但这一版多给了一个信息:等号成立当且仅当 q 等于真实后验。下界的松紧完全由 q 的质量决定,跟别的都没关系。 3.4 把 ELBO 拆成两项看 把 $\log p_\theta(x, z)$ 展开,ELBO 可以写成更有物理含义的形式: $$\mathcal{L}(\theta, \phi) = \mathbb{E}_{q_\phi}\!\left[\log p_\theta(x \mid z)\right] - \mathrm{KL}\big(q_\phi(z \mid x) \,\|\, p(z)\big)$$ 第一项是重建项:从 q 里采 z,解码器能不能还原出 x。第二项是先验对齐项:q 别离先验太远——因为生成时我们是从先验采 z 的,如果 q 把 z 放到了先验覆盖不到的地方,生成阶段就对不上了。 这里有个必须记住的区分,后面第 08 节还会回到它: $\mathrm{KL}(q \,\|\, p(z))$ —— 出现在 ELBO 里,是要被最小化的一项,物理含义是「编码分布别跑太偏」。 $\mathrm{KL}(q \,\|\, p(z \mid x))$ —— 不出现在 ELBO 里,是下界的缝隙,物理含义是「q 离真后验还差多少」。 在我们那个 D=6 的例子里,这两个数是 3.364364 和 1.985436,不是一回事,也不是简单的包含关系。β-VAE 显式改的是先验 KL 的权重,但重新训练后 q 与生成模型都会变化,所以真实后验差距也会变化,且不保证变小。 顺带一个很容易漏掉的推论:把 KL 乘上 β 之后得到的目标 $$\mathcal{L}_\beta = \mathbb{E}_{q_\phi}\!\left[\log p_\theta(x \mid z)\right] - \beta \cdot \mathrm{KL}\big(q_\phi(z \mid x) \,\|\, p(z)\big)$$ 当 $\beta\ge1$,仍有 $\mathcal L_\beta=\mathcal L-(\beta-1)\mathrm{KL}(q\Vert p)\le\mathcal L\le\log p_\theta(x)$,所以它仍是下界,只是更松;$0<\beta<1$ 时则不再保证下界。不同 β 的目标包含不同惩罚,不能直接当成同一种似然估计比较。 3.5 梯度怎么穿过采样 ELBO 对解码器参数 $\theta$ 的梯度没问题,重建项直接可导。麻烦在对 $\phi$ 的梯度:期望的分布本身依赖于 $\phi$,而采样操作不可导。设 $f(z) = \log p_\theta(x \mid z)$,要估的是 $\nabla_\phi \mathbb{E}_{q_\phi}[f(z)]$。 方法一:得分函数估计量(REINFORCE)。 直接把导数挪进期望: $$\nabla_\phi \mathbb{E}_{q_\phi}[f(z)] = \mathbb{E}_{q_\phi}\!\left[f(z) \, \nabla_\phi \log q_\phi(z \mid x)\right]$$ 在可交换微分与积分、分布支持集不随参数变化等常见正则条件下,这个式子成立;它不要求 f 对 z 可导,也不要求 z 连续——代价是方差极大,因为整个 f 的量级都被乘进了梯度里。 方法二:重参数化。 把 z 写成参数的确定性函数外加一个与参数无关的噪声: $$z = m_\phi(x) + s_\phi(x) \odot \epsilon, \quad \epsilon \sim \mathcal{N}(0, I)$$ 于是期望可以改写成对 $\epsilon$ 的期望,导数直接进去: $$\nabla_\phi \mathbb{E}_{q_\phi}[f(z)] = \mathbb{E}_{\epsilon}\!\left[\nabla_\phi f\big(m_\phi(x) + s_\phi(x) \odot \epsilon\big)\right] = \mathbb{E}_{\epsilon}\!\left[\nabla_z f(z) \cdot \nabla_\phi z\right]$$ 关键在于 $\nabla_z f$ 用到了 f 对 z 的局部形状,相当于「知道往哪个方向挪 z 会让 f 变大」,这是方法一完全没有利用的信息。 方差为什么差这么多——一个具体的机制。 在高斯解码器下 $$f(z) = -\frac{D}{2}\log(2\pi\sigma^2) - \frac{\|x - Wz - b\|^2}{2\sigma^2}$$ 第一项与 z 无关,是个常数。它对 $\nabla_z f$ 的贡献恒为 0,重参数化天然看不见它;但方法一要把整个 f(包括这个常数)乘上 $\nabla_\phi \log q$,常数按平方进方差。$D=2048$、$\sigma^2=1$ 时这个常数是 $-\frac{2048}{2}\log 2\pi \approx -1877$,平方之后就是 350 万量级——这正是第三组数字里那个 6191 万倍的来源。 实验也验证了这一点:给方法一减掉一个「预言机 baseline」(把 f 换成 $f - \mathbb{E}_q[f]$,常数的均值被减掉,期望不变),方差比从 61 916 893× 掉到 5.3×。所以差距主要来自那个常数项,不是来自「采样本身」。 3.6 下界能有多紧:IWAE 单样本 ELBO 的缝是 $\mathrm{KL}(q \,\|\, p(z \mid x))$。要缩缝有两条路:把 q 变强(换更灵活的后验族),或者换一个更紧的界。后者的经典做法是 IWAE 的 k 样本界: $$L_k = \mathbb{E}_{z_1 \dots z_k \sim q_\phi}\!\left[\log \frac{1}{k} \sum_{l=1}^{k} \frac{p_\theta(x, z_l)}{q_\phi(z_l \mid x)}\right]$$ 注意顺序:先对 k 个重要性权重取平均,再取 log。单样本 ELBO 是「log 再取平均」(每个样本的 log 比值取期望),IWAE 是「平均再 log」,Jensen 保证后者更紧。可以证明 $\log p_\theta(x) \ge L_{k+1} \ge L_k \ge L_1 = \mathcal{L}$,在重要性权重满足支持覆盖与可积性等条件下,$k \to \infty$ 时收敛到 $\log p_\theta(x)$。 「换个顺序就更紧」这件事直觉上可以这么理解:单样本 ELBO 每次只看一个 z,好坏全押在它身上;IWAE 一次看 k 个,其中只要有一个 z 的 $p_\theta(x, z)/q_\phi(z \mid x)$ 特别大,平均权重就被拉上去,而 log 是凹函数,对这个「运气好」的样本惩罚得比线性小。所以 k 越大,越有机会碰到好样本,界越紧——本质上是用更多的采样换取更少的方差,跟蒙特卡洛里加样本降噪是同一回事,但这里降的是偏差而不是方差。 代价是:梯度不再是对单个样本的简单求和,k 个权重互相耦合(每个权重的梯度都带上了其他权重的归一化因子),总计算量随 k 线性上涨;编码器梯度的信噪比还可能随 k 增大而下降(见 Rainforth 等,2018),这反而不利于训练编码器——工程上这是个明确的取舍,不是免费的午餐。 这张图要看什么:左图四根柱子——ELBO(-7.420)加 KL(q‖p(z|x))(1.985)正好等于 log p(x)(-5.434),而 KL(q‖p(z))(3.364)是另一根完全不同的柱子,别混。中图是单样本蒙特卡洛 ELBO 的 4000 次采样分布,标准差 4.31,虚线是真实上界 -5.434,40.7% 的样本落在它右边。右图是 IWAE 的 k 扫描,k 从 1 涨到 500,缝隙从 2.049 缩到 0.0026——同样是用重要性采样,换个顺序就差三个数量级。 04. 代码实现 4.1 恒等式核对:先造一个能算到机器精度的模型 要验证「差的那一项到底是什么」,必须三样东西都有解析解:边际似然、真实后验、ELBO。高斯线性模型满足这一点:$z \sim \mathcal{N}(0, I_K)$,$x \mid z \sim \mathcal{N}(Wz + b, \sigma^2 I_D)$。边际仍是高斯,后验是共轭高斯,ELBO 也能闭式算。 import numpy as np D, K, SIGMA2 = 6, 3, 0.35 ** 2 W = np.random.default_rng(0).normal(0, 1, size=(D, K)) * 0.9 b = np.random.default_rng(1).normal(0, 0.3, size=(D,)) def true_posterior(x): """p(z|x) 的闭式解:precision = I + W^T W / sigma^2。""" prec = np.eye(K) + W.T @ W / SIGMA2 cov = np.linalg.inv(prec) return cov @ (W.T @ (x - b) / SIGMA2), cov def elbo_analytic(x, m, s): """ELBO 的解析值。注意 ||x - Wz - b||^2 对 q 求期望会多出一项 trace。""" var = s ** 2 resid = x - (W @ m + b) recon = -0.5 * (D * np.log(2 * np.pi * SIGMA2) + (resid @ resid + np.trace(W.T @ W @ np.diag(var))) / SIGMA2) kl = 0.5 * np.sum(var + m ** 2 - 1.0 - np.log(var)) return float(recon - kl), float(recon), float(kl) def kl_gauss_gauss(m, s, mean2, cov2): """KL(N(m, diag(s^2)) || N(mean2, cov2)),两个高斯的闭式。""" var, diff = s ** 2, mean2 - m cov2_inv = np.linalg.inv(cov2) _, logdet2 = np.linalg.slogdet(cov2) return 0.5 * (np.trace(cov2_inv @ np.diag(var)) + diff @ cov2_inv @ diff - K + logdet2 - np.sum(np.log(var))) 重建项里那个 np.trace(W.T @ W @ np.diag(var)) 是最容易漏掉的一项:对 $z$ 求期望时,$\|x - Wz - b\|^2$ 里的 $Wz$ 项会因为 z 的随机性多出一份方差贡献。漏了它,ELBO 就不是下界了。 关键一步是故意把 q 指定错:取真实后验的均值再加扰动,协方差只留对角线并缩小 20%,让 q 比真后验更自信且忽略相关性。这样 KL(q‖p(z|x)) 严格大于 0,缝才看得见。跑起来: [1] 解析三项 E_q[log p(x|z)] = -4.055403 KL(q || p(z)) = 3.364364 ELBO = -7.419767 KL(q || p(z|x)) = 1.985436 ELBO + KL(q||p(z|x)) = -5.434331 log p(x) = -5.434331 |误差| = 8.882e-16 恒等式对到 1e-15。同时注意:KL(q‖p(z))=3.364 比 KL(q‖p(z|x))=1.985 还大——这两个量的大小关系没有必然规律,别用其中一个去猜另一个。 4.2 两种梯度估计量,方差实测 同一个模型上,$\nabla_m$ 和 $\nabla_\ell$($\ell = \log s$)的解析梯度都能写出来,所以「谁对谁错」有标准答案,剩下的差别纯粹是方差: def sample_reparam(x, m, s, W, b, sigma2, rng, n): """重参数化:z = m + s ⊙ ε,梯度经 z 反传。""" eps = rng.normal(size=(n, m.size)) z = m + s * eps score_z = (x - (z @ W.T + b)) @ W / sigma2 # ∇_z log p(x|z) return score_z, score_z * s * eps # ∂z/∂m = 1, ∂z/∂ℓ = s ⊙ ε def sample_score(x, m, s, W, b, D, sigma2, rng, n, baseline=None): """得分函数:∇ E[f] = E[f · ∇ log q],f = log p(x|z)。""" eps = rng.normal(size=(n, m.size)) z = m + s * eps r = x - (z @ W.T + b) f = -0.5 * (D * np.log(2 * np.pi * sigma2) + np.einsum("nd,nd->n", r, r) / sigma2) if baseline is not None: f = f - baseline return f[:, None] * (eps / s), f[:, None] * (eps ** 2 - 1.0) 两行代码的差别就是全部:重参数化用的是 $\nabla_z f$,得分函数用的是 $f$ 本身。跑 40000 次独立采样: 真实梯度(解析) ∇_m = [ 0.4228 4.9273 -0.6448] ∇_ℓ = [-1.9319 -2.4426 -3.8296] 估计量对比(40000 次独立采样,方差 = 6 个参数分量方差之和) 重参数化 方差和= 377.55 均值相对误差= 0.41% 达到5%需 n≈3068 得分函数(无 baseline) 方差和= 3043.11 均值相对误差= 3.95% 达到5%需 n≈24722 得分函数(预言机 baseline) 方差和= 1586.45 均值相对误差= 3.09% 达到5%需 n≈12889 方差比:得分函数 / 重参数化 = 8.1× 加 baseline 之后 = 4.2× 三个估计量的均值都对(无偏),差别全在方差。换算成「达到 5% 相对误差需要多少样本」,是 3068 对 24722。 这张图要看什么:左图是 $\nabla_m$ 第 0 个分量(真值 0.4228)在三种估计量下的分布——均值都压在真值附近,标准差分别是 5.7、16.9、12.0,尾巴长度差一个量级。右图的双对数坐标里,重参数化那条线几乎是平的(方差 2.08 → 1.84),得分函数那条严格贴着 D² 参考线往上走;而把蓝色点(加了 baseline)和橙色点对比,比值从 1188× 一路到 6191 万×,加完 baseline 却稳定在 5~7×——说明炸掉的部分是那个常数项,不是采样噪声。 4.3 β 加权:把维度死掉的过程解出来 这一节的目标是要一个能解到全局最优的玩具模型,这样「维度死了」是解析结论而不是训练运气。构造:每个坐标独立,$x_j \sim \mathcal{N}(0, \lambda_j)$,$q(z_j \mid x_j) = \mathcal{N}(a_j x_j, s_j^2)$,$p(x_j \mid z_j) = \mathcal{N}(b_j z_j, \sigma^2)$。每个维度就是一份独立副本,逐个维度单独求最优即可: SIGMA2 = 0.25 # 解码器观测噪声方差 def kl_of(a, lam, s): return 0.5 * (a ** 2 * lam + s ** 2 - 1.0 - np.log(s ** 2)) def recon_mse(a, b, lam, s): """E_{x,z}[(x - b z)^2],重建误差的期望。""" return lam * (1.0 - b * a) ** 2 + b ** 2 * s ** 2 def objective(p, lam, beta): a, b, ell = p s = np.exp(ell) return -recon_mse(a, b, lam, s) / (2 * SIGMA2) - beta * kl_of(a, lam, s) 目标对 $(a, b, \ell)$ 的梯度全部解析,脚本里再用中心差分核对一遍(实测最大偏差 5.5e-10),避免推错。优化必须多起点——$(a, b) = (0, 0)$ 就是「维度死亡」解,单起点很容易掉进去出不来。 跑完 β 扫描,把最优解和闭式预测放在一起对: λ β s²实测 βσ²/λ a²λ实测 1-s² MSE实测 βσ² 1.0000 1.00 0.25000 0.25000 0.75000 0.75000 0.25000 0.25000 1.0000 2.00 0.50000 0.50000 0.50000 0.50000 0.50000 0.50000 1.0000 3.00 0.75000 0.75000 0.25000 0.25000 0.75000 0.75000 0.2500 0.50 0.50000 0.50000 0.50000 0.50000 0.12500 0.12500 最大偏差 = 5.53e-03 → 存活时最优解确实落在闭式上 λ=4.0000 理论死亡阈值 β* = λ/σ² = 16.00 λ=1.0000 理论死亡阈值 β* = λ/σ² = 4.00 λ=0.2500 理论死亡阈值 β* = λ/σ² = 1.00 λ=0.0625 理论死亡阈值 β* = λ/σ² = 0.25 存活时最优解落在闭式 $s^2 = \beta\sigma^2/\lambda$、$a^2\lambda = 1 - s^2$、$\mathrm{MSE} = \beta\sigma^2$ 上。这个闭式有个反直觉的推论:存活时的重建误差只由 β 和解码器噪声决定,跟这个维度携带多少信息 λ 无关——λ 只决定「这个维度值不值得用」。 05. 工业级实现对照 最小实现里 KL 是两个高斯的闭式,工业代码也是闭式,但有几个地方长得不一样,值得逐个说明为什么。 1. 官方最小实现:pytorch/examples。 pytorch/examples 的 vae/main.py 里,loss_function 就是标准的 BCE 重建加闭式 KL,没有任何 β(以 2026-09 时的实现为准): def loss_function(recon_x, x, mu, logvar): BCE = F.binary_cross_entropy(recon_x, x.view(-1, 784), reduction='sum') # 0.5 * sum(1 + log(sigma^2) - mu^2 - sigma^2) KLD = -0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp()) return BCE + KLD 和本文 3.4 节那一项完全对得上:$\mathrm{KL} = \frac{1}{2}\sum(\mu^2 + s^2 - 1 - \log s^2)$,代码里写成 -0.5 * sum(1 + logvar - mu^2 - exp(logvar)),logvar 就是 $\log s^2$。要加 β 得自己动手,官方示例没有。 2. 为什么参数化 logvar 而不是 s。 方差必须为正,直接学 s 要加约束;学 $\log s^2$ 则值域是全体实数,网络怎么输出都合法。diffusers 的 DiagonalGaussianDistribution 还额外做了截断(以 2026-09 时的实现为准): self.mean, self.logvar = torch.chunk(parameters, 2, dim=1) self.logvar = torch.clamp(self.logvar, -30.0, 20.0) self.std = torch.exp(0.5 * self.logvar) self.var = torch.exp(self.logvar) clamp 到 [-30, 20] 是为了数值稳定:$\exp(20)$ 约为 4.85 亿,fp32 可表示但 fp16 已无法表示;下界 -30 避免过小方差与过大负 logvar。还需核对计算 dtype,截断本身不是所有精度下的安全保证。 3. 采样就是重参数化,一行代码。 同一个类里: def sample(self, generator=None): sample = randn_tensor(self.mean.shape, generator=generator, device=self.parameters.device, dtype=self.parameters.dtype) x = self.mean + self.std * sample return x 就是本文 3.5 的 $z = m + s \odot \epsilon$。mode() 直接返回均值——一些确定性编码场景使用 mode(),但 diffusers img2img 默认使用后验 sample(generator=...),训练和推理均可采样。 4. KL 的求和维度与归一化。 diffusers 里 kl() 写成(以 2026-09 时的实现为准): return 0.5 * torch.sum(torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar, dim=[1, 2, 3]) 对通道、高、宽三个维度求和,得到的是每个样本一个标量,然后训练脚本再对 batch 取平均。这个区别很实际:如果对全部分量求和再除以元素总数,等价于把每个维度的 KL 权重除以 $C \times H \times W$——潜变量维度一多,KL 项的有效权重会相应减小,可能使先验对齐变弱、尺度漂移;过强 KL 才更直接推动后验坍缩。 5. 训练时 KL 根本不在模型里。 值得注意的是 AutoencoderKL 的模型文件里没有 kl_loss 这个东西,它只通过 encode() 返回 DiagonalGaussianDistribution 后验,KL 由外部训练脚本算。这是工程上的职责划分:模型只负责给出分布,损失函数怎么组合(KL 权重、perceptual loss、对抗损失)交给训练配置。 6. 潜空间的缩放。 SD1.x 常见 scaling_factor=0.18215 是乘到原始 latent 上的系数,对应原始标准差约 $1/0.18215\approx5.49$,缩放后才接近 1,不能把 0.18 当成原始潜空间尺度。它来自特定 VAE 与扩散训练的尺度约定,不由 ELBO 唯一决定。 7. 最小实现和工业实现的差异。 工业代码增加 logvar 截断、dtype 控制、分布采样接口及 latent 尺度约定。KL 逐维求和来自概率定义,batch 平均来自训练目标;它们不是任意工程细节。sample() 与 mode() 都有推理用途,要依具体管线选择。 06. 代价与边界 省了什么。 把一个对 K 维积分的不可解问题,变成了「采样 + 两个可导项」,梯度能用标准反向传播算。这是 VAE、扩散模型、以及一大票潜变量模型能训练起来的全部前提。 赔了什么,四条。 一是下界不是似然。 优化 ELBO 不等于优化 log p(x),中间隔着 KL(q‖p(z|x))。在固定生成模型下,这个缝受 q 族表达能力与实际优化结果共同影响:均值场对角高斯拟合不了有相关性的真实后验,缝就永远在。本例里缝是 1.985 nats,在一个 log p(x) 只有 -5.43 的小模型上,数值上差约 36%,但连续对数密度受坐标单位影响,不宜把这个比例当作通用误差尺度。 二是 KL 的方向有偏好。 $\mathrm{KL}(q\Vert p(z\mid x))$ 是 reverse KL:当 $q$ 把概率放到真实后验很低的地方,代价很大,因此在受限的单峰近似族下常表现为 mode-seeking,可能只覆盖一个峰;反向的 $\mathrm{KL}(p\Vert q)$ 则倾向覆盖目标的质量,即 mass-covering。VAE 的模糊不能单归因于这个方向,逐像素重建目标、解码器分布与压缩瓶颈都有关。 三是推断被参数化摊平(amortized)之后又多了一层误差。 经典变分推断对每个 x 单独优化一组变分参数;VAE 用一个共享的编码器网络去输出所有 x 的 $m_\phi(x)$ 和 $s_\phi(x)$。这一步是为了快——测试时一次前向就得到后验,不用重新迭代——但它意味着 q 的可行域被限制在「神经网络能表达的那些分布」里。即便每个 x 单独看,最优的对角高斯后验就在那儿,共享编码器也可能一辈子到不了。所以总的缝其实是两笔账叠起来的:函数族的近似误差加上摊销误差。这也解释了为什么给编码器加容量在有些任务上能明显提 ELBO——在固定生成模型时,这改善的是近似推断。 四是加权 KL 会直接杀死维度。 这是第 4.3 节那张表的完整结论: λ=4.0000 在 β≤12.0 内都存活 λ=1.0000 从 β=4.0 起死亡(β=1 时 KL=0.6931) λ=0.2500 从 β=1.0 起死亡 λ=0.0625 从 β=0.25 起死亡 死亡阈值是 $\beta^\ast = \lambda / \sigma^2$。读出这个式子的含义:一个维度要活下来,它携带的信息量 λ 必须盖过「KL 的价」β 乘上「解码器噪声」σ²。β 翻倍,能活下来的维度门槛就翻倍;解码器越准(σ² 越小),越多的维度能活。这仅解释线性高斯模型的阈值。强大的自回归解码器可通过其他条件预测数据而忽略 z,不能简单等同于把固定高斯观测方差 σ² 调小。 这张图要看什么:左图四条曲线是不同 λ 下重建 MSE 随 β 的变化,实心点表示维度存活、空心方块表示已死——注意存活段的 MSE 就是 $\beta\sigma^2$ 这条直线,跟 λ 无关;而死掉的段 MSE 平在 λ 上不再变化。右图是 $(\lambda, \beta)$ 平面上的相图,斜线是 $\beta = \lambda/\sigma^2$,线右上方全死、左下方全活。调 β 之前先看一眼这个平面:你真正要判断的是「这条线上方还有多少维度」。 缓解手段:free bits。 把每组 KL 换成 $\max(\mathrm{KL},C)$,是在 KL 小于 C 时去掉进一步压低它的梯度,给重建项使用这部分容量的机会;这不是保证至少传 C nats 的硬约束。下面的线性实验确实在边界附近找到更好的重建,但一般神经网络不保证被救活。实测(β=4): λ C=0(纯 β) C=0.05 C=0.2 1.0000 a=0.000 KL=0.000 MSE=1.000 a=-0.308 KL=0.050 MSE=0.905 a=-0.574 KL=0.200 MSE=0.670 0.2500 a=0.000 KL=0.000 MSE=0.250 a=-0.617 KL=0.050 MSE=0.226 a=-1.148 KL=0.200 MSE=0.168 λ=1.0 的维度被救回来了(MSE 1.000 → 0.670),λ=0.25 的也从 0.250 降到 0.168。代价是放松了对先验 KL 的惩罚;是否改善或损害解耦需要独立评估——这是一笔明码标价的交易。 什么时候不该用。 归一化流和自回归模型在相应建模假设下可计算精确似然,不一定需要变分下界;VQ 模型可精确计算离散 token 序列的自回归概率,但一般仍不能精确边缘化得到像素似然。ELBO 也不是训练高维生成模型的唯一途径,score matching、流匹配和对抗训练是其他路线。同一数据、同一似然约定下标准 ELBO 可比较为下界,但 q 的质量影响松紧,不能直接据此断言真实似然或感知质量的排序。 07. 经典论文脉络 Kingma & Welling, 1312.6114(2013)——VAE 原文。贡献是把「变分推断 + 重参数化 + 神经网络编码器」拼成一个能用 SGD 训的东西,并给出 SGVB 估计量。它同时确立了沿用至今的 loss 形式:重建项减 KL。本文 3.2~3.5 节基本是这篇的复述。 Rezende, Mohamed & Wierstra, 1401.4082(2014)——几乎同期独立提出的重参数化,论文里叫 stochastic backpropagation。贡献是把这个方法从「VAE 的一个技巧」推广成「任何可微分概率模型上的通用推断方法」,并系统讨论了高斯之外的分布族怎么处理。想理解重参数化的适用边界(哪些分布能做、哪些只能退回得分函数),这篇比 VAE 原文讲得更清楚。 Burda, Grosse & Salakhutdinov, 1509.00519(2015)——IWAE。贡献就是本文 3.6 那个 $L_k$:把「log 再平均」换成「平均再 log」,得到一个随 k 单调变紧的界,并证明了收敛性。它澄清了一个当时普遍的误解——「多采几个样本只是降方差」,实际上换的是界本身。 Higgins et al., ICLR 2017——β-VAE。给 KL 项加权重 β > 1,换来更好的解耦表征。本文第 4.3 节那张表就是它的代价面:β 每翻一倍,重建误差也翻一倍,且维度按 $\beta^\ast = \lambda/\sigma^2$ 逐个死掉。 Bowman et al., 1511.06349(2015)——后验坍缩最早被认真对待的现场。用 VAE 做句子生成,强大的自回归解码器会直接忽略潜变量,KL 掉到 0。这篇提出的 KL annealing(训练初期把 β 从 0 慢慢涨到 1)至今还是最实用的缓解手段之一,效果依赖具体模型和优化过程,并非保证解决坍缩。 补充两篇:Kingma et al. 的 IAF(1606.04934)走的是另一条路——不改界,改 q,用可逆变换把后验族变强来直接缩小那个缝;Kingma & Welling 的综述(1906.02691)适合把上述脉络串起来通读。 08. 常见误解 误解一:「ELBO 里的 KL 就是 q 和真实后验的差距」。 ELBO 里减去的是 $\mathrm{KL}(q\Vert p(z))$;下界的差距是 $\mathrm{KL}(q\Vert p(z\mid x))$。调整 β 会通过训练改变 q 和生成模型,后者也可能变化,只是没有保证随前者一起下降。 误解二:「ELBO 是下界,所以蒙特卡洛估计值不会超过 log p(x)」。 下界性质是对期望成立的,不是对每个样本成立。本例单样本估计的标准差是 4.308,而缝只有 1.985,结果 40.7% 的样本越过了真实上界。看到自己的「ELBO」比之前算的 log p(x) 还大时,先别怀疑代码,这是正常的采样噪声。 误解三:「重参数化是为了让采样可微」。 更准确的说法是为了降低梯度估计量的方差。采样「不可微」这个表述本身就有问题:得分函数估计量里 $\nabla_\phi \log q_\phi$ 是对参数求导,完全可微,它对 f 连可导性都不要求。真正的区别是重参数化用上了 $\nabla_z f$ 这个局部信息,方差低几个数量级(D=2048 时差 6191 万倍)。 误解四:「KL 掉到 0 说明 KL 项优化到位了」。 这是后验坍缩的典型症状。本例 λ=1.0、β=4 时最优解就是 a=0、s=1、KL=0、MSE=1.000(正好等于 λ),编码器彻底放弃了这个维度。判断方法不是看 KL,而是看 $a^2\lambda$(Higgins 的 active unit 判据)——本文脚本里 active_unit() 就是干这个的。 误解五:「β 越大解耦越好,重建变差只是小代价」。 在线性模型中,最优 MSE 为 $\min(\beta\sigma^2,\lambda)$,连续增加后饱和;每维 KL 也连续趋零。离散的存活计数会出现台阶,但不意味着误差发生跳变。更强的 KL 约束还可能让所有维度关闭,不能只用 β 大小判断解耦质量。 09. 动手验证 三个脚本都只依赖 numpy,复制下来直接跑(完整代码见文末附录)。 python elbo_identity.py # 约 5 秒 python reparam_gradients.py # 约 15 秒 python beta_kl_weight.py # 约 3 分钟(多起点优化) 实验一:把 q 的均值和边际方差改成真后验的对应量,看残留差距。 改 elbo_identity.py 里的 make_q,把 var = np.diag(cov).copy() * 0.8 + 0.05 改成直接返回真实后验的对角(var = np.diag(cov).copy()),均值扰动也去掉。预期:KL(q‖p(z|x)) 会从 1.985 掉下来但不会掉到 0——因为对角 q 还是拟合不了真实后验的相关性。这个残留量就是「均值场假设的代价」,值得亲眼看一次。 实验二:给得分函数估计量换一个真实的 baseline。 reparam_gradients.py 里用的是预言机 baseline($\mathbb{E}_q[f]$ 的解析值),实战拿不到。把它换成一个滑动平均的 f(比如前 100 个样本的均值)再跑,预期方差比从 8.1× 降到接近 4.2× 的水平——这解释了为什么 REINFORCE 类算法里 baseline 是标配而不是优化项。 实验三:给 β 扫描加一个 free bits,看能救回几个维度。 beta_kl_weight.py 已经内置 free_bits 参数,把 main() 里的 optimize(lam, 4.0, free_bits=C) 的 C 从 0.05 调到 0.5 再跑。预期:这个线性实验中 C=0.5 时低方差维度也能使用 latent,KL 可能停在 C 附近;这不代表一般模型保证传满 C nats,也不能从 KL 数值单独判断解耦程度。这就是那笔交易的完整价格表。 10. 延伸阅读 按知识树的依赖关系,从这篇出发有三个方向: 往下(直接依赖本文):扩散模型的 SDE 视角——把 DDPM 的离散公式和连续 SDE 对上,会发现它的训练目标就是一层层套起来的 ELBO,本文 3.3 那个恒等式在那边换个马甲反复出现;视频 VAE 的潜空间与 loss 组合;其中 scaling_factor 的来历见 VAE 结构与训练目标。 往旁(同一个估计量,另一副马甲):策略梯度与 PPO 基础——本文 3.5 的得分函数估计量在 RL 里叫 REINFORCE,方差问题一模一样,baseline 的作用也一模一样。读过这篇再去看 PPO,会省掉一次重新理解。 已发布的相关篇目:视频 VAE 的常见 loss 组合——KL 之外还有哪些 loss 会加进来,以及它们各自的量级;FlashAttention 为什么不需要存下注意力矩阵 与 性能建模与 Profiling——本文没碰的工程侧。 完整的知识树见博客目录页 AIGC 基本功知识树,按依赖顺序排好了先修课。 附录:完整代码 09 节用到的脚本全文如下(elbo_identity.py、reparam_gradients.py、beta_kl_weight.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 elbo_identity.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """ELBO 恒等式核对:log p(x) = ELBO + KL(q || p(z|x))。 构造一个「后验有解析解」的模型,这样三样东西都能算到机器精度: * 边际似然 log p(x) —— 高斯线性模型,边际仍是高斯 * 真实后验 p(z|x) —— 共轭高斯,闭式解 * ELBO —— 重建项对 q 可解析求期望,KL 两项都是高斯闭式 有解析解才能验证「差的那一项到底是什么」,靠采样是验不出来的。 只依赖 numpy,直接 `python elbo_identity.py` 即可运行。 """ import numpy as np RNG = np.random.default_rng(20260925) # ── 模型:z ~ N(0, I_K),x|z ~ N(W z + b, sigma^2 I_D) ── D, K = 6, 3 SIGMA = 0.35 # 解码器观测噪声标准差 SIGMA2 = SIGMA ** 2 W = RNG.normal(0.0, 1.0, size=(D, K)) * 0.9 # 解码器权重 b = RNG.normal(0.0, 0.3, size=(D,)) # 解码器偏置 def make_one_x(rng=None): """从真实的边际分布里采一个 x,并返回它的解析 log p(x)。 rng 可显式传入:这样 reparam_gradients / make_figures 拿到的是同一个 x, 不会因为模块级 RNG 被别人消耗过而对不上数字。 """ if rng is None: rng = RNG z = rng.normal(size=(K,)) x = W @ z + b + SIGMA * rng.normal(size=(D,)) cov = W @ W.T + SIGMA2 * np.eye(D) sign, logdet = np.linalg.slogdet(cov) d = x - b quad = d @ np.linalg.solve(cov, d) log_px = -0.5 * (D * np.log(2 * np.pi) + logdet + quad) return x, log_px def true_posterior(x): """p(z|x) 的闭式解:precision = I + W^T W / sigma^2。""" prec = np.eye(K) + W.T @ W / SIGMA2 cov = np.linalg.inv(prec) mean = cov @ (W.T @ (x - b) / SIGMA2) return mean, cov # ── 变分后验 q(z|x) = N(m, diag(s^2)):故意用「对角」错误指定 ── def make_q(x, rng=None): """取真实后验的均值,协方差只留对角线并加一点扰动。 真实后验是有相关性的(cov 非对角),q 强行对角 => KL(q||p(z|x)) > 0, 这正是我们要留出来的那条缝。 """ if rng is None: rng = RNG mean, cov = true_posterior(x) m = mean + 0.15 * rng.normal(size=(K,)) # 均值也偏一点 var = np.diag(cov).copy() * 0.8 + 0.05 # 方差偏小 => q 过自信 return m, np.sqrt(var) def log_gauss_diag(z, mean, var, ): """log N(z; mean, diag(var))(省略常数也行,但这里算全)。""" return float(-0.5 * np.sum(np.log(2 * np.pi * var) + (z - mean) ** 2 / var)) def log_joint(x, z): """log p(x, z) = log p(x|z) + log p(z)。""" r = x - (W @ z + b) log_lik = -0.5 * (D * np.log(2 * np.pi * SIGMA2) + r @ r / SIGMA2) log_prior = -0.5 * (K * np.log(2 * np.pi) + z @ z) return log_lik + log_prior def elbo_analytic(x, m, s): """ELBO 的解析值。 E_q[log p(x|z)] 里 ||x - Wz - b||^2 对 q 求期望: ||x - Wm - b||^2 + tr(W^T W diag(s^2)) 第二项是采样噪声经过解码器放大出来的那部分,容易被漏掉。 """ var = s ** 2 resid = x - (W @ m + b) recon = -0.5 * (D * np.log(2 * np.pi * SIGMA2) + (resid @ resid + np.trace(W.T @ W @ np.diag(var))) / SIGMA2) kl = 0.5 * np.sum(var + m ** 2 - 1.0 - np.log(var)) return float(recon - kl), float(recon), float(kl) def kl_gauss_gauss(m, s, mean2, cov2): """KL(N(m, diag(s^2)) || N(mean2, cov2)),两个高斯的闭式。""" var = s ** 2 diff = mean2 - m cov2_inv = np.linalg.inv(cov2) trace = float(np.trace(cov2_inv @ np.diag(var))) quad = float(diff @ cov2_inv @ diff) _, logdet2 = np.linalg.slogdet(cov2) return 0.5 * (trace + quad - K + logdet2 - np.sum(np.log(var))) def mc_elbo(x, m, s, L, rng): """蒙特卡洛 ELBO:L 个样本取平均(L=1 就是训练时实际用的那个)。""" eps = rng.normal(size=(L, K)) z = m + s * eps log_lik = np.array([-0.5 * (D * np.log(2 * np.pi * SIGMA2) + np.sum((x - (W @ zz + b)) ** 2) / SIGMA2) for zz in z]) var = s ** 2 kl = 0.5 * np.sum(var + m ** 2 - 1.0 - np.log(var)) return float(log_lik.mean() - kl) def iwae_log_bound(x, m, s, k, rng): """IWAE 的 k 样本界:log (1/k) sum_l p(x,z_l)/q(z_l),用 logsumexp 稳算。""" eps = rng.normal(size=(k, K)) z = m + s * eps var = s ** 2 logw = np.array([log_joint(x, zz) - log_gauss_diag(zz, m, var) for zz in z]) mx = logw.max() return float(mx + np.log(np.mean(np.exp(logw - mx)))) # ── 下面两个表被 main 与 make_figures 共用,种子固定 => 数字永远一致 ── MC_LS = (1, 4, 16, 64) MC_REP = 4000 IW_KS = (1, 5, 50, 500) IW_REP = 6000 STATE_SEED = 20260925 def state(): """文章与配图共用的那一份状态:同一个 x、同一个 q。 用固定种子重放,保证单独跑本脚本与跑 make_figures 拿到同一组数字。 """ rng = np.random.default_rng(STATE_SEED) x, log_px = make_one_x(rng) m, s = make_q(x, rng) post_mean, post_cov = true_posterior(x) elbo, recon, kl = elbo_analytic(x, m, s) gap = kl_gauss_gauss(m, s, post_mean, post_cov) return dict(x=x, log_px=log_px, m=m, s=s, elbo=elbo, recon=recon, kl=kl, gap=gap) def mc_table(st): """蒙特卡洛 ELBO 的均值 / 标准差 / 越过 log p(x) 的比例。""" rng = np.random.default_rng(7) out = {} for L in MC_LS: vals = np.array([mc_elbo(st["x"], st["m"], st["s"], L, rng) for _ in range(MC_REP)]) out[L] = dict(mean=float(vals.mean()), std=float(vals.std()), frac=float(np.mean(vals > st["log_px"])), vals=vals) return out def iwae_table(st): """IWAE 的 k 样本界。""" rng = np.random.default_rng(13) out = {} for k in IW_KS: vals = np.array([iwae_log_bound(st["x"], st["m"], st["s"], k, rng) for _ in range(IW_REP)]) out[k] = dict(mean=float(vals.mean()), std=float(vals.std()), gap=float(st["log_px"] - vals.mean()), vals=vals) return out def main(): print("=" * 68) print("ELBO 恒等式核对 log p(x) = ELBO + KL(q || p(z|x))") print("=" * 68) print(f"模型: D={D}, K={K}, sigma={SIGMA}") st = state() x, log_px, m, s = st["x"], st["log_px"], st["m"], st["s"] elbo, recon, kl, gap = st["elbo"], st["recon"], st["kl"], st["gap"] print(f"\n[1] 解析三项") print(f" E_q[log p(x|z)] = {recon: .6f}") print(f" KL(q || p(z)) = {kl: .6f}") print(f" ELBO = {elbo: .6f}") print(f" KL(q || p(z|x)) = {gap: .6f}") print(f" ELBO + KL(q||p(z|x)) = {elbo + gap: .6f}") print(f" log p(x) = {log_px: .6f}") print(f" |误差| = {abs(elbo + gap - log_px): .3e}") # ── 蒙特卡洛波动:ELBO 的估计量是无偏的,不是恒小于 log p(x) ── print(f"\n[2] 蒙特卡洛估计的波动(每项 {MC_REP} 次重复)") print(f" {'L':>5s} {'mean':>12s} {'std':>10s} {'超过 log p(x) 的比例':>22s}") mc = mc_table(st) for L in MC_LS: r = mc[L] print(f" {L:5d} {r['mean']:12.5f} {r['std']:10.5f} {r['frac']:21.1%}") print(f" log p(x) = {log_px:.5f}(上界本身),ELBO(解析) = {elbo:.5f}") # ── IWAE:k 越大越紧 ── print(f"\n[3] IWAE 的 k 样本界(每项 {IW_REP} 次重复)") print(f" {'k':>5s} {'mean L_k':>12s} {'std':>9s} {'log p(x) - L_k':>16s}") iw = iwae_table(st) for k in IW_KS: r = iw[k] print(f" {k:5d} {r['mean']:12.5f} {r['std']:9.5f} {r['gap']:16.5f}") print(f" 参考:单样本 MC ELBO 的均值 = {elbo:.5f}(= ELBO 解析值,不随 k 变)") print("\n结论:k 增大时 L_k 单调逼近 log p(x),但永远不越过;" "而「取平均再 log」和「log 再取平均」是两件事。") if __name__ == "__main__": main() reparam_gradients.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """重参数化 vs 得分函数(REINFORCE)梯度估计量,方差实测。 要估的量:∇_φ E_{q_φ}[log p(x|z)],φ = (m, ℓ),ℓ = log s。 这一项没法解析求的时候只能采样,两种采样方式的方差差好几个数量级。 真实梯度在这个高斯线性模型里能解析算出来,所以「谁对谁错」有标准答案, 剩下的差别就纯粹是方差。 两部分实验: A. 固定一个小模型(D=6),看三种估计量的分布 B. 扫数据维度 D,看方差比怎么长——这里才是重参数化真正救命的地方 只依赖 numpy,直接 `python reparam_gradients.py` 即可运行。 """ import numpy as np import elbo_identity as EI from elbo_identity import ( D, K, SIGMA2, W, b, elbo_analytic, make_one_x, make_q, ) N_TRIAL = 40000 # ══════════════════════════════════════════════════════════════ # A. 小模型:三种估计量的分布 # ══════════════════════════════════════════════════════════════ def true_grad(x, m, s, Wm, bm, sigma2): """∇_m 与 ∇_ℓ 的解析梯度,ℓ = log s。""" g_m = Wm.T @ (x - (Wm @ m + bm)) / sigma2 g_ell = -(s ** 2) * np.diag(Wm.T @ Wm) / sigma2 return g_m, g_ell def log_lik_rows(x, z, Wm, bm, Dm, sigma2): """批量算 log p(x|z),z: [n, K]。""" r = x - (z @ Wm.T + bm) return -0.5 * (Dm * np.log(2 * np.pi * sigma2) + np.einsum("nd,nd->n", r, r) / sigma2) def sample_reparam(x, m, s, Wm, bm, sigma2, rng, n): """重参数化:z = m + s ⊙ ε,梯度经 z 反传。""" eps = rng.normal(size=(n, m.size)) z = m + s * eps score_z = (x - (z @ Wm.T + bm)) @ Wm / sigma2 # ∇_z log p(x|z) return score_z, score_z * s * eps # ∂z/∂m=1, ∂z/∂ℓ=s⊙ε def sample_score(x, m, s, Wm, bm, Dm, sigma2, rng, n, baseline=None): """得分函数:∇ E[f] = E[f · ∇ log q],f = log p(x|z)。""" eps = rng.normal(size=(n, m.size)) z = m + s * eps f = log_lik_rows(x, z, Wm, bm, Dm, sigma2) if baseline is not None: f = f - baseline return f[:, None] * (eps / s), f[:, None] * (eps ** 2 - 1.0) def summarize(name, gm_hat, ge_hat, g_m, g_ell): var = float(np.var(gm_hat, axis=0).sum() + np.var(ge_hat, axis=0).sum()) gnorm = float(np.linalg.norm(np.concatenate([g_m, g_ell]))) err = float(np.linalg.norm(np.concatenate([gm_hat.mean(0) - g_m, ge_hat.mean(0) - g_ell]))) / gnorm n_needed = int(np.ceil((var ** 0.5 / (0.05 * gnorm)) ** 2)) print(f" {name:<24s} 方差和={var:11.2f} 均值相对误差={err:6.2%} " f"达到5%需 n≈{n_needed}") return var def part_a(): print("=" * 68) print("A. 小模型(D=6, σ=0.35):三种估计量的分布") print("=" * 68) # 与 elbo_identity.state() 同一颗种子 => 同一个 x、同一个 q rng_state = np.random.default_rng(EI.STATE_SEED) x, _ = make_one_x(rng_state) m, s = make_q(x, rng_state) g_m, g_ell = true_grad(x, m, s, W, b, SIGMA2) _, recon, kl = elbo_analytic(x, m, s) print(f"\n真实梯度(解析)") print(f" ∇_m = {np.array2string(g_m, precision=4)}") print(f" ∇_ℓ = {np.array2string(g_ell, precision=4)}") print(f" E_q[log p(x|z)] = {recon:.4f},KL(q||p(z)) = {kl:.4f}") rng = np.random.default_rng(11) n = N_TRIAL print(f"\n估计量对比({n} 次独立采样,方差 = 6 个参数分量方差之和)") gm_r, ge_r = sample_reparam(x, m, s, W, b, SIGMA2, rng, n) gm_s, ge_s = sample_score(x, m, s, W, b, D, SIGMA2, rng, n) gm_sb, ge_sb = sample_score(x, m, s, W, b, D, SIGMA2, rng, n, baseline=recon) v_r = summarize("重参数化", gm_r, ge_r, g_m, g_ell) v_s = summarize("得分函数(无 baseline)", gm_s, ge_s, g_m, g_ell) v_sb = summarize("得分函数(预言机 baseline)", gm_sb, ge_sb, g_m, g_ell) print(f"\n 方差比:得分函数 / 重参数化 = {v_s / v_r:.1f}×") print(f" 加 baseline 之后 = {v_sb / v_r:.1f}×") print(f"\n∇_m 第 0 个分量(真值 {g_m[0]:.4f})的估计分布") for name, g in (("重参数化", gm_r[:, 0]), ("得分函数", gm_s[:, 0]), ("得分函数+baseline", gm_sb[:, 0])): print(f" {name:<20s} mean={g.mean():9.4f} std={g.std():8.3f} " f"min={g.min():9.2f} max={g.max():9.2f}") print(f"\n均值误差随样本数的收敛(∇_m 第 0 个分量)") print(f" {'n':>7s} {'重参数化':>12s} {'得分函数':>12s}") for n2 in (10, 100, 1000, 10000, N_TRIAL): print(f" {n2:7d} {abs(gm_r[:n2, 0].mean() - g_m[0]):12.5f} " f"{abs(gm_s[:n2, 0].mean() - g_m[0]):12.5f}") return dict(g_m0=g_m[0], gm_r=gm_r[:, 0], gm_s=gm_s[:, 0], gm_sb=gm_sb[:, 0], var_r=v_r, var_s=v_s, var_sb=v_sb) # ══════════════════════════════════════════════════════════════ # B. 方差比随数据维度 D 怎么长 # ══════════════════════════════════════════════════════════════ def build_model(Dm, Km, seed): """构造一个 decoder 列范数归一化的线性高斯模型,让梯度尺度不随 D 漂。""" rng = np.random.default_rng(seed) Wm = rng.normal(size=(Dm, Km)) Wm /= np.linalg.norm(Wm, axis=0, keepdims=True) # 每列范数 = 1 bm = rng.normal(0.0, 0.3, size=(Dm,)) return Wm, bm def part_b(): print("\n" + "=" * 68) print("B. 方差比随数据维度 D 的变化(σ²=1,decoder 列范数归一化)") print("=" * 68) rng = np.random.default_rng(23) rows = [] print(f"\n {'D':>6s} {'重参数化':>12s} {'得分函数':>14s} {'比值':>10s} " f"{'+baseline 方差':>14s} {'比值':>10s}") for Dm in (8, 64, 512, 2048): Wm, bm = build_model(Dm, K, seed=1000 + Dm) z_true = rng.normal(size=(K,)) x = Wm @ z_true + bm + rng.normal(size=(Dm,)) # σ² = 1 m = Wm.T @ (x - bm) # 一个合理的变分均值 s = np.full((K,), 0.6) sigma2 = 1.0 n = 4000 gm_r, ge_r = sample_reparam(x, m, s, Wm, bm, sigma2, rng, n) gm_s, ge_s = sample_score(x, m, s, Wm, bm, Dm, sigma2, rng, n) # 预言机 baseline:E_q[log p(x|z)] 的解析值 resid = x - (Wm @ m + bm) recon = -0.5 * (Dm * np.log(2 * np.pi * sigma2) + (resid @ resid + np.trace(Wm.T @ Wm @ np.diag(s ** 2))) / sigma2) gm_sb, ge_sb = sample_score(x, m, s, Wm, bm, Dm, sigma2, rng, n, baseline=recon) v_r = float(np.var(gm_r, axis=0).sum() + np.var(ge_r, axis=0).sum()) v_s = float(np.var(gm_s, axis=0).sum() + np.var(ge_s, axis=0).sum()) v_sb = float(np.var(gm_sb, axis=0).sum() + np.var(ge_sb, axis=0).sum()) rows.append((Dm, v_r, v_s, v_sb)) print(f" {Dm:6d} {v_r:12.2f} {v_s:14.2f} {v_s / v_r:9.1f}× " f"{v_sb:14.2f} {v_sb / v_r:9.1f}×") print("\n 为什么长这么快:f = log p(x|z) 的均值里有一大坨与参数无关的") print(" 「底噪」(-D/2·log 2πσ² 占了主要部分),它进到 f·∇log q 里按平方") print(" 放大;∇_z f 对它求导恒为 0,所以重参数化天然免疫。baseline 减掉的") print(" 也正是这一坨——减完比值只剩 5 倍左右,说明差距主要来自常数项,") print(" 不是来自「采样本身」。") return rows if __name__ == "__main__": a = part_a() rows = part_b() beta_kl_weight.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """β-VAE 的 KL 加权:一个能解到最优的玩具模型,看潜变量维度怎么死掉。 构造:数据每个坐标独立,x_j ~ N(0, λ_j),潜变量每个坐标也独立, q(z_j | x_j) = N(a_j x_j, s_j²),p(x_j | z_j) = N(b_j z_j, σ²) 这样每个维度就是一份独立副本,最优解可以逐个维度单独求,不用训练神经网络。 单维目标(要最大化),把 λ 与 σ² 都当作常数: L(a, b, ℓ) = -[λ(1 - b a)² + b² s²] / (2σ²) - β · KL KL = 0.5 · (a² λ + s² - 1 - log s²), s = exp(ℓ) 梯度全部解析,脚本里再用有限差分核对一遍,避免推错。 只依赖 numpy,直接 `python beta_kl_weight.py` 即可运行。 """ import numpy as np SIGMA2 = 0.25 # 解码器观测噪声方差 def kl_of(a, lam, s): """KL(N(a x, s²) || N(0, 1)),对 x ~ N(0, λ) 取期望后的形式。""" return 0.5 * (a ** 2 * lam + s ** 2 - 1.0 - np.log(s ** 2)) def recon_mse(a, b, lam, s): """E_{x,z}[(x - b z)²],重建误差的期望。""" return lam * (1.0 - b * a) ** 2 + b ** 2 * s ** 2 def objective(p, lam, beta, free_bits=None): """目标函数值(越大越好)。free_bits=C 时 KL 项取 max(KL, C)。""" a, b, ell = p s = np.exp(ell) kl = kl_of(a, lam, s) kl_eff = max(kl, free_bits) if free_bits is not None else kl return -(recon_mse(a, b, lam, s)) / (2 * SIGMA2) - beta * kl_eff def grad(p, lam, beta, free_bits=None): """解析梯度 ∇(∂L/∂a, ∂L/∂b, ∂L/∂ℓ)。""" a, b, ell = p s2 = np.exp(2 * ell) g = np.zeros(3) g[0] = lam * b * (1.0 - b * a) / SIGMA2 g[1] = (lam * a * (1.0 - b * a) - b * s2) / SIGMA2 g[2] = -b ** 2 * s2 / SIGMA2 if free_bits is None or kl_of(a, lam, np.exp(ell)) > free_bits: # KL 项对 (a, b, ℓ) 的梯度 g[0] -= beta * a * lam g[2] -= beta * (s2 - 1.0) return g def grad_fd(p, lam, beta, free_bits=None, h=1e-6): """中心差分梯度,用来核对解析梯度有没有推错。""" g = np.zeros(3) for i in range(3): e = np.zeros(3) e[i] = h g[i] = (objective(p + e, lam, beta, free_bits) - objective(p - e, lam, beta, free_bits)) / (2 * h) return g def optimize(lam, beta, free_bits=None, n_init=9, n_iter=12000, seed=0): """多起点梯度上升,返回最好的 (a, b, s)。 必须多起点:(a, b) = (0, 0) 是「维度死亡」解,单起点容易掉进去出不来。 """ rng = np.random.default_rng(seed) best_p, best_v = None, -np.inf for k in range(n_init): if k == 0: p = np.array([0.9, 0.9, np.log(0.6)]) elif k == 1: p = np.array([0.0, 0.0, 0.0]) # 死亡解,也让它试试 else: p = np.array([rng.uniform(-1.5, 1.5), rng.uniform(-1.5, 1.5), rng.uniform(-1.2, 0.5)]) lr = 0.02 for t in range(n_iter): g = grad(p, lam, beta, free_bits) g = np.clip(g, -50.0, 50.0) # 梯度裁剪:b²s²/σ² 那一项能把参数炸飞 p = p + lr * g # 数值保护,别让 exp(ℓ) 或 (a, b) 爆掉 p[0] = float(np.clip(p[0], -8.0, 8.0)) p[1] = float(np.clip(p[1], -8.0, 8.0)) p[2] = float(np.clip(p[2], -4.0, 1.0)) if t % 1000 == 999: lr *= 0.6 # 收尾再磨一遍:β 很小的时候收敛慢,不磨的话 s² 会差到 1e-3 lr = 1e-4 for t in range(6000): g = np.clip(grad(p, lam, beta, free_bits), -50.0, 50.0) p = p + lr * g v = objective(p, lam, beta, free_bits) if v > best_v: best_v, best_p = v, p a, b, ell = best_p return float(a), float(b), float(np.exp(ell)), float(best_v) def active_unit(a, lam, thr=0.01): """Higgins 的 active unit 判据:Cov_x(E_q[z]) = a²λ 是否超过阈值。""" return a ** 2 * lam > thr # 文章与配图共用的扫描范围 LAMS = [4.0, 1.0, 0.25, 0.0625] BETAS = [0.25, 0.5, 1.0, 2.0, 3.0, 4.0, 6.0, 8.0, 12.0] def sweep(lams, betas, verbose=True): """扫 (λ, β) 网格,返回 {(λ, β): (a, b, s, MSE, KL, a²λ)}。""" table = {} if verbose: print("\n[1] β 扫描:每个 (λ, β) 的最优解") print(f"\n {'λ':>7s} {'β':>5s} {'a':>8s} {'b':>8s} {'s':>8s} " f"{'重建MSE':>9s} {'KL':>8s} {'a²λ':>8s} {'存活':>5s}") for lam in lams: for beta in betas: a, b, s, _ = optimize(lam, beta, seed=17) kl = kl_of(a, lam, s) mse = recon_mse(a, b, lam, s) au = a ** 2 * lam table[(lam, beta)] = (a, b, s, mse, kl, au) if verbose: print(f" {lam:7.4f} {beta:5.2f} {a:8.4f} {b:8.4f} {s:8.4f} " f"{mse:9.4f} {kl:8.4f} {au:8.4f} " f"{'是' if active_unit(a, lam) else '死':>5s}") return table def main(): print("=" * 72) print("β-VAE 的 KL 加权:潜变量维度怎么死掉") print(f"σ² = {SIGMA2}(解码器噪声),单维独立副本,多起点梯度上升求最优") print("=" * 72) # ── 先核对解析梯度 ── print("\n[0] 解析梯度 vs 有限差分") for lam, beta in ((4.0, 1.0), (0.25, 4.0)): p = np.array([0.7, -0.4, np.log(0.8)]) ga, gf = grad(p, lam, beta), grad_fd(p, lam, beta) print(f" λ={lam:<5} β={beta:<4} 解析={np.array2string(ga, precision=5)} " f"差分={np.array2string(gf, precision=5)} 最大偏差={np.abs(ga - gf).max():.2e}") table = sweep(LAMS, BETAS) print("\n[2] 每个 λ 的「死亡阈值」:β 到多大时这个维度不再被用") for lam in LAMS: dead = [beta for beta in BETAS if not active_unit(table[(lam, beta)][0], lam)] if dead: print(f" λ={lam:<7.4f} 从 β={min(dead)} 起死亡" f"(β=1 时 KL={table[(lam, 1.0)][4]:.4f})") else: print(f" λ={lam:<7.4f} 在 β≤{max(BETAS)} 内都存活") # ── 闭式解核对:存活时 s² = βσ²/λ,a²λ = 1 - s²,重建 MSE = βσ² ── print("\n[2b] 闭式解核对(存活的格子才成立)") print(f" {'λ':>7s} {'β':>5s} {'s²实测':>9s} {'βσ²/λ':>9s} " f"{'a²λ实测':>9s} {'1-s²':>9s} {'MSE实测':>9s} {'βσ²':>9s}") worst = 0.0 for lam in LAMS: for beta in BETAS: a, b, s, _, kl, au = table[(lam, beta)] if not active_unit(a, lam): continue s2_pred = beta * SIGMA2 / lam rows = (s ** 2, s2_pred, au, 1 - s ** 2, recon_mse(a, b, lam, s), beta * SIGMA2) print(f" {lam:7.4f} {beta:5.2f} " + " ".join(f"{v:9.5f}" for v in rows)) worst = max(worst, abs(s ** 2 - s2_pred), abs(au - (1 - s ** 2)), abs(recon_mse(a, b, lam, s) - beta * SIGMA2)) print(f" 最大偏差 = {worst:.2e} → 存活时最优解确实落在闭式上") print(f" 推论:维度存活条件 s² < 1 ⟺ β < λ/σ²,即该维的信噪比要盖过 KL 的价") for lam in LAMS: print(f" λ={lam:<7.4f} 理论死亡阈值 β* = λ/σ² = {lam / SIGMA2:.2f}") # ── free bits:把 KL 压在下界,维度就不会死 ── print("\n[3] free bits 对照(β=4,KL 项取 max(KL, C))") print(f" {'λ':>7s} {'C=0(纯 β)':>22s} {'C=0.05':>22s} {'C=0.2':>22s}") for lam in LAMS: cells = [] for C in (None, 0.05, 0.2): a, b, s, _ = optimize(lam, 4.0, free_bits=C, seed=17) cells.append(f"a={a:5.3f} KL={kl_of(a, lam, s):5.3f} " f"MSE={recon_mse(a, b, lam, s):5.3f}") print(f" {lam:7.4f} " + " ".join(f"{c:>22s}" for c in cells)) print("\n[4] 一个直觉:β 大了以后,模型宁可不要这个维度") print(" λ=0.0625 时这个坐标的信息量本来就小于解码器噪声 σ²=0.25,") print(" 用它换来的重建收益抵不上 KL 的代价,最优解就是把它整条关掉。") if __name__ == "__main__": main() make_figures.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """生成「变分下界与重参数化」的三张解释图。 数值全部来自同目录的三个脚本(elbo_identity / reparam_gradients / beta_kl_weight),这里只负责把它们画出来——改了那几个脚本这里要重跑, 避免图与正文数字不一致。 三张图分别回答: 1. ELBO 离真实目标差多少,差的那一项是什么,多采样能不能补上 2. 重参数化到底省了多少方差,以及这个差距随数据维度怎么长 3. KL 项加权的 β 是怎么把潜变量维度一刀一刀切掉的 只依赖 numpy + matplotlib,直接 `python make_figures.py` 即可运行。 """ from pathlib import Path import matplotlib.pyplot as plt import numpy as np from matplotlib.patches import Rectangle import beta_kl_weight as BK import elbo_identity as EI import reparam_gradients as RG ROOT = Path(__file__).resolve().parents[1] OUT = ROOT / "figures" OUT.mkdir(exist_ok=True) plt.rcParams.update({ "font.sans-serif": ["PingFang SC", "Arial Unicode MS", "DejaVu Sans"], "axes.unicode_minus": False, "figure.dpi": 160, }) INK = "#1f2937" C_ELBO = "#2f6fb0" # ELBO:蓝 C_GAP = "#e0a03c" # 缺口:橙 C_BAD = "#d1495b" # 越界 / 死亡:红 C_OK = "#2f9e6f" # 收紧 / 存活:绿 C_SCORE = "#8b5cf6" # 得分函数:紫 def style(ax, title, xlabel=None, ylabel=None): ax.set_title(title, fontsize=13.0, weight="bold", color=INK, pad=10) if xlabel: ax.set_xlabel(xlabel, fontsize=11, color="#475569") if ylabel: ax.set_ylabel(ylabel, fontsize=11, color="#475569") ax.tick_params(labelsize=10, colors="#64748b") for s in ("top", "right"): ax.spines[s].set_visible(False) for s in ("left", "bottom"): ax.spines[s].set_color("#cbd5e1") ax.grid(axis="y", color="#eef2f7", lw=1.0) ax.set_axisbelow(True) # ══════════════════════════════════════════════════════════════ # 图 1:ELBO 离 log p(x) 差多少 # ══════════════════════════════════════════════════════════════ def fig_elbo_gap(): st = EI.state() mc = EI.mc_table(st) iw = EI.iwae_table(st) log_px, elbo, gap, kl_prior = st["log_px"], st["elbo"], st["gap"], st["kl"] fig, axes = plt.subplots(1, 3, figsize=(16.2, 4.9)) fig.patch.set_facecolor("white") # ── (a) 分解:两条横向长条 ── ax = axes[0] style(ax, "(a) log p(x) 拆成两截", "nats(对数似然,0 在右边)") rows = [ ("ELBO(能算,优化它)", elbo, C_ELBO, "white"), ("log p(x)(真想要的,算不出来)", log_px, C_BAD, "white"), ("KL(q‖p(z)):loss 里那一项", -kl_prior, "#94a3b8", "white"), ("KL(q‖p(z|x)):上面两条的差", -gap, C_GAP, "white"), ] for i, (name, v, color, tc) in enumerate(rows): y = len(rows) - 1 - i ax.barh(y, -v, left=v, height=0.52, color=color) ax.text(v + 0.18, y, f"{abs(v):.3f}", ha="left", va="center", fontsize=10, color=INK, weight="bold") ax.axvline(log_px, color=C_BAD, ls="--", lw=1.3) ax.axvline(elbo, color=C_ELBO, ls="--", lw=1.3) ax.annotate("", xy=(elbo, 3.28), xytext=(log_px, 3.28), arrowprops=dict(arrowstyle="<->", color="#a16207", lw=1.8)) ax.text((elbo + log_px) / 2, 3.75, "差 = 1.985", ha="center", va="bottom", fontsize=10.5, color="#a16207", weight="bold", bbox=dict(fc="white", ec="none", alpha=0.9, pad=1.0)) ax.text(-9.35, 0.68, "上面两条之差 = 下面橙色那条,\n不是灰色那条——这两个 KL 常被混为一谈", ha="left", va="center", fontsize=9, color="#64748b", linespacing=1.5) ax.set_yticks(range(len(rows))) ax.set_yticklabels([r[0] for r in rows][::-1], fontsize=9.5) ax.set_xlim(-9.6, 1.5) ax.set_ylim(-0.6, 3.9) ax.grid(axis="x", color="#eef2f7", lw=1.0) ax.grid(axis="y", visible=False) # ── (b) 单样本 MC 会越过上界 ── ax = axes[1] style(ax, f"(b) L=1 的 MC 估计({EI.MC_REP} 次)", "单次估计值", "频数") vals = mc[1]["vals"] ax.hist(vals, bins=70, color="#cbd5e1", edgecolor="white", lw=0.4) over = vals[vals > log_px] ax.hist(over, bins=70, color=C_BAD, alpha=0.85, label=f"{mc[1]['frac']:.1%} 越过了 log p(x)") ax.axvline(elbo, color=C_ELBO, lw=2.0, label=f"ELBO = {elbo:.3f}") ax.axvline(log_px, color=C_BAD, lw=2.0, ls="--", label=f"log p(x) = {log_px:.3f}") ax.set_xlim(-26, 7) ax.legend(fontsize=9, frameon=False, loc="upper left") ax.set_yscale("log") ax.set_ylim(0.7, 900) ax.text(0.03, 0.05, "估计量是无偏的:它在 ELBO 周围晃,\n不是「恒小于 log p(x)」", transform=ax.transAxes, ha="left", va="bottom", fontsize=9.5, color="#64748b", linespacing=1.5) # ── (c) IWAE 收紧 ── ax = axes[2] style(ax, f"(c) 多采样收紧(每项 {EI.IW_REP} 次)", "k(每次采几个 z)", "nats") ks = np.array(EI.IW_KS, dtype=float) means = np.array([iw[k]["mean"] for k in EI.IW_KS]) stds = np.array([iw[k]["std"] for k in EI.IW_KS]) ax.errorbar(ks, means, yerr=stds, fmt="o-", color=C_OK, lw=2.0, capsize=4, markersize=6) ax.axhline(log_px, color=C_BAD, ls="--", lw=1.6, label=f"log p(x) = {log_px:.3f}") ax.axhline(elbo, color=C_ELBO, ls=":", lw=1.6, label=f"ELBO = {elbo:.3f}") ax.set_xscale("log") ax.set_ylim(-9.4, -4.7) ax.set_xlim(0.7, 900) for k, m in zip(ks, means): ax.annotate(f"还差 {log_px - m:.3f}", (k, m), textcoords="offset points", xytext=(8, 12 if k > 1 else -18), fontsize=9, color="#64748b", bbox=dict(fc="white", ec="none", alpha=0.85, pad=0.6)) ax.legend(fontsize=9, frameon=False, loc="lower right") ax.text(0.03, 0.06, "k 越大越紧,但永远不越过虚线", transform=ax.transAxes, fontsize=9.5, color="#64748b") fig.suptitle("这张图要看什么:ELBO 与真实目标之间那道缝," "一半来自 q 的表达力,一半来自「只采一个样本」", fontsize=11.5, color="#64748b", y=1.02) fig.tight_layout() p = OUT / "elbo_gap.png" fig.savefig(p, bbox_inches="tight", facecolor="white") plt.close(fig) print(f" ✓ {p.name}") # ══════════════════════════════════════════════════════════════ # 图 2:重参数化省了多少方差 # ══════════════════════════════════════════════════════════════ def fig_grad_variance(): a = RG.part_a() rows = RG.part_b() fig, axes = plt.subplots(1, 2, figsize=(13.4, 5.0)) fig.patch.set_facecolor("white") # ── (a) 估计分布 ── ax = axes[0] style(ax, "(a) ∂L/∂m 第 0 个分量的估计分布", "估计值", "频数") bins = np.linspace(-70, 70, 90) for vals, color, name in ((a["gm_r"], C_ELBO, "重参数化"), (a["gm_s"], C_SCORE, "得分函数"), (a["gm_sb"], C_GAP, "得分函数+baseline")): ax.hist(vals, bins=bins, histtype="step", lw=1.8, color=color, label=f"{name} std={vals.std():.1f}") ax.axvline(a["g_m0"], color=C_OK, lw=2.2, label=f"真值 = {a['g_m0']:.3f}") ax.set_xlim(-70, 70) ax.set_yscale("log") ax.legend(fontsize=9.5, frameon=False, loc="upper left") ax.text(0.98, 0.05, "三条曲线均值都对,差的是腰围:\n" f"std = {a['gm_r'].std():.1f} / {a['gm_s'].std():.1f} / " f"{a['gm_sb'].std():.1f}", transform=ax.transAxes, ha="right", va="bottom", fontsize=9.5, color="#64748b", linespacing=1.5) # ── (b) 随维度怎么长 ── ax = axes[1] style(ax, "(b) 方差随数据维度 D 怎么长", "D(数据维度)", "梯度方差(6 个分量之和)") Ds = np.array([r[0] for r in rows], dtype=float) vr = np.array([r[1] for r in rows]) vs = np.array([r[2] for r in rows]) vsb = np.array([r[3] for r in rows]) ax.loglog(Ds, vs, "o-", color=C_SCORE, lw=2.0, markersize=7, label="得分函数(无 baseline)") ax.loglog(Ds, vsb, "s--", color=C_GAP, lw=2.0, markersize=6, label="得分函数 + baseline") ax.loglog(Ds, vr, "o-", color=C_ELBO, lw=2.0, markersize=7, label="重参数化") # 斜率 2 的参考线 ref = vs[0] * (Ds / Ds[0]) ** 2 ax.loglog(Ds, ref, ":", color="#cbd5e1", lw=2.0, label="∝ D² 参考线") for d, v, r in zip(Ds, vs, vr): ratio = v / r txt = f"{ratio:,.0f}×" if ratio < 1e4 else f"{ratio:.1e}×" ax.annotate(txt, (d, v), textcoords="offset points", xytext=(8, -14), fontsize=9.5, color=C_SCORE, weight="bold", bbox=dict(fc="white", ec="none", alpha=0.85, pad=0.6)) ax.set_ylim(2e-1, 1e10) ax.legend(fontsize=9.5, frameon=False, loc="upper left") ax.text(0.97, 0.03, "重参数化几乎与 D 无关", transform=ax.transAxes, fontsize=9.5, color=C_ELBO, ha="right", va="bottom", weight="bold") fig.suptitle("这张图要看什么:重参数化不是「更准」,是「腰围更小」——" "同样一次采样,梯度离真值近一个数量级", fontsize=11.5, color="#64748b", y=1.02) fig.tight_layout() p = OUT / "grad_variance.png" fig.savefig(p, bbox_inches="tight", facecolor="white") plt.close(fig) print(f" ✓ {p.name}") # ══════════════════════════════════════════════════════════════ # 图 3:β 怎么把潜变量维度切掉 # ══════════════════════════════════════════════════════════════ def fig_beta_tradeoff(): table = BK.sweep(BK.LAMS, BK.BETAS, verbose=False) s2 = BK.SIGMA2 fig, axes = plt.subplots(1, 2, figsize=(13.4, 5.0)) fig.patch.set_facecolor("white") # ── (a) 重建误差 vs β ── ax = axes[0] style(ax, "(a) β 越大,重建越差——直到某一维被直接放弃", "β(KL 项的权重)", "该维的重建 MSE") colors = ["#2f6fb0", "#2f9e6f", "#e0a03c", "#d1495b"] bgrid = np.linspace(0.2, 14, 100) ax.plot(bgrid, bgrid * s2, ":", color="#94a3b8", lw=1.8, label="理论:MSE = βσ²(还活着)") for lam, c in zip(BK.LAMS, colors): betas, mses, alive = [], [], [] for beta in BK.BETAS: a, b, s, mse, kl, au = table[(lam, beta)] betas.append(beta) mses.append(mse) alive.append(BK.active_unit(a, lam)) betas = np.array(betas) mses = np.array(mses) alive = np.array(alive) ax.plot(betas[alive], mses[alive], "o-", color=c, lw=2.0, markersize=6, label=f"λ={lam}") if (~alive).any(): ax.plot(betas[~alive], mses[~alive], "s", color=c, markersize=7, mfc="white", mew=1.8) ax.axhline(mses[~alive][0], color=c, ls="--", lw=1.0, alpha=0.45) ax.set_xscale("log") ax.set_yscale("log") ax.set_ylim(0.028, 7.0) ax.legend(fontsize=9, frameon=False, loc="upper left", ncol=2) ax.text(0.5, 0.012, "实心圆 = 该维还活着(MSE 贴着 βσ² 往上走);" "空心方块 = 该维已死,MSE 停在 λ", transform=ax.transAxes, ha="center", va="bottom", fontsize=9.5, color="#64748b") # ── (b) 相图 ── ax = axes[1] style(ax, "(b) 存活 / 死亡的相图", "λ / σ²(该维的信噪比)", "β") snr = np.array(BK.LAMS) / s2 grid = np.logspace(-1, 1.6, 60) ax.loglog(grid, grid, "-", color="#94a3b8", lw=2.2, label="理论分界:β = λ/σ²") ax.fill_between(grid, grid, 1e3, color=C_BAD, alpha=0.06) ax.fill_between(grid, 1e-3, grid, color=C_OK, alpha=0.06) for lam in BK.LAMS: for beta in BK.BETAS: a, b, s, mse, kl, au = table[(lam, beta)] ok = BK.active_unit(a, lam) ax.plot(lam / s2, beta, "o" if ok else "s", color=C_OK if ok else C_BAD, markersize=7, mfc=C_OK if ok else "white", mew=1.6) ax.text(0.05, 0.93, "上方:KL 太贵,维度被关掉", transform=ax.transAxes, fontsize=10, color=C_BAD, weight="bold") ax.text(0.05, 0.06, "下方:信息量盖过 KL 的价,维度存活", transform=ax.transAxes, fontsize=10, color=C_OK, weight="bold") ax.set_xlim(0.1, 40) ax.set_ylim(0.15, 20) ax.legend(fontsize=9.5, frameon=False, loc="center right") ax.grid(axis="y", color="#eef2f7", lw=1.0) fig.suptitle("这张图要看什么:β 不是在「调重建和 KL 的比例」," "是在给每个潜变量维度标一个价——信噪比不够的维度直接归零", fontsize=11.5, color="#64748b", y=1.02) fig.tight_layout() p = OUT / "beta_tradeoff.png" fig.savefig(p, bbox_inches="tight", facecolor="white") plt.close(fig) print(f" ✓ {p.name}") if __name__ == "__main__": print("生成配图 →", OUT) fig_elbo_gap() fig_grad_variance() fig_beta_tradeoff() print("完成:3 张")
2026年09月25日
2 阅读
0 评论
0 点赞
2026-09-25
AIGC 基本功|FlashAttention 为什么不需要存下注意力矩阵-FlashAttn
FlashAttention 为什么不需要存下注意力矩阵 所属方向:注意力与位置编码 | 难度:进阶 | 前置知识:自注意力机制的计算与显存账本(本篇是它的直接后继) 关键词:FlashAttention、online softmax、tiling、IO感知、显存优化 01. 为什么需要它 上一篇《自注意力机制的计算与显存账本》结尾留了一个没解的结:N=4096 的视频 DiT,按保守教学分配模型计,单层注意力为 2.39 GiB 激活,无重算训练时按同一假设累加 32 层就是 76.5 GiB,尚未包含权重、梯度、优化器与 FFN。推理的临时激活则不能直接乘层数。当时摆出了两条出路:稀疏化(算得少,但要看清赔的是什么质量)和 FlashAttention。这一篇把后者讲透。 先把一个流传很广的说法钉死:FlashAttention 不是近似注意力。它算出来的就是标准的 softmax 注意力,在实数算术下等价,浮点下允许舍入差异(第 04 节有实测:float64 下最大误差 6.1e-16,纯浮点舍入)。它快的理由也不神秘——同目录 io_ledger.py 算过一笔账,单头 d=64、N=4096 时: 朴素实现要在 HBM(显存)上搬 130 MiB:分数矩阵 S 写一次读一次,softmax 权重 P 写一次读一次,4N² 次元素搬运; FlashAttention 只搬 17.8 MiB,少了 7.3 倍。 两种实现的主导矩阵乘 FLOPs 相同,但重算和归一化开销不同。在这组 A100 参数的 roofline 模型中,朴素实现算术强度为 31.5 FLOP/byte,低于约 200 的平衡点;这支持优先优化 IO 的方向,不证明所有序列、硬件或近似注意力方法都必然更慢。本文 130 MiB、17.8 MiB 是脚本估算值,不是 GPU 访存计数器实测。 另一个结论更值钱:显式保存概率矩阵 P 的朴素实现仍需二次存储,S 本身通常不必一并保留,显存永远是二次的;FlashAttention 把这两个矩阵整个从显存里删掉了,训练激活从 O(N²) 降到 O(N)——在统一采用六份线性张量、忽略小的 LSE 与工作区的教学预算中,上面那个 76.5 GiB 降为 4.5 GiB。N=32768 的长视频任务,本文教学模型的朴素实现要约 4.54 TiB 激活,物理上不存在能装下的卡;FlashAttention 只要 36 GiB。整个长序列时代(32K、128K 上下文)就是踩在这个技巧上站起来的。 02. 最小可用理解 三句话讲完核心思想: softmax attention 每行需要两个标量和一个 d 维向量:这行的最大值 m、指数和 l、加权和 O。把分数按块流进来,每块用一条递推式把这三个量修正一次,全流完之后 O÷l 就是精确的 softmax 注意力输出——中间任何一步都不需要把整行摆在内存里。这叫 online softmax。 显式物化 N×N 经常造成较低的算术强度;受限类型取决于维度、硬件与实现。朴素实现把 N×N 的 S、P 写回显存再读回来;FlashAttention 用 tiling(分块)把它们关在片上 SRAM 里算完就扔,HBM 读写量从 Θ(N²) 降到 Θ(N²d²/M),M 是片上 SRAM 的大小。 反向传播要用的 S、P 全部重算。除 Q/K/V(或它们的重算来源)外,前向存输出 O 和每行的 logsumexp(都是 O(N)),反向时拿 Q、K 重新过一遍分块流程,把需要的局部 P 重新算出来——用重复计算换显存,这是整个方法里最「敢」的一步。 这张图要看什么:左边朴素实现的 S、P 两个 N×N 是该前向示意中的主要中间量,它们必须写回 HBM 再读回来;右边是同一个 N×N 被切成 B_r×B_c 的小块,K/V 块进 SRAM 常驻、Q 块逐行流过,跨块只有 O_i、l_i、m_i 三个 O(N) 的量一直活着。 03. 数学推导 3.1 出发点:safe softmax 为什么需要先看全一行 设一行分数为 $s_1, \dots, s_N$,softmax 的定义是 $$p_i = \frac{\exp(s_i)}{\sum_{j=1}^{N} \exp(s_j)}$$ 分子分母都是 exp 的和。直接算会溢出:s 只要有 89 左右,$\exp(s)$ 在 fp32 就到 inf 了。工程上全部改用 safe softmax——先求这行的最大值 m,再算平移后的指数: $$p_i = \frac{\exp(s_i - m)}{\sum_{j=1}^{N} \exp(s_j - m)}, \qquad m = \max_{1 \le j \le N} s_j$$ 原始分子分母同乘 $\exp(-m)$,结果不变,但指数的输入全部落在 $(-\infty, 0]$,永不溢出。问题就出在这个 m 上:m 是对整行取的 max。你必须先把 N 个分数全部看过一遍才知道 m 是多少,然后才能开始算 exp——这解释了常见实现先保存分数再做 softmax 的流程;算法并不强制保存整个 N×N,也可重算分数或逐行处理,只是 IO 和效率不同。 而注意力输出对这一行还要再多两个量:分母 $l = \sum_j \exp(s_j - m)$,以及加权和 $O = \sum_j \exp(s_j - m)\, v_j$($v_j$ 是第 j 个 token 的 Value 向量)。最终输出就是 $O / l$。 所以真正要回答的问题是:如果分数是一块一块到来的(事先不知道后面块里有什么),这三个量还能算吗? 3.2 online softmax 递推式(全文核心) 能。做法是把「以 m 为参考系」改成「以当前的 m 为参考系,m 变了就整体换算」。 设已经流过了前 t 块,维护三个量:参考系最大值 $m^{(t)}$、分母 $l^{(t)}$、未归一化加权和 $O^{(t)}$,它们满足不变式 $$O^{(t)} = \sum_{j \le t} \exp(s_j - m^{(t)})\, v_j, \qquad l^{(t)} = \sum_{j \le t} \exp(s_j - m^{(t)})$$ (这里 $j \le t$ 是「属于前 t 块的所有下标」的缩写。)现在第 $t+1$ 块到了,块内最大值是 $m_{\text{blk}}$。新的全局最大值是 $$m^{(t+1)} = \max(m^{(t)},\ m_{\text{blk}})$$ 关键一步来了:旧累积量是按 $m^{(t)}$ 为参考系记的,而新的不变式要求参考系换成 $m^{(t+1)}$。把不变式里的 $\exp(s_j - m^{(t)})$ 拆成 $\exp(s_j - m^{(t+1)}) \cdot \exp(m^{(t+1)} - m^{(t)})$,旧量的换算系数就是 $\exp(m^{(t)} - m^{(t+1)})$: $$l^{(t+1)} = \exp(m^{(t)} - m^{(t+1)})\, l^{(t)} + \sum_{j \in \text{blk}} \exp(s_j - m^{(t+1)})$$ $$O^{(t+1)} = \exp(m^{(t)} - m^{(t+1)})\, O^{(t)} + \sum_{j \in \text{blk}} \exp(s_j - m^{(t+1)})\, v_j$$ 每一步只是把指数拆成两项相乘再重新合并,等价性是代入即可验证的恒等式;所有块流完后 $O^{(T)}/l^{(T)}$ 与朴素 softmax 在实数算术下相同(差在浮点舍入,第 04 节实测 1e-16 量级)。这个「换参考系」的系数在论文和代码里叫 rescale,跨块最大值、求和与 rescale 都会带来额外的标量操作;它们不改变主导矩阵乘次数。 严谨一点可以正向验证不变式:假设第 t 步的不变式成立,那么 $$O^{(t+1)} = e^{m^{(t)} - m^{(t+1)}} \sum_{j \le t} e^{s_j - m^{(t)}} v_j + \sum_{j \in \text{blk}} e^{s_j - m^{(t+1)}} v_j = \sum_{j \le t+1} e^{s_j - m^{(t+1)}} v_j$$ (第一个等号就是递推式,第二个等号把 $e^{m^{(t)} - m^{(t+1)}}$ 乘进求和号里、指数相加后正好变回 $e^{s_j - m^{(t+1)}}$。)旧块和新块在同一个参考系下合并,不变式保持。$l$ 的证明一字不差,把 $v_j$ 去掉就行。归纳基础是初始状态 $m^{(0)} = -\infty$、$l^{(0)} = 0$、$O^{(0)} = 0$:第一块到来时换算系数按 0 处理(对应代码里 np.where(np.isneginf(m_old), 0.0, ...) 那一行),三个量直接等于第一块的局部值。 三个细节值得停一下: 这个递推对任意分块都成立,块大小可以是 1(一个 key 一个 key 地流),也可以是 128 行。块大小只影响效率,不影响结果。 $\exp(m^{(t)} - m^{(t+1)}) \le 1$ 恒成立(因为 $m$ 单调不减),rescale 永远是在把旧量缩小,数值上很安全。 $m^{(t)} + \log l^{(t)}$ 就是 logsumexp(LSE)。它是反向传播唯一需要额外记录的东西——记住这个,第 3.3 节要用。 顺带说一句因果掩码:掩码就是把被遮位置的 $s$ 设成 $-\infty$,exp 之后是 0,对 m、l、O 都没有贡献;更进一步,如果一整块都被遮住(Q 块整体在 K 块之前),这块连算都不用算,直接跳过。第 04 节实测这个「整块跳过」在 N=4096、块 64 时省掉 49.2% 的块。 3.3 反向传播:重算换显存 这里的 O 指最终归一化输出。令 $G=\partial L/\partial O$、$A=GV^\top=\partial L/\partial P$。softmax 沿每行归一化,其正确反向公式是: $$\frac{\partial L}{\partial S}=P\odot\left(A-\operatorname{rowsum}(A\odot P)\right)$$ 行和为 $N\times1$,沿 key 轴广播;它是上游梯度在概率权重下的行平均,不能写成 $1-P^\top\mathbf 1$。随后 $dQ=dS\,K/\sqrt d$、$dK=dS^\top Q/\sqrt d$、$dV=P^\top G$。朴素实现可保存 P 而无需同时保存 S;FlashAttention 则重算局部 P。 FlashAttention 的做法是:除 Q/K/V 外,前向额外存 O 和 LSE($N \times d$ 加 $N$ 个数,O(N));反向时把分块流程原样再走一遍,在每一块里用 $\exp(S_{ij} - \text{LSE})$ 把局部的那一小块 P 重新算出来,立刻用于梯度,算完就扔。整个反向里 P 从头到尾没有以 N×N 的形态存在过。 用重复计算换显存——这笔交易的换算率是:每层每个头多算一遍 QK^T 和一次 exp(具体比例取决于前后向统计口径与实现),对本文保守教学账本,这对应去掉 S、P、浮点 dropout 乘子三项;真实朴素反向通常不用同时保存 S,具体减少几份取决于实现。N=4096、本文简化模型上,单层激活 2.39 GiB → 0.141 GiB,17 倍。 3.4 IO 账:搬运量到底差多少 设单头维度 d,片上存储预算为 M 个元素(不是字节);fp16 下 192 KiB 对应 M=98304。SRAM 里要同时放下 K 块、V 块(各 $B_c \times d$)和 Q 块、输出块(各 $B_r \times d$),论文取 $$B_c = \frac{M}{4d}, \qquad B_r = \min(B_c,\ d)$$ A100 每个 SM 有 192 KB SRAM,fp16 下 d=64 时 $B_c = 384$、$B_r = 64$——这是论文 IO 模型的粗略分块预算;真实 kernel 还受 score tile、累加器、寄存器、共享内存配额与 occupancy 约束,192 KiB 也不是每个 block 可独占的共享内存。 朴素实现的搬运量(单位:元素个数):QK^T 读 Q、K 各 Nd、写 S 一次 N²;softmax 读 S 写 P 各 N²;PV 读 P 一次 N²、读 V 写 O 各 Nd。合计 $4Nd + 4N^2$,主导项 4N²。 FlashAttention 的搬运量:K、V 各进 SRAM 一次($2Nd$);Q 块每换一个 K/V 块就要重读一遍,共 $T_c \cdot Nd$($T_c = \lceil N/B_c\rceil$ 是 K/V 块数);输出 O 同理要读出写回各一遍($2 T_c Nd$);l、m 两个 O(N) 的运行量共 $4T_c N$。合计约 $2Nd + 3T_cNd + 4T_cN$,代进去: $$\text{HBM 搬运量} \;\approx\; 3 \cdot \frac{N}{B_c} \cdot Nd \;=\; \Theta\!\left(\frac{N^2 d^2}{M}\right)$$ 两个量级一比:朴素是 $\Theta(N^2)$,Flash 是 $\Theta(N^2 d^2/M)$,比值 $d^2/M$ 在 d=64、fp16、M=98304 个元素(192 KiB) 时约等于 0.04——大 O 比例省略了常数,不能直接当作 25 倍的实际收益;脚本计入 Q/O 反复读写后的模型比值为 7.3 倍。这就是整篇论文的全部:不是新数学,是把「数据在哪」当成一等公民来优化。 按本文简化模型,朴素注意力的算术强度在 N 远大于 d 时约为 d/2 FLOP/byte(fp16),所以固定 d=64 时接近 32;它会随 d 改变,并非与头维度无关。A100 示例的约 200 FLOP/byte 是 dense FP16 Tensor Core 峰值与 HBM 带宽之比,真实 softmax 的非矩阵乘指令、缓存与调度仍会影响性能。 这张图要看什么:固定 d 与 SRAM 大小时,左图两者的大 N 主导项均为 N²,比值渐近趋于常数;线性项和取整会让有限 N 的比值变化;右图斜率不同(N² 对 N),N 越大两条线离得越远,朴素教学激活曲线越过示例 80 GiB 预算线的位置,FlashAttention 还有几百倍余量。 04. 代码实现 完整脚本在文末附录(flash_online_softmax.py、io_ledger.py、make_figures.py),只依赖 numpy,全部用 /usr/local/bin/python3 实跑过,下面的数字都是真实输出。 4.1 递推式长什么样:逐块打印 先看一个能盯着看的例子(N=8、D=4、K/V 块大小 B_c=2,Q 整行一起处理): def attention_flash(Q, K, V, B_r, B_c): N, D = Q.shape scale = 1.0 / np.sqrt(D) O = np.zeros((N, D)) # 未归一化的输出累加器 l = np.zeros(N) # 归一化分母(exp 之和) m = np.full(N, -np.inf) # 到目前为止见过的最大值 for j0 in range(0, N, B_c): Kj = K[j0:j0 + B_c] # [B_c, D] Vj = V[j0:j0 + B_c] # [B_c, D] for i0 in range(0, N, B_r): Qi = Q[i0:i0 + B_r] # [B_r, D] Sij = (Qi @ Kj.T) * scale # 局部分数 [B_r, B_c];Pij 也是同尺寸临时量 m_blk = Sij.max(axis=-1) # [B_r] m_old = m[i0:i0 + B_r] m_new = np.maximum(m_old, m_blk) Pij = np.exp(Sij - m_new[:, None]) # [B_r, B_c] l_blk = Pij.sum(axis=-1) corr = np.where(np.isneginf(m_old), 0.0, np.exp(m_old - m_new)) l[i0:i0 + B_r] = l[i0:i0 + B_r] * corr + l_blk O[i0:i0 + B_r] = O[i0:i0 + B_r] * corr[:, None] + Pij @ Vj m[i0:i0 + B_r] = m_new return O / l[:, None], (m + np.log(l)) out, lse = attention_flash(Q, K, V, B_r=64, B_c=64) print(out.shape, lse.shape) # (N, D) (N,) —— 输出和 LSE 都只有 O(N) 和第 3.2 节逐符号对上:m_new 是 $m^{(t+1)}$,corr 是换参考系的 $\exp(m^{(t)} - m^{(t+1)})$,l_blk 是块内指数和,Pij @ Vj 是块内加权和。corr 里的 np.where(np.isneginf(m_old), 0.0, ...) 处理的是第一块之前 $m = -\infty$ 的情况($-\infty - (-\infty)$ 会出 nan,直接规定换算系数为 0,旧累积量本来就是 0)。最后一行把 $m + \log l$ 作为 LSE 返回,留给反向。 实跑的逐块轨迹(第 0 号 query): j= 0 m=+0.2968 l=1.7748 corr=0.0000 j= 2 m=+0.6071 l=3.0317 corr=0.7333 j= 4 m=+0.6071 l=3.9732 corr=1.0000 j= 6 m=+0.6071 l=5.3007 corr=1.0000 看两点:j=2 时新块里出现了更大的分数,m 被抬高、旧累积量被打了个 0.733 的折扣;j=4 之后 m 没再变,corr 恒为 1,rescale 白做——此时乘子在数学上为 1;具体 GPU kernel 是否跳过这些操作要核对实现,不能仅由轨迹推断。最大的单个临时矩阵只有 16 个元素([8, 2]),而 N×N 是 64;这里不等于同时存活临时数组的总元素数。 4.2 等价性:和朴素实现在实数算术下等价,浮点下允许舍入差异 attention_naive 是上一篇的标准实现(S、P 都落地),两者对拍: N D causal 最大绝对误差 相对误差 256 64 False 6.106e-16 1.311e-15 256 64 True 8.882e-16 3.520e-16 1024 64 False 6.106e-16 2.629e-15 1024 64 True 6.661e-16 2.332e-16 2048 64 False 9.437e-16 3.724e-15 2048 64 True 8.327e-16 3.570e-16 4096 64 False 7.702e-16 5.679e-15 4096 64 True 7.772e-16 2.426e-16 误差全是 1e-16 量级——float64 的舍入级别。这就是「精确注意力」四个字的实测含义:不是「误差很小」,是算法本身和朴素 softmax 完全等价。 4.3 峰值内存:O(N²) 对 O(N),实测 用 tracemalloc 量函数内新分配的峰值内存(float64、D=64、块 64×64;Q/K/V 已提前分配,不计入本表): N naive 实测 naive 理论 2N²·8 flash 实测 比值 1024 16.5 MiB 16.0 MiB 1.1 MiB 14.4x 2048 65.0 MiB 64.0 MiB 2.2 MiB 30.1x 4096 258.0 MiB 256.0 MiB 4.2 MiB 61.5x 8192 1028.0 MiB 1024.0 MiB 8.3 MiB 123.6x N 从 4096 翻到 8192:naive 峰值 ×4.0,flash 峰值 ×2.0 naive 的实测和理论列($2N^2 \times 8$ 字节,S 和 P 两个 N×N)对得上,说明量的方法可信。最后那行是阶数的直接证据:N 翻倍,naive 峰值 ×4(二次),flash 峰值 ×2(线性)。 4.4 时间:在 CPU 上它反而更慢——这一点必须诚实 N naive flash(B=64) flash/naive 1024 7.2 ms 10.2 ms 1.4x 2048 26.3 ms 42.4 ms 1.6x 4096 98.7 ms 166.3 ms 1.7x CPU + numpy 上 flash 慢约 1.6 倍。三个原因:乘加次数一样还多了 rescale 和逐块 exp;一次大矩阵乘被拆成 (N/64)² 个 64×64 小矩阵乘,BLAS 跑不满小块,Python 循环开销也进来了;CPU 也有 SRAM 缓存,但这里的 Python/NumPy 分块没有实现专门的缓存与线程优化,不能把它当作 GPU 内核速度的预测。第三张图展示 IO 模型为何支持在 GPU 上尝试这一优化。 4.5 因果掩码:整块跳过,白捡一半 朴素实现加因果掩码,N×N 还是得整块算完再往被遮的位置上写 $-\infty$,一个 FLOP 都省不下来。FlashAttention 按「Q 块整体在 K 块之前就整块跳过」处理,N=4096、块 64 时实测: 块总数(非因果) : 4096 块总数(因果跳过): 2080 跳过的块占比 : 49.2% 保留的是含对角线的下三角,块数是 $T(T+1)/2$,占比 $(T+1)/2T$,T 大时趋近一半。训练 GPT 类因果模型、以及视频 DiT 里的时序因果注意力,这半是免费的。 05. 工业级实现对照 最小实现讲清了原理,但生产 kernel 和它有四处本质差异,每处都值得知道为什么: 第一,块大小不是从公式算的,是 autotune 出来的。 第 3.4 节的 $B_c = M/4d$ 是IO 分析采用的可行块预算;真实的 flash-attention kernel(flash_attn/flash_attn_interface.py 的 flash_attn_func,以 2025-09 的实现为准)里,块大小是按(头维度、数据类型、是否因果、显存架构)在若干组预编译配置里选的,还受 warp 数量、寄存器压力、shared memory bank conflict 的影响——公式只负责告诉你「必须小于某个数」,调优负责在约束内找最快的。 第二,减少非矩阵乘操作。 FlashAttention-2 使用未归一化输出累计,减少 rescale、除法等非矩阵乘工作。corr=1 时数学上无需改变旧值,但不能笼统声称所有 kernel 都按每行最大值是否变化来分支跳过。 第三,前向的结构是「外层 Q、内层 K/V」。 我们按论文 v1 的写法外层遍历 K/V 块;FlashAttention-2 把循环反过来(外层 Q 块),好处是输出 O 常驻寄存器不用反复读写、且不同 Q 块之间天然并行,能吃满更多 SM。论文 v1 的伪代码适合理解递推,v2 的循环结构才是现在 kernel 的样子。 第四,dropout 不存掩码,存随机数种子。 朴素实现要为反向留一个 B×H×N×N 的 dropout 掩码;kernel 里只存 Philox 计数器的 seed 和 offset(几十字节),反向时用同一个种子重新生成同样的掩码。这是「重算换显存」哲学最极致的一次应用——连随机数本身都可以重算。 另外两条工程事实:PyTorch 2.0 起 F.scaled_dot_product_attention 会自动按(头维度、掩码、数据类型、硬件)在 flash / memory-efficient / math 三个后端里挑,你不写一行 CUDA 也在用它;论文报告的端到端收益是 BERT-large(seq 512)比 MLPerf 1.1 训练记录快 15%、GPT-2(seq 1K)快 3 倍、Long Range Arena(seq 1K-4K)快 2.4 倍——注意 seq 512 时只有 15%,因为那时注意力在整层里占比还小,收益随序列长度涨,这正是 IO 复杂度模型的预测。 排查问题时你会想知道「此刻到底在用哪个后端」。PyTorch 留了一个官方口子: from torch.nn.attention import sdpa_kernel, SDPBackend with sdpa_kernel(SDPBackend.FLASH_ATTENTION): out = F.scaled_dot_product_attention(q, k, v, is_causal=True) # 强制走 flash;如果这个头维度/掩码组合它不支持,这里会直接报错, # 而不是悄悄退回 math 后端——「悄悄降级」正是性能莫名掉一半时最该先查的事 这张图要看什么:横轴是 N,纵轴是「每搬 1 字节做多少次运算」,灰色虚线是机器平衡点(201 FLOP/byte)。本模型中朴素实现的强度渐近接近 32;FlashAttention 的估计从 N=2048 起超过平衡点。这是按矩阵乘峰值做的模型分类,真实瓶颈还受非矩阵乘指令、缓存、并行度和调度影响。 06. 代价与边界 FlashAttention 省下了 HBM 搬运和 N² 显存,赔进去的和没管住的也要说清楚。 代价:重计算和额外归一化操作。 反向要重算局部分数及概率,其中包含矩阵乘和逐元素运算,不能统称为固定 30% 的额外非矩阵乘开销。小序列的收益取决于 kernel、调度和硬件,没有统一的 seq<512 亏损阈值。 数值边界:实数算术等价不保证浮点逐位一致。 kernel 内部用 fp16/bf16 存储、fp32 累加,块内的归一化和朴素实现的一次性归一化在浮点上不同。对训练的影响需要结合 dtype、输入尺度与误差测试判断,但如果你在做数值敏感的分析(比如逐 token 概率对比),要知道它和参考实现差在舍入级别,不是 bug。 边界的核心一条:它没有改变复杂度,改变的常数。 显存从 $O(N^2)$ 降到 $O(N)$,但算力还是 $\Theta(N^2 d)$、HBM 搬运还是 $\Theta(N^2 d^2/M)$。N=32768 时 FlashAttention 的单头搬运是 1061 MiB——比朴素实现的 8.2 GiB 好得多,但随 N 继续平方增长这一点没变。上下文再往上涨(1M token),接力棒要交给稀疏注意力、线性注意力、状态空间模型这些真正改复杂度的方法。FlashAttention 的块结构恰恰是它们的底座:把注意力切成块之后,「整块跳过」才成为可能,第 4.5 节那个 49.2% 推广到任意稀疏模式就是块稀疏注意力。 不该用的场景:需要拿到完整注意力权重做分析或可视化的(P 从头到尾没存在过,想看它就得回到朴素实现);自定义的任意注意力偏置如果 kernel 不支持,绕过去的方法可能把优势吃掉;以及缺乏适配内核的环境;CPU 也有 SRAM 缓存,但本文 NumPy 循环没有实现专门的 CPU cache 优化(第 4.4 节的 CPU 实测就是例子)。 07. 经典论文脉络 这条线的演进关系一句话各说清: Milakov & Gimelshein, 2018(arXiv:1805.02867)Online normalizer calculation for softmax——首次提出 softmax 的 online 计算:流式更新 max 和指数和。当时的目标只是省一次对 logits 的遍历,还没人把它和注意力显存联系起来。 Rabe & Staats, 2021(arXiv:2112.05682)Self-attention Does Not Need O(n²) Memory——讨论单 query 的常数额外空间方案,并给出分块的低内存自注意力实现及速度/内存实验;不能把单 query 的额外空间结论写成整段输出总存储 O(1),也不能概括为没有实用价值。 Dao et al., 2022(arXiv:2205.14135)FlashAttention——本篇锚点。补上缺失的一环:把流式更新从「逐 token」改成「逐块」(tiling),配上 IO 复杂度分析和 GPU kernel,第一次让「精确 + 更快 + 更省显存」三者同时成立。 Dao, 2023(arXiv:2307.08691)FlashAttention-2——循环重排(外层 Q)、削减非矩阵乘 FLOPs、更好的并行度,把 v1 大约 25-40% 的峰值算力利用率推向 50-73%。 一条清晰的线:2018 年有技巧,2021 年有证明,2022 年才有产品——缺的从来不是数学,是「意识到瓶颈在 IO」这个视角。 08. 常见误解 以下几条都值得单独记住,前两条我当初也信过: 「FlashAttention 是近似注意力,所以有精度损失」。错。它是精确算法,和朴素 softmax 数学等价(第 4.2 节实测 1e-16)。真正近似的是 Linformer、Performer 那一族。两者经常被并列讨论,但一个在改算法,一个在改数据的搬运方式。 「它快是因为算了更少的 FLOPs」。反了,它的 FLOPs 略多于朴素实现(第 4.4 节 CPU 实测慢 1.6 倍)。主要收益来自减少 IO,也来自更好的并行划分与调度——这正是它给所有做系统优化的人的启示:先问数据在哪,再问算了多少。 「显存优化只对前向有用」。最大的收益在反向:除 Q/K/V 外额外存 O 和 LSE(固定 d 时为 O(N)),反向重算 P。本文三份二次张量是教学分配假设,真实反向不一定保存三份;可靠结论是分块重算避免显式保留完整 P。 「块开得越大越快」。块大小受 SRAM 硬约束:超过每个 block 的资源上限可能无法启动或编译;寄存器 spill 等情况也会引入额外访存,整个方法的根基(中间结果不落 HBM)就塌了。真实 kernel 的块大小是约束内的调优问题,不是越大越好。 「有了它就不用再关心 N²」。它优化的是常数,不是阶数——算力和搬运仍随 N 平方涨。能继续走的长上下文路线是稀疏化(改算力阶数)和状态空间/线性注意力(改注意力本身),FlashAttention 的块结构是它们的载体,不是替代品。 09. 动手验证 跑文末附录的 flash_online_softmax.py(只依赖 numpy),预期输出与正文一致: 第 2 节等价性表:非因果最大误差 6.1e-16 到 9.4e-16,因果 6.7e-16 到 8.9e-16——全是舍入级别; 第 3 节内存表:naive 峰值和理论值 $2N^2 \times 8$ 字节吻合,N 翻倍时 naive ×4.0、flash ×2.0。 再做一个一行代码的实验,直接看清递推式里 rescale 的分量:把 corr = np.where(np.isneginf(m_old), 0.0, np.exp(m_old - m_new)) 改成 corr = np.ones_like(m_old)(假装 m 永远不变,也就是退回「先见全家再算」之前的朴素流式假设),重跑等价性检验。我实测过:N=1024、D=64、随机高斯输入下,最大误差从 6.1e-16 恶化到 0.276,平均误差 0.0136——输出在量级上就是错的。这个对比说明:流式计算 softmax 时,「用新参考系换算旧累积量」这一步不是工程细节,是正确性本身。 最后改 io_ledger.py 开头的硬件常数(比如把 SRAM_PER_SM 调到 48 KB 模拟消费级卡),重跑看块大小和搬运量比值怎么变——你会看到 SRAM 越小,FlashAttention 相对朴素实现的搬运量优势越小,$d^2/M$ 里的 M 直接控制这一切。 10. 延伸阅读 按知识树的依赖关系,建议按这个顺序继续走: 前置:自注意力机制的计算与显存账本——本篇所有显存账本的出处(2.39 GiB、94% 那几笔账都在那篇里);旋转位置编码 RoPE 与 视频 DiT 里的 3D RoPE——注意力的另外两个必备零件,本篇刻意没有碰位置编码。 后继:视频生成里的稀疏注意力(规划中)——FlashAttention 的块结构是稀疏模式的执行底座,「整块跳过」从因果掩码推广到任意稀疏图;算子融合与 CUDA Graph(规划中)——把「少搬内存」推到极端就是融合,FlashAttention 是这个思路最成功的案例;KV Cache 与自回归视频生成——推理时的另一本显存账。 延伸到知识树之外:想读 kernel 源码,从 Dao-AILab/flash-attention 的 flash_attn/flash_attn_interface.py 进,先读 forward 再读 backward;想读原始推导,Milakov & Gimelshein(1805.02867)给出了独立 softmax 在线归一化的直接推导。 附录:完整代码 09 节用到的脚本全文如下(io_ledger.py、flash_online_softmax.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 io_ledger.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """IO 账本:朴素注意力和 FlashAttention 在 HBM 上到底搬了多少字节。 这篇讲的是「为什么快」。答案不在 FLOPs 上——两者的乘加次数几乎一样—— 而在于 GPU 有两级内存: HBM(显存) 带宽约 1.5 TB/s,容量 40~80 GB SRAM(片上) 带宽约 19 TB/s,但每个 SM 只有 192 KB 朴素实现把 N×N 的分数矩阵写回 HBM 再读出来,等于把数据在这条 13 倍带宽差的 通道上来回搬;FlashAttention 用 tiling 让这些中间结果根本不落 HBM。 下面的模型把每一笔读写都点清楚,参数是可调的,读者可以改硬件常数重算。 硬件数量级取自 FlashAttention 论文 Table 1(A100 40GB)。 只依赖 numpy。直接 `python io_ledger.py` 即可运行。 """ import numpy as np # ── 硬件常数(A100 40GB 量级,论文 Table 1)──────────────────── HBM_BW = 1.555e12 # HBM 带宽,字节/秒 SRAM_BW = 19.0e12 # 片上 SRAM 带宽,字节/秒 SRAM_PER_SM = 192 * 1024 # 每个 SM 的片上 SRAM,字节 FLOPS_PEAK = 312e12 # fp16 tensor core 峰值,FLOP/秒 GIB = 1024 ** 3 MIB = 1024 ** 2 # ──────────────────────────────────────────────────────────── # 块大小:SRAM 里能同时放下什么 # ──────────────────────────────────────────────────────────── def block_sizes(d, sram_bytes=SRAM_PER_SM, dtype_bytes=2): """返回 (B_r, B_c):Q 块行数与 K/V 块行数。 SRAM 里要同时放下 K_j、V_j 两个 [B_c, d] 和 Q_i、O_i 两个 [B_r, d]。 论文的取法是 B_c = M / (4d)、B_r = min(B_c, d),这里照抄: 先让 K_j+V_j 占掉一半 SRAM,Q 块则不超过 d 行(保证 softmax 按行算得下)。 """ elems = sram_bytes / dtype_bytes B_c = max(1, int(elems // (4 * d))) B_r = min(B_c, d) return B_r, B_c # ──────────────────────────────────────────────────────────── # HBM 读写量(单位:元素个数,乘 dtype_bytes 得字节) # ──────────────────────────────────────────────────────────── def hbm_elems_naive(N, d): """朴素实现:S 和 P 都要落地。 QK^T: 读 Q(Nd) + 读 K(Nd) + 写 S(N²) softmax: 读 S(N²) + 写 P(N²) PV: 读 P(N²) + 读 V(Nd) + 写 O(Nd) """ return 4 * N * d + 4 * N * N def hbm_elems_flash(N, d, B_r, B_c): """FlashAttention:外层遍历 K/V 块,内层遍历 Q 块。 K、V 各读一遍(每个 j 块进 SRAM 后,内层 i 循环里一直复用) Q 每个 j 都要重读一遍:T_c · N·d O 每个 (j,i) 都要读出来再写回去(累加器跨 j 迭代):2 · T_c · N·d l、m 两个 O(N) 的运行量同理:4 · T_c · N """ T_c = int(np.ceil(N / B_c)) return 2 * N * d + 3 * T_c * N * d + 4 * T_c * N def flops_attention(N, d): """两个 N×N 矩阵乘,一次 [M,K]×[K,N] 算 2MKN 个浮点运算。""" return 4 * N * N * d def activation_bytes(N, d, H, B=1, dtype_bytes=2): """单层注意力的训练激活(反向要用的中间张量)。 保守教学模型:6 个 [B,N,d_model] 线性项 + S/P/浮点乘子 3 个 [B,H,N,N]。 此函数 d 是总宽度 d_model,不是其他 IO 函数中的头宽度。 融合侧保留同样六份线性预算,忽略小的 LSE 与工作区;不是框架峰值测量。 """ lin = 6 * B * N * d * dtype_bytes quad = 3 * B * H * N * N * dtype_bytes return lin + quad, lin def roofline(bytes_moved, flops): """算术强度(FLOP/byte)与两个上界时间。返回 (强度, 内存时间, 算力时间, 瓶颈)。""" intensity = flops / bytes_moved t_mem = bytes_moved / HBM_BW t_comp = flops / FLOPS_PEAK bound = "内存受限" if t_mem > t_comp else "算力受限" return intensity, t_mem, t_comp, bound # ──────────────────────────────────────────────────────────── # 报表 # ──────────────────────────────────────────────────────────── def report_io(): d, b = 64, 2 # 单头维度 64,fp16 B_r, B_c = block_sizes(d) print("=" * 74) print("1. SRAM 块大小与 HBM 读写量(单头,d=64,fp16)") print("=" * 74) print(f" SRAM {SRAM_PER_SM / 1024:.0f} KB / SM,fp16 下能放 " f"{SRAM_PER_SM / b:.0f} 个元素") print(f" → B_c = M/(4d) = {B_c},B_r = min(B_c, d) = {B_r}\n") print(f"{'N':>7} {'naive HBM':>12} {'flash HBM':>12} {'比值':>8} " f"{'naive 强度':>11} {'flash 强度':>11}") rows = [] for N in (1024, 2048, 4096, 8192, 16384, 32768): nb = hbm_elems_naive(N, d) * b fb = hbm_elems_flash(N, d, B_r, B_c) * b fl = flops_attention(N, d) i_n, _, _, bound_n = roofline(nb, fl) i_f, _, _, bound_f = roofline(fb, fl) rows.append((N, nb, fb, i_n, i_f, bound_n, bound_f)) print(f"{N:>7} {nb / MIB:>10.1f} MiB {fb / MIB:>10.1f} MiB " f"{nb / fb:>7.1f}x {i_n:>9.1f} {i_f:>9.1f} ") print(f"\n 机器平衡点(峰值算力/带宽)= {FLOPS_PEAK / HBM_BW:.0f} FLOP/byte") print(" 强度低于它 → 内存受限,加算力没用;高于它 → 才开始吃算力。\n") print(f"{'N':>7} {'naive 瓶颈':>12} {'flash 瓶颈':>12}") for N, nb, fb, i_n, i_f, bn, bf in rows: print(f"{N:>7} {bn:>12} {bf:>12}") print("\n 注意:flash 的强度在 N 大时越过平衡点,模型说它变成算力受限了。") print(" 但真实 kernel 达不到峰值——softmax 的 exp 走的是特殊函数单元,") print(" 不走 tensor core,这个「非矩阵乘开销」正是 FlashAttention-2 之后") print(" 继续优化的地方。模型给出的是上界,不是承诺。") def report_activation(): d, H, B, b = 3072, 24, 1, 2 print("\n" + "=" * 74) print("2. 教学激活存储模型(简化 Transformer:d=3072, H=24, B=1, 32 层, bf16)") print("=" * 74) print(f"{'N':>7} {'朴素/层':>12} {'Flash/层':>12} {'比值':>8} " f"{'朴素 32 层':>12} {'Flash 32 层':>13}") for N in (1024, 4096, 8192, 32768): naive, flash = activation_bytes(N, d, H, B, b) print(f"{N:>7} {naive / GIB:>10.2f} GiB {flash / GIB:>10.3f} GiB " f"{naive / flash:>7.0f}x {32 * naive / GIB:>10.1f} GiB " f"{32 * flash / GIB:>11.2f} GiB") n4096, f4096 = activation_bytes(4096, d, H, B, b) print(f"\n N=4096 时,朴素实现单层 {n4096 / GIB:.2f} GiB,其中二次项占 " f"{100 * (1 - f4096 / n4096):.1f}%——") print(" 这正是 attention_basics 那篇里「三个 N×N 吃掉 94%」的那一格。") print(" FlashAttention 删掉的就是这一格,剩下的是 O(N) 的线性项。") def report_speed_limit(): d, b = 64, 2 B_r, B_c = block_sizes(d) print("\n" + "=" * 74) print("3. 理论上界:如果只受 HBM 带宽限制,两者各要多久(单头,d=64)") print("=" * 74) print(f"{'N':>7} {'FLOPs':>10} {'naive 内存时间':>16} {'flash 内存时间':>16} " f"{'纯带宽模型比':>9}") for N in (1024, 4096, 16384): nb = hbm_elems_naive(N, d) * b fb = hbm_elems_flash(N, d, B_r, B_c) * b fl = flops_attention(N, d) print(f"{N:>7} {fl / 1e9:>8.2f} G {nb / HBM_BW * 1e3:>14.3f} ms " f"{fb / HBM_BW * 1e3:>14.3f} ms {nb / fb:>8.1f}x") print("\n 此列只比较理想 HBM 时间,并非完整模型训练或实际 kernel 加速上界。") print(" 论文的端到端训练加速与此处单头 IO 模型口径不同,不能直接相比。") if __name__ == "__main__": report_io() report_activation() report_speed_limit() flash_online_softmax.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """online softmax 的正确性与显存代价实测(FlashAttention 的核心那一步)。 对比两种实现: naive —— 先算出完整的 N×N 分数矩阵 S,再逐行 softmax 得到 P,最后算 O = P V flash —— 按块流过 K/V,只维护三个 O(N) 的运行量:输出累加器 O、分母 l、 运行最大值 m(第 03 节递推式的直接实现) 两者在数学上完全等价,唯一差别是中间有没有 N×N 的矩阵落到内存里。 这份脚本回答三个问题: 1. 递推式写出来的结果,和朴素 softmax 逐位一致吗?(1.1 节) 2. 峰值内存真的差一个 N 吗?(用 tracemalloc 量,不是估的) 3. 那算力呢?——在 CPU + numpy 上 flash 是**更慢**的,这一点必须诚实讲清楚 只依赖 numpy。直接 `python flash_online_softmax.py` 即可运行。 """ import time import tracemalloc import numpy as np NEG = -np.inf # ──────────────────────────────────────────────────────────── # 两个被测实现 # ──────────────────────────────────────────────────────────── def attention_naive(Q, K, V, causal=False): """标准实现:S 和 P 都是完整的 N×N 常驻张量。""" D = Q.shape[-1] S = (Q @ K.T) / np.sqrt(D) # [N, N] ← 第一块 N×N if causal: S = np.where(np.triu(np.ones((Q.shape[0], K.shape[0])), 1) > 0, NEG, S) S -= S.max(axis=-1, keepdims=True) # safe softmax,不改变结果 P = np.exp(S) # [N, N] ← 第二块 N×N P /= P.sum(axis=-1, keepdims=True) return P @ V def attention_flash(Q, K, V, B_r, B_c, causal=False, trace=False): """按块流过 + online softmax。全程不出现 N×N 的张量。 B_r / B_c 分别是 Q 块和 K/V 块的行数,对应 SRAM 里各放得下多少行。 """ N, D = Q.shape scale = 1.0 / np.sqrt(D) O = np.zeros((N, D)) # 未归一化的输出累加器 l = np.zeros(N) # 归一化分母(exp 之和) m = np.full(N, NEG) # 到目前为止见过的最大值 peak_tmp = 0 # 记录出现过的最大临时矩阵(元素个数) for j0 in range(0, N, B_c): Kj = K[j0:j0 + B_c] # [B_c, D] Vj = V[j0:j0 + B_c] # [B_c, D] for i0 in range(0, N, B_r): # 因果掩码下,若整个 Q 块都在 K 块之前(所有 query 下标 < 所有 key # 下标),这一块全被遮掉,连算都不用算 —— 朴素实现做不到这一点。 # 条件:块内最大 query 下标 i0+B_r-1 < j0 if causal and i0 + B_r <= j0: continue Qi = Q[i0:i0 + B_r] # [B_r, D] Sij = (Qi @ Kj.T) * scale # [B_r, B_c] 局部分数;Pij 也是临时块 if causal: q_idx = i0 + np.arange(Sij.shape[0])[:, None] k_idx = j0 + np.arange(Sij.shape[1])[None, :] Sij = np.where(k_idx > q_idx, NEG, Sij) peak_tmp = max(peak_tmp, Sij.size) m_blk = Sij.max(axis=-1) # [B_r] m_old = m[i0:i0 + B_r] m_new = np.maximum(m_old, m_blk) # [B_r] Pij = np.exp(Sij - m_new[:, None]) # [B_r, B_c] l_blk = Pij.sum(axis=-1) # [B_r] # 把之前累积的量从旧的参考最大值搬到新的(关键的 rescale 一步) corr = np.where(np.isneginf(m_old), 0.0, np.exp(m_old - m_new)) l[i0:i0 + B_r] = l_old = l[i0:i0 + B_r] * corr + l_blk O[i0:i0 + B_r] = O[i0:i0 + B_r] * corr[:, None] + Pij @ Vj m[i0:i0 + B_r] = m_new if trace: print(f" j={j0:>2} i={i0:>2} m={m_new[0]:+.4f} " f"l={l_old[0]:.4f} corr={corr[0]:.4f}") return O / l[:, None], (m + np.log(l)), peak_tmp def peak_bytes(fn, *args, **kwargs): """跑一次 fn,返回 (结果, 峰值字节数)。用 tracemalloc 实测。""" tracemalloc.start() tracemalloc.reset_peak() out = fn(*args, **kwargs) _, peak = tracemalloc.get_traced_memory() tracemalloc.stop() return out, peak def timed(fn, *args, repeat=3, **kwargs): best = float("inf") out = None for _ in range(repeat): t0 = time.perf_counter() out = fn(*args, **kwargs) best = min(best, time.perf_counter() - t0) return out, best # ──────────────────────────────────────────────────────────── # 1. 递推式长什么样:一个能逐块打印的最小例子 # ──────────────────────────────────────────────────────────── def demo_recurrence(): rng = np.random.default_rng(0) N, D, B_c = 8, 4, 2 Q = rng.normal(size=(N, D)) K = rng.normal(size=(N, D)) V = rng.normal(size=(N, D)) print("=" * 68) print("1. online softmax 的递推过程(N=8, D=4, K/V 块大小 B_c=2)") print("=" * 68) print(" Q 块固定为整行(B_r=N),K/V 分成 4 块依次流过;") print(" 每行打印第 0 号 query 的 m / l / corr,看它们怎么被逐次修正:\n") _, _, peak = attention_flash(Q, K, V, B_r=N, B_c=B_c, trace=True) print(f"\n 最大的单个临时矩阵只有 {peak} 个元素(并非临时内存总和) = [{N}, {B_c}],而 N×N = {N * N}") # ──────────────────────────────────────────────────────────── # 2. 等价性:和朴素 softmax 逐位对得上吗 # ──────────────────────────────────────────────────────────── def demo_exactness(): print("\n" + "=" * 68) print("2. 数值等价性(float64,非因果 / 因果两种掩码)") print("=" * 68) print(f"{'N':>6} {'D':>4} {'causal':>7} {'最大绝对误差':>14} {'相对误差':>12}") for N in (256, 1024, 2048, 4096): D = 64 rng = np.random.default_rng(N) Q = rng.normal(size=(N, D)) K = rng.normal(size=(N, D)) V = rng.normal(size=(N, D)) for causal in (False, True): ref = attention_naive(Q, K, V, causal) got, _, _ = attention_flash(Q, K, V, B_r=64, B_c=64, causal=causal) diff = np.abs(ref - got).max() rel = diff / np.abs(ref).max() print(f"{N:>6} {D:>4} {str(causal):>7} {diff:>14.3e} {rel:>12.3e}") print("\n 误差量级是浮点舍入(1e-15),不是近似——FlashAttention 是精确算法。") # ──────────────────────────────────────────────────────────── # 3. 峰值内存:是不是真的差一个 N # ──────────────────────────────────────────────────────────── def demo_memory(): print("\n" + "=" * 68) print("3. 峰值内存实测(tracemalloc,float64,D=64,B_r=B_c=64)") print("=" * 68) print(f"{'N':>6} {'naive 实测':>12} {'naive 理论 2N²·8':>18} " f"{'flash 实测':>12} {'比值':>8}") peaks = {} for N in (1024, 2048, 4096, 8192): D = 64 rng = np.random.default_rng(N) Q = rng.normal(size=(N, D)) K = rng.normal(size=(N, D)) V = rng.normal(size=(N, D)) _, p_naive = peak_bytes(attention_naive, Q, K, V) _, p_flash = peak_bytes(attention_flash, Q, K, V, 64, 64) peaks[N] = (p_naive, p_flash) print(f"{N:>6} {p_naive / 2**20:>10.1f} MiB " f"{2 * N * N * 8 / 2**20:>16.1f} MiB " f"{p_flash / 2**20:>10.1f} MiB {p_naive / p_flash:>7.1f}x") g_n = peaks[8192][0] / peaks[4096][0] g_f = peaks[8192][1] / peaks[4096][1] print(f"\n N 从 4096 翻到 8192:naive 峰值 ×{g_n:.1f},flash 峰值 ×{g_f:.1f}") print(" 这就是 O(N²) 和 O(N) 的区别:前者翻两倍(4×),后者跟着翻倍(2×)。") print(" 朴素实现要同时留住 S 和 P 两个 N×N(理论列就是 2N²·8 字节,和实测对得上);") print(" flash 只留 B_r×B_c 的块,剩下的是 O(N·D) 的输出累加器。") # ──────────────────────────────────────────────────────────── # 4. 时间:CPU + numpy 上 flash 反而更慢,这才是重点 # ──────────────────────────────────────────────────────────── def demo_time(): print("\n" + "=" * 68) print("4. 墙钟时间(CPU + numpy,BLAS 多线程)——反直觉的一项是这个") print("=" * 68) print(f"{'N':>6} {'naive':>10} {'flash(B=64)':>13} {'flash/naive':>12}") for N in (1024, 2048, 4096): D = 64 rng = np.random.default_rng(N) Q = rng.normal(size=(N, D)) K = rng.normal(size=(N, D)) V = rng.normal(size=(N, D)) _, t_naive = timed(attention_naive, Q, K, V) _, t_flash = timed(attention_flash, Q, K, V, 64, 64) print(f"{N:>6} {t_naive * 1e3:>9.1f} ms {t_flash * 1e3:>12.1f} ms " f"{t_flash / t_naive:>11.1f}x") print("\n 以上比值是本机本次计时,不能据此预测 GPU;可能影响速度的因素包括:") print(" 1. 乘加次数一模一样,还额外多了每块的 rescale 和逐元素 exp;") print(" 2. 一次大矩阵乘被拆成 (N/B)² 次 64×64 的小矩阵乘,") print(" BLAS 在小块上根本跑不满,Python 循环开销也进来了;") print(" 3. 最关键的:它省的是**内存搬运**,不是 FLOPs,") print(" CPU 同样有 SRAM 缓存,但本示例没有专门优化缓存和线程,") print(" 因此分块可能省容量,却被 Python 循环与小 GEMM 开销抵消。") print(" 第 05 节讲 GPU 上为什么结论会反过来。") # ──────────────────────────────────────────────────────────── # 5. 因果掩码:flash 能顺手省掉一半算力,朴素实现不能 # ──────────────────────────────────────────────────────────── def demo_causal_blocks(): print("\n" + "=" * 68) print("5. 因果掩码下实际算了多少块(N=4096, B_r=B_c=64)") print("=" * 68) N, B = 4096, 64 T_r = T_c = N // B total = 0 for j in range(T_c): for i in range(T_r): if i * B + B <= j * B: # 整个 Q 块都在 K 块之前 → 全遮,跳过 continue total += 1 print(f" 块总数(非因果) : {T_r * T_c}") print(f" 块总数(因果跳过): {total}") print(f" 跳过的块占比 : {100.0 * (1 - total / (T_r * T_c)):.1f}%") print("\n 朴素实现就算加了因果掩码,N×N 的矩阵照样得整块算完再遮,") print(" 省不了一点算力;flash 是整块跳过,这是免费的一半(严格说是") print(" (T²+T)/2 / T² ≈ 一半多一点的对角块保留)。") if __name__ == "__main__": np.set_printoptions(precision=4, suppress=True) demo_recurrence() demo_exactness() demo_memory() demo_time() demo_causal_blocks() make_figures.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """生成「FlashAttention」的三张解释图。 数值全部来自同目录下的 io_ledger.py(HBM 读写量、激活显存、算术强度), 改了那个脚本的话这里要跟着重跑,避免图与正文数字不一致。 三张图分别回答: 1. tiling 到底怎么切、什么留在 SRAM、什么留在 HBM 2. 两本账随 N 怎么长(HBM 读写量、教学激活存储模型) 3. 为什么说是「内存受限」变的「不那么受限」(算术强度对机器平衡点) 只依赖 numpy + matplotlib,直接 `python make_figures.py` 即可运行。 """ from pathlib import Path import matplotlib.pyplot as plt import numpy as np from matplotlib.patches import FancyArrowPatch, FancyBboxPatch, Rectangle from io_ledger import ( FLOPS_PEAK, HBM_BW, activation_bytes, block_sizes, flops_attention, hbm_elems_flash, hbm_elems_naive, ) ROOT = Path(__file__).resolve().parents[1] OUT = ROOT / "figures" OUT.mkdir(exist_ok=True) plt.rcParams.update({ "font.sans-serif": ["PingFang SC", "Arial Unicode MS", "DejaVu Sans"], "axes.unicode_minus": False, "figure.dpi": 160, }) C_NAIVE = "#e05263" # 朴素实现:红 C_FLASH = "#2f9e6f" # FlashAttention:绿 C_SRAM = "#f0a03c" # SRAM 高亮:橙 C_GREY = "#94a3b8" C_SKIP = "#e4e9f0" INK = "#182238" MIB = 1024 ** 2 GIB = 1024 ** 3 def _style(ax, title, xlabel=None, ylabel=None): ax.set_title(title, fontsize=13.5, weight="bold", color=INK, pad=10) if xlabel: ax.set_xlabel(xlabel, fontsize=11, color="#475569") if ylabel: ax.set_ylabel(ylabel, fontsize=11, color="#475569") ax.tick_params(colors="#475569", labelsize=10) for s in ("top", "right"): ax.spines[s].set_visible(False) for s in ("left", "bottom"): ax.spines[s].set_color("#cbd5e1") ax.set_axisbelow(True) def _box(ax, x, y, w, h, text, fc, fs=10.5, tc="white"): ax.add_patch(FancyBboxPatch( (x, y), w, h, boxstyle="round,pad=0.02,rounding_size=0.12", linewidth=0, facecolor=fc, edgecolor="none")) ax.text(x + w / 2, y + h / 2, text, ha="center", va="center", fontsize=fs, color=tc, weight="bold", linespacing=1.5) def _arrow(ax, x1, y1, x2, y2, label=None, color="#475569"): ax.add_patch(FancyArrowPatch( (x1, y1), (x2, y2), arrowstyle="-|>", mutation_scale=13, linewidth=1.4, color=color, shrinkA=2, shrinkB=2)) if label: ax.text((x1 + x2) / 2, max(y1, y2) + 0.22, label, ha="center", fontsize=9, color=color) def _caption(fig, text): fig.text(0.5, 0.012, text, ha="center", fontsize=9.5, color="#94a3b8") # ──────────────────────────────────────────────────────────── # 图 1:tiling 怎么切 # ──────────────────────────────────────────────────────────── def fig_tiling(): fig = plt.figure(figsize=(13.6, 6.0), facecolor="#fbfcfe") # ── 左:朴素实现 ── ax = fig.add_axes([0.03, 0.09, 0.46, 0.80]) ax.set_xlim(0, 10); ax.set_ylim(0, 7.4); ax.axis("off") ax.text(0, 7.05, "本例朴素前向:物化 S 与 P", fontsize=13.5, weight="bold", color=C_NAIVE) _box(ax, 0.1, 5.1, 1.5, 1.1, "Q, K, V\n[N, d]", "#64748b", fs=10) _box(ax, 2.4, 4.8, 2.0, 1.7, "S = QK^T / √d\n[N, N]", C_NAIVE, fs=11) _box(ax, 5.3, 4.8, 2.0, 1.7, "P = softmax(S)\n[N, N]", C_NAIVE, fs=11) _box(ax, 8.2, 5.1, 1.6, 1.1, "O = PV\n[N, d]", "#64748b", fs=10) _arrow(ax, 1.65, 5.65, 2.35, 5.65) _arrow(ax, 4.45, 5.65, 5.25, 5.65, "读 S 写 P") _arrow(ax, 7.35, 5.65, 8.15, 5.65, "读 P") ax.text(4.4, 4.25, "N=4096, d=64, fp16:S 和 P 各 32 MiB", ha="center", fontsize=9.5, color="#475569") ax.add_patch(Rectangle((0.1, 2.2), 9.7, 1.35, facecolor="#eef2f7", edgecolor="#cbd5e1", linewidth=1.2)) ax.text(4.95, 3.2, "HBM 带宽 1.5 TB/s", ha="center", fontsize=10.5, color="#475569", weight="bold") ax.text(4.95, 2.5, "S 写一次读一次,P 写一次读一次 → 4N² 次元素搬运", ha="center", fontsize=9.5, color="#64748b") _arrow(ax, 3.4, 4.75, 3.4, 3.6, color=C_NAIVE) _arrow(ax, 6.3, 4.75, 6.3, 3.6, color=C_NAIVE) ax.text(0.1, 1.35, "代价:数据在这条通道上来回两趟,", fontsize=10.5, color=C_NAIVE, weight="bold") ax.text(0.1, 0.75, "而 HBM 比 SRAM 慢 12 倍", fontsize=10.5, color=C_NAIVE, weight="bold") # ── 右:FlashAttention ── ax2 = fig.add_axes([0.53, 0.09, 0.44, 0.80]) ax2.set_xlim(0, 10); ax2.set_ylim(0, 7.4); ax2.axis("off") ax2.text(0, 7.05, "FlashAttention:按块流过,只留 O(N)", fontsize=13.5, weight="bold", color=C_FLASH) T = 8 gx0, gy0, cell = 1.05, 3.1, 0.46 for i in range(T): for j in range(T): if i + 1 <= j: # 因果掩码下整块跳过(先判,优先级最高) fc, ec = C_SKIP, "#c3ccd8" elif j == 3: # 当前正在处理的 K/V 块列 fc, ec = "#cdeadb", C_SRAM elif i == 5: # 当前 Q 块行 fc, ec = "#dfeaf6", "#94a3b8" else: fc, ec = "#f7fafc", "#dde3ea" ax2.add_patch(Rectangle( (gx0 + j * cell, gy0 + (T - 1 - i) * cell), cell * 0.90, cell * 0.90, facecolor=fc, edgecolor=ec, linewidth=1.1)) gtop = gy0 + T * cell ax2.text(gx0 - 0.42, gy0 + T * cell / 2, "Q 块逐行流过", fontsize=9.5, color="#475569", ha="center", va="center", rotation=90) ax2.add_patch(FancyBboxPatch( (gx0 + 3 * cell - 0.06, gy0 - 0.06), cell * 1.02, T * cell + 0.12, boxstyle="round,pad=0.03,rounding_size=0.1", linewidth=1.6, edgecolor=C_SRAM, facecolor="none", linestyle="--")) ax2.text(gx0 + 3.5 * cell, gy0 - 0.42, "K_j, V_j 常驻 SRAM", ha="center", fontsize=9.5, color=C_SRAM, weight="bold") ax2.text(gx0 + T * cell + 0.45, gtop - 0.35, "一块 = [B_r, B_c]\n= [64, 64]", fontsize=9.5, color="#475569", va="top", linespacing=1.5) ax2.text(gx0 + T * cell + 0.45, gtop - 1.55, "灰色块:因果掩码下\n整块跳过,连算都不算", fontsize=9, color="#94a3b8", va="top", linespacing=1.5) ax2.text(gx0 + T * cell + 0.45, gtop - 2.95, "片上空间受限:\n驻留 K/V 与 Q/O 块,\n还要容纳分数和统计量", fontsize=9, color=C_SRAM, va="top", linespacing=1.5) by = 1.85 ax2.text(0.15, by + 0.52, "跨 K/V 块一直复用的三个量:", fontsize=10.5, color=INK, weight="bold") for k, (label, color) in enumerate([ ("O_i 输出累加器 [N, d]", C_FLASH), ("l_i 分母 [N]", "#7fb3d5"), ("m_i 运行最大值 [N]", "#c9a227")]): ax2.add_patch(Rectangle((0.15, by - k * 0.46), 4.6, 0.28, facecolor=color, edgecolor="none")) ax2.text(4.9, by - k * 0.46 + 0.14, label, fontsize=9.5, color="#475569", va="center") ax2.text(0.15, 0.12, "每来一块就 rescale 一次(乘 e^{m旧−m新}),最后 O/l 收尾", fontsize=9.5, color="#64748b") _caption(fig, "图 1:切的是「K/V 块 × Q 块」这两层循环,不是把注意力切开。" "左边突出两份二次中间量,右边只维护线性的累加量(固定 d)。") fig.savefig(OUT / "tiling.png", facecolor="#fbfcfe") plt.close(fig) print(" figures/tiling.png") # ──────────────────────────────────────────────────────────── # 图 2:两本账随 N 怎么长 # ──────────────────────────────────────────────────────────── def fig_curves(): d, b = 64, 2 B_r, B_c = block_sizes(d) Ns = np.array([1024, 2048, 4096, 8192, 16384, 32768]) naive_io = np.array([hbm_elems_naive(N, d) * b for N in Ns]) / MIB flash_io = np.array([hbm_elems_flash(N, d, B_r, B_c) * b for N in Ns]) / MIB fig, axes = plt.subplots(1, 2, figsize=(13.2, 5.1), facecolor="#fbfcfe") fig.subplots_adjust(bottom=0.20, top=0.86, wspace=0.25) ax = axes[0] ax.loglog(Ns, naive_io, "o-", color=C_NAIVE, lw=2.2, ms=6, label="朴素实现(∝N²)") ax.loglog(Ns, flash_io, "s-", color=C_FLASH, lw=2.2, ms=6, label="FlashAttention(∝N²/M)") ax.annotate(f"{naive_io[2] / flash_io[2]:.1f}×", xy=(Ns[2], naive_io[2]), xytext=(Ns[2] * 1.6, naive_io[2] * 3.0), fontsize=11, color=C_NAIVE, weight="bold", arrowprops=dict(arrowstyle="->", color=C_NAIVE, lw=1.3)) ax.legend(fontsize=10, frameon=False, loc="upper left") ax.grid(True, which="both", color="#eef2f7", lw=0.9) _style(ax, "HBM 读写量(单头 d=64, fp16)", "序列长度 N", "MiB") ax = axes[1] d2, H, B = 3072, 24, 1 lay = 32 naive_act = np.array( [activation_bytes(N, d2, H, B)[0] * lay for N in Ns]) / GIB flash_act = np.array( [activation_bytes(N, d2, H, B)[1] * lay for N in Ns]) / GIB ax.loglog(Ns, naive_act, "o-", color=C_NAIVE, lw=2.2, ms=6, label="朴素实现(∝N²)") ax.loglog(Ns, flash_act, "s-", color=C_FLASH, lw=2.2, ms=6, label="FlashAttention(∝N)") ax.axhline(80, color="#64748b", ls="--", lw=1.4) ax.text(Ns[0] * 1.15, 92, "示例预算:80 GiB(非硬件规格)", fontsize=9.5, color="#64748b") ax.annotate("76.5 GiB", xy=(4096, naive_act[2]), xytext=(4096 * 1.8, naive_act[2] * 2.2), fontsize=10.5, color=C_NAIVE, weight="bold", arrowprops=dict(arrowstyle="->", color=C_NAIVE, lw=1.3)) ax.annotate("4.5 GiB", xy=(4096, flash_act[2]), xytext=(4096 * 0.55, flash_act[2] * 6.0), fontsize=10.5, color=C_FLASH, weight="bold", arrowprops=dict(arrowstyle="->", color=C_FLASH, lw=1.3)) ax.legend(fontsize=10, frameon=False, loc="upper left") ax.grid(True, which="both", color="#eef2f7", lw=0.9) _style(ax, "教学激活存储模型(简化视频 DiT,32 层)", "序列长度 N", "GiB") _caption(fig, "图 2:两张都是对数轴,斜率就是复杂度阶数。" "左图的大 N 主导项均为 N²,有限 N 时比例受线性项和取整影响;" "右图斜率不同(N² 对 N),N 越大差距越离谱。") fig.savefig(OUT / "io_curve.png", facecolor="#fbfcfe") plt.close(fig) print(" figures/io_curve.png") # ──────────────────────────────────────────────────────────── # 图 3:算术强度 vs 机器平衡点 # ──────────────────────────────────────────────────────────── def fig_intensity(): d, b = 64, 2 B_r, B_c = block_sizes(d) Ns = np.array([1024, 2048, 4096, 8192, 16384, 32768]) naive_i = np.array([flops_attention(N, d) / (hbm_elems_naive(N, d) * b) for N in Ns]) flash_i = np.array([flops_attention(N, d) / (hbm_elems_flash(N, d, B_r, B_c) * b) for N in Ns]) balance = FLOPS_PEAK / HBM_BW fig, ax = plt.subplots(figsize=(11.4, 5.0), facecolor="#fbfcfe") fig.subplots_adjust(bottom=0.22, top=0.86) ax.axhspan(0, balance, color="#fdf1f2", zorder=0) ax.axhline(balance, color="#64748b", ls="--", lw=1.6, zorder=3) ax.text(Ns[-1] * 2.1, balance * 0.70, f"机器平衡点 {balance:.0f} FLOP/byte\n(峰值算力 ÷ HBM 带宽)", fontsize=10, color="#475569", va="center", linespacing=1.5) ax.text(Ns[0] * 1.02, balance * 0.40, "线下:本 roofline 模型由 HBM 项主导(实际瓶颈需测量)", fontsize=10, color=C_NAIVE) ax.text(Ns[0] * 1.02, balance * 1.32, "线上:本模型计算项主导", fontsize=10, color=C_FLASH) ax.semilogx(Ns, naive_i, "o-", color=C_NAIVE, lw=2.4, ms=7, label="朴素实现", zorder=4) ax.semilogx(Ns, flash_i, "s-", color=C_FLASH, lw=2.4, ms=7, label="FlashAttention", zorder=4) for x, v in zip(Ns, naive_i): ax.annotate(f"{v:.0f}", (x, v), textcoords="offset points", xytext=(0, -16), ha="center", fontsize=9, color=C_NAIVE) for x, v in zip(Ns, flash_i): ax.annotate(f"{v:.0f}", (x, v), textcoords="offset points", xytext=(0, 9), ha="center", fontsize=9, color=C_FLASH) ax.set_xticks(Ns) ax.set_xticklabels([f"{n // 1024}K" for n in Ns]) ax.set_xlim(Ns[0] * 0.75, Ns[-1] * 4.5) ax.set_ylim(0, balance * 1.85) ax.legend(fontsize=10.5, frameon=False, loc="center right") ax.grid(axis="y", color="#eef2f7", lw=0.9) _style(ax, "算术强度:每搬 1 字节能做多少次运算(单头 d=64, fp16)", "序列长度 N", "FLOP / byte") _caption(fig, "图 3:朴素实现的强度几乎不随 N 变(一直贴在 32 附近)," "在此模型中由 HBM 项主导;FlashAttention 的估计抬过平衡点," "并不单独证明实际 kernel 已充分利用计算单元。") fig.savefig(OUT / "intensity.png", facecolor="#fbfcfe") plt.close(fig) print(" figures/intensity.png") if __name__ == "__main__": print("生成配图:") fig_tiling() fig_curves() fig_intensity()
2026年09月25日
1 阅读
0 评论
0 点赞
2026-09-25
AIGC 每日速读|2026-09-25|快手开源修图登顶,GEdit仍输GPT
今日 AIGC 论文速览 今日共 10 篇 · 图像编辑与生成安全 2 篇 · 统一预训练与视频运动 2 篇 · 多模态联合生成 2 篇 · 音频生成与后训练 2 篇 · 推理加速与视频评测 2 篇 重点论文标题列表 KwaiMind(快手):180万对数据炼电商修图 UVU(清华大学):视觉监督前移,RefCOCO涨7.7 ⚡ Unite-Audio(东京科学大学):117M音频生成,CLAP登顶 MotionSpec(卡迪夫大学):频谱轨迹约束视频运动 RecCAR(巴伊兰大学):反向注意力补齐,音画更同步 今日论文速览 1. KwaiMind:180万对数据炼电商修图 KwaiMind Technical Report | 快手 | arXiv:2609.26375 关键词:图像编辑,电商,CTR奖励,Ecom-Bench,在线强化学习 ⚠️ 前序问题:电商修图不是通用编辑:商品主体要一模一样、促销文字要渲染准确、图还得有点击欲。通用编辑模型在这三件事上都没专门优化,而电商数据的采集清洗又贵又杂,缺一套能稳定产出高质量配对数据的流水线。 本文贡献:快手发布 KwaiMind:多模态 DiT 底座,Agent 数据引擎维护约 180 万高质量编辑对;继续预训练加 SFT 后走偏好优化与在线 RL,用 VLM 通用裁判加 CTR、文字渲染、商品一致性三类专用奖励训练专精策略,再经 on-policy 蒸馏合并成单一模型;配套发布覆盖 11 类商业编辑任务的 Ecom-Bench。 Overview of the end-to-end data pipeline. Coordinator Agent routes samples through filtering, generation, captioning, and post-filtering with bounded feedback loops and human review. 实验效果:开源编辑器里综合最强:ImgEdit 4.15、GEdit 5.79、REDEdit 英/中 3.75/3.73,Ecom-Bench 视觉质量第一;CTR 引导优化把生成图预测 CTR 超过原图的比例从 12.16% 提到 37.41%,线上 A/B 实际 CTR 相对提升约 2.44%。 Overall comparison on general image editing benchmarks: ImgEdit, GEdit, and the English and Chinese splits of REDEdit. Hatched bars denote closed-source models. 批判点评:最强的限定词是开源:图上闭源模型全面更高——ImgEdit 被 Seedream 4.0 的 4.27 和 GPT-Image-2 的 4.20 压着,GEdit 上 Nano Banana 2 拿 7.02、GPT-Image-2 拿 7.10,比 KwaiMind 的 5.79 高出一大截,REDEdit 两个语种同样如此。真正立住的卖点其实是 CTR 这条商业指标,而不是编辑质量本身。 2. UVU:视觉监督前移,RefCOCO涨7.7 UVU: Improving Multimodal Understanding via Vision-Language Unified Autoregressive Paradigm | 清华大学;腾讯优图实验室;南京大学;格拉斯哥大学 | arXiv:2609.27915 关键词:统一自回归,像素级codebook,连续视觉编码,视觉监督,预训练 ⚠️ 前序问题:多模态大模型的细粒度视觉理解长期靠稀疏文本监督撑着,已有的视觉监督大多加在后训练阶段,那时视觉表征已基本定型,监督信号只能当辅助约束。要重塑感知骨干,得把视觉监督搬进预训练。 本文贡献:提出 UVU 视觉语言统一自回归范式:不用向量量化,SigLIP-2 连续视觉特征直接进 LM 解码器做下一 token 预测;设计大规模迭代层次聚类算法构建 20 万词表的像素级视觉 codebook,让图像 patch 与文本 token 在预训练期共享统一监督,模型由此内化视觉重建能力。 Overview of the UVU vision-language unified autoregressive framework. Continuous visual features from SigLIP-2 are projected and integrated with textual embeddings for autoregressive next-token prediction in an LM decoder, generating pixel-level image tokens. 实验效果:同为 3B/Qwen2.5 底座,RefCOCO 从 84.1 涨到 91.8、LISA 57.4→71.7、CVBench-3D 71.9→76.6、BLINK 46.9→52.8;相对去掉视觉监督的自版 UVU*,RefCOCO +6.2、LISA +12.1。注意力热图显示 UVU 更聚焦目标区域。 Comparison of visual attention heatmaps under three paradigms: from left to right, without visual supervision, AR + VQ, and UVU (Ours). 批判点评:对手没输干净:SEEDB 74.1 仍低于 LLaVA-OV 的 75.4,MMStar 55.0 输给 56.7,ScienceQA 91.3 远低于 GLM-4v 的 96.7;MME 2201.9 对 LLaVA-OV 2146.3 的领先也只有 2.6%。统一目前只覆盖理解侧,生成侧数据还没进训练,作者自己也在文中承认。 3. Unite-Audio:117M音频生成,CLAP登顶 UNITE-AUDIO: Joint Learning of Continuous Tokenization and Latent Flow Matching for Text-to-Audio Generation | 东京科学大学;清华大学;武汉大学;东京大学;蚂蚁集团 | arXiv:2609.28206 关键词:文本生成音频,联合训练,flow matching,Flow-GRPO,连续表征 ⚠️ 前序问题:文本生成音频的主流做法是两阶段:先训重建导向的音频 tokenizer 再冻结,然后在固定隐空间训生成模型。但为重建优化的表征未必适合生成,生成目标无法反过来塑造隐空间。 本文贡献:首个把连续音频表征学习和隐空间 flow matching 联合训练的 TTA 系统:掩码波形过在线编码器,EMA 编码器提供停止梯度的目标特征,重建与自监督生成预测耦合,让生成目标直接塑造隐空间;再用 Flow-GRPO 后训练强化文本条件对齐。 Overview of the proposed framework: (a) the reconstruction path and (b) the generation path. Blue denotes regions with gradient propagation, while gray denotes regions without gradient propagation. 实验效果:117M 隐空间生成模型(不含编解码器)在 AudioCaps-886 上 CLAP 0.534 全表最高(EzAudio 0.496、TangoFlux-RL 0.480),KL 1.057 最低,RTF 0.062;NFE 压到 4 时 RTF 仅 0.009 仍保 CLAP 0.512。 Comparison between conventional separate tokenizer-generator training and our joint single-stage training. 批判点评:分布指标难看:FD_OpenL3 84.76 是 EzAudio 38.26 的两倍多,FD_16k 43.25 也落后 GenAU 的 31.94,IS 11.34 不及 TangoFlux-RL 的 12.20——语义对齐赢了,音频分布保真还差得远;首个联合训练的含金量取决于你更看重 CLAP 还是 FD。 4. MotionSpec:频谱轨迹约束视频运动 MotionSpec: Spectral Trajectory Supervision for Motion-Consistent Video Generation | 卡迪夫大学;中国科大;华东师范大学 | arXiv:2609.28095 关键词:视频生成,运动一致性,傅里叶谱,光流,微调目标 ⚠️ 前序问题:文生视频单帧越来越真,运动却常失真:动作推进不一致、复杂动作下结构崩坏。标准生成目标对运动本身几乎没有专门约束,运动演化处于欠约束状态。 本文贡献:提出 MotionSpec 运动监督:谱轨迹一致性 STC 构造锚点相对的稠密运动轨迹,经时间傅里叶变换成运动谱体积,对齐幅度与相位同时约束运动强度和时序组织;辅以局部光流一致性 LFC 稳定相邻帧过渡,两者随时间步加权加进微调目标。 Overview of the pipeline. At each training step, the video generative model takes the prompt, timestep, noise level, and real data as inputs. 实验效果:在 Wan2.1 上微调,VMBench 均分 61.57→65.07,其中物体完整性 OIS 48.34→56.53、时序一致性 TCS 97.09→98.82;VideoJAM-Bench 运动分 85.3→93.61。定性对比里 Wan2.1 的骑车段后半程自行车几乎消失,MotionSpec 保留完整。 Qualitative Results. On the left are the results of Wan2.1, and on the right are the results of ours. 批判点评:外观分几乎没动:App 仅 81.95→82.20(+0.25);PAS 上单独用 LFC(24.27)反而高于完整版(23.74),说明两个目标存在拉扯。评测全部在 Wan2.1 一个底座上做,换模型的泛化性未知。 5. RecCAR:反向注意力补齐,音画更同步 All modalities are equal, but video is more equal: Closing the Cross-Attention Gap in Joint Video Generation | 巴伊兰大学;NVIDIA | arXiv:2609.27901 关键词:联合视频生成,跨模态注意力,KL正则,音画同步,人体解剖 ⚠️ 前序问题:联合视频-动作/视频-音频生成的扩散 Transformer 里存在不对称:伴随模态到视频的对应很强,但反过来约束视频生成的对应路始终弱——视频听动作指令的能力远弱于动作看视频的能力,人体解剖崩坏、音画不同步由此而来。 本文贡献:把两个方向的跨模态对应表示成视频 token 上的分布,定义反向对应差距;提出 RecCAR,KL 正则把弱方向对齐到以视频到模态对应为固定参考的强方向,即插即用地微调 EchoMotion 这类联合生成模型。 Illustration of asymmetric reciprocal cross-modal correspondence. For a motion token corresponding to the head, the VM correspondence correctly localizes the head region in the video, while the reciprocal MV correspondence is dispersed. 实验效果:视频-动作联合生成 Human Anatomy 0.69→0.75,VMBench 多数指标超 EchoMotion、CoMoVi、FlowMo;视频-音频生成音画失同步从 0.804 降到 0.752。散点图上每个 prompt 的反向差距在微调后一致下降。 Qualitative comparison of RecCAR against EchoMotion, CoMoVi, and FlowMo across three representative scenes. 批判点评:增益以自家基线为参照:0.75 的解剖分意味着仍有约四分之一样本不达标,音频侧 0.752 的失同步也谈不上小;teaser 里对比的 LTX-2 只做了定性展示,未给出与闭源模型在同一协议下的量化差距。 6. GestureFAR:9.3毫秒每token的流式手势 GestureFAR: Streaming Co-Speech Gesture Generation with Flow Autoregression | 罗切斯特大学;东京大学;加州大学圣克鲁兹分校;加州大学洛杉矶分校;Meta | arXiv:2609.21576 关键词:手势生成,流式,flow自回归,蒸馏,实时交互 ⚠️ 前序问题:流式陪聊手势生成此前靠离散运动 token 自回归,把高维连续运动压进有限 codebook,真实感和多样性都受损;而连续方法又难以保证逐 token 因果,实时交互卡在延迟上。 本文贡献:提出流式手势生成框架 GestureFAR:因果 VAE 把全身动作编码成可流式的连续隐变量,Transformer 对音频-运动上下文自回归,每 token 配 flow-matching 头从连续分布采样下一潜变量;再冻结因果骨干,用一致性加分布匹配目标把多步 flow 头蒸馏成单步前向。 Overview of GestureFAR. GestureFAR generates co-speech gestures from streaming audio by autoregressing over continuous motion latents. 实验效果:BEAT2 流式组 FGD 3.08 全面最佳(LiveGesture 4.57、MIBURI 8.06),BC 0.741 接近离线最优;1 步蒸馏后每 token 9.3ms,是教师 8 步 24.4ms 的 2.6 倍速、MIBURI 62.9ms 的 6.8 倍速。 Qualitative comparison on BEAT2. We visualize generated full-body gestures from different methods under the same speech inputs, with the semantically salient words highlighted in red. 批判点评:BC 0.741 仍低于 MIBURI 的 0.790 和 LiveGesture 的 0.794,多样性 13.24 也略逊离线系;单步蒸馏换速度后 FGD 反而比教师还好(3.08 vs 3.17),这个反常更像评测协议的产物而非能力证明,文中未解释。 7. DeltaS:状态漂移选KV,六个基准+2.1 DeltaS: Reading the Gated Linear Attention State for KV Cache Eviction in Streaming Video | Maum AI;首尔大学;延世大学 | arXiv:2609.27470 关键词:KV缓存淘汰,流式视频,线性注意力,混合架构,训练无关 ⚠️ 前序问题:视频语言模型转向线性/全注意力混合架构后,线性注意力状态固定大小,但全注意力 KV 缓存仍随流式视频无限增长;流式场景下问题还没来就得决定淘汰谁,现有基于位置、注意力或 KV 本身的信号都拿不到这个先验。 本文贡献:提出训练无关的 DeltaS:读门控 delta 线性注意力的循环状态,用一块帧内状态的归一化变化量(状态漂移)衡量该块带来多少新信息,漂移大的块对应 KV 留存;信号计算只占前向的 1.9%,无需代理查询。 Overview of DeltaS. Each video chunk updates the recurrent linear-attention states, whose normalized changes yield a shared score for retaining the corresponding KV entries across full-attention layers. 实验效果:预算与留存策略对齐的受控比较里,状态漂移信号全面胜过位置/注意力/KV 三类基线;六个长视频基准平均超最强无查询基线 2.1 分,最长基准 LVBench 超 5.6 分。 Performance margin as a function of KV retention rate. DeltaS generally shows larger margins at lower retention rates. 批判点评:增益随留存率升高而收窄:留存 token 超过约四成后在 EgoS 上反而落后基线 0.7 分,高预算场景收益有限;hybrid 架构前提也把它限死在门控 delta 线性注意力这一类骨干上。 8. DriftAudio:一步生成后训练FAD降34% DriftAudio: Marginal Drifting for Distributional Post-Training of One-Step Text-to-Audio Generators | 悉尼科技大学 | arXiv:2609.27598 关键词:一步生成,分布后训练,文本生成音频,Drifting,FAD ⚠️ 前序问题:一步文本生成音频模型把推理成本打下来之后,生成分布仍欠火候;常规后训练依赖逐条件配对的真实样本,自由文本条件下一个 prompt 往往只有一两条真实参考,逐条件 Drifting 根本做不了。 本文贡献:提出 DriftAudio 边缘分布后训练:把 Drifting 从逐条件搬到边缘音频分布上做,在冻结的 PANNs 特征空间里用重采样真实样本、滚动生成的 FIFO 样本库和当前批次共同估计漂移场,得到 detached 训练目标只更新生成器,一步推理流程原样保留。 Overview of DriftAudio. The one-step generator remains text-conditioned, while Drifting is performed on the marginal audio distribution in feature space. 实验效果:从 MeanAudio 出发 FAD 降 33.9%、FD 降 17.6%,KL 与 CLAP 同步改善;从 FdAudio 出发 FAD 1.22→0.99,一步生成的分布质量明显抬升。 Mean and covariance components of FAD for the four evaluated models. Numbers above the bars denote total FAD. 批判点评:不是白捡的:从 FdAudio 出发时 IS 和 CLAP 出现回退,说明边缘分布对齐和条件保真之间存在交换;方法绑定冻结的 PANNs 特征空间,换更强的音频评价骨干是否还成立未验证。 9. InGuard:嵌入层安检,省一半去噪步 InGuard: Towards Generalized Inner Guardrail for Safe Text-to-Image Generation | 阿里巴巴 AAIG | arXiv:2609.27620 关键词:文生图安全,内嵌护栏,嵌入改写,潜空间检测,RevGen ⚠️ 前序问题:T2I 的安全护栏都装在外面:前置 prompt 分类器加后置图像分类器,两者不用模型自身表征——prompt 检查准头有限,图像检查要等全量去噪烧完才跑,而且违规 prompt 只能一刀切拒绝,本可改写成安全输出的也被扔掉。 本文贡献:提出 InGuard 内嵌护栏:文本编码器嵌入上直接做风险三级分类(不安全/风险/良性),SAGE 对风险 prompt 做嵌入空间软门控改写使其输出安全图而非拒答,去噪中途用一步干净潜变量估计做潜空间检测,发现风险即刻提前终止;配套 RevGen Safety 基准,1 万条 prompt 由真实图反生成并注入受控 IP 角色。 Overview of our InGuard framework. Three complementary components work in sequence: prompt embedding risk classification, SAGE embedding-space safety enhancement, and latent detection with early termination. 实验效果:五个开源 T2I 模型上安全率 97.9%-98.8%,追平或超过外置护栏;良性扰动少 57.5%-73.5%,参数少约 3.7 倍,50%-55.6% 的去噪步被跳过;潜空间检测在 44%-50% 计算量处就逼近图像级 F1。 Pretraining effects on Avg F1 across five LDMs; latent detection vs denoising-step fraction used as a compute proxy. 批判点评:SAGE 改写不是万能的:论文自己承认增强失败时只能靠潜空间检测兜底拦截,等于承认存在漏改样本;RevGen 由自家反生成流水线构造,跨基准的代表性有待第三方复检。 10. CineGuard:11种运镜的抄袭检测 Did You Steal My Shot? Pioneering Camera Motion Plagiarism Detection in Generative Videos | 河海大学 | arXiv:2609.22267 关键词:运镜抄袭,生成视频,涡度,基准,版权保护 ⚠️ 前序问题:运镜是影视级 IP,但生成视频用一句 prompt 就能复刻高价值运镜;现有相似度检测都盯着画面内容,训练数据把运镜和内容缠在一起,传统光流又表达不了复杂相机运动——运镜抄袭至今无人管。 本文贡献:构建首个运镜分析基准 CineFlow:11 种运镜风格的模拟数据集把运镜与内容解耦;提出 CineGuard 检测网络,在光流上叠加流体力学涡度作为平移不变线索,3D 卷积提 tubelet 嵌入后过 ViT 主干加全局相机适配器,对比学习聚拢同运镜视频。 Architecture of the motion plagiarism detection network. The input flow is first augmented with vorticity, then a 3D convolution extracts tubelet embeddings. 实验效果:抄袭检测超最强基线 3.02 倍;在 Veo 3.1 与即梦的商业生成视频上,余弦相似度 0.56-0.83,基线全部落在 0.35 以下,Precision@5 最高 0.73 对基线约 0.27。 Performance of plagiarism detection on generative videos. Left: cosine similarity between videos. Right: Precision@5 for retrieving videos with the same motion label. 批判点评:局限写在图里:FPV Drone 这类复合运镜相似度掉到 0.56,是全表最低,混合运镜仍是软肋;基准是简单几何场景模拟出来的,真实影视素材上的迁移只有商业视频小样本佐证。 趋势观察 编辑模型的主战场从通用能力挪向行业指标 KwaiMind 把 CTR 预测直接做成奖励信号训进编辑模型,线上 A/B 实测 +2.44%;InGuard 则把安全检查搬进生成管线内部,省下一半去噪步。两条线指向同一件事:图像生成的竞争正在从「生成质量」转向「商业约束下的生成质量」,奖励函数和护栏的构造方式开始比骨干网络更值钱。 视觉监督正在从后训练前移到预训练 UVU 用像素级 codebook 把视觉监督直接塞进预训练,RefCOCO 一项涨 7.7 分;MotionSpec 则在视频微调里用频谱约束补上运动监督的缺口。共同逻辑:等表征定型再补监督太晚了,感知质量的上限在预训练阶段就被决定。 一步生成与流式生成进入质量补课期 DriftAudio 给一步 TTA 做边缘分布后训练,FAD 降三分之一;GestureFAR 用单步蒸馏把手势生成压到每 token 9.3ms;Unite-Audio 把去噪步数压到 4 还能保住 CLAP。速度战的下一程是质量战——谁能在一步、几步的预算里把分布差距补回来,谁就拿到实时产品化的门票。 混合架构的两组记忆开始学会协作 DeltaS 证明线性注意力的循环状态可以反过来指挥全注意力 KV 缓存的淘汰,六个长视频基准平均 +2.1;Unite-Audio 则让生成目标直接塑造 tokenizer 隐空间,打破「先冻结表征再训生成」的铁律。两篇都在挑战模块化管线里各训各的默认假设。 人工智能炼丹君 整理 | 2026-09-25 更多 AIGC 论文解读,关注微信公众号「人工智能炼丹君」 每日更新 · 论文精选 · 深度解读 · 技术脉络 微信搜索 人工智能炼丹君 或扫描下方二维码关注
2026年09月25日
5 阅读
0 评论
0 点赞
2026-09-25
AIGC 基本功|性能建模与 Profiling:算力、带宽与显存账本-Roofline
性能建模与 Profiling:算力、带宽与显存账本 所属方向:推理加速 | 难度:进阶 | 前置知识:混合精度与数值稳定性、自注意力机制的计算与显存账本 关键词:性能建模、Profiling、Roofline、算术强度、memory bandwidth、FLOPs、MFU、PyTorch Profiler、Nsight 01. 为什么需要它 先看一组在同一台笔记本(Apple silicon,numpy fp32)上此前实测出来的数字(第 04 节另给本次复核快照),出自文末附录的 machine_probe.py,可以自己复现: GEMM : F = 0.998 GFLOP, D = 8.1 MB, 实测 0.498 ms(2003 GFLOP/s) 逐元素 : F = 1.000 GFLOP, D = 12000 MB, 实测 189.649 ms(5 GFLOP/s) → FLOPs 基本相同,耗时差 381 倍。 两段计算量的浮点运算次数几乎一样(都约 1 GFLOP),耗时差了 381 倍。如果拿「FLOPs 少的算子更快」这类直觉去做优化决策,在这个数字面前是反的:这里逐元素算子与 GEMM 的 FLOPs 相同,却因为每次都要把数据从内存搬进搬出,被带宽死死卡住。 再比如一个真实场景的优化评审:有人提议「把 LayerNorm 内部换成 fp8 计算单元重写,算力能翻几倍」。查一下账(第 03 节会算):LayerNorm 的算术强度只有约 1.5 FLOP/Byte(下文不含 beta 偏置的简化 LayerNorm),远低于 A100 的 ridge point 200.6,是典型的带宽受限算子——给它换更快的算力,在只改变算力峰值、访存不变的模型中加速比是 1.00,一分钱收益都没有。反过来,同是 fp8,把它用在 decode 阶段的大权重 GEMM 上却有收益,但收益来自权重字节减半,不是算力翻倍。同一笔投资,用在哪类算子上,结论完全相反。 这篇的目的就是把这套「先算账再动手」的方法补齐:两个账本——速度账(算力 vs 带宽)和容量账(显存四项)——加上一套实测手段(profiler)。量化、算子融合、KV cache 管理、并行切分,所有推理优化节点的收益判断都站在这篇的地基上。 02. 最小可用理解 三句话讲完核心: 任何算子的耗时有一个下界,由两种资源里更慢的那个决定:$T \ge \max(F/P_{\text{peak}},\ D/\beta)$。$F$ 是算子要做的浮点运算数,$P_{\text{peak}}$ 是硬件与该指令类型匹配的峰值算力;$D$ 是算子要搬运的字节数,$\beta$ 是带宽。算得再快也快不过「数据没到」。 算术强度 $I = F/D$ 决定卡在哪:与 ridge point $I^{\ast} = P_{\text{peak}}/\beta$ 比较,$I < I^{\ast}$ 是带宽受限,优化方向是少搬字节(融合、量化、FlashAttention);$I > I^{\ast}$ 是算力受限,优化方向才是少算(更好算法、更低精度计算单元)。 显存容量是另一本独立的账:权重 + KV cache + 激活 + 额外开销四项加总,才决定一张卡能塞多少并发。「权重放得下就能跑」只覆盖了四项里的第一项。 03. 数学推导 3.1 时间下界为什么取 max 一个算子要做 $F$ 个浮点运算(FLOP,口径:一次乘加记 2 个 FLOP),硬件每秒最多做 $P_{\text{peak}}$ 个——就算计算单元一刻不停,也至少要 $F/P_{\text{peak}}$ 秒。同理,算子要把 $D$ 字节的数据在内存和计算单元之间搬个来回,总线每秒最多搬 $\beta$ 字节——至少要 $D/\beta$ 秒。这两件事用的是不同资源,理想情况下可以完全重叠(算上一批数据的同时搬下一批),所以总时间的下界是两者取 max: $$T \ge \max\left(\frac{F}{P_{\text{peak}}},\ \frac{D}{\beta}\right)$$ 注意这是下界:真实 kernel 还有 kernel launch、同步、缓存未命中、TLB miss 等额外开销,实测只会更慢。后面的实验用经验标定值和估算流量代入,此时算出的只是模型估计;计时波动或缓存会产生超过 100% 的达成率,并不违反物理下界。差距也不能全算成可消除的优化空间。 3.2 算术强度与 ridge point 定义算术强度: $$I = \frac{F}{D}$$ 物理含义:每从内存搬 1 字节数据,能换来多少次浮点运算。它取决于实现与所选存储层级;实际缓存命中和重复加载又与硬件有关,不能视为完全与机器无关。再定义机器的 ridge point: $$I^{\ast} = \frac{P_{\text{peak}}}{\beta}$$ 物理含义:这台机器「算」和「搬」一样快的分界点,单位都是 FLOP/Byte,所以可以比。把 $I$ 与 $I^{\ast}$ 代回 3.1 的下界: $I < I^{\ast}$(带宽受限):$D/\beta$ 那一项更大,$T \approx D/\beta$,模型性能上界 $F/T \le I \cdot \beta$——性能与算力峰值无关,只跟 $I$ 成正比,这就是 roofline 图上那条斜线; $I \ge I^{\ast}$(算力受限):$T \approx F/P_{\text{peak}}$,性能封顶在 $P_{\text{peak}}$,这就是平顶。 $P_{\text{peak}}$ 用哪个口径要非常小心:A100 SXM 的 BF16 dense 是 312 TFLOP/s,A100 不支持原生 FP8 Tensor Core;H100 SXM 的 BF16 dense 约 989、FP8 dense 约 1979 TFLOP/s,差 6.3 倍——口径选错,受限类型的判断直接反掉(第 06 节细说)。 3.3 给几类算子记账 GEMM:$A[M,K] \times B[K,N] \to C[M,N]$。每个输出元素要做 $K$ 次乘加,共 $M N K$ 次,乘加各记一次: $$F_{\text{gemm}} = 2MNK,\qquad D_{\text{gemm}} = b\,(MK + KN + MN)$$ $b$ 是每元素字节数(bf16 取 2)。读 $A$、读 $B$、写 $C$ 各一遍,统计量这类小东西忽略。$M$ 越大,权重 $B[K,N]$ 被摊得越薄,$I$ 越高——这解释了为什么大 batch 的 GEMM 是算力受限、小 batch 的 GEMM 是带宽受限。 逐元素算子:$n$ 个元素各做 1 次运算,读 2 份写 1 份: $$F = n,\qquad D = 3 b n,\qquad I = \frac{1}{3b}$$ $I$ 是常数(fp32 下约 0.083,bf16 下约 0.17),与 $n$ 无关。在固定 dtype 与访存模型下,增大 n 本身不会提高 I,因而不会像增加 GEMM 的复用维度那样跨越 roofline 分界。 softmax 的读写口径:普通三阶段实现约 3 读 2 写。若整行能驻留片上存储,融合实现可约 1 读 1 写,得到 2.5 倍的理想流量比;独立的 online normalizer 通常先流式求归一化量,再重读输入写出概率,为 2 读 1 写,不能混为一谈。本文用约 $5nd$ 作为普通 softmax 的操作计数;max、exp 与除法并非都能跑在 Tensor Core 上,online 递推还会增加标量运算。可对照 在线归一化论文与 Triton 整行融合示例的不同前提。 单头注意力($N$ 是 token 数,$D_{h}$ 是每头维度):两次 $N \times N \times D_{h}$ 的矩阵乘加 softmax: $$F_{\text{attn}} = 4N^{2}D_{h} + 5N^{2}$$ 朴素实现的分数矩阵 $S$ 和概率矩阵 $P$ 都要落显存,二次项访存约 $4N^{2}$;FlashAttention 让 $N^{2}$ 只留在片上 SRAM,若仅统计每份输入读取一次与输出写回一次,可得到不可避免的数据流量下界;真实分块内核会反复读入 K/V 或 Q 等数据: $$D_{\text{朴素}} = b\,(4N^{2} + 4ND_{h}),\qquad D_{\text{flash,min}} = b \cdot 4ND_{h}$$ 代入 $N=4096$、$D_{h}=128$、bf16(脚本 roofline_model.py 实算):朴素 $D = 138.41\ \text{MB}$、$I = 62.7$,落在 A100($I^{\ast} = 200.6$)的带宽受限区;理想最低流量 $D_{\mathrm{flash,min}} = 4.19\ \text{MB}$、$I = 2068$,跳进算力受限区。FLOPs 一动没动(×1.000),roofline 时间下界从 0.089 ms 降到 0.028 ms(×3.2),这是按最低流量计算的理想示例,33 倍不是实际 FlashAttention 访存或速度的测量;有限片上存储下还需分块 IO 模型。FlashAttention 那篇的完整递推在知识树的下一节点展开,这里先用 roofline 把它的收益定位清楚。 图 1:这张图要看三样——散点的横坐标是各算子的算术强度 $I$,点越靠右越「算得过来」;按 I 与 I 的相对位置给点作模型分类;点到上界的距离并不能单独证明真实瓶颈;FlashAttention 会改变访存量和横坐标,无法直接把这段垂直差距视为它的加速收益。* 3.4 优化收益的上限 把算力翻倍($P_{\text{peak}} \to 2P_{\text{peak}}$),加速比是: $$S_{\text{算力}} = \frac{\max(F/P_{\text{peak}},\ D/\beta)}{\max(F/(2P_{\text{peak}}),\ D/\beta)}$$ 带宽受限时分子分母都是 $D/\beta$,$S_{\text{算力}} = 1$:白花钱。算力受限时加速至多为 2;当 $1<I/I^{\ast}<2$ 时,算力翻倍会遇到带宽上限,加速为 $I/I^{\ast}$。带宽翻倍对称地反一次。roofline_model.py 第四节把第 3.3 节的每个算子都代了一遍(A100 口径): 算子 I/I* 算力×2 的加速 带宽×2 的加速 逐元素 add [16M] 0.001 1.00x 2.00x LayerNorm [4096, 3072] 0.007 1.00x 2.00x softmax 朴素三遍 [4096, 4096] 0.002 1.00x 2.00x GEMM 512x4096x4096 2.041 2.00x 1.00x attention 朴素 [N=4096,D=128] 0.312 1.00x 2.00x attention ideal-min [N=4096,D=128] 10.307 2.00x 1.00x 这张表就是「投资之前先看图 2」的数字版:你的 kernel 在分界线哪一侧,决定哪类投资是零收益。 图 2:这张图要看什么——横轴 $I/I^{\ast}=1$ 那条竖线就是分界线:线左边算力翻倍的理想收益为 1,带宽翻倍收益在 1 到 2 之间;只有 I/I≤0.5 时完整得到 2 倍。线右边对称,I/I≥2 才完整得到算力翻倍的 2 倍。投入硬件或投入算子融合之前,先看自己在哪一侧。 3.5 显存的容量账 速度账之外另有一本容量账。下面以未分块 prefill、同时处理 BS 个 token 为例,显存分四项: $$M_{\text{total}} = N_{\text{params}} \cdot b_{w} + 2BSLd_{\text{kv}}b + BS \cdot a + \rho \cdot M_{\text{sum}}$$ 逐项说物理含义:$N_{\text{params}}$ 是参数量、$b_{w}$ 是每参数字节数(fp16 为 2)——这一项与并发无关,是常数;$B$ 是并发序列数、$S$ 是序列长度、$L$ 是层数、$d_{\text{kv}}$ 是每层 KV 总维度(GQA 模型用实际的 KV head 数乘头维度)——KV cache 每个 token 每层都要存一份 K 和一份 V,所以是 $2BSLd_{\text{kv}}b$ 字节,随并发线性增长;$a$ 是每 token 的激活峰值(推理不保留整层中间结果,但当前层十几份临时张量要同时活着,$d_{\text{model}}=4096$、fp16 时取约 128 KiB/token 是经验值,随实现差距很大);$\rho$ 是额外开销率,$M_{\text{sum}}$ 是前三项之和。逐 token decode 时通常只有 B 个活跃 token,应将激活项 BS·a 改成 B·a;KV cache 仍随 BS 增长。分块 prefill 则按实际活跃 chunk 计。第 04 节把未分块 prefill 的教学账本代进一个 7B 模型。 04. 代码实现 三个脚本全部只用 numpy,因为 roofline 的方法论不依赖 GPU:同一台机器、同一套口径,把「峰值」和「落点」都实测出来,预测和实测的差距才看得见。本次复核实测峰值:$\beta = 71.8\ \text{GB/s}$、$P = 1635\ \text{GFLOP/s}$、$I^{\ast} = 22.8\ \text{FLOP/Byte}$(注意这是「numpy 能摸到的上限」,不是芯片标称值——方法论可比的前提是口径一致)。 4.1 标定两个峰值 带宽用两输入向量加法模式标定(不是含乘法的 STREAM triad):用足够大的工作集降低缓存影响;128 MB 是否超过目标机器缓存仍需核对,测得的是该访问模式的有效带宽。 def measure_bandwidth(n: int = 32_000_000, repeat: int = 8) -> float: """z = x + y:读 2n、写 n,共 3n 个 fp32 元素。""" x = np.ones(n, dtype=np.float32) y = np.ones(n, dtype=np.float32) z = np.empty(n, dtype=np.float32) dt = _best(lambda: np.add(x, y, out=z), repeat) return 3 * n * FP32 / dt # Byte/s,_best 取 repeat 次最快 算力用足够大的方阵乘标定(4096 的方阵乘访存被摊薄,$F = 2n^{3}$)。实跑输出: 本机标定(arm64 / numpy 2.1.3 / fp32,2026-10-02 02:21:20) 实测可达带宽 beta = 71.78 GB/s 实测可达算力 P = 1635.04 GFLOP/s ridge point I* = 22.8 FLOP/Byte 4.2 把真实算子打上 roofline 关键测试对象是 matmul-softmax-matmul 微基准;此处省略 1/√D 缩放,不是完整的生产 attention。下面这七行就是「账本」本身——每一行右边标了它读写了几个 $N^{2}$ 量级的遍数,$D$ 就是这么数出来的,不是拍脑袋: def attention_naive(): np.matmul(q, kt, out=S) # 写 S 1 遍 np.max(S, axis=-1, keepdims=True, out=rowmax) # 读 S 2 遍 np.subtract(S, rowmax, out=S) # 读写 S 4 遍 np.exp(S, out=S) # 读写 S 6 遍 np.sum(S, axis=-1, keepdims=True, out=rowsum) # 读 S 7 遍 np.divide(S, rowsum, out=S) # 读写 S 9 遍 np.matmul(S, v, out=O) # 读 S 写 O 10 遍 实跑结果(本机口径,$I^{\ast} = 22.8$): 算子 I 模型分类 T估计 T实测 达成率 实测算力 逐元素 add(1 遍) 0.08 带宽 5.35ms 5.07ms 106% 6.3 GF/s add+relu 两遍(预分配) 0.10 带宽 8.92ms 16.76ms 53% 3.8 GF/s add+relu 写成一行(有中间数组) 0.10 带宽 8.92ms 26.80ms 33% 2.4 GF/s LayerNorm [8192,4096] 预分配 0.15 带宽 18.70ms 32.56ms 57% 6.2 GF/s attention 朴素 [N=2048,D=128] 12.61 带宽 2.40ms 11.41ms 21% 190.1 GF/s GEMM 512x4096x4096 204.80 算力 10.51ms 16.18ms 65% 1061.5 GF/s GEMM 4096x4096x4096 682.67 算力 84.06ms 84.82ms 99% 1620.3 GF/s 逐行解读,三种达成率各说明一件事: 逐元素 add 106%:与标定时相同访问模式,已接近本次有效带宽。超过 100% 来自经验标定与实际计时的差异,不能理解成超过硬件物理峰值。 GEMM 大矩阵 99%:已接近本机同类 GEMM 标定值;这不能证明其他实现没有改进空间。 朴素 attention 21%:实测约为模型估计时间的 4.8 倍,可能涉及缓存、指令吞吐、线程和内核调度;仅凭总时间不能分离原因,也不能直接解释成多搬了几倍字节。 图 1 使用完整脚本中的 8 个算子(正文节选了 7 行)。点到 roofline 的距离表示相对模型上界的性能差距;它本身不能区分额外访存、指令开销或同步等原因。 4.3 显存账本与 batch 拐点 未分块 prefill 的容量教学账代一个 7B 模型(32 层、$d_{\text{kv}} = 4096$、fp16): 模型:7B,fp16 权重 = 13.0 GiB,32 层,d_kv = 4096 KV cache 单价:512 KiB / token (一条 4096 长的序列 = 2.00 GiB) batch 权重 KV cache 激活 额外开销 合计 假设 78 GiB 可用预算 1 13.0G 2.00G 0.50G 4.66G 20.20G 装得下 8 13.0G 16.00G 4.00G 9.91G 42.95G 装得下 16 13.0G 32.00G 8.00G 15.91G 68.95G 装得下 24 13.0G 48.00G 12.00G 21.91G 94.95G OOM 32 13.0G 64.00G 16.00G 27.91G 120.95G OOM 两个读数:权重是常数,batch 再大都是 13.0 GiB;KV cache 按 512 KiB/token 线性涨,batch 16 时(32 GiB)已经是权重的 2.5 倍。此处 30% 是人为设定的额外开销/前三项总量比率,不是实测额外开销率;它也不同于 PagedAttention 论文中 KV cache 已分配容量的浪费比例。同一张卡只把本教学模型的额外开销率改成 4%,batch 24 从 94.95 GiB 降到 75.96 GiB,从 OOM 变装得下:不换卡、不改模型,这个假设账本跨过一个 batch 档位;真实分页收益需只对可优化的 KV 分配项建模,并检查实际可用显存。 容量账决定了「能开多大」,速度账决定「该开多大」。固定一份 64 MiB 的权重矩阵,扫 batch(每个 batch 位置喂一个 token),本机实测: B F (GFLOP) D (MB) I 模型分类 延迟 ms 吞吐 K/s 1 0.034 67.14 0.50 带宽 1.718 0.6 2 0.067 67.17 1.00 带宽 8.232 0.2 4 0.134 67.24 2.00 带宽 8.161 0.5 8 0.268 67.37 3.98 带宽 4.477 1.8 16 0.537 67.63 7.94 带宽 4.436 3.6 32 1.074 68.16 15.75 带宽 4.461 7.2 64 2.147 69.21 31.03 算力 4.865 13.2 128 4.295 71.30 60.24 算力 6.362 20.1 256 8.590 75.50 113.78 算力 8.612 29.7 延迟:B=1 时 1.718 ms,B=256 时 8.612 ms,涨 5.0 倍;吞吐涨 51.1 倍。理想带宽模型中,加大 batch 可摊薄权重读取;本机表格并不呈严格平坦段。B=32 到 64 之间越过模型分界 $I^{\ast}=22.8$,这不保证实测曲线在同一处出现锐利拐点。B=256 时 64 MiB 的权重读取被摊到每样本 0.25 MiB。 B=2 的延迟是 B=8 的 1.84 倍,可能涉及 BLAS 内核选择、线程调度或计时波动;仅凭这张时间表不能确证原因。roofline 看不见这类事情,定 batch 必须结合实测。 图 3:这张图要看什么——左边是实测延迟,右边是实测吞吐,曲线不保证出现理想平坦段;竖虚线是本次标定值代入模型后的分类交界;B=2 处那个尖是 roofline 模型解释不了的库行为,标出来是为了让你对「模型给方向、实测给结论」有体感。 图 4:这张图要看什么——左图是显存四项的堆叠:蓝色权重是常数,红色 KV cache 随并发线性涨,越过黑色可用线就是 OOM;右图是只改一个参数(额外开销率 30%→4%)的效果,batch 24 从 OOM 变装得下。 05. 工业级实现对照 最小实现是方法论,生产里测量走的是另一套工具,但问的是同一组问题。 PyTorch Profiler(pytorch/pytorch · torch/profiler/profiler.py · profile,以 2026-09 的实现为准)是最常用的第一站: from torch.profiler import profile, ProfilerActivity with profile(activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA], record_shapes=True, profile_memory=True, with_flops=True) as prof: model(x) print(prof.key_averages().table(sort_by="cuda_time_total", row_limit=15)) prof.export_chrome_trace("trace.json") 和本文方法的对应关系:key_averages().table 按算子事件键聚合 CPU/CUDA 耗时,而非直接按 GPU kernel 名称聚合,回答「时间花在哪」;with_flops=True 只为支持的算子(如矩阵乘与二维卷积)估算 FLOPs,不能视为全模型所有操作的完整计数,除以耗时和相应算力峰值得到该算子的利用率估计,并不是通常定义的模型级 MFU(Model FLOPs Utilization,模型理论 FLOPs ÷ 墙钟时间 ÷ 硬件峰值算力);profile_memory=True 抓张量级显存分配,对应容量账的激活项。训练报告里常见的整体 MFU 是同一口径的粗化:拿模型理论 FLOPs 除以墙钟时间和卡数峰值,这个数在 decode 型负载里天然上不去,原因见第 08 节第 4 条。 Nsight 全家桶是更细的一层:Nsight Systems(nsys)看时间线——kernel 之间的空隙、同步等待、通信重叠,对应「下界假设完美重叠」不成立的部分;Nsight Compute(ncu)看单个 kernel,它的 Speed Of Light 面板直接给出 Compute Throughput 和 Memory Throughput 两个百分比——那就是 roofline 的粗版:应结合指令流水线、occupancy、缓存层级和 stall 原因判断,不能仅比较两个百分比便确定瓶颈。实操顺序通常是:nsys 找到热点和空隙,ncu 对热点 kernel 看 SOL 定受限类型,再决定投算力还是投带宽。 和框架选择的关系:SDPA 会根据输入与硬件分派后端,不能断言所有 diffusers/transformers 都默认使用 FlashAttention。本文 2068 FLOP/byte 来自每份 Q/K/V 只读一次的理想下界;真实 FlashAttention 还需考虑分块重读。小 batch GEMM 的权重读取成本通常很高,但具体量化收益与核实现、反量化和 batch 都有关。 06. 代价与边界 roofline 是模型,模型有假设。四条主要假设和不成立时的样子: 假设计算与访存完全重叠。真实 kernel 在算和搬之间来回切换,还有 launch 和同步的开销。小 kernel(本机实测 B=1 约 1.7 ms 的场景)里这些固定开销占比不小,模型系统性偏乐观。 假设峰值算力是一个数。实际有 fp16/bf16/fp8/稀疏好几档,差 6 倍以上;按 BF16 dense 口径,H100 SXM 相比 A100 40GB SXM 的 ridge point 从约 201 升到 295——换新卡后更多 kernel 会落进带宽受限区,拿旧卡的直觉做判断会错。 假设 $D$ 与缓存无关。账本里的 $D$ 按落盘遍数数,但缓存命中会让有效带宽远大于 DRAM 标称值。锚点论文 Hierarchical Roofline(arXiv:2009.05257)就是把单一 roofline 扩展成每级缓存一条,用于定位数据移动发生在哪一层。 假设算子孤立。decode 阶段 GEMM 的「权重」每层都要重新读一遍,全局账(整个模型、整个请求)和单算子账结论可能不同;通信算子(allreduce)的账本里延迟和消息数占大头,照搬本文公式会算错。 什么时候不用它:动态 shape、算子间强耦合(融合边界在变)、通信密集的分布式场景——这些先上 profiler 看时间线,roofline 只对「单 kernel、口径清晰」的问题给下界。坦诚标注:本文所有「实测」都来自一台笔记本的 numpy,数字本身不可迁移,理想分段趋势也可能被缓存、内核选择和调度打破。 07. 经典论文脉络 Roofline: An Insightful Visual Performance Model(Williams et al., CACM 2009,未挂 arXiv):提出算术强度与 ridge point,一根折线把「算力受限/带宽受限」变成可判定的题。一切性能建模的原点。 Hierarchical Roofline Performance Analysis for Deep Learning Applications(Yang et al., 2020):把 roofline 按缓存层级展开,回答「多搬的字节发生在哪一级存储」,是本文锚点论文,也补了单一 roofline 在深度学习负载上最大的盲区。 FlashAttention(Dao et al., 2022):IO-aware 的代表作——FLOPs 一动不动,靠 tiling + online softmax 减少 N×N 中间量落地及相应 IO,把注意力从带宽受限拉进算力受限。roofline 视角下「优化带宽」的教科书案例。 Mixed Precision Training(Micikevicius et al., 2017):换 dtype 同时改两本账——$b$ 变小省字节,算力单元换挡提峰值。哪半边有收益取决于受限类型(已发长文专门讲数值稳定那半)。 Efficient Memory Management with PagedAttention(Kwon et al., 2023):在论文比较的服务负载中显著减少 KV cache 分配浪费;这些比例不能直接乘到权重、激活和所有显存上,虚拟内存的分页思想搬进显存管理。知识树里 KV cache 一篇的主角。 五篇连起来是一条线:先有判定工具(roofline),再按受限类型各给一把钥匙——算力侧(混合精度)、带宽侧(FlashAttention)、容量侧(PagedAttention)。 08. 常见误解 「FLOPs 少的算子更快」。本文开头的实测:同样约 1 GFLOP,GEMM 0.498 ms,逐元素 189.6 ms,差 381 倍。FLOPs 只在算力受限区才和耗时挂钩;带宽受限区里,决定耗时的是字节数。 「显存够放权重就能跑」。7B fp16 权重只要 13 GiB,但 batch 24 时 KV cache 48 GiB + 额外开销 21.9 GiB,假设有 78 GiB 可用预算的设备照样 OOM。容量账要四项加总,KV cache 那一项随并发线性涨,batch 16 就反超权重了。 「新卡算力翻倍,我的推理一定提速」。带宽受限的 kernel 加速比是 1.00(第 3.4 节表格里的 1.00x)。H100 SXM 的 FP8 dense 峰值约为 A100 BF16 的 6.3 倍,但实际换卡同时还改变带宽与内核;带宽受限时应估算字节数/实际带宽。权重量化会减少读取量,单独提高计算峰值则未必有用。 「MFU 低就是实现烂」。decode 阶段每个 token 过一遍全部权重,$I$ 天然低于 ridge point,MFU 高不了——这是负载形状决定的,不是代码烂。看 MFU 前先分清 典型 prefill 与小 batch decode 的负载形状;足够大的 decode batch、长上下文注意力或通信可能改变瓶颈。 「账本算出来的就是实际」。朴素 attention 模型估计 2.40 ms,实测 11.41 ms(差 4.8 倍);batch 扫描里 B=2 延迟是 B=8 的 1.84 倍(原因需进一步 profile)。模型给方向和上限,实测给结论——两个都要,缺一个都会做出错误决策。 09. 动手验证 三个都能在笔记本上跑(附录有完整代码): python machine_probe.py——预期:逐元素 add 达成率 ≈100%,GEMM 大矩阵 ≈100%,朴素 attention 明显低于 50%。如果你的机器上 attention 达成率反而很高,多半是缓存把 $N^{2}$ 矩阵装下了,把 N 调大一倍再看。 python memory_ledger.py——观察实际延迟/吞吐,不预设理想三段式;按 I 与 I* 标记模型分类交界(本机在 32→64 之间)。改 batch_sweep 里的 B 列表,看吞吐什么时候不再涨。 打开 roofline_model.py,仅把 A100 示例的 peak_flops 改成 1979e12、其余不变,重跑第五节(这是控制变量实验,不代表实际 H100,真实换卡还需更新带宽与 dtype)——预期:GEMM 拐点从 M∈(128, 256] 右移,更多算子被判为带宽受限。这一步会让你体感 ridge point 抬高对优化决策的影响。 10. 延伸阅读 按知识树的依赖关系,从这篇出发有三个方向: 往上游:自注意力机制的计算与显存账本——本文 3.3 节那笔注意力账的完整推导;混合精度与数值稳定性——换 dtype 的另一本账(数值范围与累加精度)。 往下游(受限类型各一把钥匙):FlashAttention(带宽侧)、KV Cache 与自回归视频生成(容量侧)、算子融合与 CUDA Graph(launch 开销侧,已排期)、量化(字节侧,已排期)。 目录:完整知识树见博客目录页 AIGC 基本功知识树,按依赖顺序排好了先修课。 继续阅读 FlashAttention 为什么不需要存下注意力矩阵:本文 3.3 节的最低流量示例省略了真实分块重读,那篇把 online softmax 的递推式一步步推出来。 附录:完整代码 09 节用到的脚本全文如下(machine_probe.py、roofline_model.py、memory_ledger.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 machine_probe.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """在同一台机器上把 roofline 的「峰值」和「落点」都实测出来。 只依赖 numpy。直接 `python machine_probe.py` 即可运行,约 20 秒。 为什么要自己测一遍 ------------------ roofline_model.py 里的 A100/H100 峰值是厂商标称值,读者手边不一定有 GPU。 这个脚本用同一套口径(F 记 2·MAC、D 记读写字节)在同一台机器上先标定 峰值算力与峰值带宽,再把真实算子打到这张 roofline 上——所以你能看到 「预测的下界」和「实测」差多少,而不只是相信规格表。 每个算子的 D 都按实现里实际的读写遍数来数: LayerNorm 预分配版 10 遍(mean / 减均值 / 平方 / 再 mean / 除 / 乘 gamma) attention 朴素预分配 10 遍(写 S / max / 减 / exp / sum / 除 / 读 S 写 O) 遍数不是拍脑袋来的,是照着下面每一行 numpy 数出来的。 输出的每个数字都来自 `time.perf_counter()` 的真实计时。 """ import json import platform import time from pathlib import Path import numpy as np FP32 = 4 # 本机跑 fp32,numpy 在 Apple 上走 Accelerate SNAPSHOT = Path(__file__).resolve().parent / "probe_results.json" def _best(fn, repeat: int): """跑 repeat 次取最快的一次:避开冷启动和调度抖动。""" fn() # 预热:第一次要分配 / 触发缺页 best = float("inf") for _ in range(repeat): t = time.perf_counter() fn() best = min(best, time.perf_counter() - t) return best # ── 一、标定可达带宽:二元 add 模式(读两份写一份)─────────────────── def measure_bandwidth(n: int = 32_000_000, repeat: int = 8) -> float: """返回实测可达带宽(Byte/s)。 每份 fp32 数组 128 MB;应核对目标机器缓存大小,这里测量此访问模式的有效带宽。 z = x + y:读 2n、写 n,共 3n 个元素。 """ x = np.ones(n, dtype=np.float32) y = np.ones(n, dtype=np.float32) z = np.empty(n, dtype=np.float32) dt = _best(lambda: np.add(x, y, out=z), repeat) return 3 * n * FP32 / dt # ── 二、标定峰值算力:大矩阵乘 ──────────────────────────────────── def measure_matmul_peak(n: int = 4096, repeat: int = 5) -> float: """返回实测可达算力(FLOP/s)。 4096 的方阵乘足够大,访存被摊薄,测出来的是算力上限。 FLOPs = 2·n³(口径同 roofline_model.py:一次乘加记 2 个浮点运算)。 """ rng = np.random.default_rng(0) a = rng.standard_normal((n, n), dtype=np.float32) b = rng.standard_normal((n, n), dtype=np.float32) c = np.empty((n, n), dtype=np.float32) dt = _best(lambda: np.matmul(a, b, out=c), repeat) return 2 * n ** 3 / dt def calibrate(verbose: bool = True): bw = measure_bandwidth() fl = measure_matmul_peak() peak = dict(peak_flops=fl, peak_bw=bw) if verbose: print("=" * 80) print(f"本机标定({platform.machine()},numpy {np.__version__},fp32)") print("=" * 80) print(f" 实测可达带宽 beta = {bw / 1e9:8.2f} GB/s") print(f" 实测可达算力 P = {fl / 1e9:8.2f} GFLOP/s") print(f" ridge point I* = {fl / bw:8.1f} FLOP/Byte") print(" 注:这是「用 numpy 能摸到」的上限,不是硬件标称值;") print(" 换 BLAS 后端、换 dtype、换线程数都会变。口径一致才有可比性。") return peak def build_ops(): """构造待实测的算子。每个元素 = (名字, F, D, 函数, repeat)。""" rng = np.random.default_rng(0) ops = [] # ── 逐元素:1 遍 vs 2 遍 vs「看起来融合了」 ── n = 32_000_000 a = rng.standard_normal(n).astype(np.float32) b = rng.standard_normal(n).astype(np.float32) o1 = np.empty(n, dtype=np.float32) tmp = np.empty(n, dtype=np.float32) def add_relu_two_pass(): # 5n 字节:tmp 读写各一次 np.add(a, b, out=tmp) np.maximum(tmp, 0, out=o1) def add_relu_one_line(): # 写法像融合了,numpy 照样分配中间数组 np.maximum(np.add(a, b), 0, out=o1) ops.append(("逐元素 add(1 遍)", 1.0 * n, FP32 * 3 * n, lambda: np.add(a, b, out=o1), 8)) ops.append(("add+relu 两遍(预分配)", 2.0 * n, FP32 * 5 * n, add_relu_two_pass, 8)) ops.append(("add+relu 写成一行(有中间数组)", 2.0 * n, FP32 * 5 * n, add_relu_one_line, 8)) # ── LayerNorm:预分配的 10 遍版 ── ln_n, ln_d = 8192, 4096 x = rng.standard_normal((ln_n, ln_d)).astype(np.float32) g = rng.standard_normal(ln_d).astype(np.float32) mu = np.empty((ln_n, 1), dtype=np.float32) var = np.empty((ln_n, 1), dtype=np.float32) xc = np.empty_like(x) xc2 = np.empty_like(x) def layernorm_buffered(): np.mean(x, axis=-1, keepdims=True, out=mu) # 读 x 1 遍 np.subtract(x, mu, out=xc) # 读 x 写 xc 3 遍 np.multiply(xc, xc, out=xc2) # 读 xc 写 xc2 5 遍 np.mean(xc2, axis=-1, keepdims=True, out=var) # 读 xc2 6 遍 np.sqrt(var + 1e-5, out=var) # 小量 np.divide(xc, var, out=xc) # 读 xc 写 xc 8 遍 np.multiply(xc, g, out=xc) # 读 xc 写 xc 10 遍 ops.append((f"LayerNorm [{ln_n},{ln_d}] 预分配", 6.0 * ln_n * ln_d, FP32 * 10 * ln_n * ln_d, layernorm_buffered, 5)) # ── 朴素注意力:S 落盘,10 遍 N² 访存 ── N, D = 2048, 128 q = rng.standard_normal((N, D)).astype(np.float32) k = rng.standard_normal((N, D)).astype(np.float32) v = rng.standard_normal((N, D)).astype(np.float32) kt = np.ascontiguousarray(k.T) S = np.empty((N, N), dtype=np.float32) O = np.empty((N, D), dtype=np.float32) rowmax = np.empty((N, 1), dtype=np.float32) rowsum = np.empty((N, 1), dtype=np.float32) def attention_naive(): # 本微基准省略 1/sqrt(D) 缩放,只测 matmul-softmax-matmul 的执行开销。 np.matmul(q, kt, out=S) # 写 S 1 遍 np.max(S, axis=-1, keepdims=True, out=rowmax) # 读 S 2 遍 np.subtract(S, rowmax, out=S) # 读写 S 4 遍 np.exp(S, out=S) # 读写 S 6 遍 np.sum(S, axis=-1, keepdims=True, out=rowsum) # 读 S 7 遍 np.divide(S, rowsum, out=S) # 读写 S 9 遍 np.matmul(S, v, out=O) # 读 S 写 O 10 遍 ops.append((f"attention 朴素 [N={N},D={D}]", 4.0 * N * N * D + 5.0 * N * N, FP32 * (10 * N * N + 4 * N * D), attention_naive, 5)) # ── GEMM 三个规模 ── for m in (512, 2048, 4096): ma = rng.standard_normal((m, 4096)).astype(np.float32) mb = rng.standard_normal((4096, 4096)).astype(np.float32) mc = np.empty((m, 4096), dtype=np.float32) ops.append((f"GEMM {m}x4096x4096", 2.0 * m * 4096 * 4096, FP32 * (m * 4096 + 4096 * 4096 + m * 4096), lambda ma=ma, mb=mb, mc=mc: np.matmul(ma, mb, out=mc), 5)) return ops def probe(peak: dict, verbose: bool = True): ops = build_ops() I_star = peak["peak_flops"] / peak["peak_bw"] if verbose: print("\n" + "=" * 80) print("三、真实算子打到这张 roofline 上:下界 vs 实测") print("=" * 80) hdr = (f"{'算子':<32s}{'I':>9s}{'受限':>8s}{'T下界':>10s}{'T实测':>10s}" f"{'达成率':>9s}{'实测算力':>14s}") print(hdr) print("-" * len(hdr)) rows = [] for name, F, D, fn, repeat in ops: dt = _best(fn, repeat) I = F / D bound = "算力" if I >= I_star else "带宽" t_pred = max(F / peak["peak_flops"], D / peak["peak_bw"]) eff = t_pred / dt rows.append(dict(name=name, I=I, bound=bound, t_pred=t_pred, t_real=dt, eff=eff, flops=F, bytes=D, gflops=F / dt / 1e9)) if verbose: print(f"{name:<32s}{I:>9.2f}{bound:>8s}{t_pred * 1e3:>8.2f}ms" f"{dt * 1e3:>8.2f}ms{eff * 100:>8.0f}%" f"{F / dt / 1e9:>11.1f} GF/s") if verbose: print("\n 达成率 = 下界 / 实测,回答的是「这个实现离 roofline 还有多远」。") print(" 低达成率可能来自指令、缓存、调度或同步;单凭此比值不能诊断额外访存。") return rows def demo_same_flops(): """同样 1 GFLOP 的算力,GEMM 和逐元素算子差多少时间。""" print("\n" + "=" * 80) print("四、同样的 FLOPs,两种算子差多少时间") print("=" * 80) rng = np.random.default_rng(3) target = 1.0e9 # GEMM:2·M·K·N = target,取 K=N=1024 M = int(target / (2 * 1024 * 1024)) ma = rng.standard_normal((M, 1024)).astype(np.float32) mb = rng.standard_normal((1024, 1024)).astype(np.float32) mc = np.empty((M, 1024), dtype=np.float32) t_gemm = _best(lambda: np.matmul(ma, mb, out=mc), 10) f_gemm = 2.0 * M * 1024 * 1024 # 逐元素:n 个元素各算 1 次,F = n n_elem = int(target) a = rng.standard_normal(n_elem).astype(np.float32) b = rng.standard_normal(n_elem).astype(np.float32) o = np.empty(n_elem, dtype=np.float32) t_elem = _best(lambda: np.add(a, b, out=o), 5) print(f" GEMM : F = {f_gemm / 1e9:.3f} GFLOP, " f"D = {FP32 * (M * 1024 + 1024 * 1024 + M * 1024) / 1e6:.1f} MB, " f"实测 {t_gemm * 1e3:.3f} ms({f_gemm / t_gemm / 1e9:.0f} GFLOP/s)") print(f" 逐元素 : F = {n_elem / 1e9:.3f} GFLOP, " f"D = {FP32 * 3 * n_elem / 1e6:.1f} MB, " f"实测 {t_elem * 1e3:.3f} ms({n_elem / t_elem / 1e9:.0f} GFLOP/s)") print(f" → FLOPs 基本相同,耗时差 {t_elem / t_gemm:.0f} 倍。") print(" 同样运算量不代表同样时间:这里逐元素算子的运算数并未") print(f" 减少,却因为 I 只有 {1 / (FP32 * 3):.2f} FLOP/Byte 而完全被带宽卡住。") def main(): peak = calibrate() rows = probe(peak) demo_same_flops() # 把这次标定存成快照:make_figures / memory_ledger 直接读它, # 保证「图上的数字」和「正文引用的数字」来自同一次运行。 SNAPSHOT.write_text(json.dumps( dict(peak=peak, rows=rows, note=f"{platform.machine()} / numpy {np.__version__} / fp32", timestamp=time.strftime("%Y-%m-%d %H:%M:%S")), ensure_ascii=False, indent=1)) print(f"\n标定快照已存到 {SNAPSHOT.name},make_figures / memory_ledger 会复用它。") print(f"这台机器的 roofline:P = {peak['peak_flops'] / 1e9:.0f} GFLOP/s, " f"beta = {peak['peak_bw'] / 1e9:.1f} GB/s, " f"I* = {peak['peak_flops'] / peak['peak_bw']:.1f} FLOP/Byte") if __name__ == "__main__": main() roofline_model.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """Roofline 模型的最小实现:把「算力账」和「带宽账」合成一条上界曲线。 只依赖 numpy(其实只是用来排版,算法本身是纯标量运算)。 直接 `python roofline_model.py` 即可运行,约 1 秒。 符号与文章第 03 节一一对应 -------------------------- F 浮点运算数(FLOP)。口径:一次乘加(MAC)记 2 FLOP D 访存字节数(Byte)。口径:读 + 写,缓存命中不重复计 I 算术强度 I = F / D,单位 FLOP/Byte P_peak 峰值算力(FLOP/s) beta 峰值带宽(Byte/s) I_star ridge point = P_peak / beta,I 低于它就是带宽受限 关键结论(代码算完会打印): T >= max(F / P_peak, D / beta) 把 F 降为一半或把 D 降为一半时的模型比值,见 optimize_headroom() """ import json from pathlib import Path def _local_spec(): """本机实测峰值:优先读 machine_probe.py 存下的标定快照。 快照不存在时退回兜底常数(先跑一次 machine_probe.py 更准)。 这样 SPECS 里的「本机实测」永远是同一次标定,不会和正文引用的数字分叉。 """ snap = Path(__file__).resolve().parent / "probe_results.json" if snap.exists(): p = json.loads(snap.read_text())["peak"] return dict(peak_flops=p["peak_flops"], peak_bw=p["peak_bw"]) return dict(peak_flops=1680e9, peak_bw=77e9) # ── 厂商标称峰值。来源:NVIDIA A100 / H100 SXM 白皮书,bf16 dense(不含稀疏) # 注意 H100 BF16 dense 算力增幅大于带宽增幅,ridge point 比 A100 高: # 带宽涨 2.15 倍、bf16 算力涨约 3.17 倍 → 更多 kernel 落在带宽受限区。 SPECS = { "A100-40GB SXM (bf16)": dict(peak_flops=312e12, peak_bw=1555e9), "H100-80GB SXM (bf16)": dict(peak_flops=989e12, peak_bw=3350e9), "H100-80GB SXM (fp8)": dict(peak_flops=1979e12, peak_bw=3350e9), # 下面这一条不是标称值,是 machine_probe.py 在同一台机器上实测出来的 "本机实测(见 machine_probe.py)": _local_spec(), } BF16 = 2 # bytes per element def roofline_time(F: float, D: float, peak_flops: float, peak_bw: float): """roofline 时间下界:算力时间与带宽时间取 max(假设两者完美重叠)。""" t_compute = F / peak_flops t_memory = D / peak_bw return max(t_compute, t_memory), t_compute, t_memory def bound_of(F: float, D: float, peak_flops: float, peak_bw: float) -> str: I = F / D I_star = peak_flops / peak_bw return "算力受限" if I >= I_star else "带宽受限" def optimize_headroom(F: float, D: float, peak_flops: float, peak_bw: float): """把运算量 F 或搬运量 D 减半,理想模型可快多少? 答案与算术强度有关:带宽受限时砍算力收益为 0,算力受限时最多 2 倍。 返回 (F 减半的加速比, D 减半的加速比),分别等价于只将峰值算力或带宽翻倍。 """ def speedup(F2, D2): t0, _, _ = roofline_time(F, D, peak_flops, peak_bw) t1, _, _ = roofline_time(F2, D2, peak_flops, peak_bw) return t0 / t1 return speedup(F / 2, D), speedup(F, D / 2) # ── 几类算子的 F / D 账本 ────────────────────────────────────────── def gemm(M: int, K: int, N: int, b: int = BF16): """矩阵乘 [M,K] @ [K,N] -> [M,N]。读 A、读 B、写 C。""" F = 2.0 * M * K * N D = b * (M * K + K * N + M * N) return F, D def elementwise(n: int, b: int = BF16, n_pass: int = 1, n_read: int = 1): """逐元素算子:读 n_read 份、写 1 份;n_pass 表示这样读写几轮。 二元运算(如 add)是 n_read=2:读两份输入、写一份输出。 """ F = 1.0 * n * n_pass # 每个元素算 1 次 D = b * (n_read + 1) * n * n_pass return F, D def layernorm(n: int, d: int, b: int = BF16): """简化 LayerNorm(无 beta):两次归约、减均值、平方、除标准差、乘 gamma,约 6 次/元素。 访存 = 读 x 一遍 + 写 y 一遍(统计量是 O(n) 的小量,忽略)。""" F = 6.0 * n * d D = b * 2 * n * d return F, D def softmax(n: int, d: int, b: int = BF16, fused: bool = True): """softmax:exp / 减最大值 / 归一化,约 5 次运算/元素。 fused=True 假设整行可驻留片上存储,读 1 遍写 1 遍;不是独立 online normalizer 的通用 IO fused=False 朴素三遍(求 max、求 exp 和、归一化),读 3 遍写 2 遍 """ F = 6.0 * n * d D = b * 2 * n * d if fused else b * (3 * n * d + 2 * n * d) return F, D def attention(N: int, D: int, b: int = BF16, flash: bool = False): """单头注意力。N 是 token 数,D 是每头维度。 flash=False(朴素):分数矩阵 S 和权重 P 都要落 HBM S = Q K^T 写 N^2 P = softmax(S) 读 N^2 写 N^2 O = P V 读 N^2 → 二次项访存 ≈ 4 N^2 b flash=True(强缓存假设下的最低流量,不是真实 FlashAttention IO):Q/K/V 读进来、O 写出去,N^2 项只留在 SRAM → 访存 ≈ 4 N D b """ F = 4.0 * N * N * D + 5.0 * N * N # 两个 N×N×D 的矩阵乘 + softmax if flash: D_bytes = b * (3 * N * D + N * D) # 读 Q K V,写 O else: D_bytes = b * (4 * N * N + 4 * N * D) # 上面再加 S/P 的读写 return F, D_bytes def fmt(x: float, unit: str = "") -> str: for u, s in (("T", 1e12), ("G", 1e9), ("M", 1e6), ("K", 1e3)): if x >= s: return f"{x / s:.2f} {u}{unit}" return f"{x:.2f} {unit}" def main(): print("=" * 78) print("一、四台「机器」的峰值与 ridge point") print("=" * 78) print(f"{'设备':<28s}{'P_peak':>14s}{'beta':>14s}{'I* = P/beta':>16s}") ridges = {} for name, s in SPECS.items(): I_star = s["peak_flops"] / s["peak_bw"] ridges[name] = (s, I_star) print(f"{name:<28s}{fmt(s['peak_flops'], 'FLOP/s'):>16s}" f"{fmt(s['peak_bw'], 'B/s'):>14s}{I_star:>12.1f} FLOP/B") a100 = SPECS["A100-40GB SXM (bf16)"] h100 = SPECS["H100-80GB SXM (bf16)"] print(f"\n A100 -> H100:带宽 ×{h100['peak_bw'] / a100['peak_bw']:.2f}," f"算力 ×{h100['peak_flops'] / a100['peak_flops']:.2f}," f"ridge point {a100['peak_flops'] / a100['peak_bw']:.0f} -> " f"{h100['peak_flops'] / h100['peak_bw']:.0f}(更高 = 更多算子落入带宽受限区)") # ── 二、算子账本 ──────────────────────────────────────────────── N, D = 4096, 128 # 一个典型的 DiT / LLM 单头规模 cases = [ ("逐元素 add [16M]", elementwise(16_000_000, n_read=2)), ("LayerNorm [4096, 3072]", layernorm(4096, 3072)), ("softmax fused [4096, 4096]", softmax(4096, 4096, fused=True)), ("softmax 朴素三遍 [4096, 4096]", softmax(4096, 4096, fused=False)), ("GEMM 512x4096x4096", gemm(512, 4096, 4096)), ("GEMM 4096x4096x4096", gemm(4096, 4096, 4096)), ("attention 朴素 [N=4096,D=128]", attention(N, D, flash=False)), ("attention ideal-min [N=4096,D=128]", attention(N, D, flash=True)), ] print("\n" + "=" * 78) print("二、算子账本(bf16,A100 口径)") print("=" * 78) hdr = (f"{'算子':<32s}{'F (FLOP)':>14s}{'D (Byte)':>14s}" f"{'I':>10s}{'受限':>10s}{'T_roof':>12s}") print(hdr) print("-" * len(hdr)) for name, (F, Db) in cases: I = F / Db T, tc, tm = roofline_time(F, Db, a100["peak_flops"], a100["peak_bw"]) print(f"{name:<32s}{fmt(F):>14s}{fmt(Db, 'B'):>14s}" f"{I:>10.2f}{bound_of(F, Db, a100['peak_flops'], a100['peak_bw']):>10s}" f"{T * 1e3:>10.3f} ms") # ── 三、FlashAttention 把算术强度抬了多少 ──────────────────────── print("\n" + "=" * 78) print("三、物化 S/P 与理想最低 IO 对比(单头 N=4096, D=128, bf16)") print("=" * 78) F_naive, D_naive = attention(N, D, flash=False) F_flash, D_flash = attention(N, D, flash=True) T_naive, _, _ = roofline_time(F_naive, D_naive, a100["peak_flops"], a100["peak_bw"]) T_flash, _, _ = roofline_time(F_flash, D_flash, a100["peak_flops"], a100["peak_bw"]) print(f" FLOPs : {fmt(F_naive)} -> {fmt(F_flash)} " f"(×{F_flash / F_naive:.3f},几乎没变)") print(f" 访存 : {fmt(D_naive, 'B')} -> {fmt(D_flash, 'B')} " f"(×{D_flash / D_naive:.4f},省 {D_naive / D_flash:.1f} 倍)") print(f" 算术强度 : {F_naive / D_naive:.1f} -> {F_flash / D_flash:.1f} FLOP/B " f"(×{(F_flash / D_flash) / (F_naive / D_naive):.1f})") print(f" 受限类型 : {bound_of(F_naive, D_naive, **a100)} -> " f"{bound_of(F_flash, D_flash, **a100)}") print(f" roofline : {T_naive * 1e3:.3f} ms -> {T_flash * 1e3:.3f} ms " f"(×{T_naive / T_flash:.2f})") print(f" 注意:加速来自「少搬 {D_naive / D_flash:.0f} 倍字节」," "不是「少算」——此处仅比较相同主导 FLOPs 与理想最低访存,不是实测内核。") # ── 四、砍算力 vs 砍带宽,谁的收益大 ──────────────────────────── print("\n" + "=" * 78) print("四、同一个 kernel,只翻倍峰值算力或带宽的理想加速") print("=" * 78) print(f"{'算子':<32s}{'I/I*':>10s}{'算力×2 的加速':>16s}{'带宽×2 的加速':>16s}") I_star = a100["peak_flops"] / a100["peak_bw"] for name, (F, Db) in cases: s_flops, s_bw = optimize_headroom(F, Db, a100["peak_flops"], a100["peak_bw"]) print(f"{name:<32s}{(F / Db) / I_star:>10.3f}" f"{s_flops:>14.2f}x{s_bw:>14.2f}x") print("\n 读法:带宽受限的行(I/I* < 1)里「算力×2」那一列几乎都是 1.00——" "\n 在只改变计算峰值且访存不变的理想模型里,该列为 1。") # ── 五、GEMM 什么时候从带宽受限翻到算力受限 ───────────────────── print("\n" + "=" * 78) print("五、GEMM [M,4096] x [4096,4096]:M 多大才开始算力受限") print("=" * 78) K = Nout = 4096 print(f"{'M':>8s}{'I (FLOP/B)':>14s}{'受限':>12s}") prev = None for M in (1, 2, 4, 8, 16, 32, 64, 128, 256, 512, 1024, 4096): F, Db = gemm(M, K, Nout) b = bound_of(F, Db, a100["peak_flops"], a100["peak_bw"]) print(f"{M:>8d}{F / Db:>14.2f}{b:>12s}") if prev == "带宽受限" and b == "算力受限": print(f" ↑ 拐点:M 在 ({M // 2}, {M}] 之间," f"I 越过 ridge point {I_star:.0f}") prev = b print("\n 推论:小 batch 推理(M=1~8)里的 GEMM 是带宽受限的," "这时候做 fp8 量化省的是字节、不是算力。") if __name__ == "__main__": main() memory_ledger.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """显存账本 + batch 扫描:容量这本账怎么算,以及它怎么反过来决定速度。 只依赖 numpy(第二部分要真跑计时)。直接 `python memory_ledger.py` 即可运行,约 10 秒。 第一部分是「容量账」:权重 / KV cache / 激活 / 额外开销四项加总,看 假设 78 GiB 可用预算 能塞多少条并发请求。第二部分是「速度账」:batch 变大之后,同一份权重被更多 token 摊薄,算术强度抬上去,kernel 从带宽受限翻到算力受限——这个拐点用真实 计时测出来,顺便告诉你 batch 该开多大。 """ import json import time from pathlib import Path import numpy as np GIB = 1024 ** 3 def _load_peak(): """优先读 machine_probe.py 存下的标定快照(保证口径一致); 没跑过 machine_probe.py 时退回兜底常数,并照常工作。""" snap = Path(__file__).resolve().parent / "probe_results.json" if snap.exists(): peak = json.loads(snap.read_text())["peak"] return peak["peak_flops"], peak["peak_bw"] return 1680.64e9, 77.21e9 # 兜底:先跑一次 machine_probe.py 更准 P_PEAK, BETA = _load_peak() I_STAR = P_PEAK / BETA def kv_cache_bytes(n_tokens: int, n_layers: int, d_kv: int, b: int = 2) -> int: """KV cache 字节数:每个 token 每层都要存一份 K 和一份 V。 n_tokens 序列长度(或并发请求的总 token 数) n_layers 层数 L d_kv 每层的 KV 总维度 = n_kv_heads × d_head(GQA 时用实际的 KV head 数) b 每个元素的字节数(fp16/bf16 = 2,fp8 = 1) """ return 2 * n_tokens * n_layers * d_kv * b def ledger(n_params: float, b_w: int, n_layers: int, d_kv: int, batch: int, seq: int, act_per_token: int, frag_rate: float): """推理显存账本的四项,单位 GiB。 act_per_token 每个 token 的激活峰值(字节)。推理不保留整层的中间结果, 但当前层的十几份临时张量要同时活着:d_model=4096、fp16 时 一份 [d] 张量 8 KiB,取 16 份 ≈ 128 KiB / token。 这是经验值,随实现(融合程度、是否分块)差距很大。 frag_rate 教学额外开销 / (权重 + KV + 激活),不是实测碎片率。 本函数按未分块 prefill 计 batch*seq 个活跃 token; 逐 token decode 的激活应改为 batch*act_per_token,KV 仍取 batch*seq。 """ w = n_params * b_w kv = kv_cache_bytes(batch * seq, n_layers, d_kv) act = batch * seq * act_per_token frag = (w + kv + act) * frag_rate return dict(weights=w / GIB, kv=kv / GIB, act=act / GIB, frag=frag / GIB, total=(w + kv + act + frag) / GIB) def part1(): print("=" * 80) print("一、显存账本:一张 假设 78 GiB 可用预算,能塞多少条并发请求") print("=" * 80) # Llama-2-7B 的配置:32 层、32 个 KV head、d_head=128 → d_kv = 4096 N_PARAMS, N_LAYERS, D_KV = 7e9, 32, 4096 SEQ = 4096 per_token = kv_cache_bytes(1, N_LAYERS, D_KV) print(f" 模型:7B,fp16 权重 = {N_PARAMS * 2 / GIB:.1f} GiB," f"{N_LAYERS} 层,d_kv = {D_KV}") print(f" KV cache 单价:{per_token / 1024:.0f} KiB / token " f"(一条 {SEQ} 长的序列 = {per_token * SEQ / GIB:.2f} GiB)") print() hdr = (f"{'batch':>6s}{'权重':>9s}{'KV cache':>10s}{'激活':>9s}" f"{'额外开销':>9s}{'合计':>9s}{'78GiB预算':>10s}") print(hdr) print("-" * len(hdr)) for batch in (1, 4, 8, 16, 24, 32): # 朴素预分配:按最大长度全预留,教学假设:以权重+KV+激活总和的 30% 估计额外开销,不是实测碎片比例 naive = ledger(N_PARAMS, 2, N_LAYERS, D_KV, batch, SEQ, act_per_token=131072, frag_rate=0.30) ok = "装得下" if naive["total"] < 78 else "OOM" print(f"{batch:>6d}{naive['weights']:>8.1f}G{naive['kv']:>9.2f}G" f"{naive['act']:>8.2f}G{naive['frag']:>8.2f}G" f"{naive['total']:>8.2f}G{ok:>10s}") # 换成低开销假设(教学假设,非 PagedAttention 实测):额外开销率降到 4% print("\n 同一张卡,把额外开销率从 30% 降到 4%(低开销假设):") print(f"{'batch':>6s}{'朴素合计':>11s}{'低开销假设':>11s}{'多出来的并发':>16s}") for batch in (8, 16, 24, 32): naive = ledger(N_PARAMS, 2, N_LAYERS, D_KV, batch, SEQ, 131072, 0.30) paged = ledger(N_PARAMS, 2, N_LAYERS, D_KV, batch, SEQ, 131072, 0.04) extra = "" if naive["total"] > 78 >= paged["total"]: extra = "从 OOM 变装得下" print(f"{batch:>6d}{naive['total']:>10.2f}G{paged['total']:>10.2f}G" f"{extra:>16s}") print("\n 读法:权重那一项是常数,batch 再大也不变;KV cache 才是随并发") print(" 线性增长的那一项,batch 16 时它已经比权重还大。") print(" 「显存够放权重就能跑」这句话漏掉了后面三项。") # 用 roofline_model 里同款口径:下界 = max(F/P, D/beta)。 # P_PEAK / BETA / I_STAR 已在文件顶部从标定快照读出,这里不再覆盖, # 否则快照刷新后正文引用的数字会和实际跑出来的对不上。 def batch_sweep(repeats: int = 15): """固定一份「权重」,扫 batch,返回每一档的实测延迟/吞吐/算术强度。""" rng = np.random.default_rng(0) d = 4096 W = rng.standard_normal((d, d)).astype(np.float32) # 64 MiB 的「权重」 W_bytes = W.nbytes FP32 = 4 def run(B): x = rng.standard_normal((B, d)).astype(np.float32) y = np.empty((B, d), dtype=np.float32) np.matmul(x, W, out=y) # 预热 best = float("inf") for _ in range(repeats): t = time.perf_counter() np.matmul(x, W, out=y) best = min(best, time.perf_counter() - t) return best rows = [] for B in (1, 2, 4, 8, 16, 32, 64, 128, 256): F = 2.0 * B * d * d D = W_bytes + 2 * B * d * FP32 I = F / D dt = run(B) rows.append(dict(B=B, F=F, D=D, I=I, lat=dt, thr=B / dt, bound="算力" if I >= I_STAR else "带宽", t_pred=max(F / P_PEAK, D / BETA), eff=max(F / P_PEAK, D / BETA) / dt)) return rows, W_bytes def part2(): """batch 扫描实测:延迟、吞吐、算术强度,以及拐点在哪。""" print("\n" + "=" * 80) print("二、batch 扫描实测:吞吐什么时候不再涨") print("=" * 80) rows, W_bytes = batch_sweep() (Path(__file__).resolve().parent / "batch_results.json").write_text( json.dumps(dict(rows=rows, weight_bytes=W_bytes, peak=dict(peak_flops=P_PEAK, peak_bw=BETA)), indent=2)) print(f" 固定「权重」W 形状 [4096, 4096],fp32 = {W_bytes / 2 ** 20:.0f} MiB;" f"batch B 就是一次喂进去的 token 数") print(f" 峰值口径(来自 machine_probe 的标定快照):" f"P = {P_PEAK / 1e9:.0f} GFLOP/s,beta = {BETA / 1e9:.1f} GB/s," f"I* = {I_STAR:.1f}") print() hdr = (f"{'B':>6s}{'F (GFLOP)':>12s}{'D (MB)':>10s}{'I':>9s}{'受限':>8s}" f"{'延迟 ms':>10s}{'吞吐 K/s':>11s}{'下界 ms':>10s}{'达成率':>9s}") print(hdr) print("-" * len(hdr)) prev_bound = None for r in rows: print(f"{r['B']:>6d}{r['F'] / 1e9:>12.3f}{r['D'] / 1e6:>10.2f}" f"{r['I']:>9.2f}{r['bound']:>8s}{r['lat'] * 1e3:>10.3f}" f"{r['thr'] / 1e3:>9.1f}K{r['t_pred'] * 1e3:>10.3f}" f"{r['eff'] * 100:>8.0f}%") if prev_bound == "带宽" and r["bound"] == "算力": print(f" ↑ 拐点:B 从 {r['B'] // 2} 到 {r['B']} 之间" f"越过 ridge point {I_STAR:.1f}") prev_bound = r["bound"] # 反常检测:batch 变大反而变慢,说明库换了代码路径(roofline 看不见这件事) for r0, r1 in zip(rows, rows[1:]): if r1["B"] == 2 * r0["B"] and r1["lat"] > r0["lat"] * 1.5: print(f" ! 反常:B={r0['B']} 只要 {r0['lat'] * 1e3:.2f} ms," f"B={r1['B']} 却要 {r1['lat'] * 1e3:.2f} ms") base = rows[0] top = rows[-1] print(f"\n 延迟:B=1 时 {base['lat'] * 1e3:.3f} ms,B={top['B']} 时 " f"{top['lat'] * 1e3:.3f} ms(涨 {top['lat'] / base['lat']:.1f} 倍)") print(f" 吞吐:B=1 时 {base['thr'] / 1e3:.1f} K/s,B={top['B']} 时 " f"{top['thr'] / 1e3:.1f} K/s(涨 {top['thr'] / base['thr']:.1f} 倍)") print(f" 单样本成本:B={top['B']} 时把 64 MiB 权重的读取摊薄到 {top['B']} 个样本," f"每个样本只摊 {W_bytes / top['B'] / 2 ** 20:.2f} MiB") print(" → 增大 batch 可摊薄权重访问;延迟是否平坦、吞吐能增加多少仍需实测。") print(" → 上面标 ! 的行是 roofline 看不见的东西:同样的公式、更大的 batch," "可能涉及内核选择、调度或计时波动,单凭本表不能确诊。") if __name__ == "__main__": part1() part2() make_figures.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """生成「性能建模与 Profiling」一文的四张解释图。 图里的计时优先读取已保存的实测快照,避免重绘时图文使用不同测量: * roofline.png 峰值与落点来自 machine_probe.calibrate() / probe() * batch_scaling.png 来自 memory_ledger.batch_sweep() * memory_ledger.png 来自 memory_ledger.ledger() * opt_gain.png 来自 roofline_model.optimize_headroom() 只依赖 numpy + matplotlib,直接 `python make_figures.py` 即可运行,约 40 秒。 改了另外三个脚本,这里要重跑,避免图与正文数字对不上。 """ import json import sys from pathlib import Path import matplotlib.pyplot as plt import numpy as np HERE = Path(__file__).resolve().parent sys.path.insert(0, str(HERE)) import machine_probe as MP # noqa: E402 import memory_ledger as ML # noqa: E402 import roofline_model as RM # noqa: E402 OUT = HERE.parent / "figures" try: # exist_ok=True 是必须的:目录已存在时 pathlib 会抛 FileExistsError, # 少数沙箱环境连 exist_ok=True 的 mkdir 也一并拦,这里再兜一层。 OUT.mkdir(exist_ok=True) except PermissionError: if not OUT.is_dir(): raise plt.rcParams.update({ "font.sans-serif": ["PingFang SC", "Arial Unicode MS", "DejaVu Sans"], "axes.unicode_minus": False, "figure.dpi": 160, "savefig.dpi": 160, }) C_MEM = "#e05263" # 带宽受限:红 C_CMP = "#3f7fbf" # 算力受限:蓝 C_ACC = "#2f9e6f" # 好结果:绿 C_GREY = "#94a3b8" # ── 图 1:roofline 曲线 + 实测落点 ──────────────────────────────── def fig_roofline(peak, rows): P = peak["peak_flops"] / 1e9 # GFLOP/s beta = peak["peak_bw"] / 1e9 # GB/s I_star = P / beta # 因为两个都除过 1e9,比值不变 I = np.logspace(-2, 4, 400) perf = np.minimum(P, beta * I) fig, ax = plt.subplots(figsize=(8.4, 5.4)) ax.loglog(I, perf, color="k", lw=2.2) ax.axvline(I_star, color=C_GREY, ls="--", lw=1.4) ax.text(I_star * 1.15, 3, f"ridge point I* = {I_star:.0f}\nFLOP/Byte", color="#475569", fontsize=9, va="bottom") ax.text(0.012, 25, f"峰值算力 P = {P:.0f} GFLOP/s", fontsize=9, color="#475569") ax.text(0.012, 8.5, "斜线 = 带宽天花板,beta = " f"{beta:.0f} GB/s", fontsize=9, color="#475569") # 落点:用实测时间反推实际性能 short = { "逐元素 add(1 遍)": "逐元素 add", "add+relu 两遍(预分配)": "add+relu 两遍", "add+relu 写成一行(有中间数组)": "add+relu 一行写法", "LayerNorm [8192,4096] 预分配": "LayerNorm", "attention 朴素 [N=2048,D=128]": "朴素 attention", "GEMM 512x4096x4096": "GEMM M=512", "GEMM 2048x4096x4096": "GEMM M=2048", "GEMM 4096x4096x4096": "GEMM M=4096", } offsets = { # 手工调过:右边缘和重叠的标签让位 "GEMM M=4096": (-10, 10), "GEMM M=2048": (10, -4), "GEMM M=512": (12, -14), "朴素 attention": (7, -3), } # 左下角四个点挤在一起,用箭头把标签拉到右边的空地上 arrows = { "逐元素 add": (0.42, 9.5), "LayerNorm": (0.42, 5.6), "add+relu 两遍": (0.42, 3.1), "add+relu 一行写法": (0.42, 1.75), } for r in rows: name = short.get(r["name"], r["name"]) y = r["gflops"] ax.scatter([r["I"]], [y], s=52, color=C_MEM if r["bound"] == "带宽" else C_CMP, zorder=5, edgecolor="white", linewidth=0.8) if name in arrows: tx, ty = arrows[name] ax.annotate(name, (r["I"], y), xytext=(tx, ty), fontsize=8.5, color="#334155", arrowprops=dict(arrowstyle="-", color=C_GREY, lw=0.9), va="center") continue dx, dy = offsets.get(name, (7, -3)) ha = "right" if dx < 0 else "left" ax.annotate(name, (r["I"], y), textcoords="offset points", xytext=(dx, dy), fontsize=8.5, color="#334155", ha=ha) ax.set_xlabel("算术强度 I = F / D (FLOP/Byte)") ax.set_ylabel("实测性能 (GFLOP/s)") ax.set_title("图 1:本机的 roofline —— 斜线是带宽,平顶是算力,点是实测落点", fontsize=11) ax.set_xlim(0.01, 5000) ax.set_ylim(0.3, 5000) ax.grid(alpha=0.25, which="both", ls=":") handles = [plt.Line2D([], [], marker="o", ls="", color=C_MEM, label="模型分类:I < I*"), plt.Line2D([], [], marker="o", ls="", color=C_CMP, label="模型分类:I >= I*")] ax.legend(handles=handles, loc="lower right", fontsize=8.5, framealpha=0.9) fig.tight_layout() fig.savefig(OUT / "roofline.png") plt.close(fig) # ── 图 3:batch 扫描的延迟与吞吐 ────────────────────────────────── def fig_batch(rows): B = [r["B"] for r in rows] lat = [r["lat"] * 1e3 for r in rows] thr = [r["thr"] / 1e3 for r in rows] fig, ax1 = plt.subplots(figsize=(8.4, 5.0)) ax1.plot(B, lat, "o-", color=C_CMP, lw=2, label="延迟(左轴)") ax1.set_xscale("log", base=2) ax1.set_xlabel("batch B(一次喂进去的 token 数)") ax1.set_ylabel("延迟 (ms)", color=C_CMP) ax1.tick_params(axis="y", labelcolor=C_CMP) ax2 = ax1.twinx() ax2.plot(B, thr, "s--", color=C_ACC, lw=2, label="吞吐(右轴)") ax2.set_ylabel("吞吐 (K token/s)", color=C_ACC) ax2.tick_params(axis="y", labelcolor=C_ACC) # 拐点:I 越过 ridge point 的地方 cross = next((r for r in rows if r["bound"] == "算力"), None) if cross: ax1.axvline(cross["B"], color=C_GREY, ls=":", lw=1.5) ax1.text(cross["B"] * 1.05, max(lat) * 0.92, f"越过 ridge point\nB ≈ {cross['B']}", fontsize=9, color="#475569") # 库的反常:延迟不随 batch 单调 b8 = next((r["lat"] for r in rows if r["B"] == 8), None) spike = next((r for r in rows if r["B"] in (2, 4) and b8 is not None and r["lat"] > 1.5 * b8), None) if spike: ax1.annotate(f"B={spike['B']} 延迟为 B=8 的 {spike['lat']/b8:.1f} 倍\n(原因需 profile)", (spike["B"], spike["lat"] * 1e3), textcoords="offset points", xytext=(14, -6), fontsize=8.5, color=C_MEM) ax1.set_title("图 3:batch 扫描 —— 实测延迟、吞吐与 roofline 模型交界", fontsize=11) ax1.grid(alpha=0.25, ls=":") lines1, lab1 = ax1.get_legend_handles_labels() lines2, lab2 = ax2.get_legend_handles_labels() ax1.legend(lines1 + lines2, lab1 + lab2, loc="upper left", fontsize=9) fig.tight_layout() fig.savefig(OUT / "batch_scaling.png") plt.close(fig) # ── 图 4:显存账本 ──────────────────────────────────────────────── def fig_ledger(): batches = [1, 4, 8, 16, 24, 32] naive = [ML.ledger(7e9, 2, 32, 4096, b, 4096, 131072, 0.30) for b in batches] paged = [ML.ledger(7e9, 2, 32, 4096, b, 4096, 131072, 0.04) for b in batches] fig, (ax, ax2) = plt.subplots(1, 2, figsize=(11.2, 4.8), gridspec_kw={"width_ratios": [1.35, 1]}) x = np.arange(len(batches)) keys = [("weights", "权重", "#3f7fbf"), ("kv", "KV cache", "#e05263"), ("act", "激活", "#e8a33d"), ("frag", "额外开销", "#94a3b8")] bottom = np.zeros(len(batches)) for k, label, color in keys: vals = np.array([d[k] for d in naive]) ax.bar(x, vals, bottom=bottom, label=label, color=color, width=0.62, edgecolor="white", linewidth=0.6) bottom += vals ax.axhline(78, color="k", ls="--", lw=1.3) ax.text(len(batches) - 0.4, 79, "教学假设:78 GiB 可用预算", fontsize=8.5, ha="right", color="#334155") for i, d in enumerate(naive): if d["total"] > 78: ax.text(i, d["total"] + 2, "OOM", ha="center", fontsize=9, color=C_MEM, fontweight="bold") ax.set_xticks(x) ax.set_xticklabels(batches) ax.set_xlabel("并发序列数 batch") ax.set_ylabel("显存 (GiB)") ax.set_title("图 4a:未分块 prefill 教学账本(额外开销率 30%)", fontsize=10.5) ax.legend(fontsize=8.5, ncol=2) ax.grid(axis="y", alpha=0.25, ls=":") w_naive = [d["total"] for d in naive] w_paged = [d["total"] for d in paged] ax2.bar(x - 0.2, w_naive, width=0.4, label="朴素", color="#e05263") ax2.bar(x + 0.2, w_paged, width=0.4, label="低开销假设", color="#2f9e6f") ax2.axhline(78, color="k", ls="--", lw=1.3) ax2.set_xticks(x) ax2.set_xticklabels(batches) ax2.set_xlabel("并发序列数 batch") ax2.set_ylabel("合计显存 (GiB)") ax2.set_title("图 4b:只把额外开销率从 30% 降到 4%", fontsize=10.5) ax2.legend(fontsize=8.5) ax2.grid(axis="y", alpha=0.25, ls=":") fig.suptitle("图 4:显存账本的四项 —— 权重是常数,KV cache 随并发线性增长", fontsize=11) fig.tight_layout() fig.savefig(OUT / "memory_ledger.png") plt.close(fig) # ── 图 2:砍算力 vs 砍带宽,各有多少钱 ──────────────────────────── def fig_opt_gain(): a100 = RM.SPECS["A100-40GB SXM (bf16)"] ratio = np.logspace(-2, 1.2, 300) sp_flops, sp_bw = [], [] for r in ratio: # 造一个算术强度恰好是 r × I* 的算子:固定 D,F 由 r 决定 D = 1e8 F = r * (a100["peak_flops"] / a100["peak_bw"]) * D s_f, s_b = RM.optimize_headroom(F, D, a100["peak_flops"], a100["peak_bw"]) sp_flops.append(s_f) sp_bw.append(s_b) fig, ax = plt.subplots(figsize=(8.4, 4.8)) ax.semilogx(ratio, sp_flops, color=C_CMP, lw=2.2, label="算力翻 2 倍能拿到的加速") ax.semilogx(ratio, sp_bw, color=C_MEM, lw=2.2, label="带宽翻 2 倍能拿到的加速") ax.axvline(1.0, color=C_GREY, ls="--", lw=1.4) ax.text(1.05, 1.02, "I = I*:分界线", fontsize=9, color="#475569") ax.fill_between(ratio, 0.98, 2.02, where=np.array(ratio) < 1, color=C_MEM, alpha=0.07) ax.fill_between(ratio, 0.98, 2.02, where=np.array(ratio) >= 1, color=C_CMP, alpha=0.07) ax.text(0.05, 1.9, "带宽受限区:\n换更快的算力单元 = 0 收益", fontsize=9, color=C_MEM) ax.text(6, 1.9, "算力受限区:\n加带宽 = 0 收益", fontsize=9, color=C_CMP) ax.set_xlabel("算术强度 / ridge point (I / I*)") ax.set_ylabel("能拿到的加速比") ax.set_ylim(0.95, 2.1) ax.set_title("图 2:投资之前先看这张图 —— 你的 kernel 在分界线哪一侧", fontsize=11) ax.grid(alpha=0.25, ls=":") ax.legend(fontsize=9, loc="center right") fig.tight_layout() fig.savefig(OUT / "opt_gain.png") plt.close(fig) def main(): # --only <图名>:只重画指定一张(roofline / batch / ledger / gain), # 避免为了改一张图的标签把所有实测算子重跑一遍、数字跟正文引用分叉。 only = None if "--only" in sys.argv: only = sys.argv[sys.argv.index("--only") + 1] # 优先读 machine_probe.py 存的标定快照,保证图上数字与正文引用同源 if MP.SNAPSHOT.exists(): snap = json.loads(MP.SNAPSHOT.read_text()) peak = snap["peak"] print(f"读取标定快照({snap['timestamp']},{snap['note']}):") else: print("没找到标定快照,现场标定一遍…") peak = MP.calibrate(verbose=False) MP.SNAPSHOT.write_text(json.dumps( dict(peak=peak, rows=MP.probe(peak, verbose=False), note="fresh", timestamp=""))) print(f" P = {peak['peak_flops'] / 1e9:.0f} GFLOP/s, " f"beta = {peak['peak_bw'] / 1e9:.1f} GB/s, " f"I* = {peak['peak_flops'] / peak['peak_bw']:.1f}") if only in (None, "roofline"): rows = snap.get("rows") if MP.SNAPSHOT.exists() and "snap" in locals() else None if not rows: rows = MP.probe(peak, verbose=False) fig_roofline(peak, rows) if only in (None, "batch"): batch_snap = HERE / "batch_results.json" if batch_snap.exists(): sweeps = json.loads(batch_snap.read_text())["rows"] else: sweeps, weight_bytes = ML.batch_sweep() batch_snap.write_text(json.dumps(dict(rows=sweeps, weight_bytes=weight_bytes, peak=peak), indent=2)) fig_batch(sweeps) if only in (None, "ledger"): fig_ledger() if only in (None, "gain"): fig_opt_gain() if only and only not in ("roofline", "batch", "ledger", "gain"): raise ValueError("unknown figure: " + only) print(f"\n图已写入 {OUT}") if __name__ == "__main__": main()
2026年09月25日
4 阅读
0 评论
0 点赞
2026-09-24
AIGC 基本功|自注意力机制的计算与显存账本-MHA
自注意力机制的计算与显存账本 所属方向:注意力与位置编码 | 难度:入门 | 前置知识:无(这是知识树注意力方向的根节点) 关键词:自注意力、multi-head attention、QKV、复杂度、显存账本、缩放因子 01. 为什么需要它 先看一个真实会撞上的场景:你拿到一个 约 3.62B 参数的简化 Transformer(d_model=3072、32 层),想在 1024×1024 的图生视频任务上做训练,并保留各层激活供反向使用。VAE 八倍下采样、patch size 为 2 之后,单帧 latent 是 128×128、patch 网格是 64×64,单帧序列长度 N=4096(视频若联合多个潜帧,N 还要乘潜帧数)。模型权重 bf16 只有 6.75 GiB,80G 的 A100/H100 看起来绰绰有余,然后第一步就 CUDA OOM。 把账摊开看(怎么算出来的见第 03、04 节):一层注意力在 N=4096 时按下述保守教学账本计为 2.39 GiB 激活,32 层就是 76.5 GiB——是权重的 11 倍。这说明只看权重大小无法判断是否 OOM;在这套训练存储假设下,激活已经超过预算,而且爆的是其中三个特定张量:分数矩阵、softmax 权重、浮点 dropout 乘子,各占 31.4%,三项合计吃掉单层激活的 94%。 按本文简化层结构,N=4096 时两个 N×N 矩阵乘只占整层前向 FLOPs 的 18.2%;它们与线性项之比在 N=18432 达到 1,只看 attention 投影则交叉点为 6144。这是运算量比例,不是运行时间比例;IO、kernel 形状与融合可能让 FLOPs 较少的部分反而更慢。 因此应同时记录 FLOPs、激活存储假设和实际时间线。N=32768 时,本文教学账本的一层存储为 145 GiB,二次项 FLOPs 占 64%;这些值用于理解增长趋势,不能替代实际框架的显存测量。 02. 最小可用理解 三句话讲完核心思想: 每个 token 拿自己的向量生成三份拷贝——Query、Key、Value,然后每个 token 拿自己的 Query 去和所有 token(包括自己)的 Key 做内积打分,分数过 softmax 变成权重,再对所有 Value 加权求和,得到这个 token 的新表示。权重由 Q/K、位置编码和掩码共同决定,输出内容还依赖 V。 计算和显存都分两笔:一笔随 N 线性(QKV 投影、输出投影、FFN),一笔随 N 二次(分数矩阵 S、softmax 权重 P、浮点 dropout 乘子各一份 B×H×N×N)。显存爆炸的几乎都是第二笔,这是显式保存中间量的教学实现假设;融合、重算或不使用 dropout 会改变这笔账。 「O(N²) 是瓶颈」有适用条件:二次项与线性项的比值是 N/(6d),占总量的比例是 N/(6d+N),N 小于 2d 时它连 attention 模块内部的一半都不到。在本例朴素存储方案下,N² 项先成为显存大头;实际耗时瓶颈仍需测量。 这张图要看什么:多头切的是 d_model 这个维度(H·D 拆成 H 份),不是把注意力复制 H 份;切分本身零算术开销,但分数矩阵从 1 个 N×N 变成 H 个 N×N,显存乘上 H。 03. 数学推导 3.1 单头:打分、归一化、加权求和 设输入序列 $X \in \mathbb{R}^{N \times d}$,N 是 token 数,d 是每个 token 的向量维度(d_model)。三个投影矩阵 $W_q, W_k, W_v \in \mathbb{R}^{d \times d}$ 把每个 token 映射成查询、键、值: $$Q = X W_q, \quad K = X W_k, \quad V = X W_v$$ 每个 token 的 Query 要和所有 token 的 Key 算相似度,写成矩阵形式就是一次 $N \times D$ 对 $D \times N$ 的矩阵乘,得到分数矩阵 $S \in \mathbb{R}^{N \times N}$,其中 $S_{ij}$ 是第 i 个 token 对第 j 个 token 的打分: $$S = \frac{Q K^{\top}}{\sqrt{D}}$$ 本小节是单头,H=1、D=d,所以 Q/K/V 都是 N×D;下一小节切成 H 头后才有 D=d/H。除以 $\sqrt{D}$ 不是装饰,推导一下就知道:假设 Q、K 的分量独立、零均值、方差为 1,那么点积的方差是 $$\mathrm{Var}(q \cdot k) = \sum_{i=1}^{D} \mathrm{Var}(q_i k_i) = D$$ 点积的标准差随 $\sqrt{D}$ 线性增长。D=128 时分数的摆动幅度是 D=1 的 11 倍,softmax 拿到这么大的输入会直接饱和:最大的那个分数吃掉几乎全部权重,输出逼近 one-hot。接近 one-hot 时,softmax Jacobian 的多数项会很小;有限 logits 的精确 softmax 通常并非严格 one-hot,浮点舍入可能进一步使梯度消失。除以 $\sqrt{D}$ 恰好把方差归一回 1。第 04 节的代码里有一张 D 从 8 扫到 256 的实测表,饱和是看得见的。 分数过 softmax 变成权重(每行归一化,行内竞争): $$P_{ij} = \frac{\exp(S_{ij})}{\sum_{j'} \exp(S_{ij'})}$$ 最后对 Value 加权求和得到输出 $O = P V$,形状和输入一样是 $N \times d$。工程实现里 softmax 前要先减去每行最大值再取指数,防止 $\exp$ 上溢——这不改变结果,因为分子分母同乘了一个常数。 3.2 多头:切的是维度,不是份数 把 d 维切成 H 段,每段 D = d/H 维当作一个独立的「头」,各算各的注意力,最后拼回来过一个输出投影: $$\mathrm{MHA}(X) = \mathrm{Concat}(\mathrm{head}_1, \dots, \mathrm{head}_H)\, W_o, \qquad \mathrm{head}_h = \mathrm{softmax}\!\left(\frac{Q_h K_h^{\top}}{\sqrt{D}}\right) V_h$$ $Q_h$ 是 $Q$ 的第 h 段 D 列。两个常被搞错的点: 固定 d 时,多头不改变两次矩阵乘的主导 FLOPs。H 个头各做 $N^2 D$ 次乘加,总共 $H \cdot N^2 D = N^2 d$,和不切头(一个 D=d 的单头)一样;softmax、调度等开销仍随头数变化。多头改的是「在多少个独立子空间里同时做注意力」,是表达能力的再分配,不是算力的加倍。 多头增加显存。分数矩阵是按头存的:H 个 $N \times N$。显存里 显式分数存储的 N² 项系数是 H 而不是 1。 3.3 算力账:二次项什么时候过半 一层 transformer 的前向 FLOPs(一次 $[M,K] \times [K,N]$ 矩阵乘算 $2MKN$ 个浮点运算): 四个 d×d 投影(Q、K、V 输入投影 + 输出投影):$8 N d^2$ FFN(升维 4d 再降回):$16 N d^2$ 注意力内部两个 N×N 矩阵乘($QK^{\top}$ 与 $PV$):$4 N^2 d$ 线性项合计 $24 N d^2$,二次项是 $4 N^2 d$,比值等于 $N / (6d)$。令比值等于 1: $$4 N^2 d = 24 N d^2 \quad \Longrightarrow \quad N^{*} = 6d$$ d=3072 时 $N^{*} = 18432$。如果只看 attention 模块内部(4 个投影对 2 个 N×N 乘),交叉点是 $N^{*} = 2d = 6144$。你日常跑的 N=4096 在两条线之下——二次项占整层算力 18.2%,占模块内 40.0%。 3.4 显存账:三个 B×H×N×N 下面采用保守的教学分配模型:同时计入 X/Q/K/V/ctx/O,以及 S、P 和一份浮点 dropout 乘子,全部按 bf16 两字节估算。这不是某个框架的峰值实测,也不是反向传播的最低存储要求;bool mask 通常只占一字节,dropout=0 时可省掉该项,S 通常不必与 P 同时保留。 线性项:X、Q、K、V、加权和、输出,共 6 个 $[B, N, d]$ 张量(不含 FFN 的话); 二次项:分数矩阵 S、softmax 权重 P、浮点 dropout 乘子,各 $[B, H, N, N]$。 $$\text{单层激活} \approx 6 B N d \cdot 2 + 3 B H N^2 \cdot 2 \;\; \text{字节}$$ 两笔相等解得 $N^{*} = 2d/H$,d=3072、H=24 时 *N=256**——序列长度刚过几百,显存就已经被 N² 项主导了。算力和显存的交叉点差 72 倍(18432 对 256),这就是「算力瓶颈来得晚、显存瓶颈来得早」的定量出处。 这张图要看什么:左边 N=4096 的堆叠条里红色三个格子(S、P、浮点 dropout 乘子)占 94%,线性项挤在边上几乎看不见;右边是对数轴,两条教学模型的比值随 N 近似线性增长,N=32768 时为 129 倍。融合侧统一按六份线性张量 6BNd·bytes 计,忽略小的 LSE 与内核工作区;这不是具体框架的峰值比。 04. 代码实现 完整脚本在文末附录(mha_minimal.py、attention_memory.py、flops_ledger.py、make_figures.py),只依赖 numpy。下面按执行顺序拆核心片段,所有数值都是 /usr/local/bin/python3 真跑出来的。 4.1 前向:六行写完 MHA def mha(X, W_q, W_k, W_v, W_o, H, causal=False): B, N, d_model = X.shape D = d_model // H Q = split_heads(X @ W_q, H) # [B, H, N, D] K = split_heads(X @ W_k, H) # [B, H, N, D] V = split_heads(X @ W_v, H) # [B, H, N, D] S = Q @ K.transpose(0, 1, 3, 2) / np.sqrt(D) # [B, H, N, N] if causal: mask = np.triu(np.ones((N, N), dtype=bool), k=1) S = np.where(mask, -np.inf, S) P = softmax(S, axis=-1) # [B, H, N, N] return merge_heads(P @ V) @ W_o, P # [B, N, d_model] 切头函数值得单独看一眼——对这里连续的输入,它是 reshape 加 transpose 的视图变换;非连续输入的 reshape 可能触发拷贝: def split_heads(t, H): B, N, d_model = t.shape D = d_model // H return t.reshape(B, N, H, D).transpose(0, 2, 1, 3) 配置 B=2、N=8、d_model=32、H=4 时各张量的真实形状: [1] 各张量形状 X (2, 8, 32) 输入 Q/K/V (2, 4, 8, 8) 切头后 [B, H, N, D] S (2, 4, 8, 8) 分数矩阵,N×N 是显存爆炸的源头 O (2, 8, 32) 输出,和输入同形 4.2 验证四条性质 脚本跑出来的校验结果: [3] 性质校验 (a) 权重行和为 1 : max|sum(P)-1| = 2.22e-16 (b) 向量化 vs 四重循环 : max|O - O_loop| = 2.22e-16 allclose=True (c) 因果掩码上三角全 0 : max = 0.00e+00 且下三角行和仍为 1 : max|sum-1| = 2.22e-16 (d) 各头权重并不相同 : 头 0 与 H 头平均的平均绝对差 = 0.0223 头之间的平均离散度 : 0.0206(0 表示所有头完全一样) (b) 是最值得做的一次校验:把公式照定义抄成四重循环(一个 token 一个 token 地打分、归一、加权),结果和向量化版本在浮点精度内完全一致。下标搞反、transpose 方向写错这类 bug,靠肉眼很难发现,靠一个慢十倍但显然正确的对照实现能当场抓住。 4.3 缩放因子的实测 D 扫描表(固定 scale=1,看 softmax 最大权重): 固定 scale=1,扫一遍 D 看 softmax 最大权重怎么变: D 分数标准差 最大权重(未缩放) 最大权重(缩放后) 8 2.8587 0.9089 0.2711 16 3.9246 0.9970 0.1718 32 5.6569 1.0000 0.2427 64 8.0203 1.0000 0.1547 128 11.3903 1.0000 0.1322 256 16.0367 1.0000 0.1619 和 3.1 节的推导对上了:分数标准差就是 $\sqrt{D}$(2.8587 ≈ √8,16.0367 ≈ √256)。本次随机样本在 D=16 时最大权重为 0.9970,之后若干行四舍五入显示 1.0000;这说明可能接近饱和,不能证明所有 token 的梯度严格为零。缩放后最大权重回落到 0.13~0.27,分布活着。另外注意主配置(D=8)下的对比:未缩放最大权重 0.5948,缩放后 0.2589——D 小的时候不缩放也能活,缩放用于控制点积随维度增长的方差;这组随机样本不构成某个 D 阈值的通用结论。 4.4 显存账本实跑 attention_memory.py 在 d_model=3072、H=24、B=1、bf16 下逐项清点: [1] 逐项账本 B=1 N=4096 bf16(2 字节) 分数矩阵 S [B, H, N, N] 805,306,368 31.4% softmax 权重 P [B, H, N, N] 805,306,368 31.4% 浮点 dropout 乘子 [B, H, N, N] 805,306,368 31.4% 输入 X / Q / K / V / ctx / O 各 25,165,824 1.0% 合计 2,566,914,048 2.39 GiB → N² 项共 2.25 GiB,线性项共 144.00 MiB,N² 项占 94.1% 随 N 的增长(一层,不含 FFN): [2] N 增长时一层 MHA 的激活显存(B=1, bf16, 含 浮点 dropout 乘子) N 线性项 N² 项 合计 融合侧教学值 倍数 1024 36.00 MiB 144.00 MiB 180.00 MiB 36.00 MiB 5.0x 2048 72.00 MiB 576.00 MiB 648.00 MiB 72.00 MiB 9.0x 4096 144.00 MiB 2.25 GiB 2.39 GiB 144.00 MiB 17.0x 8192 288.00 MiB 9.00 GiB 9.28 GiB 288.00 MiB 33.0x 16384 576.00 MiB 36.00 GiB 36.56 GiB 576.00 MiB 65.0x 32768 1.12 GiB 144.00 GiB 145.12 GiB 1.12 GiB 129.0x 若训练时按同一教学假设保留每层中间量,且不做重算,才可再乘层数:L=32 层时,N=2048 的激活是 20.25 GiB(权重的 3.0 倍),N=8192 是 297 GiB(权重的 44 倍)。推理时不应把当前层临时激活直接乘 L;需要缓存历史的自回归推理另计 KV cache 随 N 和并发数 B 都线性涨:B=32、N=32768、L=80 时光 KV cache 就要 960 GiB,这就是为什么长上下文服务都把 GQA/MLA 当标配。 4.5 算力账本实跑 flops_ledger.py 的占比表(d=3072,含 FFN): N 线性项 二次项 二次项占比 4096 927.71 GFLOPs 206.16 GFLOPs 18.2% 16384 3.71 TFLOPs 3.30 TFLOPs 47.1% 32768 7.42 TFLOPs 13.19 TFLOPs 64.0% 65536 14.84 TFLOPs 52.78 TFLOPs 78.0% 脚本末尾有一段本机 CPU 实测(numpy float32,d_t=1024、16 头,数值每台机器都不同,看趋势): N 投影 (ms) QK^T (ms) 实测比 FLOPs 比 256 1.64 2.36 1.44 0.25 1024 6.64 60.80 9.16 1.00 4096 21.36 1148.15 53.75 4.00 N=256 那行最扎眼:QK^T 的算术量只有投影的四分之一,实测却慢了 1.44 倍。原因在最后一节的算术强度(AI = FLOPs / 访存字节):投影的 AI 是 85~228(权重矩阵被整批 token 反复复用),QK^T 只有 21~31(输出是 H·N²,写完就走)。FLOPs 回答「要做多少运算」,AI 回答「能不能跑快」——这也解释了 FlashAttention 为什么省显存的同时还提速:它压根不把 N² 写回显存,等于把最贵的那笔带宽也省了。 这张图要看什么:两条曲线是二次项算力占比随 N 的爬升,红蓝两条竖虚线分别是 N=2d=6144(只算 attention 模块)和 N=6d=18432(算上 FFN);你常用的 N=4096 在两条线左边很远的位置。 05. 工业级实现对照 参考实现(以 2026-09 的 main 分支为准,上游重构频繁): huggingface/transformers → modeling_llama.py:LlamaAttention.forward huggingface/transformers → masking_utils.py / integrations/flash_attention.py:attention 实现分发 生产代码和第 04 节的最小实现有五处不一样,每一处都有理由。 5.1 不落地 N×N:eager / sdpa / flash 三条路 HF 的 attention 实现 attn_implementation 有三档: eager:与本文显式计算 S/P 的思路相同;最小代码没有 dropout,也不代表训练账本的全部分配。好处是 P 可访问,代价是二次存储;具体峰值取决于 dtype、存活期与 dropout。 sdpa:调 PyTorch 的 scaled_dot_product_attention,由 PyTorch 按硬件、dtype、mask 等选择 flash、memory-efficient 或 math 后端;不能仅凭 sdpa 名称断言没有 N×N 中间量。 flash_attention_2:在线 softmax + 分块计算,显存 O(N·d),本文统一教学账本在 N=32768 时两者为 129 倍,真实节省比例需实测。 训练长序列一律用后两档。代价是 P 不再可见——想可视化注意力图、或给 P 加自定义正则,就得回 eager 或单独导出。 5.2 因果掩码不是加 −inf 的稠密矩阵 最小实现里我建了一个 $N \times N$ 的 bool 矩阵,这本身就又是一笔 N² 显存。生产实现传 is_causal=True 让内核按位置关系现场判断,或者用范围的 sliding_window 参数,掩码矩阵完全不落地。自己手写 causal mask 矩阵是新手常见的第二处 OOM 来源。 5.3 QKV 合并成一个投影 三个 $d \times d$ 投影合并成一个 $d \times 3d$(或直接 qkv_proj),一次 GEMM 出 Q、K、V。算术量不变,但少起两次 kernel、权重读取更连续。代价是 PyTorch 里要自己 chunk(3, dim=-1) 拆回来——本次核对的 modeling_llama.py 仍保留独立的 q_proj/k_proj/v_proj;融合 QKV 是另一些架构或执行后端的选择。 5.4 KV 头数可以比 Q 头少:GQA 多头切分时 K、V 的头数用 $H_{kv} < H$(比如 8 对 32),多个 Query 头共享一组 KV。公式的改动只是把 $K_h$ 换成 $K_{\lfloor h/(H/H_{kv})\rfloor}$。它不改变 attention 核心两次矩阵乘的主导 FLOPs,但可以减少 K/V 投影的 FLOPs,省的是 KV cache 和 KV 的显存与带宽——推理时 KV cache 缩到 $H_{kv}/H$(LLaMA-3 70B 的 8/64 为 1/8)。这是在固定 Query 宽度时减少 KV 头数、压缩线性 KV 项的优化,而 FlashAttention 优化的是 N² 那一笔,两者正交,经常一起用。 5.5 buffer 化与 position_ids 无位置编码且无位置相关掩码的自注意力对 token 排列是等变的。位置信息可通过正弦表、可学习嵌入、RoPE 或掩码注入,其中 RoPE 的旋转点积体现相对位移,不应统称绝对位置编码。生产实现把 cos/sin 表注册成 buffer 预计算缓存,并用外部传入的 position_ids 而不是 arange——因为 KV cache 场景下每个新 token 的位置不是从 0 开始,packed 训练时一段序列内部还要重置。位置怎么进注意力,是 RoPE 那一篇的主题。 06. 代价与边界 把 N² 落地换来了什么,又赔了什么。 朴素实现唯一的优点是 S、P 全程可见:可视化注意力图、或使用依赖完整 P 的蒸馏损失,需要访问相应权重。加到 logits 上的结构化 bias 则不必先物化 P;若内核支持其形式,可在分块时应用。导出完整 P 仍需相应的二次输出存储。工程上常见的折中是:训练用 sdpa/flash,分析时用小 N 的 eager 导出注意力图。 二次项的 FLOPs 收益要按序列长度评估。 若各项运行时间恰好与 FLOPs 成正比,N=4096 时消除占比 18.2% 的二次项,整层理想加速约为 1.22 倍;现实中的 IO 和融合会改变时间占比,因此这不是 FlashAttention 的实测加速上限。 多头的账要两头看。 头数 H 越大,每个头的 D = d/H 越小:每头维度 D 改变会影响子空间容量,但不存在这里能够证明的 D<32 通用质量阈值;固定 d=H·D 时,标准 MHA 的 KV cache 与 KV 投影参数量不因增加 H 而线性增长;若固定 D 则另当别论。所以现代模型反而从「H 越多越好」退到「适度头数 + GQA」:LLaMA-3 70B 用 64 个 Query 头配 8 个 KV 头。 什么时候根本不该用全局自注意力。 像素级 self-attention(把 H×W 个像素当 token)在中等分辨率下 N 就上了万,N² 显存直接不可行——这是 latent diffusion 在压缩空间里计算更经济的原因之一;卷积等其他算子的成本也同时下降。高分辨率密集预测里,窗口注意力、局部注意力是常态而不是妥协。 别只优化注意力。 N 小于 2d 时(d=3072 即 N<6144),attention 模块内部的算力大头是投影;FFN 属于整层的另一部分;显存侧倒是早就归 N² 管。所以「 profiling 之前先改结构」是赌博——performance_profiling 那一篇讲的账本方法就是为此准备的。 07. 经典论文脉络 Attention Is All You Need(Vaswani et al., 2017)——本文锚点。把缩放点积注意力 + 多头定型成今天的形态,丢掉循环结构,端到端只靠注意力。 Neural Machine Translation by Jointly Learning to Align and Translate(Bahdanau et al., 2014)——注意力的史前史:注意力最初是翻译里的一组对齐权重,softmax 那一行的「分布」语义就是从这来的。 Generating Long Sequences with Sparse Transformers(Child et al., 2019)——第一条正面强攻 N² 显存的路线:既然 N² 落不下,就让大部分格子为零。稀疏化一脉的开端。 FlashAttention(Dao et al., 2022)——不动公式、只改执行:在线 softmax + 分块,让 N×N 不落地。精确注意力,不是近似——这一点和稀疏/线性路线有本质区别。 GQA: Training Generalized Multi-Query Transformer Models(Ainslie et al., 2023)——把矛头从 N² 转向 KV cache:KV 头数变少、Q 头分组共享,是许多现代模型降低 KV cache 成本的重要设计。 五篇连起来读的线索:注意力先是「一种对齐手段」(2014),再是「唯一的序列算子」(2017),然后 N² 账单到期,工程上先有人砍格子(2019),再有人改执行不砍精度(2022),最后有人发现真正贵的还有 KV 那笔线性账(2023)。 08. 常见误解 「注意力是 O(N²),所以序列不长时它也是最慢的部分。」 N=4096、d=3072 时二次项只占本文层结构前向 FLOPs 的 18.2%,但不能据此猜测热点。应结合算术强度、硬件与 profiler 判断。 「多头注意力算 H 遍,所以比单头慢 H 倍。」 H·N²·D = N²·d,多头和全维单头的算术量完全相同;多的是 H 份 N×N 显存和 H 份小矩阵乘的调度开销,不是 H 倍算力。 「显存不够就是模型太大,换个更小的模型。」 d=3072、L=32 的模型权重 6.75 GiB,N=8192 时仅激活就 297 GiB。先算激活账(6B·N·d·bytes + 3B·H·N²·bytes 乘层数),再决定动不动模型。这套无重算训练账本提示应先比较开 gradient checkpointing、换 sdpa/flash、或降分辨率。 「除以 √D 是可要可不要的数值技巧。」 本次样本中未缩放 logits 更容易接近饱和;D=16 的最大权重实际为 0.9970,显示 1.0000 也不等于数学上梯度严格为零。它控制点积方差随维度增长;其他归一化或初始化设计也能缓解饱和,不能据此宣称不缩放就一定无法训练。 「浮点 dropout 乘子不占显存。」 bool mask 通常每元素一字节,而 bf16 为两字节;本文账本计的是两字节的浮点 dropout 乘子,两者不能混称。使用逐元素 dropout 的 eager 训练通常还需保留随机掩码或等价信息;占用取决于表示与实现,不一定恰好是一份 bf16 张量。FlashAttention 可以保存随机数状态并在分块反向时重建掩码。 09. 动手验证 把附录里的 mha_minimal.py 存下来直接跑(python mha_minimal.py),对照三处输出: 形状链:X (2,8,32) → Q/K/V (2,4,8,8) → S (2,4,8,8) → O (2,8,32)。确认分数矩阵是按头存的 4 个 8×8,不是 1 个。 缩放对照:未缩放最大权重 0.5948,缩放后 0.2589;再把 D 扫描表看一遍,本次 D=16 的未缩放值是 0.9970,后续多行显示值接近 1。 等价性:向量化实现与四重循环实现的 max 误差应为 1e-16 量级(你机器上具体数字可能略有不同,但 allclose 一定是 True)。 然后做两个改动观察变化: 把 H 从 4 改成 1 再改成 8,跑性质校验 (d):头间离散度会从 0.0206 变成 0(单头没有「头间」可言)——多头不是免费的多样性,是切分带来的。 跑 flops_ledger.py 的 CPU 实测段,找到你机器上「实测比」超过「FLOPs 比」的 N 拐点,和理论交叉点 N*=2d 对一下差多少。 预期最容易翻车的是第二个:很多人会预期实测比从一开始就贴近 FLOPs 比,实际 N=256 时实测比约 1.44、FLOPs 比只有 0.25——算术强度那笔账不在 FLOPs 公式里。 10. 延伸阅读 按知识树的依赖顺序,下一步建议这么走: 旋转位置编码 RoPE 的原理与实现——注意力对位置是盲的,RoPE 是视觉/语言模型目前的主流注入方式;本文的 Q、K 在那里会被旋转一次。 视频 DiT 里的 3D RoPE 与分辨率外推——RoPE 在时间/高/宽三组频率上的拆分,视频生成的位置问题。 性能建模与 Profiling:算力、带宽与显存账本——把本文的两本账变成系统方法,roofline 分析。 序列并行——N² 显存装不下时的横向切分方案。 策略梯度与 PPO 基础 等 RL 方向的文章与本篇无直接依赖,可随时穿插。 附录:完整代码 09 节用到的脚本全文如下(mha_minimal.py、attention_memory.py、flops_ledger.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 mha_minimal.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """多头自注意力(MHA)的最小可运行实现:把公式逐行翻译成 numpy。 只依赖 numpy,直接 `python mha_minimal.py` 即可运行。 刻意不用 PyTorch:讲原理时框架的抽象反而是噪声,而且 numpy 谁都能跑。 生产实现见文章第 05 节对 transformers 的引用。 公式符号与代码变量名的对应关系 ------------------------------ X 输入序列 形状 [B, N, d_model] W_q/W_k/W_v 三个输入投影矩阵 形状 [d_model, d_model] W_o 输出投影矩阵 形状 [d_model, d_model] Q, K, V 查询 / 键 / 值 形状 [B, H, N, D] S 缩放后的注意力分数 形状 [B, H, N, N],S = Q K^T / sqrt(D) P softmax 后的权重 形状 [B, H, N, N] O 注意力输出 形状 [B, N, d_model] 运行后你会看到:每一层的真实 shape、注意力权重的数值范围、 以及四条性质校验(行和为 1 / 与循环实现等价 / 缩放系数的影响 / 因果掩码)。 """ import numpy as np rng = np.random.default_rng(0) def split_heads(t: np.ndarray, H: int) -> np.ndarray: """[B, N, H*D] -> [B, H, N, D]。 多头不是「复制 H 份再算」,而是把 d_model 这个维度切成 H 段, 每段独立算一次注意力,最后再拼回去。这一步只是 reshape + transpose, 本身不做任何算术,但它决定了后面所有矩阵乘的形状。 """ B, N, d_model = t.shape D = d_model // H return t.reshape(B, N, H, D).transpose(0, 2, 1, 3) def merge_heads(t: np.ndarray) -> np.ndarray: """[B, H, N, D] -> [B, N, H*D],split_heads 的逆操作。""" B, H, N, D = t.shape return t.transpose(0, 2, 1, 3).reshape(B, N, H * D) def softmax(x: np.ndarray, axis: int = -1) -> np.ndarray: """数值稳定的 softmax:先减去最大值再取指数,避免 exp 溢出。""" x = x - x.max(axis=axis, keepdims=True) e = np.exp(x) return e / e.sum(axis=axis, keepdims=True) def mha(X: np.ndarray, W_q, W_k, W_v, W_o, H: int, causal: bool = False): """完整的多头自注意力前向。返回 (输出, 注意力权重 P)。""" B, N, d_model = X.shape D = d_model // H # ① 三个投影:每个 token 独立地把自己映射成 query / key / value Q = split_heads(X @ W_q, H) # [B, H, N, D] K = split_heads(X @ W_k, H) # [B, H, N, D] V = split_heads(X @ W_v, H) # [B, H, N, D] # ② 打分:每个 query 和所有 key 做内积,得到 N×N 的分数矩阵 S = Q @ K.transpose(0, 1, 3, 2) / np.sqrt(D) # [B, H, N, N] # ③ 因果掩码:第 i 个 query 只能看 0..i 的 key if causal: mask = np.triu(np.ones((N, N), dtype=bool), k=1) S = np.where(mask, -np.inf, S) # ④ 归一化成权重,再对 V 做加权求和 P = softmax(S, axis=-1) # [B, H, N, N] ctx = P @ V # [B, H, N, D] # ⑤ 拼回 d_model,过输出投影 O = merge_heads(ctx) @ W_o # [B, N, d_model] return O, P def mha_naive_loop(X, W_q, W_k, W_v, W_o, H): """完全不用矩阵乘的「照着定义抄」版本:四重循环。 用来验证上面的向量化实现没有把下标搞反。逻辑等价但慢得多(O(B·H·N²·D))。 """ B, N, d_model = X.shape D = d_model // H Q = (X @ W_q).reshape(B, N, H, D).transpose(0, 2, 1, 3) K = (X @ W_k).reshape(B, N, H, D).transpose(0, 2, 1, 3) V = (X @ W_v).reshape(B, N, H, D).transpose(0, 2, 1, 3) out = np.zeros((B, H, N, D)) for b in range(B): for h in range(H): for i in range(N): # 先逐个算出这一行 N 个分数,再 softmax,再加权求和 s = np.array([float(Q[b, h, i] @ K[b, h, j]) / np.sqrt(D) for j in range(N)]) p = softmax(s) out[b, h, i] = p @ V[b, h] return merge_heads(out) @ W_o def main(): B, N, d_model, H = 2, 8, 32, 4 D = d_model // H print(f"配置: B={B} N={N} d_model={d_model} H={H} D={D}") X = rng.normal(size=(B, N, d_model)) * 0.5 W_q, W_k, W_v, W_o = (rng.normal(size=(d_model, d_model)) / np.sqrt(d_model) for _ in range(4)) # ── 1. 每一层的真实形状 ────────────────────────────────── Q = split_heads(X @ W_q, H) K = split_heads(X @ W_k, H) S = Q @ K.transpose(0, 1, 3, 2) / np.sqrt(D) print("\n[1] 各张量形状") print(f" X {X.shape} 输入") print(f" Q/K/V {Q.shape} 切头后 [B, H, N, D]") print(f" S {S.shape} 分数矩阵,N×N 是显存爆炸的源头") print(f" O {mha(X, W_q, W_k, W_v, W_o, H)[0].shape} 输出,和输入同形") # ── 2. 分数的量级:为什么要除以 sqrt(D) ────────────────── raw = Q @ K.transpose(0, 1, 3, 2) # 不除以 sqrt(D) print("\n[2] 缩放系数 sqrt(D) 的作用") print(f" D = {D}, sqrt(D) = {np.sqrt(D):.4f}") print(f" 未缩放分数的标准差 : {raw.std():.4f} 分布范围 [{raw.min():.2f}, {raw.max():.2f}]") print(f" 缩放后分数的标准差 : {S.std():.4f} 分布范围 [{S.min():.2f}, {S.max():.2f}]") print(f" 未缩放 softmax 的最大权重: {softmax(raw).max():.4f}") print(f" 缩放后 softmax 的最大权重: {softmax(S).max():.4f}") # D 越大,不缩放的分数方差越大,softmax 越容易塌成 one-hot print(" 固定 scale=1,扫一遍 D 看 softmax 最大权重怎么变:") print(" D 分数标准差 最大权重(未缩放) 最大权重(缩放后)") for D_test in (8, 16, 32, 64, 128, 256): q = rng.normal(size=(256, D_test)) k = rng.normal(size=(256, D_test)) s_raw = q @ k.T s_scaled = s_raw / np.sqrt(D_test) print(f" {D_test:4d} {s_raw.std():9.4f} " f"{softmax(s_raw).max():14.4f} {softmax(s_scaled).max():15.4f}") # ── 3. 性质校验 ──────────────────────────────────────── print("\n[3] 性质校验") O, P = mha(X, W_q, W_k, W_v, W_o, H) print(f" (a) 权重行和为 1 : max|sum(P)-1| = " f"{np.abs(P.sum(axis=-1) - 1).max():.2e}") O_loop = mha_naive_loop(X, W_q, W_k, W_v, W_o, H) print(f" (b) 向量化 vs 四重循环 : max|O - O_loop| = " f"{np.abs(O - O_loop).max():.2e} allclose=" f"{np.allclose(O, O_loop)}") _, P_causal = mha(X, W_q, W_k, W_v, W_o, H, causal=True) upper = np.triu(P_causal, k=1) print(f" (c) 因果掩码上三角全 0 : max = {upper.max():.2e}") print(f" 且下三角行和仍为 1 : max|sum-1| = " f"{np.abs(P_causal.sum(axis=-1) - 1).max():.2e}") # 各头学出来的权重并不相同,这才让「多头」有意义 P_mean = P.mean(axis=1) # [B, N, N],把 H 个头的权重平均 diff = np.abs(P[0, 0] - P_mean[0]).mean() spread = np.abs(P[0] - P[0].mean(axis=0, keepdims=True)).mean() print(f" (d) 各头权重并不相同 : 头 0 与 H 头平均的平均绝对差 = {diff:.4f}") print(f" 头之间的平均离散度 : {spread:.4f}(0 表示所有头完全一样)") # ── 4. 不看自己的极端情形:注意力塌缩成「复制」 ────────── print("\n[4] 极端情形:把 K 设成和 Q 完全一样(自注意力且 W_q=W_k)") X2 = rng.normal(size=(1, 4, d_model)) Q2 = split_heads(X2 @ W_q, H) S2 = Q2 @ Q2.transpose(0, 1, 3, 2) / np.sqrt(D) P2 = softmax(S2) diag_share = np.trace(P2[0, 0]) / 4 print(f" 对角线权重占比 = {diag_share:.4f}(1.0 表示每个 token 只关注自己)") if __name__ == "__main__": main() attention_memory.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """自注意力的显存账本:把一层 attention 的每一笔开销都算成字节。 只依赖 numpy(其实只用到它做格式化,核心就是整数乘除)。 直接 `python attention_memory.py` 即可运行。 算的是「训练时一层 MHA 需要为反向传播留下来的激活」, 这是显式保留中间量的保守教学模型,并非框架实测或反向所需最小值。默认的模型尺寸 采用简化 Transformer:d_model=3072、H=24、D=128。 """ import numpy as np BYTES = {"fp32": 4, "fp16": 2, "bf16": 2, "fp8": 1} MIB = 1024 ** 2 GIB = 1024 ** 3 def fmt(nbytes: float) -> str: if nbytes >= GIB: return f"{nbytes / GIB:8.2f} GiB" if nbytes >= MIB: return f"{nbytes / MIB:8.2f} MiB" return f"{nbytes / 1024:8.2f} KiB" def ledger(B: int, N: int, d_model: int, H: int, dtype: str = "bf16", dropout: bool = True): """返回一层 MHA 的逐项激活显存(字节)。 N² 项有三个:分数矩阵 S、softmax 后的 P、以及与所选 dtype 同宽的浮点 dropout 乘子(非 bool mask)。 这三项就是 O(N²) 显存的真身——不是「注意力复杂度是 N²」这句话, 而是本教学模型假设同时保存的三个 B×H×N×N 张量;框架可复用或省略它们。 """ b = BYTES[dtype] linear = B * N * d_model * b # 每个 [B, N, d_model] 的张量 square = B * H * N * N * b # 每个 [B, H, N, N] 的张量 return { "输入 X": linear, "Q": linear, "K": linear, "V": linear, "分数矩阵 S": square, "softmax 权重 P": square, "浮点 dropout 乘子": square if dropout else 0, "加权和 ctx": linear, "输出 O": linear, } def main(): d_model, H, D = 3072, 24, 128 print(f"模型尺寸: d_model={d_model} H={H} D={D} (H*D = {H * D})") # ── 1. 单个 N 下的逐项账本 ────────────────────────────── B, N = 1, 4096 items = ledger(B, N, d_model, H, "bf16") total = sum(items.values()) print(f"\n[1] 逐项账本 B={B} N={N} bf16(2 字节)") print(f" {'项目':<16}{'形状':<22}{'字节数':>14} 占比") for name, nb in items.items(): shape = "[B, N, d_model]" if nb == B * N * d_model * 2 else "[B, H, N, N]" print(f" {name:<16}{shape:<22}{nb:>14,} {nb / total:6.1%}") print(f" {'合计':<16}{'':<22}{total:>14,} {fmt(total)}") quad = sum(v for k, v in items.items() if "S" in k or "P" in k or "乘子" in k) lin = total - quad print(f" → N² 项共 {fmt(quad)},线性项共 {fmt(lin)},N² 项占 {quad / total:.1%}") # ── 2. N 增长时账本怎么变 ─────────────────────────────── print("\n[2] N 增长时一层 MHA 的激活显存(B=1, bf16, 含 浮点 dropout 乘子)") print(f" {'N':>7}{'线性项':>12}{'N² 项':>12}{'合计':>12} " f"{'融合侧教学值':>15} 倍数") for N in (1024, 2048, 4096, 8192, 16384, 32768): it = ledger(1, N, d_model, H, "bf16") q = it["分数矩阵 S"] + it["softmax 权重 P"] + it["浮点 dropout 乘子"] l = sum(it.values()) - q # 与 FlashAttention 文统一:六份线性量的教学预算,忽略小的 LSE 与工作区 flash = 6 * 1 * N * d_model * 2 print(f" {N:>7}{fmt(l):>12}{fmt(q):>12}{fmt(l + q):>12} " f"{fmt(flash):>15} {(l + q) / flash:5.1f}x") # ── 3. 交叉点:从哪个 N 开始 N² 项压过线性项 ──────────── # 3·B·H·N² = 6·B·N·d → N* = 2d / H N_star = 2 * d_model / H print("\n[3] 交叉点") print(f" 3·B·H·N²·bytes = 6·B·N·d·bytes 解得 N* = 2d/H = {N_star:.0f}") print(f" 也就是说 N 超过 {N_star:.0f} 之后,一层 attention 的激活显存") print(f" 就由 N² 项主导;N=4096 时已经超出 {(4096 / N_star):.0f} 倍。") # ── 4. 乘上层数:为什么 L 层比 N 更狠 ──────────────────── print("\n[4] 乘上层数 L(无重算训练,按上述教学模型保留各层中间量)") print(f" {'L':>4}{'N=2048':>12}{'N=4096':>12}{'N=8192':>12} 说明") for L in (12, 24, 32, 48): row = [] for N in (2048, 4096, 8192): tot = sum(ledger(1, N, d_model, H, "bf16").values()) row.append(fmt(tot * L)) print(f" {L:>4}{row[0]:>12}{row[1]:>12}{row[2]:>12}") print(" 以上均未计入模型权重与优化器状态;梯度检查点减少内部保存量,但仍需保留边界激活,并非通用的 1/L。") # ── 5. 推理侧的另一种账本:KV cache ───────────────────── print("\n[5] 推理侧:KV cache 的账本(随 N 线性增长,但随并发数 B 线性增长)") print(" 公式: 2 · B · N · L · H · D · bytes") print(f" {'B':>4}{'N':>7}{'L=32':>12}{'L=80':>12}") for B in (1, 8, 32): for N in (4096, 32768): r = [] for L in (32, 80): nb = 2 * B * N * L * H * D * 2 r.append(fmt(nb)) print(f" {B:>4}{N:>7}{r[0]:>12}{r[1]:>12}") # ── 6. 权重和激活谁更大 ──────────────────────────────── print("\n[6] 权重 vs 激活:别只盯着模型大小") L = 32 # 一层 transformer 的权重:4 个 d×d 投影 + 2 个 FFN 的 4d×d w_per_layer = (4 * d_model * d_model + 2 * d_model * 4 * d_model) * 2 print(f" 一层权重(4 个 d×d + FFN 的 2×4d×d,bf16): {fmt(w_per_layer)}") print(f" L={L} 层总权重 : {fmt(w_per_layer * L)}") for N in (2048, 8192): act = sum(ledger(1, N, d_model, H, "bf16").values()) * L print(f" N={N} 时 L={L} 层激活 : {fmt(act)} " f"(是权重的 {act / (w_per_layer * L):.1f} 倍)") if __name__ == "__main__": main() flops_ledger.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """自注意力的算力账本:把一层 attention 拆成「线性项」和「二次项」两笔。 只依赖 numpy。直接 `python flops_ledger.py` 即可运行。 最后有一段 CPU 实测,耗时约 10 秒;不同机器数值会不同,看趋势即可。 FLOPs 口径:一次矩阵乘 [M,K] @ [K,N] 需要 M·K·N 次乘加(MAC), 一次乘加算 2 个浮点运算,所以 FLOPs = 2·M·K·N。全文统一用这个口径。 """ import time import numpy as np def layer_flops(N: int, d: int, with_ffn: bool = True): """一层 transformer 的前向 FLOPs,拆成线性项和二次项。""" # 4 个 d×d 投影:Q、K、V 三个输入投影 + 1 个输出投影 proj = 4 * 2 * N * d * d # FFN:升到 4d 再降回来,两个矩阵乘 ffn = 2 * 2 * N * d * 4 * d if with_ffn else 0.0 # 注意力内部两个 N×N 的矩阵乘:QK^T 与 P·V attn = 2 * 2 * N * N * d return proj + ffn, attn def fmt_flops(f: float) -> str: for unit, scale in (("P", 1e15), ("T", 1e12), ("G", 1e9), ("M", 1e6)): if f >= scale: return f"{f / scale:7.2f} {unit}FLOPs" return f"{f:7.2f} FLOPs" def main(): d, H = 3072, 24 print(f"模型尺寸: d_model={d} H={H}") # ── 1. 一层里两笔账各是多少 ───────────────────────────── print("\n[1] 一层 transformer 的前向 FLOPs(N=4096, 含 FFN)") linear, quad = layer_flops(4096, d) print(f" 线性项(QKV 投影 + O 投影 + FFN): {fmt_flops(linear)}") print(f" 二次项(QK^T 与 P·V) : {fmt_flops(quad)}") print(f" 二次项占比 : {quad / (linear + quad):.1%}") # 只看 attention 内部:4 个投影 vs 2 个 N×N 矩阵乘 l_attn_only, q_attn_only = layer_flops(4096, d, with_ffn=False) print(f" 只算 attention 模块本身(去掉 FFN): 二次项占 " f"{q_attn_only / (l_attn_only + q_attn_only):.1%}") # ── 2. 交叉点:二次项什么时候超过线性项 ────────────────── # 4N²d = 24Nd² → N* = 6d print("\n[2] 交叉点") print(f" 含 FFN : 4N²d = 24Nd² → N* = 6d = {6 * d}") print(f" 只含投影: 4N²d = 8Nd² → N* = 2d = {2 * d}") print(f" 结论:N 小于 {2 * d} 时,二次项的 FLOPs 少于 attention 投影,") print(f" 这只比较运算量;是否为耗时瓶颈仍取决于访存、内核与硬件。") # ── 3. 占比怎么随 N 变化 ─────────────────────────────── print("\n[3] 二次项占比随 N 变化(d=3072,含 FFN)") print(f" {'N':>7}{'线性项':>16}{'二次项':>16}{'二次项占比':>10}") for N in (512, 1024, 2048, 4096, 8192, 16384, 32768, 65536): l, q = layer_flops(N, d) print(f" {N:>7}{fmt_flops(l):>16}{fmt_flops(q):>16}{q / (l + q):>10.1%}") # ── 4. 训练总算力:前向的三倍 ─────────────────────────── print("\n[4] 训练一个 token 的 FLOPs(前向 + 反向 ≈ 3 倍前向)") for N in (2048, 4096, 8192): l, q = layer_flops(N, d) print(f" N={N:<6} 每 token 前向 {fmt_flops((l + q) / N)}" f" 训练 {fmt_flops(3 * (l + q) / N)}") # ── 5. CPU 实测:二次项的增长是不是真的更快 ────────────── print("\n[5] CPU 实测(本机一次运行,看趋势不看绝对值;numpy float32)") d_t, H_t = 1024, 16 rng = np.random.default_rng(0) X = rng.normal(size=(4096, d_t)).astype(np.float32) W = rng.normal(size=(d_t, d_t)).astype(np.float32) / np.sqrt(d_t) def bench(fn, repeat=3): best = float("inf") for _ in range(repeat): t0 = time.perf_counter() fn() best = min(best, time.perf_counter() - t0) return best * 1000.0 print(f" {'N':>6}{'投影 (ms)':>12}{'QK^T (ms)':>12}{'实测比':>9}" f"{'FLOPs 比':>10}") measured = {} for N in (256, 512, 1024, 2048, 4096): x = X[:N] t_proj = bench(lambda: x @ W) D_t = d_t // H_t Q = (x @ W)[:, :].reshape(N, H_t, D_t).transpose(1, 0, 2) t_qk = bench(lambda: Q @ Q.transpose(0, 2, 1)) # 理论 FLOPs 比:二次项 / 一个投影 ratio_flops = (2 * N * N * d_t) / (2 * N * d_t * d_t) measured[N] = t_qk / t_proj print(f" {N:>6}{t_proj:>12.2f}{t_qk:>12.2f}" f"{t_qk / t_proj:>9.2f}{ratio_flops:>10.2f}") print(f" 注意看 N=256 那一行:QK^T 的算术量只有投影的 0.25 倍,") print(f" 本次 QK^T / 投影耗时比为 {measured[256]:.2f}。耗时比不必等于 FLOPs 比。") # ── 6. 算术强度:FLOPs 回答不了「为什么慢」 ─────────────── print("\n[6] 算术强度 AI = FLOPs / 访存字节(float32,4 字节/元素)") print(" AI 低 = 每读一个字节只做很少的运算 = 带宽先撑不住(访存受限)") print(f" {'算子':<22}{'N=256':>12}{'N=1024':>12}{'N=4096':>12}") rows = [] for N in (256, 1024, 4096): # 投影 X[N,d] @ W[d,d]:读 X 与 W,写输出 f_proj = 2 * N * d_t * d_t m_proj = (N * d_t + d_t * d_t + N * d_t) * 4 # QK^T:读 Q 与 K,写 N×N 的分数矩阵(H 份) f_qk = 2 * N * N * d_t m_qk = (2 * N * d_t + H_t * N * N) * 4 rows.append((N, f_proj / m_proj, f_qk / m_qk)) print(f" {'投影 X@W_q':<22}" + "".join(f"{r[1]:>12.1f}" for r in rows)) print(f" {'注意力 QK^T':<22}" + "".join(f"{r[2]:>12.1f}" for r in rows)) print(f" {'(投影 / QK^T)':<22}" + "".join(f"{r[1] / r[2]:>12.1f}" for r in rows)) print(" 投影的 AI 高出 4~7 倍:权重矩阵被整批 token 复用,读一次能做 N 次乘加。") print(" QK^T 的 AI 低:输出是 H·N²,写完就走,数据复用少。") print(" 这正是 FlashAttention 的第二个收益——不把 N² 写回显存,") print(" 它省的不只是容量,还有带宽。") if __name__ == "__main__": main() make_figures.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """生成「自注意力机制的计算与显存账本」的三张解释图。 数值全部来自同目录下的 attention_memory.py 与 flops_ledger.py, 改了那两个脚本的话这里要跟着重跑,避免图与正文数字不一致。 只依赖 numpy + matplotlib,直接 `python make_figures.py` 即可运行。 """ from pathlib import Path import matplotlib.pyplot as plt import numpy as np from matplotlib.patches import FancyBboxPatch ROOT = Path(__file__).resolve().parents[1] OUT = ROOT / "figures" OUT.mkdir(exist_ok=True) plt.rcParams.update({ "font.sans-serif": ["PingFang SC", "Arial Unicode MS", "DejaVu Sans"], "axes.unicode_minus": False, "figure.dpi": 160, }) C_QUAD = "#e05263" # N² 项:红 C_LIN = "#3f7fbf" # 线性项:蓝 C_FLASH = "#2f9e6f" # FlashAttention:绿 C_GREY = "#94a3b8" INK = "#182238" MIB = 1024 ** 2 GIB = 1024 ** 3 def _style(ax, title, xlabel, ylabel): ax.set_title(title, fontsize=14, weight="bold", color=INK, pad=10) ax.set_xlabel(xlabel, fontsize=11, color="#475569") ax.set_ylabel(ylabel, fontsize=11, color="#475569") ax.tick_params(colors="#475569", labelsize=10) for s in ("top", "right"): ax.spines[s].set_visible(False) for s in ("left", "bottom"): ax.spines[s].set_color("#cbd5e1") ax.grid(axis="y", color="#eef2f7", linewidth=1) ax.set_axisbelow(True) # ──────────────────────────────────────────────────────────── # 图 1:显存账本的构成 # ──────────────────────────────────────────────────────────── def fig_memory_ledger(): d_model, H = 3072, 24 b = 2 # bf16 fig, axes = plt.subplots(1, 2, figsize=(13.2, 4.9), facecolor="#fbfcfe") # 左:N=4096 时的逐项占比(一根堆叠条) ax = axes[0] N = 4096 lin_item = 1 * N * d_model * b quad_item = 1 * H * N * N * b labels = ["输入 X", "Q", "K", "V", "分数矩阵 S", "softmax 权重 P", "浮点 dropout 乘子", "加权和 ctx", "输出 O"] sizes = [lin_item, lin_item, lin_item, lin_item, quad_item, quad_item, quad_item, lin_item, lin_item] colors = [C_LIN] * 4 + [C_QUAD] * 3 + [C_LIN] * 2 total = sum(sizes) left = 0.0 for lab, s, c in zip(labels, sizes, colors): ax.barh([0], [s / GIB], left=left / GIB, color=c, edgecolor="white", linewidth=1.2, height=0.55) if s / total > 0.05: ax.text(left / GIB + s / GIB / 2, 0, f"{lab}\n{s / total:.1%}", ha="center", va="center", fontsize=9.5, color="white", weight="bold") left += s ax.set_xlim(0, total / GIB) ax.set_yticks([]) for s in ("left", "top", "right"): ax.spines[s].set_visible(False) ax.set_xlabel("单层 MHA 的教学激活存储(GiB)", fontsize=11, color="#475569") ax.set_title(f"N={N} 时,三个 N² 项吃掉 94%", fontsize=14, weight="bold", color=INK, pad=10) ax.text(0, -0.42, f"合计 {total / GIB:.2f} GiB | 模型 d_model={d_model}, " f"H={H}, bf16", ha="left", fontsize=10, color="#64748b") # 右:随 N 增长,朴素实现 vs FlashAttention ax = axes[1] Ns = np.array([1024, 2048, 4096, 8192, 16384, 32768]) lin = 6 * Ns * d_model * b # 与逐项账本相同:6 个 [B,N,d] 张量 quad = 3 * H * Ns ** 2 * b # 3 个 [B,H,N,N] 张量 naive = (lin + quad) / GIB flash = (6 * Ns * d_model * b) / GIB # 统一教学线性预算,忽略小的 LSE ax.plot(Ns, naive, "o-", color=C_QUAD, linewidth=2.4, markersize=6, label="朴素实现(落 N² 到显存)") ax.plot(Ns, flash, "s-", color=C_FLASH, linewidth=2.4, markersize=6, label="融合侧教学值(6Nd)") ax.fill_between(Ns, flash, naive, color=C_QUAD, alpha=0.10) ax.set_yscale("log") ax.set_xscale("log") ax.set_xticks(Ns) ax.set_xticklabels([f"{n:,}" for n in Ns], fontsize=9) ax.legend(fontsize=10, frameon=False, loc="upper left") _style(ax, "两者比值随 N 近似线性增长", "序列长度 N(token 数)", "单层教学激活存储(GiB,对数轴)") ax.text(0.98, 0.06, f"N=32768 时相差 {naive[-1] / flash[-1]:.0f} 倍", transform=ax.transAxes, ha="right", fontsize=10, color=C_QUAD, weight="bold") fig.suptitle("自注意力的显存账本:钱花在哪一笔", fontsize=17, weight="bold", color="#0f172a", y=1.03) fig.tight_layout() fig.savefig(OUT / "memory_ledger.png", bbox_inches="tight", facecolor=fig.get_facecolor()) plt.close(fig) # ──────────────────────────────────────────────────────────── # 图 2:算力账本里二次项的占比 # ──────────────────────────────────────────────────────────── def fig_flops_ratio(): d = 3072 Ns = np.logspace(9, 16.5, 300, base=2) # 512 ~ 92682 quad = 4 * Ns ** 2 * d lin_proj = 8 * Ns * d * d # 只含 4 个投影 lin_all = lin_proj + 16 * Ns * d * d # 再算上 FFN fig, ax = plt.subplots(figsize=(8.6, 5.0), facecolor="#fbfcfe") r_only = quad / (quad + lin_proj) r_all = quad / (quad + lin_all) ax.plot(Ns, r_only, "-", color=C_QUAD, linewidth=2.6, label="只算 attention 模块(4 个投影 vs 2 个 N×N 乘)") ax.plot(Ns, r_all, "-", color=C_LIN, linewidth=2.6, label="算上 FFN(整层的线性项)") ax.axhline(0.5, color=C_GREY, linestyle="--", linewidth=1.2) ax.text(Ns[0], 0.52, "50% 线", fontsize=10, color=C_GREY) for N_star, color, tag in ((2 * d, C_QUAD, "N* = 2d = 6,144"), (6 * d, C_LIN, "N* = 6d = 18,432")): ax.axvline(N_star, color=color, linestyle=":", linewidth=1.6) ax.text(N_star * 1.05, 0.06, tag, fontsize=10.5, color=color, weight="bold", rotation=90, va="bottom") ax.set_xscale("log", base=2) ax.set_xticks([512, 1024, 2048, 4096, 8192, 16384, 32768, 65536]) ax.set_xticklabels(["512", "1K", "2K", "4K", "8K", "16K", "32K", "64K"]) ax.set_ylim(0, 1) ax.legend(fontsize=10.5, frameon=False, loc="upper left") _style(ax, "二次项的 FLOPs 占比随 N 增长(不等于耗时占比)", "序列长度 N(token 数)", "二次项占该层前向 FLOPs 的比例") ax.text(0.98, 0.30, "常用的 N=4096:\n只算 attention 时二次项占 40%,\n算上 FFN 后只占 18%", transform=ax.transAxes, ha="right", fontsize=10, color="#475569", bbox=dict(boxstyle="round,pad=0.4", facecolor="#f1f5f9", edgecolor="#e2e8f0")) fig.suptitle("算力账本:二次项什么时候才是主角", fontsize=16, weight="bold", color="#0f172a", y=1.00) fig.tight_layout() fig.savefig(OUT / "flops_ratio.png", bbox_inches="tight", facecolor=fig.get_facecolor()) plt.close(fig) # ──────────────────────────────────────────────────────────── # 图 3:多头到底切了什么 # ──────────────────────────────────────────────────────────── def fig_head_split(): fig, ax = plt.subplots(figsize=(11.4, 5.0), facecolor="#fbfcfe") ax.set_xlim(0, 1) ax.set_ylim(0, 1) ax.axis("off") def block(x, y, w, h, title, lines, color, fs=11): ax.add_patch(FancyBboxPatch( (x, y), w, h, boxstyle="round,pad=0.01,rounding_size=0.02", facecolor=color, edgecolor="white", linewidth=1.6)) ax.text(x + w / 2, y + h * 0.74, title, ha="center", va="center", fontsize=fs + 2, weight="bold", color=INK) ax.text(x + w / 2, y + h * 0.33, "\n".join(lines), ha="center", va="center", fontsize=fs, color="#334155", linespacing=1.5) def arrow(x1, y, x2, label): ax.annotate("", xy=(x2, y), xytext=(x1, y), arrowprops=dict(arrowstyle="-|>", color=C_GREY, linewidth=2.0, mutation_scale=16)) ax.text((x1 + x2) / 2, y + 0.045, label, ha="center", fontsize=11, color="#475569", weight="bold") # X block(0.02, 0.30, 0.16, 0.42, "X", ["[B, N, d_model]", "d_model = H · D"], "#dbeafe") arrow(0.185, 0.51, 0.245, "三个投影") # Q/K/V 未切头 block(0.25, 0.30, 0.17, 0.42, "Q, K, V", ["各 [B, N, H·D]", "仍是完整宽度"], "#e0e7ff") arrow(0.425, 0.51, 0.485, "切头") # 切头后 block(0.49, 0.30, 0.17, 0.42, "切头之后", ["[B, H, N, D]", "reshape + transpose"], "#ede9fe") arrow(0.665, 0.51, 0.725, r"$QK^\top/\sqrt{D}$") # 分数矩阵 block(0.73, 0.22, 0.25, 0.58, "分数矩阵 S", ["[B, H, N, N]", "H 个 N×N,不是 1 个", "← 显存就花在这里"], "#fecaca", fs=11) # 底部说明 ax.text(0.5, 0.10, "每个头只用自己的 D = d_model / H 维去做内积,H 个头各算各的 N×N;\n" "最后把 H 份 D 维结果拼回 d_model,再过一次输出投影 W_o。", ha="center", va="center", fontsize=11.5, color="#475569", linespacing=1.6) ax.text(0.5, 0.02, "关键:分头不改变总算术量(H · N² · D = N² · d_model),它改变的是每个子空间的表达能力。", ha="center", va="center", fontsize=11, color=C_QUAD, weight="bold") fig.suptitle("多头注意力切的是 d_model 这一维,不是复制 H 份", fontsize=16, weight="bold", color="#0f172a", y=0.99) fig.tight_layout() fig.savefig(OUT / "head_split.png", bbox_inches="tight", facecolor=fig.get_facecolor()) plt.close(fig) if __name__ == "__main__": fig_memory_ledger() fig_flops_ratio() fig_head_split() print(f"已生成 3 张图到 {OUT}") for p in sorted(OUT.glob("*.png")): print(f" {p.name} {p.stat().st_size / 1024:.0f} KiB")
2026年09月24日
4 阅读
0 评论
0 点赞
2026-09-24
AIGC 每日速读|2026-09-24|阿里Qwen3.8-Omni发布,视频推理输Gemini
今日 AIGC 论文速览 今日共 10 篇 · 全模态基座与语音智能体 2 篇 · 视频编辑与生成智能体 2 篇 · 生成推理加速与量化 2 篇 · 视觉Tokenizer与音视频生成 2 篇 · 奖励模型与生成安全评测 2 篇 重点论文标题列表 Qwen3.8-Omni(Qwen Team, Alibaba Group):原生全模态智能体,百万上下文 VideoX-Qwen(Nanjing University):120万对指令视频编辑数据 QuantWM(Harbin Institute of Technology (Shenzhen)):2比特KV量化压住时序闪烁 Flash-dLLM(VILA Lab, MBZUAI):扩散LLM提速五倍且不伤精度 VideoGen-Agent(Princeton University):RL教会智能体用工具做视频 今日论文速览 1. Qwen3.8-Omni:原生全模态智能体,百万上下文 Qwen3.8-Omni: Towards Native Omni-Modal Agents | Qwen Team, Alibaba Group | arXiv:2609.25611 关键词:全模态智能体,稀疏MoE,百万上下文,视频剪辑,插件框架 前序问题:全模态模型此前主要强调感知和交互,真正的长程智能体任务——调工具、跑多步、把一段素材剪辑成片——要靠外挂 harness 拼,而现有 agent harness 对音频和视频没有原生支持。缺一个在多模态上原生 co-training、又不丢文本能力的底座。 本文贡献:发布 Qwen3.8-Omni-Flash:Thinker 继承 Qwen3.8-Next 的稀疏 MoE 骨干(GDN 与交错注意力混合),上下文扩到 100 万 token;原生多模态 co-training 把文本侧的智能体能力迁移到音视频任务,再经多教师蒸馏合并领域专精策略;配套开源 Qwen-MM-Plugins 插件框架与 Qwen-Live-Harness 实时交互框架。 Qwen3.8-Omni-Flash is a unified end-to-end model capable of processing multiple modalities 实验效果:WildClawBench-MM 从上代 34.5 涨到 71.0(+36.5)、AgenticVBench 14.5→36.8、OmniGAIA 57.2→74.0;Video-MME-v2 47.9→65.0。开智能体模式后 LVOmniBench 63.3→73.6,OmniVideoBench 单次 query 消耗 token 从 14.6 万降到 7.9 万(约 -45.7%)。 The overview of Qwen3.8-Omni-Flash-Realtime and Qwen-Live-Harness 批判点评:摘要口气很大,但表格里对手没输干净:Gemini 3.8 Flash 在三个视频推理基准的静态设置上全面更高,开智能体后 OmniVideoBench(70.1 vs 67.8)和 Video-MME-v2 仍被压着;teaser 图的 OmniCap-IF 一栏 Flash 拿 20.2,落在上代 Plus 的 28.3 和 Muse Spark 1.2 的 26.8 之后。所谓「领先」主要立在自己参与的新评测上。 2. VideoX-Qwen:120万对指令视频编辑数据 VideoX-Qwen: Data-Centric Instruction-Based Video Editing | Nanjing University;Jiutian Research | arXiv:2609.26015 关键词:指令视频编辑,数据合成,Qwen3-VL,Wan,成对监督 前序问题:通用视频编辑缺的不是模型而是监督:编辑要在执行指令的同时保住无关主体、场景结构和时序连续性,而成对(源视频、指令、目标视频)数据的规模和任务覆盖都远远不够,指令视频编辑长期落后于图像编辑。 本文贡献:一条可扩产的数据流水线:Qwen3-VL-235B 做源视频筛选与目标解析,加/删走 SAM3 加 Minimax-Remover,替换/属性走 Qwen-Image-Edit 首帧编辑加 WanAnimate 传播,自动通过率 89%,产出 120 万条方向性编辑对(每组任务超 40 万条)。模型侧统一 Qwen-Wan 编辑器:多模态语义条件加稠密源视频 latent 引导,图像-视频三阶段渐进训练。 Paired-video construction pipeline 实验效果:100 例对比 UniVideo 和 Kling O1:11 项指标里 9 项均值最优,覆盖指令遵循、编辑质量、内容保持、结构与感知相似度和视频分布质量。 Qualitative comparisons on compound and attribute editing 批判点评:「9/11 最优」的另一面是 11 项里输了 2 项,且对比规模只有 100 例;加/删两条数据路线直接复用 Minimax-Remover 和 WanAnimate,编辑对质量上限被这些现成生成器封顶。论文把「逐条校验富化后的指令」画成 proposed extension(虚线框),等于承认 120 万对里指令与实际编辑的一致性还没有被逐一验证。 3. QuantWM:2比特KV量化压住时序闪烁 QuantWM: Temporally Consistent 2-Bit KV Cache Quantization for World Models and Video Generation | Harbin Institute of Technology (Shenzhen);National University of Singapore | arXiv:2609.26425 关键词:KV cache量化,视频生成加速,世界模型,注意力补偿,训练无关 前序问题:视频生成和世界模型的 KV cache 是部署瓶颈,已有 2-bit 量化方法在 VBench 这类指标上「几乎无损」,但作者实测发现:指标没掉,画面却在闪——逐帧指标天生看不见时序问题。 本文贡献:训练无关、严格因果的 2-bit KV 量化框架。诊断先行:Key 的量化重建误差比 Value 小,输出退化却更大,因为小扰动就能改变 QK^T 注意力 logits、挪走 Query 选中的时序-空间 token。对症两招:QSAC 聚类联合历史 Query 敏感度与残差范围选 INT2 友好的 Key 质心;PSAC 用低秩投影沿 Query 主子空间补偿残余 Key 误差。 Overview of QuantWM's inference system 实验效果:Causal-Forcing、LingBot-World-v2、HY-World 1.5、Matrix-Game-2、Longcat-Video 五个底模上,KV cache 最高压缩 6.20×,图像与视频质量指标全面超过现有 2-bit 方法,额外开销有限,且保持严格因果可流式推理。 Visual comparison of BF16, K16V2, K2V16 and K2V2 on world models 批判点评:这篇最值钱的是诊断而不是方法:它解释了为什么「VBench 几乎无损」的结论靠不住——逐帧指标对时序闪烁盲。但时序改善的验收仍以看图为主,论文没给出时序稳定性的量化指标;6.20× 的压缩上限是在 93 帧生成这一档测的,更长序列的收益未报告。 4. Flash-dLLM:扩散LLM提速五倍且不伤精度 Flash-dLLM: IO-Aware KV Caching and Parallel Decoding for Fast, Memory-Efficient Diffusion LLMs | VILA Lab, MBZUAI | arXiv:2609.26796 关键词:扩散语言模型,KV cache,并行解码,Triton,推理加速 前序问题:扩散 LLM 靠并行去噪生成文本,但双向注意力让自回归那套 KV cache 直接失效;已有加速工作把 KV cache 和并行解码分开研究,两者合用时 GPU 显存 I/O 反而成了主要瓶颈。 本文贡献:训练无关两件套。Flash-Cache:把 QKV 投影、RoPE、cache 写入、注意力计算融进一个 I/O 感知的融合 Triton 核,砍掉冗余访存;Flash-Verify:让 dLLM 自己既当 drafter 又当 verifier 做草稿-验证式并行解码,不挂辅助模型。 Overview of Flash-dLLM 实验效果:LLaDA-1.5 上对此前最强的 Elastic-Cache 在 GSM8K/HumanEval 分别提速 5.1 倍和 11.0 倍;Flash-Cache 加 Flash-Verify 达到精度 83.02%、吞吐 210.6 tokens/s(Elastic-Cache 为 82.79%、41.7);batch 16 显存约 26GB,对照 Fast-dLLM 的 50GB 降约 48%,后者在 batch 24 直接 OOM。 Throughput and peak memory versus batch size on GSM8K 批判点评:赢面都在 dLLM 阵营内部:可扩展性图里自回归的 Llama3 在 batch 32 的吞吐仍高于 Flash-dLLM,「扩散 LLM 加速后追平 AR」还没有发生。精度与吞吐是硬权衡(80.2% 到 83.2% 对应 278 降到 186 tokens/s),宣传里 81 倍是对无 cache 基线的倍数,直接引用容易被误读成对 SOTA 的提速。 5. VideoGen-Agent:RL教会智能体用工具做视频 VideoGen-Agent: Reinforcing Video Generation Agents | Princeton University;Stanford University;UC Davis;MMLab, CUHK;BenchFlow;GWU | arXiv:2609.24997 关键词:生成智能体,强化学习,工具调用,视频生成,VABench 前序问题:视频生成模型面对需要专业知识、特定身份、物理一致或时序事件的 prompt 经常直接翻车:太极招式做不对、指定角色画不像、物理参数乱来。单靠扩大底模收益有限,缺一个会检索、仿真、检测来补信息的生成智能体。 本文贡献:六类任务上训一个共享策略:先在教师轨迹上 SFT 建立工具使用习惯(检索文本/图像、仿真、目标检测、深度估计等),再用类别感知混合奖励(工具调用合法性、任务适配度、视频质量)做 RL。配套 VABench:600 条 held-out prompt,覆盖程序知识、单/多实体身份、物理一致、场景合成、多镜头时序。 Overview of VideoGen-Agent 实验效果:VABench 上从底模 56.5 提到 75.6(+19.1 分);把生成工具换成更强的 Seedance 2.0 后,不重训智能体就到 86.1;人评 84.3% 偏好升级配置而非最强独立基线。 Qualitative comparison on Procedural Knowledge generation 批判点评:86.1 这档数字藏着归因问题:从 75.6 到 86.1 的增量完全来自换生成工具,智能体一行代码没动——agent 框架的收益上限被底模和工具质量强绑定,这一段「白捡的分」其实是底模的分。混合奖励里「视频质量」一项依赖 VLM 裁判打分,裁判偏差会直接写进策略,论文未讨论防 reward hacking 的机制。 6. Vorch-Human:一个模型统一说话人与歌声视频 Vorch-Human: Unified Multi-Task Human-Centric Generation via Long-Horizon Continuation | Vorch Team;Harbin Institute of Technology, Shenzhen;Tongji University | arXiv:2609.26117 关键词:音视频统一生成,数字人,扩散transformer,长视频,身份保持 前序问题:以人为核心的音视频生成被拆成好几个模型:语音驱动数字人、参考音色合成说话视频、外观加声音参考生成场景,各自训练各自维护。它们本质上是同一组模态、只是条件不同的任务。 本文贡献:双流音视频扩散 transformer,在传统 noisy 音频/视频接口之外增加 clean 条件 token 组,用 per-token 任务嵌入、时间位置类型、条件掩码和共享多模态 prompt 编码器,把驱动语音、音色样本、首帧、主体图统一进一个模型;两级数据管线做说话人聚类与跨片段身份串联;长视频靠冻结 latent 前缀递归、每段只生成新后缀。 Role-aware unified human-centric architecture 实验效果:5 分钟长生成保持身份一致与音画同步;GSB 人评对 LongCat-Video-Avatar-1.5:动作/姿态/表情净偏好 +46.9%,口型准确度 +47.5%,单人音视频驱动整体净偏好 +27.1%。 Raw GSB breakdown for LongCat-Video-Avatar-1.5 批判点评:同一张 GSB 图的另一半:视觉质量净偏好 -13.1%,多人场景再输 -9.2%——赢的两项全在「跟音频对齐」上,纯画质反而输给 LongCat,摘要里 strong identity preservation 的措辞没覆盖这个短板。署名主体写作「Vorch Team」,机构披露不完整,技术报告性质明显。 7. StableVQ:三招治住VQ码本训练崩溃 StableVQ: Practical Guidelines for Stable Vector-Quantized Tokenizer Training | Huazhong University of Science and Technology;KlingAI Research;South China Normal University | arXiv:2609.26774 关键词:VQ tokenizer,码本利用率,训练稳定,自回归生成,ImageNet 前序问题:共享投影 codebook 把 VQ tokenizer 的利用率推上去之后,训练稳定性成了新瓶颈:码本利用长期趴在低位、出现永久死码、甚至在 26 万步附近突然崩到 0。现有方法各打各的补丁,缺少对根因的解释。 本文贡献:把根因归结为 Encoder-Decoder 与 Codebook 的「纠缠训练」:两个子系统各自都无法独立完成任务,只能靠碰巧合作维持。三个零参数改动:Dynamic STE 按量化距离压低不可靠梯度;Region VQ Loss 让激活码向邻近未激活码按比例传播目标;Decoupled Schedule 给两个子系统各配独立学习率调度。 Three characteristic failure modes of VQ training 实验效果:ImageNet 256 重建:262k 码本 120 epoch rFID 0.92(对照 FVQ 1.29、SimVQ 3.16),16k 码本 1.13;鲁棒性指标 UR-AUC 60.59,比 FVQ 的 8.08、SimVQ 的 2.17 高一个量级;只用单个线性投影层就追平需要 ViTBlock 投影器的 FVQ。 Qualitative reconstruction comparison with a 16k codebook 批判点评:主张是「治崩溃」,但主表里 FVQ 和 SimVQ 在各自标准配置下利用率本来就是 100%,差距全在故意制造码本-token 分布错配的鲁棒性测试里拉开——这三招买的更像「最坏情况的保险」,常态训练收益主要是 rFID 那 0.2 到 0.4。下游生成只给「competitive」,FID 全表在附录。 8. Qwen-Audio 3.1:全双工误答率从73%压到13% Qwen-Audio-3.1-Realtime: Towards Reliable Agentic Voice Interaction | Alibaba Token Foundry, Alibaba Group | arXiv:2609.25176 关键词:全双工语音,语音智能体,GRPO,在线蒸馏,工具调用 前序问题:实时语音助手要同时干三件互相打架的事:理解还在演化的请求、执行动作、遵守对话规则(何时说、何时闭嘴、是否插话)。上一代在背景有人说话时 73% 的轮次会错误应答,全双工的「闭嘴」比「会说」更难。 本文贡献:Think/Act/Speak and Coordinate 三层:Core-Cocktail SFT 加 M2-OPD 多教师在线策略蒸馏迁移文本能力;自进化可执行环境加多粒度 rollout 的 GRPO 教工具调用与参数落地;全双工决策模型独立管「说不说」。另给 Voice Harness 前后台原型:对话在前台、任务在后台异步跑。 Qwen-Audio-3.1-Realtime model framework for continuous spoken interaction 实验效果:tau-Voice 半双工改编任务成功 78.4%→82.0%(Airline 64.0%→74.0%);Full-Duplex-Bench v1.5 背景语音应答率 73.0%→13.0%、恢复率 0.26→0.87;FLEURS 宏平均 WER 9.01→3.98;WebSearch1K 检索调用从 4.37 次降到 1.05 次(-76%)。 Qwen-Audio-3.1-Realtime post-training pipeline 批判点评:两个回撤写在正文里:EVA-A 的 Mean 分从 70.50 掉到 66.26,WebSearch1K 的 F1 从 60.87% 降到 58.61%——少说话、少调用是把双刃剑;Daily Chat 和 persona 交互的 Spoken 分也在微跌。摘要只报提升项,回撤要翻到结果表才看得见,读技术报告别只读开头。 9. RULER:六项Rubric奖励治SVG生成 RULER: Instance-aware Rubric Rewards for SVG Generation | Ant Group;The Hong Kong University of Science and Technology (Guangzhou);Independent Researcher;University of Oxford | arXiv:2609.25270 关键词:SVG生成,Rubric奖励,GRPO,reward hacking,视觉裁判 前序问题:自然语言生成 SVG 是没有绝对真值的开放任务:CLIP、Aesthetic 这些在自然图像上标定的标量指标迁到风格化矢量图上会误判,直接拿来当 RL 奖励还会触发 reward hacking——评测和优化同时失去可靠信号。 本文贡献:实例化 rubric 奖励:每条指令由前沿模型生成六项 rubric(语义、视觉、风格三轴),裁判 VLM 逐项给渲染 rollout 打分,加权满意度经 GRPO 优化。rubric 只从文本推导,不需要配对 SVG 真值、也不需要人类偏好标注。训练用 Qwen3-VL-8B 做裁判,与评测用的 GPT-5-mini 分开。 Qualitative comparison of scalable vector graphics generation between RULER and four baselines 实验效果:MMSVG-Illustration/Icon 上 rubric 分从 0.432/0.395 提到 0.693/0.683,超过专用 SVG 模型、追平大得多的 DeepSeek-V3;人评盲测对自家 Qwen3-8B 骨干胜率 89.7%、对 Qwen3-32B 66.1%;消融显示去掉语义轴掉 10.1%、换静态 rubric 掉 22.7%。 Ablation of rubric design on MMSVG-Illustration and MMSVG-Icon 批判点评:「追平 DeepSeek-V3」在可视化图里并非全赢:五个示例里同心圆靶一栏 DeepSeek-V3 的 0.94 高于 Ours 的 0.92;Ours 的 CLIP 分(0.28-0.35)与基线几乎无差,说明传统指标确实分不出好坏,但反过来也意味着 rubric 分数与人类直觉的对应完全押在裁判 VLM 上。人评规模只有 150 条 prompt。 10. Moderation Gap:三成有害视频骗过逐帧审核 The Temporal Moderation Gap: Text-to-Video Safety Filters Are Blind to Harm in Motion | National University of Singapore;University of New South Wales;Fuzhou University | arXiv:2609.26233 关键词:文生视频,安全审核,帧序,红队测试,评测方法 前序问题:T2V 服务的安全栈是从图像生成继承的:关键词过滤加随机抽 8 帧逐帧检查。这套栈对「伤害只存在于帧序」的视频天然失明——闯红灯、往转动的车床里伸手、误服药片,每一帧单看都无害。 本文贡献:理论上证明:任何忽略帧序的审核器只要放行某片段,就必然放行其良性重排。实测未修改的基准 prompt 就能让 32.7% 的 Sequential-Action 目标落入审核盲区,改写、拆场景、反馈式搜索都无显著改善(McNemar p≥0.12)——不需要任何 prompt 技巧就能触发。修复方向是读帧序:时序感知检测器 AUC 0.74,逐帧检查等于抛硬币。 The temporal moderation gap 实验效果:对 97 帧全部打分发现:约三分之一的交付片段只是把不安全帧藏起来,其余按帧序整体看仍有害,尽管每一帧都通过了审核;用户实验确认人能区分片段与其重排的安全性。 Three Sequential-Action gap clips 批判点评:论文顺带揭了一个测量陷阱:在搜索出的 prompt 上用它自己的渲染种子打分,会把 7.5% 的单次生成率虚标成 46.7%——这个坑对所有做安全评测的人都成立。但修复方案的 AUC 0.74 离可用尚远,且全部证据来自 Sequential-Action 一类动作伤害,跨帧累积的暗示性内容等其他类型未覆盖。 趋势观察 全模态模型开始从「能听会说」转向「能干活」 Qwen3.8-Omni 把上下文推到 100 万 token 并配好插件与实时 harness,Qwen-Audio-3.1-Realtime 则把「何时闭嘴」当成一等公民单独训练。两条线指向同一件事:多模态模型的竞争正在从感知指标挪向任务成功率,模型外面的 harness、插件和记忆管理层开始决定产品形态。 视频编辑与视频生成的瓶颈,先卡在数据而不是模型 VideoX-Qwen 用 120 万对合成编辑数据喂出一个统一编辑器,VideoGen-Agent 则用六类任务的教师轨迹加 RL 教模型调用工具。两者都在补同一个洞:指令视频任务缺少可靠的监督与反馈信号,谁的流水线能稳定产出高质量配对数据,谁就拿到上限。 加速这条线在「省显存」和「省访存」上分头推进,且都是训练无关 QuantWM 发现 2-bit KV 量化的真正受害者是时序稳定性而非单帧指标,用注意力补偿把闪烁压回去;Flash-dLLM 则把瓶颈定位到 GPU 显存 I/O,用融合核加自验证并行解码把吞吐推到 210.6 tokens/s。共同点是都不动训练,靠推理系统层的重新安排。 评测开始把「裁判本身」当成被审对象 RULER 证明静态标量指标在 SVG 上失灵,改用逐指令生成的 rubric 当奖励;Moderation Gap 则证明逐帧安全审核对帧序盲。两者都在说同一件事:在开放生成任务里,裁判的构造方式本身就是被评测对象,默认信任任何单一打分器都危险。 人工智能炼丹君 整理 | 2026-09-24 更多 AIGC 论文解读,关注微信公众号「人工智能炼丹君」 每日更新 · 论文精选 · 深度解读 · 技术脉络 微信搜索 人工智能炼丹君 或扫描下方二维码关注
2026年09月24日
6 阅读
0 评论
0 点赞
2026-09-23
AIGC 每日速读|2026-09-23|NVIDIA像素扩散FID 1.46,仍输latent
今日 AIGC 论文速览 今日共 10 篇 · 图像生成与像素扩散 2 篇 · 视频世界模型与 3D 一致 2 篇 · 视频生成加速与流式编辑 2 篇 · 动作与音频生成 2 篇 · 音视频因果与奖励评测 2 篇 重点论文标题列表 PixelDiT2(NVIDIA):像素扩散 FID 压到 1.46 WorldCrafter(ARC Lab, Tencent IEG):分钟级探索重访 PSNR 18.0 ⚡ GAE(HKUST):相机轨迹误差砍半 FVD 降 23% SparkDiffusion(Peking University, Melon Group):97% 稀疏下 265 倍单卡加速 MixiMotion(PTIT, Hanoi):一步动作对齐 0.835 逼近教师 今日论文速览 1. PixelDiT2:像素扩散 FID 压到 1.46 PixelDiT2: Representation-Grounded Pixel Diffusion Transformers | NVIDIA;University of Rochester | arXiv:2609.24919 关键词:像素空间扩散,表征先验,表示对齐,DINOv3,ImageNet ⚠️ 前序问题:像素空间扩散不用自编码器、直接在 RGB 上去噪,省掉了 latent 重建瓶颈,代价是收敛慢、最终画质长期落后 latent 扩散。作者的判断是缺一个显式表征先验:latent 扩散天然在紧凑结构化的隐空间里去噪,像素扩散却要一边学「好去噪的表征」一边学「怎么生成像素」,两件事挤在同一次训练里互相拖累。 本文贡献:提出 representation grounding:用一个冻结的预训练视觉基座模型在整个去噪过程里持续提供 per-patch 表征监督。带噪图像走像素去噪通路,同时冻结编码器产出 per-patch 特征,一个时间步条件化的 DiT 块 P_g 把它映射成 scaffold,再通过空间 AdaLN 调制扩散 transformer;REPA 只作为训练期辅助目标。另提出延迟 grounding dropout:训练初期保留 grounding,到 epoch 160 再开启 dropout。 PixelDiT2 at ImageNet-512x512: samples, architecture, and convergence. 实验效果:ImageNet-256×256 上 H/16 模型 600 epoch 达 FID 1.46、IS 301.6;512×512 上 680 epoch 达 FID 1.48、IS 295.7,两档都是像素侧最好(同表其余像素方法 1.61–3.94 / 1.78–3.95)。论文强调 epoch 预算:512 档 epoch 200 就超过 PixelDiT 850 epoch 的 FID,预算低 4.25 倍。 PixelDiT2 reaches FID 1.78 at epoch 200 and 1.48 at epoch 680. 批判点评:「补上表征先验」在像素侧成立,但把两张表横着读是另一面:256 上 latent 侧 RAE-XL FID 1.13、REPA-XL 1.29 都优于 1.46,512 上 RAE-XL 1.13 同样压过 1.48,IS 也不及 REPA-XL 306.3 与 JiT-G 306.8。也就是说「追平」只在像素扩散内部成立,跨阵营仍差约 0.3 FID,摘要里 narrowed the gap 的措辞容易被读成已经追平。另外延迟 dropout 的开启时机是在 512 档上调出来的,256 档用的是从头独立 dropout,两档训练协议并不一致,跨分辨率可比性打了折扣。 2. WorldCrafter:分钟级探索重访 PSNR 18.0 WorldCrafter: Consistent Video World Model with Implicit 3D-aware Memory | ARC Lab, Tencent IEG;Peking University | arXiv:2609.24984 关键词:视频世界模型,隐式3D记忆,长程一致性,相机可控,流式生成 ⚠️ 前序问题:视频世界模型做交互式探索时,长程一致性一直是硬伤:走出一段距离再回头,场景常常「认不出」是刚才那个地方。难点在于历史观测越来越多,而视频生成器的 token 预算是固定的,怎么把历史压进有限 token、又不依赖显式深度对应,是个两难——显式 3D 重建重且脆,简单堆历史帧则 token 直接爆掉。 本文贡献:核心洞察是「让被请求的视角来决定历史怎么压缩」。一个与视频生成器联合训练的 memory encoder 把历史观测编成紧凑的 3D-aware 表示,再由 pose-conditioned readout 按当前目标相机位姿读出固定数量的、目标视角专属 token,在去噪前注入,全程不需要显式深度对应。配合 max-coverage 历史检索、最近时间上下文和 few-step 蒸馏,实现从单张图或文本出发的流式分钟级探索。 Overview of the WorldCrafter pipeline. 实验效果:VBench 自建输入模式 725 段视频上 Overall 81.910 为全表第一(Alaya-EVOKE 81.406、SANA-WM 80.841、Echo-WM 80.211),主体一致性 82.695、背景一致性 90.589、运动平滑 98.787、整体一致性 26.745 均为第一。长程重访一致性上 MEt3R 0.166、PSNR 18.016、SSIM 0.517,相对次优 Lyra 2.0(0.334 / 14.050 / 0.390)近乎翻倍。 Qualitative long-horizon revisit comparison in a static scene. 批判点评:重访那张表真正的第一不是 WorldCrafter 而是它的加速版 WorldCrafter-fast(MEt3R 0.129、PSNR 20.868、SSIM 0.616)——少步蒸馏的变体反而在一致性上更好,这暗示「更多步数」本身可能带来累积漂移,论文没有就这条给出解释。视觉质量上也不是全胜:美学质量 AQ 61.432 只排第四,不及 SANA-WM 63.467;动态程度 DD 96.893 是全表最低(其余 97.087–100),说明换来的一致性里有一部分是靠画面动得更少买到的。评测均在自建输入模式下完成,未与公开榜单协议对齐。 3. GAE:相机轨迹误差砍半 FVD 降 23% GAE: Learning a Geometry-Native Latent Space for 3D-Consistent World Generation | The Hong Kong University of Science and Technology;ARC Lab, Tencent IEG;The University of Hong Kong;The University of Texas at Austin | arXiv:2609.24981 关键词:几何原生隐空间,3D一致性,自编码器表征,相机轨迹,感知生成统一 ⚠️ 前序问题:视觉生成器能出逼真画面,却保不住一个前后一致的 3D 场景。作者认为这不只是建模问题,更是表征问题:生成器演化的是外观为中心的隐空间,感知模型恢复几何用的是语义丰富、编码跨视角结构的空间,两者不在同一坐标系上。常见补丁是「再输出一个几何头」,等于把几何当副产品,治标不治本。 本文贡献:不把几何当额外输出,而是直接把一个几何基座模型的特征重参数化成紧凑、可扩散的隐空间——几何原生自编码器 GAE,其 latent 同时可解码到外观、深度、相机和点图。训练分两阶段:先训 codec 把冻结的多层几何特征压进 latent,再用条件流只在「待生成视角」的 latent 上做流匹配,参考视角 latent、相机射线和文本作为控制,让感知与生成共用同一个隐空间接口。 Overview of Geometry-Native Autoencoder (GAE). 实验效果:在固定生成器与训练协议的控制对比下:重建侧 GAE-128 用比原始几何特征少 24 倍的通道,PSNR 28.76 / LPIPS 0.036 / rFID 5.4(RealEstate10K)反而略优于原始 L0 特征;生成侧 GAE-64 最优,RealEstate10K FVD 225.7、PSNR 20.02、SSIM 0.711,DL3DV FVD 287.0、PSNR 18.00,相对最好的非 GAE 受控 latent 分别降 FVD 12.7% 与 23.1%;ATE 在 RealEstate10K 降 52.8%、DL3DV 降 23.3%,MEt3R 取到受控组最佳 0.1208 / 0.1347。 Multi-frame RGB, geometry, and camera-pose comparison across RealEstate10K scenes. 批判点评:最有价值的是控制变量设计,但同一张表也留了口子:受控组之外还列了 Gen3R 与 GLD 两个参考方法,其中 Gen3R 在 RealEstate10K 的 MEt3R 是 0.1157,优于 GAE-64 的 0.1208——「最佳」限于受控 latent 这个口径。DL3DV 的重投影误差上 GAE 只有 0.0028–0.0029,反而不如 WAN2.1 VAE 的 0.0019,说明几何原生 latent 在像素级回投影精度上并不占优。另外通道数选择本身就不一致:128 在重建上最好、64 在生成上最好,论文没给选择依据,实际部署还得自己再权衡一次。 4. SparkDiffusion:97% 稀疏下 265 倍单卡加速 SparkDiffusion: Mitigating the High-Sparsity Trap --- A Unified Framework for up to $265\times$ Single-GPU Acceleration of Visual Generation | Peking University, Melon Group;Tsinghua University;Alibaba Group;University of Electronic Science and Technology of China;Harbin Institute of Technology | arXiv:2609.23153 关键词:稀疏注意力,推理加速,少步蒸馏,FP8量化,视频扩散 ⚠️ 前序问题:视频扩散 transformer 贵在注意力要处理超长时空 token 序列,做稀疏注意力是自然的加速思路。但作者发现一个「高稀疏陷阱」:稀疏度推到极致时,逐步训练损失还在稳步下降,最终生成质量却停滞甚至倒退。诊断指出这是监督问题——主要误差来自高噪声的结构生成阶段,逐步局部训练修不动它,只有对齐终态的训练才修得动。 本文贡献:由此给出一条分阶段原则:先把稀疏架构适配成粗粒度先验,再校正终态分布。SparkDiffusion 按此串起三件事——短程稀疏 warm-up、少步 trajectory-mixed 蒸馏、FP8 量化配融合 kernel,并把三者做成可乘的加速因子。覆盖 Wan2.1/Wan2.2 骨干、T2V/I2V 任务、480P/720P 分辨率,以及 H100 与 RTX 5090 两种 GPU。 Overview of the SparkDiffusion framework. 实验效果:3 步无 CFG 推理下,Wan2.1-T2V-14B-720P 在单张 RTX 5090 上端到端 265× 加速(H100 上 220×);Wan2.1-T2V-1.3B-480P 端到端 1.3 秒出片。长序列 720P 上维持 97% 注意力稀疏,1.3B-480P 上 90%。VBench 总分从 dense 的 83.69 到 83.15,掉 0.54 分;同表延迟与显存项均为最好(18.0 / 8.0)。 End-to-end latency and speedup of SparkDiffusion over Full Attention on NVIDIA H100 and RTX 5090 GPUs. 批判点评:265× 是三个因子相乘的结果,其中大部分来自 3 步蒸馏而非稀疏本身,读标题很容易把功劳全记在稀疏上——论文在图表里做了分解,但正文强调的仍是合成后的倍数。质量代价是实打实的:83.15 低于 dense 83.69,且稀疏度从 90% 提到 97% 时分数从 83.42 掉到 83.15,说明高稀疏陷阱只是被缓解、没被消除。另外源码里仍能看到若干未清理的 TODO 占位注释(作者元数据、BF16 两列标注为占位),作为预印本部分数值还需等正式版复核。 5. MixiMotion:一步动作对齐 0.835 逼近教师 MixiMotion: One-Step Text-to-Motion Generation via Asymmetric Set Distillation | Posts and Telecommunications Institute of Technology, Hanoi, Vietnam | arXiv:2609.23010 关键词:文本到动作,一步生成,集合蒸馏,运动学监督,推理延迟 ⚠️ 前序问题:文本到动作生成质量上来了,但迭代采样要几十上百次网络评估,延迟高,交互式动画、游戏这类场景用不了。压缩成严格一步模型有两大障碍:一是文本条件下的动作天然多模态,同一提示对应多个合理动作,固定的教师-学生一一对应会把多种有效模式平均掉;二是只在归一化表征空间里对齐,解码后可能出现明显运动学瑕疵。 本文贡献:MixiMotion 用离线集合蒸馏:对每个提示让冻结的多步教师先生成 K=4 个动作存进离线教师库,学生训练与推理时都不再查教师;学生从独立噪声采样 M=8 个一步样本,用非对称双向集合匹配训练——教师到学生方向鼓励覆盖教师支持的多样动作,学生到教师方向抑制不被支持的生成,两侧权重 α=(1, 0.20) 故意不对称。再叠加解码空间的端点、轨迹与身体分组运动学监督。 MixiMotion: offline asymmetric set distillation for one-step text-to-motion generation. 实验效果:ViMoGen 上语义对齐 0.835,一步方法里第一(MotionLCM-1 0.825、naive 一步教师 0.711),逼近 50 步 HY-Motion-1.0-Lite 教师的 0.858;六个类别中有四类超过 MotionLCM-1。25 人盲评总体 4.33 对教师 4.50,同为一步/少步里最高(MotionHiFlow 4.20、MotionLCM-1 3.98)。单次网络评估延迟从 829.58 ms 降到 9.30 ms。 Qualitative comparisons between HY-Motion-1.0-Lite and MixiMotion. 批判点评:摘要里的 89.2× 是单次网络评估口径,正文自己给出的端到端加速是 6.75×——差了十几倍,文本编码与渲染等开销没算进去,读者很容易把 89.2× 当成真实提速。另外 0.835 其实仍略低于多步的 MLD 0.840(与 MDM 0.833 基本持平),「逼近教师」成立,但严格说还没超过最好的多步基线。消融里非对称权重只带来 0.829→0.835 的 0.006 增益,解码空间监督贡献 0.832→0.835,各组件边际收益都很小;学生从教师热启动、460M 参数与教师同量级,并没有变小。语义分数由 Qwen3-VL-30B 判 yes/no 问题自动给出,非人类标注。 6. SVEET:双向模型改流式 15 FPS Streaming Video Editing with Easy Adaptation | Shanghai Jiao Tong University | arXiv:2609.24788 关键词:流式视频编辑,双向转自回归,特征解耦,正交解耦训练,实时生成 ⚠️ 前序问题:预训练的双向视频扩散模型编辑质量好,但它需要看完整段视频,做不了逐帧自回归的流式编辑。直接改造会撞上两个矛盾:双向骨干的特征在时间上互相依赖,而流式推理要求条件帧彼此独立;可控性与因果性两个优化目标还会互相拉扯,联合训练通常顾此失彼。 本文贡献:先系统复盘已有 video-to-video 扩散方案,归纳出流式适配的两条原则:骨干特征解耦、条件帧独立。据此给预训练双向模型加一条辅助控制分支,用时间独立自注意力编码源视频,中间特征注入对应骨干块;为弥合双向与流式的特征空间差异,再提出解耦训练 ODT,显式约束「视频可控性」与「模型因果性」两个优化方向正交,使二者在推理时兼容,并实现跨异构骨干的零样本迁移。 Our proposed SVEET. 实验效果:视频风格迁移上 VLM 编辑准确率 7.4322 为全表第一(LiveEdit 5.8672、DayDream+CF 5.2321、SDV2 4.2297),CLIP 文本对齐 0.2287、主体一致性 0.9447、背景一致性 0.9298、运动平滑 0.9898、美学质量 0.5550 均为第一。消融中 ODT 把 2D 版从 7.1898 提到 7.4322、3D 版从 7.0653 提到 7.3315,推理期投影(6.4250)与两阶段 teacher forcing(7.0927)都不如它。单张 H100 上 15 FPS,未用任何额外加速手段。 Qualitative comparison results on stylization. 批判点评:唯一没拿第一的一列是整体一致性(0.1123 第二),而拿下这列的 Channel concat 编辑准确率只有 1.8902——几乎没在编辑,说明这个指标在这组对比里更接近「越不改越高分」,拿它当唯一失分项反而暴露了指标本身失效。实验只在风格迁移这一个任务上给了完整数字,其余编辑任务的优势主要靠定性图支撑。15 FPS 是在 H100 上测的,离真正的端侧实时还有距离,且论文没有给出不同分辨率或时长下的 FPS 曲线。 7. COT-TTS:听对话推断语气 MOS 4.15 COT-TTS: Audio Context-Aware Text-to-Speech with Chain-of-Thought Reasoning | The Hong Kong University of Science and Technology;Nanjing University;JIUTIAN Research, China Mobile;Peking University;China Mobile (Hong Kong) Innovation Research Institute | arXiv:2609.22697 关键词:上下文感知语音合成,链式推理,可编辑风格,级联对比,时长一致性 ⚠️ 前序问题:语音合成的说话方式现在基本靠用户显式下指令,但自然对话里该怎么说本应从上下文推断出来。现有系统要么让用户逐句指定,要么把 ASR、LLM、TTS 串成级联——参数动辄 30B 以上,而且级联转换会把原始语音里的副语言信息(非语言发声、节奏、情绪强度)丢掉。 本文贡献:把「上下文感知推理式 TTS」定义成新任务:给定历史对话音频、目标文本和参考语音,先理解上下文,再显式推理出一段中间 CoT,最后按指定音色合成目标语音,且这段 CoT 可被检查与编辑。配套建了大规模双语对话语音数据集(900 万训练样本,含 100 万高质量子集)和 800 条人工核验、源不重叠的评测基准,并给出 0.6B / 1.7B 端到端自回归模型,直接产出带情绪标注的转写、可编辑风格推理和语音 token。 Overview of the end-to-end CoT-guided autoregressive model. 实验效果:英文基准上 COT-TTS-0.6B 拿到 Human MOS 4.150 全表第一(级联 3.050–3.500),情绪一致性 0.954 第一,时长误差 1.155 秒远优于级联的 5.1–12.5 秒;中文基准上 0.6B Human MOS 4.200 同样第一。可编辑 CoT 变体在中文上拿到最高的 LLM 综合分 3.461 与历史理解 4.079。参数量只有级联系统(34–39B)的几十分之一。 Distributions of the normalized audio quality score, naturalness score, target-audio effective-speech ratio, and emotional expression intensity over 50K randomly sampled examples. 批判点评:「参数少几十倍、效果相当」要看是哪一列:UTMOSv2 2.943 低于 two-a3b-fish 的 3.138,DNSMOSPro 3.180 明显不及 two-qwen3omni-seedvc 的 4.240,WER 0.041 也高于 three-dia-a3b-fish 的 0.012;LLM 评测的 CoT 逻辑 4.558 不及 4.836,综合分 2.304 不及 3.601。真正领先的是人类整体评分、情绪一致性和时长误差三项。更值得注意的是两个反常:0.6B 的人类评分(4.150 / 4.200)高于 1.7B(4.050 / 4.150),规模变大反而变差;开启可编辑 CoT 后推理分数上去了,英文 Human MOS 却从 4.150 掉到 3.500——论文把这解释为改动把分布推离训练域,等于承认「可编辑」目前要付音质代价。 8. Common Cause:灰度视频让音频翻车 62% Common Cause, Not Cross-Attention: Blocking Visual Shortcuts in Audio-Video Generation | RIKEN iTHEMS;RIKEN AIP;South China University of Technology;Columbia University | arXiv:2609.22361 关键词:音视频生成,因果推断,反事实不变性,视觉捷径,共享表征失效 ⚠️ 前序问题:音视频联合生成器的训练数据里,「看起来是什么」和「听起来是什么」高度相关——某种材质、纹理或物体外观总与某种声音同时出现。这种相关往往是虚假的:模型可以只从外观预测声音,完全不碰「是什么事件」这个因果变量。一旦测试时外观与事件的关联被打破,模型就会直接合成错误事件的声音。 本文贡献:这是一篇受控的因果研究而非又一个模型。先建音视结构因果模型,让音频在构造上独立于视频的干扰外观,再系统检验流行解法——直接 cross-attention、共享 latent、bottleneck、无监督 shared/private 分解、忠实共享先验,证明这些「走公共因」的做法全都抓不住要害,因为共享 latent 不是干预,模型照样拿得到外观代理变量。真正堵住捷径需要对干扰变量做干预:在 SCM 与干预假设下,论文证明反事实不变性是识别因果预测器的充要条件。 The counterfactual intervention do(v:=v') made concrete (Colored-MNIST). 实验效果:在特征向量 SCM、程序化像素视频、真实图像配频谱音频、移动真实数字、条件生成器五个场景验证。反事实惩罚随因果线索变强从 263× 平滑坍缩到 2.7×,说明捷径只在因果线索难读时才被采用;Moving-MNIST 上直接模型 OOD 准确率 0.21 而本文方法 0.995。对真实预训练视频转音频模型 MMAudio 做输入干预:灰度化让 top-1 声音类别翻转 62%、色彩旋转 54%,最小编辑只翻 22%。 The failure, made audible. 批判点评:最尖锐的一段是「修不好」:作者用重着色一致性项把 MMAudio 的速度网络微调 1500 步,色彩旋转翻转率基本没动(62.5% → 63.9%),真正下降的是通用 OOD 敏感度(高斯噪声 83.3% → 70.8%)——这次微调买到的是整体更鲁棒,而不是论文想修的外观捷径,恰好反证诊断与修复之间还有距离。另外结论高度依赖 SCM 与干预假设成立,真实数据上的干预只能做「改写外观」这类近似;MMAudio 实验只有 24 条真实视频、其中 18 条用于微调,样本量偏小。 9. RewardVerse:先写准则再打分 PLCC 0.52 RewardVerse: Rubric-Guided Policy Optimization for Video Reward Modeling | Institute of Automation, Chinese Academy of Sciences;University of Chinese Academy of Sciences;The Hong Kong University of Science and Technology;Tencent | arXiv:2609.22947 关键词:视频奖励模型,评分准则中介,标量漂移,RGPO,强化学习奖励 ⚠️ 前序问题:用强化学习优化视频生成模型时,奖励模型是地基,但现有视频奖励模型直接把复杂主观质量映射成一个标量,没有显式评分准则,于是出现标量漂移:分数尺度在不同提示之间坍缩或整体平移,RL 拿这种奖励训练并不稳。 本文贡献:借鉴人工标注工程,在「评测请求」和「打分器」之间插入一个动态评分准则 rubric 作为中间表示:先显式生成评测标准,再按准则打分,给标量一个语义锚点。为训练这条协作流水线提出两阶段 RGPO:阶段一用自演化的种子 rubric 预热打分器(偏好奖励 + 格式奖励 + 校准损失),阶段二联合优化 rubric 生成器、产出随请求自适应的准则,同时持续把打分器对齐人类评分。 Overview of RewardVerse and its two-stage RGPO training. 实验效果:16 维 EvalVerse 基准上 Joint(Ours) 在 14 个维度取得最高 PLCC,宏平均 PLCC 0.520、SRCC 0.493 均为第一(次优 Q-Scorer 0.467 / 0.354);Logic 维度 PLCC 0.750 对次优 0.593,Action 0.566 对 0.390。漂移诊断上 RGPO 把均值偏差从 +0.610 降到 +0.164,PLCC/SRCC 从 0.240/0.222 提到 0.505/0.458,预测标准差从 0.457 扩到 1.332,被压缩的分数区间被重新拉开。换到闭源 Gemini-3.1-Pro 上,加 rubric 仍让 PLCC 从 0.490 提到 0.593。 Score distributions of six reward models on the shared 10-dimensional test subset. 批判点评:「14/16 维最高」是真的,但剩下两维恰好是最容易被低层画质解释的:Temporal 那列 Ours 0.538 不及 VideoScore2 的 0.697,SRCC 侧也有两列被 Q-Scorer 与 Raw-rubric 超过。也就是说在偏感知质量而非认知/时序的维度上,rubric 中介的优势会收窄——论文自己也点出 Q-Scorer 在光照、逻辑这类与低层画质相关的维度上仍有竞争力。Gemini 实验只对比了有/无 rubric 两个变体,且受限于 API 只能读自然语言浮点输出,无法与开源模型在同一训练协议下横比。此外 rubric 由模型自生成,其质量评估仍依赖同一套 VLM 裁判,存在自证循环的风险。 10. IMPLICIT-Bench:5493 条中性提示词测隐式偏见 IMPLICIT-Bench: Measuring Implicit Bias in Text-to-Image Models under Neutral Prompts | The University of Melbourne | arXiv:2609.24228 关键词:文生图偏见评测,隐式偏见,知识图谱三元组,中性提示,去偏见权衡 ⚠️ 前序问题:文生图模型的偏见评测大多用「一张 [职业] 的照片」这类槽位模板,只能孤立地探显性人口属性(性别、肤色)。它漏掉了更隐蔽的一类:在自然提示里,当与刻板印象相关的属性没有被指定时,模型仍然默认输出刻板结果——这种隐式偏见现有评测基本看不见。 本文贡献:用结构化知识图谱构造受控提示三元组:中性、刻板、反刻板三个变体只在单一偏见维度上不同,其余场景语义保持一致,从而能把偏见效应精确归因到模板类基准做不到的粒度。最终发布 5,493 条提示、覆盖 11 个偏见类别,并用多模型一致性、CLIP 验证和人工评测三重校验。基于它评测现有去偏见方法,揭示出偏见降低与语义保真之间的基本权衡。 Four-stage KG-grounded benchmark pipeline. 实验效果:过滤后保留的中性图显著偏向刻板端:VLM 均值从 1.75 升到 3.40(Qwen3-VL)、1.26 升到 2.65(Gemma-4),未参与筛选的 CLIP 验证器同向移动。最终基准上 CLIP 判定 67.7% 的 Δx 为正(Cohen's d = +0.39),Qwen3-VL 71.6%(d = +0.85),三个裁判在 Δx 层面正相关(Pearson r = 0.22–0.44)。12 名人工评审在 100 个样例上复现了同样排序:刻板 3.78 > 中性 3.17 > 反刻板 1.60,75.9% 的 KG 被认为确实反映社会刻板印象。 Cross-model qualitative comparison, set 1. 批判点评:三个裁判的一致性是这个方法最脆弱的一环:Gemma-4 的倾向比例只有 48.7%,几乎等于抛硬币,d 也只有 +0.29,与 CLIP 的 67.7%、Qwen3-VL 的 71.6% 不在一个量级;裁判间 Pearson r 仅 0.22–0.44,CLIP 与两个构建裁判的逐图相关性也只有 0.42 / 0.55。人工侧同样不完美:VLM 相对人类均值的 MAE 接近 0.97–0.99,且 12%–16% 的样例把刻板与反刻板的方向判反。样本分布也不均衡——physical-appearance 类只有 14 条,人工研究里 sexual-orientation 一行只落在单个样例上。作者自列的局限同样实在:KG schema 单轴、源数据带 WEIRD 偏斜、去偏见实验只在 Qwen-Image 上做过。 趋势观察 世界模型这一轮在比「记得住」,不是在比「画得好」 WorldCrafter 让目标相机位姿去决定历史怎么压成固定数量的 token,GAE 则干脆把几何基座的特征重参数化成生成用的隐空间。两者承认的是同一件事:分钟级的一致性不是把模型做大就能换来的,而是取决于历史和几何以什么形式存在。有意思的是 WorldCrafter 的加速版在重访一致性上反超了标准版,说明步数与一致性之间未必是正相关。 像素扩散与 latent 扩散的差距,正被「表征」而不是「架构」填平 PixelDiT2 不引入自编码器,而是让一个冻结的视觉编码器在整个去噪过程里持续提供 per-patch 表征,把表征学习和像素生成拆成两件事。不过两张主表横着读会发现,1.46 是像素侧最好、却仍不及 latent 侧 RAE-XL 的 1.13,IS 也不及 REPA-XL 的 306.3——差距在缩小,但没有消失。 加速方案的瓶颈已经从「算力」挪到了「监督」 SparkDiffusion 指出高稀疏下逐步损失在降、终态质量却不涨,根因是监督没有对齐终态;SVEET 则指出双向改流式的瓶颈不在算力而在两个优化方向互相干扰,于是显式约束它们正交。两者都不是靠新算子取胜,而是重新安排了训练时究竟该优化什么。 评测类工作开始把「裁判本身」当成被审对象 IMPLICIT-Bench 的三个裁判一致性只有 Pearson r = 0.22–0.44,其中一个几乎等于抛硬币;RewardVerse 则发现外部打分器在认知类维度上可以是负相关(VideoScore-v1.1 在 Logic 上 PLCC −0.245)。当奖励模型和偏见裁判都开始被自己的盲区反噬,评测可信度就成了需要单独论证的事,而不是默认前提。 人工智能炼丹君 整理 | 2026-09-23 更多 AIGC 论文解读,关注微信公众号「人工智能炼丹君」 每日更新 · 论文精选 · 深度解读 · 技术脉络 微信搜索 人工智能炼丹君 或扫描下方二维码关注
2026年09月23日
3 阅读
0 评论
0 点赞
1
2
3
4
...
18
粤ICP备2021042327号