diff --git a/Gemma/nodes.py b/Gemma/nodes.py index 55e0883..9781ea1 100644 --- a/Gemma/nodes.py +++ b/Gemma/nodes.py @@ -6,17 +6,11 @@ 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 +tenc_root = ( + folder_paths.folder_names_and_paths.get( + "text_encoders", + folder_paths.folder_names_and_paths.get("clip", [[], set()]) + ) ) dtypes = [ diff --git a/Sana/diffusers_convert.py b/Sana/diffusers_convert.py index 312ea9d..dd15b95 100644 --- a/Sana/diffusers_convert.py +++ b/Sana/diffusers_convert.py @@ -1,20 +1,8 @@ # 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 +# https://github.com/NVlabs/Sana/blob/main/tools/convert_sana_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()) @@ -30,7 +18,7 @@ def get_lora_depth(state_dict): return cnt def get_conversion_map(state_dict): - conversion_map = [ # main SD conversion map (PixArt reference, HF Diffusers) + conversion_map = [ # main SD conversion map (Sana reference, HF Diffusers) # Patch embeddings ("x_embedder.proj.weight", "pos_embed.proj.weight"), ("x_embedder.proj.bias", "pos_embed.proj.bias"), @@ -82,10 +70,7 @@ def find_prefix(state_dict, target_key): 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) + 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} @@ -109,16 +94,17 @@ def convert_state_dict(state_dict): 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(f"Sana: 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(f"Sana: UNET conversion has missing keys!") print(missing) return new_state_dict # Same as above but for LoRA weights: +# TODO: Not used yet, need to support LoRA for Sana def convert_lora_state_dict(state_dict, peft=True): # koyha rep_ak = lambda x: x.replace(".weight", ".lora_down.weight") @@ -137,18 +123,18 @@ def convert_lora_state_dict(state_dict, peft=True): rep_pp = lambda x: x.replace(".", "_")[:-7] + ".alpha" prefix = "lora_transformer_" - t5_marker = "lora_te_encoder" - t5_keys = [] + gemma_marker = "lora_te_encoder" + gemma_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") + elif gemma_marker in key: + gemma_keys.append(state_dict.pop(key)) + if len(gemma_keys) > 0: + print(f"Text Encoder not supported for Sana LoRA, ignoring {len(gemma_keys)} keys") cmap = [] - cmap_unet = get_conversion_map(state_dict) + conversion_map_ms # todo: 512 model + cmap_unet = get_conversion_map(state_dict) # todo: 512 model for k, v in cmap_unet: if v.endswith(".weight"): cmap.append((rep_ak(k), rep_ap(v))) @@ -213,11 +199,11 @@ def convert_lora_state_dict(state_dict, peft=True): pass if len(matched) < len(state_dict): - print(f"PixArt: LoRA conversion has leftover keys! ({len(matched)} vs {len(state_dict)})") + print(f"Sana: 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(f"Sana: LoRA conversion has missing keys! (probably)") print(missing) return new_state_dict diff --git a/Sana/lora.py b/Sana/lora.py deleted file mode 100644 index fca5931..0000000 --- a/Sana/lora.py +++ /dev/null @@ -1,146 +0,0 @@ -import os -import copy -import json -import torch -import comfy.lora -import comfy.model_management -from comfy.model_patcher import ModelPatcher -from .diffusers_convert import convert_lora_state_dict - -class EXM_PixArt_ModelPatcher(ModelPatcher): - def calculate_weight(self, patches, weight, key): - """ - This is almost the same as the comfy function, but stripped down to just the LoRA patch code. - The problem with the original code is the q/k/v keys being combined into one for the attention. - In the diffusers code, they're treated as separate keys, but in the reference code they're recombined (q+kv|qkv). - This means, for example, that the [1152,1152] weights become [3456,1152] in the state dict. - The issue with this is that the LoRA weights are [128,1152],[1152,128] and become [384,1162],[3456,128] instead. - - This is the best thing I could think of that would fix that, but it's very fragile. - - Check key shape to determine if it needs the fallback logic - - Cut the input into parts based on the shape (undoing the torch.cat) - - Do the matrix multiplication logic - - Recombine them to match the expected shape - """ - for p in patches: - alpha = p[0] - v = p[1] - strength_model = p[2] - if strength_model != 1.0: - weight *= strength_model - - if isinstance(v, list): - v = (self.calculate_weight(v[1:], v[0].clone(), key), ) - - if len(v) == 2: - patch_type = v[0] - v = v[1] - - if patch_type == "lora": - mat1 = comfy.model_management.cast_to_device(v[0], weight.device, torch.float32) - mat2 = comfy.model_management.cast_to_device(v[1], weight.device, torch.float32) - if v[2] is not None: - alpha *= v[2] / mat2.shape[0] - try: - mat1 = mat1.flatten(start_dim=1) - mat2 = mat2.flatten(start_dim=1) - - ch1 = mat1.shape[0] // mat2.shape[1] - ch2 = mat2.shape[0] // mat1.shape[1] - ### Fallback logic for shape mismatch ### - if mat1.shape[0] != mat2.shape[1] and ch1 == ch2 and (mat1.shape[0]/mat2.shape[1])%1 == 0: - mat1 = mat1.chunk(ch1, dim=0) - mat2 = mat2.chunk(ch1, dim=0) - weight += torch.cat( - [alpha * torch.mm(mat1[x], mat2[x]) for x in range(ch1)], - dim=0, - ).reshape(weight.shape).type(weight.dtype) - else: - weight += (alpha * torch.mm(mat1, mat2)).reshape(weight.shape).type(weight.dtype) - except Exception as e: - print("ERROR", key, e) - return weight - - def clone(self): - n = EXM_PixArt_ModelPatcher(self.model, self.load_device, self.offload_device, self.size, self.current_device, weight_inplace_update=self.weight_inplace_update) - n.patches = {} - for k in self.patches: - n.patches[k] = self.patches[k][:] - - n.object_patches = self.object_patches.copy() - n.model_options = copy.deepcopy(self.model_options) - n.model_keys = self.model_keys - return n - -def replace_model_patcher(model): - n = EXM_PixArt_ModelPatcher( - model = model.model, - size = model.size, - load_device = model.load_device, - offload_device = model.offload_device, - weight_inplace_update = model.weight_inplace_update, - ) - n.patches = {} - for k in model.patches: - n.patches[k] = model.patches[k][:] - - n.object_patches = model.object_patches.copy() - n.model_options = copy.deepcopy(model.model_options) - return n - -def find_peft_alpha(path): - def load_json(json_path): - with open(json_path) as f: - data = json.load(f) - alpha = data.get("lora_alpha") - alpha = alpha or data.get("alpha") - if not alpha: - print(" Found config but `lora_alpha` is missing!") - else: - print(f" Found config at {json_path} [alpha:{alpha}]") - return alpha - - # For some weird reason peft doesn't include the alpha in the actual model - print("PixArt: Warning! This is a PEFT LoRA. Trying to find config...") - files = [ - f"{os.path.splitext(path)[0]}.json", - f"{os.path.splitext(path)[0]}.config.json", - os.path.join(os.path.dirname(path),"adapter_config.json"), - ] - for file in files: - if os.path.isfile(file): - return load_json(file) - - print(" Missing config/alpha! assuming alpha of 8. Consider converting it/adding a config json to it.") - return 8.0 - -def load_pixart_lora(model, lora, lora_path, strength): - k_back = lambda x: x.replace(".lora_up.weight", "") - # need to convert the actual weights for this to work. - if any(True for x in lora.keys() if x.endswith("adaln_single.linear.lora_A.weight")): - lora = convert_lora_state_dict(lora, peft=True) - alpha = find_peft_alpha(lora_path) - lora.update({f"{k_back(x)}.alpha":torch.tensor(alpha) for x in lora.keys() if "lora_up" in x}) - else: # OneTrainer - lora = convert_lora_state_dict(lora, peft=False) - - key_map = {k_back(x):f"diffusion_model.{k_back(x)}.weight" for x in lora.keys() if "lora_up" in x} # fake - - loaded = comfy.lora.load_lora(lora, key_map) - if model is not None: - # switch to custom model patcher when using LoRAs - if isinstance(model, EXM_PixArt_ModelPatcher): - new_modelpatcher = model.clone() - else: - new_modelpatcher = replace_model_patcher(model) - k = new_modelpatcher.add_patches(loaded, strength) - else: - k = () - new_modelpatcher = None - - k = set(k) - for x in loaded: - if (x not in k): - print("NOT LOADED", x) - - return new_modelpatcher diff --git a/Sana/models/sana.py b/Sana/models/sana.py index 9da2d71..81c2211 100644 --- a/Sana/models/sana.py +++ b/Sana/models/sana.py @@ -226,8 +226,6 @@ class Sana(nn.Module): ) 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. diff --git a/Sana/models/sana_blocks.py b/Sana/models/sana_blocks.py index 31ac821..cdd37e4 100644 --- a/Sana/models/sana_blocks.py +++ b/Sana/models/sana_blocks.py @@ -16,10 +16,8 @@ # 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 diff --git a/Sana/models/utils.py b/Sana/models/utils.py index d74db3b..f6965d2 100644 --- a/Sana/models/utils.py +++ b/Sana/models/utils.py @@ -14,21 +14,12 @@ # # 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 +from typing import Union, Tuple 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): @@ -44,25 +35,6 @@ 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): @@ -99,478 +71,12 @@ def checkpoint_sequential(functions, step, input, *args, **kwargs): 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 @@ -582,8 +88,7 @@ def val2tuple(x: list or tuple or any, min_len: int = 1, idx_repeat: int = -1) - return tuple(x) - -def get_same_padding(kernel_size: int or tuple[int, ...]) -> int or tuple[int, ...]: +def get_same_padding(kernel_size: Union[int, Tuple[int, ...]]) -> Union[int, Tuple[int, ...]]: if isinstance(kernel_size, tuple): return tuple([get_same_padding(ks) for ks in kernel_size]) else: diff --git a/Sana/nodes.py b/Sana/nodes.py index 57b5928..8291c55 100644 --- a/Sana/nodes.py +++ b/Sana/nodes.py @@ -1,13 +1,8 @@ -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", @@ -93,7 +88,6 @@ class SanaTextEncode: return { "required": { "text": ("STRING", {"multiline": True}), - "preset_styles": (STYLE_NAMES,), "GEMMA": ("GEMMA",), } } @@ -103,19 +97,15 @@ class SanaTextEncode: CATEGORY = "ExtraModels/Sana" TITLE = "Sana Text Encode" - def encode(self, text, preset_styles, GEMMA=None): + def encode(self, text, 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]和[_]标记 + max_length = num_chi_tokens + 300 - 2 tokens = tokenizer( [full_prompt], @@ -128,82 +118,21 @@ class SanaTextEncode: 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) 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", - }, +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: ' ] -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,