mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
Merge pull request #106 from a-r-r-o-w/resume-mochi-1
feat: support checkpointing saving and loading
This commit is contained in:
@@ -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).
|
||||
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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}.")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user