1. 可视化效果 

这里以dbnet网络训练,icdar2015数据集为例。

from mmcv import Config, imdenormalize
from mmocr.datasets import build_dataset

if __name__ == '__main__':
    import cv2
    import numpy as np
    import torch

    # config = r'D:\code\python\mmocr\configs\textdet\dbnet\dbnet_r50dcnv2_fpnc_1200e_icdar2015.py'
    config = r'D:\code\python\mmocr\configs\textdet\dbnet\dbnet_r50_dcnv2.py'
    cfg = Config.fromfile(config)
    datalayer = build_dataset(cfg.data.train)
    print(len(datalayer))
    for i, data_batch in enumerate(datalayer):

        img_info = data_batch['img_metas']
        img = data_batch["img"]
        gt_shrink = data_batch["gt_shrink"]
        gt_shrink_mask = data_batch["gt_shrink_mask"]
        gt_thr = data_batch["gt_thr"]
        gt_thr_mask = data_batch["gt_thr_mask"]

        img_norm_cfg = img_info.data["img_norm_cfg"]
        img_numpy = img.data.permute(1, 2, 0).detach().cpu().numpy()
        orig_img = imdenormalize(img_numpy, mean=img_norm_cfg["mean"], std=img_norm_cfg["std"], to_bgr=img_norm_cfg["to_rgb"])

        # (h, w ,1)
        gt_shrink = gt_shrink.data.masks.transpose(1, 2, 0)  # 图片上有值的地方是文本索引值,从1开始,像分割mask
        gt_shrink_mask = gt_shrink_mask.data.masks.transpose(1, 2, 0)  # mask. 通过polygons_ignore将不需要地方填0,其他地方都是1. 用于屏蔽不清晰文本
        gt_thr = gt_thr.data.masks.transpose(1, 2, 0)
        gt_thr_mask = gt_thr_mask.data.masks.transpose(1, 2, 0)

        print("img shape: ", img_numpy.shape)
        print("gt_shrink shape: ", gt_shrink.shape)
        print("gt_shrink_mask shape: ", gt_shrink_mask.shape)
        print("gt_thr shape: ", gt_thr.shape)
        print("gt_thr_mask shape: ", gt_thr_mask.shape)

        cv2.namedWindow("orig_img", cv2.WINDOW_NORMAL), cv2.imshow("orig_img", np.uint8(orig_img))
        cv2.namedWindow("gt_shrink", cv2.WINDOW_NORMAL), cv2.imshow("gt_shrink", np.uint8(gt_shrink*255))
        cv2.namedWindow("gt_shrink_mask", cv2.WINDOW_NORMAL), cv2.imshow("gt_shrink_mask", gt_shrink_mask*255)
        cv2.namedWindow("gt_thr", cv2.WINDOW_NORMAL), cv2.imshow("gt_thr", np.uint8(gt_thr*255))
        cv2.namedWindow("gt_thr_mask", cv2.WINDOW_NORMAL), cv2.imshow("gt_thr_mask", gt_thr_mask*255), cv2.waitKey()

更多推荐