Coverage for hopwise/model/general_recommender/ract.py: 16%
148 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/2/16
2# @Author : Haoran Cheng
3# @Email : chenghaoran29@foxmail.com
5r"""RaCT
6################################################
7Reference:
8 Sam Lobel et al. "RaCT: Towards Amortized Ranking-Critical Training for Collaborative Filtering." in ICLR 2020.
10"""
12import numpy as np
13import torch
14import torch.nn.functional as F
15from torch import nn
17from hopwise.model.abstract_recommender import AutoEncoderMixin, GeneralRecommender
18from hopwise.model.init import xavier_normal_initialization
19from hopwise.utils import InputType
22class RaCT(GeneralRecommender, AutoEncoderMixin):
23 r"""RaCT is a collaborative filtering model which uses methods based on actor-critic reinforcement learning for training.
25 We implement the RaCT model with only user dataloader.
26 """ # noqa: E501
28 input_type = InputType.USERWISE
30 def __init__(self, config, dataset):
31 super().__init__(config, dataset)
33 self.layers = config["mlp_hidden_size"]
34 self.lat_dim = config["latent_dimension"]
35 self.drop_out = config["dropout_prob"]
36 self.anneal_cap = config["anneal_cap"]
37 self.total_anneal_steps = config["total_anneal_steps"]
39 self.build_histroy_items(dataset)
41 self.update = 0
43 self.encode_layer_dims = [self.n_items] + self.layers + [self.lat_dim]
44 self.decode_layer_dims = [int(self.lat_dim / 2)] + self.encode_layer_dims[::-1][1:]
46 self.encoder = self.mlp_layers(self.encode_layer_dims)
47 self.decoder = self.mlp_layers(self.decode_layer_dims)
49 self.critic_layers = config["critic_layers"]
50 self.metrics_k = config["metrics_k"]
51 self.number_of_seen_items = 0
52 self.number_of_unseen_items = 0
53 self.critic_layer_dims = [3] + self.critic_layers + [1]
55 self.input_matrix = None
56 self.predict_matrix = None
57 self.true_matrix = None
58 self.critic_net = self.construct_critic_layers(self.critic_layer_dims)
60 self.train_stage = config["train_stage"]
61 self.pre_model_path = config["pre_model_path"]
63 # parameters initialization
64 assert self.train_stage in ["actor_pretrain", "critic_pretrain", "finetune"]
65 if self.train_stage == "actor_pretrain":
66 self.apply(xavier_normal_initialization)
67 for p in self.critic_net.parameters():
68 p.requires_grad = False
69 elif self.train_stage == "critic_pretrain":
70 # load pretrained model for finetune
71 pretrained = torch.load(self.pre_model_path)
72 self.logger.info("Load pretrained model from" + self.pre_model_path)
73 self.load_state_dict(pretrained["state_dict"])
74 for p in self.encoder.parameters():
75 p.requires_grad = False
76 for p in self.decoder.parameters():
77 p.requires_grad = False
78 else:
79 # load pretrained model for finetune
80 pretrained = torch.load(self.pre_model_path)
81 self.logger.info("Load pretrained model from" + self.pre_model_path)
82 self.load_state_dict(pretrained["state_dict"])
83 for p in self.critic_net.parameters():
84 p.requires_grad = False
86 def mlp_layers(self, layer_dims):
87 mlp_modules = []
88 for i, (d_in, d_out) in enumerate(zip(layer_dims[:-1], layer_dims[1:])):
89 mlp_modules.append(nn.Linear(d_in, d_out))
90 if i != len(layer_dims[:-1]) - 1:
91 mlp_modules.append(nn.Tanh())
92 return nn.Sequential(*mlp_modules)
94 def reparameterize(self, mu, logvar):
95 if self.training:
96 std = torch.exp(0.5 * logvar)
97 epsilon = torch.zeros_like(std).normal_(mean=0, std=0.01)
98 return mu + epsilon * std
99 else:
100 return mu
102 def forward(self, rating_matrix):
103 t = F.normalize(rating_matrix)
105 h = F.dropout(t, self.drop_out, training=self.training) * (1 - self.drop_out)
106 self.input_matrix = h
107 self.number_of_seen_items = (h != 0).sum(dim=1) # network input
109 mask = (h > 0) * (t > 0)
110 self.true_matrix = t * ~mask
111 self.number_of_unseen_items = (self.true_matrix != 0).sum(dim=1) # remaining input
113 h = self.encoder(h)
115 mu = h[:, : int(self.lat_dim / 2)]
116 logvar = h[:, int(self.lat_dim / 2) :]
118 z = self.reparameterize(mu, logvar)
119 z = self.decoder(z)
120 self.predict_matrix = z
121 return z, mu, logvar
123 def calculate_actor_loss(self, interaction):
124 user = interaction[self.USER_ID]
125 rating_matrix = self.get_rating_matrix(user)
127 self.update += 1
128 if self.total_anneal_steps > 0:
129 anneal = min(self.anneal_cap, 1.0 * self.update / self.total_anneal_steps)
130 else:
131 anneal = self.anneal_cap
133 z, mu, logvar = self.forward(rating_matrix)
135 # KL loss
136 kl_loss = -0.5 * (torch.sum(1 + logvar - mu.pow(2) - logvar.exp(), dim=1)) * anneal
138 # CE loss
139 ce_loss = -(F.log_softmax(z, 1) * rating_matrix).sum(1)
141 return ce_loss + kl_loss
143 def construct_critic_input(self, actor_loss):
144 critic_inputs = []
145 critic_inputs.append(self.number_of_seen_items)
146 critic_inputs.append(self.number_of_unseen_items)
147 critic_inputs.append(actor_loss)
148 return torch.stack(critic_inputs, dim=1)
150 def construct_critic_layers(self, layer_dims):
151 mlp_modules = []
152 mlp_modules.append(nn.BatchNorm1d(3))
153 for i, (d_in, d_out) in enumerate(zip(layer_dims[:-1], layer_dims[1:])):
154 mlp_modules.append(nn.Linear(d_in, d_out))
155 if i != len(layer_dims[:-1]) - 1:
156 mlp_modules.append(nn.ReLU())
157 else:
158 mlp_modules.append(nn.Sigmoid())
159 return nn.Sequential(*mlp_modules)
161 def calculate_ndcg(self, predict_matrix, true_matrix, input_matrix, k):
162 users_num = predict_matrix.shape[0]
163 predict_matrix[input_matrix.nonzero(as_tuple=True)] = -np.inf
164 _, idx_sorted = torch.sort(predict_matrix, dim=1, descending=True)
166 topk_result = true_matrix[np.arange(users_num)[:, np.newaxis], idx_sorted[:, :k]]
168 number_non_zero = ((true_matrix > 0) * 1).sum(dim=1)
170 tp = 1.0 / torch.log2(torch.arange(2, k + 2).type(torch.FloatTensor)).to(topk_result.device)
171 DCG = (topk_result * tp).sum(dim=1)
172 IDCG = torch.Tensor([(tp[: min(n, k)]).sum() for n in number_non_zero]).to(topk_result.device)
173 IDCG = torch.maximum(0.1 * torch.ones_like(IDCG).to(IDCG.device), IDCG)
175 return DCG / IDCG
177 def critic_forward(self, actor_loss):
178 h = self.construct_critic_input(actor_loss)
179 y = self.critic_net(h)
180 y = torch.squeeze(y)
181 return y
183 def calculate_critic_loss(self, interaction):
184 actor_loss = self.calculate_actor_loss(interaction)
185 y = self.critic_forward(actor_loss)
186 score = self.calculate_ndcg(self.predict_matrix, self.true_matrix, self.input_matrix, self.metrics_k)
188 mse_loss = (y - score) ** 2
189 return mse_loss
191 def calculate_ac_loss(self, interaction):
192 actor_loss = self.calculate_actor_loss(interaction)
193 y = self.critic_forward(actor_loss)
194 return -1 * y
196 def calculate_loss(self, interaction):
197 # actor_pretrain
198 if self.train_stage == "actor_pretrain":
199 return self.calculate_actor_loss(interaction).mean()
200 # critic_pretrain
201 elif self.train_stage == "critic_pretrain":
202 return self.calculate_critic_loss(interaction).mean()
203 # finetune
204 else:
205 return self.calculate_ac_loss(interaction).mean()
207 def predict(self, interaction):
208 user = interaction[self.USER_ID]
209 item = interaction[self.ITEM_ID]
211 rating_matrix = self.get_rating_matrix(user)
213 scores, _, _ = self.forward(rating_matrix)
215 return scores[[torch.arange(len(item)).to(self.device), item]]
217 def full_sort_predict(self, interaction):
218 user = interaction[self.USER_ID]
220 rating_matrix = self.get_rating_matrix(user)
222 scores, _, _ = self.forward(rating_matrix)
224 return scores.view(-1)