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

TensorFlow2.7环境下集成TFHub模型到Keras后batch>1训练报错如何解决

报错原因

问题出在你手动写的张量展平逻辑上:
你用tf.reshape(representation_full, (-1, representation_full.shape[-1]))做展平的时候,直接调用了张量的.shape属性拿静态维度。TensorFlow的fit方法默认运行在图模式下,模型构建阶段batch维度是动态不确定的,用.shape拿静态值会导致维度推断错误,最终让最后一层Dense层生成了不符合预期的权重:按报错信息看Dense层生成了形状为[1,10]的权重,但batch size>1时实际传入Dense层的输入形状为[2,2],矩阵乘法尺寸不匹配就报错了。
此外你给Dense层额外加的input_shape参数也是多余的,反而会干扰Keras的自动维度推断。

修复方案

直接用Keras内置的Flatten层替换你手动写的reshape逻辑即可,Flatten层原生支持动态batch size的适配,不需要手动计算维度:
将原来的代码段:

flat = tf.reshape(representation_full, (-1, representation_full.shape[-1]))

# Outputs
output = Dense(10, input_shape=(-1, flat.shape[-1]),
    activation='linear')(flat)

修改为:

flat = tf.keras.layers.Flatten()(representation_full)
output = Dense(10, activation='linear')(flat)

修改后不管batch size是多少,fit、evaluate、predict方法都能正常运行。


内容的提问来源于stack exchange,提问作者Dani Pinkovich

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.26 08:36:02