如何修改tf.Dataset对象特定索引位置的元素值
问题说明
现有tf.data.Dataset对象,单个样本结构为((3, 400, 1), (3, 400, 1)),需求为将每个样本中x、y的第3行(索引为2)元素值除以10,要求优先使用map方法实现。原有实现通过as_numpy_iterator迭代后直接修改numpy数组,因numpy数组不可变抛出报错,原代码如下:
def alternate_row (dataset): xx, yy = [], [] for x, y in dataset.as_numpy_iterator(): x[2] /= 10 y[2] /= 10 xx.append(x) yy.append(y) return xx, yy
实现方案
优先方案:使用tf.data原生map方法(推荐)
全程使用TensorFlow原生算子实现,无需转numpy数组,性能更高,处理后仍返回tf.data.Dataset对象,可直接对接后续训练流水线:
import tensorflow as tf def map_func(x, y): # 提取第3行做除法运算 new_x_row = x[2] / 10.0 new_y_row = y[2] / 10.0 # 指定要更新的位置为第0维索引2的位置 update_indices = [[2]] # 替换对应行得到新张量 x_updated = tf.tensor_scatter_nd_update(x, update_indices, [new_x_row]) y_updated = tf.tensor_scatter_nd_update(y, update_indices, [new_y_row]) return x_updated, y_updated # 调用map处理数据集 processed_dataset = raw_dataset.map(map_func)
兼容原有逻辑的修正方案
如果需要保留原有返回numpy列表的逻辑,只需要在修改前对numpy数组做拷贝,避免直接修改不可变的原数组即可:
import numpy as np def alternate_row (dataset): xx, yy = [], [] for x, y in dataset.as_numpy_iterator(): # 拷贝生成可修改的新数组 x = x.copy() y = y.copy() x[2] /= 10 y[2] /= 10 xx.append(x) yy.append(y) return np.array(xx), np.array(yy)
内容的提问来源于stack exchange,提问作者noone
相关产品推荐
相关产品推荐

