使用fit_generator实现数据增强时触发TypeError:float()参数应为字符串或数字而非BatchDataset的问题求助
你遇到的TypeError核心原因很明确:tf.keras.utils.image_dataset_from_directory返回的是TensorFlow的BatchDataset对象,而ImageDataGenerator.flow()方法只接受numpy数组作为输入——它尝试把Dataset转换成numpy数组时失败,所以抛出了float() argument must be a string or a number, not 'BatchDataset'这个错误。
另外还要提醒你:fit_generator在TensorFlow 2.10版本之后已经被正式弃用了,现在直接用model.fit()就可以处理生成器或者Dataset对象,完全不需要再用fit_generator。
下面给你两种可行的解决方案,优先推荐第一种:
方案1:使用TensorFlow内置的数据增强层(最推荐)
这种方式把数据增强逻辑直接集成到模型中,不仅代码简洁,而且增强操作会和模型一起保存(后续推理时也能复用,当然推理时通常会关闭增强),同时完美兼容TF Dataset管道,效率更高。
修改你的模型定义,把数据增强层加到最前面:
model = Sequential([ # 数据增强层,放在输入之后、预处理之前 layers.RandomFlip("horizontal", input_shape=(img_height, img_width, 3)), layers.RandomRotation(0.1), # 对应原rotation_range=20(可根据需求调整弧度值) layers.RandomZoom(0.15), # 对应原zoom_range=0.15 layers.RandomTranslation(height_factor=0.2, width_factor=0.2), # 对应原height/width_shift_range=0.2 # 原有的预处理层 layers.Rescaling(1./127.5, offset=-1), # Encoder部分 layers.Conv2D(8, 3, activation='relu'), layers.MaxPooling2D(), layers.Conv2D(16, 3, activation='relu'), layers.MaxPooling2D(), layers.Conv2D(32, 3, activation='relu'), layers.Flatten(), # Decoder部分 layers.Dense(64, activation='relu'), layers.Dropout(0.5), layers.Dense(2, activation='softmax') ])
然后直接用model.fit()训练即可,不需要ImageDataGenerator和fit_generator:
# 训练网络 history = model.fit( train_ds, validation_data=val_ds, epochs=epochs )
方案2:将Dataset转换为numpy数组(适合小数据集)
如果你坚持想用ImageDataGenerator,可以把train_ds和val_ds转换成numpy数组,再传入aug.flow():
首先添加转换函数:
import numpy as np def dataset_to_numpy(ds): images = [] labels = [] for x, y in ds: images.append(x.numpy()) labels.append(y.numpy()) return np.concatenate(images), np.concatenate(labels) train_images, train_labels = dataset_to_numpy(train_ds) val_images, val_labels = dataset_to_numpy(val_ds)
然后修改训练代码:
# 训练网络 history = model.fit( aug.flow(train_images, train_labels, batch_size=batch_size), validation_data=(val_images, val_labels), steps_per_epoch=len(train_images) // batch_size, epochs=epochs )
⚠️ 注意:这种方式会把所有数据加载到内存中,如果你的数据集很大,很容易出现内存不足的问题,所以只适合小数据集使用。
额外修复点
你代码最后调用了show_history(history),但没有定义这个函数,补充一下实现:
def show_history(history): plt.figure(figsize=(12, 4)) # 绘制损失曲线 plt.subplot(1, 2, 1) plt.plot(history.history['loss'], label='Training Loss') plt.plot(history.history['val_loss'], label='Validation Loss') plt.xlabel('Epoch') plt.ylabel('Loss') plt.legend() plt.title('Loss Curve') # 绘制准确率曲线 plt.subplot(1, 2, 2) plt.plot(history.history['accuracy'], label='Training Accuracy') plt.plot(history.history['val_accuracy'], label='Validation Accuracy') plt.xlabel('Epoch') plt.ylabel('Accuracy') plt.legend() plt.title('Accuracy Curve') plt.tight_layout() plt.show()
内容的提问来源于stack exchange,提问作者Oliver

