文章目录

流模型Normalizing Flow的雅可比行列式与可逆变换设计

发布于 2026-07-20 16:41:43 · 浏览 24 次 · 评论 0 条

流模型 Normalizing Flow 的核心:可逆变换与雅可比行列式

流模型的核心思想,是把一个简单的概率分布(比如标准正态分布)通过一系列可逆变换,变成一个复杂的、我们想要的数据分布。这个过程就像“流动”一样。要实现这种变换,设计可逆变换并计算其雅可比行列式是关键。下面直接讲解如何动手设计和计算。


1. 设计可逆变换:从简单到复杂

定义:一个可逆变换 f 必须满足:给定输入 x,输出 z = f(x),同时存在一个逆变换 f^{-1} 使得 x = f^{-1}(z)。在流模型中,我们通常让变换方向为:从数据空间 x 映射到潜变量空间 z(即“标准化”方向),而生成则使用逆变换。

选择变换家族:为了确保可逆性,我们使用结构化的参数化函数。常用的设计包括:

  • 仿射变换z = a * x + b,其中 a 不为零。逆变换为 x = (z - b) / a。简单但能力有限。
  • 耦合层 (Coupling Layer):将输入 x 分成两部分 x_ax_b。让 x_a 不变直接通过,然后利用一个神经网络根据 x_a 输出缩放和平移参数,作用于 x_bz_b = s(x_a) * x_b + t(x_a)z_a = x_a。逆变换时,同样用 x_a(你已有 z_a)计算参数,再反向操作。
  • 自回归变换 (Autoregressive):输出按顺序依赖前面的输入,例如 z_i = s(x_{1:i-1}) * x_i + t(x_{1:i-1})。可逆性通过类似的顺序实现。
  • Householder 反射:使用反射矩阵保持体积(雅可比行列式绝对值为1),比较简单。

实施步骤

  1. 选择变换类型:从数据分布出发,评估数据维度(低维用仿射,高维用耦合层)。确定变换结构。
  2. 参数化:缩放因子 s 和平移 t 通常用神经网络(如 MLP)输出,但必须确保缩放 s 始终为正(例如通过 exp 激活函数输出)。
  3. 实现正向和逆向编写函数 forward 计算 zlog_det_J编写函数 inverse 计算 xz

2. 计算雅可比行列式:为什么它必不可少

在流模型的训练中,我们需要最大化数据 x 的对数似然。根据概率密度变换公式:

log p(x) = log p(z) + log |det(J)|,其中 J = dz/dx(雅可比矩阵),det(J) 是它的行列式。加号是因为我们正变换从 xz,概率密度需要乘以雅可比行列式的绝对值补回体积变化。

关键洞察:如果我们精心设计变换,使雅可比矩阵是三角矩阵,那么行列式就简化为对角线元素的乘积,计算代价降至 O(D)D 是数据维度)。耦合层和自回归层就是为此设计的——它们使得雅可比矩阵为分块三角或严格三角。

分步计算

  • 耦合层:对分割后的两部分,雅可比矩阵是分块三角的。其中,z_a = x_a 对应单位矩阵块;z_b 部分对应与 x_b 相关的对角矩阵(缩放 s 作用于元素)。因此行列式等于 prod(s(x_a))。对数行列式:sum(log(s(x_a)))
  • 自回归变换:雅可比矩阵是下三角(或上三角,取决于排序)。对角线元素是每个输入的缩放因子 s_i。对数行列式:sum(log(s_i))
  • Householder 反射:雅可比矩阵行列式绝对值为1,所以贡献为零。

实施细节:在计算对数似然的公式中,你需要 记录 正向变换时每个缩放因子 s 的对数和。示例:耦合层中,对 x_b 每个元素 log s_j 求和。

以下是耦合层的一个简单计算示例(伪代码):

def forward_coupling(x):
    x_a, x_b = split(x)
    # 用网络计算缩放和平移
    s, t = network(x_a)  # s 和 t 形状与 x_b 相同
    s = exp(s)  # 确保正数
    z_b = s * x_b + t
    z_a = x_a
    log_det = sum(log(s))
    return concat(z_a, z_b), log_det

