fix the conversation and run success;
This commit is contained in:
+15
-29
@@ -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
|
||||
|
||||
-146
@@ -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
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
+2
-497
@@ -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:
|
||||
|
||||
+11
-82
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user