This commit is contained in:
Aryan
2024-12-20 10:03:39 +01:00
parent 223add1a59
commit 66189c804b
9 changed files with 550 additions and 94 deletions
+51
View File
@@ -143,6 +143,57 @@ video = pipe("<my-awesome-prompt>").frames[0]
export_to_video(video, "output.mp4", fps=8)
```
### Memory Usage
LoRA with rank 128, batch size 1, gradient checkpointing, optimizer adamw, `49x512x768` resolution, **without precomputation**:
```
Memory before training start: {
"memory_allocated": 13.486,
"memory_reserved": 13.879,
"max_memory_allocated": 13.486,
"max_memory_reserved": 13.879
}
Training configuration: {
"trainable parameters": 117440512,
"total samples": 69,
"train epochs": 1,
"train steps": 10,
"batches per device": 1,
"total batches observed per epoch": 69,
"train batch size": 1,
"gradient accumulation steps": 1
}
Memory before validation start: {
"memory_allocated": 14.146,
"memory_reserved": 16.809,
"max_memory_allocated": 15.527,
"max_memory_reserved": 17.623
}
Memory after validation end: {
"memory_allocated": 14.146,
"memory_reserved": 14.627,
"max_memory_allocated": 15.527,
"max_memory_reserved": 17.623
}
Memory after epoch 1: {
"memory_allocated": 14.146,
"memory_reserved": 14.627,
"max_memory_allocated": 15.527,
"max_memory_reserved": 17.623
}
Memory after training end: {
"memory_allocated": 4.461,
"memory_reserved": 5.014,
"max_memory_allocated": 15.527,
"max_memory_reserved": 17.623
}
```
LoRA with rank 128, batch size 1, gradient checkpointing, optimizer adamw, `49x512x768` resolution, **with precomputation**:
TODO
</details>
<details>
+43
View File
@@ -1,6 +1,8 @@
import argparse
from typing import Any, Dict, List, Optional, Tuple
import torch
from .constants import DEFAULT_IMAGE_RESOLUTION_BUCKETS, DEFAULT_VIDEO_RESOLUTION_BUCKETS
@@ -20,6 +22,12 @@ class Args:
revision: Optional[str] = None
variant: Optional[str] = None
cache_dir: Optional[str] = None
text_encoder_dtype: torch.dtype = torch.bfloat16
text_encoder_2_dtype: torch.dtype = torch.bfloat16
text_encoder_3_dtype: torch.dtype = torch.bfloat16
transformer_dtype: torch.dtype = torch.bfloat16
unet_dtype: torch.dtype = torch.bfloat16
vae_dtype: torch.dtype = torch.bfloat16
# Dataset arguments
data_root: str = None
@@ -32,6 +40,7 @@ class Args:
video_reshape_mode: Optional[str] = None
caption_dropout_p: float = 0.00
caption_dropout_technique: str = "empty"
precompute_conditions: bool = False
# Dataloader arguments
dataloader_num_workers: int = 0
@@ -113,6 +122,12 @@ class Args:
"revision": self.revision,
"variant": self.variant,
"cache_dir": self.cache_dir,
"text_encoder_dtype": self.text_encoder_dtype,
"text_encoder_2_dtype": self.text_encoder_2_dtype,
"text_encoder_3_dtype": self.text_encoder_3_dtype,
"transformer_dtype": self.transformer_dtype,
"unet_dtype": self.unet_dtype,
"vae_dtype": self.vae_dtype,
},
"dataset_arguments": {
"data_root": self.data_root,
@@ -124,6 +139,8 @@ class Args:
"video_resolution_buckets": self.video_resolution_buckets,
"video_reshape_mode": self.video_reshape_mode,
"caption_dropout_p": self.caption_dropout_p,
"caption_dropout_technique": self.caption_dropout_technique,
"precompute_conditions": self.precompute_conditions,
},
"dataloader_arguments": {
"dataloader_num_workers": self.dataloader_num_workers,
@@ -234,6 +251,12 @@ def _add_model_arguments(parser: argparse.ArgumentParser) -> None:
default=None,
help="The directory where the downloaded models and datasets will be stored.",
)
parser.add_argument("--text_encoder_dtype", type=str, default="bf16", help="Data type for the text encoder.")
parser.add_argument("--text_encoder_2_dtype", type=str, default="bf16", help="Data type for the text encoder 2.")
parser.add_argument("--text_encoder_3_dtype", type=str, default="bf16", help="Data type for the text encoder 3.")
parser.add_argument("--transformer_dtype", type=str, default="bf16", help="Data type for the transformer model.")
parser.add_argument("--unet_dtype", type=str, default="bf16", help="Data type for the U-Net model.")
parser.add_argument("--vae_dtype", type=str, default="bf16", help="Data type for the VAE model.")
def _add_dataset_arguments(parser: argparse.ArgumentParser) -> None:
@@ -317,6 +340,11 @@ def _add_dataset_arguments(parser: argparse.ArgumentParser) -> None:
choices=["empty", "zero"],
help="Technique to use for caption dropout.",
)
parser.add_argument(
"--precompute_conditions",
action="store_true",
help="Whether or not to precompute the conditionings for the model.",
)
def _add_dataloader_arguments(parser: argparse.ArgumentParser) -> None:
@@ -645,6 +673,13 @@ def _add_miscellaneous_arguments(parser: argparse.ArgumentParser) -> None:
)
_DTYPE_MAP = {
"bf16": torch.bfloat16,
"fp16": torch.float16,
"fp32": torch.float32,
}
def _map_to_args_type(args: Dict[str, Any]) -> Args:
result_args = Args()
@@ -654,6 +689,12 @@ def _map_to_args_type(args: Dict[str, Any]) -> Args:
result_args.revision = args.revision
result_args.variant = args.variant
result_args.cache_dir = args.cache_dir
result_args.text_encoder_dtype = _DTYPE_MAP[args.text_encoder_dtype]
result_args.text_encoder_2_dtype = _DTYPE_MAP[args.text_encoder_2_dtype]
result_args.text_encoder_3_dtype = _DTYPE_MAP[args.text_encoder_3_dtype]
result_args.transformer_dtype = _DTYPE_MAP[args.transformer_dtype]
result_args.unet_dtype = _DTYPE_MAP[args.unet_dtype]
result_args.vae_dtype = _DTYPE_MAP[args.vae_dtype]
# Dataset arguments
if args.data_root is None and args.dataset_file is None:
@@ -668,6 +709,8 @@ def _map_to_args_type(args: Dict[str, Any]) -> Args:
result_args.video_resolution_buckets = args.video_resolution_buckets or DEFAULT_VIDEO_RESOLUTION_BUCKETS
result_args.video_reshape_mode = args.video_reshape_mode
result_args.caption_dropout_p = args.caption_dropout_p
result_args.caption_dropout_technique = args.caption_dropout_technique
result_args.precompute_conditions = args.precompute_conditions
# Dataloader arguments
result_args.dataloader_num_workers = args.dataloader_num_workers
+4
View File
@@ -19,6 +19,10 @@ for frames in DEFAULT_FRAME_BUCKETS:
FINETRAINERS_LOG_LEVEL = os.environ.get("FINETRAINERS_LOG_LEVEL", "INFO")
PRECOMPUTED_DIR_NAME = "precomputed"
PRECOMPUTED_CONDITIONS_DIR_NAME = "conditions"
PRECOMPUTED_LATENTS_DIR_NAME = "latents"
MODEL_DESCRIPTION = r"""
\# {model_id} {training_type} finetune
+30
View File
@@ -1,3 +1,4 @@
import os
import random
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple
@@ -19,6 +20,9 @@ import decord # isort:skip
decord.bridge.set_bridge("torch")
from .constants import PRECOMPUTED_DIR_NAME, PRECOMPUTED_CONDITIONS_DIR_NAME, PRECOMPUTED_LATENTS_DIR_NAME
logger = get_logger(__name__)
@@ -257,6 +261,32 @@ class VideoDatasetWithResizeAndRectangleCrop(VideoDataset):
return nearest_res[1], nearest_res[2]
class PrecomputedDataset(Dataset):
def __init__(self, data_root: str) -> None:
super().__init__()
self.data_root = Path(data_root)
self.latents_path = self.data_root / PRECOMPUTED_DIR_NAME / PRECOMPUTED_LATENTS_DIR_NAME
self.conditions_path = self.data_root / PRECOMPUTED_DIR_NAME / PRECOMPUTED_CONDITIONS_DIR_NAME
self.latent_conditions = sorted(os.listdir(self.latents_path))
self.other_conditions = sorted(os.listdir(self.conditions_path))
assert len(self.latent_conditions) == len(self.other_conditions), "Number of captions and videos do not match"
def __len__(self) -> int:
return len(self.latent_conditions)
def __getitem__(self, index: int) -> Dict[str, Any]:
conditions = {}
latent_path = self.latents_path / self.latent_conditions[index]
condition_path = self.conditions_path / self.other_conditions[index]
conditions["latent_conditions"] = torch.load(latent_path, map_location="cpu", weights_only=True)
conditions["other_conditions"] = torch.load(condition_path, map_location="cpu", weights_only=True)
return conditions
class BucketSampler(Sampler):
r"""
PyTorch Sampler that groups 3D data by height, width and frames.
+68 -19
View File
@@ -12,32 +12,45 @@ from PIL import Image
logger = get_logger("finetrainers") # pylint: disable=invalid-name
def load_components(
def load_condition_models(
model_id: str = "Lightricks/LTX-Video",
text_encoder_dtype: torch.dtype = torch.bfloat16,
transformer_dtype: torch.dtype = torch.bfloat16,
vae_dtype: torch.dtype = torch.bfloat16,
revision: Optional[str] = None,
cache_dir: Optional[str] = None,
**kwargs,
) -> Dict[str, nn.Module]:
tokenizer = T5Tokenizer.from_pretrained(model_id, subfolder="tokenizer", revision=revision, cache_dir=cache_dir)
text_encoder = T5EncoderModel.from_pretrained(
model_id, subfolder="text_encoder", torch_dtype=text_encoder_dtype, revision=revision, cache_dir=cache_dir
)
transformer = LTXVideoTransformer3DModel.from_pretrained(
model_id, subfolder="transformer", torch_dtype=transformer_dtype, revision=revision, cache_dir=cache_dir
)
return {"tokenizer": tokenizer, "text_encoder": text_encoder}
def load_latent_models(
model_id: str = "Lightricks/LTX-Video",
vae_dtype: torch.dtype = torch.bfloat16,
revision: Optional[str] = None,
cache_dir: Optional[str] = None,
**kwargs,
) -> Dict[str, nn.Module]:
vae = AutoencoderKLLTXVideo.from_pretrained(
model_id, subfolder="vae", torch_dtype=vae_dtype, revision=revision, cache_dir=cache_dir
)
return {"vae": vae}
def load_diffusion_models(
model_id: str = "Lightricks/LTX-Video",
transformer_dtype: torch.dtype = torch.bfloat16,
revision: Optional[str] = None,
cache_dir: Optional[str] = None,
**kwargs,
) -> Dict[str, nn.Module]:
transformer = LTXVideoTransformer3DModel.from_pretrained(
model_id, subfolder="transformer", torch_dtype=transformer_dtype, revision=revision, cache_dir=cache_dir
)
scheduler = FlowMatchEulerDiscreteScheduler()
return {
"tokenizer": tokenizer,
"text_encoder": text_encoder,
"transformer": transformer,
"vae": vae,
"scheduler": scheduler,
}
return {"transformer": transformer, "scheduler": scheduler}
def initialize_pipeline(
@@ -114,19 +127,52 @@ def prepare_latents(
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
generator: Optional[torch.Generator] = None,
precompute: bool = False,
) -> torch.Tensor:
device = device or vae.device
dtype = dtype or vae.dtype
if image_or_video.ndim == 4:
image_or_video = image_or_video.unsqueeze(2)
assert image_or_video.ndim == 5, f"Expected 5D tensor, got {image_or_video.ndim}D tensor"
image_or_video = image_or_video.to(device=device, dtype=dtype)
image_or_video = image_or_video.to(device=device, dtype=vae.dtype)
image_or_video = image_or_video.permute(0, 2, 1, 3, 4).contiguous() # [B, C, F, H, W] -> [B, F, C, H, W]
latents = vae.encode(image_or_video).latent_dist.sample(generator=generator)
_, _, num_frames, height, width = latents.shape
latents = _normalize_latents(latents, vae.latents_mean, vae.latents_std)
if not precompute:
latents = vae.encode(image_or_video).latent_dist.sample(generator=generator)
latents = latents.to(dtype=dtype)
_, _, num_frames, height, width = latents.shape
latents = _normalize_latents(latents, vae.latents_mean, vae.latents_std)
latents = _pack_latents(latents, patch_size, patch_size_t)
return {"latents": latents, "num_frames": num_frames, "height": height, "width": width}
else:
if vae.use_slicing and image_or_video.shape[0] > 1:
encoded_slices = [vae._encode(x_slice) for x_slice in image_or_video.split(1)]
h = torch.cat(encoded_slices)
else:
h = vae._encode(image_or_video)
_, _, num_frames, height, width = h.shape
# TODO(aryan): this is very very very stupid, but anything to make it work for now. refactor and design better later
return {
"latents": h,
"num_frames": num_frames,
"height": height,
"width": width,
"latents_mean": vae.latents_mean,
"latents_std": vae.latents_std,
}
def post_latent_preparation(
latents: torch.Tensor,
latents_mean: torch.Tensor,
latents_std: torch.Tensor,
num_frames: int,
height: int,
width: int,
patch_size: int = 1,
patch_size_t: int = 1,
) -> torch.Tensor:
latents = _normalize_latents(latents, latents_mean, latents_std)
latents = _pack_latents(latents, patch_size, patch_size_t)
return {"latents": latents, "num_frames": num_frames, "height": height, "width": width}
@@ -260,10 +306,13 @@ def _pack_latents(latents: torch.Tensor, patch_size: int = 1, patch_size_t: int
LTX_VIDEO_T2V_LORA_CONFIG = {
"pipeline_cls": LTXPipeline,
"load_components": load_components,
"load_condition_models": load_condition_models,
"load_latent_models": load_latent_models,
"load_diffusion_models": load_diffusion_models,
"initialize_pipeline": initialize_pipeline,
"prepare_conditions": prepare_conditions,
"prepare_latents": prepare_latents,
"post_latent_preparation": post_latent_preparation,
"collate_fn": collate_fn_t2v,
"forward_pass": forward_pass,
"validation": validation,
+297 -75
View File
@@ -29,20 +29,27 @@ from diffusers.training_utils import (
compute_density_for_timestep_sampling,
compute_loss_weighting_for_sd3,
)
from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution
from diffusers.utils import export_to_video, load_image, load_video
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 .constants import FINETRAINERS_LOG_LEVEL
from .dataset import BucketSampler, VideoDatasetWithResizing
from .constants import (
FINETRAINERS_LOG_LEVEL,
PRECOMPUTED_DIR_NAME,
PRECOMPUTED_CONDITIONS_DIR_NAME,
PRECOMPUTED_LATENTS_DIR_NAME,
)
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.memory_utils import get_memory_statistics, free_memory, make_contiguous
from .utils.torch_utils import unwrap_model
from .utils.torch_utils import unwrap_model, align_device_and_dtype
logger = get_logger("finetrainers")
@@ -73,6 +80,9 @@ class Trainer:
# Autoencoders
self.vae = None
# Scheduler
self.scheduler = None
self._init_distributed()
self._init_logging()
self._init_directories_and_repositories()
@@ -80,38 +90,8 @@ class Trainer:
self.state.model_name = self.args.model_name
self.model_config = get_config_from_model_name(self.args.model_name, self.args.training_type)
def prepare_models(self) -> None:
logger.info("Initializing models")
# TODO(aryan): refactor in future
load_components_kwargs = {
"text_encoder_dtype": torch.bfloat16,
"transformer_dtype": torch.bfloat16,
"vae_dtype": torch.bfloat16,
"revision": self.args.revision,
"cache_dir": self.args.cache_dir,
}
if self.args.pretrained_model_name_or_path is not None:
load_components_kwargs["model_id"] = self.args.pretrained_model_name_or_path
components = self._model_config_call(self.model_config["load_components"], load_components_kwargs)
self.tokenizer = components.get("tokenizer", None)
self.text_encoder = components.get("text_encoder", None)
self.tokenizer_2 = components.get("tokenizer_2", None)
self.text_encoder_2 = components.get("text_encoder_2", None)
self.transformer = components.get("transformer", None)
self.vae = components.get("vae", None)
self.scheduler = components.get("scheduler", None)
if self.vae is not None:
if self.args.enable_slicing:
self.vae.enable_slicing()
if self.args.enable_tiling:
self.vae.enable_tiling()
self.transformer_config = self.transformer.config if self.transformer is not None else None
def prepare_dataset(self) -> None:
# TODO(aryan): Make a background process for fetching
logger.info("Initializing dataset and dataloader")
self.dataset = VideoDatasetWithResizing(
@@ -131,16 +111,238 @@ class Trainer:
pin_memory=self.args.pin_memory,
)
def _get_load_components_kwargs(self) -> Dict[str, Any]:
load_component_kwargs = {
"text_encoder_dtype": self.args.text_encoder_dtype,
"text_encoder_2_dtype": self.args.text_encoder_2_dtype,
"text_encoder_3_dtype": self.args.text_encoder_3_dtype,
"transformer_dtype": self.args.transformer_dtype,
"unet_dtype": self.args.unet_dtype,
"vae_dtype": self.args.vae_dtype,
"revision": self.args.revision,
"cache_dir": self.args.cache_dir,
}
if self.args.pretrained_model_name_or_path is not None:
load_component_kwargs["model_id"] = self.args.pretrained_model_name_or_path
return load_component_kwargs
def _set_components(self, components: Dict[str, Any]) -> None:
self.tokenizer = components.get("tokenizer", self.tokenizer)
self.tokenizer_2 = components.get("tokenizer_2", self.tokenizer_2)
self.tokenizer_3 = components.get("tokenizer_3", self.tokenizer_3)
self.text_encoder = components.get("text_encoder", self.text_encoder)
self.text_encoder_2 = components.get("text_encoder_2", self.text_encoder_2)
self.text_encoder_3 = components.get("text_encoder_3", self.text_encoder_3)
self.transformer = components.get("transformer", self.transformer)
self.unet = components.get("unet", self.unet)
self.vae = components.get("vae", self.vae)
self.scheduler = components.get("scheduler", self.scheduler)
def _nuke_components(self) -> None:
self.tokenizer = None
self.tokenizer_2 = None
self.tokenizer_3 = None
self.text_encoder = None
self.text_encoder_2 = None
self.text_encoder_3 = None
self.transformer = None
self.unet = None
self.vae = None
self.scheduler = None
free_memory()
torch.cuda.synchronize(self.state.accelerator.device)
def prepare_models(self) -> None:
logger.info("Initializing models")
load_components_kwargs = self._get_load_components_kwargs()
condition_components, latent_components, diffusion_components = {}, {}, {}
if not self.args.precompute_conditions:
condition_components = self.model_config["load_condition_models"](**load_components_kwargs)
latent_components = self.model_config["load_latent_models"](**load_components_kwargs)
diffusion_components = self.model_config["load_diffusion_models"](**load_components_kwargs)
components = {}
components.update(condition_components)
components.update(latent_components)
components.update(diffusion_components)
self._set_components(components)
if self.vae is not None:
if self.args.enable_slicing:
self.vae.enable_slicing()
if self.args.enable_tiling:
self.vae.enable_tiling()
self.transformer_config = self.transformer.config if self.transformer is not None else None
def prepare_precomputations(self) -> None:
if not self.args.precompute_conditions:
return
logger.info("Initializing precomputations")
if self.args.batch_size != 1:
raise ValueError("Precomputation is only supported with batch size 1. This will be supported in future.")
def collate_fn(batch):
latent_conditions = [x["latent_conditions"] for x in batch]
other_conditions = [x["other_conditions"] for x in batch]
batched_latent_conditions = {}
batched_other_conditions = {}
for key in list(latent_conditions[0].keys()):
if torch.is_tensor(latent_conditions[0][key]):
batched_latent_conditions[key] = torch.cat([x[key] for x in latent_conditions], dim=0)
else:
# TODO(aryan): implement batch sampler for precomputed latents
batched_latent_conditions[key] = [x[key] for x in latent_conditions][0]
for key in list(other_conditions[0].keys()):
if torch.is_tensor(other_conditions[0][key]):
batched_other_conditions[key] = torch.cat([x[key] for x in other_conditions], dim=0)
else:
# TODO(aryan): implement batch sampler for precomputed latents
batched_other_conditions[key] = [x[key] for x in other_conditions][0]
return {"latent_conditions": batched_latent_conditions, "other_conditions": batched_other_conditions}
should_recompute = should_perform_precomputation(self.args.data_root)
if not should_recompute:
logger.info("Precomputed conditions and latents found. Loading precomputed data.")
self.dataloader = torch.utils.data.DataLoader(
PrecomputedDataset(self.args.data_root),
batch_size=self.args.batch_size,
shuffle=True,
collate_fn=collate_fn,
num_workers=self.args.dataloader_num_workers,
pin_memory=self.args.pin_memory,
)
return
logger.info("Precomputed conditions and latents not found. Running precomputation.")
# At this point, no models are loaded, so we need to load and precompute conditions and latents
condition_components = self.model_config["load_condition_models"](**self._get_load_components_kwargs())
self._set_components(condition_components)
self._move_components_to_device()
if self.args.caption_dropout_p > 0 and self.args.caption_dropout_technique == "empty":
logger.warning(
"Caption dropout is not supported with precomputation yet. This will be supported in the future."
)
conditions_dir = Path(self.args.data_root) / PRECOMPUTED_DIR_NAME / PRECOMPUTED_CONDITIONS_DIR_NAME
latents_dir = Path(self.args.data_root) / PRECOMPUTED_DIR_NAME / PRECOMPUTED_LATENTS_DIR_NAME
conditions_dir.mkdir(parents=True, exist_ok=True)
latents_dir.mkdir(parents=True, exist_ok=True)
# Precompute conditions
progress_bar = tqdm(
range(0, len(self.dataset)),
desc="Precomputing conditions",
disable=not self.state.accelerator.is_local_main_process,
)
index = 0
for i, data in enumerate(self.dataset):
if i % self.state.accelerator.num_processes != self.state.accelerator.process_index:
continue
logger.debug(
f"Precomputing conditions and latents for batch {i + 1}/{len(self.dataset)} on process {self.state.accelerator.process_index}"
)
other_conditions = self.model_config["prepare_conditions"](
tokenizer=self.tokenizer,
tokenizer_2=self.tokenizer_2,
tokenizer_3=self.tokenizer_3,
text_encoder=self.text_encoder,
text_encoder_2=self.text_encoder_2,
text_encoder_3=self.text_encoder_3,
prompt=data["prompt"],
device=self.state.accelerator.device,
dtype=self.state.weight_dtype,
)
filename = conditions_dir / f"conditions-{i}-{index}.pt"
torch.save(other_conditions, filename.as_posix())
index += 1
progress_bar.update(1)
self._nuke_components()
memory_statistics = get_memory_statistics()
logger.info(f"Memory after precomputing conditions: {json.dumps(memory_statistics, indent=4)}")
torch.cuda.reset_peak_memory_stats(self.state.accelerator.device)
# Precompute latents
latent_components = self.model_config["load_latent_models"](**self._get_load_components_kwargs())
self._set_components(latent_components)
self._move_components_to_device()
if self.vae is not None:
if self.args.enable_slicing:
self.vae.enable_slicing()
if self.args.enable_tiling:
self.vae.enable_tiling()
progress_bar = tqdm(
range(0, len(self.dataset)),
desc="Precomputing latents",
disable=not self.state.accelerator.is_local_main_process,
)
index = 0
for i, data in enumerate(self.dataset):
if i % self.state.accelerator.num_processes != self.state.accelerator.process_index:
continue
logger.debug(
f"Precomputing latents for batch {i + 1}/{len(self.dataset)} on process {self.state.accelerator.process_index}"
)
latent_conditions = self.model_config["prepare_latents"](
vae=self.vae,
image_or_video=data["video"].unsqueeze(0),
device=self.state.accelerator.device,
dtype=self.state.weight_dtype,
generator=self.state.generator,
precompute=True,
)
filename = latents_dir / f"latents-{i}-{index}.pt"
torch.save(latent_conditions, filename.as_posix())
index += 1
progress_bar.update(1)
self._nuke_components()
self.state.accelerator.wait_for_everyone()
logger.info("Precomputation complete")
memory_statistics = get_memory_statistics()
logger.info(f"Memory after precomputing latents: {json.dumps(memory_statistics, indent=4)}")
torch.cuda.reset_peak_memory_stats(self.state.accelerator.device)
# Update dataloader to use precomputed conditions and latents
self.dataloader = torch.utils.data.DataLoader(
PrecomputedDataset(self.args.data_root),
batch_size=self.args.batch_size,
shuffle=True,
collate_fn=collate_fn,
num_workers=self.args.dataloader_num_workers,
pin_memory=self.args.pin_memory,
)
def prepare_trainable_parameters(self) -> None:
logger.info("Initializing trainable parameters")
# TODO(aryan): refactor later. for now only lora is supported
self.text_encoder.requires_grad_(False)
self.transformer.requires_grad_(False)
self.vae.requires_grad_(False)
diffusion_components = self.model_config["load_diffusion_models"](**self._get_load_components_kwargs())
self._set_components(diffusion_components)
if self.text_encoder_2 is not None:
self.text_encoder_2.requires_grad_(False)
# TODO(aryan): refactor later. for now only lora is supported
components_to_disable_grads = [
self.text_encoder,
self.text_encoder_2,
self.text_encoder_3,
self.transformer,
self.vae,
]
for component in components_to_disable_grads:
if component is not None:
component.requires_grad_(False)
# 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.
@@ -158,12 +360,8 @@ class Trainer:
# TODO(aryan): handle torch dtype from accelerator vs model dtype; refactor
self.state.weight_dtype = weight_dtype
self.text_encoder.to(self.state.accelerator.device, dtype=weight_dtype)
self.transformer.to(self.state.accelerator.device, dtype=weight_dtype)
self.vae.to(self.state.accelerator.device, dtype=weight_dtype)
if self.text_encoder_2 is not None:
self.text_encoder_2.to(self.state.accelerator.device, dtype=weight_dtype)
self.transformer.to(dtype=weight_dtype)
self._move_components_to_device()
if self.args.gradient_checkpointing:
self.transformer.enable_gradient_checkpointing()
@@ -364,34 +562,46 @@ class Trainer:
logs = {}
with accelerator.accumulate(models_to_accumulate):
videos = batch["videos"]
prompts = batch["prompts"]
batch_size = len(prompts)
if not self.args.precompute_conditions:
videos = batch["videos"]
prompts = batch["prompts"]
batch_size = len(prompts)
if self.args.caption_dropout_technique == "empty":
if random.random() < self.args.caption_dropout_p:
prompts = [""] * batch_size
if self.args.caption_dropout_technique == "empty":
if random.random() < self.args.caption_dropout_p:
prompts = [""] * batch_size
latent_conditions = self.model_config["prepare_latents"](
vae=self.vae,
image_or_video=videos,
patch_size=self.transformer_config.patch_size,
patch_size_t=self.transformer_config.patch_size_t,
device=accelerator.device,
dtype=weight_dtype,
generator=generator,
)
other_conditions = self.model_config["prepare_conditions"](
tokenizer=self.tokenizer,
text_encoder=self.text_encoder,
tokenizer_2=self.tokenizer_2,
text_encoder_2=self.text_encoder_2,
prompt=prompts,
device=accelerator.device,
dtype=weight_dtype,
)
else:
latent_conditions = batch["latent_conditions"]
other_conditions = batch["other_conditions"]
latent_conditions["latents"] = DiagonalGaussianDistribution(
latent_conditions["latents"]
).sample(generator)
if "post_latent_preparation" in self.model_config.keys():
latent_conditions = self.model_config["post_latent_preparation"](**latent_conditions)
align_device_and_dtype(latent_conditions, accelerator.device, weight_dtype)
align_device_and_dtype(other_conditions, accelerator.device, weight_dtype)
batch_size = latent_conditions["latents"].shape[0]
latent_conditions = self.model_config["prepare_latents"](
vae=self.vae,
image_or_video=videos,
patch_size=self.transformer_config.patch_size,
patch_size_t=self.transformer_config.patch_size_t,
device=accelerator.device,
dtype=weight_dtype,
generator=generator,
)
latent_conditions = make_contiguous(latent_conditions)
other_conditions = self.model_config["prepare_conditions"](
tokenizer=self.tokenizer,
text_encoder=self.text_encoder,
tokenizer_2=self.tokenizer_2,
text_encoder_2=self.text_encoder_2,
prompt=prompts,
device=accelerator.device,
dtype=weight_dtype,
)
other_conditions = make_contiguous(other_conditions)
if self.args.caption_dropout_technique == "zero":
@@ -629,9 +839,12 @@ class Trainer:
tracker.log({"validation": all_artifacts}, step=step)
accelerator.wait_for_everyone()
free_memory()
memory_statistics = get_memory_statistics()
logger.info(f"Memory after validation end: {json.dumps(memory_statistics, indent=4)}")
torch.cuda.reset_peak_memory_stats(accelerator.device)
self.transformer.train()
def evaluate(self) -> None:
@@ -692,7 +905,16 @@ class Trainer:
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
def _model_config_call(self, fn, kwargs):
accepted_kwargs = inspect.signature(fn).parameters.keys()
kwargs = {k: v for k, v in kwargs.items() if k in accepted_kwargs}
return fn(**kwargs)
def _move_components_to_device(self):
if self.text_encoder is not None:
self.text_encoder = self.text_encoder.to(self.state.accelerator.device)
if self.text_encoder_2 is not None:
self.text_encoder_2 = self.text_encoder_2.to(self.state.accelerator.device)
if self.text_encoder_3 is not None:
self.text_encoder_3 = self.text_encoder_3.to(self.state.accelerator.device)
if self.transformer is not None:
self.transformer = self.transformer.to(self.state.accelerator.device)
if self.unet is not None:
self.unet = self.unet.to(self.state.accelerator.device)
if self.vae is not None:
self.vae = self.vae.to(self.state.accelerator.device)
+35
View File
@@ -0,0 +1,35 @@
from pathlib import Path
from typing import Union
from accelerate.logging import get_logger
from ..constants import PRECOMPUTED_DIR_NAME, PRECOMPUTED_CONDITIONS_DIR_NAME, PRECOMPUTED_LATENTS_DIR_NAME
logger = get_logger("finetrainers")
def should_perform_precomputation(data_root: Union[str, Path]) -> bool:
if isinstance(data_root, str):
data_root = Path(data_root)
conditions_dir = data_root / PRECOMPUTED_DIR_NAME / PRECOMPUTED_CONDITIONS_DIR_NAME
latents_dir = data_root / PRECOMPUTED_DIR_NAME / PRECOMPUTED_LATENTS_DIR_NAME
if conditions_dir.exists() and latents_dir.exists():
num_files_conditions = len(list(conditions_dir.glob("*.pt")))
num_files_latents = len(list(latents_dir.glob("*.pt")))
if num_files_conditions != num_files_latents:
logger.warning(
f"Number of precomputed conditions ({num_files_conditions}) does not match number of precomputed latents ({num_files_latents})."
f"Cleaning up precomputed directories and re-running precomputation."
)
# clean up precomputed directories
for file in conditions_dir.glob("*.pt"):
file.unlink()
for file in latents_dir.glob("*.pt"):
file.unlink()
return True
if num_files_conditions > 0:
logger.info(f"Found {num_files_conditions} precomputed conditions and latents.")
return False
logger.info("Precomputed data not found. Running precomputation.")
return True
+21
View File
@@ -1,3 +1,6 @@
from typing import Dict, Optional, Union
import torch
from accelerate import Accelerator
from diffusers.utils.torch_utils import is_compiled_module
@@ -6,3 +9,21 @@ def unwrap_model(accelerator: Accelerator, model):
model = accelerator.unwrap_model(model)
model = model._orig_mod if is_compiled_module(model) else model
return model
def align_device_and_dtype(
x: Union[torch.Tensor, Dict[str, torch.Tensor]],
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
):
if isinstance(x, torch.Tensor):
if device is not None:
x = x.to(device)
if dtype is not None:
x = x.to(dtype)
elif isinstance(x, dict):
if device is not None:
x = {k: align_device_and_dtype(v, device, dtype) for k, v in x.items()}
if dtype is not None:
x = {k: align_device_and_dtype(v, device, dtype) for k, v in x.items()}
return x
+1
View File
@@ -27,6 +27,7 @@ def main():
trainer.prepare_dataset()
trainer.prepare_models()
trainer.prepare_precomputations()
trainer.prepare_trainable_parameters()
trainer.prepare_optimizer()
trainer.prepare_for_training()