Create trainer.
(params: config_definitions.ExperimentConfig,
task: base_task.Task,
train: bool,
evaluate: bool,
checkpoint_exporter: Optional[BestCheckpointExporter] = None,
trainer_cls=base_trainer.Trainer)
| 265 | |
| 266 | @gin.configurable |
| 267 | def create_trainer(params: config_definitions.ExperimentConfig, |
| 268 | task: base_task.Task, |
| 269 | train: bool, |
| 270 | evaluate: bool, |
| 271 | checkpoint_exporter: Optional[BestCheckpointExporter] = None, |
| 272 | trainer_cls=base_trainer.Trainer) -> base_trainer.Trainer: |
| 273 | """Create trainer.""" |
| 274 | logging.info('Running default trainer.') |
| 275 | model = task.build_model() |
| 276 | optimizer = create_optimizer(task, params) |
| 277 | return trainer_cls( |
| 278 | params, |
| 279 | task, |
| 280 | model=model, |
| 281 | optimizer=optimizer, |
| 282 | train=train, |
| 283 | evaluate=evaluate, |
| 284 | checkpoint_exporter=checkpoint_exporter) |
| 285 | |
| 286 | |
| 287 | @dataclasses.dataclass |
nothing calls this directly
no test coverage detected