(self, x, y, xpos, ypos, return_attn=False)
| 317 | self.norm_y = norm_layer(dim) if norm_mem else nn.Identity() |
| 318 | |
| 319 | def forward(self, x, y, xpos, ypos, return_attn=False): |
| 320 | if return_attn: |
| 321 | self_attn_output, self_attn = self.attn(self.norm1(x), xpos, return_attn=True) |
| 322 | x = x + self.drop_path(self_attn_output) |
| 323 | y_ = self.norm_y(y) |
| 324 | cross_attn_output, cross_attn = self.cross_attn(self.norm2(x), y_, y_, xpos, ypos, return_attn=True) |
| 325 | x = x + self.drop_path(cross_attn_output) |
| 326 | x = x + self.drop_path(self.mlp(self.norm3(x))) |
| 327 | return x, y, self_attn, cross_attn |
| 328 | else: |
| 329 | x = x + self.drop_path(self.attn(self.norm1(x), xpos)) |
| 330 | y_ = self.norm_y(y) |
| 331 | x = x + self.drop_path(self.cross_attn(self.norm2(x), y_, y_, xpos, ypos)) |
| 332 | x = x + self.drop_path(self.mlp(self.norm3(x))) |
| 333 | return x, y, None, None |
| 334 | |
| 335 | class CustomDecoderBlock(nn.Module): |
| 336 |
nothing calls this directly
no outgoing calls
no test coverage detected