Coverage for hopwise/model/sequential_recommender/gru4reckg.py: 89%

73 statements  

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

1# @Time : 2020/10/10 

2# @Author : Shanlei Mu 

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

4 

5# UPDATE: 

6# @Time : 2020/10/19 

7# @Author : Yupeng Hou 

8# @Email : houyupeng@ruc.edu.cn 

9 

10r"""GRU4RecKG 

11################################################ 

12""" 

13 

14import torch 

15from torch import nn 

16 

17from hopwise.model.abstract_recommender import SequentialRecommender 

18from hopwise.model.init import xavier_normal_initialization 

19from hopwise.model.loss import BPRLoss 

20 

21 

22class GRU4RecKG(SequentialRecommender): 

23 r"""It is an extension of GRU4Rec, which concatenates item and its corresponding 

24 pre-trained knowledge graph embedding feature as the input. 

25 

26 """ 

27 

28 def __init__(self, config, dataset): 

29 super().__init__(config, dataset) 

30 

31 # load dataset info 

32 self.entity_embedding_matrix = dataset.get_preload_weight("entity_embedding_id") 

33 

34 # load parameters info 

35 self.embedding_size = config["embedding_size"] 

36 self.hidden_size = config["hidden_size"] 

37 self.num_layers = config["num_layers"] 

38 self.dropout = config["dropout_prob"] 

39 self.freeze_kg = config["freeze_kg"] 

40 self.loss_type = config["loss_type"] 

41 

42 # define layers and loss 

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

44 self.entity_embedding = nn.Embedding(self.n_items, self.embedding_size, padding_idx=0) 

45 self.item_emb_dropout = nn.Dropout(self.dropout) 

46 self.entity_emb_dropout = nn.Dropout(self.dropout) 

47 self.entity_embedding.weight.requires_grad = not self.freeze_kg 

48 self.item_gru_layers = nn.GRU( 

49 input_size=self.embedding_size, 

50 hidden_size=self.hidden_size, 

51 num_layers=self.num_layers, 

52 bias=False, 

53 batch_first=True, 

54 ) 

55 self.entity_gru_layers = nn.GRU( 

56 input_size=self.embedding_size, 

57 hidden_size=self.hidden_size, 

58 num_layers=self.num_layers, 

59 bias=False, 

60 batch_first=True, 

61 ) 

62 self.dense_layer = nn.Linear(self.hidden_size * 2, self.embedding_size) 

63 if self.loss_type == "BPR": 

64 self.loss_fct = BPRLoss() 

65 elif self.loss_type == "CE": 

66 self.loss_fct = nn.CrossEntropyLoss() 

67 else: 

68 raise NotImplementedError("Make sure 'loss_type' in ['BPR', 'CE']!") 

69 

70 # parameters initialization 

71 self.apply(xavier_normal_initialization) 

72 self.entity_embedding.weight.data.copy_(torch.from_numpy(self.entity_embedding_matrix[: self.n_items])) 

73 

74 def forward(self, item_seq, item_seq_len): 

75 item_emb = self.item_embedding(item_seq) 

76 entity_emb = self.entity_embedding(item_seq) 

77 item_emb = self.item_emb_dropout(item_emb) 

78 entity_emb = self.entity_emb_dropout(entity_emb) 

79 

80 item_gru_output, _ = self.item_gru_layers(item_emb) # [B Len H] 

81 entity_gru_output, _ = self.entity_gru_layers(entity_emb) 

82 

83 output_concat = torch.cat((item_gru_output, entity_gru_output), -1) # [B Len 2*H] 

84 output = self.dense_layer(output_concat) 

85 output = self.gather_indexes(output, item_seq_len - 1) # [B H] 

86 return output 

87 

88 def calculate_loss(self, interaction): 

89 item_seq = interaction[self.ITEM_SEQ] 

90 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

91 seq_output = self.forward(item_seq, item_seq_len) 

92 pos_items = interaction[self.POS_ITEM_ID] 

93 if self.loss_type == "BPR": 

94 neg_items = interaction[self.NEG_ITEM_ID] 

95 pos_items_emb = self.item_embedding(pos_items) # [B H] 

96 neg_items_emb = self.item_embedding(neg_items) # [B H] 

97 pos_score = torch.sum(seq_output * pos_items_emb, dim=-1) # [B] 

98 neg_score = torch.sum(seq_output * neg_items_emb, dim=-1) # [B] 

99 loss = self.loss_fct(pos_score, neg_score) 

100 return loss 

101 else: # self.loss_type = 'CE' 

102 test_item_emb = self.item_embedding.weight 

103 logits = torch.matmul(seq_output, test_item_emb.transpose(0, 1)) 

104 loss = self.loss_fct(logits, pos_items) 

105 return loss 

106 

107 def predict(self, interaction): 

108 item_seq = interaction[self.ITEM_SEQ] 

109 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

110 test_item = interaction[self.ITEM_ID] 

111 seq_output = self.forward(item_seq, item_seq_len) 

112 test_item_emb = self.item_embedding(test_item) 

113 scores = torch.mul(seq_output, test_item_emb).sum(dim=1) # [B] 

114 return scores 

115 

116 def full_sort_predict(self, interaction): 

117 item_seq = interaction[self.ITEM_SEQ] 

118 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

119 seq_output = self.forward(item_seq, item_seq_len) 

120 test_items_emb = self.item_embedding.weight 

121 scores = torch.matmul(seq_output, test_items_emb.transpose(0, 1)) # [B n_items] 

122 return scores