diff --git a/Sana/loader.py b/Sana/loader.py index b2aea1b..129e898 100644 --- a/Sana/loader.py +++ b/Sana/loader.py @@ -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] diff --git a/Sana/models/sana.py b/Sana/models/sana.py index 085eb65..113adfb 100644 --- a/Sana/models/sana.py +++ b/Sana/models/sana.py @@ -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 diff --git a/Sana/models/sana_multi_scale.py b/Sana/models/sana_multi_scale.py index e69e183..a00e6f5 100644 --- a/Sana/models/sana_multi_scale.py +++ b/Sana/models/sana_multi_scale.py @@ -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( [