扩散模型实战:在Colab上零门槛生成MNIST风格手写数字

最近有不少朋友问我,想体验一下扩散模型(Diffusion Model)的魅力,但又担心环境配置复杂、代码调试困难。特别是看到那些动辄需要多张GPU、训练好几天的项目,新手往往望而却步。其实,入门扩散模型完全可以更轻松、更直观。今天,我就带大家在Google Colab这个免费的云端平台上,用DDPM(Denoising Diffusion Probabilistic Models)模型,快速跑通一个生成MNIST风格手写数字的完整流程。

整个过程无需本地安装任何复杂环境,不依赖高性能显卡,重点在于理解核心流程和避开常见陷阱。无论你是想快速验证一个想法,还是学习扩散模型的基本工作原理,这篇文章都能给你一个清晰的路线图。我们会从最基础的Colab环境配置讲起,一步步完成数据加载、模型定义、训练和采样生成,并穿插我实际调试中遇到的那些“坑”和解决方案。你会发现,生成第一张属于自己的“AI手写数字”,可能比你想象的要简单得多。

1. 环境准备与Colab高效使用指南

Google Colab是体验AI模型的绝佳起点,它提供了免费的GPU资源(通常是Tesla T4或K80)和预装好的Python环境。但要想用得顺手,避免运行时中断和数据丢失,有几个小技巧必须掌握。

首先,打开Colab后,你需要确认运行时类型。在菜单栏选择 “运行时” -> “更改运行时类型”,在硬件加速器下拉菜单中,务必选择 “GPU”。这能让你后续的模型训练速度提升数十倍。连接成功后,你可以在代码单元格中执行以下命令来验证GPU是否可用:

import torch
print(f"PyTorch版本: {torch.__version__}")
print(f"CUDA是否可用: {torch.cuda.is_available()}")
print(f"当前GPU: {torch.cuda.get_device_name(0) if torch.cuda.is_available() else '无'}")

接下来是依赖安装。Colab自带的PyTorch版本可能不是最新的,且缺少一些可视化库。我习惯在开头用一个单元格集中安装所有依赖,避免后续报错。

!pip install torch torchvision --upgrade --quiet
!pip install matplotlib ipywidgets tqdm --quiet

注意:Colab的运行时是临时的。一旦浏览器标签页关闭或闲置时间过长,运行时会被回收,所有安装的包和存储在/content目录下的文件都会丢失。因此,重要代码和训练好的模型需要及时保存。

我强烈推荐将Colab与Google Drive联动。这样,你可以把数据集、代码脚本甚至训练好的模型权重直接保存到网盘,下次打开还能接着用。挂载Drive只需几行代码:

from google.colab import drive
drive.mount('/content/drive')

执行后会弹出一个授权窗口,按照提示操作即可。之后,你就可以像操作本地文件夹一样访问/content/drive/MyDrive/下的内容了。例如,把当前工作目录切换到Drive里的一个项目文件夹:

import os
project_path = '/content/drive/MyDrive/DDPM_MNIST'
os.makedirs(project_path, exist_ok=True)
os.chdir(project_path)

最后,关于资源监控。Colab的免费GPU有内存限制,长时间运行复杂模型可能爆内存。你可以用nvidia-smi命令随时查看显存使用情况:

!nvidia-smi

如果发现显存接近耗尽,可以考虑减小批次大小(batch size)或模型尺寸。一个稳定的环境是成功的第一步,这些准备工作能帮你节省大量排查环境问题的时间。

2. 理解DDPM:从噪声到图像的魔法

在动手写代码之前,我们花点时间捋清楚DDPM到底在做什么。很多人被论文里复杂的数学公式吓退,但其实它的核心思想非常直观:教会一个模型如何一步步把一张纯噪声图片“去噪”,恢复成有意义的图像。

想象一下,你有一张清晰的手写数字“7”的图片。我们通过一个“加噪”过程,逐步向这张图片添加随机噪声。经过很多步(比如1000步)之后,这张图片会变成一张看起来完全是随机像素点的噪声图。这个正向过程是确定的、可计算的。

DDPM要学习的,是上述过程的逆过程——去噪。给定一张噪声图,以及它处于加噪过程的第几步(时间步t),模型需要预测出这一步所添加的噪声是什么。如果模型能准确预测每一步的噪声,那么从一张完全随机的噪声开始,逐步减去模型预测的噪声,最终就能得到一张全新的、清晰的手写数字图片。

