1use std::collections::{HashMap, HashSet};
4use std::os::raw::{c_char, c_void};
5use std::sync::Arc;
6
7use pyo3::exceptions::{PyBufferError, PyMemoryError, PyRuntimeError, PyValueError};
8use pyo3::prelude::*;
9use pyo3::types::{PyBool, PyDict, PyInt, PyList, PyTuple};
10
11use xlog_core::{MemoryBudget, Schema};
12#[cfg(feature = "arrow-device-import")]
13use xlog_cuda::{ArrowDeviceArray, ArrowDeviceArrayOwned};
14use xlog_cuda::{CudaBuffer, CudaKernelProvider, CudaProviderBuilder, DlpackManagedTensor};
15use xlog_gpu::logic as gpu_logic;
16use xlog_logic::ast::ProbEngine;
17use xlog_neural::{NetworkRegistry, TensorSourceRegistry};
18use xlog_prob::exact::GpuConfig;
19
20use xlog_core::RelId;
21use xlog_ir::ExecutionPlan;
22use xlog_logic::ast::Program as AstProgram;
23use xlog_runtime::{Executor, RelationStore};
24
25mod neural_registry;
26use neural_registry::NeuralPredicateRegistry;
27mod diagnostic_serialization;
28mod dlpack;
29mod epistemic;
30mod ilp;
31mod ilp_exact;
32mod ilp_gpu;
33mod joint_carrier;
34mod logic;
35mod neural;
36mod program;
37mod relation_metadata;
38mod semantic_transition;
39mod training;
40mod types;
41pub(crate) use diagnostic_serialization::{pack_query_proof_traces, pack_rule_provenance};
42pub(crate) use program::{
43 CachedCircuit, CompiledProbProgram, HardFilter, InputSource, JoinPlan, NeuralGroup,
44 QuerySignature,
45};
46use relation_metadata::RelationMetadataStore;
47
48const DLPACK_CAPSULE_NAME: &[u8] = b"dltensor\0";
49const USED_DLPACK_CAPSULE_NAME: &[u8] = b"used_dltensor\0";
50const DLPACK_CUDA_LEGACY_DEFAULT_STREAM: i64 = 1;
53
54#[cfg(feature = "arrow-device-import")]
55const ARROW_DEVICE_ARRAY_CAPSULE_NAME: &[u8] = b"arrow_device_array\0";
56#[cfg(feature = "arrow-device-import")]
57const USED_ARROW_DEVICE_ARRAY_CAPSULE_NAME: &[u8] = b"used_arrow_device_array\0";
58
59unsafe extern "C" fn dlpack_capsule_destructor(capsule: *mut pyo3::ffi::PyObject) {
60 if capsule.is_null() {
61 return;
62 }
63
64 let valid =
65 pyo3::ffi::PyCapsule_IsValid(capsule, DLPACK_CAPSULE_NAME.as_ptr() as *const c_char);
66 if valid == 0 {
67 return;
68 }
69
70 let ptr =
71 pyo3::ffi::PyCapsule_GetPointer(capsule, DLPACK_CAPSULE_NAME.as_ptr() as *const c_char);
72 if ptr.is_null() {
73 pyo3::ffi::PyErr_Clear();
74 return;
75 }
76
77 let managed = ptr as *mut xlog_cuda::DLManagedTensor;
78 drop(DlpackManagedTensor::from_raw(managed));
79}
80
81pub(crate) fn dlpack_capsule_from_tensor(
82 py: Python<'_>,
83 tensor: DlpackManagedTensor,
84) -> PyResult<Py<PyAny>> {
85 let raw = tensor.into_raw();
86 let ptr = raw as *mut c_void;
87 let capsule = unsafe {
89 pyo3::ffi::PyCapsule_New(
90 ptr,
91 DLPACK_CAPSULE_NAME.as_ptr() as *const c_char,
92 Some(dlpack_capsule_destructor),
93 )
94 };
95 if capsule.is_null() {
96 unsafe {
98 drop(DlpackManagedTensor::from_raw(raw));
99 }
100 return Err(PyRuntimeError::new_err("Failed to create DLPack capsule"));
101 }
102 let obj: Py<PyAny> = unsafe { Bound::from_owned_ptr(py, capsule) }.unbind();
104 Ok(obj)
105}
106
107#[cfg(feature = "arrow-device-import")]
108unsafe extern "C" fn arrow_device_array_capsule_destructor(capsule: *mut pyo3::ffi::PyObject) {
109 if capsule.is_null() {
110 return;
111 }
112
113 let valid = pyo3::ffi::PyCapsule_IsValid(
114 capsule,
115 ARROW_DEVICE_ARRAY_CAPSULE_NAME.as_ptr() as *const c_char,
116 );
117 if valid == 0 {
118 return;
119 }
120
121 let ptr = pyo3::ffi::PyCapsule_GetPointer(
122 capsule,
123 ARROW_DEVICE_ARRAY_CAPSULE_NAME.as_ptr() as *const c_char,
124 );
125 if ptr.is_null() {
126 pyo3::ffi::PyErr_Clear();
127 return;
128 }
129
130 drop(ArrowDeviceArrayOwned::from_raw(
131 ptr as *mut ArrowDeviceArray,
132 ));
133}
134
135#[cfg(feature = "arrow-device-import")]
136pub(crate) fn arrow_device_capsule_from_device_array(
137 py: Python<'_>,
138 device_array: ArrowDeviceArrayOwned,
139) -> PyResult<Py<PyAny>> {
140 let raw = device_array.into_raw();
141 let ptr = raw as *mut c_void;
142 let capsule = unsafe {
144 pyo3::ffi::PyCapsule_New(
145 ptr,
146 ARROW_DEVICE_ARRAY_CAPSULE_NAME.as_ptr() as *const c_char,
147 Some(arrow_device_array_capsule_destructor),
148 )
149 };
150 if capsule.is_null() {
151 unsafe {
153 drop(ArrowDeviceArrayOwned::from_raw(raw));
154 }
155 return Err(PyRuntimeError::new_err(
156 "Failed to create Arrow device array capsule",
157 ));
158 }
159 let obj: Py<PyAny> = unsafe { Bound::from_owned_ptr(py, capsule) }.unbind();
161 Ok(obj)
162}
163
164#[cfg(feature = "arrow-device-import")]
165pub(crate) fn arrow_device_from_py(obj: &Bound<'_, PyAny>) -> PyResult<ArrowDeviceArrayOwned> {
166 if unsafe {
168 pyo3::ffi::PyCapsule_IsValid(
169 obj.as_ptr(),
170 ARROW_DEVICE_ARRAY_CAPSULE_NAME.as_ptr() as *const c_char,
171 )
172 } == 0
173 {
174 return Err(PyValueError::new_err(
175 "Expected an Arrow device array capsule (arrow_device_array)",
176 ));
177 }
178
179 let ptr = unsafe {
181 pyo3::ffi::PyCapsule_GetPointer(
182 obj.as_ptr(),
183 ARROW_DEVICE_ARRAY_CAPSULE_NAME.as_ptr() as *const c_char,
184 )
185 };
186 if ptr.is_null() {
187 return Err(PyRuntimeError::new_err(
188 "Failed to get Arrow device array pointer",
189 ));
190 }
191
192 let rc = unsafe {
195 pyo3::ffi::PyCapsule_SetName(
196 obj.as_ptr(),
197 USED_ARROW_DEVICE_ARRAY_CAPSULE_NAME.as_ptr() as *const c_char,
198 )
199 };
200 if rc != 0 {
201 return Err(PyRuntimeError::new_err(
202 "Failed to mark Arrow device array capsule as consumed",
203 ));
204 }
205
206 Ok(unsafe { ArrowDeviceArrayOwned::from_raw(ptr as *mut ArrowDeviceArray) })
208}
209
210pub(crate) fn provider_from_config(config: GpuConfig) -> xlog_core::Result<CudaKernelProvider> {
211 CudaProviderBuilder::new(
212 config.device_ordinal,
213 MemoryBudget::with_limit(config.memory_bytes),
214 )
215 .build()
216}
217
218pub(crate) fn enforce_call_memory_limit(
219 provider: &Arc<CudaKernelProvider>,
220 memory_mb: Option<u64>,
221) -> PyResult<()> {
222 let Some(memory_mb) = memory_mb else {
223 return Ok(());
224 };
225 if memory_mb == 0 {
226 return Err(PyValueError::new_err("memory_mb must be > 0"));
227 }
228 let memory_limit_bytes = memory_mb.saturating_mul(1024 * 1024);
229 let allocated_bytes = provider.memory().allocated_bytes();
230 if allocated_bytes > memory_limit_bytes {
231 return Err(PyMemoryError::new_err(format!(
232 "per-call memory limit exceeded before evaluation: allocated_bytes={} memory_limit_bytes={}",
233 allocated_bytes, memory_limit_bytes
234 )));
235 }
236 Ok(())
237}
238
239pub(crate) fn provider_memory_stats(
240 py: Python<'_>,
241 provider: &Arc<CudaKernelProvider>,
242) -> PyResult<Py<PyAny>> {
243 let dict = PyDict::new(py);
244 let memory = provider.memory();
245 dict.set_item("allocated_bytes", memory.allocated_bytes())?;
246 dict.set_item("memory_limit_bytes", memory.budget().device_bytes)?;
247 dict.set_item("peak_memory_bytes", memory.peak_bytes())?;
248 dict.set_item("status", "available")?;
249 Ok(dict.into())
250}
251
252pub(crate) fn parse_prob_engine_override(s: &str) -> PyResult<ProbEngine> {
253 let v = s.trim().to_ascii_lowercase();
254 match v.as_str() {
255 "exact_ddnnf" | "exact" | "ddnnf" => Ok(ProbEngine::ExactDdnnf),
256 "mc" => Ok(ProbEngine::Mc),
257 other => Err(PyValueError::new_err(format!(
258 "Unknown prob_engine '{}'; expected 'exact_ddnnf' or 'mc'",
259 other
260 ))),
261 }
262}
263
264pub(crate) fn guarded_python_callback<T>(
267 check: impl Fn() -> PyResult<()>,
268 callback: impl FnOnce() -> PyResult<T>,
269) -> PyResult<T> {
270 check()?;
271 let result = callback();
272 let after = check();
273 match result {
274 Err(error) => Err(error),
275 Ok(value) => after.map(|()| value),
276 }
277}
278
279#[cfg(test)]
281pub(crate) fn dlpack_export_for_stream(
282 obj: &Bound<'_, PyAny>,
283 stream: Option<i64>,
284) -> PyResult<Py<PyAny>> {
285 dlpack_export_for_stream_guarded(obj, stream, &|| Ok(()))
286}
287
288fn dlpack_export_for_stream_guarded(
289 obj: &Bound<'_, PyAny>,
290 stream: Option<i64>,
291 check: &dyn Fn() -> PyResult<()>,
292) -> PyResult<Py<PyAny>> {
293 check()?;
294 let kwargs = PyDict::new(obj.py());
295 if let Some(stream) = stream {
296 kwargs.set_item("stream", stream)?;
297 }
298 let sys = guarded_python_callback(check, || obj.py().import("sys"))?;
305 let modules = guarded_python_callback(check, || sys.getattr("modules"))?;
306 let mut producer = obj.clone();
307 if let Some(torch) =
308 guarded_python_callback(check, || modules.cast::<PyDict>()?.get_item("torch"))?
309 {
310 let tensor_type = guarded_python_callback(check, || torch.getattr("Tensor"))?;
311 let is_tensor = guarded_python_callback(check, || obj.is_instance(&tensor_type))?;
312 let requires_grad = if is_tensor {
313 let value = guarded_python_callback(check, || obj.getattr("requires_grad"))?;
314 guarded_python_callback(check, || value.extract::<bool>())?
315 } else {
316 false
317 };
318 if requires_grad {
319 let detach = guarded_python_callback(check, || tensor_type.getattr("detach"))?;
320 let alias = guarded_python_callback(check, || detach.call1((obj,)))?;
321 for method in ["data_ptr", "storage_offset", "stride"] {
322 let alias_method = guarded_python_callback(check, || alias.getattr(method))?;
323 let alias_value = guarded_python_callback(check, || alias_method.call0())?;
324 let original_method = guarded_python_callback(check, || obj.getattr(method))?;
325 let original_value = guarded_python_callback(check, || original_method.call0())?;
326 let equal = guarded_python_callback(check, || {
327 alias_value.rich_compare(&original_value, pyo3::basic::CompareOp::Eq)
328 })?;
329 if !guarded_python_callback(check, || equal.is_truthy())? {
330 return Err(PyBufferError::new_err(
331 "DLPack transport alias changed the original tensor storage",
332 ));
333 }
334 }
335 for attribute in ["shape", "dtype", "device"] {
336 let alias_value = guarded_python_callback(check, || alias.getattr(attribute))?;
337 let original_value = guarded_python_callback(check, || obj.getattr(attribute))?;
338 let equal = guarded_python_callback(check, || {
339 alias_value.rich_compare(&original_value, pyo3::basic::CompareOp::Eq)
340 })?;
341 if !guarded_python_callback(check, || equal.is_truthy())? {
342 return Err(PyBufferError::new_err(
343 "DLPack transport alias changed the original tensor layout",
344 ));
345 }
346 }
347 producer = alias;
348 }
349 }
350 let export = guarded_python_callback(check, || producer.getattr("__dlpack__"))?;
351 Ok(guarded_python_callback(check, || export.call((), Some(&kwargs)))?.unbind())
352}
353
354pub(crate) fn dlpack_from_py(obj: &Bound<'_, PyAny>) -> PyResult<DlpackManagedTensor> {
355 dlpack_from_py_for_stream(obj, DLPACK_CUDA_LEGACY_DEFAULT_STREAM)
356}
357
358pub(crate) fn dlpack_from_py_for_stream(
360 obj: &Bound<'_, PyAny>,
361 consumer_stream: i64,
362) -> PyResult<DlpackManagedTensor> {
363 dlpack_from_py_for_stream_guarded(obj, consumer_stream, &|| Ok(()))
364}
365
366pub(crate) fn dlpack_device_pair(device: &Bound<'_, PyAny>) -> PyResult<(i32, i32)> {
369 let invalid_pair = || {
370 PyValueError::new_err(
371 "DLPack device must be an exact tuple (integer kind, exact integer id)",
372 )
373 };
374 if !device.is_exact_instance_of::<PyTuple>() {
375 return Err(invalid_pair());
376 }
377 let fields = device.cast::<PyTuple>()?;
378 if fields.len() != 2 {
379 return Err(invalid_pair());
380 }
381 let kind = fields.get_item(0)?;
382 let ordinal = fields.get_item(1)?;
383 if kind.is_instance_of::<PyBool>()
384 || !kind.is_instance_of::<PyInt>()
385 || !ordinal.is_exact_instance_of::<PyInt>()
386 {
387 return Err(invalid_pair());
388 }
389 let device_type = kind.extract::<i32>()?;
392 let device_id = ordinal.extract::<i32>()?;
393 if device_type != xlog_cuda::dlpack::K_DLCUDA {
394 return Err(PyBufferError::new_err(format!(
395 "Unsupported DLPack producer device type {device_type} (device {device_id}); \
396 XLOG requires CUDA device memory (kDLCUDA=2)"
397 )));
398 }
399 if device_id < 0 {
400 return Err(PyValueError::new_err(
401 "DLPack device id must be nonnegative",
402 ));
403 }
404 Ok((device_type, device_id))
405}
406
407pub(crate) fn dlpack_from_py_for_stream_guarded(
408 obj: &Bound<'_, PyAny>,
409 consumer_stream: i64,
410 check: &dyn Fn() -> PyResult<()>,
411) -> PyResult<DlpackManagedTensor> {
412 if consumer_stream <= 0 || consumer_stream == 2 {
413 return Err(PyValueError::new_err(
414 "DLPack requires an explicit supported consumer stream",
415 ));
416 }
417 check()?;
418 let py = obj.py();
419
420 let capsule_obj: Bound<'_, PyAny> = if unsafe {
422 pyo3::ffi::PyCapsule_IsValid(obj.as_ptr(), DLPACK_CAPSULE_NAME.as_ptr() as *const c_char)
423 } != 0
424 {
425 obj.clone()
426 } else if guarded_python_callback(check, || obj.hasattr("__dlpack__"))? {
427 let method = guarded_python_callback(check, || obj.getattr("__dlpack_device__"))?;
428 let device = guarded_python_callback(check, || method.call0())?;
429 dlpack_device_pair(&device)?;
430 dlpack_export_for_stream_guarded(obj, Some(consumer_stream), check)?.into_bound(py)
435 } else {
436 return Err(PyValueError::new_err(
437 "Expected a DLPack capsule or an object with __dlpack__",
438 ));
439 };
440 check()?;
441
442 if unsafe {
444 pyo3::ffi::PyCapsule_IsValid(
445 capsule_obj.as_ptr(),
446 DLPACK_CAPSULE_NAME.as_ptr() as *const c_char,
447 )
448 } == 0
449 {
450 return Err(PyValueError::new_err("Invalid DLPack capsule"));
451 }
452
453 let ptr = unsafe {
455 pyo3::ffi::PyCapsule_GetPointer(
456 capsule_obj.as_ptr(),
457 DLPACK_CAPSULE_NAME.as_ptr() as *const c_char,
458 )
459 };
460 if ptr.is_null() {
461 return Err(PyRuntimeError::new_err("Failed to get DLPack pointer"));
462 }
463
464 let rc = unsafe {
466 pyo3::ffi::PyCapsule_SetName(
467 capsule_obj.as_ptr(),
468 USED_DLPACK_CAPSULE_NAME.as_ptr() as *const c_char,
469 )
470 };
471 if rc != 0 {
472 return Err(PyRuntimeError::new_err(
473 "Failed to mark DLPack capsule as consumed",
474 ));
475 }
476
477 Ok(unsafe { DlpackManagedTensor::from_raw(ptr as *mut xlog_cuda::DLManagedTensor) })
479}
480
481#[cfg(test)]
482mod dlpack_guard_tests {
483 use super::*;
484
485 #[test]
486 fn dlpack_device_callback_refusal_stops_all_later_producer_callbacks() {
487 Python::initialize();
488 Python::attach(|py| {
489 let globals = PyDict::new(py);
490 py.run(
491 cr#"
492from enum import IntEnum
493class DeviceKind(IntEnum):
494 CUDA = 2
495refused = False
496calls = []
497class Producer:
498 def __getattribute__(self, name):
499 if refused:
500 calls.append(('late_lookup', name))
501 return object.__getattribute__(self, name)
502 def __dlpack_device__(self):
503 global refused
504 calls.append('device')
505 refused = True
506 return (DeviceKind.CUDA, 0)
507 def __dlpack__(self, *, stream):
508 calls.append('export')
509 raise AssertionError('export ran after authority refusal')
510producer = Producer()
511"#,
512 Some(&globals),
513 None,
514 )
515 .unwrap();
516 let producer = globals.get_item("producer").unwrap().unwrap();
517 let check = || {
518 if globals.get_item("refused")?.unwrap().extract::<bool>()? {
519 Err(PyValueError::new_err(
520 "original producer authority was refused",
521 ))
522 } else {
523 Ok(())
524 }
525 };
526 assert!(super::dlpack_from_py_for_stream(&producer, 19).is_err());
529 py.run(
530 c"assert 'export' in calls\ncalls.clear()\nrefused = False",
531 Some(&globals),
532 None,
533 )
534 .unwrap();
535 let error = match super::dlpack_from_py_for_stream_guarded(&producer, 19, &check) {
536 Ok(_) => panic!("revoked producer handed off a native tensor"),
537 Err(error) => error,
538 };
539 assert!(error
540 .to_string()
541 .contains("original producer authority was refused"));
542 assert_eq!(
543 globals
544 .get_item("calls")
545 .unwrap()
546 .unwrap()
547 .extract::<Vec<String>>()
548 .unwrap(),
549 ["device"]
550 );
551 });
552 }
553
554 #[test]
555 fn dlpack_integer_device_kinds_reach_original_export_and_capsule_validation() {
556 Python::initialize();
557 Python::attach(|py| {
558 let globals = PyDict::new(py);
559 py.run(
560 cr#"
561from enum import IntEnum
562class DeviceKind(IntEnum):
563 CUDA = 2
564 def __int__(self):
565 raise AssertionError('device kind conversion invoked')
566 def __index__(self):
567 raise AssertionError('device kind index invoked')
568class IntegerKind(int):
569 def __int__(self):
570 raise AssertionError('integer subclass conversion invoked')
571 def __index__(self):
572 raise AssertionError('integer subclass index invoked')
573calls = []
574class Producer:
575 def __dlpack_device__(self):
576 calls.append('device')
577 return device
578 def __dlpack__(self, *, stream):
579 assert stream == 19
580 calls.append('export')
581 return None
582producer = Producer()
583devices = [(2, 0), (DeviceKind.CUDA, 0), (IntegerKind(2), 0)]
584"#,
585 Some(&globals),
586 None,
587 )
588 .unwrap();
589 let producer = globals.get_item("producer").unwrap().unwrap();
590 let devices = globals.get_item("devices").unwrap().unwrap();
591 for device in devices.cast::<PyList>().unwrap().iter() {
592 globals.set_item("device", device).unwrap();
593 py.run(c"calls.clear()", Some(&globals), None).unwrap();
594 let error = super::dlpack_from_py_for_stream_guarded(&producer, 19, &|| Ok(()))
595 .err()
596 .expect("invalid capsule must not be consumed");
597 assert!(error.to_string().contains("Invalid DLPack capsule"));
598 assert_eq!(
599 globals
600 .get_item("calls")
601 .unwrap()
602 .unwrap()
603 .extract::<Vec<String>>()
604 .unwrap(),
605 ["device", "export"]
606 );
607 }
608 });
609 }
610
611 #[test]
612 fn dlpack_device_pair_rejects_invalid_fields_before_export() {
613 Python::initialize();
614 Python::attach(|py| {
615 let globals = PyDict::new(py);
616 py.run(
617 cr#"
618class Pair(tuple):
619 def __iter__(self):
620 raise AssertionError('tuple subclass iterated')
621 def __getitem__(self, index):
622 raise AssertionError('tuple subclass indexed')
623class IntegerKind(int):
624 def __int__(self):
625 raise AssertionError('kind converted')
626 def __index__(self):
627 raise AssertionError('kind index invoked')
628calls = []
629class Producer:
630 def __dlpack_device__(self):
631 calls.append('device')
632 return device
633 def __dlpack__(self, *, stream):
634 calls.append('export')
635 return None
636producer = Producer()
637devices = [None, [2, 0], Pair((2, 0)), (), (2,), (2, 0, 0),
638 (True, 0), (2, False), (2.0, 0), (2, 0.0), ('2', 0),
639 (1, 0), (-1, 0), (2, -1), (2, 1 << 31), (1 << 31, 0),
640 (IntegerKind(1), 0), (IntegerKind(-1), 0),
641 (IntegerKind(1 << 80), 0), (2, IntegerKind(0))]
642"#,
643 Some(&globals),
644 None,
645 )
646 .unwrap();
647 let producer = globals.get_item("producer").unwrap().unwrap();
648 let devices = globals.get_item("devices").unwrap().unwrap();
649 for device in devices.cast::<PyList>().unwrap().iter() {
650 globals.set_item("device", device).unwrap();
651 py.run(c"calls.clear()", Some(&globals), None).unwrap();
652 assert!(
653 super::dlpack_from_py_for_stream_guarded(&producer, 19, &|| Ok(())).is_err()
654 );
655 assert_eq!(
656 globals
657 .get_item("calls")
658 .unwrap()
659 .unwrap()
660 .extract::<Vec<String>>()
661 .unwrap(),
662 ["device"]
663 );
664 }
665 });
666 }
667
668 #[test]
669 fn dlpack_device_metadata_never_invokes_custom_integer_conversion() {
670 Python::initialize();
671 Python::attach(|py| {
672 let globals = PyDict::new(py);
673 py.run(
674 cr#"
675calls = []
676class DeviceIndex:
677 def __int__(self):
678 calls.append('int')
679 return 2
680 def __index__(self):
681 calls.append('index')
682 return 2
683class Producer:
684 def __dlpack_device__(self):
685 calls.append('device')
686 return device
687 def __dlpack__(self, *, stream):
688 calls.append('export')
689 raise AssertionError('export ran after invalid device metadata')
690producer = Producer()
691devices = [(2, DeviceIndex()), (DeviceIndex(), 0)]
692"#,
693 Some(&globals),
694 None,
695 )
696 .unwrap();
697 let producer = globals.get_item("producer").unwrap().unwrap();
698 let devices = globals.get_item("devices").unwrap().unwrap();
699 for device in devices.cast::<PyList>().unwrap().iter() {
700 globals.set_item("device", device).unwrap();
701 py.run(c"calls.clear()", Some(&globals), None).unwrap();
702 let error = super::dlpack_from_py_for_stream_guarded(&producer, 19, &|| Ok(()))
703 .err()
704 .expect("custom device conversion must not hand off a native tensor");
705 assert!(error
706 .to_string()
707 .contains("exact tuple (integer kind, exact integer id)"));
708 assert_eq!(
709 globals
710 .get_item("calls")
711 .unwrap()
712 .unwrap()
713 .extract::<Vec<String>>()
714 .unwrap(),
715 ["device"]
716 );
717 }
718 });
719 }
720}
721
722#[pyfunction]
723fn dlpack_is_cuda(obj: &Bound<'_, PyAny>) -> PyResult<bool> {
724 if unsafe {
727 pyo3::ffi::PyCapsule_IsValid(obj.as_ptr(), DLPACK_CAPSULE_NAME.as_ptr() as *const c_char)
728 } == 0
729 {
730 return Err(PyValueError::new_err(
731 "Expected a DLPack capsule (dltensor)",
732 ));
733 }
734
735 let ptr = unsafe {
737 pyo3::ffi::PyCapsule_GetPointer(obj.as_ptr(), DLPACK_CAPSULE_NAME.as_ptr() as *const c_char)
738 };
739 if ptr.is_null() {
740 return Err(PyRuntimeError::new_err("Failed to get DLPack pointer"));
741 }
742
743 let managed = unsafe { &*(ptr as *const xlog_cuda::DLManagedTensor) };
745 Ok(managed.dl_tensor.device.device_type == xlog_cuda::dlpack::K_DLCUDA)
746}
747
748#[pyfunction]
749fn intern_symbols(symbols: Vec<String>) -> Vec<u32> {
750 symbols
751 .iter()
752 .map(|symbol| xlog_core::symbol::intern(symbol))
753 .collect()
754}
755
756#[pyfunction]
757fn resolve_symbols(symbol_ids: Vec<u32>) -> PyResult<Vec<String>> {
758 symbol_ids
759 .into_iter()
760 .enumerate()
761 .map(|(index, symbol_id)| {
762 xlog_core::symbol::resolve_checked(symbol_id).ok_or_else(|| {
763 PyValueError::new_err(format!("unknown symbol ID {symbol_id} at index {index}"))
764 })
765 })
766 .collect()
767}
768
769#[pyclass(name = "DifferentiableProofTraceMap")]
770pub struct PyDifferentiableProofTraceMap {
771 inner: xlog_logic::DifferentiableProofTraceMap,
772}
773
774fn pack_differentiable_proof_trace(
775 py: Python<'_>,
776 trace: &xlog_logic::ProofTrace,
777) -> PyResult<Py<PyAny>> {
778 let dict = PyDict::new(py);
779 dict.set_item("proof_id", trace.proof_id)?;
780 dict.set_item("answer_key", &trace.answer_key)?;
781 dict.set_item("clause_id", &trace.clause_id)?;
782 dict.set_item("support_atoms", &trace.support_atoms)?;
783 dict.set_item("weight", trace.weight)?;
784 dict.set_item("gradient", trace.gradient)?;
785 Ok(dict.into())
786}
787
788#[pymethods]
789impl PyDifferentiableProofTraceMap {
790 #[new]
791 fn new() -> Self {
792 Self {
793 inner: xlog_logic::DifferentiableProofTraceMap::new(),
794 }
795 }
796
797 fn insert(
798 &mut self,
799 answer_key: String,
800 clause_id: String,
801 support_atoms: Vec<String>,
802 initial_weight: f64,
803 ) -> PyResult<u64> {
804 if !initial_weight.is_finite() {
805 return Err(PyValueError::new_err(
806 "initial_weight must be a finite float",
807 ));
808 }
809 Ok(self.inner.insert(xlog_logic::ProofTraceSpec {
810 answer_key,
811 clause_id,
812 support_atoms,
813 initial_weight,
814 }))
815 }
816
817 fn trace(&self, py: Python<'_>, proof_id: u64) -> PyResult<Option<Py<PyAny>>> {
818 self.inner
819 .trace(proof_id)
820 .map(|trace| pack_differentiable_proof_trace(py, trace))
821 .transpose()
822 }
823
824 fn traces(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
825 let list = PyList::empty(py);
826 for trace in self.inner.traces() {
827 list.append(pack_differentiable_proof_trace(py, trace)?)?;
828 }
829 Ok(list.into())
830 }
831
832 fn accumulate_binary_logistic_gradients(
833 &mut self,
834 targets: Vec<(String, f64)>,
835 ) -> PyResult<f64> {
836 if targets.iter().any(|(_, target)| !target.is_finite()) {
837 return Err(PyValueError::new_err("targets must be finite floats"));
838 }
839 Ok(self.inner.accumulate_binary_logistic_gradients(&targets))
840 }
841
842 fn apply_gradients(&mut self, learning_rate: f64) -> PyResult<()> {
843 if !learning_rate.is_finite() || learning_rate < 0.0 {
844 return Err(PyValueError::new_err(
845 "learning_rate must be a finite non-negative float",
846 ));
847 }
848 self.inner.apply_gradients(learning_rate);
849 Ok(())
850 }
851}
852
853#[pyclass]
854pub struct Program;
855
856#[pyclass]
857pub struct CompiledProgram {
858 pub(crate) program: CompiledProbProgram,
859 pub(crate) output_provider: Arc<CudaKernelProvider>,
860 pub(crate) network_registry: NetworkRegistry,
862 pub(crate) neural_registry: NeuralPredicateRegistry,
864 pub(crate) declared_networks: HashSet<String>,
866 pub(crate) declared_network_forms: HashMap<String, bool>,
868 pub(crate) tensor_sources: TensorSourceRegistry,
870 pub(crate) domain_source: Option<String>,
875 pub(crate) domain_ids: Vec<i64>,
898 pub(crate) _source: String,
900 pub(crate) ast: xlog_logic::ast::Program,
902 pub(crate) _gpu_config: GpuConfig,
904 pub(crate) _prob_engine: ProbEngine,
906 pub(crate) query_signature_cache: HashMap<String, QuerySignature>,
908 pub(crate) circuit_cache: HashMap<String, CachedCircuit>,
910 pub(crate) circuit_cache_hits: usize,
912 pub(crate) circuit_cache_misses: usize,
914 pub(crate) template_compile_count: usize,
916 pub(crate) batch_queries: bool,
918 pub(crate) last_compile_profile: Option<xlog_prob::compilation::CircuitCompileProfile>,
920}
921
922#[pyclass]
923pub struct LogicProgram;
924
925#[pyclass]
926pub struct CompiledLogicProgram {
927 pub(crate) program: Arc<gpu_logic::LogicProgram>,
928 pub(crate) provider: Arc<CudaKernelProvider>,
929}
930
931#[pyclass(skip_from_py_object)]
936#[derive(Clone)]
937pub struct CompiledConditionedProgram {
938 #[cfg(feature = "host-io")]
939 pub(crate) program: xlog_prob::epistemic_production::PreparedConditionedProgram,
940 #[cfg(feature = "host-io")]
941 pub(crate) result_provider: Arc<CudaKernelProvider>,
942}
943
944#[pyclass]
945pub struct LogicRelationSession {
946 pub(crate) program: Arc<gpu_logic::LogicProgram>,
947 pub(crate) provider: Arc<CudaKernelProvider>,
948 pub(crate) relation_store: RelationStore,
949 pub(crate) evaluation_store: Option<gpu_logic::LogicMaterializedStore>,
950 pub(crate) session_runtime: Option<gpu_logic::LogicSessionRuntime>,
951 pub(crate) last_delta_stats: Option<LogicDeltaStats>,
952 pub(crate) relation_callbacks: Vec<RelationChangeCallback>,
953 pub(crate) next_relation_callback_id: u64,
954 pub(crate) relation_generations: HashMap<String, u64>,
955 pub(crate) relation_metadata: RelationMetadataStore,
956}
957
958pub(crate) struct RelationChangeCallback {
959 pub id: u64,
960 pub callback: Py<PyAny>,
961}
962
963#[derive(Clone, Debug)]
964pub(crate) struct LogicDeltaStats {
965 pub input_delta_count: usize,
966 pub changed_relations: usize,
967 pub changed_relation_names: Vec<String>,
968 pub insert_rows: u64,
969 pub delete_rows: u64,
970 pub has_deletes: bool,
971 pub affected_sccs: usize,
972 pub recomputed_sccs: usize,
973 pub incremental_sccs: usize,
974 pub coalesced_insert_rows: u64,
975 pub coalesced_delete_rows: u64,
976 pub canceled_rows: u64,
977 pub equivalent_to_full_recompute: Option<bool>,
978 pub planner_telemetry: gpu_logic::DeltaPlannerTelemetry,
979 pub debug_trace: Vec<String>,
980}
981
982#[pyclass]
983pub struct LogicQueryResult {
984 #[pyo3(get)]
985 pub relation_name: String,
986 #[pyo3(get)]
987 pub columns: Vec<String>,
988 #[pyo3(get)]
989 pub sort_labels: Vec<String>,
990 #[pyo3(get)]
991 pub tensors: Vec<Py<PyAny>>,
992 #[pyo3(get)]
993 pub num_rows: usize,
994 #[pyo3(get)]
995 pub is_true: bool,
996}
997
998#[pyclass]
999pub struct LogicEvalResult {
1000 #[pyo3(get)]
1001 pub queries: Vec<Py<LogicQueryResult>>,
1002}
1003
1004#[pyclass]
1005pub struct IlpTaggedCreditDeviceResult {
1006 #[pyo3(get)]
1007 pub fact_row_offsets: Py<PyAny>,
1008 #[pyo3(get)]
1009 pub entry_indices: Py<PyAny>,
1010 #[pyo3(get)]
1011 pub entry_i: Py<PyAny>,
1012 #[pyo3(get)]
1013 pub entry_j: Py<PyAny>,
1014 #[pyo3(get)]
1015 pub entry_k: Py<PyAny>,
1016}
1017
1018#[pyclass]
1019pub struct McDeviceEvalResult {
1020 #[pyo3(get)]
1022 pub query_counts: Py<PyAny>,
1023 #[pyo3(get)]
1025 pub evidence_count: Py<PyAny>,
1026 #[pyo3(get)]
1027 pub total_samples: usize,
1028 #[pyo3(get)]
1029 pub seed: u64,
1030 #[pyo3(get)]
1031 pub confidence: f64,
1032 #[pyo3(get)]
1033 pub nonmonotone_semantics: String,
1034 #[pyo3(get)]
1035 pub nonmonotone_sccs: usize,
1036 #[pyo3(get)]
1037 pub nonmonotone_cycles: usize,
1038 #[pyo3(get)]
1039 pub nonmonotone_iteration_limit_hits: usize,
1040 #[pyo3(get)]
1041 pub sampling_method: String,
1042 #[pyo3(get)]
1043 pub resident_no_host_certified: bool,
1044 #[pyo3(get)]
1045 pub resident_no_host_policy_result: String,
1046 #[pyo3(get)]
1047 pub resident_no_host_tracked_dtoh_calls: u64,
1048 #[pyo3(get)]
1049 pub resident_no_host_tracked_htod_calls: u64,
1050 #[pyo3(get)]
1051 pub resident_no_host_host_loop_iterations: u64,
1052 #[pyo3(get)]
1053 pub resident_no_host_per_sample_host_launches: u64,
1054 #[pyo3(get)]
1055 pub resident_no_host_untracked_metadata_reads: u64,
1056 #[pyo3(get)]
1057 pub resident_no_host_engine_launches: u64,
1058 #[pyo3(get)]
1059 pub resident_no_host_host_fixpoint_iterations: u64,
1060 #[pyo3(get)]
1061 pub resident_no_host_per_operator_host_allocations: u64,
1062}
1063
1064#[pyclass]
1065pub struct EvalResult {
1066 #[pyo3(get)]
1067 pub atoms: Vec<String>,
1068 #[pyo3(get)]
1069 pub prob: Py<PyAny>,
1070 #[pyo3(get)]
1071 pub log_prob: Py<PyAny>,
1072 #[pyo3(get)]
1073 pub num_vars: usize,
1074 #[pyo3(get)]
1076 pub log_z_e: Option<f64>,
1077 #[pyo3(get)]
1078 pub grad_true: Option<Vec<Py<PyAny>>>,
1079 #[pyo3(get)]
1080 pub grad_false: Option<Vec<Py<PyAny>>>,
1081 #[pyo3(get)]
1082 pub approx: bool,
1083 #[pyo3(get)]
1084 pub stderr: Option<Py<PyAny>>,
1085 #[pyo3(get)]
1086 pub ci_low: Option<Py<PyAny>>,
1087 #[pyo3(get)]
1088 pub ci_high: Option<Py<PyAny>>,
1089 #[pyo3(get)]
1090 pub samples: Option<usize>,
1091 #[pyo3(get)]
1092 pub evidence_samples: Option<usize>,
1093 #[pyo3(get)]
1094 pub seed: Option<u64>,
1095 #[pyo3(get)]
1096 pub confidence: Option<f64>,
1097 #[pyo3(get)]
1098 pub nonmonotone_semantics: Option<String>,
1099 #[pyo3(get)]
1100 pub nonmonotone_sccs: Option<usize>,
1101 #[pyo3(get)]
1102 pub nonmonotone_cycles: Option<usize>,
1103 #[pyo3(get)]
1104 pub nonmonotone_iteration_limit_hits: Option<usize>,
1105 #[pyo3(get)]
1106 pub sampling_method: Option<String>,
1107 #[pyo3(get)]
1111 pub mc_engine: Option<String>,
1112}
1113
1114#[pyclass]
1128pub struct EpistemicEvalResult {
1129 #[pyo3(get)]
1130 pub atoms: Vec<String>,
1131 #[pyo3(get)]
1132 pub prob: Py<PyAny>,
1133 #[pyo3(get)]
1134 pub log_prob: Py<PyAny>,
1135 #[pyo3(get)]
1136 pub log_z_e: f64,
1137 #[pyo3(get)]
1138 pub trace: Py<PyAny>,
1139}
1140
1141#[pyclass]
1149pub struct EpistemicEvidence {
1150 #[pyo3(get)]
1151 pub epistemic_mode: String,
1152 #[pyo3(get)]
1153 pub know_operator_count: usize,
1154 #[pyo3(get)]
1155 pub possible_operator_count: usize,
1156 #[pyo3(get)]
1157 pub accepted_candidates: usize,
1158 #[pyo3(get)]
1159 pub rejected_candidates: usize,
1160 #[pyo3(get)]
1161 pub accepted_world_views: usize,
1162 #[pyo3(get)]
1163 pub final_output_rows: usize,
1164}
1165
1166#[pyclass(skip_from_py_object)]
1174#[derive(Clone)]
1175pub struct EpochStats {
1176 #[pyo3(get)]
1178 pub avg_loss: f64,
1179 #[pyo3(get)]
1181 pub num_batches: usize,
1182 #[pyo3(get)]
1184 pub total_queries: usize,
1185}
1186
1187#[pyclass(skip_from_py_object)]
1191#[derive(Clone)]
1192pub struct TrainingHistory {
1193 #[pyo3(get)]
1195 pub epoch_losses: Vec<f64>,
1196 #[pyo3(get)]
1198 pub epoch_times: Vec<f64>,
1199 #[pyo3(get)]
1201 pub batch_losses: Vec<f64>,
1202 #[pyo3(get)]
1204 pub stopped_early: bool,
1205}
1206
1207#[pyclass]
1208pub struct IlpProgramFactory;
1209
1210#[pyclass]
1211pub struct CompiledIlpProgram {
1212 pub(crate) base_source: String,
1213 pub(crate) _learnable_source: String,
1214 pub(crate) ast: AstProgram,
1215 pub(crate) executor: Executor,
1216 pub(crate) provider: Arc<CudaKernelProvider>,
1217 pub(crate) plan: ExecutionPlan,
1218 pub(crate) rel_index: Vec<(RelId, String)>,
1219 pub(crate) schemas: HashMap<String, Schema>,
1220 pub(crate) left_keys: Vec<usize>,
1221 pub(crate) right_keys: Vec<usize>,
1222 pub(crate) head_projection: Vec<usize>,
1223 pub(crate) compiled_schema_size: usize,
1224 pub(crate) head_rel_name: String,
1225 pub(crate) max_active_rules: usize,
1226 pub(crate) candidate_map: Option<HashMap<(u32, u32, u32), u32>>,
1227 pub(crate) candidate_order: Option<Vec<(u32, u32, u32)>>,
1228 pub(crate) relation_overrides: HashMap<String, CudaBuffer>,
1229 pub(crate) coo_chunk_budget: u64,
1233 pub(crate) strict_zero_dtoh: bool,
1236 pub(crate) compile_timing: Vec<(&'static str, f64)>,
1238}
1239
1240#[pymodule]
1241#[pyo3(name = "_native")]
1242fn pyxlog(_py: Python<'_>, m: &Bound<'_, PyModule>) -> PyResult<()> {
1243 m.add("__version__", env!("CARGO_PKG_VERSION"))?;
1244 m.add_class::<Program>()?;
1245 m.add_class::<CompiledProgram>()?;
1246 m.add_class::<LogicProgram>()?;
1247 m.add_class::<CompiledLogicProgram>()?;
1248 m.add_class::<CompiledConditionedProgram>()?;
1249 m.add_class::<LogicRelationSession>()?;
1250 m.add_class::<semantic_transition::PySemanticTransitionSession>()?;
1251 m.add_class::<semantic_transition::PySemanticTransitionController>()?;
1252 m.add_class::<semantic_transition::cold_task::PySemanticTransitionColdTask>()?;
1253 m.add_class::<semantic_transition::cold_task::PySemanticTransitionFreshParent>()?;
1254 m.add_class::<semantic_transition::PySemanticTransitionTaskUse>()?;
1255 m.add_class::<semantic_transition::PySemanticInitialPrefillStage>()?;
1256 m.add_class::<semantic_transition::PySemanticInitialPrefillContentWitness>()?;
1257 m.add_class::<semantic_transition::PySemanticInitialPrefillReceipt>()?;
1258 m.add_class::<semantic_transition::PySemanticPublishedParent>()?;
1259 m.add_class::<semantic_transition::PySemanticTransitionRestoredCheckpoint>()?;
1260 m.add_class::<semantic_transition::learning_phase::PySemanticLearningPhaseRecipe>()?;
1261 m.add_class::<semantic_transition::learning_phase::PySemanticLearningPhaseTransition>()?;
1262 m.add_class::<semantic_transition::PyNativeTensorAllocation>()?;
1263 m.add_class::<semantic_transition::PySemanticPreparedStep>()?;
1264 #[cfg(feature = "semantic-policy")]
1265 {
1266 m.add(
1267 "SemanticCompletedSegmentPending",
1268 m.py()
1269 .get_type::<semantic_transition::SemanticCompletedSegmentPending>(),
1270 )?;
1271 m.add(
1272 "SemanticModelEvaluationPending",
1273 m.py()
1274 .get_type::<semantic_transition::model_evaluation::SemanticModelEvaluationPending>(
1275 ),
1276 )?;
1277 m.add_class::<semantic_transition::model_evaluation::PySemanticEvaluationCohort>()?;
1278 m.add_class::<semantic_transition::model_evaluation::PySemanticModelEvaluation>()?;
1279 m.add_class::<semantic_transition::PySemanticCompletedExecutionObservation>()?;
1280 m.add_class::<semantic_transition::model_evaluation::PySemanticCompletedModelEvaluation>()?;
1281 }
1282 #[cfg(feature = "semantic-policy")]
1283 m.add_class::<semantic_transition::PySemanticCompletedModelCarrier>()?;
1284 #[cfg(feature = "semantic-policy")]
1285 m.add_class::<semantic_transition::PySemanticCompletedMaterial>()?;
1286 #[cfg(feature = "semantic-policy")]
1287 m.add_class::<semantic_transition::PySemanticCompletedActionComponent>()?;
1288 #[cfg(feature = "semantic-policy")]
1289 m.add_class::<semantic_transition::PySemanticCompletedTextAction>()?;
1290 #[cfg(feature = "semantic-policy")]
1291 m.add_class::<semantic_transition::PySemanticCompletedDecodedEdit>()?;
1292 #[cfg(feature = "semantic-policy")]
1293 m.add_class::<semantic_transition::PySemanticCompletedActionLane>()?;
1294 #[cfg(feature = "semantic-policy")]
1295 m.add_class::<semantic_transition::PySemanticCompletedTheoryDelta>()?;
1296 #[cfg(feature = "semantic-policy")]
1297 m.add_class::<semantic_transition::PySemanticCompletedTaskFacts>()?;
1298 #[cfg(feature = "semantic-policy")]
1299 m.add_class::<semantic_transition::PySemanticCompletedLaneOutcome>()?;
1300 #[cfg(feature = "semantic-policy")]
1301 m.add_class::<semantic_transition::PySemanticCompletedTaskGround>()?;
1302 #[cfg(feature = "semantic-policy")]
1303 m.add_class::<semantic_transition::PySemanticCompletedEditSolution>()?;
1304 #[cfg(feature = "semantic-policy")]
1305 m.add_class::<semantic_transition::PySemanticCompletedActionProjection>()?;
1306 m.add_class::<semantic_transition::PySemanticGradientDelivery>()?;
1307 m.add_class::<semantic_transition::PySemanticTensorContentWitness>()?;
1308 m.add_class::<semantic_transition::PySemanticModelForwardWitness>()?;
1309 #[cfg(feature = "semantic-policy")]
1310 m.add_class::<semantic_transition::PySemanticPolicyInvocation>()?;
1311 #[cfg(feature = "semantic-policy")]
1312 m.add_class::<semantic_transition::PySemanticRetainedReplayMember>()?;
1313 m.add_class::<relation_metadata::RelationEvidence>()?;
1314 m.add_class::<LogicQueryResult>()?;
1315 m.add_class::<LogicEvalResult>()?;
1316 m.add_class::<McDeviceEvalResult>()?;
1317 m.add_class::<EvalResult>()?;
1318 m.add_class::<EpistemicEvalResult>()?;
1319 m.add_class::<EpistemicEvidence>()?;
1320 m.add_class::<PyDifferentiableProofTraceMap>()?;
1322 m.add_class::<EpochStats>()?;
1323 m.add_class::<TrainingHistory>()?;
1324 m.add_class::<IlpProgramFactory>()?;
1326 m.add_class::<CompiledIlpProgram>()?;
1327 m.add_class::<IlpTaggedCreditDeviceResult>()?;
1328 m.add_function(wrap_pyfunction!(training::train_model, m)?)?;
1329 m.add_function(wrap_pyfunction!(training::train_model_tensor, m)?)?;
1330 m.add_function(wrap_pyfunction!(dlpack::dlpack_roundtrip, m)?)?;
1331 m.add_function(wrap_pyfunction!(dlpack_is_cuda, m)?)?;
1332 m.add_function(wrap_pyfunction!(intern_symbols, m)?)?;
1333 m.add_function(wrap_pyfunction!(resolve_symbols, m)?)?;
1334 #[cfg(feature = "arrow-device-import")]
1335 m.add_function(wrap_pyfunction!(dlpack::export_arrow_device, m)?)?;
1336 #[cfg(feature = "arrow-device-import")]
1337 m.add_function(wrap_pyfunction!(dlpack::import_arrow_device, m)?)?;
1338 m.add_class::<joint_carrier::JointConstraintCarrier>()?;
1340 m.add(
1341 "CarrierRefused",
1342 _py.get_type::<joint_carrier::CarrierRefused>(),
1343 )?;
1344 m.add(
1345 "SolverResourceExhausted",
1346 _py.get_type::<joint_carrier::SolverResourceExhausted>(),
1347 )?;
1348 m.add(
1349 "RelationMetadataError",
1350 _py.get_type::<relation_metadata::RelationMetadataError>(),
1351 )?;
1352 m.add("SOLVER_ABI_IDENTITY", xlog_cuda::SOLVER_ABI_IDENTITY)?;
1353 Ok(())
1354}