(
self,
config: MixtralConfig,
parallel_config: ParallelConfig,
attention_backend: str,
linear_method: Optional[LinearMethodBase] = None,
)
| 303 | class MixtralModel(nn.Module): |
| 304 | |
| 305 | def __init__( |
| 306 | self, |
| 307 | config: MixtralConfig, |
| 308 | parallel_config: ParallelConfig, |
| 309 | attention_backend: str, |
| 310 | linear_method: Optional[LinearMethodBase] = None, |
| 311 | ) -> None: |
| 312 | super().__init__() |
| 313 | self.padding_idx = config.pad_token_id |
| 314 | self.vocab_size = config.vocab_size |
| 315 | self.parallel_config = parallel_config |
| 316 | self.embed_tokens = VocabParallelEmbedding( |
| 317 | config.vocab_size, |
| 318 | config.hidden_size, |
| 319 | ) |
| 320 | self.layers = nn.ModuleList() |
| 321 | for i in range(self.parallel_config.start, |
| 322 | self.parallel_config.end + 1): |
| 323 | self.layers.add_module( |
| 324 | f"{i}", |
| 325 | MixtralDecoderLayer(config, |
| 326 | attention_backend, |
| 327 | linear_method=linear_method)) |
| 328 | |
| 329 | if self.parallel_config.is_last: |
| 330 | self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) |
| 331 | |
| 332 | def forward( |
| 333 | self, |
nothing calls this directly
no test coverage detected