整个流程可以概括为两个阶段:

  1. 训练阶段:模型学习预测噪声。我们拿一张真实图片,随机选择一个时间步t,用公式合成第t步的加噪图片,然后让模型(通常是一个U-Net)根据这张加噪图片和t,去预测我们加入的噪声。损失函数就是预测噪声和真实加入噪声的差距。
  2. 采样(生成)阶段:模型执行去噪。我们从标准高斯分布(纯噪声)中采样一张图,作为第T步(最后一步)的图片。然后,从t=T开始,到t=0结束,循环执行:模型根据当前噪声图和当前步数t,预测噪声;然后用一个特定的更新规则,用预测的噪声计算出t-1步的(噪声更少的)图片。循环结束,就得到生成的图片。

这里涉及几个关键的超参数和概念,我整理了一个表格,方便你后续查阅:

概念符号/名称作用与常见设置说明
时间步Ttimesteps加噪/去噪的总步数,通常为1000。步数越多,过程越精细,但采样速度越慢。MNIST数据简单,500步也够用。
噪声调度beta_schedule控制每一步加噪的强度。常见有linear(线性)和cosine(余弦)。linear简单直接;cosine在开始和结束时变化平缓,中间变化快,有时效果更好。
U-NetUNetModel模型主干。负责根据带噪图片和时间步t,预测噪声。包含下采样、中间瓶颈和上采样层,通常带有残差连接和自注意力机制。
批次大小batch_size一次训练所抓取的数据样本数量。在Colab的T4 GPU上,对于MNIST(28x28小图),可以设置到64甚至128。

理解了这些,再看代码就不会觉得是一团乱麻了。接下来,我们就开始构建模型的核心部分。

3. 构建DDPM模型核心组件

我们将把模型拆解成几个模块来构建:定义噪声调度、实现前向加噪过程、搭建U-Net去噪模型。我会用代码块展示关键部分,并解释其中容易出错的细节。

首先,实现噪声调度。这里我们提供线性和余弦两种选择,你可以通过参数切换。

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

def linear_beta_schedule(timesteps):
    """线性噪声调度,beta值从很小线性增长到较大值。"""
    scale = 1000.0 / timesteps
    beta_start = scale * 0.0001
    beta_end = scale * 0.02
    return torch.linspace(beta_start, beta_end, timesteps, dtype=torch.float64)

def cosine_beta_schedule(timesteps, s=0.008):
    """余弦噪声调度,来自Improved DDPM论文。"""
    steps = timesteps + 1
    x = torch.linspace(0, timesteps, steps, dtype=torch.float64)
    alphas_cumprod = torch.cos(((x / timesteps) + s) / (1 + s) * math.pi * 0.5) ** 2
    alphas_cumprod = alphas_cumprod / alphas_cumprod[0]
    betas = 1 - (alphas_cumprod[1:] / alphas_cumprod[:-1])
    return torch.clamp(betas, 0, 0.999)

提示:对于MNIST这种结构简单的数据集,线性调度通常就足够了,且计算更快。余弦调度在生成更复杂、细节更丰富的图像(如人脸、自然场景)时可能有优势。

接下来是扩散过程的核心类GaussianDiffusion。它不包含可训练参数,主要职责是管理噪声调度,并提供前向加噪(q_sample)、计算后验分布(q_posterior_mean_variance)和反向采样(p_sample)所需的各项计算。初始化时,它会根据选择的调度预先计算好所有时间步的中间变量,这是一个典型的“用空间换时间”的优化。

