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

如何正确将自定义split函数传入TextVectorization层并正常运行

问题原因

TextVectorization层默认期望split操作返回规则稠密张量,你写的split_slash函数处理批量输入时会返回tf.RaggedTensor类型的不规则输出,未显式配置的情况下层会将每个样本拆分后的所有token整体识别为单个元素,因此出现每个输入仅对应1个ID的错误结果。额外给split函数加@tf.function装饰器也可能干扰层内部的输入形状推断逻辑。

解决方法

修改TextVectorization层的初始化参数即可,具体调整如下:

  1. 显式配置ragged=True,告诉层接受不规则长度的token输出
  2. 移除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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 10:27:00