如何在Rust中获取可对含NaN的浮点向量排序的索引?
实现含NaN浮点向量的argsort(返回排序索引)
要实现不对原向量排序、同时能处理NaN的浮点向量排序索引功能,我们可以结合你提供的两段代码逻辑:通过生成索引向量,再用处理NaN的自定义规则对索引进行排序。
完整实现代码
use std::cmp::Ordering; fn argsort_floats(data: &[f64]) -> Vec<usize> { // 生成从0到数据长度的索引向量 let mut indices = (0..data.len()).collect::<Vec<_>>(); // 对索引向量使用自定义排序规则,规则和你提供的浮点排序逻辑一致 indices.sort_by(|&i, &j| { let a = data[i]; let b = data[j]; match (a.is_nan(), b.is_nan()) { // 两个都是NaN,视为相等 (true, true) => Ordering::Equal, // 当前元素是NaN,排在后面 (true, false) => Ordering::Greater, // 另一个元素是NaN,当前元素排在前面 (false, true) => Ordering::Less, // 都不是NaN,直接用浮点比较 (false, false) => a.partial_cmp(&b).unwrap(), } }); indices }
关键说明
- 不修改原输入向量:我们只操作索引向量,原数据保持完全不变。
- 替换排序逻辑:放弃
sort_by_key(它要求键实现Ord,而浮点类型仅支持PartialOrd),改用sort_by实现自定义比较规则,和你提供的浮点排序函数逻辑完全对齐——确保NaN排在所有非NaN元素之后。 - 安全的
unwrap:在排除NaN的情况下,partial_cmp一定会返回有效的Ordering,所以这里的unwrap是安全的。
测试示例
fn main() { let nums = vec![3.1, f64::NAN, 1.5, 2.2, f64::NAN, 0.8]; let sorted_indices = argsort_floats(&nums); // 根据索引取出原向量元素,验证排序结果 let sorted_nums: Vec<f64> = sorted_indices.iter().map(|&i| nums[i]).collect(); println!("排序后的向量:{:?}", sorted_nums); // 输出:[0.8, 1.5, 2.2, 3.1, NaN, NaN] }
内容的提问来源于stack exchange,提问作者PyRsquared
相关产品推荐
相关产品推荐

