mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
address review feedback.
This commit is contained in:
@@ -48,7 +48,7 @@ from .utils.file_utils import string_to_filename
|
||||
from .utils.optimizer_utils import get_optimizer
|
||||
from .utils.memory_utils import get_memory_statistics, free_memory, make_contiguous
|
||||
from .utils.torch_utils import unwrap_model, align_device_and_dtype
|
||||
from .utils.checkpointing import sort_out_and_load_latest_ckpt_states, save_intermediate_ckpt_states
|
||||
from .utils.checkpointing import get_latest_ckpt_path_to_resume_from, get_intermediate_ckpt_path
|
||||
|
||||
|
||||
logger = get_logger("finetrainers")
|
||||
@@ -561,12 +561,13 @@ class Trainer:
|
||||
initial_global_step = 0
|
||||
|
||||
# Potentially load in the weights and states from a previous save
|
||||
initial_global_step, global_step, first_epoch = sort_out_and_load_latest_ckpt_states(
|
||||
accelerator=self.state.accelerator,
|
||||
resume_from_checkpoint_path, initial_global_step, global_step, first_epoch = get_latest_ckpt_path_to_resume_from(
|
||||
resume_from_checkpoint=self.args.resume_from_checkpoint,
|
||||
num_update_steps_per_epoch=self.state.num_update_steps_per_epoch,
|
||||
output_dir=self.args.output_dir
|
||||
)
|
||||
if resume_from_checkpoint_path:
|
||||
self.state.accelerator.load_state(resume_from_checkpoint_path)
|
||||
|
||||
progress_bar = tqdm(
|
||||
range(0, self.state.train_steps),
|
||||
@@ -699,9 +700,10 @@ class Trainer:
|
||||
# Checkpointing
|
||||
if accelerator.distributed_type == DistributedType.DEEPSPEED or accelerator.is_main_process:
|
||||
if global_step % self.args.checkpointing_steps == 0:
|
||||
save_intermediate_ckpt_states(
|
||||
accelerator=accelerator, checkpointing_limit=self.args.checkpointing_limit, step=global_step, output_dir=self.args.output_dir
|
||||
save_path = get_intermediate_ckpt_path(
|
||||
checkpointing_limit=self.args.checkpointing_limit, step=global_step, output_dir=self.args.output_dir
|
||||
)
|
||||
accelerator.save_state(save_path)
|
||||
|
||||
# Maybe run validation
|
||||
should_run_validation = (
|
||||
@@ -748,12 +750,11 @@ class Trainer:
|
||||
transformer_lora_layers=transformer_lora_layers,
|
||||
)
|
||||
|
||||
self.validate(step=global_step, final_validation=True)
|
||||
if self.args.push_to_hub:
|
||||
upload_folder(
|
||||
repo_id=self.state.repo_id, folder_path=self.args.output_dir, ignore_patterns=["checkpoint-*"]
|
||||
)
|
||||
|
||||
self.validate(step=global_step, final_validation=True)
|
||||
|
||||
del self.tokenizer, self.text_encoder, self.transformer, self.vae, self.scheduler
|
||||
free_memory()
|
||||
@@ -795,6 +796,7 @@ class Trainer:
|
||||
)
|
||||
else:
|
||||
# `torch_dtype` is manually set within `initialize_pipeline()`.
|
||||
self._delete_components()
|
||||
pipeline = self.model_config["initialize_pipeline"](
|
||||
model_id=self.args.pretrained_model_name_or_path,
|
||||
device=accelerator.device,
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import os
|
||||
from typing import Tuple
|
||||
from accelerate.logging import get_logger
|
||||
from ..constants import FINETRAINERS_LOG_LEVEL
|
||||
from ..utils.file_utils import find_files, delete_files
|
||||
@@ -6,13 +7,14 @@ from ..utils.file_utils import find_files, delete_files
|
||||
logger = get_logger("finetrainers")
|
||||
logger.setLevel(FINETRAINERS_LOG_LEVEL)
|
||||
|
||||
def sort_out_and_load_latest_ckpt_states(
|
||||
accelerator, resume_from_checkpoint, num_update_steps_per_epoch, output_dir
|
||||
):
|
||||
def get_latest_ckpt_path_to_resume_from(
|
||||
resume_from_checkpoint: str, num_update_steps_per_epoch: int, output_dir: str
|
||||
) -> Tuple[str, int, int, int]:
|
||||
if not resume_from_checkpoint:
|
||||
initial_global_step = 0
|
||||
global_step = 0
|
||||
first_epoch = 0
|
||||
resume_from_checkpoint_path = None
|
||||
else:
|
||||
if resume_from_checkpoint != "latest":
|
||||
path = os.path.basename(resume_from_checkpoint)
|
||||
@@ -29,18 +31,19 @@ def sort_out_and_load_latest_ckpt_states(
|
||||
)
|
||||
resume_from_checkpoint = None
|
||||
initial_global_step = 0
|
||||
resume_from_checkpoint_path = None
|
||||
else:
|
||||
logger.info(f"Resuming from checkpoint {path}")
|
||||
accelerator.load_state(os.path.join(output_dir, path))
|
||||
resume_from_checkpoint_path = os.path.join(output_dir, path)
|
||||
global_step = int(path.split("-")[1])
|
||||
|
||||
initial_global_step = global_step
|
||||
first_epoch = global_step // num_update_steps_per_epoch
|
||||
|
||||
return initial_global_step, global_step, first_epoch
|
||||
return resume_from_checkpoint_path, initial_global_step, global_step, first_epoch
|
||||
|
||||
|
||||
def save_intermediate_ckpt_states(accelerator, checkpointing_limit, step, output_dir):
|
||||
def get_intermediate_ckpt_path(checkpointing_limit: int, step: int, output_dir: str) -> str:
|
||||
# before saving state, check if this save would set us over the `checkpointing_limit`
|
||||
if checkpointing_limit is not None:
|
||||
checkpoints = find_files(output_dir, prefix="checkpoint")
|
||||
@@ -53,5 +56,5 @@ def save_intermediate_ckpt_states(accelerator, checkpointing_limit, step, output
|
||||
|
||||
logger.info(f"Checkpointing at step {step}")
|
||||
save_path = os.path.join(output_dir, f"checkpoint-{step}")
|
||||
accelerator.save_state(save_path)
|
||||
logger.info(f"Saved state to {save_path}")
|
||||
logger.info(f"Saving state to {save_path}")
|
||||
return save_path
|
||||
Reference in New Issue
Block a user