1use 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 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 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 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 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 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 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 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 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 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 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 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 pub fn ilp_credit_forward_f32_launch(
606 &self,
607 row_offsets: &TrackedCudaSlice<u32>,
608 col_indices: &TrackedCudaSlice<u32>,
609 cand_probs: &CudaColumn, 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 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 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 pub fn ilp_credit_forward_f64_launch(
666 &self,
667 row_offsets: &TrackedCudaSlice<u32>,
668 col_indices: &TrackedCudaSlice<u32>,
669 cand_probs: &CudaColumn, 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 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 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 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 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 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 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 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 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 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 #[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 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 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 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}