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

Class MossForCausalLM

models_jittor/model.py:347–397  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

345
346
347class 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 )

Callers 1

Calls

no outgoing calls

Tested by

no test coverage detected