| 112 | |
| 113 | |
| 114 | class AttnBlock(nn.Module): |
| 115 | def __init__(self, in_channels): |
| 116 | super().__init__() |
| 117 | self.in_channels = in_channels |
| 118 | |
| 119 | self.norm = Normalize(in_channels) |
| 120 | self.q = torch.nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0) |
| 121 | self.k = torch.nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0) |
| 122 | self.v = torch.nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0) |
| 123 | self.proj_out = torch.nn.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0) |
| 124 | |
| 125 | def forward(self, x): |
| 126 | h_ = x |
| 127 | h_ = self.norm(h_) |
| 128 | q = self.q(h_) |
| 129 | k = self.k(h_) |
| 130 | v = self.v(h_) |
| 131 | |
| 132 | # compute attention |
| 133 | b, c, h, w = q.shape |
| 134 | q = q.reshape(b, c, h * w) |
| 135 | q = q.permute(0, 2, 1) # b,hw,c |
| 136 | k = k.reshape(b, c, h * w) # b,c,hw |
| 137 | |
| 138 | # # original version, nan in fp16 |
| 139 | # w_ = torch.bmm(q,k) # b,hw,hw w[b,i,j]=sum_c q[b,i,c]k[b,c,j] |
| 140 | # w_ = w_ * (int(c)**(-0.5)) |
| 141 | # # implement c**-0.5 on q |
| 142 | q = q * (int(c) ** (-0.5)) |
| 143 | w_ = torch.bmm(q, k) # b,hw,hw w[b,i,j]=sum_c q[b,i,c]k[b,c,j] |
| 144 | |
| 145 | w_ = torch.nn.functional.softmax(w_, dim=2) |
| 146 | |
| 147 | # attend to values |
| 148 | v = v.reshape(b, c, h * w) |
| 149 | w_ = w_.permute(0, 2, 1) # b,hw,hw (first hw of k, second of q) |
| 150 | h_ = torch.bmm(v, w_) # b, c,hw (hw of q) h_[b,c,j] = sum_i v[b,c,i] w_[b,i,j] |
| 151 | h_ = h_.reshape(b, c, h, w) |
| 152 | |
| 153 | h_ = self.proj_out(h_) |
| 154 | |
| 155 | return x + h_ |
| 156 | |
| 157 | |
| 158 | class Encoder(nn.Module): |