Skip to content

Latest commit

 

History

2 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

VideoGen-Align:Wan2.1 视频生成 SFT + GRPO 后训练

本项目围绕 Wan2.1-T2V-1.3B 搭建了一套完整的两阶段视频生成 Post-training Pipeline:

  1. 基于 OpenVid-1k 进行 LoRA Supervised Fine-Tuning(SFT)
  2. 基于 VideoAlign 奖励进行 Online GRPO 强化学习后训练

项目重点不是重新训练一个视频生成模型,而是研究:

监督微调与强化学习后训练分别会如何影响视频视觉质量、运动质量和文本对齐能力,以及 RL 是否能够修复 SFT 带来的能力偏移。

实验中我们观察到:

  • SFT 能提升 Visual Quality 和 Motion Quality;
  • 但同时会明显降低 Text Alignment;
  • 因此进一步使用 VideoAlign-TA 作为 Online GRPO Reward;
  • 最终 GRPO 成功恢复并超过 Base 的文本对齐能力,同时进一步提升 Motion Quality 和 Overall Score。

1. 整体流程

                Wan2.1-T2V-1.3B
                        │
                        ▼
                LoRA Supervised FT
                   OpenVid-1k
                        │
                        ▼
                  SFT-150 Model
                        │
                        │
              ┌─────────┴─────────┐
              │                   │
              ▼                   ▼
        Visual / Motion ↑      Text Alignment ↓
                                  │
                                  ▼
                       VideoAlign-TA Reward
                                  │
                                  ▼
                           Online GRPO
                                  │
                                  ▼
                             GRPO Model
                                  │
                                  ▼
                    Text Alignment / Motion ↑

整个实验使用固定的 Prompt + Seed 进行 Base / SFT / GRPO 配对评测,以尽量排除随机采样带来的影响。


2. 核心实验结果

最终使用 VideoAlign 对固定的 100 组 Prompt-Seed 进行评测,并使用:

use_norm=True

得到:

Model VQ ↑ MQ ↑ TA ↑ Overall ↑
Wan2.1 Base -0.650449 -0.476089 -0.326040 -1.452578
SFT-150 -0.602377 -0.434135 -0.401954 -1.438467
GRPO -0.646487 -0.355632 -0.301491 -1.303609

其中:

  • VQ:Visual Quality
  • MQ:Motion Quality
  • TA:Text Alignment
  • Overall:VideoAlign 综合得分

由于这里使用的是 VideoAlign normalized score,因此出现负值是正常的,主要关注不同模型之间的相对变化。


2.1 SFT 相比 Base

Metric Δ SFT - Base
VQ +0.048072
MQ +0.041954
TA -0.075914
Overall +0.014112

可以看到:

SFT 提升了视觉质量与运动质量,但牺牲了文本-视频对齐能力。

这也是本项目进一步引入 RL Post-training 的直接动机。


2.2 GRPO 相比 SFT

Metric Δ GRPO - SFT
VQ -0.044110
MQ +0.078503
TA +0.100463
Overall +0.134858

Online GRPO 显著修复了 SFT 阶段下降的 Text Alignment,同时进一步提高 Motion Quality。


2.3 GRPO 相比 Base

Metric Δ GRPO - Base
VQ +0.003962
MQ +0.120457
TA +0.024549
Overall +0.148969

最终 GRPO:

  • VQ 基本维持 Base 水平;
  • MQ 显著高于 Base;
  • TA 超过 Base;
  • Overall 获得三个模型中最佳结果。

3. Repository Structure

VideoRL-PostTraining/
│
├── README.md
│
├── config/
│   ├── sft_openvid1k.json
│   └── grpo_videoalign_ta_1gpu.yaml
│
├── scripts/
│   ├── train_sft.sh
│   ├── infer_sft.py
│   ├── merge_sft_lora.py
│   ├── generate_eval_videos.py
│   │
│   ├── prepare_grpo_prompts.py
│   ├── infer_grpo.py
│   └── eval_videoalign.py
│
├── trainer/
│   ├── wan_trainer.py
│   └── rewards.py
│
├── data/
│   ├── grpo_ta_formal/
│   │   ├── train.txt
│   │   └── test.txt
│   │
│   └── eval_prompts.json
│
├── results/
│   └── videoalign_results.json
│
├── requirements.txt
└── .gitignore

