Skip to main content

burn_ir/
operation.rs

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