MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / __init__

Method __init__

sat/dit_video_concat.py:268–324  ·  view source on GitHub ↗
(
        self,
        height,
        width,
        compressed_num_frames,
        hidden_size,
        hidden_size_head,
        text_length,
        theta=10000,
        rot_v=False,
        pnp=False,
        learnable_pos_embed=False,
    )

Source from the content-addressed store, hash-verified

266
267class Rotary3DPositionEmbeddingMixin(BaseMixin):
268 def __init__(
269 self,
270 height,
271 width,
272 compressed_num_frames,
273 hidden_size,
274 hidden_size_head,
275 text_length,
276 theta=10000,
277 rot_v=False,
278 pnp=False,
279 learnable_pos_embed=False,
280 ):
281 super().__init__()
282 self.rot_v = rot_v
283
284 dim_t = hidden_size_head // 4
285 dim_h = hidden_size_head // 8 * 3
286 dim_w = hidden_size_head // 8 * 3
287
288 # 'lang':
289 freqs_t = 1.0 / (theta ** (torch.arange(0, dim_t, 2)[: (dim_t // 2)].float() / dim_t))
290 freqs_h = 1.0 / (theta ** (torch.arange(0, dim_h, 2)[: (dim_h // 2)].float() / dim_h))
291 freqs_w = 1.0 / (theta ** (torch.arange(0, dim_w, 2)[: (dim_w // 2)].float() / dim_w))
292
293 grid_t = torch.arange(compressed_num_frames, dtype=torch.float32)
294 grid_h = torch.arange(height, dtype=torch.float32)
295 grid_w = torch.arange(width, dtype=torch.float32)
296
297 freqs_t = torch.einsum("..., f -> ... f", grid_t, freqs_t)
298 freqs_h = torch.einsum("..., f -> ... f", grid_h, freqs_h)
299 freqs_w = torch.einsum("..., f -> ... f", grid_w, freqs_w)
300
301 freqs_t = repeat(freqs_t, "... n -> ... (n r)", r=2)
302 freqs_h = repeat(freqs_h, "... n -> ... (n r)", r=2)
303 freqs_w = repeat(freqs_w, "... n -> ... (n r)", r=2)
304
305 freqs = broadcat((freqs_t[:, None, None, :], freqs_h[None, :, None, :], freqs_w[None, None, :, :]), dim=-1)
306 # (T H W D)
307
308 self.pnp = pnp
309
310 if not self.pnp:
311 freqs = rearrange(freqs, "t h w d -> (t h w) d")
312
313 freqs = freqs.contiguous()
314 freqs_sin = freqs.sin()
315 freqs_cos = freqs.cos()
316 self.register_buffer("freqs_sin", freqs_sin)
317 self.register_buffer("freqs_cos", freqs_cos)
318
319 self.text_length = text_length
320 if learnable_pos_embed:
321 num_patches = height * width * compressed_num_frames + text_length
322 self.pos_embedding = nn.Parameter(torch.zeros(1, num_patches, int(hidden_size)), requires_grad=True)
323 else:
324 self.pos_embedding = None
325

Callers

nothing calls this directly

Calls 2

broadcatFunction · 0.70
__init__Method · 0.45

Tested by

no test coverage detected