From 804a7076cab1b614bef54cd5623fdb8b1ca54cf0 Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Sat, 14 Dec 2024 20:31:36 +0100 Subject: [PATCH] pos_emb and interpolation logic --- comfy/ldm/pixart/pixart.py | 16 +++++++++++-- comfy/ldm/pixart/pixartms.py | 46 ++++++++++++++++++++---------------- comfy/model_detection.py | 3 --- 3 files changed, 39 insertions(+), 26 deletions(-) diff --git a/comfy/ldm/pixart/pixart.py b/comfy/ldm/pixart/pixart.py index 1db1e6616..d79a59b9b 100644 --- a/comfy/ldm/pixart/pixart.py +++ b/comfy/ldm/pixart/pixart.py @@ -14,7 +14,7 @@ from .blocks import ( TimestepEmbedder, Mlp ) -from comfy.ldm.modules.diffusionmodules.mmdit import PatchEmbed +from comfy.ldm.modules.diffusionmodules.mmdit import PatchEmbed, get_1d_sincos_pos_embed_from_grid_torch class PixArtBlock(nn.Module): @@ -134,7 +134,6 @@ class PixArt(nn.Module): timestep = t.to(self.dtype) y = y.to(self.dtype) pos_embed = self.pos_embed.to(self.dtype) - self.h, self.w = x.shape[-2]//self.patch_size, x.shape[-1]//self.patch_size x = self.x_embedder(x) + pos_embed # (N, T, D), where T = H * W / patch_size ** 2 t = self.t_embedder(timestep.to(x.dtype)) # (N, D) t0 = self.t_block(t) @@ -193,7 +192,20 @@ class PixArt(nn.Module): imgs = x.reshape(shape=(x.shape[0], c, h * p, h * p)) return imgs +def get_2d_sincos_pos_embed_torch(embed_dim, w, h, pe_interpolation=1.0, base_size=16, device=None, dtype=torch.float32): + grid_h, grid_w = torch.meshgrid( + torch.arange(h, device=device, dtype=dtype) / (h/base_size) / pe_interpolation, + torch.arange(w, device=device, dtype=dtype) / (w/base_size) / pe_interpolation, + # torch.linspace(-val_h + val_center, val_h + val_center, h, device=device, dtype=dtype), + # torch.linspace(-val_w + val_center, val_w + val_center, w, device=device, dtype=dtype), + indexing='ij' + ) + emb_h = get_1d_sincos_pos_embed_from_grid_torch(embed_dim // 2, grid_h, device=device, dtype=dtype) + emb_w = get_1d_sincos_pos_embed_from_grid_torch(embed_dim // 2, grid_w, device=device, dtype=dtype) + emb = torch.cat([emb_w, emb_h], dim=1) # (H*W, D) + return emb +# Unused def get_2d_sincos_pos_embed(embed_dim, grid_size, cls_token=False, extra_tokens=0, pe_interpolation=1.0, base_size=16): """ grid_size: int of the grid height and width diff --git a/comfy/ldm/pixart/pixartms.py b/comfy/ldm/pixart/pixartms.py index 2b2fb1b1e..d5edb2ab7 100644 --- a/comfy/ldm/pixart/pixartms.py +++ b/comfy/ldm/pixart/pixartms.py @@ -5,7 +5,7 @@ import torch import torch.nn as nn 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, get_2d_sincos_pos_embed_torch class PatchEmbed(nn.Module): """ @@ -117,7 +117,6 @@ class PixArtMS(PixArt): self.hidden_size = hidden_size self.depth = depth - self.h = self.w = 0 approx_gelu = lambda: nn.GELU(approximate="tanh") self.t_block = nn.Sequential( nn.SiLU(), @@ -142,10 +141,10 @@ class PixArtMS(PixArt): 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)) + # For 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: @@ -157,7 +156,6 @@ class PixArtMS(PixArt): self.blocks = nn.ModuleList([ PixArtMSBlock( hidden_size, num_heads, mlp_ratio=mlp_ratio, drop_path=drop_path[i], - input_size=(input_size // patch_size, input_size // patch_size), 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, @@ -180,18 +178,22 @@ class PixArtMS(PixArt): ar: (N, 1): aspect ratio cs: (N ,2) size conditioning for height/width """ + B, C, H, W = x.shape + c_res = (H + W) // 2 pe_interpolation = self.pe_interpolation if pe_interpolation is None or self.pe_precision is not None: # calculate pe_interpolation on-the-fly - pe_interpolation = round((x.shape[-1]+x.shape[-2])/2.0 / (512/8.0), self.pe_precision or 0) + pe_interpolation = round(c_res / (512/8.0), self.pe_precision or 0) - self.h, self.w = x.shape[-2]//self.patch_size, x.shape[-1]//self.patch_size - pos_embed = torch.from_numpy( - get_2d_sincos_pos_embed( - self.hidden_size, (self.h, self.w), pe_interpolation=pe_interpolation, - base_size=self.base_size - ) - ).to(device=x.device, dtype=x.dtype).unsqueeze(0) + pos_embed = get_2d_sincos_pos_embed_torch( + self.hidden_size, + h=(H // self.patch_size), + w=(W // self.patch_size), + pe_interpolation=pe_interpolation, + base_size=((round(c_res / 64) * 64) // self.patch_size), + 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, x.dtype) # (N, D) @@ -215,10 +217,10 @@ class PixArtMS(PixArt): y_lens = [y.shape[2]] * y.shape[0] y = y.squeeze(1).view(1, -1, x.shape[-1]) for block in self.blocks: - x = block(x, y, t0, y_lens, (self.h, self.w), **kwargs) # (N, T, D) + x = block(x, y, t0, y_lens, (H, W), **kwargs) # (N, T, D) x = self.final_layer(x, t) # (N, T, patch_size ** 2 * out_channels) - x = self.unpatchify(x) # (N, out_channels, H, W) + x = self.unpatchify(x, H, W) # (N, out_channels, H, W) return x @@ -245,16 +247,18 @@ class PixArtMS(PixArt): return out[:, :self.in_channels] return out - def unpatchify(self, x): + def unpatchify(self, x, h, w): """ x: (N, T, patch_size**2 * C) imgs: (N, H, W, C) """ c = self.out_channels p = self.x_embedder.patch_size[0] - assert self.h * self.w == x.shape[1] + h = h // self.patch_size + w = w // self.patch_size + assert h * w == x.shape[1] - x = x.reshape(shape=(x.shape[0], self.h, self.w, p, p, c)) + x = x.reshape(shape=(x.shape[0], h, w, p, p, c)) x = torch.einsum('nhwpqc->nchpwq', x) - imgs = x.reshape(shape=(x.shape[0], c, self.h * p, self.w * p)) + imgs = x.reshape(shape=(x.shape[0], c, h * p, w * p)) return imgs diff --git a/comfy/model_detection.py b/comfy/model_detection.py index 11d3bd5ba..51d2035e4 100644 --- a/comfy/model_detection.py +++ b/comfy/model_detection.py @@ -209,9 +209,6 @@ def detect_unet_config(state_dict, key_prefix): if pe_key in state_dict_keys: dit_config["input_size"] = int(math.sqrt(state_dict[pe_key].shape[1])) * patch_size dit_config["pe_interpolation"] = dit_config["input_size"] // (512//8) # guess - else: - dit_config["input_size"] = 128 # 1024 - dit_config["pe_interpolation"] = 2 ar_key = "{}ar_embedder.mlp.0.weight".format(key_prefix) if ar_key in state_dict_keys: