使用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

