Use native ops for sana

This commit is contained in:
City
2024-12-11 22:01:11 +01:00
parent 5505ab4f40
commit 7e19ac5e43
7 changed files with 274 additions and 199 deletions
+61 -41
View File
@@ -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)