Coverage for hopwise/model/knowledge_graph_embedding_recommender/rotate.py: 72%

130 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"""RotatE 

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

7Reference: 

8 Sun et al. "RotatE: Knowledge Graph Embedding by Relational Rotation in Complex Space." in ICLR 2019. 

9 

10Reference code: 

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

12""" 

13 

14import math 

15 

16import torch 

17from torch import nn 

18 

19from hopwise.model.abstract_recommender import KnowledgeRecommender 

20from hopwise.model.init import xavier_normal_initialization 

21from hopwise.utils import InputType 

22 

23 

24class RotatE(KnowledgeRecommender): 

25 r"""RotatE models relations as rotations in a complex latent space with h, r, t belonging 

26 to the set of d-dimensional complex numbers. The embedding for r belonging to the set of d-dimensional 

27 complex numbers, is a rotation vector: in all its elements, the phase conveys the rotation along that axis, 

28 and the modulus is equal to 1. 

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.margin = config["margin"] 

42 self.device = config["device"] 

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

44 

45 # Embeddings 

46 self.user_embedding = nn.Embedding(self.n_users, self.embedding_size) 

47 self.user_embedding_im = nn.Embedding(self.n_users, self.embedding_size) 

48 

49 self.entity_embedding = nn.Embedding(self.n_entities, self.embedding_size) 

50 self.entity_embedding_im = nn.Embedding(self.n_entities, self.embedding_size) 

51 

52 self.relation_embedding = nn.Embedding(self.n_relations, self.embedding_size) 

53 

54 # Loss 

55 self.loss = nn.BCEWithLogitsLoss() 

56 

57 # Parameters initialization 

58 self.apply(xavier_normal_initialization) 

59 nn.init.uniform_(self.relation_embedding.weight, 0, 2 * math.pi) 

60 

61 def forward(self, head_re, head_im, relation, tail_re, tail_im): 

62 rel_re, rel_im = torch.cos(relation), torch.sin(relation) 

63 

64 re_score = (rel_re * head_re - rel_im * head_im) - tail_re 

65 im_score = (rel_re * head_im + rel_im * head_re) - tail_im 

66 complex_score = torch.stack([re_score, im_score], dim=2) 

67 score = torch.linalg.vector_norm(complex_score, dim=(1, 2)) 

68 

69 return self.margin - score 

70 

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

72 user_re = self.user_embedding(user) 

73 user_im = self.user_embedding_im(user) 

74 pos_item_re = self.entity_embedding(positive_items) 

75 pos_item_im = self.entity_embedding_im(positive_items) 

76 

77 neg_item_re = self.entity_embedding(negative_items) 

78 neg_item_im = self.entity_embedding_im(negative_items) 

79 

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

81 rec_r_e = self.relation_embedding(relation_user) 

82 

83 return user_re, user_im, rec_r_e, pos_item_re, pos_item_im, neg_item_re, neg_item_im 

84 

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

86 head_re = self.entity_embedding(head) 

87 head_im = self.entity_embedding_im(head) 

88 pos_tail_re = self.entity_embedding(positive_tails) 

89 pos_tail_im = self.entity_embedding_im(positive_tails) 

90 

91 neg_tail_re = self.entity_embedding(negative_tails) 

92 neg_tail_im = self.entity_embedding_im(negative_tails) 

93 

94 kg_r_e = self.relation_embedding(relation) 

95 

96 return head_re, head_im, kg_r_e, pos_tail_re, pos_tail_im, neg_tail_re, neg_tail_im 

97 

98 def calculate_loss(self, interaction): 

99 user = interaction[self.USER_ID] 

100 

101 pos_item = interaction[self.ITEM_ID] 

102 neg_item = interaction[self.NEG_ITEM_ID] 

103 

104 head = interaction[self.HEAD_ENTITY_ID] 

105 

106 relation = interaction[self.RELATION_ID] 

107 

108 pos_tail = interaction[self.TAIL_ENTITY_ID] 

109 neg_tail = interaction[self.NEG_TAIL_ENTITY_ID] 

110 

111 user_re, user_im, rec_r_e, pos_item_re, pos_item_im, neg_item_re, neg_item_im = self._get_rec_embeddings( 

112 user, pos_item, neg_item 

113 ) 

114 head_re, head_im, kg_r_e, pos_tail_re, pos_tail_im, neg_tail_re, neg_tail_im = self._get_kg_embeddings( 

115 head, relation, pos_tail, neg_tail 

116 ) 

117 

118 score_pos_users = self.forward(user_re, user_im, rec_r_e, pos_item_re, pos_item_im) 

119 score_neg_users = self.forward(user_re, user_im, rec_r_e, neg_item_re, neg_item_im) 

120 score_pos_kg = self.forward(head_re, head_im, kg_r_e, pos_tail_re, pos_tail_im) 

121 score_neg_kg = self.forward(head_re, head_im, kg_r_e, neg_tail_re, neg_tail_im) 

122 

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

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

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

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

127 

128 rec_loss = self.loss(scores_rec, labels_rec) 

129 kg_loss = self.loss(scores_kg, labels_kg) 

130 

131 return rec_loss + kg_loss 

132 

133 def predict(self, interaction): 

134 user = interaction[self.USER_ID] 

135 item = interaction[self.ITEM_ID] 

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

137 

138 user_re = self.user_embedding(user) 

139 user_im = self.user_embedding_im(user) 

140 item_re = self.entity_embedding(item) 

141 item_im = self.entity_embedding_im(item) 

142 

143 rec_r_e = self.relation_embedding(relation_user) 

144 

145 return self.forward(user_re, user_im, rec_r_e, item_re, item_im) 

146 

147 def predict_kg(self, interaction): 

148 head = interaction[self.HEAD_ENTITY_ID] 

149 relation = interaction[self.RELATION_ID] 

150 tail = interaction[self.TAIL_ENTITY_ID] 

151 

152 head_re = self.entity_embedding(head) 

153 head_im = self.entity_embedding_im(head) 

154 tail_re = self.entity_embedding(tail) 

155 tail_im = self.entity_embedding_im(tail) 

156 

157 rec_r_e = self.relation_embedding(relation) 

158 

159 return self.forward(head_re, head_im, rec_r_e, tail_re, tail_im) 

160 

161 def full_sort_predict(self, interaction): 

162 user = interaction[self.USER_ID] 

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

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

165 

166 user_re = self.user_embedding(user) 

167 user_im = self.user_embedding_im(user) 

168 

169 item_re = self.entity_embedding(item_indices) 

170 item_im = self.entity_embedding_im(item_indices) 

171 

172 rel_theta = self.relation_embedding(relation_user) 

173 

174 rel_re, rel_im = torch.cos(rel_theta), torch.sin(rel_theta) 

175 

176 user_re = user_re.unsqueeze(1).expand(-1, item_indices.shape[0], -1) 

177 user_im = user_im.unsqueeze(1).expand(-1, item_indices.shape[0], -1) 

178 

179 rel_re = rel_re.unsqueeze(1).expand(-1, item_indices.shape[0], -1) 

180 rel_im = rel_im.unsqueeze(1).expand(-1, item_indices.shape[0], -1) 

181 

182 item_re = item_re.unsqueeze(0) 

183 item_im = item_im.unsqueeze(0) 

184 

185 re_score = (rel_re * user_re - rel_im * user_im) - item_re 

186 im_score = (rel_re * user_im + rel_im * user_re) - item_im 

187 complex_score = torch.stack([re_score, im_score], dim=3) 

188 score = torch.linalg.vector_norm(complex_score, dim=(2, 3)) 

189 

190 return self.margin - score 

191 

192 def full_sort_predict_kg(self, interaction): 

193 head = interaction[self.HEAD_ENTITY_ID] 

194 relation = interaction[self.RELATION_ID] 

195 

196 head_re = self.entity_embedding(head) 

197 head_im = self.entity_embedding_im(head) 

198 

199 tail_re = self.entity_embedding.weight 

200 tail_im = self.entity_embedding_im.weight 

201 

202 rel_theta = self.relation_embedding(relation) 

203 

204 rel_re, rel_im = torch.cos(rel_theta), torch.sin(rel_theta) 

205 

206 head_re = head_re.unsqueeze(1).expand(-1, self.entity_embedding.weight.shape[0], -1) 

207 head_im = head_im.unsqueeze(1).expand(-1, self.entity_embedding.weight.shape[0], -1) 

208 

209 rel_re = rel_re.unsqueeze(1).expand(-1, self.entity_embedding.weight.shape[0], -1) 

210 rel_im = rel_im.unsqueeze(1).expand(-1, self.entity_embedding.weight.shape[0], -1) 

211 

212 tail_re = tail_re.unsqueeze(0) 

213 tail_im = tail_im.unsqueeze(0) 

214 

215 re_score = (rel_re * head_re - rel_im * head_im) - tail_re 

216 im_score = (rel_re * head_im + rel_im * head_re) - tail_im 

217 complex_score = torch.stack([re_score, im_score], dim=3) 

218 score = torch.linalg.vector_norm(complex_score, dim=(2, 3)) 

219 

220 return self.margin - score