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