You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于yolov8与easyocr的车牌识别系统:提取文本无法写入CSV

解决YOLOv8+EasyOCR车牌识别系统的CSV写入问题

问题描述

基于YOLOv8和EasyOCR开发的车牌识别系统,可从视频中检测汽车、摩托等车辆并识别车牌文本,但提取的车牌文本无法写入CSV文件,其余功能正常。

原代码

import hydra
import torch
from ultralytics.yolo.engine.predictor import BasePredictor
from ultralytics.yolo.utils import DEFAULT_CONFIG, ROOT, ops
from ultralytics.yolo.utils.checks import check_imgsz
from ultralytics.yolo.utils.plotting import Annotator, colors, save_one_box
import easyocr
import cv2


reader = easyocr.Reader(['en'], gpu=True)

def ocr_image(img, coordinates):
    x, y, w, h = int(coordinates[0]), int(coordinates[1]), int(coordinates[2]), int(coordinates[3])
    img = img[y:h, x:w]

    gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)
    result = reader.readtext(gray)
    text = ""

    for res in result:
        if len(res[1]) >= 8:  # Filter texts with length greater than or equal to 10
            text = res[1]
            break  # Break after finding the first text with length greater than or equal to 10

    return str(text)


# Define the write_to_csv function
def write_to_csv(file_path, data_dict):
    import csv

    with open(file_path, 'w', newline='') as csvfile:
        fieldnames = ['Frame', 'Car ID', 'License Plate']
        writer = csv.DictWriter(csvfile, fieldnames=fieldnames)
        writer.writeheader()

        for frame_num, car_data in data_dict.items():
            for car_id, license_plate in car_data.items():
                writer.writerow({'Frame': frame_num, 'Car ID': car_id, 'License Plate': license_plate})


class DetectionPredictor(BasePredictor):

    def __init__(self, cfg):
        super().__init__(cfg)
        self.all_outputs = []  # Store predictions

    def get_annotator(self, img):
        return Annotator(img, line_width=self.args.line_thickness, example=str(self.model.names))

    def preprocess(self, img):
        img = torch.from_numpy(img).to(self.model.device)
        img = img.half() if self.model.fp16 else img.float()  # uint8 to fp16/32
        img /= 255  # 0 - 255 to 0.0 - 1.0
        return img

    def postprocess(self, preds, img, orig_img):
        preds = ops.non_max_suppression(preds,
                                        self.args.conf,
                                        self.args.iou,
                                        agnostic=self.args.agnostic_nms,
                                        max_det=self.args.max_det)

        for i, pred in enumerate(preds):
            shape = orig_img[i].shape if self.webcam else orig_img.shape
            pred[:, :4] = ops.scale_boxes(img.shape[2:], pred[:, :4], shape).round()

        return preds

    def write_results(self, idx, preds, batch):
        p, im, im0 = batch
        log_string = ""
        if len(im.shape) == 3:
            im = im[None]  # expand for batch dim
        self.seen += 1
        im0 = im0.copy()
        if self.webcam:  # batch_size >= 1
            log_string += f'{idx}: '
            frame = self.dataset.count
        else:
            frame = getattr(self.dataset, 'frame', 0)

        self.data_path = p
        # save_path = str(self.save_dir / p.name)  # im.jpg
        self.txt_path = str(self.save_dir / 'labels' / p.stem) + ('' if self.dataset.mode == 'image' else f'_{frame}')
        log_string += '%gx%g ' % im.shape[2:]  # print string
        self.annotator = self.get_annotator(im0)

        det = preds[idx]
        self.all_outputs.append(det)
        if len(det) == 0:
            return log_string
        for c in det[:, 5].unique():
            n = (det[:, 5] == c).sum()  # detections per class
            log_string += f"{n} {self.model.names[int(c)]}{'s' * (n > 1)}, "
        # write
        gn = torch.tensor(im0.shape)[[1, 0, 1, 0]]  # normalization gain whwh
        for *xyxy, conf, cls in reversed(det):
            if self.args.save_txt:  # Write to file
                xywh = (ops.xyxy2xywh(torch.tensor(xyxy).view(1, 4)) / gn).view(-1).tolist()  # normalized xywh
                line = (cls, *xywh, conf) if self.args.save_conf else (cls, *xywh)  # label format
                with open(f'{self.txt_path}.txt', 'a') as f:
                    f.write(('%g ' * len(line)).rstrip() % line + '\n')

            if self.args.save or self.args.save_crop or self.args.show:  # Add bbox to image
                c = int(cls)  # integer class
                label = None if self.args.hide_labels else (
                    self.model.names[c] if self.args.hide_conf else f'{self.model.names[c]} {conf:.2f}')
                text_ocr = ocr_image(im0,xyxy)
                label = text_ocr              
                self.annotator.box_label(xyxy, label, color=colors(c, True))
            if self.args.save_crop:
                imc = im0.copy()
                save_one_box(xyxy,
                             imc,
                             file=self.save_dir / 'crops' / self.model.model.names[c] / f'{self.data_path.stem}.jpg',
                             BGR=True)

        return log_string



