【流匹配模型Flow Maching】流匹配模型入门理解(2)

发布时间:2026/10/1 11:36:48

【流匹配模型Flow Maching】流匹配模型入门理解(2) 目录前言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。
延伸阅读

更多相关文章

2026/10/1 11:36:48

细胞衰老的核心密码:NAD+平衡状态决定人体机能的存续时长

细胞衰老的核心密码:NAD平衡状态决定人体机能的存续时长人体的器官衰老、体能衰退、机能下滑,所有老化表现的底层核心密码,都指向同一个物质:NAD。它不是普通的营养物质,是调控细胞代谢、修复、更新、维稳的核心辅酶&a…

2026/10/1 11:31:48

Python校园外卖点餐系统:毕设选题与开发全流程解析

每年到毕业设计选题季,就有不少学弟学妹来问我:“学长,Python 的课设/毕设题目怎么选?想做个网站类的但又不想完全照搬网上的管理系统,有没有什么推荐?”说实话,校园外卖点餐系统这个题目我见过…

2026/10/1 11:31:48

字节码与机器码的区别:从JVM跨平台到JIT运行时转换

字节码和机器码到底有什么不一样?这个问题我几乎每隔一段时间就会遇到一次,尤其是新同事第一次用javap反汇编.class文件、或者第一次用objdump查看一个可执行文件的时候。我的第一句回答通常很短:机器码是 CPU 直接执行的二进制指令&#xff…

2026/10/1 12:31:51

基于Spring Boot的废旧物资预约回收系统:毕设项目全链路解析

每年帮学生复审毕业设计的Java项目,我都会遇到同一类题目:基于Spring Boot的业务管理系统。这次拿到的“瑞回宝废旧物资预约回收系统”比较有代表性——题面是一个环保回收业务,背后却串联了Spring Boot后端开发从项目初始化、数据建模、状态…

2026/10/1 12:31:51

【C++入门】编译链接模型 - 02 预处理把头文件怎样塞进源文件

博主介绍:程序喵大人 35 - 资深C/C/Rust/Android/iOS客户端开发10年大厂工作经验嵌入式/人工智能/自动驾驶/音视频/游戏开发入门级选手《C20高级编程》《C23高级编程》等多本书籍著译者更多原创精品文章,首发gzh,见文末👇&#x…

2026/10/1 12:31:51

MessageBox深度解析:从API参数到封装与高阶应用

做桌面客户端开发这些年,我发现被问得最多的问题不是高深算法,而是 MessageBox(消息提示框)这种看起来人人都会的组件。同事拿着一段弹窗代码来找我:“这个确定按钮点下去,整个界面卡住不动了,到…

2026/10/1 12:31:51

QiLink

QiLink是道息实验室发起的全球首个开源协同协议体系,是整套“道息-气链”双螺旋架构的核心技术工具层,完全由徐玉生原创定义,是连接顶层东方哲学理念与实体产业落地的核心枢纽‌。🔍 名称的专属原创内涵它的命名本身就是独创语义的…

2026/10/1 5:21:14

东莞市品牌网站建设报价常见报错与解决

东莞品牌网站建设报价单背后:一份保姆级建站教程避坑实录 网站做好了没人访问,这大概是很多老板最头疼的事。花了大几万做的品牌站,上线后流量惨淡,比路边摊还冷清。别急着骂外包公司,很多“东莞品牌网站建设报价”里藏着不少猫腻,比如用模板站冒充定制…

2026/9/29 21:48:03

如何划分训练/验证集:Spirula Studio五种eval_mode策略详解

如何划分训练/验证集:Spirula Studio五种eval_mode策略详解 【免费下载链接】spirula-studio Cross-vendor 3D Gaussian Splatting trainer - video to splat to mesh, Vulkan or CUDA. 项目地址: https://gitcode.com/GitHub_Trending/sp/spirula-studio Sp…

2026/10/1 10:48:55

SEO怎么推广速查手册新手避坑实战指南

SEO怎么推广速查手册新手避坑实战指南 模板网站太丑不够用?别急着加滤镜,那是治标不治本。很多老板盯着后台流量掉得眼红,却还在纠结首页Banner的圆角是不是3像素。这就像穿着西装去挖土,姿势不对,努力白费。我整理这份 速查手册…

还想了解更多?直接咨询顾问

免费诊断 + 免费方案 + 透明报价。

全国咨询热线400-8866-253
免费获取方案
☎咨询二维码 ☎ ↑