mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
Merge pull request #139 from a-r-r-o-w/support-deepspeed
[feat] support DeepSpeed.
This commit is contained in:
@@ -1,6 +1,6 @@
|
||||
.PHONY: quality style
|
||||
|
||||
check_dirs := training tests
|
||||
check_dirs := finetrainers tests
|
||||
|
||||
quality:
|
||||
ruff check $(check_dirs)
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -20,4 +20,4 @@ same_network: true
|
||||
tpu_env: []
|
||||
tpu_use_cluster: false
|
||||
tpu_use_sudo: false
|
||||
use_cpu: false
|
||||
use_cpu: false
|
||||
@@ -14,4 +14,4 @@ same_network: true
|
||||
tpu_env: []
|
||||
tpu_use_cluster: false
|
||||
tpu_use_sudo: false
|
||||
use_cpu: false
|
||||
use_cpu: false
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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,5 +1,4 @@
|
||||
import inspect
|
||||
import logging
|
||||
|
||||
from accelerate.logging import get_logger
|
||||
import torch
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user