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

如何用TensorFlow.Keras Sequential类实现嵌入层点积并设置权重

问题描述

尝试使用tensorflow.keras构建模型,计算两个带预定义权重(训练阶段可优化)的嵌入层的点积。其中:

  • embedding_layer_1需作为推理时的查找表,设置trainable=True,对应的weights_matrix形状为(288, 3569)
  • embedding_layer_2是embedding_layer_1的转置,形状应为(3569, 288)
    使用tensorflow==2.8.0、keras==2.8.0编写代码后报错,错误提示添加的不是Layer实例,尝试转置等方法未解决。

原代码

import tensorflow as tf
from tensorflow.keras import Sequential
from tensorflow.keras.layers import Flatten, Embedding, Dot, Flatten

input_dim=288
output_dim=3569

model_1 = Sequential()
embedding_layer_1 = Embedding(input_dim=input_dim, output_dim=output_dim, name='embedding_layer_1', dtype='float64', trainable=True,  input_length=1)
embedding_layer_1.build((None,))
model_1.add(embedding_layer_1)
model_1.layers[0].set_weights([weights_matrix])
model_1.add(Flatten())

model_2 = Sequential()
embedding_layer_2 = Embedding(input_dim=input_dim, output_dim=output_dim, name='embedding_layer_2', dtype='float64', trainable=True,  input_length=1)
embedding_layer_2.build((None,))
model_2.add(embedding_layer_2)
model_2.layers[0].set_weights([role_skill_matrix])
model_2.add(Flatten())

dot_product = Dot(axes=-1)([model_1.output, model_2.output])

model = Sequential([model_1, model_2, dot_product])
model.summary()

报错信息

model = Sequential([model_1, model_2, dot_product])
  File "/Users/ayalaallon/opt/anaconda3/envs/ml-pipeline/lib/python3.8/site-packages/tensorflow/python/training/tracking/base.py", line 629, in _method_wrapper
    result = method(self, *args, **kwargs)
  File "/Users/ayalaallon/opt/anaconda3/envs/ml-pipeline/lib/python3.8/site-packages/keras/utils/traceback_utils.py", line 67, in error_handler
    raise e.with_traceback(filtered_tb) from None
  File "/Users/ayalaallon/opt/anaconda3/envs/ml-pipeline/lib/python3.8/site-packages/keras/engine/sequential.py", line 178, in add
    raise TypeError('The added layer must be an instance of class Layer. '
TypeError: The added layer must be an instance of class Layer. Received: layer=KerasTensor(type_spec=TensorSpec(shape=(None, 1), dtype=tf.float32, name=None), name='dot/Squeeze:0', description="created by layer 'dot'") of type <class 'keras.engine.keras_tensor.KerasTensor'>

Process finished with exit code 1
错误原因分析
  1. Sequential模型的局限性:Sequential只能接收Layer实例作为堆叠元素,你传入的dot_product是Dot层计算后的KerasTensor对象,并非Layer本身;同时Sequential是线性堆叠结构,无法处理多输入分支的合并需求。
  2. 嵌入层参数错误:根据你描述的转置形状,embedding_layer_2的input_dim应为3569、output_dim应为288,但原代码中参数和embedding_layer_1一致,会导致权重形状不匹配。
  3. 冗余操作:手动调用build方法无必要,Sequential.add会自动处理层的构建逻辑。
解决方案

改用Keras函数式API构建模型,适配多输入分支结构,同时修正嵌入层参数:

import tensorflow as tf
from tensorflow.keras.layers import Input, Embedding, Flatten, Dot
from tensorflow.keras.models import Model

# 定义嵌入层维度
emb1_input_dim = 288
emb1_output_dim = 3569
# embedding_layer_2是emb1的转置,维度反转
emb2_input_dim = emb1_output_dim
emb2_output_dim = emb1_input_dim

# 定义两个输入(对应两个嵌入层的查询索引)
input1 = Input(shape=(1,), name='input1')
input2 = Input(shape=(1,), name='input2')

# 构建第一个嵌入层并加载预定义权重
emb1 = Embedding(
    input_dim=emb1_input_dim,
    output_dim=emb1_output_dim,
    name='embedding_layer_1',
    dtype='float64',
    trainable=True,
    input_length=1
)(input1)
emb1_flat = Flatten()(emb1)
# 加载预定义权重
emb1_flat._keras_history[0].set_weights([weights_matrix])

# 构建第二个嵌入层,加载转置后的权重
emb2 = Embedding(
    input_dim=emb2_input_dim,
    output_dim=emb2_output_dim,
    name='embedding_layer_2',
    dtype='float64',
    trainable=True,
    input_length=1
)(input2)
emb2_flat = Flatten()(emb2)
# 加载转置后的权重(如果role_skill_matrix是emb1权重的转置,直接使用即可)
emb2_flat._keras_history[0].set_weights([role_skill_matrix])

# 计算两个扁平嵌入向量的点积
dot_product = Dot(axes=-1, name='dot_product')([emb1_flat, emb2_flat])

# 构建完整模型
model = Model(inputs=[input1, input2], outputs=dot_product)
model.summary()

关键说明

  1. 函数式API适配多分支:通过Input定义独立输入分支,分别连接嵌入层,最后用Dot层合并输出,完美支持你的分支计算需求。
  2. 修正维度匹配:根据转置后的权重形状调整embedding_layer_2的输入输出维度,避免形状不兼容错误。
  3. 简化权重加载:通过_keras_history[0]获取嵌入层实例,直接设置预定义权重,无需手动调用build。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.09 18:40:10