Coverage for hopwise/trainer/trainer.py: 71%

1102 statements  

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

1# @Time : 2020/6/26 

2# @Author : Shanlei Mu 

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

4 

5# UPDATE: 

6# @Time : 2022/7/8, 2021/6/23, 2020/9/26, 2020/9/26, 2020/10/01, 2020/9/16 

7# @Author : Zhen Tian, Zihan Lin, Yupeng Hou, Yushuo Chen, Shanlei Mu, Xingyu Pan 

8# @Email : chenyuwuxinn@gmail.com, zhlin@ruc.edu.cn, houyupeng@ruc.edu.cn, chenyushuo@ruc.edu.cn, slmu@ruc.edu.cn, panxy@ruc.edu.cn # noqa: E501 

9 

10# UPDATE: 

11# @Time : 2020/10/8, 2020/10/15, 2020/11/20, 2021/2/20, 2021/3/3, 2021/3/5, 2021/7/18, 2022/7/11, 2023/2/11 

12# @Author : Hui Wang, Xinyan Fan, Chen Yang, Yibo Li, Lanling Xu, Haoran Cheng, Zhichao Feng, Lei Wang, Gaowei Zhang 

13# @Email : hui.wang@ruc.edu.cn, xinyan.fan@ruc.edu.cn, 254170321@qq.com, 2018202152@ruc.edu.cn, xulanling_sherry@163.com, chenghaoran29@foxmail.com, fzcbupt@gmail.com, zxcptss@gmail.com, zgw2022101006@ruc.edu.cn # noqa: E501 

14 

15# UPDATE: 

16# @Time : 2025 

17# @Author : Giacomo Medda, Alessandro Soccol 

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

19 

20"""hopwise.trainer.trainer 

21################################ 

22""" 

23 

24import os 

25from collections import defaultdict 

26from logging import getLogger 

27from time import time 

28 

29import numpy as np 

30import torch 

31from scipy import sparse 

32from torch import optim 

33from torch.nn.parallel import DistributedDataParallel 

34from torch.nn.utils.clip_grad import clip_grad_norm_ 

35 

36from hopwise.data.dataloader import FullSortLPEvalDataLoader, NegSampleDataLoader 

37from hopwise.data.interaction import Interaction 

38from hopwise.evaluator import Collector, Collector_KG, Evaluator, Evaluator_KG, ExplainableCollector 

39from hopwise.utils import ( 

40 EvaluatorType, 

41 KGDataLoaderState, 

42 KnowledgeEvaluationType, 

43 WandbLogger, 

44 calculate_valid_score, 

45 dict2str, 

46 early_stopping, 

47 ensure_dir, 

48 get_gpu_usage, 

49 get_local_time, 

50 get_tensorboard, 

51 progress_bar, 

52 set_color, 

53) 

54 

55try: 

56 grad_scaler = torch.GradScaler 

57 autocast = torch.autocast 

58except AttributeError: 

59 

60 def grad_scaler(device, **kwargs): 

61 if torch.cuda.is_available(): 

62 return torch.cuda.amp.GradScaler(**kwargs) 

63 else: 

64 return torch.cuda.amp.GradScaler(**kwargs) 

65 

66 def autocast(device_type=None, **kwargs): 

67 if torch.cuda.is_available(): 

68 return torch.cuda.amp.autocast(**kwargs) 

69 else: 

70 return torch.cpu.amp.autocast(**kwargs) 

71 

72 

73class AbstractTrainer: 

74 r"""Trainer Class is used to manage the training and evaluation processes of recommender system models. 

75 AbstractTrainer is an abstract class in which the fit() and evaluate() method should be implemented according 

76 to different training and evaluation strategies. 

77 """ 

78 

79 def __init__(self, config, model): 

80 self.config = config 

81 self.model = model 

82 if not config["single_spec"]: 

83 self.model = torch.nn.SyncBatchNorm.convert_sync_batchnorm(model) 

84 self.distributed_model = DistributedDataParallel(self.model, device_ids=[config["local_rank"]]) 

85 

86 def fit(self, train_data): 

87 r"""Train the model based on the train data.""" 

88 raise NotImplementedError("Method [next] should be implemented.") 

89 

90 def evaluate(self, eval_data): 

91 r"""Evaluate the model based on the eval data.""" 

92 raise NotImplementedError("Method [next] should be implemented.") 

93 

94 def set_reduce_hook(self): 

95 r"""Call the forward function of 'distributed_model' to apply grads 

96 reduce hook to each parameter of its module. 

97 

98 """ 

99 t = self.model.forward 

100 self.model.forward = lambda x: x 

101 self.distributed_model(torch.LongTensor([0]).to(self.device)) 

102 self.model.forward = t 

103 

104 def sync_grad_loss(self): 

105 r"""Ensure that each parameter appears to the loss function to 

106 make the grads reduce sync in each node. 

107 

108 """ 

109 sync_loss = 0 

110 for params in self.model.parameters(): 

111 sync_loss += torch.sum(params) * 0 

112 return sync_loss 

113 

114 

115class Trainer(AbstractTrainer): 

116 """The basic Trainer for basic training and evaluation strategies in recommender systems. This class defines common 

117 functions for training and evaluation processes of most recommender system models, including fit(), evaluate(), 

118 resume_checkpoint() and some other features helpful for model training and evaluation. 

119 

120 Generally speaking, this class can serve most recommender system models, If the training process of the model is to 

121 simply optimize a single loss without involving any complex training strategies, such as adversarial learning, 

122 pre-training and so on. 

123 

124 Initializing the Trainer needs two parameters: `config` and `model`. `config` records the parameters information 

125 for controlling training and evaluation, such as `learning_rate`, `epochs`, `eval_step` and so on. 

126 `model` is the instantiated object of a Model Class. 

127 

128 """ 

129 

130 def __init__(self, config, model): 

131 super().__init__(config, model) 

132 

133 self.logger = getLogger() 

134 self.tensorboard = get_tensorboard(self.logger) 

135 self.wandblogger = WandbLogger(config) 

136 self.learner = config["learner"] 

137 self.learning_rate = config["learning_rate"] 

138 self.epochs = config["epochs"] 

139 self.eval_step = min(config["eval_step"], self.epochs) 

140 self.stopping_step = config["stopping_step"] 

141 self.clip_grad_norm = config["clip_grad_norm"] 

142 self.valid_metric = config["valid_metric"].lower() 

143 self.valid_metric_bigger = config["valid_metric_bigger"] 

144 self.test_batch_size = config["eval_batch_size"] 

145 self.gpu_available = torch.cuda.is_available() and config["use_gpu"] 

146 self.device = config["device"] 

147 self.checkpoint_dir = config["checkpoint_dir"] 

148 self.enable_amp = config["enable_amp"] 

149 self.enable_scaler = torch.cuda.is_available() and config["enable_scaler"] 

150 ensure_dir(self.checkpoint_dir) 

151 saved_model_file = "{}-{}.pth".format(self.config["model"], get_local_time()) 

152 self.saved_model_file = os.path.join(self.checkpoint_dir, saved_model_file) 

153 self.weight_decay = config["weight_decay"] 

154 

155 self.start_epoch = 0 

156 self.cur_step = 0 

157 self.best_valid_score = -np.inf if self.valid_metric_bigger else np.inf 

158 self.best_valid_result = None 

159 self.train_loss_dict = dict() 

160 self.optimizer = self._build_optimizer() 

161 self.eval_type = config["eval_type"] 

162 self.eval_collector = Collector(config) 

163 self.evaluator = Evaluator(config) 

164 

165 def _build_optimizer(self, **kwargs): 

166 r"""Init the Optimizer 

167 

168 Args: 

169 params (torch.nn.Parameter, optional): The parameters to be optimized. 

170 Defaults to ``self.model.parameters()``. 

171 learner (str, optional): The name of used optimizer. Defaults to ``self.learner``. 

172 learning_rate (float, optional): Learning rate. Defaults to ``self.learning_rate``. 

173 weight_decay (float, optional): The L2 regularization weight. Defaults to ``self.weight_decay``. 

174 

175 Returns: 

176 torch.optim: the optimizer 

177 """ 

178 params = kwargs.pop("params", self.model.parameters()) 

179 learner = kwargs.pop("learner", self.learner) 

180 learning_rate = kwargs.pop("learning_rate", self.learning_rate) 

181 weight_decay = kwargs.pop("weight_decay", self.weight_decay) 

182 

183 if self.config["reg_weight"] and weight_decay and weight_decay * self.config["reg_weight"] > 0: 

184 self.logger.warning( 

185 "The parameters [weight_decay] and [reg_weight] are specified simultaneously, " 

186 "which may lead to double regularization." 

187 ) 

188 

189 if learner.lower() == "adam": 

190 optimizer = optim.Adam(params, lr=learning_rate, weight_decay=weight_decay) 

191 elif learner.lower() == "adamw": 

192 optimizer = optim.AdamW(params, lr=learning_rate, weight_decay=weight_decay) 

193 elif learner.lower() == "sgd": 

194 optimizer = optim.SGD(params, lr=learning_rate, weight_decay=weight_decay) 

195 elif learner.lower() == "adagrad": 

196 optimizer = optim.Adagrad(params, lr=learning_rate, weight_decay=weight_decay) 

197 elif learner.lower() == "rmsprop": 

198 optimizer = optim.RMSprop(params, lr=learning_rate, weight_decay=weight_decay) 

199 elif learner.lower() == "sparse_adam": 

200 optimizer = optim.SparseAdam(params, lr=learning_rate) 

201 if weight_decay > 0: 

202 self.logger.warning("Sparse Adam cannot argument received argument [{weight_decay}]") 

203 else: 

204 self.logger.warning("Received unrecognized optimizer, set default Adam optimizer") 

205 optimizer = optim.Adam(params, lr=learning_rate) 

206 return optimizer 

207 

208 def _train_epoch(self, train_data, epoch_idx, loss_func=None, show_progress=False): 

209 r"""Train the model in an epoch. 

210 

211 Args: 

212 train_data (DataLoader): The train data. 

213 epoch_idx (int): The current epoch id. 

214 loss_func (function): The loss function of :attr:`model`. If it is ``None``, the loss function will be 

215 :attr:`self.model.calculate_loss`. Defaults to ``None``. 

216 show_progress (bool): Show the progress of training epoch. Defaults to ``False``. 

217 

218 Returns: 

219 float or tuple: The sum of loss returned by all batches in this epoch. 

220 If the loss in each batch contains multiple parts and the model 

221 returns these multiple parts loss instead of the sum of loss, it 

222 will return a tuple which includes the sum of loss in each part. 

223 """ 

224 self.model.train() 

225 loss_func = loss_func or self.model.calculate_loss 

226 total_loss = None 

227 iter_data = ( 

228 progress_bar( 

229 train_data, 

230 total=len(train_data), 

231 ncols=100, 

232 desc=set_color(f"Train {epoch_idx:>5}", "magenta", progress=True), 

233 ) 

234 if show_progress 

235 else train_data 

236 ) 

237 

238 if not self.config["single_spec"] and train_data.shuffle: 

239 train_data.sampler.set_epoch(epoch_idx) 

240 

241 scaler = grad_scaler(self.device, enabled=self.enable_scaler) 

242 for batch_idx, batch_interaction in enumerate(iter_data): 

243 interaction = batch_interaction.to(self.device) 

244 self.optimizer.zero_grad() 

245 sync_loss = 0 

246 if not self.config["single_spec"]: 

247 self.set_reduce_hook() 

248 sync_loss = self.sync_grad_loss() 

249 

250 with autocast(device_type=self.device.type, enabled=self.enable_amp): 

251 losses = loss_func(interaction) 

252 

253 if isinstance(losses, tuple): 

254 loss = sum(losses) 

255 loss_tuple = tuple(per_loss.item() for per_loss in losses) 

256 total_loss = loss_tuple if total_loss is None else tuple(map(sum, zip(total_loss, loss_tuple))) 

257 else: 

258 loss = losses 

259 total_loss = losses.item() if total_loss is None else total_loss + losses.item() 

