| 88 | return x |
| 89 | |
| 90 | class Attention(nn.Module): |
| 91 | def __init__(self, query_dim, context_dim=None, |
| 92 | num_heads=8, dim_head=48, qkv_bias=False, flash=False): |
| 93 | super().__init__() |
| 94 | inner_dim = self.inner_dim = dim_head * num_heads |
| 95 | context_dim = default(context_dim, query_dim) |
| 96 | self.scale = dim_head**-0.5 |
| 97 | self.heads = num_heads |
| 98 | self.flash = flash |
| 99 | |
| 100 | self.to_q = nn.Linear(query_dim, inner_dim, bias=qkv_bias) |
| 101 | self.to_kv = nn.Linear(context_dim, inner_dim * 2, bias=qkv_bias) |
| 102 | self.to_out = nn.Linear(inner_dim, query_dim) |
| 103 | |
| 104 | def forward(self, x, context=None, attn_bias=None): |
| 105 | B, N1, _ = x.shape |
| 106 | C = self.inner_dim |
| 107 | h = self.heads |
| 108 | q = self.to_q(x).reshape(B, N1, h, C // h).permute(0, 2, 1, 3) |
| 109 | context = default(context, x) |
| 110 | k, v = self.to_kv(context).chunk(2, dim=-1) |
| 111 | |
| 112 | N2 = context.shape[1] |
| 113 | k = k.reshape(B, N2, h, C // h).permute(0, 2, 1, 3) |
| 114 | v = v.reshape(B, N2, h, C // h).permute(0, 2, 1, 3) |
| 115 | |
| 116 | with torch.autocast("cuda", enabled=True, dtype=torch.bfloat16): |
| 117 | if self.flash==False: |
| 118 | sim = (q @ k.transpose(-2, -1)) * self.scale |
| 119 | if attn_bias is not None: |
| 120 | sim = sim + attn_bias |
| 121 | if sim.abs().max()>1e2: |
| 122 | import pdb; pdb.set_trace() |
| 123 | attn = sim.softmax(dim=-1) |
| 124 | x = (attn @ v).transpose(1, 2).reshape(B, N1, C) |
| 125 | else: |
| 126 | input_args = [x.contiguous() for x in [q, k, v]] |
| 127 | x = F.scaled_dot_product_attention(*input_args).permute(0,2,1,3).reshape(B,N1,-1) # type: ignore |
| 128 | |
| 129 | if self.to_out.bias.dtype != x.dtype: |
| 130 | x = x.to(self.to_out.bias.dtype) |
| 131 | |
| 132 | return self.to_out(x) |
| 133 | |
| 134 | |
| 135 | class VGG19(nn.Module): |