1use crate::internal::*;
2use std::collections::HashSet;
3use std::fmt::Debug;
4use tract_itertools::Itertools;
5
6pub mod change_axes;
7mod concat_then_einsum;
8mod op_optim;
9mod prop_const;
10pub mod propagate_roi;
11pub mod propagate_uniform_tdim;
12mod push_split_down;
13mod slice;
14mod uniform_mask;
15
16use self::change_axes::ChangeAxes;
17use self::prop_const::PropConst;
18use self::propagate_roi::PropagateRoi;
19use self::propagate_uniform_tdim::PropagateUniformTdim;
20use self::push_split_down::PushSplitDown;
21use self::slice::PushSliceUp;
22use self::uniform_mask::FoldUniformMask;
23use op_optim::OpOptim;
24
25pub const CONST_FOLD_MEM_BUDGET: u64 = 4 << 20;
36
37pub trait TypedPass: Debug + Send + Sync + dyn_clone::DynClone {
38 fn reset(&mut self) -> TractResult<()>;
39 fn next(
40 &mut self,
41 session: &mut OptimizerSession,
42 model: &TypedModel,
43 ) -> TractResult<Option<TypedModelPatch>>;
44 fn run_direct(&mut self, _model: &mut TypedModel) -> TractResult<bool> {
46 Ok(false)
47 }
48}
49
50#[derive(Clone, Debug, Default)]
51struct MergeConsecutiveSameRoleAxes;
52
53impl TypedPass for MergeConsecutiveSameRoleAxes {
54 fn reset(&mut self) -> TractResult<()> {
55 Ok(())
56 }
57 fn next(
58 &mut self,
59 _session: &mut OptimizerSession,
60 _model: &TypedModel,
61 ) -> TractResult<Option<TypedModelPatch>> {
62 Ok(None)
63 }
64 fn run_direct(&mut self, model: &mut TypedModel) -> TractResult<bool> {
65 let before = model.nodes.len();
66 crate::ops::einsum::einsum_matmul::merge_consecutive_same_role_axes(model)?;
67 Ok(model.nodes.len() != before)
68 }
69}
70
71dyn_clone::clone_trait_object!(TypedPass);
72
73#[derive(Debug)]
74pub struct Optimizer {
75 pub passes: Vec<Box<dyn TypedPass>>,
76 pub steps: Option<usize>,
77}
78
79impl Optimizer {
80 fn passes(passes: Vec<Box<dyn TypedPass>>) -> Optimizer {
81 Optimizer { passes, steps: None }
82 }
83
84 pub fn add_pass(&mut self, idx: usize, pass: Box<dyn TypedPass>) {
85 let num_pass = self.passes.len();
86 if idx > num_pass {
87 log::warn!(
88 "Cannot add new pass {pass:?} at index {idx}. Optimizer currently as {num_pass} passes, pass will be added as the last pass."
89 );
90 self.passes.push(pass);
91 } else {
92 self.passes.insert(idx, pass);
93 }
94 }
95
96 pub fn stopping_at(self, steps: usize) -> Optimizer {
97 Optimizer { steps: Some(steps), ..self }
98 }
99
100 pub fn prop_consts() -> Optimizer {
101 Optimizer::passes(vec![Box::<PropConst>::default()])
102 }
103
104 pub fn declutter() -> Optimizer {
105 Optimizer::passes(vec![
106 Box::<PropConst>::default(),
107 Box::<PropagateUniformTdim>::default(),
108 Box::<PropagateRoi>::default(),
109 Box::<FoldUniformMask>::default(),
110 Box::new(OpOptim("declutter", TypedOp::declutter_with_session, 0)),
111 Box::new(PushSliceUp),
112 Box::new(PushSplitDown),
113 Box::<concat_then_einsum::ConcatThenEinsum>::default(),
114 Box::<ChangeAxes>::default(),
115 ])
116 }
117
118 pub fn codegen() -> Optimizer {
119 Optimizer::passes(vec![
120 Box::<PropConst>::default(),
121 Box::<MergeConsecutiveSameRoleAxes>::default(),
122 Box::new(OpOptim(
123 "codegen",
124 |op, _session, model, node| TypedOp::codegen(op, model, node),
125 0,
126 )),
127 Box::new(OpOptim("declutter", TypedOp::declutter_with_session, 0)),
128 Box::new(PushSplitDown),
129 Box::new(OpOptim(
130 "fuse",
131 |op, _session, model, node| TypedOp::fuse(op, model, node),
132 0,
133 )),
134 ])
135 }
136
137 pub fn optimize(&self, model: &mut TypedModel) -> TractResult<()> {
138 self.session().optimize(model)
139 }
140
141 pub fn session(&self) -> OptimizerSession<'_> {
142 OptimizerSession { optimizer: self, counter: 0, seen: Default::default() }
143 }
144}
145
146#[derive(Debug)]
147pub struct OptimizerSession<'o> {
148 optimizer: &'o Optimizer,
149 counter: usize,
150 seen: HashSet<String>,
151}
152
153impl OptimizerSession<'_> {
154 pub fn optimize(&mut self, model: &mut TypedModel) -> TractResult<()> {
155 let _proof_session = model.symbols.proof_cache_session();
156 model.check_consistency().context("during optimizer preflight check")?;
157 model.compact().context("during optimizer preflight compaction")?;
158 model.check_names().context("after optimizer preflight compaction")?;
159 for i in 0.. {
160 let old = self.counter;
161 self.run_all_passes(i, model)?;
162 if old == self.counter {
163 return Ok(());
164 }
165 model.compact()?;
166 }
167 unreachable!()
168 }
169
170 pub fn run_all_passes(&mut self, i: usize, model: &mut TypedModel) -> TractResult<()> {
171 let mut passes = self.optimizer.passes.clone();
172 for p in passes.iter_mut() {
173 self.run_one_pass_outer(i, p.as_mut(), model)
174 .with_context(|| format!("running pass {p:?}"))?;
175 model.compact()?;
176 model
177 .check_consistency()
178 .with_context(|| format!("consistency check after pass {p:?}"))?;
179 }
180 Ok(())
181 }
182
183 pub fn run_one_pass_outer(
184 &mut self,
185 i: usize,
186 p: &mut dyn TypedPass,
187 model: &mut TypedModel,
188 ) -> TractResult<()> {
189 loop {
190 let old_counter = self.counter;
191 self.run_one_pass_inner(i, p, model)?;
192 if self.counter == old_counter {
193 return Ok(());
194 }
195 model.compact().with_context(|| format!("after pass {p:?}"))?;
196 }
197 }
198
199 pub fn run_one_pass_inner(
200 &mut self,
201 i: usize,
202 p: &mut dyn TypedPass,
203 model: &mut TypedModel,
204 ) -> TractResult<()> {
205 p.reset()?;
206 if let Some(steps) = self.optimizer.steps
207 && self.counter >= steps
208 {
209 return Ok(());
210 }
211 while let Some(mut patch) = p.next(self, model)? {
212 patch.push_context(format!("{p:?}/{i}"));
213 patch.model.check_consistency().context("checking patch internal consistency")?;
214 model
215 .check_consistency()
216 .context("Checking target model consistency before patching")?;
217 if let Some(watchdog) = patch.dont_apply_twice.take() {
218 if self.seen.contains(&watchdog) {
219 debug!("Loop detected: {watchdog} seen before");
220 continue;
221 } else {
222 self.seen.insert(watchdog);
223 }
224 }
225 let patch_name = patch.context.iter().rev().join(" >> ");
226 debug!("applying patch #{}: {patch_name}", self.counter);
227 patch.apply(model).with_context(|| format!("Applying patch {patch_name}"))?;
228 model
229 .check_consistency()
230 .context("Checking target model consistency after patching")?;
231 self.counter += 1;
232 if let Some(steps) = self.optimizer.steps
233 && self.counter >= steps
234 {
235 return Ok(());
236 }
237 }
238 if p.run_direct(model)? {
239 model.check_consistency().with_context(|| format!("after run_direct {p:?}"))?;
240 }
241 model.check_consistency().with_context(|| format!("after pass {p:?}"))?;
242 Ok(())
243 }
244}