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

Method forward

models/sr3.py:430–458  ·  view source on GitHub ↗
(self, x, cond)

Source from the content-addressed store, hash-verified

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
445 amp_out = self.attn_op(q.abs(), k.abs(), v.abs())
446 phas_out = self.attn_op(q.angle(), k.angle(), v.angle())
447
448 # convert to complex
449 out = torch.polar(amp_out, phas_out)
450
451 # convert to rgb space
452 out = torch.fft.irfft2(out, dim=(-2, -1), norm="ortho")
453
454 attn_out = self.attn_out(out) + self.attn_res(x)
455
456 # ffn
457 ffn_out = self.ffn_drop_path(self.ffn(attn_out)) + attn_out
458 return ffn_out
459
460
461 def attn_op(self, q, k, v):

Callers

nothing calls this directly

Calls 1

attn_opMethod · 0.95

Tested by

no test coverage detected