目录前言1. 模拟数据定义2. 构造训练路径3. 速度预测网络4. 训练5. 从噪声逐步生成6. 相同起点一步与多步比较前言之前已经介绍过DDMP以及流模型见如下四篇链接: 【扩散模型DDPM】扩散模型入门理解1 【扩散模型DDPM】扩散模型入门理解2【扩散模型DDPM】扩散模型入门理解3【流匹配模型Flow Maching】流匹配模型入门理解1现在看一下流模型的代码这个代码在DDPM代码的基础上比较好理解见扩散模型理解3。1. 模拟数据定义importmathimportnumpyasnpimportmatplotlib.pyplotaspltimporttorchfromtorchimportnnfromsklearn.datasetsimportmake_moons np.random.seed(42)torch.manual_seed(42)devicetorch.device(cuda:3iftorch.cuda.is_available()elsecpu)ifdevice.typecpu:torch.set_num_threads(min(4,torch.get_num_threads()))print(device:,device)points,_make_moons(n_samples10000,noise0.05,random_state42)points(points-points.mean(axis0))/points.std(axis0)datatorch.tensor(points,dtypetorch.float32,devicedevice)defplot_points(ax,points,title):ifisinstance(points,torch.Tensor):pointspoints.detach().cpu().numpy()ax.scatter(points[:,0],points[:,1],s3,alpha0.4)ax.set_title(title)ax.set_xlim(-3.5,3.5)ax.set_ylim(-3.5,3.5)ax.set_aspect(equal)fig,axplt.subplots(figsize(4,4))plot_points(ax,data[:1500],Real data)plt.tight_layout()plt.show()2. 构造训练路径固定同一批起点和终点查看不同t tt的插值这不是训练好的模型生成的轨迹。x t ( 1 − t ) z t x d a t a x_t(1-t)zt x_{\mathrm{data}}xt​(1−t)ztxdata​这里展示的是人为构造的训练插值还不是网络生成的结果。definterpolate(x_data,t,noise):tt[:,None]# [batch] - [batch, 1]同一个 t 用于两个坐标return(1-t)*noiset*x_data x_demodata[:1500]noise_demotorch.randn_like(x_demo)fig,axesplt.subplots(1,5,figsize(15,3))forax,timeinzip(axes,[0.0,0.25,0.5,0.75,1.0]):ttorch.full((len(x_demo),),time,devicedevice)plot_points(ax,interpolate(x_demo,t,noise_demo),ft {time:.2f})plt.tight_layout()plt.show()3. 速度预测网络与之前 NoisePredictor 的结构相同t tt已位于[ 0 , 1 ] [0,1][0,1]不再除以T TT(之前 DDPM 的时间是整数编号现在 Flow Matching 的时间已经是一个 01 之间的小数。)。输出两个坐标方向的速度。classVelocityPredictor(nn.Module):def__init__(self):super().__init__()self.register_buffer(freq,torch.arange(1,9).float()*math.pi)self.netnn.Sequential(nn.Linear(18,128),nn.SiLU(),nn.Linear(128,128),nn.SiLU(),nn.Linear(128,128),nn.SiLU(),nn.Linear(128,2),)defforward(self,x_t,t):phaset.float()[:,None]*self.freq[None,:]time_embeddingtorch.cat([phase.sin(),phase.cos()],dim1)returnself.net(torch.cat([x_t,time_embedding],dim1))modelVelocityPredictor().to(device)optimizertorch.optim.Adam(model.parameters(),lr1e-3)4. 训练每个数据点随机抽取一个连续时间随机噪声与数据独立配对。这个地方注意噪声跟原始数据是独立配对的也就是说噪声和数据没有一一对应的关系是随机的这个地方可以用OT最有传输提前配对后面再处理这是另一种模型方式。路径对时间求导得到目标速度u t d x t d t x d a t a − z u_t\frac{dx_t}{dt}x_{\mathrm{data}}-zut​dtdxt​​xdata​−z因此损失为∥ v θ ( x t , t ) − ( x d a t a − z ) ∥ 2 \|v_\theta(x_t,t)-(x_{data}-z)\|^2∥vθ​(xt​,t)−(xdata​−z)∥2。不同端点可能给出冲突的速度标签所以不要求训练损失降到零。batch_size256train_steps4000loss_history[]model.train()forstepinrange(1,train_steps1):x_datadata[torch.randint(len(data),(batch_size,),devicedevice)]ttorch.rand(batch_size,devicedevice)noisetorch.randn_like(x_data)x_tinterpolate(x_data,t,noise)target_velocityx_data-noise predicted_velocitymodel(x_t,t)loss(predicted_velocity-target_velocity).square().mean()optimizer.zero_grad()loss.backward()optimizer.step()loss_history.append(loss.item())ifstep%5000:print(fstep{step:4d}| mean loss{np.mean(loss_history[-500:]):.4f})plt.figure(figsize(6,3))plt.plot(loss_history,alpha0.3,labelBatch loss)window100smoothednp.convolve(loss_history,np.ones(window)/window,modevalid)plt.plot(np.arange(window,len(loss_history)1),smoothed,label100-step mean)plt.xlabel(Training step)plt.ylabel(Velocity MSE)plt.legend()plt.tight_layout()plt.show()5. 从噪声逐步生成训练完成后我们从纯噪声出发通过100 步 Euler 采样逐步生成数据。这里每一步都重新调用同一个网络v θ v_\thetavθ​并且只在初始化时抽取一次噪声后续每一步不再额外加噪。使用最简单的 Euler 更新公式x t Δ t x t Δ t v θ ( x t , t ) \boxed{x_{t\Delta t}x_t\Delta t\,v_\theta(x_t,t)}xtΔt​xt​Δtvθ​(xt​,t)​这里设置生成过程走100 步所以时间步长Δ t 0.01 \Delta t0.01Δt0.01。注意这 100 步是采样精度的设置不是训练中离散时间步的数量。model.eval()n_steps100# 采样步数dt1.0/n_steps# 时间步长 Δt 0.01# 只在初始化时抽取一次噪声ztorch.randn(1500,2,devicedevice)x_tz.clone()withtorch.no_grad():foriinrange(n_steps):ttorch.full((len(x_t),),i*dt,devicedevice)vmodel(x_t,t)# 每一步重新调用同一个网络x_tx_tdt*v# Euler 更新fig,axplt.subplots(figsize(4,4))plot_points(ax,x_t,Generated (100-step Euler))plt.tight_layout()plt.show()从上面的代码可以看到整个生成过程就是一个确定性的 ODE 积分给定初始噪声z zz沿着网络预测的速度场v θ v_\thetavθ​走 100 步最终得到近似数据分布的样本。这个地方非常好理解比DDPM的反向高斯采样好理解多了DDPM需要推导公式再去更新流模型的更新就是当前位置加上时间乘以速度这个地方个人理解非常好非常清爽。下面的部分也说明了流模型不一定生成的快还是得一步一步的但是流模型的这种形式更加清晰简洁。现在也是非常火这个模型。6. 相同起点一步与多步比较一步使用d t 1 d_t1dt​1直线训练不保证学到的生成流能用一步准确求解。withtorch.no_grad():t_zerotorch.zeros(len(initial_noise),devicedevice)one_stepinitial_noisemodel(initial_noise,t_zero)fig,axesplt.subplots(1,3,figsize(12,4))plot_points(axes[0],data[:1500],Real data)plot_points(axes[1],one_step,1 Euler step)plot_points(axes[2],generated,100 Euler steps)plt.tight_layout()plt.show()条件流模型见链接: 【流匹配模型Flow Maching】流匹配模型入门理解3。