Merge pull request #139 from a-r-r-o-w/support-deepspeed

[feat] support DeepSpeed.
This commit is contained in:
Sayak Paul
2024-12-27 15:21:31 +05:30
committed by GitHub
11 changed files with 268 additions and 89 deletions
+1 -1
View File
@@ -1,6 +1,6 @@
.PHONY: quality style
check_dirs := training tests
check_dirs := finetrainers tests
quality:
ruff check $(check_dirs)
+14 -7
View File
@@ -30,6 +30,8 @@ huggingface-cli download \
Then launch LoRA fine-tuning. For CogVideoX and Mochi, refer to [this](./training/README.md) and [this](./training/mochi-1/README.md).
Note: It is recommended to use Pytorch 2.5.1 or above for training. Previous versions can lead to completely black videos, OOM errors, or other issues and are not tested.
<details>
<summary> LTX Video </summary>
@@ -52,6 +54,8 @@ CAPTION_COLUMN="prompts.txt"
VIDEO_COLUMN="videos.txt"
OUTPUT_DIR="/path/to/output/directory/ltx-video/ltxv_disney"
ID_TOKEN="BW_STYLE"
# Model arguments
model_cmd="--model_name ltx_video \
--pretrained_model_name_or_path Lightricks/LTX-Video"
@@ -60,7 +64,7 @@ model_cmd="--model_name ltx_video \
dataset_cmd="--data_root $DATA_ROOT \
--video_column $VIDEO_COLUMN \
--caption_column $CAPTION_COLUMN \
--id_token BW_STYLE \
--id_token $ID_TOKEN \
--video_resolution_buckets 49x512x768 \
--caption_dropout_p 0.05"
@@ -99,7 +103,7 @@ optimizer_cmd="--optimizer adamw \
--max_grad_norm 1.0"
# Validation arguments
validation_cmd="--validation_prompts \"afkx A black and white animated scene unfolds with an anthropomorphic goat surrounded by musical notes and symbols, suggesting a playful environment. Mickey Mouse appears, leaning forward in curiosity as the goat remains still. The goat then engages with Mickey, who bends down to converse or react. The dynamics shift as Mickey grabs the goat, potentially in surprise or playfulness, amidst a minimalistic background. The scene captures the evolving relationship between the two characters in a whimsical, animated setting, emphasizing their interactions and emotions.@@@49x512x768:::A woman with long brown hair and light skin smiles at another woman with long blonde hair. The woman with brown hair wears a black jacket and has a small, barely noticeable mole on her right cheek. The camera angle is a close-up, focused on the woman with brown hair's face. The lighting is warm and natural, likely from the setting sun, casting a soft glow on the scene. The scene appears to be real-life footage@@@49x512x768\" \
validation_cmd="--validation_prompts \"$ID_TOKEN A black and white animated scene unfolds with an anthropomorphic goat surrounded by musical notes and symbols, suggesting a playful environment. Mickey Mouse appears, leaning forward in curiosity as the goat remains still. The goat then engages with Mickey, who bends down to converse or react. The dynamics shift as Mickey grabs the goat, potentially in surprise or playfulness, amidst a minimalistic background. The scene captures the evolving relationship between the two characters in a whimsical, animated setting, emphasizing their interactions and emotions.@@@49x512x768:::$ID_TOKEN A woman with long brown hair and light skin smiles at another woman with long blonde hair. The woman with brown hair wears a black jacket and has a small, barely noticeable mole on her right cheek. The camera angle is a close-up, focused on the woman with brown hair's face. The lighting is warm and natural, likely from the setting sun, casting a soft glow on the scene. The scene appears to be real-life footage@@@49x512x768\" \
--num_validation_videos 1 \
--validation_steps 100"
@@ -221,6 +225,8 @@ CAPTION_COLUMN="prompts.txt"
VIDEO_COLUMN="videos.txt"
OUTPUT_DIR="/path/to/models/hunyuan-video/hunyuan-video-loras/hunyuan-video_cakify_500_3e-5_constant_with_warmup"
ID_TOKEN="afkx"
# Model arguments
model_cmd="--model_name hunyuan_video \
--pretrained_model_name_or_path hunyuanvideo-community/HunyuanVideo"
@@ -229,8 +235,8 @@ model_cmd="--model_name hunyuan_video \
dataset_cmd="--data_root $DATA_ROOT \
--video_column $VIDEO_COLUMN \
--caption_column $CAPTION_COLUMN \
--id_token afkx \
--video_resolution_buckets 17x512x768 49x512x768 61x512x768 129x512x768 \
--id_token $ID_TOKEN \
--video_resolution_buckets 17x512x768 49x512x768 61x512x768 \
--caption_dropout_p 0.05"
# Dataloader arguments
@@ -268,7 +274,7 @@ optimizer_cmd="--optimizer adamw \
--max_grad_norm 1.0"
# Validation arguments
validation_cmd="--validation_prompts \"afkx A baker carefully cuts a green bell pepper cake on a white plate against a bright yellow background, followed by a strawberry cake with a similar slice of cake being cut before the interior of the bell pepper cake is revealed with the surrounding cake-to-object sequence.@@@49x512x768:::afkx A cake shaped like a Nutella container is carefully sliced, revealing a light interior, amidst a Nutella-themed setup, showcasing deliberate cutting and preserved details for an appetizing dessert presentation on a white base with accompanying jello and cutlery, highlighting culinary skills and creative cake designs.@@@49x512x768:::afkx A cake shaped like a Nutella container is carefully sliced, revealing a light interior, amidst a Nutella-themed setup, showcasing deliberate cutting and preserved details for an appetizing dessert presentation on a white base with accompanying jello and cutlery, highlighting culinary skills and creative cake designs.@@@61x512x768:::afkx A vibrant orange cake disguised as a Nike packaging box sits on a dark surface, meticulous in its detail and design, complete with a white swoosh and 'NIKE' logo. A person's hands, holding a knife, hover over the cake, ready to make a precise cut, amidst a simple and clean background.@@@61x512x768:::afkx A vibrant orange cake disguised as a Nike packaging box sits on a dark surface, meticulous in its detail and design, complete with a white swoosh and 'NIKE' logo. A person's hands, holding a knife, hover over the cake, ready to make a precise cut, amidst a simple and clean background.@@@97x512x768:::afkx A vibrant orange cake disguised as a Nike packaging box sits on a dark surface, meticulous in its detail and design, complete with a white swoosh and 'NIKE' logo. A person's hands, holding a knife, hover over the cake, ready to make a precise cut, amidst a simple and clean background.@@@129x512x768:::A person with gloved hands carefully cuts a cake shaped like a Skittles bottle, beginning with a precise incision at the lid, followed by careful sequential cuts around the neck, eventually detaching the lid from the body, revealing the chocolate interior of the cake while showcasing the layered design's detail.@@@61x512x768:::afkx A woman with long brown hair and light skin smiles at another woman with long blonde hair. The woman with brown hair wears a black jacket and has a small, barely noticeable mole on her right cheek. The camera angle is a close-up, focused on the woman with brown hair's face. The lighting is warm and natural, likely from the setting sun, casting a soft glow on the scene. The scene appears to be real-life footage@@@61x512x768\" \
validation_cmd="--validation_prompts \"$ID_TOKEN A baker carefully cuts a green bell pepper cake on a white plate against a bright yellow background, followed by a strawberry cake with a similar slice of cake being cut before the interior of the bell pepper cake is revealed with the surrounding cake-to-object sequence.@@@49x512x768:::$ID_TOKEN A cake shaped like a Nutella container is carefully sliced, revealing a light interior, amidst a Nutella-themed setup, showcasing deliberate cutting and preserved details for an appetizing dessert presentation on a white base with accompanying jello and cutlery, highlighting culinary skills and creative cake designs.@@@49x512x768:::$ID_TOKEN A cake shaped like a Nutella container is carefully sliced, revealing a light interior, amidst a Nutella-themed setup, showcasing deliberate cutting and preserved details for an appetizing dessert presentation on a white base with accompanying jello and cutlery, highlighting culinary skills and creative cake designs.@@@61x512x768:::$ID_TOKEN A vibrant orange cake disguised as a Nike packaging box sits on a dark surface, meticulous in its detail and design, complete with a white swoosh and 'NIKE' logo. A person's hands, holding a knife, hover over the cake, ready to make a precise cut, amidst a simple and clean background.@@@61x512x768:::$ID_TOKEN A vibrant orange cake disguised as a Nike packaging box sits on a dark surface, meticulous in its detail and design, complete with a white swoosh and 'NIKE' logo. A person's hands, holding a knife, hover over the cake, ready to make a precise cut, amidst a simple and clean background.@@@97x512x768:::$ID_TOKEN A vibrant orange cake disguised as a Nike packaging box sits on a dark surface, meticulous in its detail and design, complete with a white swoosh and 'NIKE' logo. A person's hands, holding a knife, hover over the cake, ready to make a precise cut, amidst a simple and clean background.@@@129x512x768:::$ID_TOKEN A person with gloved hands carefully cuts a cake shaped like a Skittles bottle, beginning with a precise incision at the lid, followed by careful sequential cuts around the neck, eventually detaching the lid from the body, revealing the chocolate interior of the cake while showcasing the layered design's detail.@@@61x512x768:::$ID_TOKEN A woman with long brown hair and light skin smiles at another woman with long blonde hair. The woman with brown hair wears a black jacket and has a small, barely noticeable mole on her right cheek. The camera angle is a close-up, focused on the woman with brown hair's face. The lighting is warm and natural, likely from the setting sun, casting a soft glow on the scene. The scene appears to be real-life footage@@@61x512x768\" \
--num_validation_videos 1 \
--validation_steps 100"
@@ -350,7 +356,7 @@ Training configuration: {
| after epoch 1 | 39.748 | 40.910 |
| after training end | 25.288 | 40.910 |
Note: requires about `59` GB of VRAM without precomputation.
Note: requires about `59` GB of VRAM when validation is performed.
LoRA with rank 128, batch size 1, gradient checkpointing, optimizer adamw, `49x512x768` resolutions, **with precomputation**:
@@ -377,7 +383,7 @@ Training configuration: {
| after validation end | 39.558 | 46.947 |
| after training end | 24.842 | 41.039 |
Note: requires about `47` GB of VRAM with precomputation. If validation is not performed, the memory usage is reduced to about `42` GB.
Note: requires about `47` GB of VRAM with validation. If validation is not performed, the memory usage is reduced to about `42` GB.
</details>
@@ -385,6 +391,7 @@ If you would like to use a custom dataset, refer to the dataset preparation guid
> [!NOTE]
> To lower memory requirements:
> - Use a DeepSpeed config to launch training (refer to [`accelerate_configs/deepspeed.yaml`](./accelerate_configs/deepspeed.yaml) as an example).
> - Pass `--precompute_conditions` when launching training.
> - Pass `--gradient_checkpointing` when launching training.
> - Do not perform validation/testing. This saves a significant amount of memory, which can be used to focus solely on training if you're on smaller VRAM GPUs.
+1 -1
View File
@@ -20,4 +20,4 @@ same_network: true
tpu_env: []
tpu_use_cluster: false
tpu_use_sudo: false
use_cpu: false
use_cpu: false
+1 -1
View File
@@ -14,4 +14,4 @@ same_network: true
tpu_env: []
tpu_use_cluster: false
tpu_use_sudo: false
use_cpu: false
use_cpu: false
+4 -1
View File
@@ -60,7 +60,9 @@ class Args:
# Training arguments
training_type: str = None
seed: int = 42
mixed_precision: str = None
mixed_precision: str = (
None # TODO: consider removing later https://github.com/a-r-r-o-w/finetrainers/pull/139#discussion_r1897438414
)
batch_size: int = 1
train_epochs: int = 1
train_steps: int = None
@@ -675,6 +677,7 @@ _DTYPE_MAP = {
"fp16": torch.float16,
"fp32": torch.float32,
}
_INVERSE_DTYPE_MAP = {v: k for k, v in _DTYPE_MAP.items()}
def _map_to_args_type(args: Dict[str, Any]) -> Args:
+1
View File
@@ -15,6 +15,7 @@ class State:
learning_rate: float = None
train_batch_size: int = None
generator: torch.Generator = None
num_update_steps_per_epoch: int = None
# Hub state
repo_id: str = None
+178 -76
View File
@@ -1,10 +1,8 @@
import inspect
import json
import logging
import math
import os
import random
import shutil
from datetime import timedelta
from typing import Any, Dict
from pathlib import Path
@@ -35,7 +33,7 @@ from huggingface_hub import create_repo, upload_folder
from peft import LoraConfig, get_peft_model_state_dict, set_peft_model_state_dict
from tqdm import tqdm
from .args import Args, validate_args
from .args import Args, validate_args, _INVERSE_DTYPE_MAP
from .constants import (
FINETRAINERS_LOG_LEVEL,
PRECOMPUTED_DIR_NAME,
@@ -46,10 +44,11 @@ from .dataset import BucketSampler, PrecomputedDataset, VideoDatasetWithResizing
from .models import get_config_from_model_name
from .state import State
from .utils.data_utils import should_perform_precomputation
from .utils.file_utils import find_files, delete_files, string_to_filename
from .utils.optimizer_utils import get_optimizer, gradient_norm
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.torch_utils import unwrap_model, align_device_and_dtype, expand_tensor_to_dims
from .utils.checkpointing import get_latest_ckpt_path_to_resume_from, get_intermediate_ckpt_path
logger = get_logger("finetrainers")
@@ -361,11 +360,7 @@ class Trainer:
# For mixed precision training we cast all non-trainable weights (vae, text_encoder and transformer) to half-precision
# as these weights are only used for inference, keeping weights in full precision is not required.
weight_dtype = torch.float32
if self.state.accelerator.mixed_precision == "fp16":
weight_dtype = torch.float16
elif self.state.accelerator.mixed_precision == "bf16":
weight_dtype = torch.bfloat16
weight_dtype = self._get_training_dtype(accelerator=self.state.accelerator)
if torch.backends.mps.is_available() and weight_dtype == torch.bfloat16:
# due to pytorch#99272, MPS does not yet support bfloat16.
@@ -375,6 +370,11 @@ class Trainer:
# TODO(aryan): handle torch dtype from accelerator vs model dtype; refactor
self.state.weight_dtype = weight_dtype
if self.args.mixed_precision != _INVERSE_DTYPE_MAP[weight_dtype]:
logger.warning(
f"`mixed_precision` was set to {_INVERSE_DTYPE_MAP[weight_dtype]} which is different from configured argument ({self.args.mixed_precision})."
)
self.args.mixed_precision = _INVERSE_DTYPE_MAP[weight_dtype]
self.transformer.to(dtype=weight_dtype)
self._move_components_to_device()
@@ -389,7 +389,13 @@ class Trainer:
)
self.transformer.add_adapter(transformer_lora_config)
# TODO: refactor
# Enable TF32 for faster training on Ampere GPUs: https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices
if self.args.allow_tf32 and torch.cuda.is_available():
torch.backends.cuda.matmul.allow_tf32 = True
self.register_saving_loading_hooks(transformer_lora_config)
def register_saving_loading_hooks(self, transformer_lora_config):
# create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format
def save_model_hook(models, weights, output_dir):
if self.state.accelerator.is_main_process:
@@ -415,13 +421,25 @@ class Trainer:
)
def load_model_hook(models, input_dir):
transformer_ = self.model_config["pipeline_cls"].from_pretrained(
self.args.pretrained_model_name_or_path, subfolder="transformer"
)
transformer_.add_adapter(transformer_lora_config)
if not self.state.accelerator.distributed_type == DistributedType.DEEPSPEED:
while len(models) > 0:
model = models.pop()
if isinstance(
unwrap_model(self.state.accelerator, model),
type(unwrap_model(self.state.accelerator, self.transformer)),
):
transformer_ = unwrap_model(self.state.accelerator, model)
else:
raise ValueError(
f"Unexpected save model: {unwrap_model(self.state.accelerator, model).__class__}"
)
else:
transformer_ = unwrap_model(self.state.accelerator, self.transformer).__class__.from_pretrained(
self.args.pretrained_model_name_or_path, subfolder="transformer"
)
transformer_.add_adapter(transformer_lora_config)
lora_state_dict = self.model_config["pipeline_cls"].lora_state_dict(input_dir)
transformer_state_dict = {
f'{k.replace("transformer.", "")}': v
for k, v in lora_state_dict.items()
@@ -447,10 +465,6 @@ class Trainer:
self.state.accelerator.register_save_state_pre_hook(save_model_hook)
self.state.accelerator.register_load_state_pre_hook(load_model_hook)
# Enable TF32 for faster training on Ampere GPUs: https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices
if self.args.allow_tf32 and torch.cuda.is_available():
torch.backends.cuda.matmul.allow_tf32 = True
def prepare_optimizer(self) -> None:
logger.info("Initializing optimizer and lr scheduler")
@@ -479,16 +493,20 @@ class Trainer:
params_to_optimize = [transformer_parameters_with_lr]
self.state.num_trainable_parameters = sum(p.numel() for p in transformer_lora_parameters)
# TODO(aryan): add deepspeed support
use_deepspeed_opt = (
self.state.accelerator.state.deepspeed_plugin is not None
and "optimizer" in self.state.accelerator.state.deepspeed_plugin.deepspeed_config
)
optimizer = get_optimizer(
params_to_optimize=params_to_optimize,
optimizer_name=self.args.optimizer,
learning_rate=self.args.lr,
learning_rate=self.state.learning_rate,
beta1=self.args.beta1,
beta2=self.args.beta2,
beta3=self.args.beta3,
epsilon=self.args.epsilon,
weight_decay=self.args.weight_decay,
use_deepspeed=use_deepspeed_opt,
)
num_update_steps_per_epoch = math.ceil(len(self.dataloader) / self.args.gradient_accumulation_steps)
@@ -496,14 +514,31 @@ class Trainer:
self.state.train_steps = self.state.train_epochs * num_update_steps_per_epoch
self.state.overwrote_max_train_steps = True
lr_scheduler = get_scheduler(
name=self.args.lr_scheduler,
optimizer=optimizer,
num_warmup_steps=self.args.lr_warmup_steps * self.state.accelerator.num_processes,
num_training_steps=self.state.train_steps * self.state.accelerator.num_processes,
num_cycles=self.args.lr_num_cycles,
power=self.args.lr_power,
use_deepspeed_lr_scheduler = (
self.state.accelerator.state.deepspeed_plugin is not None
and "scheduler" in self.state.accelerator.state.deepspeed_plugin.deepspeed_config
)
total_training_steps = self.state.train_steps * self.state.accelerator.num_processes
num_warmup_steps = self.args.lr_warmup_steps * self.state.accelerator.num_processes
if use_deepspeed_lr_scheduler:
from accelerate.utils import DummyScheduler
lr_scheduler = DummyScheduler(
name=self.args.lr_scheduler,
optimizer=optimizer,
total_num_steps=total_training_steps,
num_warmup_steps=num_warmup_steps,
)
else:
lr_scheduler = get_scheduler(
name=self.args.lr_scheduler,
optimizer=optimizer,
num_warmup_steps=num_warmup_steps,
num_training_steps=total_training_steps,
num_cycles=self.args.lr_num_cycles,
power=self.args.lr_power,
)
self.optimizer = optimizer
self.lr_scheduler = lr_scheduler
@@ -519,6 +554,7 @@ class Trainer:
self.state.train_steps = self.state.train_epochs * num_update_steps_per_epoch
# Afterwards we recalculate our number of training epochs
self.state.train_epochs = math.ceil(self.state.train_steps / num_update_steps_per_epoch)
self.state.num_update_steps_per_epoch = num_update_steps_per_epoch
def prepare_trackers(self) -> None:
logger.info("Initializing trackers")
@@ -547,10 +583,24 @@ class Trainer:
}
logger.info(f"Training configuration: {json.dumps(info, indent=4)}")
# TODO(aryan): handle resume from checkpoint
global_step = 0
first_epoch = 0
initial_global_step = 0
# Potentially load in the weights and states from a previous save
(
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),
initial=initial_global_step,
@@ -646,6 +696,7 @@ class Trainer:
device=accelerator.device,
dtype=weight_dtype,
)
sigmas = expand_tensor_to_dims(sigmas, ndim=latent_conditions["latents"].ndim)
noisy_latents = (1.0 - sigmas) * latent_conditions["latents"] + sigmas * noise
latent_conditions.update({"noisy_latents": noisy_latents})
@@ -666,8 +717,19 @@ class Trainer:
loss = loss.mean()
accelerator.backward(loss)
if accelerator.sync_gradients and accelerator.distributed_type != DistributedType.DEEPSPEED:
grad_norm = accelerator.clip_grad_norm_(self.transformer.parameters(), self.args.max_grad_norm)
if accelerator.sync_gradients:
if accelerator.distributed_type == DistributedType.DEEPSPEED:
grad_norm = self.transformer.get_global_grad_norm()
# In some cases the grad norm may not return a float
if torch.is_tensor(grad_norm):
grad_norm = grad_norm.item()
else:
grad_norm = accelerator.clip_grad_norm_(
self.transformer.parameters(), self.args.max_grad_norm
)
if torch.is_tensor(grad_norm):
grad_norm = grad_norm.item()
logs["grad_norm"] = grad_norm
self.optimizer.step()
@@ -682,20 +744,12 @@ class Trainer:
# Checkpointing
if accelerator.distributed_type == DistributedType.DEEPSPEED or accelerator.is_main_process:
if global_step % self.args.checkpointing_steps == 0:
# before saving state, check if this save would set us over the `checkpointing_limit`
if self.args.checkpointing_limit is not None:
checkpoints = find_files(self.args.output_dir, prefix="checkpoint")
# before we save the new checkpoint, we need to have at_most `checkpoints_total_limit - 1` checkpoints
if len(checkpoints) >= self.args.checkpointing_limit:
num_to_remove = len(checkpoints) - self.args.checkpointing_limit + 1
checkpoints_to_remove = checkpoints[0:num_to_remove]
delete_files(checkpoints_to_remove)
logger.info(f"Checkpointing at step {global_step}")
save_path = os.path.join(self.args.output_dir, f"checkpoint-{global_step}")
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)
logger.info(f"Saved state to {save_path}")
# Maybe run validation
should_run_validation = (
@@ -705,7 +759,8 @@ class Trainer:
if should_run_validation:
self.validate(global_step)
logs = {"loss": loss.detach().item(), "lr": self.lr_scheduler.get_last_lr()[0]}
logs["loss"] = loss.detach().item()
logs["lr"] = self.lr_scheduler.get_last_lr()[0]
progress_bar.set_postfix(logs)
accelerator.log(logs, step=global_step)
@@ -725,15 +780,8 @@ class Trainer:
accelerator.wait_for_everyone()
if accelerator.is_main_process:
# TODO: consider factoring this out when supporting other types of training algos.
self.transformer = unwrap_model(accelerator, self.transformer)
dtype = (
torch.float16
if self.args.mixed_precision == "fp16"
else torch.bfloat16
if self.args.mixed_precision == "bf16"
else torch.float32
)
self.transformer = self.transformer.to(dtype)
transformer_lora_layers = get_peft_model_state_dict(self.transformer)
self.model_config["pipeline_cls"].save_lora_weights(
@@ -741,6 +789,14 @@ class Trainer:
transformer_lora_layers=transformer_lora_layers,
)
self.validate(step=global_step, final_validation=True)
if accelerator.is_main_process:
if self.args.push_to_hub:
upload_folder(
repo_id=self.state.repo_id, folder_path=self.args.output_dir, ignore_patterns=["checkpoint-*"]
)
del self.tokenizer, self.text_encoder, self.transformer, self.vae, self.scheduler
free_memory()
memory_statistics = get_memory_statistics()
@@ -748,7 +804,7 @@ class Trainer:
accelerator.end_training()
def validate(self, step: int) -> None:
def validate(self, step: int, final_validation: bool = False) -> None:
logger.info("Starting validation")
accelerator = self.state.accelerator
@@ -763,21 +819,35 @@ class Trainer:
memory_statistics = get_memory_statistics()
logger.info(f"Memory before validation start: {json.dumps(memory_statistics, indent=4)}")
pipeline = self.model_config["initialize_pipeline"](
model_id=self.args.pretrained_model_name_or_path,
tokenizer=self.tokenizer,
text_encoder=self.text_encoder,
tokenizer_2=self.tokenizer_2,
text_encoder_2=self.text_encoder_2,
transformer=unwrap_model(accelerator, self.transformer),
vae=self.vae,
device=accelerator.device,
revision=self.args.revision,
cache_dir=self.args.cache_dir,
enable_slicing=self.args.enable_slicing,
enable_tiling=self.args.enable_tiling,
enable_model_cpu_offload=self.args.enable_model_cpu_offload,
)
if not final_validation:
pipeline = self.model_config["initialize_pipeline"](
model_id=self.args.pretrained_model_name_or_path,
tokenizer=self.tokenizer,
text_encoder=self.text_encoder,
tokenizer_2=self.tokenizer_2,
text_encoder_2=self.text_encoder_2,
transformer=unwrap_model(accelerator, self.transformer),
vae=self.vae,
device=accelerator.device,
revision=self.args.revision,
cache_dir=self.args.cache_dir,
enable_slicing=self.args.enable_slicing,
enable_tiling=self.args.enable_tiling,
enable_model_cpu_offload=self.args.enable_model_cpu_offload,
)
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,
revision=self.args.revision,
cache_dir=self.args.cache_dir,
enable_slicing=self.args.enable_slicing,
enable_tiling=self.args.enable_tiling,
enable_model_cpu_offload=self.args.enable_model_cpu_offload,
)
pipeline.load_lora_weights(self.args.output_dir)
all_processes_artifacts = []
for i in range(num_validation_samples):
@@ -811,6 +881,7 @@ class Trainer:
num_frames=num_frames,
num_videos_per_prompt=self.args.num_validation_videos_per_prompt,
generator=self.state.generator,
# todo support passing `fps` for supported pipelines.
)
prompt_filename = string_to_filename(prompt)[:25]
@@ -841,6 +912,7 @@ class Trainer:
artifact_value = wandb.Image(filename)
elif artifact_type == "video":
logger.debug(f"Saving video to {filename}")
# TODO: this should be configurable here as well as in validation runs where we call the pipeline that has `fps`.
export_to_video(artifact_value, filename, fps=15)
artifact_value = wandb.Video(filename, caption=prompt)
@@ -849,9 +921,18 @@ class Trainer:
all_artifacts = gather_object(all_processes_artifacts)
if accelerator.is_main_process:
tracker_key = "final" if final_validation else "validation"
for tracker in accelerator.trackers:
if tracker.name == "wandb":
tracker.log({"validation": all_artifacts}, step=step)
image_artifacts = [artifact for artifact in all_artifacts if isinstance(artifact, wandb.Image)]
video_artifacts = [artifact for artifact in all_artifacts if isinstance(artifact, wandb.Video)]
tracker.log(
{
tracker_key: {"images": image_artifacts, "videos": video_artifacts},
},
step=step,
)
# Remove all hooks that might have been added during pipeline initialization to the models
pipeline.remove_all_hooks()
@@ -864,11 +945,11 @@ class Trainer:
logger.info(f"Memory after validation end: {json.dumps(memory_statistics, indent=4)}")
torch.cuda.reset_peak_memory_stats(accelerator.device)
self.transformer.train()
if not final_validation:
self.transformer.train()
def evaluate(self) -> None:
logger.info("Starting evaluation")
# TODO: implement metrics for evaluation
raise NotImplementedError
def _init_distributed(self) -> None:
logging_dir = Path(self.args.output_dir, self.args.logging_dir)
@@ -922,7 +1003,7 @@ class Trainer:
if self.args.push_to_hub:
repo_id = self.args.hub_model_id or Path(self.args.output_dir).name
self.state.repo_id = create_repo(token=self.args.hub_token, name=repo_id).repo_id
self.state.repo_id = create_repo(token=self.args.hub_token, repo_id=repo_id).repo_id
def _move_components_to_device(self):
if self.text_encoder is not None:
@@ -937,3 +1018,24 @@ class Trainer:
self.unet = self.unet.to(self.state.accelerator.device)
if self.vae is not None:
self.vae = self.vae.to(self.state.accelerator.device)
def _get_training_dtype(self, accelerator) -> torch.dtype:
weight_dtype = torch.float32
if accelerator.state.deepspeed_plugin:
# DeepSpeed is handling precision, use what's in the DeepSpeed config
if (
"fp16" in accelerator.state.deepspeed_plugin.deepspeed_config
and accelerator.state.deepspeed_plugin.deepspeed_config["fp16"]["enabled"]
):
weight_dtype = torch.float16
if (
"bf16" in accelerator.state.deepspeed_plugin.deepspeed_config
and accelerator.state.deepspeed_plugin.deepspeed_config["bf16"]["enabled"]
):
weight_dtype = torch.bfloat16
else:
if self.state.accelerator.mixed_precision == "fp16":
weight_dtype = torch.float16
elif self.state.accelerator.mixed_precision == "bf16":
weight_dtype = torch.bfloat16
return weight_dtype
+61
View File
@@ -0,0 +1,61 @@
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
logger = get_logger("finetrainers")
logger.setLevel(FINETRAINERS_LOG_LEVEL)
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)
else:
# Get the most recent checkpoint
dirs = os.listdir(output_dir)
dirs = [d for d in dirs if d.startswith("checkpoint")]
dirs = sorted(dirs, key=lambda x: int(x.split("-")[1]))
path = dirs[-1] if len(dirs) > 0 else None
if path is None:
logger.info(f"Checkpoint '{resume_from_checkpoint}' does not exist. Starting a new training run.")
resume_from_checkpoint = None
initial_global_step = 0
global_step = 0
first_epoch = 0
resume_from_checkpoint_path = None
else:
logger.info(f"Resuming from checkpoint {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 resume_from_checkpoint_path, initial_global_step, global_step, first_epoch
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")
# before we save the new checkpoint, we need to have at_most `checkpoints_total_limit - 1` checkpoints
if len(checkpoints) >= checkpointing_limit:
num_to_remove = len(checkpoints) - checkpointing_limit + 1
checkpoints_to_remove = checkpoints[0:num_to_remove]
delete_files(checkpoints_to_remove)
logger.info(f"Checkpointing at step {step}")
save_path = os.path.join(output_dir, f"checkpoint-{step}")
logger.info(f"Saving state to {save_path}")
return save_path
-1
View File
@@ -1,5 +1,4 @@
import inspect
import logging
from accelerate.logging import get_logger
import torch
+6
View File
@@ -27,3 +27,9 @@ def align_device_and_dtype(
if dtype is not None:
x = {k: align_device_and_dtype(v, device, dtype) for k, v in x.items()}
return x
def expand_tensor_to_dims(tensor, ndim):
while len(tensor.shape) < ndim:
tensor = tensor.unsqueeze(-1)
return tensor
+1 -1
View File
@@ -33,7 +33,7 @@ def main():
trainer.prepare_for_training()
trainer.prepare_trackers()
trainer.train()
trainer.evaluate()
# trainer.evaluate()
except KeyboardInterrupt:
logger.info("Received keyboard interrupt. Exiting...")