1use std::collections::HashMap as StdHashMap;
6use std::collections::{HashMap, HashSet};
7use std::sync::Arc;
8use std::time::Instant;
9
10use pyo3::exceptions::{PyRuntimeError, PyValueError};
11use pyo3::prelude::*;
12use pyo3::types::{PyDict, PySequence};
13
14use xlog_core::{RelId, ScalarType, Schema};
15use xlog_cuda::{CudaKernelProvider, JoinType};
16use xlog_ir::{ExecutionPlan, RirNode};
17use xlog_logic::ast::{Program as AstProgram, Term, TypeRef};
18use xlog_logic::ground_term_encoding::append_ground_term_bytes;
19use xlog_prob::exact::GpuConfig;
20use xlog_runtime::ilp_registry::IlpMask;
21use xlog_runtime::{read_device_row_count, Executor};
22
23use xlog_cuda::type_seam::GpuScalar;
24
25use super::{
26 dlpack_capsule_from_tensor, dlpack_from_py, ilp_gpu, provider_from_config, types,
27 CompiledIlpProgram, IlpProgramFactory, IlpTaggedCreditDeviceResult,
28};
29
30struct RelationExampleGroup {
35 relation: String,
36 query_buf: xlog_cuda::CudaBuffer,
37 num_rows: u32,
38}
39
40fn type_ref_name(typ: &TypeRef) -> String {
41 match typ {
42 TypeRef::Scalar(scalar) => types::scalar_type_name(scalar),
43 TypeRef::Domain(name) => name.clone(),
44 TypeRef::List(inner) => format!("list<{}>", type_ref_name(inner)),
45 TypeRef::Term => "term".to_string(),
46 TypeRef::Compound => "compound".to_string(),
47 TypeRef::PredRef => "predref".to_string(),
48 }
49}
50
51fn collect_dlpack_columns(
52 dlpack_columns: &Bound<'_, PyAny>,
53 err_msg: &str,
54) -> PyResult<Vec<xlog_cuda::DlpackManagedTensor>> {
55 let seq = dlpack_columns
56 .cast::<PySequence>()
57 .map_err(|_| PyValueError::new_err(err_msg.to_string()))?;
58 let mut tensors = Vec::with_capacity(seq.len()?);
59 for item in seq.try_iter()? {
60 tensors.push(dlpack_from_py(&item?)?);
61 }
62 Ok(tensors)
63}
64
65fn build_zero_typed<T>(
68 provider: &CudaKernelProvider,
69 py: Python<'_>,
70 num_cands: u32,
71 scalar_type: ScalarType,
72) -> PyResult<(Py<PyAny>, Py<PyAny>)>
73where
74 T: GpuScalar + cudarc::driver::ValidAsZeroBits,
75{
76 let mut d_grad = provider
77 .memory()
78 .alloc::<T>(num_cands as usize)
79 .map_err(|e| types::gpu_err("alloc grad", e))?;
80 if num_cands > 0 {
81 provider
82 .device()
83 .inner()
84 .memset_zeros(&mut d_grad)
85 .map_err(|e| types::gpu_err("zero grad", e))?;
86 }
87 let mut d_loss = provider
88 .memory()
89 .alloc::<T>(1)
90 .map_err(|e| types::gpu_err("alloc loss", e))?;
91 provider
92 .device()
93 .inner()
94 .memset_zeros(&mut d_loss)
95 .map_err(|e| types::gpu_err("zero loss", e))?;
96 ilp_gpu::export_loss_grad_device(provider, py, d_loss, d_grad, num_cands, scalar_type)
97}
98
99fn export_device_bool_tensor(
100 provider: &CudaKernelProvider,
101 py: Python<'_>,
102 values: xlog_cuda::memory::TrackedCudaSlice<u8>,
103 rows: usize,
104) -> PyResult<Py<PyAny>> {
105 let rows_u32 = u32::try_from(rows)
106 .map_err(|_| PyValueError::new_err(format!("Row count {} exceeds u32::MAX", rows)))?;
107
108 let mut d_num_rows = provider.memory().alloc::<u32>(1).map_err(types::xlog_err)?;
109 provider
110 .device()
111 .inner()
112 .htod_sync_copy_into(&[rows_u32], &mut d_num_rows)
113 .map_err(types::xlog_err)?;
114
115 let buffer = xlog_cuda::CudaBuffer::from_columns(
116 vec![values.into_bytes().into()],
117 rows as u64,
118 d_num_rows,
119 Schema::new(vec![("col0".to_string(), ScalarType::Bool)]),
120 );
121 let tensor = provider
122 .to_dlpack_table(buffer)
123 .column(0)
124 .map_err(types::xlog_err)?;
125 dlpack_capsule_from_tensor(py, tensor)
126}
127
128fn export_device_u32_tensor_as_i32(
129 provider: &CudaKernelProvider,
130 py: Python<'_>,
131 values: xlog_cuda::memory::TrackedCudaSlice<u32>,
132 rows: usize,
133) -> PyResult<Py<PyAny>> {
134 let rows_u32 = u32::try_from(rows)
135 .map_err(|_| PyValueError::new_err(format!("Row count {} exceeds u32::MAX", rows)))?;
136
137 let mut d_num_rows = provider.memory().alloc::<u32>(1).map_err(types::xlog_err)?;
138 provider
139 .device()
140 .inner()
141 .htod_sync_copy_into(&[rows_u32], &mut d_num_rows)
142 .map_err(types::xlog_err)?;
143
144 let buffer = xlog_cuda::CudaBuffer::from_columns(
146 vec![values.into_bytes().into()],
147 rows as u64,
148 d_num_rows,
149 Schema::new(vec![("col0".to_string(), ScalarType::I32)]),
150 );
151 let tensor = provider
152 .to_dlpack_table(buffer)
153 .column(0)
154 .map_err(types::xlog_err)?;
155 dlpack_capsule_from_tensor(py, tensor)
156}
157
158fn empty_tagged_credit_device_result(
159 provider: &CudaKernelProvider,
160 py: Python<'_>,
161 num_facts: usize,
162) -> PyResult<IlpTaggedCreditDeviceResult> {
163 let mut d_row_offsets = provider
164 .memory()
165 .alloc::<u32>(num_facts + 1)
166 .map_err(types::xlog_err)?;
167 provider
168 .device()
169 .inner()
170 .memset_zeros(&mut d_row_offsets)
171 .map_err(types::xlog_err)?;
172 let d_empty_indices = provider.memory().alloc::<u32>(0).map_err(types::xlog_err)?;
173 let d_empty_i = provider.memory().alloc::<u32>(0).map_err(types::xlog_err)?;
174 let d_empty_j = provider.memory().alloc::<u32>(0).map_err(types::xlog_err)?;
175 let d_empty_k = provider.memory().alloc::<u32>(0).map_err(types::xlog_err)?;
176
177 Ok(IlpTaggedCreditDeviceResult {
178 fact_row_offsets: export_device_u32_tensor_as_i32(
179 provider,
180 py,
181 d_row_offsets,
182 num_facts + 1,
183 )?,
184 entry_indices: export_device_u32_tensor_as_i32(provider, py, d_empty_indices, 0)?,
185 entry_i: export_device_u32_tensor_as_i32(provider, py, d_empty_i, 0)?,
186 entry_j: export_device_u32_tensor_as_i32(provider, py, d_empty_j, 0)?,
187 entry_k: export_device_u32_tensor_as_i32(provider, py, d_empty_k, 0)?,
188 })
189}
190
191pub(crate) fn pack_i64_columns_typed(
195 relation: &str,
196 facts: &[Vec<i64>],
197 schema: &Schema,
198) -> PyResult<Vec<Vec<u8>>> {
199 let arity = schema.arity();
200 for (idx, fact) in facts.iter().enumerate() {
201 if fact.len() != arity {
202 return Err(PyValueError::new_err(format!(
203 "Relation '{}': fact {} has {} values, expected {}",
204 relation,
205 idx,
206 fact.len(),
207 arity,
208 )));
209 }
210 }
211
212 let mut columns: Vec<Vec<u8>> = (0..arity)
213 .map(|col_idx| {
214 let elem_size = schema
215 .column_type(col_idx)
216 .map(|t| t.size_bytes())
217 .unwrap_or(4);
218 Vec::with_capacity(facts.len() * elem_size)
219 })
220 .collect();
221
222 for fact in facts {
223 for (col_idx, &val) in fact.iter().enumerate() {
224 let col_type = schema.column_type(col_idx);
225 let col = &mut columns[col_idx];
226 match col_type {
227 Some(ScalarType::U32) => {
228 let v = u32::try_from(val).map_err(|_| {
229 PyValueError::new_err(format!(
230 "Relation '{}' column {} (U32): value {} out of range [0, {}]",
231 relation,
232 col_idx,
233 val,
234 u32::MAX,
235 ))
236 })?;
237 col.extend_from_slice(&v.to_le_bytes());
238 }
239 Some(ScalarType::I32) => {
240 let v = i32::try_from(val).map_err(|_| {
241 PyValueError::new_err(format!(
242 "Relation '{}' column {} (I32): value {} out of range [{}, {}]",
243 relation,
244 col_idx,
245 val,
246 i32::MIN,
247 i32::MAX,
248 ))
249 })?;
250 col.extend_from_slice(&v.to_le_bytes());
251 }
252 Some(ScalarType::U64) => {
253 let v = u64::try_from(val).map_err(|_| PyValueError::new_err(format!(
254 "Relation '{}' column {} (U64): value {} is negative; U64 requires non-negative values",
255 relation, col_idx, val,
256 )))?;
257 col.extend_from_slice(&v.to_le_bytes());
258 }
259 Some(ScalarType::I64) => {
260 col.extend_from_slice(&val.to_le_bytes());
261 }
262 Some(ScalarType::Bool) => match val {
263 0 => col.push(0u8),
264 1 => col.push(1u8),
265 _ => {
266 return Err(PyValueError::new_err(format!(
267 "Relation '{}' column {} (Bool): value {} not in {{0, 1}}",
268 relation, col_idx, val,
269 )))
270 }
271 },
272 Some(ScalarType::Symbol) => {
273 let v = u32::try_from(val).map_err(|_| {
274 PyValueError::new_err(format!(
275 "Relation '{}' column {} (Symbol): value {} out of range [0, {}]",
276 relation,
277 col_idx,
278 val,
279 u32::MAX,
280 ))
281 })?;
282 col.extend_from_slice(&v.to_le_bytes());
283 }
284 Some(ScalarType::F32) => {
285 return Err(PyValueError::new_err(format!(
286 "Relation '{}' column {} (F32): float columns not supported in batch APIs",
287 relation, col_idx,
288 )));
289 }
290 Some(ScalarType::F64) => {
291 return Err(PyValueError::new_err(format!(
292 "Relation '{}' column {} (F64): float columns not supported in batch APIs",
293 relation, col_idx,
294 )));
295 }
296 None => {
297 return Err(PyValueError::new_err(format!(
298 "Relation '{}' column {}: no type in schema",
299 relation, col_idx,
300 )));
301 }
302 }
303 }
304 }
305
306 Ok(columns)
307}
308
309pub(crate) fn load_facts_into_store(
310 ast: &AstProgram,
311 provider: &CudaKernelProvider,
312 executor: &mut Executor,
313 schemas: &HashMap<String, Schema>,
314) -> xlog_core::Result<()> {
315 use xlog_core::XlogError;
316 let mut rows_by_pred: HashMap<&str, Vec<&[Term]>> = HashMap::new();
317 for fact in ast.facts() {
318 rows_by_pred
319 .entry(fact.head.predicate.as_str())
320 .or_default()
321 .push(&fact.head.terms);
322 }
323
324 for (pred, rows) in rows_by_pred {
325 let schema = schemas.get(pred).ok_or_else(|| {
326 XlogError::Execution(format!("Missing schema for fact predicate {}", pred))
327 })?;
328
329 if rows.iter().any(|r| r.len() != schema.arity()) {
330 return Err(XlogError::Execution(format!(
331 "Fact arity mismatch for {} (expected {})",
332 pred,
333 schema.arity()
334 )));
335 }
336
337 let mut columns: Vec<Vec<u8>> = vec![Vec::new(); schema.arity()];
338 for row in &rows {
339 for (col_idx, term) in row.iter().enumerate() {
340 let typ = schema.column_type(col_idx).ok_or_else(|| {
341 XlogError::Execution(format!("Missing type for col {}", col_idx))
342 })?;
343 append_ground_term_bytes(&mut columns[col_idx], term, typ).map_err(|error| {
344 xlog_core::XlogError::Execution(format!(
345 "Failed to encode fact for predicate {pred} at column {col_idx}: {error}"
346 ))
347 })?;
348 }
349 }
350
351 let slices: Vec<&[u8]> = columns.iter().map(|c| c.as_slice()).collect();
352 let fact_buf = provider.create_buffer_from_slices(&slices, schema.clone())?;
353
354 let existing = executor.store().get(pred).ok_or_else(|| {
355 XlogError::Execution(format!(
356 "Missing base relation {} while loading facts",
357 pred
358 ))
359 })?;
360 let merged = provider.union(existing, &fact_buf)?;
361 executor.store_mut().put(pred, merged);
362 }
363 Ok(())
364}
365
366struct TmjMeta {
368 left_keys: Vec<usize>,
369 right_keys: Vec<usize>,
370 head_projection: Vec<usize>,
371 schema_size: usize,
372 head_rel_name: String,
373}
374
375fn walk_tmj(node: &RirNode, target_mask: Option<&str>) -> Option<TmjMeta> {
376 match node {
377 RirNode::TensorMaskedJoin {
378 mask_name,
379 left_keys,
380 right_keys,
381 head_projection,
382 schema_size,
383 head_rel_name,
384 ..
385 } => {
386 if target_mask.is_none() || target_mask == Some(mask_name.as_str()) {
387 Some(TmjMeta {
388 left_keys: left_keys.clone(),
389 right_keys: right_keys.clone(),
390 head_projection: head_projection.clone(),
391 schema_size: *schema_size,
392 head_rel_name: head_rel_name.clone(),
393 })
394 } else {
395 None
396 }
397 }
398 RirNode::Fixpoint {
399 base, recursive, ..
400 } => walk_tmj(base, target_mask).or_else(|| walk_tmj(recursive, target_mask)),
401 RirNode::Union { inputs } => inputs.iter().find_map(|n| walk_tmj(n, target_mask)),
402 RirNode::Filter { input, .. }
403 | RirNode::Project { input, .. }
404 | RirNode::Distinct { input, .. }
405 | RirNode::GroupBy { input, .. } => walk_tmj(input, target_mask),
406 RirNode::Join { left, right, .. } | RirNode::Diff { left, right } => {
407 walk_tmj(left, target_mask).or_else(|| walk_tmj(right, target_mask))
408 }
409 RirNode::MultiWayJoin { fallback, .. } => walk_tmj(fallback, target_mask),
414 RirNode::ChainJoin { fallback, .. } => walk_tmj(fallback, target_mask),
415 _ => None,
416 }
417}
418
419fn extract_tmj_meta(plan: &ExecutionPlan) -> TmjMeta {
420 extract_tmj_meta_for_mask(plan, None)
421}
422
423fn extract_tmj_meta_for_mask(plan: &ExecutionPlan, mask_name: Option<&str>) -> TmjMeta {
424 for scc_rules in &plan.rules_by_scc {
425 for rule in scc_rules {
426 if let Some(meta) = walk_tmj(&rule.body, mask_name) {
427 return meta;
428 }
429 }
430 }
431 TmjMeta {
432 left_keys: vec![],
433 right_keys: vec![],
434 head_projection: vec![],
435 schema_size: 0,
436 head_rel_name: String::new(),
437 }
438}
439
440fn strip_learnable_declarations(source: &str) -> String {
441 source
442 .lines()
443 .filter(|line| !line.trim_start().starts_with("learnable("))
444 .collect::<Vec<_>>()
445 .join("\n")
446}
447
448fn extract_learnable_declarations(source: &str) -> String {
449 source
450 .lines()
451 .filter(|line| line.trim_start().starts_with("learnable("))
452 .collect::<Vec<_>>()
453 .join("\n")
454}
455
456#[pymethods]
461impl IlpProgramFactory {
462 #[staticmethod]
463 #[pyo3(signature = (source, device=0, memory_mb=512, max_active_rules=None))]
464 pub fn compile(
465 source: &str,
466 device: usize,
467 memory_mb: u64,
468 max_active_rules: Option<usize>,
469 ) -> PyResult<CompiledIlpProgram> {
470 if let Some(max) = max_active_rules {
472 if !(16..=128).contains(&max) {
473 return Err(PyValueError::new_err(format!(
474 "max_active_rules must be between 16 and 128, got {}",
475 max
476 )));
477 }
478 }
479
480 build_ilp_program(
483 source,
484 max_active_rules,
485 ProviderFor::New { device, memory_mb },
486 )
487 }
488}
489
490enum ProviderFor {
492 New { device: usize, memory_mb: u64 },
494 Shared(Arc<CudaKernelProvider>),
496}
497
498fn ms_since(t: Instant) -> f64 {
499 t.elapsed().as_secs_f64() * 1000.0
500}
501
502fn build_ilp_program(
505 source: &str,
506 max_active_rules: Option<usize>,
507 provider_for: ProviderFor,
508) -> PyResult<CompiledIlpProgram> {
509 let mut timing: Vec<(&'static str, f64)> = Vec::new();
510 let t = Instant::now();
511 let ast = xlog_logic::parse_program(source).map_err(types::val_err)?;
512
513 let base_source = strip_learnable_declarations(source);
514 let learnable_source = extract_learnable_declarations(source);
515
516 let mut compiler = xlog_logic::Compiler::new();
517 if let Some(max) = max_active_rules {
518 compiler.set_max_active_rules(max);
519 }
520 let plan = compiler.compile_program(&ast).map_err(types::xlog_err)?;
521
522 let mut rel_index: Vec<(RelId, String)> = compiler
523 .rel_ids()
524 .iter()
525 .map(|(name, id)| (*id, name.clone()))
526 .collect();
527 rel_index.sort_by_key(|(id, _)| id.0);
528 let schemas = compiler.schemas().clone();
529 timing.push(("frontend", ms_since(t)));
530
531 let t = Instant::now();
532 let provider = match provider_for {
533 ProviderFor::New { device, memory_mb } => {
534 let mut config = GpuConfig::default();
535 config.device_ordinal = device;
536 config.memory_bytes = memory_mb * 1024 * 1024;
537 let provider = Arc::new(provider_from_config(config).map_err(types::xlog_err)?);
538 timing.push(("provider", ms_since(t)));
539 provider
540 }
541 ProviderFor::Shared(provider) => provider,
542 };
543
544 let t = Instant::now();
545 let mut executor = Executor::new(provider.clone());
546
547 for (name, rel_id) in compiler.rel_ids() {
548 executor.register_relation(*rel_id, name);
549 }
550
551 for (name, schema) in &schemas {
552 let empty = provider
553 .create_empty_buffer(schema.clone())
554 .map_err(types::xlog_err)?;
555 executor.store_mut().put(name, empty);
556 }
557
558 load_facts_into_store(&ast, &provider, &mut executor, &schemas).map_err(types::xlog_err)?;
559 timing.push(("facts", ms_since(t)));
560
561 let t = Instant::now();
562 executor.execute_plan(&plan).map_err(types::xlog_err)?;
563 timing.push(("execute", ms_since(t)));
564
565 let tmj = extract_tmj_meta(&plan);
566
567 let active_rules = max_active_rules.unwrap_or(32);
568
569 Ok(CompiledIlpProgram {
570 base_source,
571 _learnable_source: learnable_source,
572 ast,
573 executor,
574 provider,
575 plan,
576 rel_index,
577 schemas,
578 left_keys: tmj.left_keys,
579 right_keys: tmj.right_keys,
580 head_projection: tmj.head_projection,
581 compiled_schema_size: tmj.schema_size,
582 head_rel_name: tmj.head_rel_name,
583 max_active_rules: active_rules,
584 candidate_map: None,
585 candidate_order: None,
586 relation_overrides: HashMap::new(),
587 coo_chunk_budget: 16 * 1024 * 1024,
588 strict_zero_dtoh: false,
589 compile_timing: timing,
590 })
591}
592
593#[pymethods]
598impl CompiledIlpProgram {
599 pub fn compile_variant(&self, py: Python<'_>, source: &str) -> PyResult<CompiledIlpProgram> {
618 py.detach(|| {
619 build_ilp_program(
620 source,
621 Some(self.max_active_rules),
622 ProviderFor::Shared(Arc::clone(&self.provider)),
623 )
624 })
625 }
626
627 pub fn compile_timing_ms(&self) -> StdHashMap<String, f64> {
632 self.compile_timing
633 .iter()
634 .map(|(k, v)| ((*k).to_string(), *v))
635 .collect()
636 }
637
638 pub fn set_candidate_map(&mut self, candidates: Vec<(u32, u32, u32)>) -> PyResult<()> {
640 let mut map = HashMap::with_capacity(candidates.len());
641 for (cidx, &(i, j, k)) in candidates.iter().enumerate() {
642 map.insert((i, j, k), cidx as u32);
643 }
644 self.candidate_map = Some(map);
645 self.candidate_order = Some(candidates);
646 Ok(())
647 }
648
649 pub fn candidate_map_len(&self) -> usize {
651 self.candidate_map.as_ref().map_or(0, |m| m.len())
652 }
653
654 pub fn debug_ilp_mask_kind(&self, name: String) -> Option<String> {
655 self.executor.ilp_registry().get_mask(&name).map(|mask| {
656 match mask {
657 IlpMask::Dense { .. } => "dense",
658 IlpMask::Sparse { .. } => "sparse_host",
659 IlpMask::SparseDevice { .. } => "sparse_device",
660 }
661 .to_string()
662 })
663 }
664
665 pub fn set_coo_chunk_budget(&mut self, bytes: u64) {
668 self.coo_chunk_budget = bytes;
669 }
670
671 #[allow(deprecated)]
673 pub fn set_coo_memory_cap(&mut self, bytes: u64) {
674 self.coo_chunk_budget = bytes;
675 }
676
677 pub fn set_strict_zero_dtoh(&mut self, strict: bool) {
681 self.strict_zero_dtoh = strict;
682 }
683
684 pub fn put_relation(
690 &mut self,
691 name: String,
692 dlpack_columns: &Bound<'_, PyAny>,
693 ) -> PyResult<()> {
694 if name.starts_with("__") {
695 return Err(PyValueError::new_err(format!(
696 "Relation {} is internal and cannot be stored in a compiled ILP program",
697 name
698 )));
699 }
700 let schema = self.schemas.get(&name).ok_or_else(|| {
701 PyValueError::new_err(format!(
702 "Unknown relation {} (not present in compiled schemas)",
703 name
704 ))
705 })?;
706 let tensors = collect_dlpack_columns(
707 dlpack_columns,
708 &format!("Relation {} must be a sequence of DLPack columns", name),
709 )?;
710 let buffer = self
711 .provider
712 .from_dlpack_tensors_with_schema(schema.clone(), tensors)
713 .map_err(types::xlog_err)?;
714 let live_buffer = self
715 .provider
716 .clone_buffer(&buffer)
717 .map_err(types::xlog_err)?;
718 self.relation_overrides.insert(name.clone(), buffer);
719 self.executor.put_relation(&name, live_buffer);
720 Ok(())
721 }
722
723 pub fn compute_ilp_loss_grad_gpu<'py>(
735 &mut self,
736 py: Python<'py>,
737 positives: Vec<(String, Vec<i64>)>,
738 negatives: Vec<(String, Vec<i64>)>,
739 cand_probs_obj: &Bound<'py, PyAny>,
740 ) -> PyResult<(Py<PyAny>, Py<PyAny>)> {
741 let candidate_map = self.candidate_map.clone().ok_or_else(|| {
743 PyRuntimeError::new_err(
744 "candidate_map not set — call set_candidate_map() before compute_ilp_loss_grad_gpu()",
745 )
746 })?;
747 let num_cands = candidate_map.len() as u32;
748
749 let managed = dlpack_from_py(cand_probs_obj)?;
750 let cand_buf = self
751 .provider
752 .from_dlpack_tensors(vec![managed])
753 .map_err(|e| types::gpu_err("DLPack import", e))?;
754
755 let cand_schema = cand_buf.schema().clone();
757 let cand_dtype = cand_schema
758 .column_type(0)
759 .ok_or_else(|| PyRuntimeError::new_err("cand_probs has no column type"))?;
760 let is_f64 = match cand_dtype {
761 ScalarType::F32 => false,
762 ScalarType::F64 => true,
763 other => {
764 return Err(PyValueError::new_err(format!(
765 "cand_probs must be F32 or F64, got {:?}",
766 other
767 )));
768 }
769 };
770
771 let cand_rows = read_device_row_count(&self.provider, &cand_buf)
772 .map_err(|e| types::gpu_err("row count", e))?;
773 if cand_rows != num_cands as usize {
774 return Err(PyValueError::new_err(format!(
775 "cand_probs length ({}) != candidate_map length ({})",
776 cand_rows, num_cands
777 )));
778 }
779
780 let num_pos = positives.len();
782 let num_neg = negatives.len();
783 let num_facts = (num_pos + num_neg) as u32;
784
785 struct FactInfo {
787 relation: String,
788 values: Vec<i64>,
789 is_positive: bool,
790 }
791 let mut all_facts: Vec<FactInfo> = Vec::with_capacity(num_pos + num_neg);
792 for (rel, vals) in positives {
793 all_facts.push(FactInfo {
794 relation: rel,
795 values: vals,
796 is_positive: true,
797 });
798 }
799 for (rel, vals) in negatives {
800 all_facts.push(FactInfo {
801 relation: rel,
802 values: vals,
803 is_positive: false,
804 });
805 }
806
807 if num_facts == 0 {
809 return self.build_zero_loss_grad(py, num_cands, is_f64);
810 }
811
812 let is_positive_host: Vec<u8> = all_facts
814 .iter()
815 .map(|f| if f.is_positive { 1u8 } else { 0u8 })
816 .collect();
817
818 let mut groups: HashMap<String, Vec<(u32, Vec<i64>)>> = HashMap::new();
820 for (global_idx, fact) in all_facts.iter().enumerate() {
821 groups
822 .entry(fact.relation.clone())
823 .or_default()
824 .push((global_idx as u32, fact.values.clone()));
825 }
826
827 if self.strict_zero_dtoh && self.executor.ilp_registry().has_sparse_device_mask() {
828 self.evaluate_ilp_plan(py)?;
829 }
830
831 let tagged = self
833 .executor
834 .ilp_last_result()
835 .ok_or_else(|| PyRuntimeError::new_err("No ILP result — call evaluate() first"))?;
836
837 let mut tasks: Vec<ilp_gpu::CooTask> = Vec::new();
850 let mut fact_indices_buffers: Vec<xlog_cuda::memory::TrackedCudaSlice<u32>> = Vec::new();
851
852 for (relation, facts_with_idx) in &groups {
853 let k_idx = self
854 .rel_index
855 .iter()
856 .position(|(_, name)| name == relation)
857 .ok_or_else(|| {
858 PyValueError::new_err(format!("Relation '{}' not in ILP schema", relation))
859 })? as u32;
860
861 let relevant_entries: Vec<&xlog_runtime::ilp_registry::IlpTagEntry> = tagged
862 .entries
863 .iter()
864 .filter(|e| e.k == k_idx && e.num_rows > 0 && e.buffer.is_some())
865 .collect();
866
867 if relevant_entries.is_empty() {
868 continue;
869 }
870
871 let first_buf = relevant_entries[0]
872 .buffer
873 .as_ref()
874 .ok_or_else(|| PyRuntimeError::new_err("internal: filtered entry has no buffer"))?;
875 let arity = first_buf.arity();
876 if arity == 0 {
877 continue;
878 }
879 let schema = first_buf.schema().clone();
880
881 let fact_values: Vec<Vec<i64>> =
882 facts_with_idx.iter().map(|(_, v)| v.clone()).collect();
883 let col_bytes = pack_i64_columns_typed(relation, &fact_values, &schema)?;
884 let col_slices: Vec<&[u8]> = col_bytes.iter().map(|c| c.as_slice()).collect();
885 let query_buf = self
886 .provider
887 .create_buffer_from_slices(&col_slices, schema)
888 .map_err(|e| types::gpu_err("create_buffer", e))?;
889
890 let keys: Vec<usize> = (0..arity).collect();
891 let num_query = fact_values.len() as u32;
892
893 let global_indices: Vec<u32> = facts_with_idx.iter().map(|(idx, _)| *idx).collect();
896 let mut d_fact_indices = self
897 .provider
898 .memory()
899 .alloc::<u32>(num_query as usize)
900 .map_err(|e| types::gpu_err("alloc fact_indices", e))?;
901 self.provider
902 .device()
903 .inner()
904 .htod_sync_copy_into(&global_indices, &mut d_fact_indices)
905 .map_err(|e| types::gpu_err("htod fact_indices", e))?;
906 let fi_idx = fact_indices_buffers.len();
907 fact_indices_buffers.push(d_fact_indices);
908
909 for entry in &relevant_entries {
910 let cidx = match candidate_map.get(&(entry.i, entry.j, entry.k)) {
911 Some(c) => *c,
912 None => continue,
913 };
914
915 let entry_buf = entry.buffer.as_ref().ok_or_else(|| {
916 PyRuntimeError::new_err("internal: filtered entry has no buffer")
917 })?;
918 let d_mask = self
919 .provider
920 .membership_mask_device(&query_buf, entry_buf, &keys, &keys)
921 .map_err(|e| types::gpu_err("membership_mask", e))?;
922
923 tasks.push(ilp_gpu::CooTask {
924 cidx,
925 num_query,
926 d_mask,
927 fact_indices_idx: fi_idx,
928 });
929 }
930 }
931
932 let num_tasks = tasks.len();
933 let upper_bound: u32 = tasks.iter().map(|t| t.num_query).sum();
934
935 if upper_bound == 0 || num_tasks == 0 {
936 return self.build_loss_grad_empty_coo(
937 py,
938 &is_positive_host,
939 num_facts,
940 num_cands,
941 is_f64,
942 );
943 }
944
945 let coo_bytes = (upper_bound as u64) * 8;
948 let needs_chunking = coo_bytes > self.coo_chunk_budget;
949
950 let mut d_is_positive = self
952 .provider
953 .memory()
954 .alloc::<u8>(num_facts as usize)
955 .map_err(|e| types::gpu_err("alloc is_positive", e))?;
956 self.provider
957 .device()
958 .inner()
959 .htod_sync_copy_into(&is_positive_host, &mut d_is_positive)
960 .map_err(|e| types::gpu_err("htod is_positive", e))?;
961
962 let cand_col = cand_buf
963 .column(0)
964 .ok_or_else(|| PyRuntimeError::new_err("cand_probs has no column"))?;
965
966 let (mut d_coo_facts, mut d_coo_cands, actual_nnz) = if !needs_chunking {
968 ilp_gpu::build_coo_single(
969 &self.provider,
970 &tasks,
971 &fact_indices_buffers,
972 num_facts,
973 num_cands,
974 upper_bound,
975 )?
976 } else {
977 match ilp_gpu::build_coo_chunked(
978 &self.provider,
979 &tasks,
980 &fact_indices_buffers,
981 num_facts,
982 self.coo_chunk_budget,
983 )? {
984 Some(result) => result,
985 None => {
986 return self.build_loss_grad_empty_coo(
987 py,
988 &is_positive_host,
989 num_facts,
990 num_cands,
991 is_f64,
992 );
993 }
994 }
995 };
996
997 let d_row_offsets = ilp_gpu::sort_and_build_csr(
999 &self.provider,
1000 &mut d_coo_facts,
1001 &mut d_coo_cands,
1002 actual_nnz,
1003 num_facts,
1004 )?;
1005
1006 ilp_gpu::forward_backward_reduce(
1008 &self.provider,
1009 py,
1010 &d_row_offsets,
1011 &d_coo_cands,
1012 cand_col,
1013 &d_is_positive,
1014 num_facts,
1015 num_cands,
1016 is_f64,
1017 )
1018 }
1019
1020 pub fn compute_ilp_loss_grad_gpu_relations<'py>(
1027 &mut self,
1028 py: Python<'py>,
1029 positives_by_relation: &Bound<'py, PyAny>,
1030 negatives_by_relation: &Bound<'py, PyAny>,
1031 cand_probs_obj: &Bound<'py, PyAny>,
1032 ) -> PyResult<(Py<PyAny>, Py<PyAny>)> {
1033 let candidate_map = self.candidate_map.clone().ok_or_else(|| {
1034 PyRuntimeError::new_err(
1035 "candidate_map not set — call set_candidate_map() before compute_ilp_loss_grad_gpu_relations()",
1036 )
1037 })?;
1038 let num_cands = candidate_map.len() as u32;
1039
1040 let managed = dlpack_from_py(cand_probs_obj)?;
1041 let cand_buf = self
1042 .provider
1043 .from_dlpack_tensors(vec![managed])
1044 .map_err(|e| types::gpu_err("DLPack import", e))?;
1045
1046 let cand_schema = cand_buf.schema().clone();
1047 let cand_dtype = cand_schema
1048 .column_type(0)
1049 .ok_or_else(|| PyRuntimeError::new_err("cand_probs has no column type"))?;
1050 let is_f64 = match cand_dtype {
1051 ScalarType::F32 => false,
1052 ScalarType::F64 => true,
1053 other => {
1054 return Err(PyValueError::new_err(format!(
1055 "cand_probs must be F32 or F64, got {:?}",
1056 other
1057 )));
1058 }
1059 };
1060
1061 let cand_rows = read_device_row_count(&self.provider, &cand_buf)
1062 .map_err(|e| types::gpu_err("row count", e))?;
1063 if cand_rows != num_cands as usize {
1064 return Err(PyValueError::new_err(format!(
1065 "cand_probs length ({}) != candidate_map length ({})",
1066 cand_rows, num_cands
1067 )));
1068 }
1069
1070 let positive_groups = self.collect_relation_example_groups(
1071 positives_by_relation,
1072 "positives_by_relation must be a dict[str, sequence[dlpack]]",
1073 )?;
1074 let negative_groups = self.collect_relation_example_groups(
1075 negatives_by_relation,
1076 "negatives_by_relation must be a dict[str, sequence[dlpack]]",
1077 )?;
1078
1079 let num_pos: u32 = positive_groups.iter().map(|g| g.num_rows).sum();
1080 let num_neg: u32 = negative_groups.iter().map(|g| g.num_rows).sum();
1081 let num_facts = num_pos + num_neg;
1082
1083 if num_facts == 0 {
1084 return self.build_zero_loss_grad(py, num_cands, is_f64);
1085 }
1086
1087 let mut is_positive_host = Vec::with_capacity(num_facts as usize);
1088 is_positive_host.extend(std::iter::repeat_n(1u8, num_pos as usize));
1089 is_positive_host.extend(std::iter::repeat_n(0u8, num_neg as usize));
1090
1091 if self.strict_zero_dtoh && self.executor.ilp_registry().has_sparse_device_mask() {
1092 self.evaluate_ilp_plan(py)?;
1093 }
1094
1095 let tagged = self
1096 .executor
1097 .ilp_last_result()
1098 .ok_or_else(|| PyRuntimeError::new_err("No ILP result — call evaluate() first"))?;
1099
1100 let mut tasks: Vec<ilp_gpu::CooTask> = Vec::new();
1101 let mut fact_indices_buffers: Vec<xlog_cuda::memory::TrackedCudaSlice<u32>> = Vec::new();
1102 let mut next_global_idx: u32 = 0;
1103
1104 for group in positive_groups.iter().chain(negative_groups.iter()) {
1105 if group.num_rows == 0 {
1106 continue;
1107 }
1108 let k_idx = self
1109 .rel_index
1110 .iter()
1111 .position(|(_, name)| name == &group.relation)
1112 .ok_or_else(|| {
1113 PyValueError::new_err(format!(
1114 "Relation '{}' not in ILP schema",
1115 group.relation
1116 ))
1117 })? as u32;
1118
1119 let relevant_entries: Vec<&xlog_runtime::ilp_registry::IlpTagEntry> = tagged
1120 .entries
1121 .iter()
1122 .filter(|e| e.k == k_idx && e.num_rows > 0 && e.buffer.is_some())
1123 .collect();
1124 if relevant_entries.is_empty() {
1125 next_global_idx = next_global_idx.saturating_add(group.num_rows);
1126 continue;
1127 }
1128
1129 let arity = group.query_buf.arity();
1130 if arity == 0 {
1131 next_global_idx = next_global_idx.saturating_add(group.num_rows);
1132 continue;
1133 }
1134 let keys: Vec<usize> = (0..arity).collect();
1135
1136 let global_indices: Vec<u32> =
1137 (next_global_idx..next_global_idx + group.num_rows).collect();
1138 let mut d_fact_indices = self
1139 .provider
1140 .memory()
1141 .alloc::<u32>(group.num_rows as usize)
1142 .map_err(|e| types::gpu_err("alloc fact_indices", e))?;
1143 self.provider
1144 .device()
1145 .inner()
1146 .htod_sync_copy_into(&global_indices, &mut d_fact_indices)
1147 .map_err(|e| types::gpu_err("htod fact_indices", e))?;
1148 let fi_idx = fact_indices_buffers.len();
1149 fact_indices_buffers.push(d_fact_indices);
1150
1151 for entry in &relevant_entries {
1152 let cidx = match candidate_map.get(&(entry.i, entry.j, entry.k)) {
1153 Some(c) => *c,
1154 None => continue,
1155 };
1156 let entry_buf = entry.buffer.as_ref().ok_or_else(|| {
1157 PyRuntimeError::new_err("internal: filtered entry has no buffer")
1158 })?;
1159 let d_mask = self
1160 .provider
1161 .membership_mask_device(&group.query_buf, entry_buf, &keys, &keys)
1162 .map_err(|e| types::gpu_err("membership_mask", e))?;
1163 tasks.push(ilp_gpu::CooTask {
1164 cidx,
1165 num_query: group.num_rows,
1166 d_mask,
1167 fact_indices_idx: fi_idx,
1168 });
1169 }
1170
1171 next_global_idx = next_global_idx.saturating_add(group.num_rows);
1172 }
1173
1174 let num_tasks = tasks.len();
1175 let upper_bound: u32 = tasks.iter().map(|t| t.num_query).sum();
1176 if upper_bound == 0 || num_tasks == 0 {
1177 return self.build_loss_grad_empty_coo(
1178 py,
1179 &is_positive_host,
1180 num_facts,
1181 num_cands,
1182 is_f64,
1183 );
1184 }
1185
1186 let coo_bytes = (upper_bound as u64) * 8;
1187 let needs_chunking = coo_bytes > self.coo_chunk_budget;
1188
1189 let mut d_is_positive = self
1190 .provider
1191 .memory()
1192 .alloc::<u8>(num_facts as usize)
1193 .map_err(|e| types::gpu_err("alloc is_positive", e))?;
1194 self.provider
1195 .device()
1196 .inner()
1197 .htod_sync_copy_into(&is_positive_host, &mut d_is_positive)
1198 .map_err(|e| types::gpu_err("htod is_positive", e))?;
1199
1200 let cand_col = cand_buf
1201 .column(0)
1202 .ok_or_else(|| PyRuntimeError::new_err("cand_probs has no column"))?;
1203
1204 let (mut d_coo_facts, mut d_coo_cands, actual_nnz) = if !needs_chunking {
1205 ilp_gpu::build_coo_single(
1206 &self.provider,
1207 &tasks,
1208 &fact_indices_buffers,
1209 num_facts,
1210 num_cands,
1211 upper_bound,
1212 )?
1213 } else {
1214 match ilp_gpu::build_coo_chunked(
1215 &self.provider,
1216 &tasks,
1217 &fact_indices_buffers,
1218 num_facts,
1219 self.coo_chunk_budget,
1220 )? {
1221 Some(result) => result,
1222 None => {
1223 return self.build_loss_grad_empty_coo(
1224 py,
1225 &is_positive_host,
1226 num_facts,
1227 num_cands,
1228 is_f64,
1229 );
1230 }
1231 }
1232 };
1233
1234 let d_row_offsets = ilp_gpu::sort_and_build_csr(
1235 &self.provider,
1236 &mut d_coo_facts,
1237 &mut d_coo_cands,
1238 actual_nnz,
1239 num_facts,
1240 )?;
1241
1242 ilp_gpu::forward_backward_reduce(
1243 &self.provider,
1244 py,
1245 &d_row_offsets,
1246 &d_coo_cands,
1247 cand_col,
1248 &d_is_positive,
1249 num_facts,
1250 num_cands,
1251 is_f64,
1252 )
1253 }
1254
1255 pub fn set_rule_mask(
1256 &mut self,
1257 name: String,
1258 mask_hard_flat: &Bound<'_, PyAny>,
1259 mask_soft_flat: &Bound<'_, PyAny>,
1260 schema_size: usize,
1261 ) -> PyResult<()> {
1262 if self.compiled_schema_size > 0 && schema_size != self.compiled_schema_size {
1263 return Err(PyValueError::new_err(format!(
1264 "schema_size mismatch: mask has N={} but compiled program expects N={}",
1265 schema_size, self.compiled_schema_size,
1266 )));
1267 }
1268
1269 let hard_dmt = dlpack_from_py(mask_hard_flat)?;
1270 let soft_dmt = dlpack_from_py(mask_soft_flat)?;
1271
1272 let hard_buf = self
1273 .provider
1274 .from_dlpack_tensors(vec![hard_dmt])
1275 .map_err(types::xlog_err)?;
1276 let soft_buf = self
1277 .provider
1278 .from_dlpack_tensors(vec![soft_dmt])
1279 .map_err(types::xlog_err)?;
1280
1281 self.executor
1282 .ilp_registry_mut()
1283 .insert_mask(name, hard_buf, soft_buf, schema_size);
1284 Ok(())
1285 }
1286
1287 #[pyo3(signature = (name, candidate_ids, soft_probs_dlpack, budget, allow_recursive=false))]
1297 pub fn set_rule_mask_sparse(
1298 &mut self,
1299 name: String,
1300 candidate_ids: Vec<u32>,
1301 soft_probs_dlpack: &Bound<'_, PyAny>,
1302 budget: usize,
1303 allow_recursive: bool,
1304 ) -> PyResult<()> {
1305 if self.strict_zero_dtoh {
1306 return Err(PyRuntimeError::new_err(
1307 "strict_zero_dtoh forbids legacy set_rule_mask_sparse; use set_rule_mask_sparse_selected instead",
1308 ));
1309 }
1310
1311 let tmj = extract_tmj_meta_for_mask(&self.plan, Some(&name));
1312 let n = tmj.schema_size;
1313 if n == 0 {
1314 return Err(PyValueError::new_err(format!(
1315 "no learnable mask '{}' found",
1316 name
1317 )));
1318 }
1319 if self.compiled_schema_size > 0 && n != self.compiled_schema_size {
1320 return Err(PyValueError::new_err(format!(
1321 "schema_size mismatch for '{}': plan N={} compiled N={}",
1322 name, n, self.compiled_schema_size
1323 )));
1324 }
1325
1326 let expected_c = self.expected_candidate_count(&name, allow_recursive)?;
1327 if candidate_ids.len() != expected_c {
1328 return Err(PyValueError::new_err(format!(
1329 "candidate_ids length {} != expected candidate count {}",
1330 candidate_ids.len(),
1331 expected_c
1332 )));
1333 }
1334 for (idx, &cid) in candidate_ids.iter().enumerate() {
1335 if cid != idx as u32 {
1336 return Err(PyValueError::new_err(format!(
1337 "candidate_ids must be [0..{}), got id {} at position {}",
1338 expected_c, cid, idx
1339 )));
1340 }
1341 }
1342
1343 let soft_dmt = dlpack_from_py(soft_probs_dlpack)?;
1345 let soft_buf = self
1346 .provider
1347 .from_dlpack_tensors(vec![soft_dmt])
1348 .map_err(types::xlog_err)?;
1349
1350 let soft_probs = self
1352 .provider
1353 .download_column_untracked::<f64>(&soft_buf, 0)
1354 .map_err(types::xlog_err)?;
1355
1356 if soft_probs.len() != candidate_ids.len() {
1357 return Err(PyValueError::new_err(format!(
1358 "soft_probs tensor length {} != candidate_ids length {}",
1359 soft_probs.len(),
1360 candidate_ids.len()
1361 )));
1362 }
1363
1364 let candidate_triples = self.candidate_triples_for_mask(&name, allow_recursive)?;
1365 if candidate_triples.len() != expected_c {
1366 return Err(PyRuntimeError::new_err(format!(
1367 "internal candidate count mismatch: triples={} expected={}",
1368 candidate_triples.len(),
1369 expected_c
1370 )));
1371 }
1372
1373 let active_soft: Vec<f32> = soft_probs.iter().map(|&v| v as f32).collect();
1375
1376 self.executor
1377 .ilp_registry_mut()
1378 .insert_mask_from_sparse(name, n, &candidate_triples, &active_soft, budget)
1379 .map_err(types::xlog_err)
1380 }
1381
1382 #[pyo3(signature = (name, selected_candidate_ids, selected_soft_probs_dlpack, allow_recursive=false))]
1389 pub fn set_rule_mask_sparse_selected(
1390 &mut self,
1391 name: String,
1392 selected_candidate_ids: Vec<u32>,
1393 selected_soft_probs_dlpack: &Bound<'_, PyAny>,
1394 allow_recursive: bool,
1395 ) -> PyResult<()> {
1396 if self.strict_zero_dtoh {
1397 return Err(PyRuntimeError::new_err(
1398 "strict_zero_dtoh forbids set_rule_mask_sparse_selected; use explicit compatibility export instead",
1399 ));
1400 }
1401 let tmj = extract_tmj_meta_for_mask(&self.plan, Some(&name));
1402 let n = tmj.schema_size;
1403 if n == 0 {
1404 return Err(PyValueError::new_err(format!(
1405 "no learnable mask '{}' found",
1406 name
1407 )));
1408 }
1409 if self.compiled_schema_size > 0 && n != self.compiled_schema_size {
1410 return Err(PyValueError::new_err(format!(
1411 "schema_size mismatch for '{}': plan N={} compiled N={}",
1412 name, n, self.compiled_schema_size
1413 )));
1414 }
1415
1416 let candidate_triples = self.candidate_triples_for_mask(&name, allow_recursive)?;
1417 let expected_c = candidate_triples.len();
1418
1419 let soft_dmt = dlpack_from_py(selected_soft_probs_dlpack)?;
1420 let soft_buf = self
1421 .provider
1422 .from_dlpack_tensors(vec![soft_dmt])
1423 .map_err(types::xlog_err)?;
1424 let selected_len = usize::try_from(soft_buf.num_rows())
1425 .map_err(|_| PyValueError::new_err("selected soft_probs length overflow"))?;
1426
1427 if selected_len != selected_candidate_ids.len() {
1428 return Err(PyValueError::new_err(format!(
1429 "selected soft_probs length {} != selected_candidate_ids length {}",
1430 selected_len,
1431 selected_candidate_ids.len()
1432 )));
1433 }
1434
1435 let mut seen = HashSet::with_capacity(selected_candidate_ids.len());
1436 let mut selected_entries = Vec::with_capacity(selected_candidate_ids.len());
1437 for (pos, &cid) in selected_candidate_ids.iter().enumerate() {
1438 let idx = usize::try_from(cid).map_err(|_| {
1439 PyValueError::new_err(format!("candidate id {} overflows usize", cid))
1440 })?;
1441 if idx >= expected_c {
1442 return Err(PyValueError::new_err(format!(
1443 "selected candidate id {} out of range [0, {}) at position {}",
1444 cid, expected_c, pos
1445 )));
1446 }
1447 if !seen.insert(cid) {
1448 return Err(PyValueError::new_err(format!(
1449 "duplicate selected candidate id {} at position {}",
1450 cid, pos
1451 )));
1452 }
1453 selected_entries.push(candidate_triples[idx]);
1454 }
1455
1456 self.executor
1457 .ilp_registry_mut()
1458 .insert_selected_mask(name, n, &selected_entries);
1459 Ok(())
1460 }
1461
1462 #[pyo3(signature = (name, selected_candidate_ids_dlpack, selected_soft_probs_dlpack, allow_recursive=false))]
1466 pub fn set_rule_mask_sparse_selected_device(
1467 &mut self,
1468 name: String,
1469 selected_candidate_ids_dlpack: &Bound<'_, PyAny>,
1470 selected_soft_probs_dlpack: &Bound<'_, PyAny>,
1471 allow_recursive: bool,
1472 ) -> PyResult<()> {
1473 self.set_rule_mask_sparse_selected_device_impl(
1474 name,
1475 selected_candidate_ids_dlpack,
1476 selected_soft_probs_dlpack,
1477 allow_recursive,
1478 true,
1479 )
1480 }
1481
1482 pub fn evaluate(&mut self, py: Python<'_>) -> PyResult<()> {
1483 if self.strict_zero_dtoh && self.executor.ilp_registry().has_sparse_device_mask() {
1484 return Err(PyRuntimeError::new_err(
1485 "SparseDevice evaluate() is incompatible with strict_zero_dtoh; \
1486use train_only(..., strict_gpu_native=True) or export an explicit compatibility mask first",
1487 ));
1488 }
1489 self.evaluate_ilp_plan(py)
1490 }
1491
1492 pub fn reset_runtime(&mut self, py: Python<'_>) -> PyResult<()> {
1503 let _result: xlog_core::Result<()> = py.detach(|| {
1504 self.executor.reset_for_ilp();
1506
1507 for (name, schema) in &self.schemas {
1509 let empty = self.provider.create_empty_buffer(schema.clone())?;
1510 self.executor.store_mut().put(name, empty);
1511 }
1512
1513 load_facts_into_store(&self.ast, &self.provider, &mut self.executor, &self.schemas)?;
1515
1516 self.apply_relation_overrides()?;
1518
1519 self.executor.execute_plan(&self.plan)?;
1521
1522 Ok(())
1523 });
1524 _result.map_err(types::xlog_err)?;
1525
1526 self.provider.reset_d2h_transfer_count();
1528
1529 Ok(())
1530 }
1531
1532 pub fn get_tagged_results(&self) -> PyResult<Vec<(u32, u32, u32, u32)>> {
1533 self.ensure_host_semantic_compat(
1534 "get_tagged_results()",
1535 "disable strict_zero_dtoh for host materialization",
1536 )?;
1537 match self.executor.ilp_last_result() {
1538 Some(result) => Ok(result
1539 .entries
1540 .iter()
1541 .map(|e| (e.i, e.j, e.k, e.num_rows))
1542 .collect()),
1543 None => Ok(Vec::new()),
1544 }
1545 }
1546
1547 pub fn fact_exists(&self, relation: &str, values: Vec<i64>) -> PyResult<bool> {
1548 self.ensure_host_semantic_compat("fact_exists()", "batch_fact_membership_device(...)")?;
1549 let buf =
1550 self.executor.store().get(relation).ok_or_else(|| {
1551 PyValueError::new_err(format!("Relation '{}' not found", relation))
1552 })?;
1553
1554 Self::fact_exists_in_buffer(&self.provider, buf, &values).map_err(types::xlog_err)
1555 }
1556
1557 #[pyo3(signature = (rel_name))]
1559 pub fn relation_facts(&self, rel_name: String) -> PyResult<Vec<Vec<i64>>> {
1560 self.ensure_host_semantic_compat(
1561 "relation_facts()",
1562 "disable strict_zero_dtoh for host materialization",
1563 )?;
1564 let buf =
1565 self.executor.store().get(&rel_name).ok_or_else(|| {
1566 PyValueError::new_err(format!("Relation '{}' not found", rel_name))
1567 })?;
1568
1569 let num_rows =
1570 read_device_row_count(&self.provider, buf).map_err(types::xlog_err)? as usize;
1571 if num_rows == 0 {
1572 return Ok(Vec::new());
1573 }
1574
1575 let schema = buf.schema();
1577 let mut columns: Vec<Vec<i64>> = Vec::new();
1578 for col_idx in 0..buf.arity() {
1579 let col_type = schema.column_type(col_idx).ok_or_else(|| {
1580 PyRuntimeError::new_err(format!("Column {} type not found in schema", col_idx))
1581 })?;
1582 let col_i64: Vec<i64> = match col_type {
1583 ScalarType::I64 => self
1584 .provider
1585 .download_column::<i64>(buf, col_idx)
1586 .map_err(types::xlog_err)?,
1587 ScalarType::I32 => self
1588 .provider
1589 .download_column::<i32>(buf, col_idx)
1590 .map_err(types::xlog_err)?
1591 .into_iter()
1592 .map(|v| v as i64)
1593 .collect(),
1594 ScalarType::U32 | ScalarType::Symbol => self
1595 .provider
1596 .download_column::<u32>(buf, col_idx)
1597 .map_err(types::xlog_err)?
1598 .into_iter()
1599 .map(|v| v as i64)
1600 .collect(),
1601 ScalarType::U64 => self
1602 .provider
1603 .download_column::<u64>(buf, col_idx)
1604 .map_err(types::xlog_err)?
1605 .into_iter()
1606 .map(|v| v as i64)
1607 .collect(),
1608 ScalarType::Bool => self
1609 .provider
1610 .download_column::<bool>(buf, col_idx)
1611 .map_err(types::xlog_err)?
1612 .into_iter()
1613 .map(|v| if v { 1i64 } else { 0i64 })
1614 .collect(),
1615 ScalarType::F32 | ScalarType::F64 => {
1616 return Err(PyRuntimeError::new_err(format!(
1617 "relation_facts does not support float column type {:?}",
1618 col_type
1619 )));
1620 }
1621 };
1622 columns.push(col_i64);
1623 }
1624
1625 let mut result = Vec::with_capacity(num_rows);
1626 for r in 0..num_rows {
1627 let mut row = Vec::with_capacity(buf.arity());
1628 for c in 0..buf.arity() {
1629 row.push(columns[c][r]);
1630 }
1631 result.push(row);
1632 }
1633 Ok(result)
1634 }
1635
1636 #[pyo3(signature = (head_rel, exclude, max_n))]
1641 pub fn sample_false_positives(
1642 &self,
1643 head_rel: String,
1644 exclude: Vec<(String, Vec<i64>)>,
1645 max_n: usize,
1646 ) -> PyResult<Vec<Vec<i64>>> {
1647 self.ensure_host_semantic_compat(
1648 "sample_false_positives()",
1649 "disable strict_zero_dtoh for host materialization",
1650 )?;
1651 let exclude_set: HashSet<Vec<i64>> = exclude
1653 .into_iter()
1654 .filter(|(rel, _)| rel == &head_rel)
1655 .map(|(_, vals)| vals)
1656 .collect();
1657
1658 let buf =
1660 self.executor.store().get(&head_rel).ok_or_else(|| {
1661 PyValueError::new_err(format!("Relation '{}' not found", head_rel))
1662 })?;
1663
1664 let num_rows =
1665 read_device_row_count(&self.provider, buf).map_err(types::xlog_err)? as usize;
1666 if num_rows == 0 {
1667 return Ok(Vec::new());
1668 }
1669
1670 let schema = buf.schema();
1671 let mut columns: Vec<Vec<i64>> = Vec::new();
1672 for col_idx in 0..buf.arity() {
1673 let col_type = schema.column_type(col_idx).ok_or_else(|| {
1674 PyRuntimeError::new_err(format!("Column {} type not found in schema", col_idx))
1675 })?;
1676 let col_i64: Vec<i64> = match col_type {
1677 ScalarType::I64 => self
1678 .provider
1679 .download_column::<i64>(buf, col_idx)
1680 .map_err(types::xlog_err)?,
1681 ScalarType::I32 => self
1682 .provider
1683 .download_column::<i32>(buf, col_idx)
1684 .map_err(types::xlog_err)?
1685 .into_iter()
1686 .map(|v| v as i64)
1687 .collect(),
1688 ScalarType::U32 | ScalarType::Symbol => self
1689 .provider
1690 .download_column::<u32>(buf, col_idx)
1691 .map_err(types::xlog_err)?
1692 .into_iter()
1693 .map(|v| v as i64)
1694 .collect(),
1695 ScalarType::U64 => self
1696 .provider
1697 .download_column::<u64>(buf, col_idx)
1698 .map_err(types::xlog_err)?
1699 .into_iter()
1700 .map(|v| v as i64)
1701 .collect(),
1702 ScalarType::Bool => self
1703 .provider
1704 .download_column::<bool>(buf, col_idx)
1705 .map_err(types::xlog_err)?
1706 .into_iter()
1707 .map(|v| if v { 1i64 } else { 0i64 })
1708 .collect(),
1709 ScalarType::F32 | ScalarType::F64 => {
1710 return Err(PyRuntimeError::new_err(format!(
1711 "sample_false_positives does not support float column type {:?}",
1712 col_type
1713 )));
1714 }
1715 };
1716 columns.push(col_i64);
1717 }
1718
1719 let mut result = Vec::with_capacity(max_n.min(num_rows));
1721 for r in 0..num_rows {
1722 if result.len() >= max_n {
1723 break;
1724 }
1725 let mut row = Vec::with_capacity(buf.arity());
1726 for c in 0..buf.arity() {
1727 row.push(columns[c][r]);
1728 }
1729 if !exclude_set.contains(&row) {
1730 result.push(row);
1731 }
1732 }
1733 Ok(result)
1734 }
1735
1736 pub fn tagged_entries_containing_fact(
1737 &self,
1738 relation: &str,
1739 values: Vec<i64>,
1740 ) -> PyResult<Vec<(u32, u32, u32)>> {
1741 self.ensure_host_semantic_compat(
1742 "tagged_entries_containing_fact()",
1743 "batch_tagged_credit_device(...)",
1744 )?;
1745 let k_idx = self
1746 .rel_index
1747 .iter()
1748 .position(|(_, name)| name == relation)
1749 .ok_or_else(|| {
1750 PyValueError::new_err(format!("Relation '{}' not in ILP schema", relation))
1751 })? as u32;
1752
1753 let tagged = match self.executor.ilp_last_result() {
1754 Some(t) => t,
1755 None => return Ok(Vec::new()),
1756 };
1757
1758 let mut result = Vec::new();
1759 for entry in &tagged.entries {
1760 if entry.k != k_idx || entry.num_rows == 0 {
1761 continue;
1762 }
1763
1764 let (_, left_name) = &self.rel_index[entry.i as usize];
1765 let (_, right_name) = &self.rel_index[entry.j as usize];
1766
1767 let left_buf = match self.executor.store().get(left_name) {
1768 Some(buf) if buf.arity() > 0 => buf,
1769 _ => continue,
1770 };
1771 let right_buf = match self.executor.store().get(right_name) {
1772 Some(buf) if buf.arity() > 0 => buf,
1773 _ => continue,
1774 };
1775
1776 let left_max = self.left_keys.iter().copied().max().unwrap_or(0);
1778 let right_max = self.right_keys.iter().copied().max().unwrap_or(0);
1779 if left_buf.arity() <= left_max || right_buf.arity() <= right_max {
1780 continue;
1781 }
1782
1783 let joined = self
1784 .provider
1785 .hash_join_v2(
1786 left_buf,
1787 right_buf,
1788 &self.left_keys,
1789 &self.right_keys,
1790 JoinType::Inner,
1791 )
1792 .map_err(types::xlog_err)?;
1793
1794 let found = if !self.head_projection.is_empty()
1797 && self.head_projection.len() == values.len()
1798 {
1799 Self::fact_exists_projected(&self.provider, &joined, &values, &self.head_projection)
1800 .map_err(types::xlog_err)?
1801 } else {
1802 Self::fact_exists_in_buffer(&self.provider, &joined, &values)
1803 .map_err(types::xlog_err)?
1804 };
1805
1806 if found {
1807 result.push((entry.i, entry.j, entry.k));
1808 }
1809 }
1810 Ok(result)
1811 }
1812
1813 pub fn ilp_schema_size(&self) -> usize {
1814 self.rel_index.len()
1815 }
1816
1817 pub fn ilp_relation_names(&self) -> Vec<String> {
1818 self.rel_index
1819 .iter()
1820 .map(|(_, name)| name.clone())
1821 .collect()
1822 }
1823
1824 pub fn relation_type_annotations(&self) -> Vec<(String, Vec<String>)> {
1830 self.ast
1831 .predicates
1832 .iter()
1833 .map(|pred| {
1834 let types = pred
1835 .schema_columns()
1836 .iter()
1837 .map(|column| type_ref_name(&column.typ))
1838 .collect();
1839 (pred.name.clone(), types)
1840 })
1841 .collect()
1842 }
1843
1844 #[pyo3(signature = (mask_name, allow_recursive=false))]
1856 fn valid_candidates(
1857 &self,
1858 py: Python<'_>,
1859 mask_name: String,
1860 allow_recursive: bool,
1861 ) -> PyResult<Vec<StdHashMap<String, Py<PyAny>>>> {
1862 let candidates = self.candidate_triples_for_mask(&mask_name, allow_recursive)?;
1863
1864 let result: Vec<StdHashMap<String, Py<PyAny>>> = candidates
1865 .iter()
1866 .enumerate()
1867 .map(
1868 |(id, &(i, j, k))| -> PyResult<StdHashMap<String, Py<PyAny>>> {
1869 let mut d = StdHashMap::new();
1870 d.insert("id".into(), id.into_pyobject(py)?.into_any().unbind());
1871 d.insert("i".into(), i.into_pyobject(py)?.into_any().unbind());
1872 d.insert("j".into(), j.into_pyobject(py)?.into_any().unbind());
1873 d.insert("k".into(), k.into_pyobject(py)?.into_any().unbind());
1874 d.insert(
1875 "left_name".into(),
1876 self.rel_index[i as usize]
1877 .1
1878 .clone()
1879 .into_pyobject(py)?
1880 .into_any()
1881 .unbind(),
1882 );
1883 d.insert(
1884 "right_name".into(),
1885 self.rel_index[j as usize]
1886 .1
1887 .clone()
1888 .into_pyobject(py)?
1889 .into_any()
1890 .unbind(),
1891 );
1892 d.insert(
1893 "head_name".into(),
1894 self.rel_index[k as usize]
1895 .1
1896 .clone()
1897 .into_pyobject(py)?
1898 .into_any()
1899 .unbind(),
1900 );
1901 Ok(d)
1902 },
1903 )
1904 .collect::<PyResult<_>>()?;
1905
1906 Ok(result)
1907 }
1908
1909 pub fn commit_induced_rule(&mut self, rule_source: &str) -> PyResult<()> {
1910 let new_base = format!("{}\n{}", self.base_source, rule_source);
1911
1912 let ast = xlog_logic::parse_program(&new_base).map_err(types::val_err)?;
1913 let mut compiler = xlog_logic::Compiler::new();
1914 compiler.set_max_active_rules(self.max_active_rules);
1915 let plan = compiler.compile_program(&ast).map_err(types::xlog_err)?;
1916 let schemas = compiler.schemas().clone();
1917
1918 self.executor.reset_for_mc();
1919 for (name, rel_id) in compiler.rel_ids() {
1920 self.executor.register_relation(*rel_id, name);
1921 }
1922 for (name, schema) in &schemas {
1923 let empty = self
1924 .provider
1925 .create_empty_buffer(schema.clone())
1926 .map_err(types::xlog_err)?;
1927 self.executor.store_mut().put(name, empty);
1928 }
1929 load_facts_into_store(&ast, &self.provider, &mut self.executor, &schemas)
1930 .map_err(types::xlog_err)?;
1931 self.apply_relation_overrides().map_err(types::xlog_err)?;
1932 self.executor.execute_plan(&plan).map_err(types::xlog_err)?;
1933
1934 self.base_source = new_base;
1935 self.ast = ast;
1936 let tmj = extract_tmj_meta(&plan);
1937 self.left_keys = tmj.left_keys;
1938 self.right_keys = tmj.right_keys;
1939 self.head_projection = tmj.head_projection;
1940 self.compiled_schema_size = tmj.schema_size;
1941 self.head_rel_name = tmj.head_rel_name;
1942 self.plan = plan;
1943 self.schemas = schemas;
1944 Ok(())
1945 }
1946
1947 pub fn batch_fact_membership_device(
1952 &self,
1953 py: Python<'_>,
1954 relation: &str,
1955 facts: Vec<Vec<i64>>,
1956 ) -> PyResult<Py<PyAny>> {
1957 let buf =
1958 self.executor.store().get(relation).ok_or_else(|| {
1959 PyValueError::new_err(format!("Relation '{}' not found", relation))
1960 })?;
1961
1962 if facts.is_empty() {
1963 let empty = self
1964 .provider
1965 .memory()
1966 .alloc::<u8>(0)
1967 .map_err(types::xlog_err)?;
1968 return export_device_bool_tensor(&self.provider, py, empty, 0);
1969 }
1970 if buf.arity() == 0 {
1971 let mut zeros = self
1972 .provider
1973 .memory()
1974 .alloc::<u8>(facts.len())
1975 .map_err(types::xlog_err)?;
1976 self.provider
1977 .device()
1978 .inner()
1979 .memset_zeros(&mut zeros)
1980 .map_err(types::xlog_err)?;
1981 return export_device_bool_tensor(&self.provider, py, zeros, facts.len());
1982 }
1983
1984 let col_bytes = pack_i64_columns_typed(relation, &facts, buf.schema())?;
1985 let col_slices: Vec<&[u8]> = col_bytes.iter().map(|c| c.as_slice()).collect();
1986 let query_buf = self
1987 .provider
1988 .create_buffer_from_slices(&col_slices, buf.schema().clone())
1989 .map_err(types::xlog_err)?;
1990
1991 let keys: Vec<usize> = (0..buf.arity()).collect();
1992 let mask = self
1993 .provider
1994 .membership_mask_device(&query_buf, buf, &keys, &keys)
1995 .map_err(types::xlog_err)?;
1996 export_device_bool_tensor(&self.provider, py, mask, facts.len())
1997 }
1998
1999 pub fn batch_fact_membership(
2000 &self,
2001 relation: &str,
2002 facts: Vec<Vec<i64>>,
2003 ) -> PyResult<Vec<bool>> {
2004 self.ensure_host_semantic_compat(
2005 "batch_fact_membership()",
2006 "batch_fact_membership_device(...)",
2007 )?;
2008 if facts.is_empty() {
2009 return Ok(Vec::new());
2010 }
2011
2012 let buf =
2013 self.executor.store().get(relation).ok_or_else(|| {
2014 PyValueError::new_err(format!("Relation '{}' not found", relation))
2015 })?;
2016
2017 let arity = buf.arity();
2018 if arity == 0 {
2019 return Ok(vec![false; facts.len()]);
2020 }
2021
2022 let col_bytes = pack_i64_columns_typed(relation, &facts, buf.schema())?;
2024 let col_slices: Vec<&[u8]> = col_bytes.iter().map(|c| c.as_slice()).collect();
2025 let query_buf = self
2026 .provider
2027 .create_buffer_from_slices(&col_slices, buf.schema().clone())
2028 .map_err(types::xlog_err)?;
2029
2030 let keys: Vec<usize> = (0..arity).collect();
2032
2033 self.provider
2034 .membership_mask(&query_buf, buf, &keys, &keys)
2035 .map_err(types::xlog_err)
2036 }
2037
2038 pub fn batch_tagged_credit(
2044 &self,
2045 relation: &str,
2046 facts: Vec<Vec<i64>>,
2047 ) -> PyResult<Vec<Vec<(u32, u32, u32)>>> {
2048 self.ensure_host_semantic_compat(
2049 "batch_tagged_credit()",
2050 "batch_tagged_credit_device(...)",
2051 )?;
2052 if facts.is_empty() {
2053 return Ok(Vec::new());
2054 }
2055
2056 let k_idx = self
2058 .rel_index
2059 .iter()
2060 .position(|(_, name)| name == relation)
2061 .ok_or_else(|| {
2062 PyValueError::new_err(format!("Relation '{}' not in ILP schema", relation))
2063 })? as u32;
2064
2065 let tagged = match self.executor.ilp_last_result() {
2066 Some(t) => t,
2067 None => return Ok(vec![Vec::new(); facts.len()]),
2068 };
2069
2070 let relevant_entries: Vec<&xlog_runtime::ilp_registry::IlpTagEntry> = tagged
2072 .entries
2073 .iter()
2074 .filter(|e| e.k == k_idx && e.num_rows > 0 && e.buffer.is_some())
2075 .collect();
2076
2077 if relevant_entries.is_empty() {
2078 return Ok(vec![Vec::new(); facts.len()]);
2079 }
2080
2081 let first_buf = relevant_entries[0]
2083 .buffer
2084 .as_ref()
2085 .ok_or_else(|| PyRuntimeError::new_err("internal: filtered entry has no buffer"))?;
2086 let arity = first_buf.arity();
2087 if arity == 0 {
2088 return Ok(vec![Vec::new(); facts.len()]);
2089 }
2090
2091 let schema = first_buf.schema().clone();
2093 let col_bytes = pack_i64_columns_typed(relation, &facts, &schema)?;
2094 let col_slices: Vec<&[u8]> = col_bytes.iter().map(|c| c.as_slice()).collect();
2095 let query_buf = self
2096 .provider
2097 .create_buffer_from_slices(&col_slices, schema)
2098 .map_err(types::xlog_err)?;
2099
2100 let keys: Vec<usize> = (0..arity).collect();
2101
2102 let mut per_fact_credits: Vec<Vec<(u32, u32, u32)>> = vec![Vec::new(); facts.len()];
2104
2105 for entry in &relevant_entries {
2106 let entry_buf = entry
2107 .buffer
2108 .as_ref()
2109 .ok_or_else(|| PyRuntimeError::new_err("internal: filtered entry has no buffer"))?;
2110 let mask = self
2111 .provider
2112 .membership_mask(&query_buf, entry_buf, &keys, &keys)
2113 .map_err(types::xlog_err)?;
2114
2115 for (fact_idx, &found) in mask.iter().enumerate() {
2116 if found {
2117 per_fact_credits[fact_idx].push((entry.i, entry.j, entry.k));
2118 }
2119 }
2120 }
2121
2122 Ok(per_fact_credits)
2123 }
2124
2125 pub fn batch_tagged_credit_device(
2135 &self,
2136 py: Python<'_>,
2137 relation: &str,
2138 facts: Vec<Vec<i64>>,
2139 ) -> PyResult<IlpTaggedCreditDeviceResult> {
2140 if facts.is_empty() {
2141 return empty_tagged_credit_device_result(&self.provider, py, 0);
2142 }
2143
2144 let k_idx = self
2145 .rel_index
2146 .iter()
2147 .position(|(_, name)| name == relation)
2148 .ok_or_else(|| {
2149 PyValueError::new_err(format!("Relation '{}' not in ILP schema", relation))
2150 })? as u32;
2151
2152 let tagged = match self.executor.ilp_last_result() {
2153 Some(t) => t,
2154 None => return empty_tagged_credit_device_result(&self.provider, py, facts.len()),
2155 };
2156
2157 let relevant_entries: Vec<&xlog_runtime::ilp_registry::IlpTagEntry> = tagged
2158 .entries
2159 .iter()
2160 .filter(|e| e.k == k_idx && e.num_rows > 0 && e.buffer.is_some())
2161 .collect();
2162
2163 if relevant_entries.is_empty() {
2164 return empty_tagged_credit_device_result(&self.provider, py, facts.len());
2165 }
2166
2167 let first_buf = relevant_entries[0]
2168 .buffer
2169 .as_ref()
2170 .ok_or_else(|| PyRuntimeError::new_err("internal: filtered entry has no buffer"))?;
2171 let arity = first_buf.arity();
2172 if arity == 0 {
2173 return empty_tagged_credit_device_result(&self.provider, py, facts.len());
2174 }
2175
2176 let schema = first_buf.schema().clone();
2177 let col_bytes = pack_i64_columns_typed(relation, &facts, &schema)?;
2178 let col_slices: Vec<&[u8]> = col_bytes.iter().map(|c| c.as_slice()).collect();
2179 let query_buf = self
2180 .provider
2181 .create_buffer_from_slices(&col_slices, schema)
2182 .map_err(types::xlog_err)?;
2183
2184 let keys: Vec<usize> = (0..arity).collect();
2185 let num_facts = u32::try_from(facts.len())
2186 .map_err(|_| PyValueError::new_err("facts length exceeds u32::MAX"))?;
2187 let num_entries = u32::try_from(relevant_entries.len())
2188 .map_err(|_| PyValueError::new_err("entry count exceeds u32::MAX"))?;
2189 let upper_bound = num_facts
2190 .checked_mul(num_entries)
2191 .ok_or_else(|| PyValueError::new_err("credit upper bound overflow"))?;
2192
2193 let fact_indices_host: Vec<u32> = (0..num_facts).collect();
2194 let mut d_fact_indices = self
2195 .provider
2196 .memory()
2197 .alloc::<u32>(num_facts as usize)
2198 .map_err(|e| types::gpu_err("alloc fact_indices", e))?;
2199 self.provider
2200 .device()
2201 .inner()
2202 .htod_sync_copy_into(&fact_indices_host, &mut d_fact_indices)
2203 .map_err(|e| types::gpu_err("htod fact_indices", e))?;
2204
2205 let mut entry_i_host = Vec::with_capacity(relevant_entries.len());
2206 let mut entry_j_host = Vec::with_capacity(relevant_entries.len());
2207 let mut entry_k_host = Vec::with_capacity(relevant_entries.len());
2208 let mut tasks = Vec::with_capacity(relevant_entries.len());
2209
2210 for (entry_idx, entry) in relevant_entries.iter().enumerate() {
2211 let entry_buf = entry
2212 .buffer
2213 .as_ref()
2214 .ok_or_else(|| PyRuntimeError::new_err("internal: filtered entry has no buffer"))?;
2215 let d_mask = self
2216 .provider
2217 .membership_mask_device(&query_buf, entry_buf, &keys, &keys)
2218 .map_err(|e| types::gpu_err("membership_mask", e))?;
2219 tasks.push(ilp_gpu::CooTask {
2220 d_mask,
2221 fact_indices_idx: 0,
2222 cidx: entry_idx as u32,
2223 num_query: num_facts,
2224 });
2225 entry_i_host.push(entry.i);
2226 entry_j_host.push(entry.j);
2227 entry_k_host.push(entry.k);
2228 }
2229
2230 let (mut d_coo_facts, mut d_coo_cands, actual_nnz) = ilp_gpu::build_coo_single(
2231 &self.provider,
2232 &tasks,
2233 &[d_fact_indices],
2234 num_facts,
2235 num_entries,
2236 upper_bound,
2237 )?;
2238 let d_row_offsets = ilp_gpu::sort_and_build_csr(
2239 &self.provider,
2240 &mut d_coo_facts,
2241 &mut d_coo_cands,
2242 actual_nnz,
2243 num_facts,
2244 )?;
2245
2246 let mut d_entry_i = self
2247 .provider
2248 .memory()
2249 .alloc::<u32>(entry_i_host.len())
2250 .map_err(|e| types::gpu_err("alloc entry_i", e))?;
2251 let mut d_entry_j = self
2252 .provider
2253 .memory()
2254 .alloc::<u32>(entry_j_host.len())
2255 .map_err(|e| types::gpu_err("alloc entry_j", e))?;
2256 let mut d_entry_k = self
2257 .provider
2258 .memory()
2259 .alloc::<u32>(entry_k_host.len())
2260 .map_err(|e| types::gpu_err("alloc entry_k", e))?;
2261 self.provider
2262 .device()
2263 .inner()
2264 .htod_sync_copy_into(&entry_i_host, &mut d_entry_i)
2265 .map_err(|e| types::gpu_err("htod entry_i", e))?;
2266 self.provider
2267 .device()
2268 .inner()
2269 .htod_sync_copy_into(&entry_j_host, &mut d_entry_j)
2270 .map_err(|e| types::gpu_err("htod entry_j", e))?;
2271 self.provider
2272 .device()
2273 .inner()
2274 .htod_sync_copy_into(&entry_k_host, &mut d_entry_k)
2275 .map_err(|e| types::gpu_err("htod entry_k", e))?;
2276
2277 Ok(IlpTaggedCreditDeviceResult {
2278 fact_row_offsets: export_device_u32_tensor_as_i32(
2279 &self.provider,
2280 py,
2281 d_row_offsets,
2282 facts.len() + 1,
2283 )?,
2284 entry_indices: export_device_u32_tensor_as_i32(
2285 &self.provider,
2286 py,
2287 d_coo_cands,
2288 upper_bound as usize,
2289 )?,
2290 entry_i: export_device_u32_tensor_as_i32(
2291 &self.provider,
2292 py,
2293 d_entry_i,
2294 entry_i_host.len(),
2295 )?,
2296 entry_j: export_device_u32_tensor_as_i32(
2297 &self.provider,
2298 py,
2299 d_entry_j,
2300 entry_j_host.len(),
2301 )?,
2302 entry_k: export_device_u32_tensor_as_i32(
2303 &self.provider,
2304 py,
2305 d_entry_k,
2306 entry_k_host.len(),
2307 )?,
2308 })
2309 }
2310
2311 pub fn d2h_transfer_count(&self) -> u64 {
2312 self.provider.d2h_transfer_count()
2313 }
2314
2315 pub fn reset_d2h_transfer_count(&self) {
2316 self.provider.reset_d2h_transfer_count()
2317 }
2318
2319 pub fn host_transfer_stats(&self, py: Python<'_>) -> PyResult<Py<PyAny>> {
2320 let stats = self.provider.host_transfer_stats();
2321 let dict = PyDict::new(py);
2322 dict.set_item("dtoh_bytes", stats.dtoh_bytes)?;
2323 dict.set_item("htod_bytes", stats.htod_bytes)?;
2324 dict.set_item("dtoh_calls", stats.dtoh_calls)?;
2325 dict.set_item("htod_calls", stats.htod_calls)?;
2326 Ok(dict.into())
2327 }
2328
2329 pub fn reset_host_transfer_stats(&self) {
2330 self.provider.reset_host_transfer_stats()
2331 }
2332}
2333
2334impl CompiledIlpProgram {
2339 fn collect_relation_example_groups(
2340 &self,
2341 relations_obj: &Bound<'_, PyAny>,
2342 err_msg: &str,
2343 ) -> PyResult<Vec<RelationExampleGroup>> {
2344 let dict = relations_obj
2345 .cast::<PyDict>()
2346 .map_err(|_| PyValueError::new_err(err_msg.to_string()))?;
2347 let mut groups = Vec::with_capacity(dict.len());
2348 for (name_obj, columns_obj) in dict.iter() {
2349 let relation: String = name_obj.extract()?;
2350 let schema = self.schemas.get(&relation).ok_or_else(|| {
2351 PyValueError::new_err(format!(
2352 "Unknown relation {} (not present in compiled schemas)",
2353 relation
2354 ))
2355 })?;
2356 let tensors = collect_dlpack_columns(
2357 &columns_obj,
2358 &format!("Relation {} must be a sequence of DLPack columns", relation),
2359 )?;
2360 let query_buf = self
2361 .provider
2362 .from_dlpack_tensors_with_schema(schema.clone(), tensors)
2363 .map_err(types::xlog_err)?;
2364 let num_rows = u32::try_from(query_buf.num_rows())
2365 .map_err(|_| PyValueError::new_err("relation row count exceeds u32::MAX"))?;
2366 groups.push(RelationExampleGroup {
2367 relation,
2368 query_buf,
2369 num_rows,
2370 });
2371 }
2372 Ok(groups)
2373 }
2374
2375 fn apply_relation_overrides(&mut self) -> xlog_core::Result<()> {
2376 let relation_names: Vec<String> = self.relation_overrides.keys().cloned().collect();
2377 for name in relation_names {
2378 let stored = self
2379 .relation_overrides
2380 .get(&name)
2381 .expect("relation override disappeared during runtime reset");
2382 let live = self.provider.clone_buffer(stored)?;
2383 self.executor.put_relation(&name, live);
2384 }
2385 Ok(())
2386 }
2387
2388 fn evaluate_ilp_plan(&mut self, py: Python<'_>) -> PyResult<()> {
2389 let result: xlog_core::Result<xlog_cuda::CudaBuffer> = py.detach(|| {
2390 self.executor.reset_for_mc();
2391 for (name, schema) in &self.schemas {
2392 let empty = self.provider.create_empty_buffer(schema.clone())?;
2393 self.executor.store_mut().put(name, empty);
2394 }
2395 load_facts_into_store(&self.ast, &self.provider, &mut self.executor, &self.schemas)?;
2396 self.apply_relation_overrides()?;
2397 self.executor.execute_plan(&self.plan)
2398 });
2399 result.map_err(types::xlog_err)?;
2400 Ok(())
2401 }
2402
2403 fn set_rule_mask_sparse_selected_device_impl(
2404 &mut self,
2405 name: String,
2406 selected_candidate_ids_dlpack: &Bound<'_, PyAny>,
2407 selected_soft_probs_dlpack: &Bound<'_, PyAny>,
2408 allow_recursive: bool,
2409 validate_ids: bool,
2410 ) -> PyResult<()> {
2411 let tmj = extract_tmj_meta_for_mask(&self.plan, Some(&name));
2412 let n = tmj.schema_size;
2413 if n == 0 {
2414 return Err(PyValueError::new_err(format!(
2415 "no learnable mask '{}' found",
2416 name
2417 )));
2418 }
2419 if self.compiled_schema_size > 0 && n != self.compiled_schema_size {
2420 return Err(PyValueError::new_err(format!(
2421 "schema_size mismatch for '{}': plan N={} compiled N={}",
2422 name, n, self.compiled_schema_size
2423 )));
2424 }
2425
2426 let _ = allow_recursive;
2427 let candidate_order = self.candidate_order.as_ref().ok_or_else(|| {
2428 PyRuntimeError::new_err(
2429 "candidate order not set — call set_candidate_map() before strict sparse mask updates",
2430 )
2431 })?;
2432 let expected_c = candidate_order.len();
2433
2434 let ids_dmt = dlpack_from_py(selected_candidate_ids_dlpack)?;
2435 let ids_buf = self
2436 .provider
2437 .from_dlpack_tensors(vec![ids_dmt])
2438 .map_err(types::xlog_err)?;
2439 let soft_dmt = dlpack_from_py(selected_soft_probs_dlpack)?;
2440 let soft_buf = self
2441 .provider
2442 .from_dlpack_tensors(vec![soft_dmt])
2443 .map_err(types::xlog_err)?;
2444
2445 let selected_len = usize::try_from(ids_buf.num_rows())
2446 .map_err(|_| PyValueError::new_err("selected candidate ids length overflow"))?;
2447 let soft_len = usize::try_from(soft_buf.num_rows())
2448 .map_err(|_| PyValueError::new_err("selected soft_probs length overflow"))?;
2449 if soft_len != selected_len {
2450 return Err(PyValueError::new_err(format!(
2451 "selected soft_probs length {} != selected candidate ids length {}",
2452 soft_len, selected_len
2453 )));
2454 }
2455
2456 if validate_ids {
2457 self.provider
2458 .validate_selected_ids(&ids_buf, expected_c)
2459 .map_err(types::val_err)?;
2460 }
2461
2462 let active_flags = self
2463 .provider
2464 .build_selected_id_mask(&ids_buf, expected_c)
2465 .map_err(types::xlog_err)?;
2466
2467 self.executor
2468 .ilp_registry_mut()
2469 .insert_selected_mask_device(
2470 name,
2471 n,
2472 candidate_order.clone(),
2473 active_flags,
2474 selected_len,
2475 );
2476 Ok(())
2477 }
2478
2479 fn ensure_host_semantic_compat(&self, api: &str, device_hint: &str) -> PyResult<()> {
2480 if self.strict_zero_dtoh {
2481 return Err(PyRuntimeError::new_err(format!(
2482 "strict_zero_dtoh forbids {}; use {} instead",
2483 api, device_hint
2484 )));
2485 }
2486 Ok(())
2487 }
2488
2489 fn build_zero_loss_grad(
2493 &self,
2494 py: Python<'_>,
2495 num_cands: u32,
2496 is_f64: bool,
2497 ) -> PyResult<(Py<PyAny>, Py<PyAny>)> {
2498 if is_f64 {
2499 build_zero_typed::<f64>(&self.provider, py, num_cands, ScalarType::F64)
2500 } else {
2501 build_zero_typed::<f32>(&self.provider, py, num_cands, ScalarType::F32)
2502 }
2503 }
2504
2505 fn build_loss_grad_empty_coo(
2509 &self,
2510 py: Python<'_>,
2511 is_positive_host: &[u8],
2512 num_facts: u32,
2513 num_cands: u32,
2514 is_f64: bool,
2515 ) -> PyResult<(Py<PyAny>, Py<PyAny>)> {
2516 let row_offsets = vec![0u32; (num_facts + 1) as usize];
2518 let mut d_row_offsets = self
2519 .provider
2520 .memory()
2521 .alloc::<u32>((num_facts + 1) as usize)
2522 .map_err(|e| types::gpu_err("alloc", e))?;
2523 self.provider
2524 .device()
2525 .inner()
2526 .htod_sync_copy_into(&row_offsets, &mut d_row_offsets)
2527 .map_err(|e| types::gpu_err("htod", e))?;
2528
2529 let d_col_indices = self
2531 .provider
2532 .memory()
2533 .alloc::<u32>(0)
2534 .map_err(|e| types::gpu_err("alloc", e))?;
2535
2536 let mut d_is_positive = self
2537 .provider
2538 .memory()
2539 .alloc::<u8>(num_facts as usize)
2540 .map_err(|e| types::gpu_err("alloc", e))?;
2541 self.provider
2542 .device()
2543 .inner()
2544 .htod_sync_copy_into(is_positive_host, &mut d_is_positive)
2545 .map_err(|e| types::gpu_err("htod", e))?;
2546
2547 let dummy_col: xlog_cuda::CudaColumn = if is_f64 {
2552 self.provider
2553 .memory()
2554 .alloc::<f64>(num_cands.max(1) as usize)
2555 .map_err(|e| types::gpu_err("alloc dummy", e))?
2556 .into_bytes()
2557 .into()
2558 } else {
2559 self.provider
2560 .memory()
2561 .alloc::<f32>(num_cands.max(1) as usize)
2562 .map_err(|e| types::gpu_err("alloc dummy", e))?
2563 .into_bytes()
2564 .into()
2565 };
2566
2567 ilp_gpu::forward_backward_reduce(
2568 &self.provider,
2569 py,
2570 &d_row_offsets,
2571 &d_col_indices,
2572 &dummy_col,
2573 &d_is_positive,
2574 num_facts,
2575 num_cands,
2576 is_f64,
2577 )
2578 }
2579
2580 fn candidate_triples_for_mask(
2583 &self,
2584 mask_name: &str,
2585 allow_recursive: bool,
2586 ) -> PyResult<Vec<(u32, u32, u32)>> {
2587 let tmj = extract_tmj_meta_for_mask(&self.plan, Some(mask_name));
2588 let n = tmj.schema_size;
2589 if n == 0 {
2590 return Err(PyValueError::new_err(format!(
2591 "no learnable mask '{}' found in compiled program",
2592 mask_name
2593 )));
2594 }
2595 let head_name = &tmj.head_rel_name;
2596 let k_head = self
2597 .rel_index
2598 .iter()
2599 .position(|(_, name)| name == head_name)
2600 .ok_or_else(|| {
2601 PyValueError::new_err(format!(
2602 "head relation '{}' not in rel_index for mask '{}'",
2603 head_name, mask_name
2604 ))
2605 })? as u32;
2606
2607 let has_tuples: Vec<bool> = self
2609 .rel_index
2610 .iter()
2611 .map(|(_, name)| {
2612 self.executor
2613 .store()
2614 .get(name)
2615 .map(|buf| buf.num_rows() > 0)
2616 .unwrap_or(false)
2617 })
2618 .collect();
2619
2620 let mut triples: Vec<(u32, u32, u32)> = Vec::new();
2621 for i in 0..n as u32 {
2622 for j in 0..n as u32 {
2623 let k = k_head;
2624
2625 if !has_tuples[i as usize] && !has_tuples[j as usize] {
2627 continue;
2628 }
2629
2630 if !allow_recursive && (i == k || j == k) && !has_tuples[k as usize] {
2633 continue;
2634 }
2635
2636 triples.push((i, j, k));
2637 }
2638 }
2639 triples.sort_by_key(|&(i, j, k)| (k, i, j));
2640 Ok(triples)
2641 }
2642
2643 fn expected_candidate_count(&self, mask_name: &str, allow_recursive: bool) -> PyResult<usize> {
2644 Ok(self
2645 .candidate_triples_for_mask(mask_name, allow_recursive)?
2646 .len())
2647 }
2648
2649 fn fact_exists_in_buffer(
2650 provider: &CudaKernelProvider,
2651 buf: &xlog_cuda::CudaBuffer,
2652 values: &[i64],
2653 ) -> xlog_core::Result<bool> {
2654 use xlog_core::XlogError;
2655 let num_rows = read_device_row_count(provider, buf)? as usize;
2656 if num_rows == 0 {
2657 return Ok(false);
2658 }
2659 if values.len() != buf.arity() {
2660 return Ok(false);
2661 }
2662
2663 let schema = buf.schema();
2664 let mut columns: Vec<Vec<i64>> = Vec::new();
2665 for col_idx in 0..buf.arity() {
2666 let col_type = schema.column_type(col_idx).ok_or_else(|| {
2667 XlogError::Kernel(format!("Column {} type not found in schema", col_idx))
2668 })?;
2669 let col_i64: Vec<i64> = match col_type {
2670 ScalarType::I64 => provider.download_column::<i64>(buf, col_idx)?,
2671 ScalarType::I32 => provider
2672 .download_column::<i32>(buf, col_idx)?
2673 .into_iter()
2674 .map(|v| v as i64)
2675 .collect(),
2676 ScalarType::U32 | ScalarType::Symbol => provider
2677 .download_column::<u32>(buf, col_idx)?
2678 .into_iter()
2679 .map(|v| v as i64)
2680 .collect(),
2681 ScalarType::U64 => {
2682 let col_u64 = provider.download_column::<u64>(buf, col_idx)?;
2683 col_u64.into_iter().map(|v| v as i64).collect()
2684 }
2685 ScalarType::Bool => provider
2686 .download_column::<bool>(buf, col_idx)?
2687 .into_iter()
2688 .map(|v| if v { 1i64 } else { 0i64 })
2689 .collect(),
2690 ScalarType::F32 | ScalarType::F64 => {
2691 return Err(XlogError::Kernel(format!(
2692 "fact_exists does not support float column type {:?}",
2693 col_type
2694 )));
2695 }
2696 };
2697 columns.push(col_i64);
2698 }
2699
2700 for row in 0..num_rows {
2701 let mut matches = true;
2702 for (col_idx, val) in values.iter().enumerate() {
2703 if columns[col_idx][row] != *val {
2704 matches = false;
2705 break;
2706 }
2707 }
2708 if matches {
2709 return Ok(true);
2710 }
2711 }
2712 Ok(false)
2713 }
2714
2715 fn fact_exists_projected(
2719 provider: &CudaKernelProvider,
2720 buf: &xlog_cuda::CudaBuffer,
2721 values: &[i64],
2722 projection: &[usize],
2723 ) -> xlog_core::Result<bool> {
2724 use xlog_core::XlogError;
2725 let num_rows = read_device_row_count(provider, buf)? as usize;
2726 if num_rows == 0 {
2727 return Ok(false);
2728 }
2729
2730 let schema = buf.schema();
2731 let mut columns: Vec<Vec<i64>> = Vec::new();
2732 for &col_idx in projection {
2733 if col_idx >= buf.arity() {
2734 return Ok(false);
2735 }
2736 let col_type = schema.column_type(col_idx).ok_or_else(|| {
2737 XlogError::Kernel(format!("Column {} type not found in schema", col_idx))
2738 })?;
2739 let col_i64: Vec<i64> = match col_type {
2740 ScalarType::I64 => provider.download_column::<i64>(buf, col_idx)?,
2741 ScalarType::I32 => provider
2742 .download_column::<i32>(buf, col_idx)?
2743 .into_iter()
2744 .map(|v| v as i64)
2745 .collect(),
2746 ScalarType::U32 | ScalarType::Symbol => provider
2747 .download_column::<u32>(buf, col_idx)?
2748 .into_iter()
2749 .map(|v| v as i64)
2750 .collect(),
2751 ScalarType::U64 => provider
2752 .download_column::<u64>(buf, col_idx)?
2753 .into_iter()
2754 .map(|v| v as i64)
2755 .collect(),
2756 ScalarType::Bool => provider
2757 .download_column::<bool>(buf, col_idx)?
2758 .into_iter()
2759 .map(|v| if v { 1i64 } else { 0i64 })
2760 .collect(),
2761 ScalarType::F32 | ScalarType::F64 => {
2762 return Err(XlogError::Kernel(format!(
2763 "fact_exists does not support float column type {:?}",
2764 col_type
2765 )));
2766 }
2767 };
2768 columns.push(col_i64);
2769 }
2770
2771 for row in 0..num_rows {
2772 let mut matches = true;
2773 for (i, val) in values.iter().enumerate() {
2774 if columns[i][row] != *val {
2775 matches = false;
2776 break;
2777 }
2778 }
2779 if matches {
2780 return Ok(true);
2781 }
2782 }
2783 Ok(false)
2784 }
2785}