Compare commits

...
4 Commits
Author SHA1 Message Date
JerryZhou54 38ee9dc3b4 Complete Model Config design for VAEs 2025-04-23 19:08:11 +00:00
JerryZhou54 77b013fb8a Add model config for WanVAE 2025-04-23 19:01:13 +00:00
JerryZhou54 c31efe1234 Add model config for VAE 2025-04-23 18:59:25 +00:00
JerryZhou54 c056b89aea Add preliminary design for model config 2025-04-23 18:57:13 +00:00
30 changed files with 609 additions and 357 deletions
-9
View File
@@ -1,9 +0,0 @@
from fastvideo.v1.configs.base import BaseConfig, SlidingTileAttnConfig
from fastvideo.v1.configs.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.v1.configs.registry import get_pipeline_config_cls_for_name
from fastvideo.v1.configs.wan import WanI2V480PConfig, WanT2V480PConfig
__all__ = [
"HunyuanConfig", "FastHunyuanConfig", "BaseConfig", "SlidingTileAttnConfig",
"WanT2V480PConfig", "WanI2V480PConfig", "get_pipeline_config_cls_for_name"
]
-67
View File
@@ -1,67 +0,0 @@
from dataclasses import dataclass
from typing import Optional
@dataclass
class BaseConfig:
"""Base configuration for all pipeline architectures."""
# Video parameters
height: int = 720
width: int = 1280
num_frames: int = 125
fps: int = 24
# Video generation parameters
num_inference_steps: int = 50
guidance_scale: float = 1.0
seed: int = 1024
guidance_rescale: float = 0.0
embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None
use_cpu_offload: bool = False
disable_autocast: bool = False
# Model configuration
precision: str = "bf16"
# VAE configuration
vae_precision: str = "fp16"
vae_tiling: bool = True
vae_sp: bool = True
vae_scale_factor: Optional[int] = None
# DiT configuration
num_channels_latents: Optional[int] = None
# Image encoder configuration
image_encoder_precision: str = "fp32"
# Text encoder configuration
text_encoder_precision: str = "fp16"
text_len: int = -1
hidden_state_skip_layer: int = 0
# STA (Spatial-Temporal Attention) parameters
mask_strategy_file_path: Optional[str] = None
enable_torch_compile: bool = False
neg_prompt: Optional[str] = None
@dataclass
class SlidingTileAttnConfig(BaseConfig):
"""Configuration for sliding tile attention."""
# Override any BaseConfig defaults as needed
# Add sliding tile specific parameters
window_size: int = 16
stride: int = 8
# You can provide custom defaults for inherited fields
height: int = 576
width: int = 1024
# Additional configuration specific to sliding tile attention
pad_to_square: bool = False
use_overlap_optimization: bool = True
+7
View File
@@ -0,0 +1,7 @@
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
__all__ = [
"ArchConfig", "ModelConfig",
"VAEArchConfig", "VAEConfig"
]
+47
View File
@@ -0,0 +1,47 @@
from dataclasses import dataclass, fields
from typing import Dict, Any
# 1. ArchConfig contains all fields from diffuser's/transformer's config.json (i.e. all fields related to the architecture of the model)
# 2. ArchConfig should be inherited & overriden by each model arch_config
# 3. Any field in ArchConfig is fixed upon initialization, and should be hidden away from users
@dataclass
class ArchConfig:
pass
@dataclass
class ModelConfig:
# Every model config parameter can be categorized into either ArchConfig or everything else
# Diffuser/Transformer parameters
arch_config: ArchConfig = ArchConfig()
# FastVideo-specific parameters here
# i.e. STA, quantization, teacache
# This should be used only when loading from transformers/diffusers
def update_model_arch(
self,
source_model_dict: Dict[str, Any]
) -> None:
arch_config = self.arch_config
valid_fields = {f.name for f in fields(arch_config)}
for key, value in source_model_dict.items():
if key in valid_fields:
setattr(arch_config, key, value)
else:
raise AttributeError(f"{type(arch_config).__name__} has no field '{key}'")
def update_model_config(
self,
source_model_dict: Dict[str, Any]
) -> None:
assert "arch_config" not in source_model_dict.keys(), "Source model config shouldn't contain arch_config."
valid_fields = {f.name for f in fields(self)}
for key, value in source_model_dict.items():
if key in valid_fields:
setattr(self, key, value)
else:
print(f"{type(self).__name__} does not contain field '{key}'!")
raise AttributeError(f"Invalid field: {key}")
@@ -0,0 +1,7 @@
from fastvideo.v1.configs.models.vaes.hunyuanvae import HunyuanVAEConfig, HunyuanVAEArchConfig
from fastvideo.v1.configs.models.vaes.wanvae import WanVAEConfig, WanVAEArchConfig
__all__ = [
"HunyuanVAEConfig", "HunyuanVAEArchConfig",
"WanVAEConfig", "WanVAEArchConfig"
]
+36
View File
@@ -0,0 +1,36 @@
from dataclasses import dataclass
from typing import Union
import torch
from fastvideo.v1.configs.models import ArchConfig, ModelConfig
@dataclass
class VAEArchConfig(ArchConfig):
scaling_factor: Union[float, torch.tensor] = 0
temporal_compression_ratio: int = 4
spatial_compression_ratio: int = 8
@dataclass
class VAEConfig(ModelConfig):
arch_config: VAEArchConfig = VAEArchConfig()
# FastVideoVAE-specific parameters
load_encoder: bool = True
load_decoder: bool = True
tile_sample_min_height: int = 256
tile_sample_min_width: int = 256
tile_sample_min_num_frames: int = 16
tile_sample_stride_height: int = 192
tile_sample_stride_width: int = 192
tile_sample_stride_num_frames: int = 12
blend_num_frames: int = 0
use_tiling: bool = True
use_temporal_tiling: bool = True
use_parallel_tiling: bool = True
def __post_init__(self):
self.blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
@@ -0,0 +1,37 @@
from dataclasses import dataclass
from typing import Tuple
from fastvideo.v1.configs.models.vaes.base import VAEConfig, VAEArchConfig
@dataclass
class HunyuanVAEArchConfig(VAEArchConfig):
in_channels: int = 3
out_channels: int = 3
latent_channels: int = 16
down_block_types: Tuple[str, ...] = (
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
)
up_block_types: Tuple[str, ...] = (
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
)
block_out_channels: Tuple[int, ...] = (128, 256, 512, 512)
layers_per_block: int = 2
act_fn: str = "silu"
norm_num_groups: int = 32
scaling_factor: float = 0.476986
spatial_compression_ratio: int = 8
temporal_compression_ratio: int = 4
mid_block_add_attention: bool = True
def __post_init__(self):
self.spatial_compression_ratio: int = 2**(len(self.block_out_channels)-1)
@dataclass
class HunyuanVAEConfig(VAEConfig):
arch_config: VAEArchConfig = HunyuanVAEArchConfig()
@@ -0,0 +1,72 @@
from dataclasses import dataclass
from typing import Tuple
import torch
from fastvideo.v1.configs.models.vaes.base import VAEConfig, VAEArchConfig
@dataclass
class WanVAEArchConfig(VAEArchConfig):
base_dim: int = 96
z_dim: int = 16
dim_mult: Tuple[int, ...] = (1, 2, 4, 4)
num_res_blocks: int = 2
attn_scales: Tuple[float, ...] = ()
temperal_downsample: Tuple[bool, ...] = (False, True, True)
dropout: float = 0.0
latents_mean: Tuple[float, ...] = (
-0.7571,
-0.7089,
-0.9113,
0.1075,
-0.1745,
0.9653,
-0.1517,
1.5508,
0.4134,
-0.0715,
0.5517,
-0.3632,
-0.1922,
-0.9497,
0.2503,
-0.2921,
)
latents_std: Tuple[float, ...] = (
2.8184,
1.4541,
2.3275,
2.6558,
1.2196,
1.7708,
2.6052,
2.0743,
3.2687,
2.1526,
2.8652,
1.5579,
1.6382,
1.1253,
2.8251,
1.9160,
)
temporal_compression_ratio = 4
spatial_compression_ratio = 8
def __post_init__(self):
self.scaling_factor: torch.tensor = 1.0 / torch.tensor(self.latents_std).view(
1, self.z_dim, 1, 1, 1)
self.shift_factor: torch.tensor = torch.tensor(self.latents_mean).view(
1, self.z_dim, 1, 1, 1)
@dataclass
class WanVAEConfig(VAEConfig):
arch_config: VAEArchConfig = WanVAEArchConfig()
use_feature_cache: bool = True
use_tiling: bool = False
use_temporal_tiling: bool = False
use_parallel_tiling: bool = False
def __post_init__(self):
self.blend_num_frames = (self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames) * 2
@@ -0,0 +1,9 @@
from fastvideo.v1.configs.pipelines.hunyuan import HunyuanConfig, FastHunyuanConfig
from fastvideo.v1.configs.pipelines.wan import WanT2V480PConfig, WanI2V480PConfig
from fastvideo.v1.configs.pipelines.base import BaseConfig, SlidingTileAttnConfig
from fastvideo.v1.configs.pipelines.registry import get_pipeline_config_cls_for_name
__all__ = [
"HunyuanConfig", "FastHunyuanConfig", "BaseConfig", "SlidingTileAttnConfig",
"WanT2V480PConfig", "WanI2V480PConfig", "get_pipeline_config_cls_for_name"
]
+103
View File
@@ -0,0 +1,103 @@
from dataclasses import dataclass, asdict, fields
from typing import Optional, Dict, Any
import json
from fastvideo.v1.configs.models import ModelConfig, VAEConfig
from fastvideo.v1.utils import shallow_asdict
@dataclass
class BaseConfig:
"""Base configuration for all pipeline architectures."""
# Video parameters
height: int = 720
width: int = 1280
num_frames: int = 125
fps: int = 24
# Video generation parameters
num_inference_steps: int = 50
guidance_scale: float = 1.0
seed: int = 1024
guidance_rescale: float = 0.0
embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None
use_cpu_offload: bool = False
disable_autocast: bool = False
# Model configuration
precision: str = "bf16"
# VAE configuration
vae_precision: str = "fp16"
vae_tiling: bool = True
vae_sp: bool = True
# vae_scale_factor: Optional[int] = None # Deprecated
vae_config: VAEConfig = VAEConfig()
# DiT configuration
num_channels_latents: Optional[int] = None # Deprecated
# Image encoder configuration
image_encoder_precision: str = "fp32"
# Text encoder configuration
text_encoder_precision: str = "fp16"
text_len: int = -1 # Deprecated
hidden_state_skip_layer: int = 0 # Deprecated
# STA (Spatial-Temporal Attention) parameters
mask_strategy_file_path: Optional[str] = None
enable_torch_compile: bool = False
neg_prompt: Optional[str] = None
def dump_to_json(self, file_path: str):
output_dict = shallow_asdict(self)
for key, value in output_dict.items():
if isinstance(value, ModelConfig):
model_dict = asdict(value)
# Model Arch Config should be hidden away from the users
model_dict.pop("arch_config")
output_dict[key] = model_dict
with open(file_path, "w") as f:
json.dump(output_dict, f, indent=2)
def load_from_json(self, file_path: str):
with open(file_path, "r") as f:
input_pipeline_dict = json.load(f)
self.update_pipeline_config(input_pipeline_dict)
def update_pipeline_config(
self,
source_pipeline_dict: Dict[str, Any]
) -> None:
for f in fields(self):
key = f.name
if key in source_pipeline_dict:
current_value = getattr(self, key)
new_value = source_pipeline_dict[key]
# If it's a nested ModelConfig, update it recursively
if isinstance(current_value, ModelConfig):
current_value.update_model_config(new_value)
else:
setattr(self, key, new_value)
@dataclass
class SlidingTileAttnConfig(BaseConfig):
"""Configuration for sliding tile attention."""
# Override any BaseConfig defaults as needed
# Add sliding tile specific parameters
window_size: int = 16
stride: int = 8
# You can provide custom defaults for inherited fields
height: int = 576
width: int = 1024
# Additional configuration specific to sliding tile attention
pad_to_square: bool = False
use_overlap_optimization: bool = True
@@ -1,13 +1,16 @@
from dataclasses import dataclass
from fastvideo.v1.configs.pipelines.base import BaseConfig
from fastvideo.v1.configs.base import BaseConfig
from fastvideo.v1.configs.models import VAEConfig
from fastvideo.v1.configs.models.vaes import HunyuanVAEConfig
@dataclass
class HunyuanConfig(BaseConfig):
"""Base configuration for HunYuan pipeline architecture."""
# HunyuanConfig-specific parameters with defaults
# VAE
vae_config: VAEConfig = HunyuanVAEConfig()
# Denoising stage
embedded_cfg_scale: int = 6
flow_shift: int = 7
@@ -27,6 +30,10 @@ class HunyuanConfig(BaseConfig):
text_encoder_precision_2: str = "fp16"
text_len_2: int = 77
def __post_init__(self):
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
@dataclass
class FastHunyuanConfig(HunyuanConfig):
@@ -3,9 +3,10 @@
import os
from typing import Callable, Dict, Optional, Type
from fastvideo.v1.configs.base import BaseConfig
from fastvideo.v1.configs.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.v1.configs.wan import WanI2V480PConfig, WanT2V480PConfig
from fastvideo.v1.configs.pipelines.base import BaseConfig
from fastvideo.v1.configs.pipelines.hunyuan import HunyuanConfig, FastHunyuanConfig
from fastvideo.v1.configs.pipelines.wan import WanT2V480PConfig, WanI2V480PConfig
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import (maybe_download_model_index,
verify_model_config_and_directory)
@@ -1,13 +1,19 @@
from dataclasses import dataclass
from fastvideo.v1.configs.pipelines.base import BaseConfig
from fastvideo.v1.configs.base import BaseConfig
from fastvideo.v1.configs.models import VAEConfig
from fastvideo.v1.configs.models.vaes import WanVAEConfig
@dataclass
class WanT2V480PConfig(BaseConfig):
"""Base configuration for Wan T2V 1.3B pipeline architecture."""
# WanConfig-specific parameters with defaults
# VAE
vae_config: VAEConfig = WanVAEConfig()
vae_tiling: bool = False
vae_sp: bool = False
# Video parameters
height: int = 480
width: int = 832
@@ -31,6 +37,9 @@ class WanT2V480PConfig(BaseConfig):
# WanConfig-specific added parameters
def __post_init__(self):
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
@dataclass
class WanI2V480PConfig(WanT2V480PConfig):
@@ -43,3 +52,7 @@ class WanI2V480PConfig(WanT2V480PConfig):
# Precision for each component
image_encoder_precision: str = "fp32"
def __post_init__(self):
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
@@ -0,0 +1,41 @@
{
"height": 480,
"width": 832,
"num_frames": 81,
"fps": 16,
"num_inference_steps": 50,
"guidance_scale": 3.0,
"seed": 1024,
"guidance_rescale": 0.0,
"embedded_cfg_scale": 6.0,
"flow_shift": 3,
"use_cpu_offload": true,
"disable_autocast": false,
"precision": "bf16",
"vae_precision": "fp16",
"vae_tiling": false,
"vae_sp": false,
"vae_config": {
"load_encoder": false,
"load_decoder": true,
"tile_sample_min_height": 256,
"tile_sample_min_width": 256,
"tile_sample_min_num_frames": 16,
"tile_sample_stride_height": 192,
"tile_sample_stride_width": 192,
"tile_sample_stride_num_frames": 12,
"blend_num_frames": 8,
"use_tiling": false,
"use_temporal_tiling": false,
"use_parallel_tiling": false,
"use_feature_cache": true
},
"num_channels_latents": null,
"image_encoder_precision": "fp32",
"text_encoder_precision": "fp32",
"text_len": 512,
"hidden_state_skip_layer": 0,
"mask_strategy_file_path": null,
"enable_torch_compile": false,
"neg_prompt": "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
}
@@ -0,0 +1,41 @@
{
"height": 480,
"width": 832,
"num_frames": 81,
"fps": 16,
"num_inference_steps": 40,
"guidance_scale": 5.0,
"seed": 1024,
"guidance_rescale": 0.0,
"embedded_cfg_scale": 6.0,
"flow_shift": 3,
"use_cpu_offload": true,
"disable_autocast": false,
"precision": "bf16",
"vae_precision": "fp16",
"vae_tiling": false,
"vae_sp": false,
"vae_config": {
"load_encoder": true,
"load_decoder": true,
"tile_sample_min_height": 256,
"tile_sample_min_width": 256,
"tile_sample_min_num_frames": 16,
"tile_sample_stride_height": 192,
"tile_sample_stride_width": 192,
"tile_sample_stride_num_frames": 12,
"blend_num_frames": 8,
"use_tiling": false,
"use_temporal_tiling": false,
"use_parallel_tiling": false,
"use_feature_cache": true
},
"num_channels_latents": null,
"image_encoder_precision": "fp32",
"text_encoder_precision": "fp32",
"text_len": 512,
"hidden_state_skip_layer": 0,
"mask_strategy_file_path": null,
"enable_torch_compile": false,
"neg_prompt": "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
}
+23 -10
View File
@@ -17,11 +17,11 @@ import torch
import torchvision
from einops import rearrange
from fastvideo.v1.configs import get_pipeline_config_cls_for_name
from fastvideo.v1.configs.pipelines import get_pipeline_config_cls_for_name, BaseConfig
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines import ForwardBatch
from fastvideo.v1.utils import align_to
from fastvideo.v1.utils import align_to, shallow_asdict
from fastvideo.v1.worker.executor import Executor
logger = init_logger(__name__)
@@ -52,6 +52,7 @@ class VideoGenerator:
model_path: str,
device: Optional[str] = None,
torch_dtype: Optional[torch.dtype] = None,
pipeline_config: Optional[Union[str | BaseConfig]] = None,
**kwargs) -> "VideoGenerator":
"""
Create a video generator from a pretrained model.
@@ -64,22 +65,30 @@ class VideoGenerator:
Returns:
The created video generator
Priority level: Default pipeline config < User's pipeline config < User's kwargs
"""
config = None
config_cls = get_pipeline_config_cls_for_name(model_path)
if config_cls is not None:
config = config_cls()
# 1. If users provide a pipeline config, it will override the default pipeline config
if isinstance(pipeline_config, BaseConfig):
config = pipeline_config
else:
config_cls = get_pipeline_config_cls_for_name(model_path)
if config_cls is not None:
config = config_cls()
if isinstance(pipeline_config, str):
config.load_from_json(pipeline_config)
# 2. If users also provide some kwargs, it will override the pipeline config.
# The user kwargs shouldn't contain model config parameters!
if config is None:
logger.warning("No config found for model %s, using default config",
model_path)
config_args = {}
config_args = kwargs
else:
config_args = asdict(config)
# override config_args with kwargs
config_args.update(kwargs)
config_args = shallow_asdict(config)
config_args.update(kwargs)
fastvideo_args = FastVideoArgs(
model_path=model_path,
@@ -115,6 +124,7 @@ class VideoGenerator:
def generate_video(
self,
prompt: str,
image_path: Optional[str] = None,
negative_prompt: Optional[str] = None,
output_path: Optional[str] = None,
save_video: bool = True,
@@ -155,6 +165,8 @@ class VideoGenerator:
fastvideo_args = self.fastvideo_args
# Override parameters if provided
if image_path is not None:
fastvideo_args.image_path = image_path
if negative_prompt is not None:
fastvideo_args.neg_prompt = negative_prompt
if num_inference_steps is not None:
@@ -224,6 +236,7 @@ class VideoGenerator:
device = torch.device(fastvideo_args.device_str)
batch = ForwardBatch(
prompt=prompt,
image_path=fastvideo_args.image_path,
negative_prompt=fastvideo_args.neg_prompt,
num_videos_per_prompt=fastvideo_args.num_videos,
height=fastvideo_args.height,
+6 -3
View File
@@ -10,6 +10,8 @@ from typing import List, Optional
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import FlexibleArgumentParser
from fastvideo.v1.configs.models import VAEConfig
logger = init_logger(__name__)
@@ -50,9 +52,10 @@ class FastVideoArgs:
# VAE configuration
vae_precision: str = "fp16"
vae_tiling: bool = True
vae_sp: bool = False
vae_scale_factor: Optional[int] = None
vae_tiling: bool = True # Might change in between forward passes
vae_sp: bool = False # Might change in between forward passes
# vae_scale_factor: Optional[int] = None # Deprecated
vae_config: VAEConfig = VAEConfig()
# DiT configuration
num_channels_latents: Optional[int] = None
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
import dataclasses
from dataclasses import asdict
import glob
import os
import time
@@ -310,9 +311,11 @@ class VAELoader(ComponentLoader):
assert class_name is not None, "Model config does not contain a _class_name attribute. Only diffusers format is supported."
config.pop("_diffusers_version")
vae_config = fastvideo_args.vae_config
vae_config.update_model_arch(config)
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vae = vae_cls(**config).to(fastvideo_args.device)
vae = vae_cls(vae_config).to(fastvideo_args.device)
# Find all safetensors files
safetensors_list = glob.glob(
@@ -322,7 +325,7 @@ class VAELoader(ComponentLoader):
safetensors_list
) == 1, f"Found {len(safetensors_list)} safetensors files in {model_path}"
loaded = safetensors_load_file(safetensors_list[0])
vae.load_state_dict(loaded)
vae.load_state_dict(loaded, strict=False) # We might only load encoder or decoder
dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
vae = vae.eval().to(dtype)
@@ -343,6 +346,10 @@ class TransformerLoader(ComponentLoader):
"Only diffusers format is supported.")
model_config.pop("_diffusers_version")
# Config from Diffusers supercedes fastvideo's model config
# dit_config = fastvideo_args.dit_config
# model_config.update(dit_config)
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
# Find all safetensors files
+37 -19
View File
@@ -11,6 +11,7 @@ from diffusers.utils.torch_utils import randn_tensor
from fastvideo.v1.distributed import (get_sequence_model_parallel_rank,
get_sequence_model_parallel_world_size)
from fastvideo.v1.configs.models import VAEConfig
class ParallelTiledVAE(ABC):
@@ -20,29 +21,36 @@ class ParallelTiledVAE(ABC):
tile_sample_stride_height: int
tile_sample_stride_width: int
tile_sample_stride_num_frames: int
blend_num_frames: int
use_tiling: bool
use_temporal_tiling: bool
use_parallel_tiling: bool
temporal_compression_ratio: int
spatial_compression_ratio: int
scaling_factor: Union[float, torch.tensor]
def __init__(self, *args, **kwargs) -> None:
# Check if subclass has defined all required properties
required_attributes = [
'tile_sample_min_height', 'tile_sample_min_width',
'tile_sample_min_num_frames', 'tile_sample_stride_height',
'tile_sample_stride_width', 'tile_sample_stride_num_frames',
'spatial_compression_ratio', 'temporal_compression_ratio',
'use_tiling', 'use_temporal_tiling', 'use_parallel_tiling',
'scaling_factor'
]
def __init__(self, config: VAEConfig, **kwargs) -> None:
self.config = config
self.arch_config = config.arch_config
self.tile_sample_min_height = config.tile_sample_min_height
self.tile_sample_min_width = config.tile_sample_min_width
self.tile_sample_min_num_frames = config.tile_sample_min_num_frames
self.tile_sample_stride_height = config.tile_sample_stride_height
self.tile_sample_stride_width = config.tile_sample_stride_width
self.tile_sample_stride_num_frames = config.tile_sample_stride_num_frames
self.blend_num_frames = config.blend_num_frames
self.use_tiling = config.use_tiling
self.use_temporal_tiling = config.use_temporal_tiling
self.use_parallel_tiling = config.use_parallel_tiling
for attr in required_attributes:
if not hasattr(self, attr):
raise AttributeError(
f"Subclasses of ParallelVAE must define '{attr}' property")
self.blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
@property
def temporal_compression_ratio(self) -> int:
return self.arch_config.temporal_compression_ratio
@property
def spatial_compression_ratio(self) -> int:
return self.arch_config.spatial_compression_ratio
@property
def scaling_factor(self) -> Union[float, torch.tensor]:
return self.arch_config.scaling_factor
@abstractmethod
def _encode(self, *args, **kwargs) -> torch.Tensor:
@@ -408,6 +416,10 @@ class ParallelTiledVAE(ABC):
tile_sample_stride_height: Optional[int] = None,
tile_sample_stride_width: Optional[int] = None,
tile_sample_stride_num_frames: Optional[int] = None,
blend_num_frames: Optional[int] = None,
use_tiling: Optional[bool] = None,
use_temporal_tiling: Optional[bool] = None,
use_parallel_tiling: Optional[bool] = None,
) -> None:
r"""
Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
@@ -439,7 +451,13 @@ class ParallelTiledVAE(ABC):
self.tile_sample_stride_height = tile_sample_stride_height or self.tile_sample_stride_height
self.tile_sample_stride_width = tile_sample_stride_width or self.tile_sample_stride_width
self.tile_sample_stride_num_frames = tile_sample_stride_num_frames or self.tile_sample_stride_num_frames
self.blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
if blend_num_frames is not None:
self.blend_num_frames = blend_num_frames
else:
self.blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
self.use_tiling = use_tiling or self.use_tiling
self.use_temporal_tiling = use_temporal_tiling or self.use_temporal_tiling
self.use_parallel_tiling = use_parallel_tiling or self.use_parallel_tiling
def disable_tiling(self) -> None:
r"""
+33 -75
View File
@@ -15,7 +15,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.
from typing import Optional, Tuple, Union
from typing import Optional, Tuple, Union, cast
import numpy as np
import torch
@@ -24,8 +24,8 @@ import torch.nn.functional as F
import torch.utils.checkpoint
from fastvideo.v1.layers.activation import get_act_fn
from fastvideo.v1.models.utils import auto_attributes
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
from fastvideo.v1.configs.models.vaes import HunyuanVAEConfig, HunyuanVAEArchConfig
def prepare_causal_attention_mask(
@@ -773,95 +773,53 @@ class AutoencoderKLHunyuanVideo(nn.Module, ParallelTiledVAE):
_supports_gradient_checkpointing = True
@auto_attributes
def __init__(
self,
in_channels: int = 3,
out_channels: int = 3,
latent_channels: int = 16,
down_block_types: Tuple[str, ...] = (
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
),
up_block_types: Tuple[str, ...] = (
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
),
block_out_channels: Tuple[int, ...] = (128, 256, 512, 512),
layers_per_block: int = 2,
act_fn: str = "silu",
norm_num_groups: int = 32,
scaling_factor: float = 0.476986,
spatial_compression_ratio: int = 8,
temporal_compression_ratio: int = 4,
mid_block_add_attention: bool = True,
load_encoder: bool = True,
load_decoder: bool = True,
config: HunyuanVAEConfig,
) -> None:
super().__init__()
ParallelTiledVAE.__init__(self, config)
arch_config: HunyuanVAEArchConfig = cast(HunyuanVAEArchConfig, config.arch_config)
# TODO(will): only pass in config. We do this by manually defining a
# config for hunyuan vae
self.block_out_channels = block_out_channels
if load_encoder:
self.block_out_channels = arch_config.block_out_channels
if config.load_encoder:
self.encoder = HunyuanVideoEncoder3D(
in_channels=in_channels,
out_channels=latent_channels,
down_block_types=down_block_types,
block_out_channels=block_out_channels,
layers_per_block=layers_per_block,
norm_num_groups=norm_num_groups,
act_fn=act_fn,
in_channels=arch_config.in_channels,
out_channels=arch_config.latent_channels,
down_block_types=arch_config.down_block_types,
block_out_channels=arch_config.block_out_channels,
layers_per_block=arch_config.layers_per_block,
norm_num_groups=arch_config.norm_num_groups,
act_fn=arch_config.act_fn,
double_z=True,
mid_block_add_attention=mid_block_add_attention,
temporal_compression_ratio=temporal_compression_ratio,
spatial_compression_ratio=spatial_compression_ratio,
mid_block_add_attention=arch_config.mid_block_add_attention,
temporal_compression_ratio=arch_config.temporal_compression_ratio,
spatial_compression_ratio=arch_config.spatial_compression_ratio,
)
self.quant_conv = nn.Conv3d(2 * latent_channels,
2 * latent_channels,
self.quant_conv = nn.Conv3d(2 * arch_config.latent_channels,
2 * arch_config.latent_channels,
kernel_size=1)
if load_decoder:
if config.load_decoder:
self.decoder = HunyuanVideoDecoder3D(
in_channels=latent_channels,
out_channels=out_channels,
up_block_types=up_block_types,
block_out_channels=block_out_channels,
layers_per_block=layers_per_block,
norm_num_groups=norm_num_groups,
act_fn=act_fn,
time_compression_ratio=temporal_compression_ratio,
spatial_compression_ratio=spatial_compression_ratio,
mid_block_add_attention=mid_block_add_attention,
in_channels=arch_config.latent_channels,
out_channels=arch_config.out_channels,
up_block_types=arch_config.up_block_types,
block_out_channels=arch_config.block_out_channels,
layers_per_block=arch_config.layers_per_block,
norm_num_groups=arch_config.norm_num_groups,
act_fn=arch_config.act_fn,
time_compression_ratio=arch_config.temporal_compression_ratio,
spatial_compression_ratio=arch_config.spatial_compression_ratio,
mid_block_add_attention=arch_config.mid_block_add_attention,
)
self.post_quant_conv = nn.Conv3d(latent_channels,
latent_channels,
self.post_quant_conv = nn.Conv3d(arch_config.latent_channels,
arch_config.latent_channels,
kernel_size=1)
# When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent
# frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the
# intermediate tiles together, the memory requirement can be lowered.
self.use_tiling = True
self.use_temporal_tiling = True
self.use_parallel_tiling = True
self.scaling_factor = scaling_factor
# The minimal tile height and width for spatial tiling to be used
self.tile_sample_min_height = 256
self.tile_sample_min_width = 256
self.tile_sample_min_num_frames = 16
# The minimal distance between two spatial tiles
self.tile_sample_stride_height = 192
self.tile_sample_stride_width = 192
self.tile_sample_stride_num_frames = 12
ParallelTiledVAE.__init__(self)
def _encode(self, x: torch.Tensor) -> torch.Tensor:
x = self.encoder(x)
enc = self.quant_conv(x)
+33 -94
View File
@@ -16,16 +16,15 @@
import contextvars
from contextlib import contextmanager
from typing import Optional, Tuple, Union
from typing import Optional, Tuple, Union, cast
import torch
import torch.nn as nn
import torch.nn.functional as F
from fastvideo.v1.layers.activation import get_act_fn
from fastvideo.v1.models.utils import auto_attributes
from fastvideo.v1.models.vaes.common import (DiagonalGaussianDistribution,
ParallelTiledVAE)
from fastvideo.v1.models.vaes.common import ParallelTiledVAE, DiagonalGaussianDistribution
from fastvideo.v1.configs.models.vaes import WanVAEConfig, WanVAEArchConfig
CACHE_T = 2
@@ -781,96 +780,34 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
_supports_gradient_checkpointing = False
@auto_attributes
def __init__(self,
base_dim: int = 96,
z_dim: int = 16,
dim_mult: Tuple[int, ...] = (1, 2, 4, 4),
num_res_blocks: int = 2,
attn_scales: Tuple[float, ...] = (),
temperal_downsample: Tuple[bool, ...] = (False, True, True),
dropout: float = 0.0,
latents_mean: Tuple[float, ...] = (
-0.7571,
-0.7089,
-0.9113,
0.1075,
-0.1745,
0.9653,
-0.1517,
1.5508,
0.4134,
-0.0715,
0.5517,
-0.3632,
-0.1922,
-0.9497,
0.2503,
-0.2921,
),
latents_std: Tuple[float, ...] = (
2.8184,
1.4541,
2.3275,
2.6558,
1.2196,
1.7708,
2.6052,
2.0743,
3.2687,
2.1526,
2.8652,
1.5579,
1.6382,
1.1253,
2.8251,
1.9160,
),
load_encoder: bool = True,
load_decoder: bool = True) -> None:
config: WanVAEConfig,
) -> None:
super().__init__()
ParallelTiledVAE.__init__(self, config)
self.z_dim = z_dim
self.temperal_downsample = list(temperal_downsample)
self.temperal_upsample = list(temperal_downsample)[::-1]
self.latents_mean = list(latents_mean)
self.latents_std = list(latents_std)
self.scaling_factor = 1.0 / torch.tensor(self.config.latents_std).view(
1, self.config.z_dim, 1, 1, 1)
self.shift_factor = torch.tensor(self.config.latents_mean).view(
1, self.config.z_dim, 1, 1, 1)
self.arch_config = cast(WanVAEArchConfig, self.arch_config)
self.z_dim = self.arch_config.z_dim
self.temperal_downsample = list(self.arch_config.temperal_downsample)
self.temperal_upsample = list(self.arch_config.temperal_downsample)[::-1]
self.latents_mean = list(self.arch_config.latents_mean)
self.latents_std = list(self.arch_config.latents_std)
self.shift_factor = self.arch_config.shift_factor
if load_encoder:
self.encoder = WanEncoder3d(base_dim, z_dim * 2, dim_mult,
num_res_blocks, attn_scales,
self.temperal_downsample, dropout)
self.quant_conv = WanCausalConv3d(z_dim * 2, z_dim * 2, 1)
self.post_quant_conv = WanCausalConv3d(z_dim, z_dim, 1)
if config.load_encoder:
self.encoder = WanEncoder3d(self.arch_config.base_dim, self.z_dim * 2, self.arch_config.dim_mult,
self.arch_config.num_res_blocks, self.arch_config.attn_scales,
self.temperal_downsample, self.arch_config.dropout)
self.quant_conv = WanCausalConv3d(self.z_dim * 2, self.z_dim * 2, 1)
self.post_quant_conv = WanCausalConv3d(self.z_dim, self.z_dim, 1)
if load_decoder:
self.decoder = WanDecoder3d(base_dim, z_dim, dim_mult,
num_res_blocks, attn_scales,
self.temperal_upsample, dropout)
if config.load_decoder:
self.decoder = WanDecoder3d(self.arch_config.base_dim, self.z_dim, self.arch_config.dim_mult,
self.arch_config.num_res_blocks, self.arch_config.attn_scales,
self.temperal_upsample, self.arch_config.dropout)
self.use_tiling = True
self.use_temporal_tiling = False
self.use_parallel_tiling = False
self.spatial_compression_ratio = 8
self.temporal_compression_ratio = 4
# The minimal tile height and width for spatial tiling to be used
self.tile_sample_min_height = 256
self.tile_sample_min_width = 256
self.tile_sample_min_num_frames = 16
# The minimal distance between two spatial tiles
self.tile_sample_stride_height = 192
self.tile_sample_stride_width = 192
self.tile_sample_stride_num_frames = 12
# Whether to use the feature cache algorithm used by diffusers and Wan2.1
self.use_feature_cache = True # default to True for best performance
ParallelTiledVAE.__init__(self)
self.use_feature_cache = config.use_feature_cache
def clear_cache(self) -> None:
@@ -881,13 +818,15 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
count += 1
return count
self._conv_num = _count_conv3d(self.decoder)
self._conv_idx = 0
self._feat_map = [None] * self._conv_num
if self.config.load_decoder:
self._conv_num = _count_conv3d(self.decoder)
self._conv_idx = 0
self._feat_map = [None] * self._conv_num
# cache encode
self._enc_conv_num = _count_conv3d(self.encoder)
self._enc_conv_idx = 0
self._enc_feat_map = [None] * self._enc_conv_num
if self.config.load_encoder:
self._enc_conv_num = _count_conv3d(self.encoder)
self._enc_conv_idx = 0
self._enc_feat_map = [None] * self._enc_conv_num
def encode(self, x: torch.Tensor) -> torch.Tensor:
if self.use_feature_cache:
@@ -6,8 +6,6 @@ This module contains an implementation of the Hunyuan video diffusion pipeline
using the modular pipeline architecture.
"""
from diffusers.image_processor import VaeImageProcessor
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
@@ -71,14 +69,6 @@ class HunyuanVideoPipeline(ComposedPipelineBase):
"""
Initialize the pipeline.
"""
vae_scale_factor = 2**(len(self.get_module("vae").block_out_channels) -
1)
fastvideo_args.vae_scale_factor = vae_scale_factor
self.image_processor = VaeImageProcessor(
vae_scale_factor=vae_scale_factor)
self.add_module("image_processor", self.image_processor)
num_channels_latents = self.get_module("transformer").in_channels
fastvideo_args.num_channels_latents = num_channels_latents
+3 -2
View File
@@ -10,6 +10,7 @@ from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.utils import PRECISION_TO_TYPE
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
logger = init_logger(__name__)
@@ -22,8 +23,8 @@ class DecodingStage(PipelineStage):
output format (e.g., pixel values).
"""
def __init__(self, vae) -> None:
self.vae = vae
def __init__(self, vae: ParallelTiledVAE) -> None:
self.vae: ParallelTiledVAE = vae
def forward(
self,
+3 -2
View File
@@ -15,6 +15,7 @@ from fastvideo.v1.models.vision_utils import (get_default_height_width,
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.utils import PRECISION_TO_TYPE
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
logger = init_logger(__name__)
@@ -27,8 +28,8 @@ class EncodingStage(PipelineStage):
input format (e.g., latents).
"""
def __init__(self, vae) -> None:
self.vae = vae
def __init__(self, vae: ParallelTiledVAE) -> None:
self.vae: ParallelTiledVAE = vae
def forward(
self,
@@ -6,7 +6,6 @@ from diffusers.utils.torch_utils import randn_tensor
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
@@ -21,10 +20,9 @@ class LatentPreparationStage(PipelineStage):
denoised during the diffusion process.
"""
def __init__(self, scheduler, vae=None) -> None:
def __init__(self, scheduler) -> None:
super().__init__()
self.scheduler = scheduler
self.vae = vae
def forward(
self,
@@ -44,7 +42,7 @@ class LatentPreparationStage(PipelineStage):
# Adjust video length based on VAE version if needed
if hasattr(self, 'adjust_video_length'):
batch = self.adjust_video_length(self.vae, batch, fastvideo_args)
batch = self.adjust_video_length(batch, fastvideo_args)
# Determine batch size
if isinstance(batch.prompt, list):
batch_size = len(batch.prompt)
@@ -70,15 +68,14 @@ class LatentPreparationStage(PipelineStage):
raise ValueError("Height and width must be provided")
assert fastvideo_args.num_channels_latents is not None
assert fastvideo_args.vae_scale_factor is not None
# Calculate latent shape
shape = (
batch_size,
fastvideo_args.num_channels_latents,
num_frames,
height // fastvideo_args.vae_scale_factor,
width // fastvideo_args.vae_scale_factor,
height // fastvideo_args.vae_config.arch_config.spatial_compression_ratio,
width // fastvideo_args.vae_config.arch_config.spatial_compression_ratio,
)
# Validate generator if it's a list
@@ -106,7 +103,7 @@ class LatentPreparationStage(PipelineStage):
return batch
def adjust_video_length(self, vae: ParallelTiledVAE, batch: ForwardBatch,
def adjust_video_length(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""
Adjust video length based on VAE version.
@@ -119,7 +116,7 @@ class LatentPreparationStage(PipelineStage):
The batch with adjusted video length.
"""
video_length = batch.num_frames
temporal_scale_factor = vae.temporal_compression_ratio if vae is not None else 4
temporal_scale_factor = fastvideo_args.vae_config.arch_config.temporal_compression_ratio
# TODO
batch.num_frames = (video_length - 1) // temporal_scale_factor + 1
return batch
@@ -54,8 +54,7 @@ class WanImageToVideoPipeline(ComposedPipelineBase):
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae")))
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="image_latent_preparation_stage",
stage=EncodingStage(vae=self.get_module("vae")))
@@ -72,9 +71,6 @@ class WanImageToVideoPipeline(ComposedPipelineBase):
"""
Initialize the pipeline.
"""
vae_scale_factor = self.get_module("vae").spatial_compression_ratio
fastvideo_args.vae_scale_factor = vae_scale_factor
num_channels_latents = self.get_module("transformer").out_channels
fastvideo_args.num_channels_latents = num_channels_latents
+1 -5
View File
@@ -47,8 +47,7 @@ class WanPipeline(ComposedPipelineBase):
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae")))
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
@@ -62,9 +61,6 @@ class WanPipeline(ComposedPipelineBase):
"""
Initialize the pipeline.
"""
vae_scale_factor = self.get_module("vae").spatial_compression_ratio
fastvideo_args.vae_scale_factor = vae_scale_factor
num_channels_latents = self.get_module("transformer").in_channels
fastvideo_args.num_channels_latents = num_channels_latents
+12 -16
View File
@@ -8,8 +8,11 @@ import torch
from safetensors.torch import load_file
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vaes.hunyuanvae import (
AutoencoderKLHunyuanVideo as MyHunyuanVAE)
# from fastvideo.v1.models.vaes.hunyuanvae import (
# AutoencoderKLHunyuanVideo as MyHunyuanVAE)
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.models.loader.component_loader import VAELoader
from fastvideo.v1.configs.models.vaes import HunyuanVAEConfig
from fastvideo.v1.utils import maybe_download_model
logger = init_logger(__name__)
@@ -31,21 +34,14 @@ REFERENCE_LATENT = -105.51324462890625
@pytest.mark.usefixtures("distributed_setup")
def test_hunyuan_vae():
device = torch.device("cuda:0")
# Initialize the two model implementations
config = json.load(open(CONFIG_PATH))
config.pop("_class_name")
config.pop("_diffusers_version")
model = MyHunyuanVAE(**config).to(torch.bfloat16)
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=VAE_PATH, vae_precision=precision_str)
args.device = device
args.vae_config = HunyuanVAEConfig()
loaded = load_file(os.path.join(VAE_PATH,
"diffusion_pytorch_model.safetensors"))
model.load_state_dict(loaded)
# Set model to eval mode
model.eval()
# Move to GPU
model = model.to(device)
loader = VAELoader()
model = loader.load(VAE_PATH, "", args)
model.enable_tiling(tile_sample_min_height=32,
tile_sample_min_width=32,
+9 -11
View File
@@ -9,6 +9,7 @@ from diffusers import AutoencoderKLWan
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import VAELoader
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo.v1.utils import maybe_download_model
logger = init_logger(__name__)
@@ -30,6 +31,7 @@ def test_wan_vae():
precision_str = "bf16"
args = FastVideoArgs(model_path=VAE_PATH, vae_precision=precision_str)
args.device = device
args.vae_config = WanVAEConfig()
loader = VAELoader()
model2 = loader.load(VAE_PATH, "", args)
@@ -77,23 +79,19 @@ def test_wan_vae():
# Test decoding
logger.info("Testing decoding...")
latent1_tensor = latent1.mode()
latents_mean = (torch.tensor(model1.config.latents_mean).view(
mean1 = (torch.tensor(model1.config.latents_mean).view(
1, model1.config.z_dim, 1, 1, 1).to(input_tensor.device,
input_tensor.dtype))
latents_std = 1.0 / torch.tensor(model1.config.latents_std).view(
1, model1.config.z_dim, 1, 1, 1).to(input_tensor.device,
std1 = (1.0 / torch.tensor(model1.config.latents_std).view(
1, model1.config.z_dim, 1, 1, 1)).to(input_tensor.device,
input_tensor.dtype)
latent1_tensor = latent1_tensor / latents_std + latents_mean
latent1_tensor = latent1_tensor / std1 + mean1
output1 = model1.decode(latent1_tensor).sample
mean2 = model2.config.arch_config.shift_factor.to(input_tensor.device, input_tensor.dtype)
std2 = model2.config.arch_config.scaling_factor.to(input_tensor.device, input_tensor.dtype)
latent2_tensor = latent2.mode()
latents_mean = (torch.tensor(model2.config.latents_mean).view(
1, model2.config.z_dim, 1, 1, 1).to(input_tensor.device,
input_tensor.dtype))
latents_std = 1.0 / torch.tensor(model2.config.latents_std).view(
1, model2.config.z_dim, 1, 1, 1).to(input_tensor.device,
input_tensor.dtype)
latent2_tensor = latent2_tensor / latents_std + latents_mean
latent2_tensor = latent2_tensor / std2 + mean2
output2 = model2.decode(latent2_tensor)
# Check if outputs have the same shape
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
+5 -11
View File
@@ -13,7 +13,7 @@ import signal
import sys
import tempfile
import traceback
from dataclasses import asdict, fields
from dataclasses import asdict, fields, is_dataclass
from functools import partial, wraps
from typing import (Any, Callable, Dict, List, Optional, Tuple, Type, TypeVar,
Union, cast)
@@ -550,16 +550,10 @@ def run_method(obj: Any, method: Union[str, bytes, Callable], args: tuple[Any],
func = partial(method, obj) # type: ignore
return func(*args, **kwargs)
def diff_keys(a, b):
return [k for k in asdict(a) if asdict(a)[k] != asdict(b)[k]]
def update_in_place(target, source, ignore_fields=()) -> None:
for f in fields(target):
if hasattr(source, f.name) and f.name not in list(ignore_fields):
setattr(target, f.name, getattr(source, f.name))
def shallow_asdict(obj):
if not is_dataclass(obj):
raise TypeError("Expected dataclass instance")
return {f.name: getattr(obj, f.name) for f in fields(obj)}
def kill_itself_when_parent_died() -> None:
# if sys.platform == "linux":