(self)
| 986 | precision: Optional[Union[jax.lax.Precision, str]]=None |
| 987 | |
| 988 | def setup(self): |
| 989 | self.embed_dim = self.config.hidden_size |
| 990 | |
| 991 | self.wte = nn.Embed( |
| 992 | self.config.vocab_size, |
| 993 | self.config.hidden_size, |
| 994 | embedding_init=jax.nn.initializers.normal(stddev=self.config.initializer_range), |
| 995 | dtype=self.dtype, |
| 996 | param_dtype=self.param_dtype, |
| 997 | ) |
| 998 | self.dropout = nn.Dropout(rate=self.config.embd_pdrop) |
| 999 | self.h = FlaxLLaMABlockCollection(self.config, dtype=self.dtype, param_dtype=self.param_dtype, precision=self.precision) |
| 1000 | self.ln_f = RMSNorm(self.config.hidden_size, eps=self.config.rms_norm_eps, dtype=self.dtype, param_dtype=self.param_dtype) |
| 1001 | |
| 1002 | def __call__( |
| 1003 | self, |
nothing calls this directly
no test coverage detected