Coverage for hopwise/model/general_recommender/recvae.py: 95%
111 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/28
2# @Author : Lanling Xu
3# @Email : xulanling_sherry@163.com
5r"""RecVAE
6################################################
7Reference:
8 Shenbin, Ilya, et al. "RecVAE: A new variational autoencoder for Top-N recommendations with implicit feedback." In WSDM 2020.
10Reference code:
11 https://github.com/ilya-shenbin/RecVAE
12""" # noqa: E501
14from copy import deepcopy
16import numpy as np
17import torch
18import torch.nn.functional as F
19from torch import nn
21from hopwise.model.abstract_recommender import AutoEncoderMixin, GeneralRecommender
22from hopwise.model.init import xavier_normal_initialization
23from hopwise.utils import InputType
26def swish(x):
27 r"""Swish activation function:
29 .. math::
30 \text{Swish}(x) = \frac{x}{1 + \exp(-x)}
31 """
32 return x.mul(torch.sigmoid(x))
35def log_norm_pdf(x, mu, logvar):
36 return -0.5 * (logvar + np.log(2 * np.pi) + (x - mu).pow(2) / logvar.exp())
39class CompositePrior(nn.Module):
40 def __init__(self, hidden_dim, latent_dim, input_dim, mixture_weights):
41 super().__init__()
43 self.mixture_weights = mixture_weights
45 self.mu_prior = nn.Parameter(torch.Tensor(1, latent_dim), requires_grad=False)
46 self.mu_prior.data.fill_(0)
48 self.logvar_prior = nn.Parameter(torch.Tensor(1, latent_dim), requires_grad=False)
49 self.logvar_prior.data.fill_(0)
51 self.logvar_uniform_prior = nn.Parameter(torch.Tensor(1, latent_dim), requires_grad=False)
52 self.logvar_uniform_prior.data.fill_(10)
54 self.encoder_old = Encoder(hidden_dim, latent_dim, input_dim)
55 self.encoder_old.requires_grad_(False)
57 def forward(self, x, z):
58 post_mu, post_logvar = self.encoder_old(x, 0)
60 stnd_prior = log_norm_pdf(z, self.mu_prior, self.logvar_prior)
61 post_prior = log_norm_pdf(z, post_mu, post_logvar)
62 unif_prior = log_norm_pdf(z, self.mu_prior, self.logvar_uniform_prior)
64 gaussians = [stnd_prior, post_prior, unif_prior]
65 gaussians = [g.add(np.log(w)) for g, w in zip(gaussians, self.mixture_weights)]
67 density_per_gaussian = torch.stack(gaussians, dim=-1)
69 return torch.logsumexp(density_per_gaussian, dim=-1)
72class Encoder(nn.Module):
73 def __init__(self, hidden_dim, latent_dim, input_dim, eps=1e-1):
74 super().__init__()
76 self.fc1 = nn.Linear(input_dim, hidden_dim)
77 self.ln1 = nn.LayerNorm(hidden_dim, eps=eps)
78 self.fc2 = nn.Linear(hidden_dim, hidden_dim)
79 self.ln2 = nn.LayerNorm(hidden_dim, eps=eps)
80 self.fc3 = nn.Linear(hidden_dim, hidden_dim)
81 self.ln3 = nn.LayerNorm(hidden_dim, eps=eps)
82 self.fc4 = nn.Linear(hidden_dim, hidden_dim)
83 self.ln4 = nn.LayerNorm(hidden_dim, eps=eps)
84 self.fc5 = nn.Linear(hidden_dim, hidden_dim)
85 self.ln5 = nn.LayerNorm(hidden_dim, eps=eps)
86 self.fc_mu = nn.Linear(hidden_dim, latent_dim)
87 self.fc_logvar = nn.Linear(hidden_dim, latent_dim)
89 def forward(self, x, dropout_prob):
90 x = F.normalize(x)
91 x = F.dropout(x, dropout_prob, training=self.training)
93 h1 = self.ln1(swish(self.fc1(x)))
94 h2 = self.ln2(swish(self.fc2(h1) + h1))
95 h3 = self.ln3(swish(self.fc3(h2) + h1 + h2))
96 h4 = self.ln4(swish(self.fc4(h3) + h1 + h2 + h3))
97 h5 = self.ln5(swish(self.fc5(h4) + h1 + h2 + h3 + h4))
98 return self.fc_mu(h5), self.fc_logvar(h5)
101class RecVAE(GeneralRecommender, AutoEncoderMixin):
102 r"""Collaborative Denoising Auto-Encoder (RecVAE) is a recommendation model
103 for top-N recommendation with implicit feedback.
105 We implement the model following the original author
106 """
108 input_type = InputType.USERWISE
110 def __init__(self, config, dataset):
111 super().__init__(config, dataset)
113 self.hidden_dim = config["hidden_dimension"]
114 self.latent_dim = config["latent_dimension"]
115 self.dropout_prob = config["dropout_prob"]
116 self.beta = config["beta"]
117 self.mixture_weights = config["mixture_weights"]
118 self.gamma = config["gamma"]
120 self.build_histroy_items(dataset)
122 self.encoder = Encoder(self.hidden_dim, self.latent_dim, self.n_items)
123 self.prior = CompositePrior(self.hidden_dim, self.latent_dim, self.n_items, self.mixture_weights)
124 self.decoder = nn.Linear(self.latent_dim, self.n_items)
126 # parameters initialization
127 self.apply(xavier_normal_initialization)
129 def reparameterize(self, mu, logvar):
130 if self.training:
131 std = torch.exp(0.5 * logvar)
132 epsilon = torch.zeros_like(std).normal_(mean=0, std=0.01)
133 return mu + epsilon * std
134 else:
135 return mu
137 def forward(self, rating_matrix, dropout_prob):
138 mu, logvar = self.encoder(rating_matrix, dropout_prob=dropout_prob)
139 z = self.reparameterize(mu, logvar)
140 x_pred = self.decoder(z)
141 return x_pred, mu, logvar, z
143 def calculate_loss(self, interaction, encoder_flag):
144 user = interaction[self.USER_ID]
145 rating_matrix = self.get_rating_matrix(user)
146 if encoder_flag:
147 dropout_prob = self.dropout_prob
148 else:
149 dropout_prob = 0
150 x_pred, mu, logvar, z = self.forward(rating_matrix, dropout_prob)
152 if self.gamma:
153 norm = rating_matrix.sum(dim=-1)
154 kl_weight = self.gamma * norm
155 else:
156 kl_weight = self.beta
158 mll = (F.log_softmax(x_pred, dim=-1) * rating_matrix).sum(dim=-1).mean()
159 kld = (log_norm_pdf(z, mu, logvar) - self.prior(rating_matrix, z)).sum(dim=-1).mul(kl_weight).mean()
160 negative_elbo = -(mll - kld)
162 return negative_elbo
164 def predict(self, interaction):
165 user = interaction[self.USER_ID]
166 item = interaction[self.ITEM_ID]
168 rating_matrix = self.get_rating_matrix(user)
170 scores, _, _, _ = self.forward(rating_matrix, self.dropout_prob)
172 return scores[[torch.arange(len(item)).to(self.device), item]]
174 def full_sort_predict(self, interaction):
175 user = interaction[self.USER_ID]
177 rating_matrix = self.get_rating_matrix(user)
179 scores, _, _, _ = self.forward(rating_matrix, self.dropout_prob)
181 return scores.view(-1)
183 def update_prior(self):
184 self.prior.encoder_old.load_state_dict(deepcopy(self.encoder.state_dict()))