Coverage for hopwise/model/knowledge_graph_embedding_recommender/convkb.py: 93%
99 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 : 2024/11/24
2# @Author : Alessandro Soccol
3# @Email : alessandro.soccol@unica.it
5"""ConvKB
6##################################################
7Reference:
8 Nguyen et al. "A Novel Embedding Model for Knowledge Base Completion Based on
9 Convolutional Neural Network." in NAACL 2018.
11Reference code:
12 https://github.com/torchkge-team/torchkge
13"""
15import torch
16from torch import nn
18from hopwise.model.abstract_recommender import KnowledgeRecommender
19from hopwise.model.init import xavier_normal_initialization
20from hopwise.model.loss import LogisticLoss, RegLoss
21from hopwise.utils import InputType
24class ConvKB(KnowledgeRecommender):
25 r"""ConvKB: The main differences from ConvE are that when scoring h, r and t,
26 it concatenates them into a d x 3 matrix. This output undergoes convolution
27 by a set of omega of T filters of shape 1x3, resulting in a Tx3 feature map.
28 This feature map goes through a dense layer with one neuron and weights W, resulting in the final score.
30 Note:
31 In this version, we sample recommender data and knowledge data separately, and put them together for training.
32 """
34 input_type = InputType.PAIRWISE
36 def __init__(self, config, dataset):
37 super().__init__(config, dataset)
39 # Load parameters info
40 self.embedding_size = config["embedding_size"]
41 self.device = config["device"]
42 self.out_channels = config["out_channels"]
43 self.kernel_size = config["kernel_size"]
44 self.drop_prob = config["dropout_prob"]
45 self.lmbda = config["lambda"]
46 self.ui_relation = dataset.field2token_id["relation_id"][dataset.ui_relation]
48 # Embeddings and Layers
49 self.user_embedding = nn.Embedding(self.n_users, self.embedding_size)
50 self.entity_embedding = nn.Embedding(self.n_entities, self.embedding_size)
51 self.relation_embedding = nn.Embedding(self.n_relations, self.embedding_size)
53 self.conv1_bn = nn.BatchNorm2d(1)
54 self.conv_layer = nn.Conv2d(1, self.out_channels, (self.kernel_size, 3))
55 self.conv2_bn = nn.BatchNorm2d(self.out_channels)
56 self.dropout = nn.Dropout(self.drop_prob)
57 self.non_linearity = nn.ReLU()
58 self.fc_layer = nn.Linear((self.embedding_size - self.kernel_size + 1) * self.out_channels, 1, bias=False)
60 # Loss and Regularization
61 self.loss = LogisticLoss()
62 self.reg = RegLoss()
64 # Embeddings Initialization
65 self.apply(xavier_normal_initialization)
67 def forward(self, head, relation, tail):
68 h = head.unsqueeze(1)
69 r = relation.unsqueeze(1)
70 t = tail.unsqueeze(1)
72 conv_input = torch.cat([h, r, t], 1)
73 conv_input = conv_input.transpose(1, 2)
74 conv_input = conv_input.unsqueeze(1)
75 conv_input = self.conv1_bn(conv_input)
76 out_conv = self.conv_layer(conv_input)
77 out_conv = self.conv2_bn(out_conv)
78 out_conv = self.non_linearity(out_conv)
79 out_conv = out_conv.view(-1, (self.embedding_size - self.kernel_size + 1) * self.out_channels)
80 input_fc = self.dropout(out_conv)
81 score = self.fc_layer(input_fc).view(-1)
83 return score
85 def _get_regularization(self, head, relation, tail):
86 l2_reg = torch.mean(head**2) + torch.mean(tail**2) + torch.mean(relation**2)
87 l2_reg = self.reg(self.conv_layer.parameters(), l2_reg)
88 l2_reg = self.reg(self.fc_layer.parameters(), l2_reg)
89 return self.lmbda * l2_reg
91 def _get_rec_embeddings(self, user, positive_items, negative_items):
92 relation_users = torch.tensor([self.ui_relation] * user.shape[0], device=self.device)
93 h = self.user_embedding(user)
94 r = self.relation_embedding(relation_users)
95 t_pos = self.entity_embedding(positive_items)
96 t_neg = self.entity_embedding(negative_items)
97 return h, r, t_pos, t_neg
99 def _get_kg_embeddings(self, head, relation, positive_tails, negative_tails):
100 h = self.entity_embedding(head)
101 r = self.relation_embedding(relation)
102 t_pos = self.entity_embedding(positive_tails)
103 t_neg = self.entity_embedding(negative_tails)
104 return h, r, t_pos, t_neg
106 def calculate_loss(self, interaction):
107 user = interaction[self.USER_ID]
109 pos_item = interaction[self.ITEM_ID]
110 neg_item = interaction[self.NEG_ITEM_ID]
112 head = interaction[self.HEAD_ENTITY_ID]
114 relation = interaction[self.RELATION_ID]
116 pos_tail = interaction[self.TAIL_ENTITY_ID]
117 neg_tail = interaction[self.NEG_TAIL_ENTITY_ID]
119 users_embedding, relations_user_embedding, pos_items_embedding, neg_items_embedding = self._get_rec_embeddings(
120 user, pos_item, neg_item
121 )
122 heads_embedding, relations_kg_embedding, pos_tails_embedding, neg_tails_embedding = self._get_kg_embeddings(
123 head, relation, pos_tail, neg_tail
124 )
126 score_pos_users = self.forward(users_embedding, relations_user_embedding, pos_items_embedding)
127 score_neg_users = self.forward(users_embedding, relations_user_embedding, neg_items_embedding)
128 score_pos_kg = self.forward(heads_embedding, relations_kg_embedding, pos_tails_embedding)
129 score_neg_kg = self.forward(heads_embedding, relations_kg_embedding, neg_tails_embedding)
131 pos_users_reg = self._get_regularization(users_embedding, relations_user_embedding, pos_items_embedding)
132 neg_users_reg = self._get_regularization(users_embedding, relations_user_embedding, neg_items_embedding)
133 pos_kg_reg = self._get_regularization(heads_embedding, relations_kg_embedding, pos_tails_embedding)
134 neg_kg_reg = self._get_regularization(heads_embedding, relations_kg_embedding, neg_tails_embedding)
136 rec_loss = self.loss(-score_pos_users, -score_neg_users, pos_users_reg, neg_users_reg)
137 kg_loss = self.loss(-score_pos_kg, -score_neg_kg, pos_kg_reg, neg_kg_reg)
138 return rec_loss + kg_loss
140 def predict(self, interaction):
141 users = interaction[self.USER_ID]
142 items = interaction[self.ITEM_ID]
143 relations = torch.tensor([self.ui_relation] * users.shape[0], device=self.device)
145 users_e = self.user_embedding(users)
146 relations_e = self.relation_embedding(relations)
147 items_e = self.entity_embedding(items)
149 return self.forward(users_e, relations_e, items_e)
151 def predict_kg(self, interaction):
152 heads = interaction[self.HEAD_ENTITY_ID]
153 relations = interaction[self.RELATION_ID]
154 tails = interaction[self.TAIL_ENTITY_ID]
156 heads_e = self.entity_embedding(heads)
157 relations_e = self.relation_embedding(relations)
158 tails_e = self.entity_embedding(tails)
160 return self.forward(heads_e, relations_e, tails_e)