Keras训练时如何对接非Keras词嵌入模型并获取输出?
解决方案
不要用run_eagerly=True,性能损失太大,有两个兼容图模式的可行方案,按优先级从高到低列:
方案1:预计算FastText嵌入+查表(性能最优,优先选)
你的场景里电影标题的词表是固定的,完全不需要训练时实时调用FastText计算,提前把所有标题的嵌入算好存成查找表,训练时直接查表就行,全程跑纯TF图模式,零额外开销。
实现步骤:
- 提前遍历所有电影标题,计算每个标题分词后的平均FastText嵌入,和词表对齐生成权重矩阵
import numpy as np # 假设fasttext_model是你加载好的预训练FastText模型 embed_dim = fasttext_model.wv.vector_size fasttext_embed_weights = [] for title in unique_movie_titles: words = title.strip().split() # 替换成你实际的分词逻辑 valid_word_embeds = [fasttext_model.wv[w] for w in words if w in fasttext_model.wv] if valid_word_embeds: avg_embed = np.mean(valid_word_embeds, axis=0).astype(np.float32) else: avg_embed = np.zeros(embed_dim, dtype=np.float32) fasttext_embed_weights.append(avg_embed) # 补一个mask位的嵌入,和StringLookup的输出对齐 fasttext_embed_weights.append(np.zeros(embed_dim, dtype=np.float32)) fasttext_embed_weights = np.array(fasttext_embed_weights)
- 在模型里加一个不可训练的Embedding层存预计算的FastText权重,和原有逻辑完全兼容
class MovieModel(tf.keras.Model): def __init__(self, fasttext_embed_weights): super().__init__() max_tokens = 10_000 # 共享StringLookup层,保证id映射一致 self.title_lookup = tf.keras.layers.StringLookup( vocabulary=unique_movie_titles, mask_token=None ) # 原有可训练的标题嵌入 self.title_embedding = tf.keras.Sequential([ self.title_lookup, tf.keras.layers.Embedding(len(unique_movie_titles) + 1, 32) ]) # 固定权重的FastText嵌入,不参与训练 self.fasttext_embedding = tf.keras.Sequential([ self.title_lookup, tf.keras.layers.Embedding( input_dim=len(unique_movie_titles) + 1, output_dim=fasttext_embed_weights.shape[1], weights=[fasttext_embed_weights], trainable=False ) ]) def call(self, inputs): title_input = inputs["movie_title"] return tf.concat([ self.title_embedding(title_input), self.fasttext_embedding(title_input) ], axis=1)
方案2:用tf.py_function包装实时计算逻辑(适合输入不固定的场景)
如果你的输入是动态的、没法提前预计算所有嵌入,就用tf.py_function把FastText调用逻辑包装成TF可识别的自定义op,只有FastText计算部分走Python环境,其余前向反向传播还是跑图模式,性能比run_eagerly=True高一个量级。
实现步骤:
- 先写好处理numpy格式输入的FastText批量计算函数,注意tf.py_function传入的字符串是bytes类型,需要先解码
def batch_fasttext_process(title_batch_np): # 输入是shape=(batch_size,)的numpy字节数组,输出是shape=(batch_size, embed_dim)的float32数组 embed_dim = fasttext_model.wv.vector_size batch_result = [] for title_bytes in title_batch_np: title_str = title_bytes.decode("utf-8") words = title_str.strip().split() valid_embeds = [fasttext_model.wv[w] for w in words if w in fasttext_model.wv] avg_emb = np.mean(valid_embeds, axis=0).astype(np.float32) if valid_embeds else np.zeros(embed_dim, dtype=np.float32) batch_result.append(avg_emb) return np.array(batch_result, dtype=np.float32)
- 在模型里封装调用逻辑,记得手动设置输出形状,避免图模式形状推断失败
class MovieModel(tf.keras.Model): def __init__(self, fasttext_model): super().__init__() self.fasttext_model = fasttext_model self.embed_dim = fasttext_model.wv.vector_size max_tokens = 10_000 self.title_embedding = tf.keras.Sequential([ tf.keras.layers.StringLookup( vocabulary=unique_movie_titles,mask_token=None), tf.keras.layers.Embedding(len(unique_movie_titles) + 1, 32) ]) def _get_fasttext_embeds(self, title_tensor): embeds = tf.py_function( func=batch_fasttext_process, inp=[title_tensor], Tout=tf.float32 ) # 手动补全形状信息 embeds.set_shape([None, self.embed_dim]) return embeds def call(self, inputs): title_input = inputs["movie_title"] return tf.concat([ self.title_embedding(title_input), self._get_fasttext_embeds(title_input) ], axis=1)
注意事项
- 不要尝试在图模式下直接取张量的
.numpy()值或者遍历张量内容,图构建阶段张量没有实际值,这么做只会报错 - 方案1是固定词表场景下的最优解,性能和纯Keras模型没有区别;方案2适合动态输入场景,性能损耗仅来自Python侧的FastText计算部分,远好于全局开eager模式
- 如果用方案2,注意不要在处理函数里放任何和TensorFlow相关的计算,纯numpy/Python逻辑的运行效率最高
内容的提问来源于stack exchange,提问作者heisenberg
相关产品推荐
相关产品推荐

