Use native ops for sana
This commit is contained in:
+3
-8
@@ -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
|
||||
|
||||
+2
-2
@@ -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 = {}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
+50
-50
@@ -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):
|
||||
"""
|
||||
|
||||
+95
-43
@@ -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.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()
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user