update reward training
This commit is contained in:
@@ -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
@@ -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()
|
||||
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user