使用TensorFlow预训练模型实现视频目标检测遇调用错误求助
解决TensorFlow预训练模型调用时的TypeError: 'AutoTrackable' object is not callable问题
问题原因
使用tf.compat.v2.saved_model.load加载预训练模型后,返回的是AutoTrackable对象,并非直接可调用的推理函数,直接调用会触发上述错误。需要获取模型的签名函数来执行推理。
解决步骤
获取模型的签名推理函数
加载模型后,通过signatures['serving_default']获取标准的服务签名函数(大部分TensorFlow预训练检测模型使用这个签名):model = tf.compat.v2.saved_model.load(r'C:\Users\g\Downloads\faster_rcnn_openimages_v4_inception_resnet_v2_1') infer = model.signatures['serving_default'] # 添加这一行调用推理函数执行检测
调用infer函数时,需将输入张量按模型要求的名称传入(Faster RCNN模型的输入名称通常为image_tensor):# 替换原有的outputs = model(inputs) outputs = infer(image_tensor=tf.constant(frame[np.newaxis, ...]))
修改后的完整代码
import numpy as np import tensorflow as tf import cv2 # Load pretrained model model = tf.compat.v2.saved_model.load(r'C:\Users\g\Downloads\faster_rcnn_openimages_v4_inception_resnet_v2_1') infer = model.signatures['serving_default'] # 获取推理签名函数 # Open a video file cap = cv2.VideoCapture(r'C:\Users\g\Desktop\1\training_videos\11.avi') while True: # Read a frame from the video ret, frame = cap.read() if not ret: break # Run the frame through the model inputs = tf.constant(frame[np.newaxis, ...]) outputs = infer(image_tensor=inputs) # 使用签名函数调用 # Get the object detect results boxes, scores, classes, num = outputs["detection_boxes"], outputs["detection_scores"], outputs["detection_classes"], \ outputs["num_detections"] # Draw the detect boxes on the frame for i in range(num.numpy()[0]): if scores.numpy()[0, i] > 0.5: box = boxes.numpy()[0, i] x1, y1, x2, y2 = box[1] * frame.shape[1], box[0] * frame.shape[0], box[3] * frame.shape[1], box[2] * \ frame.shape[0] cv2.rectangle(frame, (int(x1), int(y1)), (int(x2), int(y2)), (0, 0, 255), 2) cv2.putText(frame, '{}'.format(classes.numpy()[0, i]), (int(x1), int(y1)), cv2.FONT_HERSHEY_SIMPLEX, 1, (255, 0, 0), 2) # Show the frame cv2.imshow('Video', frame) if cv2.waitKey(1) & 0xFF == ord('q'): break # Release the video file and close the window cap.release() cv2.destroyAllWindows()
额外提示
如果不确定模型的签名或输入名称,可以通过以下代码查看:
print(model.signatures.keys()) # 查看所有可用签名 print(model.signatures['serving_default'].inputs) # 查看输入张量信息
内容的提问来源于stack exchange,提问作者Berglund
相关产品推荐
相关产品推荐

