Coverage for hopwise/model/sequential_recommender/shan.py: 91%
107 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/11/20 22:33
2# @Author : Shao Weiqi
3# @Reviewer : Lin Kun
4# @Email : shaoweiqi@ruc.edu.cn
6r"""SHAN
7################################################
9Reference:
10 Ying, H et al. "Sequential Recommender System based on Hierarchical Attention Network."in IJCAI 2018
13"""
15import numpy as np
16import torch
17from torch import nn
18from torch.nn.init import normal_, uniform_
20from hopwise.model.abstract_recommender import SequentialRecommender
21from hopwise.model.loss import BPRLoss
24class SHAN(SequentialRecommender):
25 r"""SHAN exploit the Hierarchical Attention Network to get the long-short term preference
26 first get the long term purpose and then fuse the long-term with recent items to get long-short term purpose
28 """
30 def __init__(self, config, dataset):
31 super().__init__(config, dataset)
33 # load the dataset information
34 self.n_users = dataset.num(self.USER_ID)
35 self.device = config["device"]
36 self.INVERSE_ITEM_SEQ = config["INVERSE_ITEM_SEQ"]
38 # load the parameter information
39 self.embedding_size = config["embedding_size"]
40 self.short_item_length = config["short_item_length"] # the length of the short session items
41 assert self.short_item_length <= self.max_seq_length, "short_item_length can't longer than the max_seq_length"
42 self.reg_weight = config["reg_weight"]
44 # define layers and loss
45 self.item_embedding = nn.Embedding(self.n_items, self.embedding_size, padding_idx=0)
46 self.user_embedding = nn.Embedding(self.n_users, self.embedding_size)
48 self.long_w = nn.Linear(self.embedding_size, self.embedding_size)
49 self.long_b = nn.Parameter(
50 uniform_(
51 tensor=torch.zeros(self.embedding_size),
52 a=-np.sqrt(3 / self.embedding_size),
53 b=np.sqrt(3 / self.embedding_size),
54 ),
55 requires_grad=True,
56 )
57 self.long_short_w = nn.Linear(self.embedding_size, self.embedding_size)
58 self.long_short_b = nn.Parameter(
59 uniform_(
60 tensor=torch.zeros(self.embedding_size),
61 a=-np.sqrt(3 / self.embedding_size),
62 b=np.sqrt(3 / self.embedding_size),
63 ),
64 requires_grad=True,
65 )
67 self.relu = nn.ReLU()
69 self.loss_type = config["loss_type"]
70 if self.loss_type == "BPR":
71 self.loss_fct = BPRLoss()
72 elif self.loss_type == "CE":
73 self.loss_fct = nn.CrossEntropyLoss()
74 else:
75 raise NotImplementedError("Make sure 'loss_type' in ['BPR', 'CE']!")
77 # init the parameter of the model
78 self.apply(self.init_weights)
80 def reg_loss(self, user_embedding, item_embedding):
81 reg_1, reg_2 = self.reg_weight
82 loss_1 = reg_1 * torch.norm(self.long_w.weight, p=2) + reg_1 * torch.norm(self.long_short_w.weight, p=2)
83 loss_2 = reg_2 * torch.norm(user_embedding, p=2) + reg_2 * torch.norm(item_embedding, p=2)
85 return loss_1 + loss_2
87 def init_weights(self, module):
88 if isinstance(module, nn.Embedding):
89 normal_(module.weight.data, 0.0, 0.01)
90 elif isinstance(module, nn.Linear):
91 uniform_(
92 module.weight.data,
93 -np.sqrt(3 / self.embedding_size),
94 np.sqrt(3 / self.embedding_size),
95 )
96 elif isinstance(module, nn.Parameter):
97 uniform_(
98 module.data,
99 -np.sqrt(3 / self.embedding_size),
100 np.sqrt(3 / self.embedding_size),
101 )
102 print(module.data)
104 def forward(self, seq_item, user):
105 seq_item_embedding = self.item_embedding(seq_item)
106 user_embedding = self.user_embedding(user)
108 # get the mask
109 mask = seq_item.data.eq(0)
110 long_term_attention_based_pooling_layer = self.long_term_attention_based_pooling_layer(
111 seq_item_embedding, user_embedding, mask
112 )
113 # batch_size * 1 * embedding_size
115 short_item_embedding = seq_item_embedding[:, -self.short_item_length :, :]
116 mask_long_short = mask[:, -self.short_item_length :]
117 batch_size = mask_long_short.size(0)
118 x = torch.zeros(size=(batch_size, 1)).eq(1).to(self.device)
119 mask_long_short = torch.cat([x, mask_long_short], dim=1)
120 # batch_size * short_item_length * embedding_size
121 long_short_item_embedding = torch.cat([long_term_attention_based_pooling_layer, short_item_embedding], dim=1)
122 # batch_size * 1_plus_short_item_length * embedding_size
124 long_short_item_embedding = self.long_and_short_term_attention_based_pooling_layer(
125 long_short_item_embedding, user_embedding, mask_long_short
126 )
127 # batch_size * embedding_size
129 return long_short_item_embedding
131 def calculate_loss(self, interaction):
132 inverse_seq_item = interaction[self.INVERSE_ITEM_SEQ]
133 user = interaction[self.USER_ID]
134 user_embedding = self.user_embedding(user)
135 seq_output = self.forward(inverse_seq_item, user)
136 pos_items = interaction[self.POS_ITEM_ID]
137 pos_items_emb = self.item_embedding(pos_items)
138 if self.loss_type == "BPR":
139 neg_items = interaction[self.NEG_ITEM_ID]
140 neg_items_emb = self.item_embedding(neg_items)
141 pos_score = torch.sum(seq_output * pos_items_emb, dim=-1)
142 neg_score = torch.sum(seq_output * neg_items_emb, dim=-1)
143 loss = self.loss_fct(pos_score, neg_score)
144 return loss + self.reg_loss(user_embedding, pos_items_emb)
145 else: # self.loss_type = 'CE'
146 test_item_emb = self.item_embedding.weight
147 logits = torch.matmul(seq_output, test_item_emb.transpose(0, 1))
148 loss = self.loss_fct(logits, pos_items)
149 return loss + self.reg_loss(user_embedding, pos_items_emb)
151 def predict(self, interaction):
152 inverse_item_seq = interaction[self.INVERSE_ITEM_SEQ]
153 test_item = interaction[self.ITEM_ID]
154 user = interaction[self.USER_ID]
155 seq_output = self.forward(inverse_item_seq, user)
156 test_item_emb = self.item_embedding(test_item)
157 scores = torch.mul(seq_output, test_item_emb).sum(dim=1)
158 return scores
160 def full_sort_predict(self, interaction):
161 inverse_item_seq = interaction[self.ITEM_SEQ]
162 user = interaction[self.USER_ID]
163 seq_output = self.forward(inverse_item_seq, user)
164 test_items_emb = self.item_embedding.weight
165 scores = torch.matmul(seq_output, test_items_emb.transpose(0, 1))
166 return scores
168 def long_and_short_term_attention_based_pooling_layer(self, long_short_item_embedding, user_embedding, mask=None):
169 """Fusing the long term purpose with the short-term preference"""
170 long_short_item_embedding_value = long_short_item_embedding
172 long_short_item_embedding = self.relu(self.long_short_w(long_short_item_embedding) + self.long_short_b)
173 long_short_item_embedding = torch.matmul(long_short_item_embedding, user_embedding.unsqueeze(2)).squeeze(-1)
174 # batch_size * seq_len
175 if mask is not None:
176 long_short_item_embedding.masked_fill_(mask, -1e9)
177 long_short_item_embedding = nn.Softmax(dim=-1)(long_short_item_embedding)
178 long_short_item_embedding = torch.mul(
179 long_short_item_embedding_value, long_short_item_embedding.unsqueeze(2)
180 ).sum(dim=1)
182 return long_short_item_embedding
184 def long_term_attention_based_pooling_layer(self, seq_item_embedding, user_embedding, mask=None):
185 """Get the long term purpose of user"""
186 seq_item_embedding_value = seq_item_embedding
188 seq_item_embedding = self.relu(self.long_w(seq_item_embedding) + self.long_b)
189 user_item_embedding = torch.matmul(seq_item_embedding, user_embedding.unsqueeze(2)).squeeze(-1)
190 # batch_size * seq_len
191 if mask is not None:
192 user_item_embedding.masked_fill_(mask, -1e9)
193 user_item_embedding = nn.Softmax(dim=1)(user_item_embedding)
194 user_item_embedding = torch.mul(seq_item_embedding_value, user_item_embedding.unsqueeze(2)).sum(
195 dim=1, keepdim=True
196 )
197 # batch_size * 1 * embedding_size
199 return user_item_embedding