Resample the weights of the patch embedding kernel to target resolution. We resample the patch embedding kernel by approximately inverting the effect of patch resizing. Code based on: https://github.com/google-research/big_vision/blob/b00544b81f8694488d5f36295aeb7972f3755ffe/big_v
(
patch_embed,
new_size: List[int],
interpolation: str = 'bicubic',
antialias: bool = True,
verbose: bool = False,
)
| 11 | |
| 12 | |
| 13 | def resample_patch_embed( |
| 14 | patch_embed, |
| 15 | new_size: List[int], |
| 16 | interpolation: str = 'bicubic', |
| 17 | antialias: bool = True, |
| 18 | verbose: bool = False, |
| 19 | ): |
| 20 | """Resample the weights of the patch embedding kernel to target resolution. |
| 21 | We resample the patch embedding kernel by approximately inverting the effect |
| 22 | of patch resizing. |
| 23 | |
| 24 | Code based on: |
| 25 | https://github.com/google-research/big_vision/blob/b00544b81f8694488d5f36295aeb7972f3755ffe/big_vision/models/proj/flexi/vit.py |
| 26 | |
| 27 | With this resizing, we can for example load a B/8 filter into a B/16 model |
| 28 | and, on 2x larger input image, the result will match. |
| 29 | |
| 30 | Args: |
| 31 | patch_embed: original parameter to be resized. |
| 32 | new_size (tuple(int, int): target shape (height, width)-only. |
| 33 | interpolation (str): interpolation for resize |
| 34 | antialias (bool): use anti-aliasing filter in resize |
| 35 | verbose (bool): log operation |
| 36 | Returns: |
| 37 | Resized patch embedding kernel. |
| 38 | """ |
| 39 | import numpy as np |
| 40 | try: |
| 41 | import functorch |
| 42 | vmap = functorch.vmap |
| 43 | except ImportError: |
| 44 | if hasattr(torch, 'vmap'): |
| 45 | vmap = torch.vmap |
| 46 | else: |
| 47 | assert False, "functorch or a version of torch with vmap is required for FlexiViT resizing." |
| 48 | |
| 49 | assert len(patch_embed.shape) == 4, "Four dimensions expected" |
| 50 | assert len(new_size) == 2, "New shape should only be hw" |
| 51 | old_size = patch_embed.shape[-2:] |
| 52 | if tuple(old_size) == tuple(new_size): |
| 53 | return patch_embed |
| 54 | |
| 55 | if verbose: |
| 56 | _logger.info(f"Resize patch embedding {patch_embed.shape} to {new_size}, w/ {interpolation} interpolation.") |
| 57 | |
| 58 | def resize(x_np, _new_size): |
| 59 | x_tf = torch.Tensor(x_np)[None, None, ...] |
| 60 | x_upsampled = F.interpolate( |
| 61 | x_tf, size=_new_size, mode=interpolation, antialias=antialias)[0, 0, ...].numpy() |
| 62 | return x_upsampled |
| 63 | |
| 64 | def get_resize_mat(_old_size, _new_size): |
| 65 | mat = [] |
| 66 | for i in range(np.prod(_old_size)): |
| 67 | basis_vec = np.zeros(_old_size) |
| 68 | basis_vec[np.unravel_index(i, _old_size)] = 1. |
| 69 | mat.append(resize(basis_vec, _new_size).reshape(-1)) |
| 70 | return np.stack(mat).T |
no test coverage detected