260 self._check_nan(loss) 

261 scaler.scale(loss + sync_loss).backward() 

262 if self.clip_grad_norm: 

263 clip_grad_norm_(self.model.parameters(), **self.clip_grad_norm) 

264 scaler.step(self.optimizer) 

265 scaler.update() 

266 if self.gpu_available and show_progress: 

267 iter_data.set_postfix_str(set_color("GPU RAM: " + get_gpu_usage(self.device), "yellow")) 

268 return total_loss 

269 

270 def _valid_epoch(self, valid_data, show_progress=False): 

271 r"""Valid the model with valid data 

272 

273 Args: 

274 valid_data (DataLoader): the valid data. 

275 show_progress (bool): Show the progress of evaluate epoch. Defaults to ``False``. 

276 

277 Returns: 

278 float: valid score 

279 dict: valid result 

280 """ 

281 

282 valid_result = self.evaluate(valid_data, load_best_model=False, show_progress=show_progress) 

283 valid_score = calculate_valid_score(valid_result, self.valid_metric) 

284 return valid_score, valid_result 

285 

286 def _save_checkpoint(self, epoch, verbose=True, **kwargs): 

287 r"""Store the model parameters information and training information. 

288 

289 Args: 

290 epoch (int): the current epoch id 

291 

292 """ 

293 if not self.config["single_spec"] and self.config["local_rank"] != 0: 

294 return 

295 saved_model_file = kwargs.pop("saved_model_file", self.saved_model_file) 

296 state = { 

297 "config": self.config, 

298 "epoch": epoch, 

299 "cur_step": self.cur_step, 

300 "best_valid_score": self.best_valid_score, 

301 "state_dict": self.model.state_dict(), 

302 "other_parameter": self.model.other_parameter(), 

303 "optimizer": self.optimizer.state_dict(), 

304 } 

305 torch.save(state, saved_model_file, pickle_protocol=4) 

306 if verbose: 

307 self.logger.info(set_color("Saving current", "blue") + f": {saved_model_file}") 

308 

309 def resume_checkpoint(self, resume_file): 

310 r"""Load the model parameters information and training information. 

311 

312 Args: 

313 resume_file (file): the checkpoint file 

314 

315 """ 

316 resume_file = str(resume_file) 

317 self.saved_model_file = resume_file 

318 checkpoint = torch.load(resume_file, map_location=self.device, weights_only=False) 

319 self.start_epoch = checkpoint["epoch"] + 1 

320 self.cur_step = checkpoint["cur_step"] 

321 self.best_valid_score = checkpoint["best_valid_score"] 

322 

323 # load architecture params from checkpoint 

324 if checkpoint["config"]["model"].lower() != self.config["model"].lower(): 

325 self.logger.warning( 

326 "Architecture configuration given in config file is different from that of checkpoint. " 

327 "This may yield an exception while state_dict is being loaded." 

328 ) 

329 self.model.load_state_dict(checkpoint["state_dict"]) 

330 self.model.load_other_parameter(checkpoint.get("other_parameter")) 

331 

332 # load optimizer state from checkpoint only when optimizer type is not changed 

333 self.optimizer.load_state_dict(checkpoint["optimizer"]) 

334 message_output = f"Checkpoint loaded. Resume training from epoch {self.start_epoch}" 

335 self.logger.info(message_output) 

336 

337 def _check_nan(self, loss): 

338 if torch.isnan(loss): 

339 raise ValueError("Training loss is nan") 

340 

341 def _generate_train_loss_output(self, epoch_idx, s_time, e_time, losses): 

342 des = self.config["loss_decimal_place"] or 4 

343 train_loss_output = ( 

344 set_color("epoch %d training", "green") + " [" + set_color("time", "blue") + ": %.2fs, " 

345 ) % (epoch_idx, e_time - s_time) 

346 if isinstance(losses, tuple): 

347 des = set_color("train_loss%d", "blue") + ": %." + str(des) + "f" 

348 train_loss_output += ", ".join(des % (idx + 1, loss) for idx, loss in enumerate(losses)) 

349 else: 

350 des = "%." + str(des) + "f" 

351 train_loss_output += set_color("train loss", "blue") + ": " + des % losses 

352 return train_loss_output + "]" 

353 

354 def _add_train_loss_to_tensorboard(self, epoch_idx, losses, tag="Loss/Train"): 

355 if isinstance(losses, tuple): 

356 for idx, loss in enumerate(losses): 

357 self.tensorboard.add_scalar(tag + str(idx), loss, epoch_idx) 

358 else: 

359 self.tensorboard.add_scalar(tag, losses, epoch_idx) 

360 

361 def _add_hparam_to_tensorboard(self, best_valid_result): 

362 # base hparam 

363 hparam_dict = { 

364 "learner": self.config["learner"], 

365 "learning_rate": self.config["learning_rate"], 

366 "train_batch_size": self.config["train_batch_size"], 

367 } 

368 # unrecorded parameter 

369 unrecorded_parameter = { 

370 parameter for parameters in self.config.parameters.values() for parameter in parameters 

371 }.union({"model", "dataset", "config_files", "device"}) 

372 # other model-specific hparam 

373 hparam_dict.update( 

374 {para: val for para, val in self.config.final_config_dict.items() if para not in unrecorded_parameter} 

375 ) 

376 for k, hparam in hparam_dict.items(): 

377 if hparam is not None and not isinstance(hparam, (bool, str, float, int)): 

378 hparam_dict[k] = str(hparam) 

379 

380 self.tensorboard.add_hparams(hparam_dict, {"hparam/best_valid_result": best_valid_result}) 

381 

382 def fit(self, train_data, valid_data=None, verbose=True, saved=True, show_progress=False, callback_fn=None): 

383 r"""Train the model based on the train data and the valid data. 

384 

385 Args: 

386 train_data (DataLoader): the train data 

387 valid_data (DataLoader, optional): the valid data, default: None. 

388 If it's None, the early_stopping is invalid. 

389 verbose (bool, optional): whether to write training and evaluation information to logger, default: True 

390 saved (bool, optional): whether to save the model parameters, default: True 

391 show_progress (bool): Show the progress of training epoch and evaluate epoch. Defaults to ``False``. 

392 callback_fn (callable): Optional callback function executed at end of epoch. 

393 Includes (epoch_idx, valid_score) input arguments. 

394 

395 Returns: 

396 (float, dict): best valid score and best valid result. If valid_data is None, it returns (-1, None) 

397 """ 

398 if saved and self.start_epoch >= self.epochs: 

399 self._save_checkpoint(-1, verbose=verbose) 

400 

401 self.eval_collector.train_data_collect(train_data) 

402 if self.config["train_neg_sample_args"].get("dynamic", False): 

403 train_data.get_model(self.model) 

404 valid_step = 0 

405 

406 for epoch_idx in range(self.start_epoch, self.epochs): 

407 # train 

408 training_start_time = time() 

409 train_loss = self._train_epoch(train_data, epoch_idx, show_progress=show_progress) 

410 self.train_loss_dict[epoch_idx] = sum(train_loss) if isinstance(train_loss, tuple) else train_loss 

411 training_end_time = time() 

412 train_loss_output = self._generate_train_loss_output( 

413 epoch_idx, training_start_time, training_end_time, train_loss 

414 ) 

415 if verbose: 

416 self.logger.info(train_loss_output) 

417 self._add_train_loss_to_tensorboard(epoch_idx, train_loss) 

418 self.wandblogger.log_metrics( 

419 {"epoch": epoch_idx, "train_loss": train_loss, "train_step": epoch_idx}, 

420 head="train", 

421 ) 

422 

423 # eval 

424 if self.eval_step <= 0 or not valid_data: 

425 if saved: 

426 self._save_checkpoint(epoch_idx, verbose=verbose) 

427 continue 

428 

429 if (epoch_idx + 1) % self.eval_step == 0: 

430 valid_start_time = time() 

431 valid_score, valid_result = self._valid_epoch(valid_data, show_progress=show_progress) 

432 ( 

433 self.best_valid_score, 

434 self.cur_step, 

435 stop_flag, 

436 update_flag, 

437 ) = early_stopping( 

438 valid_score, 

439 self.best_valid_score, 

440 self.cur_step, 

441 max_step=self.stopping_step, 

442 bigger=self.valid_metric_bigger, 

443 ) 

444 valid_end_time = time() 

445 valid_score_output = ( 

446 set_color("epoch %d evaluating", "green") 

447 + " [" 

448 + set_color("time", "blue") 

449 + ": %.2fs, " 

450 + set_color("valid_score", "blue") 

451 + ": %f]" 

452 ) % (epoch_idx, valid_end_time - valid_start_time, valid_score) 

453 valid_result_output = set_color("valid result", "blue") + ": \n" + dict2str(valid_result) 

454 if verbose: 

455 self.logger.info(valid_score_output) 

456 self.logger.info(valid_result_output) 

457 self.tensorboard.add_scalar("Valid_score", valid_score, epoch_idx) 

458 self.wandblogger.log_metrics({**valid_result, "valid_step": valid_step}, head="valid") 

459 

460 if update_flag: 

461 if saved: 

462 self._save_checkpoint(epoch_idx, verbose=verbose) 

463 self.best_valid_result = valid_result 

464 

465 if callback_fn: 

466 callback_fn(epoch_idx, valid_score) 

467 

468 if stop_flag: 

469 stop_output = "Finished training, best eval result in epoch %d" % ( 

470 epoch_idx - self.cur_step * self.eval_step 

471 ) 

472 if verbose: 

473 self.logger.info(stop_output) 

474 break 

475 

476 valid_step += 1 

477 

478 self._add_hparam_to_tensorboard(self.best_valid_score) 

479 return self.best_valid_score, self.best_valid_result 

480 

481 def _batch_eval(self, batched_data, tot_item_num, neg_sampling=False, item_tensor=None): 

482 if neg_sampling: 

483 return self._neg_sample_batch_eval(batched_data, tot_item_num=tot_item_num) 

484 else: 

485 return self._full_sort_batch_eval(batched_data, tot_item_num, item_tensor) 

486 

487 def _full_sort_batch_eval(self, batched_data, tot_item_num, item_tensor): 

488 interaction, history_index, positive_u, positive_i = batched_data 

489 try: 

490 # Note: interaction without item ids 

491 scores = self.model.full_sort_predict(interaction.to(self.device)) 

492 except NotImplementedError: 

493 inter_len = len(interaction) 

494 new_inter = interaction.to(self.device).repeat_interleave(tot_item_num) 

495 batch_size = len(new_inter) 

496 new_inter.update(item_tensor.repeat(inter_len)) 

497 if batch_size <= self.test_batch_size: 

498 scores = self.model.predict(new_inter) 

499 else: 

500 scores = self._split_predict(new_inter, batch_size) 

501 

502 scores = scores.view(-1, tot_item_num) 

503 scores[:, 0] = -np.inf 

504 if history_index is not None: 

505 scores[history_index] = -np.inf 

506 return interaction, scores, positive_u, positive_i 

507 

508 def _neg_sample_batch_eval(self, batched_data, tot_item_num=None): 

509 interaction, row_idx, positive_u, positive_i = batched_data 

510 batch_size = interaction.length 

511 if batch_size <= self.test_batch_size: 

512 origin_scores = self.model.predict(interaction.to(self.device)) 

513 else: 

514 origin_scores = self._split_predict(interaction, batch_size) 

515 

516 if self.config["eval_type"] == EvaluatorType.VALUE: 

517 return interaction, origin_scores, positive_u, positive_i 

518 elif self.config["eval_type"] == EvaluatorType.RANKING: 

519 col_idx = interaction[self.config["ITEM_ID_FIELD"]] 

520 batch_user_num = positive_u[-1] + 1 

521 scores = torch.full((batch_user_num, tot_item_num), -np.inf, device=self.device) 

522 scores[row_idx, col_idx] = origin_scores 

523 return interaction, scores, positive_u, positive_i 

524 

525 @torch.no_grad() 

526 def evaluate(self, eval_data, load_best_model=True, model_file=None, show_progress=False): 

