Forward step.
(data, model, args, timers, mems)
| 30 | |
| 31 | |
| 32 | def 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 | |
| 56 | def train_valid_datasets_provider(args, tokenizer): |