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)
| 59 | |
| 60 | |
| 61 | def 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 | |
| 104 | def _eval_during_train(model, data, epoch, output_root): |