527 r"""Evaluate the model based on the eval data. 

528 

529 Args: 

530 eval_data (DataLoader): the eval data 

531 load_best_model (bool, optional): whether load the best model in the training process, default: True. 

532 It should be set True, if users want to test the model after training. 

533 model_file (str, optional): the saved model file, default: None. If users want to test the previously 

534 trained model file, they can set this parameter. 

535 show_progress (bool): Show the progress of evaluate epoch. Defaults to ``False``. 

536 

537 Returns: 

538 collections.OrderedDict: eval result, key is the eval metric and value in the corresponding metric value. 

539 """ 

540 if not eval_data: 

541 return 

542 

543 # self.eval_collector.eval_data_collect(eval_data) 

544 

545 if load_best_model: 

546 checkpoint_file = model_file or self.saved_model_file 

547 checkpoint = torch.load(checkpoint_file, weights_only=False, map_location=self.device) 

548 missing_keys, unexpected_keys = self.model.load_state_dict(checkpoint["state_dict"], strict=False) 

549 if missing_keys: 

550 self.logger.info(set_color(f"Missing loaded keys: {missing_keys}", "red")) 

551 if unexpected_keys: 

552 self.logger.info(set_color(f"Unexpected loaded keys: {unexpected_keys}", "red")) 

553 self.model.load_other_parameter(checkpoint.get("other_parameter")) 

554 message_output = f"Loading model structure and parameters from {checkpoint_file}" 

555 self.logger.info(message_output) 

556 

557 self.model.eval() 

558 

559 item_tensor = None 

560 tot_item_num = eval_data._dataset.item_num 

561 neg_sampling = isinstance(eval_data, NegSampleDataLoader) 

562 if not neg_sampling: 

563 item_tensor = eval_data._dataset.get_item_feature().to(self.device) 

564 

565 iter_data = ( 

566 progress_bar( 

567 eval_data, 

568 total=len(eval_data), 

569 ncols=100, 

570 desc=set_color("Evaluate ", "magenta", progress=True), 

571 ) 

572 if show_progress 

573 else eval_data 

574 ) 

575 num_sample = 0 

576 for batch_idx, batched_data in enumerate(iter_data): 

577 num_sample += len(batched_data) 

578 interaction, scores, positive_u, positive_i = self._batch_eval( 

579 batched_data, tot_item_num, neg_sampling=neg_sampling, item_tensor=item_tensor 

580 ) 

581 if self.gpu_available and show_progress: 

582 iter_data.set_postfix_str(set_color("GPU RAM: " + get_gpu_usage(self.device), "yellow")) 

583 self.eval_collector.eval_batch_collect(scores, interaction, positive_u, positive_i) 

584 self.eval_collector.model_collect(self.model) 

585 struct = self.eval_collector.get_data_struct() 

586 result = self.evaluator.evaluate(struct) 

587 if not self.config["single_spec"]: 

588 result = self._map_reduce(result, num_sample) 

589 self.wandblogger.log_eval_metrics(result, head="eval") 

590 return result 

591 

592 def _map_reduce(self, result, num_sample): 

593 gather_result = {} 

594 total_sample = [torch.zeros(1).to(self.device) for _ in range(self.config["world_size"])] 

595 torch.distributed.all_gather(total_sample, torch.Tensor([num_sample]).to(self.device)) 

596 total_sample = torch.cat(total_sample, 0) 

597 total_sample = torch.sum(total_sample).item() 

598 for key, value in result.items(): 

599 result[key] = torch.Tensor([value * num_sample]).to(self.device) 

600 gather_result[key] = [ 

601 torch.zeros_like(result[key]).to(self.device) for _ in range(self.config["world_size"]) 

602 ] 

603 torch.distributed.all_gather(gather_result[key], result[key]) 

604 gather_result[key] = torch.cat(gather_result[key], dim=0) 

605 gather_result[key] = round( 

606 torch.sum(gather_result[key]).item() / total_sample, 

607 self.config["metric_decimal_place"], 

608 ) 

609 return gather_result 

610 

611 def _split_predict(self, interaction, batch_size): 

612 split_interaction = dict() 

613 for key, tensor in interaction.interaction.items(): 

614 split_interaction[key] = tensor.split(self.test_batch_size, dim=0) 

615 num_block = (batch_size + self.test_batch_size - 1) // self.test_batch_size 

616 result_list = [] 

617 for i in range(num_block): 

618 current_interaction = dict() 

619 for key, spilt_tensor in split_interaction.items(): 

620 current_interaction[key] = spilt_tensor[i] 

621 result = self.model.predict(Interaction(current_interaction).to(self.device)) 

622 if len(result.shape) == 0: 

623 result = result.unsqueeze(0) 

624 result_list.append(result) 

625 return torch.cat(result_list, dim=0) 

626 

627 

628class KGTrainer(Trainer): 

629 r"""KGTrainer is designed for Knowledge-aware recommendation methods. Some of these models need to train the 

630 recommendation related task and knowledge related task alternately. 

631 

632 """ 

633 

634 def __init__(self, config, model): 

635 super().__init__(config, model) 

636 self.train_rec_step = config["train_rec_step"] 

637 self.train_kg_step = config["train_kg_step"] 

638 self.best_valid_score_lp = -np.inf if self.valid_metric_bigger else np.inf 

639 self.best_valid_result_lp = None 

640 self.cur_step_lp = 0 

641 self.tail_tensor = None 

642 

643 if config["metrics_lp"]: 

644 self.eval_collector_kg = Collector_KG(config) 

645 self.evaluator_kg = Evaluator_KG(config) 

646 

647 def _train_epoch(self, train_data, epoch_idx, loss_func=None, show_progress=False): 

648 if self.train_rec_step is None or self.train_kg_step is None: 

649 interaction_state = KGDataLoaderState.RSKG 

650 elif epoch_idx % (self.train_rec_step + self.train_kg_step) < self.train_rec_step: 

651 interaction_state = KGDataLoaderState.RS 

652 else: 

653 interaction_state = KGDataLoaderState.KG 

654 if not self.config["single_spec"]: 

655 train_data.knowledge_shuffle(epoch_idx) 

656 train_data.set_mode(interaction_state) 

657 if interaction_state in [KGDataLoaderState.RSKG, KGDataLoaderState.RS]: 

658 return super()._train_epoch(train_data, epoch_idx, show_progress=show_progress) 

659 elif interaction_state in [KGDataLoaderState.KG]: 

660 return super()._train_epoch( 

661 train_data, 

662 epoch_idx, 

663 loss_func=self.model.calculate_kg_loss, 

664 show_progress=show_progress, 

665 ) 

666 return None 

667 

668 def _valid_epoch(self, valid_data, show_progress=False): 

669 r"""Valid the model with valid data 

670 

671 Args: 

672 valid_data (Dataloader, list[Dataloader]): the valid data. 

673 show_progress (bool): Show the progress of evaluate epoch. Defaults to ``False``. 

674 

675 Returns: 

676 float: valid score 

677 dict: valid result 

678 """ 

679 if isinstance(valid_data, list): 

680 # then we have also data for the kg part 

681 valid_data_inter = valid_data[0] 

682 valid_data_kg = valid_data[1] 

683 valid_result_inter = self.evaluate(valid_data_inter, load_best_model=False, show_progress=show_progress) 

684 valid_result_kg = self.evaluate(valid_data_kg, load_best_model=False, show_progress=show_progress) 

685 valid_score_inter = calculate_valid_score(valid_result_inter, self.valid_metric) 

686 valid_score_kg = calculate_valid_score(valid_result_kg, self.valid_metric) 

687 return { 

688 KnowledgeEvaluationType.REC: [valid_score_inter, valid_result_inter], 

689 KnowledgeEvaluationType.LP: [valid_score_kg, valid_result_kg], 

690 } 

691 else: 

692 valid_data_inter = valid_data 

693 valid_result_inter = self.evaluate(valid_data_inter, load_best_model=False, show_progress=show_progress) 

694 valid_score_inter = calculate_valid_score(valid_result_inter, self.valid_metric) 

695 return {KnowledgeEvaluationType.REC: [valid_score_inter, valid_result_inter]} 

696 

697 def _batch_eval(self, batched_data, tot_target_num, neg_sampling=False, task=None, target_tensor=None): 

698 if neg_sampling: 

699 return self._neg_sample_batch_eval(batched_data, tot_item_num=tot_target_num) 

700 else: 

701 if task == KnowledgeEvaluationType.REC: 

702 full_sort_predict_fn = self.model.full_sort_predict 

703 predict_fn = self.model.predict 

704 else: 

705 full_sort_predict_fn = self.model.full_sort_predict_kg 

706 predict_fn = self.model.predict_kg 

707 

708 return self._full_sort_batch_eval( 

709 batched_data, 

710 full_sort_predict_fn, 

711 predict_fn, 

712 target_tensor, 

713 tot_target_num, 

714 ) 

715 

716 def _full_sort_batch_eval(self, batched_data, full_sort_predict_fn, predict_fn, column_tensor, tot_column_num): 

717 # in the case of recommendation, positive_h and positive_t are the user and item ids 

718 interaction, history_index, positive_h, positive_t = batched_data 

719 try: 

720 scores = full_sort_predict_fn(interaction.to(self.device)) 

721 except NotImplementedError: 

722 inter_len = len(interaction) 

723 new_inter = interaction.to(self.device).repeat_interleave(tot_column_num) 

724 batch_size = len(new_inter) 

725 new_inter.update(column_tensor.repeat(inter_len)) 

726 if batch_size <= self.test_batch_size: 

727 scores = predict_fn(new_inter) 

728 else: 

729 scores = self._split_predict_fn(new_inter, batch_size, predict_fn) 

730 

731 scores = scores.view(-1, tot_column_num) 

732 scores[:, 0] = -np.inf 

733 if history_index is not None: 

734 scores[history_index] = -np.inf 

735 return interaction, scores, positive_h, positive_t 

736 

737 def _split_predict_fn(self, interaction, batch_size, predict_fn): 

738 split_interaction = dict() 

739 for key, tensor in interaction.interaction.items(): 

740 split_interaction[key] = tensor.split(self.test_batch_size, dim=0) 

741 num_block = (batch_size + self.test_batch_size - 1) // self.test_batch_size 

742 result_list = [] 

743 for i in range(num_block): 

744 current_interaction = dict() 

745 for key, split_tensor in split_interaction.items(): 

746 current_interaction[key] = split_tensor[i] 

747 result = predict_fn(Interaction(current_interaction).to(self.device)) 

748 if len(result.shape) == 0: 

749 result = result.unsqueeze(0) 

750 result_list.append(result) 

751 return torch.cat(result_list, dim=0) 

752 

753 @torch.no_grad() 

754 def evaluate(self, eval_data, load_best_model=True, model_file=None, show_progress=False): 

755 r"""Evaluate the model based on the eval data. 

756 

757 Args: 

758 eval_data (Dataloader, list[Dataloader]): the eval data. 

759 load_best_model (bool, optional): whether load the best model in the training process, default: True. 

760 It should be set True, if users want to test the model after training. 

761 model_file (str, optional): the saved model file, default: None. If users want to test the previously 

762 trained model file, they can set this parameter. 

763 show_progress (bool): Show the progress of evaluate epoch. Defaults to ``False``. 

764 

765 Returns: 

766 collections.OrderedDict: eval result, key is the eval metric and value in the corresponding metric value. 

767 """ 

768 if not eval_data: 

769 return 

770 

771 if load_best_model: 

772 checkpoint_file = model_file or self.saved_model_file 

773 checkpoint = torch.load(checkpoint_file, weights_only=False, map_location=self.device) 

774 self.model.load_state_dict(checkpoint["state_dict"]) 

775 self.model.load_other_parameter(checkpoint.get("other_parameter")) 

776 message_output = f"Loading model structure and parameters from {checkpoint_file}" 

777 self.logger.info(message_output) 

778 

779 self.model.eval() 

780 

781 results = dict() 

782 if isinstance(eval_data, list): 

