Team Ai
Apppublic

Dynamatrix/DiffBIR-OpenXLab

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
attention.py342 linesDownload Raw Back to modules
1from inspect import isfunction2import math3import torch4import torch.nn.functional as F5from torch import nn, einsum6from einops import rearrange, repeat7from typing import Optional, Any8 9from ldm.modules.diffusionmodules.util import checkpoint10 11 12try:13    import xformers14    import xformers.ops15    XFORMERS_IS_AVAILBLE = True16except:17    XFORMERS_IS_AVAILBLE = False18 19# CrossAttn precision handling20import os21_ATTN_PRECISION = os.environ.get("ATTN_PRECISION", "fp32")22 23def exists(val):24    return val is not None25 26 27def uniq(arr):28    return{el: True for el in arr}.keys()29 30 31def default(val, d):32    if exists(val):33        return val34    return d() if isfunction(d) else d35 36 37def max_neg_value(t):38    return -torch.finfo(t.dtype).max39 40 41def init_(tensor):42    dim = tensor.shape[-1]43    std = 1 / math.sqrt(dim)44    tensor.uniform_(-std, std)45    return tensor46 47 48# feedforward49class GEGLU(nn.Module):50    def __init__(self, dim_in, dim_out):51        super().__init__()52        self.proj = nn.Linear(dim_in, dim_out * 2)53 54    def forward(self, x):55        x, gate = self.proj(x).chunk(2, dim=-1)56        return x * F.gelu(gate)57 58 59class FeedForward(nn.Module):60    def __init__(self, dim, dim_out=None, mult=4, glu=False, dropout=0.):61        super().__init__()62        inner_dim = int(dim * mult)63        dim_out = default(dim_out, dim)64        project_in = nn.Sequential(65            nn.Linear(dim, inner_dim),66            nn.GELU()67        ) if not glu else GEGLU(dim, inner_dim)68 69        self.net = nn.Sequential(70            project_in,71            nn.Dropout(dropout),72            nn.Linear(inner_dim, dim_out)73        )74 75    def forward(self, x):76        return self.net(x)77 78 79def zero_module(module):80    """81    Zero out the parameters of a module and return it.82    """83    for p in module.parameters():84        p.detach().zero_()85    return module86 87 88def Normalize(in_channels):89    return torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True)90 91 92class SpatialSelfAttention(nn.Module):93    def __init__(self, in_channels):94        super().__init__()95        self.in_channels = in_channels96 97        self.norm = Normalize(in_channels)98        self.q = torch.nn.Conv2d(in_channels,99                                 in_channels,100                                 kernel_size=1,101                                 stride=1,102                                 padding=0)103        self.k = torch.nn.Conv2d(in_channels,104                                 in_channels,105                                 kernel_size=1,106                                 stride=1,107                                 padding=0)108        self.v = torch.nn.Conv2d(in_channels,109                                 in_channels,110                                 kernel_size=1,111                                 stride=1,112                                 padding=0)113        self.proj_out = torch.nn.Conv2d(in_channels,114                                        in_channels,115                                        kernel_size=1,116                                        stride=1,117                                        padding=0)118 119    def forward(self, x):120        h_ = x121        h_ = self.norm(h_)122        q = self.q(h_)123        k = self.k(h_)124        v = self.v(h_)125 126        # compute attention127        b,c,h,w = q.shape128        q = rearrange(q, 'b c h w -> b (h w) c')129        k = rearrange(k, 'b c h w -> b c (h w)')130        w_ = torch.einsum('bij,bjk->bik', q, k)131 132        w_ = w_ * (int(c)**(-0.5))133        w_ = torch.nn.functional.softmax(w_, dim=2)134 135        # attend to values136        v = rearrange(v, 'b c h w -> b c (h w)')137        w_ = rearrange(w_, 'b i j -> b j i')138        h_ = torch.einsum('bij,bjk->bik', v, w_)139        h_ = rearrange(h_, 'b c (h w) -> b c h w', h=h)140        h_ = self.proj_out(h_)141 142        return x+h_143 144 145class CrossAttention(nn.Module):146    def __init__(self, query_dim, context_dim=None, heads=8, dim_head=64, dropout=0.):147        super().__init__()148        inner_dim = dim_head * heads149        context_dim = default(context_dim, query_dim)150 151        self.scale = dim_head ** -0.5152        self.heads = heads153 154        self.to_q = nn.Linear(query_dim, inner_dim, bias=False)155        self.to_k = nn.Linear(context_dim, inner_dim, bias=False)156        self.to_v = nn.Linear(context_dim, inner_dim, bias=False)157 158        self.to_out = nn.Sequential(159            nn.Linear(inner_dim, query_dim),160            nn.Dropout(dropout)161        )162 163    def forward(self, x, context=None, mask=None):164        h = self.heads165 166        q = self.to_q(x)167        context = default(context, x)168        k = self.to_k(context)169        v = self.to_v(context)170 171        q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> (b h) n d', h=h), (q, k, v))172 173        # force cast to fp32 to avoid overflowing174        if _ATTN_PRECISION =="fp32":175            with torch.autocast(enabled=False, device_type = 'cuda'):176                q, k = q.float(), k.float()177                sim = einsum('b i d, b j d -> b i j', q, k) * self.scale178        else:179            sim = einsum('b i d, b j d -> b i j', q, k) * self.scale180        181        del q, k182    183        if exists(mask):184            mask = rearrange(mask, 'b ... -> b (...)')185            max_neg_value = -torch.finfo(sim.dtype).max186            mask = repeat(mask, 'b j -> (b h) () j', h=h)187            sim.masked_fill_(~mask, max_neg_value)188 189        # attention, what we cannot get enough of190        sim = sim.softmax(dim=-1)191 192        out = einsum('b i j, b j d -> b i d', sim, v)193        out = rearrange(out, '(b h) n d -> b n (h d)', h=h)194        return self.to_out(out)195 196 197class MemoryEfficientCrossAttention(nn.Module):198    # https://github.com/MatthieuTPHR/diffusers/blob/d80b531ff8060ec1ea982b65a1b8df70f73aa67c/src/diffusers/models/attention.py#L223199    def __init__(self, query_dim, context_dim=None, heads=8, dim_head=64, dropout=0.0):200        super().__init__()201        print(f"Setting up {self.__class__.__name__}. Query dim is {query_dim}, context_dim is {context_dim} and using "202              f"{heads} heads.")203        inner_dim = dim_head * heads204        context_dim = default(context_dim, query_dim)205 206        self.heads = heads207        self.dim_head = dim_head208 209        self.to_q = nn.Linear(query_dim, inner_dim, bias=False)210        self.to_k = nn.Linear(context_dim, inner_dim, bias=False)211        self.to_v = nn.Linear(context_dim, inner_dim, bias=False)212 213        self.to_out = nn.Sequential(nn.Linear(inner_dim, query_dim), nn.Dropout(dropout))214        self.attention_op: Optional[Any] = None215 216    def forward(self, x, context=None, mask=None):217        q = self.to_q(x)218        context = default(context, x)219        k = self.to_k(context)220        v = self.to_v(context)221 222        b, _, _ = q.shape223        q, k, v = map(224            lambda t: t.unsqueeze(3)225            .reshape(b, t.shape[1], self.heads, self.dim_head)226            .permute(0, 2, 1, 3)227            .reshape(b * self.heads, t.shape[1], self.dim_head)228            .contiguous(),229            (q, k, v),230        )231 232        # actually compute the attention, what we cannot get enough of233        out = xformers.ops.memory_efficient_attention(q, k, v, attn_bias=None, op=self.attention_op)234 235        if exists(mask):236            raise NotImplementedError237        out = (238            out.unsqueeze(0)239            .reshape(b, self.heads, out.shape[1], self.dim_head)240            .permute(0, 2, 1, 3)241            .reshape(b, out.shape[1], self.heads * self.dim_head)242        )243        return self.to_out(out)244 245 246class BasicTransformerBlock(nn.Module):247    ATTENTION_MODES = {248        "softmax": CrossAttention,  # vanilla attention249        "softmax-xformers": MemoryEfficientCrossAttention250    }251    def __init__(self, dim, n_heads, d_head, dropout=0., context_dim=None, gated_ff=True, checkpoint=True,252                 disable_self_attn=False):253        super().__init__()254        attn_mode = "softmax-xformers" if XFORMERS_IS_AVAILBLE else "softmax"255        assert attn_mode in self.ATTENTION_MODES256        attn_cls = self.ATTENTION_MODES[attn_mode]257        self.disable_self_attn = disable_self_attn258        self.attn1 = attn_cls(query_dim=dim, heads=n_heads, dim_head=d_head, dropout=dropout,259                              context_dim=context_dim if self.disable_self_attn else None)  # is a self-attention if not self.disable_self_attn260        self.ff = FeedForward(dim, dropout=dropout, glu=gated_ff)261        self.attn2 = attn_cls(query_dim=dim, context_dim=context_dim,262                              heads=n_heads, dim_head=d_head, dropout=dropout)  # is self-attn if context is none263        self.norm1 = nn.LayerNorm(dim)264        self.norm2 = nn.LayerNorm(dim)265        self.norm3 = nn.LayerNorm(dim)266        self.checkpoint = checkpoint267 268    def forward(self, x, context=None):269        return checkpoint(self._forward, (x, context), self.parameters(), self.checkpoint)270 271    def _forward(self, x, context=None):272        x = self.attn1(self.norm1(x), context=context if self.disable_self_attn else None) + x273        x = self.attn2(self.norm2(x), context=context) + x274        x = self.ff(self.norm3(x)) + x275        return x276 277 278class SpatialTransformer(nn.Module):279    """280    Transformer block for image-like data.281    First, project the input (aka embedding)282    and reshape to b, t, d.283    Then apply standard transformer action.284    Finally, reshape to image285    NEW: use_linear for more efficiency instead of the 1x1 convs286    """287    def __init__(self, in_channels, n_heads, d_head,288                 depth=1, dropout=0., context_dim=None,289                 disable_self_attn=False, use_linear=False,290                 use_checkpoint=True):291        super().__init__()292        if exists(context_dim) and not isinstance(context_dim, list):293            context_dim = [context_dim]294        self.in_channels = in_channels295        inner_dim = n_heads * d_head296        self.norm = Normalize(in_channels)297        if not use_linear:298            self.proj_in = nn.Conv2d(in_channels,299                                     inner_dim,300                                     kernel_size=1,301                                     stride=1,302                                     padding=0)303        else:304            self.proj_in = nn.Linear(in_channels, inner_dim)305 306        self.transformer_blocks = nn.ModuleList(307            [BasicTransformerBlock(inner_dim, n_heads, d_head, dropout=dropout, context_dim=context_dim[d],308                                   disable_self_attn=disable_self_attn, checkpoint=use_checkpoint)309                for d in range(depth)]310        )311        if not use_linear:312            self.proj_out = zero_module(nn.Conv2d(inner_dim,313                                                  in_channels,314                                                  kernel_size=1,315                                                  stride=1,316                                                  padding=0))317        else:318            self.proj_out = zero_module(nn.Linear(in_channels, inner_dim))319        self.use_linear = use_linear320 321    def forward(self, x, context=None):322        # note: if no context is given, cross-attention defaults to self-attention323        if not isinstance(context, list):324            context = [context]325        b, c, h, w = x.shape326        x_in = x327        x = self.norm(x)328        if not self.use_linear:329            x = self.proj_in(x)330        x = rearrange(x, 'b c h w -> b (h w) c').contiguous()331        if self.use_linear:332            x = self.proj_in(x)333        for i, block in enumerate(self.transformer_blocks):334            x = block(x, context=context[i])335        if self.use_linear:336            x = self.proj_out(x)337        x = rearrange(x, 'b (h w) c -> b c h w', h=h, w=w).contiguous()338        if not self.use_linear:339            x = self.proj_out(x)340        return x + x_in341 342