Keras自定义ShiftViT模型虚拟输入测试报错排查
报错原因
你自定义的ShiftViTModel为Keras子类化模型,但未重写父类要求的call()方法:你将前向传播逻辑全部耦合在仅用于训练/验证阶段计算损失的_calculate_loss方法中,当直接传入输入调用模型实例时,Keras找不到可执行的前向传播路径,就会触发该报错。
修复步骤
- 第一步:在
ShiftViTModel类中补充标准call()方法,将纯特征提取的前向传播逻辑从_calculate_loss中抽离到该方法内,代码如下:
def call(self, inputs, training=False): x = self.data_augmentation(inputs, training=training) x = self.patch_projection(x) for stage in self.stages: x = stage(x, training=training) logits = self.global_avg_pool(x) return logits
- 第二步:简化原有
_calculate_loss方法,直接调用call()获取模型输出,避免重复维护前向传播逻辑,修改后代码如下:
def _calculate_loss(self, data, training=False): (images, labels) = data logits = self(images, training=training) total_loss = self.compiled_loss(labels, logits) return total_loss, labels, logits
- 第三步:修改完成后,原有的虚拟输入测试代码即可正常执行,输出模型输出维度:
dummy_inputs = tf.ones((2, 32, 32, 3)) outputs = model(dummy_inputs, training=False) print(outputs.shape)
额外提示
- 上述修改不会改动你原有
train_step、test_step的训练逻辑,模型训练流程完全不受影响 - 补充
call()方法后,你还可以直接调用model.summary()打印完整模型结构与参数量,更方便做结构正确性校验
内容的提问来源于stack exchange,提问作者Jacob
相关产品推荐
相关产品推荐

