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

使用tf.layers.TextVectorization处理9GB医学文献遇内存错误求助

问题描述

处理约300k个总计9GB的医学文献文件,目标是统计所有token的频率并导出为CSV(格式:token, frequency)。使用layers.TextVectorization(output_mode='count')适配数据集成功,但调用vectorize_layer.get_vocabulary()时触发MemoryError,报错栈如下:

2024-04-20 22:38:56.832518: I tensorflow/core/platform/cpu_feature_guard.cc:193] This TensorFlow binary is optimized with oneAPI Deep Neural Network Library (oneDNN) to use the following CPU instructions in performance-critical operations:  AVX2
To enable them in other operations, rebuild TensorFlow with the appropriate compiler flags.
2024-04-20 22:38:56.833545: I tensorflow/core/common_runtime/process_util.cc:146] Creating new thread pool with default inter op setting: 2. Tune using inter_op_parallelism_threads for best performance.
Discoverd 323719 files
Finished adapting
Traceback (most recent call last):
  File "(file_location_of_source_code)", line 55, in <module>
    inverse_vocab = vectorize_layer.get_vocabulary()
  File "D:\Anaconda\envs\src\lib\site-packages\keras\layers\preprocessing\text_vectorization.py", line 487, in get_vocabulary
    return self._lookup_layer.get_vocabulary(include_special_tokens)
  File "D:\Anaconda\envs\src\lib\site-packages\keras\layers\preprocessing\index_lookup.py", line 385, in get_vocabulary
    self._tensor_vocab_to_numpy(vocab),
  File "D:\Anaconda\envs\src\lib\site-packages\keras\layers\preprocessing\string_lookup.py", line 416, in _tensor_vocab_to_numpy
    [tf.compat.as_text(x, self.encoding) for x in vocabulary]
  File "D:\Anaconda\envs\src\lib\site-packages\keras\layers\preprocessing\string_lookup.py", line 416, in <listcomp>
    [tf.compat.as_text(x, self.encoding) for x in vocabulary]
MemoryError

Process finished with exit code 1

相关核心代码:

files_root = pathlib.Path(r"directoryname")
files = tf.data.TextLineDataset.list_files(str(files_root/'*'))
text_ds = tf.data.TextLineDataset(files).filter(lambda x: tf.cast(tf.strings.length(x), bool))

vectorize_layer = layers.TextVectorization(
    standardize=custom_standardization,
    output_mode='count')
print(f"Discoverd {len(files)} files")
vectorize_layer.adapt(text_ds.batch(1024))
print("Finished adapting")
inverse_vocab = vectorize_layer.get_vocabulary() # 触发报错

原计划合并频率数组的代码(小数据集有效):

text_vector_ds = text_ds.batch(1024).prefetch(AUTOTUNE).map(vectorize_layer).unbatch()
it = text_vector_ds.as_numpy_iterator()
freq_arr = None
for i, entry in enumerate(text_vector_ds.as_numpy_iterator()):
    if i == 0:
        freq_arr = np.zeros(len(entry))
        freq_arr += entry.astype(int)
    else:
        freq_arr += entry.astype(int)

需求:解决内存问题,获取token-频率映射,同时得到处理大型数据集的指导,最终用于skip-gram训练。


解决方案

1. 限制词汇表大小(优先推荐)

医学文献中存在大量低频专业术语或罕见拼写,这些token对skip-gram训练价值极低,直接限制词汇表最大数量,既能解决内存问题,又能提升后续训练效率。

修改TextVectorization初始化代码,添加max_tokens参数(比如保留前10万高频token,可根据需求调整):

vectorize_layer = layers.TextVectorization(
    standardize=custom_standardization,
    output_mode='count',
    max_tokens=100000  # 只保留频率最高的10万个token
)
vectorize_layer.adapt(text_ds.batch(1024))
inverse_vocab = vectorize_layer.get_vocabulary()
# 此时词汇表大小受限,不会触发内存溢出

之后合并频率数组时,数组长度就是max_tokens+2(默认包含""和"[UNK]"两个特殊token),内存占用可控。

2. 分批导出词汇表(需保留全部词汇时)

如果必须保留所有token,可通过访问vectorize_layer._lookup_layer的底层张量,分批转换为numpy字符串,避免一次性加载所有词汇到内存:

import tensorflow as tf

def get_vocab_in_batches(layer, batch_size=10000):
    vocab_tensor = layer._lookup_layer.vocabulary()
    vocab_size = tf.shape(vocab_tensor)[0].numpy()
    vocab = []
    for start in range(0, vocab_size, batch_size):
        end = min(start + batch_size, vocab_size)
        batch = vocab_tensor[start:end].numpy()
        vocab.extend([tf.compat.as_text(x, layer.encoding) for x in batch])
    return vocab

# 调用方法获取词汇表
inverse_vocab = get_vocab_in_batches(vectorize_layer)

该方法通过分批处理词汇张量,每次只加载部分数据到内存,避免一次性转换整个大张量导致的内存溢出。

3. 换用轻量统计方式(绕过TextVectorization的内存瓶颈)

直接用Python的collections.Counter结合tf.data分批处理文本,全程控制内存占用:

from collections import Counter
import tensorflow as tf
import pathlib
import csv

def process_batch(batch):
    # 应用自定义标准化和分词
    standardized = custom_standardization(batch)
    tokens = tf.strings.split(standardized)
    # 转换为Python列表统计
    flat_tokens = tf.reshape(tokens, [-1]).numpy().astype(str)
    return Counter(flat_tokens)

files_root = pathlib.Path(r"directoryname")
files = tf.data.TextLineDataset.list_files(str(files_root/'*'))
text_ds = tf.data.TextLineDataset(files).filter(lambda x: tf.cast(tf.strings.length(x), bool))

# 分批处理并累加统计结果
total_counter = Counter()
for batch in text_ds.batch(1024).prefetch(tf.data.AUTOTUNE):
    batch_counter = process_batch(batch)
    total_counter.update(batch_counter)

# 导出为CSV
with open('token_frequencies.csv', 'w', newline='', encoding='utf-8') as f:
    writer = csv.writer(f)
    writer.writerow(['token', 'frequency'])
    for token, freq in total_counter.most_common():
        writer.writerow([token, freq])

这种方式无需依赖TextVectorization的词汇表存储,直接分批统计,内存占用完全由批次大小控制,适合超大规模数据集。

4. 针对skip-gram训练的优化指导

  • 低频token过滤:无论用哪种方式统计,建议过滤掉出现次数少于5次的token,这类token不仅增加内存负担,还会降低skip-gram的训练效果。
  • 使用TFRecord存储预处理数据:将分词后的文本或索引序列保存为TFRecord格式,后续训练时可直接加载,避免重复预处理。
  • 采用负采样:skip-gram训练时用负采样替代全量softmax,大幅降低内存和计算开销,适合大型词汇表场景。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 09:47:07