AIGC 基本功|视频 VAE 的时空压缩结构-VideoVAE

人工智能炼丹君
2026-09-30 / 0 评论 / 2 阅读 / 正在检测是否收录...

视频 VAE 的时空压缩结构

所属方向:表征与压缩 | 难度:进阶 | 前置知识:VAE 结构与训练目标(vae_basics)
关键词:视频VAE、3D因果卷积、时间压缩、分块推理、闪烁伪影、潜变量 token


01. 为什么需要它

先给四个数字,全部来自文末附录里能直接跑的脚本。

数字一:不压时间维,潜变量序列会长到做不了全注意力。 同样一段 121 帧、768×512 的视频,用逐帧图像 VAE(空间 8×8、4 通道)编码出来是 743,424 个 token;改成 4×8×8(时间再压 4 倍、16 通道,CogVideoX / Wan / HunyuanVideo 这一档)是 190,464 个;再激进一点到 8×32×32(LTX-Video 这一档)只剩 6,144 个。自注意力的代价是 $O(N^2)$,所以相对第一档,注意力分别便宜 15.2 倍和 14641 倍(04 节算账,图 4 画出来)。这个差距决定了视频 DiT 能不能做「全时空自注意力」——做不了就只能在空间上做注意力、时间维另想办法,而「时间维另想办法」正是早期视频模型帧间闪烁的根源之一。

数字二:对称时间卷积可能读取未来,本 toy 的每个输出都有这种依赖。 把编码器里的时间填充从「只补前面」换成「前后对称补」,实测 12 个潜变量帧全部都依赖未来帧,最多超前 8 帧。这意味着你没法一边生成一边往外吐——潜变量位置 $j$ 最多要等到输入位置 $4j+8$(与本 toy 总步距 4 对齐),而不是把两侧帧索引直接相加。因果卷积把这个数字压到 0(03 节证明,图 1 左右两栏对比)。

数字三:分块解码不补够上下文,一整块都是错的,不是只有接缝那一帧。 把 24 个潜变量帧切成 3 块、每块 8 帧逐块解码:上下文带 0 帧时,两块合计 64 个输出帧里有 48 帧和整段解码的结果对不上;每块每多带 1 帧上下文,就少错 4 个输出帧(两块合计 8 帧);带到 6 帧时误差精确归零(不是变小,是浮点意义上完全相等)。6 这个数字不是经验值,它等于解码器的时间感受野减 1(03 节推导,图 3 是整条曲线)。

数字四:时间压缩的额度取决于画面运动有多快。 把一个匀速移动的高斯斑点压 8 倍时间再还原,运动速度 0.25 像素/帧时重建 PSNR 是 40.81 dB,速度提到 4 像素/帧掉到 24.92 dB,差 15.9 dB。而且会交叉:压时间 2 倍的曲线在约 2 像素/帧处掉到「压空间 2 倍」这条与速度无关的基线之下——过了这个点,继续压时间不如改压空间(06 节展开,图 5)。

所以这篇文章回答四件事:时间维到底怎么压、何时需要因果卷积、逐块推理要带多少上下文才不接缝、以及时间压缩比能推到多大。


02. 最小可用理解

三句话讲完:

  1. 视频 VAE 相对图像 VAE 只多一件事:在时间轴上再压 $s_T$ 倍。 潜变量形状从 $T \times H' \times W' \times C$ 变成 $T' \times H' \times W' \times C$,其中 $T' = \lfloor (T-1)/s_T \rfloor + 1$。token 数直接除以 $s_T$,注意力代价除以 $s_T^2$。

  2. 需要严格流式处理时,时间运算应满足因果性。 无膨胀、核长 $k_t$ 的卷积可只在前面补 $k_t-1$ 帧,边界可补零或复制首帧。离线视频 VAE 也可以采用非因果结构;卷积因果还不够,归一化、注意力、池化等其他时间运算也须检查。

  3. 代价有两笔,都要记账。 一是时间压缩等价于给运动物体糊上一条长度 $(s_T - 1) \cdot v$ 像素的运动模糊($v$ 是运动速度);二是因果卷积的感受野有限,逐块推理时每块必须额外带「感受野 − 1」帧上下文,带不够就不是接缝难看,是整块算错。

图 1:因果卷积(左)与普通卷积(右)的依赖矩阵

这张图要看什么:横轴是被人为改动的输入帧,纵轴是跟着发生变化的潜变量帧,蓝点表示「这一对有依赖关系」。左图所有蓝点都落在虚线(输入帧 $= 4j$)左边或线上——没有任何一个潜变量帧看到未来;右图蓝点越过虚线,右侧那团就是泄漏的未来信息,实测最多超前 8 帧。两张图除了时间填充方式,其余完全相同。


03. 数学推导

3.1 输出帧数:为什么是 $\lfloor (T-1)/s_T \rfloor + 1$ 而不是 $\lfloor T/s_T \rfloor$

一维卷积(时间轴)在输入长度 $T_{\text{in}}$、核 $k_t$、步距 $s_t$、两端填充 $p_{\text{front}}$ 和 $p_{\text{back}}$ 下的输出长度是

$$T_{\text{out}} = \left\lfloor \frac{T_{\text{in}} + p_{\text{front}} + p_{\text{back}} - k_t}{s_t} \right\rfloor + 1$$

每个符号:$T_{\text{in}}$ 是进入这一层的帧数,$p_{\text{front}}$ 是时间轴前面补的帧数,$p_{\text{back}}$ 是后面补的,$k_t$ 是时间维核长,$s_t$ 是时间步距。

因果卷积取 $p_{\text{front}} = k_t - 1$、$p_{\text{back}} = 0$,代进去:

$$T_{\text{out}} = \left\lfloor \frac{T_{\text{in}} - 1}{s_t} \right\rfloor + 1$$

减 1 加 1 这一对不是凑出来的,它有物理含义:首帧被单独保留下来了。 前面补的那 $k_t - 1$ 帧全是零,所以第一个输出位置看到的是「$k_t - 1$ 个零 + 第 0 帧」,它天生就是第 0 帧的专属输出位。剩下 $T_{\text{in}} - 1$ 帧才按步距 $s_t$ 分组。

许多视频模型约定输入帧数为 $1+k s_T$,但要以完整实现为准。 本文 stride-conv toy 两次时间步距 2 时,49 帧的链路为:

$$49 \to \lfloor 48/2 \rfloor + 1 = 25 \to \lfloor 24/2 \rfloor + 1 = 13$$

得到 13 个 latent 帧。对本 toy,100 帧得到 25 个,最后输出所对齐的输入位置为 96,因此尾部 97—99 没被覆盖;真实模型也可能通过补帧、裁剪或专门的池化分支处理。下面引用的 CogVideoX 对奇偶长度有不同池化逻辑,不能把单层 stride-conv 公式不加条件地替代完整模型。

3.2 因果性:为什么前补零就够了

设时间填充后第 $i$ 个输入帧落在下标 $i + k_t - 1$ 上。步距为 $s_t$ 的第 $j$ 个输出取的是填充后区间 $[j s_t,\; j s_t + k_t - 1]$,对应原始输入下标

$$[j s_t - (k_t - 1),\; j s_t]$$

上界恰好是 $j s_t$——第 $j$ 个输出能用到的最新输入帧就是第 $j s_t$ 帧,未来的帧一个都进不来。这就是因果性的全部证明,它只依赖「$p_{\text{front}} = k_t - 1$、$p_{\text{back}} = 0$」这一条。

堆多层时把每层的时间步距连乘,记到第 $l$ 层输入为止的累积步距为 $S_l = \prod_{m<l} s_m$,则最终第 $j$ 个潜变量帧能用到的最新输入帧是 $S_{\text{total}} \cdot j$。脚本 causality_check() 用扰动法验过:把第 $t$ 帧之后的所有输入都改掉,凡是满足 $S_{\text{total}} \cdot j < t$ 的潜变量帧,变化量精确等于 0——不是小,是零,因为根本没连过来。

换成对称填充($p_{\text{front}} = p_{\text{back}} = (k_t - 1)/2$),上界变成 $j s_t + \lfloor (k_t - 1)/2 \rfloor$,未来的帧就漏进来了。实测 12 个潜变量帧全部泄漏,最多超前 8 帧。

3.3 时间感受野:为什么越靠后的层越值钱

$t$ 时刻的输出能看到多长的历史,叫时间感受野。标准公式是

$$\mathrm{RF} = 1 + \sum_l (k_l - 1) \cdot S_l,\qquad S_l = \prod_{m<l} s_m$$

每个符号:$k_l$ 是第 $l$ 层的时间核长,$S_l$ 是走到第 $l$ 层输入时已经累积的时间步距。关键点在 $S_l$:第 $l$ 层的一个核长 $k_l$,在原始输入上跨的是 $k_l \times S_l$ 帧,因为它的每一个位置本身就代表了 $S_l$ 个原始帧。

本工程的 toy 编码器是 4 层 $k = 3$、步距 $[1, 2, 2, 1]$:

$$\mathrm{RF} = 1 + 2\cdot 1 + 2\cdot 1 + 2\cdot 2 + 2\cdot 4 = 17$$

注意最后两项:同样一个 $k = 3$ 的核,在第 3 层贡献 4 帧,到第 4 层就贡献 8 帧——因为前面已经把时间压了 4 倍。这就是「层数不变、压缩比一上去感受野暴涨」的原因,也是下面那条上下文公式的来源。

图 2:每一层给时间感受野的贡献

这张图要看什么:四根柱子的核长全都是 3,但贡献从 2 帧一路涨到 8 帧,涨的唯一原因是柱底标注的「累积步距」从 1 变成 4。要控制感受野和分块开销,就要一起考虑各层的核长与累计步距。

脚本 receptive_field() 用逐帧扰动实测了同一个数:最大跨度 17 帧,与公式逐项吻合;同时验了因果上界(潜变量帧 $j$ 能看到的最新输入帧 $\le 4j$)零次违反。

3.4 逐块推理:需要的上下文帧数恰好是 $\mathrm{RF} - 1$

由 3.3,输出(或潜变量)帧 $j$ 依赖的输入区间长度是 $\mathrm{RF}$,右端点是 $S_{\text{total}} \cdot j$,所以左端点是 $S_{\text{total}} \cdot j - \mathrm{RF} + 1$。

现在做分块:要正确算出第 $j_0$ 块,就必须拿到它左端点往前的全部输入。缺掉的部分在普通卷积里是被零填充替掉的——零不是正确的值,于是结果错。所以:

$$\text{所需上下文} = \mathrm{RF} - 1 \quad \text{(以该侧的时间单位为帧)}$$

注意单位随你在哪一侧算:

  • 解码侧:单位是潜变量帧。本 toy 解码器的感受野实测是 7 个潜变量帧(两次最近邻上采样会把间距缩小,所以按「潜变量帧」计只有 7 而不是几十),于是需要 6 帧上下文。
  • 编码侧:单位是输入帧。编码器感受野 17 帧,于是需要 16 帧输入上下文。

