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

训练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实现按需懒加载:

  1. 第一阶段仅逐行扫描数据集,统计字符集、最大序列长度、构建字符索引映射,这部分元信息内存占用可以忽略
  2. 训练阶段每次仅读取、处理当前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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.02 05:00:45