address review feedback.

This commit is contained in:
sayakpaul
2024-12-24 21:20:37 +05:30
parent 6329e9ed07
commit 875cb72f63
2 changed files with 20 additions and 15 deletions
+9 -7
View File
@@ -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,
+11 -8
View File
@@ -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