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

Keras做One Hot Encoding训练词级LSTM时出现内存错误如何解决

核心问题原因

你碰到的内存报错完全来自to_categorical生成的独热标签:118万样本 × 32360词汇量 × int32单值4字节 = 约143GiB,这个量级的数组不可能直接加载到普通内存中。

生成器是否是正确解决方向?

是,但它不是当前场景下成本最低的解决方案,你可以先用更简单的方案解决问题。

可落地的优先级解决方案

1. 最高优先级:替换损失函数,完全规避独热编码(仅需修改2行代码)

Keras内置的sparse_categorical_crossentropy损失函数原生支持整数类型的标签,完全不需要你手动把标签转成独热编码,直接可以省掉143GiB的内存占用,是当前最适合你的方案。
修改步骤:

  • 删除one_hot_labels = to_categorical(labels, num_classes=vocab_size, dtype='int32')这行代码
  • 模型编译时把损失函数替换为sparse_categorical_crossentropy
  • 调用model.fit时直接传入原始整数标签labels即可

修改后的对应代码片段:

# 删掉原来的to_categorical行,直接用原始labels
input_sequences, labels = seq[:,:-1], seq[:,-1]

# 编译时替换损失函数
model.compile(loss='sparse_categorical_crossentropy', optimizer='adam', metrics=['accuracy'])

# 训练时传入labels而非one_hot_labels
history = model.fit(input_sequences, labels, epochs=epochs, batch_size=batch_size, callbacks=callbacks_list, verbose=1)

2. 次优先级优化(可选,进一步降低内存占用)

如果替换损失函数后,输入序列input_sequences还是占内存过大,可以做以下优化:

  • 过滤低频词:统计词频后,把出现次数低于23次的词统一替换为`<UNK>`标记,把词汇量降到1万2万区间,既可以降低内存占用,也可以减少模型过拟合概率
  • 调整序列长度:不要用全局最长序列作为统一padding长度,可以根据文本长度分布取95分位数作为padding长度,截断过长的文本,减少输入序列的维度
  • 用tf.data.Dataset构建数据流水线:比手动写生成器逻辑更简单,支持动态加载、预取、批次处理,不需要把所有输入序列一次性加载到内存中,示例代码:
import tensorflow as tf
dataset = tf.data.Dataset.from_tensor_slices((input_sequences, labels))
dataset = dataset.shuffle(10000).batch(batch_size).prefetch(tf.data.AUTOTUNE)
history = model.fit(dataset, epochs=epochs, callbacks=callbacks_list, verbose=1)

3. 生成器方案适用场景

如果你的文本规模进一步扩大,连input_sequences都无法一次性加载到内存中,再考虑自定义生成器,每批次动态读取文本、转序列、padding、生成标签即可。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.26 18:54:03