Dynamatrix/DiffBIR-OpenXLab
0
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 