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

1# @Time : 2025/05/25 

2# @Author : Alessandro Soccol 

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

4 

5r"""PEARLMGPT2 

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

15 

16import math 

17 

18import torch 

19import torch.nn.functional as F 

20from torch import nn 

21 

22from hopwise.model.abstract_recommender import ExplainablePathLanguageModelingRecommender 

23from hopwise.utils import PathLanguageModelingTokenType 

24 

25 

26class LayerNorm(nn.Module): 

27 """LayerNorm but with an optional bias. PyTorch doesn't support simply bias=False""" 

28 

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 

33 

34 def forward(self, input): 

35 return F.layer_norm(input, self.weight.shape, self.weight, self.bias, 1e-5) 

36 

37 

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

46 

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

48 

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) 

53 

54 # output projection 

55 self.out_proj = nn.Linear(self.hidden_size, self.hidden_size, bias=False) 

56 

57 # regularization 

58 self.attn_dropout = nn.Dropout(config["dropout"]) 

59 self.resid_dropout = nn.Dropout(config["dropout"]) 

60 

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 ) 

67 

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) 

73 

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) 

78 

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) 

81 

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) 

85 

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

89 

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) 

94 

95 # calculate attention scores probabilities 

96 attn_scores = F.softmax(attn_scores, dim=-1) 

97 

98 attn_scores = self.attn_dropout(attn_scores) 

99 

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

101 

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 

106 

107 

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

115 

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 

122 

123 

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) 

131 

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) 

137 

138 return x 

139 

140 

141class PEARLMGPT2(ExplainablePathLanguageModelingRecommender): 

142 """ 

143 Low-level implementation of PEARLM model based on GPT-2 architecture that does not rely on HuggingFace tools. 

144 """ 

145 

146 def __init__(self, config, dataset): 

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

148 config["context_length"] = dataset.context_length 

149 

150 self.temperature = config["temperature"] 

151 

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

161 

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

165 

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

169 

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

171 

172 # weight tying 

173 self.wte.weight = self.lm_head.weight 

174 

175 self.loss = nn.CrossEntropyLoss() 

176 

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

183 

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) 

191 

192 def forward(self, idx): 

193 pos = torch.arange(0, idx.size(1), dtype=torch.long, device=self.device) # shape (t) 

194 

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

200 

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) 

207 

208 return x 

209 

210 def calculate_loss(self, interaction): 

211 input_ids = interaction["input_ids"] 

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

213 

214 lm_output = self.forward(input_ids) 

215 

216 logits = self.lm_head(lm_output) 

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

218 

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

220 

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

225 

226 return logits