所属方向:表征 | 难度:进阶 | 前置知识:变分下界与重参数化(ELBO)、VAE 结构与训练目标(知道「重参数化」和「KL 项怎么来的」就够了)
关键词:VQ-VAE、VQGAN、码本、straight-through、码本坍缩、感知损失
先算一笔账,这笔账是 VQGAN 那篇论文(Taming Transformers, 2012.09841)的全部动机。把一张 256×256 的 RGB 图像当成序列直接做自回归:每个像素三个通道各算一个 token,序列长度是 256 × 256 × 3 = 196,608。自回归 Transformer 的注意力矩阵是序列长度的平方,也就是 3.87 × 10¹⁰ 个元素——一张图都喂不进去,更别说训练了。
我按这个口径把不同下采样率的账都算了一遍(token_budget.py,本机实跑):
| 下采样率 f | 256² 图像的 token 数 | 注意力矩阵规模 | 相对像素级 |
|---|---|---|---|
| 像素级 RGB | 196,608 | 3.87 × 10¹⁰ | 1 |
| f=4 | 4,096 | 1.68 × 10⁷ | 4.3 × 10⁻⁴ |
| f=8 | 1,024 | 1.05 × 10⁶ | 2.7 × 10⁻⁵ |
| f=16 | 256 | 6.55 × 10⁴ | 1.7 × 10⁻⁶ |

这张图要看的是两件事:左图里三条 tokenizer 曲线与像素级虚线之间的纵向鸿沟——f=16 时 256² 图像只有 256 个 token,是像素级序列的 1/768;右图是同一个事实在注意力开销上的投影,柱子从 3.87 × 10¹⁰ 掉到 6.55 × 10⁴,省了约 59 万倍。分辨率越高鸿沟越陡(1024² 图像在 f=16 下是 4,096 个 token),视频更极端:5 秒 24fps 共 121 帧,时间下采样 4 倍、空间 f=16 时只要 7,680 个 token。没有这个压缩,自回归视频生成连第一步都迈不出去。
但「短」只是必要条件,自回归还要求「离散」。语言模型之所以好训,是因为下一个 token 的预测是 K 分类交叉熵——一个良性的、方差可控的目标。连续潜变量上做自回归就得给每个位置配一个连续密度(混合密度网络那一路),训练病态且难以和大模型基建兼容。把潜变量离散成「码本里的编号」之后,图像生成就变成了「图像版的语言建模」:DALL-E 用 8192 个码字的 dVAE 加 32×32=1024 个 token 的自回归,VQGAN 用 16×16=256 个 token 的自回归,后来的 MaskGIT、VAR 走的都是这条路。
不过离散化是有代价的,而且真正容易被忽视的代价不在「码本有多大」,在「码本用得怎么样」。我在 8×8 小块的玩具上训了一个 K=512 的 VQ 自编码器(collapse_lab.py,实跑):512 个码字里只有 17 个被用到,96.7% 的码一次都没被选中;码本标称 9 bit,实际只用出去 log2(14.50) = 3.86 bit。这篇就讲三件事:量化怎么写进损失函数、码本怎么训练才不会死、以及为什么重建损失必须从 MSE 换成感知损失加对抗损失。
三句话:
VQ-VAE(Neural Discrete Representation Learning, 1711.00937)的名字里有 VAE,推导起点也确实是 ELBO:
$$\log p(x) \ge \mathbb{E}_{q(z \mid x)} \big[ \log p(x \mid z) \big] - \mathrm{KL}\big( q(z \mid x) \,\Vert\, p(z) \big)$$
每一项的含义:$q(z \mid x)$ 是编码器给出的「后验」,$p(z)$ 是我们先验地相信 latent 该有的分布,$p(x \mid z)$ 是解码器。VAE 里这三样都是连续分布,KL 项把后验往先验上压,重参数化让采样可导。VQ-VAE 把这三样全换了:
把这个确定性后验和均匀先验代进 KL:$q$ 在 $e_k$ 处为 1、其余为 0,求和只剩被选中那一项:
$$\mathrm{KL}\big( q \,\Vert\, p \big) = \sum_{j=1}^{K} q_j \log \frac{q_j}{p_j} = 1 \cdot \log \frac{1}{1/K} = \log K$$
$\log K$ 是一个常数,对梯度没有任何贡献。VQ-VAE 的 KL 项就此消失——没有 KL、没有重参数化、没有「均值方差都被压向先验」的正则,ELBO 退化成「重建项 + 两个逐样本的 L2 距离项」。名字里的 V 是历史包袱,这也是后面 08 节第一条误解的来源。
量化操作本身是:
$$z_q = e_k, \quad k = \arg\min_{j} \, \Vert z_e - e_j \Vert_2^2$$
符号含义:$z_e \in \mathbb{R}^d$ 是编码器输出(e 指 encoder),$e_j \in \mathbb{R}^d$ 是第 j 个码字,$z_q$ 是替换后的向量(q 指 quantized)。问题出在 $\arg\min$:它是分段常数函数。训练中把 $z_e$ 微动一点点,只要不跨过两个码字的垂直平分面,选中的 $k$ 根本不变,$z_q$ 不变,重建损失也不变。
这不是「梯度小」,是梯度恒等于零。我在训好的 VQ-AE 上用有限差分实测过(vq_core.py 的 [D] 段,512 个样本,沿随机单位方向扰动 $z_e$):
| 扰动步长 eps | 有限差分 $\partial L/\partial z_e$ | argmin 保持不变的样本比例 |
|---|---|---|
| 10⁻¹ | −7.6 × 10⁻⁶ | 0.9941 |
| 10⁻² | 0.0(精确为零) | 1.0000 |
| 10⁻³ | 0.0(精确为零) | 1.0000 |
| 10⁻⁴ | 0.0(精确为零) | 1.0000 |
eps=10⁻¹ 那一行有 0.59% 的样本跨过了平分面,所以差分不为零——这也正是 argmin「分段常数、边界处跳变」的直接展示。而在平分面之间的整片区域里,重建损失对编码器没有任何梯度。VQ-VAE 的解法是 straight-through:假装量化是恒等映射,把解码器对 $z_q$ 的梯度原封不动地抄给 $z_e$:
$$\frac{\partial L}{\partial z_e} \mathrel{:=} \frac{\partial L}{\partial z_q}$$
要强调的是:这不是真实梯度的估计,是替代品。真实梯度是 0,直通梯度实测平均范数 1.49 × 10⁻⁴(同一组权重),它携带的信息是「如果量化不存在,往哪边调编码器能让重建更好」。它能工作的原因是:编码器的真正职责不是让 $z_e$ 落在哪个精确位置,而是让「选出来的码字」是对的——直通梯度恰好只优化这件事。
实现上只有一行(taming-transformers 的写法,见 05 节):z_q = z + (z_q - z).detach()。前向时 $z_q$ 是真的码字,反向时梯度绕过 $(z_q - z)$ 这个常量直接流向 $z$。
直通梯度有个致命遗漏:码本自己拿不到任何来自重建损失的梯度。$z_q$ 是查表查出来的,对 $e_j$ 的导数被查表操作挡住了;直通又把全部梯度引向 $z_e$。如果什么都不加,码本会永远停在初始化的位置。我在玩具上实测过这条「什么都不加」的路线(vq_core.py [E] 段,K=64,800 步):码本位移精确为 0.000000,重建 MSE 0.185734,是正常训练的 3 倍多。
所以码本必须有独立的更新机制,VQ-VAE 给了第一条路——字典损失。把编码器输出当成常数(stop-gradient,记作 sg),把码字往它身上拉:
$$L_{\text{codebook}} = \big\Vert \mathrm{sg}[z_e] - e_k \big\Vert_2^2, \qquad \frac{\partial L_{\text{codebook}}}{\partial e_k} = 2\,(e_k - z_e)$$
第二条路是 EMA(VQ-VAE-2 之后的主流):不用梯度,直接对「每个码字被选中样本的均值」做指数滑动。记第 t 步里码字 j 被选中了 $n_j$ 次、被选中样本之和为 $s_j$,则
$$c_j \leftarrow \gamma c_j + (1-\gamma)\, n_j, \quad m_j \leftarrow \gamma m_j + (1-\gamma)\, s_j, \quad e_j \leftarrow \frac{m_j}{c_j + \epsilon_{\text{smooth}}}$$
$\gamma$ 是滑动系数(taming 里 decay=0.99),$c_j$ 是每个码字的滑动计数,$m_j$ 是滑动累加的样本和,$\epsilon_{\text{smooth}}$ 用来防止某个几乎没人用的码字除以接近零的数(taming 的平滑是 $(c_j + \epsilon) \cdot n / (n + K\epsilon)$,其中 $n$ 是全部计数之和)。EMA 的本质是把码字变成「最近被它编码过的那些向量的滑动平均」,没有学习率要调,这也是它取代字典损失的原因。
现在把绳子接上另一头。EMA 把码字往编码器输出上拉,但没有任何东西把编码器输出往码字上拉——编码器完全可以漂走,让量化误差 $\Vert z_e - e_k \Vert^2$ 失控。commitment 项就是拴住编码器的那根绳:
$$L_{\text{commit}} = \beta \,\big\Vert z_e - \mathrm{sg}[e_k] \big\Vert_2^2$$
注意方向和字典损失正好相反:字典损失动了 $e_k$($z_e$ 被 stop-gradient 冻住),commitment 动了 $z_e$($e_k$ 被冻住)。两项合起来,VQ-VAE 的完整训练目标是:
$$L = \underbrace{\log p(x \mid z_q)}_{\text{重建,经直通传给编码器}} + \underbrace{\big\Vert \mathrm{sg}[z_e] - e_k \big\Vert_2^2}_{\text{码本项或 EMA}} + \underbrace{\beta \big\Vert z_e - \mathrm{sg}[e_k] \big\Vert_2^2}_{\text{commitment}}$$
$\beta$ 不是一个可以随手抄的超参。在同一组玩具权重上,我把直通梯度和 commitment 项的梯度范数都量了一下(vq_core.py [D] 段):直通梯度平均范数 1.49 × 10⁻⁴,commitment 梯度平均范数 6.40 × 10⁻²,差 429 倍。$\beta$ 扫描的实测结果(collapse_lab.py [A] 段,K=64,1200 步):
| $\beta$ | 0 | 0.05 | 0.25 | 1.0 | 4.0 |
|---|---|---|---|---|---|
| 重建 MSE | 0.0673 | 0.0484 | 0.0600 | 0.0387 | 0.0385 |
| 活跃码数 | 14 | 14 | 16 | 17 | 18 |
在这个玩具上 $\beta=1.0$ 反而比论文默认的 0.25 好 35%。原因就藏在那个 429 倍里:当重建梯度相对太弱时,加大 $\beta$ 相当于在帮编码器「站稳」在码字附近,量化误差随之下降。$\beta$ 的最优值和你的重建损失量纲绑死——这就是为什么换损失(比如 VQGAN 换成 L1 + LPIPS)之后不能照抄别人的 $\beta$。
VQ-VAE 的重建损失是逐像素的,这有一个数学上无解的毛病(06 节用实验展开):最优解是条件均值,纹理会被平均掉。VQGAN 把重建侧换成三件套:
$$L_{\text{VQGAN}} = \underbrace{\big\Vert x - \hat x \big\Vert_1 + \lambda_{\text{lpips}} L_{\text{LPIPS}}(x, \hat x)}_{\text{感知重建}} + \underbrace{\lambda_{\text{GAN}} \big( -\log D(\hat x) \big)}_{\text{对抗}} + \underbrace{\lambda_{\text{cb}} L_{\text{codebook}} + \beta L_{\text{commit}}}_{\text{量化}}$$
$\hat x$ 是解码器输出,$D$ 是 PatchGAN 判别器,LPIPS 是在 VGG16 五层特征上算距离再加一层学出来的 1×1 卷积(lpips.py 里五个 NetLinLayer,实读源码确认)。对抗项的权重不是手调的,而是自适应的:对解码器最后一层分别求重建损失和对抗损失的梯度,取范数比 $\lambda_{\text{GAN}} \leftarrow \Vert \nabla L_{\text{rec}} \Vert / \Vert \nabla L_{\text{GAN}} \Vert$,让两边的梯度量级匹配——GAN 一开判别器就抢梯度主导权,这是压住它的办法。
三个从粗到细的指标,别混用:
两个容易踩的坑。第一,perplexity 是分布层面的量:它下降只说明使用分布变尖了,既不告诉你死码落在哪,也不等价于重建质量——要看死码分布得直接画使用次数的直方图(图 4 干的就是这件事)。第二,taming 是在当前 batch 上算 avg_probs 的,batch 越小噪声越大;小 batch 上看到 ppl 上下抖动不等于坍缩,别急着改 decay 或加重启。
最小实现不需要网络,一个「线性编码器 + VQ + 线性解码器」就能把所有机制跑出来。数据是我合成的 8×8 小块:16 个低频类心,类内加连续变化和白噪声(make_patch_data,4096 个样本,种子 0,全部结果可复现)。核心的量化与直通就这几行(完整脚本在附录):
def quantize(z_e, codebook):
d2 = ((z_e[:, None, :] - codebook[None, :, :]) ** 2).sum(-1) # [N, K]
idx = d2.argmin(axis=1) # 最近邻
return codebook[idx], idx, d2[np.arange(len(idx)), idx]
# 训练循环里(z_q 是查表结果,e_k = z_q):
x_hat = z_q @ params["Wd"] + params["bd"]
dxh = 2.0 * (x_hat - x) / (len(x) * d_in) # 对 x_hat 的梯度
dz_q = dxh @ params["Wd"].T # 解码器传给 z_q 的梯度
dz_e = dz_q + beta * 2.0 * (z_e - z_q) / (len(x) * d_lat) # 直通 + commitment
# 码本走 EMA(onehot 统计被选中的次数与样本和,见附录 vq_core.py)
第一件事:误差到底由哪几块组成。 用同一组训好的权重,把「走不走量化」作为开关(vq_core.py [C] 段,K=64,d=8):
| 通路 | 重建 MSE/像素 |
|---|---|
| 连续自编码器基线(无 VQ,单独训练到收敛) | 0.011082 |
| 同一组 VQ-AE 权重,绕开量化(直接用 $z_e$ 过解码器) | 0.023360 |
| 同一组 VQ-AE 权重,正常走量化 | 0.060345 |
两个结论都值得停下来想。第一,量化把误差从 0.0234 推到 0.0603,多出来的 0.0370(+158%)全是量化的账。第二,VQ-AE 的连续通路 0.0234 比独立训练的连续基线 0.0111 差了一倍——量化还会反过来把编码器带偏:commitment 在拉编码器,编码器为迁就码字牺牲了一部分子空间的质量。这部分「隐性代价」在只报一个重建指标时完全看不见。
第二件事:码本容量 K 的收益被码本维度 d 卡死。 固定一个训好的连续编码器,对它的潜变量做 k-means(k-means 就是「给定 K 个码字的最优最近邻量化器」的近似),在留出集上测失真随 K 的变化:
| 码本维度 d | 理论斜率 −2/d | 实测斜率(K≥16 段) | K 从 8 加到 512,失真降多少 |
|---|---|---|---|
| 2 | −1.00 | −0.87 | 66.5 倍 |
| 4 | −0.50 | −0.40 | 33.8 倍 |
| 8 | −0.25 | −0.28 | 13.9 倍 |
| 16 | −0.125 | −0.31 | 12.9 倍 |

