专题二 · 外传 · L200 精讲版
算力强,显存弱 —— 这样一颗芯片,该配什么样的活
TPU v6e 与扩散模型。L100 那十五分钟给的是结论;这一版把每一步摊开 ——
每个数怎么除出来的、拿什么锚点验过、什么时候不成立,
以及 L100 只给了结论、这里给全过程的两件事:流水线为什么又胖又瘦,
还有我们在这条路上真跑过的十个模型。
怎么读这一页
L100 是同一批材料的另一种剪法,不是这一版的前几节 ——
它十五分钟,砍掉了显微镜那几张、换进了三阶段部署与十模型时间线。
两边的节号不再一一对应;L100 的收尾有一张
「本讲某节 → 本页某节」的对照表,拿着它跳。
⛔ 全篇仍然不做端到端 benchmark 对比,理由在 §零 说清楚。
零这一讲只回答一个问题
专题二立过一条线:算力 ÷ 显存带宽。它的物理含义是
「每从显存搬一个字节,这台机器本来能算多少次」 ——
一个纯粹由硬件规格决定的数,跟你跑什么模型无关。
那一讲量了两颗:B200 是 312,TPU v7 是 313。几乎一模一样。
当时的结论是「这一代旗舰的胃口都差不多」。
这一讲把同一把尺子多量两颗。其中一颗立刻破了那个「都一样」。
零点一 为什么这个数值得单独拿出来说
因为它是一条分界线,而分界线的用途只有一个:判某一段计算落在哪一侧。
把你要跑的那段计算也算出一个数 —— 算术强度,即
「这段计算每搬一个字节,实际算了多少次」。两个数一比:
- 强度 < 屋脊点 —— 算力在等数据。此时提升算力毫无意义,
要么加带宽,要么想办法让每个字节多干点活。
- 强度 > 屋脊点 —— 数据管够,算力是唯一瓶颈。
此时加带宽也毫无意义。
就这么两句话。这一讲剩下的全部内容,都是在给这两句话填具体的数。
⛔ 先说清这一讲不做什么,以及为什么
不比 benchmark,不谈谁跑得快。屋脊点是结构量 ——
它只说「这台机器的胃口有多大」,不说「这顿饭多久吃完」。
⚠️ 而且这不是回避:我们手上唯一一组同模型双平台实测(HunyuanVideo-1.5),
本身就不支持「v6e 更快」这个结论。
与其挑一组好看的数,不如把这一讲的边界说死 ——
它讲的是结构匹配,不是快慢。
扩散模型本身的原理(加噪、去噪、latent、VAE、DiT、CFG)在专题十一,
那是一小时的深潜。这里只取那些结构事实对硬件的后果。
一同一把尺子,量四颗芯片
图 X-1 三颗旗舰挤在 295–313,v6e 是 560。下面三小节分别回答:这个数怎么除出来的、我们凭什么信它、以及「算力强显存弱」在这四颗上具体差多少。
一点一 这个数是怎么除出来的
分子是稠密 bf16 算力,分母是HBM 带宽。两个都取官方规格表上的数,
不取任何实测值 —— 这一点很重要,它保证了这条线在跑任何东西之前就能画出来。
| 芯片 | bf16 稠密算力 | HBM 带宽 |
屋脊点 = 算力 ÷ 带宽 | 出处 |
| H100 SXM | 989.5 TFLOPS | 3.35 TB/s |
295 | NVIDIA 官方数据表 |
| B200 | — | — |
312 | 专题二已推导 |
| TPU v7 | — | — |
313 | 专题二已推导 |
| TPU v6e | 918 TFLOPS | 1,638 GB/s |
560 | Cloud TPU v6e 官方规格页 |
⚠️ H100 那个 989.5 是算出来的,不是抄的
NVIDIA 的数据表上印的是 BFLOAT16 Tensor Core:1,979 TFLOPS,
但那一行下面有一句小字:with sparsity(结构化稀疏,2:4)。
稀疏算力不能跟稠密算力比 —— 它要求权重里每四个有两个是零,
而我们这里比的所有负载都不满足这个前提。
稠密值是它的一半:989.5。
⭐ 这是这张表里唯一一处需要动脑的地方,也是最容易被抄错的地方 ——
网上大量「H100 有 1979 TFLOPS」的说法,都漏了那三个词。
一点二 我们凭什么信这条公式 —— 拿一个已知答案复现它
公式再简单,也得验。验的办法不是再推一遍,而是找一个别人已经算过的例子,
看我们的算法能不能把他的答案原样复现出来。
用的是 Google 那本公开的《How to Scale Your Model》。它在讲屋顶线的那一章里,
对 TPU v5p 给出了自己的屋脊点。我们把 v5p 的官方算力和带宽代进
「算力 ÷ 带宽」,得到的数跟它印出来的一致。
⭐ 这一步为什么不能省
我们要用这条公式去给 v6e 和 v7 下判断,而那两个数没有人可以对照。
先在一个有标准答案的输入上跑通,再拿它去算没有答案的输入 ——
这样万一 v6e 那个 560 错了,错也只可能错在输入的规格数上,
不会错在公式本身。
整个错误空间被砍掉了一半,代价是十分钟。
一点三 「算力强、显存弱」具体是多强、多弱
把 v6e 跟 H100 逐项相除,三个比值:
| 比什么 | v6e | H100 SXM |
v6e ÷ H100 | 这一项算强还是弱 |
| bf16 稠密算力 | 918 TFLOPS | 989.5 TFLOPS |
93% | 基本打平 |
| HBM 带宽 | 1,638 GB/s | 3.35 TB/s |
49% | 只有一半 |
| HBM 容量 | 32 GB | 80 GB |
40% | 不到一半 |
三项里两项是短板,只有一项打平。所以「v6e 合适」这句话,
从这里开始就注定只能是有条件的 —— 条件就是那两个短板得不参与。
⭐ 屋脊点 560 其实就是这三个比值的一句话总结:
分子基本没变,分母砍了一半,商自然涨到约两倍。
560 ÷ 295 = 1.90 —— 这个 1.9 倍,跟带宽那栏的 49% 是同一件事的两种说法。
二把两颗芯片拆开
上一节的三个比值是结果。这一节看它们是从什么样的硅片布局里长出来的 ——两颗芯片在同样的三个问题上,做了完全相反的选择。
二点一 v6e:算力集中在两个大方阵里
图 X-2 一颗 v6e:1 颗 = 1 个核 = 1 个 device,两个 256×256 的 MXU,128 MiB 片上暂存,片外那道门 1,638 GB/s。⭐ 注意右栏那三条看着像减配的规格 —— 二点四节会解释它们为什么是配套的。
三件事值得记住:
- 算力集中。整颗芯片的稠密算力就装在两个 256×256 的方阵里。
要喂饱它,你的矩阵得足够大 —— 小矩阵会让方阵大半空转。
- 片上有一整块暂存。128 MiB 的 VMEM,比 H100 一个 SM 的
共享内存大三个数量级。但它全部由编译器安排,没有硬件缓存兜底。
- 通往片外的管子窄。1,638 GB/s,不到 H100 的一半 ——
这就是上一节那个 49% 的物理来源。
⭐ 这一代没有「容量除以 2」那个坑
专题二反复强调过:TPU v7 上 1 颗芯片 = 2 个 device,
所以框架日志里按 device 报的数,换算成每芯片要乘 2、看容量要除 2。
v6e 是 1:1 —— 一颗芯片就是一个 TensorCore 就是一个 device。
跨代沿用那条换算规则会算错,这是本讲需要单独提醒的一处。
二点二 H100:同样的画法,三个相反的选择
下面这张图刻意用了跟 X-2 完全一样的画法和比例尺 ——尤其是那道「门」的宽度,两张图是按带宽等比画的,可以直接叠着看。
图 X-3 一颗 H100:算力摊成 528 个小单元(132 个 SM,每个 4 个),片上多了一整层硬件自动管的 L2,门宽一倍。⭐ 门的宽度按带宽等比 —— 跟 X-2 直接可比。
| 同一个问题 | TPU v6e 的答 | H100 的答 |
这个选择在防什么 |
| 算力怎么摆 |
集中成 2 个 256×256 方阵 |
摊成 528 个小单元 |
摊开是为了不管来什么形状都有人能接;
集中是为了大矩阵上把利用率吃满 |
| 片上谁做主 |
128 MiB 全归编译器,CMEM = 0 |
50 MB L2 由硬件自动管,另有每 SM 256 KB |
硬件自动管是为了应付「我不知道你要跑什么」;
交给编译器是为了吃透「我早就知道你要跑什么」 |
| 门开多宽 |
1,638 GB/s |
3.35 TB/s |
门宽一倍,代价是屋脊点低一半 —— 它挡的是「强度不够」的活 |
⭐⭐ 两颗芯片各自都是自洽的
这三行不是「谁做得好谁做得差」,是同一道题的两个解,各自配套。
H100 那一列全部指向「负载未知」:摊开、缓存兜底、门开大 ——
三个都是为不确定性买的保险。
v6e 那一列全部指向「负载已知」:集中、编译期写死、门只开够用 ——
三个都是把保险费省下来换算力密度。
—— 所以问题从来不是谁更强,是你手上的活属于哪一种。
⚠️ 那个 CMEM = 0 是有代价的,别只看它省下的
没有硬件缓存兜底,意味着编译器排错了就没有后手。
形状是动态的、访存模式在运行时才知道的负载,在这套设计上会很难受。
这条恰好是 §三要用的:扩散模型为什么正好把这个前提喂得满满的。
二点三 一颗装得下吗 —— 那三条「减配」其实是配套的
X-2 右栏有三条规格看着像减配:只有 4 个 ICI 口、二维环面、一个 Pod 只有 256 颗(对照 v7 是 6 口、三维、9,216 颗)。把模型的体积摆出来,这三条立刻说得通。
图 X-6 上半权重体积 + 实测配置,下半是同一颗芯片上的真实预算。⛔ 右列原来写的是「1 颗 / ≥2 颗」—— 那是拿体积 ÷ 32 GB 推的。换成实测之后立刻看出两处会猜错的地方(见图下第一条落点带)。
逻辑很直接:v7 那套三维环面、九千多颗的规格,是为「一个模型摊在几千颗上」
准备的。而扩散这一族根本不打那场仗。
不打,就不用付那个成本 —— 少两个 ICI 口、少一个维度、Pod 小一个数量级,
省下来的面积和功耗,全都还给了算力密度。这三条不是减配,是不同的题面。
⛔⛔ 「装得下」和「该用几颗」是两个不同的问题
这张图右列以前是拿权重体积除以 32 GB 推出来的颗数。换成实测之后,
两处猜错立刻显形:
① S3Diff 只有 6.6 GB,按体积「一颗绰绰有余」—— 对,但重点不在这。
我们把它摊到 8 卡做张量并行,实测反而更慢(5.46 s 对 5.28 s),预热还长 15 倍。
② Wan2.1 的 28 GB 按体积「贴边能塞一颗」——
而实测从来没人这么跑,它是 v6e-8 上 dp=1、tp=8 摊开跑的。
⭐ 前者看体积就能答,后者只能实测。
⛔⛔ 更要紧的是第三条:「28 GB < 32 GB 所以能跑」分子分母都错。
分母:32 GB 是标称不是预算 —— 我们那次 OOM 时 XLA 报的是只剩
13.10 GB。分子:要放的不只是权重,还有峰值激活,
而扩散这一族激活比权重大 —— CogVideoX-5B 权重 10 GB,
它的 VAE 解码一步却要 19 GB。
⭐ 激活随分辨率与帧数涨,权重一个字节不涨:Wan2.1 的 480P 跑得动、
720P OOM,用的是同一份权重。权重决定装不装得进,激活决定跑不跑得动。
二点四 X-2 右栏那两样只标了存在的:SparseCore 与二维环面
图 X-7 SparseCore 有两个职责:① 稀疏 / 大表(推荐模型那类);② 集合通信卸载 —— 把 all-reduce、all-gather 从 TensorCore 手里接过去,不占 MXU,跟计算真正并行。⛔ 这张图推翻重写过一次,图上原样保留了错在哪 —— 那个错法本身值得讲给学员听。
⛔ 这个错法值得单独讲一遍
初版这张图写的是「扩散模型用不到 SparseCore,它是给推荐系统那类负载准备的」。
那是错的。
错误的来源:架构文档里写着
「a primary use case is accelerating recommendation models」,
我把「主用途」当成了「用途的全集」。
⭐ 判据:「a primary use case」这种措辞本身就在告诉你「还有别的」 ——
它是一个明确的不完全枚举信号,而我把它读成了定义。
⚠️ 更该警惕的是:我当时还给这个错误配了一条像模像样的推导链
(扩散没有大 embedding 表 → 没有稀疏聚合 → 它基本闲着)。
链子每一环都对,前提漏了一半,于是整条链推向了错的地方。
—— 这是这门课里最该带走的一种自我怀疑:
推理顺畅从来不是前提完整的证据。
更正后的结论整个反过来:只要跨卡,SparseCore 就在干活。
扩散确实没有大 embedding 表,所以它走的不是第 ① 条路;
但只要模型要切开,all-gather 和 reduce-scatter 就一大堆 —— 第 ② 条路它走得很勤。
⚠️ 口径:这批材料主要是第四代 SparseCore 的语境
「集合通信卸载」那批官方材料以 v7x / 第四代 SparseCore 为主,
而 Trillium 是第三代。v6e 上这条路的成熟度我们没有实测。
如实标出来,不含糊过去 —— 这一条在图上也写着。
图 X-8 4 个 ICI 口连上下左右,边缘绕回成二维环面,16×16 = 256 颗,最远 16 跳(对照 v7 的 4×4×4 是 6 跳)。走不走这条路取决于你切不切模型 —— v6e 的互联是「够扩散这一族用」的档位,不是「不用」。
三那么扩散模型是哪一种活
要判它落在轴的哪一侧,得先算出它的算术强度。而算强度之前,得先回答一个更基础的问题:这类模型为什么是「计算密集」的?—— 拿 Wan2.1-T2V-14B 的官方配置当场算一遍,一步都不跳。
图 X-9 三步:像素 → latent → token,然后拆 FLOP。下面三小节把图上这三列各自的算式摊开,并给出一个能证伪整套公式的外部锚点。
三点一 第一步:一段视频到底有多少个 token
输入是 1280×720、81 帧。它要过两道压缩才进 Transformer:
| 这一步 | 做了什么 | 形状变成 | 依据 |
| 原始像素 | — | 81 × 720 × 1280 × 3 | 输入参数 |
| VAE 压缩 |
时间 ÷4、空间 ÷8×8,通道 3 → 16 |
16 × 21 × 90 × 160 |
vae_stride = (4, 8, 8) |
| Patch 化 |
空间上每 2×2 合成一个 token |
21 × 45 × 80 |
patch_size = (1, 2, 2) |
| 序列长度 N | 21 × 45 × 80 |
= 75,600 个 token,每个 dim = 5120 维 |
对照一下:一次 LLM 解码,一步只处理 1 个 token。差了将近五个数量级 ——
而这个差距,就是后面所有结论的源头。
三点二 第二步:这 FLOP 是怎么花掉的
一层 DiT、一个去噪步,四项开销。N = 75,600,d = 5,120,FFN 中间层 13,824。
| 这一项 | 算式 | FLOP | 占比 |
| 注意力(QKᵀ 与 AV) | 4 · N² · d |
1.17 × 10¹⁴ | 72% |
| FFN | 4 · N · d · 13,824 |
2.14 × 10¹³ | 13% |
| 自注意力的 q/k/v/o 投影 | 8 · N · d² |
1.59 × 10¹³ | 10% |
| 交叉注意力(文本只有 512 个 token) | — |
8.77 × 10¹² | 5% |
| 每层每步合计 |
1.63 × 10¹⁴ | 100% |
| × 40 层 × 50 步 = 生成一段视频 |
3.26 × 10¹⁷ FLOP |
七成花在那个 N² 项上。而 N² 意味着:token 数翻倍,这一项翻四倍。
上一小节那个 75,600 之所以是源头,原因就在这个平方上。
这里可以顺手回答「一颗 v6e 要跑多久」
3.26 × 10¹⁷ ÷ 918 × 10¹² = 355 秒,约 5.9 分钟。
⛔ 但这是算术,不是性能预测 —— 它假设算力 100% 吃满,而这做不到(见三点五)。
它的正确用法是当下界:任何低于 5.9 分钟的宣称,一定有别的事情发生了
(降步数、降分辨率、缓存复用、稀疏化)。
三点三 那一行 config 才是分水岭
上面那个「注意力占七成」完全取决于一个字段:
window_size = (-1, -1) # ← 不开滑窗,全局注意力
如果它不是 −1,而是一个有限的窗口 w,注意力就从 4·N²·d 掉成
4·N·w·d —— 从平方变成线性。整套结论会怎么变:
| 假如窗口是 | 注意力 FLOP | 注意力占比 |
每层总量 | 总算力变成原来的 |
| −1(实际情况) | 1.17 × 10¹⁴ | 72% |
1.63 × 10¹⁴ | 1 × |
| 8,192 | 1.27 × 10¹³ | 22% |
5.87 × 10¹³ | 1 / 2.8 |
| 4,096 | 6.34 × 10¹² | 12% |
5.24 × 10¹³ | 1 / 3.1 |
| 2,048 | 3.17 × 10¹² | 6% |
4.92 × 10¹³ | 1 / 3.3 |
⛔ 所以「视频模型都很重」这句话不能当依据
开了滑窗,注意力从占七成掉到占一成,整个「计算密集」的结论也就跟着塌了。
⭐ 判据:这个字段必须去 config 里看,不能凭直觉推。
同一族模型里,开不开窗口是设计者逐个模型做的选择 ——
它不是「视频模型」这个类别的属性。
三点四 凭什么信这一整套算式 —— 一个能证伪它的锚点
上面那张 FLOP 表里,任何一项的系数写错 2 倍,结果看起来都一样合理。
所以必须找一个独立的、有标准答案的量来验。
用参数量。同一份 config,把一层的权重数出来:
| 一层里有什么 | 参数量 |
| 自注意力 q/k/v/o:4 · d² | 104.9 M |
| 交叉注意力 q/k/v/o:4 · d² | 104.9 M |
| FFN:2 · d · 13,824 | 141.6 M |
| 每层合计 | 351.3 M |
| × 40 层 |
14.05 B ← 官方标称「14B」 |
⭐⭐ 这一步为什么不能省 —— 它当场抓到过一个错
参数量能对上,说明我们对「这一层里到底有哪些矩阵」的理解是对的;
而 FLOP 表用的是同一批矩阵。形状对了,系数才有意义。
⭐ 写这一节时,我另起了一个脚本复算,得到的却是 9.86 B,对不上。
查下去发现:那个脚本漏掉了交叉注意力的四个投影矩阵(少了 104.9 M/层)。
—— 是锚点抓到了它,不是我看出来的。
如果当初没设这个锚点,那个脚本会安安静静地给出一套自洽但错误的数。
孤立的、对不上任何外部量的数字,才是最该害怕的。
三点五 ⚠️ 强度够高 ≠ 算力自动吃满
到这里可以算强度了:这段计算的算术强度约等于 N = 75,600,
是 v6e 那条 560 的 135 倍。带宽完全不是瓶颈。
但落在算力侧,不等于算力就用满了。
⛔ 这里的口径要说准(2026-09-09 审计改):上面整段算的是
Wan2.1,而下面这个 MFU 出自我们 Wan2.2 的 profile 分析
(tpu/Wan2.2/docs/wan_tpu_optimization_guide.md),
而且它是单个算子的数,不是整模型的:
| 口径 | MFU |
怎么来的 |
| Splash Attention 这一个算子 | 37% |
roofline 15.974 ms ÷ 实测 43.93 ms = 36.4% |
| 整模型(优化后) | 34% |
Xprof 总览;基线是 12% |
两个数都说明同一件事 —— 六成多的算力没吃到。原因不在带宽,在别处:
⚠️ head_dim = 128,而 MXU 是 256×256
注意力是按头算的,每个头的维度是 128。而 v6e 的 MXU 收缩阵列是
256×256 —— 送进去的 K 维只有 128,方阵一半的位置是空的。
⭐ 这正是 §二点一那个「算力集中」的另一面:
集中成大方阵,好处是大矩阵上利用率高,代价是喂不满时浪费也成块。
—— 屋顶线告诉你瓶颈在哪一侧,它不告诉你那一侧用得好不好。
这是两个问题,L100 只讲了第一个。
三点六 结构上的三条:它把编译期想知道的全提前说了
三点一到三点五讲的是算术上的匹配。还有结构上的一半 ——
而这一半恰好对上 §二点二那个「CMEM = 0」的设计前提。
图 X-5 没有 KV cache、形状从头到尾不变、同一段计算原样跑五十遍。⭐ 对照左栏的自回归 LLM:它每一步 KV 都更长,形状步步在变 ——编译好的那一份,下一步就不合用了。
⭐ 两层匹配叠在一起,才是「合适」的全部含义
算术上:强度 75,600 对门槛 560,窄带宽这个短板没成为瓶颈(三点五)。
结构上:形状全静态、没有历史要养、同一段跑五十遍 ——
编译优先这套打法拿到了它最想要的输入(这一小节)。
缺任何一半,「合适」都不成立。
四把活放回那两条线上
前面每算一种负载都要拆一遍 FLOP,太慢。其实有一条心算规则,
三十秒就能把任何一种负载放到轴上。
四点一 一条能心算的规则:强度 ≈ 同一份权重被多少个「位置」共用
推导只要三行。看一个权重矩阵 W,形状 in × out,bf16 存储:
- 要搬的字节:每个参数 2 字节 → 2 · in · out
- 能算的次数:batch 里有 B 个位置,每个位置每参数 2 次浮点
(一乘一加)→ 2 · B · in · out
- 相除:强度 = B
⭐ 就这么干净:强度直接等于「有多少个位置在共用这一次权重搬运」。
它跟模型多大、多少层、什么架构全都无关 —— 那些量在分子分母上同时出现,约掉了。
四点二 于是四类负载各自落在哪
图 X-4 同一根轴,四类负载。⭐ 位置在轴上的左右,完全由上一小节那个 B 决定 —— 不需要跑就能画。
| 负载 | 「位置」是什么 | 典型强度 |
对 v6e 的 560 | 瓶颈在哪 |
| decode,batch = 1 | 1 个 token | ≈ 1 |
差 560 倍 | 带宽(算力几乎全空转) |
| decode,batch = 64 | 64 个 token | ≈ 64 |
还差 8.75 倍 | 带宽 |
| prefill,8K 上下文 | 8,192 个 token 一次过 | ≈ 8,192 |
超 14.6 倍 | 算力 |
| 扩散一步 | 75,600 个 token 一起动 | ≈ 75,600 |
超 135 倍 | 算力 |
⭐ 这张表能直接读出一条采购判据
要让 v6e 的算力不空转,decode 的 batch 得堆到 560;
同一件事在 H100 上只要 295。
—— 「屋脊点高」翻译成运维语言,就是「凑批的压力大一倍」。
凑不到那个批,买来的算力就是在等数据。
⭐⭐ 「v6e 适合扩散」的精确说法
不是它跑扩散更快。是扩散把它的短板挡在了瓶颈之外 ——
强度 75,600 对门槛 560,那根窄管子根本没参与;
于是只剩下没被挡住的那一项在起作用:算力 918 对 989.5,是 H100 的 93%。
短板不参与,长板打平 —— 这就是「合适」的全部含义。
⛔ 反过来同样成立,而且更该记住
小 batch 的 decode 落在轴的最左边,那里比的全是带宽 ——
而 v6e 只有 H100 的 49%。
同一颗芯片,在那种活上就是最吃亏的那颗。
这不是缺点,是同一个设计选择的另一面:
§二点二那三个「相反的选择」,在这里一次性结清账单。
五流水线的形状:又胖又瘦
到 §四为止,这门课一直把「跑一次扩散」当成一件事。
实际部署时它是三件事:文本编码 → 五十步去噪 → VAE 解码。
而这三件事可以拆到三台机器上跑。这一节回答两个问题:
凭什么能拆(本节前半),拆完怎么摆(后半)。
⭐ 为什么这一节属于一门讲硬件的课
因为「能不能拆」完全是一道硬件账:拆开的代价是切口上的数据要搬一趟,
而搬得起搬不起,取决于那个切口有多细。
这跟 §四那条「强度 ≈ 位置数」是同一类推理 ——
都是先把量算出来,再让量去决定架构。
五点一 管子里每一处有多粗
下面这张图把一次生成从头到尾的数据体积画在一根对数轴上。
全部以 Wan2.1-T2V-14B、1280×720、81 帧为准。
拿 Wan 2.2 图生视频当例子 —— 十个模型都是这套布局:
github.com/yangwhale/gpu-tpu-pedia/tree/main/tpu/Wan2.2
tpu/Wan2.2/
├── generate_i2v_torchax.py ← ① 一体化:一个进程从头跑到尾
├── generate_diffusers_i2v_torchax_staged/ ← ② 三阶段
│ ├── stage1_encoder.py 文本 + 首帧 → embedding
│ ├── stage2_transformer.py 五十步去噪 → latent
│ ├── stage3_vae_decoder.py latent → 成片
│ ├── utils.py 落盘 / 读盘的 helper 全在这
│ └── stage_outputs/ ← ⭐ 切口就在这个目录里
│ ├── stage1_embeddings.safetensors ← 段与段之间唯一交接的,就这三个文件
│ ├── stage2_latents.safetensors
│ ├── generation_config.json
│ └── output_video.mp4 成片
├── docs/wan_tpu_optimization_guide.md 约 1,970 行迁移与优化指南
└── README.md
⚠️ 下面这张图上的字节数取自 Wan2.1 那一份同名目录
(tpu/Wan2.1/generate_diffusers_torchax_staged/stage_outputs/)——
两个模型目录结构相同,数不同。
图 X-10 三阶段之间只交接三个文件。⛔ 这张图推翻重写过一次:初版画过一根 457 GB 的「注意力分数矩阵」柱子,而那个矩阵从来没有被物化过(Flash / Splash 分块算,算完即弃)——把不存在的量画成「最粗处」,整张图的比例尺就锚在了虚构上。
三个数值得单独读:
| 读什么 | 多少 | 怎么来的 |
| VAE 压缩比 | 46.3 × |
223,948,800 个像素数 ÷ 4,838,400 个 latent 数
= 空间 64 倍 × 时间 3.86 倍 ÷ 通道摊薄 5.33 倍 |
| 跨机要传多少 | 19.4 MB |
DiT 吐出来的那份 latent。只有它需要过网 |
| 传它要多久 | 1.5 ms |
19.4 MB 走 100 Gbps;就算走 1 Gbps 也只要 155 ms |
⭐⭐ 于是「能不能切」这个问题根本不成立
被切开的那一段本身要算 229 秒,而搬一趟切口数据是 1.5 毫秒 ——
切口成本是计算量的万分之零点七。
⭐ 更值得记的是为什么最细的那一处正好就是该切开的那一处:
往前是 774 MB 的激活,往后是 448 MB 的像素,
而 latent 自己只有 19.4 MB —— 扩散模型的形状天然把切口摆在了明处。
⭐ 那个 19.4 MB 顺手把形状也验了
16 × 21 × 90 × 160 个数 × 4 字节(fp32)= 19,353,600,
加上 safetensors 的 272 字节文件头,正好等于文件的 19,353,872 —— 一个字节不差。
跟 §三点四那个 14.05 B 是同一种手法:用一个独立可测的量,去验一整套推导。
这次连误差都没有,因为字节数是精确的。
⛔ 这张图有两处极易读反
① 它画的是体积,不是时间。腰细不代表那一段轻松 ——
恰恰相反,产出那 19.4 MB 的 Stage 2 占了全程 98% 的时间。
② ⛔ 初版这里画过一根 457 GB 的「注意力分数矩阵」柱子,已删除。
那个矩阵从来没有被物化出来过 —— Flash / Splash Attention 分块算,算完即弃。
把一个不存在的量画成「最粗处」,整张图的比例尺就锚在了虚构上;
而旁边那句「从不落地」的注解抵消不了图形本身的断言。
五点二 三段的胃口完全不一样 —— 这才是分开部署的真正理由
切得动只是可行性。真正的收益来自另一件事:这三段吃的根本不是同一种资源。
图 X-11 上半是三段的资源画像,下半是两种部署拓扑。⭐ 左边一体化那栏的问题不是「慢」,是按最馋的那一段配机器,另外两段的钱就白付了。
- ① 文本编码 —— 几乎不吃算力,就是把权重读一遍。占全程 1.3%。
- ② DiT 去噪 —— 吃算力。占全程 98.3%。§三算过的那 3.26 × 10¹⁷ FLOP 全在这儿。
- ③ VAE 解码 —— 吃显存峰值。把 480 万个数展开成 2.24 亿个,占全程 0.4%。
⭐ VAE 那一段有个特别的形状:编译极贵,运行极便宜
720P 下 预热要 80 秒,之后每次只要 1 秒 —— 80 倍差。
这个形状天生该做成常驻服务:编译好放着,谁要解码谁来调,
那 80 秒摊到成千上万次调用上等于零。
而一体化脚本每起一次进程就得重付一遍。
—— 注意这是一条纯粹由「编译期做主」推出来的部署结论。
§二点二那个 CMEM = 0 的设计,代价在这里,收益也在这里。
五点三 ⛔ 一条顺口的结论,和一个当场推翻它的反例
看完上面那张图,最容易得出的结论是:「文本编码那么轻,放 CPU 就行。」
这条在 Wan2.1 上成立 —— 它用 T5,那一段只要 3 秒。
但 Flux.2 立刻推翻它。
| 模型 | ① 文本编码 | ② 去噪(TPU) |
③ VAE | 最慢的是哪一段 |
| Wan2.1(T5) | 3 秒 | 230 秒 | 1 秒 |
去噪 —— 符合直觉 |
| Flux.2(Mistral3,放 CPU) | 30 秒 | 13.5 秒 |
1.5 秒 | 文本编码 —— 反过来了 |
⛔ 照顺口结论摆,等于把瓶颈从 TPU 挪到了 CPU 上
Flux.2 那个「轻量」的第一段,反而是全程最慢的一段。
⭐ 判据不是「哪一段天生该放哪」,而是
「先量一量它在目标硬件上要多久」。
而分段的价值恰恰在这里 ——
拆开之后,每一段的账才第一次能单独算清楚。
一体化脚本里,你连「文本编码花了 30 秒」这件事都看不见。
五点四 一个节点摆几路:先看塞不塞得下,再看它够不够忙
一台 v6e-8 是 8 颗 × 32 GB = 256 GB。同一台机器,两种截然不同的用法:
| 模型 | bf16 权重 | 单颗放得下吗 |
于是怎么用这 8 颗 |
| Wan2.1-14B | 28 GB |
⚠️ 贴边,只剩 4 GB 余量 |
TP 摊到 8 颗,每颗 3.5 GB —— 40 个头 ÷ 8 = 每颗 5 个头 |
| SDXL | 7 GB | ✅ 轻松 |
开 8 路各生成各的(数据并行)—— 实测 2.40 张/秒 |
⭐⭐ 两行是同一句话
装不下就把一个模型摊开,装得下就多放几路。
⭐ 而分段之后,这个决定是按段做的:
DiT 那一层摊开、VAE 那一层多放几路,互不牵扯 ——
一体化脚本里这两个决定被绑死成了一个。
这就是「分三阶段有利于部署」这句话的全部具体含义。
六这套说法,我们自己验过吗
前五节讲的全是应该怎样:屋脊点该怎么算、活该落在轴的哪一侧、
流水线该在哪儿切。这一节讲实际怎样。
—— 同一批道理,在十个真模型上跑了五个月之后,留下了什么。
六点一 五个月,十个模型
图 X-12 每一行的起止与提交数都是 git 提交历史直接数出来的。⭐ 顶部三条色带是移植方法的更替 —— 注意 12 月 10 日那道红线:那天手写 Flax 的文件被删除,第三代从那一刻算起。
值得单独指出的是提交数的走向:
| 时期 | 代表模型 | 提交数 | 在做什么 |
| 前两周 | HunyuanVideo-1.5 | 73 |
发明方法 —— 同时在试三条移植路线 |
| 头一个月 | CogVideoX / Wan2.1 | 49 / 44 |
方法定型,并回头反哺老模型 |
| 后三个月 | SDXL / Real-ESRGAN / Flux.1 | 9 / 3 / 2 |
套用方法 —— 照着 README 抄一遍 |
⭐⭐ 提交数掉下去,才是一条工程路线走通的标志
差别不在模型难度 —— SDXL 和 Flux.1 都不比 CogVideoX 简单。
差别在于前面几个是在发明方法,后面几个是在套用方法。
—— 所以这一节的落点不是「我们做了十个模型」,
而是「接第十一个模型的成本,已经不是前十个的量级了」。
⚠️ 提交数只能这么用,别过度解读
它反映的是改动次数,不等于工作量,更不等于难度。
这里拿它做的唯一推断是「发明 vs 套用」的量级差(73 对 2),
不用它比较任意两个模型谁更难。
六点二 为什么又是 JAX 版、又是 torchax 版
这是客户最常问的一条。而它讲道理是讲不赢的 ——
「用 PyTorch 生态」和「用 JAX 原生」,两边都能说出一串听起来对的理由。
我们的答法是:三条都走完,把对照版原样留在仓库里。
图 X-13 三代移植方法逐项对照,下半是 torchax、Flax NNX、纯 JAX 的四项实测比较。⭐ 最后一项生态兼容才是决定性的 —— 前三项说的是快慢,它说的是下一个模型接不接得进来。
⭐ torchax 为什么不需要重新 tracing
Flax 和纯 JAX 都要把整个函数 trace 一遍再交给 XLA。
torchax 走的是 PyTorch 自己的 C++ dispatcher:
调用 conv3d 时,dispatcher 认出算子类型,
路由到 torchax 注册的后端,直接落到对应的 JAX 实现。
—— 它不是「又一个 tracing 框架」,它是在算子这一层换了后端。
⛔ 但别把这张图读成「第二代是弯路」
没有第二代,我们不会知道该往哪个方向优化。
分片策略、内存布局、数值对齐这些认识,全是在手写那一版里长出来的 ——
一开始就用 torchax,这些东西会被框架挡在视野之外。
⭐ 而且那批代码今天仍在服役:torchax 版算出可疑结果时,
拿纯 JAX 版跑同一个输入对一遍,是最快的定位手段。
判据:一条被换掉的路线值不值得留下,取决于它还能不能当参照物 ——
而不是取决于它现在跑得快不快。
六点三 这十个模型跟前面五节是什么关系
不是「附录里的项目清单」。每一个都在验证前面某一节的某一条断言:
| 前面哪一节说过 | 哪个模型验的 | 验出了什么 |
| §二点三 权重体积决定要几颗 | SDXL vs Wan2.1 |
7 GB 的开 8 路,28 GB 的摊到 8 颗 —— 同一台机器两种用法 |
| §三 扩散是计算密集 | Wan2.1 720P |
强度 75,600;Wan2.2 上实测算子 MFU 37% / 整模型 34%
—— 落在算力侧,但没吃满 |
| §三点五 MXU 喂不满会浪费成块 | Wan2.1 |
head_dim 128 对 MXU 256 —— 一半位置空着 |
| §五点一 腰在 latent 那一处 | 全部七个三阶段模型 |
无一例外,中间产物都是全程最小的那一份 |
| §五点三 编码器放哪要先量 | Flux.2 |
反例:Mistral3 放 CPU 30 秒,比它的 DiT 还慢 |
| 整套框架对非 Transformer 成不成立 | Real-ESRGAN |
纯卷积网,8.8M 参数 —— 切块之后比 L40S 快 1.9 倍 |
⭐ 最后一行是特意留的
Real-ESRGAN 是十个里唯一一个非 Transformer、非扩散的模型
(纯卷积,只有 8.8M 参数)。
放它进来的目的不是多凑一个,是看这套分析框架在完全不同的架构上还成不成立。
—— 答案是成立的:它同样落在算力侧,同样靠切块提高每字节的计算量,
只是「位置」从 token 变成了 tile。
收尾落点
—— 选芯片不是选「更强的那颗」,是先看你的活落在那根轴的哪一边。
落右边(扩散、prefill、大 batch)比的是算力,v6e 打平;
落左边(小 batch decode)比的是带宽,v6e 吃亏。
而这两件事,从两颗芯片的布局图上就能提前看出来 —— 不用等跑完。
⚠️ 这一讲刻意留白的地方
没有端到端性能数。本讲讲结构匹配,不讲快慢 —— 两件事需要的证据不一样,
后者要实测,而实测要连口径一起给。
想看扩散模型本身:专题十一。想看这套分析框架怎么来的:专题二。
⭐⭐ 这一版比 L100 多给了什么
L100 给的是结论:560、75,600、短板被挡在瓶颈之外。
这一版多的是四样东西 ——
① 推导链(§一点一、§三点二):每个数怎么除出来的;
② 验证(§一点二、§三点四、§五点一):拿什么外部锚点验过它,
以及锚点当场抓到过什么错;
③ 反例(§三点三的 window_size、§三点五的 MFU 37%、§五点三的 Flux.2):
这些结论什么时候不成立;
④ 落地(§五、§六):从「该怎么切」到「真的切了,十个模型五个月」。
⭐ 如果只带走一条,带走 ③ ——
一套只有顺例的说法,第一次遇到边界就会碎;
而知道边界在哪的人,才敢把它用在没见过的模型上。