Input: - memory: bs, \sum{hw}, d_model - memory_padding_mask: bs, \sum{hw} - spatial_shapes: nlevel, 2 Output: - output_memory: bs, \sum{hw}, d_model - output_proposals: bs, \sum{hw}, 4
(memory:Tensor, memory_padding_mask:Tensor, spatial_shapes:Tensor)
| 31 | |
| 32 | |
| 33 | def gen_encoder_output_proposals(memory:Tensor, memory_padding_mask:Tensor, spatial_shapes:Tensor): |
| 34 | """ |
| 35 | Input: |
| 36 | - memory: bs, \sum{hw}, d_model |
| 37 | - memory_padding_mask: bs, \sum{hw} |
| 38 | - spatial_shapes: nlevel, 2 |
| 39 | Output: |
| 40 | - output_memory: bs, \sum{hw}, d_model |
| 41 | - output_proposals: bs, \sum{hw}, 4 |
| 42 | """ |
| 43 | N_, S_, C_ = memory.shape |
| 44 | base_scale = 4.0 |
| 45 | proposals = [] |
| 46 | _cur = 0 |
| 47 | for lvl, (H_, W_) in enumerate(spatial_shapes): |
| 48 | mask_flatten_ = memory_padding_mask[:, _cur:(_cur + H_ * W_)].view(N_, H_, W_, 1) |
| 49 | valid_H = torch.sum(~mask_flatten_[:, :, 0, 0], 1) |
| 50 | valid_W = torch.sum(~mask_flatten_[:, 0, :, 0], 1) |
| 51 | |
| 52 | grid_y, grid_x = torch.meshgrid(torch.linspace(0, H_ - 1, H_, dtype=torch.float32, device=memory.device), |
| 53 | torch.linspace(0, W_ - 1, W_, dtype=torch.float32, device=memory.device)) |
| 54 | grid = torch.cat([grid_x.unsqueeze(-1), grid_y.unsqueeze(-1)], -1) |
| 55 | |
| 56 | scale = torch.cat([valid_W.unsqueeze(-1), valid_H.unsqueeze(-1)], 1).view(N_, 1, 1, 2) |
| 57 | grid = (grid.unsqueeze(0).expand(N_, -1, -1, -1) + 0.5) / scale |
| 58 | wh = torch.ones_like(grid) * 0.05 * (2.0 ** lvl) |
| 59 | proposal = torch.cat((grid, wh), -1).view(N_, -1, 4) |
| 60 | proposals.append(proposal) |
| 61 | _cur += (H_ * W_) |
| 62 | output_proposals = torch.cat(proposals, 1) |
| 63 | output_proposals_valid = ((output_proposals > 0.01) & (output_proposals < 0.99)).all(-1, keepdim=True) |
| 64 | output_proposals = torch.log(output_proposals / (1 - output_proposals)) |
| 65 | output_proposals = output_proposals.masked_fill(memory_padding_mask.unsqueeze(-1), float('inf')) |
| 66 | output_proposals = output_proposals.masked_fill(~output_proposals_valid, float('inf')) |
| 67 | |
| 68 | output_memory = memory |
| 69 | output_memory = output_memory.masked_fill(memory_padding_mask.unsqueeze(-1), float(0)) |
| 70 | output_memory = output_memory.masked_fill(~output_proposals_valid, float(0)) |
| 71 | return output_memory, output_proposals |
| 72 | |
| 73 | |
| 74 | def gen_sineembed_for_position(pos_tensor): |