Coverage for hopwise/model/path_language_modeling_recommender/pearlmllama2.py: 99%
134 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"""PEARLMLlama2
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/rasbt/LLMs-from-scratch/blob/main/ch05/07_gpt_to_llama/converting-gpt-to-llama2.ipynb
13"""
15import math
17import torch
18from torch import nn
20from hopwise.model.abstract_recommender import ExplainablePathLanguageModelingRecommender
21from hopwise.utils import PathLanguageModelingTokenType
24class AutoregressiveSelfAttention(nn.Module):
25 def __init__(self, config):
26 super().__init__()
27 self.hidden_size = config["embedding_size"]
28 self.num_heads = config["num_heads"]
29 self.dropout = config["dropout"]
30 # Reduce the projection dim to match desired output dim
31 self.head_dim = config["embedding_size"] // config["num_heads"]
33 assert config["embedding_size"] % config["num_heads"] == 0
35 # the second hidden size could be different
36 self.W_query = nn.Linear(self.hidden_size, self.hidden_size, bias=False)
37 self.W_key = nn.Linear(self.hidden_size, self.hidden_size, bias=False)
38 self.W_value = nn.Linear(self.hidden_size, self.hidden_size, bias=False)
39 self.out_proj = nn.Linear(self.hidden_size, self.hidden_size, bias=False)
41 # RoPE pe
42 self.cos, self.sin = precompute_rope_params(
43 head_dim=self.head_dim, context_length=config["context_length"], device=config["device"]
44 )
46 self.causal_mask = (
47 torch.triu(torch.ones(config["context_length"], config["context_length"]), diagonal=1)
48 .bool()
49 .to(config["device"])
50 )
52 def forward(self, x):
53 batch_size, seq_length, hidden_size = (
54 x.size()
55 ) # batch size, sequence length, embedding dimensionality (n_embd)
57 # calculate query, key, values for all heads in batch and move head forward to be the batch dim
59 # applies a linear transformation on the last dimension: xW^T = (9,256)(256,256)
60 keys = self.W_key(x) # Shape: (b, num_tokens, d_out)
61 queries = self.W_query(x)
62 values = self.W_value(x)
64 # We implicitly split the matrix by adding a `num_heads` dimension
65 # Unroll last dim: (b, num_tokens, d_out) -> (b, num_tokens, num_heads, head_dim)
67 k = keys.view(batch_size, seq_length, self.num_heads, self.head_dim).transpose(1, 2)
68 v = values.view(batch_size, seq_length, self.num_heads, self.head_dim).transpose(1, 2)
69 q = queries.view(batch_size, seq_length, self.num_heads, self.head_dim).transpose(1, 2)
71 # NOTE: Fetch ROPE PE Embeddings
72 k = compute_rope(k, self.cos, self.sin)
73 q = compute_rope(q, self.cos, self.sin)
75 # Dot product for each head
76 # Calculate attention scores
78 attn_scores = (q @ k.transpose(2, 3)) * (1.0 / math.sqrt(k.size(-1)))
80 # apply causal masking
81 causal_mask = self.causal_mask.bool()[:seq_length, :seq_length]
82 attn_scores = attn_scores.masked_fill(causal_mask, -torch.inf)
84 attn_scores = torch.nn.functional.softmax(attn_scores, dim=-1)
86 context_vec = (attn_scores @ v).transpose(1, 2)
88 context_vec = context_vec.reshape(batch_size, seq_length, self.hidden_size)
89 context_vec = self.out_proj(context_vec)
90 return context_vec
93class FeedForward(nn.Module):
94 def __init__(self, config):
95 super().__init__()
96 self.fc1 = nn.Linear(config["embedding_size"], config["embedding_size"], bias=False)
97 self.fc2 = nn.Linear(config["embedding_size"], config["embedding_size"], bias=False)
98 self.fc3 = nn.Linear(config["embedding_size"], config["embedding_size"], bias=False)
99 self.silu = nn.SiLU()
101 def forward(self, x):
102 x_fc1 = self.fc1(x)
103 x_fc2 = self.fc2(x)
104 x = self.silu(x_fc1) * x_fc2
105 return self.fc3(x)
108class Block(nn.Module):
109 def __init__(self, config):
110 super().__init__()
111 self.rmsnorm1 = nn.RMSNorm(config["embedding_size"])
112 self.causal_attn = AutoregressiveSelfAttention(config)
113 self.rmsnorm2 = nn.RMSNorm(config["embedding_size"])
114 self.feedforward = FeedForward(config)
116 def forward(self, x):
117 # NOTE: Attention Block
118 shortcut = x
119 x = self.rmsnorm1(x)
120 x = self.causal_attn(x)
121 x += shortcut
123 # NOTE: Feed Forward Block
124 shortcut = x
125 x = self.rmsnorm2(x)
126 x = self.feedforward(x)
127 x = x + shortcut
128 return x
131class PEARLMLlama2(ExplainablePathLanguageModelingRecommender):
132 """
133 Low-level implementation of PEARLM model based on Llama2 architecture.
135 Novelties:
136 - LayerNorm is replaced with RMSNorm
137 - GeLU is replaced with SiLU
138 - Feedforward is replaced with a simple linear head
139 """
141 def __init__(self, config, dataset):
142 super().__init__(config, dataset, _skip_nn_module_init=False)
143 config["context_length"] = dataset.context_length
145 self.temperature = config["temperature"]
147 spec_type = PathLanguageModelingTokenType.SPECIAL.token_id
148 ent_type = PathLanguageModelingTokenType.ENTITY.token_id
149 rel_type = PathLanguageModelingTokenType.RELATION.token_id
150 type_tokens = [spec_type, ent_type, rel_type]
151 self.type_emb_pos = torch.LongTensor(
152 # BOS + ENT + REL + ENT + REL + ... + ENT + REL + EOS
153 [spec_type, ent_type] + [rel_type, ent_type] * dataset.path_hop_length + [spec_type],
154 )
155 self.type_emb_pos = self.type_emb_pos.to(config["device"])
157 self.wte = nn.Embedding(self.n_tokens, config["embedding_size"])
158 self.wpe = nn.Embedding(len(type_tokens), config["embedding_size"])
159 self.blocks = nn.ModuleList([Block(config) for _ in range(config["num_layers"])])
160 self.rmsnorm = nn.RMSNorm(config["embedding_size"])
162 self.lm_head = nn.Linear(config["embedding_size"], self.n_tokens, bias=False)
164 # weight tying
165 self.wte.weight = self.lm_head.weight
167 self.loss = nn.CrossEntropyLoss()
169 # init all weights
170 self.apply(self._init_weights)
171 # apply special scaled init to the residual projections, per GPT-2 paper
172 for pn, p in self.named_parameters():
173 if pn.endswith("c_proj.weight"):
174 torch.nn.init.normal_(p, mean=0.0, std=0.02 / math.sqrt(2 * config["num_layers"]))
176 def _init_weights(self, module):
177 if isinstance(module, nn.Linear):
178 torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
179 if module.bias is not None:
180 torch.nn.init.zeros_(module.bias)
181 elif isinstance(module, nn.Embedding):
182 torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
184 def forward(self, idx):
185 # forward the GPT model itself
186 # token embeddings of shape (b, t, n_embd)
187 tok_emb = self.wte(idx)
189 # think about get_flops, it restrict the max length of the input
190 pos_emb = self.wpe(self.type_emb_pos)[: tok_emb.size(1)]
191 x = tok_emb + pos_emb
192 for block in self.blocks:
193 x = block(x)
194 x = self.rmsnorm(x)
196 return x
198 def calculate_loss(self, interaction):
199 input_ids = interaction["input_ids"]
200 labels = input_ids[:, 1:].contiguous()
202 lm_output = self.forward(input_ids)
204 logits = self.lm_head(lm_output)
205 logits = logits[:, :-1, :].contiguous()
207 return self.loss(logits.view(-1, logits.size(-1)), labels.view(-1))
209 def predict(self, interaction):
210 input_ids = interaction["input_ids"]
211 lm_output = self.forward(input_ids)
212 logits = self.lm_head(lm_output[:, [-1], :])
214 return logits
217# RoPE
220def precompute_rope_params(head_dim, theta_base=10000, context_length=4096, device=None):
221 assert head_dim % 2 == 0, "Embedding dimension must be even"
223 # Compute the inverse frequencies
224 inv_freq = 1.0 / (theta_base ** (torch.arange(0, head_dim, 2)[: (head_dim // 2)].float() / head_dim))
226 # Generate position indices
227 positions = torch.arange(context_length)
229 # Compute the angles
230 # Shape: (context_length, head_dim // 2)
231 angles = positions[:, None] * inv_freq[None, :]
233 # Expand angles to match the head_dim
234 # Shape: (context_length, head_dim)
235 angles = torch.cat([angles, angles], dim=1)
237 # Precompute sine and cosine
238 cos = torch.cos(angles)
239 sin = torch.sin(angles)
241 return cos.to(device), sin.to(device)
244def compute_rope(x, cos, sin):
245 # x: (batch_size, num_heads, seq_len, head_dim)
246 batch_size, num_heads, seq_len, head_dim = x.shape
247 assert head_dim % 2 == 0, "Head dimension must be even"
249 # Split x into first half and second half
250 x1 = x[..., : head_dim // 2] # First half
251 x2 = x[..., head_dim // 2 :] # Second half
253 # Adjust sin and cos shapes
254 cos = cos[:seq_len, :].unsqueeze(0).unsqueeze(0) # Shape: (1, 1, seq_len, head_dim)
255 sin = sin[:seq_len, :].unsqueeze(0).unsqueeze(0)
257 # Apply the rotary transformation
258 rotated = torch.cat((-x2, x1), dim=-1)
259 x_rotated = (x * cos) + (rotated * sin)
261 return x_rotated