Skip to main content

ruprim/
collective.rs

1//! Operators shared by device, block and logical-warp collectives.
2
3use ruda_kernel::dsl as kernel_dsl;
4use ruda_kernel::dsl::prelude::*;
5
6pub(crate) mod merge;
7pub mod radix;
8pub mod record;
9pub mod decompose;
10pub mod tile;
11
12#[ruda]
13pub trait RudaUnaryOp<T: RudaType, U: RudaType>: RudaType {
14    fn apply(&self, value: T) -> U;
15}
16
17/// A binary operator. Scans and reductions require associativity and preserve
18/// operand order; adjacent-difference operations do not require associativity.
19#[ruda]
20pub trait RudaBinaryOp<T: RudaType>: RudaType {
21    fn combine(&self, left: T, right: T) -> T;
22}
23
24/// A strict weak ordering for comparison-based algorithms.
25#[ruda]
26pub trait RudaCompare<T: RudaType>: RudaType {
27    fn before(&self, left: T, right: T) -> bool;
28}
29
30/// An equivalence relation for adjacent-key operations.
31#[ruda]
32pub trait RudaKeyEqual<T: RudaType>: RudaType {
33    fn equal(&self, left: T, right: T) -> bool;
34}
35
36#[derive(Clone, Copy, Debug, RudaType, RudaLaunch)]
37pub struct RudaSum;
38
39#[ruda]
40impl<T: Numeric> RudaBinaryOp<T> for RudaSum {
41    fn combine(&self, left: T, right: T) -> T {
42        left + right
43    }
44}
45
46#[derive(Clone, Copy, Debug, RudaType, RudaLaunch)]
47pub struct RudaProduct;
48
49#[ruda]
50impl<T: Numeric> RudaBinaryOp<T> for RudaProduct {
51    fn combine(&self, left: T, right: T) -> T {
52        left * right
53    }
54}
55
56#[derive(Clone, Copy, Debug, RudaType, RudaLaunch)]
57pub struct RudaSubtract;
58
59#[ruda]
60impl<T: Numeric> RudaBinaryOp<T> for RudaSubtract {
61    fn combine(&self, left: T, right: T) -> T { left - right }
62}
63
64#[derive(Clone, Copy, Debug, RudaType, RudaLaunch)]
65pub struct RudaMinimum;
66
67#[ruda]
68impl<T: Numeric> RudaBinaryOp<T> for RudaMinimum {
69    fn combine(&self, left: T, right: T) -> T {
70        if right < left { right } else { left }
71    }
72}
73
74#[derive(Clone, Copy, Debug, RudaType, RudaLaunch)]
75pub struct RudaMaximum;
76
77#[ruda]
78impl<T: Numeric> RudaBinaryOp<T> for RudaMaximum {
79    fn combine(&self, left: T, right: T) -> T {
80        if left < right { right } else { left }
81    }
82}
83
84#[derive(Clone, Copy, Debug, RudaType, RudaLaunch)]
85pub struct RudaAscending;
86
87#[ruda]
88impl<T: Numeric> RudaCompare<T> for RudaAscending {
89    fn before(&self, left: T, right: T) -> bool {
90        left < right
91    }
92}
93
94#[derive(Clone, Copy, Debug, RudaType, RudaLaunch)]
95pub struct RudaDescending;
96
97#[ruda]
98impl<T: Numeric> RudaCompare<T> for RudaDescending {
99    fn before(&self, left: T, right: T) -> bool {
100        left > right
101    }
102}
103
104#[derive(Clone, Copy, Debug, RudaType, RudaLaunch)]
105pub struct RudaEqual;
106
107macro_rules! clone_operator_launch {
108    ($($name:ident),* $(,)?) => {$(
109        impl<R: Runtime> Clone for $name<R> {
110            fn clone(&self) -> Self { Self::new() }
111        }
112    )*};
113}
114
115clone_operator_launch!(RudaSumLaunch, RudaProductLaunch, RudaSubtractLaunch,
116    RudaMinimumLaunch, RudaMaximumLaunch, RudaAscendingLaunch, RudaDescendingLaunch,
117    RudaEqualLaunch);
118
119#[ruda]
120impl<T: Scalar> RudaKeyEqual<T> for RudaEqual {
121    fn equal(&self, left: T, right: T) -> bool {
122        left == right
123    }
124}