Skip to content

Commit 9af882d

Browse files
authored
Merge pull request #217 from SauersML/codex/fix-failing-tests-in-pca_tests
Fix Python PCA test harness interface
2 parents 7feae3c + 0b05117 commit 9af882d

1 file changed

Lines changed: 53 additions & 10 deletions

File tree

‎tests/pca.py‎

Lines changed: 53 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -2,11 +2,44 @@
22
import sys
33
import numpy as np
44
from sklearn.decomposition import PCA
5-
from sklearn.preprocessing import StandardScaler
6-
import json # For original_main_logic
7-
import os # For original_main_logic
8-
import scipy.linalg as la # For original_main_logic
9-
import random # For original_main_logic
5+
import json # For original_main_logic
6+
import os # For original_main_logic
7+
import scipy.linalg as la # For original_main_logic
8+
import random # For original_main_logic
9+
10+
11+
NEAR_ZERO_THRESHOLD = 1e-9
12+
13+
14+
def center_and_scale_columns_np(X):
15+
"""Replicates the Rust PCA preprocessing: column-wise centering and scaling
16+
using the unbiased sample standard deviation (n-1 in the denominator) with
17+
small-value sanitization."""
18+
19+
X = np.asarray(X, dtype=float)
20+
if X.size == 0:
21+
return X.copy(), np.array([]), np.array([])
22+
23+
n_samples, n_features = X.shape
24+
25+
mean = np.mean(X, axis=0)
26+
centered = X - mean
27+
28+
if n_samples > 1:
29+
sum_sq = np.sum(centered * centered, axis=0)
30+
variance = np.maximum(sum_sq, 0.0) / float(n_samples - 1)
31+
else:
32+
variance = np.zeros(n_features)
33+
34+
std = np.sqrt(variance, dtype=float)
35+
sanitized_std = np.where(
36+
(~np.isfinite(std)) | (std <= NEAR_ZERO_THRESHOLD),
37+
1.0,
38+
std,
39+
)
40+
41+
scaled = centered / sanitized_std
42+
return scaled, mean, sanitized_std
1043

1144

1245
def print_numpy_array_for_rust(arr):
@@ -145,8 +178,7 @@ def generate_random_data_original(samples=5, features=5, random_seed=None):
145178
return np.random.randn(samples, features)
146179

147180
def manual_pca_original(X, n_components=None):
148-
scaler = StandardScaler()
149-
X_scaled = scaler.fit_transform(X)
181+
X_scaled, _, _ = center_and_scale_columns_np(X)
150182
n_samples, n_features = X_scaled.shape
151183
if n_components is None: n_components = min(n_samples, n_features)
152184
else: n_components = min(n_components, min(n_samples, n_features))
@@ -181,8 +213,7 @@ def manual_pca_original(X, n_components=None):
181213
return X_transformed, components, eigvals
182214

183215
def library_pca_original(X, n_components=None):
184-
scaler = StandardScaler()
185-
X_scaled = scaler.fit_transform(X)
216+
X_scaled, _, _ = center_and_scale_columns_np(X)
186217
n_samples, n_features = X_scaled.shape
187218
max_components = min(n_samples, n_features)
188219
if n_components is None: n_components = max_components
@@ -288,7 +319,19 @@ def parse_arguments_main():
288319

289320
# Argument for n_components, used by both modes
290321
# For --generate-reference-pca, it's required. For original_main_logic, it's optional.
291-
parser.add_argument("-k", "--n-components", type=int, help="Number of components for PCA.")
322+
# Accept both hyphenated and underscored versions of the flag. The Rust tests
323+
# currently invoke the script with "--n_components", while the original
324+
# command-line interface exposed "--n-components". Argparse treats hyphens
325+
# and underscores as distinct option names, so support both spellings here
326+
# by routing them to the same destination.
327+
parser.add_argument(
328+
"-k",
329+
"--n-components",
330+
"--n_components",
331+
dest="n_components",
332+
type=int,
333+
help="Number of components for PCA.",
334+
)
292335

293336
# Argument to switch to reference generation mode
294337
parser.add_argument("--generate-reference-pca",

0 commit comments

Comments
 (0)