Skip to main content

melior/pass/
manager.rs

1use super::OperationPassManager;
2use crate::{
3    context::Context, ir::Module, logical_result::LogicalResult, pass::Pass, string_ref::StringRef,
4    Error,
5};
6use mlir_sys::{
7    mlirPassManagerAddOwnedPass, mlirPassManagerCreate, mlirPassManagerDestroy,
8    mlirPassManagerEnableIRPrinting, mlirPassManagerEnableVerifier,
9    mlirPassManagerGetAsOpPassManager, mlirPassManagerGetNestedUnder, mlirPassManagerRunOnOp,
10    MlirPassManager,
11};
12use std::{marker::PhantomData, mem::forget};
13
14/// A pass manager.
15pub struct PassManager<'c> {
16    raw: MlirPassManager,
17    _context: PhantomData<&'c Context>,
18}
19
20impl PassManager<'_> {
21    /// Creates a pass manager.
22    pub fn new(context: &Context) -> Self {
23        Self {
24            raw: unsafe { mlirPassManagerCreate(context.to_raw()) },
25            _context: Default::default(),
26        }
27    }
28
29    /// Returns an operation pass manager for nested operations corresponding to
30    /// a given name.
31    pub fn nested_under(&self, name: &str) -> OperationPassManager {
32        let name = StringRef::new(name);
33
34        unsafe {
35            OperationPassManager::from_raw(mlirPassManagerGetNestedUnder(self.raw, name.to_raw()))
36        }
37    }
38
39    /// Adds a pass.
40    pub fn add_pass(&self, pass: Pass) {
41        unsafe { mlirPassManagerAddOwnedPass(self.raw, pass.to_raw()) }
42    }
43
44    /// Enables a verifier.
45    pub fn enable_verifier(&self, enabled: bool) {
46        unsafe { mlirPassManagerEnableVerifier(self.raw, enabled) }
47    }
48
49    /// Enables IR printing.
50    pub fn enable_ir_printing(&self) {
51        unsafe { mlirPassManagerEnableIRPrinting(self.raw) }
52    }
53
54    /// Runs passes added to a pass manager against a module.
55    pub fn run(&self, module: &mut Module) -> Result<(), Error> {
56        let result = LogicalResult::from_raw(unsafe {
57            mlirPassManagerRunOnOp(self.raw, module.as_operation().to_raw())
58        });
59
60        if result.is_success() {
61            Ok(())
62        } else {
63            Err(Error::RunPass)
64        }
65    }
66
67    /// Converts a pass manager to an operation pass manager.
68    pub fn as_operation_pass_manager(&self) -> OperationPassManager {
69        unsafe { OperationPassManager::from_raw(mlirPassManagerGetAsOpPassManager(self.raw)) }
70    }
71
72    /// Creates a PassManager from the given raw pointer.
73    ///
74    /// # Safety
75    /// Caller must ensure this is a valid PassManager pointer.
76    pub unsafe fn from_raw(raw: MlirPassManager) -> Self {
77        Self {
78            raw,
79            _context: Default::default(),
80        }
81    }
82
83    /// Gets the raw object of this pass manager.
84    pub const fn to_raw(&self) -> MlirPassManager {
85        self.raw
86    }
87
88    /// Converts a PassManager into an owned raw object.
89    pub const fn into_raw(self) -> MlirPassManager {
90        let raw = self.raw;
91        forget(self);
92        raw
93    }
94}
95
96impl Drop for PassManager<'_> {
97    fn drop(&mut self) {
98        unsafe { mlirPassManagerDestroy(self.raw) }
99    }
100}
101
102#[cfg(test)]
103mod tests {
104    use super::*;
105    use crate::{
106        ir::{Location, Module},
107        pass::{self, transform::register_print_op_stats},
108        test::create_test_context,
109        utility::parse_pass_pipeline,
110    };
111    use indoc::indoc;
112    use pretty_assertions::assert_eq;
113
114    #[test]
115    fn new() {
116        let context = create_test_context();
117
118        PassManager::new(&context);
119    }
120
121    #[test]
122    fn add_pass() {
123        let context = create_test_context();
124
125        PassManager::new(&context).add_pass(pass::conversion::create_func_to_llvm());
126    }
127
128    #[test]
129    fn enable_verifier() {
130        let context = create_test_context();
131
132        PassManager::new(&context).enable_verifier(true);
133    }
134
135    // TODO Enable this test.
136    // #[test]
137    // fn enable_ir_printing() {
138    //     let context = Context::new();
139
140    //     PassManager::new(&context).enable_ir_printing();
141    // }
142
143    #[test]
144    fn run() {
145        let context = create_test_context();
146        let manager = PassManager::new(&context);
147
148        manager.add_pass(pass::conversion::create_func_to_llvm());
149        manager
150            .run(&mut Module::new(Location::unknown(&context)))
151            .unwrap();
152    }
153
154    #[test]
155    fn run_on_function() {
156        let context = create_test_context();
157
158        let mut module = Module::parse(
159            &context,
160            indoc!(
161                "
162                func.func @foo(%arg0 : i32) -> i32 {
163                    %res = arith.addi %arg0, %arg0 : i32
164                    return %res : i32
165                }
166                "
167            ),
168        )
169        .unwrap();
170
171        let manager = PassManager::new(&context);
172        manager.add_pass(pass::transform::create_print_op_stats());
173
174        assert_eq!(manager.run(&mut module), Ok(()));
175    }
176
177    #[test]
178    fn run_on_function_in_nested_module() {
179        let context = create_test_context();
180
181        let mut module = Module::parse(
182            &context,
183            indoc!(
184                "
185                func.func @foo(%arg0 : i32) -> i32 {
186                    %res = arith.addi %arg0, %arg0 : i32
187                    return %res : i32
188                }
189
190                module {
191                    func.func @bar(%arg0 : f32) -> f32 {
192                        %res = arith.addf %arg0, %arg0 : f32
193                        return %res : f32
194                    }
195                }
196                "
197            ),
198        )
199        .unwrap();
200
201        let manager = PassManager::new(&context);
202        manager
203            .nested_under("func.func")
204            .add_pass(pass::transform::create_print_op_stats());
205
206        assert_eq!(manager.run(&mut module), Ok(()));
207
208        let manager = PassManager::new(&context);
209        manager
210            .nested_under("builtin.module")
211            .nested_under("func.func")
212            .add_pass(pass::transform::create_print_op_stats());
213
214        assert_eq!(manager.run(&mut module), Ok(()));
215    }
216
217    #[test]
218    fn print_pass_pipeline() {
219        let context = create_test_context();
220        let manager = PassManager::new(&context);
221        let function_manager = manager.nested_under("func.func");
222
223        function_manager.add_pass(pass::transform::create_print_op_stats());
224
225        assert_eq!(
226            manager.as_operation_pass_manager().to_string(),
227            "any(func.func(print-op-stats{json=false}))"
228        );
229        assert_eq!(
230            function_manager.to_string(),
231            "func.func(print-op-stats{json=false})"
232        );
233    }
234
235    #[test]
236    fn parse_pass_pipeline_() {
237        let context = Context::new();
238        let manager = PassManager::new(&context);
239
240        insta::assert_snapshot!(parse_pass_pipeline(
241            manager.as_operation_pass_manager(),
242            "builtin.module(func.func(print-op-stats{json=false}),\
243                func.func(print-op-stats{json=false}))"
244        )
245        .unwrap_err());
246
247        register_print_op_stats();
248
249        assert_eq!(
250            parse_pass_pipeline(
251                manager.as_operation_pass_manager(),
252                "builtin.module(func.func(print-op-stats{json=false}),\
253                func.func(print-op-stats{json=false}))"
254            ),
255            Ok(())
256        );
257
258        assert_eq!(
259            manager.as_operation_pass_manager().to_string(),
260            "builtin.module(func.func(print-op-stats{json=false}),\
261            func.func(print-op-stats{json=false}))"
262        );
263    }
264}