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

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本质是数据迭代器提前耗尽,核心触发点有三个:

  1. 训练数据集迭代器被提前消耗:代码中用data_train.take(1)可视化样本,TensorFlow旧版本中tf.data.Dataset是一次性迭代器,这会消耗掉第一个批次,后续model.fit时迭代器无法完成完整一轮遍历。
  2. 数据集路径/结构异常:部分类别文件夹为空、路径下存在非图片文件(如.DS_Store),导致生成的数据集样本数不足,迭代中途无数据。
  3. 验证数据集为空: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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.20 18:13:18