(cfg, dataset)
| 95 | |
| 96 | |
| 97 | def build_downstream_solver(cfg, dataset): |
| 98 | train_set, valid_set, test_set = dataset.split() |
| 99 | if comm.get_rank() == 0: |
| 100 | logger.warning(dataset) |
| 101 | logger.warning("#train: %d, #valid: %d, #test: %d" % (len(train_set), len(valid_set), len(test_set))) |
| 102 | |
| 103 | if cfg.task['class'] == 'MultipleBinaryClassification': |
| 104 | cfg.task.task = [_ for _ in range(len(dataset.tasks))] |
| 105 | else: |
| 106 | cfg.task.task = dataset.tasks |
| 107 | task = core.Configurable.load_config_dict(cfg.task) |
| 108 | if not "lr_ratio" in cfg: |
| 109 | cfg.optimizer.params = task.parameters() |
| 110 | else: |
| 111 | cfg.optimizer.params = [ |
| 112 | {'params': task.model.model.parameters(), 'lr': cfg.optimizer.lr * cfg.lr_ratio}, |
| 113 | ] |
| 114 | cfg.optimizer.params = task.parameters() |
| 115 | optimizer = core.Configurable.load_config_dict(cfg.optimizer) |
| 116 | solver = core.Engine(task, train_set, valid_set, test_set, optimizer, **cfg.engine) |
| 117 | |
| 118 | if cfg.get("checkpoint") is not None: |
| 119 | solver.load(cfg.checkpoint) |
| 120 | |
| 121 | if cfg.get("model_checkpoint") is not None: |
| 122 | if comm.get_rank() == 0: |
| 123 | logger.warning("Load checkpoint from %s" % cfg.model_checkpoint) |
| 124 | cfg.model_checkpoint = os.path.expanduser(cfg.model_checkpoint) |
| 125 | model_dict = torch.load(cfg.model_checkpoint, map_location=torch.device('cpu')) |
| 126 | task.model.load_state_dict(model_dict) |
| 127 | |
| 128 | return solver |
| 129 | |
| 130 | |
| 131 | def build_pretrain_solver(cfg, dataset): |
nothing calls this directly
no outgoing calls
no test coverage detected