Skip to main content

burn_ir/
operation.rs

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
21/// Visitor for mutating the components of an [`OperationIr`] in place.
22pub trait IrVisitorMut {
23    /// Visit a [`TensorIr`] mutably.
24    fn visit_tensor_mut(&mut self, _tensor: &mut TensorIr) {}
25    /// Visit a [`ScalarIr`] mutably.
26    fn visit_scalar_mut(&mut self, _scalar: &mut ScalarIr) {}
27    /// Visit a slice range mutably.
28    fn visit_range_mut(&mut self, _range: &mut Slice) {}
29}
30
31/// Custom operation in fusion stream, declaring its inputs, outputs and scalar arguments.
32#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
33pub struct CustomOpIr {
34    /// Unique identifier of the operation.
35    pub id: String,
36    /// Input tensors used in the custom operation.
37    pub inputs: Vec<TensorIr>,
38    /// Output tensors used in the custom operation.
39    pub outputs: Vec<TensorIr>,
40    /// Non-tensor scalar arguments, in declaration order.
41    ///
42    /// Carried through fusion's relativization like any other scalar (see the
43    /// `RelativeOps for CustomOpIr` impl), so a cached graph replays with fresh scalar values. The
44    /// remote backend relies on these to ship a custom op's scalar arguments to the server, where
45    /// the registered handler reads them back with [`ScalarIr::elem`].
46    pub scalars: Vec<ScalarIr>,
47}
48
49impl CustomOpIr {
50    /// Create a new custom operation intermediate representation (without scalar arguments).
51    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    /// Create a new custom operation intermediate representation with scalar arguments.
61    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    /// Cast the intermediate representation, and get the in and output tensors.
76    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/// Describe all tensor operations possible.
111#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
112#[allow(clippy::large_enum_variant)]
113pub enum OperationIr {
114    /// Basic operation on a float tensor.
115    BaseFloat(BaseOperationIr),
116    /// Basic operation on an int tensor.
117    BaseInt(BaseOperationIr),
118    /// Basic operation on a bool tensor.
119    BaseBool(BaseOperationIr),
120    /// Numeric operation on a float tensor.
121    NumericFloat(DType, NumericOperationIr),
122    /// Numeric operation on an int tensor.
123    NumericInt(DType, NumericOperationIr),
124    /// Operation specific to a bool tensor.
125    Bool(BoolOperationIr),
126    /// Operation specific to an int tensor.
127    Int(IntOperationIr),
128    /// Operation specific to a float tensor.
129    Float(DType, FloatOperationIr),
130    /// Module operation.
131    Module(ModuleOperationIr),
132    /// Initialize operation.
133    Init(InitOperationIr),
134    /// A custom operation.
135    Custom(CustomOpIr),
136    /// A tensor is dropped.
137    Drop(TensorIr),
138    /// Operation specific to a distributed tensor.
139    Distributed(DistributedOperationIr),
140    /// Activation function operation (relu, gelu, sigmoid, softmax, …).
141    Activation(ActivationOperationIr),
142}
143
144/// Operation intermediate representation specific to a float tensor.
145#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
146pub enum FloatOperationIr {
147    /// Operation corresponding to [exp](burn_backend::ops::FloatTensorOps::float_exp).
148    Exp(UnaryOpIr),
149    /// Operation corresponding to [log](burn_backend::ops::FloatTensorOps::float_log).
150    Log(UnaryOpIr),
151    /// Operation corresponding to [log1p](burn_backend::ops::FloatTensorOps::float_log1p).
152    Log1p(UnaryOpIr),
153    /// Operation corresponding to [erf](burn_backend::ops::FloatTensorOps::float_erf).
154    Erf(UnaryOpIr),
155    /// Operation corresponding to [powf_scalar](burn_backend::ops::FloatTensorOps::float_powf_scalar).
156    PowfScalar(ScalarOpIr),
157    /// Operation corresponding to [sqrt](burn_backend::ops::FloatTensorOps::float_sqrt).
158    Sqrt(UnaryOpIr),
159    /// Operation corresponding to [cos](burn_backend::ops::FloatTensorOps::float_cos).
160    Cos(UnaryOpIr),
161    /// Operation corresponding to [cosh](burn_backend::ops::FloatTensorOps::float_cosh).
162    Cosh(UnaryOpIr),
163    /// Operation corresponding to [sin](burn_backend::ops::FloatTensorOps::float_sin).
164    Sin(UnaryOpIr),
165    /// Operation corresponding to [sin](burn_backend::ops::FloatTensorOps::float_sinh).
166    Sinh(UnaryOpIr),
167    /// Operation corresponding to [tan](burn_backend::ops::FloatTensorOps::float_tan).
168    Tan(UnaryOpIr),
169    /// Operation corresponding to [tanh](burn_backend::ops::FloatTensorOps::float_tanh).
170    Tanh(UnaryOpIr),
171    /// Operation corresponding to [acos](burn_backend::ops::FloatTensorOps::float_acos).
172    ArcCos(UnaryOpIr),
173    /// Operation corresponding to [acosh](burn_backend::ops::FloatTensorOps::float_acosh).
174    ArcCosh(UnaryOpIr),
175    /// Operation corresponding to [asin](burn_backend::ops::FloatTensorOps::float_asin).
176    ArcSin(UnaryOpIr),
177    /// Operation corresponding to [asinh](burn_backend::ops::FloatTensorOps::float_asinh).
178    ArcSinh(UnaryOpIr),
179    /// Operation corresponding to [atan](burn_backend::ops::FloatTensorOps::float_atan).
180    ArcTan(UnaryOpIr),
181    /// Operation corresponding to [atanh](burn_backend::ops::FloatTensorOps::float_atanh).
182    ArcTanh(UnaryOpIr),
183    /// Operation corresponding to [atan2](burn_backend::ops::FloatTensorOps::float_atan2).
184    ArcTan2(BinaryOpIr),
185    /// Operation corresponding to [round](burn_backend::ops::FloatTensorOps::float_round).
186    Round(UnaryOpIr),
187    /// Operation corresponding to [floor](burn_backend::ops::FloatTensorOps::float_floor).
188    Floor(UnaryOpIr),
189    /// Operation corresponding to [ceil](burn_backend::ops::FloatTensorOps::float_ceil).
190    Ceil(UnaryOpIr),
191    /// Operation corresponding to [trunc](burn_backend::ops::FloatTensorOps::float_trunc).
192    Trunc(UnaryOpIr),
193    /// Operation corresponding to [into_int](burn_backend::ops::FloatTensorOps::float_into_int).
194    IntoInt(CastOpIr),
195    /// Operation corresponding to [matmul](burn_backend::ops::FloatTensorOps::float_matmul).
196    Matmul(MatmulOpIr),
197    /// Operation corresponding to [cross](burn_backend::ops::FloatTensorOps::float_cross).
198    Cross(CrossOpIr),
199    /// Operation corresponding to [random](burn_backend::ops::FloatTensorOps::float_random).
200    Random(RandomOpIr),
201    /// Operation corresponding to [recip](burn_backend::ops::FloatTensorOps::float_recip).
202    Recip(UnaryOpIr),
203    /// Operation corresponding to [is_nan](burn_backend::ops::FloatTensorOps::float_is_nan).
204    IsNan(UnaryOpIr),
205    /// Operation corresponding to [is_nan](burn_backend::ops::FloatTensorOps::float_is_inf).
206    IsInf(UnaryOpIr),
207    /// Operation corresponding to [quantize](burn_backend::ops::QTensorOps::quantize).
208    Quantize(QuantizeOpIr),
209    /// Operation corresponding to [dequantize](burn_backend::ops::QTensorOps::dequantize).
210    Dequantize(DequantizeOpIr),
211    /// Operation corresponding to [grid_sample_2d](burn_backend::ops::FloatTensorOps::float_grid_sample_2d).
212    GridSample2d(GridSample2dOpIr),
213    /// Operation corresponding to [powf](burn_backend::ops::FloatTensorOps::float_powi).
214    Powf(BinaryOpIr),
215    /// Operation corresponding to [hypot](burn_backend::ops::FloatTensorOps::float_hypot).
216    Hypot(BinaryOpIr),
217}
218
219/// Operation intermediate representation specific to module.
220#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
221pub enum ModuleOperationIr {
222    /// Operation corresponding to [embedding](burn_backend::ops::ModuleOps::embedding).
223    Embedding(EmbeddingOpIr),
224    /// Operation corresponding to [embedding_backward](burn_backend::ops::ModuleOps::embedding_backward).
225    EmbeddingBackward(EmbeddingBackwardOpIr),
226    /// Operation corresponding to [linear](burn_backend::ops::ModuleOps::linear).
227    Linear(LinearOpIr),
228    /// Operation corresponding to [linear_x_backward](burn_backend::ops::ModuleOps::linear_x_backward).
229    LinearXBackward(LinearXBackwardOpIr),
230    /// Operation corresponding to [linear_weight_backward](burn_backend::ops::ModuleOps::linear_weight_backward).
231    LinearWeightBackward(LinearWeightBackwardOpIr),
232    /// Operation corresponding to [linear_bias_backward](burn_backend::ops::ModuleOps::linear_bias_backward).
233    LinearBiasBackward(LinearBiasBackwardOpIr),
234    /// Operation corresponding to [conv1d](burn_backend::ops::ModuleOps::conv1d).
235    Conv1d(Conv1dOpIr),
236    /// Operation corresponding to [conv1d_x_backward](burn_backend::ops::ModuleOps::conv1d_x_backward).
237    Conv1dXBackward(Conv1dXBackwardOpIr),
238    /// Operation corresponding to [conv1d_weight_backward](burn_backend::ops::ModuleOps::conv1d_weight_backward).
239    Conv1dWeightBackward(Conv1dWeightBackwardOpIr),
240    /// Operation corresponding to [conv1d_bias_backward](burn_backend::ops::ModuleOps::conv1d_bias_backward).
241    Conv1dBiasBackward(Conv1dBiasBackwardOpIr),
242    /// Operation corresponding to [conv2d](burn_backend::ops::ModuleOps::conv2d).
243    Conv2d(Conv2dOpIr),
244    /// Operation corresponding to [conv2d_x_backward](burn_backend::ops::ModuleOps::conv2d_x_backward).
245    Conv2dXBackward(Conv2dXBackwardOpIr),
246    /// Operation corresponding to [conv2d_weight_backward](burn_backend::ops::ModuleOps::conv2d_weight_backward).
247    Conv2dWeightBackward(Conv2dWeightBackwardOpIr),
248    /// Operation corresponding to [conv2d_bias_backward](burn_backend::ops::ModuleOps::conv2d_bias_backward).
249    Conv2dBiasBackward(Conv2dBiasBackwardOpIr),
250    /// Operation corresponding to [conv3d](burn_backend::ops::ModuleOps::conv3d).
251    Conv3d(Conv3dOpIr),
252    /// Operation corresponding to [conv3d_x_backward](burn_backend::ops::ModuleOps::conv3d_x_backward).
253    Conv3dXBackward(Conv3dXBackwardOpIr),
254    /// Operation corresponding to [conv3d_weight_backward](burn_backend::ops::ModuleOps::conv3d_weight_backward).
255    Conv3dWeightBackward(Conv3dWeightBackwardOpIr),
256    /// Operation corresponding to [conv3d_bias_backward](burn_backend::ops::ModuleOps::conv3d_bias_backward).
257    Conv3dBiasBackward(Conv3dBiasBackwardOpIr),
258    /// Operation corresponding to [deform_conv2d](burn_backend::ops::ModuleOps::deform_conv2d)
259    DeformableConv2d(Box<DeformConv2dOpIr>),
260    /// Operation corresponding to [deform_conv2d_backward](burn_backend::ops::ModuleOps::deform_conv2d_backward)
261    DeformableConv2dBackward(Box<DeformConv2dBackwardOpIr>),
262    /// Operation corresponding to [conv transpose 1d](burn_backend::ops::ModuleOps::conv_transpose1d).
263    ConvTranspose1d(ConvTranspose1dOpIr),
264    /// Operation corresponding to [conv transpose 2d](burn_backend::ops::ModuleOps::conv_transpose2d).
265    ConvTranspose2d(ConvTranspose2dOpIr),
266    /// Operation corresponding to [conv transpose 3d](burn_backend::ops::ModuleOps::conv_transpose3d).
267    ConvTranspose3d(ConvTranspose3dOpIr),
268    /// Operation corresponding to [avg pool 1d](burn_backend::ops::ModuleOps::avg_pool1d).
269    AvgPool1d(AvgPool1dOpIr),
270    /// Operation corresponding to [avg pool 2d](burn_backend::ops::ModuleOps::avg_pool2d).
271    AvgPool2d(AvgPool2dOpIr),
272    /// Operation corresponding to
273    /// [avg pool 1d backward](burn_backend::ops::ModuleOps::avg_pool1d_backward).
274    AvgPool1dBackward(AvgPool1dBackwardOpIr),
275    /// Operation corresponding to
276    /// [avg pool 2d backward](burn_backend::ops::ModuleOps::avg_pool2d_backward).
277    AvgPool2dBackward(AvgPool2dBackwardOpIr),
278    /// Operation corresponding to
279    /// [adaptive avg pool 1d](burn_backend::ops::ModuleOps::adaptive_avg_pool1d).
280    AdaptiveAvgPool1d(AdaptiveAvgPool1dOpIr),
281    /// Operation corresponding to
282    /// [adaptive avg pool 2d](burn_backend::ops::ModuleOps::adaptive_avg_pool2d).
283    AdaptiveAvgPool2d(AdaptiveAvgPool2dOpIr),
284    /// Operation corresponding to
285    /// [adaptive avg pool 1d backward](burn_backend::ops::ModuleOps::adaptive_avg_pool1d_backward).
286    AdaptiveAvgPool1dBackward(AdaptiveAvgPool1dBackwardOpIr),
287    /// Operation corresponding to
288    /// [adaptive avg pool 2d backward](burn_backend::ops::ModuleOps::adaptive_avg_pool2d_backward).
289    AdaptiveAvgPool2dBackward(AdaptiveAvgPool2dBackwardOpIr),
290    /// Operation corresponding to
291    /// [adaptive avg pool 3d](burn_backend::ops::ModuleOps::adaptive_avg_pool3d).
292    AdaptiveAvgPool3d(AdaptiveAvgPool3dOpIr),
293    /// Operation corresponding to
294    /// [adaptive avg pool 3d backward](burn_backend::ops::ModuleOps::adaptive_avg_pool3d_backward).
295    AdaptiveAvgPool3dBackward(AdaptiveAvgPool3dBackwardOpIr),
296    /// Operation corresponding to
297    /// [max pool 1d](burn_backend::ops::ModuleOps::max_pool1d).
298    MaxPool1d(MaxPool1dOpIr),
299    /// Operation corresponding to
300    /// [max pool 1d with indices](burn_backend::ops::ModuleOps::max_pool1d_with_indices).
301    MaxPool1dWithIndices(MaxPool1dWithIndicesOpIr),
302    /// Operation corresponding to
303    /// [max pool 1d with indices backward](burn_backend::ops::ModuleOps::max_pool1d_with_indices_backward).
304    MaxPool1dWithIndicesBackward(MaxPool1dWithIndicesBackwardOpIr),
305    /// Operation corresponding to
306    /// [max pool 2d](burn_backend::ops::ModuleOps::max_pool1d).
307    MaxPool2d(MaxPool2dOpIr),
308    /// Operation corresponding to
309    /// [max pool 2d with indices](burn_backend::ops::ModuleOps::max_pool2d_with_indices).
310    MaxPool2dWithIndices(MaxPool2dWithIndicesOpIr),
311    /// Operation corresponding to
312    /// [max pool 2d with indices backward](burn_backend::ops::ModuleOps::max_pool2d_with_indices_backward).
313    MaxPool2dWithIndicesBackward(MaxPool2dWithIndicesBackwardOpIr),
314    /// Operation corresponding to [interpolate](burn_backend::ops::ModuleOps::interpolate).
315    Interpolate(InterpolateOpIr),
316    /// Operation corresponding to [interpolate backward](burn_backend::ops::ModuleOps::interpolate_backward).
317    InterpolateBackward(InterpolateBackwardOpIr),
318    /// Operation corresponding to [rfft](burn_backend::ops::ModuleOps::rfft)
319    Rfft(RfftOpIr),
320    /// Operation corresponding to [irfft](burn_backend::ops::ModuleOps::irfft)
321    IRfft(IRfftOpIr),
322    /// Operation corresponding to [attention](burn_backend::ops::ModuleOps::attention).
323    Attention(AttentionOpIr),
324    /// Operation corresponding to [ctc_loss](burn_backend::ops::ModuleOps::ctc_loss).
325    CtcLoss(CtcLossOpIr),
326    /// Operation corresponding to
327    /// [ctc_loss_backward](burn_backend::ops::ModuleOps::ctc_loss_backward).
328    CtcLossBackward(CtcLossBackwardOpIr),
329    /// Operation corresponding to [layer_norm](burn_backend::ops::ModuleOps::layer_norm).
330    LayerNorm(LayerNormOpIr),
331    /// Operation corresponding to [unfold4d](burn_backend::ops::ModuleOps::unfold4d).
332    Unfold4d(Unfold4dOpIr),
333    /// Operation corresponding to
334    /// [conv_transpose1d_weight_backward](burn_backend::ops::ModuleOps::conv_transpose1d_weight_backward).
335    ConvTranspose1dWeightBackward(ConvTranspose1dWeightBackwardOpIr),
336    /// Operation corresponding to
337    /// [conv_transpose1d_bias_backward](burn_backend::ops::ModuleOps::conv_transpose1d_bias_backward).
338    ConvTranspose1dBiasBackward(ConvTranspose1dBiasBackwardOpIr),
339    /// Operation corresponding to
340    /// [conv_transpose2d_weight_backward](burn_backend::ops::ModuleOps::conv_transpose2d_weight_backward).
341    ConvTranspose2dWeightBackward(ConvTranspose2dWeightBackwardOpIr),
342    /// Operation corresponding to
343    /// [conv_transpose2d_bias_backward](burn_backend::ops::ModuleOps::conv_transpose2d_bias_backward).
344    ConvTranspose2dBiasBackward(ConvTranspose2dBiasBackwardOpIr),
345    /// Operation corresponding to
346    /// [conv_transpose3d_weight_backward](burn_backend::ops::ModuleOps::conv_transpose3d_weight_backward).
347    ConvTranspose3dWeightBackward(ConvTranspose3dWeightBackwardOpIr),
348    /// Operation corresponding to
349    /// [conv_transpose3d_bias_backward](burn_backend::ops::ModuleOps::conv_transpose3d_bias_backward).
350    ConvTranspose3dBiasBackward(ConvTranspose3dBiasBackwardOpIr),
351}
352
353/// Basic operations that can be done on any tensor type.
354#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
355pub enum BaseOperationIr {
356    /// Operation corresponding to:
357    ///
358    /// Float => [reshape](burn_backend::ops::FloatTensorOps::float_reshape).
359    /// Int => [reshape](burn_backend::ops::IntTensorOps::int_reshape).
360    /// Bool => [reshape](burn_backend::ops::BoolTensorOps::bool_reshape).
361    Reshape(ShapeOpIr),
362
363    /// Operation corresponding to:
364    ///
365    /// Float => [swap_dims](burn_backend::ops::FloatTensorOps::float_swap_dims).
366    /// Int => [swap_dims](burn_backend::ops::IntTensorOps::int_swap_dims).
367    /// Bool => [swap_dims](burn_backend::ops::BoolTensorOps::bool_swap_dims).
368    SwapDims(SwapDimsOpIr),
369
370    /// Operation corresponding to:
371    ///
372    /// Float => [permute](burn_backend::ops::FloatTensorOps::float_permute).
373    /// Int => [permute](burn_backend::ops::IntTensorOps::int_permute).
374    /// Bool => [permute](burn_backend::ops::BoolTensorOps::bool_permute).
375    Permute(PermuteOpIr),
376
377    /// Operation corresponding to:
378    /// Float => [flip](burn_backend::ops::FloatTensorOps::float_flip).
379    /// Int => [flip](burn_backend::ops::IntTensorOps::int_flip).
380    /// Bool => [flip](burn_backend::ops::BoolTensorOps::bool_flip).
381    Flip(FlipOpIr),
382
383    /// Operation corresponding to:
384    ///
385    /// Float => [expand](burn_backend::ops::FloatTensorOps::float_expand).
386    /// Int => [expand](burn_backend::ops::IntTensorOps::int_expand).
387    /// Bool => [expand](burn_backend::ops::BoolTensorOps::bool_expand).
388    Expand(ShapeOpIr),
389
390    /// Unfold windows along an axis.
391    ///
392    Unfold(UnfoldOpIr),
393
394    /// Operation corresponding to:
395    ///
396    /// Float => [slice](burn_backend::ops::FloatTensorOps::float_slice).
397    /// Int => [slice](burn_backend::ops::IntTensorOps::int_slice).
398    /// Bool => [slice](burn_backend::ops::BoolTensorOps::bool_slice).
399    Slice(SliceOpIr),
400    /// Operation corresponding to:
401    ///
402    /// Float => [slice assign](burn_backend::ops::FloatTensorOps::float_slice_assign).
403    /// Int => [slice assign](burn_backend::ops::IntTensorOps::int_slice_assign).
404    /// Bool => [slice assign](burn_backend::ops::BoolTensorOps::bool_slice_assign).
405    SliceAssign(SliceAssignOpIr),
406    /// Operation corresponding to:
407    ///
408    /// Float => [select](burn_backend::ops::FloatTensorOps::float_select).
409    /// Int => [select](burn_backend::ops::IntTensorOps::int_select).
410    /// Bool => [select](burn_backend::ops::BoolTensorOps::bool_select).
411    Select(SelectOpIr),
412    /// Operation corresponding to:
413    ///
414    /// Float => [select assign](burn_backend::ops::FloatTensorOps::float_select_assign).
415    /// Int => [select assign](burn_backend::ops::IntTensorOps::int_select_assign).
416    /// Bool => [select assign](burn_backend::ops::BoolTensorOps::bool_select_or).
417    SelectAssign(SelectAssignOpIr),
418    /// Operation corresponding to:
419    ///
420    /// Float => [mask where](burn_backend::ops::FloatTensorOps::float_mask_where).
421    /// Int => [mask where](burn_backend::ops::IntTensorOps::int_mask_where).
422    /// Bool => [mask where](burn_backend::ops::BoolTensorOps::bool_mask_where).
423    MaskWhere(MaskWhereOpIr),
424    /// Operation corresponding to:
425    ///
426    /// Float => [mask fill](burn_backend::ops::FloatTensorOps::float_mask_fill).
427    /// Int => [mask fill](burn_backend::ops::IntTensorOps::int_mask_fill).
428    /// Bool => [mask fill](burn_backend::ops::BoolTensorOps::bool_mask_fill).
429    MaskFill(MaskFillOpIr),
430    /// Operation corresponding to:
431    ///
432    /// Float => [gather](burn_backend::ops::FloatTensorOps::float_gather).
433    /// Int => [gather](burn_backend::ops::IntTensorOps::int_gather).
434    /// Bool => [gather](burn_backend::ops::BoolTensorOps::bool_gather).
435    Gather(GatherOpIr),
436    /// Operation corresponding to:
437    ///
438    /// Float => [scatter](burn_backend::ops::FloatTensorOps::float_scatter_add).
439    /// Int => [scatter](burn_backend::ops::IntTensorOps::int_scatter_add).
440    /// Bool => [scatter](burn_backend::ops::BoolTensorOps::bool_scatter_or).
441    Scatter(ScatterOpIr),
442    /// Multi-dimensional scatter operation.
443    ScatterNd(ScatterNdOpIr),
444    /// Multi-dimensional gather operation.
445    GatherNd(GatherNdOpIr),
446    /// Operation corresponding to:
447    ///
448    /// Float => [equal](burn_backend::ops::FloatTensorOps::float_equal).
449    /// Int => [equal](burn_backend::ops::IntTensorOps::int_equal).
450    /// Bool => [equal](burn_backend::ops::BoolTensorOps::bool_equal).
451    Equal(BinaryOpIr),
452    /// Operation corresponding to:
453    ///
454    /// Float => [equal elem](burn_backend::ops::FloatTensorOps::float_equal_elem).
455    /// Int => [equal elem](burn_backend::ops::IntTensorOps::int_equal_elem).
456    /// Bool => [equal elem](burn_backend::ops::BoolTensorOps::bool_equal_elem).
457    EqualElem(ScalarOpIr),
458    /// Operation corresponding to:
459    ///
460    /// Float => [repeat dim](burn_backend::ops::FloatTensorOps::float_repeat_dim).
461    /// Int => [repeat dim](burn_backend::ops::IntTensorOps::int_repeat_dim).
462    /// Bool => [repeat dim](burn_backend::ops::BoolTensorOps::bool_repeat_dim).
463    RepeatDim(RepeatDimOpIr),
464    /// Operation corresponding to:
465    ///
466    /// Float => [cat](burn_backend::ops::FloatTensorOps::float_cat).
467    /// Int => [cat](burn_backend::ops::IntTensorOps::int_cat).
468    /// Bool => [cat](burn_backend::ops::BoolTensorOps::bool_cat).
469    Cat(CatOpIr),
470    /// Cast operation, no direct operation and should be supported by fusion backend.
471    Cast(CastOpIr),
472    /// Operation corresponding to:
473    ///
474    /// Float => [empty](burn_backend::ops::FloatTensorOps::float_empty).
475    /// Int => [empty](burn_backend::ops::IntTensorOps::int_empty).
476    /// Bool => [empty](burn_backend::ops::BoolTensorOps::bool_empty).
477    Empty(CreationOpIr),
478    /// Operation corresponding to:
479    ///
480    /// Float => [ones](burn_backend::ops::FloatTensorOps::float_ones).
481    /// Int => [ones](burn_backend::ops::IntTensorOps::int_ones).
482    /// Bool => [ones](burn_backend::ops::BoolTensorOps::bool_ones).
483    Ones(CreationOpIr),
484    /// Operation corresponding to:
485    ///
486    /// Float => [zeros](burn_backend::ops::FloatTensorOps::float_zeros).
487    /// Int => [zeros](burn_backend::ops::IntTensorOps::int_zeros).
488    /// Bool => [zeros](burn_backend::ops::BoolTensorOps::bool_zeros).
489    Zeros(CreationOpIr),
490    /// Operation corresponding to:
491    ///
492    /// Float => [not_equal](burn_backend::ops::FloatTensorOps::float_not_equal).
493    /// Int => [not_equal](burn_backend::ops::IntTensorOps::int_not_equal).
494    /// Bool => [not_equal](burn_backend::ops::BoolTensorOps::bool_not_equal).
495    NotEqual(BinaryOpIr),
496    /// Operation corresponding to:
497    ///
498    /// Float => [not_equal_elem](burn_backend::ops::FloatTensorOps::float_not_equal_elem).
499    /// Int => [not_equal_elem](burn_backend::ops::IntTensorOps::int_not_equal_elem).
500    /// Bool => [not_equal_elem](burn_backend::ops::BoolTensorOps::bool_not_equal_elem).
501    NotEqualElem(ScalarOpIr),
502    /// Reduce-`all` over the input tensor.
503    ///
504    /// Float/Int input is treated as a non-zero check; bool input is the value itself.
505    /// Output is a single-element bool tensor.
506    All(ReduceOpIr),
507    /// Reduce-`any` over the input tensor.
508    Any(ReduceOpIr),
509    /// Reduce-`all` along a dim.
510    AllDim(ReduceDimOpIr),
511    /// Reduce-`any` along a dim.
512    AnyDim(ReduceDimOpIr),
513}
514
515/// Numeric operations on int and float tensors.
516#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
517pub enum NumericOperationIr {
518    /// Operation corresponding to:
519    ///
520    /// Float => [add](burn_backend::ops::FloatTensorOps::float_add).
521    /// Int => [add](burn_backend::ops::IntTensorOps::int_add).
522    Add(BinaryOpIr),
523    /// Operation corresponding to:
524    ///
525    /// Float => [add scalar](burn_backend::ops::FloatTensorOps::float_add_scalar).
526    /// Int => [add scalar](burn_backend::ops::IntTensorOps::int_add_scalar).
527    AddScalar(ScalarOpIr),
528    /// Operation corresponding to:
529    ///
530    /// Float => [sub](burn_backend::ops::FloatTensorOps::float_sub).
531    /// Int => [sub](burn_backend::ops::IntTensorOps::int_sub).
532    Sub(BinaryOpIr),
533    /// Operation corresponding to:
534    ///
535    /// Float => [sub scalar](burn_backend::ops::FloatTensorOps::float_sub_scalar).
536    /// Int => [sub scalar](burn_backend::ops::IntTensorOps::int_sub_scalar).
537    SubScalar(ScalarOpIr),
538    /// Operation corresponding to:
539    ///
540    /// Float => [div](burn_backend::ops::FloatTensorOps::float_div).
541    /// Int => [div](burn_backend::ops::IntTensorOps::int_div).
542    Div(BinaryOpIr),
543    /// Operation corresponding to:
544    ///
545    /// Float => [div scalar](burn_backend::ops::FloatTensorOps::float_div_scalar).
546    /// Int => [div scalar](burn_backend::ops::IntTensorOps::int_div_scalar).
547    DivScalar(ScalarOpIr),
548    /// Operation corresponding to:
549    ///
550    /// Float => [rem](burn_backend::ops::FloatTensorOps::float_remainder).
551    /// Int => [rem](burn_backend::ops::IntTensorOps::int_remainder).
552    Rem(BinaryOpIr),
553    /// Operation corresponding to:
554    ///
555    /// Float => [rem scalar](burn_backend::ops::FloatTensorOps::float_remainder_scalar).
556    /// Int => [rem scalar](burn_backend::ops::IntTensorOps::int_remainder_scalar).
557    RemScalar(ScalarOpIr),
558    /// Operation corresponding to:
559    ///
560    /// Float => [mul](burn_backend::ops::FloatTensorOps::float_mul).
561    /// Int => [mul](burn_backend::ops::IntTensorOps::int_mul).
562    Mul(BinaryOpIr),
563    /// Operation corresponding to:
564    ///
565    /// Float => [mul scalar](burn_backend::ops::FloatTensorOps::float_mul_scalar).
566    /// Int => [mul scalar](burn_backend::ops::IntTensorOps::int_mul_scalar).
567    MulScalar(ScalarOpIr),
568    /// Operation corresponding to:
569    ///
570    /// Float => [abs](burn_backend::ops::FloatTensorOps::float_abs).
571    /// Int => [abs](burn_backend::ops::IntTensorOps::int_abs).
572    Abs(UnaryOpIr),
573    /// Operation corresponding to:
574    ///
575    /// Float => [full](burn_backend::ops::FloatTensorOps::float_full).
576    /// Int => [full](burn_backend::ops::IntTensorOps::int_full).
577    Full(FullOpIr),
578    /// Operation corresponding to:
579    ///
580    /// Float => [mean dim](burn_backend::ops::FloatTensorOps::float_mean_dim).
581    /// Int => [mean dim](burn_backend::ops::IntTensorOps::int_mean_dim).
582    MeanDim(ReduceDimOpIr),
583    /// Operation corresponding to:
584    ///
585    /// Float => [mean](burn_backend::ops::FloatTensorOps::float_mean).
586    /// Int => [mean](burn_backend::ops::IntTensorOps::int_mean).
587    Mean(ReduceOpIr),
588    /// Operation corresponding to:
589    ///
590    /// Float => [sum](burn_backend::ops::FloatTensorOps::float_sum).
591    /// Int => [sum](burn_backend::ops::IntTensorOps::int_sum).
592    Sum(ReduceOpIr),
593    /// Operation corresponding to:
594    ///
595    /// Float => [sum dim](burn_backend::ops::FloatTensorOps::float_sum_dim).
596    /// Int => [sum dim](burn_backend::ops::IntTensorOps::int_sum_dim).
597    SumDim(ReduceDimOpIr),
598    /// Operation corresponding to:
599    ///
600    /// Float => [prod](burn_backend::ops::FloatTensorOps::float_prod).
601    /// Int => [prod](burn_backend::ops::IntTensorOps::int_prod).
602    Prod(ReduceOpIr),
603    /// Operation corresponding to:
604    ///
605    /// Float => [prod dim](burn_backend::ops::FloatTensorOps::float_prod_dim).
606    /// Int => [prod dim](burn_backend::ops::IntTensorOps::int_prod_dim).
607    ProdDim(ReduceDimOpIr),
608    /// Operation corresponding to:
609    ///
610    /// Float => [greater](burn_backend::ops::FloatTensorOps::float_greater).
611    /// Int => [greater](burn_backend::ops::IntTensorOps::int_greater).
612    Greater(BinaryOpIr),
613    /// Operation corresponding to:
614    ///
615    /// Float => [greater elem](burn_backend::ops::FloatTensorOps::float_greater_elem).
616    /// Int => [greater elem](burn_backend::ops::IntTensorOps::int_greater_elem).
617    GreaterElem(ScalarOpIr),
618    /// Operation corresponding to:
619    ///
620    /// Float => [greater equal](burn_backend::ops::FloatTensorOps::float_greater_elem).
621    /// Int => [greater elem](burn_backend::ops::IntTensorOps::int_greater_elem).
622    GreaterEqual(BinaryOpIr),
623    /// Operation corresponding to:
624    ///
625    /// Float => [greater equal elem](burn_backend::ops::FloatTensorOps::float_greater_equal_elem).
626    /// Int => [greater equal elem](burn_backend::ops::IntTensorOps::int_greater_equal_elem).
627    GreaterEqualElem(ScalarOpIr),
628    /// Operation corresponding to:
629    ///
630    /// Float => [lower](burn_backend::ops::FloatTensorOps::float_lower).
631    /// Int => [lower](burn_backend::ops::IntTensorOps::int_lower).
632    Lower(BinaryOpIr),
633    /// Operation corresponding to:
634    ///
635    /// Float => [lower elem](burn_backend::ops::FloatTensorOps::float_lower_elem).
636    /// Int => [lower elem](burn_backend::ops::IntTensorOps::int_lower_elem).
637    LowerElem(ScalarOpIr),
638    /// Operation corresponding to:
639    ///
640    /// Float => [lower equal](burn_backend::ops::FloatTensorOps::float_lower_equal).
641    /// Int => [lower equal](burn_backend::ops::IntTensorOps::int_lower_equal).
642    LowerEqual(BinaryOpIr),
643    /// Operation corresponding to:
644    ///
645    /// Float => [lower equal elem](burn_backend::ops::FloatTensorOps::float_lower_equal_elem).
646    /// Int => [lower equal elem](burn_backend::ops::IntTensorOps::int_lower_equal_elem).
647    LowerEqualElem(ScalarOpIr),
648    /// Operation corresponding to:
649    ///
650    /// Float => [argmax](burn_backend::ops::FloatTensorOps::float_argmax).
651    /// Int => [argmax](burn_backend::ops::IntTensorOps::int_argmax).
652    ArgMax(ReduceDimOpIr),
653    /// Operation corresponding to:
654    ///
655    /// Float => [argtopk](burn_backend::ops::FloatTensorOps::float_argtopk).
656    /// Int => [argtopk](burn_backend::ops::IntTensorOps::int_argtopk).
657    ArgTopK(ReduceDimOpIr),
658    /// Operation corresponding to:
659    ///
660    /// Float => [topk](burn_backend::ops::FloatTensorOps::float_topk).
661    /// Int => [topk](burn_backend::ops::IntTensorOps::int_topk).
662    TopK(ReduceDimOpIr),
663    /// Operation corresponding to:
664    ///
665    /// Float => [topk with indices](burn_backend::ops::FloatTensorOps::float_topk_with_indices).
666    /// Int => [topk with indices](burn_backend::ops::IntTensorOps::int_topk_with_indices).
667    TopKWithIndices(TopKWithIndicesOpIr),
668    /// Operation corresponding to:
669    ///
670    /// Float => [argmin](burn_backend::ops::FloatTensorOps::float_argmin).
671    /// Int => [argmin](burn_backend::ops::IntTensorOps::int_argmin).
672    ArgMin(ReduceDimOpIr),
673    /// Operation corresponding to:
674    ///
675    /// Float => [max](burn_backend::ops::FloatTensorOps::float_max).
676    /// Int => [max](burn_backend::ops::IntTensorOps::int_max).
677    Max(ReduceOpIr),
678    /// Operation corresponding to:
679    ///
680    /// Float => [max dim with indices](burn_backend::ops::FloatTensorOps::float_max_dim_with_indices).
681    /// Int => [max dim with indices](burn_backend::ops::IntTensorOps::int_max_dim_with_indices).
682    MaxDimWithIndices(ReduceDimWithIndicesOpIr),
683    /// Operation corresponding to:
684    ///
685    /// Float => [min dim with indices](burn_backend::ops::FloatTensorOps::float_min_dim_with_indices).
686    /// Int => [min dim with indices](burn_backend::ops::IntTensorOps::int_min_dim_with_indices).
687    MinDimWithIndices(ReduceDimWithIndicesOpIr),
688    /// Operation corresponding to:
689    ///
690    /// Float => [min](burn_backend::ops::FloatTensorOps::float_min).
691    /// Int => [min](burn_backend::ops::IntTensorOps::int_min).
692    Min(ReduceOpIr),
693    /// Operation corresponding to:
694    ///
695    /// Float => [max dim](burn_backend::ops::FloatTensorOps::float_max_dim).
696    /// Int => [max dim](burn_backend::ops::IntTensorOps::int_max_dim).
697    MaxDim(ReduceDimOpIr),
698    /// Operation corresponding to:
699    ///
700    /// Float => [min dim](burn_backend::ops::FloatTensorOps::float_min_dim).
701    /// Int => [min dim](burn_backend::ops::IntTensorOps::int_min_dim).
702    MinDim(ReduceDimOpIr),
703    /// Operation corresponding to:
704    ///
705    /// Float => [max_abs](burn_backend::ops::FloatTensorOps::float_max_abs).
706    /// Int => [max_abs](burn_backend::ops::IntTensorOps::int_max_abs).
707    MaxAbs(ReduceOpIr),
708    /// Operation corresponding to:
709    ///
710    /// Float => [max_abs dim](burn_backend::ops::FloatTensorOps::float_max_abs_dim).
711    /// Int => [max_abs dim](burn_backend::ops::IntTensorOps::int_max_abs_dim).
712    MaxAbsDim(ReduceDimOpIr),
713    /// Operation corresponding to:
714    ///
715    /// Float => [clamp](burn_backend::ops::FloatTensorOps::float_clamp).
716    /// Int => [clamp](burn_backend::ops::IntTensorOps::int_clamp).
717    Clamp(ClampOpIr),
718    /// Operation corresponding to:
719    ///
720    /// Int => [random](burn_backend::ops::IntTensorOps::int_random).
721    IntRandom(RandomOpIr),
722    /// Operation corresponding to:
723    ///
724    /// Float => [powf](burn_backend::ops::FloatTensorOps::float_powi).
725    /// Int => [powf](burn_backend::ops::IntTensorOps::int_powi).
726    Powi(BinaryOpIr),
727    /// Operation corresponding to:
728    ///
729    /// Float => [powi_scalar](burn_backend::ops::FloatTensorOps::float_powi_scalar).
730    /// Int => [powi_scalar](burn_backend::ops::IntTensorOps::int_powi_scalar).
731    PowiScalar(ScalarOpIr),
732    /// Operation corresponding to:
733    ///
734    /// Float => [cumsum](burn_backend::ops::FloatTensorOps::float_cumsum).
735    /// Int => [cumsum](burn_backend::ops::IntTensorOps::int_cumsum).
736    CumSum(DimOpIr),
737    /// Operation corresponding to:
738    ///
739    /// Float => [cumprod](burn_backend::ops::FloatTensorOps::float_cumprod).
740    /// Int => [cumprod](burn_backend::ops::IntTensorOps::int_cumprod).
741    CumProd(DimOpIr),
742    /// Operation corresponding to:
743    ///
744    /// Float => [cummin](burn_backend::ops::FloatTensorOps::float_cummin).
745    /// Int => [cummin](burn_backend::ops::IntTensorOps::int_cummin).
746    CumMin(DimOpIr),
747    /// Operation corresponding to:
748    ///
749    /// Float => [cummax](burn_backend::ops::FloatTensorOps::float_cummax).
750    /// Int => [cummax](burn_backend::ops::IntTensorOps::int_cummax).
751    CumMax(DimOpIr),
752    /// Operation corresponding to:
753    ///
754    /// Float => [neg](burn_backend::ops::FloatTensorOps::float_neg).
755    /// Int => [neg](burn_backend::ops::IntTensorOps::int_neg).
756    Neg(UnaryOpIr),
757    /// Operation corresponding to:
758    ///
759    /// Float => [sign](burn_backend::ops::FloatTensorOps::float_sign).
760    /// Int => [sign](burn_backend::ops::IntTensorOps::int_sign).
761    Sign(UnaryOpIr),
762    /// Operation corresponding to:
763    ///
764    /// Float => [clamp_min](burn_backend::ops::FloatTensorOps::float_clamp_min).
765    /// Int => [clamp_min](burn_backend::ops::IntTensorOps::int_clamp_min).
766    ClampMin(ScalarOpIr),
767    /// Operation corresponding to:
768    ///
769    /// Float => [clamp_max](burn_backend::ops::FloatTensorOps::float_clamp_max).
770    /// Int => [clamp_max](burn_backend::ops::IntTensorOps::int_clamp_max).
771    ClampMax(ScalarOpIr),
772    /// Sort along a dim. Output shares the input's shape and dtype.
773    Sort(SortOpIr),
774    /// Sort along a dim, also returning the source indices.
775    SortWithIndices(SortWithIndicesOpIr),
776    /// Sort along a dim and return only the indices (i.e., the argsort).
777    ///
778    /// Shares [`SortOpIr`] with [`Sort`](Self::Sort); only the output dtype differs.
779    ArgSort(SortOpIr),
780}
781
782/// Operation intermediate representation specific to an int tensor.
783#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
784pub enum IntOperationIr {
785    /// Operation corresponding to [into float](burn_backend::ops::IntTensorOps::int_into_float).
786    IntoFloat(CastOpIr),
787    /// Operation corresponding to:
788    ///
789    /// Int => [bitwise and](burn_backend::ops::IntTensorOps::bitwise_and).
790    BitwiseAnd(BinaryOpIr),
791    /// Operation corresponding to:
792    ///
793    /// Int => [bitwise and scalar](burn_backend::ops::IntTensorOps::bitwise_and_scalar).
794    BitwiseAndScalar(ScalarOpIr),
795    /// Operation corresponding to:
796    ///
797    /// Int => [bitwise or](burn_backend::ops::IntTensorOps::bitwise_or).
798    BitwiseOr(BinaryOpIr),
799    /// Operation corresponding to:
800    ///
801    /// Int => [bitwise or scalar](burn_backend::ops::IntTensorOps::bitwise_or_scalar).
802    BitwiseOrScalar(ScalarOpIr),
803    /// Operation corresponding to:
804    ///
805    /// Int => [bitwise xor](burn_backend::ops::IntTensorOps::bitwise_xor).
806    BitwiseXor(BinaryOpIr),
807    /// Operation corresponding to:
808    ///
809    /// Int => [bitwise xor scalar](burn_backend::ops::IntTensorOps::bitwise_xor_scalar).
810    BitwiseXorScalar(ScalarOpIr),
811    /// Operation corresponding to:
812    ///
813    /// Int => [bitwise not](burn_backend::ops::IntTensorOps::bitwise_not).
814    BitwiseNot(UnaryOpIr),
815    /// Operation corresponding to:
816    ///
817    /// Int => [bitwise left shift](burn_backend::ops::IntTensorOps::bitwise_left_shift).
818    BitwiseLeftShift(BinaryOpIr),
819    /// Operation corresponding to:
820    ///
821    /// Int => [bitwise left shift scalar](burn_backend::ops::IntTensorOps::bitwise_left_shift_scalar).
822    BitwiseLeftShiftScalar(ScalarOpIr),
823    /// Operation corresponding to:
824    ///
825    /// Int => [bitwise right shift](burn_backend::ops::IntTensorOps::bitwise_right_shift).
826    BitwiseRightShift(BinaryOpIr),
827    /// Operation corresponding to:
828    ///
829    /// Int => [bitwise right shift scalar](burn_backend::ops::IntTensorOps::bitwise_right_shift_scalar).
830    BitwiseRightShiftScalar(ScalarOpIr),
831    /// Operation corresponding to [matmul](burn_backend::ops::IntTensorOps::int_matmul).
832    Matmul(MatmulOpIr),
833}
834
835/// Operation intermediate representation specific to a bool tensor.
836#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
837pub enum BoolOperationIr {
838    /// Operation corresponding to [into float](burn_backend::ops::BoolTensorOps::bool_into_float).
839    IntoFloat(CastOpIr),
840    /// Operation corresponding to [into int](burn_backend::ops::BoolTensorOps::bool_into_int).
841    IntoInt(CastOpIr),
842    /// Operation corresponding to [not](burn_backend::ops::BoolTensorOps::bool_not).
843    Not(UnaryOpIr),
844    /// Operation corresponding to [and](burn_backend::ops::BoolTensorOps::bool_and).
845    And(BinaryOpIr),
846    /// Operation corresponding to [or](burn_backend::ops::BoolTensorOps::bool_or).
847    Or(BinaryOpIr),
848    /// Operation corresponding to [xor](burn_backend::ops::BoolTensorOps::bool_xor).
849    Xor(BinaryOpIr),
850}
851
852/// Operations that can be done on distributed tensors.
853#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
854#[allow(clippy::large_enum_variant)]
855pub enum DistributedOperationIr {
856    /// Operation corresponding to:
857    /// [all_reduce](burn_backend::distributed::DistributedOps::all_reduce).
858    AllReduce(AllReduceOpIr),
859    /// Resolve the pending collective operations on the executing device. Corresponds to
860    /// [sync_collective](burn_backend::distributed::DistributedOps::sync_collective).
861    ///
862    /// Fire-and-forget and payload-free: it syncs whichever device the interpreter is bound to.
863    /// Modeled as an operation (not a side-channel call) so it travels the normal op stream
864    /// alongside [`AllReduce`](DistributedOperationIr::AllReduce) and is ordered against it.
865    SyncCollective,
866}
867
868/// Swap dim operation intermediate representation.
869#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
870pub struct SwapDimsOpIr {
871    /// Input tensor intermediate representation.
872    pub input: TensorIr,
873    /// Output tensor intermediate representation.
874    pub out: TensorIr,
875    /// The first dim to swap.
876    pub dim1: usize,
877    /// The second dim to swap.
878    pub dim2: usize,
879}
880
881/// Permute operation intermediate representation.
882#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
883pub struct PermuteOpIr {
884    /// Input tensor intermediate representation.
885    pub input: TensorIr,
886    /// Output tensor intermediate representation.
887    pub out: TensorIr,
888    /// The new order of the dimensions.
889    pub axes: Vec<usize>,
890}
891
892/// Shape operation intermediate representation.
893#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
894pub struct ShapeOpIr {
895    /// Input tensor intermediate representation.
896    pub input: TensorIr,
897    /// Output tensor intermediate representation with the new shape.
898    pub out: TensorIr,
899}
900
901/// Unfold operation intermediate representation.
902#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
903pub struct UnfoldOpIr {
904    /// Input tensor intermediate representation.
905    pub input: TensorIr,
906    /// Output tensor intermediate representation.
907    pub out: TensorIr,
908
909    /// The selected dim.
910    pub dim: usize,
911    /// The window size.
912    pub size: usize,
913    /// The window step along dim.
914    pub step: usize,
915}
916
917/// Flip operation intermediate representation.
918#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
919pub struct FlipOpIr {
920    /// Input tensor intermediate representation.
921    pub input: TensorIr,
922    /// Output tensor intermediate representation.
923    pub out: TensorIr,
924    /// The dimensions to flip.
925    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/// Creation operation intermediate representation.
936/// As opposed to [InitOperationIr], creation operations are lazy initialized.
937#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
938pub struct CreationOpIr {
939    /// Output tensor intermediate representation.
940    pub out: TensorIr,
941}
942
943/// Full operation intermediate representation.
944#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
945pub struct FullOpIr {
946    /// Output tensor intermediate representation.
947    pub out: TensorIr,
948    /// Fill value.
949    pub value: ScalarIr,
950}
951
952#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
953/// Declares a tensor has been initialized.
954///
955/// It is necessary to register for proper orphan detection and avoid memory leak.
956pub struct InitOperationIr {
957    /// The initialized tensor.
958    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    // TODO: Make that an enum with `Value` and `Id` variants for relative/global
998    // conversion.
999    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/// IR for operations that operate along a dimension without reducing it.
1027/// Unlike `ReduceDimOpIr`, the output shape is the same as the input shape.
1028#[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    /// How to reduce the values across the participating devices.
1161    pub op: burn_backend::distributed::ReduceOperation,
1162    /// The devices participating in the collective operation.
1163    pub device_ids: Vec<DeviceIdIr>,
1164}
1165
1166/// Serializable representation of a [device id](burn_backend::DeviceId).
1167///
1168/// The intermediate representation is part of the wire protocol (e.g. the remote backend), so it
1169/// cannot store `burn_backend::DeviceId` directly since that type is not serializable.
1170#[derive(Clone, Copy, Debug, Hash, PartialEq, Eq, Serialize, Deserialize)]
1171pub struct DeviceIdIr {
1172    /// Identifies the type of the device.
1173    pub type_id: u16,
1174    /// Identifies the device number.
1175    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)]
1204/// Like [`ReduceDimWithIndicesOpIr`], but for a top-k: the reduced axis keeps `k` entries
1205/// instead of collapsing to 1, so `k` has to be carried explicitly.
1206pub 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/// Quantization parameters intermediate representation.
1507#[derive(Debug, Clone, Hash, PartialEq, Eq, Serialize, Deserialize)]
1508pub struct QuantizationParametersIr {
1509    /// The scaling factor.
1510    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        // `n` is required to be a power of two at the public API boundary, so
1920        // the output has `n / 2 + 1` bins (matching scipy/torch for pow2 n).
1921        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    /// Get all input [tensors](TensorIr) involved with the current operation.
2149    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    /// Get all output [tensors](TensorIr) involved with the current operation.
2169    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    /// Get all [tensor](TensorIr) involved with the current operation.
2189    pub fn nodes(&self) -> Vec<&TensorIr> {
2190        self.inputs().chain(self.outputs()).collect()
2191    }
2192
2193    /// Set the given nodes that are [read write](super::TensorStatus::ReadWrite) to
2194    /// [read only](super::TensorStatus::ReadOnly) in the current operation.
2195    ///
2196    /// Returns the tensor that were updated with their original representation.
2197    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    /// Visit every mutable component of this operation in place.
2229    ///
2230    /// For each variant, visits its [tensors](TensorIr) (in the same order as
2231    /// [`inputs`](Self::inputs) then [`outputs`](Self::outputs) combined), then its
2232    /// [scalars](ScalarIr), then its slice ranges. Only `Slice`/`SliceAssign` (under the `Base*`
2233    /// variants) carry ranges; every other category visits no ranges.
2234    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
4779/// Extension trait to extract outputs when registering an operation.
4780pub trait OperationOutput<O> {
4781    /// Extract a single output.
4782    fn output(self) -> O;
4783
4784    /// Extract a fixed number of outputs.
4785    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/// Operation IR for sort along a dim. The output preserves the input shape/dtype.
4800#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4801pub struct SortOpIr {
4802    /// Input tensor.
4803    pub input: TensorIr,
4804    /// Dim along which to sort.
4805    pub dim: usize,
4806    /// Sort descending.
4807    pub descending: bool,
4808    /// Output tensor (same shape/dtype as input).
4809    pub out: TensorIr,
4810}
4811
4812/// Operation IR for sort-with-indices: returns sorted values + source indices.
4813#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4814pub struct SortWithIndicesOpIr {
4815    /// Input tensor.
4816    pub input: TensorIr,
4817    /// Dim along which to sort.
4818    pub dim: usize,
4819    /// Sort descending.
4820    pub descending: bool,
4821    /// Output tensor with sorted values.
4822    pub out: TensorIr,
4823    /// Output tensor with the indices into the original input.
4824    pub out_indices: TensorIr,
4825}
4826
4827/// Operation IR for layer normalization with optional bias.
4828#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4829pub struct LayerNormOpIr {
4830    /// Input tensor.
4831    pub input: TensorIr,
4832    /// Scale (gamma) parameter.
4833    pub gamma: TensorIr,
4834    /// Optional shift (beta) parameter.
4835    pub beta: Option<TensorIr>,
4836    /// Numerical-stability epsilon.
4837    pub epsilon: ScalarIr,
4838    /// Output tensor.
4839    pub out: TensorIr,
4840}
4841
4842/// Operation IR for unfold4d (a 4d sliding-window kernel-as-explicit-columns reshape).
4843#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4844pub struct Unfold4dOpIr {
4845    /// Input tensor of shape `[N, C, H, W]`.
4846    pub x: TensorIr,
4847    /// Kernel size `[kH, kW]`.
4848    pub kernel_size: [usize; 2],
4849    /// Conv-like options (stride, padding, dilation).
4850    pub options: Unfold4dOptionsIr,
4851    /// Output tensor.
4852    pub out: TensorIr,
4853}
4854
4855/// Options for [`Unfold4dOpIr`] — mirrors the backend's `UnfoldOptions`.
4856#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4857pub struct Unfold4dOptionsIr {
4858    /// Stride `[sH, sW]`.
4859    pub stride: [usize; 2],
4860    /// Padding `[pH, pW]`.
4861    pub padding: [usize; 2],
4862    /// Dilation `[dH, dW]`.
4863    pub dilation: [usize; 2],
4864}
4865
4866/// Operation IR for the weight-gradient of `conv_transpose1d`.
4867#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4868pub struct ConvTranspose1dWeightBackwardOpIr {
4869    /// Forward input.
4870    pub x: TensorIr,
4871    /// Convolution weights.
4872    pub weight: TensorIr,
4873    /// Upstream gradient.
4874    pub output_grad: TensorIr,
4875    /// Convolution options.
4876    pub options: ConvTranspose1dOptionsIr,
4877    /// Output: gradient w.r.t. `weight`.
4878    pub out: TensorIr,
4879}
4880
4881/// Operation IR for the bias-gradient of `conv_transpose1d`.
4882#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4883pub struct ConvTranspose1dBiasBackwardOpIr {
4884    /// Forward input.
4885    pub x: TensorIr,
4886    /// Bias tensor (sets the output shape).
4887    pub bias: TensorIr,
4888    /// Upstream gradient.
4889    pub output_grad: TensorIr,
4890    /// Output: gradient w.r.t. `bias`.
4891    pub out: TensorIr,
4892}
4893
4894/// Operation IR for the weight-gradient of `conv_transpose2d`.
4895#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4896pub struct ConvTranspose2dWeightBackwardOpIr {
4897    /// Forward input.
4898    pub x: TensorIr,
4899    /// Convolution weights.
4900    pub weight: TensorIr,
4901    /// Upstream gradient.
4902    pub output_grad: TensorIr,
4903    /// Convolution options.
4904    pub options: ConvTranspose2dOptionsIr,
4905    /// Output: gradient w.r.t. `weight`.
4906    pub out: TensorIr,
4907}
4908
4909/// Operation IR for the bias-gradient of `conv_transpose2d`.
4910#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4911pub struct ConvTranspose2dBiasBackwardOpIr {
4912    /// Forward input.
4913    pub x: TensorIr,
4914    /// Bias tensor.
4915    pub bias: TensorIr,
4916    /// Upstream gradient.
4917    pub output_grad: TensorIr,
4918    /// Output: gradient w.r.t. `bias`.
4919    pub out: TensorIr,
4920}
4921
4922/// Operation IR for the weight-gradient of `conv_transpose3d`.
4923#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4924pub struct ConvTranspose3dWeightBackwardOpIr {
4925    /// Forward input.
4926    pub x: TensorIr,
4927    /// Convolution weights.
4928    pub weight: TensorIr,
4929    /// Upstream gradient.
4930    pub output_grad: TensorIr,
4931    /// Convolution options.
4932    pub options: ConvTranspose3dOptionsIr,
4933    /// Output: gradient w.r.t. `weight`.
4934    pub out: TensorIr,
4935}
4936
4937/// Operation IR for the bias-gradient of `conv_transpose3d`.
4938#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4939pub struct ConvTranspose3dBiasBackwardOpIr {
4940    /// Forward input.
4941    pub x: TensorIr,
4942    /// Bias tensor.
4943    pub bias: TensorIr,
4944    /// Upstream gradient.
4945    pub output_grad: TensorIr,
4946    /// Output: gradient w.r.t. `bias`.
4947    pub out: TensorIr,
4948}
4949
4950/// Operation IR for the hard-sigmoid activation function, which takes two scalars
4951/// (`alpha` and `beta`) in addition to the input tensor.
4952#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4953pub struct HardSigmoidOpIr {
4954    /// Input tensor.
4955    pub tensor: TensorIr,
4956    /// Alpha — multiplied with the tensor before adding `beta`.
4957    pub alpha: ScalarIr,
4958    /// Beta — added after multiplying by `alpha`.
4959    pub beta: ScalarIr,
4960    /// Output tensor.
4961    pub out: TensorIr,
4962}
4963
4964/// Operation intermediate representation for activation functions.
4965///
4966/// Activations are kept in their own enum (rather than smeared across `FloatOperationIr`)
4967/// so backends can match a single arm to dispatch the fused implementation, instead of
4968/// receiving the function decomposed into 5–20 primitive ops.
4969#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4970pub enum ActivationOperationIr {
4971    /// [relu](burn_backend::ops::ActivationOps::relu).
4972    Relu(UnaryOpIr),
4973    /// [relu_backward](burn_backend::ops::ActivationOps::relu_backward).
4974    /// `lhs` = output of the forward pass, `rhs` = upstream gradient.
4975    ReluBackward(BinaryOpIr),
4976    /// [leaky_relu](burn_backend::ops::ActivationOps::leaky_relu).
4977    /// `lhs` = input, `rhs` = `negative_slope` scalar.
4978    LeakyRelu(ScalarOpIr),
4979    /// [prelu](burn_backend::ops::ActivationOps::prelu).
4980    /// `lhs` = input, `rhs` = alpha tensor.
4981    PRelu(BinaryOpIr),
4982    /// [gelu](burn_backend::ops::ActivationOps::gelu).
4983    Gelu(UnaryOpIr),
4984    /// [gelu_backward](burn_backend::ops::ActivationOps::gelu_backward).
4985    /// `lhs` = forward input, `rhs` = upstream gradient.
4986    GeluBackward(BinaryOpIr),
4987    /// [sigmoid](burn_backend::ops::ActivationOps::sigmoid).
4988    Sigmoid(UnaryOpIr),
4989    /// [sigmoid_backward](burn_backend::ops::ActivationOps::sigmoid_backward).
4990    /// `lhs` = output of the forward pass, `rhs` = upstream gradient.
4991    SigmoidBackward(BinaryOpIr),
4992    /// [hard_sigmoid](burn_backend::ops::ActivationOps::hard_sigmoid).
4993    HardSigmoid(HardSigmoidOpIr),
4994    /// [log_sigmoid](burn_backend::ops::ActivationOps::log_sigmoid).
4995    LogSigmoid(UnaryOpIr),
4996    /// [log_sigmoid_backward](burn_backend::ops::ActivationOps::log_sigmoid_backward).
4997    /// `lhs` = forward input, `rhs` = upstream gradient.
4998    LogSigmoidBackward(BinaryOpIr),
4999    /// [softmax](burn_backend::ops::ActivationOps::softmax).
5000    Softmax(DimOpIr),
5001    /// [log_softmax](burn_backend::ops::ActivationOps::log_softmax).
5002    LogSoftmax(DimOpIr),
5003    /// [softmin](burn_backend::ops::ActivationOps::softmin).
5004    Softmin(DimOpIr),
5005}
5006
5007/// Generate [`ActivationOperationIr`]'s `inputs`/`outputs`/`mark_read_only` from a single
5008/// variant → input-field table, so the three methods can't drift out of sync. The `out`
5009/// field is implicit (every variant's op IR has one). Each entry lists the variant and the
5010/// names of its input tensor fields (`input`, `lhs`/`rhs`, `tensor`, …).
5011macro_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    /// Bumps every visited tensor id by 100 and collects (and rewrites) every visited scalar.
5137    #[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        // NumericFloat MulScalar (a ScalarOpIr with a scalar `rhs`).
5159        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        // Bump every tensor id by 100 and rewrite the scalar to 9.0.
5169        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        // The scalar must have been visited (and was rewritable).
5182        assert_eq!(visitor.scalars, vec![ScalarIr::Float(2.0)]);
5183
5184        // Visiting again observes the rewritten scalar.
5185        let mut after = CollectVisitor::default();
5186        mul.visit_mut(&mut after);
5187        assert_eq!(after.scalars, vec![ScalarIr::Float(9.0)]);
5188
5189        // BaseFloat Reshape (a ShapeOpIr, input + out, no scalars).
5190        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}