MCPcopy Create free account
hub / github.com/DeepGraphLearning/DiffPack / build_downstream_solver

Function build_downstream_solver

util.py:97–128  ·  view source on GitHub ↗
(cfg, dataset)

Source from the content-addressed store, hash-verified

95
96
97def 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
131def build_pretrain_solver(cfg, dataset):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected