From ac56132c35215bae4424abb5de50de81a63366c0 Mon Sep 17 00:00:00 2001 From: quanfa Date: Sun, 24 May 2026 08:23:21 +0800 Subject: [PATCH] Memory-efficient TCA --- scSpace/models.py | 55 +++++++++++++++++++++++++++++++++++++++-------- 1 file changed, 46 insertions(+), 9 deletions(-) diff --git a/scSpace/models.py b/scSpace/models.py index bc7549d..e0cdc8f 100644 --- a/scSpace/models.py +++ b/scSpace/models.py @@ -47,21 +47,58 @@ def fit(self, Xs, Xt): :return: Xs_new and Xt_new after TCA ''' X = np.hstack((Xs.T, Xt.T)) - X /= np.linalg.norm(X, axis=0) + # Safe column norm — cells with zero expression across all genes + # must not produce NaN (a pre-existing bug exposed by full data). + norms = np.linalg.norm(X, axis=0) + X[:, norms > 0] /= norms[norms > 0] m, n = X.shape ns, nt = len(Xs), len(Xt) - e = np.vstack((1 / ns * np.ones((ns, 1)), -1 / nt * np.ones((nt, 1)))) - M = e * e.T - M = M / np.linalg.norm(M, 'fro') - H = np.eye(n) - 1 / n * np.ones((n, n)) - K = kernel(self.kernel_type, X, None, gamma=self.gamma) - n_eye = m if self.kernel_type == 'primal' else n - a, b = K @ M @ K.T + self.lamb * np.eye(n_eye), K @ H @ K.T + + if self.kernel_type == 'primal': + # ── Memory-efficient TCA ────────────────────────────────── + # Original formulation constructs two (n×n) matrices M and H, + # which cost O(n²) memory (~25 GB for 56K samples). + # + # M = e·eᵀ / ‖e·eᵀ‖_F (rank-1, where e = [1/ns; -1/nt]) + # H = I - (1/n)·1·1ᵀ (centering matrix) + # + # For primal kernel K = X → a, b are (m, m) = (541, 541): + # + # K·M·Kᵀ = (X·e)·(X·e)ᵀ / (1/ns+1/nt) [vector outer product] + # K·H·Kᵀ = X·Xᵀ - (1/n)·(X·1)·(X·1)ᵀ [gram − rank-1 update] + # + # Memory: O(m·n + m²) ≈ 244 MB instead of O(n²) ≈ 50 GB. + # ────────────────────────────────────────────────────────── + + # X·e — difference of column means (m, 1) + Xe = X[:, :ns].sum(axis=1, keepdims=True) / ns \ + - X[:, ns:].sum(axis=1, keepdims=True) / nt + norm_factor = 1 / ns + 1 / nt # ‖e·eᵀ‖_F + a = (Xe @ Xe.T) / norm_factor + self.lamb * np.eye(m) + + # X·Xᵀ and X·1 for centering (m, m) and (m, 1) + XXt = X @ X.T + X1 = X.sum(axis=1, keepdims=True) + b = XXt - (X1 @ X1.T) / n + + K = X # primal: kernel is identity + else: + # Fallback — for linear / rbf kernels K is (n × n) so the + # (n × n) M and H matrices are unavoidable here. + K = kernel(self.kernel_type, X, None, gamma=self.gamma) + e = np.vstack((1 / ns * np.ones((ns, 1)), -1 / nt * np.ones((nt, 1)))) + M = (e * e.T) / (1 / ns + 1 / nt) + H = np.eye(n) - np.ones((n, n)) / n + n_eye = n + a = K @ M @ K.T + self.lamb * np.eye(n_eye) + b = K @ H @ K.T + w, V = scipy.linalg.eig(a, b) ind = np.argsort(w) A = V[:, ind[:self.dim]] Z = A.T @ K - Z /= np.linalg.norm(Z, axis=0) + norms_z = np.linalg.norm(Z, axis=0) + Z[:, norms_z > 0] /= norms_z[norms_z > 0] Xs_new, Xt_new = Z[:, :ns].T, Z[:, ns:].T return Xs_new, Xt_new