Coverage for hopwise/model/path_language_modeling_recommender/pearlmllama3.py: 93%
162 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
2# @Author : Alessandro Soccol
3# @Email : alessandro.soccol@unica.it
5r"""PEARLMLlama3
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-llama2-to-llama3.ipynb
13"""
15import math
17import torch
18from torch import nn
20from hopwise.model.abstract_recommender import ExplainablePathLanguageModelingRecommender
21from hopwise.utils import PathLanguageModelingTokenType
24class AutoregressiveGroupQuerySelfAttention(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"]
31 # Reduce the projection dim to match desired output dim
32 self.head_dim = config["embedding_size"] // config["num_heads"]
34 assert config["embedding_size"] % config["num_heads"] == 0
36 # the second hidden size could be different
37 self.W_key = nn.Linear(self.hidden_size, self.hidden_size, bias=False, dtype=config["weight_precision"])
38 self.W_value = nn.Linear(self.hidden_size, self.hidden_size, bias=False, dtype=config["weight_precision"])
40 num_kv_groups = config["num_kv_groups"]
41 self.group_size = self.num_heads // num_kv_groups
43 self.W_query = nn.Linear(self.hidden_size, self.hidden_size, bias=False, dtype=config["weight_precision"])
44 self.out_proj = nn.Linear(self.hidden_size, self.hidden_size, bias=False, dtype=config["weight_precision"])
46 # RoPE pe
47 self.mask, self.cos, self.sin = SharedBuffers.get_buffers(
48 config["context_length"],
49 self.head_dim,
50 config["rope_base"],
51 config["rope_config"],
52 config["weight_precision"],
53 )
55 self.mask, self.cos, self.sin = (
56 self.mask.to(config["device"]),
57 self.cos.to(config["device"]),
58 self.sin.to(config["device"]),
59 )
61 self.causal_mask = (
62 torch.triu(torch.ones(config["context_length"], config["context_length"]), diagonal=1)
63 .bool()
64 .to(config["device"])
65 )
67 def forward(self, x):
68 batch_size, seq_length, hidden_size = (
69 x.size()
70 ) # batch size, sequence length, embedding dimensionality (n_embd)
72 # calculate query, key, values for all heads in batch and move head forward
73 # to be the batch dim
75 # applies a linear transformation on the last dimension: xW^T = (9,256)(256,256)
76 keys = self.W_key(x) # Shape: (b, num_tokens, d_out)
77 queries = self.W_query(x)
78 values = self.W_value(x)
80 # We implicitly split the matrix by adding a `num_heads` dimension
81 # Unroll last dim: (b, num_tokens, d_out) -> (b, num_tokens, num_heads, head_dim)
83 k = keys.view(batch_size, seq_length, self.num_heads, self.head_dim).transpose(1, 2)
84 v = values.view(batch_size, seq_length, self.num_heads, self.head_dim).transpose(1, 2)
85 q = queries.view(batch_size, seq_length, self.num_heads, self.head_dim).transpose(1, 2)
87 # NOTE: Fetch ROPE PE Embeddings
88 k = compute_rope(k, self.cos, self.sin)
89 q = compute_rope(q, self.cos, self.sin)
91 # Shape: (b, num_heads, num_tokens, head_dim)
92 keys = keys.repeat_interleave(self.group_size, dim=1)
93 # Shape: (b, num_heads, num_tokens, head_dim)
94 values = values.repeat_interleave(self.group_size, dim=1)
96 # Dot product for each head
97 # Calculate attention scores
99 attn_scores = (q @ k.transpose(2, 3)) * (1.0 / math.sqrt(k.size(-1)))
101 # apply causal masking
102 causal_mask = self.causal_mask.bool()[:seq_length, :seq_length]
103 attn_scores = attn_scores.masked_fill(causal_mask, -torch.inf)
105 attn_scores = torch.nn.functional.softmax(attn_scores, dim=-1)
107 context_vec = (attn_scores @ v).transpose(1, 2)
109 context_vec = context_vec.reshape(batch_size, seq_length, self.hidden_size)
110 context_vec = self.out_proj(context_vec)
111 return context_vec
114class FeedForward(nn.Module):
115 def __init__(self, config):
116 super().__init__()
117 self.fc1 = nn.Linear(config["embedding_size"], config["embedding_size"], bias=False)
118 self.fc2 = nn.Linear(config["embedding_size"], config["embedding_size"], bias=False)
119 self.fc3 = nn.Linear(config["embedding_size"], config["embedding_size"], bias=False)
120 self.silu = nn.SiLU()
122 def forward(self, x):
123 x_fc1 = self.fc1(x)
124 x_fc2 = self.fc2(x)
125 x = self.silu(x_fc1) * x_fc2
126 return self.fc3(x)
129class Block(nn.Module):
130 def __init__(self, config):
131 super().__init__()
132 self.rmsnorm1 = nn.RMSNorm(config["embedding_size"], eps=1e-5)
133 self.causal_attn = AutoregressiveGroupQuerySelfAttention(config)
134 self.rmsnorm2 = nn.RMSNorm(config["embedding_size"], eps=1e-5)
135 self.feedforward = FeedForward(config)
137 def forward(self, x):
138 # NOTE: Attention Block
139 shortcut = x
140 x = self.rmsnorm1(x)
141 x = self.causal_attn(x)
142 x += shortcut
144 # NOTE: Feed Forward Block
145 shortcut = x
146 x = self.rmsnorm2(x)
147 x = self.feedforward(x)
148 x = x + shortcut
149 return x
152class PEARLMLlama3(ExplainablePathLanguageModelingRecommender):
153 """
154 Low-level implementation of PEARLM model based on LLaMA 3 architecture.
156 With 8 kv-groups (that's how many Llama 3 8B uses), we can see that the number of rows
157 of the key and value matrices are reduced by a factor of 4
158 (because 32 attention heads divided by 8 kv-groups is 4)
159 To make the GroupedQueryAttention equivalent to standard multi-head attention,
160 you can set the number of query groups equal to the number of heads.
161 """
163 def __init__(self, config, dataset):
164 super().__init__(config, dataset, _skip_nn_module_init=False)
166 self.temperature = config["temperature"]
167 self.weight_precision = config["weight_precision"]
169 spec_type = PathLanguageModelingTokenType.SPECIAL.token_id
170 ent_type = PathLanguageModelingTokenType.ENTITY.token_id
171 rel_type = PathLanguageModelingTokenType.RELATION.token_id
172 type_tokens = [spec_type, ent_type, rel_type]
173 self.type_emb_pos = torch.LongTensor(
174 # BOS + ENT + REL + ENT + REL + ... + ENT + REL + EOS
175 [spec_type, ent_type] + [rel_type, ent_type] * dataset.path_hop_length + [spec_type],
176 )
177 self.type_emb_pos = self.type_emb_pos.to(config["device"])
179 self.wte = nn.Embedding(self.n_tokens, config["embedding_size"]).to(dtype=config["weight_precision"])
180 self.wpe = nn.Embedding(len(type_tokens), config["embedding_size"]).to(dtype=config["weight_precision"])
181 self.blocks = nn.ModuleList([Block(config) for _ in range(config["num_layers"])])
182 self.rmsnorm = nn.RMSNorm(config["embedding_size"], eps=1e-5)
184 self.lm_head = nn.Linear(config["embedding_size"], self.n_tokens, bias=False).to(
185 dtype=config["weight_precision"]
186 )
188 # weight tying
189 self.wte.weight = self.lm_head.weight
191 self.loss = nn.CrossEntropyLoss()
193 # init all weights
194 self.apply(self._init_weights)
195 # apply special scaled init to the residual projections, per GPT-2 paper
196 for pn, p in self.named_parameters():
197 if pn.endswith("c_proj.weight"):
198 torch.nn.init.normal_(p, mean=0.0, std=0.02 / math.sqrt(2 * config["num_layers"]))
200 def _init_weights(self, module):
201 if isinstance(module, nn.Linear):
202 torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
203 if module.bias is not None:
204 torch.nn.init.zeros_(module.bias)
205 elif isinstance(module, nn.Embedding):
206 torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
208 def forward(self, idx):
209 # forward the GPT model itself
210 # token embeddings of shape (b, t, n_embd)
211 tok_emb = self.wte(idx)
213 # think about get_flops, it restrict the max length of the input
214 pos_emb = self.wpe(self.type_emb_pos)[: tok_emb.size(1)]
215 x = tok_emb + pos_emb
216 for block in self.blocks:
217 x = block(x.to(self.weight_precision))
218 x = self.rmsnorm(x)
220 return x
222 def calculate_loss(self, interaction):
223 input_ids = interaction["input_ids"]
224 labels = input_ids[:, 1:].contiguous()
226 lm_output = self.forward(input_ids)
228 logits = self.lm_head(lm_output)
229 logits = logits[:, :-1, :].contiguous()
231 return self.loss(logits.view(-1, logits.size(-1)), labels.view(-1))
233 def predict(self, interaction):
234 input_ids = interaction["input_ids"]
235 lm_output = self.forward(input_ids)
236 logits = self.lm_head(lm_output[:, [-1], :])
238 return logits
241# RoPE
244def precompute_rope_params(head_dim, theta_base=10000, context_length=4096, freq_config=None, device=None):
245 assert head_dim % 2 == 0, "Embedding dimension must be even"
247 # Compute the inverse frequencies
248 inv_freq = 1.0 / (theta_base ** (torch.arange(0, head_dim, 2)[: (head_dim // 2)].float() / head_dim))
250 # Frequency adjustments used in LLaMA 3.1 and 3.2
251 if freq_config is not None:
252 low_freq_wavelen = freq_config["original_context_length"] / freq_config["low_freq_factor"]
253 high_freq_wavelen = freq_config["original_context_length"] / freq_config["high_freq_factor"]
255 wavelen = 2 * torch.pi / inv_freq
257 inv_freq_llama = torch.where(wavelen > low_freq_wavelen, inv_freq / freq_config["factor"], inv_freq)
259 smooth_factor = (freq_config["original_context_length"] / wavelen - freq_config["low_freq_factor"]) / (
260 freq_config["high_freq_factor"] - freq_config["low_freq_factor"]
261 )
263 smoothed_inv_freq = (1 - smooth_factor) * (inv_freq / freq_config["factor"]) + smooth_factor * inv_freq
265 is_medium_freq = (wavelen <= low_freq_wavelen) & (wavelen >= high_freq_wavelen)
266 inv_freq_llama = torch.where(is_medium_freq, smoothed_inv_freq, inv_freq_llama)
267 inv_freq = inv_freq_llama
269 # Generate position indices
270 positions = torch.arange(context_length)
272 # Compute the angles
273 # Shape: (context_length, head_dim // 2)
274 angles = positions[:, None] * inv_freq[None, :]
276 # Expand angles to match the head_dim
277 # Shape: (context_length, head_dim)
278 angles = torch.cat([angles, angles], dim=1)
280 # Precompute sine and cosine
281 cos = torch.cos(angles)
282 sin = torch.sin(angles)
284 return cos.to(device), sin.to(device)
287def compute_rope(x, cos, sin):
288 # x: (batch_size, num_heads, seq_len, head_dim)
289 batch_size, num_heads, seq_len, head_dim = x.shape
290 assert head_dim % 2 == 0, "Head dimension must be even"
292 # Split x into first half and second half
293 x1 = x[..., : head_dim // 2] # First half
294 x2 = x[..., head_dim // 2 :] # Second half
296 # Adjust sin and cos shapes
297 cos = cos[:seq_len, :].unsqueeze(0).unsqueeze(0) # Shape: (1, 1, seq_len, head_dim)
298 sin = sin[:seq_len, :].unsqueeze(0).unsqueeze(0)
300 # Apply the rotary transformation
301 rotated = torch.cat((-x2, x1), dim=-1)
302 x_rotated = (x * cos) + (rotated * sin)
304 return x_rotated
307class SharedBuffers:
308 _buffers = {}
310 @staticmethod
311 def get_buffers(context_length, head_dim, rope_base, freq_config, dtype=torch.float32):
312 key = (context_length, head_dim, rope_base, tuple(freq_config.values()) if freq_config else freq_config, dtype)
314 if key not in SharedBuffers._buffers:
315 # Create or fetch the buffers
316 mask = torch.triu(torch.ones(context_length, context_length), diagonal=1)
317 cos, sin = precompute_rope_params(head_dim, rope_base, context_length, freq_config)
318 if dtype is not None:
319 cos = cos.to(dtype)
320 sin = sin.to(dtype)
321 SharedBuffers._buffers[key] = (mask, cos, sin)
323 return SharedBuffers._buffers[key]