Coverage for hopwise/model/sequential_recommender/core.py: 93%
113 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
1r"""CORE
2################################################
3Reference:
4 Yupeng Hou, Binbin Hu, Zhiqiang Zhang, Wayne Xin Zhao. "CORE: Simple and Effective Session-based Recommendation within Consistent Representation Space." in SIGIR 2022.
6 https://github.com/RUCAIBox/CORE
7""" # noqa: E501
9import numpy as np
10import torch
11import torch.nn.functional as F
12from torch import nn
14from hopwise.model.abstract_recommender import SequentialRecommender
15from hopwise.model.layers import TransformerEncoder
18class TransNet(nn.Module):
19 def __init__(self, config, dataset):
20 super().__init__()
22 self.n_layers = config["n_layers"]
23 self.n_heads = config["n_heads"]
24 self.hidden_size = config["embedding_size"]
25 self.inner_size = config["inner_size"]
26 self.hidden_dropout_prob = config["hidden_dropout_prob"]
27 self.attn_dropout_prob = config["attn_dropout_prob"]
28 self.hidden_act = config["hidden_act"]
29 self.layer_norm_eps = config["layer_norm_eps"]
30 self.initializer_range = config["initializer_range"]
32 self.position_embedding = nn.Embedding(
33 dataset.field2seqlen[config["ITEM_ID_FIELD"] + config["LIST_SUFFIX"]],
34 self.hidden_size,
35 )
36 self.trm_encoder = TransformerEncoder(
37 n_layers=self.n_layers,
38 n_heads=self.n_heads,
39 hidden_size=self.hidden_size,
40 inner_size=self.inner_size,
41 hidden_dropout_prob=self.hidden_dropout_prob,
42 attn_dropout_prob=self.attn_dropout_prob,
43 hidden_act=self.hidden_act,
44 layer_norm_eps=self.layer_norm_eps,
45 )
47 self.LayerNorm = nn.LayerNorm(self.hidden_size, eps=self.layer_norm_eps)
48 self.dropout = nn.Dropout(self.hidden_dropout_prob)
49 self.fn = nn.Linear(self.hidden_size, 1)
51 self.apply(self._init_weights)
53 def get_attention_mask(self, item_seq, bidirectional=False):
54 """Generate left-to-right uni-directional or bidirectional attention mask for multi-head attention."""
55 attention_mask = item_seq != 0
56 extended_attention_mask = attention_mask.unsqueeze(1).unsqueeze(2) # torch.bool
57 if not bidirectional:
58 extended_attention_mask = torch.tril(extended_attention_mask.expand((-1, -1, item_seq.size(-1), -1)))
59 extended_attention_mask = torch.where(extended_attention_mask, 0.0, -10000.0)
60 return extended_attention_mask
62 def forward(self, item_seq, item_emb):
63 mask = item_seq.gt(0)
65 position_ids = torch.arange(item_seq.size(1), dtype=torch.long, device=item_seq.device)
66 position_ids = position_ids.unsqueeze(0).expand_as(item_seq)
67 position_embedding = self.position_embedding(position_ids)
69 input_emb = item_emb + position_embedding
70 input_emb = self.LayerNorm(input_emb)
71 input_emb = self.dropout(input_emb)
73 extended_attention_mask = self.get_attention_mask(item_seq)
75 trm_output = self.trm_encoder(input_emb, extended_attention_mask, output_all_encoded_layers=True)
76 output = trm_output[-1]
78 alpha = self.fn(output).to(torch.double)
79 alpha = torch.where(mask.unsqueeze(-1), alpha, -9e15)
80 alpha = torch.softmax(alpha, dim=1, dtype=torch.float)
81 return alpha
83 def _init_weights(self, module):
84 """Initialize the weights"""
85 if isinstance(module, (nn.Linear, nn.Embedding)):
86 # Slightly different from the TF version which uses truncated_normal for initialization
87 # cf https://github.com/pytorch/pytorch/pull/5617
88 module.weight.data.normal_(mean=0.0, std=self.initializer_range)
89 elif isinstance(module, nn.LayerNorm):
90 module.bias.data.zero_()
91 module.weight.data.fill_(1.0)
92 if isinstance(module, nn.Linear) and module.bias is not None:
93 module.bias.data.zero_()
96class CORE(SequentialRecommender):
97 r"""CORE is a simple and effective framewor, which unifies the representation spac
98 for both the encoding and decoding processes in session-based recommendation.
99 """
101 def __init__(self, config, dataset):
102 super().__init__(config, dataset)
104 # load parameters info
105 self.embedding_size = config["embedding_size"]
106 self.loss_type = config["loss_type"]
108 self.dnn_type = config["dnn_type"]
109 self.sess_dropout = nn.Dropout(config["sess_dropout"])
110 self.item_dropout = nn.Dropout(config["item_dropout"])
111 self.temperature = config["temperature"]
113 # item embedding
114 self.item_embedding = nn.Embedding(self.n_items, self.embedding_size, padding_idx=0)
116 # DNN
117 if self.dnn_type == "trm":
118 self.net = TransNet(config, dataset)
119 elif self.dnn_type == "ave":
120 self.net = self.ave_net
121 else:
122 raise ValueError(f"dnn_type should be either trm or ave, but have [{self.dnn_type}].")
124 if self.loss_type == "CE":
125 self.loss_fct = nn.CrossEntropyLoss()
126 else:
127 raise NotImplementedError("Make sure 'loss_type' in ['CE']!")
129 # parameters initialization
130 self._reset_parameters()
132 def _reset_parameters(self):
133 stdv = 1.0 / np.sqrt(self.embedding_size)
134 for weight in self.parameters():
135 weight.data.uniform_(-stdv, stdv)
137 @staticmethod
138 def ave_net(item_seq, item_emb):
139 mask = item_seq.gt(0)
140 alpha = mask.to(torch.float) / mask.sum(dim=-1, keepdim=True)
141 return alpha.unsqueeze(-1)
143 def forward(self, item_seq):
144 x = self.item_embedding(item_seq)
145 x = self.sess_dropout(x)
146 # Representation-Consistent Encoder (RCE)
147 alpha = self.net(item_seq, x)
148 seq_output = torch.sum(alpha * x, dim=1)
149 seq_output = F.normalize(seq_output, dim=-1)
150 return seq_output
152 def calculate_loss(self, interaction):
153 item_seq = interaction[self.ITEM_SEQ]
154 seq_output = self.forward(item_seq)
155 pos_items = interaction[self.POS_ITEM_ID]
157 all_item_emb = self.item_embedding.weight
158 all_item_emb = self.item_dropout(all_item_emb)
159 # Robust Distance Measuring (RDM)
160 all_item_emb = F.normalize(all_item_emb, dim=-1)
161 logits = torch.matmul(seq_output, all_item_emb.transpose(0, 1)) / self.temperature
162 loss = self.loss_fct(logits, pos_items)
163 return loss
165 def predict(self, interaction):
166 item_seq = interaction[self.ITEM_SEQ]
167 test_item = interaction[self.ITEM_ID]
168 seq_output = self.forward(item_seq)
169 test_item_emb = self.item_embedding(test_item)
170 scores = torch.mul(seq_output, test_item_emb).sum(dim=1) / self.temperature
171 return scores
173 def full_sort_predict(self, interaction):
174 item_seq = interaction[self.ITEM_SEQ]
175 seq_output = self.forward(item_seq)
176 test_item_emb = self.item_embedding.weight
177 # no dropout for evaluation
178 test_item_emb = F.normalize(test_item_emb, dim=-1)
179 scores = torch.matmul(seq_output, test_item_emb.transpose(0, 1)) / self.temperature
180 return scores