两条都被脚本验到了:解码侧 ctx 从 0 加到 6,误差在 6 处精确归零;编码侧 ctx 从 0 加到 16,误差在 16 处精确归零。这不是调参调出来的,是感受野直接算出来的。

3.5 时间压缩等价于一条多长的运动模糊

把时间下采样简化成最朴素的「$s$ 帧取平均」(CogVideoX 的下采样层真的就是 avg_pool1d,见 05 节)。一个以速度 $v$ 像素/帧平移的物体,在 $s$ 帧内走过的距离是 $(s-1)v$,所以 $s$ 帧平均等价于把它和一条长度 $L = (s-1)v$ 的盒式核做卷积。

这实际上是 $s$ 个离散平移位置的均匀平均,位移取 $0,v,\ldots,(s-1)v$,其方差为 $v^2(s^2-1)/12$。不是把端点距离直接代入连续盒核的 $L^2/12$。因此,无限画布上按二阶矩定义的宽度满足

$$\sigma_{\text{new}} = \sqrt{\sigma^2 + \frac{v^2(s^2-1)}{12}}$$

当 $s=8,v=4,\sigma=3$ 时,解析宽度是 $\sqrt{9+84}=\sqrt{93}\approx9.644$ 像素;带微噪声并截去噪声底的代码测得约 9.640。差别来自有限画布和阈值测宽,而不是神经网络实验。这只是“时间平均 + 最近邻还原”的模型,不是学习型视频 VAE 的必然误差或压缩下界。


04. 代码实现

环境只要 numpy。全部脚本在文末附录,这里放最核心的三段。

4.1 因果卷积:时间只补前面

def causal_pad(x, k_t):
    """时间轴前面补 k_t-1 帧零,后面不补。这是「因果」二字的全部实现。"""
    if k_t <= 1:
        return x
    pad = np.zeros((x.shape[0], k_t - 1, x.shape[2], x.shape[3]), dtype=x.dtype)
    return np.concatenate([pad, x], axis=1)


