Coverage for hopwise/model/path_language_modeling_recommender/kgglm.py: 100%
17 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 : 2025/06
2# @Author : Giacomo Medda, Alessandro Soccol
3# @Email : giacomo.medda@unica.it, alessandro.soccol@unica.it
5r"""KGGLM
6##################################################
7Reference:
8 Balloccu et al. "KGGLM: A Generative Language Model for Generalizable Knowledge
9 Graph Representation Learning in Recommendation." in RecSys 2024.
11Reference code:
12 https://github.com/mirkomarras/kgglm
13"""
15import os
17from transformers.trainer_utils import get_last_checkpoint
19from hopwise.model.path_language_modeling_recommender.pearlm import PEARLM
22class KGGLM(PEARLM):
23 TRAIN_STAGES = ["pretrain", "finetune"]
25 def __init__(self, config, dataset):
26 super().__init__(config, dataset)
28 self.train_stage = config["train_stage"]
29 self.pre_model_path = config["pre_model_path"]
31 assert self.train_stage in self.TRAIN_STAGES
32 if self.train_stage == "finetune":
33 # load pretrained model for finetune
34 if not os.path.exists(os.path.join(self.pre_model_path, "config.json")):
35 # if the path is not a checkpoint, we assume it contains the checkpoint
36 self.pre_model_path = get_last_checkpoint(self.pre_model_path)
38 from safetensors.torch import load_file
40 self.logger.info(f"Load pretrained model from {self.pre_model_path}")
41 weights = load_file(os.path.join(self.pre_model_path, "model.safetensors"))
42 self.load_state_dict(weights, strict=False)