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

Keras Model Subclassing双输入模型报错:NoneType与int无法比较

错误原因定位
  • 核心触发点:直接将tf.data.BatchDataset对象作为输入传给Keras模型,而非实际的张量数据。从报错日志的调用参数可以看到,传入的train_img_b、train_ans_b是数据集迭代器对象,不是具体数值张量。TimeDistributed层初始化时会校验输入张量的维度,因为无法从Dataset对象解析得到有效维度值,出现NoneType与整数比较的类型错误。
  • 其余会阻塞运行的潜在问题:
    • ArchitectureNet层的call方法定义了prev_output、dataset_embed两个必传参数,但实际调用时仅传入了dataset_embed,运行时会触发参数缺失报错。
    • 图像张量dtype为tf.float32,答案张量dtype为tf.float64,后续拼接操作会触发类型不匹配报错。
    • 分开对图像、答案数据集做batch,无法保证两个数据集的批次样本一一对应。
    • plot_model传入的是类对象StructureModel,而非实例化构建完成的模型对象,会触发参数类型错误。
    • ArchitectureNet的call方法中调用tf.make_ndarray尝试在图模式下将张量转numpy数组,会触发运行时错误。
可行修复方案

按以下步骤修改代码即可解决当前TypeError,同时规避后续运行的显性错误:

  1. 正确构造配对的训练数据集,不要分开batch两个独立数据集,压缩配对后统一张量类型、生成批次:
    替换原有数据集构造、前向传播测试的代码段:
# 替换原有的分开batch逻辑
train_dataset = tf.data.Dataset.zip((train_img, train_ans))
# 统一所有输入张量dtype为float32
train_dataset = train_dataset.map(lambda img, ans: (tf.cast(img, tf.float32), tf.cast(ans, tf.float32)))
train_dataset = train_dataset.batch(batch_size)

structuremodel = StructureModel()
# 从数据集中取1个实际的张量批次做前向传播测试,禁止直接传入Dataset对象
for test_imgs, test_ans in train_dataset.take(1):
    hnet_output, anet_output = structuremodel([test_imgs, test_ans])
    break
  1. 修正ArchitectureNet层的逻辑,匹配当前调用方式:
class ArchitectureNet(keras.layers.Layer):
    def __init__(self, anet_pred_vars, **kwargs):
        super().__init__()
        self.anet_pred_vars = anet_pred_vars
        self.dense1 = Dense(units=50, activation='relu')
        self.dense2 = Dense(units=50, activation='relu')
        self.anet_output = Dense(units=self.anet_pred_vars, name='Architecture')
        self.stopping_node = Dense(units=1, activation='sigmoid')

    def call(self, dataset_embed):
        x = self.dense1(dataset_embed)
        x = self.dense2(x)
        anet_output = self.anet_output(x)
        stop_node_output = self.stopping_node(x)
        return anet_output, stop_node_output

如果后续需要循环调用该网络、传入上一步的输出结果,再对应加回prev_output参数和拼接逻辑即可,当前单次前向传播场景下不需要该参数。
3. 对应调整StructureModel中调用anet_layer的返回值接收逻辑,以及模型最终返回值:

# anet部分
anet_output, stop_output = self.anet_layer(dataset_embed)
return hnet_output, anet_output, stop_output
  1. 修正plot_model的调用逻辑,必须传入实例化且build完成的模型对象:
# 先跑一次前向传播完成模型权重初始化、形状构建
for test_imgs, test_ans in train_dataset.take(1):
    structuremodel([test_imgs, test_ans])
    break
plot_model(structuremodel, to_file='aeu.png', show_shapes=True)

内容的提问来源于stack exchange,提问作者Hudson Liu

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.26 18:09:16