Coverage for hopwise/model/knowledge_aware_recommender/cfkg.py: 100%

58 statements  

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

1# @Time : 2020/9/14 

2# @Author : Shanlei Mu 

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

4 

5"""CFKG 

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

7Reference: 

8 Qingyao Ai et al. "Learning heterogeneous knowledge base embeddings for explainable recommendation." in MDPI 2018. 

9""" 

10 

11import torch 

12from torch import nn 

13 

14from hopwise.model.abstract_recommender import KnowledgeRecommender 

15from hopwise.model.init import xavier_normal_initialization 

16from hopwise.model.loss import InnerProductLoss 

17from hopwise.utils import InputType 

18 

19 

20class CFKG(KnowledgeRecommender): 

21 r"""CFKG is a knowledge-based recommendation model, it combines knowledge graph and the user-item interaction 

22 graph to a new graph. In this graph, user, item and related attribute are viewed as entities, and the interaction 

23 between user and item and the link between item and attribute are viewed as relations. It define a new score 

24 function as follows: 

25 

26 .. math:: 

27 d (u_i + r_{buy}, v_j) 

28 

29 Note: 

30 In the original paper, CFKG puts recommender data (u-i interaction) and knowledge data (h-r-t) together 

31 for sampling and mix them for training. In this version, we sample recommender data 

32 and knowledge data separately, and put them together for training. 

33 """ 

34 

35 input_type = InputType.PAIRWISE 

36 

37 def __init__(self, config, dataset): 

38 super().__init__(config, dataset) 

39 

40 # load parameters info 

41 self.embedding_size = config["embedding_size"] 

42 

43 # define layers and loss 

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

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

46 self.relation_embedding = nn.Embedding(self.n_relations + 1, self.embedding_size) 

47 self.rec_loss = InnerProductLoss() 

48 

49 # parameters initialization 

50 self.apply(xavier_normal_initialization) 

51 

52 def forward(self, user, item): 

53 user_e = self.user_embedding(user) 

54 item_e = self.entity_embedding(item) 

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

56 rec_r_e = rec_r_e.expand_as(user_e) 

57 score = self._get_score(user_e, item_e, rec_r_e) 

58 return score 

59 

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

61 user_e = self.user_embedding(user) 

62 pos_item_e = self.entity_embedding(pos_item) 

63 neg_item_e = self.entity_embedding(neg_item) 

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

65 rec_r_e = rec_r_e.expand_as(user_e) 

66 

67 return user_e, pos_item_e, neg_item_e, rec_r_e 

68 

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

70 head_e = self.entity_embedding(head) 

71 pos_tail_e = self.entity_embedding(pos_tail) 

72 neg_tail_e = self.entity_embedding(neg_tail) 

73 relation_e = self.relation_embedding(relation) 

74 return head_e, pos_tail_e, neg_tail_e, relation_e 

75 

76 def _get_score(self, h_e, t_e, r_e): 

77 return torch.mul(h_e + r_e, t_e).sum(dim=1) 

78 

79 def calculate_loss(self, interaction): 

80 user = interaction[self.USER_ID] 

81 pos_item = interaction[self.ITEM_ID] 

82 neg_item = interaction[self.NEG_ITEM_ID] 

83 head = interaction[self.HEAD_ENTITY_ID] 

84 relation = interaction[self.RELATION_ID] 

85 pos_tail = interaction[self.TAIL_ENTITY_ID] 

86 neg_tail = interaction[self.NEG_TAIL_ENTITY_ID] 

87 

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

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

90 

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

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

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

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

95 

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

97 

98 return loss 

99 

100 def predict(self, interaction): 

101 user = interaction[self.USER_ID] 

102 item = interaction[self.ITEM_ID] 

103 return self.forward(user, item)