Kaggle中TensorFlow/Keras LSTM因RAM过载崩溃,求数据集处理方案
解决Kaggle上LSTM训练RAM过载问题
问题根源
你的代码将整个文本数据集一次性转换成独热编码的X和y数组并加载到内存中,这是导致RAM过载的核心原因。3.95MB的文本会生成数十万条样本,每条样本对应100长度×词汇量维度的独热向量,总内存占用会远超Kaggle的可用RAM上限。
解决方案
采用动态数据生成的方式,避免一次性加载所有预处理后的数据到内存。同时用整数编码替代独热编码,配合Embedding层进一步降低内存消耗,具体步骤如下:
1. 用整数编码替代独热编码
独热编码会大幅增加数据维度,改用整数编码(每个字符对应一个索引),再通过模型的Embedding层转换成向量,能显著减少内存占用。
2. 使用tf.data.Dataset动态生成样本
TensorFlow的tf.data.Dataset支持从文本中动态切分样本、生成批次,无需提前把所有样本存入内存。
3. 可选参数优化
- 适当增大
steps(步长),减少总样本数 - 若仍有压力,减小
max_length(输入序列长度)或batch_size
修改后的完整代码
from __future__ import absolute_import, division, print_function, unicode_literals import numpy as np import tensorflow as tf from tensorflow.keras.models import Sequential from tensorflow.keras.layers import Dense, Activation, LSTM, Embedding from tensorflow.keras.optimizers import RMSprop from tensorflow.keras.callbacks import LambdaCallback, ModelCheckpoint, ReduceLROnPlateau import random import sys # 加载文本 with open('/kaggle/input/crptic-python/dataset.txt', 'r') as file: text = file.read() # 构建词汇表与映射 vocabulary = sorted(list(set(text))) vocab_size = len(vocabulary) char_to_idx = dict((c, i) for i, c in enumerate(vocabulary)) idx_to_char = dict((i, c) for i, c in enumerate(vocabulary)) # 参数设置 max_length = 100 steps = 5 batch_size = 128 # 用tf.data.Dataset动态生成样本 def create_dataset(text, max_length, steps): # 将文本转为整数序列 text_int = np.array([char_to_idx[c] for c in text]) # 生成输入序列和目标序列 sequences = [] targets = [] for i in range(0, len(text_int) - max_length, steps): sequences.append(text_int[i:i+max_length]) targets.append(text_int[i+max_length]) # 转为tf.data.Dataset dataset = tf.data.Dataset.from_tensor_slices((sequences, targets)) # 打乱、分批、预取 dataset = dataset.shuffle(len(sequences)).batch(batch_size, drop_remainder=True).prefetch(tf.data.AUTOTUNE) return dataset train_dataset = create_dataset(text, max_length, steps) # 构建模型:用Embedding层替代提前独热编码 model = Sequential() # Embedding层:vocab_size -> 64维向量(可根据需求调整) model.add(Embedding(input_dim=vocab_size, output_dim=64, input_length=max_length)) model.add(LSTM(128)) model.add(Dense(vocab_size)) model.add(Activation('softmax')) optimizer = RMSprop(learning_rate=0.01) model.compile(loss='sparse_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) # epoch结束生成文本的回调(保持原逻辑,仅调整编码部分) def on_epoch_end(epoch, logs): if epoch % 30 == 0: print() print('----- 第{}轮训练后生成文本'.format(epoch)) start_index = random.randint(0, len(text) - max_length - 1) for diversity in [0.2, 0.5, 1.0, 1.2]: print('----- 多样性参数:', diversity) generated = '' sentence = text[start_index: start_index + max_length] generated += sentence print('----- 初始文本: "{}"'.format(sentence)) sys.stdout.write(generated) for i in range(400): # 转为整数编码输入 x_pred = np.array([[char_to_idx[c] for c in sentence]]) 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_lr = ReduceLROnPlateau(monitor='loss', factor=0.2, patience=1, min_lr=0.001) callbacks = [print_callback, checkpoint, reduce_lr] # 训练模型:传入dataset而非X、y model.fit(train_dataset, epochs=28, 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.array([[char_to_idx[c] for c in sentence]]) 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(500, 0.5))
关键优化点说明
- 整数编码+Embedding层:将原先生成的高维度独热数组替换为一维整数序列,Embedding层在训练时动态转换为低维向量,内存占用至少降低一个数量级。
- tf.data.Dataset动态生成:只在需要时生成当前批次的样本,不会把所有样本加载到内存,彻底解决RAM过载问题。
- 损失函数调整:因为目标是整数而非独热向量,改用
sparse_categorical_crossentropy损失,无需对y做独热编码。
内容的提问来源于stack exchange,提问作者ProgrammerGuy
相关产品推荐
相关产品推荐

