Coverage for hopwise/model/sequential_recommender/lightsans.py: 90%
84 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/05/01
2# @Author : Xinyan Fan
3# @Email : xinyan.fan@ruc.edu.cn
5"""LightSANs
6################################################
7Reference:
8 Xin-Yan Fan et al. "Lighter and Better: Low-Rank Decomposed Self-Attention Networks for Next-Item Recommendation." in SIGIR 2021.
9Reference:
10 https://github.com/BELIEVEfxy/LightSANs
11""" # noqa: E501
13import torch
14from torch import nn
16from hopwise.model.abstract_recommender import SequentialRecommender
17from hopwise.model.layers import LightTransformerEncoder
18from hopwise.model.loss import BPRLoss
21class LightSANs(SequentialRecommender):
22 def __init__(self, config, dataset):
23 super().__init__(config, dataset)
25 # load parameters info
26 self.n_layers = config["n_layers"]
27 self.n_heads = config["n_heads"]
28 self.k_interests = config["k_interests"]
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.initializer_range = config["initializer_range"]
37 self.loss_type = config["loss_type"]
39 self.seq_len = self.max_seq_length
40 # define layers and loss
41 self.item_embedding = nn.Embedding(self.n_items, self.hidden_size, padding_idx=0)
42 self.position_embedding = nn.Embedding(self.max_seq_length, self.hidden_size)
43 self.trm_encoder = LightTransformerEncoder(
44 n_layers=self.n_layers,
45 n_heads=self.n_heads,
46 k_interests=self.k_interests,
47 hidden_size=self.hidden_size,
48 seq_len=self.seq_len,
49 inner_size=self.inner_size,
50 hidden_dropout_prob=self.hidden_dropout_prob,
51 attn_dropout_prob=self.attn_dropout_prob,
52 hidden_act=self.hidden_act,
53 layer_norm_eps=self.layer_norm_eps,
54 )
56 self.LayerNorm = nn.LayerNorm(self.hidden_size, eps=self.layer_norm_eps)
57 self.dropout = nn.Dropout(self.hidden_dropout_prob)
59 if self.loss_type == "BPR":
60 self.loss_fct = BPRLoss()
61 elif self.loss_type == "CE":
62 self.loss_fct = nn.CrossEntropyLoss()
63 else:
64 raise NotImplementedError("Make sure 'loss_type' in ['BPR', 'CE']!")
66 # parameters initialization
67 self.apply(self._init_weights)
69 def _init_weights(self, module):
70 """Initialize the weights"""
71 if isinstance(module, (nn.Linear, nn.Embedding)):
72 module.weight.data.normal_(mean=0.0, std=self.initializer_range)
73 elif isinstance(module, nn.LayerNorm):
74 module.bias.data.zero_()
75 module.weight.data.fill_(1.0)
76 if isinstance(module, nn.Linear) and module.bias is not None:
77 module.bias.data.zero_()
79 def embedding_layer(self, item_seq):
80 position_ids = torch.arange(item_seq.size(1), dtype=torch.long, device=item_seq.device)
81 position_embedding = self.position_embedding(position_ids)
82 item_emb = self.item_embedding(item_seq)
83 return item_emb, position_embedding
85 def forward(self, item_seq, item_seq_len):
86 item_emb, position_embedding = self.embedding_layer(item_seq)
87 item_emb = self.LayerNorm(item_emb)
88 item_emb = self.dropout(item_emb)
90 trm_output = self.trm_encoder(item_emb, position_embedding, output_all_encoded_layers=True)
91 output = trm_output[-1]
92 output = self.gather_indexes(output, item_seq_len - 1)
93 return output # [B H]
95 def calculate_loss(self, interaction):
96 item_seq = interaction[self.ITEM_SEQ]
97 item_seq_len = interaction[self.ITEM_SEQ_LEN]
98 seq_output = self.forward(item_seq, item_seq_len)
99 pos_items = interaction[self.POS_ITEM_ID]
100 if self.loss_type == "BPR":
101 neg_items = interaction[self.NEG_ITEM_ID]
102 pos_items_emb = self.item_embedding(pos_items)
103 neg_items_emb = self.item_embedding(neg_items)
104 pos_score = torch.sum(seq_output * pos_items_emb, dim=-1) # [B]
105 neg_score = torch.sum(seq_output * neg_items_emb, dim=-1) # [B]
106 loss = self.loss_fct(pos_score, neg_score)
107 return loss
108 else: # self.loss_type = 'CE'
109 test_item_emb = self.item_embedding.weight
110 logits = torch.matmul(seq_output, test_item_emb.transpose(0, 1))
111 loss = self.loss_fct(logits, pos_items)
112 return loss
114 def predict(self, interaction):
115 item_seq = interaction[self.ITEM_SEQ]
116 item_seq_len = interaction[self.ITEM_SEQ_LEN]
117 test_item = interaction[self.ITEM_ID]
119 seq_output = self.forward(item_seq, item_seq_len)
120 test_item_emb = self.item_embedding(test_item)
121 scores = torch.mul(seq_output, test_item_emb).sum(dim=1) # [B]
122 return scores
124 def full_sort_predict(self, interaction):
125 item_seq = interaction[self.ITEM_SEQ]
126 item_seq_len = interaction[self.ITEM_SEQ_LEN]
127 seq_output = self.forward(item_seq, item_seq_len)
128 test_items_emb = self.item_embedding.weight
129 scores = torch.matmul(seq_output, test_items_emb.transpose(0, 1)) # [B n_items]
130 return scores