TensorFlow Keras图像分类模型训练第1轮Epoch停滞故障求助
图像分类模型训练停滞&OUT_OF_RANGE报错解决
问题背景
Python 3.9.12环境下,用TensorFlow Keras构建图像分类模型,训练第1轮Epoch时停滞,抛出错误:
Local rendezvous is aborting with status: OUT_OF_RANGE: End of sequence
数据集包含约8000张图片、70个类别,代码如下:
import numpy as np import pandas as pd import matplotlib.pyplot as plt import tensorflow as tf from tensorflow import keras from tensorflow.keras import layers from tensorflow.keras.models import Sequential data_training_path = '/training_path' data_testing_path = '/testing_path' data_validation_path = '/validation_path' image_width = 180 image_height = 180 batch_size = 16 data_train = tf.keras.utils.image_dataset_from_directory( data_training_path, shuffle=True, image_size=(image_width, image_height), batch_size=batch_size, validation_split=None ) data_test = tf.keras.utils.image_dataset_from_directory( data_testing_path, shuffle=False, image_size=(image_width, image_height), batch_size=batch_size, validation_split=None ) data_val = tf.keras.utils.image_dataset_from_directory( data_validation_path, shuffle=False, image_size=(image_width, image_height), batch_size=batch_size, validation_split=None ) data_classes = data_train.class_names print("Training data classes:", data_classes) plt.figure(figsize=(10, 10)) for images, labels in data_train.take(1): print("Batch of images shape:", images.shape) print("Batch of labels shape:", labels.shape) for i in range(min(9, len(images))): plt.subplot(3, 3, i+1) plt.imshow(images[i].numpy().astype('uint8')) plt.title(data_classes[labels[i]]) plt.axis('off') plt.show() model = Sequential([ layers.Rescaling(1./255), layers.Conv2D(16, 3, padding='same', activation='relu'), layers.MaxPooling2D(), layers.Conv2D(32, 3, padding='same', activation='relu'), layers.MaxPooling2D(), layers.Conv2D(64, 3, padding='same', activation='relu'), layers.MaxPooling2D(), layers.Flatten(), layers.Dropout(0.2), layers.Dense(128, activation='relu'), layers.Dense(len(data_classes), activation='softmax') ]) model.compile(optimizer='adam', loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=False), metrics=['accuracy']) epochs_size = 25 print("Number of batches in training data:", tf.data.experimental.cardinality(data_train).numpy()) print("Number of batches in validation data:", tf.data.experimental.cardinality(data_val).numpy()) history = model.fit( data_train, validation_data=data_val, epochs=epochs_size, verbose=1)
报错原因
OUT_OF_RANGE本质是数据迭代器提前耗尽,核心触发点有三个:
- 训练数据集迭代器被提前消耗:代码中用
data_train.take(1)可视化样本,TensorFlow旧版本中tf.data.Dataset是一次性迭代器,这会消耗掉第一个批次,后续model.fit时迭代器无法完成完整一轮遍历。 - 数据集路径/结构异常:部分类别文件夹为空、路径下存在非图片文件(如
.DS_Store),导致生成的数据集样本数不足,迭代中途无数据。 - 验证数据集为空:
data_validation_path路径错误或无有效图片,验证环节直接触发序列结束报错。
解决方案
1. 修复迭代器被消耗的问题
方案A:重新加载训练数据集
在可视化代码执行后,重新加载训练集,确保迭代器是全新的:
# 可视化完成后重新加载训练集 data_train = tf.keras.utils.image_dataset_from_directory( data_training_path, shuffle=True, image_size=(image_width, image_height), batch_size=batch_size, validation_split=None )
方案B:将数据集转为可重复迭代类型
给数据集添加缓存和重复属性,支持多次迭代:
# 加载训练集时添加cache和repeat data_train = tf.keras.utils.image_dataset_from_directory( data_training_path, shuffle=True, image_size=(image_width, image_height), batch_size=batch_size, validation_split=None ).cache().repeat() # 训练时需要指定steps_per_epoch,避免无限迭代 steps_per_epoch = tf.data.experimental.cardinality(data_train).numpy() history = model.fit( data_train, validation_data=data_val, epochs=epochs_size, steps_per_epoch=steps_per_epoch, verbose=1 )
2. 检查并修复数据集结构
- 确认数据集路径下的结构符合要求:每个类别对应一个独立子文件夹,子文件夹下存放对应类别的图片。
- 删除路径下的系统隐藏文件(如
.DS_Store、Thumbs.db),避免干扰数据集读取。 - 检查每个类别文件夹下至少存在一张图片,删除空的类别文件夹。
3. 验证数据集完整性
在训练前手动检查验证集的样本数量,确认其不为空:
import os def count_total_samples(data_path): total = 0 for class_dir in os.listdir(data_path): class_full_path = os.path.join(data_path, class_dir) if os.path.isdir(class_full_path): total += len([f for f in os.listdir(class_full_path) if f.lower().endswith(('.png', '.jpg', '.jpeg'))]) return total print("训练集样本数:", count_total_samples(data_training_path)) print("验证集样本数:", count_total_samples(data_validation_path))
如果验证集样本数为0,需修正路径或补充数据。
4. 升级TensorFlow版本(可选)
旧版本TensorFlow的image_dataset_from_directory存在迭代器相关bug,升级到2.10及以上版本可避免部分问题:
pip install --upgrade tensorflow
内容的提问来源于stack exchange,提问作者akam
相关产品推荐
相关产品推荐

