Coverage for hopwise/model/general_recommender/macridvae.py: 85%

109 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-09-30 13:25 +0000

1# @Time : 2020/12/23 

2# @Author : Yihong Guo 

3# @Email : gyihong@hotmail.com 

4 

5# UPDATE 

6# @Time : 2021/6/30, 

7# @Author : Xingyu Pan 

8# @email : xy_pan@foxmail.com 

9 

10r"""MacridVAE 

11################################################ 

12Reference: 

13 Jianxin Ma et al. "Learning Disentangled Representations for Recommendation." in NeurIPS 2019. 

14 

15Reference code: 

16 https://jianxinma.github.io/disentangle-recsys.html 

17""" 

18 

19import torch 

20import torch.nn.functional as F 

21from torch import nn 

22 

23from hopwise.model.abstract_recommender import AutoEncoderMixin, GeneralRecommender 

24from hopwise.model.init import xavier_normal_initialization 

25from hopwise.model.loss import EmbLoss 

26from hopwise.utils import InputType 

27 

28 

29class MacridVAE(GeneralRecommender, AutoEncoderMixin): 

30 r"""MacridVAE is an item-based collaborative filtering model that learns disentangled representations from user 

31 behavior and simultaneously ranks all items for each user. 

32 

33 We implement the model following the original author. 

34 """ 

35 

36 input_type = InputType.USERWISE 

37 

38 def __init__(self, config, dataset): 

39 super().__init__(config, dataset) 

40 

41 self.layers = config["encoder_hidden_size"] 

42 self.embedding_size = config["embedding_size"] 

43 self.drop_out = config["dropout_prob"] 

44 self.kfac = config["kfac"] 

45 self.tau = config["tau"] 

46 self.nogb = config["nogb"] 

47 self.anneal_cap = config["anneal_cap"] 

48 self.total_anneal_steps = config["total_anneal_steps"] 

49 self.regs = config["reg_weights"] 

50 self.std = config["std"] 

51 

52 self.update = 0 

53 self.build_histroy_items(dataset) 

54 self.encode_layer_dims = [self.n_items] + self.layers + [self.embedding_size * 2] 

55 

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

57 

58 self.item_embedding = nn.Embedding(self.n_items, self.embedding_size) 

59 self.k_embedding = nn.Embedding(self.kfac, self.embedding_size) 

60 

61 self.l2_loss = EmbLoss() 

62 # parameters initialization 

63 self.apply(xavier_normal_initialization) 

64 

65 def mlp_layers(self, layer_dims): 

66 mlp_modules = [] 

67 for i, (d_in, d_out) in enumerate(zip(layer_dims[:-1], layer_dims[1:])): 

68 mlp_modules.append(nn.Linear(d_in, d_out)) 

69 if i != len(layer_dims[:-1]) - 1: 

70 mlp_modules.append(nn.Tanh()) 

71 return nn.Sequential(*mlp_modules) 

72 

73 def reparameterize(self, mu, logvar): 

74 if self.training: 

75 std = torch.exp(0.5 * logvar) 

76 epsilon = torch.zeros_like(std).normal_(mean=0, std=self.std) 

77 return mu + epsilon * std 

78 else: 

79 return mu 

80 

81 def forward(self, rating_matrix): 

82 cores = F.normalize(self.k_embedding.weight, dim=1) 

83 items = F.normalize(self.item_embedding.weight, dim=1) 

84 

85 rating_matrix = F.normalize(rating_matrix) 

86 rating_matrix = F.dropout(rating_matrix, self.drop_out, training=self.training) 

87 

88 cates_logits = torch.matmul(items, cores.transpose(0, 1)) / self.tau 

89 

90 if self.nogb: 

91 cates = torch.softmax(cates_logits, dim=-1) 

92 else: 

93 cates_sample = F.gumbel_softmax(cates_logits, tau=1, hard=False, dim=-1) 

94 cates_mode = torch.softmax(cates_logits, dim=-1) 

95 cates = self.training * cates_sample + (1 - self.training) * cates_mode 

96 

97 probs = None 

98 mulist = [] 

99 logvarlist = [] 

100 for k in range(self.kfac): 

101 cates_k = cates[:, k].reshape(1, -1) 

102 # encoder 

103 x_k = rating_matrix * cates_k 

104 h = self.encoder(x_k) 

105 mu = h[:, : self.embedding_size] 

106 mu = F.normalize(mu, dim=1) 

107 logvar = h[:, self.embedding_size :] 

108 

109 mulist.append(mu) 

110 logvarlist.append(logvar) 

111 

112 z = self.reparameterize(mu, logvar) 

113 

114 # decoder 

115 z_k = F.normalize(z, dim=1) 

116 logits_k = torch.matmul(z_k, items.transpose(0, 1)) / self.tau 

117 probs_k = torch.exp(logits_k) 

118 probs_k = probs_k * cates_k 

119 probs = probs_k if (probs is None) else (probs + probs_k) 

120 

121 logits = torch.log(probs) 

122 

123 return logits, mulist, logvarlist 

124 

125 def calculate_loss(self, interaction): 

126 user = interaction[self.USER_ID] 

127 

128 rating_matrix = self.get_rating_matrix(user) 

129 

130 self.update += 1 

131 if self.total_anneal_steps > 0: 

132 anneal = min(self.anneal_cap, 1.0 * self.update / self.total_anneal_steps) 

133 else: 

134 anneal = self.anneal_cap 

135 

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

137 kl_loss = None 

138 for i in range(self.kfac): 

139 kl_ = -0.5 * torch.mean(torch.sum(1 + logvar[i] - logvar[i].exp(), dim=1)) 

140 kl_loss = kl_ if (kl_loss is None) else (kl_loss + kl_) 

141 

142 # CE loss 

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

144 

145 if self.regs[0] != 0 or self.regs[1] != 0: 

146 return ce_loss + kl_loss * anneal + self.reg_loss() 

147 

148 return ce_loss + kl_loss * anneal 

149 

150 def reg_loss(self): 

151 r"""Calculate the L2 normalization loss of model parameters. 

152 Including embedding matrices and weight matrices of model. 

153 

154 Returns: 

155 loss(torch.FloatTensor): The L2 Loss tensor. shape of [1,] 

156 """ 

157 reg_1, reg_2 = self.regs[:2] 

158 loss_1 = reg_1 * self.item_embedding.weight.norm(2) 

159 loss_2 = reg_1 * self.k_embedding.weight.norm(2) 

160 loss_3 = 0 

161 for name, parm in self.encoder.named_parameters(): 

162 if name.endswith("weight"): 

163 loss_3 = loss_3 + reg_2 * parm.norm(2) 

164 return loss_1 + loss_2 + loss_3 

165 

166 def predict(self, interaction): 

167 user = interaction[self.USER_ID] 

168 item = interaction[self.ITEM_ID] 

169 

170 rating_matrix = self.get_rating_matrix(user) 

171 

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

173 

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

175 

176 def full_sort_predict(self, interaction): 

177 user = interaction[self.USER_ID] 

178 

179 rating_matrix = self.get_rating_matrix(user) 

180 

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

182 

183 return scores.view(-1)