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

TensorFlow Keras to_categorical内存错误:无法分配221GiB内存

内存错误:无法为形状(3324240, 17877)的float32数组分配221 GiB空间

问题场景

我用pandas加载无标点CSV文本,将高频词定义为“重要词”,随后构建Keras Seq2Seq模型的编码器和解码器输入(转为numpy数组),但运行到to_categorical步骤时触发内存错误。

核心错误提示:

Unable to allocate 221. GiB for an array with shape (3324240, 17877) and data type float32.

完整错误栈:

---------------------------------------------------------------------------
MemoryError                               Traceback (most recent call last)
Cell In[183], line 2
      1 de=decoder_inp_final[:17876]
----> 2 decoder_inp_final=to_categorical(decoder_inp_final)

File d:\Desktop AI\env\Lib\site-packages\keras\utils\np_utils.py:73, in to_categorical(y, num_classes, dtype)
     71     num_classes = np.max(y) + 1
     72 n = y.shape[0]
---> 73 categorical = np.zeros((n, num_classes), dtype=dtype)
     74 categorical[np.arange(n), y] = 1
     75 output_shape = input_shape + (num_classes,)

MemoryError: Unable to allocate 221. GiB for an array with shape (3324240, 17877) and data type float32

相关代码

from tensorflow.python.keras.models import Model
from tensorflow.keras.layers import Dense, Embedding, Input, LSTM
from keras.utils import to_categorical
from keras.utils import pad_sequences
import pandas as pd
import numpy as np

word2count = {}
for _, line in df.iterrows():
    for word in line['no_punc'].split():
        if word not in word2count:
            word2count[word] = 1
        else:
            word2count[word] += 1

important_word={}
thresh = 5
important_word['<PAD>'] = 0
word_id = 1
for key, value in word2count.items():
    if value > thresh:
        important_word[key] = word_id
        word_id += 1

questions = {}
answers = {}

for q, ans in convo_line.items():
    questions[q] = df.loc[q, 'no_punc']
    answers[ans] = df.loc[ans, 'no_punc']

for key, value in answers.items():
    answers[key] = '<SOS>' + value + '<EOS>'

tokens=['<EOS>', '<OUT>', '<SOS>']

x = len(important_word)

for t in tokens:
    important_word[t] = x
    x += 1

encoder_inp = []
for id, line in questions.items():
    lst = []
    for word in line.split():
        if word not in important_word:
            lst.append(important_word['<OUT>'])
        else:
            lst.append(important_word[word])
    encoder_inp.append(lst)

decoder_inp = []

for id, line in answers.items():
    lst=[]
    for word in line.split():
        if word not in important_word:
            lst.append(important_word['<OUT>'])
        else:
            lst.append(important_word[word])
    decoder_inp.append(lst)

encoder_inp = pad_sequences(encoder_inp, 15, padding='post', truncating='post')
decoder_inp = pad_sequences(decoder_inp, 15, padding='post', truncating='post')

decoder_inp_final = []

for i in decoder_inp:
    decoder_inp_final.append(i[1:])

decoder_inp_final = pad_sequences(decoder_inp_final, 15, padding='post', truncating='post')

de = decoder_inp_final[:17876]
decoder_inp_final = to_categorical(decoder_inp_final) 

解决方案

1. 移除预one-hot编码,改用稀疏分类损失

Seq2Seq模型的解码器输出不需要提前做one-hot编码,直接保留整数标签,配合sparse_categorical_crossentropy损失函数即可,这能直接避免超大数组的内存占用:

  • 删除decoder_inp_final = to_categorical(decoder_inp_final)这行代码
  • 模型输出层设置为Dense(len(important_word), activation='softmax')
  • 编译模型时指定损失函数:
model.compile(optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy'])

2. 精简词汇表

  • 提高高频词阈值thresh,减少important_word的总数量,缩小分类维度
  • 进一步过滤低频词,或使用词频统计工具合并语义相近的词汇

3. 用批量数据生成器加载数据

如果数据量过大,改用TensorFlow的tf.data.Dataset或Keras生成器,每次仅加载一个批次的数据,避免一次性占用全部内存:

import tensorflow as tf

# 构建数据集并设置批次大小
dataset = tf.data.Dataset.from_tensor_slices((encoder_inp, decoder_inp_final))
dataset = dataset.batch(32)  # 按需调整批次大小

# 训练时直接传入数据集
model.fit(dataset, epochs=10)

4. 降低数据类型精度(仅保留one-hot时用)

如果必须使用one-hot编码,可将数据类型从float32改为uint8(one-hot仅含0和1),大幅减少内存占用:

decoder_inp_final = to_categorical(decoder_inp_final, dtype='uint8')

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 13:17:50