Coverage for hopwise/model/sequential_recommender/stamp.py: 91%
88 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/9/8 19:24
2# @Author : Yujie Lu
3# @Email : yujielu1998@gmail.com
5# UPDATE
6# @Time : 2020/10/2
7# @Author : Yujie Lu
8# @Email : yujielu1998@gmail.com
10r"""STAMP
11################################################
13Reference:
14 Qiao Liu et al. "STAMP: Short-Term Attention/Memory Priority Model for Session-based Recommendation." in KDD 2018.
16"""
18import torch
19from torch import nn
20from torch.nn.init import normal_
22from hopwise.model.abstract_recommender import SequentialRecommender
23from hopwise.model.loss import BPRLoss
26class STAMP(SequentialRecommender):
27 r"""STAMP is capable of capturing users’ general interests from the long-term memory of a session context,
28 whilst taking into account users’ current interests from the short-term memory of the last-clicks.
31 Note:
32 According to the test results, we made a little modification to the score function mentioned in the paper,
33 and did not use the final sigmoid activation function.
35 """
37 def __init__(self, config, dataset):
38 super().__init__(config, dataset)
40 # load parameters info
41 self.embedding_size = config["embedding_size"]
43 # define layers and loss
44 self.item_embedding = nn.Embedding(self.n_items, self.embedding_size, padding_idx=0)
45 self.w1 = nn.Linear(self.embedding_size, self.embedding_size, bias=False)
46 self.w2 = nn.Linear(self.embedding_size, self.embedding_size, bias=False)
47 self.w3 = nn.Linear(self.embedding_size, self.embedding_size, bias=False)
48 self.w0 = nn.Linear(self.embedding_size, 1, bias=False)
49 self.b_a = nn.Parameter(torch.zeros(self.embedding_size), requires_grad=True)
50 self.mlp_a = nn.Linear(self.embedding_size, self.embedding_size, bias=True)
51 self.mlp_b = nn.Linear(self.embedding_size, self.embedding_size, bias=True)
52 self.sigmoid = nn.Sigmoid()
53 self.tanh = nn.Tanh()
54 self.loss_type = config["loss_type"]
55 if self.loss_type == "BPR":
56 self.loss_fct = BPRLoss()
57 elif self.loss_type == "CE":
58 self.loss_fct = nn.CrossEntropyLoss()
59 else:
60 raise NotImplementedError("Make sure 'loss_type' in ['BPR', 'CE']!")
62 # # parameters initialization
63 self.apply(self._init_weights)
65 def _init_weights(self, module):
66 if isinstance(module, nn.Embedding):
67 normal_(module.weight.data, 0, 0.002)
68 elif isinstance(module, nn.Linear):
69 normal_(module.weight.data, 0, 0.05)
70 if module.bias is not None:
71 module.bias.data.fill_(0.0)
73 def forward(self, item_seq, item_seq_len):
74 item_seq_emb = self.item_embedding(item_seq)
75 last_inputs = self.gather_indexes(item_seq_emb, item_seq_len - 1)
76 org_memory = item_seq_emb
77 ms = torch.div(torch.sum(org_memory, dim=1), item_seq_len.unsqueeze(1).float())
78 alpha = self.count_alpha(org_memory, last_inputs, ms)
79 vec = torch.matmul(alpha.unsqueeze(1), org_memory)
80 ma = vec.squeeze(1) + ms
81 hs = self.tanh(self.mlp_a(ma))
82 ht = self.tanh(self.mlp_b(last_inputs))
83 seq_output = hs * ht
84 return seq_output
86 def count_alpha(self, context, aspect, output):
87 r"""This is a function that count the attention weights
89 Args:
90 context(torch.FloatTensor): Item list embedding matrix, shape of [batch_size, time_steps, emb]
91 aspect(torch.FloatTensor): The embedding matrix of the last click item, shape of [batch_size, emb]
92 output(torch.FloatTensor): The average of the context, shape of [batch_size, emb]
94 Returns:
95 torch.Tensor:attention weights, shape of [batch_size, time_steps]
96 """
97 timesteps = context.size(1)
98 aspect_3dim = aspect.repeat(1, timesteps).view(-1, timesteps, self.embedding_size)
99 output_3dim = output.repeat(1, timesteps).view(-1, timesteps, self.embedding_size)
100 res_ctx = self.w1(context)
101 res_asp = self.w2(aspect_3dim)
102 res_output = self.w3(output_3dim)
103 res_sum = res_ctx + res_asp + res_output + self.b_a
104 res_act = self.w0(self.sigmoid(res_sum))
105 alpha = res_act.squeeze(2)
106 return alpha
108 def calculate_loss(self, interaction):
109 item_seq = interaction[self.ITEM_SEQ]
110 item_seq_len = interaction[self.ITEM_SEQ_LEN]
111 seq_output = self.forward(item_seq, item_seq_len)
112 pos_items = interaction[self.POS_ITEM_ID]
113 if self.loss_type == "BPR":
114 neg_items = interaction[self.NEG_ITEM_ID]
115 pos_items_emb = self.item_embedding(pos_items)
116 neg_items_emb = self.item_embedding(neg_items)
117 pos_score = torch.sum(seq_output * pos_items_emb, dim=-1) # [B]
118 neg_score = torch.sum(seq_output * neg_items_emb, dim=-1) # [B]
119 loss = self.loss_fct(pos_score, neg_score)
120 return loss
121 else: # self.loss_type = 'CE'
122 test_item_emb = self.item_embedding.weight
123 logits = torch.matmul(seq_output, test_item_emb.transpose(0, 1))
124 loss = self.loss_fct(logits, pos_items)
125 return loss
127 def predict(self, interaction):
128 item_seq = interaction[self.ITEM_SEQ]
129 item_seq_len = interaction[self.ITEM_SEQ_LEN]
130 test_item = interaction[self.ITEM_ID]
131 seq_output = self.forward(item_seq, item_seq_len)
132 test_item_emb = self.item_embedding(test_item)
133 scores = torch.mul(seq_output, test_item_emb).sum(dim=1) # [B]
134 return scores
136 def full_sort_predict(self, interaction):
137 item_seq = interaction[self.ITEM_SEQ]
138 item_seq_len = interaction[self.ITEM_SEQ_LEN]
139 seq_output = self.forward(item_seq, item_seq_len)
140 test_items_emb = self.item_embedding.weight
141 scores = torch.matmul(seq_output, test_items_emb.transpose(0, 1))
142 return scores