MCPcopy Create free account
hub / github.com/THUDM/GLM / seq2seq_forward_step

Function seq2seq_forward_step

tasks/seq2seq/finetune.py:32–53  ·  view source on GitHub ↗

Forward step.

(data, model, args, timers, mems)

Source from the content-addressed store, hash-verified

30
31
32def seq2seq_forward_step(data, model, args, timers, mems):
33 """Forward step."""
34
35 # Get the batch.
36 if timers is not None:
37 timers('batch generator').start()
38 tokens, labels, loss_mask, attention_mask, position_ids = get_batch(data, args)
39 if timers is not None:
40 timers('batch generator').stop()
41 # Forward model.
42 logits, *mems = model(tokens, position_ids, attention_mask, *mems)
43 # logits, loss_mask = logits[:, args.src_seq_length:], loss_mask[:, args.src_seq_length:]
44 # target_ids = target_ids[:, args.src_seq_length:]
45 losses = mpu.vocab_parallel_cross_entropy(logits.contiguous().float(), labels)
46 if args.label_smoothing > 0.0:
47 epsilon = args.label_smoothing
48 smooth_loss = -torch.nn.functional.log_softmax(logits, dim=-1).mean(dim=-1)
49 losses = (1 - epsilon) * losses + epsilon * smooth_loss
50 loss_mask = loss_mask.reshape(-1)
51 # The loss is not normalized for fair comparison
52 loss = torch.sum(losses.reshape(-1) * loss_mask) / loss_mask.sum()
53 return loss, mems, 'bert'
54
55
56def train_valid_datasets_provider(args, tokenizer):

Callers

nothing calls this directly

Calls 3

get_batchFunction · 0.90
startMethod · 0.80
stopMethod · 0.80

Tested by

no test coverage detected