aXeleRate/example_scripts/tensorflow_lite/classifier/classifier_file.py

92 lines
3.0 KiB
Python
Executable File

import time
import argparse
import os
import cv2
import numpy as np
from tqdm import tqdm
from cv_utils import init_video_file_capture, decode_classifier, draw_classification, preprocess
from tflite_runtime.interpreter import Interpreter
def load_labels(path):
with open(path, 'r') as f:
return {i: line.strip() for i, line in enumerate(f.read().replace('"','').split(','))}
class NetworkExecutor(object):
def __init__(self, model_file):
self.interpreter = Interpreter(model_file, num_threads=3)
self.interpreter.allocate_tensors()
_, self.input_height, self.input_width, _ = self.interpreter.get_input_details()[0]['shape']
self.tensor_index = self.interpreter.get_input_details()[0]['index']
def get_output_tensors(self):
output_details = self.interpreter.get_output_details()
tensor_indices = []
tensor_list = []
for output in output_details:
tensor = np.squeeze(self.interpreter.get_tensor(output['index']))
tensor_list.append(tensor)
return tensor_list
def run(self, image):
if image.shape[1:2] != (self.input_height, self.input_width):
img = cv2.resize(image, (self.input_width, self.input_height))
img = preprocess(img)
self.interpreter.set_tensor(self.tensor_index, img)
self.interpreter.invoke()
return self.get_output_tensors()
def main(args):
video, video_writer, frame_count = init_video_file_capture(args.file, 'classifier_demo')
if not os.path.exists(args.labels[0]):
labels = args.labels
else:
labels = load_labels(args.labels[0])
frame_num = len(frame_count)
times = []
for _ in tqdm(frame_count, desc='Processing frames'):
frame_present, frame = video.read()
if not frame_present:
continue
start_time = time.time()
results = classification_network.run(frame)
elapsed_ms = (time.time() - start_time) * 1000
classification = decode_classifier(netout = results, top_k = args.top_k)
draw_classification(frame, classification, labels)
times.append(elapsed_ms)
video_writer.write(frame)
print('Finished processing frames')
video.release(), video_writer.release()
print("Average time(ms): ", sum(times)//frame_num)
print("FPS: ", 1000.0 / (sum(times)//frame_num)) # FPS = 1 / time to process loop
if __name__ == "__main__" :
print("OpenCV version: {}".format(cv2. __version__))
parser = argparse.ArgumentParser(formatter_class=argparse.ArgumentDefaultsHelpFormatter)
parser.add_argument('--model', help='File path of .tflite file.', required=True)
parser.add_argument('--labels', nargs="+", help='File path of labels file.', required=True)
parser.add_argument('--top_k', help='How many top results to display', default=3)
parser.add_argument('--file', help='File path of video file', default=None)
args = parser.parse_args()
classification_network = NetworkExecutor(args.model)
main(args)