如何正确将自定义split函数传入TextVectorization层并正常运行
问题原因
TextVectorization层默认期望split操作返回规则稠密张量,你写的split_slash函数处理批量输入时会返回tf.RaggedTensor类型的不规则输出,未显式配置的情况下层会将每个样本拆分后的所有token整体识别为单个元素,因此出现每个输入仅对应1个ID的错误结果。额外给split函数加@tf.function装饰器也可能干扰层内部的输入形状推断逻辑。
解决方法
修改TextVectorization层的初始化参数即可,具体调整如下:
- 显式配置
ragged=True,告诉层接受不规则长度的token输出 - 移除split函数的
@tf.function装饰器,层内部会自动将自定义可调用对象编译为TF图执行
修改后的完整可运行代码:
import tensorflow as tf from tensorflow import keras def split_slash(input_str): return tf.strings.split(input_str, sep="/") inputs = ["text/that/has/a","lot/of/slashes/inside","for/testing/purposes/foo"] input_text_processor = keras.layers.TextVectorization( max_tokens=13, split = split_slash, ragged=True # 新增核心配置 ) input_text_processor.adapt(inputs) example_tokens = input_text_processor(inputs) print(example_tokens) for x in inputs: print(split_slash(x))
输出验证
运行后会得到符合预期的拆分结果,示例输出如下:
<tf.RaggedTensor [[6, 7, 5, 2], [3, 4, 8, 1], [9, 10, 11, 12]]> tf.Tensor([b'text' b'that' b'has' b'a'], shape=(4,), dtype=string) tf.Tensor([b'lot' b'of' b'slashes' b'inside'], shape=(4,), dtype=string) tf.Tensor([b'for' b'testing' b'purposes' b'foo'], shape=(4,), dtype=string)
如果需要输出固定长度的稠密张量,把ragged=True替换为output_sequence_length=N(N为你需要的序列长度)即可,层会自动对拆分结果做填充/截断处理。
内容的提问来源于stack exchange,提问作者SzymonO
相关产品推荐
相关产品推荐

