Skip to main content

xlog_cuda/provider/
ilp_exact.rs

1//! Launcher for the native bounded exact-induction scoring kernel.
2//!
3//! Drives `kernels/ilp_exact.cu`'s `ilp_exact_score` kernel: scores all
4//! `(topology, L, R)` triples for a single `induce_exact` call in one
5//! launch and returns the positive/negative coverage count arrays to host.
6//!
7//! Design: `docs/plans/2026-04-17-m8-ilp-exact-kernel-design.md`.
8
9use std::marker::PhantomData;
10use std::sync::atomic::Ordering;
11
12use crate::{LaunchAsync, LaunchConfig};
13use xlog_core::{Result, ScalarType, XlogError};
14
15use super::{ilp_exact_kernels, RawCudaView, ILP_EXACT_MODULE};
16use crate::memory::{CudaBuffer, TrackedCudaSlice};
17
18const ILP_EXACT_BLOCK_SIZE: u32 = 256;
19const ILP_EXACT_TOPK_FIELDS: usize = 9;
20const ENV_ILP_EXACT_CHAIN_SMEM: &str = "XLOG_ILP_EXACT_CHAIN_SMEM";
21const ENV_ILP_EXACT_CHAIN_SMEM_MIN_ROWS: &str = "XLOG_ILP_EXACT_CHAIN_SMEM_MIN_ROWS";
22const DEFAULT_ILP_EXACT_CHAIN_SMEM_MIN_ROWS: u32 = 256;
23
24#[derive(Clone, Copy, Debug, PartialEq, Eq)]
25pub struct IlpExactTopkCandidate {
26    pub topology_idx: u32,
27    pub left_idx: u32,
28    pub right_idx: u32,
29    pub positives_covered: u32,
30    pub negatives_covered: u32,
31    pub local_rank: u32,
32    pub next_positives_covered: u32,
33    pub next_negatives_covered: u32,
34    pub tie_class_size: u32,
35}
36
37struct IlpExactDeviceScores {
38    candidate_count: usize,
39    #[cfg(test)]
40    slot_count: usize,
41    pos_covered: TrackedCudaSlice<u32>,
42    neg_covered: TrackedCudaSlice<u32>,
43}
44
45#[derive(Clone, Copy, Debug, Eq, PartialEq)]
46enum ExactPairLayout {
47    U64,
48    U32,
49    Symbol,
50}
51
52impl ExactPairLayout {
53    fn elem_size(self) -> usize {
54        match self {
55            Self::U64 => std::mem::size_of::<u64>(),
56            Self::U32 | Self::Symbol => std::mem::size_of::<u32>(),
57        }
58    }
59}
60
61fn ilp_exact_chain_smem_enabled() -> bool {
62    match std::env::var(ENV_ILP_EXACT_CHAIN_SMEM) {
63        Ok(value) => !matches!(
64            value.trim().to_ascii_lowercase().as_str(),
65            "0" | "false" | "off" | "no"
66        ),
67        Err(_) => true,
68    }
69}
70
71fn chain_smem_shared_bytes(layout: ExactPairLayout) -> u32 {
72    let block = ILP_EXACT_BLOCK_SIZE as usize;
73    let bytes = (2usize * block * layout.elem_size()) + (block * std::mem::size_of::<u32>());
74    u32::try_from(bytes).expect("chain smem byte count fits in u32")
75}
76
77fn ilp_exact_chain_smem_min_rows() -> u32 {
78    std::env::var(ENV_ILP_EXACT_CHAIN_SMEM_MIN_ROWS)
79        .ok()
80        .and_then(|value| value.trim().parse::<u32>().ok())
81        .unwrap_or(DEFAULT_ILP_EXACT_CHAIN_SMEM_MIN_ROWS)
82}
83
84impl super::CudaKernelProvider {
85    /// Test-only full-score export for validating the scoring kernels.
86    ///
87    /// Returns `(pos_covered, neg_covered)`, each of length `4 * C * C`
88    /// where `C = candidate_buffers.len()`. Slot ordering:
89    /// `slot = topology * (C * C) + L * C + R`, with topology indices
90    /// `chain=0, star=1, fanout=2, fanin=3`.
91    ///
92    /// Host-side contract:
93    ///   * All buffers must be arity 2 with one matching pair type: `U64`,
94    ///     `U32`, or `Symbol`.
95    ///   * `cached_row_count()` must be populated on every buffer (DLPack
96    ///     ingest and `create_empty_buffer` both guarantee this).
97    ///   * `negatives` is always a valid buffer — the caller constructs
98    ///     an empty pair buffer matching the positive pair type when there are
99    ///     no negatives.
100    ///
101    /// D2H budget: **2** counter-tracked transfers (one per count array).
102    /// Setup H2D / D2D copies are not D2H-counted.
103    #[cfg(test)]
104    fn ilp_exact_score(
105        &self,
106        candidate_buffers: &[&CudaBuffer],
107        positives: &CudaBuffer,
108        negatives: &CudaBuffer,
109    ) -> Result<(Vec<u32>, Vec<u32>)> {
110        let scores = self.ilp_exact_score_device(candidate_buffers, positives, negatives)?;
111        let device = self.device.inner();
112        self.device.synchronize()?;
113
114        let mut pos_covered = vec![0u32; scores.slot_count];
115        self.d2h_transfer_count.fetch_add(1, Ordering::Relaxed);
116        device
117            .dtoh_sync_copy_into(&scores.pos_covered, &mut pos_covered)
118            .map_err(|e| XlogError::Kernel(format!("ilp_exact_score: dtoh pos_covered: {}", e)))?;
119
120        let mut neg_covered = vec![0u32; scores.slot_count];
121        self.d2h_transfer_count.fetch_add(1, Ordering::Relaxed);
122        device
123            .dtoh_sync_copy_into(&scores.neg_covered, &mut neg_covered)
124            .map_err(|e| XlogError::Kernel(format!("ilp_exact_score: dtoh neg_covered: {}", e)))?;
125
126        Ok((pos_covered, neg_covered))
127    }
128
129    /// Score on GPU, reduce per-topology top-K on GPU, and transfer only the
130    /// compact selected rows back to host.
131    pub fn ilp_exact_score_topk(
132        &self,
133        candidate_buffers: &[&CudaBuffer],
134        positives: &CudaBuffer,
135        negatives: &CudaBuffer,
136        k_per_topology: u32,
137    ) -> Result<Vec<IlpExactTopkCandidate>> {
138        if k_per_topology == 0 {
139            return Ok(Vec::new());
140        }
141
142        let scores = self.ilp_exact_score_device(candidate_buffers, positives, negatives)?;
143        let out_rows = 4usize
144            .checked_mul(k_per_topology as usize)
145            .ok_or_else(|| XlogError::Kernel("ilp_exact_score_topk: output row overflow".into()))?;
146        let out_words = out_rows.checked_mul(ILP_EXACT_TOPK_FIELDS).ok_or_else(|| {
147            XlogError::Kernel("ilp_exact_score_topk: output word overflow".into())
148        })?;
149        let mut selected_buf = self.memory.alloc::<u32>(out_words)?;
150        let device = self.device.inner();
151        let func = device
152            .get_func(ILP_EXACT_MODULE, ilp_exact_kernels::ILP_EXACT_SELECT_TOPK)
153            .ok_or_else(|| {
154                XlogError::Kernel(format!(
155                    "{} kernel not loaded",
156                    ilp_exact_kernels::ILP_EXACT_SELECT_TOPK
157                ))
158            })?;
159
160        unsafe {
161            func.clone().launch(
162                LaunchConfig {
163                    grid_dim: (4, 1, 1),
164                    block_dim: (1, 1, 1),
165                    shared_mem_bytes: 0,
166                },
167                (
168                    &scores.pos_covered,
169                    &scores.neg_covered,
170                    scores.candidate_count as u32,
171                    k_per_topology,
172                    &mut selected_buf,
173                ),
174            )
175        }
176        .map_err(|e| XlogError::Kernel(format!("ilp_exact_select_topk launch: {}", e)))?;
177
178        self.device.synchronize()?;
179        let mut words = vec![0u32; out_words];
180        self.d2h_transfer_count.fetch_add(1, Ordering::Relaxed);
181        device
182            .dtoh_sync_copy_into(&selected_buf, &mut words)
183            .map_err(|e| {
184                XlogError::Kernel(format!("ilp_exact_score_topk: dtoh selected: {}", e))
185            })?;
186
187        let mut selected = Vec::new();
188        for chunk in words.chunks_exact(ILP_EXACT_TOPK_FIELDS) {
189            if chunk[3] == 0 {
190                continue;
191            }
192            selected.push(IlpExactTopkCandidate {
193                topology_idx: chunk[0],
194                left_idx: chunk[1],
195                right_idx: chunk[2],
196                positives_covered: chunk[3],
197                negatives_covered: chunk[4],
198                local_rank: chunk[5],
199                next_positives_covered: chunk[6],
200                next_negatives_covered: chunk[7],
201                tie_class_size: chunk[8],
202            });
203        }
204        Ok(selected)
205    }
206
207    fn ilp_exact_score_device(
208        &self,
209        candidate_buffers: &[&CudaBuffer],
210        positives: &CudaBuffer,
211        negatives: &CudaBuffer,
212    ) -> Result<IlpExactDeviceScores> {
213        let c = candidate_buffers.len();
214        if c == 0 {
215            return Err(XlogError::Kernel(
216                "ilp_exact_score: candidate list is empty (filter at the engine)".to_string(),
217            ));
218        }
219        let c_u32 = u32::try_from(c).map_err(|_| {
220            XlogError::Kernel(format!(
221                "ilp_exact_score: candidate count {} exceeds u32::MAX",
222                c
223            ))
224        })?;
225
226        // ── Validate shapes and gather host-side row counts ────────────────
227        let layout = validate_exact_pair_buffer(positives, "positives")?;
228        require_exact_pair_layout(negatives, "negatives", layout)?;
229        let pos_rows = cached_rows(positives, "positives")?;
230        let neg_rows = cached_rows(negatives, "negatives")?;
231
232        let mut cand_rows: Vec<u32> = Vec::with_capacity(c);
233        for (i, buf) in candidate_buffers.iter().enumerate() {
234            let label = format!("candidate[{}]", i);
235            require_exact_pair_layout(buf, &label, layout)?;
236            cand_rows.push(cached_rows(buf, &label)?);
237        }
238
239        // ── Exclusive prefix sum of row counts (cand_offsets, length C+1) ─
240        let mut cand_offsets_host: Vec<u32> = Vec::with_capacity(c + 1);
241        let mut running: u32 = 0;
242        cand_offsets_host.push(0);
243        for &r in &cand_rows {
244            running = running.checked_add(r).ok_or_else(|| {
245                XlogError::Kernel("ilp_exact_score: candidate row count overflow u32".to_string())
246            })?;
247            cand_offsets_host.push(running);
248        }
249        let total_rows = running as usize;
250        let elem_size = layout.elem_size();
251        let total_bytes = total_rows * elem_size;
252
253        let device = self.device.inner();
254
255        // ── Concatenate candidate columns via D2D copies ──────────────────
256        // Setup-phase D→D; neither counted by the D2H gate nor by the
257        // transfer tracker as a host-to-device round trip.
258        let mut cand_arg0_buf = self.memory.alloc::<u8>(total_bytes)?;
259        let mut cand_arg1_buf = self.memory.alloc::<u8>(total_bytes)?;
260        if total_bytes > 0 {
261            let mut byte_offset: usize = 0;
262            for (i, buf) in candidate_buffers.iter().enumerate() {
263                let rows = cand_rows[i] as usize;
264                if rows == 0 {
265                    continue;
266                }
267                let bytes = rows * elem_size;
268
269                let src0 = buf.column(0).ok_or_else(|| {
270                    XlogError::Kernel(format!("candidate[{}] missing column 0", i))
271                })?;
272                let src1 = buf.column(1).ok_or_else(|| {
273                    XlogError::Kernel(format!("candidate[{}] missing column 1", i))
274                })?;
275                let src_view0 = self.column_bytes_view(src0, bytes)?;
276                let src_view1 = self.column_bytes_view(src1, bytes)?;
277                let mut dst0 = cand_arg0_buf.slice_mut(byte_offset..byte_offset + bytes);
278                let mut dst1 = cand_arg1_buf.slice_mut(byte_offset..byte_offset + bytes);
279                device.dtod_copy(&src_view0, &mut dst0).map_err(|e| {
280                    XlogError::Kernel(format!(
281                        "ilp_exact_score: d2d concat arg0 (candidate {}): {}",
282                        i, e
283                    ))
284                })?;
285                device.dtod_copy(&src_view1, &mut dst1).map_err(|e| {
286                    XlogError::Kernel(format!(
287                        "ilp_exact_score: d2d concat arg1 (candidate {}): {}",
288                        i, e
289                    ))
290                })?;
291                byte_offset += bytes;
292            }
293        }
294
295        // ── Upload cand_offsets (H→D, not D2H-counted) ────────────────────
296        let mut cand_offsets_buf = self.memory.alloc::<u32>(c + 1)?;
297        self.htod_sync_copy_into_tracked(&cand_offsets_host, &mut cand_offsets_buf)
298            .map_err(|e| XlogError::Kernel(format!("ilp_exact_score: h2d cand_offsets: {}", e)))?;
299
300        // ── Alloc output count arrays ─────────────────────────────────────
301        let n_slots = 4usize
302            .checked_mul(c)
303            .and_then(|v| v.checked_mul(c))
304            .ok_or_else(|| {
305                XlogError::Kernel("ilp_exact_score: n_slots = 4 * C * C overflow".to_string())
306            })?;
307        let mut pos_covered_buf = self.memory.alloc::<u32>(n_slots)?;
308        let mut neg_covered_buf = self.memory.alloc::<u32>(n_slots)?;
309        // Kernel writes every slot exactly once — no zero-init required.
310
311        let pos_col0 = positives
312            .column(0)
313            .ok_or_else(|| XlogError::Kernel("positives: missing column 0".to_string()))?;
314        let pos_col1 = positives
315            .column(1)
316            .ok_or_else(|| XlogError::Kernel("positives: missing column 1".to_string()))?;
317        let neg_col0 = negatives
318            .column(0)
319            .ok_or_else(|| XlogError::Kernel("negatives: missing column 0".to_string()))?;
320        let neg_col1 = negatives
321            .column(1)
322            .ok_or_else(|| XlogError::Kernel("negatives: missing column 1".to_string()))?;
323
324        // ── Launch ────────────────────────────────────────────────────────
325        let max_candidate_rows = cand_rows.iter().copied().max().unwrap_or(0);
326        let chain_smem_enabled =
327            ilp_exact_chain_smem_enabled() && max_candidate_rows >= ilp_exact_chain_smem_min_rows();
328        let shared_mem_bytes = if chain_smem_enabled {
329            chain_smem_shared_bytes(layout)
330        } else {
331            0
332        };
333        match layout {
334            ExactPairLayout::U64 => {
335                let cand_arg0_view = RawCudaView::<u64> {
336                    ptr: *cand_arg0_buf.device_ptr(),
337                    len: total_rows,
338                    stream: cand_arg0_buf.stream().clone(),
339                    _marker: PhantomData,
340                };
341                let cand_arg1_view = RawCudaView::<u64> {
342                    ptr: *cand_arg1_buf.device_ptr(),
343                    len: total_rows,
344                    stream: cand_arg1_buf.stream().clone(),
345                    _marker: PhantomData,
346                };
347                let pos_arg0_view = self.column_as_u64_view(pos_col0, pos_rows as usize)?;
348                let pos_arg1_view = self.column_as_u64_view(pos_col1, pos_rows as usize)?;
349                let neg_arg0_view = self.column_as_u64_view(neg_col0, neg_rows as usize)?;
350                let neg_arg1_view = self.column_as_u64_view(neg_col1, neg_rows as usize)?;
351                let kernel_name = if chain_smem_enabled {
352                    ilp_exact_kernels::ILP_EXACT_SCORE_CHAIN_SMEM
353                } else {
354                    ilp_exact_kernels::ILP_EXACT_SCORE
355                };
356                let func = device
357                    .get_func(ILP_EXACT_MODULE, kernel_name)
358                    .ok_or_else(|| {
359                        XlogError::Kernel(format!("{} kernel not loaded", kernel_name))
360                    })?;
361                unsafe {
362                    func.clone().launch(
363                        LaunchConfig {
364                            grid_dim: (c_u32, c_u32, 4),
365                            block_dim: (ILP_EXACT_BLOCK_SIZE, 1, 1),
366                            shared_mem_bytes,
367                        },
368                        (
369                            &cand_arg0_view,
370                            &cand_arg1_view,
371                            &cand_offsets_buf,
372                            c_u32,
373                            &pos_arg0_view,
374                            &pos_arg1_view,
375                            pos_rows,
376                            &neg_arg0_view,
377                            &neg_arg1_view,
378                            neg_rows,
379                            &mut pos_covered_buf,
380                            &mut neg_covered_buf,
381                        ),
382                    )
383                }
384                .map_err(|e| XlogError::Kernel(format!("ilp_exact_score launch: {}", e)))?;
385            }
386            ExactPairLayout::U32 | ExactPairLayout::Symbol => {
387                let cand_arg0_view = RawCudaView::<u32> {
388                    ptr: *cand_arg0_buf.device_ptr(),
389                    len: total_rows,
390                    stream: cand_arg0_buf.stream().clone(),
391                    _marker: PhantomData,
392                };
393                let cand_arg1_view = RawCudaView::<u32> {
394                    ptr: *cand_arg1_buf.device_ptr(),
395                    len: total_rows,
396                    stream: cand_arg1_buf.stream().clone(),
397                    _marker: PhantomData,
398                };
399                let pos_arg0_view = self.column_as_u32_view(pos_col0, pos_rows as usize)?;
400                let pos_arg1_view = self.column_as_u32_view(pos_col1, pos_rows as usize)?;
401                let neg_arg0_view = self.column_as_u32_view(neg_col0, neg_rows as usize)?;
402                let neg_arg1_view = self.column_as_u32_view(neg_col1, neg_rows as usize)?;
403                let kernel_name = if chain_smem_enabled {
404                    ilp_exact_kernels::ILP_EXACT_SCORE_CHAIN_SMEM_U32
405                } else {
406                    ilp_exact_kernels::ILP_EXACT_SCORE_U32
407                };
408                let func = device
409                    .get_func(ILP_EXACT_MODULE, kernel_name)
410                    .ok_or_else(|| {
411                        XlogError::Kernel(format!("{} kernel not loaded", kernel_name))
412                    })?;
413                unsafe {
414                    func.clone().launch(
415                        LaunchConfig {
416                            grid_dim: (c_u32, c_u32, 4),
417                            block_dim: (ILP_EXACT_BLOCK_SIZE, 1, 1),
418                            shared_mem_bytes,
419                        },
420                        (
421                            &cand_arg0_view,
422                            &cand_arg1_view,
423                            &cand_offsets_buf,
424                            c_u32,
425                            &pos_arg0_view,
426                            &pos_arg1_view,
427                            pos_rows,
428                            &neg_arg0_view,
429                            &neg_arg1_view,
430                            neg_rows,
431                            &mut pos_covered_buf,
432                            &mut neg_covered_buf,
433                        ),
434                    )
435                }
436                .map_err(|e| XlogError::Kernel(format!("ilp_exact_score_u32 launch: {}", e)))?;
437            }
438        }
439
440        Ok(IlpExactDeviceScores {
441            candidate_count: c,
442            #[cfg(test)]
443            slot_count: n_slots,
444            pos_covered: pos_covered_buf,
445            neg_covered: neg_covered_buf,
446        })
447    }
448}
449
450fn validate_exact_pair_buffer(buf: &CudaBuffer, label: &str) -> Result<ExactPairLayout> {
451    if buf.arity() != 2 {
452        return Err(XlogError::Kernel(format!(
453            "ilp_exact_score: {} buffer arity = {}, expected 2",
454            label,
455            buf.arity(),
456        )));
457    }
458    let mut layout: Option<ExactPairLayout> = None;
459    for col_idx in 0..2 {
460        let t = buf.schema().column_type(col_idx).ok_or_else(|| {
461            XlogError::Kernel(format!(
462                "ilp_exact_score: {} buffer missing column {} type",
463                label, col_idx,
464            ))
465        })?;
466        let col_layout = match t {
467            ScalarType::U64 => ExactPairLayout::U64,
468            ScalarType::U32 => ExactPairLayout::U32,
469            ScalarType::Symbol => ExactPairLayout::Symbol,
470            _ => {
471                return Err(XlogError::Kernel(format!(
472                    "ilp_exact_score: {} buffer column {} type = {:?}, expected U64, U32, or Symbol",
473                    label, col_idx, t,
474                )));
475            }
476        };
477        if let Some(expected) = layout {
478            if expected != col_layout {
479                return Err(XlogError::Kernel(format!(
480                    "ilp_exact_score: {} buffer column {} type mismatch: {:?} vs {:?}",
481                    label, col_idx, expected, col_layout,
482                )));
483            }
484        } else {
485            layout = Some(col_layout);
486        }
487    }
488    Ok(layout.expect("arity 2 loop sets layout"))
489}
490
491fn require_exact_pair_layout(
492    buf: &CudaBuffer,
493    label: &str,
494    expected: ExactPairLayout,
495) -> Result<()> {
496    let actual = validate_exact_pair_buffer(buf, label)?;
497    if actual != expected {
498        return Err(XlogError::Kernel(format!(
499            "ilp_exact_score: {} buffer type mismatch: expected {:?}, got {:?}",
500            label, expected, actual,
501        )));
502    }
503    Ok(())
504}
505
506fn cached_rows(buf: &CudaBuffer, label: &str) -> Result<u32> {
507    buf.cached_row_count().ok_or_else(|| {
508        XlogError::Kernel(format!(
509            "ilp_exact_score: {} buffer has no cached row count \
510             (DLPack ingest and create_empty_buffer both populate it)",
511            label
512        ))
513    })
514}
515
516#[cfg(test)]
517mod tests {
518    //! CUDA-gated correctness tests for the ilp_exact launcher.
519    //!
520    //! Pinned to a hand-computed fixture so the kernel's coverage arithmetic
521    //! can be verified without relying on the Python backend as oracle. The
522    //! fixture uses C=2 candidate relations so the expected flat output
523    //! (4 × C × C = 16 slots per count array) is tractable to enumerate.
524
525    use std::sync::Arc;
526
527    use xlog_core::{MemoryBudget, ScalarType, Schema};
528
529    use crate::{CudaDevice, CudaKernelProvider, GpuMemoryManager};
530
531    fn make_provider() -> Option<CudaKernelProvider> {
532        let device = Arc::new(CudaDevice::new(0).ok()?);
533        let budget = MemoryBudget::with_limit(1024 * 1024 * 1024);
534        let memory = Arc::new(GpuMemoryManager::new(device.clone(), budget));
535        CudaKernelProvider::new(device, memory).ok()
536    }
537
538    /// Build a `(u64, u64)` pair buffer from parallel host-side column arrays.
539    /// Uses `create_buffer_from_slice` per column then recombines, relying on
540    /// the provider's buffer-from-columns path to set the cached row count.
541    fn pair_buffer(provider: &CudaKernelProvider, arg0: &[u64], arg1: &[u64]) -> crate::CudaBuffer {
542        assert_eq!(arg0.len(), arg1.len());
543        let schema = Schema::new(vec![
544            ("arg0".to_string(), ScalarType::U64),
545            ("arg1".to_string(), ScalarType::U64),
546        ]);
547        if arg0.is_empty() {
548            return provider
549                .create_empty_buffer(schema)
550                .expect("empty pair buffer");
551        }
552        // Pack both columns as a single 2-column buffer by constructing
553        // byte-columns manually — mirrors what `from_dlpack_tensors_with_schema`
554        // does for the in-process launcher tests.
555        let device = provider.device().inner();
556        let arg0_bytes: Vec<u8> = arg0.iter().flat_map(|v| v.to_le_bytes()).collect();
557        let arg1_bytes: Vec<u8> = arg1.iter().flat_map(|v| v.to_le_bytes()).collect();
558        let mut col0 = provider
559            .memory()
560            .alloc::<u8>(arg0_bytes.len())
561            .expect("alloc");
562        let mut col1 = provider
563            .memory()
564            .alloc::<u8>(arg1_bytes.len())
565            .expect("alloc");
566        device
567            .htod_sync_copy_into(&arg0_bytes, &mut col0)
568            .expect("h2d arg0");
569        device
570            .htod_sync_copy_into(&arg1_bytes, &mut col1)
571            .expect("h2d arg1");
572        provider
573            .buffer_from_columns(vec![col0.into(), col1.into()], arg0.len() as u64, schema)
574            .expect("buffer_from_columns")
575    }
576
577    fn pair_buffer_u32(
578        provider: &CudaKernelProvider,
579        arg0: &[u32],
580        arg1: &[u32],
581        typ: ScalarType,
582    ) -> crate::CudaBuffer {
583        assert_eq!(arg0.len(), arg1.len());
584        assert!(matches!(typ, ScalarType::U32 | ScalarType::Symbol));
585        let schema = Schema::new(vec![("arg0".to_string(), typ), ("arg1".to_string(), typ)]);
586        if arg0.is_empty() {
587            return provider
588                .create_empty_buffer(schema)
589                .expect("empty pair buffer");
590        }
591        let device = provider.device().inner();
592        let arg0_bytes: Vec<u8> = arg0.iter().flat_map(|v| v.to_le_bytes()).collect();
593        let arg1_bytes: Vec<u8> = arg1.iter().flat_map(|v| v.to_le_bytes()).collect();
594        let mut col0 = provider
595            .memory()
596            .alloc::<u8>(arg0_bytes.len())
597            .expect("alloc");
598        let mut col1 = provider
599            .memory()
600            .alloc::<u8>(arg1_bytes.len())
601            .expect("alloc");
602        device
603            .htod_sync_copy_into(&arg0_bytes, &mut col0)
604            .expect("h2d arg0");
605        device
606            .htod_sync_copy_into(&arg1_bytes, &mut col1)
607            .expect("h2d arg1");
608        provider
609            .buffer_from_columns(vec![col0.into(), col1.into()], arg0.len() as u64, schema)
610            .expect("buffer_from_columns")
611    }
612
613    fn pair_buffer_i32(
614        provider: &CudaKernelProvider,
615        arg0: &[i32],
616        arg1: &[i32],
617    ) -> crate::CudaBuffer {
618        assert_eq!(arg0.len(), arg1.len());
619        let schema = Schema::new(vec![
620            ("arg0".to_string(), ScalarType::I32),
621            ("arg1".to_string(), ScalarType::I32),
622        ]);
623        if arg0.is_empty() {
624            return provider
625                .create_empty_buffer(schema)
626                .expect("empty pair buffer");
627        }
628        let device = provider.device().inner();
629        let arg0_bytes: Vec<u8> = arg0.iter().flat_map(|v| v.to_le_bytes()).collect();
630        let arg1_bytes: Vec<u8> = arg1.iter().flat_map(|v| v.to_le_bytes()).collect();
631        let mut col0 = provider
632            .memory()
633            .alloc::<u8>(arg0_bytes.len())
634            .expect("alloc");
635        let mut col1 = provider
636            .memory()
637            .alloc::<u8>(arg1_bytes.len())
638            .expect("alloc");
639        device
640            .htod_sync_copy_into(&arg0_bytes, &mut col0)
641            .expect("h2d arg0");
642        device
643            .htod_sync_copy_into(&arg1_bytes, &mut col1)
644            .expect("h2d arg1");
645        provider
646            .buffer_from_columns(vec![col0.into(), col1.into()], arg0.len() as u64, schema)
647            .expect("buffer_from_columns")
648    }
649
650    /// Hand-computed coverage for C=2 candidates {p_B, p_C} against positives
651    /// `{(1,4), (2,5)}` and negatives `{(7,8)}`. The only non-zero coverage
652    /// is `chain(p_B, p_C) = 2` (both positives covered via chain joins
653    /// z=2 and z=3). Everything else is zero by direct enumeration of the
654    /// four topology templates — see
655    /// `docs/plans/2026-04-17-m8-ilp-exact-kernel-design.md` for the
656    /// templates. Also exercises the negative-scoring path with one negative
657    /// that no topology-L-R combination covers.
658    #[test]
659    fn ilp_exact_score_matches_hand_computed_fixture() {
660        let provider = match make_provider() {
661            Some(p) => p,
662            None => {
663                eprintln!("Skipping test: no CUDA device available");
664                return;
665            }
666        };
667
668        // Candidate relations.
669        let p_b = pair_buffer(&provider, &[1, 2], &[2, 3]);
670        let p_c = pair_buffer(&provider, &[2, 3, 4], &[4, 5, 6]);
671
672        // Positives: {(1,4), (2,5)}. Negatives: {(7,8)}.
673        let positives = pair_buffer(&provider, &[1, 2], &[4, 5]);
674        let negatives = pair_buffer(&provider, &[7], &[8]);
675
676        let (pos, neg) = provider
677            .ilp_exact_score(&[&p_b, &p_c], &positives, &negatives)
678            .expect("ilp_exact_score launch");
679
680        // Slot layout: topology * C² + L * C + R, with C=2.
681        //   topology: chain=0, star=1, fanout=2, fanin=3.
682        //   L/R: p_B=0, p_C=1.
683        // Only chain(p_B=0, p_C=1) → slot 0*4 + 0*2 + 1 = 1 is non-zero.
684        let mut expected_pos = vec![0u32; 16];
685        expected_pos[1] = 2;
686        assert_eq!(
687            pos, expected_pos,
688            "positives coverage mismatch: expected {:?}, got {:?}",
689            expected_pos, pos,
690        );
691
692        // All negatives coverage slots are zero: no (L, R, topology) covers (7, 8).
693        let expected_neg = vec![0u32; 16];
694        assert_eq!(
695            neg, expected_neg,
696            "negatives coverage mismatch: expected {:?}, got {:?}",
697            expected_neg, neg,
698        );
699    }
700
701    #[test]
702    fn ilp_exact_score_topk_reduces_on_device_to_compact_result() {
703        let provider = match make_provider() {
704            Some(p) => p,
705            None => {
706                eprintln!("Skipping test: no CUDA device available");
707                return;
708            }
709        };
710
711        let p_b = pair_buffer(&provider, &[1, 2], &[2, 3]);
712        let p_c = pair_buffer(&provider, &[2, 3, 4], &[4, 5, 6]);
713        let positives = pair_buffer(&provider, &[1, 2], &[4, 5]);
714        let negatives = pair_buffer(&provider, &[7], &[8]);
715
716        provider.reset_d2h_transfer_count();
717        let selected = provider
718            .ilp_exact_score_topk(&[&p_b, &p_c], &positives, &negatives, 2)
719            .expect("ilp_exact_score_topk launch");
720
721        assert_eq!(provider.d2h_transfer_count(), 1);
722        assert_eq!(selected.len(), 1);
723        let winner = selected[0];
724        assert_eq!(winner.topology_idx, 0);
725        assert_eq!(winner.left_idx, 0);
726        assert_eq!(winner.right_idx, 1);
727        assert_eq!(winner.positives_covered, 2);
728        assert_eq!(winner.negatives_covered, 0);
729        assert_eq!(winner.local_rank, 0);
730        assert_eq!(winner.next_positives_covered, 0);
731        assert_eq!(winner.next_negatives_covered, 0);
732        assert_eq!(winner.tie_class_size, 1);
733    }
734
735    #[test]
736    fn ilp_exact_score_topk_preserves_rank_next_and_tie_diagnostics() {
737        let provider = match make_provider() {
738            Some(p) => p,
739            None => {
740                eprintln!("Skipping test: no CUDA device available");
741                return;
742            }
743        };
744
745        let p_all = pair_buffer(&provider, &[1, 2], &[1, 2]);
746        let p_one = pair_buffer(&provider, &[1], &[1]);
747        let p_two = pair_buffer(&provider, &[2], &[2]);
748        let positives = pair_buffer(&provider, &[1, 2], &[1, 2]);
749        let negatives = pair_buffer(&provider, &[9], &[9]);
750
751        let selected = provider
752            .ilp_exact_score_topk(&[&p_all, &p_one, &p_two], &positives, &negatives, 2)
753            .expect("ilp_exact_score_topk launch");
754
755        let star_rank0 = selected
756            .iter()
757            .find(|row| row.topology_idx == 1 && row.local_rank == 0)
758            .expect("star rank 0");
759        assert_eq!(star_rank0.left_idx, 0);
760        assert_eq!(star_rank0.right_idx, 0);
761        assert_eq!(star_rank0.positives_covered, 2);
762        assert_eq!(star_rank0.negatives_covered, 0);
763        assert_eq!(star_rank0.next_positives_covered, 1);
764        assert_eq!(star_rank0.next_negatives_covered, 0);
765        assert_eq!(star_rank0.tie_class_size, 1);
766
767        let star_rank1 = selected
768            .iter()
769            .find(|row| row.topology_idx == 1 && row.local_rank == 1)
770            .expect("star rank 1");
771        assert_eq!(star_rank1.left_idx, 0);
772        assert_eq!(star_rank1.right_idx, 1);
773        assert_eq!(star_rank1.positives_covered, 1);
774        assert_eq!(star_rank1.negatives_covered, 0);
775        assert_eq!(star_rank1.next_positives_covered, 1);
776        assert_eq!(star_rank1.next_negatives_covered, 0);
777        assert_eq!(star_rank1.tie_class_size, 6);
778    }
779
780    /// Determinism: the same inputs produce identical outputs on repeat runs.
781    /// The kernel relies on integer counts + each block owning one unique
782    /// output slot, so determinism is structural — no associativity or
783    /// floating-point ordering concerns. Still worth pinning as a regression
784    /// guard in case a future change swaps in atomics or shared state.
785    #[test]
786    fn ilp_exact_score_is_deterministic_across_runs() {
787        let provider = match make_provider() {
788            Some(p) => p,
789            None => {
790                eprintln!("Skipping test: no CUDA device available");
791                return;
792            }
793        };
794
795        let p_b = pair_buffer(&provider, &[1, 2], &[2, 3]);
796        let p_c = pair_buffer(&provider, &[2, 3, 4], &[4, 5, 6]);
797        let positives = pair_buffer(&provider, &[1, 2], &[4, 5]);
798        let negatives = pair_buffer(&provider, &[7], &[8]);
799
800        let run_a = provider
801            .ilp_exact_score(&[&p_b, &p_c], &positives, &negatives)
802            .unwrap();
803        let run_b = provider
804            .ilp_exact_score(&[&p_b, &p_c], &positives, &negatives)
805            .unwrap();
806        assert_eq!(run_a.0, run_b.0, "pos coverage drifted across runs");
807        assert_eq!(run_a.1, run_b.1, "neg coverage drifted across runs");
808    }
809
810    /// Empty negatives: when the caller supplies a zero-row negatives buffer
811    /// (the engine's normal treatment of `None`), the kernel must not
812    /// dereference the negative pointers and must leave all `neg_covered`
813    /// slots at zero.
814    #[test]
815    fn ilp_exact_score_handles_empty_negatives() {
816        let provider = match make_provider() {
817            Some(p) => p,
818            None => {
819                eprintln!("Skipping test: no CUDA device available");
820                return;
821            }
822        };
823
824        let p_b = pair_buffer(&provider, &[1, 2], &[2, 3]);
825        let p_c = pair_buffer(&provider, &[2, 3, 4], &[4, 5, 6]);
826        let positives = pair_buffer(&provider, &[1, 2], &[4, 5]);
827        let negatives = pair_buffer(&provider, &[], &[]);
828
829        let (pos, neg) = provider
830            .ilp_exact_score(&[&p_b, &p_c], &positives, &negatives)
831            .unwrap();
832
833        let mut expected_pos = vec![0u32; 16];
834        expected_pos[1] = 2;
835        assert_eq!(pos, expected_pos);
836        assert_eq!(neg, vec![0u32; 16]);
837    }
838
839    #[test]
840    fn ilp_exact_score_accepts_u32_pair_buffers() {
841        let provider = match make_provider() {
842            Some(p) => p,
843            None => {
844                eprintln!("Skipping test: no CUDA device available");
845                return;
846            }
847        };
848
849        let p_b = pair_buffer_u32(&provider, &[1, 2], &[2, 3], ScalarType::U32);
850        let p_c = pair_buffer_u32(&provider, &[2, 3, 4], &[4, 5, 6], ScalarType::U32);
851        let positives = pair_buffer_u32(&provider, &[1, 2], &[4, 5], ScalarType::U32);
852        let negatives = pair_buffer_u32(&provider, &[7], &[8], ScalarType::U32);
853
854        let (pos, neg) = provider
855            .ilp_exact_score(&[&p_b, &p_c], &positives, &negatives)
856            .expect("U32 ilp_exact_score launch");
857
858        let mut expected_pos = vec![0u32; 16];
859        expected_pos[1] = 2;
860        assert_eq!(pos, expected_pos);
861        assert_eq!(neg, vec![0u32; 16]);
862    }
863
864    #[test]
865    fn ilp_exact_score_accepts_symbol_pair_buffers() {
866        let provider = match make_provider() {
867            Some(p) => p,
868            None => {
869                eprintln!("Skipping test: no CUDA device available");
870                return;
871            }
872        };
873
874        let p_b = pair_buffer_u32(&provider, &[1, 2], &[2, 3], ScalarType::Symbol);
875        let p_c = pair_buffer_u32(&provider, &[2, 3, 4], &[4, 5, 6], ScalarType::Symbol);
876        let positives = pair_buffer_u32(&provider, &[1, 2], &[4, 5], ScalarType::Symbol);
877        let negatives = pair_buffer_u32(&provider, &[7], &[8], ScalarType::Symbol);
878
879        let (pos, neg) = provider
880            .ilp_exact_score(&[&p_b, &p_c], &positives, &negatives)
881            .expect("Symbol ilp_exact_score launch");
882
883        let mut expected_pos = vec![0u32; 16];
884        expected_pos[1] = 2;
885        assert_eq!(pos, expected_pos);
886        assert_eq!(neg, vec![0u32; 16]);
887    }
888
889    #[test]
890    fn ilp_exact_score_rejects_mixed_pair_types() {
891        let provider = match make_provider() {
892            Some(p) => p,
893            None => {
894                eprintln!("Skipping test: no CUDA device available");
895                return;
896            }
897        };
898
899        let p_b = pair_buffer_u32(&provider, &[1, 2], &[2, 3], ScalarType::U32);
900        let positives = pair_buffer(&provider, &[1, 2], &[4, 5]);
901        let negatives = pair_buffer(&provider, &[7], &[8]);
902
903        let err = provider
904            .ilp_exact_score(&[&p_b], &positives, &negatives)
905            .expect_err("mixed U64/U32 buffers must be rejected");
906        assert!(
907            err.to_string().contains("expected U64") || err.to_string().contains("type mismatch"),
908            "unexpected error: {err}"
909        );
910    }
911
912    #[test]
913    fn ilp_exact_score_rejects_unsupported_pair_types() {
914        let provider = match make_provider() {
915            Some(p) => p,
916            None => {
917                eprintln!("Skipping test: no CUDA device available");
918                return;
919            }
920        };
921
922        let p_b = pair_buffer_i32(&provider, &[1, 2], &[2, 3]);
923        let positives = pair_buffer_i32(&provider, &[1, 2], &[4, 5]);
924        let negatives = pair_buffer_i32(&provider, &[7], &[8]);
925
926        let err = provider
927            .ilp_exact_score(&[&p_b], &positives, &negatives)
928            .expect_err("I32 pair buffers must be rejected");
929        assert!(
930            err.to_string().contains("expected U64, U32, or Symbol"),
931            "unexpected error: {err}"
932        );
933    }
934}