Interpolate position embeddings for high-resolution.
(model, checkpoint_model)
| 255 | |
| 256 | |
| 257 | def interpolate_pos_embed(model, checkpoint_model): |
| 258 | """Interpolate position embeddings for high-resolution.""" |
| 259 | if 'pos_embed' in checkpoint_model: |
| 260 | pos_embed_checkpoint = checkpoint_model['pos_embed'] |
| 261 | embedding_size = pos_embed_checkpoint.shape[-1] |
| 262 | num_patches = model.patch_embed.num_patches |
| 263 | num_extra_tokens = model.pos_embed.shape[-2] - num_patches |
| 264 | orig_size = int((pos_embed_checkpoint.shape[-2] - num_extra_tokens) ** 0.5) |
| 265 | new_size = int(num_patches ** 0.5) |
| 266 | if orig_size != new_size: |
| 267 | print("Position interpolate from %dx%d to %dx%d" % (orig_size, orig_size, new_size, new_size)) |
| 268 | extra_tokens = pos_embed_checkpoint[:, :num_extra_tokens] |
| 269 | pos_tokens = pos_embed_checkpoint[:, num_extra_tokens:] |
| 270 | pos_tokens = pos_tokens.reshape(-1, orig_size, orig_size, embedding_size).permute(0, 3, 1, 2) |
| 271 | pos_tokens = torch.nn.functional.interpolate( |
| 272 | pos_tokens, size=(new_size, new_size), mode='bicubic', align_corners=False) |
| 273 | pos_tokens = pos_tokens.permute(0, 2, 3, 1).flatten(1, 2) |
| 274 | new_pos_embed = torch.cat((extra_tokens, pos_tokens), dim=1) |
| 275 | checkpoint_model['pos_embed'] = new_pos_embed |
nothing calls this directly
no outgoing calls
no test coverage detected