Coverage for hopwise/model/path_language_modeling_recommender/pearlmgpt2.py: 100%
126 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/05/25
2# @Author : Alessandro Soccol
3# @Email : alessandro.soccol@unica.it
5r"""PEARLMGPT2
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 https://github.com/karpathy/nanoGPT/blob/master/model.py
13 https://github.com/rasbt/LLMs-from-scratch/blob/main/ch05/07_gpt_to_llama/converting-gpt-to-llama2.ipynb
14"""
16import math
18import torch
19import torch.nn.functional as F
20from torch import nn
22from hopwise.model.abstract_recommender import ExplainablePathLanguageModelingRecommender
23from hopwise.utils import PathLanguageModelingTokenType
26class LayerNorm(nn.Module):
27 """LayerNorm but with an optional bias. PyTorch doesn't support simply bias=False"""
29 def __init__(self, ndim, bias):
30 super().__init__()
31 self.weight = nn.Parameter(torch.ones(ndim))
32 self.bias = nn.Parameter(torch.zeros(ndim)) if bias else None
34 def forward(self, input):
35 return F.layer_norm(input, self.weight.shape, self.weight, self.bias, 1e-5)
38class AutoregressiveSelfAttention(nn.Module):
39 def __init__(self, config):
40 super().__init__()
41 self.hidden_size = config["embedding_size"]
42 self.num_heads = config["num_heads"]
43 self.dropout = config["dropout"]
44 # Reduce the projection dim to match desired output dim
45 self.head_dim = config["embedding_size"] // config["num_heads"]
47 assert config["embedding_size"] % config["num_heads"] == 0
49 # the second hidden size could be different
50 self.W_query = nn.Linear(self.hidden_size, self.hidden_size, bias=False)
51 self.W_key = nn.Linear(self.hidden_size, self.hidden_size, bias=False)
52 self.W_value = nn.Linear(self.hidden_size, self.hidden_size, bias=False)
54 # output projection
55 self.out_proj = nn.Linear(self.hidden_size, self.hidden_size, bias=False)
57 # regularization
58 self.attn_dropout = nn.Dropout(config["dropout"])
59 self.resid_dropout = nn.Dropout(config["dropout"])
61 # causal mask to ensure that attention is only applied to the left in the input sequence
62 self.causal_mask = (
63 torch.triu(torch.ones(config["context_length"], config["context_length"]), diagonal=1)
64 .bool()
65 .to(config["device"])
66 )
68 def forward(self, x):
69 # B,C,T = x.size() # batch size, sequence length, embedding dimensionality (n_embd)
70 batch_size, seq_length, hidden_size = (
71 x.size()
72 ) # batch size, sequence length, embedding dimensionality (n_embd)
74 # applies a linear transformation on the last dimension: xW^T = (9,256)(256,256)
75 keys = self.W_key(x) # Shape: (b, num_tokens, d_out)
76 queries = self.W_query(x)
77 values = self.W_value(x)
79 # We implicitly split the matrix by adding a `num_heads` dimension
80 # Unroll last dim: (b, num_tokens, d_out) -> (b, num_tokens, num_heads, head_dim)
82 k = keys.view(batch_size, seq_length, self.num_heads, self.head_dim).transpose(1, 2) # 4096,4,9,64
83 v = values.view(batch_size, seq_length, self.num_heads, self.head_dim).transpose(1, 2)
84 q = queries.view(batch_size, seq_length, self.num_heads, self.head_dim).transpose(1, 2)
86 # Dot product for each head
87 # Calculate attention scores
88 attn_scores = (q @ k.transpose(2, 3)) * (1.0 / math.sqrt(k.size(-1)))
90 # truncating the mask to the current sequence length
91 causal_mask = self.causal_mask.bool()[:seq_length, :seq_length]
92 # apply causal masking
93 attn_scores = attn_scores.masked_fill(causal_mask, -torch.inf)
95 # calculate attention scores probabilities
96 attn_scores = F.softmax(attn_scores, dim=-1)
98 attn_scores = self.attn_dropout(attn_scores)
100 context_vec = (attn_scores @ v).transpose(1, 2)
102 context_vec = context_vec.reshape(batch_size, seq_length, self.hidden_size)
103 context_vec = self.out_proj(context_vec) # optional projection
104 context_vec = self.resid_dropout(context_vec)
105 return context_vec
108class FeedForward(nn.Module):
109 def __init__(self, config):
110 super().__init__()
111 self.c_fc = nn.Linear(config["embedding_size"], 4 * config["embedding_size"], bias=config["bias"])
112 self.silu = nn.GELU()
113 self.c_proj = nn.Linear(4 * config["embedding_size"], config["embedding_size"], bias=config["bias"])
114 self.dropout = nn.Dropout(config["dropout"])
116 def forward(self, x):
117 x = self.c_fc(x)
118 x = self.silu(x)
119 x = self.c_proj(x)
120 x = self.dropout(x)
121 return x
124class Block(nn.Module):
125 def __init__(self, config):
126 super().__init__()
127 self.layernorm_1 = LayerNorm(config["embedding_size"], bias=config["bias"])
128 self.causal_attn = AutoregressiveSelfAttention(config)
129 self.layernorm_2 = LayerNorm(config["embedding_size"], bias=config["bias"])
130 self.feedforward = FeedForward(config)
132 def forward(self, x):
133 x = self.layernorm_1(x)
134 x = x + self.causal_attn(x)
135 x = self.layernorm_2(x)
136 x = x + self.feedforward(x)
138 return x
141class PEARLMGPT2(ExplainablePathLanguageModelingRecommender):
142 """
143 Low-level implementation of PEARLM model based on GPT-2 architecture that does not rely on HuggingFace tools.
144 """
146 def __init__(self, config, dataset):
147 super().__init__(config, dataset, _skip_nn_module_init=False)
148 config["context_length"] = dataset.context_length
150 self.temperature = config["temperature"]
152 spec_type = PathLanguageModelingTokenType.SPECIAL.token_id
153 ent_type = PathLanguageModelingTokenType.ENTITY.token_id
154 rel_type = PathLanguageModelingTokenType.RELATION.token_id
155 type_tokens = [spec_type, ent_type, rel_type]
156 self.type_emb_pos = torch.LongTensor(
157 # BOS + ENT + REL + ENT + REL + ... + ENT + REL + EOS
158 [spec_type, ent_type] + [rel_type, ent_type] * dataset.path_hop_length + [spec_type],
159 )
160 self.type_emb_pos = self.type_emb_pos.to(config["device"])
162 self.wte = nn.Embedding(self.n_tokens, config["embedding_size"])
163 self.wpe = nn.Embedding(dataset.context_length, config["embedding_size"])
164 self.wp_type_e = nn.Embedding(len(type_tokens), config["embedding_size"])
166 self.blocks = nn.ModuleList([Block(config) for _ in range(config["num_layers"])])
167 self.layernorm = nn.LayerNorm(config["embedding_size"], bias=config["bias"])
168 self.dropout = nn.Dropout(config["dropout"])
170 self.lm_head = nn.Linear(config["embedding_size"], self.n_tokens, bias=False)
172 # weight tying
173 self.wte.weight = self.lm_head.weight
175 self.loss = nn.CrossEntropyLoss()
177 # init all weights
178 self.apply(self._init_weights)
179 # apply special scaled init to the residual projections, per GPT-2 paper
180 for pn, p in self.named_parameters():
181 if pn.endswith("c_proj.weight"):
182 torch.nn.init.normal_(p, mean=0.0, std=0.02 / math.sqrt(2 * config["num_layers"]))
184 def _init_weights(self, module):
185 if isinstance(module, nn.Linear):
186 torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
187 if module.bias is not None:
188 torch.nn.init.zeros_(module.bias)
189 elif isinstance(module, nn.Embedding):
190 torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
192 def forward(self, idx):
193 pos = torch.arange(0, idx.size(1), dtype=torch.long, device=self.device) # shape (t)
195 # forward the GPT model itself
196 # token embeddings of shape (b, t, n_embd)
197 tok_emb = self.wte(idx)
198 pos_emb = self.wpe(pos)[: tok_emb.size(1)]
199 type_emb = self.wp_type_e(self.type_emb_pos)[: tok_emb.size(1)]
201 # think about get_flops, it restrict the max length of the input
202 x = tok_emb + pos_emb + type_emb
203 x = self.dropout(tok_emb + pos_emb)
204 for block in self.blocks:
205 x = block(x)
206 x = self.layernorm(x)
208 return x
210 def calculate_loss(self, interaction):
211 input_ids = interaction["input_ids"]
212 labels = input_ids[:, 1:].contiguous()
214 lm_output = self.forward(input_ids)
216 logits = self.lm_head(lm_output)
217 logits = logits[:, :-1, :].contiguous()
219 return self.loss(logits.view(-1, logits.size(-1)), labels.view(-1))
221 def predict(self, interaction):
222 input_ids = interaction["input_ids"]
223 lm_output = self.forward(input_ids)
224 logits = self.lm_head(lm_output[:, [-1], :])
226 return logits