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:
hkz
2024-11-28 13:59:53 +08:00
committed by GitHub
co-authored by bubbliiiing
parent f419bf850b
commit 17c1f02ba8
15 changed files with 2912 additions and 27 deletions
+3
View File
@@ -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>
+3
View File
@@ -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>
+2
View File
@@ -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>
+1
View File
@@ -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)
+385
View File
@@ -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)}")
+34 -27
View File
@@ -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
+299
View File
@@ -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
+36
View File
@@ -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