Skip to main content

tract_core/optim/
mod.rs

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
25/// Byte allowance for the growth an eager constant fold may cause.
26///
27/// A fold is admissible when its output is no larger than
28/// `input_mem.max(BUDGET)`: under the allowance it is cheap enough to pay for
29/// at load time, above it only a fold that does not grow its inputs is worth
30/// the copy. A caller that can see a constant's consumers also weighs the
31/// allowance against turning one shared buffer into one buffer per consumer.
32///
33/// Shape and index arithmetic sits far below it, decoded weight tensors far
34/// above.
35pub 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    /// In-place model mutation hook. Returns true if the model was changed.
45    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}