Coverage for hopwise/model/sequential_recommender/sasrecf.py: 84%
104 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/18 11:32
2# @Author : Hui Wang
3# @Email : hui.wang@ruc.edu.cn
5r"""SASRecF
6################################################
7"""
9import torch
10from torch import nn
12from hopwise.model.abstract_recommender import SequentialRecommender
13from hopwise.model.layers import FeatureSeqEmbLayer, TransformerEncoder
14from hopwise.model.loss import BPRLoss
15from hopwise.utils import FeatureType
18class SASRecF(SequentialRecommender):
19 """This is an extension of SASRec, which concatenates item representations and item attribute representations
20 as the input to the model.
21 """
23 def __init__(self, config, dataset):
24 super().__init__(config, dataset)
26 # load parameters info
27 self.n_layers = config["n_layers"]
28 self.n_heads = config["n_heads"]
29 self.hidden_size = config["hidden_size"] # same as embedding_size
30 self.inner_size = config["inner_size"] # the dimensionality in feed-forward layer
31 self.hidden_dropout_prob = config["hidden_dropout_prob"]
32 self.attn_dropout_prob = config["attn_dropout_prob"]
33 self.hidden_act = config["hidden_act"]
34 self.layer_norm_eps = config["layer_norm_eps"]
36 self.selected_features = config["selected_features"]
37 self.pooling_mode = config["pooling_mode"]
38 self.device = config["device"]
39 self.num_feature_field = sum(
40 (1 if dataset.field2type[field] != FeatureType.FLOAT_SEQ else dataset.num(field))
41 for field in config["selected_features"]
42 )
44 self.initializer_range = config["initializer_range"]
45 self.loss_type = config["loss_type"]
47 # define layers and loss
48 self.item_embedding = nn.Embedding(self.n_items, self.hidden_size, padding_idx=0)
49 self.position_embedding = nn.Embedding(self.max_seq_length, self.hidden_size)
50 self.feature_embed_layer = FeatureSeqEmbLayer(
51 dataset,
52 self.hidden_size,
53 self.selected_features,
54 self.pooling_mode,
55 self.device,
56 )
58 self.trm_encoder = TransformerEncoder(
59 n_layers=self.n_layers,
60 n_heads=self.n_heads,
61 hidden_size=self.hidden_size,
62 inner_size=self.inner_size,
63 hidden_dropout_prob=self.hidden_dropout_prob,
64 attn_dropout_prob=self.attn_dropout_prob,
65 hidden_act=self.hidden_act,
66 layer_norm_eps=self.layer_norm_eps,
67 )
69 self.concat_layer = nn.Linear(self.hidden_size * (1 + self.num_feature_field), self.hidden_size)
71 self.LayerNorm = nn.LayerNorm(self.hidden_size, eps=self.layer_norm_eps)
72 self.dropout = nn.Dropout(self.hidden_dropout_prob)
74 if self.loss_type == "BPR":
75 self.loss_fct = BPRLoss()
76 elif self.loss_type == "CE":
77 self.loss_fct = nn.CrossEntropyLoss()
78 else:
79 raise NotImplementedError("Make sure 'loss_type' in ['BPR', 'CE']!")
81 # parameters initialization
82 self.apply(self._init_weights)
83 self.other_parameter_name = ["feature_embed_layer"]
85 def _init_weights(self, module):
86 """Initialize the weights"""
87 if isinstance(module, (nn.Linear, nn.Embedding)):
88 # Slightly different from the TF version which uses truncated_normal for initialization
89 # cf https://github.com/pytorch/pytorch/pull/5617
90 module.weight.data.normal_(mean=0.0, std=self.initializer_range)
91 elif isinstance(module, nn.LayerNorm):
92 module.bias.data.zero_()
93 module.weight.data.fill_(1.0)
94 if isinstance(module, nn.Linear) and module.bias is not None:
95 module.bias.data.zero_()
97 def forward(self, item_seq, item_seq_len):
98 item_emb = self.item_embedding(item_seq)
100 # position embedding
101 position_ids = torch.arange(item_seq.size(1), dtype=torch.long, device=item_seq.device)
102 position_ids = position_ids.unsqueeze(0).expand_as(item_seq)
103 position_embedding = self.position_embedding(position_ids)
105 sparse_embedding, dense_embedding = self.feature_embed_layer(None, item_seq)
106 sparse_embedding = sparse_embedding["item"]
107 dense_embedding = dense_embedding["item"]
108 # concat the sparse embedding and float embedding
109 feature_table = []
110 if sparse_embedding is not None:
111 feature_table.append(sparse_embedding)
112 if dense_embedding is not None:
113 feature_table.append(dense_embedding)
115 feature_table = torch.cat(feature_table, dim=-2)
116 table_shape = feature_table.shape
117 feat_num, embedding_size = table_shape[-2], table_shape[-1]
118 feature_emb = feature_table.view(table_shape[:-2] + (feat_num * embedding_size,))
119 input_concat = torch.cat((item_emb, feature_emb), -1) # [B 1+field_num*H]
121 input_emb = self.concat_layer(input_concat)
122 input_emb = input_emb + position_embedding
123 input_emb = self.LayerNorm(input_emb)
124 input_emb = self.dropout(input_emb)
126 extended_attention_mask = self.get_attention_mask(item_seq)
127 trm_output = self.trm_encoder(input_emb, extended_attention_mask, output_all_encoded_layers=True)
128 output = trm_output[-1]
129 seq_output = self.gather_indexes(output, item_seq_len - 1)
130 return seq_output # [B H]
132 def calculate_loss(self, interaction):
133 item_seq = interaction[self.ITEM_SEQ]
134 item_seq_len = interaction[self.ITEM_SEQ_LEN]
135 seq_output = self.forward(item_seq, item_seq_len)
136 pos_items = interaction[self.POS_ITEM_ID]
137 if self.loss_type == "BPR":
138 neg_items = interaction[self.NEG_ITEM_ID]
139 pos_items_emb = self.item_embedding(pos_items)
140 neg_items_emb = self.item_embedding(neg_items)
141 pos_score = torch.sum(seq_output * pos_items_emb, dim=-1) # [B]
142 neg_score = torch.sum(seq_output * neg_items_emb, dim=-1) # [B]
143 loss = self.loss_fct(pos_score, neg_score)
144 return loss
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
151 def predict(self, interaction):
152 item_seq = interaction[self.ITEM_SEQ]
153 item_seq_len = interaction[self.ITEM_SEQ_LEN]
154 test_item = interaction[self.ITEM_ID]
155 seq_output = self.forward(item_seq, item_seq_len)
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 item_seq = interaction[self.ITEM_SEQ]
162 item_seq_len = interaction[self.ITEM_SEQ_LEN]
163 seq_output = self.forward(item_seq, item_seq_len)
164 test_items_emb = self.item_embedding.weight
165 scores = torch.matmul(seq_output, test_items_emb.transpose(0, 1)) # [B, item_num]
166 return scores