这张图要看的是实测线(实线)和理论斜率(点线)的贴合程度:d=2、4、8 都贴得不错(量化理论的 Zador 渐近:最优失真 $\propto K^{-2/d}$),d=16 在大 K 端因为每个码字只剩 8 个训练样本而偏离。直白地说:d=8 时把码本从 8 加到 512(64 倍),量化失真只降到 1/13.9——这就是为什么后面 LFQ、FSQ、残差量化都要在「维度」上做文章,而不是无脑加码本。
第三件事:码本坍缩与两副解药。 K=512、随机初始化、EMA 更新,训练 1200 步:
| 配置 | 重建 MSE | 活跃码数 | perplexity | 有效 bit(log2 ppl) |
|---|---|---|---|---|
| 随机初始化 | 0.049067 | 17 / 512 | 14.50 | 3.86 |
| k-means 初始化 | 0.019695 | 325 / 512 | 217.26 | 7.76 |
| 死码重采样重启 | 0.018368 | 500 / 512 | 426.71 | 8.74 |

这张图要看的是三条线的起点和走势:红线(随机初始化)从第 200 步起就钉死在 20 以下——坍缩发生在训练极早期,一旦码字没被选中过,它就再也没有机会被选中(EMA 的计数是零,均值是零除零);蓝线(k-means 初始化)起点就是 261;绿线(死码重采样,把长期没人用的码字随机替换成真实的编码器输出)稳定在 470 上下。

这张图要看的是横轴超过 20 之后红线的缺席:随机初始化的码本只有 17 个码被用过,剩下的在 log 轴上根本画不出来;绿线的使用占比分布虽然仍是长尾(最热的码占 0.13%),但整条尾巴被抬起来了。两个数字合起来读:坍缩让 9 bit 的码本只传出 3.86 bit,而这是在重建 MSE 上实打实付了 2.7 倍代价换来的(0.0491 对 0.0184)。
顺带一提,如果不用 EMA 而用字典损失的梯度更新码本,学习率极其敏感(码本是普通 SGD,不走 Adam):在同一个玩具上把码本学习率从 0.02 加到 10,重建 MSE 从 0.2418 降到 0.0538,活跃码从 8 涨到 14——始终追不上 EMA。这就是主流实现全部转向 EMA 的实证原因。
上面是最小实现,生产代码在 taming-transformers(VQGAN 官方仓库,本文对照的是 taming/modules/vqvae/quantize.py 与 taming/modules/losses/vqperceptual.py,见 GitHub 源码,以下均以 2026-10 时的 master 为准)。逐个对照:
距离计算用展开式。 最小实现里我直接算 (z[:, None] - cb[None]) ** 2,会物化一个 N×K×d 的中间张量;taming 用 $\Vert z - e \Vert^2 = \Vert z \Vert^2 + \Vert e \Vert^2 - 2 z \cdot e$ 展开,只算 N×K 的矩阵乘:
d = torch.sum(z_flattened ** 2, dim=1, keepdim=True) \
+ torch.sum(self.embedding.weight**2, dim=1) \
- 2 * torch.einsum('bd,dn->bn', z_flattened, ...)
图像 tokenizer 的 N 是 batch × H/16 × W/16,K 上万,这个展开是省显存的关键。
直通就一行,和推导完全一致:z_q = z + (z_q - z).detach()。
taming 保留了一个带 bug 的版本。 VectorQuantizer2 有个 legacy 开关,legacy=True(默认)时损失是 loss = mean((z_q.detach()-z)**2) + beta * mean((z_q - z.detach())**2)——$\beta$ 被加在了码本项上,而不是 commitment 项,与论文公式相反。源码注释直说这是历史 bug,为了兼容旧 checkpoint 保留默认。读老代码、对老权重时要注意这个坑。
EMA 码本的权重是 requires_grad=False 的。EmbeddingEMA 里 weight、cluster_size、embed_avg 全是关闭梯度的参数,更新完全靠滑动平均——和我 3.3 节的推导一致。taming 还顺手算了 perplexity:exp(-sum(avg_probs * log(avg_probs))),这正是我用来诊断坍缩的量。
decay 不是「越接近 1 越稳」,它是一个时间常数。 衰减系数 $\gamma$ 决定码字「记得多久以前的样本」:$\gamma=0.99$ 时一个历史样本的贡献半衰期是 $\log 0.5 / \log 0.99 \approx 69$ 步,458 步后只剩 1%;$\gamma=0.999$ 半衰期拉到 693 步,$\gamma=0.95$ 只有 13.5 步。这个数要和 batch 里的 token 数对着调:VQGAN f=16 时一张 256² 图给 256 个 token,batch=8 一共 2048 个 token,摊到 K=16384 的码本上平均每步每个码字只被选中 0.125 次。也就是说大码本下 EMA 的计数极其稀疏,decay 太小会让码字在两次命中之间就被洗回零——这是「大码本更容易坍缩」的一条工程解释,也是 LFQ/FSQ 从结构上绕开它的动机。
损失在 losses/vqperceptual.py:L1 + LPIPS + hinge GAN。 三个值得抄的细节:一是 adopt_weight,判别器从第 disc_start 步才介入(先让重建学好,再上对抗);二是 calculate_adaptive_weight,取 $\Vert \nabla_{\text{last}} L_{\text{rec}} \Vert / \Vert \nabla_{\text{last}} L_{\text{GAN}} \Vert$ 并 clamp 到 10⁴;三是判别器是 NLayerPatchGAN,按 patch 判真假,这样高分辨率下判别器参数量不随分辨率爆炸。
官方数字(仓库 README 的重建 FID 表):
| 模型 | f | 码本 K | 重建 rFID |
|---|---|---|---|
| DALL-E dVAE(Gumbel) | 8 | 8192 | 33.88 |
| VQGAN ImageNet | 16 | 1024 | 10.54 |
| VQGAN ImageNet | 16 | 16384 | 7.41 |
| VQGAN OpenImages | 8 | 256 | 1.49 |
| VQGAN OpenImages | 8 | 16384 | 1.14 |
两行读法:同为 f=16,K 从 1024 加到 16384,rFID 从 10.54 降到 7.41——码本容量确实有用,但注意这是在 d=256 的码本维度上(vq_model.py 里 quant_conv 把通道投影到 embed_dim=256),按 $K^{-2/d}$ 的规律这个收益已经非常温和。另一个对照更惊人:VQGAN f=8 K=256 的 rFID(1.49)比 DALL-E dVAE f=8 K=8192(33.88)好了 22 倍——感知损失加对抗带来的提升,比码本大 32 倍带来的提升大得多。这就是 VQGAN 论文标题里 "taming" 的真正含义。
代价一:量化误差是硬地板,而且维度惩罚很重。 04 节的表已经给了 d=8 的数字:K=512 时量化误差仍把重建误差推高 80.5%(相对连续通路 0.011378)。想压低这块,有两条数学上已知的路:降有效维度(FSQ 干的事:把码本从「K 个 d 维向量」换成「每维只有少量取值」,维度语义变了但利用率为 100%),或者分层量化(残差 VQ、VAR 用的 RQ-VAE:一层量化不完的残差给下一层)。蛮力加 K 是最差的一条路。
代价二:码本坍缩几乎必然发生,解药都有副作用。 实测里随机初始化有 96.7% 死码;k-means 初始化要额外跑一次 k-means 且只在训练初期有用;死码重采样最有效,但它等价于「用随机重启换利用率」——被重启的码字携带的信息丢了,而且重启阈值又是一个新超参。工业界还有第三条路:直接改量化方式让坍缩在结构上不可能发生(LFQ 把每个维度独立二值化/多值化,FSQ 同理),这超出了本文范围,07 节给出处。
代价三:感知重建会「编」细节。 这是 3.5 节埋的伏笔,用一个能算清楚的玩具展开(perceptual_lab.py)。构造:一个 latent 对应两种等概率的纹理 $x = m \pm p$(m 是低频内容,p 是棋盘纹理,纹理占 83.5% 的梯度能量)。在候选输出 $m + c \cdot p$ 上比较两种损失:
| 候选输出 | MSE | PSNR | 梯度能量保留 | 到最近模态的距离 |
|---|---|---|---|---|
| $c=0$(MSE 最优,条件均值) | 0.250000 | 5.30 dB | 18.6% | 0.250 |
| $c=-0.79$(特征空间最优) | 0.404056 | 3.22 dB | 70.7% | 0.012 |
| $c=1$(选一个清晰模态) | 0.500000 | 2.29 dB | 97.5% | 0.000 |

