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

使用TensorFlow子类API训练Fashion MNIST时遇形状不匹配ValueError

TensorFlow子类API处理Fashion MNIST时的形状不兼容错误

问题描述

使用TensorFlow子类API处理Fashion MNIST批量数据时,执行model.fit()触发形状不兼容错误,报错信息如下:

ValueError: Input 0 of layer "dense" is incompatible with the layer: expected axis -1 of input shape to have value 28, but received input with shape (None, 784)

完整代码如下:

数据加载代码

import tensorflow as tf
from tensorflow import keras

import os
import sys
assert sys.version_info >= (3, 5)
import sklearn
assert sklearn.__version__ >= "0.20"

import numpy as np
import matplotlib

fashion_mnist = keras.datasets.fashion_mnist
(X_train_full, y_train_full), (X_test, y_test) = fashion_mnist.load_data()

X_valid, X_train = X_train_full[:5000] / 255., X_train_full[5000:] / 255.
y_valid, y_train = y_train_full[:5000], y_train_full[5000:]
X_test = X_test / 255.

模型定义代码

class Modeling(keras.Model):
  def __init__(self, **kwargs):
      super().__init__(**kwargs)
      self.Flatten = keras.layers.Flatten(input_shape=[28,28])
      self.dense1 = keras.layers.Dense(30, activation="relu")
      self.out = keras.layers.Dense(10, activation="softmax")

  def call(self, inputs):
    Z = self.dense1(self.Flatten(inputs))
    return self.out(Z)

模型初始化与训练代码

model = Modeling()
model.build([28,28])
model.summary()

model.compile(optimizer="sgd", loss="sparse_categorical_crossentropy", metrics=['accuracy'])
model.fit(X_train, y_train, epochs=1, validation_data=(X_valid, y_valid))

问题根源

错误出在model.build([28,28])这一行:

  • 手动指定的输入形状未包含批量维度,导致Dense层初始化时错误认为输入最后一维是28
  • 实际训练时输入为批量数据,形状是(None,28,28),经过Flatten层后变为(None,784),与Dense层期望的输入形状冲突

解决方案

有两种简单修复方式:

方式一:移除手动build调用

TensorFlow会在第一次调用模型(如model.fit())时自动推断输入形状并完成构建,无需手动调用build():

model = Modeling()
# 移除model.build([28,28])
# 若要查看summary,可先传入一个样本让模型推断形状:model(X_train[:1])
model.summary()
model.compile(optimizer="sgd", loss="sparse_categorical_crossentropy", metrics=['accuracy'])
model.fit(X_train, y_train, epochs=1, validation_data=(X_valid, y_valid))

方式二:正确指定build的输入形状

如果需要手动build,需包含批量维度(用None表示任意批量大小):

model = Modeling()
model.build((None,28,28))  # 改为包含批量维度的形状
model.summary()
model.compile(optimizer="sgd", loss="sparse_categorical_crossentropy", metrics=['accuracy'])
model.fit(X_train, y_train, epochs=1, validation_data=(X_valid, y_valid))

额外优化:移除Flatten层的input_shape参数

子类API中无需给Flatten层指定input_shape,模型会自动处理输入形状,避免不必要的形状绑定:

class Modeling(keras.Model):
  def __init__(self, **kwargs):
      super().__init__(**kwargs)
      self.Flatten = keras.layers.Flatten()  # 移除input_shape参数
      self.dense1 = keras.layers.Dense(30, activation="relu")
      self.out = keras.layers.Dense(10, activation="softmax")

  def call(self, inputs):
    Z = self.dense1(self.Flatten(inputs))
    return self.out(Z)

内容的提问来源于stack exchange,提问作者Sylvius Meynert

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.26 08:03:57