use std::path::{Path, PathBuf};
use std::sync::Arc;
use arrow::array::{
ArrayRef, Float64Builder, Int32Builder, Int64Builder, RecordBatch, StringBuilder,
};
use super::{IterationRecord, TrainingOutput, WorkerTimingRecord};
use crate::output::atomic::write_parquet_atomic;
use crate::output::error::OutputError;
use crate::output::parquet_config::ParquetWriterConfig;
use crate::output::schemas::{convergence_schema, iteration_timing_schema};
pub struct TrainingParquetWriter {
output_dir: PathBuf,
config: ParquetWriterConfig,
}
impl TrainingParquetWriter {
pub fn new(output_dir: &Path, config: &ParquetWriterConfig) -> Result<Self, OutputError> {
let training_dir = output_dir.join("training");
let timing_dir = output_dir.join("training/timing");
if !training_dir.exists() {
return Err(OutputError::io(
&training_dir,
std::io::Error::new(
std::io::ErrorKind::NotFound,
"training/ directory does not exist",
),
));
}
if !timing_dir.exists() {
return Err(OutputError::io(
&timing_dir,
std::io::Error::new(
std::io::ErrorKind::NotFound,
"training/timing/ directory does not exist",
),
));
}
Ok(Self {
output_dir: output_dir.to_path_buf(),
config: config.clone(),
})
}
pub fn write(&self, training_output: &TrainingOutput) -> Result<(), OutputError> {
let records = &training_output.convergence_records;
let convergence_batch =
build_convergence_batch(records, &training_output.final_upper_bound_kind)?;
let convergence_path = self.output_dir.join("training/convergence.parquet");
write_parquet_atomic(&convergence_path, &convergence_batch, &self.config)?;
let timing_batch = build_iteration_timing_batch(&training_output.worker_timing_records)?;
let timing_path = self.output_dir.join("training/timing/iterations.parquet");
write_parquet_atomic(&timing_path, &timing_batch, &self.config)?;
Ok(())
}
}
#[allow(
clippy::cast_possible_wrap,
clippy::cast_possible_truncation,
clippy::cast_sign_loss
)]
fn build_convergence_batch(
records: &[IterationRecord],
upper_bound_kind: &str,
) -> Result<RecordBatch, OutputError> {
let schema = Arc::new(convergence_schema());
let n = records.len();
let is_exact = upper_bound_kind == "exact";
let mut iteration = Int32Builder::with_capacity(n);
let mut lower_bound = Float64Builder::with_capacity(n);
let mut upper_bound = Float64Builder::with_capacity(n);
let mut upper_bound_std = Float64Builder::with_capacity(n);
let mut upper_bound_kind_col = StringBuilder::with_capacity(n, n * 12);
let mut gap_percent = Float64Builder::with_capacity(n);
let mut cuts_added = Int32Builder::with_capacity(n);
let mut cuts_removed = Int32Builder::with_capacity(n);
let mut cuts_active = Int64Builder::with_capacity(n);
let mut time_forward_ms = Int64Builder::with_capacity(n);
let mut time_backward_ms = Int64Builder::with_capacity(n);
let mut time_total_ms = Int64Builder::with_capacity(n);
let mut forward_passes = Int32Builder::with_capacity(n);
let mut lp_solves = Int64Builder::with_capacity(n);
let mut mean_rows_in_lp = Float64Builder::with_capacity(n);
for rec in records {
iteration.append_value(rec.iteration as i32);
lower_bound.append_value(rec.lower_bound);
upper_bound.append_value(rec.upper_bound);
if is_exact {
upper_bound_std.append_null();
} else {
upper_bound_std.append_value(rec.upper_bound_std);
}
upper_bound_kind_col.append_value(upper_bound_kind);
gap_percent.append_option(rec.gap_percent);
cuts_added.append_value(rec.cuts_added as i32);
cuts_removed.append_value(rec.cuts_removed as i32);
cuts_active.append_value(i64::from(rec.cuts_active));
time_forward_ms.append_value(rec.time_forward_ms as i64);
time_backward_ms.append_value(rec.time_backward_ms as i64);
time_total_ms.append_value(rec.time_total_ms as i64);
forward_passes.append_value(rec.forward_passes as i32);
lp_solves.append_value(i64::from(rec.lp_solves));
mean_rows_in_lp.append_value(rec.mean_rows_in_lp);
}
RecordBatch::try_new(
schema,
vec![
Arc::new(iteration.finish()),
Arc::new(lower_bound.finish()),
Arc::new(upper_bound.finish()),
Arc::new(upper_bound_std.finish()),
Arc::new(upper_bound_kind_col.finish()),
Arc::new(gap_percent.finish()),
Arc::new(cuts_added.finish()),
Arc::new(cuts_removed.finish()),
Arc::new(cuts_active.finish()),
Arc::new(time_forward_ms.finish()),
Arc::new(time_backward_ms.finish()),
Arc::new(time_total_ms.finish()),
Arc::new(forward_passes.finish()),
Arc::new(lp_solves.finish()),
Arc::new(mean_rows_in_lp.finish()),
],
)
.map_err(|e| OutputError::serialization("convergence", e.to_string()))
}
#[allow(clippy::cast_possible_wrap)]
fn build_iteration_timing_batch(
records: &[WorkerTimingRecord],
) -> Result<RecordBatch, OutputError> {
let schema = Arc::new(iteration_timing_schema());
let n = records.len();
let mut iteration = Int32Builder::with_capacity(n);
let mut rank = Int32Builder::with_capacity(n);
let mut worker_id = Int32Builder::with_capacity(n);
let mut forward_wall_ms = Int64Builder::with_capacity(n);
let mut backward_wall_ms = Int64Builder::with_capacity(n);
let mut cut_selection_ms = Int64Builder::with_capacity(n);
let mut mpi_allreduce_ms = Int64Builder::with_capacity(n);
let mut cut_sync_ms = Int64Builder::with_capacity(n);
let mut lower_bound_ms = Int64Builder::with_capacity(n);
let mut state_exchange_ms = Int64Builder::with_capacity(n);
let mut cut_batch_build_ms = Int64Builder::with_capacity(n);
let mut bwd_setup_ms = Int64Builder::with_capacity(n);
let mut bwd_load_imbalance_ms = Int64Builder::with_capacity(n);
let mut bwd_scheduling_overhead_ms = Int64Builder::with_capacity(n);
let mut fwd_setup_ms = Int64Builder::with_capacity(n);
let mut fwd_load_imbalance_ms = Int64Builder::with_capacity(n);
let mut fwd_scheduling_overhead_ms = Int64Builder::with_capacity(n);
let mut overhead_ms = Int64Builder::with_capacity(n);
let mut lazy_scoring_ms = Int64Builder::with_capacity(n);
for rec in records {
iteration.append_value(rec.iteration as i32);
rank.append_value(rec.rank);
worker_id.append_option(rec.worker_id);
forward_wall_ms.append_value(rec.timings[0] as i64);
backward_wall_ms.append_value(rec.timings[1] as i64);
cut_selection_ms.append_value(rec.timings[2] as i64);
mpi_allreduce_ms.append_value(rec.timings[3] as i64);
cut_sync_ms.append_value(rec.timings[4] as i64);
lower_bound_ms.append_value(rec.timings[5] as i64);
state_exchange_ms.append_value(rec.timings[6] as i64);
cut_batch_build_ms.append_value(rec.timings[7] as i64);
bwd_setup_ms.append_value(rec.timings[8] as i64);
bwd_load_imbalance_ms.append_value(rec.timings[9] as i64);
bwd_scheduling_overhead_ms.append_value(rec.timings[10] as i64);
fwd_setup_ms.append_value(rec.timings[11] as i64);
fwd_load_imbalance_ms.append_value(rec.timings[12] as i64);
fwd_scheduling_overhead_ms.append_value(rec.timings[13] as i64);
overhead_ms.append_value(rec.timings[14] as i64);
lazy_scoring_ms.append_value(rec.timings[15] as i64);
}
RecordBatch::try_new(
schema,
vec![
Arc::new(iteration.finish()),
Arc::new(rank.finish()),
Arc::new(worker_id.finish()),
Arc::new(forward_wall_ms.finish()),
Arc::new(backward_wall_ms.finish()),
Arc::new(cut_selection_ms.finish()),
Arc::new(mpi_allreduce_ms.finish()),
Arc::new(cut_sync_ms.finish()),
Arc::new(lower_bound_ms.finish()),
Arc::new(state_exchange_ms.finish()),
Arc::new(cut_batch_build_ms.finish()),
Arc::new(bwd_setup_ms.finish()),
Arc::new(bwd_load_imbalance_ms.finish()),
Arc::new(bwd_scheduling_overhead_ms.finish()),
Arc::new(fwd_setup_ms.finish()),
Arc::new(fwd_load_imbalance_ms.finish()),
Arc::new(fwd_scheduling_overhead_ms.finish()),
Arc::new(overhead_ms.finish()),
Arc::new(lazy_scoring_ms.finish()),
],
)
.map_err(|e| OutputError::serialization("iteration_timing", e.to_string()))
}
pub fn write_row_selection_records(
output_dir: &Path,
records: &[super::RowSelectionRecord],
config: &ParquetWriterConfig,
) -> Result<(), OutputError> {
if records.is_empty() {
return Ok(());
}
let dir = output_dir.join("training/cut_selection");
std::fs::create_dir_all(&dir).map_err(|e| OutputError::io(&dir, e))?;
let schema = Arc::new(super::schemas::row_selection_schema());
let n = records.len();
let mut iteration_builder = Int32Builder::with_capacity(n);
let mut stage_builder = Int32Builder::with_capacity(n);
let mut populated_builder = Int32Builder::with_capacity(n);
let mut active_before_builder = Int32Builder::with_capacity(n);
let mut deactivated_builder = Int32Builder::with_capacity(n);
let mut reactivated_builder = Int32Builder::with_capacity(n);
let mut active_after_builder = Int32Builder::with_capacity(n);
let mut selection_time_builder = Float64Builder::with_capacity(n);
let mut budget_evicted_builder = Int32Builder::with_capacity(n);
let mut active_after_budget_builder = Int32Builder::with_capacity(n);
#[allow(clippy::cast_possible_truncation, clippy::cast_possible_wrap)]
for r in records {
iteration_builder.append_value(r.iteration as i32);
stage_builder.append_value(r.stage as i32);
populated_builder.append_value(r.cuts_populated as i32);
active_before_builder.append_value(r.cuts_active_before as i32);
deactivated_builder.append_value(r.cuts_deactivated as i32);
reactivated_builder.append_value(r.cuts_reactivated as i32);
active_after_builder.append_value(r.cuts_active_after as i32);
selection_time_builder.append_value(r.selection_time_ms);
budget_evicted_builder.append_option(r.budget_evicted.map(|v| v as i32));
active_after_budget_builder.append_option(r.active_after_budget.map(|v| v as i32));
}
let columns: Vec<ArrayRef> = vec![
Arc::new(iteration_builder.finish()),
Arc::new(stage_builder.finish()),
Arc::new(populated_builder.finish()),
Arc::new(active_before_builder.finish()),
Arc::new(deactivated_builder.finish()),
Arc::new(reactivated_builder.finish()),
Arc::new(active_after_builder.finish()),
Arc::new(selection_time_builder.finish()),
Arc::new(budget_evicted_builder.finish()),
Arc::new(active_after_budget_builder.finish()),
];
let batch = RecordBatch::try_new(Arc::clone(&schema), columns)
.map_err(|e| OutputError::serialization("cut_selection", e.to_string()))?;
write_parquet_atomic(&dir.join("iterations.parquet"), &batch, config)
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::expect_used,
clippy::panic,
clippy::float_cmp
)]
mod tests {
use super::*;
use crate::MetadataTrainingSolveStats;
use crate::output::{RowPoolStatistics, TrainingOutput};
fn make_record(iteration: u32, gap: Option<f64>) -> IterationRecord {
IterationRecord {
iteration,
lower_bound: f64::from(iteration) * 10.0,
upper_bound: f64::from(iteration) * 11.0,
upper_bound_std: 0.5,
gap_percent: gap,
cuts_added: 5,
cuts_removed: 1,
cuts_active: 4,
time_forward_ms: 100,
time_backward_ms: 200,
time_total_ms: 300,
time_forward_wall_ms: 100,
time_backward_wall_ms: 200,
time_cut_selection_ms: 0,
time_mpi_allreduce_ms: 0,
time_cut_sync_ms: 0,
time_lower_bound_ms: 0,
time_state_exchange_ms: 0,
time_cut_batch_build_ms: 0,
time_bwd_setup_ms: 0,
time_bwd_load_imbalance_ms: 0,
time_bwd_scheduling_overhead_ms: 0,
time_fwd_setup_ms: 0,
time_fwd_load_imbalance_ms: 0,
time_fwd_scheduling_overhead_ms: 0,
time_overhead_ms: 0,
forward_passes: 4,
lp_solves: 40,
solve_time_ms: 0.0,
mean_rows_in_lp: 0.0,
}
}
fn make_training_output(records: Vec<IterationRecord>) -> TrainingOutput {
TrainingOutput {
convergence_records: records,
final_lower_bound: 99.5,
final_upper_bound: Some(101.0),
final_gap_percent: Some(1.51),
final_upper_bound_std: Some(0.5),
final_upper_bound_kind: "statistical".to_string(),
iterations_completed: 0,
converged: true,
termination_reason: "gap tolerance reached".to_string(),
total_time_ms: 5_000,
cut_stats: RowPoolStatistics {
total_generated: 200,
total_active: 80,
peak_active: 95,
cuts_active: 80,
rows_in_lp_total: 0,
rows_in_lp_solve_count: 0,
rows_in_lp_max: 0,
total_loaded: 0,
},
cut_selection_records: vec![],
worker_timing_records: vec![],
training_solve_stats: MetadataTrainingSolveStats::default(),
}
}
fn make_worker_timing_record(iteration: u32) -> WorkerTimingRecord {
WorkerTimingRecord {
iteration,
rank: 0,
worker_id: None,
timings: [100, 200, 10, 5, 8, 3, 4, 7, 50, 60, 20, 40, 30, 15, 12, 0],
}
}
#[test]
fn convergence_batch_from_empty_records() {
let batch = build_convergence_batch(&[], "statistical").expect("empty batch must succeed");
assert_eq!(batch.num_rows(), 0, "empty records yield 0 rows");
assert_eq!(batch.num_columns(), 15, "convergence schema has 15 columns");
}
#[test]
fn convergence_batch_field_count_and_types() {
let records: Vec<IterationRecord> = (1..=3).map(|i| make_record(i, Some(5.0))).collect();
let batch = build_convergence_batch(&records, "statistical").expect("batch must be built");
assert_eq!(batch.num_rows(), 3);
assert_eq!(batch.num_columns(), 15);
let expected_schema = convergence_schema();
assert_eq!(
batch.schema().fields(),
expected_schema.fields(),
"schema must match convergence_schema()"
);
}
#[test]
fn convergence_batch_nullable_columns() {
let records = vec![
make_record(1, Some(10.0)),
make_record(2, Some(5.0)),
make_record(3, None), ];
let batch = build_convergence_batch(&records, "statistical").expect("batch must be built");
let gap_col = batch
.column_by_name("gap_percent")
.expect("gap_percent column must exist");
assert!(!gap_col.is_null(0), "row 0: Some(10.0) must not be null");
assert!(!gap_col.is_null(1), "row 1: Some(5.0) must not be null");
assert!(gap_col.is_null(2), "row 2: None must be null");
}
#[test]
fn iteration_timing_batch_field_count() {
let records: Vec<WorkerTimingRecord> = (1..=3).map(make_worker_timing_record).collect();
let batch = build_iteration_timing_batch(&records).expect("timing batch must be built");
assert_eq!(batch.num_rows(), 3, "3 records yield 3 rows");
assert_eq!(
batch.num_columns(),
19,
"iteration_timing schema has 19 columns (16 timings + iteration + rank + worker_id)"
);
let expected_schema = iteration_timing_schema();
assert_eq!(
batch.schema().fields(),
expected_schema.fields(),
"schema must match iteration_timing_schema()"
);
}
#[test]
#[allow(clippy::too_many_lines)]
fn iteration_timing_columns_six_decomposed_overhead() {
use arrow::array::Int64Array;
use parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder;
let tmp = tempfile::tempdir().expect("tempdir must succeed");
std::fs::create_dir_all(tmp.path().join("training/timing")).unwrap();
let config = ParquetWriterConfig::default();
let records: Vec<IterationRecord> = (1u32..=3)
.map(|i| IterationRecord {
iteration: i,
lower_bound: f64::from(i) * 10.0,
upper_bound: f64::from(i) * 11.0,
upper_bound_std: 0.5,
gap_percent: Some(5.0),
cuts_added: 5,
cuts_removed: 1,
cuts_active: 4,
time_forward_ms: 100,
time_backward_ms: 200,
time_total_ms: 300,
time_forward_wall_ms: 100,
time_backward_wall_ms: 200,
time_cut_selection_ms: 0,
time_mpi_allreduce_ms: 0,
time_cut_sync_ms: 0,
time_lower_bound_ms: 0,
time_state_exchange_ms: 0,
time_cut_batch_build_ms: 0,
time_bwd_setup_ms: u64::from(i) * 10,
time_bwd_load_imbalance_ms: u64::from(i) * 20,
time_bwd_scheduling_overhead_ms: u64::from(i) * 30,
time_fwd_setup_ms: u64::from(i) * 40,
time_fwd_load_imbalance_ms: u64::from(i) * 50,
time_fwd_scheduling_overhead_ms: u64::from(i) * 60,
time_overhead_ms: 0,
forward_passes: 4,
lp_solves: 40,
solve_time_ms: 0.0,
mean_rows_in_lp: 0.0,
})
.collect();
let worker_records: Vec<WorkerTimingRecord> = (1u32..=3)
.map(|i| WorkerTimingRecord {
iteration: i,
rank: 0,
worker_id: None,
timings: [
100,
200,
0,
0,
0,
0,
0,
0,
u64::from(i) * 10,
u64::from(i) * 20,
u64::from(i) * 30,
u64::from(i) * 40,
u64::from(i) * 50,
u64::from(i) * 60,
0,
0,
],
})
.collect();
let mut training = make_training_output(records);
training.worker_timing_records = worker_records;
let writer = TrainingParquetWriter::new(tmp.path(), &config).expect("new must succeed");
writer.write(&training).expect("write must succeed");
let timing_path = tmp.path().join("training/timing/iterations.parquet");
assert!(timing_path.exists(), "iterations.parquet must exist");
let file = std::fs::File::open(&timing_path).expect("file must open");
let mut reader = ParquetRecordBatchReaderBuilder::try_new(file)
.expect("builder")
.build()
.expect("reader");
let batch = reader.next().expect("must have rows").expect("batch Ok");
assert_eq!(batch.num_rows(), 3);
assert_eq!(batch.num_columns(), 19);
assert!(batch.column_by_name("bwd_rayon_overhead_ms").is_none());
assert!(batch.column_by_name("fwd_rayon_overhead_ms").is_none());
let expected_schema = iteration_timing_schema();
assert_eq!(
batch.schema().fields(),
expected_schema.fields(),
"schema must match iteration_timing_schema()"
);
for (col_name, expected_row1_val) in &[
("bwd_setup_ms", 10_i64),
("bwd_load_imbalance_ms", 20_i64),
("bwd_scheduling_overhead_ms", 30_i64),
("fwd_setup_ms", 40_i64),
("fwd_load_imbalance_ms", 50_i64),
("fwd_scheduling_overhead_ms", 60_i64),
] {
let col = batch
.column_by_name(col_name)
.unwrap_or_else(|| panic!("column {col_name} must exist"));
let arr = col
.as_any()
.downcast_ref::<Int64Array>()
.unwrap_or_else(|| panic!("{col_name} must be Int64Array"));
assert_eq!(
arr.value(0),
*expected_row1_val,
"{col_name} row 0 must be {expected_row1_val}"
);
assert_eq!(
arr.value(1),
expected_row1_val * 2,
"{col_name} row 1 must be {}",
expected_row1_val * 2
);
}
}
#[test]
fn write_convergence_parquet_roundtrip() {
use parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder;
let records: Vec<IterationRecord> = (1..=5).map(|i| make_record(i, Some(1.0))).collect();
let batch = build_convergence_batch(&records, "statistical").expect("batch must be built");
let tmp = tempfile::tempdir().expect("tempdir must succeed");
let path = tmp.path().join("convergence.parquet");
let config = ParquetWriterConfig::default();
write_parquet_atomic(&path, &batch, &config).expect("write must succeed");
assert!(path.exists(), "convergence.parquet must exist after write");
let file = std::fs::File::open(&path).expect("file must open");
let builder =
ParquetRecordBatchReaderBuilder::try_new(file).expect("builder must be created");
let mut reader = builder.build().expect("reader must be built");
let read_batch = reader
.next()
.expect("must have at least one batch")
.expect("batch must be Ok");
assert_eq!(read_batch.num_rows(), 5, "must have 5 rows");
let expected_schema = convergence_schema();
assert_eq!(
read_batch.schema().fields(),
expected_schema.fields(),
"schema must match convergence_schema()"
);
let iteration_col = read_batch
.column_by_name("iteration")
.expect("iteration column must exist");
let iteration_arr = iteration_col
.as_any()
.downcast_ref::<arrow::array::Int32Array>()
.expect("iteration must be Int32Array");
let iteration_values: Vec<i32> = (0..5).map(|i| iteration_arr.value(i)).collect();
assert_eq!(iteration_values, vec![1, 2, 3, 4, 5]);
let lb_col = read_batch
.column_by_name("lower_bound")
.expect("lower_bound column must exist");
let lb_arr = lb_col
.as_any()
.downcast_ref::<arrow::array::Float64Array>()
.expect("lower_bound must be Float64Array");
for (i, rec) in records.iter().enumerate() {
assert_eq!(
lb_arr.value(i),
rec.lower_bound,
"lower_bound mismatch at row {i}"
);
}
}
#[test]
fn write_convergence_parquet_atomic_rename() {
let records: Vec<IterationRecord> = (1..=2).map(|i| make_record(i, None)).collect();
let batch = build_convergence_batch(&records, "statistical").expect("batch must be built");
let tmp = tempfile::tempdir().expect("tempdir must succeed");
let path = tmp.path().join("convergence.parquet");
let config = ParquetWriterConfig::default();
write_parquet_atomic(&path, &batch, &config).expect("write must succeed");
let tmp_path = path.with_extension("parquet.tmp");
assert!(
!tmp_path.exists(),
".tmp file must not exist after successful atomic rename"
);
assert!(path.exists(), "final file must exist");
}
#[test]
fn writer_fails_if_training_dir_missing() {
let tmp = tempfile::tempdir().expect("tempdir must succeed");
let config = ParquetWriterConfig::default();
let result = TrainingParquetWriter::new(tmp.path(), &config);
assert!(result.is_err(), "new() must fail when training/ is missing");
}
#[test]
fn writer_fails_if_timing_dir_missing() {
let tmp = tempfile::tempdir().expect("tempdir must succeed");
let config = ParquetWriterConfig::default();
std::fs::create_dir_all(tmp.path().join("training")).unwrap();
let result = TrainingParquetWriter::new(tmp.path(), &config);
assert!(
result.is_err(),
"new() must fail when training/timing/ is missing"
);
}
#[test]
fn writer_writes_empty_training_output() {
use parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder;
let tmp = tempfile::tempdir().expect("tempdir must succeed");
std::fs::create_dir_all(tmp.path().join("training/timing")).unwrap();
let config = ParquetWriterConfig::default();
let writer = TrainingParquetWriter::new(tmp.path(), &config).expect("new must succeed");
let training = make_training_output(vec![]);
writer.write(&training).expect("write must succeed");
let conv_path = tmp.path().join("training/convergence.parquet");
assert!(conv_path.exists(), "convergence.parquet must exist");
let timing_path = tmp.path().join("training/timing/iterations.parquet");
assert!(timing_path.exists(), "iterations.parquet must exist");
let file = std::fs::File::open(&conv_path).expect("file must open");
let builder = ParquetRecordBatchReaderBuilder::try_new(file).expect("builder created");
let schema = builder.schema().clone();
let reader = builder.build().expect("reader built");
let total_rows: usize = reader
.map(|b| b.expect("batch must be Ok").num_rows())
.sum();
assert_eq!(total_rows, 0, "empty training output must produce 0 rows");
let expected_schema = convergence_schema();
assert_eq!(
schema.fields(),
expected_schema.fields(),
"schema must match convergence_schema()"
);
}
#[test]
fn writer_writes_five_records() {
use parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder;
let tmp = tempfile::tempdir().expect("tempdir must succeed");
std::fs::create_dir_all(tmp.path().join("training/timing")).unwrap();
let config = ParquetWriterConfig::default();
let records: Vec<IterationRecord> = (1..=5).map(|i| make_record(i, Some(1.0))).collect();
let mut training = make_training_output(records);
training.worker_timing_records = (1u32..=5).map(make_worker_timing_record).collect();
let writer = TrainingParquetWriter::new(tmp.path(), &config).expect("new must succeed");
writer.write(&training).expect("write must succeed");
let conv_path = tmp.path().join("training/convergence.parquet");
let file = std::fs::File::open(&conv_path).expect("file must open");
let mut reader = ParquetRecordBatchReaderBuilder::try_new(file)
.expect("builder")
.build()
.expect("reader");
let batch = reader.next().expect("must have rows").expect("batch Ok");
assert_eq!(batch.num_rows(), 5);
assert_eq!(batch.num_columns(), 15);
let timing_path = tmp.path().join("training/timing/iterations.parquet");
let file = std::fs::File::open(&timing_path).expect("file must open");
let mut reader = ParquetRecordBatchReaderBuilder::try_new(file)
.expect("builder")
.build()
.expect("reader");
let batch = reader.next().expect("must have rows").expect("batch Ok");
assert_eq!(batch.num_rows(), 5);
assert_eq!(batch.num_columns(), 19, "timing schema has 19 columns");
}
#[test]
fn writer_gap_percent_null_at_correct_row() {
use parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder;
let tmp = tempfile::tempdir().expect("tempdir must succeed");
std::fs::create_dir_all(tmp.path().join("training/timing")).unwrap();
let config = ParquetWriterConfig::default();
let records = vec![
make_record(1, Some(10.0)),
make_record(2, Some(5.0)),
make_record(3, None), make_record(4, Some(2.0)),
make_record(5, Some(1.0)),
];
let training = make_training_output(records);
let writer = TrainingParquetWriter::new(tmp.path(), &config).expect("new must succeed");
writer.write(&training).expect("write must succeed");
let conv_path = tmp.path().join("training/convergence.parquet");
let file = std::fs::File::open(&conv_path).expect("file must open");
let mut reader = ParquetRecordBatchReaderBuilder::try_new(file)
.expect("builder")
.build()
.expect("reader");
let batch = reader.next().expect("must have rows").expect("batch Ok");
let gap_col = batch
.column_by_name("gap_percent")
.expect("gap_percent column must exist");
assert!(!gap_col.is_null(0), "row 0: Some(10.0) must not be null");
assert!(!gap_col.is_null(1), "row 1: Some(5.0) must not be null");
assert!(gap_col.is_null(2), "row 2: None must be null");
assert!(!gap_col.is_null(3), "row 3: Some(2.0) must not be null");
assert!(!gap_col.is_null(4), "row 4: Some(1.0) must not be null");
}
#[test]
fn write_cut_selection_empty_is_noop() {
let tmp = tempfile::tempdir().unwrap();
let config = ParquetWriterConfig::default();
write_row_selection_records(tmp.path(), &[], &config).unwrap();
assert!(
!tmp.path()
.join("training/cut_selection/iterations.parquet")
.exists()
);
}
#[test]
fn write_cut_selection_roundtrip() {
use super::super::RowSelectionRecord;
let tmp = tempfile::tempdir().unwrap();
let config = ParquetWriterConfig::default();
let records = vec![
RowSelectionRecord {
iteration: 3,
stage: 0,
cuts_populated: 10,
cuts_active_before: 10,
cuts_deactivated: 0,
cuts_reactivated: 0,
cuts_active_after: 10,
selection_time_ms: 0.0,
budget_evicted: None,
active_after_budget: None,
},
RowSelectionRecord {
iteration: 3,
stage: 1,
cuts_populated: 8,
cuts_active_before: 8,
cuts_deactivated: 2,
cuts_reactivated: 0,
cuts_active_after: 6,
selection_time_ms: 1.5,
budget_evicted: None,
active_after_budget: None,
},
];
write_row_selection_records(tmp.path(), &records, &config).unwrap();
let path = tmp.path().join("training/cut_selection/iterations.parquet");
assert!(path.exists());
let file = std::fs::File::open(&path).unwrap();
let reader = parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder::try_new(file)
.unwrap()
.build()
.unwrap();
let batch: RecordBatch = reader.into_iter().next().unwrap().unwrap();
assert_eq!(batch.num_rows(), 2);
assert_eq!(batch.num_columns(), 10);
}
#[test]
fn write_cut_selection_with_budget_columns_roundtrip() {
use super::super::RowSelectionRecord;
use parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder;
let tmp = tempfile::tempdir().unwrap();
let config = ParquetWriterConfig::default();
let records = vec![
RowSelectionRecord {
iteration: 5,
stage: 0,
cuts_populated: 20,
cuts_active_before: 20,
cuts_deactivated: 0,
cuts_reactivated: 0,
cuts_active_after: 20,
selection_time_ms: 0.0,
budget_evicted: Some(3),
active_after_budget: Some(15),
},
RowSelectionRecord {
iteration: 5,
stage: 1,
cuts_populated: 15,
cuts_active_before: 15,
cuts_deactivated: 2,
cuts_reactivated: 1,
cuts_active_after: 13,
selection_time_ms: 2.0,
budget_evicted: None,
active_after_budget: None,
},
];
write_row_selection_records(tmp.path(), &records, &config).unwrap();
let path = tmp.path().join("training/cut_selection/iterations.parquet");
assert!(path.exists());
let file = std::fs::File::open(&path).unwrap();
let mut reader = ParquetRecordBatchReaderBuilder::try_new(file)
.unwrap()
.build()
.unwrap();
let batch = reader.next().unwrap().unwrap();
assert_eq!(batch.num_rows(), 2);
assert_eq!(batch.num_columns(), 10);
let budget_evicted_col = batch.column_by_name("budget_evicted").unwrap();
assert!(
!budget_evicted_col.is_null(0),
"row 0: budget_evicted Some(3) must not be null"
);
assert!(
budget_evicted_col.is_null(1),
"row 1: budget_evicted None must be null"
);
let budget_col = batch.column_by_name("active_after_budget").unwrap();
assert!(
!budget_col.is_null(0),
"row 0: active_after_budget Some(15) must not be null"
);
assert!(
budget_col.is_null(1),
"row 1: active_after_budget None must be null"
);
let budget_evicted_arr = budget_evicted_col
.as_any()
.downcast_ref::<arrow::array::Int32Array>()
.unwrap();
assert_eq!(budget_evicted_arr.value(0), 3);
let budget_arr = budget_col
.as_any()
.downcast_ref::<arrow::array::Int32Array>()
.unwrap();
assert_eq!(budget_arr.value(0), 15);
}
#[test]
fn parquet_schema_includes_cuts_active_column() {
use arrow::datatypes::DataType;
use parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder;
use super::super::RowSelectionRecord;
let tmp = tempfile::tempdir().unwrap();
let config = ParquetWriterConfig::default();
let records = vec![RowSelectionRecord {
iteration: 1,
stage: 0,
cuts_populated: 5,
cuts_active_before: 5,
cuts_deactivated: 0,
cuts_reactivated: 0,
cuts_active_after: 5,
selection_time_ms: 0.0,
budget_evicted: None,
active_after_budget: None,
}];
write_row_selection_records(tmp.path(), &records, &config).unwrap();
let path = tmp.path().join("training/cut_selection/iterations.parquet");
let file = std::fs::File::open(&path).unwrap();
let builder = ParquetRecordBatchReaderBuilder::try_new(file).unwrap();
let schema = builder.schema().clone();
let cuts_reactivated = schema
.field_with_name("cuts_reactivated")
.expect("cuts_reactivated column must exist in schema");
assert_eq!(
cuts_reactivated.data_type(),
&DataType::Int32,
"cuts_reactivated must be Int32"
);
assert!(
!cuts_reactivated.is_nullable(),
"cuts_reactivated must not be nullable"
);
assert!(
schema.field_with_name("cuts_in_lp").is_err(),
"cuts_in_lp column must not be present in schema"
);
}
}