cubecl_cpp/metal/
plane.rs1use cubecl::prelude::*;
2use cubecl_core as cubecl;
3use cubecl_core::ir::{cube_op, dialect::plane::*};
4use pliron::{
5 builtin::types::{IntegerType, Signedness},
6 derive::op_interface_impl,
7 value::Value,
8};
9
10use crate::{
11 metal::metal_op_with_out,
12 shared::{lowering::LowerOp, unroll::unrolling},
13 target::Metal,
14};
15
16metal_op_with_out!(BroadcastOp, |op, ctx| {
17 let val = op.input(ctx).name(ctx);
18 let lane = op.lane(ctx).0;
19 format!("simd_shuffle({val}, {lane});")
20});
21
22metal_op_with_out!(ShuffleOp, |op, ctx| {
23 let val = op.input(ctx).name(ctx);
24 let lane = op.lane(ctx).name(ctx);
25 format!("simd_shuffle({val}, {lane});")
26});
27
28metal_op_with_out!(ShuffleXorOp, |op, ctx| {
29 let val = op.input(ctx).name(ctx);
30 let mask = op.mask(ctx).name(ctx);
31 format!("simd_shuffle_xor({val}, {mask});")
32});
33
34metal_op_with_out!(ShuffleUpOp, |op, ctx| {
35 let val = op.input(ctx).name(ctx);
36 let delta = op.delta(ctx).name(ctx);
37 format!("simd_shuffle_up({val}, {delta});")
38});
39
40metal_op_with_out!(ShuffleDownOp, |op, ctx| {
41 let val = op.input(ctx).name(ctx);
42 let delta = op.delta(ctx).name(ctx);
43 format!("simd_shuffle_down({val}, {delta});")
44});
45
46metal_op_with_out!(ElectOp, |_, _| { "simd_is_first()".into() });
47
48metal_op_with_out!(AllOp, |op, ctx| {
49 let val = op.input(ctx).name(ctx);
50 format!("simd_all({val});")
51});
52
53metal_op_with_out!(AnyOp, |op, ctx| {
54 let val = op.input(ctx).name(ctx);
55 format!("simd_any({val});")
56});
57
58#[cube_op(name = "msl.ballot")]
59#[result_ty(fixed = IntegerType::get(ctx, 64, Signedness::Unsigned).to_handle())]
60pub struct MslBallotOp {
61 input: Value,
62}
63
64metal_op_with_out!(MslBallotOp, |op, ctx| {
65 let val = op.input(ctx).name(ctx);
66 format!("uint64_t(simd_ballot({val}));")
67});
68
69#[cube]
70fn msl_ballot(value: bool) -> u64 {
71 intrinsic!(|scope| {
72 let value = value.read_value(scope);
73 let ballot = MslBallotOp::new(scope.ctx_mut(), value);
74 scope.register_with_result(&ballot).into()
75 })
76}
77
78#[cube]
79fn ballot(value: bool) -> Vector<u32, Const<4>> {
80 let mut out = Vector::<u64, Const<2>>::zero();
81 out.insert(0usize, msl_ballot(value));
82 Vector::reinterpret(out)
83}
84
85#[op_interface_impl]
86impl LowerOp<Metal> for BallotOp {
87 fn lower(&self, scope: &Scope) -> Vec<Value> {
88 let value = self.input(scope.ctx()).into();
89 vec![ballot::expand(scope, value).read_value(scope)]
90 }
91}
92
93unrolling!(ISumOp);
94metal_op_with_out!(ISumOp, |op, ctx| {
95 let val = op.input(ctx).name(ctx);
96 format!("simd_sum({val})")
97});
98unrolling!(FSumOp);
99metal_op_with_out!(FSumOp, |op, ctx| {
100 let val = op.input(ctx).name(ctx);
101 format!("simd_sum({val})")
102});
103
104unrolling!(IProdOp);
105metal_op_with_out!(IProdOp, |op, ctx| {
106 let val = op.input(ctx).name(ctx);
107 format!("simd_product({val})")
108});
109unrolling!(FProdOp);
110metal_op_with_out!(FProdOp, |op, ctx| {
111 let val = op.input(ctx).name(ctx);
112 format!("simd_product({val})")
113});
114
115unrolling!(SMinOp);
116metal_op_with_out!(SMinOp, |op, ctx| {
117 let val = op.input(ctx).name(ctx);
118 format!("simd_min({val})")
119});
120unrolling!(UMinOp);
121metal_op_with_out!(UMinOp, |op, ctx| {
122 let val = op.input(ctx).name(ctx);
123 format!("simd_min({val})")
124});
125unrolling!(FMinOp);
126metal_op_with_out!(FMinOp, |op, ctx| {
127 let val = op.input(ctx).name(ctx);
128 format!("simd_min({val})")
129});
130
131unrolling!(SMaxOp);
132metal_op_with_out!(SMaxOp, |op, ctx| {
133 let val = op.input(ctx).name(ctx);
134 format!("simd_max({val})")
135});
136unrolling!(UMaxOp);
137metal_op_with_out!(UMaxOp, |op, ctx| {
138 let val = op.input(ctx).name(ctx);
139 format!("simd_max({val})")
140});
141unrolling!(FMaxOp);
142metal_op_with_out!(FMaxOp, |op, ctx| {
143 let val = op.input(ctx).name(ctx);
144 format!("simd_max({val})")
145});
146
147unrolling!(InclusiveISumOp);
148metal_op_with_out!(InclusiveISumOp, |op, ctx| {
149 let val = op.input(ctx).name(ctx);
150 format!("simd_prefix_inclusive_sum({val})")
151});
152unrolling!(InclusiveFSumOp);
153metal_op_with_out!(InclusiveFSumOp, |op, ctx| {
154 let val = op.input(ctx).name(ctx);
155 format!("simd_prefix_inclusive_sum({val})")
156});
157
158unrolling!(InclusiveIProdOp);
159metal_op_with_out!(InclusiveIProdOp, |op, ctx| {
160 let val = op.input(ctx).name(ctx);
161 format!("simd_prefix_inclusive_product({val})")
162});
163unrolling!(InclusiveFProdOp);
164metal_op_with_out!(InclusiveFProdOp, |op, ctx| {
165 let val = op.input(ctx).name(ctx);
166 format!("simd_prefix_inclusive_product({val})")
167});
168
169unrolling!(ExclusiveISumOp);
170metal_op_with_out!(ExclusiveISumOp, |op, ctx| {
171 let val = op.input(ctx).name(ctx);
172 format!("simd_prefix_exclusive_sum({val})")
173});
174unrolling!(ExclusiveFSumOp);
175metal_op_with_out!(ExclusiveFSumOp, |op, ctx| {
176 let val = op.input(ctx).name(ctx);
177 format!("simd_prefix_exclusive_sum({val})")
178});
179
180unrolling!(ExclusiveIProdOp);
181metal_op_with_out!(ExclusiveIProdOp, |op, ctx| {
182 let val = op.input(ctx).name(ctx);
183 format!("simd_prefix_exclusive_product({val})")
184});
185unrolling!(ExclusiveFProdOp);
186metal_op_with_out!(ExclusiveFProdOp, |op, ctx| {
187 let val = op.input(ctx).name(ctx);
188 format!("simd_prefix_exclusive_product({val})")
189});