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
« 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
5# UPDATE
6# @Time : 2021/6/30,
7# @Author : Xingyu Pan
8# @email : xy_pan@foxmail.com
10r"""MacridVAE
11################################################
12Reference:
13 Jianxin Ma et al. "Learning Disentangled Representations for Recommendation." in NeurIPS 2019.
15Reference code:
16 https://jianxinma.github.io/disentangle-recsys.html
17"""
19import torch
20import torch.nn.functional as F
21from torch import nn
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
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.
33 We implement the model following the original author.
34 """
36 input_type = InputType.USERWISE
38 def __init__(self, config, dataset):
39 super().__init__(config, dataset)
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"]
52 self.update = 0
53 self.build_histroy_items(dataset)
54 self.encode_layer_dims = [self.n_items] + self.layers + [self.embedding_size * 2]
56 self.encoder = self.mlp_layers(self.encode_layer_dims)
58 self.item_embedding = nn.Embedding(self.n_items, self.embedding_size)
59 self.k_embedding = nn.Embedding(self.kfac, self.embedding_size)
61 self.l2_loss = EmbLoss()
62 # parameters initialization
63 self.apply(xavier_normal_initialization)
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)
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
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)
85 rating_matrix = F.normalize(rating_matrix)
86 rating_matrix = F.dropout(rating_matrix, self.drop_out, training=self.training)
88 cates_logits = torch.matmul(items, cores.transpose(0, 1)) / self.tau
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
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 :]
109 mulist.append(mu)
110 logvarlist.append(logvar)
112 z = self.reparameterize(mu, logvar)
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)
121 logits = torch.log(probs)
123 return logits, mulist, logvarlist
125 def calculate_loss(self, interaction):
126 user = interaction[self.USER_ID]
128 rating_matrix = self.get_rating_matrix(user)
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
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_)
142 # CE loss
143 ce_loss = -(F.log_softmax(z, 1) * rating_matrix).sum(1).mean()
145 if self.regs[0] != 0 or self.regs[1] != 0:
146 return ce_loss + kl_loss * anneal + self.reg_loss()
148 return ce_loss + kl_loss * anneal
150 def reg_loss(self):
151 r"""Calculate the L2 normalization loss of model parameters.
152 Including embedding matrices and weight matrices of model.
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
166 def predict(self, interaction):
167 user = interaction[self.USER_ID]
168 item = interaction[self.ITEM_ID]
170 rating_matrix = self.get_rating_matrix(user)
172 scores, _, _ = self.forward(rating_matrix)
174 return scores[[torch.arange(len(item)).to(self.device), item]]
176 def full_sort_predict(self, interaction):
177 user = interaction[self.USER_ID]
179 rating_matrix = self.get_rating_matrix(user)
181 scores, _, _ = self.forward(rating_matrix)
183 return scores.view(-1)