use crate::backend::Formula;
use crate::graph::Symbol;
use super::pattern::Pattern;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PatternKind {
WindowProduct,
ReduceWindow,
BatchNormTraining,
BatchNormInference,
}
#[derive(Debug, Clone)]
pub struct PatternMatch {
pub(crate) kind: PatternKind,
pub(crate) root: Symbol,
pub(crate) nodes: Vec<Symbol>,
}
impl PatternMatch {
pub fn kind(&self) -> PatternKind {
self.kind
}
pub fn root(&self) -> Symbol {
self.root
}
pub fn nodes(&self) -> &[Symbol] {
&self.nodes
}
pub fn formula(&self) -> Formula {
match self.kind {
PatternKind::WindowProduct => Formula::WindowProduct,
PatternKind::ReduceWindow => Formula::ReduceWindow,
PatternKind::BatchNormTraining => Formula::BatchNormTraining,
PatternKind::BatchNormInference => Formula::BatchNormInference,
}
}
}
impl Pattern {
pub(crate) fn kind(&self) -> PatternKind {
match self {
Pattern::WindowProduct(_) => PatternKind::WindowProduct,
Pattern::ReduceWindow(_) => PatternKind::ReduceWindow,
Pattern::BatchNormTraining(_) => PatternKind::BatchNormTraining,
Pattern::BatchNormInference(_) => PatternKind::BatchNormInference,
}
}
}