class GaussianDiffusion:
    def __init__(self, timesteps=1000, beta_schedule='linear'):
        self.timesteps = timesteps
        # 选择并计算beta序列
        if beta_schedule == 'linear':
            betas = linear_beta_schedule(timesteps)
        elif beta_schedule == 'cosine':
            betas = cosine_beta_schedule(timesteps)
        else:
            raise ValueError(f'未知的beta调度: {beta_schedule}')
        self.betas = betas.to(torch.float32)  # 确保为float32,与模型精度匹配

        # 计算一系列派生变量,用于后续公式
        self.alphas = 1. - self.betas
        self.alphas_cumprod = torch.cumprod(self.alphas, dim=0)
        self.alphas_cumprod_prev = F.pad(self.alphas_cumprod[:-1], (1, 0), value=1.0)

        # 用于q(x_t | x_0)的计算
        self.sqrt_alphas_cumprod = torch.sqrt(self.alphas_cumprod)
        self.sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - self.alphas_cumprod)

        # 用于后验q(x_{t-1} | x_t, x_0)的计算
        self.posterior_variance = self.betas * (1. - self.alphas_cumprod_prev) / (1. - self.alphas_cumprod)
        self.posterior_log_variance_clipped = torch.log(torch.clamp(self.posterior_variance, min=1e-20))
        self.posterior_mean_coef1 = self.betas * torch.sqrt(self.alphas_cumprod_prev) / (1. - self.alphas_cumprod)
        self.posterior_mean_coef2 = (1. - self.alphas_cumprod_prev) * torch.sqrt(self.alphas) / (1. - self.alphas_cumprod)

    def q_sample(self, x_start, t, noise=None):
        """前向扩散:根据x_0和t,计算加噪后的x_t。"""
        if noise is None:
            noise = torch.randn_like(x_start)
        sqrt_alphas_cumprod_t = self._extract(self.sqrt_alphas_cumprod, t, x_start.shape)
        sqrt_one_minus_alphas_cumprod_t = self._extract(self.sqrt_one_minus_alphas_cumprod, t, x_start.shape)
        return sqrt_alphas_cumprod_t * x_start + sqrt_one_minus_alphas_cumprod_t * noise

    def _extract(self, a, t, x_shape):
        """辅助函数:从张量a中提取对应时间步t的值,并广播到x_shape形状。"""
        batch_size = t.shape[0]
        out = a.to(t.device).gather(0, t.long())
        return out.view(batch_size, *((1,) * (len(x_shape) - 1)))

模型的心脏是U-Net。它的输入是带噪图像x_t和时间步t的嵌入(embedding),输出是对噪声的预测。下面是一个简化但功能完整的U-Net实现,包含了残差块、注意力机制和下/上采样。

class ResidualBlock(nn.Module):
    """残差块,融合时间步信息。"""
    def __init__(self, in_channels, out_channels, time_channels, dropout=0.1):
        super().__init__()
        # 第一个卷积组
        self.conv1 = nn.Sequential(
            nn.GroupNorm(32, in_channels),
            nn.SiLU(),
            nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1)
        )
        # 时间步嵌入的投影层
        self.time_emb = nn.Sequential(
            nn.SiLU(),
            nn.Linear(time_channels, out_channels)
        )
        # 第二个卷积组
        self.conv2 = nn.Sequential(
            nn.GroupNorm(32, out_channels),
            nn.SiLU(),
            nn.Dropout(p=dropout),
            nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1)
        )
        # 快捷连接,如果输入输出通道数不同,用1x1卷积对齐
        self.shortcut = nn.Conv2d(in_channels, out_channels, kernel_size=1) if in_channels != out_channels else nn.Identity()

    def forward(self, x, t_emb):
        h = self.conv1(x)
        # 将时间步嵌入加到特征图上,需调整维度 [B, C] -> [B, C, 1, 1]
        h = h + self.time_emb(t_emb)[:, :, None, None]
        h = self.conv2(h)
        return h + self.shortcut(x)

构建完整的U-Net时,需要仔细设计下采样和上采样的通道数变化。对于MNIST(单通道,28x28),一个较小的网络就足够。下面是一个适配MNIST的UNetModel初始化示例,其中channel_mult=(1, 2, 2)意味着在下采样过程中,通道数依次变为 model_channels 的 1倍、2倍、2倍。

# 在Colab中初始化模型
device = "cuda" if torch.cuda.is_available() else "cpu"
model = UNetModel(
    in_channels=1,        # MNIST是灰度图,单通道
    model_channels=64,     # 基础通道数,可调整以改变模型大小
    out_channels=1,        # 预测的噪声图也是单通道
    channel_mult=(1, 2, 2), # 下采样各阶段通道倍增因子
    attention_resolutions=[], # MNIST简单,可以不用注意力机制以加速
    num_heads=4,
    dropout=0.1
).to(device)
print(f"模型参数量: {sum(p.numel() for p in model.parameters()):,}")

model_channels从96降低到64,并移除attention_resolutions,能显著减少模型参数量和计算量,在Colab上训练更快,且对MNIST的生成质量影响很小。这是针对具体任务和资源的有效调优。

4. 训练与采样全流程实操

有了模型组件,我们就可以把它们组装起来,开始真正的训练和生成过程了。这部分我会结合代码,详细说明每一步的目的,并标注出我踩过坑的地方。

第一步:准备MNIST数据。 PyTorch的torchvision库让这个过程变得极其简单。

from torchvision import datasets, transforms
from torch.utils.data import DataLoader

# 数据预处理:将像素值从[0,255]归一化到[-1, 1],这是DDPM常用的输入范围
transform = transforms.Compose([
    transforms.ToTensor(),  # 转换为Tensor,并缩放到[0,1]
    transforms.Normalize(mean=[0.5], std=[0.5])  # 归一化到[-1, 1]
])

