MCPcopy Create free account
hub / github.com/UVA-Computer-Vision-Lab/FrameINO / __init__

Method __init__

architecture/embeddings.py:462–515  ·  view source on GitHub ↗
(
        self,
        height=224,
        width=224,
        patch_size=16,
        in_channels=3,
        embed_dim=768,
        layer_norm=False,
        flatten=True,
        bias=True,
        interpolation_scale=1,
        pos_embed_type="sincos",
        pos_embed_max_size=None,  # For SD3 cropping
    )

Source from the content-addressed store, hash-verified

460 """
461
462 def __init__(
463 self,
464 height=224,
465 width=224,
466 patch_size=16,
467 in_channels=3,
468 embed_dim=768,
469 layer_norm=False,
470 flatten=True,
471 bias=True,
472 interpolation_scale=1,
473 pos_embed_type="sincos",
474 pos_embed_max_size=None, # For SD3 cropping
475 ):
476 super().__init__()
477
478 num_patches = (height // patch_size) * (width // patch_size)
479 self.flatten = flatten
480 self.layer_norm = layer_norm
481 self.pos_embed_max_size = pos_embed_max_size
482
483 self.proj = nn.Conv2d(
484 in_channels, embed_dim, kernel_size=(patch_size, patch_size), stride=patch_size, bias=bias
485 )
486 if layer_norm:
487 self.norm = nn.LayerNorm(embed_dim, elementwise_affine=False, eps=1e-6)
488 else:
489 self.norm = None
490
491 self.patch_size = patch_size
492 self.height, self.width = height // patch_size, width // patch_size
493 self.base_size = height // patch_size
494 self.interpolation_scale = interpolation_scale
495
496 # Calculate positional embeddings based on max size or default
497 if pos_embed_max_size:
498 grid_size = pos_embed_max_size
499 else:
500 grid_size = int(num_patches**0.5)
501
502 if pos_embed_type is None:
503 self.pos_embed = None
504 elif pos_embed_type == "sincos":
505 pos_embed = get_2d_sincos_pos_embed(
506 embed_dim,
507 grid_size,
508 base_size=self.base_size,
509 interpolation_scale=self.interpolation_scale,
510 output_type="pt",
511 )
512 persistent = True if pos_embed_max_size else False
513 self.register_buffer("pos_embed", pos_embed.float().unsqueeze(0), persistent=persistent)
514 else:
515 raise ValueError(f"Unsupported pos_embed_type: {pos_embed_type}")
516
517 def cropped_pos_embed(self, height, width):
518 """Crops positional embeddings for SD3 compatibility."""

Callers

nothing calls this directly

Calls 2

get_2d_sincos_pos_embedFunction · 0.70
__init__Method · 0.45

Tested by

no test coverage detected