TensorFlow Hub的MoveNet模型:固定输入形状等三类技术问题咨询
MoveNet模型问题解决方案
问题1:固定输入形状为(1,256,256,3)并生成TensorFlow Lite模型
直接加载的MoveNet模型输入形状是动态的,需要通过包装固定输入签名后再保存,才能生成固定形状的TFLite模型:
import tensorflow as tf import tensorflow_hub as tfhub input_size = 256 model_url = "https://www.kaggle.com/models/google/movenet/TensorFlow2/multipose-lightning/1" model = tfhub.load(model_url) movenet = model.signatures['serving_default'] # 定义带固定输入签名的函数 @tf.function(input_signature=[tf.TensorSpec(shape=(1, input_size, input_size, 3), dtype=tf.int32, name='input')]) def fixed_shape_infer(input): return movenet(input=input) # 保存带固定签名的模型 tf.saved_model.save(model, './fixed_shape_movenet', signatures={'serving_default': fixed_shape_infer}) # 转换为TensorFlow Lite模型 converter = tf.lite.TFLiteConverter.from_saved_model('./fixed_shape_movenet') tflite_model = converter.convert() with open('movenet_fixed_256.tflite', 'wb') as f: f.write(tflite_model)
问题2:无法显示模型摘要的解决方法
tfhub.load返回的是_UserObject类型,不是标准Keras模型,因此没有summary()方法。可以将其包装为Keras模型来查看摘要:
# 包装为Keras模型 input_layer = tf.keras.layers.Input(shape=(input_size, input_size, 3), batch_size=1, dtype=tf.int32, name='input') output_tensor = model(input_layer) keras_model = tf.keras.Model(inputs=input_layer, outputs=output_tensor) keras_model.summary()
也可以通过以下方式查看模型的变量信息:
# 查看所有可训练变量 print("可训练变量列表:") for var in model.trainable_variables: print(f"{var.name}: {var.shape}")
问题3:加载saved_model后查看输入输出形状
直接保存原始_UserObject可能会丢失签名信息,需要按问题1的方式保存带指定签名的模型。加载后查看输入输出形状的方法如下:
# 加载固定形状的模型 loaded_model = tf.saved_model.load('./fixed_shape_movenet') # 查看所有可用签名 print("可用签名:", list(loaded_model.signatures.keys())) # 获取签名并查看输入输出 infer_func = loaded_model.signatures['serving_default'] print("输入签名:", infer_func.structured_input_signature) print("输出签名:", infer_func.structured_outputs)
如果加载的模型没有显式签名,可尝试查看模型的Concrete函数:
loaded_model = tf.saved_model.load('./saved_model') # 遍历所有Concrete函数 for attr_name in dir(loaded_model): attr = getattr(loaded_model, attr_name) if isinstance(attr, tf.types.experimental.ConcreteFunction): print(f"函数{attr_name}输入形状:", attr.structured_input_signature) print(f"函数{attr_name}输出形状:", attr.structured_outputs)
内容的提问来源于stack exchange,提问作者Shinobu HUYUGIRI
相关产品推荐
相关产品推荐

