(self)
| 259 | precision: Optional[Union[jax.lax.Precision, str]]=None |
| 260 | |
| 261 | def setup(self): |
| 262 | self.embed_dim = self.config.hidden_size |
| 263 | |
| 264 | self.vte = nn.Embed( |
| 265 | self.config.vision_vocab_size, |
| 266 | self.config.hidden_size, |
| 267 | embedding_init=jax.nn.initializers.normal(stddev=self.config.initializer_range), |
| 268 | dtype=self.dtype, |
| 269 | param_dtype=self.param_dtype, |
| 270 | ) |
| 271 | |
| 272 | self.wte = nn.Embed( |
| 273 | self.config.vocab_size, |
| 274 | self.config.hidden_size, |
| 275 | embedding_init=jax.nn.initializers.normal(stddev=self.config.initializer_range), |
| 276 | dtype=self.dtype, |
| 277 | param_dtype=self.param_dtype, |
| 278 | ) |
| 279 | self.dropout = nn.Dropout(rate=self.config.embd_pdrop) |
| 280 | self.h = FlaxLLaMABlockCollection(self.config, dtype=self.dtype, param_dtype=self.param_dtype, precision=self.precision) |
| 281 | self.ln_f = RMSNorm(self.config.hidden_size, eps=self.config.rms_norm_eps, dtype=self.dtype, param_dtype=self.param_dtype) |
| 282 | |
| 283 | def __call__( |
| 284 | self, |
nothing calls this directly
no test coverage detected