news 2026/9/16 14:39:21

Sana × Cosmos-RL 后训练实战:图像与视频扩散模型的 SFT / RL 配置、训练与本地实现解析

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Sana × Cosmos-RL 后训练实战:图像与视频扩散模型的 SFT / RL 配置、训练与本地实现解析

Sana × Cosmos-RL 后训练实战:图像与视频扩散模型的 SFT / RL 配置、训练与本地实现解析

【免费下载链接】SanaSANA: Efficient High-Resolution Image Synthesis with Linear Diffusion Transformer项目地址: https://gitcode.com/GitHub_Trending/sana/Sana

本指南以 docs/sana_cosmos_rl.md 为骨架,结合仓库内 Sol-RL 后训练模块(configs/sol_rl、train_scripts/sol_rl、diffusion/post_training)进行源码级扩充。读者读完本文后,将掌握:SANA 与 Cosmos-RL 联合后训练的算法全景(SFT / LoRA / DiffusionNFT / Flow-GRPO)、完整配置预设清单与异步奖励服务部署方式、可直接运行的 SFT 与 RL 训练命令,以及仓库内 DiffusionNFT(BON + Preview Rollout)参考实现的底层原理与关键参数。

背景:当高效扩散模型遇上通用 RL 基础设施

SANA是面向高分辨率图像与视频生成的高效代码库(线性注意力 DiT 架构),而Cosmos-RL是 NVIDIA 推出的灵活、可扩展的强化学习框架。两者通过官方合作打通,为 SANA 提供了完整的后训练(Post-Training)基础设施,覆盖:

  • SFT(监督微调):图像与视频的 Full Fine-Tuning 与 LoRA 微调;
  • RL(强化学习):如DiffusionNFTFlow-GRPO,支持图像与视频,搭配异步奖励服务(async reward service)可配置数据集

这条链路的意义在于:在预训练模型已经具备强大生成能力的基础上,通过后训练让模型对齐人类偏好(审美、文本跟随、指令遵循等),这是从"能生成"走向"生成得好"的关键一步。

支持的算法与特性

Cosmos-RL 面向不同模态提供了一组 SOTA 算法:

模态算法说明
LLMGRPO、DAPO语言模型主流 RL 算法(GRPO 为分组相对策略优化)
扩散 / 世界模型FlowGRPO、DDRL、DiffusionNFT面向扩散过程的 RL 算法

SANA 是 Cosmos-RL 的原生支持对象(natively supported),这意味着 Cosmos-RL 已内置 SANA 的模型接入、采样器与数据流适配。完整的算法细节请参阅 Cosmos-RL 官方文档及"扩散模型后训练"(post-training of diffusion models)专题。

配置体系:预设(Presets)与参数入口

配置文件位置:Cosmos-RL 仓库的configs/sana目录下维护了 SANA 的全部预设配置(本仓库未内置这些.toml,因为其托管在 Cosmos-RL 侧;下方命令中的./configs/sana/...均指 Cosmos-RL 仓库内路径)。预设清单如下:

任务图像视频
SFTsana-image-sftsana-image-sft-lorasana-video-sftsana-video-sft-lora
RLsana-image-nftsana-video-nft

这些预设统一遵循 Cosmos-RL 的配置规范,参数细节见其 Configuration 文档。命名规律为<模态>-<任务>-<变体>,其中-lora后缀表示只训练 LoRA 适配器而非全量权重,-nft后缀表示走 DiffusionNFT 强化学习流程。

本仓库的对应本地实现:如果你希望在本仓库内直接体验与 Cosmos-RL 同源的 DiffusionNFT 式 RL 训练(无需外部框架),可参考 Sol-RL 后训练模块。其配置命名遵循<model>_<family>_<reward>模式,例如:

  • sana_diffusionnft_pickscore
  • sana_compile_hpsv2
  • sana_sol_rl_imagereward

详见 configs/sol_rl/sana.py(另有 flux1.py、sd3.py 对应 FLUX.1 与 SD3.5-L)。

奖励服务(Reward Service)

Cosmos-RL 推荐使用独立的异步奖励服务(async reward service)来并行计算奖励,训练器与奖励服务解耦。训练侧需要配置三个环境变量:

环境变量作用
REMOTE_REWARD_TOKEN奖励服务的鉴权令牌
REMOTE_REWARD_ENQUEUE_URL奖励任务入队(enqueue)地址
REMOTE_REWARD_FETCH_URL奖励结果拉取(fetch)地址

