Coverage for hopwise/model/path_language_modeling_recommender/pearlm.py: 93%

57 statements  

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

1# @Time : 2025 

2# @Author : Giacomo Medda 

3# @Email : giacomo.medda@unica.it 

4 

5r"""PEARLM 

6################################################## 

7Reference: 

8 Balloccu et al. "Faithful Path Language Modeling for Explainable Recommendation over Knowledge Graph." - preprint. 

9 

10Reference code: 

11 https://github.com/Chris1nexus/pearlm 

12""" 

13 

14from typing import Optional, Union 

15 

16import torch 

17from torch import nn 

18from transformers import AutoConfig, GPT2LMHeadModel 

19from transformers.modeling_outputs import CausalLMOutputWithCrossAttentions 

20 

21from hopwise.data import Interaction 

22from hopwise.model.abstract_recommender import ExplainablePathLanguageModelingRecommender 

23from hopwise.utils import PathLanguageModelingTokenType 

24 

25 

26class PEARLM(ExplainablePathLanguageModelingRecommender, GPT2LMHeadModel): 

27 """PEARLM is a path-language-modeling recommender. It learns the sequence of entity-relation triplets 

28 as paths extracted from a knowledge graph. It is trained to predict the next token in a sequence of tokens 

29 representing a path. The model extends PLM by adding a constrained graph decoding mechanism to ensure that 

30 the generated paths are valid according to the knowledge graph structure. The model can be used for 

31 explainable recommendation by generating paths that explain the recommendations made by the model. 

32 """ 

33 

34 def __init__(self, config, dataset): 

35 ExplainablePathLanguageModelingRecommender.__init__(self, config, dataset) 

36 

37 self.use_kg_token_types = config["use_kg_token_types"] 

38 

39 transformers_config = AutoConfig.from_pretrained( 

40 "distilgpt2", 

41 **{ 

42 "vocab_size": self.n_tokens, 

43 "n_ctx": config["context_length"], 

44 "n_positions": config["context_length"], 

45 "pad_token_id": dataset.tokenizer.pad_token_id, 

46 "bos_token_id": dataset.tokenizer.bos_token_id, 

47 "eos_token_id": dataset.tokenizer.eos_token_id, 

48 "n_embd": config["embedding_size"], 

49 "n_head": config["num_heads"], 

50 "n_layer": config["num_layers"], 

51 }, 

52 ) 

53 GPT2LMHeadModel.__init__(self, transformers_config) 

54 self.to(config["device"]) 

55 

56 # Add type ids template 

57 if self.use_kg_token_types: 

58 prev_vocab_size = self.n_tokens 

59 spec_type, spec_type_id = PathLanguageModelingTokenType.SPECIAL.value 

60 ent_type, ent_type_id = PathLanguageModelingTokenType.ENTITY.value 

61 rel_type, rel_type_id = PathLanguageModelingTokenType.RELATION.value 

62 

63 token_types = [f"<Token-Type.{token_type}>" for token_type in [spec_type, ent_type, rel_type]] 

64 for token_type in token_types: 

65 dataset.tokenizer.add_tokens(token_type) 

66 self.n_tokens = len(dataset.tokenizer) # Update the vocabulary size after adding new tokens 

67 

68 spec_type_id, ent_type_id, rel_type_id = ( 

69 spec_type_id + prev_vocab_size, 

70 ent_type_id + prev_vocab_size, 

71 rel_type_id + prev_vocab_size, 

72 ) 

73 self.token_type_ids = torch.LongTensor( 

74 # BOS + ENT + REL + ENT + REL + ... + ENT + REL + EOS 

75 [spec_type_id, ent_type_id] + [rel_type_id, ent_type_id] * dataset.path_hop_length + [spec_type_id] 

76 ) 

77 self.token_type_ids = self.token_type_ids.to(config["device"]) 

78 

79 self.transformer.resize_token_embeddings(self.n_tokens) 

80 

81 self.loss = nn.CrossEntropyLoss() 

82 self.post_init() 

83 

