cubecl_ir/dialect/
plane.rs1use 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);