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

如何编写可随TensorFlow模型一同保存的自定义文本预处理逻辑

如何编写可作为模型组成部分被保存的自定义文本预处理逻辑

不要用TextVectorization自带的standardize回调实现这类需求,这个回调只支持纯TensorFlow原生string算子组合的逻辑,没法灵活实现自定义纠错、动态加词的复杂映射。直接写自定义Keras层实现预处理逻辑即可,层会作为模型的一部分随整个模型一起保存加载,不需要单独维护预处理规则。

核心实现逻辑

  • 自定义层继承tf.keras.layers.Layer,把纠错映射表、查询扩展规则定义为层的初始化属性,模型保存时这些属性会被一起序列化,不会丢失。
  • 字符串处理逻辑写在层的call方法里:如果逻辑简单可以直接用tf原生string算子实现,性能更好;如果是复杂的自定义Python逻辑(比如调用第三方纠错库),用tf.py_function包裹即可,不影响模型保存。
  • 自定义预处理层的输出是字符串张量,可以直接对接TextVectorization层,只要提前把loc_city这类扩展词元加入TextVectorization的词表,后续Embedding层会自动为新增词元训练对应权重,和普通词元没有区别。

可运行代码示例

import tensorflow as tf

class CustomTextPreprocess(tf.keras.layers.Layer):
    def __init__(self, correction_map=None, expand_rules=None, **kwargs):
        super().__init__(**kwargs)
        # 纠错映射、扩展规则作为层属性存储,随模型一起保存
        self.correction_map = correction_map or {"fli": "fly"}
        self.expand_rules = expand_rules or {"london": ["london", "loc_city"]}
    
    def _process_single_text(self, text):
        # 单条文本的自定义处理逻辑,可按需替换为任意纠错、扩展规则
        text = text.numpy().decode("utf-8").lower()
        tokens = text.strip().split()
        # 第一步:词汇纠错
        corrected_tokens = [self.correction_map.get(tok, tok) for tok in tokens]
        # 第二步:查询扩展
        expanded_tokens = []
        for tok in corrected_tokens:
            expanded_tokens.extend(self.expand_rules.get(tok, [tok]))
        return " ".join(expanded_tokens)
    
    def call(self, inputs):
        # 批处理逻辑
        processed = tf.map_fn(
            lambda x: tf.py_function(self._process_single_text, inp=[x], Tout=tf.string),
            inputs,
            fn_output_signature=tf.TensorSpec(shape=(), dtype=tf.string)
        )
        return processed

# 功能测试
preprocess_layer = CustomTextPreprocess()
# 纠错场景测试:输入fli to london,输出fly to london loc_city
print(preprocess_layer(tf.constant(["fli to london"])))
# 扩展场景测试:输入fly to london,输出fly to london loc_city
print(preprocess_layer(tf.constant(["fly to london"])))

# 搭建完整可训练模型
# 提前将扩展词元加入词表
vectorizer = tf.keras.layers.TextVectorization(
    output_sequence_length=10,
    vocabulary=["[UNK]", "fly", "to", "london", "loc_city"]
)

model = tf.keras.Sequential([
    tf.keras.layers.Input(shape=(), dtype=tf.string),
    CustomTextPreprocess(), # 预处理层作为模型第一层
    vectorizer,
    tf.keras.layers.Embedding(input_dim=len(vectorizer.get_vocabulary()), output_dim=32),
    # 后续接下游任务层即可
])

# 整个模型可直接保存,预处理逻辑会被一并存储
model.save("text_model_with_preprocess")
# 加载后可直接推理,不需要额外还原预处理步骤
loaded_model = tf.keras.models.load_model("text_model_with_preprocess")
print(loaded_model(tf.constant(["fli to london"])))

注意事项

如果你的纠错、扩展逻辑依赖外部大词典或模型文件,不要在call方法里写死文件路径,把词典内容、规则参数作为初始化参数传入层中,避免模型迁移到其他环境时加载失败。如果对推理性能要求高,可以把Python实现的处理逻辑替换为TensorFlow原生算子实现,去掉tf.py_function的包裹,推理速度会有明显提升。

关于对接问题:预处理后的输出完全可以直接馈入TextVectorization/Embedding层,新增词元对应的Embedding权重会在训练过程中正常参与梯度更新,不需要额外做特殊处理。

内容的提问来源于stack exchange,提问作者John

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.26 19:09:24