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

1# @Time : 2020/7/17 

2# @Author : Shanlei Mu 

3# @Email : slmu@ruc.edu.cn 

4 

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 

9 

10# UPDATE 

11# @Time : 2025 

12# @Author : Alessandro Soccol 

13# @Email : alessandro.soccol@unica.it 

14 

15"""hopwise.utils.utils 

16################################ 

17""" 

18 

19import copy 

20import datetime 

21import importlib 

22import os 

23import random 

24from dataclasses import dataclass 

25 

26import numpy as np 

27import pandas as pd 

28import torch 

29from texttable import Texttable 

30from torch import nn 

31from torch.utils.tensorboard import SummaryWriter 

32 

33from hopwise.utils.enum_type import ModelType 

34 

35 

36def get_local_time(): 

37 r"""Get current time 

38 

39 Returns: 

40 str: current time 

41 """ 

42 cur = datetime.datetime.now() 

43 cur = cur.strftime("%b-%d-%Y_%H-%M-%S") 

44 

45 return cur 

46 

47 

48def ensure_dir(dir_path): 

49 r"""Make sure the directory exists, if it does not exist, create it 

50 

51 Args: 

52 dir_path (str): directory path 

53 

54 """ 

55 os.makedirs(dir_path, exist_ok=True) 

56 

57 

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 

66 

67 

68def get_model(model_name): 

69 r"""Automatically select model class based on model name 

70 

71 Args: 

72 model_name (str): model name 

73 

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 ] 

86 

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 

94 

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 

99 

100 

101def get_trainer(model_type, model_name): 

102 r"""Automatically select trainer class based on model type and model name 

103 

104 Args: 

105 model_type (ModelType): model type 

106 model_name (str): model name 

107 

108 Returns: 

109 Trainer: trainer class 

110 """ 

111 register_table = { 

112 "PEARLMGPT2": "PEARLMfromscratchTrainer", 

113 "PEARLMLlama2": "PEARLMfromscratchTrainer", 

114 "PEARLMLlama3": "PEARLMfromscratchTrainer", 

115 } 

116 

117 try: 

118 if model_name.startswith("User"): 

119 real_model_name = model_name.replace("User", "") 

120 else: 

121 real_model_name = model_name 

122 

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") 

135 

136 

137def early_stopping(value, best, cur_step, max_step, bigger=True): 

138 r"""validation-based early stopping 

139 

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 

146 

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 

178 

179 

180def calculate_valid_score(valid_result, valid_metric=None): 

181 r"""Return valid score from valid result 

182 

183 Args: 

184 valid_result (dict): valid result 

185 valid_metric (str, optional): the selected metric in valid result for valid score 

186 

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"] 

194 

195 

196def dict2str(result_dict): 

197 r"""Convert result dict to str 

198 

199 Args: 

200 result_dict (dict): result dict 

201 

202 Returns: 

203 str: result str 

204 """ 

205 return " ".join([str(metric) + " : " + str(value) for metric, value in result_dict.items()]) 

206 

207 

208def init_seed(seed, reproducibility): 

209 r"""Init random seed for random functions in numpy, torch, cuda and cudnn 

210 

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 

226 

227 

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. 

232 

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. 

236 

237 Returns: 

238 SummaryWriter: it will write out events and summaries to the event file. 

239 """ 

240 base_path = "log_tensorboard" 

241 

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()) 

249 

250 dir_path = os.path.join(base_path, dir_name) 

251 writer = SummaryWriter(dir_path) 

252 return writer 

253 

254 

255def get_gpu_usage(device=None): 

256 r"""Return the reserved memory and total memory of given device in a string. 

257 

258 Args: 

259 device: cuda.device. It is the device that the model run on. 

260 

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 

266 

267 return f"{reserved:.2f} G/{total:.2f} G" 

268 

269 

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. 

273 

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. 

279 

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 

287 

288 model = copy.deepcopy(model) 

289 

290 def count_normalization(m, x, y): 

291 x = x[0] 

292 flops = torch.DoubleTensor([2 * x.numel()]) 

293 m.total_ops += flops 

294 

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 

300 

301 class TracingAdapter(torch.nn.Module): 

302 def __init__(self, rec_model): 

303 super().__init__() 

304 self.model = rec_model 

305 

306 def forward(self, interaction): 

307 return self.model.predict(interaction) 

308 

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 

319 

320 handler_collection = {} 

321 fn_handles = [] 

322 params_handles = [] 

323 types_collection = set() 

324 if custom_ops is None: 

325 custom_ops = {} 

326 

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)) 

330 

331 m_type = type(m) 

332 

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) 

344 

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) 

355 

356 prev_training_status = wrapper.training 

357 

358 wrapper.eval() 

359 wrapper.apply(add_hooks) 

360 

361 with torch.no_grad(): 

362 wrapper(*inputs) 

363 

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 

376 

377 return total_ops, total_params, ret_dict 

378 

379 total_ops, total_params, ret_dict = dfs_count(wrapper) 

380 

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() 

389 

390 return total_ops 

391 

392 

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] 

401 

402 df_result = pd.DataFrame.from_dict(result, orient="index").T 

403 

404 if len(subset_columns) == 0: 

405 tex = df_result.to_latex(index=False) 

406 return df_result, tex 

407 

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) 

413 

414 style = df_result.style 

415 style.apply(bold_func, bigger_flag=bigger_flag, subset=subset_columns) 

416 style.format(precision=4) 

417 

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 ) 

427 

428 return df_result, tex 

429 

430 

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" 

433 

434 import psutil 

435 

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 ]""" 

445 

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 ) 

457 

458 return table 

459 

460 

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 ) 

469 

470 

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 ) 

480 

481 

482@dataclass 

483class GenerationOutputs(dict): 

484 r"""Dataclass to hold the outputs of the generation process. 

485 

486 Attributes: 

487 sequences (torch.Tensor): The generated sequences. 

488 scores (torch.Tensor): The scores for each generated token. 

489 """ 

490 

491 sequences: torch.Tensor 

492 scores: torch.Tensor 

493 

494 def __post_init__(self): 

495 self.update({"sequences": self.sequences, "scores": self.scores})