奖励服务的部署细节请参见 Cosmos-RL 仓库的reward_service/README.md。这种异步设计使得 rollout 采样的奖励计算不阻塞训练主循环,是支撑大规模 RL 吞吐的关键。本仓库的本地实现同样体现了"采样与打分手解耦"的思想:在 train_scripts/sol_rl/train_sana.py 中,rollout 产生的样本通过ThreadPoolExecutor异步提交奖励计算(executor.submit(reward_fn, ...)),训练循环在需要时才result()取回奖励,与 Cosmos-RL 的异步奖励服务在架构意图上一致。

训练实操:SFT 与 RL 的命令与流程

SFT(以图像 LoRA 为例)

cosmos-rl --config ./configs/sana/sana-image-sft-lora.toml cosmos_rl.tools.dataset.diffusers_dataset

要点:

  • --config指向预设 TOML;
  • 末尾的cosmos_rl.tools.dataset.diffusers_dataset指定使用 diffusers 兼容的数据集工具加载本地数据;
  • 若要做全量微调(非 LoRA),改用sana-image-sft/sana-video-sft预设。

RL(图像 DiffusionNFT)

cosmos-rl --config ./configs/sana/sana-image-nft.toml cosmos_rl.tools.dataset.diffusion_nft

要点:

  • cosmos_rl.tools.dataset.diffusion_nft是 DiffusionNFT 专用数据集工具;
  • 对应视频任务替换为sana-video-nft预设。

