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