You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.27 00:24:04