From 17c1f02ba8fbaef00132dbe4ecf53c03aa4fe070 Mon Sep 17 00:00:00 2001 From: hkz Date: Thu, 28 Nov 2024 13:59:53 +0800 Subject: [PATCH] Add reward LoRA training (#160) * Add reward lora * Fix bug in lora utils && Update Readme --------- Co-authored-by: bubbliiiing <3323290568@qq.com> --- README.md | 3 + README_ja-JP.md | 3 + README_zh-CN.md | 2 + easyanimate/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 + easyanimate/reward/reward_fn.py | 385 +++++ easyanimate/utils/lora_utils.py | 61 +- scripts/README_TRAIN_REWARD.md | 299 ++++ scripts/train_reward_lora.py | 1501 +++++++++++++++++ scripts/train_reward_lora.sh | 36 + 15 files changed, 2912 insertions(+), 27 deletions(-) create mode 100644 easyanimate/reward/MPS/README.md create mode 100644 easyanimate/reward/MPS/trainer/models/base_model.py create mode 100644 easyanimate/reward/MPS/trainer/models/clip_model.py create mode 100644 easyanimate/reward/MPS/trainer/models/cross_modeling.py create mode 100644 easyanimate/reward/aesthetic_predictor_v2_5/__init__.py create mode 100644 easyanimate/reward/aesthetic_predictor_v2_5/siglip_v2_5.py create mode 100644 easyanimate/reward/improved_aesthetic_predictor.py create mode 100644 easyanimate/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 882319f..f70adaa 100644 --- a/README.md +++ b/README.md @@ -31,6 +31,7 @@ EasyAnimate is a pipeline based on the transformer architecture, designed for ge We will support quick pull-ups from different platforms, refer to [Quick Start](#quick-start). **New Features:** +- Use reward backpropagation to train Lora and optimize the video, aligning it better with human preferences, detailes in [here](scripts/README_TRAIN_REWARD.md). EasyAnimateV5-7b is released now. [2024.11.27] - **Updated to v5**, supporting video generation up to 1024x1024, 49 frames, 6s, 8fps, with expanded model scale to 12B, incorporating the MMDIT structure, and enabling control models with diverse inputs; supports bilingual predictions in Chinese and English. [2024.11.08] - **Updated to v4**, allowing for video generation up to 1024x1024, 144 frames, 6s, 24fps; supports video generation from text, image, and video, with a single model handling resolutions from 512 to 1280; bilingual predictions in Chinese and English enabled. [2024.08.15] - **Updated to v3**, supporting video generation up to 960x960, 144 frames, 6s, 24fps, from text and image. [2024.07.01] @@ -414,6 +415,7 @@ EasyAnimateV5: |--|--|--|--|--|--| | EasyAnimateV5-7b-zh-InP | EasyAnimateV5 | 22 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-7b-zh-InP) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-7b-zh-InP) | Official 7B image-to-video weights. Supports video prediction at multiple resolutions (512, 768, 1024), trained with 49 frames at 8 frames per second, and supports bilingual prediction in Chinese and English. | | EasyAnimateV5-7b-zh | EasyAnimateV5 | 22 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-7b-zh) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-7b-zh) | Official 7B text-to-video weights. Supports video prediction at multiple resolutions (512, 768, 1024), trained with 49 frames at 8 frames per second, and supports bilingual prediction in Chinese and English. | +| EasyAnimateV5-Reward-LoRAs | - | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-Reward-LoRAs) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-Reward-LoRAs) | The official reward backpropagation technology model optimizes the videos generated by EasyAnimateV5-12b to better match human preferences. | 12B: | Name | Type | Storage Space | Hugging Face | Model Scope | Description | @@ -421,6 +423,7 @@ EasyAnimateV5: | EasyAnimateV5-12b-zh-InP | EasyAnimateV5 | 34 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh-InP) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh-InP) | Official image-to-video weights. Supports video prediction at multiple resolutions (512, 768, 1024), trained with 49 frames at 8 frames per second, and supports bilingual prediction in Chinese and English. | | EasyAnimateV5-12b-zh-Control | EasyAnimateV5 | 34 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh-Control) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh-Control) | Official video control weights, supporting various control conditions such as Canny, Depth, Pose, MLSD, etc. Supports video prediction at multiple resolutions (512, 768, 1024) and is trained with 49 frames at 8 frames per second. Bilingual prediction in Chinese and English is supported. | | EasyAnimateV5-12b-zh | EasyAnimateV5 | 34 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh) | Official text-to-video weights. Supports video prediction at multiple resolutions (512, 768, 1024), trained with 49 frames at 8 frames per second, and supports bilingual prediction in Chinese and English. | +| EasyAnimateV5-Reward-LoRAs | - | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-Reward-LoRAs) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-Reward-LoRAs) | The official reward backpropagation technology model optimizes the videos generated by EasyAnimateV5-12b to better match human preferences. |
(Obsolete) EasyAnimateV4: diff --git a/README_ja-JP.md b/README_ja-JP.md index b6bc319..265c59e 100644 --- a/README_ja-JP.md +++ b/README_ja-JP.md @@ -31,6 +31,7 @@ EasyAnimateは、トランスフォーマーアーキテクチャに基づいた 異なるプラットフォームからのクイックプルアップをサポートします。詳細は[クイックスタート](#クイックスタート)を参照してください。 **新機能:** +- インセンティブ逆伝播を使用してLoraを訓練し、人間の好みに合うようにビデオを最適化します。詳細は、[ここ](scripts/README _ train _ REVARD.md)を参照してください。EasyAnimateV 5-7 bがリリースされました。[2024.11.27] - **v5に更新**、1024x1024までの動画生成をサポート、49フレーム、6秒、8fps、モデルスケールを12Bに拡張、MMDIT構造を組み込み、さまざまな入力を持つ制御モデルをサポート。中国語と英語のバイリンガル予測をサポート。[2024.11.08] - **v4に更新**、1024x1024までの動画生成をサポート、144フレーム、6秒、24fps、テキスト、画像、動画からの動画生成をサポート、512から1280までの解像度を単一モデルで処理。中国語と英語のバイリンガル予測をサポート。[2024.08.15] - **v3に更新**、960x960までの動画生成をサポート、144フレーム、6秒、24fps、テキストと画像からの動画生成をサポート。[2024.07.01] @@ -408,6 +409,7 @@ EasyAnimateV5: |--|--|--|--|--|--| | EasyAnimateV5-7b-zh-InP | EasyAnimateV5 | 22 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-7b-zh-InP) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-7b-zh-InP) | 公式の画像から動画への重み。複数の解像度(512、768、1024)での動画予測をサポートし、49フレーム、毎秒8フレームでトレーニングされ、中国語と英語のバイリンガル予測をサポートします。 | | EasyAnimateV5-7b-zh | EasyAnimateV5 | 22 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-7b-zh) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-7b-zh) | 公式のテキストから動画への重み。複数の解像度(512、768、1024)での動画予測をサポートし、49フレーム、毎秒8フレームでトレーニングされ、中国語と英語のバイリンガル予測をサポートします。 | +| EasyAnimateV5-Reward-LoRAs | - | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-Reward-LoRAs) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-Reward-LoRAs) | 公式インバース伝播技術モデルによるEasyAnimateV 5-12 b生成ビデオの最適化によるヒト選好の最適化| 12B: | 名前 | 種類 | ストレージスペース | Hugging Face | Model Scope | 説明 | @@ -415,6 +417,7 @@ EasyAnimateV5: | EasyAnimateV5-12b-zh-InP | EasyAnimateV5 | 34 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh-InP) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh-InP) | 公式の画像から動画への重み。複数の解像度(512、768、1024)での動画予測をサポートし、49フレーム、毎秒8フレームでトレーニングされ、中国語と英語のバイリンガル予測をサポートします。 | | EasyAnimateV5-12b-zh-Control | EasyAnimateV5 | 34 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh-Control) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh-Control) | 公式の動画制御重み。Canny、Depth、Pose、MLSDなどのさまざまな制御条件をサポートします。複数の解像度(512、768、1024)での動画予測をサポートし、49フレーム、毎秒8フレームでトレーニングされ、中国語と英語のバイリンガル予測をサポートします。 | | EasyAnimateV5-12b-zh | EasyAnimateV5 | 34 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh) | 公式のテキストから動画への重み。複数の解像度(512、768、1024)での動画予測をサポートし、49フレーム、毎秒8フレームでトレーニングされ、中国語と英語のバイリンガル予測をサポートします。 | +| EasyAnimateV5-Reward-LoRAs | - | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-Reward-LoRAs) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-Reward-LoRAs) | 公式インバース伝播技術モデルによるEasyAnimateV 5-12 b生成ビデオの最適化によるヒト選好の最適化|
(Obsolete) EasyAnimateV4: diff --git a/README_zh-CN.md b/README_zh-CN.md index 2b514f3..d454bbb 100644 --- a/README_zh-CN.md +++ b/README_zh-CN.md @@ -31,6 +31,7 @@ EasyAnimate是一个基于transformer结构的pipeline,可用于生成AI图片 我们会逐渐支持从不同平台快速启动,请参阅 [快速启动](#快速启动)。 新特性: +- 使用奖励反向传播来训练Lora并优化视频,使其更好地符合人类偏好,详细信息请参见[此处](scripts/README_train_REVARD.md)。EasyAnimateV5-7b现已发布。[ 2024.11.27 ] - 更新到v5版本,最大支持1024x1024,49帧, 6s, 8fps视频生成,拓展模型规模到12B,应用MMDIT结构,支持不同输入的控制模型,支持中文与英文双语预测。[ 2024.11.08 ] - 更新到v4版本,最大支持1024x1024,144帧, 6s, 24fps视频生成,支持文、图、视频生视频,单个模型可支持512到1280任意分辨率,支持中文与英文双语预测。[ 2024.08.15 ] - 更新到v3版本,最大支持960x960,144帧,6s, 24fps视频生成,支持文与图生视频模型。[ 2024.07.01 ] @@ -415,6 +416,7 @@ EasyAnimateV5: | EasyAnimateV5-12b-zh-InP | EasyAnimateV5 | 34 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh-InP) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh-InP)| 官方的图生视频权重。支持多分辨率(512,768,1024)的视频预测,支持多分辨率(512,768,1024)的视频预测,以49帧、每秒8帧进行训练,支持中文与英文双语预测 | | EasyAnimateV5-12b-zh-Control | EasyAnimateV5 | 34 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh-Control) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh-Control)| 官方的视频控制权重,支持不同的控制条件,如Canny、Depth、Pose、MLSD等。支持多分辨率(512,768,1024)的视频预测,支持多分辨率(512,768,1024)的视频预测,以49帧、每秒8帧进行训练,支持中文与英文双语预测 | | EasyAnimateV5-12b-zh | EasyAnimateV5 | 34 GB | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-12b-zh) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-12b-zh)| 官方的文生视频权重。可用于进行下游任务的fientune。支持多分辨率(512,768,1024)的视频预测,支持多分辨率(512,768,1024)的视频预测,以49帧、每秒8帧进行训练,支持中文与英文双语预测 | +| EasyAnimateV5-Reward-LoRAs | - | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-Reward-LoRAs) | [😄Link](https://modelscope.cn/models/PAI/EasyAnimateV5-Reward-LoRAs) | 通过奖励反向传播技术,优化了EasyAnimateV5-12b生成的视频,以更好地匹配人类偏好|
(Obsolete) EasyAnimateV4: diff --git a/easyanimate/reward/MPS/README.md b/easyanimate/reward/MPS/README.md new file mode 100644 index 0000000..d66d2ee --- /dev/null +++ b/easyanimate/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/easyanimate/reward/MPS/trainer/models/base_model.py b/easyanimate/reward/MPS/trainer/models/base_model.py new file mode 100644 index 0000000..df7907f --- /dev/null +++ b/easyanimate/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/easyanimate/reward/MPS/trainer/models/clip_model.py b/easyanimate/reward/MPS/trainer/models/clip_model.py new file mode 100644 index 0000000..003bb5d --- /dev/null +++ b/easyanimate/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/easyanimate/reward/MPS/trainer/models/cross_modeling.py b/easyanimate/reward/MPS/trainer/models/cross_modeling.py new file mode 100644 index 0000000..6822329 --- /dev/null +++ b/easyanimate/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/easyanimate/reward/aesthetic_predictor_v2_5/__init__.py b/easyanimate/reward/aesthetic_predictor_v2_5/__init__.py new file mode 100644 index 0000000..2d3d8f1 --- /dev/null +++ b/easyanimate/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/easyanimate/reward/aesthetic_predictor_v2_5/siglip_v2_5.py b/easyanimate/reward/aesthetic_predictor_v2_5/siglip_v2_5.py new file mode 100644 index 0000000..867f429 --- /dev/null +++ b/easyanimate/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/easyanimate/reward/improved_aesthetic_predictor.py b/easyanimate/reward/improved_aesthetic_predictor.py new file mode 100644 index 0000000..43037b9 --- /dev/null +++ b/easyanimate/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/easyanimate/reward/reward_fn.py b/easyanimate/reward/reward_fn.py new file mode 100644 index 0000000..6526919 --- /dev/null +++ b/easyanimate/reward/reward_fn.py @@ -0,0 +1,385 @@ +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 + # TODO: [transforms.Resize(224), transforms.CenterCrop(224)] for any aspect ratio. + 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 + # TODO: [transforms.Resize(224), transforms.CenterCrop(224)] for any aspect ratio. + 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 + # TODO: [transforms.Resize(224), transforms.CenterCrop(224)] for any aspect ratio. + 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/easyanimate/utils/lora_utils.py b/easyanimate/utils/lora_utils.py index b50d11b..778fecc 100644 --- a/easyanimate/utils/lora_utils.py +++ b/easyanimate/utils/lora_utils.py @@ -369,7 +369,6 @@ def create_network( def merge_lora(pipeline, lora_path, multiplier, device='cpu', dtype=torch.float32, state_dict=None, transformer_only=False): LORA_PREFIX_TRANSFORMER = "lora_unet" LORA_PREFIX_TEXT_ENCODER = "lora_te" - SPECIAL_LAYER_NAME = ["text_proj_t5"] if state_dict is None: state_dict = load_file(lora_path, device=device) else: @@ -410,20 +409,22 @@ def merge_lora(pipeline, lora_path, multiplier, device='cpu', dtype=torch.float3 else: temp_name = layer_infos.pop(0) - weight_up = elems['lora_up.weight'].to(dtype) - weight_down = elems['lora_down.weight'].to(dtype) + origin_dtype = curr_layer.weight.data.dtype + weight_up = elems['lora_up.weight'].to(curr_layer.weight.data.device, dtype) + weight_down = elems['lora_down.weight'].to(curr_layer.weight.data.device, dtype) + curr_layer = curr_layer.to(dtype) if 'alpha' in elems.keys(): alpha = elems['alpha'].item() / weight_up.shape[1] else: alpha = 1.0 - curr_layer.weight.data = curr_layer.weight.data.to(device) if len(weight_up.shape) == 4: - curr_layer.weight.data += multiplier * alpha * torch.mm(weight_up.squeeze(3).squeeze(2), - weight_down.squeeze(3).squeeze(2)).unsqueeze( - 2).unsqueeze(3) + curr_layer.weight.data += multiplier * alpha * torch.mm( + weight_up.squeeze(3).squeeze(2), weight_down.squeeze(3).squeeze(2) + ).unsqueeze(2).unsqueeze(3) else: curr_layer.weight.data += multiplier * alpha * torch.mm(weight_up, weight_down) + curr_layer = curr_layer.to(origin_dtype) return pipeline @@ -448,25 +449,29 @@ def unmerge_lora(pipeline, lora_path, multiplier=1, device="cpu", dtype=torch.fl layer_infos = layer.split(LORA_PREFIX_UNET + "_")[-1].split("_") curr_layer = pipeline.transformer - temp_name = layer_infos.pop(0) - print(layer, curr_layer) - while len(layer_infos) > -1: - try: - curr_layer = curr_layer.__getattr__(temp_name) - if len(layer_infos) > 0: - temp_name = layer_infos.pop(0) - elif len(layer_infos) == 0: - break - except Exception: - if len(layer_infos) == 0: - print('Error loading layer') - if len(temp_name) > 0: - temp_name += "_" + layer_infos.pop(0) - else: - temp_name = layer_infos.pop(0) + try: + curr_layer = curr_layer.__getattr__("_".join(layer_infos[1:])) + except Exception: + temp_name = layer_infos.pop(0) + while len(layer_infos) > -1: + try: + curr_layer = curr_layer.__getattr__(temp_name) + if len(layer_infos) > 0: + temp_name = layer_infos.pop(0) + elif len(layer_infos) == 0: + break + except Exception: + if len(layer_infos) == 0: + print('Error loading layer') + if len(temp_name) > 0: + temp_name += "_" + layer_infos.pop(0) + else: + temp_name = layer_infos.pop(0) - weight_up = elems['lora_up.weight'].to(dtype) - weight_down = elems['lora_down.weight'].to(dtype) + origin_dtype = curr_layer.weight.data.dtype + weight_up = elems['lora_up.weight'].to(curr_layer.weight.data.device, dtype) + weight_down = elems['lora_down.weight'].to(curr_layer.weight.data.device, dtype) + curr_layer = curr_layer.to(dtype) if 'alpha' in elems.keys(): alpha = elems['alpha'].item() / weight_up.shape[1] else: @@ -474,9 +479,11 @@ def unmerge_lora(pipeline, lora_path, multiplier=1, device="cpu", dtype=torch.fl curr_layer.weight.data = curr_layer.weight.data.to(device) if len(weight_up.shape) == 4: - curr_layer.weight.data -= multiplier * alpha * torch.mm(weight_up.squeeze(3).squeeze(2), - weight_down.squeeze(3).squeeze(2)).unsqueeze(2).unsqueeze(3) + curr_layer.weight.data -= multiplier * alpha * torch.mm( + weight_up.squeeze(3).squeeze(2), weight_down.squeeze(3).squeeze(2) + ).unsqueeze(2).unsqueeze(3) else: curr_layer.weight.data -= multiplier * alpha * torch.mm(weight_up, weight_down) + curr_layer = curr_layer.to(origin_dtype) return pipeline \ 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..cd05f37 --- /dev/null +++ b/scripts/README_TRAIN_REWARD.md @@ -0,0 +1,299 @@ +# Enhance EasyAnimate with Reward Backpropagation (Preference Optimization) +We explore the Reward Backpropagation technique [1](#ref1) [2](#ref2) to optimized the generated videos by [EasyAnimateV5](https://github.com/aigc-apps/EasyAnimate/tree/main/easyanimate) for better alignment with human preferences. +We provide pre-trained models (i.e. LoRAs) along with the training script. You can use these LoRAs to enhance the corresponding base model as a plug-in or train your own reward LoRA. + +- [Enhance EasyAnimate with Reward Backpropagation (Preference Optimization)](#enhance-easyanimate-with-reward-backpropagation-preference-optimization) + - [Demo](#demo) + - [EasyAnimateV5-12b-zh-InP](#easyanimatev5-12b-zh-inp) + - [EasyAnimateV5-7b-zh-InP](#easyanimatev5-7b-zh-inp) + - [Model Zoo](#model-zoo) + - [Inference](#inference) + - [Training](#training) + - [Setup](#setup) + - [Important Args](#important-args) + - [Limitations](#limitations) + - [References](#references) + + +## Demo +### EasyAnimateV5-12b-zh-InP + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
PromptEasyAnimateV5-12b-zh-InPEasyAnimateV5-12b-zh-InP
HPSv2.1 Reward LoRA
EasyAnimateV5-12b-zh-InP
MPS Reward LoRA
+ Porcelain rabbit hopping by a golden cactus + + + + + + +
+ Yellow rubber duck floating next to a blue bath towel + + + + + + +
+ An elephant sprays water with its trunk, a lion sitting nearby + + + + + + +
+ A fish swims gracefully in a tank as a horse gallops outside + + + + + + +
+ +### EasyAnimateV5-7b-zh-InP + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
PromptEasyAnimateV5-7b-zh-InPEasyAnimateV5-7b-zh-InP
HPSv2.1 Reward LoRA
EasyAnimateV5-7b-zh-InP
MPS Reward LoRA
+ Crystal cake shimmering beside a metal apple + + + + + + +
+ Elderly artist with a white beard painting on a white canvas + + + + + + +
+ Porcelain rabbit hopping by a golden cactus + + + + + + +
+ Green parrot perching on a brown chair + + + + + + +
+ +> [!NOTE] +> The above test prompts are from T2V-CompBench. All videos are generated with lora weight 0.7. + +## Model Zoo +| Name | Base Model | Reward Model | Hugging Face | Description | +|--|--|--|--|--| +| EasyAnimateV5-12b-zh-InP-HPS2.1.safetensors | EasyAnimateV5-12b-zh-InP | [HPS v2.1](https://github.com/tgxs002/HPSv2) | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-Reward-LoRAs/resolve/main/EasyAnimateV5-12b-zh-InP-HPS2.1.safetensors) | Official HPS v2.1 reward LoRA (`rank=128` and `network_alpha=64`) for EasyAnimateV5-12b-zh-InP. It is trained with a batch size of 8 for 2,500 steps.| +| EasyAnimateV5-7b-zh-InP-HPS2.1.safetensors | EasyAnimateV5-7b-zh-InP | [HPS v2.1](https://github.com/tgxs002/HPSv2) | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-Reward-LoRAs/resolve/main/EasyAnimateV5-7b-zh-InP-HPS2.1.safetensors) | Official HPS v2.1 reward LoRA (`rank=128` and `network_alpha=64`) for EasyAnimateV5-7b-zh-InP. It is trained with a batch size of 8 for 3,500 steps.| +| EasyAnimateV5-12b-zh-InP-MPS.safetensors | EasyAnimateV5-12b-zh-InP | [MPS](https://github.com/Kwai-Kolors/MPS) | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-Reward-LoRAs/resolve/main/EasyAnimateV5-12b-zh-InP-MPS.safetensors) | Official MPS reward LoRA (`rank=128` and `network_alpha=64`) for EasyAnimateV5-12b-zh-InP. It is trained with a batch size of 8 for 2,500 steps.| +| EasyAnimateV5-7b-zh-InP-MPS.safetensors | EasyAnimateV5-7b-zh-InP | [MPS](https://github.com/Kwai-Kolors/MPS) | [🤗Link](https://huggingface.co/alibaba-pai/EasyAnimateV5-Reward-LoRAs/resolve/main/EasyAnimateV5-7b-zh-InP-MPS.safetensors) | Official MPS reward LoRA (`rank=128` and `network_alpha=64`) for EasyAnimateV5-7b-zh-InP. It is trained with a batch size of 8 for 2,000 steps.| + +## Inference +We provide an example inference code to run EasyAnimateV5-12b-zh-InP with its HPS2.1 reward LoRA. + +```python +import torch +from diffusers import DDIMScheduler +from omegaconf import OmegaConf +from transformers import BertModel, BertTokenizer, T5EncoderModel, T5Tokenizer + +from easyanimate.models import AutoencoderKLMagvit, EasyAnimateTransformer3DModel +from easyanimate.pipeline.pipeline_easyanimate_multi_text_encoder_inpaint import EasyAnimatePipeline_Multi_Text_Encoder_Inpaint +from easyanimate.utils.lora_utils import merge_lora +from easyanimate.utils.utils import get_image_to_video_latent, save_videos_grid +from easyanimate.utils.fp8_optimization import convert_weight_dtype_wrapper + +# GPU memory mode, which can be choosen in [model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload]. +GPU_memory_mode = "model_cpu_offload" +# Download from https://raw.githubusercontent.com/aigc-apps/EasyAnimate/refs/heads/main/config/easyanimate_video_v5_magvit_multi_text_encoder.yaml +config_path = "config/easyanimate_video_v5_magvit_multi_text_encoder.yaml" +model_path = "alibaba-pai/EasyAnimateV5-12b-zh-InP" +lora_path = "alibaba-pai/EasyAnimateV5-Reward-LoRAs/EasyAnimateV5-12b-zh-InP-HPS2.1.safetensors" +weight_dtype = torch.bfloat16 +lora_weight = 0.7 + +prompt = "A panda eats bamboo while a monkey swings from branch to branch" +sample_size = [512, 512] +video_length = 49 + +config = OmegaConf.load(config_path) +transformer_additional_kwargs = OmegaConf.to_container(config['transformer_additional_kwargs']) +if weight_dtype == torch.float16: + transformer_additional_kwargs["upcast_attention"] = True +transformer = EasyAnimateTransformer3DModel.from_pretrained_2d( + model_path, + subfolder="transformer", + transformer_additional_kwargs=transformer_additional_kwargs, + torch_dtype=torch.float8_e4m3fn if GPU_memory_mode == "model_cpu_offload_and_qfloat8" else weight_dtype, + low_cpu_mem_usage=True, +) +vae = AutoencoderKLMagvit.from_pretrained( + model_path, subfolder="vae", vae_additional_kwargs=OmegaConf.to_container(config['vae_kwargs']) +).to(weight_dtype) +if config['vae_kwargs'].get('vae_type', 'AutoencoderKL') == 'AutoencoderKLMagvit' and weight_dtype == torch.float16: + vae.upcast_vae = True + +pipeline = EasyAnimatePipeline_Multi_Text_Encoder_Inpaint.from_pretrained( + model_path, + text_encoder=BertModel.from_pretrained(model_path, subfolder="text_encoder").to(weight_dtype), + text_encoder_2=T5EncoderModel.from_pretrained(model_path, subfolder="text_encoder_2").to(weight_dtype), + tokenizer=BertTokenizer.from_pretrained(model_path, subfolder="tokenizer"), + tokenizer_2=T5Tokenizer.from_pretrained(model_path, subfolder="tokenizer_2"), + vae=vae, + transformer=transformer, + scheduler=DDIMScheduler.from_pretrained(model_path, subfolder="scheduler"), + torch_dtype=weight_dtype +) +if GPU_memory_mode == "sequential_cpu_offload": + pipeline.enable_sequential_cpu_offload() +elif GPU_memory_mode == "model_cpu_offload_and_qfloat8": + pipeline.enable_model_cpu_offload() + convert_weight_dtype_wrapper(pipeline.transformer, weight_dtype) +else: + 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, + video_length = 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 EasyAnimateV5-12b-zh-InP. + +### 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/factory.py https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/package/patches/hpsv2_src_open_clip_factory_patches.py +wget -O $site_packages/hpsv2/src/open_clip/ https://github.com/tgxs002/HPSv2/raw/refs/heads/master/hpsv2/src/open_clip/bpe_simple_vocab_16e6.txt.gz +``` + +> [!NOTE] +> Since some models will be downloaded automatically from HuggingFace, Please run `HF_ENDPOINT=https://hf-mirror.com sh scripts/train_reward_lora.sh` if you cannot access to huggingface.com. + +### 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. +Due to the resize and crop preprocessing operations, we suggest using a 1:1 aspect ratio. ++ `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..5ab7d83 --- /dev/null +++ b/scripts/train_reward_lora.py @@ -0,0 +1,1501 @@ +"""Modified from EasyAnimate/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 pickle +import shutil +import sys +import json +from contextlib import contextmanager + +import random +from typing import Optional, List + +import accelerate +import diffusers +import numpy as np +import torch +import torch.utils.checkpoint +import torchvision.transforms as transforms +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 +from diffusers.optimization import get_scheduler +from diffusers.utils import check_min_version, deprecate, is_wandb_available +from diffusers.utils.import_utils import is_xformers_available +from diffusers.utils.torch_utils import is_compiled_module +from decord import VideoReader +from einops import rearrange +from omegaconf import OmegaConf +from packaging import version +from tqdm.auto import tqdm +from transformers import BertModel, BertTokenizer, T5EncoderModel, T5Tokenizer +from transformers.utils import ContextManagers + +import datasets + +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 + +from transformers import T5EncoderModel, T5Tokenizer +from transformers.utils import ContextManagers + +import easyanimate.reward.reward_fn as reward_fn +from easyanimate.models import (name_to_autoencoder_magvit, + name_to_transformer3d) +from easyanimate.pipeline.pipeline_easyanimate_multi_text_encoder import get_3d_rotary_pos_embed, get_resize_crop_region_for_grid +from easyanimate.pipeline.pipeline_easyanimate_multi_text_encoder_inpaint import EasyAnimatePipeline_Multi_Text_Encoder_Inpaint +from easyanimate.utils import gaussian_diffusion as gd +from easyanimate.utils.lora_utils import create_network, merge_lora +from easyanimate.utils.respace import SpacedDiffusion, space_timesteps +from easyanimate.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, text_encoder_2, tokenizer, tokenizer_2, transformer3d, network, + loss_fn, config, args, accelerator, weight_dtype, global_step, validation_prompts_idx +): + logger.info("Running validation... ") + + # Get New Transformer + Choosen_Transformer3DModel = name_to_transformer3d[ + config['transformer_additional_kwargs'].get('transformer_type', 'Transformer3DModel') + ] + + transformer3d_val = Choosen_Transformer3DModel.from_pretrained_2d( + args.pretrained_model_name_or_path, subfolder="transformer", + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']) + ).to(weight_dtype) + transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + + scheduler = DDIMScheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="scheduler") + pipeline = EasyAnimatePipeline_Multi_Text_Encoder_Inpaint.from_pretrained( + args.pretrained_model_name_or_path, + vae=accelerator.unwrap_model(vae).to(weight_dtype), + text_encoder=accelerator.unwrap_model(text_encoder), + text_encoder_2=accelerator.unwrap_model(text_encoder_2), + tokenizer=tokenizer, + tokenizer_2=tokenizer_2, + transformer=transformer3d_val, + scheduler=scheduler, + 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 = transforms.ToTensor() + validation_loss, validation_reward = 0, 0 + + if args.enable_xformers_memory_efficient_attention \ + and config['transformer_additional_kwargs'].get('transformer_type', 'Transformer3DModel') == 'Transformer3DModel': + pipeline.enable_xformers_memory_efficient_attention() + + for i in range(len(validation_prompts_idx)): + validation_idx, validation_prompt = validation_prompts_idx[i] + with torch.no_grad(): + with torch.autocast("cuda", dtype=weight_dtype): + if vae.cache_mag_vae: + video_length = int((args.video_length - 1) // vae.mini_batch_encoder * vae.mini_batch_encoder) + 1 if args.video_length != 1 else 1 + else: + video_length = int(args.video_length // vae.mini_batch_encoder * vae.mini_batch_encoder) if args.video_length != 1 else 1 + sample_size = [args.validation_sample_height, args.validation_sample_width] + input_video, input_video_mask, clip_image = get_image_to_video_latent( + None, None, video_length=args.video_length, sample_size=sample_size + ) + + if args.seed is None: + generator = None + else: + generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + + sample = pipeline( + validation_prompt, + video_length = video_length, + negative_prompt = "bad detailed", + height = args.validation_sample_height, + width = args.validation_sample_width, + guidance_scale = 7, + generator = generator, + + video = input_video, + mask_video = input_video_mask, + clip_image = clip_image, + ).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 EasyAnimatePipeline_Multi_Text_Encoder_Inpaint.encode_prompt +def encode_prompt( + tokenizer, + tokenizer_2, + text_encoder, + text_encoder_2, + prompt: str, + device: torch.device, + dtype: torch.dtype, + num_images_per_prompt: int = 1, + do_classifier_free_guidance: bool = True, + negative_prompt: Optional[str] = None, + prompt_embeds: Optional[torch.Tensor] = None, + negative_prompt_embeds: Optional[torch.Tensor] = None, + prompt_attention_mask: Optional[torch.Tensor] = None, + negative_prompt_attention_mask: Optional[torch.Tensor] = None, + max_sequence_length: Optional[int] = None, + text_encoder_index: int = 0, + actual_max_sequence_length: int = 256, + enable_text_attention_mask: bool = False, +): + tokenizers = [tokenizer, tokenizer_2] + text_encoders = [text_encoder, text_encoder_2] + + tokenizer = tokenizers[text_encoder_index] + text_encoder = text_encoders[text_encoder_index] + + if max_sequence_length is None: + if text_encoder_index == 0: + max_length = min(tokenizer.model_max_length, actual_max_sequence_length) + if text_encoder_index == 1: + max_length = min(tokenizer_2.model_max_length, actual_max_sequence_length) + else: + max_length = max_sequence_length + + if prompt is not None and isinstance(prompt, str): + batch_size = 1 + elif prompt is not None and isinstance(prompt, list): + batch_size = len(prompt) + else: + batch_size = prompt_embeds.shape[0] + + if prompt_embeds is None: + text_inputs = tokenizer( + prompt, + padding="max_length", + max_length=max_length, + truncation=True, + return_attention_mask=True, + return_tensors="pt", + ) + text_input_ids = text_inputs.input_ids + if text_input_ids.shape[-1] > actual_max_sequence_length: + reprompt = tokenizer.batch_decode(text_input_ids[:, :actual_max_sequence_length], skip_special_tokens=True) + text_inputs = tokenizer( + reprompt, + padding="max_length", + max_length=max_length, + truncation=True, + return_attention_mask=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 + ): + _actual_max_sequence_length = min(tokenizer.model_max_length, actual_max_sequence_length) + removed_text = tokenizer.batch_decode(untruncated_ids[:, _actual_max_sequence_length - 1 : -1]) + logger.warning( + "The following part of your input was truncated because CLIP can only handle sequences up to" + f" {_actual_max_sequence_length} tokens: {removed_text}" + ) + prompt_attention_mask = text_inputs.attention_mask.to(device) + if enable_text_attention_mask: + prompt_embeds = text_encoder( + text_input_ids.to(device), + attention_mask=prompt_attention_mask, + ) + else: + prompt_embeds = text_encoder( + text_input_ids.to(device) + ) + prompt_embeds = prompt_embeds[0] + prompt_attention_mask = prompt_attention_mask.repeat(num_images_per_prompt, 1) + + prompt_embeds = prompt_embeds.to(dtype=dtype, device=device) + + bs_embed, seq_len, _ = prompt_embeds.shape + # duplicate text embeddings for each generation per prompt, using mps friendly method + prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1) + prompt_embeds = prompt_embeds.view(bs_embed * num_images_per_prompt, seq_len, -1) + + # get unconditional embeddings for classifier free guidance + if do_classifier_free_guidance and negative_prompt_embeds is None: + uncond_tokens: List[str] + if negative_prompt is None: + uncond_tokens = [""] * batch_size + elif 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 isinstance(negative_prompt, str): + uncond_tokens = [negative_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`." + ) + else: + uncond_tokens = negative_prompt + + max_length = prompt_embeds.shape[1] + uncond_input = tokenizer( + uncond_tokens, + padding="max_length", + max_length=max_length, + truncation=True, + return_tensors="pt", + ) + uncond_input_ids = uncond_input.input_ids + if uncond_input_ids.shape[-1] > actual_max_sequence_length: + reuncond_tokens = tokenizer.batch_decode(uncond_input_ids[:, :actual_max_sequence_length], skip_special_tokens=True) + uncond_input = tokenizer( + reuncond_tokens, + padding="max_length", + max_length=max_length, + truncation=True, + return_attention_mask=True, + return_tensors="pt", + ) + uncond_input_ids = uncond_input.input_ids + + negative_prompt_attention_mask = uncond_input.attention_mask.to(device) + if enable_text_attention_mask: + negative_prompt_embeds = text_encoder( + uncond_input.input_ids.to(device), + attention_mask=negative_prompt_attention_mask, + ) + else: + negative_prompt_embeds = text_encoder( + uncond_input.input_ids.to(device) + ) + negative_prompt_embeds = negative_prompt_embeds[0] + negative_prompt_attention_mask = negative_prompt_attention_mask.repeat(num_images_per_prompt, 1) + + if do_classifier_free_guidance: + # duplicate unconditional embeddings for each generation per prompt, using mps friendly method + seq_len = negative_prompt_embeds.shape[1] + + negative_prompt_embeds = negative_prompt_embeds.to(dtype=dtype, device=device) + + negative_prompt_embeds = negative_prompt_embeds.repeat(1, num_images_per_prompt, 1) + negative_prompt_embeds = negative_prompt_embeds.view(batch_size * num_images_per_prompt, seq_len, -1) + + return prompt_embeds, negative_prompt_embeds, prompt_attention_mask, negative_prompt_attention_mask + + +# Modified from EasyAnimatePipeline_Multi_Text_Encoder_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 + + +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=100) + 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 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( + "--dataloader_num_workers", + type=int, + default=0, + help=( + "Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process." + ), + ) + 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("--push_to_hub", action="store_true", help="Whether or not to push the model to the Hub.") + parser.add_argument("--hub_token", type=str, default=None, help="The token to use to push to the Model Hub.") + parser.add_argument( + "--prediction_type", + type=str, + default=None, + help="The prediction_type that shall be used for training. Choose between 'epsilon' or 'v_prediction' or leave `None`. If left to `None` the default prediction type of the scheduler: `noise_scheduler.config.prediciton_type` is chosen.", + ) + parser.add_argument( + "--hub_model_id", + type=str, + default=None, + help="The name of the repository to keep in sync with the local `output_dir`.", + ) + 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( + "--report_model_info", action="store_true", help="Whether or not to report more info about model (such as norm, grad)." + ) + 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( + "--enable_xformers_memory_efficient_attention", action="store_true", help="Whether or not to use xformers." + ) + parser.add_argument("--noise_offset", type=float, default=0, help="The scale of noise offset.") + parser.add_argument( + "--validation_epochs", + type=int, + 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( + "--token_sample_size", + type=int, + default=512, + help="Sample size of the token.", + ) + parser.add_argument( + "--video_sample_n_frames", + type=int, + default=17, + help="Num frame of video.", + ) + parser.add_argument( + "--config_path", + type=str, + default=None, + help=( + "The config of the model in training." + ), + ) + 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="aesthetic_loss_fn", + 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) + + config = OmegaConf.load(args.config_path) + 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) + + if config['text_encoder_kwargs'].get('enable_multi_text_encoder', False): + print("Init BertTokenizer") + tokenizer = BertTokenizer.from_pretrained( + args.pretrained_model_name_or_path, subfolder="tokenizer", revision=args.revision + ) + print("Init T5Tokenizer") + tokenizer_2 = T5Tokenizer.from_pretrained( + args.pretrained_model_name_or_path, subfolder="tokenizer_2", revision=args.revision + ) + else: + print("Init T5Tokenizer") + tokenizer = T5Tokenizer.from_pretrained( + args.pretrained_model_name_or_path, subfolder="tokenizer", revision=args.revision + ) + tokenizer_2 = None + + 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()): + if config['text_encoder_kwargs'].get('enable_multi_text_encoder', False): + text_encoder = BertModel.from_pretrained( + args.pretrained_model_name_or_path, subfolder="text_encoder", revision=args.revision, variant=args.variant, + torch_dtype=weight_dtype + ) + text_encoder_2 = T5EncoderModel.from_pretrained( + args.pretrained_model_name_or_path, subfolder="text_encoder_2", revision=args.revision, variant=args.variant, + torch_dtype=weight_dtype + ) + else: + text_encoder = T5EncoderModel.from_pretrained( + args.pretrained_model_name_or_path, subfolder="text_encoder", revision=args.revision, variant=args.variant, + torch_dtype=weight_dtype + ) + text_encoder_2 = None + + # Get Vae + Choosen_AutoencoderKL = name_to_autoencoder_magvit[ + config['vae_kwargs'].get('vae_type', 'AutoencoderKL') + ] + vae = Choosen_AutoencoderKL.from_pretrained( + args.pretrained_model_name_or_path, subfolder="vae", revision=args.revision, variant=args.variant, + vae_additional_kwargs=OmegaConf.to_container(config['vae_kwargs']) + ) + + # Get Transformer + Choosen_Transformer3DModel = name_to_transformer3d[ + config['transformer_additional_kwargs'].get('transformer_type', 'Transformer3DModel') + ] + transformer3d = Choosen_Transformer3DModel.from_pretrained_2d( + args.pretrained_model_name_or_path, subfolder="transformer", + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']) + ) + + # Freeze vae and text_encoder and set transformer3d to trainable + vae.requires_grad_(False) + vae.eval() + text_encoder.requires_grad_(False) + if config['text_encoder_kwargs'].get('enable_multi_text_encoder', False): + text_encoder_2.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 and not args.training_with_video_token_length, True) + + # Load transformer and vae from path if it needs. + 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)}") + + 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)}") + + if args.enable_xformers_memory_efficient_attention \ + and config['transformer_additional_kwargs'].get('transformer_type', 'Transformer3DModel') == 'Transformer3DModel': + if is_xformers_available(): + import xformers + + xformers_version = version.parse(xformers.__version__) + if xformers_version == version.parse("0.0.16"): + logger.warn( + "xFormers 0.0.16 cannot be used for training in some GPUs. If you observe problems during training, please update xFormers to at least 0.0.17. See https://huggingface.co/docs/diffusers/main/en/optimization/xformers for more details." + ) + transformer3d.enable_xformers_memory_efficient_attention() + else: + raise ValueError("xformers is not available. Make sure it is installed correctly") + + # `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() + + # 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) + + # Init optimizer + 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) + if config['text_encoder_kwargs'].get('enable_multi_text_encoder', False): + text_encoder_2.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.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: + global_step = int(path.split("-")[1]) + + initial_global_step = global_step + + pkl_path = os.path.join(os.path.join(args.output_dir, path), "sampler_pos_start.pkl") + if os.path.exists(pkl_path): + with open(pkl_path, 'rb') as file: + _, first_epoch = pickle.load(file) + else: + first_epoch = global_step // num_update_steps_per_epoch + print(f"Load pkl from {pkl_path}. Get first_epoch = {first_epoch}.") + + from safetensors.torch import load_file, safe_open + state_dict = load_file(os.path.join(os.path.join(args.output_dir, path), "lora_diffusion_pytorch_model.safetensors")) + m, u = accelerator.unwrap_model(network).load_state_dict(state_dict, strict=False) + 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)) + 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 + # `EasyAnimatePipeline_Multi_Text_Encoder_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}") + + # default height and width + height = int(args.train_sample_height // 16 * 16) + width = int(args.train_sample_width // 16 * 16) + + # 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, + prompt_attention_mask, + negative_prompt_attention_mask, + ) = encode_prompt( + tokenizer, + tokenizer_2, + text_encoder, + text_encoder_2, + train_prompt, + device=accelerator.device, + dtype=weight_dtype, + do_classifier_free_guidance=do_classifier_free_guidance, + negative_prompt=[""] * len(train_prompt), + text_encoder_index=0, + enable_text_attention_mask=transformer3d.config.enable_text_attention_mask, + ) + ( + prompt_embeds_2, + negative_prompt_embeds_2, + prompt_attention_mask_2, + negative_prompt_attention_mask_2, + ) = encode_prompt( + tokenizer, + tokenizer_2, + text_encoder, + text_encoder_2, + train_prompt, + device=accelerator.device, + dtype=weight_dtype, + do_classifier_free_guidance=do_classifier_free_guidance, + negative_prompt=[""] * len(train_prompt), + text_encoder_index=1, + enable_text_attention_mask=transformer3d.config.enable_text_attention_mask, + ) + + # Prepare timesteps + timesteps = noise_scheduler.timesteps + + # Prepare latent variables + num_channels_latents = vae.config.latent_channels + num_channels_transformer = transformer3d.config.in_channels + vae_scale_factor = 2 ** (len(vae.config.block_out_channels) - 1) + latent_shape = [ + args.train_batch_size, + vae.config.latent_channels, + int((args.video_length - 1) // vae.mini_batch_encoder * vae.mini_batch_decoder + 1) if args.video_length != 1 else 1, + args.train_sample_height // vae_scale_factor, + args.train_sample_width // vae_scale_factor, + ] + + 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 + + # Prepare inpaint latents if it needs. + # Use zero latents if we want to t2v. + 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=1).to(latents.dtype) + + # Check that sizes of mask, masked image and latents match + if num_channels_transformer != num_channels_latents: + num_channels_mask = mask_latents.shape[1] + num_channels_masked_image = masked_video_latents.shape[1] + if num_channels_latents + num_channels_mask + num_channels_masked_image != transformer3d.config.in_channels: + raise ValueError( + f"Incorrect configuration settings! The config of `pipeline.transformer`: {transformer3d.config} expects" + f" {transformer3d.config.in_channels} but received `num_channels_latents`: {num_channels_latents} +" + f" `num_channels_mask`: {num_channels_mask} + `num_channels_masked_image`: {num_channels_masked_image}" + f" = {num_channels_latents+num_channels_masked_image+num_channels_mask}. Please verify the config of" + " `pipeline.transformer` or your `mask_image` or `image` input." + ) + + 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 image_rotary_emb, style embedding & time ids + grid_height = height // 8 // transformer3d.config.patch_size + grid_width = width // 8 // transformer3d.config.patch_size + base_size_width = 720 // 8 // transformer3d.config.patch_size + base_size_height = 480 // 8 // transformer3d.config.patch_size + grid_crops_coords = get_resize_crop_region_for_grid( + (grid_height, grid_width), base_size_width, base_size_height + ) + image_rotary_emb = get_3d_rotary_pos_embed( + transformer3d.config.attention_head_dim, grid_crops_coords, grid_size=(grid_height, grid_width), + temporal_size=latents.size(2), use_real=True, + ) + + # Get other hunyuan params + style = torch.tensor([0], device=accelerator.device) + + original_size = (1024, 1024) + crops_coords_top_left = (0, 0) + target_size = (height, width) + add_time_ids = list(original_size + target_size + crops_coords_top_left) + add_time_ids = torch.tensor([add_time_ids], dtype=prompt_embeds.dtype) + + if do_classifier_free_guidance: + prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds]) + prompt_attention_mask = torch.cat([negative_prompt_attention_mask, prompt_attention_mask]) + prompt_embeds_2 = torch.cat([negative_prompt_embeds_2, prompt_embeds_2]) + prompt_attention_mask_2 = torch.cat([negative_prompt_attention_mask_2, prompt_attention_mask_2]) + add_time_ids = torch.cat([add_time_ids] * 2, dim=0) + style = torch.cat([style] * 2, dim=0) + + prompt_embeds = prompt_embeds.to(device=accelerator.device) + prompt_attention_mask = prompt_attention_mask.to(device=accelerator.device) + prompt_embeds_2 = prompt_embeds_2.to(device=accelerator.device) + prompt_attention_mask_2 = prompt_attention_mask_2.to(device=accelerator.device) + add_time_ids = add_time_ids.to(dtype=prompt_embeds.dtype, device=accelerator.device).repeat( + args.train_batch_size, 1 + ) + style = style.to(device=accelerator.device).repeat(args.train_batch_size) + + # Denoising loop + for i, t in enumerate(tqdm(timesteps)): + # 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 + 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 + + # expand the latents if we are doing classifier free guidance + 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) + + # expand scalar t to 1-D tensor to match the 1st dim of latent_model_input + t_expand = torch.tensor([t] * latent_model_input.shape[0], device=accelerator.device).to( + dtype=latent_model_input.dtype + ) + + # predict the noise residual + noise_pred = transformer3d( + latent_model_input, + t_expand, + encoder_hidden_states=prompt_embeds, + text_embedding_mask=prompt_attention_mask, + encoder_hidden_states_t5=prompt_embeds_2, + text_embedding_mask_t5=prompt_attention_mask_2, + image_meta_size=add_time_ids, + style=style, + image_rotary_emb=image_rotary_emb, + inpaint_latents=inpaint_latents, + clip_encoder_hidden_states=None, + clip_attention_mask=None, + return_dict=False, + )[0] + + # perform guidance + if do_classifier_free_guidance: + noise_pred_uncond, noise_pred_text = noise_pred.chunk(2) + noise_pred = noise_pred_uncond + args.guidance_scale * (noise_pred_text - noise_pred_uncond) + + # compute the previous noisy sample x_t -> x_t-1 + latents = noise_scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0] + + # 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. + # TODO: Decode all latents but keep a portion of the decoding operation within the computational graph. + sampled_latent_indices = list(range(args.num_decoded_latents)) + sampled_latents = latents[:, :, sampled_latent_indices, :, :] + sampled_latents = 1 / vae.config.scaling_factor * sampled_latents + sampled_frames = vae.decode(sampled_latents)[0] + sampled_frames = sampled_frames.clamp(-1, 1) + 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, + text_encoder_2, + tokenizer, + tokenizer_2, + transformer3d, + network, + loss_fn, + config, + 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() + accelerator.print(avg_validation_loss, avg_validation_reward) + if accelerator.is_main_process: + accelerator.log({"validation_loss": avg_validation_loss, "validation_reward": avg_validation_reward}, step=global_step) + + accelerator.wait_for_everyone() + + 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..deb0e88 --- /dev/null +++ b/scripts/train_reward_lora.sh @@ -0,0 +1,36 @@ +export MODEL_NAME="models/Diffusion_Transformer/EasyAnimateV5-12b-zh-InP" +export TRAIN_PROMPT_PATH="MovieGenVideoBench_train.txt" +# Performing validation simultaneously with training will increase time and GPU memory usage. +export VALIDATION_PROMPT_PATH="MovieGenVideoBench_val.txt" + +export NCCL_IB_DISABLE=1 +export NCCL_P2P_DISABLE=1 +NCCL_DEBUG=INFO + +# When train model with multi machines, use "--config_file accelerate.yaml" instead of "--mixed_precision='bf16'". +accelerate launch --num_processes=8 --mixed_precision="bf16" --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json scripts/train_reward_lora.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --config_path="config/easyanimate_video_v5_magvit_multi_text_encoder.yaml" \ + --train_batch_size=1 \ + --gradient_accumulation_steps=1 \ + --max_train_steps=10000 \ + --checkpointing_steps=100 \ + --learning_rate=1e-05 \ + --seed=42 \ + --output_dir="output_dir" \ + --gradient_checkpointing \ + --mixed_precision="bf16" \ + --adam_weight_decay=3e-2 \ + --adam_epsilon=1e-10 \ + --max_grad_norm=0.3 \ + --prompt_path=$TRAIN_PROMPT_PATH \ + --train_sample_height=256 \ + --train_sample_width=256 \ + --video_length=49 \ + --validation_prompt_path=$VALIDATION_PROMPT_PATH \ + --validation_steps=100 \ + --validation_batch_size=8 \ + --num_decoded_latents=1 \ + --reward_fn="HPSReward" \ + --reward_fn_kwargs='{"version": "v2.1"}' \ + --backprop \ No newline at end of file