Coverage for hopwise/quick_start/quick_start.py: 55%

136 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-09-30 13:25 +0000

1# @Time : 2020/10/6, 2022/7/18 

2# @Author : Shanlei Mu, Lei Wang 

3# @Email : slmu@ruc.edu.cn, zxcptss@gmail.com 

4 

5# UPDATE: 

6# @Time : 2022/7/8, 2022/07/10, 2022/07/13, 2023/2/11 

7# @Author : Zhen Tian, Junjie Zhang, Gaowei Zhang 

8# @Email : chenyuwuxinn@gmail.com, zjj001128@163.com, zgw15630559577@163.com 

9 

10# @Time : 2025 

11# @Author : Alessandro Soccol, Giacomo Medda 

12# @Email : alessandro.soccol@unica.it, giacomo.medda@unica.it 

13 

14"""hopwise.quick_start 

15######################## 

16""" 

17 

18import logging 

19import sys 

20from collections.abc import MutableMapping 

21from logging import getLogger 

22 

23import torch 

24import torch.distributed as dist 

25 

26from hopwise.config import Config 

27from hopwise.data import construct_transform, create_dataset, data_preparation 

28from hopwise.trainer import HFPathLanguageModelingTrainer 

29from hopwise.utils import ( 

30 KnowledgeEvaluationType, 

31 ModelType, 

32 calculate_valid_score, 

33 deep_dict_update, 

34 get_environment, 

35 get_flops, 

36 get_model, 

37 get_trainer, 

38 init_logger, 

39 init_seed, 

40 set_color, 

41) 

42 

43 

44def run( 

45 model, 

46 dataset, 

47 run="train", 

48 checkpoint=None, 

49 config_file_list=None, 

50 config_dict=None, 

51 saved=True, 

52 nproc=1, 

53 world_size=-1, 

54 ip="localhost", 

55 port="5678", 

56 group_offset=0, 

57): 

58 if nproc == 1 and world_size <= 0: 

59 res = run_hopwise( 

60 model=model, 

61 dataset=dataset, 

62 run=run, 

63 checkpoint=checkpoint, 

64 config_file_list=config_file_list, 

65 config_dict=config_dict, 

66 saved=saved, 

67 ) 

68 else: 

69 if world_size == -1: 

70 world_size = nproc 

71 import torch.multiprocessing as mp 

72 

73 # Refer to https://discuss.pytorch.org/t/problems-with-torch-multiprocess-spawn-and-simplequeue/69674/2 

74 # https://discuss.pytorch.org/t/return-from-mp-spawn/94302/2 

75 queue = mp.get_context("spawn").SimpleQueue() 

76 

77 config_dict = config_dict or {} 

78 config_dict.update( 

79 { 

80 "world_size": world_size, 

81 "ip": ip, 

82 "port": port, 

83 "nproc": nproc, 

84 "offset": group_offset, 

85 } 

86 ) 

87 kwargs = { 

88 "config_dict": config_dict, 

89 "queue": queue, 

90 } 

91 

92 mp.spawn( 

93 run_hopwises, 

94 args=(model, dataset, run, checkpoint, config_file_list, kwargs), 

95 nprocs=nproc, 

96 join=True, 

97 ) 

98 

99 # Normally, there should be only one item in the queue 

100 res = None if queue.empty() else queue.get() 

101 return res 

102 

103 

104def run_hopwise( 

105 model=None, 

106 dataset=None, 

107 run="train", 

108 checkpoint=None, 

109 config_file_list=None, 

110 config_dict=None, 

111 saved=True, 

112 queue=None, 

113): 

114 r"""A fast running api, which includes the complete process of 

115 training and testing a model on a specified dataset 

116 

117 Args: 

118 model (str, optional): Model name. Defaults to ``None``. 

119 dataset (str, optional): Dataset name. Defaults to ``None``. 

120 run (str, optional): The running mode, 'train' or 'evaluate'. Defaults to ``'train'``. 

121 checkpoint (str, optional): The path of the saved model file. Defaults to ``None``. 

122 config_file_list (list, optional): Config files used to modify experiment parameters. Defaults to ``None``. 

123 config_dict (dict, optional): Parameters dictionary used to modify experiment parameters. Defaults to ``None``. 

124 saved (bool, optional): Whether to save the model. Defaults to ``True``. 

125 queue (torch.multiprocessing.Queue, optional): The queue used to pass the result to the main process. Defaults to ``None``. 

126 """ # noqa: E501 

127 

128 # Initialize configuration 

129 config = Config( 

130 model=model, 

131 dataset=dataset, 

132 config_file_list=config_file_list, 

133 config_dict=config_dict, 

134 ) 

135 

136 if checkpoint is not None: 

137 config, model, dataset, train_data, valid_data, test_data = load_data_and_model( 

138 model_file=checkpoint, updating_config=config 

139 ) 

