训练200万样本字符级循环seq2seq模型报OOM错误如何解决
问题根因
你遇到的内存不足问题和模型训练流程无关,核心出在数据预处理逻辑:你在训练开始前就把全量200万条样本全部转成了稠密one-hot格式的numpy数组常驻内存。
简单核算内存占用:假设输入、输出字符集各100个字符,最长序列长度各200,你创建的三个三维数组(encoder输入、decoder输入、decoder标签)每个元素为占4字节的float32类型,总内存需求约为2000000 * 200 * 100 * 4 * 3 ≈ 480GB,普通消费级硬件完全无法承载。
另外你原代码中f.read().split("\n")的读文件方式也会把全量文本一次性加载进内存,会进一步加剧内存占用。
解决方案
核心思路是放弃全量数据预加载,用tf.data.Dataset实现按需懒加载:
- 第一阶段仅逐行扫描数据集,统计字符集、最大序列长度、构建字符索引映射,这部分元信息内存占用可以忽略
- 训练阶段每次仅读取、处理当前batch需要的样本,实时完成编码、padding、one-hot转换,用完立即释放内存,不会长期驻留
适配后的可运行代码
import tensorflow as tf from tensorflow import keras import numpy as np # 基础超参数和原代码保持一致 batch_size = 64 epochs = 100 latent_dim = 256 num_samples = 2000000 data_path = "es.txt" # -------------------------- # 第一步:仅扫描元信息,不加载全量文本 # -------------------------- input_characters = set() target_characters = set() max_encoder_seq_length = 0 max_decoder_seq_length = 0 valid_sample_count = 0 with open(data_path, "r", encoding="utf-8") as f: for line in f: if valid_sample_count >= num_samples: break line = line.strip("\n") if not line: continue parts = line.split("\t") input_text = parts[0] target_text = "\t" + parts[1] + "\n" # 加目标序列起止符 # 更新统计值 max_encoder_seq_length = max(max_encoder_seq_length, len(input_text)) max_decoder_seq_length = max(max_decoder_seq_length, len(target_text)) input_characters.update(list(input_text)) target_characters.update(list(target_text)) valid_sample_count += 1 # 构建字符索引映射 input_characters = sorted(list(input_characters)) target_characters = sorted(list(target_characters)) num_encoder_tokens = len(input_characters) num_decoder_tokens = len(target_characters) input_token_index = {c:i for i,c in enumerate(input_characters)} target_token_index = {c:i for i,c in enumerate(target_characters)} pad_enc_id = input_token_index[" "] pad_dec_id = target_token_index[" "] # 打印统计信息确认 print(f"有效样本数: {valid_sample_count}") print(f"输入字符集大小: {num_encoder_tokens}, 输出字符集大小: {num_decoder_tokens}") print(f"输入最大序列长度: {max_encoder_seq_length}, 输出最大序列长度: {max_decoder_seq_length}") # -------------------------- # 第二步:构建tf.data流水线 # -------------------------- def process_single_line(line): parts = tf.strings.split(line, sep="\t", maxsplit=2) input_text = parts[0] target_text = tf.strings.join(["\t", parts[1], "\n"]) # 字符转索引 enc_chars = tf.strings.unicode_split(input_text, "UTF-8") enc_ids = tf.map_fn( lambda c: input_token_index[c.numpy().decode("utf-8")], enc_chars, fn_output_signature=tf.int32 ) dec_chars = tf.strings.unicode_split(target_text, "UTF-8") dec_ids = tf.map_fn( lambda c: target_token_index[c.numpy().decode("utf-8")], dec_chars, fn_output_signature=tf.int32 ) # 构造decoder输入和标签:标签比输入偏移一个时间步 dec_input_ids = dec_ids dec_target_ids = tf.concat([dec_ids[1:], [pad_dec_id]], axis=0) return enc_ids, dec_input_ids, dec_target_ids def tf_wrap_process(line): enc_ids, dec_in_ids, dec_tgt_ids = tf.py_function( process_single_line, inp=[line], Tout=[tf.int32, tf.int32, tf.int32] ) enc_ids.set_shape([None]) dec_in_ids.set_shape([None]) dec_tgt_ids.set_shape([None]) return (enc_ids, dec_in_ids), dec_tgt_ids def build_dataset(is_train=True): # 逐行读取文件,不加载全量到内存 ds = tf.data.TextLineDataset(data_path) ds = ds.take(num_samples) # 按8:2拆分训练/验证集,和原代码validation_split=0.2逻辑一致 if is_train: ds = ds.take(int(num_samples * 0.8)) else: ds = ds.skip(int(num_samples * 0.8)) # 逐行处理样本 ds = ds.map(tf_wrap_process, num_parallel_calls=tf.data.AUTOTUNE) # 按batch做动态padding ds = ds.padded_batch( batch_size, padded_shapes=( ([max_encoder_seq_length], [max_decoder_seq_length]), [max_decoder_seq_length] ), padding_values=((pad_enc_id, pad_dec_id), pad_dec_id), drop_remainder=False ) # 仅对当前batch做one-hot转换,避免全量转换占用内存 def batch_to_onehot(batch_x, batch_y): enc_x, dec_x = batch_x enc_x = tf.one_hot(enc_x, depth=num_encoder_tokens, dtype=tf.float32) dec_x = tf.one_hot(dec_x, depth=num_decoder_tokens, dtype=tf.float32) batch_y = tf.one_hot(batch_y, depth=num_decoder_tokens, dtype=tf.float32) return (enc_x, dec_x), batch_y ds = ds.map(batch_to_onehot, num_parallel_calls=tf.data.AUTOTUNE) # 预加载下一批次数据,提升训练速度 ds = ds.prefetch(tf.data.AUTOTUNE) return ds train_ds = build_dataset(is_train=True) val_ds = build_dataset(is_train=False) # -------------------------- # 第三步:模型构建与训练(结构和原代码完全一致) # -------------------------- encoder_inputs = keras.Input(shape=(None, num_encoder_tokens)) encoder = keras.layers.LSTM(latent_dim, return_state=True) encoder_outputs, state_h, state_c = encoder(encoder_inputs) encoder_states = [state_h, state_c] decoder_inputs = keras.Input(shape=(None, num_decoder_tokens)) decoder_lstm = keras.layers.LSTM(latent_dim, return_sequences=True, return_state=True) decoder_outputs, _, _ = decoder_lstm(decoder_inputs, initial_state=encoder_states) decoder_dense = keras.layers.Dense(num_decoder_tokens, activation="softmax") decoder_outputs = decoder_dense(decoder_outputs) model = keras.Model([encoder_inputs, decoder_inputs], decoder_outputs) model.compile(optimizer="rmsprop", loss="categorical_crossentropy", metrics=["accuracy"]) # 训练时直接传入数据集对象,不需要传入全量numpy数组 model.fit( train_ds, validation_data=val_ds, epochs=epochs ) model.save("ess2s")
额外优化建议
- 如果运行时仍出现显存不足,直接把
batch_size下调到32或16即可 - 可以先统计句子长度的分位数,过滤掉长度超过99分位的极长句子,减少无效padding占用的内存
- 后续迭代可以把one-hot输入替换为
Embedding层,直接输入字符索引即可,不需要做one-hot转换,能进一步降低内存占用、提升训练速度
内容的提问来源于stack exchange,提问作者Laz22434
相关产品推荐
相关产品推荐

