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
相关产品推荐
相关产品推荐

