使用tensor_scatter_nd_update更新批量3阶张量时遇形状不匹配错误
解决tf.tensor_scatter_nd_update更新3阶张量的形状不匹配问题
你之前用tensor_scatter_nd_update更新2阶张量没问题,是因为indices的每个元素对应2阶张量的行索引(长度1,匹配张量前1个维度),updates的每个元素对应行的完整内容(形状(2,),匹配张量最后一个维度)。
但到3阶张量(4,3,2)时,问题出在indices的格式不对:你需要给每个更新位置指定完整的坐标,也就是包含batch索引和行索引(因为要替换的是每个batch里的整行),而不是每个batch单独给行索引。
你的错误写法中,indices是(4,2,1)的形状,updates是(4,2,2),但tensor_scatter_nd_update要求:
- indices的最后一维长度等于目标张量的阶数减去要替换的维度数(这里要替换整行,目标张量最后一维是2,所以坐标需要前2个维度:batch和行,即indices最后一维长度为2)
- updates的形状应该是
(N, 2),其中N是所有要更新的位置总数(这里4个batch×2个位置=8个)
正确实现代码
import tensorflow as tf # 初始化3阶目标张量 tensor = tf.zeros((4, 3, 2)) # 每个batch需要更新的行索引 row_indices = [[0, 2], [1, 0], [0, 1], [0, 2]] # 生成batch索引:每个batch对应2个更新位置,所以重复每个batch号2次 batch_indices = tf.repeat(tf.range(4), repeats=2) # 拼接batch索引和行索引,得到形状为(8, 2)的完整坐标 indices = tf.stack([batch_indices, tf.reshape(row_indices, [-1])], axis=1) # 将updates展平为(8, 2),对应每个更新位置的内容 updates = tf.reshape( [[[5, 5], [10, 10]], [[1, 1], [7, 7]], [[3, 3], [2, 2]], [[5, 5], [1, 1]]], [-1, 2] ) # 执行更新 output = tf.tensor_scatter_nd_update(tensor, indices, updates) print(output)
运行结果说明
输出的每个batch都会按要求更新指定行:
- 第0个batch:第0行是
[5,5],第2行是[10,10] - 第1个batch:第1行是
[1,1],第0行是[7,7] - 第2个batch:第0行是
[3,3],第1行是[2,2] - 第3个batch:第0行是
[5,5],第2行是[1,1]
内容的提问来源于stack exchange,提问作者Le_Coeur
相关产品推荐
相关产品推荐

