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
« 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
5# UPDATE:
6# @Time : 2020/10/19
7# @Author : Yupeng Hou
8# @Email : houyupeng@ruc.edu.cn
10r"""GRU4RecKG
11################################################
12"""
14import torch
15from torch import nn
17from hopwise.model.abstract_recommender import SequentialRecommender
18from hopwise.model.init import xavier_normal_initialization
19from hopwise.model.loss import BPRLoss
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.
26 """
28 def __init__(self, config, dataset):
29 super().__init__(config, dataset)
31 # load dataset info
32 self.entity_embedding_matrix = dataset.get_preload_weight("entity_embedding_id")
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"]
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']!")
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]))
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)
80 item_gru_output, _ = self.item_gru_layers(item_emb) # [B Len H]
81 entity_gru_output, _ = self.entity_gru_layers(entity_emb)
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
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
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
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