tf.unsortedSegmentSum未聚合全部有效索引?TensorFlow JS问题咨询
环境信息
- Tensorflow JS:4.16.0 版本
- 后端:webgpu
代码实现
// https://js.tensorflow.org/api/4.17.0/#unsortedSegmentSum const onesTensor = tf.ones([size * size], 'float32'); const countTensor = tf.unsortedSegmentSum(onesTensor, indexTensor, size * size); console.log("onesTensor shape =", onesTensor.shape, "; dtype =", onesTensor.dtype); console.log("indexTensor shape =", indexTensor.shape, "; dtype =", indexTensor.dtype); console.log("countTensor shape =", countTensor.shape, "; dtype =", countTensor.dtype); console.log("size * size =", size * size); const indexBuffer = indexTensor.arraySync(); console.log("indexBuffer valid element number =" , indexBuffer.map((v) => (v >= 0 && v < size * size)? 1 : 0).reduce((a, b) => a + b) ); const data = countTensor.arraySync(); console.log("total =", data.reduce((a, b) => a + b));
控制台输出
onesTensor shape = [88804] ; dtype = float32 indexTensor shape = [1000000] ; dtype = int32 countTensor shape = [88804] ; dtype = float32 size * size = 88804 indexBuffer valid element number = 1000000 total = 88804
问题描述
我用上述代码调用tf.unsortedSegmentSum,基于indexTensor里的索引对全1张量做求和统计。控制台显示indexBuffer里有1000000个有效索引,但countTensor的求和结果只有88804。已经确认88804个索引的聚合逻辑是对的,为什么1000000个有效索引没有被完整聚合?
解决方案
问题根源是对tf.unsortedSegmentSum的参数逻辑理解有误:该API要求第一个参数data(也就是你的onesTensor)的长度,必须和第二个参数segmentIds(你的indexTensor)的长度完全一致。
你的代码里,onesTensor只有88804个元素,但indexTensor有1000000个索引。unsortedSegmentSum的逻辑是把data中每个元素,按照segmentIds对应位置的索引加到结果里——现在data只有88804个元素,所以只会处理indexTensor的前88804个索引,剩下的911196个索引没有对应的data元素,自然不会被统计。
要实现“统计1000000个索引的出现次数”,只需把onesTensor的长度改成和indexTensor一致即可:
const onesTensor = tf.ones([1000000], 'float32'); // 与indexTensor长度匹配 const countTensor = tf.unsortedSegmentSum(onesTensor, indexTensor, size * size);
修改后countTensor的求和结果就会等于1000000,所有有效索引都会被正确聚合。
内容的提问来源于stack exchange,提问作者Mr.Wang from Next Door
相关产品推荐
相关产品推荐

