使用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
相关产品推荐
相关产品推荐

