Coverage for hopwise/model/path_language_modeling_recommender/plm.py: 92%
79 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/5/29
2# @Author : Giacomo Medda
3# @Email : giacomo.medda@unica.it
5r"""PLM
6################################################
7Reference:
8 Shijie Geng et al. "Path Language Modeling over Knowledge Graphs for Explainable Recommendation." in WWW 2022.
10Reference code:
11 https://github.com/mirkomarras/kgglm
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 PLM(ExplainablePathLanguageModelingRecommender, GPT2LMHeadModel):
27 """PLM is a path-language-modeling recommender. It learns the sequence of entity-relation triplets
28 from a knowledge graph as a next-token prediction task and employs two feature transformations separately
29 for entities and relations. Its decoding process is unbounded, meaning that it can generate paths that are
30 not faithful to the knowledge graph structure, i.e., it can generate paths that do not exist in the KG.
31 """
33 def __init__(self, config, dataset):
34 ExplainablePathLanguageModelingRecommender.__init__(self, config, dataset)
36 self.use_kg_token_types = config["use_kg_token_types"]
37 transformers_config = AutoConfig.from_pretrained(
38 "distilgpt2",
39 **{
40 "vocab_size": self.n_tokens,
41 "n_ctx": config["context_length"],
42 "n_positions": config["context_length"],
43 "pad_token_id": dataset.tokenizer.pad_token_id,
44 "bos_token_id": dataset.tokenizer.bos_token_id,
45 "eos_token_id": dataset.tokenizer.eos_token_id,
46 "n_embd": config["embedding_size"],
47 "n_head": config["num_heads"],
48 "n_layer": config["num_layers"],
49 },
50 )
51 GPT2LMHeadModel.__init__(self, transformers_config)
53 # Add type ids template
54 prev_vocab_size = self.n_tokens
55 spec_type, spec_type_id = PathLanguageModelingTokenType.SPECIAL.value
56 ent_type, ent_type_id = PathLanguageModelingTokenType.ENTITY.value
57 rel_type, rel_type_id = PathLanguageModelingTokenType.RELATION.value
59 token_types = [f"<Token-Type.{token_type}>" for token_type in [spec_type, ent_type, rel_type]]
60 for token_type in token_types:
61 dataset.tokenizer.add_tokens(token_type)
62 self.n_tokens = len(dataset.tokenizer) # Update the vocabulary size after adding new tokens
64 spec_type_id, ent_type_id, rel_type_id = (
65 spec_type_id + prev_vocab_size,
66 ent_type_id + prev_vocab_size,
67 rel_type_id + prev_vocab_size,
68 )
69 self.token_type_ids = torch.LongTensor(
70 # BOS + ENT + REL + ENT + REL + ... + ENT + REL + EOS
71 [spec_type_id, ent_type_id] + [rel_type_id, ent_type_id] * dataset.path_hop_length + [spec_type_id]
72 )
73 self.token_entity_type_id = ent_type_id
74 self.token_relation_type_id = rel_type_id
75 self.token_type_ids = self.token_type_ids.to(config["device"])
77 self.transformer.resize_token_embeddings(self.n_tokens)
79 vocab_inv = {v: k for k, v in dataset.tokenizer.get_vocab().items()}
80 relation_mask = torch.tensor(
81 [vocab_inv[i].startswith(rel_type) for i in range(self.config.vocab_size)],
82 dtype=torch.float32,
83 )
84 user_type = PathLanguageModelingTokenType.USER.token
85 item_type = PathLanguageModelingTokenType.ITEM.token
86 entity_mask = torch.tensor(
87 [
88 vocab_inv[i].startswith(user_type)
89 or vocab_inv[i].startswith(item_type)
90 or vocab_inv[i].startswith(ent_type)
91 for i in range(self.config.vocab_size)
92 ],
93 dtype=torch.float32,
94 )
96 self.entity_head = torch.nn.Linear(self.config.n_embd, self.config.vocab_size, bias=False)
97 self.relation_head = torch.nn.Linear(self.config.n_embd, self.config.vocab_size, bias=False)
99 self.entity_loss = nn.CrossEntropyLoss(weight=entity_mask)
100 self.relation_loss = nn.CrossEntropyLoss(weight=relation_mask)
102 self.to(config["device"])
103 self.post_init()
105 def forward(
106 self,
107 input_ids: Optional[torch.LongTensor] = None,
108 past_key_values: Optional[tuple[tuple[torch.Tensor]]] = None,
109 attention_mask: Optional[torch.FloatTensor] = None,
110 token_type_ids: Optional[torch.LongTensor] = None,
111 position_ids: Optional[torch.LongTensor] = None,
112 head_mask: Optional[torch.FloatTensor] = None,
113 inputs_embeds: Optional[torch.FloatTensor] = None,
114 encoder_hidden_states: Optional[torch.Tensor] = None,
115 encoder_attention_mask: Optional[torch.FloatTensor] = None,
116 labels: Optional[torch.LongTensor] = None,
117 use_cache: Optional[bool] = None,
118 output_attentions: Optional[bool] = None,
119 output_hidden_states: Optional[bool] = None,
120 return_dict: Optional[bool] = None,
121 **kwargs, # Additional arguments for compatibility with HuggingFace Trainer
122 ) -> Union[tuple, CausalLMOutputWithCrossAttentions]:
123 if isinstance(input_ids, Interaction):
124 token_type_ids = input_ids["token_type_ids"]
125 attention_mask = input_ids["attention_mask"]
126 input_ids = input_ids["input_ids"]
128 token_type_ids = self.token_type_ids[[*range(input_ids.shape[1] - 1), -1]]
129 token_type_ids = token_type_ids.expand(input_ids.shape[0], -1)
131 transformer_outputs = self.transformer(
132 input_ids,
133 past_key_values=past_key_values,
134 attention_mask=attention_mask,
135 token_type_ids=token_type_ids,
136 position_ids=position_ids,
137 head_mask=head_mask,
138 inputs_embeds=inputs_embeds,
139 encoder_hidden_states=encoder_hidden_states,
140 encoder_attention_mask=encoder_attention_mask,
141 use_cache=use_cache and labels is None,
142 output_attentions=output_attentions,
143 output_hidden_states=output_hidden_states,
144 return_dict=return_dict or self.config.use_return_dict,
145 )
147 sequence_output = transformer_outputs[0]
148 if self.model_parallel:
149 torch.cuda.set_device(self.transformer.first_device)
150 sequence_output = sequence_output.to(self.entity_head.weight.device)
152 # Get logits from the two heads, first based on entity tokens, then on relation tokens
153 lm_entity_scores = self.entity_head(sequence_output)
154 lm_relation_scores = self.relation_head(sequence_output)
156 lm_loss = None
157 if labels is not None:
158 # entity head loss
159 entity_scores_mask = (token_type_ids[0, :-1] != self.token_entity_type_id).to(lm_entity_scores.device)
160 entity_labels_mask = (token_type_ids[0, 1:-1] != self.token_relation_type_id).to(lm_entity_scores.device)
161 shifted_entity_scores = lm_entity_scores[:, :-1, :][:, entity_scores_mask, :].contiguous()
162 entity_labels = labels[:, 1:-1][:, entity_labels_mask].contiguous()
163 lm_loss = self.entity_loss(shifted_entity_scores.view(-1, self.config.vocab_size), entity_labels.view(-1))
165 # relation head loss
166 relations_scores_mask = (token_type_ids[0, 1:-1] != self.token_relation_type_id).to(
167 lm_entity_scores.device
168 )
169 relation_labels_mask = (token_type_ids[0, 1:] != self.token_entity_type_id).to(lm_entity_scores.device)
170 shifted_relation_scores = lm_relation_scores[:, 1:-1, :][:, relations_scores_mask, :].contiguous()
171 relation_labels = labels[:, 1:][:, relation_labels_mask].contiguous()
172 lm_loss += self.relation_loss(
173 shifted_relation_scores.view(-1, self.config.vocab_size), relation_labels.view(-1)
174 )
176 lm_scores = lm_entity_scores
177 relation_idx = torch.arange(1, input_ids.shape[1], 2, device=input_ids.device)
178 lm_scores[:, relation_idx] = lm_relation_scores[:, relation_idx]
180 if not return_dict:
181 output = (lm_scores,) + transformer_outputs[2:]
182 return ((lm_loss,) + output) if lm_loss is not None else output
184 return CausalLMOutputWithCrossAttentions(
185 loss=lm_loss,
186 logits=lm_scores,
187 past_key_values=transformer_outputs.past_key_values,
188 hidden_states=transformer_outputs.hidden_states,
189 attentions=transformer_outputs.attentions,
190 cross_attentions=transformer_outputs.cross_attentions,
191 )
193 def predict(self, input_ids, **kwargs):
194 return self.forward(input_ids, **kwargs)
196 def generate(self, inputs, **kwargs):
197 kwargs["logits_processor"] = self.logits_processor_list
198 kwargs["num_return_sequences"] = kwargs.pop("paths_per_user")
199 # Diverse/group beam search was moved to a `custom_generate` repo in transformers >=4.57.
200 # We trust the transformers-community repositories, so allow loading the remote generation code.
201 kwargs.setdefault("trust_remote_code", True)
202 return super(GPT2LMHeadModel, self).generate(**inputs, **kwargs)