1414"""
1515from time import time
1616
17+ import csv
18+ import gzip
1719import numpy as np
18- import pandas as pd
1920import os
2021import sys
2122import harmonypy as hm
@@ -29,6 +30,35 @@ def pearsonr(x, y):
2930 return r
3031
3132
33+ def read_tsv (path ):
34+ """Read a TSV file, return (header, columns_dict).
35+
36+ Numeric columns are returned as float64 arrays.
37+ Non-numeric columns are returned as string arrays.
38+ """
39+ with gzip .open (path , "rt" ) if path .endswith (".gz" ) else open (path ) as f :
40+ reader = csv .reader (f , delimiter = "\t " )
41+ header = next (reader )
42+ rows = list (reader )
43+ cols = {}
44+ for i , name in enumerate (header ):
45+ values = [row [i ] for row in rows ]
46+ try :
47+ cols [name ] = np .array (values , dtype = np .float64 )
48+ except ValueError :
49+ cols [name ] = np .array (values )
50+ return header , cols
51+
52+
53+ def cols_to_matrix (cols , header ):
54+ """Extract numeric columns into a matrix (N x d)."""
55+ numeric = []
56+ for name in header :
57+ if cols [name ].dtype == np .float64 :
58+ numeric .append (cols [name ])
59+ return np .column_stack (numeric )
60+
61+
3262def _get_current_rss_mb ():
3363 """Get current RSS (resident set size) in MB. Works on macOS and Linux.
3464
@@ -89,27 +119,27 @@ def test_random_seed():
89119 print ("TEST: test_random_seed" )
90120 print ("=" * 60 )
91121
92- meta_data = pd .read_csv ("data/pbmc_3500_meta.tsv.gz" , sep = "\t " )
93- data_mat = pd .read_csv ("data/pbmc_3500_pcs.tsv.gz" , sep = "\t " )
122+ _ , meta_cols = read_tsv ("data/pbmc_3500_meta.tsv.gz" )
123+ _ , pcs_cols = read_tsv ("data/pbmc_3500_pcs.tsv.gz" )
124+ data_mat = cols_to_matrix (pcs_cols , list (pcs_cols .keys ()))
94125
95126 def run (random_state ):
96127 ho = hm .run_harmony (data_mat ,
97- meta_data , ['donor' ],
128+ meta_cols , ['donor' ],
98129 max_iter_harmony = 2 ,
99130 max_iter_kmeans = 2 ,
100131 verbose = False ,
101132 random_state = random_state )
102133 return ho .Z_corr
103134
104135 # Assert same results when random_state is set.
105- # Note: MPS (Apple Silicon) has slight non-determinism, so we use relaxed tolerance
106136 print ("\n --- Testing reproducibility with random_state=42 ---" )
107137 result1 = run (42 )
108138 result2 = run (42 )
109139 diff_same_seed = np .abs (result1 - result2 ).sum ()
110140 print (f"Difference between two runs with same seed: { diff_same_seed :.6f} " )
111141 np .testing .assert_allclose (result1 , result2 , rtol = 1e-3 , atol = 1e-4 )
112- print ("✓ Same seed produces similar results (PASSED) " )
142+ print ("PASSED: Same seed produces similar results" )
113143
114144 # Assert different values when random_state is different
115145 print ("\n --- Testing variability with different seeds ---" )
@@ -118,7 +148,7 @@ def run(random_state):
118148 diff_diff_seed = np .abs (result3 - result4 ).sum ()
119149 print (f"Difference between runs with different seeds: { diff_diff_seed :.2f} " )
120150 assert diff_diff_seed > 1000 , f"Expected diff > 1000, got { diff_diff_seed } "
121- print ("✓ Different seeds produce different results (PASSED) " )
151+ print ("PASSED: Different seeds produce different results" )
122152
123153
124154def run_harmony (meta_tsv , pcs_tsv , harmonized_tsv , batch_var ):
@@ -130,27 +160,27 @@ def run_harmony(meta_tsv, pcs_tsv, harmonized_tsv, batch_var):
130160 return {"time" : 0 , "rss_delta_mb" : 0 }
131161
132162 # Load input data
133- meta_data = pd . read_csv (meta_tsv , sep = " \t " , low_memory = False )
134- data_mat = pd . read_csv (pcs_tsv , sep = " \t " , low_memory = False )
135- data_mat = data_mat . select_dtypes ( include = [ np . number ] )
163+ meta_header , meta_cols = read_tsv (meta_tsv )
164+ pcs_header , pcs_cols = read_tsv (pcs_tsv )
165+ data_mat = cols_to_matrix ( pcs_cols , pcs_header )
136166
167+ N = len (meta_cols [batch_var ])
168+ unique_batches = np .unique (meta_cols [batch_var ])
137169 print ("\n --- Input Data ---" )
138- print (f"data_mat shape: { data_mat .shape } (cells × PCs)" )
139- print (f"meta_data shape: { meta_data .shape } " )
140- print (f"meta_data columns: { list (meta_data .columns )[:10 ]} ..." )
141- print (f"Batch variable '{ batch_var } ' unique values: { meta_data [batch_var ].unique ()} " )
142- print (f"Cells per { batch_var } :\n { meta_data [batch_var ].value_counts ()} " )
170+ print (f"data_mat shape: { data_mat .shape } (cells x PCs)" )
171+ print (f"meta_data columns: { meta_header } " )
172+ print (f"Batch variable '{ batch_var } ' unique values: { unique_batches } " )
143173
144174 print ("\n --- Running Harmony ---" )
145- import gc , sys
175+ import gc
146176 gc .collect ()
147177 rss_before = _get_current_rss_mb ()
148178 start = time ()
149- ho = hm .run_harmony (data_mat , meta_data , [batch_var ])
179+ ho = hm .run_harmony (data_mat , meta_cols , [batch_var ])
150180 end = time ()
151181 rss_after = _get_current_rss_mb ()
152182 rss_delta_mb = rss_after - rss_before
153- print (f"\n ✓ Harmony completed in { end - start :.2f} seconds" )
183+ print (f"\n Harmony completed in { end - start :.2f} seconds" )
154184 print (f" RSS before: { rss_before :.1f} MB" )
155185 print (f" RSS after: { rss_after :.1f} MB" )
156186 print (f" RSS delta: { rss_delta_mb :.1f} MB" )
@@ -159,34 +189,32 @@ def run_harmony(meta_tsv, pcs_tsv, harmonized_tsv, batch_var):
159189 print (f"Number of clusters (K): { ho .K } " )
160190 print (f"Number of harmony iterations: { len (ho .objective_harmony )} " )
161191 print (f"K-means rounds per iteration: { ho .kmeans_rounds } " )
162- print (f"Z_corr shape: { ho .Z_corr .shape } (cells × PCs)" )
192+ print (f"Z_corr shape: { ho .Z_corr .shape } (cells x PCs)" )
163193 print (f"Z_orig shape: { ho .Z_orig .shape } " )
164194
165195 # Check convergence
166196 print ("\n --- Convergence ---" )
167197 print (f"Objective (harmony) history: { [f'{ x :.2f} ' for x in ho .objective_harmony ]} " )
168198
169- # Z_corr is now cells × PCs (same as input)
170- res = pd .DataFrame (ho .Z_corr )
171- res .columns = ['PC{}' .format (i + 1 ) for i in range (res .shape [1 ])]
172-
173199 # Compare to expected results from R
174- harm = pd .read_csv (harmonized_tsv , sep = "\t " )
175- harm = harm .select_dtypes (include = [np .number ])
200+ res = ho .Z_corr # cells x PCs
201+ harm_header , harm_cols = read_tsv (harmonized_tsv )
202+ harm = cols_to_matrix (harm_cols , harm_header )
176203 print ("\n --- Comparison with R Results ---" )
177204 print (f"Expected result shape: { harm .shape } " )
178205
206+ n_pcs = min (res .shape [1 ], harm .shape [1 ])
179207 cors_values = []
180- for i in range (res . shape [ 1 ] ):
181- cors_values .append (pearsonr (res . iloc [:, i ]. values , harm . iloc [:, i ]. values ))
208+ for i in range (n_pcs ):
209+ cors_values .append (pearsonr (res [:, i ], harm [:, i ]))
182210 print (f"Correlations (Python vs R) per PC: { [f'{ x :.3f} ' for x in cors_values ]} " )
183211 print (f"Min correlation: { min (cors_values ):.3f} " )
184212 print (f"Mean correlation: { np .mean (cors_values ):.3f} " )
185213
186214 # Correlation between test PCs and observed PCs is high
187215 assert np .all (np .array (cors_values ) >= 0.9 ), f"Some correlations < 0.9: { cors_values } "
188- print ("✓ All correlations >= 0.9 (PASSED) " )
189-
216+ print ("PASSED: All correlations >= 0.9" )
217+
190218 return {"time" : end - start , "rss_delta_mb" : rss_delta_mb }
191219
192220
@@ -211,34 +239,34 @@ def download_data():
211239 print ("# Running harmonypy tests" )
212240 print ("#" * 60 )
213241 print ()
214-
242+
215243 timings = {}
216244
217245 download_data ()
218-
246+
219247 timings ['small' ] = run_harmony (
220248 meta_tsv = "data/pbmc_3500_meta.tsv.gz" ,
221249 pcs_tsv = "data/pbmc_3500_pcs.tsv.gz" ,
222250 harmonized_tsv = "data/pbmc_3500_pcs_harmonized.tsv.gz" ,
223251 batch_var = "donor"
224252 )
225-
253+
226254 timings ['medium' ] = run_harmony (
227255 meta_tsv = "data/ircolitis_blood_cd8_obs.tsv.gz" ,
228256 pcs_tsv = "data/ircolitis_blood_cd8_pcs.tsv.gz" ,
229257 harmonized_tsv = "data/ircolitis_blood_cd8_pcs_harmonized.tsv.gz" ,
230258 batch_var = "batch"
231259 )
232-
260+
233261 timings ['large' ] = run_harmony (
234262 meta_tsv = "data/acute_myeloid_obs.tsv.gz" ,
235263 pcs_tsv = "data/acute_myeloid_pcs.tsv.gz" ,
236264 harmonized_tsv = "data/acute_myeloid_pcs_harmonized.tsv.gz" ,
237265 batch_var = "batch"
238266 )
239-
267+
240268 test_random_seed ()
241-
269+
242270 print ("\n " + "#" * 60 )
243271 print ("# Performance Summary" )
244272 print ("#" * 60 )
0 commit comments