# 下载并加载训练集
train_dataset = datasets.MNIST(root='./data', train=True, download=True, transform=transform)
train_loader = DataLoader(train_dataset, batch_size=128, shuffle=True, num_workers=2, pin_memory=True)

print(f"训练集样本数: {len(train_dataset)}")
print(f"批次大小: {train_loader.batch_size}")
print(f"图像形状示例: {train_dataset[0][0].shape}")  # 应为 torch.Size([1, 28, 28])

第二步:配置训练循环。 这是DDPM训练最核心的部分,但代码逻辑非常清晰。

import torch.optim as optim
from tqdm.auto import tqdm

# 初始化扩散模型和噪声调度
timesteps = 500  # 对于MNIST,500步足够
diffusion = GaussianDiffusion(timesteps=timesteps, beta_schedule='linear')
diffusion = diffusion.to(device)

# 定义优化器
optimizer = optim.AdamW(model.parameters(), lr=1e-4)  # 使用AdamW,更稳定
epochs = 20  # 训练轮数

model.train()
for epoch in range(epochs):
    epoch_loss = 0.0
    progress_bar = tqdm(train_loader, desc=f'Epoch {epoch+1}/{epochs}')
    for batch_idx, (images, _) in enumerate(progress_bar):
        images = images.to(device)
        batch_size = images.shape[0]

        # 1. 为批次中的每个样本随机采样一个时间步t
        t = torch.randint(0, timesteps, (batch_size,), device=device).long()

        # 2. 从标准正态分布采样噪声
        noise = torch.randn_like(images)

        # 3. 前向加噪,得到x_t
        x_t = diffusion.q_sample(images, t, noise)

        # 4. 模型预测噪声
        predicted_noise = model(x_t, t)

        # 5. 计算均方误差损失
        loss = F.mse_loss(predicted_noise, noise)

        # 6. 反向传播与优化
        optimizer.zero_grad()
        loss.backward()
        # 可选:梯度裁剪,防止训练不稳定
        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
        optimizer.step()

        epoch_loss += loss.item()
        progress_bar.set_postfix({'Batch Loss': loss.item()})

    avg_loss = epoch_loss / len(train_loader)
    print(f'Epoch {epoch+1} 平均损失: {avg_loss:.4f}')

    # 每5个epoch保存一次模型检查点
    if (epoch + 1) % 5 == 0:
        checkpoint_path = f'/content/drive/MyDrive/DDPM_MNIST/checkpoint_epoch_{epoch+1}.pt'
        torch.save({
            'epoch': epoch,
            'model_state_dict': model.state_dict(),
            'optimizer_state_dict': optimizer.state_dict(),
            'loss': avg_loss,
        }, checkpoint_path)
        print(f'模型已保存至: {checkpoint_path}')

注意:损失值(MSE)会从较高的值(如0.9)开始,并随着训练稳步下降。对于MNIST,当平均损失降至0.02以下时,模型通常就能生成比较清晰的数字了。如果损失不下降或出现NaN,请检查学习率是否过高,或尝试加入梯度裁剪。

第三步:采样生成新图像。 训练完成后,最激动人心的部分来了——从随机噪声中创造新的手写数字。

@torch.no_grad()
def sample_loop(model, diffusion, image_size, batch_size=16, channels=1):
    """完整的反向去噪采样循环。"""
    model.eval()
    shape = (batch_size, channels, image_size, image_size)
    # 1. 从标准正态分布采样初始噪声 x_T
    img = torch.randn(shape, device=device)

    # 2. 从t=T到t=1循环去噪
    for i in tqdm(reversed(range(0, timesteps)), desc='采样进度', total=timesteps):
        t = torch.full((batch_size,), i, device=device, dtype=torch.long)
        # 3. 预测噪声,并计算前一步的x_{t-1}
        img = diffusion.p_sample(model, img, t, clip_denoised=True)
    # 4. 循环结束,得到生成的x_0
    return img

# 生成64张图像
generated_images = sample_loop(model, diffusion, image_size=28, batch_size=64, channels=1)
# generated_images 形状为 [64, 1, 28, 28],值范围约为[-1, 1]

第四步:可视化结果。 将生成的张量转换回图片格式并显示。

import matplotlib.pyplot as plt
import numpy as np

