如何在tf.data中对CSV格式输入与标签波高数据执行同步数据增强
问题解决方法
错误根因
tf.data.Dataset.map默认运行在TensorFlow图执行模式下,传入自定义函数的路径参数是张量类型,而pandas.read_csv仅支持字符串类型的本地路径,不识别张量格式的路径输入,因此触发类型错误。
可行实现方案
方案1:纯TensorFlow原生加载(推荐,性能更高,无图模式兼容问题)
完全使用TensorFlow内置接口读取CSV,避免跨框架调用的兼容性问题,适合大数据量流式加载:
import tensorflow as tf # 固定数据集形状 SHAPE = (160, 160) def load_csv_tensor(file_path): # 读取文件原始字节并解码为字符串 content = tf.io.read_file(file_path) # 按换行符拆分每一行,过滤末尾空行 lines = tf.strings.split(content, '\n') lines = tf.boolean_mask(lines, tf.strings.length(lines) > 0) # 按逗号拆分每个单元格,统一转换为float32格式 rows = tf.map_fn( lambda x: tf.strings.to_number(tf.strings.split(x, ','), out_type=tf.float32), lines, fn_output_signature=tf.TensorSpec(shape=(SHAPE[1],), dtype=tf.float32) ) # 转为指定形状的张量 return tf.reshape(rows, SHAPE) def add_offset(img_pair): offset = tf.random.uniform([1], 0, 5) * tf.ones((2, *SHAPE)) return img_pair + offset def augment(input_path, label_path): input_img = load_csv_tensor(input_path) label_img = load_csv_tensor(label_path) # 堆叠输入和标签保证增强逻辑完全同步 img_pair = tf.stack([input_img, label_img]) # 按概率触发增强 if tf.random.uniform(()) < 0.1: img_pair = add_offset(img_pair) return img_pair[0], img_pair[1] # 构造数据集 input_files = ['./Data/input_{}.csv'.format(i) for i in range(1, 200)] label_files = ['./Data/label_{}.csv'.format(i) for i in range(1, 200)] train_data = tf.data.Dataset.from_tensor_slices((input_files, label_files)) # 开启并行加载提升吞吐量 train_data = train_data.map(augment, num_parallel_calls=tf.data.AUTOTUNE).batch(1)
方案2:封装pandas加载逻辑(适合需要复杂CSV预处理的场景)
如果一定要用pandas处理CSV读取逻辑,可以用tf.py_function把Python逻辑包装成TensorFlow图可识别的操作,注意需要手动指定输出类型和形状:
import pandas as pd import tensorflow as tf SHAPE = (160, 160) def load_with_pandas(file_path): # 把张量路径转为numpy字符串 path = file_path.numpy().decode('utf-8') img = pd.read_csv(path, header=None).values return tf.convert_to_tensor(img, dtype=tf.float32) def augment(input_path, label_path): # 用py_function包装Python原生逻辑 input_img = tf.py_function(load_with_pandas, inp=[input_path], Tout=tf.float32) label_img = tf.py_function(load_with_pandas, inp=[label_path], Tout=tf.float32) # 固定形状避免后续图编译报错 input_img.set_shape(SHAPE) label_img.set_shape(SHAPE) img_pair = tf.stack([input_img, label_img]) if tf.random.uniform(()) < 0.1: offset = tf.random.uniform([1], 0, 5) * tf.ones((2, *SHAPE)) img_pair = img_pair + offset return img_pair[0], img_pair[1] # 数据集构造逻辑和方案1一致 input_files = ['./Data/input_{}.csv'.format(i) for i in range(1, 200)] label_files = ['./Data/label_{}.csv'.format(i) for i in range(1, 200)] train_data = tf.data.Dataset.from_tensor_slices((input_files, label_files)) train_data = train_data.map(augment, num_parallel_calls=tf.data.AUTOTUNE).batch(1)
注意事项
- 如果后续要添加翻转、裁剪等其他增强操作,都可以直接对堆叠后的
img_pair执行,保证输入和标签的增强参数完全一致 - 数据量较大时建议优先使用方案1,原生TensorFlow接口的并行加载和调度效率远高于跨框架的py_function调用
内容的提问来源于stack exchange,提问作者Jannik Kühn
相关产品推荐
相关产品推荐

