first run sucessfull with text encoder mask bug not fix;
This commit is contained in:
+128
@@ -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",
|
||||
}
|
||||
@@ -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"],
|
||||
})
|
||||
@@ -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
|
||||
+100
@@ -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
|
||||
+146
@@ -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
|
||||
@@ -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")
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
+223
@@ -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",
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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']
|
||||
|
||||
|
||||
Reference in New Issue
Block a user