如何在Webcam的TensorFlow Lite车牌检测代码中集成OCR?
车牌检测集成OCR解决方案
问题背景
已通过Colab训练车牌检测模型并导出为TFLite文件,现有基于Webcam的实时车牌检测代码可正常检测车牌,但多次尝试集成OCR识别车牌字符未成功,需可行的集成方案。
方案选择:EasyOCR
选用EasyOCR作为OCR工具,它无需额外训练,支持多语言,轻量适配实时场景,可直接处理裁剪后的车牌区域。
实施步骤
安装依赖
运行以下命令安装所需工具:pip install easyocr opencv-python numpy核心修改点
- 导入EasyOCR识别类
- 在检测到车牌 bounding box 后,裁剪出目标区域
- 对车牌区域做灰度转换、二值化、降噪预处理,提升识别准确率
- 调用EasyOCR识别字符,将结果绘制到视频帧上
完整修改后代码
import os import argparse import cv2 import numpy as np import sys import time from threading import Thread import importlib.util import easyocr # Define VideoStream class to handle streaming of video from webcam in separate processing thread class VideoStream: """Camera object that controls video streaming from the webcam""" def __init__(self,resolution=(640,480),framerate=30): # Initialize the camera image stream self.stream = cv2.VideoCapture(0) # 适配默认摄像头,若你的设备摄像头索引为1可改回 ret = self.stream.set(cv2.CAP_PROP_FOURCC, cv2.VideoWriter_fourcc(*'MJPG')) ret = self.stream.set(3,resolution[0]) ret = self.stream.set(4,resolution[1]) # Read first frame from the stream (self.grabbed, self.frame) = self.stream.read() # Variable to control when the camera is stopped self.stopped = False def start(self): # Start the thread that reads frames from the video stream Thread(target=self.update,args=()).start() return self def update(self): # Keep looping indefinitely until the thread is stopped while True: # If the camera is stopped, stop the thread if self.stopped: # Close camera resources self.stream.release() return # Otherwise, grab the next frame from the stream (self.grabbed, self.frame) = self.stream.read() def read(self): # Return the most recent frame return self.frame def stop(self): # Indicate that the camera and thread should be stopped self.stopped = True # Define and parse input arguments parser = argparse.ArgumentParser() parser.add_argument('--modeldir', help='Folder the .tflite file is located in', required=True) parser.add_argument('--graph', help='Name of the .tflite file, if different than detect.tflite', default='detect.tflite') parser.add_argument('--labels', help='Name of the labelmap file, if different than label_map.pbtxt', default='label_map.pbtxt') parser.add_argument('--threshold', help='Minimum confidence threshold for displaying detected objects', default=0.5) parser.add_argument('--resolution', help='Desired webcam resolution in WxH. If the webcam does not support the resolution entered, errors may occur.', default='1280x720') parser.add_argument('--edgetpu', help='Use Coral Edge TPU Accelerator to speed up detection', action='store_true') args = parser.parse_args() MODEL_NAME = args.modeldir GRAPH_NAME = args.graph LABELMAP_NAME = args.labels min_conf_threshold = float(args.threshold) resW, resH = args.resolution.split('x') imW, imH = int(resW), int(resH) use_TPU = args.edgetpu # Initialize EasyOCR reader - 根据车牌语言调整,国内车牌用['ch_sim','en'] reader = easyocr.Reader(['ch_sim','en'], gpu=False) # 有GPU可改为True加速 # Import TensorFlow libraries pkg = importlib.util.find_spec('tflite_runtime') if pkg: from tflite_runtime.interpreter import Interpreter if use_TPU: from tflite_runtime.interpreter import load_delegate else: from tensorflow.lite.python.interpreter import Interpreter if use_TPU: from tensorflow.lite.python.interpreter import load_delegate # If using Edge TPU, assign filename for Edge TPU model if use_TPU: if (GRAPH_NAME == 'detect.tflite'): GRAPH_NAME = 'edgetpu.tflite' # Get path to current working directory CWD_PATH = os.getcwd() # Path to .tflite file PATH_TO_CKPT = os.path.join(CWD_PATH,MODEL_NAME,GRAPH_NAME) # Path to label map file PATH_TO_LABELS = os.path.join(CWD_PATH,MODEL_NAME,LABELMAP_NAME) # Load the label map with open(PATH_TO_LABELS, 'r') as f: labels = [line.strip() for line in f.readlines()] # Fix COCO label map issue if labels[0] == '???': del(labels[0]) # Load the Tensorflow Lite model if use_TPU: interpreter = Interpreter(model_path=PATH_TO_CKPT, experimental_delegates=[load_delegate('libedgetpu.so.1.0')]) print(PATH_TO_CKPT) else: interpreter = Interpreter(model_path=PATH_TO_CKPT) interpreter.allocate_tensors() # Get model details input_details = interpreter.get_input_details() output_details = interpreter.get_output_details() height = input_details[0]['shape'][1] width = input_details[0]['shape'][2] floating_model = (input_details[0]['dtype'] == np.float32) input_mean = 127.5 input_std = 127.5 # Initialize frame rate calculation frame_rate_calc = 1 freq = cv2.getTickFrequency() # Initialize video stream videostream = VideoStream(resolution=(imW,imH),framerate=30).start() time.sleep(1) while True: # Start timer (for calculating frame rate) t1 = cv2.getTickCount() # Grab frame from video stream frame1 = videostream.read() # Acquire frame and resize to expected shape [1xHxWx3] frame = frame1.copy() frame_rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) frame_resized = cv2.resize(frame_rgb, (width, height)) input_data = np.expand_dims(frame_resized, axis=0) # Normalize pixel values if using a floating model if floating_model: input_data = (np.float32(input_data) - input_mean) / input_std # Perform detection interpreter.set_tensor(input_details[0]['index'],input_data) interpreter.invoke() # Retrieve detection results boxes = interpreter.get_tensor(output_details[1]['index'])[0] classes = interpreter.get_tensor(output_details[3]['index'])[0] scores = interpreter.get_tensor(output_details[0]['index'])[0] # Loop over all detections for i in range(len(scores)): if ((scores[i] > min_conf_threshold) and (scores[i] <= 1.0)): # Get bounding box coordinates ymin = int(max(1,(boxes[i][0] * imH))) xmin = int(max(1,(boxes[i][1] * imW))) ymax = int(min(imH,(boxes[i][2] * imH))) xmax = int(min(imW,(boxes[i][3] * imW))) cv2.rectangle(frame, (xmin,ymin), (xmax,ymax), (10, 255, 0), 2) # 裁剪并预处理车牌区域 plate_region = frame[ymin:ymax, xmin:xmax] gray_plate = cv2.cvtColor(plate_region, cv2.COLOR_BGR2GRAY) thresh_plate = cv2.adaptiveThreshold(gray_plate, 255, cv2.ADAPTIVE_THRESH_GAUSSIAN_C, cv2.THRESH_BINARY_INV, 11, 2) thresh_plate = cv2.medianBlur(thresh_plate, 3) # 识别车牌字符 result = reader.readtext(thresh_plate) plate_text = "" if result: for (bbox, text, prob) in result: plate_text += text.strip() # 绘制检测标签和OCR结果 object_name = labels[int(classes[i])] label = '%s: %d%%' % (object_name, int(scores[i]*100)) labelSize, baseLine = cv2.getTextSize(label, cv2.FONT_HERSHEY_SIMPLEX, 0.7, 2) label_ymin = max(ymin, labelSize[1] + 10) cv2.rectangle(frame, (xmin, label_ymin-labelSize[1]-10), (xmin+labelSize[0], label_ymin+baseLine-10), (255, 255, 255), cv2.FILLED) cv2.putText(frame, label, (xmin, label_ymin-7), cv2.FONT_HERSHEY_SIMPLEX, 0.7, (0, 0, 0), 2) if plate_text: ocr_label = 'Plate: %s' % plate_text ocr_labelSize, ocr_baseLine = cv2.getTextSize(ocr_label, cv2.FONT_HERSHEY_SIMPLEX, 0.7, 2) ocr_ymin = max(label_ymin, ocr_labelSize[1] + 20) cv2.rectangle(frame, (xmin, ocr_ymin-ocr_labelSize[1]-10), (xmin+ocr_labelSize[0], ocr_ymin+ocr_baseLine-10), (255, 255, 255), cv2.FILLED) cv2.putText(frame, ocr_label, (xmin, ocr_ymin-7), cv2.FONT_HERSHEY_SIMPLEX, 0.7, (0, 0, 255), 2) # Draw framerate cv2.putText(frame,'FPS: {0:.2f}'.format(frame_rate_calc),(30,50),cv2.FONT_HERSHEY_SIMPLEX,1,(255,255,0),2,cv2.LINE_AA) # Display frame cv2.imshow('Object detector with OCR', frame) # Calculate framerate t2 = cv2.getTickCount() time1 = (t2-t1)/freq frame_rate_calc= 1/time1 # Exit on 'q' press if cv2.waitKey(1) == ord('q'): break # Clean up cv2.destroyAllWindows() videostream.stop()
注意事项
- 摄像头索引:若默认摄像头无法启动,将
VideoStream类中的cv2.VideoCapture(0)改为cv2.VideoCapture(1) - GPU加速:有可用GPU时,将EasyOCR初始化的
gpu=False改为gpu=True,可大幅提升识别速度 - 语言配置:根据车牌实际语言调整EasyOCR的参数,纯英文车牌可改为
['en']
内容的提问来源于stack exchange,提问作者Abdulrahman Q
相关产品推荐
相关产品推荐

