Add Reward LoRA Training (#71)

This commit is contained in:
hkz
2024-11-22 09:45:59 +08:00
committed by GitHub
parent 0788cd4e8b
commit 2114d906df
14 changed files with 2740 additions and 7 deletions
+42 -1
View File
@@ -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
View File
@@ -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;">
+5 -5
View File
@@ -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):
+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)
+382
View File
@@ -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)}")
+258
View File
@@ -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
+62
View File
@@ -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