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
« 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
4r"""ADMMSLIM
5################################################
6Reference:
7 Steck et al. ADMM SLIM: Sparse Recommendations for Many Users. https://doi.org/10.1145/3336191.3371774
9"""
11import numpy as np
12import torch
14from hopwise.model.abstract_recommender import GeneralRecommender
15from hopwise.utils import InputType, ModelType
18def soft_threshold(x, threshold):
19 return (np.abs(x) > threshold) * (np.abs(x) - threshold) * np.sign(x)
22def zero_mean_columns(a):
23 return a - np.mean(a, axis=0)
26def add_noise(t, mag=1e-5):
27 return t + mag * torch.rand(t.shape)
30class ADMMSLIM(GeneralRecommender):
31 input_type = InputType.POINTWISE
32 type = ModelType.TRADITIONAL
34 def __init__(self, config, dataset):
35 super().__init__(config, dataset)
37 # need at least one param
38 self.dummy_param = torch.nn.Parameter(torch.zeros(1))
40 X = dataset.inter_matrix(form="csr").astype(np.float32)
42 num_users, num_items = X.shape
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()
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()
61 diag = lambda2 * np.diag(np.power(self.item_means, alpha)) + rho * np.identity(num_items)
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)
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
83 def forward(self):
84 pass
86 def calculate_loss(self, interaction):
87 return torch.nn.Parameter(torch.zeros(1))
89 def predict(self, interaction):
90 user = interaction[self.USER_ID].cpu().numpy()
91 item = interaction[self.ITEM_ID].cpu().numpy()
93 user_interactions = self.interaction_matrix[user, :].toarray()
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()
102 return add_noise(torch.from_numpy(r)).to(self.device)
104 def full_sort_predict(self, interaction):
105 user = interaction[self.USER_ID].cpu().numpy()
107 user_interactions = self.interaction_matrix[user, :].toarray()
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()
114 return add_noise(torch.from_numpy(r))