Faster model initialization
don't initialize base Sana model for SanaMS
This commit is contained in:
@@ -1,3 +1,4 @@
|
||||
import math
|
||||
import logging
|
||||
|
||||
import comfy.utils
|
||||
@@ -76,6 +77,12 @@ def model_config_from_unet(sd):
|
||||
}
|
||||
config["depth"] = sum([key.endswith(".point_conv.conv.weight") for key in sd.keys()]) or 28
|
||||
|
||||
if "pos_embed" in sd:
|
||||
config["input_size"] = int(math.sqrt(sd["pos_embed"].shape[1])) * config["patch_size"]
|
||||
else:
|
||||
# TODO: this isn't optimal though most models don't use it
|
||||
config["use_pe"] = False
|
||||
|
||||
if "x_embedder.proj.bias" in sd:
|
||||
config["hidden_size"] = sd["x_embedder.proj.bias"].shape[0]
|
||||
|
||||
|
||||
+16
-8
@@ -34,7 +34,7 @@ from .sana_blocks import (
|
||||
t2i_modulate,
|
||||
)
|
||||
from .norms import RMSNorm
|
||||
from .utils import auto_grad_checkpoint, to_2tuple
|
||||
from .utils import to_2tuple
|
||||
|
||||
|
||||
class SanaBlock(nn.Module):
|
||||
@@ -166,7 +166,7 @@ class Sana(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
input_size=32,
|
||||
input_size=None,
|
||||
patch_size=1,
|
||||
in_channels=32,
|
||||
hidden_size=1152,
|
||||
@@ -196,7 +196,7 @@ class Sana(nn.Module):
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.dtype = torch.float32
|
||||
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
|
||||
@@ -215,10 +215,18 @@ class Sana(nn.Module):
|
||||
dtype=dtype, device=device, operations=operations
|
||||
)
|
||||
self.t_embedder = TimestepEmbedder(hidden_size, dtype=dtype, device=device, operations=operations)
|
||||
|
||||
if input_size is not None:
|
||||
self.base_size = input_size // self.patch_size
|
||||
else:
|
||||
self.base_size = None
|
||||
|
||||
num_patches = self.x_embedder.num_patches
|
||||
self.base_size = input_size // self.patch_size
|
||||
# Will use fixed sin-cos embedding:
|
||||
self.register_buffer("pos_embed", torch.zeros(1, num_patches, hidden_size))
|
||||
if self.use_pe and num_patches is not None:
|
||||
#Will use fixed sin-cos embedding:
|
||||
self.register_buffer("pos_embed", torch.zeros(1, num_patches, hidden_size))
|
||||
else:
|
||||
self.pos_embed = None
|
||||
|
||||
approx_gelu = lambda: nn.GELU(approximate="tanh")
|
||||
self.t_block = nn.Sequential(
|
||||
@@ -269,9 +277,9 @@ class Sana(nn.Module):
|
||||
y: (N, 1, 120, C) tensor of class labels
|
||||
"""
|
||||
y = context # remap comfy cond name
|
||||
pos_embed = self.pos_embed.to(x.dtype)
|
||||
self.h, self.w = x.shape[-2] // self.patch_size, x.shape[-1] // self.patch_size
|
||||
if self.use_pe:
|
||||
pos_embed = self.pos_embed.to(x.dtype)
|
||||
x = self.x_embedder(x) + pos_embed # (N, T, D), where T = H * W / patch_size ** 2
|
||||
else:
|
||||
x = self.x_embedder(x)
|
||||
@@ -290,7 +298,7 @@ class Sana(nn.Module):
|
||||
y_lens = [y.shape[2]] * y.shape[0]
|
||||
y = y.squeeze(1).view(1, -1, x.shape[-1])
|
||||
for block in self.blocks:
|
||||
x = auto_grad_checkpoint(block, x, y, t0, y_lens) # (N, T, D) #support grad checkpoint
|
||||
x = block(x, y, t0, y_lens) # (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)
|
||||
return x
|
||||
|
||||
@@ -23,6 +23,7 @@ from .sana import Sana, get_2d_sincos_pos_embed
|
||||
from .sana_blocks import (
|
||||
Attention,
|
||||
CaptionEmbedder,
|
||||
TimestepEmbedder,
|
||||
FlashAttention,
|
||||
LiteLA,
|
||||
MultiHeadCrossAttention,
|
||||
@@ -30,8 +31,7 @@ from .sana_blocks import (
|
||||
T2IFinalLayer,
|
||||
t2i_modulate,
|
||||
)
|
||||
from .utils import auto_grad_checkpoint
|
||||
|
||||
from .norms import RMSNorm
|
||||
|
||||
class SanaMSBlock(nn.Module):
|
||||
"""
|
||||
@@ -195,44 +195,33 @@ class SanaMS(Sana):
|
||||
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,
|
||||
caption_channels=caption_channels,
|
||||
pe_interpolation=pe_interpolation,
|
||||
config=config,
|
||||
model_max_length=model_max_length,
|
||||
qk_norm=qk_norm,
|
||||
y_norm=y_norm,
|
||||
norm_eps=norm_eps,
|
||||
attn_type=attn_type,
|
||||
ffn_type=ffn_type,
|
||||
use_pe=use_pe,
|
||||
y_norm_scale_factor=y_norm_scale_factor,
|
||||
patch_embed_kernel=patch_embed_kernel,
|
||||
mlp_acts=mlp_acts,
|
||||
linear_head_dim=linear_head_dim,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
operations=operations,
|
||||
**kwargs,
|
||||
)
|
||||
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.depth = depth
|
||||
self.use_pe = use_pe
|
||||
self.y_norm = y_norm
|
||||
self.model_max_length = model_max_length
|
||||
self.fp32_attention = kwargs.get("use_fp32_attention", False)
|
||||
self.h = self.w = 0
|
||||
|
||||
approx_gelu = lambda: nn.GELU(approximate="tanh")
|
||||
self.t_block = nn.Sequential(
|
||||
nn.SiLU(), operations.Linear(hidden_size, 6 * hidden_size, bias=True, dtype=dtype, device=device)
|
||||
)
|
||||
|
||||
self.t_embedder = TimestepEmbedder(hidden_size, dtype=dtype, device=device, operations=operations)
|
||||
|
||||
self.pos_embed_ms = None
|
||||
if input_size is not None:
|
||||
self.base_size = input_size // self.patch_size
|
||||
else:
|
||||
self.base_size = None
|
||||
|
||||
kernel_size = patch_embed_kernel or patch_size
|
||||
self.x_embedder = PatchEmbedMS(
|
||||
@@ -249,6 +238,9 @@ class SanaMS(Sana):
|
||||
device=device,
|
||||
operations=operations,
|
||||
)
|
||||
if self.y_norm:
|
||||
self.attention_y_norm = RMSNorm(hidden_size, scale_factor=y_norm_scale_factor, eps=norm_eps)
|
||||
|
||||
drop_path = [x.item() for x in torch.linspace(0, drop_path, depth)] # stochastic depth decay rule
|
||||
self.blocks = nn.ModuleList(
|
||||
[
|
||||
|
||||
Reference in New Issue
Block a user