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

1# @Time : 2025/5/29 

2# @Author : Giacomo Medda 

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

4 

5r"""PLM 

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

7Reference: 

8 Shijie Geng et al. "Path Language Modeling over Knowledge Graphs for Explainable Recommendation." in WWW 2022. 

9 

10Reference code: 

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

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 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 """ 

32 

33 def __init__(self, config, dataset): 

34 ExplainablePathLanguageModelingRecommender.__init__(self, config, dataset) 

35 

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) 

52 

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 

58 

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 

63 

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"]) 

76 

77 self.transformer.resize_token_embeddings(self.n_tokens) 

78 

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 ) 

95 

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) 

98 

99 self.entity_loss = nn.CrossEntropyLoss(weight=entity_mask) 

100 self.relation_loss = nn.CrossEntropyLoss(weight=relation_mask) 

101 

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

103 self.post_init() 

104 

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"] 

127 

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) 

130 

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 ) 

146 

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) 

151 

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) 

155 

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)) 

164 

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 ) 

175 

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] 

179 

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 

183 

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 ) 

192 

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

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

195 

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)