Skip to main content

cubecl_ir/dialect/
plane.rs

1use cubecl_macros_internal::{cube_op, op_traits};
2use pliron::{
3    builtin::types::{IntegerType, Signedness},
4    r#type::TypeHandle,
5};
6
7use crate::{
8    CanMaterialize, NoMemoryEffect,
9    attributes::IndexAttr,
10    dialect::{ptr_value_ty, synchronization::SyncScope},
11    interfaces::{
12        TriviallyUnrollable, synchronizes,
13        uniformity::{UniformOpInterface, Uniformity},
14    },
15    prelude::*,
16    types::{VectorType, scalar::BoolType},
17};
18
19#[cube_op(name = "plane.elect")]
20#[result_ty(fixed = BoolType::get(ctx).into())]
21#[op_traits(CanMaterialize, NoMemoryEffect)]
22pub struct ElectOp {}
23synchronizes!(ElectOp, SyncScope::Plane);
24
25#[op_interface_impl]
26impl UniformOpInterface for ElectOp {
27    fn uniformity(&self, _ctx: &Context, _operands: &[Uniformity]) -> Uniformity {
28        Uniformity::None
29    }
30}
31
32macro_rules! unary_plane_op {
33    ($name: literal, $ty: ident) => {
34        #[cube_op(name = $name)]
35        #[result_ty(same_as = input)]
36        #[op_interfaces(TriviallyUnrollable)]
37        #[op_traits(CanMaterialize, NoMemoryEffect)]
38        pub struct $ty {
39            pub input: Value,
40        }
41        synchronizes!($ty, SyncScope::Plane);
42
43        #[op_interface_impl]
44        impl UniformOpInterface for $ty {
45            fn uniformity(&self, _ctx: &Context, operands: &[Uniformity]) -> Uniformity {
46                operands[0].max(Uniformity::Plane)
47            }
48        }
49    };
50}
51
52macro_rules! nonuniform_unary_plane_op {
53    ($name: literal, $ty: ident) => {
54        #[cube_op(name = $name)]
55        #[result_ty(same_as = input)]
56        #[op_interfaces(TriviallyUnrollable)]
57        #[op_traits(CanMaterialize, NoMemoryEffect)]
58        pub struct $ty {
59            pub input: Value,
60        }
61        synchronizes!($ty, SyncScope::Plane);
62
63        #[op_interface_impl]
64        impl UniformOpInterface for $ty {
65            fn uniformity(&self, _ctx: &Context, _operands: &[Uniformity]) -> Uniformity {
66                Uniformity::None
67            }
68        }
69    };
70}
71
72unary_plane_op!("plane.all", AllOp);
73unary_plane_op!("plane.any", AnyOp);
74unary_plane_op!("plane.i_sum", ISumOp);
75unary_plane_op!("plane.f_sum", FSumOp);
76nonuniform_unary_plane_op!("plane.inclusive_i_sum", InclusiveISumOp);
77nonuniform_unary_plane_op!("plane.inclusive_f_sum", InclusiveFSumOp);
78nonuniform_unary_plane_op!("plane.exclusive_i_sum", ExclusiveISumOp);
79nonuniform_unary_plane_op!("plane.exclusive_f_sum", ExclusiveFSumOp);
80unary_plane_op!("plane.i_prod", IProdOp);
81unary_plane_op!("plane.f_prod", FProdOp);
82nonuniform_unary_plane_op!("plane.inclusive_i_prod", InclusiveIProdOp);
83nonuniform_unary_plane_op!("plane.inclusive_f_prod", InclusiveFProdOp);
84nonuniform_unary_plane_op!("plane.exclusive_i_prod", ExclusiveIProdOp);
85nonuniform_unary_plane_op!("plane.exclusive_f_prod", ExclusiveFProdOp);
86unary_plane_op!("plane.s_min", SMinOp);
87unary_plane_op!("plane.u_min", UMinOp);
88unary_plane_op!("plane.f_min", FMinOp);
89unary_plane_op!("plane.s_max", SMaxOp);
90unary_plane_op!("plane.u_max", UMaxOp);
91unary_plane_op!("plane.f_max", FMaxOp);
92
93#[cube_op(name = "plane.ballot")]
94#[result_ty(fixed = ballot_ty(ctx))]
95#[op_interfaces(TriviallyUnrollable)]
96#[op_traits(CanMaterialize, NoMemoryEffect)]
97pub struct BallotOp {
98    pub input: Value,
99}
100synchronizes!(BallotOp, SyncScope::Plane);
101
102#[op_interface_impl]
103impl UniformOpInterface for BallotOp {
104    fn uniformity(&self, _ctx: &Context, operands: &[Uniformity]) -> Uniformity {
105        operands[0].max(Uniformity::Plane)
106    }
107}
108
109fn ballot_ty(ctx: &Context) -> TypeHandle {
110    let u32 = IntegerType::get(ctx, 32, Signedness::Unsigned);
111    VectorType::get(ctx, u32.into(), 4).into()
112}
113
114#[cube_op(name = "plane.broadcast")]
115#[result_ty(same_as = input)]
116#[op_interfaces(TriviallyUnrollable)]
117#[op_traits(CanMaterialize, NoMemoryEffect)]
118pub struct BroadcastOp {
119    pub input: Value,
120    pub lane: IndexAttr,
121}
122synchronizes!(BroadcastOp, SyncScope::Plane);
123
124#[op_interface_impl]
125impl UniformOpInterface for BroadcastOp {
126    fn uniformity(&self, _ctx: &Context, operands: &[Uniformity]) -> Uniformity {
127        operands[0].max(Uniformity::Plane)
128    }
129}
130
131#[cube_op(name = "plane.shuffle")]
132#[result_ty(same_as = input)]
133#[op_interfaces(TriviallyUnrollable)]
134#[op_traits(CanMaterialize, NoMemoryEffect)]
135pub struct ShuffleOp {
136    pub input: Value,
137    pub lane: Value,
138}
139synchronizes!(ShuffleOp, SyncScope::Plane);
140
141#[op_interface_impl]
142impl UniformOpInterface for ShuffleOp {
143    fn uniformity(&self, _ctx: &Context, operands: &[Uniformity]) -> Uniformity {
144        operands[0]
145    }
146}
147
148#[cube_op(name = "plane.shuffle_xor")]
149#[result_ty(same_as = input)]
150#[op_interfaces(TriviallyUnrollable)]
151#[op_traits(CanMaterialize, NoMemoryEffect)]
152pub struct ShuffleXorOp {
153    pub input: Value,
154    pub mask: Value,
155}
156synchronizes!(ShuffleXorOp, SyncScope::Plane);
157
158#[op_interface_impl]
159impl UniformOpInterface for ShuffleXorOp {
160    fn uniformity(&self, _ctx: &Context, operands: &[Uniformity]) -> Uniformity {
161        operands[0]
162    }
163}
164
165#[cube_op(name = "plane.shuffle_up")]
166#[result_ty(same_as = input)]
167#[op_interfaces(TriviallyUnrollable)]
168#[op_traits(CanMaterialize, NoMemoryEffect)]
169pub struct ShuffleUpOp {
170    pub input: Value,
171    pub delta: Value,
172}
173synchronizes!(ShuffleUpOp, SyncScope::Plane);
174
175#[op_interface_impl]
176impl UniformOpInterface for ShuffleUpOp {
177    fn uniformity(&self, _ctx: &Context, operands: &[Uniformity]) -> Uniformity {
178        operands[0]
179    }
180}
181
182#[cube_op(name = "plane.shuffle_down")]
183#[result_ty(same_as = input)]
184#[op_interfaces(TriviallyUnrollable)]
185#[op_traits(CanMaterialize, NoMemoryEffect)]
186pub struct ShuffleDownOp {
187    pub input: Value,
188    pub delta: Value,
189}
190synchronizes!(ShuffleDownOp, SyncScope::Plane);
191
192#[op_interface_impl]
193impl UniformOpInterface for ShuffleDownOp {
194    fn uniformity(&self, _ctx: &Context, operands: &[Uniformity]) -> Uniformity {
195        operands[0]
196    }
197}
198
199#[cube_op(name = "plane.uniform_load")]
200#[result_ty(from_inputs = ptr_value_ty)]
201#[op_interfaces(TriviallyUnrollable)]
202#[op_traits(CanMaterialize)]
203pub struct UniformLoadOp {
204    #[operand(ptr_read)]
205    pub ptr: Value,
206}
207synchronizes!(UniformLoadOp, SyncScope::Plane);
208
209#[op_interface_impl]
210impl UniformOpInterface for UniformLoadOp {
211    fn uniformity(&self, _ctx: &Context, operands: &[Uniformity]) -> Uniformity {
212        operands[0].max(Uniformity::Plane)
213    }
214}
215
216#[cube_op(name = "plane.atomic_uniform_load")]
217#[result_ty(from_inputs = ptr_value_ty)]
218#[op_interfaces(TriviallyUnrollable)]
219#[op_traits(CanMaterialize)]
220pub struct AtomicUniformLoadOp {
221    #[operand(ptr_read)]
222    pub ptr: Value,
223}
224synchronizes!(AtomicUniformLoadOp, SyncScope::Plane);
225
226#[op_interface_impl]
227impl UniformOpInterface for AtomicUniformLoadOp {
228    fn uniformity(&self, _ctx: &Context, _operands: &[Uniformity]) -> Uniformity {
229        Uniformity::Plane
230    }
231}