Initialize parts of embedding module that are dependent on dimension of tokens. Should be called when setting up FourM. Args: dim_tokens: Dimension of tokens init_std: Standard deviation of init
(self, dim_tokens: int = 768, init_std=0.02)
| 54 | self.init(dim_tokens=dim_tokens) |
| 55 | |
| 56 | def init(self, dim_tokens: int = 768, init_std=0.02): |
| 57 | """ |
| 58 | Initialize parts of embedding module that are dependent on dimension of tokens. |
| 59 | Should be called when setting up FourM. |
| 60 | |
| 61 | Args: |
| 62 | dim_tokens: Dimension of tokens |
| 63 | init_std: Standard deviation of init |
| 64 | """ |
| 65 | self.dim_tokens = dim_tokens |
| 66 | |
| 67 | # Task embedding identifying from which task a given token comes from |
| 68 | # Fixed-size positional embeddings. Can be interpolated to different input sizes |
| 69 | |
| 70 | if self.sincos_pos_emb: |
| 71 | if self.max_length > self.max_sincos_pos_emb: |
| 72 | raise ValueError(f"Max length ({self.max_length}) is greater than the number of posembs ({self.max_sincos_pos_emb}") |
| 73 | # Get all posembs, than truncate up to max length |
| 74 | pos_emb = build_1d_sincos_posemb(max_len=self.max_sincos_pos_emb, embed_dim=self.dim_tokens)[:self.max_length] |
| 75 | self.register_buffer("pos_emb", pos_emb) |
| 76 | else: |
| 77 | self.pos_emb = nn.Parameter(torch.zeros(1, self.max_length, self.dim_tokens)) |
| 78 | nn.init.normal_(self.pos_emb, std=init_std) |
| 79 | |
| 80 | self.mod_emb = nn.Parameter(torch.zeros(1, 1, self.dim_tokens)) |
| 81 | nn.init.normal_(self.mod_emb, std=init_std) |
| 82 | |
| 83 | # Token embedding |
| 84 | self.token_emb = nn.Embedding(num_embeddings=self.vocab_size, embedding_dim=self.dim_tokens, padding_idx=self.padding_idx) |
| 85 | |
| 86 | # Output projection layer |
| 87 | self.to_logits = nn.Linear(self.dim_tokens, self.vocab_size, bias=False) |
| 88 | |
| 89 | if self.share_embedding: |
| 90 | # Share input and output embedding weights |
| 91 | self.to_logits.weight = self.token_emb.weight |
| 92 | |
| 93 | |
| 94 | @torch.jit.ignore |
no test coverage detected