Update Wan2.2 reward lora (#310)
* Add Reward LoRAs for Wan2.2-Fun * Update README_TRAIN_REWARD.md
This commit is contained in:
@@ -224,7 +224,7 @@ Please read the [quick-start](https://github.com/aigc-apps/CogVideoX-Fun/blob/ma
|
||||
pip install hpsv2
|
||||
site_packages=$(python -c "import site; print(site.getsitepackages()[0])")
|
||||
wget -O $site_packages/hpsv2/src/open_clip/factory.py https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/package/patches/hpsv2_src_open_clip_factory_patches.py
|
||||
wget -O $site_packages/hpsv2/src/open_clip/ https://github.com/tgxs002/HPSv2/raw/refs/heads/master/hpsv2/src/open_clip/bpe_simple_vocab_16e6.txt.gz
|
||||
wget -O $site_packages/hpsv2/src/open_clip/bpe_simple_vocab_16e6.txt.gz https://github.com/tgxs002/HPSv2/raw/refs/heads/master/hpsv2/src/open_clip/bpe_simple_vocab_16e6.txt.gz
|
||||
```
|
||||
|
||||
> [!NOTE]
|
||||
|
||||
Executable
+41
@@ -0,0 +1,41 @@
|
||||
# Wan2.2-Fun-Reward-LoRAs
|
||||
## Introduction
|
||||
We explore the Reward Backpropagation technique <sup>[1](#ref1) [2](#ref2)</sup> to optimized the generated videos by [Wan2.2-Fun](https://github.com/aigc-apps/VideoX-Fun) for better alignment with human preferences.
|
||||
We provide the following pre-trained models (i.e. LoRAs) along with [the training script](https://github.com/aigc-apps/VideoX-Fun/blob/main/scripts/wan2.2_fun/train_reward_lora.py). You can use these LoRAs to enhance the corresponding base model as a plug-in or train your own reward LoRA.
|
||||
|
||||
For more details, please refer to our [GitHub repo](https://github.com/aigc-apps/VideoX-Fun).
|
||||
|
||||
| Name | Base Model | Reward Model | Hugging Face | Description |
|
||||
|--|--|--|--|--|
|
||||
| Wan2.2-Fun-A14B-InP-high-noise-HPS2.1.safetensors | [Wan2.2-Fun-A14B-InP (high noise)](https://huggingface.co/alibaba-pai/Wan2.2-Fun-A14B-InP/tree/main/high_noise_model) | [HPS v2.1](https://github.com/tgxs002/HPSv2) | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.2-Fun-Reward-LoRAs/resolve/main/Wan2.2-Fun-A14B-InP-high-noise-HPS2.1.safetensors) | Official HPS v2.1 reward LoRA (`rank=128` and `network_alpha=64`) for Wan2.2-Fun-A14B-InP (high noise). It is trained with a batch size of 8 for 5,000 steps.|
|
||||
| Wan2.2-Fun-A14B-InP-low-noise-HPS2.1.safetensors | [Wan2.2-Fun-A14B-InP (low noise)](https://huggingface.co/alibaba-pai/Wan2.2-Fun-A14B-InP/tree/main/low_noise_model) | [MPS](https://github.com/Kwai-Kolors/MPS) | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.2-Fun-Reward-LoRAs/resolve/main/Wan2.2-Fun-A14B-InP-low-noise-HPS2.1.safetensors) | Official HPS v2.1 reward LoRA (`rank=128` and `network_alpha=64`) for Wan2.2-Fun-A14B-InP (low noise). It is trained with a batch size of 8 for 2,700 steps.|
|
||||
| Wan2.2-Fun-A14B-InP-high-noise-MPS.safetensors | [Wan2.2-Fun-A14B-InP (high noise)](https://huggingface.co/alibaba-pai/Wan2.2-Fun-A14B-InP/tree/main/high_noise_model) | [HPS v2.1](https://github.com/tgxs002/HPSv2) | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.2-Fun-Reward-LoRAs/resolve/main/Wan2.2-Fun-A14B-InP-high-noise-MPS.safetensors) | Official MPS reward LoRA (`rank=128` and `network_alpha=64`) for Wan2.2-Fun-A14B-InP (high noise). It is trained with a batch size of 8 for 5,000 steps.|
|
||||
| Wan2.2-Fun-A14B-InP-low-noise-MPS.safetensors | [Wan2.2-Fun-A14B-InP (low noise)](https://huggingface.co/alibaba-pai/Wan2.2-Fun-A14B-InP/tree/main/low_noise_model) | [MPS](https://github.com/Kwai-Kolors/MPS) | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-Reward-LoRAs/resolve/main/Wan2.1-Fun-14B-InP-MPS.safetensors) | Official MPS reward LoRA (`rank=128` and `network_alpha=64`) for Wan2.2-Fun-A14B-InP (low noise). It is trained with a batch size of 8 for xxx steps.|
|
||||
|
||||
> [!NOTE]
|
||||
> We found that, MPS reward LoRA for the low-noise model converges significantly more slowly than on the other models, and may not deliver satisfactory results. Therefore, for the low-noise model, we recommend using HPSv2.1 reward LoRA.
|
||||
|
||||
## Demo
|
||||
Please refer to [here](https://huggingface.co/alibaba-pai/Wan2.2-Fun-Reward-LoRAs#demo).
|
||||
|
||||
## Quick Start
|
||||
Set `lora_path` along with `lora_weight` for the low noise reward LoRA, while specifying `lora_high_path` and `lora_high_weight` for high noise reward LoRA in [examples/wan2.2_fun/predict_t2v.py](https://github.com/aigc-apps/VideoX-Fun/blob/main/examples/wan2.1_fun/predict_t2v.py).
|
||||
|
||||
## Training
|
||||
The training code is based on [train_lora.py](./train_lora.py). We provide a shell script to train the HPS v2.1 reward LoRA for the low noise model of Wan2.2-Fun-A14B-InP, which can be trained on a single 8*A100 node with 80GB VRAM. To train reward LoRA for the high noise model, Deepspeed Zero3 with CPU offload is required.
|
||||
|
||||
Please refer to [Setup](https://github.com/aigc-apps/VideoX-Fun/blob/main/scripts/cogvideox_fun/README_TRAIN_REWARD.md#setup) and [Important Args](https://github.com/aigc-apps/VideoX-Fun/blob/main/scripts/cogvideox_fun/README_TRAIN_REWARD.md#important-args) before training.
|
||||
|
||||
|
||||
## Limitations
|
||||
1. We observe after training to a certain extent, the reward continues to increase, but the quality of the generated videos does not further improve.
|
||||
The model trickly learns some shortcuts (by adding artifacts in the background, i.e., adversarial patches) to increase the reward.
|
||||
2. Currently, there is still a lack of suitable preference models for video generation. Directly using image preference models cannot
|
||||
evaluate preferences along the temporal dimension (such as dynamism and consistency). Further more, We find using image preference models leads to a decrease
|
||||
in the dynamism of generated videos. Although this can be mitigated by computing the reward using only the first frame of the decoded video, the impact still persists.
|
||||
|
||||
## Reference
|
||||
<ol>
|
||||
<li id="ref1">Clark, Kevin, et al. "Directly fine-tuning diffusion models on differentiable rewards.". In ICLR 2024.</li>
|
||||
<li id="ref2">Prabhudesai, Mihir, et al. "Aligning text-to-image diffusion models with reward backpropagation." arXiv preprint arXiv:2310.03739 (2023).</li>
|
||||
</ol>
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,65 @@
|
||||
export MODEL_NAME="models/Diffusion_Transformer/Wan2.2-Fun-A14B-InP"
|
||||
export TRAIN_PROMPT_PATH="MovieGenVideoBench_train.txt"
|
||||
|
||||
# Train HPSv2.1 reward LoRA for the low noise model of Wan2.2-Fun-A14B-InP
|
||||
accelerate launch --mixed_precision="bf16" --num-processes=8 --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json scripts/wan2.2_fun/train_reward_lora.py \
|
||||
--config_path="config/wan2.2/wan_civitai_i2v.yaml" \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--train_batch_size=1 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--max_train_steps=10000 \
|
||||
--checkpointing_steps=100 \
|
||||
--learning_rate=1e-05 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--max_grad_norm=0.3 \
|
||||
--boundary_type="low" \
|
||||
--lora_skip_name="ffn" \
|
||||
--low_vram \
|
||||
--use_deepspeed \
|
||||
--prompt_path=$TRAIN_PROMPT_PATH \
|
||||
--train_sample_height=256 \
|
||||
--train_sample_width=256 \
|
||||
--num_inference_steps=40 \
|
||||
--video_length=81 \
|
||||
--num_decoded_latents=1 \
|
||||
--reward_fn="HPSReward" \
|
||||
--reward_fn_kwargs='{"version": "v2.1"}' \
|
||||
--backprop_strategy="tail" \
|
||||
--backprop_num_steps=1 \
|
||||
--backprop
|
||||
|
||||
# Train MPS reward LoRA for the high noise model of Wan2.2-Fun-A14B-InP
|
||||
# accelerate launch --mixed_precision="bf16" --num-processes=8 --use_deepspeed --deepspeed_config_file config/zero_stage3_config_cpu_offload.json scripts/wan2.2_fun/train_reward_lora.py \
|
||||
# --config_path="config/wan2.2/wan_civitai_i2v.yaml" \
|
||||
# --pretrained_model_name_or_path=$MODEL_NAME \
|
||||
# --train_batch_size=1 \
|
||||
# --gradient_accumulation_steps=1 \
|
||||
# --max_train_steps=10000 \
|
||||
# --checkpointing_steps=100 \
|
||||
# --learning_rate=1e-05 \
|
||||
# --seed=42 \
|
||||
# --output_dir="output_dir" \
|
||||
# --gradient_checkpointing \
|
||||
# --mixed_precision="bf16" \
|
||||
# --adam_weight_decay=3e-2 \
|
||||
# --adam_epsilon=1e-10 \
|
||||
# --max_grad_norm=0.3 \
|
||||
# --boundary_type="high" \
|
||||
# --lora_skip_name="ffn" \
|
||||
# --low_vram \
|
||||
# --use_deepspeed \
|
||||
# --prompt_path=$TRAIN_PROMPT_PATH \
|
||||
# --train_sample_height=256 \
|
||||
# --train_sample_width=256 \
|
||||
# --num_inference_steps=40 \
|
||||
# --video_length=81 \
|
||||
# --num_decoded_latents=1 \
|
||||
# --reward_fn="MPSReward" \
|
||||
# --backprop_strategy="tail" \
|
||||
# --backprop_num_steps=1 \
|
||||
# --backprop
|
||||
Reference in New Issue
Block a user