adaption for CogVideoX1.5 (#92)

* adaption for CogVideoX1.5

* add patch_size_t in full finetuning of T2V and lora finetuning of I2V

* Update training/args.py

Co-authored-by: Sayak Paul <spsayakpaul@gmail.com>

---------

Co-authored-by: Sayak Paul <spsayakpaul@gmail.com>
This commit is contained in:
jiashenggu
2024-11-23 08:58:20 +08:00
committed by GitHub
parent d63a826f37
commit 0b80dba1de
5 changed files with 26 additions and 7 deletions
+1
View File
@@ -78,6 +78,7 @@ def _get_dataset_args(parser: argparse.ArgumentParser) -> None:
nargs="+",
type=int,
default=[49],
help="CogVideoX1.5 need to guarantee that ((num_frames - 1) // self.vae_scale_factor_temporal + 1) % patch_size_t != 0, such as 53"
)
parser.add_argument(
"--load_tensors",
@@ -787,6 +787,7 @@ def main(args):
num_frames=num_frames,
vae_scale_factor_spatial=VAE_SCALE_FACTOR_SPATIAL,
patch_size=model_config.patch_size,
patch_size_t=model_config.patch_size_t if hasattr(model_config, "patch_size_t") else None,
attention_head_dim=model_config.attention_head_dim,
device=accelerator.device,
)
+1
View File
@@ -696,6 +696,7 @@ def main(args):
num_frames=num_frames,
vae_scale_factor_spatial=VAE_SCALE_FACTOR_SPATIAL,
patch_size=model_config.patch_size,
patch_size_t=model_config.patch_size_t if hasattr(model_config, "patch_size_t") else None,
attention_head_dim=model_config.attention_head_dim,
device=accelerator.device,
)
+1
View File
@@ -662,6 +662,7 @@ def main(args):
num_frames=num_frames,
vae_scale_factor_spatial=VAE_SCALE_FACTOR_SPATIAL,
patch_size=model_config.patch_size,
patch_size_t=model_config.patch_size_t if hasattr(model_config, "patch_size_t") else None,
attention_head_dim=model_config.attention_head_dim,
device=accelerator.device,
)
+22 -7
View File
@@ -198,6 +198,7 @@ def prepare_rotary_positional_embeddings(
num_frames: int,
vae_scale_factor_spatial: int = 8,
patch_size: int = 2,
patch_size_t: int = None,
attention_head_dim: int = 64,
device: Optional[torch.device] = None,
base_height: int = 480,
@@ -207,14 +208,28 @@ def prepare_rotary_positional_embeddings(
grid_width = width // (vae_scale_factor_spatial * patch_size)
base_size_width = base_width // (vae_scale_factor_spatial * patch_size)
base_size_height = base_height // (vae_scale_factor_spatial * patch_size)
if patch_size_t is None:
# CogVideoX 1.0
grid_crops_coords = get_resize_crop_region_for_grid((grid_height, grid_width), base_size_width, base_size_height)
freqs_cos, freqs_sin = get_3d_rotary_pos_embed(
embed_dim=attention_head_dim,
crops_coords=grid_crops_coords,
grid_size=(grid_height, grid_width),
temporal_size=num_frames,
)
else:
# CogVideoX 1.5
base_num_frames = (num_frames + patch_size_t - 1) // patch_size_t
freqs_cos, freqs_sin = get_3d_rotary_pos_embed(
embed_dim=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),
)
grid_crops_coords = get_resize_crop_region_for_grid((grid_height, grid_width), base_size_width, base_size_height)
freqs_cos, freqs_sin = get_3d_rotary_pos_embed(
embed_dim=attention_head_dim,
crops_coords=grid_crops_coords,
grid_size=(grid_height, grid_width),
temporal_size=num_frames,
)
freqs_cos = freqs_cos.to(device=device)
freqs_sin = freqs_sin.to(device=device)