TensorFlow自定义Vision Transformer模型predict调用失败求助
明确
predict()和直接调用call()/fit()的核心差异:predict()会通过tf.function将模型推理逻辑转为静态图执行,对形状推断、TensorFlow ops的规范性要求更严格;而直接call()是动态图模式,fit()训练时的图追踪逻辑也和推理模式有区别,这就是前两者正常但predict()报错的关键原因。检查
get_patches方法中的形状计算逻辑:
如果你在tf.reshape中用了静态形状(比如直接取images.shape[0]作为batch维度),当predict()处理动态batch(shape为None)时会触发错误。要改用TensorFlow的动态形状API获取维度,示例修改:def get_patches(self, images): # 提取patch的逻辑保持不变 patches = tf.image.extract_patches( images=images, sizes=[1, self.patch_size, self.patch_size, 1], strides=[1, self.patch_size, self.patch_size, 1], rates=[1,1,1,1], padding="VALID" ) # 用tf.shape获取动态形状,避免依赖静态的None维度 patch_dynamic_shape = tf.shape(patches) # 重新reshape,用动态batch维度替代静态的None patches = tf.reshape(patches, (patch_dynamic_shape[0], -1, self.patch_dim)) return patches排查
get_patches中的非TensorFlow操作:如果方法里混合了Python原生运算(比如len()、numpy函数)或非TF控制流(比如普通if/else而不是tf.cond),在静态图追踪时会导致形状推断失败。把所有形状计算、数据处理逻辑都替换成TensorFlow原生操作。显式指定模型输入形状:在模型实例化后,调用
model.build((None, img_height, img_width, img_channels)),给TensorFlow明确的输入静态形状参考,帮助它正确推断后续层的输出形状,避免predict()时出现形状歧义。添加形状断言调试:在
get_patches方法中加入形状断言,快速定位形状异常的位置:def get_patches(self, images): patches = tf.image.extract_patches(...) # 断言提取后的patch形状符合预期,根据你的实际逻辑调整维度 tf.debugging.assert_shape(patches, (None, None, None, self.patch_dim)) patches = tf.reshape(...) return patches运行
predict()时如果触发断言,就能明确是提取patch的步骤还是reshape步骤出了问题。检查自定义层的
call方法是否有条件分支:如果call里存在基于Python变量的分支(而非TensorFlow的tf.cond),静态图追踪时可能无法覆盖所有分支的形状情况,导致形状推断失败。尽量统一分支的输出形状,或改用TF原生控制流。
内容的提问来源于stack exchange,提问作者stschn

