| |
| |
|
|
| |
| |
| |
| |
| |
|
|
| import util.logging as logging |
| import numpy as np |
| import torch |
|
|
|
|
| logger = logging.get_logger(__name__) |
|
|
|
|
| |
| |
| |
| |
| |
| def interpolate_pos_embed(model, checkpoint_model): |
| if "pos_embed" in checkpoint_model: |
| pos_embed_checkpoint = checkpoint_model["pos_embed"] |
| embedding_size = pos_embed_checkpoint.shape[-1] |
| num_patches = model.patch_embed.num_patches |
| num_extra_tokens = model.pos_embed.shape[-2] - num_patches |
| |
| orig_size = int((pos_embed_checkpoint.shape[-2] - num_extra_tokens) ** 0.5) |
| |
| new_size = int(num_patches**0.5) |
| |
| if orig_size != new_size: |
| print( |
| "Position interpolate from %dx%d to %dx%d" |
| % (orig_size, orig_size, new_size, new_size) |
| ) |
| extra_tokens = pos_embed_checkpoint[:, :num_extra_tokens] |
| |
| pos_tokens = pos_embed_checkpoint[:, num_extra_tokens:] |
| pos_tokens = pos_tokens.reshape( |
| -1, orig_size, orig_size, embedding_size |
| ).permute(0, 3, 1, 2) |
| pos_tokens = torch.nn.functional.interpolate( |
| pos_tokens, |
| size=(new_size, new_size), |
| mode="bicubic", |
| align_corners=False, |
| ) |
| pos_tokens = pos_tokens.permute(0, 2, 3, 1).flatten(1, 2) |
| new_pos_embed = torch.cat((extra_tokens, pos_tokens), dim=1) |
| checkpoint_model["pos_embed"] = new_pos_embed |
|
|