| 165 | share_embedding: Set to True to share input and output embedding weights |
| 166 | """ |
| 167 | def __init__(self, |
| 168 | vocab_size: int, |
| 169 | patch_size: Union[int, Tuple[int,int]] = 16, |
| 170 | dim_tokens: Optional[int] = None, |
| 171 | sincos_pos_emb: bool = True, |
| 172 | image_size: Union[int, Tuple[int]] = 224, |
| 173 | share_embedding: bool = True, |
| 174 | **kwargs): |
| 175 | super().__init__() |
| 176 | self.vocab_size = vocab_size |
| 177 | self.patch_size = pair(patch_size) |
| 178 | self.dim_tokens = dim_tokens |
| 179 | self.sincos_pos_emb = sincos_pos_emb |
| 180 | self.image_size = pair(image_size) |
| 181 | self.num_patches = (self.image_size[0] // self.patch_size[0]) * (self.image_size[1] // self.patch_size[1]) |
| 182 | self.share_embedding = share_embedding |
| 183 | |
| 184 | if self.dim_tokens is not None: |
| 185 | self.init(dim_tokens=dim_tokens) |
| 186 | |
| 187 | def init(self, dim_tokens: int = 768, init_std=0.02): |
| 188 | """ |