如何在Keras中使用Tokenizer类将单词列表编码为INT8类型独热矩阵
How to Generate INT8 One-Hot Encoding Matrices with Keras Tokenizer
Great question—switching to int8 for one-hot encoding is a smart way to cut down on memory usage, since float64 takes 8x more space than int8. Here's a straightforward approach to achieve this:
Step 1: Initialize and Fit the Tokenizer
First, set up your Tokenizer and train it on your text data to build the vocabulary:
from tensorflow.keras.preprocessing.text import Tokenizer import numpy as np # Example text data texts = [ "The cat sat on the mat", "The dog chased the cat", "The cat and dog are friends" ] # Initialize Tokenizer (adjust num_words based on your vocabulary size) tokenizer = Tokenizer(num_words=100) # Fit on texts to build word index tokenizer.fit_on_texts(texts)
Step 2: Generate One-Hot Matrix and Cast to INT8
The texts_to_matrix method defaults to float64, but since one-hot encoding (using mode='binary') only uses 0s and 1s, we can safely cast the result to np.int8 without losing any information:
# Generate binary one-hot matrix (default dtype is float64) one_hot_float = tokenizer.texts_to_matrix(texts, mode='binary') # Convert to int8 to save memory one_hot_int8 = one_hot_float.astype(np.int8)
Key Notes:
- Mode Selection: Use
mode='binary'for true one-hot encoding (indicates presence/absence of words). If you're using other modes like'count'(word frequency), ensure your count values stay within the int8 range (-128 to 127) before casting—this works for most small-to-medium datasets. - Sequence-Level One-Hot Encoding: If you need one-hot vectors for individual words in sequences (instead of whole-document vectors), you can do this:
from tensorflow.keras.preprocessing.sequence import pad_sequences # Get integer sequences from texts sequences = tokenizer.texts_to_sequences(texts) # Pad sequences to a uniform length padded_sequences = pad_sequences(sequences, maxlen=10) # Convert each sequence to one-hot and cast to int8 one_hot_sequences = np.array([np.eye(len(tokenizer.word_index) + 1)[seq] for seq in padded_sequences]).astype(np.int8)
This approach will slash your memory usage significantly—perfect for avoiding out-of-memory errors when working with large text corpora!
内容的提问来源于stack exchange,提问作者Shashi Tunga
相关产品推荐
相关产品推荐

