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
« 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
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
10# @Time : 2025
11# @Author : Alessandro Soccol, Giacomo Medda
12# @Email : alessandro.soccol@unica.it, giacomo.medda@unica.it
14"""hopwise.quick_start
15########################
16"""
18import logging
19import sys
20from collections.abc import MutableMapping
21from logging import getLogger
23import torch
24import torch.distributed as dist
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)
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
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()
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 }
92 mp.spawn(
93 run_hopwises,
94 args=(model, dataset, run, checkpoint, config_file_list, kwargs),
95 nprocs=nproc,
96 join=True,
97 )
99 # Normally, there should be only one item in the queue
100 res = None if queue.empty() else queue.get()
101 return res
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
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
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 )
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 )
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)
149 # dataset splitting
150 train_data, valid_data, test_data = data_preparation(config, dataset)
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"])
158 logger.info(model)
160 transform = construct_transform(config)
162 flops = get_flops(model, dataset, config["device"], logger, transform)
163 logger.info(set_color("FLOPs", "blue") + f": {flops}")
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)
171 best_valid_score, best_valid_result = trainer.fit(
172 train_data, valid_data, saved=saved, show_progress=config["show_progress"]
173 )
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)
180 if isinstance(trainer, HFPathLanguageModelingTrainer):
181 trainer.init_hf_trainer(train_data, valid_data, show_progress=config["show_progress"])
182 trainer.resume_checkpoint(checkpoint)
184 best_valid_result = trainer.evaluate(
185 valid_data, load_best_model=False, model_file=checkpoint, show_progress=config["show_progress"]
186 )
188 best_valid_score = calculate_valid_score(best_valid_result, trainer.valid_metric)
189 else:
190 raise ValueError(f"Invalid run mode: {run}")
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)}")
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 )
207 environment_tb = get_environment(config)
208 logger.info("The running environment of this training is as follows:\n" + environment_tb.draw())
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)}")
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 }
226 if not config["single_spec"]:
227 dist.destroy_process_group()
229 if config["local_rank"] == 0 and queue is not None:
230 queue.put(result) # for multiprocessing, e.g., mp.spawn
232 return result # for the single process
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)
243 return logger
246def format_metrics(metrics):
247 formatted_str = "".join([f"[{key}]: {value} " for key, value in metrics.items()])
248 return formatted_str
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 )
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
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 """
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)
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 }
310def load_data_and_model(model_file, load_only_data=False, updating_config=None):
311 r"""Load filtered dataset, split dataloaders and saved model.
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``.
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"]
332 if updating_config is not None:
333 deep_dict_update(config.final_config_dict, updating_config.final_config_dict)
335 init_seed(config["seed"], config["reproducibility"])
336 init_logger(config)
337 logger = getLogger()
338 logger.info(config)
340 dataset = create_dataset(config)
341 logger.info(dataset)
342 train_data, valid_data, test_data = data_preparation(config, dataset)
344 init_seed(config["seed"], config["reproducibility"])
345 model = get_model(config["model"])(config, train_data.dataset).to(config["device"])
347 if not load_only_data:
348 if config["MODEL_TYPE"] == ModelType.PATH_LANGUAGE_MODELING:
349 from transformers.modeling_utils import PreTrainedModel
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