MCPcopy Create free account
hub / github.com/DeepGraphLearning/S3F / build_solver

Function build_solver

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

Source from the content-addressed store, hash-verified

84
85
86def 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
116def parse_args():

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected