【问题标题】:What's the fastest way of finding the index of the maximum value in an array?在数组中找到最大值索引的最快方法是什么?
【发布时间】:2019-09-05 23:11:46
【问题描述】:

我有一个f32 类型的二维数组(来自ndarray::ArrayView2),我想在每一行中找到最大值的索引,并将索引值放入另一个数组中。

Python 中的等价物是这样的:

import numpy as np

for i in range (0, max_val, batch_size):
   sims = xp.dot(batch, vectors.T) 
   # sims is the dot product of batch and vectors.T
   # the shape is, for example, (1024, 10000)

   best_rows[i: i+batch_size] = sims.argmax(axis = 1)

在 Python 中,函数 .argmax 非常快,但我在 Rust 中没有看到类似的函数。最快的方法是什么?

【问题讨论】:

  • 它实际上是一个数组还是Vec
  • 你需要在这里展示一些 Rust 代码作为上下文。
  • @jhpratt 正如问题中提到的,它是一个数组。
  • 是的,但是很多时候人们在实际使用向量时会说数组。不过,您添加实际类型的说明很有帮助。

标签: arrays rust


【解决方案1】:

考虑一般Ord 类型的简单情况:答案会略有不同,具体取决于您是否知道值是Copy,但这里是代码:

fn position_max_copy<T: Ord + Copy>(slice: &[T]) -> Option<usize> {
    slice.iter().enumerate().max_by_key(|(_, &value)| value).map(|(idx, _)| idx)
}

fn position_max<T: Ord>(slice: &[T]) -> Option<usize> {
    slice.iter().enumerate().max_by(|(_, value0), (_, value1)| value0.cmp(value1)).map(|(idx, _)| idx)
}

基本思想是我们将数组中的每个项目(实际上是一个切片——不管它是 Vec 还是数组或其他更奇特的东西)与其索引配对,使用 @987654328 @函数只根据值(不是索引)找到最大值,然后只返回索引。如果切片为空 None 将被返回。根据文档,将返回最右边的索引;如果你需要最左边的,请在rev() 之后 enumerate()

rev()enumerate()max_by_key()max_by() 记录在 hereslice::iter() 已记录在 here 中(但作为 rust 开发人员,您需要在没有文档的情况下将其列入您的候选清单); mapOption::map() 记录在案的 here(同上)。哦,cmpOrd::cmp,但大多数时候您可以使用不需要它的 Copy 版本(例如,如果您正在比较整数)。


现在要注意了:f32 不是 Ord,因为 IEEE 浮点数的工作方式。大多数语言都忽略了这一点,并且有一些微妙的错误算法。在Ord 上提供总订单(通过声明所有NaN 相等,并且大于所有数字)的最受欢迎的板条箱似乎是ordered-float。假设它被正确实施,它应该是非常轻量级的。它确实引入了num_traits,但这是最流行的数字库的一部分,因此很可能已经被其他依赖项引入了。

在这种情况下,您可以通过在切片迭代器 (slice.iter().map(ordered_float::OrderedFloat)) 上映射 ordered_float::OrderedFloat(元组类型的“构造函数”)来使用它。由于您只想要最大元素的位置,因此无需提取 f32。

【讨论】:

  • 请注意,这是一维向量,但 OP 使用的是二维数组,因此他需要遍历其数组的行并为每一行调用 position_max
  • 对于一维情况,另一个选项是(0..slice.len()).max_by_key(|i| &amp;slice[i])。 (我没有对此进行测试,但无论T: Copy是否都应该工作。)
  • 是的,这可能更容易理解 TBH。 max_by_key 对于T: !Copy 的问题在于它(隐式)要求返回类型为Ord + 'static;我不确定这对于算法是否真的有必要,或者可能是疏忽。
【解决方案2】:

approach from @David A 很酷,但如前所述,有一个问题:f32f64 不实现 Ord::cmp。 (这真的很痛苦。)

有多种解决方法:您可以自己实现cmp,也可以使用ordered-float等。

就我而言,这是一个更大项目的一部分,我们在使用外部包时非常小心。此外,我很确定我们没有任何 NaN 值。因此我更喜欢使用fold,如果您仔细查看max_by_key 源代码,他们也一直在使用它。

for (i, row) in matrix.axis_iter(Axis(1)).enumerate() {
    let (max_idx, max_val) =
        row.iter()
            .enumerate()
            .fold((0, row[0]), |(idx_max, val_max), (idx, val)| {
                if &val_max > val {
                    (idx_max, val_max)
                } else {
                    (idx, *val)
                }
            });
}

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2011-01-10
    • 2017-03-15
    • 2021-03-23
    相关资源
    最近更新 更多