MCPcopy Create free account
hub / github.com/cvg/NoPoSplat / resample_patch_embed

Function resample_patch_embed

src/misc/weight_modify.py:13–84  ·  view source on GitHub ↗

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,
)

Source from the content-addressed store, hash-verified

11
12
13def 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

Callers 1

checkpoint_filter_fnFunction · 0.85

Calls 1

get_resize_matFunction · 0.85

Tested by

no test coverage detected