mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
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:
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user