Coverage for hopwise/model/general_recommender/admmslim.py: 80%

66 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-09-30 13:25 +0000

1# @Time : 2021/01/09 

2# @Author : Deklan Webster 

3 

4r"""ADMMSLIM 

5################################################ 

6Reference: 

7 Steck et al. ADMM SLIM: Sparse Recommendations for Many Users. https://doi.org/10.1145/3336191.3371774 

8 

9""" 

10 

11import numpy as np 

12import torch 

13 

14from hopwise.model.abstract_recommender import GeneralRecommender 

15from hopwise.utils import InputType, ModelType 

16 

17 

18def soft_threshold(x, threshold): 

19 return (np.abs(x) > threshold) * (np.abs(x) - threshold) * np.sign(x) 

20 

21 

22def zero_mean_columns(a): 

23 return a - np.mean(a, axis=0) 

24 

25 

26def add_noise(t, mag=1e-5): 

27 return t + mag * torch.rand(t.shape) 

28 

29 

30class ADMMSLIM(GeneralRecommender): 

31 input_type = InputType.POINTWISE 

32 type = ModelType.TRADITIONAL 

33 

34 def __init__(self, config, dataset): 

35 super().__init__(config, dataset) 

36 

37 # need at least one param 

38 self.dummy_param = torch.nn.Parameter(torch.zeros(1)) 

39 

40 X = dataset.inter_matrix(form="csr").astype(np.float32) 

41 

42 num_users, num_items = X.shape 

43 

44 lambda1 = config["lambda1"] 

45 lambda2 = config["lambda2"] 

46 alpha = config["alpha"] 

47 rho = config["rho"] 

48 k = config["k"] 

49 positive_only = config["positive_only"] 

50 self.center_columns = config["center_columns"] 

51 self.item_means = X.mean(axis=0).getA1() 

52 

53 if self.center_columns: 

54 zero_mean_X = X.toarray() - self.item_means 

55 G = zero_mean_X.T @ zero_mean_X 

56 # large memory cost because we need to make X dense to subtract mean, delete asap 

57 del zero_mean_X 

58 else: 

59 G = (X.T @ X).toarray() 

60 

61 diag = lambda2 * np.diag(np.power(self.item_means, alpha)) + rho * np.identity(num_items) 

62 

63 P = np.linalg.inv(G + diag).astype(np.float32) 

64 B_aux = (P @ G).astype(np.float32) 

65 # initialize 

66 Gamma = np.zeros_like(G, dtype=np.float32) 

67 C = np.zeros_like(G, dtype=np.float32) 

68 

69 del diag, G 

70 # fixed number of iterations 

71 for _ in range(k): 

72 B_tilde = B_aux + P @ (rho * C - Gamma) 

73 gamma = np.diag(B_tilde) / (np.diag(P) + 1e-7) 

74 B = B_tilde - P * gamma 

75 C = soft_threshold(B + Gamma / rho, lambda1 / rho) 

76 if positive_only: 

77 C = (C > 0) * C 

78 Gamma += rho * (B - C) 

79 # torch doesn't support sparse tensor slicing, so will do everything with np/scipy 

80 self.item_similarity = C 

81 self.interaction_matrix = X 

82 

83 def forward(self): 

84 pass 

85 

86 def calculate_loss(self, interaction): 

87 return torch.nn.Parameter(torch.zeros(1)) 

88 

89 def predict(self, interaction): 

90 user = interaction[self.USER_ID].cpu().numpy() 

91 item = interaction[self.ITEM_ID].cpu().numpy() 

92 

93 user_interactions = self.interaction_matrix[user, :].toarray() 

94 

95 if self.center_columns: 

96 r = ( 

97 ((user_interactions - self.item_means) * self.item_similarity[:, item].T).sum(axis=1) 

98 ).flatten() + self.item_means[item] 

99 else: 

100 r = (user_interactions * self.item_similarity[:, item].T).sum(axis=1).flatten() 

101 

102 return add_noise(torch.from_numpy(r)).to(self.device) 

103 

104 def full_sort_predict(self, interaction): 

105 user = interaction[self.USER_ID].cpu().numpy() 

106 

107 user_interactions = self.interaction_matrix[user, :].toarray() 

108 

109 if self.center_columns: 

110 r = ((user_interactions - self.item_means) @ self.item_similarity + self.item_means).flatten() 

111 else: 

112 r = (user_interactions @ self.item_similarity).flatten() 

113 

114 return add_noise(torch.from_numpy(r))