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

1# @Time : 2020/10/2 

2# @Author : Changxin Tian 

3# @Email : cx.tian@outlook.com 

4 

5"""SpectralCF 

6################################################ 

7 

8Reference: 

9 Lei Zheng et al. "Spectral collaborative filtering." in RecSys 2018. 

10 

11Reference code: 

12 https://github.com/lzheng21/SpectralCF 

13""" 

14 

15import torch 

16 

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 

21 

22 

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. 

26 

27 The spectral convolution operation with C input channels and F filters is shown as the following: 

28 

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) 

34 

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. 

38 

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 """ 

44 

45 input_type = InputType.PAIRWISE 

46 

47 def __init__(self, config, dataset): 

48 super().__init__(config, dataset) 

49 

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"] 

54 

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) 

61 

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 ) 

74 

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 

80 

81 self.other_parameter_name = ["restore_user_e", "restore_item_e"] 

82 # parameters initialization 

83 self.apply(xavier_uniform_initialization) 

84 

85 def get_ego_embeddings(self): 

86 r"""Get the embedding of users and items and combine to an embedding matrix. 

87 

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 

95 

96 def forward(self): 

97 all_embeddings = self.get_ego_embeddings() 

98 embeddings_list = [all_embeddings] 

99 

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) 

104 

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 

108 

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 

112 

113 user = interaction[self.USER_ID] 

114 pos_item = interaction[self.ITEM_ID] 

115 neg_item = interaction[self.NEG_ITEM_ID] 

116 

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) 

123 

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 

127 

128 return loss 

129 

130 def predict(self, interaction): 

131 user = interaction[self.USER_ID] 

132 item = interaction[self.ITEM_ID] 

133 

134 user_all_embeddings, item_all_embeddings = self.forward() 

135 

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 

140 

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] 

146 

147 scores = torch.matmul(u_embeddings, self.restore_item_e.transpose(0, 1)) 

148 return scores.view(-1)