use crate::matrix::traits::{IoOps, MeltOps};
use crate::param::traits::*;
use parquet::basic::Type as ParquetType;
use parquet::basic::{Compression, ConvertedType, ZstdLevel};
use parquet::data_type::{ByteArray, ByteArrayType, FloatType};
use parquet::file::properties::WriterProperties;
use parquet::file::writer::SerializedFileWriter;
use parquet::schema::types::Type;
use std::fs::File;
use std::sync::Arc;
fn precompute_name_bytes(names: Option<&[Box<str>]>, count: usize) -> Vec<ByteArray> {
match names {
Some(n) => n.iter().map(|s| ByteArray::from(s.as_ref())).collect(),
None => (0..count)
.map(|i| ByteArray::from(i.to_string().as_str()))
.collect(),
}
}
fn build_parquet_schema(
row_title: &str,
col_title: &str,
include_factor: bool,
) -> anyhow::Result<Arc<Type>> {
let mut fields: Vec<(&str, ParquetType, ConvertedType)> = vec![
(row_title, ParquetType::BYTE_ARRAY, ConvertedType::UTF8),
(col_title, ParquetType::BYTE_ARRAY, ConvertedType::UTF8),
];
if include_factor {
fields.push(("factor", ParquetType::BYTE_ARRAY, ConvertedType::UTF8));
}
fields.extend([
("mean", ParquetType::FLOAT, ConvertedType::NONE),
("sd", ParquetType::FLOAT, ConvertedType::NONE),
("log_mean", ParquetType::FLOAT, ConvertedType::NONE),
("log_sd", ParquetType::FLOAT, ConvertedType::NONE),
]);
Ok(Arc::new(
Type::group_type_builder("GammaMatrix")
.with_fields(
fields
.into_iter()
.map(|(name, parquet_type, converted_type)| {
Arc::new(
Type::primitive_type_builder(name, parquet_type)
.with_repetition(parquet::basic::Repetition::REQUIRED)
.with_converted_type(converted_type)
.build()
.unwrap(),
)
})
.collect(),
)
.build()?,
))
}
pub trait ParamIo: Inference
where
f32: From<<<Self as Inference>::Mat as MeltOps>::Scalar>,
{
type Mat: IoOps + MeltOps;
fn to_tsv(&self, header: &str) -> anyhow::Result<()> {
self.posterior_log_mean()
.to_tsv(&(header.to_string() + ".log_mean.gz"))?;
self.posterior_log_sd()
.to_tsv(&(header.to_string() + ".log_sd.gz"))?;
self.posterior_mean()
.to_tsv(&(header.to_string() + ".mean.gz"))?;
self.posterior_sd()
.to_tsv(&(header.to_string() + ".sd.gz"))?;
Ok(())
}
fn to_melted_parquet(
&self,
file_path: &str,
row_names: (Option<&[Box<str>]>, Option<&str>),
column_names: (Option<&[Box<str>]>, Option<&str>),
) -> anyhow::Result<()> {
let row_names_slice = row_names.0;
let row_title = row_names.1.unwrap_or("row");
let col_title = column_names.1.unwrap_or("column");
let schema = build_parquet_schema(row_title, col_title, false)?;
let row_bytes = precompute_name_bytes(row_names_slice, self.nrows());
let col_bytes = precompute_name_bytes(column_names.0, self.ncols());
let mat_mean = self.posterior_mean();
let (mean_scalars, idx) = mat_mean.melt_with_indexes();
let mean: Vec<f32> = mean_scalars.into_iter().map(|x| x.into()).collect();
let nelem = mean.len();
let melt_or_zeros = |m: &<Self as Inference>::Mat| -> Vec<f32> {
let v: Vec<f32> = m.melt().into_iter().map(|x| x.into()).collect();
if v.len() == nelem {
v
} else {
vec![0.0; nelem]
}
};
let sd = melt_or_zeros(self.posterior_sd());
let log_mean = melt_or_zeros(self.posterior_log_mean());
let log_sd = melt_or_zeros(self.posterior_log_sd());
let rows: Vec<_> = idx[0].iter().map(|&i| row_bytes[i].clone()).collect();
let cols: Vec<_> = idx[1].iter().map(|&i| col_bytes[i].clone()).collect();
let nelem = mean.len();
assert_eq!(nelem, sd.len());
assert_eq!(nelem, log_sd.len());
assert_eq!(nelem, log_mean.len());
let file = File::create(file_path)?;
let zstd_level = ZstdLevel::try_new(5)?; let writer_properties = Arc::new(
WriterProperties::builder()
.set_compression(Compression::ZSTD(zstd_level))
.build(),
);
let mut writer = SerializedFileWriter::new(file, schema, writer_properties)?;
let mut row_group_writer = writer.next_row_group()?;
let name_columns = vec![&rows, &cols];
for data in name_columns {
if let Some(mut column_writer) = row_group_writer.next_column()? {
let typed_writer = column_writer.typed::<ByteArrayType>();
typed_writer.write_batch(data, None, None)?;
column_writer.close()?;
}
}
let val_columns: Vec<&[f32]> = vec![
mean.as_slice(),
sd.as_slice(),
log_mean.as_slice(),
log_sd.as_slice(),
];
for data in val_columns {
if let Some(mut column_writer) = row_group_writer.next_column()? {
let typed_writer = column_writer.typed::<FloatType>();
typed_writer.write_batch(data, None, None)?;
column_writer.close()?;
}
}
row_group_writer.close()?;
writer.close()?;
Ok(())
}
fn to_parquet(&self, file_path: &str) -> anyhow::Result<()> {
self.to_melted_parquet(file_path, (None, None), (None, None))
}
}
pub fn to_parquet<Param: Inference>(
parameters: &[Param],
row_names: (Option<&[Box<str>]>, Option<&str>),
column_names: (Option<&[Box<str>]>, Option<&str>),
factor_names: Option<&[Box<str>]>,
file_path: &str,
) -> anyhow::Result<()>
where
f32: From<<<Param as Inference>::Mat as MeltOps>::Scalar>,
{
let factor_names: Vec<Box<str>> = match factor_names {
Some(x) => x.to_vec(),
_ => (0..parameters.len())
.map(|x| x.to_string().into_boxed_str())
.collect(),
};
if parameters.is_empty() {
return Err(anyhow::anyhow!("parameters cannot be empty"));
}
if factor_names.len() != parameters.len() {
return Err(anyhow::anyhow!(
"number of the parameters and factor names should match"
));
}
let row_title = row_names.1.unwrap_or("row");
let col_title = column_names.1.unwrap_or("column");
let schema = build_parquet_schema(row_title, col_title, true)?;
let file = File::create(file_path)?;
let zstd_level = ZstdLevel::try_new(5)?;
let writer_properties = Arc::new(
WriterProperties::builder()
.set_compression(Compression::ZSTD(zstd_level))
.build(),
);
let mut writer = SerializedFileWriter::new(file, schema, writer_properties)?;
let first_param = ¶meters[0];
let row_bytes = precompute_name_bytes(row_names.0, first_param.nrows());
let col_bytes = precompute_name_bytes(column_names.0, first_param.ncols());
for (factor_idx, param) in parameters.iter().enumerate() {
let mat_mean = param.posterior_mean();
let (mean_scalars, idx) = mat_mean.melt_with_indexes();
let mean: Vec<f32> = mean_scalars.into_iter().map(|x| x.into()).collect();
let nelem = mean.len();
let melt_or_zeros = |m: &<Param as Inference>::Mat| -> Vec<f32> {
let v: Vec<f32> = m.melt().into_iter().map(|x| x.into()).collect();
if v.len() == nelem {
v
} else {
vec![0.0; nelem]
}
};
let sd = melt_or_zeros(param.posterior_sd());
let log_mean = melt_or_zeros(param.posterior_log_mean());
let log_sd = melt_or_zeros(param.posterior_log_sd());
let factor_name = factor_names[factor_idx].clone();
let factor_label = ByteArray::from(factor_name.as_bytes());
let rows: Vec<_> = idx[0].iter().map(|&i| row_bytes[i].clone()).collect();
let cols: Vec<_> = idx[1].iter().map(|&i| col_bytes[i].clone()).collect();
let nelem = mean.len();
assert_eq!(nelem, sd.len());
assert_eq!(nelem, log_sd.len());
assert_eq!(nelem, log_mean.len());
let mut row_group_writer = writer.next_row_group()?;
let name_columns = vec![rows, cols, vec![factor_label; nelem]];
for data in name_columns {
if let Some(mut column_writer) = row_group_writer.next_column()? {
let typed_writer = column_writer.typed::<ByteArrayType>();
typed_writer.write_batch(&data, None, None)?;
column_writer.close()?;
}
}
let val_columns: Vec<&[f32]> = vec![
mean.as_slice(),
sd.as_slice(),
log_mean.as_slice(),
log_sd.as_slice(),
];
for data in val_columns {
if let Some(mut column_writer) = row_group_writer.next_column()? {
let typed_writer = column_writer.typed::<FloatType>();
typed_writer.write_batch(data, None, None)?;
column_writer.close()?;
}
}
row_group_writer.close()?;
}
writer.close()?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::param::dmatrix_gamma::GammaMatrix;
use parquet::file::reader::{FileReader, SerializedFileReader};
use parquet::record::RowAccessor;
use rustc_hash::FxHashMap as HashMap;
#[test]
fn test_param_io_to_parquet() -> anyhow::Result<()> {
let nrows = 3;
let ncols = 2;
let mut gamma = GammaMatrix::new((nrows, ncols), 2.0, 1.0);
gamma.calibrate();
let temp_dir = tempfile::tempdir()?;
let file_path = temp_dir.path().join("test_output.parquet");
let file_path_str = file_path.to_str().unwrap();
let row_names: Vec<Box<str>> = vec!["r0".into(), "r1".into(), "r2".into()];
let col_names: Vec<Box<str>> = vec!["c0".into(), "c1".into()];
gamma.to_melted_parquet(
file_path_str,
(Some(row_names.as_slice()), None),
(Some(col_names.as_slice()), None),
)?;
let file = File::open(&file_path)?;
let reader = SerializedFileReader::new(file)?;
let iter = reader.get_row_iter(None)?;
let mut results: HashMap<(String, String), (f32, f32, f32, f32)> = Default::default();
for row in iter {
let row = row?;
let row_name = row.get_string(0)?.to_string();
let col_name = row.get_string(1)?.to_string();
let mean = row.get_float(2)?;
let sd = row.get_float(3)?;
let log_mean = row.get_float(4)?;
let log_sd = row.get_float(5)?;
results.insert((row_name, col_name), (mean, sd, log_mean, log_sd));
}
assert_eq!(results.len(), nrows * ncols);
for r in &row_names {
for c in &col_names {
assert!(
results.contains_key(&(r.to_string(), c.to_string())),
"Missing entry for ({}, {})",
r,
c
);
}
}
let mean_mat = gamma.posterior_mean();
let sd_mat = gamma.posterior_sd();
let log_mean_mat = gamma.posterior_log_mean();
let log_sd_mat = gamma.posterior_log_sd();
for (ri, r) in row_names.iter().enumerate() {
for (ci, c) in col_names.iter().enumerate() {
let (mean, sd, log_mean, log_sd) =
results.get(&(r.to_string(), c.to_string())).unwrap();
let expected_mean = mean_mat[(ri, ci)];
let expected_sd = sd_mat[(ri, ci)];
let expected_log_mean = log_mean_mat[(ri, ci)];
let expected_log_sd = log_sd_mat[(ri, ci)];
assert!(
(mean - expected_mean).abs() < 1e-6,
"mean mismatch at ({}, {}): {} vs {}",
r,
c,
mean,
expected_mean
);
assert!(
(sd - expected_sd).abs() < 1e-6,
"sd mismatch at ({}, {}): {} vs {}",
r,
c,
sd,
expected_sd
);
assert!(
(log_mean - expected_log_mean).abs() < 1e-6,
"log_mean mismatch at ({}, {}): {} vs {}",
r,
c,
log_mean,
expected_log_mean
);
assert!(
(log_sd - expected_log_sd).abs() < 1e-6,
"log_sd mismatch at ({}, {}): {} vs {}",
r,
c,
log_sd,
expected_log_sd
);
}
}
Ok(())
}
#[test]
fn test_param_io_to_parquet_without_names() -> anyhow::Result<()> {
let nrows = 2;
let ncols = 3;
let mut gamma = GammaMatrix::new((nrows, ncols), 1.5, 0.5);
gamma.calibrate();
let temp_dir = tempfile::tempdir()?;
let file_path = temp_dir.path().join("test_no_names.parquet");
let file_path_str = file_path.to_str().unwrap();
gamma.to_parquet(file_path_str)?;
let file = File::open(&file_path)?;
let reader = SerializedFileReader::new(file)?;
let iter = reader.get_row_iter(None)?;
let mut count = 0;
for row in iter {
let row = row?;
let row_idx: usize = row.get_string(0)?.parse()?;
let col_idx: usize = row.get_string(1)?.parse()?;
assert!(row_idx < nrows, "row index out of bounds: {}", row_idx);
assert!(col_idx < ncols, "col index out of bounds: {}", col_idx);
count += 1;
}
assert_eq!(count, nrows * ncols);
Ok(())
}
#[test]
fn mean_only_param_serializes_with_zero_aux_planes() -> anyhow::Result<()> {
let (nrows, ncols) = (3usize, 4usize);
let mut gamma = GammaMatrix::new((nrows, ncols), 2.0, 1.0);
gamma.calibrate_with(crate::param::traits::CalibrateTarget::MeanOnly);
assert_eq!(gamma.posterior_sd().nrows(), 0, "aux plane should be lazy");
let temp_dir = tempfile::tempdir()?;
let file_path = temp_dir.path().join("mean_only.parquet");
gamma.to_parquet(file_path.to_str().unwrap())?;
let file = File::open(&file_path)?;
let reader = SerializedFileReader::new(file)?;
let mut count = 0;
for row in reader.get_row_iter(None)? {
let row = row?;
assert_eq!(row.get_float(3)?, 0.0, "sd should be zero under MeanOnly");
assert_eq!(row.get_float(4)?, 0.0, "log_mean should be zero");
assert_eq!(row.get_float(5)?, 0.0, "log_sd should be zero");
count += 1;
}
assert_eq!(count, nrows * ncols);
Ok(())
}
#[test]
fn test_to_parquet_multiple_factors() -> anyhow::Result<()> {
let nrows = 2;
let ncols = 2;
let n_factors = 3;
let mut params: Vec<GammaMatrix> = Vec::new();
for i in 0..n_factors {
let mut gamma = GammaMatrix::new((nrows, ncols), 1.0 + i as f32, 0.5 + i as f32 * 0.1);
gamma.calibrate();
params.push(gamma);
}
let temp_dir = tempfile::tempdir()?;
let file_path = temp_dir.path().join("test_multi_factor.parquet");
let file_path_str = file_path.to_str().unwrap();
let row_names: Vec<Box<str>> = vec!["gene1".into(), "gene2".into()];
let col_names: Vec<Box<str>> = vec!["cell1".into(), "cell2".into()];
let factor_names: Vec<Box<str>> =
vec!["factor0".into(), "factor1".into(), "factor2".into()];
to_parquet(
¶ms,
(Some(&row_names), None),
(Some(&col_names), None),
Some(&factor_names),
file_path_str,
)?;
let file = File::open(&file_path)?;
let reader = SerializedFileReader::new(file)?;
let iter = reader.get_row_iter(None)?;
#[allow(clippy::type_complexity)]
let mut results: HashMap<(String, String, String), (f32, f32, f32, f32)> =
Default::default();
for row in iter {
let row = row?;
let row_name = row.get_string(0)?.to_string();
let col_name = row.get_string(1)?.to_string();
let factor_name = row.get_string(2)?.to_string();
let mean = row.get_float(3)?;
let sd = row.get_float(4)?;
let log_mean = row.get_float(5)?;
let log_sd = row.get_float(6)?;
results.insert(
(row_name, col_name, factor_name),
(mean, sd, log_mean, log_sd),
);
}
assert_eq!(results.len(), nrows * ncols * n_factors);
for (fi, param) in params.iter().enumerate() {
let factor = &factor_names[fi];
let mean_mat = param.posterior_mean();
let sd_mat = param.posterior_sd();
for (ri, r) in row_names.iter().enumerate() {
for (ci, c) in col_names.iter().enumerate() {
let key = (r.to_string(), c.to_string(), factor.to_string());
let (mean, sd, _, _) = results.get(&key).expect("Missing entry");
let expected_mean = mean_mat[(ri, ci)];
let expected_sd = sd_mat[(ri, ci)];
assert!(
(mean - expected_mean).abs() < 1e-6,
"mean mismatch for factor {} at ({}, {})",
factor,
r,
c
);
assert!(
(sd - expected_sd).abs() < 1e-6,
"sd mismatch for factor {} at ({}, {})",
factor,
r,
c
);
}
}
}
Ok(())
}
#[test]
fn test_to_parquet_empty_parameters() {
let params: Vec<GammaMatrix> = vec![];
let temp_dir = tempfile::tempdir().unwrap();
let file_path = temp_dir.path().join("test_empty.parquet");
let file_path_str = file_path.to_str().unwrap();
let result =
to_parquet::<GammaMatrix>(¶ms, (None, None), (None, None), None, file_path_str);
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("empty"));
}
}