faster model loading

This commit is contained in:
kijai
2024-03-22 10:54:46 +02:00
parent 0ce54ffcaf
commit 062e9b7a67
2 changed files with 36 additions and 22 deletions
+26 -19
View File
@@ -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):
+10 -3
View File
@@ -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):