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