Coverage for hopwise/model/knowledge_graph_embedding_recommender/complex.py: 71%

123 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-09-30 13:25 +0000

1# @Time : 2024/11/20 

2# @Author : Alessandro Soccol 

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

4 

5"""ComplEx 

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

7Reference: 

8 Trouillon et al. "Complex embeddings for simple link prediction." in ICML'16. 

9 

10Reference code: 

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

12""" 

13 

14import torch 

15from torch import nn 

16 

17from hopwise.model.abstract_recommender import KnowledgeRecommender 

18from hopwise.model.init import xavier_normal_initialization 

19from hopwise.utils import InputType 

20 

21 

22class ComplEx(KnowledgeRecommender): 

23 r"""ComplEx extends DistMult by introducing complex-valued embeddings. 

24 

25 Note: 

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

27 """ 

28 

29 input_type = InputType.PAIRWISE 

30 

31 def __init__(self, config, dataset): 

32 super().__init__(config, dataset) 

33 

34 # load parameters info 

35 self.embedding_size = config["embedding_size"] 

36 self.device = config["device"] 

37 self.ui_relation = dataset.field2token_id["relation_id"][dataset.ui_relation] 

38 # define layers and loss 

39 self.user_re_embedding = nn.Embedding(self.n_users, self.embedding_size) 

40 self.user_im_embedding = nn.Embedding(self.n_users, self.embedding_size) 

41 

42 self.entity_re_embedding = nn.Embedding(self.n_entities, self.embedding_size) 

43 self.entity_im_embedding = nn.Embedding(self.n_entities, self.embedding_size) 

44 

45 self.relation_re_embedding = nn.Embedding(self.n_relations, self.embedding_size) 

46 self.relation_im_embedding = nn.Embedding(self.n_relations, self.embedding_size) 

47 

48 self.loss = nn.BCEWithLogitsLoss() 

49 

50 # parameters initialization 

51 self.apply(xavier_normal_initialization) 

52 

53 def forward(self, head_re_e, head_im_e, rec_r_re_e, rec_r_im_e, tail_re_e, tail_im_e): 

54 return ( 

55 self.triple_dot(head_re_e, rec_r_re_e, tail_re_e) 

56 + self.triple_dot(head_im_e, rec_r_re_e, tail_im_e) 

57 + self.triple_dot(head_re_e, rec_r_im_e, tail_im_e) 

58 - self.triple_dot(head_im_e, rec_r_im_e, tail_im_e) 

59 ) 

60 

61 def triple_dot(self, x, y, z): 

62 return (x * y * z).sum(dim=-1) 

63 

64 def _get_rec_embeddings(self, user, positive_items, negative_items): 

65 relation_users = torch.tensor([self.ui_relation] * user.shape[0], device=self.device) 

66 user_re_e = self.user_re_embedding(user) 

67 user_im_e = self.user_im_embedding(user) 

68 

69 pos_item_re_e = self.entity_re_embedding(positive_items) 

70 pos_item_im_e = self.entity_im_embedding(positive_items) 

71 

72 neg_item_re_e = self.entity_re_embedding(negative_items) 

73 neg_item_im_e = self.entity_im_embedding(negative_items) 

74 

75 rec_r_re_e = self.relation_re_embedding(relation_users) 

76 rec_r_im_e = self.relation_im_embedding(relation_users) 

77 

78 return user_re_e, user_im_e, rec_r_re_e, rec_r_im_e, pos_item_re_e, pos_item_im_e, neg_item_re_e, neg_item_im_e 

79 

80 def _get_kg_embeddings(self, head, relation, positive_tails, negative_tails): 

81 head_re_e = self.entity_re_embedding(head) 

82 head_im_e = self.entity_im_embedding(head) 

83 

84 pos_tail_re_e = self.entity_re_embedding(positive_tails) 

85 pos_tail_im_e = self.entity_im_embedding(positive_tails) 

86 

87 neg_tail_re_e = self.entity_re_embedding(negative_tails) 

