use laddu_expr::{ExprGraph, ExprId, ExprMetadata, ExprNode};
use crate::{CompileResult, facts::NodeFacts};
mod canonical;
mod pipeline;
mod rewrite;
mod rules;
pub struct OptimizationPipeline {
passes: Vec<Box<dyn OptimizationPass>>,
max_iterations: usize,
}
pub struct OptimizationPassOutcome {
pub graph: ExprGraph,
pub changed: bool,
}
pub trait OptimizationPass: Send + Sync {
fn name(&self) -> &'static str;
fn run(&self, graph: ExprGraph) -> CompileResult<ExprGraph>;
fn run_with_change(&self, graph: ExprGraph) -> CompileResult<OptimizationPassOutcome> {
let before = pipeline::graph_fingerprint(&graph);
let graph = self.run(graph)?;
let changed = before != pipeline::graph_fingerprint(&graph);
Ok(OptimizationPassOutcome { graph, changed })
}
}
pub struct CostGatePass {
name: &'static str,
candidate: OptimizationPipeline,
}
#[derive(Copy, Clone, Debug, Default)]
pub struct CanonicalCsePass;
pub struct RewritePass {
name: &'static str,
rules: Vec<Box<dyn RewriteRule>>,
}
pub trait RewriteRule: Send + Sync {
fn name(&self) -> &'static str;
fn rewrite(
&self,
node: &ExprNode,
metadata: &ExprMetadata,
context: &RewriteContext<'_>,
) -> CompileResult<Rewrite>;
}
#[derive(Clone, Debug, PartialEq)]
pub enum Rewrite {
Keep,
Alias(ExprId),
Replace {
node: ExprNode,
metadata: ExprMetadata,
},
ReplaceMany {
nodes: Vec<(ExprNode, ExprMetadata)>,
},
}
pub struct RewriteContext<'a> {
nodes: &'a [ExprNode],
metadata: &'a [ExprMetadata],
facts: &'a [NodeFacts],
}
#[derive(Copy, Clone, Debug, Default)]
pub struct ConstantFoldScalarRule;
#[derive(Copy, Clone, Debug, Default)]
pub struct AlgebraicIdentityRule;
#[derive(Copy, Clone, Debug, Default)]
pub struct NormSqrReductionRule;
#[derive(Copy, Clone, Debug, Default)]
pub struct TrigIdentityRule;
#[derive(Copy, Clone, Debug, Default)]
pub struct CombineLikeTermsRule;
#[derive(Copy, Clone, Debug, Default)]
pub struct FactorCommonProductRule;
#[derive(Copy, Clone, Debug, Default)]
pub struct ExponentialRule;
#[derive(Copy, Clone, Debug, Default)]
pub struct NormSqrExpansionRule;
#[derive(Copy, Clone, Debug, Default)]
pub struct ConjugationRule;
#[derive(Copy, Clone, Debug, Default)]
pub struct MatrixVectorRule;
#[derive(Copy, Clone, Debug, Default)]
pub struct ComplexFactRule;