MCPcopy Create free account
hub / github.com/294coder/Dif-PAN / forward

Method forward

models/sr3_dwt.py:442–470  ·  view source on GitHub ↗
(self, x, cond)

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 1

attn_opMethod · 0.95

Tested by

no test coverage detected