Merge branch 'main' into mochi-1-tuner

This commit is contained in:
sayakpaul
2024-11-19 10:29:48 +05:30
6 changed files with 162 additions and 96 deletions
+16 -2
View File
@@ -53,9 +53,23 @@ video = pipe("<my-awesome-prompt>").frames[0]
export_to_video(video, "output.mp4", fps=8)
```
You can also check if your LoRA is correctly mounted [here](tests/test_lora_inference.py).
For Image-to-Video LoRAs trained with multiresolution videos, one must also add the following lines (see [this](https://github.com/a-r-r-o-w/cogvideox-factory/issues/26) Issue for more details):
**Note:** For Image-to-Video finetuning, you must install diffusers from [this](https://github.com/huggingface/diffusers/pull/9482) branch (which adds lora loading support in CogVideoX image-to-video) until it is merged.
```python
from diffusers import CogVideoXImageToVideoPipeline
pipe = CogVideoXImageToVideoPipeline.from_pretrained(
"THUDM/CogVideoX-5b-I2V", torch_dtype=torch.bfloat16
).to("cuda")
# ...
del pipe.transformer.patch_embed.pos_embedding
pipe.transformer.patch_embed.use_learned_positional_embeddings = False
pipe.transformer.config.use_learned_positional_embeddings = False
```
You can also check if your LoRA is correctly mounted [here](tests/test_lora_inference.py).
Below we provide additional sections detailing on more options explored in this repository. They all attempt to make fine-tuning for video models as accessible as possible by reducing memory requirements as much as possible.
+17 -1
View File
@@ -131,7 +131,13 @@ def _get_validation_args(parser: argparse.ArgumentParser) -> None:
parser.add_argument(
"--validation_epochs",
type=int,
default=50,
default=None,
help="Run validation every X training epochs. Validation consists of running the validation prompt `args.num_validation_videos` times.",
)
parser.add_argument(
"--validation_steps",
type=int,
default=None,
help="Run validation every X training steps. Validation consists of running the validation prompt `args.num_validation_videos` times.",
)
parser.add_argument(
@@ -323,6 +329,16 @@ def _get_training_args(parser: argparse.ArgumentParser) -> None:
default=0.05,
help="Image condition dropout probability when finetuning image-to-video.",
)
parser.add_argument(
"--ignore_learned_positional_embeddings",
action="store_true",
default=False,
help=(
"Whether to ignore the learned positional embeddings when training CogVideoX Image-to-Video. This setting "
"should be used when performing multi-resolution training, because CogVideoX-I2V does not support it "
"otherwise. Please read the comments in https://github.com/a-r-r-o-w/cogvideox-factory/issues/26 to understand why."
),
)
def _get_optimizer_args(parser: argparse.ArgumentParser) -> None:
+91 -65
View File
@@ -46,7 +46,6 @@ from diffusers.optimization import get_scheduler
from diffusers.training_utils import cast_training_params
from diffusers.utils import convert_unet_state_dict_to_peft, export_to_video, load_image
from diffusers.utils.hub_utils import load_or_create_model_card, populate_model_card
from diffusers.utils.torch_utils import is_compiled_module
from huggingface_hub import create_repo, upload_folder
from peft import LoraConfig, get_peft_model_state_dict, set_peft_model_state_dict
from torch.utils.data import DataLoader
@@ -57,7 +56,14 @@ from transformers import AutoTokenizer, T5EncoderModel
from args import get_args # isort:skip
from dataset import BucketSampler, VideoDatasetWithResizing, VideoDatasetWithResizeAndRectangleCrop # isort:skip
from text_encoder import compute_prompt_embeddings # isort:skip
from utils import get_gradient_norm, get_optimizer, prepare_rotary_positional_embeddings, print_memory, reset_memory # isort:skip
from utils import (
get_gradient_norm,
get_optimizer,
prepare_rotary_positional_embeddings,
print_memory,
reset_memory,
unwrap_model,
)
logger = get_logger(__name__)
@@ -155,7 +161,6 @@ def log_validation(
pipe: CogVideoXImageToVideoPipeline,
args: Dict[str, Any],
pipeline_args: Dict[str, Any],
epoch,
is_final_validation: bool = False,
):
logger.info(
@@ -201,6 +206,64 @@ def log_validation(
return videos
def run_validation(
args: Dict[str, Any],
accelerator: Accelerator,
transformer,
scheduler,
model_config: Dict[str, Any],
weight_dtype: torch.dtype,
) -> None:
accelerator.print("===== Memory before validation =====")
print_memory(accelerator.device)
torch.cuda.synchronize(accelerator.device)
pipe = CogVideoXImageToVideoPipeline.from_pretrained(
args.pretrained_model_name_or_path,
transformer=unwrap_model(accelerator, transformer),
scheduler=scheduler,
revision=args.revision,
variant=args.variant,
torch_dtype=weight_dtype,
)
if args.enable_slicing:
pipe.vae.enable_slicing()
if args.enable_tiling:
pipe.vae.enable_tiling()
if args.enable_model_cpu_offload:
pipe.enable_model_cpu_offload()
validation_prompts = args.validation_prompt.split(args.validation_prompt_separator)
validation_images = args.validation_images.split(args.validation_prompt_separator)
for validation_image, validation_prompt in zip(validation_images, validation_prompts):
pipeline_args = {
"image": load_image(validation_image),
"prompt": validation_prompt,
"guidance_scale": args.guidance_scale,
"use_dynamic_cfg": args.use_dynamic_cfg,
"height": args.height,
"width": args.width,
"max_sequence_length": model_config.max_text_seq_length,
}
log_validation(
pipe=pipe,
args=args,
accelerator=accelerator,
pipeline_args=pipeline_args,
)
accelerator.print("===== Memory after validation =====")
print_memory(accelerator.device)
reset_memory(accelerator.device)
del pipe
gc.collect()
torch.cuda.empty_cache()
torch.cuda.synchronize(accelerator.device)
class CollateFunction:
def __init__(self, weight_dtype: torch.dtype, load_tensors: bool) -> None:
self.weight_dtype = weight_dtype
@@ -308,6 +371,12 @@ def main(args):
variant=args.variant,
)
# These changes will also be required when trying to run inference with the trained lora
if args.ignore_learned_positional_embeddings:
del transformer.patch_embed.pos_embedding
transformer.patch_embed.use_learned_positional_embeddings = False
transformer.config.use_learned_positional_embeddings = False
vae = AutoencoderKLCogVideoX.from_pretrained(
args.pretrained_model_name_or_path,
subfolder="vae",
@@ -373,19 +442,14 @@ def main(args):
)
transformer.add_adapter(transformer_lora_config)
def unwrap_model(model):
model = accelerator.unwrap_model(model)
model = model._orig_mod if is_compiled_module(model) else model
return model
# create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format
def save_model_hook(models, weights, output_dir):
if accelerator.is_main_process:
transformer_lora_layers_to_save = None
for model in models:
if isinstance(unwrap_model(model), type(unwrap_model(transformer))):
model = unwrap_model(model)
if isinstance(unwrap_model(accelerator, model), type(unwrap_model(accelerator, transformer))):
model = unwrap_model(accelerator, model)
transformer_lora_layers_to_save = get_peft_model_state_dict(model)
else:
raise ValueError(f"unexpected save model: {model.__class__}")
@@ -407,10 +471,10 @@ def main(args):
while len(models) > 0:
model = models.pop()
if isinstance(unwrap_model(model), type(unwrap_model(transformer))):
transformer_ = unwrap_model(model)
if isinstance(unwrap_model(accelerator, model), type(unwrap_model(accelerator, transformer))):
transformer_ = unwrap_model(accelerator, model)
else:
raise ValueError(f"Unexpected save model: {unwrap_model(model).__class__}")
raise ValueError(f"Unexpected save model: {unwrap_model(accelerator, model).__class__}")
else:
transformer_ = CogVideoXTransformer3DModel.from_pretrained(
args.pretrained_model_name_or_path, subfolder="transformer"
@@ -776,6 +840,7 @@ def main(args):
progress_bar.update(1)
global_step += 1
# Checkpointing
if accelerator.is_main_process or accelerator.distributed_type == DistributedType.DEEPSPEED:
if global_step % args.checkpointing_steps == 0:
# _before_ saving state, check if this save would set us over the `checkpoints_total_limit`
@@ -802,6 +867,13 @@ def main(args):
accelerator.save_state(save_path)
logger.info(f"Saved state to {save_path}")
# Validation
should_run_validation = args.validation_prompt is not None and (
args.validation_steps is not None and global_step % args.validation_steps == 0
)
if should_run_validation:
run_validation(args, accelerator, transformer, scheduler, model_config, weight_dtype)
last_lr = lr_scheduler.get_last_lr()[0] if lr_scheduler is not None else args.learning_rate
logs = {"loss": loss.detach().item(), "lr": last_lr}
# gradnorm + deepspeed: https://github.com/microsoft/DeepSpeed/issues/4555
@@ -819,61 +891,16 @@ def main(args):
break
if accelerator.is_main_process:
if args.validation_prompt is not None and (epoch + 1) % args.validation_epochs == 0:
accelerator.print("===== Memory before validation =====")
print_memory(accelerator.device)
torch.cuda.synchronize(accelerator.device)
pipe = CogVideoXImageToVideoPipeline.from_pretrained(
args.pretrained_model_name_or_path,
transformer=unwrap_model(transformer),
scheduler=scheduler,
revision=args.revision,
variant=args.variant,
torch_dtype=weight_dtype,
)
if args.enable_slicing:
pipe.vae.enable_slicing()
if args.enable_tiling:
pipe.vae.enable_tiling()
if args.enable_model_cpu_offload:
pipe.enable_model_cpu_offload()
validation_prompts = args.validation_prompt.split(args.validation_prompt_separator)
validation_images = args.validation_images.split(args.validation_prompt_separator)
for validation_image, validation_prompt in zip(validation_images, validation_prompts):
pipeline_args = {
"image": load_image(validation_image),
"prompt": validation_prompt,
"guidance_scale": args.guidance_scale,
"use_dynamic_cfg": args.use_dynamic_cfg,
"height": args.height,
"width": args.width,
"max_sequence_length": model_config.max_text_seq_length,
}
log_validation(
pipe=pipe,
args=args,
accelerator=accelerator,
pipeline_args=pipeline_args,
epoch=epoch,
)
accelerator.print("===== Memory after validation =====")
print_memory(accelerator.device)
reset_memory(accelerator.device)
del pipe
gc.collect()
torch.cuda.empty_cache()
torch.cuda.synchronize(accelerator.device)
should_run_validation = args.validation_prompt is not None and (
args.validation_epochs is not None and (epoch + 1) % args.validation_epochs == 0
)
if should_run_validation:
run_validation(args, accelerator, transformer, scheduler, model_config, weight_dtype)
accelerator.wait_for_everyone()
if accelerator.is_main_process:
transformer = unwrap_model(transformer)
transformer = unwrap_model(accelerator, transformer)
dtype = (
torch.float16
if args.mixed_precision == "fp16"
@@ -944,7 +971,6 @@ def main(args):
pipe=pipe,
args=args,
pipeline_args=pipeline_args,
epoch=epoch,
is_final_validation=True,
)
validation_outputs.extend(video)
+15 -14
View File
@@ -45,7 +45,6 @@ from diffusers.optimization import get_scheduler
from diffusers.training_utils import cast_training_params
from diffusers.utils import convert_unet_state_dict_to_peft, export_to_video
from diffusers.utils.hub_utils import load_or_create_model_card, populate_model_card
from diffusers.utils.torch_utils import is_compiled_module
from huggingface_hub import create_repo, upload_folder
from peft import LoraConfig, get_peft_model_state_dict, set_peft_model_state_dict
from torch.utils.data import DataLoader
@@ -56,7 +55,14 @@ from transformers import AutoTokenizer, T5EncoderModel
from args import get_args # isort:skip
from dataset import BucketSampler, VideoDatasetWithResizing, VideoDatasetWithResizeAndRectangleCrop # isort:skip
from text_encoder import compute_prompt_embeddings # isort:skip
from utils import get_gradient_norm, get_optimizer, prepare_rotary_positional_embeddings, print_memory, reset_memory # isort:skip
from utils import (
get_gradient_norm,
get_optimizer,
prepare_rotary_positional_embeddings,
print_memory,
reset_memory,
unwrap_model,
) # isort:skip
logger = get_logger(__name__)
@@ -366,19 +372,14 @@ def main(args):
)
transformer.add_adapter(transformer_lora_config)
def unwrap_model(model):
model = accelerator.unwrap_model(model)
model = model._orig_mod if is_compiled_module(model) else model
return model
# create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format
def save_model_hook(models, weights, output_dir):
if accelerator.is_main_process:
transformer_lora_layers_to_save = None
for model in models:
if isinstance(unwrap_model(model), type(unwrap_model(transformer))):
model = unwrap_model(model)
if isinstance(unwrap_model(accelerator, model), type(unwrap_model(accelerator, transformer))):
model = unwrap_model(accelerator, model)
transformer_lora_layers_to_save = get_peft_model_state_dict(model)
else:
raise ValueError(f"unexpected save model: {model.__class__}")
@@ -400,10 +401,10 @@ def main(args):
while len(models) > 0:
model = models.pop()
if isinstance(unwrap_model(model), type(unwrap_model(transformer))):
transformer_ = unwrap_model(model)
if isinstance(unwrap_model(accelerator, model), type(unwrap_model(accelerator, transformer))):
transformer_ = unwrap_model(accelerator, model)
else:
raise ValueError(f"Unexpected save model: {unwrap_model(model).__class__}")
raise ValueError(f"Unexpected save model: {unwrap_model(accelerator, model).__class__}")
else:
transformer_ = CogVideoXTransformer3DModel.from_pretrained(
args.pretrained_model_name_or_path, subfolder="transformer"
@@ -797,7 +798,7 @@ def main(args):
pipe = CogVideoXPipeline.from_pretrained(
args.pretrained_model_name_or_path,
transformer=unwrap_model(transformer),
transformer=unwrap_model(accelerator, transformer),
scheduler=scheduler,
revision=args.revision,
variant=args.variant,
@@ -842,7 +843,7 @@ def main(args):
accelerator.wait_for_everyone()
if accelerator.is_main_process:
transformer = unwrap_model(transformer)
transformer = unwrap_model(accelerator, transformer)
dtype = (
torch.float16
if args.mixed_precision == "fp16"
+15 -14
View File
@@ -45,7 +45,6 @@ from diffusers.optimization import get_scheduler
from diffusers.training_utils import cast_training_params
from diffusers.utils import export_to_video
from diffusers.utils.hub_utils import load_or_create_model_card, populate_model_card
from diffusers.utils.torch_utils import is_compiled_module
from huggingface_hub import create_repo, upload_folder
from torch.utils.data import DataLoader
from tqdm.auto import tqdm
@@ -55,7 +54,14 @@ from transformers import AutoTokenizer, T5EncoderModel
from args import get_args # isort:skip
from dataset import BucketSampler, VideoDatasetWithResizing, VideoDatasetWithResizeAndRectangleCrop # isort:skip
from text_encoder import compute_prompt_embeddings # isort:skip
from utils import get_gradient_norm, get_optimizer, prepare_rotary_positional_embeddings, print_memory, reset_memory # isort:skip
from utils import (
get_gradient_norm,
get_optimizer,
prepare_rotary_positional_embeddings,
print_memory,
reset_memory,
unwrap_model,
) # isort:skip
logger = get_logger(__name__)
@@ -346,18 +352,13 @@ def main(args):
if args.gradient_checkpointing:
transformer.enable_gradient_checkpointing()
def unwrap_model(model):
model = accelerator.unwrap_model(model)
model = model._orig_mod if is_compiled_module(model) else model
return model
# create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format
def save_model_hook(models, weights, output_dir):
if accelerator.is_main_process:
for model in models:
if isinstance(unwrap_model(model), type(unwrap_model(transformer))):
if isinstance(unwrap_model(accelerator, model), type(unwrap_model(accelerator, transformer))):
model: CogVideoXTransformer3DModel
model = unwrap_model(model)
model = unwrap_model(accelerator, model)
model.save_pretrained(
os.path.join(output_dir, "transformer"), safe_serialization=True, max_shard_size="5GB"
)
@@ -377,10 +378,10 @@ def main(args):
while len(models) > 0:
model = models.pop()
if isinstance(unwrap_model(model), type(unwrap_model(transformer))):
transformer_ = unwrap_model(model)
if isinstance(unwrap_model(accelerator, model), type(unwrap_model(accelerator, transformer))):
transformer_ = unwrap_model(accelerator, model)
else:
raise ValueError(f"Unexpected save model: {unwrap_model(model).__class__}")
raise ValueError(f"Unexpected save model: {unwrap_model(accelerator, model).__class__}")
else:
with init_empty_weights():
transformer_ = CogVideoXTransformer3DModel.from_config(
@@ -763,7 +764,7 @@ def main(args):
pipe = CogVideoXPipeline.from_pretrained(
args.pretrained_model_name_or_path,
transformer=unwrap_model(transformer),
transformer=unwrap_model(accelerator, transformer),
scheduler=scheduler,
revision=args.revision,
variant=args.variant,
@@ -809,7 +810,7 @@ def main(args):
accelerator.wait_for_everyone()
if accelerator.is_main_process:
transformer = unwrap_model(transformer)
transformer = unwrap_model(accelerator, transformer)
dtype = (
torch.float16
if args.mixed_precision == "fp16"
+8
View File
@@ -3,8 +3,10 @@ import inspect
from typing import Optional, Tuple, Union
import torch
from accelerate import Accelerator
from accelerate.logging import get_logger
from diffusers.models.embeddings import get_3d_rotary_pos_embed
from diffusers.utils.torch_utils import is_compiled_module
logger = get_logger(__name__)
@@ -233,3 +235,9 @@ def print_memory(device: Union[str, torch.device]) -> None:
print(f"{memory_allocated=:.3f} GB")
print(f"{max_memory_allocated=:.3f} GB")
print(f"{max_memory_reserved=:.3f} GB")
def unwrap_model(accelerator: Accelerator, model):
model = accelerator.unwrap_model(model)
model = model._orig_mod if is_compiled_module(model) else model
return model