TensorFlow2.6训练CRNN时Model.fit报str无shape属性错误
报错原因
- 直接触发点:你在生成器返回的模型输入字典里加入了
source_str字段,对应值是字符串列表。这个字段仅用于结果可视化,不属于CRNN模型输入层定义的接收参数。Keras处理输入时会遍历字典内所有值,读取.shape属性做格式校验,字符串类型没有该属性,直接抛出对应错误。 - 隐性用法错误:你将数据生成器类继承了
keras.callbacks.Callback,还把生成器实例直接传入了callbacks回调列表。Callback是训练生命周期钩子的基类,数据生成器不需要继承该类,也不能作为回调传入fit,否则后续会触发其他逻辑错误。
修复步骤
- 第一步:清理生成器返回的输入字典,移除非模型输入字段
把next_batch()方法里构造inputs字典的代码中'source_str': source_str这一行删掉,不要把可视化用的字符串数据混在模型输入里。修改后的输入构造代码:inputs = { 'img_input': X_data, 'ground_truth_labels': Y_data, 'input_length': input_length, 'label_length': label_length } outputs = {'ctc': np.zeros([self.batch_size])} yield (inputs, outputs) - 第二步:修正类继承关系和fit传参
- 把类定义从
class DataGenerator(keras.callbacks.Callback):改为class DataGenerator:,数据生成器不需要继承回调基类。 - 从
fit()方法的callbacks列表中移除train_gene、val_gen两个生成器实例,回调列表只保留真正的回调对象,修改后的训练代码:img_text_recog.fit( x = train_gene.next_batch(), steps_per_epoch=int(train_gene.n / batch_size), epochs=20, callbacks=[viz_cb_train,viz_cb_val,tensorboard_callback,early_stop,model_chk_pt], validation_data=val_gen.next_batch(), validation_steps=int(val_gen.n / batch_size) )
- 把类定义从
- 第三步:(可选)适配可视化逻辑
如果训练过程中需要展示原始识别文本,可以在自定义可视化回调内部单独获取批次对应的原始文本,不要通过模型输入通道传递字符串数据。
前置校验方法
正式启动训练前,可以先拉取一批生成器数据做格式校验,避免反复触发训练报错:
test_inputs, test_outputs = next(train_gene.next_batch()) for key, val in test_inputs.items(): print(f"输入字段名:{key},数据类型:{type(val)},数据形状:{val.shape}")
确认所有输入字段均为数值型numpy数组、shape符合模型输入要求后,再启动训练即可。
内容的提问来源于stack exchange,提问作者htbi127
相关产品推荐
相关产品推荐

