news 2026/9/30 20:56:53

扩散模型(DDPM)的训练与推理

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
扩散模型(DDPM)的训练与推理

扩散模型(Diffusion Model)作为近年来生成式 AI 领域的核心技术,已广泛应用于图像生成、语音合成、文本生成等场景。不同于 GAN 的对抗训练模式,扩散模型通过模拟 “逐步加噪 - 逐步去噪” 的物理过程实现数据生成,具有训练稳定、生成质量高的显著优势。本文将从原理入手,拆解扩散模型的训练与推理全过程,并提供可落地的伪代码和 Python 实现示例。

一、扩散模型训练过程

1.1 训练核心逻辑

训练的本质是让模型(噪声预测器)尽可能准确地预测前向扩散过程中添加的噪声。具体步骤如下:
1. 采样一批真实数据;
2. 随机采样步长;
3. 采样高斯噪声;
4. 根据前向扩散公式计算;
5. 将和输入模型,得到预测噪声;
6. 计算预测噪声与真实噪声的 MSE 损失,反向传播更新模型参数。

1.2 训练代码示例

import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader import numpy as np T = 1000 beta = torch.linspace(1e-4, 0.02, T) alpha = 1 - beta alpha_bar = torch.cumprod(alpha, dim=0) # 累计乘积 def train_diffusion(model, dataloader, epochs, device): model.to(device) criterion = nn.MSELoss() optimizer = optim.Adam(model.parameters(), lr=1e-4) alpha_bar = alpha_bar.to(device) for epoch in range(epochs): model.train() total_loss = 0.0 for batch in dataloader: x0 = batch.to(device) # 假设输入为3通道图像,shape=(B,3,H,W) B = x0.shape[0] # 随机采样步长t t = torch.randint(1, T+1, (B,), device=device) # 采样高斯噪声 eps = torch.randn_like(x0) # 计算xt alpha_bar_t = alpha_bar[t-1].reshape(B, 1, 1, 1) # t从1开始,索引从0开始 xt = torch.sqrt(alpha_bar_t) * x0 + torch.sqrt(1 - alpha_bar_t) * eps # 预测噪声 eps_theta = model(xt, t-1) # 嵌入层索引从0开始 # 计算损失并更新 loss = criterion(eps_theta, eps) optimizer.zero_grad() loss.backward() optimizer.step() total_loss += loss.item() print(f"Epoch {epoch+1}, Loss: {total_loss/len(dataloader):.4f}")

二、扩散模型推理过程

2.1 推理核心逻辑

推理(生成)是反向扩散的过程:从纯噪声出发,利用训练好的模型逐步去噪,最终得到接近真实数据的。具体步骤如下:
1. 初始化为随机高斯噪声;
2. 反向迭代:
将和输入模型,得到预测噪声;
根据反向扩散公式计算;
3. 迭代结束后,即为生成的结果。

2.2 推理代码示例

def sample_diffusion(model, shape, device): """ 扩散模型推理函数 :param model: 训练好的噪声预测模型 :param shape: 生成数据的形状,如(B, 3, 32, 32) :param device: 运行设备 :return: 生成的数据x0 """ model.eval() alpha = alpha.to(device) alpha_bar = alpha_bar.to(device) beta = beta.to(device) # 初始化xT为纯噪声 xt = torch.randn(shape, device=device) with torch.no_grad(): # 推理阶段禁用梯度 for t in range(T, 0, -1): # 1. 准备时间步(适配嵌入层) t_tensor = torch.tensor([t-1], device=device).repeat(shape[0]) # 2. 预测噪声 eps_theta = model(xt, t_tensor) # 3. 计算反向扩散的均值和方差 alpha_t = alpha[t-1] alpha_bar_t = alpha_bar[t-1] beta_t = beta[t-1] # 均值μ_t mean = (1 / torch.sqrt(alpha_t)) * ( xt - (1 - alpha_t) / torch.sqrt(1 - alpha_bar_t) * eps_theta ) # 方差σ_t(简化为sqrt(β_t)) std = torch.sqrt(beta_t) # 4. 采样噪声(t=1时为0) if t > 1: z = torch.randn_like(xt) else: z = torch.zeros_like(xt) # 5. 更新xt为x_{t-1} xt = mean + std * z # 归一化到[0,1](可选,根据数据分布调整) x0 = torch.clamp(xt, -1, 1) x0 = (x0 + 1) / 2 return x0
版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/8/23 9:50:18

手把手教你用Coppeliasim实现UR5机械臂关节控制(附PID参数调试指南)

从零掌握Coppeliasim中UR5机械臂的高精度关节控制 在工业自动化领域,六轴协作机械臂已成为智能制造的核心设备之一。UR5作为Universal Robots的经典机型,凭借其高灵活性和易用性广受欢迎。而Coppeliasim(原V-REP)作为一款功能强大…

作者头像 李华
网站建设 2026/8/23 9:50:18

终极防撤回指南:一键破解微信QQ消息撤回限制

终极防撤回指南:一键破解微信QQ消息撤回限制 【免费下载链接】RevokeMsgPatcher :trollface: A hex editor for WeChat/QQ/TIM - PC版微信/QQ/TIM防撤回补丁(我已经看到了,撤回也没用了) 项目地址: https://gitcode.com/GitHub_…

作者头像 李华
网站建设 2026/8/23 9:50:18

深度测评:2026年YOLO计算机视觉模型横评!目标检测哪家强?

点击上方“小白学视觉”,选择加"星标"或“置顶” 重磅干货,第一时间送达文章来源于微信公众号:漠岩yggg本文仅用于学术分享,如有侵权,请联系后台作删文处理——目标检测哪家强?一篇帮你搞懂所有Y…

作者头像 李华
网站建设 2026/8/23 9:50:18

新手必看!PasteMD智能美化工具从安装到上手指南

新手必看!PasteMD智能美化工具从安装到上手指南 1. 为什么你需要PasteMD 在日常工作和学习中,我们经常遇到这样的困扰: 从会议录音转写的笔记杂乱无章,需要手动整理结构复制的代码片段粘贴到Markdown中失去格式和高亮同事发来的…

作者头像 李华
网站建设 2026/8/23 9:50:18

配电网有功电压控制:多智能体强化学习的奇妙之旅

配电网有功电压控制的多智能体强化学习(代码) 针对电压主动控制问题的不同场景,采用7种最先进的MARL算法进行了大规模实验,将电压约束转化为势垒函数,并从实验结果中观察到设计合适的电压势垒函数的重要性。 主动电压控…

作者头像 李华
网站建设 2026/8/23 9:50:18

JQ6500_Serial库详解:Arduino控制MP3模块全指南

1. JQ6500_Serial 库深度解析:面向嵌入式工程师的 MP3 模块全功能控制指南JQ6500_Serial 是一个专为 Arduino 平台设计的轻量级、高可靠性的串口通信库,用于完整控制 JQ6500 系列 MP3 解码模块(包括 JQ6500-28P 和 JQ6500-16P)。该…

作者头像 李华