1use burn_backend::ops::AttentionModuleOptions;
2use burn_backend::tensor::IndexingUpdateOp;
3use core::hash::Hash;
4use serde::{Deserialize, Serialize};
5
6use alloc::borrow::ToOwned;
7use alloc::boxed::Box;
8use alloc::{string::String, vec::Vec};
9
10use burn_backend::{
11 DType, Distribution, Slice,
12 ops::{
13 ConvOptions, ConvTransposeOptions, DeformConvOptions, GridSampleOptions,
14 GridSamplePaddingMode, InterpolateMode, InterpolateOptions,
15 },
16 quantization::QuantScheme,
17};
18
19use crate::{ScalarIr, TensorId, TensorIr, TensorStatus};
20
21pub trait IrVisitorMut {
23 fn visit_tensor_mut(&mut self, _tensor: &mut TensorIr) {}
25 fn visit_scalar_mut(&mut self, _scalar: &mut ScalarIr) {}
27 fn visit_range_mut(&mut self, _range: &mut Slice) {}
29}
30
31#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
33pub struct CustomOpIr {
34 pub id: String,
36 pub inputs: Vec<TensorIr>,
38 pub outputs: Vec<TensorIr>,
40 pub scalars: Vec<ScalarIr>,
47}
48
49impl CustomOpIr {
50 pub fn new(id: &'static str, inputs: &[TensorIr], outputs: &[TensorIr]) -> Self {
52 Self {
53 id: id.to_owned(),
54 inputs: inputs.to_vec(),
55 outputs: outputs.to_vec(),
56 scalars: Vec::new(),
57 }
58 }
59
60 pub fn with_scalars(
62 id: &'static str,
63 inputs: &[TensorIr],
64 outputs: &[TensorIr],
65 scalars: Vec<ScalarIr>,
66 ) -> Self {
67 Self {
68 id: id.to_owned(),
69 inputs: inputs.to_vec(),
70 outputs: outputs.to_vec(),
71 scalars,
72 }
73 }
74
75 pub fn as_fixed<const N_IN: usize, const N_OUT: usize>(
77 &self,
78 ) -> (&[TensorIr; N_IN], &[TensorIr; N_OUT]) {
79 (
80 self.inputs.as_slice().try_into().expect(
81 "Wrong number of inputs expected (expected {D}, is {}), check your implementation",
82 ),
83 self.outputs.as_slice().try_into().expect(
84 "Wrong number of outputs expected (expected {D}, is {}), check your implementation",
85 ),
86 )
87 }
88
89 fn inputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
90 Box::new(self.inputs.iter())
91 }
92
93 fn outputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
94 Box::new(self.outputs.iter())
95 }
96
97 fn visit_mut(&mut self, v: &mut impl IrVisitorMut) {
98 for t in self.inputs.iter_mut() {
99 v.visit_tensor_mut(t);
100 }
101 for t in self.outputs.iter_mut() {
102 v.visit_tensor_mut(t);
103 }
104 for s in self.scalars.iter_mut() {
105 v.visit_scalar_mut(s);
106 }
107 }
108}
109
110#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
112#[allow(clippy::large_enum_variant)]
113pub enum OperationIr {
114 BaseFloat(BaseOperationIr),
116 BaseInt(BaseOperationIr),
118 BaseBool(BaseOperationIr),
120 NumericFloat(DType, NumericOperationIr),
122 NumericInt(DType, NumericOperationIr),
124 Bool(BoolOperationIr),
126 Int(IntOperationIr),
128 Float(DType, FloatOperationIr),
130 Module(ModuleOperationIr),
132 Init(InitOperationIr),
134 Custom(CustomOpIr),
136 Drop(TensorIr),
138 Distributed(DistributedOperationIr),
140 Activation(ActivationOperationIr),
142}
143
144#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
146pub enum FloatOperationIr {
147 Exp(UnaryOpIr),
149 Log(UnaryOpIr),
151 Log1p(UnaryOpIr),
153 Erf(UnaryOpIr),
155 PowfScalar(ScalarOpIr),
157 Sqrt(UnaryOpIr),
159 Cos(UnaryOpIr),
161 Cosh(UnaryOpIr),
163 Sin(UnaryOpIr),
165 Sinh(UnaryOpIr),
167 Tan(UnaryOpIr),
169 Tanh(UnaryOpIr),
171 ArcCos(UnaryOpIr),
173 ArcCosh(UnaryOpIr),
175 ArcSin(UnaryOpIr),
177 ArcSinh(UnaryOpIr),
179 ArcTan(UnaryOpIr),
181 ArcTanh(UnaryOpIr),
183 ArcTan2(BinaryOpIr),
185 Round(UnaryOpIr),
187 Floor(UnaryOpIr),
189 Ceil(UnaryOpIr),
191 Trunc(UnaryOpIr),
193 IntoInt(CastOpIr),
195 Matmul(MatmulOpIr),
197 Cross(CrossOpIr),
199 Random(RandomOpIr),
201 Recip(UnaryOpIr),
203 IsNan(UnaryOpIr),
205 IsInf(UnaryOpIr),
207 Quantize(QuantizeOpIr),
209 Dequantize(DequantizeOpIr),
211 GridSample2d(GridSample2dOpIr),
213 Powf(BinaryOpIr),
215 Hypot(BinaryOpIr),
217}
218
219#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
221pub enum ModuleOperationIr {
222 BatchNorm(BatchNormOpIr),
225 Embedding(EmbeddingOpIr),
227 EmbeddingBackward(EmbeddingBackwardOpIr),
229 Linear(LinearOpIr),
231 LinearXBackward(LinearXBackwardOpIr),
233 LinearWeightBackward(LinearWeightBackwardOpIr),
235 LinearBiasBackward(LinearBiasBackwardOpIr),
237 Conv1d(Conv1dOpIr),
239 Conv1dXBackward(Conv1dXBackwardOpIr),
241 Conv1dWeightBackward(Conv1dWeightBackwardOpIr),
243 Conv1dBiasBackward(Conv1dBiasBackwardOpIr),
245 Conv2d(Conv2dOpIr),
247 Conv2dXBackward(Conv2dXBackwardOpIr),
249 Conv2dWeightBackward(Conv2dWeightBackwardOpIr),
251 Conv2dBiasBackward(Conv2dBiasBackwardOpIr),
253 Conv3d(Conv3dOpIr),
255 Conv3dXBackward(Conv3dXBackwardOpIr),
257 Conv3dWeightBackward(Conv3dWeightBackwardOpIr),
259 Conv3dBiasBackward(Conv3dBiasBackwardOpIr),
261 DeformableConv2d(Box<DeformConv2dOpIr>),
263 DeformableConv2dBackward(Box<DeformConv2dBackwardOpIr>),
265 ConvTranspose1d(ConvTranspose1dOpIr),
267 ConvTranspose2d(ConvTranspose2dOpIr),
269 ConvTranspose3d(ConvTranspose3dOpIr),
271 AvgPool1d(AvgPool1dOpIr),
273 AvgPool2d(AvgPool2dOpIr),
275 AvgPool1dBackward(AvgPool1dBackwardOpIr),
278 AvgPool2dBackward(AvgPool2dBackwardOpIr),
281 AdaptiveAvgPool1d(AdaptiveAvgPool1dOpIr),
284 AdaptiveAvgPool2d(AdaptiveAvgPool2dOpIr),
287 AdaptiveAvgPool1dBackward(AdaptiveAvgPool1dBackwardOpIr),
290 AdaptiveAvgPool2dBackward(AdaptiveAvgPool2dBackwardOpIr),
293 AdaptiveAvgPool3d(AdaptiveAvgPool3dOpIr),
296 AdaptiveAvgPool3dBackward(AdaptiveAvgPool3dBackwardOpIr),
299 MaxPool1d(MaxPool1dOpIr),
302 MaxPool1dWithIndices(MaxPool1dWithIndicesOpIr),
305 MaxPool1dWithIndicesBackward(MaxPool1dWithIndicesBackwardOpIr),
308 MaxPool2d(MaxPool2dOpIr),
311 MaxPool2dWithIndices(MaxPool2dWithIndicesOpIr),
314 MaxPool2dWithIndicesBackward(MaxPool2dWithIndicesBackwardOpIr),
317 Interpolate(InterpolateOpIr),
319 InterpolateBackward(InterpolateBackwardOpIr),
321 Rfft(RfftOpIr),
323 IRfft(IRfftOpIr),
325 Attention(AttentionOpIr),
327 CtcLoss(CtcLossOpIr),
329 CtcLossBackward(CtcLossBackwardOpIr),
332 LayerNorm(LayerNormOpIr),
334 Unfold4d(Unfold4dOpIr),
336 ConvTranspose1dWeightBackward(ConvTranspose1dWeightBackwardOpIr),
339 ConvTranspose1dBiasBackward(ConvTranspose1dBiasBackwardOpIr),
342 ConvTranspose2dWeightBackward(ConvTranspose2dWeightBackwardOpIr),
345 ConvTranspose2dBiasBackward(ConvTranspose2dBiasBackwardOpIr),
348 ConvTranspose3dWeightBackward(ConvTranspose3dWeightBackwardOpIr),
351 ConvTranspose3dBiasBackward(ConvTranspose3dBiasBackwardOpIr),
354}
355
356#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
358pub enum BaseOperationIr {
359 Reshape(ShapeOpIr),
365
366 SwapDims(SwapDimsOpIr),
372
373 Permute(PermuteOpIr),
379
380 Flip(FlipOpIr),
385
386 Expand(ShapeOpIr),
392
393 Unfold(UnfoldOpIr),
396
397 Slice(SliceOpIr),
403 SliceAssign(SliceAssignOpIr),
409 Select(SelectOpIr),
415 SelectAssign(SelectAssignOpIr),
421 MaskWhere(MaskWhereOpIr),
427 MaskFill(MaskFillOpIr),
433 Gather(GatherOpIr),
439 Scatter(ScatterOpIr),
445 ScatterNd(ScatterNdOpIr),
447 GatherNd(GatherNdOpIr),
449 Equal(BinaryOpIr),
455 EqualElem(ScalarOpIr),
461 RepeatDim(RepeatDimOpIr),
467 Cat(CatOpIr),
473 Cast(CastOpIr),
475 Empty(CreationOpIr),
481 Ones(CreationOpIr),
487 Zeros(CreationOpIr),
493 NotEqual(BinaryOpIr),
499 NotEqualElem(ScalarOpIr),
505 All(ReduceOpIr),
510 Any(ReduceOpIr),
512 AllDim(ReduceDimOpIr),
514 AnyDim(ReduceDimOpIr),
516}
517
518#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
520pub enum NumericOperationIr {
521 Add(BinaryOpIr),
526 AddScalar(ScalarOpIr),
531 Sub(BinaryOpIr),
536 SubScalar(ScalarOpIr),
541 Div(BinaryOpIr),
546 DivScalar(ScalarOpIr),
551 Rem(BinaryOpIr),
556 RemScalar(ScalarOpIr),
561 Mul(BinaryOpIr),
566 MulScalar(ScalarOpIr),
571 Abs(UnaryOpIr),
576 Full(FullOpIr),
581 MeanDim(ReduceDimOpIr),
586 Mean(ReduceOpIr),
591 Sum(ReduceOpIr),
596 SumDim(ReduceDimOpIr),
601 Prod(ReduceOpIr),
606 ProdDim(ReduceDimOpIr),
611 Greater(BinaryOpIr),
616 GreaterElem(ScalarOpIr),
621 GreaterEqual(BinaryOpIr),
626 GreaterEqualElem(ScalarOpIr),
631 Lower(BinaryOpIr),
636 LowerElem(ScalarOpIr),
641 LowerEqual(BinaryOpIr),
646 LowerEqualElem(ScalarOpIr),
651 ArgMax(ReduceDimOpIr),
656 ArgTopK(ReduceDimOpIr),
661 TopK(ReduceDimOpIr),
666 TopKWithIndices(TopKWithIndicesOpIr),
671 ArgMin(ReduceDimOpIr),
676 Max(ReduceOpIr),
681 MaxDimWithIndices(ReduceDimWithIndicesOpIr),
686 MinDimWithIndices(ReduceDimWithIndicesOpIr),
691 Min(ReduceOpIr),
696 MaxDim(ReduceDimOpIr),
701 MinDim(ReduceDimOpIr),
706 MaxAbs(ReduceOpIr),
711 MaxAbsDim(ReduceDimOpIr),
716 Clamp(ClampOpIr),
721 IntRandom(RandomOpIr),
725 Powi(BinaryOpIr),
730 PowiScalar(ScalarOpIr),
735 CumSum(DimOpIr),
740 CumProd(DimOpIr),
745 CumMin(DimOpIr),
750 CumMax(DimOpIr),
755 Neg(UnaryOpIr),
760 Sign(UnaryOpIr),
765 ClampMin(ScalarOpIr),
770 ClampMax(ScalarOpIr),
775 Sort(SortOpIr),
777 SortWithIndices(SortWithIndicesOpIr),
779 ArgSort(SortOpIr),
783}
784
785#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
787pub enum IntOperationIr {
788 IntoFloat(CastOpIr),
790 BitwiseAnd(BinaryOpIr),
794 BitwiseAndScalar(ScalarOpIr),
798 BitwiseOr(BinaryOpIr),
802 BitwiseOrScalar(ScalarOpIr),
806 BitwiseXor(BinaryOpIr),
810 BitwiseXorScalar(ScalarOpIr),
814 BitwiseNot(UnaryOpIr),
818 BitwiseLeftShift(BinaryOpIr),
822 BitwiseLeftShiftScalar(ScalarOpIr),
826 BitwiseRightShift(BinaryOpIr),
830 BitwiseRightShiftScalar(ScalarOpIr),
834 Matmul(MatmulOpIr),
836}
837
838#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
840pub enum BoolOperationIr {
841 IntoFloat(CastOpIr),
843 IntoInt(CastOpIr),
845 Not(UnaryOpIr),
847 And(BinaryOpIr),
849 Or(BinaryOpIr),
851 Xor(BinaryOpIr),
853}
854
855#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
857#[allow(clippy::large_enum_variant)]
858pub enum DistributedOperationIr {
859 AllReduce(AllReduceOpIr),
862 SyncCollective,
869}
870
871#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
873pub struct SwapDimsOpIr {
874 pub input: TensorIr,
876 pub out: TensorIr,
878 pub dim1: usize,
880 pub dim2: usize,
882}
883
884#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
886pub struct PermuteOpIr {
887 pub input: TensorIr,
889 pub out: TensorIr,
891 pub axes: Vec<usize>,
893}
894
895#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
897pub struct ShapeOpIr {
898 pub input: TensorIr,
900 pub out: TensorIr,
902}
903
904#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
906pub struct UnfoldOpIr {
907 pub input: TensorIr,
909 pub out: TensorIr,
911
912 pub dim: usize,
914 pub size: usize,
916 pub step: usize,
918}
919
920#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
922pub struct FlipOpIr {
923 pub input: TensorIr,
925 pub out: TensorIr,
927 pub axes: Vec<usize>,
929}
930
931#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
932#[allow(missing_docs)]
933pub struct RandomOpIr {
934 pub out: TensorIr,
935 pub distribution: Distribution,
936}
937
938#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
941pub struct CreationOpIr {
942 pub out: TensorIr,
944}
945
946#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
948pub struct FullOpIr {
949 pub out: TensorIr,
951 pub value: ScalarIr,
953}
954
955#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
956pub struct InitOperationIr {
960 pub out: TensorIr,
962}
963
964#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
965#[allow(missing_docs)]
966pub struct BinaryOpIr {
967 pub lhs: TensorIr,
968 pub rhs: TensorIr,
969 pub out: TensorIr,
970}
971
972#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
973#[allow(missing_docs)]
974pub struct MatmulOpIr {
975 pub lhs: TensorIr,
976 pub rhs: TensorIr,
977 pub out: TensorIr,
978}
979
980#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
981#[allow(missing_docs)]
982pub struct CrossOpIr {
983 pub lhs: TensorIr,
984 pub rhs: TensorIr,
985 pub out: TensorIr,
986 pub dim: usize,
987}
988
989#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
990#[allow(missing_docs)]
991pub struct UnaryOpIr {
992 pub input: TensorIr,
993 pub out: TensorIr,
994}
995
996#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
997#[allow(missing_docs)]
998pub struct ScalarOpIr {
999 pub lhs: TensorIr,
1000 pub rhs: ScalarIr,
1003 pub out: TensorIr,
1004}
1005
1006#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, Hash)]
1007#[allow(missing_docs)]
1008pub struct ReduceOpIr {
1009 pub input: TensorIr,
1010 pub out: TensorIr,
1011}
1012
1013#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, Hash)]
1014#[allow(missing_docs)]
1015pub struct ReduceDimOpIr {
1016 pub input: TensorIr,
1017 pub out: TensorIr,
1018 pub axis: usize,
1019 pub accumulator_len: usize,
1020}
1021
1022#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1023#[allow(missing_docs)]
1024pub struct CastOpIr {
1025 pub input: TensorIr,
1026 pub out: TensorIr,
1027}
1028
1029#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, Hash)]
1032#[allow(missing_docs)]
1033pub struct DimOpIr {
1034 pub input: TensorIr,
1035 pub out: TensorIr,
1036 pub axis: usize,
1037}
1038
1039#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1040#[allow(missing_docs)]
1041pub struct GatherOpIr {
1042 pub tensor: TensorIr,
1043 pub dim: usize,
1044 pub indices: TensorIr,
1045 pub out: TensorIr,
1046}
1047
1048#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1049#[allow(missing_docs)]
1050pub struct ScatterOpIr {
1051 pub tensor: TensorIr,
1052 pub dim: usize,
1053 pub indices: TensorIr,
1054 pub value: TensorIr,
1055 pub update: IndexingUpdateOp,
1056 pub out: TensorIr,
1057}
1058
1059#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1060#[allow(missing_docs)]
1061pub struct ScatterNdOpIr {
1062 pub data: TensorIr,
1063 pub indices: TensorIr,
1064 pub values: TensorIr,
1065 pub reduction: IndexingUpdateOp,
1066 pub out: TensorIr,
1067}
1068
1069#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1070#[allow(missing_docs)]
1071pub struct GatherNdOpIr {
1072 pub data: TensorIr,
1073 pub indices: TensorIr,
1074 pub out: TensorIr,
1075}
1076
1077#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1078#[allow(missing_docs)]
1079pub struct SelectOpIr {
1080 pub tensor: TensorIr,
1081 pub dim: usize,
1082 pub indices: TensorIr,
1083 pub out: TensorIr,
1084}
1085
1086#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1087#[allow(missing_docs)]
1088pub struct SelectAssignOpIr {
1089 pub tensor: TensorIr,
1090 pub dim: usize,
1091 pub indices: TensorIr,
1092 pub value: TensorIr,
1093 pub update: IndexingUpdateOp,
1094 pub out: TensorIr,
1095}
1096
1097#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1098#[allow(missing_docs)]
1099pub struct SliceOpIr {
1100 pub tensor: TensorIr,
1101 pub ranges: Vec<Slice>,
1102 pub out: TensorIr,
1103}
1104
1105#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1106#[allow(missing_docs)]
1107pub struct SliceAssignOpIr {
1108 pub tensor: TensorIr,
1109 pub ranges: Vec<burn_backend::Slice>,
1110 pub value: TensorIr,
1111 pub out: TensorIr,
1112}
1113
1114#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1115#[allow(missing_docs)]
1116pub struct MaskWhereOpIr {
1117 pub tensor: TensorIr,
1118 pub mask: TensorIr,
1119 pub value: TensorIr,
1120 pub out: TensorIr,
1121}
1122
1123#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1124#[allow(missing_docs)]
1125pub struct MaskFillOpIr {
1126 pub tensor: TensorIr,
1127 pub mask: TensorIr,
1128 pub value: ScalarIr,
1129 pub out: TensorIr,
1130}
1131
1132#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1133#[allow(missing_docs)]
1134pub struct ClampOpIr {
1135 pub tensor: TensorIr,
1136 pub min: ScalarIr,
1137 pub max: ScalarIr,
1138 pub out: TensorIr,
1139}
1140
1141#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1142#[allow(missing_docs)]
1143pub struct RepeatDimOpIr {
1144 pub tensor: TensorIr,
1145 pub dim: usize,
1146 pub times: usize,
1147 pub out: TensorIr,
1148}
1149
1150#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1151#[allow(missing_docs)]
1152pub struct CatOpIr {
1153 pub tensors: Vec<TensorIr>,
1154 pub dim: usize,
1155 pub out: TensorIr,
1156}
1157
1158#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1159#[allow(missing_docs)]
1160pub struct AllReduceOpIr {
1161 pub tensor: TensorIr,
1162 pub out: TensorIr,
1163 pub op: burn_backend::distributed::ReduceOperation,
1165 pub device_ids: Vec<DeviceIdIr>,
1167}
1168
1169#[derive(Clone, Copy, Debug, Hash, PartialEq, Eq, Serialize, Deserialize)]
1174pub struct DeviceIdIr {
1175 pub type_id: u16,
1177 pub index_id: u16,
1179}
1180
1181impl From<burn_backend::DeviceId> for DeviceIdIr {
1182 fn from(value: burn_backend::DeviceId) -> Self {
1183 Self {
1184 type_id: value.type_id,
1185 index_id: value.index_id,
1186 }
1187 }
1188}
1189
1190impl From<DeviceIdIr> for burn_backend::DeviceId {
1191 fn from(value: DeviceIdIr) -> Self {
1192 burn_backend::DeviceId::new(value.type_id, value.index_id)
1193 }
1194}
1195
1196#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1197#[allow(missing_docs)]
1198pub struct ReduceDimWithIndicesOpIr {
1199 pub tensor: TensorIr,
1200 pub dim: usize,
1201 pub out: TensorIr,
1202 pub out_indices: TensorIr,
1203}
1204
1205#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1206#[allow(missing_docs)]
1207pub struct TopKWithIndicesOpIr {
1210 pub tensor: TensorIr,
1211 pub dim: usize,
1212 pub k: usize,
1213 pub out: TensorIr,
1214 pub out_indices: TensorIr,
1215}
1216
1217#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1218#[allow(missing_docs)]
1219pub struct EmbeddingOpIr {
1220 pub weights: TensorIr,
1221 pub indices: TensorIr,
1222 pub out: TensorIr,
1223}
1224
1225#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1226#[allow(missing_docs)]
1227pub struct EmbeddingBackwardOpIr {
1228 pub weights: TensorIr,
1229 pub out_grad: TensorIr,
1230 pub indices: TensorIr,
1231 pub out: TensorIr,
1232}
1233
1234#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1238#[allow(missing_docs)]
1239pub struct BatchNormOpIr {
1240 pub x: TensorIr,
1241 pub gamma: TensorIr,
1242 pub beta: TensorIr,
1243 pub mean: TensorIr,
1244 pub variance: TensorIr,
1245 pub epsilon: ScalarIr,
1246 pub out: TensorIr,
1247}
1248
1249#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1250#[allow(missing_docs)]
1251pub struct LinearOpIr {
1252 pub x: TensorIr,
1253 pub weight: TensorIr,
1254 pub bias: Option<TensorIr>,
1255 pub out: TensorIr,
1256}
1257
1258#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1259#[allow(missing_docs)]
1260pub struct LinearXBackwardOpIr {
1261 pub weight: TensorIr,
1262 pub output_grad: TensorIr,
1263 pub out: TensorIr,
1264}
1265
1266#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1267#[allow(missing_docs)]
1268pub struct LinearWeightBackwardOpIr {
1269 pub x: TensorIr,
1270 pub output_grad: TensorIr,
1271 pub out: TensorIr,
1272}
1273
1274#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1275#[allow(missing_docs)]
1276pub struct LinearBiasBackwardOpIr {
1277 pub output_grad: TensorIr,
1278 pub out: TensorIr,
1279}
1280
1281#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1282#[allow(missing_docs)]
1283pub struct Conv1dOpIr {
1284 pub x: TensorIr,
1285 pub weight: TensorIr,
1286 pub bias: Option<TensorIr>,
1287 pub options: Conv1dOptionsIr,
1288 pub out: TensorIr,
1289}
1290
1291#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1292#[allow(missing_docs)]
1293pub struct Conv1dXBackwardOpIr {
1294 pub x: TensorIr,
1295 pub weight: TensorIr,
1296 pub output_grad: TensorIr,
1297 pub options: Conv1dOptionsIr,
1298 pub out: TensorIr,
1299}
1300
1301#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1302#[allow(missing_docs)]
1303pub struct Conv1dWeightBackwardOpIr {
1304 pub x: TensorIr,
1305 pub weight: TensorIr,
1306 pub output_grad: TensorIr,
1307 pub options: Conv1dOptionsIr,
1308 pub out: TensorIr,
1309}
1310
1311#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1312#[allow(missing_docs)]
1313pub struct Conv1dBiasBackwardOpIr {
1314 pub x: TensorIr,
1315 pub bias: TensorIr,
1316 pub output_grad: TensorIr,
1317 pub out: TensorIr,
1318}
1319
1320#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1321#[allow(missing_docs)]
1322pub struct Conv2dOpIr {
1323 pub x: TensorIr,
1324 pub weight: TensorIr,
1325 pub bias: Option<TensorIr>,
1326 pub options: Conv2dOptionsIr,
1327 pub out: TensorIr,
1328}
1329
1330#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1331#[allow(missing_docs)]
1332pub struct Conv2dXBackwardOpIr {
1333 pub x: TensorIr,
1334 pub weight: TensorIr,
1335 pub output_grad: TensorIr,
1336 pub options: Conv2dOptionsIr,
1337 pub out: TensorIr,
1338}
1339
1340#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1341#[allow(missing_docs)]
1342pub struct Conv2dWeightBackwardOpIr {
1343 pub x: TensorIr,
1344 pub weight: TensorIr,
1345 pub output_grad: TensorIr,
1346 pub options: Conv2dOptionsIr,
1347 pub out: TensorIr,
1348}
1349
1350#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1351#[allow(missing_docs)]
1352pub struct Conv2dBiasBackwardOpIr {
1353 pub x: TensorIr,
1354 pub bias: TensorIr,
1355 pub output_grad: TensorIr,
1356 pub out: TensorIr,
1357}
1358
1359#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1360#[allow(missing_docs)]
1361pub struct DeformConv2dOpIr {
1362 pub x: TensorIr,
1363 pub offset: TensorIr,
1364 pub weight: TensorIr,
1365 pub mask: Option<TensorIr>,
1366 pub bias: Option<TensorIr>,
1367 pub options: DeformableConv2dOptionsIr,
1368 pub out: TensorIr,
1369}
1370
1371#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1372#[allow(missing_docs)]
1373pub struct DeformConv2dBackwardOpIr {
1374 pub x: TensorIr,
1375 pub offset: TensorIr,
1376 pub weight: TensorIr,
1377 pub mask: Option<TensorIr>,
1378 pub bias: Option<TensorIr>,
1379 pub out_grad: TensorIr,
1380 pub options: DeformableConv2dOptionsIr,
1381 pub input_grad: TensorIr,
1382 pub offset_grad: TensorIr,
1383 pub weight_grad: TensorIr,
1384 pub mask_grad: Option<TensorIr>,
1385 pub bias_grad: Option<TensorIr>,
1386}
1387
1388#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1389#[allow(missing_docs)]
1390pub struct Conv3dOpIr {
1391 pub x: TensorIr,
1392 pub weight: TensorIr,
1393 pub bias: Option<TensorIr>,
1394 pub options: Conv3dOptionsIr,
1395 pub out: TensorIr,
1396}
1397
1398#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1399#[allow(missing_docs)]
1400pub struct Conv3dXBackwardOpIr {
1401 pub x: TensorIr,
1402 pub weight: TensorIr,
1403 pub output_grad: TensorIr,
1404 pub options: Conv3dOptionsIr,
1405 pub out: TensorIr,
1406}
1407
1408#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1409#[allow(missing_docs)]
1410pub struct Conv3dWeightBackwardOpIr {
1411 pub x: TensorIr,
1412 pub weight: TensorIr,
1413 pub output_grad: TensorIr,
1414 pub options: Conv3dOptionsIr,
1415 pub out: TensorIr,
1416}
1417
1418#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1419#[allow(missing_docs)]
1420pub struct Conv3dBiasBackwardOpIr {
1421 pub x: TensorIr,
1422 pub bias: TensorIr,
1423 pub output_grad: TensorIr,
1424 pub out: TensorIr,
1425}
1426
1427#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1428#[allow(missing_docs)]
1429pub struct ConvTranspose1dOpIr {
1430 pub x: TensorIr,
1431 pub weight: TensorIr,
1432 pub bias: Option<TensorIr>,
1433 pub options: ConvTranspose1dOptionsIr,
1434 pub out: TensorIr,
1435}
1436
1437#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1438#[allow(missing_docs)]
1439pub struct ConvTranspose2dOpIr {
1440 pub x: TensorIr,
1441 pub weight: TensorIr,
1442 pub bias: Option<TensorIr>,
1443 pub options: ConvTranspose2dOptionsIr,
1444 pub out: TensorIr,
1445}
1446
1447#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1448#[allow(missing_docs)]
1449pub struct ConvTranspose3dOpIr {
1450 pub x: TensorIr,
1451 pub weight: TensorIr,
1452 pub bias: Option<TensorIr>,
1453 pub options: ConvTranspose3dOptionsIr,
1454 pub out: TensorIr,
1455}
1456
1457#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1458#[allow(missing_docs)]
1459pub struct Conv1dOptionsIr {
1460 pub stride: [usize; 1],
1461 pub padding: [usize; 1],
1462 pub dilation: [usize; 1],
1463 pub groups: usize,
1464}
1465
1466#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1467#[allow(missing_docs)]
1468pub struct Conv2dOptionsIr {
1469 pub stride: [usize; 2],
1470 pub padding: [usize; 2],
1471 pub dilation: [usize; 2],
1472 pub groups: usize,
1473}
1474
1475#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1476#[allow(missing_docs)]
1477pub struct DeformableConv2dOptionsIr {
1478 pub stride: [usize; 2],
1479 pub padding: [usize; 2],
1480 pub dilation: [usize; 2],
1481 pub weight_groups: usize,
1482 pub offset_groups: usize,
1483}
1484
1485#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1486#[allow(missing_docs)]
1487pub struct Conv3dOptionsIr {
1488 pub stride: [usize; 3],
1489 pub padding: [usize; 3],
1490 pub dilation: [usize; 3],
1491 pub groups: usize,
1492}
1493
1494#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1495#[allow(missing_docs)]
1496pub struct ConvTranspose1dOptionsIr {
1497 pub stride: [usize; 1],
1498 pub padding: [usize; 1],
1499 pub padding_out: [usize; 1],
1500 pub dilation: [usize; 1],
1501 pub groups: usize,
1502}
1503
1504#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1505#[allow(missing_docs)]
1506pub struct ConvTranspose2dOptionsIr {
1507 pub stride: [usize; 2],
1508 pub padding: [usize; 2],
1509 pub padding_out: [usize; 2],
1510 pub dilation: [usize; 2],
1511 pub groups: usize,
1512}
1513
1514#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1515#[allow(missing_docs)]
1516pub struct ConvTranspose3dOptionsIr {
1517 pub stride: [usize; 3],
1518 pub padding: [usize; 3],
1519 pub padding_out: [usize; 3],
1520 pub dilation: [usize; 3],
1521 pub groups: usize,
1522}
1523
1524#[derive(Debug, Clone, Hash, PartialEq, Eq, Serialize, Deserialize)]
1526pub struct QuantizationParametersIr {
1527 pub scales: TensorIr,
1529 pub global: Option<TensorIr>,
1532}
1533
1534#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1535#[allow(missing_docs)]
1536pub struct QuantizeOpIr {
1537 pub tensor: TensorIr,
1538 pub qparams: QuantizationParametersIr,
1539 pub scheme: QuantScheme,
1540 pub out: TensorIr,
1541}
1542
1543#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1544#[allow(missing_docs)]
1545pub struct DequantizeOpIr {
1546 pub input: TensorIr,
1547 pub out: TensorIr,
1548}
1549
1550impl From<ConvOptions<1>> for Conv1dOptionsIr {
1551 fn from(value: ConvOptions<1>) -> Self {
1552 Self {
1553 stride: value.stride,
1554 padding: value.padding,
1555 dilation: value.dilation,
1556 groups: value.groups,
1557 }
1558 }
1559}
1560
1561impl From<ConvOptions<2>> for Conv2dOptionsIr {
1562 fn from(value: ConvOptions<2>) -> Self {
1563 Self {
1564 stride: value.stride,
1565 padding: value.padding,
1566 dilation: value.dilation,
1567 groups: value.groups,
1568 }
1569 }
1570}
1571
1572impl From<ConvOptions<3>> for Conv3dOptionsIr {
1573 fn from(value: ConvOptions<3>) -> Self {
1574 Self {
1575 stride: value.stride,
1576 padding: value.padding,
1577 dilation: value.dilation,
1578 groups: value.groups,
1579 }
1580 }
1581}
1582
1583impl From<DeformConvOptions<2>> for DeformableConv2dOptionsIr {
1584 fn from(value: DeformConvOptions<2>) -> Self {
1585 Self {
1586 stride: value.stride,
1587 padding: value.padding,
1588 dilation: value.dilation,
1589 weight_groups: value.weight_groups,
1590 offset_groups: value.offset_groups,
1591 }
1592 }
1593}
1594
1595impl From<ConvTransposeOptions<1>> for ConvTranspose1dOptionsIr {
1596 fn from(value: ConvTransposeOptions<1>) -> Self {
1597 Self {
1598 stride: value.stride,
1599 padding: value.padding,
1600 padding_out: value.padding_out,
1601 dilation: value.dilation,
1602 groups: value.groups,
1603 }
1604 }
1605}
1606
1607impl From<ConvTransposeOptions<2>> for ConvTranspose2dOptionsIr {
1608 fn from(value: ConvTransposeOptions<2>) -> Self {
1609 Self {
1610 stride: value.stride,
1611 padding: value.padding,
1612 padding_out: value.padding_out,
1613 dilation: value.dilation,
1614 groups: value.groups,
1615 }
1616 }
1617}
1618
1619impl From<ConvTransposeOptions<3>> for ConvTranspose3dOptionsIr {
1620 fn from(value: ConvTransposeOptions<3>) -> Self {
1621 Self {
1622 stride: value.stride,
1623 padding: value.padding,
1624 padding_out: value.padding_out,
1625 dilation: value.dilation,
1626 groups: value.groups,
1627 }
1628 }
1629}
1630
1631impl From<Conv1dOptionsIr> for ConvOptions<1> {
1632 fn from(val: Conv1dOptionsIr) -> Self {
1633 ConvOptions {
1634 stride: val.stride,
1635 padding: val.padding,
1636 dilation: val.dilation,
1637 groups: val.groups,
1638 }
1639 }
1640}
1641
1642impl From<Conv2dOptionsIr> for ConvOptions<2> {
1643 fn from(val: Conv2dOptionsIr) -> Self {
1644 ConvOptions {
1645 stride: val.stride,
1646 padding: val.padding,
1647 dilation: val.dilation,
1648 groups: val.groups,
1649 }
1650 }
1651}
1652
1653impl From<Conv3dOptionsIr> for ConvOptions<3> {
1654 fn from(val: Conv3dOptionsIr) -> Self {
1655 ConvOptions {
1656 stride: val.stride,
1657 padding: val.padding,
1658 dilation: val.dilation,
1659 groups: val.groups,
1660 }
1661 }
1662}
1663
1664impl From<DeformableConv2dOptionsIr> for DeformConvOptions<2> {
1665 fn from(value: DeformableConv2dOptionsIr) -> Self {
1666 DeformConvOptions {
1667 stride: value.stride,
1668 padding: value.padding,
1669 dilation: value.dilation,
1670 weight_groups: value.weight_groups,
1671 offset_groups: value.offset_groups,
1672 }
1673 }
1674}
1675
1676impl From<ConvTranspose1dOptionsIr> for ConvTransposeOptions<1> {
1677 fn from(val: ConvTranspose1dOptionsIr) -> Self {
1678 ConvTransposeOptions {
1679 stride: val.stride,
1680 padding: val.padding,
1681 padding_out: val.padding_out,
1682 dilation: val.dilation,
1683 groups: val.groups,
1684 }
1685 }
1686}
1687
1688impl From<ConvTranspose2dOptionsIr> for ConvTransposeOptions<2> {
1689 fn from(val: ConvTranspose2dOptionsIr) -> Self {
1690 ConvTransposeOptions {
1691 stride: val.stride,
1692 padding: val.padding,
1693 padding_out: val.padding_out,
1694 dilation: val.dilation,
1695 groups: val.groups,
1696 }
1697 }
1698}
1699
1700impl From<ConvTranspose3dOptionsIr> for ConvTransposeOptions<3> {
1701 fn from(val: ConvTranspose3dOptionsIr) -> Self {
1702 ConvTransposeOptions {
1703 stride: val.stride,
1704 padding: val.padding,
1705 padding_out: val.padding_out,
1706 dilation: val.dilation,
1707 groups: val.groups,
1708 }
1709 }
1710}
1711
1712#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1713#[allow(missing_docs)]
1714pub struct AvgPool1dOpIr {
1715 pub x: TensorIr,
1716 pub kernel_size: usize,
1717 pub stride: usize,
1718 pub padding: usize,
1719 pub count_include_pad: bool,
1720 pub ceil_mode: bool,
1721 pub out: TensorIr,
1722}
1723
1724#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1725#[allow(missing_docs)]
1726pub struct AvgPool2dOpIr {
1727 pub x: TensorIr,
1728 pub kernel_size: [usize; 2],
1729 pub stride: [usize; 2],
1730 pub padding: [usize; 2],
1731 pub count_include_pad: bool,
1732 pub ceil_mode: bool,
1733 pub out: TensorIr,
1734}
1735
1736#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1737#[allow(missing_docs)]
1738pub struct AvgPool1dBackwardOpIr {
1739 pub x: TensorIr,
1740 pub grad: TensorIr,
1741 pub kernel_size: usize,
1742 pub stride: usize,
1743 pub padding: usize,
1744 pub count_include_pad: bool,
1745 pub ceil_mode: bool,
1746 pub out: TensorIr,
1747}
1748
1749#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1750#[allow(missing_docs)]
1751pub struct AvgPool2dBackwardOpIr {
1752 pub x: TensorIr,
1753 pub grad: TensorIr,
1754 pub kernel_size: [usize; 2],
1755 pub stride: [usize; 2],
1756 pub padding: [usize; 2],
1757 pub count_include_pad: bool,
1758 pub ceil_mode: bool,
1759 pub out: TensorIr,
1760}
1761
1762#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1763#[allow(missing_docs)]
1764pub struct AdaptiveAvgPool1dOpIr {
1765 pub x: TensorIr,
1766 pub output_size: usize,
1767 pub out: TensorIr,
1768}
1769
1770#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1771#[allow(missing_docs)]
1772pub struct AdaptiveAvgPool2dOpIr {
1773 pub x: TensorIr,
1774 pub output_size: [usize; 2],
1775 pub out: TensorIr,
1776}
1777
1778#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1779#[allow(missing_docs)]
1780pub struct AdaptiveAvgPool1dBackwardOpIr {
1781 pub x: TensorIr,
1782 pub grad: TensorIr,
1783 pub out: TensorIr,
1784}
1785
1786#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1787#[allow(missing_docs)]
1788pub struct AdaptiveAvgPool2dBackwardOpIr {
1789 pub x: TensorIr,
1790 pub grad: TensorIr,
1791 pub out: TensorIr,
1792}
1793
1794#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1795#[allow(missing_docs)]
1796pub struct AdaptiveAvgPool3dOpIr {
1797 pub x: TensorIr,
1798 pub output_size: [usize; 3],
1799 pub out: TensorIr,
1800}
1801
1802#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1803#[allow(missing_docs)]
1804pub struct AdaptiveAvgPool3dBackwardOpIr {
1805 pub x: TensorIr,
1806 pub grad: TensorIr,
1807 pub out: TensorIr,
1808}
1809
1810#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1811#[allow(missing_docs)]
1812pub struct MaxPool1dOpIr {
1813 pub x: TensorIr,
1814 pub kernel_size: usize,
1815 pub stride: usize,
1816 pub padding: usize,
1817 pub dilation: usize,
1818 pub ceil_mode: bool,
1819 pub out: TensorIr,
1820}
1821
1822#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1823#[allow(missing_docs)]
1824pub struct MaxPool1dWithIndicesOpIr {
1825 pub x: TensorIr,
1826 pub kernel_size: usize,
1827 pub stride: usize,
1828 pub padding: usize,
1829 pub dilation: usize,
1830 pub ceil_mode: bool,
1831 pub out: TensorIr,
1832 pub out_indices: TensorIr,
1833}
1834
1835#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1836#[allow(missing_docs)]
1837pub struct MaxPool1dWithIndicesBackwardOpIr {
1838 pub x: TensorIr,
1839 pub grad: TensorIr,
1840 pub indices: TensorIr,
1841 pub kernel_size: usize,
1842 pub stride: usize,
1843 pub padding: usize,
1844 pub dilation: usize,
1845 pub ceil_mode: bool,
1846 pub out: TensorIr,
1847}
1848
1849#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1850#[allow(missing_docs)]
1851pub struct MaxPool2dOpIr {
1852 pub x: TensorIr,
1853 pub kernel_size: [usize; 2],
1854 pub stride: [usize; 2],
1855 pub padding: [usize; 2],
1856 pub dilation: [usize; 2],
1857 pub ceil_mode: bool,
1858 pub out: TensorIr,
1859}
1860
1861#[allow(missing_docs)]
1862#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1863pub struct MaxPool2dWithIndicesOpIr {
1864 pub x: TensorIr,
1865 pub kernel_size: [usize; 2],
1866 pub stride: [usize; 2],
1867 pub padding: [usize; 2],
1868 pub dilation: [usize; 2],
1869 pub ceil_mode: bool,
1870 pub out: TensorIr,
1871 pub out_indices: TensorIr,
1872}
1873
1874#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1875#[allow(missing_docs)]
1876pub struct MaxPool2dWithIndicesBackwardOpIr {
1877 pub x: TensorIr,
1878 pub grad: TensorIr,
1879 pub indices: TensorIr,
1880 pub kernel_size: [usize; 2],
1881 pub stride: [usize; 2],
1882 pub padding: [usize; 2],
1883 pub dilation: [usize; 2],
1884 pub ceil_mode: bool,
1885 pub out: TensorIr,
1886}
1887
1888#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1889#[allow(missing_docs)]
1890pub enum InterpolateModeIr {
1891 Nearest,
1892 NearestExact,
1893 Bilinear,
1894 Bicubic,
1895 Lanczos3,
1896}
1897
1898#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1899#[allow(missing_docs)]
1900pub struct InterpolateOptionsIr {
1901 pub mode: InterpolateModeIr,
1902 pub align_corners: bool,
1903}
1904
1905#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1906#[allow(missing_docs)]
1907pub struct InterpolateOpIr {
1908 pub x: TensorIr,
1909 pub output_size: [usize; 2],
1910 pub options: InterpolateOptionsIr,
1911 pub out: TensorIr,
1912}
1913
1914#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1915#[allow(missing_docs)]
1916pub struct RfftOpIr {
1917 pub signal: TensorIr,
1918 pub dim: usize,
1919 pub n: Option<usize>,
1920 pub out_re: TensorIr,
1921 pub out_im: TensorIr,
1922}
1923
1924#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1925#[allow(missing_docs)]
1926pub struct IRfftOpIr {
1927 pub input_re: TensorIr,
1928 pub input_im: TensorIr,
1929 pub dim: usize,
1930 pub n: Option<usize>,
1931 pub out_signal: TensorIr,
1932}
1933
1934#[allow(missing_docs)]
1935impl RfftOpIr {
1936 pub fn create<F>(signal: TensorIr, dim: usize, n: Option<usize>, mut new_id: F) -> Self
1937 where
1938 F: FnMut() -> crate::TensorId,
1939 {
1940 let mut shape = signal.shape.clone();
1943 let fft_len = n.unwrap_or(shape[dim]);
1944 shape[dim] = fft_len / 2 + 1;
1945 let dtype = signal.dtype;
1946
1947 Self {
1948 signal,
1949 dim,
1950 n,
1951 out_re: TensorIr::uninit(new_id(), shape.clone(), dtype),
1952 out_im: TensorIr::uninit(new_id(), shape, dtype),
1953 }
1954 }
1955}
1956
1957#[allow(missing_docs)]
1958impl IRfftOpIr {
1959 pub fn create<F>(
1960 input_re: TensorIr,
1961 input_im: TensorIr,
1962 dim: usize,
1963 n: Option<usize>,
1964 mut new_id: F,
1965 ) -> Self
1966 where
1967 F: FnMut() -> crate::TensorId,
1968 {
1969 debug_assert!(
1970 input_re.shape[dim] >= 1,
1971 "IRfftOpIr: input spectrum dimension must be >= 1"
1972 );
1973 debug_assert!(
1974 !matches!(n, Some(0)),
1975 "IRfftOpIr: n must be >= 1 when specified"
1976 );
1977 let mut shape = input_re.shape.clone();
1978 shape[dim] = n.unwrap_or((shape[dim] - 1) * 2);
1979 let dtype = input_re.dtype;
1980
1981 Self {
1982 input_re,
1983 input_im,
1984 dim,
1985 n,
1986 out_signal: TensorIr::uninit(new_id(), shape, dtype),
1987 }
1988 }
1989}
1990
1991#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1992#[allow(missing_docs)]
1993pub struct AttentionOptionsIr {
1994 pub scale: Option<ScalarIr>,
1995 pub softcap: Option<ScalarIr>,
1996 pub is_causal: bool,
1997}
1998
1999impl From<AttentionOptionsIr> for AttentionModuleOptions {
2000 fn from(ir: AttentionOptionsIr) -> Self {
2001 AttentionModuleOptions {
2002 scale: ir.scale.map(|s| s.elem()),
2003 softcap: ir.softcap.map(|s| s.elem()),
2004 is_causal: ir.is_causal,
2005 }
2006 }
2007}
2008
2009impl From<AttentionModuleOptions> for AttentionOptionsIr {
2010 fn from(ir: AttentionModuleOptions) -> Self {
2011 AttentionOptionsIr {
2012 scale: ir.scale.map(ScalarIr::Float),
2013 softcap: ir.softcap.map(ScalarIr::Float),
2014 is_causal: ir.is_causal,
2015 }
2016 }
2017}
2018
2019#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
2020#[allow(missing_docs)]
2021pub struct AttentionOpIr {
2022 pub query: TensorIr,
2023 pub key: TensorIr,
2024 pub value: TensorIr,
2025 pub mask: Option<TensorIr>,
2026 pub attn_bias: Option<TensorIr>,
2027 pub options: AttentionOptionsIr,
2028 pub out: TensorIr,
2029}
2030
2031#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
2032#[allow(missing_docs)]
2033pub struct CtcLossOpIr {
2034 pub log_probs: TensorIr,
2035 pub targets: TensorIr,
2036 pub input_lengths: TensorIr,
2037 pub target_lengths: TensorIr,
2038 pub blank: usize,
2039 pub out: TensorIr,
2040}
2041
2042#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
2043#[allow(missing_docs)]
2044pub struct CtcLossBackwardOpIr {
2045 pub log_probs: TensorIr,
2046 pub targets: TensorIr,
2047 pub input_lengths: TensorIr,
2048 pub target_lengths: TensorIr,
2049 pub grad_loss: TensorIr,
2050 pub blank: usize,
2051 pub out: TensorIr,
2052}
2053
2054impl From<InterpolateModeIr> for InterpolateMode {
2055 fn from(val: InterpolateModeIr) -> Self {
2056 match val {
2057 InterpolateModeIr::Nearest => Self::Nearest,
2058 InterpolateModeIr::NearestExact => Self::NearestExact,
2059 InterpolateModeIr::Bilinear => Self::Bilinear,
2060 InterpolateModeIr::Bicubic => Self::Bicubic,
2061 InterpolateModeIr::Lanczos3 => Self::Lanczos3,
2062 }
2063 }
2064}
2065
2066impl From<InterpolateOptionsIr> for InterpolateOptions {
2067 fn from(val: InterpolateOptionsIr) -> Self {
2068 Self::new(val.mode.into()).with_align_corners(val.align_corners)
2069 }
2070}
2071
2072impl From<InterpolateMode> for InterpolateModeIr {
2073 fn from(val: InterpolateMode) -> Self {
2074 match val {
2075 InterpolateMode::Nearest => Self::Nearest,
2076 InterpolateMode::NearestExact => Self::NearestExact,
2077 InterpolateMode::Bilinear => Self::Bilinear,
2078 InterpolateMode::Bicubic => Self::Bicubic,
2079 InterpolateMode::Lanczos3 => Self::Lanczos3,
2080 }
2081 }
2082}
2083
2084impl From<InterpolateOptions> for InterpolateOptionsIr {
2085 fn from(val: InterpolateOptions) -> Self {
2086 Self {
2087 mode: val.mode.into(),
2088 align_corners: val.align_corners,
2089 }
2090 }
2091}
2092
2093#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
2094#[allow(missing_docs)]
2095pub struct InterpolateBackwardOpIr {
2096 pub x: TensorIr,
2097 pub grad: TensorIr,
2098 pub output_size: [usize; 2],
2099 pub options: InterpolateOptionsIr,
2100 pub out: TensorIr,
2101}
2102
2103#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
2104#[allow(missing_docs)]
2105pub enum GridSamplePaddingModeIr {
2106 Zeros,
2107 Border,
2108 Reflection,
2109}
2110
2111#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
2112#[allow(missing_docs)]
2113pub struct GridSampleOptionsIr {
2114 pub mode: InterpolateModeIr,
2115 pub padding_mode: GridSamplePaddingModeIr,
2116 pub align_corners: bool,
2117}
2118
2119#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
2120#[allow(missing_docs)]
2121pub struct GridSample2dOpIr {
2122 pub tensor: TensorIr,
2123 pub grid: TensorIr,
2124 pub options: GridSampleOptionsIr,
2125 pub out: TensorIr,
2126}
2127
2128impl From<GridSamplePaddingModeIr> for GridSamplePaddingMode {
2129 fn from(val: GridSamplePaddingModeIr) -> Self {
2130 match val {
2131 GridSamplePaddingModeIr::Zeros => Self::Zeros,
2132 GridSamplePaddingModeIr::Border => Self::Border,
2133 GridSamplePaddingModeIr::Reflection => Self::Reflection,
2134 }
2135 }
2136}
2137
2138impl From<GridSamplePaddingMode> for GridSamplePaddingModeIr {
2139 fn from(val: GridSamplePaddingMode) -> Self {
2140 match val {
2141 GridSamplePaddingMode::Zeros => Self::Zeros,
2142 GridSamplePaddingMode::Border => Self::Border,
2143 GridSamplePaddingMode::Reflection => Self::Reflection,
2144 }
2145 }
2146}
2147
2148impl From<GridSampleOptionsIr> for GridSampleOptions {
2149 fn from(val: GridSampleOptionsIr) -> Self {
2150 Self {
2151 mode: val.mode.into(),
2152 padding_mode: val.padding_mode.into(),
2153 align_corners: val.align_corners,
2154 }
2155 }
2156}
2157
2158impl From<GridSampleOptions> for GridSampleOptionsIr {
2159 fn from(val: GridSampleOptions) -> Self {
2160 Self {
2161 mode: val.mode.into(),
2162 padding_mode: val.padding_mode.into(),
2163 align_corners: val.align_corners,
2164 }
2165 }
2166}
2167
2168impl OperationIr {
2169 pub fn inputs(&self) -> impl Iterator<Item = &TensorIr> {
2171 match self {
2172 OperationIr::BaseFloat(repr) => repr.inputs(),
2173 OperationIr::BaseInt(repr) => repr.inputs(),
2174 OperationIr::BaseBool(repr) => repr.inputs(),
2175 OperationIr::NumericFloat(_dtype, repr) => repr.inputs(),
2176 OperationIr::NumericInt(_dtype, repr) => repr.inputs(),
2177 OperationIr::Bool(repr) => repr.inputs(),
2178 OperationIr::Int(repr) => repr.inputs(),
2179 OperationIr::Float(_dtype, repr) => repr.inputs(),
2180 OperationIr::Module(repr) => repr.inputs(),
2181 OperationIr::Init(repr) => repr.inputs(),
2182 OperationIr::Custom(repr) => repr.inputs(),
2183 OperationIr::Drop(repr) => Box::new([repr].into_iter()),
2184 OperationIr::Distributed(repr) => repr.inputs(),
2185 OperationIr::Activation(repr) => repr.inputs(),
2186 }
2187 }
2188
2189 pub fn outputs(&self) -> impl Iterator<Item = &TensorIr> {
2191 match self {
2192 OperationIr::BaseFloat(repr) => repr.outputs(),
2193 OperationIr::BaseInt(repr) => repr.outputs(),
2194 OperationIr::BaseBool(repr) => repr.outputs(),
2195 OperationIr::NumericFloat(_dtype, repr) => repr.outputs(),
2196 OperationIr::NumericInt(_dtype, repr) => repr.outputs(),
2197 OperationIr::Bool(repr) => repr.outputs(),
2198 OperationIr::Int(repr) => repr.outputs(),
2199 OperationIr::Float(_dtype, repr) => repr.outputs(),
2200 OperationIr::Module(repr) => repr.outputs(),
2201 OperationIr::Init(repr) => repr.outputs(),
2202 OperationIr::Custom(repr) => repr.outputs(),
2203 OperationIr::Drop(_repr) => Box::new([].into_iter()),
2204 OperationIr::Distributed(repr) => repr.outputs(),
2205 OperationIr::Activation(repr) => repr.outputs(),
2206 }
2207 }
2208
2209 pub fn nodes(&self) -> Vec<&TensorIr> {
2211 self.inputs().chain(self.outputs()).collect()
2212 }
2213
2214 pub fn mark_read_only(&mut self, nodes: &[TensorId]) -> Vec<TensorIr> {
2219 match self {
2220 OperationIr::BaseFloat(repr) => repr.mark_read_only(nodes),
2221 OperationIr::BaseInt(repr) => repr.mark_read_only(nodes),
2222 OperationIr::BaseBool(repr) => repr.mark_read_only(nodes),
2223 OperationIr::NumericFloat(_dtype, repr) => repr.mark_read_only(nodes),
2224 OperationIr::NumericInt(_dtype, repr) => repr.mark_read_only(nodes),
2225 OperationIr::Bool(repr) => repr.mark_read_only(nodes),
2226 OperationIr::Int(repr) => repr.mark_read_only(nodes),
2227 OperationIr::Float(_dtype, repr) => repr.mark_read_only(nodes),
2228 OperationIr::Module(repr) => repr.mark_read_only(nodes),
2229 OperationIr::Init(_) => Vec::new(),
2230 OperationIr::Drop(repr) => {
2231 let mut output = Vec::new();
2232 repr.mark_read_only(nodes, &mut output);
2233 output
2234 }
2235 OperationIr::Custom(repr) => {
2236 let mut output = Vec::new();
2237
2238 for input in repr.inputs.iter_mut() {
2239 input.mark_read_only(nodes, &mut output);
2240 }
2241
2242 output
2243 }
2244 OperationIr::Distributed(repr) => repr.mark_read_only(nodes),
2245 OperationIr::Activation(repr) => repr.mark_read_only(nodes),
2246 }
2247 }
2248
2249 pub fn visit_mut(&mut self, v: &mut impl IrVisitorMut) {
2256 match self {
2257 OperationIr::BaseFloat(repr) => repr.visit_mut(v),
2258 OperationIr::BaseInt(repr) => repr.visit_mut(v),
2259 OperationIr::BaseBool(repr) => repr.visit_mut(v),
2260 OperationIr::NumericFloat(_dtype, repr) => repr.visit_mut(v),
2261 OperationIr::NumericInt(_dtype, repr) => repr.visit_mut(v),
2262 OperationIr::Bool(repr) => repr.visit_mut(v),
2263 OperationIr::Int(repr) => repr.visit_mut(v),
2264 OperationIr::Float(_dtype, repr) => repr.visit_mut(v),
2265 OperationIr::Module(repr) => repr.visit_mut(v),
2266 OperationIr::Init(repr) => repr.visit_mut(v),
2267 OperationIr::Custom(repr) => repr.visit_mut(v),
2268 OperationIr::Drop(repr) => v.visit_tensor_mut(repr),
2269 OperationIr::Distributed(repr) => repr.visit_mut(v),
2270 OperationIr::Activation(repr) => repr.visit_mut(v),
2271 }
2272 }
2273}
2274
2275impl BaseOperationIr {
2276 fn inputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
2277 match self {
2278 BaseOperationIr::Reshape(repr) => Box::new([&repr.input].into_iter()),
2279 BaseOperationIr::SwapDims(repr) => Box::new([&repr.input].into_iter()),
2280 BaseOperationIr::Permute(repr) => Box::new([&repr.input].into_iter()),
2281 BaseOperationIr::Expand(repr) => Box::new([&repr.input].into_iter()),
2282 BaseOperationIr::Flip(repr) => Box::new([&repr.input].into_iter()),
2283 BaseOperationIr::Slice(repr) => Box::new([&repr.tensor].into_iter()),
2284 BaseOperationIr::SliceAssign(repr) => Box::new([&repr.tensor, &repr.value].into_iter()),
2285 BaseOperationIr::Gather(repr) => Box::new([&repr.tensor, &repr.indices].into_iter()),
2286 BaseOperationIr::Scatter(repr) => {
2287 Box::new([&repr.tensor, &repr.indices, &repr.value].into_iter())
2288 }
2289 BaseOperationIr::ScatterNd(repr) => {
2290 Box::new([&repr.data, &repr.indices, &repr.values].into_iter())
2291 }
2292 BaseOperationIr::GatherNd(repr) => Box::new([&repr.data, &repr.indices].into_iter()),
2293 BaseOperationIr::Select(repr) => Box::new([&repr.tensor, &repr.indices].into_iter()),
2294 BaseOperationIr::SelectAssign(repr) => {
2295 Box::new([&repr.tensor, &repr.indices, &repr.value].into_iter())
2296 }
2297 BaseOperationIr::MaskWhere(repr) => {
2298 Box::new([&repr.tensor, &repr.mask, &repr.value].into_iter())
2299 }
2300 BaseOperationIr::MaskFill(repr) => Box::new([&repr.tensor, &repr.mask].into_iter()),
2301 BaseOperationIr::Equal(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2302 BaseOperationIr::EqualElem(repr) => Box::new([&repr.lhs].into_iter()),
2303 BaseOperationIr::RepeatDim(repr) => Box::new([&repr.tensor].into_iter()),
2304 BaseOperationIr::Cat(repr) => Box::new(repr.tensors.iter()),
2305 BaseOperationIr::Cast(repr) => Box::new([&repr.input].into_iter()),
2306 BaseOperationIr::Unfold(repr) => Box::new([&repr.input].into_iter()),
2307 BaseOperationIr::Empty(_repr) => Box::new([].into_iter()),
2308 BaseOperationIr::Ones(_repr) => Box::new([].into_iter()),
2309 BaseOperationIr::Zeros(_repr) => Box::new([].into_iter()),
2310 BaseOperationIr::NotEqual(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2311 BaseOperationIr::NotEqualElem(repr) => Box::new([&repr.lhs].into_iter()),
2312 BaseOperationIr::All(repr) => Box::new([&repr.input].into_iter()),
2313 BaseOperationIr::Any(repr) => Box::new([&repr.input].into_iter()),
2314 BaseOperationIr::AllDim(repr) => Box::new([&repr.input].into_iter()),
2315 BaseOperationIr::AnyDim(repr) => Box::new([&repr.input].into_iter()),
2316 }
2317 }
2318
2319 fn outputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
2320 match self {
2321 BaseOperationIr::Reshape(repr) => Box::new([&repr.out].into_iter()),
2322 BaseOperationIr::SwapDims(repr) => Box::new([&repr.out].into_iter()),
2323 BaseOperationIr::Permute(repr) => Box::new([&repr.out].into_iter()),
2324 BaseOperationIr::Expand(repr) => Box::new([&repr.out].into_iter()),
2325 BaseOperationIr::Flip(repr) => Box::new([&repr.out].into_iter()),
2326 BaseOperationIr::Slice(repr) => Box::new([&repr.out].into_iter()),
2327 BaseOperationIr::SliceAssign(repr) => Box::new([&repr.out].into_iter()),
2328 BaseOperationIr::Gather(repr) => Box::new([&repr.out].into_iter()),
2329 BaseOperationIr::Scatter(repr) => Box::new([&repr.out].into_iter()),
2330 BaseOperationIr::ScatterNd(repr) => Box::new([&repr.out].into_iter()),
2331 BaseOperationIr::GatherNd(repr) => Box::new([&repr.out].into_iter()),
2332 BaseOperationIr::Select(repr) => Box::new([&repr.out].into_iter()),
2333 BaseOperationIr::SelectAssign(repr) => Box::new([&repr.out].into_iter()),
2334 BaseOperationIr::MaskWhere(repr) => Box::new([&repr.out].into_iter()),
2335 BaseOperationIr::MaskFill(repr) => Box::new([&repr.out].into_iter()),
2336 BaseOperationIr::Equal(repr) => Box::new([&repr.out].into_iter()),
2337 BaseOperationIr::EqualElem(repr) => Box::new([&repr.out].into_iter()),
2338 BaseOperationIr::RepeatDim(repr) => Box::new([&repr.out].into_iter()),
2339 BaseOperationIr::Cat(repr) => Box::new([&repr.out].into_iter()),
2340 BaseOperationIr::Cast(repr) => Box::new([&repr.out].into_iter()),
2341 BaseOperationIr::Unfold(repr) => Box::new([&repr.out].into_iter()),
2342 BaseOperationIr::Empty(repr) => Box::new([&repr.out].into_iter()),
2343 BaseOperationIr::Ones(repr) => Box::new([&repr.out].into_iter()),
2344 BaseOperationIr::Zeros(repr) => Box::new([&repr.out].into_iter()),
2345 BaseOperationIr::NotEqual(repr) => Box::new([&repr.out].into_iter()),
2346 BaseOperationIr::NotEqualElem(repr) => Box::new([&repr.out].into_iter()),
2347 BaseOperationIr::All(repr) => Box::new([&repr.out].into_iter()),
2348 BaseOperationIr::Any(repr) => Box::new([&repr.out].into_iter()),
2349 BaseOperationIr::AllDim(repr) => Box::new([&repr.out].into_iter()),
2350 BaseOperationIr::AnyDim(repr) => Box::new([&repr.out].into_iter()),
2351 }
2352 }
2353
2354 fn mark_read_only(&mut self, nodes: &[TensorId]) -> Vec<TensorIr> {
2355 let mut output = Vec::new();
2356
2357 match self {
2358 BaseOperationIr::Reshape(repr) => {
2359 repr.input.mark_read_only(nodes, &mut output);
2360 }
2361 BaseOperationIr::SwapDims(repr) => {
2362 repr.input.mark_read_only(nodes, &mut output);
2363 }
2364 BaseOperationIr::Permute(repr) => {
2365 repr.input.mark_read_only(nodes, &mut output);
2366 }
2367
2368 BaseOperationIr::Expand(repr) => {
2369 repr.input.mark_read_only(nodes, &mut output);
2370 }
2371
2372 BaseOperationIr::Flip(repr) => {
2373 repr.input.mark_read_only(nodes, &mut output);
2374 }
2375 BaseOperationIr::Slice(repr) => {
2376 repr.tensor.mark_read_only(nodes, &mut output);
2377 }
2378 BaseOperationIr::SliceAssign(repr) => {
2379 repr.tensor.mark_read_only(nodes, &mut output);
2380 repr.value.mark_read_only(nodes, &mut output);
2381 }
2382 BaseOperationIr::Gather(repr) => {
2383 repr.tensor.mark_read_only(nodes, &mut output);
2384 repr.indices.mark_read_only(nodes, &mut output);
2385 }
2386 BaseOperationIr::Scatter(repr) => {
2387 repr.tensor.mark_read_only(nodes, &mut output);
2388 repr.indices.mark_read_only(nodes, &mut output);
2389 repr.value.mark_read_only(nodes, &mut output);
2390 }
2391 BaseOperationIr::ScatterNd(repr) => {
2392 repr.data.mark_read_only(nodes, &mut output);
2393 repr.indices.mark_read_only(nodes, &mut output);
2394 repr.values.mark_read_only(nodes, &mut output);
2395 }
2396 BaseOperationIr::GatherNd(repr) => {
2397 repr.data.mark_read_only(nodes, &mut output);
2398 repr.indices.mark_read_only(nodes, &mut output);
2399 }
2400 BaseOperationIr::Select(repr) => {
2401 repr.tensor.mark_read_only(nodes, &mut output);
2402 repr.indices.mark_read_only(nodes, &mut output);
2403 }
2404 BaseOperationIr::SelectAssign(repr) => {
2405 repr.tensor.mark_read_only(nodes, &mut output);
2406 repr.indices.mark_read_only(nodes, &mut output);
2407 repr.value.mark_read_only(nodes, &mut output);
2408 }
2409 BaseOperationIr::MaskWhere(repr) => {
2410 repr.tensor.mark_read_only(nodes, &mut output);
2411 repr.mask.mark_read_only(nodes, &mut output);
2412 repr.value.mark_read_only(nodes, &mut output);
2413 }
2414 BaseOperationIr::MaskFill(repr) => {
2415 repr.tensor.mark_read_only(nodes, &mut output);
2416 repr.mask.mark_read_only(nodes, &mut output);
2417 }
2418 BaseOperationIr::Equal(repr) => {
2419 repr.lhs.mark_read_only(nodes, &mut output);
2420 repr.rhs.mark_read_only(nodes, &mut output);
2421 }
2422 BaseOperationIr::EqualElem(repr) => {
2423 repr.lhs.mark_read_only(nodes, &mut output);
2424 }
2425 BaseOperationIr::RepeatDim(repr) => {
2426 repr.tensor.mark_read_only(nodes, &mut output);
2427 }
2428 BaseOperationIr::Cat(repr) => {
2429 for t in repr.tensors.iter_mut() {
2430 t.mark_read_only(nodes, &mut output);
2431 }
2432 }
2433 BaseOperationIr::Cast(repr) => {
2434 repr.input.mark_read_only(nodes, &mut output);
2435 }
2436 BaseOperationIr::Unfold(repr) => {
2437 repr.input.mark_read_only(nodes, &mut output);
2438 }
2439 BaseOperationIr::Empty(_) => {}
2440 BaseOperationIr::Zeros(_) => {}
2441 BaseOperationIr::Ones(_) => {}
2442 BaseOperationIr::NotEqual(repr) => {
2443 repr.lhs.mark_read_only(nodes, &mut output);
2444 repr.rhs.mark_read_only(nodes, &mut output);
2445 }
2446 BaseOperationIr::NotEqualElem(repr) => {
2447 repr.lhs.mark_read_only(nodes, &mut output);
2448 }
2449 BaseOperationIr::All(repr) => {
2450 repr.input.mark_read_only(nodes, &mut output);
2451 }
2452 BaseOperationIr::Any(repr) => {
2453 repr.input.mark_read_only(nodes, &mut output);
2454 }
2455 BaseOperationIr::AllDim(repr) => {
2456 repr.input.mark_read_only(nodes, &mut output);
2457 }
2458 BaseOperationIr::AnyDim(repr) => {
2459 repr.input.mark_read_only(nodes, &mut output);
2460 }
2461 };
2462
2463 output
2464 }
2465
2466 fn visit_mut(&mut self, v: &mut impl IrVisitorMut) {
2467 match self {
2468 BaseOperationIr::Reshape(repr) => {
2469 v.visit_tensor_mut(&mut repr.input);
2470 v.visit_tensor_mut(&mut repr.out);
2471 }
2472 BaseOperationIr::SwapDims(repr) => {
2473 v.visit_tensor_mut(&mut repr.input);
2474 v.visit_tensor_mut(&mut repr.out);
2475 }
2476 BaseOperationIr::Permute(repr) => {
2477 v.visit_tensor_mut(&mut repr.input);
2478 v.visit_tensor_mut(&mut repr.out);
2479 }
2480 BaseOperationIr::Expand(repr) => {
2481 v.visit_tensor_mut(&mut repr.input);
2482 v.visit_tensor_mut(&mut repr.out);
2483 }
2484 BaseOperationIr::Flip(repr) => {
2485 v.visit_tensor_mut(&mut repr.input);
2486 v.visit_tensor_mut(&mut repr.out);
2487 }
2488 BaseOperationIr::Slice(repr) => {
2489 v.visit_tensor_mut(&mut repr.tensor);
2490 v.visit_tensor_mut(&mut repr.out);
2491 repr.ranges.iter_mut().for_each(|r| v.visit_range_mut(r));
2492 }
2493 BaseOperationIr::SliceAssign(repr) => {
2494 v.visit_tensor_mut(&mut repr.tensor);
2495 v.visit_tensor_mut(&mut repr.value);
2496 v.visit_tensor_mut(&mut repr.out);
2497 repr.ranges.iter_mut().for_each(|r| v.visit_range_mut(r));
2498 }
2499 BaseOperationIr::Gather(repr) => {
2500 v.visit_tensor_mut(&mut repr.tensor);
2501 v.visit_tensor_mut(&mut repr.indices);
2502 v.visit_tensor_mut(&mut repr.out);
2503 }
2504 BaseOperationIr::Scatter(repr) => {
2505 v.visit_tensor_mut(&mut repr.tensor);
2506 v.visit_tensor_mut(&mut repr.indices);
2507 v.visit_tensor_mut(&mut repr.value);
2508 v.visit_tensor_mut(&mut repr.out);
2509 }
2510 BaseOperationIr::ScatterNd(repr) => {
2511 v.visit_tensor_mut(&mut repr.data);
2512 v.visit_tensor_mut(&mut repr.indices);
2513 v.visit_tensor_mut(&mut repr.values);
2514 v.visit_tensor_mut(&mut repr.out);
2515 }
2516 BaseOperationIr::GatherNd(repr) => {
2517 v.visit_tensor_mut(&mut repr.data);
2518 v.visit_tensor_mut(&mut repr.indices);
2519 v.visit_tensor_mut(&mut repr.out);
2520 }
2521 BaseOperationIr::Select(repr) => {
2522 v.visit_tensor_mut(&mut repr.tensor);
2523 v.visit_tensor_mut(&mut repr.indices);
2524 v.visit_tensor_mut(&mut repr.out);
2525 }
2526 BaseOperationIr::SelectAssign(repr) => {
2527 v.visit_tensor_mut(&mut repr.tensor);
2528 v.visit_tensor_mut(&mut repr.indices);
2529 v.visit_tensor_mut(&mut repr.value);
2530 v.visit_tensor_mut(&mut repr.out);
2531 }
2532 BaseOperationIr::MaskWhere(repr) => {
2533 v.visit_tensor_mut(&mut repr.tensor);
2534 v.visit_tensor_mut(&mut repr.mask);
2535 v.visit_tensor_mut(&mut repr.value);
2536 v.visit_tensor_mut(&mut repr.out);
2537 }
2538 BaseOperationIr::MaskFill(repr) => {
2539 v.visit_tensor_mut(&mut repr.tensor);
2540 v.visit_tensor_mut(&mut repr.mask);
2541 v.visit_tensor_mut(&mut repr.out);
2542 v.visit_scalar_mut(&mut repr.value);
2543 }
2544 BaseOperationIr::Equal(repr) => {
2545 v.visit_tensor_mut(&mut repr.lhs);
2546 v.visit_tensor_mut(&mut repr.rhs);
2547 v.visit_tensor_mut(&mut repr.out);
2548 }
2549 BaseOperationIr::EqualElem(repr) => {
2550 v.visit_tensor_mut(&mut repr.lhs);
2551 v.visit_tensor_mut(&mut repr.out);
2552 v.visit_scalar_mut(&mut repr.rhs);
2553 }
2554 BaseOperationIr::RepeatDim(repr) => {
2555 v.visit_tensor_mut(&mut repr.tensor);
2556 v.visit_tensor_mut(&mut repr.out);
2557 }
2558 BaseOperationIr::Cat(repr) => {
2559 for t in repr.tensors.iter_mut() {
2560 v.visit_tensor_mut(t);
2561 }
2562 v.visit_tensor_mut(&mut repr.out);
2563 }
2564 BaseOperationIr::Cast(repr) => {
2565 v.visit_tensor_mut(&mut repr.input);
2566 v.visit_tensor_mut(&mut repr.out);
2567 }
2568 BaseOperationIr::Unfold(repr) => {
2569 v.visit_tensor_mut(&mut repr.input);
2570 v.visit_tensor_mut(&mut repr.out);
2571 }
2572 BaseOperationIr::Empty(repr) => {
2573 v.visit_tensor_mut(&mut repr.out);
2574 }
2575 BaseOperationIr::Ones(repr) => {
2576 v.visit_tensor_mut(&mut repr.out);
2577 }
2578 BaseOperationIr::Zeros(repr) => {
2579 v.visit_tensor_mut(&mut repr.out);
2580 }
2581 BaseOperationIr::NotEqual(repr) => {
2582 v.visit_tensor_mut(&mut repr.lhs);
2583 v.visit_tensor_mut(&mut repr.rhs);
2584 v.visit_tensor_mut(&mut repr.out);
2585 }
2586 BaseOperationIr::NotEqualElem(repr) => {
2587 v.visit_tensor_mut(&mut repr.lhs);
2588 v.visit_tensor_mut(&mut repr.out);
2589 v.visit_scalar_mut(&mut repr.rhs);
2590 }
2591 BaseOperationIr::All(repr) => {
2592 v.visit_tensor_mut(&mut repr.input);
2593 v.visit_tensor_mut(&mut repr.out);
2594 }
2595 BaseOperationIr::Any(repr) => {
2596 v.visit_tensor_mut(&mut repr.input);
2597 v.visit_tensor_mut(&mut repr.out);
2598 }
2599 BaseOperationIr::AllDim(repr) => {
2600 v.visit_tensor_mut(&mut repr.input);
2601 v.visit_tensor_mut(&mut repr.out);
2602 }
2603 BaseOperationIr::AnyDim(repr) => {
2604 v.visit_tensor_mut(&mut repr.input);
2605 v.visit_tensor_mut(&mut repr.out);
2606 }
2607 }
2608 }
2609}
2610
2611impl NumericOperationIr {
2612 fn inputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
2613 match self {
2614 NumericOperationIr::Add(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2615 NumericOperationIr::AddScalar(repr) => Box::new([&repr.lhs].into_iter()),
2616 NumericOperationIr::Sub(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2617 NumericOperationIr::SubScalar(repr) => Box::new([&repr.lhs].into_iter()),
2618 NumericOperationIr::Mul(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2619 NumericOperationIr::MulScalar(repr) => Box::new([&repr.lhs].into_iter()),
2620 NumericOperationIr::Div(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2621 NumericOperationIr::DivScalar(repr) => Box::new([&repr.lhs].into_iter()),
2622 NumericOperationIr::Rem(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2623 NumericOperationIr::RemScalar(repr) => Box::new([&repr.lhs].into_iter()),
2624 NumericOperationIr::GreaterElem(repr) => Box::new([&repr.lhs].into_iter()),
2625 NumericOperationIr::GreaterEqualElem(repr) => Box::new([&repr.lhs].into_iter()),
2626 NumericOperationIr::LowerElem(repr) => Box::new([&repr.lhs].into_iter()),
2627 NumericOperationIr::LowerEqualElem(repr) => Box::new([&repr.lhs].into_iter()),
2628 NumericOperationIr::Greater(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2629 NumericOperationIr::GreaterEqual(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2630 NumericOperationIr::Lower(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2631 NumericOperationIr::LowerEqual(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2632 NumericOperationIr::ArgMax(repr) => Box::new([&repr.input].into_iter()),
2633 NumericOperationIr::ArgTopK(repr) => Box::new([&repr.input].into_iter()),
2634 NumericOperationIr::TopK(repr) => Box::new([&repr.input].into_iter()),
2635 NumericOperationIr::ArgMin(repr) => Box::new([&repr.input].into_iter()),
2636 NumericOperationIr::Clamp(repr) => Box::new([&repr.tensor].into_iter()),
2637 NumericOperationIr::Abs(repr) => Box::new([&repr.input].into_iter()),
2638 NumericOperationIr::Full(_repr) => Box::new([].into_iter()),
2639 NumericOperationIr::MeanDim(repr) => Box::new([&repr.input].into_iter()),
2640 NumericOperationIr::Mean(repr) => Box::new([&repr.input].into_iter()),
2641 NumericOperationIr::Sum(repr) => Box::new([&repr.input].into_iter()),
2642 NumericOperationIr::SumDim(repr) => Box::new([&repr.input].into_iter()),
2643 NumericOperationIr::Prod(repr) => Box::new([&repr.input].into_iter()),
2644 NumericOperationIr::ProdDim(repr) => Box::new([&repr.input].into_iter()),
2645 NumericOperationIr::Max(repr) => Box::new([&repr.input].into_iter()),
2646 NumericOperationIr::MaxDimWithIndices(repr) => Box::new([&repr.tensor].into_iter()),
2647 NumericOperationIr::TopKWithIndices(repr) => Box::new([&repr.tensor].into_iter()),
2648 NumericOperationIr::MinDimWithIndices(repr) => Box::new([&repr.tensor].into_iter()),
2649 NumericOperationIr::Min(repr) => Box::new([&repr.input].into_iter()),
2650 NumericOperationIr::MaxDim(repr) => Box::new([&repr.input].into_iter()),
2651 NumericOperationIr::MinDim(repr) => Box::new([&repr.input].into_iter()),
2652 NumericOperationIr::MaxAbs(repr) => Box::new([&repr.input].into_iter()),
2653 NumericOperationIr::MaxAbsDim(repr) => Box::new([&repr.input].into_iter()),
2654 NumericOperationIr::IntRandom(_repr) => Box::new([].into_iter()),
2655 NumericOperationIr::Powi(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2656 NumericOperationIr::PowiScalar(repr) => Box::new([&repr.lhs].into_iter()),
2657 NumericOperationIr::CumMin(repr) => Box::new([&repr.input].into_iter()),
2658 NumericOperationIr::CumMax(repr) => Box::new([&repr.input].into_iter()),
2659 NumericOperationIr::CumProd(repr) => Box::new([&repr.input].into_iter()),
2660 NumericOperationIr::CumSum(repr) => Box::new([&repr.input].into_iter()),
2661 NumericOperationIr::Neg(repr) => Box::new([&repr.input].into_iter()),
2662 NumericOperationIr::Sign(repr) => Box::new([&repr.input].into_iter()),
2663 NumericOperationIr::ClampMin(repr) => Box::new([&repr.lhs].into_iter()),
2664 NumericOperationIr::ClampMax(repr) => Box::new([&repr.lhs].into_iter()),
2665 NumericOperationIr::Sort(repr) => Box::new([&repr.input].into_iter()),
2666 NumericOperationIr::SortWithIndices(repr) => Box::new([&repr.input].into_iter()),
2667 NumericOperationIr::ArgSort(repr) => Box::new([&repr.input].into_iter()),
2668 }
2669 }
2670
2671 fn outputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
2672 match self {
2673 NumericOperationIr::Add(repr) => Box::new([&repr.out].into_iter()),
2674 NumericOperationIr::AddScalar(repr) => Box::new([&repr.out].into_iter()),
2675 NumericOperationIr::Sub(repr) => Box::new([&repr.out].into_iter()),
2676 NumericOperationIr::SubScalar(repr) => Box::new([&repr.out].into_iter()),
2677 NumericOperationIr::Mul(repr) => Box::new([&repr.out].into_iter()),
2678 NumericOperationIr::MulScalar(repr) => Box::new([&repr.out].into_iter()),
2679 NumericOperationIr::Div(repr) => Box::new([&repr.out].into_iter()),
2680 NumericOperationIr::DivScalar(repr) => Box::new([&repr.out].into_iter()),
2681 NumericOperationIr::Rem(repr) => Box::new([&repr.out].into_iter()),
2682 NumericOperationIr::RemScalar(repr) => Box::new([&repr.out].into_iter()),
2683 NumericOperationIr::GreaterElem(repr) => Box::new([&repr.out].into_iter()),
2684 NumericOperationIr::GreaterEqualElem(repr) => Box::new([&repr.out].into_iter()),
2685 NumericOperationIr::LowerElem(repr) => Box::new([&repr.out].into_iter()),
2686 NumericOperationIr::LowerEqualElem(repr) => Box::new([&repr.out].into_iter()),
2687 NumericOperationIr::Greater(repr) => Box::new([&repr.out].into_iter()),
2688 NumericOperationIr::GreaterEqual(repr) => Box::new([&repr.out].into_iter()),
2689 NumericOperationIr::Lower(repr) => Box::new([&repr.out].into_iter()),
2690 NumericOperationIr::LowerEqual(repr) => Box::new([&repr.out].into_iter()),
2691 NumericOperationIr::ArgMax(repr) => Box::new([&repr.out].into_iter()),
2692 NumericOperationIr::ArgTopK(repr) => Box::new([&repr.out].into_iter()),
2693 NumericOperationIr::TopK(repr) => Box::new([&repr.out].into_iter()),
2694 NumericOperationIr::ArgMin(repr) => Box::new([&repr.out].into_iter()),
2695 NumericOperationIr::Clamp(repr) => Box::new([&repr.out].into_iter()),
2696 NumericOperationIr::Abs(repr) => Box::new([&repr.out].into_iter()),
2697 NumericOperationIr::Full(repr) => Box::new([&repr.out].into_iter()),
2698 NumericOperationIr::MeanDim(repr) => Box::new([&repr.out].into_iter()),
2699 NumericOperationIr::Mean(repr) => Box::new([&repr.out].into_iter()),
2700 NumericOperationIr::Sum(repr) => Box::new([&repr.out].into_iter()),
2701 NumericOperationIr::SumDim(repr) => Box::new([&repr.out].into_iter()),
2702 NumericOperationIr::Prod(repr) => Box::new([&repr.out].into_iter()),
2703 NumericOperationIr::ProdDim(repr) => Box::new([&repr.out].into_iter()),
2704 NumericOperationIr::Max(repr) => Box::new([&repr.out].into_iter()),
2705 NumericOperationIr::MaxDimWithIndices(repr) => {
2706 Box::new([&repr.out, &repr.out_indices].into_iter())
2707 }
2708 NumericOperationIr::TopKWithIndices(repr) => {
2709 Box::new([&repr.out, &repr.out_indices].into_iter())
2710 }
2711 NumericOperationIr::MinDimWithIndices(repr) => {
2712 Box::new([&repr.out, &repr.out_indices].into_iter())
2713 }
2714 NumericOperationIr::Min(repr) => Box::new([&repr.out].into_iter()),
2715 NumericOperationIr::MaxDim(repr) => Box::new([&repr.out].into_iter()),
2716 NumericOperationIr::MinDim(repr) => Box::new([&repr.out].into_iter()),
2717 NumericOperationIr::MaxAbs(repr) => Box::new([&repr.out].into_iter()),
2718 NumericOperationIr::MaxAbsDim(repr) => Box::new([&repr.out].into_iter()),
2719 NumericOperationIr::IntRandom(repr) => Box::new([&repr.out].into_iter()),
2720 NumericOperationIr::Powi(repr) => Box::new([&repr.out].into_iter()),
2721 NumericOperationIr::PowiScalar(repr) => Box::new([&repr.out].into_iter()),
2722 NumericOperationIr::CumMin(repr) => Box::new([&repr.out].into_iter()),
2723 NumericOperationIr::CumMax(repr) => Box::new([&repr.out].into_iter()),
2724 NumericOperationIr::CumProd(repr) => Box::new([&repr.out].into_iter()),
2725 NumericOperationIr::CumSum(repr) => Box::new([&repr.out].into_iter()),
2726 NumericOperationIr::Neg(repr) => Box::new([&repr.out].into_iter()),
2727 NumericOperationIr::Sign(repr) => Box::new([&repr.out].into_iter()),
2728 NumericOperationIr::ClampMin(repr) => Box::new([&repr.out].into_iter()),
2729 NumericOperationIr::ClampMax(repr) => Box::new([&repr.out].into_iter()),
2730 NumericOperationIr::Sort(repr) => Box::new([&repr.out].into_iter()),
2731 NumericOperationIr::SortWithIndices(repr) => {
2732 Box::new([&repr.out, &repr.out_indices].into_iter())
2733 }
2734 NumericOperationIr::ArgSort(repr) => Box::new([&repr.out].into_iter()),
2735 }
2736 }
2737 fn mark_read_only(&mut self, nodes: &[TensorId]) -> Vec<TensorIr> {
2738 let mut output = Vec::new();
2739
2740 match self {
2741 NumericOperationIr::Add(repr) => {
2742 repr.lhs.mark_read_only(nodes, &mut output);
2743 repr.rhs.mark_read_only(nodes, &mut output);
2744 }
2745 NumericOperationIr::AddScalar(repr) => {
2746 repr.lhs.mark_read_only(nodes, &mut output);
2747 }
2748 NumericOperationIr::Sub(repr) => {
2749 repr.lhs.mark_read_only(nodes, &mut output);
2750 repr.rhs.mark_read_only(nodes, &mut output);
2751 }
2752 NumericOperationIr::SubScalar(repr) => {
2753 repr.lhs.mark_read_only(nodes, &mut output);
2754 }
2755 NumericOperationIr::Mul(repr) => {
2756 repr.lhs.mark_read_only(nodes, &mut output);
2757 repr.rhs.mark_read_only(nodes, &mut output);
2758 }
2759 NumericOperationIr::MulScalar(repr) => {
2760 repr.lhs.mark_read_only(nodes, &mut output);
2761 }
2762 NumericOperationIr::Div(repr) => {
2763 repr.lhs.mark_read_only(nodes, &mut output);
2764 repr.rhs.mark_read_only(nodes, &mut output);
2765 }
2766 NumericOperationIr::DivScalar(repr) => {
2767 repr.lhs.mark_read_only(nodes, &mut output);
2768 }
2769 NumericOperationIr::Rem(repr) => {
2770 repr.lhs.mark_read_only(nodes, &mut output);
2771 repr.rhs.mark_read_only(nodes, &mut output);
2772 }
2773 NumericOperationIr::RemScalar(repr) => {
2774 repr.lhs.mark_read_only(nodes, &mut output);
2775 }
2776 NumericOperationIr::GreaterElem(repr) => {
2777 repr.lhs.mark_read_only(nodes, &mut output);
2778 }
2779 NumericOperationIr::GreaterEqualElem(repr) => {
2780 repr.lhs.mark_read_only(nodes, &mut output);
2781 }
2782 NumericOperationIr::LowerElem(repr) => {
2783 repr.lhs.mark_read_only(nodes, &mut output);
2784 }
2785 NumericOperationIr::LowerEqualElem(repr) => {
2786 repr.lhs.mark_read_only(nodes, &mut output);
2787 }
2788 NumericOperationIr::Greater(repr) => {
2789 repr.lhs.mark_read_only(nodes, &mut output);
2790 repr.rhs.mark_read_only(nodes, &mut output);
2791 }
2792 NumericOperationIr::GreaterEqual(repr) => {
2793 repr.lhs.mark_read_only(nodes, &mut output);
2794 repr.rhs.mark_read_only(nodes, &mut output);
2795 }
2796 NumericOperationIr::Lower(repr) => {
2797 repr.lhs.mark_read_only(nodes, &mut output);
2798 repr.rhs.mark_read_only(nodes, &mut output);
2799 }
2800 NumericOperationIr::LowerEqual(repr) => {
2801 repr.lhs.mark_read_only(nodes, &mut output);
2802 repr.rhs.mark_read_only(nodes, &mut output);
2803 }
2804 NumericOperationIr::ArgMax(repr) => {
2805 repr.input.mark_read_only(nodes, &mut output);
2806 }
2807 NumericOperationIr::ArgTopK(repr) => {
2808 repr.input.mark_read_only(nodes, &mut output);
2809 }
2810 NumericOperationIr::TopK(repr) => {
2811 repr.input.mark_read_only(nodes, &mut output);
2812 }
2813 NumericOperationIr::ArgMin(repr) => {
2814 repr.input.mark_read_only(nodes, &mut output);
2815 }
2816 NumericOperationIr::Clamp(repr) => {
2817 repr.tensor.mark_read_only(nodes, &mut output);
2818 }
2819 NumericOperationIr::Abs(repr) => {
2820 repr.input.mark_read_only(nodes, &mut output);
2821 }
2822 NumericOperationIr::Full(_) => {}
2823 NumericOperationIr::MeanDim(repr) => {
2824 repr.input.mark_read_only(nodes, &mut output);
2825 }
2826 NumericOperationIr::Mean(repr) => {
2827 repr.input.mark_read_only(nodes, &mut output);
2828 }
2829 NumericOperationIr::Sum(repr) => {
2830 repr.input.mark_read_only(nodes, &mut output);
2831 }
2832 NumericOperationIr::SumDim(repr) => {
2833 repr.input.mark_read_only(nodes, &mut output);
2834 }
2835 NumericOperationIr::Prod(repr) => {
2836 repr.input.mark_read_only(nodes, &mut output);
2837 }
2838 NumericOperationIr::ProdDim(repr) => {
2839 repr.input.mark_read_only(nodes, &mut output);
2840 }
2841 NumericOperationIr::Max(repr) => {
2842 repr.input.mark_read_only(nodes, &mut output);
2843 }
2844 NumericOperationIr::MaxDimWithIndices(repr) => {
2845 repr.tensor.mark_read_only(nodes, &mut output);
2846 }
2847 NumericOperationIr::TopKWithIndices(repr) => {
2848 repr.tensor.mark_read_only(nodes, &mut output);
2849 }
2850 NumericOperationIr::MinDimWithIndices(repr) => {
2851 repr.tensor.mark_read_only(nodes, &mut output);
2852 }
2853 NumericOperationIr::Min(repr) => {
2854 repr.input.mark_read_only(nodes, &mut output);
2855 }
2856 NumericOperationIr::MaxDim(repr) => {
2857 repr.input.mark_read_only(nodes, &mut output);
2858 }
2859 NumericOperationIr::MinDim(repr) => {
2860 repr.input.mark_read_only(nodes, &mut output);
2861 }
2862 NumericOperationIr::MaxAbs(repr) => {
2863 repr.input.mark_read_only(nodes, &mut output);
2864 }
2865 NumericOperationIr::MaxAbsDim(repr) => {
2866 repr.input.mark_read_only(nodes, &mut output);
2867 }
2868 NumericOperationIr::IntRandom(_) => {}
2869 NumericOperationIr::Powi(repr) => {
2870 repr.lhs.mark_read_only(nodes, &mut output);
2871 repr.rhs.mark_read_only(nodes, &mut output);
2872 }
2873 NumericOperationIr::PowiScalar(repr) => {
2874 repr.lhs.mark_read_only(nodes, &mut output);
2875 }
2876 NumericOperationIr::CumSum(repr) => {
2877 repr.input.mark_read_only(nodes, &mut output);
2878 }
2879 NumericOperationIr::CumProd(repr) => {
2880 repr.input.mark_read_only(nodes, &mut output);
2881 }
2882 NumericOperationIr::CumMin(repr) => {
2883 repr.input.mark_read_only(nodes, &mut output);
2884 }
2885 NumericOperationIr::CumMax(repr) => {
2886 repr.input.mark_read_only(nodes, &mut output);
2887 }
2888 NumericOperationIr::Neg(repr) => {
2889 repr.input.mark_read_only(nodes, &mut output);
2890 }
2891 NumericOperationIr::Sign(repr) => {
2892 repr.input.mark_read_only(nodes, &mut output);
2893 }
2894 NumericOperationIr::ClampMin(repr) => {
2895 repr.lhs.mark_read_only(nodes, &mut output);
2896 }
2897 NumericOperationIr::ClampMax(repr) => {
2898 repr.lhs.mark_read_only(nodes, &mut output);
2899 }
2900 NumericOperationIr::Sort(repr) => {
2901 repr.input.mark_read_only(nodes, &mut output);
2902 }
2903 NumericOperationIr::SortWithIndices(repr) => {
2904 repr.input.mark_read_only(nodes, &mut output);
2905 }
2906 NumericOperationIr::ArgSort(repr) => {
2907 repr.input.mark_read_only(nodes, &mut output);
2908 }
2909 };
2910
2911 output
2912 }
2913
2914 fn visit_mut(&mut self, v: &mut impl IrVisitorMut) {
2915 match self {
2916 NumericOperationIr::Add(repr) => {
2917 v.visit_tensor_mut(&mut repr.lhs);
2918 v.visit_tensor_mut(&mut repr.rhs);
2919 v.visit_tensor_mut(&mut repr.out);
2920 }
2921 NumericOperationIr::AddScalar(repr) => {
2922 v.visit_tensor_mut(&mut repr.lhs);
2923 v.visit_tensor_mut(&mut repr.out);
2924 v.visit_scalar_mut(&mut repr.rhs);
2925 }
2926 NumericOperationIr::Sub(repr) => {
2927 v.visit_tensor_mut(&mut repr.lhs);
2928 v.visit_tensor_mut(&mut repr.rhs);
2929 v.visit_tensor_mut(&mut repr.out);
2930 }
2931 NumericOperationIr::SubScalar(repr) => {
2932 v.visit_tensor_mut(&mut repr.lhs);
2933 v.visit_tensor_mut(&mut repr.out);
2934 v.visit_scalar_mut(&mut repr.rhs);
2935 }
2936 NumericOperationIr::Mul(repr) => {
2937 v.visit_tensor_mut(&mut repr.lhs);
2938 v.visit_tensor_mut(&mut repr.rhs);
2939 v.visit_tensor_mut(&mut repr.out);
2940 }
2941 NumericOperationIr::MulScalar(repr) => {
2942 v.visit_tensor_mut(&mut repr.lhs);
2943 v.visit_tensor_mut(&mut repr.out);
2944 v.visit_scalar_mut(&mut repr.rhs);
2945 }
2946 NumericOperationIr::Div(repr) => {
2947 v.visit_tensor_mut(&mut repr.lhs);
2948 v.visit_tensor_mut(&mut repr.rhs);
2949 v.visit_tensor_mut(&mut repr.out);
2950 }
2951 NumericOperationIr::DivScalar(repr) => {
2952 v.visit_tensor_mut(&mut repr.lhs);
2953 v.visit_tensor_mut(&mut repr.out);
2954 v.visit_scalar_mut(&mut repr.rhs);
2955 }
2956 NumericOperationIr::Rem(repr) => {
2957 v.visit_tensor_mut(&mut repr.lhs);
2958 v.visit_tensor_mut(&mut repr.rhs);
2959 v.visit_tensor_mut(&mut repr.out);
2960 }
2961 NumericOperationIr::RemScalar(repr) => {
2962 v.visit_tensor_mut(&mut repr.lhs);
2963 v.visit_tensor_mut(&mut repr.out);
2964 v.visit_scalar_mut(&mut repr.rhs);
2965 }
2966 NumericOperationIr::GreaterElem(repr) => {
2967 v.visit_tensor_mut(&mut repr.lhs);
2968 v.visit_tensor_mut(&mut repr.out);
2969 v.visit_scalar_mut(&mut repr.rhs);
2970 }
2971 NumericOperationIr::GreaterEqualElem(repr) => {
2972 v.visit_tensor_mut(&mut repr.lhs);
2973 v.visit_tensor_mut(&mut repr.out);
2974 v.visit_scalar_mut(&mut repr.rhs);
2975 }
2976 NumericOperationIr::LowerElem(repr) => {
2977 v.visit_tensor_mut(&mut repr.lhs);
2978 v.visit_tensor_mut(&mut repr.out);
2979 v.visit_scalar_mut(&mut repr.rhs);
2980 }
2981 NumericOperationIr::LowerEqualElem(repr) => {
2982 v.visit_tensor_mut(&mut repr.lhs);
2983 v.visit_tensor_mut(&mut repr.out);
2984 v.visit_scalar_mut(&mut repr.rhs);
2985 }
2986 NumericOperationIr::Greater(repr) => {
2987 v.visit_tensor_mut(&mut repr.lhs);
2988 v.visit_tensor_mut(&mut repr.rhs);
2989 v.visit_tensor_mut(&mut repr.out);
2990 }
2991 NumericOperationIr::GreaterEqual(repr) => {
2992 v.visit_tensor_mut(&mut repr.lhs);
2993 v.visit_tensor_mut(&mut repr.rhs);
2994 v.visit_tensor_mut(&mut repr.out);
2995 }
2996 NumericOperationIr::Lower(repr) => {
2997 v.visit_tensor_mut(&mut repr.lhs);
2998 v.visit_tensor_mut(&mut repr.rhs);
2999 v.visit_tensor_mut(&mut repr.out);
3000 }
3001 NumericOperationIr::LowerEqual(repr) => {
3002 v.visit_tensor_mut(&mut repr.lhs);
3003 v.visit_tensor_mut(&mut repr.rhs);
3004 v.visit_tensor_mut(&mut repr.out);
3005 }
3006 NumericOperationIr::ArgMax(repr) => {
3007 v.visit_tensor_mut(&mut repr.input);
3008 v.visit_tensor_mut(&mut repr.out);
3009 }
3010 NumericOperationIr::ArgTopK(repr) => {
3011 v.visit_tensor_mut(&mut repr.input);
3012 v.visit_tensor_mut(&mut repr.out);
3013 }
3014 NumericOperationIr::TopK(repr) => {
3015 v.visit_tensor_mut(&mut repr.input);
3016 v.visit_tensor_mut(&mut repr.out);
3017 }
3018 NumericOperationIr::ArgMin(repr) => {
3019 v.visit_tensor_mut(&mut repr.input);
3020 v.visit_tensor_mut(&mut repr.out);
3021 }
3022 NumericOperationIr::Clamp(repr) => {
3023 v.visit_tensor_mut(&mut repr.tensor);
3024 v.visit_tensor_mut(&mut repr.out);
3025 v.visit_scalar_mut(&mut repr.min);
3026 v.visit_scalar_mut(&mut repr.max);
3027 }
3028 NumericOperationIr::Abs(repr) => {
3029 v.visit_tensor_mut(&mut repr.input);
3030 v.visit_tensor_mut(&mut repr.out);
3031 }
3032 NumericOperationIr::Full(repr) => {
3033 v.visit_tensor_mut(&mut repr.out);
3034 v.visit_scalar_mut(&mut repr.value);
3035 }
3036 NumericOperationIr::MeanDim(repr) => {
3037 v.visit_tensor_mut(&mut repr.input);
3038 v.visit_tensor_mut(&mut repr.out);
3039 }
3040 NumericOperationIr::Mean(repr) => {
3041 v.visit_tensor_mut(&mut repr.input);
3042 v.visit_tensor_mut(&mut repr.out);
3043 }
3044 NumericOperationIr::Sum(repr) => {
3045 v.visit_tensor_mut(&mut repr.input);
3046 v.visit_tensor_mut(&mut repr.out);
3047 }
3048 NumericOperationIr::SumDim(repr) => {
3049 v.visit_tensor_mut(&mut repr.input);
3050 v.visit_tensor_mut(&mut repr.out);
3051 }
3052 NumericOperationIr::Prod(repr) => {
3053 v.visit_tensor_mut(&mut repr.input);
3054 v.visit_tensor_mut(&mut repr.out);
3055 }
3056 NumericOperationIr::ProdDim(repr) => {
3057 v.visit_tensor_mut(&mut repr.input);
3058 v.visit_tensor_mut(&mut repr.out);
3059 }
3060 NumericOperationIr::Max(repr) => {
3061 v.visit_tensor_mut(&mut repr.input);
3062 v.visit_tensor_mut(&mut repr.out);
3063 }
3064 NumericOperationIr::MaxDimWithIndices(repr) => {
3065 v.visit_tensor_mut(&mut repr.tensor);
3066 v.visit_tensor_mut(&mut repr.out);
3067 v.visit_tensor_mut(&mut repr.out_indices);
3068 }
3069 NumericOperationIr::TopKWithIndices(repr) => {
3070 v.visit_tensor_mut(&mut repr.tensor);
3071 v.visit_tensor_mut(&mut repr.out);
3072 v.visit_tensor_mut(&mut repr.out_indices);
3073 }
3074 NumericOperationIr::MinDimWithIndices(repr) => {
3075 v.visit_tensor_mut(&mut repr.tensor);
3076 v.visit_tensor_mut(&mut repr.out);
3077 v.visit_tensor_mut(&mut repr.out_indices);
3078 }
3079 NumericOperationIr::Min(repr) => {
3080 v.visit_tensor_mut(&mut repr.input);
3081 v.visit_tensor_mut(&mut repr.out);
3082 }
3083 NumericOperationIr::MaxDim(repr) => {
3084 v.visit_tensor_mut(&mut repr.input);
3085 v.visit_tensor_mut(&mut repr.out);
3086 }
3087 NumericOperationIr::MinDim(repr) => {
3088 v.visit_tensor_mut(&mut repr.input);
3089 v.visit_tensor_mut(&mut repr.out);
3090 }
3091 NumericOperationIr::MaxAbs(repr) => {
3092 v.visit_tensor_mut(&mut repr.input);
3093 v.visit_tensor_mut(&mut repr.out);
3094 }
3095 NumericOperationIr::MaxAbsDim(repr) => {
3096 v.visit_tensor_mut(&mut repr.input);
3097 v.visit_tensor_mut(&mut repr.out);
3098 }
3099 NumericOperationIr::IntRandom(repr) => {
3100 v.visit_tensor_mut(&mut repr.out);
3101 }
3102 NumericOperationIr::Powi(repr) => {
3103 v.visit_tensor_mut(&mut repr.lhs);
3104 v.visit_tensor_mut(&mut repr.rhs);
3105 v.visit_tensor_mut(&mut repr.out);
3106 }
3107 NumericOperationIr::PowiScalar(repr) => {
3108 v.visit_tensor_mut(&mut repr.lhs);
3109 v.visit_tensor_mut(&mut repr.out);
3110 v.visit_scalar_mut(&mut repr.rhs);
3111 }
3112 NumericOperationIr::CumMin(repr) => {
3113 v.visit_tensor_mut(&mut repr.input);
3114 v.visit_tensor_mut(&mut repr.out);
3115 }
3116 NumericOperationIr::CumMax(repr) => {
3117 v.visit_tensor_mut(&mut repr.input);
3118 v.visit_tensor_mut(&mut repr.out);
3119 }
3120 NumericOperationIr::CumProd(repr) => {
3121 v.visit_tensor_mut(&mut repr.input);
3122 v.visit_tensor_mut(&mut repr.out);
3123 }
3124 NumericOperationIr::CumSum(repr) => {
3125 v.visit_tensor_mut(&mut repr.input);
3126 v.visit_tensor_mut(&mut repr.out);
3127 }
3128 NumericOperationIr::Neg(repr) => {
3129 v.visit_tensor_mut(&mut repr.input);
3130 v.visit_tensor_mut(&mut repr.out);
3131 }
3132 NumericOperationIr::Sign(repr) => {
3133 v.visit_tensor_mut(&mut repr.input);
3134 v.visit_tensor_mut(&mut repr.out);
3135 }
3136 NumericOperationIr::ClampMin(repr) => {
3137 v.visit_tensor_mut(&mut repr.lhs);
3138 v.visit_tensor_mut(&mut repr.out);
3139 v.visit_scalar_mut(&mut repr.rhs);
3140 }
3141 NumericOperationIr::ClampMax(repr) => {
3142 v.visit_tensor_mut(&mut repr.lhs);
3143 v.visit_tensor_mut(&mut repr.out);
3144 v.visit_scalar_mut(&mut repr.rhs);
3145 }
3146 NumericOperationIr::Sort(repr) => {
3147 v.visit_tensor_mut(&mut repr.input);
3148 v.visit_tensor_mut(&mut repr.out);
3149 }
3150 NumericOperationIr::SortWithIndices(repr) => {
3151 v.visit_tensor_mut(&mut repr.input);
3152 v.visit_tensor_mut(&mut repr.out);
3153 v.visit_tensor_mut(&mut repr.out_indices);
3154 }
3155 NumericOperationIr::ArgSort(repr) => {
3156 v.visit_tensor_mut(&mut repr.input);
3157 v.visit_tensor_mut(&mut repr.out);
3158 }
3159 }
3160 }
3161}
3162
3163impl FloatOperationIr {
3164 fn inputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
3165 match self {
3166 FloatOperationIr::Matmul(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3167 FloatOperationIr::Cross(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3168 FloatOperationIr::Random(_repr) => Box::new([].into_iter()),
3169 FloatOperationIr::Exp(repr) => Box::new([&repr.input].into_iter()),
3170 FloatOperationIr::Log(repr) => Box::new([&repr.input].into_iter()),
3171 FloatOperationIr::Log1p(repr) => Box::new([&repr.input].into_iter()),
3172 FloatOperationIr::Erf(repr) => Box::new([&repr.input].into_iter()),
3173 FloatOperationIr::Recip(repr) => Box::new([&repr.input].into_iter()),
3174 FloatOperationIr::PowfScalar(repr) => Box::new([&repr.lhs].into_iter()),
3175 FloatOperationIr::Sqrt(repr) => Box::new([&repr.input].into_iter()),
3176 FloatOperationIr::Cos(repr) => Box::new([&repr.input].into_iter()),
3177 FloatOperationIr::Sin(repr) => Box::new([&repr.input].into_iter()),
3178 FloatOperationIr::Tanh(repr) => Box::new([&repr.input].into_iter()),
3179 FloatOperationIr::Round(repr) => Box::new([&repr.input].into_iter()),
3180 FloatOperationIr::Floor(repr) => Box::new([&repr.input].into_iter()),
3181 FloatOperationIr::Ceil(repr) => Box::new([&repr.input].into_iter()),
3182 FloatOperationIr::Trunc(repr) => Box::new([&repr.input].into_iter()),
3183 FloatOperationIr::IntoInt(repr) => Box::new([&repr.input].into_iter()),
3184 FloatOperationIr::Quantize(repr) => Box::new(
3185 [&repr.tensor, &repr.qparams.scales]
3186 .into_iter()
3187 .chain(repr.qparams.global.iter()),
3188 ),
3189 FloatOperationIr::Dequantize(repr) => Box::new([&repr.input].into_iter()),
3190 FloatOperationIr::IsNan(repr) => Box::new([&repr.input].into_iter()),
3191 FloatOperationIr::IsInf(repr) => Box::new([&repr.input].into_iter()),
3192 FloatOperationIr::GridSample2d(repr) => {
3193 Box::new([&repr.tensor, &repr.grid].into_iter())
3194 }
3195 FloatOperationIr::Tan(repr) => Box::new([&repr.input].into_iter()),
3196 FloatOperationIr::Cosh(repr) => Box::new([&repr.input].into_iter()),
3197 FloatOperationIr::Sinh(repr) => Box::new([&repr.input].into_iter()),
3198 FloatOperationIr::ArcCos(repr) => Box::new([&repr.input].into_iter()),
3199 FloatOperationIr::ArcCosh(repr) => Box::new([&repr.input].into_iter()),
3200 FloatOperationIr::ArcSin(repr) => Box::new([&repr.input].into_iter()),
3201 FloatOperationIr::ArcSinh(repr) => Box::new([&repr.input].into_iter()),
3202 FloatOperationIr::ArcTan(repr) => Box::new([&repr.input].into_iter()),
3203 FloatOperationIr::ArcTanh(repr) => Box::new([&repr.input].into_iter()),
3204 FloatOperationIr::ArcTan2(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3205 FloatOperationIr::Powf(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3206 FloatOperationIr::Hypot(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3207 }
3208 }
3209 fn outputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
3210 match self {
3211 FloatOperationIr::Matmul(repr) => Box::new([&repr.out].into_iter()),
3212 FloatOperationIr::Cross(repr) => Box::new([&repr.out].into_iter()),
3213 FloatOperationIr::Random(repr) => Box::new([&repr.out].into_iter()),
3214 FloatOperationIr::Exp(repr) => Box::new([&repr.out].into_iter()),
3215 FloatOperationIr::Log(repr) => Box::new([&repr.out].into_iter()),
3216 FloatOperationIr::Log1p(repr) => Box::new([&repr.out].into_iter()),
3217 FloatOperationIr::Erf(repr) => Box::new([&repr.out].into_iter()),
3218 FloatOperationIr::Recip(repr) => Box::new([&repr.out].into_iter()),
3219 FloatOperationIr::PowfScalar(repr) => Box::new([&repr.out].into_iter()),
3220 FloatOperationIr::Sqrt(repr) => Box::new([&repr.out].into_iter()),
3221 FloatOperationIr::Cos(repr) => Box::new([&repr.out].into_iter()),
3222 FloatOperationIr::Sin(repr) => Box::new([&repr.out].into_iter()),
3223 FloatOperationIr::Tanh(repr) => Box::new([&repr.out].into_iter()),
3224 FloatOperationIr::Round(repr) => Box::new([&repr.out].into_iter()),
3225 FloatOperationIr::Floor(repr) => Box::new([&repr.out].into_iter()),
3226 FloatOperationIr::Ceil(repr) => Box::new([&repr.out].into_iter()),
3227 FloatOperationIr::Trunc(repr) => Box::new([&repr.out].into_iter()),
3228 FloatOperationIr::IntoInt(repr) => Box::new([&repr.out].into_iter()),
3229 FloatOperationIr::Quantize(repr) => Box::new([&repr.out].into_iter()),
3230 FloatOperationIr::Dequantize(repr) => Box::new([&repr.out].into_iter()),
3231 FloatOperationIr::IsNan(repr) => Box::new([&repr.out].into_iter()),
3232 FloatOperationIr::IsInf(repr) => Box::new([&repr.out].into_iter()),
3233 FloatOperationIr::GridSample2d(repr) => Box::new([&repr.out].into_iter()),
3234 FloatOperationIr::Tan(repr) => Box::new([&repr.out].into_iter()),
3235 FloatOperationIr::Cosh(repr) => Box::new([&repr.out].into_iter()),
3236 FloatOperationIr::Sinh(repr) => Box::new([&repr.out].into_iter()),
3237 FloatOperationIr::ArcCos(repr) => Box::new([&repr.out].into_iter()),
3238 FloatOperationIr::ArcCosh(repr) => Box::new([&repr.out].into_iter()),
3239 FloatOperationIr::ArcSin(repr) => Box::new([&repr.out].into_iter()),
3240 FloatOperationIr::ArcSinh(repr) => Box::new([&repr.out].into_iter()),
3241 FloatOperationIr::ArcTan(repr) => Box::new([&repr.out].into_iter()),
3242 FloatOperationIr::ArcTanh(repr) => Box::new([&repr.out].into_iter()),
3243 FloatOperationIr::ArcTan2(repr) => Box::new([&repr.out].into_iter()),
3244 FloatOperationIr::Powf(repr) => Box::new([&repr.out].into_iter()),
3245 FloatOperationIr::Hypot(repr) => Box::new([&repr.out].into_iter()),
3246 }
3247 }
3248
3249 fn mark_read_only(&mut self, nodes: &[TensorId]) -> Vec<TensorIr> {
3250 let mut output = Vec::new();
3251
3252 match self {
3253 FloatOperationIr::Matmul(repr) => {
3254 repr.lhs.mark_read_only(nodes, &mut output);
3255 repr.rhs.mark_read_only(nodes, &mut output);
3256 }
3257 FloatOperationIr::Cross(repr) => {
3258 repr.lhs.mark_read_only(nodes, &mut output);
3259 repr.rhs.mark_read_only(nodes, &mut output);
3260 }
3261 FloatOperationIr::Random(_) => {}
3262 FloatOperationIr::Exp(repr) => {
3263 repr.input.mark_read_only(nodes, &mut output);
3264 }
3265 FloatOperationIr::Log(repr) => {
3266 repr.input.mark_read_only(nodes, &mut output);
3267 }
3268 FloatOperationIr::Log1p(repr) => {
3269 repr.input.mark_read_only(nodes, &mut output);
3270 }
3271 FloatOperationIr::Erf(repr) => {
3272 repr.input.mark_read_only(nodes, &mut output);
3273 }
3274 FloatOperationIr::Recip(repr) => {
3275 repr.input.mark_read_only(nodes, &mut output);
3276 }
3277 FloatOperationIr::PowfScalar(repr) => {
3278 repr.lhs.mark_read_only(nodes, &mut output);
3279 }
3280 FloatOperationIr::Sqrt(repr) => {
3281 repr.input.mark_read_only(nodes, &mut output);
3282 }
3283 FloatOperationIr::Cos(repr) => {
3284 repr.input.mark_read_only(nodes, &mut output);
3285 }
3286 FloatOperationIr::Sin(repr) => {
3287 repr.input.mark_read_only(nodes, &mut output);
3288 }
3289 FloatOperationIr::Tanh(repr) => {
3290 repr.input.mark_read_only(nodes, &mut output);
3291 }
3292 FloatOperationIr::Round(repr) => {
3293 repr.input.mark_read_only(nodes, &mut output);
3294 }
3295 FloatOperationIr::Floor(repr) => {
3296 repr.input.mark_read_only(nodes, &mut output);
3297 }
3298 FloatOperationIr::Ceil(repr) => {
3299 repr.input.mark_read_only(nodes, &mut output);
3300 }
3301 FloatOperationIr::Trunc(repr) => {
3302 repr.input.mark_read_only(nodes, &mut output);
3303 }
3304 FloatOperationIr::Quantize(repr) => {
3305 repr.tensor.mark_read_only(nodes, &mut output);
3306 repr.qparams.scales.mark_read_only(nodes, &mut output);
3307 if let Some(global) = &mut repr.qparams.global {
3308 global.mark_read_only(nodes, &mut output);
3309 }
3310 }
3311 FloatOperationIr::Dequantize(repr) => {
3312 repr.input.mark_read_only(nodes, &mut output);
3313 }
3314 FloatOperationIr::IntoInt(repr) => {
3315 repr.input.mark_read_only(nodes, &mut output);
3316 }
3317 FloatOperationIr::IsNan(repr) => {
3318 repr.input.mark_read_only(nodes, &mut output);
3319 }
3320 FloatOperationIr::IsInf(repr) => {
3321 repr.input.mark_read_only(nodes, &mut output);
3322 }
3323 FloatOperationIr::GridSample2d(repr) => {
3324 repr.tensor.mark_read_only(nodes, &mut output);
3325 repr.grid.mark_read_only(nodes, &mut output);
3326 }
3327 FloatOperationIr::Tan(repr) => repr.input.mark_read_only(nodes, &mut output),
3328 FloatOperationIr::Cosh(repr) => repr.input.mark_read_only(nodes, &mut output),
3329 FloatOperationIr::Sinh(repr) => repr.input.mark_read_only(nodes, &mut output),
3330 FloatOperationIr::ArcCos(repr) => repr.input.mark_read_only(nodes, &mut output),
3331 FloatOperationIr::ArcCosh(repr) => repr.input.mark_read_only(nodes, &mut output),
3332 FloatOperationIr::ArcSin(repr) => repr.input.mark_read_only(nodes, &mut output),
3333 FloatOperationIr::ArcSinh(repr) => repr.input.mark_read_only(nodes, &mut output),
3334 FloatOperationIr::ArcTan(repr) => repr.input.mark_read_only(nodes, &mut output),
3335 FloatOperationIr::ArcTanh(repr) => repr.input.mark_read_only(nodes, &mut output),
3336 FloatOperationIr::ArcTan2(repr) => {
3337 repr.lhs.mark_read_only(nodes, &mut output);
3338 repr.rhs.mark_read_only(nodes, &mut output);
3339 }
3340 FloatOperationIr::Powf(repr) => {
3341 repr.lhs.mark_read_only(nodes, &mut output);
3342 repr.rhs.mark_read_only(nodes, &mut output);
3343 }
3344 FloatOperationIr::Hypot(repr) => {
3345 repr.lhs.mark_read_only(nodes, &mut output);
3346 repr.rhs.mark_read_only(nodes, &mut output);
3347 }
3348 };
3349
3350 output
3351 }
3352
3353 fn visit_mut(&mut self, v: &mut impl IrVisitorMut) {
3354 match self {
3355 FloatOperationIr::Matmul(repr) => {
3356 v.visit_tensor_mut(&mut repr.lhs);
3357 v.visit_tensor_mut(&mut repr.rhs);
3358 v.visit_tensor_mut(&mut repr.out);
3359 }
3360 FloatOperationIr::Cross(repr) => {
3361 v.visit_tensor_mut(&mut repr.lhs);
3362 v.visit_tensor_mut(&mut repr.rhs);
3363 v.visit_tensor_mut(&mut repr.out);
3364 }
3365 FloatOperationIr::Random(repr) => {
3366 v.visit_tensor_mut(&mut repr.out);
3367 }
3368 FloatOperationIr::Exp(repr) => {
3369 v.visit_tensor_mut(&mut repr.input);
3370 v.visit_tensor_mut(&mut repr.out);
3371 }
3372 FloatOperationIr::Log(repr) => {
3373 v.visit_tensor_mut(&mut repr.input);
3374 v.visit_tensor_mut(&mut repr.out);
3375 }
3376 FloatOperationIr::Log1p(repr) => {
3377 v.visit_tensor_mut(&mut repr.input);
3378 v.visit_tensor_mut(&mut repr.out);
3379 }
3380 FloatOperationIr::Erf(repr) => {
3381 v.visit_tensor_mut(&mut repr.input);
3382 v.visit_tensor_mut(&mut repr.out);
3383 }
3384 FloatOperationIr::Recip(repr) => {
3385 v.visit_tensor_mut(&mut repr.input);
3386 v.visit_tensor_mut(&mut repr.out);
3387 }
3388 FloatOperationIr::PowfScalar(repr) => {
3389 v.visit_tensor_mut(&mut repr.lhs);
3390 v.visit_tensor_mut(&mut repr.out);
3391 v.visit_scalar_mut(&mut repr.rhs);
3392 }
3393 FloatOperationIr::Sqrt(repr) => {
3394 v.visit_tensor_mut(&mut repr.input);
3395 v.visit_tensor_mut(&mut repr.out);
3396 }
3397 FloatOperationIr::Cos(repr) => {
3398 v.visit_tensor_mut(&mut repr.input);
3399 v.visit_tensor_mut(&mut repr.out);
3400 }
3401 FloatOperationIr::Sin(repr) => {
3402 v.visit_tensor_mut(&mut repr.input);
3403 v.visit_tensor_mut(&mut repr.out);
3404 }
3405 FloatOperationIr::Tanh(repr) => {
3406 v.visit_tensor_mut(&mut repr.input);
3407 v.visit_tensor_mut(&mut repr.out);
3408 }
3409 FloatOperationIr::Round(repr) => {
3410 v.visit_tensor_mut(&mut repr.input);
3411 v.visit_tensor_mut(&mut repr.out);
3412 }
3413 FloatOperationIr::Floor(repr) => {
3414 v.visit_tensor_mut(&mut repr.input);
3415 v.visit_tensor_mut(&mut repr.out);
3416 }
3417 FloatOperationIr::Ceil(repr) => {
3418 v.visit_tensor_mut(&mut repr.input);
3419 v.visit_tensor_mut(&mut repr.out);
3420 }
3421 FloatOperationIr::Trunc(repr) => {
3422 v.visit_tensor_mut(&mut repr.input);
3423 v.visit_tensor_mut(&mut repr.out);
3424 }
3425 FloatOperationIr::IntoInt(repr) => {
3426 v.visit_tensor_mut(&mut repr.input);
3427 v.visit_tensor_mut(&mut repr.out);
3428 }
3429 FloatOperationIr::Quantize(repr) => {
3430 v.visit_tensor_mut(&mut repr.tensor);
3431 v.visit_tensor_mut(&mut repr.qparams.scales);
3432 if let Some(global) = &mut repr.qparams.global {
3433 v.visit_tensor_mut(global);
3434 }
3435 v.visit_tensor_mut(&mut repr.out);
3436 }
3437 FloatOperationIr::Dequantize(repr) => {
3438 v.visit_tensor_mut(&mut repr.input);
3439 v.visit_tensor_mut(&mut repr.out);
3440 }
3441 FloatOperationIr::IsNan(repr) => {
3442 v.visit_tensor_mut(&mut repr.input);
3443 v.visit_tensor_mut(&mut repr.out);
3444 }
3445 FloatOperationIr::IsInf(repr) => {
3446 v.visit_tensor_mut(&mut repr.input);
3447 v.visit_tensor_mut(&mut repr.out);
3448 }
3449 FloatOperationIr::GridSample2d(repr) => {
3450 v.visit_tensor_mut(&mut repr.tensor);
3451 v.visit_tensor_mut(&mut repr.grid);
3452 v.visit_tensor_mut(&mut repr.out);
3453 }
3454 FloatOperationIr::Tan(repr) => {
3455 v.visit_tensor_mut(&mut repr.input);
3456 v.visit_tensor_mut(&mut repr.out);
3457 }
3458 FloatOperationIr::Cosh(repr) => {
3459 v.visit_tensor_mut(&mut repr.input);
3460 v.visit_tensor_mut(&mut repr.out);
3461 }
3462 FloatOperationIr::Sinh(repr) => {
3463 v.visit_tensor_mut(&mut repr.input);
3464 v.visit_tensor_mut(&mut repr.out);
3465 }
3466 FloatOperationIr::ArcCos(repr) => {
3467 v.visit_tensor_mut(&mut repr.input);
3468 v.visit_tensor_mut(&mut repr.out);
3469 }
3470 FloatOperationIr::ArcCosh(repr) => {
3471 v.visit_tensor_mut(&mut repr.input);
3472 v.visit_tensor_mut(&mut repr.out);
3473 }
3474 FloatOperationIr::ArcSin(repr) => {
3475 v.visit_tensor_mut(&mut repr.input);
3476 v.visit_tensor_mut(&mut repr.out);
3477 }
3478 FloatOperationIr::ArcSinh(repr) => {
3479 v.visit_tensor_mut(&mut repr.input);
3480 v.visit_tensor_mut(&mut repr.out);
3481 }
3482 FloatOperationIr::ArcTan(repr) => {
3483 v.visit_tensor_mut(&mut repr.input);
3484 v.visit_tensor_mut(&mut repr.out);
3485 }
3486 FloatOperationIr::ArcTanh(repr) => {
3487 v.visit_tensor_mut(&mut repr.input);
3488 v.visit_tensor_mut(&mut repr.out);
3489 }
3490 FloatOperationIr::ArcTan2(repr) => {
3491 v.visit_tensor_mut(&mut repr.lhs);
3492 v.visit_tensor_mut(&mut repr.rhs);
3493 v.visit_tensor_mut(&mut repr.out);
3494 }
3495 FloatOperationIr::Powf(repr) => {
3496 v.visit_tensor_mut(&mut repr.lhs);
3497 v.visit_tensor_mut(&mut repr.rhs);
3498 v.visit_tensor_mut(&mut repr.out);
3499 }
3500 FloatOperationIr::Hypot(repr) => {
3501 v.visit_tensor_mut(&mut repr.lhs);
3502 v.visit_tensor_mut(&mut repr.rhs);
3503 v.visit_tensor_mut(&mut repr.out);
3504 }
3505 }
3506 }
3507}
3508
3509impl IntOperationIr {
3510 fn inputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
3511 match self {
3512 IntOperationIr::Matmul(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3513 IntOperationIr::IntoFloat(repr) => Box::new([&repr.input].into_iter()),
3514 IntOperationIr::BitwiseAnd(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3515 IntOperationIr::BitwiseAndScalar(repr) => Box::new([&repr.lhs].into_iter()),
3516 IntOperationIr::BitwiseOr(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3517 IntOperationIr::BitwiseOrScalar(repr) => Box::new([&repr.lhs].into_iter()),
3518 IntOperationIr::BitwiseXor(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3519 IntOperationIr::BitwiseXorScalar(repr) => Box::new([&repr.lhs].into_iter()),
3520 IntOperationIr::BitwiseNot(repr) => Box::new([&repr.input].into_iter()),
3521 IntOperationIr::BitwiseLeftShift(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3522 IntOperationIr::BitwiseLeftShiftScalar(repr) => Box::new([&repr.lhs].into_iter()),
3523 IntOperationIr::BitwiseRightShift(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3524 IntOperationIr::BitwiseRightShiftScalar(repr) => Box::new([&repr.lhs].into_iter()),
3525 }
3526 }
3527
3528 fn outputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
3529 match self {
3530 IntOperationIr::Matmul(repr) => Box::new([&repr.out].into_iter()),
3531 IntOperationIr::IntoFloat(repr) => Box::new([&repr.out].into_iter()),
3532 IntOperationIr::BitwiseAnd(repr) => Box::new([&repr.out].into_iter()),
3533 IntOperationIr::BitwiseAndScalar(repr) => Box::new([&repr.out].into_iter()),
3534 IntOperationIr::BitwiseOr(repr) => Box::new([&repr.out].into_iter()),
3535 IntOperationIr::BitwiseOrScalar(repr) => Box::new([&repr.out].into_iter()),
3536 IntOperationIr::BitwiseXor(repr) => Box::new([&repr.out].into_iter()),
3537 IntOperationIr::BitwiseXorScalar(repr) => Box::new([&repr.out].into_iter()),
3538 IntOperationIr::BitwiseNot(repr) => Box::new([&repr.out].into_iter()),
3539 IntOperationIr::BitwiseLeftShift(repr) => Box::new([&repr.out].into_iter()),
3540 IntOperationIr::BitwiseLeftShiftScalar(repr) => Box::new([&repr.out].into_iter()),
3541 IntOperationIr::BitwiseRightShift(repr) => Box::new([&repr.out].into_iter()),
3542 IntOperationIr::BitwiseRightShiftScalar(repr) => Box::new([&repr.out].into_iter()),
3543 }
3544 }
3545
3546 fn mark_read_only(&mut self, nodes: &[TensorId]) -> Vec<TensorIr> {
3547 let mut output = Vec::new();
3548
3549 match self {
3550 IntOperationIr::Matmul(repr) => {
3551 repr.lhs.mark_read_only(nodes, &mut output);
3552 repr.rhs.mark_read_only(nodes, &mut output);
3553 }
3554 IntOperationIr::IntoFloat(repr) => {
3555 repr.input.mark_read_only(nodes, &mut output);
3556 }
3557 IntOperationIr::BitwiseAnd(repr) => {
3558 repr.lhs.mark_read_only(nodes, &mut output);
3559 repr.rhs.mark_read_only(nodes, &mut output);
3560 }
3561 IntOperationIr::BitwiseAndScalar(repr) => {
3562 repr.lhs.mark_read_only(nodes, &mut output);
3563 }
3564 IntOperationIr::BitwiseOr(repr) => {
3565 repr.lhs.mark_read_only(nodes, &mut output);
3566 repr.rhs.mark_read_only(nodes, &mut output);
3567 }
3568 IntOperationIr::BitwiseOrScalar(repr) => {
3569 repr.lhs.mark_read_only(nodes, &mut output);
3570 }
3571 IntOperationIr::BitwiseXor(repr) => {
3572 repr.lhs.mark_read_only(nodes, &mut output);
3573 repr.rhs.mark_read_only(nodes, &mut output);
3574 }
3575 IntOperationIr::BitwiseXorScalar(repr) => {
3576 repr.lhs.mark_read_only(nodes, &mut output);
3577 }
3578 IntOperationIr::BitwiseNot(repr) => {
3579 repr.input.mark_read_only(nodes, &mut output);
3580 }
3581 IntOperationIr::BitwiseLeftShift(repr) => {
3582 repr.lhs.mark_read_only(nodes, &mut output);
3583 repr.rhs.mark_read_only(nodes, &mut output);
3584 }
3585 IntOperationIr::BitwiseLeftShiftScalar(repr) => {
3586 repr.lhs.mark_read_only(nodes, &mut output);
3587 }
3588 IntOperationIr::BitwiseRightShift(repr) => {
3589 repr.lhs.mark_read_only(nodes, &mut output);
3590 repr.rhs.mark_read_only(nodes, &mut output);
3591 }
3592 IntOperationIr::BitwiseRightShiftScalar(repr) => {
3593 repr.lhs.mark_read_only(nodes, &mut output);
3594 }
3595 };
3596
3597 output
3598 }
3599
3600 fn visit_mut(&mut self, v: &mut impl IrVisitorMut) {
3601 match self {
3602 IntOperationIr::Matmul(repr) => {
3603 v.visit_tensor_mut(&mut repr.lhs);
3604 v.visit_tensor_mut(&mut repr.rhs);
3605 v.visit_tensor_mut(&mut repr.out);
3606 }
3607 IntOperationIr::IntoFloat(repr) => {
3608 v.visit_tensor_mut(&mut repr.input);
3609 v.visit_tensor_mut(&mut repr.out);
3610 }
3611 IntOperationIr::BitwiseAnd(repr) => {
3612 v.visit_tensor_mut(&mut repr.lhs);
3613 v.visit_tensor_mut(&mut repr.rhs);
3614 v.visit_tensor_mut(&mut repr.out);
3615 }
3616 IntOperationIr::BitwiseAndScalar(repr) => {
3617 v.visit_tensor_mut(&mut repr.lhs);
3618 v.visit_tensor_mut(&mut repr.out);
3619 v.visit_scalar_mut(&mut repr.rhs);
3620 }
3621 IntOperationIr::BitwiseOr(repr) => {
3622 v.visit_tensor_mut(&mut repr.lhs);
3623 v.visit_tensor_mut(&mut repr.rhs);
3624 v.visit_tensor_mut(&mut repr.out);
3625 }
3626 IntOperationIr::BitwiseOrScalar(repr) => {
3627 v.visit_tensor_mut(&mut repr.lhs);
3628 v.visit_tensor_mut(&mut repr.out);
3629 v.visit_scalar_mut(&mut repr.rhs);
3630 }
3631 IntOperationIr::BitwiseXor(repr) => {
3632 v.visit_tensor_mut(&mut repr.lhs);
3633 v.visit_tensor_mut(&mut repr.rhs);
3634 v.visit_tensor_mut(&mut repr.out);
3635 }
3636 IntOperationIr::BitwiseXorScalar(repr) => {
3637 v.visit_tensor_mut(&mut repr.lhs);
3638 v.visit_tensor_mut(&mut repr.out);
3639 v.visit_scalar_mut(&mut repr.rhs);
3640 }
3641 IntOperationIr::BitwiseNot(repr) => {
3642 v.visit_tensor_mut(&mut repr.input);
3643 v.visit_tensor_mut(&mut repr.out);
3644 }
3645 IntOperationIr::BitwiseLeftShift(repr) => {
3646 v.visit_tensor_mut(&mut repr.lhs);
3647 v.visit_tensor_mut(&mut repr.rhs);
3648 v.visit_tensor_mut(&mut repr.out);
3649 }
3650 IntOperationIr::BitwiseLeftShiftScalar(repr) => {
3651 v.visit_tensor_mut(&mut repr.lhs);
3652 v.visit_tensor_mut(&mut repr.out);
3653 v.visit_scalar_mut(&mut repr.rhs);
3654 }
3655 IntOperationIr::BitwiseRightShift(repr) => {
3656 v.visit_tensor_mut(&mut repr.lhs);
3657 v.visit_tensor_mut(&mut repr.rhs);
3658 v.visit_tensor_mut(&mut repr.out);
3659 }
3660 IntOperationIr::BitwiseRightShiftScalar(repr) => {
3661 v.visit_tensor_mut(&mut repr.lhs);
3662 v.visit_tensor_mut(&mut repr.out);
3663 v.visit_scalar_mut(&mut repr.rhs);
3664 }
3665 }
3666 }
3667}
3668
3669impl BoolOperationIr {
3670 fn inputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
3671 match self {
3672 BoolOperationIr::IntoFloat(repr) => Box::new([&repr.input].into_iter()),
3673 BoolOperationIr::IntoInt(repr) => Box::new([&repr.input].into_iter()),
3674 BoolOperationIr::Not(repr) => Box::new([&repr.input].into_iter()),
3675 BoolOperationIr::And(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3676 BoolOperationIr::Or(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3677 BoolOperationIr::Xor(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3678 }
3679 }
3680 fn outputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
3681 match self {
3682 BoolOperationIr::IntoFloat(repr) => Box::new([&repr.out].into_iter()),
3683 BoolOperationIr::IntoInt(repr) => Box::new([&repr.out].into_iter()),
3684 BoolOperationIr::Not(repr) => Box::new([&repr.out].into_iter()),
3685 BoolOperationIr::And(repr) => Box::new([&repr.out].into_iter()),
3686 BoolOperationIr::Or(repr) => Box::new([&repr.out].into_iter()),
3687 BoolOperationIr::Xor(repr) => Box::new([&repr.out].into_iter()),
3688 }
3689 }
3690 fn mark_read_only(&mut self, nodes: &[TensorId]) -> Vec<TensorIr> {
3691 let mut output = Vec::new();
3692
3693 match self {
3694 BoolOperationIr::IntoFloat(repr) => {
3695 repr.input.mark_read_only(nodes, &mut output);
3696 }
3697 BoolOperationIr::IntoInt(repr) => {
3698 repr.input.mark_read_only(nodes, &mut output);
3699 }
3700 BoolOperationIr::Not(repr) => {
3701 repr.input.mark_read_only(nodes, &mut output);
3702 }
3703 BoolOperationIr::And(repr) => {
3704 repr.lhs.mark_read_only(nodes, &mut output);
3705 repr.rhs.mark_read_only(nodes, &mut output);
3706 }
3707 BoolOperationIr::Or(repr) => {
3708 repr.lhs.mark_read_only(nodes, &mut output);
3709 repr.rhs.mark_read_only(nodes, &mut output);
3710 }
3711 BoolOperationIr::Xor(repr) => {
3712 repr.lhs.mark_read_only(nodes, &mut output);
3713 repr.rhs.mark_read_only(nodes, &mut output);
3714 }
3715 };
3716
3717 output
3718 }
3719
3720 fn visit_mut(&mut self, v: &mut impl IrVisitorMut) {
3721 match self {
3722 BoolOperationIr::IntoFloat(repr) => {
3723 v.visit_tensor_mut(&mut repr.input);
3724 v.visit_tensor_mut(&mut repr.out);
3725 }
3726 BoolOperationIr::IntoInt(repr) => {
3727 v.visit_tensor_mut(&mut repr.input);
3728 v.visit_tensor_mut(&mut repr.out);
3729 }
3730 BoolOperationIr::Not(repr) => {
3731 v.visit_tensor_mut(&mut repr.input);
3732 v.visit_tensor_mut(&mut repr.out);
3733 }
3734 BoolOperationIr::And(repr) => {
3735 v.visit_tensor_mut(&mut repr.lhs);
3736 v.visit_tensor_mut(&mut repr.rhs);
3737 v.visit_tensor_mut(&mut repr.out);
3738 }
3739 BoolOperationIr::Or(repr) => {
3740 v.visit_tensor_mut(&mut repr.lhs);
3741 v.visit_tensor_mut(&mut repr.rhs);
3742 v.visit_tensor_mut(&mut repr.out);
3743 }
3744 BoolOperationIr::Xor(repr) => {
3745 v.visit_tensor_mut(&mut repr.lhs);
3746 v.visit_tensor_mut(&mut repr.rhs);
3747 v.visit_tensor_mut(&mut repr.out);
3748 }
3749 }
3750 }
3751}
3752
3753impl ModuleOperationIr {
3754 fn inputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
3755 match self {
3756 ModuleOperationIr::BatchNorm(repr) => {
3757 Box::new([&repr.x, &repr.gamma, &repr.beta, &repr.mean, &repr.variance].into_iter())
3758 }
3759 ModuleOperationIr::Embedding(repr) => {
3760 Box::new([&repr.weights, &repr.indices].into_iter())
3761 }
3762 ModuleOperationIr::EmbeddingBackward(repr) => {
3763 Box::new([&repr.weights, &repr.out_grad, &repr.indices].into_iter())
3764 }
3765 ModuleOperationIr::Linear(repr) => {
3766 if let Some(bias) = &repr.bias {
3767 Box::new([&repr.x, &repr.weight, bias].into_iter())
3768 } else {
3769 Box::new([&repr.x, &repr.weight].into_iter())
3770 }
3771 }
3772 ModuleOperationIr::LinearXBackward(repr) => {
3773 Box::new([&repr.weight, &repr.output_grad].into_iter())
3774 }
3775 ModuleOperationIr::LinearWeightBackward(repr) => {
3776 Box::new([&repr.x, &repr.output_grad].into_iter())
3777 }
3778 ModuleOperationIr::LinearBiasBackward(repr) => {
3779 Box::new([&repr.output_grad].into_iter())
3780 }
3781 ModuleOperationIr::Conv1d(repr) => {
3782 if let Some(bias) = &repr.bias {
3783 Box::new([&repr.x, &repr.weight, bias].into_iter())
3784 } else {
3785 Box::new([&repr.x, &repr.weight].into_iter())
3786 }
3787 }
3788 ModuleOperationIr::Conv1dXBackward(repr) => {
3789 Box::new([&repr.x, &repr.weight, &repr.output_grad].into_iter())
3790 }
3791 ModuleOperationIr::Conv1dWeightBackward(repr) => {
3792 Box::new([&repr.x, &repr.weight, &repr.output_grad].into_iter())
3793 }
3794 ModuleOperationIr::Conv1dBiasBackward(repr) => {
3795 Box::new([&repr.x, &repr.bias, &repr.output_grad].into_iter())
3796 }
3797 ModuleOperationIr::Conv2d(repr) => {
3798 if let Some(bias) = &repr.bias {
3799 Box::new([&repr.x, &repr.weight, bias].into_iter())
3800 } else {
3801 Box::new([&repr.x, &repr.weight].into_iter())
3802 }
3803 }
3804 ModuleOperationIr::Conv2dXBackward(repr) => {
3805 Box::new([&repr.x, &repr.weight, &repr.output_grad].into_iter())
3806 }
3807 ModuleOperationIr::Conv2dWeightBackward(repr) => {
3808 Box::new([&repr.x, &repr.weight, &repr.output_grad].into_iter())
3809 }
3810 ModuleOperationIr::Conv2dBiasBackward(repr) => {
3811 Box::new([&repr.x, &repr.bias, &repr.output_grad].into_iter())
3812 }
3813 ModuleOperationIr::Conv3d(repr) => {
3814 if let Some(bias) = &repr.bias {
3815 Box::new([&repr.x, &repr.weight, bias].into_iter())
3816 } else {
3817 Box::new([&repr.x, &repr.weight].into_iter())
3818 }
3819 }
3820 ModuleOperationIr::Conv3dXBackward(repr) => {
3821 Box::new([&repr.x, &repr.weight, &repr.output_grad].into_iter())
3822 }
3823 ModuleOperationIr::Conv3dWeightBackward(repr) => {
3824 Box::new([&repr.x, &repr.weight, &repr.output_grad].into_iter())
3825 }
3826 ModuleOperationIr::Conv3dBiasBackward(repr) => {
3827 Box::new([&repr.x, &repr.bias, &repr.output_grad].into_iter())
3828 }
3829 ModuleOperationIr::DeformableConv2d(repr) => match (&repr.mask, &repr.bias) {
3830 (Some(mask), Some(bias)) => {
3831 Box::new([&repr.x, &repr.offset, &repr.weight, mask, bias].into_iter())
3832 }
3833 (Some(mask), None) => {
3834 Box::new([&repr.x, &repr.offset, &repr.weight, mask].into_iter())
3835 }
3836 (None, Some(bias)) => {
3837 Box::new([&repr.x, &repr.offset, &repr.weight, bias].into_iter())
3838 }
3839 (None, None) => Box::new([&repr.x, &repr.offset, &repr.weight].into_iter()),
3840 },
3841 ModuleOperationIr::DeformableConv2dBackward(repr) => match (&repr.mask, &repr.bias) {
3842 (Some(mask), Some(bias)) => Box::new(
3843 [
3844 &repr.x,
3845 &repr.offset,
3846 &repr.weight,
3847 &repr.out_grad,
3848 mask,
3849 bias,
3850 ]
3851 .into_iter(),
3852 ),
3853 (Some(mask), None) => Box::new(
3854 [&repr.x, &repr.offset, &repr.weight, &repr.out_grad, mask].into_iter(),
3855 ),
3856 (None, Some(bias)) => Box::new(
3857 [&repr.x, &repr.offset, &repr.weight, &repr.out_grad, bias].into_iter(),
3858 ),
3859 (None, None) => {
3860 Box::new([&repr.x, &repr.offset, &repr.weight, &repr.out_grad].into_iter())
3861 }
3862 },
3863 ModuleOperationIr::ConvTranspose1d(repr) => {
3864 if let Some(bias) = &repr.bias {
3865 Box::new([&repr.x, &repr.weight, bias].into_iter())
3866 } else {
3867 Box::new([&repr.x, &repr.weight].into_iter())
3868 }
3869 }
3870 ModuleOperationIr::ConvTranspose2d(repr) => {
3871 if let Some(bias) = &repr.bias {
3872 Box::new([&repr.x, &repr.weight, bias].into_iter())
3873 } else {
3874 Box::new([&repr.x, &repr.weight].into_iter())
3875 }
3876 }
3877 ModuleOperationIr::ConvTranspose3d(repr) => {
3878 if let Some(bias) = &repr.bias {
3879 Box::new([&repr.x, &repr.weight, bias].into_iter())
3880 } else {
3881 Box::new([&repr.x, &repr.weight].into_iter())
3882 }
3883 }
3884 ModuleOperationIr::AvgPool1d(repr) => Box::new([&repr.x].into_iter()),
3885 ModuleOperationIr::AvgPool2d(repr) => Box::new([&repr.x].into_iter()),
3886 ModuleOperationIr::AvgPool1dBackward(repr) => {
3887 Box::new([&repr.x, &repr.grad].into_iter())
3888 }
3889 ModuleOperationIr::AvgPool2dBackward(repr) => {
3890 Box::new([&repr.x, &repr.grad].into_iter())
3891 }
3892 ModuleOperationIr::AdaptiveAvgPool1d(repr) => Box::new([&repr.x].into_iter()),
3893 ModuleOperationIr::AdaptiveAvgPool2d(repr) => Box::new([&repr.x].into_iter()),
3894 ModuleOperationIr::AdaptiveAvgPool1dBackward(repr) => {
3895 Box::new([&repr.x, &repr.grad].into_iter())
3896 }
3897 ModuleOperationIr::AdaptiveAvgPool2dBackward(repr) => {
3898 Box::new([&repr.x, &repr.grad].into_iter())
3899 }
3900 ModuleOperationIr::AdaptiveAvgPool3d(repr) => Box::new([&repr.x].into_iter()),
3901 ModuleOperationIr::AdaptiveAvgPool3dBackward(repr) => {
3902 Box::new([&repr.x, &repr.grad].into_iter())
3903 }
3904 ModuleOperationIr::MaxPool1d(repr) => Box::new([&repr.x].into_iter()),
3905 ModuleOperationIr::MaxPool1dWithIndices(repr) => Box::new([&repr.x].into_iter()),
3906 ModuleOperationIr::MaxPool1dWithIndicesBackward(repr) => {
3907 Box::new([&repr.x, &repr.indices, &repr.grad].into_iter())
3908 }
3909 ModuleOperationIr::MaxPool2d(repr) => Box::new([&repr.x].into_iter()),
3910 ModuleOperationIr::MaxPool2dWithIndices(repr) => Box::new([&repr.x].into_iter()),
3911 ModuleOperationIr::MaxPool2dWithIndicesBackward(repr) => {
3912 Box::new([&repr.x, &repr.indices, &repr.grad].into_iter())
3913 }
3914 ModuleOperationIr::Interpolate(repr) => Box::new([&repr.x].into_iter()),
3915 ModuleOperationIr::InterpolateBackward(repr) => {
3916 Box::new([&repr.x, &repr.grad].into_iter())
3917 }
3918 ModuleOperationIr::Rfft(repr) => Box::new([&repr.signal].into_iter()),
3919 ModuleOperationIr::IRfft(repr) => {
3920 Box::new([&repr.input_re, &repr.input_im].into_iter())
3921 }
3922 ModuleOperationIr::Attention(repr) => {
3923 if let Some(mask) = &repr.mask {
3924 if let Some(attn_bias) = &repr.attn_bias {
3925 Box::new([&repr.query, &repr.key, &repr.value, mask, attn_bias].into_iter())
3926 } else {
3927 Box::new([&repr.query, &repr.key, &repr.value, mask].into_iter())
3928 }
3929 } else if let Some(attn_bias) = &repr.attn_bias {
3930 Box::new([&repr.query, &repr.key, &repr.value, attn_bias].into_iter())
3931 } else {
3932 Box::new([&repr.query, &repr.key, &repr.value].into_iter())
3933 }
3934 }
3935 ModuleOperationIr::CtcLoss(repr) => Box::new(
3936 [
3937 &repr.log_probs,
3938 &repr.targets,
3939 &repr.input_lengths,
3940 &repr.target_lengths,
3941 ]
3942 .into_iter(),
3943 ),
3944 ModuleOperationIr::CtcLossBackward(repr) => Box::new(
3945 [
3946 &repr.log_probs,
3947 &repr.targets,
3948 &repr.input_lengths,
3949 &repr.target_lengths,
3950 &repr.grad_loss,
3951 ]
3952 .into_iter(),
3953 ),
3954 ModuleOperationIr::LayerNorm(repr) => match &repr.beta {
3955 Some(beta) => Box::new([&repr.input, &repr.gamma, beta].into_iter()),
3956 None => Box::new([&repr.input, &repr.gamma].into_iter()),
3957 },
3958 ModuleOperationIr::Unfold4d(repr) => Box::new([&repr.x].into_iter()),
3959 ModuleOperationIr::ConvTranspose1dWeightBackward(repr) => {
3960 Box::new([&repr.x, &repr.weight, &repr.output_grad].into_iter())
3961 }
3962 ModuleOperationIr::ConvTranspose1dBiasBackward(repr) => {
3963 Box::new([&repr.x, &repr.bias, &repr.output_grad].into_iter())
3964 }
3965 ModuleOperationIr::ConvTranspose2dWeightBackward(repr) => {
3966 Box::new([&repr.x, &repr.weight, &repr.output_grad].into_iter())
3967 }
3968 ModuleOperationIr::ConvTranspose2dBiasBackward(repr) => {
3969 Box::new([&repr.x, &repr.bias, &repr.output_grad].into_iter())
3970 }
3971 ModuleOperationIr::ConvTranspose3dWeightBackward(repr) => {
3972 Box::new([&repr.x, &repr.weight, &repr.output_grad].into_iter())
3973 }
3974 ModuleOperationIr::ConvTranspose3dBiasBackward(repr) => {
3975 Box::new([&repr.x, &repr.bias, &repr.output_grad].into_iter())
3976 }
3977 }
3978 }
3979 fn outputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
3980 match self {
3981 ModuleOperationIr::BatchNorm(repr) => Box::new([&repr.out].into_iter()),
3982 ModuleOperationIr::Embedding(repr) => Box::new([&repr.out].into_iter()),
3983 ModuleOperationIr::EmbeddingBackward(repr) => Box::new([&repr.out].into_iter()),
3984 ModuleOperationIr::Linear(repr) => Box::new([&repr.out].into_iter()),
3985 ModuleOperationIr::LinearXBackward(repr) => Box::new([&repr.out].into_iter()),
3986 ModuleOperationIr::LinearWeightBackward(repr) => Box::new([&repr.out].into_iter()),
3987 ModuleOperationIr::LinearBiasBackward(repr) => Box::new([&repr.out].into_iter()),
3988 ModuleOperationIr::Conv1d(repr) => Box::new([&repr.out].into_iter()),
3989 ModuleOperationIr::Conv1dXBackward(repr) => Box::new([&repr.out].into_iter()),
3990 ModuleOperationIr::Conv1dWeightBackward(repr) => Box::new([&repr.out].into_iter()),
3991 ModuleOperationIr::Conv1dBiasBackward(repr) => Box::new([&repr.out].into_iter()),
3992 ModuleOperationIr::Conv2d(repr) => Box::new([&repr.out].into_iter()),
3993 ModuleOperationIr::Conv2dXBackward(repr) => Box::new([&repr.out].into_iter()),
3994 ModuleOperationIr::Conv2dWeightBackward(repr) => Box::new([&repr.out].into_iter()),
3995 ModuleOperationIr::Conv2dBiasBackward(repr) => Box::new([&repr.out].into_iter()),
3996 ModuleOperationIr::Conv3d(repr) => Box::new([&repr.out].into_iter()),
3997 ModuleOperationIr::Conv3dXBackward(repr) => Box::new([&repr.out].into_iter()),
3998 ModuleOperationIr::Conv3dWeightBackward(repr) => Box::new([&repr.out].into_iter()),
3999 ModuleOperationIr::Conv3dBiasBackward(repr) => Box::new([&repr.out].into_iter()),
4000 ModuleOperationIr::DeformableConv2d(repr) => Box::new([&repr.out].into_iter()),
4001 ModuleOperationIr::DeformableConv2dBackward(repr) => {
4002 match (&repr.mask_grad, &repr.bias_grad) {
4003 (Some(mask_grad), Some(bias_grad)) => Box::new(
4004 [
4005 &repr.input_grad,
4006 &repr.offset_grad,
4007 &repr.weight_grad,
4008 mask_grad,
4009 bias_grad,
4010 ]
4011 .into_iter(),
4012 ),
4013 (Some(mask_grad), None) => Box::new(
4014 [
4015 &repr.input_grad,
4016 &repr.offset_grad,
4017 &repr.weight_grad,
4018 mask_grad,
4019 ]
4020 .into_iter(),
4021 ),
4022 (None, Some(bias_grad)) => Box::new(
4023 [
4024 &repr.input_grad,
4025 &repr.offset_grad,
4026 &repr.weight_grad,
4027 bias_grad,
4028 ]
4029 .into_iter(),
4030 ),
4031 (None, None) => Box::new(
4032 [&repr.input_grad, &repr.offset_grad, &repr.weight_grad].into_iter(),
4033 ),
4034 }
4035 }
4036 ModuleOperationIr::ConvTranspose1d(repr) => Box::new([&repr.out].into_iter()),
4037 ModuleOperationIr::ConvTranspose2d(repr) => Box::new([&repr.out].into_iter()),
4038 ModuleOperationIr::ConvTranspose3d(repr) => Box::new([&repr.out].into_iter()),
4039 ModuleOperationIr::AvgPool1d(repr) => Box::new([&repr.out].into_iter()),
4040 ModuleOperationIr::AvgPool2d(repr) => Box::new([&repr.out].into_iter()),
4041 ModuleOperationIr::AvgPool1dBackward(repr) => Box::new([&repr.out].into_iter()),
4042 ModuleOperationIr::AvgPool2dBackward(repr) => Box::new([&repr.out].into_iter()),
4043 ModuleOperationIr::AdaptiveAvgPool1d(repr) => Box::new([&repr.out].into_iter()),
4044 ModuleOperationIr::AdaptiveAvgPool2d(repr) => Box::new([&repr.out].into_iter()),
4045 ModuleOperationIr::AdaptiveAvgPool1dBackward(repr) => Box::new([&repr.out].into_iter()),
4046 ModuleOperationIr::AdaptiveAvgPool2dBackward(repr) => Box::new([&repr.out].into_iter()),
4047 ModuleOperationIr::AdaptiveAvgPool3d(repr) => Box::new([&repr.out].into_iter()),
4048 ModuleOperationIr::AdaptiveAvgPool3dBackward(repr) => Box::new([&repr.out].into_iter()),
4049 ModuleOperationIr::MaxPool1d(repr) => Box::new([&repr.out].into_iter()),
4050 ModuleOperationIr::MaxPool1dWithIndices(repr) => {
4051 Box::new([&repr.out, &repr.out_indices].into_iter())
4052 }
4053 ModuleOperationIr::MaxPool1dWithIndicesBackward(repr) => {
4054 Box::new([&repr.out].into_iter())
4055 }
4056 ModuleOperationIr::MaxPool2d(repr) => Box::new([&repr.out].into_iter()),
4057 ModuleOperationIr::MaxPool2dWithIndices(repr) => {
4058 Box::new([&repr.out, &repr.out_indices].into_iter())
4059 }
4060 ModuleOperationIr::MaxPool2dWithIndicesBackward(repr) => {
4061 Box::new([&repr.out].into_iter())
4062 }
4063 ModuleOperationIr::Interpolate(repr) => Box::new([&repr.out].into_iter()),
4064 ModuleOperationIr::InterpolateBackward(repr) => Box::new([&repr.out].into_iter()),
4065 ModuleOperationIr::Rfft(repr) => Box::new([&repr.out_re, &repr.out_im].into_iter()),
4066 ModuleOperationIr::IRfft(repr) => Box::new([&repr.out_signal].into_iter()),
4067 ModuleOperationIr::Attention(repr) => Box::new([&repr.out].into_iter()),
4068 ModuleOperationIr::CtcLoss(repr) => Box::new([&repr.out].into_iter()),
4069 ModuleOperationIr::CtcLossBackward(repr) => Box::new([&repr.out].into_iter()),
4070 ModuleOperationIr::LayerNorm(repr) => Box::new([&repr.out].into_iter()),
4071 ModuleOperationIr::Unfold4d(repr) => Box::new([&repr.out].into_iter()),
4072 ModuleOperationIr::ConvTranspose1dWeightBackward(repr) => {
4073 Box::new([&repr.out].into_iter())
4074 }
4075 ModuleOperationIr::ConvTranspose1dBiasBackward(repr) => {
4076 Box::new([&repr.out].into_iter())
4077 }
4078 ModuleOperationIr::ConvTranspose2dWeightBackward(repr) => {
4079 Box::new([&repr.out].into_iter())
4080 }
4081 ModuleOperationIr::ConvTranspose2dBiasBackward(repr) => {
4082 Box::new([&repr.out].into_iter())
4083 }
4084 ModuleOperationIr::ConvTranspose3dWeightBackward(repr) => {
4085 Box::new([&repr.out].into_iter())
4086 }
4087 ModuleOperationIr::ConvTranspose3dBiasBackward(repr) => {
4088 Box::new([&repr.out].into_iter())
4089 }
4090 }
4091 }
4092
4093 fn mark_read_only(&mut self, nodes: &[TensorId]) -> Vec<TensorIr> {
4094 let mut output = Vec::new();
4095
4096 match self {
4097 ModuleOperationIr::BatchNorm(repr) => {
4098 repr.x.mark_read_only(nodes, &mut output);
4099 repr.gamma.mark_read_only(nodes, &mut output);
4100 repr.beta.mark_read_only(nodes, &mut output);
4101 repr.mean.mark_read_only(nodes, &mut output);
4102 repr.variance.mark_read_only(nodes, &mut output);
4103 }
4104 ModuleOperationIr::Embedding(repr) => {
4105 repr.weights.mark_read_only(nodes, &mut output);
4106 repr.indices.mark_read_only(nodes, &mut output);
4107 }
4108 ModuleOperationIr::EmbeddingBackward(repr) => {
4109 repr.weights.mark_read_only(nodes, &mut output);
4110 repr.out_grad.mark_read_only(nodes, &mut output);
4111 repr.indices.mark_read_only(nodes, &mut output);
4112 }
4113 ModuleOperationIr::Linear(repr) => {
4114 repr.x.mark_read_only(nodes, &mut output);
4115 repr.weight.mark_read_only(nodes, &mut output);
4116
4117 if let Some(bias) = &mut repr.bias {
4118 bias.mark_read_only(nodes, &mut output);
4119 }
4120 }
4121 ModuleOperationIr::LinearXBackward(repr) => {
4122 repr.weight.mark_read_only(nodes, &mut output);
4123 repr.output_grad.mark_read_only(nodes, &mut output);
4124 }
4125 ModuleOperationIr::LinearWeightBackward(repr) => {
4126 repr.x.mark_read_only(nodes, &mut output);
4127 repr.output_grad.mark_read_only(nodes, &mut output);
4128 }
4129 ModuleOperationIr::LinearBiasBackward(repr) => {
4130 repr.output_grad.mark_read_only(nodes, &mut output);
4131 }
4132 ModuleOperationIr::Conv1d(repr) => {
4133 repr.x.mark_read_only(nodes, &mut output);
4134 repr.weight.mark_read_only(nodes, &mut output);
4135
4136 if let Some(bias) = &mut repr.bias {
4137 bias.mark_read_only(nodes, &mut output);
4138 }
4139 }
4140 ModuleOperationIr::Conv1dXBackward(repr) => {
4141 repr.x.mark_read_only(nodes, &mut output);
4142 repr.weight.mark_read_only(nodes, &mut output);
4143 repr.output_grad.mark_read_only(nodes, &mut output);
4144 }
4145 ModuleOperationIr::Conv1dWeightBackward(repr) => {
4146 repr.x.mark_read_only(nodes, &mut output);
4147 repr.weight.mark_read_only(nodes, &mut output);
4148 repr.output_grad.mark_read_only(nodes, &mut output);
4149 }
4150 ModuleOperationIr::Conv1dBiasBackward(repr) => {
4151 repr.x.mark_read_only(nodes, &mut output);
4152 repr.bias.mark_read_only(nodes, &mut output);
4153 repr.output_grad.mark_read_only(nodes, &mut output);
4154 }
4155 ModuleOperationIr::Conv2d(repr) => {
4156 repr.x.mark_read_only(nodes, &mut output);
4157 repr.weight.mark_read_only(nodes, &mut output);
4158
4159 if let Some(bias) = &mut repr.bias {
4160 bias.mark_read_only(nodes, &mut output);
4161 }
4162 }
4163 ModuleOperationIr::Conv2dXBackward(repr) => {
4164 repr.x.mark_read_only(nodes, &mut output);
4165 repr.weight.mark_read_only(nodes, &mut output);
4166 repr.output_grad.mark_read_only(nodes, &mut output);
4167 }
4168 ModuleOperationIr::Conv2dWeightBackward(repr) => {
4169 repr.x.mark_read_only(nodes, &mut output);
4170 repr.weight.mark_read_only(nodes, &mut output);
4171 repr.output_grad.mark_read_only(nodes, &mut output);
4172 }
4173 ModuleOperationIr::Conv2dBiasBackward(repr) => {
4174 repr.x.mark_read_only(nodes, &mut output);
4175 repr.bias.mark_read_only(nodes, &mut output);
4176 repr.output_grad.mark_read_only(nodes, &mut output);
4177 }
4178 ModuleOperationIr::Conv3d(repr) => {
4179 repr.x.mark_read_only(nodes, &mut output);
4180 repr.weight.mark_read_only(nodes, &mut output);
4181
4182 if let Some(bias) = &mut repr.bias {
4183 bias.mark_read_only(nodes, &mut output);
4184 }
4185 }
4186 ModuleOperationIr::Conv3dXBackward(repr) => {
4187 repr.x.mark_read_only(nodes, &mut output);
4188 repr.weight.mark_read_only(nodes, &mut output);
4189 repr.output_grad.mark_read_only(nodes, &mut output);
4190 }
4191 ModuleOperationIr::Conv3dWeightBackward(repr) => {
4192 repr.x.mark_read_only(nodes, &mut output);
4193 repr.weight.mark_read_only(nodes, &mut output);
4194 repr.output_grad.mark_read_only(nodes, &mut output);
4195 }
4196 ModuleOperationIr::Conv3dBiasBackward(repr) => {
4197 repr.x.mark_read_only(nodes, &mut output);
4198 repr.bias.mark_read_only(nodes, &mut output);
4199 repr.output_grad.mark_read_only(nodes, &mut output);
4200 }
4201 ModuleOperationIr::DeformableConv2d(repr) => {
4202 repr.x.mark_read_only(nodes, &mut output);
4203 repr.weight.mark_read_only(nodes, &mut output);
4204 repr.offset.mark_read_only(nodes, &mut output);
4205
4206 match (&mut repr.mask, &mut repr.bias) {
4207 (Some(mask), Some(bias)) => {
4208 mask.mark_read_only(nodes, &mut output);
4209 bias.mark_read_only(nodes, &mut output);
4210 }
4211 (Some(mask), None) => {
4212 mask.mark_read_only(nodes, &mut output);
4213 }
4214 (None, Some(bias)) => {
4215 bias.mark_read_only(nodes, &mut output);
4216 }
4217 (None, None) => {}
4218 };
4219 }
4220 ModuleOperationIr::DeformableConv2dBackward(repr) => {
4221 repr.x.mark_read_only(nodes, &mut output);
4222 repr.weight.mark_read_only(nodes, &mut output);
4223 repr.offset.mark_read_only(nodes, &mut output);
4224 repr.out_grad.mark_read_only(nodes, &mut output);
4225
4226 if let Some(mask) = repr.mask.as_mut() {
4227 mask.mark_read_only(nodes, &mut output);
4228 }
4229 if let Some(bias) = repr.bias.as_mut() {
4230 bias.mark_read_only(nodes, &mut output);
4231 }
4232 }
4233 ModuleOperationIr::ConvTranspose1d(repr) => {
4234 repr.x.mark_read_only(nodes, &mut output);
4235 repr.weight.mark_read_only(nodes, &mut output);
4236
4237 if let Some(bias) = &mut repr.bias {
4238 bias.mark_read_only(nodes, &mut output);
4239 }
4240 }
4241 ModuleOperationIr::ConvTranspose2d(repr) => {
4242 repr.x.mark_read_only(nodes, &mut output);
4243 repr.weight.mark_read_only(nodes, &mut output);
4244
4245 if let Some(bias) = &mut repr.bias {
4246 bias.mark_read_only(nodes, &mut output);
4247 }
4248 }
4249 ModuleOperationIr::ConvTranspose3d(repr) => {
4250 repr.x.mark_read_only(nodes, &mut output);
4251 repr.weight.mark_read_only(nodes, &mut output);
4252
4253 if let Some(bias) = &mut repr.bias {
4254 bias.mark_read_only(nodes, &mut output);
4255 }
4256 }
4257 ModuleOperationIr::AvgPool1d(repr) => {
4258 repr.x.mark_read_only(nodes, &mut output);
4259 }
4260 ModuleOperationIr::AvgPool2d(repr) => {
4261 repr.x.mark_read_only(nodes, &mut output);
4262 }
4263 ModuleOperationIr::AvgPool1dBackward(repr) => {
4264 repr.x.mark_read_only(nodes, &mut output);
4265 repr.grad.mark_read_only(nodes, &mut output);
4266 }
4267 ModuleOperationIr::AvgPool2dBackward(repr) => {
4268 repr.x.mark_read_only(nodes, &mut output);
4269 repr.grad.mark_read_only(nodes, &mut output);
4270 }
4271 ModuleOperationIr::AdaptiveAvgPool1d(repr) => {
4272 repr.x.mark_read_only(nodes, &mut output);
4273 }
4274 ModuleOperationIr::AdaptiveAvgPool2d(repr) => {
4275 repr.x.mark_read_only(nodes, &mut output);
4276 }
4277 ModuleOperationIr::AdaptiveAvgPool1dBackward(repr) => {
4278 repr.x.mark_read_only(nodes, &mut output);
4279 repr.grad.mark_read_only(nodes, &mut output);
4280 }
4281 ModuleOperationIr::AdaptiveAvgPool2dBackward(repr) => {
4282 repr.x.mark_read_only(nodes, &mut output);
4283 repr.grad.mark_read_only(nodes, &mut output);
4284 }
4285 ModuleOperationIr::AdaptiveAvgPool3d(repr) => {
4286 repr.x.mark_read_only(nodes, &mut output);
4287 }
4288 ModuleOperationIr::AdaptiveAvgPool3dBackward(repr) => {
4289 repr.x.mark_read_only(nodes, &mut output);
4290 repr.grad.mark_read_only(nodes, &mut output);
4291 }
4292 ModuleOperationIr::MaxPool1d(repr) => {
4293 repr.x.mark_read_only(nodes, &mut output);
4294 }
4295 ModuleOperationIr::MaxPool1dWithIndices(repr) => {
4296 repr.x.mark_read_only(nodes, &mut output);
4297 }
4298 ModuleOperationIr::MaxPool1dWithIndicesBackward(repr) => {
4299 repr.x.mark_read_only(nodes, &mut output);
4300 repr.grad.mark_read_only(nodes, &mut output);
4301 }
4302 ModuleOperationIr::MaxPool2d(repr) => {
4303 repr.x.mark_read_only(nodes, &mut output);
4304 }
4305 ModuleOperationIr::MaxPool2dWithIndices(repr) => {
4306 repr.x.mark_read_only(nodes, &mut output);
4307 }
4308 ModuleOperationIr::MaxPool2dWithIndicesBackward(repr) => {
4309 repr.x.mark_read_only(nodes, &mut output);
4310 repr.grad.mark_read_only(nodes, &mut output);
4311 }
4312 ModuleOperationIr::Interpolate(repr) => {
4313 repr.x.mark_read_only(nodes, &mut output);
4314 }
4315 ModuleOperationIr::InterpolateBackward(repr) => {
4316 repr.x.mark_read_only(nodes, &mut output);
4317 repr.grad.mark_read_only(nodes, &mut output);
4318 }
4319 ModuleOperationIr::Rfft(repr) => {
4320 repr.signal.mark_read_only(nodes, &mut output);
4321 }
4322 ModuleOperationIr::IRfft(repr) => {
4323 repr.input_re.mark_read_only(nodes, &mut output);
4324 repr.input_im.mark_read_only(nodes, &mut output);
4325 }
4326 ModuleOperationIr::Attention(repr) => {
4327 repr.query.mark_read_only(nodes, &mut output);
4328 repr.key.mark_read_only(nodes, &mut output);
4329 repr.value.mark_read_only(nodes, &mut output);
4330 if let Some(mask) = &mut repr.mask {
4331 mask.mark_read_only(nodes, &mut output);
4332 }
4333 if let Some(attn_bias) = &mut repr.attn_bias {
4334 attn_bias.mark_read_only(nodes, &mut output);
4335 }
4336 }
4337 ModuleOperationIr::CtcLoss(repr) => {
4338 repr.log_probs.mark_read_only(nodes, &mut output);
4339 repr.targets.mark_read_only(nodes, &mut output);
4340 repr.input_lengths.mark_read_only(nodes, &mut output);
4341 repr.target_lengths.mark_read_only(nodes, &mut output);
4342 }
4343 ModuleOperationIr::CtcLossBackward(repr) => {
4344 repr.log_probs.mark_read_only(nodes, &mut output);
4345 repr.targets.mark_read_only(nodes, &mut output);
4346 repr.input_lengths.mark_read_only(nodes, &mut output);
4347 repr.target_lengths.mark_read_only(nodes, &mut output);
4348 repr.grad_loss.mark_read_only(nodes, &mut output);
4349 }
4350 ModuleOperationIr::LayerNorm(repr) => {
4351 repr.input.mark_read_only(nodes, &mut output);
4352 repr.gamma.mark_read_only(nodes, &mut output);
4353 if let Some(beta) = &mut repr.beta {
4354 beta.mark_read_only(nodes, &mut output);
4355 }
4356 }
4357 ModuleOperationIr::Unfold4d(repr) => {
4358 repr.x.mark_read_only(nodes, &mut output);
4359 }
4360 ModuleOperationIr::ConvTranspose1dWeightBackward(repr) => {
4361 repr.x.mark_read_only(nodes, &mut output);
4362 repr.weight.mark_read_only(nodes, &mut output);
4363 repr.output_grad.mark_read_only(nodes, &mut output);
4364 }
4365 ModuleOperationIr::ConvTranspose1dBiasBackward(repr) => {
4366 repr.x.mark_read_only(nodes, &mut output);
4367 repr.bias.mark_read_only(nodes, &mut output);
4368 repr.output_grad.mark_read_only(nodes, &mut output);
4369 }
4370 ModuleOperationIr::ConvTranspose2dWeightBackward(repr) => {
4371 repr.x.mark_read_only(nodes, &mut output);
4372 repr.weight.mark_read_only(nodes, &mut output);
4373 repr.output_grad.mark_read_only(nodes, &mut output);
4374 }
4375 ModuleOperationIr::ConvTranspose2dBiasBackward(repr) => {
4376 repr.x.mark_read_only(nodes, &mut output);
4377 repr.bias.mark_read_only(nodes, &mut output);
4378 repr.output_grad.mark_read_only(nodes, &mut output);
4379 }
4380 ModuleOperationIr::ConvTranspose3dWeightBackward(repr) => {
4381 repr.x.mark_read_only(nodes, &mut output);
4382 repr.weight.mark_read_only(nodes, &mut output);
4383 repr.output_grad.mark_read_only(nodes, &mut output);
4384 }
4385 ModuleOperationIr::ConvTranspose3dBiasBackward(repr) => {
4386 repr.x.mark_read_only(nodes, &mut output);
4387 repr.bias.mark_read_only(nodes, &mut output);
4388 repr.output_grad.mark_read_only(nodes, &mut output);
4389 }
4390 };
4391
4392 output
4393 }
4394
4395 fn visit_mut(&mut self, v: &mut impl IrVisitorMut) {
4396 match self {
4397 ModuleOperationIr::BatchNorm(repr) => {
4398 v.visit_tensor_mut(&mut repr.x);
4399 v.visit_tensor_mut(&mut repr.gamma);
4400 v.visit_tensor_mut(&mut repr.beta);
4401 v.visit_tensor_mut(&mut repr.mean);
4402 v.visit_tensor_mut(&mut repr.variance);
4403 v.visit_scalar_mut(&mut repr.epsilon);
4404 v.visit_tensor_mut(&mut repr.out);
4405 }
4406 ModuleOperationIr::Embedding(repr) => {
4407 v.visit_tensor_mut(&mut repr.weights);
4408 v.visit_tensor_mut(&mut repr.indices);
4409 v.visit_tensor_mut(&mut repr.out);
4410 }
4411 ModuleOperationIr::EmbeddingBackward(repr) => {
4412 v.visit_tensor_mut(&mut repr.weights);
4413 v.visit_tensor_mut(&mut repr.out_grad);
4414 v.visit_tensor_mut(&mut repr.indices);
4415 v.visit_tensor_mut(&mut repr.out);
4416 }
4417 ModuleOperationIr::Linear(repr) => {
4418 v.visit_tensor_mut(&mut repr.x);
4419 v.visit_tensor_mut(&mut repr.weight);
4420 if let Some(bias) = &mut repr.bias {
4421 v.visit_tensor_mut(bias);
4422 }
4423 v.visit_tensor_mut(&mut repr.out);
4424 }
4425 ModuleOperationIr::LinearXBackward(repr) => {
4426 v.visit_tensor_mut(&mut repr.weight);
4427 v.visit_tensor_mut(&mut repr.output_grad);
4428 v.visit_tensor_mut(&mut repr.out);
4429 }
4430 ModuleOperationIr::LinearWeightBackward(repr) => {
4431 v.visit_tensor_mut(&mut repr.x);
4432 v.visit_tensor_mut(&mut repr.output_grad);
4433 v.visit_tensor_mut(&mut repr.out);
4434 }
4435 ModuleOperationIr::LinearBiasBackward(repr) => {
4436 v.visit_tensor_mut(&mut repr.output_grad);
4437 v.visit_tensor_mut(&mut repr.out);
4438 }
4439 ModuleOperationIr::Conv1d(repr) => {
4440 v.visit_tensor_mut(&mut repr.x);
4441 v.visit_tensor_mut(&mut repr.weight);
4442 if let Some(bias) = &mut repr.bias {
4443 v.visit_tensor_mut(bias);
4444 }
4445 v.visit_tensor_mut(&mut repr.out);
4446 }
4447 ModuleOperationIr::Conv1dXBackward(repr) => {
4448 v.visit_tensor_mut(&mut repr.x);
4449 v.visit_tensor_mut(&mut repr.weight);
4450 v.visit_tensor_mut(&mut repr.output_grad);
4451 v.visit_tensor_mut(&mut repr.out);
4452 }
4453 ModuleOperationIr::Conv1dWeightBackward(repr) => {
4454 v.visit_tensor_mut(&mut repr.x);
4455 v.visit_tensor_mut(&mut repr.weight);
4456 v.visit_tensor_mut(&mut repr.output_grad);
4457 v.visit_tensor_mut(&mut repr.out);
4458 }
4459 ModuleOperationIr::Conv1dBiasBackward(repr) => {
4460 v.visit_tensor_mut(&mut repr.x);
4461 v.visit_tensor_mut(&mut repr.bias);
4462 v.visit_tensor_mut(&mut repr.output_grad);
4463 v.visit_tensor_mut(&mut repr.out);
4464 }
4465 ModuleOperationIr::Conv2d(repr) => {
4466 v.visit_tensor_mut(&mut repr.x);
4467 v.visit_tensor_mut(&mut repr.weight);
4468 if let Some(bias) = &mut repr.bias {
4469 v.visit_tensor_mut(bias);
4470 }
4471 v.visit_tensor_mut(&mut repr.out);
4472 }
4473 ModuleOperationIr::Conv2dXBackward(repr) => {
4474 v.visit_tensor_mut(&mut repr.x);
4475 v.visit_tensor_mut(&mut repr.weight);
4476 v.visit_tensor_mut(&mut repr.output_grad);
4477 v.visit_tensor_mut(&mut repr.out);
4478 }
4479 ModuleOperationIr::Conv2dWeightBackward(repr) => {
4480 v.visit_tensor_mut(&mut repr.x);
4481 v.visit_tensor_mut(&mut repr.weight);
4482 v.visit_tensor_mut(&mut repr.output_grad);
4483 v.visit_tensor_mut(&mut repr.out);
4484 }
4485 ModuleOperationIr::Conv2dBiasBackward(repr) => {
4486 v.visit_tensor_mut(&mut repr.x);
4487 v.visit_tensor_mut(&mut repr.bias);
4488 v.visit_tensor_mut(&mut repr.output_grad);
4489 v.visit_tensor_mut(&mut repr.out);
4490 }
4491 ModuleOperationIr::Conv3d(repr) => {
4492 v.visit_tensor_mut(&mut repr.x);
4493 v.visit_tensor_mut(&mut repr.weight);
4494 if let Some(bias) = &mut repr.bias {
4495 v.visit_tensor_mut(bias);
4496 }
4497 v.visit_tensor_mut(&mut repr.out);
4498 }
4499 ModuleOperationIr::Conv3dXBackward(repr) => {
4500 v.visit_tensor_mut(&mut repr.x);
4501 v.visit_tensor_mut(&mut repr.weight);
4502 v.visit_tensor_mut(&mut repr.output_grad);
4503 v.visit_tensor_mut(&mut repr.out);
4504 }
4505 ModuleOperationIr::Conv3dWeightBackward(repr) => {
4506 v.visit_tensor_mut(&mut repr.x);
4507 v.visit_tensor_mut(&mut repr.weight);
4508 v.visit_tensor_mut(&mut repr.output_grad);
4509 v.visit_tensor_mut(&mut repr.out);
4510 }
4511 ModuleOperationIr::Conv3dBiasBackward(repr) => {
4512 v.visit_tensor_mut(&mut repr.x);
4513 v.visit_tensor_mut(&mut repr.bias);
4514 v.visit_tensor_mut(&mut repr.output_grad);
4515 v.visit_tensor_mut(&mut repr.out);
4516 }
4517 ModuleOperationIr::DeformableConv2d(repr) => {
4518 v.visit_tensor_mut(&mut repr.x);
4519 v.visit_tensor_mut(&mut repr.offset);
4520 v.visit_tensor_mut(&mut repr.weight);
4521 if let Some(mask) = &mut repr.mask {
4522 v.visit_tensor_mut(mask);
4523 }
4524 if let Some(bias) = &mut repr.bias {
4525 v.visit_tensor_mut(bias);
4526 }
4527 v.visit_tensor_mut(&mut repr.out);
4528 }
4529 ModuleOperationIr::DeformableConv2dBackward(repr) => {
4530 v.visit_tensor_mut(&mut repr.x);
4531 v.visit_tensor_mut(&mut repr.offset);
4532 v.visit_tensor_mut(&mut repr.weight);
4533 v.visit_tensor_mut(&mut repr.out_grad);
4534 if let Some(mask) = &mut repr.mask {
4535 v.visit_tensor_mut(mask);
4536 }
4537 if let Some(bias) = &mut repr.bias {
4538 v.visit_tensor_mut(bias);
4539 }
4540 v.visit_tensor_mut(&mut repr.input_grad);
4541 v.visit_tensor_mut(&mut repr.offset_grad);
4542 v.visit_tensor_mut(&mut repr.weight_grad);
4543 if let Some(mask_grad) = &mut repr.mask_grad {
4544 v.visit_tensor_mut(mask_grad);
4545 }
4546 if let Some(bias_grad) = &mut repr.bias_grad {
4547 v.visit_tensor_mut(bias_grad);
4548 }
4549 }
4550 ModuleOperationIr::ConvTranspose1d(repr) => {
4551 v.visit_tensor_mut(&mut repr.x);
4552 v.visit_tensor_mut(&mut repr.weight);
4553 if let Some(bias) = &mut repr.bias {
4554 v.visit_tensor_mut(bias);
4555 }
4556 v.visit_tensor_mut(&mut repr.out);
4557 }
4558 ModuleOperationIr::ConvTranspose2d(repr) => {
4559 v.visit_tensor_mut(&mut repr.x);
4560 v.visit_tensor_mut(&mut repr.weight);
4561 if let Some(bias) = &mut repr.bias {
4562 v.visit_tensor_mut(bias);
4563 }
4564 v.visit_tensor_mut(&mut repr.out);
4565 }
4566 ModuleOperationIr::ConvTranspose3d(repr) => {
4567 v.visit_tensor_mut(&mut repr.x);
4568 v.visit_tensor_mut(&mut repr.weight);
4569 if let Some(bias) = &mut repr.bias {
4570 v.visit_tensor_mut(bias);
4571 }
4572 v.visit_tensor_mut(&mut repr.out);
4573 }
4574 ModuleOperationIr::AvgPool1d(repr) => {
4575 v.visit_tensor_mut(&mut repr.x);
4576 v.visit_tensor_mut(&mut repr.out);
4577 }
4578 ModuleOperationIr::AvgPool2d(repr) => {
4579 v.visit_tensor_mut(&mut repr.x);
4580 v.visit_tensor_mut(&mut repr.out);
4581 }
4582 ModuleOperationIr::AvgPool1dBackward(repr) => {
4583 v.visit_tensor_mut(&mut repr.x);
4584 v.visit_tensor_mut(&mut repr.grad);
4585 v.visit_tensor_mut(&mut repr.out);
4586 }
4587 ModuleOperationIr::AvgPool2dBackward(repr) => {
4588 v.visit_tensor_mut(&mut repr.x);
4589 v.visit_tensor_mut(&mut repr.grad);
4590 v.visit_tensor_mut(&mut repr.out);
4591 }
4592 ModuleOperationIr::AdaptiveAvgPool1d(repr) => {
4593 v.visit_tensor_mut(&mut repr.x);
4594 v.visit_tensor_mut(&mut repr.out);
4595 }
4596 ModuleOperationIr::AdaptiveAvgPool2d(repr) => {
4597 v.visit_tensor_mut(&mut repr.x);
4598 v.visit_tensor_mut(&mut repr.out);
4599 }
4600 ModuleOperationIr::AdaptiveAvgPool1dBackward(repr) => {
4601 v.visit_tensor_mut(&mut repr.x);
4602 v.visit_tensor_mut(&mut repr.grad);
4603 v.visit_tensor_mut(&mut repr.out);
4604 }
4605 ModuleOperationIr::AdaptiveAvgPool2dBackward(repr) => {
4606 v.visit_tensor_mut(&mut repr.x);
4607 v.visit_tensor_mut(&mut repr.grad);
4608 v.visit_tensor_mut(&mut repr.out);
4609 }
4610 ModuleOperationIr::AdaptiveAvgPool3d(repr) => {
4611 v.visit_tensor_mut(&mut repr.x);
4612 v.visit_tensor_mut(&mut repr.out);
4613 }
4614 ModuleOperationIr::AdaptiveAvgPool3dBackward(repr) => {
4615 v.visit_tensor_mut(&mut repr.x);
4616 v.visit_tensor_mut(&mut repr.grad);
4617 v.visit_tensor_mut(&mut repr.out);
4618 }
4619 ModuleOperationIr::MaxPool1d(repr) => {
4620 v.visit_tensor_mut(&mut repr.x);
4621 v.visit_tensor_mut(&mut repr.out);
4622 }
4623 ModuleOperationIr::MaxPool1dWithIndices(repr) => {
4624 v.visit_tensor_mut(&mut repr.x);
4625 v.visit_tensor_mut(&mut repr.out);
4626 v.visit_tensor_mut(&mut repr.out_indices);
4627 }
4628 ModuleOperationIr::MaxPool1dWithIndicesBackward(repr) => {
4629 v.visit_tensor_mut(&mut repr.x);
4630 v.visit_tensor_mut(&mut repr.indices);
4631 v.visit_tensor_mut(&mut repr.grad);
4632 v.visit_tensor_mut(&mut repr.out);
4633 }
4634 ModuleOperationIr::MaxPool2d(repr) => {
4635 v.visit_tensor_mut(&mut repr.x);
4636 v.visit_tensor_mut(&mut repr.out);
4637 }
4638 ModuleOperationIr::MaxPool2dWithIndices(repr) => {
4639 v.visit_tensor_mut(&mut repr.x);
4640 v.visit_tensor_mut(&mut repr.out);
4641 v.visit_tensor_mut(&mut repr.out_indices);
4642 }
4643 ModuleOperationIr::MaxPool2dWithIndicesBackward(repr) => {
4644 v.visit_tensor_mut(&mut repr.x);
4645 v.visit_tensor_mut(&mut repr.indices);
4646 v.visit_tensor_mut(&mut repr.grad);
4647 v.visit_tensor_mut(&mut repr.out);
4648 }
4649 ModuleOperationIr::Interpolate(repr) => {
4650 v.visit_tensor_mut(&mut repr.x);
4651 v.visit_tensor_mut(&mut repr.out);
4652 }
4653 ModuleOperationIr::InterpolateBackward(repr) => {
4654 v.visit_tensor_mut(&mut repr.x);
4655 v.visit_tensor_mut(&mut repr.grad);
4656 v.visit_tensor_mut(&mut repr.out);
4657 }
4658 ModuleOperationIr::Rfft(repr) => {
4659 v.visit_tensor_mut(&mut repr.signal);
4660 v.visit_tensor_mut(&mut repr.out_re);
4661 v.visit_tensor_mut(&mut repr.out_im);
4662 }
4663 ModuleOperationIr::IRfft(repr) => {
4664 v.visit_tensor_mut(&mut repr.input_re);
4665 v.visit_tensor_mut(&mut repr.input_im);
4666 v.visit_tensor_mut(&mut repr.out_signal);
4667 }
4668 ModuleOperationIr::Attention(repr) => {
4669 v.visit_tensor_mut(&mut repr.query);
4670 v.visit_tensor_mut(&mut repr.key);
4671 v.visit_tensor_mut(&mut repr.value);
4672 if let Some(mask) = &mut repr.mask {
4673 v.visit_tensor_mut(mask);
4674 }
4675 if let Some(attn_bias) = &mut repr.attn_bias {
4676 v.visit_tensor_mut(attn_bias);
4677 }
4678 v.visit_tensor_mut(&mut repr.out);
4679 if let Some(scale) = &mut repr.options.scale {
4680 v.visit_scalar_mut(scale);
4681 }
4682 if let Some(softcap) = &mut repr.options.softcap {
4683 v.visit_scalar_mut(softcap);
4684 }
4685 }
4686 ModuleOperationIr::CtcLoss(repr) => {
4687 v.visit_tensor_mut(&mut repr.log_probs);
4688 v.visit_tensor_mut(&mut repr.targets);
4689 v.visit_tensor_mut(&mut repr.input_lengths);
4690 v.visit_tensor_mut(&mut repr.target_lengths);
4691 v.visit_tensor_mut(&mut repr.out);
4692 }
4693 ModuleOperationIr::CtcLossBackward(repr) => {
4694 v.visit_tensor_mut(&mut repr.log_probs);
4695 v.visit_tensor_mut(&mut repr.targets);
4696 v.visit_tensor_mut(&mut repr.input_lengths);
4697 v.visit_tensor_mut(&mut repr.target_lengths);
4698 v.visit_tensor_mut(&mut repr.grad_loss);
4699 v.visit_tensor_mut(&mut repr.out);
4700 }
4701 ModuleOperationIr::LayerNorm(repr) => {
4702 v.visit_tensor_mut(&mut repr.input);
4703 v.visit_tensor_mut(&mut repr.gamma);
4704 if let Some(beta) = &mut repr.beta {
4705 v.visit_tensor_mut(beta);
4706 }
4707 v.visit_tensor_mut(&mut repr.out);
4708 v.visit_scalar_mut(&mut repr.epsilon);
4709 }
4710 ModuleOperationIr::Unfold4d(repr) => {
4711 v.visit_tensor_mut(&mut repr.x);
4712 v.visit_tensor_mut(&mut repr.out);
4713 }
4714 ModuleOperationIr::ConvTranspose1dWeightBackward(repr) => {
4715 v.visit_tensor_mut(&mut repr.x);
4716 v.visit_tensor_mut(&mut repr.weight);
4717 v.visit_tensor_mut(&mut repr.output_grad);
4718 v.visit_tensor_mut(&mut repr.out);
4719 }
4720 ModuleOperationIr::ConvTranspose1dBiasBackward(repr) => {
4721 v.visit_tensor_mut(&mut repr.x);
4722 v.visit_tensor_mut(&mut repr.bias);
4723 v.visit_tensor_mut(&mut repr.output_grad);
4724 v.visit_tensor_mut(&mut repr.out);
4725 }
4726 ModuleOperationIr::ConvTranspose2dWeightBackward(repr) => {
4727 v.visit_tensor_mut(&mut repr.x);
4728 v.visit_tensor_mut(&mut repr.weight);
4729 v.visit_tensor_mut(&mut repr.output_grad);
4730 v.visit_tensor_mut(&mut repr.out);
4731 }
4732 ModuleOperationIr::ConvTranspose2dBiasBackward(repr) => {
4733 v.visit_tensor_mut(&mut repr.x);
4734 v.visit_tensor_mut(&mut repr.bias);
4735 v.visit_tensor_mut(&mut repr.output_grad);
4736 v.visit_tensor_mut(&mut repr.out);
4737 }
4738 ModuleOperationIr::ConvTranspose3dWeightBackward(repr) => {
4739 v.visit_tensor_mut(&mut repr.x);
4740 v.visit_tensor_mut(&mut repr.weight);
4741 v.visit_tensor_mut(&mut repr.output_grad);
4742 v.visit_tensor_mut(&mut repr.out);
4743 }
4744 ModuleOperationIr::ConvTranspose3dBiasBackward(repr) => {
4745 v.visit_tensor_mut(&mut repr.x);
4746 v.visit_tensor_mut(&mut repr.bias);
4747 v.visit_tensor_mut(&mut repr.output_grad);
4748 v.visit_tensor_mut(&mut repr.out);
4749 }
4750 }
4751 }
4752}
4753
4754impl DistributedOperationIr {
4755 fn inputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
4756 match self {
4757 DistributedOperationIr::AllReduce(repr) => Box::new([&repr.tensor].into_iter()),
4758 DistributedOperationIr::SyncCollective => Box::new([].into_iter()),
4759 }
4760 }
4761
4762 fn outputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
4763 match self {
4764 DistributedOperationIr::AllReduce(repr) => Box::new([&repr.out].into_iter()),
4765 DistributedOperationIr::SyncCollective => Box::new([].into_iter()),
4766 }
4767 }
4768
4769 fn mark_read_only(&mut self, nodes: &[TensorId]) -> Vec<TensorIr> {
4770 let mut output = Vec::new();
4771
4772 match self {
4773 DistributedOperationIr::AllReduce(repr) => {
4774 repr.tensor.mark_read_only(nodes, &mut output);
4775 }
4776 DistributedOperationIr::SyncCollective => {}
4777 }
4778
4779 output
4780 }
4781
4782 fn visit_mut(&mut self, v: &mut impl IrVisitorMut) {
4783 match self {
4784 DistributedOperationIr::AllReduce(repr) => {
4785 v.visit_tensor_mut(&mut repr.tensor);
4786 v.visit_tensor_mut(&mut repr.out);
4787 }
4788 DistributedOperationIr::SyncCollective => {}
4789 }
4790 }
4791}
4792
4793impl InitOperationIr {
4794 fn inputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
4795 Box::new([].into_iter())
4796 }
4797 fn outputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
4798 Box::new([&self.out].into_iter())
4799 }
4800
4801 fn visit_mut(&mut self, v: &mut impl IrVisitorMut) {
4802 v.visit_tensor_mut(&mut self.out);
4803 }
4804}
4805
4806impl TensorIr {
4807 fn mark_read_only(&mut self, nodes: &[TensorId], output: &mut Vec<TensorIr>) {
4808 if self.status == TensorStatus::ReadWrite && nodes.contains(&self.id) {
4809 output.push(self.clone());
4810 self.status = TensorStatus::ReadOnly;
4811 }
4812 }
4813}
4814
4815impl core::hash::Hash for RandomOpIr {
4816 fn hash<H: core::hash::Hasher>(&self, state: &mut H) {
4817 self.out.hash(state);
4818
4819 match self.distribution {
4820 Distribution::Default => 1u8.hash(state),
4821 Distribution::Bernoulli(_) => 2u8.hash(state),
4822 Distribution::Uniform(_, _) => 3u8.hash(state),
4823 Distribution::Normal(_, _) => 4u8.hash(state),
4824 }
4825 }
4826}
4827
4828pub trait OperationOutput<O> {
4830 fn output(self) -> O;
4832
4833 fn outputs<const N: usize>(self) -> [O; N];
4835}
4836
4837impl<O: core::fmt::Debug> OperationOutput<O> for Vec<O> {
4838 fn output(self) -> O {
4839 let [tensor] = self.outputs();
4840 tensor
4841 }
4842
4843 fn outputs<const N: usize>(self) -> [O; N] {
4844 self.try_into().unwrap()
4845 }
4846}
4847
4848#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4850pub struct SortOpIr {
4851 pub input: TensorIr,
4853 pub dim: usize,
4855 pub descending: bool,
4857 pub out: TensorIr,
4859}
4860
4861#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4863pub struct SortWithIndicesOpIr {
4864 pub input: TensorIr,
4866 pub dim: usize,
4868 pub descending: bool,
4870 pub out: TensorIr,
4872 pub out_indices: TensorIr,
4874}
4875
4876#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4878pub struct LayerNormOpIr {
4879 pub input: TensorIr,
4881 pub gamma: TensorIr,
4883 pub beta: Option<TensorIr>,
4885 pub epsilon: ScalarIr,
4887 pub out: TensorIr,
4889}
4890
4891#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4893pub struct Unfold4dOpIr {
4894 pub x: TensorIr,
4896 pub kernel_size: [usize; 2],
4898 pub options: Unfold4dOptionsIr,
4900 pub out: TensorIr,
4902}
4903
4904#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4906pub struct Unfold4dOptionsIr {
4907 pub stride: [usize; 2],
4909 pub padding: [usize; 2],
4911 pub dilation: [usize; 2],
4913}
4914
4915#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4917pub struct ConvTranspose1dWeightBackwardOpIr {
4918 pub x: TensorIr,
4920 pub weight: TensorIr,
4922 pub output_grad: TensorIr,
4924 pub options: ConvTranspose1dOptionsIr,
4926 pub out: TensorIr,
4928}
4929
4930#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4932pub struct ConvTranspose1dBiasBackwardOpIr {
4933 pub x: TensorIr,
4935 pub bias: TensorIr,
4937 pub output_grad: TensorIr,
4939 pub out: TensorIr,
4941}
4942
4943#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4945pub struct ConvTranspose2dWeightBackwardOpIr {
4946 pub x: TensorIr,
4948 pub weight: TensorIr,
4950 pub output_grad: TensorIr,
4952 pub options: ConvTranspose2dOptionsIr,
4954 pub out: TensorIr,
4956}
4957
4958#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4960pub struct ConvTranspose2dBiasBackwardOpIr {
4961 pub x: TensorIr,
4963 pub bias: TensorIr,
4965 pub output_grad: TensorIr,
4967 pub out: TensorIr,
4969}
4970
4971#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4973pub struct ConvTranspose3dWeightBackwardOpIr {
4974 pub x: TensorIr,
4976 pub weight: TensorIr,
4978 pub output_grad: TensorIr,
4980 pub options: ConvTranspose3dOptionsIr,
4982 pub out: TensorIr,
4984}
4985
4986#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4988pub struct ConvTranspose3dBiasBackwardOpIr {
4989 pub x: TensorIr,
4991 pub bias: TensorIr,
4993 pub output_grad: TensorIr,
4995 pub out: TensorIr,
4997}
4998
4999#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
5002pub struct HardSigmoidOpIr {
5003 pub tensor: TensorIr,
5005 pub alpha: ScalarIr,
5007 pub beta: ScalarIr,
5009 pub out: TensorIr,
5011}
5012
5013#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
5019pub enum ActivationOperationIr {
5020 Relu(UnaryOpIr),
5022 ReluBackward(BinaryOpIr),
5025 LeakyRelu(ScalarOpIr),
5028 PRelu(BinaryOpIr),
5031 Gelu(UnaryOpIr),
5033 GeluBackward(BinaryOpIr),
5036 Sigmoid(UnaryOpIr),
5038 SigmoidBackward(BinaryOpIr),
5041 HardSigmoid(HardSigmoidOpIr),
5043 LogSigmoid(UnaryOpIr),
5045 LogSigmoidBackward(BinaryOpIr),
5048 Softmax(DimOpIr),
5050 LogSoftmax(DimOpIr),
5052 Softmin(DimOpIr),
5054}
5055
5056macro_rules! activation_ir_tensor_access {
5061 ($( $variant:ident => [ $($field:ident),+ ] ),+ $(,)?) => {
5062 impl ActivationOperationIr {
5063 fn inputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
5064 match self {
5065 $( Self::$variant(repr) => Box::new([$(&repr.$field),+].into_iter()), )+
5066 }
5067 }
5068
5069 fn outputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
5070 match self {
5071 $( Self::$variant(repr) => Box::new([&repr.out].into_iter()), )+
5072 }
5073 }
5074
5075 fn mark_read_only(&mut self, nodes: &[TensorId]) -> Vec<TensorIr> {
5076 let mut output = Vec::new();
5077 match self {
5078 $( Self::$variant(repr) => {
5079 $( repr.$field.mark_read_only(nodes, &mut output); )+
5080 } )+
5081 }
5082 output
5083 }
5084 }
5085 };
5086}
5087
5088impl ActivationOperationIr {
5089 fn visit_mut(&mut self, v: &mut impl IrVisitorMut) {
5090 match self {
5091 ActivationOperationIr::Relu(repr) => {
5092 v.visit_tensor_mut(&mut repr.input);
5093 v.visit_tensor_mut(&mut repr.out);
5094 }
5095 ActivationOperationIr::ReluBackward(repr) => {
5096 v.visit_tensor_mut(&mut repr.lhs);
5097 v.visit_tensor_mut(&mut repr.rhs);
5098 v.visit_tensor_mut(&mut repr.out);
5099 }
5100 ActivationOperationIr::LeakyRelu(repr) => {
5101 v.visit_tensor_mut(&mut repr.lhs);
5102 v.visit_tensor_mut(&mut repr.out);
5103 v.visit_scalar_mut(&mut repr.rhs);
5104 }
5105 ActivationOperationIr::PRelu(repr) => {
5106 v.visit_tensor_mut(&mut repr.lhs);
5107 v.visit_tensor_mut(&mut repr.rhs);
5108 v.visit_tensor_mut(&mut repr.out);
5109 }
5110 ActivationOperationIr::Gelu(repr) => {
5111 v.visit_tensor_mut(&mut repr.input);
5112 v.visit_tensor_mut(&mut repr.out);
5113 }
5114 ActivationOperationIr::GeluBackward(repr) => {
5115 v.visit_tensor_mut(&mut repr.lhs);
5116 v.visit_tensor_mut(&mut repr.rhs);
5117 v.visit_tensor_mut(&mut repr.out);
5118 }
5119 ActivationOperationIr::Sigmoid(repr) => {
5120 v.visit_tensor_mut(&mut repr.input);
5121 v.visit_tensor_mut(&mut repr.out);
5122 }
5123 ActivationOperationIr::SigmoidBackward(repr) => {
5124 v.visit_tensor_mut(&mut repr.lhs);
5125 v.visit_tensor_mut(&mut repr.rhs);
5126 v.visit_tensor_mut(&mut repr.out);
5127 }
5128 ActivationOperationIr::HardSigmoid(repr) => {
5129 v.visit_tensor_mut(&mut repr.tensor);
5130 v.visit_tensor_mut(&mut repr.out);
5131 v.visit_scalar_mut(&mut repr.alpha);
5132 v.visit_scalar_mut(&mut repr.beta);
5133 }
5134 ActivationOperationIr::LogSigmoid(repr) => {
5135 v.visit_tensor_mut(&mut repr.input);
5136 v.visit_tensor_mut(&mut repr.out);
5137 }
5138 ActivationOperationIr::LogSigmoidBackward(repr) => {
5139 v.visit_tensor_mut(&mut repr.lhs);
5140 v.visit_tensor_mut(&mut repr.rhs);
5141 v.visit_tensor_mut(&mut repr.out);
5142 }
5143 ActivationOperationIr::Softmax(repr) => {
5144 v.visit_tensor_mut(&mut repr.input);
5145 v.visit_tensor_mut(&mut repr.out);
5146 }
5147 ActivationOperationIr::LogSoftmax(repr) => {
5148 v.visit_tensor_mut(&mut repr.input);
5149 v.visit_tensor_mut(&mut repr.out);
5150 }
5151 ActivationOperationIr::Softmin(repr) => {
5152 v.visit_tensor_mut(&mut repr.input);
5153 v.visit_tensor_mut(&mut repr.out);
5154 }
5155 }
5156 }
5157}
5158
5159activation_ir_tensor_access! {
5160 Relu => [input],
5161 ReluBackward => [lhs, rhs],
5162 LeakyRelu => [lhs],
5163 PRelu => [lhs, rhs],
5164 Gelu => [input],
5165 GeluBackward => [lhs, rhs],
5166 Sigmoid => [input],
5167 SigmoidBackward => [lhs, rhs],
5168 HardSigmoid => [tensor],
5169 LogSigmoid => [input],
5170 LogSigmoidBackward => [lhs, rhs],
5171 Softmax => [input],
5172 LogSoftmax => [input],
5173 Softmin => [input],
5174}
5175
5176#[cfg(test)]
5177mod visit_mut_tests {
5178 use super::*;
5179 use burn_backend::{DType, Shape};
5180
5181 fn tensor(id: u64) -> TensorIr {
5182 TensorIr::uninit(TensorId::new(id), Shape::from([2, 2]), DType::F32)
5183 }
5184
5185 #[derive(Default)]
5187 struct CollectVisitor {
5188 scalars: Vec<ScalarIr>,
5189 rewrite_scalar: Option<ScalarIr>,
5190 }
5191
5192 impl IrVisitorMut for CollectVisitor {
5193 fn visit_tensor_mut(&mut self, tensor: &mut TensorIr) {
5194 tensor.id = TensorId::new(tensor.id.value() + 100);
5195 }
5196
5197 fn visit_scalar_mut(&mut self, scalar: &mut ScalarIr) {
5198 self.scalars.push(*scalar);
5199 if let Some(value) = self.rewrite_scalar {
5200 *scalar = value;
5201 }
5202 }
5203 }
5204
5205 #[test]
5206 fn visit_mut_visits_all_tensors_and_scalars() {
5207 let mut mul = OperationIr::NumericFloat(
5209 DType::F32,
5210 NumericOperationIr::MulScalar(ScalarOpIr {
5211 lhs: tensor(1),
5212 rhs: ScalarIr::Float(2.0),
5213 out: tensor(2),
5214 }),
5215 );
5216
5217 let mut visitor = CollectVisitor {
5219 rewrite_scalar: Some(ScalarIr::Float(9.0)),
5220 ..Default::default()
5221 };
5222 mul.visit_mut(&mut visitor);
5223
5224 let ids: Vec<u64> = mul
5225 .inputs()
5226 .chain(mul.outputs())
5227 .map(|t| t.id.value())
5228 .collect();
5229 assert_eq!(ids, vec![101, 102]);
5230 assert_eq!(visitor.scalars, vec![ScalarIr::Float(2.0)]);
5232
5233 let mut after = CollectVisitor::default();
5235 mul.visit_mut(&mut after);
5236 assert_eq!(after.scalars, vec![ScalarIr::Float(9.0)]);
5237
5238 let mut reshape = OperationIr::BaseFloat(BaseOperationIr::Reshape(ShapeOpIr {
5240 input: tensor(10),
5241 out: tensor(11),
5242 }));
5243
5244 let mut visitor = CollectVisitor::default();
5245 reshape.visit_mut(&mut visitor);
5246 let ids: Vec<u64> = reshape
5247 .inputs()
5248 .chain(reshape.outputs())
5249 .map(|t| t.id.value())
5250 .collect();
5251 assert_eq!(ids, vec![110, 111]);
5252 assert_eq!(visitor.scalars.len(), 0);
5253 }
5254}