ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

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

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 训练命令以及仓库内 DiffusionNFTBON Preview Rollout参考实现的底层原理与关键参数。背景当高效扩散模型遇上通用 RL 基础设施SANA是面向高分辨率图像与视频生成的高效代码库线性注意力 DiT 架构而Cosmos-RL是 NVIDIA 推出的灵活、可扩展的强化学习框架。两者通过官方合作打通为 SANA 提供了完整的后训练Post-Training基础设施覆盖SFT监督微调图像与视频的 Full Fine-Tuning 与 LoRA 微调RL强化学习如DiffusionNFT与Flow-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-sft、sana-image-sft-lorasana-video-sft、sana-video-sft-loraRLsana-image-nftsana-video-nft这些预设统一遵循 Cosmos-RL 的配置规范参数细节见其 Configuration 文档。命名规律为模态-任务-变体其中-lora后缀表示只训练 LoRA 适配器而非全量权重-nft后缀表示走 DiffusionNFT 强化学习流程。本仓库的对应本地实现如果你希望在本仓库内直接体验与 Cosmos-RL 同源的 DiffusionNFT 式 RL 训练无需外部框架可参考 Sol-RL 后训练模块。其配置命名遵循model_family_reward模式例如sana_diffusionnft_pickscoresana_compile_hpsv2sana_sol_rl_imagereward详见 configs/sol_rl/sana.py另有 flux1.py、sd3.py 对应 FLUX.1 与 SD3.5-L。奖励服务Reward ServiceCosmos-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图像 DiffusionNFTcosmos-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 指南。仓库内参考实现DiffusionNFTBON Preview Rollout原理剖析与本仓库 Sol-RL 模块直接对应的是DiffusionNFT算法族。下面以源码为依据说明其运行机制帮助你理解 Cosmos-RL 中*-nft预设背后的通用范式。配置族与 rollout 形状configs/sol_rl/sana.py 定义了五个配置族对应不同的 rollout 成本与量化策略Family含义Rollout 形状in-N / best-of-MTE / NVFP4diffusionnftPEFT 推理基线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_scaling为preview_modelpeft、fullrollout_modelpeftcompile为fullrollout_modelcompilenaive_quant为fullrollout_modelcompile_nvfp4sol_rl则为preview_step6、preview_modelcompile_nvfp4、fullrollout_modelcompile。推荐首次运行的配置sana_diffusionnft_pickscore、sd3_diffusionnft_pickscore、flux1_diffusionnft_pickscore均为 24-in-24 的 PEFT 基线无需编译与量化即可验证全链路。两阶段解耦FP4 探索 BF16 再生sol_rl族是仓库中吞吐优化最激进的方案其核心是把广泛探索与精细再生分离见 train_scripts/sol_rl/train_sana.py 中_rollout_for_one_prompt的实现Stage 1FP4 探索用 NVFP4 量化的 compile 草稿模型draft model以 6 步preview_step6采样 96 张候选图并先用奖励函数打分only_strictTrue的严格模式Stage 2BF16 再生按stage1_select_modebest_worst从草稿池中选出候选种子换用 BF16 compile 模型以完整 10 步rollout_sample_num_steps10重新生成得到高质量最终样本用于训练。源码中_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_SPECconfigs/sol_rl/sana.py:sana_diffusionnft_pickscore \ bash train_scripts/sol_rl/run_sana_single_node_8gpu.sh脚本要点默认CONFIG_SPECconfigs/sol_rl/sana.py:sana_diffusionnft_pickscore即不传环境变量也能直接跑通基线可覆盖的环境变量包括NPROC_PER_NODE默认 8、CUDA_VISIBLE_DEVICES、MASTER_PORT、NATIVE_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-enginepython -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供后续按步训练优势估计默认启用PerPromptStatTrackerdiffusion/post_training/stat_tracking.py做 per-prompt 的奖励统计归一化global_stdTrue并计算 zero-std 比例等诊断指标奖励向量会在时间步维度复制展开策略损失对每个时间步构造正/负预测positive_prediction与implicit_negative_prediction由config.beta控制插值以采样奖励加权后的 MSE 形式r * positive_loss (1-r) * negative_loss构成policy_lossKL 约束叠加当前模型与参考模型禁用适配器预测差异的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-b79Kyuvalkirstain/PickScore_v1首次使用自动下载clipscoreopenai/clip-vit-large-patch14首次使用自动下载hpsv2open_clip_pytorch_model.binHPS_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_rate3e-4、batch_size1每 GPU、gradient_accumulation_steps1、max_grad_norm0.002、num_inner_epochs1、adv_clip_max5、timestep_fraction0.6、beta0.0001、mixed_precisionbf16采样num_steps40、eval_num_steps40、num_image_per_prompt24、best_of_n24、noise_level1.0、test_batch_size1LoRAlora_rank32、lora_alpha64、lora_init_weightsTrueSANA 具体目标模块见 configs/sol_rl/sana.py 的lora_target_modulesattn.qkv、attn.proj、cross_attn.q_linear、cross_attn.kv_linear、cross_attn.projRollout 相关rollout_sample_num_steps10、preview_step0默认关闭两阶段、rollout_sample_guidance_scale1.0、compile_modemax-autotune-no-cudagraphsNVFP4nvfp4_skip_modulesSANA 跳过t_embedder、y_embedder、x_embedder、final_layer、attn.qkv等敏感层与nvfp4_min_dim2240控制量化粒度。SANA 原生模型结构由 configs/sol_rl/Sana1.0_1600M_linear.yaml 描述SanaMSLinearFFN_1600M_P1_D20线性注意力架构、attn_type: linear、ffn_type: glumbconv_linear、Gemma-2-2B-it 文本编码器、DC-AEdc-ae-f32c32-sana-1.1-diffusersVAE、线性 flow 调度。RL 训练时通过pyrallis解析该 YAML 构建原生 Transformer并复用 diffusersSanaPipeline的文本编码器与 VAE仅在推理时切换训练好的 LoRA / 编译 / 量化模型副本。注意事项与适用前提Cosmos-RL 预设的归属configs/sana/*.toml与reward_service/位于 Cosmos-RL 仓库本文中标注为仓库本地实现的内容configs/sol_rl、train_scripts/sol_rl、diffusion/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/mean、reward/max、zero_std_ratio、policy_loss、kl_div_loss、grad_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.pySol-RL 训练配方借鉴了 Advantage Weighted Matching 与 DiffusionNFT 两套公开工作见 docs/sol_rl.md 的致谢部分通过本文你可以同时掌握两条路径在 Cosmos-RL 生态中使用官方预设快速开启 SANA 的 SFT / RL 训练以及在本仓库内基于 DiffusionNFTBON Preview Rollout参考实现逐参数理解图像扩散模型 RL 后训练的完整工程细节。【免费下载链接】SanaSANA: Efficient High-Resolution Image Synthesis with Linear Diffusion Transformer项目地址: https://gitcode.com/GitHub_Trending/sana/Sana创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表