| 21 | |
| 22 | |
| 23 | def get_abs_pos(abs_pos, tgt_size): |
| 24 | # abs_pos: L, C |
| 25 | # tgt_size: M |
| 26 | # return: M, C |
| 27 | src_size = int(math.sqrt(abs_pos.size(0))) |
| 28 | tgt_size = int(math.sqrt(tgt_size)) |
| 29 | dtype = abs_pos.dtype |
| 30 | |
| 31 | if src_size != tgt_size: |
| 32 | return F.interpolate( |
| 33 | abs_pos.float().reshape(1, src_size, src_size, -1).permute(0, 3, 1, 2), |
| 34 | size=(tgt_size, tgt_size), |
| 35 | mode="bicubic", |
| 36 | align_corners=False, |
| 37 | ).permute(0, 2, 3, 1).flatten(0, 2).to(dtype=dtype) |
| 38 | else: |
| 39 | return abs_pos |
| 40 | |
| 41 | # https://github.com/facebookresearch/mae/blob/efb2a8062c206524e35e47d04501ed4f544c0ae8/util/pos_embed.py#L20 |
| 42 | def get_2d_sincos_pos_embed(embed_dim, grid_size, cls_token=False): |