MCPcopy Create free account
hub / github.com/PaddlePaddle/Research / epoch_train

Function epoch_train

NLP/Text2SQL-BASELINE/text2sql/launch/trainer.py:61–101  ·  view source on GitHub ↗

train for one epoch Args: model (TYPE): NULL optimizer (TYPE): NULL epoch (TYPE): NULL train_data (TYPE): NULL Returns: TODO Raises: NULL

(config, model, optimizer, epoch, train_data, is_debug=False)

Source from the content-addressed store, hash-verified

59
60
61def epoch_train(config, model, optimizer, epoch, train_data, is_debug=False):
62 """train for one epoch
63
64 Args:
65 model (TYPE): NULL
66 optimizer (TYPE): NULL
67 epoch (TYPE): NULL
68 train_data (TYPE): NULL
69
70 Returns: TODO
71
72 Raises: NULL
73 """
74 model.train()
75
76 total_loss = 0
77 steps_loss = []
78 timer = utils.Timer()
79 batch_id= 1
80 for batch_id, (inputs, labels) in enumerate(train_data(), start=1):
81 loss = model(inputs, labels)
82
83 #if trainer_num > 1:
84 # loss = model.scale_loss(loss)
85 # loss.backward()
86 # model.apply_collective_grads()
87 #else:
88 loss.backward()
89 optimizer.step()
90 optimizer.clear_grad()
91 ## trick,这里的 _learning_rate 实际是 scheduler
92 if type(optimizer._learning_rate) is not float:
93 optimizer._learning_rate.step()
94
95 total_loss += loss.numpy().item()
96 steps_loss.append(loss.numpy().item())
97 if batch_id % config.train.log_steps == 0 or is_debug:
98 log_train_step(epoch, batch_id, steps_loss, timer.interval())
99 log_train_step(epoch, batch_id, steps_loss, timer.interval())
100
101 return total_loss / batch_id
102
103
104def _eval_during_train(model, data, epoch, output_root):

Callers 1

trainFunction · 0.85

Calls 5

intervalMethod · 0.95
log_train_stepFunction · 0.85
trainMethod · 0.45
stepMethod · 0.45
appendMethod · 0.45

Tested by

no test coverage detected