use std::path::Path;
use std::sync::Arc;
use arrow::array::{Float64Array, Int32Array, StringBuilder, UInt32Array, UInt64Array};
use arrow::record_batch::RecordBatch;
use parquet::arrow::ArrowWriter;
use parquet::basic::{Compression, ZstdLevel};
use parquet::file::properties::WriterProperties;
use super::error::OutputError;
use super::schemas::{retry_histogram_schema, solver_iterations_schema};
#[derive(Debug, Clone)]
pub struct SolverStatsRow {
pub iteration: u32,
pub phase: String,
pub stage: i32,
pub opening: Option<i32>,
pub rank: Option<i32>,
pub worker_id: Option<i32>,
pub lp_solves: u32,
pub lp_successes: u32,
pub lp_retries: u32,
pub lp_failures: u32,
pub retry_attempts: u32,
pub basis_offered: u32,
pub basis_consistency_failures: u32,
pub simplex_iterations: u64,
pub solve_time_ms: f64,
pub load_model_time_ms: f64,
pub set_bounds_time_ms: f64,
pub basis_set_time_ms: f64,
pub basis_reconstructions: u64,
pub retry_level_histogram: Vec<u64>,
}
pub fn write_solver_stats(output_dir: &Path, rows: &[SolverStatsRow]) -> Result<(), OutputError> {
write_solver_stats_to(&output_dir.join("training/solver"), rows)
}
pub fn write_simulation_solver_stats(
output_dir: &Path,
rows: &[SolverStatsRow],
) -> Result<(), OutputError> {
write_solver_stats_to(&output_dir.join("simulation/solver"), rows)
}
fn build_iterations_columns(rows: &[SolverStatsRow]) -> Vec<Arc<dyn arrow::array::Array>> {
let n = rows.len();
let iteration_arr = UInt32Array::from(rows.iter().map(|r| r.iteration).collect::<Vec<_>>());
let mut phase_builder = StringBuilder::with_capacity(n, n * 10);
for r in rows {
phase_builder.append_value(&r.phase);
}
let phase_arr = phase_builder.finish();
let stage_arr = Int32Array::from(rows.iter().map(|r| r.stage).collect::<Vec<_>>());
let opening_arr =
Int32Array::from(rows.iter().map(|r| r.opening).collect::<Vec<Option<i32>>>());
let rank_arr = Int32Array::from(rows.iter().map(|r| r.rank).collect::<Vec<Option<i32>>>());
let worker_id_arr = Int32Array::from(
rows.iter()
.map(|r| r.worker_id)
.collect::<Vec<Option<i32>>>(),
);
let lp_solves_arr = UInt32Array::from(rows.iter().map(|r| r.lp_solves).collect::<Vec<_>>());
let lp_successes_arr =
UInt32Array::from(rows.iter().map(|r| r.lp_successes).collect::<Vec<_>>());
let lp_retries_arr = UInt32Array::from(rows.iter().map(|r| r.lp_retries).collect::<Vec<_>>());
let lp_failures_arr = UInt32Array::from(rows.iter().map(|r| r.lp_failures).collect::<Vec<_>>());
let retry_attempts_arr =
UInt32Array::from(rows.iter().map(|r| r.retry_attempts).collect::<Vec<_>>());
let basis_offered_arr =
UInt32Array::from(rows.iter().map(|r| r.basis_offered).collect::<Vec<_>>());
let basis_consistency_failures_arr = UInt32Array::from(
rows.iter()
.map(|r| r.basis_consistency_failures)
.collect::<Vec<_>>(),
);
let simplex_iter_arr = UInt64Array::from(
rows.iter()
.map(|r| r.simplex_iterations)
.collect::<Vec<_>>(),
);
let solve_time_arr =
Float64Array::from(rows.iter().map(|r| r.solve_time_ms).collect::<Vec<_>>());
let load_model_time_arr = Float64Array::from(
rows.iter()
.map(|r| r.load_model_time_ms)
.collect::<Vec<_>>(),
);
let set_bounds_time_arr = Float64Array::from(
rows.iter()
.map(|r| r.set_bounds_time_ms)
.collect::<Vec<_>>(),
);
let basis_set_time_arr =
Float64Array::from(rows.iter().map(|r| r.basis_set_time_ms).collect::<Vec<_>>());
let basis_reconstructions_arr = UInt64Array::from(
rows.iter()
.map(|r| r.basis_reconstructions)
.collect::<Vec<_>>(),
);
vec![
Arc::new(iteration_arr),
Arc::new(phase_arr),
Arc::new(stage_arr),
Arc::new(opening_arr),
Arc::new(rank_arr),
Arc::new(worker_id_arr),
Arc::new(lp_solves_arr),
Arc::new(lp_successes_arr),
Arc::new(lp_retries_arr),
Arc::new(lp_failures_arr),
Arc::new(retry_attempts_arr),
Arc::new(basis_offered_arr),
Arc::new(basis_consistency_failures_arr),
Arc::new(simplex_iter_arr),
Arc::new(solve_time_arr),
Arc::new(load_model_time_arr),
Arc::new(set_bounds_time_arr),
Arc::new(basis_set_time_arr),
Arc::new(basis_reconstructions_arr),
]
}
fn build_retry_histogram_batch(rows: &[SolverStatsRow]) -> Result<RecordBatch, OutputError> {
let mut iterations = Vec::new();
let mut phases = Vec::new();
let mut stages = Vec::new();
let mut levels = Vec::new();
let mut counts = Vec::new();
for r in rows {
#[allow(clippy::cast_possible_truncation)]
for (level, &count) in r.retry_level_histogram.iter().enumerate() {
if count > 0 {
iterations.push(r.iteration);
phases.push(r.phase.as_str());
stages.push(r.stage);
levels.push(level as u32);
counts.push(count);
}
}
}
let n = iterations.len();
let mut phase_builder = StringBuilder::with_capacity(n, n * 10);
for &p in &phases {
phase_builder.append_value(p);
}
let schema = Arc::new(retry_histogram_schema());
RecordBatch::try_new(
schema,
vec![
Arc::new(UInt32Array::from(iterations)),
Arc::new(phase_builder.finish()),
Arc::new(Int32Array::from(stages)),
Arc::new(UInt32Array::from(levels)),
Arc::new(UInt64Array::from(counts)),
],
)
.map_err(|e| OutputError::serialization("retry_histogram", format!("RecordBatch: {e}")))
}
fn write_parquet(
path: &Path,
schema: &Arc<arrow::datatypes::Schema>,
batch: &RecordBatch,
) -> Result<(), OutputError> {
let tmp_path = path.with_extension("parquet.tmp");
let file = std::fs::File::create(&tmp_path).map_err(|e| OutputError::io(&tmp_path, e))?;
let props = WriterProperties::builder()
.set_compression(Compression::ZSTD(ZstdLevel::default()))
.build();
let mut writer = ArrowWriter::try_new(file, Arc::clone(schema), Some(props))
.map_err(|e| OutputError::serialization("solver_stats", format!("ArrowWriter: {e}")))?;
writer
.write(batch)
.map_err(|e| OutputError::serialization("solver_stats", format!("write: {e}")))?;
writer
.close()
.map_err(|e| OutputError::serialization("solver_stats", format!("close: {e}")))?;
std::fs::rename(&tmp_path, path).map_err(|e| OutputError::io(path, e))?;
Ok(())
}
fn write_solver_stats_to(dir: &Path, rows: &[SolverStatsRow]) -> Result<(), OutputError> {
std::fs::create_dir_all(dir).map_err(|e| OutputError::io(dir, e))?;
let iter_schema = Arc::new(solver_iterations_schema());
let columns = build_iterations_columns(rows);
let iter_batch = RecordBatch::try_new(Arc::clone(&iter_schema), columns)
.map_err(|e| OutputError::serialization("solver_stats", format!("RecordBatch: {e}")))?;
write_parquet(&dir.join("iterations.parquet"), &iter_schema, &iter_batch)?;
let hist_schema = Arc::new(retry_histogram_schema());
let hist_batch = build_retry_histogram_batch(rows)?;
write_parquet(
&dir.join("retry_histogram.parquet"),
&hist_schema,
&hist_batch,
)?;
Ok(())
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::float_cmp)]
mod tests {
use super::*;
use arrow::array::{Array, Float64Array, Int32Array, UInt32Array, UInt64Array};
use parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder;
fn make_rows() -> Vec<SolverStatsRow> {
vec![
SolverStatsRow {
iteration: 1,
phase: "forward".to_string(),
stage: 0,
opening: None,
rank: None,
worker_id: None,
lp_solves: 100,
lp_successes: 98,
lp_retries: 2,
lp_failures: 0,
retry_attempts: 4,
basis_offered: 90,
basis_consistency_failures: 3,
simplex_iterations: 5000,
solve_time_ms: 42.5,
load_model_time_ms: 0.0,
set_bounds_time_ms: 0.0,
basis_set_time_ms: 0.0,
basis_reconstructions: 0,
retry_level_histogram: vec![0; 12],
},
SolverStatsRow {
iteration: 1,
phase: "backward".to_string(),
stage: 2,
opening: Some(0),
rank: None,
worker_id: None,
lp_solves: 200,
lp_successes: 200,
lp_retries: 0,
lp_failures: 0,
retry_attempts: 0,
basis_offered: 180,
basis_consistency_failures: 1,
simplex_iterations: 10000,
solve_time_ms: 85.0,
load_model_time_ms: 0.0,
set_bounds_time_ms: 0.0,
basis_set_time_ms: 0.0,
basis_reconstructions: 0,
retry_level_histogram: vec![0; 12],
},
]
}
fn read_parquet(path: &std::path::Path) -> RecordBatch {
let file = std::fs::File::open(path).unwrap();
let builder = ParquetRecordBatchReaderBuilder::try_new(file).unwrap();
let mut reader = builder.build().unwrap();
reader.next().unwrap().unwrap()
}
#[test]
fn write_and_read_back() {
let dir = tempfile::TempDir::new().unwrap();
let rows = make_rows();
write_solver_stats(dir.path(), &rows).unwrap();
let iter_path = dir.path().join("training/solver/iterations.parquet");
assert!(iter_path.exists());
let batch = read_parquet(&iter_path);
assert_eq!(batch.num_rows(), 2);
assert_eq!(batch.num_columns(), 19);
let iteration_col = batch
.column(0)
.as_any()
.downcast_ref::<UInt32Array>()
.unwrap();
assert_eq!(iteration_col.value(0), 1);
assert_eq!(iteration_col.value(1), 1);
let solve_time_col = batch
.column(14)
.as_any()
.downcast_ref::<Float64Array>()
.unwrap();
assert!((solve_time_col.value(0) - 42.5).abs() < 1e-10);
assert!((solve_time_col.value(1) - 85.0).abs() < 1e-10);
let simplex_col = batch
.column(13)
.as_any()
.downcast_ref::<UInt64Array>()
.unwrap();
assert_eq!(simplex_col.value(0), 5000);
let hist_path = dir.path().join("training/solver/retry_histogram.parquet");
assert!(hist_path.exists());
let file = std::fs::File::open(&hist_path).unwrap();
let builder = ParquetRecordBatchReaderBuilder::try_new(file).unwrap();
assert_eq!(builder.schema().fields().len(), 5);
let total_rows: usize = builder
.build()
.unwrap()
.flatten()
.map(|b| b.num_rows())
.sum();
assert_eq!(total_rows, 0);
}
#[test]
fn write_empty_rows() {
let dir = tempfile::TempDir::new().unwrap();
write_solver_stats(dir.path(), &[]).unwrap();
let iter_path = dir.path().join("training/solver/iterations.parquet");
assert!(iter_path.exists());
let file = std::fs::File::open(&iter_path).unwrap();
let builder = ParquetRecordBatchReaderBuilder::try_new(file).unwrap();
assert_eq!(builder.schema().fields().len(), 19);
let hist_path = dir.path().join("training/solver/retry_histogram.parquet");
assert!(hist_path.exists());
let file = std::fs::File::open(&hist_path).unwrap();
let builder = ParquetRecordBatchReaderBuilder::try_new(file).unwrap();
assert_eq!(builder.schema().fields().len(), 5);
}
#[test]
fn retry_histogram_sparse_encoding() {
let dir = tempfile::TempDir::new().unwrap();
let rows = vec![
SolverStatsRow {
iteration: 1,
phase: "forward".to_string(),
stage: 0,
opening: None,
rank: None,
worker_id: None,
lp_solves: 50,
lp_successes: 48,
lp_retries: 2,
lp_failures: 0,
retry_attempts: 3,
basis_offered: 40,
basis_consistency_failures: 0,
simplex_iterations: 2000,
solve_time_ms: 10.0,
load_model_time_ms: 0.0,
set_bounds_time_ms: 0.0,
basis_set_time_ms: 0.0,
basis_reconstructions: 0,
retry_level_histogram: vec![5, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0],
},
SolverStatsRow {
iteration: 1,
phase: "backward".to_string(),
stage: 0,
opening: Some(0),
rank: None,
worker_id: None,
lp_solves: 100,
lp_successes: 100,
lp_retries: 0,
lp_failures: 0,
retry_attempts: 0,
basis_offered: 80,
basis_consistency_failures: 0,
simplex_iterations: 5000,
solve_time_ms: 20.0,
load_model_time_ms: 0.0,
set_bounds_time_ms: 0.0,
basis_set_time_ms: 0.0,
basis_reconstructions: 0,
retry_level_histogram: vec![0; 12],
},
];
write_solver_stats(dir.path(), &rows).unwrap();
let hist_path = dir.path().join("training/solver/retry_histogram.parquet");
let batch = read_parquet(&hist_path);
assert_eq!(batch.num_rows(), 2);
let level_col = batch
.column(3)
.as_any()
.downcast_ref::<UInt32Array>()
.unwrap();
assert_eq!(level_col.value(0), 0);
assert_eq!(level_col.value(1), 2);
let count_col = batch
.column(4)
.as_any()
.downcast_ref::<UInt64Array>()
.unwrap();
assert_eq!(count_col.value(0), 5);
assert_eq!(count_col.value(1), 1);
}
#[test]
fn test_solver_stats_row_builds_with_none_opening() {
let dir = tempfile::TempDir::new().unwrap();
let rows = vec![SolverStatsRow {
iteration: 1,
phase: "forward".to_string(),
stage: 0,
opening: None,
rank: None,
worker_id: None,
lp_solves: 10,
lp_successes: 10,
lp_retries: 0,
lp_failures: 0,
retry_attempts: 0,
basis_offered: 8,
basis_consistency_failures: 0,
simplex_iterations: 500,
solve_time_ms: 1.0,
load_model_time_ms: 0.0,
set_bounds_time_ms: 0.0,
basis_set_time_ms: 0.0,
basis_reconstructions: 0,
retry_level_histogram: vec![0; 12],
}];
write_solver_stats(dir.path(), &rows).unwrap();
let iter_path = dir.path().join("training/solver/iterations.parquet");
let batch = read_parquet(&iter_path);
assert_eq!(batch.num_rows(), 1);
let opening_col = batch
.column(3)
.as_any()
.downcast_ref::<Int32Array>()
.unwrap();
assert!(opening_col.is_null(0), "forward row must have NULL opening");
}
#[test]
fn test_solver_stats_row_builds_with_none_rank_and_worker_id() {
let dir = tempfile::TempDir::new().unwrap();
let rows = vec![SolverStatsRow {
iteration: 1,
phase: "backward".to_string(),
stage: 0,
opening: Some(3),
rank: None,
worker_id: None,
lp_solves: 1,
lp_successes: 1,
lp_retries: 0,
lp_failures: 0,
retry_attempts: 0,
basis_offered: 0,
basis_consistency_failures: 0,
simplex_iterations: 0,
solve_time_ms: 0.0,
load_model_time_ms: 0.0,
set_bounds_time_ms: 0.0,
basis_set_time_ms: 0.0,
basis_reconstructions: 0,
retry_level_histogram: vec![0; 12],
}];
write_solver_stats(dir.path(), &rows).unwrap();
let iter_path = dir.path().join("training/solver/iterations.parquet");
let batch = read_parquet(&iter_path);
assert_eq!(batch.num_columns(), 19);
let rank_col = batch.column_by_name("rank").unwrap();
assert_eq!(
rank_col.null_count(),
1,
"rank must be NULL for rank-aggregated rows"
);
let worker_col = batch.column_by_name("worker_id").unwrap();
assert_eq!(
worker_col.null_count(),
1,
"worker_id must be NULL for rank-aggregated rows"
);
let opening_col = batch
.column(3)
.as_any()
.downcast_ref::<Int32Array>()
.unwrap();
assert!(!opening_col.is_null(0), "opening must be non-NULL");
assert_eq!(opening_col.value(0), 3, "opening value must be 3");
}
#[allow(clippy::too_many_lines)]
#[test]
fn test_opening_column_sum_invariant() {
let dir = tempfile::TempDir::new().unwrap();
let rows = vec![
SolverStatsRow {
iteration: 1,
phase: "forward".to_string(),
stage: 0,
opening: None,
rank: None,
worker_id: None,
lp_solves: 50,
lp_successes: 50,
lp_retries: 0,
lp_failures: 0,
retry_attempts: 0,
basis_offered: 40,
basis_consistency_failures: 0,
simplex_iterations: 1000,
solve_time_ms: 5.0,
load_model_time_ms: 0.0,
set_bounds_time_ms: 0.0,
basis_set_time_ms: 0.0,
basis_reconstructions: 0,
retry_level_histogram: vec![0; 12],
},
SolverStatsRow {
iteration: 1,
phase: "backward".to_string(),
stage: 0,
opening: Some(0),
rank: None,
worker_id: None,
lp_solves: 10,
lp_successes: 10,
lp_retries: 0,
lp_failures: 0,
retry_attempts: 0,
basis_offered: 8,
basis_consistency_failures: 0,
simplex_iterations: 200,
solve_time_ms: 2.0,
load_model_time_ms: 0.0,
set_bounds_time_ms: 0.0,
basis_set_time_ms: 0.0,
basis_reconstructions: 0,
retry_level_histogram: vec![0; 12],
},
SolverStatsRow {
iteration: 1,
phase: "backward".to_string(),
stage: 0,
opening: Some(1),
rank: None,
worker_id: None,
lp_solves: 20,
lp_successes: 20,
lp_retries: 0,
lp_failures: 0,
retry_attempts: 0,
basis_offered: 18,
basis_consistency_failures: 0,
simplex_iterations: 400,
solve_time_ms: 4.0,
load_model_time_ms: 0.0,
set_bounds_time_ms: 0.0,
basis_set_time_ms: 0.0,
basis_reconstructions: 0,
retry_level_histogram: vec![0; 12],
},
SolverStatsRow {
iteration: 1,
phase: "backward".to_string(),
stage: 0,
opening: Some(2),
rank: None,
worker_id: None,
lp_solves: 30,
lp_successes: 30,
lp_retries: 0,
lp_failures: 0,
retry_attempts: 0,
basis_offered: 28,
basis_consistency_failures: 0,
simplex_iterations: 600,
solve_time_ms: 6.0,
load_model_time_ms: 0.0,
set_bounds_time_ms: 0.0,
basis_set_time_ms: 0.0,
basis_reconstructions: 0,
retry_level_histogram: vec![0; 12],
},
];
write_solver_stats(dir.path(), &rows).unwrap();
let iter_path = dir.path().join("training/solver/iterations.parquet");
let batch = read_parquet(&iter_path);
assert_eq!(batch.num_rows(), 4);
let lp_col = batch
.column(6)
.as_any()
.downcast_ref::<UInt32Array>()
.unwrap();
let backward_sum: u32 = (0..4)
.filter(|&i| {
batch
.column(1)
.as_any()
.downcast_ref::<arrow::array::StringArray>()
.unwrap()
.value(i)
== "backward"
})
.map(|i| lp_col.value(i))
.sum();
assert_eq!(
backward_sum, 60,
"SUM(lp_solves) for backward stage 0 must equal 60"
);
let opening_col = batch
.column(3)
.as_any()
.downcast_ref::<Int32Array>()
.unwrap();
assert!(opening_col.is_null(0), "forward row must have NULL opening");
assert_eq!(opening_col.value(1), 0, "backward opening[0] must be 0");
assert_eq!(opening_col.value(2), 1, "backward opening[1] must be 1");
assert_eq!(opening_col.value(3), 2, "backward opening[2] must be 2");
}
#[test]
fn test_forward_rows_are_per_stage_in_parquet() {
fn make_forward_row(stage: i32, lp_solves: u32) -> SolverStatsRow {
SolverStatsRow {
iteration: 1,
phase: "forward".to_string(),
stage,
opening: None, rank: None,
worker_id: None,
lp_solves,
lp_successes: lp_solves,
lp_retries: 0,
lp_failures: 0,
retry_attempts: 0,
basis_offered: 0,
basis_consistency_failures: 0,
simplex_iterations: u64::from(lp_solves) * 5,
solve_time_ms: f64::from(lp_solves) * 0.5,
load_model_time_ms: 0.0,
set_bounds_time_ms: 0.0,
basis_set_time_ms: 0.0,
basis_reconstructions: 0,
retry_level_histogram: vec![0; 12],
}
}
let dir = tempfile::TempDir::new().unwrap();
let rows = vec![
make_forward_row(0, 10), make_forward_row(1, 20), make_forward_row(2, 30), ];
write_solver_stats(dir.path(), &rows).unwrap();
let iter_path = dir.path().join("training/solver/iterations.parquet");
let batch = read_parquet(&iter_path);
assert_eq!(batch.num_rows(), 3, "one forward row per stage");
let opening_col = batch
.column(3)
.as_any()
.downcast_ref::<Int32Array>()
.unwrap();
let stage_col = batch
.column(2)
.as_any()
.downcast_ref::<arrow::array::Int32Array>()
.unwrap();
let lp_col = batch
.column(6)
.as_any()
.downcast_ref::<UInt32Array>()
.unwrap();
for row in 0..3 {
assert!(
opening_col.is_null(row),
"forward row {row} must have NULL opening"
);
assert_eq!(
stage_col.value(row),
i32::try_from(row).unwrap(),
"forward row {row} must have stage = {row}"
);
}
for row in 0..3 {
assert_ne!(
stage_col.value(row),
-1,
"forward rows must not use stage = -1"
);
}
assert_eq!(lp_col.value(0), 10, "stage 0: 10 lp_solves");
assert_eq!(lp_col.value(1), 20, "stage 1: 20 lp_solves");
assert_eq!(lp_col.value(2), 30, "stage 2: 30 lp_solves");
}
}