diff --git a/qdp/qdp-core/src/pipeline_runner.rs b/qdp/qdp-core/src/pipeline_runner.rs index 71b95d5d15..e9374dbf21 100644 --- a/qdp/qdp-core/src/pipeline_runner.rs +++ b/qdp/qdp-core/src/pipeline_runner.rs @@ -26,7 +26,7 @@ use crate::dlpack::DLManagedTensor; use crate::error::{MahoutError, Result}; use crate::gpu::memory::Precision; use crate::io; -use crate::reader::{NullHandling, StreamingDataReader}; +use crate::reader::{FloatElem, NullHandling, StreamingDataReader}; use crate::readers::ParquetStreamingReader; use crate::types::Encoding; @@ -116,6 +116,45 @@ pub enum BatchData { F64(Vec), } +impl BatchData { + #[allow(clippy::len_without_is_empty)] + pub fn len(&self) -> usize { + match self { + Self::F32(v) => v.len(), + Self::F64(v) => v.len(), + } + } +} + +pub(crate) trait ToBatchData: FloatElem { + fn wrap(v: Vec) -> BatchData; + fn from_recycled(b: BatchData) -> Option>; +} + +impl ToBatchData for f32 { + fn wrap(v: Vec) -> BatchData { + BatchData::F32(v) + } + fn from_recycled(b: BatchData) -> Option> { + match b { + BatchData::F32(v) => Some(v), + _ => None, + } + } +} + +impl ToBatchData for f64 { + fn wrap(v: Vec) -> BatchData { + BatchData::F64(v) + } + fn from_recycled(b: BatchData) -> Option> { + match b { + BatchData::F64(v) => Some(v), + _ => None, + } + } +} + pub struct PrefetchedBatch { pub data: BatchData, pub batch_n: usize, @@ -232,8 +271,8 @@ impl BatchProducer for SyntheticProducer { } } -pub struct InMemoryProducer { - pub data: Vec, +pub struct InMemoryProducer { + pub data: Vec, pub cursor: usize, pub sample_size: usize, pub batch_size: usize, @@ -242,7 +281,7 @@ pub struct InMemoryProducer { pub batch_limit: usize, } -impl BatchProducer for InMemoryProducer { +impl BatchProducer for InMemoryProducer { fn produce(&mut self, recycled: Option) -> Result> { if self.batches_yielded >= self.batch_limit { return Ok(None); @@ -259,13 +298,13 @@ impl BatchProducer for InMemoryProducer { self.batches_yielded += 1; let slice = &self.data[start..end]; - let data = match recycled { - Some(BatchData::F64(mut buf)) => { + let data = match recycled.and_then(T::from_recycled) { + Some(mut buf) => { buf.clear(); buf.extend_from_slice(slice); - BatchData::F64(buf) + T::wrap(buf) } - _ => BatchData::F64(slice.to_vec()), + None => T::wrap(slice.to_vec()), }; Ok(Some(PrefetchedBatch { @@ -277,11 +316,11 @@ impl BatchProducer for InMemoryProducer { } } -pub struct StreamingProducer { - pub reader: ParquetStreamingReader, - pub buffer: Vec, +pub struct StreamingProducer { + pub reader: ParquetStreamingReader, + pub buffer: Vec, pub buffer_cursor: usize, - pub read_chunk_scratch: Vec, + pub read_chunk_scratch: Vec, pub sample_size: usize, pub batch_size: usize, pub num_qubits: usize, @@ -289,7 +328,7 @@ pub struct StreamingProducer { pub batch_limit: usize, } -impl BatchProducer for StreamingProducer { +impl BatchProducer for StreamingProducer { fn produce(&mut self, recycled: Option) -> Result> { if self.batches_yielded >= self.batch_limit { return Ok(None); @@ -316,13 +355,13 @@ impl BatchProducer for StreamingProducer { self.buffer_cursor = end; self.batches_yielded += 1; - let data = match recycled { - Some(BatchData::F64(mut buf)) => { + let data = match recycled.and_then(T::from_recycled) { + Some(mut buf) => { buf.clear(); buf.extend_from_slice(&self.buffer[start..end]); - BatchData::F64(buf) + T::wrap(buf) } - _ => BatchData::F64(self.buffer[start..end].to_vec()), + None => T::wrap(self.buffer[start..end].to_vec()), }; if self.buffer_cursor >= self.buffer.len() / BUFFER_COMPACT_DENOM { @@ -389,33 +428,106 @@ fn path_extension_lower(path: &Path) -> Option { .map(|s| s.to_lowercase()) } -/// Dispatches by path extension to the appropriate io reader. Returns (data, num_samples, sample_size). -/// Unsupported or missing extension returns Err with message listing supported formats. +/// Largest integer exactly representable by an f32 (its mantissa is 24 bits): `2^24`. +/// Basis indices above this would silently change value when narrowed to f32. +const MAX_EXACT_F32_INT: f64 = (1u64 << 24) as f64; + +/// Reject an explicit f32 request for a basis file whose indices exceed f32's exact +/// integer range. Basis values are integer state indices, not floats, so narrowing +/// `16_777_217` to `16_777_216` would encode the wrong state. Rather than silently +/// corrupt (or silently widen back to f64), surface the conflict to the caller. +fn reject_basis_f32_out_of_range(data: &BatchData) -> Result<()> { + if let BatchData::F64(v) = data + && let Some(&bad) = v.iter().find(|&&x| x > MAX_EXACT_F32_INT) + { + return Err(MahoutError::InvalidInput(format!( + "basis index {bad:.0} exceeds f32's exact integer range ({MAX_EXACT_F32_INT:.0}); \ + narrowing to f32 would encode the wrong state. Use dtype='float64' for this basis file." + ))); + } + Ok(()) +} + +/// f64→f32 narrowing cast: values outside f32 range silently become ±Inf. +fn cast_f64_to_batch_data( + data: Vec, + n: usize, + s: usize, + dtype: Precision, + fmt: &str, +) -> (BatchData, usize, usize) { + if matches!(dtype, Precision::Float32) { + log::warn!( + "{fmt} file loaded as f64, casting to f32: values outside f32 range become ±Inf." + ); + ( + BatchData::F32(data.iter().map(|&x| x as f32).collect()), + n, + s, + ) + } else { + (BatchData::F64(data), n, s) + } +} + fn read_file_by_extension( path: &Path, null_handling: NullHandling, -) -> Result<(Vec, usize, usize)> { + dtype: Precision, + encoding: Encoding, +) -> Result<(BatchData, usize, usize)> { + use crate::reader::DataReader; + // Basis values are integer state indices, not floats; f32's 24-bit mantissa + // corrupts indices above 2^24. Always read basis as f64, then reject an explicit + // f32 request whose indices would not survive the narrowing (see below). + let basis = matches!(encoding, Encoding::Basis); + let read_dtype = if basis { Precision::Float64 } else { dtype }; let ext_lower = path_extension_lower(path); let ext = ext_lower.as_deref(); - match ext { + let result = match ext { Some("parquet") => { - use crate::reader::DataReader; - let mut reader = crate::readers::ParquetReader::new(path, None, null_handling)?; - reader.read_batch() + if matches!(read_dtype, Precision::Float32) { + let mut reader = + crate::readers::ParquetReader::::new(path, None, null_handling)?; + let (data, n, s) = reader.read_batch()?; + (BatchData::F32(data), n, s) + } else { + let mut reader = + crate::readers::ParquetReader::::new(path, None, null_handling)?; + let (data, n, s) = reader.read_batch()?; + (BatchData::F64(data), n, s) + } } Some("arrow") | Some("feather") | Some("ipc") => { - use crate::reader::DataReader; let mut reader = crate::readers::ArrowIPCReader::new(path, null_handling)?; - reader.read_batch() + let (data, n, s) = reader.read_batch()?; + cast_f64_to_batch_data(data, n, s, read_dtype, "Arrow IPC") } - Some("npy") => io::read_numpy_batch(path), - Some("pt") | Some("pth") => io::read_torch_batch(path), - Some("pb") => io::read_tensorflow_batch(path), - _ => Err(MahoutError::InvalidInput(format!( - "Unsupported file extension {:?}. Supported: .parquet, .arrow, .feather, .ipc, .npy, .pt, .pth, .pb", - path.extension() - ))), + Some("npy") => { + let (data, n, s) = io::read_numpy_batch(path)?; + cast_f64_to_batch_data(data, n, s, read_dtype, "NumPy") + } + Some("pt") | Some("pth") => { + let (data, n, s) = io::read_torch_batch(path)?; + cast_f64_to_batch_data(data, n, s, read_dtype, "PyTorch") + } + Some("pb") => { + let (data, n, s) = io::read_tensorflow_batch(path)?; + cast_f64_to_batch_data(data, n, s, read_dtype, "TensorFlow") + } + _ => { + return Err(MahoutError::InvalidInput(format!( + "Unsupported file extension {:?}. Supported: .parquet, .arrow, .feather, .ipc, .npy, .pt, .pth, .pb", + path.extension() + ))); + } + }; + // basis is always read as f64 above; reject the f32 request only when the indices + // would actually have been corrupted by the narrowing the caller asked for. + if basis && matches!(dtype, Precision::Float32) { + reject_basis_f32_out_of_range(&result.0)?; } + Ok(result) } /// Stateful iterator that yields one batch DLPack at a time for Python `for` loop consumption. @@ -450,6 +562,89 @@ impl Drop for PipelineIterator { } } +/// Spawn an in-memory producer over already-loaded `data`. Mirrors +/// [`build_streaming_producer`] so the f32/f64 dispatch in `new_from_file` is a single +/// generic call instead of two byte-for-byte-identical match arms. +fn build_inmemory_producer( + data: Vec, + sample_size: usize, + config: &PipelineConfig, + batch_limit: usize, +) -> Result +where + T: FloatElem + ToBatchData, +{ + spawn_producer( + InMemoryProducer:: { + data, + cursor: 0, + sample_size, + batch_size: config.batch_size, + num_qubits: config.num_qubits as usize, + batches_yielded: 0, + batch_limit, + }, + config.prefetch_depth, + ) +} + +fn build_streaming_producer( + path: &Path, + config: &PipelineConfig, + batch_limit: usize, +) -> Result +where + T: FloatElem + ToBatchData, +{ + let mut reader = ParquetStreamingReader::::new( + path, + Some(DEFAULT_PARQUET_ROW_GROUP_SIZE), + config.null_handling, + )?; + let vector_len = vector_len(config.num_qubits, config.encoding); + + const INITIAL_CHUNK_CAP: usize = 64 * 1024; + // Buffer must hold at least one complete sample; for amplitude encoding with 17+ qubits + // vector_len (2^n) exceeds INITIAL_CHUNK_CAP and read_chunk would return 0 on a valid file. + let initial_cap = INITIAL_CHUNK_CAP.max(vector_len); + let mut buffer = vec![T::default(); initial_cap]; + let written = reader.read_chunk(&mut buffer)?; + if written == 0 { + return Err(MahoutError::InvalidInput( + "Parquet file is empty or contains no data.".to_string(), + )); + } + let sample_size = reader.get_sample_size().ok_or_else(|| { + MahoutError::InvalidInput( + "Parquet streaming reader did not set sample_size after first chunk.".to_string(), + ) + })?; + if sample_size != vector_len { + return Err(MahoutError::InvalidInput(format!( + "File feature length {} does not match vector_len {} for num_qubits={}, encoding={}", + sample_size, + vector_len, + config.num_qubits, + config.encoding.as_str() + ))); + } + + buffer.truncate(written); + let read_chunk_scratch = vec![T::default(); initial_cap]; + let producer = StreamingProducer:: { + reader, + buffer, + buffer_cursor: 0, + read_chunk_scratch, + sample_size, + batch_size: config.batch_size, + num_qubits: config.num_qubits as usize, + batches_yielded: 0, + batch_limit, + }; + spawn_producer(producer, config.prefetch_depth) +} + impl PipelineIterator { pub fn new_synthetic(engine: QdpEngine, mut config: PipelineConfig) -> Result { config.normalize(); @@ -479,10 +674,11 @@ impl PipelineIterator { ) -> Result { config.normalize(); let path = path.as_ref(); - let (data, num_samples, sample_size) = read_file_by_extension(path, config.null_handling)?; + let (batch_data, num_samples, sample_size) = + read_file_by_extension(path, config.null_handling, config.dtype, config.encoding)?; let vector_len = vector_len(config.num_qubits, config.encoding); - // Dimension validation at construction. + // Dimension validation before moving batch_data. if sample_size != vector_len { return Err(MahoutError::InvalidInput(format!( "File feature length {} does not match vector_len {} for num_qubits={}, encoding={}", @@ -492,26 +688,23 @@ impl PipelineIterator { config.encoding.as_str() ))); } - if data.len() != num_samples * sample_size { + if batch_data.len() != num_samples * sample_size { return Err(MahoutError::InvalidInput(format!( "File data length {} is not num_samples ({}) * sample_size ({})", - data.len(), + batch_data.len(), num_samples, sample_size ))); } - let producer = InMemoryProducer { - data, - cursor: 0, - sample_size, - batch_size: config.batch_size, - num_qubits: config.num_qubits as usize, - batches_yielded: 0, - batch_limit, + let (rx, recycle_tx, _producer_handle) = match batch_data { + BatchData::F32(data) => { + build_inmemory_producer::(data, sample_size, &config, batch_limit)? + } + BatchData::F64(data) => { + build_inmemory_producer::(data, sample_size, &config, batch_limit)? + } }; - let prefetch_depth = config.prefetch_depth; - let (rx, recycle_tx, _producer_handle) = spawn_producer(producer, prefetch_depth)?; Ok(Self { engine, config, @@ -539,54 +732,29 @@ impl PipelineIterator { ))); } - let mut reader = ParquetStreamingReader::new( - path, - Some(DEFAULT_PARQUET_ROW_GROUP_SIZE), - config.null_handling, - )?; - let vector_len = vector_len(config.num_qubits, config.encoding); - - // Read first chunk to learn sample_size; reuse as initial buffer. - const INITIAL_CHUNK_CAP: usize = 64 * 1024; - let mut buffer = vec![0.0; INITIAL_CHUNK_CAP]; - let written = reader.read_chunk(&mut buffer)?; - if written == 0 { - return Err(MahoutError::InvalidInput( - "Parquet file is empty or contains no data.".to_string(), - )); - } - let sample_size = reader.get_sample_size().ok_or_else(|| { - MahoutError::InvalidInput( - "Parquet streaming reader did not set sample_size after first chunk.".to_string(), - ) - })?; - - if sample_size != vector_len { - return Err(MahoutError::InvalidInput(format!( - "File feature length {} does not match vector_len {} for num_qubits={}, encoding={}", - sample_size, - vector_len, - config.num_qubits, - config.encoding.as_str() - ))); + // Basis values are integer state indices; f32 narrowing corrupts indices above + // 2^24 (see read_file_by_extension). The streaming reader is chunked so we cannot + // pre-scan the whole file, so for basis we always stream as f64 (lossless) rather + // than honoring an f32 request that could silently change states mid-stream. + // + // The non-streaming loader can pre-scan and so *rejects* an out-of-range f32 basis + // request; the chunked streaming path cannot, so it downgrades to f64 instead. Warn + // so the caller is not left believing the stream ran in f32. + let f32_requested = matches!(config.dtype, Precision::Float32); + let is_basis = matches!(config.encoding, Encoding::Basis); + if f32_requested && is_basis { + log::warn!( + "float32 requested for streaming basis file; basis indices are integers and f32 \ + cannot represent indices above {MAX_EXACT_F32_INT:.0} exactly, so this stream is \ + read as f64. Use dtype='float64' to silence this warning." + ); } - - buffer.truncate(written); - let read_chunk_scratch = vec![0.0; INITIAL_CHUNK_CAP]; - - let producer = StreamingProducer { - reader, - buffer, - buffer_cursor: 0, - read_chunk_scratch, - sample_size, - batch_size: config.batch_size, - num_qubits: config.num_qubits as usize, - batches_yielded: 0, - batch_limit, + let use_f32 = f32_requested && !is_basis; + let (rx, recycle_tx, _producer_handle) = if use_f32 { + build_streaming_producer::(path, &config, batch_limit)? + } else { + build_streaming_producer::(path, &config, batch_limit)? }; - let prefetch_depth = config.prefetch_depth; - let (rx, recycle_tx, _producer_handle) = spawn_producer(producer, prefetch_depth)?; Ok(Self { engine, config, @@ -1365,4 +1533,381 @@ mod tests { "normalize() must set prefetch_depth > 0" ); } + + // ------------------------------------------------------------------------- + // dtype file-load tests + // + // These tests verify that PipelineConfig.dtype is respected when loading + // from file sources. They stop at the BatchData variant boundary — the + // encode kernel (encode_batch_f32_for_pipeline) is CUDA-gated and cannot + // be exercised in CPU-only CI. BatchData::F32 is the observable proxy that + // confirms the f32 kernel would be called on a GPU host; this mirrors the + // existing convention in test_synthetic_producer_f32_*. + // ------------------------------------------------------------------------- + + mod dtype_file_tests { + use super::*; + use arrow::array::{ArrayRef, FixedSizeListBuilder, Float32Builder, RecordBatch}; + use arrow::datatypes::{DataType, Field, Schema}; + use parquet::arrow::ArrowWriter; + use std::fs; + use std::sync::Arc; + + fn write_f32_parquet(path: &std::path::Path) { + // 8 samples, each 4 features — matches amplitude encoding with 2 qubits (2^2=4) + let item_field = Arc::new(Field::new("item", DataType::Float32, true)); + let list_field = Field::new("data", DataType::FixedSizeList(item_field, 4), true); + let schema = Arc::new(Schema::new(vec![list_field])); + let mut builder = FixedSizeListBuilder::new(Float32Builder::new(), 4); + for _ in 0..8 { + builder.values().append_slice(&[0.25_f32, 0.5, 0.75, 1.0]); + builder.append(true); + } + let array = Arc::new(builder.finish()) as ArrayRef; + let batch = RecordBatch::try_new(schema.clone(), vec![array]).unwrap(); + let file = fs::File::create(path).unwrap(); + let mut writer = ArrowWriter::try_new(file, schema, None).unwrap(); + writer.write(&batch).unwrap(); + writer.close().unwrap(); + } + + static FILE_COUNTER: std::sync::atomic::AtomicUsize = + std::sync::atomic::AtomicUsize::new(0); + + fn temp_parquet_path(tag: &str) -> std::path::PathBuf { + let n = FILE_COUNTER.fetch_add(1, std::sync::atomic::Ordering::Relaxed); + std::env::temp_dir().join(format!( + "mahout_pipeline_dtype_{tag}_{pid}_{n}.parquet", + pid = std::process::id(), + )) + } + + #[test] + fn test_read_file_by_extension_f32_parquet_returns_f32_batch_data() { + let path = temp_parquet_path("f32"); + write_f32_parquet(&path); + let result = read_file_by_extension( + &path, + NullHandling::FillZero, + Precision::Float32, + Encoding::Amplitude, + ); + let _ = fs::remove_file(&path); + let (batch_data, num_samples, sample_size) = result.unwrap(); + // Assert the actual values, not just the variant: a zeroed/garbled F32 batch + // would still match `BatchData::F32(_)` but is not what we loaded. + match batch_data { + BatchData::F32(buf) => { + assert_eq!(buf.len(), 8 * 4); + assert!((buf[0] - 0.25).abs() < 1e-6, "first value must round-trip"); + assert_eq!(&buf[..4], &[0.25_f32, 0.5, 0.75, 1.0]); + } + other => { + panic!("dtype=Float32 + f32 Parquet must yield BatchData::F32, got {other:?}") + } + } + assert_eq!(num_samples, 8); + assert_eq!(sample_size, 4); + } + + #[test] + fn test_read_file_by_extension_f64_parquet_returns_f64_batch_data() { + use arrow::array::Float64Builder; + use arrow::datatypes::DataType; + let item_field = Arc::new(Field::new("item", DataType::Float64, true)); + let list_field = Field::new("data", DataType::FixedSizeList(item_field, 4), true); + let schema = Arc::new(Schema::new(vec![list_field])); + let mut builder = FixedSizeListBuilder::new(Float64Builder::new(), 4); + for _ in 0..8 { + builder.values().append_slice(&[0.1_f64, 0.2, 0.3, 0.4]); + builder.append(true); + } + let array = Arc::new(builder.finish()) as ArrayRef; + let batch = RecordBatch::try_new(schema.clone(), vec![array]).unwrap(); + let path = temp_parquet_path("f64"); + let file = fs::File::create(&path).unwrap(); + let mut writer = ArrowWriter::try_new(file, schema, None).unwrap(); + writer.write(&batch).unwrap(); + writer.close().unwrap(); + + let result = read_file_by_extension( + &path, + NullHandling::FillZero, + Precision::Float64, + Encoding::Amplitude, + ); + let _ = fs::remove_file(&path); + let (batch_data, num_samples, sample_size) = result.unwrap(); + match batch_data { + BatchData::F64(buf) => { + assert_eq!(buf.len(), 8 * 4); + assert!((buf[0] - 0.1).abs() < 1e-12, "first value must round-trip"); + assert_eq!(&buf[..4], &[0.1_f64, 0.2, 0.3, 0.4]); + } + other => { + panic!("dtype=Float64 must yield BatchData::F64 (regression), got {other:?}") + } + } + assert_eq!(num_samples, 8); + assert_eq!(sample_size, 4); + } + + /// Basis encoding has sample_size 1 (one integer state index per sample). + fn write_basis_parquet(path: &std::path::Path, indices: &[f64]) { + use arrow::array::Float64Builder; + let item_field = Arc::new(Field::new("item", DataType::Float64, true)); + let list_field = Field::new("data", DataType::FixedSizeList(item_field, 1), true); + let schema = Arc::new(Schema::new(vec![list_field])); + let mut builder = FixedSizeListBuilder::new(Float64Builder::new(), 1); + for &idx in indices { + builder.values().append_value(idx); + builder.append(true); + } + let array = Arc::new(builder.finish()) as ArrayRef; + let batch = RecordBatch::try_new(schema.clone(), vec![array]).unwrap(); + let file = fs::File::create(path).unwrap(); + let mut writer = ArrowWriter::try_new(file, schema, None).unwrap(); + writer.write(&batch).unwrap(); + writer.close().unwrap(); + } + + #[test] + fn test_basis_f32_request_kept_as_f64_when_indices_fit() { + // Indices <= 2^24 are exactly representable in f32, so an f32 request is + // lossless — but basis is always kept as f64 to route through the integer + // path. Values must round-trip exactly. + let path = temp_parquet_path("basis_small"); + write_basis_parquet(&path, &[0.0, 1.0, 1024.0, 16_777_216.0]); + let result = read_file_by_extension( + &path, + NullHandling::FillZero, + Precision::Float32, + Encoding::Basis, + ); + let _ = fs::remove_file(&path); + match result.unwrap().0 { + BatchData::F64(buf) => { + assert_eq!(buf, vec![0.0, 1.0, 1024.0, 16_777_216.0]); + } + other => panic!("basis must be kept as F64 to preserve indices, got {other:?}"), + } + } + + #[test] + fn test_basis_f32_request_rejected_when_index_exceeds_f32_range() { + // 16_777_217 = 2^24 + 1 cannot be represented exactly in f32 (becomes + // 16_777_216), so an explicit f32 request must be rejected, not silently + // corrupted. This is the bug from PR #1407 review item #1. + let path = temp_parquet_path("basis_big_f32"); + write_basis_parquet(&path, &[1.0, 16_777_217.0]); + let result = read_file_by_extension( + &path, + NullHandling::FillZero, + Precision::Float32, + Encoding::Basis, + ); + let _ = fs::remove_file(&path); + let err = result.expect_err("f32 basis with index > 2^24 must error"); + assert!( + matches!(err, MahoutError::InvalidInput(_)), + "expected InvalidInput, got {err:?}" + ); + } + + #[test] + fn test_basis_f64_large_index_loads_exactly() { + // The same large index under an f64 request loads fine and keeps its value. + let path = temp_parquet_path("basis_big_f64"); + write_basis_parquet(&path, &[1.0, 16_777_217.0]); + let result = read_file_by_extension( + &path, + NullHandling::FillZero, + Precision::Float64, + Encoding::Basis, + ); + let _ = fs::remove_file(&path); + match result.unwrap().0 { + BatchData::F64(buf) => assert_eq!(buf, vec![1.0, 16_777_217.0]), + other => panic!("f64 basis must yield BatchData::F64, got {other:?}"), + } + } + + #[test] + fn test_inmemory_producer_f32_produce_yields_f32_batch_data() { + let data = vec![0.25_f32, 0.5, 0.75, 1.0, 0.1, 0.2, 0.3, 0.4]; // 2 samples × 4 + let mut producer = InMemoryProducer:: { + data, + cursor: 0, + sample_size: 4, + batch_size: 2, + num_qubits: 2, + batches_yielded: 0, + batch_limit: 10, + }; + let batch = producer.produce(None).unwrap().unwrap(); + assert!( + matches!(batch.data, BatchData::F32(_)), + "InMemoryProducer::.produce() must yield BatchData::F32 \ + so that next_batch() routes to encode_batch_f32_for_pipeline" + ); + assert_eq!(batch.batch_n, 2); + assert_eq!(batch.sample_size, 4); + } + + #[test] + fn test_streaming_producer_f32_produce_yields_f32_batch_data() { + let path = temp_parquet_path("streaming_f32"); + write_f32_parquet(&path); + + let result = (|| -> Result { + let mut reader = ParquetStreamingReader::::new( + &path, + Some(DEFAULT_PARQUET_ROW_GROUP_SIZE), + NullHandling::FillZero, + )?; + const CAP: usize = 64 * 1024; + let mut buffer = vec![0.0_f32; CAP]; + let written = reader.read_chunk(&mut buffer)?; + if written == 0 { + return Err(MahoutError::InvalidInput("empty file".into())); + } + let sample_size = reader.get_sample_size().unwrap(); + buffer.truncate(written); + let scratch = vec![0.0_f32; CAP]; + let mut producer = StreamingProducer:: { + reader, + buffer, + buffer_cursor: 0, + read_chunk_scratch: scratch, + sample_size, + batch_size: 4, + num_qubits: 2, + batches_yielded: 0, + batch_limit: 10, + }; + let batch = producer.produce(None)?.unwrap(); + Ok(batch.data) + })(); + + let _ = fs::remove_file(&path); + assert!( + matches!(result.unwrap(), BatchData::F32(_)), + "StreamingProducer::.produce() must yield BatchData::F32" + ); + } + + /// Direct unit test of the shared f64→f32 narrowing helper used by every + /// non-Parquet format (Arrow IPC, NumPy, PyTorch, TensorFlow). Covers both the + /// variant switch and the documented ±Inf behavior for values outside f32 range. + #[test] + fn test_cast_f64_to_batch_data_narrows_and_overflows() { + // f32 request: in-range values cast exactly, out-of-range overflow to +Inf. + let (batch, n, s) = cast_f64_to_batch_data( + vec![0.25, 0.5, 1e40, -1e40], + 1, + 4, + Precision::Float32, + "test", + ); + match batch { + BatchData::F32(buf) => { + assert_eq!(buf[0], 0.25_f32); + assert_eq!(buf[1], 0.5_f32); + assert!( + buf[2].is_infinite() && buf[2] > 0.0, + "1e40 must overflow to +Inf" + ); + assert!( + buf[3].is_infinite() && buf[3] < 0.0, + "-1e40 must overflow to -Inf" + ); + } + other => panic!("Float32 request must yield BatchData::F32, got {other:?}"), + } + assert_eq!((n, s), (1, 4)); + + // f64 request: passthrough, no cast. + let (batch, _, _) = + cast_f64_to_batch_data(vec![0.1, 0.2], 1, 2, Precision::Float64, "test"); + assert_eq!(batch, BatchData::F64(vec![0.1, 0.2])); + } + + /// End-to-end coverage of the Arrow IPC → f32 path through `read_file_by_extension`: + /// the reader produces f64, the cast helper narrows it to BatchData::F32. + #[test] + fn test_read_file_by_extension_arrow_f32_narrows_to_f32_batch_data() { + use arrow::array::Float64Builder; + use arrow::ipc::writer::FileWriter as ArrowIpcFileWriter; + + let item_field = Arc::new(Field::new("item", DataType::Float64, true)); + let list_field = Field::new("data", DataType::FixedSizeList(item_field, 4), true); + let schema = Arc::new(Schema::new(vec![list_field])); + let mut builder = FixedSizeListBuilder::new(Float64Builder::new(), 4); + for _ in 0..3 { + builder.values().append_slice(&[0.25_f64, 0.5, 0.75, 1.0]); + builder.append(true); + } + let array = Arc::new(builder.finish()) as ArrayRef; + let batch = RecordBatch::try_new(schema.clone(), vec![array]).unwrap(); + + let n = FILE_COUNTER.fetch_add(1, std::sync::atomic::Ordering::Relaxed); + let path = std::env::temp_dir().join(format!( + "mahout_pipeline_dtype_arrow_{pid}_{n}.arrow", + pid = std::process::id(), + )); + { + let file = fs::File::create(&path).unwrap(); + let mut writer = ArrowIpcFileWriter::try_new(file, &schema).unwrap(); + writer.write(&batch).unwrap(); + writer.finish().unwrap(); + } + + let result = read_file_by_extension( + &path, + NullHandling::FillZero, + Precision::Float32, + Encoding::Amplitude, + ); + let _ = fs::remove_file(&path); + let (batch_data, num_samples, sample_size) = result.unwrap(); + match batch_data { + BatchData::F32(buf) => { + assert_eq!(buf.len(), 3 * 4); + assert_eq!(&buf[..4], &[0.25_f32, 0.5, 0.75, 1.0]); + } + other => { + panic!("Arrow IPC + dtype=Float32 must narrow to BatchData::F32, got {other:?}") + } + } + assert_eq!(num_samples, 3); + assert_eq!(sample_size, 4); + } + + #[test] + fn test_build_streaming_producer_f32_channel_yields_f32_batch_data() { + // BatchData::F32 is the observable proxy for the f32 GPU kernel path; the kernel + // itself is CUDA-gated and cannot be called in CPU-only CI. + let path = temp_parquet_path("build_sp_f32"); + write_f32_parquet(&path); + // 2 qubits + Amplitude → vector_len = 2^2 = 4, matching write_f32_parquet's 4 features. + let config = PipelineConfig { + num_qubits: 2, + encoding: Encoding::Amplitude, + batch_size: 4, + dtype: Precision::Float32, + prefetch_depth: 1, + ..PipelineConfig::default() + }; + let result = (|| -> Result { + let (rx, _recycle_tx, _handle) = + build_streaming_producer::(&path, &config, 1)?; + rx.recv().unwrap().map(|b| b.data) + })(); + let _ = fs::remove_file(&path); + assert!( + matches!(result.unwrap(), BatchData::F32(_)), + "build_streaming_producer:: must deliver BatchData::F32 through the channel" + ); + } + } } diff --git a/qdp/qdp-python/qumat_qdp/loader.py b/qdp/qdp-python/qumat_qdp/loader.py index cc01a405cb..a70233650c 100644 --- a/qdp/qdp-python/qumat_qdp/loader.py +++ b/qdp/qdp-python/qumat_qdp/loader.py @@ -44,6 +44,10 @@ # Seed must fit Rust u64: 0 <= seed <= 2^64 - 1. _U64_MAX = 2**64 - 1 +# Accepted dtype aliases for .dtype(); forwarded verbatim to the native loader, +# which parses them case-insensitively (Dtype::from_str_ci). +_VALID_DTYPES = frozenset({"float32", "f32", "float64", "f64"}) + # Canonical encoding names (must match Encoding enum in qdp-core/src/types.rs). _VALID_ENCODINGS: frozenset[str] = frozenset( {"amplitude", "angle", "basis", "iqp", "iqp-z", "phase"} @@ -302,6 +306,7 @@ def __init__( self._synthetic_requested = False # set True only by source_synthetic() self._file_requested = False self._null_handling: str | None = None + self._dtype: str | None = None # None -> native default (float64, lossless) self._backend_name: str = _BACKEND_RUST def qubits(self, n: int) -> QuantumDataLoader: @@ -410,6 +415,13 @@ def source_file(self, path: str, streaming: bool = False) -> QuantumDataLoader: Remote ``s3://`` and ``gs://`` paths are accepted when the native remote I/O feature is enabled; remote query strings and fragments are rejected. + Element precision is controlled by :meth:`dtype`. By default file input is + loaded as ``float64`` (lossless). Selecting ``dtype("float32")`` narrows f64 + file contents to f32 on load; values outside the f32 range become ``±Inf``. + ``basis`` encoding is exempt: its values are integer state indices, so it is + always loaded as f64 and an explicit ``float32`` request is rejected when an + index exceeds f32's exact integer range (``2**24``). + :param path: Local or supported remote input path. :param streaming: Whether to request native streaming file loading. :returns: ``self`` for fluent builder chaining. @@ -435,6 +447,25 @@ def source_file(self, path: str, streaming: bool = False) -> QuantumDataLoader: self._streaming_requested = streaming return self + def dtype(self, name: str) -> QuantumDataLoader: + """Set the element precision used when loading file sources. + + Applies to the native :meth:`source_file` path. ``"float64"`` (the default) + loads file contents losslessly; ``"float32"`` narrows them to f32 on load. + See :meth:`source_file` for the cast caveats, including the ``basis`` + exemption. + + :param name: ``"float32"``/``"f32"`` or ``"float64"``/``"f64"`` (case-insensitive). + :returns: ``self`` for fluent builder chaining. + :raises ValueError: If ``name`` is not a recognized dtype. + """ + if not isinstance(name, str) or name.strip().lower() not in _VALID_DTYPES: + raise ValueError( + f"dtype must be one of {sorted(_VALID_DTYPES)}, got {name!r}" + ) + self._dtype = name.strip().lower() + return self + def seed(self, s: int | None = None) -> QuantumDataLoader: """Set or clear the synthetic data seed. @@ -580,6 +611,7 @@ def _create_rust_iterator(self, QdpEngine, use_synthetic: bool) -> Iterator[obje encoding_method=self._encoding_method, batch_limit=None, null_handling=self._null_handling, + dtype=self._dtype, ) ) create_synthetic_loader = getattr(engine, "create_synthetic_loader", None) diff --git a/qdp/qdp-python/src/engine.rs b/qdp/qdp-python/src/engine.rs index 8297bf5e78..e3b9a9bd7a 100644 --- a/qdp/qdp-python/src/engine.rs +++ b/qdp/qdp-python/src/engine.rs @@ -25,7 +25,9 @@ use pyo3::prelude::*; use qdp_core::{Dtype, Encoding, QdpEngine as CoreEngine}; #[cfg(target_os = "linux")] -use crate::loader::{PyQuantumLoader, config_from_args, parse_null_handling, path_from_py}; +use crate::loader::{ + PyQuantumLoader, config_from_args, parse_dtype, parse_null_handling, path_from_py, +}; /// PyO3 wrapper for QdpEngine /// @@ -535,6 +537,10 @@ impl QdpEngine { null_handling: Option<&str>, ) -> PyResult { let nh = parse_null_handling(null_handling)?; + // Synthetic data is generated in-process for throughput benchmarking, so it + // defaults to f32 (PipelineConfig::normalize downgrades to f64 for encodings + // without an f32 batch path). This is deliberate and unrelated to the file + // loaders, which default to f64 to keep user-supplied data lossless. let config = config_from_args( &self.engine, batch_size, @@ -554,7 +560,7 @@ impl QdpEngine { #[cfg(target_os = "linux")] /// Create a file-backed pipeline iterator (full read then batch; for QuantumDataLoader.source_file(path)). #[allow(clippy::too_many_arguments)] - #[pyo3(signature = (path, batch_size, num_qubits, encoding_method, batch_limit=None, null_handling=None))] + #[pyo3(signature = (path, batch_size, num_qubits, encoding_method, batch_limit=None, null_handling=None, dtype=None))] fn create_file_loader( &self, py: Python<'_>, @@ -564,10 +570,12 @@ impl QdpEngine { encoding_method: &str, batch_limit: Option, null_handling: Option<&str>, + dtype: Option<&str>, ) -> PyResult { let path_str = path_from_py(path)?; let batch_limit = batch_limit.unwrap_or(usize::MAX); let nh = parse_null_handling(null_handling)?; + let dt = parse_dtype(dtype)?; let config = config_from_args( &self.engine, batch_size, @@ -576,7 +584,7 @@ impl QdpEngine { 0, None, nh, - Dtype::Float32, + dt, )?; let engine = self.engine.clone(); // Resolve remote URLs before detaching from GIL. The _resolved guard keeps the @@ -603,7 +611,7 @@ impl QdpEngine { #[cfg(target_os = "linux")] /// Create a streaming Parquet pipeline iterator (for QuantumDataLoader.source_file(path, streaming=True)). #[allow(clippy::too_many_arguments)] - #[pyo3(signature = (path, batch_size, num_qubits, encoding_method, batch_limit=None, null_handling=None))] + #[pyo3(signature = (path, batch_size, num_qubits, encoding_method, batch_limit=None, null_handling=None, dtype=None))] fn create_streaming_file_loader( &self, py: Python<'_>, @@ -613,10 +621,12 @@ impl QdpEngine { encoding_method: &str, batch_limit: Option, null_handling: Option<&str>, + dtype: Option<&str>, ) -> PyResult { let path_str = path_from_py(path)?; let batch_limit = batch_limit.unwrap_or(usize::MAX); let nh = parse_null_handling(null_handling)?; + let dt = parse_dtype(dtype)?; let config = config_from_args( &self.engine, batch_size, @@ -625,7 +635,7 @@ impl QdpEngine { 0, None, nh, - Dtype::Float32, + dt, )?; let engine = self.engine.clone(); // Resolve remote URLs before detaching from GIL. The _resolved guard keeps the diff --git a/qdp/qdp-python/src/loader.rs b/qdp/qdp-python/src/loader.rs index a43f947944..466d7ea640 100644 --- a/qdp/qdp-python/src/loader.rs +++ b/qdp/qdp-python/src/loader.rs @@ -83,6 +83,19 @@ mod loader_impl { } } + /// Parse an optional Python dtype string into the Rust enum. Defaults to f64 so + /// file loads are lossless unless the caller explicitly opts into f32 narrowing. + pub fn parse_dtype(s: Option<&str>) -> PyResult { + match s { + None => Ok(Dtype::Float64), + Some(v) => Dtype::from_str_ci(v).map_err(|e| { + pyo3::exceptions::PyValueError::new_err(format!( + "Invalid dtype '{v}': {e}. Expected 'float32'/'f32' or 'float64'/'f64'." + )) + }), + } + } + /// Build PipelineConfig from Python args. device_id is 0 (engine does not expose it); iterator uses engine clone with correct device. #[allow(clippy::too_many_arguments)] pub fn config_from_args( @@ -121,4 +134,6 @@ mod loader_impl { } #[cfg(target_os = "linux")] -pub use loader_impl::{PyQuantumLoader, config_from_args, parse_null_handling, path_from_py}; +pub use loader_impl::{ + PyQuantumLoader, config_from_args, parse_dtype, parse_null_handling, path_from_py, +};