Skip to main content

legume_numeric/param/
io.rs

1use crate::matrix::traits::{IoOps, MeltOps};
2use crate::param::traits::*;
3
4use parquet::basic::Type as ParquetType;
5use parquet::basic::{Compression, ConvertedType, ZstdLevel};
6use parquet::data_type::{ByteArray, ByteArrayType, FloatType};
7use parquet::file::properties::WriterProperties;
8use parquet::file::writer::SerializedFileWriter;
9use parquet::schema::types::Type;
10use std::fs::File;
11use std::sync::Arc;
12
13/// Pre-compute ByteArray lookup table for names.
14/// If names are provided, converts them to ByteArray.
15/// Otherwise, generates numeric strings "0", "1", "2", ... for the given count.
16fn precompute_name_bytes(names: Option<&[Box<str>]>, count: usize) -> Vec<ByteArray> {
17    match names {
18        Some(n) => n.iter().map(|s| ByteArray::from(s.as_ref())).collect(),
19        None => (0..count)
20            .map(|i| ByteArray::from(i.to_string().as_str()))
21            .collect(),
22    }
23}
24
25/// Build parquet schema for parameter matrices.
26/// If `include_factor` is true, includes a "factor" column between "column" and "mean".
27fn build_parquet_schema(
28    row_title: &str,
29    col_title: &str,
30    include_factor: bool,
31) -> anyhow::Result<Arc<Type>> {
32    let mut fields: Vec<(&str, ParquetType, ConvertedType)> = vec![
33        (row_title, ParquetType::BYTE_ARRAY, ConvertedType::UTF8),
34        (col_title, ParquetType::BYTE_ARRAY, ConvertedType::UTF8),
35    ];
36
37    if include_factor {
38        fields.push(("factor", ParquetType::BYTE_ARRAY, ConvertedType::UTF8));
39    }
40
41    fields.extend([
42        ("mean", ParquetType::FLOAT, ConvertedType::NONE),
43        ("sd", ParquetType::FLOAT, ConvertedType::NONE),
44        ("log_mean", ParquetType::FLOAT, ConvertedType::NONE),
45        ("log_sd", ParquetType::FLOAT, ConvertedType::NONE),
46    ]);
47
48    Ok(Arc::new(
49        Type::group_type_builder("GammaMatrix")
50            .with_fields(
51                fields
52                    .into_iter()
53                    .map(|(name, parquet_type, converted_type)| {
54                        Arc::new(
55                            Type::primitive_type_builder(name, parquet_type)
56                                .with_repetition(parquet::basic::Repetition::REQUIRED)
57                                .with_converted_type(converted_type)
58                                .build()
59                                .unwrap(),
60                        )
61                    })
62                    .collect(),
63            )
64            .build()?,
65    ))
66}
67
68/// consolidated input and output
69pub trait ParamIo: Inference
70where
71    f32: From<<<Self as Inference>::Mat as MeltOps>::Scalar>,
72{
73    type Mat: IoOps + MeltOps;
74
75    fn to_tsv(&self, header: &str) -> anyhow::Result<()> {
76        self.posterior_log_mean()
77            .to_tsv(&(header.to_string() + ".log_mean.gz"))?;
78
79        self.posterior_log_sd()
80            .to_tsv(&(header.to_string() + ".log_sd.gz"))?;
81
82        self.posterior_mean()
83            .to_tsv(&(header.to_string() + ".mean.gz"))?;
84
85        self.posterior_sd()
86            .to_tsv(&(header.to_string() + ".sd.gz"))?;
87
88        Ok(())
89    }
90
91    fn to_melted_parquet(
92        &self,
93        file_path: &str,
94        row_names: (Option<&[Box<str>]>, Option<&str>),
95        column_names: (Option<&[Box<str>]>, Option<&str>),
96    ) -> anyhow::Result<()> {
97        let row_names_slice = row_names.0;
98        let row_title = row_names.1.unwrap_or("row");
99        let col_title = column_names.1.unwrap_or("column");
100        let schema = build_parquet_schema(row_title, col_title, false)?;
101
102        // Pre-compute name ByteArrays once for efficient lookup
103        let row_bytes = precompute_name_bytes(row_names_slice, self.nrows());
104        let col_bytes = precompute_name_bytes(column_names.0, self.ncols());
105
106        // The mean plane defines the canonical (row, col) order and element
107        // count. Auxiliary planes (sd / log_mean / log_sd) may be lazily
108        // unallocated (0×0) when the parameter was only mean-calibrated
109        // (e.g. CalibrateTarget::MeanOnly); emit zeros of the right length in
110        // that case so serialization always succeeds — matching the pre-lazy
111        // behavior of writing zeros for never-computed planes.
112        let mat_mean = self.posterior_mean();
113        let (mean_scalars, idx) = mat_mean.melt_with_indexes();
114        let mean: Vec<f32> = mean_scalars.into_iter().map(|x| x.into()).collect();
115        let nelem = mean.len();
116        let melt_or_zeros = |m: &<Self as Inference>::Mat| -> Vec<f32> {
117            let v: Vec<f32> = m.melt().into_iter().map(|x| x.into()).collect();
118            if v.len() == nelem {
119                v
120            } else {
121                vec![0.0; nelem]
122            }
123        };
124        let sd = melt_or_zeros(self.posterior_sd());
125        let log_mean = melt_or_zeros(self.posterior_log_mean());
126        let log_sd = melt_or_zeros(self.posterior_log_sd());
127
128        // Map indices to pre-computed ByteArrays
129        let rows: Vec<_> = idx[0].iter().map(|&i| row_bytes[i].clone()).collect();
130        let cols: Vec<_> = idx[1].iter().map(|&i| col_bytes[i].clone()).collect();
131
132        let nelem = mean.len();
133        assert_eq!(nelem, sd.len());
134        assert_eq!(nelem, log_sd.len());
135        assert_eq!(nelem, log_mean.len());
136
137        // write data to parquet
138        let file = File::create(file_path)?;
139        let zstd_level = ZstdLevel::try_new(5)?; // Specify ZSTD compression level (e.g., 5)
140        let writer_properties = Arc::new(
141            WriterProperties::builder()
142                .set_compression(Compression::ZSTD(zstd_level))
143                .build(),
144        );
145        let mut writer = SerializedFileWriter::new(file, schema, writer_properties)?;
146
147        let mut row_group_writer = writer.next_row_group()?;
148
149        let name_columns = vec![&rows, &cols];
150
151        for data in name_columns {
152            if let Some(mut column_writer) = row_group_writer.next_column()? {
153                let typed_writer = column_writer.typed::<ByteArrayType>();
154                typed_writer.write_batch(data, None, None)?;
155                column_writer.close()?;
156            }
157        }
158
159        let val_columns: Vec<&[f32]> = vec![
160            mean.as_slice(),
161            sd.as_slice(),
162            log_mean.as_slice(),
163            log_sd.as_slice(),
164        ];
165
166        for data in val_columns {
167            if let Some(mut column_writer) = row_group_writer.next_column()? {
168                let typed_writer = column_writer.typed::<FloatType>();
169                typed_writer.write_batch(data, None, None)?;
170                column_writer.close()?;
171            }
172        }
173
174        row_group_writer.close()?;
175        writer.close()?;
176
177        Ok(())
178    }
179
180    /// Write to parquet with default names
181    fn to_parquet(&self, file_path: &str) -> anyhow::Result<()> {
182        self.to_melted_parquet(file_path, (None, None), (None, None))
183    }
184}
185
186/// Write down a vector of matrix parameters into one parquet file.
187///
188/// * `parameters`: a vector of row x column parameters (factors)
189/// * `row_names`: (values, optional title) — title defaults to "row"
190/// * `column_names`: (values, optional title) — title defaults to "column"
191/// * `factor_names`: a vector of factor names
192/// * `file_path`
193pub fn to_parquet<Param: Inference>(
194    parameters: &[Param],
195    row_names: (Option<&[Box<str>]>, Option<&str>),
196    column_names: (Option<&[Box<str>]>, Option<&str>),
197    factor_names: Option<&[Box<str>]>,
198    file_path: &str,
199) -> anyhow::Result<()>
200where
201    f32: From<<<Param as Inference>::Mat as MeltOps>::Scalar>,
202{
203    let factor_names: Vec<Box<str>> = match factor_names {
204        Some(x) => x.to_vec(),
205        _ => (0..parameters.len())
206            .map(|x| x.to_string().into_boxed_str())
207            .collect(),
208    };
209
210    if parameters.is_empty() {
211        return Err(anyhow::anyhow!("parameters cannot be empty"));
212    }
213
214    if factor_names.len() != parameters.len() {
215        return Err(anyhow::anyhow!(
216            "number of the parameters and factor names should match"
217        ));
218    }
219
220    let row_title = row_names.1.unwrap_or("row");
221    let col_title = column_names.1.unwrap_or("column");
222    let schema = build_parquet_schema(row_title, col_title, true)?;
223
224    // Write data to parquet
225    let file = File::create(file_path)?;
226    let zstd_level = ZstdLevel::try_new(5)?;
227    let writer_properties = Arc::new(
228        WriterProperties::builder()
229            .set_compression(Compression::ZSTD(zstd_level))
230            .build(),
231    );
232    let mut writer = SerializedFileWriter::new(file, schema, writer_properties)?;
233
234    // Pre-compute name ByteArrays once (reused across all factors)
235    let first_param = &parameters[0];
236    let row_bytes = precompute_name_bytes(row_names.0, first_param.nrows());
237    let col_bytes = precompute_name_bytes(column_names.0, first_param.ncols());
238
239    for (factor_idx, param) in parameters.iter().enumerate() {
240        // Mean defines the canonical order/element count; auxiliary planes may
241        // be lazily unallocated (0×0) under mean-only calibration, in which
242        // case we emit zeros of the right length so serialization succeeds.
243        let mat_mean = param.posterior_mean();
244        let (mean_scalars, idx) = mat_mean.melt_with_indexes();
245        let mean: Vec<f32> = mean_scalars.into_iter().map(|x| x.into()).collect();
246        let nelem = mean.len();
247        let melt_or_zeros = |m: &<Param as Inference>::Mat| -> Vec<f32> {
248            let v: Vec<f32> = m.melt().into_iter().map(|x| x.into()).collect();
249            if v.len() == nelem {
250                v
251            } else {
252                vec![0.0; nelem]
253            }
254        };
255        let sd = melt_or_zeros(param.posterior_sd());
256        let log_mean = melt_or_zeros(param.posterior_log_mean());
257        let log_sd = melt_or_zeros(param.posterior_log_sd());
258
259        let factor_name = factor_names[factor_idx].clone();
260        let factor_label = ByteArray::from(factor_name.as_bytes());
261
262        // Map indices to pre-computed ByteArrays
263        let rows: Vec<_> = idx[0].iter().map(|&i| row_bytes[i].clone()).collect();
264        let cols: Vec<_> = idx[1].iter().map(|&i| col_bytes[i].clone()).collect();
265
266        let nelem = mean.len();
267        assert_eq!(nelem, sd.len());
268        assert_eq!(nelem, log_sd.len());
269        assert_eq!(nelem, log_mean.len());
270
271        // Start a new row group for this inference
272        let mut row_group_writer = writer.next_row_group()?;
273
274        // Write the "inference", "row", and "column" columns
275        let name_columns = vec![rows, cols, vec![factor_label; nelem]];
276
277        for data in name_columns {
278            if let Some(mut column_writer) = row_group_writer.next_column()? {
279                let typed_writer = column_writer.typed::<ByteArrayType>();
280                typed_writer.write_batch(&data, None, None)?;
281                column_writer.close()?;
282            }
283        }
284
285        // Write the "mean", "sd", "log_mean", and "log_sd" columns
286        let val_columns: Vec<&[f32]> = vec![
287            mean.as_slice(),
288            sd.as_slice(),
289            log_mean.as_slice(),
290            log_sd.as_slice(),
291        ];
292
293        for data in val_columns {
294            if let Some(mut column_writer) = row_group_writer.next_column()? {
295                let typed_writer = column_writer.typed::<FloatType>();
296                typed_writer.write_batch(data, None, None)?;
297                column_writer.close()?;
298            }
299        }
300
301        row_group_writer.close()?;
302    }
303
304    // Close the writer
305    writer.close()?;
306
307    Ok(())
308}
309
310#[cfg(test)]
311mod tests {
312    use super::*;
313    use crate::param::dmatrix_gamma::GammaMatrix;
314    use parquet::file::reader::{FileReader, SerializedFileReader};
315    use parquet::record::RowAccessor;
316    use rustc_hash::FxHashMap as HashMap;
317
318    #[test]
319    fn test_param_io_to_parquet() -> anyhow::Result<()> {
320        // Create a small GammaMatrix
321        let nrows = 3;
322        let ncols = 2;
323        let mut gamma = GammaMatrix::new((nrows, ncols), 2.0, 1.0);
324        gamma.calibrate();
325
326        // Write to a temp file
327        let temp_dir = tempfile::tempdir()?;
328        let file_path = temp_dir.path().join("test_output.parquet");
329        let file_path_str = file_path.to_str().unwrap();
330
331        let row_names: Vec<Box<str>> = vec!["r0".into(), "r1".into(), "r2".into()];
332        let col_names: Vec<Box<str>> = vec!["c0".into(), "c1".into()];
333
334        gamma.to_melted_parquet(
335            file_path_str,
336            (Some(row_names.as_slice()), None),
337            (Some(col_names.as_slice()), None),
338        )?;
339
340        // Read back and verify
341        let file = File::open(&file_path)?;
342        let reader = SerializedFileReader::new(file)?;
343        let iter = reader.get_row_iter(None)?;
344
345        // Collect all rows into a map keyed by (row, col)
346        let mut results: HashMap<(String, String), (f32, f32, f32, f32)> = Default::default();
347        for row in iter {
348            let row = row?;
349            let row_name = row.get_string(0)?.to_string();
350            let col_name = row.get_string(1)?.to_string();
351            let mean = row.get_float(2)?;
352            let sd = row.get_float(3)?;
353            let log_mean = row.get_float(4)?;
354            let log_sd = row.get_float(5)?;
355            results.insert((row_name, col_name), (mean, sd, log_mean, log_sd));
356        }
357
358        // Should have nrows * ncols entries
359        assert_eq!(results.len(), nrows * ncols);
360
361        // Verify all row/col combinations exist
362        for r in &row_names {
363            for c in &col_names {
364                assert!(
365                    results.contains_key(&(r.to_string(), c.to_string())),
366                    "Missing entry for ({}, {})",
367                    r,
368                    c
369                );
370            }
371        }
372
373        // Verify values match the posterior estimates
374        let mean_mat = gamma.posterior_mean();
375        let sd_mat = gamma.posterior_sd();
376        let log_mean_mat = gamma.posterior_log_mean();
377        let log_sd_mat = gamma.posterior_log_sd();
378
379        for (ri, r) in row_names.iter().enumerate() {
380            for (ci, c) in col_names.iter().enumerate() {
381                let (mean, sd, log_mean, log_sd) =
382                    results.get(&(r.to_string(), c.to_string())).unwrap();
383
384                let expected_mean = mean_mat[(ri, ci)];
385                let expected_sd = sd_mat[(ri, ci)];
386                let expected_log_mean = log_mean_mat[(ri, ci)];
387                let expected_log_sd = log_sd_mat[(ri, ci)];
388
389                assert!(
390                    (mean - expected_mean).abs() < 1e-6,
391                    "mean mismatch at ({}, {}): {} vs {}",
392                    r,
393                    c,
394                    mean,
395                    expected_mean
396                );
397                assert!(
398                    (sd - expected_sd).abs() < 1e-6,
399                    "sd mismatch at ({}, {}): {} vs {}",
400                    r,
401                    c,
402                    sd,
403                    expected_sd
404                );
405                assert!(
406                    (log_mean - expected_log_mean).abs() < 1e-6,
407                    "log_mean mismatch at ({}, {}): {} vs {}",
408                    r,
409                    c,
410                    log_mean,
411                    expected_log_mean
412                );
413                assert!(
414                    (log_sd - expected_log_sd).abs() < 1e-6,
415                    "log_sd mismatch at ({}, {}): {} vs {}",
416                    r,
417                    c,
418                    log_sd,
419                    expected_log_sd
420                );
421            }
422        }
423
424        Ok(())
425    }
426
427    #[test]
428    fn test_param_io_to_parquet_without_names() -> anyhow::Result<()> {
429        // Test with numeric indices instead of names
430        let nrows = 2;
431        let ncols = 3;
432        let mut gamma = GammaMatrix::new((nrows, ncols), 1.5, 0.5);
433        gamma.calibrate();
434
435        let temp_dir = tempfile::tempdir()?;
436        let file_path = temp_dir.path().join("test_no_names.parquet");
437        let file_path_str = file_path.to_str().unwrap();
438
439        gamma.to_parquet(file_path_str)?;
440
441        // Read back and verify
442        let file = File::open(&file_path)?;
443        let reader = SerializedFileReader::new(file)?;
444        let iter = reader.get_row_iter(None)?;
445
446        let mut count = 0;
447        for row in iter {
448            let row = row?;
449            let row_idx: usize = row.get_string(0)?.parse()?;
450            let col_idx: usize = row.get_string(1)?.parse()?;
451
452            assert!(row_idx < nrows, "row index out of bounds: {}", row_idx);
453            assert!(col_idx < ncols, "col index out of bounds: {}", col_idx);
454
455            count += 1;
456        }
457
458        assert_eq!(count, nrows * ncols);
459
460        Ok(())
461    }
462
463    #[test]
464    fn mean_only_param_serializes_with_zero_aux_planes() -> anyhow::Result<()> {
465        // With lazy GammaMatrix, a MeanOnly-calibrated param leaves
466        // sd/log_mean/log_sd unallocated (0×0). to_parquet must still succeed,
467        // emitting zeros for those planes (regression guard for the lazy change).
468        let (nrows, ncols) = (3usize, 4usize);
469        let mut gamma = GammaMatrix::new((nrows, ncols), 2.0, 1.0);
470        gamma.calibrate_with(crate::param::traits::CalibrateTarget::MeanOnly);
471        assert_eq!(gamma.posterior_sd().nrows(), 0, "aux plane should be lazy");
472
473        let temp_dir = tempfile::tempdir()?;
474        let file_path = temp_dir.path().join("mean_only.parquet");
475        gamma.to_parquet(file_path.to_str().unwrap())?; // must not panic
476
477        let file = File::open(&file_path)?;
478        let reader = SerializedFileReader::new(file)?;
479        let mut count = 0;
480        for row in reader.get_row_iter(None)? {
481            let row = row?;
482            // schema: row, col, mean, sd, log_mean, log_sd
483            assert_eq!(row.get_float(3)?, 0.0, "sd should be zero under MeanOnly");
484            assert_eq!(row.get_float(4)?, 0.0, "log_mean should be zero");
485            assert_eq!(row.get_float(5)?, 0.0, "log_sd should be zero");
486            count += 1;
487        }
488        assert_eq!(count, nrows * ncols);
489        Ok(())
490    }
491
492    #[test]
493    fn test_to_parquet_multiple_factors() -> anyhow::Result<()> {
494        let nrows = 2;
495        let ncols = 2;
496        let n_factors = 3;
497
498        // Create multiple GammaMatrix parameters with different hyperparameters
499        let mut params: Vec<GammaMatrix> = Vec::new();
500        for i in 0..n_factors {
501            let mut gamma = GammaMatrix::new((nrows, ncols), 1.0 + i as f32, 0.5 + i as f32 * 0.1);
502            gamma.calibrate();
503            params.push(gamma);
504        }
505
506        let temp_dir = tempfile::tempdir()?;
507        let file_path = temp_dir.path().join("test_multi_factor.parquet");
508        let file_path_str = file_path.to_str().unwrap();
509
510        let row_names: Vec<Box<str>> = vec!["gene1".into(), "gene2".into()];
511        let col_names: Vec<Box<str>> = vec!["cell1".into(), "cell2".into()];
512        let factor_names: Vec<Box<str>> =
513            vec!["factor0".into(), "factor1".into(), "factor2".into()];
514
515        to_parquet(
516            &params,
517            (Some(&row_names), None),
518            (Some(&col_names), None),
519            Some(&factor_names),
520            file_path_str,
521        )?;
522
523        // Read back and verify
524        let file = File::open(&file_path)?;
525        let reader = SerializedFileReader::new(file)?;
526        let iter = reader.get_row_iter(None)?;
527
528        // Collect results keyed by (row, col, factor)
529        #[allow(clippy::type_complexity)]
530        let mut results: HashMap<(String, String, String), (f32, f32, f32, f32)> =
531            Default::default();
532        for row in iter {
533            let row = row?;
534            let row_name = row.get_string(0)?.to_string();
535            let col_name = row.get_string(1)?.to_string();
536            let factor_name = row.get_string(2)?.to_string();
537            let mean = row.get_float(3)?;
538            let sd = row.get_float(4)?;
539            let log_mean = row.get_float(5)?;
540            let log_sd = row.get_float(6)?;
541            results.insert(
542                (row_name, col_name, factor_name),
543                (mean, sd, log_mean, log_sd),
544            );
545        }
546
547        // Should have nrows * ncols * n_factors entries
548        assert_eq!(results.len(), nrows * ncols * n_factors);
549
550        // Verify values for each factor
551        for (fi, param) in params.iter().enumerate() {
552            let factor = &factor_names[fi];
553            let mean_mat = param.posterior_mean();
554            let sd_mat = param.posterior_sd();
555
556            for (ri, r) in row_names.iter().enumerate() {
557                for (ci, c) in col_names.iter().enumerate() {
558                    let key = (r.to_string(), c.to_string(), factor.to_string());
559                    let (mean, sd, _, _) = results.get(&key).expect("Missing entry");
560
561                    let expected_mean = mean_mat[(ri, ci)];
562                    let expected_sd = sd_mat[(ri, ci)];
563
564                    assert!(
565                        (mean - expected_mean).abs() < 1e-6,
566                        "mean mismatch for factor {} at ({}, {})",
567                        factor,
568                        r,
569                        c
570                    );
571                    assert!(
572                        (sd - expected_sd).abs() < 1e-6,
573                        "sd mismatch for factor {} at ({}, {})",
574                        factor,
575                        r,
576                        c
577                    );
578                }
579            }
580        }
581
582        Ok(())
583    }
584
585    #[test]
586    fn test_to_parquet_empty_parameters() {
587        let params: Vec<GammaMatrix> = vec![];
588        let temp_dir = tempfile::tempdir().unwrap();
589        let file_path = temp_dir.path().join("test_empty.parquet");
590        let file_path_str = file_path.to_str().unwrap();
591
592        let result =
593            to_parquet::<GammaMatrix>(&params, (None, None), (None, None), None, file_path_str);
594        assert!(result.is_err());
595        assert!(result.unwrap_err().to_string().contains("empty"));
596    }
597}