如何保存无输入的TensorFlow矩阵乘法基准测试模型?
问题解决:TensorFlow保存无输入矩阵乘法模型报错的方案
报错原因
Keras模型保存时需要明确的输入签名来生成完整计算图元数据,你的代码中模型没有定义输入张量,导致无法识别前向传播的输入规范,因此报错“前向传播未定义”。
解决方法
方法1:给Keras模型添加Dummy输入,重写call方法
给模型添加一个不参与计算的占位符输入,让Keras能识别输入签名:
import tensorflow as tf class GEMMBenchmark(tf.keras.Model): def __init__(self, m, n, k): super(GEMMBenchmark, self).__init__() self.A = tf.Variable(tf.random.normal((m, k)), trainable=False) self.B = tf.Variable(tf.random.normal((k, n)), trainable=False) def call(self, inputs): # inputs为占位符输入,不参与实际计算 return tf.matmul(self.A, self.B) # 定义dummy输入,形状可任意设置 dummy_input = tf.keras.Input(shape=()) GEMM = GEMMBenchmark(1, 1, 1) # 通过dummy输入构建计算图 _ = GEMM(dummy_input) # 保存模型 GEMM.save("GEMM") # 加载测试 loaded_model = tf.keras.models.load_model("GEMM") print(loaded_model(tf.constant(0)))
方法2:直接用tf.saved_model.save保存函数(无需Keras Model)
如果不需要Keras Model的特性,直接封装矩阵乘法函数并保存:
import tensorflow as tf m, n, k = 1, 1, 1 A = tf.Variable(tf.random.normal((m, k)), trainable=False) B = tf.Variable(tf.random.normal((k, n)), trainable=False) @tf.function(input_signature=[]) def gemm_benchmark(): return tf.matmul(A, B) # 保存模型 tf.saved_model.save(gemm_benchmark, "GEMM") # 加载测试 loaded = tf.saved_model.load("GEMM") print(loaded())
方法3:用tf.keras.layers.Lambda构建无输入模型
通过Lambda层封装计算,结合占位符输入构建模型:
import tensorflow as tf m, n, k = 1, 1, 1 A = tf.Variable(tf.random.normal((m, k)), trainable=False) B = tf.Variable(tf.random.normal((k, n)), trainable=False) # 定义dummy输入 dummy_input = tf.keras.Input(shape=()) # Lambda层执行矩阵乘法 output = tf.keras.layers.Lambda(lambda x: tf.matmul(A, B))(dummy_input) model = tf.keras.Model(inputs=dummy_input, outputs=output) model.save("GEMM") # 加载测试 loaded_model = tf.keras.models.load_model("GEMM") print(loaded_model(tf.constant(0)))
注意事项
- 方法1和3中的dummy输入仅用于满足Keras的输入要求,推理时传入任意符合形状的张量即可
- 方法2更适合无输入场景,加载后可直接调用,无需传入参数
内容的提问来源于stack exchange,提问作者InputBlackBoxOutput
相关产品推荐
相关产品推荐

