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

1# @Time : 2025/06 

2# @Author : Giacomo Medda, Alessandro Soccol 

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

4 

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. 

10 

11Reference code: 

12 https://github.com/mirkomarras/kgglm 

13""" 

14 

15import os 

16 

17from transformers.trainer_utils import get_last_checkpoint 

18 

19from hopwise.model.path_language_modeling_recommender.pearlm import PEARLM 

20 

21 

22class KGGLM(PEARLM): 

23 TRAIN_STAGES = ["pretrain", "finetune"] 

24 

25 def __init__(self, config, dataset): 

26 super().__init__(config, dataset) 

27 

28 self.train_stage = config["train_stage"] 

29 self.pre_model_path = config["pre_model_path"] 

30 

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) 

37 

38 from safetensors.torch import load_file 

39 

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)