//! Stain vector estimation for histopathology images //! //! Provides standard stain vectors and estimation methods for //! common staining protocols like H&E and IHC-DAB. use crate::config::EstimationMethod; use crate::optical_density::{extract_foreground_pixels, get_tissue_mask}; use anyhow::{Result, bail}; use ndarray::{Array1, Array2, Array3, Axis}; use slidescope_shared::StainType; /// Standard H&E stain vectors (Ruifrok & Johnston, 2001) /// /// Rows are RGB components, columns are stains: /// - Column 0: Hematoxylin (blue-purple) /// - Column 1: Eosin (pink-red) /// - Column 2: Residual/background pub fn he_stain_matrix() -> Array2 { let mut matrix = Array2::from_shape_vec( (3, 3), vec![ 0.650, 0.072, 0.268, // R: Hematoxylin, Eosin, Residual 0.704, 0.990, 0.570, // G 0.286, 0.105, 0.776, // B ], ) .unwrap(); // Normalize each column normalize_columns(&mut matrix); matrix } /// Standard IHC-DAB stain vectors /// /// - Column 0: Hematoxylin (counterstain) /// - Column 1: DAB (brown chromogen) /// - Column 2: Residual pub fn ihc_dab_stain_matrix() -> Array2 { let mut matrix = Array2::from_shape_vec( (3, 3), vec![ 0.650, 0.268, 0.0, // R 0.704, 0.570, 0.0, // G 0.286, 0.776, 0.0, // B ], ) .unwrap(); normalize_columns(&mut matrix); matrix } /// Normalize each column of the matrix to unit length fn normalize_columns(matrix: &mut Array2) { for mut col in matrix.columns_mut() { let norm: f32 = col.iter().map(|x| x * x).sum::().sqrt(); if norm > 0.0 { col.mapv_inplace(|x| x / norm); } } } /// Get standard stain matrix for a stain type pub fn get_standard_matrix(stain_type: StainType) -> Array2 { match stain_type { StainType::HE => he_stain_matrix(), StainType::IhcDab => ihc_dab_stain_matrix(), StainType::Custom => Array2::eye(3), // Identity matrix for custom } } /// Result of stain vector estimation #[derive(Debug, Clone)] pub struct StainEstimationResult { /// Estimated stain matrix (3 x n_components) pub stain_matrix: Array2, /// Quality score (0.0 - 1.0) pub quality_score: f32, /// Detected stain type pub stain_type: StainType, /// Number of samples used for estimation pub n_samples: usize, } /// Stain vector estimator pub struct StainEstimator { method: EstimationMethod, background_threshold: f32, sample_size: Option, } impl StainEstimator { /// Create a new stain estimator pub fn new( method: EstimationMethod, background_threshold: f32, sample_size: Option, ) -> Self { Self { method, background_threshold, sample_size, } } /// Estimate stain vectors from an optical density image /// /// # Arguments /// * `od_image` - Optical density image (3 channels) /// * `n_components` - Number of stain components to extract pub fn estimate( &self, od_image: &Array3, n_components: usize, ) -> Result { if !(2..=3).contains(&n_components) { bail!("n_components must be 2 or 3"); } match self.method { EstimationMethod::Ruifrok => self.estimate_ruifrok(n_components), EstimationMethod::Macenko => self.estimate_macenko(od_image, n_components), EstimationMethod::Robust => self.estimate_robust(od_image, n_components), } } /// Use pre-defined Ruifrok vectors fn estimate_ruifrok(&self, n_components: usize) -> Result { let full_matrix = he_stain_matrix(); // Extract only the requested number of components let stain_matrix = full_matrix .slice(ndarray::s![.., 0..n_components]) .to_owned(); Ok(StainEstimationResult { stain_matrix, quality_score: 0.95, // High confidence for standard vectors stain_type: StainType::HE, n_samples: 0, }) } /// Macenko's method using PCA fn estimate_macenko( &self, od_image: &Array3, n_components: usize, ) -> Result { // Get tissue mask let mask = get_tissue_mask(od_image, self.background_threshold); // Extract foreground pixels let pixels = extract_foreground_pixels(od_image, &mask); let n_samples = pixels.nrows(); if n_samples < 100 { // Fall back to standard vectors if not enough tissue return self.estimate_ruifrok(n_components); } // Optionally subsample for speed let samples = if let Some(max_samples) = self.sample_size { if n_samples > max_samples { subsample_rows(&pixels, max_samples) } else { pixels } } else { pixels }; // Compute covariance matrix let cov = compute_covariance(&samples)?; // Compute eigenvectors (simple power iteration for top 2) let eigenvectors = power_iteration_top_k(&cov, n_components)?; // Build stain matrix let mut stain_matrix = Array2::zeros((3, n_components)); for (i, eigvec) in eigenvectors.iter().enumerate() { stain_matrix.column_mut(i).assign(eigvec); } // Ensure positive orientation for mut col in stain_matrix.columns_mut() { if col.sum() < 0.0 { col.mapv_inplace(|x| -x); } } normalize_columns(&mut stain_matrix); // Compute quality score based on eigenvalue separation let quality_score = 0.8; // Simplified Ok(StainEstimationResult { stain_matrix, quality_score, stain_type: StainType::Custom, n_samples: samples.nrows(), }) } /// Robust estimation combining methods fn estimate_robust( &self, od_image: &Array3, n_components: usize, ) -> Result { // Try Macenko first let macenko_result = self.estimate_macenko(od_image, n_components)?; // Get standard Ruifrok vectors let ruifrok_result = self.estimate_ruifrok(n_components)?; // Compare with standard and choose best let similarity = matrix_similarity(&macenko_result.stain_matrix, &ruifrok_result.stain_matrix); if similarity > 0.7 && macenko_result.quality_score > 0.6 { // Macenko result is consistent with standard Ok(macenko_result) } else { // Fall back to standard vectors Ok(ruifrok_result) } } } /// Subsample rows from a matrix fn subsample_rows(matrix: &Array2, n_samples: usize) -> Array2 { use rand::seq::SliceRandom; let mut rng = rand::thread_rng(); let n_rows = matrix.nrows(); let mut indices: Vec = (0..n_rows).collect(); indices.shuffle(&mut rng); indices.truncate(n_samples); let mut result = Array2::zeros((n_samples, matrix.ncols())); for (i, &idx) in indices.iter().enumerate() { result.row_mut(i).assign(&matrix.row(idx)); } result } /// Compute covariance matrix fn compute_covariance(data: &Array2) -> Result> { let n_samples = data.nrows() as f32; if n_samples < 2.0 { bail!("Need at least 2 samples for covariance"); } let mean = data.mean_axis(Axis(0)).unwrap(); let centered = data - &mean; let cov = centered.t().dot(¢ered) / (n_samples - 1.0); Ok(cov) } /// Power iteration to find top k eigenvectors fn power_iteration_top_k(matrix: &Array2, k: usize) -> Result>> { let n = matrix.nrows(); let mut eigenvectors = Vec::with_capacity(k); let mut deflated = matrix.clone(); for _ in 0..k { // Random initial vector let mut v = Array1::from_iter((0..n).map(|i| (i as f32 + 1.0).sin())); // Power iteration for _ in 0..100 { let new_v = deflated.dot(&v); let norm: f32 = new_v.iter().map(|x| x * x).sum::().sqrt(); if norm < 1e-10 { break; } v = new_v / norm; } eigenvectors.push(v.clone()); // Deflate matrix let eigenvalue = v.dot(&deflated.dot(&v)); let outer = outer_product(&v, &v); deflated = deflated - eigenvalue * outer; } Ok(eigenvectors) } /// Compute outer product of two vectors fn outer_product(a: &Array1, b: &Array1) -> Array2 { let n = a.len(); let m = b.len(); let mut result = Array2::zeros((n, m)); for i in 0..n { for j in 0..m { result[[i, j]] = a[i] * b[j]; } } result } /// Compute similarity between two matrices (cosine of angle) fn matrix_similarity(a: &Array2, b: &Array2) -> f32 { if a.shape() != b.shape() { return 0.0; } let dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum(); let norm_a: f32 = a.iter().map(|x| x * x).sum::().sqrt(); let norm_b: f32 = b.iter().map(|x| x * x).sum::().sqrt(); if norm_a < 1e-10 || norm_b < 1e-10 { return 0.0; } (dot / (norm_a * norm_b)).abs() } #[cfg(test)] mod tests { use super::*; #[test] fn test_he_matrix_normalized() { let matrix = he_stain_matrix(); // Each column should have unit norm for col in matrix.columns() { let norm: f32 = col.iter().map(|x| x * x).sum::().sqrt(); assert!((norm - 1.0).abs() < 1e-5); } } #[test] fn test_ruifrok_estimation() { let estimator = StainEstimator::new(EstimationMethod::Ruifrok, 0.8, None); let result = estimator.estimate_ruifrok(2).unwrap(); assert_eq!(result.stain_matrix.shape(), &[3, 2]); assert!(result.quality_score > 0.9); } #[test] fn test_covariance() { let data = Array2::from_shape_vec( (4, 3), vec![1.0, 2.0, 3.0, 2.0, 3.0, 4.0, 3.0, 4.0, 5.0, 4.0, 5.0, 6.0], ) .unwrap(); let cov = compute_covariance(&data).unwrap(); assert_eq!(cov.shape(), &[3, 3]); // Covariance matrix should be symmetric for i in 0..3 { for j in 0..3 { assert!((cov[[i, j]] - cov[[j, i]]).abs() < 1e-5); } } } }