From f5182acd1c9d252926173b5880ea00f85d5175c5 Mon Sep 17 00:00:00 2001 From: Kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 12 Feb 2024 19:17:16 +0200 Subject: [PATCH] rework --- animatediff/utils/util.py | 2 +- configs/ad_unet_config.yaml | 73 +++++++++++ configs/text_encoder_config.json | 25 ++++ configs/tokenizer_config.json | 34 +++++ configs/v1-inference.yaml | 70 ++++++++++ nodes.py | 207 +++++++++++++++++++++++------- temp.py | 213 +++++++++++++++++++++++++++++++ 7 files changed, 579 insertions(+), 45 deletions(-) create mode 100644 configs/ad_unet_config.yaml create mode 100755 configs/text_encoder_config.json create mode 100755 configs/tokenizer_config.json create mode 100644 configs/v1-inference.yaml create mode 100644 temp.py diff --git a/animatediff/utils/util.py b/animatediff/utils/util.py index 08da96c..892d334 100644 --- a/animatediff/utils/util.py +++ b/animatediff/utils/util.py @@ -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) diff --git a/configs/ad_unet_config.yaml b/configs/ad_unet_config.yaml new file mode 100644 index 0000000..86c0138 --- /dev/null +++ b/configs/ad_unet_config.yaml @@ -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' \ No newline at end of file diff --git a/configs/text_encoder_config.json b/configs/text_encoder_config.json new file mode 100755 index 0000000..4d3e873 --- /dev/null +++ b/configs/text_encoder_config.json @@ -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 +} diff --git a/configs/tokenizer_config.json b/configs/tokenizer_config.json new file mode 100755 index 0000000..5ba7bf7 --- /dev/null +++ b/configs/tokenizer_config.json @@ -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 + } +} diff --git a/configs/v1-inference.yaml b/configs/v1-inference.yaml new file mode 100644 index 0000000..d4effe5 --- /dev/null +++ b/configs/v1-inference.yaml @@ -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 diff --git a/nodes.py b/nodes.py index 23fdc07..fab8b8d 100644 --- a/nodes.py +++ b/nodes.py @@ -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" } \ No newline at end of file diff --git a/temp.py b/temp.py new file mode 100644 index 0000000..e866100 --- /dev/null +++ b/temp.py @@ -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)) +)