# ... (other imports and code)

@hydra.main(version_base=None, config_path=str(DEFAULT_CONFIG.parent), config_name=DEFAULT_CONFIG.name)
def predict(cfg):
    cfg.model = cfg.model or "yolov8n.pt"
    cfg.imgsz = check_imgsz(cfg.imgsz, min_dim=2)  # check image size
    cfg.source = cfg.source if cfg.source is not None else ROOT / "assets"
    predictor = DetectionPredictor(cfg)
    predictor()

    # Store detected license plate text in a CSV file
    license_plate_data = {}
    for frame_num, frame_data in enumerate(predictor.all_outputs):  # Access stored predictions
        license_plate_data[frame_num] = {}
        if isinstance(frame_data, torch.Tensor):
            frame_data = frame_data.cpu().numpy()  # Convert tensor to numpy array
        for car_info in frame_data:
            car_dict = {
                'class': int(car_info[0]),  # Convert to integer (class index)
                'xyxy': car_info[1:5],
                'confidence': car_info[5],
                'license_plate': {'text': car_info[6].item()} if len(car_info) >= 7 else None
            }

            if car_dict['license_plate'] and len(car_dict['license_plate']['text']) >= 8:
                license_plate_data[frame_num][car_dict['class']] = car_dict

    write_to_csv('./detected_license_plates.csv', license_plate_data)

if __name__ == "__main__":
    predict()

问题根源

  1. 车牌文本未存入检测结果:write_results中计算出的text_ocr仅用于图片标注,未添加到self.all_outputs存储的检测张量里,后续无法读取。
  2. 数据结构错误:predict函数错误假设car_info包含第7个元素(车牌文本),但原始检测结果只有6个元素(xyxy+conf+cls),导致license_plate始终为None。
  3. CSV写入逻辑不匹配:write_to_csv期望car_data的值是车牌文本,但代码中存入的是car_dict字典,结构不符。

修改方案

第一步:修改DetectionPredictor的write_results方法,保存车牌文本

将计算出的车牌文本添加到检测结果张量中:

