Skip to main content

xlog_cuda/provider/
ilp.rs

1//! ILP (Inductive Logic Programming) kernel operations: credit/loss, COO fill, CSR histogram, reduce_sum.
2
3use std::marker::PhantomData;
4
5use crate::{DeviceSlice, LaunchAsync, LaunchConfig};
6use xlog_core::{Result, ScalarType, Schema, XlogError};
7
8use super::{ilp_credit_kernels, ilp_kernels, RawCudaView, ILP_CREDIT_MODULE, ILP_MODULE};
9use crate::memory::{CudaBuffer, CudaColumn, TrackedCudaSlice};
10
11impl super::CudaKernelProvider {
12    fn ilp_i32_view<'a>(
13        &self,
14        col: &'a CudaColumn,
15        num_elements: usize,
16    ) -> Result<RawCudaView<'a, i32>> {
17        let required_bytes = num_elements * std::mem::size_of::<i32>();
18        if col.num_bytes() < required_bytes {
19            return Err(XlogError::Kernel(format!(
20                "Column has {} bytes but {} required for {} i32 elements",
21                col.num_bytes(),
22                required_bytes,
23                num_elements
24            )));
25        }
26        let ptr = *col.device_ptr();
27        if !(ptr as usize).is_multiple_of(std::mem::align_of::<i32>()) {
28            return Err(XlogError::Kernel(
29                "Column device pointer is not i32-aligned".to_string(),
30            ));
31        }
32        Ok(RawCudaView {
33            ptr,
34            len: num_elements,
35            stream: col.stream().clone(),
36            _marker: PhantomData,
37        })
38    }
39
40    fn ilp_i64_view<'a>(
41        &self,
42        col: &'a CudaColumn,
43        num_elements: usize,
44    ) -> Result<RawCudaView<'a, i64>> {
45        let required_bytes = num_elements * std::mem::size_of::<i64>();
46        if col.num_bytes() < required_bytes {
47            return Err(XlogError::Kernel(format!(
48                "Column has {} bytes but {} required for {} i64 elements",
49                col.num_bytes(),
50                required_bytes,
51                num_elements
52            )));
53        }
54        let ptr = *col.device_ptr();
55        if !(ptr as usize).is_multiple_of(std::mem::align_of::<i64>()) {
56            return Err(XlogError::Kernel(
57                "Column device pointer is not i64-aligned".to_string(),
58            ));
59        }
60        Ok(RawCudaView {
61            ptr,
62            len: num_elements,
63            stream: col.stream().clone(),
64            _marker: PhantomData,
65        })
66    }
67
68    pub fn build_selected_id_mask(
69        &self,
70        ids_buf: &CudaBuffer,
71        candidate_count: usize,
72    ) -> Result<CudaBuffer> {
73        let selected_len = usize::try_from(ids_buf.num_rows())
74            .map_err(|_| XlogError::Kernel("selected id row count overflow".to_string()))?;
75        let candidate_count_u32 = u32::try_from(candidate_count).map_err(|_| {
76            XlogError::Kernel(format!(
77                "candidate count {} exceeds u32::MAX for strict sparse mask",
78                candidate_count
79            ))
80        })?;
81
82        let mut active_flags = self.memory.alloc::<u32>(candidate_count)?;
83        if candidate_count > 0 {
84            self.device
85                .inner()
86                .memset_zeros(&mut active_flags)
87                .map_err(|e| XlogError::Kernel(format!("zero strict sparse mask: {}", e)))?;
88        }
89
90        if selected_len > 0 {
91            let selected_len_u32 = u32::try_from(selected_len).map_err(|_| {
92                XlogError::Kernel(format!(
93                    "selected id count {} exceeds u32::MAX for strict sparse mask",
94                    selected_len
95                ))
96            })?;
97            let block_size = 256u32;
98            let grid_size = selected_len_u32.div_ceil(block_size);
99            let ids_col = ids_buf
100                .column(0)
101                .ok_or_else(|| XlogError::Kernel("selected id buffer has no column".to_string()))?;
102            match ids_buf.schema().column_type(0).ok_or_else(|| {
103                XlogError::Kernel("selected id buffer has no schema type".to_string())
104            })? {
105                ScalarType::U32 | ScalarType::Symbol => {
106                    let ids_view = self.column_as_u32_view(ids_col, selected_len)?;
107                    let func = self
108                        .device
109                        .inner()
110                        .get_func(ILP_MODULE, ilp_kernels::ILP_MARK_SELECTED_IDS_U32)
111                        .ok_or_else(|| {
112                            XlogError::Kernel(
113                                "ilp_mark_selected_ids_u32 kernel not found".to_string(),
114                            )
115                        })?;
116                    // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
117                    unsafe {
118                        func.clone().launch(
119                            LaunchConfig {
120                                grid_dim: (grid_size, 1, 1),
121                                block_dim: (block_size, 1, 1),
122                                shared_mem_bytes: 0,
123                            },
124                            (
125                                &ids_view,
126                                selected_len_u32,
127                                candidate_count_u32,
128                                &mut active_flags,
129                            ),
130                        )
131                    }
132                    .map_err(|e| {
133                        XlogError::Kernel(format!(
134                            "strict sparse selected-id scatter failed: {}",
135                            e
136                        ))
137                    })?;
138                }
139                ScalarType::I32 => {
140                    let ids_view = self.ilp_i32_view(ids_col, selected_len)?;
141                    let func = self
142                        .device
143                        .inner()
144                        .get_func(ILP_MODULE, ilp_kernels::ILP_MARK_SELECTED_IDS_I32)
145                        .ok_or_else(|| {
146                            XlogError::Kernel(
147                                "ilp_mark_selected_ids_i32 kernel not found".to_string(),
148                            )
149                        })?;
150                    // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
151                    unsafe {
152                        func.clone().launch(
153                            LaunchConfig {
154                                grid_dim: (grid_size, 1, 1),
155                                block_dim: (block_size, 1, 1),
156                                shared_mem_bytes: 0,
157                            },
158                            (
159                                &ids_view,
160                                selected_len_u32,
161                                candidate_count_u32,
162                                &mut active_flags,
163                            ),
164                        )
165                    }
166                    .map_err(|e| {
167                        XlogError::Kernel(format!(
168                            "strict sparse selected-id scatter failed: {}",
169                            e
170                        ))
171                    })?;
172                }
173                ScalarType::I64 => {
174                    let ids_view = self.ilp_i64_view(ids_col, selected_len)?;
175                    let func = self
176                        .device
177                        .inner()
178                        .get_func(ILP_MODULE, ilp_kernels::ILP_MARK_SELECTED_IDS_I64)
179                        .ok_or_else(|| {
180                            XlogError::Kernel(
181                                "ilp_mark_selected_ids_i64 kernel not found".to_string(),
182                            )
183                        })?;
184                    // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
185                    unsafe {
186                        func.clone().launch(
187                            LaunchConfig {
188                                grid_dim: (grid_size, 1, 1),
189                                block_dim: (block_size, 1, 1),
190                                shared_mem_bytes: 0,
191                            },
192                            (
193                                &ids_view,
194                                selected_len_u32,
195                                candidate_count_u32,
196                                &mut active_flags,
197                            ),
198                        )
199                    }
200                    .map_err(|e| {
201                        XlogError::Kernel(format!(
202                            "strict sparse selected-id scatter failed: {}",
203                            e
204                        ))
205                    })?;
206                }
207                ScalarType::U64 => {
208                    let ids_view = self.column_as_u64_view(ids_col, selected_len)?;
209                    let func = self
210                        .device
211                        .inner()
212                        .get_func(ILP_MODULE, ilp_kernels::ILP_MARK_SELECTED_IDS_U64)
213                        .ok_or_else(|| {
214                            XlogError::Kernel(
215                                "ilp_mark_selected_ids_u64 kernel not found".to_string(),
216                            )
217                        })?;
218                    // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
219                    unsafe {
220                        func.clone().launch(
221                            LaunchConfig {
222                                grid_dim: (grid_size, 1, 1),
223                                block_dim: (block_size, 1, 1),
224                                shared_mem_bytes: 0,
225                            },
226                            (
227                                &ids_view,
228                                selected_len_u32,
229                                candidate_count_u32,
230                                &mut active_flags,
231                            ),
232                        )
233                    }
234                    .map_err(|e| {
235                        XlogError::Kernel(format!(
236                            "strict sparse selected-id scatter failed: {}",
237                            e
238                        ))
239                    })?;
240                }
241                other => {
242                    return Err(XlogError::Kernel(format!(
243                        "selected candidate ids must be I32/I64/U32/U64, got {:?}",
244                        other
245                    )));
246                }
247            }
248
249            self.device
250                .synchronize()
251                .map_err(|e| XlogError::Kernel(format!("strict sparse scatter sync: {}", e)))?;
252        }
253
254        let d_num_rows = self.upload_device_row_count(candidate_count_u32)?;
255        Ok(CudaBuffer::from_columns_with_host_count(
256            vec![active_flags.into_bytes().into()],
257            candidate_count as u64,
258            d_num_rows,
259            Schema::new(vec![("active".to_string(), ScalarType::U32)]),
260            candidate_count_u32,
261        ))
262    }
263
264    pub fn validate_selected_ids(
265        &self,
266        ids_buf: &CudaBuffer,
267        candidate_count: usize,
268    ) -> Result<()> {
269        let selected_len = usize::try_from(ids_buf.num_rows())
270            .map_err(|_| XlogError::Kernel("selected id row count overflow".to_string()))?;
271        let candidate_count_u32 = u32::try_from(candidate_count).map_err(|_| {
272            XlogError::Kernel(format!(
273                "candidate count {} exceeds u32::MAX for strict sparse mask",
274                candidate_count
275            ))
276        })?;
277
278        if selected_len == 0 {
279            return Ok(());
280        }
281
282        let selected_len_u32 = u32::try_from(selected_len).map_err(|_| {
283            XlogError::Kernel(format!(
284                "selected id count {} exceeds u32::MAX for strict sparse mask",
285                selected_len
286            ))
287        })?;
288        let block_size = 256u32;
289        let grid_size = selected_len_u32.div_ceil(block_size);
290        let ids_col = ids_buf
291            .column(0)
292            .ok_or_else(|| XlogError::Kernel("selected id buffer has no column".to_string()))?;
293
294        let mut seen_flags = self.memory.alloc::<u32>(candidate_count)?;
295        if candidate_count > 0 {
296            self.device
297                .inner()
298                .memset_zeros(&mut seen_flags)
299                .map_err(|e| {
300                    XlogError::Kernel(format!("zero strict sparse validation flags: {}", e))
301                })?;
302        }
303
304        let mut error_code = self.memory.alloc::<u32>(1)?;
305        let mut error_pos = self.memory.alloc::<u32>(1)?;
306        self.device
307            .inner()
308            .memset_zeros(&mut error_code)
309            .map_err(|e| XlogError::Kernel(format!("zero strict sparse error code: {}", e)))?;
310        self.device
311            .inner()
312            .memset_zeros(&mut error_pos)
313            .map_err(|e| XlogError::Kernel(format!("zero strict sparse error pos: {}", e)))?;
314
315        match ids_buf
316            .schema()
317            .column_type(0)
318            .ok_or_else(|| XlogError::Kernel("selected id buffer has no schema type".to_string()))?
319        {
320            ScalarType::U32 | ScalarType::Symbol => {
321                let ids_view = self.column_as_u32_view(ids_col, selected_len)?;
322                let func = self
323                    .device
324                    .inner()
325                    .get_func(ILP_MODULE, ilp_kernels::ILP_VALIDATE_SELECTED_IDS_U32)
326                    .ok_or_else(|| {
327                        XlogError::Kernel(
328                            "ilp_validate_selected_ids_u32 kernel not found".to_string(),
329                        )
330                    })?;
331                // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
332                unsafe {
333                    func.clone().launch(
334                        LaunchConfig {
335                            grid_dim: (grid_size, 1, 1),
336                            block_dim: (block_size, 1, 1),
337                            shared_mem_bytes: 0,
338                        },
339                        (
340                            &ids_view,
341                            selected_len_u32,
342                            candidate_count_u32,
343                            &mut seen_flags,
344                            &mut error_code,
345                            &mut error_pos,
346                        ),
347                    )
348                }
349                .map_err(|e| {
350                    XlogError::Kernel(format!(
351                        "strict sparse selected-id validation failed: {}",
352                        e
353                    ))
354                })?;
355            }
356            ScalarType::I32 => {
357                let ids_view = self.ilp_i32_view(ids_col, selected_len)?;
358                let func = self
359                    .device
360                    .inner()
361                    .get_func(ILP_MODULE, ilp_kernels::ILP_VALIDATE_SELECTED_IDS_I32)
362                    .ok_or_else(|| {
363                        XlogError::Kernel(
364                            "ilp_validate_selected_ids_i32 kernel not found".to_string(),
365                        )
366                    })?;
367                // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
368                unsafe {
369                    func.clone().launch(
370                        LaunchConfig {
371                            grid_dim: (grid_size, 1, 1),
372                            block_dim: (block_size, 1, 1),
373                            shared_mem_bytes: 0,
374                        },
375                        (
376                            &ids_view,
377                            selected_len_u32,
378                            candidate_count_u32,
379                            &mut seen_flags,
380                            &mut error_code,
381                            &mut error_pos,
382                        ),
383                    )
384                }
385                .map_err(|e| {
386                    XlogError::Kernel(format!(
387                        "strict sparse selected-id validation failed: {}",
388                        e
389                    ))
390                })?;
391            }
392            ScalarType::I64 => {
393                let ids_view = self.ilp_i64_view(ids_col, selected_len)?;
394                let func = self
395                    .device
396                    .inner()
397                    .get_func(ILP_MODULE, ilp_kernels::ILP_VALIDATE_SELECTED_IDS_I64)
398                    .ok_or_else(|| {
399                        XlogError::Kernel(
400                            "ilp_validate_selected_ids_i64 kernel not found".to_string(),
401                        )
402                    })?;
403                // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
404                unsafe {
405                    func.clone().launch(
406                        LaunchConfig {
407                            grid_dim: (grid_size, 1, 1),
408                            block_dim: (block_size, 1, 1),
409                            shared_mem_bytes: 0,
410                        },
411                        (
412                            &ids_view,
413                            selected_len_u32,
414                            candidate_count_u32,
415                            &mut seen_flags,
416                            &mut error_code,
417                            &mut error_pos,
418                        ),
419                    )
420                }
421                .map_err(|e| {
422                    XlogError::Kernel(format!(
423                        "strict sparse selected-id validation failed: {}",
424                        e
425                    ))
426                })?;
427            }
428            ScalarType::U64 => {
429                let ids_view = self.column_as_u64_view(ids_col, selected_len)?;
430                let func = self
431                    .device
432                    .inner()
433                    .get_func(ILP_MODULE, ilp_kernels::ILP_VALIDATE_SELECTED_IDS_U64)
434                    .ok_or_else(|| {
435                        XlogError::Kernel(
436                            "ilp_validate_selected_ids_u64 kernel not found".to_string(),
437                        )
438                    })?;
439                // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
440                unsafe {
441                    func.clone().launch(
442                        LaunchConfig {
443                            grid_dim: (grid_size, 1, 1),
444                            block_dim: (block_size, 1, 1),
445                            shared_mem_bytes: 0,
446                        },
447                        (
448                            &ids_view,
449                            selected_len_u32,
450                            candidate_count_u32,
451                            &mut seen_flags,
452                            &mut error_code,
453                            &mut error_pos,
454                        ),
455                    )
456                }
457                .map_err(|e| {
458                    XlogError::Kernel(format!(
459                        "strict sparse selected-id validation failed: {}",
460                        e
461                    ))
462                })?;
463            }
464            other => {
465                return Err(XlogError::Kernel(format!(
466                    "selected candidate ids must be I32/I64/U32/U64, got {:?}",
467                    other
468                )));
469            }
470        }
471
472        self.device
473            .synchronize()
474            .map_err(|e| XlogError::Kernel(format!("strict sparse validation sync: {}", e)))?;
475
476        let error_code_host = self.dtoh_scalar_untracked(&error_code, 0)?;
477        if error_code_host == 0 {
478            return Ok(());
479        }
480        let error_pos_host = self.dtoh_scalar_untracked(&error_pos, 0)?;
481        match error_code_host {
482            1 => Err(XlogError::Kernel(format!(
483                "selected candidate id out of range at position {}",
484                error_pos_host
485            ))),
486            2 => Err(XlogError::Kernel(format!(
487                "duplicate selected candidate id at position {}",
488                error_pos_host
489            ))),
490            code => Err(XlogError::Kernel(format!(
491                "strict sparse selected-id validation failed with error code {}",
492                code
493            ))),
494        }
495    }
496
497    pub fn filter_buffer_by_candidate_flag(
498        &self,
499        input: &CudaBuffer,
500        candidate_flags: &CudaBuffer,
501        candidate_idx: usize,
502    ) -> Result<CudaBuffer> {
503        if input.is_empty() {
504            return self.create_empty_buffer(input.schema().clone());
505        }
506        if candidate_idx >= candidate_flags.num_rows() as usize {
507            return Err(XlogError::Kernel(format!(
508                "candidate flag index {} out of range [0, {})",
509                candidate_idx,
510                candidate_flags.num_rows()
511            )));
512        }
513
514        let flag_col = candidate_flags
515            .column(0)
516            .ok_or_else(|| XlogError::Kernel("candidate flag buffer has no column".to_string()))?;
517        let flag_view = self.column_as_u32_view(flag_col, candidate_flags.num_rows() as usize)?;
518        let row_count = u32::try_from(input.num_rows()).map_err(|_| {
519            XlogError::Kernel(format!(
520                "strict sparse row count {} exceeds u32::MAX",
521                input.num_rows()
522            ))
523        })?;
524        let candidate_idx_u32 = u32::try_from(candidate_idx).map_err(|_| {
525            XlogError::Kernel(format!(
526                "candidate flag index {} exceeds u32::MAX",
527                candidate_idx
528            ))
529        })?;
530
531        let mut row_mask = self.memory.alloc::<u8>(row_count as usize)?;
532        let func = self
533            .device
534            .inner()
535            .get_func(ILP_MODULE, ilp_kernels::ILP_BROADCAST_CANDIDATE_FLAG)
536            .ok_or_else(|| {
537                XlogError::Kernel("ilp_broadcast_candidate_flag kernel not found".to_string())
538            })?;
539        let block_size = 256u32;
540        let grid_size = row_count.div_ceil(block_size);
541        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
542        unsafe {
543            func.clone().launch(
544                LaunchConfig {
545                    grid_dim: (grid_size, 1, 1),
546                    block_dim: (block_size, 1, 1),
547                    shared_mem_bytes: 0,
548                },
549                (&flag_view, candidate_idx_u32, row_count, &mut row_mask),
550            )
551        }
552        .map_err(|e| XlogError::Kernel(format!("strict sparse flag broadcast failed: {}", e)))?;
553
554        self.filter_by_device_mask(input, &row_mask)
555    }
556
557    // ─── ILP credit kernel launchers ───────────────────────────────────
558
559    /// Launch `ilp_coo_fill` kernel: writes `(compacted_fact_indices[i], cidx)`
560    /// pairs at `coo_fact[offset..]` and `coo_cand[offset..]`.
561    pub fn ilp_coo_fill_launch(
562        &self,
563        compacted_fact_indices: &TrackedCudaSlice<u32>,
564        cidx: u32,
565        count: u32,
566        offset: u32,
567        coo_fact: &mut TrackedCudaSlice<u32>,
568        coo_cand: &mut TrackedCudaSlice<u32>,
569    ) -> Result<()> {
570        if count == 0 {
571            return Ok(());
572        }
573        let func = self
574            .device
575            .inner()
576            .get_func(ILP_CREDIT_MODULE, ilp_credit_kernels::ILP_COO_FILL)
577            .ok_or_else(|| XlogError::Kernel("ilp_coo_fill kernel not found".to_string()))?;
578        let block_size = 256u32;
579        let grid_size = count.div_ceil(block_size);
580        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
581        unsafe {
582            func.clone().launch(
583                LaunchConfig {
584                    grid_dim: (grid_size, 1, 1),
585                    block_dim: (block_size, 1, 1),
586                    shared_mem_bytes: 0,
587                },
588                (
589                    compacted_fact_indices,
590                    cidx,
591                    count,
592                    offset,
593                    coo_fact,
594                    coo_cand,
595                ),
596            )
597        }
598        .map_err(|e| XlogError::Kernel(format!("ilp_coo_fill failed: {}", e)))?;
599        self.device.synchronize()?;
600        Ok(())
601    }
602
603    /// Launch `ilp_credit_forward_f32`: CSR credit gather + clamp + NLL loss.
604    /// Returns `(credit_out, loss_contrib)` device slices of length `num_facts`.
605    pub fn ilp_credit_forward_f32_launch(
606        &self,
607        row_offsets: &TrackedCudaSlice<u32>,
608        col_indices: &TrackedCudaSlice<u32>,
609        cand_probs: &CudaColumn, // raw byte column from CudaBuffer
610        is_positive: &TrackedCudaSlice<u8>,
611        num_facts: u32,
612        eps: f32,
613    ) -> Result<(TrackedCudaSlice<f32>, TrackedCudaSlice<f32>)> {
614        let mut credit_out = self.memory.alloc::<f32>(num_facts as usize)?;
615        let mut loss_contrib = self.memory.alloc::<f32>(num_facts as usize)?;
616        if num_facts == 0 {
617            return Ok((credit_out, loss_contrib));
618        }
619        let func = self
620            .device
621            .inner()
622            .get_func(
623                ILP_CREDIT_MODULE,
624                ilp_credit_kernels::ILP_CREDIT_FORWARD_F32,
625            )
626            .ok_or_else(|| {
627                XlogError::Kernel("ilp_credit_forward_f32 kernel not found".to_string())
628            })?;
629        let block_size = 256u32;
630        let grid_size = num_facts.div_ceil(block_size);
631        // reinterpret the u8 byte column as f32 for the kernel
632        let cand_view = RawCudaView::<f32> {
633            ptr: *cand_probs.device_ptr(),
634            len: cudarc::driver::DeviceSlice::len(cand_probs) / 4,
635            stream: cand_probs.stream().clone(),
636            _marker: PhantomData,
637        };
638        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
639        unsafe {
640            func.clone().launch(
641                LaunchConfig {
642                    grid_dim: (grid_size, 1, 1),
643                    block_dim: (block_size, 1, 1),
644                    shared_mem_bytes: 0,
645                },
646                (
647                    row_offsets,
648                    col_indices,
649                    &cand_view,
650                    is_positive,
651                    num_facts,
652                    eps,
653                    &mut credit_out,
654                    &mut loss_contrib,
655                ),
656            )
657        }
658        .map_err(|e| XlogError::Kernel(format!("ilp_credit_forward_f32 failed: {}", e)))?;
659        self.device.synchronize()?;
660        Ok((credit_out, loss_contrib))
661    }
662
663    /// Launch `ilp_credit_forward_f64`: CSR credit gather + clamp + NLL loss.
664    /// Returns `(credit_out, loss_contrib)` device slices of length `num_facts`.
665    pub fn ilp_credit_forward_f64_launch(
666        &self,
667        row_offsets: &TrackedCudaSlice<u32>,
668        col_indices: &TrackedCudaSlice<u32>,
669        cand_probs: &CudaColumn, // raw byte column from CudaBuffer
670        is_positive: &TrackedCudaSlice<u8>,
671        num_facts: u32,
672        eps: f64,
673    ) -> Result<(TrackedCudaSlice<f64>, TrackedCudaSlice<f64>)> {
674        let mut credit_out = self.memory.alloc::<f64>(num_facts as usize)?;
675        let mut loss_contrib = self.memory.alloc::<f64>(num_facts as usize)?;
676        if num_facts == 0 {
677            return Ok((credit_out, loss_contrib));
678        }
679        let func = self
680            .device
681            .inner()
682            .get_func(
683                ILP_CREDIT_MODULE,
684                ilp_credit_kernels::ILP_CREDIT_FORWARD_F64,
685            )
686            .ok_or_else(|| {
687                XlogError::Kernel("ilp_credit_forward_f64 kernel not found".to_string())
688            })?;
689        let block_size = 256u32;
690        let grid_size = num_facts.div_ceil(block_size);
691        let cand_view = RawCudaView::<f64> {
692            ptr: *cand_probs.device_ptr(),
693            len: cudarc::driver::DeviceSlice::len(cand_probs) / 8,
694            stream: cand_probs.stream().clone(),
695            _marker: PhantomData,
696        };
697        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
698        unsafe {
699            func.clone().launch(
700                LaunchConfig {
701                    grid_dim: (grid_size, 1, 1),
702                    block_dim: (block_size, 1, 1),
703                    shared_mem_bytes: 0,
704                },
705                (
706                    row_offsets,
707                    col_indices,
708                    &cand_view,
709                    is_positive,
710                    num_facts,
711                    eps,
712                    &mut credit_out,
713                    &mut loss_contrib,
714                ),
715            )
716        }
717        .map_err(|e| XlogError::Kernel(format!("ilp_credit_forward_f64 failed: {}", e)))?;
718        self.device.synchronize()?;
719        Ok((credit_out, loss_contrib))
720    }
721
722    /// Launch `ilp_credit_backward_f32`: gradient scatter via CSR + atomicAdd.
723    /// Returns `d_cand_probs` gradient of length `num_cands` (zeroed, then accumulated).
724    pub fn ilp_credit_backward_f32_launch(
725        &self,
726        row_offsets: &TrackedCudaSlice<u32>,
727        col_indices: &TrackedCudaSlice<u32>,
728        credit_out: &TrackedCudaSlice<f32>,
729        is_positive: &TrackedCudaSlice<u8>,
730        num_facts: u32,
731        num_cands: u32,
732    ) -> Result<TrackedCudaSlice<f32>> {
733        let mut d_grad = self.memory.alloc::<f32>(num_cands as usize)?;
734        self.device
735            .inner()
736            .memset_zeros(&mut d_grad)
737            .map_err(|e| XlogError::Kernel(format!("Failed to zero grad: {}", e)))?;
738        if num_facts == 0 {
739            return Ok(d_grad);
740        }
741        let func = self
742            .device
743            .inner()
744            .get_func(
745                ILP_CREDIT_MODULE,
746                ilp_credit_kernels::ILP_CREDIT_BACKWARD_F32,
747            )
748            .ok_or_else(|| {
749                XlogError::Kernel("ilp_credit_backward_f32 kernel not found".to_string())
750            })?;
751        let block_size = 256u32;
752        let grid_size = num_facts.div_ceil(block_size);
753        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
754        unsafe {
755            func.clone().launch(
756                LaunchConfig {
757                    grid_dim: (grid_size, 1, 1),
758                    block_dim: (block_size, 1, 1),
759                    shared_mem_bytes: 0,
760                },
761                (
762                    row_offsets,
763                    col_indices,
764                    credit_out,
765                    is_positive,
766                    num_facts,
767                    &mut d_grad,
768                ),
769            )
770        }
771        .map_err(|e| XlogError::Kernel(format!("ilp_credit_backward_f32 failed: {}", e)))?;
772        self.device.synchronize()?;
773        Ok(d_grad)
774    }
775
776    /// Launch `ilp_credit_backward_f64`: gradient scatter via CSR + atomicAdd.
777    /// Returns `d_cand_probs` gradient of length `num_cands` (zeroed, then accumulated).
778    pub fn ilp_credit_backward_f64_launch(
779        &self,
780        row_offsets: &TrackedCudaSlice<u32>,
781        col_indices: &TrackedCudaSlice<u32>,
782        credit_out: &TrackedCudaSlice<f64>,
783        is_positive: &TrackedCudaSlice<u8>,
784        num_facts: u32,
785        num_cands: u32,
786    ) -> Result<TrackedCudaSlice<f64>> {
787        let mut d_grad = self.memory.alloc::<f64>(num_cands as usize)?;
788        self.device
789            .inner()
790            .memset_zeros(&mut d_grad)
791            .map_err(|e| XlogError::Kernel(format!("Failed to zero grad: {}", e)))?;
792        if num_facts == 0 {
793            return Ok(d_grad);
794        }
795        let func = self
796            .device
797            .inner()
798            .get_func(
799                ILP_CREDIT_MODULE,
800                ilp_credit_kernels::ILP_CREDIT_BACKWARD_F64,
801            )
802            .ok_or_else(|| {
803                XlogError::Kernel("ilp_credit_backward_f64 kernel not found".to_string())
804            })?;
805        let block_size = 256u32;
806        let grid_size = num_facts.div_ceil(block_size);
807        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
808        unsafe {
809            func.clone().launch(
810                LaunchConfig {
811                    grid_dim: (grid_size, 1, 1),
812                    block_dim: (block_size, 1, 1),
813                    shared_mem_bytes: 0,
814                },
815                (
816                    row_offsets,
817                    col_indices,
818                    credit_out,
819                    is_positive,
820                    num_facts,
821                    &mut d_grad,
822                ),
823            )
824        }
825        .map_err(|e| XlogError::Kernel(format!("ilp_credit_backward_f64 failed: {}", e)))?;
826        self.device.synchronize()?;
827        Ok(d_grad)
828    }
829
830    /// GPU-side sum reduction (f32).
831    ///
832    /// Sums `n` elements of `input` on device and returns a single-element
833    /// device buffer containing the result.  The caller must zero the output
834    /// buffer *before* launching the kernel — this function handles that.
835    pub fn ilp_reduce_sum_f32_launch(
836        &self,
837        input: &TrackedCudaSlice<f32>,
838        n: u32,
839    ) -> Result<TrackedCudaSlice<f32>> {
840        let mut d_result = self.memory.alloc::<f32>(1)?;
841        self.device
842            .inner()
843            .memset_zeros(&mut d_result)
844            .map_err(|e| XlogError::Kernel(format!("ilp_reduce_sum_f32 zero result: {}", e)))?;
845
846        if n == 0 {
847            return Ok(d_result);
848        }
849
850        let func = self
851            .device
852            .inner()
853            .get_func(ILP_MODULE, ilp_kernels::ILP_REDUCE_SUM_F32)
854            .ok_or_else(|| XlogError::Kernel("ilp_reduce_sum_f32 not found".to_string()))?;
855        let block_size = 256u32;
856        let grid_size = n.div_ceil(block_size);
857        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
858        unsafe {
859            func.clone().launch(
860                LaunchConfig {
861                    grid_dim: (grid_size, 1, 1),
862                    block_dim: (block_size, 1, 1),
863                    shared_mem_bytes: 0,
864                },
865                (input, n, &mut d_result),
866            )
867        }
868        .map_err(|e| XlogError::Kernel(format!("ilp_reduce_sum_f32: {}", e)))?;
869        self.device.synchronize()?;
870        Ok(d_result)
871    }
872
873    /// GPU-side sum reduction (f64).
874    ///
875    /// Sums `n` elements of `input` on device and returns a single-element
876    /// device buffer containing the result.  Requires sm_60+ for double
877    /// atomicAdd (this project targets sm_75 baseline).
878    pub fn ilp_reduce_sum_f64_launch(
879        &self,
880        input: &TrackedCudaSlice<f64>,
881        n: u32,
882    ) -> Result<TrackedCudaSlice<f64>> {
883        let mut d_result = self.memory.alloc::<f64>(1)?;
884        self.device
885            .inner()
886            .memset_zeros(&mut d_result)
887            .map_err(|e| XlogError::Kernel(format!("ilp_reduce_sum_f64 zero result: {}", e)))?;
888
889        if n == 0 {
890            return Ok(d_result);
891        }
892
893        let func = self
894            .device
895            .inner()
896            .get_func(ILP_MODULE, ilp_kernels::ILP_REDUCE_SUM_F64)
897            .ok_or_else(|| XlogError::Kernel("ilp_reduce_sum_f64 not found".to_string()))?;
898        let block_size = 256u32;
899        let grid_size = n.div_ceil(block_size);
900        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
901        unsafe {
902            func.clone().launch(
903                LaunchConfig {
904                    grid_dim: (grid_size, 1, 1),
905                    block_dim: (block_size, 1, 1),
906                    shared_mem_bytes: 0,
907                },
908                (input, n, &mut d_result),
909            )
910        }
911        .map_err(|e| XlogError::Kernel(format!("ilp_reduce_sum_f64: {}", e)))?;
912        self.device.synchronize()?;
913        Ok(d_result)
914    }
915
916    /// Fill COO arrays from a device-side mask and prefix-sum.
917    ///
918    /// For each set bit in `mask`, writes the corresponding `fact_indices` entry
919    /// into `coo_fact` and `cand_value` into `coo_cand` at the position
920    /// determined by `d_offsets[offset_idx] + prefix_sum[tid]`.
921    ///
922    /// Parameters:
923    /// - `offset_idx`: index into `d_offsets` for the write base position
924    /// - `cand_value`: actual candidate index to write into `coo_cand`
925    ///
926    /// This keeps COO assembly fully on device, eliminating the mask D2H transfer.
927    #[allow(clippy::too_many_arguments)]
928    pub fn ilp_coo_fill_from_mask_launch(
929        &self,
930        mask: &TrackedCudaSlice<u8>,
931        prefix_sum: &TrackedCudaSlice<u32>,
932        fact_indices: &TrackedCudaSlice<u32>,
933        offset_idx: u32,
934        cand_value: u32,
935        num_query: u32,
936        d_offsets: &TrackedCudaSlice<u32>,
937        coo_fact: &mut TrackedCudaSlice<u32>,
938        coo_cand: &mut TrackedCudaSlice<u32>,
939    ) -> Result<()> {
940        if num_query == 0 {
941            return Ok(());
942        }
943        let func = self
944            .device()
945            .inner()
946            .get_func(ILP_MODULE, ilp_kernels::ILP_COO_FILL_FROM_MASK)
947            .ok_or_else(|| XlogError::Kernel("ilp_coo_fill_from_mask not found".to_string()))?;
948        let block_size = 256u32;
949        let grid_size = num_query.div_ceil(block_size);
950        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
951        unsafe {
952            func.clone().launch(
953                LaunchConfig {
954                    grid_dim: (grid_size, 1, 1),
955                    block_dim: (block_size, 1, 1),
956                    shared_mem_bytes: 0,
957                },
958                (
959                    mask,
960                    prefix_sum,
961                    fact_indices,
962                    offset_idx,
963                    cand_value,
964                    num_query,
965                    d_offsets,
966                    coo_fact,
967                    coo_cand,
968                ),
969            )
970        }
971        .map_err(|e| XlogError::Kernel(format!("ilp_coo_fill_from_mask: {}", e)))?;
972        self.device()
973            .inner()
974            .synchronize()
975            .map_err(|e| XlogError::Kernel(format!("ilp_coo_fill_from_mask sync: {}", e)))?;
976        Ok(())
977    }
978
979    /// Build a histogram of fact indices from sorted COO data.
980    ///
981    /// For each entry in `sorted_facts[0..nnz]`, atomically increments
982    /// the corresponding bin in the output histogram. The result is a
983    /// device-side count array of length `num_facts`, suitable for
984    /// prefix-sum to produce CSR `row_offsets`.
985    ///
986    /// The caller provides sorted fact indices; the histogram is
987    /// zero-initialized internally.
988    pub fn ilp_csr_histogram_launch(
989        &self,
990        sorted_facts: &TrackedCudaSlice<u32>,
991        nnz: u32,
992        num_facts: u32,
993    ) -> Result<TrackedCudaSlice<u32>> {
994        let mut d_hist = self.memory().alloc::<u32>(num_facts as usize)?;
995        self.device()
996            .inner()
997            .memset_zeros(&mut d_hist)
998            .map_err(|e| XlogError::Kernel(format!("ilp_csr_histogram zero hist: {}", e)))?;
999
1000        if nnz == 0 {
1001            return Ok(d_hist);
1002        }
1003
1004        let func = self
1005            .device()
1006            .inner()
1007            .get_func(ILP_MODULE, ilp_kernels::ILP_CSR_HISTOGRAM)
1008            .ok_or_else(|| XlogError::Kernel("ilp_csr_histogram kernel not found".to_string()))?;
1009
1010        let block_size = 256u32;
1011        let grid_size = nnz.div_ceil(block_size);
1012
1013        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1014        unsafe {
1015            func.clone()
1016                .launch(
1017                    cudarc::driver::LaunchConfig {
1018                        grid_dim: (grid_size, 1, 1),
1019                        block_dim: (block_size, 1, 1),
1020                        shared_mem_bytes: 0,
1021                    },
1022                    (sorted_facts, nnz, num_facts, &mut d_hist),
1023                )
1024                .map_err(|e| XlogError::Kernel(format!("ilp_csr_histogram launch: {}", e)))?;
1025        }
1026
1027        self.device()
1028            .inner()
1029            .synchronize()
1030            .map_err(|e| XlogError::Kernel(format!("ilp_csr_histogram sync: {}", e)))?;
1031
1032        Ok(d_hist)
1033    }
1034}