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

1# @Time : 2025/05 

2# @Author : Alessandro Soccol 

3# @Email : alessandro.soccol@unica.it 

4 

5r"""PEARLMLlama3 

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

7Reference: 

8 Balloccu et al. "Faithful Path Language Modeling for Explainable Recommendation over Knowledge Graph." - preprint. 

9 

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

14 

15import math 

16 

17import torch 

18from torch import nn 

19 

20from hopwise.model.abstract_recommender import ExplainablePathLanguageModelingRecommender 

21from hopwise.utils import PathLanguageModelingTokenType 

22 

23 

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

30 

31 # Reduce the projection dim to match desired output dim 

32 self.head_dim = config["embedding_size"] // config["num_heads"] 

33 

34 assert config["embedding_size"] % config["num_heads"] == 0 

35 

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

39 

40 num_kv_groups = config["num_kv_groups"] 

41 self.group_size = self.num_heads // num_kv_groups 

42 

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

45 

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 ) 

54 

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 ) 

60 

61 self.causal_mask = ( 

62 torch.triu(torch.ones(config["context_length"], config["context_length"]), diagonal=1) 

63 .bool() 

64 .to(config["device"]) 

65 ) 

66 

67 def forward(self, x): 

68 batch_size, seq_length, hidden_size = ( 

69 x.size() 

70 ) # batch size, sequence length, embedding dimensionality (n_embd) 

71 

72 # calculate query, key, values for all heads in batch and move head forward 

73 # to be the batch dim 

74 

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) 

79 

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) 

82 

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) 

86 

87 # NOTE: Fetch ROPE PE Embeddings 

88 k = compute_rope(k, self.cos, self.sin) 

89 q = compute_rope(q, self.cos, self.sin) 

90 

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) 

95 

96 # Dot product for each head 

97 # Calculate attention scores 

98 

99 attn_scores = (q @ k.transpose(2, 3)) * (1.0 / math.sqrt(k.size(-1))) 

100 

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) 

104 

105 attn_scores = torch.nn.functional.softmax(attn_scores, dim=-1) 

106 

107 context_vec = (attn_scores @ v).transpose(1, 2) 

108 

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 

112 

113 

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

121 

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) 

127 

128 

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) 

136 

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 

143 

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 

150 

151 

152class PEARLMLlama3(ExplainablePathLanguageModelingRecommender): 

153 """ 

154 Low-level implementation of PEARLM model based on LLaMA 3 architecture. 

155 

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

162 

163 def __init__(self, config, dataset): 

164 super().__init__(config, dataset, _skip_nn_module_init=False) 

165 

166 self.temperature = config["temperature"] 

167 self.weight_precision = config["weight_precision"] 

168 

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

178 

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) 

183 

184 self.lm_head = nn.Linear(config["embedding_size"], self.n_tokens, bias=False).to( 

185 dtype=config["weight_precision"] 

186 ) 

187 

188 # weight tying 

189 self.wte.weight = self.lm_head.weight 

190 

191 self.loss = nn.CrossEntropyLoss() 

192 

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

199 

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) 

207 

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) 

212 

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) 

219 

220 return x 

221 

222 def calculate_loss(self, interaction): 

223 input_ids = interaction["input_ids"] 

224 labels = input_ids[:, 1:].contiguous() 

225 

226 lm_output = self.forward(input_ids) 

227 

228 logits = self.lm_head(lm_output) 

229 logits = logits[:, :-1, :].contiguous() 

230 

231 return self.loss(logits.view(-1, logits.size(-1)), labels.view(-1)) 

232 

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], :]) 

237 

238 return logits 

239 

240 

241# RoPE 

242 

243 

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" 

246 

247 # Compute the inverse frequencies 

248 inv_freq = 1.0 / (theta_base ** (torch.arange(0, head_dim, 2)[: (head_dim // 2)].float() / head_dim)) 

249 

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

254 

255 wavelen = 2 * torch.pi / inv_freq 

256 

257 inv_freq_llama = torch.where(wavelen > low_freq_wavelen, inv_freq / freq_config["factor"], inv_freq) 

258 

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 ) 

262 

263 smoothed_inv_freq = (1 - smooth_factor) * (inv_freq / freq_config["factor"]) + smooth_factor * inv_freq 

264 

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 

268 

269 # Generate position indices 

270 positions = torch.arange(context_length) 

271 

272 # Compute the angles 

273 # Shape: (context_length, head_dim // 2) 

274 angles = positions[:, None] * inv_freq[None, :] 

275 

276 # Expand angles to match the head_dim 

277 # Shape: (context_length, head_dim) 

278 angles = torch.cat([angles, angles], dim=1) 

279 

280 # Precompute sine and cosine 

281 cos = torch.cos(angles) 

282 sin = torch.sin(angles) 

283 

284 return cos.to(device), sin.to(device) 

285 

286 

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" 

291 

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 

295 

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) 

299 

300 # Apply the rotary transformation 

301 rotated = torch.cat((-x2, x1), dim=-1) 

302 x_rotated = (x * cos) + (rotated * sin) 

303 

304 return x_rotated 

305 

306 

307class SharedBuffers: 

308 _buffers = {} 

309 

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) 

313 

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) 

322 

323 return SharedBuffers._buffers[key]