leochen085 commited on
Commit
3b29d80
·
verified ·
1 Parent(s): 5ddd9db

Remove old util/ directory

Browse files
Files changed (2) hide show
  1. util/__init__.py +0 -0
  2. util/pos_embed.py +0 -246
util/__init__.py DELETED
File without changes
util/pos_embed.py DELETED
@@ -1,246 +0,0 @@
1
- import numpy as np
2
- import torch
3
-
4
-
5
- def get_1d_sincos_pos_embed(embed_dim, length, cls_token=False):
6
- """
7
- Create 1D sine-cosine positional embeddings.
8
-
9
- Args:
10
- embed_dim (int): Dimension of the embedding (must be even)
11
- length (int): Number of positions (sequence length)
12
- cls_token (bool): Whether to include an extra zero vector for [CLS] token
13
-
14
- Returns:
15
- np.ndarray of shape (length, embed_dim) or (1+length, embed_dim) if cls_token=True
16
- """
17
- # position indices 0 ... length-1
18
- pos = np.arange(length, dtype=np.float32)
19
-
20
- # get embedding from grid
21
- pos_embed = get_1d_sincos_pos_embed_from_grid(embed_dim, pos) # (L, D)
22
-
23
- # optionally add CLS token embedding
24
- if cls_token:
25
- pos_embed = np.concatenate([np.zeros([1, embed_dim]), pos_embed], axis=0)
26
- return pos_embed
27
-
28
- def get_2d_sincos_pos_embed(embed_dim, grid_size, cls_token=False):
29
- # --------------------------------------------------------
30
- # 2D sine-cosine position embedding
31
- # References:
32
- # Transformer: https://github.com/tensorflow/models/blob/master/official/nlp/transformer/model_utils.py
33
- # MoCo v3: https://github.com/facebookresearch/moco-v3
34
- # --------------------------------------------------------
35
-
36
- grid_h = np.arange(grid_size[0], dtype=np.float32)
37
- grid_w = np.arange(grid_size[1], dtype=np.float32)
38
- grid = np.meshgrid(grid_w, grid_h) # here w goes first
39
- grid = np.stack(grid, axis=0)
40
-
41
- grid = grid.reshape([2, 1, grid_size[0], grid_size[1]])
42
- pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid)
43
- if cls_token:
44
- pos_embed = np.concatenate([np.zeros([1, embed_dim]), pos_embed], axis=0)
45
- return pos_embed
46
-
47
- def get_2d_sincos_pos_embed_from_grid(embed_dim, grid):
48
- assert embed_dim % 2 == 0
49
-
50
- # use half of dimensions to encode grid_h
51
- emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0]) #changed(H*W, D/2)
52
- emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1]) #changed (H*W, D/2)
53
-
54
- emb = np.concatenate([emb_h, emb_w], axis=1) # (H*W, D)
55
- return emb
56
-
57
- def get_1d_sincos_pos_embed_from_grid(embed_dim, pos):
58
- """
59
- embed_dim: output dimension for each position
60
- pos: a list of positions to be encoded: size (M,)
61
- out: (M, D)
62
- """
63
- assert embed_dim % 2 == 0
64
- omega = np.arange(embed_dim // 2, dtype=np.float32)
65
- omega /= embed_dim / 2.
66
- omega = 1. / 10000**omega # (D/2,)
67
-
68
- pos = pos.reshape(-1) # (M,)
69
- out = np.einsum('m,d->md', pos, omega) # (M, D/2), outer product
70
-
71
- emb_sin = np.sin(out) # (M, D/2)
72
- emb_cos = np.cos(out) # (M, D/2)
73
-
74
- emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D)
75
- return emb
76
-
77
- def interpolate_pos_embed(model, checkpoint_model, orig_size, new_size):
78
- '''
79
- Input: model: the class is definging for downstream
80
- checkpoint_model: pre-train weight
81
- orig_size = patch size in the ckpt
82
- new_size = patch size in the current model
83
- '''
84
-
85
- if 'pos_embed' in checkpoint_model:
86
- pos_embed_checkpoint = checkpoint_model['pos_embed'] # 1 x 560 x 768 (1 x num_patches x E)
87
- embedding_size = pos_embed_checkpoint.shape[-1] # 768
88
-
89
- # number of special tokens (e.g. in this case num_extra_tokens = 1 for the cls token)
90
- num_patches = model.patch_embed.num_patches
91
- num_extra_tokens = model.pos_embed.shape[-2] - num_patches
92
-
93
- if orig_size != new_size:
94
- print("Position interpolate from %dx%d to %dx%d" % (orig_size[0], orig_size[1], new_size[0], new_size[1]))
95
- extra_tokens = pos_embed_checkpoint[:, :num_extra_tokens]
96
- # only the position tokens are interpolated
97
- pos_tokens = pos_embed_checkpoint[:, num_extra_tokens:] # old positions
98
- pos_tokens = pos_tokens.reshape(-1, orig_size[0], orig_size[1], embedding_size).permute(0, 3, 1, 2)
99
- pos_tokens = torch.nn.functional.interpolate(
100
- pos_tokens, size=(new_size[0], new_size[1]), mode='bicubic', align_corners=False)
101
- pos_tokens = pos_tokens.permute(0, 2, 3, 1).flatten(1, 2)
102
- new_pos_embed = torch.cat((extra_tokens, pos_tokens), dim=1)
103
- checkpoint_model['pos_embed'] = new_pos_embed
104
-
105
-
106
-
107
- # RoPE: https://huggingface.co/thuml/sundial-base-128m/blob/main/modeling_sundial.py
108
- class RotaryEmbedding(torch.nn.Module):
109
- def __init__(self, dim, max_position_embeddings=10000, base=10000, device=None):
110
- super().__init__()
111
- self.dim = dim
112
- self.max_position_embeddings = max_position_embeddings
113
- self.base = base
114
- inv_freq = 1.0 / (self.base ** (torch.arange(0, self.dim,
115
- 2, dtype=torch.int64).float().to(device) / self.dim))
116
- self.register_buffer("inv_freq", inv_freq, persistent=False)
117
-
118
- # Build here to make `torch.jit.trace` work.
119
- self._set_cos_sin_cache(
120
- seq_len=max_position_embeddings, device=self.inv_freq.device, dtype=torch.get_default_dtype()
121
- )
122
-
123
- def _set_cos_sin_cache(self, seq_len, device, dtype):
124
- self.max_seq_len_cached = seq_len
125
- t = torch.arange(self.max_seq_len_cached, device=device,
126
- dtype=torch.int64).type_as(self.inv_freq)
127
-
128
- freqs = torch.outer(t, self.inv_freq)
129
- # Different from paper, but it uses a different permutation in order to obtain the same calculation
130
- emb = torch.cat((freqs, freqs), dim=-1)
131
- self.register_buffer(
132
- "cos_cached", emb.cos().to(dtype), persistent=False)
133
- self.register_buffer(
134
- "sin_cached", emb.sin().to(dtype), persistent=False)
135
-
136
- def forward(self, x, seq_len=None):
137
- # x: [bs, num_attention_heads, seq_len, head_size]
138
- if seq_len > self.max_seq_len_cached:
139
- self._set_cos_sin_cache(
140
- seq_len=seq_len, device=x.device, dtype=x.dtype)
141
-
142
- return (
143
- self.cos_cached[:seq_len].to(dtype=x.dtype),
144
- self.sin_cached[:seq_len].to(dtype=x.dtype),
145
- )
146
-
147
- def rotate_half(x):
148
- x1 = x[..., : x.shape[-1] // 2]
149
- x2 = x[..., x.shape[-1] // 2:]
150
- return torch.cat((-x2, x1), dim=-1)
151
-
152
-
153
- def apply_rotary_pos_emb(q, k, cos, sin, position_ids, unsqueeze_dim=1):
154
- cos = cos[position_ids].unsqueeze(unsqueeze_dim)
155
- sin = sin[position_ids].unsqueeze(unsqueeze_dim)
156
- q_embed = (q * cos) + (rotate_half(q) * sin)
157
- k_embed = (k * cos) + (rotate_half(k) * sin)
158
- return q_embed, k_embed
159
-
160
- # two dimensional version
161
- def apply_rotary_pos_emb_2d(q, k,
162
- cos_h, sin_h,
163
- cos_w, sin_w,
164
- pos_h, pos_w,
165
- unsqueeze_dim=1):
166
- """
167
- q, k: [B, heads, N, Dh]
168
- cos_h, sin_h: caches from 1D rotary with dim = Dh // 2 for the first axis
169
- cos_w, sin_w: caches from 1D rotary with dim = Dh // 2 for the second axis
170
- pos_h, pos_w: [B, N] integer positions for each token along the two axes
171
- returns q_out, k_out with same shape as q, k
172
- """
173
- Dh = q.shape[-1]
174
- assert Dh % 4 == 0, "head dim must be divisible by 4 so each half is even for rotate_half"
175
-
176
- # split channel dim into two halves
177
- q_h, q_w = q.split(Dh // 2, dim=-1)
178
- k_h, k_w = k.split(Dh // 2, dim=-1)
179
-
180
- # apply 1D RoPE on each half with its own positions
181
- pos_h = pos_h.long()
182
- pos_w = pos_w.long()
183
- q_h, k_h = apply_rotary_pos_emb(q_h, k_h, cos_h, sin_h, pos_h, unsqueeze_dim=unsqueeze_dim)
184
- q_w, k_w = apply_rotary_pos_emb(q_w, k_w, cos_w, sin_w, pos_w, unsqueeze_dim=unsqueeze_dim)
185
-
186
- # concat back
187
- q_out = torch.cat([q_h, q_w], dim=-1)
188
- k_out = torch.cat([k_h, k_w], dim=-1)
189
- return q_out, k_out
190
-
191
-
192
- def build_2d_position_ids(attention_mask: torch.Tensor,
193
- flatten: bool = True):
194
- """
195
- attention_mask: Tensor [BS, nvar, num_p] with 1 for valid patches, 0 for padding.
196
-
197
- Returns:
198
- If flatten is True:
199
- pos_var_flat: LongTensor [BS, nvar*num_p]
200
- pos_patch_flat: LongTensor [BS, nvar*num_p]
201
- Else:
202
- pos_var: LongTensor [BS, nvar, num_p]
203
- pos_patch: LongTensor [BS, nvar, num_p]
204
- """
205
- assert attention_mask.dim() == 3, "attention_mask must be [BS, nvar, num_p]"
206
- B, V, P = attention_mask.shape
207
- mask = attention_mask.to(dtype=torch.long)
208
-
209
- # per patch index within each variable, ignores padding
210
- pos_patch = (mask.cumsum(dim=-1) - 1) * mask # [B, V, P]
211
-
212
- # per variable index, ignores variables that are entirely padded
213
- var_valid = mask.any(dim=-1).to(dtype=torch.long) # [B, V]
214
- pos_var_base = (var_valid.cumsum(dim=1) - 1) * var_valid # [B, V]
215
- pos_var = pos_var_base.unsqueeze(-1).expand(B, V, P) * mask # [B, V, P]
216
-
217
- if flatten:
218
- return pos_var.reshape(B, V * P).long(), pos_patch.reshape(B, V * P).long()
219
-
220
- return pos_var.long(), pos_patch.long()
221
-
222
- def build_1d_position_ids(attention_mask: torch.Tensor):
223
- """
224
- Build 1D position ids for [BS, nvar, num_p],
225
- output shape [BS * nvar, num_p].
226
-
227
- Each (batch, variable) pair gets its own 1D position index sequence
228
- along the patch axis, skipping padded positions.
229
-
230
- Args:
231
- attention_mask: Tensor [BS, nvar, num_p], 1 for valid, 0 for padding.
232
-
233
- Returns:
234
- pos_ids: LongTensor [BS * nvar, num_p]
235
- """
236
- assert attention_mask.dim() == 3, "attention_mask must be [BS, nvar, num_p]"
237
- B, V, P = attention_mask.shape
238
- mask = attention_mask.to(dtype=torch.long)
239
-
240
- # Compute per-variable cumulative index
241
- pos_ids = (mask.cumsum(dim=-1) - 1) * mask # [B, V, P]
242
-
243
- # Reshape to [BS * nvar, num_p]
244
- pos_ids = pos_ids.view(B * V, P).long()
245
-
246
- return pos_ids