Args: model (TYPE): NULL optimizer (TYPE): NULL epochs (TYPE): NULL train_data (TYPE): NULL dev_data (TYPE): NULL test_data (TYPE): Default is None Returns: TODO Raises: NULL
(config, model, optimizer, epochs, train_data, dev_data, test_data=None)
| 141 | |
| 142 | |
| 143 | def train(config, model, optimizer, epochs, train_data, dev_data, test_data=None): |
| 144 | """ |
| 145 | |
| 146 | Args: |
| 147 | model (TYPE): NULL |
| 148 | optimizer (TYPE): NULL |
| 149 | epochs (TYPE): NULL |
| 150 | train_data (TYPE): NULL |
| 151 | dev_data (TYPE): NULL |
| 152 | test_data (TYPE): Default is None |
| 153 | |
| 154 | Returns: TODO |
| 155 | |
| 156 | Raises: NULL |
| 157 | """ |
| 158 | best_acc = -1e10 |
| 159 | best_epoch = 0 |
| 160 | timer = utils.Timer() |
| 161 | for epoch in range(1, epochs + 1): |
| 162 | loss = epoch_train(config, model, optimizer, epoch, train_data, config.general.is_debug) |
| 163 | cost_time = timer.interval() |
| 164 | logging.info(f'[train] epoch {epoch}/{epochs} loss is {loss:.6f}, cost {cost_time:.2f}s.') |
| 165 | |
| 166 | dev_loss, dev_acc = _eval_during_train(model, dev_data, epoch, config.data.output) |
| 167 | log_str = f'[eval] dev loss {dev_loss:.6f}, acc {dev_acc:.4f}.' |
| 168 | if test_data is not None: |
| 169 | test_loss, test_acc = _eval_during_train(model, test_data, epoch, config.data.output) |
| 170 | log_str += f' test loss {test_loss:.6f}, acc {test_acc:.4f}.' |
| 171 | |
| 172 | if dev_acc > best_acc: |
| 173 | best_acc, best_epoch = dev_acc, epoch |
| 174 | save_path = os.path.join(config.data.output, f'epoch{epoch:03d}_acc{best_acc:.4f}', 'model') |
| 175 | io.save(model, optimizer, save_path) |
| 176 | log_str += ' got best and saved.' |
| 177 | else: |
| 178 | log_str += f' best acc is {best_acc} on epoch {best_epoch}.' |
| 179 | |
| 180 | cost_time = timer.interval() |
| 181 | log_str += f' cost [{cost_time:.2f}s]' |
| 182 | logging.info(log_str) |
| 183 | |
| 184 | |
| 185 | if __name__ == "__main__": |
nothing calls this directly
no test coverage detected