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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.01 02:42:31