4. Stage I:LoRA-SFT

4.1 Base Model

SFT 使用:

Wan2.1-T2V-1.3B-Diffusers

建议将模型放置于:

checkpoints/
└── Wan2.1-T2V-1.3B-Diffusers/

模型权重不包含在本仓库中,需要单独下载。


4.2 SFT Dataset

SFT 使用 FineTainers 提供的:

finetrainers/OpenVid-1k-split

该数据集是 OpenVid-1M 的一个约 1K 视频子集,主要用于低成本的视频生成微调实验。

数据以 WebDataset/TAR 风格存储,例如:

sample_xxx.video
sample_xxx.caption
sample_xxx.aesthetic_score
sample_xxx.motion_score
sample_xxx.temporal_consistency_score
sample_xxx.camera_motion
...

本地目录建议组织为:

data/
└── openvid-1k/
    ├── README.md
    └── dataset_0000.tar

本项目使用的数据配置:

config/sft_openvid1k.json

内容:

{
  "datasets": [
    {
      "data_root": "./data/openvid-1k",
      "dataset_type": "video",
      "video_resolution_buckets": [
        [49, 480, 832]
      ],
      "reshape_mode": "bicubic",
      "remove_common_llm_caption_prefixes": true
    }
  ]
}

即训练视频统一使用:

49 Frames
480 × 832
Bicubic Resize

5. SFT Training

SFT 基于 FineTainers 完成。

核心训练配置如下:

Parameter Value
Base Model Wan2.1-T2V-1.3B
Training Type LoRA
Planned Training Steps 200
Selected Checkpoint 150
Batch Size 8
Gradient Accumulation 1
LoRA Rank 32
LoRA Alpha 32
Learning Rate 5e-5
Optimizer AdamW
LR Scheduler Constant with Warmup
Warmup Steps 10
Weight Decay 1e-4
Gradient Checkpointing Enabled
Seed 42
Checkpoint Interval 50

LoRA Target Modules:

blocks.*(to_q|to_k|to_v|to_out.0)

启动训练:

export FINETRAINERS_ROOT=/path/to/finetrainers
export WAN_MODEL_PATH=/path/to/Wan2.1-T2V-1.3B-Diffusers

bash scripts/train_sft.sh

默认使用:

config/sft_openvid1k.json

作为 Dataset Config。


5.1 SFT Checkpoint

训练过程中每 50 Steps 保存一次 LoRA。

本项目后续实验使用:

lora_weights/
└── 000150/
    └── pytorch_lora_weights.safetensors

即 SFT-150。


6. SFT Inference

可以使用:

scripts/infer_sft.py

对 Base 和 SFT 模型进行同 Prompt / 同 Seed 的生成对比。

例如:

python scripts/infer_sft.py \
  --model_dir ./checkpoints/Wan2.1-T2V-1.3B-Diffusers \
  --lora_dir ./outputs/sft/openvid1k_sft/lora_weights/000150 \
  --prompt "a dog running." \
  --seed 42

默认生成配置:

Parameter Value
Resolution 832 × 480
Frames 49
Sampling Steps 50
Guidance Scale 5.0
FPS 16

Base 与 SFT 使用完全相同的采样配置与随机种子。


7. Merge SFT LoRA

后续 GRPO 并不是直接从 Base Wan2.1 开始,而是从 SFT-150 merged checkpoint 初始化。

使用:

python scripts/merge_sft_lora.py \
  --base-model ./checkpoints/Wan2.1-T2V-1.3B-Diffusers \
  --lora-dir ./outputs/sft/openvid1k_sft/lora_weights/000150 \
  --output-dir ./checkpoints/models_sft150_merged

得到:

checkpoints/
└── models_sft150_merged/

该模型将作为第二阶段 Online GRPO 的初始化模型。


8. Paired Base / SFT Generation

为了保证定量评测公平,本项目使用固定的 Prompt + Seed 对 Base 和 SFT 进行 paired generation。

运行:

