diff --git a/crates/pyxlog/python/pyxlog/_native.pyi b/crates/pyxlog/python/pyxlog/_native.pyi index 7098d9dd4..f457267ad 100644 --- a/crates/pyxlog/python/pyxlog/_native.pyi +++ b/crates/pyxlog/python/pyxlog/_native.pyi @@ -348,6 +348,10 @@ class LogicRelationSession: """Snapshot DLPack columns into a persistent session relation.""" ... + def put_relation_rows(self, name: str, rows: Sequence[Sequence[str]]) -> None: + """Parse typed lexical rows and upload them into a persistent relation.""" + ... + def put_relation_with_provenance( self, name: str, @@ -512,6 +516,10 @@ class LogicRelationSession: """Export the named relation as a list of DLPack column capsules.""" ... + def export_relation_rows(self, name: str) -> list[list[str]]: + """Download one stored or materialized relation as typed lexical rows.""" + ... + def export_relation_with_provenance( self, name: str ) -> _RelationProvenanceExport: diff --git a/crates/pyxlog/src/logic.rs b/crates/pyxlog/src/logic.rs index a09bdc406..13e4b47df 100644 --- a/crates/pyxlog/src/logic.rs +++ b/crates/pyxlog/src/logic.rs @@ -7,6 +7,7 @@ use pyo3::exceptions::{PyKeyError, PyRuntimeError, PyValueError}; use pyo3::prelude::*; use pyo3::types::{PyDict, PySequence}; +use xlog_core::{symbol, ScalarType}; use xlog_cuda::DlpackManagedTensor; use xlog_gpu::logic as gpu_logic; use xlog_logic::ast::ProbEngine; @@ -282,6 +283,12 @@ impl LogicRelationSession { self.commit_relation_replacement(name, buffer, RelationReplacementMetadata::Clear) } + pub fn put_relation_rows(&mut self, name: String, rows: Vec>) -> PyResult<()> { + let schema = self.relation_replacement_schema(&name)?.clone(); + let buffer = self.lexical_relation_buffer(&name, &schema, rows)?; + self.commit_relation_replacement(name, buffer, RelationReplacementMetadata::Clear) + } + #[pyo3(signature = (name, dlpack_columns, *, roles, facts))] pub fn put_relation_with_provenance( &mut self, @@ -833,6 +840,20 @@ impl LogicRelationSession { export_buffer_columns(py, &self.provider, buffer) } + pub fn export_relation_rows(&self, name: &str) -> PyResult>> { + let buffer = self + .evaluation_store + .as_ref() + .and_then(|store| store.as_relation_store().get(name)) + .or_else(|| self.relation_store.get(name)) + .ok_or_else(|| { + PyValueError::new_err(format!( + "Relation '{name}' is unavailable for lexical export" + )) + })?; + lexical_relation_rows(&self.provider, buffer) + } + pub fn export_relation_with_provenance( &mut self, py: Python<'_>, @@ -952,6 +973,55 @@ impl LogicRelationSession { .map_err(types::xlog_err) } + fn lexical_relation_buffer( + &self, + name: &str, + schema: &xlog_core::Schema, + rows: Vec>, + ) -> PyResult { + if rows.iter().any(|row| row.len() != schema.arity()) { + return Err(PyValueError::new_err(format!( + "Relation {name} lexical row arity differs from compiled arity {}", + schema.arity() + ))); + } + if schema.arity() == 0 { + let row_count = u32::try_from(rows.len()).map_err(|_| { + PyValueError::new_err(format!("Relation {name} row count exceeds u32::MAX")) + })?; + return self + .provider + .create_zero_arity_buffer(schema.clone(), row_count) + .map_err(types::xlog_err); + } + let mut columns = schema + .columns + .iter() + .map(|_| Vec::::new()) + .collect::>(); + for (row_index, row) in rows.iter().enumerate() { + for (column_index, value) in row.iter().enumerate() { + let scalar_type = schema.column_type(column_index).ok_or_else(|| { + PyRuntimeError::new_err(format!( + "Relation {name} column {column_index} schema disappeared" + )) + })?; + append_lexical_scalar( + &mut columns[column_index], + value, + scalar_type, + name, + row_index, + column_index, + )?; + } + } + let slices = columns.iter().map(Vec::as_slice).collect::>(); + self.provider + .create_buffer_from_slices(&slices, schema.clone()) + .map_err(types::xlog_err) + } + fn detached_relation_replacement_buffer( &self, name: &str, @@ -1385,6 +1455,158 @@ fn export_buffer_columns( Ok(tensors) } +fn lexical_parse_error( + relation: &str, + row_index: usize, + column_index: usize, + scalar_type: ScalarType, + value: &str, +) -> PyErr { + PyValueError::new_err(format!( + "Relation {relation} row {row_index} column {column_index} value {value:?} is not valid {scalar_type:?}" + )) +} + +fn append_lexical_scalar( + output: &mut Vec, + value: &str, + scalar_type: ScalarType, + relation: &str, + row_index: usize, + column_index: usize, +) -> PyResult<()> { + macro_rules! parse_scalar { + ($target:ty) => {{ + let parsed = value.parse::<$target>().map_err(|_| { + lexical_parse_error(relation, row_index, column_index, scalar_type, value) + })?; + output.extend_from_slice(&parsed.to_le_bytes()); + }}; + } + match scalar_type { + ScalarType::U32 => parse_scalar!(u32), + ScalarType::U64 => parse_scalar!(u64), + ScalarType::I32 => parse_scalar!(i32), + ScalarType::I64 => parse_scalar!(i64), + ScalarType::F32 => { + let parsed = value.parse::().map_err(|_| { + lexical_parse_error(relation, row_index, column_index, scalar_type, value) + })?; + if !parsed.is_finite() { + return Err(lexical_parse_error( + relation, + row_index, + column_index, + scalar_type, + value, + )); + } + output.extend_from_slice(&parsed.to_le_bytes()); + } + ScalarType::F64 => { + let parsed = value.parse::().map_err(|_| { + lexical_parse_error(relation, row_index, column_index, scalar_type, value) + })?; + if !parsed.is_finite() { + return Err(lexical_parse_error( + relation, + row_index, + column_index, + scalar_type, + value, + )); + } + output.extend_from_slice(&parsed.to_le_bytes()); + } + ScalarType::Bool => match value { + "true" => output.push(1), + "false" => output.push(0), + _ => { + return Err(lexical_parse_error( + relation, + row_index, + column_index, + scalar_type, + value, + )) + } + }, + ScalarType::Symbol => output.extend_from_slice(&symbol::intern(value).to_le_bytes()), + } + Ok(()) +} + +fn lexical_relation_rows( + provider: &Arc, + buffer: &xlog_cuda::CudaBuffer, +) -> PyResult>> { + let row_count = provider + .validated_logical_row_count(buffer) + .map_err(types::xlog_err)?; + let mut rows = vec![Vec::with_capacity(buffer.arity()); row_count]; + for column_index in 0..buffer.arity() { + let scalar_type = buffer.schema().column_type(column_index).ok_or_else(|| { + PyRuntimeError::new_err(format!( + "Relation export column {column_index} schema disappeared" + )) + })?; + let values: Vec = match scalar_type { + ScalarType::U32 => provider + .download_column::(buffer, column_index) + .map(|values| values.into_iter().map(|value| value.to_string()).collect()), + ScalarType::U64 => provider + .download_column::(buffer, column_index) + .map(|values| values.into_iter().map(|value| value.to_string()).collect()), + ScalarType::I32 => provider + .download_column::(buffer, column_index) + .map(|values| values.into_iter().map(|value| value.to_string()).collect()), + ScalarType::I64 => provider + .download_column::(buffer, column_index) + .map(|values| values.into_iter().map(|value| value.to_string()).collect()), + ScalarType::F32 => provider + .download_column::(buffer, column_index) + .map(|values| values.into_iter().map(|value| value.to_string()).collect()), + ScalarType::F64 => provider + .download_column::(buffer, column_index) + .map(|values| values.into_iter().map(|value| value.to_string()).collect()), + ScalarType::Bool => { + provider + .download_column::(buffer, column_index) + .map(|values| { + values + .into_iter() + .map(|value| if value { "true" } else { "false" }.to_string()) + .collect() + }) + } + ScalarType::Symbol => provider + .download_column::(buffer, column_index) + .and_then(|values| { + values + .into_iter() + .map(|value| { + symbol::resolve_checked(value).ok_or_else(|| { + xlog_core::XlogError::Execution(format!( + "Relation export contains unknown symbol ID {value}" + )) + }) + }) + .collect() + }), + } + .map_err(types::xlog_err)?; + if values.len() != row_count { + return Err(PyRuntimeError::new_err(format!( + "Relation export column {column_index} row count changed" + ))); + } + for (row, value) in rows.iter_mut().zip(values) { + row.push(value); + } + } + Ok(rows) +} + fn pack_logic_result_with_provider( py: Python<'_>, provider: &Arc, diff --git a/python/tests/test_pyxlog_host_rows.py b/python/tests/test_pyxlog_host_rows.py new file mode 100644 index 000000000..e1bbf49aa --- /dev/null +++ b/python/tests/test_pyxlog_host_rows.py @@ -0,0 +1,80 @@ +from __future__ import annotations + +import os + +import pytest + + +pyxlog = pytest.importorskip("pyxlog") + + +def _compile_or_skip(entrypoint, tmp_path): + try: + return pyxlog.LogicProgram.compile_file( + entrypoint, + module_paths=[tmp_path], + device=0, + memory_mb=512, + ) + except RuntimeError as exc: + if os.environ.get("XLOG_REQUIRE_CUDA") == "1": + raise + pytest.skip(f"CUDA is unavailable to Pyxlog: {exc}") + + +def test_session_host_rows_keep_exact_evaluation_inside_zero_transfer_window( + tmp_path, +) -> None: + entrypoint = tmp_path / "main.xlog" + entrypoint.write_text( + "pred source(symbol, u32, u64, i32, i64, f32, f64, bool).\n" + "pred result(symbol, u32, u64, i32, i64, f32, f64, bool).\n" + "result(A, B, C, D, E, F, G, H) :- source(A, B, C, D, E, F, G, H).\n" + "?- result(A, B, C, D, E, F, G, H).\n", + encoding="utf-8", + ) + compiled = _compile_or_skip(entrypoint, tmp_path) + session = compiled.session() + + with pytest.raises(ValueError, match="lexical row arity"): + session.put_relation_rows("source", [["too", "short"]]) + with pytest.raises(ValueError, match="not valid F32"): + session.put_relation_rows( + "source", + [["alpha", "1", "2", "-3", "-4", "nan", "2.5", "true"]], + ) + with pytest.raises(ValueError, match="not valid Bool"): + session.put_relation_rows( + "source", + [["alpha", "1", "2", "-3", "-4", "1.5", "2.5", "yes"]], + ) + + session.put_relation_rows( + "source", + [ + ["alpha", "1", "2", "-3", "-4", "1.5", "2.5", "true"], + ["beta", "5", "6", "-7", "-8", "-1.25", "-2.75", "false"], + ], + ) + + session.reset_host_transfer_stats() + session.set_strict_deterministic_d2h(True) + session.reset_deterministic_d2h_violations() + evaluated = session.evaluate() + stats = session.host_transfer_stats() + + assert len(evaluated.queries) == 1 + assert stats == { + "dtoh_bytes": 0, + "htod_bytes": 0, + "dtoh_calls": 0, + "htod_calls": 0, + } + assert session.strict_deterministic_d2h_enabled() is True + assert session.deterministic_d2h_violation_count() == 0 + + session.set_strict_deterministic_d2h(False) + assert session.export_relation_rows("__xlog_query_0") == [ + ["alpha", "1", "2", "-3", "-4", "1.5", "2.5", "true"], + ["beta", "5", "6", "-7", "-8", "-1.25", "-2.75", "false"], + ]