You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.14 10:15:34