TensorFlow中基于top_k修改张量后重构原形状张量的问题
解决方法
首先纠正你提取数据时的笔误:你代码里的fs = tf.gather(X, ti[:, 0], axis=-1)应该改为从原始data张量中提取,因为你要操作的是数据而非分数。正确的提取代码可以简化为:
# 从data中提取每个样本top95索引对应的数据,形状(None, 3969, 95) fs = tf.gather(data, Y, axis=-1) Z = fs * 0.7 # 修改后的数据
接下来核心问题是把修改后的Z放回原张量的对应位置,以下是具体实现步骤:
步骤1:构造更新所需的完整索引
要将Z中的每个值精准放回data的原始位置,需要构造包含批次索引、3969维度索引、128维度索引的三维索引组:
# 获取动态批次大小(因为batch_size是None,不能用Y.shape[0]) batch_size = tf.shape(Y)[0] seq_len = tf.shape(data)[1] # 即3969 # 生成批次维度的索引,扩展为(batch_size, 1, 1)用于广播 batch_indices = tf.range(batch_size)[:, tf.newaxis, tf.newaxis] # 生成3969维度的索引,扩展为(1, 3969, 1)用于广播 seq_indices = tf.range(seq_len)[tf.newaxis, :, tf.newaxis] # 将Y扩展为(batch_size, 1, 95),匹配前两个维度的广播要求 y_indices = Y[:, tf.newaxis, :] # 拼接成完整的更新索引,形状为(batch_size, 3969, 95, 3) # 每个元素格式为[批次号, 3969位置号, 128维度的top_k索引] update_indices = tf.concat([batch_indices, seq_indices, y_indices], axis=-1)
步骤2:展平索引和数据以适配scatter操作
tf.tensor_scatter_nd_update要求索引是二维数组(每行一个索引),值是一维数组,因此需要展平:
# 展平索引为(batch_size*3969*95, 3) update_indices_flat = tf.reshape(update_indices, [-1, 3]) # 展平Z为(batch_size*3969*95,) Z_flat = tf.reshape(Z, [-1])
步骤3:执行更新得到最终张量F
复制原始data并更新指定位置的值:
# 注意:tf.tensor_scatter_nd_update会返回新张量,不会修改原data F = tf.tensor_scatter_nd_update(data, update_indices_flat, Z_flat)
最终得到的F形状为(None, 3969, 128),其中所有top95分数对应位置的数据是修改后的Z值,其余位置保留原始data的内容,且完全保持原始顺序。
内容的提问来源于stack exchange,提问作者Surya Majumder
相关产品推荐
相关产品推荐

