ControlNet/marlin_vit_base_ytf
1112
1import math2import warnings3from typing import Union, Optional, Callable, Tuple, List, Sequence4 5import torch6from einops.layers.torch import Rearrange7from torch import Tensor, nn, Size8from torch.nn import Conv3d, ModuleList9from torch.nn import functional as F10 11Shape = Union[Size, List[int], Tuple[int, ...]]12ModuleFactory = Union[Callable[[], nn.Module], Callable[[int], nn.Module]]13 14 15class PatchEmbedding3d(nn.Module):16 17 def __init__(self, input_size: Shape, patch_size: Union[int, Shape], embedding: int,18 strides: Optional[Union[int, Shape]] = None,19 build_normalization: Optional[ModuleFactory] = None20 ):21 super().__init__()22 # channel, time, height, width23 c, t, h, w = input_size24 # patch_time, patch_height, patch_width25 pt, ph, pw = (patch_size, patch_size, patch_size) if type(patch_size) is int else patch_size26 27 # configure the strides for conv3d28 if strides is None:29 # no specified means no overlap and gap between patches30 strides = (pt, ph, pw)31 elif type(strides) is int:32 # transform the side length of strides to 3D33 strides = (strides, strides, strides)34 35 self.projection = Conv3d(c, embedding, kernel_size=(pt, ph, pw), stride=strides)36 self.has_norm = build_normalization is not None37 if self.has_norm:38 self.normalization = build_normalization()39 self.rearrange = Rearrange("b d nt nh nw -> b (nt nh nw) d")40 41 def forward(self, x: Tensor) -> Tensor:42 x = self.projection(x)43 x = self.rearrange(x)44 if self.has_norm:45 x = self.normalization(x)46 return x47 48 49class Linear(nn.Module):50 51 def __init__(self, in_features: int, out_features: int, bias: bool = True,52 build_activation: Optional[ModuleFactory] = None,53 build_normalization: Optional[ModuleFactory] = None,54 normalization_after_activation: bool = False,55 dropout_rate: float = 0.56 ):57 super().__init__()58 self.linear = nn.Linear(in_features, out_features, bias)59 60 self.has_act = build_activation is not None61 if self.has_act:62 self.activation = build_activation()63 else:64 self.activation = None65 66 self.has_norm = build_normalization is not None67 if self.has_norm:68 self.normalization = build_normalization()69 self.norm_after_act = normalization_after_activation70 else:71 self.normalization = None72 73 self.has_dropout = dropout_rate > 074 if self.has_dropout:75 self.dropout = nn.Dropout(dropout_rate)76 77 def forward(self, x: Tensor) -> Tensor:78 x = self.linear(x)79 if self.has_act and self.has_norm:80 if self.norm_after_act:81 x = self.activation(x)82 x = self.normalization(x)83 else:84 x = self.normalization(x)85 x = self.activation(x)86 elif self.has_act and not self.has_norm:87 x = self.activation(x)88 elif not self.has_act and self.has_norm:89 x = self.normalization(x)90 91 if self.has_dropout:92 x = self.dropout(x)93 return x94 95 96class MLP(nn.Module):97 98 def __init__(self, neurons: Sequence[int],99 build_activation: Optional[ModuleFactory] = None, dropout_rate: float = 0.100 ):101 super().__init__()102 n_features = neurons[1:]103 self.layers: ModuleList[Linear] = ModuleList(104 [Linear(neurons[i], neurons[i + 1], True, build_activation, None,105 False, dropout_rate106 ) for i in range(len(n_features) - 1)107 ] + [108 Linear(neurons[-2], neurons[-1], True)109 ]110 )111 112 def forward(self, x: Tensor) -> Tensor:113 for layer in self.layers:114 x = layer(x)115 return x116 117 118class Attention(nn.Module):119 120 def __init__(121 self, dim, num_heads=8, qkv_bias=False, qk_scale=None, attn_drop=0.,122 proj_drop=0., attn_head_dim=None123 ):124 super().__init__()125 self.num_heads = num_heads126 head_dim = dim // num_heads127 if attn_head_dim is not None:128 head_dim = attn_head_dim129 all_head_dim = head_dim * self.num_heads130 self.scale = qk_scale or head_dim ** -0.5131 132 self.qkv = nn.Linear(dim, all_head_dim * 3, bias=False)133 if qkv_bias:134 self.q_bias = nn.Parameter(torch.zeros(all_head_dim))135 self.v_bias = nn.Parameter(torch.zeros(all_head_dim))136 else:137 self.q_bias = None138 self.v_bias = None139 140 self.attn_drop = nn.Dropout(attn_drop)141 self.proj = nn.Linear(all_head_dim, dim)142 self.proj_drop = nn.Dropout(proj_drop)143 144 def forward(self, x):145 B, N, C = x.shape146 qkv_bias = None147 if self.q_bias is not None:148 qkv_bias = torch.cat((self.q_bias, torch.zeros_like(self.v_bias, requires_grad=False), self.v_bias))149 # qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)150 qkv = F.linear(input=x, weight=self.qkv.weight, bias=qkv_bias)151 qkv = qkv.reshape(B, N, 3, self.num_heads, -1).permute(2, 0, 3, 1, 4)152 q, k, v = qkv[0], qkv[1], qkv[2] # make torchscript happy (cannot use tensor as tuple)153 154 q = q * self.scale155 attn = (q @ k.transpose(-2, -1))156 157 attn = attn.softmax(dim=-1)158 attn = self.attn_drop(attn)159 160 x = (attn @ v).transpose(1, 2).reshape(B, N, -1)161 x = self.proj(x)162 x = self.proj_drop(x)163 return x164 165 166class Block(nn.Module):167 168 def __init__(self, dim, num_heads, mlp_ratio=4., qkv_bias=False, qk_scale=None, drop=0., attn_drop=0.,169 init_values=None, act_layer=nn.GELU, norm_layer=nn.LayerNorm,170 attn_head_dim=None171 ):172 super().__init__()173 self.norm1 = norm_layer(dim)174 self.attn = Attention(175 dim, num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale,176 attn_drop=attn_drop, proj_drop=drop, attn_head_dim=attn_head_dim)177 self.norm2 = norm_layer(dim)178 mlp_hidden_dim = int(dim * mlp_ratio)179 self.mlp = MLP(180 neurons=[dim, mlp_hidden_dim, dim],181 build_activation=act_layer,182 dropout_rate=drop183 )184 185 if init_values > 0:186 self.gamma_1 = nn.Parameter(init_values * torch.ones((dim)), requires_grad=True)187 self.gamma_2 = nn.Parameter(init_values * torch.ones((dim)), requires_grad=True)188 else:189 self.gamma_1, self.gamma_2 = None, None190 191 def forward(self, x):192 if self.gamma_1 is None:193 x = x + self.attn(self.norm1(x))194 x = x + self.mlp(self.norm2(x))195 else:196 x = x + (self.gamma_1 * self.attn(self.norm1(x)))197 x = x + (self.gamma_2 * self.mlp(self.norm2(x)))198 return x199 200 201def no_grad_trunc_normal_(tensor, mean, std, a, b):202 # Cut & paste from PyTorch official master until it's in a few official releases - RW203 # Method based on https://people.sc.fsu.edu/~jburkardt/presentations/truncated_normal.pdf204 def norm_cdf(x):205 # Computes standard normal cumulative distribution function206 return (1. + math.erf(x / math.sqrt(2.))) / 2.207 208 if (mean < a - 2 * std) or (mean > b + 2 * std):209 warnings.warn("mean is more than 2 std from [a, b] in nn.init.trunc_normal_. "210 "The distribution of values may be incorrect.",211 stacklevel=2)212 213 with torch.no_grad():214 # Values are generated by using a truncated uniform distribution and215 # then using the inverse CDF for the normal distribution.216 # Get upper and lower cdf values217 l = norm_cdf((a - mean) / std)218 u = norm_cdf((b - mean) / std)219 220 # Uniformly fill tensor with values from [l, u], then translate to221 # [2l-1, 2u-1].222 tensor.uniform_(2 * l - 1, 2 * u - 1)223 224 # Use inverse cdf transform for normal distribution to get truncated225 # standard normal226 tensor.erfinv_()227 228 # Transform to proper mean, std229 tensor.mul_(std * math.sqrt(2.))230 tensor.add_(mean)231 232 # Clamp to ensure it's in the proper range233 tensor.clamp_(min=a, max=b)234 return tensor235 