Coverage for hopwise/model/general_recommender/spectralcf.py: 89%
72 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/10/2
2# @Author : Changxin Tian
3# @Email : cx.tian@outlook.com
5"""SpectralCF
6################################################
8Reference:
9 Lei Zheng et al. "Spectral collaborative filtering." in RecSys 2018.
11Reference code:
12 https://github.com/lzheng21/SpectralCF
13"""
15import torch
17from hopwise.model.abstract_recommender import GeneralRecommender
18from hopwise.model.init import xavier_uniform_initialization
19from hopwise.model.loss import BPRLoss, EmbLoss
20from hopwise.utils import InputType
23class SpectralCF(GeneralRecommender):
24 r"""SpectralCF is a spectral convolution model that directly learns latent factors of users and items
25 from the spectral domain for recommendation.
27 The spectral convolution operation with C input channels and F filters is shown as the following:
29 .. math::
30 \left[\begin{array} {c} X_{new}^{u} \\
31 X_{new}^{i} \end{array}\right]=\sigma\left(\left(U U^{\top}+U \Lambda U^{\top}\right)
32 \left[\begin{array}{c} X^{u} \\
33 X^{i} \end{array}\right] \Theta^{\prime}\right)
35 where :math:`X_{new}^{u} \in R^{n_{users} \times F}` and :math:`X_{new}^{i} \in R^{n_{items} \times F}`
36 denote convolution results learned with F filters from the spectral domain for users and items, respectively;
37 :math:`\sigma` denotes the logistic sigmoid function.
39 Note:
40 Our implementation is a improved version which is different from the original paper.
41 For a better stability, we replace :math:`U U^T` with identity matrix :math:`I` and
42 replace :math:`U \Lambda U^T` with laplace matrix :math:`L`.
43 """
45 input_type = InputType.PAIRWISE
47 def __init__(self, config, dataset):
48 super().__init__(config, dataset)
50 # load parameters info
51 self.n_layers = config["n_layers"]
52 self.emb_dim = config["embedding_size"]
53 self.reg_weight = config["reg_weight"]
55 # generate intermediate data
56 # "A_hat = I + L" is equivalent to "A_hat = U U^T + U \Lambda U^T"
57 I = dataset._create_eye_matrix() # noqa: E741
58 L = I - dataset._create_norm_adjacency_matrix(symmetric=False)
59 A_hat = I + L
60 self.A_hat = A_hat.to(self.device)
62 # define layers and loss
63 self.user_embedding = torch.nn.Embedding(num_embeddings=self.n_users, embedding_dim=self.emb_dim)
64 self.item_embedding = torch.nn.Embedding(num_embeddings=self.n_items, embedding_dim=self.emb_dim)
65 self.filters = torch.nn.ParameterList(
66 [
67 torch.nn.Parameter(
68 torch.normal(mean=0.01, std=0.02, size=(self.emb_dim, self.emb_dim)),
69 requires_grad=True,
70 )
71 for _ in range(self.n_layers)
72 ]
73 )
75 self.sigmoid = torch.nn.Sigmoid()
76 self.mf_loss = BPRLoss()
77 self.reg_loss = EmbLoss()
78 self.restore_user_e = None
79 self.restore_item_e = None
81 self.other_parameter_name = ["restore_user_e", "restore_item_e"]
82 # parameters initialization
83 self.apply(xavier_uniform_initialization)
85 def get_ego_embeddings(self):
86 r"""Get the embedding of users and items and combine to an embedding matrix.
88 Returns:
89 Tensor of the embedding matrix. Shape of (n_items+n_users, embedding_dim)
90 """
91 user_embeddings = self.user_embedding.weight
92 item_embeddings = self.item_embedding.weight
93 ego_embeddings = torch.cat([user_embeddings, item_embeddings], dim=0)
94 return ego_embeddings
96 def forward(self):
97 all_embeddings = self.get_ego_embeddings()
98 embeddings_list = [all_embeddings]
100 for k in range(self.n_layers):
101 all_embeddings = torch.sparse.mm(self.A_hat, all_embeddings)
102 all_embeddings = self.sigmoid(torch.mm(all_embeddings, self.filters[k]))
103 embeddings_list.append(all_embeddings)
105 new_embeddings = torch.cat(embeddings_list, dim=1)
106 user_all_embeddings, item_all_embeddings = torch.split(new_embeddings, [self.n_users, self.n_items])
107 return user_all_embeddings, item_all_embeddings
109 def calculate_loss(self, interaction):
110 if self.restore_user_e is not None or self.restore_item_e is not None:
111 self.restore_user_e, self.restore_item_e = None, None
113 user = interaction[self.USER_ID]
114 pos_item = interaction[self.ITEM_ID]
115 neg_item = interaction[self.NEG_ITEM_ID]
117 user_all_embeddings, item_all_embeddings = self.forward()
118 u_embeddings = user_all_embeddings[user]
119 pos_embeddings = item_all_embeddings[pos_item]
120 neg_embeddings = item_all_embeddings[neg_item]
121 pos_scores = torch.mul(u_embeddings, pos_embeddings).sum(dim=1)
122 neg_scores = torch.mul(u_embeddings, neg_embeddings).sum(dim=1)
124 mf_loss = self.mf_loss(pos_scores, neg_scores)
125 reg_loss = self.reg_loss(u_embeddings, pos_embeddings, neg_embeddings)
126 loss = mf_loss + self.reg_weight * reg_loss
128 return loss
130 def predict(self, interaction):
131 user = interaction[self.USER_ID]
132 item = interaction[self.ITEM_ID]
134 user_all_embeddings, item_all_embeddings = self.forward()
136 u_embeddings = user_all_embeddings[user]
137 i_embeddings = item_all_embeddings[item]
138 scores = torch.mul(u_embeddings, i_embeddings).sum(dim=1)
139 return scores
141 def full_sort_predict(self, interaction):
142 user = interaction[self.USER_ID]
143 if self.restore_user_e is None or self.restore_item_e is None:
144 self.restore_user_e, self.restore_item_e = self.forward()
145 u_embeddings = self.restore_user_e[user]
147 scores = torch.matmul(u_embeddings, self.restore_item_e.transpose(0, 1))
148 return scores.view(-1)