Merge pull request #106 from a-r-r-o-w/resume-mochi-1

feat: support checkpointing saving and loading
This commit is contained in:
Sayak Paul
2024-12-01 20:43:25 +05:30
committed by GitHub
3 changed files with 59 additions and 11 deletions
+4 -1
View File
@@ -15,6 +15,10 @@ Now you can make Mochi-1 your own with `diffusers`, too 🤗 🧨
We provide a minimal and faithful reimplementation of the [Mochi-1 original fine-tuner](https://github.com/genmoai/mochi/tree/aba74c1b5e0755b1fa3343d9e4bd22e89de77ab1/demos/fine_tuner). As usual, we leverage `peft` for things LoRA in our implementation.
**Updates**
December 1 2024: Support for checkpoint saving and loading.
## Getting started
Install the dependencies: `pip install -r requirements.txt`. Also make sure your `diffusers` installation is from the current `main`.
@@ -98,7 +102,6 @@ export_to_video(video)
Our script currently doesn't leverage `accelerate` and some of its consequences are detailed below:
* No support for distributed training.
* No intermediate checkpoint saving and loading support.
* `train_batch_size > 1` are supported but can potentially lead to OOMs because we currently don't have gradient accumulation support.
* No support for 8bit optimizers (but should be relatively easy to add).
+11 -6
View File
@@ -197,6 +197,16 @@ def _get_training_args(parser: argparse.ArgumentParser) -> None:
default=200,
help="Number of steps for the warmup in the lr scheduler.",
)
parser.add_argument(
"--checkpointing_steps",
type=int,
default=None,
)
parser.add_argument(
"--resume_from_checkpoint",
type=str,
default=None,
)
def _get_optimizer_args(parser: argparse.ArgumentParser) -> None:
@@ -242,12 +252,7 @@ def _get_configuration_args(parser: argparse.ArgumentParser) -> None:
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
),
)
parser.add_argument(
"--report_to",
type=str,
default=None,
help="If logging to wandb."
)
parser.add_argument("--report_to", type=str, default=None, help="If logging to wandb.")
def get_args():
+44 -4
View File
@@ -31,7 +31,7 @@ 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 huggingface_hub import create_repo, upload_folder
from peft import LoraConfig, get_peft_model_state_dict
from peft import LoraConfig, get_peft_model_state_dict, set_peft_model_state_dict
from torch.utils.data import DataLoader
from tqdm.auto import tqdm
@@ -215,6 +215,19 @@ def cast_dit(model, dtype):
return model
def save_checkpoint(model, optimizer, lr_scheduler, global_step, checkpoint_path):
lora_state_dict = get_peft_model_state_dict(model)
torch.save(
{
"state_dict": lora_state_dict,
"optimizer": optimizer.state_dict(),
"lr_scheduler": lr_scheduler.state_dict(),
"global_step": global_step,
},
checkpoint_path,
)
class CollateFunction:
def __init__(self, caption_dropout: float = None) -> None:
self.caption_dropout = caption_dropout
@@ -244,7 +257,7 @@ class CollateFunction:
def main(args):
if not torch.cuda.is_available():
raise ValueError("Not supported without CUDA.")
if args.report_to == "wandb" and args.hub_token is not None:
raise ValueError(
"You cannot use both --report_to=wandb and --hub_token due to a security risk of exposing your token."
@@ -346,6 +359,23 @@ def main(args):
tracker_name = args.tracker_name or "mochi-1-lora"
wandb_run = wandb.init(project=tracker_name, config=vars(args))
# Resume from checkpoint if specified
if args.resume_from_checkpoint:
checkpoint = torch.load(args.resume_from_checkpoint, map_location="cpu", weights_only=True)
if "global_step" in checkpoint:
global_step = checkpoint["global_step"]
if "optimizer" in checkpoint:
optimizer.load_state_dict(checkpoint["optimizer"])
if "lr_scheduler" in checkpoint:
lr_scheduler.load_state_dict(checkpoint["lr_scheduler"])
set_peft_model_state_dict(transformer, checkpoint["state_dict"])
print(f"Resuming from checkpoint: {args.resume_from_checkpoint}")
print(f"Resuming from global step: {global_step}")
else:
global_step = 0
print("===== Memory before training =====")
reset_memory("cuda")
print_memory("cuda")
@@ -361,7 +391,6 @@ def main(args):
print(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}")
print(f" Total optimization steps = {args.max_train_steps}")
global_step = 0
first_epoch = 0
progress_bar = tqdm(
range(0, args.max_train_steps),
@@ -415,6 +444,17 @@ def main(args):
if wandb_run:
wandb_run.log(logs, step=global_step)
if args.checkpointing_steps is not None and global_step % args.checkpointing_steps == 0:
print(f"Saving checkpoint at step {global_step}")
checkpoint_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.pt")
save_checkpoint(
transformer,
optimizer,
lr_scheduler,
global_step,
checkpoint_path,
)
if global_step >= args.max_train_steps:
break
@@ -546,7 +586,7 @@ def main(args):
repo_id=repo_id,
folder_path=args.output_dir,
commit_message="End of training",
ignore_patterns=["step_*", "epoch_*", "*.bin", "*.pt"],
ignore_patterns=["*.bin"],
)
print(f"Params pushed to {repo_id}.")