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
« 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
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
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
15# UPDATE:
16# @Time : 2025
17# @Author : Giacomo Medda, Alessandro Soccol
18# @Email : giacomo.medda@unica.it, alessandro.soccol@unica.it
20"""hopwise.trainer.trainer
21################################
22"""
24import os
25from collections import defaultdict
26from logging import getLogger
27from time import time
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_
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)
55try:
56 grad_scaler = torch.GradScaler
57 autocast = torch.autocast
58except AttributeError:
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)
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)
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 """
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"]])
86 def fit(self, train_data):
87 r"""Train the model based on the train data."""
88 raise NotImplementedError("Method [next] should be implemented.")
90 def evaluate(self, eval_data):
91 r"""Evaluate the model based on the eval data."""
92 raise NotImplementedError("Method [next] should be implemented.")
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.
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
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.
108 """
109 sync_loss = 0
110 for params in self.model.parameters():
111 sync_loss += torch.sum(params) * 0
112 return sync_loss
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.
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.
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.
128 """
130 def __init__(self, config, model):
131 super().__init__(config, model)
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"]
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)
165 def _build_optimizer(self, **kwargs):
166 r"""Init the Optimizer
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``.
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)
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 )
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
208 def _train_epoch(self, train_data, epoch_idx, loss_func=None, show_progress=False):
209 r"""Train the model in an epoch.
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``.
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 )
238 if not self.config["single_spec"] and train_data.shuffle:
239 train_data.sampler.set_epoch(epoch_idx)
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()
250 with autocast(device_type=self.device.type, enabled=self.enable_amp):
251 losses = loss_func(interaction)
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
270 def _valid_epoch(self, valid_data, show_progress=False):
271 r"""Valid the model with valid data
273 Args:
274 valid_data (DataLoader): the valid data.
275 show_progress (bool): Show the progress of evaluate epoch. Defaults to ``False``.
277 Returns:
278 float: valid score
279 dict: valid result
280 """
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
286 def _save_checkpoint(self, epoch, verbose=True, **kwargs):
287 r"""Store the model parameters information and training information.
289 Args:
290 epoch (int): the current epoch id
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}")
309 def resume_checkpoint(self, resume_file):
310 r"""Load the model parameters information and training information.
312 Args:
313 resume_file (file): the checkpoint file
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"]
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"))
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)
337 def _check_nan(self, loss):
338 if torch.isnan(loss):
339 raise ValueError("Training loss is nan")
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 + "]"
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)
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)
380 self.tensorboard.add_hparams(hparam_dict, {"hparam/best_valid_result": best_valid_result})
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.
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.
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)
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
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 )
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
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")
460 if update_flag:
461 if saved:
462 self._save_checkpoint(epoch_idx, verbose=verbose)
463 self.best_valid_result = valid_result
465 if callback_fn:
466 callback_fn(epoch_idx, valid_score)
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
476 valid_step += 1
478 self._add_hparam_to_tensorboard(self.best_valid_score)
479 return self.best_valid_score, self.best_valid_result
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)
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)
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
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)
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
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.
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``.
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
543 # self.eval_collector.eval_data_collect(eval_data)
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)
557 self.model.eval()
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)
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
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
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)
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.
632 """
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
643 if config["metrics_lp"]:
644 self.eval_collector_kg = Collector_KG(config)
645 self.evaluator_kg = Evaluator_KG(config)
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
668 def _valid_epoch(self, valid_data, show_progress=False):
669 r"""Valid the model with valid data
671 Args:
672 valid_data (Dataloader, list[Dataloader]): the valid data.
673 show_progress (bool): Show the progress of evaluate epoch. Defaults to ``False``.
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]}
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
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 )
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)
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
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)
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.
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``.
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
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)
779 self.model.eval()
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
790 task_eval_data = {kg_eval_type: eval_data}
792 # REC task
793 if KnowledgeEvaluationType.REC in task_eval_data:
794 task = KnowledgeEvaluationType.REC
795 rec_eval_data = task_eval_data[task]
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)
803 results[task] = self.evaluate_data_loop(
804 rec_eval_data, task, tot_item_num, item_tensor, show_progress=show_progress
805 )
807 # LP task
808 if KnowledgeEvaluationType.LP in task_eval_data:
809 task = KnowledgeEvaluationType.LP
810 kg_eval_data = task_eval_data[task]
812 tot_entity_num = kg_eval_data._dataset.entity_num
813 tail_tensor = kg_eval_data._dataset.get_tail_feature().to(self.device)
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 )
823 if isinstance(eval_data, list):
824 return results
825 else:
826 return results[kg_eval_type]
828 def evaluate_data_loop(self, eval_data, task, tot_target_num, target_tensor, show_progress=True):
829 neg_sampling = isinstance(eval_data, NegSampleDataLoader)
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
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 )
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
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.
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.
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)
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
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 )
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
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)
926 best_valid = defaultdict(dict)
928 if KnowledgeEvaluationType.LP in return_data:
929 kg_valid_scores = list()
930 kg_valid_results = list()
932 update_flag = False
933 stop_flag = False
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 )
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)
976 if verbose:
977 self.logger.info(valid_score_output)
978 self.logger.info(valid_result_output)
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")
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
988 if callback_fn:
989 callback_fn(epoch_idx, valid_score)
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
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
1010 best_valid["score"][KnowledgeEvaluationType.REC] = self.best_valid_score
1011 best_valid["result"][KnowledgeEvaluationType.REC] = self.best_valid_result
1013 self._add_hparam_to_tensorboard(self.best_valid_score)
1014 return best_valid["score"], best_valid["result"]
1017class ExplainableTrainer(Trainer):
1018 """ExplainableTrainer is designed for explainable recommendation methods."""
1020 def __init__(self, config, model):
1021 super().__init__(config, model)
1022 self.eval_collector = ExplainableCollector(config)
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
1027 scores, paths = self.model.explain(interaction.to(self.device), **kwargs)
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
1034 return interaction, (scores, paths), positive_u, positive_i
1037class PGPRTrainer(ExplainableTrainer):
1038 r"""PGPRTrainer is designed for PGPR, which is a knowledge-aware recommendation method."""
1040 def __init__(self, config, model):
1041 super().__init__(config, model)
1044class CAFETrainer(ExplainableTrainer):
1045 r"""CAFETrainer is designed for CAFE, which is a knowledge-aware recommendation method."""
1047 def __init__(self, config, model):
1048 super().__init__(config, model)
1051class KGATTrainer(Trainer):
1052 r"""KGATTrainer is designed for KGAT, which is a knowledge-aware recommendation method."""
1054 def __init__(self, config, model):
1055 super().__init__(config, model)
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)
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 )
1073 # update A
1074 self.model.eval()
1075 with torch.no_grad():
1076 self.model.update_attentive_A()
1078 return rs_total_loss, kg_total_loss
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 """
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"]
1091 def save_pretrained_model(self, epoch, saved_model_file):
1092 r"""Store the model parameters information and training information.
1094 Args:
1095 epoch (int): the current epoch id
1096 saved_model_file (str): file name for saved pretrained model
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
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 )
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)
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)
1137 return self.best_valid_score, self.best_valid_result
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.
1144 """
1146 def __init__(self, config, model):
1147 super().__init__(config, model)
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'!")
1166class TPRecTrainer(PretrainTrainer):
1167 """
1168 TPRecTrainer is designed for TPRec, which is a knowledge-aware recommendation method.
1169 """
1171 def __init__(self, config, model):
1172 super().__init__(config, model)
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'!")
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)
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
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.
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``.
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
1234 # self.eval_collector.eval_data_collect(eval_data)
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)
1244 self.model.eval()
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)
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)
1266 if self.model.train_stage == "policy":
1267 batched_data = (batched_data, eval_data.temporal_weights) # noqa: PLW2901
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
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)
1287 paths = None
1288 batched_data, temporal_weights = batched_data
1290 interaction, history_index, positive_u, positive_i = batched_data
1292 # Note: interaction without item ids
1293 scores = self.model.full_sort_predict((interaction.to(self.device), temporal_weights))
1295 if isinstance(scores, tuple):
1296 # then the first is the score, the second are paths
1297 scores, paths = scores
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
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
1310class MKRTrainer(Trainer):
1311 r"""MKRTrainer is designed for MKR, which is a knowledge-aware recommendation method."""
1313 def __init__(self, config, model):
1314 super().__init__(config, model)
1315 self.kge_interval = config["kge_interval"]
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
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 )
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 )
1341 return rs_total_loss, kg_total_loss
1344class TraditionalTrainer(Trainer):
1345 """TraditionalTrainer is designed for Traditional model(Pop,ItemKNN),
1346 which set the epoch to 1 whatever the config."""
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
1353class DecisionTreeTrainer(AbstractTrainer):
1354 """DecisionTreeTrainer is designed for DecisionTree model."""
1356 def __init__(self, config, model):
1357 super().__init__(config, model)
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"]
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)
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)
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)
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)
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
1390 def _interaction_to_sparse(self, dataloader):
1391 r"""Convert data format from interaction to sparse or numpy
1393 Args:
1394 dataloader (DecisionTreeDataLoader): DecisionTreeDataLoader dataloader.
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))
1413 if self.convert_token_to_onehot:
1414 convert_col_list = dataloader._dataset.convert_col_list
1415 hash_count = dataloader._dataset.hash_count
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))
1422 cur_j = 0
1423 new_j = 0
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
1436 cur_data = sparse.csc_matrix(onehot_data)
1438 return cur_data, interaction_np[self.label_field]
1440 def _interaction_to_lib_datatype(self, dataloader):
1441 pass
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
1451 def _save_checkpoint(self, epoch):
1452 r"""Store the model parameters information and training information.
1454 Args:
1455 epoch (int): the current epoch id
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)
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)
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)
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 )
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)
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
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
1521 return self.best_valid_score, self.best_valid_result
1523 def evaluate(self, eval_data, load_best_model=True, model_file=None, show_progress=False):
1524 raise NotImplementedError
1526 def _train_at_once(self, train_data, valid_data):
1527 raise NotImplementedError
1530class XGBoostTrainer(DecisionTreeTrainer):
1531 """XGBoostTrainer is designed for XGBOOST."""
1533 def __init__(self, config, model):
1534 super().__init__(config, model)
1536 self.xgb = __import__("xgboost")
1537 self.boost_model = config["xgb_model"]
1538 self.silent = config["xgb_silent"]
1539 self.nthread = config["xgb_nthread"]
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
1552 def _interaction_to_lib_datatype(self, dataloader):
1553 r"""Convert data format from interaction to DMatrix
1555 Args:
1556 dataloader (DecisionTreeDataLoader): xgboost dataloader.
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)
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 )
1584 self.model.save_model(self.temp_file)
1585 self.boost_model = self.temp_file
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)
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))
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
1604class LightGBMTrainer(DecisionTreeTrainer):
1605 """LightGBMTrainer is designed for LightGBM."""
1607 def __init__(self, config, model):
1608 super().__init__(config, model)
1610 self.lgb = __import__("lightgbm")
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
1619 def _interaction_to_lib_datatype(self, dataloader):
1620 r"""Convert data format from interaction to Dataset
1622 Args:
1623 dataloader (DecisionTreeDataLoader): xgboost dataloader.
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)
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)
1641 self.model.save_model(self.temp_file)
1642 self.boost_model = self.temp_file
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)
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))
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
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.
1665 """
1667 def __init__(self, config, model):
1668 super().__init__(config, model)
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 )
1691class RecVAETrainer(Trainer):
1692 r"""RecVAETrainer is designed for RecVAE, which is a general recommender."""
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"]
1699 self.optimizer_encoder = self._build_optimizer(params=self.model.encoder.parameters())
1700 self.optimizer_decoder = self._build_optimizer(params=self.model.decoder.parameters())
1702 def _train_epoch(self, train_data, epoch_idx, loss_func=None, show_progress=False):
1703 self.optimizer = self.optimizer_encoder
1705 def encoder_loss_func(data):
1706 return self.model.calculate_loss(data, encoder_flag=True)
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 )
1716 self.model.update_prior()
1717 loss = 0.0
1718 self.optimizer = self.optimizer_decoder
1720 def decoder_loss_func(data):
1721 return self.model.calculate_loss(data, encoder_flag=False)
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
1733class NCLTrainer(Trainer):
1734 def __init__(self, config, model):
1735 super().__init__(config, model)
1737 self.num_m_step = config["m_step"]
1738 assert self.num_m_step is not None
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.
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.
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)
1767 self.eval_collector.train_data_collect(train_data)
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)
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)
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)
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
1833 if callback_fn:
1834 callback_fn(epoch_idx, valid_score)
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
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``.
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)
1875 if not self.config["single_spec"] and train_data.shuffle:
1876 train_data.sampler.set_epoch(epoch_idx)
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()
1886 with autocast(device_type=self.device.type, enabled=self.enable_amp):
1887 losses = loss_func(interaction)
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()
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
1910class PEARLMfromscratchTrainer(ExplainableTrainer):
1911 def __init__(self, config, model):
1912 super().__init__(config, model)
1914 self.path_generation_args = config["path_generation_args"]
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)
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 """
1926 HOPWISE_SAVE_PATH_SUFFIX = "hopwise-"
1927 HUGGINGFACE_SAVE_PATH_SUFFIX = "huggingface-"
1929 def __init__(self, config, model):
1930 super().__init__(config, model)
1932 self.path_generation_args = self.config["path_generation_args"]
1934 self.HOPWISE_SAVE_PATH_SUFFIX += f"{config['base_model']}-"
1935 self.HUGGINGFACE_SAVE_PATH_SUFFIX += f"{config['base_model']}-"
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)
1940 def prepare_hf_args(self, **kwargs):
1941 from transformers import TrainingArguments
1943 output_dir = self.saved_model_file.replace(self.HOPWISE_SAVE_PATH_SUFFIX, self.HUGGINGFACE_SAVE_PATH_SUFFIX)
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)
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
1982 training_args = training_args or {}
1984 hf_callbacks = hf_callbacks or []
1985 hf_args = self.prepare_hf_args(**training_args)
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 ]
2002 self.hf_trainer = HFPathTrainer(self.model, callbacks, train_data=train_data, args=hf_args)
2004 @property
2005 def processing_class(self):
2006 if hasattr(self, "hf_trainer"):
2007 return self.hf_trainer.processing_class
2008 return None
2010 def _save_checkpoint(self, epoch, verbose=True, **kwargs):
2011 r"""Store the model parameters information and training information.
2013 Args:
2014 epoch (int): the current epoch id
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}")
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.
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
2044 if not hasattr(self, "hf_trainer"):
2045 raise ValueError("The HuggingFace Trainer has not been initialized. Please call `init_hf_trainer` first.")
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.")
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"]
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)
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)
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 )
2086 self.hf_trainer.train()
2087 self.hf_trainer.save_model()
2089 return self.best_valid_score, self.best_valid_result
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 )
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
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)
2112 return super().evaluate(eval_data, load_best_model=False, model_file=None, show_progress=show_progress)
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 """
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 )
2128 def pretrain(self, train_data, verbose=True, show_progress=False):
2129 from transformers import TrainerCallback
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 )
2140 class PretrainSaveCallback(TrainerCallback):
2141 def __init__(self, hopwise_trainer):
2142 self.hopwise_trainer = hopwise_trainer
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
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 )
2159 self.hf_trainer.train()
2161 return self.best_valid_score, self.best_valid_result
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
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 )
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}]!")