cubecl_cpp/metal/
binary.rs1use cubecl_core::{
2 frontend::polyfills::*,
3 ir::{
4 dialect::{cmp::*, math::*},
5 interfaces::TypedExt,
6 prelude::*,
7 },
8 prelude::*,
9};
10
11use crate::{
12 metal::metal_op_with_out,
13 shared::{lowering::LowerOp, unroll::unrolling},
14 target::Metal,
15};
16
17unrolling!(SaturatingSAddOp);
18metal_op_with_out!(SaturatingSAddOp, |op, ctx| {
19 let lhs = op.lhs(ctx).name(ctx);
20 let rhs = op.rhs(ctx).name(ctx);
21 format!("addsat({lhs}, {rhs})")
22});
23unrolling!(SaturatingUAddOp);
24metal_op_with_out!(SaturatingUAddOp, |op, ctx| {
25 let lhs = op.lhs(ctx).name(ctx);
26 let rhs = op.rhs(ctx).name(ctx);
27 format!("addsat({lhs}, {rhs})")
28});
29
30unrolling!(SaturatingSSubOp);
31metal_op_with_out!(SaturatingSSubOp, |op, ctx| {
32 let lhs = op.lhs(ctx).name(ctx);
33 let rhs = op.rhs(ctx).name(ctx);
34 format!("subsat({lhs}, {rhs})")
35});
36unrolling!(SaturatingUSubOp);
37metal_op_with_out!(SaturatingUSubOp, |op, ctx| {
38 let lhs = op.lhs(ctx).name(ctx);
39 let rhs = op.rhs(ctx).name(ctx);
40 format!("subsat({lhs}, {rhs})")
41});
42
43metal_op_with_out!(SMinOp, |op, ctx| {
44 let lhs = op.lhs(ctx).name(ctx);
45 let rhs = op.rhs(ctx).name(ctx);
46 format!("min({lhs}, {rhs})")
47});
48metal_op_with_out!(UMinOp, |op, ctx| {
49 let lhs = op.lhs(ctx).name(ctx);
50 let rhs = op.rhs(ctx).name(ctx);
51 format!("min({lhs}, {rhs})")
52});
53metal_op_with_out!(FMinOp, |op, ctx| {
54 let lhs = op.lhs(ctx).name(ctx);
55 let rhs = op.rhs(ctx).name(ctx);
56 format!("min({lhs}, {rhs})")
57});
58
59metal_op_with_out!(SMaxOp, |op, ctx| {
60 let lhs = op.lhs(ctx).name(ctx);
61 let rhs = op.rhs(ctx).name(ctx);
62 format!("max({lhs}, {rhs})")
63});
64metal_op_with_out!(UMaxOp, |op, ctx| {
65 let lhs = op.lhs(ctx).name(ctx);
66 let rhs = op.rhs(ctx).name(ctx);
67 format!("max({lhs}, {rhs})")
68});
69metal_op_with_out!(FMaxOp, |op, ctx| {
70 let lhs = op.lhs(ctx).name(ctx);
71 let rhs = op.rhs(ctx).name(ctx);
72 format!("max({lhs}, {rhs})")
73});
74
75metal_op_with_out!(PowfOp, |op, ctx| {
76 let lhs = op.lhs(ctx).name(ctx);
77 let rhs = op.rhs(ctx).name(ctx);
78 format!("pow({lhs}, {rhs})")
79});
80
81metal_op_with_out!(PowiOp, |op, ctx| {
82 let lhs = op.lhs(ctx).name(ctx);
83 let rhs = op.rhs(ctx).name(ctx);
84 format!("pow({lhs}, {rhs})")
85});
86
87metal_op_with_out!(HypotOp, |op, ctx| {
88 let lhs = op.lhs(ctx).name(ctx);
89 let rhs = op.rhs(ctx).name(ctx);
90 format!("length(float2({lhs}, {rhs}))")
91});
92
93metal_op_with_out!(RhypotOp, |op, ctx| {
94 let lhs = op.lhs(ctx).name(ctx);
95 let rhs = op.rhs(ctx).name(ctx);
96 format!("rsqrt({lhs} * {lhs} + {rhs} * {rhs})")
97});
98
99#[op_interface_impl]
100impl LowerOp<Metal> for SMulHiOp {
101 fn lower(&self, scope: &Scope) -> Vec<Value> {
102 let ctx = scope.ctx();
103 let lhs = self.lhs(ctx);
104 let val = if lhs.size_bits(ctx) == 32 {
105 expand_s_himul_64(scope, lhs, self.rhs(ctx))
106 } else {
107 expand_himul_sim(scope, lhs, self.rhs(ctx))
108 };
109 vec![val]
110 }
111}
112
113#[op_interface_impl]
114impl LowerOp<Metal> for UMulHiOp {
115 fn lower(&self, scope: &Scope) -> Vec<Value> {
116 let ctx = scope.ctx();
117 let lhs = self.lhs(ctx);
118 let val = if lhs.size_bits(ctx) == 32 {
119 expand_u_himul_64(scope, lhs, self.rhs(ctx))
120 } else {
121 expand_himul_sim(scope, lhs, self.rhs(ctx))
122 };
123 vec![val]
124 }
125}