140 

141 logger = get_logger(config) 

142 logger.info(set_color(f"A checkpoint is provided from which to resume training {checkpoint}", "red")) 

143 else: 

144 logger = get_logger(config) 

145 # dataset filtering 

146 dataset = create_dataset(config) 

147 logger.info(dataset) 

148 

149 # dataset splitting 

150 train_data, valid_data, test_data = data_preparation(config, dataset) 

151 

152 # model loading and initialization 

153 init_seed(config["seed"] + config["local_rank"], config["reproducibility"]) 

154 model = get_model(config["model"])(config, train_data.dataset) 

155 if isinstance(model, torch.nn.Module): 

156 model = model.to(device=config["device"], dtype=config["weight_precision"]) 

157 

158 logger.info(model) 

159 

160 transform = construct_transform(config) 

161 

162 flops = get_flops(model, dataset, config["device"], logger, transform) 

163 logger.info(set_color("FLOPs", "blue") + f": {flops}") 

164 

165 # trainer loading and initialization 

166 trainer = get_trainer(config["MODEL_TYPE"], config["model"])(config, model) 

167 if run == "train": 

168 if checkpoint is not None: 

169 trainer.resume_checkpoint(checkpoint) 

170 

171 best_valid_score, best_valid_result = trainer.fit( 

172 train_data, valid_data, saved=saved, show_progress=config["show_progress"] 

173 ) 

174 

175 elif run == "evaluate": 

176 if checkpoint is None: 

177 raise ValueError("Checkpoint is needed for evaluation") 

178 trainer.eval_collector.train_data_collect(train_data) 

179 

180 if isinstance(trainer, HFPathLanguageModelingTrainer): 

181 trainer.init_hf_trainer(train_data, valid_data, show_progress=config["show_progress"]) 

182 trainer.resume_checkpoint(checkpoint) 

183 

184 best_valid_result = trainer.evaluate( 

185 valid_data, load_best_model=False, model_file=checkpoint, show_progress=config["show_progress"] 

186 ) 

187 

188 best_valid_score = calculate_valid_score(best_valid_result, trainer.valid_metric) 

189 else: 

190 raise ValueError(f"Invalid run mode: {run}") 

191 

192 if best_valid_result is not None: 

193 if KnowledgeEvaluationType.REC in best_valid_result or KnowledgeEvaluationType.LP in best_valid_result: 

194 for task, result in best_valid_result.items(): 

195 logger.info(set_color(f"[{task}] best valid ", "yellow") + f": {format_metrics(result)}") 

196 else: 

197 logger.info(set_color("best valid result", "yellow") + f": {format_metrics(best_valid_result)}") 

198 

199 # model evaluation 

200 test_result = trainer.evaluate( 

201 test_data, 

202 load_best_model=saved and run != "evaluate", 

203 model_file=checkpoint, 

204 show_progress=config["show_progress"], 

205 ) 

206 

207 environment_tb = get_environment(config) 

208 logger.info("The running environment of this training is as follows:\n" + environment_tb.draw()) 

209 

210 if test_result is not None: 

211 if KnowledgeEvaluationType.REC in test_result or KnowledgeEvaluationType.LP in test_result: 

212 for task, result in test_result.items(): 

213 logger.info(set_color(f"[{task}] test result ", "yellow") + f": {format_metrics(result)}") 

214 else: 

215 logger.info(set_color("test result", "yellow") + f": {format_metrics(test_result)}") 

216 

217 # In the case of KG-aware tasks, we don't care about the final "best_valid_score" 

218 # format because it is not used anywhere. 

219 result = { 

220 "best_valid_score": best_valid_score, 

221 "valid_score_bigger": config["valid_metric_bigger"], 

222 "best_valid_result": best_valid_result, 

223 "test_result": test_result, 

224 } 

225 

226 if not config["single_spec"]: 

227 dist.destroy_process_group() 

228 

229 if config["local_rank"] == 0 and queue is not None: 

230 queue.put(result) # for multiprocessing, e.g., mp.spawn 

231 

232 return result # for the single process 

233 

234 

235def get_logger(config): 

236 init_seed(config["seed"], config["reproducibility"]) 

237 # logger initialization 

238 init_logger(config) 

239 logger = getLogger() 

240 logger.info(sys.argv) 

241 logger.info(config) 

242 

243 return logger 

244 

245 

246def format_metrics(metrics): 

247 formatted_str = "".join([f"[{key}]: {value} " for key, value in metrics.items()]) 

248 return formatted_str 

249 

250 

251def run_hopwises(rank, *args): 

252 kwargs = args[-1] 

253 if not isinstance(kwargs, MutableMapping): 

254 raise ValueError(f"The last argument of run_hopwises should be a dict, but got {type(kwargs)}") 

