| 37 | |
| 38 | |
| 39 | class AttnBlock(nn.Module): |
| 40 | def __init__(self, in_channels): |
| 41 | super().__init__() |
| 42 | self.in_channels = in_channels |
| 43 | |
| 44 | self.norm = Normalize(in_channels) |
| 45 | self.q = torch.nn.Conv2d(in_channels, |
| 46 | in_channels, |
| 47 | kernel_size=1, |
| 48 | stride=1, |
| 49 | padding=0) |
| 50 | self.k = torch.nn.Conv2d(in_channels, |
| 51 | in_channels, |
| 52 | kernel_size=1, |
| 53 | stride=1, |
| 54 | padding=0) |
| 55 | self.v = torch.nn.Conv2d(in_channels, |
| 56 | in_channels, |
| 57 | kernel_size=1, |
| 58 | stride=1, |
| 59 | padding=0) |
| 60 | self.proj_out = torch.nn.Conv2d(in_channels, |
| 61 | in_channels, |
| 62 | kernel_size=1, |
| 63 | stride=1, |
| 64 | padding=0) |
| 65 | |
| 66 | def forward(self, x): |
| 67 | h_ = x |
| 68 | h_ = self.norm(h_) |
| 69 | q = self.q(h_) |
| 70 | k = self.k(h_) |
| 71 | v = self.v(h_) |
| 72 | |
| 73 | # compute attention |
| 74 | b, c, h, w = q.shape |
| 75 | q = q.reshape(b, c, h * w) |
| 76 | q = q.permute(0, 2, 1) # b,hw,c |
| 77 | k = k.reshape(b, c, h * w) # b,c,hw |
| 78 | w_ = torch.bmm(q, k) # b,hw,hw w[b,i,j]=sum_c q[b,i,c]k[b,c,j] |
| 79 | w_ = w_ * (int(c)**(-0.5)) |
| 80 | w_ = torch.nn.functional.softmax(w_, dim=2) |
| 81 | |
| 82 | # attend to values |
| 83 | v = v.reshape(b, c, h * w) |
| 84 | w_ = w_.permute(0, 2, 1) # b,hw,hw (first hw of k, second of q) |
| 85 | # b, c,hw (hw of q) h_[b,c,j] = sum_i v[b,c,i] w_[b,i,j] |
| 86 | h_ = torch.bmm(v, w_) |
| 87 | h_ = h_.reshape(b, c, h, w) |
| 88 | |
| 89 | h_ = self.proj_out(h_) |
| 90 | |
| 91 | return x + h_ |
nothing calls this directly
no outgoing calls
no test coverage detected