Train in one epoch
(args, model, optimizer, loader, epoch)
| 101 | |
| 102 | |
| 103 | def epoch_train(args, model, optimizer, loader, epoch): |
| 104 | """Train in one epoch""" |
| 105 | model.train() |
| 106 | total_loss = 0 |
| 107 | pad_index = args.pad_index |
| 108 | bos_index = args.bos_index |
| 109 | eos_index = args.eos_index |
| 110 | |
| 111 | for batch, inputs in enumerate(loader(), start=1): |
| 112 | model.clear_gradients() |
| 113 | |
| 114 | if args.encoding_model.startswith("ernie"): |
| 115 | words, arcs, rels = inputs |
| 116 | s_arc, s_rel, words = model(words) |
| 117 | else: |
| 118 | words, feats, arcs, rels = inputs |
| 119 | s_arc, s_rel, words = model(words, feats) |
| 120 | |
| 121 | mask = layers.logical_and( |
| 122 | layers.logical_and(words != pad_index, words != bos_index), |
| 123 | words != eos_index, |
| 124 | ) |
| 125 | |
| 126 | loss = loss_function(s_arc, s_rel, arcs, rels, mask) |
| 127 | loss.backward() |
| 128 | |
| 129 | optimizer.minimize(loss) |
| 130 | total_loss += loss.numpy().item() |
| 131 | logging.info("epoch: {}, batch: {}/{}, batch_size: {}, loss: {:.4f}".format( |
| 132 | epoch, batch, math.ceil(len(loader)), len(words), |
| 133 | loss.numpy().item())) |
| 134 | total_loss /= len(loader) |
| 135 | return total_loss |
| 136 | |
| 137 | |
| 138 | @dygraph.no_grad |
no test coverage detected