使用tf.tensor_scatter_add在张量第一维度插入切片报错求助
解决
tf.tensor_scatter_add的报错问题 咱们先一步步拆解你代码里的问题,从明显的拼写错误到核心的维度匹配问题:
1. 先修正代码里的拼写小错误
你这段代码里有两处变量名拼写不一致的问题:
- 初始化Session时用的是
se,但后面调用却写成了ses - 最后打印的变量是
scatter,但实际要运行的是定义好的updated
修正后的基础代码如下:
indices = tf.constant([[0], [2]]) updates = tf.constant([[[5, 5, 5, 5], [6, 6, 6, 6], [7, 7, 7, 7], [8, 8, 8, 8]], [[5, 5, 5, 5], [6, 6, 6, 6], [7, 7, 7, 7], [8, 8, 8, 8]]]) tensor = tf.ones([4, 5, 4]) updated = tf.tensor_scatter_add(tensor, indices, updates) with tf.Session() as ses: print(ses.run(updated))
2. 核心报错原因:维度不匹配
就算修正了拼写,这段代码还是会报错,因为**tf.tensor_scatter_add要求updates的形状必须和indices指向的原张量切片形状完全一致**。
咱们理清楚各部分的维度:
- 原
tensor的形状是[4, 5, 4]:第一维度有4个元素,每个元素是5行4列的矩阵 - 你的
indices = [[0], [2]]:表示要对原张量的第0个和第2个第一维度元素做加法更新 - 但你的
updates形状是[2, 4, 4]:每个更新切片是4行4列的矩阵,和原张量的5行4列切片形状不匹配,这就是报错的根源。
如果你是想给现有位置的矩阵做加法更新
那只需要把updates的形状改成[2,5,4],让每个更新切片和原张量的第一维度元素形状一致就行,示例代码:
indices = tf.constant([[0], [2]]) # 调整updates为每个切片是5行4列,和原张量的第一维度元素匹配 updates = tf.constant([[[5]*4]*5, [[5]*4]*5]) tensor = tf.ones([4, 5, 4]) updated = tf.tensor_scatter_add(tensor, indices, updates) with tf.Session() as ses: print(ses.run(updated))
这样运行后,原张量的第0和第2个5x4矩阵就会加上对应的更新矩阵。
如果你是想在第一维度插入新的矩阵切片(不是加法更新)
那tf.tensor_scatter_add做不到这个功能——它的作用是对现有位置的元素做加法,不是插入新元素。这时候你需要用tf.concat来拆分原张量并拼接新切片,举个例子:
假设你想在原张量的索引0和2位置插入新的4x4矩阵(注意要保证插入切片和原张量的后两个维度匹配,这里我把原张量调整为[4,4,4]来适配你的更新切片):
tensor = tf.ones([4, 4, 4]) # 要插入的两个4x4矩阵切片 insert_slices = tf.constant([[[5,5,5,5],[6,6,6,6],[7,7,7,7],[8,8,8,8]], [[5,5,5,5],[6,6,6,6],[7,7,7,7],[8,8,8,8]]]) # 拆分原张量并按位置拼接插入切片 result = tf.concat([ insert_slices[0:1], # 在开头插入第一个切片 tensor[0:2], # 原张量的前2个元素 insert_slices[1:2], # 在索引2位置插入第二个切片 tensor[2:] # 原张量剩下的元素 ], axis=0) with tf.Session() as ses: print(ses.run(result).shape) # 输出 (6,4,4),符合插入后的维度预期
内容的提问来源于stack exchange,提问作者user12486020
相关产品推荐
相关产品推荐

