Skip to content
huc
Go back

DDCM:把 DDPM 变成离散编码器

目录 1 / 5

《生成扩散模型漫谈(二十九):用DDPM来离散编码》 —— 苏剑林

连续噪声序列就是 DDPM 的隐变量

写出 DDPM 的随机反向一步:

xt1=μt(xt)+σtεt,εtN(0,I).(1)x_{t-1}=\mu_t(x_t)+\sigma_t\varepsilon_t, \qquad \varepsilon_t\sim\mathcal N(0,I). \tag{1}

连同初始 xTx_T,一条生成轨迹共使用 T+1T+1 个高维高斯变量。DDCM 为每一步预采样独立码本

Ct={ct,1,,ct,K},ct,ki.i.d.N(0,I),\mathcal C_t=\{c_{t,1},\ldots,c_{t,K}\}, \qquad c_{t,k}\overset{\text{i.i.d.}}\sim\mathcal N(0,I),

并把公式 (1) 中的噪声限制为 εtCt\varepsilon_t\in\mathcal C_t。生成过程便成为离散序列到图像的映射:

(iT,,i1){1,,K}Tx0.(2)(i_T,\ldots,i_1)\in\{1,\ldots,K\}^{T} \longmapsto x_0. \tag{2}

每个时间步应使用独立码本,因为不同 tt 的噪声扮演不同局部修正角色;共享码本会显著降低可选方向的组合多样性,通常需要更大的 KK 才能补偿。

有限码本的经验前提是:从高斯分布抽取的一组固定方向已足以近似每一步所需随机性。它不表示离散分布严格等于高斯;KK 增大时,经验分布才逐渐逼近连续分布。

编码不能从 x0x_0 直接倒推

给定目标图像 x0x_0^*,希望找到索引序列使解码结果接近它。直接由

x0=μ1(x1)+σ1ε1x_0^*=\mu_1(x_1)+\sigma_1\varepsilon_1

倒推会同时遇到未知 x1x_1 和未知码字,形成组合搜索。DDCM 改为从固定 xTx_T 正向执行反向采样,并在每一步选择最能把当前预测推向 x0x_0^* 的码字。

μˉt(xt)\bar\mu_t(x_t) 表示模型由 xtx_t 预测的干净数据。DDPM 的反向均值可写成

μt(xt)=atxt+btμˉt(xt),\mu_t(x_t) =a_t x_t+b_t\bar\mu_t(x_t),

其中在常用记号下

at=αtβˉt12βˉt2,bt=αˉt1βt2βˉt2.a_t=\frac{\alpha_t\bar\beta_{t-1}^2}{\bar\beta_t^2}, \qquad b_t=\frac{\bar\alpha_{t-1}\beta_t^2}{\bar\beta_t^2}.

目标残差为 rt=x0μˉt(xt)r_t=x_0^*-\bar\mu_t(x_t)。原方法选择

εt=argmaxcCtc,rt,(3)\varepsilon_t =\mathop{\operatorname{argmax}}_{c\in\mathcal C_t} \langle c,r_t\rangle, \tag{3}

再用公式 (1) 更新。内积最大意味着码字方向最能补偿当前干净图预测的残差。每一步记录 cc 的索引,最终得到天然有序的一维离散编码。

从条件后验理解选择规则

已知真实 x0x_0^* 时,DDPM 的解析后验为

q(xt1xt,x0)=N(atxt+btx0,σt2I).q(x_{t-1}\mid x_t,x_0^*) =\mathcal N\left( a_t x_t+b_t x_0^*, \sigma_t^2I \right).

x0=μˉt(xt)+rtx_0^*=\bar\mu_t(x_t)+r_t 代入,其均值等于无条件均值再加 btrtb_t r_t。所以条件采样希望噪声修正 σtc\sigma_t c 接近 btrtb_t r_t

ct=argmincCtbtrtσtc2.(4)c_t^* =\mathop{\operatorname{argmin}}_{c\in\mathcal C_t} \left\|b_t r_t-\sigma_t c\right\|^2. \tag{4}

若码字范数近似相等,展开平方后与公式 (3) 等价,因为只剩最大化 c,rt\langle c,r_t\rangle。这说明 argmax 规则是有限候选集上的条件引导近似,而不是任意的启发式最近邻。

为什么还需要随机化版本

纯 argmax 在 KK\to\infty 时会选择越来越极端地对齐残差的方向,不会退化回原 DDPM 随机采样。更一致的有限近似是在码本上按

pt(crt)exp(12cbtσtrt2),cCt(5)p_t(c\mid r_t) \propto \exp\left( -\frac12\left\|c-\frac{b_t}{\sigma_t}r_t\right\|^2 \right), \qquad c\in\mathcal C_t \tag{5}

采样。残差为零时,它近似原标准高斯的有限经验采样;随着 KK 增大,也更有希望恢复连续条件高斯。argmax 是公式 (5) 的 MAP 选择,重建更确定,但牺牲了连续极限与多样性。

编码长度、码本大小与压缩率

若有效记录 TT' 个索引,每个索引需要 log2K\log_2K bit,总码长约为 Tlog2KT'\log_2K。减少采样步数同时缩短编码,却会提高压缩率并增加重建误差;因此普通扩散加速技巧不能无代价地移植到 DDCM。

相比保留二维网格的 VQ 编码,DDCM 的索引按扩散时间天然排列成一维序列,便于接入自回归模型。代价是编码也必须运行完整的多步扩散过程,速度与原 DDPM 同阶,并且码本质量、时间步数和随机/确定选择共同决定率失真权衡。