use crate::tensor::mem_formats::layout::Layout;
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum Sign {
Plus,
Minus,
}
#[derive(Clone, Debug, PartialEq)]
pub enum OpKindScalar<T> {
AxBy(T, T),
Exp,
Ln,
Log2,
Inv, ReLU,
Tanh,
}
#[derive(Clone, Debug, PartialEq)]
pub enum OpKind<T> {
NoOp,
ScalarOp(OpKindScalar<T>),
FusedScalar(Box<[OpKindScalar<T>]>),
View(Layout),
Slice(Layout),
Transpose,
TransposeAxes(Layout),
Broadcast(Layout),
MatMul(T), MatMulSum(T, T, Sign), AsContiguous,
Add,
Sub,
Mul,
Div,
Sum,
SumAxis(isize, bool),
Mean,
MeanAxis(isize, bool),
Max,
MaxAxis(isize, bool),
}
impl<T: Copy> OpKind<T> {
pub fn as_str(&self) -> &'static str {
match self {
OpKind::NoOp => "NoOp",
OpKind::ScalarOp(_) => "ScalarOp",
OpKind::FusedScalar(_) => "FusedScalar",
OpKind::View(_) => "View",
OpKind::Slice(_) => "Slice",
OpKind::Transpose => "Transpose",
OpKind::TransposeAxes(_) => "TransposeAxes",
OpKind::Broadcast(_) => "Broadcast",
OpKind::MatMul(_) => "MatMul",
OpKind::MatMulSum(_, _, _) => "MatMulSum",
OpKind::AsContiguous => "AsContiguous",
OpKind::Add => "Add",
OpKind::Sub => "Sub",
OpKind::Mul => "Mul",
OpKind::Div => "Div",
OpKind::Sum => "Sum",
OpKind::SumAxis(_, _) => "SumAxis",
OpKind::Mean => "Mean",
OpKind::MeanAxis(_, _) => "MeanAxis",
OpKind::Max => "Max",
OpKind::MaxAxis(_, _) => "MaxAxis",
}
}
}