783 task_eval_data = {KnowledgeEvaluationType.REC: eval_data[0], KnowledgeEvaluationType.LP: eval_data[1]} 

784 else: 

785 if isinstance(eval_data, FullSortLPEvalDataLoader): 

786 kg_eval_type = KnowledgeEvaluationType.LP 

787 else: 

788 kg_eval_type = KnowledgeEvaluationType.REC 

789 

790 task_eval_data = {kg_eval_type: eval_data} 

791 

792 # REC task 

793 if KnowledgeEvaluationType.REC in task_eval_data: 

794 task = KnowledgeEvaluationType.REC 

795 rec_eval_data = task_eval_data[task] 

796 

797 item_tensor = None 

798 tot_item_num = rec_eval_data._dataset.item_num 

799 neg_sampling = isinstance(rec_eval_data, NegSampleDataLoader) 

800 if not neg_sampling: 

801 item_tensor = rec_eval_data._dataset.get_item_feature().to(self.device) 

802 

803 results[task] = self.evaluate_data_loop( 

804 rec_eval_data, task, tot_item_num, item_tensor, show_progress=show_progress 

805 ) 

806 

807 # LP task 

808 if KnowledgeEvaluationType.LP in task_eval_data: 

809 task = KnowledgeEvaluationType.LP 

810 kg_eval_data = task_eval_data[task] 

811 

812 tot_entity_num = kg_eval_data._dataset.entity_num 

813 tail_tensor = kg_eval_data._dataset.get_tail_feature().to(self.device) 

814 

815 results[task] = self.evaluate_data_loop( 

816 kg_eval_data, 

817 task, 

818 tot_entity_num, 

819 tail_tensor, 

820 show_progress=show_progress, 

821 ) 

822 

823 if isinstance(eval_data, list): 

824 return results 

825 else: 

826 return results[kg_eval_type] 

827 

828 def evaluate_data_loop(self, eval_data, task, tot_target_num, target_tensor, show_progress=True): 

829 neg_sampling = isinstance(eval_data, NegSampleDataLoader) 

830 

831 if task == KnowledgeEvaluationType.REC: 

832 eval_collector = self.eval_collector 

833 evaluator = self.evaluator 

834 else: 

835 eval_collector = self.eval_collector_kg 

836 evaluator = self.evaluator_kg 

837 

838 iter_data = ( 

839 progress_bar( 

840 eval_data, 

841 total=len(eval_data), 

842 ncols=100, 

843 desc=set_color(f"Evaluate {task}", "magenta", progress=True), 

844 ) 

845 if show_progress 

846 else eval_data 

847 ) 

848 

849 num_sample = 0 

850 for batch_idx, batched_data in enumerate(iter_data): 

851 num_sample += len(batched_data) 

852 interaction, scores, positive_u, positive_i = self._batch_eval( 

853 batched_data, 

854 tot_target_num, 

855 neg_sampling=neg_sampling, 

856 task=task, 

857 target_tensor=target_tensor, 

858 ) 

859 if self.gpu_available and show_progress: 

860 iter_data.set_postfix_str(set_color("GPU RAM: " + get_gpu_usage(self.device), "yellow")) 

861 eval_collector.eval_batch_collect(scores, interaction, positive_u, positive_i) 

862 eval_collector.model_collect(self.model) 

863 struct = eval_collector.get_data_struct() 

864 result = evaluator.evaluate(struct) 

865 if not self.config["single_spec"]: 

866 result = self._map_reduce(result, num_sample) 

867 self.wandblogger.log_eval_metrics(result, head="eval") 

868 return result 

869 

870 def fit(self, train_data, valid_data=None, verbose=True, saved=True, show_progress=False, callback_fn=None): 

871 r"""Train the model based on the train data and the valid data. 

872 

873 Args: 

874 train_data (DataLoader): the train data 

875 valid_data (DataLoader, optional): the valid data, default: None. 

876 If it's None, the early_stopping is invalid. 

877 verbose (bool, optional): whether to write training and evaluation information to logger, default: True 

878 saved (bool, optional): whether to save the model parameters, default: True 

879 show_progress (bool): Show the progress of training epoch and evaluate epoch. Defaults to ``False``. 

880 callback_fn (callable): Optional callback function executed at end of epoch. 

881 Includes (epoch_idx, valid_score) input arguments. 

882 

883 Returns: 

884 (float, dict): best valid score and best valid result. If valid_data is None, it returns (-1, None) 

885 """ 

886 if saved and self.start_epoch >= self.epochs: 

887 self._save_checkpoint(-1, verbose=verbose) 

888 

889 self.eval_collector.train_data_collect(train_data) 

890 if self.config["train_neg_sample_args"].get("dynamic", False): 

891 train_data.get_model(self.model) 

892 valid_step = 0 

893 

894 for epoch_idx in range(self.start_epoch, self.epochs): 

895 # train 

896 training_start_time = time() 

897 train_loss = self._train_epoch(train_data, epoch_idx, show_progress=show_progress) 

898 self.train_loss_dict[epoch_idx] = sum(train_loss) if isinstance(train_loss, tuple) else train_loss 

899 training_end_time = time() 

900 train_loss_output = self._generate_train_loss_output( 

901 epoch_idx, training_start_time, training_end_time, train_loss 

902 ) 

903 if verbose: 

904 self.logger.info(train_loss_output) 

905 self._add_train_loss_to_tensorboard(epoch_idx, train_loss) 

906 self.wandblogger.log_metrics( 

907 {"epoch": epoch_idx, "train_loss": train_loss, "train_step": epoch_idx}, 

908 head="train", 

909 ) 

910 

911 # eval 

912 if not isinstance(valid_data, list): 

913 if self.eval_step <= 0 or not valid_data: 

914 if saved: 

915 self._save_checkpoint(epoch_idx, verbose=verbose) 

916 continue 

917 elif self.eval_step <= 0 or not valid_data[0] or not valid_data[1]: 

918 if saved: 

919 self._save_checkpoint(epoch_idx, verbose=verbose) 

920 continue 

921 

922 if (epoch_idx + 1) % self.eval_step == 0: 

923 valid_start_time = time() 

924 return_data = self._valid_epoch(valid_data, show_progress=show_progress) 

925 

926 best_valid = defaultdict(dict) 

927 

928 if KnowledgeEvaluationType.LP in return_data: 

929 kg_valid_scores = list() 

930 kg_valid_results = list() 

931 

932 update_flag = False 

933 stop_flag = False 

934 

935 for task, (valid_score, valid_result) in return_data.items(): 

936 # TODO 

937 # - add early stopping for KG 

938 if task == KnowledgeEvaluationType.REC: 

939 ( 

940 self.best_valid_score, 

941 self.cur_step, 

942 stop_flag, 

943 update_flag, 

944 ) = early_stopping( 

945 valid_score, 

946 self.best_valid_score, 

947 self.cur_step, 

948 max_step=self.stopping_step, 

949 bigger=self.valid_metric_bigger, 

950 ) 

951 else: 

952 ( 

953 self.best_valid_score_lp, 

954 self.cur_step_lp, 

955 _, 

956 _, 

957 ) = early_stopping( 

958 valid_score, 

959 self.best_valid_score_lp, 

960 self.cur_step_lp, 

961 max_step=self.stopping_step, 

962 bigger=self.valid_metric_bigger, 

963 ) 

964 

965 valid_end_time = time() 

966 valid_score_output = ( 

967 set_color(f"epoch %d evaluating {task}", "green") 

968 + " [" 

969 + set_color("time", "blue") 

970 + ": %.2fs, " 

971 + set_color("valid_score", "blue") 

972 + ": %f]" 

973 ) % (epoch_idx, valid_end_time - valid_start_time, valid_score) 

974 valid_result_output = set_color("valid result ", "blue") + ": \n" + dict2str(valid_result) 

975 

976 if verbose: 

977 self.logger.info(valid_score_output) 

978 self.logger.info(valid_result_output) 

979 

980 self.tensorboard.add_scalar(f"Valid_score_{task}", valid_score, epoch_idx) 

981 self.wandblogger.log_metrics({**valid_result, f"valid_step_{task}": valid_step}, head="valid") 

982 

983 if task == KnowledgeEvaluationType.REC and update_flag: 

984 if saved: 

985 self._save_checkpoint(epoch_idx, verbose=verbose) 

986 self.best_valid_result = valid_result 

987 

988 if callback_fn: 

989 callback_fn(epoch_idx, valid_score) 

990 

991 if task == KnowledgeEvaluationType.REC and stop_flag: 

992 stop_output = "Finished training, best eval result in epoch %d" % ( 

993 epoch_idx - self.cur_step * self.eval_step 

994 ) 

995 if verbose: 

996 self.logger.info(stop_output) 

997 break 

998 

999 if task == KnowledgeEvaluationType.LP: 

1000 # Track valid scores and results 

1001 kg_valid_scores.append(valid_score) 

1002 kg_valid_results.append(valid_result) 

1003 # Track best valid scores and results 

1004 best_valid["score"][KnowledgeEvaluationType.LP] = max(kg_valid_scores) 

1005 best_valid["result"][KnowledgeEvaluationType.LP] = kg_valid_results[ 

1006 kg_valid_scores.index(best_valid["score"][KnowledgeEvaluationType.LP]) 

1007 ] 

1008 valid_step += 1 

1009 

1010 best_valid["score"][KnowledgeEvaluationType.REC] = self.best_valid_score 

1011 best_valid["result"][KnowledgeEvaluationType.REC] = self.best_valid_result 

1012 

1013 self._add_hparam_to_tensorboard(self.best_valid_score) 

1014 return best_valid["score"], best_valid["result"] 

1015 

1016 

1017class ExplainableTrainer(Trainer): 

1018 """ExplainableTrainer is designed for explainable recommendation methods.""" 

1019 

1020 def __init__(self, config, model): 

1021 super().__init__(config, model) 

1022 self.eval_collector = ExplainableCollector(config) 

1023 

1024 def _full_sort_batch_eval(self, batched_data, tot_item_num, item_tensor, **kwargs): 

1025 interaction, history_index, positive_u, positive_i = batched_data 

1026 

1027 scores, paths = self.model.explain(interaction.to(self.device), **kwargs) 

1028 

1029 scores = scores.view(-1, tot_item_num) 

1030 scores[:, 0] = -np.inf 

1031 if history_index is not None: 

1032 scores[history_index] = -np.inf 

1033 

1034 return interaction, (scores, paths), positive_u, positive_i 

1035 

1036 

1037class PGPRTrainer(ExplainableTrainer): 

1038 r"""PGPRTrainer is designed for PGPR, which is a knowledge-aware recommendation method.""" 

1039 

1040 def __init__(self, config, model): 

1041 super().__init__(config, model) 

1042 

1043 

1044class CAFETrainer(ExplainableTrainer): 

1045 r"""CAFETrainer is designed for CAFE, which is a knowledge-aware recommendation method.""" 

1046 

1047 def __init__(self, config, model): 

1048 super().__init__(config, model) 

1049 

1050 

1051class KGATTrainer(Trainer): 

1052 r"""KGATTrainer is designed for KGAT, which is a knowledge-aware recommendation method.""" 

1053 

1054 def __init__(self, config, model): 

1055 super().__init__(config, model) 

1056 

1057 def _train_epoch(self, train_data, epoch_idx, loss_func=None, show_progress=False): 

1058 # train rs 

1059 if not self.config["single_spec"]: 

1060 train_data.knowledge_shuffle(epoch_idx) 

1061 train_data.set_mode(KGDataLoaderState.RS) 

1062 rs_total_loss = super()._train_epoch(train_data, epoch_idx, show_progress=show_progress) 

1063 

1064 # train kg 

1065 train_data.set_mode(KGDataLoaderState.KG) 

1066 kg_total_loss = super()._train_epoch( 

1067 train_data, 

1068 epoch_idx, 

1069 loss_func=self.model.calculate_kg_loss, 

1070 show_progress=show_progress, 

1071 ) 

1072 

1073 # update A 

1074 self.model.eval() 

1075 with torch.no_grad(): 

