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

如何解决输入与Dense层形状不兼容的ValueError问题?

问题分析与解决

错误根源

你处理测试集时犯了个拼写错误:把ds_test = ds_test.batch(128)写成了de_test = ds_test.batch(128),导致ds_test没有被批量处理,仍然是单个样本的数据集(每个样本形状为(28,28,1))。

模型训练时用的是批量输入(BATCH_SIZE=64),经过Conv2D(32,3,padding="same")后,每个样本的形状变成(28,28,32),再经过Flatten层会被展平成28*28*32=25088,所以Dense层默认期望输入最后一维是25088。但测试时输入是单个样本,经过Conv2D后是(28,28,32),Flatten后变成(28,896),和Dense层的预期形状不匹配,触发了这个ValueError。

修复方法

修正测试集处理里的变量名错误,确保ds_test被正确批量处理:

修正后的测试集处理代码

ds_test = ds_test.map(normalize_img, num_parallel_calls=AUTOTUNE)
ds_test = ds_test.batch(128)  # 把错误的de_test改成ds_test
ds_test = ds_test.prefetch(AUTOTUNE)

完整修正代码

import tensorflow as tf
from tensorflow import keras
from tensorflow.keras import layers

(ds_train, ds_test), ds_info = tfds.load(
    "mnist",
    split=["train", "test"],
    shuffle_files=True,
    as_supervised=True,
    with_info=True
)

def normalize_img(image, label):
    return tf.cast(image, tf.float32)/255.0, label

AUTOTUNE = tf.data.experimental.AUTOTUNE
BATCH_SIZE = 64

ds_train = ds_train.map(normalize_img, num_parallel_calls=AUTOTUNE)
ds_train = ds_train.cache()
ds_train = ds_train.shuffle(ds_info.splits["train"].num_examples)
ds_train = ds_train.batch(BATCH_SIZE)
ds_train = ds_train.prefetch(AUTOTUNE)

ds_test = ds_test.map(normalize_img, num_parallel_calls=AUTOTUNE)
ds_test = ds_test.batch(128)  # 修正变量名错误
ds_test = ds_test.prefetch(AUTOTUNE)

model = keras.Sequential ([
    keras.Input(shape=[28, 28, 1],),
    layers.Conv2D(32, 3, activation='relu', padding="same"),
    layers.Flatten(),
    layers.Dense(10),
])

model.compile(
    optimizer=keras.optimizers.Adam(learning_rate=0.001),
    loss = keras.losses.SparseCategoricalCrossentropy(from_logits=True),
    metrics=["accuracy"],
)

model.fit(ds_train, epochs=5, verbose=2)
model.evaluate(ds_test)

额外排查技巧

如果想验证各层输入输出形状是否正确,可以在模型中添加一个打印形状的Lambda层:

model = keras.Sequential ([
    keras.Input(shape=[28, 28, 1],),
    layers.Conv2D(32, 3, activation='relu', padding="same"),
    layers.Lambda(lambda x: print(x.shape)),  # 打印Conv2D输出形状
    layers.Flatten(),
    layers.Dense(10),
])

这样能直观看到每一层的形状变化,方便快速定位类似的形状不匹配问题。

内容的提问来源于stack exchange,提问作者Isesele Victor

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 04:13:11