(cfg, dataset)
| 84 | |
| 85 | |
| 86 | def build_solver(cfg, dataset): |
| 87 | generator = torch.Generator().manual_seed(0) |
| 88 | lengths = [int(len(dataset) * cfg.split[0]), int(len(dataset) * cfg.split[1])] |
| 89 | lengths.append(len(dataset) - sum(lengths)) |
| 90 | train_set, valid_set, test_set = torch_data.random_split(dataset, lengths, generator=generator) |
| 91 | if comm.get_rank() == 0: |
| 92 | logger.warning("#train: %d, #valid: %d, #test: %d" % (len(train_set), len(valid_set), len(test_set))) |
| 93 | |
| 94 | task = core.Configurable.load_config_dict(cfg.task) |
| 95 | |
| 96 | if "fix_sequence_model" in cfg: |
| 97 | model = task.model |
| 98 | assert cfg.task.model ["class"] == "FusionNetwork" |
| 99 | for p in model.sequence_model.parameters(): |
| 100 | p.requires_grad = False |
| 101 | cfg.optimizer.params = [p for p in task.parameters() if p.requires_grad] |
| 102 | else: |
| 103 | cfg.optimizer.params = task.parameters() |
| 104 | optimizer = core.Configurable.load_config_dict(cfg.optimizer) |
| 105 | |
| 106 | solver = core.Engine(task, train_set, valid_set, test_set, optimizer, **cfg.engine) |
| 107 | |
| 108 | if cfg.get("checkpoint") is not None: |
| 109 | if comm.get_rank() == 0: |
| 110 | logger.warning("Load checkpoint from %s" % cfg.checkpoint) |
| 111 | solver.load(cfg.checkpoint) |
| 112 | |
| 113 | return solver |
| 114 | |
| 115 | |
| 116 | def parse_args(): |
nothing calls this directly
no outgoing calls
no test coverage detected