如何在TensorFlow中使用TextVectorization预处理层实现字符串按字符拆分
Great question! The problem with using a pure Python lambda like lambda x: list(x) here is that TensorFlow's TextVectorization layer operates on TensorFlow tensors, not raw Python strings. That lambda works in pure Python because it handles regular string objects, but it’ll throw errors in TensorFlow’s graph mode since it can’t process tensor inputs natively.
The TensorFlow-Compatible Solution
Instead, we can use TensorFlow’s built-in tf.strings.unicode_split function—it’s designed to split string tensors into individual character tensors, and it’s fully compatible with TensorFlow’s computation graph. This is the safe, efficient way to handle character-level splitting for the layer.
Here’s the corrected, working code:
import tensorflow as tf from tensorflow.keras import layers data = tf.constant( [ "The Brain is wider than the Sky", "For put them side by side", "The one the other will contain", "With ease and You beside", ] ) # Use TensorFlow's unicode_split for tensor-safe character splitting text_vectorizer = layers.TextVectorization( output_mode="int", split=lambda x: tf.strings.unicode_split(x, input_encoding='UTF-8') ) # Build the vocabulary index from your data text_vectorizer.adapt(data) # Test the vectorizer to verify results vectorized_data = text_vectorizer(data) print("Character Vocabulary:", text_vectorizer.get_vocabulary()) print("\nVectorized first sample:\n", vectorized_data[0].numpy())
Key Details
tf.strings.unicode_split(x, input_encoding='UTF-8')takes each string tensorxand splits it into a tensor of individual Unicode characters. This works seamlessly in both eager and graph modes, and supports GPU/TPU execution.- After calling
adapt(), your vocabulary will include all unique individual characters from your input data (plus the default empty string token for padding, if required). - The output will be integer sequences where each number maps directly to a single character from your vocabulary.
Why Avoid the Pure Python Lambda?
When you use a Python lambda that calls list(x), TensorFlow has to convert tensors to raw Python strings to process them—this is inefficient, breaks graph serialization, and won’t work reliably on accelerated hardware. Using TensorFlow-native functions keeps everything within the framework’s ecosystem for optimal performance and compatibility.
内容的提问来源于stack exchange,提问作者Oleg

