一句话总结

不需要迭代去噪,不需要 Transformer,只用 IMLE 训练目标 + 卷积网络 + 几个关键技巧,在 ImageNet 256 上达到 FID 2.56,推理速度接近 GAN。

背景:扩散模型的成功与代价

做生成模型,现在的默认答案是扩散模型(Diffusion)或流匹配(Flow Matching)。FID 漂亮,样本质量高,各种 benchmark 领先。但这背后有代价:

  • 推理需要几十到几百步迭代去噪,慢
  • 架构越来越重,现在标配 Transformer(DiT)
  • 训练目标包含数值积分,实现细节多
  • 有一个隐含假设:必须渐进式地从噪声转换到数据

ROMS-IMLE 问的是:这个”渐进式变换”真的必要吗?

答案出乎意料地是:不必要。

算法原理

先从各类生成模型的痛点说起

MLE:最大化 $\sum \log p_\theta(x_i)$。问题是 $p_\theta(x)$ 必须 tractable,逼着大家用 VAE 的变分下界或 normalizing flow。

GAN:判别器 + 生成器对抗训练。样本清晰,但容易 mode collapse,训练不稳定。

扩散模型:绕开 $p_\theta(x)$ 的 intractability,但多步采样代价高。

这三类都有各自的架构包袱。IMLE 提出一个不同的视角。

IMLE:最近邻匹配的直觉

IMLE(Implicit Maximum Likelihood Estimation)的核心思路极简单:

对数据集里每个真实样本 $x_i$,从噪声中采样一堆候选 ${z_1, …, z_k}$,找生成器 $G_\theta$ 中最接近 $x_i$ 的那个 $z^$,然后把 $G_\theta(z^)$ 往 $x_i$ 方向拉。

\[z_i^* = \arg\min_{z_j \in \mathcal{S}} \|G_\theta(z_j) - x_i\|_2\] \[\mathcal{L}_\text{IMLE} = \frac{1}{n}\sum_{i=1}^n \|G_\theta(z_i^*) - x_i\|_2^2\]

为什么这能避免两种经典失败模式?

  • 避免 mode averaging:每个 $x_i$ 只被一个 $z^*$ 负责,不存在把多个 mode 平均的情况(VAE 的模糊感来自这里)
  • 避免 mode collapse:如果生成器只覆盖少数几个 mode,大量 $x_i$ 找不到近邻,损失就会爆炸

ROMS:正交潜空间采样

论文名中的 ROMS 指的是一种更均匀的潜空间采样策略。普通高斯采样在高维空间中存在”球壳集中”现象——大量采样点分布在超球面附近,中心区域稀疏,导致最近邻匹配的候选池覆盖不均匀。

ROMS 用随机正交化解决这个问题:先采样高斯向量,再通过 QR 分解做正交化,确保候选集在方向上均匀分布。

实现

最小可运行版本

import torch
import torch.nn as nn
import torch.nn.functional as F

class Generator(nn.Module):
    """轻量级卷积生成器,4x4 -> 64x64"""
    def __init__(self, latent_dim=512):
        super().__init__()
        self.fc = nn.Linear(latent_dim, 512 * 4 * 4)
        self.blocks = nn.ModuleList([
            nn.Sequential(
                nn.Upsample(scale_factor=2),
                nn.Conv2d(512 >> i, 512 >> (i + 1), 3, padding=1),
                nn.BatchNorm2d(512 >> (i + 1)),
                nn.GELU(),
            )
            for i in range(4)  # 通道: 512->256->128->64->32
        ])
        self.to_rgb = nn.Conv2d(32, 3, 1)

    def forward(self, z):
        x = self.fc(z).view(-1, 512, 4, 4)
        for block in self.blocks:
            x = block(x)
        return torch.tanh(self.to_rgb(x))


def roms_sample(n, latent_dim, device):
    """正交化采样,让候选 z 覆盖更均匀"""
    z = torch.randn(n, latent_dim, device=device)
    if n <= latent_dim:
        q, _ = torch.linalg.qr(z.T)
        z = q.T * (latent_dim ** 0.5)  # 还原尺度
    return z


def imle_step(generator, real_batch, feat_extractor, num_candidates=128):
    """IMLE 核心:在特征空间做最近邻匹配"""
    device = real_batch.device
    z_pool = roms_sample(num_candidates, 512, device)

    with torch.no_grad():
        gen_pool = generator(z_pool)
        # 在特征空间匹配比像素空间效果好得多
        real_feat = feat_extractor(real_batch).flatten(1)
        gen_feat = feat_extractor(gen_pool).flatten(1)

    # 每个真实样本找最近邻候选
    dists = torch.cdist(real_feat, gen_feat)     # (B, num_candidates)
    best_idx = dists.argmin(dim=1)               # (B,)
    matched_z = z_pool[best_idx]                 # (B, latent_dim)

    gen_matched = generator(matched_z)
    return F.mse_loss(gen_matched, real_batch)

