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,同时规避后续运行的显性错误:
- 正确构造配对的训练数据集,不要分开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
- 修正
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
- 修正
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
相关产品推荐
相关产品推荐

