Skip to main content

melior/ir/type/
mem_ref.rs

1use super::{shaped_type_like::ShapedTypeLike, TypeLike};
2use crate::{
3    ir::{affine_map::AffineMap, attribute::AttributeLike, Attribute, Location, Type},
4    Error,
5};
6use mlir_sys::{
7    mlirMemRefTypeGet, mlirMemRefTypeGetAffineMap, mlirMemRefTypeGetChecked,
8    mlirMemRefTypeGetLayout, mlirMemRefTypeGetMemorySpace, MlirType,
9};
10
11/// A mem-ref type.
12#[derive(Clone, Copy, Debug)]
13pub struct MemRefType<'c> {
14    r#type: Type<'c>,
15}
16
17impl<'c> MemRefType<'c> {
18    /// Creates a mem-ref type.
19    pub fn new(
20        r#type: Type<'c>,
21        dimensions: &[i64],
22        layout: Option<Attribute<'c>>,
23        memory_space: Option<Attribute<'c>>,
24    ) -> Self {
25        unsafe {
26            Self::from_raw(mlirMemRefTypeGet(
27                r#type.to_raw(),
28                dimensions.len() as _,
29                dimensions.as_ptr() as *const _,
30                layout.unwrap_or_else(|| Attribute::null()).to_raw(),
31                memory_space.unwrap_or_else(|| Attribute::null()).to_raw(),
32            ))
33        }
34    }
35
36    /// Creates a mem-ref type with diagnostics.
37    pub fn checked(
38        location: Location<'c>,
39        r#type: Type<'c>,
40        dimensions: &[u64],
41        layout: Attribute<'c>,
42        memory_space: Attribute<'c>,
43    ) -> Option<Self> {
44        unsafe {
45            Self::from_option_raw(mlirMemRefTypeGetChecked(
46                location.to_raw(),
47                r#type.to_raw(),
48                dimensions.len() as isize,
49                dimensions.as_ptr() as *const i64,
50                layout.to_raw(),
51                memory_space.to_raw(),
52            ))
53        }
54    }
55
56    /// Returns a layout.
57    pub fn layout(&self) -> Attribute<'c> {
58        unsafe { Attribute::from_raw(mlirMemRefTypeGetLayout(self.r#type.to_raw())) }
59    }
60
61    /// Returns an affine map.
62    pub fn affine_map(&self) -> AffineMap<'c> {
63        unsafe { AffineMap::from_raw(mlirMemRefTypeGetAffineMap(self.r#type.to_raw())) }
64    }
65
66    /// Returns a memory space.
67    pub fn memory_space(&self) -> Option<Attribute<'c>> {
68        unsafe { Attribute::from_option_raw(mlirMemRefTypeGetMemorySpace(self.r#type.to_raw())) }
69    }
70
71    unsafe fn from_option_raw(raw: MlirType) -> Option<Self> {
72        if raw.ptr.is_null() {
73            None
74        } else {
75            Some(Self::from_raw(raw))
76        }
77    }
78}
79
80impl<'c> ShapedTypeLike<'c> for MemRefType<'c> {}
81
82type_traits!(MemRefType, is_mem_ref, "mem ref");
83
84#[cfg(test)]
85mod tests {
86    use super::*;
87    use crate::Context;
88
89    #[test]
90    fn new() {
91        let context = Context::new();
92
93        assert_eq!(
94            Type::from(MemRefType::new(Type::float64(&context), &[42], None, None,)),
95            Type::parse(&context, "memref<42xf64>").unwrap()
96        );
97    }
98
99    #[test]
100    fn dynamic_dimension() {
101        let context = Context::new();
102
103        assert_eq!(
104            Type::from(MemRefType::new(
105                Type::float64(&context),
106                &[i64::MIN],
107                None,
108                None,
109            )),
110            Type::parse(&context, "memref<?xf64>").unwrap()
111        );
112    }
113
114    #[test]
115    fn layout() {
116        let context = Context::new();
117
118        assert_eq!(
119            MemRefType::new(Type::index(&context), &[42, 42], None, None,).layout(),
120            Attribute::parse(&context, "affine_map<(d0, d1) -> (d0, d1)>").unwrap(),
121        );
122    }
123
124    #[test]
125    fn affine_map() {
126        let context = Context::new();
127
128        assert_eq!(
129            MemRefType::new(Type::index(&context), &[42, 42], None, None,)
130                .affine_map()
131                .to_string(),
132            "(d0, d1) -> (d0, d1)"
133        );
134    }
135
136    #[test]
137    fn memory_space() {
138        let context = Context::new();
139
140        assert_eq!(
141            MemRefType::new(Type::index(&context), &[42, 42], None, None).memory_space(),
142            None,
143        );
144    }
145}