This commit is contained in:
Kijai
2024-02-12 19:17:16 +02:00
parent b35bf9d989
commit f5182acd1c
7 changed files with 579 additions and 45 deletions
+1 -1
View File
@@ -127,7 +127,7 @@ def load_weights(
dreambooth_state_dict[key] = f.get_tensor(key)
elif dreambooth_model_path.endswith(".ckpt"):
dreambooth_state_dict = torch.load(dreambooth_model_path, map_location="cpu")
# 1. vae
converted_vae_checkpoint = convert_ldm_vae_checkpoint(dreambooth_state_dict, animation_pipeline.vae.config)
animation_pipeline.vae.load_state_dict(converted_vae_checkpoint, strict=False)
+73
View File
@@ -0,0 +1,73 @@
sample_size: 64
in_channels: 4
out_channels: 4
center_input_sample: false
flip_sin_to_cos: true
freq_shift: 0
down_block_types:
- CrossAttnDownBlock3D
- CrossAttnDownBlock3D
- CrossAttnDownBlock3D
- DownBlock3D
mid_block_type: UNetMidBlock3DCrossAttn
up_block_types:
- UpBlock3D
- CrossAttnUpBlock3D
- CrossAttnUpBlock3D
- CrossAttnUpBlock3D
only_cross_attention: false
block_out_channels:
- 320
- 640
- 1280
- 1280
layers_per_block: 2
downsample_padding: 1
mid_block_scale_factor: 1
act_fn: silu
norm_num_groups: 32
norm_eps: 1e-05
cross_attention_dim: 768
attention_head_dim: 8
dual_cross_attention: false
use_linear_projection: false
class_embed_type: null
num_class_embeds: null
upcast_attention: false
resnet_time_scale_shift: default
use_inflated_groupnorm: true
use_motion_module: true
motion_module_resolutions:
- 1
- 2
- 4
- 8
motion_module_mid_block: false
motion_module_decoder_only: false
motion_module_type: Vanilla
motion_module_kwargs:
num_attention_heads: 8
num_transformer_block: 1
attention_block_types:
- Temporal_Self
- Temporal_Self
temporal_position_encoding: true
temporal_position_encoding_max_len: 32
temporal_attention_dim_div: 1
zero_initialize: true
unet_use_cross_frame_attention: false
unet_use_temporal_attention: false
_use_default_values:
- resnet_time_scale_shift
- only_cross_attention
- mid_block_type
- unet_use_cross_frame_attention
- class_embed_type
- unet_use_temporal_attention
- dual_cross_attention
- num_class_embeds
- upcast_attention
- use_linear_projection
- motion_module_decoder_only
_class_name: UNet3DConditionModel
_diffusers_version: '0.6.0'
+25
View File
@@ -0,0 +1,25 @@
{
"_name_or_path": "openai/clip-vit-large-patch14",
"architectures": [
"CLIPTextModel"
],
"attention_dropout": 0.0,
"bos_token_id": 0,
"dropout": 0.0,
"eos_token_id": 2,
"hidden_act": "quick_gelu",
"hidden_size": 768,
"initializer_factor": 1.0,
"initializer_range": 0.02,
"intermediate_size": 3072,
"layer_norm_eps": 1e-05,
"max_position_embeddings": 77,
"model_type": "clip_text_model",
"num_attention_heads": 12,
"num_hidden_layers": 12,
"pad_token_id": 1,
"projection_dim": 768,
"torch_dtype": "float32",
"transformers_version": "4.22.0.dev0",
"vocab_size": 49408
}
+34
View File
@@ -0,0 +1,34 @@
{
"add_prefix_space": false,
"bos_token": {
"__type": "AddedToken",
"content": "<|startoftext|>",
"lstrip": false,
"normalized": true,
"rstrip": false,
"single_word": false
},
"do_lower_case": true,
"eos_token": {
"__type": "AddedToken",
"content": "<|endoftext|>",
"lstrip": false,
"normalized": true,
"rstrip": false,
"single_word": false
},
"errors": "replace",
"model_max_length": 77,
"name_or_path": "openai/clip-vit-large-patch14",
"pad_token": "<|endoftext|>",
"special_tokens_map_file": "./special_tokens_map.json",
"tokenizer_class": "CLIPTokenizer",
"unk_token": {
"__type": "AddedToken",
"content": "<|endoftext|>",
"lstrip": false,
"normalized": true,
"rstrip": false,
"single_word": false
}
}
+70
View File
@@ -0,0 +1,70 @@
model:
base_learning_rate: 1.0e-04
target: ldm.models.diffusion.ddpm.LatentDiffusion
params:
linear_start: 0.00085
linear_end: 0.0120
num_timesteps_cond: 1
log_every_t: 200
timesteps: 1000
first_stage_key: "jpg"
cond_stage_key: "txt"
image_size: 64
channels: 4
cond_stage_trainable: false # Note: different from the one we trained before
conditioning_key: crossattn
monitor: val/loss_simple_ema
scale_factor: 0.18215
use_ema: False
scheduler_config: # 10000 warmup steps
target: ldm.lr_scheduler.LambdaLinearScheduler
params:
warm_up_steps: [ 10000 ]
cycle_lengths: [ 10000000000000 ] # incredibly large number to prevent corner cases
f_start: [ 1.e-6 ]
f_max: [ 1. ]
f_min: [ 1. ]
unet_config:
target: ldm.modules.diffusionmodules.openaimodel.UNetModel
params:
image_size: 32 # unused
in_channels: 4
out_channels: 4
model_channels: 320
attention_resolutions: [ 4, 2, 1 ]
num_res_blocks: 2
channel_mult: [ 1, 2, 4, 4 ]
num_heads: 8
use_spatial_transformer: True
transformer_depth: 1
context_dim: 768
use_checkpoint: True
legacy: False
first_stage_config:
target: ldm.models.autoencoder.AutoencoderKL
params:
embed_dim: 4
monitor: val/rec_loss
ddconfig:
double_z: true
z_channels: 4
resolution: 256
in_channels: 3
out_ch: 3
ch: 128
ch_mult:
- 1
- 2
- 4
- 4
num_res_blocks: 2
attn_resolutions: []
dropout: 0.0
lossconfig:
target: torch.nn.Identity
cond_stage_config:
target: ldm.modules.encoders.modules.FrozenCLIPEmbedder
+163 -44
View File
@@ -468,37 +468,44 @@ class DiffusersLoaderForTraining:
path = os.path.join(search_path, model_path)
if os.path.exists(path):
model_path = path
break
break
config = OmegaConf.load(os.path.join(script_directory, f"configs/training/motion_director/training.yaml"))
vae = AutoencoderKL.from_pretrained(model_path, subfolder="vae")
tokenizer = CLIPTokenizer.from_pretrained(model_path, subfolder="tokenizer")
text_encoder = CLIPTextModel.from_pretrained(model_path, subfolder="text_encoder")
unet_additional_kwargs = config.unet_additional_kwargs
unet = UNet3DConditionModel.from_pretrained_2d(
model_path, subfolder="unet",
unet_additional_kwargs=unet_additional_kwargs
)
# Load scheduler, tokenizer and models.
noise_scheduler_kwargs = config.noise_scheduler_kwargs
noise_scheduler_kwargs.update({"steps_offset": 1})
noise_scheduler = DDIMScheduler(**OmegaConf.to_container(noise_scheduler_kwargs))
del noise_scheduler_kwargs["steps_offset"]
noise_scheduler_kwargs = {
'num_train_timesteps': 1000,
'beta_start': 0.00085,
'beta_end': 0.012,
'beta_schedule': "linear",
'clip_sample': False,
'steps_offset': 1
}
# Determine the scheduler class based on the scheduler variable
SchedulerClass = DDPMScheduler if scheduler == "DDPMScheduler" else DDIMScheduler
print(f"using {SchedulerClass.__name__} for training")
if scheduler == "DDPMScheduler":
print("using DDPMScheduler for training")
noise_scheduler_kwargs['beta_schedule'] = 'scaled_linear'
train_noise_scheduler_spatial = DDPMScheduler(**OmegaConf.to_container(noise_scheduler_kwargs))
noise_scheduler_kwargs['beta_schedule'] = 'linear'
train_noise_scheduler = DDPMScheduler(**OmegaConf.to_container(noise_scheduler_kwargs))
else:
print("using DDIMScheduler for training")
noise_scheduler_kwargs['beta_schedule'] = 'scaled_linear'
train_noise_scheduler_spatial = DDIMScheduler(**OmegaConf.to_container(noise_scheduler_kwargs))
noise_scheduler_kwargs['beta_schedule'] = 'linear'
train_noise_scheduler = DDIMScheduler(**OmegaConf.to_container(noise_scheduler_kwargs))
# Set the beta_schedule and create the default noise scheduler
noise_scheduler_kwargs['beta_schedule'] = 'linear'
noise_scheduler = SchedulerClass(**noise_scheduler_kwargs)
# Set the beta_schedule for the spatial noise scheduler and create it
noise_scheduler_kwargs['beta_schedule'] = 'scaled_linear'
train_noise_scheduler_spatial = SchedulerClass(**noise_scheduler_kwargs)
# Reset the beta_schedule for the linear noise scheduler and create it
noise_scheduler_kwargs['beta_schedule'] = 'linear'
train_noise_scheduler = SchedulerClass(**noise_scheduler_kwargs)
# Freeze all models for LoRA training
unet.requires_grad_(False)
@@ -512,24 +519,18 @@ class DiffusersLoaderForTraining:
# Enable gradient checkpointing
unet.enable_gradient_checkpointing()
# Move models to GPU
vae.to(device)
text_encoder.to(device)
unet.to(device=device)
text_encoder.to(device=device)
# Validation pipeline
validation_pipeline = AnimationPipeline(
unet=unet, vae=vae, tokenizer=tokenizer, text_encoder=text_encoder, scheduler=noise_scheduler,
).to(device)
motion_module_path, domain_adapter_path, unet_checkpoint_path = validation_models
motion_module_path, domain_adapter_path = validation_models
validation_pipeline = load_weights(
validation_pipeline,
motion_module_path=motion_module_path,
adapter_lora_path=domain_adapter_path,
dreambooth_model_path=unet_checkpoint_path
dreambooth_model_path=""
)
validation_pipeline.enable_vae_slicing()
@@ -546,19 +547,140 @@ class DiffusersLoaderForTraining:
return (pipeline,)
class CheckpointLoaderForTraining:
#@classmethod
#def IS_CHANGED(s):
# return ""
@classmethod
def INPUT_TYPES(cls):
class ValidationModelSelect:
return {"required":
{
"validation_models": ("VALIDATION_MODELS", ),
"ckpt_name": (folder_paths.get_filename_list("checkpoints"), ),
"scheduler": (
[
'DDIMScheduler',
'DDPMScheduler',
], {
"default": 'DDIMScheduler'
}),
"use_xformers": ("BOOLEAN", {"default": False}),
},
}
RETURN_TYPES = ("PIPELINE",)
FUNCTION = "load_checkpoint"
CATEGORY = "AD_MotionDirector"
def load_checkpoint(self, scheduler, use_xformers, validation_models, ckpt_name):
with torch.inference_mode(False):
model_path = folder_paths.get_full_path("checkpoints", ckpt_name)
original_config = OmegaConf.load(os.path.join(script_directory, f"configs/v1-inference.yaml"))
ad_unet_config = OmegaConf.load(os.path.join(script_directory, f"configs/ad_unet_config.yaml"))
from diffusers.loaders.single_file_utils import (convert_ldm_vae_checkpoint, convert_ldm_unet_checkpoint, create_text_encoder_from_ldm_clip_checkpoint, create_vae_diffusers_config, create_unet_diffusers_config)
from safetensors import safe_open
if model_path.endswith(".safetensors"):
dreambooth_state_dict = {}
with safe_open(model_path, framework="pt", device="cpu") as f:
for key in f.keys():
dreambooth_state_dict[key] = f.get_tensor(key)
elif model_path.endswith(".ckpt"):
dreambooth_state_dict = torch.load(model_path, map_location="cpu")
tokenizer = CLIPTokenizer.from_pretrained("openai/clip-vit-large-patch14")
text_encoder = create_text_encoder_from_ldm_clip_checkpoint("openai/clip-vit-large-patch14",dreambooth_state_dict)
noise_scheduler_kwargs = {
'num_train_timesteps': 1000,
'beta_start': 0.00085,
'beta_end': 0.012,
'beta_schedule': "linear",
'clip_sample': False,
'steps_offset': 1
}
#Determine the scheduler class based on the scheduler variable
SchedulerClass = DDPMScheduler if scheduler == "DDPMScheduler" else DDIMScheduler
print(f"using {SchedulerClass.__name__} for training")
# Set the beta_schedule and create the default noise scheduler
noise_scheduler_kwargs['beta_schedule'] = 'linear'
noise_scheduler = SchedulerClass(**noise_scheduler_kwargs)
# Set the beta_schedule for the spatial noise scheduler and create it
noise_scheduler_kwargs['beta_schedule'] = 'scaled_linear'
train_noise_scheduler_spatial = SchedulerClass(**noise_scheduler_kwargs)
# Reset the beta_schedule for the linear noise scheduler and create it
noise_scheduler_kwargs['beta_schedule'] = 'linear'
train_noise_scheduler = SchedulerClass(**noise_scheduler_kwargs)
# 1. vae
converted_vae_config = create_vae_diffusers_config(original_config, image_size=512)
converted_vae = convert_ldm_vae_checkpoint(dreambooth_state_dict, converted_vae_config)
vae = AutoencoderKL(**converted_vae_config)
vae.load_state_dict(converted_vae, strict=False)
# 2. unet
converted_unet_config = create_unet_diffusers_config(original_config, image_size=512)
converted_unet = convert_ldm_unet_checkpoint(dreambooth_state_dict, converted_unet_config)
unet = UNet3DConditionModel(**ad_unet_config)
unet.load_state_dict(converted_unet, strict=False)
del dreambooth_state_dict
# Validation pipeline
validation_pipeline = AnimationPipeline(
unet=unet, vae=vae, tokenizer=tokenizer, text_encoder=text_encoder, scheduler=noise_scheduler,
)
# Freeze all models for LoRA training
unet.requires_grad_(False)
vae.requires_grad_(False)
text_encoder.requires_grad_(False)
#xformers
if use_xformers:
unet.enable_xformers_memory_efficient_attention()
# Enable gradient checkpointing
unet.enable_gradient_checkpointing()
motion_module_path, domain_adapter_path = validation_models
validation_pipeline = load_weights(
validation_pipeline,
motion_module_path=motion_module_path,
adapter_lora_path=domain_adapter_path,
dreambooth_model_path=""
)
validation_pipeline.enable_vae_slicing()
pipeline = {
'validation_pipeline': validation_pipeline,
'train_noise_scheduler': train_noise_scheduler,
'train_noise_scheduler_spatial': train_noise_scheduler_spatial,
'unet': unet,
'vae': vae,
'text_encoder': text_encoder,
'tokenizer': tokenizer
}
return (pipeline,)
class AdditionalModelSelect:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"motion_module": (folder_paths.get_filename_list("animatediff_models"),),
"use_adapter_lora": ("BOOLEAN", {"default": True}),
"use_dreambooth_model": ("BOOLEAN", {"default": False}),
"use_adapter_lora": ("BOOLEAN", {"default": True}),
},
"optional": {
"optional_adapter_lora": (folder_paths.get_filename_list("loras"),),
"optional_model": (folder_paths.get_filename_list("checkpoints"),),
"optional_adapter_lora": (folder_paths.get_filename_list("loras"),),
}
}
RETURN_TYPES = ("VALIDATION_MODELS",)
@@ -567,7 +689,7 @@ class ValidationModelSelect:
CATEGORY = "AD_MotionDirector"
def select_models(self, motion_module, use_adapter_lora, use_dreambooth_model, optional_adapter_lora="", optional_model=""):
def select_models(self, motion_module, use_adapter_lora, optional_adapter_lora=""):
validation_models = []
motion_module_path = folder_paths.get_full_path("animatediff_models", motion_module)
@@ -575,14 +697,9 @@ class ValidationModelSelect:
adapter_lora_path = folder_paths.get_full_path("loras", optional_adapter_lora)
else:
adapter_lora_path = ""
if use_dreambooth_model:
model_path = folder_paths.get_full_path("checkpoints", optional_model)
else:
model_path = ""
validation_models.append(motion_module_path)
validation_models.append(adapter_lora_path)
validation_models.append(model_path)
validation_models.append(adapter_lora_path)
return (validation_models,)
class ValidationSettings:
@@ -893,18 +1010,20 @@ class TrainMotionDirectorLora:
NODE_CLASS_MAPPINGS = {
"AD_MotionDirector_train": AD_MotionDirector_train,
"DiffusersLoaderForTraining": DiffusersLoaderForTraining,
"ValidationModelSelect": ValidationModelSelect,
"AdditionalModelSelect": AdditionalModelSelect,
"ValidationSettings": ValidationSettings,
"AD_MotionLoraLoader": AD_MotionLoraLoader,
"SaveMotionDirectorLora": SaveMotionDirectorLora,
"TrainMotionDirectorLora": TrainMotionDirectorLora
"TrainMotionDirectorLora": TrainMotionDirectorLora,
"CheckpointLoaderForTraining": CheckpointLoaderForTraining
}
NODE_DISPLAY_NAME_MAPPINGS = {
"AD_MotionDirector_train": "AD_MotionDirector_train",
"DiffusersLoaderForTraining": "DiffusersLoaderForTraining",
"ValidationModelSelect": "ValidationModelSelect",
"AdditionalModelSelect": "AdditionalModelSelect",
"ValidationSettings": "ValidationSettings",
"AD_MotionLoraLoader": "AD_MotionLoraLoader",
"SaveMotionDirectorLora": "SaveMotionDirectorLora",
"TrainMotionDirectorLora": "TrainMotionDirectorLora"
"TrainMotionDirectorLora": "TrainMotionDirectorLora",
"CheckpointLoaderForTraining": "CheckpointLoaderForTraining"
}
+213
View File
@@ -0,0 +1,213 @@
AutoencoderKL(
(encoder): Encoder(
(conv_in): Conv2d(3, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(down_blocks): ModuleList(
(0): DownEncoderBlock2D(
(resnets): ModuleList(
(0-1): 2 x ResnetBlock2D(
(norm1): GroupNorm(32, 128, eps=1e-06, affine=True)
(conv1): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(norm2): GroupNorm(32, 128, eps=1e-06, affine=True)
(dropout): Dropout(p=0.0, inplace=False)
(conv2): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(nonlinearity): SiLU()
)
)
(downsamplers): ModuleList(
(0): Downsample2D(
(conv): Conv2d(128, 128, kernel_size=(3, 3), stride=(2, 2))
)
)
)
(1): DownEncoderBlock2D(
(resnets): ModuleList(
(0): ResnetBlock2D(
(norm1): GroupNorm(32, 128, eps=1e-06, affine=True)
(conv1): Conv2d(128, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(norm2): GroupNorm(32, 256, eps=1e-06, affine=True)
(dropout): Dropout(p=0.0, inplace=False)
(conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(nonlinearity): SiLU()
(conv_shortcut): Conv2d(128, 256, kernel_size=(1, 1), stride=(1, 1))
)
(1): ResnetBlock2D(
(norm1): GroupNorm(32, 256, eps=1e-06, affine=True)
(conv1): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(norm2): GroupNorm(32, 256, eps=1e-06, affine=True)
(dropout): Dropout(p=0.0, inplace=False)
(conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(nonlinearity): SiLU()
)
)
(downsamplers): ModuleList(
(0): Downsample2D(
(conv): Conv2d(256, 256, kernel_size=(3, 3), stride=(2, 2))
)
)
)
(2): DownEncoderBlock2D(
(resnets): ModuleList(
(0): ResnetBlock2D(
(norm1): GroupNorm(32, 256, eps=1e-06, affine=True)
(conv1): Conv2d(256, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(norm2): GroupNorm(32, 512, eps=1e-06, affine=True)
(dropout): Dropout(p=0.0, inplace=False)
(conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(nonlinearity): SiLU()
(conv_shortcut): Conv2d(256, 512, kernel_size=(1, 1), stride=(1, 1))
)
(1): ResnetBlock2D(
(norm1): GroupNorm(32, 512, eps=1e-06, affine=True)
(conv1): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(norm2): GroupNorm(32, 512, eps=1e-06, affine=True)
(dropout): Dropout(p=0.0, inplace=False)
(conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(nonlinearity): SiLU()
)
)
(downsamplers): ModuleList(
(0): Downsample2D(
(conv): Conv2d(512, 512, kernel_size=(3, 3), stride=(2, 2))
)
)
)
(3): DownEncoderBlock2D(
(resnets): ModuleList(
(0-1): 2 x ResnetBlock2D(
(norm1): GroupNorm(32, 512, eps=1e-06, affine=True)
(conv1): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(norm2): GroupNorm(32, 512, eps=1e-06, affine=True)
(dropout): Dropout(p=0.0, inplace=False)
(conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(nonlinearity): SiLU()
)
)
)
)
(mid_block): UNetMidBlock2D(
(attentions): ModuleList(
(0): Attention(
(group_norm): GroupNorm(32, 512, eps=1e-06, affine=True)
(to_q): Linear(in_features=512, out_features=512, bias=True)
(to_k): Linear(in_features=512, out_features=512, bias=True)
(to_v): Linear(in_features=512, out_features=512, bias=True)
(to_out): ModuleList(
(0): Linear(in_features=512, out_features=512, bias=True)
(1): Dropout(p=0.0, inplace=False)
)
)
)
(resnets): ModuleList(
(0-1): 2 x ResnetBlock2D(
(norm1): GroupNorm(32, 512, eps=1e-06, affine=True)
(conv1): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(norm2): GroupNorm(32, 512, eps=1e-06, affine=True)
(dropout): Dropout(p=0.0, inplace=False)
(conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(nonlinearity): SiLU()
)
)
)
(conv_norm_out): GroupNorm(32, 512, eps=1e-06, affine=True)
(conv_act): SiLU()
(conv_out): Conv2d(512, 8, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
)
(decoder): Decoder(
(conv_in): Conv2d(4, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(up_blocks): ModuleList(
(0-1): 2 x UpDecoderBlock2D(
(resnets): ModuleList(
(0-2): 3 x ResnetBlock2D(
(norm1): GroupNorm(32, 512, eps=1e-06, affine=True)
(conv1): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(norm2): GroupNorm(32, 512, eps=1e-06, affine=True)
(dropout): Dropout(p=0.0, inplace=False)
(conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(nonlinearity): SiLU()
)
)
(upsamplers): ModuleList(
(0): Upsample2D(
(conv): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
)
)
)
(2): UpDecoderBlock2D(
(resnets): ModuleList(
(0): ResnetBlock2D(
(norm1): GroupNorm(32, 512, eps=1e-06, affine=True)
(conv1): Conv2d(512, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(norm2): GroupNorm(32, 256, eps=1e-06, affine=True)
(dropout): Dropout(p=0.0, inplace=False)
(conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(nonlinearity): SiLU()
(conv_shortcut): Conv2d(512, 256, kernel_size=(1, 1), stride=(1, 1))
)
(1-2): 2 x ResnetBlock2D(
(norm1): GroupNorm(32, 256, eps=1e-06, affine=True)
(conv1): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(norm2): GroupNorm(32, 256, eps=1e-06, affine=True)
(dropout): Dropout(p=0.0, inplace=False)
(conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(nonlinearity): SiLU()
)
)
(upsamplers): ModuleList(
(0): Upsample2D(
(conv): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
)
)
)
(3): UpDecoderBlock2D(
(resnets): ModuleList(
(0): ResnetBlock2D(
(norm1): GroupNorm(32, 256, eps=1e-06, affine=True)
(conv1): Conv2d(256, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(norm2): GroupNorm(32, 128, eps=1e-06, affine=True)
(dropout): Dropout(p=0.0, inplace=False)
(conv2): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(nonlinearity): SiLU()
(conv_shortcut): Conv2d(256, 128, kernel_size=(1, 1), stride=(1, 1))
)
(1-2): 2 x ResnetBlock2D(
(norm1): GroupNorm(32, 128, eps=1e-06, affine=True)
(conv1): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(norm2): GroupNorm(32, 128, eps=1e-06, affine=True)
(dropout): Dropout(p=0.0, inplace=False)
(conv2): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(nonlinearity): SiLU()
)
)
)
)
(mid_block): UNetMidBlock2D(
(attentions): ModuleList(
(0): Attention(
(group_norm): GroupNorm(32, 512, eps=1e-06, affine=True)
(to_q): Linear(in_features=512, out_features=512, bias=True)
(to_k): Linear(in_features=512, out_features=512, bias=True)
(to_v): Linear(in_features=512, out_features=512, bias=True)
(to_out): ModuleList(
(0): Linear(in_features=512, out_features=512, bias=True)
(1): Dropout(p=0.0, inplace=False)
)
)
)
(resnets): ModuleList(
(0-1): 2 x ResnetBlock2D(
(norm1): GroupNorm(32, 512, eps=1e-06, affine=True)
(conv1): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(norm2): GroupNorm(32, 512, eps=1e-06, affine=True)
(dropout): Dropout(p=0.0, inplace=False)
(conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
(nonlinearity): SiLU()
)
)
)
(conv_norm_out): GroupNorm(32, 128, eps=1e-06, affine=True)
(conv_act): SiLU()
(conv_out): Conv2d(128, 3, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
)
(quant_conv): Conv2d(8, 8, kernel_size=(1, 1), stride=(1, 1))
(post_quant_conv): Conv2d(4, 4, kernel_size=(1, 1), stride=(1, 1))
)