(self, src, mask, query_embed, pos_embed)
| 59 | nn.init.xavier_uniform_(p) |
| 60 | |
| 61 | def forward(self, src, mask, query_embed, pos_embed): |
| 62 | # flatten NxCxHxW to HWxNxC |
| 63 | bs, c, h, w = src.shape |
| 64 | src = src.flatten(2).permute(2, 0, 1) |
| 65 | pos_embed = pos_embed.flatten(2).permute(2, 0, 1) |
| 66 | query_embed = query_embed.unsqueeze(1).repeat(1, bs, 1) |
| 67 | if mask is not None: |
| 68 | mask = mask.flatten(1) |
| 69 | |
| 70 | tgt = torch.zeros_like(query_embed) |
| 71 | memory = self.encoder(src, src_key_padding_mask=mask, pos=pos_embed) |
| 72 | hs = self.decoder( |
| 73 | tgt, memory, memory_key_padding_mask=mask, pos=pos_embed, query_pos=query_embed |
| 74 | ) |
| 75 | return hs.transpose(1, 2), memory.permute(1, 2, 0).view(bs, c, h, w) |
| 76 | |
| 77 | |
| 78 | class TransformerEncoder(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected