Coverage for hopwise/model/knowledge_graph_embedding_recommender/transr.py: 78%

97 statements  

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

1# @Time : 2024/11/12 

2# @Author : Alessandro Soccol 

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

4 

5"""TransE 

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

7Reference: 

8 Bordes. A et al. "Translating Embeddings for Modeling Multi-relational Data." in NeurIPS 2013. 

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

23 r"""TransR Rather than introducing relation-specific hyperplanes, it introduces relation-specific spaces. 

24 The scoring functions is the same as TransH but h and t are projected into the space specific to relation 

25 

26 Note: 

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

28 """ 

29 

30 input_type = InputType.PAIRWISE 

31 

32 def __init__(self, config, dataset): 

33 super().__init__(config, dataset) 

34 

35 # Load parameters info 

36 self.embedding_size = config["embedding_size"] 

37 self.margin = config["margin"] 

38 self.device = config["device"] 

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

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 self.proj_mat_e = nn.Embedding(self.n_relations, self.embedding_size * self.embedding_size) 

46 

47 # Loss 

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

49 

50 # Parameters initialization 

51 self.apply(xavier_normal_initialization) 

52 

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

54 user_e = self.user_embedding(user) 

55 pos_item_e = self.entity_embedding(pos_item) 

56 neg_item_e = self.entity_embedding(neg_item) 

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

58 rec_r_e = rec_r_e.expand_as(user_e) 

59 return user_e, pos_item_e, neg_item_e, rec_r_e 

60 

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

62 head_e = self.entity_embedding(head) 

63 pos_tail_e = self.entity_embedding(pos_tail) 

64 neg_tail_e = self.entity_embedding(neg_tail) 

65 relation_e = self.relation_embedding(relation) 

66 return head_e, pos_tail_e, neg_tail_e, relation_e 

67 

68 def forward(self, ent, proj_mat): 

69 proj_e = torch.matmul(proj_mat, ent.unsqueeze(2)) 

70 proj_e = proj_e.squeeze(-1) 

71 return proj_e 

72 

73 def calculate_loss(self, interaction): 

74 user = interaction[self.USER_ID] 

75 

76 pos_item = interaction[self.ITEM_ID] 

77 neg_item = interaction[self.NEG_ITEM_ID] 

78 

79 head = interaction[self.HEAD_ENTITY_ID] 

80 

81 relation = interaction[self.RELATION_ID] 

82 

83 pos_tail = interaction[self.TAIL_ENTITY_ID] 

84 neg_tail = interaction[self.NEG_TAIL_ENTITY_ID] 

85 

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

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

88 

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

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

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

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

93 

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

95 relation = torch.cat([rec_rel, relation]) 

96 

97 proj_mat = self.proj_mat_e(relation).view(h_e.shape[0], self.embedding_size, self.embedding_size) 

98 

99 h_e_proj = self.forward(h_e, proj_mat) 

100 pos_t_e_proj = self.forward(pos_t_e, proj_mat) 

101 neg_t_e_proj = self.forward(neg_t_e, proj_mat) 

102 

103 loss = self.loss(h_e_proj + r_e, pos_t_e_proj, neg_t_e_proj) 

104 

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 

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

115 rec_r_e = rec_r_e.expand_as(user_e) 

116 

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

118 proj_mat = self.proj_mat_e(relation).view(user_e.shape[0], self.embedding_size, self.embedding_size) 

119 

120 user_e_proj = self.forward(user_e, proj_mat) 

121 item_e_proj = self.forward(item_e, proj_mat) 

122 

123 return -torch.norm(user_e_proj + rec_r_e - item_e_proj, p=2, dim=1) 

124 

125 def predict_kg(self, interaction): 

126 head = interaction[self.HEAD_ENTITY_ID] 

127 relation = interaction[self.RELATION_ID] 

128 tail = interaction[self.TAIL_ENTITY_ID] 

129 

130 head_e = self.entity_embedding(head) 

131 tail_e = self.entity_embedding(tail) 

132 

133 rec_r_e = self.relation_embedding(relation) 

134 

135 proj_mat = self.proj_mat_e(relation).view(head_e.shape[0], self.embedding_size, self.embedding_size) 

136 

137 head_e_proj = self.forward(head_e, proj_mat) 

138 tail_e_proj = self.forward(tail_e, proj_mat) 

139 

140 return -torch.norm(head_e_proj + rec_r_e - tail_e_proj, p=2, dim=1) 

141 

142 def full_sort_predict(self, interaction): 

143 user = interaction[self.USER_ID] 

144 user_e = self.user_embedding(user) 

145 

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

147 

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

149 all_item_e = self.entity_embedding.weight[item_indices] 

150 

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

152 relation_items = torch.tensor([self.ui_relation] * all_item_e.shape[0], device=self.device) 

153 

154 proj_mat_user = self.proj_mat_e(relation_users).view(user.shape[0], self.embedding_size, self.embedding_size) 

155 proj_mat_items = self.proj_mat_e(relation_items).view( 

156 all_item_e.shape[0], self.embedding_size, self.embedding_size 

157 ) 

158 

159 user_e_proj = self.forward(user_e, proj_mat_user) 

160 item_e_proj = self.forward(all_item_e, proj_mat_items) 

161 

162 user_e_proj = user_e_proj.unsqueeze(1).expand(-1, item_e_proj.shape[0], -1) 

163 rec_r_e = rec_r_e.unsqueeze(0).expand(1, item_e_proj.shape[0], -1) 

164 item_e_proj = item_e_proj.unsqueeze(0) 

165 

166 return -torch.norm(user_e_proj + rec_r_e - item_e_proj, p=2, dim=2)