使用随机索引更新张量值触发TensorScatterUpdate报错如何解决
错误触发原因
这个错误来自tf.tensor_scatter_nd_update接口的入参格式要求不匹配:
- 你当前生成的
indices是形状为(6,)的一维张量,存储的是6个标量形式的索引值 - 但该接口要求,不管待更新的张量是多少维,
indices的最后一维必须等于待更新张量的维度数。你当前待更新的tensor_testing是1维张量,所以indices的形状应该为(N, 1),N是要更新的位置数量,每个元素是长度为1的数组,对应1维张量的位置坐标。
参数形状不匹配就会触发你看到的维度校验报错。
修复方案
只需要给indices额外添加一个维度即可,两种实现方式可选:
- 生成indices的时候直接指定正确形状
把生成indices的代码修改为:
indices = tf.random.uniform(shape=[size_for_layer_submodel[index], 1], minval=0, maxval=size_for_layer[index], dtype=tf.dtypes.int64, seed=seed, name=None)
- 生成后扩展维度
在调用tf.tensor_scatter_nd_update前添加代码:
indices = tf.expand_dims(indices, axis=-1)
修改后示例代码可正常运行,最终会得到一个形状为(32,)的张量,其中随机6个位置的值为15.4,其余位置为0。
内容的提问来源于stack exchange,提问作者Fanto
相关产品推荐
相关产品推荐

