(
self,
input_ids: Optional[jt.Var] = None,
past_key_values: Optional[Tuple[Tuple[jt.Var]]] = None,
attention_mask: Optional[jt.Var] = None,
token_type_ids: Optional[jt.Var] = None,
position_ids: Optional[jt.Var] = None,
head_mask: Optional[jt.Var] = None,
inputs_embeds: Optional[jt.Var] = None,
labels: Optional[jt.Var] = None,
use_cache: Optional[bool] = None,
)
| 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 | ) |
nothing calls this directly
no outgoing calls
no test coverage detected