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

如何降低AI训练时Data Generator的RAM占用(Kaggle环境)

大文本语料Keras训练RAM耗尽问题解决方案

问题背景

在Kaggle笔记本(13GB RAM上限)中处理1GB+文本语料,使用Python+Keras构建字符级LSTM模型时,原有的Data Generator失效,RAM直接耗尽。已尝试调小batch_size、hidden_size,调大step间隔,但问题未解决。

核心问题分析

原代码中提前将整个文本切分为所有子序列并存储在sentences列表中,这是内存耗尽的关键原因:1GB文本按101字符的子序列切分后,会生成数百万个子序列,全部加载到内存中直接占满RAM,Data Generator的按需生成机制完全没发挥作用。

具体解决方案

1. 重构Data Generator,按需生成子序列

直接基于原始文本字符串生成batch数据,不再提前存储所有子序列。修改后的TextDataGenerator会在每个batch请求时,从预先生成的起始位置索引中截取对应文本片段,避免一次性加载所有子序列:

from __future__ import absolute_import, division, print_function, unicode_literals
import numpy as np
import tensorflow as tf
from keras.utils import Sequence
from keras.models import Sequential
from keras.layers import Dense, Activation, LSTM
from keras.optimizers import RMSprop
from keras.callbacks import LambdaCallback, ModelCheckpoint, ReduceLROnPlateau
import random
import sys

class TextDataGenerator(Sequence):
    def __init__(self, text, vocab, char_to_idx, max_length, batch_size, step=1):
        self.text = text
        self.vocab = vocab
        self.char_to_idx = char_to_idx
        self.max_length = max_length
        self.batch_size = batch_size
        self.step = step
        # 计算所有有效起始位置数量
        self.total_positions = (len(text) - max_length - 1) // step
        self.steps = self.total_positions // batch_size
        # 预先生成起始位置索引,洗牌时仅操作索引,不触碰大文本
        self.positions = [i * step for i in range(self.total_positions)]

    def __len__(self):
        return self.steps

    def __getitem__(self, idx):
        batch_positions = self.positions[idx*self.batch_size : (idx+1)*self.batch_size]
        # 用uint8替代bool,内存占用一致且兼容性更好
        X = np.zeros((self.batch_size, self.max_length, len(self.vocab)), dtype=np.uint8)
        y = np.zeros((self.batch_size, len(self.vocab)), dtype=np.uint8)
        
        for i, pos in enumerate(batch_positions):
            # 截取输入序列和目标字符
            seq = self.text[pos : pos+self.max_length]
            target_char = self.text[pos+self.max_length]
            # 填充one-hot矩阵
            for t, char in enumerate(seq):
                X[i, t, self.char_to_idx[char]] = 1
            y[i, self.char_to_idx[target_char]] = 1
        return X, y

    def on_epoch_end(self):
        # 仅洗牌起始位置索引,大幅降低内存开销
        random.shuffle(self.positions)

# 读取文本(1GB文本读入内存在13GB环境下可行)
with open('/kaggle/input/crptic-python/python.txt', 'r') as file:
    text = file.read()

vocab = sorted(list(set(text)))
char_to_idx = {c:i for i,c in enumerate(vocab)}
idx_to_char = {i:c for i,c in enumerate(vocab)}

max_length = 100
batch_size = 32
step = 10  # 保持原有步长,减少训练样本数量

# 构建模型(可根据需求进一步调小hidden_size至64)
model = Sequential()
model.add(LSTM(64, input_shape=(max_length, len(vocab))))
model.add(Dense(len(vocab)))
model.add(Activation('softmax'))
optimizer = RMSprop(lr=0.01)
model.compile(loss='categorical_crossentropy', optimizer=optimizer)

# 采样函数保持原逻辑
def sample_index(preds, temperature=1.0):
    preds = np.asarray(preds).astype('float64')
    preds = np.log(preds) / temperature
    exp_preds = np.exp(preds)
    preds = exp_preds / np.sum(exp_preds)
    probas = np.random.multinomial(1, preds, 1)
    return np.argmax(probas)