def inverse_coupling(z):
    z_a, z_b = split(z)
    s, t = network(z_a)  # 同样网络(参数共享)
    s = exp(s)
    x_b = (z_b - t) / s
    x_a = z_a
    return concat(x_a, x_b)

3. 多步级联:深层的流模型

单个变换表达能力有限。堆叠多个变换形成深层流,总雅可比行列式为每一步的对数行列式之和:log |det(J_total)| = sum_i log |det(J_i)|。然后通过正向依次应用每个变换,累积对数行列式。

设计策略

  • 交替分割:在耦合层之间,需要打乱特征顺序,否则某些维度永远不会相互作用。使用随机置换或固定但非平凡的排列(如逆序)。
  • 加入非线性:除了仿射耦合,还可以引入通过 1x1 卷积(对图像)或线性变换(对向量)整型混合特征,这些变换自身需要计算行列式(例如 1x1 卷积的行列式很简单)。
  • 批标准化:某些流模型引入可逆批标准化(ActNorm)来稳定训练,其雅可比行列式计算类似仿射变换。

示例:一个简单的多步流:

  1. 仿射耦合层(输入分割,一半不变,一半变换)。
  2. 置换(打乱顺序)。
  3. 仿射耦合层(在新的分割下)。
  4. 置换
  5. 重复若干次。

计算总对数行列式:在每一步,记录累加 对数行列式。最后,总似然为 log p(z_last) + sum_log_det,其中 p(z_last) 通常是标准正态分布的对数密度。


4. 避免常见陷阱

  • 可逆性中断:确保缩放因子 s 永远不为零(使用 expsigmoid 加微小偏移)。
  • 数值稳定性使用 log |det| 而不是行列式本身,防止数值溢出或下溢。在计算 log sum(s) 时,保证 s 为正。
  • 越界值:某些变换(如仿射)可能导致潜变量范围无限,但标准正态分布对极端值惩罚大,模型会自主调整。无需额外约束。
  • 数据预处理:对于图像等连续数据,通常需要 应用 反量化(dequantization)来避免离散分布的零密度问题。

5. 验证实现:小规模测试

测试一个单步耦合层,使用 2D 数据(方便视觉检查)。生成随机数据点,计算正向和逆向,确认 x == inverse(forward(x)) 成立(数值误差可容忍)。计算对数行列式并与数值差分估计比较。确保梯度反向传播正常。

代码示例(PyTorch 风格)

import torch
import torch.nn as nn

class AffineCoupling(nn.Module):
    def __init__(self, in_dim, hidden_dim):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(in_dim // 2, hidden_dim),
            nn.ReLU(),
            nn.Linear(hidden_dim, in_dim // 2 * 2)  # 输出缩放和平移
        )

    def forward(self, x):
        x_a, x_b = x.chunk(2, dim=-1)
        params = self.net(x_a)
        s, t = params.chunk(2, dim=-1)
        s = torch.exp(s)  # 正缩放
        z_b = s * x_b + t
        z_a = x_a
        log_det = s.log().sum(-1)  # 每个样本的对数行列式
        return torch.cat([z_a, z_b], dim=-1), log_det

    def inverse(self, z):
        z_a, z_b = z.chunk(2, dim=-1)
        params = self.net(z_a)
        s, t = params.chunk(2, dim=-1)
        s = torch.exp(s)
        x_b = (z_b - t) / s
        x_a = z_a
        return torch.cat([x_a, x_b], dim=-1)

测试代码:

model = AffineCoupling(in_dim=4, hidden_dim=32)
x = torch.randn(10, 4)
z, log_det = model.forward(x)
x_recon = model.inverse(z)
print(torch.allclose(x, x_recon, atol=1e-6))  # 应输出 True

通过以上步骤,你就能自行设计并实现一个可逆变换,并正确计算其雅可比行列式,为构建完整的归一化流模型打下基础。

评论 (0)

暂无评论,快来抢沙发吧!

扫一扫,手机查看

扫描上方二维码,在手机上查看本文