完整训练循环

def train(generator, dataloader, feat_extractor, epochs=100, device='cuda'):
    optimizer = torch.optim.Adam(generator.parameters(), lr=1e-4)
    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
        optimizer, T_max=epochs
    )

    for epoch in range(epochs):
        for real_imgs, _ in dataloader:
            real_imgs = real_imgs.to(device)
            optimizer.zero_grad()

            loss = imle_step(generator, real_imgs, feat_extractor)
            loss.backward()
            # IMLE 偶尔梯度爆炸,裁剪是必要的
            torch.nn.utils.clip_grad_norm_(generator.parameters(), 1.0)
            optimizer.step()

        scheduler.step()
        if epoch % 10 == 0:
            print(f"Epoch {epoch:3d} | Loss: {loss.item():.4f}")

# ... (FID 评估代码省略)

关键 Trick(不写就跑不起来)

1. 候选数量要足够

候选太少,最近邻匹配质量差,生成器学不到东西:

num_candidates = 16   # 太少:匹配几乎随机
num_candidates = 128  # 可接受的下限
num_candidates = 512  # 更好,但显存压力大

2. 特征空间匹配,不要像素空间

像素 MSE 对人眼感知不敏感,匹配到的 $z^*$ 质量差。用预训练 VGG 中间层特征做距离计算,效果大幅提升。这是论文消融里影响最大的单个设计(去掉后 FID +2.1)。

3. 定期重采样,不要每步都重新匹配

# 每 N 步重新做一次全量最近邻匹配
# 太频繁:计算量爆炸;太少:匹配过时,梯度方向错误
resample_every = 20  # 根据数据集大小调整

4. 正交采样(ROMS)不是可选项

在高维潜空间(512 维以上),普通高斯采样候选点”集中在超球面”,导致覆盖不均匀。去掉 ROMS 换回普通采样,FID 大约变差 0.8。

实验结论

FID 对比(ImageNet 256)

方法 FID ↓ 推理步数 架构
ADM 3.94 ~250 步 CNN
LDM 3.60 ~250 步 Transformer
DiT-XL/2 2.27 250 步 Transformer
ROMS-IMLE 2.56 1 步 CNN

单步生成接近 250 步扩散模型,这个对比是有说服力的。

消融:哪些东西真正重要

去掉的组件 FID 变化 结论
特征空间匹配 → 像素匹配 +2.1 关键
ROMS → 普通高斯采样 +0.8 重要
CNN → Transformer 无显著提升 可选
单步 → 多步迭代 无显著提升 不必要

最后两行是论文最有价值的发现。

调试指南

常见问题

1. 图像模糊(mode averaging 残余)

最大可能是在像素空间做匹配。换成感知特征空间:

# 用 VGG perceptual feature 替代原始像素
from torchvision.models import vgg16
vgg = vgg16(pretrained=True).features[:16].eval().cuda()

2. 多样性差(生成图像大量重复)

候选数量不够,增大 num_candidates,同时检查 ROMS 采样是否正确实现。

3. 训练震荡,loss 突然飙升

通常是学习率过高或梯度裁剪阈值太宽松:

# 保守设置起步
lr = 3e-5
clip_norm = 0.5

超参数敏感度

参数 推荐初值 敏感度 优先级
num_candidates 128 第一个调
lr 1e-4 过大直接崩
resample_every 20 影响训练速度
latent_dim 512 256-1024 差别不大

什么时候用 / 不用?

适用场景 不适用场景
推理延迟敏感(实时、移动端) 需要文本条件控制
计算资源有限,跑不起扩散步骤 需要 inpainting/editing 扩展
研究”简单方法的上限” 已有成熟的扩散 pipeline
探索 IMLE 及其变体 对 FID 有极致追求

我的观点

ROMS-IMLE 最有价值的不是 FID 2.56 这个数字,而是它提供的一个认识论层面的反驳:

“扩散模型好是因为迭代去噪” — 这个因果关系可能是错的。

更可能的解释是:扩散模型好是因为规模、数据和精心设计的网络结构,迭代只是其中一种实现手段,而不是根本原因。

实用建议

  • 做学术研究、想探索简单生成框架:IMLE 值得深入
  • 做应用系统、需要文本控制:扩散生态(ControlNet、LoRA)更务实
  • 做 single-step distillation(把扩散蒸馏成单步):IMLE 的思想可以借鉴

IMLE 这个方向早在 2018 年就有人提出(Ke Li 等人),但一直没有进入主流视野。这篇论文是个好的提醒:旧技术有时候只是需要被重新发现。

论文链接:https://arxiv.org/abs/2607.19332v1