def plot_images(images, nrows=8, ncols=8):
    """将一批图像以网格形式绘制出来。"""
    fig, axes = plt.subplots(nrows, ncols, figsize=(ncols, nrows))
    images = images.cpu().numpy()
    # 将值从[-1, 1]转换回[0, 1]以便显示
    images = (images + 1) / 2.0
    for idx, ax in enumerate(axes.flat):
        if idx < len(images):
            ax.imshow(images[idx].squeeze(), cmap='gray')
        ax.axis('off')
    plt.tight_layout()
    plt.show()

# 绘制生成的64张图像
plot_images(generated_images, nrows=8, ncols=8)

如果一切顺利,你将看到一个8x8的网格,里面是模型生成的、形态各异的手写数字。第一次看到自己训练的模型“无中生有”出这些图像,成就感是非常足的。

5. 避坑指南与效果优化技巧

在实际操作中,你可能会遇到一些意料之外的问题。这里我总结了几类最常见的情况及其解决方法,希望能帮你快速排雷。

问题一:CUDA内存不足(Out Of Memory, OOM) 这是Colab上最常见的问题。症状是运行单元格时内核崩溃或报错CUDA out of memory

  • 首要解决方案:减小批次大小(Batch Size)。将train_loaderbatch_size从128降到64或32。
  • 其次,简化模型。降低UNetModel中的model_channels(如从64降到32),或减少channel_mult中的乘数(如(1, 2))。
  • 采样时也需注意。生成图像时,sample_loop函数中的batch_size也不要设得太大,尤其是在高分辨率图像生成中。

问题二:生成的图像全黑、全灰或全是噪声 这通常意味着训练出了问题。

  • 检查数据归一化。确保输入图像的像素值被正确归一化到[-1, 1]。MNIST的ToTensor()得到[0,1],经过Normalize(mean=[0.5], std=[0.5])后才是[-1,1]。如果归一化范围不对,模型学习的目标会混乱。
  • 检查损失曲线。损失是否在持续下降?如果损失震荡剧烈或几乎不变,尝试降低学习率(如从1e-4降到5e-5),或使用学习率预热(learning rate warmup)。
  • 确认时间步t的输入。在训练时,随机采样的时间步t需要被正确嵌入并输入U-Net。确保你的ResidualBlock正确地将时间步嵌入加到了特征图上。
  • 验证前向加噪公式。可以写一个简单的测试,对一张图像进行q_sample加噪,并可视化t=0, t=250, t=499(最后一步)的结果。t=0应该接近原图,t=499应该接近纯噪声。如果中间状态不对,说明sqrt_alphas_cumprod等预计算变量有误。

问题三:训练速度慢 Colab的免费GPU算力有限。

  • 使用混合精度训练。PyTorch的AMP(Automatic Mixed Precision)可以大幅减少显存占用并加速训练,尤其对支持Tensor Core的T4 GPU效果明显。
    from torch.cuda.amp import autocast, GradScaler
    scaler = GradScaler()
    # 在训练循环中
    with autocast():
        predicted_noise = model(x_t, t)
        loss = F.mse_loss(predicted_noise, noise)
    scaler.scale(loss).backward()
    scaler.step(optimizer)
    scaler.update()
    
  • 降低timesteps。对于MNIST,将timesteps从1000降到500甚至250,可以成倍缩短采样时间,而对生成质量影响有限。
  • 关闭进度条tqdm等进度条会带来少量开销。在调试完成后,可以暂时关闭以获得极致的速度。

问题四:生成的数字多样性不足或模式单一 如果生成的数字总是那几种样子,可能是模型容量不足或训练数据覆盖不全。

  • 增加模型容量。适当增加model_channels或使用更深的网络(调整channel_mult,如(1,2,4,8))。
  • 引入注意力机制。在UNetModelattention_resolutions参数中加入分辨率(如[16]),让模型在特定层级关注图像的全局关系。
  • 延长训练时间。MNIST虽然简单,但训练轮数(epochs)太少也可能导致欠拟合。可以尝试训练更多轮次,观察损失是否还能继续下降。

最后,分享一个我常用的效果微调技巧:在采样时,可以尝试调整diffusion.p_sample中的clip_denoised参数。将其设置为True(默认)会强制将预测的x_0裁剪到[-1,1],这通常更稳定。但在某些情况下,设置为False可能会产生对比度更高、笔触更鲜明的数字,当然也可能引入一些异常像素点,需要你根据生成结果进行权衡。

整个流程跑通后,你可以尝试修改噪声调度(试试cosine)、调整U-Net结构、甚至换用其他更简单的数据集(如Fashion-MNIST)来生成服装图片,从而更深入地理解每个组件的作用。扩散模型的门槛并没有那么高,从这个小项目开始,你已经掌握了它的核心工作流程。

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