| 345 | |
| 346 | |
| 347 | class MossForCausalLM(Module): |
| 348 | |
| 349 | def __init__(self, config): |
| 350 | super(MossForCausalLM, self).__init__() |
| 351 | self.config = config |
| 352 | self.transformer = MossModel(config) |
| 353 | self.lm_head = nn.Linear(config.n_embd, config.vocab_size) |
| 354 | |
| 355 | # Initialize weights and apply final processing |
| 356 | self.apply(partial(_init_weights, config)) |
| 357 | |
| 358 | def execute( |
| 359 | self, |
| 360 | input_ids: Optional[jt.Var] = None, |
| 361 | past_key_values: Optional[Tuple[Tuple[jt.Var]]] = None, |
| 362 | attention_mask: Optional[jt.Var] = None, |
| 363 | token_type_ids: Optional[jt.Var] = None, |
| 364 | position_ids: Optional[jt.Var] = None, |
| 365 | head_mask: Optional[jt.Var] = None, |
| 366 | inputs_embeds: Optional[jt.Var] = None, |
| 367 | labels: Optional[jt.Var] = None, |
| 368 | use_cache: Optional[bool] = None, |
| 369 | ): |
| 370 | |
| 371 | hidden_states, presents = self.transformer( |
| 372 | input_ids, |
| 373 | past_key_values=past_key_values, |
| 374 | attention_mask=attention_mask, |
| 375 | token_type_ids=token_type_ids, |
| 376 | position_ids=position_ids, |
| 377 | head_mask=head_mask, |
| 378 | inputs_embeds=inputs_embeds, |
| 379 | use_cache=use_cache, |
| 380 | ) |
| 381 | |
| 382 | lm_logits = self.lm_head(hidden_states).to('float32') |
| 383 | |
| 384 | loss = None |
| 385 | if labels is not None: |
| 386 | shift_logits = lm_logits[..., :-1, :].contiguous() |
| 387 | shift_labels = labels[..., 1:].contiguous() |
| 388 | loss_fct = nn.CrossEntropyLoss() |
| 389 | loss = loss_fct(shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1)) |
| 390 | |
| 391 | loss = loss.to(hidden_states.dtype) |
| 392 | |
| 393 | return dict( |
| 394 | loss=loss, |
| 395 | logits=lm_logits, |
| 396 | past_key_values=presents |
| 397 | ) |
no outgoing calls
no test coverage detected