python scripts/generate_eval_videos.py \
  --model-dir ./checkpoints/Wan2.1-T2V-1.3B-Diffusers \
  --lora-dir ./outputs/sft/openvid1k_sft/lora_weights/000150 \
  --manifest ./data/eval_prompts.json \
  --output-root ./outputs/eval \
  --variant both

输出目录:

outputs/eval/
├── videos/
│   ├── base/
│   └── sft150/
│
└── metadata/
    ├── base_generation.jsonl
    ├── sft150_generation.jsonl
    └── inference_config.json

对于同一条样本:

Base:
0000_seed883525423.mp4

SFT:
0000_seed883525423.mp4

两者具有完全相同的:

Prompt
Seed
Resolution
Frame Number
Sampling Steps
Guidance Scale

唯一变化为模型本身。


9. Stage II:Online GRPO

SFT 实验显示:

VQ ↑
MQ ↑
TA ↓

因此第二阶段不再继续做普通监督训练,而是针对 Text Alignment 进行强化学习后训练。

本项目基于 GenRL 搭建 Wan2.1 Online GRPO Pipeline,并使用 VideoAlign 作为 Reward Model。


10. VideoAlign Reward

VideoAlign 提供:

VQ:Visual Quality
MQ:Motion Quality
TA:Text Alignment

正式 GRPO 实验中仅使用:

reward_fn:
  videoalign_ta: 1.0

即:

直接优化 Text Alignment,而不是将 VQ / MQ / TA 全部放入 Reward。

这样可以更明确地验证:

RL 是否能够针对 SFT 阶段出现的 Alignment Regression 进行定向修复。


11. GRPO Configuration

正式 GRPO 配置:

Parameter Value
Initialization SFT-150 Merged
Reward VideoAlign-TA
Training Rounds 300
Videos per Prompt 2
Online Rollouts ~600
Resolution 832 × 480
Frames 49
Sampling Steps 50
Guidance Scale 5.0
LoRA Rank 32
LoRA Alpha 64
Learning Rate 5e-5
KL Beta 3e-4
Train Micro Batch 1
Gradient Accumulation 2
Gradient Checkpointing Enabled
EMA Disabled

完整配置位于:

config/grpo_videoalign_ta_1gpu.yaml

12. Single-GPU Video GRPO

正式实验在单张:

NVIDIA A100 80GB

上运行。

由于高分辨率视频 Online RL 显存开销较大,本项目针对 GenRL 的 Wan Trainer 做了若干单卡适配。


12.1 Explicit Gradient Checkpointing

在 Wan Transformer 上显式开启:

pipeline.transformer.enable_gradient_checkpointing()

以降低视频反向传播阶段的显存占用。


12.2 Training Micro-Batch

原始逻辑中 rollout batch 与 training micro-batch 耦合较紧。

本项目将实际训练 micro-batch 改为:

micoe_batch = cfg.train.batch_size

正式设置:

Train Micro Batch = 1
Gradient Accumulation = 2

12.3 Single-GPU PEFT / KL Compatibility

在 KL Reference Forward 中,分布式环境下 Transformer 可能表现为:

self.transformer.module

而单 GPU PEFT 模型则可能直接为:

self.transformer

因此修改为同时兼容:

transformer = (
    self.transformer.module
    if hasattr(self.transformer, "module")
    else self.transformer
)

with transformer.disable_adapter():
    ...

相关实现位于:

trainer/wan_trainer.py

13. GRPO Prompt Preparation

GRPO 使用独立的 Prompt Pool。

通过:

python scripts/prepare_grpo_prompts.py \
  --input /path/to/prompt_pool \
  --output-dir ./data/grpo_ta_formal \
  --train-size 512 \
  --test-size 32 \
  --seed 42

生成:

data/
└── grpo_ta_formal/
    ├── train.txt
    └── test.txt

正式设置:

Train Prompts = 512
Test Prompts  = 32
Seed          = 42

14. GRPO Training

GRPO 的核心训练配置位于:

config/grpo_videoalign_ta_1gpu.yaml

Trainer 和 Reward 相关修改位于:

trainer/
├── wan_trainer.py
└── rewards.py

其中 trainer/ 下代码基于 GenRL 进行修改,主要用于:

  • Wan2.1 单 GPU PEFT 兼容;
  • KL Reference Forward;
  • Video GRPO Micro-Batching;
  • Gradient Checkpointing;
  • VideoAlign Reward Integration。