def write_results(self, idx, preds, batch):
    p, im, im0 = batch
    log_string = ""
    if len(im.shape) == 3:
        im = im[None]  # expand for batch dim
    self.seen += 1
    im0 = im0.copy()
    if self.webcam:  # batch_size >= 1
        log_string += f'{idx}: '
        frame = self.dataset.count
    else:
        frame = getattr(self.dataset, 'frame', 0)

    self.data_path = p
    self.txt_path = str(self.save_dir / 'labels' / p.stem) + ('' if self.dataset.mode == 'image' else f'_{frame}')
    log_string += '%gx%g ' % im.shape[2:]  # print string
    self.annotator = self.get_annotator(im0)

    det = preds[idx]
    license_texts = []
    if len(det) > 0:
        # 先遍历所有检测框获取车牌文本
        for *xyxy, conf, cls in reversed(det):
            text_ocr = ocr_image(im0, xyxy)
            license_texts.append(text_ocr)
        # 将车牌文本转为张量并添加到det末尾
        license_tensor = torch.tensor(license_texts, device=det.device).unsqueeze(1)
        # 反转张量匹配det的原始顺序
        license_tensor = torch.flip(license_tensor, [0])
        det = torch.cat([det, license_tensor], dim=1)
    # 保存包含车牌文本的检测结果
    self.all_outputs.append(det)
    
    if len(det) == 0:
        return log_string
    for c in det[:, 5].unique():
        n = (det[:, 5] == c).sum()  # detections per class
        log_string += f"{n} {self.model.names[int(c)]}{'s' * (n > 1)}, "
    # write
    gn = torch.tensor(im0.shape)[[1, 0, 1, 0]]  # normalization gain whwh
    # 遍历包含车牌文本的det
    for i, (*xyxy, conf, cls, text_ocr) in enumerate(reversed(det)):
        if self.args.save_txt:  # Write to file
            xywh = (ops.xyxy2xywh(torch.tensor(xyxy).view(1, 4)) / gn).view(-1).tolist()  # normalized xywh
            line = (cls, *xywh, conf) if self.args.save_conf else (cls, *xywh)  # label format
            with open(f'{self.txt_path}.txt', 'a') as f:
                f.write(('%g ' * len(line)).rstrip() % line + '\n')

        if self.args.save or self.args.save_crop or self.args.show:  # Add bbox to image
            c = int(cls)  # integer class
            label = None if self.args.hide_labels else (
                self.model.names[c] if self.args.hide_conf else f'{self.model.names[c]} {conf:.2f}')
            label = text_ocr              
            self.annotator.box_label(xyxy, label, color=colors(c, True))
        if self.args.save_crop:
            imc = im0.copy()
            save_one_box(xyxy,
                         imc,
                         file=self.save_dir / 'crops' / self.model.model.names[c] / f'{self.data_path.stem}.jpg',
                         BGR=True)

    return log_string

第二步:修改predict函数的数据处理逻辑

正确提取车牌文本并整理成CSV适配格式:

@hydra.main(version_base=None, config_path=str(DEFAULT_CONFIG.parent), config_name=DEFAULT_CONFIG.name)
def predict(cfg):
    cfg.model = cfg.model or "yolov8n.pt"
    cfg.imgsz = check_imgsz(cfg.imgsz, min_dim=2)  # check image size
    cfg.source = cfg.source if cfg.source is not None else ROOT / "assets"
    predictor = DetectionPredictor(cfg)
    predictor()

    # Store detected license plate text in a CSV file
    license_plate_data = {}
    for frame_num, frame_data in enumerate(predictor.all_outputs):  # Access stored predictions
        license_plate_data[frame_num] = {}
        if isinstance(frame_data, torch.Tensor):
            frame_data = frame_data.cpu().numpy()  # Convert tensor to numpy array
        # 给同帧车辆分配唯一ID,避免同类别重复
        for car_idx, car_info in enumerate(frame_data):
            license_text = car_info[6] if len(car_info) >=7 else ""
            if len(license_text) >=8:
                license_plate_data[frame_num][f"{frame_num}_{car_idx}"] = license_text

    write_to_csv('./detected_license_plates.csv', license_plate_data)

关键说明

  • 用frame_num+car_idx作为车辆唯一ID,避免同帧同类别车辆ID冲突导致数据覆盖。
  • 反转车牌文本张量,保证和检测结果的顺序一致。
  • 保留原代码中仅存储长度≥8的车牌文本的过滤逻辑。

内容的提问来源于stack exchange,提问作者Rajkumar

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.13 18:05:57