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
14pub struct PassManager<'c> {
16 raw: MlirPassManager,
17 _context: PhantomData<&'c Context>,
18}
19
20impl PassManager<'_> {
21 pub fn new(context: &Context) -> Self {
23 Self {
24 raw: unsafe { mlirPassManagerCreate(context.to_raw()) },
25 _context: Default::default(),
26 }
27 }
28
29 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 pub fn add_pass(&self, pass: Pass) {
41 unsafe { mlirPassManagerAddOwnedPass(self.raw, pass.to_raw()) }
42 }
43
44 pub fn enable_verifier(&self, enabled: bool) {
46 unsafe { mlirPassManagerEnableVerifier(self.raw, enabled) }
47 }
48
49 pub fn enable_ir_printing(&self) {
51 unsafe { mlirPassManagerEnableIRPrinting(self.raw) }
52 }
53
54 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 pub fn as_operation_pass_manager(&self) -> OperationPassManager {
69 unsafe { OperationPassManager::from_raw(mlirPassManagerGetAsOpPassManager(self.raw)) }
70 }
71
72 pub unsafe fn from_raw(raw: MlirPassManager) -> Self {
77 Self {
78 raw,
79 _context: Default::default(),
80 }
81 }
82
83 pub const fn to_raw(&self) -> MlirPassManager {
85 self.raw
86 }
87
88 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 #[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}