From 38413aa167caaed31fe0ed081576abfd55223726 Mon Sep 17 00:00:00 2001 From: Aryan Date: Mon, 6 Jan 2025 18:52:54 +0530 Subject: [PATCH] Allow images; Remove LLM generated prefixes; Allow JSON/JSONL; Fix bugs (#158) * update * update * fix * Update finetrainers/dataset.py * update * argument for enabling remove of common llm prefixes * update * make new generator for validation * update * update * update * update * update --- docs/dataset/README.md | 130 ++++++++++++------ finetrainers/args.py | 60 +++++++- finetrainers/constants.py | 26 ++++ finetrainers/dataset.py | 115 ++++++++++++++-- .../hunyuan_video/hunyuan_video_lora.py | 12 +- finetrainers/trainer.py | 31 +++-- 6 files changed, 301 insertions(+), 73 deletions(-) diff --git a/docs/dataset/README.md b/docs/dataset/README.md index 190340d..07a45c8 100644 --- a/docs/dataset/README.md +++ b/docs/dataset/README.md @@ -1,38 +1,8 @@ ## Dataset Format -### Prompt Dataset Requirements +Dataset loading format support is very limited at the moment. This will be improved in the future. For now, we support the following formats: -Create a `prompt.txt` file, which should contain prompts separated by lines. Please note that the prompts must be in English, and it is recommended to use the [prompt refinement script](https://github.com/THUDM/CogVideo/blob/main/inference/convert_demo.py) for better prompts. Alternatively, you can use [CogVideo-caption](https://huggingface.co/THUDM/cogvlm2-llama3-caption) for data annotation: - -``` -A black and white animated sequence featuring a rabbit, named Rabbity Ribfried, and an anthropomorphic goat in a musical, playful environment, showcasing their evolving interaction. -A black and white animated sequence on a ship’s deck features a bulldog character, named Bully Bulldoger, showcasing exaggerated facial expressions and body language... -... -``` - -### Video Dataset Requirements - -The framework supports resolutions and frame counts that meet the following conditions: - -- **Supported Resolutions (Width * Height)**: - - Any resolution as long as it is divisible by 32. For example, `720 * 480`, `1920 * 1020`, etc. - -- **Supported Frame Counts (Frames)**: - - Must be `4 * k` or `4 * k + 1` (example: 16, 32, 49, 81) - -It is recommended to place all videos in a single folder. - -Next, create a `videos.txt` file. The `videos.txt` file should contain the video file paths, separated by lines. Please note that the paths must be relative to the `--data_root` directory. The format is as follows: - -``` -videos/00000.mp4 -videos/00001.mp4 -... -``` - -For developers interested in more details, you can refer to the relevant `BucketSampler` code. - -### Dataset Structure +#### Two file format Your dataset structure should look like this. Running the `tree` command, you should see: @@ -41,21 +11,97 @@ dataset ├── prompt.txt ├── videos.txt ├── videos - ├── videos/00000.mp4 - ├── videos/00001.mp4 + ├── 00000.mp4 + ├── 00001.mp4 ├── ... ``` -### Using the Dataset - -When using this format, the `--caption_column` should be set to `prompt.txt`, and the `--video_column` should be set to `videos.txt`. If your data is stored in a CSV file, you can also specify `--dataset_file` as the path to the CSV file, with `--caption_column` and `--video_column` set to the actual column names in the CSV. Please refer to the [test_dataset](../tests/test_dataset.py) file for some simple examples. - -For instance, you can fine-tune using [this](https://huggingface.co/datasets/Wild-Heart/Disney-VideoGeneration-Dataset) Disney dataset. The download can be done via the 🤗 Hugging Face CLI: +For this format, you would specify arguments as follows: ``` -huggingface-cli download --repo-type dataset Wild-Heart/Disney-VideoGeneration-Dataset --local-dir video-dataset-disney +--data_root /path/to/dataset --caption_column prompt.txt --video_column videos.txt ``` -This dataset has been prepared in the expected format and can be used directly. However, directly using the video dataset may cause Out of Memory (OOM) issues on GPUs with smaller VRAM because it requires loading the [VAE](https://huggingface.co/THUDM/CogVideoX-5b/tree/main/vae) (which encodes videos into latent space) and the large [T5-XXL](https://huggingface.co/google/t5-v1_1-xxl/) text encoder. To reduce memory usage, you can use the `training/prepare_dataset.py` script to precompute latents and embeddings. +#### CSV format -Fill or modify the parameters in `prepare_dataset.sh` and execute it to get precomputed latents and embeddings (make sure to specify `--save_latents_and_embeddings` to save the precomputed artifacts). If preparing for image-to-video training, make sure to pass `--save_image_latents`, which encodes and saves image latents along with videos. When using these artifacts during training, ensure that you specify the `--load_tensors` flag, or else the videos will be used directly, requiring the text encoder and VAE to be loaded. The script also supports PyTorch DDP so that large datasets can be encoded in parallel across multiple GPUs (modify the `NUM_GPUS` parameter). +``` +dataset +├── dataset.csv +├── videos + ├── 00000.mp4 + ├── 00001.mp4 + ├── ... +``` + +The CSV can contain any number of columns, but due to limited support at the moment, we only make use of prompt and video columns. The CSV should look like this: + +``` +caption,video_file,other_column1,other_column2 +A black and white animated sequence featuring a rabbit, named Rabbity Ribfried, and an anthropomorphic goat in a musical, playful environment, showcasing their evolving interaction.,videos/00000.mp4,...,... +``` + +For this format, you would specify arguments as follows: + +``` +--data_root /path/to/dataset --caption_column caption --video_column video_file +``` + +### JSON format + +``` +dataset +├── dataset.json +├── videos + ├── 00000.mp4 + ├── 00001.mp4 + ├── ... +``` + +The JSON can contain any number of attributes, but due to limited support at the moment, we only make use of prompt and video columns. The JSON should look like this: + +```json +[ + { + "short_prompt": "A black and white animated sequence featuring a rabbit, named Rabbity Ribfried, and an anthropomorphic goat in a musical, playful environment, showcasing their evolving interaction.", + "filename": "videos/00000.mp4" + } +] +``` + +For this format, you would specify arguments as follows: + +``` +--data_root /path/to/dataset --caption_column short_prompt --video_column filename +``` + +### JSONL format + +``` +dataset +├── dataset.jsonl +├── videos + ├── 00000.mp4 + ├── 00001.mp4 + ├── ... +``` + +The JSONL can contain any number of attributes, but due to limited support at the moment, we only make use of prompt and video columns. The JSONL should look like this: + +```json +{"llm_prompt": "A black and white animated sequence featuring a rabbit, named Rabbity Ribfried, and an anthropomorphic goat in a musical, playful environment, showcasing their evolving interaction.", "filename": "videos/00000.mp4"} +{"llm_prompt": "A black and white animated sequence on a ship’s deck features a bulldog character, named Bully Bulldoger, showcasing exaggerated facial expressions and body language.", "filename": "videos/00001.mp4"} +... +``` + +For this format, you would specify arguments as follows: + +``` +--data_root /path/to/dataset --caption_column llm_prompt --video_column filename +``` + +> ![NOTE] +> Using images for finetuning is also supported. The dataset format remains the same as above. Find an example [here](https://huggingface.co/datasets/a-r-r-o-w/flux-retrostyle-dataset-mini). +> +> For example, to finetune with `512x512` resolution images, one must specify `--video_resolution_buckets 1x512x512` and point to the image files correctly. + +If you are using LLM-captioned videos, it is common to see many unwanted starting phrases like "In this video, ...", "This video features ...", etc. To remove a simple subset of these phrases, you can specify `--remove_common_llm_caption_prefixes` when starting training. diff --git a/finetrainers/args.py b/finetrainers/args.py index 2608ee3..d1c0715 100644 --- a/finetrainers/args.py +++ b/finetrainers/args.py @@ -41,6 +41,7 @@ class Args: caption_dropout_p: float = 0.00 caption_dropout_technique: str = "empty" precompute_conditions: bool = False + remove_common_llm_caption_prefixes: bool = False # Dataloader arguments dataloader_num_workers: int = 0 @@ -48,8 +49,8 @@ class Args: # Diffusion arguments flow_resolution_shifting: bool = False - flow_base_image_seq_len: int = 256 - flow_max_image_seq_len: int = 4096 + flow_base_seq_len: int = 256 + flow_max_seq_len: int = 4096 flow_base_shift: float = 0.5 flow_max_shift: float = 1.15 flow_shift: float = 1.0 @@ -143,11 +144,24 @@ class Args: "caption_dropout_p": self.caption_dropout_p, "caption_dropout_technique": self.caption_dropout_technique, "precompute_conditions": self.precompute_conditions, + "remove_common_llm_caption_prefixes": self.remove_common_llm_caption_prefixes, }, "dataloader_arguments": { "dataloader_num_workers": self.dataloader_num_workers, "pin_memory": self.pin_memory, }, + "diffusion_arguments": { + "flow_resolution_shifting": self.flow_resolution_shifting, + "flow_base_seq_len": self.flow_base_seq_len, + "flow_max_seq_len": self.flow_max_seq_len, + "flow_base_shift": self.flow_base_shift, + "flow_max_shift": self.flow_max_shift, + "flow_shift": self.flow_shift, + "flow_weighting_scheme": self.flow_weighting_scheme, + "flow_logit_mean": self.flow_logit_mean, + "flow_logit_std": self.flow_logit_std, + "flow_mode_scale": self.flow_mode_scale, + }, "training_arguments": { "training_type": self.training_type, "seed": self.seed, @@ -351,6 +365,11 @@ def _add_dataset_arguments(parser: argparse.ArgumentParser) -> None: action="store_true", help="Whether or not to precompute the conditionings for the model.", ) + parser.add_argument( + "--remove_common_llm_caption_prefixes", + action="store_true", + help="Whether or not to remove common LLM caption prefixes.", + ) def _add_dataloader_arguments(parser: argparse.ArgumentParser) -> None: @@ -373,6 +392,36 @@ def _add_diffusion_arguments(parser: argparse.ArgumentParser) -> None: action="store_true", help="Resolution-dependant shifting of timestep schedules.", ) + parser.add_argument( + "--flow_base_seq_len", + type=int, + default=256, + help="Base image/video sequence length for the diffusion model.", + ) + parser.add_argument( + "--flow_max_seq_len", + type=int, + default=4096, + help="Maximum image/video sequence length for the diffusion model.", + ) + parser.add_argument( + "--flow_base_shift", + type=float, + default=0.5, + help="Base shift as described in [Scaling Rectified Flow Transformers for High-Resolution Image Synthesis](https://arxiv.org/abs/2403.03206)", + ) + parser.add_argument( + "--flow_max_shift", + type=float, + default=1.15, + help="Maximum shift as described in [Scaling Rectified Flow Transformers for High-Resolution Image Synthesis](https://arxiv.org/abs/2403.03206)", + ) + parser.add_argument( + "--flow_shift", + type=float, + default=1.0, + help="Shift value to use for the flow matching timestep schedule.", + ) parser.add_argument( "--flow_weighting_scheme", type=str, @@ -722,6 +771,7 @@ def _map_to_args_type(args: Dict[str, Any]) -> Args: 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 + result_args.remove_common_llm_caption_prefixes = args.remove_common_llm_caption_prefixes # Dataloader arguments result_args.dataloader_num_workers = args.dataloader_num_workers @@ -729,6 +779,11 @@ def _map_to_args_type(args: Dict[str, Any]) -> Args: # Diffusion arguments result_args.flow_resolution_shifting = args.flow_resolution_shifting + result_args.flow_base_seq_len = args.flow_base_seq_len + result_args.flow_max_seq_len = args.flow_max_seq_len + result_args.flow_base_shift = args.flow_base_shift + result_args.flow_max_shift = args.flow_max_shift + result_args.flow_shift = args.flow_shift result_args.flow_weighting_scheme = args.flow_weighting_scheme result_args.flow_logit_mean = args.flow_logit_mean result_args.flow_logit_std = args.flow_logit_std @@ -743,6 +798,7 @@ def _map_to_args_type(args: Dict[str, Any]) -> Args: result_args.train_steps = args.train_steps result_args.rank = args.rank result_args.lora_alpha = args.lora_alpha + result_args.target_modules = args.target_modules result_args.gradient_accumulation_steps = args.gradient_accumulation_steps result_args.gradient_checkpointing = args.gradient_checkpointing result_args.checkpointing_steps = args.checkpointing_steps diff --git a/finetrainers/constants.py b/finetrainers/constants.py index 26d64ff..f6318f4 100644 --- a/finetrainers/constants.py +++ b/finetrainers/constants.py @@ -52,3 +52,29 @@ For more details, including weighting, merging and fusing LoRAs, check the [docu Please adhere to the license of the base model. """.strip() + +_COMMON_BEGINNING_PHRASES = ( + "This video", + "The video", + "This clip", + "The clip", + "The animation", + "This image", + "The image", + "This picture", + "The picture", +) +_COMMON_CONTINUATION_WORDS = ("shows", "depicts", "features", "captures", "highlights", "introduces", "presents") + +COMMON_LLM_START_PHRASES = ( + "In the video,", + "In this video,", + "In this video clip,", + "In the clip,", + "Caption:", + *( + f"{beginning} {continuation}" + for beginning in _COMMON_BEGINNING_PHRASES + for continuation in _COMMON_CONTINUATION_WORDS + ), +) diff --git a/finetrainers/dataset.py b/finetrainers/dataset.py index b1be6ef..6054e49 100644 --- a/finetrainers/dataset.py +++ b/finetrainers/dataset.py @@ -1,3 +1,4 @@ +import json import os import random from pathlib import Path @@ -7,6 +8,7 @@ import numpy as np import pandas as pd import torch import torchvision.transforms as TT +import torchvision.transforms.functional as TTF from accelerate.logging import get_logger from torch.utils.data import Dataset, Sampler from torchvision import transforms @@ -20,13 +22,25 @@ import decord # isort:skip decord.bridge.set_bridge("torch") -from .constants import PRECOMPUTED_CONDITIONS_DIR_NAME, PRECOMPUTED_DIR_NAME, PRECOMPUTED_LATENTS_DIR_NAME # noqa +from .constants import ( # noqa + COMMON_LLM_START_PHRASES, + PRECOMPUTED_CONDITIONS_DIR_NAME, + PRECOMPUTED_DIR_NAME, + PRECOMPUTED_LATENTS_DIR_NAME, +) logger = get_logger(__name__) -class VideoDataset(Dataset): +# TODO(aryan): This needs a refactor with separation of concerns. +# Images should be handled separately. Videos should be handled separately. +# Loading should be handled separately. +# Preprocessing (aspect ratio, resizing) should be handled separately. +# URL loading should be handled. +# Parquet format should be handled. +# Loading from ZIP should be handled. +class ImageOrVideoDataset(Dataset): def __init__( self, data_root: str, @@ -35,6 +49,7 @@ class VideoDataset(Dataset): resolution_buckets: List[Tuple[int, int, int]], dataset_file: Optional[str] = None, id_token: Optional[str] = None, + remove_llm_prefixes: bool = False, ) -> None: super().__init__() @@ -45,28 +60,52 @@ class VideoDataset(Dataset): self.id_token = f"{id_token.strip()} " if id_token else "" self.resolution_buckets = resolution_buckets - # Two methods of loading data are supported. + # Four methods of loading data are supported. # - Using a CSV: caption_column and video_column must be some column in the CSV. One could # make use of other columns too, such as a motion score or aesthetic score, by modifying the # logic in CSV processing. # - Using two files containing line-separate captions and relative paths to videos. + # - Using a JSON file containing a list of dictionaries, where each dictionary has a `caption_column` and `video_column` key. + # - Using a JSONL file containing a list of line-separated dictionaries, where each dictionary has a `caption_column` and `video_column` key. # For a more detailed explanation about preparing dataset format, checkout the README. if dataset_file is None: ( self.prompts, self.video_paths, ) = self._load_dataset_from_local_path() - else: + elif dataset_file.endswith(".csv"): ( self.prompts, self.video_paths, ) = self._load_dataset_from_csv() + elif dataset_file.endswith(".json"): + ( + self.prompts, + self.video_paths, + ) = self._load_dataset_from_json() + elif dataset_file.endswith(".jsonl"): + ( + self.prompts, + self.video_paths, + ) = self._load_dataset_from_jsonl() + else: + raise ValueError( + "Expected `--dataset_file` to be a path to a CSV file or a directory containing line-separated text prompts and video paths." + ) if len(self.video_paths) != len(self.prompts): raise ValueError( f"Expected length of prompts and videos to be the same but found {len(self.prompts)=} and {len(self.video_paths)=}. Please ensure that the number of caption prompts and videos match in your dataset." ) + # Clean LLM start phrases + if remove_llm_prefixes: + for i in range(len(self.prompts)): + self.prompts[i] = self.prompts[i].strip() + for phrase in COMMON_LLM_START_PHRASES: + if self.prompts[i].startswith(phrase): + self.prompts[i] = self.prompts[i].removeprefix(phrase).strip() + self.video_transforms = transforms.Compose( [ transforms.Lambda(self.scale_transform), @@ -95,7 +134,12 @@ class VideoDataset(Dataset): return index prompt = self.id_token + self.prompts[index] - video = self._preprocess_video(self.video_paths[index]) + + video_path: Path = self.video_paths[index] + if video_path.suffix.lower() in [".png", ".jpg", ".jpeg"]: + video = self._preprocess_image(video_path) + else: + video = self._preprocess_video(video_path) return { "prompt": prompt, @@ -148,12 +192,47 @@ class VideoDataset(Dataset): return prompts, video_paths + def _load_dataset_from_json(self) -> Tuple[List[str], List[str]]: + with open(self.dataset_file, "r", encoding="utf-8") as file: + data = json.load(file) + + prompts = [entry[self.caption_column] for entry in data] + video_paths = [self.data_root.joinpath(entry[self.video_column].strip()) for entry in data] + + if any(not path.is_file() for path in video_paths): + raise ValueError( + f"Expected `{self.video_column=}` to be a path to a file in `{self.data_root=}` containing line-separated paths to video data but found atleast one path that is not a valid file." + ) + + return prompts, video_paths + + def _load_dataset_from_jsonl(self) -> Tuple[List[str], List[str]]: + with open(self.dataset_file, "r", encoding="utf-8") as file: + data = [json.loads(line) for line in file] + + prompts = [entry[self.caption_column] for entry in data] + video_paths = [self.data_root.joinpath(entry[self.video_column].strip()) for entry in data] + + if any(not path.is_file() for path in video_paths): + raise ValueError( + f"Expected `{self.video_column=}` to be a path to a file in `{self.data_root=}` containing line-separated paths to video data but found atleast one path that is not a valid file." + ) + + return prompts, video_paths + + def _preprocess_image(self, path: Path) -> torch.Tensor: + # TODO(aryan): Support alpha channel in future by whitening background + image = TTF.Image.open(path.as_posix()).convert("RGB") + image = TTF.to_tensor(image) + image = image * 2.0 - 1.0 + image = image.unsqueeze(0).contiguous() # [C, H, W] -> [1, C, H, W] (1-frame video) + return image + def _preprocess_video(self, path: Path) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: r""" Loads a single video, or latent and prompt embedding, based on initialization parameters. - If returning a video, returns a [F, C, H, W] video tensor, and None for the prompt embedding. Here, - F, C, H and W are the frames, channels, height and width of the input video. + Returns a [F, C, H, W] video tensor. """ video_reader = decord.VideoReader(uri=path.as_posix()) video_num_frames = len(video_reader) @@ -166,12 +245,24 @@ class VideoDataset(Dataset): return frames -class VideoDatasetWithResizing(VideoDataset): +class ImageOrVideoDatasetWithResizing(ImageOrVideoDataset): def __init__(self, *args, **kwargs) -> None: super().__init__(*args, **kwargs) self.max_num_frames = max(self.resolution_buckets, key=lambda x: x[0])[0] + def _preprocess_image(self, path: Path) -> torch.Tensor: + # TODO(aryan): Support alpha channel in future by whitening background + image = TTF.Image.open(path.as_posix()).convert("RGB") + image = TTF.to_tensor(image) + + nearest_res = self._find_nearest_resolution(image.shape[1], image.shape[2]) + image = resize(image, nearest_res) + + image = image * 2.0 - 1.0 + image = image.unsqueeze(0).contiguous() + return image + def _preprocess_video(self, path: Path) -> torch.Tensor: video_reader = decord.VideoReader(uri=path.as_posix()) video_num_frames = len(video_reader) @@ -198,7 +289,7 @@ class VideoDatasetWithResizing(VideoDataset): return nearest_res[1], nearest_res[2] -class VideoDatasetWithResizeAndRectangleCrop(VideoDataset): +class ImageOrVideoDatasetWithResizeAndRectangleCrop(ImageOrVideoDataset): def __init__(self, video_reshape_mode: str = "center", *args, **kwargs) -> None: super().__init__(*args, **kwargs) @@ -292,8 +383,8 @@ class BucketSampler(Sampler): PyTorch Sampler that groups 3D data by height, width and frames. Args: - data_source (`VideoDataset`): - A PyTorch dataset object that is an instance of `VideoDataset`. + data_source (`ImageOrVideoDataset`): + A PyTorch dataset object that is an instance of `ImageOrVideoDataset`. batch_size (`int`, defaults to `8`): The batch size to use for training. shuffle (`bool`, defaults to `True`): @@ -306,7 +397,7 @@ class BucketSampler(Sampler): """ def __init__( - self, data_source: VideoDataset, batch_size: int = 8, shuffle: bool = True, drop_last: bool = False + self, data_source: ImageOrVideoDataset, batch_size: int = 8, shuffle: bool = True, drop_last: bool = False ) -> None: self.data_source = data_source self.batch_size = batch_size diff --git a/finetrainers/hunyuan_video/hunyuan_video_lora.py b/finetrainers/hunyuan_video/hunyuan_video_lora.py index dc0ecf7..9bfea53 100644 --- a/finetrainers/hunyuan_video/hunyuan_video_lora.py +++ b/finetrainers/hunyuan_video/hunyuan_video_lora.py @@ -58,6 +58,7 @@ def load_latent_models( def load_diffusion_models( model_id: str = "hunyuanvideo-community/HunyuanVideo", transformer_dtype: torch.dtype = torch.bfloat16, + shift: float = 1.0, revision: Optional[str] = None, cache_dir: Optional[str] = None, **kwargs, @@ -65,7 +66,7 @@ def load_diffusion_models( transformer = HunyuanVideoTransformer3DModel.from_pretrained( model_id, subfolder="transformer", torch_dtype=transformer_dtype, revision=revision, cache_dir=cache_dir ) - scheduler = FlowMatchEulerDiscreteScheduler() + scheduler = FlowMatchEulerDiscreteScheduler(shift=shift) return {"transformer": transformer, "scheduler": scheduler} @@ -195,10 +196,15 @@ def prepare_latents( h = torch.cat(encoded_slices) else: h = vae._encode(image_or_video) - return {"latents": h} + return {"latents": h} -def post_latent_preparation(latents: torch.Tensor, **kwargs) -> torch.Tensor: +def post_latent_preparation( + vae_config: Dict[str, Any], + latents: torch.Tensor, + **kwargs, +) -> torch.Tensor: + latents = latents * vae_config.scaling_factor return {"latents": latents} diff --git a/finetrainers/trainer.py b/finetrainers/trainer.py index 6ac59aa..8a9cc01 100644 --- a/finetrainers/trainer.py +++ b/finetrainers/trainer.py @@ -37,7 +37,7 @@ from .constants import ( PRECOMPUTED_DIR_NAME, PRECOMPUTED_LATENTS_DIR_NAME, ) -from .dataset import BucketSampler, PrecomputedDataset, VideoDatasetWithResizing +from .dataset import BucketSampler, ImageOrVideoDatasetWithResizing, PrecomputedDataset from .models import get_config_from_model_name from .state import State from .utils.checkpointing import get_intermediate_ckpt_path, get_latest_ckpt_path_to_resume_from @@ -103,13 +103,14 @@ class Trainer: # TODO(aryan): Make a background process for fetching logger.info("Initializing dataset and dataloader") - self.dataset = VideoDatasetWithResizing( + self.dataset = ImageOrVideoDatasetWithResizing( data_root=self.args.data_root, caption_column=self.args.caption_column, video_column=self.args.video_column, resolution_buckets=self.args.video_resolution_buckets, dataset_file=self.args.dataset_file, id_token=self.args.id_token, + remove_llm_prefixes=self.args.remove_common_llm_caption_prefixes, ) self.dataloader = torch.utils.data.DataLoader( self.dataset, @@ -127,6 +128,7 @@ class Trainer: "text_encoder_3_dtype": self.args.text_encoder_3_dtype, "transformer_dtype": self.args.transformer_dtype, "vae_dtype": self.args.vae_dtype, + "shift": self.args.flow_shift, "revision": self.args.revision, "cache_dir": self.args.cache_dir, } @@ -825,13 +827,13 @@ class Trainer: ) accelerator.save_state(save_path) - # Maybe run validation - should_run_validation = ( - self.args.validation_every_n_steps is not None - and global_step % self.args.validation_every_n_steps == 0 - ) - if should_run_validation: - self.validate(global_step) + # Maybe run validation + should_run_validation = ( + self.args.validation_every_n_steps is not None + and global_step % self.args.validation_every_n_steps == 0 + ) + if should_run_validation: + self.validate(global_step) logs["loss"] = loss.detach().item() logs["lr"] = self.lr_scheduler.get_last_lr()[0] @@ -871,8 +873,7 @@ class Trainer: 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() + self._delete_components() memory_statistics = get_memory_statistics() logger.info(f"Memory after training end: {json.dumps(memory_statistics, indent=4)}") @@ -955,7 +956,9 @@ class Trainer: width=width, num_frames=num_frames, num_videos_per_prompt=self.args.num_validation_videos_per_prompt, - generator=self.state.generator, + generator=torch.Generator(device=accelerator.device).manual_seed( + self.args.seed if self.args.seed is not None else 0 + ), # todo support passing `fps` for supported pipelines. ) @@ -971,7 +974,7 @@ class Trainer: main_process_only=False, ) - for key, value in list(artifacts.items()): + for index, (key, value) in enumerate(list(artifacts.items())): artifact_type = value["type"] artifact_value = value["value"] if artifact_type not in ["image", "video"] or artifact_value is None: @@ -979,7 +982,7 @@ class Trainer: extension = "png" if artifact_type == "image" else "mp4" filename = "validation-" if not final_validation else "final-" - filename += f"{step}-{accelerator.process_index}-{prompt_filename}.{extension}" + filename += f"{step}-{accelerator.process_index}-{index}-{prompt_filename}.{extension}" if accelerator.is_main_process and extension == "mp4": prompts_to_filenames[prompt] = filename filename = os.path.join(self.args.output_dir, filename)