Coverage for hopwise/model/knowledge_aware_recommender/cke.py: 91%

82 statements  

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

1# @Time : 2020/8/6 

2# @Author : Shanlei Mu 

3# @Email : slmu@ruc.edu.cn 

4 

5r"""CKE 

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

7Reference: 

8 Fuzheng Zhang et al. "Collaborative Knowledge Base Embedding for Recommender Systems." in SIGKDD 2016. 

9""" 

10 

11import torch 

12import torch.nn.functional as F 

13from torch import nn 

14 

15from hopwise.model.abstract_recommender import KnowledgeRecommender 

16from hopwise.model.init import xavier_normal_initialization 

17from hopwise.model.loss import BPRLoss, EmbLoss 

18from hopwise.utils import InputType 

19 

20 

21class CKE(KnowledgeRecommender): 

22 r"""CKE is a knowledge-based recommendation model, it can incorporate KG and other information such as corresponding 

23 images to enrich the representation of items for item recommendations. 

24 

25 Note: 

26 In the original paper, CKE used structural knowledge, textual knowledge and visual knowledge. In our 

27 implementation, we only used structural knowledge. Meanwhile, the version we implemented uses a simpler 

28 regular way which can get almost the same result (even better) as the original regular way. 

29 """ # noqa: E501 

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

39 self.reg_weights = config["reg_weights"] 

40 

41 # define layers and loss 

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

43 self.item_embedding = nn.Embedding(self.n_items, self.embedding_size) 

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

45 self.relation_embedding = nn.Embedding(self.n_relations, self.kg_embedding_size) 

46 self.trans_w = nn.Embedding(self.n_relations, self.embedding_size * self.kg_embedding_size) 

47 self.rec_loss = BPRLoss() 

48 self.kg_loss = BPRLoss() 

49 self.reg_loss = EmbLoss() 

50 

51 # parameters initialization 

52 self.apply(xavier_normal_initialization) 

53 

54 def _get_kg_embedding(self, h, r, pos_t, neg_t): 

55 h_e = self.entity_embedding(h).unsqueeze(1) 

56 pos_t_e = self.entity_embedding(pos_t).unsqueeze(1) 

57 neg_t_e = self.entity_embedding(neg_t).unsqueeze(1) 

58 r_e = self.relation_embedding(r) 

59 r_trans_w = self.trans_w(r).view(r.size(0), self.embedding_size, self.kg_embedding_size) 

60 

61 h_e = torch.bmm(h_e, r_trans_w).squeeze(1) 

62 pos_t_e = torch.bmm(pos_t_e, r_trans_w).squeeze(1) 

63 neg_t_e = torch.bmm(neg_t_e, r_trans_w).squeeze(1) 

64 

65 r_e = F.normalize(r_e, p=2, dim=1) 

66 h_e = F.normalize(h_e, p=2, dim=1) 

67 pos_t_e = F.normalize(pos_t_e, p=2, dim=1) 

68 neg_t_e = F.normalize(neg_t_e, p=2, dim=1) 

69 

70 return h_e, r_e, pos_t_e, neg_t_e, r_trans_w 

71 

72 def forward(self, user, item): 

73 u_e = self.user_embedding(user) 

74 i_e = self.item_embedding(item) + self.entity_embedding(item) 

75 score = torch.mul(u_e, i_e).sum(dim=1) 

76 return score 

77 

78 def _get_rec_loss(self, user_e, pos_e, neg_e): 

79 pos_score = torch.mul(user_e, pos_e).sum(dim=1) 

80 neg_score = torch.mul(user_e, neg_e).sum(dim=1) 

81 rec_loss = self.rec_loss(pos_score, neg_score) 

82 return rec_loss 

83 

84 def _get_kg_loss(self, h_e, r_e, pos_e, neg_e): 

85 pos_tail_score = ((h_e + r_e - pos_e) ** 2).sum(dim=1) 

86 neg_tail_score = ((h_e + r_e - neg_e) ** 2).sum(dim=1) 

87 kg_loss = self.kg_loss(neg_tail_score, pos_tail_score) 

88 return kg_loss 

89 

90 def calculate_loss(self, interaction): 

91 user = interaction[self.USER_ID] 

92 pos_item = interaction[self.ITEM_ID] 

93 neg_item = interaction[self.NEG_ITEM_ID] 

94 h = interaction[self.HEAD_ENTITY_ID] 

95 r = interaction[self.RELATION_ID] 

96 pos_t = interaction[self.TAIL_ENTITY_ID] 

97 neg_t = interaction[self.NEG_TAIL_ENTITY_ID] 

98 

99 user_e = self.user_embedding(user) 

100 pos_item_e = self.item_embedding(pos_item) 

101 neg_item_e = self.item_embedding(neg_item) 

102 pos_item_kg_e = self.entity_embedding(pos_item) 

103 neg_item_kg_e = self.entity_embedding(neg_item) 

104 pos_item_final_e = pos_item_e + pos_item_kg_e 

105 neg_item_final_e = neg_item_e + neg_item_kg_e 

106 

107 rec_loss = self._get_rec_loss(user_e, pos_item_final_e, neg_item_final_e) 

108 

109 h_e, r_e, pos_t_e, neg_t_e, r_trans_w = self._get_kg_embedding(h, r, pos_t, neg_t) 

110 kg_loss = self._get_kg_loss(h_e, r_e, pos_t_e, neg_t_e) 

111 

112 reg_loss = self.reg_weights[0] * self.reg_loss(user_e, pos_item_final_e, neg_item_final_e) + self.reg_weights[ 

113 1 

114 ] * self.reg_loss(h_e, r_e, pos_t_e, neg_t_e) 

115 

116 return rec_loss, kg_loss, reg_loss 

117 

118 def predict(self, interaction): 

119 user = interaction[self.USER_ID] 

120 item = interaction[self.ITEM_ID] 

121 return self.forward(user, item) 

122 

123 def full_sort_predict(self, interaction): 

124 user = interaction[self.USER_ID] 

125 user_e = self.user_embedding(user) 

126 all_item_e = self.item_embedding.weight + self.entity_embedding.weight[: self.n_items] 

127 score = torch.matmul(user_e, all_item_e.transpose(0, 1)) 

128 return score.view(-1)