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
相关产品推荐
相关产品推荐