1076 self.model.update_attentive_A() 

1077 

1078 return rs_total_loss, kg_total_loss 

1079 

1080 

1081class PretrainTrainer(Trainer): 

1082 r"""PretrainTrainer is designed for pre-training. 

1083 It can be inherited by the trainer which needs pre-training and fine-tuning. 

1084 """ 

1085 

1086 def __init__(self, config, model): 

1087 super().__init__(config, model) 

1088 self.pretrain_epochs = self.config["pretrain_epochs"] 

1089 self.save_step = self.config["save_step"] 

1090 

1091 def save_pretrained_model(self, epoch, saved_model_file): 

1092 r"""Store the model parameters information and training information. 

1093 

1094 Args: 

1095 epoch (int): the current epoch id 

1096 saved_model_file (str): file name for saved pretrained model 

1097 

1098 """ 

1099 state = { 

1100 "config": self.config, 

1101 "epoch": epoch, 

1102 "state_dict": self.model.state_dict(), 

1103 "optimizer": self.optimizer.state_dict(), 

1104 "other_parameter": self.model.other_parameter(), 

1105 } 

1106 torch.save(state, saved_model_file) 

1107 self.saved_model_file = saved_model_file 

1108 

1109 def _get_pretrained_model_path(self, epoch_label=None): 

1110 epoch_label = str(epoch_label) if epoch_label is not None else "pretrained" 

1111 return os.path.join( 

1112 self.checkpoint_dir, 

1113 "{}-{}-{}.pth".format(self.config["model"], self.config["dataset"], epoch_label), 

1114 ) 

1115 

1116 def pretrain(self, train_data, verbose=True, show_progress=False): 

1117 for epoch_idx in range(self.start_epoch, self.pretrain_epochs): 

1118 # train 

1119 training_start_time = time() 

1120 train_loss = self._train_epoch(train_data, epoch_idx, show_progress=show_progress) 

1121 self.train_loss_dict[epoch_idx] = sum(train_loss) if isinstance(train_loss, tuple) else train_loss 

1122 training_end_time = time() 

1123 train_loss_output = self._generate_train_loss_output( 

1124 epoch_idx, training_start_time, training_end_time, train_loss 

1125 ) 

1126 if verbose: 

1127 self.logger.info(train_loss_output) 

1128 self._add_train_loss_to_tensorboard(epoch_idx, train_loss) 

1129 

1130 if (epoch_idx + 1) % self.save_step == 0: 

1131 saved_model_file = self._get_pretrained_model_path(epoch_idx + 1) 

1132 self.save_pretrained_model(epoch_idx, saved_model_file) 

1133 update_output = set_color("Saving current", "blue") + ": %s" % saved_model_file 

1134 if verbose: 

1135 self.logger.info(update_output) 

1136 

1137 return self.best_valid_score, self.best_valid_result 

1138 

1139 

1140class S3RecTrainer(PretrainTrainer): 

1141 r"""S3RecTrainer is designed for S3Rec, which is a self-supervised learning based sequential recommenders. 

1142 It includes two training stages: pre-training ang fine-tuning. 

1143 

1144 """ 

1145 

1146 def __init__(self, config, model): 

1147 super().__init__(config, model) 

1148 

1149 def fit( 

1150 self, 

1151 train_data, 

1152 valid_data=None, 

1153 verbose=True, 

1154 saved=True, 

1155 show_progress=False, 

1156 callback_fn=None, 

1157 ): 

1158 if self.model.train_stage == "pretrain": 

1159 return self.pretrain(train_data, verbose, show_progress) 

1160 elif self.model.train_stage == "finetune": 

1161 return super().fit(train_data, valid_data, verbose, saved, show_progress, callback_fn) 

1162 else: 

1163 raise ValueError("Please make sure that the 'train_stage' is 'pretrain' or 'finetune'!") 

1164 

1165 

1166class TPRecTrainer(PretrainTrainer): 

1167 """ 

1168 TPRecTrainer is designed for TPRec, which is a knowledge-aware recommendation method. 

1169 """ 

1170 

1171 def __init__(self, config, model): 

1172 super().__init__(config, model) 

1173 

1174 def fit( 

1175 self, 

1176 train_data, 

1177 valid_data=None, 

1178 verbose=True, 

1179 saved=True, 

1180 show_progress=False, 

1181 callback_fn=None, 

1182 ): 

1183 if self.model.train_stage == "pretrain": 

1184 return self.pretrain(train_data, verbose, show_progress) 

1185 elif self.model.train_stage == "policy": 

1186 return super().fit(train_data, valid_data, verbose, saved, show_progress, callback_fn) 

1187 else: 

1188 raise ValueError("Please make sure that the 'train_stage' is 'pretrain' or 'finetune'!") 

1189 

1190 def _train_epoch(self, train_data, epoch_idx, loss_func=None, show_progress=False): 

1191 if self.config["train_stage"] == "policy": 

1192 return super()._train_epoch(train_data, epoch_idx, loss_func=loss_func, show_progress=show_progress) 

1193 

1194 if self.config["train_rec_step"] is None or self.config["train_kg_step"] is None: 

1195 interaction_state = KGDataLoaderState.RSKG 

1196 elif ( 

1197 epoch_idx % (self.config["train_rec_step"] + self.config["train_kg_step"]) < self.config["train_rec_step"] 

1198 ): 

1199 interaction_state = KGDataLoaderState.RS 

1200 else: 

1201 interaction_state = KGDataLoaderState.KG 

1202 if not self.config["single_spec"]: 

1203 train_data.knowledge_shuffle(epoch_idx) 

1204 train_data.set_mode(interaction_state) 

1205 if interaction_state in [KGDataLoaderState.RSKG, KGDataLoaderState.RS]: 

1206 return super()._train_epoch(train_data, epoch_idx, show_progress=show_progress) 

1207 elif interaction_state in [KGDataLoaderState.KG]: 

1208 return super()._train_epoch( 

1209 train_data, 

1210 epoch_idx, 

1211 loss_func=self.model.calculate_loss, 

1212 show_progress=show_progress, 

1213 ) 

1214 return None 

1215 

1216 @torch.no_grad() 

1217 def evaluate(self, eval_data, load_best_model=True, model_file=None, show_progress=False): 

1218 r"""Evaluate the model based on the eval data. 

1219 

1220 Args: 

1221 eval_data (DataLoader): the eval data 

1222 load_best_model (bool, optional): whether load the best model in the training process, default: True. 

1223 It should be set True, if users want to test the model after training. 

1224 model_file (str, optional): the saved model file, default: None. If users want to test the previously 

1225 trained model file, they can set this parameter. 

1226 show_progress (bool): Show the progress of evaluate epoch. Defaults to ``False``. 

1227 

1228 Returns: 

1229 collections.OrderedDict: eval result, key is the eval metric and value in the corresponding metric value. 

1230 """ 

1231 if not eval_data: 

1232 return 

1233 

1234 # self.eval_collector.eval_data_collect(eval_data) 

1235 

1236 if load_best_model: 

1237 checkpoint_file = model_file or self.saved_model_file 

1238 checkpoint = torch.load(checkpoint_file, weights_only=False, map_location=self.device) 

1239 self.model.load_state_dict(checkpoint["state_dict"]) 

1240 self.model.load_other_parameter(checkpoint.get("other_parameter")) 

1241 message_output = f"Loading model structure and parameters from {checkpoint_file}" 

1242 self.logger.info(message_output) 

1243 

1244 self.model.eval() 

1245 

1246 item_tensor = None 

1247 tot_item_num = eval_data._dataset.item_num 

1248 neg_sampling = isinstance(eval_data, NegSampleDataLoader) 

1249 if not neg_sampling: 

1250 item_tensor = eval_data._dataset.get_item_feature().to(self.device) 

1251 

1252 iter_data = ( 

1253 progress_bar( 

1254 eval_data, 

1255 total=len(eval_data), 

1256 ncols=100, 

1257 desc=set_color("Evaluate ", "magenta", progress=True), 

1258 ) 

1259 if show_progress 

1260 else eval_data 

1261 ) 

1262 num_sample = 0 

1263 for batch_idx, batched_data in enumerate(iter_data): 

1264 num_sample += len(batched_data) 

1265 

1266 if self.model.train_stage == "policy": 

1267 batched_data = (batched_data, eval_data.temporal_weights) # noqa: PLW2901 

1268 

1269 interaction, scores, positive_u, positive_i = self._batch_eval( 

1270 batched_data, tot_item_num, neg_sampling=neg_sampling, item_tensor=item_tensor 

1271 ) 

1272 if self.gpu_available and show_progress: 

1273 iter_data.set_postfix_str(set_color("GPU RAM: " + get_gpu_usage(self.device), "yellow")) 

1274 self.eval_collector.eval_batch_collect(scores, interaction, positive_u, positive_i) 

1275 self.eval_collector.model_collect(self.model) 

1276 struct = self.eval_collector.get_data_struct() 

1277 result = self.evaluator.evaluate(struct) 

1278 if not self.config["single_spec"]: 

1279 result = self._map_reduce(result, num_sample) 

1280 self.wandblogger.log_eval_metrics(result, head="eval") 

1281 return result 

1282 

1283 def _full_sort_batch_eval(self, batched_data, tot_item_num, item_tensor): 

1284 if self.model.train_stage == "pretrain": 

1285 return super()._full_sort_batch_eval(batched_data, tot_item_num, item_tensor) 

1286 

1287 paths = None 

1288 batched_data, temporal_weights = batched_data 

1289 

1290 interaction, history_index, positive_u, positive_i = batched_data 

1291 

1292 # Note: interaction without item ids 

1293 scores = self.model.full_sort_predict((interaction.to(self.device), temporal_weights)) 

1294 

1295 if isinstance(scores, tuple): 

1296 # then the first is the score, the second are paths 

1297 scores, paths = scores 

1298 

1299 scores = scores.view(-1, tot_item_num) 

1300 scores[:, 0] = -np.inf 

1301 if history_index is not None: 

1302 scores[history_index] = -np.inf 

1303 

1304 if paths is not None: 

1305 return interaction, (scores, paths), positive_u, positive_i 

1306 else: 

1307 return interaction, scores, positive_u, positive_i 

1308 

1309 

1310class MKRTrainer(Trainer): 

1311 r"""MKRTrainer is designed for MKR, which is a knowledge-aware recommendation method.""" 

1312 

1313 def __init__(self, config, model): 

1314 super().__init__(config, model) 

1315 self.kge_interval = config["kge_interval"] 

1316 

1317 def _train_epoch(self, train_data, epoch_idx, loss_func=None, show_progress=False): 

1318 rs_total_loss, kg_total_loss = 0.0, 0.0 

1319 

1320 # train rs 

1321 self.logger.info("Train RS") 

1322 train_data.set_mode(KGDataLoaderState.RS) 

1323 rs_total_loss = super()._train_epoch( 

1324 train_data, 

1325 epoch_idx, 

1326 loss_func=self.model.calculate_rs_loss, 

1327 show_progress=show_progress, 

1328 ) 

1329 

1330 # train kg 

1331 if epoch_idx % self.kge_interval == 0: 

1332 self.logger.info("Train KG") 

1333 train_data.set_mode(KGDataLoaderState.KG) 

1334 kg_total_loss = super()._train_epoch( 

1335 train_data, 

1336 epoch_idx, 

1337 loss_func=self.model.calculate_kg_loss, 

1338 show_progress=show_progress, 

1339 ) 

1340 

1341 return rs_total_loss, kg_total_loss 

1342 

1343 

1344class TraditionalTrainer(Trainer): 

1345 """TraditionalTrainer is designed for Traditional model(Pop,ItemKNN), 

1346 which set the epoch to 1 whatever the config.""" 

1347 

1348 def __init__(self, config, model): 

1349 super().__init__(config, model) 

1350 self.epochs = 1 # Set the epoch to 1 when running memory based model 

1351 

1352 

1353class DecisionTreeTrainer(AbstractTrainer): 

1354 """DecisionTreeTrainer is designed for DecisionTree model.""" 

