| 453 | return len(self.blocks) |
| 454 | |
| 455 | def forward_features(self, x): |
| 456 | B, C, H, W = x.shape |
| 457 | x, (Hp, Wp) = self.patch_embed(x) |
| 458 | batch_size, seq_len, _ = x.size() |
| 459 | |
| 460 | x = x + self.pos_embed[:, 1:] + self.pos_embed[:, :1] |
| 461 | |
| 462 | # if self.test_pos_mode is False: |
| 463 | # # x = x + self.pos_embed |
| 464 | # x = x + self.pos_embed[:, 1:] + self.pos_embed[:, :1] |
| 465 | # elif self.test_pos_mode == 'regenerate': |
| 466 | # pos_embed = get_2d_sincos_pos_embed(self.pos_embed.shape[-1], (Hp, Wp), cls_token=False) |
| 467 | # x = x + torch.from_numpy(pos_embed).float().unscqueeze(0).cuda() |
| 468 | # elif self.test_pos_mode == 'scaled_regenerate': |
| 469 | # patch_shape = (Hp, Wp) |
| 470 | # orig_size = (math.ceil(Hp/20)*7, math.ceil(Wp/20)*7) |
| 471 | |
| 472 | # # as in original scale |
| 473 | # pos_embed = get_2d_sincos_pos_embed(self.pos_embed.shape[-1], orig_size, cls_token=False) |
| 474 | # pos_embed = torch.from_numpy(pos_embed).float().unsqueeze(0).cuda() |
| 475 | |
| 476 | # # as in finetuning scale |
| 477 | # pos_embed = pos_embed.reshape(-1, orig_size[0], orig_size[1], self.pos_embed.shape[-1]).permute(0, 3, 1, 2) |
| 478 | # pos_embed = torch.nn.functional.interpolate(pos_embed, size=(orig_size[0]//7*20, orig_size[1]//7*20), |
| 479 | # mode='bicubic', align_corners=False) |
| 480 | |
| 481 | # # as in test image |
| 482 | # pos_embed = pos_embed[:, :, :patch_shape[0], :patch_shape[1]].permute(0, 2, 3, 1).flatten(1, 2) |
| 483 | |
| 484 | # x = x + pos_embed |
| 485 | # elif self.test_pos_mode == 'simple_interpolate': |
| 486 | # patch_shape = (Hp, Wp) |
| 487 | # orig_size = (14, 14) |
| 488 | |
| 489 | # # as in original scale |
| 490 | # pos_embed = get_2d_sincos_pos_embed(self.pos_embed.shape[-1], orig_size, cls_token=False) |
| 491 | # pos_embed = torch.from_numpy(pos_embed).float().unsqueeze(0).cuda() |
| 492 | |
| 493 | # # as in finetuning scale |
| 494 | # pos_embed = pos_embed.reshape(-1, orig_size[0], orig_size[1], self.pos_embed.shape[-1]).permute(0, 3, 1, 2) |
| 495 | # pos_embed = torch.nn.functional.interpolate(pos_embed, size=patch_shape, mode='bicubic', align_corners=False) |
| 496 | # pos_embed = pos_embed.permute(0, 2, 3, 1).flatten(1, 2) |
| 497 | |
| 498 | # x = x + pos_embed |
| 499 | # else: |
| 500 | # raise NotImplementedError |
| 501 | |
| 502 | x = self.pos_drop(x) |
| 503 | |
| 504 | for i, blk in enumerate(self.blocks): |
| 505 | x = blk(x, Hp, Wp) |
| 506 | |
| 507 | x = self.norm(x) |
| 508 | return x.permute(0, 2, 1).reshape(B, -1, Hp, Wp) |
| 509 | |
| 510 | def forward(self, input_var): |
| 511 | output = {} |