Skip to main content

pyxlog/
ilp.rs

1// ---------------------------------------------------------------------------
2// ILP (Inductive Logic Programming) Python bindings
3// ---------------------------------------------------------------------------
4
5use 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
30// ---------------------------------------------------------------------------
31// Helper functions
32// ---------------------------------------------------------------------------
33
34struct 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
65/// Allocate zero-initialized loss (scalar) and grad (num_cands) on GPU and
66/// export both as DLPack capsules. Shared by both f32 and f64 paths.
67fn 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    // PyTorch does not support unsigned 32-bit DLPack tensors; export as i32.
145    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
191/// Pack `i64` fact values into typed byte columns according to schema.
192/// Returns one `Vec<u8>` per column with correctly-encoded LE bytes.
193/// Rejects F32/F64 columns (not supported in batch APIs).
194pub(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
366/// Extracted TensorMaskedJoin metadata from the execution plan.
367struct 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        // Descend into the fallback so a TMJ wrapped beneath a promoted
410        // MultiWayJoin is still discoverable. The promoter does not currently
411        // wrap TMJ-bearing trees, but the explicit arm documents the contract
412        // instead of relying on the catch-all.
413        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// ---------------------------------------------------------------------------
457// IlpProgramFactory
458// ---------------------------------------------------------------------------
459
460#[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        // Validate max_active_rules range
471        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        // The frontend runs before the provider is built, so an unparsable
481        // source still fails with a ValueError without initialising CUDA.
482        build_ilp_program(
483            source,
484            max_active_rules,
485            ProviderFor::New { device, memory_mb },
486        )
487    }
488}
489
490/// Where a program being built gets its CUDA provider from.
491enum ProviderFor {
492    /// Build a new provider (and pay for a fresh CUDA context).
493    New { device: usize, memory_mb: u64 },
494    /// Share an existing program's provider.
495    Shared(Arc<CudaKernelProvider>),
496}
497
498fn ms_since(t: Instant) -> f64 {
499    t.elapsed().as_secs_f64() * 1000.0
500}
501
502/// Shared body of [`IlpProgramFactory::compile`] and
503/// [`CompiledIlpProgram::compile_variant`].
504fn 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// ---------------------------------------------------------------------------
594// CompiledIlpProgram — #[pymethods] block
595// ---------------------------------------------------------------------------
596
597#[pymethods]
598impl CompiledIlpProgram {
599    /// Compile `source` on this program's CUDA provider instead of a new one.
600    ///
601    /// The result is a separate program with its own relation store, in the
602    /// state `IlpProgramFactory.compile(source, ...)` would produce — the
603    /// frontend runs again, the facts are loaded again and the plan is
604    /// executed once. What is skipped is only the per-compile CUDA provider
605    /// setup (a fresh context, allocator and memory pool). Intended for ILP
606    /// loops that compile many one-rule variants of one program (hold-out
607    /// folds, candidate scans).
608    ///
609    /// Because the provider is shared, so are its device-memory budget
610    /// (`memory_mb` given to the base compile covers this program and every
611    /// variant alive at the same time), its streams and its host-transfer
612    /// counters — `variant.fact_exists()` bumps the counter
613    /// `base.d2h_transfer_count()` reads. Relation overrides uploaded to the
614    /// base with `put_relation` are NOT carried over: the variant sees only
615    /// the facts in `source`. The variant is compiled with the same
616    /// `max_active_rules` as this program.
617    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    /// Wall-clock breakdown of the compile that produced this program, in
628    /// milliseconds per phase (`provider`, `frontend`, `facts`, `execute`).
629    /// A variant has no `provider` entry. Phase names are diagnostic and may
630    /// change; do not build on the exact set.
631    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    /// Upload candidate (i,j,k) -> index mapping. Called once per attempt.
639    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    /// Length of current candidate map (0 if not set).
650    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    /// Set the per-chunk temp allocation budget in bytes. The final merged
666    /// COO buffer is exact-NNZ sized and may exceed this budget. Default: 16 MB.
667    pub fn set_coo_chunk_budget(&mut self, bytes: u64) {
668        self.coo_chunk_budget = bytes;
669    }
670
671    /// Deprecated: use `set_coo_chunk_budget`. Kept for one release cycle.
672    #[allow(deprecated)]
673    pub fn set_coo_memory_cap(&mut self, bytes: u64) {
674        self.coo_chunk_budget = bytes;
675    }
676
677    /// Enable strict zero-D2H mode. When true, raises RuntimeError instead
678    /// of falling back to the chunked COO path (which uses D2H transfers).
679    /// Use for zero-D2H benchmarks and CI gates.
680    pub fn set_strict_zero_dtoh(&mut self, strict: bool) {
681        self.strict_zero_dtoh = strict;
682    }
683
684    /// Upload a named relation as DLPack tensor columns into the ILP program store.
685    ///
686    /// This enables tensor-native ILP data upload: compile a schema-only source
687    /// (predicate declarations only, no facts), then upload GPU tensor data via
688    /// DLPack without host materialization.
689    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    /// GPU-resident ILP loss + gradient computation.
724    ///
725    /// Builds a sparse CSR structure from retained per-entry membership masks
726    /// and launches forward (credit gather + NLL loss) and backward (gradient
727    /// scatter) CUDA kernels.  Returns `(loss_dlpack, grad_dlpack)` where both
728    /// are GPU-resident tensors exported via DLPack.
729    ///
730    /// # Arguments
731    /// * `positives` — list of `(relation, [col_values])` positive examples
732    /// * `negatives` — list of `(relation, [col_values])` negative examples
733    /// * `cand_probs_obj` — DLPack/PyTorch tensor of candidate probabilities on GPU
734    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        // Validate inputs and import candidate probabilities via DLPack.
742        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        // Determine dtype from imported buffer
756        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        // Build the fact list and group examples by relation.
781        let num_pos = positives.len();
782        let num_neg = negatives.len();
783        let num_facts = (num_pos + num_neg) as u32;
784
785        // Build all_facts: (relation, values, is_positive, global_fact_idx)
786        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        // Handle empty facts edge case: return zero loss + zero grad
808        if num_facts == 0 {
809            return self.build_zero_loss_grad(py, num_cands, is_f64);
810        }
811
812        // Build is_positive array for upload
813        let is_positive_host: Vec<u8> = all_facts
814            .iter()
815            .map(|f| if f.is_positive { 1u8 } else { 0u8 })
816            .collect();
817
818        // Group facts by relation, preserving global fact_idx
819        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        // Get ILP tagged result
832        let tagged = self
833            .executor
834            .ilp_last_result()
835            .ok_or_else(|| PyRuntimeError::new_err("No ILP result — call evaluate() first"))?;
836
837        // Build the COO structure on the device without D2H reads.
838        //
839        // Two-pass approach over (relation, candidate) tasks:
840        //   Pass 1: compute GPU membership masks
841        //   Pass 2: scatter COO entries at host-computed offsets via device kernel
842        //
843        // The key insight is that each task's num_query is known on the host
844        // (it equals the number of facts in that relation group), so we can
845        // compute COO write offsets entirely on the host without any D2H reads.
846        // We over-allocate COO arrays at upper_bound (sum of all num_query)
847        // and fill sentinel values for unused slots.
848
849        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            // Upload global fact indices for this relation group (H2D, allowed).
894            // Shared across all entries in this relation group via index.
895            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        // Check if COO allocation would exceed memory cap.
946        // Each COO entry uses 8 bytes (4 for fact_idx + 4 for cand_idx).
947        let coo_bytes = (upper_bound as u64) * 8;
948        let needs_chunking = coo_bytes > self.coo_chunk_budget;
949
950        // Upload is_positive once (shared across all paths, H2D allowed)
951        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        // Construct COO entries.
967        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        // Sort COO entries and build CSR row offsets on the device.
998        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        // Run forward, backward, and device-side reduction kernels.
1007        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    /// GPU-resident ILP loss + gradient computation from relation-native
1021    /// positive/negative examples stored as DLPack column groups.
1022    ///
1023    /// `positives_by_relation` and `negatives_by_relation` must be Python dicts
1024    /// mapping relation name -> sequence of DLPack columns with schema-matching
1025    /// dtypes. No host fact tuple materialization occurs on this path.
1026    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    /// Sparse mask API: candidate IDs + DLPack soft probabilities (GPU tensor).
1288    ///
1289    /// `candidate_ids` must be exactly `[0..C)` where C is the candidate count
1290    /// for this mask under the provided recursion policy.
1291    ///
1292    /// `soft_probs_dlpack` is a DLPack capsule (CUDA f64 tensor) passed from PyTorch.
1293    /// Rust imports it zero-copy, downloads values (not counted by the D2H counter),
1294    /// performs deterministic top-k (desc soft value, then lower id), and
1295    /// stores a sparse IlpMask (no dense N^3 materialization).
1296    #[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        // Import DLPack tensor (zero-copy, stays on GPU)
1344        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        // Download f64 values (control-plane, NOT tracked by D2H counter)
1351        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        // Convert to f32 for top-k ranking in insert_mask_from_sparse
1374        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    /// Preferred sparse mask API: caller preselects candidate IDs and passes
1383    /// only the selected subset plus aligned soft probabilities.
1384    ///
1385    /// This path avoids the full-vector soft-probability download in
1386    /// `set_rule_mask_sparse(...)`. The selected IDs are mapped directly to the
1387    /// existing candidate triples and stored in sparse-mask order.
1388    #[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    /// Strict sparse mask API: selected candidate IDs stay as a device tensor on
1463    /// the Python side. Rust resolves them against the fixed candidate order
1464    /// uploaded via `set_candidate_map(...)`.
1465    #[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    /// Reset mutable runtime state for ILP attempt reuse.
1493    ///
1494    /// Clears ILP registry (masks/tagged results), executor store,
1495    /// join index cache, stats, and profiler. Then re-registers schemas
1496    /// with empty buffers, reloads base facts from AST, and re-executes
1497    /// the plan. Preserves all immutable compile artifacts (AST, plan,
1498    /// schemas, rel_index, provider, TMJ metadata, max_active_rules).
1499    ///
1500    /// After reset, the program is in the same state as a fresh compile()
1501    /// with the same source — ready for set_rule_mask / evaluate cycles.
1502    pub fn reset_runtime(&mut self, py: Python<'_>) -> PyResult<()> {
1503        let _result: xlog_core::Result<()> = py.detach(|| {
1504            // 1. Clear all mutable state (ILP registry, store, caches, stats)
1505            self.executor.reset_for_ilp();
1506
1507            // 2. Re-register schemas with empty buffers
1508            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            // 3. Reload base facts from preserved AST
1514            load_facts_into_store(&self.ast, &self.provider, &mut self.executor, &self.schemas)?;
1515
1516            // 3b. Reapply persistent relation uploads on top of AST/base facts.
1517            self.apply_relation_overrides()?;
1518
1519            // 4. Re-execute plan (populates derived relations)
1520            self.executor.execute_plan(&self.plan)?;
1521
1522            Ok(())
1523        });
1524        _result.map_err(types::xlog_err)?;
1525
1526        // 5. Reset D2H transfer counter
1527        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    /// Return all facts in the named relation as a list of int lists.
1558    #[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        // Download all columns (reuse fact_exists_in_buffer pattern)
1576        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    /// Sample up to `max_n` derived facts for `head_rel` that are NOT in `exclude`.
1637    ///
1638    /// Returns `list[list[int]]` — each inner list is a tuple of column values.
1639    /// Uses the same column-download pattern as `relation_facts`.
1640    #[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        // Build exclude set: only consider tuples for the requested relation
1652        let exclude_set: HashSet<Vec<i64>> = exclude
1653            .into_iter()
1654            .filter(|(rel, _)| rel == &head_rel)
1655            .map(|(_, vals)| vals)
1656            .collect();
1657
1658        // Download all facts using the same pattern as relation_facts
1659        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        // Filter out excluded tuples and cap at max_n
1720        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            // Arity guard: same as executor (skip if join keys exceed columns)
1777            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            // Apply head_projection: check projected columns against values,
1795            // not the raw join output (which has more columns than the head).
1796            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    /// Return declared predicate types from source `pred` declarations.
1825    ///
1826    /// Output is a list of `(name, types)` tuples so callers can
1827    /// deterministically inspect whether metadata is available for
1828    /// relations used during promotion.
1829    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    /// Return the set of valid (i,j,k) candidates for the given learnable mask.
1845    ///
1846    /// Pruning rules:
1847    /// - k must be the head relation for this mask
1848    /// - At least one of (i,j) must have nonzero tuples in the store
1849    /// - Template+template body pairs (both have zero tuples) are pruned
1850    /// - If allow_recursive is false: i==k_head or j==k_head are pruned
1851    ///   (unless head already has base facts)
1852    ///
1853    /// Returns list of dicts: [{id, i, j, k, left_name, right_name, head_name}]
1854    /// IDs assigned 0..C-1 after sorting by (k, i, j) ascending.
1855    #[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    /// GPU-side batch fact membership check.
1948    /// Uploads `facts` (list of value-lists) to a temporary CudaBuffer,
1949    /// semi-joins against the named relation, returns per-fact boolean mask.
1950    /// Zero download_column_* calls — only downloads the u8 mask.
1951    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        // Schema-aware typed upload
2023        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        // All columns are keys (full-tuple match)
2031        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    /// GPU-side batch credit assignment.
2039    ///
2040    /// For each fact in `facts`, returns the list of (i,j,k) entries whose
2041    /// join result contains that fact. Uses membership_mask against retained
2042    /// per-entry buffers — zero download_column_* calls.
2043    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        // Find k index for this relation
2057        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        // Filter entries to those matching target relation k, with retained buffers
2071        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        // Determine arity from the first entry's buffer
2082        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        // Schema-aware typed upload
2092        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        // For each relevant entry, compute membership mask against query facts
2103        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    /// GPU-side batch credit assignment.
2126    ///
2127    /// Returns a CSR-style device representation:
2128    /// - `fact_row_offsets`: len = num_facts + 1
2129    /// - `entry_indices`: COO candidate indices, sorted by fact row
2130    /// - `entry_i/j/k`: metadata arrays indexed by `entry_indices`
2131    ///
2132    /// Zero DTOH calls on the query path. Uses the non-chunked COO builder
2133    /// to avoid reading device-side nnz metadata back to the host.
2134    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
2334// ---------------------------------------------------------------------------
2335// CompiledIlpProgram — plain impl block (GPU export helpers + internal)
2336// ---------------------------------------------------------------------------
2337
2338impl 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    // ─── GPU loss/grad export helpers ──────────────────────────────────
2490
2491    /// Build zero loss (scalar) + zero grad (num_cands) on GPU, export via DLPack.
2492    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    /// Handle the case where COO is empty but we have facts.
2506    /// Facts with no covering entries get: positive -> -log(eps), negative -> -log(1.0) = 0.
2507    /// We still run the kernels with an empty CSR for correctness.
2508    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        // Build CSR with all-zero row_offsets (every row has 0 non-zeros)
2517        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        // Empty col_indices
2530        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        // Build a single-column CudaBuffer with 0 elements to represent an empty cand_probs
2548        // for the kernel launch (won't be read since row ranges are all empty).
2549        // We need a dummy CudaColumn. Use the actual cand count = num_cands.
2550        // Since COO is empty the kernel won't access cand_probs, but we need to pass something.
2551        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    /// Returns sorted (i,j,k) candidate triples for the given learnable mask.
2581    /// Pruning logic must stay aligned with `valid_candidates`.
2582    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        // Identify which relations currently have nonzero tuples in store.
2608        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                // Prune template+template (both no tuples).
2626                if !has_tuples[i as usize] && !has_tuples[j as usize] {
2627                    continue;
2628                }
2629
2630                // Keep behavior aligned with existing alpha candidate pruning:
2631                // recursive body refs are allowed only if head already has tuples.
2632                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    /// Like fact_exists_in_buffer but checks only the projected columns.
2716    /// `projection[i]` is the column index in `buf` that corresponds to
2717    /// `values[i]` in the head relation.
2718    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}