Skip to main content

cubecl_cpp/metal/
binary.rs

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