具体训练入口沿用 GenRL 的训练框架。


15. GRPO Inference

GRPO LoRA 是基于 SFT-150 merged model 训练得到的,因此推理时必须加载到相同的 SFT 初始化模型上。

运行:

python scripts/infer_grpo.py \
  --sft-model ./checkpoints/models_sft150_merged \
  --grpo-lora ./checkpoints/grpo_final \
  --manifest ./data/eval_prompts.json \
  --output-dir ./outputs/grpo \
  --local-files-only

默认推理设置:

Resolution      832 × 480
Frames          49
Sampling Steps  50
Guidance Scale  5.0
FPS             16

这些参数与 Base / SFT 的最终评测设置保持一致。


16. Evaluation Protocol

最终评测使用固定:

100 Prompt-Seed Pairs

Manifest:

data/eval_prompts.json

示例:

{
  "id": "0000",
  "category": "scenery",
  "prompt": "view of the sea from an abandoned building",
  "seed": 883525423
}

Base、SFT 和 GRPO 都使用完全相同的:

Prompt
Seed
Resolution
Frames
Sampling Steps
Guidance Scale

从而形成严格的 paired evaluation。


17. VideoAlign Evaluation

使用:

scripts/eval_videoalign.py

对生成视频计算:

VQ
MQ
TA
Overall

运行示例:

python scripts/eval_videoalign.py \
  --video-dir ./outputs/grpo \
  --manifest ./data/eval_prompts.json \
  --videoalign-dir /path/to/VideoAlign \
  --checkpoint /path/to/VideoAlign/checkpoints \
  --output-dir ./results/videoalign_grpo \
  --model-name grpo \
  --batch-size 2

正式实验设置:

Batch Size = 2
use_norm   = True

评测脚本支持:

  • Batch Evaluation
  • JSONL Cache
  • 中断续评
  • Per-video CSV
  • Mean / Std / Min / Max 汇总

18. 实验结论

本项目得到的一个比较明确的现象是:

Supervised Fine-Tuning
        │
        ├── Visual Quality ↑
        ├── Motion Quality ↑
        │
        └── Text Alignment ↓

这说明视频生成 SFT 并不一定会在所有维度上同步提高模型能力。

因此进一步使用:

VideoAlign-TA
     ↓
Online GRPO

进行定向后训练。

最终:

GRPO vs SFT

TA       +0.100463
MQ       +0.078503
Overall  +0.134858

同时:

GRPO VQ  = -0.646487
Base VQ  = -0.650449

表明在显著提升对齐能力的同时,模型并未出现明显的视觉质量崩塌。


19. 模型权重与数据

本仓库主要提供:

Training Config
SFT Scripts
GRPO Trainer Adaptation
Inference Scripts
Evaluation Scripts
Prompt Splits
Experiment Results

以下大文件不直接提交至 Git:

Wan2.1 Base Checkpoint
SFT LoRA Weights
Merged SFT Model
GRPO LoRA Weights
VideoAlign Checkpoint
OpenVid Video Dataset
Generated Evaluation Videos

请根据相应项目说明单独下载或生成。


20. Acknowledgements

本项目基于以下开源项目构建:

  • Wan2.1
  • Hugging Face Diffusers
  • PEFT
  • FineTainers
  • GenRL
  • VideoAlign
  • OpenVid

其中:

trainer/wan_trainer.py
trainer/rewards.py

基于 GenRL 相关实现进行适配修改,主要用于本项目的 Wan2.1 单 GPU Video GRPO 实验。

本仓库不主张对第三方项目原始实现的所有权,请同时遵循对应项目的开源许可证。


21. 项目定位

本项目聚焦于:

视频生成模型的 SFT + RL 两阶段后训练流程。

相比单纯进行一次 LoRA 微调,本项目更关注模型能力在不同 Post-training Stage 之间的变化:

Base
  ↓
SFT
  ↓
发现 Alignment Trade-off
  ↓
Targeted RL Post-training
  ↓
GRPO

最终展示了一条完整的视频生成模型后训练实践路径:

监督微调 → 能力诊断 → Reward Design → Online RL → 配对评测。

About

No description, website, or topics provided.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages