diff --git a/depthfm/unet/attention.py b/depthfm/unet/attention.py index 6a30ab1..b8e9f6b 100644 --- a/depthfm/unet/attention.py +++ b/depthfm/unet/attention.py @@ -8,6 +8,13 @@ from typing import Optional, Any from util import checkpoint +class Conv2d(torch.nn.Conv2d): + def reset_parameters(self): + return None + +class Linear(torch.nn.Linear): + def reset_parameters(self): + return None try: import xformers @@ -52,7 +59,7 @@ def init_(tensor): class GEGLU(nn.Module): def __init__(self, dim_in, dim_out): super().__init__() - self.proj = nn.Linear(dim_in, dim_out * 2) + self.proj = Linear(dim_in, dim_out * 2) def forward(self, x): x, gate = self.proj(x).chunk(2, dim=-1) @@ -65,13 +72,13 @@ class FeedForward(nn.Module): inner_dim = int(dim * mult) dim_out = default(dim_out, dim) project_in = ( - nn.Sequential(nn.Linear(dim, inner_dim), nn.GELU()) + nn.Sequential(Linear(dim, inner_dim), nn.GELU()) if not glu else GEGLU(dim, inner_dim) ) self.net = nn.Sequential( - project_in, nn.Dropout(dropout), nn.Linear(inner_dim, dim_out) + project_in, nn.Dropout(dropout), Linear(inner_dim, dim_out) ) def forward(self, x): @@ -99,16 +106,16 @@ class SpatialSelfAttention(nn.Module): self.in_channels = in_channels self.norm = Normalize(in_channels) - self.q = torch.nn.Conv2d( + self.q = torch.Conv2d( in_channels, in_channels, kernel_size=1, stride=1, padding=0 ) - self.k = torch.nn.Conv2d( + self.k = torch.Conv2d( in_channels, in_channels, kernel_size=1, stride=1, padding=0 ) - self.v = torch.nn.Conv2d( + self.v = torch.Conv2d( in_channels, in_channels, kernel_size=1, stride=1, padding=0 ) - self.proj_out = torch.nn.Conv2d( + self.proj_out = torch.Conv2d( in_channels, in_channels, kernel_size=1, stride=1, padding=0 ) @@ -149,12 +156,12 @@ class CrossAttention(nn.Module): self.scale = dim_head**-0.5 self.heads = heads - self.to_q = nn.Linear(query_dim, inner_dim, bias=False) - self.to_k = nn.Linear(context_dim, inner_dim, bias=False) - self.to_v = nn.Linear(context_dim, inner_dim, bias=False) + self.to_q = Linear(query_dim, inner_dim, bias=False) + self.to_k = Linear(context_dim, inner_dim, bias=False) + self.to_v = Linear(context_dim, inner_dim, bias=False) self.to_out = nn.Sequential( - nn.Linear(inner_dim, query_dim), nn.Dropout(dropout) + Linear(inner_dim, query_dim), nn.Dropout(dropout) ) def forward(self, x, context=None, mask=None, rescale_attention=True): @@ -195,12 +202,12 @@ class MemoryEfficientCrossAttention(nn.Module): self.heads = heads self.dim_head = dim_head - self.to_q = nn.Linear(query_dim, inner_dim, bias=False) - self.to_k = nn.Linear(context_dim, inner_dim, bias=False) - self.to_v = nn.Linear(context_dim, inner_dim, bias=False) + self.to_q = Linear(query_dim, inner_dim, bias=False) + self.to_k = Linear(context_dim, inner_dim, bias=False) + self.to_v = Linear(context_dim, inner_dim, bias=False) self.to_out = nn.Sequential( - nn.Linear(inner_dim, query_dim), nn.Dropout(dropout) + Linear(inner_dim, query_dim), nn.Dropout(dropout) ) self.attention_op: Optional[Any] = None @@ -324,11 +331,11 @@ class SpatialTransformer(nn.Module): inner_dim = n_heads * d_head self.norm = Normalize(in_channels) if not use_linear: - self.proj_in = nn.Conv2d( + self.proj_in = Conv2d( in_channels, inner_dim, kernel_size=1, stride=1, padding=0 ) else: - self.proj_in = nn.Linear(in_channels, inner_dim) + self.proj_in = Linear(in_channels, inner_dim) self.transformer_blocks = nn.ModuleList( [ @@ -346,10 +353,10 @@ class SpatialTransformer(nn.Module): ) if not use_linear: self.proj_out = zero_module( - nn.Conv2d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0) + Conv2d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0) ) else: - self.proj_out = zero_module(nn.Linear(in_channels, inner_dim)) + self.proj_out = zero_module(Linear(in_channels, inner_dim)) self.use_linear = use_linear def forward(self, x, context=None): diff --git a/depthfm/unet/util.py b/depthfm/unet/util.py index 56690f7..f57c31a 100644 --- a/depthfm/unet/util.py +++ b/depthfm/unet/util.py @@ -141,7 +141,14 @@ class GroupNorm32(nn.GroupNorm): def forward(self, x): return super().forward(x.float()).type(x.dtype) - +class Conv2d(torch.nn.Conv2d): + def reset_parameters(self): + return None + +class Linear(torch.nn.Linear): + def reset_parameters(self): + return None + def conv_nd(dims, *args, **kwargs): """ Create a 1D, 2D, or 3D convolution module. @@ -149,7 +156,7 @@ def conv_nd(dims, *args, **kwargs): if dims == 1: return nn.Conv1d(*args, **kwargs) elif dims == 2: - return nn.Conv2d(*args, **kwargs) + return Conv2d(*args, **kwargs) elif dims == 3: return nn.Conv3d(*args, **kwargs) raise ValueError(f"unsupported dimensions: {dims}") @@ -159,7 +166,7 @@ def linear(*args, **kwargs): """ Create a linear module. """ - return nn.Linear(*args, **kwargs) + return Linear(*args, **kwargs) def avg_pool_nd(dims, *args, **kwargs):