Coverage for hopwise/model/knowledge_graph_embedding_recommender/rescal.py: 74%

91 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"""RESCAL 

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

7Reference: 

8 Nickel et al. "A three-way model for collective learning on multi-relational data." in ICML 2011. 

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

23 r"""RESCAL associates each entity with a vector to capture its latent semantics. 

24 Each relation is represented as a matrix which models pairwise interactions between latent vectors 

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 

40 # Embeddings 

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

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

43 self.relation_embedding = nn.Embedding(self.n_relations, self.embedding_size**2) 

44 

45 # Loss 

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

47 

48 # Parameters initialization 

49 self.apply(xavier_normal_initialization) 

50 

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

52 hr = torch.matmul(head.view(-1, 1, 1, self.embedding_size), relation) 

53 hr = hr.view(-1, self.embedding_size) 

54 return (hr * tail).sum(dim=1) 

55 

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

57 user_e = self.user_embedding(user) 

58 pos_item_e = self.entity_embedding(pos_item) 

59 neg_item_e = self.entity_embedding(neg_item) 

60 rec_r_e = self.relation_embedding.weight[-1].view(1, 1, self.embedding_size, self.embedding_size) 

61 rec_r_e = rec_r_e.repeat(user_e.shape[0], 1, 1, 1) 

62 

63 return user_e, pos_item_e, neg_item_e, rec_r_e 

64 

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

66 head_e = self.entity_embedding(head) 

67 pos_tail_e = self.entity_embedding(pos_tail) 

68 neg_tail_e = self.entity_embedding(neg_tail) 

69 relation_e = self.relation_embedding(relation).view(-1, 1, self.embedding_size, self.embedding_size) 

70 

71 return head_e, pos_tail_e, neg_tail_e, relation_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, pos_tail, neg_tail, relation) 

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 pos_score = self.forward(h_e, r_e, pos_t_e) 

95 neg_score = self.forward(h_e, r_e, neg_t_e) 

96 

97 loss = self.loss(pos_score, neg_score, torch.ones_like(pos_score).to(self.device)) 

98 

99 return loss 

100 

101 def predict(self, interaction): 

102 user = interaction[self.USER_ID] 

103 item = interaction[self.ITEM_ID] 

104 

105 user_e = self.user_embedding(user) 

106 item_e = self.entity_embedding(item) 

107 rec_r_e = self.relation_embedding.weight[-1].view(1, 1, self.embedding_size, self.embedding_size) 

108 rec_r_e = rec_r_e.repeat(user_e.shape[0], 1, 1, 1) 

109 

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

111 

112 def predict_kg(self, interaction): 

113 head = interaction[self.HEAD_ENTITY_ID] 

114 relation = interaction[self.RELATION_ID] 

115 tail = interaction[self.TAIL_ENTITY_ID] 

116 

117 head_e = self.entity_embedding(head) 

118 tail_e = self.entity_embedding(tail) 

119 rec_r_e = self.relation_embedding(relation).view( 

120 relation.shape[0], 1, self.embedding_size, self.embedding_size 

121 ) 

122 

123 return self.forward(head_e, rec_r_e, tail_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].view(1, 1, self.embedding_size, self.embedding_size) 

130 rec_r_e = rec_r_e.repeat(user_e.shape[0], 1, 1, 1) 

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 user_e = user_e.view(-1, 1, 1, self.embedding_size) 

136 hr = torch.matmul(user_e, rec_r_e) 

137 

138 scores = torch.matmul(hr.squeeze(2), all_item_e.T) 

139 scores = scores.squeeze(1) 

140 return scores 

141 

142 def full_sort_predict_kg(self, interaction): 

143 head = interaction[self.HEAD_ENTITY_ID] 

144 relation = interaction[self.RELATION_ID] 

145 

146 head_e = self.entity_embedding(head) 

147 

148 rec_r_e = self.relation_embedding(relation).view( 

149 relation.shape[0], 1, self.embedding_size, self.embedding_size 

150 ) 

151 

152 all_tail_e = self.entity_embedding.weight 

153 

154 head_e = head_e.view(-1, 1, 1, self.embedding_size) 

155 hr = torch.matmul(head_e, rec_r_e) 

156 

157 scores = torch.matmul(hr.squeeze(2), all_tail_e.T) 

158 scores = scores.squeeze(1) 

159 return scores