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

图像字幕模型训练报错,寻求替代代码及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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 18:06:03