图像字幕模型训练报错,寻求替代代码及Keras fit_generator内存不足训练方案
图像字幕模型训练内存不足问题解决方案
你遇到的报错属于资源耗尽错误,核心原因是一次性加载全量数据集到内存、batch size设置过大等导致硬件资源被占满。以下是附带可运行代码的渐进式训练方案,及相关优化建议:
可运行替代代码
自定义数据集生成器
继承Keras官方Sequence类实现线程安全的渐进式数据加载,无需提前把全量数据集载入内存:
import numpy as np import gc from tensorflow.keras.utils import Sequence from tensorflow.keras.preprocessing.sequence import pad_sequences class ImageCaptionGenerator(Sequence): def __init__(self, img_feat_paths, encoded_captions, vocab_size, max_caption_len, batch_size=8, shuffle=True): self.img_feat_paths = img_feat_paths self.encoded_captions = encoded_captions self.vocab_size = vocab_size self.max_caption_len = max_caption_len self.batch_size = batch_size self.shuffle = shuffle self.on_epoch_end() def __len__(self): return int(np.floor(len(self.img_feat_paths) / self.batch_size)) def __getitem__(self, idx): batch_idx = self.indexes[idx*self.batch_size : (idx+1)*self.batch_size] batch_img_paths = [self.img_feat_paths[i] for i in batch_idx] batch_captions = [self.encoded_captions[i] for i in batch_idx] return self.__generate_batch_data(batch_img_paths, batch_captions) def on_epoch_end(self): self.indexes = np.arange(len(self.img_feat_paths)) if self.shuffle: np.random.shuffle(self.indexes) gc.collect() def __generate_batch_data(self, batch_img_paths, batch_captions): X_img, X_seq, y = [], [], [] for img_path, cap in zip(batch_img_paths, batch_captions): # 按需加载单张图像的预提取特征,避免训练时实时计算特征占显存 img_feat = np.load(img_path, allow_pickle=True) for i in range(1, len(cap)): input_seq = pad_sequences([cap[:i]], maxlen=self.max_caption_len)[0] output_token = np.zeros(self.vocab_size, dtype=np.float16) output_token[cap[i]] = 1 X_img.append(img_feat) X_seq.append(input_seq) y.append(output_token) return [np.array(X_img), np.array(X_seq)], np.array(y)
训练调用代码
from tensorflow import keras # 全局启用混合精度训练,显存占用直接减半 keras.mixed_precision.set_global_policy('mixed_float16') # 配置训练参数 VOCAB_SIZE = 8000 # 替换为你的实际词典大小 MAX_CAPTION_LEN = 30 # 替换为你的字幕最大长度 BATCH_SIZE = 4 # 内存不足可继续调小到2 EPOCHS = 30 # 初始化生成器,仅传入图像特征路径和编码后的字幕,不加载全量数据 train_gen = ImageCaptionGenerator( img_feat_paths=train_img_feature_paths, # 替换为你的训练集图像特征npy文件路径列表 encoded_captions=train_encoded_captions, # 替换为你的训练集编码后字幕列表 vocab_size=VOCAB_SIZE, max_caption_len=MAX_CAPTION_LEN, batch_size=BATCH_SIZE ) val_gen = ImageCaptionGenerator( img_feat_paths=val_img_feature_paths, # 替换为你的验证集图像特征npy文件路径列表 encoded_captions=val_encoded_captions, # 替换为你的验证集编码后字幕列表 vocab_size=VOCAB_SIZE, max_caption_len=MAX_CAPTION_LEN, batch_size=BATCH_SIZE, shuffle=False ) # 模型训练,老版本Keras将fit替换为fit_generator即可 model.fit( train_gen, steps_per_epoch=len(train_gen), epochs=EPOCHS, validation_data=val_gen, validation_steps=len(val_gen) )
额外优化建议
- 提前用ResNet、EfficientNet等CNN模型提取所有图像的特征,存储为本地npy文件,训练阶段直接加载特征,不要在训练时同步做图像特征提取,可降低70%以上的显存占用
- 词典裁剪:只保留训练集中出现次数≥3的词,降低vocab size,减少onehot编码的内存占用
- 小batch训练时可搭配梯度累加:每累计N个batch的梯度后再更新一次权重,弥补小batch带来的训练不稳定问题
- 训练过程中及时删除不再使用的临时变量,调用
gc.collect()手动回收内存
内容的提问来源于stack exchange,提问作者Aravindhan maran
相关产品推荐
相关产品推荐

