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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.03 07:01:26