Coverage for hopwise/model/sequential_recommender/gru4rec.py: 88%
68 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/17 19:38
2# @Author : Yujie Lu
3# @Email : yujielu1998@gmail.com
5# UPDATE:
6# @Time : 2020/8/19, 2020/10/2
7# @Author : Yupeng Hou, Yujie Lu
8# @Email : houyupeng@ruc.edu.cn, yujielu1998@gmail.com
10r"""GRU4Rec
11################################################
13Reference:
14 Yong Kiam Tan et al. "Improved Recurrent Neural Networks for Session-based Recommendations." in DLRS 2016.
16"""
18import torch
19from torch import nn
20from torch.nn.init import xavier_normal_, xavier_uniform_
22from hopwise.model.abstract_recommender import SequentialRecommender
23from hopwise.model.loss import BPRLoss
26class GRU4Rec(SequentialRecommender):
27 r"""GRU4Rec is a model that incorporate RNN for recommendation.
29 Note:
30 Regarding the innovation of this article,we can only achieve the data augmentation mentioned
31 in the paper and directly output the embedding of the item,
32 in order that the generation method we used is common to other sequential models.
33 """
35 def __init__(self, config, dataset):
36 super().__init__(config, dataset)
38 # load parameters info
39 self.embedding_size = config["embedding_size"]
40 self.hidden_size = config["hidden_size"]
41 self.loss_type = config["loss_type"]
42 self.num_layers = config["num_layers"]
43 self.dropout_prob = config["dropout_prob"]
45 # define layers and loss
46 self.item_embedding = nn.Embedding(self.n_items, self.embedding_size, padding_idx=0)
47 self.emb_dropout = nn.Dropout(self.dropout_prob)
48 self.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.dense = nn.Linear(self.hidden_size, self.embedding_size)
56 if self.loss_type == "BPR":
57 self.loss_fct = BPRLoss()
58 elif self.loss_type == "CE":
59 self.loss_fct = nn.CrossEntropyLoss()
60 else:
61 raise NotImplementedError("Make sure 'loss_type' in ['BPR', 'CE']!")
63 # parameters initialization
64 self.apply(self._init_weights)
66 def _init_weights(self, module):
67 if isinstance(module, nn.Embedding):
68 xavier_normal_(module.weight)
69 elif isinstance(module, nn.GRU):
70 xavier_uniform_(module.weight_hh_l0)
71 xavier_uniform_(module.weight_ih_l0)
73 def forward(self, item_seq, item_seq_len):
74 item_seq_emb = self.item_embedding(item_seq)
75 item_seq_emb_dropout = self.emb_dropout(item_seq_emb)
76 gru_output, _ = self.gru_layers(item_seq_emb_dropout)
77 gru_output = self.dense(gru_output)
78 # the embedding of the predicted item, shape of (batch_size, embedding_size)
79 seq_output = self.gather_indexes(gru_output, item_seq_len - 1)
80 return seq_output
82 def calculate_loss(self, interaction):
83 item_seq = interaction[self.ITEM_SEQ]
84 item_seq_len = interaction[self.ITEM_SEQ_LEN]
85 seq_output = self.forward(item_seq, item_seq_len)
86 pos_items = interaction[self.POS_ITEM_ID]
87 if self.loss_type == "BPR":
88 neg_items = interaction[self.NEG_ITEM_ID]
89 pos_items_emb = self.item_embedding(pos_items)
90 neg_items_emb = self.item_embedding(neg_items)
91 pos_score = torch.sum(seq_output * pos_items_emb, dim=-1) # [B]
92 neg_score = torch.sum(seq_output * neg_items_emb, dim=-1) # [B]
93 loss = self.loss_fct(pos_score, neg_score)
94 return loss
95 else: # self.loss_type = 'CE'
96 test_item_emb = self.item_embedding.weight
97 logits = torch.matmul(seq_output, test_item_emb.transpose(0, 1))
98 loss = self.loss_fct(logits, pos_items)
99 return loss
101 def predict(self, interaction):
102 item_seq = interaction[self.ITEM_SEQ]
103 item_seq_len = interaction[self.ITEM_SEQ_LEN]
104 test_item = interaction[self.ITEM_ID]
105 seq_output = self.forward(item_seq, item_seq_len)
106 test_item_emb = self.item_embedding(test_item)
107 scores = torch.mul(seq_output, test_item_emb).sum(dim=1) # [B]
108 return scores
110 def full_sort_predict(self, interaction):
111 item_seq = interaction[self.ITEM_SEQ]
112 item_seq_len = interaction[self.ITEM_SEQ_LEN]
113 seq_output = self.forward(item_seq, item_seq_len)
114 test_items_emb = self.item_embedding.weight
115 scores = torch.matmul(seq_output, test_items_emb.transpose(0, 1)) # [B, n_items]
116 return scores