From 7352e257efddaddadc8d69e0769ae468b8c9a245 Mon Sep 17 00:00:00 2001 From: primorLee Date: Sat, 22 Aug 2026 05:31:19 +0800 Subject: [PATCH] fix: align CogVideoX training rotary embeddings --- finetune/cogvideox_rope.py | 48 ++++++++++ finetune/models/cogvideox_i2v/lora_trainer.py | 27 ++---- finetune/models/cogvideox_t2v/lora_trainer.py | 26 ++--- tests/finetune/test_cogvideox_rope.py | 96 +++++++++++++++++++ 4 files changed, 158 insertions(+), 39 deletions(-) create mode 100644 finetune/cogvideox_rope.py create mode 100644 tests/finetune/test_cogvideox_rope.py diff --git a/finetune/cogvideox_rope.py b/finetune/cogvideox_rope.py new file mode 100644 index 0000000..d2e979f --- /dev/null +++ b/finetune/cogvideox_rope.py @@ -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, + ) diff --git a/finetune/models/cogvideox_i2v/lora_trainer.py b/finetune/models/cogvideox_i2v/lora_trainer.py index 793cf76..1efe559 100644 --- a/finetune/models/cogvideox_i2v/lora_trainer.py +++ b/finetune/models/cogvideox_i2v/lora_trainer.py @@ -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) diff --git a/finetune/models/cogvideox_t2v/lora_trainer.py b/finetune/models/cogvideox_t2v/lora_trainer.py index 5f0ec1c..088c984 100644 --- a/finetune/models/cogvideox_t2v/lora_trainer.py +++ b/finetune/models/cogvideox_t2v/lora_trainer.py @@ -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) diff --git a/tests/finetune/test_cogvideox_rope.py b/tests/finetune/test_cogvideox_rope.py new file mode 100644 index 0000000..5515e64 --- /dev/null +++ b/tests/finetune/test_cogvideox_rope.py @@ -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()