Coverage for hopwise/utils/utils.py: 80%
242 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/7/17
2# @Author : Shanlei Mu
3# @Email : slmu@ruc.edu.cn
5# UPDATE
6# @Time : 2021/3/8, 2022/7/12, 2023/2/11
7# @Author : Jiawei Guan, Lei Wang, Gaowei Zhang
8# @Email : guanjw@ruc.edu.cn, zxcptss@gmail.com, zgw2022101006@ruc.edu.cn
10# UPDATE
11# @Time : 2025
12# @Author : Alessandro Soccol
13# @Email : alessandro.soccol@unica.it
15"""hopwise.utils.utils
16################################
17"""
19import copy
20import datetime
21import importlib
22import os
23import random
24from dataclasses import dataclass
26import numpy as np
27import pandas as pd
28import torch
29from texttable import Texttable
30from torch import nn
31from torch.utils.tensorboard import SummaryWriter
33from hopwise.utils.enum_type import ModelType
36def get_local_time():
37 r"""Get current time
39 Returns:
40 str: current time
41 """
42 cur = datetime.datetime.now()
43 cur = cur.strftime("%b-%d-%Y_%H-%M-%S")
45 return cur
48def ensure_dir(dir_path):
49 r"""Make sure the directory exists, if it does not exist, create it
51 Args:
52 dir_path (str): directory path
54 """
55 os.makedirs(dir_path, exist_ok=True)
58def deep_dict_update(updated_dict, updating_dict):
59 overwrite_keys = ["split"]
60 if isinstance(updated_dict, dict) and isinstance(updating_dict, dict):
61 for key, value in updating_dict.items():
62 if isinstance(value, dict) and isinstance(updated_dict.get(key), dict) and key not in overwrite_keys:
63 deep_dict_update(updated_dict[key], value)
64 else:
65 updated_dict[key] = value
68def get_model(model_name):
69 r"""Automatically select model class based on model name
71 Args:
72 model_name (str): model name
74 Returns:
75 Recommender: model class
76 """
77 model_submodule = [
78 "general_recommender",
79 "context_aware_recommender",
80 "sequential_recommender",
81 "knowledge_aware_recommender",
82 "knowledge_graph_embedding_recommender",
83 "path_language_modeling_recommender",
84 "exlib_recommender",
85 ]
87 model_file_name = model_name.lower()
88 model_module = None
89 for submodule in model_submodule:
90 module_path = ".".join(["hopwise.model", submodule, model_file_name])
91 if importlib.util.find_spec(module_path, __name__):
92 model_module = importlib.import_module(module_path, __name__)
93 break
95 if model_module is None:
96 raise ValueError(f"`model_name` [{model_name}] is not the name of an existing model.")
97 model_class = getattr(model_module, model_name)
98 return model_class
101def get_trainer(model_type, model_name):
102 r"""Automatically select trainer class based on model type and model name
104 Args:
105 model_type (ModelType): model type
106 model_name (str): model name
108 Returns:
109 Trainer: trainer class
110 """
111 register_table = {
112 "PEARLMGPT2": "PEARLMfromscratchTrainer",
113 "PEARLMLlama2": "PEARLMfromscratchTrainer",
114 "PEARLMLlama3": "PEARLMfromscratchTrainer",
115 }
117 try:
118 if model_name.startswith("User"):
119 real_model_name = model_name.replace("User", "")
120 else:
121 real_model_name = model_name
123 return getattr(importlib.import_module("hopwise.trainer"), real_model_name + "Trainer")
124 except AttributeError:
125 if model_name in register_table:
126 return getattr(importlib.import_module("hopwise.trainer"), register_table[model_name])
127 elif model_type == ModelType.PATH_LANGUAGE_MODELING:
128 return getattr(importlib.import_module("hopwise.trainer"), "HFPathLanguageModelingTrainer")
129 elif model_type == ModelType.KNOWLEDGE:
130 return getattr(importlib.import_module("hopwise.trainer"), "KGTrainer")
131 elif model_type == ModelType.TRADITIONAL:
132 return getattr(importlib.import_module("hopwise.trainer"), "TraditionalTrainer")
133 else:
134 return getattr(importlib.import_module("hopwise.trainer"), "Trainer")
137def early_stopping(value, best, cur_step, max_step, bigger=True):
138 r"""validation-based early stopping
140 Args:
141 value (float): current result
142 best (float): best result
143 cur_step (int): the number of consecutive steps that did not exceed the best result
144 max_step (int): threshold steps for stopping
145 bigger (bool, optional): whether the bigger the better
147 Returns:
148 tuple:
149 - float,
150 best result after this step
151 - int,
152 the number of consecutive steps that did not exceed the best result after this step
153 - bool,
154 whether to stop
155 - bool,
156 whether to update
157 """
158 stop_flag = False
159 update_flag = False
160 if bigger:
161 if value >= best:
162 cur_step = 0
163 best = value
164 update_flag = True
165 else:
166 cur_step += 1
167 if cur_step > max_step:
168 stop_flag = True
169 elif value <= best:
170 cur_step = 0
171 best = value
172 update_flag = True
173 else:
174 cur_step += 1
175 if cur_step > max_step:
176 stop_flag = True
177 return best, cur_step, stop_flag, update_flag
180def calculate_valid_score(valid_result, valid_metric=None):
181 r"""Return valid score from valid result
183 Args:
184 valid_result (dict): valid result
185 valid_metric (str, optional): the selected metric in valid result for valid score
187 Returns:
188 float: valid score
189 """
190 if valid_metric:
191 return valid_result[valid_metric]
192 else:
193 return valid_result["Recall@10"]
196def dict2str(result_dict):
197 r"""Convert result dict to str
199 Args:
200 result_dict (dict): result dict
202 Returns:
203 str: result str
204 """
205 return " ".join([str(metric) + " : " + str(value) for metric, value in result_dict.items()])
208def init_seed(seed, reproducibility):
209 r"""Init random seed for random functions in numpy, torch, cuda and cudnn
211 Args:
212 seed (int): random seed
213 reproducibility (bool): Whether to require reproducibility
214 """
215 random.seed(seed)
216 np.random.seed(seed)
217 torch.manual_seed(seed)
218 torch.cuda.manual_seed(seed)
219 torch.cuda.manual_seed_all(seed)
220 if reproducibility:
221 torch.backends.cudnn.benchmark = False
222 torch.backends.cudnn.deterministic = True
223 else:
224 torch.backends.cudnn.benchmark = True
225 torch.backends.cudnn.deterministic = False
228def get_tensorboard(logger):
229 r"""Creates a SummaryWriter of Tensorboard that can log PyTorch models and metrics into a directory for
230 visualization within the TensorBoard UI.
231 For the convenience of the user, the naming rule of the SummaryWriter's log_dir is the same as the logger.
233 Args:
234 logger: its output filename is used to name the SummaryWriter's log_dir.
235 If the filename is not available, we will name the log_dir according to the current time.
237 Returns:
238 SummaryWriter: it will write out events and summaries to the event file.
239 """
240 base_path = "log_tensorboard"
242 dir_name = None
243 for handler in logger.handlers:
244 if hasattr(handler, "baseFilename"):
245 dir_name = os.path.basename(getattr(handler, "baseFilename")).split(".")[0]
246 break
247 if dir_name is None:
248 dir_name = "{}-{}".format("model", get_local_time())
250 dir_path = os.path.join(base_path, dir_name)
251 writer = SummaryWriter(dir_path)
252 return writer
255def get_gpu_usage(device=None):
256 r"""Return the reserved memory and total memory of given device in a string.
258 Args:
259 device: cuda.device. It is the device that the model run on.
261 Returns:
262 str: it contains the info about reserved memory and total memory of given device.
263 """
264 reserved = torch.cuda.max_memory_reserved(device) / 1024**3
265 total = torch.cuda.get_device_properties(device).total_memory / 1024**3
267 return f"{reserved:.2f} G/{total:.2f} G"
270def get_flops(model, dataset, device, logger, transform, verbose=False):
271 r"""Given a model and dataset to the model, compute the per-operator flops
272 of the given model.
274 Args:
275 model: the model to compute flop counts.
276 dataset: dataset that are passed to `model` to count flops.
277 device: cuda.device. It is the device that the model run on.
278 verbose: whether to print information of modules.
280 Returns:
281 total_ops: the number of flops for each operation.
282 """
283 if model.type == ModelType.DECISIONTREE:
284 return 1
285 if model.__class__.__name__ == "Pop":
286 return 1
288 model = copy.deepcopy(model)
290 def count_normalization(m, x, y):
291 x = x[0]
292 flops = torch.DoubleTensor([2 * x.numel()])
293 m.total_ops += flops
295 def count_embedding(m, x, y):
296 x = x[0]
297 nelements = x.numel()
298 hiddensize = y.shape[-1]
299 m.total_ops += nelements * hiddensize
301 class TracingAdapter(torch.nn.Module):
302 def __init__(self, rec_model):
303 super().__init__()
304 self.model = rec_model
306 def forward(self, interaction):
307 return self.model.predict(interaction)
309 custom_ops = {
310 torch.nn.Embedding: count_embedding,
311 torch.nn.LayerNorm: count_normalization,
312 }
313 wrapper = TracingAdapter(model)
314 inter = dataset[torch.tensor([1])].to(device)
315 inter = transform(dataset, inter)
316 inputs = (inter,)
317 from thop.profile import register_hooks
318 from thop.vision.basic_hooks import count_parameters
320 handler_collection = {}
321 fn_handles = []
322 params_handles = []
323 types_collection = set()
324 if custom_ops is None:
325 custom_ops = {}
327 def add_hooks(m: nn.Module):
328 m.register_buffer("total_ops", torch.zeros(1, dtype=torch.float64))
329 m.register_buffer("total_params", torch.zeros(1, dtype=torch.float64))
331 m_type = type(m)
333 fn = None
334 if m_type in custom_ops:
335 fn = custom_ops[m_type]
336 if m_type not in types_collection and verbose:
337 logger.info("Customize rule %s() %s." % (fn.__qualname__, m_type))
338 elif m_type in register_hooks:
339 fn = register_hooks[m_type]
340 if m_type not in types_collection and verbose:
341 logger.info("Register %s() for %s." % (fn.__qualname__, m_type))
342 elif m_type not in types_collection and verbose:
343 logger.warning("[WARN] Cannot find rule for %s. Treat it as zero Macs and zero Params." % m_type)
345 if fn is not None:
346 handle_fn = m.register_forward_hook(fn)
347 handle_paras = m.register_forward_hook(count_parameters)
348 handler_collection[m] = (
349 handle_fn,
350 handle_paras,
351 )
352 fn_handles.append(handle_fn)
353 params_handles.append(handle_paras)
354 types_collection.add(m_type)
356 prev_training_status = wrapper.training
358 wrapper.eval()
359 wrapper.apply(add_hooks)
361 with torch.no_grad():
362 wrapper(*inputs)
364 def dfs_count(module: nn.Module, prefix="\t"):
365 total_ops, total_params = module.total_ops.item(), 0
366 ret_dict = {}
367 for n, m in module.named_children():
368 next_dict = {}
369 if m in handler_collection and not isinstance(m, (nn.Sequential, nn.ModuleList)):
370 m_ops, m_params = m.total_ops.item(), m.total_params.item()
371 else:
372 m_ops, m_params, next_dict = dfs_count(m, prefix=prefix + "\t")
373 ret_dict[n] = (m_ops, m_params, next_dict)
374 total_ops += m_ops
375 total_params += m_params
377 return total_ops, total_params, ret_dict
379 total_ops, total_params, ret_dict = dfs_count(wrapper)
381 # reset wrapper to original status
382 wrapper.train(prev_training_status)
383 for m, (op_handler, params_handler) in handler_collection.items():
384 m._buffers.pop("total_ops")
385 m._buffers.pop("total_params")
386 for i in range(len(fn_handles)):
387 fn_handles[i].remove()
388 params_handles[i].remove()
390 return total_ops
393def list_to_latex(convert_list, bigger_flag=True, subset_columns=[]):
394 result = {}
395 for d in convert_list:
396 for key, value in d.items():
397 if key in result:
398 result[key].append(value)
399 else:
400 result[key] = [value]
402 df_result = pd.DataFrame.from_dict(result, orient="index").T
404 if len(subset_columns) == 0:
405 tex = df_result.to_latex(index=False)
406 return df_result, tex
408 def bold_func(x, bigger_flag):
409 if bigger_flag:
410 return np.where(x == np.max(x.to_numpy()), "font-weight:bold", None)
411 else:
412 return np.where(x == np.min(x.to_numpy()), "font-weight:bold", None)
414 style = df_result.style
415 style.apply(bold_func, bigger_flag=bigger_flag, subset=subset_columns)
416 style.format(precision=4)
418 num_column = len(df_result.columns)
419 column_format = "c" * num_column
420 tex = style.hide(axis="index").to_latex(
421 caption="Result Table",
422 label="Result Table",
423 convert_css=True,
424 hrules=True,
425 column_format=column_format,
426 )
428 return df_result, tex
431def get_environment(config):
432 gpu_usage = get_gpu_usage(config["device"]) if torch.cuda.is_available() and config["use_gpu"] else "0.0 / 0.0"
434 import psutil
436 memory_used = psutil.Process(os.getpid()).memory_info().rss / 1024**3
437 memory_total = psutil.virtual_memory()[0] / 1024**3
438 memory_usage = f"{memory_used:.2f} G/{memory_total:.2f} G"
439 cpu_usage = f"{psutil.cpu_percent(interval=1):.2f} %"
440 """environment_data = [
441 {"Environment": "CPU", "Usage": cpu_usage,},
442 {"Environment": "GPU", "Usage": gpu_usage, },
443 {"Environment": "Memory", "Usage": memory_usage, },
444 ]"""
446 table = Texttable()
447 table.set_cols_align(["l", "c"])
448 table.set_cols_valign(["m", "m"])
449 table.add_rows(
450 [
451 ["Environment", "Usage"],
452 ["CPU", cpu_usage],
453 ["GPU", gpu_usage],
454 ["Memory", memory_usage],
455 ]
456 )
458 return table
461def get_sequence_postprocessor(postprocessor_name):
462 try:
463 postprocessor_name = postprocessor_name + "SequenceScorePostProcessor"
464 return getattr(importlib.import_module("hopwise.model.sequence_postprocessor"), postprocessor_name)
465 except AttributeError:
466 return getattr(
467 importlib.import_module("hopwise.model.sequence_postprocessor"), "BeamSearchSequenceScorePostProcessor"
468 )
471def get_logits_processor(model_name):
472 try:
473 return getattr(
474 importlib.import_module("hopwise.model.logits_processor"), model_name + "LogitsProcessorWordLevel"
475 )
476 except AttributeError:
477 return getattr(
478 importlib.import_module("hopwise.model.logits_processor"), "ConstrainedLogitsProcessorWordLevel"
479 )
482@dataclass
483class GenerationOutputs(dict):
484 r"""Dataclass to hold the outputs of the generation process.
486 Attributes:
487 sequences (torch.Tensor): The generated sequences.
488 scores (torch.Tensor): The scores for each generated token.
489 """
491 sequences: torch.Tensor
492 scores: torch.Tensor
494 def __post_init__(self):
495 self.update({"sequences": self.sequences, "scores": self.scores})