如何降低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
相关产品推荐
相关产品推荐

