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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.08 16:02:36