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    /// Batch normalization with explicitly supplied statistics, corresponding
223    /// to [batch_norm](burn_backend::ops::ModuleOps::batch_norm).
224    BatchNorm(BatchNormOpIr),
225    /// Operation corresponding to [embedding](burn_backend::ops::ModuleOps::embedding).
226    Embedding(EmbeddingOpIr),
227    /// Operation corresponding to [embedding_backward](burn_backend::ops::ModuleOps::embedding_backward).
228    EmbeddingBackward(EmbeddingBackwardOpIr),
229    /// Operation corresponding to [linear](burn_backend::ops::ModuleOps::linear).
230    Linear(LinearOpIr),
231    /// Operation corresponding to [linear_x_backward](burn_backend::ops::ModuleOps::linear_x_backward).
232    LinearXBackward(LinearXBackwardOpIr),
233    /// Operation corresponding to [linear_weight_backward](burn_backend::ops::ModuleOps::linear_weight_backward).
234    LinearWeightBackward(LinearWeightBackwardOpIr),
235    /// Operation corresponding to [linear_bias_backward](burn_backend::ops::ModuleOps::linear_bias_backward).
236    LinearBiasBackward(LinearBiasBackwardOpIr),
237    /// Operation corresponding to [conv1d](burn_backend::ops::ModuleOps::conv1d).
238    Conv1d(Conv1dOpIr),
239    /// Operation corresponding to [conv1d_x_backward](burn_backend::ops::ModuleOps::conv1d_x_backward).
240    Conv1dXBackward(Conv1dXBackwardOpIr),
241    /// Operation corresponding to [conv1d_weight_backward](burn_backend::ops::ModuleOps::conv1d_weight_backward).
242    Conv1dWeightBackward(Conv1dWeightBackwardOpIr),
243    /// Operation corresponding to [conv1d_bias_backward](burn_backend::ops::ModuleOps::conv1d_bias_backward).
244    Conv1dBiasBackward(Conv1dBiasBackwardOpIr),
245    /// Operation corresponding to [conv2d](burn_backend::ops::ModuleOps::conv2d).
246    Conv2d(Conv2dOpIr),
247    /// Operation corresponding to [conv2d_x_backward](burn_backend::ops::ModuleOps::conv2d_x_backward).
248    Conv2dXBackward(Conv2dXBackwardOpIr),
249    /// Operation corresponding to [conv2d_weight_backward](burn_backend::ops::ModuleOps::conv2d_weight_backward).
250    Conv2dWeightBackward(Conv2dWeightBackwardOpIr),
251    /// Operation corresponding to [conv2d_bias_backward](burn_backend::ops::ModuleOps::conv2d_bias_backward).
252    Conv2dBiasBackward(Conv2dBiasBackwardOpIr),
253    /// Operation corresponding to [conv3d](burn_backend::ops::ModuleOps::conv3d).
254    Conv3d(Conv3dOpIr),
255    /// Operation corresponding to [conv3d_x_backward](burn_backend::ops::ModuleOps::conv3d_x_backward).
256    Conv3dXBackward(Conv3dXBackwardOpIr),
257    /// Operation corresponding to [conv3d_weight_backward](burn_backend::ops::ModuleOps::conv3d_weight_backward).
258    Conv3dWeightBackward(Conv3dWeightBackwardOpIr),
259    /// Operation corresponding to [conv3d_bias_backward](burn_backend::ops::ModuleOps::conv3d_bias_backward).
260    Conv3dBiasBackward(Conv3dBiasBackwardOpIr),
261    /// Operation corresponding to [deform_conv2d](burn_backend::ops::ModuleOps::deform_conv2d)
262    DeformableConv2d(Box<DeformConv2dOpIr>),
263    /// Operation corresponding to [deform_conv2d_backward](burn_backend::ops::ModuleOps::deform_conv2d_backward)
264    DeformableConv2dBackward(Box<DeformConv2dBackwardOpIr>),
265    /// Operation corresponding to [conv transpose 1d](burn_backend::ops::ModuleOps::conv_transpose1d).
266    ConvTranspose1d(ConvTranspose1dOpIr),
267    /// Operation corresponding to [conv transpose 2d](burn_backend::ops::ModuleOps::conv_transpose2d).
268    ConvTranspose2d(ConvTranspose2dOpIr),
269    /// Operation corresponding to [conv transpose 3d](burn_backend::ops::ModuleOps::conv_transpose3d).
270    ConvTranspose3d(ConvTranspose3dOpIr),
271    /// Operation corresponding to [avg pool 1d](burn_backend::ops::ModuleOps::avg_pool1d).
272    AvgPool1d(AvgPool1dOpIr),
273    /// Operation corresponding to [avg pool 2d](burn_backend::ops::ModuleOps::avg_pool2d).
274    AvgPool2d(AvgPool2dOpIr),
275    /// Operation corresponding to
276    /// [avg pool 1d backward](burn_backend::ops::ModuleOps::avg_pool1d_backward).
277    AvgPool1dBackward(AvgPool1dBackwardOpIr),
278    /// Operation corresponding to
279    /// [avg pool 2d backward](burn_backend::ops::ModuleOps::avg_pool2d_backward).
280    AvgPool2dBackward(AvgPool2dBackwardOpIr),
281    /// Operation corresponding to
282    /// [adaptive avg pool 1d](burn_backend::ops::ModuleOps::adaptive_avg_pool1d).
283    AdaptiveAvgPool1d(AdaptiveAvgPool1dOpIr),
284    /// Operation corresponding to
285    /// [adaptive avg pool 2d](burn_backend::ops::ModuleOps::adaptive_avg_pool2d).
286    AdaptiveAvgPool2d(AdaptiveAvgPool2dOpIr),
287    /// Operation corresponding to
288    /// [adaptive avg pool 1d backward](burn_backend::ops::ModuleOps::adaptive_avg_pool1d_backward).
289    AdaptiveAvgPool1dBackward(AdaptiveAvgPool1dBackwardOpIr),
290    /// Operation corresponding to
291    /// [adaptive avg pool 2d backward](burn_backend::ops::ModuleOps::adaptive_avg_pool2d_backward).
292    AdaptiveAvgPool2dBackward(AdaptiveAvgPool2dBackwardOpIr),
293    /// Operation corresponding to
294    /// [adaptive avg pool 3d](burn_backend::ops::ModuleOps::adaptive_avg_pool3d).
295    AdaptiveAvgPool3d(AdaptiveAvgPool3dOpIr),
296    /// Operation corresponding to
297    /// [adaptive avg pool 3d backward](burn_backend::ops::ModuleOps::adaptive_avg_pool3d_backward).
298    AdaptiveAvgPool3dBackward(AdaptiveAvgPool3dBackwardOpIr),
299    /// Operation corresponding to
300    /// [max pool 1d](burn_backend::ops::ModuleOps::max_pool1d).
301    MaxPool1d(MaxPool1dOpIr),
302    /// Operation corresponding to
303    /// [max pool 1d with indices](burn_backend::ops::ModuleOps::max_pool1d_with_indices).
304    MaxPool1dWithIndices(MaxPool1dWithIndicesOpIr),
305    /// Operation corresponding to
306    /// [max pool 1d with indices backward](burn_backend::ops::ModuleOps::max_pool1d_with_indices_backward).
307    MaxPool1dWithIndicesBackward(MaxPool1dWithIndicesBackwardOpIr),
308    /// Operation corresponding to
309    /// [max pool 2d](burn_backend::ops::ModuleOps::max_pool1d).
310    MaxPool2d(MaxPool2dOpIr),
311    /// Operation corresponding to
312    /// [max pool 2d with indices](burn_backend::ops::ModuleOps::max_pool2d_with_indices).
313    MaxPool2dWithIndices(MaxPool2dWithIndicesOpIr),
314    /// Operation corresponding to
315    /// [max pool 2d with indices backward](burn_backend::ops::ModuleOps::max_pool2d_with_indices_backward).
316    MaxPool2dWithIndicesBackward(MaxPool2dWithIndicesBackwardOpIr),
317    /// Operation corresponding to [interpolate](burn_backend::ops::ModuleOps::interpolate).
318    Interpolate(InterpolateOpIr),
319    /// Operation corresponding to [interpolate backward](burn_backend::ops::ModuleOps::interpolate_backward).
320    InterpolateBackward(InterpolateBackwardOpIr),
321    /// Operation corresponding to [rfft](burn_backend::ops::ModuleOps::rfft)
322    Rfft(RfftOpIr),
323    /// Operation corresponding to [irfft](burn_backend::ops::ModuleOps::irfft)
324    IRfft(IRfftOpIr),
325    /// Operation corresponding to [attention](burn_backend::ops::ModuleOps::attention).
326    Attention(AttentionOpIr),
327    /// Operation corresponding to [ctc_loss](burn_backend::ops::ModuleOps::ctc_loss).
328    CtcLoss(CtcLossOpIr),
329    /// Operation corresponding to
330    /// [ctc_loss_backward](burn_backend::ops::ModuleOps::ctc_loss_backward).
331    CtcLossBackward(CtcLossBackwardOpIr),
332    /// Operation corresponding to [layer_norm](burn_backend::ops::ModuleOps::layer_norm).
333    LayerNorm(LayerNormOpIr),
334    /// Operation corresponding to [unfold4d](burn_backend::ops::ModuleOps::unfold4d).
335    Unfold4d(Unfold4dOpIr),
336    /// Operation corresponding to
337    /// [conv_transpose1d_weight_backward](burn_backend::ops::ModuleOps::conv_transpose1d_weight_backward).
338    ConvTranspose1dWeightBackward(ConvTranspose1dWeightBackwardOpIr),
339    /// Operation corresponding to
340    /// [conv_transpose1d_bias_backward](burn_backend::ops::ModuleOps::conv_transpose1d_bias_backward).
341    ConvTranspose1dBiasBackward(ConvTranspose1dBiasBackwardOpIr),
342    /// Operation corresponding to
343    /// [conv_transpose2d_weight_backward](burn_backend::ops::ModuleOps::conv_transpose2d_weight_backward).
344    ConvTranspose2dWeightBackward(ConvTranspose2dWeightBackwardOpIr),
345    /// Operation corresponding to
346    /// [conv_transpose2d_bias_backward](burn_backend::ops::ModuleOps::conv_transpose2d_bias_backward).
347    ConvTranspose2dBiasBackward(ConvTranspose2dBiasBackwardOpIr),
348    /// Operation corresponding to
349    /// [conv_transpose3d_weight_backward](burn_backend::ops::ModuleOps::conv_transpose3d_weight_backward).
350    ConvTranspose3dWeightBackward(ConvTranspose3dWeightBackwardOpIr),
351    /// Operation corresponding to
352    /// [conv_transpose3d_bias_backward](burn_backend::ops::ModuleOps::conv_transpose3d_bias_backward).
353    ConvTranspose3dBiasBackward(ConvTranspose3dBiasBackwardOpIr),
354}
355
356/// Basic operations that can be done on any tensor type.
357#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
358pub enum BaseOperationIr {
359    /// Operation corresponding to:
360    ///
361    /// Float => [reshape](burn_backend::ops::FloatTensorOps::float_reshape).
362    /// Int => [reshape](burn_backend::ops::IntTensorOps::int_reshape).
363    /// Bool => [reshape](burn_backend::ops::BoolTensorOps::bool_reshape).
364    Reshape(ShapeOpIr),
365
366    /// Operation corresponding to:
367    ///
368    /// Float => [swap_dims](burn_backend::ops::FloatTensorOps::float_swap_dims).
369    /// Int => [swap_dims](burn_backend::ops::IntTensorOps::int_swap_dims).
370    /// Bool => [swap_dims](burn_backend::ops::BoolTensorOps::bool_swap_dims).
371    SwapDims(SwapDimsOpIr),
372
373    /// Operation corresponding to:
374    ///
375    /// Float => [permute](burn_backend::ops::FloatTensorOps::float_permute).
376    /// Int => [permute](burn_backend::ops::IntTensorOps::int_permute).
377    /// Bool => [permute](burn_backend::ops::BoolTensorOps::bool_permute).
378    Permute(PermuteOpIr),
379
380    /// Operation corresponding to:
381    /// Float => [flip](burn_backend::ops::FloatTensorOps::float_flip).
382    /// Int => [flip](burn_backend::ops::IntTensorOps::int_flip).
383    /// Bool => [flip](burn_backend::ops::BoolTensorOps::bool_flip).
384    Flip(FlipOpIr),
385
386    /// Operation corresponding to:
387    ///
388    /// Float => [expand](burn_backend::ops::FloatTensorOps::float_expand).
389    /// Int => [expand](burn_backend::ops::IntTensorOps::int_expand).
390    /// Bool => [expand](burn_backend::ops::BoolTensorOps::bool_expand).
391    Expand(ShapeOpIr),
392
393    /// Unfold windows along an axis.
394    ///
395    Unfold(UnfoldOpIr),
396
397    /// Operation corresponding to:
398    ///
399    /// Float => [slice](burn_backend::ops::FloatTensorOps::float_slice).
400    /// Int => [slice](burn_backend::ops::IntTensorOps::int_slice).
401    /// Bool => [slice](burn_backend::ops::BoolTensorOps::bool_slice).
402    Slice(SliceOpIr),
403    /// Operation corresponding to:
404    ///
405    /// Float => [slice assign](burn_backend::ops::FloatTensorOps::float_slice_assign).
406    /// Int => [slice assign](burn_backend::ops::IntTensorOps::int_slice_assign).
407    /// Bool => [slice assign](burn_backend::ops::BoolTensorOps::bool_slice_assign).
408    SliceAssign(SliceAssignOpIr),
409    /// Operation corresponding to:
410    ///
411    /// Float => [select](burn_backend::ops::FloatTensorOps::float_select).
412    /// Int => [select](burn_backend::ops::IntTensorOps::int_select).
413    /// Bool => [select](burn_backend::ops::BoolTensorOps::bool_select).
414    Select(SelectOpIr),
415    /// Operation corresponding to:
416    ///
417    /// Float => [select assign](burn_backend::ops::FloatTensorOps::float_select_assign).
418    /// Int => [select assign](burn_backend::ops::IntTensorOps::int_select_assign).
419    /// Bool => [select assign](burn_backend::ops::BoolTensorOps::bool_select_or).
420    SelectAssign(SelectAssignOpIr),
421    /// Operation corresponding to:
422    ///
423    /// Float => [mask where](burn_backend::ops::FloatTensorOps::float_mask_where).
424    /// Int => [mask where](burn_backend::ops::IntTensorOps::int_mask_where).
425    /// Bool => [mask where](burn_backend::ops::BoolTensorOps::bool_mask_where).
426    MaskWhere(MaskWhereOpIr),
427    /// Operation corresponding to:
428    ///
429    /// Float => [mask fill](burn_backend::ops::FloatTensorOps::float_mask_fill).
430    /// Int => [mask fill](burn_backend::ops::IntTensorOps::int_mask_fill).
431    /// Bool => [mask fill](burn_backend::ops::BoolTensorOps::bool_mask_fill).
432    MaskFill(MaskFillOpIr),
433    /// Operation corresponding to:
434    ///
435    /// Float => [gather](burn_backend::ops::FloatTensorOps::float_gather).
436    /// Int => [gather](burn_backend::ops::IntTensorOps::int_gather).
437    /// Bool => [gather](burn_backend::ops::BoolTensorOps::bool_gather).
438    Gather(GatherOpIr),
439    /// Operation corresponding to:
440    ///
441    /// Float => [scatter](burn_backend::ops::FloatTensorOps::float_scatter_add).
442    /// Int => [scatter](burn_backend::ops::IntTensorOps::int_scatter_add).
443    /// Bool => [scatter](burn_backend::ops::BoolTensorOps::bool_scatter_or).
444    Scatter(ScatterOpIr),
445    /// Multi-dimensional scatter operation.
446    ScatterNd(ScatterNdOpIr),
447    /// Multi-dimensional gather operation.
448    GatherNd(GatherNdOpIr),
449    /// Operation corresponding to:
450    ///
451    /// Float => [equal](burn_backend::ops::FloatTensorOps::float_equal).
452    /// Int => [equal](burn_backend::ops::IntTensorOps::int_equal).
453    /// Bool => [equal](burn_backend::ops::BoolTensorOps::bool_equal).
454    Equal(BinaryOpIr),
455    /// Operation corresponding to:
456    ///
457    /// Float => [equal elem](burn_backend::ops::FloatTensorOps::float_equal_elem).
458    /// Int => [equal elem](burn_backend::ops::IntTensorOps::int_equal_elem).
459    /// Bool => [equal elem](burn_backend::ops::BoolTensorOps::bool_equal_elem).
460    EqualElem(ScalarOpIr),
461    /// Operation corresponding to:
462    ///
463    /// Float => [repeat dim](burn_backend::ops::FloatTensorOps::float_repeat_dim).
464    /// Int => [repeat dim](burn_backend::ops::IntTensorOps::int_repeat_dim).
465    /// Bool => [repeat dim](burn_backend::ops::BoolTensorOps::bool_repeat_dim).
466    RepeatDim(RepeatDimOpIr),
467    /// Operation corresponding to:
468    ///
469    /// Float => [cat](burn_backend::ops::FloatTensorOps::float_cat).
470    /// Int => [cat](burn_backend::ops::IntTensorOps::int_cat).
471    /// Bool => [cat](burn_backend::ops::BoolTensorOps::bool_cat).
472    Cat(CatOpIr),
473    /// Cast operation, no direct operation and should be supported by fusion backend.
474    Cast(CastOpIr),
475    /// Operation corresponding to:
476    ///
477    /// Float => [empty](burn_backend::ops::FloatTensorOps::float_empty).
478    /// Int => [empty](burn_backend::ops::IntTensorOps::int_empty).
479    /// Bool => [empty](burn_backend::ops::BoolTensorOps::bool_empty).
480    Empty(CreationOpIr),
481    /// Operation corresponding to:
482    ///
483    /// Float => [ones](burn_backend::ops::FloatTensorOps::float_ones).
484    /// Int => [ones](burn_backend::ops::IntTensorOps::int_ones).
485    /// Bool => [ones](burn_backend::ops::BoolTensorOps::bool_ones).
486    Ones(CreationOpIr),
487    /// Operation corresponding to:
488    ///
489    /// Float => [zeros](burn_backend::ops::FloatTensorOps::float_zeros).
490    /// Int => [zeros](burn_backend::ops::IntTensorOps::int_zeros).
491    /// Bool => [zeros](burn_backend::ops::BoolTensorOps::bool_zeros).
492    Zeros(CreationOpIr),
493    /// Operation corresponding to:
494    ///
495    /// Float => [not_equal](burn_backend::ops::FloatTensorOps::float_not_equal).
496    /// Int => [not_equal](burn_backend::ops::IntTensorOps::int_not_equal).
497    /// Bool => [not_equal](burn_backend::ops::BoolTensorOps::bool_not_equal).
498    NotEqual(BinaryOpIr),
499    /// Operation corresponding to:
500    ///
501    /// Float => [not_equal_elem](burn_backend::ops::FloatTensorOps::float_not_equal_elem).
502    /// Int => [not_equal_elem](burn_backend::ops::IntTensorOps::int_not_equal_elem).
503    /// Bool => [not_equal_elem](burn_backend::ops::BoolTensorOps::bool_not_equal_elem).
504    NotEqualElem(ScalarOpIr),
505    /// Reduce-`all` over the input tensor.
506    ///
507    /// Float/Int input is treated as a non-zero check; bool input is the value itself.
508    /// Output is a single-element bool tensor.
509    All(ReduceOpIr),
510    /// Reduce-`any` over the input tensor.
511    Any(ReduceOpIr),
512    /// Reduce-`all` along a dim.
513    AllDim(ReduceDimOpIr),
514    /// Reduce-`any` along a dim.
515    AnyDim(ReduceDimOpIr),
516}
517
518/// Numeric operations on int and float tensors.
519#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
520pub enum NumericOperationIr {
521    /// Operation corresponding to:
522    ///
523    /// Float => [add](burn_backend::ops::FloatTensorOps::float_add).
524    /// Int => [add](burn_backend::ops::IntTensorOps::int_add).
525    Add(BinaryOpIr),
526    /// Operation corresponding to:
527    ///
528    /// Float => [add scalar](burn_backend::ops::FloatTensorOps::float_add_scalar).
529    /// Int => [add scalar](burn_backend::ops::IntTensorOps::int_add_scalar).
530    AddScalar(ScalarOpIr),
531    /// Operation corresponding to:
532    ///
533    /// Float => [sub](burn_backend::ops::FloatTensorOps::float_sub).
534    /// Int => [sub](burn_backend::ops::IntTensorOps::int_sub).
535    Sub(BinaryOpIr),
536    /// Operation corresponding to:
537    ///
538    /// Float => [sub scalar](burn_backend::ops::FloatTensorOps::float_sub_scalar).
539    /// Int => [sub scalar](burn_backend::ops::IntTensorOps::int_sub_scalar).
540    SubScalar(ScalarOpIr),
541    /// Operation corresponding to:
542    ///
543    /// Float => [div](burn_backend::ops::FloatTensorOps::float_div).
544    /// Int => [div](burn_backend::ops::IntTensorOps::int_div).
545    Div(BinaryOpIr),
546    /// Operation corresponding to:
547    ///
548    /// Float => [div scalar](burn_backend::ops::FloatTensorOps::float_div_scalar).
549    /// Int => [div scalar](burn_backend::ops::IntTensorOps::int_div_scalar).
550    DivScalar(ScalarOpIr),
551    /// Operation corresponding to:
552    ///
553    /// Float => [rem](burn_backend::ops::FloatTensorOps::float_remainder).
554    /// Int => [rem](burn_backend::ops::IntTensorOps::int_remainder).
555    Rem(BinaryOpIr),
556    /// Operation corresponding to:
557    ///
558    /// Float => [rem scalar](burn_backend::ops::FloatTensorOps::float_remainder_scalar).
559    /// Int => [rem scalar](burn_backend::ops::IntTensorOps::int_remainder_scalar).
560    RemScalar(ScalarOpIr),
561    /// Operation corresponding to:
562    ///
563    /// Float => [mul](burn_backend::ops::FloatTensorOps::float_mul).
564    /// Int => [mul](burn_backend::ops::IntTensorOps::int_mul).
565    Mul(BinaryOpIr),
566    /// Operation corresponding to:
567    ///
568    /// Float => [mul scalar](burn_backend::ops::FloatTensorOps::float_mul_scalar).
569    /// Int => [mul scalar](burn_backend::ops::IntTensorOps::int_mul_scalar).
570    MulScalar(ScalarOpIr),
571    /// Operation corresponding to:
572    ///
573    /// Float => [abs](burn_backend::ops::FloatTensorOps::float_abs).
574    /// Int => [abs](burn_backend::ops::IntTensorOps::int_abs).
575    Abs(UnaryOpIr),
576    /// Operation corresponding to:
577    ///
578    /// Float => [full](burn_backend::ops::FloatTensorOps::float_full).
579    /// Int => [full](burn_backend::ops::IntTensorOps::int_full).
580    Full(FullOpIr),
581    /// Operation corresponding to:
582    ///
583    /// Float => [mean dim](burn_backend::ops::FloatTensorOps::float_mean_dim).
584    /// Int => [mean dim](burn_backend::ops::IntTensorOps::int_mean_dim).
585    MeanDim(ReduceDimOpIr),
586    /// Operation corresponding to:
587    ///
588    /// Float => [mean](burn_backend::ops::FloatTensorOps::float_mean).
589    /// Int => [mean](burn_backend::ops::IntTensorOps::int_mean).
590    Mean(ReduceOpIr),
591    /// Operation corresponding to:
592    ///
593    /// Float => [sum](burn_backend::ops::FloatTensorOps::float_sum).
594    /// Int => [sum](burn_backend::ops::IntTensorOps::int_sum).
595    Sum(ReduceOpIr),
596    /// Operation corresponding to:
597    ///
598    /// Float => [sum dim](burn_backend::ops::FloatTensorOps::float_sum_dim).
599    /// Int => [sum dim](burn_backend::ops::IntTensorOps::int_sum_dim).
600    SumDim(ReduceDimOpIr),
601    /// Operation corresponding to:
602    ///
603    /// Float => [prod](burn_backend::ops::FloatTensorOps::float_prod).
604    /// Int => [prod](burn_backend::ops::IntTensorOps::int_prod).
605    Prod(ReduceOpIr),
606    /// Operation corresponding to:
607    ///
608    /// Float => [prod dim](burn_backend::ops::FloatTensorOps::float_prod_dim).
609    /// Int => [prod dim](burn_backend::ops::IntTensorOps::int_prod_dim).
610    ProdDim(ReduceDimOpIr),
611    /// Operation corresponding to:
612    ///
613    /// Float => [greater](burn_backend::ops::FloatTensorOps::float_greater).
614    /// Int => [greater](burn_backend::ops::IntTensorOps::int_greater).
615    Greater(BinaryOpIr),
616    /// Operation corresponding to:
617    ///
618    /// Float => [greater elem](burn_backend::ops::FloatTensorOps::float_greater_elem).
619    /// Int => [greater elem](burn_backend::ops::IntTensorOps::int_greater_elem).
620    GreaterElem(ScalarOpIr),
621    /// Operation corresponding to:
622    ///
623    /// Float => [greater equal](burn_backend::ops::FloatTensorOps::float_greater_elem).
624    /// Int => [greater elem](burn_backend::ops::IntTensorOps::int_greater_elem).
625    GreaterEqual(BinaryOpIr),
626    /// Operation corresponding to:
627    ///
628    /// Float => [greater equal elem](burn_backend::ops::FloatTensorOps::float_greater_equal_elem).
629    /// Int => [greater equal elem](burn_backend::ops::IntTensorOps::int_greater_equal_elem).
630    GreaterEqualElem(ScalarOpIr),
631    /// Operation corresponding to:
632    ///
633    /// Float => [lower](burn_backend::ops::FloatTensorOps::float_lower).
634    /// Int => [lower](burn_backend::ops::IntTensorOps::int_lower).
635    Lower(BinaryOpIr),
636    /// Operation corresponding to:
637    ///
638    /// Float => [lower elem](burn_backend::ops::FloatTensorOps::float_lower_elem).
639    /// Int => [lower elem](burn_backend::ops::IntTensorOps::int_lower_elem).
640    LowerElem(ScalarOpIr),
641    /// Operation corresponding to:
642    ///
643    /// Float => [lower equal](burn_backend::ops::FloatTensorOps::float_lower_equal).
644    /// Int => [lower equal](burn_backend::ops::IntTensorOps::int_lower_equal).
645    LowerEqual(BinaryOpIr),
646    /// Operation corresponding to:
647    ///
648    /// Float => [lower equal elem](burn_backend::ops::FloatTensorOps::float_lower_equal_elem).
649    /// Int => [lower equal elem](burn_backend::ops::IntTensorOps::int_lower_equal_elem).
650    LowerEqualElem(ScalarOpIr),
651    /// Operation corresponding to:
652    ///
653    /// Float => [argmax](burn_backend::ops::FloatTensorOps::float_argmax).
654    /// Int => [argmax](burn_backend::ops::IntTensorOps::int_argmax).
655    ArgMax(ReduceDimOpIr),
656    /// Operation corresponding to:
657    ///
658    /// Float => [argtopk](burn_backend::ops::FloatTensorOps::float_argtopk).
659    /// Int => [argtopk](burn_backend::ops::IntTensorOps::int_argtopk).
660    ArgTopK(ReduceDimOpIr),
661    /// Operation corresponding to:
662    ///
663    /// Float => [topk](burn_backend::ops::FloatTensorOps::float_topk).
664    /// Int => [topk](burn_backend::ops::IntTensorOps::int_topk).
665    TopK(ReduceDimOpIr),
666    /// Operation corresponding to:
667    ///
668    /// Float => [topk with indices](burn_backend::ops::FloatTensorOps::float_topk_with_indices).
669    /// Int => [topk with indices](burn_backend::ops::IntTensorOps::int_topk_with_indices).
670    TopKWithIndices(TopKWithIndicesOpIr),
671    /// Operation corresponding to:
672    ///
673    /// Float => [argmin](burn_backend::ops::FloatTensorOps::float_argmin).
674    /// Int => [argmin](burn_backend::ops::IntTensorOps::int_argmin).
675    ArgMin(ReduceDimOpIr),
676    /// Operation corresponding to:
677    ///
678    /// Float => [max](burn_backend::ops::FloatTensorOps::float_max).
679    /// Int => [max](burn_backend::ops::IntTensorOps::int_max).
680    Max(ReduceOpIr),
681    /// Operation corresponding to:
682    ///
683    /// Float => [max dim with indices](burn_backend::ops::FloatTensorOps::float_max_dim_with_indices).
684    /// Int => [max dim with indices](burn_backend::ops::IntTensorOps::int_max_dim_with_indices).
685    MaxDimWithIndices(ReduceDimWithIndicesOpIr),
686    /// Operation corresponding to:
687    ///
688    /// Float => [min dim with indices](burn_backend::ops::FloatTensorOps::float_min_dim_with_indices).
689    /// Int => [min dim with indices](burn_backend::ops::IntTensorOps::int_min_dim_with_indices).
690    MinDimWithIndices(ReduceDimWithIndicesOpIr),
691    /// Operation corresponding to:
692    ///
693    /// Float => [min](burn_backend::ops::FloatTensorOps::float_min).
694    /// Int => [min](burn_backend::ops::IntTensorOps::int_min).
695    Min(ReduceOpIr),
696    /// Operation corresponding to:
697    ///
698    /// Float => [max dim](burn_backend::ops::FloatTensorOps::float_max_dim).
699    /// Int => [max dim](burn_backend::ops::IntTensorOps::int_max_dim).
700    MaxDim(ReduceDimOpIr),
701    /// Operation corresponding to:
702    ///
703    /// Float => [min dim](burn_backend::ops::FloatTensorOps::float_min_dim).
704    /// Int => [min dim](burn_backend::ops::IntTensorOps::int_min_dim).
705    MinDim(ReduceDimOpIr),
706    /// Operation corresponding to:
707    ///
708    /// Float => [max_abs](burn_backend::ops::FloatTensorOps::float_max_abs).
709    /// Int => [max_abs](burn_backend::ops::IntTensorOps::int_max_abs).
710    MaxAbs(ReduceOpIr),
711    /// Operation corresponding to:
712    ///
713    /// Float => [max_abs dim](burn_backend::ops::FloatTensorOps::float_max_abs_dim).
714    /// Int => [max_abs dim](burn_backend::ops::IntTensorOps::int_max_abs_dim).
715    MaxAbsDim(ReduceDimOpIr),
716    /// Operation corresponding to:
717    ///
718    /// Float => [clamp](burn_backend::ops::FloatTensorOps::float_clamp).
719    /// Int => [clamp](burn_backend::ops::IntTensorOps::int_clamp).
720    Clamp(ClampOpIr),
721    /// Operation corresponding to:
722    ///
723    /// Int => [random](burn_backend::ops::IntTensorOps::int_random).
724    IntRandom(RandomOpIr),
725    /// Operation corresponding to:
726    ///
727    /// Float => [powf](burn_backend::ops::FloatTensorOps::float_powi).
728    /// Int => [powf](burn_backend::ops::IntTensorOps::int_powi).
729    Powi(BinaryOpIr),
730    /// Operation corresponding to:
731    ///
732    /// Float => [powi_scalar](burn_backend::ops::FloatTensorOps::float_powi_scalar).
733    /// Int => [powi_scalar](burn_backend::ops::IntTensorOps::int_powi_scalar).
734    PowiScalar(ScalarOpIr),
735    /// Operation corresponding to:
736    ///
737    /// Float => [cumsum](burn_backend::ops::FloatTensorOps::float_cumsum).
738    /// Int => [cumsum](burn_backend::ops::IntTensorOps::int_cumsum).
739    CumSum(DimOpIr),
740    /// Operation corresponding to:
741    ///
742    /// Float => [cumprod](burn_backend::ops::FloatTensorOps::float_cumprod).
743    /// Int => [cumprod](burn_backend::ops::IntTensorOps::int_cumprod).
744    CumProd(DimOpIr),
745    /// Operation corresponding to:
746    ///
747    /// Float => [cummin](burn_backend::ops::FloatTensorOps::float_cummin).
748    /// Int => [cummin](burn_backend::ops::IntTensorOps::int_cummin).
749    CumMin(DimOpIr),
750    /// Operation corresponding to:
751    ///
752    /// Float => [cummax](burn_backend::ops::FloatTensorOps::float_cummax).
753    /// Int => [cummax](burn_backend::ops::IntTensorOps::int_cummax).
754    CumMax(DimOpIr),
755    /// Operation corresponding to:
756    ///
757    /// Float => [neg](burn_backend::ops::FloatTensorOps::float_neg).
758    /// Int => [neg](burn_backend::ops::IntTensorOps::int_neg).
759    Neg(UnaryOpIr),
760    /// Operation corresponding to:
761    ///
762    /// Float => [sign](burn_backend::ops::FloatTensorOps::float_sign).
763    /// Int => [sign](burn_backend::ops::IntTensorOps::int_sign).
764    Sign(UnaryOpIr),
765    /// Operation corresponding to:
766    ///
767    /// Float => [clamp_min](burn_backend::ops::FloatTensorOps::float_clamp_min).
768    /// Int => [clamp_min](burn_backend::ops::IntTensorOps::int_clamp_min).
769    ClampMin(ScalarOpIr),
770    /// Operation corresponding to:
771    ///
772    /// Float => [clamp_max](burn_backend::ops::FloatTensorOps::float_clamp_max).
773    /// Int => [clamp_max](burn_backend::ops::IntTensorOps::int_clamp_max).
774    ClampMax(ScalarOpIr),
775    /// Sort along a dim. Output shares the input's shape and dtype.
776    Sort(SortOpIr),
777    /// Sort along a dim, also returning the source indices.
778    SortWithIndices(SortWithIndicesOpIr),
779    /// Sort along a dim and return only the indices (i.e., the argsort).
780    ///
781    /// Shares [`SortOpIr`] with [`Sort`](Self::Sort); only the output dtype differs.
782    ArgSort(SortOpIr),
783}
784
785/// Operation intermediate representation specific to an int tensor.
786#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
787pub enum IntOperationIr {
788    /// Operation corresponding to [into float](burn_backend::ops::IntTensorOps::int_into_float).
789    IntoFloat(CastOpIr),
790    /// Operation corresponding to:
791    ///
792    /// Int => [bitwise and](burn_backend::ops::IntTensorOps::bitwise_and).
793    BitwiseAnd(BinaryOpIr),
794    /// Operation corresponding to:
795    ///
796    /// Int => [bitwise and scalar](burn_backend::ops::IntTensorOps::bitwise_and_scalar).
797    BitwiseAndScalar(ScalarOpIr),
798    /// Operation corresponding to:
799    ///
800    /// Int => [bitwise or](burn_backend::ops::IntTensorOps::bitwise_or).
801    BitwiseOr(BinaryOpIr),
802    /// Operation corresponding to:
803    ///
804    /// Int => [bitwise or scalar](burn_backend::ops::IntTensorOps::bitwise_or_scalar).
805    BitwiseOrScalar(ScalarOpIr),
806    /// Operation corresponding to:
807    ///
808    /// Int => [bitwise xor](burn_backend::ops::IntTensorOps::bitwise_xor).
809    BitwiseXor(BinaryOpIr),
810    /// Operation corresponding to:
811    ///
812    /// Int => [bitwise xor scalar](burn_backend::ops::IntTensorOps::bitwise_xor_scalar).
813    BitwiseXorScalar(ScalarOpIr),
814    /// Operation corresponding to:
815    ///
816    /// Int => [bitwise not](burn_backend::ops::IntTensorOps::bitwise_not).
817    BitwiseNot(UnaryOpIr),
818    /// Operation corresponding to:
819    ///
820    /// Int => [bitwise left shift](burn_backend::ops::IntTensorOps::bitwise_left_shift).
821    BitwiseLeftShift(BinaryOpIr),
822    /// Operation corresponding to:
823    ///
824    /// Int => [bitwise left shift scalar](burn_backend::ops::IntTensorOps::bitwise_left_shift_scalar).
825    BitwiseLeftShiftScalar(ScalarOpIr),
826    /// Operation corresponding to:
827    ///
828    /// Int => [bitwise right shift](burn_backend::ops::IntTensorOps::bitwise_right_shift).
829    BitwiseRightShift(BinaryOpIr),
830    /// Operation corresponding to:
831    ///
832    /// Int => [bitwise right shift scalar](burn_backend::ops::IntTensorOps::bitwise_right_shift_scalar).
833    BitwiseRightShiftScalar(ScalarOpIr),
834    /// Operation corresponding to [matmul](burn_backend::ops::IntTensorOps::int_matmul).
835    Matmul(MatmulOpIr),
836}
837
838/// Operation intermediate representation specific to a bool tensor.
839#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
840pub enum BoolOperationIr {
841    /// Operation corresponding to [into float](burn_backend::ops::BoolTensorOps::bool_into_float).
842    IntoFloat(CastOpIr),
843    /// Operation corresponding to [into int](burn_backend::ops::BoolTensorOps::bool_into_int).
844    IntoInt(CastOpIr),
845    /// Operation corresponding to [not](burn_backend::ops::BoolTensorOps::bool_not).
846    Not(UnaryOpIr),
847    /// Operation corresponding to [and](burn_backend::ops::BoolTensorOps::bool_and).
848    And(BinaryOpIr),
849    /// Operation corresponding to [or](burn_backend::ops::BoolTensorOps::bool_or).
850    Or(BinaryOpIr),
851    /// Operation corresponding to [xor](burn_backend::ops::BoolTensorOps::bool_xor).
852    Xor(BinaryOpIr),
853}
854
855/// Operations that can be done on distributed tensors.
856#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
857#[allow(clippy::large_enum_variant)]
858pub enum DistributedOperationIr {
859    /// Operation corresponding to:
860    /// [all_reduce](burn_backend::distributed::DistributedOps::all_reduce).
861    AllReduce(AllReduceOpIr),
862    /// Resolve the pending collective operations on the executing device. Corresponds to
863    /// [sync_collective](burn_backend::distributed::DistributedOps::sync_collective).
864    ///
865    /// Fire-and-forget and payload-free: it syncs whichever device the interpreter is bound to.
866    /// Modeled as an operation (not a side-channel call) so it travels the normal op stream
867    /// alongside [`AllReduce`](DistributedOperationIr::AllReduce) and is ordered against it.
868    SyncCollective,
869}
870
871/// Swap dim operation intermediate representation.
872#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
873pub struct SwapDimsOpIr {
874    /// Input tensor intermediate representation.
875    pub input: TensorIr,
876    /// Output tensor intermediate representation.
877    pub out: TensorIr,
878    /// The first dim to swap.
879    pub dim1: usize,
880    /// The second dim to swap.
881    pub dim2: usize,
882}
883
884/// Permute operation intermediate representation.
885#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
886pub struct PermuteOpIr {
887    /// Input tensor intermediate representation.
888    pub input: TensorIr,
889    /// Output tensor intermediate representation.
890    pub out: TensorIr,
891    /// The new order of the dimensions.
892    pub axes: Vec<usize>,
893}
894
895/// Shape operation intermediate representation.
896#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
897pub struct ShapeOpIr {
898    /// Input tensor intermediate representation.
899    pub input: TensorIr,
900    /// Output tensor intermediate representation with the new shape.
901    pub out: TensorIr,
902}
903
904/// Unfold operation intermediate representation.
905#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
906pub struct UnfoldOpIr {
907    /// Input tensor intermediate representation.
908    pub input: TensorIr,
909    /// Output tensor intermediate representation.
910    pub out: TensorIr,
911
912    /// The selected dim.
913    pub dim: usize,
914    /// The window size.
915    pub size: usize,
916    /// The window step along dim.
917    pub step: usize,
918}
919
920/// Flip operation intermediate representation.
921#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
922pub struct FlipOpIr {
923    /// Input tensor intermediate representation.
924    pub input: TensorIr,
925    /// Output tensor intermediate representation.
926    pub out: TensorIr,
927    /// The dimensions to flip.
928    pub axes: Vec<usize>,
929}
930
931#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
932#[allow(missing_docs)]
933pub struct RandomOpIr {
934    pub out: TensorIr,
935    pub distribution: Distribution,
936}
937
938/// Creation operation intermediate representation.
939/// As opposed to [InitOperationIr], creation operations are lazy initialized.
940#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
941pub struct CreationOpIr {
942    /// Output tensor intermediate representation.
943    pub out: TensorIr,
944}
945
946/// Full operation intermediate representation.
947#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
948pub struct FullOpIr {
949    /// Output tensor intermediate representation.
950    pub out: TensorIr,
951    /// Fill value.
952    pub value: ScalarIr,
953}
954
955#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
956/// Declares a tensor has been initialized.
957///
958/// It is necessary to register for proper orphan detection and avoid memory leak.
959pub struct InitOperationIr {
960    /// The initialized tensor.
961    pub out: TensorIr,
962}
963
964#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
965#[allow(missing_docs)]
966pub struct BinaryOpIr {
967    pub lhs: TensorIr,
968    pub rhs: TensorIr,
969    pub out: TensorIr,
970}
971
972#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
973#[allow(missing_docs)]
974pub struct MatmulOpIr {
975    pub lhs: TensorIr,
976    pub rhs: TensorIr,
977    pub out: TensorIr,
978}
979
980#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
981#[allow(missing_docs)]
982pub struct CrossOpIr {
983    pub lhs: TensorIr,
984    pub rhs: TensorIr,
985    pub out: TensorIr,
986    pub dim: usize,
987}
988
989#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
990#[allow(missing_docs)]
991pub struct UnaryOpIr {
992    pub input: TensorIr,
993    pub out: TensorIr,
994}
995
996#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
997#[allow(missing_docs)]
998pub struct ScalarOpIr {
999    pub lhs: TensorIr,
1000    // TODO: Make that an enum with `Value` and `Id` variants for relative/global
1001    // conversion.
1002    pub rhs: ScalarIr,
1003    pub out: TensorIr,
1004}
1005
1006#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, Hash)]
1007#[allow(missing_docs)]
1008pub struct ReduceOpIr {
1009    pub input: TensorIr,
1010    pub out: TensorIr,
1011}
1012
1013#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, Hash)]
1014#[allow(missing_docs)]
1015pub struct ReduceDimOpIr {
1016    pub input: TensorIr,
1017    pub out: TensorIr,
1018    pub axis: usize,
1019    pub accumulator_len: usize,
1020}
1021
1022#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1023#[allow(missing_docs)]
1024pub struct CastOpIr {
1025    pub input: TensorIr,
1026    pub out: TensorIr,
1027}
1028
1029/// IR for operations that operate along a dimension without reducing it.
1030/// Unlike `ReduceDimOpIr`, the output shape is the same as the input shape.
1031#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, Hash)]
1032#[allow(missing_docs)]
1033pub struct DimOpIr {
1034    pub input: TensorIr,
1035    pub out: TensorIr,
1036    pub axis: usize,
1037}
1038
1039#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1040#[allow(missing_docs)]
1041pub struct GatherOpIr {
1042    pub tensor: TensorIr,
1043    pub dim: usize,
1044    pub indices: TensorIr,
1045    pub out: TensorIr,
1046}
1047
1048#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1049#[allow(missing_docs)]
1050pub struct ScatterOpIr {
1051    pub tensor: TensorIr,
1052    pub dim: usize,
1053    pub indices: TensorIr,
1054    pub value: TensorIr,
1055    pub update: IndexingUpdateOp,
1056    pub out: TensorIr,
1057}
1058
1059#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1060#[allow(missing_docs)]
1061pub struct ScatterNdOpIr {
1062    pub data: TensorIr,
1063    pub indices: TensorIr,
1064    pub values: TensorIr,
1065    pub reduction: IndexingUpdateOp,
1066    pub out: TensorIr,
1067}
1068
1069#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1070#[allow(missing_docs)]
1071pub struct GatherNdOpIr {
1072    pub data: TensorIr,
1073    pub indices: TensorIr,
1074    pub out: TensorIr,
1075}
1076
1077#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1078#[allow(missing_docs)]
1079pub struct SelectOpIr {
1080    pub tensor: TensorIr,
1081    pub dim: usize,
1082    pub indices: TensorIr,
1083    pub out: TensorIr,
1084}
1085
1086#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1087#[allow(missing_docs)]
1088pub struct SelectAssignOpIr {
1089    pub tensor: TensorIr,
1090    pub dim: usize,
1091    pub indices: TensorIr,
1092    pub value: TensorIr,
1093    pub update: IndexingUpdateOp,
1094    pub out: TensorIr,
1095}
1096
1097#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1098#[allow(missing_docs)]
1099pub struct SliceOpIr {
1100    pub tensor: TensorIr,
1101    pub ranges: Vec<Slice>,
1102    pub out: TensorIr,
1103}
1104
1105#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1106#[allow(missing_docs)]
1107pub struct SliceAssignOpIr {
1108    pub tensor: TensorIr,
1109    pub ranges: Vec<burn_backend::Slice>,
1110    pub value: TensorIr,
1111    pub out: TensorIr,
1112}
1113
1114#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1115#[allow(missing_docs)]
1116pub struct MaskWhereOpIr {
1117    pub tensor: TensorIr,
1118    pub mask: TensorIr,
1119    pub value: TensorIr,
1120    pub out: TensorIr,
1121}
1122
1123#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1124#[allow(missing_docs)]
1125pub struct MaskFillOpIr {
1126    pub tensor: TensorIr,
1127    pub mask: TensorIr,
1128    pub value: ScalarIr,
1129    pub out: TensorIr,
1130}
1131
1132#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1133#[allow(missing_docs)]
1134pub struct ClampOpIr {
1135    pub tensor: TensorIr,
1136    pub min: ScalarIr,
1137    pub max: ScalarIr,
1138    pub out: TensorIr,
1139}
1140
1141#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1142#[allow(missing_docs)]
1143pub struct RepeatDimOpIr {
1144    pub tensor: TensorIr,
1145    pub dim: usize,
1146    pub times: usize,
1147    pub out: TensorIr,
1148}
1149
1150#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1151#[allow(missing_docs)]
1152pub struct CatOpIr {
1153    pub tensors: Vec<TensorIr>,
1154    pub dim: usize,
1155    pub out: TensorIr,
1156}
1157
1158#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1159#[allow(missing_docs)]
1160pub struct AllReduceOpIr {
1161    pub tensor: TensorIr,
1162    pub out: TensorIr,
1163    /// How to reduce the values across the participating devices.
1164    pub op: burn_backend::distributed::ReduceOperation,
1165    /// The devices participating in the collective operation.
1166    pub device_ids: Vec<DeviceIdIr>,
1167}
1168
1169/// Serializable representation of a [device id](burn_backend::DeviceId).
1170///
1171/// The intermediate representation is part of the wire protocol (e.g. the remote backend), so it
1172/// cannot store `burn_backend::DeviceId` directly since that type is not serializable.
1173#[derive(Clone, Copy, Debug, Hash, PartialEq, Eq, Serialize, Deserialize)]
1174pub struct DeviceIdIr {
1175    /// Identifies the type of the device.
1176    pub type_id: u16,
1177    /// Identifies the device number.
1178    pub index_id: u16,
1179}
1180
1181impl From<burn_backend::DeviceId> for DeviceIdIr {
1182    fn from(value: burn_backend::DeviceId) -> Self {
1183        Self {
1184            type_id: value.type_id,
1185            index_id: value.index_id,
1186        }
1187    }
1188}
1189
1190impl From<DeviceIdIr> for burn_backend::DeviceId {
1191    fn from(value: DeviceIdIr) -> Self {
1192        burn_backend::DeviceId::new(value.type_id, value.index_id)
1193    }
1194}
1195
1196#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1197#[allow(missing_docs)]
1198pub struct ReduceDimWithIndicesOpIr {
1199    pub tensor: TensorIr,
1200    pub dim: usize,
1201    pub out: TensorIr,
1202    pub out_indices: TensorIr,
1203}
1204
1205#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1206#[allow(missing_docs)]
1207/// Like [`ReduceDimWithIndicesOpIr`], but for a top-k: the reduced axis keeps `k` entries
1208/// instead of collapsing to 1, so `k` has to be carried explicitly.
1209pub struct TopKWithIndicesOpIr {
1210    pub tensor: TensorIr,
1211    pub dim: usize,
1212    pub k: usize,
1213    pub out: TensorIr,
1214    pub out_indices: TensorIr,
1215}
1216
1217#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1218#[allow(missing_docs)]
1219pub struct EmbeddingOpIr {
1220    pub weights: TensorIr,
1221    pub indices: TensorIr,
1222    pub out: TensorIr,
1223}
1224
1225#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1226#[allow(missing_docs)]
1227pub struct EmbeddingBackwardOpIr {
1228    pub weights: TensorIr,
1229    pub out_grad: TensorIr,
1230    pub indices: TensorIr,
1231    pub out: TensorIr,
1232}
1233
1234/// Batch normalization using explicitly supplied channel statistics.
1235///
1236/// This operation neither calculates nor updates the supplied statistics.
1237#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1238#[allow(missing_docs)]
1239pub struct BatchNormOpIr {
1240    pub x: TensorIr,
1241    pub gamma: TensorIr,
1242    pub beta: TensorIr,
1243    pub mean: TensorIr,
1244    pub variance: TensorIr,
1245    pub epsilon: ScalarIr,
1246    pub out: TensorIr,
1247}
1248
1249#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1250#[allow(missing_docs)]
1251pub struct LinearOpIr {
1252    pub x: TensorIr,
1253    pub weight: TensorIr,
1254    pub bias: Option<TensorIr>,
1255    pub out: TensorIr,
1256}
1257
1258#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1259#[allow(missing_docs)]
1260pub struct LinearXBackwardOpIr {
1261    pub weight: TensorIr,
1262    pub output_grad: TensorIr,
1263    pub out: TensorIr,
1264}
1265
1266#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1267#[allow(missing_docs)]
1268pub struct LinearWeightBackwardOpIr {
1269    pub x: TensorIr,
1270    pub output_grad: TensorIr,
1271    pub out: TensorIr,
1272}
1273
1274#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1275#[allow(missing_docs)]
1276pub struct LinearBiasBackwardOpIr {
1277    pub output_grad: TensorIr,
1278    pub out: TensorIr,
1279}
1280
1281#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1282#[allow(missing_docs)]
1283pub struct Conv1dOpIr {
1284    pub x: TensorIr,
1285    pub weight: TensorIr,
1286    pub bias: Option<TensorIr>,
1287    pub options: Conv1dOptionsIr,
1288    pub out: TensorIr,
1289}
1290
1291#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1292#[allow(missing_docs)]
1293pub struct Conv1dXBackwardOpIr {
1294    pub x: TensorIr,
1295    pub weight: TensorIr,
1296    pub output_grad: TensorIr,
1297    pub options: Conv1dOptionsIr,
1298    pub out: TensorIr,
1299}
1300
1301#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1302#[allow(missing_docs)]
1303pub struct Conv1dWeightBackwardOpIr {
1304    pub x: TensorIr,
1305    pub weight: TensorIr,
1306    pub output_grad: TensorIr,
1307    pub options: Conv1dOptionsIr,
1308    pub out: TensorIr,
1309}
1310
1311#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1312#[allow(missing_docs)]
1313pub struct Conv1dBiasBackwardOpIr {
1314    pub x: TensorIr,
1315    pub bias: TensorIr,
1316    pub output_grad: TensorIr,
1317    pub out: TensorIr,
1318}
1319
1320#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1321#[allow(missing_docs)]
1322pub struct Conv2dOpIr {
1323    pub x: TensorIr,
1324    pub weight: TensorIr,
1325    pub bias: Option<TensorIr>,
1326    pub options: Conv2dOptionsIr,
1327    pub out: TensorIr,
1328}
1329
1330#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1331#[allow(missing_docs)]
1332pub struct Conv2dXBackwardOpIr {
1333    pub x: TensorIr,
1334    pub weight: TensorIr,
1335    pub output_grad: TensorIr,
1336    pub options: Conv2dOptionsIr,
1337    pub out: TensorIr,
1338}
1339
1340#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1341#[allow(missing_docs)]
1342pub struct Conv2dWeightBackwardOpIr {
1343    pub x: TensorIr,
1344    pub weight: TensorIr,
1345    pub output_grad: TensorIr,
1346    pub options: Conv2dOptionsIr,
1347    pub out: TensorIr,
1348}
1349
1350#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1351#[allow(missing_docs)]
1352pub struct Conv2dBiasBackwardOpIr {
1353    pub x: TensorIr,
1354    pub bias: TensorIr,
1355    pub output_grad: TensorIr,
1356    pub out: TensorIr,
1357}
1358
1359#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1360#[allow(missing_docs)]
1361pub struct DeformConv2dOpIr {
1362    pub x: TensorIr,
1363    pub offset: TensorIr,
1364    pub weight: TensorIr,
1365    pub mask: Option<TensorIr>,
1366    pub bias: Option<TensorIr>,
1367    pub options: DeformableConv2dOptionsIr,
1368    pub out: TensorIr,
1369}
1370
1371#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1372#[allow(missing_docs)]
1373pub struct DeformConv2dBackwardOpIr {
1374    pub x: TensorIr,
1375    pub offset: TensorIr,
1376    pub weight: TensorIr,
1377    pub mask: Option<TensorIr>,
1378    pub bias: Option<TensorIr>,
1379    pub out_grad: TensorIr,
1380    pub options: DeformableConv2dOptionsIr,
1381    pub input_grad: TensorIr,
1382    pub offset_grad: TensorIr,
1383    pub weight_grad: TensorIr,
1384    pub mask_grad: Option<TensorIr>,
1385    pub bias_grad: Option<TensorIr>,
1386}
1387
1388#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1389#[allow(missing_docs)]
1390pub struct Conv3dOpIr {
1391    pub x: TensorIr,
1392    pub weight: TensorIr,
1393    pub bias: Option<TensorIr>,
1394    pub options: Conv3dOptionsIr,
1395    pub out: TensorIr,
1396}
1397
1398#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1399#[allow(missing_docs)]
1400pub struct Conv3dXBackwardOpIr {
1401    pub x: TensorIr,
1402    pub weight: TensorIr,
1403    pub output_grad: TensorIr,
1404    pub options: Conv3dOptionsIr,
1405    pub out: TensorIr,
1406}
1407
1408#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1409#[allow(missing_docs)]
1410pub struct Conv3dWeightBackwardOpIr {
1411    pub x: TensorIr,
1412    pub weight: TensorIr,
1413    pub output_grad: TensorIr,
1414    pub options: Conv3dOptionsIr,
1415    pub out: TensorIr,
1416}
1417
1418#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1419#[allow(missing_docs)]
1420pub struct Conv3dBiasBackwardOpIr {
1421    pub x: TensorIr,
1422    pub bias: TensorIr,
1423    pub output_grad: TensorIr,
1424    pub out: TensorIr,
1425}
1426
1427#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1428#[allow(missing_docs)]
1429pub struct ConvTranspose1dOpIr {
1430    pub x: TensorIr,
1431    pub weight: TensorIr,
1432    pub bias: Option<TensorIr>,
1433    pub options: ConvTranspose1dOptionsIr,
1434    pub out: TensorIr,
1435}
1436
1437#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1438#[allow(missing_docs)]
1439pub struct ConvTranspose2dOpIr {
1440    pub x: TensorIr,
1441    pub weight: TensorIr,
1442    pub bias: Option<TensorIr>,
1443    pub options: ConvTranspose2dOptionsIr,
1444    pub out: TensorIr,
1445}
1446
1447#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1448#[allow(missing_docs)]
1449pub struct ConvTranspose3dOpIr {
1450    pub x: TensorIr,
1451    pub weight: TensorIr,
1452    pub bias: Option<TensorIr>,
1453    pub options: ConvTranspose3dOptionsIr,
1454    pub out: TensorIr,
1455}
1456
1457#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1458#[allow(missing_docs)]
1459pub struct Conv1dOptionsIr {
1460    pub stride: [usize; 1],
1461    pub padding: [usize; 1],
1462    pub dilation: [usize; 1],
1463    pub groups: usize,
1464}
1465
1466#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1467#[allow(missing_docs)]
1468pub struct Conv2dOptionsIr {
1469    pub stride: [usize; 2],
1470    pub padding: [usize; 2],
1471    pub dilation: [usize; 2],
1472    pub groups: usize,
1473}
1474
1475#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1476#[allow(missing_docs)]
1477pub struct DeformableConv2dOptionsIr {
1478    pub stride: [usize; 2],
1479    pub padding: [usize; 2],
1480    pub dilation: [usize; 2],
1481    pub weight_groups: usize,
1482    pub offset_groups: usize,
1483}
1484
1485#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1486#[allow(missing_docs)]
1487pub struct Conv3dOptionsIr {
1488    pub stride: [usize; 3],
1489    pub padding: [usize; 3],
1490    pub dilation: [usize; 3],
1491    pub groups: usize,
1492}
1493
1494#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1495#[allow(missing_docs)]
1496pub struct ConvTranspose1dOptionsIr {
1497    pub stride: [usize; 1],
1498    pub padding: [usize; 1],
1499    pub padding_out: [usize; 1],
1500    pub dilation: [usize; 1],
1501    pub groups: usize,
1502}
1503
1504#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1505#[allow(missing_docs)]
1506pub struct ConvTranspose2dOptionsIr {
1507    pub stride: [usize; 2],
1508    pub padding: [usize; 2],
1509    pub padding_out: [usize; 2],
1510    pub dilation: [usize; 2],
1511    pub groups: usize,
1512}
1513
1514#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1515#[allow(missing_docs)]
1516pub struct ConvTranspose3dOptionsIr {
1517    pub stride: [usize; 3],
1518    pub padding: [usize; 3],
1519    pub padding_out: [usize; 3],
1520    pub dilation: [usize; 3],
1521    pub groups: usize,
1522}
1523
1524/// Quantization parameters intermediate representation.
1525#[derive(Debug, Clone, Hash, PartialEq, Eq, Serialize, Deserialize)]
1526pub struct QuantizationParametersIr {
1527    /// The scaling factor, one per block or a single one for a per-tensor level.
1528    pub scales: TensorIr,
1529    /// The per-tensor scale that [`scales`](Self::scales) are expressed relative to, for a
1530    /// two-level scheme.
1531    pub global: Option<TensorIr>,
1532}
1533
1534#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1535#[allow(missing_docs)]
1536pub struct QuantizeOpIr {
1537    pub tensor: TensorIr,
1538    pub qparams: QuantizationParametersIr,
1539    pub scheme: QuantScheme,
1540    pub out: TensorIr,
1541}
1542
1543#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1544#[allow(missing_docs)]
1545pub struct DequantizeOpIr {
1546    pub input: TensorIr,
1547    pub out: TensorIr,
1548}
1549
1550impl From<ConvOptions<1>> for Conv1dOptionsIr {
1551    fn from(value: ConvOptions<1>) -> Self {
1552        Self {
1553            stride: value.stride,
1554            padding: value.padding,
1555            dilation: value.dilation,
1556            groups: value.groups,
1557        }
1558    }
1559}
1560
1561impl From<ConvOptions<2>> for Conv2dOptionsIr {
1562    fn from(value: ConvOptions<2>) -> Self {
1563        Self {
1564            stride: value.stride,
1565            padding: value.padding,
1566            dilation: value.dilation,
1567            groups: value.groups,
1568        }
1569    }
1570}
1571
1572impl From<ConvOptions<3>> for Conv3dOptionsIr {
1573    fn from(value: ConvOptions<3>) -> Self {
1574        Self {
1575            stride: value.stride,
1576            padding: value.padding,
1577            dilation: value.dilation,
1578            groups: value.groups,
1579        }
1580    }
1581}
1582
1583impl From<DeformConvOptions<2>> for DeformableConv2dOptionsIr {
1584    fn from(value: DeformConvOptions<2>) -> Self {
1585        Self {
1586            stride: value.stride,
1587            padding: value.padding,
1588            dilation: value.dilation,
1589            weight_groups: value.weight_groups,
1590            offset_groups: value.offset_groups,
1591        }
1592    }
1593}
1594
1595impl From<ConvTransposeOptions<1>> for ConvTranspose1dOptionsIr {
1596    fn from(value: ConvTransposeOptions<1>) -> Self {
1597        Self {
1598            stride: value.stride,
1599            padding: value.padding,
1600            padding_out: value.padding_out,
1601            dilation: value.dilation,
1602            groups: value.groups,
1603        }
1604    }
1605}
1606
1607impl From<ConvTransposeOptions<2>> for ConvTranspose2dOptionsIr {
1608    fn from(value: ConvTransposeOptions<2>) -> Self {
1609        Self {
1610            stride: value.stride,
1611            padding: value.padding,
1612            padding_out: value.padding_out,
1613            dilation: value.dilation,
1614            groups: value.groups,
1615        }
1616    }
1617}
1618
1619impl From<ConvTransposeOptions<3>> for ConvTranspose3dOptionsIr {
1620    fn from(value: ConvTransposeOptions<3>) -> Self {
1621        Self {
1622            stride: value.stride,
1623            padding: value.padding,
1624            padding_out: value.padding_out,
1625            dilation: value.dilation,
1626            groups: value.groups,
1627        }
1628    }
1629}
1630
1631impl From<Conv1dOptionsIr> for ConvOptions<1> {
1632    fn from(val: Conv1dOptionsIr) -> Self {
1633        ConvOptions {
1634            stride: val.stride,
1635            padding: val.padding,
1636            dilation: val.dilation,
1637            groups: val.groups,
1638        }
1639    }
1640}
1641
1642impl From<Conv2dOptionsIr> for ConvOptions<2> {
1643    fn from(val: Conv2dOptionsIr) -> Self {
1644        ConvOptions {
1645            stride: val.stride,
1646            padding: val.padding,
1647            dilation: val.dilation,
1648            groups: val.groups,
1649        }
1650    }
1651}
1652
1653impl From<Conv3dOptionsIr> for ConvOptions<3> {
1654    fn from(val: Conv3dOptionsIr) -> Self {
1655        ConvOptions {
1656            stride: val.stride,
1657            padding: val.padding,
1658            dilation: val.dilation,
1659            groups: val.groups,
1660        }
1661    }
1662}
1663
1664impl From<DeformableConv2dOptionsIr> for DeformConvOptions<2> {
1665    fn from(value: DeformableConv2dOptionsIr) -> Self {
1666        DeformConvOptions {
1667            stride: value.stride,
1668            padding: value.padding,
1669            dilation: value.dilation,
1670            weight_groups: value.weight_groups,
1671            offset_groups: value.offset_groups,
1672        }
1673    }
1674}
1675
1676impl From<ConvTranspose1dOptionsIr> for ConvTransposeOptions<1> {
1677    fn from(val: ConvTranspose1dOptionsIr) -> Self {
1678        ConvTransposeOptions {
1679            stride: val.stride,
1680            padding: val.padding,
1681            padding_out: val.padding_out,
1682            dilation: val.dilation,
1683            groups: val.groups,
1684        }
1685    }
1686}
1687
1688impl From<ConvTranspose2dOptionsIr> for ConvTransposeOptions<2> {
1689    fn from(val: ConvTranspose2dOptionsIr) -> Self {
1690        ConvTransposeOptions {
1691            stride: val.stride,
1692            padding: val.padding,
1693            padding_out: val.padding_out,
1694            dilation: val.dilation,
1695            groups: val.groups,
1696        }
1697    }
1698}
1699
1700impl From<ConvTranspose3dOptionsIr> for ConvTransposeOptions<3> {
1701    fn from(val: ConvTranspose3dOptionsIr) -> Self {
1702        ConvTransposeOptions {
1703            stride: val.stride,
1704            padding: val.padding,
1705            padding_out: val.padding_out,
1706            dilation: val.dilation,
1707            groups: val.groups,
1708        }
1709    }
1710}
1711
1712#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1713#[allow(missing_docs)]
1714pub struct AvgPool1dOpIr {
1715    pub x: TensorIr,
1716    pub kernel_size: usize,
1717    pub stride: usize,
1718    pub padding: usize,
1719    pub count_include_pad: bool,
1720    pub ceil_mode: bool,
1721    pub out: TensorIr,
1722}
1723
1724#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1725#[allow(missing_docs)]
1726pub struct AvgPool2dOpIr {
1727    pub x: TensorIr,
1728    pub kernel_size: [usize; 2],
1729    pub stride: [usize; 2],
1730    pub padding: [usize; 2],
1731    pub count_include_pad: bool,
1732    pub ceil_mode: bool,
1733    pub out: TensorIr,
1734}
1735
1736#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1737#[allow(missing_docs)]
1738pub struct AvgPool1dBackwardOpIr {
1739    pub x: TensorIr,
1740    pub grad: TensorIr,
1741    pub kernel_size: usize,
1742    pub stride: usize,
1743    pub padding: usize,
1744    pub count_include_pad: bool,
1745    pub ceil_mode: bool,
1746    pub out: TensorIr,
1747}
1748
1749#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1750#[allow(missing_docs)]
1751pub struct AvgPool2dBackwardOpIr {
1752    pub x: TensorIr,
1753    pub grad: TensorIr,
1754    pub kernel_size: [usize; 2],
1755    pub stride: [usize; 2],
1756    pub padding: [usize; 2],
1757    pub count_include_pad: bool,
1758    pub ceil_mode: bool,
1759    pub out: TensorIr,
1760}
1761
1762#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1763#[allow(missing_docs)]
1764pub struct AdaptiveAvgPool1dOpIr {
1765    pub x: TensorIr,
1766    pub output_size: usize,
1767    pub out: TensorIr,
1768}
1769
1770#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1771#[allow(missing_docs)]
1772pub struct AdaptiveAvgPool2dOpIr {
1773    pub x: TensorIr,
1774    pub output_size: [usize; 2],
1775    pub out: TensorIr,
1776}
1777
1778#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1779#[allow(missing_docs)]
1780pub struct AdaptiveAvgPool1dBackwardOpIr {
1781    pub x: TensorIr,
1782    pub grad: TensorIr,
1783    pub out: TensorIr,
1784}
1785
1786#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1787#[allow(missing_docs)]
1788pub struct AdaptiveAvgPool2dBackwardOpIr {
1789    pub x: TensorIr,
1790    pub grad: TensorIr,
1791    pub out: TensorIr,
1792}
1793
1794#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1795#[allow(missing_docs)]
1796pub struct AdaptiveAvgPool3dOpIr {
1797    pub x: TensorIr,
1798    pub output_size: [usize; 3],
1799    pub out: TensorIr,
1800}
1801
1802#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1803#[allow(missing_docs)]
1804pub struct AdaptiveAvgPool3dBackwardOpIr {
1805    pub x: TensorIr,
1806    pub grad: TensorIr,
1807    pub out: TensorIr,
1808}
1809
1810#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1811#[allow(missing_docs)]
1812pub struct MaxPool1dOpIr {
1813    pub x: TensorIr,
1814    pub kernel_size: usize,
1815    pub stride: usize,
1816    pub padding: usize,
1817    pub dilation: usize,
1818    pub ceil_mode: bool,
1819    pub out: TensorIr,
1820}
1821
1822#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1823#[allow(missing_docs)]
1824pub struct MaxPool1dWithIndicesOpIr {
1825    pub x: TensorIr,
1826    pub kernel_size: usize,
1827    pub stride: usize,
1828    pub padding: usize,
1829    pub dilation: usize,
1830    pub ceil_mode: bool,
1831    pub out: TensorIr,
1832    pub out_indices: TensorIr,
1833}
1834
1835#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1836#[allow(missing_docs)]
1837pub struct MaxPool1dWithIndicesBackwardOpIr {
1838    pub x: TensorIr,
1839    pub grad: TensorIr,
1840    pub indices: TensorIr,
1841    pub kernel_size: usize,
1842    pub stride: usize,
1843    pub padding: usize,
1844    pub dilation: usize,
1845    pub ceil_mode: bool,
1846    pub out: TensorIr,
1847}
1848
1849#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1850#[allow(missing_docs)]
1851pub struct MaxPool2dOpIr {
1852    pub x: TensorIr,
1853    pub kernel_size: [usize; 2],
1854    pub stride: [usize; 2],
1855    pub padding: [usize; 2],
1856    pub dilation: [usize; 2],
1857    pub ceil_mode: bool,
1858    pub out: TensorIr,
1859}
1860
1861#[allow(missing_docs)]
1862#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1863pub struct MaxPool2dWithIndicesOpIr {
1864    pub x: TensorIr,
1865    pub kernel_size: [usize; 2],
1866    pub stride: [usize; 2],
1867    pub padding: [usize; 2],
1868    pub dilation: [usize; 2],
1869    pub ceil_mode: bool,
1870    pub out: TensorIr,
1871    pub out_indices: TensorIr,
1872}
1873
1874#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1875#[allow(missing_docs)]
1876pub struct MaxPool2dWithIndicesBackwardOpIr {
1877    pub x: TensorIr,
1878    pub grad: TensorIr,
1879    pub indices: TensorIr,
1880    pub kernel_size: [usize; 2],
1881    pub stride: [usize; 2],
1882    pub padding: [usize; 2],
1883    pub dilation: [usize; 2],
1884    pub ceil_mode: bool,
1885    pub out: TensorIr,
1886}
1887
1888#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1889#[allow(missing_docs)]
1890pub enum InterpolateModeIr {
1891    Nearest,
1892    NearestExact,
1893    Bilinear,
1894    Bicubic,
1895    Lanczos3,
1896}
1897
1898#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1899#[allow(missing_docs)]
1900pub struct InterpolateOptionsIr {
1901    pub mode: InterpolateModeIr,
1902    pub align_corners: bool,
1903}
1904
1905#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1906#[allow(missing_docs)]
1907pub struct InterpolateOpIr {
1908    pub x: TensorIr,
1909    pub output_size: [usize; 2],
1910    pub options: InterpolateOptionsIr,
1911    pub out: TensorIr,
1912}
1913
1914#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1915#[allow(missing_docs)]
1916pub struct RfftOpIr {
1917    pub signal: TensorIr,
1918    pub dim: usize,
1919    pub n: Option<usize>,
1920    pub out_re: TensorIr,
1921    pub out_im: TensorIr,
1922}
1923
1924#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1925#[allow(missing_docs)]
1926pub struct IRfftOpIr {
1927    pub input_re: TensorIr,
1928    pub input_im: TensorIr,
1929    pub dim: usize,
1930    pub n: Option<usize>,
1931    pub out_signal: TensorIr,
1932}
1933
1934#[allow(missing_docs)]
1935impl RfftOpIr {
1936    pub fn create<F>(signal: TensorIr, dim: usize, n: Option<usize>, mut new_id: F) -> Self
1937    where
1938        F: FnMut() -> crate::TensorId,
1939    {
1940        // `n` is required to be a power of two at the public API boundary, so
1941        // the output has `n / 2 + 1` bins (matching scipy/torch for pow2 n).
1942        let mut shape = signal.shape.clone();
1943        let fft_len = n.unwrap_or(shape[dim]);
1944        shape[dim] = fft_len / 2 + 1;
1945        let dtype = signal.dtype;
1946
1947        Self {
1948            signal,
1949            dim,
1950            n,
1951            out_re: TensorIr::uninit(new_id(), shape.clone(), dtype),
1952            out_im: TensorIr::uninit(new_id(), shape, dtype),
1953        }
1954    }
1955}
1956
1957#[allow(missing_docs)]
1958impl IRfftOpIr {
1959    pub fn create<F>(
1960        input_re: TensorIr,
1961        input_im: TensorIr,
1962        dim: usize,
1963        n: Option<usize>,
1964        mut new_id: F,
1965    ) -> Self
1966    where
1967        F: FnMut() -> crate::TensorId,
1968    {
1969        debug_assert!(
1970            input_re.shape[dim] >= 1,
1971            "IRfftOpIr: input spectrum dimension must be >= 1"
1972        );
1973        debug_assert!(
1974            !matches!(n, Some(0)),
1975            "IRfftOpIr: n must be >= 1 when specified"
1976        );
1977        let mut shape = input_re.shape.clone();
1978        shape[dim] = n.unwrap_or((shape[dim] - 1) * 2);
1979        let dtype = input_re.dtype;
1980
1981        Self {
1982            input_re,
1983            input_im,
1984            dim,
1985            n,
1986            out_signal: TensorIr::uninit(new_id(), shape, dtype),
1987        }
1988    }
1989}
1990
1991#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
1992#[allow(missing_docs)]
1993pub struct AttentionOptionsIr {
1994    pub scale: Option<ScalarIr>,
1995    pub softcap: Option<ScalarIr>,
1996    pub is_causal: bool,
1997}
1998
1999impl From<AttentionOptionsIr> for AttentionModuleOptions {
2000    fn from(ir: AttentionOptionsIr) -> Self {
2001        AttentionModuleOptions {
2002            scale: ir.scale.map(|s| s.elem()),
2003            softcap: ir.softcap.map(|s| s.elem()),
2004            is_causal: ir.is_causal,
2005        }
2006    }
2007}
2008
2009impl From<AttentionModuleOptions> for AttentionOptionsIr {
2010    fn from(ir: AttentionModuleOptions) -> Self {
2011        AttentionOptionsIr {
2012            scale: ir.scale.map(ScalarIr::Float),
2013            softcap: ir.softcap.map(ScalarIr::Float),
2014            is_causal: ir.is_causal,
2015        }
2016    }
2017}
2018
2019#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
2020#[allow(missing_docs)]
2021pub struct AttentionOpIr {
2022    pub query: TensorIr,
2023    pub key: TensorIr,
2024    pub value: TensorIr,
2025    pub mask: Option<TensorIr>,
2026    pub attn_bias: Option<TensorIr>,
2027    pub options: AttentionOptionsIr,
2028    pub out: TensorIr,
2029}
2030
2031#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
2032#[allow(missing_docs)]
2033pub struct CtcLossOpIr {
2034    pub log_probs: TensorIr,
2035    pub targets: TensorIr,
2036    pub input_lengths: TensorIr,
2037    pub target_lengths: TensorIr,
2038    pub blank: usize,
2039    pub out: TensorIr,
2040}
2041
2042#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
2043#[allow(missing_docs)]
2044pub struct CtcLossBackwardOpIr {
2045    pub log_probs: TensorIr,
2046    pub targets: TensorIr,
2047    pub input_lengths: TensorIr,
2048    pub target_lengths: TensorIr,
2049    pub grad_loss: TensorIr,
2050    pub blank: usize,
2051    pub out: TensorIr,
2052}
2053
2054impl From<InterpolateModeIr> for InterpolateMode {
2055    fn from(val: InterpolateModeIr) -> Self {
2056        match val {
2057            InterpolateModeIr::Nearest => Self::Nearest,
2058            InterpolateModeIr::NearestExact => Self::NearestExact,
2059            InterpolateModeIr::Bilinear => Self::Bilinear,
2060            InterpolateModeIr::Bicubic => Self::Bicubic,
2061            InterpolateModeIr::Lanczos3 => Self::Lanczos3,
2062        }
2063    }
2064}
2065
2066impl From<InterpolateOptionsIr> for InterpolateOptions {
2067    fn from(val: InterpolateOptionsIr) -> Self {
2068        Self::new(val.mode.into()).with_align_corners(val.align_corners)
2069    }
2070}
2071
2072impl From<InterpolateMode> for InterpolateModeIr {
2073    fn from(val: InterpolateMode) -> Self {
2074        match val {
2075            InterpolateMode::Nearest => Self::Nearest,
2076            InterpolateMode::NearestExact => Self::NearestExact,
2077            InterpolateMode::Bilinear => Self::Bilinear,
2078            InterpolateMode::Bicubic => Self::Bicubic,
2079            InterpolateMode::Lanczos3 => Self::Lanczos3,
2080        }
2081    }
2082}
2083
2084impl From<InterpolateOptions> for InterpolateOptionsIr {
2085    fn from(val: InterpolateOptions) -> Self {
2086        Self {
2087            mode: val.mode.into(),
2088            align_corners: val.align_corners,
2089        }
2090    }
2091}
2092
2093#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
2094#[allow(missing_docs)]
2095pub struct InterpolateBackwardOpIr {
2096    pub x: TensorIr,
2097    pub grad: TensorIr,
2098    pub output_size: [usize; 2],
2099    pub options: InterpolateOptionsIr,
2100    pub out: TensorIr,
2101}
2102
2103#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
2104#[allow(missing_docs)]
2105pub enum GridSamplePaddingModeIr {
2106    Zeros,
2107    Border,
2108    Reflection,
2109}
2110
2111#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
2112#[allow(missing_docs)]
2113pub struct GridSampleOptionsIr {
2114    pub mode: InterpolateModeIr,
2115    pub padding_mode: GridSamplePaddingModeIr,
2116    pub align_corners: bool,
2117}
2118
2119#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
2120#[allow(missing_docs)]
2121pub struct GridSample2dOpIr {
2122    pub tensor: TensorIr,
2123    pub grid: TensorIr,
2124    pub options: GridSampleOptionsIr,
2125    pub out: TensorIr,
2126}
2127
2128impl From<GridSamplePaddingModeIr> for GridSamplePaddingMode {
2129    fn from(val: GridSamplePaddingModeIr) -> Self {
2130        match val {
2131            GridSamplePaddingModeIr::Zeros => Self::Zeros,
2132            GridSamplePaddingModeIr::Border => Self::Border,
2133            GridSamplePaddingModeIr::Reflection => Self::Reflection,
2134        }
2135    }
2136}
2137
2138impl From<GridSamplePaddingMode> for GridSamplePaddingModeIr {
2139    fn from(val: GridSamplePaddingMode) -> Self {
2140        match val {
2141            GridSamplePaddingMode::Zeros => Self::Zeros,
2142            GridSamplePaddingMode::Border => Self::Border,
2143            GridSamplePaddingMode::Reflection => Self::Reflection,
2144        }
2145    }
2146}
2147
2148impl From<GridSampleOptionsIr> for GridSampleOptions {
2149    fn from(val: GridSampleOptionsIr) -> Self {
2150        Self {
2151            mode: val.mode.into(),
2152            padding_mode: val.padding_mode.into(),
2153            align_corners: val.align_corners,
2154        }
2155    }
2156}
2157
2158impl From<GridSampleOptions> for GridSampleOptionsIr {
2159    fn from(val: GridSampleOptions) -> Self {
2160        Self {
2161            mode: val.mode.into(),
2162            padding_mode: val.padding_mode.into(),
2163            align_corners: val.align_corners,
2164        }
2165    }
2166}
2167
2168impl OperationIr {
2169    /// Get all input [tensors](TensorIr) involved with the current operation.
2170    pub fn inputs(&self) -> impl Iterator<Item = &TensorIr> {
2171        match self {
2172            OperationIr::BaseFloat(repr) => repr.inputs(),
2173            OperationIr::BaseInt(repr) => repr.inputs(),
2174            OperationIr::BaseBool(repr) => repr.inputs(),
2175            OperationIr::NumericFloat(_dtype, repr) => repr.inputs(),
2176            OperationIr::NumericInt(_dtype, repr) => repr.inputs(),
2177            OperationIr::Bool(repr) => repr.inputs(),
2178            OperationIr::Int(repr) => repr.inputs(),
2179            OperationIr::Float(_dtype, repr) => repr.inputs(),
2180            OperationIr::Module(repr) => repr.inputs(),
2181            OperationIr::Init(repr) => repr.inputs(),
2182            OperationIr::Custom(repr) => repr.inputs(),
2183            OperationIr::Drop(repr) => Box::new([repr].into_iter()),
2184            OperationIr::Distributed(repr) => repr.inputs(),
2185            OperationIr::Activation(repr) => repr.inputs(),
2186        }
2187    }
2188
2189    /// Get all output [tensors](TensorIr) involved with the current operation.
2190    pub fn outputs(&self) -> impl Iterator<Item = &TensorIr> {
2191        match self {
2192            OperationIr::BaseFloat(repr) => repr.outputs(),
2193            OperationIr::BaseInt(repr) => repr.outputs(),
2194            OperationIr::BaseBool(repr) => repr.outputs(),
2195            OperationIr::NumericFloat(_dtype, repr) => repr.outputs(),
2196            OperationIr::NumericInt(_dtype, repr) => repr.outputs(),
2197            OperationIr::Bool(repr) => repr.outputs(),
2198            OperationIr::Int(repr) => repr.outputs(),
2199            OperationIr::Float(_dtype, repr) => repr.outputs(),
2200            OperationIr::Module(repr) => repr.outputs(),
2201            OperationIr::Init(repr) => repr.outputs(),
2202            OperationIr::Custom(repr) => repr.outputs(),
2203            OperationIr::Drop(_repr) => Box::new([].into_iter()),
2204            OperationIr::Distributed(repr) => repr.outputs(),
2205            OperationIr::Activation(repr) => repr.outputs(),
2206        }
2207    }
2208
2209    /// Get all [tensor](TensorIr) involved with the current operation.
2210    pub fn nodes(&self) -> Vec<&TensorIr> {
2211        self.inputs().chain(self.outputs()).collect()
2212    }
2213
2214    /// Set the given nodes that are [read write](super::TensorStatus::ReadWrite) to
2215    /// [read only](super::TensorStatus::ReadOnly) in the current operation.
2216    ///
2217    /// Returns the tensor that were updated with their original representation.
2218    pub fn mark_read_only(&mut self, nodes: &[TensorId]) -> Vec<TensorIr> {
2219        match self {
2220            OperationIr::BaseFloat(repr) => repr.mark_read_only(nodes),
2221            OperationIr::BaseInt(repr) => repr.mark_read_only(nodes),
2222            OperationIr::BaseBool(repr) => repr.mark_read_only(nodes),
2223            OperationIr::NumericFloat(_dtype, repr) => repr.mark_read_only(nodes),
2224            OperationIr::NumericInt(_dtype, repr) => repr.mark_read_only(nodes),
2225            OperationIr::Bool(repr) => repr.mark_read_only(nodes),
2226            OperationIr::Int(repr) => repr.mark_read_only(nodes),
2227            OperationIr::Float(_dtype, repr) => repr.mark_read_only(nodes),
2228            OperationIr::Module(repr) => repr.mark_read_only(nodes),
2229            OperationIr::Init(_) => Vec::new(),
2230            OperationIr::Drop(repr) => {
2231                let mut output = Vec::new();
2232                repr.mark_read_only(nodes, &mut output);
2233                output
2234            }
2235            OperationIr::Custom(repr) => {
2236                let mut output = Vec::new();
2237
2238                for input in repr.inputs.iter_mut() {
2239                    input.mark_read_only(nodes, &mut output);
2240                }
2241
2242                output
2243            }
2244            OperationIr::Distributed(repr) => repr.mark_read_only(nodes),
2245            OperationIr::Activation(repr) => repr.mark_read_only(nodes),
2246        }
2247    }
2248
2249    /// Visit every mutable component of this operation in place.
2250    ///
2251    /// For each variant, visits its [tensors](TensorIr) (in the same order as
2252    /// [`inputs`](Self::inputs) then [`outputs`](Self::outputs) combined), then its
2253    /// [scalars](ScalarIr), then its slice ranges. Only `Slice`/`SliceAssign` (under the `Base*`
2254    /// variants) carry ranges; every other category visits no ranges.
2255    pub fn visit_mut(&mut self, v: &mut impl IrVisitorMut) {
2256        match self {
2257            OperationIr::BaseFloat(repr) => repr.visit_mut(v),
2258            OperationIr::BaseInt(repr) => repr.visit_mut(v),
2259            OperationIr::BaseBool(repr) => repr.visit_mut(v),
2260            OperationIr::NumericFloat(_dtype, repr) => repr.visit_mut(v),
2261            OperationIr::NumericInt(_dtype, repr) => repr.visit_mut(v),
2262            OperationIr::Bool(repr) => repr.visit_mut(v),
2263            OperationIr::Int(repr) => repr.visit_mut(v),
2264            OperationIr::Float(_dtype, repr) => repr.visit_mut(v),
2265            OperationIr::Module(repr) => repr.visit_mut(v),
2266            OperationIr::Init(repr) => repr.visit_mut(v),
2267            OperationIr::Custom(repr) => repr.visit_mut(v),
2268            OperationIr::Drop(repr) => v.visit_tensor_mut(repr),
2269            OperationIr::Distributed(repr) => repr.visit_mut(v),
2270            OperationIr::Activation(repr) => repr.visit_mut(v),
2271        }
2272    }
2273}
2274
2275impl BaseOperationIr {
2276    fn inputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
2277        match self {
2278            BaseOperationIr::Reshape(repr) => Box::new([&repr.input].into_iter()),
2279            BaseOperationIr::SwapDims(repr) => Box::new([&repr.input].into_iter()),
2280            BaseOperationIr::Permute(repr) => Box::new([&repr.input].into_iter()),
2281            BaseOperationIr::Expand(repr) => Box::new([&repr.input].into_iter()),
2282            BaseOperationIr::Flip(repr) => Box::new([&repr.input].into_iter()),
2283            BaseOperationIr::Slice(repr) => Box::new([&repr.tensor].into_iter()),
2284            BaseOperationIr::SliceAssign(repr) => Box::new([&repr.tensor, &repr.value].into_iter()),
2285            BaseOperationIr::Gather(repr) => Box::new([&repr.tensor, &repr.indices].into_iter()),
2286            BaseOperationIr::Scatter(repr) => {
2287                Box::new([&repr.tensor, &repr.indices, &repr.value].into_iter())
2288            }
2289            BaseOperationIr::ScatterNd(repr) => {
2290                Box::new([&repr.data, &repr.indices, &repr.values].into_iter())
2291            }
2292            BaseOperationIr::GatherNd(repr) => Box::new([&repr.data, &repr.indices].into_iter()),
2293            BaseOperationIr::Select(repr) => Box::new([&repr.tensor, &repr.indices].into_iter()),
2294            BaseOperationIr::SelectAssign(repr) => {
2295                Box::new([&repr.tensor, &repr.indices, &repr.value].into_iter())
2296            }
2297            BaseOperationIr::MaskWhere(repr) => {
2298                Box::new([&repr.tensor, &repr.mask, &repr.value].into_iter())
2299            }
2300            BaseOperationIr::MaskFill(repr) => Box::new([&repr.tensor, &repr.mask].into_iter()),
2301            BaseOperationIr::Equal(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2302            BaseOperationIr::EqualElem(repr) => Box::new([&repr.lhs].into_iter()),
2303            BaseOperationIr::RepeatDim(repr) => Box::new([&repr.tensor].into_iter()),
2304            BaseOperationIr::Cat(repr) => Box::new(repr.tensors.iter()),
2305            BaseOperationIr::Cast(repr) => Box::new([&repr.input].into_iter()),
2306            BaseOperationIr::Unfold(repr) => Box::new([&repr.input].into_iter()),
2307            BaseOperationIr::Empty(_repr) => Box::new([].into_iter()),
2308            BaseOperationIr::Ones(_repr) => Box::new([].into_iter()),
2309            BaseOperationIr::Zeros(_repr) => Box::new([].into_iter()),
2310            BaseOperationIr::NotEqual(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2311            BaseOperationIr::NotEqualElem(repr) => Box::new([&repr.lhs].into_iter()),
2312            BaseOperationIr::All(repr) => Box::new([&repr.input].into_iter()),
2313            BaseOperationIr::Any(repr) => Box::new([&repr.input].into_iter()),
2314            BaseOperationIr::AllDim(repr) => Box::new([&repr.input].into_iter()),
2315            BaseOperationIr::AnyDim(repr) => Box::new([&repr.input].into_iter()),
2316        }
2317    }
2318
2319    fn outputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
2320        match self {
2321            BaseOperationIr::Reshape(repr) => Box::new([&repr.out].into_iter()),
2322            BaseOperationIr::SwapDims(repr) => Box::new([&repr.out].into_iter()),
2323            BaseOperationIr::Permute(repr) => Box::new([&repr.out].into_iter()),
2324            BaseOperationIr::Expand(repr) => Box::new([&repr.out].into_iter()),
2325            BaseOperationIr::Flip(repr) => Box::new([&repr.out].into_iter()),
2326            BaseOperationIr::Slice(repr) => Box::new([&repr.out].into_iter()),
2327            BaseOperationIr::SliceAssign(repr) => Box::new([&repr.out].into_iter()),
2328            BaseOperationIr::Gather(repr) => Box::new([&repr.out].into_iter()),
2329            BaseOperationIr::Scatter(repr) => Box::new([&repr.out].into_iter()),
2330            BaseOperationIr::ScatterNd(repr) => Box::new([&repr.out].into_iter()),
2331            BaseOperationIr::GatherNd(repr) => Box::new([&repr.out].into_iter()),
2332            BaseOperationIr::Select(repr) => Box::new([&repr.out].into_iter()),
2333            BaseOperationIr::SelectAssign(repr) => Box::new([&repr.out].into_iter()),
2334            BaseOperationIr::MaskWhere(repr) => Box::new([&repr.out].into_iter()),
2335            BaseOperationIr::MaskFill(repr) => Box::new([&repr.out].into_iter()),
2336            BaseOperationIr::Equal(repr) => Box::new([&repr.out].into_iter()),
2337            BaseOperationIr::EqualElem(repr) => Box::new([&repr.out].into_iter()),
2338            BaseOperationIr::RepeatDim(repr) => Box::new([&repr.out].into_iter()),
2339            BaseOperationIr::Cat(repr) => Box::new([&repr.out].into_iter()),
2340            BaseOperationIr::Cast(repr) => Box::new([&repr.out].into_iter()),
2341            BaseOperationIr::Unfold(repr) => Box::new([&repr.out].into_iter()),
2342            BaseOperationIr::Empty(repr) => Box::new([&repr.out].into_iter()),
2343            BaseOperationIr::Ones(repr) => Box::new([&repr.out].into_iter()),
2344            BaseOperationIr::Zeros(repr) => Box::new([&repr.out].into_iter()),
2345            BaseOperationIr::NotEqual(repr) => Box::new([&repr.out].into_iter()),
2346            BaseOperationIr::NotEqualElem(repr) => Box::new([&repr.out].into_iter()),
2347            BaseOperationIr::All(repr) => Box::new([&repr.out].into_iter()),
2348            BaseOperationIr::Any(repr) => Box::new([&repr.out].into_iter()),
2349            BaseOperationIr::AllDim(repr) => Box::new([&repr.out].into_iter()),
2350            BaseOperationIr::AnyDim(repr) => Box::new([&repr.out].into_iter()),
2351        }
2352    }
2353
2354    fn mark_read_only(&mut self, nodes: &[TensorId]) -> Vec<TensorIr> {
2355        let mut output = Vec::new();
2356
2357        match self {
2358            BaseOperationIr::Reshape(repr) => {
2359                repr.input.mark_read_only(nodes, &mut output);
2360            }
2361            BaseOperationIr::SwapDims(repr) => {
2362                repr.input.mark_read_only(nodes, &mut output);
2363            }
2364            BaseOperationIr::Permute(repr) => {
2365                repr.input.mark_read_only(nodes, &mut output);
2366            }
2367
2368            BaseOperationIr::Expand(repr) => {
2369                repr.input.mark_read_only(nodes, &mut output);
2370            }
2371
2372            BaseOperationIr::Flip(repr) => {
2373                repr.input.mark_read_only(nodes, &mut output);
2374            }
2375            BaseOperationIr::Slice(repr) => {
2376                repr.tensor.mark_read_only(nodes, &mut output);
2377            }
2378            BaseOperationIr::SliceAssign(repr) => {
2379                repr.tensor.mark_read_only(nodes, &mut output);
2380                repr.value.mark_read_only(nodes, &mut output);
2381            }
2382            BaseOperationIr::Gather(repr) => {
2383                repr.tensor.mark_read_only(nodes, &mut output);
2384                repr.indices.mark_read_only(nodes, &mut output);
2385            }
2386            BaseOperationIr::Scatter(repr) => {
2387                repr.tensor.mark_read_only(nodes, &mut output);
2388                repr.indices.mark_read_only(nodes, &mut output);
2389                repr.value.mark_read_only(nodes, &mut output);
2390            }
2391            BaseOperationIr::ScatterNd(repr) => {
2392                repr.data.mark_read_only(nodes, &mut output);
2393                repr.indices.mark_read_only(nodes, &mut output);
2394                repr.values.mark_read_only(nodes, &mut output);
2395            }
2396            BaseOperationIr::GatherNd(repr) => {
2397                repr.data.mark_read_only(nodes, &mut output);
2398                repr.indices.mark_read_only(nodes, &mut output);
2399            }
2400            BaseOperationIr::Select(repr) => {
2401                repr.tensor.mark_read_only(nodes, &mut output);
2402                repr.indices.mark_read_only(nodes, &mut output);
2403            }
2404            BaseOperationIr::SelectAssign(repr) => {
2405                repr.tensor.mark_read_only(nodes, &mut output);
2406                repr.indices.mark_read_only(nodes, &mut output);
2407                repr.value.mark_read_only(nodes, &mut output);
2408            }
2409            BaseOperationIr::MaskWhere(repr) => {
2410                repr.tensor.mark_read_only(nodes, &mut output);
2411                repr.mask.mark_read_only(nodes, &mut output);
2412                repr.value.mark_read_only(nodes, &mut output);
2413            }
2414            BaseOperationIr::MaskFill(repr) => {
2415                repr.tensor.mark_read_only(nodes, &mut output);
2416                repr.mask.mark_read_only(nodes, &mut output);
2417            }
2418            BaseOperationIr::Equal(repr) => {
2419                repr.lhs.mark_read_only(nodes, &mut output);
2420                repr.rhs.mark_read_only(nodes, &mut output);
2421            }
2422            BaseOperationIr::EqualElem(repr) => {
2423                repr.lhs.mark_read_only(nodes, &mut output);
2424            }
2425            BaseOperationIr::RepeatDim(repr) => {
2426                repr.tensor.mark_read_only(nodes, &mut output);
2427            }
2428            BaseOperationIr::Cat(repr) => {
2429                for t in repr.tensors.iter_mut() {
2430                    t.mark_read_only(nodes, &mut output);
2431                }
2432            }
2433            BaseOperationIr::Cast(repr) => {
2434                repr.input.mark_read_only(nodes, &mut output);
2435            }
2436            BaseOperationIr::Unfold(repr) => {
2437                repr.input.mark_read_only(nodes, &mut output);
2438            }
2439            BaseOperationIr::Empty(_) => {}
2440            BaseOperationIr::Zeros(_) => {}
2441            BaseOperationIr::Ones(_) => {}
2442            BaseOperationIr::NotEqual(repr) => {
2443                repr.lhs.mark_read_only(nodes, &mut output);
2444                repr.rhs.mark_read_only(nodes, &mut output);
2445            }
2446            BaseOperationIr::NotEqualElem(repr) => {
2447                repr.lhs.mark_read_only(nodes, &mut output);
2448            }
2449            BaseOperationIr::All(repr) => {
2450                repr.input.mark_read_only(nodes, &mut output);
2451            }
2452            BaseOperationIr::Any(repr) => {
2453                repr.input.mark_read_only(nodes, &mut output);
2454            }
2455            BaseOperationIr::AllDim(repr) => {
2456                repr.input.mark_read_only(nodes, &mut output);
2457            }
2458            BaseOperationIr::AnyDim(repr) => {
2459                repr.input.mark_read_only(nodes, &mut output);
2460            }
2461        };
2462
2463        output
2464    }
2465
2466    fn visit_mut(&mut self, v: &mut impl IrVisitorMut) {
2467        match self {
2468            BaseOperationIr::Reshape(repr) => {
2469                v.visit_tensor_mut(&mut repr.input);
2470                v.visit_tensor_mut(&mut repr.out);
2471            }
2472            BaseOperationIr::SwapDims(repr) => {
2473                v.visit_tensor_mut(&mut repr.input);
2474                v.visit_tensor_mut(&mut repr.out);
2475            }
2476            BaseOperationIr::Permute(repr) => {
2477                v.visit_tensor_mut(&mut repr.input);
2478                v.visit_tensor_mut(&mut repr.out);
2479            }
2480            BaseOperationIr::Expand(repr) => {
2481                v.visit_tensor_mut(&mut repr.input);
2482                v.visit_tensor_mut(&mut repr.out);
2483            }
2484            BaseOperationIr::Flip(repr) => {
2485                v.visit_tensor_mut(&mut repr.input);
2486                v.visit_tensor_mut(&mut repr.out);
2487            }
2488            BaseOperationIr::Slice(repr) => {
2489                v.visit_tensor_mut(&mut repr.tensor);
2490                v.visit_tensor_mut(&mut repr.out);
2491                repr.ranges.iter_mut().for_each(|r| v.visit_range_mut(r));
2492            }
2493            BaseOperationIr::SliceAssign(repr) => {
2494                v.visit_tensor_mut(&mut repr.tensor);
2495                v.visit_tensor_mut(&mut repr.value);
2496                v.visit_tensor_mut(&mut repr.out);
2497                repr.ranges.iter_mut().for_each(|r| v.visit_range_mut(r));
2498            }
2499            BaseOperationIr::Gather(repr) => {
2500                v.visit_tensor_mut(&mut repr.tensor);
2501                v.visit_tensor_mut(&mut repr.indices);
2502                v.visit_tensor_mut(&mut repr.out);
2503            }
2504            BaseOperationIr::Scatter(repr) => {
2505                v.visit_tensor_mut(&mut repr.tensor);
2506                v.visit_tensor_mut(&mut repr.indices);
2507                v.visit_tensor_mut(&mut repr.value);
2508                v.visit_tensor_mut(&mut repr.out);
2509            }
2510            BaseOperationIr::ScatterNd(repr) => {
2511                v.visit_tensor_mut(&mut repr.data);
2512                v.visit_tensor_mut(&mut repr.indices);
2513                v.visit_tensor_mut(&mut repr.values);
2514                v.visit_tensor_mut(&mut repr.out);
2515            }
2516            BaseOperationIr::GatherNd(repr) => {
2517                v.visit_tensor_mut(&mut repr.data);
2518                v.visit_tensor_mut(&mut repr.indices);
2519                v.visit_tensor_mut(&mut repr.out);
2520            }
2521            BaseOperationIr::Select(repr) => {
2522                v.visit_tensor_mut(&mut repr.tensor);
2523                v.visit_tensor_mut(&mut repr.indices);
2524                v.visit_tensor_mut(&mut repr.out);
2525            }
2526            BaseOperationIr::SelectAssign(repr) => {
2527                v.visit_tensor_mut(&mut repr.tensor);
2528                v.visit_tensor_mut(&mut repr.indices);
2529                v.visit_tensor_mut(&mut repr.value);
2530                v.visit_tensor_mut(&mut repr.out);
2531            }
2532            BaseOperationIr::MaskWhere(repr) => {
2533                v.visit_tensor_mut(&mut repr.tensor);
2534                v.visit_tensor_mut(&mut repr.mask);
2535                v.visit_tensor_mut(&mut repr.value);
2536                v.visit_tensor_mut(&mut repr.out);
2537            }
2538            BaseOperationIr::MaskFill(repr) => {
2539                v.visit_tensor_mut(&mut repr.tensor);
2540                v.visit_tensor_mut(&mut repr.mask);
2541                v.visit_tensor_mut(&mut repr.out);
2542                v.visit_scalar_mut(&mut repr.value);
2543            }
2544            BaseOperationIr::Equal(repr) => {
2545                v.visit_tensor_mut(&mut repr.lhs);
2546                v.visit_tensor_mut(&mut repr.rhs);
2547                v.visit_tensor_mut(&mut repr.out);
2548            }
2549            BaseOperationIr::EqualElem(repr) => {
2550                v.visit_tensor_mut(&mut repr.lhs);
2551                v.visit_tensor_mut(&mut repr.out);
2552                v.visit_scalar_mut(&mut repr.rhs);
2553            }
2554            BaseOperationIr::RepeatDim(repr) => {
2555                v.visit_tensor_mut(&mut repr.tensor);
2556                v.visit_tensor_mut(&mut repr.out);
2557            }
2558            BaseOperationIr::Cat(repr) => {
2559                for t in repr.tensors.iter_mut() {
2560                    v.visit_tensor_mut(t);
2561                }
2562                v.visit_tensor_mut(&mut repr.out);
2563            }
2564            BaseOperationIr::Cast(repr) => {
2565                v.visit_tensor_mut(&mut repr.input);
2566                v.visit_tensor_mut(&mut repr.out);
2567            }
2568            BaseOperationIr::Unfold(repr) => {
2569                v.visit_tensor_mut(&mut repr.input);
2570                v.visit_tensor_mut(&mut repr.out);
2571            }
2572            BaseOperationIr::Empty(repr) => {
2573                v.visit_tensor_mut(&mut repr.out);
2574            }
2575            BaseOperationIr::Ones(repr) => {
2576                v.visit_tensor_mut(&mut repr.out);
2577            }
2578            BaseOperationIr::Zeros(repr) => {
2579                v.visit_tensor_mut(&mut repr.out);
2580            }
2581            BaseOperationIr::NotEqual(repr) => {
2582                v.visit_tensor_mut(&mut repr.lhs);
2583                v.visit_tensor_mut(&mut repr.rhs);
2584                v.visit_tensor_mut(&mut repr.out);
2585            }
2586            BaseOperationIr::NotEqualElem(repr) => {
2587                v.visit_tensor_mut(&mut repr.lhs);
2588                v.visit_tensor_mut(&mut repr.out);
2589                v.visit_scalar_mut(&mut repr.rhs);
2590            }
2591            BaseOperationIr::All(repr) => {
2592                v.visit_tensor_mut(&mut repr.input);
2593                v.visit_tensor_mut(&mut repr.out);
2594            }
2595            BaseOperationIr::Any(repr) => {
2596                v.visit_tensor_mut(&mut repr.input);
2597                v.visit_tensor_mut(&mut repr.out);
2598            }
2599            BaseOperationIr::AllDim(repr) => {
2600                v.visit_tensor_mut(&mut repr.input);
2601                v.visit_tensor_mut(&mut repr.out);
2602            }
2603            BaseOperationIr::AnyDim(repr) => {
2604                v.visit_tensor_mut(&mut repr.input);
2605                v.visit_tensor_mut(&mut repr.out);
2606            }
2607        }
2608    }
2609}
2610
2611impl NumericOperationIr {
2612    fn inputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
2613        match self {
2614            NumericOperationIr::Add(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2615            NumericOperationIr::AddScalar(repr) => Box::new([&repr.lhs].into_iter()),
2616            NumericOperationIr::Sub(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2617            NumericOperationIr::SubScalar(repr) => Box::new([&repr.lhs].into_iter()),
2618            NumericOperationIr::Mul(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2619            NumericOperationIr::MulScalar(repr) => Box::new([&repr.lhs].into_iter()),
2620            NumericOperationIr::Div(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2621            NumericOperationIr::DivScalar(repr) => Box::new([&repr.lhs].into_iter()),
2622            NumericOperationIr::Rem(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2623            NumericOperationIr::RemScalar(repr) => Box::new([&repr.lhs].into_iter()),
2624            NumericOperationIr::GreaterElem(repr) => Box::new([&repr.lhs].into_iter()),
2625            NumericOperationIr::GreaterEqualElem(repr) => Box::new([&repr.lhs].into_iter()),
2626            NumericOperationIr::LowerElem(repr) => Box::new([&repr.lhs].into_iter()),
2627            NumericOperationIr::LowerEqualElem(repr) => Box::new([&repr.lhs].into_iter()),
2628            NumericOperationIr::Greater(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2629            NumericOperationIr::GreaterEqual(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2630            NumericOperationIr::Lower(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2631            NumericOperationIr::LowerEqual(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2632            NumericOperationIr::ArgMax(repr) => Box::new([&repr.input].into_iter()),
2633            NumericOperationIr::ArgTopK(repr) => Box::new([&repr.input].into_iter()),
2634            NumericOperationIr::TopK(repr) => Box::new([&repr.input].into_iter()),
2635            NumericOperationIr::ArgMin(repr) => Box::new([&repr.input].into_iter()),
2636            NumericOperationIr::Clamp(repr) => Box::new([&repr.tensor].into_iter()),
2637            NumericOperationIr::Abs(repr) => Box::new([&repr.input].into_iter()),
2638            NumericOperationIr::Full(_repr) => Box::new([].into_iter()),
2639            NumericOperationIr::MeanDim(repr) => Box::new([&repr.input].into_iter()),
2640            NumericOperationIr::Mean(repr) => Box::new([&repr.input].into_iter()),
2641            NumericOperationIr::Sum(repr) => Box::new([&repr.input].into_iter()),
2642            NumericOperationIr::SumDim(repr) => Box::new([&repr.input].into_iter()),
2643            NumericOperationIr::Prod(repr) => Box::new([&repr.input].into_iter()),
2644            NumericOperationIr::ProdDim(repr) => Box::new([&repr.input].into_iter()),
2645            NumericOperationIr::Max(repr) => Box::new([&repr.input].into_iter()),
2646            NumericOperationIr::MaxDimWithIndices(repr) => Box::new([&repr.tensor].into_iter()),
2647            NumericOperationIr::TopKWithIndices(repr) => Box::new([&repr.tensor].into_iter()),
2648            NumericOperationIr::MinDimWithIndices(repr) => Box::new([&repr.tensor].into_iter()),
2649            NumericOperationIr::Min(repr) => Box::new([&repr.input].into_iter()),
2650            NumericOperationIr::MaxDim(repr) => Box::new([&repr.input].into_iter()),
2651            NumericOperationIr::MinDim(repr) => Box::new([&repr.input].into_iter()),
2652            NumericOperationIr::MaxAbs(repr) => Box::new([&repr.input].into_iter()),
2653            NumericOperationIr::MaxAbsDim(repr) => Box::new([&repr.input].into_iter()),
2654            NumericOperationIr::IntRandom(_repr) => Box::new([].into_iter()),
2655            NumericOperationIr::Powi(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
2656            NumericOperationIr::PowiScalar(repr) => Box::new([&repr.lhs].into_iter()),
2657            NumericOperationIr::CumMin(repr) => Box::new([&repr.input].into_iter()),
2658            NumericOperationIr::CumMax(repr) => Box::new([&repr.input].into_iter()),
2659            NumericOperationIr::CumProd(repr) => Box::new([&repr.input].into_iter()),
2660            NumericOperationIr::CumSum(repr) => Box::new([&repr.input].into_iter()),
2661            NumericOperationIr::Neg(repr) => Box::new([&repr.input].into_iter()),
2662            NumericOperationIr::Sign(repr) => Box::new([&repr.input].into_iter()),
2663            NumericOperationIr::ClampMin(repr) => Box::new([&repr.lhs].into_iter()),
2664            NumericOperationIr::ClampMax(repr) => Box::new([&repr.lhs].into_iter()),
2665            NumericOperationIr::Sort(repr) => Box::new([&repr.input].into_iter()),
2666            NumericOperationIr::SortWithIndices(repr) => Box::new([&repr.input].into_iter()),
2667            NumericOperationIr::ArgSort(repr) => Box::new([&repr.input].into_iter()),
2668        }
2669    }
2670
2671    fn outputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
2672        match self {
2673            NumericOperationIr::Add(repr) => Box::new([&repr.out].into_iter()),
2674            NumericOperationIr::AddScalar(repr) => Box::new([&repr.out].into_iter()),
2675            NumericOperationIr::Sub(repr) => Box::new([&repr.out].into_iter()),
2676            NumericOperationIr::SubScalar(repr) => Box::new([&repr.out].into_iter()),
2677            NumericOperationIr::Mul(repr) => Box::new([&repr.out].into_iter()),
2678            NumericOperationIr::MulScalar(repr) => Box::new([&repr.out].into_iter()),
2679            NumericOperationIr::Div(repr) => Box::new([&repr.out].into_iter()),
2680            NumericOperationIr::DivScalar(repr) => Box::new([&repr.out].into_iter()),
2681            NumericOperationIr::Rem(repr) => Box::new([&repr.out].into_iter()),
2682            NumericOperationIr::RemScalar(repr) => Box::new([&repr.out].into_iter()),
2683            NumericOperationIr::GreaterElem(repr) => Box::new([&repr.out].into_iter()),
2684            NumericOperationIr::GreaterEqualElem(repr) => Box::new([&repr.out].into_iter()),
2685            NumericOperationIr::LowerElem(repr) => Box::new([&repr.out].into_iter()),
2686            NumericOperationIr::LowerEqualElem(repr) => Box::new([&repr.out].into_iter()),
2687            NumericOperationIr::Greater(repr) => Box::new([&repr.out].into_iter()),
2688            NumericOperationIr::GreaterEqual(repr) => Box::new([&repr.out].into_iter()),
2689            NumericOperationIr::Lower(repr) => Box::new([&repr.out].into_iter()),
2690            NumericOperationIr::LowerEqual(repr) => Box::new([&repr.out].into_iter()),
2691            NumericOperationIr::ArgMax(repr) => Box::new([&repr.out].into_iter()),
2692            NumericOperationIr::ArgTopK(repr) => Box::new([&repr.out].into_iter()),
2693            NumericOperationIr::TopK(repr) => Box::new([&repr.out].into_iter()),
2694            NumericOperationIr::ArgMin(repr) => Box::new([&repr.out].into_iter()),
2695            NumericOperationIr::Clamp(repr) => Box::new([&repr.out].into_iter()),
2696            NumericOperationIr::Abs(repr) => Box::new([&repr.out].into_iter()),
2697            NumericOperationIr::Full(repr) => Box::new([&repr.out].into_iter()),
2698            NumericOperationIr::MeanDim(repr) => Box::new([&repr.out].into_iter()),
2699            NumericOperationIr::Mean(repr) => Box::new([&repr.out].into_iter()),
2700            NumericOperationIr::Sum(repr) => Box::new([&repr.out].into_iter()),
2701            NumericOperationIr::SumDim(repr) => Box::new([&repr.out].into_iter()),
2702            NumericOperationIr::Prod(repr) => Box::new([&repr.out].into_iter()),
2703            NumericOperationIr::ProdDim(repr) => Box::new([&repr.out].into_iter()),
2704            NumericOperationIr::Max(repr) => Box::new([&repr.out].into_iter()),
2705            NumericOperationIr::MaxDimWithIndices(repr) => {
2706                Box::new([&repr.out, &repr.out_indices].into_iter())
2707            }
2708            NumericOperationIr::TopKWithIndices(repr) => {
2709                Box::new([&repr.out, &repr.out_indices].into_iter())
2710            }
2711            NumericOperationIr::MinDimWithIndices(repr) => {
2712                Box::new([&repr.out, &repr.out_indices].into_iter())
2713            }
2714            NumericOperationIr::Min(repr) => Box::new([&repr.out].into_iter()),
2715            NumericOperationIr::MaxDim(repr) => Box::new([&repr.out].into_iter()),
2716            NumericOperationIr::MinDim(repr) => Box::new([&repr.out].into_iter()),
2717            NumericOperationIr::MaxAbs(repr) => Box::new([&repr.out].into_iter()),
2718            NumericOperationIr::MaxAbsDim(repr) => Box::new([&repr.out].into_iter()),
2719            NumericOperationIr::IntRandom(repr) => Box::new([&repr.out].into_iter()),
2720            NumericOperationIr::Powi(repr) => Box::new([&repr.out].into_iter()),
2721            NumericOperationIr::PowiScalar(repr) => Box::new([&repr.out].into_iter()),
2722            NumericOperationIr::CumMin(repr) => Box::new([&repr.out].into_iter()),
2723            NumericOperationIr::CumMax(repr) => Box::new([&repr.out].into_iter()),
2724            NumericOperationIr::CumProd(repr) => Box::new([&repr.out].into_iter()),
2725            NumericOperationIr::CumSum(repr) => Box::new([&repr.out].into_iter()),
2726            NumericOperationIr::Neg(repr) => Box::new([&repr.out].into_iter()),
2727            NumericOperationIr::Sign(repr) => Box::new([&repr.out].into_iter()),
2728            NumericOperationIr::ClampMin(repr) => Box::new([&repr.out].into_iter()),
2729            NumericOperationIr::ClampMax(repr) => Box::new([&repr.out].into_iter()),
2730            NumericOperationIr::Sort(repr) => Box::new([&repr.out].into_iter()),
2731            NumericOperationIr::SortWithIndices(repr) => {
2732                Box::new([&repr.out, &repr.out_indices].into_iter())
2733            }
2734            NumericOperationIr::ArgSort(repr) => Box::new([&repr.out].into_iter()),
2735        }
2736    }
2737    fn mark_read_only(&mut self, nodes: &[TensorId]) -> Vec<TensorIr> {
2738        let mut output = Vec::new();
2739
2740        match self {
2741            NumericOperationIr::Add(repr) => {
2742                repr.lhs.mark_read_only(nodes, &mut output);
2743                repr.rhs.mark_read_only(nodes, &mut output);
2744            }
2745            NumericOperationIr::AddScalar(repr) => {
2746                repr.lhs.mark_read_only(nodes, &mut output);
2747            }
2748            NumericOperationIr::Sub(repr) => {
2749                repr.lhs.mark_read_only(nodes, &mut output);
2750                repr.rhs.mark_read_only(nodes, &mut output);
2751            }
2752            NumericOperationIr::SubScalar(repr) => {
2753                repr.lhs.mark_read_only(nodes, &mut output);
2754            }
2755            NumericOperationIr::Mul(repr) => {
2756                repr.lhs.mark_read_only(nodes, &mut output);
2757                repr.rhs.mark_read_only(nodes, &mut output);
2758            }
2759            NumericOperationIr::MulScalar(repr) => {
2760                repr.lhs.mark_read_only(nodes, &mut output);
2761            }
2762            NumericOperationIr::Div(repr) => {
2763                repr.lhs.mark_read_only(nodes, &mut output);
2764                repr.rhs.mark_read_only(nodes, &mut output);
2765            }
2766            NumericOperationIr::DivScalar(repr) => {
2767                repr.lhs.mark_read_only(nodes, &mut output);
2768            }
2769            NumericOperationIr::Rem(repr) => {
2770                repr.lhs.mark_read_only(nodes, &mut output);
2771                repr.rhs.mark_read_only(nodes, &mut output);
2772            }
2773            NumericOperationIr::RemScalar(repr) => {
2774                repr.lhs.mark_read_only(nodes, &mut output);
2775            }
2776            NumericOperationIr::GreaterElem(repr) => {
2777                repr.lhs.mark_read_only(nodes, &mut output);
2778            }
2779            NumericOperationIr::GreaterEqualElem(repr) => {
2780                repr.lhs.mark_read_only(nodes, &mut output);
2781            }
2782            NumericOperationIr::LowerElem(repr) => {
2783                repr.lhs.mark_read_only(nodes, &mut output);
2784            }
2785            NumericOperationIr::LowerEqualElem(repr) => {
2786                repr.lhs.mark_read_only(nodes, &mut output);
2787            }
2788            NumericOperationIr::Greater(repr) => {
2789                repr.lhs.mark_read_only(nodes, &mut output);
2790                repr.rhs.mark_read_only(nodes, &mut output);
2791            }
2792            NumericOperationIr::GreaterEqual(repr) => {
2793                repr.lhs.mark_read_only(nodes, &mut output);
2794                repr.rhs.mark_read_only(nodes, &mut output);
2795            }
2796            NumericOperationIr::Lower(repr) => {
2797                repr.lhs.mark_read_only(nodes, &mut output);
2798                repr.rhs.mark_read_only(nodes, &mut output);
2799            }
2800            NumericOperationIr::LowerEqual(repr) => {
2801                repr.lhs.mark_read_only(nodes, &mut output);
2802                repr.rhs.mark_read_only(nodes, &mut output);
2803            }
2804            NumericOperationIr::ArgMax(repr) => {
2805                repr.input.mark_read_only(nodes, &mut output);
2806            }
2807            NumericOperationIr::ArgTopK(repr) => {
2808                repr.input.mark_read_only(nodes, &mut output);
2809            }
2810            NumericOperationIr::TopK(repr) => {
2811                repr.input.mark_read_only(nodes, &mut output);
2812            }
2813            NumericOperationIr::ArgMin(repr) => {
2814                repr.input.mark_read_only(nodes, &mut output);
2815            }
2816            NumericOperationIr::Clamp(repr) => {
2817                repr.tensor.mark_read_only(nodes, &mut output);
2818            }
2819            NumericOperationIr::Abs(repr) => {
2820                repr.input.mark_read_only(nodes, &mut output);
2821            }
2822            NumericOperationIr::Full(_) => {}
2823            NumericOperationIr::MeanDim(repr) => {
2824                repr.input.mark_read_only(nodes, &mut output);
2825            }
2826            NumericOperationIr::Mean(repr) => {
2827                repr.input.mark_read_only(nodes, &mut output);
2828            }
2829            NumericOperationIr::Sum(repr) => {
2830                repr.input.mark_read_only(nodes, &mut output);
2831            }
2832            NumericOperationIr::SumDim(repr) => {
2833                repr.input.mark_read_only(nodes, &mut output);
2834            }
2835            NumericOperationIr::Prod(repr) => {
2836                repr.input.mark_read_only(nodes, &mut output);
2837            }
2838            NumericOperationIr::ProdDim(repr) => {
2839                repr.input.mark_read_only(nodes, &mut output);
2840            }
2841            NumericOperationIr::Max(repr) => {
2842                repr.input.mark_read_only(nodes, &mut output);
2843            }
2844            NumericOperationIr::MaxDimWithIndices(repr) => {
2845                repr.tensor.mark_read_only(nodes, &mut output);
2846            }
2847            NumericOperationIr::TopKWithIndices(repr) => {
2848                repr.tensor.mark_read_only(nodes, &mut output);
2849            }
2850            NumericOperationIr::MinDimWithIndices(repr) => {
2851                repr.tensor.mark_read_only(nodes, &mut output);
2852            }
2853            NumericOperationIr::Min(repr) => {
2854                repr.input.mark_read_only(nodes, &mut output);
2855            }
2856            NumericOperationIr::MaxDim(repr) => {
2857                repr.input.mark_read_only(nodes, &mut output);
2858            }
2859            NumericOperationIr::MinDim(repr) => {
2860                repr.input.mark_read_only(nodes, &mut output);
2861            }
2862            NumericOperationIr::MaxAbs(repr) => {
2863                repr.input.mark_read_only(nodes, &mut output);
2864            }
2865            NumericOperationIr::MaxAbsDim(repr) => {
2866                repr.input.mark_read_only(nodes, &mut output);
2867            }
2868            NumericOperationIr::IntRandom(_) => {}
2869            NumericOperationIr::Powi(repr) => {
2870                repr.lhs.mark_read_only(nodes, &mut output);
2871                repr.rhs.mark_read_only(nodes, &mut output);
2872            }
2873            NumericOperationIr::PowiScalar(repr) => {
2874                repr.lhs.mark_read_only(nodes, &mut output);
2875            }
2876            NumericOperationIr::CumSum(repr) => {
2877                repr.input.mark_read_only(nodes, &mut output);
2878            }
2879            NumericOperationIr::CumProd(repr) => {
2880                repr.input.mark_read_only(nodes, &mut output);
2881            }
2882            NumericOperationIr::CumMin(repr) => {
2883                repr.input.mark_read_only(nodes, &mut output);
2884            }
2885            NumericOperationIr::CumMax(repr) => {
2886                repr.input.mark_read_only(nodes, &mut output);
2887            }
2888            NumericOperationIr::Neg(repr) => {
2889                repr.input.mark_read_only(nodes, &mut output);
2890            }
2891            NumericOperationIr::Sign(repr) => {
2892                repr.input.mark_read_only(nodes, &mut output);
2893            }
2894            NumericOperationIr::ClampMin(repr) => {
2895                repr.lhs.mark_read_only(nodes, &mut output);
2896            }
2897            NumericOperationIr::ClampMax(repr) => {
2898                repr.lhs.mark_read_only(nodes, &mut output);
2899            }
2900            NumericOperationIr::Sort(repr) => {
2901                repr.input.mark_read_only(nodes, &mut output);
2902            }
2903            NumericOperationIr::SortWithIndices(repr) => {
2904                repr.input.mark_read_only(nodes, &mut output);
2905            }
2906            NumericOperationIr::ArgSort(repr) => {
2907                repr.input.mark_read_only(nodes, &mut output);
2908            }
2909        };
2910
2911        output
2912    }
2913
2914    fn visit_mut(&mut self, v: &mut impl IrVisitorMut) {
2915        match self {
2916            NumericOperationIr::Add(repr) => {
2917                v.visit_tensor_mut(&mut repr.lhs);
2918                v.visit_tensor_mut(&mut repr.rhs);
2919                v.visit_tensor_mut(&mut repr.out);
2920            }
2921            NumericOperationIr::AddScalar(repr) => {
2922                v.visit_tensor_mut(&mut repr.lhs);
2923                v.visit_tensor_mut(&mut repr.out);
2924                v.visit_scalar_mut(&mut repr.rhs);
2925            }
2926            NumericOperationIr::Sub(repr) => {
2927                v.visit_tensor_mut(&mut repr.lhs);
2928                v.visit_tensor_mut(&mut repr.rhs);
2929                v.visit_tensor_mut(&mut repr.out);
2930            }
2931            NumericOperationIr::SubScalar(repr) => {
2932                v.visit_tensor_mut(&mut repr.lhs);
2933                v.visit_tensor_mut(&mut repr.out);
2934                v.visit_scalar_mut(&mut repr.rhs);
2935            }
2936            NumericOperationIr::Mul(repr) => {
2937                v.visit_tensor_mut(&mut repr.lhs);
2938                v.visit_tensor_mut(&mut repr.rhs);
2939                v.visit_tensor_mut(&mut repr.out);
2940            }
2941            NumericOperationIr::MulScalar(repr) => {
2942                v.visit_tensor_mut(&mut repr.lhs);
2943                v.visit_tensor_mut(&mut repr.out);
2944                v.visit_scalar_mut(&mut repr.rhs);
2945            }
2946            NumericOperationIr::Div(repr) => {
2947                v.visit_tensor_mut(&mut repr.lhs);
2948                v.visit_tensor_mut(&mut repr.rhs);
2949                v.visit_tensor_mut(&mut repr.out);
2950            }
2951            NumericOperationIr::DivScalar(repr) => {
2952                v.visit_tensor_mut(&mut repr.lhs);
2953                v.visit_tensor_mut(&mut repr.out);
2954                v.visit_scalar_mut(&mut repr.rhs);
2955            }
2956            NumericOperationIr::Rem(repr) => {
2957                v.visit_tensor_mut(&mut repr.lhs);
2958                v.visit_tensor_mut(&mut repr.rhs);
2959                v.visit_tensor_mut(&mut repr.out);
2960            }
2961            NumericOperationIr::RemScalar(repr) => {
2962                v.visit_tensor_mut(&mut repr.lhs);
2963                v.visit_tensor_mut(&mut repr.out);
2964                v.visit_scalar_mut(&mut repr.rhs);
2965            }
2966            NumericOperationIr::GreaterElem(repr) => {
2967                v.visit_tensor_mut(&mut repr.lhs);
2968                v.visit_tensor_mut(&mut repr.out);
2969                v.visit_scalar_mut(&mut repr.rhs);
2970            }
2971            NumericOperationIr::GreaterEqualElem(repr) => {
2972                v.visit_tensor_mut(&mut repr.lhs);
2973                v.visit_tensor_mut(&mut repr.out);
2974                v.visit_scalar_mut(&mut repr.rhs);
2975            }
2976            NumericOperationIr::LowerElem(repr) => {
2977                v.visit_tensor_mut(&mut repr.lhs);
2978                v.visit_tensor_mut(&mut repr.out);
2979                v.visit_scalar_mut(&mut repr.rhs);
2980            }
2981            NumericOperationIr::LowerEqualElem(repr) => {
2982                v.visit_tensor_mut(&mut repr.lhs);
2983                v.visit_tensor_mut(&mut repr.out);
2984                v.visit_scalar_mut(&mut repr.rhs);
2985            }
2986            NumericOperationIr::Greater(repr) => {
2987                v.visit_tensor_mut(&mut repr.lhs);
2988                v.visit_tensor_mut(&mut repr.rhs);
2989                v.visit_tensor_mut(&mut repr.out);
2990            }
2991            NumericOperationIr::GreaterEqual(repr) => {
2992                v.visit_tensor_mut(&mut repr.lhs);
2993                v.visit_tensor_mut(&mut repr.rhs);
2994                v.visit_tensor_mut(&mut repr.out);
2995            }
2996            NumericOperationIr::Lower(repr) => {
2997                v.visit_tensor_mut(&mut repr.lhs);
2998                v.visit_tensor_mut(&mut repr.rhs);
2999                v.visit_tensor_mut(&mut repr.out);
3000            }
3001            NumericOperationIr::LowerEqual(repr) => {
3002                v.visit_tensor_mut(&mut repr.lhs);
3003                v.visit_tensor_mut(&mut repr.rhs);
3004                v.visit_tensor_mut(&mut repr.out);
3005            }
3006            NumericOperationIr::ArgMax(repr) => {
3007                v.visit_tensor_mut(&mut repr.input);
3008                v.visit_tensor_mut(&mut repr.out);
3009            }
3010            NumericOperationIr::ArgTopK(repr) => {
3011                v.visit_tensor_mut(&mut repr.input);
3012                v.visit_tensor_mut(&mut repr.out);
3013            }
3014            NumericOperationIr::TopK(repr) => {
3015                v.visit_tensor_mut(&mut repr.input);
3016                v.visit_tensor_mut(&mut repr.out);
3017            }
3018            NumericOperationIr::ArgMin(repr) => {
3019                v.visit_tensor_mut(&mut repr.input);
3020                v.visit_tensor_mut(&mut repr.out);
3021            }
3022            NumericOperationIr::Clamp(repr) => {
3023                v.visit_tensor_mut(&mut repr.tensor);
3024                v.visit_tensor_mut(&mut repr.out);
3025                v.visit_scalar_mut(&mut repr.min);
3026                v.visit_scalar_mut(&mut repr.max);
3027            }
3028            NumericOperationIr::Abs(repr) => {
3029                v.visit_tensor_mut(&mut repr.input);
3030                v.visit_tensor_mut(&mut repr.out);
3031            }
3032            NumericOperationIr::Full(repr) => {
3033                v.visit_tensor_mut(&mut repr.out);
3034                v.visit_scalar_mut(&mut repr.value);
3035            }
3036            NumericOperationIr::MeanDim(repr) => {
3037                v.visit_tensor_mut(&mut repr.input);
3038                v.visit_tensor_mut(&mut repr.out);
3039            }
3040            NumericOperationIr::Mean(repr) => {
3041                v.visit_tensor_mut(&mut repr.input);
3042                v.visit_tensor_mut(&mut repr.out);
3043            }
3044            NumericOperationIr::Sum(repr) => {
3045                v.visit_tensor_mut(&mut repr.input);
3046                v.visit_tensor_mut(&mut repr.out);
3047            }
3048            NumericOperationIr::SumDim(repr) => {
3049                v.visit_tensor_mut(&mut repr.input);
3050                v.visit_tensor_mut(&mut repr.out);
3051            }
3052            NumericOperationIr::Prod(repr) => {
3053                v.visit_tensor_mut(&mut repr.input);
3054                v.visit_tensor_mut(&mut repr.out);
3055            }
3056            NumericOperationIr::ProdDim(repr) => {
3057                v.visit_tensor_mut(&mut repr.input);
3058                v.visit_tensor_mut(&mut repr.out);
3059            }
3060            NumericOperationIr::Max(repr) => {
3061                v.visit_tensor_mut(&mut repr.input);
3062                v.visit_tensor_mut(&mut repr.out);
3063            }
3064            NumericOperationIr::MaxDimWithIndices(repr) => {
3065                v.visit_tensor_mut(&mut repr.tensor);
3066                v.visit_tensor_mut(&mut repr.out);
3067                v.visit_tensor_mut(&mut repr.out_indices);
3068            }
3069            NumericOperationIr::TopKWithIndices(repr) => {
3070                v.visit_tensor_mut(&mut repr.tensor);
3071                v.visit_tensor_mut(&mut repr.out);
3072                v.visit_tensor_mut(&mut repr.out_indices);
3073            }
3074            NumericOperationIr::MinDimWithIndices(repr) => {
3075                v.visit_tensor_mut(&mut repr.tensor);
3076                v.visit_tensor_mut(&mut repr.out);
3077                v.visit_tensor_mut(&mut repr.out_indices);
3078            }
3079            NumericOperationIr::Min(repr) => {
3080                v.visit_tensor_mut(&mut repr.input);
3081                v.visit_tensor_mut(&mut repr.out);
3082            }
3083            NumericOperationIr::MaxDim(repr) => {
3084                v.visit_tensor_mut(&mut repr.input);
3085                v.visit_tensor_mut(&mut repr.out);
3086            }
3087            NumericOperationIr::MinDim(repr) => {
3088                v.visit_tensor_mut(&mut repr.input);
3089                v.visit_tensor_mut(&mut repr.out);
3090            }
3091            NumericOperationIr::MaxAbs(repr) => {
3092                v.visit_tensor_mut(&mut repr.input);
3093                v.visit_tensor_mut(&mut repr.out);
3094            }
3095            NumericOperationIr::MaxAbsDim(repr) => {
3096                v.visit_tensor_mut(&mut repr.input);
3097                v.visit_tensor_mut(&mut repr.out);
3098            }
3099            NumericOperationIr::IntRandom(repr) => {
3100                v.visit_tensor_mut(&mut repr.out);
3101            }
3102            NumericOperationIr::Powi(repr) => {
3103                v.visit_tensor_mut(&mut repr.lhs);
3104                v.visit_tensor_mut(&mut repr.rhs);
3105                v.visit_tensor_mut(&mut repr.out);
3106            }
3107            NumericOperationIr::PowiScalar(repr) => {
3108                v.visit_tensor_mut(&mut repr.lhs);
3109                v.visit_tensor_mut(&mut repr.out);
3110                v.visit_scalar_mut(&mut repr.rhs);
3111            }
3112            NumericOperationIr::CumMin(repr) => {
3113                v.visit_tensor_mut(&mut repr.input);
3114                v.visit_tensor_mut(&mut repr.out);
3115            }
3116            NumericOperationIr::CumMax(repr) => {
3117                v.visit_tensor_mut(&mut repr.input);
3118                v.visit_tensor_mut(&mut repr.out);
3119            }
3120            NumericOperationIr::CumProd(repr) => {
3121                v.visit_tensor_mut(&mut repr.input);
3122                v.visit_tensor_mut(&mut repr.out);
3123            }
3124            NumericOperationIr::CumSum(repr) => {
3125                v.visit_tensor_mut(&mut repr.input);
3126                v.visit_tensor_mut(&mut repr.out);
3127            }
3128            NumericOperationIr::Neg(repr) => {
3129                v.visit_tensor_mut(&mut repr.input);
3130                v.visit_tensor_mut(&mut repr.out);
3131            }
3132            NumericOperationIr::Sign(repr) => {
3133                v.visit_tensor_mut(&mut repr.input);
3134                v.visit_tensor_mut(&mut repr.out);
3135            }
3136            NumericOperationIr::ClampMin(repr) => {
3137                v.visit_tensor_mut(&mut repr.lhs);
3138                v.visit_tensor_mut(&mut repr.out);
3139                v.visit_scalar_mut(&mut repr.rhs);
3140            }
3141            NumericOperationIr::ClampMax(repr) => {
3142                v.visit_tensor_mut(&mut repr.lhs);
3143                v.visit_tensor_mut(&mut repr.out);
3144                v.visit_scalar_mut(&mut repr.rhs);
3145            }
3146            NumericOperationIr::Sort(repr) => {
3147                v.visit_tensor_mut(&mut repr.input);
3148                v.visit_tensor_mut(&mut repr.out);
3149            }
3150            NumericOperationIr::SortWithIndices(repr) => {
3151                v.visit_tensor_mut(&mut repr.input);
3152                v.visit_tensor_mut(&mut repr.out);
3153                v.visit_tensor_mut(&mut repr.out_indices);
3154            }
3155            NumericOperationIr::ArgSort(repr) => {
3156                v.visit_tensor_mut(&mut repr.input);
3157                v.visit_tensor_mut(&mut repr.out);
3158            }
3159        }
3160    }
3161}
3162
3163impl FloatOperationIr {
3164    fn inputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
3165        match self {
3166            FloatOperationIr::Matmul(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3167            FloatOperationIr::Cross(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3168            FloatOperationIr::Random(_repr) => Box::new([].into_iter()),
3169            FloatOperationIr::Exp(repr) => Box::new([&repr.input].into_iter()),
3170            FloatOperationIr::Log(repr) => Box::new([&repr.input].into_iter()),
3171            FloatOperationIr::Log1p(repr) => Box::new([&repr.input].into_iter()),
3172            FloatOperationIr::Erf(repr) => Box::new([&repr.input].into_iter()),
3173            FloatOperationIr::Recip(repr) => Box::new([&repr.input].into_iter()),
3174            FloatOperationIr::PowfScalar(repr) => Box::new([&repr.lhs].into_iter()),
3175            FloatOperationIr::Sqrt(repr) => Box::new([&repr.input].into_iter()),
3176            FloatOperationIr::Cos(repr) => Box::new([&repr.input].into_iter()),
3177            FloatOperationIr::Sin(repr) => Box::new([&repr.input].into_iter()),
3178            FloatOperationIr::Tanh(repr) => Box::new([&repr.input].into_iter()),
3179            FloatOperationIr::Round(repr) => Box::new([&repr.input].into_iter()),
3180            FloatOperationIr::Floor(repr) => Box::new([&repr.input].into_iter()),
3181            FloatOperationIr::Ceil(repr) => Box::new([&repr.input].into_iter()),
3182            FloatOperationIr::Trunc(repr) => Box::new([&repr.input].into_iter()),
3183            FloatOperationIr::IntoInt(repr) => Box::new([&repr.input].into_iter()),
3184            FloatOperationIr::Quantize(repr) => Box::new(
3185                [&repr.tensor, &repr.qparams.scales]
3186                    .into_iter()
3187                    .chain(repr.qparams.global.iter()),
3188            ),
3189            FloatOperationIr::Dequantize(repr) => Box::new([&repr.input].into_iter()),
3190            FloatOperationIr::IsNan(repr) => Box::new([&repr.input].into_iter()),
3191            FloatOperationIr::IsInf(repr) => Box::new([&repr.input].into_iter()),
3192            FloatOperationIr::GridSample2d(repr) => {
3193                Box::new([&repr.tensor, &repr.grid].into_iter())
3194            }
3195            FloatOperationIr::Tan(repr) => Box::new([&repr.input].into_iter()),
3196            FloatOperationIr::Cosh(repr) => Box::new([&repr.input].into_iter()),
3197            FloatOperationIr::Sinh(repr) => Box::new([&repr.input].into_iter()),
3198            FloatOperationIr::ArcCos(repr) => Box::new([&repr.input].into_iter()),
3199            FloatOperationIr::ArcCosh(repr) => Box::new([&repr.input].into_iter()),
3200            FloatOperationIr::ArcSin(repr) => Box::new([&repr.input].into_iter()),
3201            FloatOperationIr::ArcSinh(repr) => Box::new([&repr.input].into_iter()),
3202            FloatOperationIr::ArcTan(repr) => Box::new([&repr.input].into_iter()),
3203            FloatOperationIr::ArcTanh(repr) => Box::new([&repr.input].into_iter()),
3204            FloatOperationIr::ArcTan2(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3205            FloatOperationIr::Powf(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3206            FloatOperationIr::Hypot(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3207        }
3208    }
3209    fn outputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
3210        match self {
3211            FloatOperationIr::Matmul(repr) => Box::new([&repr.out].into_iter()),
3212            FloatOperationIr::Cross(repr) => Box::new([&repr.out].into_iter()),
3213            FloatOperationIr::Random(repr) => Box::new([&repr.out].into_iter()),
3214            FloatOperationIr::Exp(repr) => Box::new([&repr.out].into_iter()),
3215            FloatOperationIr::Log(repr) => Box::new([&repr.out].into_iter()),
3216            FloatOperationIr::Log1p(repr) => Box::new([&repr.out].into_iter()),
3217            FloatOperationIr::Erf(repr) => Box::new([&repr.out].into_iter()),
3218            FloatOperationIr::Recip(repr) => Box::new([&repr.out].into_iter()),
3219            FloatOperationIr::PowfScalar(repr) => Box::new([&repr.out].into_iter()),
3220            FloatOperationIr::Sqrt(repr) => Box::new([&repr.out].into_iter()),
3221            FloatOperationIr::Cos(repr) => Box::new([&repr.out].into_iter()),
3222            FloatOperationIr::Sin(repr) => Box::new([&repr.out].into_iter()),
3223            FloatOperationIr::Tanh(repr) => Box::new([&repr.out].into_iter()),
3224            FloatOperationIr::Round(repr) => Box::new([&repr.out].into_iter()),
3225            FloatOperationIr::Floor(repr) => Box::new([&repr.out].into_iter()),
3226            FloatOperationIr::Ceil(repr) => Box::new([&repr.out].into_iter()),
3227            FloatOperationIr::Trunc(repr) => Box::new([&repr.out].into_iter()),
3228            FloatOperationIr::IntoInt(repr) => Box::new([&repr.out].into_iter()),
3229            FloatOperationIr::Quantize(repr) => Box::new([&repr.out].into_iter()),
3230            FloatOperationIr::Dequantize(repr) => Box::new([&repr.out].into_iter()),
3231            FloatOperationIr::IsNan(repr) => Box::new([&repr.out].into_iter()),
3232            FloatOperationIr::IsInf(repr) => Box::new([&repr.out].into_iter()),
3233            FloatOperationIr::GridSample2d(repr) => Box::new([&repr.out].into_iter()),
3234            FloatOperationIr::Tan(repr) => Box::new([&repr.out].into_iter()),
3235            FloatOperationIr::Cosh(repr) => Box::new([&repr.out].into_iter()),
3236            FloatOperationIr::Sinh(repr) => Box::new([&repr.out].into_iter()),
3237            FloatOperationIr::ArcCos(repr) => Box::new([&repr.out].into_iter()),
3238            FloatOperationIr::ArcCosh(repr) => Box::new([&repr.out].into_iter()),
3239            FloatOperationIr::ArcSin(repr) => Box::new([&repr.out].into_iter()),
3240            FloatOperationIr::ArcSinh(repr) => Box::new([&repr.out].into_iter()),
3241            FloatOperationIr::ArcTan(repr) => Box::new([&repr.out].into_iter()),
3242            FloatOperationIr::ArcTanh(repr) => Box::new([&repr.out].into_iter()),
3243            FloatOperationIr::ArcTan2(repr) => Box::new([&repr.out].into_iter()),
3244            FloatOperationIr::Powf(repr) => Box::new([&repr.out].into_iter()),
3245            FloatOperationIr::Hypot(repr) => Box::new([&repr.out].into_iter()),
3246        }
3247    }
3248
3249    fn mark_read_only(&mut self, nodes: &[TensorId]) -> Vec<TensorIr> {
3250        let mut output = Vec::new();
3251
3252        match self {
3253            FloatOperationIr::Matmul(repr) => {
3254                repr.lhs.mark_read_only(nodes, &mut output);
3255                repr.rhs.mark_read_only(nodes, &mut output);
3256            }
3257            FloatOperationIr::Cross(repr) => {
3258                repr.lhs.mark_read_only(nodes, &mut output);
3259                repr.rhs.mark_read_only(nodes, &mut output);
3260            }
3261            FloatOperationIr::Random(_) => {}
3262            FloatOperationIr::Exp(repr) => {
3263                repr.input.mark_read_only(nodes, &mut output);
3264            }
3265            FloatOperationIr::Log(repr) => {
3266                repr.input.mark_read_only(nodes, &mut output);
3267            }
3268            FloatOperationIr::Log1p(repr) => {
3269                repr.input.mark_read_only(nodes, &mut output);
3270            }
3271            FloatOperationIr::Erf(repr) => {
3272                repr.input.mark_read_only(nodes, &mut output);
3273            }
3274            FloatOperationIr::Recip(repr) => {
3275                repr.input.mark_read_only(nodes, &mut output);
3276            }
3277            FloatOperationIr::PowfScalar(repr) => {
3278                repr.lhs.mark_read_only(nodes, &mut output);
3279            }
3280            FloatOperationIr::Sqrt(repr) => {
3281                repr.input.mark_read_only(nodes, &mut output);
3282            }
3283            FloatOperationIr::Cos(repr) => {
3284                repr.input.mark_read_only(nodes, &mut output);
3285            }
3286            FloatOperationIr::Sin(repr) => {
3287                repr.input.mark_read_only(nodes, &mut output);
3288            }
3289            FloatOperationIr::Tanh(repr) => {
3290                repr.input.mark_read_only(nodes, &mut output);
3291            }
3292            FloatOperationIr::Round(repr) => {
3293                repr.input.mark_read_only(nodes, &mut output);
3294            }
3295            FloatOperationIr::Floor(repr) => {
3296                repr.input.mark_read_only(nodes, &mut output);
3297            }
3298            FloatOperationIr::Ceil(repr) => {
3299                repr.input.mark_read_only(nodes, &mut output);
3300            }
3301            FloatOperationIr::Trunc(repr) => {
3302                repr.input.mark_read_only(nodes, &mut output);
3303            }
3304            FloatOperationIr::Quantize(repr) => {
3305                repr.tensor.mark_read_only(nodes, &mut output);
3306                repr.qparams.scales.mark_read_only(nodes, &mut output);
3307                if let Some(global) = &mut repr.qparams.global {
3308                    global.mark_read_only(nodes, &mut output);
3309                }
3310            }
3311            FloatOperationIr::Dequantize(repr) => {
3312                repr.input.mark_read_only(nodes, &mut output);
3313            }
3314            FloatOperationIr::IntoInt(repr) => {
3315                repr.input.mark_read_only(nodes, &mut output);
3316            }
3317            FloatOperationIr::IsNan(repr) => {
3318                repr.input.mark_read_only(nodes, &mut output);
3319            }
3320            FloatOperationIr::IsInf(repr) => {
3321                repr.input.mark_read_only(nodes, &mut output);
3322            }
3323            FloatOperationIr::GridSample2d(repr) => {
3324                repr.tensor.mark_read_only(nodes, &mut output);
3325                repr.grid.mark_read_only(nodes, &mut output);
3326            }
3327            FloatOperationIr::Tan(repr) => repr.input.mark_read_only(nodes, &mut output),
3328            FloatOperationIr::Cosh(repr) => repr.input.mark_read_only(nodes, &mut output),
3329            FloatOperationIr::Sinh(repr) => repr.input.mark_read_only(nodes, &mut output),
3330            FloatOperationIr::ArcCos(repr) => repr.input.mark_read_only(nodes, &mut output),
3331            FloatOperationIr::ArcCosh(repr) => repr.input.mark_read_only(nodes, &mut output),
3332            FloatOperationIr::ArcSin(repr) => repr.input.mark_read_only(nodes, &mut output),
3333            FloatOperationIr::ArcSinh(repr) => repr.input.mark_read_only(nodes, &mut output),
3334            FloatOperationIr::ArcTan(repr) => repr.input.mark_read_only(nodes, &mut output),
3335            FloatOperationIr::ArcTanh(repr) => repr.input.mark_read_only(nodes, &mut output),
3336            FloatOperationIr::ArcTan2(repr) => {
3337                repr.lhs.mark_read_only(nodes, &mut output);
3338                repr.rhs.mark_read_only(nodes, &mut output);
3339            }
3340            FloatOperationIr::Powf(repr) => {
3341                repr.lhs.mark_read_only(nodes, &mut output);
3342                repr.rhs.mark_read_only(nodes, &mut output);
3343            }
3344            FloatOperationIr::Hypot(repr) => {
3345                repr.lhs.mark_read_only(nodes, &mut output);
3346                repr.rhs.mark_read_only(nodes, &mut output);
3347            }
3348        };
3349
3350        output
3351    }
3352
3353    fn visit_mut(&mut self, v: &mut impl IrVisitorMut) {
3354        match self {
3355            FloatOperationIr::Matmul(repr) => {
3356                v.visit_tensor_mut(&mut repr.lhs);
3357                v.visit_tensor_mut(&mut repr.rhs);
3358                v.visit_tensor_mut(&mut repr.out);
3359            }
3360            FloatOperationIr::Cross(repr) => {
3361                v.visit_tensor_mut(&mut repr.lhs);
3362                v.visit_tensor_mut(&mut repr.rhs);
3363                v.visit_tensor_mut(&mut repr.out);
3364            }
3365            FloatOperationIr::Random(repr) => {
3366                v.visit_tensor_mut(&mut repr.out);
3367            }
3368            FloatOperationIr::Exp(repr) => {
3369                v.visit_tensor_mut(&mut repr.input);
3370                v.visit_tensor_mut(&mut repr.out);
3371            }
3372            FloatOperationIr::Log(repr) => {
3373                v.visit_tensor_mut(&mut repr.input);
3374                v.visit_tensor_mut(&mut repr.out);
3375            }
3376            FloatOperationIr::Log1p(repr) => {
3377                v.visit_tensor_mut(&mut repr.input);
3378                v.visit_tensor_mut(&mut repr.out);
3379            }
3380            FloatOperationIr::Erf(repr) => {
3381                v.visit_tensor_mut(&mut repr.input);
3382                v.visit_tensor_mut(&mut repr.out);
3383            }
3384            FloatOperationIr::Recip(repr) => {
3385                v.visit_tensor_mut(&mut repr.input);
3386                v.visit_tensor_mut(&mut repr.out);
3387            }
3388            FloatOperationIr::PowfScalar(repr) => {
3389                v.visit_tensor_mut(&mut repr.lhs);
3390                v.visit_tensor_mut(&mut repr.out);
3391                v.visit_scalar_mut(&mut repr.rhs);
3392            }
3393            FloatOperationIr::Sqrt(repr) => {
3394                v.visit_tensor_mut(&mut repr.input);
3395                v.visit_tensor_mut(&mut repr.out);
3396            }
3397            FloatOperationIr::Cos(repr) => {
3398                v.visit_tensor_mut(&mut repr.input);
3399                v.visit_tensor_mut(&mut repr.out);
3400            }
3401            FloatOperationIr::Sin(repr) => {
3402                v.visit_tensor_mut(&mut repr.input);
3403                v.visit_tensor_mut(&mut repr.out);
3404            }
3405            FloatOperationIr::Tanh(repr) => {
3406                v.visit_tensor_mut(&mut repr.input);
3407                v.visit_tensor_mut(&mut repr.out);
3408            }
3409            FloatOperationIr::Round(repr) => {
3410                v.visit_tensor_mut(&mut repr.input);
3411                v.visit_tensor_mut(&mut repr.out);
3412            }
3413            FloatOperationIr::Floor(repr) => {
3414                v.visit_tensor_mut(&mut repr.input);
3415                v.visit_tensor_mut(&mut repr.out);
3416            }
3417            FloatOperationIr::Ceil(repr) => {
3418                v.visit_tensor_mut(&mut repr.input);
3419                v.visit_tensor_mut(&mut repr.out);
3420            }
3421            FloatOperationIr::Trunc(repr) => {
3422                v.visit_tensor_mut(&mut repr.input);
3423                v.visit_tensor_mut(&mut repr.out);
3424            }
3425            FloatOperationIr::IntoInt(repr) => {
3426                v.visit_tensor_mut(&mut repr.input);
3427                v.visit_tensor_mut(&mut repr.out);
3428            }
3429            FloatOperationIr::Quantize(repr) => {
3430                v.visit_tensor_mut(&mut repr.tensor);
3431                v.visit_tensor_mut(&mut repr.qparams.scales);
3432                if let Some(global) = &mut repr.qparams.global {
3433                    v.visit_tensor_mut(global);
3434                }
3435                v.visit_tensor_mut(&mut repr.out);
3436            }
3437            FloatOperationIr::Dequantize(repr) => {
3438                v.visit_tensor_mut(&mut repr.input);
3439                v.visit_tensor_mut(&mut repr.out);
3440            }
3441            FloatOperationIr::IsNan(repr) => {
3442                v.visit_tensor_mut(&mut repr.input);
3443                v.visit_tensor_mut(&mut repr.out);
3444            }
3445            FloatOperationIr::IsInf(repr) => {
3446                v.visit_tensor_mut(&mut repr.input);
3447                v.visit_tensor_mut(&mut repr.out);
3448            }
3449            FloatOperationIr::GridSample2d(repr) => {
3450                v.visit_tensor_mut(&mut repr.tensor);
3451                v.visit_tensor_mut(&mut repr.grid);
3452                v.visit_tensor_mut(&mut repr.out);
3453            }
3454            FloatOperationIr::Tan(repr) => {
3455                v.visit_tensor_mut(&mut repr.input);
3456                v.visit_tensor_mut(&mut repr.out);
3457            }
3458            FloatOperationIr::Cosh(repr) => {
3459                v.visit_tensor_mut(&mut repr.input);
3460                v.visit_tensor_mut(&mut repr.out);
3461            }
3462            FloatOperationIr::Sinh(repr) => {
3463                v.visit_tensor_mut(&mut repr.input);
3464                v.visit_tensor_mut(&mut repr.out);
3465            }
3466            FloatOperationIr::ArcCos(repr) => {
3467                v.visit_tensor_mut(&mut repr.input);
3468                v.visit_tensor_mut(&mut repr.out);
3469            }
3470            FloatOperationIr::ArcCosh(repr) => {
3471                v.visit_tensor_mut(&mut repr.input);
3472                v.visit_tensor_mut(&mut repr.out);
3473            }
3474            FloatOperationIr::ArcSin(repr) => {
3475                v.visit_tensor_mut(&mut repr.input);
3476                v.visit_tensor_mut(&mut repr.out);
3477            }
3478            FloatOperationIr::ArcSinh(repr) => {
3479                v.visit_tensor_mut(&mut repr.input);
3480                v.visit_tensor_mut(&mut repr.out);
3481            }
3482            FloatOperationIr::ArcTan(repr) => {
3483                v.visit_tensor_mut(&mut repr.input);
3484                v.visit_tensor_mut(&mut repr.out);
3485            }
3486            FloatOperationIr::ArcTanh(repr) => {
3487                v.visit_tensor_mut(&mut repr.input);
3488                v.visit_tensor_mut(&mut repr.out);
3489            }
3490            FloatOperationIr::ArcTan2(repr) => {
3491                v.visit_tensor_mut(&mut repr.lhs);
3492                v.visit_tensor_mut(&mut repr.rhs);
3493                v.visit_tensor_mut(&mut repr.out);
3494            }
3495            FloatOperationIr::Powf(repr) => {
3496                v.visit_tensor_mut(&mut repr.lhs);
3497                v.visit_tensor_mut(&mut repr.rhs);
3498                v.visit_tensor_mut(&mut repr.out);
3499            }
3500            FloatOperationIr::Hypot(repr) => {
3501                v.visit_tensor_mut(&mut repr.lhs);
3502                v.visit_tensor_mut(&mut repr.rhs);
3503                v.visit_tensor_mut(&mut repr.out);
3504            }
3505        }
3506    }
3507}
3508
3509impl IntOperationIr {
3510    fn inputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
3511        match self {
3512            IntOperationIr::Matmul(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3513            IntOperationIr::IntoFloat(repr) => Box::new([&repr.input].into_iter()),
3514            IntOperationIr::BitwiseAnd(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3515            IntOperationIr::BitwiseAndScalar(repr) => Box::new([&repr.lhs].into_iter()),
3516            IntOperationIr::BitwiseOr(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3517            IntOperationIr::BitwiseOrScalar(repr) => Box::new([&repr.lhs].into_iter()),
3518            IntOperationIr::BitwiseXor(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3519            IntOperationIr::BitwiseXorScalar(repr) => Box::new([&repr.lhs].into_iter()),
3520            IntOperationIr::BitwiseNot(repr) => Box::new([&repr.input].into_iter()),
3521            IntOperationIr::BitwiseLeftShift(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3522            IntOperationIr::BitwiseLeftShiftScalar(repr) => Box::new([&repr.lhs].into_iter()),
3523            IntOperationIr::BitwiseRightShift(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3524            IntOperationIr::BitwiseRightShiftScalar(repr) => Box::new([&repr.lhs].into_iter()),
3525        }
3526    }
3527
3528    fn outputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
3529        match self {
3530            IntOperationIr::Matmul(repr) => Box::new([&repr.out].into_iter()),
3531            IntOperationIr::IntoFloat(repr) => Box::new([&repr.out].into_iter()),
3532            IntOperationIr::BitwiseAnd(repr) => Box::new([&repr.out].into_iter()),
3533            IntOperationIr::BitwiseAndScalar(repr) => Box::new([&repr.out].into_iter()),
3534            IntOperationIr::BitwiseOr(repr) => Box::new([&repr.out].into_iter()),
3535            IntOperationIr::BitwiseOrScalar(repr) => Box::new([&repr.out].into_iter()),
3536            IntOperationIr::BitwiseXor(repr) => Box::new([&repr.out].into_iter()),
3537            IntOperationIr::BitwiseXorScalar(repr) => Box::new([&repr.out].into_iter()),
3538            IntOperationIr::BitwiseNot(repr) => Box::new([&repr.out].into_iter()),
3539            IntOperationIr::BitwiseLeftShift(repr) => Box::new([&repr.out].into_iter()),
3540            IntOperationIr::BitwiseLeftShiftScalar(repr) => Box::new([&repr.out].into_iter()),
3541            IntOperationIr::BitwiseRightShift(repr) => Box::new([&repr.out].into_iter()),
3542            IntOperationIr::BitwiseRightShiftScalar(repr) => Box::new([&repr.out].into_iter()),
3543        }
3544    }
3545
3546    fn mark_read_only(&mut self, nodes: &[TensorId]) -> Vec<TensorIr> {
3547        let mut output = Vec::new();
3548
3549        match self {
3550            IntOperationIr::Matmul(repr) => {
3551                repr.lhs.mark_read_only(nodes, &mut output);
3552                repr.rhs.mark_read_only(nodes, &mut output);
3553            }
3554            IntOperationIr::IntoFloat(repr) => {
3555                repr.input.mark_read_only(nodes, &mut output);
3556            }
3557            IntOperationIr::BitwiseAnd(repr) => {
3558                repr.lhs.mark_read_only(nodes, &mut output);
3559                repr.rhs.mark_read_only(nodes, &mut output);
3560            }
3561            IntOperationIr::BitwiseAndScalar(repr) => {
3562                repr.lhs.mark_read_only(nodes, &mut output);
3563            }
3564            IntOperationIr::BitwiseOr(repr) => {
3565                repr.lhs.mark_read_only(nodes, &mut output);
3566                repr.rhs.mark_read_only(nodes, &mut output);
3567            }
3568            IntOperationIr::BitwiseOrScalar(repr) => {
3569                repr.lhs.mark_read_only(nodes, &mut output);
3570            }
3571            IntOperationIr::BitwiseXor(repr) => {
3572                repr.lhs.mark_read_only(nodes, &mut output);
3573                repr.rhs.mark_read_only(nodes, &mut output);
3574            }
3575            IntOperationIr::BitwiseXorScalar(repr) => {
3576                repr.lhs.mark_read_only(nodes, &mut output);
3577            }
3578            IntOperationIr::BitwiseNot(repr) => {
3579                repr.input.mark_read_only(nodes, &mut output);
3580            }
3581            IntOperationIr::BitwiseLeftShift(repr) => {
3582                repr.lhs.mark_read_only(nodes, &mut output);
3583                repr.rhs.mark_read_only(nodes, &mut output);
3584            }
3585            IntOperationIr::BitwiseLeftShiftScalar(repr) => {
3586                repr.lhs.mark_read_only(nodes, &mut output);
3587            }
3588            IntOperationIr::BitwiseRightShift(repr) => {
3589                repr.lhs.mark_read_only(nodes, &mut output);
3590                repr.rhs.mark_read_only(nodes, &mut output);
3591            }
3592            IntOperationIr::BitwiseRightShiftScalar(repr) => {
3593                repr.lhs.mark_read_only(nodes, &mut output);
3594            }
3595        };
3596
3597        output
3598    }
3599
3600    fn visit_mut(&mut self, v: &mut impl IrVisitorMut) {
3601        match self {
3602            IntOperationIr::Matmul(repr) => {
3603                v.visit_tensor_mut(&mut repr.lhs);
3604                v.visit_tensor_mut(&mut repr.rhs);
3605                v.visit_tensor_mut(&mut repr.out);
3606            }
3607            IntOperationIr::IntoFloat(repr) => {
3608                v.visit_tensor_mut(&mut repr.input);
3609                v.visit_tensor_mut(&mut repr.out);
3610            }
3611            IntOperationIr::BitwiseAnd(repr) => {
3612                v.visit_tensor_mut(&mut repr.lhs);
3613                v.visit_tensor_mut(&mut repr.rhs);
3614                v.visit_tensor_mut(&mut repr.out);
3615            }
3616            IntOperationIr::BitwiseAndScalar(repr) => {
3617                v.visit_tensor_mut(&mut repr.lhs);
3618                v.visit_tensor_mut(&mut repr.out);
3619                v.visit_scalar_mut(&mut repr.rhs);
3620            }
3621            IntOperationIr::BitwiseOr(repr) => {
3622                v.visit_tensor_mut(&mut repr.lhs);
3623                v.visit_tensor_mut(&mut repr.rhs);
3624                v.visit_tensor_mut(&mut repr.out);
3625            }
3626            IntOperationIr::BitwiseOrScalar(repr) => {
3627                v.visit_tensor_mut(&mut repr.lhs);
3628                v.visit_tensor_mut(&mut repr.out);
3629                v.visit_scalar_mut(&mut repr.rhs);
3630            }
3631            IntOperationIr::BitwiseXor(repr) => {
3632                v.visit_tensor_mut(&mut repr.lhs);
3633                v.visit_tensor_mut(&mut repr.rhs);
3634                v.visit_tensor_mut(&mut repr.out);
3635            }
3636            IntOperationIr::BitwiseXorScalar(repr) => {
3637                v.visit_tensor_mut(&mut repr.lhs);
3638                v.visit_tensor_mut(&mut repr.out);
3639                v.visit_scalar_mut(&mut repr.rhs);
3640            }
3641            IntOperationIr::BitwiseNot(repr) => {
3642                v.visit_tensor_mut(&mut repr.input);
3643                v.visit_tensor_mut(&mut repr.out);
3644            }
3645            IntOperationIr::BitwiseLeftShift(repr) => {
3646                v.visit_tensor_mut(&mut repr.lhs);
3647                v.visit_tensor_mut(&mut repr.rhs);
3648                v.visit_tensor_mut(&mut repr.out);
3649            }
3650            IntOperationIr::BitwiseLeftShiftScalar(repr) => {
3651                v.visit_tensor_mut(&mut repr.lhs);
3652                v.visit_tensor_mut(&mut repr.out);
3653                v.visit_scalar_mut(&mut repr.rhs);
3654            }
3655            IntOperationIr::BitwiseRightShift(repr) => {
3656                v.visit_tensor_mut(&mut repr.lhs);
3657                v.visit_tensor_mut(&mut repr.rhs);
3658                v.visit_tensor_mut(&mut repr.out);
3659            }
3660            IntOperationIr::BitwiseRightShiftScalar(repr) => {
3661                v.visit_tensor_mut(&mut repr.lhs);
3662                v.visit_tensor_mut(&mut repr.out);
3663                v.visit_scalar_mut(&mut repr.rhs);
3664            }
3665        }
3666    }
3667}
3668
3669impl BoolOperationIr {
3670    fn inputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
3671        match self {
3672            BoolOperationIr::IntoFloat(repr) => Box::new([&repr.input].into_iter()),
3673            BoolOperationIr::IntoInt(repr) => Box::new([&repr.input].into_iter()),
3674            BoolOperationIr::Not(repr) => Box::new([&repr.input].into_iter()),
3675            BoolOperationIr::And(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3676            BoolOperationIr::Or(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3677            BoolOperationIr::Xor(repr) => Box::new([&repr.lhs, &repr.rhs].into_iter()),
3678        }
3679    }
3680    fn outputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
3681        match self {
3682            BoolOperationIr::IntoFloat(repr) => Box::new([&repr.out].into_iter()),
3683            BoolOperationIr::IntoInt(repr) => Box::new([&repr.out].into_iter()),
3684            BoolOperationIr::Not(repr) => Box::new([&repr.out].into_iter()),
3685            BoolOperationIr::And(repr) => Box::new([&repr.out].into_iter()),
3686            BoolOperationIr::Or(repr) => Box::new([&repr.out].into_iter()),
3687            BoolOperationIr::Xor(repr) => Box::new([&repr.out].into_iter()),
3688        }
3689    }
3690    fn mark_read_only(&mut self, nodes: &[TensorId]) -> Vec<TensorIr> {
3691        let mut output = Vec::new();
3692
3693        match self {
3694            BoolOperationIr::IntoFloat(repr) => {
3695                repr.input.mark_read_only(nodes, &mut output);
3696            }
3697            BoolOperationIr::IntoInt(repr) => {
3698                repr.input.mark_read_only(nodes, &mut output);
3699            }
3700            BoolOperationIr::Not(repr) => {
3701                repr.input.mark_read_only(nodes, &mut output);
3702            }
3703            BoolOperationIr::And(repr) => {
3704                repr.lhs.mark_read_only(nodes, &mut output);
3705                repr.rhs.mark_read_only(nodes, &mut output);
3706            }
3707            BoolOperationIr::Or(repr) => {
3708                repr.lhs.mark_read_only(nodes, &mut output);
3709                repr.rhs.mark_read_only(nodes, &mut output);
3710            }
3711            BoolOperationIr::Xor(repr) => {
3712                repr.lhs.mark_read_only(nodes, &mut output);
3713                repr.rhs.mark_read_only(nodes, &mut output);
3714            }
3715        };
3716
3717        output
3718    }
3719
3720    fn visit_mut(&mut self, v: &mut impl IrVisitorMut) {
3721        match self {
3722            BoolOperationIr::IntoFloat(repr) => {
3723                v.visit_tensor_mut(&mut repr.input);
3724                v.visit_tensor_mut(&mut repr.out);
3725            }
3726            BoolOperationIr::IntoInt(repr) => {
3727                v.visit_tensor_mut(&mut repr.input);
3728                v.visit_tensor_mut(&mut repr.out);
3729            }
3730            BoolOperationIr::Not(repr) => {
3731                v.visit_tensor_mut(&mut repr.input);
3732                v.visit_tensor_mut(&mut repr.out);
3733            }
3734            BoolOperationIr::And(repr) => {
3735                v.visit_tensor_mut(&mut repr.lhs);
3736                v.visit_tensor_mut(&mut repr.rhs);
3737                v.visit_tensor_mut(&mut repr.out);
3738            }
3739            BoolOperationIr::Or(repr) => {
3740                v.visit_tensor_mut(&mut repr.lhs);
3741                v.visit_tensor_mut(&mut repr.rhs);
3742                v.visit_tensor_mut(&mut repr.out);
3743            }
3744            BoolOperationIr::Xor(repr) => {
3745                v.visit_tensor_mut(&mut repr.lhs);
3746                v.visit_tensor_mut(&mut repr.rhs);
3747                v.visit_tensor_mut(&mut repr.out);
3748            }
3749        }
3750    }
3751}
3752
3753impl ModuleOperationIr {
3754    fn inputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
3755        match self {
3756            ModuleOperationIr::BatchNorm(repr) => {
3757                Box::new([&repr.x, &repr.gamma, &repr.beta, &repr.mean, &repr.variance].into_iter())
3758            }
3759            ModuleOperationIr::Embedding(repr) => {
3760                Box::new([&repr.weights, &repr.indices].into_iter())
3761            }
3762            ModuleOperationIr::EmbeddingBackward(repr) => {
3763                Box::new([&repr.weights, &repr.out_grad, &repr.indices].into_iter())
3764            }
3765            ModuleOperationIr::Linear(repr) => {
3766                if let Some(bias) = &repr.bias {
3767                    Box::new([&repr.x, &repr.weight, bias].into_iter())
3768                } else {
3769                    Box::new([&repr.x, &repr.weight].into_iter())
3770                }
3771            }
3772            ModuleOperationIr::LinearXBackward(repr) => {
3773                Box::new([&repr.weight, &repr.output_grad].into_iter())
3774            }
3775            ModuleOperationIr::LinearWeightBackward(repr) => {
3776                Box::new([&repr.x, &repr.output_grad].into_iter())
3777            }
3778            ModuleOperationIr::LinearBiasBackward(repr) => {
3779                Box::new([&repr.output_grad].into_iter())
3780            }
3781            ModuleOperationIr::Conv1d(repr) => {
3782                if let Some(bias) = &repr.bias {
3783                    Box::new([&repr.x, &repr.weight, bias].into_iter())
3784                } else {
3785                    Box::new([&repr.x, &repr.weight].into_iter())
3786                }
3787            }
3788            ModuleOperationIr::Conv1dXBackward(repr) => {
3789                Box::new([&repr.x, &repr.weight, &repr.output_grad].into_iter())
3790            }
3791            ModuleOperationIr::Conv1dWeightBackward(repr) => {
3792                Box::new([&repr.x, &repr.weight, &repr.output_grad].into_iter())
3793            }
3794            ModuleOperationIr::Conv1dBiasBackward(repr) => {
3795                Box::new([&repr.x, &repr.bias, &repr.output_grad].into_iter())
3796            }
3797            ModuleOperationIr::Conv2d(repr) => {
3798                if let Some(bias) = &repr.bias {
3799                    Box::new([&repr.x, &repr.weight, bias].into_iter())
3800                } else {
3801                    Box::new([&repr.x, &repr.weight].into_iter())
3802                }
3803            }
3804            ModuleOperationIr::Conv2dXBackward(repr) => {
3805                Box::new([&repr.x, &repr.weight, &repr.output_grad].into_iter())
3806            }
3807            ModuleOperationIr::Conv2dWeightBackward(repr) => {
3808                Box::new([&repr.x, &repr.weight, &repr.output_grad].into_iter())
3809            }
3810            ModuleOperationIr::Conv2dBiasBackward(repr) => {
3811                Box::new([&repr.x, &repr.bias, &repr.output_grad].into_iter())
3812            }
3813            ModuleOperationIr::Conv3d(repr) => {
3814                if let Some(bias) = &repr.bias {
3815                    Box::new([&repr.x, &repr.weight, bias].into_iter())
3816                } else {
3817                    Box::new([&repr.x, &repr.weight].into_iter())
3818                }
3819            }
3820            ModuleOperationIr::Conv3dXBackward(repr) => {
3821                Box::new([&repr.x, &repr.weight, &repr.output_grad].into_iter())
3822            }
3823            ModuleOperationIr::Conv3dWeightBackward(repr) => {
3824                Box::new([&repr.x, &repr.weight, &repr.output_grad].into_iter())
3825            }
3826            ModuleOperationIr::Conv3dBiasBackward(repr) => {
3827                Box::new([&repr.x, &repr.bias, &repr.output_grad].into_iter())
3828            }
3829            ModuleOperationIr::DeformableConv2d(repr) => match (&repr.mask, &repr.bias) {
3830                (Some(mask), Some(bias)) => {
3831                    Box::new([&repr.x, &repr.offset, &repr.weight, mask, bias].into_iter())
3832                }
3833                (Some(mask), None) => {
3834                    Box::new([&repr.x, &repr.offset, &repr.weight, mask].into_iter())
3835                }
3836                (None, Some(bias)) => {
3837                    Box::new([&repr.x, &repr.offset, &repr.weight, bias].into_iter())
3838                }
3839                (None, None) => Box::new([&repr.x, &repr.offset, &repr.weight].into_iter()),
3840            },
3841            ModuleOperationIr::DeformableConv2dBackward(repr) => match (&repr.mask, &repr.bias) {
3842                (Some(mask), Some(bias)) => Box::new(
3843                    [
3844                        &repr.x,
3845                        &repr.offset,
3846                        &repr.weight,
3847                        &repr.out_grad,
3848                        mask,
3849                        bias,
3850                    ]
3851                    .into_iter(),
3852                ),
3853                (Some(mask), None) => Box::new(
3854                    [&repr.x, &repr.offset, &repr.weight, &repr.out_grad, mask].into_iter(),
3855                ),
3856                (None, Some(bias)) => Box::new(
3857                    [&repr.x, &repr.offset, &repr.weight, &repr.out_grad, bias].into_iter(),
3858                ),
3859                (None, None) => {
3860                    Box::new([&repr.x, &repr.offset, &repr.weight, &repr.out_grad].into_iter())
3861                }
3862            },
3863            ModuleOperationIr::ConvTranspose1d(repr) => {
3864                if let Some(bias) = &repr.bias {
3865                    Box::new([&repr.x, &repr.weight, bias].into_iter())
3866                } else {
3867                    Box::new([&repr.x, &repr.weight].into_iter())
3868                }
3869            }
3870            ModuleOperationIr::ConvTranspose2d(repr) => {
3871                if let Some(bias) = &repr.bias {
3872                    Box::new([&repr.x, &repr.weight, bias].into_iter())
3873                } else {
3874                    Box::new([&repr.x, &repr.weight].into_iter())
3875                }
3876            }
3877            ModuleOperationIr::ConvTranspose3d(repr) => {
3878                if let Some(bias) = &repr.bias {
3879                    Box::new([&repr.x, &repr.weight, bias].into_iter())
3880                } else {
3881                    Box::new([&repr.x, &repr.weight].into_iter())
3882                }
3883            }
3884            ModuleOperationIr::AvgPool1d(repr) => Box::new([&repr.x].into_iter()),
3885            ModuleOperationIr::AvgPool2d(repr) => Box::new([&repr.x].into_iter()),
3886            ModuleOperationIr::AvgPool1dBackward(repr) => {
3887                Box::new([&repr.x, &repr.grad].into_iter())
3888            }
3889            ModuleOperationIr::AvgPool2dBackward(repr) => {
3890                Box::new([&repr.x, &repr.grad].into_iter())
3891            }
3892            ModuleOperationIr::AdaptiveAvgPool1d(repr) => Box::new([&repr.x].into_iter()),
3893            ModuleOperationIr::AdaptiveAvgPool2d(repr) => Box::new([&repr.x].into_iter()),
3894            ModuleOperationIr::AdaptiveAvgPool1dBackward(repr) => {
3895                Box::new([&repr.x, &repr.grad].into_iter())
3896            }
3897            ModuleOperationIr::AdaptiveAvgPool2dBackward(repr) => {
3898                Box::new([&repr.x, &repr.grad].into_iter())
3899            }
3900            ModuleOperationIr::AdaptiveAvgPool3d(repr) => Box::new([&repr.x].into_iter()),
3901            ModuleOperationIr::AdaptiveAvgPool3dBackward(repr) => {
3902                Box::new([&repr.x, &repr.grad].into_iter())
3903            }
3904            ModuleOperationIr::MaxPool1d(repr) => Box::new([&repr.x].into_iter()),
3905            ModuleOperationIr::MaxPool1dWithIndices(repr) => Box::new([&repr.x].into_iter()),
3906            ModuleOperationIr::MaxPool1dWithIndicesBackward(repr) => {
3907                Box::new([&repr.x, &repr.indices, &repr.grad].into_iter())
3908            }
3909            ModuleOperationIr::MaxPool2d(repr) => Box::new([&repr.x].into_iter()),
3910            ModuleOperationIr::MaxPool2dWithIndices(repr) => Box::new([&repr.x].into_iter()),
3911            ModuleOperationIr::MaxPool2dWithIndicesBackward(repr) => {
3912                Box::new([&repr.x, &repr.indices, &repr.grad].into_iter())
3913            }
3914            ModuleOperationIr::Interpolate(repr) => Box::new([&repr.x].into_iter()),
3915            ModuleOperationIr::InterpolateBackward(repr) => {
3916                Box::new([&repr.x, &repr.grad].into_iter())
3917            }
3918            ModuleOperationIr::Rfft(repr) => Box::new([&repr.signal].into_iter()),
3919            ModuleOperationIr::IRfft(repr) => {
3920                Box::new([&repr.input_re, &repr.input_im].into_iter())
3921            }
3922            ModuleOperationIr::Attention(repr) => {
3923                if let Some(mask) = &repr.mask {
3924                    if let Some(attn_bias) = &repr.attn_bias {
3925                        Box::new([&repr.query, &repr.key, &repr.value, mask, attn_bias].into_iter())
3926                    } else {
3927                        Box::new([&repr.query, &repr.key, &repr.value, mask].into_iter())
3928                    }
3929                } else if let Some(attn_bias) = &repr.attn_bias {
3930                    Box::new([&repr.query, &repr.key, &repr.value, attn_bias].into_iter())
3931                } else {
3932                    Box::new([&repr.query, &repr.key, &repr.value].into_iter())
3933                }
3934            }
3935            ModuleOperationIr::CtcLoss(repr) => Box::new(
3936                [
3937                    &repr.log_probs,
3938                    &repr.targets,
3939                    &repr.input_lengths,
3940                    &repr.target_lengths,
3941                ]
3942                .into_iter(),
3943            ),
3944            ModuleOperationIr::CtcLossBackward(repr) => Box::new(
3945                [
3946                    &repr.log_probs,
3947                    &repr.targets,
3948                    &repr.input_lengths,
3949                    &repr.target_lengths,
3950                    &repr.grad_loss,
3951                ]
3952                .into_iter(),
3953            ),
3954            ModuleOperationIr::LayerNorm(repr) => match &repr.beta {
3955                Some(beta) => Box::new([&repr.input, &repr.gamma, beta].into_iter()),
3956                None => Box::new([&repr.input, &repr.gamma].into_iter()),
3957            },
3958            ModuleOperationIr::Unfold4d(repr) => Box::new([&repr.x].into_iter()),
3959            ModuleOperationIr::ConvTranspose1dWeightBackward(repr) => {
3960                Box::new([&repr.x, &repr.weight, &repr.output_grad].into_iter())
3961            }
3962            ModuleOperationIr::ConvTranspose1dBiasBackward(repr) => {
3963                Box::new([&repr.x, &repr.bias, &repr.output_grad].into_iter())
3964            }
3965            ModuleOperationIr::ConvTranspose2dWeightBackward(repr) => {
3966                Box::new([&repr.x, &repr.weight, &repr.output_grad].into_iter())
3967            }
3968            ModuleOperationIr::ConvTranspose2dBiasBackward(repr) => {
3969                Box::new([&repr.x, &repr.bias, &repr.output_grad].into_iter())
3970            }
3971            ModuleOperationIr::ConvTranspose3dWeightBackward(repr) => {
3972                Box::new([&repr.x, &repr.weight, &repr.output_grad].into_iter())
3973            }
3974            ModuleOperationIr::ConvTranspose3dBiasBackward(repr) => {
3975                Box::new([&repr.x, &repr.bias, &repr.output_grad].into_iter())
3976            }
3977        }
3978    }
3979    fn outputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
3980        match self {
3981            ModuleOperationIr::BatchNorm(repr) => Box::new([&repr.out].into_iter()),
3982            ModuleOperationIr::Embedding(repr) => Box::new([&repr.out].into_iter()),
3983            ModuleOperationIr::EmbeddingBackward(repr) => Box::new([&repr.out].into_iter()),
3984            ModuleOperationIr::Linear(repr) => Box::new([&repr.out].into_iter()),
3985            ModuleOperationIr::LinearXBackward(repr) => Box::new([&repr.out].into_iter()),
3986            ModuleOperationIr::LinearWeightBackward(repr) => Box::new([&repr.out].into_iter()),
3987            ModuleOperationIr::LinearBiasBackward(repr) => Box::new([&repr.out].into_iter()),
3988            ModuleOperationIr::Conv1d(repr) => Box::new([&repr.out].into_iter()),
3989            ModuleOperationIr::Conv1dXBackward(repr) => Box::new([&repr.out].into_iter()),
3990            ModuleOperationIr::Conv1dWeightBackward(repr) => Box::new([&repr.out].into_iter()),
3991            ModuleOperationIr::Conv1dBiasBackward(repr) => Box::new([&repr.out].into_iter()),
3992            ModuleOperationIr::Conv2d(repr) => Box::new([&repr.out].into_iter()),
3993            ModuleOperationIr::Conv2dXBackward(repr) => Box::new([&repr.out].into_iter()),
3994            ModuleOperationIr::Conv2dWeightBackward(repr) => Box::new([&repr.out].into_iter()),
3995            ModuleOperationIr::Conv2dBiasBackward(repr) => Box::new([&repr.out].into_iter()),
3996            ModuleOperationIr::Conv3d(repr) => Box::new([&repr.out].into_iter()),
3997            ModuleOperationIr::Conv3dXBackward(repr) => Box::new([&repr.out].into_iter()),
3998            ModuleOperationIr::Conv3dWeightBackward(repr) => Box::new([&repr.out].into_iter()),
3999            ModuleOperationIr::Conv3dBiasBackward(repr) => Box::new([&repr.out].into_iter()),
4000            ModuleOperationIr::DeformableConv2d(repr) => Box::new([&repr.out].into_iter()),
4001            ModuleOperationIr::DeformableConv2dBackward(repr) => {
4002                match (&repr.mask_grad, &repr.bias_grad) {
4003                    (Some(mask_grad), Some(bias_grad)) => Box::new(
4004                        [
4005                            &repr.input_grad,
4006                            &repr.offset_grad,
4007                            &repr.weight_grad,
4008                            mask_grad,
4009                            bias_grad,
4010                        ]
4011                        .into_iter(),
4012                    ),
4013                    (Some(mask_grad), None) => Box::new(
4014                        [
4015                            &repr.input_grad,
4016                            &repr.offset_grad,
4017                            &repr.weight_grad,
4018                            mask_grad,
4019                        ]
4020                        .into_iter(),
4021                    ),
4022                    (None, Some(bias_grad)) => Box::new(
4023                        [
4024                            &repr.input_grad,
4025                            &repr.offset_grad,
4026                            &repr.weight_grad,
4027                            bias_grad,
4028                        ]
4029                        .into_iter(),
4030                    ),
4031                    (None, None) => Box::new(
4032                        [&repr.input_grad, &repr.offset_grad, &repr.weight_grad].into_iter(),
4033                    ),
4034                }
4035            }
4036            ModuleOperationIr::ConvTranspose1d(repr) => Box::new([&repr.out].into_iter()),
4037            ModuleOperationIr::ConvTranspose2d(repr) => Box::new([&repr.out].into_iter()),
4038            ModuleOperationIr::ConvTranspose3d(repr) => Box::new([&repr.out].into_iter()),
4039            ModuleOperationIr::AvgPool1d(repr) => Box::new([&repr.out].into_iter()),
4040            ModuleOperationIr::AvgPool2d(repr) => Box::new([&repr.out].into_iter()),
4041            ModuleOperationIr::AvgPool1dBackward(repr) => Box::new([&repr.out].into_iter()),
4042            ModuleOperationIr::AvgPool2dBackward(repr) => Box::new([&repr.out].into_iter()),
4043            ModuleOperationIr::AdaptiveAvgPool1d(repr) => Box::new([&repr.out].into_iter()),
4044            ModuleOperationIr::AdaptiveAvgPool2d(repr) => Box::new([&repr.out].into_iter()),
4045            ModuleOperationIr::AdaptiveAvgPool1dBackward(repr) => Box::new([&repr.out].into_iter()),
4046            ModuleOperationIr::AdaptiveAvgPool2dBackward(repr) => Box::new([&repr.out].into_iter()),
4047            ModuleOperationIr::AdaptiveAvgPool3d(repr) => Box::new([&repr.out].into_iter()),
4048            ModuleOperationIr::AdaptiveAvgPool3dBackward(repr) => Box::new([&repr.out].into_iter()),
4049            ModuleOperationIr::MaxPool1d(repr) => Box::new([&repr.out].into_iter()),
4050            ModuleOperationIr::MaxPool1dWithIndices(repr) => {
4051                Box::new([&repr.out, &repr.out_indices].into_iter())
4052            }
4053            ModuleOperationIr::MaxPool1dWithIndicesBackward(repr) => {
4054                Box::new([&repr.out].into_iter())
4055            }
4056            ModuleOperationIr::MaxPool2d(repr) => Box::new([&repr.out].into_iter()),
4057            ModuleOperationIr::MaxPool2dWithIndices(repr) => {
4058                Box::new([&repr.out, &repr.out_indices].into_iter())
4059            }
4060            ModuleOperationIr::MaxPool2dWithIndicesBackward(repr) => {
4061                Box::new([&repr.out].into_iter())
4062            }
4063            ModuleOperationIr::Interpolate(repr) => Box::new([&repr.out].into_iter()),
4064            ModuleOperationIr::InterpolateBackward(repr) => Box::new([&repr.out].into_iter()),
4065            ModuleOperationIr::Rfft(repr) => Box::new([&repr.out_re, &repr.out_im].into_iter()),
4066            ModuleOperationIr::IRfft(repr) => Box::new([&repr.out_signal].into_iter()),
4067            ModuleOperationIr::Attention(repr) => Box::new([&repr.out].into_iter()),
4068            ModuleOperationIr::CtcLoss(repr) => Box::new([&repr.out].into_iter()),
4069            ModuleOperationIr::CtcLossBackward(repr) => Box::new([&repr.out].into_iter()),
4070            ModuleOperationIr::LayerNorm(repr) => Box::new([&repr.out].into_iter()),
4071            ModuleOperationIr::Unfold4d(repr) => Box::new([&repr.out].into_iter()),
4072            ModuleOperationIr::ConvTranspose1dWeightBackward(repr) => {
4073                Box::new([&repr.out].into_iter())
4074            }
4075            ModuleOperationIr::ConvTranspose1dBiasBackward(repr) => {
4076                Box::new([&repr.out].into_iter())
4077            }
4078            ModuleOperationIr::ConvTranspose2dWeightBackward(repr) => {
4079                Box::new([&repr.out].into_iter())
4080            }
4081            ModuleOperationIr::ConvTranspose2dBiasBackward(repr) => {
4082                Box::new([&repr.out].into_iter())
4083            }
4084            ModuleOperationIr::ConvTranspose3dWeightBackward(repr) => {
4085                Box::new([&repr.out].into_iter())
4086            }
4087            ModuleOperationIr::ConvTranspose3dBiasBackward(repr) => {
4088                Box::new([&repr.out].into_iter())
4089            }
4090        }
4091    }
4092
4093    fn mark_read_only(&mut self, nodes: &[TensorId]) -> Vec<TensorIr> {
4094        let mut output = Vec::new();
4095
4096        match self {
4097            ModuleOperationIr::BatchNorm(repr) => {
4098                repr.x.mark_read_only(nodes, &mut output);
4099                repr.gamma.mark_read_only(nodes, &mut output);
4100                repr.beta.mark_read_only(nodes, &mut output);
4101                repr.mean.mark_read_only(nodes, &mut output);
4102                repr.variance.mark_read_only(nodes, &mut output);
4103            }
4104            ModuleOperationIr::Embedding(repr) => {
4105                repr.weights.mark_read_only(nodes, &mut output);
4106                repr.indices.mark_read_only(nodes, &mut output);
4107            }
4108            ModuleOperationIr::EmbeddingBackward(repr) => {
4109                repr.weights.mark_read_only(nodes, &mut output);
4110                repr.out_grad.mark_read_only(nodes, &mut output);
4111                repr.indices.mark_read_only(nodes, &mut output);
4112            }
4113            ModuleOperationIr::Linear(repr) => {
4114                repr.x.mark_read_only(nodes, &mut output);
4115                repr.weight.mark_read_only(nodes, &mut output);
4116
4117                if let Some(bias) = &mut repr.bias {
4118                    bias.mark_read_only(nodes, &mut output);
4119                }
4120            }
4121            ModuleOperationIr::LinearXBackward(repr) => {
4122                repr.weight.mark_read_only(nodes, &mut output);
4123                repr.output_grad.mark_read_only(nodes, &mut output);
4124            }
4125            ModuleOperationIr::LinearWeightBackward(repr) => {
4126                repr.x.mark_read_only(nodes, &mut output);
4127                repr.output_grad.mark_read_only(nodes, &mut output);
4128            }
4129            ModuleOperationIr::LinearBiasBackward(repr) => {
4130                repr.output_grad.mark_read_only(nodes, &mut output);
4131            }
4132            ModuleOperationIr::Conv1d(repr) => {
4133                repr.x.mark_read_only(nodes, &mut output);
4134                repr.weight.mark_read_only(nodes, &mut output);
4135
4136                if let Some(bias) = &mut repr.bias {
4137                    bias.mark_read_only(nodes, &mut output);
4138                }
4139            }
4140            ModuleOperationIr::Conv1dXBackward(repr) => {
4141                repr.x.mark_read_only(nodes, &mut output);
4142                repr.weight.mark_read_only(nodes, &mut output);
4143                repr.output_grad.mark_read_only(nodes, &mut output);
4144            }
4145            ModuleOperationIr::Conv1dWeightBackward(repr) => {
4146                repr.x.mark_read_only(nodes, &mut output);
4147                repr.weight.mark_read_only(nodes, &mut output);
4148                repr.output_grad.mark_read_only(nodes, &mut output);
4149            }
4150            ModuleOperationIr::Conv1dBiasBackward(repr) => {
4151                repr.x.mark_read_only(nodes, &mut output);
4152                repr.bias.mark_read_only(nodes, &mut output);
4153                repr.output_grad.mark_read_only(nodes, &mut output);
4154            }
4155            ModuleOperationIr::Conv2d(repr) => {
4156                repr.x.mark_read_only(nodes, &mut output);
4157                repr.weight.mark_read_only(nodes, &mut output);
4158
4159                if let Some(bias) = &mut repr.bias {
4160                    bias.mark_read_only(nodes, &mut output);
4161                }
4162            }
4163            ModuleOperationIr::Conv2dXBackward(repr) => {
4164                repr.x.mark_read_only(nodes, &mut output);
4165                repr.weight.mark_read_only(nodes, &mut output);
4166                repr.output_grad.mark_read_only(nodes, &mut output);
4167            }
4168            ModuleOperationIr::Conv2dWeightBackward(repr) => {
4169                repr.x.mark_read_only(nodes, &mut output);
4170                repr.weight.mark_read_only(nodes, &mut output);
4171                repr.output_grad.mark_read_only(nodes, &mut output);
4172            }
4173            ModuleOperationIr::Conv2dBiasBackward(repr) => {
4174                repr.x.mark_read_only(nodes, &mut output);
4175                repr.bias.mark_read_only(nodes, &mut output);
4176                repr.output_grad.mark_read_only(nodes, &mut output);
4177            }
4178            ModuleOperationIr::Conv3d(repr) => {
4179                repr.x.mark_read_only(nodes, &mut output);
4180                repr.weight.mark_read_only(nodes, &mut output);
4181
4182                if let Some(bias) = &mut repr.bias {
4183                    bias.mark_read_only(nodes, &mut output);
4184                }
4185            }
4186            ModuleOperationIr::Conv3dXBackward(repr) => {
4187                repr.x.mark_read_only(nodes, &mut output);
4188                repr.weight.mark_read_only(nodes, &mut output);
4189                repr.output_grad.mark_read_only(nodes, &mut output);
4190            }
4191            ModuleOperationIr::Conv3dWeightBackward(repr) => {
4192                repr.x.mark_read_only(nodes, &mut output);
4193                repr.weight.mark_read_only(nodes, &mut output);
4194                repr.output_grad.mark_read_only(nodes, &mut output);
4195            }
4196            ModuleOperationIr::Conv3dBiasBackward(repr) => {
4197                repr.x.mark_read_only(nodes, &mut output);
4198                repr.bias.mark_read_only(nodes, &mut output);
4199                repr.output_grad.mark_read_only(nodes, &mut output);
4200            }
4201            ModuleOperationIr::DeformableConv2d(repr) => {
4202                repr.x.mark_read_only(nodes, &mut output);
4203                repr.weight.mark_read_only(nodes, &mut output);
4204                repr.offset.mark_read_only(nodes, &mut output);
4205
4206                match (&mut repr.mask, &mut repr.bias) {
4207                    (Some(mask), Some(bias)) => {
4208                        mask.mark_read_only(nodes, &mut output);
4209                        bias.mark_read_only(nodes, &mut output);
4210                    }
4211                    (Some(mask), None) => {
4212                        mask.mark_read_only(nodes, &mut output);
4213                    }
4214                    (None, Some(bias)) => {
4215                        bias.mark_read_only(nodes, &mut output);
4216                    }
4217                    (None, None) => {}
4218                };
4219            }
4220            ModuleOperationIr::DeformableConv2dBackward(repr) => {
4221                repr.x.mark_read_only(nodes, &mut output);
4222                repr.weight.mark_read_only(nodes, &mut output);
4223                repr.offset.mark_read_only(nodes, &mut output);
4224                repr.out_grad.mark_read_only(nodes, &mut output);
4225
4226                if let Some(mask) = repr.mask.as_mut() {
4227                    mask.mark_read_only(nodes, &mut output);
4228                }
4229                if let Some(bias) = repr.bias.as_mut() {
4230                    bias.mark_read_only(nodes, &mut output);
4231                }
4232            }
4233            ModuleOperationIr::ConvTranspose1d(repr) => {
4234                repr.x.mark_read_only(nodes, &mut output);
4235                repr.weight.mark_read_only(nodes, &mut output);
4236
4237                if let Some(bias) = &mut repr.bias {
4238                    bias.mark_read_only(nodes, &mut output);
4239                }
4240            }
4241            ModuleOperationIr::ConvTranspose2d(repr) => {
4242                repr.x.mark_read_only(nodes, &mut output);
4243                repr.weight.mark_read_only(nodes, &mut output);
4244
4245                if let Some(bias) = &mut repr.bias {
4246                    bias.mark_read_only(nodes, &mut output);
4247                }
4248            }
4249            ModuleOperationIr::ConvTranspose3d(repr) => {
4250                repr.x.mark_read_only(nodes, &mut output);
4251                repr.weight.mark_read_only(nodes, &mut output);
4252
4253                if let Some(bias) = &mut repr.bias {
4254                    bias.mark_read_only(nodes, &mut output);
4255                }
4256            }
4257            ModuleOperationIr::AvgPool1d(repr) => {
4258                repr.x.mark_read_only(nodes, &mut output);
4259            }
4260            ModuleOperationIr::AvgPool2d(repr) => {
4261                repr.x.mark_read_only(nodes, &mut output);
4262            }
4263            ModuleOperationIr::AvgPool1dBackward(repr) => {
4264                repr.x.mark_read_only(nodes, &mut output);
4265                repr.grad.mark_read_only(nodes, &mut output);
4266            }
4267            ModuleOperationIr::AvgPool2dBackward(repr) => {
4268                repr.x.mark_read_only(nodes, &mut output);
4269                repr.grad.mark_read_only(nodes, &mut output);
4270            }
4271            ModuleOperationIr::AdaptiveAvgPool1d(repr) => {
4272                repr.x.mark_read_only(nodes, &mut output);
4273            }
4274            ModuleOperationIr::AdaptiveAvgPool2d(repr) => {
4275                repr.x.mark_read_only(nodes, &mut output);
4276            }
4277            ModuleOperationIr::AdaptiveAvgPool1dBackward(repr) => {
4278                repr.x.mark_read_only(nodes, &mut output);
4279                repr.grad.mark_read_only(nodes, &mut output);
4280            }
4281            ModuleOperationIr::AdaptiveAvgPool2dBackward(repr) => {
4282                repr.x.mark_read_only(nodes, &mut output);
4283                repr.grad.mark_read_only(nodes, &mut output);
4284            }
4285            ModuleOperationIr::AdaptiveAvgPool3d(repr) => {
4286                repr.x.mark_read_only(nodes, &mut output);
4287            }
4288            ModuleOperationIr::AdaptiveAvgPool3dBackward(repr) => {
4289                repr.x.mark_read_only(nodes, &mut output);
4290                repr.grad.mark_read_only(nodes, &mut output);
4291            }
4292            ModuleOperationIr::MaxPool1d(repr) => {
4293                repr.x.mark_read_only(nodes, &mut output);
4294            }
4295            ModuleOperationIr::MaxPool1dWithIndices(repr) => {
4296                repr.x.mark_read_only(nodes, &mut output);
4297            }
4298            ModuleOperationIr::MaxPool1dWithIndicesBackward(repr) => {
4299                repr.x.mark_read_only(nodes, &mut output);
4300                repr.grad.mark_read_only(nodes, &mut output);
4301            }
4302            ModuleOperationIr::MaxPool2d(repr) => {
4303                repr.x.mark_read_only(nodes, &mut output);
4304            }
4305            ModuleOperationIr::MaxPool2dWithIndices(repr) => {
4306                repr.x.mark_read_only(nodes, &mut output);
4307            }
4308            ModuleOperationIr::MaxPool2dWithIndicesBackward(repr) => {
4309                repr.x.mark_read_only(nodes, &mut output);
4310                repr.grad.mark_read_only(nodes, &mut output);
4311            }
4312            ModuleOperationIr::Interpolate(repr) => {
4313                repr.x.mark_read_only(nodes, &mut output);
4314            }
4315            ModuleOperationIr::InterpolateBackward(repr) => {
4316                repr.x.mark_read_only(nodes, &mut output);
4317                repr.grad.mark_read_only(nodes, &mut output);
4318            }
4319            ModuleOperationIr::Rfft(repr) => {
4320                repr.signal.mark_read_only(nodes, &mut output);
4321            }
4322            ModuleOperationIr::IRfft(repr) => {
4323                repr.input_re.mark_read_only(nodes, &mut output);
4324                repr.input_im.mark_read_only(nodes, &mut output);
4325            }
4326            ModuleOperationIr::Attention(repr) => {
4327                repr.query.mark_read_only(nodes, &mut output);
4328                repr.key.mark_read_only(nodes, &mut output);
4329                repr.value.mark_read_only(nodes, &mut output);
4330                if let Some(mask) = &mut repr.mask {
4331                    mask.mark_read_only(nodes, &mut output);
4332                }
4333                if let Some(attn_bias) = &mut repr.attn_bias {
4334                    attn_bias.mark_read_only(nodes, &mut output);
4335                }
4336            }
4337            ModuleOperationIr::CtcLoss(repr) => {
4338                repr.log_probs.mark_read_only(nodes, &mut output);
4339                repr.targets.mark_read_only(nodes, &mut output);
4340                repr.input_lengths.mark_read_only(nodes, &mut output);
4341                repr.target_lengths.mark_read_only(nodes, &mut output);
4342            }
4343            ModuleOperationIr::CtcLossBackward(repr) => {
4344                repr.log_probs.mark_read_only(nodes, &mut output);
4345                repr.targets.mark_read_only(nodes, &mut output);
4346                repr.input_lengths.mark_read_only(nodes, &mut output);
4347                repr.target_lengths.mark_read_only(nodes, &mut output);
4348                repr.grad_loss.mark_read_only(nodes, &mut output);
4349            }
4350            ModuleOperationIr::LayerNorm(repr) => {
4351                repr.input.mark_read_only(nodes, &mut output);
4352                repr.gamma.mark_read_only(nodes, &mut output);
4353                if let Some(beta) = &mut repr.beta {
4354                    beta.mark_read_only(nodes, &mut output);
4355                }
4356            }
4357            ModuleOperationIr::Unfold4d(repr) => {
4358                repr.x.mark_read_only(nodes, &mut output);
4359            }
4360            ModuleOperationIr::ConvTranspose1dWeightBackward(repr) => {
4361                repr.x.mark_read_only(nodes, &mut output);
4362                repr.weight.mark_read_only(nodes, &mut output);
4363                repr.output_grad.mark_read_only(nodes, &mut output);
4364            }
4365            ModuleOperationIr::ConvTranspose1dBiasBackward(repr) => {
4366                repr.x.mark_read_only(nodes, &mut output);
4367                repr.bias.mark_read_only(nodes, &mut output);
4368                repr.output_grad.mark_read_only(nodes, &mut output);
4369            }
4370            ModuleOperationIr::ConvTranspose2dWeightBackward(repr) => {
4371                repr.x.mark_read_only(nodes, &mut output);
4372                repr.weight.mark_read_only(nodes, &mut output);
4373                repr.output_grad.mark_read_only(nodes, &mut output);
4374            }
4375            ModuleOperationIr::ConvTranspose2dBiasBackward(repr) => {
4376                repr.x.mark_read_only(nodes, &mut output);
4377                repr.bias.mark_read_only(nodes, &mut output);
4378                repr.output_grad.mark_read_only(nodes, &mut output);
4379            }
4380            ModuleOperationIr::ConvTranspose3dWeightBackward(repr) => {
4381                repr.x.mark_read_only(nodes, &mut output);
4382                repr.weight.mark_read_only(nodes, &mut output);
4383                repr.output_grad.mark_read_only(nodes, &mut output);
4384            }
4385            ModuleOperationIr::ConvTranspose3dBiasBackward(repr) => {
4386                repr.x.mark_read_only(nodes, &mut output);
4387                repr.bias.mark_read_only(nodes, &mut output);
4388                repr.output_grad.mark_read_only(nodes, &mut output);
4389            }
4390        };
4391
4392        output
4393    }
4394
4395    fn visit_mut(&mut self, v: &mut impl IrVisitorMut) {
4396        match self {
4397            ModuleOperationIr::BatchNorm(repr) => {
4398                v.visit_tensor_mut(&mut repr.x);
4399                v.visit_tensor_mut(&mut repr.gamma);
4400                v.visit_tensor_mut(&mut repr.beta);
4401                v.visit_tensor_mut(&mut repr.mean);
4402                v.visit_tensor_mut(&mut repr.variance);
4403                v.visit_scalar_mut(&mut repr.epsilon);
4404                v.visit_tensor_mut(&mut repr.out);
4405            }
4406            ModuleOperationIr::Embedding(repr) => {
4407                v.visit_tensor_mut(&mut repr.weights);
4408                v.visit_tensor_mut(&mut repr.indices);
4409                v.visit_tensor_mut(&mut repr.out);
4410            }
4411            ModuleOperationIr::EmbeddingBackward(repr) => {
4412                v.visit_tensor_mut(&mut repr.weights);
4413                v.visit_tensor_mut(&mut repr.out_grad);
4414                v.visit_tensor_mut(&mut repr.indices);
4415                v.visit_tensor_mut(&mut repr.out);
4416            }
4417            ModuleOperationIr::Linear(repr) => {
4418                v.visit_tensor_mut(&mut repr.x);
4419                v.visit_tensor_mut(&mut repr.weight);
4420                if let Some(bias) = &mut repr.bias {
4421                    v.visit_tensor_mut(bias);
4422                }
4423                v.visit_tensor_mut(&mut repr.out);
4424            }
4425            ModuleOperationIr::LinearXBackward(repr) => {
4426                v.visit_tensor_mut(&mut repr.weight);
4427                v.visit_tensor_mut(&mut repr.output_grad);
4428                v.visit_tensor_mut(&mut repr.out);
4429            }
4430            ModuleOperationIr::LinearWeightBackward(repr) => {
4431                v.visit_tensor_mut(&mut repr.x);
4432                v.visit_tensor_mut(&mut repr.output_grad);
4433                v.visit_tensor_mut(&mut repr.out);
4434            }
4435            ModuleOperationIr::LinearBiasBackward(repr) => {
4436                v.visit_tensor_mut(&mut repr.output_grad);
4437                v.visit_tensor_mut(&mut repr.out);
4438            }
4439            ModuleOperationIr::Conv1d(repr) => {
4440                v.visit_tensor_mut(&mut repr.x);
4441                v.visit_tensor_mut(&mut repr.weight);
4442                if let Some(bias) = &mut repr.bias {
4443                    v.visit_tensor_mut(bias);
4444                }
4445                v.visit_tensor_mut(&mut repr.out);
4446            }
4447            ModuleOperationIr::Conv1dXBackward(repr) => {
4448                v.visit_tensor_mut(&mut repr.x);
4449                v.visit_tensor_mut(&mut repr.weight);
4450                v.visit_tensor_mut(&mut repr.output_grad);
4451                v.visit_tensor_mut(&mut repr.out);
4452            }
4453            ModuleOperationIr::Conv1dWeightBackward(repr) => {
4454                v.visit_tensor_mut(&mut repr.x);
4455                v.visit_tensor_mut(&mut repr.weight);
4456                v.visit_tensor_mut(&mut repr.output_grad);
4457                v.visit_tensor_mut(&mut repr.out);
4458            }
4459            ModuleOperationIr::Conv1dBiasBackward(repr) => {
4460                v.visit_tensor_mut(&mut repr.x);
4461                v.visit_tensor_mut(&mut repr.bias);
4462                v.visit_tensor_mut(&mut repr.output_grad);
4463                v.visit_tensor_mut(&mut repr.out);
4464            }
4465            ModuleOperationIr::Conv2d(repr) => {
4466                v.visit_tensor_mut(&mut repr.x);
4467                v.visit_tensor_mut(&mut repr.weight);
4468                if let Some(bias) = &mut repr.bias {
4469                    v.visit_tensor_mut(bias);
4470                }
4471                v.visit_tensor_mut(&mut repr.out);
4472            }
4473            ModuleOperationIr::Conv2dXBackward(repr) => {
4474                v.visit_tensor_mut(&mut repr.x);
4475                v.visit_tensor_mut(&mut repr.weight);
4476                v.visit_tensor_mut(&mut repr.output_grad);
4477                v.visit_tensor_mut(&mut repr.out);
4478            }
4479            ModuleOperationIr::Conv2dWeightBackward(repr) => {
4480                v.visit_tensor_mut(&mut repr.x);
4481                v.visit_tensor_mut(&mut repr.weight);
4482                v.visit_tensor_mut(&mut repr.output_grad);
4483                v.visit_tensor_mut(&mut repr.out);
4484            }
4485            ModuleOperationIr::Conv2dBiasBackward(repr) => {
4486                v.visit_tensor_mut(&mut repr.x);
4487                v.visit_tensor_mut(&mut repr.bias);
4488                v.visit_tensor_mut(&mut repr.output_grad);
4489                v.visit_tensor_mut(&mut repr.out);
4490            }
4491            ModuleOperationIr::Conv3d(repr) => {
4492                v.visit_tensor_mut(&mut repr.x);
4493                v.visit_tensor_mut(&mut repr.weight);
4494                if let Some(bias) = &mut repr.bias {
4495                    v.visit_tensor_mut(bias);
4496                }
4497                v.visit_tensor_mut(&mut repr.out);
4498            }
4499            ModuleOperationIr::Conv3dXBackward(repr) => {
4500                v.visit_tensor_mut(&mut repr.x);
4501                v.visit_tensor_mut(&mut repr.weight);
4502                v.visit_tensor_mut(&mut repr.output_grad);
4503                v.visit_tensor_mut(&mut repr.out);
4504            }
4505            ModuleOperationIr::Conv3dWeightBackward(repr) => {
4506                v.visit_tensor_mut(&mut repr.x);
4507                v.visit_tensor_mut(&mut repr.weight);
4508                v.visit_tensor_mut(&mut repr.output_grad);
4509                v.visit_tensor_mut(&mut repr.out);
4510            }
4511            ModuleOperationIr::Conv3dBiasBackward(repr) => {
4512                v.visit_tensor_mut(&mut repr.x);
4513                v.visit_tensor_mut(&mut repr.bias);
4514                v.visit_tensor_mut(&mut repr.output_grad);
4515                v.visit_tensor_mut(&mut repr.out);
4516            }
4517            ModuleOperationIr::DeformableConv2d(repr) => {
4518                v.visit_tensor_mut(&mut repr.x);
4519                v.visit_tensor_mut(&mut repr.offset);
4520                v.visit_tensor_mut(&mut repr.weight);
4521                if let Some(mask) = &mut repr.mask {
4522                    v.visit_tensor_mut(mask);
4523                }
4524                if let Some(bias) = &mut repr.bias {
4525                    v.visit_tensor_mut(bias);
4526                }
4527                v.visit_tensor_mut(&mut repr.out);
4528            }
4529            ModuleOperationIr::DeformableConv2dBackward(repr) => {
4530                v.visit_tensor_mut(&mut repr.x);
4531                v.visit_tensor_mut(&mut repr.offset);
4532                v.visit_tensor_mut(&mut repr.weight);
4533                v.visit_tensor_mut(&mut repr.out_grad);
4534                if let Some(mask) = &mut repr.mask {
4535                    v.visit_tensor_mut(mask);
4536                }
4537                if let Some(bias) = &mut repr.bias {
4538                    v.visit_tensor_mut(bias);
4539                }
4540                v.visit_tensor_mut(&mut repr.input_grad);
4541                v.visit_tensor_mut(&mut repr.offset_grad);
4542                v.visit_tensor_mut(&mut repr.weight_grad);
4543                if let Some(mask_grad) = &mut repr.mask_grad {
4544                    v.visit_tensor_mut(mask_grad);
4545                }
4546                if let Some(bias_grad) = &mut repr.bias_grad {
4547                    v.visit_tensor_mut(bias_grad);
4548                }
4549            }
4550            ModuleOperationIr::ConvTranspose1d(repr) => {
4551                v.visit_tensor_mut(&mut repr.x);
4552                v.visit_tensor_mut(&mut repr.weight);
4553                if let Some(bias) = &mut repr.bias {
4554                    v.visit_tensor_mut(bias);
4555                }
4556                v.visit_tensor_mut(&mut repr.out);
4557            }
4558            ModuleOperationIr::ConvTranspose2d(repr) => {
4559                v.visit_tensor_mut(&mut repr.x);
4560                v.visit_tensor_mut(&mut repr.weight);
4561                if let Some(bias) = &mut repr.bias {
4562                    v.visit_tensor_mut(bias);
4563                }
4564                v.visit_tensor_mut(&mut repr.out);
4565            }
4566            ModuleOperationIr::ConvTranspose3d(repr) => {
4567                v.visit_tensor_mut(&mut repr.x);
4568                v.visit_tensor_mut(&mut repr.weight);
4569                if let Some(bias) = &mut repr.bias {
4570                    v.visit_tensor_mut(bias);
4571                }
4572                v.visit_tensor_mut(&mut repr.out);
4573            }
4574            ModuleOperationIr::AvgPool1d(repr) => {
4575                v.visit_tensor_mut(&mut repr.x);
4576                v.visit_tensor_mut(&mut repr.out);
4577            }
4578            ModuleOperationIr::AvgPool2d(repr) => {
4579                v.visit_tensor_mut(&mut repr.x);
4580                v.visit_tensor_mut(&mut repr.out);
4581            }
4582            ModuleOperationIr::AvgPool1dBackward(repr) => {
4583                v.visit_tensor_mut(&mut repr.x);
4584                v.visit_tensor_mut(&mut repr.grad);
4585                v.visit_tensor_mut(&mut repr.out);
4586            }
4587            ModuleOperationIr::AvgPool2dBackward(repr) => {
4588                v.visit_tensor_mut(&mut repr.x);
4589                v.visit_tensor_mut(&mut repr.grad);
4590                v.visit_tensor_mut(&mut repr.out);
4591            }
4592            ModuleOperationIr::AdaptiveAvgPool1d(repr) => {
4593                v.visit_tensor_mut(&mut repr.x);
4594                v.visit_tensor_mut(&mut repr.out);
4595            }
4596            ModuleOperationIr::AdaptiveAvgPool2d(repr) => {
4597                v.visit_tensor_mut(&mut repr.x);
4598                v.visit_tensor_mut(&mut repr.out);
4599            }
4600            ModuleOperationIr::AdaptiveAvgPool1dBackward(repr) => {
4601                v.visit_tensor_mut(&mut repr.x);
4602                v.visit_tensor_mut(&mut repr.grad);
4603                v.visit_tensor_mut(&mut repr.out);
4604            }
4605            ModuleOperationIr::AdaptiveAvgPool2dBackward(repr) => {
4606                v.visit_tensor_mut(&mut repr.x);
4607                v.visit_tensor_mut(&mut repr.grad);
4608                v.visit_tensor_mut(&mut repr.out);
4609            }
4610            ModuleOperationIr::AdaptiveAvgPool3d(repr) => {
4611                v.visit_tensor_mut(&mut repr.x);
4612                v.visit_tensor_mut(&mut repr.out);
4613            }
4614            ModuleOperationIr::AdaptiveAvgPool3dBackward(repr) => {
4615                v.visit_tensor_mut(&mut repr.x);
4616                v.visit_tensor_mut(&mut repr.grad);
4617                v.visit_tensor_mut(&mut repr.out);
4618            }
4619            ModuleOperationIr::MaxPool1d(repr) => {
4620                v.visit_tensor_mut(&mut repr.x);
4621                v.visit_tensor_mut(&mut repr.out);
4622            }
4623            ModuleOperationIr::MaxPool1dWithIndices(repr) => {
4624                v.visit_tensor_mut(&mut repr.x);
4625                v.visit_tensor_mut(&mut repr.out);
4626                v.visit_tensor_mut(&mut repr.out_indices);
4627            }
4628            ModuleOperationIr::MaxPool1dWithIndicesBackward(repr) => {
4629                v.visit_tensor_mut(&mut repr.x);
4630                v.visit_tensor_mut(&mut repr.indices);
4631                v.visit_tensor_mut(&mut repr.grad);
4632                v.visit_tensor_mut(&mut repr.out);
4633            }
4634            ModuleOperationIr::MaxPool2d(repr) => {
4635                v.visit_tensor_mut(&mut repr.x);
4636                v.visit_tensor_mut(&mut repr.out);
4637            }
4638            ModuleOperationIr::MaxPool2dWithIndices(repr) => {
4639                v.visit_tensor_mut(&mut repr.x);
4640                v.visit_tensor_mut(&mut repr.out);
4641                v.visit_tensor_mut(&mut repr.out_indices);
4642            }
4643            ModuleOperationIr::MaxPool2dWithIndicesBackward(repr) => {
4644                v.visit_tensor_mut(&mut repr.x);
4645                v.visit_tensor_mut(&mut repr.indices);
4646                v.visit_tensor_mut(&mut repr.grad);
4647                v.visit_tensor_mut(&mut repr.out);
4648            }
4649            ModuleOperationIr::Interpolate(repr) => {
4650                v.visit_tensor_mut(&mut repr.x);
4651                v.visit_tensor_mut(&mut repr.out);
4652            }
4653            ModuleOperationIr::InterpolateBackward(repr) => {
4654                v.visit_tensor_mut(&mut repr.x);
4655                v.visit_tensor_mut(&mut repr.grad);
4656                v.visit_tensor_mut(&mut repr.out);
4657            }
4658            ModuleOperationIr::Rfft(repr) => {
4659                v.visit_tensor_mut(&mut repr.signal);
4660                v.visit_tensor_mut(&mut repr.out_re);
4661                v.visit_tensor_mut(&mut repr.out_im);
4662            }
4663            ModuleOperationIr::IRfft(repr) => {
4664                v.visit_tensor_mut(&mut repr.input_re);
4665                v.visit_tensor_mut(&mut repr.input_im);
4666                v.visit_tensor_mut(&mut repr.out_signal);
4667            }
4668            ModuleOperationIr::Attention(repr) => {
4669                v.visit_tensor_mut(&mut repr.query);
4670                v.visit_tensor_mut(&mut repr.key);
4671                v.visit_tensor_mut(&mut repr.value);
4672                if let Some(mask) = &mut repr.mask {
4673                    v.visit_tensor_mut(mask);
4674                }
4675                if let Some(attn_bias) = &mut repr.attn_bias {
4676                    v.visit_tensor_mut(attn_bias);
4677                }
4678                v.visit_tensor_mut(&mut repr.out);
4679                if let Some(scale) = &mut repr.options.scale {
4680                    v.visit_scalar_mut(scale);
4681                }
4682                if let Some(softcap) = &mut repr.options.softcap {
4683                    v.visit_scalar_mut(softcap);
4684                }
4685            }
4686            ModuleOperationIr::CtcLoss(repr) => {
4687                v.visit_tensor_mut(&mut repr.log_probs);
4688                v.visit_tensor_mut(&mut repr.targets);
4689                v.visit_tensor_mut(&mut repr.input_lengths);
4690                v.visit_tensor_mut(&mut repr.target_lengths);
4691                v.visit_tensor_mut(&mut repr.out);
4692            }
4693            ModuleOperationIr::CtcLossBackward(repr) => {
4694                v.visit_tensor_mut(&mut repr.log_probs);
4695                v.visit_tensor_mut(&mut repr.targets);
4696                v.visit_tensor_mut(&mut repr.input_lengths);
4697                v.visit_tensor_mut(&mut repr.target_lengths);
4698                v.visit_tensor_mut(&mut repr.grad_loss);
4699                v.visit_tensor_mut(&mut repr.out);
4700            }
4701            ModuleOperationIr::LayerNorm(repr) => {
4702                v.visit_tensor_mut(&mut repr.input);
4703                v.visit_tensor_mut(&mut repr.gamma);
4704                if let Some(beta) = &mut repr.beta {
4705                    v.visit_tensor_mut(beta);
4706                }
4707                v.visit_tensor_mut(&mut repr.out);
4708                v.visit_scalar_mut(&mut repr.epsilon);
4709            }
4710            ModuleOperationIr::Unfold4d(repr) => {
4711                v.visit_tensor_mut(&mut repr.x);
4712                v.visit_tensor_mut(&mut repr.out);
4713            }
4714            ModuleOperationIr::ConvTranspose1dWeightBackward(repr) => {
4715                v.visit_tensor_mut(&mut repr.x);
4716                v.visit_tensor_mut(&mut repr.weight);
4717                v.visit_tensor_mut(&mut repr.output_grad);
4718                v.visit_tensor_mut(&mut repr.out);
4719            }
4720            ModuleOperationIr::ConvTranspose1dBiasBackward(repr) => {
4721                v.visit_tensor_mut(&mut repr.x);
4722                v.visit_tensor_mut(&mut repr.bias);
4723                v.visit_tensor_mut(&mut repr.output_grad);
4724                v.visit_tensor_mut(&mut repr.out);
4725            }
4726            ModuleOperationIr::ConvTranspose2dWeightBackward(repr) => {
4727                v.visit_tensor_mut(&mut repr.x);
4728                v.visit_tensor_mut(&mut repr.weight);
4729                v.visit_tensor_mut(&mut repr.output_grad);
4730                v.visit_tensor_mut(&mut repr.out);
4731            }
4732            ModuleOperationIr::ConvTranspose2dBiasBackward(repr) => {
4733                v.visit_tensor_mut(&mut repr.x);
4734                v.visit_tensor_mut(&mut repr.bias);
4735                v.visit_tensor_mut(&mut repr.output_grad);
4736                v.visit_tensor_mut(&mut repr.out);
4737            }
4738            ModuleOperationIr::ConvTranspose3dWeightBackward(repr) => {
4739                v.visit_tensor_mut(&mut repr.x);
4740                v.visit_tensor_mut(&mut repr.weight);
4741                v.visit_tensor_mut(&mut repr.output_grad);
4742                v.visit_tensor_mut(&mut repr.out);
4743            }
4744            ModuleOperationIr::ConvTranspose3dBiasBackward(repr) => {
4745                v.visit_tensor_mut(&mut repr.x);
4746                v.visit_tensor_mut(&mut repr.bias);
4747                v.visit_tensor_mut(&mut repr.output_grad);
4748                v.visit_tensor_mut(&mut repr.out);
4749            }
4750        }
4751    }
4752}
4753
4754impl DistributedOperationIr {
4755    fn inputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
4756        match self {
4757            DistributedOperationIr::AllReduce(repr) => Box::new([&repr.tensor].into_iter()),
4758            DistributedOperationIr::SyncCollective => Box::new([].into_iter()),
4759        }
4760    }
4761
4762    fn outputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
4763        match self {
4764            DistributedOperationIr::AllReduce(repr) => Box::new([&repr.out].into_iter()),
4765            DistributedOperationIr::SyncCollective => Box::new([].into_iter()),
4766        }
4767    }
4768
4769    fn mark_read_only(&mut self, nodes: &[TensorId]) -> Vec<TensorIr> {
4770        let mut output = Vec::new();
4771
4772        match self {
4773            DistributedOperationIr::AllReduce(repr) => {
4774                repr.tensor.mark_read_only(nodes, &mut output);
4775            }
4776            DistributedOperationIr::SyncCollective => {}
4777        }
4778
4779        output
4780    }
4781
4782    fn visit_mut(&mut self, v: &mut impl IrVisitorMut) {
4783        match self {
4784            DistributedOperationIr::AllReduce(repr) => {
4785                v.visit_tensor_mut(&mut repr.tensor);
4786                v.visit_tensor_mut(&mut repr.out);
4787            }
4788            DistributedOperationIr::SyncCollective => {}
4789        }
4790    }
4791}
4792
4793impl InitOperationIr {
4794    fn inputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
4795        Box::new([].into_iter())
4796    }
4797    fn outputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
4798        Box::new([&self.out].into_iter())
4799    }
4800
4801    fn visit_mut(&mut self, v: &mut impl IrVisitorMut) {
4802        v.visit_tensor_mut(&mut self.out);
4803    }
4804}
4805
4806impl TensorIr {
4807    fn mark_read_only(&mut self, nodes: &[TensorId], output: &mut Vec<TensorIr>) {
4808        if self.status == TensorStatus::ReadWrite && nodes.contains(&self.id) {
4809            output.push(self.clone());
4810            self.status = TensorStatus::ReadOnly;
4811        }
4812    }
4813}
4814
4815impl core::hash::Hash for RandomOpIr {
4816    fn hash<H: core::hash::Hasher>(&self, state: &mut H) {
4817        self.out.hash(state);
4818
4819        match self.distribution {
4820            Distribution::Default => 1u8.hash(state),
4821            Distribution::Bernoulli(_) => 2u8.hash(state),
4822            Distribution::Uniform(_, _) => 3u8.hash(state),
4823            Distribution::Normal(_, _) => 4u8.hash(state),
4824        }
4825    }
4826}
4827
4828/// Extension trait to extract outputs when registering an operation.
4829pub trait OperationOutput<O> {
4830    /// Extract a single output.
4831    fn output(self) -> O;
4832
4833    /// Extract a fixed number of outputs.
4834    fn outputs<const N: usize>(self) -> [O; N];
4835}
4836
4837impl<O: core::fmt::Debug> OperationOutput<O> for Vec<O> {
4838    fn output(self) -> O {
4839        let [tensor] = self.outputs();
4840        tensor
4841    }
4842
4843    fn outputs<const N: usize>(self) -> [O; N] {
4844        self.try_into().unwrap()
4845    }
4846}
4847
4848/// Operation IR for sort along a dim. The output preserves the input shape/dtype.
4849#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4850pub struct SortOpIr {
4851    /// Input tensor.
4852    pub input: TensorIr,
4853    /// Dim along which to sort.
4854    pub dim: usize,
4855    /// Sort descending.
4856    pub descending: bool,
4857    /// Output tensor (same shape/dtype as input).
4858    pub out: TensorIr,
4859}
4860
4861/// Operation IR for sort-with-indices: returns sorted values + source indices.
4862#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4863pub struct SortWithIndicesOpIr {
4864    /// Input tensor.
4865    pub input: TensorIr,
4866    /// Dim along which to sort.
4867    pub dim: usize,
4868    /// Sort descending.
4869    pub descending: bool,
4870    /// Output tensor with sorted values.
4871    pub out: TensorIr,
4872    /// Output tensor with the indices into the original input.
4873    pub out_indices: TensorIr,
4874}
4875
4876/// Operation IR for layer normalization with optional bias.
4877#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4878pub struct LayerNormOpIr {
4879    /// Input tensor.
4880    pub input: TensorIr,
4881    /// Scale (gamma) parameter.
4882    pub gamma: TensorIr,
4883    /// Optional shift (beta) parameter.
4884    pub beta: Option<TensorIr>,
4885    /// Numerical-stability epsilon.
4886    pub epsilon: ScalarIr,
4887    /// Output tensor.
4888    pub out: TensorIr,
4889}
4890
4891/// Operation IR for unfold4d (a 4d sliding-window kernel-as-explicit-columns reshape).
4892#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4893pub struct Unfold4dOpIr {
4894    /// Input tensor of shape `[N, C, H, W]`.
4895    pub x: TensorIr,
4896    /// Kernel size `[kH, kW]`.
4897    pub kernel_size: [usize; 2],
4898    /// Conv-like options (stride, padding, dilation).
4899    pub options: Unfold4dOptionsIr,
4900    /// Output tensor.
4901    pub out: TensorIr,
4902}
4903
4904/// Options for [`Unfold4dOpIr`] — mirrors the backend's `UnfoldOptions`.
4905#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4906pub struct Unfold4dOptionsIr {
4907    /// Stride `[sH, sW]`.
4908    pub stride: [usize; 2],
4909    /// Padding `[pH, pW]`.
4910    pub padding: [usize; 2],
4911    /// Dilation `[dH, dW]`.
4912    pub dilation: [usize; 2],
4913}
4914
4915/// Operation IR for the weight-gradient of `conv_transpose1d`.
4916#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4917pub struct ConvTranspose1dWeightBackwardOpIr {
4918    /// Forward input.
4919    pub x: TensorIr,
4920    /// Convolution weights.
4921    pub weight: TensorIr,
4922    /// Upstream gradient.
4923    pub output_grad: TensorIr,
4924    /// Convolution options.
4925    pub options: ConvTranspose1dOptionsIr,
4926    /// Output: gradient w.r.t. `weight`.
4927    pub out: TensorIr,
4928}
4929
4930/// Operation IR for the bias-gradient of `conv_transpose1d`.
4931#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4932pub struct ConvTranspose1dBiasBackwardOpIr {
4933    /// Forward input.
4934    pub x: TensorIr,
4935    /// Bias tensor (sets the output shape).
4936    pub bias: TensorIr,
4937    /// Upstream gradient.
4938    pub output_grad: TensorIr,
4939    /// Output: gradient w.r.t. `bias`.
4940    pub out: TensorIr,
4941}
4942
4943/// Operation IR for the weight-gradient of `conv_transpose2d`.
4944#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4945pub struct ConvTranspose2dWeightBackwardOpIr {
4946    /// Forward input.
4947    pub x: TensorIr,
4948    /// Convolution weights.
4949    pub weight: TensorIr,
4950    /// Upstream gradient.
4951    pub output_grad: TensorIr,
4952    /// Convolution options.
4953    pub options: ConvTranspose2dOptionsIr,
4954    /// Output: gradient w.r.t. `weight`.
4955    pub out: TensorIr,
4956}
4957
4958/// Operation IR for the bias-gradient of `conv_transpose2d`.
4959#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4960pub struct ConvTranspose2dBiasBackwardOpIr {
4961    /// Forward input.
4962    pub x: TensorIr,
4963    /// Bias tensor.
4964    pub bias: TensorIr,
4965    /// Upstream gradient.
4966    pub output_grad: TensorIr,
4967    /// Output: gradient w.r.t. `bias`.
4968    pub out: TensorIr,
4969}
4970
4971/// Operation IR for the weight-gradient of `conv_transpose3d`.
4972#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4973pub struct ConvTranspose3dWeightBackwardOpIr {
4974    /// Forward input.
4975    pub x: TensorIr,
4976    /// Convolution weights.
4977    pub weight: TensorIr,
4978    /// Upstream gradient.
4979    pub output_grad: TensorIr,
4980    /// Convolution options.
4981    pub options: ConvTranspose3dOptionsIr,
4982    /// Output: gradient w.r.t. `weight`.
4983    pub out: TensorIr,
4984}
4985
4986/// Operation IR for the bias-gradient of `conv_transpose3d`.
4987#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
4988pub struct ConvTranspose3dBiasBackwardOpIr {
4989    /// Forward input.
4990    pub x: TensorIr,
4991    /// Bias tensor.
4992    pub bias: TensorIr,
4993    /// Upstream gradient.
4994    pub output_grad: TensorIr,
4995    /// Output: gradient w.r.t. `bias`.
4996    pub out: TensorIr,
4997}
4998
4999/// Operation IR for the hard-sigmoid activation function, which takes two scalars
5000/// (`alpha` and `beta`) in addition to the input tensor.
5001#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
5002pub struct HardSigmoidOpIr {
5003    /// Input tensor.
5004    pub tensor: TensorIr,
5005    /// Alpha — multiplied with the tensor before adding `beta`.
5006    pub alpha: ScalarIr,
5007    /// Beta — added after multiplying by `alpha`.
5008    pub beta: ScalarIr,
5009    /// Output tensor.
5010    pub out: TensorIr,
5011}
5012
5013/// Operation intermediate representation for activation functions.
5014///
5015/// Activations are kept in their own enum (rather than smeared across `FloatOperationIr`)
5016/// so backends can match a single arm to dispatch the fused implementation, instead of
5017/// receiving the function decomposed into 5–20 primitive ops.
5018#[derive(Clone, Debug, Hash, PartialEq, Serialize, Deserialize)]
5019pub enum ActivationOperationIr {
5020    /// [relu](burn_backend::ops::ActivationOps::relu).
5021    Relu(UnaryOpIr),
5022    /// [relu_backward](burn_backend::ops::ActivationOps::relu_backward).
5023    /// `lhs` = output of the forward pass, `rhs` = upstream gradient.
5024    ReluBackward(BinaryOpIr),
5025    /// [leaky_relu](burn_backend::ops::ActivationOps::leaky_relu).
5026    /// `lhs` = input, `rhs` = `negative_slope` scalar.
5027    LeakyRelu(ScalarOpIr),
5028    /// [prelu](burn_backend::ops::ActivationOps::prelu).
5029    /// `lhs` = input, `rhs` = alpha tensor.
5030    PRelu(BinaryOpIr),
5031    /// [gelu](burn_backend::ops::ActivationOps::gelu).
5032    Gelu(UnaryOpIr),
5033    /// [gelu_backward](burn_backend::ops::ActivationOps::gelu_backward).
5034    /// `lhs` = forward input, `rhs` = upstream gradient.
5035    GeluBackward(BinaryOpIr),
5036    /// [sigmoid](burn_backend::ops::ActivationOps::sigmoid).
5037    Sigmoid(UnaryOpIr),
5038    /// [sigmoid_backward](burn_backend::ops::ActivationOps::sigmoid_backward).
5039    /// `lhs` = output of the forward pass, `rhs` = upstream gradient.
5040    SigmoidBackward(BinaryOpIr),
5041    /// [hard_sigmoid](burn_backend::ops::ActivationOps::hard_sigmoid).
5042    HardSigmoid(HardSigmoidOpIr),
5043    /// [log_sigmoid](burn_backend::ops::ActivationOps::log_sigmoid).
5044    LogSigmoid(UnaryOpIr),
5045    /// [log_sigmoid_backward](burn_backend::ops::ActivationOps::log_sigmoid_backward).
5046    /// `lhs` = forward input, `rhs` = upstream gradient.
5047    LogSigmoidBackward(BinaryOpIr),
5048    /// [softmax](burn_backend::ops::ActivationOps::softmax).
5049    Softmax(DimOpIr),
5050    /// [log_softmax](burn_backend::ops::ActivationOps::log_softmax).
5051    LogSoftmax(DimOpIr),
5052    /// [softmin](burn_backend::ops::ActivationOps::softmin).
5053    Softmin(DimOpIr),
5054}
5055
5056/// Generate [`ActivationOperationIr`]'s `inputs`/`outputs`/`mark_read_only` from a single
5057/// variant → input-field table, so the three methods can't drift out of sync. The `out`
5058/// field is implicit (every variant's op IR has one). Each entry lists the variant and the
5059/// names of its input tensor fields (`input`, `lhs`/`rhs`, `tensor`, …).
5060macro_rules! activation_ir_tensor_access {
5061    ($( $variant:ident => [ $($field:ident),+ ] ),+ $(,)?) => {
5062        impl ActivationOperationIr {
5063            fn inputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
5064                match self {
5065                    $( Self::$variant(repr) => Box::new([$(&repr.$field),+].into_iter()), )+
5066                }
5067            }
5068
5069            fn outputs(&self) -> Box<dyn Iterator<Item = &TensorIr> + '_> {
5070                match self {
5071                    $( Self::$variant(repr) => Box::new([&repr.out].into_iter()), )+
5072                }
5073            }
5074
5075            fn mark_read_only(&mut self, nodes: &[TensorId]) -> Vec<TensorIr> {
5076                let mut output = Vec::new();
5077                match self {
5078                    $( Self::$variant(repr) => {
5079                        $( repr.$field.mark_read_only(nodes, &mut output); )+
5080                    } )+
5081                }
5082                output
5083            }
5084        }
5085    };
5086}
5087
5088impl ActivationOperationIr {
5089    fn visit_mut(&mut self, v: &mut impl IrVisitorMut) {
5090        match self {
5091            ActivationOperationIr::Relu(repr) => {
5092                v.visit_tensor_mut(&mut repr.input);
5093                v.visit_tensor_mut(&mut repr.out);
5094            }
5095            ActivationOperationIr::ReluBackward(repr) => {
5096                v.visit_tensor_mut(&mut repr.lhs);
5097                v.visit_tensor_mut(&mut repr.rhs);
5098                v.visit_tensor_mut(&mut repr.out);
5099            }
5100            ActivationOperationIr::LeakyRelu(repr) => {
5101                v.visit_tensor_mut(&mut repr.lhs);
5102                v.visit_tensor_mut(&mut repr.out);
5103                v.visit_scalar_mut(&mut repr.rhs);
5104            }
5105            ActivationOperationIr::PRelu(repr) => {
5106                v.visit_tensor_mut(&mut repr.lhs);
5107                v.visit_tensor_mut(&mut repr.rhs);
5108                v.visit_tensor_mut(&mut repr.out);
5109            }
5110            ActivationOperationIr::Gelu(repr) => {
5111                v.visit_tensor_mut(&mut repr.input);
5112                v.visit_tensor_mut(&mut repr.out);
5113            }
5114            ActivationOperationIr::GeluBackward(repr) => {
5115                v.visit_tensor_mut(&mut repr.lhs);
5116                v.visit_tensor_mut(&mut repr.rhs);
5117                v.visit_tensor_mut(&mut repr.out);
5118            }
5119            ActivationOperationIr::Sigmoid(repr) => {
5120                v.visit_tensor_mut(&mut repr.input);
5121                v.visit_tensor_mut(&mut repr.out);
5122            }
5123            ActivationOperationIr::SigmoidBackward(repr) => {
5124                v.visit_tensor_mut(&mut repr.lhs);
5125                v.visit_tensor_mut(&mut repr.rhs);
5126                v.visit_tensor_mut(&mut repr.out);
5127            }
5128            ActivationOperationIr::HardSigmoid(repr) => {
5129                v.visit_tensor_mut(&mut repr.tensor);
5130                v.visit_tensor_mut(&mut repr.out);
5131                v.visit_scalar_mut(&mut repr.alpha);
5132                v.visit_scalar_mut(&mut repr.beta);
5133            }
5134            ActivationOperationIr::LogSigmoid(repr) => {
5135                v.visit_tensor_mut(&mut repr.input);
5136                v.visit_tensor_mut(&mut repr.out);
5137            }
5138            ActivationOperationIr::LogSigmoidBackward(repr) => {
5139                v.visit_tensor_mut(&mut repr.lhs);
5140                v.visit_tensor_mut(&mut repr.rhs);
5141                v.visit_tensor_mut(&mut repr.out);
5142            }
5143            ActivationOperationIr::Softmax(repr) => {
5144                v.visit_tensor_mut(&mut repr.input);
5145                v.visit_tensor_mut(&mut repr.out);
5146            }
5147            ActivationOperationIr::LogSoftmax(repr) => {
5148                v.visit_tensor_mut(&mut repr.input);
5149                v.visit_tensor_mut(&mut repr.out);
5150            }
5151            ActivationOperationIr::Softmin(repr) => {
5152                v.visit_tensor_mut(&mut repr.input);
5153                v.visit_tensor_mut(&mut repr.out);
5154            }
5155        }
5156    }
5157}
5158
5159activation_ir_tensor_access! {
5160    Relu => [input],
5161    ReluBackward => [lhs, rhs],
5162    LeakyRelu => [lhs],
5163    PRelu => [lhs, rhs],
5164    Gelu => [input],
5165    GeluBackward => [lhs, rhs],
5166    Sigmoid => [input],
5167    SigmoidBackward => [lhs, rhs],
5168    HardSigmoid => [tensor],
5169    LogSigmoid => [input],
5170    LogSigmoidBackward => [lhs, rhs],
5171    Softmax => [input],
5172    LogSoftmax => [input],
5173    Softmin => [input],
5174}
5175
5176#[cfg(test)]
5177mod visit_mut_tests {
5178    use super::*;
5179    use burn_backend::{DType, Shape};
5180
5181    fn tensor(id: u64) -> TensorIr {
5182        TensorIr::uninit(TensorId::new(id), Shape::from([2, 2]), DType::F32)
5183    }
5184
5185    /// Bumps every visited tensor id by 100 and collects (and rewrites) every visited scalar.
5186    #[derive(Default)]
5187    struct CollectVisitor {
5188        scalars: Vec<ScalarIr>,
5189        rewrite_scalar: Option<ScalarIr>,
5190    }
5191
5192    impl IrVisitorMut for CollectVisitor {
5193        fn visit_tensor_mut(&mut self, tensor: &mut TensorIr) {
5194            tensor.id = TensorId::new(tensor.id.value() + 100);
5195        }
5196
5197        fn visit_scalar_mut(&mut self, scalar: &mut ScalarIr) {
5198            self.scalars.push(*scalar);
5199            if let Some(value) = self.rewrite_scalar {
5200                *scalar = value;
5201            }
5202        }
5203    }
5204
5205    #[test]
5206    fn visit_mut_visits_all_tensors_and_scalars() {
5207        // NumericFloat MulScalar (a ScalarOpIr with a scalar `rhs`).
5208        let mut mul = OperationIr::NumericFloat(
5209            DType::F32,
5210            NumericOperationIr::MulScalar(ScalarOpIr {
5211                lhs: tensor(1),
5212                rhs: ScalarIr::Float(2.0),
5213                out: tensor(2),
5214            }),
5215        );
5216
5217        // Bump every tensor id by 100 and rewrite the scalar to 9.0.
5218        let mut visitor = CollectVisitor {
5219            rewrite_scalar: Some(ScalarIr::Float(9.0)),
5220            ..Default::default()
5221        };
5222        mul.visit_mut(&mut visitor);
5223
5224        let ids: Vec<u64> = mul
5225            .inputs()
5226            .chain(mul.outputs())
5227            .map(|t| t.id.value())
5228            .collect();
5229        assert_eq!(ids, vec![101, 102]);
5230        // The scalar must have been visited (and was rewritable).
5231        assert_eq!(visitor.scalars, vec![ScalarIr::Float(2.0)]);
5232
5233        // Visiting again observes the rewritten scalar.
5234        let mut after = CollectVisitor::default();
5235        mul.visit_mut(&mut after);
5236        assert_eq!(after.scalars, vec![ScalarIr::Float(9.0)]);
5237
5238        // BaseFloat Reshape (a ShapeOpIr, input + out, no scalars).
5239        let mut reshape = OperationIr::BaseFloat(BaseOperationIr::Reshape(ShapeOpIr {
5240            input: tensor(10),
5241            out: tensor(11),
5242        }));
5243
5244        let mut visitor = CollectVisitor::default();
5245        reshape.visit_mut(&mut visitor);
5246        let ids: Vec<u64> = reshape
5247            .inputs()
5248            .chain(reshape.outputs())
5249            .map(|t| t.id.value())
5250            .collect();
5251        assert_eq!(ids, vec![110, 111]);
5252        assert_eq!(visitor.scalars.len(), 0);
5253    }
5254}