Rust Arrayfire如何用数组索引提取目标对应预测置信度?
使用Rust Arrayfire实现类Numpy的索引提取正确置信度
可以用Arrayfire内置的gather函数高效实现类似Numpy的索引操作,完全适配大样本量场景,替代效率低下的Seqs向量方案,具体实现如下:
核心实现步骤
- 利用Arrayfire的列主序特性,先转置预测数组调整维度顺序,方便按样本索引提取对应类别置信度
- 构造样本行索引序列,结合目标标签数组,通过
gather直接提取每个样本的正确类别置信度 - 基于提取结果计算负自然对数的均值得到交叉熵损失
示例代码
use arrayfire::{self as af, Dim4}; fn cross_entropy_loss(predictions: &af::Array<f32>, targets: &af::Array<u32>) -> f32 { // 获取样本总数:predictions默认维度为(类别数, 样本数),取第二维度值 let num_samples = predictions.dims()[1] as u64; // 生成样本行索引序列:0到num_samples-1 let row_indices = af::range(Dim4::new(&[num_samples, 1, 1, 1]), 0, f32::default()); // 转置预测数组为(样本数, 类别数),适配gather的索引逻辑 let pred_transposed = af::transpose(predictions, false); // 按目标标签索引提取每个样本的正确类别置信度 let correct_confidences = af::gather(&pred_transposed, targets, 1); // 计算负对数并求均值 let neg_log = af::neg(af::log(&correct_confidences)); af::mean_all(&neg_log).0 }
关键细节说明
- Arrayfire默认采用列主序存储,所以初始预测数组维度为
(类别数, 样本数),转置后变为(样本数, 类别数),才能通过gather按第二维度(类别维度)用目标标签索引取值 gather是Arrayfire底层优化的索引操作,支持GPU加速,性能远优于手动构造Seqs向量的方式,完全适配大样本量场景- 目标标签数组必须是无符号整数类型(如
u32),符合gather的索引参数要求
内容的提问来源于stack exchange,提问作者user20407415
相关产品推荐
相关产品推荐

