From d5a47fa3e5f8ff3b552d15e7545b8170f06ab984 Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Wed, 11 Dec 2024 00:52:48 +0100 Subject: [PATCH] Use native ops for PixArt --- PixArt/config.py | 3 +- PixArt/loader.py | 12 +--- PixArt/model/blocks.py | 150 ++++++++++++++++++++------------------- PixArt/model/pixartms.py | 149 +++++++++++++++++++------------------- 4 files changed, 157 insertions(+), 157 deletions(-) diff --git a/PixArt/config.py b/PixArt/config.py index 14975f7..e11e26a 100644 --- a/PixArt/config.py +++ b/PixArt/config.py @@ -93,8 +93,7 @@ def model_config_from_unet(sd): config["pe_interpolation"] = 1 model_config = PixArtConfig(config) model_config.unet_class = model_class - logging.info(f"Detected PixArt model as [{model_class}]") - logging.info(f"PixArt config:\n{config}") + logging.debug(f"PixArt config: {model_class}\n{config}") return model_config resolutions = { diff --git a/PixArt/loader.py b/PixArt/loader.py index b3627d6..d63b1d8 100644 --- a/PixArt/loader.py +++ b/PixArt/loader.py @@ -19,14 +19,8 @@ class PixArtModel(comfy.model_base.BaseModel): def extra_conds(self, **kwargs): out = super().extra_conds(**kwargs) - img_hw = kwargs.get("img_hw", None) - if img_hw is not None: - 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)) - + for name in ["width", "height", "aspect_ratio", "img_hw"]: # TODO: remove last one + out[name] = comfy.conds.CONDRegular(torch.tensor(name)) return out 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() 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) unet_weight_dtype = list(model_config.supported_inference_dtypes) diff --git a/PixArt/model/blocks.py b/PixArt/model/blocks.py index df2eae0..3f3c2d9 100644 --- a/PixArt/model/blocks.py +++ b/PixArt/model/blocks.py @@ -12,9 +12,10 @@ import math import torch import torch.nn as nn import torch.nn.functional as F -from timm.models.vision_transformer import Mlp, Attention as Attention_ from einops import rearrange +from .utils import to_2tuple + sdpa_32b = None Q_4GB_LIMIT = 32000000 """If q is greater than this, the operation will likely require >4GB VRAM, which will fail on Intel Arc Alchemist GPUs without a workaround.""" @@ -44,7 +45,7 @@ def t2i_modulate(x, shift, scale): return x * (1 + scale) + shift class MultiHeadCrossAttention(nn.Module): - def __init__(self, d_model, num_heads, attn_drop=0., 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__() 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.head_dim = d_model // num_heads - self.q_linear = nn.Linear(d_model, d_model) - self.kv_linear = nn.Linear(d_model, d_model*2) + self.q_linear = operations.Linear(d_model, d_model, dtype=dtype, device=device) + self.kv_linear = operations.Linear(d_model, d_model*2, dtype=dtype, device=device) 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) def forward(self, x, cond, mask=None): @@ -111,7 +112,7 @@ class MultiHeadCrossAttention(nn.Module): return x -class AttentionKVCompress(Attention_): +class AttentionKVCompress(nn.Module): """Multi-head Attention block with KV token compression and qk norm.""" def __init__( @@ -122,6 +123,9 @@ class AttentionKVCompress(Attention_): sampling='conv', sr_ratio=1, qk_norm=False, + dtype=None, + device=None, + operations=None, **block_kwargs, ): """ @@ -130,19 +134,26 @@ class AttentionKVCompress(Attention_): num_heads (int): Number of attention heads. qkv_bias (bool: If True, add a learnable bias to query, key, value. """ - super().__init__(dim, num_heads=num_heads, qkv_bias=qkv_bias, **block_kwargs) + 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.sr_ratio = sr_ratio if sr_ratio > 1 and sampling == 'conv': # Avg Conv Init. - self.sr = nn.Conv2d(dim, dim, groups=dim, kernel_size=sr_ratio, stride=sr_ratio) - self.sr.weight.data.fill_(1/sr_ratio**2) - self.sr.bias.data.zero_() - self.norm = nn.LayerNorm(dim) + 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.bias.data.zero_() + self.norm = operations.LayerNorm(dim, dtype=dtype, device=device) if qk_norm: - self.q_norm = nn.LayerNorm(dim) - self.k_norm = nn.LayerNorm(dim) + self.q_norm = operations.LayerNorm(dim, dtype=dtype, device=device) + self.k_norm = operations.LayerNorm(dim, dtype=dtype, device=device) else: self.q_norm = nn.Identity() self.k_norm = nn.Identity() @@ -204,14 +215,12 @@ class AttentionKVCompress(Attention_): if model_management.xformers_enabled(): x = xformers.ops.memory_efficient_attention( q, k, v, - p=self.attn_drop.p, + p=0, attn_bias=attn_bias ) else: q, k, v = map(lambda t: t.transpose(1, 2),(q, k, v),) - - p = getattr(self.attn_drop, "p", 0) # IPEX.optimize() will turn attn_drop into an Identity() - + p = 0 if sdpa_32b is not None and (q.element_size() * q.nelement()) > Q_4GB_LIMIT: sdpa = sdpa_32b else: @@ -224,30 +233,6 @@ class AttentionKVCompress(Attention_): ).transpose(1, 2).contiguous() x = x.view(B, N, C) 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 @@ -256,13 +241,13 @@ class FinalLayer(nn.Module): 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__() - self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True) + self.norm_final = operations.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device) + self.linear = operations.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True, dtype=dtype, device=device) self.adaLN_modulation = nn.Sequential( 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): @@ -271,23 +256,23 @@ class FinalLayer(nn.Module): x = self.linear(x) return x - class T2IFinalLayer(nn.Module): """ 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__() - self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True) + self.norm_final = operations.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device) + 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.out_channels = out_channels def forward(self, x, t): + dtype = x.dtype shift, scale = (self.scale_shift_table[None] + t[:, None]).chunk(2, dim=1) x = t2i_modulate(self.norm_final(x), shift, scale) - x = self.linear(x) + x = self.linear(x.to(dtype)) return x @@ -296,13 +281,13 @@ class MaskFinalLayer(nn.Module): 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__() - self.norm_final = nn.LayerNorm(final_hidden_size, elementwise_affine=False, eps=1e-6) - self.linear = nn.Linear(final_hidden_size, patch_size * patch_size * out_channels, bias=True) + self.norm_final = operations.LayerNorm(final_hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device) + self.linear = operations.Linear(final_hidden_size, patch_size * patch_size * out_channels, bias=True, dtype=dtype, device=device) self.adaLN_modulation = nn.Sequential( 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): shift, scale = self.adaLN_modulation(t).chunk(2, dim=1) @@ -316,13 +301,13 @@ class DecoderLayer(nn.Module): 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__() - self.norm_decoder = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - self.linear = nn.Linear(hidden_size, decoder_hidden_size, bias=True) + self.norm_decoder = operations.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device) + self.linear = operations.Linear(hidden_size, decoder_hidden_size, bias=True, dtype=dtype, device=device) self.adaLN_modulation = nn.Sequential( 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): shift, scale = self.adaLN_modulation(t).chunk(2, dim=1) @@ -339,12 +324,12 @@ class TimestepEmbedder(nn.Module): 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__() 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.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 @@ -368,9 +353,9 @@ class TimestepEmbedder(nn.Module): embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1) return embedding - def forward(self, t): + def forward(self, t, dtype): 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 @@ -379,12 +364,12 @@ class SizeEmbedder(TimestepEmbedder): Embeds scalar timesteps into vector representations. """ - def __init__(self, hidden_size, frequency_embedding_size=256): - super().__init__(hidden_size=hidden_size, frequency_embedding_size=frequency_embedding_size) + 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, operations=operations) 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.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.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. """ - 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__() 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.dropout_prob = dropout_prob @@ -434,15 +419,31 @@ class LabelEmbedder(nn.Module): embeddings = self.embedding_table(labels) 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): """ 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__() - 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.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. """ - 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__() - 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.y_embedding = nn.Parameter(torch.randn(token_num, in_channels) / 10 ** 0.5) self.uncond_prob = uncond_prob diff --git a/PixArt/model/pixartms.py b/PixArt/model/pixartms.py index 0bb5b36..a9e7035 100644 --- a/PixArt/model/pixartms.py +++ b/PixArt/model/pixartms.py @@ -11,11 +11,9 @@ import torch import torch.nn as nn 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 .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 @@ -31,12 +29,15 @@ class PatchEmbed(nn.Module): norm_layer=None, flatten=True, bias=True, + dtype=None, + device=None, + operations=None ): super().__init__() patch_size = to_2tuple(patch_size) self.patch_size = patch_size 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() def forward(self, x): @@ -52,29 +53,34 @@ class PixArtMSBlock(nn.Module): 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, - 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__() 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( 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.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + self.cross_attn = MultiHeadCrossAttention( + 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 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.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity() + self.mlp = Mlp( + 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) def forward(self, x, y, t, mask=None, HW=None, **kwargs): 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) - x = x + self.drop_path(gate_msa * self.attn(t2i_modulate(self.norm1(x), shift_msa, scale_msa), HW=HW)) + 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 + (gate_msa * self.attn(t2i_modulate(self.norm1(x), shift_msa, scale_msa), HW=HW)) x = x + self.cross_attn(x, y, mask) - x = x + self.drop_path(gate_mlp * self.mlp(t2i_modulate(self.norm2(x), shift_mlp, scale_mlp))) + x = x + (gate_mlp * self.mlp(t2i_modulate(self.norm2(x), shift_mlp, scale_mlp))) return x @@ -105,40 +111,52 @@ class PixArtMS(PixArt): micro_condition=True, qk_norm=False, kv_compress_config=None, + dtype=None, + device=None, + operations=None, **kwargs, ): - super().__init__( - input_size=input_size, - patch_size=patch_size, - in_channels=in_channels, - hidden_size=hidden_size, - depth=depth, - num_heads=num_heads, - mlp_ratio=mlp_ratio, - class_dropout_prob=class_dropout_prob, - learn_sigma=learn_sigma, - pred_sigma=pred_sigma, - drop_path=drop_path, - 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() + nn.Module.__init__(self) + self.dtype = dtype + self.pred_sigma = pred_sigma + self.in_channels = in_channels + self.out_channels = in_channels * 2 if pred_sigma else in_channels + self.patch_size = patch_size + self.num_heads = num_heads + self.pe_interpolation = pe_interpolation + self.pe_precision = pe_precision + self.depth = depth + self.h = self.w = 0 approx_gelu = lambda: nn.GELU(approximate="tanh") self.t_block = nn.Sequential( nn.SiLU(), - nn.Linear(hidden_size, 6 * hidden_size, bias=True) + 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.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) + self.x_embedder = PatchEmbed( + 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 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 if kv_compress_config is None: kv_compress_config = { @@ -153,12 +171,17 @@ class PixArtMS(PixArt): sampling=kv_compress_config['sampling'], sr_ratio=int(kv_compress_config['scale_factor']) if i in kv_compress_config['kv_compress_layer'] else 1, qk_norm=qk_norm, + dtype=dtype, + device=device, + operations=operations, ) 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. 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 """ bs = x.shape[0] - x = x.to(self.dtype) - timestep = t.to(self.dtype) - y = y.to(self.dtype) pe_interpolation = self.pe_interpolation 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, 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 - t = self.t_embedder(timestep) # (N, D) + t = self.t_embedder(timestep, x.dtype) # (N, D) if self.micro_conditioning: 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 - def forward(self, x, timesteps, context, img_hw=None, aspect_ratio=None, **kwargs): - """ - Forward pass that adapts comfy input to original forward function - x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images) - timesteps: (N,) tensor of diffusion timesteps - context: (N, 1, 120, C) conditioning - img_hw: height|width conditioning - aspect_ratio: aspect ratio conditioning - """ + def forward(self, x, timesteps, context, width=None, height=None, img_hw=None, aspect_ratio=None, **kwargs): + bs, c, h, w = x.shape + dtype = self.dtype + device = x.device + ## size/ar from cond with fallback based on the latent image shape. bs = x.shape[0] data_info = {} if img_hw is None: - data_info["img_hw"] = torch.tensor( - [[x.shape[2]*8, x.shape[3]*8]], - dtype=self.dtype, - device=x.device - ).repeat(bs, 1) + data_info["img_hw"] = torch.tensor([h*8, w*8], dtype=dtype, device=device).repeat(bs, 1) else: data_info["img_hw"] = img_hw.to(dtype=x.dtype, device=x.device) - if aspect_ratio is None or True: - data_info["aspect_ratio"] = torch.tensor( - [[x.shape[2]/x.shape[3]]], - dtype=self.dtype, - device=x.device - ).repeat(bs, 1) + if aspect_ratio is None: + data_info["aspect_ratio"] = torch.tensor([h/w], dtype=dtype, device=device).repeat(bs, 1) 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 if len(context.shape) == 3: context = context.unsqueeze(1) ## run original forward pass - out = self.forward_raw( - x = x.to(self.dtype), - t = timesteps.to(self.dtype), - y = context.to(self.dtype), - data_info=data_info, - ) + out = self.forward_orig(x, timesteps, context, data_info=data_info) ## only return EPS out = out.to(torch.float)