Coverage for hopwise/model/sequential_recommender/sine.py: 94%
129 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 : 2021/11/23 11:10
2# @Author : Jingqi Gao
3# @Email : jgaoaz@connect.ust.hk
5r"""SINE
6################################################
8Reference:
9 Qiaoyu Tan et al. "Sparse-Interest Network for Sequential Recommendation." in WSDM 2021.
11"""
13import numpy as np
14import torch
15import torch.nn.functional as F
16from torch import nn
17from torch.nn.init import xavier_normal_
19from hopwise.model.abstract_recommender import SequentialRecommender
20from hopwise.model.loss import BPRLoss
21from hopwise.utils import InputType
23torch.autograd.set_detect_anomaly(True)
26class SINE(SequentialRecommender):
27 input_type = InputType.PAIRWISE
29 def __init__(self, config, dataset):
30 super().__init__(config, dataset)
32 # load dataset info
33 self.n_users = dataset.user_num
34 self.n_items = dataset.item_num
36 # load parameters info
37 self.device = config["device"]
38 self.embedding_size = config["embedding_size"]
39 self.loss_type = config["loss_type"]
40 self.layer_norm_eps = config["layer_norm_eps"]
42 if self.loss_type == "BPR":
43 self.loss_fct = BPRLoss()
44 elif self.loss_type == "CE":
45 self.loss_fct = nn.CrossEntropyLoss()
46 elif self.loss_type == "NLL":
47 self.loss_fct = nn.NLLLoss()
48 else:
49 raise NotImplementedError("Make sure 'loss_type' in ['BPR', 'CE', 'NLL']!")
51 self.D = config["embedding_size"]
52 self.L = config["prototype_size"] # 500 for movie-len dataset
53 self.k = config["interest_size"] # 4 for movie-len dataset
54 self.tau = config["tau_ratio"] # 0.1 in paper
55 self.reg_loss_ratio = config["reg_loss_ratio"] # 0.1 in paper
57 self.initializer_range = 0.01
59 self.w1 = self._init_weight((self.D, self.D))
60 self.w2 = self._init_weight(self.D)
61 self.w3 = self._init_weight((self.D, self.D))
62 self.w4 = self._init_weight(self.D)
64 self.C = nn.Embedding(self.L, self.D)
66 self.w_k_1 = self._init_weight((self.k, self.D, self.D))
67 self.w_k_2 = self._init_weight((self.k, self.D))
68 self.item_embedding = nn.Embedding(self.n_items, self.D, padding_idx=0)
69 self.ln2 = nn.LayerNorm(self.embedding_size, eps=self.layer_norm_eps)
70 self.ln4 = nn.LayerNorm(self.embedding_size, eps=self.layer_norm_eps)
72 # parameters initialization
73 self.apply(self._init_weights)
75 def _init_weight(self, shape):
76 mat = torch.FloatTensor(np.random.normal(0, self.initializer_range, shape))
77 return nn.Parameter(mat, requires_grad=True)
79 def _init_weights(self, module):
80 if isinstance(module, nn.Embedding):
81 xavier_normal_(module.weight)
82 elif isinstance(module, nn.LayerNorm):
83 module.bias.data.zero_()
84 module.weight.data.fill_(1.0)
86 def calculate_loss(self, interaction):
87 item_seq = interaction[self.ITEM_SEQ]
88 item_seq_len = interaction[self.ITEM_SEQ_LEN]
89 seq_output = self.forward(item_seq, item_seq_len)
90 pos_items = interaction[self.POS_ITEM_ID]
92 if self.loss_type == "BPR":
93 neg_items = interaction[self.NEG_ITEM_ID]
94 pos_items_emb = self.item_embedding(pos_items)
95 neg_items_emb = self.item_embedding(neg_items)
96 pos_score = torch.sum(seq_output * pos_items_emb, dim=-1) # [B]
97 neg_score = torch.sum(seq_output * neg_items_emb, dim=-1) # [B]
98 loss = self.loss_fct(pos_score, neg_score)
99 return loss
100 elif self.loss_type == "CE":
101 test_item_emb = self.item_embedding.weight
102 logits = torch.matmul(seq_output, test_item_emb.transpose(0, 1))
103 loss = self.loss_fct(logits, pos_items)
104 return loss
105 else:
106 test_item_emb = self.item_embedding.weight
107 logits = torch.matmul(seq_output, test_item_emb.transpose(0, 1))
108 logits = F.log_softmax(logits, dim=1)
109 loss = self.loss_fct(logits, pos_items)
110 return loss + self.calculate_reg_loss() * self.reg_loss_ratio
112 def calculate_reg_loss(self):
113 C_mean = torch.mean(self.C.weight, dim=1, keepdim=True)
114 C_reg = self.C.weight - C_mean
115 C_reg = C_reg.matmul(C_reg.T) / self.D
116 return (torch.norm(C_reg) ** 2 - torch.norm(torch.diag(C_reg)) ** 2) / 2
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 forward(self, item_seq, item_seq_len):
128 x_u = self.item_embedding(item_seq).to(self.device) # [B, N, D]
130 # concept activation
131 # sort by inner product
132 x = torch.matmul(x_u, self.w1)
133 x = torch.tanh(x)
134 x = torch.matmul(x, self.w2)
135 a = F.softmax(x, dim=1)
136 z_u = torch.matmul(a.unsqueeze(2).transpose(1, 2), x_u).transpose(1, 2)
137 s_u = torch.matmul(self.C.weight, z_u)
138 s_u = s_u.squeeze(2)
139 idx = s_u.argsort(1)[:, -self.k :]
140 s_u_idx = s_u.sort(1)[0][:, -self.k :]
141 c_u = self.C(idx)
142 sigs = torch.sigmoid(s_u_idx.unsqueeze(2).repeat(1, 1, self.embedding_size))
143 C_u = c_u.mul(sigs)
145 # intention assignment
146 # use matrix multiplication instead of cos()
147 w3_x_u_norm = F.normalize(x_u.matmul(self.w3), p=2, dim=2)
148 C_u_norm = self.ln2(C_u)
149 P_k_t = torch.bmm(w3_x_u_norm, C_u_norm.transpose(1, 2))
150 P_k_t_b = F.softmax(P_k_t, dim=2)
151 P_k_t_b_t = P_k_t_b.transpose(1, 2)
153 # attention weighting
154 a_k = x_u.unsqueeze(1).repeat(1, self.k, 1, 1).matmul(self.w_k_1)
155 P_t_k = F.softmax(
156 torch.tanh(a_k).matmul(self.w_k_2.reshape(self.k, self.embedding_size, 1)).squeeze(3),
157 dim=2,
158 )
160 # interest embedding generation
161 mul_p = P_k_t_b_t.mul(P_t_k)
162 x_u_re = x_u.unsqueeze(1).repeat(1, self.k, 1, 1)
163 mul_p_re = mul_p.unsqueeze(3)
164 delta_k = x_u_re.mul(mul_p_re).sum(2)
165 delta_k = F.normalize(delta_k, p=2, dim=2)
167 # prototype sequence
168 x_u_bar = P_k_t_b.matmul(C_u)
169 C_apt = F.softmax(torch.tanh(x_u_bar.matmul(self.w3)).matmul(self.w4), dim=1)
170 C_apt = C_apt.reshape(-1, 1, self.max_seq_length).matmul(x_u_bar)
171 C_apt = self.ln4(C_apt)
173 # aggregation weight
174 e_k = delta_k.bmm(C_apt.reshape(-1, self.embedding_size, 1)) / self.tau
175 e_k_u = F.softmax(e_k.squeeze(2), dim=1)
176 v_u = e_k_u.unsqueeze(2).mul(delta_k).sum(dim=1)
178 return v_u
180 def full_sort_predict(self, interaction):
181 item_seq = interaction[self.ITEM_SEQ]
182 item_seq_len = interaction[self.ITEM_SEQ_LEN]
183 seq_output = self.forward(item_seq, item_seq_len)
184 test_items_emb = self.item_embedding.weight
185 scores = torch.matmul(seq_output, test_items_emb.transpose(0, 1)) # [B, n_items]
186 return scores