(cfg: DictConfig, return_trainer: bool = False, do_train: bool = True)
| 241 | |
| 242 | |
| 243 | def train(cfg: DictConfig, return_trainer: bool = False, do_train: bool = True) -> Optional[Trainer]: |
| 244 | print("Training using config: ") |
| 245 | print(om.to_yaml(cfg)) |
| 246 | reproducibility.seed_all(cfg.seed) |
| 247 | |
| 248 | # Get batch size info |
| 249 | cfg = update_batch_size_info(cfg) |
| 250 | |
| 251 | # Build Model |
| 252 | print("Initializing model...") |
| 253 | model = build_model(cfg.model) |
| 254 | n_params = sum(p.numel() for p in model.parameters()) |
| 255 | print(f"{n_params=:.4e}") |
| 256 | |
| 257 | # Dataloaders |
| 258 | print("Building train loader...") |
| 259 | train_loader = build_my_dataloader( |
| 260 | cfg.train_loader, |
| 261 | cfg.global_train_batch_size // dist.get_world_size(), |
| 262 | ) |
| 263 | print("Building eval loader...") |
| 264 | global_eval_batch_size = cfg.get("global_eval_batch_size", cfg.global_train_batch_size) |
| 265 | eval_loader = build_my_dataloader( |
| 266 | cfg.eval_loader, |
| 267 | cfg.get("device_eval_batch_size", global_eval_batch_size // dist.get_world_size()), |
| 268 | ) |
| 269 | eval_evaluator = Evaluator( |
| 270 | label="eval", |
| 271 | dataloader=eval_loader, |
| 272 | device_eval_microbatch_size=cfg.get("device_eval_microbatch_size", None), |
| 273 | ) |
| 274 | |
| 275 | # Optimizer |
| 276 | optimizer = build_optimizer(cfg.optimizer, model) |
| 277 | |
| 278 | # Scheduler |
| 279 | scheduler = build_scheduler(cfg.scheduler) |
| 280 | |
| 281 | # Loggers |
| 282 | loggers = [build_logger(name, logger_cfg) for name, logger_cfg in cfg.get("loggers", {}).items()] |
| 283 | |
| 284 | # Callbacks |
| 285 | callbacks = [build_callback(name, callback_cfg) for name, callback_cfg in cfg.get("callbacks", {}).items()] |
| 286 | |
| 287 | # Algorithms |
| 288 | algorithms = [build_algorithm(name, algorithm_cfg) for name, algorithm_cfg in cfg.get("algorithms", {}).items()] |
| 289 | |
| 290 | if cfg.get("run_name") is None: |
| 291 | cfg.run_name = os.environ.get("COMPOSER_RUN_NAME", "sequence-classification") |
| 292 | |
| 293 | # Build the Trainer |
| 294 | trainer = Trainer( |
| 295 | run_name=cfg.run_name, |
| 296 | seed=cfg.seed, |
| 297 | model=model, |
| 298 | algorithms=algorithms, |
| 299 | train_dataloader=train_loader, |
| 300 | eval_dataloader=eval_evaluator, |
no test coverage detected