Download nanosaur2_support/model.py from levzalt/Nanosaur2-Inpaint-ControlNet: direct link, hf CLI and curl.
- Browser
- Download file 14.5 kB
-
https://huggingface.co/levzalt/Nanosaur2-Inpaint-ControlNet/resolve/main/nanosaur2_support/model.py
- Command line
-
hf download hf://levzalt/Nanosaur2-Inpaint-ControlNet/nanosaur2_support/model.py
-
curl -L -o model.py https://huggingface.co/levzalt/Nanosaur2-Inpaint-ControlNet/resolve/main/nanosaur2_support/model.py
14.5 kB
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| import comfy.model_management | |
| import comfy.ops | |
| import comfy.patcher_extension | |
| import comfy.quant_ops | |
| from comfy.ldm.modules.attention import optimized_attention | |
| from comfy.ldm.modules.diffusionmodules.mmdit import TimestepEmbedder | |
| SPRINT_NUM_F = 2 | |
| SPRINT_NUM_H = 2 | |
| TEXT_EMBED_DIM = 640 | |
| def rope_2d(head_dim, height, width, device): | |
| """Split-half 2D RoPE over centered token coordinates (x frequencies, then y), as (1, N, 1, head_dim // 2, 2, 2) rotations.""" | |
| axis_dim = head_dim // 2 | |
| inv_freq = 1.0 / (10000.0 ** (torch.arange(0, axis_dim, 2, dtype=torch.float32, device=device) / axis_dim)) | |
| y = torch.arange(height, dtype=torch.float32, device=device) - (height - 1) / 2 | |
| x = torch.arange(width, dtype=torch.float32, device=device) - (width - 1) / 2 | |
| y, x = torch.meshgrid(y, x, indexing="ij") | |
| angles = torch.cat([torch.outer(x.flatten(), inv_freq), torch.outer(y.flatten(), inv_freq)], dim=-1) | |
| cos, sin = torch.cos(angles), torch.sin(angles) | |
| return torch.stack([cos, -sin, sin, cos], dim=-1).view(1, height * width, 1, axis_dim, 2, 2) | |
| def modulate(x, shift, scale): | |
| return torch.addcmul(shift, x, 1 + scale) | |
| def rope_split_half(x, pe): | |
| """Differentiable ck.apply_rope_split_half1: pair k is (x[k], x[k + head_dim // 2]).""" | |
| x_ = x.reshape(*x.shape[:-1], 2, -1).movedim(-2, -1).unsqueeze(-2).to(pe.dtype) | |
| return (pe[..., 0] * x_[..., 0] + pe[..., 1] * x_[..., 1]).movedim(-1, -2).reshape(x.shape).type_as(x) | |
| def norm_rope(q, k, q_norm, k_norm, pe): | |
| if not comfy.model_management.in_training: | |
| q_scale, _, q_offload_stream = comfy.ops.cast_bias_weight(q_norm, q, offloadable=True) | |
| k_scale, _, k_offload_stream = comfy.ops.cast_bias_weight(k_norm, k, offloadable=True) | |
| q, k = comfy.quant_ops.ck.rms_rope_split_half(q, k, pe, q_scale, k_scale, q_norm.eps) | |
| comfy.ops.uncast_bias_weight(q_norm, q_scale, None, q_offload_stream) | |
| comfy.ops.uncast_bias_weight(k_norm, k_scale, None, k_offload_stream) | |
| return q, k | |
| return rope_split_half(q_norm(q), pe), rope_split_half(k_norm(k), pe) | |
| class Embed(nn.Module): | |
| def __init__(self, in_dim, hidden_size, norm=False, dtype=None, device=None, operations=None): | |
| super().__init__() | |
| self.proj = operations.Linear(in_dim, hidden_size, bias=True, dtype=dtype, device=device) | |
| self.norm = operations.RMSNorm(hidden_size, eps=1e-6, dtype=dtype, device=device) if norm else nn.Identity() | |
| def forward(self, x): | |
| return self.norm(self.proj(x)) | |
| class FeedForward(nn.Module): | |
| def __init__(self, dim, hidden_dim, dtype=None, device=None, operations=None): | |
| super().__init__() | |
| self.w12 = operations.Linear(dim, hidden_dim * 2, bias=False, dtype=dtype, device=device) | |
| self.w3 = operations.Linear(hidden_dim, dim, bias=False, dtype=dtype, device=device) | |
| def forward(self, x): | |
| x1, x2 = self.w12(x).chunk(2, dim=-1) | |
| return self.w3(F.silu(x1) * x2) | |
| class Attention(nn.Module): | |
| """Self-attention over image tokens with 2D RoPE, optionally joint with (non-rotated) text keys and values.""" | |
| def __init__(self, dim, num_heads, cross_attention, dtype=None, device=None, operations=None): | |
| super().__init__() | |
| self.num_heads = num_heads | |
| self.head_dim = dim // num_heads | |
| self.qkv_x = operations.Linear(dim, dim * 3, bias=False, dtype=dtype, device=device) | |
| self.kv_y = operations.Linear(dim, dim * 2, bias=False, dtype=dtype, device=device) if cross_attention else None | |
| 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, bias=True, dtype=dtype, device=device) | |
| def forward(self, x, txt, pe, txt_bias, transformer_options={}): | |
| b, n, c = x.shape | |
| q, k, v = self.qkv_x(x).view(b, n, 3, self.num_heads, self.head_dim).unbind(2) | |
| q, k = norm_rope(q, k, self.q_norm, self.k_norm, pe) | |
| mask = None | |
| if self.kv_y is not None: | |
| ky, vy = self.kv_y(txt).view(b, -1, 2, self.num_heads, self.head_dim).unbind(2) | |
| k = torch.cat([k, self.k_norm(ky)], dim=1) | |
| v = torch.cat([v, vy], dim=1) | |
| # Image keys are always visible; text keys carry log(emphasis weight), -inf for padding. | |
| mask = F.pad(txt_bias, (n, 0))[:, None, None, :] | |
| x = optimized_attention(q.reshape(b, n, c), k.reshape(b, -1, c), v.reshape(b, -1, c), self.num_heads, mask=mask, transformer_options=transformer_options) | |
| return self.proj(x) | |
| class DiTBlock(nn.Module): | |
| """Image block; its adaLN modulation is computed by the caller (adaLN-single).""" | |
| def __init__(self, hidden_size, num_heads, mlp_hidden, cross_attention, dtype=None, device=None, operations=None): | |
| super().__init__() | |
| self.norm1 = operations.RMSNorm(hidden_size, eps=1e-6, dtype=dtype, device=device) | |
| self.attn = Attention(hidden_size, num_heads, cross_attention, dtype=dtype, device=device, operations=operations) | |
| self.norm2 = operations.RMSNorm(hidden_size, eps=1e-6, dtype=dtype, device=device) | |
| self.mlp = FeedForward(hidden_size, mlp_hidden, dtype=dtype, device=device, operations=operations) | |
| def forward(self, x, txt, pe, mod, txt_bias, transformer_options={}): | |
| shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = mod.chunk(6, dim=-1) | |
| attention_input = modulate(self.norm1(x), shift_msa, scale_msa) | |
| for adapter, features, strength in transformer_options.get("nanosaur2_inpaint", []): | |
| attention_input = attention_input + strength * adapter.inject[2 * transformer_options["nanosaur2_block"]](attention_input, features) | |
| x = torch.addcmul(x, gate_msa, self.attn(attention_input, txt, pe, txt_bias, transformer_options=transformer_options)) | |
| mlp_input = modulate(self.norm2(x), shift_mlp, scale_mlp) | |
| for adapter, features, strength in transformer_options.get("nanosaur2_inpaint", []): | |
| mlp_input = mlp_input + strength * adapter.inject[2 * transformer_options["nanosaur2_block"] + 1](mlp_input, features) | |
| return torch.addcmul(x, gate_mlp, self.mlp(mlp_input)) | |
| class TextRefineAttention(nn.Module): | |
| def __init__(self, dim, num_heads, dtype=None, device=None, operations=None): | |
| super().__init__() | |
| self.num_heads = num_heads | |
| self.head_dim = dim // num_heads | |
| self.qkv = operations.Linear(dim, dim * 3, bias=False, 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, bias=True, dtype=dtype, device=device) | |
| def forward(self, x, mask, transformer_options={}): | |
| b, n, c = x.shape | |
| q, k, v = self.qkv(x).view(b, n, 3, self.num_heads, self.head_dim).unbind(2) | |
| x = optimized_attention(self.q_norm(q).reshape(b, n, c), self.k_norm(k).reshape(b, n, c), v.reshape(b, n, c), self.num_heads, mask=mask, transformer_options=transformer_options) | |
| return self.proj(x) | |
| class TextRefineBlock(nn.Module): | |
| def __init__(self, hidden_size, num_heads, mlp_hidden, dtype=None, device=None, operations=None): | |
| super().__init__() | |
| self.norm1 = operations.RMSNorm(hidden_size, eps=1e-6, dtype=dtype, device=device) | |
| self.attn = TextRefineAttention(hidden_size, num_heads, dtype=dtype, device=device, operations=operations) | |
| self.norm2 = operations.RMSNorm(hidden_size, eps=1e-6, dtype=dtype, device=device) | |
| self.mlp = FeedForward(hidden_size, mlp_hidden, dtype=dtype, device=device, operations=operations) | |
| self.adaLN_modulation = nn.Sequential(operations.Linear(hidden_size, 6 * hidden_size, bias=True, dtype=dtype, device=device)) | |
| def forward(self, x, c, mask, keep, transformer_options={}): | |
| shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(c).chunk(6, dim=-1) | |
| x = torch.addcmul(x, gate_msa, self.attn(modulate(self.norm1(x), shift_msa, scale_msa), mask, transformer_options=transformer_options)) | |
| x = torch.addcmul(x, gate_mlp, self.mlp(modulate(self.norm2(x), shift_mlp, scale_mlp))) | |
| return x * keep | |
| class FinalLayer(nn.Module): | |
| def __init__(self, hidden_size, out_channels, dtype=None, device=None, operations=None): | |
| super().__init__() | |
| self.norm_final = operations.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device) | |
| self.adaLN_modulation = operations.Linear(hidden_size, 2 * hidden_size, bias=True, dtype=dtype, device=device) | |
| self.linear = operations.Linear(hidden_size, out_channels, bias=True, dtype=dtype, device=device) | |
| def forward(self, x, c): | |
| shift, scale = self.adaLN_modulation(c).chunk(2, dim=-1) | |
| if comfy.model_management.in_training: | |
| return self.linear(modulate(self.norm_final(x), shift, scale)) | |
| return self.linear(comfy.quant_ops.ck.adaln(x, scale, shift, self.norm_final.eps)) | |
| class Nanosaur2Transformer2DModel(nn.Module): | |
| def __init__(self, in_channels=64, hidden_size=1536, num_heads=16, num_blocks=18, num_text_blocks=2, mlp_hidden=4096, dtype=None, device=None, operations=None): | |
| super().__init__() | |
| self.dtype = dtype | |
| self.in_channels = in_channels | |
| self.num_heads = num_heads | |
| self.head_dim = hidden_size // num_heads | |
| self.s_embedder = Embed(in_channels, hidden_size, dtype=dtype, device=device, operations=operations) | |
| self.t_embedder = TimestepEmbedder(hidden_size, dtype=dtype, device=device, operations=operations) | |
| self.y_embedder = Embed(TEXT_EMBED_DIM, hidden_size, norm=True, dtype=dtype, device=device, operations=operations) | |
| self.shared_encoder_adaLN = nn.Sequential(operations.Linear(hidden_size, 6 * hidden_size, bias=True, dtype=dtype, device=device)) | |
| self.encoder_adaLN_offsets = nn.ParameterList([nn.Parameter(torch.empty(6 * hidden_size, dtype=dtype, device=device)) for _ in range(num_blocks)]) | |
| self.y_pool_proj = operations.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device) | |
| self.sprint_out_proj = operations.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device) | |
| self.blocks = nn.ModuleList([DiTBlock(hidden_size, num_heads, mlp_hidden, i % 2 == 0, dtype=dtype, device=device, operations=operations) for i in range(num_blocks)]) | |
| self.text_refine_blocks = nn.ModuleList([TextRefineBlock(hidden_size, num_heads, mlp_hidden, dtype=dtype, device=device, operations=operations) for _ in range(num_text_blocks)]) | |
| self.final_layer = FinalLayer(hidden_size, in_channels, dtype=dtype, device=device, operations=operations) | |
| def run_block(self, i, x, txt, pe, mod, txt_bias, transformer_options): | |
| offset = comfy.model_management.cast_to(self.encoder_adaLN_offsets[i], dtype=mod.dtype, device=mod.device) | |
| transformer_options = {**transformer_options, "nanosaur2_block": i} | |
| return self.blocks[i](x, txt, pe, mod + offset, txt_bias, transformer_options=transformer_options) | |
| def sparse_path(self, x, txt, pe, mod, txt_bias, transformer_options): | |
| g = x | |
| for i in range(SPRINT_NUM_F, len(self.blocks) - SPRINT_NUM_H): | |
| g = self.run_block(i, g, txt, pe, mod, txt_bias, transformer_options) | |
| return x + self.sprint_out_proj(g - x) | |
| def forward(self, x, timestep, context, token_weights, sparse_skip=None, transformer_options={}, **kwargs): | |
| return comfy.patcher_extension.WrapperExecutor.new_class_executor( | |
| self._forward, | |
| self, | |
| comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, transformer_options) | |
| ).execute(x, timestep, context, token_weights, sparse_skip, transformer_options, **kwargs) | |
| def _forward(self, x, timestep, context, token_weights, sparse_skip=None, transformer_options={}, **kwargs): | |
| """sparse_skip: optional per-row flags; flagged rows skip the sparse middle blocks.""" | |
| b, _, h, w = x.shape | |
| keep = token_weights > 0 | |
| txt_bias = torch.log(token_weights.clamp(min=1e-4)).masked_fill(~keep, float("-inf")) | |
| # Refine attention is a hard mask; a row with no text key keeps its first one and is zeroed after. | |
| refine_keep = keep.clone() | |
| refine_keep[:, 0] |= ~keep.any(dim=1) | |
| refine_mask = torch.zeros_like(txt_bias).masked_fill(~refine_keep, float("-inf"))[:, None, None, :] | |
| keep = keep.unsqueeze(-1).to(x.dtype) | |
| t = self.t_embedder(timestep * 1000.0, x.dtype).unsqueeze(1) | |
| txt = self.y_embedder(context) | |
| time_condition = F.silu(t) | |
| for block in self.text_refine_blocks: | |
| txt = block(txt, time_condition, refine_mask, keep, transformer_options=transformer_options) | |
| weights = token_weights.unsqueeze(-1) | |
| pooled = (txt * weights).sum(dim=1) / weights.sum(dim=1).clamp(min=1.0) | |
| condition = F.silu(t + self.y_pool_proj(pooled).unsqueeze(1)) | |
| mod = self.shared_encoder_adaLN(condition) | |
| pe = rope_2d(self.head_dim, h, w, x.device) | |
| s = self.s_embedder(x.flatten(2).transpose(1, 2)) | |
| for i in range(SPRINT_NUM_F): | |
| s = self.run_block(i, s, txt, pe, mod, txt_bias, transformer_options) | |
| if sparse_skip is None or not any(sparse_skip): | |
| s = self.sparse_path(s, txt, pe, mod, txt_bias, transformer_options) | |
| elif not all(sparse_skip): | |
| rows = [i for i, skip in enumerate(sparse_skip) if not skip] | |
| sparse_options = transformer_options.copy() | |
| sparse_options["nanosaur2_inpaint"] = [(adapter, features[rows], strength) | |
| for adapter, features, strength in transformer_options.get("nanosaur2_inpaint", [])] | |
| s[rows] = self.sparse_path(s[rows], txt[rows], pe, mod[rows], txt_bias[rows], sparse_options) | |
| for i in range(len(self.blocks) - SPRINT_NUM_H, len(self.blocks)): | |
| s = self.run_block(i, s, txt, pe, mod, txt_bias, transformer_options) | |
| x0 = self.final_layer(s, condition).transpose(1, 2).reshape(b, self.in_channels, h, w) | |
| return (x - x0) / timestep.view(-1, 1, 1, 1) | |