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

TensorFlow 2中如何保存使用StaticVocabularyTable的.pb格式模型

TF2 兼容字符串查找表模型实现方案

核心问题原因

TF2默认开启即时执行模式,直接将Keras Input输出的符号张量传入全局定义的StaticVocabularyTable.lookup会触发类型校验失败;同时TF2废弃了会话和simple_save接口,需要使用原生SavedModel导出逻辑适配资源型变量(如查找表)的保存。

可直接运行的兼容代码

import tensorflow as tf
import numpy as np

# 自定义层封装查找表和嵌入计算逻辑,保证资源变量被正确追踪
class TokenLookupEncoder(tf.keras.layers.Layer):
    def __init__(self, vocabulary, token_embeddings, **kwargs):
        super().__init__(**kwargs)
        # 初始化静态查找表
        self.lookup_table = tf.lookup.StaticVocabularyTable(
            initializer=tf.lookup.KeyValueTensorInitializer(
                keys=vocabulary,
                values=np.arange(len(vocabulary), dtype=np.int64)
            ),
            num_oov_buckets=1
        )
        # 拼接OOV token的零嵌入
        self.token_embeddings = tf.convert_to_tensor(
            np.vstack([token_embeddings, np.zeros(token_embeddings.shape[1])]),
            dtype=tf.float32
        )

    def call(self, input_tokens):
        # 输入为二维字符串张量 [batch_size, token_count]
        token_indices = self.lookup_table.lookup(input_tokens)
        # 转换为one hot后求和得到词袋编码
        one_hot = tf.one_hot(token_indices, depth=tf.cast(self.lookup_table.size(), tf.int32))
        bag_of_tokens = tf.reduce_sum(one_hot, axis=1)
        # 计算嵌入并归一化
        embedded = tf.matmul(bag_of_tokens, self.token_embeddings)
        normed_embedded = embedded / tf.norm(embedded, ord=2, axis=-1, keepdims=True)
        return normed_embedded

# -------------- 配置参数 --------------
vocabulary = ['one', 'two', 'three', 'four', 'five', 'six']
embedding_dim = 512
n_tokens = len(vocabulary)
token_embeddings = np.random.random((n_tokens, embedding_dim))
# 检索用的匹配矩阵
match_matrix = tf.convert_to_tensor(np.random.random((100, embedding_dim)), dtype=tf.float32)

# -------------- 构建模型 --------------
# 输入层对应原TF1的placeholder,维度为[batch_size, 任意长度token序列]
model_input = tf.keras.Input(shape=(None,), dtype=tf.string, name="input")
encoder = TokenLookupEncoder(vocabulary, token_embeddings)
normed_emb = encoder(model_input)
# 计算相似度输出
model_output = tf.matmul(normed_emb, match_matrix, transpose_b=True, name="output")

model = tf.keras.Model(inputs=model_input, outputs=model_output)

# -------------- 验证和导出 --------------
# 测试模型推理
test_input = np.array([["one", "three"], ["two", "four"]])
print(model(test_input).shape) # 输出应为(2, 100),和TF1版本逻辑一致

# 导出为TF Serving兼容的SavedModel格式
model.save(
    "./serving_model",
    save_format="tf",
    signatures=model.call.get_concrete_function(
        tf.TensorSpec(shape=[None, None], dtype=tf.string, name="input")
    )
)

关键适配说明

  • 所有资源型变量(静态查找表)必须封装在tf.keras.layers.Layer或tf.Module中,才能被模型自动追踪,解决符号张量传入lookup的报错问题
  • 导出时显式指定输入签名,保证输入输出的名称、维度和TF1版本完全一致,无需修改原有TensorFlow Serving的部署和调用逻辑
  • 导出后的serving_model目录下包含saved_model.pb文件和变量目录,直接挂载到TensorFlow Serving镜像即可正常加载,无需额外配置初始化算子

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 18:06:08