Merge 7352e257efddaddadc8d69e0769ae468b8c9a245 into 7a1af7154511e0ce4e4be8d62faa8c5e5a3532d2

This commit is contained in:
primorLee 2026-08-22 05:41:51 +08:00 committed by GitHub
commit 5d3ca45461
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
4 changed files with 158 additions and 39 deletions

View 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,
)

View File

@ -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)

View File

@ -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)

View 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()