CodeByMySelf-Diffusion

所有之前:CodeByMySelf 系列之 Diffusion Models 的实现记录,变体不会在这篇里更新。 Diffusion Models 原理 具体的一些公式推导请看李宏毅老师的视频: 李宏毅Diffusion 基本概念 (以图像生成为例)首先Diffusion的基本运作过程是,首先Sample出一个全是噪声的图片,然后经过一步一步的Denoise最终得到图片,Denoise的步数是事先确定好的。 在Diffusion中,深度学习模型的唯一作用就是预测噪声,其他过程均由数学公式推理得到 进一步说,不同情况的噪声需要的处理肯定是不一样的,一个Denoise模型肯定无法做到处理不同的情况的噪声,因此Denoise模型需要一个“噪声严重程度”作为输入,在Diffusion中称为时间步。 Denoise内部的工作是,接受带Noise的图片,对噪声进行预测,然后使用带Noise的图片减去预测噪声那么就可以得到。 我们清楚Noise predictor模型的输入输出之后,该怎么训练这个模型呢?已知模型的输入-输出对是Noise图片加上“噪声程度”-Noise数据对,我们需要人为创造这个数据对出来,具体来说,从image list中拿一张图片,然后从GasuionN中采样出一张纯噪声的图片,加了一定的步数之后,就得到了很多数据对(Noise图片、时间步–噪声),这个过程称为Diffusion Process。 现在,已经可以从一个纯噪声得到一张图片,但是我们不能根据噪声随机生成图片,所以我们需要把文字考虑进来,让模型考虑进文字之后再预测噪声。 完整算法 首先是Training的过程,从immage list中取出一张干净的图,之后从1-T中随机取一个值作为时间步,然后从Gaussion中取样出一个noise,之后就是一步一步的梯度更新训练模型。这种方法和我们的想法有些不一样,我们想像中noise是一步一步加进去的,但是在实际中是一下子直接加进去的。 具体的加噪公式推导过程如下: 然后是推理过程,也就是Sampling这个过程。首先,采样得到一个纯噪声图片X_T,然后进行T次公式中的循环,$\epsilon_\theta$表示预测噪声的模型。 但是为什么要加这个z噪声:每次都要重新采样出一个噪声,但是只有当t>1的时候.这个z的使用和生成语句中一个问题类似:为什么每次总要SAMPLE而不是直接取Mean(概率最大)?在The curious case of nerual text degeneration论文中有一个分析,当直接取Mean时候,会一直重复相同的回答和出现跳帧现象,并且他们分析得到人在写文章的时候并不会一定会选概率最大的词。 简单代码实现 Diffusion已经有了许多变体,下面是一个简单的实现,变体不会在这篇里更新。 import torch from torch import nn import math import torch.nn.functional as F import numpy as np import time class WeightedLoss(nn.Module): def __init__(self): super(WeightedLoss,self).__init__() def forward(self,pred,target,weighted=1.0): loss = self._loss(pred,target) weighted_loss = (loss * weighted).mean() # mean by batch return weighted_loss class L1Loss(WeightedLoss): def _loss(self,pred,target): return torch.abs(pred - target) class L2Loss(WeightedLoss): def _loss(self,pred,target): return (pred - target) ** 2 Losses = { 'l1':L1Loss, 'l2':L2Loss } class SinusoidalPosEmb(nn.Module): def __init__(self,dim): super(SinusoidalPosEmb,self).__init__() self.dim = dim def forward(self,t): ''' 扩散模型时间步的位置编码公式和Transformer中的位置编码公式是类似的。 ''' device = t.device half_dim = self.dim // 2 emb = math.log(10000) / (half_dim - 1) emb = torch.exp(torch.arange(half_dim,device=device) * -emb) emb = t[:,None] * emb[None,:] emb = torch.cat((emb.sin(),emb.cos()),dim=-1) return emb def extract(a,t,x_shape): ''' a: (T,) t: (B,) x_shape: (B,...) ''' batch_size = t.shape[0] out = a.gather(-1,t) # (B,) return out.view(batch_size, *((1,) * (len(x_shape) - 1))) # (B,1,1,1...) class MLP(nn.Module): ''' 这是一个简单的多层感知机(MLP)模型,结合了时间步的位置编码,用于处理状态和动作的输入(强化学习) ''' def __init__(self,state_dim,action_dim,hidden_dim,device,t_dim): super(MLP,self).__init__() self.device = device self.t_dim = t_dim self.a_tim = action_dim self.time_mlp = nn.Sequential( SinusoidalPosEmb(t_dim), nn.Linear(t_dim,t_dim*2), nn.Mish(), nn.Linear(t_dim*2,t_dim) ) input_dim = state_dim + action_dim + t_dim self.mid_layer = nn.Sequential( nn.Linear(input_dim,hidden_dim), nn.Mish(), nn.Linear(hidden_dim,hidden_dim), nn.Mish(), nn.Linear(hidden_dim,hidden_dim), nn.Mish(), ) self.f_layer = nn.Linear(hidden_dim,action_dim) def init_weights(self): ''' 初始化模型权重,好的权重初始化可以帮助模型更快收敛,提升训练效果。 ''' for m in self.modules(): if isinstance(m,nn.Linear): nn.init.xavier_normal_(m.weight) if m.bias is not None: nn.init.zeros_(m.bias) def forward(self,state,action,t): t_emb = self.time_mlp(t) # (batch,t_dim) x = torch.cat((state,action,t_emb),dim=-1) # (batch,state_dim+action_dim+t_dim) x = self.mid_layer(x) # (batch,hidden_dim) action_delta = self.f_layer(x) # (batch,action_dim) return action_delta class Diffusion(nn.Module): def __init__(self,loss_type,beta_schedule='linear',clip_denoise=True,predict_epsilon=True,**kwargs): ''' loss_type: 损失函数类型 beta_schedule: beta(时间)的离散方式,常见的有'linear'(线性)和'cosine'(余弦)。 clip_denoise: 是否在去噪过程中裁剪输出,以防止数值过大。 ''' super(Diffusion,self).__init__() self.device = kwargs.get('device','cpu') self.t_dim = kwargs.get('t_dim',16) self.state_dim = kwargs.get('state_dim',10) self.action_dim = kwargs.get('action_dim',2) self.hidden_dim = kwargs.get('hidden_dim',256) self.loss_type = loss_type self.beta_schedule = beta_schedule self.clip_denoise = clip_denoise self.T = kwargs.get('timesteps',1000) self.model = MLP(self.state_dim,self.action_dim,self.hidden_dim,self.device,self.t_dim).to(self.device) self.kwargs = kwargs self.model.init_weights() if beta_schedule == 'linear': betas = torch.linspace(1e-4,0.02,self.kwargs.get('timesteps', 1000),dtype=torch.float32) alphas = 1. - betas alphas_cumprod = torch.cumprod(alphas,dim=0) # [1,2,3] -> [1,1*2,1*2*3] alphas_cumprod_prev = torch.cat((torch.tensor([1.],dtype=torch.float32),alphas_cumprod[:-1]),dim=0) self.register_buffer('betas',betas) self.register_buffer('alphas',alphas) self.register_buffer('alphas_cumprod',alphas_cumprod) self.register_buffer('alphas_cumprod_prev',alphas_cumprod_prev) self.register_buffer('sqrt_alphas_cumprod',torch.sqrt(alphas_cumprod)) # (前向过程) self.register_buffer('sqrt_alphas_cumprob',torch.sqrt(alphas_cumprod)) self.register_buffer('sqrt_one_minus_alphas_cumprod',torch.sqrt(1. - alphas_cumprod)) # (反向过程) posterior_variance = betas * (1. - alphas_cumprod_prev) / (1. - alphas_cumprod) self.register_buffer('posterior_variance',posterior_variance) # 用于从xt求得x0 self.register_buffer('sqrt_recip_alphas_cumprod',torch.sqrt(1. / alphas_cumprod)) self.register_buffer('sqrt_recipm1_alphas_cumprod',torch.sqrt(1. / alphas_cumprod - 1)) self.register_buffer('posterior_mean_coef1',betas *torch.sqrt(alphas_cumprod_prev) / (1. - alphas_cumprod)) self.register_buffer('posterior_mean_coef2', (1. - alphas_cumprod_prev) * torch.sqrt(alphas) / (1. - alphas_cumprod)) self.loss_func = Losses[loss_type]() def forward(self,state,*args,**kwargs): return self.sample(state,*args,**kwargs) def sample(self,state,*args,**kwargs): batch_size = state.shape[0] shape = [batch_size,self.action_dim] action = self.p_sample_loop(state,shape,*args,**kwargs) return action.clamp(-1.,1.) def q_posterior(self,x0,xt,t): ''' 计算后验分布q(x_{t-1}|x_t,x_0)的均值和方差 ''' posterior_mean = (extract(self.posterior_mean_coef1,t,x0.shape) * x0 + extract(self.posterior_mean_coef2,t,x0.shape) * xt) posterior_variance = extract(self.posterior_variance,t,x0.shape) posterior_log_variance = torch.log(posterior_variance.clamp(min=1e-20)) return posterior_mean,posterior_variance,posterior_log_variance def p_sample_loop(self,state,shape,*args,**kwargs): device = self.kwargs.get('device','cpu') batch_size = shape[0] x = torch.randn(shape,device=device,requires_grad=False) # 标准DDPM噪声不需要梯度 for i in reversed(range(0,self.kwargs.get('timesteps',1000))): t = torch.full((batch_size,),i,dtype=torch.long,device=device) x = self.p_sample(x,t,state) return x def predict_x0_from_noise(self,x,t,noise): ''' 根据公式计算x0 ''' return (extract(self.sqrt_recip_alphas_cumprod,t,x.shape) * x - extract(self.sqrt_recipm1_alphas_cumprod,t,x.shape) * noise) def p_mean_variance(self,x,t,state): pred_noise = self.model(state,x,t) x_0 = self.predict_x0_from_noise(x,t,pred_noise) x_0.clamp_(-1.,1.) if self.clip_denoise else x_0 model_mean,posterior_variance,posterior_log_variance = self.q_posterior(x_0,x,t) return model_mean,posterior_log_variance def p_sample(self,x,t,state): ''' x: 当前的噪声状态xt t: 当前的时间步t state: 状态 ''' model_mean,model_log_variance = self.p_mean_variance(x,t,state) noise = torch.randn_like(x) if t.sum() > 0 else 0. # 如果t=0,则不添加噪声 return model_mean + torch.exp(0.5 * model_log_variance) * noise def q_sample(self,x0,t,noise): sample = (extract(self.sqrt_alphas_cumprod,t,x0.shape) * x0 + extract(self.sqrt_one_minus_alphas_cumprod,t,x0.shape) * noise) return sample def p_losses(self,x0,state,t,weights): noise = torch.randn_like(x0) x_noisy = self.q_sample(x0,t,noise) x_recon = self.model(state,x_noisy,t) loss = self.loss_func(x_recon,noise,weights) return loss def loss(self,x,state,weights=1.0): batch_size = x.shape[0] t = torch.randint(0,self.T,(batch_size,),device=self.device).long() return self.p_losses(x,state,t,weights) if __name__ == '__main__': x = torch.randn(256,2) state = torch.randn(256,11) model = Diffusion(loss_type='l2',beta_schedule='linear',clip_denoise=True,predict_epsilon=True, device='cpu',t_dim=16,state_dim=11,action_dim=2,hidden_dim=256,timesteps=100) result = model(state) print(result) loss = model.loss(x,state) print(f"loss: {loss.item()}")

September 24, 2025 · 3 min