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
13fn 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
25fn 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
68pub 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 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 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 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 let file = File::create(file_path)?;
139 let zstd_level = ZstdLevel::try_new(5)?; 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 fn to_parquet(&self, file_path: &str) -> anyhow::Result<()> {
182 self.to_melted_parquet(file_path, (None, None), (None, None))
183 }
184}
185
186pub 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 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 let first_param = ¶meters[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 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 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 let mut row_group_writer = writer.next_row_group()?;
273
274 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 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 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 let nrows = 3;
322 let ncols = 2;
323 let mut gamma = GammaMatrix::new((nrows, ncols), 2.0, 1.0);
324 gamma.calibrate();
325
326 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 let file = File::open(&file_path)?;
342 let reader = SerializedFileReader::new(file)?;
343 let iter = reader.get_row_iter(None)?;
344
345 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 assert_eq!(results.len(), nrows * ncols);
360
361 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 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 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 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 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())?; 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 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 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 ¶ms,
517 (Some(&row_names), None),
518 (Some(&col_names), None),
519 Some(&factor_names),
520 file_path_str,
521 )?;
522
523 let file = File::open(&file_path)?;
525 let reader = SerializedFileReader::new(file)?;
526 let iter = reader.get_row_iter(None)?;
527
528 #[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 assert_eq!(results.len(), nrows * ncols * n_factors);
549
550 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>(¶ms, (None, None), (None, None), None, file_path_str);
594 assert!(result.is_err());
595 assert!(result.unwrap_err().to_string().contains("empty"));
596 }
597}