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

使用tf.data加载水果数据集训练无效果,请求排查问题

问题:tf.data加载水果数据集训练无提升排查

数据集信息

  • 存储结构:datasets/fruits/{class_name}/*.jpg
  • 包含类别:strawberry、mango、broccoli、grape、apple、lemon、orange

问题描述

使用tf.data.Dataset加载数据集训练时,验证集损失(val_loss)和准确率(val_accuracy)完全没有提升;但用普通Python方法加载同一模型训练,能得到合理的训练结果。怀疑问题出在图像加载处理环节。

原始代码

import tensorflow as tf
import random
from tensorflow.data import AUTOTUNE
from tensforflow.keras.optimizers import SGD
from tensorflow.keras.utils import to_categorical
from tensorflow.keras.layers.experimental import preprocessing
from imutils.import paths
from sklearn.preprocessing import LabelEncoder

INIT_LR = 1e-2 # learning rate
BS = 32 # batch size
EPOCHS = 50 # number of epochs

# load images with tensorflow
def load_images(imagePath, label):
    image = tf.io.read_file(imagePath)
    image = tf.image.decode_jpeg(image, channels=3)
    image = tf.image.convert_image_dtype(image, dtype=tf.float32)
    image = tf.image.resize(image, (64, 64))
    return (image, label)

# augment helper function
def augment(image, label, aug):
    image = aug(image)
    return (image, label)

# get all image paths and save them as strings with format **/{class_name}/*jpg
allImages = list(paths.list_images("datasets/fruits"))
random.shuffle(allImages) # shuffle the images

# perform 0.75/0.25 train/test split
i = int(len(allImages) * 0.25)
trainPaths = allImages[i:]

# get labels by getting {class_name} from **/{class_name}/*jpg
trainLabels = [p.split(os.path.sep)[-2] for p in trainPaths]
testPaths = allImages[:i]
testLabels = [p.split(os.path.sep)[-2] for p in testPaths]

# use LabelEncoder to one-hot encode the class names
labelEncoder = LabelEncoder()
labelEncoder = labelEncoder.fit(trainLabels)
trainLabels = labelEncoder.transform(trainLabels)
trainLabels = to_categorical(trainLabels)
testLabels = labelEncoder.transform(testLabels)
testLabels = to_categorical(testLabels)

# load the train and test data into a tf.data.Dataset
trainDS = tf.data.Dataset.from_tensor_slices((trainPaths, trainLabels))
trainDS = (
    trainDS
    .shuffle(32, seed=42)
    .map(load_images, num_parallel_calls=AUTOTUNE)
    .batch(BS)
    .cache()
)

# rescale the pixels from [0, 1]
trainAug = tf.keras.Sequential(
    [
        preprocessing.Rescaling(scale=1.0/255),
    ]
)

trainDS = (
    trainDS
    .map(lambda x, y: augment(x, y, trainAug), num_parallel_calls=AUTOTUNE)
    .prefetch(AUTOTUNE)
)

testDS = tf.data.Dataset.from_tensor_slices(( testPaths, testLabels ))
testDS = (
    testDS
    .shuffle(32)
    .map(load_images, num_parallel_calls=AUTOTUNE)
    .batch(BS)
    .cache()
)

testAug = tf.keras.Sequential(
    [
        preprocessing.Rescaling(scale=1.0/255),
    ]
)

testDS = (
    testDS
    .map(lambda x, y: augment(x, y, testAug), num_parallel_calls=AUTOTUNE)
    .prefetch(AUTOTUNE)
)

# I don't think there is any issues with this part, but here I am setting up the optimizer and model for training
sgd = SGD(learning_rate=INIT_LR, momentum=0.9, weight_decay=INIT_LR/EPOCHS)
model = MiniVGGNet.build(64, 64, 3, num_classes=num_classes)
model.compile(loss="categorical_crossentropy", optimizer=sgd, metrics=["accuracy"])
training_history = model.fit(
    x=trainDS,
    validation_data=testDS,
    epochs=EPOCHS
)

核心问题排查与修正

1. 重复像素缩放导致数据异常

load_images函数中tf.image.convert_image_dtype(image, dtype=tf.float32)会自动将像素值从[0,255]缩放到[0,1],后续又通过Rescaling(scale=1.0/255)再次缩放,最终像素值范围变成[0, 1/255],模型无法有效学习特征。

2. 代码笔误

  • 导入错误:from tensforflow.keras.optimizers import SGD → 修正为tensorflow
  • 导入错误:from imutils.import paths → 修正为from imutils import paths
  • 缺少import os,代码中用到了os.path.sep
  • 未定义num_classes,需设置为7

3. 数据处理流程顺序错误

  • cache()应放在数据增强之后,否则缓存的是未正确处理的数据
  • 测试集不需要shuffle,测试阶段仅需一次有序遍历评估

4. Shuffle缓冲区过小

shuffle(32)的缓冲区远小于训练集规模,无法充分打乱数据,影响模型泛化能力,建议设置为训练集样本数的1/10或至少1000。

修正后代码示例

import tensorflow as tf
import random
import os
from tensorflow.data import AUTOTUNE
from tensorflow.keras.optimizers import SGD
from tensorflow.keras.utils import to_categorical
from tensorflow.keras.layers.experimental import preprocessing
from imutils import paths
from sklearn.preprocessing import LabelEncoder

INIT_LR = 1e-2 # learning rate
BS = 32 # batch size
EPOCHS = 50 # number of epochs
NUM_CLASSES = 7 # 明确类别数量

def load_images(imagePath, label):
    image = tf.io.read_file(imagePath)
    image = tf.image.decode_jpeg(image, channels=3)
    image = tf.image.resize(image, (64, 64))
    # 仅保留一次像素缩放,转成float32范围[0,1]
    image = tf.image.convert_image_dtype(image, dtype=tf.float32)
    return (image, label)

def augment(image, label, aug):
    image = aug(image)
    return (image, label)

# 加载并打乱图像路径
allImages = list(paths.list_images("datasets/fruits"))
random.shuffle(allImages)

# 划分训练/测试集
i = int(len(allImages) * 0.25)
trainPaths = allImages[i:]
trainLabels = [p.split(os.path.sep)[-2] for p in trainPaths]
testPaths = allImages[:i]
testLabels = [p.split(os.path.sep)[-2] for p in testPaths]

# 标签编码
labelEncoder = LabelEncoder()
labelEncoder.fit(trainLabels)
trainLabels = to_categorical(labelEncoder.transform(trainLabels))
testLabels = to_categorical(labelEncoder.transform(testLabels))

# 构建训练数据集:调整流程顺序,增大shuffle缓冲区
trainDS = tf.data.Dataset.from_tensor_slices((trainPaths, trainLabels))
trainDS = (
    trainDS
    .shuffle(len(trainPaths), seed=42)
    .map(load_images, num_parallel_calls=AUTOTUNE)
    .batch(BS)
    .map(lambda x, y: augment(x, y, preprocessing.Rescaling(scale=1.0/255)), num_parallel_calls=AUTOTUNE)
    .cache()
    .prefetch(AUTOTUNE)
)

# 构建测试数据集:移除shuffle,简化流程
testDS = tf.data.Dataset.from_tensor_slices((testPaths, testLabels))
testDS = (
    testDS
    .map(load_images, num_parallel_calls=AUTOTUNE)
    .batch(BS)
    .map(lambda x, y: augment(x, y, preprocessing.Rescaling(scale=1.0/255)), num_parallel_calls=AUTOTUNE)
    .cache()
    .prefetch(AUTOTUNE)
)

# 模型训练
sgd = SGD(learning_rate=INIT_LR, momentum=0.9, weight_decay=INIT_LR/EPOCHS)
model = MiniVGGNet.build(64, 64, 3, num_classes=NUM_CLASSES)
model.compile(loss="categorical_crossentropy", optimizer=sgd, metrics=["accuracy"])
training_history = model.fit(
    x=trainDS,
    validation_data=testDS,
    epochs=EPOCHS
)

额外验证建议

  • 抽取一个batch的训练数据,打印像素值范围,确认是[0,1]
  • 打印标签的one-hot编码,验证与类别对应关系正确
  • 对比两种加载方式下的输入数据维度、像素范围是否完全一致

内容的提问来源于stack exchange,提问作者Kuan Chen

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 17:14:54