Skip to main content

cubecl_cpp/metal/
plane.rs

1use 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});