Coverage for hopwise/model/knowledge_graph_embedding_recommender/transh.py: 86%

85 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"""TransH 

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

7Reference: 

8 Wang Z. et al. "Knowledge Graph Embedding by Translating on Hyperplanes." in AAAI 2014. 

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

23 r"""TransH Have been invented to overcome the disadvantages of TransE, 

24 allowing an entity to have distinct representations when involved in different relations. 

25 It introduces relation-specific hyperplanes. 

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 self.ui_relation = dataset.field2token_id["relation_id"][dataset.ui_relation] 

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

46 

47 # Loss 

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

49 

50 # Parameters initialization 

51 self.apply(xavier_normal_initialization) 

52 

53 def forward(self, head, relation, tail, relation_ids): 

54 head_proj = self.project(head, relation_ids) 

55 tail_proj = self.project(tail, relation_ids) 

56 score = -torch.norm(head_proj + relation - tail_proj, p=2, dim=1) 

57 return score 

58 

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

60 user_e = self.user_embedding(user) 

61 pos_item_e = self.entity_embedding(pos_item) 

62 neg_item_e = self.entity_embedding(neg_item) 

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

64 rec_r_e = rec_r_e.expand_as(user_e) 

65 

66 return user_e, pos_item_e, neg_item_e, rec_r_e 

67 

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

69 head_e = self.entity_embedding(head) 

70 pos_tail_e = self.entity_embedding(pos_tail) 

71 neg_tail_e = self.entity_embedding(neg_tail) 

72 relation_e = self.relation_embedding(relation) 

73 return head_e, pos_tail_e, neg_tail_e, relation_e 

74 

75 def project(self, ent, rel): 

76 return ent - (ent * self.norm_vec(rel).sum(1).view(-1, 1)) * self.norm_vec(rel) 

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 relation_user = torch.tensor([self.ui_relation] * user_e.shape[0], device=self.device) 

95 # Projections 

96 user_e = self.project(user_e, relation_user) 

97 head_e = self.project(head_e, relation) 

98 pos_item_e = self.project(pos_item_e, relation_user) 

99 pos_tail_e = self.project(pos_tail_e, relation) 

100 neg_item_e = self.project(neg_item_e, relation_user) 

101 neg_tail_e = self.project(neg_tail_e, relation) 

102 

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

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

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

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

107 

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

109 

110 return loss 

111 

112 def predict(self, interaction): 

113 user = interaction[self.USER_ID] 

114 item = interaction[self.ITEM_ID] 

115 

116 user_e = self.user_embedding(user) 

117 

118 item_e = self.entity_embedding(item) 

119 

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

121 rec_r_e = rec_r_e.expand_as(user_e) 

122 

123 relation_ids = torch.tensor([self.ui_relation] * user_e.shape[0], device=self.device) 

124 

125 return self.forward(user_e, rec_r_e, item_e, relation_ids) 

126 

127 def full_sort_predict(self, interaction): 

128 user = interaction[self.USER_ID] 

129 user_e = self.user_embedding(user) 

130 

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

132 rec_r_e = rec_r_e.expand_as(user_e) 

133 

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

135 all_item_e = self.entity_embedding.weight[item_indices] 

136 

137 relation_ids_user = torch.tensor([self.ui_relation] * user_e.shape[0], device=self.device) 

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

139 

140 h_r = self.project(user_e, relation_ids_user) + rec_r_e 

141 h_r = h_r.unsqueeze(1).expand(-1, all_item_e.shape[0], -1) 

142 

143 t = self.project(all_item_e, relation_ids_item) 

144 t = t.unsqueeze(0) 

145 return -torch.norm(h_r - t, p=2, dim=2)