MCPcopy Create free account
hub / github.com/AnswerDotAI/ModernBERT / tile_embedding

Function tile_embedding

src/bert_layers/initialization.py:515–551  ·  view source on GitHub ↗

Tile the weights of an embedding layer to a new, larger embedding dimension. Args: pretrained_embedding (nn.Embedding): The original embedding layer new_embedding (nn.Embedding): The new embedding layer with larger embedding_dim tile_mode (Union[str, TileMode]): The Phi-style w

(
    pretrained_embedding: nn.Embedding,
    new_embedding: nn.Embedding,
    mode: Union[str, TileMode] = TileMode.tile_weights_from_middle,
)

Source from the content-addressed store, hash-verified

513
514
515def tile_embedding(
516 pretrained_embedding: nn.Embedding,
517 new_embedding: nn.Embedding,
518 mode: Union[str, TileMode] = TileMode.tile_weights_from_middle,
519) -> nn.Embedding:
520 """
521 Tile the weights of an embedding layer to a new, larger embedding dimension.
522
523 Args:
524 pretrained_embedding (nn.Embedding): The original embedding layer
525 new_embedding (nn.Embedding): The new embedding layer with larger embedding_dim
526 tile_mode (Union[str, TileMode]): The Phi-style weight tiling mode to use
527
528 Returns:
529 nn.Embedding: The new embedding layer with tiled weights
530 """
531 with torch.no_grad():
532 # Ensure vocabulary size remains the same
533 if pretrained_embedding.num_embeddings != new_embedding.num_embeddings:
534 raise ValueError("Vocabulary size (num_embeddings) must remain constant")
535
536 # Ensure new embedding dimension is larger
537 if new_embedding.embedding_dim <= pretrained_embedding.embedding_dim:
538 raise ValueError("New embedding_dim must be larger than the old embedding_dim")
539
540 # Tile the weights
541 new_embedding.weight.data = nn.Parameter(
542 tile_weight(pretrained_embedding.weight, new_embedding.weight, mode=mode),
543 requires_grad=new_embedding.weight.requires_grad,
544 )
545
546 # Handle padding_idx if it exists
547 if pretrained_embedding.padding_idx is not None:
548 if new_embedding.padding_idx is None:
549 new_embedding.padding_idx = pretrained_embedding.padding_idx
550 else:
551 assert new_embedding.padding_idx == pretrained_embedding.padding_idx, "padding_idx must remain the same"

Callers 1

Calls 1

tile_weightFunction · 0.85

Tested by

no test coverage detected