(pos_embed_checkpoint, patch_shape, num_extra_tokens)
| 781 | |
| 782 | |
| 783 | def interpolate_pos_embed(pos_embed_checkpoint, patch_shape, num_extra_tokens): |
| 784 | embedding_size = pos_embed_checkpoint.shape[-1] |
| 785 | orig_size = to_2tuple(int((pos_embed_checkpoint.shape[-2] - num_extra_tokens) ** 0.5)) |
| 786 | # class_token and dist_token are kept unchanged |
| 787 | print(f"[rank {dist.get_rank()}] Position interpolate from {orig_size} to {patch_shape}") |
| 788 | # only the position tokens are interpolated |
| 789 | pos_tokens = pos_embed_checkpoint[:, num_extra_tokens:] if pos_embed_checkpoint.size(0) == 1 else pos_embed_checkpoint[num_extra_tokens:] |
| 790 | pos_tokens = pos_tokens.reshape(-1, orig_size[0], orig_size[1], embedding_size).permute(0, 3, 1, 2) |
| 791 | pos_tokens = torch.nn.functional.interpolate(pos_tokens, size=patch_shape, mode='bicubic', align_corners=False) |
| 792 | new_pos_embed = pos_tokens.permute(0, 2, 3, 1).flatten(1, 2) # (b, h*w, c) |
| 793 | return new_pos_embed |
| 794 | |
| 795 | |
| 796 | def interpolate_pos_embed_with_cls_token(pos_embed_checkpoint, patch_shape, num_extra_tokens): |
no test coverage detected