Coverage for hopwise/model/knowledge_aware_recommender/kglrr.py: 81%
380 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
1import logging
2import os
4import numpy as np
5import torch
6import torch.nn.functional as F
7from torch import nn
9from hopwise.model.abstract_recommender import KnowledgeRecommender
10from hopwise.utils import InputType
13class GraphAttentionLayer(nn.Module):
14 def __init__(self, in_features, out_features, dropout, alpha, concat=True):
15 super().__init__()
16 self.dropout = dropout
17 self.in_features = in_features
18 self.out_features = out_features
19 self.alpha = alpha
20 self.concat = concat
22 self.W = nn.Parameter(torch.empty(size=(in_features, out_features)))
23 nn.init.xavier_uniform_(self.W.data, gain=1.414)
24 self.a = nn.Parameter(torch.empty(size=(2 * out_features, 1)))
25 nn.init.xavier_uniform_(self.a.data, gain=1.414)
26 self.fc = nn.Linear(2 * out_features, out_features)
28 self.leakyrelu = nn.LeakyReLU(self.alpha)
30 def forward_relation(self, item_embs, entity_embs, relations, adj):
31 # item_embs: N, dim
32 # entity_embs: N, e_num, dim
33 # relations: N, e_num, r_dim
34 # adj: N, e_num
36 # N, e_num, dim
37 Wh = item_embs.unsqueeze(1).expand(entity_embs.size())
38 # N, e_num, dim
39 We = entity_embs
40 a_input = torch.cat((Wh, We), dim=-1) # (N, e_num, 2*dim)
41 # N,e,2dim -> N,e,dim
42 e_input = torch.multiply(self.fc(a_input), relations).sum(-1) # N,e
43 e = self.leakyrelu(e_input) # (N, e_num)
45 zero_vec = -9e15 * torch.ones_like(e)
46 attention = torch.where(adj > 0, e, zero_vec)
47 attention = F.softmax(attention, dim=1)
48 attention = F.dropout(attention, self.dropout, training=self.training) # N, e_num
49 # (N, 1, e_num) * (N, e_num, out_features) -> N, out_features
50 entity_emb_weighted = torch.bmm(attention.unsqueeze(1), entity_embs).squeeze()
51 h_prime = entity_emb_weighted + item_embs
53 if self.concat:
54 return F.elu(h_prime)
55 else:
56 return h_prime
58 def forward(self, item_embs, entity_embs, adj):
59 Wh = torch.mm(item_embs, self.W) # h.shape: (N, in_features), Wh.shape: (N, out_features)
60 We = torch.matmul(
61 entity_embs, self.W
62 ) # entity_embs: (N, e_num, in_features), We.shape: (N, e_num, out_features)
63 a_input = self._prepare_cat(Wh, We) # (N, e_num, 2*out_features)
64 e = self.leakyrelu(torch.matmul(a_input, self.a).squeeze(2)) # (N, e_num)
66 zero_vec = -9e15 * torch.ones_like(e)
67 attention = torch.where(adj > 0, e, zero_vec)
68 attention = F.softmax(attention, dim=1)
69 attention = F.dropout(attention, self.dropout, training=self.training) # N, e_num
70 # (N, 1, e_num) * (N, e_num, out_features) -> N, out_features
71 entity_emb_weighted = torch.bmm(attention.unsqueeze(1), entity_embs).squeeze()
72 h_prime = entity_emb_weighted + item_embs
74 if self.concat:
75 return F.elu(h_prime)
76 else:
77 return h_prime
79 def _prepare_cat(self, Wh, We):
80 Wh = Wh.unsqueeze(1).expand(We.size()) # (N, e_num, out_features)
81 return torch.cat((Wh, We), dim=-1) # (N, e_num, 2*out_features)
84class KGEncoder(nn.Module):
85 def __init__(self, config, dataset, kg_dataset):
86 super().__init__()
88 self.user_history_matrix = dataset.history_item_matrix()[0].to(config["device"])
90 self.maxhis = config["maxhis"]
91 self.kgcn = config["kgcn"]
92 self.dropout = config["dropout"]
93 self.keep_prob = 1 - self.dropout # Added
94 self.A_split = config["A_split"]
95 self.device = config["device"]
97 self.latent_dim = config["latent_dim_rec"]
98 self.n_layers = config["lightGCN_n_layers"]
99 self.max_entities_per_user = config["max_entities_per_user"]
100 self.kg_dataset = kg_dataset
101 self.gat = GAT(self.latent_dim, self.latent_dim, dropout=0.4, alpha=0.2).train()
103 self.inter_feat = dataset.inter_feat
104 self.num_users = dataset.user_num
105 self.num_items = dataset.item_num
107 self.__init_weight(dataset)
108 self.config = config
110 def __init_weight(self, dataset):
111 self.entity_count = dataset.entity_num
112 self.relation_count = dataset.relation_num
114 self.embedding_user = torch.nn.Embedding(num_embeddings=self.num_users, embedding_dim=self.latent_dim)
115 # item and kg entity
116 self.embedding_item = torch.nn.Embedding(num_embeddings=self.num_items, embedding_dim=self.latent_dim)
117 self.embedding_entity = torch.nn.Embedding(num_embeddings=self.entity_count + 1, embedding_dim=self.latent_dim)
118 self.embedding_relation = torch.nn.Embedding(
119 num_embeddings=self.relation_count + 1, embedding_dim=self.latent_dim
120 )
121 # relation weights
122 self.W_R = nn.Parameter(torch.Tensor(self.relation_count, self.latent_dim, self.latent_dim))
123 nn.init.xavier_uniform_(self.W_R, gain=nn.init.calculate_gain("relu"))
125 nn.init.normal_(self.embedding_user.weight, std=0.1)
126 nn.init.normal_(self.embedding_item.weight, std=0.1)
127 nn.init.normal_(self.embedding_entity.weight, std=0.1)
128 nn.init.normal_(self.embedding_relation.weight, std=0.1)
130 self.f = nn.Sigmoid()
131 self.Graph = dataset.norm_adjacency_matrix().coalesce().to(self.device)
132 self.kg_dict, self.item2relations = self.get_kg_dict(self.num_items)
134 def get_kg_dict(self, item_num):
135 i2es = dict()
136 i2rs = dict()
137 for item in range(item_num):
138 rts = self.kg_dataset.get(item, False)
139 if rts:
140 tails = list(set([ent for tail_list in rts.values() for ent in tail_list]))
141 relations = list(rts.keys())
142 if len(tails) > self.max_entities_per_user:
143 i2es[item] = torch.IntTensor(tails).to(self.device)[: self.max_entities_per_user]
144 i2rs[item] = torch.IntTensor(relations).to(self.device)[: self.max_entities_per_user]
145 else:
146 # last embedding pos as padding idx
147 tails.extend([self.dataset.entity_count] * (self.max_entities_per_user - len(tails)))
148 relations.extend([self.dataset.relation_count] * (self.max_entities_per_user - len(relations)))
149 i2es[item] = torch.IntTensor(tails).to(self.device)
150 i2rs[item] = torch.IntTensor(relations).to(self.device)
151 else:
152 i2es[item] = torch.IntTensor([self.num_items] * self.max_entities_per_user).to(self.device)
153 i2rs[item] = torch.IntTensor([self.relation_count] * self.max_entities_per_user).to(self.device)
154 return i2es, i2rs
156 def computer(self):
157 with torch.no_grad():
158 users_emb = self.embedding_user.weight
159 items_emb = self.cal_item_embedding_from_kg(self.kg_dict)
160 all_emb = torch.cat([users_emb, items_emb])
161 embs = [all_emb]
162 if self.dropout:
163 if self.training:
164 g_droped = self.__dropout(self.keep_prob)
165 else:
166 g_droped = self.Graph
167 else:
168 g_droped = self.Graph
170 for layer in range(self.n_layers):
171 all_emb = torch.sparse.mm(g_droped, all_emb)
172 embs.append(all_emb)
174 embs = torch.stack(embs, dim=1)
175 light_out = torch.mean(embs, dim=1)
176 users, items = torch.split(light_out, [self.num_users, self.num_items])
177 return users, items
179 def __dropout_x(self, x, keep_prob):
180 size = x.size()
181 index = x.indices().t()
182 values = x.values()
183 random_index = torch.rand(len(values)) + keep_prob
184 random_index = random_index.int().bool()
185 index = index[random_index]
186 values = values[random_index] / keep_prob
187 g = torch.sparse_coo_tensor(index.t(), values, size)
188 return g
190 def __dropout(self, keep_prob):
191 if self.A_split:
192 graph = []
193 for g in self.Graph:
194 graph.append(self.__dropout_x(g, keep_prob))
195 else:
196 graph = self.__dropout_x(self.Graph, keep_prob)
197 return graph
199 def cal_item_embedding_from_kg(self, kg: dict):
200 if kg is None:
201 kg = self.kg_dict
203 if self.kgcn == "GAT":
204 return self.cal_item_embedding_gat(kg)
205 elif self.kgcn == "RGAT":
206 return self.cal_item_embedding_rgat(kg)
207 elif self.kgcn == "MEAN":
208 raise NotImplementedError("The 'MEAN' option for kgcn is not yet implemented.")
209 elif self.kgcn == "NO":
210 return self.embedding_item.weight
212 def cal_item_embedding_gat(self, kg: dict):
213 item_embs = self.embedding_item(torch.IntTensor(list(kg.keys())).to(self.device)) # item_num, emb_dim
214 # item_num, entity_num_each
215 item_entities = torch.stack(list(kg.values()))
216 # item_num, entity_num_each, emb_dim
217 entity_embs = self.embedding_entity(item_entities)
218 # item_num, entity_num_each
219 padding_mask = torch.where(
220 item_entities != self.entity_count, torch.ones_like(item_entities), torch.zeros_like(item_entities)
221 ).float()
222 return self.gat(item_embs, entity_embs, padding_mask)
224 def cal_item_embedding_rgat(self, kg: dict):
225 item_embs = self.embedding_item(torch.IntTensor(list(kg.keys())).to(self.device)) # item_num, emb_dim
226 # item_num, entity_num_each
227 item_entities = torch.stack(list(kg.values()))
228 item_relations = torch.stack(list(self.item2relations.values()))
229 # item_num, entity_num_each, emb_dim
230 entity_embs = self.embedding_entity(item_entities)
231 relation_embs = self.embedding_relation(item_relations) # item_num, entity_num_each, emb_dim
232 # w_r = self.W_R[relation_embs] # item_num, entity_num_each, emb_dim, emb_dim
233 # item_num, entity_num_each
234 padding_mask = torch.where(
235 item_entities != self.entity_count, torch.ones_like(item_entities), torch.zeros_like(item_entities)
236 ).float()
237 return self.gat.forward_relation(item_embs, entity_embs, relation_embs, padding_mask)
240class KGLRR(KnowledgeRecommender):
241 """
242 KGLRR: Reinforced logical reasoning over KGs for interpretable recommendation system
243 """
245 input_type = InputType.PAIRWISE
247 def __init__(self, config, dataset) -> None:
248 super().__init__(config, dataset)
250 self.kg_dataset = dataset.ckg_dict_graph()
252 self.encoder = KGEncoder(config, dataset, self.kg_dataset)
253 self.latent_dim = config["latent_dim_rec"]
255 self.r_logic = config["r_logic"]
256 self.r_length = config["r_length"]
257 self.layers = config["layers"]
258 self.sim_scale = config["sim_scale"]
259 self.loss_sum = config["loss_sum"]
260 self.l2s_weight = config["l2_loss"]
261 self.is_explain = config["explain"]
263 self.num_items = dataset.item_num
265 self._init_weights()
266 self.bceloss = nn.BCEWithLogitsLoss()
268 def _init_weights(self):
269 self.true = torch.nn.Parameter(
270 torch.from_numpy(np.random.uniform(0, 1, size=[1, self.latent_dim]).astype(np.float32)),
271 requires_grad=False,
272 )
274 self.and_layer = torch.nn.Linear(self.latent_dim * 2, self.latent_dim)
275 for i in range(self.layers):
276 setattr(self, "and_layer_%d" % i, torch.nn.Linear(self.latent_dim * 2, self.latent_dim * 2))
278 self.or_layer = torch.nn.Linear(self.latent_dim * 2, self.latent_dim)
279 for i in range(self.layers):
280 setattr(self, "or_layer_%d" % i, torch.nn.Linear(self.latent_dim * 2, self.latent_dim * 2))
282 def logic_or(self, vector1, vector2, train=False):
283 vector1, vector2 = self.uniform_size(vector1, vector2, train)
284 vector = torch.cat((vector1, vector2), dim=-1)
285 for i in range(self.layers):
286 vector = F.relu(getattr(self, "or_layer_%d" % i)(vector))
287 vector = self.or_layer(vector)
288 return vector
290 def logic_and(self, vector1, vector2, train=False):
291 vector1, vector2 = self.uniform_size(vector1, vector2, train)
292 vector = torch.cat((vector1, vector2), dim=-1)
293 for i in range(self.layers):
294 vector = F.relu(getattr(self, "and_layer_%d" % i)(vector))
295 vector = self.and_layer(vector)
296 return vector
298 def logic_regularizer(self, train: bool, check_list: list, constraint, constraint_valid):
299 # This function calculates the gap between logical expressions and the real world
301 # length
302 r_length = constraint.norm(dim=2).sum()
303 check_list.append(("r_length", r_length))
305 # and
306 r_and_true = 1 - self.similarity(self.logic_and(constraint, self.true, train=train), constraint)
307 r_and_true = (r_and_true * constraint_valid).sum()
308 check_list.append(("r_and_true", r_and_true))
310 r_and_self = 1 - self.similarity(self.logic_and(constraint, constraint, train=train), constraint)
311 r_and_self = (r_and_self * constraint_valid).sum()
312 check_list.append(("r_and_self", r_and_self))
314 # or
315 r_or_true = 1 - self.similarity(self.logic_or(constraint, self.true, train=train), self.true)
316 r_or_true = (r_or_true * constraint_valid).sum()
317 check_list.append(("r_or_true", r_or_true))
319 r_or_self = 1 - self.similarity(self.logic_or(constraint, constraint, train=train), constraint)
320 r_or_self = (r_or_self * constraint_valid).sum()
321 check_list.append(("r_or_self", r_or_self))
323 r_loss = r_and_true + r_and_self + r_or_true + r_or_self
325 if self.r_logic > 0:
326 r_loss = r_loss * self.r_logic
327 else:
328 r_loss = torch.from_numpy(np.array(0.0, dtype=np.float32)).to(self.device)
329 r_loss.requires_grad = True
331 r_loss += r_length * self.r_length
332 check_list.append(("r_loss", r_loss))
333 return r_loss
335 def similarity(self, vector1, vector2, sigmoid=True):
336 result = F.cosine_similarity(vector1, vector2, dim=-1)
337 result = result * self.sim_scale
338 if sigmoid:
339 return result.sigmoid()
340 return result
342 def uniform_size(self, vector1, vector2, train=False):
343 # Removed vector size normalization
344 if len(vector1.size()) < len(vector2.size()):
345 vector1 = vector1.expand_as(vector2)
346 elif vector2.size() != vector1.size():
347 vector2 = vector2.expand_as(vector1)
348 if train:
349 r12 = torch.Tensor(vector1.size()[:-1]).to(self.device).uniform_(0, 1).bernoulli()
350 r12 = r12.unsqueeze(-1)
351 new_v1 = r12 * vector1 + (-r12 + 1) * vector2
352 new_v2 = r12 * vector2 + (-r12 + 1) * vector1
353 return new_v1, new_v2
354 return vector1, vector2
356 def predict(self, interaction):
357 users = interaction[self.USER_ID]
359 history = self.encoder.user_history_matrix[users, : self.encoder.maxhis] # B * H
360 item_embed = self.encoder.computer()[1] # item_num * V
362 his_valid = history.ge(0).float() # B * H
364 maxlen = int(his_valid.sum(dim=1).max().item())
366 elements = item_embed[history] * his_valid.unsqueeze(-1) # B * H * V
368 tmp_o = None
369 for i in range(maxlen):
370 tmp_o_valid = his_valid[:, i].unsqueeze(-1)
371 if tmp_o is None:
372 tmp_o = elements[:, i, :] * tmp_o_valid # B * V
373 else:
374 # Only perform OR operation if valid; otherwise, if the history is not that long (not valid),
375 # keep the original content unchanged
376 tmp_o = self.logic_or(tmp_o, elements[:, i, :]) * tmp_o_valid + tmp_o * (-tmp_o_valid + 1) # B * V
377 or_vector = tmp_o # B * V
378 left_valid = his_valid[:, 0].unsqueeze(-1) # B * 1
380 prediction = []
381 for i in range(users.size(0)):
382 sent_vector = (
383 left_valid[i] * self.logic_and(or_vector[i].unsqueeze(0).repeat(self.num_items, 1), item_embed)
384 + (-left_valid[i] + 1) * item_embed
385 ) # item_size * V
386 ithpred = self.similarity(sent_vector, self.true, sigmoid=True) # item_size
387 prediction.append(ithpred)
389 prediction = torch.stack(prediction).to(self.device) # [B, item_size]
391 return prediction
393 def explain(self, users, history, items):
394 bs = users.size(0)
395 _, item_embed = self.encoder.computer() # user_num/item_num * V
397 his_valid = history.ge(0).float() # B * H
398 elements = item_embed[history.abs()] * his_valid.unsqueeze(-1) # B * H * V
400 similarity_rlt = []
401 for i in range(bs):
402 tmp_a_valid = his_valid[i, :].unsqueeze(-1) # H
403 tmp_item = items[i].unsqueeze(0).expand(elements[i].size(0), -1) # [H, V]
404 tmp_a = self.logic_and(tmp_item, elements[i]) * tmp_a_valid
405 similarity_rlt.append(self.similarity(tmp_a, self.true))
407 return torch.stack(similarity_rlt).to(self.device) # [H, V]
409 def full_sort_predict(self, interaction):
410 r"""Full sort prediction function.
411 Given users, calculate the scores between users and all candidate items.
413 Args:
414 interaction (Interaction): Interaction class of the batch.
416 Returns:
417 torch.Tensor: Predicted scores for given users and all candidate items,
418 shape: [n_batch_users, n_candidate_items]
419 """
420 # The predict function already does what is needed (users vs all items)
421 prediction = self.predict(interaction)
422 return prediction
424 def predict_or_and(self, users, pos, neg, history):
425 # Store content for checking: logic regularization
426 # Compute L2 regularization on embeddings
427 check_list = []
428 bs = users.size(0)
429 users_embed, item_embed = self.encoder.computer()
431 # Each item in the history is marked as positive, but the latter part of the history may be -1,
432 # indicating it is not that long
433 his_valid = history.ge(0).float() # B * H
434 maxlen = int(his_valid.sum(dim=1).max().item())
435 elements = item_embed[history.abs()] * his_valid.unsqueeze(-1) # B * H * V
437 # For later validation, each vector should satisfy the corresponding constraint in the logical
438 # expression; 'valid' indicates the validity of the corresponding element in the constraint vector
439 constraint = [elements.view([bs, -1, self.latent_dim])] # B * H * V
440 constraint_valid = [his_valid.view([bs, -1])] # B * H
442 tmp_o = None
443 for i in range(maxlen):
444 tmp_o_valid = his_valid[:, i].unsqueeze(-1)
445 if tmp_o is None:
446 tmp_o = elements[:, i, :] * tmp_o_valid # B * V
447 else:
448 # Only perform OR operation if valid; otherwise, if the history is not that long (not valid),
449 # keep the original content unchanged
450 tmp_o = self.logic_or(tmp_o, elements[:, i, :]) * tmp_o_valid + tmp_o * (-tmp_o_valid + 1) # B * V
451 constraint.append(tmp_o.view([bs, 1, self.latent_dim])) # B * 1 * V
452 constraint_valid.append(tmp_o_valid) # B * 1
453 or_vector = tmp_o # B * V
454 left_valid = his_valid[:, 0].unsqueeze(-1) # B * 1
456 right_vector_true = item_embed[pos] # B * V
457 right_vector_false = item_embed[neg] # B * V
459 constraint.append(right_vector_true.view([bs, 1, self.latent_dim])) # B * 1 * V
460 constraint_valid.append(
461 torch.ones((bs, 1)).to(self.device)
462 ) # B * 1 # Indicates that all items to be judged are valid
463 constraint.append(right_vector_false.view([bs, 1, self.latent_dim])) # B * 1 * V
464 constraint_valid.append(torch.ones((bs, 1)).to(self.device)) # B * 1
466 sent_vector = (
467 self.logic_and(or_vector, right_vector_true) * left_valid + (-left_valid + 1) * right_vector_true
468 ) # B * V
469 constraint.append(sent_vector.view([bs, 1, self.latent_dim])) # B * 1 * V
470 constraint_valid.append(left_valid) # B * 1
471 prediction_true = self.similarity(sent_vector, self.true, sigmoid=False).view([-1]) # B
472 check_list.append(("prediction_true", prediction_true))
474 sent_vector = (
475 self.logic_and(or_vector, right_vector_false) * left_valid + (-left_valid + 1) * right_vector_false
476 ) # B * V
477 constraint.append(sent_vector.view([bs, 1, self.latent_dim])) # B * 1 * V
478 constraint_valid.append(left_valid) # B * 1
479 prediction_false = self.similarity(sent_vector, self.true, sigmoid=False).view([-1]) # B
480 check_list.append(("prediction_false", prediction_false))
482 constraint = torch.cat(tuple(constraint), dim=1)
483 constraint_valid = torch.cat(tuple(constraint_valid), dim=1)
485 return prediction_true, prediction_false, check_list, constraint, constraint_valid
487 def calculate_loss(self, interaction):
488 """
489 Calculates the total loss by combining:
490 - BCE Loss (rloss)
491 - Entropy Loss (tloss)
492 - L2 Loss (l2loss)
493 """
494 # Extraction of tensors from the interaction dictionary
495 batch_users = interaction[self.USER_ID]
496 batch_pos = interaction[self.ITEM_ID]
497 batch_neg = interaction[self.NEG_ITEM_ID]
499 # Build the history in the same way as in predict
500 batch_history = self.encoder.user_history_matrix[batch_users, : self.encoder.maxhis]
502 # Forward of the model with the 3 loss components
503 rloss, tloss, l2loss = self.forward(False, 0, batch_users, batch_pos, batch_neg, batch_history)
505 # Combination of the 3 components
506 total_loss = rloss + tloss + l2loss
508 return total_loss
510 def triple_loss(self, TItemScore, FItemScore):
511 bce_loss = self.bceloss(TItemScore.sigmoid(), torch.ones_like(TItemScore)) + self.bceloss(
512 FItemScore.sigmoid(), torch.zeros_like(FItemScore)
513 )
514 # Input positive and negative example scores, maximizing the score difference
515 if self.loss_sum:
516 loss = torch.sum(F.softplus(-(TItemScore - FItemScore)))
517 else:
518 loss = torch.mean(F.softplus(-(TItemScore - FItemScore)))
519 return (loss + bce_loss) * 0.5
521 def l2_loss(self, users, pos, neg, history):
522 users_embed, item_embed = self.encoder.computer()
523 users_emb = users_embed[users]
524 pos_emb = item_embed[pos]
525 neg_emb = item_embed[neg]
526 his_valid = history.ge(0).float() # B * H
527 elements = item_embed[history.abs()] * his_valid.unsqueeze(-1) # B * H * V
528 # L2 regularization loss
529 reg_loss = (
530 (1 / 2)
531 * (users_emb.norm(2).pow(2) + pos_emb.norm(2).pow(2) + neg_emb.norm(2).pow(2) + elements.norm(2).pow(2))
532 / float(len(users))
533 )
534 if not self.loss_sum:
535 reg_loss /= users.size(0)
536 return reg_loss * self.l2s_weight
538 def check(self, check_list):
539 """Logs the shape and contents of tensors in the provided check_list.
541 Each element in check_list is expected to be a tuple where the first item
542 is a string (label) and the second item is a tensor. For each tuple, this
543 function converts the tensor to a NumPy array after detaching it from the
544 computation graph and moving it to CPU, then logs the label, shape and
545 array contents with a threshold of 20 elements for display.
547 Args:
548 check_list (list of tuple): List of (label, tensor) pairs to be logged for inspection.
549 """
551 logging.info(os.linesep)
552 for t in check_list:
553 d = np.array(t[1].detach().cpu())
554 logging.info(os.linesep.join([t[0] + "\t" + str(d.shape), np.array2string(d, threshold=20)]) + os.linesep)
556 def forward(self, print_check: bool, return_pred: bool, *args, **kwards):
557 prediction1, prediction0, check_list, constraint, constraint_valid = self.predict_or_and(*args, **kwards)
558 rloss = self.logic_regularizer(False, check_list, constraint, constraint_valid)
559 tloss = self.triple_loss(prediction1, prediction0)
560 l2loss = self.l2_loss(*args, **kwards)
562 if print_check:
563 self.check(check_list)
565 if return_pred:
566 return prediction1, rloss + tloss + l2loss
567 return rloss, tloss, l2loss
570class GAT(nn.Module):
571 def __init__(self, nfeat, nhid, dropout, alpha):
572 """Dense version of GAT."""
573 super().__init__()
574 self.dropout = dropout
576 self.layer = GraphAttentionLayer(nfeat, nhid, dropout=dropout, alpha=alpha, concat=False)
578 def forward(self, item_embs, entity_embs, adj):
579 x = F.dropout(item_embs, self.dropout, training=self.training)
580 y = F.dropout(entity_embs, self.dropout, training=self.training)
581 x = self.layer(x, y, adj)
582 x = F.dropout(x, self.dropout, training=self.training)
583 return x
585 def forward_relation(self, item_embs, entity_embs, w_r, adj):
586 x = F.dropout(item_embs, self.dropout, training=self.training)
587 y = F.dropout(entity_embs, self.dropout, training=self.training)
588 x = self.layer.forward_relation(x, y, w_r, adj)
589 x = F.dropout(x, self.dropout, training=self.training)
590 return x