diff --git a/Sana/loader.py b/Sana/loader.py index 60eea7a..b2aea1b 100644 --- a/Sana/loader.py +++ b/Sana/loader.py @@ -31,16 +31,11 @@ class SanaConfig(comfy.supported_models_base.BASE): return comfy.model_base.ModelType.FLOW def get_model(self, state_dict, prefix="", device=None): - return SanaModel( - model_config=self, - model_type=comfy.model_base.ModelType.FLOW, - unet_model=self.unet_class, - device=device - ) + return SanaModel(model_config=self, unet_model=self.unet_class, device=device) class SanaModel(comfy.model_base.BaseModel): - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) + def __init__(self, *args, model_type=comfy.model_base.ModelType.FLOW, unet_model=SanaMS, **kwargs): + super().__init__(*args, model_type=model_type, unet_model=unet_model, **kwargs) def load_sana_state_dict(sd, model_options={}): # prefix / format diff --git a/Sana/models/act.py b/Sana/models/act.py index 9df6a7a..7993065 100644 --- a/Sana/models/act.py +++ b/Sana/models/act.py @@ -37,7 +37,7 @@ REGISTERED_ACT_DICT: dict[str, tuple[type, dict[str, any]]] = { } -def build_act(name: str or None, **kwargs) -> nn.Module or None: +def build_act(name, **kwargs): if name in REGISTERED_ACT_DICT: act_cls, default_args = copy.deepcopy(REGISTERED_ACT_DICT[name]) for key in default_args: @@ -50,7 +50,7 @@ def build_act(name: str or None, **kwargs) -> nn.Module or None: raise ValueError(f"do not support: {name}") -def get_act_name(act: nn.Module or None) -> str or None: +def get_act_name(act): if act is None: return None module2name = {} diff --git a/Sana/models/basic_modules.py b/Sana/models/basic_modules.py index ece579a..6b279a4 100644 --- a/Sana/models/basic_modules.py +++ b/Sana/models/basic_modules.py @@ -17,13 +17,31 @@ # This file is modified from https://github.com/PixArt-alpha/PixArt-sigma import torch import torch.nn as nn -from timm.models.vision_transformer import Mlp +#from timm.models.vision_transformer import Mlp from .act import build_act, get_act_name from .norms import build_norm, get_norm_name from .utils import get_same_padding, val2tuple +class Mlp(nn.Module): + def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, bias=True, drop=None, dtype=None, device=None, operations=None) -> None: + super().__init__() + out_features = out_features or in_features + hidden_features = hidden_features or in_features + + self.fc1 = operations.Linear(in_features, hidden_features, bias=bias, dtype=dtype, device=device) + self.act = act_layer() + self.fc2 = operations.Linear(hidden_features, out_features, bias=bias, dtype=dtype, device=device) + + self.drop1 = nn.Identity() + self.drop2 = nn.Identity() + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = self.act(self.fc1(x)) + return self.fc2(x) + + class ConvLayer(nn.Module): def __init__( self, @@ -33,11 +51,14 @@ class ConvLayer(nn.Module): stride=1, dilation=1, groups=1, - padding: int or None = None, + padding=None, use_bias=False, dropout=0.0, norm="bn2d", act="relu", + dtype=None, + device=None, + operations=None, ): super().__init__() if padding is None: @@ -54,7 +75,7 @@ class ConvLayer(nn.Module): self.use_bias = use_bias self.dropout = nn.Dropout2d(dropout, inplace=False) if dropout > 0 else None - self.conv = nn.Conv2d( + self.conv = operations.Conv2d( in_dim, out_dim, kernel_size=(kernel_size, kernel_size), @@ -63,6 +84,8 @@ class ConvLayer(nn.Module): dilation=(dilation, dilation), groups=groups, bias=use_bias, + dtype=dtype, + device=device, ) self.norm = build_norm(norm, num_features=out_dim) self.act = build_act(act) @@ -86,11 +109,14 @@ class GLUMBConv(nn.Module): out_feature=None, kernel_size=3, stride=1, - padding: int or None = None, + padding=None, use_bias=False, norm=(None, None, None), act=("silu", "silu", None), dilation=1, + dtype=None, + device=None, + operations=None, ): out_feature = out_feature or in_features super().__init__() @@ -106,6 +132,9 @@ class GLUMBConv(nn.Module): use_bias=use_bias[0], norm=norm[0], act=act[0], + dtype=dtype, + device=device, + operations=operations, ) self.depth_conv = ConvLayer( hidden_features * 2, @@ -118,6 +147,9 @@ class GLUMBConv(nn.Module): norm=norm[1], act=None, dilation=dilation, + dtype=dtype, + device=device, + operations=operations, ) self.point_conv = ConvLayer( hidden_features, @@ -126,6 +158,9 @@ class GLUMBConv(nn.Module): use_bias=use_bias[2], norm=norm[2], act=act[2], + dtype=dtype, + device=device, + operations=operations, ) # from IPython import embed; embed(header='debug dilate conv') @@ -189,10 +224,13 @@ class MBConvPreGLU(nn.Module): stride=1, mid_dim=None, expand=6, - padding: int or None = None, + padding=None, use_bias=False, norm=(None, None, "ln2d"), act=("silu", "silu", None), + dtype=None, + device=None, + operations=None, ): super().__init__() use_bias = val2tuple(use_bias, 3) @@ -208,6 +246,9 @@ class MBConvPreGLU(nn.Module): use_bias=use_bias[0], norm=norm[0], act=None, + dtype=dtype, + device=device, + operations=operations, ) self.glu_act = build_act(act[0], inplace=False) self.depth_conv = ConvLayer( @@ -220,6 +261,9 @@ class MBConvPreGLU(nn.Module): use_bias=use_bias[1], norm=norm[1], act=act[1], + dtype=dtype, + device=device, + operations=operations, ) self.point_conv = ConvLayer( mid_dim, @@ -228,6 +272,9 @@ class MBConvPreGLU(nn.Module): use_bias=use_bias[2], norm=norm[2], act=act[2], + dtype=dtype, + device=device, + operations=operations, ) def forward(self, x: torch.Tensor, HW=None) -> torch.Tensor: @@ -283,6 +330,9 @@ class DWMlp(Mlp): stride=1, dilation=1, padding=None, + dtype=None, + device=None, + operations=None, ): super().__init__( in_features=in_features, @@ -291,6 +341,9 @@ class DWMlp(Mlp): act_layer=act_layer, bias=bias, drop=drop, + dtype=dtype, + device=device, + operations=operations, ) hidden_features = hidden_features or in_features self.hidden_features = hidden_features @@ -298,7 +351,7 @@ class DWMlp(Mlp): padding = get_same_padding(kernel_size) padding *= dilation - self.conv = nn.Conv2d( + self.conv = operations.Conv2d( hidden_features, hidden_features, kernel_size=(kernel_size, kernel_size), @@ -307,6 +360,8 @@ class DWMlp(Mlp): dilation=(dilation, dilation), groups=hidden_features, bias=bias, + dtype=dtype, + device=device, ) def forward(self, x, HW=None): @@ -324,38 +379,3 @@ class DWMlp(Mlp): x = self.fc2(x) x = self.drop2(x) return x - - -class Mlp(Mlp): - """MLP as used in Vision Transformer, MLP-Mixer and related networks""" - - def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, bias=True, drop=0.0): - super().__init__( - in_features=in_features, - hidden_features=hidden_features, - out_features=out_features, - act_layer=act_layer, - bias=bias, - drop=drop, - ) - - def forward(self, x, HW=None): - x = self.fc1(x) - x = self.act(x) - x = self.drop1(x) - x = self.fc2(x) - x = self.drop2(x) - return x - - -if __name__ == "__main__": - model = GLUMBConv( - 1152, - 1152 * 4, - 1152, - use_bias=(True, True, False), - norm=(None, None, None), - act=("silu", "silu", None), - ).cuda() - input = torch.randn(4, 256, 1152).cuda() - output = model(input) diff --git a/Sana/models/norms.py b/Sana/models/norms.py index 4731d69..5174fdc 100644 --- a/Sana/models/norms.py +++ b/Sana/models/norms.py @@ -48,7 +48,7 @@ REGISTERED_NORMALIZATION_DICT: dict[str, tuple[type, dict[str, any]]] = { } -def build_norm(name="bn2d", num_features=None, affine=True, **kwargs) -> nn.Module or None: +def build_norm(name="bn2d", num_features=None, affine=True, **kwargs): if name in ["ln", "ln2d"]: kwargs["normalized_shape"] = num_features kwargs["elementwise_affine"] = affine @@ -67,7 +67,7 @@ def build_norm(name="bn2d", num_features=None, affine=True, **kwargs) -> nn.Modu raise ValueError("do not support: %s" % name) -def get_norm_name(norm: nn.Module or None) -> str or None: +def get_norm_name(norm): if norm is None: return None module2name = {} @@ -171,7 +171,7 @@ def remove_bn(model: nn.Module) -> None: m.forward = lambda x: x -def set_norm_eps(model: nn.Module, eps: float or None = None, momentum: float or None = None) -> None: +def set_norm_eps(model, eps=None, momentum=None): for m in model.modules(): if isinstance(m, (nn.GroupNorm, nn.LayerNorm, _BatchNorm)): if eps is not None: diff --git a/Sana/models/sana.py b/Sana/models/sana.py index 81c2211..085eb65 100644 --- a/Sana/models/sana.py +++ b/Sana/models/sana.py @@ -20,7 +20,6 @@ import os import numpy as np import torch import torch.nn as nn -from timm.models.layers import DropPath from .basic_modules import DWMlp, GLUMBConv, MBConvPreGLU, Mlp from .sana_blocks import ( @@ -55,10 +54,13 @@ class SanaBlock(nn.Module): ffn_type="mlp", mlp_acts=("silu", "silu", None), linear_head_dim=32, + dtype=None, + device=None, + operations=None, **block_kwargs, ): super().__init__() - self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + self.norm1 = operations.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device) if attn_type == "flash": # flash self attention self.attn = FlashAttention( @@ -66,20 +68,28 @@ class SanaBlock(nn.Module): num_heads=num_heads, qkv_bias=True, qk_norm=qk_norm, + dtype=dtype, + device=device, + operations=operations, **block_kwargs, ) elif attn_type == "linear": # linear self attention # TODO: Here the num_heads set to 36 for tmp used self_num_heads = hidden_size // linear_head_dim - self.attn = LiteLA(hidden_size, hidden_size, heads=self_num_heads, eps=1e-8, qk_norm=qk_norm) + self.attn = LiteLA( + hidden_size, hidden_size, heads=self_num_heads, eps=1e-8, qk_norm=qk_norm, + dtype=dtype, device=device, operations=operations, + ) elif attn_type == "vanilla": # vanilla self attention - self.attn = Attention(hidden_size, num_heads=num_heads, qkv_bias=True) + self.attn = Attention( + hidden_size, num_heads=num_heads, qkv_bias=True, dtype=dtype, device=device, operations=operations, + ) else: raise ValueError(f"{attn_type} type is not defined.") - self.cross_attn = MultiHeadCrossAttention(hidden_size, num_heads, **block_kwargs) + self.cross_attn = MultiHeadCrossAttention(hidden_size, num_heads, dtype=dtype, device=device, operations=operations, **block_kwargs) self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) # to be compatible with lower version pytorch if ffn_type == "dwmlp": @@ -94,6 +104,9 @@ class SanaBlock(nn.Module): use_bias=(True, True, False), norm=(None, None, None), act=mlp_acts, + dtype=dtype, + device=device, + operations=operations, ) elif ffn_type == "glumbconv_dilate": self.mlp = GLUMBConv( @@ -103,6 +116,9 @@ class SanaBlock(nn.Module): norm=(None, None, None), act=mlp_acts, dilation=2, + dtype=dtype, + device=device, + operations=operations, ) elif ffn_type == "mbconvpreglu": self.mlp = MBConvPreGLU( @@ -112,15 +128,19 @@ class SanaBlock(nn.Module): use_bias=(True, True, False), norm=None, act=("silu", "silu", None), + dtype=dtype, + device=device, + operations=operations, ) elif ffn_type == "mlp": approx_gelu = lambda: nn.GELU(approximate="tanh") self.mlp = Mlp( - in_features=hidden_size, hidden_features=int(hidden_size * mlp_ratio), act_layer=approx_gelu, drop=0 + in_features=hidden_size, hidden_features=int(hidden_size * mlp_ratio), act_layer=approx_gelu, drop=0, + # dtype=dtype, device=device, operations=operations, ) else: raise ValueError(f"{ffn_type} type is not defined.") - self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity() + self.drop_path = nn.Identity() #DropPath(drop_path) if drop_path > 0.0 else nn.Identity() self.scale_shift_table = nn.Parameter(torch.randn(6, hidden_size) / hidden_size**0.5) def forward(self, x, y, t, mask=None, **kwargs): @@ -170,9 +190,13 @@ class Sana(nn.Module): patch_embed_kernel=None, mlp_acts=("silu", "silu", None), linear_head_dim=32, + dtype=None, + device=None, + operations=None, **kwargs, ): super().__init__() + self.dtype = torch.float32 self.pred_sigma = pred_sigma self.in_channels = in_channels self.out_channels = in_channels * 2 if pred_sigma else in_channels @@ -187,22 +211,28 @@ class Sana(nn.Module): kernel_size = patch_embed_kernel or patch_size self.x_embedder = PatchEmbed( - input_size, patch_size, in_channels, hidden_size, kernel_size=kernel_size, bias=True + input_size, patch_size, in_channels, hidden_size, kernel_size=kernel_size, bias=True, + dtype=dtype, device=device, operations=operations ) - self.t_embedder = TimestepEmbedder(hidden_size) + self.t_embedder = TimestepEmbedder(hidden_size, dtype=dtype, device=device, operations=operations) 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)) approx_gelu = lambda: nn.GELU(approximate="tanh") - self.t_block = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 6 * hidden_size, bias=True)) + self.t_block = nn.Sequential( + nn.SiLU(), operations.Linear(hidden_size, 6 * hidden_size, bias=True, dtype=dtype, device=device) + ) self.y_embedder = CaptionEmbedder( in_channels=caption_channels, hidden_size=hidden_size, uncond_prob=class_dropout_prob, act_layer=approx_gelu, token_num=model_max_length, + dtype=dtype, + device=device, + operations=operations ) if self.y_norm: self.attention_y_norm = RMSNorm(hidden_size, scale_factor=y_norm_scale_factor, eps=norm_eps) @@ -220,29 +250,32 @@ class Sana(nn.Module): ffn_type=ffn_type, mlp_acts=mlp_acts, linear_head_dim=linear_head_dim, + dtype=dtype, + device=device, + operations=operations ) for i in range(depth) ] ) - self.final_layer = T2IFinalLayer(hidden_size, patch_size, self.out_channels) + self.final_layer = T2IFinalLayer( + hidden_size, patch_size, self.out_channels, dtype=dtype, device=device, operations=operations + ) - def forward(self, x, timestep, y, mask=None, data_info=None, **kwargs): + def forward(self, x, timestep, context, mask=None, data_info=None, **kwargs): """ Forward pass of Sana. x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images) t: (N,) tensor of diffusion timesteps y: (N, 1, 120, C) tensor of class labels """ - x = x.to(self.dtype) - timestep = timestep.to(self.dtype) - y = y.to(self.dtype) - pos_embed = self.pos_embed.to(self.dtype) + 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: x = self.x_embedder(x) + pos_embed # (N, T, D), where T = H * W / patch_size ** 2 else: x = self.x_embedder(x) - t = self.t_embedder(timestep.to(x.dtype)) # (N, D) + t = self.t_embedder(timestep, x.dtype) # (N, D) t0 = self.t_block(t) y = self.y_embedder(y, self.training) # (N, 1, L, D) if self.y_norm: @@ -292,39 +325,6 @@ class Sana(nn.Module): imgs = x.reshape(shape=(x.shape[0], c, h * p, h * p)) return imgs - def initialize_weights(self): - # Initialize transformer layers: - def _basic_init(module): - if isinstance(module, nn.Linear): - torch.nn.init.xavier_uniform_(module.weight) - if module.bias is not None: - nn.init.constant_(module.bias, 0) - - self.apply(_basic_init) - - if self.use_pe: - # Initialize (and freeze) pos_embed by sin-cos embedding: - pos_embed = get_2d_sincos_pos_embed( - self.pos_embed.shape[-1], - int(self.x_embedder.num_patches**0.5), - pe_interpolation=self.pe_interpolation, - base_size=self.base_size, - ) - self.pos_embed.data.copy_(torch.from_numpy(pos_embed).float().unsqueeze(0)) - - # Initialize patch_embed like nn.Linear (instead of nn.Conv2d): - w = self.x_embedder.proj.weight.data - nn.init.xavier_uniform_(w.view([w.shape[0], -1])) - - # Initialize timestep embedding MLP: - nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02) - nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02) - nn.init.normal_(self.t_block[1].weight, std=0.02) - - # Initialize caption embedding MLP: - nn.init.normal_(self.y_embedder.y_proj.fc1.weight, std=0.02) - nn.init.normal_(self.y_embedder.y_proj.fc2.weight, std=0.02) - def get_2d_sincos_pos_embed(embed_dim, grid_size, cls_token=False, extra_tokens=0, pe_interpolation=1.0, base_size=16): """ diff --git a/Sana/models/sana_blocks.py b/Sana/models/sana_blocks.py index cdd37e4..d6350ad 100644 --- a/Sana/models/sana_blocks.py +++ b/Sana/models/sana_blocks.py @@ -22,12 +22,11 @@ import torch import torch.nn as nn import torch.nn.functional as F from einops import rearrange -from timm.models.vision_transformer import Attention as Attention_ -from timm.models.vision_transformer import Mlp from transformers import AutoModelForCausalLM from .norms import RMSNorm from .utils import get_same_padding, to_2tuple +from .basic_modules import Mlp sdpa_32b = None Q_4GB_LIMIT = 32000000 @@ -42,7 +41,7 @@ if model_management.xformers_enabled(): import xformers.ops else: if model_management.xpu_available: - import intel_extension_for_pytorch as ipex + import intel_extension_for_pytorch as ipex # type: ignore import os if not torch.xpu.has_fp64_dtype() and not os.environ.get('IPEX_FORCE_ATTENTION_SLICE', None): from ...utils.IPEX.attention import scaled_dot_product_attention_32_bit @@ -61,7 +60,7 @@ def t2i_modulate(x, shift, scale): class MultiHeadCrossAttention(nn.Module): - def __init__(self, d_model, num_heads, attn_drop=0.0, proj_drop=0.0, qk_norm=False, **block_kwargs): + def __init__(self, d_model, num_heads, attn_drop=0.0, proj_drop=0.0, qk_norm=False, dtype=None, device=None, operations=None, **block_kwargs): super().__init__() assert d_model % num_heads == 0, "d_model must be divisible by num_heads" @@ -69,10 +68,10 @@ class MultiHeadCrossAttention(nn.Module): self.num_heads = num_heads self.head_dim = d_model // num_heads - self.q_linear = nn.Linear(d_model, d_model) - self.kv_linear = nn.Linear(d_model, d_model * 2) + self.q_linear = operations.Linear(d_model, d_model, dtype=dtype, device=device) + self.kv_linear = operations.Linear(d_model, d_model * 2, dtype=dtype, device=device) self.attn_drop = nn.Dropout(attn_drop) - self.proj = nn.Linear(d_model, d_model) + self.proj = operations.Linear(d_model, d_model, dtype=dtype, device=device) self.proj_drop = nn.Dropout(proj_drop) if qk_norm: # not used for now @@ -135,7 +134,7 @@ class MultiHeadCrossAttention(nn.Module): return x -class LiteLA(Attention_): +class LiteLA(torch.nn.Module): # from attention r"""Lightweight linear attention""" PAD_VAL = 1 @@ -151,9 +150,20 @@ class LiteLA(Attention_): use_bias=False, qk_norm=False, norm_eps=1e-5, + dtype=None, + device=None, + operations=None, ): + super().__init__() heads = heads or int(out_dim // dim * heads_ratio) - super().__init__(in_dim, num_heads=heads, qkv_bias=use_bias) + + # assert dim % heads == 0, 'dim should be divisible by num_heads' + self.num_heads = heads + self.head_dim = in_dim // heads + self.scale = self.head_dim ** -0.5 + + self.qkv = operations.Linear(in_dim, in_dim * 3, bias=use_bias, dtype=dtype, device=device) + self.proj = operations.Linear(in_dim, in_dim, dtype=dtype, device=device) self.in_dim = in_dim self.out_dim = out_dim @@ -346,7 +356,7 @@ class SelfAttnProcessorLiteLA: return out -class FlashAttention(Attention_): +class FlashAttention(torch.nn.Module): # from attention """Multi-head Flash Attention block with qk norm.""" def __init__( @@ -395,7 +405,7 @@ class FlashAttention(Attention_): attn_bias = torch.zeros([B * self.num_heads, q.shape[1], k.shape[1]], dtype=q.dtype, device=q.device) attn_bias.masked_fill_(mask.squeeze(1).repeat(self.num_heads, 1, 1) == 0, float("-inf")) - if _xformers_available: + if model_management.xformers_enabled(): x = xformers.ops.memory_efficient_attention(q, k, v, p=self.attn_drop.p, attn_bias=attn_bias) else: q, k, v = q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2) @@ -418,7 +428,31 @@ class FlashAttention(Attention_): ################################################################################# # AMP attention with fp32 softmax to fix loss NaN problem during training # ################################################################################# -class Attention(Attention_): +class Attention(torch.nn.Module): + def __init__( + self, + dim, + num_heads=8, + qkv_bias=True, + sampling='conv', + sr_ratio=1, + qk_norm=False, + dtype=None, + device=None, + operations=None, + **block_kwargs, + ): + super().__init__() + assert dim % num_heads == 0, 'dim should be divisible by num_heads' + self.num_heads = num_heads + self.head_dim = dim // num_heads + self.scale = self.head_dim ** -0.5 + + self.qkv = operations.Linear(dim, dim * 3, bias=qkv_bias, dtype=dtype, device=device) + self.q_norm = nn.Identity() + self.k_norm = nn.Identity() + self.proj = operations.Linear(dim, dim, dtype=dtype, device=device) + def forward(self, x, HW=None): B, N, C = x.shape qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4) @@ -432,11 +466,11 @@ class Attention(Attention_): attn = (q @ k.transpose(-2, -1)) * self.scale attn = attn.softmax(dim=-1) - attn = self.attn_drop(attn) + #attn = self.attn_drop(attn) x = (attn @ v).transpose(1, 2).reshape(B, N, C) x = self.proj(x) - x = self.proj_drop(x) + #x = self.proj_drop(x) return x @@ -445,11 +479,14 @@ class FinalLayer(nn.Module): The final layer of Sana. """ - def __init__(self, hidden_size, patch_size, out_channels): + def __init__(self, hidden_size, patch_size, out_channels, dtype=None, device=None, operations=None): super().__init__() - self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True) - self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True)) + self.norm_final = operations.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device) + self.linear = operations.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True, dtype=dtype, device=device) + self.adaLN_modulation = nn.Sequential( + nn.SiLU(), + operations.Linear(hidden_size, 2 * hidden_size, bias=True, dtype=dtype, device=device) + ) def forward(self, x, c): shift, scale = self.adaLN_modulation(c).chunk(2, dim=1) @@ -463,17 +500,18 @@ class T2IFinalLayer(nn.Module): The final layer of Sana. """ - def __init__(self, hidden_size, patch_size, out_channels): + def __init__(self, hidden_size, patch_size, out_channels, dtype=None, device=None, operations=None): super().__init__() - self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True) - self.scale_shift_table = nn.Parameter(torch.randn(2, hidden_size) / hidden_size**0.5) + self.norm_final = operations.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device) + self.linear = operations.Linear(hidden_size, patch_size * patch_size * out_channels, bias=True, dtype=dtype, device=device) + self.scale_shift_table = nn.Parameter(torch.randn(2, hidden_size) / hidden_size ** 0.5) self.out_channels = out_channels def forward(self, x, t): + dtype = x.dtype shift, scale = (self.scale_shift_table[None] + t[:, None]).chunk(2, dim=1) x = t2i_modulate(self.norm_final(x), shift, scale) - x = self.linear(x) + x = self.linear(x.to(dtype)) return x @@ -482,12 +520,14 @@ class MaskFinalLayer(nn.Module): The final layer of Sana. """ - def __init__(self, final_hidden_size, c_emb_size, patch_size, out_channels): + def __init__(self, final_hidden_size, c_emb_size, patch_size, out_channels, dtype=None, device=None, operations=None): super().__init__() - self.norm_final = nn.LayerNorm(final_hidden_size, elementwise_affine=False, eps=1e-6) - self.linear = nn.Linear(final_hidden_size, patch_size * patch_size * out_channels, bias=True) - self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(c_emb_size, 2 * final_hidden_size, bias=True)) - + self.norm_final = operations.LayerNorm(final_hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device) + self.linear = operations.Linear(final_hidden_size, patch_size * patch_size * out_channels, bias=True, dtype=dtype, device=device) + self.adaLN_modulation = nn.Sequential( + nn.SiLU(), + operations.Linear(c_emb_size, 2 * final_hidden_size, bias=True, dtype=dtype, device=device) + ) def forward(self, x, t): shift, scale = self.adaLN_modulation(t).chunk(2, dim=1) x = modulate(self.norm_final(x), shift, scale) @@ -500,12 +540,14 @@ class DecoderLayer(nn.Module): The final layer of Sana. """ - def __init__(self, hidden_size, decoder_hidden_size): + def __init__(self, hidden_size, decoder_hidden_size, dtype=None, device=None, operations=None): super().__init__() - self.norm_decoder = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - self.linear = nn.Linear(hidden_size, decoder_hidden_size, bias=True) - self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True)) - + self.norm_decoder = operations.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device) + self.linear = operations.Linear(hidden_size, decoder_hidden_size, bias=True, dtype=dtype, device=device) + self.adaLN_modulation = nn.Sequential( + nn.SiLU(), + operations.Linear(hidden_size, 2 * hidden_size, bias=True, dtype=dtype, device=device) + ) def forward(self, x, t): shift, scale = self.adaLN_modulation(t).chunk(2, dim=1) x = modulate(self.norm_decoder(x), shift, scale) @@ -521,12 +563,12 @@ class TimestepEmbedder(nn.Module): Embeds scalar timesteps into vector representations. """ - def __init__(self, hidden_size, frequency_embedding_size=256): + def __init__(self, hidden_size, frequency_embedding_size=256, dtype=None, device=None, operations=None): super().__init__() self.mlp = nn.Sequential( - nn.Linear(frequency_embedding_size, hidden_size, bias=True), + operations.Linear(frequency_embedding_size, hidden_size, bias=True, dtype=dtype, device=device), nn.SiLU(), - nn.Linear(hidden_size, hidden_size, bias=True), + operations.Linear(hidden_size, hidden_size, bias=True, dtype=dtype, device=device), ) self.frequency_embedding_size = frequency_embedding_size @@ -551,9 +593,9 @@ class TimestepEmbedder(nn.Module): embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1) return embedding - def forward(self, t): - t_freq = self.timestep_embedding(t, self.frequency_embedding_size).to(self.dtype) - t_emb = self.mlp(t_freq) + def forward(self, t, dtype): + t_freq = self.timestep_embedding(t, self.frequency_embedding_size) + t_emb = self.mlp(t_freq.to(dtype)) return t_emb @property @@ -644,10 +686,14 @@ class CaptionEmbedder(nn.Module): uncond_prob, act_layer=nn.GELU(approximate="tanh"), token_num=120, + dtype=None, + device=None, + operations=None, ): super().__init__() self.y_proj = Mlp( - in_features=in_channels, hidden_features=hidden_size, out_features=hidden_size, act_layer=act_layer, drop=0 + in_features=in_channels, hidden_features=hidden_size, out_features=hidden_size, act_layer=act_layer, drop=0, + dtype=dtype, device=device, operations=operations ) self.register_buffer("y_embedding", nn.Parameter(torch.randn(token_num, in_channels) / in_channels**0.5)) self.uncond_prob = uncond_prob @@ -734,6 +780,9 @@ class PatchEmbed(nn.Module): norm_layer=None, flatten=True, bias=True, + dtype=None, + device=None, + operations=None, ): super().__init__() kernel_size = kernel_size or patch_size @@ -746,8 +795,8 @@ class PatchEmbed(nn.Module): self.flatten = flatten if not padding and kernel_size % 2 > 0: padding = get_same_padding(kernel_size) - self.proj = nn.Conv2d( - in_chans, embed_dim, kernel_size=kernel_size, stride=patch_size, padding=padding, bias=bias + self.proj = operations.Conv2d( + in_chans, embed_dim, kernel_size=kernel_size, stride=patch_size, padding=padding, bias=bias, dtype=dtype, device=device ) self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity() @@ -775,6 +824,9 @@ class PatchEmbedMS(nn.Module): norm_layer=None, flatten=True, bias=True, + dtype=None, + device=None, + operations=None, ): super().__init__() kernel_size = kernel_size or patch_size @@ -783,8 +835,8 @@ class PatchEmbedMS(nn.Module): self.flatten = flatten if not padding and kernel_size % 2 > 0: padding = get_same_padding(kernel_size) - self.proj = nn.Conv2d( - in_chans, embed_dim, kernel_size=kernel_size, stride=patch_size, padding=padding, bias=bias + self.proj = operations.Conv2d( + in_chans, embed_dim, kernel_size=kernel_size, stride=patch_size, padding=padding, bias=bias, dtype=dtype, device=device ) self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity() diff --git a/Sana/models/sana_multi_scale.py b/Sana/models/sana_multi_scale.py index 2d5452c..e69e183 100644 --- a/Sana/models/sana_multi_scale.py +++ b/Sana/models/sana_multi_scale.py @@ -17,7 +17,6 @@ # This file is modified from https://github.com/PixArt-alpha/PixArt-sigma import torch import torch.nn as nn -from timm.models.layers import DropPath from .basic_modules import DWMlp, GLUMBConv, MBConvPreGLU, Mlp from .sana import Sana, get_2d_sincos_pos_embed @@ -52,11 +51,14 @@ class SanaMSBlock(nn.Module): mlp_acts=("silu", "silu", None), linear_head_dim=32, cross_norm=False, + dtype=None, + device=None, + operations=None, **block_kwargs, ): super().__init__() self.hidden_size = hidden_size - self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + self.norm1 = operations.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device) if attn_type == "flash": # flash self attention self.attn = FlashAttention( @@ -64,25 +66,34 @@ class SanaMSBlock(nn.Module): num_heads=num_heads, qkv_bias=True, qk_norm=qk_norm, + dtype=dtype, + device=device, + operations=operations, **block_kwargs, ) elif attn_type == "linear": # linear self attention # TODO: Here the num_heads set to 36 for tmp used self_num_heads = hidden_size // linear_head_dim - self.attn = LiteLA(hidden_size, hidden_size, heads=self_num_heads, eps=1e-8, qk_norm=qk_norm) + self.attn = LiteLA( + hidden_size, hidden_size, heads=self_num_heads, eps=1e-8, qk_norm=qk_norm, + dtype=dtype, device=device, operations=operations, + ) elif attn_type == "vanilla": # vanilla self attention - self.attn = Attention(hidden_size, num_heads=num_heads, qkv_bias=True) + self.attn = Attention( + hidden_size, num_heads=num_heads, qkv_bias=True, dtype=dtype, device=device, operations=operations, + ) else: raise ValueError(f"{attn_type} type is not defined.") - self.cross_attn = MultiHeadCrossAttention(hidden_size, num_heads, qk_norm=cross_norm, **block_kwargs) + self.cross_attn = MultiHeadCrossAttention(hidden_size, num_heads, qk_norm=cross_norm, dtype=dtype, device=device, operations=operations, **block_kwargs) self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) if ffn_type == "dwmlp": approx_gelu = lambda: nn.GELU(approximate="tanh") self.mlp = DWMlp( - in_features=hidden_size, hidden_features=int(hidden_size * mlp_ratio), act_layer=approx_gelu, drop=0 + in_features=hidden_size, hidden_features=int(hidden_size * mlp_ratio), act_layer=approx_gelu, drop=0, + dtype=dtype, device=device, operations=operations, ) elif ffn_type == "glumbconv": self.mlp = GLUMBConv( @@ -91,6 +102,9 @@ class SanaMSBlock(nn.Module): use_bias=(True, True, False), norm=(None, None, None), act=mlp_acts, + dtype=dtype, + device=device, + operations=operations, ) elif ffn_type == "glumbconv_dilate": self.mlp = GLUMBConv( @@ -100,11 +114,15 @@ class SanaMSBlock(nn.Module): norm=(None, None, None), act=mlp_acts, dilation=2, + dtype=dtype, + device=device, + operations=operations, ) elif ffn_type == "mlp": approx_gelu = lambda: nn.GELU(approximate="tanh") self.mlp = Mlp( - in_features=hidden_size, hidden_features=int(hidden_size * mlp_ratio), act_layer=approx_gelu, drop=0 + in_features=hidden_size, hidden_features=int(hidden_size * mlp_ratio), act_layer=approx_gelu, drop=0, + dtype=dtype, device=device, operations=operations, ) elif ffn_type == "mbconvpreglu": self.mlp = MBConvPreGLU( @@ -114,17 +132,20 @@ class SanaMSBlock(nn.Module): use_bias=(True, True, False), norm=None, act=mlp_acts, + dtype=dtype, + device=device, + operations=operations, ) else: raise ValueError(f"{ffn_type} type is not defined.") - self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity() + self.drop_path = nn.Identity() # DropPath(drop_path) if drop_path > 0.0 else nn.Identity() self.scale_shift_table = nn.Parameter(torch.randn(6, hidden_size) / hidden_size**0.5) def forward(self, x, y, t, mask=None, HW=None, **kwargs): B, N, C = x.shape shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = ( - self.scale_shift_table[None] + t.reshape(B, 6, -1) + self.scale_shift_table[None].to(x.dtype) + t.reshape(B, 6, -1) ).chunk(6, dim=1) x = x + self.drop_path(gate_msa * self.attn(t2i_modulate(self.norm1(x), shift_msa, scale_msa), HW=HW)) x = x + self.cross_attn(x, y, mask) @@ -169,6 +190,9 @@ class SanaMS(Sana): mlp_acts=("silu", "silu", None), linear_head_dim=32, cross_norm=False, + dtype=None, + device=None, + operations=None, **kwargs, ): super().__init__( @@ -197,22 +221,33 @@ class SanaMS(Sana): patch_embed_kernel=patch_embed_kernel, mlp_acts=mlp_acts, linear_head_dim=linear_head_dim, + dtype=dtype, + device=device, + operations=operations, **kwargs, ) - self.dtype = torch.get_default_dtype() + self.dtype = dtype self.h = self.w = 0 approx_gelu = lambda: nn.GELU(approximate="tanh") - self.t_block = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 6 * hidden_size, bias=True)) + self.t_block = nn.Sequential( + nn.SiLU(), operations.Linear(hidden_size, 6 * hidden_size, bias=True, dtype=dtype, device=device) + ) self.pos_embed_ms = None kernel_size = patch_embed_kernel or patch_size - self.x_embedder = PatchEmbedMS(patch_size, in_channels, hidden_size, kernel_size=kernel_size, bias=True) + self.x_embedder = PatchEmbedMS( + patch_size, in_channels, hidden_size, kernel_size=kernel_size, bias=True, + dtype=dtype, device=device, operations=operations, + ) self.y_embedder = CaptionEmbedder( in_channels=caption_channels, hidden_size=hidden_size, uncond_prob=class_dropout_prob, act_layer=approx_gelu, token_num=model_max_length, + dtype=dtype, + device=device, + operations=operations, ) drop_path = [x.item() for x in torch.linspace(0, drop_path, depth)] # stochastic depth decay rule self.blocks = nn.ModuleList( @@ -229,13 +264,16 @@ class SanaMS(Sana): mlp_acts=mlp_acts, linear_head_dim=linear_head_dim, cross_norm=cross_norm, + dtype=dtype, + device=device, + operations=operations, ) for i in range(depth) ] ) - self.final_layer = T2IFinalLayer(hidden_size, patch_size, self.out_channels) - - self.initialize() + self.final_layer = T2IFinalLayer( + hidden_size, patch_size, self.out_channels, dtype=dtype, device=device, operations=operations + ) def forward(self, x, timesteps, context, **kwargs): """ @@ -251,7 +289,7 @@ class SanaMS(Sana): context = context.unsqueeze(1) ## run original forward pass - out = self.forward_raw( + out = self.forward_orig( x = x.to(self.dtype), timestep = timesteps.to(self.dtype), y = context.to(self.dtype), @@ -262,7 +300,7 @@ class SanaMS(Sana): return out - def forward_raw(self, x, timestep, y, mask=None, data_info=None, **kwargs): + def forward(self, x, timestep, context, mask=None, data_info=None, **kwargs): """ Forward pass of Sana. x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images) @@ -270,9 +308,7 @@ class SanaMS(Sana): y: (N, 1, 120, C) tensor of class labels """ bs = x.shape[0] - x = x.to(self.dtype) - timestep = timestep.to(self.dtype) - y = y.to(self.dtype) + y = context self.h, self.w = x.shape[-2] // self.patch_size, x.shape[-1] // self.patch_size if self.use_pe: x = self.x_embedder(x) @@ -285,16 +321,13 @@ class SanaMS(Sana): pe_interpolation=self.pe_interpolation, base_size=self.base_size, ) - ) - .unsqueeze(0) - .to(x.device) - .to(self.dtype) + ).unsqueeze(0).to(x.device).to(x.dtype) ) x += self.pos_embed_ms # (N, T, D), where T = H * W / patch_size ** 2 else: x = self.x_embedder(x) - t = self.t_embedder(timestep) # (N, D) + t = self.t_embedder(timestep, x.dtype) # (N, D) y_lens = ((y != 0).sum(dim=3) > 0).sum(dim=2).squeeze().tolist() y_lens = [y_lens[1]] * bs @@ -311,9 +344,7 @@ class SanaMS(Sana): y = y.squeeze(1).masked_select(mask.unsqueeze(-1).bool()).view(1, -1, y.shape[-1]) for block in self.blocks: - x = auto_grad_checkpoint( - block, x, y, t0, y_lens, (self.h, self.w), **kwargs - ) # (N, T, D) #support grad checkpoint + x = block(x, y, t0, y_lens, (self.h, self.w), **kwargs) # (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) @@ -348,26 +379,3 @@ class SanaMS(Sana): x = torch.einsum("nhwpqc->nchpwq", x) imgs = x.reshape(shape=(x.shape[0], c, self.h * p, self.w * p)) return imgs - - def initialize(self): - # Initialize transformer layers: - def _basic_init(module): - if isinstance(module, nn.Linear): - torch.nn.init.xavier_uniform_(module.weight) - if module.bias is not None: - nn.init.constant_(module.bias, 0) - - self.apply(_basic_init) - - # Initialize patch_embed like nn.Linear (instead of nn.Conv2d): - w = self.x_embedder.proj.weight.data - nn.init.xavier_uniform_(w.view([w.shape[0], -1])) - - # Initialize timestep embedding MLP: - nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02) - nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02) - nn.init.normal_(self.t_block[1].weight, std=0.02) - - # Initialize caption embedding MLP: - nn.init.normal_(self.y_embedder.y_proj.fc1.weight, std=0.02) - nn.init.normal_(self.y_embedder.y_proj.fc2.weight, std=0.02)