Use native ops for PixArt
This commit is contained in:
+1
-2
@@ -93,8 +93,7 @@ def model_config_from_unet(sd):
|
|||||||
config["pe_interpolation"] = 1
|
config["pe_interpolation"] = 1
|
||||||
model_config = PixArtConfig(config)
|
model_config = PixArtConfig(config)
|
||||||
model_config.unet_class = model_class
|
model_config.unet_class = model_class
|
||||||
logging.info(f"Detected PixArt model as [{model_class}]")
|
logging.debug(f"PixArt config: {model_class}\n{config}")
|
||||||
logging.info(f"PixArt config:\n{config}")
|
|
||||||
return model_config
|
return model_config
|
||||||
|
|
||||||
resolutions = {
|
resolutions = {
|
||||||
|
|||||||
+3
-9
@@ -19,14 +19,8 @@ class PixArtModel(comfy.model_base.BaseModel):
|
|||||||
def extra_conds(self, **kwargs):
|
def extra_conds(self, **kwargs):
|
||||||
out = super().extra_conds(**kwargs)
|
out = super().extra_conds(**kwargs)
|
||||||
|
|
||||||
img_hw = kwargs.get("img_hw", None)
|
for name in ["width", "height", "aspect_ratio", "img_hw"]: # TODO: remove last one
|
||||||
if img_hw is not None:
|
out[name] = comfy.conds.CONDRegular(torch.tensor(name))
|
||||||
out["img_hw"] = comfy.conds.CONDRegular(torch.tensor(img_hw))
|
|
||||||
|
|
||||||
aspect_ratio = kwargs.get("aspect_ratio", None)
|
|
||||||
if aspect_ratio is not None:
|
|
||||||
out["aspect_ratio"] = comfy.conds.CONDRegular(torch.tensor(aspect_ratio))
|
|
||||||
|
|
||||||
return out
|
return out
|
||||||
|
|
||||||
def load_pixart_state_dict(sd, model_options={}):
|
def load_pixart_state_dict(sd, model_options={}):
|
||||||
@@ -49,7 +43,7 @@ def load_pixart_state_dict(sd, model_options={}):
|
|||||||
load_device = model_management.get_torch_device()
|
load_device = model_management.get_torch_device()
|
||||||
offload_device = comfy.model_management.unet_offload_device()
|
offload_device = comfy.model_management.unet_offload_device()
|
||||||
|
|
||||||
dtype = model_options.get("dtype", torch.float16) # TODO: fix this
|
dtype = model_options.get("dtype", None)
|
||||||
weight_dtype = comfy.utils.weight_dtype(sd)
|
weight_dtype = comfy.utils.weight_dtype(sd)
|
||||||
unet_weight_dtype = list(model_config.supported_inference_dtypes)
|
unet_weight_dtype = list(model_config.supported_inference_dtypes)
|
||||||
|
|
||||||
|
|||||||
+77
-73
@@ -12,9 +12,10 @@ import math
|
|||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from timm.models.vision_transformer import Mlp, Attention as Attention_
|
|
||||||
from einops import rearrange
|
from einops import rearrange
|
||||||
|
|
||||||
|
from .utils import to_2tuple
|
||||||
|
|
||||||
sdpa_32b = None
|
sdpa_32b = None
|
||||||
Q_4GB_LIMIT = 32000000
|
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."""
|
"""If q is greater than this, the operation will likely require >4GB VRAM, which will fail on Intel Arc Alchemist GPUs without a workaround."""
|
||||||
@@ -44,7 +45,7 @@ def t2i_modulate(x, shift, scale):
|
|||||||
return x * (1 + scale) + shift
|
return x * (1 + scale) + shift
|
||||||
|
|
||||||
class MultiHeadCrossAttention(nn.Module):
|
class MultiHeadCrossAttention(nn.Module):
|
||||||
def __init__(self, d_model, num_heads, attn_drop=0., proj_drop=0., **block_kwargs):
|
def __init__(self, d_model, num_heads, attn_drop=0., proj_drop=0., dtype=None, device=None, operations=None, **block_kwargs):
|
||||||
super(MultiHeadCrossAttention, self).__init__()
|
super(MultiHeadCrossAttention, self).__init__()
|
||||||
assert d_model % num_heads == 0, "d_model must be divisible by num_heads"
|
assert d_model % num_heads == 0, "d_model must be divisible by num_heads"
|
||||||
|
|
||||||
@@ -52,10 +53,10 @@ class MultiHeadCrossAttention(nn.Module):
|
|||||||
self.num_heads = num_heads
|
self.num_heads = num_heads
|
||||||
self.head_dim = d_model // num_heads
|
self.head_dim = d_model // num_heads
|
||||||
|
|
||||||
self.q_linear = nn.Linear(d_model, d_model)
|
self.q_linear = operations.Linear(d_model, d_model, dtype=dtype, device=device)
|
||||||
self.kv_linear = nn.Linear(d_model, d_model*2)
|
self.kv_linear = operations.Linear(d_model, d_model*2, dtype=dtype, device=device)
|
||||||
self.attn_drop = nn.Dropout(attn_drop)
|
self.attn_drop = nn.Dropout(attn_drop)
|
||||||
self.proj = nn.Linear(d_model, d_model)
|
self.proj = operations.Linear(d_model, d_model, dtype=dtype, device=device)
|
||||||
self.proj_drop = nn.Dropout(proj_drop)
|
self.proj_drop = nn.Dropout(proj_drop)
|
||||||
|
|
||||||
def forward(self, x, cond, mask=None):
|
def forward(self, x, cond, mask=None):
|
||||||
@@ -111,7 +112,7 @@ class MultiHeadCrossAttention(nn.Module):
|
|||||||
return x
|
return x
|
||||||
|
|
||||||
|
|
||||||
class AttentionKVCompress(Attention_):
|
class AttentionKVCompress(nn.Module):
|
||||||
"""Multi-head Attention block with KV token compression and qk norm."""
|
"""Multi-head Attention block with KV token compression and qk norm."""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -122,6 +123,9 @@ class AttentionKVCompress(Attention_):
|
|||||||
sampling='conv',
|
sampling='conv',
|
||||||
sr_ratio=1,
|
sr_ratio=1,
|
||||||
qk_norm=False,
|
qk_norm=False,
|
||||||
|
dtype=None,
|
||||||
|
device=None,
|
||||||
|
operations=None,
|
||||||
**block_kwargs,
|
**block_kwargs,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
@@ -130,19 +134,26 @@ class AttentionKVCompress(Attention_):
|
|||||||
num_heads (int): Number of attention heads.
|
num_heads (int): Number of attention heads.
|
||||||
qkv_bias (bool: If True, add a learnable bias to query, key, value.
|
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)
|
super().__init__()
|
||||||
|
assert dim % num_heads == 0, 'dim should be divisible by num_heads'
|
||||||
|
self.num_heads = num_heads
|
||||||
|
self.head_dim = dim // num_heads
|
||||||
|
self.scale = self.head_dim ** -0.5
|
||||||
|
|
||||||
|
self.qkv = operations.Linear(dim, dim * 3, bias=qkv_bias, dtype=dtype, device=device)
|
||||||
|
self.proj = operations.Linear(dim, dim, dtype=dtype, device=device)
|
||||||
|
|
||||||
self.sampling=sampling # ['conv', 'ave', 'uniform', 'uniform_every']
|
self.sampling=sampling # ['conv', 'ave', 'uniform', 'uniform_every']
|
||||||
self.sr_ratio = sr_ratio
|
self.sr_ratio = sr_ratio
|
||||||
if sr_ratio > 1 and sampling == 'conv':
|
if sr_ratio > 1 and sampling == 'conv':
|
||||||
# Avg Conv Init.
|
# Avg Conv Init.
|
||||||
self.sr = nn.Conv2d(dim, dim, groups=dim, kernel_size=sr_ratio, stride=sr_ratio)
|
self.sr = operations.Conv2d(dim, dim, groups=dim, kernel_size=sr_ratio, stride=sr_ratio, dtype=dtype, device=device)
|
||||||
self.sr.weight.data.fill_(1/sr_ratio**2)
|
# self.sr.weight.data.fill_(1/sr_ratio**2)
|
||||||
self.sr.bias.data.zero_()
|
# self.sr.bias.data.zero_()
|
||||||
self.norm = nn.LayerNorm(dim)
|
self.norm = operations.LayerNorm(dim, dtype=dtype, device=device)
|
||||||
if qk_norm:
|
if qk_norm:
|
||||||
self.q_norm = nn.LayerNorm(dim)
|
self.q_norm = operations.LayerNorm(dim, dtype=dtype, device=device)
|
||||||
self.k_norm = nn.LayerNorm(dim)
|
self.k_norm = operations.LayerNorm(dim, dtype=dtype, device=device)
|
||||||
else:
|
else:
|
||||||
self.q_norm = nn.Identity()
|
self.q_norm = nn.Identity()
|
||||||
self.k_norm = nn.Identity()
|
self.k_norm = nn.Identity()
|
||||||
@@ -204,14 +215,12 @@ class AttentionKVCompress(Attention_):
|
|||||||
if model_management.xformers_enabled():
|
if model_management.xformers_enabled():
|
||||||
x = xformers.ops.memory_efficient_attention(
|
x = xformers.ops.memory_efficient_attention(
|
||||||
q, k, v,
|
q, k, v,
|
||||||
p=self.attn_drop.p,
|
p=0,
|
||||||
attn_bias=attn_bias
|
attn_bias=attn_bias
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
q, k, v = map(lambda t: t.transpose(1, 2),(q, k, v),)
|
q, k, v = map(lambda t: t.transpose(1, 2),(q, k, v),)
|
||||||
|
p = 0
|
||||||
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:
|
if sdpa_32b is not None and (q.element_size() * q.nelement()) > Q_4GB_LIMIT:
|
||||||
sdpa = sdpa_32b
|
sdpa = sdpa_32b
|
||||||
else:
|
else:
|
||||||
@@ -224,30 +233,6 @@ class AttentionKVCompress(Attention_):
|
|||||||
).transpose(1, 2).contiguous()
|
).transpose(1, 2).contiguous()
|
||||||
x = x.view(B, N, C)
|
x = x.view(B, N, C)
|
||||||
x = self.proj(x)
|
x = self.proj(x)
|
||||||
x = self.proj_drop(x)
|
|
||||||
return x
|
|
||||||
|
|
||||||
|
|
||||||
#################################################################################
|
|
||||||
# AMP attention with fp32 softmax to fix loss NaN problem during training #
|
|
||||||
#################################################################################
|
|
||||||
class Attention(Attention_):
|
|
||||||
def forward(self, x):
|
|
||||||
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)
|
|
||||||
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
|
return x
|
||||||
|
|
||||||
|
|
||||||
@@ -256,13 +241,13 @@ class FinalLayer(nn.Module):
|
|||||||
The final layer of PixArt.
|
The final layer of PixArt.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, hidden_size, patch_size, out_channels):
|
def __init__(self, hidden_size, patch_size, out_channels, dtype=None, device=None, operations=None):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
self.norm_final = operations.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device)
|
||||||
self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True)
|
self.linear = operations.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True, dtype=dtype, device=device)
|
||||||
self.adaLN_modulation = nn.Sequential(
|
self.adaLN_modulation = nn.Sequential(
|
||||||
nn.SiLU(),
|
nn.SiLU(),
|
||||||
nn.Linear(hidden_size, 2 * hidden_size, bias=True)
|
operations.Linear(hidden_size, 2 * hidden_size, bias=True, dtype=dtype, device=device)
|
||||||
)
|
)
|
||||||
|
|
||||||
def forward(self, x, c):
|
def forward(self, x, c):
|
||||||
@@ -271,23 +256,23 @@ class FinalLayer(nn.Module):
|
|||||||
x = self.linear(x)
|
x = self.linear(x)
|
||||||
return x
|
return x
|
||||||
|
|
||||||
|
|
||||||
class T2IFinalLayer(nn.Module):
|
class T2IFinalLayer(nn.Module):
|
||||||
"""
|
"""
|
||||||
The final layer of PixArt.
|
The final layer of PixArt.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, hidden_size, patch_size, out_channels):
|
def __init__(self, hidden_size, patch_size, out_channels, dtype=None, device=None, operations=None):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
self.norm_final = operations.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device)
|
||||||
self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True)
|
self.linear = operations.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True, dtype=dtype, device=device)
|
||||||
self.scale_shift_table = nn.Parameter(torch.randn(2, hidden_size) / hidden_size ** 0.5)
|
self.scale_shift_table = nn.Parameter(torch.randn(2, hidden_size) / hidden_size ** 0.5)
|
||||||
self.out_channels = out_channels
|
self.out_channels = out_channels
|
||||||
|
|
||||||
def forward(self, x, t):
|
def forward(self, x, t):
|
||||||
|
dtype = x.dtype
|
||||||
shift, scale = (self.scale_shift_table[None] + t[:, None]).chunk(2, dim=1)
|
shift, scale = (self.scale_shift_table[None] + t[:, None]).chunk(2, dim=1)
|
||||||
x = t2i_modulate(self.norm_final(x), shift, scale)
|
x = t2i_modulate(self.norm_final(x), shift, scale)
|
||||||
x = self.linear(x)
|
x = self.linear(x.to(dtype))
|
||||||
return x
|
return x
|
||||||
|
|
||||||
|
|
||||||
@@ -296,13 +281,13 @@ class MaskFinalLayer(nn.Module):
|
|||||||
The final layer of PixArt.
|
The final layer of PixArt.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, final_hidden_size, c_emb_size, patch_size, out_channels):
|
def __init__(self, final_hidden_size, c_emb_size, patch_size, out_channels, dtype=None, device=None, operations=None):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.norm_final = nn.LayerNorm(final_hidden_size, elementwise_affine=False, eps=1e-6)
|
self.norm_final = operations.LayerNorm(final_hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device)
|
||||||
self.linear = nn.Linear(final_hidden_size, patch_size * patch_size * out_channels, bias=True)
|
self.linear = operations.Linear(final_hidden_size, patch_size * patch_size * out_channels, bias=True, dtype=dtype, device=device)
|
||||||
self.adaLN_modulation = nn.Sequential(
|
self.adaLN_modulation = nn.Sequential(
|
||||||
nn.SiLU(),
|
nn.SiLU(),
|
||||||
nn.Linear(c_emb_size, 2 * final_hidden_size, bias=True)
|
operations.Linear(c_emb_size, 2 * final_hidden_size, bias=True, dtype=dtype, device=device)
|
||||||
)
|
)
|
||||||
def forward(self, x, t):
|
def forward(self, x, t):
|
||||||
shift, scale = self.adaLN_modulation(t).chunk(2, dim=1)
|
shift, scale = self.adaLN_modulation(t).chunk(2, dim=1)
|
||||||
@@ -316,13 +301,13 @@ class DecoderLayer(nn.Module):
|
|||||||
The final layer of PixArt.
|
The final layer of PixArt.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, hidden_size, decoder_hidden_size):
|
def __init__(self, hidden_size, decoder_hidden_size, dtype=None, device=None, operations=None):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.norm_decoder = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
self.norm_decoder = operations.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device)
|
||||||
self.linear = nn.Linear(hidden_size, decoder_hidden_size, bias=True)
|
self.linear = operations.Linear(hidden_size, decoder_hidden_size, bias=True, dtype=dtype, device=device)
|
||||||
self.adaLN_modulation = nn.Sequential(
|
self.adaLN_modulation = nn.Sequential(
|
||||||
nn.SiLU(),
|
nn.SiLU(),
|
||||||
nn.Linear(hidden_size, 2 * hidden_size, bias=True)
|
operations.Linear(hidden_size, 2 * hidden_size, bias=True, dtype=dtype, device=device)
|
||||||
)
|
)
|
||||||
def forward(self, x, t):
|
def forward(self, x, t):
|
||||||
shift, scale = self.adaLN_modulation(t).chunk(2, dim=1)
|
shift, scale = self.adaLN_modulation(t).chunk(2, dim=1)
|
||||||
@@ -339,12 +324,12 @@ class TimestepEmbedder(nn.Module):
|
|||||||
Embeds scalar timesteps into vector representations.
|
Embeds scalar timesteps into vector representations.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, hidden_size, frequency_embedding_size=256):
|
def __init__(self, hidden_size, frequency_embedding_size=256, dtype=None, device=None, operations=None):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.mlp = nn.Sequential(
|
self.mlp = nn.Sequential(
|
||||||
nn.Linear(frequency_embedding_size, hidden_size, bias=True),
|
operations.Linear(frequency_embedding_size, hidden_size, bias=True, dtype=dtype, device=device),
|
||||||
nn.SiLU(),
|
nn.SiLU(),
|
||||||
nn.Linear(hidden_size, hidden_size, bias=True),
|
operations.Linear(hidden_size, hidden_size, bias=True, dtype=dtype, device=device),
|
||||||
)
|
)
|
||||||
self.frequency_embedding_size = frequency_embedding_size
|
self.frequency_embedding_size = frequency_embedding_size
|
||||||
|
|
||||||
@@ -368,9 +353,9 @@ class TimestepEmbedder(nn.Module):
|
|||||||
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
|
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
|
||||||
return embedding
|
return embedding
|
||||||
|
|
||||||
def forward(self, t):
|
def forward(self, t, dtype):
|
||||||
t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
|
t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
|
||||||
t_emb = self.mlp(t_freq.to(t.dtype))
|
t_emb = self.mlp(t_freq.to(dtype))
|
||||||
return t_emb
|
return t_emb
|
||||||
|
|
||||||
|
|
||||||
@@ -379,12 +364,12 @@ class SizeEmbedder(TimestepEmbedder):
|
|||||||
Embeds scalar timesteps into vector representations.
|
Embeds scalar timesteps into vector representations.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, hidden_size, frequency_embedding_size=256):
|
def __init__(self, hidden_size, frequency_embedding_size=256, dtype=None, device=None, operations=None):
|
||||||
super().__init__(hidden_size=hidden_size, frequency_embedding_size=frequency_embedding_size)
|
super().__init__(hidden_size=hidden_size, frequency_embedding_size=frequency_embedding_size, operations=operations)
|
||||||
self.mlp = nn.Sequential(
|
self.mlp = nn.Sequential(
|
||||||
nn.Linear(frequency_embedding_size, hidden_size, bias=True),
|
operations.Linear(frequency_embedding_size, hidden_size, bias=True, dtype=dtype, device=device),
|
||||||
nn.SiLU(),
|
nn.SiLU(),
|
||||||
nn.Linear(hidden_size, hidden_size, bias=True),
|
operations.Linear(hidden_size, hidden_size, bias=True, dtype=dtype, device=device),
|
||||||
)
|
)
|
||||||
self.frequency_embedding_size = frequency_embedding_size
|
self.frequency_embedding_size = frequency_embedding_size
|
||||||
self.outdim = hidden_size
|
self.outdim = hidden_size
|
||||||
@@ -409,10 +394,10 @@ class LabelEmbedder(nn.Module):
|
|||||||
Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance.
|
Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, num_classes, hidden_size, dropout_prob):
|
def __init__(self, num_classes, hidden_size, dropout_prob, dtype=None, device=None, operations=None):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
use_cfg_embedding = dropout_prob > 0
|
use_cfg_embedding = dropout_prob > 0
|
||||||
self.embedding_table = nn.Embedding(num_classes + use_cfg_embedding, hidden_size)
|
self.embedding_table = operations.Embedding(num_classes + use_cfg_embedding, hidden_size, dtype=dtype, device=device),
|
||||||
self.num_classes = num_classes
|
self.num_classes = num_classes
|
||||||
self.dropout_prob = dropout_prob
|
self.dropout_prob = dropout_prob
|
||||||
|
|
||||||
@@ -434,15 +419,31 @@ class LabelEmbedder(nn.Module):
|
|||||||
embeddings = self.embedding_table(labels)
|
embeddings = self.embedding_table(labels)
|
||||||
return embeddings
|
return embeddings
|
||||||
|
|
||||||
|
class Mlp(nn.Module):
|
||||||
|
def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, dtype=None, device=None, operations=None) -> None:
|
||||||
|
super().__init__()
|
||||||
|
out_features = out_features or in_features
|
||||||
|
hidden_features = hidden_features or in_features
|
||||||
|
|
||||||
|
self.fc1 = operations.Linear(in_features, hidden_features, bias=True, dtype=dtype, device=device)
|
||||||
|
self.act = act_layer()
|
||||||
|
self.fc2 = operations.Linear(hidden_features, out_features, bias=True, dtype=dtype, device=device)
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
x = self.act(self.fc1(x))
|
||||||
|
return self.fc2(x)
|
||||||
|
|
||||||
class CaptionEmbedder(nn.Module):
|
class CaptionEmbedder(nn.Module):
|
||||||
"""
|
"""
|
||||||
Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance.
|
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):
|
def __init__(self, in_channels, hidden_size, uncond_prob, act_layer=nn.GELU(approximate='tanh'), token_num=120, dtype=None, device=None, operations=None):
|
||||||
super().__init__()
|
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.y_proj = Mlp(
|
||||||
|
in_features=in_channels, hidden_features=hidden_size, out_features=hidden_size, act_layer=act_layer,
|
||||||
|
dtype=dtype, device=device, operations=operations,
|
||||||
|
)
|
||||||
self.register_buffer("y_embedding", nn.Parameter(torch.randn(token_num, in_channels) / in_channels ** 0.5))
|
self.register_buffer("y_embedding", nn.Parameter(torch.randn(token_num, in_channels) / in_channels ** 0.5))
|
||||||
self.uncond_prob = uncond_prob
|
self.uncond_prob = uncond_prob
|
||||||
|
|
||||||
@@ -472,9 +473,12 @@ class CaptionEmbedderDoubleBr(nn.Module):
|
|||||||
Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance.
|
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):
|
def __init__(self, in_channels, hidden_size, uncond_prob, act_layer=nn.GELU(approximate='tanh'), token_num=120, dtype=None, device=None, operations=None):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.proj = Mlp(in_features=in_channels, hidden_features=hidden_size, out_features=hidden_size, act_layer=act_layer, drop=0)
|
self.proj = Mlp(
|
||||||
|
in_features=in_channels, hidden_features=hidden_size, out_features=hidden_size, act_layer=act_layer,
|
||||||
|
dtype=dtype, device=device, operations=operations,
|
||||||
|
)
|
||||||
self.embedding = nn.Parameter(torch.randn(1, in_channels) / 10 ** 0.5)
|
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.y_embedding = nn.Parameter(torch.randn(token_num, in_channels) / 10 ** 0.5)
|
||||||
self.uncond_prob = uncond_prob
|
self.uncond_prob = uncond_prob
|
||||||
|
|||||||
+76
-73
@@ -11,11 +11,9 @@
|
|||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
from timm.models.layers import DropPath
|
|
||||||
from timm.models.vision_transformer import Mlp
|
|
||||||
|
|
||||||
from .utils import auto_grad_checkpoint, to_2tuple
|
from .utils import auto_grad_checkpoint, to_2tuple
|
||||||
from .blocks import t2i_modulate, CaptionEmbedder, AttentionKVCompress, MultiHeadCrossAttention, T2IFinalLayer, TimestepEmbedder, SizeEmbedder
|
from .blocks import t2i_modulate, CaptionEmbedder, AttentionKVCompress, MultiHeadCrossAttention, T2IFinalLayer, TimestepEmbedder, SizeEmbedder, Mlp
|
||||||
from .pixart import PixArt, get_2d_sincos_pos_embed
|
from .pixart import PixArt, get_2d_sincos_pos_embed
|
||||||
|
|
||||||
|
|
||||||
@@ -31,12 +29,15 @@ class PatchEmbed(nn.Module):
|
|||||||
norm_layer=None,
|
norm_layer=None,
|
||||||
flatten=True,
|
flatten=True,
|
||||||
bias=True,
|
bias=True,
|
||||||
|
dtype=None,
|
||||||
|
device=None,
|
||||||
|
operations=None
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
patch_size = to_2tuple(patch_size)
|
patch_size = to_2tuple(patch_size)
|
||||||
self.patch_size = patch_size
|
self.patch_size = patch_size
|
||||||
self.flatten = flatten
|
self.flatten = flatten
|
||||||
self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size, bias=bias)
|
self.proj = operations.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size, bias=bias, dtype=dtype, device=device)
|
||||||
self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()
|
self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
@@ -52,29 +53,34 @@ class PixArtMSBlock(nn.Module):
|
|||||||
A PixArt block with adaptive layer norm zero (adaLN-Zero) conditioning.
|
A PixArt block with adaptive layer norm zero (adaLN-Zero) conditioning.
|
||||||
"""
|
"""
|
||||||
def __init__(self, hidden_size, num_heads, mlp_ratio=4.0, drop_path=0., input_size=None,
|
def __init__(self, hidden_size, num_heads, mlp_ratio=4.0, drop_path=0., input_size=None,
|
||||||
sampling=None, sr_ratio=1, qk_norm=False, **block_kwargs):
|
sampling=None, sr_ratio=1, qk_norm=False, dtype=None, device=None, operations=None, **block_kwargs):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.hidden_size = hidden_size
|
self.hidden_size = hidden_size
|
||||||
self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
self.norm1 = operations.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device)
|
||||||
self.attn = AttentionKVCompress(
|
self.attn = AttentionKVCompress(
|
||||||
hidden_size, num_heads=num_heads, qkv_bias=True, sampling=sampling, sr_ratio=sr_ratio,
|
hidden_size, num_heads=num_heads, qkv_bias=True, sampling=sampling, sr_ratio=sr_ratio,
|
||||||
qk_norm=qk_norm, **block_kwargs
|
qk_norm=qk_norm, dtype=dtype, device=device, operations=operations, **block_kwargs
|
||||||
)
|
)
|
||||||
self.cross_attn = MultiHeadCrossAttention(hidden_size, num_heads, **block_kwargs)
|
self.cross_attn = MultiHeadCrossAttention(
|
||||||
self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
hidden_size, num_heads, dtype=dtype, device=device, operations=operations, **block_kwargs
|
||||||
|
)
|
||||||
|
self.norm2 = operations.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device)
|
||||||
# to be compatible with lower version pytorch
|
# to be compatible with lower version pytorch
|
||||||
approx_gelu = lambda: nn.GELU(approximate="tanh")
|
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)
|
self.mlp = Mlp(
|
||||||
self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()
|
in_features=hidden_size, hidden_features=int(hidden_size * mlp_ratio), act_layer=approx_gelu,
|
||||||
|
dtype=dtype, device=device, operations=operations
|
||||||
|
)
|
||||||
self.scale_shift_table = nn.Parameter(torch.randn(6, hidden_size) / hidden_size ** 0.5)
|
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):
|
def forward(self, x, y, t, mask=None, HW=None, **kwargs):
|
||||||
B, N, C = x.shape
|
B, N, C = x.shape
|
||||||
|
dtype = x.dtype
|
||||||
|
|
||||||
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)
|
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (self.scale_shift_table[None].to(x.dtype) + 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 + (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.cross_attn(x, y, mask)
|
||||||
x = x + self.drop_path(gate_mlp * self.mlp(t2i_modulate(self.norm2(x), shift_mlp, scale_mlp)))
|
x = x + (gate_mlp * self.mlp(t2i_modulate(self.norm2(x), shift_mlp, scale_mlp)))
|
||||||
|
|
||||||
return x
|
return x
|
||||||
|
|
||||||
@@ -105,40 +111,52 @@ class PixArtMS(PixArt):
|
|||||||
micro_condition=True,
|
micro_condition=True,
|
||||||
qk_norm=False,
|
qk_norm=False,
|
||||||
kv_compress_config=None,
|
kv_compress_config=None,
|
||||||
|
dtype=None,
|
||||||
|
device=None,
|
||||||
|
operations=None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
super().__init__(
|
nn.Module.__init__(self)
|
||||||
input_size=input_size,
|
self.dtype = dtype
|
||||||
patch_size=patch_size,
|
self.pred_sigma = pred_sigma
|
||||||
in_channels=in_channels,
|
self.in_channels = in_channels
|
||||||
hidden_size=hidden_size,
|
self.out_channels = in_channels * 2 if pred_sigma else in_channels
|
||||||
depth=depth,
|
self.patch_size = patch_size
|
||||||
num_heads=num_heads,
|
self.num_heads = num_heads
|
||||||
mlp_ratio=mlp_ratio,
|
self.pe_interpolation = pe_interpolation
|
||||||
class_dropout_prob=class_dropout_prob,
|
self.pe_precision = pe_precision
|
||||||
learn_sigma=learn_sigma,
|
self.depth = depth
|
||||||
pred_sigma=pred_sigma,
|
|
||||||
drop_path=drop_path,
|
|
||||||
pe_interpolation=pe_interpolation,
|
|
||||||
config=config,
|
|
||||||
model_max_length=model_max_length,
|
|
||||||
qk_norm=qk_norm,
|
|
||||||
kv_compress_config=kv_compress_config,
|
|
||||||
**kwargs,
|
|
||||||
)
|
|
||||||
self.dtype = torch.get_default_dtype()
|
|
||||||
self.h = self.w = 0
|
self.h = self.w = 0
|
||||||
approx_gelu = lambda: nn.GELU(approximate="tanh")
|
approx_gelu = lambda: nn.GELU(approximate="tanh")
|
||||||
self.t_block = nn.Sequential(
|
self.t_block = nn.Sequential(
|
||||||
nn.SiLU(),
|
nn.SiLU(),
|
||||||
nn.Linear(hidden_size, 6 * hidden_size, bias=True)
|
operations.Linear(hidden_size, 6 * hidden_size, bias=True, dtype=dtype, device=device)
|
||||||
)
|
)
|
||||||
self.x_embedder = PatchEmbed(patch_size, in_channels, hidden_size, bias=True)
|
self.x_embedder = PatchEmbed(
|
||||||
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)
|
patch_size, in_channels, hidden_size, bias=True,
|
||||||
|
dtype=dtype, device=device, operations=operations
|
||||||
|
)
|
||||||
|
self.t_embedder = TimestepEmbedder(
|
||||||
|
hidden_size, dtype=dtype, device=device, operations=operations,
|
||||||
|
)
|
||||||
|
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,
|
||||||
|
dtype=dtype, device=device, operations=operations,
|
||||||
|
)
|
||||||
|
|
||||||
self.micro_conditioning = micro_condition
|
self.micro_conditioning = micro_condition
|
||||||
if self.micro_conditioning:
|
if self.micro_conditioning:
|
||||||
self.csize_embedder = SizeEmbedder(hidden_size//3) # c_size embed
|
|
||||||
self.ar_embedder = SizeEmbedder(hidden_size//3) # aspect ratio embed
|
self.csize_embedder = SizeEmbedder(hidden_size//3, dtype=dtype, device=device, operations=operations)
|
||||||
|
self.ar_embedder = SizeEmbedder(hidden_size//3, dtype=dtype, device=device, operations=operations)
|
||||||
|
|
||||||
|
# Will use fixed sin-cos embedding:
|
||||||
|
num_patches = (input_size // patch_size) * (input_size // patch_size)
|
||||||
|
self.base_size = input_size // self.patch_size
|
||||||
|
self.register_buffer("pos_embed", torch.zeros(1, num_patches, hidden_size))
|
||||||
|
|
||||||
drop_path = [x.item() for x in torch.linspace(0, drop_path, depth)] # stochastic depth decay rule
|
drop_path = [x.item() for x in torch.linspace(0, drop_path, depth)] # stochastic depth decay rule
|
||||||
if kv_compress_config is None:
|
if kv_compress_config is None:
|
||||||
kv_compress_config = {
|
kv_compress_config = {
|
||||||
@@ -153,12 +171,17 @@ class PixArtMS(PixArt):
|
|||||||
sampling=kv_compress_config['sampling'],
|
sampling=kv_compress_config['sampling'],
|
||||||
sr_ratio=int(kv_compress_config['scale_factor']) if i in kv_compress_config['kv_compress_layer'] else 1,
|
sr_ratio=int(kv_compress_config['scale_factor']) if i in kv_compress_config['kv_compress_layer'] else 1,
|
||||||
qk_norm=qk_norm,
|
qk_norm=qk_norm,
|
||||||
|
dtype=dtype,
|
||||||
|
device=device,
|
||||||
|
operations=operations,
|
||||||
)
|
)
|
||||||
for i in range(depth)
|
for i in range(depth)
|
||||||
])
|
])
|
||||||
self.final_layer = T2IFinalLayer(hidden_size, patch_size, self.out_channels)
|
self.final_layer = T2IFinalLayer(
|
||||||
|
hidden_size, patch_size, self.out_channels, dtype=dtype, device=device, operations=operations
|
||||||
|
)
|
||||||
|
|
||||||
def forward_raw(self, x, t, y, mask=None, data_info=None, **kwargs):
|
def forward_orig(self, x, timestep, y, mask=None, data_info=None, **kwargs):
|
||||||
"""
|
"""
|
||||||
Original forward pass of PixArt.
|
Original forward pass of PixArt.
|
||||||
x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images)
|
x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images)
|
||||||
@@ -166,9 +189,6 @@ class PixArtMS(PixArt):
|
|||||||
y: (N, 1, 120, C) tensor of class labels
|
y: (N, 1, 120, C) tensor of class labels
|
||||||
"""
|
"""
|
||||||
bs = x.shape[0]
|
bs = x.shape[0]
|
||||||
x = x.to(self.dtype)
|
|
||||||
timestep = t.to(self.dtype)
|
|
||||||
y = y.to(self.dtype)
|
|
||||||
|
|
||||||
pe_interpolation = self.pe_interpolation
|
pe_interpolation = self.pe_interpolation
|
||||||
if pe_interpolation is None or self.pe_precision is not None:
|
if pe_interpolation is None or self.pe_precision is not None:
|
||||||
@@ -181,10 +201,10 @@ class PixArtMS(PixArt):
|
|||||||
self.pos_embed.shape[-1], (self.h, self.w), pe_interpolation=pe_interpolation,
|
self.pos_embed.shape[-1], (self.h, self.w), pe_interpolation=pe_interpolation,
|
||||||
base_size=self.base_size
|
base_size=self.base_size
|
||||||
)
|
)
|
||||||
).unsqueeze(0).to(device=x.device, dtype=self.dtype)
|
).to(device=x.device, dtype=x.dtype).unsqueeze(0)
|
||||||
|
|
||||||
x = self.x_embedder(x) + pos_embed # (N, T, D), where T = H * W / patch_size ** 2
|
x = self.x_embedder(x) + pos_embed # (N, T, D), where T = H * W / patch_size ** 2
|
||||||
t = self.t_embedder(timestep) # (N, D)
|
t = self.t_embedder(timestep, x.dtype) # (N, D)
|
||||||
|
|
||||||
if self.micro_conditioning:
|
if self.micro_conditioning:
|
||||||
c_size, ar = data_info['img_hw'].to(self.dtype), data_info['aspect_ratio'].to(self.dtype)
|
c_size, ar = data_info['img_hw'].to(self.dtype), data_info['aspect_ratio'].to(self.dtype)
|
||||||
@@ -212,46 +232,29 @@ class PixArtMS(PixArt):
|
|||||||
|
|
||||||
return x
|
return x
|
||||||
|
|
||||||
def forward(self, x, timesteps, context, img_hw=None, aspect_ratio=None, **kwargs):
|
def forward(self, x, timesteps, context, width=None, height=None, img_hw=None, aspect_ratio=None, **kwargs):
|
||||||
"""
|
bs, c, h, w = x.shape
|
||||||
Forward pass that adapts comfy input to original forward function
|
dtype = self.dtype
|
||||||
x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images)
|
device = x.device
|
||||||
timesteps: (N,) tensor of diffusion timesteps
|
|
||||||
context: (N, 1, 120, C) conditioning
|
|
||||||
img_hw: height|width conditioning
|
|
||||||
aspect_ratio: aspect ratio conditioning
|
|
||||||
"""
|
|
||||||
## size/ar from cond with fallback based on the latent image shape.
|
## size/ar from cond with fallback based on the latent image shape.
|
||||||
bs = x.shape[0]
|
bs = x.shape[0]
|
||||||
data_info = {}
|
data_info = {}
|
||||||
if img_hw is None:
|
if img_hw is None:
|
||||||
data_info["img_hw"] = torch.tensor(
|
data_info["img_hw"] = torch.tensor([h*8, w*8], dtype=dtype, device=device).repeat(bs, 1)
|
||||||
[[x.shape[2]*8, x.shape[3]*8]],
|
|
||||||
dtype=self.dtype,
|
|
||||||
device=x.device
|
|
||||||
).repeat(bs, 1)
|
|
||||||
else:
|
else:
|
||||||
data_info["img_hw"] = img_hw.to(dtype=x.dtype, device=x.device)
|
data_info["img_hw"] = img_hw.to(dtype=x.dtype, device=x.device)
|
||||||
if aspect_ratio is None or True:
|
if aspect_ratio is None:
|
||||||
data_info["aspect_ratio"] = torch.tensor(
|
data_info["aspect_ratio"] = torch.tensor([h/w], dtype=dtype, device=device).repeat(bs, 1)
|
||||||
[[x.shape[2]/x.shape[3]]],
|
|
||||||
dtype=self.dtype,
|
|
||||||
device=x.device
|
|
||||||
).repeat(bs, 1)
|
|
||||||
else:
|
else:
|
||||||
data_info["aspect_ratio"] = aspect_ratio.to(dtype=x.dtype, device=x.device)
|
data_info["aspect_ratio"] = aspect_ratio.to(dtype=dtype, device=device)
|
||||||
|
|
||||||
## Still accepts the input w/o that dim but returns garbage
|
## Still accepts the input w/o that dim but returns garbage
|
||||||
if len(context.shape) == 3:
|
if len(context.shape) == 3:
|
||||||
context = context.unsqueeze(1)
|
context = context.unsqueeze(1)
|
||||||
|
|
||||||
## run original forward pass
|
## run original forward pass
|
||||||
out = self.forward_raw(
|
out = self.forward_orig(x, timesteps, context, data_info=data_info)
|
||||||
x = x.to(self.dtype),
|
|
||||||
t = timesteps.to(self.dtype),
|
|
||||||
y = context.to(self.dtype),
|
|
||||||
data_info=data_info,
|
|
||||||
)
|
|
||||||
|
|
||||||
## only return EPS
|
## only return EPS
|
||||||
out = out.to(torch.float)
|
out = out.to(torch.float)
|
||||||
|
|||||||
Reference in New Issue
Block a user