| 440 | self.ffn_drop_path = DropPath(drop_prob=drop_path_prob) |
| 441 | |
| 442 | def forward(self, x, cond): |
| 443 | x = self.prenorm_x(x) |
| 444 | # cond = self.prenorm_cond(cond) |
| 445 | |
| 446 | q = self.q(x) |
| 447 | k, v = self.kv(cond).chunk(2, dim=1) |
| 448 | |
| 449 | q, k, v = map(lambda in_qkv: F.normalize(in_qkv, dim=1), (q, k, v)) |
| 450 | |
| 451 | # convert to freq space |
| 452 | q = torch.fft.rfft2(q, dim=(-2, -1), norm="ortho") # b, c, h, w/2+1 |
| 453 | k = torch.fft.rfft2(k, dim=(-2, -1), norm="ortho") # b, c, h, w/2+1 |
| 454 | v = torch.fft.rfft2(v, dim=(-2, -1), norm="ortho") # b, c, h, w/2+1 |
| 455 | |
| 456 | # amp and phas attention |
| 457 | amp_out = self.attn_op(q.abs(), k.abs(), v.abs()) |
| 458 | phas_out = self.attn_op(q.angle(), k.angle(), v.angle()) |
| 459 | |
| 460 | # convert to complex |
| 461 | out = torch.polar(amp_out, phas_out) |
| 462 | |
| 463 | # convert to rgb space |
| 464 | out = torch.fft.irfft2(out, dim=(-2, -1), norm="ortho") |
| 465 | |
| 466 | attn_out = self.attn_out(out) + self.attn_res(x) |
| 467 | |
| 468 | # ffn |
| 469 | ffn_out = self.ffn_drop_path(self.ffn(attn_out)) + attn_out |
| 470 | return ffn_out |
| 471 | |
| 472 | def attn_op(self, q, k, v): |
| 473 | b, c, xf, yf = q.shape |