1355 

1356 def __init__(self, config, model): 

1357 super().__init__(config, model) 

1358 

1359 self.logger = getLogger() 

1360 self.tensorboard = get_tensorboard(self.logger) 

1361 self.label_field = config["LABEL_FIELD"] 

1362 self.convert_token_to_onehot = self.config["convert_token_to_onehot"] 

1363 

1364 # evaluator 

1365 self.eval_type = config["eval_type"] 

1366 self.epochs = config["epochs"] 

1367 self.eval_step = min(config["eval_step"], self.epochs) 

1368 self.valid_metric = config["valid_metric"].lower() 

1369 self.eval_collector = Collector(config) 

1370 self.evaluator = Evaluator(config) 

1371 

1372 # model saved 

1373 self.checkpoint_dir = config["checkpoint_dir"] 

1374 ensure_dir(self.checkpoint_dir) 

1375 temp_file = "{}-{}-temp.pth".format(self.config["model"], get_local_time()) 

1376 self.temp_file = os.path.join(self.checkpoint_dir, temp_file) 

1377 

1378 temp_best_file = "{}-{}-temp-best.pth".format(self.config["model"], get_local_time()) 

1379 self.temp_best_file = os.path.join(self.checkpoint_dir, temp_best_file) 

1380 

1381 saved_model_file = "{}-{}.pth".format(self.config["model"], get_local_time()) 

1382 self.saved_model_file = os.path.join(self.checkpoint_dir, saved_model_file) 

1383 

1384 self.stopping_step = config["stopping_step"] 

1385 self.valid_metric_bigger = config["valid_metric_bigger"] 

1386 self.cur_step = 0 

1387 self.best_valid_score = -np.inf if self.valid_metric_bigger else np.inf 

1388 self.best_valid_result = None 

1389 

1390 def _interaction_to_sparse(self, dataloader): 

1391 r"""Convert data format from interaction to sparse or numpy 

1392 

1393 Args: 

1394 dataloader (DecisionTreeDataLoader): DecisionTreeDataLoader dataloader. 

1395 

1396 Returns: 

1397 cur_data (sparse or numpy): data. 

1398 interaction_np[self.label_field] (numpy): label. 

1399 """ 

1400 interaction = dataloader._dataset[:] 

1401 interaction_np = interaction.numpy() 

1402 cur_data = np.array([]) 

1403 columns = [] 

1404 for key, interaction_value in interaction_np.items(): 

1405 value = np.resize(interaction_value, (interaction_value.shape[0], 1)) 

1406 if key != self.label_field: 

1407 columns.append(key) 

1408 if cur_data.shape[0] == 0: 

1409 cur_data = value 

1410 else: 

1411 cur_data = np.hstack((cur_data, value)) 

1412 

1413 if self.convert_token_to_onehot: 

1414 convert_col_list = dataloader._dataset.convert_col_list 

1415 hash_count = dataloader._dataset.hash_count 

1416 

1417 new_col = cur_data.shape[1] - len(convert_col_list) 

1418 for key, values in hash_count.items(): 

1419 new_col = new_col + values 

1420 onehot_data = sparse.dok_matrix((cur_data.shape[0], new_col)) 

1421 

1422 cur_j = 0 

1423 new_j = 0 

1424 

1425 for key in columns: 

1426 if key in convert_col_list: 

1427 for i in range(cur_data.shape[0]): 

1428 onehot_data[i, int(new_j + cur_data[i, cur_j])] = 1 

1429 new_j = new_j + hash_count[key] - 1 

1430 else: 

1431 for i in range(cur_data.shape[0]): 

1432 onehot_data[i, new_j] = cur_data[i, cur_j] 

1433 cur_j = cur_j + 1 

1434 new_j = new_j + 1 

1435 

1436 cur_data = sparse.csc_matrix(onehot_data) 

1437 

1438 return cur_data, interaction_np[self.label_field] 

1439 

1440 def _interaction_to_lib_datatype(self, dataloader): 

1441 pass 

1442 

1443 def _valid_epoch(self, valid_data): 

1444 r"""Args: 

1445 valid_data (DecisionTreeDataLoader): DecisionTreeDataLoader, which is the same with GeneralDataLoader. 

1446 """ 

1447 valid_result = self.evaluate(valid_data, load_best_model=False) 

1448 valid_score = calculate_valid_score(valid_result, self.valid_metric) 

1449 return valid_score, valid_result 

1450 

1451 def _save_checkpoint(self, epoch): 

1452 r"""Store the model parameters information and training information. 

1453 

1454 Args: 

1455 epoch (int): the current epoch id 

1456 

1457 """ 

1458 state = { 

1459 "config": self.config, 

1460 "epoch": epoch, 

1461 "cur_step": self.cur_step, 

1462 "best_valid_score": self.best_valid_score, 

1463 "state_dict": self.temp_best_file, 

1464 "other_parameter": None, 

1465 } 

1466 torch.save(state, self.saved_model_file) 

1467 

1468 def fit(self, train_data, valid_data=None, verbose=True, saved=True, show_progress=False, callback_fn=None): 

1469 for epoch_idx in range(self.epochs): 

1470 self._train_at_once(train_data, valid_data) 

1471 

1472 if (epoch_idx + 1) % self.eval_step == 0: 

1473 # evaluate 

1474 valid_start_time = time() 

1475 valid_score, valid_result = self._valid_epoch(valid_data) 

1476 

1477 ( 

1478 self.best_valid_score, 

1479 self.cur_step, 

1480 stop_flag, 

1481 update_flag, 

1482 ) = early_stopping( 

1483 valid_score, 

1484 self.best_valid_score, 

1485 self.cur_step, 

1486 max_step=self.stopping_step, 

1487 bigger=self.valid_metric_bigger, 

1488 ) 

1489 

1490 valid_end_time = time() 

1491 valid_score_output = ( 

1492 set_color("epoch %d evaluating", "green") 

1493 + " [" 

1494 + set_color("time", "blue") 

1495 + ": %.2fs, " 

1496 + set_color("valid_score", "blue") 

1497 + ": %f]" 

1498 ) % (epoch_idx, valid_end_time - valid_start_time, valid_score) 

1499 valid_result_output = set_color("valid result", "blue") + ": \n" + dict2str(valid_result) 

1500 if verbose: 

1501 self.logger.info(valid_score_output) 

1502 self.logger.info(valid_result_output) 

1503 self.tensorboard.add_scalar("Valid_score", valid_score, epoch_idx) 

1504 

1505 if update_flag: 

1506 if saved: 

1507 self.model.save_model(self.temp_best_file) 

1508 self._save_checkpoint(epoch_idx) 

1509 self.best_valid_result = valid_result 

1510 

1511 if stop_flag: 

1512 stop_output = "Finished training, best eval result in epoch %d" % ( 

1513 epoch_idx - self.cur_step * self.eval_step 

1514 ) 

1515 if self.temp_file: 

1516 os.remove(self.temp_file) 

1517 if verbose: 

1518 self.logger.info(stop_output) 

1519 break 

1520 

1521 return self.best_valid_score, self.best_valid_result 

1522 

1523 def evaluate(self, eval_data, load_best_model=True, model_file=None, show_progress=False): 

1524 raise NotImplementedError 

1525 

1526 def _train_at_once(self, train_data, valid_data): 

1527 raise NotImplementedError 

1528 

1529 

1530class XGBoostTrainer(DecisionTreeTrainer): 

1531 """XGBoostTrainer is designed for XGBOOST.""" 

1532 

1533 def __init__(self, config, model): 

1534 super().__init__(config, model) 

1535 

1536 self.xgb = __import__("xgboost") 

1537 self.boost_model = config["xgb_model"] 

1538 self.silent = config["xgb_silent"] 

1539 self.nthread = config["xgb_nthread"] 

1540 

1541 # train params 

1542 self.params = config["xgb_params"] 

1543 self.num_boost_round = config["xgb_num_boost_round"] 

1544 self.evals = () 

1545 self.early_stopping_rounds = config["xgb_early_stopping_rounds"] 

1546 self.evals_result = {} 

1547 self.verbose_eval = config["xgb_verbose_eval"] 

1548 self.callbacks = None 

1549 self.deval = None 

1550 self.eval_pred = self.eval_true = None 

1551 

1552 def _interaction_to_lib_datatype(self, dataloader): 

1553 r"""Convert data format from interaction to DMatrix 

1554 

1555 Args: 

1556 dataloader (DecisionTreeDataLoader): xgboost dataloader. 

1557 

1558 Returns: 

1559 DMatrix: Data in the form of 'DMatrix'. 

1560 """ 

1561 data, label = self._interaction_to_sparse(dataloader) 

1562 return self.xgb.DMatrix(data=data, label=label, silent=self.silent, nthread=self.nthread) 

1563 

1564 def _train_at_once(self, train_data, valid_data): 

1565 r"""Args: 

1566 train_data (DecisionTreeDataLoader): DecisionTreeDataLoader, which is the same with GeneralDataLoader. 

1567 valid_data (DecisionTreeDataLoader): DecisionTreeDataLoader, which is the same with GeneralDataLoader. 

1568 """ 

1569 self.dtrain = self._interaction_to_lib_datatype(train_data) 

1570 self.dvalid = self._interaction_to_lib_datatype(valid_data) 

1571 self.evals = [(self.dtrain, "train"), (self.dvalid, "valid")] 

1572 self.model = self.xgb.train( 

1573 self.params, 

1574 self.dtrain, 

1575 self.num_boost_round, 

1576 self.evals, 

1577 early_stopping_rounds=self.early_stopping_rounds, 

1578 evals_result=self.evals_result, 

1579 verbose_eval=self.verbose_eval, 

1580 xgb_model=self.boost_model, 

1581 callbacks=self.callbacks, 

1582 ) 

1583 

1584 self.model.save_model(self.temp_file) 

1585 self.boost_model = self.temp_file 

1586 

1587 def evaluate(self, eval_data, load_best_model=True, model_file=None, show_progress=False): 

1588 if load_best_model: 

1589 if model_file: 

1590 checkpoint_file = model_file 

1591 else: 

1592 checkpoint_file = self.temp_best_file 

1593 self.model.load_model(checkpoint_file) 

1594 

1595 self.deval = self._interaction_to_lib_datatype(eval_data) 

1596 self.eval_true = torch.Tensor(self.deval.get_label()) 

1597 self.eval_pred = torch.Tensor(self.model.predict(self.deval)) 

1598 

1599 self.eval_collector.eval_collect(self.eval_pred, self.eval_true) 

1600 result = self.evaluator.evaluate(self.eval_collector.get_data_struct()) 

1601 return result 

1602 

1603 

1604class LightGBMTrainer(DecisionTreeTrainer): 

1605 """LightGBMTrainer is designed for LightGBM.""" 

1606 

1607 def __init__(self, config, model): 

1608 super().__init__(config, model) 

1609 

1610 self.lgb = __import__("lightgbm") 

1611 

1612 # train params 

1613 self.params = config["lgb_params"] 

1614 self.num_boost_round = config["lgb_num_boost_round"] 

1615 self.evals = () 

1616 self.deval_data = self.deval_label = None 

1617 self.eval_pred = self.eval_true = None 

1618 

1619 def _interaction_to_lib_datatype(self, dataloader): 

1620 r"""Convert data format from interaction to Dataset 

1621 

1622 Args: 

1623 dataloader (DecisionTreeDataLoader): xgboost dataloader. 

1624 

1625 Returns: 

1626 dataset(lgb.Dataset): Data in the form of 'lgb.Dataset'. 

1627 """ 

1628 data, label = self._interaction_to_sparse(dataloader) 

1629 return self.lgb.Dataset(data=data, label=label) 

1630 

1631 def _train_at_once(self, train_data, valid_data): 

1632 r"""Args: 

1633 train_data (DecisionTreeDataLoader): DecisionTreeDataLoader, which is the same with GeneralDataLoader. 

1634 valid_data (DecisionTreeDataLoader): DecisionTreeDataLoader, which is the same with GeneralDataLoader. 

1635 """ 

