tract_core/ops/array/
trilu.rs1use 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}