Skip to main content

hir_ty/mir/eval/shim/
simd.rs

1//! Shim implementation for simd intrinsics
2
3use std::cmp::Ordering;
4
5use crate::consteval::try_const_usize;
6
7use super::*;
8
9impl<'a, 'db> Evaluator<'a, 'db> {
10    fn detect_simd_ty(&self, ty: Ty<'db>) -> Result<'db, (usize, Ty<'db>)> {
11        match ty.kind() {
12            TyKind::Adt(adt_def, subst) => {
13                let len = match subst.as_slice().get(1).and_then(|it| it.konst()) {
14                    Some(len) => len,
15                    _ => {
16                        if let AdtId::StructId(id) = adt_def.def_id() {
17                            let struct_data = id.fields(self.db);
18                            let fields = struct_data.fields();
19                            let Some((first_field, _)) = fields.iter().next() else {
20                                not_supported!("simd type with no field");
21                            };
22                            let field_ty = self.db.field_types(id.into())[first_field]
23                                .ty()
24                                .instantiate(self.interner(), subst)
25                                .skip_norm_wip();
26                            return Ok((fields.len(), field_ty));
27                        }
28                        return Err(MirEvalError::InternalError(
29                            "simd type with no len param".into(),
30                        ));
31                    }
32                };
33                match try_const_usize(self.db, len) {
34                    Some(len) => {
35                        let Some(ty) = subst.as_slice().first().and_then(|it| it.ty()) else {
36                            return Err(MirEvalError::InternalError(
37                                "simd type with no ty param".into(),
38                            ));
39                        };
40                        Ok((len as usize, ty))
41                    }
42                    None => Err(MirEvalError::InternalError(
43                        "simd type with unevaluatable len param".into(),
44                    )),
45                }
46            }
47            _ => Err(MirEvalError::InternalError("simd type which is not a struct".into())),
48        }
49    }
50
51    pub(super) fn exec_simd_intrinsic(
52        &mut self,
53        name: &str,
54        args: &[IntervalAndTy<'db>],
55        _generic_args: GenericArgs<'db>,
56        destination: Interval,
57        _locals: &Locals<'a, 'db>,
58        _span: MirSpan,
59    ) -> Result<'db, ()> {
60        match name {
61            "and" | "or" | "xor" => {
62                let [left, right] = args else {
63                    return Err(MirEvalError::InternalError(
64                        "simd bit op args are not provided".into(),
65                    ));
66                };
67                let result = left
68                    .get(self)?
69                    .iter()
70                    .zip(right.get(self)?)
71                    .map(|(&it, &y)| match name {
72                        "and" => it & y,
73                        "or" => it | y,
74                        "xor" => it ^ y,
75                        _ => unreachable!(),
76                    })
77                    .collect::<Vec<_>>();
78                destination.write_from_bytes(self, &result)
79            }
80            "eq" | "ne" | "lt" | "le" | "gt" | "ge" => {
81                let [left, right] = args else {
82                    return Err(MirEvalError::InternalError("simd args are not provided".into()));
83                };
84                let (len, ty) = self.detect_simd_ty(left.ty)?;
85                let is_signed = matches!(ty.kind(), TyKind::Int(_));
86                let size = left.interval.size / len;
87                let dest_size = destination.size / len;
88                let mut destination_bytes = vec![];
89                let vector = left.get(self)?.chunks(size).zip(right.get(self)?.chunks(size));
90                for (l, r) in vector {
91                    let mut result = Ordering::Equal;
92                    for (l, r) in l.iter().zip(r).rev() {
93                        let it = l.cmp(r);
94                        if it != Ordering::Equal {
95                            result = it;
96                            break;
97                        }
98                    }
99                    if is_signed
100                        && let Some((&l, &r)) = l.iter().zip(r).next_back()
101                        && l != r
102                    {
103                        result = (l as i8).cmp(&(r as i8));
104                    }
105                    let result = match result {
106                        Ordering::Less => ["lt", "le", "ne"].contains(&name),
107                        Ordering::Equal => ["ge", "le", "eq"].contains(&name),
108                        Ordering::Greater => ["ge", "gt", "ne"].contains(&name),
109                    };
110                    let result = if result { 255 } else { 0 };
111                    destination_bytes.extend(std::iter::repeat_n(result, dest_size));
112                }
113
114                destination.write_from_bytes(self, &destination_bytes)
115            }
116            "bitmask" => {
117                let [op] = args else {
118                    return Err(MirEvalError::InternalError(
119                        "simd_bitmask args are not provided".into(),
120                    ));
121                };
122                let (op_len, _) = self.detect_simd_ty(op.ty)?;
123                let op_count = op.interval.size / op_len;
124                let mut result: u64 = 0;
125                for (i, val) in op.get(self)?.chunks(op_count).enumerate() {
126                    if !val.iter().all(|&it| it == 0) {
127                        result |= 1 << i;
128                    }
129                }
130                destination.write_from_bytes(self, &result.to_le_bytes()[0..destination.size])
131            }
132            "shuffle" => {
133                let [left, right, index] = args else {
134                    return Err(MirEvalError::InternalError(
135                        "simd_shuffle args are not provided".into(),
136                    ));
137                };
138                let TyKind::Array(_, index_len) = index.ty.kind() else {
139                    return Err(MirEvalError::InternalError(
140                        "simd_shuffle index argument has non-array type".into(),
141                    ));
142                };
143                let index_len = match try_const_usize(self.db, index_len) {
144                    Some(it) => it as usize,
145                    None => {
146                        return Err(MirEvalError::InternalError(
147                            "simd type with unevaluatable len param".into(),
148                        ));
149                    }
150                };
151                let (left_len, _) = self.detect_simd_ty(left.ty)?;
152                let left_size = left.interval.size / left_len;
153                let vector =
154                    left.get(self)?.chunks(left_size).chain(right.get(self)?.chunks(left_size));
155                let mut result = vec![];
156                for index in index.get(self)?.chunks(index.interval.size / index_len) {
157                    let index = from_bytes!(u32, index) as usize;
158                    let val = match vector.clone().nth(index) {
159                        Some(it) => it,
160                        None => {
161                            return Err(MirEvalError::InternalError(
162                                "out of bound access in simd shuffle".into(),
163                            ));
164                        }
165                    };
166                    result.extend(val);
167                }
168                destination.write_from_bytes(self, &result)
169            }
170            _ => not_supported!("unknown simd intrinsic {name}"),
171        }
172    }
173}