1636 self.dtrain = self._interaction_to_lib_datatype(train_data) 

1637 self.dvalid = self._interaction_to_lib_datatype(valid_data) 

1638 self.evals = [self.dtrain, self.dvalid] 

1639 self.model = self.lgb.train(self.params, self.dtrain, self.num_boost_round, self.evals) 

1640 

1641 self.model.save_model(self.temp_file) 

1642 self.boost_model = self.temp_file 

1643 

1644 def evaluate(self, eval_data, load_best_model=True, model_file=None, show_progress=False): 

1645 if load_best_model: 

1646 if model_file: 

1647 checkpoint_file = model_file 

1648 else: 

1649 checkpoint_file = self.temp_best_file 

1650 self.model = self.lgb.Booster(model_file=checkpoint_file) 

1651 

1652 self.deval_data, self.deval_label = self._interaction_to_sparse(eval_data) 

1653 self.eval_true = torch.Tensor(self.deval_label) 

1654 self.eval_pred = torch.Tensor(self.model.predict(self.deval_data)) 

1655 

1656 self.eval_collector.eval_collect(self.eval_pred, self.eval_true) 

1657 result = self.evaluator.evaluate(self.eval_collector.get_data_struct()) 

1658 return result 

1659 

1660 

1661class RaCTTrainer(PretrainTrainer): 

1662 r"""RaCTTrainer is designed for RaCT, which is an actor-critic reinforcement learning based general recommenders. 

1663 It includes three training stages: actor pre-training, critic pre-training and actor-critic training. 

1664 

1665 """ 

1666 

1667 def __init__(self, config, model): 

1668 super().__init__(config, model) 

1669 

1670 def fit( 

1671 self, 

1672 train_data, 

1673 valid_data=None, 

1674 verbose=True, 

1675 saved=True, 

1676 show_progress=False, 

1677 callback_fn=None, 

1678 ): 

1679 if self.model.train_stage == "actor_pretrain": 

1680 return self.pretrain(train_data, verbose, show_progress) 

1681 elif self.model.train_stage == "critic_pretrain": 

1682 return self.pretrain(train_data, verbose, show_progress) 

1683 elif self.model.train_stage == "finetune": 

1684 return super().fit(train_data, valid_data, verbose, saved, show_progress, callback_fn) 

1685 else: 

1686 raise ValueError( 

1687 "Please make sure that the 'train_stage' is 'actor_pretrain', 'critic_pretrain' or 'finetune'!" 

1688 ) 

1689 

1690 

1691class RecVAETrainer(Trainer): 

1692 r"""RecVAETrainer is designed for RecVAE, which is a general recommender.""" 

1693 

1694 def __init__(self, config, model): 

1695 super().__init__(config, model) 

1696 self.n_enc_epochs = config["n_enc_epochs"] 

1697 self.n_dec_epochs = config["n_dec_epochs"] 

1698 

1699 self.optimizer_encoder = self._build_optimizer(params=self.model.encoder.parameters()) 

1700 self.optimizer_decoder = self._build_optimizer(params=self.model.decoder.parameters()) 

1701 

1702 def _train_epoch(self, train_data, epoch_idx, loss_func=None, show_progress=False): 

1703 self.optimizer = self.optimizer_encoder 

1704 

1705 def encoder_loss_func(data): 

1706 return self.model.calculate_loss(data, encoder_flag=True) 

1707 

1708 for epoch in range(self.n_enc_epochs): 

1709 super()._train_epoch( 

1710 train_data, 

1711 epoch_idx, 

1712 loss_func=encoder_loss_func, 

1713 show_progress=show_progress, 

1714 ) 

1715 

1716 self.model.update_prior() 

1717 loss = 0.0 

1718 self.optimizer = self.optimizer_decoder 

1719 

1720 def decoder_loss_func(data): 

1721 return self.model.calculate_loss(data, encoder_flag=False) 

1722 

1723 for epoch in range(self.n_dec_epochs): 

1724 loss += super()._train_epoch( 

1725 train_data, 

1726 epoch_idx, 

1727 loss_func=decoder_loss_func, 

1728 show_progress=show_progress, 

1729 ) 

1730 return loss 

1731 

1732 

1733class NCLTrainer(Trainer): 

1734 def __init__(self, config, model): 

1735 super().__init__(config, model) 

1736 

1737 self.num_m_step = config["m_step"] 

1738 assert self.num_m_step is not None 

1739 

1740 def fit( 

1741 self, 

1742 train_data, 

1743 valid_data=None, 

1744 verbose=True, 

1745 saved=True, 

1746 show_progress=False, 

1747 callback_fn=None, 

1748 ): 

1749 r"""Train the model based on the train data and the valid data. 

1750 

1751 Args: 

1752 train_data (DataLoader): the train data. 

1753 valid_data (DataLoader, optional): the valid data, default: None. 

1754 If it's None, the early_stopping is invalid. 

1755 verbose (bool, optional): whether to write training and evaluation information to logger, default: True 

1756 saved (bool, optional): whether to save the model parameters, default: True 

1757 show_progress (bool): Show the progress of training epoch and evaluate epoch. Defaults to ``False``. 

1758 callback_fn (callable): Optional callback function executed at end of epoch. 

1759 Includes (epoch_idx, valid_score) input arguments. 

1760 

1761 Returns: 

1762 (float, dict): best valid score and best valid result. If valid_data is None, it returns (-1, None) 

1763 """ 

1764 if saved and self.start_epoch >= self.epochs: 

1765 self._save_checkpoint(-1) 

1766 

1767 self.eval_collector.train_data_collect(train_data) 

1768 

1769 for epoch_idx in range(self.start_epoch, self.epochs): 

1770 # only differences from the original trainer 

1771 if epoch_idx % self.num_m_step == 0: 

1772 self.logger.info("Running E-step ! ") 

1773 self.model.e_step() 

1774 # train 

1775 training_start_time = time() 

1776 train_loss = self._train_epoch(train_data, epoch_idx, show_progress=show_progress) 

1777 self.train_loss_dict[epoch_idx] = sum(train_loss) if isinstance(train_loss, tuple) else train_loss 

1778 training_end_time = time() 

1779 train_loss_output = self._generate_train_loss_output( 

1780 epoch_idx, training_start_time, training_end_time, train_loss 

1781 ) 

1782 if verbose: 

1783 self.logger.info(train_loss_output) 

1784 self._add_train_loss_to_tensorboard(epoch_idx, train_loss) 

1785 

1786 # eval 

1787 if self.eval_step <= 0 or not valid_data: 

1788 if saved: 

1789 self._save_checkpoint(epoch_idx) 

1790 update_output = set_color("Saving current", "blue") + ": %s" % self.saved_model_file 

1791 if verbose: 

1792 self.logger.info(update_output) 

1793 continue 

1794 if (epoch_idx + 1) % self.eval_step == 0: 

1795 valid_start_time = time() 

1796 valid_score, valid_result = self._valid_epoch(valid_data, show_progress=show_progress) 

1797 

1798 ( 

1799 self.best_valid_score, 

1800 self.cur_step, 

1801 stop_flag, 

1802 update_flag, 

1803 ) = early_stopping( 

1804 valid_score, 

1805 self.best_valid_score, 

1806 self.cur_step, 

1807 max_step=self.stopping_step, 

1808 bigger=self.valid_metric_bigger, 

1809 ) 

1810 valid_end_time = time() 

1811 valid_score_output = ( 

1812 set_color("epoch %d evaluating", "green") 

1813 + " [" 

1814 + set_color("time", "blue") 

1815 + ": %.2fs, " 

1816 + set_color("valid_score", "blue") 

1817 + ": %f]" 

1818 ) % (epoch_idx, valid_end_time - valid_start_time, valid_score) 

1819 valid_result_output = set_color("valid result", "blue") + ": \n" + dict2str(valid_result) 

1820 if verbose: 

1821 self.logger.info(valid_score_output) 

1822 self.logger.info(valid_result_output) 

1823 self.tensorboard.add_scalar("Valid_score", valid_score, epoch_idx) 

1824 

1825 if update_flag: 

1826 if saved: 

1827 self._save_checkpoint(epoch_idx) 

1828 update_output = set_color("Saving current best", "blue") + ": %s" % self.saved_model_file 

1829 if verbose: 

1830 self.logger.info(update_output) 

1831 self.best_valid_result = valid_result 

1832 

1833 if callback_fn: 

1834 callback_fn(epoch_idx, valid_score) 

1835 

1836 if stop_flag: 

1837 stop_output = "Finished training, best eval result in epoch %d" % ( 

1838 epoch_idx - self.cur_step * self.eval_step 

1839 ) 

1840 if verbose: 

1841 self.logger.info(stop_output) 

1842 break 

1843 self._add_hparam_to_tensorboard(self.best_valid_score) 

1844 return self.best_valid_score, self.best_valid_result 

1845 

1846 def _train_epoch(self, train_data, epoch_idx, loss_func=None, show_progress=False): 

1847 r"""Train the model in an epoch 

1848 Args: 

1849 train_data (DataLoader): The train data. 

1850 epoch_idx (int): The current epoch id. 

1851 loss_func (function): The loss function of :attr:`model`. If it is ``None``, the loss function will be 

1852 :attr:`self.model.calculate_loss`. Defaults to ``None``. 

1853 show_progress (bool): Show the progress of training epoch. Defaults to ``False``. 

1854 

1855 Returns: 

1856 float/tuple: The sum of loss returned by all batches in this epoch. If the loss in each batch contains 

1857 multiple parts and the model return these multiple parts loss instead of the sum of loss, it will return a 

1858 tuple which includes the sum of loss in each part. 

1859 """ 

1860 self.model.train() 

1861 loss_func = loss_func or self.model.calculate_loss 

1862 total_loss = None 

1863 iter_data = ( 

1864 progress_bar( 

1865 train_data, 

1866 total=len(train_data), 

1867 ncols=100, 

1868 desc=set_color(f"Train {epoch_idx:>5}", "magenta", progress=True), 

1869 ) 

1870 if show_progress 

1871 else train_data 

1872 ) 

1873 scaler = grad_scaler(enabled=self.enable_scaler) 

1874 

1875 if not self.config["single_spec"] and train_data.shuffle: 

1876 train_data.sampler.set_epoch(epoch_idx) 

1877 

1878 for batch_idx, batch_interaction in enumerate(iter_data): 

1879 interaction = batch_interaction.to(self.device) 

1880 self.optimizer.zero_grad() 

1881 sync_loss = 0 

1882 if not self.config["single_spec"]: 

1883 self.set_reduce_hook() 

1884 sync_loss = self.sync_grad_loss() 

1885 

1886 with autocast(device_type=self.device.type, enabled=self.enable_amp): 

1887 losses = loss_func(interaction) 

1888 

1889 if isinstance(losses, tuple): 

1890 if epoch_idx < self.config["warm_up_step"]: 

1891 losses = losses[:-1] 

1892 loss = sum(losses) 

1893 loss_tuple = tuple(per_loss.item() for per_loss in losses) 

1894 total_loss = loss_tuple if total_loss is None else tuple(map(sum, zip(total_loss, loss_tuple))) 

1895 else: 

1896 loss = losses 

1897 total_loss = losses.item() if total_loss is None else total_loss + losses.item() 

1898 self._check_nan(loss) 

1899 scaler.scale(loss + sync_loss).backward() 

1900 

1901 if self.clip_grad_norm: 

1902 clip_grad_norm_(self.model.parameters(), **self.clip_grad_norm) 

1903 scaler.step(self.optimizer) 

1904 scaler.update() 

1905 if self.gpu_available and show_progress: 

1906 iter_data.set_postfix_str(set_color("GPU RAM: " + get_gpu_usage(self.device), "yellow")) 

1907 return total_loss 

1908 

1909 

1910class PEARLMfromscratchTrainer(ExplainableTrainer): 

1911 def __init__(self, config, model): 

1912 super().__init__(config, model) 

1913 

1914 self.path_generation_args = config["path_generation_args"] 

1915 

