mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
Support CogVideoX T2V (#165)
* support cog t2v. * generator. * updates * style * fixes * fix padding frames. Co-authored-by: zRzRzRzRzRzRzR <Yuxuan.Zhang2104@student.xjtlu.edu.cn> * revert changes related to generator. * refactor a lot of things. * accept revision and cache_dir. * remove unused var * refactoring fixes. * refactor * update * update --------- Co-authored-by: zRzRzRzRzRzRzR <Yuxuan.Zhang2104@student.xjtlu.edu.cn> Co-authored-by: Aryan <aryan@huggingface.co>
This commit is contained in:
@@ -170,5 +170,6 @@ wandb/
|
||||
dump*
|
||||
outputs*
|
||||
*.slurm
|
||||
.vscode/
|
||||
|
||||
!requirements.txt
|
||||
|
||||
@@ -4,6 +4,7 @@ from typing import Any, Dict, List, Optional, Tuple
|
||||
import torch
|
||||
|
||||
from .constants import DEFAULT_IMAGE_RESOLUTION_BUCKETS, DEFAULT_VIDEO_RESOLUTION_BUCKETS
|
||||
from .models import SUPPORTED_MODEL_CONFIGS
|
||||
|
||||
|
||||
class Args:
|
||||
@@ -226,7 +227,11 @@ def validate_args(args: Args):
|
||||
|
||||
def _add_model_arguments(parser: argparse.ArgumentParser) -> None:
|
||||
parser.add_argument(
|
||||
"--model_name", type=str, required=True, choices=["hunyuan_video", "ltx_video"], help="Name of model to train."
|
||||
"--model_name",
|
||||
type=str,
|
||||
required=True,
|
||||
choices=list(SUPPORTED_MODEL_CONFIGS.keys()),
|
||||
help="Name of model to train.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--pretrained_model_name_or_path",
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
from .cogvideox_lora import COGVIDEOX_T2V_LORA_CONFIG
|
||||
@@ -0,0 +1,327 @@
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
import torch
|
||||
from diffusers import AutoencoderKLCogVideoX, CogVideoXDDIMScheduler, CogVideoXPipeline, CogVideoXTransformer3DModel
|
||||
from PIL import Image
|
||||
from transformers import T5EncoderModel, T5Tokenizer
|
||||
|
||||
from .utils import prepare_rotary_positional_embeddings
|
||||
|
||||
|
||||
def load_condition_models(
|
||||
model_id: str = "THUDM/CogVideoX-5b",
|
||||
text_encoder_dtype: torch.dtype = torch.bfloat16,
|
||||
revision: Optional[str] = None,
|
||||
cache_dir: Optional[str] = None,
|
||||
**kwargs,
|
||||
):
|
||||
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
|
||||
)
|
||||
return {"tokenizer": tokenizer, "text_encoder": text_encoder}
|
||||
|
||||
|
||||
def load_latent_models(
|
||||
model_id: str = "THUDM/CogVideoX-5b",
|
||||
vae_dtype: torch.dtype = torch.bfloat16,
|
||||
revision: Optional[str] = None,
|
||||
cache_dir: Optional[str] = None,
|
||||
**kwargs,
|
||||
):
|
||||
vae = AutoencoderKLCogVideoX.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 = "THUDM/CogVideoX-5b",
|
||||
transformer_dtype: torch.dtype = torch.bfloat16,
|
||||
revision: Optional[str] = None,
|
||||
cache_dir: Optional[str] = None,
|
||||
**kwargs,
|
||||
):
|
||||
transformer = CogVideoXTransformer3DModel.from_pretrained(
|
||||
model_id, subfolder="transformer", torch_dtype=transformer_dtype, revision=revision, cache_dir=cache_dir
|
||||
)
|
||||
scheduler = CogVideoXDDIMScheduler.from_pretrained(model_id, subfolder="scheduler")
|
||||
return {"transformer": transformer, "scheduler": scheduler}
|
||||
|
||||
|
||||
def initialize_pipeline(
|
||||
model_id: str = "THUDM/CogVideoX-5b",
|
||||
text_encoder_dtype: torch.dtype = torch.bfloat16,
|
||||
transformer_dtype: torch.dtype = torch.bfloat16,
|
||||
vae_dtype: torch.dtype = torch.bfloat16,
|
||||
tokenizer: Optional[T5Tokenizer] = None,
|
||||
text_encoder: Optional[T5EncoderModel] = None,
|
||||
transformer: Optional[CogVideoXTransformer3DModel] = None,
|
||||
vae: Optional[AutoencoderKLCogVideoX] = None,
|
||||
scheduler: Optional[CogVideoXDDIMScheduler] = None,
|
||||
device: Optional[torch.device] = None,
|
||||
revision: Optional[str] = None,
|
||||
cache_dir: Optional[str] = None,
|
||||
enable_slicing: bool = False,
|
||||
enable_tiling: bool = False,
|
||||
enable_model_cpu_offload: bool = False,
|
||||
**kwargs,
|
||||
) -> CogVideoXPipeline:
|
||||
component_name_pairs = [
|
||||
("tokenizer", tokenizer),
|
||||
("text_encoder", text_encoder),
|
||||
("transformer", transformer),
|
||||
("vae", vae),
|
||||
("scheduler", scheduler),
|
||||
]
|
||||
components = {}
|
||||
for name, component in component_name_pairs:
|
||||
if component is not None:
|
||||
components[name] = component
|
||||
|
||||
pipe = CogVideoXPipeline.from_pretrained(model_id, **components, revision=revision, cache_dir=cache_dir)
|
||||
pipe.text_encoder = pipe.text_encoder.to(dtype=text_encoder_dtype)
|
||||
pipe.transformer = pipe.transformer.to(dtype=transformer_dtype)
|
||||
pipe.vae = pipe.vae.to(dtype=vae_dtype)
|
||||
|
||||
if enable_slicing:
|
||||
pipe.vae.enable_slicing()
|
||||
if enable_tiling:
|
||||
pipe.vae.enable_tiling()
|
||||
|
||||
if enable_model_cpu_offload:
|
||||
pipe.enable_model_cpu_offload(device=device)
|
||||
else:
|
||||
pipe.to(device=device)
|
||||
|
||||
return pipe
|
||||
|
||||
|
||||
def prepare_conditions(
|
||||
tokenizer,
|
||||
text_encoder,
|
||||
prompt: Union[str, List[str]],
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
max_sequence_length: int = 226, # TODO: this should be configurable
|
||||
**kwargs,
|
||||
):
|
||||
device = device or text_encoder.device
|
||||
dtype = dtype or text_encoder.dtype
|
||||
return _get_t5_prompt_embeds(
|
||||
tokenizer=tokenizer,
|
||||
text_encoder=text_encoder,
|
||||
prompt=prompt,
|
||||
max_sequence_length=max_sequence_length,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
|
||||
def prepare_latents(
|
||||
vae: AutoencoderKLCogVideoX,
|
||||
image_or_video: torch.Tensor,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
generator: Optional[torch.Generator] = None,
|
||||
precompute: bool = False,
|
||||
**kwargs,
|
||||
) -> 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=vae.dtype)
|
||||
image_or_video = image_or_video.permute(0, 2, 1, 3, 4) # [B, C, F, H, W]
|
||||
if not precompute:
|
||||
latents = vae.encode(image_or_video).latent_dist.sample(generator=generator)
|
||||
if not vae.config.invert_scale_latents:
|
||||
latents = latents * vae.config.scaling_factor
|
||||
# For training Cog 1.5, we don't need to handle the scaling factor here.
|
||||
# The CogVideoX team forgot to multiply here, so we should not do it too. Invert scale latents
|
||||
# is probably only needed for image-to-video training.
|
||||
# TODO(aryan): investigate this
|
||||
# else:
|
||||
# latents = 1 / vae.config.scaling_factor * latents
|
||||
latents = latents.to(dtype=dtype)
|
||||
return {"latents": latents}
|
||||
else:
|
||||
# handle vae scaling in the `train()` method directly.
|
||||
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)
|
||||
return {"latents": h}
|
||||
|
||||
|
||||
def post_latent_preparation(
|
||||
vae_config: Dict[str, Any], latents: torch.Tensor, patch_size_t: Optional[int] = None, **kwargs
|
||||
) -> torch.Tensor:
|
||||
if not vae_config.invert_scale_latents:
|
||||
latents = latents * vae_config.scaling_factor
|
||||
# For training Cog 1.5, we don't need to handle the scaling factor here.
|
||||
# The CogVideoX team forgot to multiply here, so we should not do it too. Invert scale latents
|
||||
# is probably only needed for image-to-video training.
|
||||
# TODO(aryan): investigate this
|
||||
# else:
|
||||
# latents = 1 / vae_config.scaling_factor * latents
|
||||
latents = _pad_frames(latents, patch_size_t)
|
||||
latents = latents.permute(0, 2, 1, 3, 4) # [B, F, C, H, W]
|
||||
return {"latents": latents}
|
||||
|
||||
|
||||
def collate_fn_t2v(batch: List[List[Dict[str, torch.Tensor]]]) -> Dict[str, torch.Tensor]:
|
||||
return {
|
||||
"prompts": [x["prompt"] for x in batch[0]],
|
||||
"videos": torch.stack([x["video"] for x in batch[0]]),
|
||||
}
|
||||
|
||||
|
||||
def calculate_noisy_latents(
|
||||
scheduler: CogVideoXDDIMScheduler,
|
||||
noise: torch.Tensor,
|
||||
latents: torch.Tensor,
|
||||
timesteps: torch.LongTensor,
|
||||
) -> torch.Tensor:
|
||||
noisy_latents = scheduler.add_noise(latents, noise, timesteps)
|
||||
return noisy_latents
|
||||
|
||||
|
||||
def forward_pass(
|
||||
transformer: CogVideoXTransformer3DModel,
|
||||
scheduler: CogVideoXDDIMScheduler,
|
||||
prompt_embeds: torch.Tensor,
|
||||
latents: torch.Tensor,
|
||||
noisy_latents: torch.Tensor,
|
||||
timesteps: torch.LongTensor,
|
||||
ofs_emb: Optional[torch.Tensor] = None,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
# Just hardcode for now. In Diffusers, we will refactor such that RoPE would be handled within the model itself.
|
||||
VAE_SPATIAL_SCALE_FACTOR = 8
|
||||
transformer_config = transformer.module.config if hasattr(transformer, "module") else transformer.config
|
||||
batch_size, num_frames, num_channels, height, width = noisy_latents.shape
|
||||
rope_base_height = transformer_config.sample_height * VAE_SPATIAL_SCALE_FACTOR
|
||||
rope_base_width = transformer_config.sample_width * VAE_SPATIAL_SCALE_FACTOR
|
||||
|
||||
image_rotary_emb = (
|
||||
prepare_rotary_positional_embeddings(
|
||||
height=height * VAE_SPATIAL_SCALE_FACTOR,
|
||||
width=width * VAE_SPATIAL_SCALE_FACTOR,
|
||||
num_frames=num_frames,
|
||||
vae_scale_factor_spatial=VAE_SPATIAL_SCALE_FACTOR,
|
||||
patch_size=transformer_config.patch_size,
|
||||
patch_size_t=transformer_config.patch_size_t if hasattr(transformer_config, "patch_size_t") else None,
|
||||
attention_head_dim=transformer_config.attention_head_dim,
|
||||
device=transformer.device,
|
||||
base_height=rope_base_height,
|
||||
base_width=rope_base_width,
|
||||
)
|
||||
if transformer_config.use_rotary_positional_embeddings
|
||||
else None
|
||||
)
|
||||
ofs_emb = None if transformer_config.ofs_embed_dim is None else latents.new_full((batch_size,), fill_value=2.0)
|
||||
|
||||
velocity = transformer(
|
||||
hidden_states=noisy_latents,
|
||||
timestep=timesteps,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
ofs=ofs_emb,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
# For CogVideoX, the transformer predicts the velocity. The denoised output is calculated by applying the same
|
||||
# code paths as scheduler.get_velocity(), which can be confusing to understand.
|
||||
denoised_latents = scheduler.get_velocity(velocity, noisy_latents, timesteps)
|
||||
|
||||
return {"latents": denoised_latents}
|
||||
|
||||
|
||||
def validation(
|
||||
pipeline: CogVideoXPipeline,
|
||||
prompt: str,
|
||||
image: Optional[Image.Image] = None,
|
||||
video: Optional[List[Image.Image]] = None,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
num_frames: Optional[int] = None,
|
||||
num_videos_per_prompt: int = 1,
|
||||
generator: Optional[torch.Generator] = None,
|
||||
**kwargs,
|
||||
):
|
||||
generation_kwargs = {
|
||||
"prompt": prompt,
|
||||
"height": height,
|
||||
"width": width,
|
||||
"num_frames": num_frames,
|
||||
"num_videos_per_prompt": num_videos_per_prompt,
|
||||
"generator": generator,
|
||||
"return_dict": True,
|
||||
"output_type": "pil",
|
||||
}
|
||||
generation_kwargs = {k: v for k, v in generation_kwargs.items() if v is not None}
|
||||
output = pipeline(**generation_kwargs).frames[0]
|
||||
return [("video", output)]
|
||||
|
||||
|
||||
def _get_t5_prompt_embeds(
|
||||
tokenizer: T5Tokenizer,
|
||||
text_encoder: T5EncoderModel,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
max_sequence_length: int = 226,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
|
||||
text_inputs = tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=max_sequence_length,
|
||||
truncation=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
text_input_ids = text_inputs.input_ids
|
||||
|
||||
prompt_embeds = text_encoder(text_input_ids.to(device))[0]
|
||||
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
|
||||
|
||||
return {"prompt_embeds": prompt_embeds}
|
||||
|
||||
|
||||
def _pad_frames(latents: torch.Tensor, patch_size_t: int):
|
||||
if patch_size_t is None or patch_size_t == 1:
|
||||
return latents
|
||||
|
||||
# `latents` should be of the following format: [B, C, F, H, W].
|
||||
# For CogVideoX 1.5, the latent frames should be padded to make it divisible by patch_size_t
|
||||
latent_num_frames = latents.shape[2]
|
||||
additional_frames = patch_size_t - latent_num_frames % patch_size_t
|
||||
|
||||
if additional_frames > 0:
|
||||
last_frame = latents[:, :, -1:, :, :]
|
||||
padding_frames = last_frame.repeat(1, 1, additional_frames, 1, 1)
|
||||
latents = torch.cat([latents, padding_frames], dim=2)
|
||||
|
||||
return latents
|
||||
|
||||
|
||||
COGVIDEOX_T2V_LORA_CONFIG = {
|
||||
"pipeline_cls": CogVideoXPipeline,
|
||||
"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,
|
||||
"calculate_noisy_latents": calculate_noisy_latents,
|
||||
"forward_pass": forward_pass,
|
||||
"validation": validation,
|
||||
}
|
||||
@@ -0,0 +1,51 @@
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
from diffusers.models.embeddings import get_3d_rotary_pos_embed
|
||||
from diffusers.pipelines.cogvideo.pipeline_cogvideox import get_resize_crop_region_for_grid
|
||||
|
||||
|
||||
def prepare_rotary_positional_embeddings(
|
||||
height: int,
|
||||
width: int,
|
||||
num_frames: int,
|
||||
vae_scale_factor_spatial: int = 8,
|
||||
patch_size: int = 2,
|
||||
patch_size_t: int = None,
|
||||
attention_head_dim: int = 64,
|
||||
device: Optional[torch.device] = None,
|
||||
base_height: int = 480,
|
||||
base_width: int = 720,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
grid_height = height // (vae_scale_factor_spatial * patch_size)
|
||||
grid_width = width // (vae_scale_factor_spatial * patch_size)
|
||||
base_size_width = base_width // (vae_scale_factor_spatial * patch_size)
|
||||
base_size_height = base_height // (vae_scale_factor_spatial * patch_size)
|
||||
|
||||
if patch_size_t is None:
|
||||
# CogVideoX 1.0
|
||||
grid_crops_coords = get_resize_crop_region_for_grid(
|
||||
(grid_height, grid_width), base_size_width, base_size_height
|
||||
)
|
||||
freqs_cos, freqs_sin = get_3d_rotary_pos_embed(
|
||||
embed_dim=attention_head_dim,
|
||||
crops_coords=grid_crops_coords,
|
||||
grid_size=(grid_height, grid_width),
|
||||
temporal_size=num_frames,
|
||||
)
|
||||
else:
|
||||
# CogVideoX 1.5
|
||||
base_num_frames = (num_frames + patch_size_t - 1) // patch_size_t
|
||||
|
||||
freqs_cos, freqs_sin = get_3d_rotary_pos_embed(
|
||||
embed_dim=attention_head_dim,
|
||||
crops_coords=None,
|
||||
grid_size=(grid_height, grid_width),
|
||||
temporal_size=base_num_frames,
|
||||
grid_type="slice",
|
||||
max_size=(base_size_height, base_size_width),
|
||||
)
|
||||
|
||||
freqs_cos = freqs_cos.to(device=device)
|
||||
freqs_sin = freqs_sin.to(device=device)
|
||||
return freqs_cos, freqs_sin
|
||||
@@ -198,10 +198,7 @@ def prepare_latents(
|
||||
return {"latents": h}
|
||||
|
||||
|
||||
def post_latent_preparation(
|
||||
latents: torch.Tensor,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
def post_latent_preparation(latents: torch.Tensor, **kwargs) -> torch.Tensor:
|
||||
return {"latents": latents}
|
||||
|
||||
|
||||
@@ -221,6 +218,7 @@ def forward_pass(
|
||||
latents: torch.Tensor,
|
||||
noisy_latents: torch.Tensor,
|
||||
timesteps: torch.LongTensor,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
denoised_latents = transformer(
|
||||
hidden_states=noisy_latents,
|
||||
|
||||
@@ -173,6 +173,7 @@ def post_latent_preparation(
|
||||
width: int,
|
||||
patch_size: int = 1,
|
||||
patch_size_t: int = 1,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
latents = _normalize_latents(latents, latents_mean, latents_std)
|
||||
latents = _pack_latents(latents, patch_size, patch_size_t)
|
||||
@@ -196,6 +197,7 @@ def forward_pass(
|
||||
num_frames: int,
|
||||
height: int,
|
||||
width: int,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
# TODO(aryan): make configurable
|
||||
rope_interpolation_scale = [1 / 25, 32, 32]
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
from typing import Any, Dict
|
||||
|
||||
from .cogvideox import COGVIDEOX_T2V_LORA_CONFIG
|
||||
from .hunyuan_video import HUNYUAN_VIDEO_T2V_LORA_CONFIG
|
||||
from .ltx_video import LTX_VIDEO_T2V_LORA_CONFIG
|
||||
|
||||
@@ -11,6 +12,9 @@ SUPPORTED_MODEL_CONFIGS = {
|
||||
"ltx_video": {
|
||||
"lora": LTX_VIDEO_T2V_LORA_CONFIG,
|
||||
},
|
||||
"cogvideox": {
|
||||
"lora": COGVIDEOX_T2V_LORA_CONFIG,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
|
||||
+110
-48
@@ -3,7 +3,7 @@ import logging
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
from datetime import timedelta
|
||||
from datetime import datetime, timedelta
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict
|
||||
|
||||
@@ -11,7 +11,6 @@ import diffusers
|
||||
import torch
|
||||
import torch.backends
|
||||
import transformers
|
||||
import wandb
|
||||
from accelerate import Accelerator, DistributedType
|
||||
from accelerate.logging import get_logger
|
||||
from accelerate.utils import (
|
||||
@@ -21,18 +20,17 @@ from accelerate.utils import (
|
||||
gather_object,
|
||||
set_seed,
|
||||
)
|
||||
from diffusers.configuration_utils import FrozenDict
|
||||
from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution
|
||||
from diffusers.optimization import get_scheduler
|
||||
from diffusers.training_utils import (
|
||||
cast_training_params,
|
||||
compute_density_for_timestep_sampling,
|
||||
compute_loss_weighting_for_sd3,
|
||||
)
|
||||
from diffusers.training_utils import cast_training_params
|
||||
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
|
||||
|
||||
import wandb
|
||||
|
||||
from .args import _INVERSE_DTYPE_MAP, Args, validate_args
|
||||
from .constants import (
|
||||
FINETRAINERS_LOG_LEVEL,
|
||||
@@ -45,10 +43,18 @@ 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
|
||||
from .utils.data_utils import should_perform_precomputation
|
||||
from .utils.diffusion_utils import (
|
||||
get_scheduler_alphas,
|
||||
get_scheduler_sigmas,
|
||||
prepare_loss_weights,
|
||||
prepare_sigmas,
|
||||
prepare_target,
|
||||
)
|
||||
from .utils.file_utils import string_to_filename
|
||||
from .utils.memory_utils import free_memory, get_memory_statistics, make_contiguous
|
||||
from .utils.model_utils import resolve_vae_cls_from_ckpt_path
|
||||
from .utils.optimizer_utils import get_optimizer
|
||||
from .utils.torch_utils import align_device_and_dtype, expand_tensor_to_dims, unwrap_model
|
||||
from .utils.torch_utils import align_device_and_dtype, expand_tensor_dims, unwrap_model
|
||||
|
||||
|
||||
logger = get_logger("finetrainers")
|
||||
@@ -60,6 +66,7 @@ class Trainer:
|
||||
validate_args(args)
|
||||
|
||||
self.args = args
|
||||
self.args.seed = self.args.seed or datetime.now().year
|
||||
self.state = State()
|
||||
|
||||
# Tokenizers
|
||||
@@ -82,6 +89,9 @@ class Trainer:
|
||||
# Scheduler
|
||||
self.scheduler = None
|
||||
|
||||
self.transformer_config = None
|
||||
self.vae_config = None
|
||||
|
||||
self._init_distributed()
|
||||
self._init_logging()
|
||||
self._init_directories_and_repositories()
|
||||
@@ -125,6 +135,7 @@ class Trainer:
|
||||
return load_component_kwargs
|
||||
|
||||
def _set_components(self, components: Dict[str, Any]) -> None:
|
||||
# Set models
|
||||
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)
|
||||
@@ -136,6 +147,10 @@ class Trainer:
|
||||
self.vae = components.get("vae", self.vae)
|
||||
self.scheduler = components.get("scheduler", self.scheduler)
|
||||
|
||||
# Set configs
|
||||
self.transformer_config = self.transformer.config if self.transformer is not None else self.transformer_config
|
||||
self.vae_config = self.vae.config if self.vae is not None else self.vae_config
|
||||
|
||||
def _delete_components(self) -> None:
|
||||
self.tokenizer = None
|
||||
self.tokenizer_2 = None
|
||||
@@ -172,8 +187,6 @@ class Trainer:
|
||||
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
|
||||
@@ -242,19 +255,21 @@ class Trainer:
|
||||
conditions_dir.mkdir(parents=True, exist_ok=True)
|
||||
latents_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
accelerator = self.state.accelerator
|
||||
|
||||
# Precompute conditions
|
||||
progress_bar = tqdm(
|
||||
range(0, len(self.dataset)),
|
||||
range(0, (len(self.dataset) + accelerator.num_processes - 1) // accelerator.num_processes),
|
||||
desc="Precomputing conditions",
|
||||
disable=not self.state.accelerator.is_local_main_process,
|
||||
disable=not 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:
|
||||
if i % accelerator.num_processes != 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}"
|
||||
f"Precomputing conditions for batch {i + 1}/{len(self.dataset)} on process {accelerator.process_index}"
|
||||
)
|
||||
|
||||
text_conditions = self.model_config["prepare_conditions"](
|
||||
@@ -265,10 +280,10 @@ class Trainer:
|
||||
text_encoder_2=self.text_encoder_2,
|
||||
text_encoder_3=self.text_encoder_3,
|
||||
prompt=data["prompt"],
|
||||
device=self.state.accelerator.device,
|
||||
device=accelerator.device,
|
||||
dtype=self.state.weight_dtype,
|
||||
)
|
||||
filename = conditions_dir / f"conditions-{i}-{index}.pt"
|
||||
filename = conditions_dir / f"conditions-{accelerator.process_index}-{index}.pt"
|
||||
torch.save(text_conditions, filename.as_posix())
|
||||
index += 1
|
||||
progress_bar.update(1)
|
||||
@@ -276,7 +291,7 @@ class Trainer:
|
||||
|
||||
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)
|
||||
torch.cuda.reset_peak_memory_stats(accelerator.device)
|
||||
|
||||
# Precompute latents
|
||||
latent_components = self.model_config["load_latent_models"](**self._get_load_components_kwargs())
|
||||
@@ -296,39 +311,39 @@ class Trainer:
|
||||
self.vae.enable_tiling()
|
||||
|
||||
progress_bar = tqdm(
|
||||
range(0, len(self.dataset)),
|
||||
range(0, (len(self.dataset) + accelerator.num_processes - 1) // accelerator.num_processes),
|
||||
desc="Precomputing latents",
|
||||
disable=not self.state.accelerator.is_local_main_process,
|
||||
disable=not 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:
|
||||
if i % accelerator.num_processes != accelerator.process_index:
|
||||
continue
|
||||
|
||||
logger.debug(
|
||||
f"Precomputing latents for batch {i + 1}/{len(self.dataset)} on process {self.state.accelerator.process_index}"
|
||||
f"Precomputing latents for batch {i + 1}/{len(self.dataset)} on process {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,
|
||||
device=accelerator.device,
|
||||
dtype=self.state.weight_dtype,
|
||||
generator=self.state.generator,
|
||||
precompute=True,
|
||||
)
|
||||
filename = latents_dir / f"latents-{self.state.accelerator.process_index}-{index}.pt"
|
||||
filename = latents_dir / f"latents-{accelerator.process_index}-{index}.pt"
|
||||
torch.save(latent_conditions, filename.as_posix())
|
||||
index += 1
|
||||
progress_bar.update(1)
|
||||
self._delete_components()
|
||||
|
||||
self.state.accelerator.wait_for_everyone()
|
||||
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)
|
||||
torch.cuda.reset_peak_memory_stats(accelerator.device)
|
||||
|
||||
# Update dataloader to use precomputed conditions and latents
|
||||
self.dataloader = torch.utils.data.DataLoader(
|
||||
@@ -569,6 +584,20 @@ class Trainer:
|
||||
memory_statistics = get_memory_statistics()
|
||||
logger.info(f"Memory before training start: {json.dumps(memory_statistics, indent=4)}")
|
||||
|
||||
if self.vae_config is None:
|
||||
# If we've precomputed conditions and latents already, and are now re-using it, we will never load
|
||||
# the VAE so self.vae_config will not be set. So, we need to load it here.
|
||||
vae_cls_name = resolve_vae_cls_from_ckpt_path(
|
||||
self.args.pretrained_model_name_or_path, revision=self.args.revision, cache_dir=self.args.cache_dir
|
||||
)
|
||||
vae_config = vae_cls_name.load_config(
|
||||
self.args.pretrained_model_name_or_path,
|
||||
subfolder="vae",
|
||||
revision=self.args.revision,
|
||||
cache_dir=self.args.cache_dir,
|
||||
)
|
||||
self.vae_config = FrozenDict(**vae_config)
|
||||
|
||||
self.state.train_batch_size = (
|
||||
self.args.batch_size * self.state.accelerator.num_processes * self.args.gradient_accumulation_steps
|
||||
)
|
||||
@@ -611,12 +640,14 @@ class Trainer:
|
||||
|
||||
accelerator = self.state.accelerator
|
||||
weight_dtype = self.state.weight_dtype
|
||||
scheduler_sigmas = self.scheduler.sigmas.clone().to(device=accelerator.device, dtype=weight_dtype)
|
||||
generator = torch.Generator(device=accelerator.device)
|
||||
if self.args.seed is not None:
|
||||
generator = generator.manual_seed(self.args.seed)
|
||||
self.state.generator = generator
|
||||
|
||||
scheduler_sigmas = get_scheduler_sigmas(self.scheduler).to(device=accelerator.device, dtype=torch.float32)
|
||||
scheduler_alphas = get_scheduler_alphas(self.scheduler).to(device=accelerator.device, dtype=torch.float32)
|
||||
|
||||
for epoch in range(first_epoch, self.state.train_epochs):
|
||||
logger.debug(f"Starting epoch ({epoch + 1}/{self.state.train_epochs})")
|
||||
|
||||
@@ -644,7 +675,7 @@ class Trainer:
|
||||
patch_size_t=self.transformer_config.patch_size_t,
|
||||
device=accelerator.device,
|
||||
dtype=weight_dtype,
|
||||
generator=generator,
|
||||
generator=self.state.generator,
|
||||
)
|
||||
text_conditions = self.model_config["prepare_conditions"](
|
||||
tokenizer=self.tokenizer,
|
||||
@@ -660,9 +691,16 @@ class Trainer:
|
||||
text_conditions = batch["text_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)
|
||||
).sample(self.state.generator)
|
||||
|
||||
# This method should only be called for precomputed latents.
|
||||
# TODO(aryan): rename this in separate PR
|
||||
latent_conditions = self.model_config["post_latent_preparation"](
|
||||
vae_config=self.vae_config,
|
||||
patch_size=self.transformer_config.patch_size,
|
||||
patch_size_t=self.transformer_config.patch_size_t,
|
||||
**latent_conditions,
|
||||
)
|
||||
align_device_and_dtype(latent_conditions, accelerator.device, weight_dtype)
|
||||
align_device_and_dtype(text_conditions, accelerator.device, weight_dtype)
|
||||
batch_size = latent_conditions["latents"].shape[0]
|
||||
@@ -679,40 +717,64 @@ class Trainer:
|
||||
if "pooled_prompt_embeds" in text_conditions:
|
||||
text_conditions["pooled_prompt_embeds"].fill_(0)
|
||||
|
||||
# These weighting schemes use a uniform timestep sampling and instead post-weight the loss
|
||||
weights = compute_density_for_timestep_sampling(
|
||||
weighting_scheme=self.args.flow_weighting_scheme,
|
||||
sigmas = prepare_sigmas(
|
||||
scheduler=self.scheduler,
|
||||
sigmas=scheduler_sigmas,
|
||||
batch_size=batch_size,
|
||||
logit_mean=self.args.flow_logit_mean,
|
||||
logit_std=self.args.flow_logit_std,
|
||||
mode_scale=self.args.flow_mode_scale,
|
||||
num_train_timesteps=self.scheduler.config.num_train_timesteps,
|
||||
flow_weighting_scheme=self.args.flow_weighting_scheme,
|
||||
flow_logit_mean=self.args.flow_logit_mean,
|
||||
flow_logit_std=self.args.flow_logit_std,
|
||||
flow_mode_scale=self.args.flow_mode_scale,
|
||||
device=accelerator.device,
|
||||
generator=self.state.generator,
|
||||
)
|
||||
indices = (weights * self.scheduler.config.num_train_timesteps).long()
|
||||
sigmas = scheduler_sigmas[indices]
|
||||
timesteps = (sigmas * 1000.0).long()
|
||||
|
||||
noise = torch.randn(
|
||||
latent_conditions["latents"].shape,
|
||||
generator=generator,
|
||||
generator=self.state.generator,
|
||||
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
|
||||
sigmas = expand_tensor_dims(sigmas, ndim=noise.ndim)
|
||||
|
||||
# TODO(aryan): We probably don't need calculate_noisy_latents because we can determine the type of
|
||||
# scheduler and calculate the noisy latents accordingly. Look into this later.
|
||||
if "calculate_noisy_latents" in self.model_config.keys():
|
||||
noisy_latents = self.model_config["calculate_noisy_latents"](
|
||||
scheduler=self.scheduler,
|
||||
noise=noise,
|
||||
latents=latent_conditions["latents"],
|
||||
timesteps=timesteps,
|
||||
)
|
||||
else:
|
||||
# Default to flow-matching noise addition
|
||||
noisy_latents = (1.0 - sigmas) * latent_conditions["latents"] + sigmas * noise
|
||||
|
||||
latent_conditions.update({"noisy_latents": noisy_latents})
|
||||
|
||||
# These weighting schemes use a uniform timestep sampling and instead post-weight the loss
|
||||
weights = compute_loss_weighting_for_sd3(
|
||||
weighting_scheme=self.args.flow_weighting_scheme, sigmas=sigmas
|
||||
weights = prepare_loss_weights(
|
||||
scheduler=self.scheduler,
|
||||
alphas=scheduler_alphas[timesteps] if scheduler_alphas is not None else None,
|
||||
sigmas=sigmas,
|
||||
flow_weighting_scheme=self.args.flow_weighting_scheme,
|
||||
)
|
||||
weights = expand_tensor_dims(weights, noise.ndim)
|
||||
|
||||
pred = self.model_config["forward_pass"](
|
||||
transformer=self.transformer, timesteps=timesteps, **latent_conditions, **text_conditions
|
||||
transformer=self.transformer,
|
||||
scheduler=self.scheduler,
|
||||
timesteps=timesteps,
|
||||
**latent_conditions,
|
||||
**text_conditions,
|
||||
)
|
||||
target = prepare_target(
|
||||
scheduler=self.scheduler, noise=noise, latents=latent_conditions["latents"]
|
||||
)
|
||||
target = noise - latent_conditions["latents"]
|
||||
|
||||
loss = weights.float() * (pred["latents"].float() - target.float()).pow(2)
|
||||
# Average loss across channel dimension
|
||||
# Average loss across all but batch dimension
|
||||
loss = loss.mean(list(range(1, loss.ndim)))
|
||||
# Average loss across batch dimension
|
||||
loss = loss.mean()
|
||||
@@ -949,7 +1011,7 @@ class Trainer:
|
||||
self.transformer.train()
|
||||
|
||||
def evaluate(self) -> None:
|
||||
raise NotImplementedError
|
||||
raise NotImplementedError("Evaluation has not been implemented yet.")
|
||||
|
||||
def _init_distributed(self) -> None:
|
||||
logging_dir = Path(self.args.output_dir, self.args.logging_dir)
|
||||
|
||||
@@ -1,4 +1,12 @@
|
||||
from .diffusion_utils import default_flow_shift, resolution_dependant_timestep_flow_shift
|
||||
from .diffusion_utils import (
|
||||
default_flow_shift,
|
||||
get_scheduler_alphas,
|
||||
get_scheduler_sigmas,
|
||||
prepare_loss_weights,
|
||||
prepare_sigmas,
|
||||
prepare_target,
|
||||
resolution_dependant_timestep_flow_shift,
|
||||
)
|
||||
from .file_utils import delete_files, find_files
|
||||
from .memory_utils import bytes_to_gigabytes, free_memory, get_memory_statistics, make_contiguous
|
||||
from .optimizer_utils import get_optimizer, gradient_norm, max_gradient
|
||||
|
||||
@@ -1,4 +1,9 @@
|
||||
import math
|
||||
from typing import Optional, Union
|
||||
|
||||
import torch
|
||||
from diffusers import CogVideoXDDIMScheduler, FlowMatchEulerDiscreteScheduler
|
||||
from diffusers.training_utils import compute_loss_weighting_for_sd3
|
||||
|
||||
|
||||
# Default values copied from https://github.com/huggingface/diffusers/blob/8957324363d8b239d82db4909fbf8c0875683e3d/src/diffusers/schedulers/scheduling_flow_match_euler_discrete.py#L47
|
||||
@@ -28,3 +33,113 @@ def resolution_dependant_timestep_flow_shift(
|
||||
def default_flow_shift(sigmas: torch.Tensor, shift: float = 1.0) -> torch.Tensor:
|
||||
sigmas = (sigmas * shift) / (1 + (shift - 1) * sigmas)
|
||||
return sigmas
|
||||
|
||||
|
||||
def compute_density_for_timestep_sampling(
|
||||
weighting_scheme: str,
|
||||
batch_size: int,
|
||||
logit_mean: float = None,
|
||||
logit_std: float = None,
|
||||
mode_scale: float = None,
|
||||
device: torch.device = torch.device("cpu"),
|
||||
generator: Optional[torch.Generator] = None,
|
||||
) -> torch.Tensor:
|
||||
r"""
|
||||
Compute the density for sampling the timesteps when doing SD3 training.
|
||||
|
||||
Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528.
|
||||
|
||||
SD3 paper reference: https://arxiv.org/abs/2403.03206v1.
|
||||
"""
|
||||
if weighting_scheme == "logit_normal":
|
||||
# See 3.1 in the SD3 paper ($rf/lognorm(0.00,1.00)$).
|
||||
u = torch.normal(mean=logit_mean, std=logit_std, size=(batch_size,), device=device, generator=generator)
|
||||
u = torch.nn.functional.sigmoid(u)
|
||||
elif weighting_scheme == "mode":
|
||||
u = torch.rand(size=(batch_size,), device=device, generator=generator)
|
||||
u = 1 - u - mode_scale * (torch.cos(math.pi * u / 2) ** 2 - 1 + u)
|
||||
else:
|
||||
u = torch.rand(size=(batch_size,), device=device, generator=generator)
|
||||
return u
|
||||
|
||||
|
||||
def get_scheduler_alphas(scheduler: Union[CogVideoXDDIMScheduler, FlowMatchEulerDiscreteScheduler]) -> torch.Tensor:
|
||||
if isinstance(scheduler, FlowMatchEulerDiscreteScheduler):
|
||||
return None
|
||||
elif isinstance(scheduler, CogVideoXDDIMScheduler):
|
||||
return scheduler.alphas_cumprod.clone()
|
||||
else:
|
||||
raise ValueError(f"Unsupported scheduler type {type(scheduler)}")
|
||||
|
||||
|
||||
def get_scheduler_sigmas(scheduler: Union[CogVideoXDDIMScheduler, FlowMatchEulerDiscreteScheduler]) -> torch.Tensor:
|
||||
if isinstance(scheduler, FlowMatchEulerDiscreteScheduler):
|
||||
return scheduler.sigmas.clone()
|
||||
elif isinstance(scheduler, CogVideoXDDIMScheduler):
|
||||
return scheduler.timesteps.clone().float() / float(scheduler.config.num_train_timesteps)
|
||||
else:
|
||||
raise ValueError(f"Unsupported scheduler type {type(scheduler)}")
|
||||
|
||||
|
||||
def prepare_sigmas(
|
||||
scheduler: Union[CogVideoXDDIMScheduler, FlowMatchEulerDiscreteScheduler],
|
||||
sigmas: torch.Tensor,
|
||||
batch_size: int,
|
||||
num_train_timesteps: int,
|
||||
flow_weighting_scheme: str = "none",
|
||||
flow_logit_mean: float = 0.0,
|
||||
flow_logit_std: float = 1.0,
|
||||
flow_mode_scale: float = 1.29,
|
||||
device: torch.device = torch.device("cpu"),
|
||||
generator: Optional[torch.Generator] = None,
|
||||
) -> torch.Tensor:
|
||||
if isinstance(scheduler, FlowMatchEulerDiscreteScheduler):
|
||||
weights = compute_density_for_timestep_sampling(
|
||||
weighting_scheme=flow_weighting_scheme,
|
||||
batch_size=batch_size,
|
||||
logit_mean=flow_logit_mean,
|
||||
logit_std=flow_logit_std,
|
||||
mode_scale=flow_mode_scale,
|
||||
device=device,
|
||||
generator=generator,
|
||||
)
|
||||
indices = (weights * num_train_timesteps).long()
|
||||
elif isinstance(scheduler, CogVideoXDDIMScheduler):
|
||||
# TODO(aryan): Currently, only uniform sampling is supported. Add more sampling schemes.
|
||||
weights = torch.rand(size=(batch_size,), device=device, generator=generator)
|
||||
indices = (weights * num_train_timesteps).long()
|
||||
else:
|
||||
raise ValueError(f"Unsupported scheduler type {type(scheduler)}")
|
||||
|
||||
return sigmas[indices]
|
||||
|
||||
|
||||
def prepare_loss_weights(
|
||||
scheduler: Union[CogVideoXDDIMScheduler, FlowMatchEulerDiscreteScheduler],
|
||||
alphas: Optional[torch.Tensor] = None,
|
||||
sigmas: Optional[torch.Tensor] = None,
|
||||
flow_weighting_scheme: str = "none",
|
||||
) -> torch.Tensor:
|
||||
if isinstance(scheduler, FlowMatchEulerDiscreteScheduler):
|
||||
return compute_loss_weighting_for_sd3(sigmas, weighting_scheme=flow_weighting_scheme)
|
||||
elif isinstance(scheduler, CogVideoXDDIMScheduler):
|
||||
# SNR is computed as (alphas / (1 - alphas)), but for some reason CogVideoX uses 1 / (1 - alphas).
|
||||
# TODO(aryan): Experiment if using alphas / (1 - alphas) gives better results.
|
||||
return 1 / (1 - alphas)
|
||||
else:
|
||||
raise ValueError(f"Unsupported scheduler type {type(scheduler)}")
|
||||
|
||||
|
||||
def prepare_target(
|
||||
scheduler: Union[CogVideoXDDIMScheduler, FlowMatchEulerDiscreteScheduler],
|
||||
noise: torch.Tensor,
|
||||
latents: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
if isinstance(scheduler, FlowMatchEulerDiscreteScheduler):
|
||||
target = noise - latents
|
||||
elif isinstance(scheduler, CogVideoXDDIMScheduler):
|
||||
target = latents
|
||||
else:
|
||||
raise ValueError(f"Unsupported scheduler type {type(scheduler)}")
|
||||
|
||||
return target
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
import importlib
|
||||
import json
|
||||
import os
|
||||
|
||||
from huggingface_hub import hf_hub_download
|
||||
|
||||
|
||||
def resolve_vae_cls_from_ckpt_path(ckpt_path, **kwargs):
|
||||
ckpt_path = str(ckpt_path)
|
||||
if os.path.exists(str(ckpt_path)) and os.path.isdir(ckpt_path):
|
||||
index_path = os.path.join(ckpt_path, "model_index.json")
|
||||
else:
|
||||
revision = kwargs.get("revision", None)
|
||||
cache_dir = kwargs.get("cache_dir", None)
|
||||
index_path = hf_hub_download(
|
||||
repo_id=ckpt_path, filename="model_index.json", revision=revision, cache_dir=cache_dir
|
||||
)
|
||||
|
||||
with open(index_path, "r") as f:
|
||||
model_index_dict = json.load(f)
|
||||
assert "vae" in model_index_dict, "No VAE found in the modelx index dict."
|
||||
|
||||
vae_cls_config = model_index_dict["vae"]
|
||||
library = importlib.import_module(vae_cls_config[0])
|
||||
return getattr(library, vae_cls_config[1])
|
||||
@@ -29,7 +29,7 @@ def align_device_and_dtype(
|
||||
return x
|
||||
|
||||
|
||||
def expand_tensor_to_dims(tensor, ndim):
|
||||
def expand_tensor_dims(tensor, ndim):
|
||||
while len(tensor.shape) < ndim:
|
||||
tensor = tensor.unsqueeze(-1)
|
||||
return tensor
|
||||
|
||||
Reference in New Issue
Block a user