案例 17:Fashion-MNIST 小型去噪扩散概率模型——从前向加噪到逐步生成#
学习目标与记号#
扩散模型把一次困难的“从噪声生成图像”拆成许多个逐步去噪的小问题。本案例不使用预训练权重,从 Fashion-MNIST 原始 IDX 文件出发,完整展示去噪扩散概率模型(denoising diffusion probabilistic model,DDPM)的训练与采样过程。本案例的学习目标包括:
高斯前向扩散与任意时间步的解析采样;
线性噪声日程及 \(\alpha_t,\bar\alpha_t\) 的含义与维度;
时间条件卷积网络预测加入的噪声;
简化噪声预测损失与 DDPM 逐步反向采样;
零噪声预测器基线、按时间步误差和噪声日程消融;
区分“快速可运行演示”和“足以评价生成质量的完整训练”。
默认的快速模式只用于检查代码能否按照正文公式完整运行。只训练一个轮次(epoch)的小模型通常只能生成粗糙轮廓,不能把展示样本当作高质量生成结果。
数据来源、许可与运行预算#
数据集:Fashion-MNIST,说明页 zalandoresearch/fashion-mnist
匿名 HTTPS 直链:
https://storage.googleapis.com/tensorflow/tf-keras-datasets/train-images-idx3-ubyte.gz
https://storage.googleapis.com/tensorflow/tf-keras-datasets/train-labels-idx1-ubyte.gz
https://storage.googleapis.com/tensorflow/tf-keras-datasets/t10k-images-idx3-ubyte.gz
https://storage.googleapis.com/tensorflow/tf-keras-datasets/t10k-labels-idx1-ubyte.gz
许可:MIT。
数据缓存在 AI_COURSE_DATA_DIR 指定的目录中;未设置该变量时使用用户缓存目录。快速模式使用 4096 个训练样本、512 个验证样本、40 个扩散步和 1 个训练轮次,可用中央处理器(central processing unit,CPU)运行。将代码变量 FAST_MODE 设为 False 后,程序使用完整训练集、200 个扩散步和 20 个训练轮次,建议使用图形处理器(graphics processing unit,GPU)。它仍然只是小型教学模型,不是对原论文实验的复现。
from pathlib import Path
import copy
import gzip
import hashlib
import os
import random
import struct
import urllib.request
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
import torch
from torch.utils.data import DataLoader, TensorDataset
SEED = 42
FAST_MODE = os.getenv("AI_COURSE_FAST_MODE", os.getenv("FAST_MODE", "1")) != "0"
random.seed(SEED)
np.random.seed(SEED)
torch.manual_seed(SEED)
cache_root = Path(
os.environ.get(
"AI_COURSE_DATA_DIR",
Path(os.environ.get("XDG_CACHE_HOME", Path.home() / ".cache")) / "ai-course-cases",
)
).expanduser()
cache_root.mkdir(parents=True, exist_ok=True)
device = torch.device(
"cuda" if torch.cuda.is_available() else
"mps" if getattr(torch.backends, "mps", None) and torch.backends.mps.is_available() else
"cpu"
)
if FAST_MODE:
device = torch.device("cpu") # 发布验证保持跨平台确定且不占用 GPU
print({"快速模式": FAST_MODE, "计算设备": str(device), "缓存目录": str(cache_root)})
{'快速模式': True, '计算设备': 'cpu', '缓存目录': '/private/tmp/ai-course-case-data'}
第 1 步:下载并解析 IDX 原始格式#
我们不依赖 torchvision 的隐藏下载逻辑,而是显式解析 IDX 文件头。图像文件头给出样本数、行数和列数;标签文件给出样本数。读取后立即断言二者数量一致、图像为 \(28\times28\)。
下载先写 .part 临时文件,成功后原子替换;打印 SHA-256 便于记录数据快照。课程源码目录中不会产生 CSV、模型权重或缓存。
URLS = {
"train_images": "https://storage.googleapis.com/tensorflow/tf-keras-datasets/train-images-idx3-ubyte.gz",
"train_labels": "https://storage.googleapis.com/tensorflow/tf-keras-datasets/train-labels-idx1-ubyte.gz",
"test_images": "https://storage.googleapis.com/tensorflow/tf-keras-datasets/t10k-images-idx3-ubyte.gz",
"test_labels": "https://storage.googleapis.com/tensorflow/tf-keras-datasets/t10k-labels-idx1-ubyte.gz",
}
def download_cached(url, target):
if not target.exists():
temporary = target.with_suffix(target.suffix + ".part")
request = urllib.request.Request(url, headers={"User-Agent": "ai-course-case/1.0"})
with urllib.request.urlopen(request, timeout=120) as response, temporary.open("wb") as out:
while chunk := response.read(1024 * 1024):
out.write(chunk)
temporary.replace(target)
print("文件:", target.name, ";SHA-256:", hashlib.sha256(target.read_bytes()).hexdigest())
return target
paths = {name: download_cached(url, cache_root / Path(url).name) for name, url in URLS.items()}
def read_idx_images(path):
with gzip.open(path, "rb") as handle:
magic, count, rows, cols = struct.unpack(">IIII", handle.read(16))
if magic != 2051:
raise ValueError(f"图像IDX魔数错误:{magic}")
data = np.frombuffer(handle.read(), dtype=np.uint8)
return data.reshape(count, rows, cols).copy()
def read_idx_labels(path):
with gzip.open(path, "rb") as handle:
magic, count = struct.unpack(">II", handle.read(8))
if magic != 2049:
raise ValueError(f"标签IDX魔数错误:{magic}")
data = np.frombuffer(handle.read(), dtype=np.uint8)
return data.reshape(count).copy()
train_images_all = read_idx_images(paths["train_images"])
train_labels_all = read_idx_labels(paths["train_labels"])
test_images_all = read_idx_images(paths["test_images"])
test_labels_all = read_idx_labels(paths["test_labels"])
assert train_images_all.shape == (60000, 28, 28)
assert test_images_all.shape == (10000, 28, 28)
assert len(train_images_all) == len(train_labels_all)
print("训练图像、训练标签和测试图像的维度:", train_images_all.shape, train_labels_all.shape, test_images_all.shape)
文件: train-images-idx3-ubyte.gz ;SHA-256: 3aede38d61863908ad78613f6a32ed271626dd12800ba2636569512369268a84
文件: train-labels-idx1-ubyte.gz ;SHA-256: a04f17134ac03560a47e3764e11b92fc97de4d1bfaf8ba1a3aa29af54cc90845
文件: t10k-images-idx3-ubyte.gz ;SHA-256: 346e55b948d973a97e58d2351dde16a484bd415d4595297633bb08f03db6a073
文件: t10k-labels-idx1-ubyte.gz ;SHA-256: 67da17c76eaffca5446c3361aaab5c3cd6d1c2608764d35dfb1850b086bf8dd5
训练图像、训练标签和测试图像的维度: (60000, 28, 28) (60000,) (10000, 28, 28)
第 2 步:在建模前看见数据#
像素原始范围是 0–255。网络与高斯过程使用 \([-1,1]\),使“0 附近”对应中灰,采样结果最终再映射回 \([0,1]\) 显示。标签只用于检查类别覆盖和可视化,不进入这个无条件 DDPM。
观察服装类别和灰度分布能帮助发现读取时的字节顺序、数组维度或归一化错误;如果图像旋转、全黑或取值范围异常,后续损失即使下降也没有意义。
class_names = ["T恤", "裤子", "套衫", "连衣裙", "外套", "凉鞋", "衬衫", "运动鞋", "包", "短靴"]
fig, axes = plt.subplots(2, 5, figsize=(10, 4))
for label, axis in enumerate(axes.flat):
idx = int(np.flatnonzero(train_labels_all == label)[0])
axis.imshow(train_images_all[idx], cmap="gray", vmin=0, vmax=255)
axis.set_title(class_names[label])
axis.axis("off")
plt.tight_layout()
counts = np.bincount(train_labels_all, minlength=10)
display(pd.DataFrame({"类别": class_names, "训练样本数": counts}))
print({"像素最小值": int(train_images_all.min()), "像素最大值": int(train_images_all.max()),
"像素均值": float(train_images_all.mean())})
| 类别 | 训练样本数 | |
|---|---|---|
| 0 | T恤 | 6000 |
| 1 | 裤子 | 6000 |
| 2 | 套衫 | 6000 |
| 3 | 连衣裙 | 6000 |
| 4 | 外套 | 6000 |
| 5 | 凉鞋 | 6000 |
| 6 | 衬衫 | 6000 |
| 7 | 运动鞋 | 6000 |
| 8 | 包 | 6000 |
| 9 | 短靴 | 6000 |
{'像素最小值': 0, '像素最大值': 255, '像素均值': 72.94035223214286}
第 3 步:先划分,再定义训练预算#
Fashion-MNIST 官方测试集只用于最后评价噪声预测误差。我们从官方训练集固定抽出验证集,用于选择训练轮次和比较日程;快速模式再从剩余训练索引抽样。标签不参与抽样,所以不会刻意平衡生成分布。
张量变为 \((B,1,28,28)\),像素线性映射到 \([-1,1]\)。训练数据加载器(DataLoader)使用固定随机数生成器;这样批次顺序可复现。
rng = np.random.default_rng(SEED)
all_indices = rng.permutation(len(train_images_all))
valid_count = 512 if FAST_MODE else 5000
train_count = 4096 if FAST_MODE else len(train_images_all) - valid_count
valid_idx = all_indices[:valid_count]
train_idx = all_indices[valid_count:valid_count + train_count]
test_idx = np.arange(512 if FAST_MODE else len(test_images_all))
def image_tensor(array, indices):
x = torch.from_numpy(array[indices]).float().unsqueeze(1) / 127.5 - 1.0
return x
x_train = image_tensor(train_images_all, train_idx)
x_valid = image_tensor(train_images_all, valid_idx)
x_test = image_tensor(test_images_all, test_idx)
generator = torch.Generator().manual_seed(SEED)
train_loader = DataLoader(
TensorDataset(x_train), batch_size=64 if FAST_MODE else 128,
shuffle=True, generator=generator, num_workers=0,
)
valid_loader = DataLoader(TensorDataset(x_valid), batch_size=128, shuffle=False)
test_loader = DataLoader(TensorDataset(x_test), batch_size=128, shuffle=False)
print({"训练集": tuple(x_train.shape), "验证集": tuple(x_valid.shape), "测试集": tuple(x_test.shape)})
assert x_train.min() >= -1 and x_train.max() <= 1
{'训练集': (4096, 1, 28, 28), '验证集': (512, 1, 28, 28), '测试集': (512, 1, 28, 28)}
第 4 步:噪声日程与任意时间步的解析采样#
离散前向过程为
利用高斯变量的可加性,不必从 1 循环到 \(t\),可直接采样
extract 会按照每个样本的时间步 \(t\),从长度为 \(T\) 的系数序列中取出对应系数,并将其变为 \((B,1,1,1)\)。程序再利用广播机制,对图像的通道、高度和宽度方向使用同一个时间系数。
T = 40 if FAST_MODE else 200
def linear_schedule(steps):
scale = 1000 / steps
return torch.linspace(scale * 1e-4, min(scale * 0.02, 0.30), steps).clamp(max=0.30)
betas = linear_schedule(T).to(device)
alphas = 1.0 - betas
alpha_bar = torch.cumprod(alphas, dim=0)
alpha_bar_previous = torch.cat([torch.ones(1, device=device), alpha_bar[:-1]])
posterior_variance = betas * (1.0 - alpha_bar_previous) / (1.0 - alpha_bar)
def extract(coefficients, t, x):
return coefficients.gather(0, t).reshape(-1, 1, 1, 1).to(x.dtype)
def q_sample(x0, t, noise=None):
noise = torch.randn_like(x0) if noise is None else noise
return (
extract(alpha_bar.sqrt(), t, x0) * x0
+ extract((1.0 - alpha_bar).sqrt(), t, x0) * noise
), noise
assert torch.all((betas > 0) & (betas < 1))
assert torch.all(alpha_bar[1:] < alpha_bar[:-1])
print({"扩散总步数": T, "首步噪声方差": float(betas[0]), "末步噪声方差": float(betas[-1]),
"末步累计信号系数": float(alpha_bar[-1])})
{'扩散总步数': 40, '首步噪声方差': 0.0024999999441206455, '末步噪声方差': 0.30000001192092896, '末步累计信号系数': 0.0011396778281778097}
第 5 步:公式到图像——信噪比怎样随 t 变化#
对同一张 \(\boldsymbol{x}_0\) 和同一份噪声 \(\boldsymbol{\epsilon}\),只改变时间步 \(t\),就能直观看到 \(\sqrt{\bar\alpha_t}\) 控制保留多少原始信号,\(\sqrt{1-\bar\alpha_t}\) 控制加入多少噪声。相应的信噪比为
早期时间步的输入接近原图,噪声难以从图像内容中分离;末期若 \(\bar\alpha_T\) 仍太大,起点就不是近似高斯。
x0_demo = x_train[:1].to(device)
shared_noise = torch.randn(x0_demo.shape, generator=torch.Generator().manual_seed(SEED)).to(device)
chosen = torch.tensor([0, T // 4, T // 2, 3 * T // 4, T - 1], device=device)
fig, axes = plt.subplots(1, len(chosen), figsize=(11, 2.3))
for axis, time_index in zip(axes, chosen):
noisy, _ = q_sample(x0_demo, time_index.reshape(1), shared_noise)
axis.imshow(((noisy[0, 0].cpu() + 1) / 2).clamp(0, 1), cmap="gray")
axis.set_title(f"时间步={int(time_index)}")
axis.axis("off")
plt.tight_layout()
snr = alpha_bar / (1 - alpha_bar)
plt.figure(figsize=(6, 3))
plt.semilogy(snr.cpu())
plt.xlabel("时间步"); plt.ylabel("信噪比(对数轴)"); plt.title("线性日程的信噪比");
第 6 步:零预测器基线#
训练时随机采样时间步 \(t\) 和噪声 \(\boldsymbol{\epsilon}\),让网络 \(\boldsymbol{\epsilon}_\theta(\boldsymbol{x}_t,t)\) 最小化
若始终预测零向量,期望均方误差(mean squared error,MSE)约为 \(\mathbb{E}[\epsilon^2]=1\)。训练后的模型应优于这个可解释的基线;否则可能是模型尚未学到有效规律、时间条件的实现有误,或训练次数太少。这个简化损失与证据下界(evidence lower bound,ELBO)有联系,但去掉了随时间变化的权重,不能把其数值直接解释为对数似然。
@torch.no_grad()
def zero_baseline(loader, batches=4):
values = []
local_generator = torch.Generator().manual_seed(SEED + 1)
for batch_no, (x0,) in enumerate(loader):
if batch_no >= batches:
break
noise = torch.randn(x0.shape, generator=local_generator)
values.append(float((noise ** 2).mean()))
return float(np.mean(values))
baseline_mse = zero_baseline(valid_loader)
print("零噪声预测器验证均方误差:", baseline_mse)
零噪声预测器验证均方误差: 1.0012334883213043
第 7 步:时间条件卷积去噪器#
同一个带噪图像数值可能对应不同噪声等级,因此模型必须知道 t。正弦嵌入把 t 映射为一组不同频率的正弦与余弦值,再经多层感知机(MLP)投影后广播成 \((B,C,1,1)\) 加到特征图。
教学网络保持 \(28\times28\) 分辨率,输入和输出的维度都为 \((B,1,28,28)\)。论文中的 DDPM 常使用多尺度 U-Net、残差块和注意力;这里使用小型卷积网络,是为了让关键的反向扩散过程能够在 CPU 上运行。
class SinusoidalTimeEmbedding(torch.nn.Module):
def __init__(self, dimension):
super().__init__()
self.dimension = dimension
def forward(self, t):
half = self.dimension // 2
frequency = torch.exp(
-np.log(10000) * torch.arange(half, device=t.device) / max(half - 1, 1)
)
angles = t.float()[:, None] * frequency[None]
embedding = torch.cat([angles.sin(), angles.cos()], dim=1)
if embedding.shape[1] < self.dimension:
embedding = torch.nn.functional.pad(embedding, (0, 1))
return embedding
class TinyDenoiser(torch.nn.Module):
def __init__(self, channels=32):
super().__init__()
self.time = torch.nn.Sequential(
SinusoidalTimeEmbedding(channels),
torch.nn.Linear(channels, channels),
torch.nn.SiLU(),
torch.nn.Linear(channels, channels),
)
self.in_conv = torch.nn.Conv2d(1, channels, 3, padding=1)
self.block = torch.nn.Sequential(
torch.nn.GroupNorm(4, channels),
torch.nn.SiLU(),
torch.nn.Conv2d(channels, channels, 3, padding=1),
torch.nn.GroupNorm(4, channels),
torch.nn.SiLU(),
torch.nn.Conv2d(channels, channels, 3, padding=1),
)
self.out = torch.nn.Conv2d(channels, 1, 3, padding=1)
def forward(self, x, t):
h = self.in_conv(x)
time = self.time(t)[:, :, None, None]
h = h + time
h = h + self.block(h)
return self.out(torch.nn.functional.silu(h))
shape_model = TinyDenoiser(16).to(device)
shape_x = x_train[:5].to(device)
shape_t = torch.tensor([0, 1, 2, 3, T - 1], device=device)
assert shape_model(shape_x, shape_t).shape == shape_x.shape
print("输入与输出的维度核验:", tuple(shape_x.shape), "→", tuple(shape_model(shape_x, shape_t).shape))
输入与输出的维度核验: (5, 1, 28, 28) → (5, 1, 28, 28)
第 8 步:随机时间步训练与验证早停#
每个批次为每张图独立采样 t 和噪声,因此网络在一次训练轮次中同时看见不同难度。优化只读取训练集;验证集用于保存最低噪声预测均方误差的状态。测试集此时不参与。
快速模式只训练 1 个轮次,目标是验证完整实现;完整模式才训练 20 个轮次。若要严肃比较生成质量,应增加网络容量、训练步数和多随机种子,并保存执行环境,而不是只关闭快速模式就宣称复现 DDPM。
model = TinyDenoiser(24 if FAST_MODE else 64).to(device)
optimizer = torch.optim.AdamW(model.parameters(), lr=2e-3 if FAST_MODE else 2e-4)
epochs = 1 if FAST_MODE else 20
def noise_loss(model, x0):
t = torch.randint(0, T, (len(x0),), device=x0.device)
xt, noise = q_sample(x0, t)
prediction = model(xt, t)
return torch.nn.functional.mse_loss(prediction, noise)
@torch.no_grad()
def evaluate_noise_mse(model, loader, max_batches=None):
model.eval()
values = []
eval_generator = torch.Generator(device=device).manual_seed(SEED + 7)
for batch_no, (x0,) in enumerate(loader):
if max_batches is not None and batch_no >= max_batches:
break
x0 = x0.to(device)
t = torch.randint(0, T, (len(x0),), generator=eval_generator, device=device)
noise = torch.randn(x0.shape, generator=eval_generator, device=device)
xt, _ = q_sample(x0, t, noise)
values.append(torch.nn.functional.mse_loss(model(xt, t), noise).item())
return float(np.mean(values))
history, best_state, best_valid = [], None, float("inf")
for epoch in range(epochs):
model.train()
train_values = []
for (x0,) in train_loader:
x0 = x0.to(device)
optimizer.zero_grad()
loss = noise_loss(model, x0)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
train_values.append(loss.item())
valid_mse = evaluate_noise_mse(model, valid_loader, 4 if FAST_MODE else None)
history.append({"训练轮次": epoch, "训练均方误差": np.mean(train_values), "验证均方误差": valid_mse})
if valid_mse < best_valid:
best_valid = valid_mse
best_state = copy.deepcopy(model.state_dict())
model.load_state_dict(best_state)
display(pd.DataFrame(history))
| 训练轮次 | 训练均方误差 | 验证均方误差 | |
|---|---|---|---|
| 0 | 0 | 0.215193 | 0.145876 |
第 9 步:最终测试误差与分时间步诊断#
确定模型与训练方案后,才在官方测试图像上评价。总体 MSE 会混合不同时间步 \(t\) 的结果;分组曲线更容易解释:早期时间步的 \(\boldsymbol{x}_t\) 接近原图,模型需要预测很弱、难以辨认的噪声;后期噪声更强,但图像结构也更少。
噪声 MSE 只是训练代理指标,不等价于感知质量或样本多样性。它主要用于发现网络是否优于零预测器、哪些时间段最困难。
test_mse = evaluate_noise_mse(model, test_loader, 4 if FAST_MODE else None)
print({"零预测器基线": baseline_mse, "测试噪声均方误差": test_mse})
@torch.no_grad()
def mse_by_time(model, x0, time_points):
local_generator = torch.Generator(device=device).manual_seed(SEED + 9)
rows = []
x0 = x0.to(device)
for time_index in time_points:
t = torch.full((len(x0),), int(time_index), dtype=torch.long, device=device)
noise = torch.randn(x0.shape, generator=local_generator, device=device)
xt, _ = q_sample(x0, t, noise)
mse = torch.nn.functional.mse_loss(model(xt, t), noise).item()
rows.append({"时间步": int(time_index), "噪声均方误差": mse, "信噪比": float(snr[time_index])})
return pd.DataFrame(rows)
time_diagnostics = mse_by_time(model, x_test[:128], np.linspace(0, T - 1, 8, dtype=int))
display(time_diagnostics)
time_diagnostics.plot(x="时间步", y="噪声均方误差", marker="o", figsize=(6, 3), title="分时间步噪声均方误差");
{'零预测器基线': 1.0012334883213043, '测试噪声均方误差': 0.14724910259246826}
| 时间步 | 噪声均方误差 | 信噪比 | |
|---|---|---|---|
| 0 | 0 | 0.846928 | 399.000397 |
| 1 | 5 | 0.255387 | 7.122963 |
| 2 | 11 | 0.134459 | 1.363004 |
| 3 | 16 | 0.088470 | 0.477658 |
| 4 | 22 | 0.050798 | 0.137343 |
| 5 | 27 | 0.033457 | 0.043036 |
| 6 | 33 | 0.027676 | 0.008321 |
| 7 | 39 | 0.031602 | 0.001141 |
第 10 步:去噪扩散概率模型的反向均值与逐步采样#
给定噪声预测,常用反向均值写成
再加入后验方差 \(\tilde\beta_t=\beta_t(1-\bar\alpha_{t-1})/(1-\bar\alpha_t)\) 对应的随机噪声;最后一步不再加噪。采样从 \(\boldsymbol{x}_T\sim\mathcal{N}(\boldsymbol{0},\boldsymbol{I})\) 开始,按时间步倒序循环到 0,因此扩散步数直接决定生成样本需要的计算量。
@torch.no_grad()
def p_sample(model, x, t):
predicted_noise = model(x, t)
beta_t = extract(betas, t, x)
alpha_t = extract(alphas, t, x)
alpha_bar_t = extract(alpha_bar, t, x)
mean = (x - beta_t / torch.sqrt(1.0 - alpha_bar_t) * predicted_noise) / torch.sqrt(alpha_t)
variance = extract(posterior_variance.clamp_min(1e-20), t, x)
noise = torch.randn_like(x)
nonzero = (t > 0).float().reshape(-1, 1, 1, 1)
return mean + nonzero * torch.sqrt(variance) * noise
@torch.no_grad()
def sample_ddpm(model, count=16):
model.eval()
x = torch.randn(count, 1, 28, 28, device=device)
snapshots = {}
for time_index in reversed(range(T)):
t = torch.full((count,), time_index, dtype=torch.long, device=device)
x = p_sample(model, x, t)
if time_index in {T - 1, T // 2, 0}:
snapshots[time_index] = x.detach().cpu()
return x.clamp(-1, 1).cpu(), snapshots
generated, snapshots = sample_ddpm(model, 16)
fig, axes = plt.subplots(2, 8, figsize=(12, 3.2))
for image, axis in zip(generated, axes.flat):
axis.imshow((image[0] + 1) / 2, cmap="gray", vmin=0, vmax=1)
axis.axis("off")
plt.suptitle("快速模式样本仅验证采样流程;不代表成熟生成质量")
plt.tight_layout()
第 11 步:比较线性与余弦噪声日程#
不重新训练也可以先比较两个日程保留信号的过程。余弦日程直接设计 \(\bar\alpha_t\),通常使信号衰减更平缓;线性 \(\beta_t\) 日程还会受扩散总步数的影响。下面在其他条件相同的情况下,只改变噪声日程,比较 \(\bar\alpha_t\) 与信噪比(signal-to-noise ratio,SNR)。这个比较不能单独证明生成质量的高低;若要公平比较生成质量,必须分别使用相同训练次数和计算量训练模型。
这个分析也暴露一个重要边界:只有 40 步时,照搬为 1000 步设计的 beta 数值可能使最终状态仍保留太多图像,或单步噪声过强。
def cosine_alpha_bar(steps, offset=0.008):
grid = torch.linspace(0, steps, steps + 1)
curve = torch.cos(((grid / steps + offset) / (1 + offset)) * np.pi / 2) ** 2
curve = curve / curve[0]
return curve[1:].clamp(min=1e-5)
cosine_bar = cosine_alpha_bar(T)
schedule_frame = pd.DataFrame({
"时间步": np.arange(T),
"线性日程累计系数": alpha_bar.cpu().numpy(),
"余弦日程累计系数": cosine_bar.numpy(),
})
schedule_frame.plot(x="时间步", y=["线性日程累计系数", "余弦日程累计系数"],
figsize=(7, 3.5), title="只比较前向信号保留,不比较生成质量")
plt.ylabel("累计信号系数");
结论、局限与常见错误#
本案例把解析前向采样、噪声预测损失和反向采样连成可运行的完整过程,但快速模式生成的图像通常较粗糙。局限包括:
小网络没有多尺度 U-Net、注意力、EMA 和长时间训练;
单次运行无法衡量随机性,展示 16 张图也不足以评价覆盖度;
噪声 MSE 不等价于 FID、精确率/召回率或人工感知质量;
线性与余弦日程只比较了前向曲线,没有完成等预算重训练;
Fashion-MNIST 是低分辨率灰度教学数据,结论不能直接外推到真实图像生成。
常见错误:混淆 \(\alpha_t\) 与 \(\bar\alpha_t\)、对整个批次只采样一个时间步 \(t\)、在最后一步仍加噪、把模型输出当作 \(\boldsymbol{x}_0\) 而代码却按噪声参数化采样、忘记把时间系数重塑为 \((B,1,1,1)\),以及把快速演示样本宣传为完整训练结果。
综合练习#
用蒙特卡洛方法验证:固定 \(\boldsymbol{x}_0\) 和 \(t\) 时,\(\boldsymbol{x}_t\) 的经验均值和方差符合正文给出的解析采样公式。
分别使用线性与余弦日程,在完全相同预算下训练 3 个种子,比较验证 MSE。
添加 EMA 参数并只用 EMA 模型采样,解释训练参数与采样参数的职责。
把网络改成小型 U-Net,逐层写出 \((B,C,H,W)\) 的维度,并检查跳连两端的维度是否一致。
完整训练后计算生成图与测试图的特征分布距离,并同时检查最近邻,避免只展示“最好看”的样本。