数据集准备

  • SFT:使用本地目录,目录内包含*.json(提示词/元数据)+*.jpg(图像)/*.mp4(视频);
  • RL 图像:内置了 pickscore、ocr、geneval 等常用数据集;
  • RL 视频:支持过滤后的 VidProM 数据集。

自定义数据集可在cosmos_rl/tools/dataset/diffusion_nft.py基础上扩展,并参考 Cosmos-RL 的 Customization 指南。

仓库内参考实现:DiffusionNFT(BON + Preview Rollout)原理剖析

与本仓库 Sol-RL 模块直接对应的是DiffusionNFT算法族。下面以源码为依据说明其运行机制,帮助你理解 Cosmos-RL 中*-nft预设背后的通用范式。

配置族与 rollout 形状

configs/sol_rl/sana.py 定义了五个配置族,对应不同的 rollout 成本与量化策略:

Family含义Rollout 形状(in-N / best-of-M)TE / NVFP4
diffusionnftPEFT 推理基线24-in-24
naive_scalingPEFT 暴力扩展24-in-96
compileBF16 编译加速24-in-96
naive_quant直接 NVFP4 量化 rollout24-in-96
sol_rl两阶段解耦 rollout24-in-96

各族对应的模型角色:diffusionnft/naive_scalingpreview_model="peft"fullrollout_model="peft"compilefullrollout_model="compile"naive_quantfullrollout_model="compile_nvfp4"sol_rl则为preview_step=6preview_model="compile_nvfp4"fullrollout_model="compile"

推荐首次运行的配置:sana_diffusionnft_pickscoresd3_diffusionnft_pickscoreflux1_diffusionnft_pickscore(均为 24-in-24 的 PEFT 基线,无需编译与量化即可验证全链路)。

两阶段解耦:FP4 探索 + BF16 再生

sol_rl族是仓库中吞吐优化最激进的方案,其核心是把"广泛探索"与"精细再生"分离(见 train_scripts/sol_rl/train_sana.py 中_rollout_for_one_prompt的实现):

  1. Stage 1(FP4 探索):用 NVFP4 量化的 compile 草稿模型(draft model)以 6 步(preview_step=6)采样 96 张候选图,并先用奖励函数打分(only_strict=True的严格模式);
  2. Stage 2(BF16 再生):按stage1_select_mode="best_worst"从草稿池中选出候选种子,换用 BF16 compile 模型以完整 10 步(rollout_sample_num_steps=10)重新生成,得到高质量最终样本用于训练。

源码中_select_inference_transformer根据modecompile_nvfp4/compile/peft)在多个推理模型副本间切换,这正是"草稿-再生"两阶段的关键调度点。选出的样本经过select_indices_by_mode二次筛选(best_of_n)后进入训练。

启动脚本与模型权重管理

单节点 8 卡启动脚本 train_scripts/sol_rl/run_sana_single_node_8gpu.sh 的用法:

CONFIG_SPEC=configs/sol_rl/sana.py:sana_diffusionnft_pickscore \ bash train_scripts/sol_rl/run_sana_single_node_8gpu.sh

脚本要点:

  • 默认CONFIG_SPEC=configs/sol_rl/sana.py:sana_diffusionnft_pickscore,即不传环境变量也能直接跑通基线;
  • 可覆盖的环境变量包括NPROC_PER_NODE(默认 8)、CUDA_VISIBLE_DEVICESMASTER_PORTNATIVE_CONFIG等;
  • SANA 原生权重路径默认output/pretrained_models/SANA_LinearFFN.pth,缺失时自动从hf://yitongl/SANA_LinearFFN/SANA_LinearFFN.pth下载(经sana.tools.hf_download_or_fpath解析);
  • 最终以torchrun --standalone --nproc_per_node拉起 train_sana.py,并透传--config与可选的--native_config

NVFP4 前置条件:若要走*_naive_quant_**_sol_rl_*路径,需要与torchrun使用同一 Python 解释器安装transformer-engine

python -m pip install --no-build-isolation "transformer-engine[pytorch]"

否则训练脚本会在构建 NVFP4 推理模型时报错(train_sana.py中显式检查_HAS_TE)。

训练循环与损失:DiffusionNFT 的源码级视角

train_sana.py的训练主循环体现了 DiffusionNFT 的核心训练信号构造:

  • 采样阶段:每个 prompt 按per_prompt_iter_num × rollout_batch_size生成候选,全部采样在torch.no_grad()下完成,同时保存latents_clean(最终干净潜变量)与完整时间步轨迹(timesteps/next_timesteps),供后续按步训练;
  • 优势估计:默认启用PerPromptStatTracker(diffusion/post_training/stat_tracking.py)做 per-prompt 的奖励统计归一化(global_std=True),并计算 zero-std 比例等诊断指标;奖励向量会在时间步维度复制展开;
  • 策略损失:对每个时间步构造正/负预测(positive_predictionimplicit_negative_prediction,由config.beta控制插值),以采样奖励加权后的 MSE 形式(r * positive_loss + (1-r) * negative_loss)构成policy_loss
  • KL 约束:叠加当前模型与参考模型(禁用适配器)预测差异的kl_div_loss,防止策略漂移;
  • EMA 与老模型更新:可选的EMAModuleWrapper维护 EMA 权重;每轮末尾按return_decay(global_step, decay_type)计算的衰减因子将训练权重指数滑动合并进 "old" 适配器,供下一轮 rollout 使用。

奖励模型:在线打分器实现

diffusion/post_training/rewards.py 提供了与配置族一一对应的在线奖励实现:

配置后缀奖励模型权重获取方式
pickscorelaion/CLIP-ViT-H-14-laion2B-s32B-b79K+yuvalkirstain/PickScore_v1首次使用自动下载
clipscoreopenai/clip-vit-large-patch14首次使用自动下载
hpsv2open_clip_pytorch_model.bin+HPS_v2.1_compressed.pt需手动放入reward_ckpts/
imagerewardImageReward-v1.0首次使用自动下载

其中 HPSv2 需要手动准备本地权重:

mkdir -p reward_ckpts cd reward_ckpts wget https://huggingface.co/laion/CLIP-ViT-H-14-laion2B-s32B-b79K/resolve/main/open_clip_pytorch_model.bin wget https://huggingface.co/xswu/HPSv2/resolve/main/HPS_v2.1_compressed.pt cd ..

multi_score调度器支持多奖励加权组合:配置如{"pickscore": 1.0}表示单一奖励,权重非 1 时可构造组合奖励,最终输出avg汇总分数供训练与评估使用。此外仓库还内置了 ImageReward 对 transformers >= 5.0 的兼容补丁(_patch_imagereward_compat)。

关键参数速查(仓库 base 配置)

configs/sol_rl/base.py 定义了所有 Sol-RL 训练共用的默认参数,理解它们有助于你调整 Cosmos-RL 或本仓库配置:

  • 训练learning_rate=3e-4batch_size=1(每 GPU)、gradient_accumulation_steps=1max_grad_norm=0.002num_inner_epochs=1adv_clip_max=5timestep_fraction=0.6beta=0.0001mixed_precision="bf16"
  • 采样num_steps=40eval_num_steps=40num_image_per_prompt=24best_of_n=24noise_level=1.0test_batch_size=1
  • LoRAlora_rank=32lora_alpha=64lora_init_weights=True(SANA 具体目标模块见 configs/sol_rl/sana.py 的lora_target_modulesattn.qkvattn.projcross_attn.q_linearcross_attn.kv_linearcross_attn.proj);
  • Rollout 相关rollout_sample_num_steps=10preview_step=0(默认关闭两阶段)、rollout_sample_guidance_scale=1.0compile_mode="max-autotune-no-cudagraphs"
  • NVFP4nvfp4_skip_modules(SANA 跳过t_embeddery_embedderx_embedderfinal_layerattn.qkv等敏感层)与nvfp4_min_dim=2240控制量化粒度。

SANA 原生模型结构由 configs/sol_rl/Sana1.0_1600M_linear.yaml 描述:SanaMSLinearFFN_1600M_P1_D20线性注意力架构、attn_type: linearffn_type: glumbconv_linear、Gemma-2-2B-it 文本编码器、DC-AE(dc-ae-f32c32-sana-1.1-diffusers)VAE、线性 flow 调度。RL 训练时通过pyrallis解析该 YAML 构建原生 Transformer,并复用 diffusersSanaPipeline的文本编码器与 VAE,仅在推理时切换训练好的 LoRA / 编译 / 量化模型副本。

注意事项与适用前提

  • Cosmos-RL 预设的归属configs/sana/*.tomlreward_service/位于 Cosmos-RL 仓库;本文中标注为仓库本地实现的内容(configs/sol_rltrain_scripts/sol_rldiffusion/post_training)可在本仓库内直接复现 DiffusionNFT 风格的 RL 训练;
  • 硬件前提:NVFP4 路径依赖 Transformer Engine 且需与torchrun解释器一致;compile族需要较新的 PyTorch 与 CUDA 环境;
  • 数据集格式:SFT 请严格遵循"本地目录 +*.json+*.jpg/*.mp4"的约定;RL 图像数据集(pickscore / ocr / geneval 等)在本仓库对应diffusion/post_training/dataset/下的 drawbench / geneval / ocr / pickscore 子目录,训练时由build_datasets_and_loaders统一装配;
  • 监控:训练通过 wandb 记录reward/meanreward/maxzero_std_ratiopolicy_losskl_div_lossgrad_norm等指标,便于观察奖励坍缩与策略漂移。

延伸阅读

  • 仓库内完整指南:docs/sol_rl.md(含 FLUX.1、SD3.5-L 的模型专属说明与 NVFP4 安装步骤)
  • Sol-RL 训练脚本族:train_scripts/sol_rl/train_sana.py、train_scripts/sol_rl/train_utils.py
  • 奖励实现:diffusion/post_training/rewards.py
  • 提示词数据集与统计跟踪:diffusion/post_training/prompt_dataset.py、diffusion/post_training/stat_tracking.py
  • Sol-RL 训练配方借鉴了 Advantage Weighted Matching 与 DiffusionNFT 两套公开工作(见 docs/sol_rl.md 的致谢部分)

通过本文,你可以同时掌握两条路径:在 Cosmos-RL 生态中使用官方预设快速开启 SANA 的 SFT / RL 训练;以及在本仓库内基于 DiffusionNFT(BON + Preview Rollout)参考实现,逐参数理解图像扩散模型 RL 后训练的完整工程细节。

【免费下载链接】SanaSANA: Efficient High-Resolution Image Synthesis with Linear Diffusion Transformer项目地址: https://gitcode.com/GitHub_Trending/sana/Sana

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/16 14:39:21

小米5刷LineageOS 18.1:Android 11老旗舰刷机实践与避坑指南

简介&#xff1a;小米5安卓11最新系统lineage-18.1小米5.zip是一份为小米5机型定制的第三方系统固件&#xff0c;基于Android 11稳定分支的LineageOS 18.1&#xff0c;适合希望突破官方更新限制、自行刷入新系统以获得性能提升和个性化体验的刷机用户。压缩包共13个文件&#x…

作者头像 李华
网站建设 2026/9/16 14:38:36

TDOA/FDOA联合定位中TSWLS与ICWLS算法性能对比

1. 项目概述在无线定位技术领域&#xff0c;TDOA&#xff08;到达时间差&#xff09;和FDOA&#xff08;到达频率差&#xff09;是两种常用的被动定位方法。最近我在实际项目中遇到了一个有趣的对比场景&#xff1a;需要评估TSWLS&#xff08;Two-Stage Weighted Least Squares…

作者头像 李华
网站建设 2026/9/16 14:38:30

风电场电-氢混合储能容量优化:Matlab多时间尺度协同配置方法

简介&#xff1a;本资源面向电力系统优化、新能源并网及储能技术研究领域的高校师生与工程技术人员&#xff0c;聚焦风电波动平抑这一实际工程痛点&#xff0c;提供基于MATLAB的电-氢混合储能系统容量优化配置完整实现方案。资源包共668个文件&#xff0c;含295个核心MATLAB脚本…

作者头像 李华
网站建设 2026/9/16 14:38:14

STM32音频采样与存储:I2S+DMA双缓冲+FatFS实现高保真录音

简介&#xff1a;该工程以STM32微控制器为核心&#xff0c;围绕音频信号采集、处理与存储展开设计&#xff0c;覆盖ADC模拟量采样、定时器触发、缓存管理和文件写入等实现细节&#xff0c;非常适合课程设计、毕业设计或嵌入式进阶练习。资源压缩包内共包含297个文件&#xff0c…

作者头像 李华