Coverage for hopwise/model/sequential_recommender/fossil.py: 83%
95 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/21 20:00
2# @Author : Shao Weiqi
3# @Reviewer : Lin Kun
4# @Email : shaoweiqi@ruc.edu.cn
6r"""FOSSIL
7################################################
9Reference:
10 Ruining He et al. "Fusing Similarity Models with Markov Chains for Sparse Sequential Recommendation." in ICDM 2016.
13"""
15import torch
16from torch import nn
17from torch.nn.init import xavier_normal_
19from hopwise.model.abstract_recommender import SequentialRecommender
20from hopwise.model.loss import BPRLoss
23class FOSSIL(SequentialRecommender):
24 r"""FOSSIL uses similarity of the items as main purpose and uses high MC as a way of sequential preference improve of
25 ability of sequential recommendation
27 """ # noqa: E501
29 def __init__(self, config, dataset):
30 super().__init__(config, dataset)
32 # load the dataset information
33 self.n_users = dataset.num(self.USER_ID)
34 self.device = config["device"]
36 # load the parameters
37 self.embedding_size = config["embedding_size"]
38 self.order_len = config["order_len"]
39 assert self.order_len <= self.max_seq_length, "order_len can't longer than the max_seq_length"
40 self.reg_weight = config["reg_weight"]
41 self.alpha = config["alpha"]
43 # define the layers and loss type
44 self.item_embedding = nn.Embedding(self.n_items, self.embedding_size, padding_idx=0)
45 self.user_lambda = nn.Embedding(self.n_users, self.order_len)
46 self.lambda_ = nn.Parameter(torch.zeros(self.order_len))
48 self.loss_type = config["loss_type"]
49 if self.loss_type == "BPR":
50 self.loss_fct = BPRLoss()
51 elif self.loss_type == "CE":
52 self.loss_fct = nn.CrossEntropyLoss()
53 else:
54 raise NotImplementedError("Make sure 'loss_type' in ['BPR', 'CE']!")
56 # init the parameters of the model
57 self.apply(self.init_weights)
59 def inverse_seq_item_embedding(self, seq_item_embedding, seq_item_len):
60 """Inverse seq_item_embedding like this (simple to 2-dim):
62 [1,2,3,0,0,0] -- ??? -- >> [0,0,0,1,2,3]
64 first: [0,0,0,0,0,0] concat [1,2,3,0,0,0]
66 using gather_indexes: to get one by one
68 first get 3,then 2,last 1
69 """
70 zeros = torch.zeros_like(seq_item_embedding, dtype=torch.float).to(self.device)
71 # batch_size * seq_len * embedding_size
72 item_embedding_zeros = torch.cat([zeros, seq_item_embedding], dim=1)
73 # batch_size * 2_mul_seq_len * embedding_size
74 embedding_list = list()
75 for i in range(self.order_len):
76 embedding = self.gather_indexes(
77 item_embedding_zeros,
78 self.max_seq_length + seq_item_len - self.order_len + i,
79 )
80 embedding_list.append(embedding.unsqueeze(1))
81 short_item_embedding = torch.cat(embedding_list, dim=1)
82 # batch_size * short_len * embedding_size
84 return short_item_embedding
86 def reg_loss(self, user_embedding, item_embedding, seq_output):
87 reg_1 = self.reg_weight
88 loss_1 = (
89 reg_1 * torch.norm(user_embedding, p=2)
90 + reg_1 * torch.norm(item_embedding, p=2)
91 + reg_1 * torch.norm(seq_output, p=2)
92 )
94 return loss_1
96 def init_weights(self, module):
97 if isinstance(module, nn.Embedding) or isinstance(module, nn.Linear):
98 xavier_normal_(module.weight.data)
100 def forward(self, seq_item, seq_item_len, user):
101 seq_item_embedding = self.item_embedding(seq_item)
103 high_order_seq_item_embedding = self.inverse_seq_item_embedding(seq_item_embedding, seq_item_len)
104 # batch_size * order_len * embedding
106 high_order = self.get_high_order_Markov(high_order_seq_item_embedding, user)
107 similarity = self.get_similarity(seq_item_embedding, seq_item_len)
109 return high_order + similarity
111 def get_high_order_Markov(self, high_order_item_embedding, user):
112 """In order to get the inference of past items and the user's taste to the current predict item"""
113 user_lambda = self.user_lambda(user).unsqueeze(dim=2)
114 # batch_size * order_len * 1
115 lambda_ = self.lambda_.unsqueeze(dim=0).unsqueeze(dim=2)
116 # 1 * order_len * 1
117 lambda_ = torch.add(user_lambda, lambda_)
118 # batch_size * order_len * 1
119 high_order_item_embedding = torch.mul(high_order_item_embedding, lambda_)
120 # batch_size * order_len * embedding_size
121 high_order_item_embedding = high_order_item_embedding.sum(dim=1)
122 # batch_size * embedding_size
124 return high_order_item_embedding
126 def get_similarity(self, seq_item_embedding, seq_item_len):
127 """In order to get the inference of past items to the current predict item"""
128 coeff = torch.pow(seq_item_len.unsqueeze(1), -self.alpha).float()
129 # batch_size * 1
130 similarity = torch.mul(coeff, seq_item_embedding.sum(dim=1))
131 # batch_size * embedding_size
133 return similarity
135 def calculate_loss(self, interaction):
136 seq_item = interaction[self.ITEM_SEQ]
137 user = interaction[self.USER_ID]
138 seq_item_len = interaction[self.ITEM_SEQ_LEN]
139 seq_output = self.forward(seq_item, seq_item_len, user)
140 pos_items = interaction[self.POS_ITEM_ID]
141 pos_items_emb = self.item_embedding(pos_items)
143 user_lambda = self.user_lambda(user)
144 pos_items_embedding = self.item_embedding(pos_items)
145 if self.loss_type == "BPR":
146 neg_items = interaction[self.NEG_ITEM_ID]
147 neg_items_emb = self.item_embedding(neg_items)
148 pos_score = torch.sum(seq_output * pos_items_emb, dim=-1)
149 neg_score = torch.sum(seq_output * neg_items_emb, dim=-1)
150 loss = self.loss_fct(pos_score, neg_score)
151 return loss + self.reg_loss(user_lambda, pos_items_embedding, seq_output)
152 else: # self.loss_type = 'CE'
153 test_item_emb = self.item_embedding.weight
154 logits = torch.matmul(seq_output, test_item_emb.transpose(0, 1))
155 loss = self.loss_fct(logits, pos_items)
156 return loss + self.reg_loss(user_lambda, pos_items_embedding, seq_output)
158 def predict(self, interaction):
159 item_seq = interaction[self.ITEM_SEQ]
160 item_seq_len = interaction[self.ITEM_SEQ_LEN]
161 test_item = interaction[self.ITEM_ID]
162 user = interaction[self.USER_ID]
163 seq_output = self.forward(item_seq, item_seq_len, user)
164 test_item_emb = self.item_embedding(test_item)
165 scores = torch.mul(seq_output, test_item_emb).sum(dim=1)
166 return scores
168 def full_sort_predict(self, interaction):
169 item_seq = interaction[self.ITEM_SEQ]
170 user = interaction[self.USER_ID]
171 item_seq_len = interaction[self.ITEM_SEQ_LEN]
172 seq_output = self.forward(item_seq, item_seq_len, user)
173 test_items_emb = self.item_embedding.weight
174 scores = torch.matmul(seq_output, test_items_emb.transpose(0, 1))
175 return scores