88 neg_tail_im_e = self.entity_im_embedding(negative_tails) 

89 

90 kg_r_re_e = self.relation_re_embedding(relation) 

91 kg_r_im_e = self.relation_im_embedding(relation) 

92 

93 return head_re_e, head_im_e, kg_r_re_e, kg_r_im_e, pos_tail_re_e, pos_tail_im_e, neg_tail_re_e, neg_tail_im_e 

94 

95 def calculate_loss(self, interaction): 

96 user = interaction[self.USER_ID] 

97 

98 pos_item = interaction[self.ITEM_ID] 

99 neg_item = interaction[self.NEG_ITEM_ID] 

100 

101 relation = interaction[self.RELATION_ID] 

102 

103 head = interaction[self.HEAD_ENTITY_ID] 

104 

105 pos_tail = interaction[self.TAIL_ENTITY_ID] 

106 neg_tail = interaction[self.NEG_TAIL_ENTITY_ID] 

107 

108 user_re_e, user_im_e, rec_r_re_e, rec_r_im_e, pos_item_re_e, pos_item_im_e, neg_item_re_e, neg_item_im_e = ( 

109 self._get_rec_embeddings(user, pos_item, neg_item) 

110 ) 

111 head_re_e, head_im_e, kg_r_re_e, kg_r_im_e, pos_tail_re_e, pos_tail_im_e, neg_tail_re_e, neg_tail_im_e = ( 

112 self._get_kg_embeddings(head, relation, pos_tail, neg_tail) 

113 ) 

114 

115 score_pos_users = self.forward(user_re_e, user_im_e, rec_r_re_e, rec_r_im_e, pos_item_re_e, pos_item_im_e) 

116 score_neg_users = self.forward(user_re_e, user_im_e, rec_r_re_e, rec_r_im_e, neg_item_re_e, neg_item_im_e) 

117 score_pos_kg = self.forward(head_re_e, head_im_e, kg_r_re_e, kg_r_im_e, pos_tail_re_e, pos_tail_im_e) 

118 score_neg_kg = self.forward(head_re_e, head_im_e, kg_r_re_e, kg_r_im_e, neg_tail_re_e, neg_tail_im_e) 

119 

120 scores_rec = torch.cat([score_pos_users, score_neg_users], dim=0) 

121 scores_kg = torch.cat([score_pos_kg, score_neg_kg], dim=0) 

122 labels_rec = torch.cat([torch.ones_like(score_pos_users), torch.zeros_like(score_neg_users)], dim=0) 

123 labels_kg = torch.cat([torch.ones_like(score_pos_kg), torch.zeros_like(score_neg_kg)], dim=0) 

124 

125 rec_loss = self.loss(scores_rec, labels_rec) 

126 kg_loss = self.loss(scores_kg, labels_kg) 

127 

128 return rec_loss + kg_loss 

129 

130 def predict(self, interaction): 

131 user = interaction[self.USER_ID] 

132 item = interaction[self.ITEM_ID] 

133 relation = torch.tensor([self.ui_relation] * user.shape[0], device=self.device) 

134 

135 user_re_e = self.user_re_embedding(user) 

136 user_im_e = self.user_im_embedding(user) 

137 

138 item_re_e = self.entity_re_embedding(item) 

139 item_im_e = self.entity_im_embedding(item) 

140 

141 rec_r_re_e = self.relation_re_embedding(relation) 

142 rec_r_im_e = self.relation_im_embedding(relation) 

143 

144 return self.forward(user_re_e, user_im_e, rec_r_re_e, rec_r_im_e, item_re_e, item_im_e) 

145 

146 def predict_kg(self, interaction): 

147 head = interaction[self.HEAD_ENTITY_ID] 

148 relation = interaction[self.RELATION_ID] 

149 tail = interaction[self.TAIL_ENTITY_ID] 

150 

151 head_re_e = self.entity_re_embedding(head) 

152 head_im_e = self.entity_im_embedding(head) 

153 

154 tail_re_e = self.entity_re_embedding(tail) 

