Skip to main content

tract_core/ops/array/
trilu.rs

1use crate::internal::*;
2
3#[derive(Debug, Clone, PartialEq, Eq)]
4pub struct Trilu {
5    pub upper: bool,
6}
7
8impl Op for Trilu {
9    fn name(&self) -> StaticName {
10        "Trilu".into()
11    }
12
13    op_as_typed_op!();
14}
15
16impl EvalOp for Trilu {
17    op_out_of_plan!();
18
19    fn eval(&self, _ctx: &EvalContext, inputs: TVec<TValue>) -> TractResult<TVec<TValue>> {
20        let (input, k) = args_2!(inputs);
21        let mut input = input.into_tensor();
22        let k = *k.try_as_plain()?.to_scalar::<i64>()?;
23        fn eval_t<T: Datum>(tensor: &mut Tensor, upper: bool, k: i64) -> TractResult<()> {
24            let mut tensor_plain = tensor.try_as_plain_mut()?;
25            let mut view = tensor_plain.to_array_view_mut::<T>()?;
26            for coords in tract_ndarray::indices(view.shape()) {
27                let row = coords[view.ndim() - 2] as i64;
28                let col = coords[view.ndim() - 1] as i64;
29                if upper {
30                    if col < row + k {
31                        view[coords] = T::default();
32                    }
33                } else if col > row + k {
34                    view[coords] = T::default();
35                }
36            }
37            Ok(())
38        }
39        dispatch_datum!(eval_t(input.datum_type())(&mut input, self.upper, k))?;
40        Ok(tvec!(input.into_tvalue()))
41    }
42}
43
44impl TypedOp for Trilu {
45    fn output_facts(&self, inputs: &[&TypedFact]) -> TractResult<TVec<TypedFact>> {
46        Ok(tvec!(inputs[0].without_value()))
47    }
48
49    as_op!();
50}