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

JavaScript多维数组处理:如何对Float32Array求指定轴的max与argmax

JavaScript处理多维数组与Float32Array张量的max/argmax实现

需求说明

需要在浏览器端处理多维数组,针对shape为(1,17,512)的Float32Array格式3D张量,沿最后一个轴(轴索引2)计算每个位置的最大值及其对应索引,最终得到17组结果,等价于mmpose对应Python函数的功能。

现有尝试的问题

使用numjs库时遇到两个限制:

  • max()方法仅能计算全局最大值,不支持指定轴参数
  • get()方法只能获取单个元素,无法批量访问子数组

用户的示例代码如下:

import * as ort from "onnxruntime-web";
import * as img from "$lib/utils/imageHelper";
import nj from "@d4c/numjs/build/module/numjs.min.js";

const session = await ort.InferenceSession.create("rtmpose.onnx", {
    executionProviders: ["wasm"],
});

const data = await img.getImageTensorFromPath(
    "https://i.imgur.com/CzXTtJV.jpg", // image url,
    [1, 3, 256, 192]
);

// prepare feeds. use model input names as keys.
const feeds = { input: data };

// feed inputs and run
const results = await session.run(feeds);
const output = results[session.outputNames[0]];

// output.dims is [1,17,512]
let arr = nj.array(output.data).reshape(output.dims);

// only returns one maximum
let maxes = arr.max(2);

// returns undefined
arr.get(0,0)

解决方案

方案1:手动实现(针对固定shape的高效方案)

由于张量shape固定为(1,17,512),可以直接通过遍历Float32Array实现,无需依赖第三方库,性能最优:

const output = results[session.outputNames[0]];
const data = output.data; // Float32Array
const [batchSize, numKeypoints, numValues] = output.dims; // [1,17,512]

const maxValues = [];
const maxIndices = [];

// 遍历每个关键点(忽略batch维度,因为batchSize=1)
for (let k = 0; k < numKeypoints; k++) {
    let maxVal = -Infinity;
    let maxIdx = -1;
    const startIdx = k * numValues;
    // 遍历当前关键点的所有512个值
    for (let i = 0; i < numValues; i++) {
        const currentVal = data[startIdx + i];
        if (currentVal > maxVal) {
            maxVal = currentVal;
            maxIdx = i;
        }
    }
    maxValues.push(maxVal);
    maxIndices.push(maxIdx);
}

// maxValues是长度为17的数组,对应每个关键点的最大值
// maxIndices是长度为17的数组,对应每个最大值在最后一轴的索引

方案2:使用支持轴操作的浏览器端数值计算库

如果需要处理更灵活的张量shape,推荐使用以下工具:

  • TensorFlow.js:可以将Float32Array转换为tf.Tensor,然后使用tf.max()和tf.argMax()方法指定axis参数:
    import * as tf from '@tensorflow/tfjs';
    
    const output = results[session.outputNames[0]];
    const tensor = tf.tensor3d(output.data, output.dims);
    // 沿最后一轴计算max
    const maxTensor = tf.max(tensor, 2);
    // 沿最后一轴计算argmax
    const argmaxTensor = tf.argMax(tensor, 2);
    // 转换为普通数组
    const maxValues = await maxTensor.squeeze().array(); // squeeze去掉batch维度,得到长度17的数组
    const maxIndices = await argmaxTensor.squeeze().array();
    
  • ndarray + ndarray-ops:轻量级的多维数组库,配合ndarray-ops可以实现轴上的聚合操作:
    import ndarray from 'ndarray';
    import * as ops from 'ndarray-ops';
    
    const output = results[session.outputNames[0]];
    const arr = ndarray(output.data, output.dims);
    const maxArr = ndarray(new Float32Array(17), [1,17]);
    
    // 沿最后一轴计算max
    ops.max(maxArr, arr);
    // 如需argmax,可使用ndarray-argmax库
    

内容的提问来源于stack exchange,提问作者KarelPrdel

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.12 14:33:14