Add reward LoRA training (#160)
* Add reward lora * Fix bug in lora utils && Update Readme --------- Co-authored-by: bubbliiiing <3323290568@qq.com>
This commit is contained in:
@@ -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. |
|
||||
|
||||
<details>
|
||||
<summary>(Obsolete) EasyAnimateV4:</summary>
|
||||
|
||||
@@ -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生成ビデオの最適化によるヒト選好の最適化|
|
||||
|
||||
<details>
|
||||
<summary>(Obsolete) EasyAnimateV4:</summary>
|
||||
|
||||
@@ -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生成的视频,以更好地匹配人类偏好|
|
||||
|
||||
<details>
|
||||
<summary>(Obsolete) EasyAnimateV4:</summary>
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
This folder is modified from the official [MPS](https://github.com/Kwai-Kolors/MPS/tree/main) repository.
|
||||
@@ -0,0 +1,7 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
|
||||
@dataclass
|
||||
class BaseModelConfig:
|
||||
pass
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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)}")
|
||||
@@ -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
|
||||
@@ -0,0 +1,299 @@
|
||||
# Enhance EasyAnimate with Reward Backpropagation (Preference Optimization)
|
||||
We explore the Reward Backpropagation technique <sup>[1](#ref1) [2](#ref2)</sup> to optimized the generated videos by [EasyAnimateV5](https://github.com/aigc-apps/EasyAnimate/tree/main/easyanimate) for better alignment with human preferences.
|
||||
We provide pre-trained models (i.e. LoRAs) along with the training script. You can use these LoRAs to enhance the corresponding base model as a plug-in or train your own reward LoRA.
|
||||
|
||||
- [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
|
||||
|
||||
<table border="0" style="width: 100%; text-align: center; margin-top: 20px;">
|
||||
<thead>
|
||||
<tr>
|
||||
<th style="text-align: center;" width="10%">Prompt</sup></th>
|
||||
<th style="text-align: center;" width="30%">EasyAnimateV5-12b-zh-InP</th>
|
||||
<th style="text-align: center;" width="30%">EasyAnimateV5-12b-zh-InP <br> HPSv2.1 Reward LoRA</th>
|
||||
<th style="text-align: center;" width="30%">EasyAnimateV5-12b-zh-InP <br> MPS Reward LoRA</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tr>
|
||||
<td>
|
||||
Porcelain rabbit hopping by a golden cactus
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/c7ee83b2-0329-4853-b47d-e8e1550f1164" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/1fea5b95-05dd-44cf-aec2-5c104e3afa8d" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/de14593a-daae-4a3e-8231-7df2108065d5" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>
|
||||
Yellow rubber duck floating next to a blue bath towel
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/c146fe30-ddcc-4e26-8659-885efd48136f" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/bd4a0a5c-cfe0-4a04-835b-1a3613926a6d" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/f5076984-9661-4670-9ca5-abc33b7d66c0" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>
|
||||
An elephant sprays water with its trunk, a lion sitting nearby
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/139bc722-d8bb-42cb-b043-99334f320496" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/87edf580-f1f3-4be2-931e-e53306ca9087" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/a38581c2-f4b3-4905-93af-debb3aec6488" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>
|
||||
A fish swims gracefully in a tank as a horse gallops outside
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/0383cdd5-1d9c-4b62-bde9-7a0423c8f863" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/efaee3eb-c361-4167-8952-92853a13df24" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/4cd406e3-8348-4589-8c07-43379547e1e1" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
### EasyAnimateV5-7b-zh-InP
|
||||
|
||||
<table border="0" style="width: 100%; text-align: center; margin-top: 20px;">
|
||||
<thead>
|
||||
<tr>
|
||||
<th style="text-align: center;" width="10%">Prompt</th>
|
||||
<th style="text-align: center;" width="30%">EasyAnimateV5-7b-zh-InP</th>
|
||||
<th style="text-align: center;" width="30%">EasyAnimateV5-7b-zh-InP <br> HPSv2.1 Reward LoRA</th>
|
||||
<th style="text-align: center;" width="30%">EasyAnimateV5-7b-zh-InP <br> MPS Reward LoRA</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tr>
|
||||
<td>
|
||||
Crystal cake shimmering beside a metal apple
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/25ae8abe-2e53-4557-b3f0-a72c247603e2" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/26f47c9b-e8f6-4768-978f-56fb47de4f2f" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/56166d66-4645-409e-b236-48ea25e8400b" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>
|
||||
Elderly artist with a white beard painting on a white canvas
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/7e0d7153-036a-4a40-b726-218760837ce7" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/314a68e8-57e3-437e-9acc-656da5f73853" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/d045e3e8-c9bd-4833-9a00-6decd50047d9" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>
|
||||
Porcelain rabbit hopping by a golden cactus
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/93890751-2ae7-4d55-82dc-7f992c8ad9b4" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/932ef7e4-c8a9-4153-94a8-8975d872701e" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/be0a01aa-a0c7-45a1-9db2-3b718c0be272" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>
|
||||
Green parrot perching on a brown chair
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/74a41dd4-8375-44be-8242-11287037c484" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/fd76e645-4ae3-427f-ac7b-9712e6dae4dd" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/6a7a0c11-1a78-4d51-90c4-814d1f4fb338" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
> [!NOTE]
|
||||
> The above test prompts are from <a href="https://github.com/KaiyueSun98/T2V-CompBench">T2V-CompBench</a>. 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
|
||||
<ol>
|
||||
<li id="ref1">Clark, Kevin, et al. "Directly fine-tuning diffusion models on differentiable rewards.". In ICLR 2024.</li>
|
||||
<li id="ref2">Prabhudesai, Mihir, et al. "Aligning text-to-image diffusion models with reward backpropagation." arXiv preprint arXiv:2310.03739 (2023).</li>
|
||||
</ol>
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,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
|
||||
Reference in New Issue
Block a user