Coverage for hopwise/model/knowledge_graph_embedding_recommender/hole.py: 76%

94 statements  

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

1# @Time : 2024/11/21 

2# @Author : Alessandro Soccol 

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

4 

5"""HolE 

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

7Reference: 

8 Nickel et al. "Holographic embeddings of knowledge graphs." in AAAI 2016. 

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

23 r"""HoLE combines the expressive power of RESCAL with the efficiency and simplicity of DistMult. 

24 The entity representations are composed into h ⋆ t in the set of real numbers, 

25 with the circular correlation operator. 

26 

27 Note: 

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

29 """ 

30 

31 input_type = InputType.PAIRWISE 

32 

33 def __init__(self, config, dataset): 

34 super().__init__(config, dataset) 

35 

36 # Load parameters info 

37 self.embedding_size = config["embedding_size"] 

38 self.margin = config["margin"] 

39 self.device = config["device"] 

40 

41 # Embeddings 

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

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

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

45 

46 # Loss 

47 self.sigmoid = nn.Sigmoid() 

48 self.loss = nn.MarginRankingLoss(margin=self.margin) 

49 

50 # Embeddings Initialization 

51 self.apply(xavier_normal_initialization) 

52 

53 def forward(self, h, r, t): 

54 r_e = self.get_rolling_matrix(r) 

55 hr = torch.matmul(h.view(-1, 1, self.embedding_size), r_e) 

56 return (hr.view(-1, self.embedding_size) * t).sum(dim=1) 

57 

58 def get_rolling_matrix(self, x): 

59 b_size, dim = x.shape 

60 x = x.view(b_size, 1, dim) 

61 return torch.cat([x.roll(i, dims=2) for i in range(dim)], dim=1) 

62 

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

64 user_e = self.user_embedding(user) 

65 pos_item_e = self.entity_embedding(pos_item) 

66 neg_item_e = self.entity_embedding(neg_item) 

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

68 rec_r_e = rec_r_e.expand_as(user_e) 

69 

70 return user_e, pos_item_e, neg_item_e, rec_r_e 

71 

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

73 head_e = self.entity_embedding(head) 

74 pos_tail_e = self.entity_embedding(pos_tail) 

75 neg_tail_e = self.entity_embedding(neg_tail) 

76 relation_e = self.relation_embedding(relation) 

77 return head_e, pos_tail_e, neg_tail_e, relation_e 

78 

79 def calculate_loss(self, interaction): 

80 user = interaction[self.USER_ID] 

81 

82 pos_item = interaction[self.ITEM_ID] 

83 neg_item = interaction[self.NEG_ITEM_ID] 

84 

85 head = interaction[self.HEAD_ENTITY_ID] 

86 

87 relation = interaction[self.RELATION_ID] 

88 

89 pos_tail = interaction[self.TAIL_ENTITY_ID] 

90 neg_tail = interaction[self.NEG_TAIL_ENTITY_ID] 

91 

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

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

94 

95 pos_score_users = self.forward(user_e, rec_r_e, pos_item_e) 

96 neg_score_users = self.forward(user_e, rec_r_e, neg_item_e) 

97 

98 pos_score_entities = self.forward(head_e, relation_e, pos_tail_e) 

99 neg_score_entities = self.forward(head_e, relation_e, neg_tail_e) 

100 

101 pos_scores = torch.cat([pos_score_users, pos_score_entities]) 

102 neg_scores = torch.cat([neg_score_users, neg_score_entities]) 

103 

104 loss = self.loss(self.sigmoid(pos_scores), self.sigmoid(neg_scores), torch.ones_like(pos_scores)) 

105 return loss 

106 

107 def predict(self, interaction): 

108 user = interaction[self.USER_ID] 

109 item = interaction[self.ITEM_ID] 

110 

111 user_e = self.user_embedding(user) 

112 item_e = self.entity_embedding(item) 

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

114 rec_r_e = rec_r_e.expand_as(user_e) 

115 

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

117 

118 def predict_kg(self, interaction): 

119 head = interaction[self.HEAD_ENTITY_ID] 

120 relation = interaction[self.RELATION_ID] 

121 tail = interaction[self.TAIL_ENTITY_ID] 

122 

123 head_e = self.entity_embedding(head) 

124 tail_e = self.entity_embedding(tail) 

125 rec_r_e = self.relation_embedding(relation) 

126 

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

128 

129 def full_sort_predict(self, interaction): 

130 user = interaction[self.USER_ID] 

131 user_e = self.user_embedding(user) 

132 

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

134 rec_r_e = rec_r_e.expand_as(user_e) 

135 

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

137 all_item_e = self.entity_embedding.weight[item_indices] 

138 

139 r_e = self.get_rolling_matrix(rec_r_e) 

140 

141 h_e = user_e.view(user_e.shape[0], 1, self.embedding_size) 

142 hr = torch.matmul(h_e, r_e).view(user_e.shape[0], self.embedding_size, 1) 

143 

144 return torch.matmul(hr.squeeze(2), all_item_e.T) 

145 

146 def full_sort_predict_kg(self, interaction): 

147 head = interaction[self.HEAD_ENTITY_ID] 

148 relation = interaction[self.RELATION_ID] 

149 

150 head_e = self.entity_embedding(head) 

151 rec_r_e = self.relation_embedding(relation) 

152 

153 all_item_e = self.entity_embedding.weight 

154 

155 r_e = self.get_rolling_matrix(rec_r_e) 

156 

157 h_e = head_e.view(head_e.shape[0], 1, self.embedding_size) 

158 hr = torch.matmul(h_e, r_e).view(head_e.shape[0], self.embedding_size, 1) 

159 

160 return torch.matmul(hr.squeeze(2), all_item_e.T)