Skip to main content

tract_core/ops/array/
topk.rs

1use std::cmp::Ordering;
2
3use tract_data::itertools::Itertools;
4use tract_ndarray::{ArrayViewMutD, Axis, Dimension};
5
6use crate::internal::*;
7
8#[derive(Debug, Clone, new, Default, Hash, PartialEq, Eq)]
9pub struct Topk {
10    pub axis: usize,
11    pub largest: bool,
12    pub fallback_k: TDim,
13}
14
15impl Op for Topk {
16    fn name(&self) -> StaticName {
17        "Topk".into()
18    }
19
20    op_as_typed_op!();
21}
22
23impl EvalOp for Topk {
24    op_out_of_plan!();
25
26    fn eval(&self, _ctx: &EvalContext, inputs: TVec<TValue>) -> TractResult<TVec<TValue>> {
27        let (input, k) = args_2!(inputs);
28        let mut output_shape: TVec<usize> = input.shape().into();
29        let k = k.cast_to_scalar::<i64>()? as usize;
30        output_shape[self.axis] = k;
31        let dt = input.datum_type();
32        let mut output_values = Tensor::zero_dt(dt, &output_shape)?;
33        let mut output_indices = Tensor::zero::<i64>(&output_shape)?;
34        let mut iterating_shape = output_shape.clone();
35        iterating_shape[self.axis] = 1;
36        let mut output_indices_plain = output_indices.try_as_plain_mut()?;
37        let mut output_indices_view = output_indices_plain.to_array_view_mut::<i64>()?;
38        for coords in tract_ndarray::indices(&*iterating_shape) {
39            let mut coords: TVec<usize> = coords.as_array_view().as_slice().unwrap().into();
40            dispatch_numbers!(Self::inner_loop_t(dt)(
41                self,
42                &mut coords,
43                &input,
44                &mut output_values,
45                &mut output_indices_view,
46                k
47            ))?;
48        }
49        Ok(tvec!(output_values.into_tvalue(), output_indices.into_tvalue()))
50    }
51}
52
53impl Topk {
54    fn inner_loop_t<T: Datum + PartialOrd>(
55        &self,
56        coords: &mut [usize],
57        input: &Tensor,
58        output_values: &mut Tensor,
59        output_indices_view: &mut ArrayViewMutD<i64>,
60        k: usize,
61    ) -> TractResult<()> {
62        let mut output_values_plain = output_values.try_as_plain_mut()?;
63        let mut output_values_view = output_values_plain.to_array_view_mut::<T>()?;
64        let mut view = input.to_plain_array_view::<T>()?;
65        for (ix, x) in coords.iter().enumerate() {
66            if ix != self.axis {
67                view.collapse_axis(Axis(ix), *x);
68            }
69        }
70        for (ix, (argmax, max)) in view
71            .iter()
72            .cloned()
73            .enumerate()
74            .sorted_by(|a, b| {
75                let ord = { a.1.partial_cmp(&b.1).unwrap_or(Ordering::Less) };
76                if self.largest { ord.reverse() } else { ord }
77            })
78            .take(k)
79            .enumerate()
80        {
81            coords[self.axis] = ix;
82            output_values_view[&*coords] = max;
83            output_indices_view[&*coords] = argmax as i64;
84        }
85        Ok(())
86    }
87}
88
89impl TypedOp for Topk {
90    fn output_facts(&self, inputs: &[&TypedFact]) -> TractResult<TVec<TypedFact>> {
91        let mut fact_values = inputs[0].without_value();
92        let mut fact_indices = inputs[0].without_value();
93        let k: TDim = if let Some(k) = &inputs[1].konst {
94            k.cast_to::<TDim>()?.try_as_plain()?.to_scalar::<TDim>()?.clone()
95        } else {
96            self.fallback_k.clone()
97        };
98        fact_values.shape.set(self.axis, k.clone());
99        fact_indices.shape.set(self.axis, k);
100        fact_indices.datum_type = i64::datum_type();
101        Ok(tvec!(fact_values, fact_indices))
102    }
103
104    fn set_symbols(
105        &self,
106        _source: &TypedModel,
107        node: &TypedNode,
108        target: &mut TypedModel,
109        mapping: &HashMap<OutletId, OutletId>,
110        subs: &HashMap<Symbol, TDim>,
111    ) -> TractResult<TVec<OutletId>> {
112        let op = Topk {
113            axis: self.axis,
114            largest: self.largest,
115            fallback_k: self.fallback_k.substitute_all(subs)?,
116        };
117        let inputs = node.inputs.iter().map(|i| mapping[i]).collect::<TVec<_>>();
118        target.wire_node(&node.name, op, &inputs)
119    }
120
121    as_op!();
122}