TensorFlow自定义Estimator结合Dataset API的NMT任务嵌入查找服务问题
搞定NMT自定义Estimator服务时的Feature Columns难题
兄弟,我太懂你这糟心事儿了——用Dataset API喂数据搞自定义NMT Estimator,结果部署服务的时候被feature_columns卡得死死的,毕竟你的任务说白了就一个输入语句特征,这玩意儿怎么凑feature_columns啊?别慌,给你捋两个实用的解决思路,都是我踩过坑后总结的:
方案一:词表场景用分类列+嵌入列
如果你的输入语句已经分词,而且有固定的词表,直接这么整:
- 先定义对应词表的分类列:
# 假设你有个包含所有词汇的vocab_list,比如["<PAD>", "<START>", "hello", "world"...] vocab_column = tf.feature_column.categorical_column_with_vocabulary_list( key="input_sentence", # 这个键名要和服务请求里的字段完全对应,别写错! vocabulary_list=vocab_list, dtype=tf.string # 如果是词ID就换成tf.int64 ) - 再把分类列转成嵌入列(毕竟NMT肯定要词嵌入嘛):
embedding_feature_column = tf.feature_column.embedding_column( categorical_column=vocab_column, dimension=256 # 嵌入维度根据你模型调,常用256/512 ) - 然后在自定义
model_fn里用input_layer解析输入:def model_fn(features, labels, mode, params): # 用feature columns把输入转成模型能认的格式 input_layer = tf.feature_column.input_layer(features, params['feature_columns']) # 转成序列形状喂给编码器:(batch_size, 序列长度, 嵌入维度) seq_input = tf.reshape(input_layer, [-1, params['max_seq_len'], 256]) # 接下来就是你熟门熟路的编码器、解码器逻辑了... - 最后创建Estimator的时候把feature_columns传进去就行:
estimator = tf.estimator.Estimator( model_fn=model_fn, params={ 'feature_columns': [embedding_feature_column], 'max_seq_len': 50, # 你的最大输入序列长度 # 其他模型参数比如隐藏层大小啥的自己加 } )
方案二:预处理好的数值序列直接用数值列
如果你的输入已经是处理好的词ID序列(比如形状是[batch_size, max_seq_len]的数值张量),那就更简单了,直接定义数值列:
input_seq_column = tf.feature_column.numeric_column( key="input_sentence", shape=[50], # 填你的最大序列长度 dtype=tf.int64 )
然后在model_fn里直接拿这个特征用就行,甚至不用过input_layer:
def model_fn(features, labels, mode, params): input_seq = features['input_sentence'] # 查嵌入表喂给编码器 embedding_table = tf.get_variable('embedding', [params['vocab_size'], 256]) embedded_input = tf.nn.embedding_lookup(embedding_table, input_seq) # 后面该咋写咋写
几个要注意的坑
- 服务请求的输入字段名必须和feature_column定义的
key完全一致!比如你设的是input_sentence,那请求就得传{"input_sentence": ["hello", "world"]}或者对应的数值序列。 - 如果你的输入是原始未分词的字符串,建议在服务前的预处理环节先分词,别把逻辑塞模型里。真要在模型里处理的话,可以用哈希桶列:
但这种方式不如固定词表靠谱,慎用。text_column = tf.feature_column.categorical_column_with_hash_bucket( key="input_sentence", hash_bucket_size=100000 # 哈希桶大小根据你的词汇量估 )
内容的提问来源于stack exchange,提问作者Andrew Schenck
相关产品推荐
相关产品推荐

