Add Reward LoRA Training (#71)
This commit is contained in:
@@ -23,7 +23,7 @@ CogVideoX-Fun is a modified pipeline based on the CogVideoX structure, designed
|
||||
We will support quick pull-ups from different platforms, refer to [Quick Start](#quick-start).
|
||||
|
||||
What's New:
|
||||
- Upload a new version of the control model that supports different control conditions such as Canny, Depth, Pose, MLSD, etc. [2024.11.16]
|
||||
- Use reinforcement learning with reward backpropagation to train Lora and optimize the video, aligning it better with human preferences, detailes in [here](scripts/README_TRAIN_REWARD.md). A new version of the control model supports various conditions (e.g., Canny, Depth, Pose, MLSD, etc.). [2024.11.21]
|
||||
- CogVideoX-Fun Control is now supported in diffusers. Thanks to [a-r-r-o-w](https://github.com/a-r-r-o-w) who contributed the support in this [PR](https://github.com/huggingface/diffusers/pull/9671). Check out the [docs](https://huggingface.co/docs/diffusers/main/en/api/pipelines/cogvideox) to know more. [ 2024.10.16 ]
|
||||
- Retrain the i2v model and add noise to increase the motion amplitude of the video. Upload the control model training code and control model. [ 2024.09.29 ]
|
||||
- Create code! Now supporting Windows and Linux. Supports 2b and 5b models. Supports video generation at any resolution from 256x256x49 to 1024x1024x49. [ 2024.09.18 ]
|
||||
@@ -175,6 +175,47 @@ Resolution-512
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
### CogVideoX-Fun-V1.1-5B with Reward Backpropagation
|
||||
|
||||
<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%">CogVideoX-Fun-V1.1-5B</th>
|
||||
<th style="text-align: center;" width="30%">CogVideoX-Fun-V1.1-5B <br> HPSv2.1 Reward LoRA</th>
|
||||
<th style="text-align: center;" width="30%">CogVideoX-Fun-V1.1-5B <br> MPS Reward LoRA</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tr>
|
||||
<td>
|
||||
Pig with wings flying above a diamond mountain
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/6682f507-4ca2-45e9-9d76-86e2d709efb3" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/ec9219a2-96b3-44dd-b918-8176b2beb3b0" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/a75c6a6a-0b69-4448-afc0-fda3c7955ba0" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>
|
||||
A dog runs through a field while a cat climbs a tree
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/0392d632-2ec3-46b4-8867-0da1db577b6d" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/7d8c729d-6afb-408e-b812-67c40c3aaa96" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/dcd1343c-7435-4558-b602-9c0fa08cbd59" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
### CogVideoX-Fun-V1.1-5B-Control
|
||||
|
||||
<table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
|
||||
|
||||
+42
-1
@@ -23,7 +23,7 @@ CogVideoX-Fun是一个基于CogVideoX结构修改后的的pipeline,是一个
|
||||
我们会逐渐支持从不同平台快速启动,请参阅 [快速启动](#快速启动)。
|
||||
|
||||
新特性:
|
||||
- 上传新版本的控制模型,支持不同的控制条件,如Canny、Depth、Pose、MLSD等。[2024.11.16]
|
||||
- 通过奖励反向传播技术训练Lora,以优化生成的视频,使其更好地与人类偏好保持一致,[更多信息](scripts/README_TRAIN_REWARD.md)。新版本的控制模型,支持不同的控制条件,如Canny、Depth、Pose、MLSD等。[2024.11.21]
|
||||
- CogVideoX-Fun Control现在在diffusers中得到了支持。感谢 [a-r-r-o-w](https://github.com/a-r-r-o-w)在这个 [PR](https://github.com/huggingface/diffusers/pull/9671)中贡献了支持。查看[文档](https://huggingface.co/docs/diffusers/main/en/api/pipelines/cogvideox)以了解更多信息。[2024.10.16]
|
||||
- 重新训练i2v模型,添加Noise,使得视频的运动幅度更大。上传控制模型训练代码与Control模型。[ 2024.09.29 ]
|
||||
- 创建代码!现在支持 Windows 和 Linux。支持2b与5b最大256x256x49到1024x1024x49的任意分辨率的视频生成。[ 2024.09.18 ]
|
||||
@@ -173,6 +173,47 @@ Resolution-512
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
### CogVideoX-Fun-V1.1-5B with Reward Backpropagation
|
||||
|
||||
<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%">CogVideoX-Fun-V1.1-5B</th>
|
||||
<th style="text-align: center;" width="30%">CogVideoX-Fun-V1.1-5B <br> HPSv2.1 Reward LoRA</th>
|
||||
<th style="text-align: center;" width="30%">CogVideoX-Fun-V1.1-5B <br> MPS Reward LoRA</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tr>
|
||||
<td>
|
||||
Pig with wings flying above a diamond mountain
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/6682f507-4ca2-45e9-9d76-86e2d709efb3" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/ec9219a2-96b3-44dd-b918-8176b2beb3b0" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/a75c6a6a-0b69-4448-afc0-fda3c7955ba0" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>
|
||||
A dog runs through a field while a cat climbs a tree
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/0392d632-2ec3-46b4-8867-0da1db577b6d" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/7d8c729d-6afb-408e-b812-67c40c3aaa96" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/dcd1343c-7435-4558-b602-9c0fa08cbd59" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
### CogVideoX-Fun-V1.1-5B-Control
|
||||
|
||||
<table border="0" style="width: 100%; text-align: left; margin-top: 20px;">
|
||||
|
||||
@@ -394,7 +394,7 @@ class CogVideoXDownBlock3D(nn.Module):
|
||||
zq: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
for resnet in self.resnets:
|
||||
if self.training and self.gradient_checkpointing:
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
|
||||
def create_custom_forward(module):
|
||||
def create_forward(*inputs):
|
||||
@@ -482,7 +482,7 @@ class CogVideoXMidBlock3D(nn.Module):
|
||||
zq: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
for resnet in self.resnets:
|
||||
if self.training and self.gradient_checkpointing:
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
|
||||
def create_custom_forward(module):
|
||||
def create_forward(*inputs):
|
||||
@@ -587,7 +587,7 @@ class CogVideoXUpBlock3D(nn.Module):
|
||||
) -> torch.Tensor:
|
||||
r"""Forward method of the `CogVideoXUpBlock3D` class."""
|
||||
for resnet in self.resnets:
|
||||
if self.training and self.gradient_checkpointing:
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
|
||||
def create_custom_forward(module):
|
||||
def create_forward(*inputs):
|
||||
@@ -709,7 +709,7 @@ class CogVideoXEncoder3D(nn.Module):
|
||||
r"""The forward method of the `CogVideoXEncoder3D` class."""
|
||||
hidden_states = self.conv_in(sample)
|
||||
|
||||
if self.training and self.gradient_checkpointing:
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
|
||||
def create_custom_forward(module):
|
||||
def custom_forward(*inputs):
|
||||
@@ -850,7 +850,7 @@ class CogVideoXDecoder3D(nn.Module):
|
||||
r"""The forward method of the `CogVideoXDecoder3D` class."""
|
||||
hidden_states = self.conv_in(sample)
|
||||
|
||||
if self.training and self.gradient_checkpointing:
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
|
||||
def create_custom_forward(module):
|
||||
def custom_forward(*inputs):
|
||||
|
||||
@@ -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,382 @@
|
||||
import os
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
import torch
|
||||
import torchvision.transforms as transforms
|
||||
from einops import rearrange
|
||||
from torchvision.datasets.utils import download_url
|
||||
from typing import Optional, Tuple
|
||||
|
||||
|
||||
# All reward models.
|
||||
__all__ = ["AestheticReward", "HPSReward", "PickScoreReward", "MPSReward"]
|
||||
|
||||
|
||||
class BaseReward(ABC):
|
||||
"""An base class for reward models. A custom Reward class must implement two functions below.
|
||||
"""
|
||||
def __init__(self):
|
||||
"""Define your reward model and image transformations (optional) here.
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def __call__(self, batch_frames: torch.Tensor, batch_prompt: Optional[list[str]]=None) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Given batch frames with shape `[B, C, T, H, W]` extracted from a list of videos and a list of prompts
|
||||
(optional) correspondingly, return the loss and reward computed by your reward model (reduction by mean).
|
||||
"""
|
||||
pass
|
||||
|
||||
class AestheticReward(BaseReward):
|
||||
"""Aesthetic Predictor [V2](https://github.com/christophschuhmann/improved-aesthetic-predictor)
|
||||
and [V2.5](https://github.com/discus0434/aesthetic-predictor-v2-5) reward model.
|
||||
"""
|
||||
def __init__(
|
||||
self,
|
||||
encoder_path="openai/clip-vit-large-patch14",
|
||||
predictor_path=None,
|
||||
version="v2",
|
||||
device="cpu",
|
||||
dtype=torch.float16,
|
||||
max_reward=10,
|
||||
loss_scale=0.1,
|
||||
):
|
||||
from .improved_aesthetic_predictor import ImprovedAestheticPredictor
|
||||
from ..video_caption.utils.siglip_v2_5 import convert_v2_5_from_siglip
|
||||
|
||||
self.encoder_path = encoder_path
|
||||
self.predictor_path = predictor_path
|
||||
self.version = version
|
||||
self.device = device
|
||||
self.dtype = dtype
|
||||
self.max_reward = max_reward
|
||||
self.loss_scale = loss_scale
|
||||
|
||||
if self.version != "v2" and self.version != "v2.5":
|
||||
raise ValueError("Only v2 and v2.5 are supported.")
|
||||
if self.version == "v2":
|
||||
assert "clip-vit-large-patch14" in encoder_path.lower()
|
||||
self.model = ImprovedAestheticPredictor(encoder_path=self.encoder_path, predictor_path=self.predictor_path)
|
||||
# https://huggingface.co/openai/clip-vit-large-patch14/blob/main/preprocessor_config.json
|
||||
self.transform = transforms.Compose([
|
||||
transforms.Resize((224, 224), interpolation=transforms.InterpolationMode.BICUBIC),
|
||||
transforms.Normalize(mean=[0.48145466, 0.4578275, 0.40821073], std=[0.26862954, 0.26130258, 0.27577711]),
|
||||
])
|
||||
elif self.version == "v2.5":
|
||||
assert "siglip-so400m-patch14-384" in encoder_path.lower()
|
||||
self.model, _ = convert_v2_5_from_siglip(encoder_model_name=self.encoder_path)
|
||||
# https://huggingface.co/google/siglip-so400m-patch14-384/blob/main/preprocessor_config.json
|
||||
self.transform = transforms.Compose([
|
||||
transforms.Resize((384, 384), interpolation=transforms.InterpolationMode.BICUBIC),
|
||||
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]),
|
||||
])
|
||||
|
||||
self.model.to(device=self.device, dtype=self.dtype)
|
||||
self.model.requires_grad_(False)
|
||||
|
||||
|
||||
def __call__(self, batch_frames: torch.Tensor, batch_prompt: Optional[list[str]]=None) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
batch_frames = rearrange(batch_frames, "b c t h w -> t b c h w")
|
||||
batch_loss, batch_reward = 0, 0
|
||||
for frames in batch_frames:
|
||||
pixel_values = torch.stack([self.transform(frame) for frame in frames])
|
||||
pixel_values = pixel_values.to(self.device, dtype=self.dtype)
|
||||
if self.version == "v2":
|
||||
reward = self.model(pixel_values)
|
||||
elif self.version == "v2.5":
|
||||
reward = self.model(pixel_values).logits.squeeze()
|
||||
# Convert reward to loss in [0, 1].
|
||||
if self.max_reward is None:
|
||||
loss = (-1 * reward) * self.loss_scale
|
||||
else:
|
||||
loss = abs(reward - self.max_reward) * self.loss_scale
|
||||
batch_loss, batch_reward = batch_loss + loss.mean(), batch_reward + reward.mean()
|
||||
|
||||
return batch_loss / batch_frames.shape[0], batch_reward / batch_frames.shape[0]
|
||||
|
||||
|
||||
class HPSReward(BaseReward):
|
||||
"""[HPS](https://github.com/tgxs002/HPSv2) v2 and v2.1 reward model.
|
||||
"""
|
||||
def __init__(
|
||||
self,
|
||||
model_path=None,
|
||||
version="v2.0",
|
||||
device="cpu",
|
||||
dtype=torch.float16,
|
||||
max_reward=1,
|
||||
loss_scale=1,
|
||||
):
|
||||
from hpsv2.src.open_clip import create_model_and_transforms, get_tokenizer
|
||||
|
||||
self.model_path = model_path
|
||||
self.version = version
|
||||
self.device = device
|
||||
self.dtype = dtype
|
||||
self.max_reward = max_reward
|
||||
self.loss_scale = loss_scale
|
||||
|
||||
self.model, _, _ = create_model_and_transforms(
|
||||
"ViT-H-14",
|
||||
"laion2B-s32B-b79K",
|
||||
precision=self.dtype,
|
||||
device=self.device,
|
||||
jit=False,
|
||||
force_quick_gelu=False,
|
||||
force_custom_text=False,
|
||||
force_patch_dropout=False,
|
||||
force_image_size=None,
|
||||
pretrained_image=False,
|
||||
image_mean=None,
|
||||
image_std=None,
|
||||
light_augmentation=True,
|
||||
aug_cfg={},
|
||||
output_dict=True,
|
||||
with_score_predictor=False,
|
||||
with_region_predictor=False,
|
||||
)
|
||||
self.tokenizer = get_tokenizer("ViT-H-14")
|
||||
|
||||
# https://huggingface.co/laion/CLIP-ViT-H-14-laion2B-s32B-b79K/blob/main/preprocessor_config.json
|
||||
self.transform = transforms.Compose([
|
||||
transforms.Resize((224, 224), interpolation=transforms.InterpolationMode.BICUBIC),
|
||||
transforms.Normalize(mean=[0.48145466, 0.4578275, 0.40821073], std=[0.26862954, 0.26130258, 0.27577711]),
|
||||
])
|
||||
|
||||
if version == "v2.0":
|
||||
url = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Third_Party/HPS_v2_compressed.pt"
|
||||
filename = "HPS_v2_compressed.pt"
|
||||
md5 = "fd9180de357abf01fdb4eaad64631db4"
|
||||
elif version == "v2.1":
|
||||
url = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Third_Party/HPS_v2.1_compressed.pt"
|
||||
filename = "HPS_v2.1_compressed.pt"
|
||||
md5 = "4067542e34ba2553a738c5ac6c1d75c0"
|
||||
else:
|
||||
raise ValueError("Only v2.0 and v2.1 are supported.")
|
||||
if self.model_path is None or not os.path.exists(self.model_path):
|
||||
download_url(url, torch.hub.get_dir(), md5=md5)
|
||||
model_path = os.path.join(torch.hub.get_dir(), filename)
|
||||
|
||||
state_dict = torch.load(model_path, map_location="cpu")["state_dict"]
|
||||
self.model.load_state_dict(state_dict)
|
||||
self.model.to(device=self.device, dtype=self.dtype)
|
||||
self.model.requires_grad_(False)
|
||||
self.model.eval()
|
||||
|
||||
def __call__(self, batch_frames: torch.Tensor, batch_prompt: list[str]) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
assert batch_frames.shape[0] == len(batch_prompt)
|
||||
# Compute batch reward and loss in frame-wise.
|
||||
batch_frames = rearrange(batch_frames, "b c t h w -> t b c h w")
|
||||
batch_loss, batch_reward = 0, 0
|
||||
for frames in batch_frames:
|
||||
image_inputs = torch.stack([self.transform(frame) for frame in frames])
|
||||
image_inputs = image_inputs.to(device=self.device, dtype=self.dtype)
|
||||
text_inputs = self.tokenizer(batch_prompt).to(device=self.device)
|
||||
outputs = self.model(image_inputs, text_inputs)
|
||||
|
||||
image_features, text_features = outputs["image_features"], outputs["text_features"]
|
||||
logits = image_features @ text_features.T
|
||||
reward = torch.diagonal(logits)
|
||||
# Convert reward to loss in [0, 1].
|
||||
if self.max_reward is None:
|
||||
loss = (-1 * reward) * self.loss_scale
|
||||
else:
|
||||
loss = abs(reward - self.max_reward) * self.loss_scale
|
||||
|
||||
batch_loss, batch_reward = batch_loss + loss.mean(), batch_reward + reward.mean()
|
||||
|
||||
return batch_loss / batch_frames.shape[0], batch_reward / batch_frames.shape[0]
|
||||
|
||||
|
||||
class PickScoreReward(BaseReward):
|
||||
"""[PickScore](https://github.com/yuvalkirstain/PickScore) reward model.
|
||||
"""
|
||||
def __init__(
|
||||
self,
|
||||
model_path="yuvalkirstain/PickScore_v1",
|
||||
device="cpu",
|
||||
dtype=torch.float16,
|
||||
max_reward=1,
|
||||
loss_scale=1,
|
||||
):
|
||||
from transformers import AutoProcessor, AutoModel
|
||||
|
||||
self.model_path = model_path
|
||||
self.device = device
|
||||
self.dtype = dtype
|
||||
self.max_reward = max_reward
|
||||
self.loss_scale = loss_scale
|
||||
|
||||
# https://huggingface.co/yuvalkirstain/PickScore_v1/blob/main/preprocessor_config.json
|
||||
self.transform = transforms.Compose([
|
||||
transforms.Resize(224, interpolation=transforms.InterpolationMode.BICUBIC),
|
||||
transforms.CenterCrop(224),
|
||||
transforms.Normalize(mean=[0.48145466, 0.4578275, 0.40821073], std=[0.26862954, 0.26130258, 0.27577711]),
|
||||
])
|
||||
self.processor = AutoProcessor.from_pretrained("laion/CLIP-ViT-H-14-laion2B-s32B-b79K", torch_dtype=self.dtype)
|
||||
self.model = AutoModel.from_pretrained(model_path, torch_dtype=self.dtype).eval().to(device)
|
||||
self.model.requires_grad_(False)
|
||||
self.model.eval()
|
||||
|
||||
def __call__(self, batch_frames: torch.Tensor, batch_prompt: list[str]) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
assert batch_frames.shape[0] == len(batch_prompt)
|
||||
# Compute batch reward and loss in frame-wise.
|
||||
batch_frames = rearrange(batch_frames, "b c t h w -> t b c h w")
|
||||
batch_loss, batch_reward = 0, 0
|
||||
for frames in batch_frames:
|
||||
image_inputs = torch.stack([self.transform(frame) for frame in frames])
|
||||
image_inputs = image_inputs.to(device=self.device, dtype=self.dtype)
|
||||
text_inputs = self.processor(
|
||||
text=batch_prompt,
|
||||
padding=True,
|
||||
truncation=True,
|
||||
max_length=77,
|
||||
return_tensors="pt",
|
||||
).to(self.device)
|
||||
image_features = self.model.get_image_features(pixel_values=image_inputs)
|
||||
text_features = self.model.get_text_features(**text_inputs)
|
||||
image_features = image_features / torch.norm(image_features, dim=-1, keepdim=True)
|
||||
text_features = text_features / torch.norm(text_features, dim=-1, keepdim=True)
|
||||
|
||||
logits = image_features @ text_features.T
|
||||
reward = torch.diagonal(logits)
|
||||
# Convert reward to loss in [0, 1].
|
||||
if self.max_reward is None:
|
||||
loss = (-1 * reward) * self.loss_scale
|
||||
else:
|
||||
loss = abs(reward - self.max_reward) * self.loss_scale
|
||||
|
||||
batch_loss, batch_reward = batch_loss + loss.mean(), batch_reward + reward.mean()
|
||||
|
||||
return batch_loss / batch_frames.shape[0], batch_reward / batch_frames.shape[0]
|
||||
|
||||
|
||||
class MPSReward(BaseReward):
|
||||
"""[MPS](https://github.com/Kwai-Kolors/MPS) reward model.
|
||||
"""
|
||||
def __init__(
|
||||
self,
|
||||
model_path=None,
|
||||
device="cpu",
|
||||
dtype=torch.float16,
|
||||
max_reward=1,
|
||||
loss_scale=1,
|
||||
):
|
||||
from transformers import AutoTokenizer, AutoConfig
|
||||
from .MPS.trainer.models.clip_model import CLIPModel
|
||||
|
||||
self.model_path = model_path
|
||||
self.device = device
|
||||
self.dtype = dtype
|
||||
self.condition = "light, color, clarity, tone, style, ambiance, artistry, shape, face, hair, hands, limbs, structure, instance, texture, quantity, attributes, position, number, location, word, things."
|
||||
self.max_reward = max_reward
|
||||
self.loss_scale = loss_scale
|
||||
|
||||
processor_name_or_path = "laion/CLIP-ViT-H-14-laion2B-s32B-b79K"
|
||||
# https://huggingface.co/laion/CLIP-ViT-H-14-laion2B-s32B-b79K/blob/main/preprocessor_config.json
|
||||
self.transform = transforms.Compose([
|
||||
transforms.Resize((224, 224), interpolation=transforms.InterpolationMode.BICUBIC),
|
||||
transforms.Normalize(mean=[0.48145466, 0.4578275, 0.40821073], std=[0.26862954, 0.26130258, 0.27577711]),
|
||||
])
|
||||
|
||||
# We convert the original [ckpt](http://drive.google.com/file/d/17qrK_aJkVNM75ZEvMEePpLj6L867MLkN/view?usp=sharing)
|
||||
# (contains the entire model) to a `state_dict`.
|
||||
url = "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/Third_Party/MPS_overall.pth"
|
||||
filename = "MPS_overall.pth"
|
||||
md5 = "1491cbbbd20565747fe07e7572e2ac56"
|
||||
if self.model_path is None or not os.path.exists(self.model_path):
|
||||
download_url(url, torch.hub.get_dir(), md5=md5)
|
||||
model_path = os.path.join(torch.hub.get_dir(), filename)
|
||||
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(processor_name_or_path, trust_remote_code=True)
|
||||
config = AutoConfig.from_pretrained(processor_name_or_path)
|
||||
self.model = CLIPModel(config)
|
||||
state_dict = torch.load(model_path, map_location="cpu")
|
||||
self.model.load_state_dict(state_dict, strict=False)
|
||||
self.model.to(device=self.device, dtype=self.dtype)
|
||||
self.model.requires_grad_(False)
|
||||
self.model.eval()
|
||||
|
||||
def _tokenize(self, caption):
|
||||
input_ids = self.tokenizer(
|
||||
caption,
|
||||
max_length=self.tokenizer.model_max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_tensors="pt"
|
||||
).input_ids
|
||||
|
||||
return input_ids
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
batch_frames: torch.Tensor,
|
||||
batch_prompt: list[str],
|
||||
batch_condition: Optional[list[str]] = None
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
if batch_condition is None:
|
||||
batch_condition = [self.condition] * len(batch_prompt)
|
||||
batch_frames = rearrange(batch_frames, "b c t h w -> t b c h w")
|
||||
batch_loss, batch_reward = 0, 0
|
||||
for frames in batch_frames:
|
||||
image_inputs = torch.stack([self.transform(frame) for frame in frames])
|
||||
image_inputs = image_inputs.to(device=self.device, dtype=self.dtype)
|
||||
text_inputs = self._tokenize(batch_prompt).to(self.device)
|
||||
condition_inputs = self._tokenize(batch_condition).to(device=self.device)
|
||||
text_features, image_features = self.model(text_inputs, image_inputs, condition_inputs)
|
||||
|
||||
text_features = text_features / text_features.norm(dim=-1, keepdim=True)
|
||||
image_features = image_features / image_features.norm(dim=-1, keepdim=True)
|
||||
# reward = self.model.logit_scale.exp() * torch.diag(torch.einsum('bd,cd->bc', text_features, image_features))
|
||||
logits = image_features @ text_features.T
|
||||
reward = torch.diagonal(logits)
|
||||
# Convert reward to loss in [0, 1].
|
||||
if self.max_reward is None:
|
||||
loss = (-1 * reward) * self.loss_scale
|
||||
else:
|
||||
loss = abs(reward - self.max_reward) * self.loss_scale
|
||||
|
||||
batch_loss, batch_reward = batch_loss + loss.mean(), batch_reward + reward.mean()
|
||||
|
||||
return batch_loss / batch_frames.shape[0], batch_reward / batch_frames.shape[0]
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import numpy as np
|
||||
from decord import VideoReader
|
||||
|
||||
video_path_list = ["your_video_path_1.mp4", "your_video_path_2.mp4"]
|
||||
prompt_list = ["your_prompt_1", "your_prompt_2"]
|
||||
num_sampled_frames = 8
|
||||
|
||||
to_tensor = transforms.ToTensor()
|
||||
|
||||
sampled_frames_list = []
|
||||
for video_path in video_path_list:
|
||||
vr = VideoReader(video_path)
|
||||
sampled_frame_indices = np.linspace(0, len(vr), num_sampled_frames, endpoint=False, dtype=int)
|
||||
sampled_frames = vr.get_batch(sampled_frame_indices).asnumpy()
|
||||
sampled_frames = torch.stack([to_tensor(frame) for frame in sampled_frames])
|
||||
sampled_frames_list.append(sampled_frames)
|
||||
sampled_frames = torch.stack(sampled_frames_list)
|
||||
sampled_frames = rearrange(sampled_frames, "b t c h w -> b c t h w")
|
||||
|
||||
aesthetic_reward_v2 = AestheticReward(device="cuda", dtype=torch.bfloat16)
|
||||
print(f"aesthetic_reward_v2: {aesthetic_reward_v2(sampled_frames)}")
|
||||
|
||||
aesthetic_reward_v2_5 = AestheticReward(
|
||||
encoder_path="google/siglip-so400m-patch14-384", version="v2.5", device="cuda", dtype=torch.bfloat16
|
||||
)
|
||||
print(f"aesthetic_reward_v2_5: {aesthetic_reward_v2_5(sampled_frames)}")
|
||||
|
||||
hps_reward_v2 = HPSReward(device="cuda", dtype=torch.bfloat16)
|
||||
print(f"hps_reward_v2: {hps_reward_v2(sampled_frames, prompt_list)}")
|
||||
|
||||
hps_reward_v2_1 = HPSReward(version="v2.1", device="cuda", dtype=torch.bfloat16)
|
||||
print(f"hps_reward_v2_1: {hps_reward_v2_1(sampled_frames, prompt_list)}")
|
||||
|
||||
pick_score = PickScoreReward(device="cuda", dtype=torch.bfloat16)
|
||||
print(f"pick_score_reward: {pick_score(sampled_frames, prompt_list)}")
|
||||
|
||||
mps_score = MPSReward(device="cuda", dtype=torch.bfloat16)
|
||||
print(f"mps_reward: {mps_score(sampled_frames, prompt_list)}")
|
||||
@@ -0,0 +1,258 @@
|
||||
# Enhance CogVideoX-Fun with Reward Backpropagation (Preference Optimization)
|
||||
We explore the Reward Backpropagation technique <sup>[1](#ref1) [2](#ref2)</sup> to optimized the generated videos by [CogVideoX-Fun-V1.1](https://github.com/aigc-apps/CogVideoX-Fun) for better alignment with human preferences.
|
||||
We provide pre-trained models (i.e. LoRAs) along with the training script. You can use these LoRAs to enhance the corresponding base model as a plug-in or train your own reward LoRA.
|
||||
|
||||
- [Enhance CogVideoX-Fun with Reward Backpropagation (Preference Optimization)](#enhance-cogvideox-fun-with-reward-backpropagation-preference-optimization)
|
||||
- [Demo](#demo)
|
||||
- [CogVideoX-Fun-V1.1-5B](#cogvideox-fun-v11-5b)
|
||||
- [CogVideoX-Fun-V1.1-2B](#cogvideox-fun-v11-2b)
|
||||
- [Model Zoo](#model-zoo)
|
||||
- [Inference](#inference)
|
||||
- [Training](#training)
|
||||
- [Setup](#setup)
|
||||
- [Important Args](#important-args)
|
||||
- [Limitations](#limitations)
|
||||
- [References](#references)
|
||||
|
||||
|
||||
## Demo
|
||||
### CogVideoX-Fun-V1.1-5B
|
||||
|
||||
<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%">CogVideoX-Fun-V1.1-5B</th>
|
||||
<th style="text-align: center;" width="30%">CogVideoX-Fun-V1.1-5B <br> HPSv2.1 Reward LoRA</th>
|
||||
<th style="text-align: center;" width="30%">CogVideoX-Fun-V1.1-5B <br> MPS Reward LoRA</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tr>
|
||||
<td>
|
||||
Pig with wings flying above a diamond mountain
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/6682f507-4ca2-45e9-9d76-86e2d709efb3" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/ec9219a2-96b3-44dd-b918-8176b2beb3b0" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/a75c6a6a-0b69-4448-afc0-fda3c7955ba0" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>
|
||||
A dog runs through a field while a cat climbs a tree
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/0392d632-2ec3-46b4-8867-0da1db577b6d" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/7d8c729d-6afb-408e-b812-67c40c3aaa96" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/dcd1343c-7435-4558-b602-9c0fa08cbd59" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>
|
||||
Crystal cake shimmering beside a metal apple
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/af0df8e0-1edb-4e2c-9a87-70df2b564aef" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/59b840f7-d33c-4972-8024-11a097f1c419" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/4a1d0af0-54e3-455c-9930-0789e2346fa0" 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/99e44f9d-c770-48ce-8cc5-69fe36d757bc" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/9c106677-e4cb-4970-a1a2-a013fa6ce903" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/0a7b57ab-36a8-4fb6-bcfa-75e3878c55b7" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
### CogVideoX-Fun-V1.1-2B
|
||||
|
||||
<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%">CogVideoX-Fun-V1.1-2B</th>
|
||||
<th style="text-align: center;" width="30%">CogVideoX-Fun-V1.1-2B <br> HPSv2.1 Reward LoRA</th>
|
||||
<th style="text-align: center;" width="30%">CogVideoX-Fun-V1.1-2B <br> MPS Reward LoRA</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tr>
|
||||
<td>
|
||||
A blue car drives past a white picket fence on a sunny day
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/274b0873-4fbd-4afa-94c0-22b23168f0a1" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/730f2ba3-4c54-44ce-ad5b-4eeca7ae844e" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/1b8eb777-0f17-46ef-9e7e-c8be7636e157" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>
|
||||
Blue jay swooping near a red maple tree
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/a14778d2-38ea-42c3-89a2-18164c48f3cf" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/90af433f-ab01-4341-9977-c675041d76d0" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/dafe8bf6-77ac-4934-8c9c-61c25088f80b" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>
|
||||
Yellow curtains swaying near a blue sofa
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/e8a445a4-781b-4b3f-899b-2cc24201f247" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/318cfb00-8bd1-407f-aaee-8d4220573b82" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/6b90e8a4-1754-42f4-b454-73510ed0701d" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>
|
||||
White tractor plowing near a green farmhouse
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/42d35282-e964-4c8b-aae9-a1592178493a" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/c9704bd4-d88d-41a1-8e5b-b7980df57a4a" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
<td>
|
||||
<video src="https://github.com/user-attachments/assets/7a785b34-4a5d-4491-9e03-c40cf953a1dc" width="100%" controls autoplay loop></video>
|
||||
</td>
|
||||
</tr>
|
||||
</table>
|
||||
|
||||
> [!NOTE]
|
||||
> The above test prompts are from <a href="https://github.com/Vchitect/VBench/tree/master/prompts">VBench</a>. All videos are generated with lora weight 0.7.
|
||||
|
||||
## Model Zoo
|
||||
| Name | Base Model | Reward Model | Hugging Face | Description |
|
||||
|--|--|--|--|--|
|
||||
| CogVideoX-Fun-V1.1-5b-InP-HPS2.1.safetensors | CogVideoX-Fun-V1.1-5b | [HPS v2.1](https://github.com/tgxs002/HPSv2) | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-Reward-LoRAs/resolve/main/CogVideoX-Fun-V1.1-5b-InP-HPS2.1.safetensors) | Official HPS v2.1 reward LoRA (`rank=128` and `network_alpha=64`) for CogVideoX-Fun-V1.1-5b-InP. It is trained with a batch size of 8 for 1,500 steps.|
|
||||
| CogVideoX-Fun-V1.1-2b-InP-HPS2.1.safetensors | CogVideoX-Fun-V1.1-2b | [HPS v2.1](https://github.com/tgxs002/HPSv2) | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-Reward-LoRAs/resolve/main/CogVideoX-Fun-V1.1-2b-InP-HPS2.1.safetensors) | Official HPS v2.1 reward LoRA (`rank=128` and `network_alpha=64`) for CogVideoX-Fun-V1.1-2b-InP. It is trained with a batch size of 8 for 3,000 steps.|
|
||||
| CogVideoX-Fun-V1.1-5b-InP-MPS.safetensors | CogVideoX-Fun-V1.1-5b | [MPS](https://github.com/Kwai-Kolors/MPS) | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-Reward-LoRAs/resolve/main/CogVideoX-Fun-V1.1-5b-InP-MPS.safetensors) | Official MPS reward LoRA (`rank=128` and `network_alpha=64`) for CogVideoX-Fun-V1.1-5b-InP. It is trained with a batch size of 8 for 5,500 steps.|
|
||||
| CogVideoX-Fun-V1.1-2b-InP-MPS.safetensors | CogVideoX-Fun-V1.1-2b | [MPS](https://github.com/Kwai-Kolors/MPS) | [🤗Link](https://huggingface.co/alibaba-pai/CogVideoX-Fun-V1.1-Reward-LoRAs/resolve/main/CogVideoX-Fun-V1.1-2b-InP-MPS.safetensors) | Official MPS reward LoRA (`rank=128` and `network_alpha=64`) for CogVideoX-Fun-V1.1-2b-InP. It is trained with a batch size of 8 for 16,000 steps.|
|
||||
|
||||
## Inference
|
||||
We provide an example inference code to run CogVideoX-Fun-V1.1-5b-InP with its HPS2.1 reward LoRA.
|
||||
|
||||
```python
|
||||
import torch
|
||||
from diffusers import CogVideoXDDIMScheduler
|
||||
|
||||
from cogvideox.models.transformer3d import CogVideoXTransformer3DModel
|
||||
from cogvideox.pipeline.pipeline_cogvideox_inpaint import CogVideoX_Fun_Pipeline_Inpaint
|
||||
from cogvideox.utils.lora_utils import merge_lora
|
||||
from cogvideox.utils.utils import get_image_to_video_latent, save_videos_grid
|
||||
|
||||
model_path = "alibaba-pai/CogVideoX-Fun-V1.1-5b-InP"
|
||||
lora_path = "alibaba-pai/CogVideoX-Fun-V1.1-Reward-LoRAs/CogVideoX-Fun-V1.1-5b-InP-HPS2.1.safetensors"
|
||||
lora_weight = 0.7
|
||||
|
||||
prompt = "Pig with wings flying above a diamond mountain"
|
||||
sample_size = [512, 512]
|
||||
video_length = 49
|
||||
|
||||
transformer = CogVideoXTransformer3DModel.from_pretrained_2d(model_path, subfolder="transformer").to(torch.bfloat16)
|
||||
scheduler = CogVideoXDDIMScheduler.from_pretrained(model_path, subfolder="scheduler")
|
||||
pipeline = CogVideoX_Fun_Pipeline_Inpaint.from_pretrained(
|
||||
model_path, transformer=transformer, scheduler=scheduler, torch_dtype=torch.bfloat16
|
||||
)
|
||||
pipeline.enable_model_cpu_offload()
|
||||
pipeline = merge_lora(pipeline, lora_path, lora_weight)
|
||||
|
||||
generator = torch.Generator(device="cuda").manual_seed(42)
|
||||
input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=sample_size)
|
||||
sample = pipeline(
|
||||
prompt,
|
||||
num_frames = video_length,
|
||||
negative_prompt = "bad detailed",
|
||||
height = sample_size[0],
|
||||
width = sample_size[1],
|
||||
generator = generator,
|
||||
guidance_scale = 7.0,
|
||||
num_inference_steps = 50,
|
||||
video = input_video,
|
||||
mask_video = input_video_mask,
|
||||
).videos
|
||||
|
||||
save_videos_grid(sample, "samples/output.mp4", fps=8)
|
||||
```
|
||||
|
||||
## Training
|
||||
The [training code](./train_reward_lora.py) is based on [train_lora.py](./train_lora.py).
|
||||
We provide [a shell script](./train_reward_lora.sh) to train the HPS v2.1 reward LoRA for CogVideoX-Fun-V1.1-2b-InP,
|
||||
which can be trained on a single A10 with 24GB VRAM. To further reduce the VRAM requirement, please read [Important Args](#important-args).
|
||||
|
||||
### Setup
|
||||
Please read the [quick-start](https://github.com/aigc-apps/CogVideoX-Fun/blob/main/README.md#quick-start) section to setup the CogVideoX-Fun environment.
|
||||
**If you're playing with HPS reward model**, please run the following script to install the dependencies:
|
||||
```bash
|
||||
# For HPS reward model only
|
||||
pip install hpsv2
|
||||
site_packages=$(python -c "import site; print(site.getsitepackages()[0])")
|
||||
wget -O $site_packages/hpsv2/src/open_clip/ https://github.com/tgxs002/HPSv2/raw/refs/heads/master/hpsv2/src/open_clip/bpe_simple_vocab_16e6.txt.gz
|
||||
```
|
||||
|
||||
### Important Args
|
||||
+ `rank`: The size of LoRA model. The higher the LoRA rank, the more parameters it has, and the more it can learn (including some unnecessary information).
|
||||
Bt default, we set the rank to 128. You can lower this value to reduce training GPU memory and the LoRA file size.
|
||||
+ `network_alpha`: A scaling factor changes how the LoRA affect the base model weight. In general, it can be set to half of the `rank`.
|
||||
+ `prompt_path`: The path to the prompt file (in txt format, each line is a prompt) for sampling training videos.
|
||||
We randomly selected 701 prompts from [MovieGenBench](https://github.com/facebookresearch/MovieGenBench/blob/main/benchmark/MovieGenVideoBench.txt).
|
||||
+ `train_sample_height` and `train_sample_width`: The resolution of the sampled training videos. We found
|
||||
training at a 256x256 resolution can generalize to any other resolution. Reducing the resolution can save GPU memory
|
||||
during training, but it is recommended that the resolution should be equal to or greater than the image input resolution of the reward model.
|
||||
+ `reward_fn` and `reward_fn_kwargs`: The reward model name and its keyword arguments. All supported reward models
|
||||
(Aesthetic Predictor [v2](https://github.com/christophschuhmann/improved-aesthetic-predictor)/[v2.5](https://github.com/discus0434/aesthetic-predictor-v2-5),
|
||||
[HPS](https://github.com/tgxs002/HPSv2) v2/v2.1, [PickScore](https://github.com/yuvalkirstain/PickScore) and [MPS](https://github.com/Kwai-Kolors/MPS))
|
||||
can be found in [reward_fn.py](../cogvideox/reward/reward_fn.py).
|
||||
You can also customize your own reward model (e.g., combining aesthetic predictor with HPS).
|
||||
+ `num_decoded_latents` and `num_sampled_frames`: The number of decoded latents (for VAE) and sampled frames (for the reward model).
|
||||
Since CogVideoX-Fun adopts the 3D casual VAE, we found decoding only the first latent to obtain the first frame for computing the reward
|
||||
not only reduces training memory usage but also prevents excessive reward optimization and maintains the dynamics of generated videos.
|
||||
|
||||
## Limitations
|
||||
1. We observe after training to a certain extent, the reward continues to increase, but the quality of the generated videos does not further improve.
|
||||
The model trickly learns some shortcuts (by adding artifacts in the background, i.e., adversarial patches) to increase the reward.
|
||||
2. Currently, there is still a lack of suitable preference models for video generation. Directly using image preference models cannot
|
||||
evaluate preferences along the temporal dimension (such as dynamism and consistency). Further more, We find using image preference models leads to a decrease
|
||||
in the dynamism of generated videos. Although this can be mitigated by computing the reward using only the first frame of the decoded video, the impact still persists.
|
||||
|
||||
## References
|
||||
<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,62 @@
|
||||
export MODEL_NAME="alibaba-pai/CogVideoX-Fun-V1.1-2b-InP"
|
||||
export TRAIN_PROMPT_PATH="MovieGenVideoBench_train.txt"
|
||||
# Performing validation simultaneously with training will increase time and GPU memory usage.
|
||||
export VALIDATION_PROMPT_PATH="MovieGenVideoBench_val.txt"
|
||||
|
||||
# When train model with multi machines, use "--config_file accelerate.yaml" instead of "--mixed_precision='bf16'".
|
||||
accelerate launch --num_processes=1 --mixed_precision="bf16" scripts/train_reward_lora.py \
|
||||
--pretrained_model_name_or_path=$MODEL_NAME \
|
||||
--rank 32 \
|
||||
--network_alpha 16 \
|
||||
--train_batch_size=1 \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--max_train_steps=10000 \
|
||||
--checkpointing_steps=100 \
|
||||
--learning_rate=1e-05 \
|
||||
--seed=42 \
|
||||
--output_dir="output_dir" \
|
||||
--gradient_checkpointing \
|
||||
--vae_gradient_checkpointing \
|
||||
--mixed_precision="bf16" \
|
||||
--adam_weight_decay=3e-2 \
|
||||
--adam_epsilon=1e-10 \
|
||||
--max_grad_norm=0.3 \
|
||||
--prompt_path $TRAIN_PROMPT_PATH \
|
||||
--train_sample_height 224 \
|
||||
--train_sample_width 224 \
|
||||
--video_length 49 \
|
||||
--num_decoded_latents 1 \
|
||||
--num_sampled_frames 1 \
|
||||
--reward_fn "HPSReward" \
|
||||
--reward_fn_kwargs '{"version": "v2.1"}' \
|
||||
--backprop
|
||||
|
||||
# Training command for CogVideoX-Fun-V1.1-2b-InP-HPS2.1.safetensors (with 8 A100 GPUs)
|
||||
# accelerate launch --num_processes=8 --mixed_precision="bf16" --use_deepspeed --deepspeed_config_file config/zero_stage2_config.json scripts/train_reward_lora.py \
|
||||
# --pretrained_model_name_or_path=$MODEL_NAME \
|
||||
# --rank 128 \
|
||||
# --network_alpha 64 \
|
||||
# --train_batch_size=1 \
|
||||
# --gradient_accumulation_steps=1 \
|
||||
# --max_train_steps=10000 \
|
||||
# --checkpointing_steps=100 \
|
||||
# --learning_rate=1e-05 \
|
||||
# --seed=42 \
|
||||
# --output_dir="output_dir" \
|
||||
# --gradient_checkpointing \
|
||||
# --mixed_precision="bf16" \
|
||||
# --adam_weight_decay=3e-2 \
|
||||
# --adam_epsilon=1e-10 \
|
||||
# --max_grad_norm=0.3 \
|
||||
# --prompt_path $TRAIN_PROMPT_PATH \
|
||||
# --train_sample_height 256 \
|
||||
# --train_sample_width 256 \
|
||||
# --video_length 49 \
|
||||
# --validation_prompt_path $VALIDATION_PROMPT_PATH \
|
||||
# --validation_steps 100 \
|
||||
# --validation_batch_size 8 \
|
||||
# --num_decoded_latents 1 \
|
||||
# --num_sampled_frames 1 \
|
||||
# --reward_fn "HPSReward" \
|
||||
# --reward_fn_kwargs '{"version": "v2.1"}' \
|
||||
# --backprop
|
||||
Reference in New Issue
Block a user