为什么绘制Keras(TensorFlow 2.0)模型图时未包含矩阵乘法相关变量
Keras模型可视化时inputP节点丢失的原因及解决方案
问题原因
- 导入路径不兼容:代码中混用了两种Keras引用,首先将
tf.keras赋值给了变量keras,后续导入层和模型类时又直接使用独立的keras库(非TF内置的tf.keras),两个不同的Keras实例对张量的计算图追踪逻辑不互通,导致inputP的张量依赖关系没有被模型正确识别。 - 原生TensorFlow算子未封装:代码中直接使用
tf.transpose、tf.matmul这类TensorFlow原生算子,没有封装为Keras层,而keras.utils.plot_model是基于Keras层的拓扑关系绘图,不是追踪底层TensorFlow计算图,因此inputP的路径没有被可视化工具捕获。
修复方案
- 统一使用
tf.keras的导入路径:
import tensorflow as tf from tensorflow.keras.layers import Input, Dense, Lambda from tensorflow.keras.models import Model
- 将原生TensorFlow算子封装为Keras Lambda层,确保拓扑关系被正确追踪:
a = Input(shape=(138,7), name='inputP') b = Input(shape=(138,7), name='inputQ') # 封装transpose操作 c = Lambda(lambda x: tf.transpose(x, [0,2,1]), name='transpose_q')(b) # 封装matmul操作 d = Lambda(lambda x: tf.matmul(x[0], x[1]), name='matmul_pq')([c, a]) e = Dense(15,activation = 'relu')(d) model = Model([a,b],e) tf.keras.utils.plot_model(model)
修改后重新执行即可在可视化图中看到两个输入节点。
内容的提问来源于stack exchange,提问作者Shivam Pande
相关产品推荐
相关产品推荐

