From 2114d906df8324b4c78fee61ce6f02a384c1f699 Mon Sep 17 00:00:00 2001 From: hkz Date: Fri, 22 Nov 2024 09:45:59 +0800 Subject: [PATCH] Add Reward LoRA Training (#71) --- README.md | 43 +- README_zh-CN.md | 43 +- cogvideox/models/autoencoder_magvit.py | 10 +- cogvideox/reward/MPS/README.md | 1 + .../reward/MPS/trainer/models/base_model.py | 7 + .../reward/MPS/trainer/models/clip_model.py | 154 ++ .../MPS/trainer/models/cross_modeling.py | 292 ++++ .../aesthetic_predictor_v2_5/__init__.py | 13 + .../aesthetic_predictor_v2_5/siglip_v2_5.py | 133 ++ .../reward/improved_aesthetic_predictor.py | 49 + cogvideox/reward/reward_fn.py | 382 +++++ scripts/README_TRAIN_REWARD.md | 258 ++++ scripts/train_reward_lora.py | 1300 +++++++++++++++++ scripts/train_reward_lora.sh | 62 + 14 files changed, 2740 insertions(+), 7 deletions(-) create mode 100644 cogvideox/reward/MPS/README.md create mode 100644 cogvideox/reward/MPS/trainer/models/base_model.py create mode 100644 cogvideox/reward/MPS/trainer/models/clip_model.py create mode 100644 cogvideox/reward/MPS/trainer/models/cross_modeling.py create mode 100644 cogvideox/reward/aesthetic_predictor_v2_5/__init__.py create mode 100644 cogvideox/reward/aesthetic_predictor_v2_5/siglip_v2_5.py create mode 100644 cogvideox/reward/improved_aesthetic_predictor.py create mode 100644 cogvideox/reward/reward_fn.py create mode 100644 scripts/README_TRAIN_REWARD.md create mode 100644 scripts/train_reward_lora.py create mode 100644 scripts/train_reward_lora.sh diff --git a/README.md b/README.md index 053455f..d801990 100644 --- a/README.md +++ b/README.md @@ -23,7 +23,7 @@ CogVideoX-Fun is a modified pipeline based on the CogVideoX structure, designed We will support quick pull-ups from different platforms, refer to [Quick Start](#quick-start). What's New: -- Upload a new version of the control model that supports different control conditions such as Canny, Depth, Pose, MLSD, etc. [2024.11.16] +- Use reinforcement learning with reward backpropagation to train Lora and optimize the video, aligning it better with human preferences, detailes in [here](scripts/README_TRAIN_REWARD.md). A new version of the control model supports various conditions (e.g., Canny, Depth, Pose, MLSD, etc.). [2024.11.21] - CogVideoX-Fun Control is now supported in diffusers. Thanks to [a-r-r-o-w](https://github.com/a-r-r-o-w) who contributed the support in this [PR](https://github.com/huggingface/diffusers/pull/9671). Check out the [docs](https://huggingface.co/docs/diffusers/main/en/api/pipelines/cogvideox) to know more. [ 2024.10.16 ] - Retrain the i2v model and add noise to increase the motion amplitude of the video. Upload the control model training code and control model. [ 2024.09.29 ] - Create code! Now supporting Windows and Linux. Supports 2b and 5b models. Supports video generation at any resolution from 256x256x49 to 1024x1024x49. [ 2024.09.18 ] @@ -175,6 +175,47 @@ Resolution-512 +### CogVideoX-Fun-V1.1-5B with Reward Backpropagation + + + + + + + + + + + + + + + + + + + + + + +
PromptCogVideoX-Fun-V1.1-5BCogVideoX-Fun-V1.1-5B
HPSv2.1 Reward LoRA
CogVideoX-Fun-V1.1-5B
MPS Reward LoRA
+ Pig with wings flying above a diamond mountain + + + + + + +
+ A dog runs through a field while a cat climbs a tree + + + + + + +
+ ### CogVideoX-Fun-V1.1-5B-Control diff --git a/README_zh-CN.md b/README_zh-CN.md index 10d0604..88d8f41 100644 --- a/README_zh-CN.md +++ b/README_zh-CN.md @@ -23,7 +23,7 @@ CogVideoX-Fun是一个基于CogVideoX结构修改后的的pipeline,是一个 我们会逐渐支持从不同平台快速启动,请参阅 [快速启动](#快速启动)。 新特性: -- 上传新版本的控制模型,支持不同的控制条件,如Canny、Depth、Pose、MLSD等。[2024.11.16] +- 通过奖励反向传播技术训练Lora,以优化生成的视频,使其更好地与人类偏好保持一致,[更多信息](scripts/README_TRAIN_REWARD.md)。新版本的控制模型,支持不同的控制条件,如Canny、Depth、Pose、MLSD等。[2024.11.21] - CogVideoX-Fun Control现在在diffusers中得到了支持。感谢 [a-r-r-o-w](https://github.com/a-r-r-o-w)在这个 [PR](https://github.com/huggingface/diffusers/pull/9671)中贡献了支持。查看[文档](https://huggingface.co/docs/diffusers/main/en/api/pipelines/cogvideox)以了解更多信息。[2024.10.16] - 重新训练i2v模型,添加Noise,使得视频的运动幅度更大。上传控制模型训练代码与Control模型。[ 2024.09.29 ] - 创建代码!现在支持 Windows 和 Linux。支持2b与5b最大256x256x49到1024x1024x49的任意分辨率的视频生成。[ 2024.09.18 ] @@ -173,6 +173,47 @@ Resolution-512
+### CogVideoX-Fun-V1.1-5B with Reward Backpropagation + + + + + + + + + + + + + + + + + + + + + + +
PromptCogVideoX-Fun-V1.1-5BCogVideoX-Fun-V1.1-5B
HPSv2.1 Reward LoRA
CogVideoX-Fun-V1.1-5B
MPS Reward LoRA
+ Pig with wings flying above a diamond mountain + + + + + + +
+ A dog runs through a field while a cat climbs a tree + + + + + + +
+ ### CogVideoX-Fun-V1.1-5B-Control diff --git a/cogvideox/models/autoencoder_magvit.py b/cogvideox/models/autoencoder_magvit.py index 9c2b906..a1ac2ec 100644 --- a/cogvideox/models/autoencoder_magvit.py +++ b/cogvideox/models/autoencoder_magvit.py @@ -394,7 +394,7 @@ class CogVideoXDownBlock3D(nn.Module): zq: Optional[torch.Tensor] = None, ) -> torch.Tensor: for resnet in self.resnets: - if self.training and self.gradient_checkpointing: + if torch.is_grad_enabled() and self.gradient_checkpointing: def create_custom_forward(module): def create_forward(*inputs): @@ -482,7 +482,7 @@ class CogVideoXMidBlock3D(nn.Module): zq: Optional[torch.Tensor] = None, ) -> torch.Tensor: for resnet in self.resnets: - if self.training and self.gradient_checkpointing: + if torch.is_grad_enabled() and self.gradient_checkpointing: def create_custom_forward(module): def create_forward(*inputs): @@ -587,7 +587,7 @@ class CogVideoXUpBlock3D(nn.Module): ) -> torch.Tensor: r"""Forward method of the `CogVideoXUpBlock3D` class.""" for resnet in self.resnets: - if self.training and self.gradient_checkpointing: + if torch.is_grad_enabled() and self.gradient_checkpointing: def create_custom_forward(module): def create_forward(*inputs): @@ -709,7 +709,7 @@ class CogVideoXEncoder3D(nn.Module): r"""The forward method of the `CogVideoXEncoder3D` class.""" hidden_states = self.conv_in(sample) - if self.training and self.gradient_checkpointing: + if torch.is_grad_enabled() and self.gradient_checkpointing: def create_custom_forward(module): def custom_forward(*inputs): @@ -850,7 +850,7 @@ class CogVideoXDecoder3D(nn.Module): r"""The forward method of the `CogVideoXDecoder3D` class.""" hidden_states = self.conv_in(sample) - if self.training and self.gradient_checkpointing: + if torch.is_grad_enabled() and self.gradient_checkpointing: def create_custom_forward(module): def custom_forward(*inputs): diff --git a/cogvideox/reward/MPS/README.md b/cogvideox/reward/MPS/README.md new file mode 100644 index 0000000..d66d2ee --- /dev/null +++ b/cogvideox/reward/MPS/README.md @@ -0,0 +1 @@ +This folder is modified from the official [MPS](https://github.com/Kwai-Kolors/MPS/tree/main) repository. \ No newline at end of file diff --git a/cogvideox/reward/MPS/trainer/models/base_model.py b/cogvideox/reward/MPS/trainer/models/base_model.py new file mode 100644 index 0000000..df7907f --- /dev/null +++ b/cogvideox/reward/MPS/trainer/models/base_model.py @@ -0,0 +1,7 @@ +from dataclasses import dataclass + + + +@dataclass +class BaseModelConfig: + pass \ No newline at end of file diff --git a/cogvideox/reward/MPS/trainer/models/clip_model.py b/cogvideox/reward/MPS/trainer/models/clip_model.py new file mode 100644 index 0000000..003bb5d --- /dev/null +++ b/cogvideox/reward/MPS/trainer/models/clip_model.py @@ -0,0 +1,154 @@ +from dataclasses import dataclass +from transformers import CLIPModel as HFCLIPModel +from transformers import AutoTokenizer + +from torch import nn, einsum + +# Modified: import +# from trainer.models.base_model import BaseModelConfig +from .base_model import BaseModelConfig + +from transformers import CLIPConfig +from typing import Any, Optional, Tuple, Union +import torch + +# Modified: import +# from trainer.models.cross_modeling import Cross_model +from .cross_modeling import Cross_model + +import gc + +class XCLIPModel(HFCLIPModel): + def __init__(self, config: CLIPConfig): + super().__init__(config) + + def get_text_features( + self, + input_ids: Optional[torch.Tensor] = None, + attention_mask: Optional[torch.Tensor] = None, + position_ids: Optional[torch.Tensor] = None, + output_attentions: Optional[bool] = None, + output_hidden_states: Optional[bool] = None, + return_dict: Optional[bool] = None, + ) -> torch.FloatTensor: + + # Use CLIP model's config for some fields (if specified) instead of those of vision & text components. + output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions + output_hidden_states = ( + output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states + ) + return_dict = return_dict if return_dict is not None else self.config.use_return_dict + + text_outputs = self.text_model( + input_ids=input_ids, + attention_mask=attention_mask, + position_ids=position_ids, + output_attentions=output_attentions, + output_hidden_states=output_hidden_states, + return_dict=return_dict, + ) + + # pooled_output = text_outputs[1] + # text_features = self.text_projection(pooled_output) + last_hidden_state = text_outputs[0] + text_features = self.text_projection(last_hidden_state) + + pooled_output = text_outputs[1] + text_features_EOS = self.text_projection(pooled_output) + + + # del last_hidden_state, text_outputs + # gc.collect() + + return text_features, text_features_EOS + + def get_image_features( + self, + pixel_values: Optional[torch.FloatTensor] = None, + output_attentions: Optional[bool] = None, + output_hidden_states: Optional[bool] = None, + return_dict: Optional[bool] = None, + ) -> torch.FloatTensor: + + # Use CLIP model's config for some fields (if specified) instead of those of vision & text components. + output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions + output_hidden_states = ( + output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states + ) + return_dict = return_dict if return_dict is not None else self.config.use_return_dict + + vision_outputs = self.vision_model( + pixel_values=pixel_values, + output_attentions=output_attentions, + output_hidden_states=output_hidden_states, + return_dict=return_dict, + ) + + # pooled_output = vision_outputs[1] # pooled_output + # image_features = self.visual_projection(pooled_output) + last_hidden_state = vision_outputs[0] + image_features = self.visual_projection(last_hidden_state) + + return image_features + + + +@dataclass +class ClipModelConfig(BaseModelConfig): + _target_: str = "trainer.models.clip_model.CLIPModel" + pretrained_model_name_or_path: str ="openai/clip-vit-base-patch32" + + +class CLIPModel(nn.Module): + def __init__(self, config): + super().__init__() + # Modified: We convert the original ckpt (contains the entire model) to a `state_dict`. + # self.model = XCLIPModel.from_pretrained(ckpt) + self.model = XCLIPModel(config) + self.cross_model = Cross_model(dim=1024, layer_num=4, heads=16) + + def get_text_features(self, *args, **kwargs): + return self.model.get_text_features(*args, **kwargs) + + def get_image_features(self, *args, **kwargs): + return self.model.get_image_features(*args, **kwargs) + + def forward(self, text_inputs=None, image_inputs=None, condition_inputs=None): + outputs = () + + text_f, text_EOS = self.model.get_text_features(text_inputs) # B*77*1024 + outputs += text_EOS, + + image_f = self.model.get_image_features(image_inputs.half()) # 2B*257*1024 + # [B, 77, 1024] + condition_f, _ = self.model.get_text_features(condition_inputs) # B*5*1024 + + sim_text_condition = einsum('b i d, b j d -> b j i', text_f, condition_f) + sim_text_condition = torch.max(sim_text_condition, dim=1, keepdim=True)[0] + sim_text_condition = sim_text_condition / sim_text_condition.max() + mask = torch.where(sim_text_condition > 0.01, 0, float('-inf')) # B*1*77 + + # Modified: Support both torch.float16 and torch.bfloat16 + # mask = mask.repeat(1,image_f.shape[1],1) # B*257*77 + model_dtype = next(self.cross_model.parameters()).dtype + mask = mask.repeat(1,image_f.shape[1],1).to(model_dtype) # B*257*77 + # bc = int(image_f.shape[0]/2) + + # Modified: The original input consists of a (batch of) text and two (batches of) images, + # primarily used to compute which (batch of) image is more consistent with the text. + # The modified input consists of a (batch of) text and a (batch of) images. + # sim0 = self.cross_model(image_f[:bc,:,:], text_f,mask.half()) + # sim1 = self.cross_model(image_f[bc:,:,:], text_f,mask.half()) + # outputs += sim0[:,0,:], + # outputs += sim1[:,0,:], + sim = self.cross_model(image_f, text_f,mask) + outputs += sim[:,0,:], + + return outputs + + @property + def logit_scale(self): + return self.model.logit_scale + + def save(self, path): + self.model.save_pretrained(path) diff --git a/cogvideox/reward/MPS/trainer/models/cross_modeling.py b/cogvideox/reward/MPS/trainer/models/cross_modeling.py new file mode 100644 index 0000000..6822329 --- /dev/null +++ b/cogvideox/reward/MPS/trainer/models/cross_modeling.py @@ -0,0 +1,292 @@ +import torch +from torch import einsum, nn +import torch.nn.functional as F +from einops import rearrange, repeat + +# helper functions + +def exists(val): + return val is not None + +def default(val, d): + return val if exists(val) else d + +# normalization +# they use layernorm without bias, something that pytorch does not offer + + +class LayerNorm(nn.Module): + def __init__(self, dim): + super().__init__() + self.weight = nn.Parameter(torch.ones(dim)) + self.register_buffer("bias", torch.zeros(dim)) + + def forward(self, x): + return F.layer_norm(x, x.shape[-1:], self.weight, self.bias) + +# residual + + +class Residual(nn.Module): + def __init__(self, fn): + super().__init__() + self.fn = fn + + def forward(self, x, *args, **kwargs): + return self.fn(x, *args, **kwargs) + x + + +# rotary positional embedding +# https://arxiv.org/abs/2104.09864 + + +class RotaryEmbedding(nn.Module): + def __init__(self, dim): + super().__init__() + inv_freq = 1.0 / (10000 ** (torch.arange(0, dim, 2).float() / dim)) + self.register_buffer("inv_freq", inv_freq) + + def forward(self, max_seq_len, *, device): + seq = torch.arange(max_seq_len, device=device, dtype=self.inv_freq.dtype) + freqs = einsum("i , j -> i j", seq, self.inv_freq) + return torch.cat((freqs, freqs), dim=-1) + + +def rotate_half(x): + x = rearrange(x, "... (j d) -> ... j d", j=2) + x1, x2 = x.unbind(dim=-2) + return torch.cat((-x2, x1), dim=-1) + + +def apply_rotary_pos_emb(pos, t): + return (t * pos.cos()) + (rotate_half(t) * pos.sin()) + + +# classic Noam Shazeer paper, except here they use SwiGLU instead of the more popular GEGLU for gating the feedforward +# https://arxiv.org/abs/2002.05202 + + +class SwiGLU(nn.Module): + def forward(self, x): + x, gate = x.chunk(2, dim=-1) + return F.silu(gate) * x + + +# parallel attention and feedforward with residual +# discovered by Wang et al + EleutherAI from GPT-J fame + +class ParallelTransformerBlock(nn.Module): + def __init__(self, dim, dim_head=64, heads=8, ff_mult=4): + super().__init__() + self.norm = LayerNorm(dim) + + attn_inner_dim = dim_head * heads + ff_inner_dim = dim * ff_mult + self.fused_dims = (attn_inner_dim, dim_head, dim_head, (ff_inner_dim * 2)) + + self.heads = heads + self.scale = dim_head**-0.5 + self.rotary_emb = RotaryEmbedding(dim_head) + + self.fused_attn_ff_proj = nn.Linear(dim, sum(self.fused_dims), bias=False) + self.attn_out = nn.Linear(attn_inner_dim, dim, bias=False) + + self.ff_out = nn.Sequential( + SwiGLU(), + nn.Linear(ff_inner_dim, dim, bias=False) + ) + + self.register_buffer("pos_emb", None, persistent=False) + + + def get_rotary_embedding(self, n, device): + if self.pos_emb is not None and self.pos_emb.shape[-2] >= n: + return self.pos_emb[:n] + + pos_emb = self.rotary_emb(n, device=device) + self.register_buffer("pos_emb", pos_emb, persistent=False) + return pos_emb + + def forward(self, x, attn_mask=None): + """ + einstein notation + b - batch + h - heads + n, i, j - sequence length (base sequence length, source, target) + d - feature dimension + """ + + n, device, h = x.shape[1], x.device, self.heads + + # pre layernorm + + x = self.norm(x) + + # attention queries, keys, values, and feedforward inner + + q, k, v, ff = self.fused_attn_ff_proj(x).split(self.fused_dims, dim=-1) + + # split heads + # they use multi-query single-key-value attention, yet another Noam Shazeer paper + # they found no performance loss past a certain scale, and more efficient decoding obviously + # https://arxiv.org/abs/1911.02150 + + q = rearrange(q, "b n (h d) -> b h n d", h=h) + + # rotary embeddings + + positions = self.get_rotary_embedding(n, device) + q, k = map(lambda t: apply_rotary_pos_emb(positions, t), (q, k)) + + # scale + + q = q * self.scale + + # similarity + + sim = einsum("b h i d, b j d -> b h i j", q, k) + + + # extra attention mask - for masking out attention from text CLS token to padding + + if exists(attn_mask): + attn_mask = rearrange(attn_mask, 'b i j -> b 1 i j') + sim = sim.masked_fill(~attn_mask, -torch.finfo(sim.dtype).max) + + # attention + + sim = sim - sim.amax(dim=-1, keepdim=True).detach() + attn = sim.softmax(dim=-1) + + # aggregate values + + out = einsum("b h i j, b j d -> b h i d", attn, v) + + # merge heads + + out = rearrange(out, "b h n d -> b n (h d)") + return self.attn_out(out) + self.ff_out(ff) + +# cross attention - using multi-query + one-headed key / values as in PaLM w/ optional parallel feedforward + +class CrossAttention(nn.Module): + def __init__( + self, + dim, + *, + context_dim=None, + dim_head=64, + heads=12, + parallel_ff=False, + ff_mult=4, + norm_context=False + ): + super().__init__() + self.heads = heads + self.scale = dim_head ** -0.5 + inner_dim = heads * dim_head + context_dim = default(context_dim, dim) + + self.norm = LayerNorm(dim) + self.context_norm = LayerNorm(context_dim) if norm_context else nn.Identity() + + self.to_q = nn.Linear(dim, inner_dim, bias=False) + self.to_kv = nn.Linear(context_dim, dim_head * 2, bias=False) + self.to_out = nn.Linear(inner_dim, dim, bias=False) + + # whether to have parallel feedforward + + ff_inner_dim = ff_mult * dim + + self.ff = nn.Sequential( + nn.Linear(dim, ff_inner_dim * 2, bias=False), + SwiGLU(), + nn.Linear(ff_inner_dim, dim, bias=False) + ) if parallel_ff else None + + def forward(self, x, context, mask): + """ + einstein notation + b - batch + h - heads + n, i, j - sequence length (base sequence length, source, target) + d - feature dimension + """ + + # pre-layernorm, for queries and context + + x = self.norm(x) + context = self.context_norm(context) + + # get queries + + q = self.to_q(x) + q = rearrange(q, 'b n (h d) -> b h n d', h = self.heads) + + # scale + + q = q * self.scale + + # get key / values + + k, v = self.to_kv(context).chunk(2, dim=-1) + + # query / key similarity + + sim = einsum('b h i d, b j d -> b h i j', q, k) + + # attention + mask = mask.unsqueeze(1).repeat(1,self.heads,1,1) + sim = sim + mask # context mask + sim = sim - sim.amax(dim=-1, keepdim=True) + attn = sim.softmax(dim=-1) + + # aggregate + + out = einsum('b h i j, b j d -> b h i d', attn, v) + + # merge and combine heads + + out = rearrange(out, 'b h n d -> b n (h d)') + out = self.to_out(out) + + # add parallel feedforward (for multimodal layers) + + if exists(self.ff): + out = out + self.ff(x) + + return out + + +class Cross_model(nn.Module): + def __init__( + self, + dim=512, + layer_num=4, + dim_head=64, + heads=8, + ff_mult=4 + ): + super().__init__() + + self.layers = nn.ModuleList([]) + + + for ind in range(layer_num): + self.layers.append(nn.ModuleList([ + Residual(CrossAttention(dim=dim, dim_head=dim_head, heads=heads, parallel_ff=True, ff_mult=ff_mult)), + Residual(ParallelTransformerBlock(dim=dim, dim_head=dim_head, heads=heads, ff_mult=ff_mult)) + ])) + + def forward( + self, + query_tokens, + context_tokens, + mask + ): + print(mask.dtype) + for cross_attn, self_attn_ff in self.layers: + query_tokens = cross_attn(query_tokens, context_tokens,mask) + query_tokens = self_attn_ff(query_tokens) + + return query_tokens \ No newline at end of file diff --git a/cogvideox/reward/aesthetic_predictor_v2_5/__init__.py b/cogvideox/reward/aesthetic_predictor_v2_5/__init__.py new file mode 100644 index 0000000..2d3d8f1 --- /dev/null +++ b/cogvideox/reward/aesthetic_predictor_v2_5/__init__.py @@ -0,0 +1,13 @@ +from .siglip_v2_5 import ( + AestheticPredictorV2_5Head, + AestheticPredictorV2_5Model, + AestheticPredictorV2_5Processor, + convert_v2_5_from_siglip, +) + +__all__ = [ + "AestheticPredictorV2_5Head", + "AestheticPredictorV2_5Model", + "AestheticPredictorV2_5Processor", + "convert_v2_5_from_siglip", +] \ No newline at end of file diff --git a/cogvideox/reward/aesthetic_predictor_v2_5/siglip_v2_5.py b/cogvideox/reward/aesthetic_predictor_v2_5/siglip_v2_5.py new file mode 100644 index 0000000..867f429 --- /dev/null +++ b/cogvideox/reward/aesthetic_predictor_v2_5/siglip_v2_5.py @@ -0,0 +1,133 @@ +# Borrowed from https://github.com/discus0434/aesthetic-predictor-v2-5/blob/3125a9e/src/aesthetic_predictor_v2_5/siglip_v2_5.py +import os +from collections import OrderedDict +from os import PathLike +from typing import Final + +import torch +import torch.nn as nn +import torchvision.transforms as transforms +from transformers import ( + SiglipImageProcessor, + SiglipVisionConfig, + SiglipVisionModel, + logging, +) +from transformers.image_processing_utils import BatchFeature +from transformers.modeling_outputs import ImageClassifierOutputWithNoAttention + +logging.set_verbosity_error() + +URL: Final[str] = ( + "https://github.com/discus0434/aesthetic-predictor-v2-5/raw/main/models/aesthetic_predictor_v2_5.pth" +) + + +class AestheticPredictorV2_5Head(nn.Module): + def __init__(self, config: SiglipVisionConfig) -> None: + super().__init__() + self.scoring_head = nn.Sequential( + nn.Linear(config.hidden_size, 1024), + nn.Dropout(0.5), + nn.Linear(1024, 128), + nn.Dropout(0.5), + nn.Linear(128, 64), + nn.Dropout(0.5), + nn.Linear(64, 16), + nn.Dropout(0.2), + nn.Linear(16, 1), + ) + + def forward(self, image_embeds: torch.Tensor) -> torch.Tensor: + return self.scoring_head(image_embeds) + + +class AestheticPredictorV2_5Model(SiglipVisionModel): + PATCH_SIZE = 14 + + def __init__(self, config: SiglipVisionConfig, *args, **kwargs) -> None: + super().__init__(config, *args, **kwargs) + self.layers = AestheticPredictorV2_5Head(config) + self.post_init() + self.transforms = transforms.Compose([ + transforms.Resize((384, 384)), + transforms.ToTensor(), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]), + ]) + + def forward( + self, + pixel_values: torch.FloatTensor | None = None, + labels: torch.Tensor | None = None, + return_dict: bool | None = None, + ) -> tuple | ImageClassifierOutputWithNoAttention: + return_dict = ( + return_dict if return_dict is not None else self.config.use_return_dict + ) + + outputs = super().forward( + pixel_values=pixel_values, + return_dict=return_dict, + ) + image_embeds = outputs.pooler_output + image_embeds_norm = image_embeds / image_embeds.norm(dim=-1, keepdim=True) + prediction = self.layers(image_embeds_norm) + + loss = None + if labels is not None: + loss_fct = nn.MSELoss() + loss = loss_fct() + + if not return_dict: + return (loss, prediction, image_embeds) + + return ImageClassifierOutputWithNoAttention( + loss=loss, + logits=prediction, + hidden_states=image_embeds, + ) + + +class AestheticPredictorV2_5Processor(SiglipImageProcessor): + def __init__(self, *args, **kwargs) -> None: + super().__init__(*args, **kwargs) + + def __call__(self, *args, **kwargs) -> BatchFeature: + return super().__call__(*args, **kwargs) + + @classmethod + def from_pretrained( + self, + pretrained_model_name_or_path: str + | PathLike = "google/siglip-so400m-patch14-384", + *args, + **kwargs, + ) -> "AestheticPredictorV2_5Processor": + return super().from_pretrained(pretrained_model_name_or_path, *args, **kwargs) + + +def convert_v2_5_from_siglip( + predictor_name_or_path: str | PathLike | None = None, + encoder_model_name: str = "google/siglip-so400m-patch14-384", + *args, + **kwargs, +) -> tuple[AestheticPredictorV2_5Model, AestheticPredictorV2_5Processor]: + model = AestheticPredictorV2_5Model.from_pretrained( + encoder_model_name, *args, **kwargs + ) + + processor = AestheticPredictorV2_5Processor.from_pretrained( + encoder_model_name, *args, **kwargs + ) + + if predictor_name_or_path is None or not os.path.exists(predictor_name_or_path): + state_dict = torch.hub.load_state_dict_from_url(URL, map_location="cpu") + else: + state_dict = torch.load(predictor_name_or_path, map_location="cpu") + + assert isinstance(state_dict, OrderedDict) + + model.layers.load_state_dict(state_dict) + model.eval() + + return model, processor \ No newline at end of file diff --git a/cogvideox/reward/improved_aesthetic_predictor.py b/cogvideox/reward/improved_aesthetic_predictor.py new file mode 100644 index 0000000..43037b9 --- /dev/null +++ b/cogvideox/reward/improved_aesthetic_predictor.py @@ -0,0 +1,49 @@ +import os + +import torch +import torch.nn as nn +from transformers import CLIPModel +from torchvision.datasets.utils import download_url + +URL = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Third_Party/sac%2Blogos%2Bava1-l14-linearMSE.pth" +FILENAME = "sac+logos+ava1-l14-linearMSE.pth" +MD5 = "b1047fd767a00134b8fd6529bf19521a" + + +class MLP(nn.Module): + def __init__(self): + super().__init__() + self.layers = nn.Sequential( + nn.Linear(768, 1024), + nn.Dropout(0.2), + nn.Linear(1024, 128), + nn.Dropout(0.2), + nn.Linear(128, 64), + nn.Dropout(0.1), + nn.Linear(64, 16), + nn.Linear(16, 1), + ) + + + def forward(self, embed): + return self.layers(embed) + + +class ImprovedAestheticPredictor(nn.Module): + def __init__(self, encoder_path="openai/clip-vit-large-patch14", predictor_path=None): + super().__init__() + self.encoder = CLIPModel.from_pretrained(encoder_path) + self.predictor = MLP() + if predictor_path is None or not os.path.exists(predictor_path): + download_url(URL, torch.hub.get_dir(), FILENAME, md5=MD5) + predictor_path = os.path.join(torch.hub.get_dir(), FILENAME) + state_dict = torch.load(predictor_path, map_location="cpu") + self.predictor.load_state_dict(state_dict) + self.eval() + + + def forward(self, pixel_values): + embed = self.encoder.get_image_features(pixel_values=pixel_values) + embed = embed / torch.linalg.vector_norm(embed, dim=-1, keepdim=True) + + return self.predictor(embed).squeeze(1) diff --git a/cogvideox/reward/reward_fn.py b/cogvideox/reward/reward_fn.py new file mode 100644 index 0000000..84795c3 --- /dev/null +++ b/cogvideox/reward/reward_fn.py @@ -0,0 +1,382 @@ +import os +from abc import ABC, abstractmethod + +import torch +import torchvision.transforms as transforms +from einops import rearrange +from torchvision.datasets.utils import download_url +from typing import Optional, Tuple + + +# All reward models. +__all__ = ["AestheticReward", "HPSReward", "PickScoreReward", "MPSReward"] + + +class BaseReward(ABC): + """An base class for reward models. A custom Reward class must implement two functions below. + """ + def __init__(self): + """Define your reward model and image transformations (optional) here. + """ + pass + + @abstractmethod + def __call__(self, batch_frames: torch.Tensor, batch_prompt: Optional[list[str]]=None) -> Tuple[torch.Tensor, torch.Tensor]: + """Given batch frames with shape `[B, C, T, H, W]` extracted from a list of videos and a list of prompts + (optional) correspondingly, return the loss and reward computed by your reward model (reduction by mean). + """ + pass + +class AestheticReward(BaseReward): + """Aesthetic Predictor [V2](https://github.com/christophschuhmann/improved-aesthetic-predictor) + and [V2.5](https://github.com/discus0434/aesthetic-predictor-v2-5) reward model. + """ + def __init__( + self, + encoder_path="openai/clip-vit-large-patch14", + predictor_path=None, + version="v2", + device="cpu", + dtype=torch.float16, + max_reward=10, + loss_scale=0.1, + ): + from .improved_aesthetic_predictor import ImprovedAestheticPredictor + from ..video_caption.utils.siglip_v2_5 import convert_v2_5_from_siglip + + self.encoder_path = encoder_path + self.predictor_path = predictor_path + self.version = version + self.device = device + self.dtype = dtype + self.max_reward = max_reward + self.loss_scale = loss_scale + + if self.version != "v2" and self.version != "v2.5": + raise ValueError("Only v2 and v2.5 are supported.") + if self.version == "v2": + assert "clip-vit-large-patch14" in encoder_path.lower() + self.model = ImprovedAestheticPredictor(encoder_path=self.encoder_path, predictor_path=self.predictor_path) + # https://huggingface.co/openai/clip-vit-large-patch14/blob/main/preprocessor_config.json + self.transform = transforms.Compose([ + transforms.Resize((224, 224), interpolation=transforms.InterpolationMode.BICUBIC), + transforms.Normalize(mean=[0.48145466, 0.4578275, 0.40821073], std=[0.26862954, 0.26130258, 0.27577711]), + ]) + elif self.version == "v2.5": + assert "siglip-so400m-patch14-384" in encoder_path.lower() + self.model, _ = convert_v2_5_from_siglip(encoder_model_name=self.encoder_path) + # https://huggingface.co/google/siglip-so400m-patch14-384/blob/main/preprocessor_config.json + self.transform = transforms.Compose([ + transforms.Resize((384, 384), interpolation=transforms.InterpolationMode.BICUBIC), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]), + ]) + + self.model.to(device=self.device, dtype=self.dtype) + self.model.requires_grad_(False) + + + def __call__(self, batch_frames: torch.Tensor, batch_prompt: Optional[list[str]]=None) -> Tuple[torch.Tensor, torch.Tensor]: + batch_frames = rearrange(batch_frames, "b c t h w -> t b c h w") + batch_loss, batch_reward = 0, 0 + for frames in batch_frames: + pixel_values = torch.stack([self.transform(frame) for frame in frames]) + pixel_values = pixel_values.to(self.device, dtype=self.dtype) + if self.version == "v2": + reward = self.model(pixel_values) + elif self.version == "v2.5": + reward = self.model(pixel_values).logits.squeeze() + # Convert reward to loss in [0, 1]. + if self.max_reward is None: + loss = (-1 * reward) * self.loss_scale + else: + loss = abs(reward - self.max_reward) * self.loss_scale + batch_loss, batch_reward = batch_loss + loss.mean(), batch_reward + reward.mean() + + return batch_loss / batch_frames.shape[0], batch_reward / batch_frames.shape[0] + + +class HPSReward(BaseReward): + """[HPS](https://github.com/tgxs002/HPSv2) v2 and v2.1 reward model. + """ + def __init__( + self, + model_path=None, + version="v2.0", + device="cpu", + dtype=torch.float16, + max_reward=1, + loss_scale=1, + ): + from hpsv2.src.open_clip import create_model_and_transforms, get_tokenizer + + self.model_path = model_path + self.version = version + self.device = device + self.dtype = dtype + self.max_reward = max_reward + self.loss_scale = loss_scale + + self.model, _, _ = create_model_and_transforms( + "ViT-H-14", + "laion2B-s32B-b79K", + precision=self.dtype, + device=self.device, + jit=False, + force_quick_gelu=False, + force_custom_text=False, + force_patch_dropout=False, + force_image_size=None, + pretrained_image=False, + image_mean=None, + image_std=None, + light_augmentation=True, + aug_cfg={}, + output_dict=True, + with_score_predictor=False, + with_region_predictor=False, + ) + self.tokenizer = get_tokenizer("ViT-H-14") + + # https://huggingface.co/laion/CLIP-ViT-H-14-laion2B-s32B-b79K/blob/main/preprocessor_config.json + self.transform = transforms.Compose([ + transforms.Resize((224, 224), interpolation=transforms.InterpolationMode.BICUBIC), + transforms.Normalize(mean=[0.48145466, 0.4578275, 0.40821073], std=[0.26862954, 0.26130258, 0.27577711]), + ]) + + if version == "v2.0": + url = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Third_Party/HPS_v2_compressed.pt" + filename = "HPS_v2_compressed.pt" + md5 = "fd9180de357abf01fdb4eaad64631db4" + elif version == "v2.1": + url = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Third_Party/HPS_v2.1_compressed.pt" + filename = "HPS_v2.1_compressed.pt" + md5 = "4067542e34ba2553a738c5ac6c1d75c0" + else: + raise ValueError("Only v2.0 and v2.1 are supported.") + if self.model_path is None or not os.path.exists(self.model_path): + download_url(url, torch.hub.get_dir(), md5=md5) + model_path = os.path.join(torch.hub.get_dir(), filename) + + state_dict = torch.load(model_path, map_location="cpu")["state_dict"] + self.model.load_state_dict(state_dict) + self.model.to(device=self.device, dtype=self.dtype) + self.model.requires_grad_(False) + self.model.eval() + + def __call__(self, batch_frames: torch.Tensor, batch_prompt: list[str]) -> Tuple[torch.Tensor, torch.Tensor]: + assert batch_frames.shape[0] == len(batch_prompt) + # Compute batch reward and loss in frame-wise. + batch_frames = rearrange(batch_frames, "b c t h w -> t b c h w") + batch_loss, batch_reward = 0, 0 + for frames in batch_frames: + image_inputs = torch.stack([self.transform(frame) for frame in frames]) + image_inputs = image_inputs.to(device=self.device, dtype=self.dtype) + text_inputs = self.tokenizer(batch_prompt).to(device=self.device) + outputs = self.model(image_inputs, text_inputs) + + image_features, text_features = outputs["image_features"], outputs["text_features"] + logits = image_features @ text_features.T + reward = torch.diagonal(logits) + # Convert reward to loss in [0, 1]. + if self.max_reward is None: + loss = (-1 * reward) * self.loss_scale + else: + loss = abs(reward - self.max_reward) * self.loss_scale + + batch_loss, batch_reward = batch_loss + loss.mean(), batch_reward + reward.mean() + + return batch_loss / batch_frames.shape[0], batch_reward / batch_frames.shape[0] + + +class PickScoreReward(BaseReward): + """[PickScore](https://github.com/yuvalkirstain/PickScore) reward model. + """ + def __init__( + self, + model_path="yuvalkirstain/PickScore_v1", + device="cpu", + dtype=torch.float16, + max_reward=1, + loss_scale=1, + ): + from transformers import AutoProcessor, AutoModel + + self.model_path = model_path + self.device = device + self.dtype = dtype + self.max_reward = max_reward + self.loss_scale = loss_scale + + # https://huggingface.co/yuvalkirstain/PickScore_v1/blob/main/preprocessor_config.json + self.transform = transforms.Compose([ + transforms.Resize(224, interpolation=transforms.InterpolationMode.BICUBIC), + transforms.CenterCrop(224), + transforms.Normalize(mean=[0.48145466, 0.4578275, 0.40821073], std=[0.26862954, 0.26130258, 0.27577711]), + ]) + self.processor = AutoProcessor.from_pretrained("laion/CLIP-ViT-H-14-laion2B-s32B-b79K", torch_dtype=self.dtype) + self.model = AutoModel.from_pretrained(model_path, torch_dtype=self.dtype).eval().to(device) + self.model.requires_grad_(False) + self.model.eval() + + def __call__(self, batch_frames: torch.Tensor, batch_prompt: list[str]) -> Tuple[torch.Tensor, torch.Tensor]: + assert batch_frames.shape[0] == len(batch_prompt) + # Compute batch reward and loss in frame-wise. + batch_frames = rearrange(batch_frames, "b c t h w -> t b c h w") + batch_loss, batch_reward = 0, 0 + for frames in batch_frames: + image_inputs = torch.stack([self.transform(frame) for frame in frames]) + image_inputs = image_inputs.to(device=self.device, dtype=self.dtype) + text_inputs = self.processor( + text=batch_prompt, + padding=True, + truncation=True, + max_length=77, + return_tensors="pt", + ).to(self.device) + image_features = self.model.get_image_features(pixel_values=image_inputs) + text_features = self.model.get_text_features(**text_inputs) + image_features = image_features / torch.norm(image_features, dim=-1, keepdim=True) + text_features = text_features / torch.norm(text_features, dim=-1, keepdim=True) + + logits = image_features @ text_features.T + reward = torch.diagonal(logits) + # Convert reward to loss in [0, 1]. + if self.max_reward is None: + loss = (-1 * reward) * self.loss_scale + else: + loss = abs(reward - self.max_reward) * self.loss_scale + + batch_loss, batch_reward = batch_loss + loss.mean(), batch_reward + reward.mean() + + return batch_loss / batch_frames.shape[0], batch_reward / batch_frames.shape[0] + + +class MPSReward(BaseReward): + """[MPS](https://github.com/Kwai-Kolors/MPS) reward model. + """ + def __init__( + self, + model_path=None, + device="cpu", + dtype=torch.float16, + max_reward=1, + loss_scale=1, + ): + from transformers import AutoTokenizer, AutoConfig + from .MPS.trainer.models.clip_model import CLIPModel + + self.model_path = model_path + self.device = device + self.dtype = dtype + self.condition = "light, color, clarity, tone, style, ambiance, artistry, shape, face, hair, hands, limbs, structure, instance, texture, quantity, attributes, position, number, location, word, things." + self.max_reward = max_reward + self.loss_scale = loss_scale + + processor_name_or_path = "laion/CLIP-ViT-H-14-laion2B-s32B-b79K" + # https://huggingface.co/laion/CLIP-ViT-H-14-laion2B-s32B-b79K/blob/main/preprocessor_config.json + self.transform = transforms.Compose([ + transforms.Resize((224, 224), interpolation=transforms.InterpolationMode.BICUBIC), + transforms.Normalize(mean=[0.48145466, 0.4578275, 0.40821073], std=[0.26862954, 0.26130258, 0.27577711]), + ]) + + # We convert the original [ckpt](http://drive.google.com/file/d/17qrK_aJkVNM75ZEvMEePpLj6L867MLkN/view?usp=sharing) + # (contains the entire model) to a `state_dict`. + url = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Third_Party/MPS_overall.pth" + filename = "MPS_overall.pth" + md5 = "1491cbbbd20565747fe07e7572e2ac56" + if self.model_path is None or not os.path.exists(self.model_path): + download_url(url, torch.hub.get_dir(), md5=md5) + model_path = os.path.join(torch.hub.get_dir(), filename) + + self.tokenizer = AutoTokenizer.from_pretrained(processor_name_or_path, trust_remote_code=True) + config = AutoConfig.from_pretrained(processor_name_or_path) + self.model = CLIPModel(config) + state_dict = torch.load(model_path, map_location="cpu") + self.model.load_state_dict(state_dict, strict=False) + self.model.to(device=self.device, dtype=self.dtype) + self.model.requires_grad_(False) + self.model.eval() + + def _tokenize(self, caption): + input_ids = self.tokenizer( + caption, + max_length=self.tokenizer.model_max_length, + padding="max_length", + truncation=True, + return_tensors="pt" + ).input_ids + + return input_ids + + def __call__( + self, + batch_frames: torch.Tensor, + batch_prompt: list[str], + batch_condition: Optional[list[str]] = None + ) -> Tuple[torch.Tensor, torch.Tensor]: + if batch_condition is None: + batch_condition = [self.condition] * len(batch_prompt) + batch_frames = rearrange(batch_frames, "b c t h w -> t b c h w") + batch_loss, batch_reward = 0, 0 + for frames in batch_frames: + image_inputs = torch.stack([self.transform(frame) for frame in frames]) + image_inputs = image_inputs.to(device=self.device, dtype=self.dtype) + text_inputs = self._tokenize(batch_prompt).to(self.device) + condition_inputs = self._tokenize(batch_condition).to(device=self.device) + text_features, image_features = self.model(text_inputs, image_inputs, condition_inputs) + + text_features = text_features / text_features.norm(dim=-1, keepdim=True) + image_features = image_features / image_features.norm(dim=-1, keepdim=True) + # reward = self.model.logit_scale.exp() * torch.diag(torch.einsum('bd,cd->bc', text_features, image_features)) + logits = image_features @ text_features.T + reward = torch.diagonal(logits) + # Convert reward to loss in [0, 1]. + if self.max_reward is None: + loss = (-1 * reward) * self.loss_scale + else: + loss = abs(reward - self.max_reward) * self.loss_scale + + batch_loss, batch_reward = batch_loss + loss.mean(), batch_reward + reward.mean() + + return batch_loss / batch_frames.shape[0], batch_reward / batch_frames.shape[0] + + +if __name__ == "__main__": + import numpy as np + from decord import VideoReader + + video_path_list = ["your_video_path_1.mp4", "your_video_path_2.mp4"] + prompt_list = ["your_prompt_1", "your_prompt_2"] + num_sampled_frames = 8 + + to_tensor = transforms.ToTensor() + + sampled_frames_list = [] + for video_path in video_path_list: + vr = VideoReader(video_path) + sampled_frame_indices = np.linspace(0, len(vr), num_sampled_frames, endpoint=False, dtype=int) + sampled_frames = vr.get_batch(sampled_frame_indices).asnumpy() + sampled_frames = torch.stack([to_tensor(frame) for frame in sampled_frames]) + 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") + + aesthetic_reward_v2 = AestheticReward(device="cuda", dtype=torch.bfloat16) + print(f"aesthetic_reward_v2: {aesthetic_reward_v2(sampled_frames)}") + + aesthetic_reward_v2_5 = AestheticReward( + encoder_path="google/siglip-so400m-patch14-384", version="v2.5", device="cuda", dtype=torch.bfloat16 + ) + print(f"aesthetic_reward_v2_5: {aesthetic_reward_v2_5(sampled_frames)}") + + hps_reward_v2 = HPSReward(device="cuda", dtype=torch.bfloat16) + print(f"hps_reward_v2: {hps_reward_v2(sampled_frames, prompt_list)}") + + hps_reward_v2_1 = HPSReward(version="v2.1", device="cuda", dtype=torch.bfloat16) + print(f"hps_reward_v2_1: {hps_reward_v2_1(sampled_frames, prompt_list)}") + + pick_score = PickScoreReward(device="cuda", dtype=torch.bfloat16) + print(f"pick_score_reward: {pick_score(sampled_frames, prompt_list)}") + + mps_score = MPSReward(device="cuda", dtype=torch.bfloat16) + print(f"mps_reward: {mps_score(sampled_frames, prompt_list)}") \ No newline at end of file diff --git a/scripts/README_TRAIN_REWARD.md b/scripts/README_TRAIN_REWARD.md new file mode 100644 index 0000000..3b4bc07 --- /dev/null +++ b/scripts/README_TRAIN_REWARD.md @@ -0,0 +1,258 @@ +# Enhance CogVideoX-Fun with Reward Backpropagation (Preference Optimization) +We explore the Reward Backpropagation technique [1](#ref1) [2](#ref2) to optimized the generated videos by [CogVideoX-Fun-V1.1](https://github.com/aigc-apps/CogVideoX-Fun) 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. + +- [Enhance CogVideoX-Fun with Reward Backpropagation (Preference Optimization)](#enhance-cogvideox-fun-with-reward-backpropagation-preference-optimization) + - [Demo](#demo) + - [CogVideoX-Fun-V1.1-5B](#cogvideox-fun-v11-5b) + - [CogVideoX-Fun-V1.1-2B](#cogvideox-fun-v11-2b) + - [Model Zoo](#model-zoo) + - [Inference](#inference) + - [Training](#training) + - [Setup](#setup) + - [Important Args](#important-args) + - [Limitations](#limitations) + - [References](#references) + + +## Demo +### CogVideoX-Fun-V1.1-5B + +
+ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
PromptCogVideoX-Fun-V1.1-5BCogVideoX-Fun-V1.1-5B
HPSv2.1 Reward LoRA
CogVideoX-Fun-V1.1-5B
MPS Reward LoRA
+ Pig with wings flying above a diamond mountain + + + + + + +
+ A dog runs through a field while a cat climbs a tree + + + + + + +
+ Crystal cake shimmering beside a metal apple + + + + + + +
+ Elderly artist with a white beard painting on a white canvas + + + + + + +
+ +### CogVideoX-Fun-V1.1-2B + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
PromptCogVideoX-Fun-V1.1-2BCogVideoX-Fun-V1.1-2B
HPSv2.1 Reward LoRA
CogVideoX-Fun-V1.1-2B
MPS Reward LoRA
+ A blue car drives past a white picket fence on a sunny day + + + + + + +
+ Blue jay swooping near a red maple tree + + + + + + +
+ Yellow curtains swaying near a blue sofa + + + + + + +
+ White tractor plowing near a green farmhouse + + + + + + +
+ +> [!NOTE] +> The above test prompts are from VBench. All videos are generated with lora weight 0.7. + +## Model Zoo +| Name | Base Model | Reward Model | Hugging Face | Description | +|--|--|--|--|--| +| CogVideoX-Fun-V1.1-5b-InP-HPS2.1.safetensors | CogVideoX-Fun-V1.1-5b | [HPS v2.1](https://github.com/tgxs002/HPSv2) | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-Reward-LoRAs/resolve/main/CogVideoX-Fun-V1.1-5b-InP-HPS2.1.safetensors) | Official HPS v2.1 reward LoRA (`rank=128` and `network_alpha=64`) for CogVideoX-Fun-V1.1-5b-InP. It is trained with a batch size of 8 for 1,500 steps.| +| CogVideoX-Fun-V1.1-2b-InP-HPS2.1.safetensors | CogVideoX-Fun-V1.1-2b | [HPS v2.1](https://github.com/tgxs002/HPSv2) | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-Reward-LoRAs/resolve/main/CogVideoX-Fun-V1.1-2b-InP-HPS2.1.safetensors) | Official HPS v2.1 reward LoRA (`rank=128` and `network_alpha=64`) for CogVideoX-Fun-V1.1-2b-InP. It is trained with a batch size of 8 for 3,000 steps.| +| CogVideoX-Fun-V1.1-5b-InP-MPS.safetensors | CogVideoX-Fun-V1.1-5b | [MPS](https://github.com/Kwai-Kolors/MPS) | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-Reward-LoRAs/resolve/main/CogVideoX-Fun-V1.1-5b-InP-MPS.safetensors) | Official MPS reward LoRA (`rank=128` and `network_alpha=64`) for CogVideoX-Fun-V1.1-5b-InP. It is trained with a batch size of 8 for 5,500 steps.| +| CogVideoX-Fun-V1.1-2b-InP-MPS.safetensors | CogVideoX-Fun-V1.1-2b | [MPS](https://github.com/Kwai-Kolors/MPS) | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-Reward-LoRAs/resolve/main/CogVideoX-Fun-V1.1-2b-InP-MPS.safetensors) | Official MPS reward LoRA (`rank=128` and `network_alpha=64`) for CogVideoX-Fun-V1.1-2b-InP. It is trained with a batch size of 8 for 16,000 steps.| + +## Inference +We provide an example inference code to run CogVideoX-Fun-V1.1-5b-InP with its HPS2.1 reward LoRA. + +```python +import torch +from diffusers import CogVideoXDDIMScheduler + +from cogvideox.models.transformer3d import CogVideoXTransformer3DModel +from cogvideox.pipeline.pipeline_cogvideox_inpaint import CogVideoX_Fun_Pipeline_Inpaint +from cogvideox.utils.lora_utils import merge_lora +from cogvideox.utils.utils import get_image_to_video_latent, save_videos_grid + +model_path = "alibaba-pai/CogVideoX-Fun-V1.1-5b-InP" +lora_path = "alibaba-pai/CogVideoX-Fun-V1.1-Reward-LoRAs/CogVideoX-Fun-V1.1-5b-InP-HPS2.1.safetensors" +lora_weight = 0.7 + +prompt = "Pig with wings flying above a diamond mountain" +sample_size = [512, 512] +video_length = 49 + +transformer = CogVideoXTransformer3DModel.from_pretrained_2d(model_path, subfolder="transformer").to(torch.bfloat16) +scheduler = CogVideoXDDIMScheduler.from_pretrained(model_path, subfolder="scheduler") +pipeline = CogVideoX_Fun_Pipeline_Inpaint.from_pretrained( + model_path, transformer=transformer, scheduler=scheduler, torch_dtype=torch.bfloat16 +) +pipeline.enable_model_cpu_offload() +pipeline = merge_lora(pipeline, lora_path, lora_weight) + +generator = torch.Generator(device="cuda").manual_seed(42) +input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=sample_size) +sample = pipeline( + prompt, + num_frames = video_length, + negative_prompt = "bad detailed", + height = sample_size[0], + width = sample_size[1], + generator = generator, + guidance_scale = 7.0, + num_inference_steps = 50, + video = input_video, + mask_video = input_video_mask, +).videos + +save_videos_grid(sample, "samples/output.mp4", fps=8) +``` + +## Training +The [training code](./train_reward_lora.py) is based on [train_lora.py](./train_lora.py). +We provide [a shell script](./train_reward_lora.sh) to train the HPS v2.1 reward LoRA for CogVideoX-Fun-V1.1-2b-InP, +which can be trained on a single A10 with 24GB VRAM. To further reduce the VRAM requirement, please read [Important Args](#important-args). + +### Setup +Please read the [quick-start](https://github.com/aigc-apps/CogVideoX-Fun/blob/main/README.md#quick-start) section to setup the CogVideoX-Fun environment. +**If you're playing with HPS reward model**, please run the following script to install the dependencies: +```bash +# For HPS reward model only +pip install hpsv2 +site_packages=$(python -c "import site; print(site.getsitepackages()[0])") +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 +``` + +### Important Args ++ `rank`: The size of LoRA model. The higher the LoRA rank, the more parameters it has, and the more it can learn (including some unnecessary information). +Bt default, we set the rank to 128. You can lower this value to reduce training GPU memory and the LoRA file size. ++ `network_alpha`: A scaling factor changes how the LoRA affect the base model weight. In general, it can be set to half of the `rank`. ++ `prompt_path`: The path to the prompt file (in txt format, each line is a prompt) for sampling training videos. +We randomly selected 701 prompts from [MovieGenBench](https://github.com/facebookresearch/MovieGenBench/blob/main/benchmark/MovieGenVideoBench.txt). ++ `train_sample_height` and `train_sample_width`: The resolution of the sampled training videos. We found +training at a 256x256 resolution can generalize to any other resolution. Reducing the resolution can save GPU memory +during training, but it is recommended that the resolution should be equal to or greater than the image input resolution of the reward model. ++ `reward_fn` and `reward_fn_kwargs`: The reward model name and its keyword arguments. All supported reward models +(Aesthetic Predictor [v2](https://github.com/christophschuhmann/improved-aesthetic-predictor)/[v2.5](https://github.com/discus0434/aesthetic-predictor-v2-5), +[HPS](https://github.com/tgxs002/HPSv2) v2/v2.1, [PickScore](https://github.com/yuvalkirstain/PickScore) and [MPS](https://github.com/Kwai-Kolors/MPS)) +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. + +## 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. + +## References +
    +
  1. Clark, Kevin, et al. "Directly fine-tuning diffusion models on differentiable rewards.". In ICLR 2024.
  2. +
  3. Prabhudesai, Mihir, et al. "Aligning text-to-image diffusion models with reward backpropagation." arXiv preprint arXiv:2310.03739 (2023).
  4. +
\ No newline at end of file diff --git a/scripts/train_reward_lora.py b/scripts/train_reward_lora.py new file mode 100644 index 0000000..6e785eb --- /dev/null +++ b/scripts/train_reward_lora.py @@ -0,0 +1,1300 @@ +"""Modified from CogVideoX-Fun/scripts/train_lora.py +""" +#!/usr/bin/env python +# coding=utf-8 +# Copyright 2024 The HuggingFace Inc. team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and + +import argparse +import gc +import logging +import math +import os +import shutil +import sys +import json +import random +from contextlib import contextmanager +from typing import List, Optional, Union + +import accelerate +import datasets +import diffusers +import numpy as np +import torch +import torch.utils.checkpoint +import torchvision +import transformers +from accelerate import Accelerator +from accelerate.logging import get_logger +from accelerate.state import AcceleratorState +from accelerate.utils import ProjectConfiguration, set_seed +from diffusers import DDIMScheduler, CogVideoXDPMScheduler +from diffusers.models.embeddings import get_3d_rotary_pos_embed +from diffusers.optimization import get_scheduler +from diffusers.utils import check_min_version, deprecate, is_wandb_available +from diffusers.utils.torch_utils import is_compiled_module +from decord import VideoReader +from einops import rearrange +from packaging import version +from tqdm.auto import tqdm +from transformers import T5EncoderModel, T5Tokenizer +from transformers.utils import ContextManagers + +current_file_path = os.path.abspath(__file__) +project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path))] +for project_root in project_roots: + sys.path.insert(0, project_root) if project_root not in sys.path else None + +import cogvideox.reward.reward_fn as reward_fn +from cogvideox.models.autoencoder_magvit import AutoencoderKLCogVideoX +from cogvideox.models.transformer3d import CogVideoXTransformer3DModel +from cogvideox.pipeline.pipeline_cogvideox_inpaint import \ + CogVideoX_Fun_Pipeline_Inpaint, get_resize_crop_region_for_grid +from cogvideox.utils.lora_utils import create_network, merge_lora +from cogvideox.utils.utils import get_image_to_video_latent, save_videos_grid + +if is_wandb_available(): + import wandb + + +# Will error if the minimal version of diffusers is not installed. Remove at your own risks. +check_min_version("0.18.0.dev0") + +logger = get_logger(__name__, log_level="INFO") + + +@contextmanager +def video_reader(*args, **kwargs): + """A context manager to solve the memory leak of decord. + """ + vr = VideoReader(*args, **kwargs) + try: + yield vr + finally: + del vr + gc.collect() + + +def log_validation(vae, text_encoder, tokenizer, transformer3d, network, loss_fn, args, accelerator, weight_dtype, global_step, validation_prompts_idx): + logger.info("Running validation... ") + + transformer3d_val = CogVideoXTransformer3DModel.from_pretrained_2d( + args.pretrained_model_name_or_path, subfolder="transformer", + ).to(weight_dtype) + transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + scheduler = DDIMScheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="scheduler") + + if args.vae_gradient_checkpointing: + # Initialize a new vae if gradient checkpointing is enabled. + vae = AutoencoderKLCogVideoX.from_pretrained( + args.pretrained_model_name_or_path, subfolder="vae", revision=args.revision, variant=args.variant + ).to(weight_dtype) + pipeline = CogVideoX_Fun_Pipeline_Inpaint.from_pretrained( + args.pretrained_model_name_or_path, + vae=vae if args.vae_gradient_checkpointing else accelerator.unwrap_model(vae).to(weight_dtype), + text_encoder=accelerator.unwrap_model(text_encoder), + tokenizer=tokenizer, + transformer=transformer3d_val, + scheduler=scheduler, + torch_dtype=weight_dtype, + ) + + pipeline = pipeline.to(accelerator.device) + pipeline = merge_lora( + pipeline, None, 1, accelerator.device, state_dict=accelerator.unwrap_model(network).state_dict(), transformer_only=True + ) + + to_tensor = torchvision.transforms.ToTensor() + validation_loss, validation_reward = 0, 0 + for i in range(len(validation_prompts_idx)): + validation_idx, validation_prompt = validation_prompts_idx[i] + logger.info(f"Process index: {accelerator.process_index}, validation_idx: {validation_idx}, validation_prompt: {validation_prompt}") + with torch.no_grad(): + with torch.autocast("cuda", dtype=weight_dtype): + video_length = int((args.video_length - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_length != 1 else 1 + input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.validation_sample_height, args.validation_sample_width]) + + if args.seed is None: + generator = None + else: + generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + + sample = pipeline( + validation_prompt, + num_frames = 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, + ).videos + 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) + + del pipeline + del transformer3d_val + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + + return validation_loss, validation_reward + + +def load_prompts(prompt_path, prompt_column="prompt", start_idx=None, end_idx=None): + prompt_list = [] + if prompt_path.endswith(".txt"): + with open(prompt_path, "r") as f: + for line in f: + prompt_list.append(line.strip()) + elif prompt_path.endswith(".jsonl"): + with open(prompt_path, "r") as f: + for line in f.readlines(): + item = json.loads(line) + prompt_list.append(item[prompt_column]) + else: + raise ValueError("The prompt_path must end with .txt or .jsonl.") + prompt_list = prompt_list[start_idx:end_idx] + + return prompt_list + + +# Modified from cogvideox.pipeline.pipeline_cogvideox_inpaint.CogVideoX_Fun_Pipeline_Inpaint._get_t5_prompt_embeds +def get_t5_prompt_embeds( + tokenizer: T5Tokenizer, + text_encoder: T5EncoderModel, + prompt: Union[str, List[str]] = None, + num_videos_per_prompt: int = 1, + max_sequence_length: int = 226, + device: Optional[torch.device] = None, + dtype: Optional[torch.dtype] = None, +): + prompt = [prompt] if isinstance(prompt, str) else prompt + batch_size = len(prompt) + + text_inputs = tokenizer( + prompt, + padding="max_length", + max_length=max_sequence_length, + truncation=True, + add_special_tokens=True, + return_tensors="pt", + ) + text_input_ids = text_inputs.input_ids + untruncated_ids = tokenizer(prompt, padding="longest", return_tensors="pt").input_ids + + if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids): + removed_text = tokenizer.batch_decode(untruncated_ids[:, max_sequence_length - 1 : -1]) + logger.warning( + "The following part of your input was truncated because `max_sequence_length` is set to " + f" {max_sequence_length} tokens: {removed_text}" + ) + + prompt_embeds = text_encoder(text_input_ids.to(device))[0] + prompt_embeds = prompt_embeds.to(dtype=dtype, device=device) + + # duplicate text embeddings for each generation per prompt, using mps friendly method + _, seq_len, _ = prompt_embeds.shape + prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1) + prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt, seq_len, -1) + + return prompt_embeds + + +# Modified from cogvideox.pipeline.pipeline_cogvideox_inpaint.CogVideoX_Fun_Pipeline_Inpaint.encode_prompt +def encode_prompt( + tokenizer: T5Tokenizer, + text_encoder: T5EncoderModel, + prompt: Union[str, List[str]], + negative_prompt: Optional[Union[str, List[str]]] = None, + do_classifier_free_guidance: bool = True, + num_videos_per_prompt: int = 1, + prompt_embeds: Optional[torch.Tensor] = None, + negative_prompt_embeds: Optional[torch.Tensor] = None, + max_sequence_length: int = 226, + device: Optional[torch.device] = None, + dtype: Optional[torch.dtype] = None, +): + r""" + Encodes the prompt into text encoder hidden states. + """ + prompt = [prompt] if isinstance(prompt, str) else prompt + if prompt is not None: + batch_size = len(prompt) + else: + batch_size = prompt_embeds.shape[0] + + if prompt_embeds is None: + prompt_embeds = get_t5_prompt_embeds( + tokenizer, + text_encoder, + prompt=prompt, + num_videos_per_prompt=num_videos_per_prompt, + max_sequence_length=max_sequence_length, + device=device, + dtype=dtype, + ) + + if do_classifier_free_guidance and negative_prompt_embeds is None: + negative_prompt = negative_prompt or "" + negative_prompt = batch_size * [negative_prompt] if isinstance(negative_prompt, str) else negative_prompt + + if prompt is not None and type(prompt) is not type(negative_prompt): + raise TypeError( + f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !=" + f" {type(prompt)}." + ) + elif batch_size != len(negative_prompt): + raise ValueError( + f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:" + f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches" + " the batch size of `prompt`." + ) + + negative_prompt_embeds = get_t5_prompt_embeds( + tokenizer, + text_encoder, + prompt=negative_prompt, + num_videos_per_prompt=num_videos_per_prompt, + max_sequence_length=max_sequence_length, + device=device, + dtype=dtype, + ) + + return prompt_embeds, negative_prompt_embeds + + +# Modified from cogvideox.pipeline.pipeline_cogvideox_inpaint.CogVideoX_Fun_Pipeline_Inpaint.prepare_extra_step_kwargs +def prepare_extra_step_kwargs(scheduler, generator, eta): + # prepare extra kwargs for the scheduler step, since not all schedulers have the same signature + # eta (η) is only used with the DDIMScheduler, it will be ignored for other schedulers. + # eta corresponds to η in DDIM paper: https://arxiv.org/abs/2010.02502 + # and should be between [0, 1] + import inspect + + accepts_eta = "eta" in set(inspect.signature(scheduler.step).parameters.keys()) + extra_step_kwargs = {} + if accepts_eta: + extra_step_kwargs["eta"] = eta + + # check if the scheduler accepts generator + accepts_generator = "generator" in set(inspect.signature(scheduler.step).parameters.keys()) + if accepts_generator: + extra_step_kwargs["generator"] = generator + return extra_step_kwargs + + +# Modified from cogvideox.pipeline.pipeline_cogvideox_inpaint.CogVideoX_Fun_Pipeline_Inpaint._prepare_rotary_positional_embeddings +def prepare_rotary_positional_embeddings( + height: int, + width: int, + num_frames: int, + vae_scale_factor_spatial: int = 8, + patch_size: int = 2, + attention_head_dim: int = 64, + device: torch.device = "cpu" +): + grid_height = height // (vae_scale_factor_spatial * patch_size) + grid_width = width // (vae_scale_factor_spatial * patch_size) + base_size_width = 720 // (vae_scale_factor_spatial * patch_size) + base_size_height = 480 // (vae_scale_factor_spatial * patch_size) + + grid_crops_coords = get_resize_crop_region_for_grid( + (grid_height, grid_width), base_size_width, base_size_height + ) + freqs_cos, freqs_sin = get_3d_rotary_pos_embed( + embed_dim=attention_head_dim, + crops_coords=grid_crops_coords, + grid_size=(grid_height, grid_width), + temporal_size=num_frames, + use_real=True, + ) + + freqs_cos = freqs_cos.to(device=device) + freqs_sin = freqs_sin.to(device=device) + return freqs_cos, freqs_sin + + +def parse_args(): + parser = argparse.ArgumentParser(description="Simple example of a training script.") + parser.add_argument( + "--pretrained_model_name_or_path", + type=str, + default=None, + required=True, + help="Path to pretrained model or model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--revision", + type=str, + default=None, + required=False, + help="Revision of pretrained model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--variant", + type=str, + default=None, + help="Variant of the model files of the pretrained model identifier from huggingface.co/models, 'e.g.' fp16", + ) + parser.add_argument( + "--validation_prompt_path", + type=str, + default=None, + help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), + ) + parser.add_argument( + "--validation_prompts", + type=str, + default=None, + nargs="+", + help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), + ) + parser.add_argument( + "--validation_batch_size", + type=int, + default=1, + help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), + ) + parser.add_argument( + "--validation_sample_height", + type=int, + default=512, + help="The height of sampling videos in validation.", + ) + parser.add_argument( + "--validation_sample_width", + type=int, + default=512, + help="The width of sampling videos in validation.", + ) + parser.add_argument( + "--output_dir", + type=str, + default="sd-model-finetuned", + help="The output directory where the model predictions and checkpoints will be written.", + ) + parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.") + parser.add_argument( + "--use_came", + action="store_true", + help="whether to use came", + ) + parser.add_argument( + "--train_batch_size", type=int, default=16, help="Batch size (per device) for the training dataloader." + ) + parser.add_argument("--num_train_epochs", type=int, default=200) + parser.add_argument( + "--max_train_steps", + type=int, + default=None, + help="Total number of training steps to perform. If provided, overrides num_train_epochs.", + ) + parser.add_argument( + "--gradient_accumulation_steps", + type=int, + default=1, + help="Number of updates steps to accumulate before performing a backward/update pass.", + ) + parser.add_argument( + "--gradient_checkpointing", + action="store_true", + help="Whether or not to use gradient checkpointing (for DiT) to save memory at the expense of slower backward pass.", + ) + parser.add_argument( + "--vae_gradient_checkpointing", + action="store_true", + help="Whether or not to use gradient checkpointing (for VAE) to save memory at the expense of slower backward pass.", + ) + parser.add_argument( + "--learning_rate", + type=float, + default=1e-4, + help="Initial learning rate (after the potential warmup period) to use.", + ) + parser.add_argument( + "--scale_lr", + action="store_true", + default=False, + help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.", + ) + parser.add_argument( + "--lr_scheduler", + type=str, + default="constant", + help=( + 'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",' + ' "constant", "constant_with_warmup"]' + ), + ) + parser.add_argument( + "--lr_warmup_steps", type=int, default=500, help="Number of steps for the warmup in the lr scheduler." + ) + parser.add_argument( + "--use_8bit_adam", action="store_true", help="Whether or not to use 8-bit Adam from bitsandbytes." + ) + parser.add_argument( + "--allow_tf32", + action="store_true", + help=( + "Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see" + " 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("--adam_beta1", type=float, default=0.9, help="The beta1 parameter for the Adam optimizer.") + parser.add_argument("--adam_beta2", type=float, default=0.999, help="The beta2 parameter for the Adam optimizer.") + parser.add_argument("--adam_weight_decay", type=float, default=1e-2, help="Weight decay to use.") + parser.add_argument("--adam_epsilon", type=float, default=1e-08, help="Epsilon value for the Adam optimizer") + parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.") + 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( + "--logging_dir", + type=str, + default="logs", + help=( + "[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to" + " *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***." + ), + ) + parser.add_argument( + "--mixed_precision", + type=str, + default=None, + choices=["no", "fp16", "bf16"], + help=( + "Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >=" + " 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the" + " flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config." + ), + ) + parser.add_argument( + "--report_to", + type=str, + default="tensorboard", + help=( + 'The integration to report the results and logs to. Supported platforms are `"tensorboard"`' + ' (default), `"wandb"` and `"comet_ml"`. Use `"all"` to report to all integrations.' + ), + ) + parser.add_argument("--local_rank", type=int, default=-1, help="For distributed training: local_rank") + parser.add_argument( + "--checkpointing_steps", + type=int, + default=500, + help=( + "Save a checkpoint of the training state every X updates. These checkpoints are only suitable for resuming" + " training using `--resume_from_checkpoint`." + ), + ) + parser.add_argument( + "--checkpoints_total_limit", + type=int, + default=None, + help=("Max number of checkpoints to store."), + ) + parser.add_argument( + "--resume_from_checkpoint", + type=str, + default=None, + help=( + "Whether training should be resumed from a previous checkpoint. Use a path saved by" + ' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.' + ), + ) + parser.add_argument( + "--validation_epochs", + type=int, + default=5, + help="Run validation every X epochs.", + ) + parser.add_argument( + "--validation_steps", + type=int, + default=2000, + help="Run validation every X steps.", + ) + parser.add_argument( + "--tracker_project_name", + type=str, + default="text2image-fine-tune", + help=( + "The `project_name` argument passed to Accelerator.init_trackers for" + " more information see https://huggingface.co/docs/accelerate/v0.17.0/en/package_reference/accelerator#accelerate.Accelerator" + ), + ) + + parser.add_argument( + "--rank", + type=int, + default=128, + help=("The dimension of the LoRA update matrices."), + ) + parser.add_argument( + "--network_alpha", + type=int, + default=64, + help=("The dimension of the LoRA update matrices."), + ) + parser.add_argument( + "--train_text_encoder", + action="store_true", + help="Whether to train the text encoder. If set, the text encoder should be float32 precision.", + ) + parser.add_argument( + "--transformer_path", + type=str, + default=None, + help=("If you want to load the weight from other transformers, input its path."), + ) + parser.add_argument( + "--vae_path", + type=str, + default=None, + help=("If you want to load the weight from other vaes, input its path."), + ) + parser.add_argument("--save_state", action="store_true", help="Whether or not to save state.") + + parser.add_argument( + "--use_deepspeed", action="store_true", help="Whether or not to use deepspeed." + ) + + parser.add_argument( + "--prompt_path", + type=str, + default="normal", + help="The path to the training prompt file.", + ) + parser.add_argument( + '--train_sample_height', + type=int, + default=384, + help='The height of sampling videos in training' + ) + parser.add_argument( + '--train_sample_width', + type=int, + default=672, + help='The width of sampling videos in training' + ) + parser.add_argument( + "--video_length", + type=int, + default=49, + help="The number of frames to generate in training and validation." + ) + parser.add_argument( + '--eta', + type=float, + default=0.0, + help='eta parameter for the DDIM sampler. this controls the amount of noise injected into the sampling process, ' + 'with 0.0 being fully deterministic and 1.0 being equivalent to the DDPM sampler.' + ) + parser.add_argument( + "--guidance_scale", + type=float, + default=6.0, + help="The classifier-free diffusion guidance." + ) + parser.add_argument( + "--num_inference_steps", + type=int, + default=50, + help="The number of denoising steps in training and validation." + ) + parser.add_argument( + "--num_decoded_latents", + type=int, + default=3, + help="The number of latents to be decoded." + ) + parser.add_argument( + "--num_sampled_frames", + type=int, + default=None, + help="The number of sampled frames for the reward function." + ) + parser.add_argument( + "--reward_fn", + type=str, + default="HPSReward", + help='The reward function.' + ) + parser.add_argument( + "--reward_fn_kwargs", + type=str, + default=None, + help='The keyword arguments of the reward function.' + ) + parser.add_argument( + "--backprop", + action="store_true", + default=False, + help="Whether to use the backprop training mode.", + ) + + args = parser.parse_args() + env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) + 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 + + +def main(): + args = parse_args() + + if args.report_to == "wandb" and args.hub_token is not None: + raise ValueError( + "You cannot use both --report_to=wandb and --hub_token due to a security risk of exposing your token." + " 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) + + accelerator_project_config = ProjectConfiguration(project_dir=args.output_dir, logging_dir=logging_dir) + + accelerator = Accelerator( + gradient_accumulation_steps=args.gradient_accumulation_steps, + mixed_precision=args.mixed_precision, + 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( + format="%(asctime)s - %(levelname)s - %(name)s - %(message)s", + datefmt="%m/%d/%Y %H:%M:%S", + level=logging.INFO, + ) + logger.info(accelerator.state, main_process_only=False) + if accelerator.is_local_main_process: + datasets.utils.logging.set_verbosity_warning() + transformers.utils.logging.set_verbosity_warning() + diffusers.utils.logging.set_verbosity_info() + else: + datasets.utils.logging.set_verbosity_error() + transformers.utils.logging.set_verbosity_error() + diffusers.utils.logging.set_verbosity_error() + + # If passed along, set the training seed now. + if args.seed is not None: + set_seed(args.seed, device_specific=True) + + # Handle the repository creation + if accelerator.is_main_process: + if args.output_dir is not None: + os.makedirs(args.output_dir, exist_ok=True) + + # For mixed precision training we cast all non-trainable weigths (vae, non-lora text_encoder and non-lora transformer3d) to half-precision + # as these weights are only used for inference, keeping weights in full precision is not required. + weight_dtype = torch.float32 + if accelerator.mixed_precision == "fp16": + weight_dtype = torch.float16 + args.mixed_precision = accelerator.mixed_precision + elif accelerator.mixed_precision == "bf16": + weight_dtype = torch.bfloat16 + args.mixed_precision = accelerator.mixed_precision + + # Load scheduler, tokenizer and models. + # Use DDIM instead of DDPM to sample training videos. + noise_scheduler = DDIMScheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="scheduler") + noise_scheduler.set_timesteps(args.num_inference_steps, device=accelerator.device) + + tokenizer = T5Tokenizer.from_pretrained( + args.pretrained_model_name_or_path, subfolder="tokenizer", revision=args.revision + ) + + def deepspeed_zero_init_disabled_context_manager(): + """ + returns either a context list that includes one that will disable zero.Init or an empty context list + """ + deepspeed_plugin = AcceleratorState().deepspeed_plugin if accelerate.state.is_initialized() else None + if deepspeed_plugin is None: + return [] + + return [deepspeed_plugin.zero3_init_context_manager(enable=False)] + + # Currently Accelerate doesn't know how to handle multiple models under Deepspeed ZeRO stage 3. + # For this to work properly all models must be run through `accelerate.prepare`. But accelerate + # will try to assign the same optimizer with the same weights to all models during + # `deepspeed.initialize`, which of course doesn't work. + # + # For now the following workaround will partially support Deepspeed ZeRO-3, by excluding the 2 + # frozen models from being partitioned during `zero.Init` which gets called during + # `from_pretrained` So CLIPTextModel and AutoencoderKL will not enjoy the parameter sharding + # across multiple gpus and only UNet2DConditionModel will get ZeRO sharded. + with ContextManagers(deepspeed_zero_init_disabled_context_manager()): + text_encoder = T5EncoderModel.from_pretrained( + args.pretrained_model_name_or_path, subfolder="text_encoder", revision=args.revision, variant=args.variant, + torch_dtype=weight_dtype + ) + + vae = AutoencoderKLCogVideoX.from_pretrained( + args.pretrained_model_name_or_path, subfolder="vae", revision=args.revision, variant=args.variant + ) + + transformer3d = CogVideoXTransformer3DModel.from_pretrained_2d( + args.pretrained_model_name_or_path, subfolder="transformer" + ) + + # Freeze vae and text_encoder and set transformer3d to trainable + vae.requires_grad_(False) + text_encoder.requires_grad_(False) + transformer3d.requires_grad_(False) + + # Lora will work with this... + network = create_network( + 1.0, + args.rank, + args.network_alpha, + text_encoder, + transformer3d, + neuron_dropout=None, + add_lora_in_attn_temporal=True, + ) + network.apply_to(text_encoder, transformer3d, args.train_text_encoder, True) + + if args.transformer_path is not None: + print(f"From checkpoint: {args.transformer_path}") + if args.transformer_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(args.transformer_path) + else: + state_dict = torch.load(args.transformer_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = transformer3d.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + assert len(u) == 0 + + if args.vae_path is not None: + print(f"From checkpoint: {args.vae_path}") + if args.vae_path.endswith("safetensors"): + from safetensors.torch import load_file, safe_open + state_dict = load_file(args.vae_path) + else: + state_dict = torch.load(args.vae_path, map_location="cpu") + state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict + + m, u = vae.load_state_dict(state_dict, strict=False) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + assert len(u) == 0 + + vae_scale_factor_spatial = 2 ** (len(vae.config.block_out_channels) - 1) + vae_scale_factor_temporal = vae.config.temporal_compression_ratio + num_channels_latent = vae.config.latent_channels + + # `accelerate` 0.16.0 will have better support for customized saving + if version.parse(accelerate.__version__) >= version.parse("0.16.0"): + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + def save_model_hook(models, weights, output_dir): + if accelerator.is_main_process: + safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) + if not args.use_deepspeed: + for _ in range(len(weights)): + weights.pop() + + accelerator.register_save_state_pre_hook(save_model_hook) + # accelerator.register_load_state_pre_hook(load_model_hook) + + if args.gradient_checkpointing: + transformer3d.enable_gradient_checkpointing() + + if args.vae_gradient_checkpointing: + # Since 3D casual VAE need a cache to decode all latents autoregressively, .Thus, gradient checkpointing can only be + # enabled when decoding the first batch (i.e. the first three) of latents, in which case the cache is not being used. + if args.num_decoded_latents > 3: + raise ValueError("The vae_gradient_checkpointing is not supported for num_decoded_latents > 3.") + vae.enable_gradient_checkpointing() + + # Enable TF32 for faster training on Ampere GPUs, + # cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices + if args.allow_tf32: + torch.backends.cuda.matmul.allow_tf32 = True + + if args.scale_lr: + args.learning_rate = ( + args.learning_rate * args.gradient_accumulation_steps * args.train_batch_size * accelerator.num_processes + ) + + # Initialize the optimizer + if args.use_8bit_adam: + try: + import bitsandbytes as bnb + except ImportError: + raise ImportError( + "Please install bitsandbytes to use 8-bit Adam. You can do so by running `pip install bitsandbytes`" + ) + + optimizer_cls = bnb.optim.AdamW8bit + elif args.use_came: + try: + from came_pytorch import CAME + except: + raise ImportError( + "Please install came_pytorch to use CAME. You can do so by running `pip install came_pytorch`" + ) + + optimizer_cls = CAME + else: + optimizer_cls = torch.optim.AdamW + + logging.info("Add network parameters") + trainable_params = list(filter(lambda p: p.requires_grad, network.parameters())) + trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate) + + if args.use_came: + optimizer = optimizer_cls( + trainable_params_optim, + lr=args.learning_rate, + betas=(0.9, 0.999, 0.9999), + eps=(1e-30, 1e-16) + ) + else: + optimizer = optimizer_cls( + trainable_params_optim, + lr=args.learning_rate, + betas=(args.adam_beta1, args.adam_beta2), + weight_decay=args.adam_weight_decay, + eps=args.adam_epsilon, + ) + + # loss function + reward_fn_kwargs = {} + if args.reward_fn_kwargs is not None: + reward_fn_kwargs = json.loads(args.reward_fn_kwargs) + if accelerator.is_main_process: + # Check if the model is downloaded in the main process. + loss_fn = getattr(reward_fn, args.reward_fn)(device="cpu", dtype=weight_dtype, **reward_fn_kwargs) + accelerator.wait_for_everyone() + loss_fn = getattr(reward_fn, args.reward_fn)(device=accelerator.device, dtype=weight_dtype, **reward_fn_kwargs) + + # Get RL training prompts + prompt_list = load_prompts(args.prompt_path) + + # Scheduler and math around the number of training steps. + overrode_max_train_steps = False + num_update_steps_per_epoch = math.ceil(len(prompt_list) / args.gradient_accumulation_steps) + if args.max_train_steps is None: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + overrode_max_train_steps = True + + lr_scheduler = get_scheduler( + args.lr_scheduler, + optimizer=optimizer, + num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes, + num_training_steps=args.max_train_steps * accelerator.num_processes, + ) + + # Prepare everything with our `accelerator`. + network, optimizer, lr_scheduler = accelerator.prepare( + network, optimizer, lr_scheduler + ) + + # Move text_encode and vae to gpu and cast to weight_dtype + vae.to(accelerator.device, dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + text_encoder.to(accelerator.device) + + # We need to recalculate our total training steps as the size of the training dataloader may have changed. + num_update_steps_per_epoch = math.ceil(len(prompt_list) / args.train_batch_size / args.gradient_accumulation_steps) + if overrode_max_train_steps: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + # Afterwards we recalculate our number of training epochs + args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch) + + # We need to initialize the trackers we use, and also store our configuration. + # The trackers initializes automatically on the main process. + if accelerator.is_main_process: + tracker_config = dict(vars(args)) + tracker_config.pop("validation_prompts") + accelerator.init_trackers(args.tracker_project_name, tracker_config) + + # Function for unwrapping if model was compiled with `torch.compile`. + def unwrap_model(model): + model = accelerator.unwrap_model(model) + model = model._orig_mod if is_compiled_module(model) else model + return model + + # Train! + total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps + + logger.info("***** Running training *****") + logger.info(f" Num examples = {len(prompt_list)}") + logger.info(f" Num Epochs = {args.num_train_epochs}") + logger.info(f" Instantaneous batch size per device = {args.train_batch_size}") + logger.info(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}") + logger.info(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}") + logger.info(f" Total optimization steps = {args.max_train_steps}") + global_step = 0 + first_epoch = 0 + + # Potentially load in the weights and states from a previous save + if args.resume_from_checkpoint: + if args.resume_from_checkpoint != "latest": + path = os.path.basename(args.resume_from_checkpoint) + else: + # Get the most recent checkpoint + dirs = os.listdir(args.output_dir) + dirs = [d for d in dirs if d.startswith("checkpoint")] + dirs = sorted(dirs, key=lambda x: int(x.split("-")[1])) + path = dirs[-1] if len(dirs) > 0 else None + + if path is None: + accelerator.print( + f"Checkpoint '{args.resume_from_checkpoint}' does not exist. Starting a new training run." + ) + args.resume_from_checkpoint = None + initial_global_step = 0 + else: + 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) + print(f"missing keys: {len(m)}, unexpected keys: {len(u)}") + + accelerator.print(f"Resuming from checkpoint {path}") + accelerator.load_state(os.path.join(args.output_dir, path)) + global_step = int(path.split("-")[1]) + + initial_global_step = global_step + first_epoch = global_step // num_update_steps_per_epoch + else: + initial_global_step = 0 + + # function for saving/removing + def save_model(ckpt_file, unwrapped_nw): + os.makedirs(args.output_dir, exist_ok=True) + accelerator.print(f"\nsaving checkpoint: {ckpt_file}") + unwrapped_nw.save_weights(ckpt_file, weight_dtype, None) + + progress_bar = tqdm( + range(0, args.max_train_steps), + initial=initial_global_step, + desc="Steps", + # Only show the progress bar once on each machine. + disable=not accelerator.is_local_main_process, + ) + + for epoch in range(first_epoch, args.num_train_epochs): + train_dataloader_iterations = 100 + train_loss = 0.0 + train_reward = 0.0 + # In the following training loop, randomly select training prompts and use the + # `CogVideoX_Fun_Pipeline_Inpaint` to sample videos, calculate rewards, and update the network. + for _ in range(train_dataloader_iterations): + # train_prompt = random.sample(prompt_list, args.train_batch_size) + train_prompt = random.choices(prompt_list, k=args.train_batch_size) + logger.info(f"train_prompt: {train_prompt}") + + # here `guidance_scale` is defined analog to the guidance weight `w` of equation (2) + # of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1` + # corresponds to doing no classifier free guidance. + do_classifier_free_guidance = args.guidance_scale > 1.0 + + # Encode input prompt + prompt_embeds, negative_prompt_embeds = encode_prompt( + tokenizer, + text_encoder, + train_prompt, + do_classifier_free_guidance=do_classifier_free_guidance, + negative_prompt="", + dtype=weight_dtype, + device=accelerator.device, + ) + if do_classifier_free_guidance: + prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0) + + # Prepare timesteps + timesteps = noise_scheduler.timesteps + + # Prepare latents + latent_shape = [ + len(train_prompt), + (args.video_length - 1) // vae_scale_factor_temporal + 1, + num_channels_latent, + args.train_sample_height // vae_scale_factor_spatial, + args.train_sample_width // vae_scale_factor_spatial, + ] + + with accelerator.accumulate(transformer3d): + with accelerator.autocast(): + latents = torch.randn(*latent_shape, device=accelerator.device, dtype=weight_dtype) + latents = latents * noise_scheduler.init_noise_sigma + + mask_latents = torch.zeros_like(latents)[:, :, :1].to(latents.device, latents.dtype) + masked_video_latents = torch.zeros_like(latents).to(latents.device, latents.dtype) + mask_input = torch.cat([mask_latents] * 2) if do_classifier_free_guidance else mask_latents + masked_video_latents_input = ( + torch.cat([masked_video_latents] * 2) if do_classifier_free_guidance else masked_video_latents + ) + inpaint_latents = torch.cat([mask_input, masked_video_latents_input], dim=2).to(latents.dtype) + + generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + # Prepare extra step kwargs. + extra_step_kwargs = prepare_extra_step_kwargs(noise_scheduler, generator, args.eta) + + # Create rotary embeds if required + image_rotary_emb = ( + prepare_rotary_positional_embeddings( + args.train_sample_height, + args.train_sample_width, + latents.size(1), + vae_scale_factor_spatial, + unwrap_model(transformer3d).config.patch_size, + device=accelerator.device + ) + if unwrap_model(transformer3d).config.use_rotary_positional_embeddings + else None + ) + + # Denoising loop + for i, t in enumerate(tqdm(timesteps)): + # DRaFT-K: the reward gradient is back propagated only for the last K steps. + if args.backprop: + # backprop_cutoff_idx = random.randint(0, args.num_sampling_steps - 1) # random + # Simply setting K=1 results in the best reward vs. compute tradeoff. + backprop_cutoff_idx = args.num_inference_steps - 1 # last + if i >= backprop_cutoff_idx: + for param in network.parameters(): + param.requires_grad = True + else: + for param in network.parameters(): + param.requires_grad = False + + # for DPM-solver++ + old_pred_original_sample = None + + latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents + latent_model_input = noise_scheduler.scale_model_input(latent_model_input, t) + + # broadcast to batch dimension in a way that's compatible with ONNX/Core ML + timestep = t.expand(latent_model_input.shape[0]) + + # predict noise model_output + noise_pred = transformer3d( + hidden_states=latent_model_input, + encoder_hidden_states=prompt_embeds, + timestep=timestep, + image_rotary_emb=image_rotary_emb, + return_dict=False, + inpaint_latents=inpaint_latents + )[0] + noise_pred = noise_pred.float() + + # perform guidance + guidance_scale = args.guidance_scale + # if args.use_dynamic_cfg: + # guidance_scale = 1 + guidance_scale * ( + # (1 - math.cos(math.pi * ((args.num_inference_steps - t.item()) / args.num_inference_steps) ** 5.0)) / 2 + # ) + if do_classifier_free_guidance: + noise_pred_uncond, noise_pred_text = noise_pred.chunk(2) + noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond) + + # compute the previous noisy sample x_t -> x_t-1 + if not isinstance(noise_scheduler, CogVideoXDPMScheduler): + latents = noise_scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0] + else: + latents, old_pred_original_sample = noise_scheduler.step( + noise_pred, + old_pred_original_sample, + t, + timesteps[i - 1] if i > 0 else None, + latents, + **extra_step_kwargs, + return_dict=False, + ) + latents = latents.to(prompt_embeds.dtype) + + # decode latents (tensor) + latents = latents.permute(0, 2, 1, 3, 4) # [B, C, T, H, W] + # Since the casual VAE decoding consumes a large amount of VRAM, and we need to keep the decoding + # operation within the computational graph. Thus, we only decode the first args.num_decoded_latents + # to calculate the reward. + sampled_frame_indices = list(range(args.num_decoded_latents)) + sampled_latents = latents[:, :, sampled_frame_indices, :, :] + sampled_latents = 1 / vae.config.scaling_factor * sampled_latents + sampled_frames = vae.decode(sampled_latents).sample + sampled_frames = (sampled_frames / 2 + 0.5).clamp(0, 1) # [-1, 1] -> [0, 1] + + if global_step % args.checkpointing_steps == 0: + saved_file = f"sample-{global_step}-{accelerator.process_index}.mp4" + save_videos_grid( + sampled_frames.to(torch.float32).detach().cpu(), + os.path.join(args.output_dir, "train_sample", saved_file), + fps=8 + ) + + if args.num_sampled_frames is not None: + num_frames = sampled_frames.size(2) - 1 + sampled_frames_indices = torch.linspace(0, num_frames, steps=args.num_sampled_frames).long() + sampled_frames = sampled_frames[:, :, sampled_frames_indices, :, :] + # compute loss and reward + loss, reward = loss_fn(sampled_frames, train_prompt) + + # Gather the losses and rewards across all processes for logging (if we use distributed training). + avg_loss = accelerator.gather(loss.repeat(args.train_batch_size)).mean() + avg_reward = accelerator.gather(reward.repeat(args.train_batch_size)).mean() + train_loss += avg_loss.item() / args.gradient_accumulation_steps + train_reward += avg_reward.item() / args.gradient_accumulation_steps + + # Backpropagate + 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 not args.use_deepspeed: + accelerator.log({"total_norm": total_norm}, step=global_step) + optimizer.step() + lr_scheduler.step() + optimizer.zero_grad() + + # Checks if the accelerator has performed an optimization step behind the scenes + if accelerator.sync_gradients: + progress_bar.update(1) + global_step += 1 + accelerator.log({"train_loss": train_loss, "train_reward": train_reward}, step=global_step) + train_loss = 0.0 + train_reward = 0.0 + + if global_step % args.checkpointing_steps == 0: + # DeepSpeed requires saving weights on every device; saving weights only on the main process would cause issues. + if args.use_deepspeed or accelerator.is_main_process: + # _before_ saving state, check if this save would set us over the `checkpoints_total_limit` + if args.checkpoints_total_limit is not None: + checkpoints = os.listdir(args.output_dir) + checkpoints = [d for d in checkpoints if d.startswith("checkpoint")] + checkpoints = sorted(checkpoints, key=lambda x: int(x.split("-")[1])) + + # before we save the new checkpoint, we need to have at _most_ `checkpoints_total_limit - 1` checkpoints + if len(checkpoints) >= args.checkpoints_total_limit: + num_to_remove = len(checkpoints) - args.checkpoints_total_limit + 1 + removing_checkpoints = checkpoints[0:num_to_remove] + + logger.info( + f"{len(checkpoints)} checkpoints already exist, removing {len(removing_checkpoints)} checkpoints" + ) + logger.info(f"removing checkpoints: {', '.join(removing_checkpoints)}") + + for removing_checkpoint in removing_checkpoints: + removing_checkpoint = os.path.join(args.output_dir, removing_checkpoint) + shutil.rmtree(removing_checkpoint) + + if not args.save_state: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + logger.info(f"Saved safetensor to {safetensor_save_path}") + else: + accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(accelerator_save_path) + logger.info(f"Saved state to {accelerator_save_path}") + + # Validation (distributed) + if do_validation and (global_step % args.validation_steps) == 0: + if args.validation_prompts is None and args.validation_prompt_path.endswith(".txt"): + validation_prompts = [] + with open(args.validation_prompt_path, "r") as f: + for line in f: + validation_prompts.append(line.strip()) + # Do not select randomly to ensure that `args.validation_prompts` is the same for each process. + args.validation_prompts = validation_prompts[:args.validation_batch_size] + + validation_prompts_idx = [(i, p) for i, p in enumerate(args.validation_prompts)] + + accelerator.wait_for_everyone() + with accelerator.split_between_processes(validation_prompts_idx) as splitted_prompts_idx: + validation_loss, validation_reward = log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + network, + loss_fn, + args, + accelerator, + weight_dtype, + global_step, + splitted_prompts_idx + ) + avg_validation_loss = accelerator.gather(validation_loss).mean() + avg_validation_reward = accelerator.gather(validation_reward).mean() + if accelerator.is_main_process: + accelerator.log({"validation_loss": avg_validation_loss, "validation_reward": avg_validation_reward}, step=global_step) + accelerator.wait_for_everyone() + + logs = {"step_loss": loss.detach().item(), "step_reward": reward.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} + progress_bar.set_postfix(**logs) + + if global_step >= args.max_train_steps: + break + + +if __name__ == "__main__": + main() diff --git a/scripts/train_reward_lora.sh b/scripts/train_reward_lora.sh new file mode 100644 index 0000000..b97a060 --- /dev/null +++ b/scripts/train_reward_lora.sh @@ -0,0 +1,62 @@ +export MODEL_NAME="alibaba-pai/CogVideoX-Fun-V1.1-2b-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" + +# When train model with multi machines, use "--config_file accelerate.yaml" instead of "--mixed_precision='bf16'". +accelerate launch --num_processes=1 --mixed_precision="bf16" scripts/train_reward_lora.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --rank 32 \ + --network_alpha 16 \ + --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 \ + --vae_gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --max_grad_norm=0.3 \ + --prompt_path $TRAIN_PROMPT_PATH \ + --train_sample_height 224 \ + --train_sample_width 224 \ + --video_length 49 \ + --num_decoded_latents 1 \ + --num_sampled_frames 1 \ + --reward_fn "HPSReward" \ + --reward_fn_kwargs '{"version": "v2.1"}' \ + --backprop + +# Training command for CogVideoX-Fun-V1.1-2b-InP-HPS2.1.safetensors (with 8 A100 GPUs) +# 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 \ +# --rank 128 \ +# --network_alpha 64 \ +# --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 \ +# --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 \ +# --num_sampled_frames 1 \ +# --reward_fn "HPSReward" \ +# --reward_fn_kwargs '{"version": "v2.1"}' \ +# --backprop \ No newline at end of file