# 回调函数优化:减少生成文本长度,降低内存开销
def on_epoch_end(epoch, logs):
    if epoch % 1 == 0:
        print(f"\n----- 生成文本(Epoch: {epoch})")
        start_index = random.randint(0, len(text) - max_length - 1)
        for diversity in [0.1, 0.3, 0.5]:
            print(f"----- diversity: {diversity}")
            generated = ''
            sentence = text[start_index: start_index + max_length]
            generated += sentence
            print(f"----- 初始文本: \"{sentence}\"")
            sys.stdout.write(generated)

            # 生成长度从400降至200,减少内存占用
            for i in range(200):
                x_pred = np.zeros((1, max_length, len(vocab)), dtype=np.uint8)
                for t, char in enumerate(sentence):
                    x_pred[0, t, char_to_idx[char]] = 1.
                preds = model.predict(x_pred, verbose=0)[0]
                next_index = sample_index(preds, diversity)
                next_char = idx_to_char[next_index]
                generated += next_char
                sentence = sentence[1:] + next_char
                sys.stdout.write(next_char)
                sys.stdout.flush()
            print()

print_callback = LambdaCallback(on_epoch_end=on_epoch_end)
filepath = "weights.hdf5"
checkpoint = ModelCheckpoint(filepath, monitor='loss', verbose=1, save_best_only=True, mode='min')
reduce_alpha = ReduceLROnPlateau(monitor='loss', factor=0.2, patience=1, min_lr=0.001)
callbacks = [print_callback, checkpoint, reduce_alpha]

# 初始化生成器并训练
data_generator = TextDataGenerator(text, vocab, char_to_idx, max_length, batch_size, step)
model.fit(data_generator, epochs=2, callbacks=callbacks)

def generate_text(length, diversity):
    start_index = random.randint(0, len(text) - max_length - 1)
    generated = ''
    sentence = text[start_index: start_index + max_length]
    generated += sentence
    for i in range(length):
        x_pred = np.zeros((1, max_length, len(vocab)), dtype=np.uint8)
        for t, char in enumerate(sentence):
            x_pred[0, t, char_to_idx[char]] = 1.
        preds = model.predict(x_pred, verbose=0)[0]
        next_index = sample_index(preds, diversity)
        next_char = idx_to_char[next_index]
        generated += next_char
        sentence = sentence[1:] + next_char
    return generated

print(generate_text(300, 0.5))

2. 进一步优化内存:用索引替代one-hot矩阵存储

若词汇表较大,可改为存储字符索引,在生成器中转换为one-hot矩阵,将内存占用降低len(vocab)倍:

修改生成器的__getitem__方法:

def __getitem__(self, idx):
    batch_positions = self.positions[idx*self.batch_size : (idx+1)*self.batch_size]
    # 存储字符索引,内存占用从(batch, max_len, vocab)降至(batch, max_len)
    X = np.zeros((self.batch_size, self.max_length), dtype=np.int32)
    y = np.zeros((self.batch_size,), dtype=np.int32)
    
    for i, pos in enumerate(batch_positions):
        seq = self.text[pos : pos+self.max_length]
        target_char = self.text[pos+self.max_length]
        X[i] = [self.char_to_idx[char] for char in seq]
        y[i] = self.char_to_idx[target_char]
    # 转换为one-hot矩阵
    X = tf.one_hot(X, depth=len(self.vocab), dtype=tf.uint8)
    y = tf.one_hot(y, depth=len(self.vocab), dtype=tf.uint8)
    return X.numpy(), y.numpy()

3. 其他辅助优化

  • 调大step间隔:将step从10调到20,进一步减少总训练样本数
  • 开启混合精度训练:减少模型参数的内存占用
    from keras.mixed_precision import set_global_policy
    set_global_policy('mixed_float16')
    

内容的提问来源于stack exchange,提问作者ProgrammerGuy

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 05:55:03