这张图要看三处:上排四张 8×8 小图里,MSE 最优那格的棋盘纹理消失了(两种极性平均成平色),特征最优那格纹理回来了;左下柱状图说明这个纹理消失在 PSNR 上是加分的(5.30 dB 最高);右下的曲线说明只要特征里给高频一点权重($\alpha > 0.74$,实跑二分定位),最优解就会离开模糊均值往清晰模态走。三条事实合起来的结论是:PSNR 和「看起来真」在这类问题上方向相反——VQGAN 之后没人用 PSNR 报告 tokenizer 的重建质量,rFID 成了标配,原因就在这。
但这笔账的另一面是:对抗训练出来的「细节」不保证是真的。判别器只关心「像不像真图」,不关心「是不是这张图」,所以感知+GAN 的重建会补出数据集里常见的纹理——做压缩、做医学影像这类需要像素保真的任务时,这条路要慎走(GAN 训练不稳的代价也真实存在,05 节的 disc_start 和自适应权重都是为此付的工程税)。
边界:什么时候不该用。 下游不是自回归/离散先验时,离散化没有收益只有损失——Stable Diffusion 的第一-stage 就是连续 VAE(latent_diffusion 那篇讲过),因为扩散模型要的是连续潜空间上的 score,不是离散 token。同样,「码本利用率」这一整章在连续 VAE 里没有对应物。选型的判据一句话:下游要对离散序列做自回归或掩码预测,才需要 VQ tokenizer。
最后坦诚标注边界:本文所有数字来自 8×8 合成小块上的线性自编码器玩具(numpy 实跑,可复现),规模和真实 VQGAN(f=16、d=256、K=16384、百万级图像)差几个数量级;机制层面(直通、EMA、坍缩、感知损失的偏好)我认为可以直接外推,但具体数字(比如 $\beta=1.0$ 更好)不能外推,它依赖我的损失量纲。
误解一:「VQ-VAE 是 VAE 的一种」。 3.1 节推过:后验是确定性的 one-hot,先验是均匀分布,KL 精确等于 $\log K$,是常数,对训练没有任何贡献。没有变分、没有重参数化、没有 KL 正则——ELBO 在这里只是叙事起点,不是训练目标。
误解二:「码本是梯度下降学出来的」。 码本拿不到重建损失的梯度(查表挡住了,直通又把梯度全引向编码器),它的更新要么靠字典损失这一项、要么靠 EMA。taming 的 EMA 实现里码本权重干脆是 requires_grad=False 的。
误解三:「straight-through 是梯度的无偏估计」。 实测(04 节表):真实有限差分精确为 0,直通梯度非零。它不是对真实梯度的估计,是「假装量化不存在」的替代品;真正把它扶正的是 commitment 项,否则编码器会漂走。
误解四:「码本越大越好」。 $K^{-2/d}$ 的维度惩罚加上坍缩风险,让大码本的收益远低于直觉——K=1024 到 16384 在 d=256 上只把 rFID 从 10.54 拉到 7.41(05 节表)。利用率不到 100% 时,标称 bit 和有效 bit 的差距更离谱(本文玩具里 9 bit 只传出 3.86 bit)。
误解五:「$\beta=0.25$ 是默认值,照抄就行」。 $\beta$ 控制的是 commitment 梯度和重建直通梯度的量级比,本文玩具里两者天然差 429 倍,$\beta$ 扫描显示 1.0 反而最好。换了重建损失(MSE 换 L1+LPIPS)量纲就变,$\beta$ 必须重调。
三个实验都可以在附录代码里一键复现(numpy only,不需要 GPU):
python vq_core.py,看 [D] 段——把 $z_e$ 沿随机方向扰动 10⁻² 到 10⁻⁴,有限差分精确为 0,argmin 保持不变的样本比例是 1.0000;对比直通梯度平均范数 1.49 × 10⁻⁴。python collapse_lab.py,看 [B][C] 段——K=512 随机初始化最终只有 17 个活码(96.7% 死码)、重建 MSE 0.0491;换死码重采样后 500 个活码、MSE 0.0184。你也可以把 restart_dead=True 关掉再跑一遍,确认结果回到 17。python perceptual_lab.py,把文件顶部的 AMP(纹理幅度)从 0.5 改成 0.1 再跑——你会发现梯度能量保留率对纹理幅度极其敏感,而 MSE 最优解的纹理保留永远是 0(条件均值把 ±纹理精确抵消)。python codebook_size_law.py,读 [A] 段输出——d=16 时 K 从 8 加到 512 只把失真降到 1/12.9,d=2 时同样 64 倍的预算能降到 1/66.5。再把 make_patch_data 的 n_cluster 从 16 改成 4(类更少、潜变量更集中)重跑,斜率会明显变陡:失真下降的速度由数据分布本身决定,码本参数只是顺着它走。09 节用到的脚本全文如下(token_budget.py、collapse_lab.py、vq_core.py、perceptual_lab.py、codebook_size_law.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。
"""token_budget.py —— 离散化到底换来多少「序列长度」上的便宜。
VQGAN 那篇论文的动机只有一句话:Transformer 的注意力和序列长度是平方关系,
在像素上做自回归根本不可能,所以必须先把图像压成一串短的离散 token。
本脚本把这笔账算成具体的数:
[A] 不同分辨率 / 下采样率下的 token 数与注意力规模
[B] 码本大小 K 决定每个 token 多少 bit,折算成每张图多少 bpp
[C] 视频的 token 数(时空压缩一起算)
[D] 码本坍缩要付的比特代价:perplexity 才是真实容量
运行:/usr/local/bin/python3 token_budget.py
"""
from __future__ import annotations
import numpy as np
def n_tokens(h, w, f, t=1, f_t=1):
return (t // f_t) * (h // f) * (w // f)
def main():
print("[A] 图像:边长 H 与下采样率 f 决定 token 数(注意力按 n^2 涨)")
print(" H f token 数 注意力矩阵 n^2 相对像素级")
for H in (256, 512, 1024):
base = H * H * 3 # 像素级(RGB 逐通道)
for f in (4, 8, 16):
n = n_tokens(H, H, f)
print(" %-5d %-4d %-12d %-18.3e %.4g" % (
H, f, n, float(n) ** 2, float(n * n) / float(base * base)))
print(" (像素级 RGB 序列 %d,n^2 = %.3e)" % (base, float(base) ** 2))
print()
print("[B] 码本容量 K 决定每个 token 的 bit 数(256x256 图像)")
print(" K bit/token f=8: bit/图 bpp f=16: bit/图 bpp")
for K in (256, 512, 1024, 4096, 16384, 262144):
bit = np.log2(K)
row = [K, bit]
for f in (8, 16):
n = n_tokens(256, 256, f)
total = n * bit
row += [total, total / (256 * 256)]
print(" %-8d %-10.2f %-12.0f %-7.3f %-12.0f %.3f" % tuple(row))
print(" bpp = bit per pixel。作为参照,JPEG 在中等质量下大约 0.5~1 bpp。")
print()
print("[C] 视频:时间维也要压(5 秒 24fps = 121 帧,256x256)")
print(" f_t f token 数 上下文长度对比(相对 121x256x256x3 像素)")
pixel = 121 * 256 * 256 * 3
for f_t in (1, 4, 8):
for f in (8, 16):
n = n_tokens(256, 256, f, t=120, f_t=f_t)
print(" %-5d %-4d %-14d %.5g" % (
f_t, f, n, float(n) / pixel))
print()
print("[D] 码本坍缩的比特代价:真实容量看 perplexity,不是看 K")
print(" K perplexity 标称 bit 有效 bit 浪费")
for K, ppl in [(512, 14.50), (512, 217.26), (512, 426.71),
(16384, 14.50), (16384, 1000.0), (16384, 16384.0)]:
nominal = np.log2(K)
eff = np.log2(ppl)
print(" %-6d %-11.2f %-10.2f %-10.2f %.1f%%" % (
K, ppl, nominal, eff, 100 * (nominal - eff) / nominal))
print(" 前两行来自 collapse_lab.py 的实测:K=512 随机初始化只有 17 个码活着,")
print(" 9 bit 的码本只传出 3.86 bit;死码重启后回到 8.74 bit。")
print()
print("[E] 一句话总结")
f16 = n_tokens(256, 256, 16)
print(" 256x256 图像在 f=16 下是 %d 个 token,是像素级 RGB 序列的 1/%.0f;"
% (f16, (256 * 256 * 3) / f16))
print(" 注意力规模从 %.2e 降到 %.2e,省了 %.0f 倍——这才是必须先做 tokenizer 的原因。"
% (float(256 * 256 * 3) ** 2, float(f16) ** 2,
float(256 * 256 * 3) ** 2 / float(f16) ** 2))
if __name__ == "__main__":
main()
"""collapse_lab.py —— 码本坍缩(codebook collapse)是怎么发生的,能救回来多少。
码本坍缩指的是:K 个码字里只有一小撮被用到,剩下的从头到尾一次都没被选中
(死码)。花了一整个 K×d 的码本,只买到 log2(perplexity) 比特的表达力。
[A] commitment 权重 beta 怎么影响坍缩程度
[B] 大码本在训练过程中怎么一步步坍缩(活跃码数曲线)
[C] 两种常用解药有多大用:k-means 初始化码本 / 死码重采样重启
运行:/usr/local/bin/python3 collapse_lab.py
"""
from __future__ import annotations
import argparse
import numpy as np
from vq_core import kmeans, make_patch_data, perplexity, quantize, train_ae
BIG_K = 512
def stats(X, params, codebook):
"""给定训练好的权重与码本,算重建误差与使用分布统计量。"""
z_e = X @ params["We"]
z_q, idx, _ = quantize(z_e, codebook)
cnt = np.bincount(idx, minlength=len(codebook)).astype(np.float64)
rec = float((((z_q @ params["Wd"] + params["bd"]) - X) ** 2).mean())
frac = np.sort(cnt / cnt.sum())[::-1]
return {
"rec": rec,
"alive": int((cnt > 0).sum()),
"ppl": perplexity(cnt),
"top1": float(frac[0]),
"bottom_half": float(frac[len(frac) // 2:].sum()),
"counts": cnt,
}
def run_vq(X, K=BIG_K, beta=0.25, steps=1200, seed=0, mode="ema",
init_cb=None, restart=False, log_every=100):
"""跑一次 VQ-AE 训练,返回 (stats, hist)。"""
p, c, h, _ = train_ae(X, 8, mode, K=K, beta=beta, steps=steps, seed=seed,
restart_dead=restart, init_cb=init_cb,
log_every=log_every)
return stats(X, p, c), h
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--steps", type=int, default=1200)
ap.add_argument("--K", type=int, default=BIG_K)
ap.add_argument("--seed", type=int, default=0)
args = ap.parse_args()
X, _, _, _ = make_patch_data(n=4096, seed=args.seed)
print("数据 X%s,训练 %d 步" % (X.shape, args.steps))
print()
# ---------------- [A] beta 扫描 ----------------
print("[A] commitment 权重 beta 的影响(K=64,码本 EMA,%d 步)" % args.steps)
print(" beta 重建MSE 活跃码 perplexity")
for beta in (0.0, 0.05, 0.25, 1.0, 4.0):
s, _ = run_vq(X, K=64, beta=beta, steps=args.steps, seed=args.seed)
print(" %-9.2f %.6f %4d %.2f" % (
beta, s["rec"], s["alive"], s["ppl"]))
print(" -> beta=0 时编码器完全不被拉向码本,量化误差最大;beta 加大能压住")
print(" 误差,但也把编码器往码本上拽,两头都不免费(本玩具上 1.0 最好)。")
print()
# ---------------- [B] 大码本的坍缩过程 ----------------
print("[B] 大码本 K=%d 的坍缩过程(beta=0.25,EMA,随机初始化)" % args.K)
s0, h0 = run_vq(X, K=args.K, steps=args.steps, seed=args.seed, log_every=200)
print(" 步数 重建MSE 活跃码 perplexity")
for i in range(len(h0["step"])):
print(" %5d %.6f %4d %.2f" % (
h0["step"][i], h0["recon"][i], h0["alive"][i], h0["ppl"][i]))
print(" 最终:%d 个码字里只有 %d 个活着;perplexity %.2f 对应 %.2f bit,"
"而码本容量是 %.2f bit" % (
args.K, s0["alive"], s0["ppl"], np.log2(max(s0["ppl"], 1e-12)),
np.log2(args.K)))
print(" 最热的 1 个码占 %.3f 的使用量,最冷的一半码一共只占 %.4f"
% (s0["top1"], s0["bottom_half"]))
print()
# ---------------- [C] 解药 ----------------
print("[C] 两种解药(同为 K=%d,%d 步)" % (args.K, args.steps))
p_init, _, _, _ = train_ae(X, 8, "continuous", steps=800, seed=args.seed)
init_cb = kmeans(X @ p_init["We"], args.K, seed=args.seed, iters=20)
s1, h1 = run_vq(X, K=args.K, steps=args.steps, seed=args.seed,
init_cb=init_cb, log_every=200)
s2, h2 = run_vq(X, K=args.K, steps=args.steps, seed=args.seed,
restart=True, log_every=200)
print(" 配置 重建MSE 活跃码 perplexity 有效bit")
for tag, s in [("随机初始化", s0), ("k-means 初始化", s1),
("死码重采样重启", s2)]:
print(" %-16s %.6f %4d/%4d %6.2f %.2f" % (
tag, s["rec"], s["alive"], args.K, s["ppl"],
np.log2(max(s["ppl"], 1e-12))))
print()
print(" 注:本玩具的真实类心只有 16 个,活跃码数的上界本来就远小于 K,")
print(" 所以这里看的是「死码能不能被救活」,不是表达力真的翻了多少倍。")
print()
print("[D] 活跃码数随训练步数的变化(供配图)")
print(" 随机初始化 :", list(zip(h0["step"], h0["alive"])))
print(" k-means 初始化:", list(zip(h1["step"], h1["alive"])))
print(" 死码重采样 :", list(zip(h2["step"], h2["alive"])))
if __name__ == "__main__":
main()
"""vq_core.py —— 向量量化器(VQ)的最小实现,外加一个真跑得起来的训练实验。
无 torch 依赖,纯 numpy。本机解释器:/usr/local/bin/python3(3.10.5)。
运行:
/usr/local/bin/python3 vq_core.py
/usr/local/bin/python3 vq_core.py --steps 1200 --K 64 --beta 0.25
输出分五段:
[A] 连续自编码器基线(无量化)——量化误差的下界
[B] VQ-AE 训练结果:重建误差 / 量化误差 / 码本使用率 / perplexity
[C] 误差分解:总误差 = 子空间残差 + 量化误差(与 [A] 对拍)
[D] straight-through 的梯度核验:真实有限差分 vs 直通梯度
[E] 码本更新方式对照(不更新 / 梯度 / EMA / EMA+死码重启)
"""
from __future__ import annotations
import argparse
import os
import numpy as np
HERE = os.path.dirname(os.path.abspath(__file__))
# ------------------------------------------------------------------ 数据
def cos_basis(patch: int = 8, n_freq: int = 3):
"""8x8 小块的低频余弦基,共 n_freq^2 = 9 个,已按行归一化。"""
t = np.arange(patch) + 0.5
one_d = [np.ones(patch)]
for k in range(1, n_freq):
one_d.append(np.cos(np.pi * k * t / patch))
B = np.stack([np.outer(a, b).ravel() for a in one_d for b in one_d])
return B / np.linalg.norm(B, axis=1, keepdims=True)
def make_patch_data(n=4096, patch=8, n_cluster=16, seed=0,
center_scale=2.0, intra=0.35, noise=0.05):
"""合成一批 8x8 小块:16 个类心 + 类内低频连续变化 + 白噪声。
返回 X[N, 64]、类标、类心、基。类间方差远大于类内,所以码本有机会学到
「块类别」这种离散结构;白噪声部分不可压缩,构成误差地板。
"""
rng = np.random.default_rng(seed)
B = cos_basis(patch) # [9, 64]
dim_b = B.shape[0]
centers = center_scale * (rng.normal(size=(n_cluster, dim_b)) @ B)
labels = rng.integers(0, n_cluster, size=n)
X = centers[labels].copy()
X += intra * (rng.normal(size=(n, dim_b)) @ B) # 类内连续变化
X += noise * rng.normal(size=(n, patch * patch)) # 不可压缩噪声
return X, labels, centers, B
# ------------------------------------------------------------------ 量化
def quantize(z_e, codebook):
"""最近邻量化。z_e [N, d],codebook [K, d]。
返回 z_q[N, d]、索引 idx[N]、量化误差(每样本 d 维平方和)。
"""
d2 = ((z_e[:, None, :] - codebook[None, :, :]) ** 2).sum(-1) # [N, K]
idx = d2.argmin(axis=1)
return codebook[idx], idx, d2[np.arange(len(idx)), idx]
def kmeans(data, k, seed=0, iters=25):
"""Lloyd 迭代 + kmeans++ 初始化。空簇用「当前最差点」补齐。"""
rng = np.random.default_rng(seed)
n, d = data.shape
k = min(k, n)
centers = np.empty((k, d), dtype=data.dtype)
centers[0] = data[rng.integers(n)]
closest = ((data - centers[0]) ** 2).sum(1)
for j in range(1, k):
tot = closest.sum()
if tot <= 0:
centers[j] = data[rng.integers(n)]
else:
centers[j] = data[rng.choice(n, p=closest / tot)]
closest = np.minimum(closest, ((data - centers[j]) ** 2).sum(1))
for _ in range(iters):
_, assign, dist = quantize(data, centers)
dist = dist.copy()
for j in range(k):
mask = assign == j
if mask.any():
centers[j] = data[mask].mean(0)
else: # 空簇:拿当前最差的点填
j_worst = int(np.argmax(dist))
centers[j] = data[j_worst]
dist[j_worst] = -1.0
return centers
def perplexity(counts):
"""码本 perplexity = exp(使用分布的熵),上界是码本大小 K。"""
p = np.asarray(counts, dtype=np.float64)
p = p / p.sum()
nz = p[p > 0]
return float(np.exp(-(nz * np.log(nz)).sum()))
# ------------------------------------------------------------------ 训练
class Adam:
"""够用的 Adam,只处理一组 numpy 参数。"""
def __init__(self, params, lr=0.02, b1=0.9, b2=0.999, eps=1e-8):
self.p = params
self.m = {k: np.zeros_like(v) for k, v in params.items()}
self.v = {k: np.zeros_like(v) for k, v in params.items()}
self.lr, self.b1, self.b2, self.eps = lr, b1, b2, eps
self.t = 0
def step(self, grads):
self.t += 1
for k, g in grads.items():
self.m[k] = self.b1 * self.m[k] + (1 - self.b1) * g
self.v[k] = self.b2 * self.v[k] + (1 - self.b2) * (g * g)
mhat = self.m[k] / (1 - self.b1 ** self.t)
vhat = self.v[k] / (1 - self.b2 ** self.t)
self.p[k] -= self.lr * mhat / (np.sqrt(vhat) + self.eps)
def make_params(d_in, d_lat, seed=0):
rng = np.random.default_rng(seed)
return {
"We": rng.normal(scale=0.1, size=(d_in, d_lat)),
"Wd": rng.normal(scale=0.1, size=(d_lat, d_in)),
"bd": np.zeros(d_in),
}
def evaluate(X, params, codebook, eval_idx):
"""在固定评估集上算重建 MSE、量化误差、活跃码数、perplexity。"""
x = X[eval_idx]
z_e = x @ params["We"]
d_lat = params["We"].shape[1]
if codebook is None:
z_q, qerr, alive, ppl, counts = z_e, 0.0, 0, 0.0, None
else:
z_q, idx, qerr = quantize(z_e, codebook)
counts = np.bincount(idx, minlength=len(codebook)).astype(np.float64)
alive = int((counts > 0).sum())
ppl = perplexity(counts)
x_hat = z_q @ params["Wd"] + params["bd"]
return {
"recon": float(((x_hat - x) ** 2).mean()),
"quant": float(np.mean(qerr)) / d_lat if codebook is not None else 0.0,
"alive": alive,
"ppl": ppl,
"counts": counts,
}
def train_ae(X, d_lat, codebook_mode, K=64, beta=0.25, steps=800, batch=512,
lr=0.02, seed=0, decay=0.99, restart_dead=False, log_every=100,
cb_lr=10.0, init_cb=None):
"""训练「线性编码器 + VQ + 线性解码器」。
codebook_mode:
"continuous" —— 无量化,连续自编码器基线
"none" —— 码本完全不更新(只有 straight-through)
"grad" —— 码本用 ||sg[z_e] - e||^2 的梯度更新
"ema" —— 码本用指数滑动平均更新(cluster_size + embed_avg)
"""
rng = np.random.default_rng(seed)
eval_rng = np.random.default_rng(1000 + seed)
eval_idx = eval_rng.choice(len(X), size=min(2048, len(X)), replace=False)
n, d_in = X.shape
params = make_params(d_in, d_lat, seed=seed)
opt = Adam(params, lr=lr)
has_vq = codebook_mode != "continuous"
if has_vq:
# 码本初始化:小方差(呼应 taming 里 uniform(-1/n_e, 1/n_e) 的量级)
codebook = (rng.normal(scale=1.0 / K, size=(K, d_lat))
if init_cb is None else np.asarray(init_cb).copy())
init_codebook = codebook.copy()
cluster_size = np.ones(K, dtype=np.float64) # EMA 用
embed_avg = codebook.copy() # EMA 用
else:
codebook = init_codebook = None
cluster_size = embed_avg = None
hist = {"step": [], "recon": [], "quant": [], "alive": [], "ppl": []}
for step in range(1, steps + 1):
b = rng.choice(n, size=min(batch, n), replace=False)
x = X[b]
z_e = x @ params["We"] # [B, d]
if has_vq:
z_q, idx, _ = quantize(z_e, codebook)
e_k = z_q
x_hat = z_q @ params["Wd"] + params["bd"]
# 反传:解码器对 z_q 的梯度,原封不动地当作对 z_e 的梯度
dxh = 2.0 * (x_hat - x) / (len(x) * d_in) # [B, d_in]
dz_q = dxh @ params["Wd"].T # [B, d]
dz_e = dz_q + beta * 2.0 * (z_e - e_k) / (len(x) * d_lat)
opt.step({"We": x.T @ dz_e, "Wd": z_q.T @ dxh,
"bd": dxh.sum(0)})
if codebook_mode == "grad":
# d/de ||sg[z_e] - e||^2 = 2 (e - z_e)。注意码本是普通 SGD,
# 不走 Adam,所以它的有效步长要单独调(见 [E] 的 cb_lr 扫描)。
g = 2.0 * (e_k - z_e) / (len(x) * d_lat)
np.add.at(codebook, idx, -cb_lr * g)
elif codebook_mode == "ema":
onehot = np.zeros((len(x), K))
onehot[np.arange(len(x)), idx] = 1.0
cnt = onehot.sum(0)
cluster_size = decay * cluster_size + (1 - decay) * cnt
embed_avg = decay * embed_avg + (1 - decay) * (onehot.T @ z_e)
tot = cluster_size.sum()
smoothed = (cluster_size + 1e-5) / (tot + K * 1e-5) * tot
codebook = embed_avg / smoothed[:, None]
if restart_dead:
dead = np.where(cluster_size < 1.0)[0]
for j in dead: # 死码重采样为真实的编码器输出
codebook[j] = z_e[rng.integers(len(z_e))]
cluster_size[j] = 1.0
embed_avg[j] = codebook[j]
else:
x_hat = z_e @ params["Wd"] + params["bd"]
dxh = 2.0 * (x_hat - x) / (len(x) * d_in)
opt.step({"We": x.T @ (dxh @ params["Wd"].T),
"Wd": z_e.T @ dxh, "bd": dxh.sum(0)})
if step % log_every == 0 or step == steps:
m = evaluate(X, params, codebook, eval_idx)
hist["step"].append(step)
for k in ("recon", "quant", "alive", "ppl"):
hist[k].append(m[k])
shift = None
if has_vq:
shift = float(np.abs(codebook - init_codebook).mean())
return params, codebook, hist, shift
# ------------------------------------------------------------------ 主流程
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--steps", type=int, default=800)
ap.add_argument("--K", type=int, default=64)
ap.add_argument("--d", type=int, default=8)
ap.add_argument("--beta", type=float, default=0.25)
ap.add_argument("--seed", type=int, default=0)
args = ap.parse_args()
X, labels, centers, B = make_patch_data(n=4096, seed=args.seed)
n, d_in = X.shape
print("数据: X%s 样本范数均值 %.4f" % (X.shape, np.linalg.norm(X, axis=1).mean()))
print("码本 K=%d, 潜维度 d=%d, beta=%.2f" % (args.K, args.d, args.beta))
print()
# ---------------- [A] 连续自编码器基线 ----------------
pc, _, hc, _ = train_ae(X, args.d, "continuous", steps=args.steps,
seed=args.seed)
recon_c = hc["recon"][-1]
print("[A] 连续自编码器(无量化,最优线性重建)")
print(" 重建 MSE/像素 = %.6f" % recon_c)
print()
# ---------------- [B] VQ-AE ----------------
pv, cb, hv, _ = train_ae(X, args.d, "ema", K=args.K, beta=args.beta,
steps=args.steps, seed=args.seed)
print("[B] VQ-AE(码本 EMA 更新, decay=0.99)训练曲线")
print(" 步数 重建MSE 量化误差/d 活跃码 perplexity")
for i in range(len(hv["step"])):
print(" %5d %.6f %.6f %5d %.2f" % (
hv["step"][i], hv["recon"][i], hv["quant"][i],
hv["alive"][i], hv["ppl"][i]))
counts = np.bincount(quantize(X @ pv["We"], cb)[1],
minlength=args.K).astype(np.float64)
print(" 最终活跃码 %d / %d,perplexity %.2f(上界 %d)" % (
(counts > 0).sum(), args.K, perplexity(counts), args.K))
print()
# ---------------- [C] 误差分解 ----------------
z_e = X @ pv["We"]
z_q, _, qerr = quantize(z_e, cb)
rec_vq = (((z_q @ pv["Wd"] + pv["bd"]) - X) ** 2).mean()
rec_cont = (((z_e @ pv["Wd"] + pv["bd"]) - X) ** 2).mean()
print("[C] 误差分解(同一组 VQ-AE 权重,只换「走不走量化」)")
print(" 连续通路重建 MSE = %.6f (子空间残差,与 [A] 同量级)" % rec_cont)
print(" 量化后重建 MSE = %.6f" % rec_vq)
print(" 差值(量化引入) = %.6f (%.1f%%)" % (
rec_vq - rec_cont, 100 * (rec_vq - rec_cont) / rec_cont))
print(" [A] 连续基线 = %.6f" % recon_c)
print(" 量化误差 E||z_e-e||^2/d = %.6f" % (qerr.mean() / args.d))
print()
# ---------------- [D] straight-through 梯度核验 ----------------
# 用 [B] 训好的 VQ-AE,在真实工作点上做有限差分。
rng = np.random.default_rng(7)
x0 = X[:512]
z_e0 = x0 @ pv["We"]
z_q0, idx0, _ = quantize(z_e0, cb)
xh0 = z_q0 @ pv["Wd"] + pv["bd"]
L0 = ((xh0 - x0) ** 2).mean()
direction = rng.normal(size=args.d)
direction /= np.linalg.norm(direction)
print("[D] straight-through 梯度核验(512 样本,取 [B] 训好的权重)")
print(" 有限差分:把 z_e 沿随机单位方向微扰 eps,看 L 怎么变")
for eps in (1e-1, 1e-2, 1e-3, 1e-4):
z_p = z_e0 + eps * direction
z_qp, idxp, _ = quantize(z_p, cb)
Lp = ((z_qp @ pv["Wd"] + pv["bd"] - x0) ** 2).mean()
same = float((idxp == idx0).mean())
print(" eps=%.0e : dL/dz_e = %+.10f (argmin 保持不变的样本 %.4f)" % (
eps, (Lp - L0) / eps, same))
dxh = 2.0 * (xh0 - x0) / (len(x0) * d_in)
dz_q = dxh @ pv["Wd"].T
st_grad = float((dz_q * direction).mean())
st_norm = float(np.linalg.norm(dz_q, axis=1).mean())
commit = 2.0 * args.beta * (z_e0 - z_q0) / args.d
commit_norm = float(np.linalg.norm(commit, axis=1).mean())
print(" 直通梯度 dL/dz_q 投影到同一方向 = %+.10f" % st_grad)
print(" 直通梯度平均范数 ||dL/dz_q|| = %.6e" % st_norm)
print(" commitment 项平均范数 = %.6e" % commit_norm)
print(" -> 真实梯度恒为 0(argmin 是分段常数),直通梯度非零;")
print(" 它是「假装量化是恒等映射」的替代品,不是真实梯度的估计。")
print()
# ---------------- [E] 码本更新方式对照 ----------------
print("[E] 码本更新方式对照(同为 %d 步,K=%d,beta=%.2f)" % (
args.steps, args.K, args.beta))
print(" 模式 重建MSE 活跃码 perplexity 码本位移")
for mode, rst, cbl in [("none", False, 10.0), ("grad", False, 10.0),
("ema", False, 10.0), ("ema", True, 10.0)]:
p, c, h, shift = train_ae(X, args.d, mode, K=args.K, beta=args.beta,
steps=args.steps, seed=args.seed,
restart_dead=rst, cb_lr=cbl)
cnt = np.bincount(quantize(X @ p["We"], c)[1],
minlength=args.K).astype(np.float64)
tag = mode + ("+restart" if rst else "")
print(" %-18s %.6f %5d %.2f %.6f" % (
tag, h["recon"][-1], int((cnt > 0).sum()), perplexity(cnt), shift))
print()
print(" 附:grad 模式对码本步长极敏感(码本是普通 SGD,不走 Adam)")
for cbl in (0.02, 0.1, 1.0, 10.0):
p, c, h, _ = train_ae(X, args.d, "grad", K=args.K, beta=args.beta,
steps=args.steps, seed=args.seed, cb_lr=cbl)
cnt = np.bincount(quantize(X @ p["We"], c)[1],
minlength=args.K).astype(np.float64)
print(" cb_lr=%-6.2f 重建MSE %.6f 活跃码 %3d perplexity %.2f" % (
cbl, h["recon"][-1], int((cnt > 0).sum()), perplexity(cnt)))
print()
if __name__ == "__main__":
main()
"""perceptual_lab.py —— 为什么 VQGAN 不能只用 MSE:一个能算清的玩具。
VQ-VAE 用 MSE(或像素空间的似然)训练解码器。MSE 的最优解是条件均值,
而条件均值会把「同一 latent 对应多种合理细节」平均掉 —— 这就是重建发糊的
数学根源,不是玄学。VQGAN 的解法是换掉重建损失:改成特征空间距离(LPIPS)
外加一个对抗项。
本脚本用一个双模态玩具把这个机制算清楚:
[A] 构造:同一个 latent 对应两种等概率的纹理(+p 和 -p)
[B] MSE 最优 = 条件均值,纹理被平均掉;高频(梯度)能量只剩多少
[C] 换成非线性特征空间后,最优解跳到清晰模态;求阈值 alpha 并与解析式对拍
[D] 三种候选输出的三项指标对比:PSNR / 梯度能量保留 / 到最近模态的距离
注意:这里的特征映射 phi(x) = [x, alpha * |grad x|] 是我手工造的非线性特征,
用来演示「非线性」这一步为什么关键;真实的 LPIPS 用的是 VGG16 五层特征加
一层学习的 1x1 卷积,机制相同但权重是学出来的。
运行:/usr/local/bin/python3 perceptual_lab.py
"""
from __future__ import annotations
import numpy as np
PATCH = 8
AMP = 0.5 # 纹理幅度
SEED = 0
def gradient_magnitude(img):
"""前向差分后取绝对值:|grad x| 的展平向量(水平 56 个 + 垂直 56 个)。"""
gx = np.diff(img, axis=1)
gy = np.diff(img, axis=0)
return np.concatenate([np.abs(gx).ravel(), np.abs(gy).ravel()])
def phi(img, alpha):
"""非线性特征映射:图像本身 + alpha 乘梯度幅值。"""
return np.concatenate([img.ravel(), alpha * gradient_magnitude(img)])
def make_toy():
"""低频内容 m + 等概率的 ±棋盘纹理 p。"""
rng = np.random.default_rng(SEED)
u = np.arange(PATCH) + 0.5
m = np.outer(np.cos(np.pi * u / PATCH), np.cos(np.pi * u / PATCH)) * 1.5
m += 0.3 * rng.normal(size=(PATCH, PATCH)) # 让内容不那么对称
ii, jj = np.meshgrid(np.arange(PATCH), np.arange(PATCH), indexing="ij")
p = AMP * ((-1.0) ** (ii + jj))
return m, p
def evaluate_candidate(m, p, c, alpha):
"""候选输出 x_hat = m + c * p,返回三项指标。"""
x_hat = m + c * p
x_plus, x_minus = m + p, m - p
# MSE(对两种真值取期望)
mse = 0.5 * (((x_hat - x_plus) ** 2).mean() + ((x_hat - x_minus) ** 2).mean())
# 特征空间距离(对两种真值取期望)
f_hat = phi(x_hat, alpha)
feat = 0.5 * (((f_hat - phi(x_plus, alpha)) ** 2).mean()
+ ((f_hat - phi(x_minus, alpha)) ** 2).mean())
# 梯度能量保留(相对真值的期望梯度能量)
g_hat = np.concatenate([np.diff(x_hat, axis=1).ravel(),
np.diff(x_hat, axis=0).ravel()])
g_true = 0.5 * (np.concatenate([np.diff(x_plus, axis=1).ravel(),
np.diff(x_plus, axis=0).ravel()]) ** 2).sum()
g_true += 0.5 * (np.concatenate([np.diff(x_minus, axis=1).ravel(),
np.diff(x_minus, axis=0).ravel()]) ** 2).sum()
grad_keep = float((g_hat ** 2).sum() / g_true)
# 到最近模态的距离(越小说明落在数据流形上)
to_mode = float(min(((x_hat - x_plus) ** 2).mean(),
((x_hat - x_minus) ** 2).mean()))
return {"c": c, "mse": float(mse), "feat": float(feat),
"grad_keep": grad_keep, "to_mode": to_mode}
def best_c(m, p, alpha, grid=None):
if grid is None:
grid = np.linspace(-1.5, 1.5, 601)
losses = np.array([evaluate_candidate(m, p, c, alpha)["feat"] for c in grid])
return float(grid[int(np.argmin(losses))]), grid, losses
def main():
m, p = make_toy()
signal_power = float(((0.5 * ((m + p) ** 2).mean() + 0.5 * ((m - p) ** 2).mean())))
print("[A] 玩具构造:8x8 patch,x = m + s * p,s = +1 / -1 各 0.5")
print(" 低频内容 m: 范数 %.4f,梯度能量 %.4f" % (
np.linalg.norm(m), (np.concatenate([np.diff(m, axis=1).ravel(),
np.diff(m, axis=0).ravel()]) ** 2).sum()))
print(" 棋盘纹理 p: 范数 %.4f,梯度能量 %.4f(纹理占了 %.1f%% 的梯度能量)" % (
np.linalg.norm(p),
(np.concatenate([np.diff(p, axis=1).ravel(),
np.diff(p, axis=0).ravel()]) ** 2).sum(),
100 * (np.concatenate([np.diff(p, axis=1).ravel(),
np.diff(p, axis=0).ravel()]) ** 2).sum()
/ (np.concatenate([np.diff(m + p, axis=1).ravel(),
np.diff(m + p, axis=0).ravel()]) ** 2).sum()))
print()
print("[B] MSE 最优 = 条件均值(c=0,纹理被平均掉)")
c_star, grid, losses = best_c(m, p, 0.0)
r0 = evaluate_candidate(m, p, 0.0, 0.0)
psnr0 = 10 * np.log10(signal_power / r0["mse"])
print(" MSE 最优的 c* = %.3f(理论值 0)" % c_star)
print(" 重建 MSE = %.6f,PSNR = %.2f dB" % (r0["mse"], psnr0))
print(" 梯度能量保留 = %.4f(条件均值必然丢高频:E||grad x||^2 = "
"||grad E[x]||^2 + E||grad(x - E[x])||^2)" % r0["grad_keep"])
print()
print("[C] 换成特征空间后,最优解跳到清晰模态")
print(" alpha 最优 c* 该点的 MSE PSNR(dB) 梯度能量保留")
for alpha in (0.0, 0.2, 0.3, 0.378, 0.4, 0.5, 1.0, 2.0):
c_a, _, _ = best_c(m, p, alpha)
r = evaluate_candidate(m, p, c_a, alpha)
psnr = 10 * np.log10(signal_power / r["mse"])
print(" %-8.3f %-10.3f %-13.6f %-10.2f %.4f" % (
alpha, c_a, r["mse"], psnr, r["grad_keep"]))
# 解析阈值:alpha^2 > ||p||^2 / ||grad p||^2
gp = np.concatenate([np.diff(p, axis=1).ravel(), np.diff(p, axis=0).ravel()])
thresh = float(np.sqrt((p ** 2).sum() / (gp ** 2).sum()))
print(" 粗略解析估计 alpha* = sqrt(||p||^2 / ||grad p||^2) = %.4f"
"(假设纹理梯度远大于内容梯度)" % thresh)
# 精确的「两个候选」比较:L(c) = (c^2+1)||p||^2 + alpha^2 * G(c),
# 于是「清晰模态 L(1)」优于「模糊均值 L(0)」的条件是 alpha^2 > ||p||^2/(G(0)-G(1))
g0 = 0.5 * ((gradient_magnitude(m) - gradient_magnitude(m + p)) ** 2).sum()
g0 += 0.5 * ((gradient_magnitude(m) - gradient_magnitude(m - p)) ** 2).sum()
g1 = 0.5 * ((gradient_magnitude(m + p) - gradient_magnitude(m - p)) ** 2).sum()
exact = float(np.sqrt((p ** 2).sum() / (g0 - g1)))
print(" 精确阈值:G(0)=%.4f, G(1)=%.4f,L(1)<L(0) 要求 alpha > %.4f"
% (g0, g1, exact))
print(" -> alpha 超过 %.3f 之后,「输出一个清晰模态」在特征损失上严格优于"
"「输出模糊均值」;" % exact)
print(" alpha 继续加大,最优 c 沿坐标轴继续往 ±1 移(不是跳变,"
"因为 |grad(m+c*p)| 关于 c 连续)。")
print()
print("[D] 四种候选输出的指标对比(PSNR 用信号能量 %.4f 作基准)" % signal_power)
print(" 候选 MSE PSNR(dB) 梯度保留 到最近模态距离")
c_feat, _, _ = best_c(m, p, 1.0)
for tag, c in [("MSE 最优 (c=0)", 0.0), ("折中 (c=0.5)", 0.5),
("特征最优 (c=%.3f)" % c_feat, c_feat), ("清晰模态 (c=1)", 1.0)]:
r = evaluate_candidate(m, p, c, 1.0)
psnr = 10 * np.log10(signal_power / r["mse"])
print(" %-20s %.6f %-10.2f %-9.4f %.6f" % (
tag, r["mse"], psnr, r["grad_keep"], r["to_mode"]))
print()
r_mean = evaluate_candidate(m, p, 0.0, 1.0)
r_sharp = evaluate_candidate(m, p, 1.0, 1.0)
r_feat = evaluate_candidate(m, p, c_feat, 1.0)
print(" PSNR 的绝对值很小,是因为这个玩具里纹理完全无法从 latent 预测,")
print(" 要看的是相对差:")
print(" · 特征最优解 (c=%.3f) 的 MSE 是模糊均值的 %.2f 倍(PSNR 低 %.2f dB);"
% (c_feat, r_feat["mse"] / r_mean["mse"],
-10 * np.log10(r_mean["mse"] / r_feat["mse"])))
print(" · 但它保留了 %.0f%% 的梯度能量,模糊均值只保留 %.0f%%;"
% (100 * r_feat["grad_keep"], 100 * r_mean["grad_keep"]))
print(" · 完全选一个模态 (c=1) 时 MSE 是均值的 %.1f 倍(PSNR 低 %.2f dB),"
"梯度能量保留 %.0f%%,且到最近模态距离为 0 —— 它落在数据流形上。"
% (r_sharp["mse"] / r_mean["mse"],
-10 * np.log10(r_mean["mse"] / r_sharp["mse"]),
100 * r_sharp["grad_keep"]))
print(" LPIPS/FID 站在后者一边,PSNR 站在前者一边 —— 这就是 VQGAN 之后")
print(" 没人再用 PSNR 报告 tokenizer 重建质量的原因。")
if __name__ == "__main__":
main()
"""codebook_size_law.py —— 码本容量 K 的收益到底有多大?
核心问题:把码本从 512 加到 16384,重建能好多少?直觉是「容量越大越好」,
但高维量化的经典结论(Zador 定理的渐近形式)说:最优量化失真随码本大小
只能按 K^(-2/d) 衰减,d 是码本向量维度。d=256 时指数只有 -1/128,也就是
K 翻一倍、失真只降 0.5%。
本脚本的实测方式:先训一个连续自编码器把潜分布固定住,再对潜变量做 k-means
(k-means 就是「给定 K 个码字的最优最近邻量化器」的近似),在**留出集**上
测失真(避免用训练集测失真造成的过拟合假象)。
[A] 不同 d 下,量化失真 vs K 的双对数斜率,与 -2/d 对拍
[B] 量化误差什么时候降到「子空间残差」之下(继续加 K 的收益拐点)
运行:/usr/local/bin/python3 codebook_size_law.py
"""
from __future__ import annotations
import argparse
import numpy as np
from vq_core import kmeans, make_patch_data, quantize, train_ae
def distortion_curve(z_fit, z_test, k_list, seed=0, iters=20):
"""在 z_fit 上拟合 k-means,在 z_test 上测失真(留出集估计)。"""
out = []
for k in k_list:
c = kmeans(z_fit, k, seed=seed, iters=iters)
_, _, d2 = quantize(z_test, c)
out.append(float(d2.mean()) / z_fit.shape[1])
return np.array(out)
def fit_slope(k_list, dist):
"""双对数线性拟合:log D = a * log K + b,返回斜率 a。"""
lk = np.log(np.asarray(k_list, dtype=np.float64))
ld = np.log(np.asarray(dist, dtype=np.float64))
a, b = np.polyfit(lk, ld, 1)
return float(a), float(b)
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--steps", type=int, default=1500, help="连续自编码器训练步数")
ap.add_argument("--seed", type=int, default=0)
args = ap.parse_args()
X, _, _, _ = make_patch_data(n=8192, seed=args.seed)
d_in = X.shape[1]
# 前一半拟合码本,后一半当留出集测失真
X_fit, X_test = X[:4096], X[4096:]
k_list = [2, 4, 8, 16, 32, 64, 128, 256, 512]
print("数据 X%s(前 4096 拟合码本,后 4096 留出评估)" % (X.shape,))
print()
print("[A] 量化失真 D(K) = E||z - q(z)||^2 / d 的双对数斜率")
print(" d 理论 -2/d 实测斜率(K>=16) D(8) D(512) 衰减倍数")
results = {}
for d in (2, 4, 8, 16):
params, _, hist, _ = train_ae(X_fit, d, "continuous", steps=args.steps,
seed=args.seed)
z_fit = X_fit @ params["We"]
z_test = X_test @ params["We"]
dist = distortion_curve(z_fit, z_test, k_list, seed=args.seed)
slope, _ = fit_slope(k_list[3:], dist[3:]) # 只拟合 K>=16 的渐近段
results[d] = {"dist": dist, "slope": slope, "params": params,
"z_fit": z_fit, "z_test": z_test}
print(" %-4d %-11.4f %-17.4f %.4e %.4e %.1fx" % (
d, -2.0 / d, slope, dist[2], dist[-1], dist[2] / dist[-1]))
print()
print(" K 从 8 加到 512(64 倍):d=2 失真降 %.1f 倍,d=16 只降 %.1f 倍。"
% (results[2]["dist"][2] / results[2]["dist"][-1],
results[16]["dist"][2] / results[16]["dist"][-1]))
print(" d=2/4 的斜率与 -2/d 吻合;d>=8 时 K=512 已经逼近每个码字 8 个样本,")
print(" 斜率被有限样本抬高(同样的码本在训练集上测会更陡),真值应更接近理论。")
print()
# ---------------- [B] 拐点:量化误差 vs 子空间残差 ----------------
params = results[8]["params"]
z_fit, z_test = results[8]["z_fit"], results[8]["z_test"]
rec_cont = (((z_test @ params["Wd"] + params["bd"]) - X_test) ** 2).mean()
print("[B] 码本大到什么程度,量化误差才降到子空间残差之下(d=8,留出集)")
print(" 连续通路重建 MSE(子空间残差) = %.6f" % rec_cont)
print(" K 量化失真/d 量化后重建MSE 相对残差涨幅")
for k in k_list:
c = kmeans(z_fit, k, seed=args.seed, iters=20)
z_q, _, d2 = quantize(z_test, c)
rec = (((z_q @ params["Wd"] + params["bd"]) - X_test) ** 2).mean()
print(" %-7d %.4e %.6f %+.1f%%" % (
k, d2.mean() / 8, rec, 100 * (rec - rec_cont) / rec_cont))
print()
if __name__ == "__main__":
main()
"""make_figures.py —— 画正文用到的五张示意图。
所有数字都现场重算(不读缓存),来源是同目录下的实验脚本:
token_budget.py -> 图 1(序列长度与注意力规模)
codebook_size_law.py-> 图 2(码本容量 K 的收益曲线)
collapse_lab.py -> 图 3、图 4(坍缩过程与使用分布)
perceptual_lab.py -> 图 5(MSE 最优 vs 特征空间最优)
运行:/usr/local/bin/python3 make_figures.py
输出:../figures/*.png
"""
from __future__ import annotations
import os
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import codebook_size_law as CSL
import collapse_lab as CL
import perceptual_lab as PL
import token_budget as TB
from vq_core import kmeans, make_patch_data, quantize, train_ae
HERE = os.path.dirname(os.path.abspath(__file__))
FIGDIR = os.path.join(HERE, os.pardir, "figures")
os.makedirs(FIGDIR, exist_ok=True)
plt.rcParams["font.sans-serif"] = ["PingFang SC", "Heiti TC", "Arial Unicode MS"]
plt.rcParams["axes.unicode_minus"] = False
C_A = "#2E6DB4" # 蓝:主曲线 / 基线
C_B = "#C0504D" # 红:第二种配置
C_C = "#4FA96B" # 绿:第三种配置 / 好的一方
C_D = "#E08A2E" # 橙:强调
C_GREY = "#8C8C8C"
def _save(fig, name):
path = os.path.join(FIGDIR, name)
fig.savefig(path, dpi=140, bbox_inches="tight", facecolor="white")
plt.close(fig)
print(" -> %s" % path)
# ---------------------------------------------------------------- 图 1
def fig_token_budget():
sizes = np.array([128, 256, 512, 1024, 2048], dtype=float)
fig, axes = plt.subplots(1, 2, figsize=(11.5, 4.3))
ax = axes[0]
for f, col, mk in [(4, C_D, "o"), (8, C_A, "s"), (16, C_C, "^")]:
n = (sizes / f) ** 2
ax.plot(sizes, n, color=col, marker=mk, lw=1.8,
label=r"tokenizer $f$=%d" % f)
ax.plot(sizes, sizes ** 2 * 3, color=C_GREY, marker="d", lw=1.8, ls="--",
label="像素级 RGB")
ax.set_xscale("log", base=2)
ax.set_yscale("log")
ax.set_xlabel("边长 H (像素)")
ax.set_ylabel("序列长度(token 数)")
ax.set_title("(a) 压缩率决定序列长度")
ax.grid(alpha=0.3, which="both")
ax.legend(fontsize=9)
ax.annotate("H=256, f=16\n只有 256 个 token",
xy=(256, 256), xytext=(300, 60), fontsize=9, color=C_C,
arrowprops=dict(arrowstyle="->", color=C_C, lw=1.2))
ax = axes[1]
labels = ["像素级\nRGB", "f=4", "f=8", "f=16"]
vals = [(256 * 256 * 3) ** 2, ((256 / 4) ** 2) ** 2,
((256 / 8) ** 2) ** 2, ((256 / 16) ** 2) ** 2]
colors = [C_GREY, C_D, C_A, C_C]
bars = ax.bar(labels, vals, color=colors)
ax.set_yscale("log")
ax.set_ylabel(r"注意力矩阵规模 $n^2$")
ax.set_title("(b) 256x256 图像的自注意力开销")
ax.grid(alpha=0.3, axis="y")
for b, v in zip(bars, vals):
ax.text(b.get_x() + b.get_width() / 2, v * 1.6, "%.1e" % v,
ha="center", fontsize=9)
ax.annotate("省 5.9e5 倍",
xy=(3, vals[3]), xytext=(2.55, 1e8), fontsize=9, color=C_C,
arrowprops=dict(arrowstyle="->", color=C_C, lw=1.2))
fig.suptitle("图 1:为什么必须先做 tokenizer(数字来自 token_budget.py)",
fontsize=11)
fig.tight_layout()
_save(fig, "token_budget.png")
# ---------------------------------------------------------------- 图 2
def fig_codebook_law():
X, _, _, _ = make_patch_data(n=8192, seed=0)
X_fit, X_test = X[:4096], X[4096:]
k_list = [2, 4, 8, 16, 32, 64, 128, 256, 512]
fig, ax = plt.subplots(figsize=(7.6, 5.2))
color_map = {2: C_C, 4: C_A, 8: C_D, 16: C_B}
slopes = {}
for d in (2, 4, 8, 16):
params, _, _, _ = train_ae(X_fit, d, "continuous", steps=1500, seed=0)
z_fit = X_fit @ params["We"]
z_test = X_test @ params["We"]
dist = CSL.distortion_curve(z_fit, z_test, k_list, seed=0)
slope, _ = CSL.fit_slope(k_list[3:], dist[3:])
slopes[d] = slope
ax.plot(k_list, dist, color=color_map[d], marker="o", lw=1.8,
label=r"$d$=%d 实测斜率 %.2f" % (d, slope))
# 理论斜率 -2/d,锚定在 K=16 处
anchor = dist[3]
theo = anchor * (np.array(k_list[3:], dtype=float) / 16.0) ** (-2.0 / d)
ax.plot(k_list[3:], theo, color=color_map[d], lw=1.0, ls=":",
alpha=0.75)
ax.set_xscale("log", base=2)
ax.set_yscale("log")
ax.set_xlabel(r"码本大小 $K$")
ax.set_ylabel(r"量化失真 $E\Vert z-q(z)\Vert^2/d$(留出集)")
ax.set_title("图 2:码本容量 K 的收益被码本维度 d 卡死\n"
"实线=实测,点线=理论斜率 -2/d(数字来自 codebook_size_law.py)",
fontsize=11)
ax.grid(alpha=0.3, which="both")
ax.legend(fontsize=9)
_save(fig, "codebook_size_law.png")
return slopes
# ---------------------------------------------------------------- 图 3、4
def fig_collapse():
X, _, _, _ = make_patch_data(n=4096, seed=0)
K = CL.BIG_K
s_rand, h_rand = CL.run_vq(X, K=K, steps=1200, seed=0, log_every=100)
p_init, _, _, _ = train_ae(X, 8, "continuous", steps=800, seed=0)
init_cb = kmeans(X @ p_init["We"], K, seed=0, iters=20)
s_km, h_km = CL.run_vq(X, K=K, steps=1200, seed=0, init_cb=init_cb,
log_every=100)
s_rs, h_rs = CL.run_vq(X, K=K, steps=1200, seed=0, restart=True,
log_every=100)
fig, ax = plt.subplots(figsize=(7.6, 5.0))
for h, col, mk, tag in [(h_rand, C_B, "o", "随机初始化"),
(h_km, C_A, "s", "k-means 初始化"),
(h_rs, C_C, "^", "死码重采样重启")]:
ax.plot(h["step"], h["alive"], color=col, marker=mk, lw=1.8, label=tag)
ax.axhline(K, color=C_GREY, ls="--", lw=1.2)
ax.text(max(h_rand["step"]) * 0.55, K * 1.08,
"码本容量 K=%d(全活)" % K, color=C_GREY, fontsize=9)
ax.set_xlabel("训练步数")
ax.set_ylabel("活跃码数(至少被选中过一次)")
ax.set_title("图 3:码本坍缩过程\n"
"K=%d 的码本,随机初始化下最终只有 %d 个码活着(%.1f%%)"
% (K, s_rand["alive"], 100.0 * s_rand["alive"] / K),
fontsize=11)
ax.grid(alpha=0.3)
ax.legend(fontsize=9)
ax.set_ylim(0, K * 1.25)
_save(fig, "collapse_curve.png")
fig, ax = plt.subplots(figsize=(7.6, 4.8))
for s, col, tag in [(s_rand, C_B, "随机初始化"), (s_rs, C_C, "死码重采样重启")]:
frac = np.sort(s["counts"] / s["counts"].sum())[::-1]
ax.plot(np.arange(1, len(frac) + 1), frac, color=col, lw=1.8,
label="%s:%d 个活码,perplexity %.1f" % (tag, s["alive"], s["ppl"]))
ax.set_yscale("log")
ax.set_xlabel("码字按使用频次从高到低排序")
ax.set_ylabel("使用占比")
ax.set_title("图 4:使用分布的长尾与死码\n"
"随机初始化有 %.1f%% 的码一次都没被用到" %
(100.0 * (1 - s_rand["alive"] / K)), fontsize=11)
ax.grid(alpha=0.3, which="both")
ax.legend(fontsize=9)
_save(fig, "usage_hist.png")
return s_rand, s_km, s_rs
# ---------------------------------------------------------------- 图 5
def fig_perceptual():
m, p = PL.make_toy()
signal_power = float(0.5 * ((m + p) ** 2).mean() + 0.5 * ((m - p) ** 2).mean())
c_feat, _, _ = PL.best_c(m, p, 1.0)
cands = [("真值样本 s=+1", m + p), ("真值样本 s=-1", m - p),
("MSE 最优 c=0", m), ("特征最优 c=%.2f" % c_feat, m + c_feat * p)]
fig, axes = plt.subplots(2, 4, figsize=(12.0, 5.6),
gridspec_kw={"height_ratios": [1.0, 0.92]})
vmin = min(a.min() for _, a in cands)
vmax = max(a.max() for _, a in cands)
for ax, (tag, img) in zip(axes[0], cands):
im = ax.imshow(img, cmap="gray", vmin=vmin, vmax=vmax, interpolation="nearest")
ax.set_title(tag, fontsize=10)
ax.set_xticks([])
ax.set_yticks([])
axes[0][0].set_ylabel("8x8 patch", fontsize=9)
rows = [("MSE 最优 c=0", 0.0), ("折中 c=0.5", 0.5),
("特征最优 c=%.2f" % c_feat, c_feat), ("清晰模态 c=1", 1.0)]
psnrs, keeps = [], []
for tag, c in rows:
r = PL.evaluate_candidate(m, p, c, 1.0)
psnrs.append(10 * np.log10(signal_power / r["mse"]))
keeps.append(100 * r["grad_keep"])
xpos = np.arange(len(rows))
ax = axes[1][0]
bs = ax.bar(xpos, psnrs, color=[C_A, C_GREY, C_D, C_C])
ax.set_xticks(xpos)
ax.set_xticklabels(["MSE\n最优", "折中", "特征\n最优", "清晰\n模态"], fontsize=8)
ax.set_ylabel("PSNR (dB)")
ax.set_title("(e) 像素误差:模糊均值最好", fontsize=10)
ax.grid(alpha=0.3, axis="y")
ax.set_ylim(0, max(psnrs) * 1.3)
for b, v in zip(bs, psnrs):
ax.text(b.get_x() + b.get_width() / 2, v + 0.08, "%.2f" % v,
ha="center", fontsize=8)
ax = axes[1][1]
bs = ax.bar(xpos, keeps, color=[C_A, C_GREY, C_D, C_C])
ax.set_xticks(xpos)
ax.set_xticklabels(["MSE\n最优", "折中", "特征\n最优", "清晰\n模态"], fontsize=8)
ax.set_ylabel("梯度能量保留 (%)")
ax.set_title("(f) 细节保留:清晰模态最好", fontsize=10)
ax.grid(alpha=0.3, axis="y")
ax.set_ylim(0, 118)
for b, v in zip(bs, keeps):
ax.text(b.get_x() + b.get_width() / 2, v + 1.5, "%.0f%%" % v,
ha="center", fontsize=8)
# 右侧两格合并成一条 alpha 扫描曲线(先把占位的两个空轴关掉)
axes[1][2].axis("off")
axes[1][3].axis("off")
ax = plt.subplot2grid((2, 4), (1, 2), colspan=2)
alphas = np.linspace(0.0, 2.0, 41)
cs = [PL.best_c(m, p, a)[0] for a in alphas]
ax.plot(alphas, np.abs(cs), color=C_A, lw=1.8)
ax.axhline(1.0, color=C_GREY, ls="--", lw=1.0)
ax.axhline(0.0, color=C_GREY, ls="--", lw=1.0)
ax.set_xlabel(r"特征里高频项的权重 $\alpha$")
ax.set_ylabel("最优输出的纹理系数 |c|")
ax.set_title(r"(g) $\alpha$ 越大,最优解越靠近清晰模态", fontsize=10)
ax.grid(alpha=0.3)
ax.set_ylim(-0.05, 1.15)
ax.set_xticks([0.0, 0.5, 1.0, 1.5, 2.0])
fig.suptitle("图 5:MSE 最优必然糊,特征空间最优不糊"
"(数字来自 perceptual_lab.py)", fontsize=11)
fig.tight_layout()
_save(fig, "perceptual_tradeoff.png")
def main():
print("画图:所有数字现场重算,不读缓存")
print("[1/5] token_budget.png")
fig_token_budget()
print("[2/5] codebook_size_law.png")
slopes = fig_codebook_law()
print(" 实测斜率:", {k: round(v, 3) for k, v in slopes.items()})
print("[3/5] collapse_curve.png + usage_hist.png")
s_rand, s_km, s_rs = fig_collapse()
print(" 随机初始化: alive=%d ppl=%.2f rec=%.6f" % (
s_rand["alive"], s_rand["ppl"], s_rand["rec"]))
print(" k-means 初始化: alive=%d ppl=%.2f rec=%.6f" % (
s_km["alive"], s_km["ppl"], s_km["rec"]))
print(" 死码重采样: alive=%d ppl=%.2f rec=%.6f" % (
s_rs["alive"], s_rs["ppl"], s_rs["rec"]))
print("[5/5] perceptual_tradeoff.png")
fig_perceptual()
print("完成,图在 %s" % os.path.abspath(FIGDIR))
if __name__ == "__main__":
main()
评论 (0)