| 245 | return args |
| 246 | |
| 247 | def main(args): |
| 248 | if args.config_args: |
| 249 | config = json.loads(_jsonnet.evaluate_file(args.config, tla_codes={'args': args.config_args})) |
| 250 | else: |
| 251 | config = json.loads(_jsonnet.evaluate_file(args.config)) |
| 252 | |
| 253 | if 'model_name' in config: |
| 254 | args.logdir = os.path.join(args.logdir, config['model_name']) |
| 255 | |
| 256 | # Initialize the logger |
| 257 | reopen_to_flush = config.get('log', {}).get('reopen_to_flush') |
| 258 | logger = Logger(os.path.join(args.logdir, 'log.txt'), reopen_to_flush) |
| 259 | |
| 260 | # Save the config info |
| 261 | with open(os.path.join(args.logdir, |
| 262 | 'config-{}.json'.format( |
| 263 | datetime.datetime.now().strftime('%Y%m%dT%H%M%S%Z'))), 'w') as f: |
| 264 | json.dump(config, f, sort_keys=True, indent=4) |
| 265 | |
| 266 | logger.log('Logging to {}'.format(args.logdir)) |
| 267 | |
| 268 | # Construct trainer and do training |
| 269 | trainer = Trainer(logger, config) |
| 270 | trainer.train(config, modeldir=args.logdir) |
| 271 | |
| 272 | if __name__ == '__main__': |
| 273 | args = add_parser() |