| 385 | |
| 386 | |
| 387 | class FreqCondInjection(nn.Module): |
| 388 | def __init__( |
| 389 | self, |
| 390 | fea_dim, |
| 391 | cond_dim, |
| 392 | qkv_dim, |
| 393 | dim_out, |
| 394 | groups=32, |
| 395 | nheads=8, |
| 396 | drop_path_prob=0.2, |
| 397 | ) -> None: |
| 398 | super().__init__() |
| 399 | assert fea_dim % nheads == 0, "@dim must be divisible by @nheads" |
| 400 | |
| 401 | self.prenorm_x = nn.GroupNorm(groups, fea_dim) |
| 402 | # self.prenorm_cond = nn.GroupNorm(groups // 4, cond_dim) |
| 403 | |
| 404 | self.q = nn.Sequential( |
| 405 | nn.Conv2d(fea_dim, fea_dim, 3, 1, 1, bias=False, groups=fea_dim), |
| 406 | nn.Conv2d(fea_dim, qkv_dim, 1, bias=True), |
| 407 | ) |
| 408 | self.kv = nn.Sequential( |
| 409 | nn.Conv2d(cond_dim, cond_dim, 3, 1, 1, bias=False, groups=cond_dim), |
| 410 | nn.Conv2d(cond_dim, qkv_dim * 2, 1, bias=True), |
| 411 | ) |
| 412 | self.nheads = nheads |
| 413 | self.scale = 1 / math.sqrt(qkv_dim // nheads) |
| 414 | |
| 415 | self.attn_out = nn.Conv2d(qkv_dim, dim_out, 1, bias=True) |
| 416 | self.attn_res = ( |
| 417 | nn.Conv2d(fea_dim, dim_out, 1, bias=True) |
| 418 | if fea_dim != dim_out |
| 419 | else nn.Identity() |
| 420 | ) |
| 421 | |
| 422 | self.ffn = nn.Sequential( |
| 423 | nn.Conv2d(dim_out, dim_out * 2, 3, 1, 1, bias=False), |
| 424 | nn.SiLU(), |
| 425 | nn.Conv2d(dim_out * 2, dim_out, 3, 1, 1, bias=False), |
| 426 | nn.Conv2d(dim_out, dim_out, 1, bias=True), |
| 427 | ) |
| 428 | self.ffn_drop_path = DropPath(drop_prob=drop_path_prob) |
| 429 | |
| 430 | def forward(self, x, cond): |
| 431 | x = self.prenorm_x(x) |
| 432 | # cond = self.prenorm_cond(cond) |
| 433 | |
| 434 | q = self.q(x) |
| 435 | k, v = self.kv(cond).chunk(2, dim=1) |
| 436 | |
| 437 | q, k, v = map(lambda in_qkv: F.normalize(in_qkv, dim=1), (q, k, v)) |
| 438 | |
| 439 | # convert to freq space |
| 440 | q = torch.fft.rfft2(q, dim=(-2, -1), norm="ortho") # b, c, h, w/2+1 |
| 441 | k = torch.fft.rfft2(k, dim=(-2, -1), norm="ortho") # b, c, h, w/2+1 |
| 442 | v = torch.fft.rfft2(v, dim=(-2, -1), norm="ortho") # b, c, h, w/2+1 |
| 443 | |
| 444 | # amp and phas attention |