1use 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 #[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 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 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 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 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 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 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 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 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 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 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 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 #[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 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 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 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 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 #[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 #[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}