diff --git a/extensions-builtin/PiD/nn/mmdit.py b/extensions-builtin/PiD/nn/mmdit.py index b9960db7..7cf64a35 100644 --- a/extensions-builtin/PiD/nn/mmdit.py +++ b/extensions-builtin/PiD/nn/mmdit.py @@ -1,56 +1,55 @@ -from functools import partial -from typing import Optional +# https://github.com/Comfy-Org/ComfyUI/blob/v0.24.1/comfy/ldm/modules/diffusionmodules/mmdit.py import torch import torch.nn as nn import torch.nn.functional as F +from backend.nn.lumina import timestep_embedding + class FeedForwardSwiGLU(nn.Module): - def __init__(self, dim: int, hidden_dim: int, multiple_of: int = 256, ffn_dim_multiplier: Optional[float] = None, dtype=None, device=None, operations=None): + def __init__(self, dim: int, hidden_dim: int, multiple_of: int = 256): super().__init__() + hidden_dim = int(2 * hidden_dim / 3) - # custom dim factor multiplier - if ffn_dim_multiplier is not None: - hidden_dim = int(ffn_dim_multiplier * hidden_dim) hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of) - self.w1 = operations.Linear(dim, hidden_dim, bias=False, dtype=dtype, device=device) - self.w2 = operations.Linear(hidden_dim, dim, bias=False, dtype=dtype, device=device) - self.w3 = operations.Linear(dim, hidden_dim, bias=False, dtype=dtype, device=device) + self.w1 = nn.Linear(dim, hidden_dim, bias=False) + self.w2 = nn.Linear(hidden_dim, dim, bias=False) + self.w3 = nn.Linear(dim, hidden_dim, bias=False) def forward(self, x): return self.w2(F.silu(self.w1(x)) * self.w3(x)) -class Mlp(nn.Module): - - def __init__( - self, - in_features, - hidden_features=None, - out_features=None, - act_layer=nn.GELU, - norm_layer=None, - bias=True, - drop=0.0, - use_conv=False, - dtype=None, - device=None, - operations=None, - ): +class TimestepEmbedder(nn.Module): + def __init__(self, hidden_size, frequency_embedding_size=256, max_period=10000): super().__init__() - out_features = out_features or in_features - hidden_features = hidden_features or in_features - drop_probs = drop - linear_layer = partial(operations.Conv2d, kernel_size=1) if use_conv else operations.Linear - self.fc1 = linear_layer(in_features, hidden_features, bias=bias, dtype=dtype, device=device) + self.mlp = nn.Sequential( + nn.Linear(frequency_embedding_size, hidden_size, bias=True), + nn.SiLU(), + nn.Linear(hidden_size, hidden_size, bias=True), + ) + self.frequency_embedding_size = frequency_embedding_size + self.max_period = max_period + + def forward(self, t, dtype, **kwargs): + t_freq = timestep_embedding(t, self.frequency_embedding_size, max_period=self.max_period).to(dtype) + t_emb = self.mlp(t_freq) + return t_emb + + +class Mlp(nn.Module): + def __init__(self, in_features, hidden_features, act_layer=nn.GELU, bias=True, drop=0.0): + super().__init__() + + self.fc1 = nn.Linear(in_features, hidden_features, bias=bias) self.act = act_layer() - self.drop1 = nn.Dropout(drop_probs) - self.norm = norm_layer(hidden_features) if norm_layer is not None else nn.Identity() - self.fc2 = linear_layer(hidden_features, out_features, bias=bias, dtype=dtype, device=device) - self.drop2 = nn.Dropout(drop_probs) + self.drop1 = nn.Dropout(drop) + self.norm = nn.Identity() + self.fc2 = nn.Linear(hidden_features, in_features, bias=bias) + self.drop2 = nn.Dropout(drop) def forward(self, x): x = self.fc1(x) @@ -65,10 +64,10 @@ class Mlp(nn.Module): def get_1d_sincos_pos_embed_from_grid_torch(embed_dim, pos, device=None, dtype=torch.float32): omega = torch.arange(embed_dim // 2, device=device, dtype=dtype) omega /= embed_dim / 2.0 - omega = 1.0 / 10000**omega # (D/2,) - pos = pos.reshape(-1) # (M,) - out = torch.einsum("m,d->md", pos, omega) # (M, D/2), outer product - emb_sin = torch.sin(out) # (M, D/2) - emb_cos = torch.cos(out) # (M, D/2) - emb = torch.cat([emb_sin, emb_cos], dim=1) # (M, D) + omega = 1.0 / 10000**omega + pos = pos.reshape(-1) + out = torch.einsum("m,d->md", pos, omega) + emb_sin = torch.sin(out) + emb_cos = torch.cos(out) + emb = torch.cat([emb_sin, emb_cos], dim=1) return emb diff --git a/extensions-builtin/PiD/nn/model.py b/extensions-builtin/PiD/nn/model.py index 5b1562b7..3855f278 100644 --- a/extensions-builtin/PiD/nn/model.py +++ b/extensions-builtin/PiD/nn/model.py @@ -1,11 +1,12 @@ +# https://github.com/Comfy-Org/ComfyUI/blob/v0.24.1/comfy/ldm/pixeldit/model.py + import torch import torch.nn as nn import torch.nn.functional as F -from backend.nn.lumina import TimestepEmbedder from backend.utils import pad_to_patch_size -from .mmdit import FeedForwardSwiGLU +from .mmdit import FeedForwardSwiGLU, TimestepEmbedder from .modules import ( FinalLayer, PatchTokenEmbedder, @@ -20,23 +21,23 @@ from .modules import ( class MMDiTJointAttention(nn.Module): - - def __init__(self, dim, num_heads=8, qkv_bias=False, dtype=None, device=None, operations=None): + def __init__(self, dim, num_heads=8, qkv_bias=False): super().__init__() + assert dim % num_heads == 0 self.num_heads = num_heads self.head_dim = dim // num_heads - self.qkv_x = operations.Linear(dim, dim * 3, bias=qkv_bias, dtype=dtype, device=device) - self.qkv_y = operations.Linear(dim, dim * 3, bias=qkv_bias, dtype=dtype, device=device) + self.qkv_x = nn.Linear(dim, dim * 3, bias=qkv_bias) + self.qkv_y = nn.Linear(dim, dim * 3, bias=qkv_bias) - self.q_norm_x = operations.RMSNorm(self.head_dim, eps=1e-6, dtype=dtype, device=device) - self.k_norm_x = operations.RMSNorm(self.head_dim, eps=1e-6, dtype=dtype, device=device) - self.q_norm_y = operations.RMSNorm(self.head_dim, eps=1e-6, dtype=dtype, device=device) - self.k_norm_y = operations.RMSNorm(self.head_dim, eps=1e-6, dtype=dtype, device=device) + self.q_norm_x = nn.RMSNorm(self.head_dim, eps=1e-6) + self.k_norm_x = nn.RMSNorm(self.head_dim, eps=1e-6) + self.q_norm_y = nn.RMSNorm(self.head_dim, eps=1e-6) + self.k_norm_y = nn.RMSNorm(self.head_dim, eps=1e-6) - self.proj_x = operations.Linear(dim, dim, dtype=dtype, device=device) - self.proj_y = operations.Linear(dim, dim, dtype=dtype, device=device) + self.proj_x = nn.Linear(dim, dim) + self.proj_y = nn.Linear(dim, dim) def forward(self, x, y, pos_img, pos_txt=None, attn_mask=None, transformer_options={}): B, Nx, _ = x.shape @@ -80,18 +81,19 @@ class MMDiTJointAttention(nn.Module): class MMDiTBlockT2I(nn.Module): - def __init__(self, hidden_size, groups, mlp_ratio=4.0, dtype=None, device=None, operations=None): + def __init__(self, hidden_size, groups, mlp_ratio=4.0): super().__init__() - self.norm_x1 = operations.RMSNorm(hidden_size, eps=1e-6, dtype=dtype, device=device) - self.norm_y1 = operations.RMSNorm(hidden_size, eps=1e-6, dtype=dtype, device=device) - self.attn = MMDiTJointAttention(hidden_size, num_heads=groups, qkv_bias=False, dtype=dtype, device=device, operations=operations) - self.norm_x2 = operations.RMSNorm(hidden_size, eps=1e-6, dtype=dtype, device=device) - self.norm_y2 = operations.RMSNorm(hidden_size, eps=1e-6, dtype=dtype, device=device) + + self.norm_x1 = nn.RMSNorm(hidden_size, eps=1e-6) + self.norm_y1 = nn.RMSNorm(hidden_size, eps=1e-6) + self.attn = MMDiTJointAttention(hidden_size, num_heads=groups, qkv_bias=False) + self.norm_x2 = nn.RMSNorm(hidden_size, eps=1e-6) + self.norm_y2 = nn.RMSNorm(hidden_size, eps=1e-6) mlp_hidden_dim = int(hidden_size * mlp_ratio) - self.mlp_x = FeedForwardSwiGLU(hidden_size, mlp_hidden_dim, multiple_of=1, dtype=dtype, device=device, operations=operations) - self.mlp_y = FeedForwardSwiGLU(hidden_size, mlp_hidden_dim, multiple_of=1, dtype=dtype, device=device, operations=operations) - self.adaLN_modulation_img = nn.Sequential(operations.Linear(hidden_size, 6 * hidden_size, bias=True, dtype=dtype, device=device)) - self.adaLN_modulation_txt = nn.Sequential(operations.Linear(hidden_size, 6 * hidden_size, bias=True, dtype=dtype, device=device)) + self.mlp_x = FeedForwardSwiGLU(hidden_size, mlp_hidden_dim, multiple_of=1) + self.mlp_y = FeedForwardSwiGLU(hidden_size, mlp_hidden_dim, multiple_of=1) + self.adaLN_modulation_img = nn.Sequential(nn.Linear(hidden_size, 6 * hidden_size, bias=True)) + self.adaLN_modulation_txt = nn.Sequential(nn.Linear(hidden_size, 6 * hidden_size, bias=True)) def forward(self, x, y, c, pos_img, pos_txt=None, attn_mask=None, transformer_options={}): shift_msa_x, scale_msa_x, gate_msa_x, shift_mlp_x, scale_mlp_x, gate_mlp_x = self.adaLN_modulation_img(c).chunk(6, dim=-1) @@ -100,16 +102,16 @@ class MMDiTBlockT2I(nn.Module): x_norm = apply_adaln_(self.norm_x1(x), shift_msa_x, scale_msa_x) y_norm = apply_adaln_(self.norm_y1(y), shift_msa_y, scale_msa_y) attn_x, attn_y = self.attn(x_norm, y_norm, pos_img, pos_txt, attn_mask, transformer_options=transformer_options) + x = torch.addcmul(x, gate_msa_x, attn_x) y = torch.addcmul(y, gate_msa_y, attn_y) - x = torch.addcmul(x, gate_mlp_x, self.mlp_x(apply_adaln_(self.norm_x2(x), shift_mlp_x, scale_mlp_x))) y = torch.addcmul(y, gate_mlp_y, self.mlp_y(apply_adaln_(self.norm_y2(y), shift_mlp_y, scale_mlp_y))) + return x, y class PixDiT_T2I(nn.Module): - def __init__( self, in_channels=3, @@ -125,14 +127,10 @@ class PixDiT_T2I(nn.Module): txt_max_length=300, use_text_rope=True, text_rope_theta=10000.0, - image_model=None, - dtype=None, - device=None, - operations=None, pixel_mlp_chunks=2, ): super().__init__() - self.dtype = dtype + self.in_channels = in_channels self.out_channels = in_channels self.hidden_size = hidden_size @@ -148,13 +146,13 @@ class PixDiT_T2I(nn.Module): self.use_text_rope = use_text_rope self.text_rope_theta = text_rope_theta - self.pixel_embedder = PixelTokenEmbedder(self.in_channels, self.pixel_hidden_size, dtype=dtype, device=device, operations=operations) - self.s_embedder = PatchTokenEmbedder(self.in_channels * self.patch_size**2, self.hidden_size, bias=True, dtype=dtype, device=device, operations=operations) - self.t_embedder = TimestepEmbedder(self.hidden_size, dtype=dtype, device=device, operations=operations, max_period=10) - self.y_embedder = PatchTokenEmbedder(self.txt_embed_dim, self.hidden_size, bias=True, use_norm=True, dtype=dtype, device=device, operations=operations) - self.y_pos_embedding = nn.Parameter(torch.empty(1, self.txt_max_length, self.hidden_size, dtype=dtype, device=device)) + self.pixel_embedder = PixelTokenEmbedder(self.in_channels, self.pixel_hidden_size) + self.s_embedder = PatchTokenEmbedder(self.in_channels * self.patch_size**2, self.hidden_size, bias=True) + self.t_embedder = TimestepEmbedder(self.hidden_size, max_period=10) + self.y_embedder = PatchTokenEmbedder(self.txt_embed_dim, self.hidden_size, bias=True, use_norm=True) + self.y_pos_embedding = nn.Parameter(torch.empty(1, self.txt_max_length, self.hidden_size)) - self.patch_blocks = nn.ModuleList([MMDiTBlockT2I(self.hidden_size, self.num_groups, dtype=dtype, device=device, operations=operations) for _ in range(self.patch_depth)]) + self.patch_blocks = nn.ModuleList([MMDiTBlockT2I(self.hidden_size, self.num_groups) for _ in range(self.patch_depth)]) self.pixel_blocks = nn.ModuleList( [ PiTBlock( @@ -164,16 +162,13 @@ class PixDiT_T2I(nn.Module): num_heads=self.num_groups, attn_hidden_size=self.pixel_attn_hidden_size, attn_num_heads=self.pixel_num_groups, - dtype=dtype, - device=device, - operations=operations, mlp_chunks=pixel_mlp_chunks, ) for _ in range(self.pixel_depth) ] ) - self.final_layer = FinalLayer(self.pixel_hidden_size, self.out_channels, dtype=dtype, device=device, operations=operations) + self.final_layer = FinalLayer(self.pixel_hidden_size, self.out_channels) def _fetch_patch_pos(self, height, width, device, dtype, **rope_opts): return precompute_freqs_cis_2d(self.hidden_size // self.num_groups, height, width, device=device, dtype=dtype, **rope_opts) @@ -202,7 +197,7 @@ class PixDiT_T2I(nn.Module): Ltxt = min(context.shape[1], self.txt_max_length) y = context[:, :Ltxt, :] y_emb = self.y_embedder(y).view(B, Ltxt, self.hidden_size) - y_emb = y_emb + self.y_pos_embedding[:, :Ltxt, :].to(y_emb) # y_pos_embedding is a raw nn.Parameter + y_emb = y_emb + self.y_pos_embedding[:, :Ltxt, :].to(y_emb) condition = F.silu(t_emb) pos_txt = self._fetch_text_pos(Ltxt, x.device, x.dtype) if self.use_text_rope else None @@ -223,4 +218,5 @@ class PixDiT_T2I(nn.Module): P2 = self.patch_size * self.patch_size x_pixels = x_pixels.view(B, L, P2, C_out).permute(0, 3, 2, 1).reshape(B, C_out * P2, L) out = F.fold(x_pixels, (H, W), kernel_size=self.patch_size, stride=self.patch_size) + return out[:, :, :H_orig, :W_orig] diff --git a/extensions-builtin/PiD/nn/modules.py b/extensions-builtin/PiD/nn/modules.py index 952018f9..d428059e 100644 --- a/extensions-builtin/PiD/nn/modules.py +++ b/extensions-builtin/PiD/nn/modules.py @@ -1,3 +1,5 @@ +# https://github.com/Comfy-Org/ComfyUI/blob/v0.24.1/comfy/ldm/pixeldit/modules.py + import torch import torch.nn as nn @@ -39,16 +41,17 @@ def get_2d_sincos_pos_embed(embed_dim, height, width, device=None, dtype=torch.f class RotaryAttention(nn.Module): - - def __init__(self, dim, num_heads=8, qkv_bias=False, dtype=None, device=None, operations=None): + def __init__(self, dim, num_heads=8, qkv_bias=False): super().__init__() + assert dim % num_heads == 0 self.num_heads = num_heads self.head_dim = dim // num_heads - self.qkv = operations.Linear(dim, dim * 3, bias=qkv_bias, dtype=dtype, device=device) - self.q_norm = operations.RMSNorm(self.head_dim, eps=1e-6, dtype=dtype, device=device) - self.k_norm = operations.RMSNorm(self.head_dim, eps=1e-6, dtype=dtype, device=device) - self.proj = operations.Linear(dim, dim, dtype=dtype, device=device) + + self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias) + self.q_norm = nn.RMSNorm(self.head_dim, eps=1e-6) + self.k_norm = nn.RMSNorm(self.head_dim, eps=1e-6) + self.proj = nn.Linear(dim, dim) def forward(self, x, pos, mask=None, transformer_options={}): B, N, C = x.shape @@ -58,37 +61,39 @@ class RotaryAttention(nn.Module): q, k, v = qkv.unbind(0) q, k = apply_rope(self.q_norm(q), self.k_norm(k), pos[None, None]) x = attention_function(q, k, v, H, mask=mask, skip_reshape=True, transformer_options=transformer_options) + return self.proj(x) class FinalLayer(nn.Module): - def __init__(self, hidden_size, out_channels, dtype=None, device=None, operations=None): + def __init__(self, hidden_size, out_channels): super().__init__() - self.norm = operations.RMSNorm(hidden_size, eps=1e-6, dtype=dtype, device=device) - self.linear = operations.Linear(hidden_size, out_channels, bias=True, dtype=dtype, device=device) + + self.norm = nn.RMSNorm(hidden_size, eps=1e-6) + self.linear = nn.Linear(hidden_size, out_channels, bias=True) def forward(self, x): return self.linear(self.norm(x)) class PatchTokenEmbedder(nn.Module): - - def __init__(self, in_chans, embed_dim, use_norm=False, bias=True, dtype=None, device=None, operations=None): + def __init__(self, in_chans, embed_dim, use_norm=False, bias=True): super().__init__() - self.proj = operations.Linear(in_chans, embed_dim, bias=bias, dtype=dtype, device=device) - self.norm = operations.RMSNorm(embed_dim, eps=1e-6, dtype=dtype, device=device) if use_norm else nn.Identity() + + self.proj = nn.Linear(in_chans, embed_dim, bias=bias) + self.norm = nn.RMSNorm(embed_dim, eps=1e-6) if use_norm else nn.Identity() def forward(self, x): return self.norm(self.proj(x)) class PixelTokenEmbedder(nn.Module): - - def __init__(self, in_channels, hidden_size_output, dtype=None, device=None, operations=None): + def __init__(self, in_channels, hidden_size_output): super().__init__() + self.in_channels = in_channels self.hidden_size_output = hidden_size_output - self.proj = operations.Linear(self.in_channels, self.hidden_size_output, bias=True, dtype=dtype, device=device) + self.proj = nn.Linear(self.in_channels, self.hidden_size_output, bias=True) def forward(self, inputs, patch_size): B, _, H, W = inputs.shape @@ -103,9 +108,9 @@ class PixelTokenEmbedder(nn.Module): class PiTBlock(nn.Module): - - def __init__(self, pixel_hidden_size, patch_hidden_size, patch_size, num_heads, mlp_ratio=4.0, attn_hidden_size=None, attn_num_heads=None, dtype=None, device=None, operations=None, mlp_chunks=1): + def __init__(self, pixel_hidden_size, patch_hidden_size, patch_size, num_heads, mlp_ratio=4.0, attn_hidden_size=None, attn_num_heads=None, mlp_chunks=1): super().__init__() + self.pixel_dim = pixel_hidden_size self.context_dim = patch_hidden_size self.attn_dim = attn_hidden_size if attn_hidden_size is not None else patch_hidden_size @@ -113,16 +118,16 @@ class PiTBlock(nn.Module): assert self.attn_dim % self.num_heads == 0 p2 = patch_size * patch_size - self.compress_to_attn = operations.Linear(p2 * self.pixel_dim, self.attn_dim, bias=True, dtype=dtype, device=device) - self.expand_from_attn = operations.Linear(self.attn_dim, p2 * self.pixel_dim, bias=True, dtype=dtype, device=device) + self.compress_to_attn = nn.Linear(p2 * self.pixel_dim, self.attn_dim, bias=True) + self.expand_from_attn = nn.Linear(self.attn_dim, p2 * self.pixel_dim, bias=True) - self.norm1 = operations.RMSNorm(self.pixel_dim, eps=1e-6, dtype=dtype, device=device) - self.attn = RotaryAttention(self.attn_dim, num_heads=self.num_heads, qkv_bias=False, dtype=dtype, device=device, operations=operations) - self.norm2 = operations.RMSNorm(self.pixel_dim, eps=1e-6, dtype=dtype, device=device) - self.mlp = Mlp(self.pixel_dim, hidden_features=int(self.pixel_dim * mlp_ratio), dtype=dtype, device=device, operations=operations) + self.norm1 = nn.RMSNorm(self.pixel_dim, eps=1e-6) + self.attn = RotaryAttention(self.attn_dim, num_heads=self.num_heads, qkv_bias=False) + self.norm2 = nn.RMSNorm(self.pixel_dim, eps=1e-6) + self.mlp = Mlp(self.pixel_dim, hidden_features=int(self.pixel_dim * mlp_ratio)) - self.adaLN_modulation_msa = operations.Linear(self.context_dim, 3 * self.pixel_dim * p2, bias=True, dtype=dtype, device=device) - self.adaLN_modulation_mlp = operations.Linear(self.context_dim, 3 * self.pixel_dim * p2, bias=True, dtype=dtype, device=device) + self.adaLN_modulation_msa = nn.Linear(self.context_dim, 3 * self.pixel_dim * p2, bias=True) + self.adaLN_modulation_mlp = nn.Linear(self.context_dim, 3 * self.pixel_dim * p2, bias=True) self._rope_fn = precompute_freqs_cis_2d self.mlp_chunks = max(1, int(mlp_chunks)) @@ -136,7 +141,6 @@ class PiTBlock(nn.Module): L = Hs * Ws B = BL // L - # Attention path uses only msa params; compute, use, free before mlp params allocate. msa_params = self.adaLN_modulation_msa(s_cond).view(BL, P2, 3 * self.pixel_dim) shift_msa, scale_msa, gate_msa = msa_params.chunk(3, dim=-1) @@ -153,13 +157,13 @@ class PiTBlock(nn.Module): mlp_params = self.adaLN_modulation_mlp(s_cond).view(BL, P2, 3 * self.pixel_dim) shift_mlp, scale_mlp, gate_mlp = mlp_params.chunk(3, dim=-1) - gate_mlp = gate_mlp.contiguous() # detach from mlp_params so the del below frees shift+scale storage before the MLP + gate_mlp = gate_mlp.contiguous() mlp_input = apply_adaln_(self.norm2(x), shift_mlp, scale_mlp) del mlp_params, shift_mlp, scale_mlp - # MLP in chunks since the peak memory usage is huge here chunk_size = (BL + self.mlp_chunks - 1) // self.mlp_chunks for s in range(0, BL, chunk_size): e = min(s + chunk_size, BL) x[s:e].addcmul_(gate_mlp[s:e], self.mlp(mlp_input[s:e])) + return x diff --git a/extensions-builtin/PiD/nn/pid.py b/extensions-builtin/PiD/nn/pid.py index 95e37a1f..6161f7af 100644 --- a/extensions-builtin/PiD/nn/pid.py +++ b/extensions-builtin/PiD/nn/pid.py @@ -1,3 +1,5 @@ +# https://github.com/Comfy-Org/ComfyUI/blob/v0.24.1/comfy/ldm/pixeldit/pid.py + import torch import torch.nn as nn import torch.nn.functional as F @@ -7,15 +9,14 @@ from .modules import precompute_freqs_cis_2d class SigmaAwareGatePerTokenPerDim(nn.Module): - - def __init__(self, dim: int, dtype=None, device=None, operations=None): + def __init__(self, dim: int): super().__init__() - self.content_proj = operations.Linear(dim * 2, dim, dtype=dtype, device=device) - self.log_alpha = nn.Parameter(torch.empty((), dtype=dtype, device=device)) + + self.content_proj = nn.Linear(dim * 2, dim) + self.log_alpha = nn.Parameter(torch.empty(())) def forward(self, x: torch.Tensor, lq: torch.Tensor, sigma: torch.Tensor) -> torch.Tensor: content_logit = self.content_proj(torch.cat([x, lq], dim=-1)) - # log_alpha is a raw nn.Parameter -> doesn't auto-cast under dynamic VRAM. log_alpha = self.log_alpha.to(device=x.device, dtype=torch.float32) sigma_offset = -log_alpha.exp() * sigma.float().view(-1, 1, 1) gate = torch.sigmoid(content_logit + sigma_offset) @@ -23,16 +24,16 @@ class SigmaAwareGatePerTokenPerDim(nn.Module): class ResBlock(nn.Module): - - def __init__(self, channels: int, num_groups: int = 4, dtype=None, device=None, operations=None): + def __init__(self, channels: int, num_groups: int = 4): super().__init__() + self.block = nn.Sequential( - operations.GroupNorm(num_groups, channels, dtype=dtype, device=device), + nn.GroupNorm(num_groups, channels), nn.SiLU(), - operations.Conv2d(channels, channels, kernel_size=3, padding=1, dtype=dtype, device=device), - operations.GroupNorm(num_groups, channels, dtype=dtype, device=device), + nn.Conv2d(channels, channels, kernel_size=3, padding=1), + nn.GroupNorm(num_groups, channels), nn.SiLU(), - operations.Conv2d(channels, channels, kernel_size=3, padding=1, dtype=dtype, device=device), + nn.Conv2d(channels, channels, kernel_size=3, padding=1), ) def forward(self, x: torch.Tensor) -> torch.Tensor: @@ -40,7 +41,6 @@ class ResBlock(nn.Module): class LQProjection2D(nn.Module): - def __init__( self, latent_channels: int, @@ -52,11 +52,9 @@ class LQProjection2D(nn.Module): num_res_blocks: int = 4, num_outputs: int = 7, interval: int = 2, - dtype=None, - device=None, - operations=None, ): super().__init__() + self.latent_channels = latent_channels self.hidden_dim = hidden_dim self.out_dim = out_dim @@ -78,16 +76,18 @@ class LQProjection2D(nn.Module): latent_proj_in_ch = latent_channels * fold_factor * fold_factor layers = [ - operations.Conv2d(latent_proj_in_ch, hidden_dim, kernel_size=3, padding=1, dtype=dtype, device=device), + nn.Conv2d(latent_proj_in_ch, hidden_dim, kernel_size=3, padding=1), nn.SiLU(), - operations.Conv2d(hidden_dim, hidden_dim, kernel_size=3, padding=1, dtype=dtype, device=device), + nn.Conv2d(hidden_dim, hidden_dim, kernel_size=3, padding=1), ] + for _ in range(num_res_blocks): - layers.append(ResBlock(hidden_dim, dtype=dtype, device=device, operations=operations)) + layers.append(ResBlock(hidden_dim)) + self.latent_proj = nn.Sequential(*layers) - self.output_heads = nn.ModuleList([operations.Linear(hidden_dim, out_dim, dtype=dtype, device=device) for _ in range(num_outputs)]) - self.gate_modules = nn.ModuleList([SigmaAwareGatePerTokenPerDim(out_dim, dtype=dtype, device=device, operations=operations) for _ in range(num_outputs)]) + self.output_heads = nn.ModuleList([nn.Linear(hidden_dim, out_dim) for _ in range(num_outputs)]) + self.gate_modules = nn.ModuleList([SigmaAwareGatePerTokenPerDim(out_dim) for _ in range(num_outputs)]) def is_gate_active(self, block_idx: int) -> bool: return block_idx % self.interval == 0 @@ -122,7 +122,6 @@ class LQProjection2D(nn.Module): class PidNet(PixDiT_T2I): - def __init__( self, lq_latent_channels: int = 16, @@ -131,20 +130,15 @@ class PidNet(PixDiT_T2I): lq_interval: int = 2, sr_scale: int = 4, latent_spatial_down_factor: int = 8, - rope_ref_h: int = 1024, # NTK ref resolution in PIXEL units: 1024px / patch=16 -> grid_ref=64. + rope_ref_h: int = 1024, rope_ref_w: int = 1024, - image_model=None, - dtype=None, - device=None, - operations=None, **pixdit_kwargs, ): - super().__init__(dtype=dtype, device=device, operations=operations, **pixdit_kwargs) + super().__init__(**pixdit_kwargs) self.rope_ref_grid_h = rope_ref_h // self.patch_size self.rope_ref_grid_w = rope_ref_w // self.patch_size - # Parent's PiTBlocks were built with plain RoPE — swap in NTK-aware. def _pit_rope_fn(head_dim, h, w, device=None, dtype=torch.float32, **rope_opts): return precompute_freqs_cis_2d(head_dim, h, w, ref_grid_h=self.rope_ref_grid_h, ref_grid_w=self.rope_ref_grid_w, device=device, dtype=dtype, **rope_opts) @@ -162,9 +156,6 @@ class PidNet(PixDiT_T2I): num_res_blocks=lq_num_res_blocks, num_outputs=num_lq_outputs, interval=lq_interval, - dtype=dtype, - device=device, - operations=operations, ) def _fetch_patch_pos(self, height, width, device, dtype, **rope_opts): @@ -188,13 +179,12 @@ class PidNet(PixDiT_T2I): return self.lq_proj.gate(s, pid_lq_features[out_idx], pid_degrade_sigma, out_idx) def forward(self, x, timesteps, context=None, attention_mask=None, transformer_options={}, lq_latent=None, degrade_sigma=None, **kwargs): - if lq_latent is None: - raise ValueError("PidNet requires lq_latent — attach via PiDConditioning") + assert lq_latent is not None + expected_c = self.lq_proj.latent_channels - if lq_latent.shape[1] != expected_c: - raise ValueError(f"Input latent has {lq_latent.shape[1]} channels, this model variant expects {expected_c}. " f"Flux1/SD3 = 16 channels, Flux2 = 128 channels.") + assert lq_latent.shape[1] == expected_c + B = x.shape[0] - # Match the backbone's pad_to_patch_size (round up) so the LQ grid lines up with the patch stream. Hs = -(-x.shape[2] // self.patch_size) Ws = -(-x.shape[3] // self.patch_size)