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,
)
| 513 | |
| 514 | |
| 515 | def 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" |
no test coverage detected