MCPcopy Create free account
hub / github.com/OpenMOSS/MOSS / execute

Method execute

models_jittor/model.py:358–397  ·  view source on GitHub ↗
(
        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,
    )

Source from the content-addressed store, hash-verified

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 )

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected