Keras 2.1.2中batch_size与steps_per_epoch冲突的解决方法咨询
在Keras 2.1.2里,这个报错的原因很明确:当你设置steps_per_epoch时,Keras要求batch_size必须设为None——因为steps_per_epoch是用来指定每个epoch要执行多少个批次的,这时候模型不会自动根据batch_size去划分你的输入数据,而是依赖你提供的批次生成逻辑来控制每个批次的大小。
要实现你“每个epoch处理固定100个批次、每个批次256个样本”的需求,有两种靠谱的解决办法:
方法一:使用数据生成器(推荐)
Keras的生成器机制天生适配steps_per_epoch的使用场景,你可以通过自定义生成器或者内置的生成工具来控制每个批次的大小,同时实现epoch前的数据洗牌。
自定义生成器示例(适用于numpy数组数据)
如果你的训练数据是numpy数组,可以写一个简单的生成器函数:
import numpy as np def custom_generator(x, y, batch_size): while True: # 每个epoch开始前洗牌数据 shuffle_indices = np.random.permutation(len(x)) x_shuffled = x[shuffle_indices] y_shuffled = y[shuffle_indices] # 按批次生成数据 for i in range(0, len(x_shuffled), batch_size): yield x_shuffled[i:i+batch_size], y_shuffled[i:i+batch_size]
然后修改你的fit调用,把原始数据换成生成器,同时设置batch_size=None:
model.fit( custom_generator(xtrain, ytrain, 256), epochs=100, callbacks=get_callbacks(patience=10), validation_data=(xval, yval), class_weight=get_class_weight_trainset(ytrain), steps_per_epoch=100, batch_size=None )
这个生成器会在每次循环(对应一个epoch)时自动洗牌数据,每次输出256个样本,steps_per_epoch=100会让模型每个epoch跑100个这样的批次,正好符合你的需求。
内置生成器(适用于图像数据)
如果你的数据是图像,可以用ImageDataGenerator的flow方法快速实现:
from keras.preprocessing.image import ImageDataGenerator datagen = ImageDataGenerator() train_generator = datagen.flow(xtrain, ytrain, batch_size=256, shuffle=True) model.fit( train_generator, epochs=100, callbacks=get_callbacks(patience=10), validation_data=(xval, yval), class_weight=get_class_weight_trainset(ytrain), steps_per_epoch=100, batch_size=None )
flow方法默认会在每个epoch前洗牌数据,同样能满足你的要求。
方法二:手动控制每个epoch的数据
如果你不想用生成器,可以手动在每个epoch前洗牌并截取固定数量的样本,然后用常规的fit调用(不设置steps_per_epoch):
for epoch in range(100): # 洗牌并截取100*256=25600个样本 shuffle_indices = np.random.permutation(len(xtrain)) xtrain_subset = xtrain[shuffle_indices[:25600]] ytrain_subset = ytrain[shuffle_indices[:25600]] # 训练一个epoch model.fit( xtrain_subset, ytrain_subset, batch_size=256, epochs=1, callbacks=get_callbacks(patience=10), validation_data=(xval, yval), class_weight=get_class_weight_trainset(ytrain) )
这种方式需要手动循环每个epoch,注意如果你的回调(比如EarlyStopping)需要跨epoch监控指标,可能需要额外调整回调的逻辑,所以更推荐第一种生成器的方法。
内容的提问来源于stack exchange,提问作者Marvin Lerousseau

