From 9ec31c864fab6c77a4be3f589c68c88270160a65 Mon Sep 17 00:00:00 2001 From: junsong Date: Sat, 30 Nov 2024 11:50:15 -0800 Subject: [PATCH] first run sucessfull with text encoder mask bug not fix; --- Gemma/nodes.py | 128 +++++ Sana/conf.py | 98 ++++ Sana/diffusers_convert.py | 223 +++++++++ Sana/loader.py | 100 ++++ Sana/lora.py | 146 ++++++ Sana/models/act.py | 59 +++ Sana/models/basic_modules.py | 361 +++++++++++++++ Sana/models/norms.py | 225 +++++++++ Sana/models/sana.py | 379 +++++++++++++++ Sana/models/sana_blocks.py | 798 ++++++++++++++++++++++++++++++++ Sana/models/sana_multi_scale.py | 374 +++++++++++++++ Sana/models/utils.py | 591 +++++++++++++++++++++++ Sana/nodes.py | 223 +++++++++ VAE/nodes.py | 32 ++ __init__.py | 9 + 15 files changed, 3746 insertions(+) create mode 100644 Gemma/nodes.py create mode 100644 Sana/conf.py create mode 100644 Sana/diffusers_convert.py create mode 100644 Sana/loader.py create mode 100644 Sana/lora.py create mode 100644 Sana/models/act.py create mode 100644 Sana/models/basic_modules.py create mode 100644 Sana/models/norms.py create mode 100644 Sana/models/sana.py create mode 100644 Sana/models/sana_blocks.py create mode 100644 Sana/models/sana_multi_scale.py create mode 100644 Sana/models/utils.py create mode 100644 Sana/nodes.py diff --git a/Gemma/nodes.py b/Gemma/nodes.py new file mode 100644 index 0000000..55a8a03 --- /dev/null +++ b/Gemma/nodes.py @@ -0,0 +1,128 @@ +import os +import torch +import folder_paths +from transformers import AutoTokenizer, AutoModelForCausalLM +from ..utils.dtype import string_to_dtype +from huggingface_hub import snapshot_download + + +# 初始化自定义文件夹路径 +os.makedirs( + os.path.join(folder_paths.models_dir, "text_encoders"), + exist_ok=True +) +folder_paths.folder_names_and_paths["text_encoders"] = ( + [ + os.path.join(folder_paths.models_dir, "text_encoders"), + *folder_paths.folder_names_and_paths.get("text_encoders", [[],set()])[0] + ], + folder_paths.supported_pt_extensions +) + +dtypes = [ + "default", + "auto (comfy)", + "BF16", + "FP32", + "FP16", +] +try: torch.float8_e5m2 +except AttributeError: print("Torch版本过旧,不支持FP8") +else: dtypes += ["FP8 E4M3", "FP8 E5M2"] + +class GemmaLoader: + @classmethod + def INPUT_TYPES(s): + devices = ["auto", "cpu", "cuda"] + # 支持多GPU + for k in range(1, torch.cuda.device_count()): + devices.append(f"cuda:{k}") + return { + "required": { + "model_name": (["google/gemma-2-2b-it", "unsloth/gemma-2-2b-it-bnb-4bit"],), + "device": (devices, {"default":"cpu"}), + "dtype": (dtypes,), + } + } + RETURN_TYPES = ("GEMMA",) + FUNCTION = "load_model" + CATEGORY = "ExtraModels/Gemma" + TITLE = "Gemma Loader" + + def load_model(self, model_name, device, dtype): + dtype = string_to_dtype(dtype, "text_encoder") + if device == "cpu": + assert dtype in [None, torch.float32], f"Can't use dtype '{dtype}' with CPU! Set dtype to 'default'." + + if model_name == 'google/gemma-2-2b-it': + text_encoder_dir = os.path.join(folder_paths.models_dir, 'text_encoders', 'models--google--gemma-2-2b-it') + if not os.path.exists(os.path.join(text_encoder_dir, 'model.safetensors')): + snapshot_download('google/gemma-2-2b-it', local_dir=text_encoder_dir) + elif model_name == 'unsloth/gemma-2-2b-it-bnb-4bit': + text_encoder_dir = os.path.join(folder_paths.models_dir, 'text_encoders', 'models--unsloth--gemma-2-2b-it-bnb-4bit') + if not os.path.exists(os.path.join(text_encoder_dir, 'model.safetensors')): + snapshot_download('unsloth/gemma-2-2b-it-bnb-4bit', local_dir=text_encoder_dir) + else: + raise ValueError('Not implemented!') + + tokenizer = AutoTokenizer.from_pretrained(model_name) + text_encoder_model = AutoModelForCausalLM.from_pretrained(model_name, torch_dtype=dtype) + tokenizer.padding_side = "right" + text_encoder = text_encoder_model.get_decoder() + + if device != "cpu": + text_encoder = text_encoder.to(device) + + return ({ + "tokenizer": tokenizer, + "text_encoder": text_encoder, + "text_encoder_model": text_encoder_model + },) + + +class GemmaTextEncode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "text": ("STRING", {"multiline": True}), + "GEMMA": ("GEMMA",), + } + } + + RETURN_TYPES = ("CONDITIONING",) + FUNCTION = "encode" + CATEGORY = "ExtraModels/Gemma" + TITLE = "Gemma Text Encode" + + def encode(self, text, GEMMA=None): + print(text) + tokenizer = GEMMA["tokenizer"] + text_encoder = GEMMA["text_encoder"] + + with torch.no_grad(): + tokens = tokenizer( + text, + max_length=300, + padding="max_length", + truncation=True, + return_tensors="pt" + ).to(text_encoder.device) + + cond = text_encoder(tokens.input_ids, tokens.attention_mask)[0][:, None] + emb_masks = tokens.attention_mask + + # 利用emb_masks将有效的cond选出来,其他置零 + # cond = cond * emb_masks.unsqueeze(-1) + + return ([[cond, {}]], ) + +NODE_CLASS_MAPPINGS = { + "GemmaLoader": GemmaLoader, + "GemmaTextEncode": GemmaTextEncode, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "GemmaLoader": "Gemma Loader", + "GemmaTextEncode": "Gemma Text Encode", +} diff --git a/Sana/conf.py b/Sana/conf.py new file mode 100644 index 0000000..7f5a046 --- /dev/null +++ b/Sana/conf.py @@ -0,0 +1,98 @@ +""" +List of all Sana model types / settings +""" + +sampling_settings = { + "shift": 3.0, +} + +sana_conf = { + "SanaMS_600M_P1_D28": { + "target": "SanaMS", + "unet_config": { + "in_channels": 32, + "depth": 28, + "hidden_size": 1152, + "patch_size": 1, + "num_heads": 36, + "linear_head_dim": 32, + "model_max_length": 300, + "y_norm": True, + "attn_type": "linear", + "ffn_type": "glumbconv", + "mlp_ratio": 2.5, + "mlp_acts": ["silu", "silu", None], + "use_pe": False, + "pred_sigma": False, + "learn_sigma": False, + "fp32_attention": True, + }, + "sampling_settings" : sampling_settings, + }, + "SanaMS_1600M_P1_D20": { + "target": "SanaMS", + "unet_config": { + "in_channels": 32, + "depth": 20, + "hidden_size": 2240, + "patch_size": 1, + "num_heads": 70, + "linear_head_dim": 32, + "model_max_length": 300, + "y_norm": True, + "attn_type": "linear", + "ffn_type": "glumbconv", + "mlp_ratio": 2.5, + "mlp_acts": ["silu", "silu", None], + "use_pe": False, + "pred_sigma": False, + "learn_sigma": False, + "fp32_attention": True, + }, + "sampling_settings" : sampling_settings, + }, +} + +sana_res = { + "1024px": { # models/SanaMS 1024x1024 + '0.25': [512, 2048], '0.26': [512, 1984], '0.27': [512, 1920], '0.28': [512, 1856], + '0.32': [576, 1792], '0.33': [576, 1728], '0.35': [576, 1664], '0.40': [640, 1600], + '0.42': [640, 1536], '0.48': [704, 1472], '0.50': [704, 1408], '0.52': [704, 1344], + '0.57': [768, 1344], '0.60': [768, 1280], '0.68': [832, 1216], '0.72': [832, 1152], + '0.78': [896, 1152], '0.82': [896, 1088], '0.88': [960, 1088], '0.94': [960, 1024], + '1.00': [1024,1024], '1.07': [1024, 960], '1.13': [1088, 960], '1.21': [1088, 896], + '1.29': [1152, 896], '1.38': [1152, 832], '1.46': [1216, 832], '1.67': [1280, 768], + '1.75': [1344, 768], '2.00': [1408, 704], '2.09': [1472, 704], '2.40': [1536, 640], + '2.50': [1600, 640], '2.89': [1664, 576], '3.00': [1728, 576], '3.11': [1792, 576], + '3.62': [1856, 512], '3.75': [1920, 512], '3.88': [1984, 512], '4.00': [2048, 512], + }, + "512px": { # models/SanaMS 512x512 + '0.25': [256,1024], '0.26': [256, 992], '0.27': [256, 960], '0.28': [256, 928], + '0.32': [288, 896], '0.33': [288, 864], '0.35': [288, 832], '0.40': [320, 800], + '0.42': [320, 768], '0.48': [352, 736], '0.50': [352, 704], '0.52': [352, 672], + '0.57': [384, 672], '0.60': [384, 640], '0.68': [416, 608], '0.72': [416, 576], + '0.78': [448, 576], '0.82': [448, 544], '0.88': [480, 544], '0.94': [480, 512], + '1.00': [512, 512], '1.07': [512, 480], '1.13': [544, 480], '1.21': [544, 448], + '1.29': [576, 448], '1.38': [576, 416], '1.46': [608, 416], '1.67': [640, 384], + '1.75': [672, 384], '2.00': [704, 352], '2.09': [736, 352], '2.40': [768, 320], + '2.50': [800, 320], '2.89': [832, 288], '3.00': [864, 288], '3.11': [896, 288], + '3.62': [928, 256], '3.75': [960, 256], '3.88': [992, 256], '4.00': [1024,256] + }, + "2K": { + '0.25': [1024, 4096], '0.26': [1024, 3968], '0.27': [1024, 3840], '0.28': [1024, 3712], + '0.32': [1152, 3584], '0.33': [1152, 3456], '0.35': [1152, 3328], '0.40': [1280, 3200], + '0.42': [1280, 3072], '0.48': [1408, 2944], '0.50': [1408, 2816], '0.52': [1408, 2688], + '0.57': [1536, 2688], '0.60': [1536, 2560], '0.68': [1664, 2432], '0.72': [1664, 2304], + '0.78': [1792, 2304], '0.82': [1792, 2176], '0.88': [1920, 2176], '0.94': [1920, 2048], + '1.00': [2048, 2048], '1.07': [2048, 1920], '1.13': [2176, 1920], '1.21': [2176, 1792], + '1.29': [2304, 1792], '1.38': [2304, 1664], '1.46': [2432, 1664], '1.67': [2560, 1536], + '1.75': [2688, 1536], '2.00': [2816, 1408], '2.09': [2944, 1408], '2.40': [3072, 1280], + '2.50': [3200, 1280], '2.89': [3328, 1152], '3.00': [3456, 1152], '3.11': [3584, 1152], + '3.62': [3712, 1024], '3.75': [3840, 1024], '3.88': [3968, 1024], '4.00': [4096, 1024] + } +} +# These should be the same +sana_res.update({ + "SanaMS_600M_P1_D28": sana_res["1024px"], + "SanaMS_1600M_P1_D20": sana_res["1024px"], +}) diff --git a/Sana/diffusers_convert.py b/Sana/diffusers_convert.py new file mode 100644 index 0000000..312ea9d --- /dev/null +++ b/Sana/diffusers_convert.py @@ -0,0 +1,223 @@ +# For using the diffusers format weights +# Based on the original ComfyUI function + +# https://github.com/PixArt-alpha/PixArt-alpha/blob/master/tools/convert_pixart_alpha_to_diffusers.py +import torch + +conversion_map_ms = [ # for multi_scale_train (MS) + # Resolution + ("csize_embedder.mlp.0.weight", "adaln_single.emb.resolution_embedder.linear_1.weight"), + ("csize_embedder.mlp.0.bias", "adaln_single.emb.resolution_embedder.linear_1.bias"), + ("csize_embedder.mlp.2.weight", "adaln_single.emb.resolution_embedder.linear_2.weight"), + ("csize_embedder.mlp.2.bias", "adaln_single.emb.resolution_embedder.linear_2.bias"), + # Aspect ratio + ("ar_embedder.mlp.0.weight", "adaln_single.emb.aspect_ratio_embedder.linear_1.weight"), + ("ar_embedder.mlp.0.bias", "adaln_single.emb.aspect_ratio_embedder.linear_1.bias"), + ("ar_embedder.mlp.2.weight", "adaln_single.emb.aspect_ratio_embedder.linear_2.weight"), + ("ar_embedder.mlp.2.bias", "adaln_single.emb.aspect_ratio_embedder.linear_2.bias"), +] + +def get_depth(state_dict): + return sum(key.endswith('.attn1.to_k.bias') for key in state_dict.keys()) + +def get_lora_depth(state_dict): + cnt = max([ + sum(key.endswith('.attn1.to_k.lora_A.weight') for key in state_dict.keys()), + sum(key.endswith('_attn1_to_k.lora_A.weight') for key in state_dict.keys()), + sum(key.endswith('.attn1.to_k.lora_up.weight') for key in state_dict.keys()), + sum(key.endswith('_attn1_to_k.lora_up.weight') for key in state_dict.keys()), + ]) + assert cnt > 0, "Unable to detect model depth!" + return cnt + +def get_conversion_map(state_dict): + conversion_map = [ # main SD conversion map (PixArt reference, HF Diffusers) + # Patch embeddings + ("x_embedder.proj.weight", "pos_embed.proj.weight"), + ("x_embedder.proj.bias", "pos_embed.proj.bias"), + # Caption projection + ("y_embedder.y_embedding", "caption_projection.y_embedding"), + ("y_embedder.y_proj.fc1.weight", "caption_projection.linear_1.weight"), + ("y_embedder.y_proj.fc1.bias", "caption_projection.linear_1.bias"), + ("y_embedder.y_proj.fc2.weight", "caption_projection.linear_2.weight"), + ("y_embedder.y_proj.fc2.bias", "caption_projection.linear_2.bias"), + # AdaLN-single LN + ("t_embedder.mlp.0.weight", "adaln_single.emb.timestep_embedder.linear_1.weight"), + ("t_embedder.mlp.0.bias", "adaln_single.emb.timestep_embedder.linear_1.bias"), + ("t_embedder.mlp.2.weight", "adaln_single.emb.timestep_embedder.linear_2.weight"), + ("t_embedder.mlp.2.bias", "adaln_single.emb.timestep_embedder.linear_2.bias"), + # Shared norm + ("t_block.1.weight", "adaln_single.linear.weight"), + ("t_block.1.bias", "adaln_single.linear.bias"), + # Final block + ("final_layer.linear.weight", "proj_out.weight"), + ("final_layer.linear.bias", "proj_out.bias"), + ("final_layer.scale_shift_table", "scale_shift_table"), + ] + + # Add actual transformer blocks + for depth in range(get_depth(state_dict)): + # Transformer blocks + conversion_map += [ + (f"blocks.{depth}.scale_shift_table", f"transformer_blocks.{depth}.scale_shift_table"), + # Projection + (f"blocks.{depth}.attn.proj.weight", f"transformer_blocks.{depth}.attn1.to_out.0.weight"), + (f"blocks.{depth}.attn.proj.bias", f"transformer_blocks.{depth}.attn1.to_out.0.bias"), + # Feed-forward + (f"blocks.{depth}.mlp.fc1.weight", f"transformer_blocks.{depth}.ff.net.0.proj.weight"), + (f"blocks.{depth}.mlp.fc1.bias", f"transformer_blocks.{depth}.ff.net.0.proj.bias"), + (f"blocks.{depth}.mlp.fc2.weight", f"transformer_blocks.{depth}.ff.net.2.weight"), + (f"blocks.{depth}.mlp.fc2.bias", f"transformer_blocks.{depth}.ff.net.2.bias"), + # Cross-attention (proj) + (f"blocks.{depth}.cross_attn.proj.weight" ,f"transformer_blocks.{depth}.attn2.to_out.0.weight"), + (f"blocks.{depth}.cross_attn.proj.bias" ,f"transformer_blocks.{depth}.attn2.to_out.0.bias"), + ] + return conversion_map + +def find_prefix(state_dict, target_key): + prefix = "" + for k in state_dict.keys(): + if k.endswith(target_key): + prefix = k.split(target_key)[0] + break + return prefix + +def convert_state_dict(state_dict): + if "adaln_single.emb.resolution_embedder.linear_1.weight" in state_dict.keys(): + cmap = get_conversion_map(state_dict) + conversion_map_ms + else: + cmap = get_conversion_map(state_dict) + + missing = [k for k,v in cmap if v not in state_dict] + new_state_dict = {k: state_dict[v] for k,v in cmap if k not in missing} + matched = list(v for k,v in cmap if v in state_dict.keys()) + + for depth in range(get_depth(state_dict)): + for wb in ["weight", "bias"]: + # Self Attention + key = lambda a: f"transformer_blocks.{depth}.attn1.to_{a}.{wb}" + new_state_dict[f"blocks.{depth}.attn.qkv.{wb}"] = torch.cat(( + state_dict[key('q')], state_dict[key('k')], state_dict[key('v')] + ), dim=0) + matched += [key('q'), key('k'), key('v')] + + # Cross-attention (linear) + key = lambda a: f"transformer_blocks.{depth}.attn2.to_{a}.{wb}" + new_state_dict[f"blocks.{depth}.cross_attn.q_linear.{wb}"] = state_dict[key('q')] + new_state_dict[f"blocks.{depth}.cross_attn.kv_linear.{wb}"] = torch.cat(( + state_dict[key('k')], state_dict[key('v')] + ), dim=0) + matched += [key('q'), key('k'), key('v')] + + if len(matched) < len(state_dict): + print(f"PixArt: UNET conversion has leftover keys! ({len(matched)} vs {len(state_dict)})") + print(list( set(state_dict.keys()) - set(matched) )) + + if len(missing) > 0: + print(f"PixArt: UNET conversion has missing keys!") + print(missing) + + return new_state_dict + +# Same as above but for LoRA weights: +def convert_lora_state_dict(state_dict, peft=True): + # koyha + rep_ak = lambda x: x.replace(".weight", ".lora_down.weight") + rep_bk = lambda x: x.replace(".weight", ".lora_up.weight") + rep_pk = lambda x: x.replace(".weight", ".alpha") + if peft: # peft + rep_ap = lambda x: x.replace(".weight", ".lora_A.weight") + rep_bp = lambda x: x.replace(".weight", ".lora_B.weight") + rep_pp = lambda x: x.replace(".weight", ".alpha") + + prefix = find_prefix(state_dict, "adaln_single.linear.lora_A.weight") + state_dict = {k[len(prefix):]:v for k,v in state_dict.items()} + else: # OneTrainer + rep_ap = lambda x: x.replace(".", "_")[:-7] + ".lora_down.weight" + rep_bp = lambda x: x.replace(".", "_")[:-7] + ".lora_up.weight" + rep_pp = lambda x: x.replace(".", "_")[:-7] + ".alpha" + + prefix = "lora_transformer_" + t5_marker = "lora_te_encoder" + t5_keys = [] + for key in list(state_dict.keys()): + if key.startswith(prefix): + state_dict[key[len(prefix):]] = state_dict.pop(key) + elif t5_marker in key: + t5_keys.append(state_dict.pop(key)) + if len(t5_keys) > 0: + print(f"Text Encoder not supported for PixArt LoRA, ignoring {len(t5_keys)} keys") + + cmap = [] + cmap_unet = get_conversion_map(state_dict) + conversion_map_ms # todo: 512 model + for k, v in cmap_unet: + if v.endswith(".weight"): + cmap.append((rep_ak(k), rep_ap(v))) + cmap.append((rep_bk(k), rep_bp(v))) + if not peft: + cmap.append((rep_pk(k), rep_pp(v))) + + missing = [k for k,v in cmap if v not in state_dict] + new_state_dict = {k: state_dict[v] for k,v in cmap if k not in missing} + matched = list(v for k,v in cmap if v in state_dict.keys()) + + lora_depth = get_lora_depth(state_dict) + for fp, fk in ((rep_ap, rep_ak),(rep_bp, rep_bk)): + for depth in range(lora_depth): + # Self Attention + key = lambda a: fp(f"transformer_blocks.{depth}.attn1.to_{a}.weight") + new_state_dict[fk(f"blocks.{depth}.attn.qkv.weight")] = torch.cat(( + state_dict[key('q')], state_dict[key('k')], state_dict[key('v')] + ), dim=0) + + matched += [key('q'), key('k'), key('v')] + if not peft: + akey = lambda a: rep_pp(f"transformer_blocks.{depth}.attn1.to_{a}.weight") + new_state_dict[rep_pk((f"blocks.{depth}.attn.qkv.weight"))] = state_dict[akey("q")] + matched += [akey('q'), akey('k'), akey('v')] + + # Self Attention projection? + key = lambda a: fp(f"transformer_blocks.{depth}.attn1.to_{a}.weight") + new_state_dict[fk(f"blocks.{depth}.attn.proj.weight")] = state_dict[key('out.0')] + matched += [key('out.0')] + + # Cross-attention (linear) + key = lambda a: fp(f"transformer_blocks.{depth}.attn2.to_{a}.weight") + new_state_dict[fk(f"blocks.{depth}.cross_attn.q_linear.weight")] = state_dict[key('q')] + new_state_dict[fk(f"blocks.{depth}.cross_attn.kv_linear.weight")] = torch.cat(( + state_dict[key('k')], state_dict[key('v')] + ), dim=0) + matched += [key('q'), key('k'), key('v')] + if not peft: + akey = lambda a: rep_pp(f"transformer_blocks.{depth}.attn2.to_{a}.weight") + new_state_dict[rep_pk((f"blocks.{depth}.cross_attn.q_linear.weight"))] = state_dict[akey("q")] + new_state_dict[rep_pk((f"blocks.{depth}.cross_attn.kv_linear.weight"))] = state_dict[akey("k")] + matched += [akey('q'), akey('k'), akey('v')] + + # Cross Attention projection? + key = lambda a: fp(f"transformer_blocks.{depth}.attn2.to_{a}.weight") + new_state_dict[fk(f"blocks.{depth}.cross_attn.proj.weight")] = state_dict[key('out.0')] + matched += [key('out.0')] + + try: + key = fp(f"transformer_blocks.{depth}.ff.net.0.proj.weight") + new_state_dict[fk(f"blocks.{depth}.mlp.fc1.weight")] = state_dict[key] + matched += [key] + except KeyError: + pass + + try: + key = fp(f"transformer_blocks.{depth}.ff.net.2.weight") + new_state_dict[fk(f"blocks.{depth}.mlp.fc2.weight")] = state_dict[key] + matched += [key] + except KeyError: + pass + + if len(matched) < len(state_dict): + print(f"PixArt: LoRA conversion has leftover keys! ({len(matched)} vs {len(state_dict)})") + print(list( set(state_dict.keys()) - set(matched) )) + + if len(missing) > 0: + print(f"PixArt: LoRA conversion has missing keys! (probably)") + print(missing) + + return new_state_dict diff --git a/Sana/loader.py b/Sana/loader.py new file mode 100644 index 0000000..34ba85f --- /dev/null +++ b/Sana/loader.py @@ -0,0 +1,100 @@ +import comfy.supported_models_base +import comfy.latent_formats +import comfy.model_patcher +import comfy.model_base +import comfy.utils +import comfy.conds +import torch +import math +from comfy import model_management +from comfy.latent_formats import LatentFormat +from .diffusers_convert import convert_state_dict + + +class SanaLatent(LatentFormat): + latent_channels = 32 + def __init__(self): + self.scale_factor = 0.41407 + + +class EXM_Sana(comfy.supported_models_base.BASE): + unet_config = {} + unet_extra_config = {} + latent_format = SanaLatent + + def __init__(self, model_conf): + self.model_target = model_conf.get("target") + self.unet_config = model_conf.get("unet_config", {}) + self.sampling_settings = model_conf.get("sampling_settings", {}) + self.latent_format = self.latent_format() + # UNET is handled by extension + self.unet_config["disable_unet_model_creation"] = True + + def model_type(self, state_dict, prefix=""): + return comfy.model_base.ModelType.FLOW + + +class EXM_Sana_Model(comfy.model_base.BaseModel): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + + def extra_conds(self, **kwargs): + out = super().extra_conds(**kwargs) + + cn_hint = kwargs.get("cn_hint", None) + if cn_hint is not None: + out["cn_hint"] = comfy.conds.CONDRegular(cn_hint) + + return out + + +def load_sana(model_path, model_conf, dtype): + state_dict = comfy.utils.load_torch_file(model_path) + state_dict = state_dict.get("model", state_dict) + + # prefix + for prefix in ["model.diffusion_model.",]: + if any(True for x in state_dict if x.startswith(prefix)): + state_dict = {k[len(prefix):]:v for k,v in state_dict.items()} + + # diffusers + if "adaln_single.linear.weight" in state_dict: + state_dict = convert_state_dict(state_dict) # Diffusers + + parameters = comfy.utils.calculate_parameters(state_dict) + unet_dtype = dtype + load_device = comfy.model_management.get_torch_device() + offload_device = comfy.model_management.unet_offload_device() + + # ignore fp8/etc and use directly for now + manual_cast_dtype = model_management.unet_manual_cast(unet_dtype, load_device) + if manual_cast_dtype: + print(f"Sana: falling back to {manual_cast_dtype}") + unet_dtype = manual_cast_dtype + + model_conf = EXM_Sana(model_conf) # convert to object + model = EXM_Sana_Model( # same as comfy.model_base.BaseModel + model_conf, + model_type=comfy.model_base.ModelType.FLOW, + device=model_management.get_torch_device() + ) + + if model_conf.model_target == "SanaMS": + from .models.sana_multi_scale import SanaMS + model.diffusion_model = SanaMS(**model_conf.unet_config) + else: + raise NotImplementedError(f"Unknown model target '{model_conf.model_target}'") + + m, u = model.diffusion_model.load_state_dict(state_dict, strict=False) + if len(m) > 0: print("Missing UNET keys", m) + if len(u) > 0: print("Leftover UNET keys", u) + model.diffusion_model.dtype = unet_dtype + model.diffusion_model.eval() + model.diffusion_model.to(unet_dtype) + + model_patcher = comfy.model_patcher.ModelPatcher( + model, + load_device = load_device, + offload_device = offload_device, + ) + return model_patcher diff --git a/Sana/lora.py b/Sana/lora.py new file mode 100644 index 0000000..fca5931 --- /dev/null +++ b/Sana/lora.py @@ -0,0 +1,146 @@ +import os +import copy +import json +import torch +import comfy.lora +import comfy.model_management +from comfy.model_patcher import ModelPatcher +from .diffusers_convert import convert_lora_state_dict + +class EXM_PixArt_ModelPatcher(ModelPatcher): + def calculate_weight(self, patches, weight, key): + """ + This is almost the same as the comfy function, but stripped down to just the LoRA patch code. + The problem with the original code is the q/k/v keys being combined into one for the attention. + In the diffusers code, they're treated as separate keys, but in the reference code they're recombined (q+kv|qkv). + This means, for example, that the [1152,1152] weights become [3456,1152] in the state dict. + The issue with this is that the LoRA weights are [128,1152],[1152,128] and become [384,1162],[3456,128] instead. + + This is the best thing I could think of that would fix that, but it's very fragile. + - Check key shape to determine if it needs the fallback logic + - Cut the input into parts based on the shape (undoing the torch.cat) + - Do the matrix multiplication logic + - Recombine them to match the expected shape + """ + for p in patches: + alpha = p[0] + v = p[1] + strength_model = p[2] + if strength_model != 1.0: + weight *= strength_model + + if isinstance(v, list): + v = (self.calculate_weight(v[1:], v[0].clone(), key), ) + + if len(v) == 2: + patch_type = v[0] + v = v[1] + + if patch_type == "lora": + mat1 = comfy.model_management.cast_to_device(v[0], weight.device, torch.float32) + mat2 = comfy.model_management.cast_to_device(v[1], weight.device, torch.float32) + if v[2] is not None: + alpha *= v[2] / mat2.shape[0] + try: + mat1 = mat1.flatten(start_dim=1) + mat2 = mat2.flatten(start_dim=1) + + ch1 = mat1.shape[0] // mat2.shape[1] + ch2 = mat2.shape[0] // mat1.shape[1] + ### Fallback logic for shape mismatch ### + if mat1.shape[0] != mat2.shape[1] and ch1 == ch2 and (mat1.shape[0]/mat2.shape[1])%1 == 0: + mat1 = mat1.chunk(ch1, dim=0) + mat2 = mat2.chunk(ch1, dim=0) + weight += torch.cat( + [alpha * torch.mm(mat1[x], mat2[x]) for x in range(ch1)], + dim=0, + ).reshape(weight.shape).type(weight.dtype) + else: + weight += (alpha * torch.mm(mat1, mat2)).reshape(weight.shape).type(weight.dtype) + except Exception as e: + print("ERROR", key, e) + return weight + + def clone(self): + n = EXM_PixArt_ModelPatcher(self.model, self.load_device, self.offload_device, self.size, self.current_device, weight_inplace_update=self.weight_inplace_update) + n.patches = {} + for k in self.patches: + n.patches[k] = self.patches[k][:] + + n.object_patches = self.object_patches.copy() + n.model_options = copy.deepcopy(self.model_options) + n.model_keys = self.model_keys + return n + +def replace_model_patcher(model): + n = EXM_PixArt_ModelPatcher( + model = model.model, + size = model.size, + load_device = model.load_device, + offload_device = model.offload_device, + weight_inplace_update = model.weight_inplace_update, + ) + n.patches = {} + for k in model.patches: + n.patches[k] = model.patches[k][:] + + n.object_patches = model.object_patches.copy() + n.model_options = copy.deepcopy(model.model_options) + return n + +def find_peft_alpha(path): + def load_json(json_path): + with open(json_path) as f: + data = json.load(f) + alpha = data.get("lora_alpha") + alpha = alpha or data.get("alpha") + if not alpha: + print(" Found config but `lora_alpha` is missing!") + else: + print(f" Found config at {json_path} [alpha:{alpha}]") + return alpha + + # For some weird reason peft doesn't include the alpha in the actual model + print("PixArt: Warning! This is a PEFT LoRA. Trying to find config...") + files = [ + f"{os.path.splitext(path)[0]}.json", + f"{os.path.splitext(path)[0]}.config.json", + os.path.join(os.path.dirname(path),"adapter_config.json"), + ] + for file in files: + if os.path.isfile(file): + return load_json(file) + + print(" Missing config/alpha! assuming alpha of 8. Consider converting it/adding a config json to it.") + return 8.0 + +def load_pixart_lora(model, lora, lora_path, strength): + k_back = lambda x: x.replace(".lora_up.weight", "") + # need to convert the actual weights for this to work. + if any(True for x in lora.keys() if x.endswith("adaln_single.linear.lora_A.weight")): + lora = convert_lora_state_dict(lora, peft=True) + alpha = find_peft_alpha(lora_path) + lora.update({f"{k_back(x)}.alpha":torch.tensor(alpha) for x in lora.keys() if "lora_up" in x}) + else: # OneTrainer + lora = convert_lora_state_dict(lora, peft=False) + + key_map = {k_back(x):f"diffusion_model.{k_back(x)}.weight" for x in lora.keys() if "lora_up" in x} # fake + + loaded = comfy.lora.load_lora(lora, key_map) + if model is not None: + # switch to custom model patcher when using LoRAs + if isinstance(model, EXM_PixArt_ModelPatcher): + new_modelpatcher = model.clone() + else: + new_modelpatcher = replace_model_patcher(model) + k = new_modelpatcher.add_patches(loaded, strength) + else: + k = () + new_modelpatcher = None + + k = set(k) + for x in loaded: + if (x not in k): + print("NOT LOADED", x) + + return new_modelpatcher diff --git a/Sana/models/act.py b/Sana/models/act.py new file mode 100644 index 0000000..9df6a7a --- /dev/null +++ b/Sana/models/act.py @@ -0,0 +1,59 @@ +# Copyright 2024 NVIDIA CORPORATION & AFFILIATES +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# SPDX-License-Identifier: Apache-2.0 + +import copy + +import torch.nn as nn + +__all__ = ["build_act", "get_act_name"] + +# register activation function here +# name: module, kwargs with default values +REGISTERED_ACT_DICT: dict[str, tuple[type, dict[str, any]]] = { + "relu": (nn.ReLU, {"inplace": True}), + "relu6": (nn.ReLU6, {"inplace": True}), + "hswish": (nn.Hardswish, {"inplace": True}), + "hsigmoid": (nn.Hardsigmoid, {"inplace": True}), + "swish": (nn.SiLU, {"inplace": True}), + "silu": (nn.SiLU, {"inplace": True}), + "tanh": (nn.Tanh, {}), + "sigmoid": (nn.Sigmoid, {}), + "gelu": (nn.GELU, {"approximate": "tanh"}), + "mish": (nn.Mish, {"inplace": True}), + "identity": (nn.Identity, {}), +} + + +def build_act(name: str or None, **kwargs) -> nn.Module or None: + if name in REGISTERED_ACT_DICT: + act_cls, default_args = copy.deepcopy(REGISTERED_ACT_DICT[name]) + for key in default_args: + if key in kwargs: + default_args[key] = kwargs[key] + return act_cls(**default_args) + elif name is None or name.lower() == "none": + return None + else: + raise ValueError(f"do not support: {name}") + + +def get_act_name(act: nn.Module or None) -> str or None: + if act is None: + return None + module2name = {} + for key, config in REGISTERED_ACT_DICT.items(): + module2name[config[0].__name__] = key + return module2name.get(type(act).__name__, "unknown") diff --git a/Sana/models/basic_modules.py b/Sana/models/basic_modules.py new file mode 100644 index 0000000..ece579a --- /dev/null +++ b/Sana/models/basic_modules.py @@ -0,0 +1,361 @@ +# Copyright 2024 NVIDIA CORPORATION & AFFILIATES +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# SPDX-License-Identifier: Apache-2.0 + +# This file is modified from https://github.com/PixArt-alpha/PixArt-sigma +import torch +import torch.nn as nn +from timm.models.vision_transformer import Mlp + +from .act import build_act, get_act_name +from .norms import build_norm, get_norm_name +from .utils import get_same_padding, val2tuple + + +class ConvLayer(nn.Module): + def __init__( + self, + in_dim: int, + out_dim: int, + kernel_size=3, + stride=1, + dilation=1, + groups=1, + padding: int or None = None, + use_bias=False, + dropout=0.0, + norm="bn2d", + act="relu", + ): + super().__init__() + if padding is None: + padding = get_same_padding(kernel_size) + padding *= dilation + + self.in_dim = in_dim + self.out_dim = out_dim + self.kernel_size = kernel_size + self.stride = stride + self.dilation = dilation + self.groups = groups + self.padding = padding + self.use_bias = use_bias + + self.dropout = nn.Dropout2d(dropout, inplace=False) if dropout > 0 else None + self.conv = nn.Conv2d( + in_dim, + out_dim, + kernel_size=(kernel_size, kernel_size), + stride=(stride, stride), + padding=padding, + dilation=(dilation, dilation), + groups=groups, + bias=use_bias, + ) + self.norm = build_norm(norm, num_features=out_dim) + self.act = build_act(act) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + if self.dropout is not None: + x = self.dropout(x) + x = self.conv(x) + if self.norm: + x = self.norm(x) + if self.act: + x = self.act(x) + return x + + +class GLUMBConv(nn.Module): + def __init__( + self, + in_features: int, + hidden_features: int, + out_feature=None, + kernel_size=3, + stride=1, + padding: int or None = None, + use_bias=False, + norm=(None, None, None), + act=("silu", "silu", None), + dilation=1, + ): + out_feature = out_feature or in_features + super().__init__() + use_bias = val2tuple(use_bias, 3) + norm = val2tuple(norm, 3) + act = val2tuple(act, 3) + + self.glu_act = build_act(act[1], inplace=False) + self.inverted_conv = ConvLayer( + in_features, + hidden_features * 2, + 1, + use_bias=use_bias[0], + norm=norm[0], + act=act[0], + ) + self.depth_conv = ConvLayer( + hidden_features * 2, + hidden_features * 2, + kernel_size, + stride=stride, + groups=hidden_features * 2, + padding=padding, + use_bias=use_bias[1], + norm=norm[1], + act=None, + dilation=dilation, + ) + self.point_conv = ConvLayer( + hidden_features, + out_feature, + 1, + use_bias=use_bias[2], + norm=norm[2], + act=act[2], + ) + # from IPython import embed; embed(header='debug dilate conv') + + def forward(self, x: torch.Tensor, HW=None) -> torch.Tensor: + B, N, C = x.shape + if HW is None: + H = W = int(N**0.5) + else: + H, W = HW + + x = x.reshape(B, H, W, C).permute(0, 3, 1, 2) + x = self.inverted_conv(x) + x = self.depth_conv(x) + + x, gate = torch.chunk(x, 2, dim=1) + gate = self.glu_act(gate) + x = x * gate + + x = self.point_conv(x) + x = x.reshape(B, C, N).permute(0, 2, 1) + + return x + + +class SlimGLUMBConv(GLUMBConv): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + + # 移除 self.inverted_conv 层 + del self.inverted_conv + self.out_dim = self.point_conv.out_dim + + def forward(self, x: torch.Tensor, HW=None) -> torch.Tensor: + B, N, C = x.shape + if HW is None: + H = W = int(N**0.5) + else: + H, W = HW + + # 直接使用 x,跳过 self.inverted_conv 层的调用 + x = x.reshape(B, H, W, C).permute(0, 3, 1, 2) + # x = self.inverted_conv(x) + x = self.depth_conv(x) + + x, gate = torch.chunk(x, 2, dim=1) + gate = self.glu_act(gate) + x = x * gate + + x = self.point_conv(x) + x = x.reshape(B, self.out_dim, N).permute(0, 2, 1) + + return x + + +class MBConvPreGLU(nn.Module): + def __init__( + self, + in_dim: int, + out_dim: int, + kernel_size=3, + stride=1, + mid_dim=None, + expand=6, + padding: int or None = None, + use_bias=False, + norm=(None, None, "ln2d"), + act=("silu", "silu", None), + ): + super().__init__() + use_bias = val2tuple(use_bias, 3) + norm = val2tuple(norm, 3) + act = val2tuple(act, 3) + + mid_dim = mid_dim or round(in_dim * expand) + + self.inverted_conv = ConvLayer( + in_dim, + mid_dim * 2, + 1, + use_bias=use_bias[0], + norm=norm[0], + act=None, + ) + self.glu_act = build_act(act[0], inplace=False) + self.depth_conv = ConvLayer( + mid_dim, + mid_dim, + kernel_size, + stride=stride, + groups=mid_dim, + padding=padding, + use_bias=use_bias[1], + norm=norm[1], + act=act[1], + ) + self.point_conv = ConvLayer( + mid_dim, + out_dim, + 1, + use_bias=use_bias[2], + norm=norm[2], + act=act[2], + ) + + def forward(self, x: torch.Tensor, HW=None) -> torch.Tensor: + B, N, C = x.shape + if HW is None: + H = W = int(N**0.5) + else: + H, W = HW + + x = x.reshape(B, H, W, C).permute(0, 3, 1, 2) + + x = self.inverted_conv(x) + x, gate = torch.chunk(x, 2, dim=1) + gate = self.glu_act(gate) + x = x * gate + + x = self.depth_conv(x) + x = self.point_conv(x) + + x = x.reshape(B, C, N).permute(0, 2, 1) + return x + + @property + def module_str(self) -> str: + _str = f"{self.depth_conv.kernel_size}{type(self).__name__}(" + _str += f"in={self.inverted_conv.in_dim},mid={self.depth_conv.in_dim},out={self.point_conv.out_dim},s={self.depth_conv.stride}" + _str += ( + f",norm={get_norm_name(self.inverted_conv.norm)}" + f"+{get_norm_name(self.depth_conv.norm)}" + f"+{get_norm_name(self.point_conv.norm)}" + ) + _str += ( + f",act={get_act_name(self.inverted_conv.act)}" + f"+{get_act_name(self.depth_conv.act)}" + f"+{get_act_name(self.point_conv.act)}" + ) + _str += f",glu_act={get_act_name(self.glu_act)})" + return _str + + +class DWMlp(Mlp): + """MLP as used in Vision Transformer, MLP-Mixer and related networks""" + + def __init__( + self, + in_features, + hidden_features=None, + out_features=None, + act_layer=nn.GELU, + bias=True, + drop=0.0, + kernel_size=3, + stride=1, + dilation=1, + padding=None, + ): + super().__init__( + in_features=in_features, + hidden_features=hidden_features, + out_features=out_features, + act_layer=act_layer, + bias=bias, + drop=drop, + ) + hidden_features = hidden_features or in_features + self.hidden_features = hidden_features + if padding is None: + padding = get_same_padding(kernel_size) + padding *= dilation + + self.conv = nn.Conv2d( + hidden_features, + hidden_features, + kernel_size=(kernel_size, kernel_size), + stride=(stride, stride), + padding=padding, + dilation=(dilation, dilation), + groups=hidden_features, + bias=bias, + ) + + def forward(self, x, HW=None): + B, N, C = x.shape + if HW is None: + H = W = int(N**0.5) + else: + H, W = HW + x = self.fc1(x) + x = self.act(x) + x = self.drop1(x) + x = x.reshape(B, H, W, self.hidden_features).permute(0, 3, 1, 2) + x = self.conv(x) + x = x.reshape(B, self.hidden_features, N).permute(0, 2, 1) + x = self.fc2(x) + x = self.drop2(x) + return x + + +class Mlp(Mlp): + """MLP as used in Vision Transformer, MLP-Mixer and related networks""" + + def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, bias=True, drop=0.0): + super().__init__( + in_features=in_features, + hidden_features=hidden_features, + out_features=out_features, + act_layer=act_layer, + bias=bias, + drop=drop, + ) + + def forward(self, x, HW=None): + x = self.fc1(x) + x = self.act(x) + x = self.drop1(x) + x = self.fc2(x) + x = self.drop2(x) + return x + + +if __name__ == "__main__": + model = GLUMBConv( + 1152, + 1152 * 4, + 1152, + use_bias=(True, True, False), + norm=(None, None, None), + act=("silu", "silu", None), + ).cuda() + input = torch.randn(4, 256, 1152).cuda() + output = model(input) diff --git a/Sana/models/norms.py b/Sana/models/norms.py new file mode 100644 index 0000000..4731d69 --- /dev/null +++ b/Sana/models/norms.py @@ -0,0 +1,225 @@ +# Copyright 2024 NVIDIA CORPORATION & AFFILIATES +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# SPDX-License-Identifier: Apache-2.0 + +import copy +import warnings + +import torch +import torch.nn as nn +from torch.nn.modules.batchnorm import _BatchNorm + +__all__ = ["LayerNorm2d", "build_norm", "get_norm_name", "reset_bn", "remove_bn", "set_norm_eps"] + + +class LayerNorm2d(nn.LayerNorm): + rmsnorm = False + + def forward(self, x: torch.Tensor) -> torch.Tensor: + out = x if LayerNorm2d.rmsnorm else x - torch.mean(x, dim=1, keepdim=True) + out = out / torch.sqrt(torch.square(out).mean(dim=1, keepdim=True) + self.eps) + if self.elementwise_affine: + out = out * self.weight.view(1, -1, 1, 1) + self.bias.view(1, -1, 1, 1) + return out + + def extra_repr(self) -> str: + return f"{self.normalized_shape}, eps={self.eps}, elementwise_affine={self.elementwise_affine}, rmsnorm={self.rmsnorm}" + + +# register normalization function here +# name: module, kwargs with default values +REGISTERED_NORMALIZATION_DICT: dict[str, tuple[type, dict[str, any]]] = { + "bn2d": (nn.BatchNorm2d, {"num_features": None, "eps": 1e-5, "momentum": 0.1, "affine": True}), + "syncbn": (nn.SyncBatchNorm, {"num_features": None, "eps": 1e-5, "momentum": 0.1, "affine": True}), + "ln": (nn.LayerNorm, {"normalized_shape": None, "eps": 1e-5, "elementwise_affine": True}), + "ln2d": (LayerNorm2d, {"normalized_shape": None, "eps": 1e-5, "elementwise_affine": True}), +} + + +def build_norm(name="bn2d", num_features=None, affine=True, **kwargs) -> nn.Module or None: + if name in ["ln", "ln2d"]: + kwargs["normalized_shape"] = num_features + kwargs["elementwise_affine"] = affine + else: + kwargs["num_features"] = num_features + kwargs["affine"] = affine + if name in REGISTERED_NORMALIZATION_DICT: + norm_cls, default_args = copy.deepcopy(REGISTERED_NORMALIZATION_DICT[name]) + for key in default_args: + if key in kwargs: + default_args[key] = kwargs[key] + return norm_cls(**default_args) + elif name is None or name.lower() == "none": + return None + else: + raise ValueError("do not support: %s" % name) + + +def get_norm_name(norm: nn.Module or None) -> str or None: + if norm is None: + return None + module2name = {} + for key, config in REGISTERED_NORMALIZATION_DICT.items(): + module2name[config[0].__name__] = key + return module2name.get(type(norm).__name__, "unknown") + + +def reset_bn( + model: nn.Module, + data_loader: list, + sync=True, + progress_bar=False, +) -> None: + import copy + + import torch.nn.functional as F + from packages.apps.utils import AverageMeter, is_master, sync_tensor + from packages.models.utils import get_device, list_join + from tqdm import tqdm + + bn_mean = {} + bn_var = {} + + tmp_model = copy.deepcopy(model) + for name, m in tmp_model.named_modules(): + if isinstance(m, _BatchNorm): + bn_mean[name] = AverageMeter(is_distributed=False) + bn_var[name] = AverageMeter(is_distributed=False) + + def new_forward(bn, mean_est, var_est): + def lambda_forward(x): + x = x.contiguous() + if sync: + batch_mean = x.mean(0, keepdim=True).mean(2, keepdim=True).mean(3, keepdim=True) # 1, C, 1, 1 + batch_mean = sync_tensor(batch_mean, reduce="cat") + batch_mean = torch.mean(batch_mean, dim=0, keepdim=True) + + batch_var = (x - batch_mean) * (x - batch_mean) + batch_var = batch_var.mean(0, keepdim=True).mean(2, keepdim=True).mean(3, keepdim=True) + batch_var = sync_tensor(batch_var, reduce="cat") + batch_var = torch.mean(batch_var, dim=0, keepdim=True) + else: + batch_mean = x.mean(0, keepdim=True).mean(2, keepdim=True).mean(3, keepdim=True) # 1, C, 1, 1 + batch_var = (x - batch_mean) * (x - batch_mean) + batch_var = batch_var.mean(0, keepdim=True).mean(2, keepdim=True).mean(3, keepdim=True) + + batch_mean = torch.squeeze(batch_mean) + batch_var = torch.squeeze(batch_var) + + mean_est.update(batch_mean.data, x.size(0)) + var_est.update(batch_var.data, x.size(0)) + + # bn forward using calculated mean & var + _feature_dim = batch_mean.shape[0] + return F.batch_norm( + x, + batch_mean, + batch_var, + bn.weight[:_feature_dim], + bn.bias[:_feature_dim], + False, + 0.0, + bn.eps, + ) + + return lambda_forward + + m.forward = new_forward(m, bn_mean[name], bn_var[name]) + + # skip if there is no batch normalization layers in the network + if len(bn_mean) == 0: + return + + tmp_model.eval() + with torch.inference_mode(): + with tqdm(total=len(data_loader), desc="reset bn", disable=not progress_bar or not is_master()) as t: + for images in data_loader: + images = images.to(get_device(tmp_model)) + tmp_model(images) + t.set_postfix( + { + "bs": images.size(0), + "res": list_join(images.shape[-2:], "x"), + } + ) + t.update() + + for name, m in model.named_modules(): + if name in bn_mean and bn_mean[name].count > 0: + feature_dim = bn_mean[name].avg.size(0) + assert isinstance(m, _BatchNorm) + m.running_mean.data[:feature_dim].copy_(bn_mean[name].avg) + m.running_var.data[:feature_dim].copy_(bn_var[name].avg) + + +def remove_bn(model: nn.Module) -> None: + for m in model.modules(): + if isinstance(m, _BatchNorm): + m.weight = m.bias = None + m.forward = lambda x: x + + +def set_norm_eps(model: nn.Module, eps: float or None = None, momentum: float or None = None) -> None: + for m in model.modules(): + if isinstance(m, (nn.GroupNorm, nn.LayerNorm, _BatchNorm)): + if eps is not None: + m.eps = eps + if momentum is not None: + m.momentum = momentum + + +class RMSNorm(torch.nn.Module): + def __init__(self, dim: int, scale_factor=1.0, eps: float = 1e-6): + """ + Initialize the RMSNorm normalization layer. + + Args: + dim (int): The dimension of the input tensor. + eps (float, optional): A small value added to the denominator for numerical stability. Default is 1e-6. + + Attributes: + eps (float): A small value added to the denominator for numerical stability. + weight (nn.Parameter): Learnable scaling parameter. + + """ + super().__init__() + self.eps = eps + self.weight = nn.Parameter(torch.ones(dim) * scale_factor) + + def _norm(self, x): + """ + Apply the RMSNorm normalization to the input tensor. + + Args: + x (torch.Tensor): The input tensor. + + Returns: + torch.Tensor: The normalized tensor. + + """ + return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) + + def forward(self, x): + """ + Forward pass through the RMSNorm layer. + + Args: + x (torch.Tensor): The input tensor. + + Returns: + torch.Tensor: The output tensor after applying RMSNorm. + + """ + return (self.weight * self._norm(x.float())).type_as(x) diff --git a/Sana/models/sana.py b/Sana/models/sana.py new file mode 100644 index 0000000..0dd6551 --- /dev/null +++ b/Sana/models/sana.py @@ -0,0 +1,379 @@ +# Copyright 2024 NVIDIA CORPORATION & AFFILIATES +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# SPDX-License-Identifier: Apache-2.0 + +# This file is modified from https://github.com/PixArt-alpha/PixArt-sigma +import os + +import numpy as np +import torch +import torch.nn as nn +from timm.models.layers import DropPath + +from .basic_modules import DWMlp, GLUMBConv, MBConvPreGLU, Mlp +from .sana_blocks import ( + Attention, + CaptionEmbedder, + FlashAttention, + LiteLA, + MultiHeadCrossAttention, + PatchEmbed, + T2IFinalLayer, + TimestepEmbedder, + t2i_modulate, +) +from .norms import RMSNorm +from .utils import auto_grad_checkpoint, to_2tuple + + +class SanaBlock(nn.Module): + """ + A Sana block with global shared adaptive layer norm (adaLN-single) conditioning. + """ + + def __init__( + self, + hidden_size, + num_heads, + mlp_ratio=4.0, + drop_path=0, + input_size=None, + qk_norm=False, + attn_type="flash", + ffn_type="mlp", + mlp_acts=("silu", "silu", None), + linear_head_dim=32, + **block_kwargs, + ): + super().__init__() + self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + if attn_type == "flash": + # flash self attention + self.attn = FlashAttention( + hidden_size, + num_heads=num_heads, + qkv_bias=True, + qk_norm=qk_norm, + **block_kwargs, + ) + elif attn_type == "linear": + # linear self attention + # TODO: Here the num_heads set to 36 for tmp used + self_num_heads = hidden_size // linear_head_dim + self.attn = LiteLA(hidden_size, hidden_size, heads=self_num_heads, eps=1e-8, qk_norm=qk_norm) + elif attn_type == "vanilla": + # vanilla self attention + self.attn = Attention(hidden_size, num_heads=num_heads, qkv_bias=True) + else: + raise ValueError(f"{attn_type} type is not defined.") + + self.cross_attn = MultiHeadCrossAttention(hidden_size, num_heads, **block_kwargs) + self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + # to be compatible with lower version pytorch + if ffn_type == "dwmlp": + approx_gelu = lambda: nn.GELU(approximate="tanh") + self.mlp = DWMlp( + in_features=hidden_size, hidden_features=int(hidden_size * mlp_ratio), act_layer=approx_gelu, drop=0 + ) + elif ffn_type == "glumbconv": + self.mlp = GLUMBConv( + in_features=hidden_size, + hidden_features=int(hidden_size * mlp_ratio), + use_bias=(True, True, False), + norm=(None, None, None), + act=mlp_acts, + ) + elif ffn_type == "glumbconv_dilate": + self.mlp = GLUMBConv( + in_features=hidden_size, + hidden_features=int(hidden_size * mlp_ratio), + use_bias=(True, True, False), + norm=(None, None, None), + act=mlp_acts, + dilation=2, + ) + elif ffn_type == "mbconvpreglu": + self.mlp = MBConvPreGLU( + in_dim=hidden_size, + out_dim=hidden_size, + mid_dim=int(hidden_size * mlp_ratio), + use_bias=(True, True, False), + norm=None, + act=("silu", "silu", None), + ) + elif ffn_type == "mlp": + approx_gelu = lambda: nn.GELU(approximate="tanh") + self.mlp = Mlp( + in_features=hidden_size, hidden_features=int(hidden_size * mlp_ratio), act_layer=approx_gelu, drop=0 + ) + else: + raise ValueError(f"{ffn_type} type is not defined.") + self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity() + self.scale_shift_table = nn.Parameter(torch.randn(6, hidden_size) / hidden_size**0.5) + + def forward(self, x, y, t, mask=None, **kwargs): + B, N, C = x.shape + + shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = ( + self.scale_shift_table[None] + t.reshape(B, 6, -1) + ).chunk(6, dim=1) + x = x + self.drop_path(gate_msa * self.attn(t2i_modulate(self.norm1(x), shift_msa, scale_msa)).reshape(B, N, C)) + x = x + self.cross_attn(x, y, mask) + x = x + self.drop_path(gate_mlp * self.mlp(t2i_modulate(self.norm2(x), shift_mlp, scale_mlp))) + + return x + + +############################################################################# +# Core Sana Model # +################################################################################# +class Sana(nn.Module): + """ + Diffusion model with a Transformer backbone. + """ + + def __init__( + self, + input_size=32, + patch_size=1, + in_channels=32, + hidden_size=1152, + depth=28, + num_heads=36, + mlp_ratio=2.5, + class_dropout_prob=0.1, + pred_sigma=False, + drop_path: float = 0.0, + caption_channels=2304, + pe_interpolation=1.0, + config=None, + model_max_length=120, + qk_norm=False, + y_norm=False, + norm_eps=1e-5, + attn_type="flash", + ffn_type="mlp", + use_pe=False, + y_norm_scale_factor=1.0, + patch_embed_kernel=None, + mlp_acts=("silu", "silu", None), + linear_head_dim=32, + **kwargs, + ): + super().__init__() + self.pred_sigma = pred_sigma + self.in_channels = in_channels + self.out_channels = in_channels * 2 if pred_sigma else in_channels + self.patch_size = patch_size + self.num_heads = num_heads + self.pe_interpolation = pe_interpolation + self.depth = depth + self.use_pe = use_pe + self.y_norm = y_norm + self.fp32_attention = kwargs.get("use_fp32_attention", False) + + kernel_size = patch_embed_kernel or patch_size + self.x_embedder = PatchEmbed( + input_size, patch_size, in_channels, hidden_size, kernel_size=kernel_size, bias=True + ) + self.t_embedder = TimestepEmbedder(hidden_size) + num_patches = self.x_embedder.num_patches + self.base_size = input_size // self.patch_size + # Will use fixed sin-cos embedding: + self.register_buffer("pos_embed", torch.zeros(1, num_patches, hidden_size)) + + approx_gelu = lambda: nn.GELU(approximate="tanh") + self.t_block = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 6 * hidden_size, bias=True)) + self.y_embedder = CaptionEmbedder( + in_channels=caption_channels, + hidden_size=hidden_size, + uncond_prob=class_dropout_prob, + act_layer=approx_gelu, + token_num=model_max_length, + ) + if self.y_norm: + self.attention_y_norm = RMSNorm(hidden_size, scale_factor=y_norm_scale_factor, eps=norm_eps) + drop_path = [x.item() for x in torch.linspace(0, drop_path, depth)] # stochastic depth decay rule + self.blocks = nn.ModuleList( + [ + SanaBlock( + hidden_size, + num_heads, + mlp_ratio=mlp_ratio, + drop_path=drop_path[i], + input_size=(input_size // patch_size, input_size // patch_size), + qk_norm=qk_norm, + attn_type=attn_type, + ffn_type=ffn_type, + mlp_acts=mlp_acts, + linear_head_dim=linear_head_dim, + ) + for i in range(depth) + ] + ) + self.final_layer = T2IFinalLayer(hidden_size, patch_size, self.out_channels) + + self.initialize_weights() + + def forward(self, x, timestep, y, mask=None, data_info=None, **kwargs): + """ + Forward pass of Sana. + x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images) + t: (N,) tensor of diffusion timesteps + y: (N, 1, 120, C) tensor of class labels + """ + x = x.to(self.dtype) + timestep = timestep.to(self.dtype) + y = y.to(self.dtype) + pos_embed = self.pos_embed.to(self.dtype) + self.h, self.w = x.shape[-2] // self.patch_size, x.shape[-1] // self.patch_size + if self.use_pe: + x = self.x_embedder(x) + pos_embed # (N, T, D), where T = H * W / patch_size ** 2 + else: + x = self.x_embedder(x) + t = self.t_embedder(timestep.to(x.dtype)) # (N, D) + t0 = self.t_block(t) + y = self.y_embedder(y, self.training) # (N, 1, L, D) + if self.y_norm: + y = self.attention_y_norm(y) + if mask is not None: + if mask.shape[0] != y.shape[0]: + mask = mask.repeat(y.shape[0] // mask.shape[0], 1) + mask = mask.squeeze(1).squeeze(1) + y = y.squeeze(1).masked_select(mask.unsqueeze(-1) != 0).view(1, -1, x.shape[-1]) + y_lens = mask.sum(dim=1).tolist() + else: + y_lens = [y.shape[2]] * y.shape[0] + y = y.squeeze(1).view(1, -1, x.shape[-1]) + for block in self.blocks: + x = auto_grad_checkpoint(block, x, y, t0, y_lens) # (N, T, D) #support grad checkpoint + x = self.final_layer(x, t) # (N, T, patch_size ** 2 * out_channels) + x = self.unpatchify(x) # (N, out_channels, H, W) + return x + + def __call__(self, *args, **kwargs): + """ + This method allows the object to be called like a function. + It simply calls the forward method. + """ + return self.forward(*args, **kwargs) + + def forward_with_dpmsolver(self, x, timestep, y, mask=None, **kwargs): + """ + dpm solver donnot need variance prediction + """ + # https://github.com/openai/glide-text2im/blob/main/notebooks/text2im.ipynb + model_out = self.forward(x, timestep, y, mask) + return model_out.chunk(2, dim=1)[0] if self.pred_sigma else model_out + + def unpatchify(self, x): + """ + x: (N, T, patch_size**2 * C) + imgs: (N, H, W, C) + """ + c = self.out_channels + p = self.x_embedder.patch_size[0] + h = w = int(x.shape[1] ** 0.5) + assert h * w == x.shape[1] + + x = x.reshape(shape=(x.shape[0], h, w, p, p, c)) + x = torch.einsum("nhwpqc->nchpwq", x) + imgs = x.reshape(shape=(x.shape[0], c, h * p, h * p)) + return imgs + + def initialize_weights(self): + # Initialize transformer layers: + def _basic_init(module): + if isinstance(module, nn.Linear): + torch.nn.init.xavier_uniform_(module.weight) + if module.bias is not None: + nn.init.constant_(module.bias, 0) + + self.apply(_basic_init) + + if self.use_pe: + # Initialize (and freeze) pos_embed by sin-cos embedding: + pos_embed = get_2d_sincos_pos_embed( + self.pos_embed.shape[-1], + int(self.x_embedder.num_patches**0.5), + pe_interpolation=self.pe_interpolation, + base_size=self.base_size, + ) + self.pos_embed.data.copy_(torch.from_numpy(pos_embed).float().unsqueeze(0)) + + # Initialize patch_embed like nn.Linear (instead of nn.Conv2d): + w = self.x_embedder.proj.weight.data + nn.init.xavier_uniform_(w.view([w.shape[0], -1])) + + # Initialize timestep embedding MLP: + nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02) + nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02) + nn.init.normal_(self.t_block[1].weight, std=0.02) + + # Initialize caption embedding MLP: + nn.init.normal_(self.y_embedder.y_proj.fc1.weight, std=0.02) + nn.init.normal_(self.y_embedder.y_proj.fc2.weight, std=0.02) + + +def get_2d_sincos_pos_embed(embed_dim, grid_size, cls_token=False, extra_tokens=0, pe_interpolation=1.0, base_size=16): + """ + grid_size: int of the grid height and width + return: + pos_embed: [grid_size*grid_size, embed_dim] or [1+grid_size*grid_size, embed_dim] (w/ or w/o cls_token) + """ + if isinstance(grid_size, int): + grid_size = to_2tuple(grid_size) + grid_h = np.arange(grid_size[0], dtype=np.float32) / (grid_size[0] / base_size) / pe_interpolation + grid_w = np.arange(grid_size[1], dtype=np.float32) / (grid_size[1] / base_size) / pe_interpolation + grid = np.meshgrid(grid_w, grid_h) # here w goes first + grid = np.stack(grid, axis=0) + grid = grid.reshape([2, 1, grid_size[1], grid_size[0]]) + + pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid) + if cls_token and extra_tokens > 0: + pos_embed = np.concatenate([np.zeros([extra_tokens, embed_dim]), pos_embed], axis=0) + return pos_embed + + +def get_2d_sincos_pos_embed_from_grid(embed_dim, grid): + assert embed_dim % 2 == 0 + + # use half of dimensions to encode grid_h + emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0]) # (H*W, D/2) + emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1]) # (H*W, D/2) + + emb = np.concatenate([emb_h, emb_w], axis=1) # (H*W, D) + return emb + + +def get_1d_sincos_pos_embed_from_grid(embed_dim, pos): + """ + embed_dim: output dimension for each position + pos: a list of positions to be encoded: size (M,) + out: (M, D) + """ + assert embed_dim % 2 == 0 + omega = np.arange(embed_dim // 2, dtype=np.float64) + omega /= embed_dim / 2.0 + omega = 1.0 / 10000**omega # (D/2,) + + pos = pos.reshape(-1) # (M,) + out = np.einsum("m,d->md", pos, omega) # (M, D/2), outer product + + emb_sin = np.sin(out) # (M, D/2) + emb_cos = np.cos(out) # (M, D/2) + + emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D) + return emb \ No newline at end of file diff --git a/Sana/models/sana_blocks.py b/Sana/models/sana_blocks.py new file mode 100644 index 0000000..31ac821 --- /dev/null +++ b/Sana/models/sana_blocks.py @@ -0,0 +1,798 @@ +# Copyright 2024 NVIDIA CORPORATION & AFFILIATES +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# SPDX-License-Identifier: Apache-2.0 + +# This file is modified from https://github.com/PixArt-alpha/PixArt-sigma +import math +import os +from typing import Optional + +import xformers.ops +import torch +import torch.nn as nn +import torch.nn.functional as F +from einops import rearrange +from timm.models.vision_transformer import Attention as Attention_ +from timm.models.vision_transformer import Mlp +from transformers import AutoModelForCausalLM + +from .norms import RMSNorm +from .utils import get_same_padding, to_2tuple + +sdpa_32b = None +Q_4GB_LIMIT = 32000000 +"""If q is greater than this, the operation will likely require >4GB VRAM, which will fail on Intel Arc Alchemist GPUs without a workaround.""" +# 2k = 37 748 736 +# 1024 = 9 437 184 +# 2k model goes very slightly over 4GB + +from comfy import model_management +if model_management.xformers_enabled(): + import xformers + import xformers.ops +else: + if model_management.xpu_available: + import intel_extension_for_pytorch as ipex + import os + if not torch.xpu.has_fp64_dtype() and not os.environ.get('IPEX_FORCE_ATTENTION_SLICE', None): + from ...utils.IPEX.attention import scaled_dot_product_attention_32_bit + sdpa_32b = scaled_dot_product_attention_32_bit + print("Using IPEX 4GB SDPA workaround") + else: + print("No IPEX 4GB workaround") + + +def modulate(x, shift, scale): + return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1) + + +def t2i_modulate(x, shift, scale): + return x * (1 + scale) + shift + + +class MultiHeadCrossAttention(nn.Module): + def __init__(self, d_model, num_heads, attn_drop=0.0, proj_drop=0.0, qk_norm=False, **block_kwargs): + super().__init__() + assert d_model % num_heads == 0, "d_model must be divisible by num_heads" + + self.d_model = d_model + self.num_heads = num_heads + self.head_dim = d_model // num_heads + + self.q_linear = nn.Linear(d_model, d_model) + self.kv_linear = nn.Linear(d_model, d_model * 2) + self.attn_drop = nn.Dropout(attn_drop) + self.proj = nn.Linear(d_model, d_model) + self.proj_drop = nn.Dropout(proj_drop) + if qk_norm: + # not used for now + self.q_norm = RMSNorm(d_model, scale_factor=1.0, eps=1e-6) + self.k_norm = RMSNorm(d_model, scale_factor=1.0, eps=1e-6) + else: + self.q_norm = nn.Identity() + self.k_norm = nn.Identity() + + def forward(self, x, cond, mask=None): + # query/value: img tokens; key: condition; mask: if padding tokens + B, N, C = x.shape + + q = self.q_linear(x).view(1, -1, self.num_heads, self.head_dim) + kv = self.kv_linear(cond).view(1, -1, 2, self.num_heads, self.head_dim) + k, v = kv.unbind(2) + + if model_management.xformers_enabled(): + attn_bias = None + if mask is not None: + attn_bias = xformers.ops.fmha.BlockDiagonalMask.from_seqlens([N] * B, mask) + x = xformers.ops.memory_efficient_attention( + q, k, v, + p=self.attn_drop.p, + attn_bias=attn_bias + ) + else: + q, k, v = map(lambda t: t.permute(0, 2, 1, 3),(q, k, v),) + attn_mask = None + if mask is not None and len(mask) > 1: + + # Create equivalent of xformer diagonal block mask, still only correct for square masks + # But depth doesn't matter as tensors can expand in that dimension + attn_mask_template = torch.ones( + [q.shape[2] // B, mask[0]], + dtype=torch.bool, + device=q.device + ) + attn_mask = torch.block_diag(attn_mask_template) + + # create a mask on the diagonal for each mask in the batch + for n in range(B - 1): + attn_mask = torch.block_diag(attn_mask, attn_mask_template) + + p = getattr(self.attn_drop, "p", 0) # IPEX.optimize() will turn attn_drop into an Identity() + + if sdpa_32b is not None and (q.element_size() * q.nelement()) > Q_4GB_LIMIT: + sdpa = sdpa_32b + else: + sdpa = torch.nn.functional.scaled_dot_product_attention + + x = sdpa( + q, k, v, + attn_mask=attn_mask, + dropout_p=p + ).permute(0, 2, 1, 3).contiguous() + x = x.view(B, -1, C) + x = self.proj(x) + x = self.proj_drop(x) + return x + + +class LiteLA(Attention_): + r"""Lightweight linear attention""" + + PAD_VAL = 1 + + def __init__( + self, + in_dim: int, + out_dim: int, + heads: Optional[int] = None, + heads_ratio: float = 1.0, + dim=32, + eps=1e-15, + use_bias=False, + qk_norm=False, + norm_eps=1e-5, + ): + heads = heads or int(out_dim // dim * heads_ratio) + super().__init__(in_dim, num_heads=heads, qkv_bias=use_bias) + + self.in_dim = in_dim + self.out_dim = out_dim + self.heads = heads + self.dim = out_dim // heads # TODO: need some change + self.eps = eps + + self.kernel_func = nn.ReLU(inplace=False) + if qk_norm: + self.q_norm = RMSNorm(in_dim, scale_factor=1.0, eps=norm_eps) + self.k_norm = RMSNorm(in_dim, scale_factor=1.0, eps=norm_eps) + else: + self.q_norm = nn.Identity() + self.k_norm = nn.Identity() + + def attn_matmul(self, q, k, v: torch.Tensor) -> torch.Tensor: + # lightweight linear attention + q = self.kernel_func(q) # B, h, h_d, N + k = self.kernel_func(k) + + q, k, v = q.float(), k.float(), v.float() + + v = F.pad(v, (0, 0, 0, 1), mode="constant", value=LiteLA.PAD_VAL) + vk = torch.matmul(v, k) + out = torch.matmul(vk, q) + + if out.dtype in [torch.float16, torch.bfloat16]: + out = out.float() + out = out[:, :, :-1] / (out[:, :, -1:] + self.eps) + + return out + + def forward(self, x: torch.Tensor, mask=None, HW=None, block_id=None) -> torch.Tensor: + B, N, C = x.shape + + qkv = self.qkv(x).reshape(B, N, 3, C) + q, k, v = qkv.unbind(2) # B, N, 3, C --> B, N, C + dtype = q.dtype + + q = self.q_norm(q).transpose(-1, -2) # (B, N, C) -> (B, C, N) + k = self.k_norm(k).transpose(-1, -2) # (B, N, C) -> (B, C, N) + v = v.transpose(-1, -2) + + q = q.reshape(B, C // self.dim, self.dim, N) # (B, h, h_d, N) + k = k.reshape(B, C // self.dim, self.dim, N).transpose(-1, -2) # (B, h, N, h_d) + v = v.reshape(B, C // self.dim, self.dim, N) # (B, h, h_d, N) + + out = self.attn_matmul(q, k, v).to(dtype) + + out = out.view(B, C, N).permute(0, 2, 1) # B, N, C + out = self.proj(out) + + if torch.get_autocast_gpu_dtype() == torch.float16: + out = out.clip(-65504, 65504) + + return out + + @property + def module_str(self) -> str: + _str = type(self).__name__ + "(" + eps = f"{self.eps:.1E}" + _str += f"i={self.in_dim},o={self.out_dim},h={self.heads},d={self.dim},eps={eps}" + return _str + + def __repr__(self): + return f"EPS{self.eps}-" + super().__repr__() + + +class PAGCFGIdentitySelfAttnProcessorLiteLA: + r"""Self Attention with Perturbed Attention & CFG Guidance""" + + def __init__(self, attn): + self.attn = attn + + def __call__(self, x: torch.Tensor, mask=None, HW=None, block_id=None) -> torch.Tensor: + x_uncond, x_org, x_ptb = x.chunk(3) + x_org = torch.cat([x_uncond, x_org]) + B, N, C = x_org.shape + + qkv = self.attn.qkv(x_org).reshape(B, N, 3, C) + # B, N, 3, C --> B, N, C + q, k, v = qkv.unbind(2) + dtype = q.dtype + q = self.attn.q_norm(q).transpose(-1, -2) # (B, N, C) -> (B, C, N) + k = self.attn.k_norm(k).transpose(-1, -2) # (B, N, C) -> (B, C, N) + v = v.transpose(-1, -2) + + q = q.reshape(B, C // self.attn.dim, self.attn.dim, N) # (B, h, h_d, N) + k = k.reshape(B, C // self.attn.dim, self.attn.dim, N).transpose(-1, -2) # (B, h, N, h_d) + v = v.reshape(B, C // self.attn.dim, self.attn.dim, N) # (B, h, h_d, N) + + out = self.attn.attn_matmul(q, k, v).to(dtype) + + out = out.view(B, C, N).permute(0, 2, 1) # B, N, C + out = self.attn.proj(out) + + # perturbed path (identity attention) + v_weight = self.attn.qkv.weight[C * 2 : C * 3, :] # Shape: (dim, dim) + if self.attn.qkv.bias: + v_bias = self.attn.qkv.bias[C * 2 : C * 3] # Shape: (dim,) + x_ptb = (torch.matmul(x_ptb, v_weight.t()) + v_bias).to(dtype) + else: + x_ptb = torch.matmul(x_ptb, v_weight.t()).to(dtype) + x_ptb = self.attn.proj(x_ptb) + + out = torch.cat([out, x_ptb]) + + if torch.get_autocast_gpu_dtype() == torch.float16: + out = out.clip(-65504, 65504) + + return out + + +class PAGIdentitySelfAttnProcessorLiteLA: + r"""Self Attention with Perturbed Attention Guidance""" + + def __init__(self, attn): + self.attn = attn + + def __call__(self, x: torch.Tensor, mask=None, HW=None, block_id=None) -> torch.Tensor: + x_org, x_ptb = x.chunk(2) + B, N, C = x_org.shape + + qkv = self.attn.qkv(x_org).reshape(B, N, 3, C) + # B, N, 3, C --> B, N, C + q, k, v = qkv.unbind(2) + dtype = q.dtype + q = self.attn.q_norm(q).transpose(-1, -2) # (B, N, C) -> (B, C, N) + k = self.attn.k_norm(k).transpose(-1, -2) # (B, N, C) -> (B, C, N) + v = v.transpose(-1, -2) + + q = q.reshape(B, C // self.attn.dim, self.attn.dim, N) # (B, h, h_d, N) + k = k.reshape(B, C // self.attn.dim, self.attn.dim, N).transpose(-1, -2) # (B, h, N, h_d) + v = v.reshape(B, C // self.attn.dim, self.attn.dim, N) # (B, h, h_d, N) + + out = self.attn.attn_matmul(q, k, v).to(dtype) + + out = out.view(B, C, N).permute(0, 2, 1) # B, N, C + out = self.attn.proj(out) + + # perturbed path (identity attention) + v_weight = self.attn.qkv.weight[C * 2 : C * 3, :] # Shape: (dim, dim) + if self.attn.qkv.bias: + v_bias = self.attn.qkv.bias[C * 2 : C * 3] # Shape: (dim,) + x_ptb = (torch.matmul(x_ptb, v_weight.t()) + v_bias).to(dtype) + else: + x_ptb = torch.matmul(x_ptb, v_weight.t()).to(dtype) + x_ptb = self.attn.proj(x_ptb) + + out = torch.cat([out, x_ptb]) + + if torch.get_autocast_gpu_dtype() == torch.float16: + out = out.clip(-65504, 65504) + + return out + + +class SelfAttnProcessorLiteLA: + r"""Self Attention with Lite Linear Attention""" + + def __init__(self, attn): + self.attn = attn + + def __call__(self, x: torch.Tensor, mask=None, HW=None, block_id=None) -> torch.Tensor: + B, N, C = x.shape + if HW is None: + H = W = int(N**0.5) + else: + H, W = HW + qkv = self.attn.qkv(x).reshape(B, N, 3, C) + # B, N, 3, C --> B, N, C + q, k, v = qkv.unbind(2) + dtype = q.dtype + q = self.attn.q_norm(q).transpose(-1, -2) # (B, N, C) -> (B, C, N) + k = self.attn.k_norm(k).transpose(-1, -2) # (B, N, C) -> (B, C, N) + v = v.transpose(-1, -2) + + q = q.reshape(B, C // self.attn.dim, self.attn.dim, N) # (B, h, h_d, N) + k = k.reshape(B, C // self.attn.dim, self.attn.dim, N).transpose(-1, -2) # (B, h, N, h_d) + v = v.reshape(B, C // self.attn.dim, self.attn.dim, N) # (B, h, h_d, N) + + out = self.attn.attn_matmul(q, k, v).to(dtype) + + out = out.view(B, C, N).permute(0, 2, 1) # B, N, C + out = self.attn.proj(out) + + if torch.get_autocast_gpu_dtype() == torch.float16: + out = out.clip(-65504, 65504) + + return out + + +class FlashAttention(Attention_): + """Multi-head Flash Attention block with qk norm.""" + + def __init__( + self, + dim, + num_heads=8, + qkv_bias=True, + qk_norm=False, + **block_kwargs, + ): + """ + Args: + dim (int): Number of input channels. + num_heads (int): Number of attention heads. + qkv_bias (bool: If True, add a learnable bias to query, key, value. + """ + super().__init__(dim, num_heads=num_heads, qkv_bias=qkv_bias, **block_kwargs) + + if qk_norm: + self.q_norm = nn.LayerNorm(dim) + self.k_norm = nn.LayerNorm(dim) + else: + self.q_norm = nn.Identity() + self.k_norm = nn.Identity() + + def forward(self, x, mask=None, HW=None, block_id=None): + B, N, C = x.shape + + qkv = self.qkv(x).reshape(B, N, 3, C) + q, k, v = qkv.unbind(2) + dtype = q.dtype + + q = self.q_norm(q) + k = self.k_norm(k) + + q = q.reshape(B, N, self.num_heads, C // self.num_heads).to(dtype) + k = k.reshape(B, N, self.num_heads, C // self.num_heads).to(dtype) + v = v.reshape(B, N, self.num_heads, C // self.num_heads).to(dtype) + + use_fp32_attention = getattr(self, "fp32_attention", False) # necessary for NAN loss + if use_fp32_attention: + q, k, v = q.float(), k.float(), v.float() + + attn_bias = None + if mask is not None: + attn_bias = torch.zeros([B * self.num_heads, q.shape[1], k.shape[1]], dtype=q.dtype, device=q.device) + attn_bias.masked_fill_(mask.squeeze(1).repeat(self.num_heads, 1, 1) == 0, float("-inf")) + + if _xformers_available: + x = xformers.ops.memory_efficient_attention(q, k, v, p=self.attn_drop.p, attn_bias=attn_bias) + else: + q, k, v = q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2) + if mask is not None and mask.ndim == 2: + mask = (1 - mask.to(x.dtype)) * -10000.0 + mask = mask[:, None, None].repeat(1, self.num_heads, 1, 1) + x = F.scaled_dot_product_attention(q, k, v, attn_mask=mask, dropout_p=0.0, is_causal=False) + x = x.transpose(1, 2) + + x = x.view(B, N, C) + x = self.proj(x) + x = self.proj_drop(x) + + if torch.get_autocast_gpu_dtype() == torch.float16: + x = x.clip(-65504, 65504) + + return x + + +################################################################################# +# AMP attention with fp32 softmax to fix loss NaN problem during training # +################################################################################# +class Attention(Attention_): + def forward(self, x, HW=None): + B, N, C = x.shape + qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4) + # B,N,3,H,C -> B,H,N,C + q, k, v = qkv.unbind(0) # make torchscript happy (cannot use tensor as tuple) + use_fp32_attention = getattr(self, "fp32_attention", False) + if use_fp32_attention: + q, k = q.float(), k.float() + + with torch.cuda.amp.autocast(enabled=not use_fp32_attention): + attn = (q @ k.transpose(-2, -1)) * self.scale + attn = attn.softmax(dim=-1) + + attn = self.attn_drop(attn) + + x = (attn @ v).transpose(1, 2).reshape(B, N, C) + x = self.proj(x) + x = self.proj_drop(x) + return x + + +class FinalLayer(nn.Module): + """ + The final layer of Sana. + """ + + def __init__(self, hidden_size, patch_size, out_channels): + super().__init__() + self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True) + self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True)) + + def forward(self, x, c): + shift, scale = self.adaLN_modulation(c).chunk(2, dim=1) + x = modulate(self.norm_final(x), shift, scale) + x = self.linear(x) + return x + + +class T2IFinalLayer(nn.Module): + """ + The final layer of Sana. + """ + + def __init__(self, hidden_size, patch_size, out_channels): + super().__init__() + self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True) + self.scale_shift_table = nn.Parameter(torch.randn(2, hidden_size) / hidden_size**0.5) + self.out_channels = out_channels + + def forward(self, x, t): + shift, scale = (self.scale_shift_table[None] + t[:, None]).chunk(2, dim=1) + x = t2i_modulate(self.norm_final(x), shift, scale) + x = self.linear(x) + return x + + +class MaskFinalLayer(nn.Module): + """ + The final layer of Sana. + """ + + def __init__(self, final_hidden_size, c_emb_size, patch_size, out_channels): + super().__init__() + self.norm_final = nn.LayerNorm(final_hidden_size, elementwise_affine=False, eps=1e-6) + self.linear = nn.Linear(final_hidden_size, patch_size * patch_size * out_channels, bias=True) + self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(c_emb_size, 2 * final_hidden_size, bias=True)) + + def forward(self, x, t): + shift, scale = self.adaLN_modulation(t).chunk(2, dim=1) + x = modulate(self.norm_final(x), shift, scale) + x = self.linear(x) + return x + + +class DecoderLayer(nn.Module): + """ + The final layer of Sana. + """ + + def __init__(self, hidden_size, decoder_hidden_size): + super().__init__() + self.norm_decoder = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + self.linear = nn.Linear(hidden_size, decoder_hidden_size, bias=True) + self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True)) + + def forward(self, x, t): + shift, scale = self.adaLN_modulation(t).chunk(2, dim=1) + x = modulate(self.norm_decoder(x), shift, scale) + x = self.linear(x) + return x + + +################################################################################# +# Embedding Layers for Timesteps and Class Labels # +################################################################################# +class TimestepEmbedder(nn.Module): + """ + Embeds scalar timesteps into vector representations. + """ + + def __init__(self, hidden_size, frequency_embedding_size=256): + super().__init__() + self.mlp = nn.Sequential( + nn.Linear(frequency_embedding_size, hidden_size, bias=True), + nn.SiLU(), + nn.Linear(hidden_size, hidden_size, bias=True), + ) + self.frequency_embedding_size = frequency_embedding_size + + @staticmethod + def timestep_embedding(t, dim, max_period=10000): + """ + Create sinusoidal timestep embeddings. + :param t: a 1-D Tensor of N indices, one per batch element. + These may be fractional. + :param dim: the dimension of the output. + :param max_period: controls the minimum frequency of the embeddings. + :return: an (N, D) Tensor of positional embeddings. + """ + # https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py + half = dim // 2 + freqs = torch.exp( + -math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32, device=t.device) / half + ) + args = t[:, None].float() * freqs[None] + embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1) + if dim % 2: + embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1) + return embedding + + def forward(self, t): + t_freq = self.timestep_embedding(t, self.frequency_embedding_size).to(self.dtype) + t_emb = self.mlp(t_freq) + return t_emb + + @property + def dtype(self): + try: + return next(self.parameters()).dtype + except StopIteration: + return torch.float32 + + +class SizeEmbedder(TimestepEmbedder): + """ + Embeds scalar timesteps into vector representations. + """ + + def __init__(self, hidden_size, frequency_embedding_size=256): + super().__init__(hidden_size=hidden_size, frequency_embedding_size=frequency_embedding_size) + self.mlp = nn.Sequential( + nn.Linear(frequency_embedding_size, hidden_size, bias=True), + nn.SiLU(), + nn.Linear(hidden_size, hidden_size, bias=True), + ) + self.frequency_embedding_size = frequency_embedding_size + self.outdim = hidden_size + + def forward(self, s, bs): + if s.ndim == 1: + s = s[:, None] + assert s.ndim == 2 + if s.shape[0] != bs: + s = s.repeat(bs // s.shape[0], 1) + assert s.shape[0] == bs + b, dims = s.shape[0], s.shape[1] + s = rearrange(s, "b d -> (b d)") + s_freq = self.timestep_embedding(s, self.frequency_embedding_size).to(self.dtype) + s_emb = self.mlp(s_freq) + s_emb = rearrange(s_emb, "(b d) d2 -> b (d d2)", b=b, d=dims, d2=self.outdim) + return s_emb + + @property + def dtype(self): + try: + return next(self.parameters()).dtype + except StopIteration: + return torch.float32 + + +class LabelEmbedder(nn.Module): + """ + Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance. + """ + + def __init__(self, num_classes, hidden_size, dropout_prob): + super().__init__() + use_cfg_embedding = dropout_prob > 0 + self.embedding_table = nn.Embedding(num_classes + use_cfg_embedding, hidden_size) + self.num_classes = num_classes + self.dropout_prob = dropout_prob + + def token_drop(self, labels, force_drop_ids=None): + """ + Drops labels to enable classifier-free guidance. + """ + if force_drop_ids is None: + drop_ids = torch.rand(labels.shape[0]).cuda() < self.dropout_prob + else: + drop_ids = force_drop_ids == 1 + labels = torch.where(drop_ids, self.num_classes, labels) + return labels + + def forward(self, labels, train, force_drop_ids=None): + use_dropout = self.dropout_prob > 0 + if (train and use_dropout) or (force_drop_ids is not None): + labels = self.token_drop(labels, force_drop_ids) + embeddings = self.embedding_table(labels) + return embeddings + + +class CaptionEmbedder(nn.Module): + """ + Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance. + """ + + def __init__( + self, + in_channels, + hidden_size, + uncond_prob, + act_layer=nn.GELU(approximate="tanh"), + token_num=120, + ): + super().__init__() + self.y_proj = Mlp( + in_features=in_channels, hidden_features=hidden_size, out_features=hidden_size, act_layer=act_layer, drop=0 + ) + self.register_buffer("y_embedding", nn.Parameter(torch.randn(token_num, in_channels) / in_channels**0.5)) + self.uncond_prob = uncond_prob + + def initialize_gemma_params(self, model_name="google/gemma-2b-it"): + num_layers = len(self.custom_gemma_layers) + text_encoder = AutoModelForCausalLM.from_pretrained(model_name).get_decoder() + pretrained_layers = text_encoder.layers[-num_layers:] + for custom_layer, pretrained_layer in zip(self.custom_gemma_layers, pretrained_layers): + info = custom_layer.load_state_dict(pretrained_layer.state_dict(), strict=False) + print(f"**** {info} ****") + print(f"**** Initialized {num_layers} Gemma layers from pretrained model: {model_name} ****") + + def token_drop(self, caption, force_drop_ids=None): + """ + Drops labels to enable classifier-free guidance. + """ + if force_drop_ids is None: + drop_ids = torch.rand(caption.shape[0]).cuda() < self.uncond_prob + else: + drop_ids = force_drop_ids == 1 + caption = torch.where(drop_ids[:, None, None, None], self.y_embedding, caption) + return caption + + def forward(self, caption, train, force_drop_ids=None, mask=None): + if train: + assert caption.shape[2:] == self.y_embedding.shape + use_dropout = self.uncond_prob > 0 + if (train and use_dropout) or (force_drop_ids is not None): + caption = self.token_drop(caption, force_drop_ids) + + caption = self.y_proj(caption) + + return caption + + +class CaptionEmbedderDoubleBr(nn.Module): + """ + Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance. + """ + + def __init__(self, in_channels, hidden_size, uncond_prob, act_layer=nn.GELU(approximate="tanh"), token_num=120): + super().__init__() + self.proj = Mlp( + in_features=in_channels, hidden_features=hidden_size, out_features=hidden_size, act_layer=act_layer, drop=0 + ) + self.embedding = nn.Parameter(torch.randn(1, in_channels) / 10**0.5) + self.y_embedding = nn.Parameter(torch.randn(token_num, in_channels) / 10**0.5) + self.uncond_prob = uncond_prob + + def token_drop(self, global_caption, caption, force_drop_ids=None): + """ + Drops labels to enable classifier-free guidance. + """ + if force_drop_ids is None: + drop_ids = torch.rand(global_caption.shape[0]).cuda() < self.uncond_prob + else: + drop_ids = force_drop_ids == 1 + global_caption = torch.where(drop_ids[:, None], self.embedding, global_caption) + caption = torch.where(drop_ids[:, None, None, None], self.y_embedding, caption) + return global_caption, caption + + def forward(self, caption, train, force_drop_ids=None): + assert caption.shape[2:] == self.y_embedding.shape + global_caption = caption.mean(dim=2).squeeze() + use_dropout = self.uncond_prob > 0 + if (train and use_dropout) or (force_drop_ids is not None): + global_caption, caption = self.token_drop(global_caption, caption, force_drop_ids) + y_embed = self.proj(global_caption) + return y_embed, caption + + +class PatchEmbed(nn.Module): + """2D Image to Patch Embedding""" + + def __init__( + self, + img_size=224, + patch_size=16, + in_chans=3, + embed_dim=768, + kernel_size=None, + padding=0, + norm_layer=None, + flatten=True, + bias=True, + ): + super().__init__() + kernel_size = kernel_size or patch_size + img_size = to_2tuple(img_size) + patch_size = to_2tuple(patch_size) + self.img_size = img_size + self.patch_size = patch_size + self.grid_size = (img_size[0] // patch_size[0], img_size[1] // patch_size[1]) + self.num_patches = self.grid_size[0] * self.grid_size[1] + self.flatten = flatten + if not padding and kernel_size % 2 > 0: + padding = get_same_padding(kernel_size) + self.proj = nn.Conv2d( + in_chans, embed_dim, kernel_size=kernel_size, stride=patch_size, padding=padding, bias=bias + ) + self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity() + + def forward(self, x): + B, C, H, W = x.shape + assert (H == self.img_size[0], f"Input image height ({H}) doesn't match model ({self.img_size[0]}).") + assert (W == self.img_size[1], f"Input image width ({W}) doesn't match model ({self.img_size[1]}).") + x = self.proj(x) + if self.flatten: + x = x.flatten(2).transpose(1, 2) # BCHW -> BNC + x = self.norm(x) + return x + + +class PatchEmbedMS(nn.Module): + """2D Image to Patch Embedding""" + + def __init__( + self, + patch_size=16, + in_chans=3, + embed_dim=768, + kernel_size=None, + padding=0, + norm_layer=None, + flatten=True, + bias=True, + ): + super().__init__() + kernel_size = kernel_size or patch_size + patch_size = to_2tuple(patch_size) + self.patch_size = patch_size + self.flatten = flatten + if not padding and kernel_size % 2 > 0: + padding = get_same_padding(kernel_size) + self.proj = nn.Conv2d( + in_chans, embed_dim, kernel_size=kernel_size, stride=patch_size, padding=padding, bias=bias + ) + self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity() + + def forward(self, x): + x = self.proj(x) + if self.flatten: + x = x.flatten(2).transpose(1, 2) # BCHW -> BNC + x = self.norm(x) + return x diff --git a/Sana/models/sana_multi_scale.py b/Sana/models/sana_multi_scale.py new file mode 100644 index 0000000..7cc3745 --- /dev/null +++ b/Sana/models/sana_multi_scale.py @@ -0,0 +1,374 @@ +# Copyright 2024 NVIDIA CORPORATION & AFFILIATES +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# SPDX-License-Identifier: Apache-2.0 + +# This file is modified from https://github.com/PixArt-alpha/PixArt-sigma +import torch +import torch.nn as nn +from timm.models.layers import DropPath + +from .basic_modules import DWMlp, GLUMBConv, MBConvPreGLU, Mlp +from .sana import Sana, get_2d_sincos_pos_embed +from .sana_blocks import ( + Attention, + CaptionEmbedder, + FlashAttention, + LiteLA, + MultiHeadCrossAttention, + PatchEmbedMS, + T2IFinalLayer, + t2i_modulate, +) +from .utils import auto_grad_checkpoint + + +class SanaMSBlock(nn.Module): + """ + A Sana block with global shared adaptive layer norm zero (adaLN-Zero) conditioning. + """ + + def __init__( + self, + hidden_size, + num_heads, + mlp_ratio=4.0, + drop_path=0.0, + input_size=None, + qk_norm=False, + attn_type="flash", + ffn_type="mlp", + mlp_acts=("silu", "silu", None), + linear_head_dim=32, + cross_norm=False, + **block_kwargs, + ): + super().__init__() + self.hidden_size = hidden_size + self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + if attn_type == "flash": + # flash self attention + self.attn = FlashAttention( + hidden_size, + num_heads=num_heads, + qkv_bias=True, + qk_norm=qk_norm, + **block_kwargs, + ) + elif attn_type == "linear": + # linear self attention + # TODO: Here the num_heads set to 36 for tmp used + self_num_heads = hidden_size // linear_head_dim + self.attn = LiteLA(hidden_size, hidden_size, heads=self_num_heads, eps=1e-8, qk_norm=qk_norm) + elif attn_type == "vanilla": + # vanilla self attention + self.attn = Attention(hidden_size, num_heads=num_heads, qkv_bias=True) + else: + raise ValueError(f"{attn_type} type is not defined.") + + self.cross_attn = MultiHeadCrossAttention(hidden_size, num_heads, qk_norm=cross_norm, **block_kwargs) + self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + if ffn_type == "dwmlp": + approx_gelu = lambda: nn.GELU(approximate="tanh") + self.mlp = DWMlp( + in_features=hidden_size, hidden_features=int(hidden_size * mlp_ratio), act_layer=approx_gelu, drop=0 + ) + elif ffn_type == "glumbconv": + self.mlp = GLUMBConv( + in_features=hidden_size, + hidden_features=int(hidden_size * mlp_ratio), + use_bias=(True, True, False), + norm=(None, None, None), + act=mlp_acts, + ) + elif ffn_type == "glumbconv_dilate": + self.mlp = GLUMBConv( + in_features=hidden_size, + hidden_features=int(hidden_size * mlp_ratio), + use_bias=(True, True, False), + norm=(None, None, None), + act=mlp_acts, + dilation=2, + ) + elif ffn_type == "mlp": + approx_gelu = lambda: nn.GELU(approximate="tanh") + self.mlp = Mlp( + in_features=hidden_size, hidden_features=int(hidden_size * mlp_ratio), act_layer=approx_gelu, drop=0 + ) + elif ffn_type == "mbconvpreglu": + self.mlp = MBConvPreGLU( + in_dim=hidden_size, + out_dim=hidden_size, + mid_dim=int(hidden_size * mlp_ratio), + use_bias=(True, True, False), + norm=None, + act=mlp_acts, + ) + else: + raise ValueError(f"{ffn_type} type is not defined.") + self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity() + self.scale_shift_table = nn.Parameter(torch.randn(6, hidden_size) / hidden_size**0.5) + + def forward(self, x, y, t, mask=None, HW=None, **kwargs): + B, N, C = x.shape + + shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = ( + self.scale_shift_table[None] + t.reshape(B, 6, -1) + ).chunk(6, dim=1) + x = x + self.drop_path(gate_msa * self.attn(t2i_modulate(self.norm1(x), shift_msa, scale_msa), HW=HW)) + x = x + self.cross_attn(x, y, mask) + x = x + self.drop_path(gate_mlp * self.mlp(t2i_modulate(self.norm2(x), shift_mlp, scale_mlp), HW=HW)) + + return x + + +############################################################################# +# Core Sana Model # +################################################################################# +class SanaMS(Sana): + """ + Diffusion model with a Transformer backbone. + """ + + def __init__( + self, + input_size=32, + patch_size=2, + in_channels=32, + hidden_size=1152, + depth=28, + num_heads=16, + mlp_ratio=4.0, + class_dropout_prob=0.1, + learn_sigma=False, + pred_sigma=False, + drop_path: float = 0.0, + caption_channels=2304, + pe_interpolation=1.0, + config=None, + model_max_length=300, + qk_norm=False, + y_norm=False, + norm_eps=1e-5, + attn_type="linear", + ffn_type="glumbconv", + use_pe=False, + y_norm_scale_factor=1.0, + patch_embed_kernel=None, + mlp_acts=("silu", "silu", None), + linear_head_dim=32, + cross_norm=False, + **kwargs, + ): + super().__init__( + input_size=input_size, + patch_size=patch_size, + in_channels=in_channels, + hidden_size=hidden_size, + depth=depth, + num_heads=num_heads, + mlp_ratio=mlp_ratio, + class_dropout_prob=class_dropout_prob, + learn_sigma=learn_sigma, + pred_sigma=pred_sigma, + drop_path=drop_path, + caption_channels=caption_channels, + pe_interpolation=pe_interpolation, + config=config, + model_max_length=model_max_length, + qk_norm=qk_norm, + y_norm=y_norm, + norm_eps=norm_eps, + attn_type=attn_type, + ffn_type=ffn_type, + use_pe=use_pe, + y_norm_scale_factor=y_norm_scale_factor, + patch_embed_kernel=patch_embed_kernel, + mlp_acts=mlp_acts, + linear_head_dim=linear_head_dim, + **kwargs, + ) + self.dtype = torch.get_default_dtype() + self.h = self.w = 0 + approx_gelu = lambda: nn.GELU(approximate="tanh") + self.t_block = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 6 * hidden_size, bias=True)) + self.pos_embed_ms = None + + kernel_size = patch_embed_kernel or patch_size + self.x_embedder = PatchEmbedMS(patch_size, in_channels, hidden_size, kernel_size=kernel_size, bias=True) + self.y_embedder = CaptionEmbedder( + in_channels=caption_channels, + hidden_size=hidden_size, + uncond_prob=class_dropout_prob, + act_layer=approx_gelu, + token_num=model_max_length, + ) + drop_path = [x.item() for x in torch.linspace(0, drop_path, depth)] # stochastic depth decay rule + self.blocks = nn.ModuleList( + [ + SanaMSBlock( + hidden_size, + num_heads, + mlp_ratio=mlp_ratio, + drop_path=drop_path[i], + input_size=(input_size // patch_size, input_size // patch_size), + qk_norm=qk_norm, + attn_type=attn_type, + ffn_type=ffn_type, + mlp_acts=mlp_acts, + linear_head_dim=linear_head_dim, + cross_norm=cross_norm, + ) + for i in range(depth) + ] + ) + self.final_layer = T2IFinalLayer(hidden_size, patch_size, self.out_channels) + + self.initialize() + + def forward(self, x, timesteps, context, **kwargs): + """ + Forward pass that adapts comfy input to original forward function + x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images) + timesteps: (N,) tensor of diffusion timesteps + context: (N, 1, 120, C) conditioning + """ + ## size/ar from cond with fallback based on the latent image shape. + bs = x.shape[0] + ## Still accepts the input w/o that dim but returns garbage + if len(context.shape) == 3: + context = context.unsqueeze(1) + + ## run original forward pass + out = self.forward_raw( + x = x.to(self.dtype), + timestep = timesteps.to(self.dtype), + y = context.to(self.dtype), + ) + + ## only return EPS + out = out.to(torch.float) + + return out + + def forward_raw(self, x, timestep, y, mask=None, data_info=None, **kwargs): + """ + Forward pass of Sana. + x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images) + t: (N,) tensor of diffusion timesteps + y: (N, 1, 120, C) tensor of class labels + """ + bs = x.shape[0] + x = x.to(self.dtype) + timestep = timestep.to(self.dtype) + y = y.to(self.dtype) + self.h, self.w = x.shape[-2] // self.patch_size, x.shape[-1] // self.patch_size + if self.use_pe: + x = self.x_embedder(x) + if self.pos_embed_ms is None or self.pos_embed_ms.shape[1:] != x.shape[1:]: + self.pos_embed_ms = ( + torch.from_numpy( + get_2d_sincos_pos_embed( + self.pos_embed.shape[-1], + (self.h, self.w), + pe_interpolation=self.pe_interpolation, + base_size=self.base_size, + ) + ) + .unsqueeze(0) + .to(x.device) + .to(self.dtype) + ) + x += self.pos_embed_ms # (N, T, D), where T = H * W / patch_size ** 2 + else: + x = self.x_embedder(x) + + t = self.t_embedder(timestep) # (N, D) + + t0 = self.t_block(t) + y = self.y_embedder(y, self.training, mask=mask) # (N, D) + if self.y_norm: + y = self.attention_y_norm(y) + + if mask is not None: + if mask.shape[0] != y.shape[0]: + mask = mask.repeat(y.shape[0] // mask.shape[0], 1) + mask = mask.squeeze(1).squeeze(1) + y = y.squeeze(1).masked_select(mask.unsqueeze(-1) != 0).view(1, -1, x.shape[-1]) + y_lens = mask.sum(dim=1).tolist() + else: + y_lens = [y.shape[2]] * y.shape[0] + y = y.squeeze(1).view(1, -1, x.shape[-1]) + + for block in self.blocks: + x = auto_grad_checkpoint( + block, x, y, t0, y_lens, (self.h, self.w), **kwargs + ) # (N, T, D) #support grad checkpoint + + x = self.final_layer(x, t) # (N, T, patch_size ** 2 * out_channels) + x = self.unpatchify(x) # (N, out_channels, H, W) + + return x + + def __call__(self, *args, **kwargs): + """ + This method allows the object to be called like a function. + It simply calls the forward method. + """ + return self.forward(*args, **kwargs) + + def forward_with_dpmsolver(self, x, timestep, y, data_info, **kwargs): + """ + dpm solver donnot need variance prediction + """ + # https://github.com/openai/glide-text2im/blob/main/notebooks/text2im.ipynb + model_out = self.forward(x, timestep, y, data_info=data_info, **kwargs) + return model_out.chunk(2, dim=1)[0] if self.pred_sigma else model_out + + def unpatchify(self, x): + """ + x: (N, T, patch_size**2 * C) + imgs: (N, H, W, C) + """ + c = self.out_channels + p = self.x_embedder.patch_size[0] + assert self.h * self.w == x.shape[1] + + x = x.reshape(shape=(x.shape[0], self.h, self.w, p, p, c)) + x = torch.einsum("nhwpqc->nchpwq", x) + imgs = x.reshape(shape=(x.shape[0], c, self.h * p, self.w * p)) + return imgs + + def initialize(self): + # Initialize transformer layers: + def _basic_init(module): + if isinstance(module, nn.Linear): + torch.nn.init.xavier_uniform_(module.weight) + if module.bias is not None: + nn.init.constant_(module.bias, 0) + + self.apply(_basic_init) + + # Initialize patch_embed like nn.Linear (instead of nn.Conv2d): + w = self.x_embedder.proj.weight.data + nn.init.xavier_uniform_(w.view([w.shape[0], -1])) + + # Initialize timestep embedding MLP: + nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02) + nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02) + nn.init.normal_(self.t_block[1].weight, std=0.02) + + # Initialize caption embedding MLP: + nn.init.normal_(self.y_embedder.y_proj.fc1.weight, std=0.02) + nn.init.normal_(self.y_embedder.y_proj.fc2.weight, std=0.02) diff --git a/Sana/models/utils.py b/Sana/models/utils.py new file mode 100644 index 0000000..d74db3b --- /dev/null +++ b/Sana/models/utils.py @@ -0,0 +1,591 @@ +# Copyright 2024 NVIDIA CORPORATION & AFFILIATES +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# SPDX-License-Identifier: Apache-2.0 + +import math +import os +import random +import re +import sys +from collections.abc import Iterable +from itertools import repeat + +import torch +import torch.distributed as dist +import torch.nn as nn +import torch.nn.functional as F +from PIL import Image +from torch.utils.checkpoint import checkpoint, checkpoint_sequential +from torchvision import transforms as T + + +def _ntuple(n): + def parse(x): + if isinstance(x, Iterable) and not isinstance(x, str): + return x + return tuple(repeat(x, n)) + + return parse + + +to_1tuple = _ntuple(1) +to_2tuple = _ntuple(2) + + +def set_grad_checkpoint(model, gc_step=1): + assert isinstance(model, nn.Module) + + def set_attr(module): + module.grad_checkpointing = True + module.grad_checkpointing_step = gc_step + + model.apply(set_attr) + + +def set_fp32_attention(model): + assert isinstance(model, nn.Module) + + def set_attr(module): + module.fp32_attention = True + + model.apply(set_attr) + + +def auto_grad_checkpoint(module, *args, **kwargs): + if getattr(module, "grad_checkpointing", False): + if isinstance(module, Iterable): + gc_step = module[0].grad_checkpointing_step + return checkpoint_sequential(module, gc_step, *args, **kwargs) + else: + return checkpoint(module, *args, **kwargs) + return module(*args, **kwargs) + + +def checkpoint_sequential(functions, step, input, *args, **kwargs): + + # Hack for keyword-only parameter in a python 2.7-compliant way + preserve = kwargs.pop("preserve_rng_state", True) + if kwargs: + raise ValueError("Unexpected keyword arguments: " + ",".join(arg for arg in kwargs)) + + def run_function(start, end, functions): + def forward(input): + for j in range(start, end + 1): + input = functions[j](input, *args) + return input + + return forward + + if isinstance(functions, torch.nn.Sequential): + functions = list(functions.children()) + + # the last chunk has to be non-volatile + end = -1 + segment = len(functions) // step + for start in range(0, step * (segment - 1), step): + end = start + step - 1 + input = checkpoint(run_function(start, end, functions), input, preserve_rng_state=preserve) + return run_function(end + 1, len(functions) - 1, functions)(input) + + +def window_partition(x, window_size): + """ + Partition into non-overlapping windows with padding if needed. + Args: + x (tensor): input tokens with [B, H, W, C]. + window_size (int): window size. + + Returns: + windows: windows after partition with [B * num_windows, window_size, window_size, C]. + (Hp, Wp): padded height and width before partition + """ + B, H, W, C = x.shape + + pad_h = (window_size - H % window_size) % window_size + pad_w = (window_size - W % window_size) % window_size + if pad_h > 0 or pad_w > 0: + x = F.pad(x, (0, 0, 0, pad_w, 0, pad_h)) + Hp, Wp = H + pad_h, W + pad_w + + x = x.view(B, Hp // window_size, window_size, Wp // window_size, window_size, C) + windows = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, C) + return windows, (Hp, Wp) + + +def window_unpartition(windows, window_size, pad_hw, hw): + """ + Window unpartition into original sequences and removing padding. + Args: + x (tensor): input tokens with [B * num_windows, window_size, window_size, C]. + window_size (int): window size. + pad_hw (Tuple): padded height and width (Hp, Wp). + hw (Tuple): original height and width (H, W) before padding. + + Returns: + x: unpartitioned sequences with [B, H, W, C]. + """ + Hp, Wp = pad_hw + H, W = hw + B = windows.shape[0] // (Hp * Wp // window_size // window_size) + x = windows.view(B, Hp // window_size, Wp // window_size, window_size, window_size, -1) + x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, Hp, Wp, -1) + + if Hp > H or Wp > W: + x = x[:, :H, :W, :].contiguous() + return x + + +def get_rel_pos(q_size, k_size, rel_pos): + """ + Get relative positional embeddings according to the relative positions of + query and key sizes. + Args: + q_size (int): size of query q. + k_size (int): size of key k. + rel_pos (Tensor): relative position embeddings (L, C). + + Returns: + Extracted positional embeddings according to relative positions. + """ + max_rel_dist = int(2 * max(q_size, k_size) - 1) + # Interpolate rel pos if needed. + if rel_pos.shape[0] != max_rel_dist: + # Interpolate rel pos. + rel_pos_resized = F.interpolate( + rel_pos.reshape(1, rel_pos.shape[0], -1).permute(0, 2, 1), + size=max_rel_dist, + mode="linear", + ) + rel_pos_resized = rel_pos_resized.reshape(-1, max_rel_dist).permute(1, 0) + else: + rel_pos_resized = rel_pos + + # Scale the coords with short length if shapes for q and k are different. + q_coords = torch.arange(q_size)[:, None] * max(k_size / q_size, 1.0) + k_coords = torch.arange(k_size)[None, :] * max(q_size / k_size, 1.0) + relative_coords = (q_coords - k_coords) + (k_size - 1) * max(q_size / k_size, 1.0) + + return rel_pos_resized[relative_coords.long()] + + +def add_decomposed_rel_pos(attn, q, rel_pos_h, rel_pos_w, q_size, k_size): + """ + Calculate decomposed Relative Positional Embeddings from :paper:`mvitv2`. + https://github.com/facebookresearch/mvit/blob/19786631e330df9f3622e5402b4a419a263a2c80/mvit/models/attention.py # noqa B950 + Args: + attn (Tensor): attention map. + q (Tensor): query q in the attention layer with shape (B, q_h * q_w, C). + rel_pos_h (Tensor): relative position embeddings (Lh, C) for height axis. + rel_pos_w (Tensor): relative position embeddings (Lw, C) for width axis. + q_size (Tuple): spatial sequence size of query q with (q_h, q_w). + k_size (Tuple): spatial sequence size of key k with (k_h, k_w). + + Returns: + attn (Tensor): attention map with added relative positional embeddings. + """ + q_h, q_w = q_size + k_h, k_w = k_size + Rh = get_rel_pos(q_h, k_h, rel_pos_h) + Rw = get_rel_pos(q_w, k_w, rel_pos_w) + + B, _, dim = q.shape + r_q = q.reshape(B, q_h, q_w, dim) + rel_h = torch.einsum("bhwc,hkc->bhwk", r_q, Rh) + rel_w = torch.einsum("bhwc,wkc->bhwk", r_q, Rw) + + attn = (attn.view(B, q_h, q_w, k_h, k_w) + rel_h[:, :, :, :, None] + rel_w[:, :, :, None, :]).view( + B, q_h * q_w, k_h * k_w + ) + + return attn + + +def mean_flat(tensor): + return tensor.mean(dim=list(range(1, tensor.ndim))) + + +################################################################################# +# Token Masking and Unmasking # +################################################################################# +def get_mask(batch, length, mask_ratio, device, mask_type=None, data_info=None, extra_len=0): + """ + Get the binary mask for the input sequence. + Args: + - batch: batch size + - length: sequence length + - mask_ratio: ratio of tokens to mask + - data_info: dictionary with info for reconstruction + return: + mask_dict with following keys: + - mask: binary mask, 0 is keep, 1 is remove + - ids_keep: indices of tokens to keep + - ids_restore: indices to restore the original order + """ + assert mask_type in ["random", "fft", "laplacian", "group"] + mask = torch.ones([batch, length], device=device) + len_keep = int(length * (1 - mask_ratio)) - extra_len + + if mask_type == "random" or mask_type == "group": + noise = torch.rand(batch, length, device=device) # noise in [0, 1] + ids_shuffle = torch.argsort(noise, dim=1) # ascend: small is keep, large is remove + ids_restore = torch.argsort(ids_shuffle, dim=1) + # keep the first subset + ids_keep = ids_shuffle[:, :len_keep] + ids_removed = ids_shuffle[:, len_keep:] + + elif mask_type in ["fft", "laplacian"]: + if "strength" in data_info: + strength = data_info["strength"] + + else: + N = data_info["N"][0] + img = data_info["ori_img"] + # 获取原图的尺寸信息 + _, C, H, W = img.shape + if mask_type == "fft": + # 对图片进行reshape,将其变为patch (3, H/N, N, W/N, N) + reshaped_image = img.reshape((batch, -1, H // N, N, W // N, N)) + fft_image = torch.fft.fftn(reshaped_image, dim=(3, 5)) + # 取绝对值并求和获取频率强度 + strength = torch.sum(torch.abs(fft_image), dim=(1, 3, 5)).reshape( + ( + batch, + -1, + ) + ) + elif type == "laplacian": + laplacian_kernel = torch.tensor([[-1, -1, -1], [-1, 8, -1], [-1, -1, -1]], dtype=torch.float32).reshape( + 1, 1, 3, 3 + ) + laplacian_kernel = laplacian_kernel.repeat(C, 1, 1, 1) + # 对图片进行reshape,将其变为patch (3, H/N, N, W/N, N) + reshaped_image = img.reshape(-1, C, H // N, N, W // N, N).permute(0, 2, 4, 1, 3, 5).reshape(-1, C, N, N) + laplacian_response = F.conv2d(reshaped_image, laplacian_kernel, padding=1, groups=C) + strength = laplacian_response.sum(dim=[1, 2, 3]).reshape( + ( + batch, + -1, + ) + ) + + # 对频率强度进行归一化,然后使用torch.multinomial进行采样 + probabilities = strength / (strength.max(dim=1)[0][:, None] + 1e-5) + ids_shuffle = torch.multinomial(probabilities.clip(1e-5, 1), length, replacement=False) + ids_keep = ids_shuffle[:, :len_keep] + ids_restore = torch.argsort(ids_shuffle, dim=1) + ids_removed = ids_shuffle[:, len_keep:] + + mask[:, :len_keep] = 0 + mask = torch.gather(mask, dim=1, index=ids_restore) + + return {"mask": mask, "ids_keep": ids_keep, "ids_restore": ids_restore, "ids_removed": ids_removed} + + +def mask_out_token(x, ids_keep, ids_removed=None): + """ + Mask out the tokens specified by ids_keep. + Args: + - x: input sequence, [N, L, D] + - ids_keep: indices of tokens to keep + return: + - x_masked: masked sequence + """ + N, L, D = x.shape # batch, length, dim + x_remain = torch.gather(x, dim=1, index=ids_keep.unsqueeze(-1).repeat(1, 1, D)) + if ids_removed is not None: + x_masked = torch.gather(x, dim=1, index=ids_removed.unsqueeze(-1).repeat(1, 1, D)) + return x_remain, x_masked + else: + return x_remain + + +def mask_tokens(x, mask_ratio): + """ + Perform per-sample random masking by per-sample shuffling. + Per-sample shuffling is done by argsort random noise. + x: [N, L, D], sequence + """ + N, L, D = x.shape # batch, length, dim + len_keep = int(L * (1 - mask_ratio)) + + noise = torch.rand(N, L, device=x.device) # noise in [0, 1] + + # sort noise for each sample + ids_shuffle = torch.argsort(noise, dim=1) # ascend: small is keep, large is remove + ids_restore = torch.argsort(ids_shuffle, dim=1) + + # keep the first subset + ids_keep = ids_shuffle[:, :len_keep] + x_masked = torch.gather(x, dim=1, index=ids_keep.unsqueeze(-1).repeat(1, 1, D)) + + # generate the binary mask: 0 is keep, 1 is remove + mask = torch.ones([N, L], device=x.device) + mask[:, :len_keep] = 0 + mask = torch.gather(mask, dim=1, index=ids_restore) + + return x_masked, mask, ids_restore + + +def unmask_tokens(x, ids_restore, mask_token): + # x: [N, T, D] if extras == 0 (i.e., no cls token) else x: [N, T+1, D] + mask_tokens = mask_token.repeat(x.shape[0], ids_restore.shape[1] - x.shape[1], 1) + x = torch.cat([x, mask_tokens], dim=1) + x = torch.gather(x, dim=1, index=ids_restore.unsqueeze(-1).repeat(1, 1, x.shape[2])) # unshuffle + return x + + +# Parse 'None' to None and others to float value +def parse_float_none(s): + assert isinstance(s, str) + return None if s == "None" else float(s) + + +# ---------------------------------------------------------------------------- +# Parse a comma separated list of numbers or ranges and return a list of ints. +# Example: '1,2,5-10' returns [1, 2, 5, 6, 7, 8, 9, 10] + + +def parse_int_list(s): + if isinstance(s, list): + return s + ranges = [] + range_re = re.compile(r"^(\d+)-(\d+)$") + for p in s.split(","): + m = range_re.match(p) + if m: + ranges.extend(range(int(m.group(1)), int(m.group(2)) + 1)) + else: + ranges.append(int(p)) + return ranges + + +def init_processes(fn, args): + """Initialize the distributed environment.""" + os.environ["MASTER_ADDR"] = args.master_address + os.environ["MASTER_PORT"] = str(random.randint(2000, 6000)) + print(f'MASTER_ADDR = {os.environ["MASTER_ADDR"]}') + print(f'MASTER_PORT = {os.environ["MASTER_PORT"]}') + torch.cuda.set_device(args.local_rank) + dist.init_process_group(backend="nccl", init_method="env://", rank=args.global_rank, world_size=args.global_size) + fn(args) + if args.global_size > 1: + cleanup() + + +def mprint(*args, **kwargs): + """ + Print only from rank 0. + """ + if dist.get_rank() == 0: + print(*args, **kwargs) + + +def cleanup(): + """ + End DDP training. + """ + dist.barrier() + mprint("Done!") + dist.barrier() + dist.destroy_process_group() + + +# ---------------------------------------------------------------------------- +# logging info. +class Logger: + """ + Redirect stderr to stdout, optionally print stdout to a file, + and optionally force flushing on both stdout and the file. + """ + + def __init__(self, file_name=None, file_mode="w", should_flush=True): + self.file = None + + if file_name is not None: + self.file = open(file_name, file_mode) + + self.should_flush = should_flush + self.stdout = sys.stdout + self.stderr = sys.stderr + + sys.stdout = self + sys.stderr = self + + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_value, traceback): + self.close() + + def write(self, text): + """Write text to stdout (and a file) and optionally flush.""" + if len(text) == 0: # workaround for a bug in VSCode debugger: sys.stdout.write(''); sys.stdout.flush() => crash + return + + if self.file is not None: + self.file.write(text) + + self.stdout.write(text) + + if self.should_flush: + self.flush() + + def flush(self): + """Flush written text to both stdout and a file, if open.""" + if self.file is not None: + self.file.flush() + + self.stdout.flush() + + def close(self): + """Flush, close possible files, and remove stdout/stderr mirroring.""" + self.flush() + + # if using multiple loggers, prevent closing in wrong order + if sys.stdout is self: + sys.stdout = self.stdout + if sys.stderr is self: + sys.stderr = self.stderr + + if self.file is not None: + self.file.close() + + +class StackedRandomGenerator: + def __init__(self, device, seeds): + super().__init__() + self.generators = [torch.Generator(device).manual_seed(int(seed) % (1 << 32)) for seed in seeds] + + def randn(self, size, **kwargs): + assert size[0] == len(self.generators) + return torch.stack([torch.randn(size[1:], generator=gen, **kwargs) for gen in self.generators]) + + def randn_like(self, input): + return self.randn(input.shape, dtype=input.dtype, layout=input.layout, device=input.device) + + def randint(self, *args, size, **kwargs): + assert size[0] == len(self.generators) + return torch.stack([torch.randint(*args, size=size[1:], generator=gen, **kwargs) for gen in self.generators]) + + +def prepare_prompt_ar(prompt, ratios, device="cpu", show=True): + # get aspect_ratio or ar + aspect_ratios = re.findall(r"--aspect_ratio\s+(\d+:\d+)", prompt) + ars = re.findall(r"--ar\s+(\d+:\d+)", prompt) + custom_hw = re.findall(r"--hw\s+(\d+:\d+)", prompt) + if show: + print("aspect_ratios:", aspect_ratios, "ars:", ars, "hws:", custom_hw) + prompt_clean = prompt.split("--aspect_ratio")[0].split("--ar")[0].split("--hw")[0] + if len(aspect_ratios) + len(ars) + len(custom_hw) == 0 and show: + print( + "Wrong prompt format. Set to default ar: 1. change your prompt into format '--ar h:w or --hw h:w' for correct generating" + ) + if len(aspect_ratios) != 0: + ar = float(aspect_ratios[0].split(":")[0]) / float(aspect_ratios[0].split(":")[1]) + elif len(ars) != 0: + ar = float(ars[0].split(":")[0]) / float(ars[0].split(":")[1]) + else: + ar = 1.0 + closest_ratio = min(ratios.keys(), key=lambda ratio: abs(float(ratio) - ar)) + if len(custom_hw) != 0: + custom_hw = [float(custom_hw[0].split(":")[0]), float(custom_hw[0].split(":")[1])] + else: + custom_hw = ratios[closest_ratio] + default_hw = ratios[closest_ratio] + prompt_show = f"prompt: {prompt_clean.strip()}\nSize: --ar {closest_ratio}, --bin hw {ratios[closest_ratio]}, --custom hw {custom_hw}" + return ( + prompt_clean, + prompt_show, + torch.tensor(default_hw, device=device)[None], + torch.tensor([float(closest_ratio)], device=device)[None], + torch.tensor(custom_hw, device=device)[None], + ) + + +def resize_and_crop_tensor(samples: torch.Tensor, new_width: int, new_height: int) -> torch.Tensor: + orig_height, orig_width = samples.shape[2], samples.shape[3] + + # Check if resizing is needed + if orig_height != new_height or orig_width != new_width: + ratio = max(new_height / orig_height, new_width / orig_width) + resized_width = int(orig_width * ratio) + resized_height = int(orig_height * ratio) + + # Resize + samples = F.interpolate(samples, size=(resized_height, resized_width), mode="bilinear", align_corners=False) + + # Center Crop + start_x = (resized_width - new_width) // 2 + end_x = start_x + new_width + start_y = (resized_height - new_height) // 2 + end_y = start_y + new_height + samples = samples[:, :, start_y:end_y, start_x:end_x] + + return samples + + +def resize_and_crop_img(img: Image, new_width, new_height): + orig_width, orig_height = img.size + + ratio = max(new_width / orig_width, new_height / orig_height) + resized_width = int(orig_width * ratio) + resized_height = int(orig_height * ratio) + + img = img.resize((resized_width, resized_height), Image.LANCZOS) + + left = (resized_width - new_width) / 2 + top = (resized_height - new_height) / 2 + right = (resized_width + new_width) / 2 + bottom = (resized_height + new_height) / 2 + + img = img.crop((left, top, right, bottom)) + + return img + + +def mask_feature(emb, mask): + if emb.shape[0] == 1: + keep_index = mask.sum().item() + return emb[:, :, :keep_index, :], keep_index + else: + masked_feature = emb * mask[:, None, :, None] + return masked_feature, emb.shape[2] + + +def val2list(x: list or tuple or any, repeat_time=1) -> list: # type: ignore + """Repeat `val` for `repeat_time` times and return the list or val if list/tuple.""" + if isinstance(x, (list, tuple)): + return list(x) + return [x for _ in range(repeat_time)] + + +def val2tuple(x: list or tuple or any, min_len: int = 1, idx_repeat: int = -1) -> tuple: # type: ignore + """Return tuple with min_len by repeating element at idx_repeat.""" + # convert to list first + x = val2list(x) + + # repeat elements if necessary + if len(x) > 0: + x[idx_repeat:idx_repeat] = [x[idx_repeat] for _ in range(min_len - len(x))] + + return tuple(x) + + +def get_same_padding(kernel_size: int or tuple[int, ...]) -> int or tuple[int, ...]: + if isinstance(kernel_size, tuple): + return tuple([get_same_padding(ks) for ks in kernel_size]) + else: + assert kernel_size % 2 > 0, f"kernel size {kernel_size} should be odd number" + return kernel_size // 2 diff --git a/Sana/nodes.py b/Sana/nodes.py new file mode 100644 index 0000000..f400962 --- /dev/null +++ b/Sana/nodes.py @@ -0,0 +1,223 @@ +import os +import json +import torch +import folder_paths + +from comfy.model_management import get_torch_device, soft_empty_cache, text_encoder_offload_device +from comfy import utils +from .conf import sana_conf, sana_res +from .loader import load_sana +from ..utils.dtype import string_to_dtype + +dtypes = [ + "auto", + "FP32", + "FP16", + "BF16" +] + +class SanaCheckpointLoader: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "ckpt_name": (folder_paths.get_filename_list("checkpoints"),), + "model": (list(sana_conf.keys()),), + "dtype": (dtypes,), + } + } + RETURN_TYPES = ("MODEL",) + RETURN_NAMES = ("model",) + FUNCTION = "load_checkpoint" + CATEGORY = "ExtraModels/Sana" + TITLE = "Sana Checkpoint Loader" + + def load_checkpoint(self, ckpt_name, model, dtype): + ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name) + model_conf = sana_conf[model] + model = load_sana( + model_path = ckpt_path, + model_conf = model_conf, + dtype = string_to_dtype(dtype, "text_encoder") + ) + return (model,) + + +class SanaResolutionSelect(): + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": (list(sana_res.keys()),), + "ratio": (list(sana_res["1024px"].keys()),{"default":"1.00"}), + } + } + RETURN_TYPES = ("INT","INT") + RETURN_NAMES = ("width","height") + FUNCTION = "get_res" + CATEGORY = "ExtraModels/Sana" + TITLE = "Sana Resolution Select" + + def get_res(self, model, ratio): + width, height = sana_res[model][ratio] + return (width,height) + + +class SanaResolutionCond: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "cond": ("CONDITIONING", ), + "width": ("INT", {"default": 1024.0, "min": 0, "max": 8192}), + "height": ("INT", {"default": 1024.0, "min": 0, "max": 8192}), + } + } + + RETURN_TYPES = ("CONDITIONING",) + RETURN_NAMES = ("cond",) + FUNCTION = "add_cond" + CATEGORY = "ExtraModels/Sana" + TITLE = "Sana Resolution Conditioning" + + def add_cond(self, cond, width, height): + for c in range(len(cond)): + cond[c][1].update({ + "img_hw": [[height, width]], + "aspect_ratio": [[height/width]], + }) + return (cond,) + + +class SanaTextEncode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "text": ("STRING", {"multiline": True}), + "preset_styles": (STYLE_NAMES,), + "GEMMA": ("GEMMA",), + } + } + + RETURN_TYPES = ("CONDITIONING",) + FUNCTION = "encode" + CATEGORY = "ExtraModels/Sana" + TITLE = "Sana Text Encode" + + def encode(self, text, preset_styles, GEMMA=None): + tokenizer = GEMMA["tokenizer"] + text_encoder = GEMMA["text_encoder"] + + # 应用预设样式 - 只使用正面提示词部分 + text, _ = apply_style(preset_styles, text) + + with torch.no_grad(): + # 处理正面提示词 + chi_prompt = "\n".join(preset_te_prompt) + full_prompt = chi_prompt + text + num_chi_tokens = len(tokenizer.encode(chi_prompt)) + max_length = num_chi_tokens + 300 - 2 # 减去[bos]和[_]标记 + + tokens = tokenizer( + [full_prompt], + max_length=max_length, + padding="max_length", + truncation=True, + return_tensors="pt" + ).to(text_encoder.device) + + select_idx = [0] + list(range(-300 + 1, 0)) + embs = text_encoder(tokens.input_ids, tokens.attention_mask)[0][:, None][:, :, select_idx] + emb_masks = tokens.attention_mask[:, select_idx] + # 利用emb_masks将有效的embs选出来,其他置零 + embs = embs * emb_masks.unsqueeze(-1) + # import IPython + # IPython.embed() + + return ([[embs, {}]], ) + +# 需要添加style相关的辅助函数 +style_list = [ + { + "name": "(No style)", + "prompt": "{prompt}", + "negative_prompt": "", + }, + { + "name": "Cinematic", + "prompt": "cinematic still {prompt} . emotional, harmonious, vignette, highly detailed, high budget, bokeh, " + "cinemascope, moody, epic, gorgeous, film grain, grainy", + "negative_prompt": "anime, cartoon, graphic, text, painting, crayon, graphite, abstract, glitch, deformed, mutated, ugly, disfigured", + }, + { + "name": "Photographic", + "prompt": "cinematic photo {prompt} . 35mm photograph, film, bokeh, professional, 4k, highly detailed", + "negative_prompt": "drawing, painting, crayon, sketch, graphite, impressionist, noisy, blurry, soft, deformed, ugly", + }, + { + "name": "Anime", + "prompt": "anime artwork {prompt} . anime style, key visual, vibrant, studio anime, highly detailed", + "negative_prompt": "photo, deformed, black and white, realism, disfigured, low contrast", + }, + { + "name": "Manga", + "prompt": "manga style {prompt} . vibrant, high-energy, detailed, iconic, Japanese comic style", + "negative_prompt": "ugly, deformed, noisy, blurry, low contrast, realism, photorealistic, Western comic style", + }, + { + "name": "Digital Art", + "prompt": "concept art {prompt} . digital artwork, illustrative, painterly, matte painting, highly detailed", + "negative_prompt": "photo, photorealistic, realism, ugly", + }, + { + "name": "Pixel art", + "prompt": "pixel-art {prompt} . low-res, blocky, pixel art style, 8-bit graphics", + "negative_prompt": "sloppy, messy, blurry, noisy, highly detailed, ultra textured, photo, realistic", + }, + { + "name": "Fantasy art", + "prompt": "ethereal fantasy concept art of {prompt} . magnificent, celestial, ethereal, painterly, epic, " + "majestic, magical, fantasy art, cover art, dreamy", + "negative_prompt": "photographic, realistic, realism, 35mm film, dslr, cropped, frame, text, deformed, " + "glitch, noise, noisy, off-center, deformed, cross-eyed, closed eyes, bad anatomy, ugly, " + "disfigured, sloppy, duplicate, mutated, black and white", + }, + { + "name": "Neonpunk", + "prompt": "neonpunk style {prompt} . cyberpunk, vaporwave, neon, vibes, vibrant, stunningly beautiful, crisp, " + "detailed, sleek, ultramodern, magenta highlights, dark purple shadows, high contrast, cinematic, " + "ultra detailed, intricate, professional", + "negative_prompt": "painting, drawing, illustration, glitch, deformed, mutated, cross-eyed, ugly, disfigured", + }, + { + "name": "3D Model", + "prompt": "professional 3d model {prompt} . octane render, highly detailed, volumetric, dramatic lighting", + "negative_prompt": "ugly, deformed, noisy, low poly, blurry, painting", + }, +] + +styles = {k["name"]: (k["prompt"], k["negative_prompt"]) for k in style_list} +STYLE_NAMES = list(styles.keys()) + +def apply_style(style_name: str, positive: str, negative: str = "") -> tuple[str, str]: + p, n = styles.get(style_name, styles[style_name]) + if not negative: + negative = "" + return p.replace("{prompt}", positive), n + negative + +preset_te_prompt = ['Given a user prompt, generate an "Enhanced prompt" that provides detailed visual descriptions suitable for image generation. Evaluate the level of detail in the user prompt:', '- If the prompt is simple, focus on adding specifics about colors, shapes, sizes, textures, and spatial relationships to create vivid and concrete scenes.', '- If the prompt is already detailed, refine and enhance the existing details slightly without overcomplicating.', 'Here are examples of how to transform or refine prompts:', '- User Prompt: A cat sleeping -> Enhanced: A small, fluffy white cat curled up in a round shape, sleeping peacefully on a warm sunny windowsill, surrounded by pots of blooming red flowers.', '- User Prompt: A busy city street -> Enhanced: A bustling city street scene at dusk, featuring glowing street lamps, a diverse crowd of people in colorful clothing, and a double-decker bus passing by towering glass skyscrapers.', 'Please generate only the enhanced description for the prompt below and avoid including any additional commentary or evaluations:', 'User Prompt: '] + +NODE_CLASS_MAPPINGS = { + "SanaCheckpointLoader" : SanaCheckpointLoader, + "SanaResolutionSelect" : SanaResolutionSelect, + "SanaTextEncode" : SanaTextEncode, + "SanaResolutionCond" : SanaResolutionCond, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "Sana Checkpoint Loader": "SanaCheckpointLoader", + "Sana Resolution Select": "SanaResolutionSelect", + "Sana Text Encoder": "SanaTextEncode", + "Sana Resolution Cond": "SanaResolutionCond", +} diff --git a/VAE/nodes.py b/VAE/nodes.py index b3639ae..00be226 100644 --- a/VAE/nodes.py +++ b/VAE/nodes.py @@ -1,4 +1,6 @@ import folder_paths +import torch +import comfy from .conf import vae_conf from .loader import EXVAE @@ -12,6 +14,8 @@ dtypes = [ "BF16" ] +MAX_RESOLUTION=16384 + class ExtraVAELoader: @classmethod def INPUT_TYPES(s): @@ -33,6 +37,34 @@ class ExtraVAELoader: vae = EXVAE(model_path, model_conf, string_to_dtype(dtype, "vae")) return (vae,) + +class EmptyDCAELatentImage: + def __init__(self): + self.device = comfy.model_management.intermediate_device() + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "width": ("INT", {"default": 512, "min": 16, "max": MAX_RESOLUTION, "step": 8, "tooltip": "The width of the latent images in pixels."}), + "height": ("INT", {"default": 512, "min": 16, "max": MAX_RESOLUTION, "step": 8, "tooltip": "The height of the latent images in pixels."}), + "batch_size": ("INT", {"default": 1, "min": 1, "max": 4096, "tooltip": "The number of latent images in the batch."}) + } + } + RETURN_TYPES = ("LATENT",) + OUTPUT_TOOLTIPS = ("The empty latent image batch.",) + FUNCTION = "generate" + TITLE = "Empty DCAE Latent Image" + + CATEGORY = "latent" + DESCRIPTION = "Create a new batch of empty latent images to be denoised via sampling." + + def generate(self, width, height, batch_size=1): + latent = torch.zeros([batch_size, 32, height // 32, width // 32], device=self.device) + return ({"samples":latent}, ) + + NODE_CLASS_MAPPINGS = { "ExtraVAELoader" : ExtraVAELoader, + "EmptyDCAELatentImage" : EmptyDCAELatentImage, } diff --git a/__init__.py b/__init__.py index b4260bb..1fff84c 100644 --- a/__init__.py +++ b/__init__.py @@ -38,5 +38,14 @@ else: from .utils.nodes import NODE_CLASS_MAPPINGS as Extra_Nodes NODE_CLASS_MAPPINGS.update(Extra_Nodes) + # Sana + from .Sana.nodes import NODE_CLASS_MAPPINGS as Sana_Nodes + NODE_CLASS_MAPPINGS.update(Sana_Nodes) + + # Gemma + from .Gemma.nodes import NODE_CLASS_MAPPINGS as Gemma_Nodes + NODE_CLASS_MAPPINGS.update(Gemma_Nodes) + NODE_DISPLAY_NAME_MAPPINGS = {k:v.TITLE for k,v in NODE_CLASS_MAPPINGS.items()} __all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] +