TensorFlow训练木薯数据集模型时输入名称不匹配报错求助
解决TensorFlow木薯数据集训练输入不匹配问题
报错原因
tfds加载的Cassava数据集是字典结构(包含image、label等键),但你的Sequential模型默认期望接收单张量输入,导致输入名称不匹配,出现Missing data for input "flatten_input"错误。
关键修复步骤
1. 修正预处理函数,返回(特征, 标签)元组
把原本返回字典的预处理逻辑改成直接返回图像张量和标签,让模型能正确识别输入:
def preprocess_fn(data): image = data['image'] # 归一化[0,255]到[0,1] image = tf.cast(image, tf.float32) image = image / 255. # 调整图像尺寸到224x224 image = tf.image.resize(image, (224, 224)) # 返回(图像, 标签)元组,而非字典 return image, data['label']
2. 修正模型输入形状的笔误
你在预处理里把图像resize到224x224,但模型输入形状写的是(244,244,3),这会导致尺寸不匹配,修正为:
def create_model(type="default", n_classes=6): if type == "something": pass else: model = K.Sequential() # 修正输入形状为(224,224,3) model.add(K.layers.Flatten(input_shape=(224, 224, 3))) model.add(K.layers.Dense(512, activation="relu")) model.add(K.layers.Dense(256, activation="relu")) model.add(K.layers.Dense(128, activation="relu")) model.add(K.layers.Dense(64, activation="relu")) model.add(K.layers.Dense(n_classes, activation="softmax")) model.compile(loss='sparse_categorical_crossentropy', optimizer=K.optimizers.Adam(0.01), metrics=['accuracy']) return model
3. 对训练集应用预处理、分批与打乱
在调用fit前,需要对数据集做预处理、分批和打乱操作,保证训练效率和数据随机性:
# 处理训练集:预处理+打乱+分批 train_dataset = dataset["train"].map(preprocess_fn).shuffle(1000).batch(32) model = create_model() model.fit(train_dataset, epochs=5)
完整修复后的代码
# tensorflow 2.x core api import logging from mlflow.models import infer_signature import numpy as np import tensorflow as tf import tensorflow_datasets as tfds from tensorflow import keras as K logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) ############################################################################################################# from matplotlib import pyplot as plt def plot(examples, predictions=None): # Get the images, labels, and optionally predictions images = examples['image'] labels = examples['label'] batch_size = len(images) if predictions is None: predictions = batch_size * [None] # Configure the layout of the grid x = np.ceil(np.sqrt(batch_size)) y = np.ceil(batch_size / x) fig = plt.figure(figsize=(x * 6, y * 7)) for i, (image, label, prediction) in enumerate(zip(images, labels, predictions)): # Render the image ax = fig.add_subplot(int(x), int(y), i+1) ax.imshow(image, aspect='auto') ax.grid(False) ax.set_xticks([]) ax.set_yticks([]) # Display the label and optionally prediction x_label = 'Label: ' + name_map[class_names[label]] if prediction is not None: x_label = 'Prediction: ' + name_map[class_names[prediction]] + '\n' + x_label ax.xaxis.label.set_color('green' if label == prediction else 'red') ax.set_xlabel(x_label) plt.show() dataset, info = tfds.load("cassava", shuffle_files=True, with_info=True) print("INFO:\n", info) # Extend the cassava dataset classes with 'unknown' class_names = info.features['label'].names + ['unknown'] # Map the class names to human readable names name_map = dict( cmd='Mosaic Disease', cbb='Bacterial Blight', cgm='Green Mite', cbsd='Brown Streak Disease', healthy='Healthy', unknown='Unknown') print(len(class_names), 'classes:') print(class_names) print([name_map[name] for name in class_names]) def preprocess_fn(data): image = data['image'] # Normalize [0, 255] to [0, 1] image = tf.cast(image, tf.float32) image = image / 255. # Resize the images to 224 x 224 image = tf.image.resize(image, (224, 224)) # 返回(图像, 标签)元组 return image, data['label'] def create_model(type="default", n_classes=6): if type == "something": pass else: model = K.Sequential() # 修正输入形状 model.add(K.layers.Flatten(input_shape=(224, 224, 3))) model.add(K.layers.Dense(512, activation="relu")) model.add(K.layers.Dense(256, activation="relu")) model.add(K.layers.Dense(128, activation="relu")) model.add(K.layers.Dense(64, activation="relu")) model.add(K.layers.Dense(n_classes, activation="softmax")) model.compile(loss='sparse_categorical_crossentropy', optimizer=K.optimizers.Adam(0.01), metrics=['accuracy']) return model # 可选:验证可视化代码 # batch = dataset['validation'].map(lambda x: x).batch(25).as_numpy_iterator() # examples = next(batch) # plot(examples) print(tf.__version__) model = create_model() # 处理训练集 train_dataset = dataset["train"].map(preprocess_fn).shuffle(1000).batch(32) model.fit(train_dataset, epochs=5)
内容的提问来源于stack exchange,提问作者user3002166
相关产品推荐
相关产品推荐