155 tail_im_e = self.entity_im_embedding(tail) 

156 

157 rec_r_re_e = self.relation_re_embedding(relation) 

158 rec_r_im_e = self.relation_im_embedding(relation) 

159 

160 return self.forward(head_re_e, head_im_e, rec_r_re_e, rec_r_im_e, tail_re_e, tail_im_e) 

161 

162 def full_sort_predict(self, interaction): 

163 user = interaction[self.USER_ID] 

164 user_re_e = self.user_re_embedding(user) 

165 user_im_e = self.user_im_embedding(user) 

166 

167 rec_r_re_e = self.relation_re_embedding.weight[-1] 

168 rec_r_im_e = self.relation_im_embedding.weight[-1] 

169 rec_r_re_e = rec_r_re_e.expand_as(user_re_e) 

170 rec_r_im_e = rec_r_im_e.expand_as(user_re_e) 

171 

172 item_indices = torch.tensor(range(self.n_items)).to(self.device) 

173 all_item_re_e = self.entity_re_embedding.weight[item_indices] 

174 all_item_im_e = self.entity_im_embedding.weight[item_indices] 

175 

176 user_re_e = user_re_e.unsqueeze(1).expand(-1, all_item_re_e.shape[0], -1) 

177 user_im_e = user_im_e.unsqueeze(1).expand(-1, all_item_re_e.shape[0], -1) 

178 

179 rec_r_re_e = rec_r_re_e.unsqueeze(1).expand(-1, all_item_re_e.shape[0], -1) 

180 rec_r_im_e = rec_r_im_e.unsqueeze(1).expand(-1, all_item_re_e.shape[0], -1) 

181 

182 all_item_re_e = all_item_re_e.unsqueeze(0) 

183 all_item_im_e = all_item_im_e.unsqueeze(0) 

184 

185 return ( 

186 self.triple_dot(user_re_e, rec_r_re_e, all_item_re_e) 

187 + self.triple_dot(user_im_e, rec_r_re_e, all_item_im_e) 

188 + self.triple_dot(user_re_e, rec_r_im_e, all_item_im_e) 

189 - self.triple_dot(user_im_e, rec_r_im_e, all_item_im_e) 

190 ) 

191 

192 def full_sort_predict_kg(self, interaction): 

193 head = interaction[self.HEAD_ENTITY_ID] 

194 relation = interaction[self.RELATION_ID] 

195 head_re_e = self.entity_re_embedding(head) 

196 head_im_e = self.entity_im_embedding(head) 

197 

198 rec_r_re_e = self.relation_re_embedding(relation) 

199 rec_r_im_e = self.relation_im_embedding(relation) 

200 

201 entity_indices = torch.tensor(range(self.n_entities)).to(self.device) 

202 all_entity_re_e = self.entity_re_embedding.weight[entity_indices] 

203 all_entity_im_e = self.entity_im_embedding.weight[entity_indices] 

204 

205 head_re_e = head_re_e.unsqueeze(1).expand(-1, all_entity_re_e.shape[0], -1) 

206 head_im_e = head_im_e.unsqueeze(1).expand(-1, all_entity_re_e.shape[0], -1) 

207 

208 rec_r_re_e = rec_r_re_e.unsqueeze(1).expand(-1, all_entity_re_e.shape[0], -1) 

209 rec_r_im_e = rec_r_im_e.unsqueeze(1).expand(-1, all_entity_re_e.shape[0], -1) 

210 

211 all_entity_re_e = all_entity_re_e.unsqueeze(0) 

212 all_entity_im_e = all_entity_im_e.unsqueeze(0) 

213 

214 return ( 

215 self.triple_dot(head_re_e, rec_r_re_e, all_entity_re_e) 

216 + self.triple_dot(head_im_e, rec_r_re_e, all_entity_im_e) 

217 + self.triple_dot(head_re_e, rec_r_im_e, all_entity_im_e) 

218 - self.triple_dot(head_im_e, rec_r_im_e, all_entity_im_e) 

219 )