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
« 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
5r"""PEARLM
6##################################################
7Reference:
8 Balloccu et al. "Faithful Path Language Modeling for Explainable Recommendation over Knowledge Graph." - preprint.
10Reference code:
11 https://github.com/Chris1nexus/pearlm
12"""
14from typing import Optional, Union
16import torch
17from torch import nn
18from transformers import AutoConfig, GPT2LMHeadModel
19from transformers.modeling_outputs import CausalLMOutputWithCrossAttentions
21from hopwise.data import Interaction
22from hopwise.model.abstract_recommender import ExplainablePathLanguageModelingRecommender
23from hopwise.utils import PathLanguageModelingTokenType
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 """
34 def __init__(self, config, dataset):
35 ExplainablePathLanguageModelingRecommender.__init__(self, config, dataset)
37 self.use_kg_token_types = config["use_kg_token_types"]
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"])
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
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
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"])
79 self.transformer.resize_token_embeddings(self.n_tokens)
81 self.loss = nn.CrossEntropyLoss()
82 self.post_init()
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"]
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)
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 )
128 sequence_output = transformer_outputs[0]
129 prediction_scores = self.lm_head(sequence_output)
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))
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
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 )
151 def predict(self, input_ids, **kwargs):
152 return self.forward(input_ids, **kwargs)
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)