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

1# @Time : 2021/2/16 

2# @Author : Haoran Cheng 

3# @Email : chenghaoran29@foxmail.com 

4 

5r"""RaCT 

6################################################ 

7Reference: 

8 Sam Lobel et al. "RaCT: Towards Amortized Ranking-Critical Training for Collaborative Filtering." in ICLR 2020. 

9 

10""" 

11 

12import numpy as np 

13import torch 

14import torch.nn.functional as F 

15from torch import nn 

16 

17from hopwise.model.abstract_recommender import AutoEncoderMixin, GeneralRecommender 

18from hopwise.model.init import xavier_normal_initialization 

19from hopwise.utils import InputType 

20 

21 

22class RaCT(GeneralRecommender, AutoEncoderMixin): 

23 r"""RaCT is a collaborative filtering model which uses methods based on actor-critic reinforcement learning for training. 

24 

25 We implement the RaCT model with only user dataloader. 

26 """ # noqa: E501 

27 

28 input_type = InputType.USERWISE 

29 

30 def __init__(self, config, dataset): 

31 super().__init__(config, dataset) 

32 

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"] 

38 

39 self.build_histroy_items(dataset) 

40 

41 self.update = 0 

42 

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:] 

45 

46 self.encoder = self.mlp_layers(self.encode_layer_dims) 

47 self.decoder = self.mlp_layers(self.decode_layer_dims) 

48 

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] 

54 

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) 

59 

60 self.train_stage = config["train_stage"] 

61 self.pre_model_path = config["pre_model_path"] 

62 

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 

85 

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) 

93 

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 

101 

102 def forward(self, rating_matrix): 

103 t = F.normalize(rating_matrix) 

104 

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 

108 

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 

112 

113 h = self.encoder(h) 

114 

115 mu = h[:, : int(self.lat_dim / 2)] 

116 logvar = h[:, int(self.lat_dim / 2) :] 

117 

118 z = self.reparameterize(mu, logvar) 

119 z = self.decoder(z) 

120 self.predict_matrix = z 

121 return z, mu, logvar 

122 

123 def calculate_actor_loss(self, interaction): 

124 user = interaction[self.USER_ID] 

125 rating_matrix = self.get_rating_matrix(user) 

126 

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 

132 

133 z, mu, logvar = self.forward(rating_matrix) 

134 

135 # KL loss 

136 kl_loss = -0.5 * (torch.sum(1 + logvar - mu.pow(2) - logvar.exp(), dim=1)) * anneal 

137 

138 # CE loss 

139 ce_loss = -(F.log_softmax(z, 1) * rating_matrix).sum(1) 

140 

141 return ce_loss + kl_loss 

142 

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) 

149 

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) 

160 

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) 

165 

166 topk_result = true_matrix[np.arange(users_num)[:, np.newaxis], idx_sorted[:, :k]] 

167 

168 number_non_zero = ((true_matrix > 0) * 1).sum(dim=1) 

169 

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) 

174 

175 return DCG / IDCG 

176 

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 

182 

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) 

187 

188 mse_loss = (y - score) ** 2 

189 return mse_loss 

190 

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 

195 

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() 

206 

207 def predict(self, interaction): 

208 user = interaction[self.USER_ID] 

209 item = interaction[self.ITEM_ID] 

210 

211 rating_matrix = self.get_rating_matrix(user) 

212 

213 scores, _, _ = self.forward(rating_matrix) 

214 

215 return scores[[torch.arange(len(item)).to(self.device), item]] 

216 

217 def full_sort_predict(self, interaction): 

218 user = interaction[self.USER_ID] 

219 

220 rating_matrix = self.get_rating_matrix(user) 

221 

222 scores, _, _ = self.forward(rating_matrix) 

223 

224 return scores.view(-1)