Coverage for hopwise/model/general_recommender/multivae.py: 91%
66 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/14
2# @Author : Yihong Guo
3# @Email : gyihong@hotmail.com
5r"""MultiVAE
6################################################
7Reference:
8 Dawen Liang et al. "Variational Autoencoders for Collaborative Filtering." in WWW 2018.
10"""
12import torch
13import torch.nn.functional as F
14from torch import nn
16from hopwise.model.abstract_recommender import AutoEncoderMixin, GeneralRecommender
17from hopwise.model.init import xavier_normal_initialization
18from hopwise.utils import InputType
21class MultiVAE(GeneralRecommender, AutoEncoderMixin):
22 r"""MultiVAE is an item-based collaborative filtering model that simultaneously ranks all items for each user.
24 We implement the MultiVAE model with only user dataloader.
25 """
27 input_type = InputType.USERWISE
29 def __init__(self, config, dataset):
30 super().__init__(config, dataset)
32 self.layers = config["mlp_hidden_size"]
33 self.lat_dim = config["latent_dimension"]
34 self.drop_out = config["dropout_prob"]
35 self.anneal_cap = config["anneal_cap"]
36 self.total_anneal_steps = config["total_anneal_steps"]
38 self.build_histroy_items(dataset)
40 self.update = 0
42 self.encode_layer_dims = [self.n_items] + self.layers + [self.lat_dim]
43 self.decode_layer_dims = [int(self.lat_dim / 2)] + self.encode_layer_dims[::-1][1:]
45 self.encoder = self.mlp_layers(self.encode_layer_dims)
46 self.decoder = self.mlp_layers(self.decode_layer_dims)
48 # parameters initialization
49 self.apply(xavier_normal_initialization)
51 def mlp_layers(self, layer_dims):
52 mlp_modules = []
53 for i, (d_in, d_out) in enumerate(zip(layer_dims[:-1], layer_dims[1:])):
54 mlp_modules.append(nn.Linear(d_in, d_out))
55 if i != len(layer_dims[:-1]) - 1:
56 mlp_modules.append(nn.Tanh())
57 return nn.Sequential(*mlp_modules)
59 def reparameterize(self, mu, logvar):
60 if self.training:
61 std = torch.exp(0.5 * logvar)
62 epsilon = torch.zeros_like(std).normal_(mean=0, std=0.01)
63 return mu + epsilon * std
64 else:
65 return mu
67 def forward(self, rating_matrix):
68 h = F.normalize(rating_matrix)
70 h = F.dropout(h, self.drop_out, training=self.training)
72 h = self.encoder(h)
74 mu = h[:, : int(self.lat_dim / 2)]
75 logvar = h[:, int(self.lat_dim / 2) :]
77 z = self.reparameterize(mu, logvar)
78 z = self.decoder(z)
79 return z, mu, logvar
81 def calculate_loss(self, interaction):
82 user = interaction[self.USER_ID]
83 rating_matrix = self.get_rating_matrix(user)
85 self.update += 1
86 if self.total_anneal_steps > 0:
87 anneal = min(self.anneal_cap, 1.0 * self.update / self.total_anneal_steps)
88 else:
89 anneal = self.anneal_cap
91 z, mu, logvar = self.forward(rating_matrix)
93 # KL loss
94 kl_loss = -0.5 * torch.mean(torch.sum(1 + logvar - mu.pow(2) - logvar.exp(), dim=1)) * anneal
96 # CE loss
97 ce_loss = -(F.log_softmax(z, 1) * rating_matrix).sum(1).mean()
99 return ce_loss + kl_loss
101 def predict(self, interaction):
102 user = interaction[self.USER_ID]
103 item = interaction[self.ITEM_ID]
105 rating_matrix = self.get_rating_matrix(user)
107 scores, _, _ = self.forward(rating_matrix)
109 return scores[[torch.arange(len(item)).to(self.device), item]]
111 def full_sort_predict(self, interaction):
112 user = interaction[self.USER_ID]
114 rating_matrix = self.get_rating_matrix(user)
116 scores, _, _ = self.forward(rating_matrix)
118 return scores.view(-1)