Coverage for hopwise/model/sequential_recommender/narm.py: 89%
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/25 19:56
2# @Author : Yujie Lu
3# @Email : yujielu1998@gmail.com
5# UPDATE
6# @Time : 2020/9/15, 2020/10/2
7# @Author : Yupeng Hou, Yujie Lu
8# @Email : houyupeng@ruc.edu.cn, yujielu1998@gmail.com
10r"""NARM
11################################################
13Reference:
14 Jing Li et al. "Neural Attentive Session-based Recommendation." in CIKM 2017.
16Reference code:
17 https://github.com/Wang-Shuo/Neural-Attentive-Session-Based-Recommendation-PyTorch
19"""
21import torch
22from torch import nn
23from torch.nn.init import constant_, xavier_normal_
25from hopwise.model.abstract_recommender import SequentialRecommender
26from hopwise.model.loss import BPRLoss
29class NARM(SequentialRecommender):
30 r"""NARM explores a hybrid encoder with an attention mechanism to model the user’s sequential behavior,
31 and capture the user’s main purpose in the current session.
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.n_layers = config["n_layers"]
42 self.dropout_probs = config["dropout_probs"]
43 self.device = config["device"]
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_probs[0])
48 self.gru = nn.GRU(
49 self.embedding_size,
50 self.hidden_size,
51 self.n_layers,
52 bias=False,
53 batch_first=True,
54 )
55 self.a_1 = nn.Linear(self.hidden_size, self.hidden_size, bias=False)
56 self.a_2 = nn.Linear(self.hidden_size, self.hidden_size, bias=False)
57 self.v_t = nn.Linear(self.hidden_size, 1, bias=False)
58 self.ct_dropout = nn.Dropout(self.dropout_probs[1])
59 self.b = nn.Linear(2 * self.hidden_size, self.embedding_size, bias=False)
60 self.loss_type = config["loss_type"]
61 if self.loss_type == "BPR":
62 self.loss_fct = BPRLoss()
63 elif self.loss_type == "CE":
64 self.loss_fct = nn.CrossEntropyLoss()
65 else:
66 raise NotImplementedError("Make sure 'loss_type' in ['BPR', 'CE']!")
68 # parameters initialization
69 self.apply(self._init_weights)
71 def _init_weights(self, module):
72 if isinstance(module, nn.Embedding):
73 xavier_normal_(module.weight.data)
74 elif isinstance(module, nn.Linear):
75 xavier_normal_(module.weight.data)
76 if module.bias is not None:
77 constant_(module.bias.data, 0)
79 def forward(self, item_seq, item_seq_len):
80 item_seq_emb = self.item_embedding(item_seq)
81 item_seq_emb_dropout = self.emb_dropout(item_seq_emb)
82 gru_out, _ = self.gru(item_seq_emb_dropout)
84 # fetch the last hidden state of last timestamp
85 c_global = ht = self.gather_indexes(gru_out, item_seq_len - 1)
86 # avoid the influence of padding
87 mask = item_seq.gt(0).unsqueeze(2).expand_as(gru_out)
88 q1 = self.a_1(gru_out)
89 q2 = self.a_2(ht)
90 q2_expand = q2.unsqueeze(1).expand_as(q1)
91 # calculate weighted factors α
92 alpha = self.v_t(mask * torch.sigmoid(q1 + q2_expand))
93 c_local = torch.sum(alpha.expand_as(gru_out) * gru_out, 1)
94 c_t = torch.cat([c_local, c_global], 1)
95 c_t = self.ct_dropout(c_t)
96 seq_output = self.b(c_t)
97 return seq_output
99 def calculate_loss(self, interaction):
100 item_seq = interaction[self.ITEM_SEQ]
101 item_seq_len = interaction[self.ITEM_SEQ_LEN]
102 seq_output = self.forward(item_seq, item_seq_len)
103 pos_items = interaction[self.POS_ITEM_ID]
104 if self.loss_type == "BPR":
105 neg_items = interaction[self.NEG_ITEM_ID]
106 pos_items_emb = self.item_embedding(pos_items)
107 neg_items_emb = self.item_embedding(neg_items)
108 pos_score = torch.sum(seq_output * pos_items_emb, dim=-1) # [B]
109 neg_score = torch.sum(seq_output * neg_items_emb, dim=-1) # [B]
110 loss = self.loss_fct(pos_score, neg_score)
111 return loss
112 else: # self.loss_type = 'CE'
113 test_item_emb = self.item_embedding.weight
114 logits = torch.matmul(seq_output, test_item_emb.transpose(0, 1))
115 loss = self.loss_fct(logits, pos_items)
116 return loss
118 def predict(self, interaction):
119 item_seq = interaction[self.ITEM_SEQ]
120 item_seq_len = interaction[self.ITEM_SEQ_LEN]
121 test_item = interaction[self.ITEM_ID]
122 seq_output = self.forward(item_seq, item_seq_len)
123 test_item_emb = self.item_embedding(test_item)
124 scores = torch.mul(seq_output, test_item_emb).sum(dim=1) # [B]
125 return scores
127 def full_sort_predict(self, interaction):
128 item_seq = interaction[self.ITEM_SEQ]
129 item_seq_len = interaction[self.ITEM_SEQ_LEN]
130 seq_output = self.forward(item_seq, item_seq_len)
131 test_items_emb = self.item_embedding.weight
132 scores = torch.matmul(seq_output, test_items_emb.transpose(0, 1))
133 return scores