mirror of
https://github.com/THUDM/CogVideo.git
synced 2026-09-11 11:15:53 +08:00
Merge 7352e257efddaddadc8d69e0769ae468b8c9a245 into 7a1af7154511e0ce4e4be8d62faa8c5e5a3532d2
This commit is contained in:
commit
5d3ca45461
48
finetune/cogvideox_rope.py
Normal file
48
finetune/cogvideox_rope.py
Normal file
@ -0,0 +1,48 @@
|
||||
from typing import Any, Tuple
|
||||
|
||||
import torch
|
||||
from diffusers.models.embeddings import get_3d_rotary_pos_embed
|
||||
from diffusers.pipelines.cogvideo.pipeline_cogvideox import get_resize_crop_region_for_grid
|
||||
|
||||
|
||||
def prepare_cogvideox_rotary_positional_embeddings(
|
||||
height: int,
|
||||
width: int,
|
||||
num_frames: int,
|
||||
transformer_config: Any,
|
||||
vae_scale_factor_spatial: int,
|
||||
device: torch.device,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Build the same rotary embeddings used by the CogVideoX inference pipelines."""
|
||||
patch_size = transformer_config.patch_size
|
||||
patch_size_t = transformer_config.patch_size_t
|
||||
grid_height = height // (vae_scale_factor_spatial * patch_size)
|
||||
grid_width = width // (vae_scale_factor_spatial * patch_size)
|
||||
|
||||
base_size_width = transformer_config.sample_width // patch_size
|
||||
base_size_height = transformer_config.sample_height // patch_size
|
||||
|
||||
if patch_size_t is None:
|
||||
# CogVideoX 1.0 interpolates the pretrained spatial grid for other resolutions.
|
||||
grid_crops_coords = get_resize_crop_region_for_grid(
|
||||
(grid_height, grid_width), base_size_width, base_size_height
|
||||
)
|
||||
return get_3d_rotary_pos_embed(
|
||||
embed_dim=transformer_config.attention_head_dim,
|
||||
crops_coords=grid_crops_coords,
|
||||
grid_size=(grid_height, grid_width),
|
||||
temporal_size=num_frames,
|
||||
device=device,
|
||||
)
|
||||
|
||||
# CogVideoX 1.5 slices its pretrained spatial grid after temporal patching.
|
||||
base_num_frames = (num_frames + patch_size_t - 1) // patch_size_t
|
||||
return get_3d_rotary_pos_embed(
|
||||
embed_dim=transformer_config.attention_head_dim,
|
||||
crops_coords=None,
|
||||
grid_size=(grid_height, grid_width),
|
||||
temporal_size=base_num_frames,
|
||||
grid_type="slice",
|
||||
max_size=(base_size_height, base_size_width),
|
||||
device=device,
|
||||
)
|
||||
@ -7,12 +7,12 @@ from diffusers import (
|
||||
CogVideoXImageToVideoPipeline,
|
||||
CogVideoXTransformer3DModel,
|
||||
)
|
||||
from diffusers.models.embeddings import get_3d_rotary_pos_embed
|
||||
from PIL import Image
|
||||
from numpy import dtype
|
||||
from transformers import AutoTokenizer, T5EncoderModel
|
||||
from typing_extensions import override
|
||||
|
||||
from finetune.cogvideox_rope import prepare_cogvideox_rotary_positional_embeddings
|
||||
from finetune.schemas import Components
|
||||
from finetune.trainer import Trainer
|
||||
from finetune.utils import unwrap_model
|
||||
@ -246,27 +246,14 @@ class CogVideoXI2VLoraTrainer(Trainer):
|
||||
vae_scale_factor_spatial: int,
|
||||
device: torch.device,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
grid_height = height // (vae_scale_factor_spatial * transformer_config.patch_size)
|
||||
grid_width = width // (vae_scale_factor_spatial * transformer_config.patch_size)
|
||||
|
||||
if transformer_config.patch_size_t is None:
|
||||
base_num_frames = num_frames
|
||||
else:
|
||||
base_num_frames = (
|
||||
num_frames + transformer_config.patch_size_t - 1
|
||||
) // transformer_config.patch_size_t
|
||||
|
||||
freqs_cos, freqs_sin = get_3d_rotary_pos_embed(
|
||||
embed_dim=transformer_config.attention_head_dim,
|
||||
crops_coords=None,
|
||||
grid_size=(grid_height, grid_width),
|
||||
temporal_size=base_num_frames,
|
||||
grid_type="slice",
|
||||
max_size=(grid_height, grid_width),
|
||||
return prepare_cogvideox_rotary_positional_embeddings(
|
||||
height=height,
|
||||
width=width,
|
||||
num_frames=num_frames,
|
||||
transformer_config=transformer_config,
|
||||
vae_scale_factor_spatial=vae_scale_factor_spatial,
|
||||
device=device,
|
||||
)
|
||||
|
||||
return freqs_cos, freqs_sin
|
||||
|
||||
|
||||
register("cogvideox-i2v", "lora", CogVideoXI2VLoraTrainer)
|
||||
|
||||
@ -7,11 +7,11 @@ from diffusers import (
|
||||
CogVideoXPipeline,
|
||||
CogVideoXTransformer3DModel,
|
||||
)
|
||||
from diffusers.models.embeddings import get_3d_rotary_pos_embed
|
||||
from PIL import Image
|
||||
from transformers import AutoTokenizer, T5EncoderModel
|
||||
from typing_extensions import override
|
||||
|
||||
from finetune.cogvideox_rope import prepare_cogvideox_rotary_positional_embeddings
|
||||
from finetune.schemas import Components
|
||||
from finetune.trainer import Trainer
|
||||
from finetune.utils import unwrap_model
|
||||
@ -203,26 +203,14 @@ class CogVideoXT2VLoraTrainer(Trainer):
|
||||
vae_scale_factor_spatial: int,
|
||||
device: torch.device,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
grid_height = height // (vae_scale_factor_spatial * transformer_config.patch_size)
|
||||
grid_width = width // (vae_scale_factor_spatial * transformer_config.patch_size)
|
||||
|
||||
if transformer_config.patch_size_t is None:
|
||||
base_num_frames = num_frames
|
||||
else:
|
||||
base_num_frames = (
|
||||
num_frames + transformer_config.patch_size_t - 1
|
||||
) // transformer_config.patch_size_t
|
||||
freqs_cos, freqs_sin = get_3d_rotary_pos_embed(
|
||||
embed_dim=transformer_config.attention_head_dim,
|
||||
crops_coords=None,
|
||||
grid_size=(grid_height, grid_width),
|
||||
temporal_size=base_num_frames,
|
||||
grid_type="slice",
|
||||
max_size=(grid_height, grid_width),
|
||||
return prepare_cogvideox_rotary_positional_embeddings(
|
||||
height=height,
|
||||
width=width,
|
||||
num_frames=num_frames,
|
||||
transformer_config=transformer_config,
|
||||
vae_scale_factor_spatial=vae_scale_factor_spatial,
|
||||
device=device,
|
||||
)
|
||||
|
||||
return freqs_cos, freqs_sin
|
||||
|
||||
|
||||
register("cogvideox-t2v", "lora", CogVideoXT2VLoraTrainer)
|
||||
|
||||
96
tests/finetune/test_cogvideox_rope.py
Normal file
96
tests/finetune/test_cogvideox_rope.py
Normal file
@ -0,0 +1,96 @@
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
import torch
|
||||
from diffusers.models.embeddings import get_3d_rotary_pos_embed
|
||||
from diffusers.pipelines.cogvideo.pipeline_cogvideox import get_resize_crop_region_for_grid
|
||||
|
||||
from finetune.cogvideox_rope import prepare_cogvideox_rotary_positional_embeddings
|
||||
|
||||
|
||||
class CogVideoXRotaryEmbeddingTests(unittest.TestCase):
|
||||
def test_cogvideox_1_0_non_base_resolution_matches_inference(self):
|
||||
config = SimpleNamespace(
|
||||
patch_size=2,
|
||||
patch_size_t=None,
|
||||
sample_height=60,
|
||||
sample_width=90,
|
||||
attention_head_dim=64,
|
||||
)
|
||||
height, width, num_frames = 576, 768, 5
|
||||
|
||||
actual = prepare_cogvideox_rotary_positional_embeddings(
|
||||
height,
|
||||
width,
|
||||
num_frames,
|
||||
config,
|
||||
vae_scale_factor_spatial=8,
|
||||
device=torch.device("cpu"),
|
||||
)
|
||||
|
||||
grid_height = height // (8 * config.patch_size)
|
||||
grid_width = width // (8 * config.patch_size)
|
||||
base_size_width = config.sample_width // config.patch_size
|
||||
base_size_height = config.sample_height // config.patch_size
|
||||
grid_crops_coords = get_resize_crop_region_for_grid(
|
||||
(grid_height, grid_width), base_size_width, base_size_height
|
||||
)
|
||||
expected = get_3d_rotary_pos_embed(
|
||||
embed_dim=config.attention_head_dim,
|
||||
crops_coords=grid_crops_coords,
|
||||
grid_size=(grid_height, grid_width),
|
||||
temporal_size=num_frames,
|
||||
device=torch.device("cpu"),
|
||||
)
|
||||
|
||||
legacy = get_3d_rotary_pos_embed(
|
||||
embed_dim=config.attention_head_dim,
|
||||
crops_coords=None,
|
||||
grid_size=(grid_height, grid_width),
|
||||
temporal_size=num_frames,
|
||||
grid_type="slice",
|
||||
max_size=(grid_height, grid_width),
|
||||
device=torch.device("cpu"),
|
||||
)
|
||||
|
||||
for actual_tensor, expected_tensor, legacy_tensor in zip(actual, expected, legacy):
|
||||
self.assertTrue(torch.equal(actual_tensor, expected_tensor))
|
||||
self.assertFalse(torch.equal(actual_tensor, legacy_tensor))
|
||||
|
||||
def test_cogvideox_1_5_uses_pretrained_grid_and_temporal_patching(self):
|
||||
config = SimpleNamespace(
|
||||
patch_size=2,
|
||||
patch_size_t=2,
|
||||
sample_height=96,
|
||||
sample_width=170,
|
||||
attention_head_dim=64,
|
||||
)
|
||||
sentinel = (object(), object())
|
||||
|
||||
with patch(
|
||||
"finetune.cogvideox_rope.get_3d_rotary_pos_embed", return_value=sentinel
|
||||
) as get_rope:
|
||||
actual = prepare_cogvideox_rotary_positional_embeddings(
|
||||
720,
|
||||
1280,
|
||||
81,
|
||||
config,
|
||||
vae_scale_factor_spatial=8,
|
||||
device=torch.device("cpu"),
|
||||
)
|
||||
|
||||
self.assertIs(actual, sentinel)
|
||||
get_rope.assert_called_once_with(
|
||||
embed_dim=64,
|
||||
crops_coords=None,
|
||||
grid_size=(45, 80),
|
||||
temporal_size=41,
|
||||
grid_type="slice",
|
||||
max_size=(48, 85),
|
||||
device=torch.device("cpu"),
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Loading…
x
Reference in New Issue
Block a user