| 220 | |
| 221 | @torch.no_grad() |
| 222 | def forward(self, results, outputs, orig_target_sizes, max_target_sizes): |
| 223 | assert len(orig_target_sizes) == len(max_target_sizes) |
| 224 | max_h, max_w = max_target_sizes.max(0)[0].tolist() |
| 225 | outputs_masks = outputs["pred_masks"].squeeze(2) |
| 226 | outputs_masks = F.interpolate(outputs_masks, size=(max_h, max_w), mode="bilinear", align_corners=False) |
| 227 | outputs_masks = (outputs_masks.sigmoid() > self.threshold).cpu() |
| 228 | |
| 229 | for i, (cur_mask, t, tt) in enumerate(zip(outputs_masks, max_target_sizes, orig_target_sizes)): |
| 230 | img_h, img_w = t[0], t[1] |
| 231 | results[i]["masks"] = cur_mask[:, :img_h, :img_w].unsqueeze(1) |
| 232 | results[i]["masks"] = F.interpolate( |
| 233 | results[i]["masks"].float(), size=tuple(tt.tolist()), mode="nearest" |
| 234 | ).byte() |
| 235 | |
| 236 | return results |
| 237 | |
| 238 | |
| 239 | class PostProcessPanoptic(nn.Module): |