Coverage for hopwise/model/knowledge_graph_embedding_recommender/toruse.py: 64%

106 statements  

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

1# @Time : 2024/11/14 

2# @Author : Alessandro Soccol 

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

4 

5"""TorusE 

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

7Reference: 

8 Takuma Ebisu and Ryutaro Ichise. "TorusE: Knowledge Graph Embedding on a Lie Group." in AAAI 2018. 

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 TorusE(KnowledgeRecommender): 

23 r"""TorusE projects each point in a Torus. 

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

37 self.device = config["device"] 

38 

39 # Embeddings 

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

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

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

43 

44 # Loss 

45 self.loss = nn.TripletMarginLoss(margin=self.margin, p=2, reduction="mean") 

46 

47 # Parameters initialization 

48 self.apply(xavier_normal_initialization) 

49 

50 def _get_rec_embedding(self, user, pos_item, neg_item): 

51 user_e = self.user_embedding(user) 

52 pos_item_e = self.entity_embedding(pos_item) 

53 neg_item_e = self.entity_embedding(neg_item) 

54 rec_r_e = self.relation_embedding.weight[-1] 

55 rec_r_e = rec_r_e.expand_as(user_e) 

56 

57 return user_e, pos_item_e, neg_item_e, rec_r_e 

58 

59 def _get_kg_embedding(self, head, pos_tail, neg_tail, relation): 

60 head_e = self.entity_embedding(head) 

61 pos_tail_e = self.entity_embedding(pos_tail) 

62 neg_tail_e = self.entity_embedding(neg_tail) 

63 relation_e = self.relation_embedding(relation) 

64 return head_e, pos_tail_e, neg_tail_e, relation_e 

65 

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

67 h_e = head.clone() 

68 r_e = relation.clone() 

69 t_e = tail.clone() 

70 

71 h_e.data.frac_() 

72 r_e.data.frac_() 

73 t_e.data.frac_() 

74 

75 h_r = h_e + r_e 

76 return -(4 * torch.min((h_r - t_e) ** 2, 1 - (h_r - t_e) ** 2).sum(dim=-1)) 

77 

78 def calculate_loss(self, interaction): 

79 user = interaction[self.USER_ID] 

80 

81 pos_item = interaction[self.ITEM_ID] 

82 neg_item = interaction[self.NEG_ITEM_ID] 

83 

84 head = interaction[self.HEAD_ENTITY_ID] 

85 

86 relation = interaction[self.RELATION_ID] 

87 

88 pos_tail = interaction[self.TAIL_ENTITY_ID] 

89 neg_tail = interaction[self.NEG_TAIL_ENTITY_ID] 

90 

91 user_e, pos_item_e, neg_item_e, rec_r_e = self._get_rec_embedding(user, pos_item, neg_item) 

92 head_e, pos_tail_e, neg_tail_e, relation_e = self._get_kg_embedding(head, pos_tail, neg_tail, relation) 

93 

94 h_e = torch.cat([user_e, head_e]) 

95 r_e = torch.cat([rec_r_e, relation_e]) 

96 pos_t_e = torch.cat([pos_item_e, pos_tail_e]) 

97 neg_t_e = torch.cat([neg_item_e, neg_tail_e]) 

98 

99 loss = self.loss(h_e + r_e, pos_t_e, neg_t_e) 

100 

101 return loss 

102 

103 def predict_kg(self, interaction): 

104 head = interaction[self.HEAD_ENTITY_ID] 

105 relation = interaction[self.RELATION_ID] 

106 tail = interaction[self.TAIL_ENTITY_ID] 

107 

108 head_e = self.entity_embedding(head) 

109 tail_e = self.entity_embedding(tail) 

110 rec_r_e = self.relation_embedding(relation) 

111 

112 return self.forward(head_e, rec_r_e, tail_e) 

113 

114 def predict(self, interaction): 

115 user = interaction[self.USER_ID] 

116 item = interaction[self.ITEM_ID] 

117 

118 user_e = self.user_embedding(user) 

119 item_e = self.entity_embedding(item) 

120 rec_r_e = self.relation_embedding.weight[-1] 

121 rec_r_e = rec_r_e.expand_as(user_e) 

122 

123 return self.forward(user_e, rec_r_e, item_e) 

124 

125 def full_sort_predict(self, interaction): 

126 user = interaction[self.USER_ID] 

127 user_e = self.user_embedding(user) 

128 

129 rec_r_e = self.relation_embedding.weight[-1] 

130 rec_r_e = rec_r_e.expand_as(user_e) 

131 

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

133 all_item_e = self.entity_embedding.weight[item_indices] 

134 

135 h_e = user_e.unsqueeze(1).expand(-1, all_item_e.shape[0], -1) 

136 r_e = rec_r_e.unsqueeze(1).expand(-1, all_item_e.shape[0], -1) 

137 t = all_item_e.unsqueeze(0) 

138 

139 h_e = h_e.clone() 

140 r_e = r_e.clone() 

141 t = t.clone() 

142 

143 h_e.data.frac_() 

144 r_e.data.frac_() 

145 t.data.frac_() 

146 

147 h_r = h_e + r_e 

148 return -(4 * torch.min((h_r - t) ** 2, 1 - (h_r - t) ** 2).sum(dim=-1)) 

149 

150 def full_sort_predict_kg(self, interaction): 

151 head = interaction[self.HEAD_ENTITY_ID] 

152 relation = interaction[self.RELATION_ID] 

153 

154 head_e = self.entity_embedding(head) 

155 rec_r_e = self.relation_embedding(relation) 

156 

157 all_tail_e = self.entity_embedding.weight 

158 

159 h_e = head_e.unsqueeze(1).expand(-1, all_tail_e.shape[0], -1) 

160 r_e = rec_r_e.unsqueeze(1).expand(-1, all_tail_e.shape[0], -1) 

161 t = all_tail_e.unsqueeze(0) 

162 

163 h_e = h_e.clone() 

164 r_e = r_e.clone() 

165 t = t.clone() 

166 

167 h_e.data.frac_() 

168 r_e.data.frac_() 

169 t.data.frac_() 

170 

171 h_r = h_e + r_e 

172 return -(4 * torch.min((h_r - t) ** 2, 1 - (h_r - t) ** 2).sum(dim=-1))