| 33 | share_embedding: Set to True to share input and output embedding weights |
| 34 | """ |
| 35 | def __init__(self, |
| 36 | vocab_size: int, |
| 37 | max_length: int, |
| 38 | dim_tokens: Optional[int] = None, |
| 39 | sincos_pos_emb: bool = True, |
| 40 | max_sincos_pos_emb: int = 512, |
| 41 | padding_idx: int = 0, |
| 42 | share_embedding: bool = True, |
| 43 | **kwargs): |
| 44 | super().__init__() |
| 45 | self.vocab_size = vocab_size |
| 46 | self.max_length = max_length |
| 47 | self.dim_tokens = dim_tokens |
| 48 | self.sincos_pos_emb = sincos_pos_emb |
| 49 | self.padding_idx = padding_idx |
| 50 | self.max_sincos_pos_emb = max_sincos_pos_emb |
| 51 | self.share_embedding = share_embedding |
| 52 | |
| 53 | if self.dim_tokens is not None: |
| 54 | self.init(dim_tokens=dim_tokens) |
| 55 | |
| 56 | def init(self, dim_tokens: int = 768, init_std=0.02): |
| 57 | """ |