Team Ai
Modelpublic

ControlNet/marlin_vit_base_ytf

sourceHugging Faceccupdated 2y agoView on Hugging Face
1likes112downloads
modules.py235 linesDownload Raw Back to root
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