MCPcopy Create free account
hub / github.com/OpenGVLab/HumanBench / forward_features

Method forward_features

PATH/core/models/backbones/vit.py:455–508  ·  view source on GitHub ↗
(self, x)

Source from the content-addressed store, hash-verified

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 = {}

Callers 2

forwardMethod · 0.95
forwardMethod · 0.45

Calls 2

sizeMethod · 0.80
permuteMethod · 0.80

Tested by

no test coverage detected