| 161 | ) |
| 162 | |
| 163 | def forward(self, x, context=None, mask=None): |
| 164 | h = self.heads |
| 165 | |
| 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 overflowing |
| 174 | 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.scale |
| 178 | else: |
| 179 | sim = einsum('b i d, b j d -> b i j', q, k) * self.scale |
| 180 | |
| 181 | del q, k |
| 182 | |
| 183 | if exists(mask): |
| 184 | mask = rearrange(mask, 'b ... -> b (...)') |
| 185 | max_neg_value = -torch.finfo(sim.dtype).max |
| 186 | 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 of |
| 190 | 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 | |
| 197 | class MemoryEfficientCrossAttention(nn.Module): |