def conv3d_causal(x, w, stride=(1, 1, 1), time_pad="causal"):
    """3D 卷积,时间填充方式可选:

    time_pad="causal"     前面补 k_t-1 帧、后面不补(只看过去)
    time_pad="symmetric"  前后各补 (k_t-1)//2 帧(会看到未来)
    """
    kt, kh, kw = w.shape[2], w.shape[3], w.shape[4]
    if time_pad == "causal":
        x = causal_pad(x, kt)
    else:
        pad = ((0, 0), ((kt - 1) // 2, (kt - 1) // 2), (0, 0), (0, 0))
        x = np.pad(x, pad, mode="constant")
    x = sym_pad_hw(x, kh, kw)
    return conv3d_valid(x, w, stride)


def out_frames(t_in, k_t=3, s_t=1):
    """因果卷积的时间输出帧数。注意不是 floor(T_in / s_t)。"""
    return (t_in - 1) // s_t + 1

toy 使用无 batch 的 [C,T,H,W];PyTorch Conv3d 常用带 batch 的 [B,C,T,H,W],也支持无 batch 输入。真实输出:

  T_in=  49 -> s_t=1: 预测  49 / 实测  49 OK  s_t=2: 预测  25 / 实测  25 OK  s_t=4: 预测  13 / 实测  13 OK
  两级时间压缩 2(CogVideoX 口径):49 帧 -> 25 -> 13 潜变量帧

4.2 因果性检验:改未来,看过去

def causality_check():
    """改未来的帧,看过去的潜变量有没有跟着变。变了就是漏了未来信息。"""
    net = ToyVideoVAE()
    v = make_video(t=48)
    z_ref = net.encode(v)

    worst = 0.0
    for t_edit in range(0, 48, 4):
        v2 = v.copy()
        v2[:, t_edit:] += 3.0            # 从第 t_edit 帧起全部改动
        z2 = net.encode(v2)
        # 总时间步距 4,所以潜变量帧 j 只应该看到输入帧 <= 4j
        for j in range(z2.shape[1]):
            if 4 * j < t_edit:
                diff = np.abs(z2[:, j] - z_ref[:, j]).max()
                worst = max(worst, diff)
    print(f"  应该完全不受影响的潜变量帧上,最大变化量 = {worst:.3e}")

真实输出:

  应该完全不受影响的潜变量帧上,最大变化量 = 0.000e+00

把 ToyVideoVAE() 换成 ToyVideoVAE(time_pad="symmetric") 再跑一次 dependency_matrix,会得到「12 个潜变量帧全部泄漏、最多超前 8 帧」——这是 09 节留给读者的第一个动手验证。

4.3 分块解码:接缝是算出来的,不是看出来的

def decode_chunked(net, z, chunk=CHUNK, ctx=0):
    """按 chunk 个潜变量帧一块解码,每块前面带 ctx 帧上下文。"""
    nz = z.shape[1]
    pieces, starts = [], []
    n_chunks = (nz + chunk - 1) // chunk
    for i in range(n_chunks):
        lo = i * chunk
        hi = min(lo + chunk, nz)
        lo_in = max(0, lo - ctx)
        out = net.decode(z[:, lo_in:hi])
        keep = (hi - lo) * T_STRIDE        # 上下文部分的输出要丢掉
        pieces.append(out[:, -keep:])
        starts.append(lo * T_STRIDE)
    return np.concatenate(pieces, axis=1), starts

96 帧输入 $\to$ 24 个潜变量帧 $\to$ 切 3 块,真实输出(误差按输出标准差归一):

  ctx | 块头误差 | 块尾误差 | 块内被污染帧数 | 边界跳变放大
  ----+---------+---------+----------------+------------
     0 | 2.550e+00 | 3.206e+00 |             48 | 1.90x
     1 | 2.094e+00 | 1.814e+00 |             40 | 2.19x
     2 | 1.658e+00 | 1.399e+00 |             32 | 1.88x
     3 | 1.313e+00 | 8.267e-01 |             24 | 2.19x
     4 | 8.700e-01 | 4.549e-01 |             16 | 1.15x
     5 | 4.051e-01 | 0.000e+00 |              8 | 1.07x
     6 | 0.000e+00 | 0.000e+00 |              0 | 1.00x

三件事要读出来:

  • 误差在 ctx = 6 处精确归零,不是渐近变小。少 1 帧(ctx = 5)还剩 8 个输出帧是错的。
  • 每块每少带 1 帧上下文,就多错 4 个输出帧(表中两块合计 8 帧),正好等于 1 个潜变量帧的输出跨度($s_T = 4$)。所以「接缝」根本不是一条线,是一段区域。
  • 块尾误差比块头误差先归零(ctx = 5 时块尾已经是 0 而块头还有 0.405)。这符合感受野的形状:越靠块尾,缺的上下文越少。

编码侧同理,真实输出(每块 32 个输入帧,误差按潜变量标准差归一):

  ctx | 块头潜变量误差 | 块尾潜变量误差 | 被污染潜变量帧数
  ----+----------------+----------------+--------------
     0 | 3.245e+00 | 2.984e+00 |              4
     8 | 1.297e+00 | 0.000e+00 |              2
    12 | 6.090e-01 | 0.000e+00 |              1
    16 | 0.000e+00 | 0.000e+00 |              0

阈值 16,正是编码器感受野 $17 - 1$。

图 3:分块解码的接缝误差随上下文帧数的变化

这张图要看什么:左图两条线(块头、块尾)在 ctx = 6 处一起掉到 0,纵坐标是对数轴——前面那段下降看着平缓,其实是从「两个标准差」这种肉眼可见的错误降到零。右图的阶梯更直白:被污染帧数每块减少 4 帧、图中两块合计减少 8 帧,斜率就是 $s_T$。

4.4 token 账本

def latent_frames(t_in, s_t):
    """因果卷积下的潜变量帧数:首帧单独占位,所以是 floor((T-1)/s)+1。"""
    return (t_in - 1) // s_t + 1

121 帧 / 768×512 / RGB 的真实输出:

  配置                                  | 潜变量形状 (T'xH'xW'xC) | token 数 N | N^2 相对代价 | 便宜倍数 | 每 token 覆盖像素
  图像 VAE 逐帧(SVD / AnimateDiff 口径)| 121 x  64 x  96 x   4 |    743,424 |   1.000e+00 |     1.0x |      192
  CogVideoX / Wan / HunyuanVideo(4x8x8)|  31 x  64 x  96 x  16 |    190,464 |   6.564e-02 |    15.2x |      768
  LTX-Video(8x32x32)                  |  16 x  16 x  24 x 128 |      6,144 |   6.830e-05 | 14641.0x |   24,576

只动时间维的对照(固定空间 8×8、通道 16):时间压缩 1→2→4→8,token 数 743,424→374,784→190,464→98,304,长序列下近似每翻一倍注意力代价降到 1/4,首帧保留会带来取整差异。

图 4:同一个输入下的 token 数与注意力代价

这张图要看什么:左图是序列长度(对数轴),右图是注意力的相对代价——右图的差距比左图大得多,因为代价是 $N^2$。6,144 个 token 意味着 LTX-Video 可以在这个分辨率上直接做全时空自注意力,而 743,424 个 token 连存 attention map 都存不下。


05. 工业级实现对照

最小实现只有一条时间前补零,生产代码多了五处,每处都有原因。以下均以 huggingface/diffusers 2026-09 的实现为准(上游会重构,引用时请对照 repo/path#symbol):

  • src/diffusers/models/autoencoders/autoencoder_kl_cogvideox.py#CogVideoXCausalConv3d
  • src/diffusers/models/autoencoders/autoencoder_kl_wan.py#WanCausalConv3d
  • src/diffusers/models/downsampling.py#CogVideoXDownsample3D
  • src/diffusers/models/upsampling.py#CogVideoXUpsample3D

5.1 填充元组:时间只补前面

src/diffusers/models/autoencoders/autoencoder_kl_cogvideox.py#CogVideoXCausalConv3d:

time_pad = time_kernel_size - 1
self.time_causal_padding = (width_pad, width_pad, height_pad, height_pad, time_pad, 0)

F.pad 的元组是从最后一维往前写的,所以这六个数的含义是 $(W_{\text{left}}, W_{\text{right}}, H_{\text{left}}, H_{\text{right}}, T_{\text{left}}, T_{\text{right}})$。最后两个是 $(k_t - 1,\ 0)$——和 3.2 节推导的 $p_{\text{front}} = k_t - 1$、$p_{\text{back}} = 0$ 完全一致。空间两维是对称的,只有时间轴不对称。

Wan 的 WanCausalConv3d(autoencoder_kl_wan.py)写成另一种形式:

self._padding = (self.padding[2], self.padding[2], self.padding[1], self.padding[1],
                 2 * self.padding[0], 0)
self.padding = (0, 0, 0)

它先把 nn.Conv3d 自带的 padding 清零,改成自己手动 pad。$k_t = 3$ 时 padding[0] = 1,于是 $2 \cdot 1 = 2 = k_t - 1$,和 CogVideoX 殊途同归。

5.2 上下文缓存:分块推理靠它

CogVideoX 的 fake_context_parallel_forward:

kernel_size = self.time_kernel_size
if kernel_size > 1:
    cached_inputs = [conv_cache] if conv_cache is not None else [inputs[:, :, :1]] * (kernel_size - 1)
    inputs = torch.cat(cached_inputs + [inputs], dim=2)
return inputs

读法:有缓存就把上一块最后 $k_t - 1$ 帧拼到前面;是第一块(没缓存)就把第 0 帧复制 $k_t - 1$ 次。这正好对应 3.4 节的两种情形——「缺上下文就用别的东西填」,而复制首帧比补零更合理(这是 pad_mode 的一个选项)。紧接着:

conv_cache = inputs[:, :, -self.time_kernel_size + 1:].clone()

把拼接后输入的最后 $k_t - 1$ 帧存起来,给下一块用。整个编码器把这些缓存收进一个以层名索引的字典(conv_in、down_block_0……)逐层传递。Wan 的对应逻辑更紧凑,直接把缓存长度从待补的填充里减掉:

if cache_x is not None and self._padding[4] > 0:
    x = torch.cat([cache_x, x], dim=2)
    padding[4] -= cache_x.shape[2]

这就是 04 节那个 ctx 在生产代码里的样子。 这里有个容易看错的地方:工业实现每层只缓存 $k_t - 1 = 2$ 帧,而 04 节实测需要 6 帧,直觉上会觉得「2 帧不够」。其实够——因为每层缓存的是该层自己输入分辨率下的帧,越往后的层分辨率越高,同样 2 帧折回潜变量帧单位就越小。把解码器 d0…d4 逐层折算后累加:

层 输入相对潜变量的帧间距 缓存 $(k_t - 1)$ 帧折回潜变量帧
d0 1.00 2.00
d1 1.00 2.00
d2 0.50 1.00
d3 0.25 0.50
d4 0.25 0.50
合计 6.00

逐层累加 = 6.00,和 04 节实测需要的 6 帧完全吻合——逐层缓存本来就等价于整条链路的感受野减 1,前提是每一层都缓存。真正的陷阱是只在网络入口缓存一次,那样只有 2 帧,接缝一定在。

5.3 时间下采样:CogVideoX 用的是平均池化,不是学习出来的卷积

src/diffusers/models/downsampling.py#CogVideoXDownsample3D:

x = x.permute(0, 3, 4, 1, 2).reshape(batch_size * height * width, channels, frames)
if x.shape[-1] % 2 == 1:
    x_first, x_rest = x[..., 0], x[..., 1:]
    x_rest = F.avg_pool1d(x_rest, kernel_size=2, stride=2)
    x = torch.cat([x_first[..., None], x_rest], dim=-1)
else:
    x = F.avg_pool1d(x, kernel_size=2, stride=2)

两个细节值得记住:

  1. 时间维是 avg_pool1d,不是 stride 卷积。 空间维倒是用 nn.Conv2d(stride=2)(注意它把 (B, T, C, H, W) 折叠成 (B*T, C, H, W) 后用 2D 卷积做的,目的是省 Conv3d 的显存)。所以 3.5 节把时间下采样建模成平均池化不是偷懒,它就是这个实现。
  2. 帧数为奇数时首帧单独保留,不参与池化——和 3.1 节「首帧单独占一个输出位」是同一件事的两种写法。

上采样端(upsampling.py#CogVideoXUpsample3D)对称地用 F.interpolate 做时间 2 倍、空间 2 倍,并且同样对首帧单独处理。

另外编码器里有一行决定「在哪几层压时间」:

temporal_compress_level = int(np.log2(temporal_compression_ratio))
compress_time = i < temporal_compress_level

压缩比 4 → level = 2 → 只有前两个下采样块压时间。这也解释了图 2 那件事:压缩集中在前段时,后段的层会以更大的累积步距去看历史,感受野涨得最快。

5.4 与最小实现的五处差距

差距 生产代码 为什么
归一化 GroupNorm / CogVideoXSpatialNorm3D 调节中间激活的尺度;不能代替 KL,也不保证 latent 标准正态。若统计量跨时间,还须单独检查因果性
激活与结构 ResNet block + SiLU,不是单层卷积 单层卷积的表达能力撑不起 8×8 的空间压缩
卷积实现 CogVideoXSafeConv3d(分块跑的 Conv3d) 避免长视频上 Conv3d 一次性分配大显存
缓存粒度 每层一个 key 的字典 分块推理要逐层续接,不是只在入口续接
精度与缓存 dtype 由模型配置决定;缓存可 .clone() bf16 降低存储,clone 控制缓存所有权;二者是不同问题

06. 代价与边界

数值个数、序列长度与信息量是三件事。 这张表最容易看错:

  配置                                  | 潜变量数值总数 | 压缩比(像素值 / 潜变量数值)
  图像 VAE 逐帧(SVD / AnimateDiff 口径)|      2,973,696 |     48.0 : 1
  CogVideoX / Wan / HunyuanVideo(4x8x8)|      3,047,424 |     46.8 : 1
  LTX-Video(8x32x32)                  |        786,432 |    181.5 : 1

4×8×8 这一档的数值元素压缩比和逐帧图像 VAE 几乎一样(46.8 vs 48.0)。它在这个数值元素口径下没有更小,但并非无损重排:学习型编码器把 743,424 个 token 重排成 190,464 个更「厚」的 token(通道 4 → 16)。数值个数接近不代表信息容量相同;这张表直接说明的是序列长度。 真正开始省容量的是 8×32×32 这一档,代价是 LTX-Video 摘要里自己写的那句:"the high compression inherently limits the representation of fine details"(高压缩本质上限制了细节的表达)。

时间压缩的额度由运动速度决定。 图 5 的三组曲线:

图 5:运动速度 vs 时间压缩的重建质量

这张图要看什么:实线是压时间,虚线是压空间(与速度无关,当基线用)。同一颜色实线往下掉、虚线不动,交叉之后「再压时间」就不如「改压空间」。实测交叉点:压 2 倍和 4 倍在约 2 像素/帧处,压 8 倍在约 4 像素/帧处。

  速度 v |  时间 2x / 空间 2x    时间 4x / 空间 4x    时间 8x / 空间 8x
   0.25  |   53.49 /  39.00      46.87 /  32.35      40.81 /  26.58
   1.00  |   41.97 /  39.00      35.28 /  32.35      30.05 /  26.58
   2.00  |   36.15 /  39.00      30.18 /  32.35      26.62 /  26.58
   4.00  |   30.83 /  39.00      26.69 /  32.29      24.92 /  26.57

这些交叉点只属于本实验的斑点、尺寸、噪声和滤波器。尤其“时间 s 倍”减少 s 倍元素,而“空间两轴各 s 倍”减少 s² 倍元素,并非等码率比较,不能据此给出“超过 2 像素/帧就不能压缩”的部署阈值。真实 VAE 需在相同码率、动作数据和感知/时序指标下做验证。

首帧总是吃亏的。 因果卷积前面补的是零(或复制首帧),第一个输出位置天然只能看到 1 帧输入。图 1 左图第 0 行只有一个蓝点就是这个意思。所以「图生视频」任务里把首帧单独处理、或者编码时多给一帧,是常见做法。

上下文是纯开销。 带够 $\mathrm{RF} - 1$ 帧上下文,意味着每块要多算:

  解码块大小 | 需要上下文 | 额外算力占比
     8 潜变量帧 |   6 帧 | 42.9%
    16 潜变量帧 |   6 帧 | 27.3%
    24 潜变量帧 |   6 帧 | 20.0%
  编码块大小 | 需要上下文 | 额外算力占比
    32 输入帧 |  16 帧 | 33.3%
   128 输入帧 |  16 帧 | 11.1%

上下文的绝对量是固定的,块越大摊得越薄。这就是「块不能切太小」的定量理由:块切成 8 帧,42.9% 的算力花在重复计算上,此时「分块省显存」的收益已经被吃掉了三分之一。

离线场景可以考虑非因果结构。 它能使用前后帧,但是否提高质量要由实验决定。非因果网络也能分块,只是需要左右上下文或接受输出延迟。因果性本身不保证无闪烁,也不保证完整模型严格流式:跨时间归一化或全局注意力仍可能破坏性质。


07. 经典论文脉络

五篇,每篇一句话说清它对「时间维怎么处理」的贡献:

  1. VQ-VAE(arXiv:1711.00937,2017) — 提出「先压成离散 token 再建模」的范式。它本身是图像的,但整套 video tokenizer 都是从它长出来的。
  2. MAGVIT(arXiv:2212.05199,2022) — 把 3D 卷积 + 时间下采样正式带进视频 tokenizer,用 masked modeling 训练,是这一方向的一项代表工作。(更早的 VideoGPT,arXiv:2104.10157,是另一条路:VQ-VAE + 自回归 Transformer。)
  3. MAGVIT-v2(arXiv:2310.05737,2023) — 换掉矢量量化,用 lookup-free 量化把词表做大,让离散视频 token 的质量第一次追上扩散。标题那句 "Tokenizer is Key" 就是这一支的纲领。
  4. CogVideoX(arXiv:2408.06072,2024) — 本文的锚点。把 3D 因果 VAE 和 Expert Transformer 绑在一起,明确以「提高压缩率同时保住保真度」为目标;它也是这套因果缓存实现被广泛复用的起点。同期的 HunyuanVideo(arXiv:2412.03603)与 Wan(arXiv:2503.20314)都走 4×8×8 这一档。
  5. LTX-Video(arXiv:2501.00103,2025) — 把压缩比推到 8×32×32(摘要自述 1:192),代价是细节;它的解法是让解码器顺手把最后一步去噪也做了,直接在像素空间出结果。

还有一条不压时间的路线:Stable Video Diffusion 使用空间压缩的图像自编码器并加入时序解码层;Align your Latents 是较早的 Video LDM 工作,不能把两个标题和论文 ID 混在一起。逐帧 latent 的长度更大,但生成器是否采用全时空、分离式注意力或其他结构会改变真实成本。


08. 常见误解

误解一:「时间压缩 4 倍时,48 帧一定输出 13 帧。」 本文单侧填充 stride-conv 公式给出 $\lfloor47/4\rfloor+1=12$,49 帧才给 13。许多真实模型限定 $1+k s_T$ 帧,其他输入长度会经过补齐、裁剪或不同池化分支,必须核对代码,不能把帧数规则泛化。

误解二:「只要卷积前补零,整个 VAE 就严格因果。」 单层卷积的因果范围可以逐项证明,但还要检查膨胀、池化对齐、归一化和注意力的时间依赖。分块既可逐层缓存中间特征,也可在入口重叠足够大的上下文后裁剪;后者在本 toy 需要 6 个 latent 帧,而不是只缓存 2 帧。两种方式都可正确,只是计算和内存开销不同。

误解三:「分块解码只要重叠 1 帧就够,接缝最多难看一点点。」 实测:缺上下文时不是接缝那一帧错,是整块错。块大小 8 个潜变量帧、上下文 0 帧时,32 个输出帧里 24 帧是错的;每少带 1 帧上下文多错 4 个输出帧。而且这个错误不是「视觉上略糊」,是数值上完全跑偏(相对标准差 2.55 倍)。

误解四:「压缩比越高越好,反正都是 VAE 重建。」 4×8×8 相对逐帧图像 VAE 的数值元素压缩比几乎没变(46.8:1 vs 48.0:1),它省的是序列长度(06 节第一张表)。真正省容量的是 8×32×32 那一档,而论文自己承认细节受限。所以「压缩比」这个词在这件事上有两个含义,混着用会得出完全相反的结论。

误解五:「视频 VAE 就是把图像 VAE 的 2D 卷积换成 3D 卷积。」 三处不一样:(a)CogVideoX 的时间池化本身没有可学习参数,但其前后的 3D 时间卷积有;(b)空间卷积被折叠成 2D 卷积做,为的是省 Conv3d 的显存;(c)temporal_compress_level = int(log2(ratio)) 决定只有前几层压时间,不是每层都压。想从图像 VAE 权重 inflate 一个视频 VAE,这三处都得单独处理。


09. 动手验证

五个脚本都在文末附录;数值实验依赖 numpy,make_figures.py 还依赖 matplotlib。运行时间取决于机器。

python causal_conv3d.py     # 形状公式 + 因果性 + 感受野(约 11 秒)
python chunk_decode.py      # 分块编解码的接缝(约 8 秒)
python temporal_budget.py   # 运动速度 vs 时间压缩
python token_ledger.py      # token 账本
python make_figures.py      # 重画本文 5 张图

预期结果:

  • causal_conv3d.py:形状公式 6 组输入、3 组步距预测值全部等于实测值;49 -> 25 -> 13;因果性违反量 0.000e+00;感受野公式 17 = 实测 17,因果上界 0 次违反。
  • chunk_decode.py:解码侧 ctx = 6 时误差 0.000e+00;编码侧 ctx = 16 时误差 0.000e+00;被污染帧数随 ctx 每 +1 减 8(两块合计)。
  • temporal_budget.py:s=8 那一列 PSNR 从 40.81 dB 掉到 24.92 dB;宽度增幅实测约 3.24 倍,修正离散方差后的解析值也约 3.24 倍;交叉点报在 2 / 2 / 4 像素/帧。

三个可以自己改的小实验:

  1. 把因果改成非因果:dependency_matrix(t_in=48, time_pad="symmetric"),看所有 12 个潜变量帧都越过 $4j$ 那条线,最多超前 8 帧。
  2. 把解码器的 latent 级卷积从 2 层加到 4 层:感受野会变成 11 个潜变量帧,所需上下文从 6 涨到 10,重跑 chunk_decode.py 会看到阈值跟着移动。这条最能验证「上下文 = 感受野 − 1」不是巧合。
  3. 把 toy 编码器的时间步距从 $[1,2,2,1]$ 改成 $[1,1,2,2]$:总压缩比不变(还是 4),但累积步距的分布变了,感受野从 17 变成 $1+2+2+2+4=11$。压缩比相同,上下文开销不同——这就是选架构时真正该看的量。

10. 延伸阅读

按依赖顺序:

  • vae_basics(VAE 结构与训练目标) — 本文默认你已经知道 KL 项和重参数化在干什么。若要接着问「潜变量为什么要近似标准正态」,看它。
  • vae_losses(视频 VAE 的常见 loss 组合) — 本文只讲了结构没讲训练目标。L1 + KL + LPIPS + GAN 四项权重怎么配、谁在管伪影,是那篇的话题;帧间闪烁的根因有一半在那儿,不在本文的因果性里。
  • latent_diffusion(潜空间扩散) — 潜变量一旦定了,扩散过程就全在它上面做,包括那个 scaling factor。
  • dit(DiT:用 Transformer 替掉 UNet) — 01 节那个 $N^2$ 账单最后是要 DiT 来付的,token 数直接决定它能不能做全时空注意力。
  • autoregressive_video(自回归视频生成与 Forcing 范式) — 本文讲的是「逐块解码不接缝」;那篇讲的是更进一步:干脆按帧自回归地生成,因果性从 VAE 一路贯穿到生成模型。

附录:完整代码

09 节用到的脚本全文如下(make_figures.py、causal_conv3d.py、chunk_decode.py、temporal_budget.py、token_ledger.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。

make_figures.py

# -*- coding: utf-8 -*-
"""画配图。所有数字都来自同目录的实验脚本,不另算一遍。

运行: python make_figures.py [--only 图名]

标签一律用 Unicode(σ、×)而不是 mathtext,省得反斜杠踩到转义检查。
"""

import os
import sys

import numpy as np
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt

sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))

import causal_conv3d as cc                            # noqa: E402
import chunk_decode as cd                             # noqa: E402
import token_ledger as tl                             # noqa: E402
import temporal_budget as tb                          # noqa: E402

HERE = os.path.dirname(os.path.abspath(__file__))
FIGDIR = os.path.join(os.path.dirname(HERE), "figures")
os.makedirs(FIGDIR, exist_ok=True)

C0 = "#2f4b7c"
C1 = "#d45087"
C2 = "#f0a35e"
C3 = "#4c9f70"
CGREY = "#8a8a8a"

plt.rcParams["font.sans-serif"] = ["Heiti TC", "Arial Unicode MS", "DejaVu Sans"]
plt.rcParams["axes.unicode_minus"] = False


def _save(fig, name):
    p = os.path.join(FIGDIR, name)
    fig.savefig(p, dpi=150, bbox_inches="tight", facecolor="white")
    plt.close(fig)
    print(f"  [ok] {name}  ({os.path.getsize(p)} bytes)")


# ─────────────── 图 1:因果 vs 非因果的依赖矩阵 ───────────────
def fig_dependency():
    dep_c = cc.dependency_matrix(t_in=48, time_pad="causal")
    dep_n = cc.dependency_matrix(t_in=48, time_pad="symmetric")
    fig, axes = plt.subplots(1, 2, figsize=(11.2, 4.6), sharey=True)
    for ax, dep, title in ((axes[0], dep_c, "因果卷积:只看过去"),
                           (axes[1], dep_n, "普通卷积:前后都看")):
        ax.imshow(dep.T, aspect="auto", cmap="Blues", interpolation="nearest",
                  origin="lower")
        nz = dep.shape[1]
        jj = np.arange(nz)
        ax.plot(4 * jj, jj + 0.5, color=C1, lw=1.6, ls="--",
                label="输出帧 j 对应的输入帧 4j")
        ax.set_xlabel("输入帧 t")
        ax.set_title(title, fontsize=12)
        ax.legend(loc="upper left", fontsize=9)
    axes[0].set_ylabel("潜变量帧 j")
    # 标一处「未来泄漏」
    leak = []
    for j in range(dep_n.shape[1]):
        idx = np.where(dep_n[:, j])[0]
        if len(idx) and idx.max() > 4 * j:
            leak.append(int(idx.max()) - 4 * j)
    if leak:
        axes[1].annotate(f"这里看到了未来\n最多超前 {max(leak)} 帧",
                         xy=(30, 6), xytext=(24, 9.5), fontsize=10, color="#333333",
                         arrowprops=dict(arrowstyle="->", color="#999999", lw=1))
    fig.suptitle("图 1:谁依赖谁——横轴是被改动的输入帧,纵轴是跟着变的潜变量帧",
                 fontsize=13)
    fig.tight_layout()
    _save(fig, "fig1_dependency.png")
    return dep_c, dep_n


# ─────────────── 图 2:感受野为什么会被步距放大 ───────────────
def fig_rf():
    kernels = [3, 3, 3, 3]
    strides = [1, 2, 2, 1]
    contrib, cum = [], []
    c = 1
    for k, s in zip(kernels, strides):
        contrib.append((k - 1) * c)
        cum.append(c)
        c *= s
    rf = 1 + sum(contrib)
    fig, ax = plt.subplots(figsize=(7.6, 4.4))
    x = np.arange(len(kernels))
    ax.bar(x, contrib, color=[C0, C0, C2, C1], width=0.62)
    for i, (v, cu) in enumerate(zip(contrib, cum)):
        ax.text(i, v + 0.25, f"+{v}", ha="center", fontsize=10)
        ax.text(i, -0.9, f"累积步距 {cu}", ha="center", fontsize=9, color="#555555")
    ax.axhline(0, color="#333333", lw=0.8)
    ax.set_xticks(x)
    ax.set_xticklabels([f"第 {i + 1} 层\nk={k}, s={s}"
                        for i, (k, s) in enumerate(zip(kernels, strides))])
    ax.set_ylabel("这一层给感受野贡献的帧数")
    ax.set_ylim(-1.6, max(contrib) + 1.6)
    ax.set_title(f"图 2:同样的 k=3,越靠后的层贡献越大(总感受野 = 1 + 各项 = {rf} 帧)",
                 fontsize=12)
    ax.grid(alpha=0.25, axis="y")
    fig.tight_layout()
    _save(fig, "fig2_rf.png")
    return contrib, rf


# ─────────────── 图 3:分块解码的接缝 ───────────────
def fig_chunk():
    net, v, z, ref = cd.build()
    scale = float(ref.std())
    ctxs = list(range(0, 9))
    head, tail, contam = [], [], []
    for ctx in ctxs:
        got, starts = cd.decode_chunked(net, z, cd.CHUNK, ctx)
        diff = np.abs(got - ref).max(axis=(0, 2, 3)) / scale
        hi, ti, cn = [], [], 0
        for s in starts[1:]:
            hi += list(range(s, s + cd.T_STRIDE))
            ti += list(range(s + cd.T_STRIDE, s + cd.CHUNK * cd.T_STRIDE))
            cn += int((diff[s:s + cd.CHUNK * cd.T_STRIDE] > 1e-9).sum())
        head.append(float(diff[hi].max()))
        tail.append(float(diff[ti].max()))
        contam.append(cn)

    fig, axes = plt.subplots(1, 2, figsize=(11.6, 4.5))
    ax = axes[0]
    ax.plot(ctxs, head, "o-", color=C1, lw=2.2, label="块头误差(块的前 4 个输出帧)")
    ax.plot(ctxs, tail, "s--", color=C0, lw=2.0, label="块尾误差(其余输出帧)")
    ax.set_yscale("symlog", linthresh=1e-3)
    ax.set_xlabel("每块前面带的潜变量上下文帧数 ctx")
    ax.set_ylabel("相对整段解码的最大误差(按输出标准差归一)")
    ax.axvline(6, color=C3, lw=1.6, ls=":")
    ax.annotate("ctx=6 起误差精确归零", xy=(6, 1e-1), xytext=(6.4, 4e-1),
                fontsize=10, color=C3,
                arrowprops=dict(arrowstyle="->", color=C3, lw=1))
    ax.set_title("左:误差随上下文帧数的变化", fontsize=12)
    ax.legend(fontsize=9)
    ax.grid(alpha=0.25)

    ax = axes[1]
    ax.plot(ctxs, contam, "o-", color=C2, lw=2.2)
    ax.set_xlabel("每块前面带的潜变量上下文帧数 ctx")
    ax.set_ylabel("被污染的块内输出帧数(两块合计)")
    for i, c in enumerate(contam):
        ax.annotate(str(c), (ctxs[i], c), textcoords="offset points",
                    xytext=(0, 7), ha="center", fontsize=9)
    ax.set_title("右:少带一帧上下文,就多错 4 个输出帧", fontsize=12)
    ax.grid(alpha=0.25)
    fig.suptitle("图 3:分块解码的接缝——上下文不是越多越好,而是有个精确阈值",
                 fontsize=13)
    fig.tight_layout()
    _save(fig, "fig3_chunk.png")
    return ctxs, head, tail, contam


# ─────────────── 图 4:token 账本 ───────────────
def fig_ledger():
    rows = [tl._row(n, a, b, c) for n, a, b, c in tl.REAL]
    labels = ["逐帧图像 VAE\n(1x8x8)", "4x8x8\n(CogVideoX/Wan)", "8x32x32\n(LTX-Video)"]
    ns = [r["n"] for r in rows]
    rel = [(r["n"] / ns[0]) ** 2 for r in rows]
    fig, axes = plt.subplots(1, 2, figsize=(11.6, 4.5))
    ax = axes[0]
    bars = ax.bar(labels, ns, color=[CGREY, C0, C3], width=0.58)
    ax.set_yscale("log")
    ax.set_ylabel("潜变量 token 数 N(对数轴)")
    for b, n in zip(bars, ns):
        ax.text(b.get_x() + b.get_width() / 2, n * 1.15, f"{n:,}",
                ha="center", fontsize=9)
    ax.set_ylim(1e3, 3e6)
    ax.set_title("左:序列长度", fontsize=12)
    ax.grid(alpha=0.25, axis="y")

    ax = axes[1]
    bars = ax.bar(labels, [1.0 / r for r in rel], color=[CGREY, C0, C3], width=0.58)
    ax.set_yscale("log")
    ax.set_ylabel("注意力代价相对「逐帧图像 VAE」便宜多少倍")
    for b, r in zip(bars, rel):
        ax.text(b.get_x() + b.get_width() / 2, (1.0 / r) * 1.15,
                f"{1.0 / r:.0f}x", ha="center", fontsize=9)
    ax.set_ylim(0.8, 6e4)
    ax.set_title("右:因为代价是 N 的平方,差距被放大", fontsize=12)
    ax.grid(alpha=0.25, axis="y")
    fig.suptitle("图 4:同一个 121 帧 768x512 输入,压缩比决定了 DiT 能不能做全注意力",
                 fontsize=13)
    fig.tight_layout()
    _save(fig, "fig4_ledger.png")
    return ns, rel


# ─────────────── 图 5:运动速度 vs 时间压缩 ───────────────
def fig_motion():
    speeds = tb.SPEEDS
    p_t = {2: [], 4: [], 8: []}
    p_s = {2: [], 4: [], 8: []}
    for v in speeds:
        x = tb.moving_blob(v)
        for s in (2, 4, 8):
            p_t[s].append(tb.psnr(x, tb.repeat_time(tb.avg_pool_time(x, s), s)))
            p_s[s].append(tb.psnr(x, tb.repeat_space(tb.avg_pool_space(x, s), s)))
    fig, ax = plt.subplots(figsize=(8.2, 5.0))
    colors = {2: C0, 4: C2, 8: C1}
    for s in (2, 4, 8):
        ax.plot(speeds, p_t[s], "o-", color=colors[s], lw=2.2,
                label=f"压时间 {s} 倍")
    for s in (2, 4, 8):
        ax.plot(speeds, p_s[s], ls="--", color=colors[s], lw=1.6, alpha=0.75,
                label=f"空间两轴各 {s} 倍(元素压 {s*s} 倍)")
    ax.set_xscale("log")
    ax.set_xticks(speeds)
    ax.set_xticklabels([str(v) for v in speeds])
    ax.set_xlabel("运动速度(像素 / 帧,对数轴)")
    ax.set_ylabel("重建 PSNR(dB)")
    ax.set_title("图 5:时间平均与空间平均的教学对照(并非等码率)",
                 fontsize=12)
    ax.legend(fontsize=9, ncol=2)
    ax.grid(alpha=0.25)
    fig.tight_layout()
    _save(fig, "fig5_motion.png")
    return p_t, p_s


if __name__ == "__main__":
    only = sys.argv[2] if len(sys.argv) > 2 and sys.argv[1] == "--only" else None
    jobs = {
        "dependency": fig_dependency,
        "rf": fig_rf,
        "chunk": fig_chunk,
        "ledger": fig_ledger,
        "motion": fig_motion,
    }
    for k, fn in jobs.items():
        if only and k != only:
            continue
        print(f"[draw] {k}")
        fn()
    print("done")

causal_conv3d.py

# -*- coding: utf-8 -*-
"""3D 因果卷积的 numpy 最小实现,以及形状公式、因果性、感受野的实测。

运行: python causal_conv3d.py

约定张量布局 [C, T, H, W](通道优先,和 torch 的 Conv3d 一致)。
因果卷积的全部秘密只有一条:时间轴**只在前面**补 k_t - 1 帧,后面一帧都不补。
"""

import numpy as np


def base_rng(seed=20260930):
    """全工程的种子入口。所有随机权重都从这里分叉,保证数字可复现。"""
    return np.random.default_rng(seed)


# ─────────────── 基本算子 ───────────────

def causal_pad(x, k_t):
    """时间轴前面补 k_t-1 帧零,后面不补。这是「因果」二字的全部实现。"""
    if k_t <= 1:
        return x
    pad = np.zeros((x.shape[0], k_t - 1, x.shape[2], x.shape[3]), dtype=x.dtype)
    return np.concatenate([pad, x], axis=1)


def sym_pad_hw(x, k_h, k_w):
    """空间轴对称补零,和 torch 的 Conv3d 默认行为一致。"""
    pad = ((0, 0), (0, 0),
           ((k_h - 1) // 2, k_h // 2), ((k_w - 1) // 2, k_w // 2))
    return np.pad(x, pad, mode="constant")


def conv3d_valid(x, w, stride=(1, 1, 1)):
    """x [C_in,T,H,W] × w [C_out,C_in,kT,kH,kW] -> [C_out,To,Ho,Wo],只做 valid 部分。"""
    c_in, t, h, wd = x.shape
    c_out, c_in2, kt, kh, kw = w.shape
    assert c_in == c_in2, "通道数对不上"
    st, sh, sw = stride
    to = (t - kt) // st + 1
    ho = (h - kh) // sh + 1
    wo = (wd - kw) // sw + 1
    out = np.zeros((c_out, to, ho, wo), dtype=np.float64)
    for o in range(c_out):
        wo_kernel = w[o]
        for i in range(to):
            for j in range(ho):
                for k in range(wo):
                    patch = x[:, i * st:i * st + kt,
                              j * sh:j * sh + kh,
                              k * sw:k * sw + kw]
                    out[o, i, j, k] = np.sum(patch * wo_kernel)
    return out


def conv3d_causal(x, w, stride=(1, 1, 1), time_pad="causal"):
    """3D 卷积,时间填充方式可选:

    time_pad="causal"   前面补 k_t-1 帧、后面不补(只看过去)
    time_pad="symmetric" 前后各补 (k_t-1)//2 帧(会看到未来)
    """
    kt, kh, kw = w.shape[2], w.shape[3], w.shape[4]
    if time_pad == "causal":
        x = causal_pad(x, kt)
    else:
        pad = ((0, 0), ((kt - 1) // 2, (kt - 1) // 2), (0, 0), (0, 0))
        x = np.pad(x, pad, mode="constant")
    x = sym_pad_hw(x, kh, kw)
    return conv3d_valid(x, w, stride)


def out_frames(t_in, k_t=3, s_t=1):
    """因果卷积的时间输出帧数。

    T_out = floor((T_in - 1) / s_t) + 1
    注意不是 floor(T_in / s_t):因为首帧被保留下来单独占了一个输出位。
    """
    return (t_in - 1) // s_t + 1


# ─────────────── 一个可复用的 toy 视频 VAE ───────────────

class CausalConv3d:
    """带激活的因果卷积层。权重由 base_rng 派生,fix 住后所有实验共用。

    act 默认用 tanh 而不是 ReLU:随机权重下 ReLU 会把一大半通道关死,
    实测感受野会被「死通道」削小,接缝误差也随之小到 1e-7 量级,看不出问题。
    真实 VAE 里有 GroupNorm 兜着,绝大多数通道是活的,tanh 更接近那个状态。
    """

    def __init__(self, c_in, c_out, k=(3, 3, 3), s=(1, 1, 1), seed=0, act="tanh",
                 time_pad="causal"):
        rng = base_rng(1000 + seed)
        scale = 0.9 / np.sqrt(c_in * k[0] * k[1] * k[2])
        self.w = rng.normal(0.0, scale, size=(c_out, c_in) + tuple(k))
        self.b = rng.normal(0.0, 0.02, size=c_out)
        self.k = tuple(k)
        self.s = tuple(s)
        self.act = act
        self.time_pad = time_pad

    def __call__(self, x):
        y = conv3d_causal(x, self.w, self.s, time_pad=self.time_pad)
        y = y + self.b.reshape(-1, 1, 1, 1)
        if self.act == "tanh":
            y = np.tanh(y)
        elif self.act is True or self.act == "relu":
            y = np.maximum(y, 0.0)
        return y


def time_upsample(x, factor=2):
    """时间维最近邻上采样:每个潜变量帧重复 factor 次。

    这是纯复制,不引入新的时间依赖,所以不改感受野的「帧数」,
    只改感受野在输出帧单位下的跨度。
    """
    return np.repeat(x, factor, axis=1)


class ToyVideoVAE:
    """一个能跑的因果视频自编码器(权重是随机的,不追求重建质量)。

    编码器把时间压 4 倍、空间压 4 倍;解码器用最近邻上采样还原。
    它的唯一用途是让「谁依赖谁」这件事可以被实测。
    """

    def __init__(self, time_pad="causal"):
        self.e0 = CausalConv3d(1, 4, k=(3, 3, 3), s=(1, 1, 1), seed=1,
                               time_pad=time_pad)
        self.e1 = CausalConv3d(4, 8, k=(3, 3, 3), s=(2, 2, 2), seed=2,
                               time_pad=time_pad)
        self.e2 = CausalConv3d(8, 8, k=(3, 3, 3), s=(2, 2, 2), seed=3,
                               time_pad=time_pad)
        self.e3 = CausalConv3d(8, 8, k=(3, 3, 3), s=(1, 1, 1), seed=4,
                               time_pad=time_pad)
        self.d0 = CausalConv3d(8, 8, k=(3, 3, 3), s=(1, 1, 1), seed=5)
        self.d1 = CausalConv3d(8, 8, k=(3, 3, 3), s=(1, 1, 1), seed=6)
        self.d2 = CausalConv3d(8, 8, k=(3, 3, 3), s=(1, 1, 1), seed=7)
        self.d3 = CausalConv3d(8, 4, k=(3, 3, 3), s=(1, 1, 1), seed=8)
        self.d4 = CausalConv3d(4, 1, k=(3, 3, 3), s=(1, 1, 1), seed=9, act=None)

    def encode(self, x):
        x = self.e0(x)
        x = self.e1(x)
        x = self.e2(x)
        return self.e3(x)

    def decode(self, z):
        x = self.d0(z)
        x = self.d1(x)
        x = time_upsample(x, 2)
        x = self.d2(x)
        x = time_upsample(x, 2)
        x = self.d3(x)
        return self.d4(x)

    def __call__(self, x):
        return self.decode(self.encode(x))


def make_video(t=48, h=16, w=16, seed=7):
    """一段有运动的合成视频:一个高斯斑点匀速横移 + 缓慢明暗变化。"""
    rng = base_rng(seed)
    ys, xs = np.mgrid[0:h, 0:w]
    frames = []
    speed = 0.6
    cy = h / 2.0 + rng.normal(0, 0.3)
    for ti in range(t):
        cx = (w / 4.0 + speed * ti) % w
        blob = np.exp(-(((ys - cy) ** 2 + (xs - cx) ** 2) / 6.0))
        bg = 0.15 * np.sin(2 * np.pi * ti / 24.0) * np.ones_like(blob)
        frames.append(blob + bg + 0.02 * rng.normal(size=blob.shape))
    v = np.stack(frames)[None]           # [1, T, H, W]
    return v.astype(np.float64)


# ─────────────── 实验 1:形状公式 ───────────────

def shape_table():
    print("── 实验 1:因果卷积的时间帧数公式 ──")
    print("公式:T_out = floor((T_in - 1) / s_t) + 1    (因果,前补 k_t-1)")
    rows = []
    for t_in in (16, 17, 32, 48, 49, 121):
        row = [t_in]
        for s_t in (1, 2, 4):
            pred = out_frames(t_in, s_t=s_t)
            w = np.zeros((2, 1, 3, 3, 3))
            w[:] = 0.1
            x = np.zeros((1, t_in, 8, 8))
            got = conv3d_causal(x, w, stride=(s_t, 1, 1)).shape[1]
            row.append((s_t, pred, got))
        rows.append(row)
        print(f"  T_in={t_in:4d} -> " +
              "  ".join(f"s_t={s}: 预测 {p:3d} / 实测 {g:3d} {'OK' if p == g else 'FAIL'}"
                        for s, p, g in row[1:]))
    # CogVideoX 的真实数字:49 帧进,两级时间压缩 2
    t = 49
    t1 = out_frames(t, s_t=2)
    t2 = out_frames(t1, s_t=2)
    print(f"  两级时间压缩 2(CogVideoX 口径):49 帧 -> {t1} -> {t2} 潜变量帧")
    return rows, (t, t1, t2)


# ─────────────── 实验 2:因果性 ───────────────

def causality_check():
    """改未来的帧,看过去的潜变量有没有跟着变。变了就是漏了未来信息。"""
    print("\n── 实验 2:因果性检验(改未来,看过去)──")
    net = ToyVideoVAE()
    v = make_video(t=48)
    z_ref = net.encode(v)

    worst = 0.0
    for t_edit in range(0, 48, 4):
        v2 = v.copy()
        v2[:, t_edit:] += 3.0            # 从第 t_edit 帧起全部改动
        z2 = net.encode(v2)
        # 总时间步距 4,所以潜变量帧 j 只应该看到输入帧 <= 4j
        for j in range(z2.shape[1]):
            if 4 * j < t_edit:           # 这个潜变量帧不该看到第 t_edit 帧及之后
                diff = np.abs(z2[:, j] - z_ref[:, j]).max()
                worst = max(worst, diff)
    print(f"  应该完全不受影响的潜变量帧上,最大变化量 = {worst:.3e}")
    print(f"  判据:等于 0 则因果性成立(浮点意义上 < 1e-12 即通过)")
    return worst


# ─────────────── 实验 3:感受野 ───────────────

def encoder_rf_formula(kernels, strides):
    """RF = 1 + sum_l (k_l - 1) * prod_{m<l} s_m

    prod_{m<l} s_m 是「到第 l 层输入为止累积的时间步距」,
    所以越靠后的层,一个 kernel 覆盖的原始帧数越多。
    """
    rf = 1
    cum = 1
    for k, s in zip(kernels, strides):
        rf += (k - 1) * cum
        cum *= s
    return rf, cum


def encoder_rf_measured():
    """扰动法实测:逐帧加扰动,看哪些潜变量帧跟着动。"""
    net = ToyVideoVAE()
    v = make_video(t=48)
    z_ref = net.encode(v)
    nz = z_ref.shape[1]
    depends = np.zeros((48, nz), dtype=bool)
    for t_edit in range(48):
        v2 = v.copy()
        v2[:, t_edit:t_edit + 1] += 2.0
        z2 = net.encode(v2)
        depends[t_edit] = np.abs(z2 - z_ref).max(axis=(0, 2, 3)) > 1e-12
    return depends, z_ref


def receptive_field():
    kernels = [3, 3, 3, 3]
    strides = [1, 2, 2, 1]
    rf_pred, cum = encoder_rf_formula(kernels, strides)
    print("\n── 实验 3:编码器的时间感受野 ──")
    print(f"  公式:RF = 1 + sum (k_l - 1) * prod_(m<l) s_m")
    print(f"  本 toy 的层配置 k={kernels}, s={strides}")
    print(f"  公式预测 RF = {rf_pred} 帧,累积时间步距 = {cum}")

    depends, z_ref = encoder_rf_measured()
    nz = z_ref.shape[1]
    spans = []
    for j in range(nz):
        idx = np.where(depends[:, j])[0]
        if len(idx) == 0:
            spans.append((j, None, None, 0))
            continue
        spans.append((j, int(idx.min()), int(idx.max()), int(idx.max()) - int(idx.min()) + 1))
    print("  实测(逐帧扰动):潜变量帧 j -> 受影响的输入帧区间")
    for j, lo, hi, sp in spans[:6]:
        if lo is None:
            print(f"    j={j:2d}: 无任何输入帧影响它(权重恰好全被 ReLU 关掉)")
        else:
            print(f"    j={j:2d}: 输入帧 [{lo:2d}, {hi:2d}],跨度 {sp:2d} 帧")

    # 因果性上界:潜变量帧 j 能看到的最新输入帧
    print("  因果性上界检查:潜变量帧 j 能看到的最新输入帧应当 <= 步距 * j")
    bad = 0
    for j, lo, hi, sp in spans:
        if hi is not None and hi > cum * j:
            bad += 1
    print(f"    违反次数 = {bad}(0 表示因果性严格成立)")

    max_span = max(s for _, _, _, s in spans)
    print(f"  实测最大跨度 = {max_span} 帧,公式预测 = {rf_pred} 帧")
    return rf_pred, cum, spans, max_span


def dependency_matrix(t_in=48, time_pad="causal"):
    """扰动法得到「输入帧 t 是否影响潜变量帧 j」的布尔矩阵 [T_in, n_z]。

    这张矩阵就是因果性的可视化:因果卷积下它必须落在 j*s_t 这条对角线以下。
    """
    net = ToyVideoVAE(time_pad=time_pad)
    v = make_video(t=t_in)
    z_ref = net.encode(v)
    dep = np.zeros((t_in, z_ref.shape[1]), dtype=bool)
    for t_edit in range(t_in):
        v2 = v.copy()
        v2[:, t_edit:t_edit + 1] += 2.0
        dep[t_edit] = np.abs(net.encode(v2) - z_ref).max(axis=(0, 2, 3)) > 1e-12
    return dep


if __name__ == "__main__":
    shape_table()
    causality_check()
    receptive_field()

chunk_decode.py

# -*- coding: utf-8 -*-
"""分块编解码的接缝实验:到底要带多少帧上下文,逐块跑才能和整段跑一模一样。

运行: python chunk_decode.py

这是整篇文章最核心的一个实验。设置:
  - 96 帧输入 -> 编码器压成 24 个潜变量帧(时间步距 4)
  - 整段编解码得到参考输出
  - 然后把潜变量切成 3 块,每块前面额外喂 ctx 个潜变量帧当上下文,
    解码后把上下文那部分输出丢掉,只留本块
  - 比较「分块结果」与「整段结果」在每个位置的差

编码侧同样要分块(长视频不可能一次装进显存),所以下面两件事都测:
  1. 解码侧:需要几个潜变量帧上下文
  2. 编码侧:需要几个输入帧上下文
"""

import numpy as np

from causal_conv3d import ToyVideoVAE, make_video

T_FRAMES = 96          # 输入帧数
T_STRIDE = 4           # 编码器总时间步距:1 个潜变量帧对应 4 个输出帧
CHUNK = 8              # 每块 8 个潜变量帧 = 32 个输出帧


def build():
    net = ToyVideoVAE()
    v = make_video(t=T_FRAMES)
    z = net.encode(v)                     # [8, 24, 4, 4]
    ref = net.decode(z)                   # [1, 96, 16, 16]
    return net, v, z, ref


# ─────────────── 解码侧 ───────────────

def decode_chunked(net, z, chunk=CHUNK, ctx=0):
    """按 chunk 个潜变量帧一块解码,每块前面带 ctx 帧上下文。"""
    nz = z.shape[1]
    pieces, starts = [], []
    n_chunks = (nz + chunk - 1) // chunk
    for i in range(n_chunks):
        lo = i * chunk
        hi = min(lo + chunk, nz)
        lo_in = max(0, lo - ctx)
        out = net.decode(z[:, lo_in:hi])
        keep = (hi - lo) * T_STRIDE        # 上下文部分的输出要丢掉
        pieces.append(out[:, -keep:])
        starts.append(lo * T_STRIDE)
    return np.concatenate(pieces, axis=1), starts


def decode_sweep():
    net, v, z, ref = build()
    scale = float(ref.std())
    print("── 解码侧:上下文帧数 vs 接缝误差 ──")
    print(f"  {T_FRAMES} 帧输入 -> {z.shape[1]} 个潜变量帧(时间步距 {T_STRIDE})"
          f" -> {ref.shape[1]} 帧输出")
    print(f"  误差已按输出标准差归一化(std = {scale:.4f})")
    print("")
    print("  ctx | 块头误差 | 块尾误差 | 块内被污染帧数 | 边界跳变放大")
    print("  ----+---------+---------+----------------+------------")
    table = []
    for ctx in range(0, 9):
        got, starts = decode_chunked(net, z, CHUNK, ctx)
        diff = np.abs(got - ref).max(axis=(0, 2, 3)) / scale
        head_idx, tail_idx, contam = [], [], 0
        for s in starts[1:]:
            head_idx += list(range(s, s + T_STRIDE))
            tail_idx += list(range(s + T_STRIDE, s + CHUNK * T_STRIDE))
            contam += int((diff[s:s + CHUNK * T_STRIDE] > 1e-9).sum())
        head = float(diff[head_idx].max())
        tail = float(diff[tail_idx].max())
        jump = boundary_jump(ref, got, starts)
        table.append((ctx, head, tail, contam, jump))
        print(f"  {ctx:4d} | {head:.3e} | {tail:.3e} | {contam:14d} | {jump:.2f}x")
    return table


def boundary_jump(ref, got, starts):
    """同一位置处,分块结果的帧间跳变相对整段结果放大了多少倍。

    分母取「整段解码在同一帧的跳变」而不是全场平均——视频在某些帧本来就变化快,
    拿全场平均当分母会把正常内容误判成接缝。
    """
    ratios = []
    for s in starts[1:]:
        j_got = float(np.abs(got[:, s] - got[:, s - 1]).max())
        j_ref = float(np.abs(ref[:, s] - ref[:, s - 1]).max())
        if j_ref > 1e-12:
            ratios.append(j_got / j_ref)
    return max(ratios) if ratios else float("nan")


def decode_profile():
    """误差在块内是怎么衰减的:看第 2 块前 24 个输出帧。"""
    net, v, z, ref = build()
    scale = float(ref.std())
    print("")
    print("── 第 2 块内的误差衰减(输出帧 32 起,取前 12 帧,按 std 归一)──")
    out = {}
    for ctx in (0, 2, 4, 6):
        got, starts = decode_chunked(net, z, CHUNK, ctx)
        diff = np.abs(got - ref).max(axis=(0, 2, 3)) / scale
        seg = diff[32:32 + 24]
        out[ctx] = seg
        print(f"  ctx={ctx}: " + " ".join(f"{x:.1e}" for x in seg[:12]))
    return out


# ─────────────── 编码侧 ───────────────

def encode_chunked(net, v, chunk_in=32, ctx=0):
    """按 chunk_in 个输入帧一块编码,每块前面带 ctx 个输入帧上下文。"""
    t = v.shape[1]
    pieces = []
    n_chunks = (t + chunk_in - 1) // chunk_in
    for i in range(n_chunks):
        lo = i * chunk_in
        hi = min(lo + chunk_in, t)
        lo_in = max(0, lo - ctx)
        z_i = net.encode(v[:, lo_in:hi])
        keep = (hi - lo) // T_STRIDE       # 上下文换来的多余潜变量帧要丢掉
        pieces.append(z_i[:, -keep:] if keep > 0 else z_i[:, :0])
    return np.concatenate(pieces, axis=1)


def encode_sweep():
    net, v, z, ref = build()
    scale = float(z.std())
    print("")
    print("── 编码侧:上下文帧数 vs 潜变量误差 ──")
    print(f"  每块 32 个输入帧(= 8 个潜变量帧),误差按潜变量标准差归一"
          f"(std = {scale:.4f})")
    print("")
    print("  ctx | 块头潜变量误差 | 块尾潜变量误差 | 被污染潜变量帧数")
    print("  ----+----------------+----------------+--------------")
    table = []
    for ctx in (0, 4, 8, 12, 16, 20, 24):
        zh = encode_chunked(net, v, 32, ctx)
        if zh.shape[1] != z.shape[1]:
            print(f"  ctx={ctx}: 潜变量帧数对不上 {zh.shape[1]} vs {z.shape[1]},跳过")
            continue
        d = np.abs(zh - z).max(axis=(0, 2, 3)) / scale
        head = float(d[8:10].max())
        tail = float(d[10:16].max())
        contam = int((d[8:16] > 1e-9).sum())
        table.append((ctx, head, tail, contam))
        print(f"  {ctx:4d} | {head:.3e} | {tail:.3e} | {contam:14d}")
    return table


# ─────────────── 直接问:谁依赖谁 ───────────────

def dependency_table():
    """扰动法:直接问「输出帧 o 依赖哪些潜变量帧」,从而算出必需的上下文帧数。"""
    net, v, z, ref = build()
    nz = z.shape[1]
    z_ref = net.decode(z)
    dep = np.zeros((nz, z_ref.shape[1]), dtype=bool)
    for j in range(nz):
        z2 = z.copy()
        z2[:, j] += 1.0
        dep[j] = np.abs(net.decode(z2) - z_ref).max(axis=(0, 2, 3)) > 1e-12

    print("")
    print("── 直接问:块头输出帧依赖哪些潜变量帧 ──")
    print("  块起点(输出帧) | 依赖的最早潜变量帧 | 需要的上下文帧数")
    need = []
    n_chunks = (nz + CHUNK - 1) // CHUNK
    for i in range(1, n_chunks):
        j0 = i * CHUNK
        o0 = j0 * T_STRIDE
        idx = np.where(dep[:, o0])[0]
        if len(idx) == 0:
            continue
        need.append((o0, int(idx.min()), j0 - int(idx.min())))
        print(f"   {o0:3d} | {int(idx.min()):3d} | {j0 - int(idx.min())}")

    widths = []
    for o in range(z_ref.shape[1]):
        idx = np.where(dep[:, o])[0]
        widths.append(int(idx.max()) - int(idx.min()) + 1 if len(idx) else 0)
    rf_dec = max(widths)
    print(f"  解码器时间感受野(潜变量帧数,取所有输出帧的最大值)= {rf_dec}")
    print(f"  推论:所需上下文 = 感受野 - 1 = {rf_dec - 1} 个潜变量帧")
    return need, rf_dec


def context_decomposition():
    """逐层只缓存 (k_t - 1) 帧,折算回潜变量帧单位后累加起来等于什么?

    这是检验「工业实现里每层只缓存 k_t-1 帧够不够」的关键一笔:
    够不够取决于你把缓存折算回哪一级的单位。逐层缓存时,第 l 层的
    k_t-1 = 2 帧是该层输入分辨率下的 2 帧,折回潜变量帧要乘上该层
    相对潜变量的帧间距(上采样会把间距缩小)。
    """
    # 解码器 d0..d4 的输入相对潜变量的时间帧间距
    spacing = {"d0": 1.0, "d1": 1.0, "d2": 0.5, "d3": 0.25, "d4": 0.25}
    k_t = 3
    print("")
    print("── 逐层缓存 (k_t - 1) 帧,累加起来是多少 ──")
    print("")
    print("  层  | 输入相对潜变量的帧间距 | 缓存 (k_t-1) 帧折回潜变量帧")
    print("  ----+------------------------+----------------------------")
    total = 0.0
    for name, sp in spacing.items():
        c = (k_t - 1) * sp
        total += c
        print(f"  {name} | {sp:22.2f} | {c:12.2f}")
    print(f"  合计 | {'':22} | {total:12.2f}")
    print("")
    print(f"  实测需要的上下文(dependency_table)= 6 个潜变量帧")
    print(f"  逐层缓存累加 = {total:.2f} 个潜变量帧 -> "
          f"{'完全吻合' if abs(total - 6) < 1e-9 else '不吻合'}")
    print("")
    print("  结论:工业实现里「每层只缓存 k_t-1 帧」是**正确的**,因为逐层")
    print("  累加后恰好等于感受野 - 1。真正的陷阱是只在网络入口缓存一次——")
    print("  那样只有 (k_t-1) = 2 帧,差得远。")
    return total


def overhead(rf_dec=7, rf_enc=17):
    """带上下文要多算多少:额外算的量占本块的比例。"""
    print("")
    print("── 上下文的开销 ──")
    print("  解码块大小 | 需要上下文 | 额外算力占比")
    rows = []
    for chunk in (8, 12, 16, 24):
        ctx = rf_dec - 1
        frac = ctx / (chunk + ctx)
        rows.append(("decode", chunk, ctx, frac))
        print(f"   {chunk:3d} 潜变量帧 | {ctx:3d} 帧 | {frac * 100:.1f}%")
    print("  编码块大小 | 需要上下文 | 额外算力占比")
    for chunk in (32, 64, 128):
        ctx = rf_enc - 1
        frac = ctx / (chunk + ctx)
        rows.append(("encode", chunk, ctx, frac))
        print(f"   {chunk:3d} 输入帧 | {ctx:3d} 帧 | {frac * 100:.1f}%")
    return rows


def state():
    net, v, z, ref = build()
    return dict(z=z, ref=ref)


if __name__ == "__main__":
    decode_sweep()
    decode_profile()
    dependency_table()
    context_decomposition()
    encode_sweep()
    overhead()

temporal_budget.py

# -*- coding: utf-8 -*-
"""时间压缩比的预算:运动多快的时候,压 s 倍时间就开始糊。

运行: python temporal_budget.py

把「时间下采样」简化成最朴素的 s 帧平均 + 最近邻还原(CogVideoX 的下采样层
真的就是 avg_pool1d,见 05 节)。学习到的时间卷积会比平均聪明,但这个教学模型
只演示量级和规律,不是神经 VAE 的误差下界,而且它有一条可以验的解析预期:
    平均 s 帧 == 给运动物体糊上一条长度 (s-1) * v 像素的运动模糊
"""

import numpy as np

from causal_conv3d import base_rng

# 画布要够宽:最快的斑点(4 px/帧 x 32 帧 = 128 px)必须全程留在画面内,
# 否则斑点在后面几帧整个飘出去,加权宽度会因为分母趋于 0 而算成 nan。
T, H, W = 32, 32, 192
SIGMA = 3.0                      # 斑点的高斯半径(像素)
SPEEDS = [0.25, 0.5, 1.0, 2.0, 4.0]
FACTORS = [1, 2, 4, 8]


def moving_blob(speed, t=T, h=H, w=W, sigma=SIGMA, seed=11):
    """一个匀速横移的高斯斑点。不环绕,避免边界跳变污染 PSNR。"""
    rng = base_rng(seed)
    ys, xs = np.mgrid[0:h, 0:w]
    cy = h / 2.0
    cx0 = w * 0.25
    frames = []
    for ti in range(t):
        cx = cx0 + speed * ti
        frames.append(np.exp(-(((ys - cy) ** 2 + (xs - cx) ** 2) / (2 * sigma ** 2))))
    v = np.stack(frames)[None]                       # [1, T, H, W]
    return v + 0.001 * rng.normal(size=v.shape)      # 微量噪声,避免除零


def avg_pool_time(x, s):
    """每 s 帧平均成 1 帧(非重叠分组)。"""
    n = x.shape[1] // s
    return x[:, :n * s].reshape(x.shape[0], n, s, H, W).mean(axis=2)


def repeat_time(z, s, t_out=T):
    """最近邻还原:每个潜变量帧重复 s 次。"""
    return np.repeat(z, s, axis=1)[:, :t_out]


def avg_pool_space(x, s):
    """空间 sxs 块平均。"""
    n_h, n_w = H // s, W // s
    y = x[:, :, :n_h * s, :n_w * s]
    y = y.reshape(x.shape[0], T, n_h, s, n_w, s)
    return y.mean(axis=(3, 5))


def repeat_space(z, s):
    """空间最近邻还原。"""
    return np.repeat(np.repeat(z, s, axis=2), s, axis=3)


def psnr(a, b, peak=1.0):
    mse = float(np.mean((a - b) ** 2))
    return 99.0 if mse <= 1e-20 else 10.0 * np.log10(peak ** 2 / mse)


def width_along_x(x):
    """斑点沿运动方向的强度加权标准差,用来量「被糊成多宽」。

    两个坑都踩过:
    1. 只能在空间维度上归约(沿 y 求和、保留 x)。把所有轴一起归约会把 x
       也压掉,得到一个不随压缩变化的常数。
    2. 要先掐掉噪声底。加权方差里 (x - mean)^2 是杠杆,画面远端一个
       -0.026 的噪声像素能贡献 -9 的「方差」,把真实值直接打成负数。
    """
    xs = np.arange(W)[None, None, :]                 # [1, 1, W]
    prof = x.sum(axis=2)                             # [1, T, W],沿 y 归约
    prof = np.clip(prof, 0.0, None)
    thr = 0.005 * prof.max(axis=2, keepdims=True)    # 掐掉噪声底,保留斑点
    prof = np.where(prof < thr, 0.0, prof)
    wsum = prof.sum(axis=2, keepdims=True) + 1e-12   # [1, T, 1]
    mx = (prof * xs).sum(axis=2, keepdims=True) / wsum
    var = (prof * (xs - mx) ** 2).sum(axis=2, keepdims=True) / wsum
    return float(np.sqrt(var).mean())


def psnr_table():
    print("── 时间压缩 s 倍之后,运动速度 vs 重建 PSNR(dB)──")
    print("")
    print("  速度 v |" + "".join(f"   s={s:<2d}       " for s in FACTORS))
    print("  --------+" + "".join("---------------" for _ in FACTORS))
    table = {}
    for v in SPEEDS:
        x = moving_blob(v)
        row = []
        for s in FACTORS:
            row.append(psnr(x, repeat_time(avg_pool_time(x, s), s)))
        table[v] = row
        print(f"  {v:5.2f}  |" + "".join(f"  {p:8.2f} dB   " for p in row))
    print("")
    print("  读法:s=1 是原图(噪声极小,PSNR 到顶);同一列往下看,运动越快越糊。")
    print(f"        s=8 那一列从最慢到最快掉了 "
          f"{table[SPEEDS[0]][-1] - table[SPEEDS[-1]][-1]:.1f} dB。")
    return table


def blur_law():
    """验那条解析规律:糊掉的长度应该是 (s-1) * v。"""
    print("")
    print("── 验规律:平均 s 帧 ≈ 加一条长度 (s-1)*v 的运动模糊 ──")
    print("")
    print("  速度 v |  s  | 斑点宽度 原图 -> 压缩后 | 实测增幅 | 解析预期")
    print("  --------+-----+-------------------------+----------+----------")
    rows = []
    for v in (1.0, 2.0, 4.0):
        x = moving_blob(v)
        w0 = width_along_x(x)
        for s in (2, 4, 8):
            r = repeat_time(avg_pool_time(x, s), s)
            w1 = width_along_x(r)
            pred = np.sqrt(SIGMA ** 2 + v ** 2 * (s ** 2 - 1) / 12.0)
            rows.append((v, s, w0, w1, w1 / w0, pred / w0))
            print(f"  {v:5.2f}  |  {s:2d}  | {w0:7.3f} -> {w1:7.3f}         "
                  f"| {w1 / w0:7.3f}x | {pred / w0:7.3f}x")
    print("")
    print("  解析预期 = sqrt(sigma^2 + v^2*(s^2-1)/12) / 实测原宽度:")
    print("  平均 s 个等间隔平移量,其离散方差是 v^2*(s^2-1)/12。")
    return rows


def time_vs_space():
    """时间压 s 倍与空间两轴各压 s 倍(总 s^2 倍),并非等码率对照。"""
    print("")
    print("── 不同元素压缩率:时间 s 倍 vs 空间 s^2 倍(PSNR,dB)──")
    print("")
    print("  速度 v |" + "".join(f"  时间 {s}x / 空间 {s}x  " for s in (2, 4, 8)))
    print("  --------+" + "".join("-----------------------" for _ in (2, 4, 8)))
    rows = []
    for v in SPEEDS:
        x = moving_blob(v)
        cells = []
        for s in (2, 4, 8):
            p_t = psnr(x, repeat_time(avg_pool_time(x, s), s))
            p_s = psnr(x, repeat_space(avg_pool_space(x, s), s))
            cells.append((p_t, p_s))
        rows.append((v, cells))
        line = "  {:5.2f}  |".format(v)
        for p_t, p_s in cells:
            line += f"  {p_t:6.2f} / {p_s:6.2f}   "
        print(line)
    print("")
    print("  读法:空间那一列基本与速度无关(压空间就是把细节磨掉,一视同仁);")
    print("        时间那一列随运动变快而下降。两条线会交叉,交叉之后")
    print("        这只比较两种滤波失真,不能从非等码率交叉点推出实际压缩策略。")
    crossings = []
    for si, s in enumerate((2, 4, 8)):
        prev_sign = None
        for v, cells in rows:
            p_t, p_s = cells[si]
            sign = p_t > p_s
            if prev_sign is not None and sign != prev_sign:
                crossings.append((s, v))
            prev_sign = sign
    if crossings:
        print("")
        print("  交叉点(压时间开始不如压空间的速度):")
        for s, v in crossings:
            print(f"    s={s}x:约 {v} px/帧 附近")
    return rows


def state():
    return dict(psnr=psnr_table(), blur=blur_law(), ts=time_vs_space())


if __name__ == "__main__":
    psnr_table()
    blur_law()
    time_vs_space()

token_ledger.py

# -*- coding: utf-8 -*-
"""潜变量 token 账本:不同的时空压缩比,到底把 Transformer 的序列变成多长。

运行: python token_ledger.py

自注意力的代价是 O(N^2),所以真正决定「视频 DiT 能不能做全时空注意力」的
不是像素总量,而是**潜变量 token 数 N**。这个脚本只做算术,但它是选压缩比的依据。
"""

# 统一的输入:121 帧(5 秒 @ 24fps)、768 x 512、RGB
T_IN, H_IN, W_IN, C_IN = 121, 512, 768, 3

# 表一:真实模型在用的配置(通道数取各模型公开权重的值)
REAL = [
    ("图像 VAE 逐帧(SVD / AnimateDiff 口径)", 1, 8, 4),
    ("CogVideoX / Wan / HunyuanVideo(4x8x8)", 4, 8, 16),
    ("LTX-Video(8x32x32)", 8, 32, 128),
]

# 表二:固定空间 8x8、通道 16,只扫时间压缩比——把时间这一维的贡献单独拎出来
TIME_SWEEP = [(s_t, 8, 16) for s_t in (1, 2, 4, 8)]


def latent_frames(t_in, s_t):
    """因果卷积下的潜变量帧数:首帧单独占位,所以是 floor((T-1)/s)+1。"""
    return (t_in - 1) // s_t + 1


def _row(name, s_t, s_h, s_c):
    t_out = latent_frames(T_IN, s_t)
    h_out = H_IN // s_h
    w_out = W_IN // s_h
    n = t_out * h_out * w_out
    return dict(name=name, s_t=s_t, s_h=s_h, ch=s_c,
                shape=(t_out, h_out, w_out), n=n, vals=n * s_c,
                cover=s_t * s_h * s_h * C_IN)


def table_real():
    pix = T_IN * H_IN * W_IN * C_IN
    print(f"输入:{T_IN} 帧(5 秒 @ 24fps)、{W_IN} x {H_IN}、RGB,"
          f"像素值总数 = {pix:,}")
    rows = [_row(n, a, b, c) for n, a, b, c in REAL]
    base = rows[0]["n"]
    print("")
    print("── 表一:三个真实档位 ──")
    print("")
    print("  配置                                  | 潜变量形状 (T'xH'xW'xC) "
          "| token 数 N | N^2 相对代价 | 便宜倍数 | 每 token 覆盖像素")
    print("  --------------------------------------+-----------------------"
          "+------------+-------------+---------+----------------")
    for r in rows:
        rel = (r["n"] / base) ** 2
        r["rel"] = rel
        t, h, w = r["shape"]
        print(f"  {r['name']:<38} | {t:3d} x {h:3d} x {w:3d} x {r['ch']:3d} "
              f"| {r['n']:10,d} | {rel:11.3e} | {1 / rel:7.1f}x | {r['cover']:8,d}")
    return rows


def table_time_sweep():
    print("")
    print("── 表二:固定空间 8x8 / 通道 16,只动时间压缩比 ──")
    print("")
    print("  时间压缩 | 潜变量帧数 | token 数 N | 相对上一步便宜 | 相对 1x 累计")
    print("  ---------+------------+------------+----------------+-------------")
    rows = []
    base = None
    for s_t, s_h, s_c in TIME_SWEEP:
        r = _row(f"s_t={s_t}", s_t, s_h, s_c)
        if base is None:
            base = r["n"]
        r["rel"] = (r["n"] / base) ** 2
        rows.append(r)
    for i, r in enumerate(rows):
        step = (rows[i - 1]["n"] / r["n"]) ** 2 if i > 0 else 1.0
        print(f"   {r['s_t']:3d} x    | {r['shape'][0]:10d} | {r['n']:10,d} "
              f"| {step:13.1f}x | {1 / r['rel']:11.1f}x")
    print("")
    print("  规律:时间压缩每翻一倍,token 数减半,注意力代价降到 1/4。")
    return rows


def why_channels_grow():
    """压缩的是序列长度,不是信息量——通道数必须补回来。"""
    pix = T_IN * H_IN * W_IN * C_IN
    print("")
    print("── 为什么压缩比上去了,潜变量通道数也要跟着涨 ──")
    print("")
    print("  配置                                  | 潜变量数值总数 | 压缩比(像素值 / 潜变量数值)")
    print("  --------------------------------------+----------------+----------------------------")
    for name, s_t, s_h, s_c in REAL:
        r = _row(name, s_t, s_h, s_c)
        print(f"  {name:<38} | {r['vals']:14,d} | {pix / r['vals']:8.1f} : 1")
    print("")
    print("  数字要看懂:4x8x8 那一行和「逐帧图像 VAE」的压缩比几乎一样(46.8 vs 48.0),")
    print("  它省下的不是信息量,而是**序列长度**——190464 个 token 变成能做注意力的规模。")
    print("  LTX-Video 摘要里自述总压缩比 1:192,和上表最后一行的 181.5:1 是同一量级")
    print("  (差别来自它把 patchify 挪进 VAE 的口径)。")


def frame_count_rule():
    """帧数必须满足什么条件,才能整除不被裁掉。"""
    print("")
    print("── 帧数该怎么选:T = 1 + k * s_T ──")
    for s_t in (4, 8):
        print(f"  时间压缩 {s_t}x:合法帧数 1, {1 + s_t}, {1 + 2 * s_t}, ... 即 T = 1 + k*{s_t}")
        print(f"    例:121 帧 -> {latent_frames(121, s_t)} 个潜变量帧"
              f"({'整除,不丢帧' if (121 - 1) % s_t == 0 else '不整除,会向下取整'})")
        print(f"    例:100 帧 -> {latent_frames(100, s_t)} 个潜变量帧"
              f"({'整除,不丢帧' if (100 - 1) % s_t == 0 else '不整除,会向下取整'})")


def state():
    return dict(real=table_real(), sweep=table_time_sweep())


if __name__ == "__main__":
    table_real()
    table_time_sweep()
    why_channels_grow()
    frame_count_rule()
0

评论 (0)

取消
粤ICP备2021042327号