--- created: 2025-08-03 21:19:11 tags: - "Research" - "基础" - "Diffusion" - "DDPM" - "PyTorch" - "代码实现" --- # 1. 准备工作 ```python import torch import torch.nn as nn import torch.nn.functional as F import torchvision import cv2 import numpy as np import einops from torch.utils.data import DataLoader from torchvision.transforms import Compose, Lambda, ToTensor ``` 下载并查看MNIST数据集的信息 ```python from torchvision.transforms import ToTensor def download_dataset(): mnist = torchvision.datasets.MNIST(root='data/mnist', download=True) print('length of MNIST', len(mnist)) id = 4 # 在数据集中取出编号为4的图片 img, label = mnist[id] # 打印它的图片和标签 print(img) print(label) # On computer with monitor # img.show() tensor = ToTensor()(img) print(tensor.shape) print(tensor.max()) print(tensor.min()) if __name__ == '__main__': download_dataset() ``` 第一行输出表明,MNIST数据集里有60000张图片。而从第二行和第三行输出中,我们发现每一项数据由图片和标签组成,图片是大小为28x28的PIL格式的图片,标签表明该图片是哪个数字。我们可以用`torchvision`里的`ToTensor()`把PIL图片转成PyTorch张量,进一步查看图片的信息。最后三行输出表明,每一张图片都是**单通道图片(灰度图)**,**颜色值的取值范围是0~1**。 我们可以用下面的代码预处理数据并创建`DataLoader`。由于DDPM会把图像和正态分布关联起来,我们更希望图像颜色值的取值范围是`[-1, 1]`。为此,我们可以对图像做一个线性变换,**减0.5再乘2**。 ```python def get_dataloader(batch_size: int): # Compose: 一个用于组合多个变换的类, 它可以将多个变换操作连接起来称为一个系列的变换, 然后用于dataset # ToTensor: 将图像转换为tensor张量. 转换后的张量形状为 (C, H, W),其中 C 是通道数,H 是高度,W 是宽度。此外,像素值会被归一化到 [0, 1] 区间。 # Lambda: 定义了一个匿名函数, 输入为x, 将x的每个像素值减去 0.5,然后乘以 2。这样做的目的是将像素值从 [0, 1] 映射到 [-1, 1] 区间 transform = Compose([ToTensor(), Lambda(lambda x: (x - 0.5) * 2)]) dataset = torchvision.datasets.MNIST(root='./data/mnist', transform=transform) return DataLoader(dataset, batch_size=batch_size, shuffle=True) ``` # 2. DDPM类实现 在代码中,我们要实现一个DDPM类。它**维护了扩散过程中的一些常量**(比如 $\alpha$ ),并且可以**计算正向过程和反向过程的结果**。 先来实现一下DDPM类的初始化函数。一开始,我们遵从论文的配置,用`torch.linspace(min_beta, max_beta, n_steps)`从`min_beta`到`max_beta`线性地生成`n_steps`个时刻的 $\beta$ 。接着,我们根据公式 $\alpha_t=1-\beta_t,\bar{\alpha}_t=\prod_{i=1}^t\alpha_i$ ,计算每个时刻的`alpha`和`alpha_bar`。注意,为了方便实现,我们让`t`的取值**从0开始**,要比论文里的 $t$ 少1。 ```python class DDPM(): # n_steps 就是论文中的 T def __init__(self, device, n_steps: int, min_beta: float = 0.0001, max_beta: float = 0.02): # 初始化beta, 这里选择 **线性变化** 从0.0001到0.002 betas = torch.linspace(min_beta, max_beta, n_steps).to(device) # 初始化alpha, alpha = 1 - beta, 在笔记中写的很清楚了 alphas = 1 - betas # 每个时刻的累乘都要存起来, 所以要先开个数组, 这里还没开始计算, 所以开了个和alphas数组大小相等的空数组 alpha_bars = torch.empty_like(alphas) product = 1 # 累乘先初始化为一 ~~这是常识~~ # 开始计算累乘 for i, alpha in enumerate(alphas): pro *= alpha alpha_bars[i] = product # 将这些变量作为类的属性存储,可以在类的其他方法中直接访问和使用这些变量。 self.betas = betas self.n_steps = n_steps self.alphas = alphas self.alpha_bars = alpha_bars ``` # 3. 实现[[Diffusion#4.1 前向过程|前向过程]] 根据我们推导出的 $$x_{t}=\sqrt{\overline{\alpha}_t}x_0+\sqrt{1-\overline{\alpha}_t}z_t $$ 我们可以通过已知的 $x_0$ 和前面准备的 $\alpha_t$ 的累乘, 求出任意时刻的图像分布 $x_t$ ```python def sample_forward(self, x, t, eps=None): # alpha_bars本来是一个一维数组, 形状为(batch_size) # 图像是一个四维张量, 形状是(batch_size, channels, height, width) # 噪音要加到图像上去, 所以alpha_bar也要和图像的的形状一样 # 这里的reshape用到了 **广播技术** alpha_bar = self.alpha_bars[t].reshape(-1,1,1,1) if eps is None: # 同样, eps的形状也是(batch_size, channels, height, width) eps = torch.randn_like(x) # eps 就是公式中的 $z_t$ , 也就是服从高斯分布的噪音 res = eps * torch.sqrt(1 - alpha_bar) + torch.sqrt(alpha_bar) * x # res即为加完噪音后的 $x_t$ return res ``` # 4. 实现[[Diffusion#4.2 逆向过程|反向过程]] 在反向过程中, DDPM 会使用神经网络 (可以是[[U-net]]也可以是其他) 预测每一轮去噪的均值, 把 $x_t$ 还原回 $x_0$, 这样就实现了图像生成 算法即为论文中提到的 $\text{Algorithm2 \; Sampling}$ 其中核心的公式是: $$ \mathbf{x}_{t-1}=\frac{1}{\sqrt{\alpha_{t}}}\left(\mathbf{x}_{t}-\frac{1-\alpha_{t}}{\sqrt{1-\bar\alpha_{t}}}\boldsymbol{\epsilon}_{\theta}(\mathbf{x}_{t},t)\right)+\sigma_{t}\mathbf{z} $$ 其中前面一大坨是均值, 也就是代码中的 `mean` 公式如下: $$ \tilde{\mu}_t=\frac1{\sqrt{a_t}}(x_t-\frac{\beta_t}{\sqrt{1-\overline{a_t}}}z_{t}) $$ 后面一小撮是加的噪音, 在代码中的 `noise` 会有体现 ```python def sample_backward(self, img_shape, net, device, simple_var=True): """ - img_shape: 生成图像的形状, (batch_size, channels, height, width) - net: 用于预测噪音的神经网络 - device: 设备 (CPU 或 GPU), 用于将数据和模型移动到相应的设备上 - simple_var: 布尔值,决定是否使用简单的方差公式, 默认为 True """ # x是随机采样得到的噪音图像, 它的形状需要和我们所需要输出的图像的形状相同, 因为我们就是对这个x不断进行逆采样得到输出图像的 x = torch.randn(img_shape).to(device) net = net.to(device) # 从 self.n_steps - 1 开始,逐步递减到 0,进行逆向采样。 for t in range(self.n_steps - 1, -1, -1): x = self.sample_backward_step(x, t, net, simple_var) # 返回最终生成的图像结果 return x def sample_backward_step(self, x_t, t, net, simple_var=True): """ - x_t: 当前时刻t的图像分布, 形状为 (batch_size, channels, height, width) - t: 当前时刻 - net: 用于预测噪音的神经网络 - simple_var: bool值, 决定是否使用简单的方差公式 """ # x_t.shape[0]: 获取 x_t 的第一个维度的大小, 即批量大小 batch_size n = x_t.shape[0] # [t] * n 是创建一个大小为n的一位张量, 里面内容全是t. dtype=torch.long指定t的数据类型是long. 例如,如果 t = 10 且 n = 4,那么 [t] * n 的结果是 [10, 10, 10, 10] # unsqueeze(1): 在时刻的第二个维度上增加一个维度,使其形状变为 (batch_size, 1). 例如,torch.tensor([10, 10, 10, 10], dtype=torch.long).unsqueeze(1) 的形状为 (4, 1) # 这样做是为了确保当前时刻的张量可以和图像的张量x_t的形状相兼容 t_tensor = torch.tensor([t] * n, dtype=torch.long).to(x_t.device).unsqueeze(1) # eps就是我们通过网络预测出来的噪音, 网络的输入是当前时刻t的图像分布x_t, 和当前时刻t_tensor, 输出预测的噪音 eps = net(x_t, t_tensor) # 如果是到了最后一步, 也就是要输出x_0的时候, 是不需要加噪音的 if t == 0: noise = 0 else: # 是否需要复杂的标准擦 if simple_var: var = self.betas[t] else: var = (1 - self.alpha_bars[t - 1]) / (1 - self.alpha_bars[t]) * self.betas[t] # 生成一个服从正态分布的, 形状和x_t一样的噪音. noise = torch.randn_like(x_t) noise *= torch.sqrt(var) # 最后的noise就是核心公式中后面的一小撮 # mean就是均值, 也就是核心公式的前面一大坨 mean = (x_t - (1 - self.alphas[t]) / torch.sqrt(1 - self.alpha_bars[t]) * eps) / torch.sqrt(self.alphas[t]) x_t = mean + noise # 合起来就是核心公式 # 返回更新后的第t时刻的图像分布x_t return x_t ``` **上面几个函数都是定义在DDPM里面的, 在nb中为了方便就放在外面的块了** ------- 下面是训练的部分, 可以放在另外一个.py文件写 # 5. 实现[[Diffusion#Training|训练算法]] 训练算法即为论文中提到的 $\text{Algorithm1 \; Training}$ 再回顾一遍伪代码。首先,我们要随机选取训练图片 $x_0$ ,随机生成当前要训练的时刻 $t$ ,以及随机生成一个生成 $x_t$ 的高斯噪声。之后,我们把 $x_t$ 和 $t$ 输入进神经网络,尝试预测噪声。最后,我们以预测噪声和实际噪声的均方误差为损失函数做梯度下降。 下面第一个块是实现了两个工具函数~~其中一个是上面的get_dataloader, 之前忘记了~~, 第二个块是正式实现训练算法 ```python def get_img_shape(): return (1, 28, 28) # 因为MNIST数据集的图像是单通道的28*28的形状, 所以返回(1, 28, 28) ``` ```python batch_size = 512 n_epochs = 100 def train(ddpm: DDPM, net, device, ckpt_path): """ - ddpm: 一个 DDPM 类的实例,包含了扩散模型的相关参数和方法 - net: 用于预测噪声的网络 - device: 设备(CPU 或 GPU) , 用于将数据和模型移动到相应的设备上 - ckpt_path: 模型权重保存的路径 - batch_size: 批量大小, 表示每次从数据集中加载多少个样本 - n_epochs: 训练的总轮数 """ # n_steps 就是公式中的 T # net 就是某个神经网络, 在之后会实现U-net n_steps = ddpm.n_steps dataloader = get_dataloader(batch_size) net = net.to(device) loss_fn = nn.MSEloss() # 使用均方误差损失(MSE Loss) optimizer = torch.optim.Adam(net.parameters(), 1e-3) # 使用 Adam 优化器,学习率为 1e-3 for e in range(n_epochs): # 随机选取训练图片x_0, 随机交给dataloader完成, 每轮遍历得到的x就是训练数据 # 注意: _ 表示忽略标签, 因为我们只关心输入图像x (因为MNIST数据集中都是有标签的) for x, _ in dataloader: # 获取当前批量的大小 cur_batch_size = x.shape[0] # 前面也提到了x是是个四维张量, 第一维是batch_size x = x.to(device) # t从[0, n_steps)随机取一个数, 并且t的形状是(current_batch_size,) # tip: 逗号后面没有值的括号表示创建一个元组,即使只有一个元素。这种写法确保了 size 参数是一个元组,即使它只有一个元素 # 例如,(current_batch_size,) 是一个包含一个元素的元组,而 (current_batch_size) 只是一个普通的括号表达式,不会创建元组 t = torch.randint(0, n_steps, (cur_batch_size, )).to(device) # 随机采样生成高斯噪声, 形状要和训练的图像形状一致 eps = torch.randnn_like(x).to(device) # 前面实现了的前向传播加噪音, 加到第t时刻, (t是之前随机生成的), 得到第t时刻的带噪音的图像分布x_t x_t = ddpm.sample_foward(x, t, eps) # 将第t时刻的图像分布x_t和当前时刻t输入到网络中, 得到输出记为eps_theta # t要reshape成(cur_batch_size, 1) eps_theta = net(x_t, t.reshape(cur_batch_size, 1)) # 计算网络预测的噪音eps_theta与真实随机采样得到的噪音eps之间的均方差损失 loss = loss_fn(eps_theta, eps) optimizer.zero_grand() # 清零梯度, 防止梯度累加 loss.backward() # 反向传播, 计算梯度 optimizer.step() # 更新模型参数 torch.save(net.state_dict(), ckpt_path) # 保存模型权重 ``` # 6. 实现去噪神经网络 (这里使用[[U-net]]) 在DDPM中,理论上我们可以用任意一种神经网络架构。但由于DDPM任务十分接近图像去噪任务,而U-Net又是去噪任务中最常见的网络架构,因此绝大多数DDPM都会使用基于U-Net的神经网络。 代码中大部分内容都和普通的U-Net无异。唯一要注意的地方就是**时序编码**。去噪网络的输入除了图像外,还有一个**时间戳t**。我们要考虑怎么把t的信息和输入图像信息融合起来。大部分人的做法是对t进行Transformer中的**位置编码**,把该编码加到图像的每一处上。 ### 6.1 实现位置编码 位置编码的目的是为每个位置生成一个唯一的向量,以便模型能够区分序列中的不同位置。为了实现这一点,通常使用正弦和余弦函数来生成这些向量。具体来说,对于每个位置 $pos$ 和每个特征维度 $i$ , 位置编码 $PE_{pos,i}$ 可以表示为: $$ PE_{pos,2i}=\sin\left(\frac{pos}{10000^{2i/d_{\mathrm{model}}}}\right)\\PE_{pos,2i+1}=\cos\left(\frac{pos}{10000^{2i/d_{\mathrm{model}}}}\right) $$ 这里 $pos$ 是**位置索引**, $i$ 是**特征索引**, $d_{model}$ 是模型的维度 ```python class PositionalEncoding(nn.Module): def __init__(self, max_seq_len: int, d_model: int): super().__init__() # 确保d_model是偶数, 因为后续的计算中要将d_model分成两部分, 分别计算正弦值和余弦值 # assert后面的条件如果为True, 程序继续执行, 如果为False, 就会抛出AssertionError并停止执行 assert d_model % 2 == 0 # 初始化 **位置编码矩阵** 形状为(max_seq_len, d_model), 里面填充0, pe用于存储位置编码 pe = torch.zeros(max_seq_len, d_model) # i_seq是从0到max_seq_len-1的线性等间距序列, 表示 **位置索引** i_seq = torch.linspace(0, max_seq_len - 1, max_seq_len) # j_seq是0到d_model-2的等间距序列, 表示 **特征索引** (这里将d_model分成了两部分) j_seq = torch.linspace(0, d_model - 2, d_model // 2) # 生成两个网格矩阵 # pos的形状为(max_seq_len,d_model // 2) 表示 **位置索引** # two_i的形状为(max_seq_len, d_model // 2) 表示 **特征索引** pos, two_i = torch.meshgrid(i_seq, j_seq) # 计算位置索引pos的正弦值, 公式为sin(pos / 10000^(two_i / d_model)) pe_2i = torch.sin(pos / 10000 ** (two_i / d_model)) # 计算位置索引的余弦值, cos(pos / 10000^(two_i / d_model)) pe_2i_1 = torch.cos(pos / 1000 ** (two_i / d_model)) # 合并正弦和余弦值 # 使用 torch.stack 将 pe_2i 和 pe_2i_1 沿着第三个维度堆叠起来,形成一个形状为 (max_seq_len, d_model // 2, 2) 的张量 # 再reshape成(max_seq_len, d_model)的形状, 得到最终的位置编码pe pe = torch.stack((pe_2i, pe_2i_1), 2).reshape(max_seq_len, d_model) # 创建嵌入层 # 先创建一个嵌入层, 用于将位置索引映射到位置编码向量 self.embedding = nn.Embedding(max_seq_len, d_model) # 将计算好的位置编码举证pe付给嵌入层的权重 self.embedding.weight.data = pe # **冻结嵌入层的权重, 使其再训练过程中不再更新** self.embedding.requires_grad_(False) def forward(self, t): return self.embedding(t) ``` ### 6.2 实现[[残差链接 (Residual Skip Connect)|残差连接块]] [[U-net]]通过[[残差链接 (Residual Skip Connect)]]将编码器中的高分辨率特征与解码器中的对应层直接连接,这有助于在[[卷积和反卷积#反卷积|上采样]]过程中恢复图像的细节信息,并且减少了训练过程中的信息丢失。 在示意图中编码器和解码器(就是U型的两端)之间的连线就是通过残差连接的 ```python class ResidualBlock(nn.Module): def __init__(self, in_c: int, out_c: int): super().__init__() self.conv1 = nn.Conv2d(in_c, out_c, 3, 1, 1) # 输入通道数为 in_c,输出通道数为 out_c,卷积核大小为 3,步幅为 1,填充为 1。这样可以保持输入和输出的尺寸不变。 self.bn1 = nn.BatchNorm2d(out_c) # 批量归一化层 self.activation1 = nn.ReLU() self.conv2 = nn.Conv2d(out_c, out_c, 3, 1, 1) self.bn2 = nn.BatchNorm2d(out_c) self.activation2 = nn.ReLU() # 如果输入通道数 in_c 不等于输出通道数 out_c,则需要通过一个 1x1 卷积层和批量归一化层来调整输入的通道数,使其与输出的通道数匹配。 if in_c != out_c: self.shortcut = nn.Sequential(nn.Conv2d(in_c, out_c, 1), nn.BatchNorm2d(out_c)) # 1x1 卷积层,用于改变通道数。 else: # 如果输入通道数等于输出通道数, 则使用nn.Identity(), 表示不做任何操作 self.shortcut = nn.Identity() def foward(self, input): x = self.conv1(input) x = self.bn1(x) x = self.activation1(x) x = self.conv2(x) x = self.bn2(x) x += self.shortcut(input) # shortcut所谓残差跳过就在于shortcut的输入是input, 如果输入输出通道相等也不会经过处理(不相等则会调整通道数使其匹配) 然后直接加到x上 x = self.activation2(x) return x ``` ### 6.3 实现卷积层 U-net是一个纯卷积的网络, 这里将需要用到的卷积层定义好, 之后在U-net类中进行拼接就可以了 这里的卷积层其实是**位置编码**, **残差连接**和**卷积层**的结合体, 然后形成了解码器/编码器中的一个块 ```python class ConvNet(nn.Module): def __init__(self, n_steps, intermediate_channels=[10, 20, 40], pe_dim=10, insert_t_to_all_layers=False): """ - n_steps: 论文中的T - intermediate_channels: 中间层的通道数列表 - pe_dim: 位置编码的维度 - insert_t_to_all_layers: 是否在所有的残差块插入t """ super.__init__() C, H, W = get_img_shape() # 这里使用的是MNIST数据集, 固定返回(1, 28, 28) self.pe = PositionalEncoding(n_steps, pe_dim) # 初始化t的位置编码 # 初始化位置编码线性层 self.pe_linears = nn.ModuleList() # pe_linears, 用于存储位置编码的线性层的ModuleList self.all_t = insert_t_to_all_layers # 表示是否在所有残差块中插入t的信息 if not insert_t_to_all_layers: # 如果不在所有的残差块中插入t的信息 self.pe_linears.append(nn.Linear(pe_dim, C)) # 就添加一个线性层, 将位置编码从pe_dim维度转换为C维度 (保证加t的编码和不加之后形状一样) # 初始化残差块 self.residual_blocks = nn.ModuleList() # 存储残差块的ModuleList prev_channel = C # 记录当前层的通道数 for channel in intermediate_channels: # 遍历中间层的通道数列表 self.residual_blocks.append(ResidualBlock(prev_channel, channel)) # 添加一个残差块, 输入通道数为prev_channel, 输出通道数为channel if insert_t_to_all_layers: # 如果在所有的残差块中插入时间t的信息, 则添加一个线性层, 将位置编码从pe_dim转换成prev_channel的维度 self.pe_linears.append(nn.Linear(pe_dim, prev_channel)) else: # 如果不加 self.pe_linears.append(None) # 那就不加 prev_channel = channel # 更新当前层的通道数 # 真正的卷积层 self.output_layer = nn.Conv2d(prev_channel, C, 3, 1, 1) # 输出层就是一个卷积层, 输入通道为prev_channel, 输出通道为C, 卷积核大小为3, 步幅为1, 填充为1 def forward(self, x, t): """输入x和t, 输出更新后的x""" n = t.shape[0] # 获取批量大小 t = self.pe(t) # 将t转换为位置编码 for m_x, m_t in zip(self.residual_blocks, self.pe_linears): # 同时遍历残差块和位置编码线性层 if m_t is not None: # 如果位置编码部位None pe = m_t(t).reshape(n, -1, 1, 1) # 就将位置编码转换为形状为(n, C, 1, 1)的张量 x = x + pe # 然后和图像的张量相加, 就算是将位置编码加入到图像的信息当中去了 x = m_x(x) # 通过当前的残差块处理输入x x = self.output_layer(x) # 通过输出层处理x return x # 返回最终的输出 ``` ### 6.4 实现U-net的基础块 ```python class UnetBlock(nn.Module): def __init__(self, shape, in_c, out_c, residual=False): """ - shape: 输入特征图的形状, 用于初始化LayerNorm - in_c: 输入通道数 - out_c: 输出通道数 - residual: 是否使用残差连接, 默认为False """ super().__init__() self.ln = nn.LayerNorm(shape) # 对输入的特征图进行层归一化 (可以稳定训练过程, 加快模型收敛和性能) self.conv1 = nn.Conv2d(in_c, out_c, 3, 1, 1) # 通过两个卷积层, 用于提取特征 self.conv2 = nn.Conv2d(out_c, out_c, 3, 1, 1) self.activation = nn.ReLU() self.residual = residual if residual: # 如果使用残差连接 if in_c == out_c: # 且当输入输出通道数相等, 就直接使用跳过连接 self.residual_conv = nn.Identity() else: # 不相等 self.residual_conv = nn.Conv2d(in_c, out_c, 1) # 如果不相等就通过一个1*1的卷积层, 将输入通道数调整为与输出通道数相同 def forward(self, x): out = self.ln(x) out = self.conv1(out) out = self.activation(out) out = self.conv2(out) if self.residual: out += self.residual_conv(x) out = self.activation(out) return out ``` ### 6.5 实现U-net主体 ```python class UNet(nn.Module): def __init__(self, n_steps, channels=[10, 20, 40, 80], pe_dim=10, residual=False) -> None: # 这里的参数也就是之前四个类所需要的所有参数的并集 """ - n_steps: 依旧是论文中的t - channels: 指定每层卷积的通道数, 控制模型的深度 (feature maps 的数量) - pe_dim: 位置编码的维度, 用于在图像中引入每个时刻的信息. 通常较小 - residual: 是否使用残差连接 """ super().__init__() # 获取输入图像的形状------------------- C, H, W = get_img_shape() # 获取输入图像的形状, 这里使用的是MNIST数据集, 这个自定义的函数默认返回(1, 28, 28) layers = len(channels) # 通过channels来计算网络的层数(U-net的层数由channels的长度决定) **(why)** Hs = [H] # 存储每一层的高值 Ws = [W] # 存储每一层的宽值 cH = H # cH和cW是当前层给高度和宽度, 初始化为输入尺寸 cW = W # 计算每层特征图的尺寸----------------- for _ in range(layers - 1): # 在编码器中(也就是U型的左边), 特征图尺寸会逐层减小, 通过将每层的高度和宽度除以2来逐步缩小尺寸 # 这样就模拟了 **逐层下采样** 的过程 # 通过整除操作, 表示特征图经过下采样(例如 **步幅为二的卷积** 或者池化操作) cH //= 2 cW //= 2 # 记录每层的特征图尺寸 Hs.append(cH) # 将更新后的cH和cW存入Hs和Ws Ws.append(cW) # 位置编码---------------------------- # 将t映射到一个固定维度(pe_dim)的向量, 用于在网络中输入时刻信息 self.pe = PositionalEncoding(n_steps, pe_dim) # 各种模块组件的列表初始化------------- self.encoders = nn.ModuleList() # 存储编码器部分的网络 self.decoders = nn.ModuleList() # 存储解码器部分的网络 self.pe_linears_en = nn.ModuleList() # 存储位置编码的线性变换层 self.pe_linears_de = nn.ModuleList() self.downs = nn.ModuleList() # down和up分别存储用于下采样和上采样的卷积层 self.ups = nn.ModuleList() # 构建编码器部分----------------------- prev_channel = C # C是输入图像的通道数 for channel, cH, cW in zip(channels[0:-1], Hs[0:-1], Ws[0:-1]): # **详细解释一下这个遍历?** # 在编码器中将t的位置编码映射到与当前通道数相匹配的维度 (通过这个层之后, t的位置编码的通道数就会变成图像的通道数C) self.pe_linears_en.append(nn.Sequential( nn.Linear(pe_dim, prev_channel), nn.ReLU(), nn.Linear(prev_channel, prev_channel))) # **那为什么又要第二个Linear呢?** # 编码器模块 self.encoders.append(nn.Sequential( # 第一层Unetblock: 将输入通道从prev_channel变成channel UnetBlock((prev_channel, cH, cW), prev_channel, channel, residual=residual), # 第二层Unetblock: 保持特征图的通道数不变 # **和上面同样的问题, 既然通道数不变为什么还要有这一层呢?** UnetBlock((channel, cH, cW), channel, channel, residual=residual))) # 这一层的encoder算完了就要进行下采样, 送进下一层的encoder self.downs.append(nn.Conv2d(channel, channel, 2, 2)) # 使用步幅为2的卷积, 将特征图尺寸减半, 同时保持通道数不变 prev_channel = channel # 更新通道数, 为下一层的encoder做准备 # 构建中间层---------------------------- self.pe_mid = nn.Linear(pe_dim, prev_channel) # 同样使用一个线性层来讲位置编码的通道数统一到当前图的通道数 channel = channels[-1] # **为啥?** # 用两个UnetBlock处理中间特征图 self.mid = nn.Sequential( # 通道数在第一层是增加了的 UnetBlock((prev_channel, Hs[-1], Ws[-1]), prev_channel, channel, residual=residual), # 在第二层通道数保持不变 UnetBlock((channel, Hs[-1], Ws[-1]), channel, channel, residual=residual)) # 构建解码器---------------------------- prev_channel = channel # 初始化解码器的输入通道, 初始化为编码器的最后一层的通道数 # **反向**遍历, 通过反向遍历可以逐步恢复图像的空间尺寸 for channel, cH, cW in zip(channels[-2::-1], Hs[-2::-1], Ws[-2::-1]): # 同编码器一样,解码器也需t的位置编码的线性变换,将t的位置编码映射到当前通道数, (通过这个层之后, t的位置编码的通道数就会变成图像的通道数prev_channel) self.pe_linears_de.append(nn.Linear(pe_dim, prev_channel)) # 上采样, 使用**反卷积**, 通道数从prev_channel变为channel, 同时图像尺寸变大一倍 self.ups.append(nn.ConvTranspose2d(prev_channel, channel, 2, 2)) # 解码器模块 self.decoders.append(nn.Sequential( # 第一层UnetBlock将输入特征图的通道数为 channel * 2 UnetBlock((channel * 2, cH, cW), channel * 2, channel, residual=residual), # 第二层输入特征图的通道数为 channel,并且维持通道数不变 UnetBlock((channel, cH, cW), channel, channel, residual=residual))) # 更新通道数, 为下一层解码器准备输入的通道数 prev_channel = channel # 输出层---------------------------------- # 最后的卷积层将解码器输出的特征图映射到最终的输出图像。 self.conv_out = nn.Conv2d(prev_channel, C, 3, 1, 1) # 输入通道数是解码器的最后一个输出通道数, 输出通道数是和图片相同的通道数 # 卷积核的大小为 3x3,步幅为 1,填充为 1,这意味着输出图像的尺寸和输入图像相同 def forward(self, x, t): n = t.shape[0] t = self.pe(t) encoder_outs = [] for pe_linear, encoder, down in zip (self.pe_linears_en, self.encoders, self.downs): pe = pe_linear(t).reshape(n, -1, 1, 1) x = encoder(x + pe) encoder_outs.append(x) x = down(x) pe = self.pe_mid(x + pe) for pe_linear, decoder, up, encoder_out in zip(self.pe_linears_de, self.decoders, self.ups, encoder_outs[::-1]): pe = pe_linear(t).reshape(n, -1, 1, 1) x = up(x) pad_x = encoder_out.shape[2] - x.shape[2] pad_y = encoder_out.shape[3] - x.shape[3] x = F.pad(x, (pad_x // 2, pad_x - pad_x // 2, pad_y // 2, pad_y - pad_y // 2)) x = torch.ca((encoder_out, x), dim=1) x - decoder(x + pe) x = self.conv_out(x) return x ```