255 kwargs["config_dict"] = kwargs.get("config_dict", {}) 

256 kwargs["config_dict"]["local_rank"] = rank 

257 run_hopwise( 

258 *args[:5], 

259 **kwargs, 

260 ) 

261 

262 

263def objective_function(config_dict=None, config_file_list=None, saved=True, show_progress=False, callback_fn=None): 

264 r"""The default objective_function used in HyperTuning 

265 

266 Args: 

267 config_dict (dict, optional): Parameters dictionary used to modify experiment parameters. Defaults to ``None``. 

268 config_file_list (list, optional): Config files used to modify experiment parameters. Defaults to ``None``. 

269 saved (bool, optional): Whether to save the model. Defaults to ``True``. 

270 """ 

271 

272 config = Config(config_dict=config_dict, config_file_list=config_file_list) 

273 init_seed(config["seed"], config["reproducibility"]) 

274 logger = getLogger() 

275 for hdlr in logger.handlers[:]: # remove all old handlers 

276 logger.removeHandler(hdlr) 

277 init_logger(config) 

278 logging.basicConfig(level=logging.ERROR) 

279 dataset = create_dataset(config) 

280 train_data, valid_data, test_data = data_preparation(config, dataset) 

281 init_seed(config["seed"], config["reproducibility"]) 

282 model_name = config["model"] 

283 model = get_model(model_name)(config, train_data.dataset).to(config["device"]) 

284 trainer = get_trainer(config["MODEL_TYPE"], config["model"])(config, model) 

285 best_valid_score, best_valid_result = trainer.fit( 

286 train_data, 

287 valid_data, 

288 verbose=show_progress, 

289 show_progress=show_progress, 

290 saved=saved, 

291 callback_fn=callback_fn, 

292 ) 

293 if best_valid_result is not None: 

294 if KnowledgeEvaluationType.REC in best_valid_result and KnowledgeEvaluationType.REC in best_valid_score: 

295 best_valid_score, best_valid_result = ( 

296 best_valid_score[KnowledgeEvaluationType.REC], 

297 best_valid_result[KnowledgeEvaluationType.REC], 

298 ) 

299 test_result = trainer.evaluate(test_data, load_best_model=saved) 

300 

301 return { 

302 "model": model_name, 

303 "best_valid_score": best_valid_score, 

304 "valid_score_bigger": config["valid_metric_bigger"], 

305 "best_valid_result": best_valid_result, 

306 "test_result": test_result, 

307 } 

308 

309 

310def load_data_and_model(model_file, load_only_data=False, updating_config=None): 

311 r"""Load filtered dataset, split dataloaders and saved model. 

312 

313 Args: 

314 model_file (str): The path of saved model file. 

315 load_only_data (bool, optional): Whether to load only the dataset and dataloaders without the model. 

316 Defaults to ``False``. 

317 updating_config (Config, optional): A Config object to update the config parameters loaded from checkpoint. 

318 Defaults to ``None``. 

319 

320 Returns: 

321 tuple: 

322 - config (Config): An instance object of Config, which record parameter information in :attr:`model_file`. 

323 - model (AbstractRecommender): The model load from :attr:`model_file`. 

324 - dataset (Dataset): The filtered dataset. 

325 - train_data (AbstractDataLoader): The dataloader for training. 

326 - valid_data (AbstractDataLoader): The dataloader for validation. 

327 - test_data (AbstractDataLoader): The dataloader for testing. 

328 """ 

329 checkpoint = torch.load(model_file, weights_only=False) 

330 config = checkpoint["config"] 

331 

332 if updating_config is not None: 

333 deep_dict_update(config.final_config_dict, updating_config.final_config_dict) 

334 

335 init_seed(config["seed"], config["reproducibility"]) 

336 init_logger(config) 

337 logger = getLogger() 

338 logger.info(config) 

339 

340 dataset = create_dataset(config) 

341 logger.info(dataset) 

342 train_data, valid_data, test_data = data_preparation(config, dataset) 

343 

344 init_seed(config["seed"], config["reproducibility"]) 

345 model = get_model(config["model"])(config, train_data.dataset).to(config["device"]) 

346 

347 if not load_only_data: 

348 if config["MODEL_TYPE"] == ModelType.PATH_LANGUAGE_MODELING: 

349 from transformers.modeling_utils import PreTrainedModel 

350 

351 model_class = get_model(config["model"]) 

352 if not issubclass(model_class, PreTrainedModel): 

353 model.load_state_dict(checkpoint["state_dict"]) 

354 model.load_other_parameter(checkpoint.get("other_parameter")) 

355 else: 

356 model.load_state_dict(checkpoint["state_dict"]) 

357 model.load_other_parameter(checkpoint.get("other_parameter")) 

358 return config, model, dataset, train_data, valid_data, test_data