#[path = "python_bootstrap.rs"]
mod python_bootstrap;
use python_bootstrap::ensure_python_packages_installed;
use efficient_pca::eigensnp::{
reorder_array_owned, reorder_columns_owned, EigenSNPCoreAlgorithm, EigenSNPCoreAlgorithmConfig,
EigenSNPCoreOutput, LdBlockSpecification, PcaReadyGenotypeAccessor, PcaSnpId, PcaSnpMetadata,
QcSampleId, ThreadSafeStdError,
};
use ndarray::{arr2, s, Array1, Array2, ArrayView1, ArrayView2, Axis}; use ndarray_rand::rand_distr::{Normal, StandardNormal, Uniform}; use ndarray_rand::RandomExt;
use rand::Rng; use rand::SeedableRng; use rand_chacha::ChaCha8Rng; use std::fmt::Write as FmtWrite;
use std::fs::{self, File}; use std::io::Write; use std::path::PathBuf;
use std::process::{Command, Stdio};
use std::str::FromStr; use std::path::Path; use lazy_static::lazy_static;
use std::fmt::Display;
use crate::eigensnp_integration_tests::generate_structured_data;
use crate::eigensnp_integration_tests::get_python_reference_pca;
use crate::eigensnp_integration_tests::TestDataAccessor;
use crate::eigensnp_integration_tests::TestResultRecord;
use crate::eigensnp_integration_tests::TEST_RESULTS;
const DEFAULT_FLOAT_TOLERANCE_F32: f32 = 1e-4; const DEFAULT_FLOAT_TOLERANCE_F64: f64 = 1e-4;
fn assert_f64_arrays_are_close(
arr1: ArrayView1<f64>, arr2: ArrayView1<f64>, tolerance: f64,
context: &str,
) {
assert_eq!(
arr1.dim(),
arr2.dim(),
"Array dimensions differ for {}. Left: {:?}, Right: {:?}",
context,
arr1.dim(),
arr2.dim()
);
for (i, val1) in arr1.iter().enumerate() {
let val2 = arr2[i]; assert!(
(val1 - val2).abs() < tolerance,
"Mismatch at index {} for {}: {} vs {} (diff: {})",
i,
context,
val1,
val2,
(val1 - val2).abs()
);
}
}
fn assert_f32_arrays_are_close_with_sign_flips(
arr1: ndarray::ArrayView2<f32>, arr2: ndarray::ArrayView2<f32>, tolerance: f32,
context: &str,
) {
assert_eq!(
arr1.dim(),
arr2.dim(),
"Array dimensions differ for {}. Left: {:?}, Right: {:?}",
context,
arr1.dim(),
arr2.dim()
);
if arr1.ncols() == 0 && arr2.ncols() == 0 {
return;
}
if arr1.ncols() == 0 || arr2.ncols() == 0 {
panic!("Array column count mismatch for {}: Left: {}, Right: {}. Both must be empty or non-empty.", context, arr1.ncols(), arr2.ncols());
}
for c_idx in 0..arr1.ncols() {
let col1 = arr1.column(c_idx); let col2 = arr2.column(c_idx);
let mut direct_match = true;
for r_idx in 0..col1.len() {
if (col1[r_idx] - col2[r_idx]).abs() >= tolerance {
direct_match = false;
break;
}
}
if direct_match {
continue;
}
let mut flipped_match = true;
for r_idx in 0..col1.len() {
if (col1[r_idx] - (-col2[r_idx])).abs() >= tolerance {
flipped_match = false;
break;
}
}
assert!(
flipped_match,
"Column {} mismatch for {} (even with sign flip check). Max diff: {}. First elements: {} vs {}",
c_idx, context,
col1.iter().zip(col2.iter()).map(|(a,b)| (a-b).abs().max((a-(-b)).abs())).fold(0.0f32, f32::max),
col1.get(0).unwrap_or(&0.0f32), col2.get(0).unwrap_or(&0.0f32)
);
}
}
fn standardize_features_across_samples(mut data: Array2<f32>) -> Array2<f32> {
if data.ncols() <= 1 {
if data.ncols() == 1 && data.nrows() > 0 {
data.fill(0.0);
}
return data;
}
for mut feature_row in data.axis_iter_mut(Axis(0)) {
let mean = feature_row.mean().unwrap_or(0.0);
feature_row.mapv_inplace(|x| x - mean);
let std_dev = feature_row.std(0.0);
if std_dev.abs() > 1e-7 {
feature_row.mapv_inplace(|x| x / std_dev);
} else {
feature_row.fill(0.0);
}
}
data
}
fn save_matrix_to_tsv<T: Display>(
matrix: &ArrayView2<T>,
dir_path: &str,
file_name: &str,
) -> Result<(), Box<dyn std::error::Error>> {
let full_path = Path::new(dir_path).join(file_name);
fs::create_dir_all(Path::new(dir_path))?; let mut file = File::create(full_path)?;
for row_idx in 0..matrix.nrows() {
for col_idx in 0..matrix.ncols() {
write!(file, "{}", matrix[[row_idx, col_idx]])?;
if col_idx < matrix.ncols() - 1 {
write!(file, " ")?; }
}
writeln!(file)?;
}
Ok(())
}
fn save_vector_to_tsv<T: Display>(
vector: &ArrayView1<T>,
dir_path: &str,
file_name: &str,
) -> Result<(), Box<dyn std::error::Error>> {
let full_path = Path::new(dir_path).join(file_name);
fs::create_dir_all(Path::new(dir_path))?; let mut file = File::create(full_path)?;
for i in 0..vector.len() {
writeln!(file, "{}", vector[i])?;
}
Ok(())
}
#[cfg(test)]
mod eigensnp_integration_tests {
use super::*;
use ctor::dtor;
use std::fs::OpenOptions;
use std::io::Write;
use std::path::Path; use std::sync::Mutex;
#[derive(Clone, Debug)] pub struct TestResultRecord {
pub test_name: String,
pub num_features_d: usize,
pub num_samples_n: usize,
pub num_pcs_requested_k: usize,
pub num_pcs_computed: usize,
pub success: bool,
pub outcome_details: String,
pub notes: String,
}
lazy_static! {
pub static ref TEST_RESULTS: Mutex<Vec<TestResultRecord>> = Mutex::new(Vec::new());
}
fn write_summary_file_impl() -> Result<(), std::io::Error> {
let results_guard = TEST_RESULTS.lock().unwrap();
if results_guard.is_empty() {
println!("[SUMMARY_WRITER] TEST_RESULTS is empty. No summary file will be written.");
return Ok(()); }
let artifact_dir = "target/test_artifacts";
let tsv_path = Path::new(artifact_dir).join("eigensnp_summary_results.tsv");
println!(
"[SUMMARY_WRITER] Attempting to write {} records to {:?}",
results_guard.len(),
tsv_path
);
if let Err(e) = std::fs::create_dir_all(artifact_dir) {
eprintln!(
"[SUMMARY_WRITER] Error creating artifact_dir '{}': {:?}",
artifact_dir, e
);
return Err(e);
}
let mut file = OpenOptions::new()
.write(true)
.create(true)
.truncate(true)
.open(&tsv_path)?;
writeln!(file, "TestName NumFeatures_D NumSamples_N NumPCsRequested_K NumPCsComputed Success OutcomeDetails Notes")?;
for record in results_guard.iter() {
writeln!(
file,
"{} {} {} {} {} {} {} {}",
record.test_name,
record.num_features_d,
record.num_samples_n,
record.num_pcs_requested_k,
record.num_pcs_computed,
record.success,
record.outcome_details.replace(" ", " ").replace("\n", "; "), record.notes.replace(" ", " ").replace("\n", "; ") )?;
}
println!(
"[SUMMARY_WRITER] Successfully wrote summary to {:?}",
tsv_path
);
Ok(())
}
#[dtor]
fn final_summary_writer() {
println!(
"[SUMMARY_WRITER_DTOR] Test execution finished. Running summary writer destructor."
);
if let Err(e) = write_summary_file_impl() {
eprintln!("[SUMMARY_WRITER_DTOR] CRITICAL: Failed to write eigensnp_summary_results.tsv: {:?}", e);
}
}
pub fn get_python_reference_pca(
standardized_data: &Array2<f32>,
k_components_to_request: usize,
artifact_dir_prefix: &str,
) -> Result<(Array2<f32>, Array2<f32>, Array1<f64>), Box<dyn std::error::Error>> {
let mut stdin_data = String::new();
for i in 0..standardized_data.nrows() {
for j in 0..standardized_data.ncols() {
stdin_data.push_str(&standardized_data[[i, j]].to_string());
if j < standardized_data.ncols() - 1 {
stdin_data.push(' ');
}
}
stdin_data.push('\n');
}
let mut script_path = PathBuf::from(env!("CARGO_MANIFEST_DIR"));
script_path.push("tests/pca.py");
ensure_python_packages_installed();
let mut process = Command::new("python3")
.arg(script_path.to_str().ok_or("Invalid script path")?)
.arg("--generate-reference-pca")
.arg("-k")
.arg(k_components_to_request.to_string())
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.spawn()?;
if let Some(mut stdin_pipe) = process.stdin.take() {
std::thread::spawn(move || {
if let Err(e) = stdin_pipe.write_all(stdin_data.as_bytes()) {
eprintln!("Failed to write to stdin of pca.py: {}", e); }
});
} else {
return Err(Box::new(std::io::Error::new(
std::io::ErrorKind::Other,
"Failed to open stdin pipe for pca.py",
)));
}
let py_cmd_output = process.wait_with_output()?;
let stdout_str = String::from_utf8_lossy(&py_cmd_output.stdout);
let stderr_str = String::from_utf8_lossy(&py_cmd_output.stderr);
if !py_cmd_output.status.success() {
let error_artifact_dir_name = format!("{}_py_error", artifact_dir_prefix);
let error_artifact_path =
Path::new("target/test_artifacts").join(error_artifact_dir_name);
fs::create_dir_all(&error_artifact_path)?;
let stdout_path = error_artifact_path.join("pca_stdout.txt");
let stderr_path = error_artifact_path.join("pca_stderr.txt");
fs::write(&stdout_path, stdout_str.as_bytes())?;
fs::write(&stderr_path, stderr_str.as_bytes())?;
return Err(format!(
"Python script pca.py failed with status {}. Stdout saved to '{}', Stderr saved to '{}'. Stderr Preview: {}",
py_cmd_output.status,
stdout_path.display(),
stderr_path.display(),
stderr_str.chars().take(500).collect::<String>() ).into());
}
parse_pca_py_output(&stdout_str).map_err(|e| {
format!(
"Failed to parse pca.py output: {}. Output:\n{}",
e, stdout_str
)
.into()
})
}
pub fn orthonormalize_columns(matrix: &mut Array2<f32>) {
if matrix.ncols() == 0 || matrix.nrows() == 0 {
return;
}
for j in 0..matrix.ncols() {
for i in 0..j {
let col_i_owned = matrix.column(i).to_owned();
let mut col_j_view = matrix.column_mut(j);
let dot_product = col_j_view.view().dot(&col_i_owned);
let scaled_col_i = col_i_owned.mapv(|x| x * dot_product);
col_j_view.zip_mut_with(&scaled_col_i, |cj_val, scaled_ci_val| {
*cj_val -= scaled_ci_val
});
}
let mut col_j_for_norm = matrix.column_mut(j);
let norm = col_j_for_norm.mapv(|x| x.powi(2)).sum().sqrt();
if norm > 1e-6 {
col_j_for_norm.mapv_inplace(|x| x / norm);
} else {
col_j_for_norm.fill(0.0);
}
}
}
pub fn generate_structured_data(
d_total_snps: usize, n_samples: usize, k_true_components: usize, signal_strength: f32, noise_std_dev: f32, seed: u64,
) -> Array2<f32> {
let mut rng = ChaCha8Rng::seed_from_u64(seed);
let k_eff = k_true_components.min(d_total_snps).min(n_samples);
let mut latent_components_n_x_k =
Array2::random_using((n_samples, k_eff), StandardNormal, &mut rng) * signal_strength;
orthonormalize_columns(&mut latent_components_n_x_k);
let mut loadings_d_x_k = Array2::zeros((d_total_snps, k_eff));
for i in 0..k_eff.min(d_total_snps) {
loadings_d_x_k[[i, i]] = 1.0;
}
if d_total_snps > k_eff {
for i in k_eff..d_total_snps {
for j in 0..k_eff {
if i % (j + 1) == 0 {
loadings_d_x_k[[i, j]] = rng.sample(Normal::new(0.0, 0.5).unwrap());
}
}
}
}
let signal_data_d_x_n = loadings_d_x_k.dot(&latent_components_n_x_k.t());
let noise_data_d_x_n = Array2::random_using(
(d_total_snps, n_samples),
Normal::new(0.0, noise_std_dev).unwrap(),
&mut rng,
);
let combined_data_d_x_n = signal_data_d_x_n + noise_data_d_x_n;
standardize_features_across_samples(combined_data_d_x_n)
}
#[derive(Clone)]
pub struct TestDataAccessor {
standardized_data: Array2<f32>,
}
impl TestDataAccessor {
pub fn new(standardized_data: Array2<f32>) -> Self {
Self { standardized_data }
}
pub fn new_empty(num_pca_snps: usize, num_qc_samples: usize) -> Self {
let standardized_data = Array2::zeros((num_pca_snps, num_qc_samples));
Self { standardized_data }
}
}
impl PcaReadyGenotypeAccessor for TestDataAccessor {
fn get_standardized_snp_sample_block(
&self,
snp_ids: &[PcaSnpId],
sample_ids: &[QcSampleId],
) -> Result<Array2<f32>, ThreadSafeStdError> {
if snp_ids.is_empty() {
return Ok(Array2::zeros((0, sample_ids.len())));
}
if sample_ids.is_empty() {
return Ok(Array2::zeros((snp_ids.len(), 0)));
}
if self.standardized_data.nrows() == 0 && !snp_ids.is_empty() {
return Err(Box::new(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
format!(
"Requested {} SNPs from an accessor with 0 SNPs.",
snp_ids.len()
),
)));
}
if self.standardized_data.ncols() == 0 && !sample_ids.is_empty() {
return Err(Box::new(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
format!(
"Requested {} samples from an accessor with 0 samples.",
sample_ids.len()
),
)));
}
let mut result_block = Array2::zeros((snp_ids.len(), sample_ids.len()));
for (i, pca_snp_id) in snp_ids.iter().enumerate() {
let target_row_idx = pca_snp_id.0;
if target_row_idx >= self.standardized_data.nrows() {
return Err(Box::new(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
format!(
"SNP ID PcaSnpId({}) out of bounds for {} SNPs",
target_row_idx,
self.standardized_data.nrows()
),
)));
}
for (j, qc_sample_id) in sample_ids.iter().enumerate() {
let target_col_idx = qc_sample_id.0;
if target_col_idx >= self.standardized_data.ncols() {
return Err(Box::new(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
format!(
"Sample ID QcSampleId({}) out of bounds for {} samples",
target_col_idx,
self.standardized_data.ncols()
),
)));
}
result_block[[i, j]] = self.standardized_data[[target_row_idx, target_col_idx]];
}
}
Ok(result_block)
}
fn num_pca_snps(&self) -> usize {
self.standardized_data.nrows()
}
fn num_qc_samples(&self) -> usize {
self.standardized_data.ncols()
}
}
fn parse_section<T: FromStr>(
lines: &mut std::iter::Peekable<std::str::Lines<'_>>,
expected_dim2: Option<usize>,
) -> Result<Array2<T>, String>
where
<T as FromStr>::Err: std::fmt::Debug,
{
let mut data_vec = Vec::new();
let mut current_dim2 = None;
loop {
match lines.peek() {
Some(line_peek) => {
if line_peek.is_empty()
|| line_peek.starts_with("LOADINGS:")
|| line_peek.starts_with("SCORES:")
|| line_peek.starts_with("EIGENVALUES:")
{
break;
}
let line = lines.next().unwrap();
let row: Vec<T> = line
.split_whitespace()
.map(|s| {
s.parse::<T>().map_err(|e| {
format!("Failed to parse value: {:?}, error: {:?}", s, e)
})
})
.collect::<Result<Vec<T>, String>>()?;
if let Some(d2) = current_dim2 {
if row.len() != d2 {
return Err(format!(
"Inconsistent row length. Expected {}, got {}",
d2,
row.len()
));
}
} else {
current_dim2 = Some(row.len());
if let Some(exp_d2) = expected_dim2 {
if !row.is_empty() && row.len() != exp_d2 {
return Err(format!(
"Unexpected row length for section. Expected {}, got {}",
exp_d2,
row.len()
));
}
}
}
data_vec.extend(row);
}
None => {
break;
}
}
}
let actual_dim2 = current_dim2.unwrap_or_else(|| expected_dim2.unwrap_or(0));
let num_rows = if actual_dim2 == 0 {
0
} else {
data_vec.len() / actual_dim2
};
Array2::from_shape_vec((num_rows, actual_dim2), data_vec)
.map_err(|e| format!("Failed to create Array2: {}", e))
}
pub fn parse_pca_py_output(
output_str: &str,
) -> Result<(Array2<f32>, Array2<f32>, Array1<f64>), String> {
let mut lines = output_str.lines().peekable();
let mut py_loadings: Option<Array2<f32>> = None;
let mut py_scores: Option<Array2<f32>> = None;
let mut py_eigenvalues: Option<Array1<f64>> = None;
while let Some(line_peek) = lines.peek() {
let current_line_is_empty = line_peek.is_empty();
if line_peek.starts_with("LOADINGS:") {
lines.next(); py_loadings = Some(parse_section(&mut lines, None)?);
} else if line_peek.starts_with("SCORES:") {
lines.next(); py_scores = Some(parse_section(&mut lines, None)?);
} else if line_peek.starts_with("EIGENVALUES:") {
lines.next(); let eig_array2 = parse_section(&mut lines, Some(1))?;
let eig_len = eig_array2.len();
py_eigenvalues = Some(
eig_array2
.into_shape_with_order((eig_len,))
.expect("Failed to reshape py_eigenvalues"),
);
} else if current_line_is_empty {
lines.next(); } else {
return Err(format!(
"Unexpected content in pca.py output. Line: '{}'",
line_peek
));
}
}
Ok((
py_loadings.ok_or_else(|| "LOADINGS section not found".to_string())?,
py_scores.ok_or_else(|| "SCORES section not found".to_string())?,
py_eigenvalues.ok_or_else(|| "EIGENVALUES section not found".to_string())?,
))
}
pub fn run_pc_scores_orthogonality_test(
test_name_str: &str,
num_snps: usize,
num_samples: usize,
num_pcs_target: usize,
seed: u64,
) {
let test_name = test_name_str.to_string();
let mut test_successful = true;
let mut outcome_details = String::new();
let notes = format!(
"Matrix: {}x{}, PCs: {}",
num_snps, num_samples, num_pcs_target
);
let mut max_off_diagonal_cov = 0.0f64;
let mut max_diag_eigenvalue_diff = 0.0f64;
let output_result_tuple = std::panic::catch_unwind(|| {
let mut rng = ChaCha8Rng::seed_from_u64(seed);
let raw_genos =
Array2::random_using((num_snps, num_samples), Uniform::new(0.0, 3.0), &mut rng);
let standardized_genos = standardize_features_across_samples(raw_genos);
let test_data = TestDataAccessor::new(standardized_genos);
let config = EigenSNPCoreAlgorithmConfig {
target_num_global_pcs: num_pcs_target,
subset_factor_for_local_basis_learning: 0.5, min_subset_size_for_local_basis_learning: (num_samples / 4)
.max(1)
.min(num_samples.max(1)),
max_subset_size_for_local_basis_learning: (num_samples / 2)
.max(10)
.min(num_samples.max(1)),
components_per_ld_block: 10
.min(num_snps.min((num_samples / 2).max(10).min(num_samples.max(1)))),
random_seed: seed,
..Default::default()
};
let algorithm = EigenSNPCoreAlgorithm::new(config);
let ld_blocks = vec![LdBlockSpecification {
user_defined_block_tag: "block1".to_string(),
pca_snp_ids_in_block: (0..num_snps).map(PcaSnpId).collect(),
}];
let snp_metadata = create_dummy_snp_metadata(num_snps);
algorithm.compute_pca(&test_data, &ld_blocks, &snp_metadata)
});
match output_result_tuple {
Ok(Ok((output, _))) => {
if output.num_principal_components_computed != num_pcs_target {
test_successful = false;
outcome_details.push_str(&format!(
"Did not compute target PCs. Expected: {}, Got: {}. ",
num_pcs_target, output.num_principal_components_computed
));
}
let scores = &output.final_sample_principal_component_scores;
if scores.nrows() != num_samples {
test_successful = false;
outcome_details.push_str(&format!(
"Scores nrows mismatch. Expected: {}, Got: {}. ",
num_samples,
scores.nrows()
));
}
if scores.ncols() != output.num_principal_components_computed {
test_successful = false;
outcome_details.push_str(&format!(
"Scores ncols mismatch. Expected: {}, Got: {}. ",
output.num_principal_components_computed,
scores.ncols()
));
}
if num_samples <= 1 || output.num_principal_components_computed == 0 {
outcome_details.push_str("Test condition (num_samples <=1 or num_pcs_computed == 0) means no further checks performed. ");
} else if test_successful {
let scores_f64 = scores.mapv(|x| x as f64);
let denominator = if output.num_qc_samples_used > 1 {
output.num_qc_samples_used as f64 - 1.0
} else {
1.0 };
if denominator == 0.0 {
test_successful = false;
outcome_details
.push_str("Denominator for covariance calculation is zero. ");
} else {
let covariance_matrix = scores_f64.t().dot(&scores_f64) / denominator;
let k_eff = output.num_principal_components_computed;
for r in 0..k_eff {
for c in 0..k_eff {
if r == c {
let diff = (covariance_matrix[[r, c]]
- output.final_principal_component_eigenvalues[r])
.abs();
if diff > max_diag_eigenvalue_diff {
max_diag_eigenvalue_diff = diff;
}
if diff >= DEFAULT_FLOAT_TOLERANCE_F64 * 100.0 {
test_successful = false;
outcome_details.push_str(&format!(
"Covariance diagonal [{},{}] {} does not match eigenvalue {}. Diff: {}. ",
r, c, covariance_matrix[[r, c]], output.final_principal_component_eigenvalues[r], diff
));
}
} else {
let off_diag_val = covariance_matrix[[r, c]].abs();
if off_diag_val > max_off_diagonal_cov {
max_off_diagonal_cov = off_diag_val;
}
if off_diag_val >= DEFAULT_FLOAT_TOLERANCE_F64 * 100.0 {
test_successful = false;
outcome_details.push_str(&format!(
"Covariance off-diagonal [{},{}] {} is not close to 0. Value: {}. ",
r, c, covariance_matrix[[r, c]], off_diag_val
));
}
}
}
}
if test_successful {
outcome_details.push_str(&format!("All orthogonality checks passed. Max off-diag cov: {:.2e}, Max diag-eigenvalue diff: {:.2e}. ", max_off_diagonal_cov, max_diag_eigenvalue_diff));
} else {
outcome_details.push_str(&format!("Orthogonality checks failed. Max off-diag cov: {:.2e}, Max diag-eigenvalue diff: {:.2e}. ", max_off_diagonal_cov, max_diag_eigenvalue_diff));
}
}
}
let record = TestResultRecord {
test_name: test_name.clone(),
num_features_d: num_snps,
num_samples_n: num_samples,
num_pcs_requested_k: num_pcs_target,
num_pcs_computed: output.num_principal_components_computed,
success: test_successful,
outcome_details: outcome_details.clone(),
notes,
};
eigensnp_integration_tests::TEST_RESULTS
.lock()
.unwrap()
.push(record);
}
Ok(Err(e)) => {
test_successful = false;
outcome_details = format!("PCA computation failed: {:?}", e);
let record = TestResultRecord {
test_name: test_name.clone(),
num_features_d: num_snps,
num_samples_n: num_samples,
num_pcs_requested_k: num_pcs_target,
num_pcs_computed: 0, success: test_successful,
outcome_details: outcome_details.clone(),
notes,
};
eigensnp_integration_tests::TEST_RESULTS
.lock()
.unwrap()
.push(record);
}
Err(e) => {
test_successful = false;
outcome_details = format!("Test panicked: {:?}", e);
let record = TestResultRecord {
test_name: test_name.clone(),
num_features_d: num_snps,
num_samples_n: num_samples,
num_pcs_requested_k: num_pcs_target,
num_pcs_computed: 0, success: test_successful,
outcome_details: outcome_details.clone(),
notes,
};
eigensnp_integration_tests::TEST_RESULTS
.lock()
.unwrap()
.push(record);
}
}
assert!(
test_successful,
"Test {} failed. Max off-diag: {:.2e}, Max diag-eig diff: {:.2e}. Details: {}",
test_name, max_off_diagonal_cov, max_diag_eigenvalue_diff, outcome_details
);
}
#[test]
fn test_pc_scores_orthogonality_large_500x100() {
run_pc_scores_orthogonality_test(
"test_pc_scores_orthogonality_large_500x100",
500,
100,
5,
123,
);
}
#[test]
fn test_pc_scores_orthogonality_large_1000x200() {
run_pc_scores_orthogonality_test(
"test_pc_scores_orthogonality_large_1000x200",
1000,
200,
10,
124,
);
}
pub fn run_snp_loadings_orthonormality_test(
test_name_str: &str,
num_snps: usize,
num_samples: usize,
num_pcs_target: usize,
seed: u64,
) {
let test_name = test_name_str.to_string();
let mut test_successful = true;
let mut outcome_details = String::new();
let notes = format!(
"Matrix: {}x{}, PCs: {}",
num_snps, num_samples, num_pcs_target
);
let mut max_deviation_from_identity = 0.0f32;
let output_result_tuple = std::panic::catch_unwind(|| {
let mut rng = ChaCha8Rng::seed_from_u64(seed);
let raw_genos =
Array2::random_using((num_snps, num_samples), Uniform::new(0.0, 3.0), &mut rng);
let standardized_genos = standardize_features_across_samples(raw_genos);
let test_data = TestDataAccessor::new(standardized_genos);
let config = EigenSNPCoreAlgorithmConfig {
target_num_global_pcs: num_pcs_target,
random_seed: seed,
subset_factor_for_local_basis_learning: 0.5,
min_subset_size_for_local_basis_learning: (num_samples / 4)
.max(1)
.min(num_samples.max(1)),
max_subset_size_for_local_basis_learning: (num_samples / 2)
.max(10)
.min(num_samples.max(1)),
components_per_ld_block: 10
.min(num_snps.min((num_samples / 2).max(10).min(num_samples.max(1)))),
..Default::default()
};
let algorithm = EigenSNPCoreAlgorithm::new(config);
let ld_blocks = vec![LdBlockSpecification {
user_defined_block_tag: "block1".to_string(),
pca_snp_ids_in_block: (0..num_snps).map(PcaSnpId).collect(),
}];
let snp_metadata = create_dummy_snp_metadata(num_snps);
algorithm.compute_pca(&test_data, &ld_blocks, &snp_metadata)
});
match output_result_tuple {
Ok(Ok((output, _))) => {
if output.num_principal_components_computed != num_pcs_target {
test_successful = false;
outcome_details.push_str(&format!(
"Did not compute target PCs. Expected: {}, Got: {}. ",
num_pcs_target, output.num_principal_components_computed
));
}
let loadings = &output.final_snp_principal_component_loadings;
if loadings.nrows() != num_snps {
test_successful = false;
outcome_details.push_str(&format!(
"Loadings nrows mismatch. Expected: {}, Got: {}. ",
num_snps,
loadings.nrows()
));
}
if loadings.ncols() != output.num_principal_components_computed {
test_successful = false;
outcome_details.push_str(&format!(
"Loadings ncols mismatch. Expected: {}, Got: {}. ",
output.num_principal_components_computed,
loadings.ncols()
));
}
if output.num_principal_components_computed == 0 {
outcome_details.push_str("No PCs computed, skipping orthonormality check. ");
} else if test_successful {
let check_identity = loadings.t().dot(loadings);
let k_eff = output.num_principal_components_computed;
for r_idx in 0..k_eff {
for c_idx in 0..k_eff {
let expected_val = if r_idx == c_idx { 1.0 } else { 0.0 };
let deviation = (check_identity[[r_idx, c_idx]] - expected_val).abs();
if deviation > max_deviation_from_identity {
max_deviation_from_identity = deviation;
}
if deviation >= DEFAULT_FLOAT_TOLERANCE_F32 {
test_successful = false;
outcome_details.push_str(&format!(
"Loadings orthonormality check: Identity matrix mismatch at [{},{}]. Expected {}, Got {}. Diff: {}. ",
r_idx, c_idx, expected_val, check_identity[[r_idx, c_idx]], deviation
));
}
}
}
if test_successful {
outcome_details.push_str(&format!("All orthonormality checks passed. Max deviation from identity: {:.2e}. ", max_deviation_from_identity));
} else {
outcome_details.push_str(&format!(
"Orthonormality checks failed. Max deviation from identity: {:.2e}. ",
max_deviation_from_identity
));
}
}
let record = eigensnp_integration_tests::TestResultRecord {
test_name: test_name.clone(),
num_features_d: num_snps,
num_samples_n: num_samples,
num_pcs_requested_k: num_pcs_target,
num_pcs_computed: output.num_principal_components_computed,
success: test_successful,
outcome_details: outcome_details.clone(),
notes,
};
TEST_RESULTS.lock().unwrap().push(record);
}
Ok(Err(e)) => {
test_successful = false;
outcome_details = format!("PCA computation failed: {:?}", e);
let record = TestResultRecord {
test_name: test_name.clone(),
num_features_d: num_snps,
num_samples_n: num_samples,
num_pcs_requested_k: num_pcs_target,
num_pcs_computed: 0,
success: test_successful,
outcome_details: outcome_details.clone(),
notes,
};
TEST_RESULTS.lock().unwrap().push(record);
}
Err(e) => {
test_successful = false;
outcome_details = format!("Test panicked: {:?}", e);
let record = TestResultRecord {
test_name: test_name.clone(),
num_features_d: num_snps,
num_samples_n: num_samples,
num_pcs_requested_k: num_pcs_target,
num_pcs_computed: 0,
success: test_successful,
outcome_details: outcome_details.clone(),
notes,
};
TEST_RESULTS.lock().unwrap().push(record);
}
}
assert!(
test_successful,
"Test {} failed. Max deviation from identity: {:.2e}. Details: {}",
test_name, max_deviation_from_identity, outcome_details
);
}
#[test]
fn test_snp_loadings_orthonormality_large_500x100() {
run_snp_loadings_orthonormality_test(
"test_snp_loadings_orthonormality_large_500x100",
500,
100,
5,
456,
);
}
#[test]
fn test_snp_loadings_orthonormality_large_1000x200() {
run_snp_loadings_orthonormality_test(
"test_snp_loadings_orthonormality_large_1000x200",
1000,
200,
10,
457,
);
}
pub fn run_eigenvalue_score_variance_correspondence_test(
test_name_str: &str,
num_snps: usize,
num_samples: usize,
num_pcs_target: usize,
seed: u64,
) {
let test_name = test_name_str.to_string();
let mut test_successful = true;
let mut outcome_details = String::new();
let notes = format!(
"Matrix: {}x{}, PCs: {}",
num_snps, num_samples, num_pcs_target
);
let mut max_variance_eigenvalue_diff = 0.0f64;
let output_result_tuple = std::panic::catch_unwind(|| {
let mut rng = ChaCha8Rng::seed_from_u64(seed);
let raw_genos =
Array2::random_using((num_snps, num_samples), Uniform::new(0.0, 3.0), &mut rng);
let standardized_genos = standardize_features_across_samples(raw_genos);
let test_data = TestDataAccessor::new(standardized_genos);
let config = EigenSNPCoreAlgorithmConfig {
target_num_global_pcs: num_pcs_target,
random_seed: seed,
subset_factor_for_local_basis_learning: 0.5,
min_subset_size_for_local_basis_learning: (num_samples / 4)
.max(1)
.min(num_samples.max(1)),
max_subset_size_for_local_basis_learning: (num_samples / 2)
.max(10)
.min(num_samples.max(1)),
components_per_ld_block: 10
.min(num_snps.min((num_samples / 2).max(10).min(num_samples.max(1)))),
..Default::default()
};
let algorithm = EigenSNPCoreAlgorithm::new(config);
let ld_blocks = vec![LdBlockSpecification {
user_defined_block_tag: "block1".to_string(),
pca_snp_ids_in_block: (0..num_snps).map(PcaSnpId).collect(),
}];
let snp_metadata = create_dummy_snp_metadata(num_snps);
algorithm.compute_pca(&test_data, &ld_blocks, &snp_metadata)
});
match output_result_tuple {
Ok(Ok((output, _))) => {
if output.num_principal_components_computed != num_pcs_target {
test_successful = false;
outcome_details.push_str(&format!(
"Did not compute target PCs. Expected: {}, Got: {}. ",
num_pcs_target, output.num_principal_components_computed
));
}
if num_samples <= 1 || output.num_principal_components_computed == 0 {
outcome_details.push_str("Test condition (num_samples <=1 or num_pcs_computed == 0) means no further checks performed. ");
} else if test_successful {
let scores = &output.final_sample_principal_component_scores;
let eigenvalues = &output.final_principal_component_eigenvalues;
let k_eff = output.num_principal_components_computed;
let denominator = if output.num_qc_samples_used > 1 {
output.num_qc_samples_used as f64 - 1.0
} else {
1.0 };
if denominator == 0.0 {
test_successful = false;
outcome_details.push_str("Denominator for variance calculation is zero. ");
} else {
for k_idx in 0..k_eff {
let score_column_k = scores.column(k_idx);
let sum_sq_f64 = score_column_k
.iter()
.map(|&x| (x as f64).powi(2))
.sum::<f64>();
let variance_of_score_k = sum_sq_f64 / denominator;
let diff = (variance_of_score_k - eigenvalues[k_idx]).abs();
if diff > max_variance_eigenvalue_diff {
max_variance_eigenvalue_diff = diff;
}
if diff >= DEFAULT_FLOAT_TOLERANCE_F64 * 100.0 {
test_successful = false;
outcome_details.push_str(&format!(
"Variance of score column {} ({}) does not match eigenvalue {} ({}). Diff: {}. ",
k_idx, variance_of_score_k, k_idx, eigenvalues[k_idx], diff
));
}
}
if test_successful {
outcome_details.push_str(&format!(
"All variance-eigenvalue checks passed. Max diff: {:.2e}. ",
max_variance_eigenvalue_diff
));
} else {
outcome_details.push_str(&format!(
"Variance-eigenvalue checks failed. Max diff: {:.2e}. ",
max_variance_eigenvalue_diff
));
}
}
}
let record = eigensnp_integration_tests::TestResultRecord {
test_name: test_name.clone(),
num_features_d: num_snps,
num_samples_n: num_samples,
num_pcs_requested_k: num_pcs_target,
num_pcs_computed: output.num_principal_components_computed,
success: test_successful,
outcome_details: outcome_details.clone(),
notes,
};
TEST_RESULTS.lock().unwrap().push(record);
}
Ok(Err(e)) => {
test_successful = false;
outcome_details = format!("PCA computation failed: {:?}", e);
let record = TestResultRecord {
test_name: test_name.clone(),
num_features_d: num_snps,
num_samples_n: num_samples,
num_pcs_requested_k: num_pcs_target,
num_pcs_computed: 0,
success: test_successful,
outcome_details: outcome_details.clone(),
notes,
};
TEST_RESULTS.lock().unwrap().push(record);
}
Err(e) => {
test_successful = false;
outcome_details = format!("Test panicked: {:?}", e);
let record = TestResultRecord {
test_name: test_name.clone(),
num_features_d: num_snps,
num_samples_n: num_samples,
num_pcs_requested_k: num_pcs_target,
num_pcs_computed: 0,
success: test_successful,
outcome_details: outcome_details.clone(),
notes,
};
TEST_RESULTS.lock().unwrap().push(record);
}
}
assert!(
test_successful,
"Test {} failed. Max variance-eigenvalue diff: {:.2e}. Details: {}",
test_name, max_variance_eigenvalue_diff, outcome_details
);
}
#[test]
fn test_eigenvalue_score_variance_correspondence_large_500x100() {
run_eigenvalue_score_variance_correspondence_test(
"test_eigenvalue_score_variance_correspondence_large_500x100",
500,
100,
5,
789,
);
}
#[test]
fn test_eigenvalue_score_variance_correspondence_large_1000x200() {
run_eigenvalue_score_variance_correspondence_test(
"test_eigenvalue_score_variance_correspondence_large_1000x200",
1000,
200,
10,
790,
);
}
#[test]
fn test_pca_zero_snps() {
let num_samples = 10;
let num_snps = 0; let k_requested = 2;
let test_data = TestDataAccessor::new_empty(num_snps, num_samples);
let config = EigenSNPCoreAlgorithmConfig {
target_num_global_pcs: k_requested,
..Default::default()
};
let algorithm = EigenSNPCoreAlgorithm::new(config);
let ld_blocks = vec![];
let snp_metadata = create_dummy_snp_metadata(num_snps);
let (output, _) = algorithm
.compute_pca(&test_data, &ld_blocks, &snp_metadata)
.expect("PCA with 0 SNPs failed");
assert_eq!(output.num_pca_snps_used, 0);
assert_eq!(output.num_qc_samples_used, num_samples);
assert_eq!(output.num_principal_components_computed, 0);
assert_eq!(output.final_snp_principal_component_loadings.nrows(), 0);
assert_eq!(output.final_snp_principal_component_loadings.ncols(), 0);
assert_eq!(
output.final_sample_principal_component_scores.nrows(),
num_samples
);
assert_eq!(output.final_sample_principal_component_scores.ncols(), 0);
assert_eq!(output.final_principal_component_eigenvalues.len(), 0);
let record = TestResultRecord {
test_name: "test_pca_zero_snps".to_string(),
num_features_d: num_snps,
num_samples_n: num_samples,
num_pcs_requested_k: k_requested,
num_pcs_computed: output.num_principal_components_computed,
success: true, outcome_details: "Validated behavior with zero SNPs. All assertions passed."
.to_string(),
notes: "Edge case test.".to_string(),
};
TEST_RESULTS.lock().unwrap().push(record);
}
#[test]
fn test_pca_zero_samples() {
let num_snps = 20;
let num_samples = 0; let k_requested = 2;
let test_data = TestDataAccessor::new_empty(num_snps, num_samples);
let config = EigenSNPCoreAlgorithmConfig {
target_num_global_pcs: k_requested,
..Default::default()
};
let algorithm = EigenSNPCoreAlgorithm::new(config);
let ld_blocks = vec![LdBlockSpecification {
user_defined_block_tag: "block1".to_string(),
pca_snp_ids_in_block: (0..num_snps).map(PcaSnpId).collect(),
}];
let snp_metadata = create_dummy_snp_metadata(num_snps);
let (output, _) = algorithm
.compute_pca(&test_data, &ld_blocks, &snp_metadata)
.expect("PCA with 0 samples failed");
assert_eq!(output.num_qc_samples_used, 0);
assert_eq!(output.num_pca_snps_used, num_snps);
assert_eq!(output.num_principal_components_computed, 0);
assert_eq!(
output.final_snp_principal_component_loadings.nrows(),
num_snps
);
assert_eq!(output.final_snp_principal_component_loadings.ncols(), 0);
assert_eq!(output.final_sample_principal_component_scores.nrows(), 0);
assert_eq!(output.final_sample_principal_component_scores.ncols(), 0);
assert_eq!(output.final_principal_component_eigenvalues.len(), 0);
let record = TestResultRecord {
test_name: "test_pca_zero_samples".to_string(),
num_features_d: num_snps,
num_samples_n: num_samples,
num_pcs_requested_k: k_requested,
num_pcs_computed: output.num_principal_components_computed,
success: true, outcome_details: "Validated behavior with zero samples. All assertions passed."
.to_string(),
notes: "Edge case test.".to_string(),
};
TEST_RESULTS.lock().unwrap().push(record);
}
#[test]
fn test_pca_more_components_requested_than_rank_d_gt_n() {
let mut overall_test_successful = true;
let mut outcome_details = String::new();
let mut notes = String::new();
notes.push_str("Testing k_requested > true rank, with D > N (150x50). ");
let mut rust_pca_computation_phase_ok = true;
let mut rust_eigenvalue_check_phase_ok = true;
let mut python_eigenvalue_check_phase_ok = true;
let mut loadings_comparison_phase_ok = true;
let mut scores_comparison_phase_ok = true;
let mut eigenvalues_comparison_phase_ok = true;
let num_samples = 50; let num_true_rank_snps = 20; let num_total_snps = 150; let k_components_requested = 30;
let mut raw_genos = Array2::<f32>::zeros((num_total_snps, num_samples));
let mut rng = ChaCha8Rng::seed_from_u64(321);
for r in 0..num_true_rank_snps {
for c in 0..num_samples {
raw_genos[[r, c]] = rng.sample(Uniform::new(0.0, 3.0));
}
}
for r in num_true_rank_snps..num_total_snps {
let source_row_idx = r % num_true_rank_snps; let factor = rng.sample(Uniform::new(0.3, 0.7));
let noise: f32 = rng.sample(Uniform::new(-0.01, 0.01)); for c in 0..num_samples {
raw_genos[[r, c]] = raw_genos[[source_row_idx, c]] * factor + noise;
}
}
let standardized_genos = standardize_features_across_samples(raw_genos.clone());
let test_data = TestDataAccessor::new(standardized_genos.clone());
let config = EigenSNPCoreAlgorithmConfig {
target_num_global_pcs: k_components_requested,
components_per_ld_block: num_total_snps.min(num_samples),
random_seed: 321,
subset_factor_for_local_basis_learning: 1.0,
min_subset_size_for_local_basis_learning: num_samples.max(1), max_subset_size_for_local_basis_learning: num_samples.max(10), ..Default::default()
};
let algorithm = EigenSNPCoreAlgorithm::new(config);
let ld_blocks = vec![LdBlockSpecification {
user_defined_block_tag: "block1".to_string(),
pca_snp_ids_in_block: (0..num_total_snps).map(PcaSnpId).collect(),
}];
let snp_metadata = create_dummy_snp_metadata(num_total_snps);
let rust_output_result_tuple = algorithm.compute_pca(&test_data, &ld_blocks, &snp_metadata);
let rust_output = match rust_output_result_tuple {
Ok((output, _)) => {
outcome_details.push_str("Rust PCA computation: SUCCESS. ");
output
}
Err(e) => {
rust_pca_computation_phase_ok = false;
overall_test_successful = false;
outcome_details.push_str(&format!("Rust PCA computation: FAILED. Error: {}. ", e));
EigenSNPCoreOutput {
final_snp_principal_component_loadings: Array2::zeros((0, 0)),
final_sample_principal_component_scores: Array2::zeros((0, 0)),
final_principal_component_eigenvalues: Array1::zeros(0),
num_principal_components_computed: 0,
num_pca_snps_used: num_total_snps, num_qc_samples_used: num_samples, }
}
};
let effective_k_rust = rust_output.num_principal_components_computed;
let mut py_loadings_k_x_d: Array2<f32> = Array2::zeros((0, 0)); let mut py_scores_n_x_k: Array2<f32> = Array2::zeros((0, 0));
let mut py_eigenvalues_k: Array1<f64> = Array1::zeros(0);
let mut effective_k_py = 0;
let mut python_script_phase_ok = true;
match get_python_reference_pca(
&standardized_genos,
k_components_requested,
"pca_low_rank_D_gt_N_py_ref",
) {
Ok((loadings_from_py, scores_from_py, eigenvalues_from_py)) => {
py_loadings_k_x_d = loadings_from_py; py_scores_n_x_k = scores_from_py;
py_eigenvalues_k = eigenvalues_from_py;
effective_k_py = py_eigenvalues_k.len();
outcome_details.push_str(&format!(
"Python script execution: SUCCESS. Rust k_eff: {}, Py k_eff: {}. ",
effective_k_rust, effective_k_py
));
}
Err(e) => {
python_script_phase_ok = false;
outcome_details
.push_str(&format!("Python script execution: FAILED. Error: {}. ", e));
}
}
overall_test_successful &= python_script_phase_ok;
if overall_test_successful && rust_pca_computation_phase_ok {
if effective_k_rust > num_true_rank_snps {
let mut all_rust_eigenvalues_small = true;
for i in num_true_rank_snps..effective_k_rust {
if rust_output.final_principal_component_eigenvalues[i] > 1e-3 {
all_rust_eigenvalues_small = false;
rust_eigenvalue_check_phase_ok = false;
outcome_details.push_str(&format!("Rust Eigenvalue Check: FAILED. PC {} ({}) beyond true rank ({}) is too large ({}). ",
i, rust_output.final_principal_component_eigenvalues[i], num_true_rank_snps, 1e-3));
break;
}
}
if all_rust_eigenvalues_small {
outcome_details.push_str(
"Rust Eigenvalue Check: SUCCESS (eigenvalues beyond true rank are small). ",
);
}
} else {
outcome_details.push_str(
"Rust Eigenvalue Check: SKIPPED (k_eff_rust <= num_true_rank_snps). ",
);
}
} else if rust_pca_computation_phase_ok {
outcome_details.push_str("Rust Eigenvalue Check: SKIPPED (prior failure). ");
}
overall_test_successful &= rust_eigenvalue_check_phase_ok;
if overall_test_successful && python_script_phase_ok {
if effective_k_py > 0 && effective_k_py > num_true_rank_snps {
let mut all_py_eigenvalues_small = true;
for i in num_true_rank_snps..effective_k_py {
if py_eigenvalues_k.get(i).map_or(false, |&val| val > 1e-3) {
all_py_eigenvalues_small = false;
python_eigenvalue_check_phase_ok = false;
outcome_details.push_str(&format!("Python Eigenvalue Check: FAILED. PC {} ({}) beyond true rank ({}) is too large ({}). ",
i, py_eigenvalues_k.get(i).unwrap_or(&0.0), num_true_rank_snps, 1e-3));
break;
}
}
if all_py_eigenvalues_small {
outcome_details.push_str("Python Eigenvalue Check: SUCCESS (eigenvalues beyond true rank are small). ");
}
} else {
outcome_details.push_str("Python Eigenvalue Check: SKIPPED (k_eff_py <= num_true_rank_snps or k_eff_py is 0). ");
}
} else if python_script_phase_ok {
outcome_details.push_str("Python Eigenvalue Check: SKIPPED (prior failure). ");
}
overall_test_successful &= python_eigenvalue_check_phase_ok;
let py_loadings_d_x_k = py_loadings_k_x_d.t().into_owned();
let artifact_dir = "target/test_artifacts/pca_low_rank_D_gt_N";
if rust_output.num_principal_components_computed > 0 {
save_matrix_to_tsv(
&rust_output.final_snp_principal_component_loadings.view(),
artifact_dir,
"rust_loadings.tsv",
)
.expect("Failed to save rust_loadings.tsv");
save_matrix_to_tsv(
&rust_output.final_sample_principal_component_scores.view(),
artifact_dir,
"rust_scores.tsv",
)
.expect("Failed to save rust_scores.tsv");
save_vector_to_tsv(
&rust_output.final_principal_component_eigenvalues.view(),
artifact_dir,
"rust_eigenvalues.tsv",
)
.expect("Failed to save rust_eigenvalues.tsv");
}
if effective_k_py > 0 {
save_matrix_to_tsv(
&py_loadings_d_x_k.view(),
artifact_dir,
"python_loadings.tsv",
)
.expect("Failed to save python_loadings.tsv");
save_matrix_to_tsv(&py_scores_n_x_k.view(), artifact_dir, "python_scores.tsv")
.expect("Failed to save python_scores.tsv");
save_vector_to_tsv(
&py_eigenvalues_k.view(),
artifact_dir,
"python_eigenvalues.tsv",
)
.expect("Failed to save python_eigenvalues.tsv");
}
let record = TestResultRecord {
test_name: "test_pca_more_components_requested_than_rank_D_gt_N".to_string(),
num_features_d: num_total_snps,
num_samples_n: num_samples,
num_pcs_requested_k: k_components_requested,
num_pcs_computed: effective_k_rust,
success: overall_test_successful, outcome_details: outcome_details.clone(),
notes,
};
TEST_RESULTS.lock().unwrap().push(record);
if overall_test_successful {
let num_pcs_to_compare = effective_k_rust
.min(effective_k_py)
.min(num_true_rank_snps + 2)
.min(k_components_requested);
if num_pcs_to_compare > 0 {
let loadings_comparison_result = std::panic::catch_unwind(|| {
assert_f32_arrays_are_close_with_sign_flips(
rust_output
.final_snp_principal_component_loadings
.slice(s![.., 0..num_pcs_to_compare]),
py_loadings_d_x_k.slice(s![.., 0..num_pcs_to_compare]),
1.5f32,
"SNP Loadings (Low-Rank D>N)",
);
});
if loadings_comparison_result.is_err() {
loadings_comparison_phase_ok = false;
outcome_details
.push_str("Loadings Comparison: FAILED. Details captured by assert. ");
} else {
outcome_details.push_str("Loadings Comparison: SUCCESS. ");
}
overall_test_successful &= loadings_comparison_phase_ok;
if overall_test_successful {
let scores_comparison_result = std::panic::catch_unwind(|| {
assert_f32_arrays_are_close_with_sign_flips(
rust_output
.final_sample_principal_component_scores
.slice(s![.., 0..num_pcs_to_compare]),
py_scores_n_x_k.slice(s![.., 0..num_pcs_to_compare]),
DEFAULT_FLOAT_TOLERANCE_F32 * 10.0,
"Sample Scores (Low-Rank D>N)",
);
});
if scores_comparison_result.is_err() {
scores_comparison_phase_ok = false;
outcome_details
.push_str("Scores Comparison: FAILED. Details captured by assert. ");
} else {
outcome_details.push_str("Scores Comparison: SUCCESS. ");
}
overall_test_successful &= scores_comparison_phase_ok;
}
if overall_test_successful {
let eigenvalues_comparison_result = std::panic::catch_unwind(|| {
assert_f64_arrays_are_close(
rust_output
.final_principal_component_eigenvalues
.slice(s![0..num_pcs_to_compare]),
py_eigenvalues_k.slice(s![0..num_pcs_to_compare]),
DEFAULT_FLOAT_TOLERANCE_F64 * 10.0,
"Eigenvalues (Low-Rank D>N)",
);
});
if eigenvalues_comparison_result.is_err() {
eigenvalues_comparison_phase_ok = false;
outcome_details.push_str(
"Eigenvalues Comparison: FAILED. Details captured by assert. ",
);
} else {
outcome_details.push_str("Eigenvalues Comparison: SUCCESS. ");
}
overall_test_successful &= eigenvalues_comparison_phase_ok;
}
} else {
outcome_details
.push_str("Detailed Comparisons: SKIPPED (num_pcs_to_compare is 0). ");
if effective_k_rust != effective_k_py {
overall_test_successful = false; outcome_details.push_str(&format!(
"Effective k mismatch (Rust: {}, Py: {}), leading to 0 comparable PCs. ",
effective_k_rust, effective_k_py
));
}
}
if !overall_test_successful {
TEST_RESULTS
.lock()
.unwrap()
.last_mut()
.map(|rec| rec.success = false);
TEST_RESULTS
.lock()
.unwrap()
.last_mut()
.map(|rec| rec.outcome_details = outcome_details.clone());
}
} else {
outcome_details
.push_str("Detailed Comparisons: SKIPPED (due to prior phase failures). ");
TEST_RESULTS
.lock()
.unwrap()
.last_mut()
.map(|rec| rec.outcome_details = outcome_details.clone());
}
assert!(
overall_test_successful,
"Test 'test_pca_more_components_requested_than_rank_D_gt_N' failed. Details: {}",
outcome_details
);
}
}
pub fn pearson_correlation(v1: ArrayView1<f32>, v2: ArrayView1<f32>) -> Option<f32> {
if v1.len() != v2.len() || v1.is_empty() {
return None;
}
let _n = v1.len() as f32;
let mean1 = v1.mean().unwrap_or(0.0);
let mean2 = v2.mean().unwrap_or(0.0);
let mut cov = 0.0;
let mut std_dev1_sq = 0.0;
let mut std_dev2_sq = 0.0;
for i in 0..v1.len() {
let d1 = v1[i] - mean1;
let d2 = v2[i] - mean2;
cov += d1 * d2;
std_dev1_sq += d1 * d1;
std_dev2_sq += d2 * d2;
}
if std_dev1_sq <= 1e-6 || std_dev2_sq <= 1e-6 {
if (std_dev1_sq - std_dev2_sq).abs() < 1e-6 && cov.abs() < 1e-6 {
let mut all_v1_same = true;
let mut all_v2_same = true;
if v1.len() > 1 {
for i in 1..v1.len() {
if (v1[i] - v1[0]).abs() > 1e-6 {
all_v1_same = false;
break;
}
if (v2[i] - v2[0]).abs() > 1e-6 {
all_v2_same = false;
break;
}
}
}
if all_v1_same && all_v2_same && (v1[0] - v2[0]).abs() < 1e-6 {
return Some(1.0);
}
}
return None;
}
Some(cov / (std_dev1_sq.sqrt() * std_dev2_sq.sqrt()))
}
pub fn run_pc_correlation_with_truth_set_test(
test_name_str: &str,
num_snps: usize, num_samples: usize, k_components: usize, seed: u64,
) {
let test_name = test_name_str.to_string();
let mut test_successful = true;
let mut outcome_details = String::new();
let mut notes = format!(
"Matrix D_snps x N_samples: {}x{}, k_requested: {}. ",
num_snps, num_samples, k_components
);
let mut rng = ChaCha8Rng::seed_from_u64(seed);
let raw_genos = Array2::random_using((num_snps, num_samples), Uniform::new(0.0, 3.0), &mut rng);
let standardized_genos_snps_x_samples = standardize_features_across_samples(raw_genos.clone());
let artifact_dir_suffix = format!(
"pc_correlation_{}x{}_k{}",
num_snps, num_samples, k_components
);
let artifact_dir = Path::new("target/test_artifacts").join(artifact_dir_suffix);
if let Err(e) = fs::create_dir_all(&artifact_dir) {
notes.push_str(&format!("Failed to create artifact dir: {}. ", e));
}
let mut py_loadings_d_x_k: Array2<f32> = Array2::zeros((0, 0)); let mut py_scores_n_x_k: Array2<f32> = Array2::zeros((0, 0)); let mut effective_k_py = 0;
let python_pca_result = get_python_reference_pca(
&standardized_genos_snps_x_samples,
k_components,
&format!(
"pc_correlation_{}x{}_k{}_py_ref",
num_snps, num_samples, k_components
),
);
match python_pca_result {
Ok((loadings_k_x_d_from_py, scores_from_py, _eigenvalues_from_py)) => {
py_loadings_d_x_k = loadings_k_x_d_from_py.t().into_owned();
py_scores_n_x_k = scores_from_py;
effective_k_py = py_loadings_d_x_k.ncols();
save_matrix_to_tsv(
&py_loadings_d_x_k.view(),
artifact_dir.to_str().unwrap_or("."),
"python_loadings.tsv",
)
.unwrap_or_default();
save_matrix_to_tsv(
&py_scores_n_x_k.view(),
artifact_dir.to_str().unwrap_or("."),
"python_scores.tsv",
)
.unwrap_or_default();
outcome_details.push_str(&format!(
"Python PCA successful. effective_k_py: {}. ",
effective_k_py
));
}
Err(e) => {
test_successful = false;
outcome_details.push_str(&format!("Python reference PCA failed: {}. ", e));
}
}
let test_data_accessor = TestDataAccessor::new(standardized_genos_snps_x_samples.clone());
let config = EigenSNPCoreAlgorithmConfig {
target_num_global_pcs: k_components,
random_seed: seed,
subset_factor_for_local_basis_learning: 0.5,
min_subset_size_for_local_basis_learning: (num_samples / 4).max(1).min(num_samples.max(1)),
max_subset_size_for_local_basis_learning: (num_samples / 2).max(10).min(num_samples.max(1)),
components_per_ld_block: 10
.min(num_snps.min((num_samples / 2).max(10).min(num_samples.max(1)))),
..Default::default()
};
let algorithm = EigenSNPCoreAlgorithm::new(config);
let ld_blocks = vec![LdBlockSpecification {
user_defined_block_tag: "block1".to_string(),
pca_snp_ids_in_block: (0..num_snps).map(PcaSnpId).collect(),
}];
let mut rust_pcs_computed = 0;
let snp_metadata = create_dummy_snp_metadata(num_snps);
match algorithm.compute_pca(&test_data_accessor, &ld_blocks, &snp_metadata) {
Ok((rust_result, _)) => {
rust_pcs_computed = rust_result.num_principal_components_computed;
save_matrix_to_tsv(
&rust_result.final_snp_principal_component_loadings.view(),
artifact_dir.to_str().unwrap_or("."),
"rust_loadings.tsv",
)
.unwrap_or_default();
save_matrix_to_tsv(
&rust_result.final_sample_principal_component_scores.view(),
artifact_dir.to_str().unwrap_or("."),
"rust_scores.tsv",
)
.unwrap_or_default();
if test_successful {
let k_to_compare = rust_result
.num_principal_components_computed
.min(effective_k_py);
if k_to_compare == 0 {
outcome_details
.push_str("No components to compare (Rust or Python computed 0 PCs). ");
if rust_result.num_principal_components_computed != effective_k_py {
test_successful = false;
outcome_details.push_str(&format!(
"Mismatch in computed k (Rust: {}, Py: {}). ",
rust_result.num_principal_components_computed, effective_k_py
));
}
} else {
let mut min_loading_abs_corr = 1.0f32;
let mut min_score_abs_corr = 1.0f32;
let mut correlations_summary = String::new();
for pc_idx in 0..k_to_compare {
let rust_loading_col = rust_result
.final_snp_principal_component_loadings
.column(pc_idx);
let py_loading_col = py_loadings_d_x_k.column(pc_idx);
let loading_corr =
pearson_correlation(rust_loading_col.view(), py_loading_col.view())
.map_or(0.0, |c| c.abs());
if loading_corr < min_loading_abs_corr {
min_loading_abs_corr = loading_corr;
}
correlations_summary
.push_str(&format!("PC{}_Load_absR={:.4}; ", pc_idx, loading_corr));
if loading_corr < 0.95 {
test_successful = false;
outcome_details.push_str(&format!(
"Low loading correlation for PC {}: {:.4}. ",
pc_idx, loading_corr
));
}
let rust_score_col = rust_result
.final_sample_principal_component_scores
.column(pc_idx);
let py_score_col = py_scores_n_x_k.column(pc_idx);
let score_corr =
pearson_correlation(rust_score_col.view(), py_score_col.view())
.map_or(0.0, |c| c.abs());
if score_corr < min_score_abs_corr {
min_score_abs_corr = score_corr;
}
correlations_summary
.push_str(&format!("PC{}_Score_absR={:.4}; ", pc_idx, score_corr));
if score_corr < 0.95 {
test_successful = false;
outcome_details.push_str(&format!(
"Low score correlation for PC {}: {:.4}. ",
pc_idx, score_corr
));
}
}
outcome_details.push_str(&format!(
"Compared {} PCs. Min loading_absR={:.4}, Min score_absR={:.4}. Full: {}. ",
k_to_compare,
min_loading_abs_corr,
min_score_abs_corr,
correlations_summary.trim_end_matches("; ")
));
}
}
}
Err(e) => {
test_successful = false;
outcome_details.push_str(&format!("Rust PCA computation failed: {}. ", e));
}
}
let record = TestResultRecord {
test_name,
num_features_d: num_snps,
num_samples_n: num_samples,
num_pcs_requested_k: k_components,
num_pcs_computed: rust_pcs_computed,
success: test_successful,
outcome_details: outcome_details.clone(),
notes,
};
TEST_RESULTS.lock().unwrap().push(record);
assert!(
test_successful,
"Test '{}' failed. Check TSV. Details: {}",
test_name_str, outcome_details
);
}
#[test]
fn test_pc_correlation_with_truth_set_large_1000x200() {
run_pc_correlation_with_truth_set_test(
"test_pc_correlation_with_truth_set_large_1000x200",
1000, 200, 10, 202401, );
}
#[test]
fn test_pc_correlation_structured_1000snps_200samples_5truepcs() {
let num_snps = 1000; let num_samples = 200; let k_components_to_request = 10; let seed = 202408;
let num_true_pcs = 5; let signal_strength = 5.0;
let noise_std_dev = 1.0;
let test_name = "test_pc_correlation_structured_1000snps_200samples_5truepcs".to_string();
let mut test_successful = true;
let mut outcome_details = String::new();
let mut notes = format!(
"Structured Data Test: D_snps={}, N_samples={}, k_requested={}, k_true={}. ",
num_snps, num_samples, k_components_to_request, num_true_pcs
);
let structured_standardized_genos_snps_x_samples = generate_structured_data(
num_snps,
num_samples,
num_true_pcs,
signal_strength,
noise_std_dev,
seed,
);
let artifact_dir_suffix = format!(
"pc_corr_structured_{}x{}_k{}_true{}",
num_snps, num_samples, k_components_to_request, num_true_pcs
);
let artifact_dir = Path::new("target/test_artifacts").join(artifact_dir_suffix);
if let Err(e) = fs::create_dir_all(&artifact_dir) {
notes.push_str(&format!("Failed to create artifact dir: {}. ", e));
}
let mut py_loadings_d_x_k: Array2<f32> = Array2::zeros((0, 0)); let mut py_scores_n_x_k: Array2<f32> = Array2::zeros((0, 0)); let mut _py_eigenvalues_k: Array1<f64> = Array1::zeros(0); let mut effective_k_py = 0;
let python_pca_prefix = format!(
"pc_corr_structured_{}x{}_k{}_true{}_py_ref",
num_snps, num_samples, k_components_to_request, num_true_pcs
);
match get_python_reference_pca(
&structured_standardized_genos_snps_x_samples,
k_components_to_request,
&python_pca_prefix,
) {
Ok((loadings_k_x_d_py, scores_n_x_k_py, eigenvalues_k_py)) => {
py_loadings_d_x_k = loadings_k_x_d_py.t().into_owned(); py_scores_n_x_k = scores_n_x_k_py; _py_eigenvalues_k = eigenvalues_k_py; effective_k_py = py_loadings_d_x_k.ncols(); outcome_details.push_str(&format!(
"Python PCA successful. effective_k_py: {}. ",
effective_k_py
));
save_matrix_to_tsv(
&py_loadings_d_x_k.view(),
artifact_dir.to_str().unwrap_or("."),
"python_loadings.tsv",
)
.unwrap_or_default();
save_matrix_to_tsv(
&py_scores_n_x_k.view(),
artifact_dir.to_str().unwrap_or("."),
"python_scores.tsv",
)
.unwrap_or_default();
save_vector_to_tsv(
&_py_eigenvalues_k.view(),
artifact_dir.to_str().unwrap_or("."),
"python_eigenvalues.tsv",
)
.unwrap_or_default();
}
Err(e) => {
test_successful = false;
outcome_details.push_str(&format!("Python reference PCA failed: {}. ", e));
}
}
let test_data_accessor =
TestDataAccessor::new(structured_standardized_genos_snps_x_samples.clone());
let min_subset_size = (num_samples / 4).max(1).min(num_samples.max(1));
let max_subset_size = (num_samples / 2).max(10).min(num_samples.max(1));
let components_per_block = 10.min(num_snps.min(max_subset_size));
let config = EigenSNPCoreAlgorithmConfig {
target_num_global_pcs: k_components_to_request,
random_seed: seed,
subset_factor_for_local_basis_learning: 0.5, min_subset_size_for_local_basis_learning: min_subset_size,
max_subset_size_for_local_basis_learning: max_subset_size,
components_per_ld_block: components_per_block,
..Default::default()
};
let algorithm = EigenSNPCoreAlgorithm::new(config);
let ld_blocks = vec![LdBlockSpecification {
user_defined_block_tag: "block_structured".to_string(),
pca_snp_ids_in_block: (0..num_snps).map(PcaSnpId).collect(),
}];
let mut rust_pcs_computed = 0;
let snp_metadata = create_dummy_snp_metadata(num_snps);
match algorithm.compute_pca(&test_data_accessor, &ld_blocks, &snp_metadata) {
Ok((rust_result, _)) => {
rust_pcs_computed = rust_result.num_principal_components_computed;
outcome_details.push_str(&format!(
"eigensnp PCA successful. rust_pcs_computed: {}. ",
rust_pcs_computed
));
save_matrix_to_tsv(
&rust_result.final_snp_principal_component_loadings.view(),
artifact_dir.to_str().unwrap_or("."),
"rust_loadings.tsv",
)
.unwrap_or_default();
save_matrix_to_tsv(
&rust_result.final_sample_principal_component_scores.view(),
artifact_dir.to_str().unwrap_or("."),
"rust_scores.tsv",
)
.unwrap_or_default();
save_vector_to_tsv(
&rust_result.final_principal_component_eigenvalues.view(),
artifact_dir.to_str().unwrap_or("."),
"rust_eigenvalues.tsv",
)
.unwrap_or_default();
if test_successful {
let k_to_compare = rust_pcs_computed
.min(effective_k_py)
.min(num_true_pcs + 2)
.min(k_components_to_request);
outcome_details.push_str(&format!(
"Comparing up to {} PCs. True PCs: {}. ",
k_to_compare, num_true_pcs
));
if k_to_compare == 0 {
outcome_details.push_str(
"No components to compare (Rust or Python computed 0 relevant PCs). ",
);
if rust_pcs_computed != effective_k_py
&& (rust_pcs_computed == 0 || effective_k_py == 0)
{
test_successful = false; outcome_details.push_str(&format!(
"Mismatch in computed k (Rust: {}, Py: {}). ",
rust_pcs_computed, effective_k_py
));
}
} else {
let mut min_true_loading_abs_corr = 1.0f32;
let mut min_true_score_abs_corr = 1.0f32;
let mut correlations_summary = String::new();
for pc_idx in 0..k_to_compare {
let rust_loading_col = rust_result
.final_snp_principal_component_loadings
.column(pc_idx);
let py_loading_col = py_loadings_d_x_k.column(pc_idx);
let loading_corr_opt =
pearson_correlation(rust_loading_col.view(), py_loading_col.view());
let loading_abs_corr = loading_corr_opt.map_or(0.0, |c| c.abs());
correlations_summary
.push_str(&format!("PC{}_Load_absR={:.4}; ", pc_idx, loading_abs_corr));
let rust_score_col = rust_result
.final_sample_principal_component_scores
.column(pc_idx);
let py_score_col = py_scores_n_x_k.column(pc_idx);
let score_corr_opt =
pearson_correlation(rust_score_col.view(), py_score_col.view());
let score_abs_corr = score_corr_opt.map_or(0.0, |c| c.abs());
correlations_summary
.push_str(&format!("PC{}_Score_absR={:.4}; ", pc_idx, score_abs_corr));
if pc_idx < num_true_pcs {
if loading_abs_corr < min_true_loading_abs_corr {
min_true_loading_abs_corr = loading_abs_corr;
}
if score_abs_corr < min_true_score_abs_corr {
min_true_score_abs_corr = score_abs_corr;
}
if loading_abs_corr < 0.98 {
test_successful = false;
outcome_details.push_str(&format!(
"Low loading correlation for true PC {}: {:.4}. ",
pc_idx, loading_abs_corr
));
}
if score_abs_corr < 0.98 {
test_successful = false;
outcome_details.push_str(&format!(
"Low score correlation for true PC {}: {:.4}. ",
pc_idx, score_abs_corr
));
}
} else {
if loading_abs_corr < 0.70 {
notes.push_str(&format!(
"Note: Loading correlation for non-true PC {} is {:.4}. ",
pc_idx, loading_abs_corr
));
}
if score_abs_corr < 0.70 {
notes.push_str(&format!(
"Note: Score correlation for non-true PC {} is {:.4}. ",
pc_idx, score_abs_corr
));
}
}
}
outcome_details.push_str(&format!(
"For first {} true PCs: Min_Load_absR={:.4}, Min_Score_absR={:.4}. Full_Corr_Summary: {}. ",
num_true_pcs.min(k_to_compare), min_true_loading_abs_corr,
min_true_score_abs_corr,
correlations_summary.trim_end_matches("; ")
));
}
}
}
Err(e) => {
test_successful = false;
outcome_details.push_str(&format!("eigensnp PCA computation failed: {}. ", e));
}
}
let record = TestResultRecord {
test_name: test_name.clone(),
num_features_d: num_snps,
num_samples_n: num_samples,
num_pcs_requested_k: k_components_to_request,
num_pcs_computed: rust_pcs_computed, success: test_successful,
outcome_details: outcome_details.clone(),
notes,
};
TEST_RESULTS.lock().unwrap().push(record);
assert!(
test_successful,
"Test '{}' failed. Check artifacts in '{}'. Details: {}",
test_name,
artifact_dir.display(),
outcome_details
);
}
pub fn run_generic_large_matrix_test(
test_name_str: &str,
num_snps: usize, num_samples: usize, k_components: usize, seed: u64,
config_modifier: Option<fn(EigenSNPCoreAlgorithmConfig) -> EigenSNPCoreAlgorithmConfig>,
) {
let mut outcome_details = String::new(); let test_name = test_name_str.to_string();
let mut test_successful = true;
let mut notes = format!(
"Matrix D_snps x N_samples: {}x{}, k_requested: {}. ",
num_snps, num_samples, k_components
);
let artifact_dir_suffix = format!(
"generic_large_matrix_{}x{}_k{}",
num_snps, num_samples, k_components
);
let artifact_dir = Path::new("target/test_artifacts").join(artifact_dir_suffix);
if let Err(e) = fs::create_dir_all(&artifact_dir) {
notes.push_str(&format!("Failed to create artifact dir: {}. ", e));
}
let mut rng = ChaCha8Rng::seed_from_u64(seed);
let raw_genos = Array2::random_using((num_snps, num_samples), Uniform::new(0.0, 3.0), &mut rng);
let standardized_genos = standardize_features_across_samples(raw_genos);
let test_data_accessor = TestDataAccessor::new(standardized_genos);
let mut base_config = EigenSNPCoreAlgorithmConfig {
target_num_global_pcs: k_components,
random_seed: seed,
..Default::default()
};
if let Some(modifier) = config_modifier {
base_config = modifier(base_config);
notes.push_str("Custom config modifier applied. ");
}
let algorithm = EigenSNPCoreAlgorithm::new(base_config);
let ld_blocks = vec![LdBlockSpecification {
user_defined_block_tag: "block_generic".to_string(),
pca_snp_ids_in_block: (0..num_snps).map(PcaSnpId).collect(),
}];
let mut rust_pcs_computed = 0;
let snp_metadata = create_dummy_snp_metadata(num_snps);
match algorithm.compute_pca(&test_data_accessor, &ld_blocks, &snp_metadata) {
Ok((output, _)) => {
rust_pcs_computed = output.num_principal_components_computed;
write!(
&mut outcome_details,
"eigensnp successful. Computed {} PCs. First eigenvalue: {:.4}. ",
output.num_principal_components_computed,
output
.final_principal_component_eigenvalues
.get(0)
.unwrap_or(&0.0)
)
.unwrap_or_default();
if rust_pcs_computed == 0 && k_components > 0 {
test_successful = false;
write!(
&mut outcome_details,
"Warning: 0 PCs computed when k_requested > 0. "
)
.unwrap_or_default();
}
if rust_pcs_computed > k_components {
test_successful = false;
write!(
&mut outcome_details,
"Warning: More PCs computed ({}) than requested ({}). ",
rust_pcs_computed, k_components
)
.unwrap_or_default();
}
save_matrix_to_tsv(
&output.final_snp_principal_component_loadings.view(),
artifact_dir.to_str().unwrap_or("."),
"rust_loadings.tsv",
)
.unwrap_or_default();
save_matrix_to_tsv(
&output.final_sample_principal_component_scores.view(),
artifact_dir.to_str().unwrap_or("."),
"rust_scores.tsv",
)
.unwrap_or_default();
save_vector_to_tsv(
&output.final_principal_component_eigenvalues.view(),
artifact_dir.to_str().unwrap_or("."),
"rust_eigenvalues.tsv",
)
.unwrap_or_default();
}
Err(e) => {
test_successful = false;
write!(
&mut outcome_details,
"eigensnp PCA computation failed: {}. ",
e
)
.unwrap_or_default();
}
}
let record = TestResultRecord {
test_name,
num_features_d: num_snps,
num_samples_n: num_samples,
num_pcs_requested_k: k_components,
num_pcs_computed: rust_pcs_computed,
success: test_successful,
outcome_details: outcome_details.clone(),
notes,
};
TEST_RESULTS.lock().unwrap().push(record);
assert!(
test_successful,
"Test '{}' failed. Check TSV. Details: {}",
test_name_str, outcome_details
);
}
#[test]
fn test_large_matrix_2000x200_k10() {
run_generic_large_matrix_test(
"test_large_matrix_2000x200_k10",
2000, 200, 10, 202405, None,
);
}
#[test]
fn test_large_matrix_5000x500_k20() {
run_generic_large_matrix_test(
"test_large_matrix_5000x500_k20",
5000, 500, 20, 202406, None,
);
}
#[test]
fn test_large_matrix_1000x100_k5_blocksize_variation() {
run_generic_large_matrix_test(
"test_large_matrix_1000x100_k5_blocksize_variation",
1000, 100, 5, 202407, Some(|mut cfg: EigenSNPCoreAlgorithmConfig| {
cfg.components_per_ld_block = 20; cfg.subset_factor_for_local_basis_learning = 0.8;
cfg
}),
);
}
pub fn run_sample_projection_accuracy_test(
test_name_str: &str,
num_snps: usize, num_samples_total: usize, num_samples_train: usize, k_components: usize, seed: u64,
) {
let test_name = test_name_str.to_string();
let mut test_successful = true;
let mut outcome_details: String;
let num_samples_test = num_samples_total - num_samples_train;
let mut notes = format!(
"Matrix D_snps x N_total_samples (N_train_samples / N_test_samples): {}x{} ({} / {}), k_requested: {}. ",
num_snps, num_samples_total, num_samples_train, num_samples_test, k_components
);
assert!(num_samples_train > 0, "num_samples_train must be > 0");
assert!(
num_samples_total > num_samples_train,
"num_samples_total must be > num_samples_train"
);
let artifact_dir_suffix = format!(
"sample_projection_{}x{}_k{}",
num_snps, num_samples_train, k_components
);
let artifact_dir = Path::new("target/test_artifacts").join(artifact_dir_suffix);
if let Err(e) = fs::create_dir_all(&artifact_dir) {
notes.push_str(&format!("Failed to create artifact dir: {}. ", e));
}
let mut rng = ChaCha8Rng::seed_from_u64(seed);
let raw_genos_total = Array2::random_using(
(num_snps, num_samples_total),
Uniform::new(0.0, 3.0),
&mut rng,
);
let standardized_genos_total_snps_x_samples =
standardize_features_across_samples(raw_genos_total);
let train_data_snps_x_samples = standardized_genos_total_snps_x_samples
.slice(s![.., 0..num_samples_train])
.to_owned();
let test_data_snps_x_samples = standardized_genos_total_snps_x_samples
.slice(s![.., num_samples_train..])
.to_owned();
let test_data_accessor_train = TestDataAccessor::new(train_data_snps_x_samples.clone());
let config_train = EigenSNPCoreAlgorithmConfig {
target_num_global_pcs: k_components,
random_seed: seed,
..Default::default()
};
let algorithm_train = EigenSNPCoreAlgorithm::new(config_train);
let ld_blocks_train = vec![LdBlockSpecification {
user_defined_block_tag: "block_train".to_string(),
pca_snp_ids_in_block: (0..num_snps).map(PcaSnpId).collect(),
}];
let mut rust_pca_output_option: Option<EigenSNPCoreOutput> = None; let mut k_eff_rust = 0;
let snp_metadata = create_dummy_snp_metadata(num_snps);
match algorithm_train.compute_pca(&test_data_accessor_train, &ld_blocks_train, &snp_metadata) {
Ok((output_struct, _)) => {
k_eff_rust = output_struct.num_principal_components_computed;
save_matrix_to_tsv(
&output_struct.final_snp_principal_component_loadings.view(),
artifact_dir.to_str().unwrap_or("."),
"rust_train_loadings.tsv",
)
.unwrap_or_default();
save_matrix_to_tsv(
&output_struct.final_sample_principal_component_scores.view(),
artifact_dir.to_str().unwrap_or("."),
"rust_train_scores.tsv",
)
.unwrap_or_default();
rust_pca_output_option = Some(output_struct);
outcome_details = format!(
"eigensnp on train data successful. k_eff_rust: {}. ",
k_eff_rust
);
}
Err(e) => {
test_successful = false;
outcome_details = format!("eigensnp on train data failed: {}. ", e);
}
}
let mut s_projected_option: Option<Array2<f32>> = None;
if let Some(ref rust_pca_output) = rust_pca_output_option {
if k_eff_rust > 0 {
let loadings_l = &rust_pca_output.final_snp_principal_component_loadings; if test_data_snps_x_samples.nrows() == loadings_l.nrows() {
let projected_scores = test_data_snps_x_samples.t().dot(loadings_l);
save_matrix_to_tsv(
&projected_scores.view(),
artifact_dir.to_str().unwrap_or("."),
"rust_projected_test_scores.tsv",
)
.unwrap_or_default();
s_projected_option = Some(projected_scores);
} else {
test_successful = false;
outcome_details.push_str(&format!("Dimension mismatch for projection: test_data_snps_x_samples.nrows() ({}) != loadings_l.nrows() ({}). ",
test_data_snps_x_samples.nrows(), loadings_l.nrows()));
}
} else {
outcome_details.push_str("Rust PCA computed 0 components, skipping projection. ");
}
}
let mut py_test_scores_ref_option: Option<Array2<f32>> = None;
if test_successful {
let python_total_data_prefix = format!(
"sample_projection_{}x{}_k{}_py_total_ref",
num_snps, num_samples_total, k_components
);
match get_python_reference_pca(
&standardized_genos_total_snps_x_samples,
k_components, &python_total_data_prefix,
) {
Ok((_py_loadings_total_k_x_d, py_scores_total_n_x_k, _py_eigenvalues_total)) => {
let k_py_total = _py_loadings_total_k_x_d.nrows(); if py_scores_total_n_x_k.nrows() == num_samples_total
&& py_scores_total_n_x_k.ncols() >= k_components.min(k_py_total)
{
let num_cols_to_slice = k_eff_rust.min(py_scores_total_n_x_k.ncols());
if num_cols_to_slice > 0 {
let py_test_scores_ref = py_scores_total_n_x_k
.slice(s![num_samples_train.., 0..num_cols_to_slice])
.to_owned();
save_matrix_to_tsv(
&py_test_scores_ref.view(),
artifact_dir.to_str().unwrap_or("."),
"python_ref_test_scores.tsv",
)
.unwrap_or_default();
py_test_scores_ref_option = Some(py_test_scores_ref);
outcome_details.push_str(&format!("Python on total data successful. k_py_total: {}. Sliced to {} cols for comparison. ", k_py_total, num_cols_to_slice));
} else {
outcome_details.push_str(
"Python on total data: 0 relevant components to slice for comparison. ",
);
if k_eff_rust > 0 {
test_successful = false;
outcome_details.push_str("Mismatch: Rust produced PCs but Python reference had 0 comparable PCs. ");
}
}
} else {
test_successful = false;
outcome_details.push_str(&format!(
"Python (total data) scores dimensions mismatch. Expected N_total x >=k_eff_py ({}x{}), Got {}x{}. ",
num_samples_total, k_components.min(k_py_total),
py_scores_total_n_x_k.nrows(), py_scores_total_n_x_k.ncols()
));
}
}
Err(e) => {
test_successful = false;
outcome_details.push_str(&format!(
"Python reference PCA on total data failed: {}. ",
e
));
}
}
}
if test_successful && s_projected_option.is_some() && py_test_scores_ref_option.is_some() {
let s_projected = s_projected_option.as_ref().unwrap();
let py_test_scores_ref = py_test_scores_ref_option.as_ref().unwrap();
let k_compare = k_eff_rust.min(py_test_scores_ref.ncols());
outcome_details.push_str(&format!("Comparing {} PCs for projection. ", k_compare));
if k_compare == 0 {
if k_eff_rust != py_test_scores_ref.ncols() {
test_successful = false;
outcome_details.push_str(&format!(
"Mismatch in comparable k (Rust_eff_k: {}, Py_ref_k: {}). ",
k_eff_rust,
py_test_scores_ref.ncols()
));
} else {
outcome_details.push_str("Both Rust and Py_ref have 0 PCs to compare. ");
}
} else {
let mut min_abs_corr = 1.0f32;
let mut max_mse = 0.0f32;
let mut correlations_summary = String::new();
let mut mses_summary = String::new();
for pc_idx in 0..k_compare {
let projected_col = s_projected.column(pc_idx);
let ref_col = py_test_scores_ref.column(pc_idx);
let abs_corr = pearson_correlation(projected_col.view(), ref_col.view())
.map_or(0.0, |c| c.abs());
if abs_corr < min_abs_corr {
min_abs_corr = abs_corr;
}
correlations_summary.push_str(&format!("PC{}_absR={:.4}; ", pc_idx, abs_corr));
if abs_corr < 0.95 {
test_successful = false;
outcome_details.push_str(&format!(
"Low projection correlation for PC {}: {:.4}. ",
pc_idx, abs_corr
));
}
let mse = (projected_col.to_owned() - ref_col.to_owned())
.mapv(|x| x * x)
.mean()
.unwrap_or(f32::MAX);
if mse > max_mse {
max_mse = mse;
}
mses_summary.push_str(&format!("PC{}_MSE={:.4e}; ", pc_idx, mse));
if mse > 0.1 {
test_successful = false;
outcome_details.push_str(&format!(
"High projection MSE for PC {}: {:.4e}. ",
pc_idx, mse
));
}
}
outcome_details.push_str(&format!(
"Min_abs_correlation: {:.4}, Max_MSE: {:.4e}. Correlations: {}. MSEs: {}. ",
min_abs_corr,
max_mse,
correlations_summary.trim_end_matches("; "),
mses_summary.trim_end_matches("; ")
));
}
} else if test_successful {
test_successful = false; outcome_details.push_str("Comparison skipped due to missing projected or reference scores despite earlier success. ");
}
let record = TestResultRecord {
test_name,
num_features_d: num_snps,
num_samples_n: num_samples_total, num_pcs_requested_k: k_components,
num_pcs_computed: k_eff_rust, success: test_successful,
outcome_details: outcome_details.clone(),
notes,
};
TEST_RESULTS.lock().unwrap().push(record);
assert!(
test_successful,
"Test '{}' failed. Check TSV. Details: {}",
test_name_str, outcome_details
);
}
#[test]
fn test_sample_projection_accuracy_large_config1() {
run_sample_projection_accuracy_test(
"test_sample_projection_accuracy_large_config1",
1000, 250, 200, 10, 202403, );
}
#[test]
fn test_sample_projection_accuracy_large_config2() {
run_sample_projection_accuracy_test(
"test_sample_projection_accuracy_large_config2",
2000, 400, 300, 15, 202404, );
}
#[test]
fn test_pc_correlation_with_truth_set_large_2000x300() {
run_pc_correlation_with_truth_set_test(
"test_pc_correlation_with_truth_set_large_2000x300",
2000, 300, 15, 202402, );
}
#[test]
fn test_reorder_array_basic() {
let original = Array1::from(vec![10, 20, 30, 40]);
let order = vec![2, 0, 3, 1];
let expected = Array1::from(vec![30, 10, 40, 20]);
let reordered = reorder_array_owned(&original, &order);
assert_eq!(reordered, expected);
}
#[test]
fn test_reorder_array_empty_array_empty_order() {
let original = Array1::<i32>::from(vec![]);
let order_empty = vec![];
let expected = Array1::<i32>::from(vec![]);
let reordered_empty_order = reorder_array_owned(&original, &order_empty);
assert_eq!(reordered_empty_order, expected);
}
#[test]
#[should_panic]
fn test_reorder_array_select_from_zero_elements_panics() {
let original = Array1::<i32>::from(vec![]);
let order = vec![0];
reorder_array_owned(&original, &order);
}
#[test]
fn test_reorder_array_empty_order() {
let original = Array1::from(vec![10, 20, 30]);
let order = vec![];
let expected = Array1::<i32>::from(vec![]);
let reordered = reorder_array_owned(&original, &order);
assert_eq!(reordered, expected);
}
#[test]
fn test_reorder_array_repeated_indices() {
let original = Array1::from(vec![10, 20, 30]);
let order = vec![0, 1, 0, 2, 1, 1];
let expected = Array1::from(vec![10, 20, 10, 30, 20, 20]);
let reordered = reorder_array_owned(&original, &order);
assert_eq!(reordered, expected);
}
#[test]
fn test_reorder_columns_basic() {
let original = arr2(&[[1, 2, 3], [4, 5, 6]]);
let order = vec![2, 0, 1];
let expected = arr2(&[[3, 1, 2], [6, 4, 5]]);
let reordered = reorder_columns_owned(&original, &order);
assert_eq!(reordered, expected);
}
#[allow(clippy::too_many_arguments)] fn run_refinement_improvement_test<F>(
test_name: &str,
standardized_structured_data: &Array2<f32>,
ld_block_specs: &[LdBlockSpecification],
k_components_to_request: usize,
python_reference_output: &(Array2<f32>, Array2<f32>, Array1<f64>), pass_count_less_refined: usize,
pass_count_more_refined: usize,
metric_evaluator: F,
metric_name: &str,
seed: u64,
) -> Result<(), String>
where
F: Fn(
&EigenSNPCoreOutput,
&(Array2<f32>, Array2<f32>, Array1<f64>), // Python ref: (loadings_k_x_d, scores_n_x_k, eigenvalues_k)
&str, // Context string
) -> f64, {
let full_test_name = format!("{}_{}", test_name, metric_name);
let artifact_dir_name = full_test_name.replace(|c: char| !c.is_alphanumeric() && c != '_', "_"); let artifact_dir = Path::new("target/test_artifacts").join(artifact_dir_name);
fs::create_dir_all(&artifact_dir).map_err(|e| {
format!(
"Failed to create artifact directory '{}': {}",
artifact_dir.display(),
e
)
})?;
let mut outcome_details = String::new();
let mut notes = format!("Seed: {}. ", seed);
let mut overall_test_success = true;
save_matrix_to_tsv(
&python_reference_output.0.view(),
artifact_dir.to_str().unwrap(),
"python_ref_loadings_k_x_d.tsv",
)
.map_err(|e| format!("Failed to save python_ref_loadings.tsv: {}", e))?;
save_matrix_to_tsv(
&python_reference_output.1.view(),
artifact_dir.to_str().unwrap(),
"python_ref_scores_n_x_k.tsv",
)
.map_err(|e| format!("Failed to save python_ref_scores.tsv: {}", e))?;
save_vector_to_tsv(
&python_reference_output.2.view(),
artifact_dir.to_str().unwrap(),
"python_ref_eigenvalues_k.tsv",
)
.map_err(|e| format!("Failed to save python_ref_eigenvalues.tsv: {}", e))?;
let config_a = EigenSNPCoreAlgorithmConfig {
target_num_global_pcs: k_components_to_request,
refine_pass_count: pass_count_less_refined,
random_seed: seed,
subset_factor_for_local_basis_learning: 0.5,
min_subset_size_for_local_basis_learning: (standardized_structured_data.ncols() / 4)
.max(1)
.min(standardized_structured_data.ncols().max(1)),
max_subset_size_for_local_basis_learning: (standardized_structured_data.ncols() / 2)
.max(10)
.min(standardized_structured_data.ncols().max(1)),
components_per_ld_block: 10.min(
standardized_structured_data.nrows().min(
(standardized_structured_data.ncols() / 2)
.max(10)
.min(standardized_structured_data.ncols().max(1)),
),
),
..Default::default()
};
let config_b = EigenSNPCoreAlgorithmConfig {
refine_pass_count: pass_count_more_refined,
..config_a.clone() };
let test_data_accessor = TestDataAccessor::new(standardized_structured_data.clone());
let algorithm_a = EigenSNPCoreAlgorithm::new(config_a);
let algorithm_b = EigenSNPCoreAlgorithm::new(config_b);
let snp_metadata = create_dummy_snp_metadata(standardized_structured_data.nrows());
let output_a = match algorithm_a.compute_pca(&test_data_accessor, ld_block_specs, &snp_metadata)
{
Ok((out, _)) => {
writeln!(
outcome_details,
"EigenSnp (Less Refined, {} passes): SUCCESS.",
pass_count_less_refined
)
.unwrap_or_default();
save_matrix_to_tsv(
&out.final_snp_principal_component_loadings.view(),
artifact_dir.to_str().unwrap(),
"eigensnp_less_refined_loadings.tsv",
)
.unwrap_or_default();
save_matrix_to_tsv(
&out.final_sample_principal_component_scores.view(),
artifact_dir.to_str().unwrap(),
"eigensnp_less_refined_scores.tsv",
)
.unwrap_or_default();
save_vector_to_tsv(
&out.final_principal_component_eigenvalues.view(),
artifact_dir.to_str().unwrap(),
"eigensnp_less_refined_eigenvalues.tsv",
)
.unwrap_or_default();
out
}
Err(e) => {
writeln!(
outcome_details,
"EigenSnp (Less Refined, {} passes): FAILED. Error: {}",
pass_count_less_refined, e
)
.unwrap_or_default();
notes.push_str(&format!("EigenSnp A failed: {}. ", e));
overall_test_success = false;
EigenSNPCoreOutput {
final_snp_principal_component_loadings: Array2::zeros((0, 0)),
final_sample_principal_component_scores: Array2::zeros((0, 0)),
final_principal_component_eigenvalues: Array1::zeros(0),
num_principal_components_computed: 0,
num_pca_snps_used: standardized_structured_data.nrows(),
num_qc_samples_used: standardized_structured_data.ncols(),
}
}
};
let output_b = match algorithm_b.compute_pca(&test_data_accessor, ld_block_specs, &snp_metadata)
{
Ok((out, _)) => {
writeln!(
outcome_details,
"EigenSnp (More Refined, {} passes): SUCCESS.",
pass_count_more_refined
)
.unwrap_or_default();
save_matrix_to_tsv(
&out.final_snp_principal_component_loadings.view(),
artifact_dir.to_str().unwrap(),
"eigensnp_more_refined_loadings.tsv",
)
.unwrap_or_default();
save_matrix_to_tsv(
&out.final_sample_principal_component_scores.view(),
artifact_dir.to_str().unwrap(),
"eigensnp_more_refined_scores.tsv",
)
.unwrap_or_default();
save_vector_to_tsv(
&out.final_principal_component_eigenvalues.view(),
artifact_dir.to_str().unwrap(),
"eigensnp_more_refined_eigenvalues.tsv",
)
.unwrap_or_default();
out
}
Err(e) => {
writeln!(
outcome_details,
"EigenSnp (More Refined, {} passes): FAILED. Error: {}",
pass_count_more_refined, e
)
.unwrap_or_default();
notes.push_str(&format!("EigenSnp B failed: {}. ", e));
overall_test_success = false;
EigenSNPCoreOutput {
final_snp_principal_component_loadings: Array2::zeros((0, 0)),
final_sample_principal_component_scores: Array2::zeros((0, 0)),
final_principal_component_eigenvalues: Array1::zeros(0),
num_principal_components_computed: 0,
num_pca_snps_used: standardized_structured_data.nrows(),
num_qc_samples_used: standardized_structured_data.ncols(),
}
}
};
let num_pcs_computed_for_log =
if overall_test_success || output_b.num_principal_components_computed > 0 {
output_b.num_principal_components_computed
} else {
output_a.num_principal_components_computed
};
let score_a = if !notes.contains("EigenSnp A failed") {
metric_evaluator(&output_a, python_reference_output, "LessRefined_vs_Ref")
} else {
f64::NAN };
let score_b = if !notes.contains("EigenSnp B failed") {
metric_evaluator(&output_b, python_reference_output, "MoreRefined_vs_Ref")
} else {
f64::NAN
};
writeln!(
outcome_details,
"{}: Less Refined ({} passes) score = {:.6e}, More Refined ({} passes) score = {:.6e}.",
metric_name, pass_count_less_refined, score_a, pass_count_more_refined, score_b
)
.unwrap_or_default();
let tolerance = 1e-9;
let improvement_observed = score_b >= score_a - tolerance;
if score_a.is_nan() || score_b.is_nan() {
writeln!(
outcome_details,
"Improvement check SKIPPED due to metric computation failure for one or both runs."
)
.unwrap_or_default();
overall_test_success = false; } else if improvement_observed {
writeln!(
outcome_details,
"Improvement (or non-degradation within tolerance) OBSERVED."
)
.unwrap_or_default();
} else {
writeln!(
outcome_details,
"Improvement NOT OBSERVED. Score B ({:.6e}) < Score A ({:.6e}) - tolerance.",
score_b, score_a
)
.unwrap_or_default();
overall_test_success = false; }
let record = TestResultRecord {
test_name: full_test_name,
num_features_d: standardized_structured_data.nrows(),
num_samples_n: standardized_structured_data.ncols(),
num_pcs_requested_k: k_components_to_request,
num_pcs_computed: num_pcs_computed_for_log,
success: overall_test_success,
outcome_details: outcome_details.clone(), notes,
};
TEST_RESULTS.lock().unwrap().push(record);
if !overall_test_success {
return Err(format!(
"Test '{}' failed. Details: {}",
test_name, outcome_details
));
}
Ok(())
}
pub fn evaluate_pc_score_correlation(
eigensnp_output: &EigenSNPCoreOutput,
reference_output: &(Array2<f32>, Array2<f32>, Array1<f64>), _context: &str, ) -> f64 {
let eigensnp_scores = &eigensnp_output.final_sample_principal_component_scores; let ref_scores = &reference_output.1;
let k_compare = eigensnp_scores.ncols().min(ref_scores.ncols());
if k_compare == 0 {
return 0.0; }
let mut total_abs_correlation = 0.0;
for pc_idx in 0..k_compare {
let eigensnp_score_col = eigensnp_scores.column(pc_idx);
let ref_score_col = ref_scores.column(pc_idx);
let corr =
pearson_correlation(eigensnp_score_col.view(), ref_score_col.view()).unwrap_or(0.0);
let abs_corr = corr.abs();
total_abs_correlation += abs_corr as f64; }
total_abs_correlation / (k_compare as f64) }
pub fn evaluate_snp_loading_correlation(
eigensnp_output: &EigenSNPCoreOutput,
reference_output: &(Array2<f32>, Array2<f32>, Array1<f64>), _context: &str, ) -> f64 {
let eigensnp_loadings = &eigensnp_output.final_snp_principal_component_loadings; let ref_loadings_k_x_d = &reference_output.0;
let k_compare = eigensnp_loadings.ncols().min(ref_loadings_k_x_d.nrows());
if k_compare == 0 {
return 0.0; }
if eigensnp_loadings.nrows() != ref_loadings_k_x_d.ncols() && k_compare > 0 {
eprintln!(
"SNP dimension mismatch in evaluate_snp_loading_correlation: eigensnp D={}, ref D={}",
eigensnp_loadings.nrows(),
ref_loadings_k_x_d.ncols()
);
return f64::NAN; }
let mut total_abs_correlation = 0.0;
for pc_idx in 0..k_compare {
let eigensnp_loading_col = eigensnp_loadings.column(pc_idx);
let ref_loading_vector_for_pc = ref_loadings_k_x_d.row(pc_idx);
let corr = pearson_correlation(
eigensnp_loading_col.view(),
ref_loading_vector_for_pc.view(),
)
.unwrap_or(0.0);
let abs_corr = corr.abs();
total_abs_correlation += abs_corr as f64;
}
total_abs_correlation / (k_compare as f64) }
pub fn evaluate_eigenvalue_accuracy(
eigensnp_output: &EigenSNPCoreOutput,
reference_output: &(Array2<f32>, Array2<f32>, Array1<f64>), _context: &str, ) -> f64 {
let eigensnp_eigenvalues = &eigensnp_output.final_principal_component_eigenvalues; let ref_eigenvalues = &reference_output.2;
let k_compare = eigensnp_eigenvalues.len().min(ref_eigenvalues.len());
if k_compare == 0 {
return 0.0;
}
let mut total_squared_relative_error = 0.0;
for i in 0..k_compare {
let e_eigensnp = eigensnp_eigenvalues[i];
let e_ref = ref_eigenvalues[i];
let squared_relative_error: f64;
if e_ref.abs() < 1e-9 {
if e_eigensnp.abs() < 1e-9 {
squared_relative_error = 0.0; } else {
squared_relative_error = 1.0e6; }
} else {
let relative_error = (e_eigensnp - e_ref) / e_ref;
squared_relative_error = relative_error.powi(2);
}
total_squared_relative_error += squared_relative_error.min(1.0e12); }
let mean_squared_relative_error = total_squared_relative_error / (k_compare as f64);
-mean_squared_relative_error }
#[test]
fn test_refinement_score_correlation() {
let d_total_snps = 5000;
let n_samples = 500;
let k_true_components = 10;
let signal_strength = 3.0;
let noise_std_dev = 1.0;
let seed = 20241001;
let standardized_structured_data = generate_structured_data(
d_total_snps,
n_samples,
k_true_components,
signal_strength,
noise_std_dev,
seed,
);
let k_components_to_request_pca = k_true_components + 5;
let python_reference_output_result = get_python_reference_pca(
&standardized_structured_data,
k_components_to_request_pca,
"refinement_score_corr_py_ref",
);
let python_reference_output = match python_reference_output_result {
Ok(output) => output,
Err(e) => {
panic!(
"Failed to get Python reference PCA for test_refinement_score_correlation: {}",
e
);
}
};
let ld_block_specs = vec![LdBlockSpecification {
user_defined_block_tag: "full_block".to_string(),
pca_snp_ids_in_block: (0..d_total_snps).map(PcaSnpId).collect(),
}];
let result = run_refinement_improvement_test(
"refinement_score_correlation",
&standardized_structured_data,
&ld_block_specs,
k_true_components, &python_reference_output,
1, 2, evaluate_pc_score_correlation,
"PCScoreCorrelation",
seed, );
if let Err(e) = result {
panic!("test_refinement_score_correlation failed: {}", e);
}
}
#[test]
fn test_refinement_loading_correlation() {
let d_total_snps = 5000;
let n_samples = 500;
let k_true_components = 10;
let signal_strength = 3.0;
let noise_std_dev = 1.0;
let seed = 20241002;
let standardized_structured_data = generate_structured_data(
d_total_snps,
n_samples,
k_true_components,
signal_strength,
noise_std_dev,
seed,
);
let k_components_to_request_pca = k_true_components + 5;
let python_reference_output_result = get_python_reference_pca(
&standardized_structured_data,
k_components_to_request_pca,
"refinement_loading_corr_py_ref", );
let python_reference_output = match python_reference_output_result {
Ok(output) => output,
Err(e) => {
panic!(
"Failed to get Python reference PCA for test_refinement_loading_correlation: {}",
e
);
}
};
let ld_block_specs = vec![LdBlockSpecification {
user_defined_block_tag: "full_block".to_string(),
pca_snp_ids_in_block: (0..d_total_snps).map(PcaSnpId).collect(),
}];
let result = run_refinement_improvement_test(
"refinement_loading_correlation",
&standardized_structured_data,
&ld_block_specs,
k_true_components, &python_reference_output,
1, 2, evaluate_snp_loading_correlation,
"SNPLoadingCorrelation",
seed, );
if let Err(e) = result {
panic!("test_refinement_loading_correlation failed: {}", e);
}
}
#[test]
fn test_refinement_eigenvalue_accuracy() {
let d_total_snps = 5000;
let n_samples = 500;
let k_true_components = 10;
let signal_strength = 3.0;
let noise_std_dev = 1.0;
let seed = 20241003;
let standardized_structured_data = generate_structured_data(
d_total_snps,
n_samples,
k_true_components,
signal_strength,
noise_std_dev,
seed,
);
let k_components_to_request_pca = k_true_components + 5;
let python_reference_output_result = get_python_reference_pca(
&standardized_structured_data,
k_components_to_request_pca,
"refinement_eigenvalue_acc_py_ref", );
let python_reference_output = match python_reference_output_result {
Ok(output) => output,
Err(e) => {
panic!(
"Failed to get Python reference PCA for test_refinement_eigenvalue_accuracy: {}",
e
);
}
};
let ld_block_specs = vec![LdBlockSpecification {
user_defined_block_tag: "full_block".to_string(),
pca_snp_ids_in_block: (0..d_total_snps).map(PcaSnpId).collect(),
}];
let result = run_refinement_improvement_test(
"refinement_eigenvalue_accuracy",
&standardized_structured_data,
&ld_block_specs,
k_true_components, &python_reference_output,
1, 2, evaluate_eigenvalue_accuracy,
"EigenvalueAccuracy",
seed, );
if let Err(e) = result {
panic!("test_refinement_eigenvalue_accuracy failed: {}", e);
}
}
struct QualityThresholds {
min_score_correlation: f64,
min_loading_correlation: f64,
max_neg_eigenvalue_accuracy: f64, }
#[test]
fn test_min_passes_for_quality_convergence() {
let test_logging_name = "test_min_passes_for_quality_convergence";
let d_total_snps = 5000;
let n_samples = 500;
let k_true_components = 10;
let signal_strength = 3.0;
let noise_std_dev = 1.0;
let seed = 20241005;
let standardized_structured_data = generate_structured_data(
d_total_snps,
n_samples,
k_true_components,
signal_strength,
noise_std_dev,
seed,
);
let artifact_dir = Path::new("target/test_artifacts").join(test_logging_name);
fs::create_dir_all(&artifact_dir).unwrap_or_else(|e| {
panic!(
"Failed to create artifact directory '{}': {}",
artifact_dir.display(),
e
)
});
let k_components_to_request_pca = k_true_components + 5;
let python_reference_output_result = get_python_reference_pca(
&standardized_structured_data,
k_components_to_request_pca,
&format!("{}_py_ref", test_logging_name),
);
let python_reference_output = match python_reference_output_result {
Ok(output) => output,
Err(e) => {
panic!(
"Failed to get Python reference PCA for {}: {}",
test_logging_name, e
);
}
};
let thresholds = QualityThresholds {
min_score_correlation: 0.998,
min_loading_correlation: 0.995,
max_neg_eigenvalue_accuracy: -0.01, };
let ld_block_specs = vec![LdBlockSpecification {
user_defined_block_tag: "full_block".to_string(),
pca_snp_ids_in_block: (0..d_total_snps).map(PcaSnpId).collect(),
}];
let snp_metadata = create_dummy_snp_metadata(d_total_snps);
let max_passes_to_test = 5;
let mut min_passes_found: i32 = -1;
let mut overall_outcome_details = String::new();
writeln!(overall_outcome_details, "Test: {}", test_logging_name).unwrap_or_default();
let mut num_pcs_computed_at_convergence = 0;
for current_pass_count in 1..=max_passes_to_test {
writeln!(
overall_outcome_details,
"\n--- Evaluating with {} refinement pass(es) ---",
current_pass_count
)
.unwrap_or_default();
let config = EigenSNPCoreAlgorithmConfig {
target_num_global_pcs: k_true_components,
refine_pass_count: current_pass_count,
random_seed: seed,
subset_factor_for_local_basis_learning: 0.5,
min_subset_size_for_local_basis_learning: (n_samples / 4).max(1).min(n_samples.max(1)),
max_subset_size_for_local_basis_learning: (n_samples / 2).max(10).min(n_samples.max(1)),
components_per_ld_block: 10
.min(d_total_snps.min((n_samples / 2).max(10).min(n_samples.max(1)))),
..Default::default()
};
let test_data_accessor = TestDataAccessor::new(standardized_structured_data.clone());
let algorithm = EigenSNPCoreAlgorithm::new(config);
match algorithm.compute_pca(&test_data_accessor, &ld_block_specs, &snp_metadata) {
Ok((eigensnp_output_current_pass, _)) => {
if min_passes_found == -1 {
num_pcs_computed_at_convergence =
eigensnp_output_current_pass.num_principal_components_computed;
}
let pass_artifact_dir_name = format!("eigensnp_pass_{}", current_pass_count);
let pass_artifact_dir = artifact_dir.join(pass_artifact_dir_name);
fs::create_dir_all(&pass_artifact_dir)
.unwrap_or_else(|e| eprintln!("Failed to create pass artifact dir: {}", e));
save_matrix_to_tsv(
&eigensnp_output_current_pass
.final_snp_principal_component_loadings
.view(),
pass_artifact_dir.to_str().unwrap_or("."),
"loadings.tsv",
)
.unwrap_or_default();
save_matrix_to_tsv(
&eigensnp_output_current_pass
.final_sample_principal_component_scores
.view(),
pass_artifact_dir.to_str().unwrap_or("."),
"scores.tsv",
)
.unwrap_or_default();
save_vector_to_tsv(
&eigensnp_output_current_pass
.final_principal_component_eigenvalues
.view(),
pass_artifact_dir.to_str().unwrap_or("."),
"eigenvalues.tsv",
)
.unwrap_or_default();
let score_corr = evaluate_pc_score_correlation(
&eigensnp_output_current_pass,
&python_reference_output,
"ScoreCorr",
);
let loading_corr = evaluate_snp_loading_correlation(
&eigensnp_output_current_pass,
&python_reference_output,
"LoadingCorr",
);
let eigen_acc = evaluate_eigenvalue_accuracy(
&eigensnp_output_current_pass,
&python_reference_output,
"EigenAcc",
);
writeln!(
overall_outcome_details,
" PC Score Correlation: {:.6e}",
score_corr
)
.unwrap_or_default();
writeln!(
overall_outcome_details,
" SNP Loading Correlation: {:.6e}",
loading_corr
)
.unwrap_or_default();
writeln!(
overall_outcome_details,
" Eigenvalue Accuracy (-MSRE): {:.6e}",
eigen_acc
)
.unwrap_or_default();
if min_passes_found == -1 {
if score_corr >= thresholds.min_score_correlation
&& loading_corr >= thresholds.min_loading_correlation
&& eigen_acc >= thresholds.max_neg_eigenvalue_accuracy
{
min_passes_found = current_pass_count as i32;
writeln!(overall_outcome_details, " SUCCESS: All quality thresholds MET at {} pass(es). PCs in this run: {}.", current_pass_count, num_pcs_computed_at_convergence).unwrap_or_default();
} else {
writeln!(
overall_outcome_details,
" INFO: Quality thresholds NOT MET at {} pass(es).",
current_pass_count
)
.unwrap_or_default();
}
} else {
writeln!(overall_outcome_details, " INFO: Thresholds previously met at {} passes. Current pass {} metrics recorded.", min_passes_found, current_pass_count).unwrap_or_default();
if !(score_corr >= thresholds.min_score_correlation
&& loading_corr >= thresholds.min_loading_correlation
&& eigen_acc >= thresholds.max_neg_eigenvalue_accuracy)
{
writeln!(overall_outcome_details, " WARNING: Quality REGRESSED at {} passes after prior convergence at {} passes.", current_pass_count, min_passes_found).unwrap_or_default();
}
}
}
Err(e) => {
writeln!(
overall_outcome_details,
" FAILURE: EigenSnp compute_pca failed for {} pass(es): {}",
current_pass_count, e
)
.unwrap_or_default();
if min_passes_found != -1 {
writeln!(overall_outcome_details, " WARNING: PCA computation FAILED at {} passes after prior convergence at {} passes.", current_pass_count, min_passes_found).unwrap_or_default();
}
}
}
}
if min_passes_found == -1 {
writeln!(
overall_outcome_details,
"\n--- High quality NOT ACHIEVED within {} passes. ---",
max_passes_to_test
)
.unwrap_or_default();
} else {
writeln!(
overall_outcome_details,
"\n--- Minimum passes for convergence: {}. PCs computed in that run: {} ---",
min_passes_found, num_pcs_computed_at_convergence
)
.unwrap_or_default();
}
let expected_max_passes_for_convergence = 2;
let success = min_passes_found != -1 && min_passes_found <= expected_max_passes_for_convergence;
let record = TestResultRecord {
test_name: test_logging_name.to_string(),
num_features_d: d_total_snps,
num_samples_n: n_samples,
num_pcs_requested_k: k_true_components,
num_pcs_computed: num_pcs_computed_at_convergence,
success,
outcome_details: overall_outcome_details.clone(),
notes: format!(
"Min passes found for convergence: {}. Expected <= {}. Thresholds: ScoreCor >= {:.3}, LoadCor >= {:.3}, EigAcc (-MSRE) >= {:.3e}",
min_passes_found, expected_max_passes_for_convergence,
thresholds.min_score_correlation, thresholds.min_loading_correlation, thresholds.max_neg_eigenvalue_accuracy
),
};
TEST_RESULTS.lock().unwrap().push(record);
assert!(
success,
"Minimum passes for quality convergence test failed. Min passes found: {}. Expected <= {}. Details:\n{}",
min_passes_found, expected_max_passes_for_convergence, overall_outcome_details
);
}
#[test]
fn test_refinement_projection_accuracy() {
let test_logging_name = "test_refinement_projection_accuracy";
let d_total_snps = 5000;
let n_samples_total = 600;
let k_true_components = 10;
let n_samples_train = 500;
let n_samples_test = n_samples_total - n_samples_train;
let signal_strength = 3.0;
let noise_std_dev = 1.0;
let seed = 20241004;
let structured_standardized_data_total = generate_structured_data(
d_total_snps,
n_samples_total,
k_true_components,
signal_strength,
noise_std_dev,
seed,
);
let train_data = structured_standardized_data_total
.slice(s![.., 0..n_samples_train])
.to_owned();
let test_data_snps_x_samples = structured_standardized_data_total
.slice(s![.., n_samples_train..])
.to_owned();
let artifact_dir = Path::new("target/test_artifacts").join(test_logging_name);
fs::create_dir_all(&artifact_dir).unwrap_or_else(|e| {
panic!(
"Failed to create artifact directory '{}': {}",
artifact_dir.display(),
e
)
});
let k_components_to_request_pca = k_true_components + 5;
let python_total_pca_result = get_python_reference_pca(
&structured_standardized_data_total,
k_components_to_request_pca,
&format!("{}_py_ref_total", test_logging_name),
);
let (_py_total_loadings_k_x_d, py_total_scores_n_x_k, _py_total_eigenvalues_k) =
match python_total_pca_result {
Ok(output) => output,
Err(e) => {
panic!("Failed to get Python reference PCA on total data: {}", e);
}
};
if py_total_scores_n_x_k.nrows() < n_samples_total {
panic!(
"Python reference scores have fewer rows ({}) than n_samples_total ({}). Cannot slice test set scores.",
py_total_scores_n_x_k.nrows(), n_samples_total
);
}
if py_total_scores_n_x_k.ncols() == 0 && k_components_to_request_pca > 0 {
eprintln!("Warning: Python reference PCA on total data resulted in 0 components.");
}
let ref_projected_scores_for_test_set = py_total_scores_n_x_k
.slice(s![n_samples_train.., ..])
.to_owned();
save_matrix_to_tsv(
&ref_projected_scores_for_test_set.view(),
artifact_dir.to_str().unwrap(),
"python_ref_projected_test_scores.tsv",
)
.expect("Failed to save python_ref_projected_test_scores.tsv");
let run_eigensnp_and_project = |pass_count: usize,
run_tag: &str|
-> Result<(Array2<f32>, EigenSNPCoreOutput), String> {
let config = EigenSNPCoreAlgorithmConfig {
target_num_global_pcs: k_true_components, refine_pass_count: pass_count,
random_seed: seed,
subset_factor_for_local_basis_learning: 0.5,
min_subset_size_for_local_basis_learning: (n_samples_train / 4)
.max(1)
.min(n_samples_train.max(1)),
max_subset_size_for_local_basis_learning: (n_samples_train / 2)
.max(10)
.min(n_samples_train.max(1)),
components_per_ld_block: 10
.min(d_total_snps.min((n_samples_train / 2).max(10).min(n_samples_train.max(1)))),
..Default::default()
};
let test_data_accessor_train = TestDataAccessor::new(train_data.clone()); let ld_block_specs_train = vec![LdBlockSpecification {
user_defined_block_tag: "full_block_train".to_string(),
pca_snp_ids_in_block: (0..d_total_snps).map(PcaSnpId).collect(),
}];
let algorithm = EigenSNPCoreAlgorithm::new(config);
let snp_metadata = create_dummy_snp_metadata(d_total_snps);
match algorithm.compute_pca(
&test_data_accessor_train,
&ld_block_specs_train,
&snp_metadata,
) {
Ok((eigensnp_train_output_struct, _)) => {
save_matrix_to_tsv(
&eigensnp_train_output_struct
.final_snp_principal_component_loadings
.view(),
artifact_dir.to_str().unwrap(),
&format!("eigensnp_train_loadings_{}.tsv", run_tag),
)
.map_err(|e| format!("Failed to save train_loadings for {}: {}", run_tag, e))?;
save_matrix_to_tsv(
&eigensnp_train_output_struct
.final_sample_principal_component_scores
.view(),
artifact_dir.to_str().unwrap(),
&format!("eigensnp_train_scores_{}.tsv", run_tag),
)
.map_err(|e| format!("Failed to save train_scores for {}: {}", run_tag, e))?;
let projected_scores = if eigensnp_train_output_struct
.final_snp_principal_component_loadings
.ncols()
> 0
{
test_data_snps_x_samples
.t()
.dot(&eigensnp_train_output_struct.final_snp_principal_component_loadings)
} else {
Array2::zeros((n_samples_test, 0))
};
save_matrix_to_tsv(
&projected_scores.view(),
artifact_dir.to_str().unwrap(),
&format!("eigensnp_projected_test_scores_{}.tsv", run_tag),
)
.map_err(|e| format!("Failed to save projected_scores for {}: {}", run_tag, e))?;
Ok((projected_scores, eigensnp_train_output_struct))
}
Err(e) => Err(format!(
"EigenSnp compute_pca failed for {}: {}",
run_tag, e
)),
}
};
let (projected_scores_a, _output_a) =
run_eigensnp_and_project(1, "pass1") .unwrap_or_else(|e| panic!("Eigensnp run/projection A (pass1) failed: {}", e));
let (projected_scores_b, _output_b) =
run_eigensnp_and_project(2, "pass2") .unwrap_or_else(|e| panic!("Eigensnp run/projection B (pass2) failed: {}", e));
let calculate_avg_abs_correlation =
|projected_scores: &Array2<f32>, ref_scores: &Array2<f32>, k_compare_max: usize| -> f64 {
let k_compare = projected_scores
.ncols()
.min(ref_scores.ncols())
.min(k_compare_max);
if k_compare == 0 {
return 0.0;
}
let mut total_abs_corr = 0.0;
for i in 0..k_compare {
let proj_col = projected_scores.column(i);
let ref_col = ref_scores.column(i);
total_abs_corr += pearson_correlation(proj_col.view(), ref_col.view())
.map_or(0.0, |c| c.abs() as f64);
}
total_abs_corr / (k_compare as f64)
};
let corr_a = calculate_avg_abs_correlation(
&projected_scores_a,
&ref_projected_scores_for_test_set,
k_true_components,
);
let corr_b = calculate_avg_abs_correlation(
&projected_scores_b,
&ref_projected_scores_for_test_set,
k_true_components,
);
let success = corr_b >= corr_a - 1e-9;
let mut outcome_details = String::new();
writeln!(outcome_details, "Projection Accuracy Test Results:").unwrap_or_default();
writeln!(outcome_details, " Correlation A (1 pass): {:.6e}", corr_a).unwrap_or_default();
writeln!(
outcome_details,
" Correlation B (2 passes): {:.6e}",
corr_b
)
.unwrap_or_default();
if success {
writeln!(
outcome_details,
" Improvement OBSERVED or non-degradation within tolerance."
)
.unwrap_or_default();
} else {
writeln!(
outcome_details,
" Improvement NOT OBSERVED. Corr_B < Corr_A - tolerance."
)
.unwrap_or_default();
}
let num_pcs_computed_log = _output_b
.num_principal_components_computed
.max(_output_a.num_principal_components_computed);
let record = TestResultRecord {
test_name: test_logging_name.to_string(),
num_features_d: d_total_snps,
num_samples_n: n_samples_total, num_pcs_requested_k: k_true_components, num_pcs_computed: num_pcs_computed_log, success,
outcome_details: outcome_details.clone(),
notes: format!(
"Train_N={}, Test_N={}, Seed={}. Python_Ref_PCs_Requested={}. Corr_A={:.4e}, Corr_B={:.4e}.",
n_samples_train, n_samples_test, seed, k_components_to_request_pca, corr_a, corr_b
),
};
TEST_RESULTS.lock().unwrap().push(record);
assert!(
success,
"Projection accuracy did not improve with more refinement. Corr_A: {:.6e}, Corr_B: {:.6e}. Details: {}",
corr_a, corr_b, outcome_details
);
}
#[test]
fn test_reorder_columns_empty_matrix_variants() {
let original_0_rows = Array2::<i32>::zeros((0, 3));
let order = vec![1, 0, 2];
let expected_0_rows = Array2::<i32>::zeros((0, 3));
let reordered_0_rows = reorder_columns_owned(&original_0_rows, &order);
assert_eq!(reordered_0_rows, expected_0_rows);
let original_0_cols = Array2::<i32>::zeros((2, 0));
let order_empty = vec![];
let expected_0_cols_empty_order = Array2::<i32>::zeros((2, 0));
let reordered_empty_order = reorder_columns_owned(&original_0_cols, &order_empty);
assert_eq!(reordered_empty_order, expected_0_cols_empty_order);
}
#[test]
#[should_panic]
fn test_reorder_columns_select_from_zero_cols_panics() {
let original_0_cols = Array2::<i32>::zeros((2, 0));
let order_for_0_cols = vec![0];
reorder_columns_owned(&original_0_cols, &order_for_0_cols);
}
#[test]
fn test_reorder_columns_empty_order() {
let original = arr2(&[[1, 2, 3], [4, 5, 6]]);
let order = vec![];
let expected = Array2::<i32>::zeros((2, 0));
let reordered = reorder_columns_owned(&original, &order);
assert_eq!(reordered, expected);
}
#[test]
fn test_reorder_columns_repeated_indices() {
let original = arr2(&[[1, 2], [3, 4]]);
let order = vec![0, 1, 0, 0];
let expected = arr2(&[[1, 2, 1, 1], [3, 4, 3, 3]]);
let reordered = reorder_columns_owned(&original, &order);
assert_eq!(reordered, expected);
}
use std::sync::Arc;
fn create_dummy_snp_metadata(num_snps: usize) -> Vec<PcaSnpMetadata> {
(0..num_snps)
.map(|i| PcaSnpMetadata {
id: Arc::new(format!("snp_{}", i)),
chr: Arc::new("chr1".to_string()),
pos: i as u64 * 1000 + 100000, })
.collect()
}