使用TF-slim构建TensorFlow数据集报错及优化方案咨询
报错原因
你遇到的这个错误本质上是API体系不兼容导致的:slim.dataset_data_provider.DatasetDataProvider是TF-Slim框架专门为它自己定义的slim.data.Dataset类设计的工具,但你用tf.data.Dataset.from_tensor_slices创建的是TensorFlow原生的TensorSliceDataset对象——这俩完全不是一个体系的东西,原生tf.data.Dataset根本没有data_sources这个属性,自然会抛出AttributeError。
简单说就是你把TF-Slim的老API和TensorFlow原生的现代数据管道API混在一起用了,肯定会出问题。
修复方案
给你两种思路,选一种符合你需求的就行:
思路1:用原生tf.data API直接处理批次(推荐)
既然已经用了原生的tf.data.Dataset,完全没必要再绕回TF-Slim的工具,直接用原生API的批处理方法更简单高效,这也是现在TensorFlow推荐的做法:
# 假设你已经创建了dataset = tf.data.Dataset.from_tensor_slices({"image": inputs,"label": labels}) # 直接链式调用API完成批处理、打乱、预取等操作 dataset = dataset.shuffle(buffer_size=1000) # 可选:打乱数据,buffer_size设为样本量的1/10左右即可 dataset = dataset.batch(BATCH_SIZE, drop_remainder=False) # drop_remainder=False对应你之前的allow_smaller_final_batch=True dataset = dataset.prefetch(tf.data.AUTOTUNE) # 预取数据,提升训练速度 # 之后直接迭代数据集就行 for batch in dataset: images = batch["image"] labels = batch["label"] # 你的训练逻辑写在这里
思路2:改用TF-Slim的完整数据管道(不推荐,TF-Slim维护已减少)
如果你一定要用TF-Slim的DatasetDataProvider,那必须按照TF-Slim的规范重新构建它要求的Dataset对象,而不是用原生tf.data.Dataset。大致步骤如下:
import tensorflow as tf from tensorflow.contrib.slim.python.slim.data import dataset, dataset_data_provider from tensorflow.contrib.slim.python.slim.data import tfexample_decoder # 1. 自定义解码器:实现从txt文件读取并解析图像、标签的逻辑 def build_custom_decoder(height, width, channels): keys_to_features = { # 假设你的txt每行是扁平化的图像数据+标签,需要对应定义特征 'image_raw': tf.FixedLenFeature([height*width*channels], tf.float32), 'label': tf.FixedLenFeature([], tf.float32), } items_to_handlers = { 'image': tfexample_decoder.Tensor('image_raw'), 'label': tfexample_decoder.Tensor('label'), } return tfexample_decoder.TFExampleDecoder(keys_to_features, items_to_handlers) # 2. 构建TF-Slim的Dataset对象 slim_dataset = dataset.Dataset( data_sources="path/to/your/*.txt", # 你的txt文件路径 decoder=build_custom_decoder(height, width, channels), reader=tf.TextLineReader, num_samples=N_matrices, # 总样本数 items_to_descriptions={'image': 'Input image', 'label': 'Corresponding label'} ) # 3. 使用DatasetDataProvider读取数据 data_provider = dataset_data_provider.DatasetDataProvider(slim_dataset, shuffle=True) image, label = data_provider.get(['image', 'label']) # 调整图像形状 image = tf.reshape(image, [height, width, channels]) # 构建批次 images, labels = tf.train.batch([image, label], batch_size=BATCH_SIZE, allow_smaller_final_batch=True)
这种方式需要写很多底层解析逻辑,而且TF-Slim现在已经不是TensorFlow的主流维护方向,所以更推荐第一种思路。
从.txt读取数据构建TF数据集的更优方案
你当前的做法是先把所有数据读到numpy数组再转成TF数据集,当数据量很大时(比如超过内存),这种方法会直接爆内存。更优的方案是用tf.data API流式读取txt文件,不需要一次性加载所有数据到内存:
场景1:每个txt文件每行对应一个样本(图像+标签)
import tensorflow as tf import glob # 获取所有txt文件路径 txt_files = glob.glob("path/to/your/*.txt") # 定义解析函数:把每行文本转成图像张量和标签 def parse_line(line, height, width, channels): # 分割每行的字符串 parts = tf.strings.split(line, sep=",") # 提取图像数据(前height*width*channels个元素)和标签(最后一个元素) image_flat = tf.strings.to_number(parts[:-1], tf.float32) label = tf.strings.to_number(parts[-1], tf.float32) # 把扁平化的图像转成目标形状 image = tf.reshape(image_flat, [height, width, channels]) return {"image": image, "label": label} # 构建流式数据集 dataset = tf.data.TextLineDataset(txt_files) # 读取所有txt的每行数据 # 并行解析数据,提升速度 dataset = dataset.map( lambda x: parse_line(x, height=LENGTH_INPUT, width=LENGTH_INPUT, channels=N_channels), num_parallel_calls=tf.data.AUTOTUNE ) dataset = dataset.shuffle(buffer_size=1000) # 打乱数据 dataset = dataset.batch(BATCH_SIZE, drop_remainder=False) # 批处理 dataset = dataset.prefetch(tf.data.AUTOTUNE) # 预取数据,加速训练 # 迭代使用数据集 for batch in dataset: images = batch["image"] labels = batch["label"] # 训练逻辑
场景2:每个txt文件对应一个完整的图像矩阵
如果你的每个txt文件存储的是单独的图像矩阵,调整解析逻辑即可:
def parse_txt_file(file_path, height, width, channels): # 读取整个txt文件内容 content = tf.io.read_file(file_path) # 分割成每行并过滤空行 lines = tf.strings.split(content, sep="\n") lines = tf.boolean_mask(lines, tf.strings.length(lines) > 0) # 把每行转成数值并扁平化 image_flat = tf.strings.to_number(tf.strings.split(lines, sep=","), tf.float32) # 转成图像形状 image = tf.reshape(image_flat, [height, width, channels]) # 从文件名提取标签(示例:文件名格式为"image_0.txt",标签是0) label = tf.strings.to_number(tf.strings.split(tf.strings.split(file_path, sep="_")[-1], sep=".")[0], tf.int32) return {"image": image, "label": label} # 构建数据集 dataset = tf.data.Dataset.from_tensor_slices(txt_files) dataset = dataset.map( lambda x: parse_txt_file(x, height=LENGTH_INPUT, width=LENGTH_INPUT, channels=N_channels), num_parallel_calls=tf.data.AUTOTUNE ) # 后续的shuffle、batch、prefetch和之前一致
这种流式方案的优势:
- 不占用大量内存,适合大数据量场景
- 支持并行解析和预取,大幅提升数据处理速度
- 完全使用TensorFlow原生API,兼容性和可维护性更好
内容的提问来源于stack exchange,提问作者Hakbijl