1916 def _full_sort_batch_eval(self, batched_data, tot_item_num, item_tensor): 

1917 return super()._full_sort_batch_eval(batched_data, tot_item_num, item_tensor, **self.path_generation_args) 

1918 

1919 

1920class HFPathLanguageModelingTrainer(ExplainableTrainer): 

1921 r"""HFPathLanguageModelingTrainer is designed for path-based knowledge-aware recommendation methods. 

1922 It is specifically designed to communicate with the Hugging Face Trainer to use language models and functionalities 

1923 as tokenizers and beam search. 

1924 """ 

1925 

1926 HOPWISE_SAVE_PATH_SUFFIX = "hopwise-" 

1927 HUGGINGFACE_SAVE_PATH_SUFFIX = "huggingface-" 

1928 

1929 def __init__(self, config, model): 

1930 super().__init__(config, model) 

1931 

1932 self.path_generation_args = self.config["path_generation_args"] 

1933 

1934 self.HOPWISE_SAVE_PATH_SUFFIX += f"{config['base_model']}-" 

1935 self.HUGGINGFACE_SAVE_PATH_SUFFIX += f"{config['base_model']}-" 

1936 

1937 dirname, basename = os.path.split(self.saved_model_file) 

1938 self.saved_model_file = os.path.join(dirname, self.HOPWISE_SAVE_PATH_SUFFIX + basename) 

1939 

1940 def prepare_hf_args(self, **kwargs): 

1941 from transformers import TrainingArguments 

1942 

1943 output_dir = self.saved_model_file.replace(self.HOPWISE_SAVE_PATH_SUFFIX, self.HUGGINGFACE_SAVE_PATH_SUFFIX) 

1944 

1945 hf_args = dict( 

1946 output_dir=output_dir, 

1947 eval_strategy="epoch", 

1948 save_strategy="epoch", 

1949 eval_steps=self.eval_step, 

1950 learning_rate=self.learning_rate, 

1951 weight_decay=self.weight_decay, 

1952 bf16=False, 

1953 fp16=self.enable_amp, 

1954 num_train_epochs=self.epochs, 

1955 per_device_train_batch_size=self.config["train_batch_size"], 

1956 per_device_eval_batch_size=self.test_batch_size, 

1957 warmup_steps=self.config["warmup_steps"], 

1958 save_steps=self.eval_step, 

1959 save_total_limit=1, 

1960 load_best_model_at_end=True, 

1961 metric_for_best_model=self.valid_metric, 

1962 greater_is_better=self.valid_metric_bigger, 

1963 seed=self.config["seed"], 

1964 report_to="none", 

1965 ) 

1966 hf_args.update(kwargs) 

1967 return TrainingArguments(**hf_args) 

1968 

1969 def init_hf_trainer( 

1970 self, 

1971 train_data, 

1972 valid_data=None, 

1973 verbose=True, 

1974 saved=True, 

1975 show_progress=False, 

1976 hf_callbacks=None, 

1977 callback_fn=None, 

1978 training_args=None, 

1979 ): 

1980 from hopwise.trainer.hf_path_trainer import HFPathTrainer, hopwiseCallback 

1981 

1982 training_args = training_args or {} 

1983 

1984 hf_callbacks = hf_callbacks or [] 

1985 hf_args = self.prepare_hf_args(**training_args) 

1986 

1987 callbacks = [ 

1988 hopwiseCallback( 

1989 self, 

1990 train_data, 

1991 valid_data=valid_data, 

1992 verbose=verbose, 

1993 saved=saved, 

1994 show_progress=show_progress, 

1995 callback_fn=callback_fn, 

1996 model=self.model, 

1997 model_name=self.model.__class__.__name__, 

1998 ), 

1999 *hf_callbacks, 

2000 ] 

2001 

2002 self.hf_trainer = HFPathTrainer(self.model, callbacks, train_data=train_data, args=hf_args) 

2003 

2004 @property 

2005 def processing_class(self): 

2006 if hasattr(self, "hf_trainer"): 

2007 return self.hf_trainer.processing_class 

2008 return None 

2009 

2010 def _save_checkpoint(self, epoch, verbose=True, **kwargs): 

2011 r"""Store the model parameters information and training information. 

2012 

2013 Args: 

2014 epoch (int): the current epoch id 

2015 

2016 """ 

2017 if not self.config["single_spec"] and self.config["local_rank"] != 0: 

2018 return 

2019 saved_model_file = kwargs.pop("saved_model_file", self.saved_model_file) 

2020 state = { 

2021 "config": self.config, 

2022 "epoch": epoch, 

2023 "cur_step": self.cur_step, 

2024 "best_valid_score": self.best_valid_score, 

2025 } 

2026 torch.save(state, saved_model_file, pickle_protocol=4) 

2027 if verbose: 

2028 self.logger.info(set_color("Saving current", "blue") + f": {saved_model_file}") 

2029 hf_output_dir = self.hf_trainer.args.output_dir 

2030 self.logger.info(set_color("HuggingFace model is saved at", "blue") + f": {hf_output_dir}") 

2031 

2032 def resume_checkpoint(self, resume_file): 

2033 """ 

2034 Load the model parameters and training information based on the directory name, 

2035 and navigate into subdirectories if necessary. 

2036 Also handles both HuggingFace and hopwise formats by reading corresponding files. 

2037 

2038 Args: 

2039 resume_file (str): the path to the directory containing the checkpoint files or subdirectories 

2040 """ 

2041 from safetensors.torch import load_file 

2042 from transformers import AutoTokenizer 

2043 

2044 if not hasattr(self, "hf_trainer"): 

2045 raise ValueError("The HuggingFace Trainer has not been initialized. Please call `init_hf_trainer` first.") 

2046 

2047 if os.path.basename(resume_file).startswith(self.HUGGINGFACE_SAVE_PATH_SUFFIX): 

2048 hf_resume_file = resume_file 

2049 hopwise_resume_file = resume_file.replace(self.HUGGINGFACE_SAVE_PATH_SUFFIX, self.HOPWISE_SAVE_PATH_SUFFIX) 

2050 elif os.path.basename(resume_file).startswith(self.HOPWISE_SAVE_PATH_SUFFIX): 

2051 hopwise_resume_file = resume_file 

2052 hf_resume_file = resume_file.replace(self.HOPWISE_SAVE_PATH_SUFFIX, self.HUGGINGFACE_SAVE_PATH_SUFFIX) 

2053 else: 

2054 raise ValueError(f"The directory name [{resume_file}] does not indicate a HuggingFace or hopwise model.") 

2055 

2056 checkpoint = torch.load(hopwise_resume_file, map_location=self.device, weights_only=False) 

2057 self.start_epoch = checkpoint["epoch"] + 1 

2058 self.cur_step = checkpoint["cur_step"] 

2059 self.best_valid_score = checkpoint["best_valid_score"] 

2060 

2061 weights = load_file(os.path.join(hf_resume_file, "model.safetensors")) 

2062 self.model.load_state_dict(weights, strict=False) 

2063 self.processing_class.tokenizer = AutoTokenizer.from_pretrained(hf_resume_file) 

2064 

2065 def fit( 

2066 self, 

2067 train_data, 

2068 valid_data=None, 

2069 verbose=True, 

2070 saved=True, 

2071 show_progress=False, 

2072 callback_fn=None, 

2073 ): 

2074 self.eval_collector.train_data_collect(train_data) 

2075 

2076 if not hasattr(self, "hf_trainer"): 

2077 self.init_hf_trainer( 

2078 train_data, 

2079 valid_data=valid_data, 

2080 verbose=verbose, 

2081 saved=saved, 

2082 show_progress=show_progress, 

2083 callback_fn=callback_fn, 

2084 ) 

2085 

2086 self.hf_trainer.train() 

2087 self.hf_trainer.save_model() 

2088 

2089 return self.best_valid_score, self.best_valid_result 

2090 

2091 def _full_sort_batch_eval(self, batched_data, tot_item_num, item_tensor): 

2092 return super()._full_sort_batch_eval( 

2093 batched_data, 

2094 tot_item_num, 

2095 item_tensor, 

2096 return_dict_in_generate=True, 

2097 output_scores=True, 

2098 **self.path_generation_args, 

2099 ) 

2100 

2101 @torch.no_grad() 

2102 def evaluate(self, eval_data, load_best_model=True, model_file=None, show_progress=False): 

2103 if not eval_data: 

2104 return 

2105 

2106 if load_best_model: 

2107 self.hf_trainer._load_best_model() 

2108 best_model_checkpoint_path = self.hf_trainer.state.best_model_checkpoint 

2109 message_output = f"Loading model structure and parameters from {best_model_checkpoint_path}" 

2110 self.logger.info(message_output) 

2111 

2112 return super().evaluate(eval_data, load_best_model=False, model_file=None, show_progress=show_progress) 

2113 

2114 

2115class KGGLMTrainer(HFPathLanguageModelingTrainer, PretrainTrainer): 

2116 r"""KGGLMTrainer is designed for KGGLM, which is a path-based language model for knowledge-aware recommendation. 

2117 It includes two training stages: link prediction pre-training and recommendation path generation fine-tuning. 

2118 """ 

2119 

2120 def _get_pretrained_model_path(self, epoch_label=None): 

2121 epoch_label = f"pretrained-{epoch_label}" if epoch_label is not None else "pretrained" 

2122 return os.path.join( 

2123 self.checkpoint_dir, 

2124 self.HUGGINGFACE_SAVE_PATH_SUFFIX 

2125 + "{}-{}-{}.pth".format(self.config["model"], self.config["dataset"], epoch_label), 

2126 ) 

2127 

2128 def pretrain(self, train_data, verbose=True, show_progress=False): 

2129 from transformers import TrainerCallback 

2130 

2131 pretrain_path = self._get_pretrained_model_path() 

2132 pretrain_args = dict( 

2133 output_dir=pretrain_path, 

2134 num_train_epochs=self.pretrain_epochs, 

2135 save_steps=self.save_step, 

2136 eval_strategy="no", 

2137 load_best_model_at_end=False, 

2138 ) 

2139 

2140 class PretrainSaveCallback(TrainerCallback): 

2141 def __init__(self, hopwise_trainer): 

2142 self.hopwise_trainer = hopwise_trainer 

2143 

2144 def on_epoch_end(self, args, state, control, **kwargs): 

2145 if control.should_save: 

2146 epoch_idx = int(state.epoch) 

2147 pretrain_path = self.hopwise_trainer._get_pretrained_model_path(epoch_idx) 

2148 self.hopwise_trainer.hf_trainer.args.output_dir = pretrain_path 

2149 

2150 self.init_hf_trainer( 

2151 train_data, 

2152 verbose=verbose, 

2153 saved=True, 

2154 show_progress=show_progress, 

2155 hf_callbacks=[PretrainSaveCallback(self)], 

2156 training_args=pretrain_args, 

2157 ) 

2158 

2159 self.hf_trainer.train() 

2160 

2161 return self.best_valid_score, self.best_valid_result 

2162 

2163 def evaluate(self, eval_data, load_best_model=True, model_file=None, show_progress=False): 

2164 if load_best_model and self.model.train_stage == "pretrain": 

2165 self.hf_trainer.state.best_model_checkpoint = self.hf_trainer.args.output_dir 

2166 

2167 return super().evaluate( 

2168 eval_data, 

2169 load_best_model=load_best_model, 

2170 model_file=model_file, 

2171 show_progress=show_progress, 

2172 ) 

2173 

2174 def fit( 

2175 self, 

2176 train_data, 

2177 valid_data=None, 

2178 verbose=True, 

2179 saved=True, 

2180 show_progress=False, 

2181 callback_fn=None, 

2182 ): 

2183 if self.model.train_stage == "pretrain": 

2184 return self.pretrain(train_data, verbose, show_progress) 

2185 elif self.model.train_stage == "finetune": 

2186 return super().fit(train_data, valid_data, verbose, saved, show_progress, callback_fn) 

2187 else: 

2188 raise ValueError(f"Please make sure that the 'train_stage' is in [{self.model.TRAIN_STAGES}]!")