MCPcopy Create free account
hub / github.com/baidu/DDParser / epoch_train

Function epoch_train

ddparser/parser/model.py:103–135  ·  view source on GitHub ↗

Train in one epoch

(args, model, optimizer, loader, epoch)

Source from the content-addressed store, hash-verified

101
102
103def 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

Callers 1

trainFunction · 0.90

Calls 2

loss_functionFunction · 0.85
trainMethod · 0.80

Tested by

no test coverage detected