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::{TriviallyUnrollable, synchronizes},
12    prelude::*,
13    types::{VectorType, scalar::BoolType},
14};
15
16#[cube_op(name = "plane.elect")]
17#[result_ty(fixed = BoolType::get(ctx).into())]
18#[op_traits(CanMaterialize, NoMemoryEffect)]
19pub struct ElectOp {}
20synchronizes!(ElectOp, SyncScope::Plane);
21
22macro_rules! unary_plane_op {
23    ($name: literal, $ty: ident) => {
24        #[cube_op(name = $name)]
25        #[result_ty(same_as = input)]
26        #[op_interfaces(TriviallyUnrollable)]
27        #[op_traits(CanMaterialize, NoMemoryEffect)]
28        pub struct $ty {
29            pub input: Value,
30        }
31        synchronizes!($ty, SyncScope::Plane);
32    };
33}
34
35unary_plane_op!("plane.all", AllOp);
36unary_plane_op!("plane.any", AnyOp);
37unary_plane_op!("plane.i_sum", ISumOp);
38unary_plane_op!("plane.f_sum", FSumOp);
39unary_plane_op!("plane.inclusive_i_sum", InclusiveISumOp);
40unary_plane_op!("plane.inclusive_f_sum", InclusiveFSumOp);
41unary_plane_op!("plane.exclusive_i_sum", ExclusiveISumOp);
42unary_plane_op!("plane.exclusive_f_sum", ExclusiveFSumOp);
43unary_plane_op!("plane.i_prod", IProdOp);
44unary_plane_op!("plane.f_prod", FProdOp);
45unary_plane_op!("plane.inclusive_i_prod", InclusiveIProdOp);
46unary_plane_op!("plane.inclusive_f_prod", InclusiveFProdOp);
47unary_plane_op!("plane.exclusive_i_prod", ExclusiveIProdOp);
48unary_plane_op!("plane.exclusive_f_prod", ExclusiveFProdOp);
49unary_plane_op!("plane.s_min", SMinOp);
50unary_plane_op!("plane.u_min", UMinOp);
51unary_plane_op!("plane.f_min", FMinOp);
52unary_plane_op!("plane.s_max", SMaxOp);
53unary_plane_op!("plane.u_max", UMaxOp);
54unary_plane_op!("plane.f_max", FMaxOp);
55
56#[cube_op(name = "plane.ballot")]
57#[result_ty(fixed = ballot_ty(ctx))]
58#[op_interfaces(TriviallyUnrollable)]
59#[op_traits(CanMaterialize, NoMemoryEffect)]
60pub struct BallotOp {
61    pub input: Value,
62}
63synchronizes!(BallotOp, SyncScope::Plane);
64
65fn ballot_ty(ctx: &Context) -> TypeHandle {
66    let u32 = IntegerType::get(ctx, 32, Signedness::Unsigned);
67    VectorType::get(ctx, u32.into(), 4).into()
68}
69
70#[cube_op(name = "plane.broadcast")]
71#[result_ty(same_as = input)]
72#[op_interfaces(TriviallyUnrollable)]
73#[op_traits(CanMaterialize, NoMemoryEffect)]
74pub struct BroadcastOp {
75    pub input: Value,
76    pub lane: IndexAttr,
77}
78synchronizes!(BroadcastOp, SyncScope::Plane);
79
80#[cube_op(name = "plane.shuffle")]
81#[result_ty(same_as = input)]
82#[op_interfaces(TriviallyUnrollable)]
83#[op_traits(CanMaterialize, NoMemoryEffect)]
84pub struct ShuffleOp {
85    pub input: Value,
86    pub lane: Value,
87}
88synchronizes!(ShuffleOp, SyncScope::Plane);
89
90#[cube_op(name = "plane.shuffle_xor")]
91#[result_ty(same_as = input)]
92#[op_interfaces(TriviallyUnrollable)]
93#[op_traits(CanMaterialize, NoMemoryEffect)]
94pub struct ShuffleXorOp {
95    pub input: Value,
96    pub mask: Value,
97}
98synchronizes!(ShuffleXorOp, SyncScope::Plane);
99
100#[cube_op(name = "plane.shuffle_up")]
101#[result_ty(same_as = input)]
102#[op_interfaces(TriviallyUnrollable)]
103#[op_traits(CanMaterialize, NoMemoryEffect)]
104pub struct ShuffleUpOp {
105    pub input: Value,
106    pub delta: Value,
107}
108synchronizes!(ShuffleUpOp, SyncScope::Plane);
109
110#[cube_op(name = "plane.shuffle_down")]
111#[result_ty(same_as = input)]
112#[op_interfaces(TriviallyUnrollable)]
113#[op_traits(CanMaterialize, NoMemoryEffect)]
114pub struct ShuffleDownOp {
115    pub input: Value,
116    pub delta: Value,
117}
118synchronizes!(ShuffleDownOp, SyncScope::Plane);
119
120#[cube_op(name = "plane.uniform_load")]
121#[result_ty(from_inputs = ptr_value_ty)]
122#[op_interfaces(TriviallyUnrollable)]
123#[op_traits(CanMaterialize)]
124pub struct UniformLoadOp {
125    #[operand(ptr_read)]
126    pub ptr: Value,
127}
128synchronizes!(UniformLoadOp, SyncScope::Plane);
129
130#[cube_op(name = "plane.atomic_uniform_load")]
131#[result_ty(from_inputs = ptr_value_ty)]
132#[op_interfaces(TriviallyUnrollable)]
133#[op_traits(CanMaterialize)]
134pub struct AtomicUniformLoadOp {
135    #[operand(ptr_read)]
136    pub ptr: Value,
137}
138synchronizes!(AtomicUniformLoadOp, SyncScope::Plane);