Use native ops for sana
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user