| 459 | |
| 460 | |
| 461 | def attn_op(self, q, k, v): |
| 462 | b, c, xf, yf = q.shape |
| 463 | |
| 464 | q, k, v = map( |
| 465 | lambda in_x: rearrange( |
| 466 | in_x, "b (h c) xf yf -> b h c (xf yf)", h=self.nheads |
| 467 | ), |
| 468 | (q, k, v), |
| 469 | ) |
| 470 | # n x n attn map |
| 471 | sim = torch.einsum("b h c m, b h c n -> b h m n", q, k) * self.scale |
| 472 | sim = sim.softmax(-1) |
| 473 | # h w fused feature map |
| 474 | out = torch.einsum("b h m n, b h c n-> b h c m", sim, v) |
| 475 | out = rearrange( |
| 476 | out, "n h c (xf yf) -> n (h c) xf yf", xf=xf, yf=yf, h=self.nheads |
| 477 | ) |
| 478 | |
| 479 | return out |
| 480 | |
| 481 | |
| 482 | |