如何用tf.tensor_scatter_nd_update批量更新3D张量的指定通道?
问题解答
可以用tf.tensor_scatter_nd_update实现你的需求,错误核心是传入的索引缺少batch维度的定位信息,导致TensorFlow无法正确匹配待更新的位置与更新值。
错误原因拆解
tf.tensor_scatter_nd_update要求:
indices的最后一维长度,需对应你要定位的张量维度数。你要更新的是每个batch下指定patch的完整特征向量(即3D张量中[batch_idx, patch_idx, :]的切片),所以每个索引需要包含batch_idx和patch_idx两个维度信息,也就是indices最后一维长度为2。updates的形状需满足:updates.shape = indices.shape[:-1] + 待替换切片的形状。你的updateTensor是(64,4,768),正好对应(batch数, 每个batch要更新的patch数, 特征长度),是符合要求的。
而你原代码中的indices仅为(64,4),只包含了每个batch内的patch索引,缺少batch维度的标识,TensorFlow无法关联每个patch所属的batch,因此触发形状不匹配的错误。
修正后的代码
import tensorflow as tf featureSize = 768 batchSize = 64 patchCount = 8 toUpdatePatchCount = 4 inputTensor = tf.random.normal([batchSize,patchCount,featureSize]) # 生成每个batch对应的索引,形状(64,1) batch_indices = tf.range(batchSize)[:, tf.newaxis] # 拼接batch索引与patch索引,得到完整定位坐标,形状(64,4,2) full_indices = tf.concat([batch_indices, indices[..., tf.newaxis]], axis=-1) updateTensor = tf.random.normal([batchSize,toUpdatePatchCount,featureSize]) # 执行更新操作 outputTensor = tf.tensor_scatter_nd_update(inputTensor, full_indices, updateTensor) print(outputTensor.shape) # 输出: (64, 8, 768),与输入形状一致
关键说明
full_indices的形状为(64,4,2),每个元素是[batch_idx, patch_idx],明确指定了每个待更新patch在3D张量中的位置。updateTensor的(64,4,768)形状与full_indices的前两维(64,4)完全对应,每个位置的768维特征向量会精准替换输入张量中对应[batch_idx,patch_idx]位置的特征向量。
内容的提问来源于stack exchange,提问作者RabbitBadger
相关产品推荐
相关产品推荐

