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
« 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
5r"""CKE
6##################################################
7Reference:
8 Fuzheng Zhang et al. "Collaborative Knowledge Base Embedding for Recommender Systems." in SIGKDD 2016.
9"""
11import torch
12import torch.nn.functional as F
13from torch import nn
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
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.
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
31 input_type = InputType.PAIRWISE
33 def __init__(self, config, dataset):
34 super().__init__(config, dataset)
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"]
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()
51 # parameters initialization
52 self.apply(xavier_normal_initialization)
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)
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)
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)
70 return h_e, r_e, pos_t_e, neg_t_e, r_trans_w
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
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
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
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]
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
107 rec_loss = self._get_rec_loss(user_e, pos_item_final_e, neg_item_final_e)
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)
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)
116 return rec_loss, kg_loss, reg_loss
118 def predict(self, interaction):
119 user = interaction[self.USER_ID]
120 item = interaction[self.ITEM_ID]
121 return self.forward(user, item)
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)