use super::*;
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
pub struct CustomOpIr {
pub id: String,
pub inputs: Vec<TensorIr>,
pub outputs: Vec<TensorIr>,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
pub struct SwapDimsOpIr {
pub input: TensorIr,
pub out: TensorIr,
pub dim1: usize,
pub dim2: usize,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
pub struct PermuteOpIr {
pub input: TensorIr,
pub out: TensorIr,
pub axes: Vec<usize>,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
pub struct ShapeOpIr {
pub input: TensorIr,
pub out: TensorIr,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
pub struct UnfoldOpIr {
pub input: TensorIr,
pub out: TensorIr,
pub dim: usize,
pub size: usize,
pub step: usize,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
pub struct FlipOpIr {
pub input: TensorIr,
pub out: TensorIr,
pub axes: Vec<usize>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[allow(missing_docs)]
pub struct RandomOpIr {
pub out: TensorIr,
pub distribution: Distribution,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
pub struct CreationOpIr {
pub out: TensorIr,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
pub struct FullOpIr {
pub out: TensorIr,
pub value: ScalarIr,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
pub struct InitOperationIr {
pub out: TensorIr,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
#[allow(missing_docs)]
pub struct BinaryOpIr {
pub lhs: TensorIr,
pub rhs: TensorIr,
pub out: TensorIr,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
#[allow(missing_docs)]
pub struct UnaryOpIr {
pub input: TensorIr,
pub out: TensorIr,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
#[allow(missing_docs)]
pub struct ScalarOpIr {
pub lhs: TensorIr,
pub rhs: ScalarIr,
pub out: TensorIr,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, Hash)]
#[allow(missing_docs)]
pub struct ReduceOpIr {
pub input: TensorIr,
pub out: TensorIr,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, Hash)]
#[allow(missing_docs)]
pub struct ReduceDimOpIr {
pub input: TensorIr,
pub out: TensorIr,
pub axis: usize,
pub accumulator_len: usize,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
#[allow(missing_docs)]
pub struct CastOpIr {
pub input: TensorIr,
pub out: TensorIr,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, Hash)]
#[allow(missing_docs)]
pub struct DimOpIr {
pub input: TensorIr,
pub out: TensorIr,
pub axis: usize,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
#[allow(missing_docs)]
pub struct GatherOpIr {
pub tensor: TensorIr,
pub dim: usize,
pub indices: TensorIr,
pub out: TensorIr,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
#[allow(missing_docs)]
pub struct ScatterOpIr {
pub tensor: TensorIr,
pub dim: usize,
pub indices: TensorIr,
pub value: TensorIr,
pub update: IndexingUpdateOp,
pub out: TensorIr,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
#[allow(missing_docs)]
pub struct ScatterNdOpIr {
pub data: TensorIr,
pub indices: TensorIr,
pub values: TensorIr,
pub reduction: IndexingUpdateOp,
pub out: TensorIr,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
#[allow(missing_docs)]
pub struct GatherNdOpIr {
pub data: TensorIr,
pub indices: TensorIr,
pub out: TensorIr,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
#[allow(missing_docs)]
pub struct SelectOpIr {
pub tensor: TensorIr,
pub dim: usize,
pub indices: TensorIr,
pub out: TensorIr,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
#[allow(missing_docs)]
pub struct SelectAssignOpIr {
pub tensor: TensorIr,
pub dim: usize,
pub indices: TensorIr,
pub value: TensorIr,
pub update: IndexingUpdateOp,
pub out: TensorIr,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
#[allow(missing_docs)]
pub struct SliceOpIr {
pub tensor: TensorIr,
pub ranges: Vec<Slice>,
pub out: TensorIr,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
#[allow(missing_docs)]
pub struct SliceAssignOpIr {
pub tensor: TensorIr,
pub ranges: Vec<crate::Slice>,
pub value: TensorIr,
pub out: TensorIr,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
#[allow(missing_docs)]
pub struct MaskWhereOpIr {
pub tensor: TensorIr,
pub mask: TensorIr,
pub value: TensorIr,
pub out: TensorIr,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
#[allow(missing_docs)]
pub struct MaskFillOpIr {
pub tensor: TensorIr,
pub mask: TensorIr,
pub value: ScalarIr,
pub out: TensorIr,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
#[allow(missing_docs)]
pub struct ClampOpIr {
pub tensor: TensorIr,
pub min: ScalarIr,
pub max: ScalarIr,
pub out: TensorIr,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
#[allow(missing_docs)]
pub struct RepeatDimOpIr {
pub tensor: TensorIr,
pub dim: usize,
pub times: usize,
pub out: TensorIr,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
#[allow(missing_docs)]
pub struct CatOpIr {
pub tensors: Vec<TensorIr>,
pub dim: usize,
pub out: TensorIr,
}
#[cfg(feature = "graph-distributed")]
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
#[allow(missing_docs)]
pub struct AllReduceOpIr {
pub tensor: TensorIr,
pub out: TensorIr,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
#[allow(missing_docs)]
pub struct ReduceDimWithIndicesOpIr {
pub tensor: TensorIr,
pub dim: usize,
pub out: TensorIr,
pub out_indices: TensorIr,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
#[allow(missing_docs)]
pub struct EmbeddingOpIr {
pub weights: TensorIr,
pub indices: TensorIr,
pub out: TensorIr,
}
#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
#[allow(missing_docs)]
pub struct EmbeddingBackwardOpIr {
pub weights: TensorIr,
pub out_grad: TensorIr,
pub indices: TensorIr,
pub out: TensorIr,
}