84 def forward( 

85 self, 

86 input_ids: Optional[torch.LongTensor] = None, 

87 past_key_values: Optional[tuple[tuple[torch.Tensor]]] = None, 

88 attention_mask: Optional[torch.FloatTensor] = None, 

89 token_type_ids: Optional[torch.LongTensor] = None, 

90 position_ids: Optional[torch.LongTensor] = None, 

91 head_mask: Optional[torch.FloatTensor] = None, 

92 inputs_embeds: Optional[torch.FloatTensor] = None, 

93 encoder_hidden_states: Optional[torch.Tensor] = None, 

94 encoder_attention_mask: Optional[torch.FloatTensor] = None, 

95 labels: Optional[torch.LongTensor] = None, 

96 use_cache: Optional[bool] = None, 

97 output_attentions: Optional[bool] = None, 

98 output_hidden_states: Optional[bool] = None, 

99 return_dict: Optional[bool] = None, 

100 **kwargs, # Additional arguments for compatibility with HuggingFace Trainer 

101 ) -> Union[tuple, CausalLMOutputWithCrossAttentions]: 

102 if isinstance(input_ids, Interaction): 

103 token_type_ids = input_ids.get("token_type_ids", None) 

104 attention_mask = input_ids.get("attention_mask", None) 

105 input_ids = input_ids["input_ids"] 

106 

107 if self.use_kg_token_types: 

108 # indexing is only relevant during generation to match token types length with input_ids 

109 token_type_ids = self.token_type_ids[[*range(input_ids.shape[1] - 1), -1]] 

110 token_type_ids = token_type_ids.expand(input_ids.shape[0], -1) 

111 

112 transformer_outputs = self.transformer( 

113 input_ids, 

114 past_key_values=past_key_values, 

115 attention_mask=attention_mask, 

116 token_type_ids=token_type_ids, 

117 position_ids=position_ids, 

118 head_mask=head_mask, 

119 inputs_embeds=inputs_embeds, 

120 encoder_hidden_states=encoder_hidden_states, 

121 encoder_attention_mask=encoder_attention_mask, 

122 use_cache=use_cache and labels is None, 

123 output_attentions=output_attentions, 

124 output_hidden_states=output_hidden_states, 

125 return_dict=return_dict or self.config.use_return_dict, 

126 ) 

127 

128 sequence_output = transformer_outputs[0] 

129 prediction_scores = self.lm_head(sequence_output) 

130 

131 lm_loss = None 

132 if labels is not None: 

133 # we are doing next-token prediction; shift prediction scores and input ids by one 

134 shifted_prediction_scores = prediction_scores[:, :-1, :].contiguous() 

135 labels = labels[:, 1:].contiguous() 

136 lm_loss = self.loss(shifted_prediction_scores.view(-1, self.config.vocab_size), labels.view(-1)) 

137 

138 if not return_dict: 

139 output = (prediction_scores,) + transformer_outputs[2:] 

140 return ((lm_loss,) + output) if lm_loss is not None else output 

141 

142 return CausalLMOutputWithCrossAttentions( 

143 loss=lm_loss, 

144 logits=prediction_scores, 

145 past_key_values=transformer_outputs.past_key_values, 

146 hidden_states=transformer_outputs.hidden_states, 

147 attentions=transformer_outputs.attentions, 

148 cross_attentions=transformer_outputs.cross_attentions, 

149 ) 

150 

151 def predict(self, input_ids, **kwargs): 

152 return self.forward(input_ids, **kwargs) 

153 

154 def generate(self, inputs, **kwargs): 

155 kwargs["logits_processor"] = self.logits_processor_list 

156 kwargs["num_return_sequences"] = kwargs.pop("paths_per_user") 

157 # Diverse/group beam search was moved to a `custom_generate` repo in transformers >=4.57. 

158 # We trust the transformers-community repositories, so allow loading the remote generation code. 

159 kwargs.setdefault("trust_remote_code", True) 

160 return super(GPT2LMHeadModel, self).generate(**inputs, **kwargs)