TensorFlow2.17运行《Python深度学习》示例10.2.3报InvalidArgumentError
TensorFlow 2.17与2.15版本差异导致《Python深度学习》示例报错的解决方案
问题背景
运行François Chollet所著《Python深度学习》示例10.2.3时,使用TensorFlow 2.17执行history = model.fit(...)会触发InvalidArgumentError,但降级到2.15版本则正常运行。报错核心信息为:
Only one input size may be -1, not both 0 and 1
错误出现在Flatten层的Reshape操作节点中。
版本变更原因
TensorFlow 2.16到2.17版本间,Keras的Flatten层及底层Reshape操作对动态形状的校验逻辑做了严格升级:
- 在2.15及更早版本中,
Flatten层可以兼容输入中存在的未知维度(比如数据集的batch维度为动态值,或sequence_length未明确为固定整数),自动推导扁平化后的形状。 - 2.17版本中,Reshape操作要求仅能有一个维度用
-1(自动计算),如果输入形状中同时存在未确定的维度(标记为0)和-1,就会触发上述参数冲突错误。示例中Input的sequence_length若为动态值,或数据集输出的batch维度未固定,就会触发这个校验。
针对TensorFlow 2.17及以上版本的修改方案
方案1:显式指定扁平化后的维度(推荐)
不再依赖Flatten层自动推导,手动计算扁平化后的特征数,用Reshape层明确指定目标形状:
from tensorflow import keras from tensorflow.keras import layers # 提前计算扁平化后的总特征数 flattened_features = sequence_length * raw_data.shape[-1] inputs = keras.Input(shape=(sequence_length, raw_data.shape[-1])) # 用Reshape替代Flatten,明确形状 x = layers.Reshape((flattened_features,))(inputs) x = layers.Dense(16, activation='relu')(x) outputs = layers.Dense(1)(x) model = keras.Model(inputs, outputs) callbacks = [ keras.callbacks.ModelCheckpoint('jena_dense.keras', save_best_only=True) ] model.compile(optimizer='rmsprop', loss='mse', metrics=['mae']) history = model.fit( train_dataset, epochs=10, validation_data=val_dataset, callbacks=callbacks) model = keras.models.load_model('jena_dense.keras') print(f"Test MAE: {model.evaluate(test_dataset)[1]:.2f}")
方案2:固定输入形状的所有维度
确保sequence_length是明确的整数类型,避免输入形状中存在未知维度:
from tensorflow import keras from tensorflow.keras import layers # 强制将sequence_length转为固定整数 sequence_length = int(sequence_length) inputs = keras.Input(shape=(sequence_length, raw_data.shape[-1])) x = layers.Flatten()(inputs) x = layers.Dense(16, activation='relu')(x) outputs = layers.Dense(1)(x) # 后续编译、训练代码保持不变
内容的提问来源于stack exchange,提问作者Mike M. Lin
相关产品推荐
相关产品推荐

