Faster model initialization

don't initialize base Sana model for SanaMS
This commit is contained in:
City
2024-12-11 22:33:57 +01:00
parent 7e19ac5e43
commit 2beff5a4aa
3 changed files with 48 additions and 41 deletions
+7
View File
@@ -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
View File
@@ -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
+25 -33
View File
@@ -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(
[