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

1# @Time : 2024/11/24 

2# @Author : Alessandro Soccol 

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

4 

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. 

10 

11Reference code: 

12 https://github.com/torchkge-team/torchkge 

13""" 

14 

15import torch 

16from torch import nn 

17 

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 

22 

23 

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. 

29 

30 Note: 

31 In this version, we sample recommender data and knowledge data separately, and put them together for training. 

32 """ 

33 

34 input_type = InputType.PAIRWISE 

35 

36 def __init__(self, config, dataset): 

37 super().__init__(config, dataset) 

38 

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] 

47 

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) 

52 

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) 

59 

60 # Loss and Regularization 

61 self.loss = LogisticLoss() 

62 self.reg = RegLoss() 

63 

64 # Embeddings Initialization 

65 self.apply(xavier_normal_initialization) 

66 

67 def forward(self, head, relation, tail): 

68 h = head.unsqueeze(1) 

69 r = relation.unsqueeze(1) 

70 t = tail.unsqueeze(1) 

71 

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) 

82 

83 return score 

84 

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 

90 

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 

98 

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 

105 

106 def calculate_loss(self, interaction): 

107 user = interaction[self.USER_ID] 

108 

109 pos_item = interaction[self.ITEM_ID] 

110 neg_item = interaction[self.NEG_ITEM_ID] 

111 

112 head = interaction[self.HEAD_ENTITY_ID] 

113 

114 relation = interaction[self.RELATION_ID] 

115 

116 pos_tail = interaction[self.TAIL_ENTITY_ID] 

117 neg_tail = interaction[self.NEG_TAIL_ENTITY_ID] 

118 

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 ) 

125 

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) 

130 

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) 

135 

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 

139 

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) 

144 

145 users_e = self.user_embedding(users) 

146 relations_e = self.relation_embedding(relations) 

147 items_e = self.entity_embedding(items) 

148 

149 return self.forward(users_e, relations_e, items_e) 

150 

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] 

155 

156 heads_e = self.entity_embedding(heads) 

157 relations_e = self.relation_embedding(relations) 

158 tails_e = self.entity_embedding(tails) 

159 

160 return self.forward(heads_e, relations_e, tails_e)