Skip to main content

xlog_cuda/
kernel_manifest_data.rs

1// Single source of truth for CUDA kernel modules.
2//
3// This file is consumed by both `build.rs` (via include!()) and `lib.rs`
4// (via `pub mod kernel_manifest_data`) to avoid the kernel list being
5// duplicated in two places.
6// NOTE: Use regular comments (//), NOT inner doc comments (//!), because
7// include!() in build.rs would interpret //! as documenting the wrong item.
8
9/// Module names matching the .cu filenames (without extension).
10/// Order matches provider/mod.rs load order. All 24 modules listed.
11pub const KERNEL_CU_NAMES: &[&str] = &[
12    "join",
13    "dedup",
14    "groupby",
15    "scan",
16    "sort",
17    "filter",
18    "set_ops",
19    "pack",
20    "pir",
21    "cnf",
22    "cache",
23    "weights",
24    "circuit",
25    "mc_sample",
26    "mc_eval",
27    "arith",
28    "sat",
29    "d4",
30    "neural",
31    "ilp",
32    "ilp_credit",
33    "ilp_exact",
34    "epistemic",
35    "wcoj",
36    "mc_resident",
37    "joint_solve",
38];
39
40/// Describes a single CUDA module: the .cu file name, the runtime module name
41/// used by cudarc, and the list of kernel function entry points within.
42pub struct KernelModuleSpec {
43    pub cu_name: &'static str,
44    pub module_name: &'static str,
45    pub kernels: &'static [&'static str],
46}
47
48/// All kernel modules with their entry-point function names.
49/// Order and cu_names match `KERNEL_CU_NAMES`.
50pub const KERNEL_MODULES: &[KernelModuleSpec] = &[
51    KernelModuleSpec {
52        cu_name: "join",
53        module_name: "xlog_join",
54        kernels: &[
55            "hash_join_build",
56            "hash_join_probe",
57            "compute_composite_hash",
58            "hash_join_bucket_count_v2",
59            "hash_join_scatter_v2",
60            "hash_join_probe_v2",
61            "hash_join_probe_v2_count_per_row",
62            "hash_join_probe_v2_materialize",
63            "hash_join_total_from_scan",
64            "hash_join_csm_unmatched_mask",
65            "hash_join_semi",
66            "hash_join_anti",
67            "init_hash_table",
68            // Nested-loop inner join production operator (emit-pairs design).
69            "nested_loop_join_inner_u32_1key_pairs",
70            // Sort-merge inner join provider-level operator (emit-pairs design,
71            // caller-asserted pre-sorted inputs).
72            "sort_merge_join_inner_u32_1key_pairs",
73        ],
74    },
75    KernelModuleSpec {
76        cu_name: "dedup",
77        module_name: "xlog_dedup",
78        kernels: &[
79            "mark_duplicates",
80            "mark_unique_columnar",
81            "mark_unique_and_scan_columnar",
82            "compact_rows",
83            "mark_unique_full_row_bytewise",
84            "mark_diff_full_row_typed_sorted",
85            "small_sort_full_row_indices_typed",
86        ],
87    },
88    KernelModuleSpec {
89        cu_name: "groupby",
90        module_name: "xlog_groupby",
91        kernels: &[
92            "detect_group_boundaries",
93            "detect_boundaries",
94            "extract_group_keys",
95            "group_ids_from_boundaries",
96            "group_start_indices",
97            "capture_num_groups",
98            "groupby_count",
99            "groupby_sum",
100            "groupby_sum_u64",
101            "groupby_min",
102            "groupby_min_u64",
103            "groupby_max",
104            "groupby_max_u64",
105            "groupby_logsumexp_max",
106            "groupby_logsumexp_sumexp",
107            "groupby_logsumexp_final",
108        ],
109    },
110    KernelModuleSpec {
111        cu_name: "scan",
112        module_name: "xlog_scan",
113        kernels: &[
114            "block_inclusive_scan",
115            "add_block_offsets",
116            "exclusive_scan_mask",
117            "count_mask",
118            "multiblock_scan_phase1",
119            "multiblock_scan_u32_phase1",
120            "multiblock_scan_phase2",
121            "multiblock_scan_phase3",
122        ],
123    },
124    KernelModuleSpec {
125        cu_name: "sort",
126        module_name: "xlog_sort",
127        kernels: &[
128            "radix_histogram",
129            "radix_scatter",
130            "compute_ranks",
131            "radix_scatter_stable",
132            "compute_digit_prefix_sums",
133            "init_indices",
134            "apply_permutation_u32",
135            "apply_permutation_bytes",
136            "gather_keys_i32_ordered_u32",
137            "gather_keys_f32_ordered_u32",
138            "gather_keys_bool_ordered_u32",
139            "gather_keys_u64_lo_u32",
140            "gather_keys_u64_hi_u32",
141            "gather_keys_i64_lo_u32",
142            "gather_keys_i64_hi_u32",
143            "gather_keys_f64_lo_u32",
144            "gather_keys_f64_hi_u32",
145            // Sort-merge sortedness-detection kernel used by provider-level
146            // callers before invoking the sort-merge join.
147            "check_ascending_sorted_u32",
148        ],
149    },
150    KernelModuleSpec {
151        cu_name: "filter",
152        module_name: "xlog_filter",
153        kernels: &[
154            "filter_compare_u32",
155            "filter_compare_i64",
156            "filter_compare_f64",
157            "filter_compare_i32",
158            "filter_compare_u64",
159            "filter_compare_f32",
160            "filter_compare_u8",
161            "filter_compare_u32_scan_phase1",
162            "filter_compare_f64_scan_phase1",
163            "filter_compare_f32_scan_phase1",
164            "filter_compare_u32_col",
165            "filter_compare_i32_col",
166            "filter_compare_i64_col",
167            "filter_compare_u64_col",
168            "filter_compare_f32_col",
169            "filter_compare_f64_col",
170            "filter_compare_u8_col",
171            "fill_u32_iota",
172            "fill_u32_const",
173            "mark_random_vars",
174            "random_var_to_bit_from_list",
175            "check_random_var_count",
176            "compact_u32_by_mask",
177            "compact_i64_by_mask",
178            "compact_f64_by_mask",
179            "compact_bytes_by_mask",
180            "capture_compact_count",
181            "mask_clamp_rows",
182            "mask_and",
183            "mask_or",
184            "mask_not",
185        ],
186    },
187    KernelModuleSpec {
188        cu_name: "set_ops",
189        module_name: "xlog_set_ops",
190        kernels: &["concat_u32", "concat_bytes", "sorted_diff_mark"],
191    },
192    KernelModuleSpec {
193        cu_name: "pack",
194        module_name: "xlog_pack",
195        kernels: &[
196            "pack_keys",
197            "hash_packed_keys",
198            "pack_and_hash_keys",
199            "pack_and_hash_keys_generic",
200            "pack_keys_aligned",
201            "unpack_column",
202            "unpack_column_counted",
203            "gather_packed_rows",
204            "gather_packed_rows_counted",
205            "scatter_packed_rows",
206            "compare_packed_keys",
207            "pack_bools_to_bitmap",
208        ],
209    },
210    KernelModuleSpec {
211        cu_name: "pir",
212        module_name: "xlog_pir",
213        kernels: &[
214            "pir_pack_keys",
215            "pir_hash_keys",
216            "pir_mark_unique",
217            "pir_find_existing",
218            "pir_mark_new_groups",
219            "pir_build_group_ids",
220            "pir_fill_child_parents",
221            "pir_mark_unique_pairs",
222            "pir_compact_pairs",
223            "pir_count_children",
224            "pir_write_child_offsets",
225            "pir_gather_children",
226            "pir_build_graph_child_counts",
227            "pir_sum_counts",
228            "pir_emit_nodes_and_ids",
229            "pir_update_counts",
230        ],
231    },
232    KernelModuleSpec {
233        cu_name: "cnf",
234        module_name: "xlog_cnf",
235        kernels: &[
236            "cnf_reachability_init",
237            "cnf_reachability_bfs",
238            "cnf_mark_leaf_choice",
239            "cnf_assign_leaf_var",
240            "cnf_assign_choice_var",
241            "cnf_mark_node_vars",
242            "cnf_count_clauses",
243            "cnf_capture_last_counts",
244            "cnf_compute_leaf_choice_totals",
245            "cnf_compute_totals",
246            "cnf_assign_node_var",
247            "cnf_emit_clauses",
248            "cnf_set_clause_end",
249        ],
250    },
251    KernelModuleSpec {
252        cu_name: "cache",
253        module_name: "xlog_cache",
254        kernels: &[
255            "cache_cnf_hash",
256            "cache_lookup_or_insert",
257            "cache_evict_lru",
258            "cache_store_u8",
259            "cache_store_u32",
260            "cache_store_i32",
261            "cache_store_f64",
262            "cache_store_meta",
263        ],
264    },
265    KernelModuleSpec {
266        cu_name: "weights",
267        module_name: "xlog_weights",
268        kernels: &[
269            "weights_fill_leaf",
270            "weights_fill_choice",
271            "weights_count_lift_exact",
272            "weights_set_evidence_from_nodes",
273            "weights_apply_evidence",
274            "weights_map_nodes_to_vars",
275            "weights_force_var_false",
276            "weights_restore_var_false",
277            "weights_force_var_true",
278            "weights_restore_var_true",
279            "weights_copy_slot_to_batch",
280            "weights_apply_query_vars",
281            "weights_restore_query_vars",
282            "weights_apply_query_vars_false_batched",
283            "weights_restore_query_vars_false_batched",
284            "weights_apply_query_vars_true_batched",
285            "weights_restore_query_vars_true_batched",
286        ],
287    },
288    KernelModuleSpec {
289        cu_name: "circuit",
290        module_name: "xlog_circuit",
291        kernels: &[
292            "xgcf_forward_level",
293            "xgcf_backward_level_propagate",
294            "xgcf_backward_level_decision_grad",
295            "xgcf_backward_level_lit_grad",
296            "xgcf_free_var_apply_grad",
297            "xgcf_free_var_reduce_stage",
298            "xgcf_add_scalar",
299            "xgcf_forward_level_cached",
300            "xgcf_eval_all_levels_cached",
301            "xgcf_eval_all_levels_cached_batched",
302            "xgcf_backward_level_propagate_cached",
303            "xgcf_backward_level_decision_grad_cached",
304            "xgcf_backward_level_lit_grad_cached",
305            "xgcf_backward_all_levels_cached",
306            "xgcf_backward_all_levels_cached_batched",
307            "xgcf_free_var_apply_grad_cached",
308            "xgcf_free_var_reduce_stage_cached",
309            "xgcf_add_scalar_cached",
310            "xgcf_set_root_adj_cached_batched",
311            "xgcf_copy_root_cached",
312            "xgcf_copy_root_cached_meta",
313            "xgcf_copy_root_cached_meta_batched",
314        ],
315    },
316    KernelModuleSpec {
317        cu_name: "mc_sample",
318        module_name: "xlog_mc_sample",
319        kernels: &["mc_sample_bernoulli"],
320    },
321    KernelModuleSpec {
322        cu_name: "mc_eval",
323        module_name: "xlog_mc_eval",
324        kernels: &[
325            "mc_eval_mask_var",
326            "mc_eval_mask_ad_choice",
327            "mc_eval_query_evidence_truth",
328            "mc_accumulate_counts",
329        ],
330    },
331    KernelModuleSpec {
332        cu_name: "arith",
333        module_name: "xlog_arith",
334        kernels: &[
335            "arith_binary_i64",
336            "arith_binary_i32",
337            "arith_binary_u64",
338            "arith_binary_u32",
339            "arith_binary_f64",
340            "arith_binary_f32",
341            "arith_abs_i64",
342            "arith_abs_i32",
343            "arith_abs_f64",
344            "arith_abs_f32",
345            "arith_pow_f64",
346            "arith_cast",
347            "arith_fill_const_u32",
348            "arith_fill_const_u64",
349            "arith_fill_const_i64",
350            "arith_fill_const_i32",
351            "arith_fill_const_f64",
352            "arith_fill_const_f32",
353            "arith_fill_const_u8",
354            "arith_select_i64",
355            "arith_select_i32",
356            "arith_select_u64",
357            "arith_select_u32",
358            "arith_select_f64",
359            "arith_select_f32",
360        ],
361    },
362    KernelModuleSpec {
363        cu_name: "sat",
364        module_name: "xlog_sat",
365        kernels: &[
366            "sat_cdcl_solve",
367            "sat_check_model",
368            "sat_proof_mark_needed",
369            "sat_proof_check",
370            "sat_assert_status",
371            "sat_assert_ok",
372            "sat_xgcf_cnf_counts",
373            "sat_xgcf_cnf_emit",
374            "sat_xgcf_cnf_capture_last_counts",
375            "sat_xgcf_cnf_compute_totals",
376            "sat_cnf_write_terminator",
377            "sat_cnf_copy_into",
378            "sat_shift_offsets",
379            "sat_xgcf_write_root_unit_clause",
380            "sat_not_phi_counts",
381            "sat_emit_not_phi",
382        ],
383    },
384    KernelModuleSpec {
385        cu_name: "d4",
386        module_name: "xlog_d4",
387        kernels: &[
388            "d4_validate_cnf",
389            "d4_levelize_counts",
390            "d4_levelize_emit",
391            "d4_frontier_prepare",
392            "d4_frontier_expand",
393            "d4_frontier_prepare_dense",
394            "d4_frontier_expand_dense",
395            "d4_compile_count",
396            "d4_compile_emit",
397            "d4_capture_emit_meta",
398            "d4_support_level",
399            "d4_support_set_root_bits",
400            "d4_smooth_count",
401            "d4_smooth_wrapper_counts",
402            "d4_smooth_wrapper_edge_counts_or",
403            "d4_smooth_wrapper_edge_counts_dec",
404            "d4_smooth_init_nodes",
405            "d4_smooth_emit_level",
406            "d4_smooth_check_edge_cap",
407            "d4_mark_vars_in_clauses",
408            "d4_mark_vars_in_circuit",
409            "d4_build_free_var_mask",
410            "d4_assert_u32_eq",
411            "d4_assert_bitset_var",
412            "d4_assert_dense_var",
413            "d4_assert_leaf_root_and_degree",
414        ],
415    },
416    KernelModuleSpec {
417        cu_name: "neural",
418        module_name: "xlog_neural",
419        kernels: &[
420            "neural_fill_ad_chain_f32",
421            "neural_scatter_ad_chain_grads_f32",
422        ],
423    },
424    KernelModuleSpec {
425        cu_name: "ilp",
426        module_name: "xlog_ilp",
427        kernels: &[
428            "extract_nonzero_indices",
429            "ilp_mark_selected_ids_u32",
430            "ilp_mark_selected_ids_i32",
431            "ilp_mark_selected_ids_i64",
432            "ilp_mark_selected_ids_u64",
433            "ilp_validate_selected_ids_u32",
434            "ilp_validate_selected_ids_i32",
435            "ilp_validate_selected_ids_i64",
436            "ilp_validate_selected_ids_u64",
437            "ilp_broadcast_candidate_flag",
438            "ilp_coo_fill_from_mask",
439            "ilp_csr_histogram",
440            "ilp_reduce_sum_f32",
441            "ilp_reduce_sum_f64",
442        ],
443    },
444    KernelModuleSpec {
445        cu_name: "ilp_credit",
446        module_name: "xlog_ilp_credit",
447        kernels: &[
448            "ilp_coo_fill",
449            "ilp_credit_forward_f32",
450            "ilp_credit_forward_f64",
451            "ilp_credit_backward_f32",
452            "ilp_credit_backward_f64",
453        ],
454    },
455    KernelModuleSpec {
456        cu_name: "ilp_exact",
457        module_name: "xlog_ilp_exact",
458        kernels: &[
459            "ilp_exact_score",
460            "ilp_exact_score_u32",
461            "ilp_exact_score_chain_smem",
462            "ilp_exact_score_chain_smem_u32",
463            "ilp_exact_select_topk",
464        ],
465    },
466    KernelModuleSpec {
467        cu_name: "epistemic",
468        module_name: "xlog_epistemic",
469        kernels: &[
470            "epistemic_generate_candidate_assumptions_u8",
471            "epistemic_propagate_candidates_u8",
472            "epistemic_validate_candidate_bits_u8",
473            "epistemic_populate_model_membership_u8",
474            "epistemic_populate_model_membership_from_tuple_source_u8",
475            "epistemic_populate_model_membership_from_tuple_source_arity1_u8",
476            "epistemic_populate_model_membership_from_tuple_source_arity2_u8",
477            "epistemic_populate_model_membership_from_tuple_source_arity3_u8",
478            "epistemic_populate_model_membership_from_tuple_source_arity_n_u8",
479            "epistemic_validate_world_views_u8",
480            "epistemic_validate_constraints_u8",
481            "epistemic_materialize_accepted_candidates_u8",
482            "epistemic_materialize_final_result_flags_u8",
483            "epistemic_build_final_tuple_row_map_u8",
484            "epistemic_close_final_tuple_rejections_u8",
485            "epistemic_materialize_final_tuple_column_u8",
486        ],
487    },
488    KernelModuleSpec {
489        cu_name: "wcoj",
490        module_name: "xlog_wcoj",
491        kernels: &[
492            "wcoj_build_metadata_mark_boundaries_u32",
493            "wcoj_build_metadata_mark_boundaries_u64",
494            "wcoj_build_metadata_scatter_u32",
495            "wcoj_build_metadata_scatter_u64",
496            "wcoj_triangle_build_hg_work_plan_u32",
497            "wcoj_triangle_count_hg_u32",
498            "wcoj_triangle_groupby_root_count_hg_u32",
499            "wcoj_triangle_groupby_root_sum_hg_u32",
500            "wcoj_triangle_groupby_root_min_hg_u32",
501            "wcoj_triangle_groupby_root_max_hg_u32",
502            "wcoj_triangle_materialize_hg_u32",
503            "wcoj_triangle_build_hg_work_plan_u64",
504            "wcoj_triangle_count_hg_u64",
505            "wcoj_triangle_groupby_root_count_hg_u64",
506            "wcoj_triangle_groupby_root_sum_hg_u64",
507            "wcoj_triangle_groupby_root_min_hg_u64",
508            "wcoj_triangle_groupby_root_max_hg_u64",
509            "wcoj_groupby_root_segment_sum_counts_u32",
510            "wcoj_groupby_root_segment_sum_values_u64",
511            "wcoj_groupby_root_segment_min_values_u64",
512            "wcoj_groupby_root_segment_max_values_u64",
513            "wcoj_triangle_materialize_hg_u64",
514            "wcoj_triangle_count_hg_cached_u32",
515            "wcoj_triangle_materialize_hg_cached_u32",
516            "wcoj_scan_hg_block_counts_u32",
517            "wcoj_compute_total",
518            "wcoj_layout_check_sorted_unique_u32",
519            "wcoj_layout_check_sorted_unique_u64",
520            "wcoj_4cycle_build_e2_work_prefix_u32",
521            "wcoj_4cycle_build_hg_work_plan_u32",
522            "wcoj_4cycle_count_hg_u32",
523            "wcoj_4cycle_groupby_root_count_hg_u32",
524            "wcoj_4cycle_groupby_root_sum_hg_u32",
525            "wcoj_4cycle_groupby_root_min_hg_u32",
526            "wcoj_4cycle_groupby_root_max_hg_u32",
527            "wcoj_4cycle_materialize_hg_u32",
528            "wcoj_4cycle_build_e2_work_prefix_u64",
529            "wcoj_4cycle_build_hg_work_plan_u64",
530            "wcoj_4cycle_count_hg_u64",
531            "wcoj_4cycle_groupby_root_count_hg_u64",
532            "wcoj_4cycle_materialize_hg_u64",
533            // General-arity WCOJ clique kernel family (k=5..8 from
534            // a single C++ template; ABI wrappers below are
535            // template-call-only for source-auditability).
536            "wcoj_clique5_count_hg_u32",
537            "wcoj_clique5_materialize_hg_u32",
538            "wcoj_clique5_count_hg_u64",
539            "wcoj_clique5_materialize_hg_u64",
540            "wcoj_clique6_count_hg_u32",
541            "wcoj_clique6_materialize_hg_u32",
542            "wcoj_clique6_count_hg_u64",
543            "wcoj_clique6_materialize_hg_u64",
544            "wcoj_clique7_count_hg_u32",
545            "wcoj_clique7_materialize_hg_u32",
546            "wcoj_clique7_count_hg_u64",
547            "wcoj_clique7_materialize_hg_u64",
548            "wcoj_clique8_count_hg_u32",
549            "wcoj_clique8_materialize_hg_u32",
550            "wcoj_clique8_count_hg_u64",
551            "wcoj_clique8_materialize_hg_u64",
552            // Aggregate-fused K-clique group-by-root count kernels (u32
553            // width-class, K=5/6).
554            "wcoj_clique5_groupby_root_count_hg_u32",
555            "wcoj_clique6_groupby_root_count_hg_u32",
556            // GPU Free Join level-synchronous frontier engine primitives.
557            "fj_expand_work_prefix_u32",
558            "fj_expand_count_u32",
559            "fj_expand_emit_u32",
560            "fj_probe_refine_u32",
561            // u64 width-class twins (work prefix is
562            // width-agnostic and shared).
563            "fj_expand_count_u64",
564            "fj_expand_emit_u64",
565            "fj_probe_refine_u64",
566            // Factorized count epilogue (width-agnostic).
567            "fj_count_multiplicity",
568            // D3 S3 spike — factorized recursive delta novel-set
569            // pipeline (dense-domain bitmap union–diff).
570            "fj_delta_range_u32",
571            "fj_delta_mark_u32",
572            "fj_delta_subtract_u32",
573            "fj_delta_popcount",
574            "fj_delta_emit_u32",
575            "fj_delta_max_u32",
576            // D3 sparse-domain spike — hash-set novel pipeline.
577            "fj_delta_sparse_estimate",
578            "fj_delta_sparse_load_r",
579            "fj_delta_sparse_insert_candidates",
580            "fj_delta_sparse_mark",
581            "fj_delta_sparse_emit",
582        ],
583    },
584    KernelModuleSpec {
585        cu_name: "mc_resident",
586        module_name: "xlog_mc_resident",
587        kernels: &["mc_resident_engine"],
588    },
589    KernelModuleSpec {
590        cu_name: "joint_solve",
591        module_name: "xlog_joint_solve",
592        kernels: &[
593            "joint_label_feasibility",
594            "joint_label_top2",
595            "joint_component_enumerate",
596            "joint_label_memoized",
597        ],
598    },
599];
600
601#[cfg(test)]
602mod tests {
603    use super::*;
604
605    #[test]
606    fn kernel_modules_matches_cu_names() {
607        assert_eq!(
608            KERNEL_MODULES.len(),
609            KERNEL_CU_NAMES.len(),
610            "KERNEL_MODULES length ({}) != KERNEL_CU_NAMES length ({})",
611            KERNEL_MODULES.len(),
612            KERNEL_CU_NAMES.len(),
613        );
614        for (i, spec) in KERNEL_MODULES.iter().enumerate() {
615            assert_eq!(
616                spec.cu_name, KERNEL_CU_NAMES[i],
617                "KERNEL_MODULES[{}].cu_name = {:?}, expected {:?}",
618                i, spec.cu_name, KERNEL_CU_NAMES[i],
619            );
620        }
621    }
622
623    #[test]
624    fn kernel_modules_count_is_26() {
625        assert_eq!(KERNEL_MODULES.len(), 26);
626    }
627
628    #[test]
629    fn all_kernel_entries_are_non_empty() {
630        for spec in KERNEL_MODULES {
631            assert!(
632                !spec.kernels.is_empty(),
633                "module {:?} has no kernel entries",
634                spec.cu_name,
635            );
636            assert!(
637                !spec.module_name.is_empty(),
638                "module {:?} has empty module_name",
639                spec.cu_name,
640            );
641        }
642    }
643}