update reward training

This commit is contained in:
hkunzhe
2025-01-21 14:37:28 +08:00
parent 286b9617eb
commit 68b01135cf
3 changed files with 177 additions and 175 deletions
+12 -5
View File
@@ -2,6 +2,9 @@
We explore the Reward Backpropagation technique <sup>[1](#ref1) [2](#ref2)</sup> to optimized the generated videos by [EasyAnimateV5](https://github.com/aigc-apps/EasyAnimate/tree/main/easyanimate) for better alignment with human preferences.
We provide pre-trained models (i.e. LoRAs) along with the training script. You can use these LoRAs to enhance the corresponding base model as a plug-in or train your own reward LoRA.
> [!NOTE]
> For EasyAnimateV5.1, we have merged the reward LoRAs into the base model. Please use the base model directly.
- [Enhance EasyAnimate with Reward Backpropagation (Preference Optimization)](#enhance-easyanimate-with-reward-backpropagation-preference-optimization)
- [Demo](#demo)
- [EasyAnimateV5-12b-zh-InP](#easyanimatev5-12b-zh-inp)
@@ -282,18 +285,22 @@ Due to the resize and crop preprocessing operations, we suggest using a 1:1 aspe
can be found in [reward_fn.py](../cogvideox/reward/reward_fn.py).
You can also customize your own reward model (e.g., combining aesthetic predictor with HPS).
+ `num_decoded_latents` and `num_sampled_frames`: The number of decoded latents (for VAE) and sampled frames (for the reward model).
Since CogVideoX-Fun adopts the 3D casual VAE, we found decoding only the first latent to obtain the first frame for computing the reward
not only reduces training memory usage but also prevents excessive reward optimization and maintains the dynamics of generated videos.
Since EasyAnimate adopts the 3D casual VAE, we found decoding only the first latent to obtain the first frame for computing the reward
not only reduces training GPU memory usage but also prevents excessive reward optimization and maintains the dynamics of generated videos.
> [!NOTE]
> In EasyAnimateV5, we only retained the gradient of the last step in the denoising process to reduce GPU memory usage. However, for V5.1, we found that if we only perform reward backpropagation on the last step, the gradient norm becomes very small (usually below 0.001), making it difficult for reward training to converge. This might be due to V5.1 adopts the flow-matching sampling in the training and inference. Therefore, in pratice, we retain the gradients of the last several steps for V5.1.
## 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.
The model trickly learns some shortcuts (by adding artifacts in the background, i.e., reward hacking) 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.
## References
<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>
<li id="ref1">Wu, Xiaoshi, et al. "Deep reward supervisions for tuning text-to-image diffusion models." In ECCV 2025.</li>
<li id="ref2">Clark, Kevin, et al. "Directly fine-tuning diffusion models on differentiable rewards.". In ICLR 2024.</li>
<li id="ref3">Prabhudesai, Mihir, et al. "Aligning text-to-image diffusion models with reward backpropagation." arXiv preprint arXiv:2310.03739 (2023).</li>
</ol>
+119 -166
View File
@@ -17,22 +17,21 @@
import argparse
import gc
import json
import logging
import math
import os
import pickle
import random
import shutil
import sys
import json
from contextlib import contextmanager
from typing import List, Optional
import random
from typing import Optional, List
import accelerate
import diffusers
import numpy as np
import torch
import torch.nn.functional as F
import torch.utils.checkpoint
import torchvision.transforms as transforms
import transformers
@@ -40,27 +39,16 @@ from accelerate import Accelerator
from accelerate.logging import get_logger
from accelerate.state import AcceleratorState
from accelerate.utils import ProjectConfiguration, set_seed
from decord import VideoReader
from diffusers import (AutoencoderKL, DDIMScheduler, DDPMScheduler,
FlowMatchEulerDiscreteScheduler)
from diffusers import DDIMScheduler, FlowMatchEulerDiscreteScheduler
from diffusers.optimization import get_scheduler
from diffusers.training_utils import EMAModel
from diffusers.utils import check_min_version, deprecate, is_wandb_available
from diffusers.utils import check_min_version, is_wandb_available
from diffusers.utils.import_utils import is_xformers_available
from diffusers.utils.torch_utils import is_compiled_module
from decord import VideoReader
from einops import rearrange
from huggingface_hub import create_repo, upload_folder
from omegaconf import OmegaConf
from packaging import version
from PIL import Image
from torch.utils.data import RandomSampler
from torch.utils.tensorboard import SummaryWriter
from torchvision import transforms
from tqdm.auto import tqdm
from transformers import (AutoTokenizer, BertModel, BertTokenizer,
CLIPImageProcessor, CLIPVisionModelWithProjection,
Qwen2Tokenizer, Qwen2VLForConditionalGeneration,
T5EncoderModel, T5Tokenizer)
from transformers import BertModel, BertTokenizer, Qwen2Tokenizer, Qwen2VLForConditionalGeneration, T5EncoderModel, T5Tokenizer
from transformers.utils import ContextManagers
import datasets
@@ -76,7 +64,7 @@ from transformers.utils import ContextManagers
import easyanimate.reward.reward_fn as reward_fn
from easyanimate.models import (name_to_autoencoder_magvit,
name_to_transformer3d)
from easyanimate.pipeline.pipeline_easyanimate import get_3d_rotary_pos_embed, get_resize_crop_region_for_grid
from easyanimate.pipeline.pipeline_easyanimate_inpaint import get_3d_rotary_pos_embed, get_resize_crop_region_for_grid
from easyanimate.pipeline.pipeline_easyanimate_inpaint import EasyAnimateInpaintPipeline
from easyanimate.utils.lora_utils import create_network, merge_lora
from easyanimate.utils.utils import get_image_to_video_latent, save_videos_grid
@@ -106,101 +94,112 @@ def log_validation(
vae, text_encoder, text_encoder_2, tokenizer, tokenizer_2, transformer3d, network,
loss_fn, config, args, accelerator, weight_dtype, global_step, validation_prompts_idx
):
logger.info("Running validation... ")
try:
logger.info("Running validation... ")
# Get New Transformer
Choosen_Transformer3DModel = name_to_transformer3d[
config['transformer_additional_kwargs'].get('transformer_type', 'Transformer3DModel')
]
# Get New Transformer
Choosen_Transformer3DModel = name_to_transformer3d[
config['transformer_additional_kwargs'].get('transformer_type', 'Transformer3DModel')
]
transformer3d_val = Choosen_Transformer3DModel.from_pretrained_2d(
args.pretrained_model_name_or_path, subfolder="transformer",
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs'])
).to(weight_dtype)
transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
transformer3d_val = Choosen_Transformer3DModel.from_pretrained_2d(
args.pretrained_model_name_or_path, subfolder="transformer",
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs'])
).to(weight_dtype)
transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
if "EasyAnimateV5.1" in args.pretrained_model_name_or_path:
scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="scheduler")
else:
scheduler = DDIMScheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="scheduler")
if "EasyAnimateV5.1" in args.pretrained_model_name_or_path:
scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="scheduler")
else:
scheduler = DDIMScheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="scheduler")
pipeline = EasyAnimateInpaintPipeline(
vae=accelerator.unwrap_model(vae).to(weight_dtype),
text_encoder=accelerator.unwrap_model(text_encoder),
text_encoder_2=accelerator.unwrap_model(text_encoder_2),
tokenizer=tokenizer,
tokenizer_2=tokenizer_2,
transformer=transformer3d_val,
scheduler=scheduler,
)
pipeline = pipeline.to(weight_dtype, accelerator.device)
pipeline = merge_lora(
pipeline, None, 1, accelerator.device, state_dict=accelerator.unwrap_model(network).state_dict(), transformer_only=True
)
to_tensor = transforms.ToTensor()
validation_loss, validation_reward = 0, 0
pipeline = EasyAnimateInpaintPipeline(
vae=accelerator.unwrap_model(vae).to(weight_dtype),
text_encoder=accelerator.unwrap_model(text_encoder),
text_encoder_2=accelerator.unwrap_model(text_encoder_2),
tokenizer=tokenizer,
tokenizer_2=tokenizer_2,
transformer=transformer3d_val,
scheduler=scheduler,
)
pipeline = pipeline.to(dtype=weight_dtype)
if args.low_vram:
pipeline.enable_model_cpu_offload()
else:
pipeline = pipeline.to(device=accelerator.device)
pipeline = merge_lora(
pipeline, None, 1, accelerator.device, state_dict=accelerator.unwrap_model(network).state_dict(), transformer_only=True
)
to_tensor = transforms.ToTensor()
validation_loss, validation_reward = 0, 0
if args.enable_xformers_memory_efficient_attention \
and config['transformer_additional_kwargs'].get('transformer_type', 'Transformer3DModel') == 'Transformer3DModel':
pipeline.enable_xformers_memory_efficient_attention()
if args.enable_xformers_memory_efficient_attention \
and config['transformer_additional_kwargs'].get('transformer_type', 'Transformer3DModel') == 'Transformer3DModel':
pipeline.enable_xformers_memory_efficient_attention()
for i in range(len(validation_prompts_idx)):
validation_idx, validation_prompt = validation_prompts_idx[i]
with torch.no_grad():
with torch.autocast("cuda", dtype=weight_dtype):
if vae.cache_mag_vae:
video_length = int((args.video_length - 1) // vae.mini_batch_encoder * vae.mini_batch_encoder) + 1 if args.video_length != 1 else 1
else:
video_length = int(args.video_length // vae.mini_batch_encoder * vae.mini_batch_encoder) if args.video_length != 1 else 1
sample_size = [args.validation_sample_height, args.validation_sample_width]
input_video, input_video_mask, clip_image = get_image_to_video_latent(
None, None, video_length=args.video_length, sample_size=sample_size
)
for i in range(len(validation_prompts_idx)):
validation_idx, validation_prompt = validation_prompts_idx[i]
with torch.no_grad():
with torch.autocast("cuda", dtype=weight_dtype):
if vae.cache_mag_vae:
video_length = int((args.video_length - 1) // vae.mini_batch_encoder * vae.mini_batch_encoder) + 1 if args.video_length != 1 else 1
else:
video_length = int(args.video_length // vae.mini_batch_encoder * vae.mini_batch_encoder) if args.video_length != 1 else 1
sample_size = [args.validation_sample_height, args.validation_sample_width]
input_video, input_video_mask, clip_image = get_image_to_video_latent(
None, None, video_length=args.video_length, sample_size=sample_size
)
if args.seed is None:
generator = None
else:
generator = torch.Generator(device=accelerator.device).manual_seed(args.seed)
if args.seed is None:
generator = None
else:
generator = torch.Generator(device=accelerator.device).manual_seed(args.seed)
sample = pipeline(
validation_prompt,
video_length = video_length,
negative_prompt = "bad detailed",
height = args.validation_sample_height,
width = args.validation_sample_width,
guidance_scale = 7,
generator = generator,
sample = pipeline(
validation_prompt,
video_length = video_length,
negative_prompt = "bad detailed",
height = args.validation_sample_height,
width = args.validation_sample_width,
guidance_scale = 7,
generator = generator,
video = input_video,
mask_video = input_video_mask,
clip_image = clip_image,
).frames
sample_saved_path = os.path.join(args.output_dir, f"validation_sample/sample-{global_step}-{validation_idx}.mp4")
save_videos_grid(sample, sample_saved_path, fps=8)
video = input_video,
mask_video = input_video_mask,
clip_image = clip_image,
).frames
sample_saved_path = os.path.join(args.output_dir, f"validation_sample/sample-{global_step}-{validation_idx}.mp4")
save_videos_grid(sample, sample_saved_path, fps=8)
num_sampled_frames = 4
sampled_frames_list = []
with video_reader(sample_saved_path) as vr:
sampled_frame_idx_list = np.linspace(0, len(vr), num_sampled_frames, endpoint=False, dtype=int)
sampled_frame_list = vr.get_batch(sampled_frame_idx_list).asnumpy()
sampled_frames = torch.stack([to_tensor(frame) for frame in sampled_frame_list], dim=0)
sampled_frames_list.append(sampled_frames)
sampled_frames = torch.stack(sampled_frames_list)
sampled_frames = rearrange(sampled_frames, "b t c h w -> b c t h w")
loss, reward = loss_fn(sampled_frames, [validation_prompt])
validation_loss, validation_reward = validation_loss + loss, validation_reward + reward
validation_loss = validation_loss / len(validation_prompts_idx)
validation_reward = validation_reward / len(validation_prompts_idx)
num_sampled_frames = 4
sampled_frames_list = []
with video_reader(sample_saved_path) as vr:
sampled_frame_idx_list = np.linspace(0, len(vr), num_sampled_frames, endpoint=False, dtype=int)
sampled_frame_list = vr.get_batch(sampled_frame_idx_list).asnumpy()
sampled_frames = torch.stack([to_tensor(frame) for frame in sampled_frame_list], dim=0)
sampled_frames_list.append(sampled_frames)
sampled_frames = torch.stack(sampled_frames_list)
sampled_frames = rearrange(sampled_frames, "b t c h w -> b c t h w")
loss, reward = loss_fn(sampled_frames, [validation_prompt])
validation_loss, validation_reward = validation_loss + loss, validation_reward + reward
validation_loss = validation_loss / len(validation_prompts_idx)
validation_reward = validation_reward / len(validation_prompts_idx)
del pipeline
del transformer3d_val
gc.collect()
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
del pipeline
del transformer3d_val
gc.collect()
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
return validation_loss, validation_reward
return validation_loss, validation_reward
except Exception as e:
gc.collect()
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
print(f"Eval error with info {e}")
return None, None
def load_prompts(prompt_path, prompt_column="prompt", start_idx=None, end_idx=None):
@@ -609,17 +608,6 @@ def parse_args():
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
),
)
parser.add_argument("--use_ema", action="store_true", help="Whether to use EMA model.")
parser.add_argument(
"--non_ema_revision",
type=str,
default=None,
required=False,
help=(
"Revision of pretrained non-ema model identifier. Must be a branch, tag or git identifier of the local or"
" remote repository specified with --pretrained_model_name_or_path."
),
)
parser.add_argument(
"--dataloader_num_workers",
type=int,
@@ -635,12 +623,6 @@ def parse_args():
parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.")
parser.add_argument("--push_to_hub", action="store_true", help="Whether or not to push the model to the Hub.")
parser.add_argument("--hub_token", type=str, default=None, help="The token to use to push to the Model Hub.")
parser.add_argument(
"--prediction_type",
type=str,
default=None,
help="The prediction_type that shall be used for training. Choose between 'epsilon' or 'v_prediction' or leave `None`. If left to `None` the default prediction type of the scheduler: `noise_scheduler.config.prediciton_type` is chosen.",
)
parser.add_argument(
"--hub_model_id",
type=str,
@@ -707,7 +689,6 @@ def parse_args():
parser.add_argument(
"--enable_xformers_memory_efficient_attention", action="store_true", help="Whether or not to use xformers."
)
parser.add_argument("--noise_offset", type=float, default=0, help="The scale of noise offset.")
parser.add_argument(
"--validation_epochs",
type=int,
@@ -747,18 +728,6 @@ def parse_args():
action="store_true",
help="Whether to train the text encoder. If set, the text encoder should be float32 precision.",
)
parser.add_argument(
"--token_sample_size",
type=int,
default=512,
help="Sample size of the token.",
)
parser.add_argument(
"--video_sample_n_frames",
type=int,
default=17,
help="Num frame of video.",
)
parser.add_argument(
"--config_path",
type=str,
@@ -904,10 +873,6 @@ def parse_args():
if env_local_rank != -1 and env_local_rank != args.local_rank:
args.local_rank = env_local_rank
# default to using the same revision for the non-ema model if not specified
if args.non_ema_revision is None:
args.non_ema_revision = args.revision
return args
@@ -920,15 +885,6 @@ def main():
" Please use `huggingface-cli login` to authenticate with the Hub."
)
if args.non_ema_revision is not None:
deprecate(
"non_ema_revision!=None",
"0.15.0",
message=(
"Downloading 'non_ema' weights from revision branches of the Hub is deprecated. Please make sure to"
" use `--variant=non_ema` instead."
),
)
logging_dir = os.path.join(args.output_dir, args.logging_dir)
config = OmegaConf.load(args.config_path)
@@ -940,14 +896,6 @@ def main():
log_with=args.report_to,
project_config=accelerator_project_config,
)
# Sanity check for validation
do_validation = (args.validation_prompt_path is not None or args.validation_prompts is not None)
if do_validation:
if not (os.path.exists(args.validation_prompt_path) or args.validation_prompt_path.endswith(".txt")):
raise ValueError("The `--validation_prompt_path` must be a txt file containing prompts.")
if args.validation_batch_size < accelerator.num_processes or args.validation_batch_size % accelerator.num_processes != 0:
raise ValueError("The `--validation_batch_size` must be divisible by the number of processes.")
# Make one log on every process with the configuration for debugging.
logging.basicConfig(
@@ -982,7 +930,7 @@ def main():
)
assert any(step <= args.num_inference_steps - 1 for step in args.backprop_step_list)
else:
if args.backprop_strategy in set("tail", "uniform", "random"):
if args.backprop_strategy in set(["tail", "uniform", "random"]):
assert args.backprop_num_steps <= args.num_inference_steps - 1
if args.backprop_strategy == "random":
assert args.backprop_random_start_step <= args.backprop_random_end_step
@@ -1322,8 +1270,6 @@ def main():
initial_global_step = global_step
first_epoch = global_step // num_update_steps_per_epoch
from safetensors.torch import load_file, safe_open
state_dict = load_file(os.path.join(os.path.join(args.output_dir, path), "lora_diffusion_pytorch_model.safetensors"))
m, u = accelerator.unwrap_model(network).load_state_dict(state_dict, strict=False)
@@ -1353,7 +1299,7 @@ def main():
train_reward = 0.0
# In the following training loop, randomly select training prompts and use the
# `EasyAnimatePipeline_Multi_Text_Encoder_Inpaint` to sample videos, calculate rewards, and update the network.
# `EasyAnimatePipelineInpaint` to sample videos, calculate rewards, and update the network.
for _ in range(num_update_steps_per_epoch):
# train_prompt = random.sample(prompt_list, args.train_batch_size)
train_prompt = random.choices(prompt_list, k=args.train_batch_size)
@@ -1527,7 +1473,7 @@ def main():
if args.backprop_strategy == "last":
backprop_step_list = [args.num_inference_steps - 1]
elif args.backprop_strategy == "tail":
backprop_step_list = list(range(args.num_inference_steps - 1))[-args.backprop_num_steps:]
backprop_step_list = list(range(args.num_inference_steps))[-args.backprop_num_steps:]
elif args.backprop_strategy == "uniform":
interval = args.num_inference_steps // args.backprop_num_steps
random_start = random.randint(0, interval)
@@ -1624,9 +1570,12 @@ def main():
accelerator.backward(loss)
if accelerator.sync_gradients:
total_norm = accelerator.clip_grad_norm_(trainable_params, args.max_grad_norm)
# If `args.use_deepspeed` is enabled, `total_norm` cannot be logged by accelerator.
# If use_deepspeed, `total_norm` cannot be logged by accelerator.
if not args.use_deepspeed:
accelerator.log({"total_norm": total_norm}, step=global_step)
else:
if hasattr(optimizer, "optimizer") and hasattr(optimizer.optimizer, "_global_grad_norm"):
accelerator.log({"total_norm": optimizer.optimizer._global_grad_norm}, step=global_step)
optimizer.step()
lr_scheduler.step()
optimizer.zero_grad()
@@ -1700,11 +1649,15 @@ def main():
global_step,
splitted_prompts_idx
)
avg_validation_loss = accelerator.gather(validation_loss).mean()
avg_validation_reward = accelerator.gather(validation_reward).mean()
accelerator.print(avg_validation_loss, avg_validation_reward)
if accelerator.is_main_process:
accelerator.log({"validation_loss": avg_validation_loss, "validation_reward": avg_validation_reward}, step=global_step)
if validation_loss is not None and validation_reward is not None:
avg_validation_loss = accelerator.gather(validation_loss).mean()
avg_validation_reward = accelerator.gather(validation_reward).mean()
accelerator.print(avg_validation_loss, avg_validation_reward)
if accelerator.is_main_process:
accelerator.log(
{"validation_loss": avg_validation_loss, "validation_reward": avg_validation_reward},
step=global_step
)
accelerator.wait_for_everyone()
+46 -4
View File
@@ -1,4 +1,4 @@
export MODEL_NAME="models/Diffusion_Transformer/EasyAnimateV5.1-12b-zh-InP"
export MODEL_NAME="models/Diffusion_Transformer/EasyAnimateV5-12b-zh-InP"
export TRAIN_PROMPT_PATH="MovieGenVideoBench_train.txt"
# Performing validation simultaneously with training will increase time and GPU memory usage.
export VALIDATION_PROMPT_PATH="MovieGenVideoBench_val.txt"
@@ -10,7 +10,7 @@ NCCL_DEBUG=INFO
# When train model with multi machines, use "--config_file accelerate.yaml" instead of "--mixed_precision='bf16'".
accelerate launch --num_processes=8 --mixed_precision="bf16" --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json scripts/train_reward_lora.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--config_path="config/easyanimate_video_v5.1_magvit_qwen.yaml" \
--config_path="config/easyanimate_video_v5_magvit_multi_text_encoder.yaml" \
--train_batch_size=1 \
--gradient_accumulation_steps=1 \
--max_train_steps=10000 \
@@ -23,6 +23,8 @@ accelerate launch --num_processes=8 --mixed_precision="bf16" --use_deepspeed --d
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--max_grad_norm=0.3 \
--low_vram \
--use_deepspeed \
--prompt_path=$TRAIN_PROMPT_PATH \
--train_sample_height=256 \
--train_sample_width=256 \
@@ -31,7 +33,47 @@ accelerate launch --num_processes=8 --mixed_precision="bf16" --use_deepspeed --d
--validation_steps=100 \
--validation_batch_size=8 \
--num_decoded_latents=1 \
--use_deepspeed \
--reward_fn="HPSReward" \
--reward_fn_kwargs='{"version": "v2.1"}' \
--backprop
--backprop
# For V5.1
# export MODEL_NAME="models/Diffusion_Transformer/EasyAnimateV5.1-12b-zh-InP"
# export TRAIN_PROMPT_PATH="MovieGenVideoBench_train.txt"
# # Performing validation simultaneously with training will increase time and GPU memory usage.
# export VALIDATION_PROMPT_PATH="MovieGenVideoBench_val.txt"
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
# NCCL_DEBUG=INFO
# # When train model with multi machines, use "--config_file accelerate.yaml" instead of "--mixed_precision='bf16'".
# accelerate launch --num_processes=8 --mixed_precision="bf16" --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json scripts/train_reward_lora.py \
# --pretrained_model_name_or_path=$MODEL_NAME \
# --config_path="config/easyanimate_video_v5.1_magvit_qwen.yaml" \
# --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 \
# --low_vram \
# --use_deepspeed \
# --prompt_path=$TRAIN_PROMPT_PATH \
# --train_sample_height=256 \
# --train_sample_width=256 \
# --video_length=49 \
# --validation_prompt_path=$VALIDATION_PROMPT_PATH \
# --validation_steps=100 \
# --validation_batch_size=8 \
# --num_decoded_latents=1 \
# --reward_fn="HPSReward" \
# --reward_fn_kwargs='{"version": "v2.1"}' \
# --backprop_strategy "tail" \
# --backprop_num_steps 10