From d721649148fb200adefcb2d6ccd14068e012fa44 Mon Sep 17 00:00:00 2001 From: wangjingyeye <1025993141@qq.com> Date: Wed, 24 Aug 2022 12:00:51 +0000 Subject: [PATCH 1/7] db++ doc --- doc/doc_ch/algorithm_overview.md | 4 +++- doc/doc_en/algorithm_det_db_en.md | 22 ++++++++++++++++++++-- doc/doc_en/algorithm_overview_en.md | 3 ++- 3 files changed, 25 insertions(+), 4 deletions(-) diff --git a/doc/doc_ch/algorithm_overview.md b/doc/doc_ch/algorithm_overview.md index ef96f6ec12..cda8b7a927 100755 --- a/doc/doc_ch/algorithm_overview.md +++ b/doc/doc_ch/algorithm_overview.md @@ -17,7 +17,7 @@ ### 1.1 文本检测算法 已支持的文本检测算法列表(戳链接获取使用教程): -- [x] [DB](./algorithm_det_db.md) +- [x] [DB与DB++](./algorithm_det_db.md) - [x] [EAST](./algorithm_det_east.md) - [x] [SAST](./algorithm_det_sast.md) - [x] [PSENet](./algorithm_det_psenet.md) @@ -34,6 +34,8 @@ |SAST|ResNet50_vd|91.39%|83.77%|87.42%|[训练模型](https://paddleocr.bj.bcebos.com/dygraph_v2.0/en/det_r50_vd_sast_icdar15_v2.0_train.tar)| |PSE|ResNet50_vd|85.81%|79.53%|82.55%|[训练模型](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/det_r50_vd_pse_v2.0_train.tar)| |PSE|MobileNetV3|82.20%|70.48%|75.89%|[训练模型](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/det_mv3_pse_v2.0_train.tar)| +|DB|ResNet50|86.41%|78.72%|82.38%|[训练模型](https://paddleocr.bj.bcebos.com/dygraph_v2.0/en/det_r50_vd_db_v2.0_train.tar)| +|DB++|ResNet50|90.89%|82.66%|86.58%|[合成数据预训练模型](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/ResNet50_dcn_asf_synthtext_pretrained.pdparams)/[训练模型](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/det_r50_db%2B%2B_icdar15_train.tar)| 在Total-text文本检测公开数据集上,算法效果如下: diff --git a/doc/doc_en/algorithm_det_db_en.md b/doc/doc_en/algorithm_det_db_en.md index f5f333a039..0bd0152ce3 100644 --- a/doc/doc_en/algorithm_det_db_en.md +++ b/doc/doc_en/algorithm_det_db_en.md @@ -1,4 +1,4 @@ -# DB +# DB and DB++ - [1. Introduction](#1) - [2. Environment](#2) @@ -21,13 +21,23 @@ Paper: > Liao, Minghui and Wan, Zhaoyi and Yao, Cong and Chen, Kai and Bai, Xiang > AAAI, 2020 +> [Real-Time Scene Text Detection with Differentiable Binarization and Adaptive Scale Fusion](https://arxiv.org/abs/2202.10304) +> Liao, Minghui and Zou, Zhisheng and Wan, Zhaoyi and Yao, Cong and Bai, Xiang +> TPAMI, 2022 + On the ICDAR2015 dataset, the text detection result is as follows: |Model|Backbone|Configuration|Precision|Recall|Hmean|Download| | --- | --- | --- | --- | --- | --- | --- | |DB|ResNet50_vd|[configs/det/det_r50_vd_db.yml](../../configs/det/det_r50_vd_db.yml)|86.41%|78.72%|82.38%|[trained model](https://paddleocr.bj.bcebos.com/dygraph_v2.0/en/det_r50_vd_db_v2.0_train.tar)| |DB|MobileNetV3|[configs/det/det_mv3_db.yml](../../configs/det/det_mv3_db.yml)|77.29%|73.08%|75.12%|[trained model](https://paddleocr.bj.bcebos.com/dygraph_v2.0/en/det_mv3_db_v2.0_train.tar)| +|DB++|ResNet50|[configs/det/det_r50_db++_ic15.yml](../../configs/det/det_r50_db++_ic15.yml)|90.89%|82.66%|86.58%|[pretrained model](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/ResNet50_dcn_asf_synthtext_pretrained.pdparams)/[trained model](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/det_r50_db%2B%2B_icdar15_train.tar)| +On the TD_TR dataset, the text detection result is as follows: + +|Model|Backbone|Configuration|Precision|Recall|Hmean|Download| +| --- | --- | --- | --- | --- | --- | --- | +|DB++|ResNet50|[configs/det/det_r50_db++_td_tr.yml](../../configs/det/det_r50_db++_td_tr.yml)|92.92%|86.48%|89.58%|[合成数据预训练模型](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/ResNet50_dcn_asf_synthtext_pretrained.pdparams)/[训练模型](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/det_r50_db%2B%2B_td_tr_train.tar)| ## 2. Environment @@ -96,4 +106,12 @@ More deployment schemes supported for DB: pages={11474--11481}, year={2020} } -``` \ No newline at end of file + +@article{liao2022real, + title={Real-Time Scene Text Detection with Differentiable Binarization and Adaptive Scale Fusion}, + author={Liao, Minghui and Zou, Zhisheng and Wan, Zhaoyi and Yao, Cong and Bai, Xiang}, + journal={IEEE Transactions on Pattern Analysis and Machine Intelligence}, + year={2022}, + publisher={IEEE} +} +``` diff --git a/doc/doc_en/algorithm_overview_en.md b/doc/doc_en/algorithm_overview_en.md index bc96cdf235..d8f9428aae 100755 --- a/doc/doc_en/algorithm_overview_en.md +++ b/doc/doc_en/algorithm_overview_en.md @@ -17,7 +17,7 @@ This tutorial lists the OCR algorithms supported by PaddleOCR, as well as the mo ### 1.1 Text Detection Algorithms Supported text detection algorithms (Click the link to get the tutorial): -- [x] [DB](./algorithm_det_db_en.md) +- [x] [DB and DB++](./algorithm_det_db_en.md) - [x] [EAST](./algorithm_det_east_en.md) - [x] [SAST](./algorithm_det_sast_en.md) - [x] [PSENet](./algorithm_det_psenet_en.md) @@ -34,6 +34,7 @@ On the ICDAR2015 dataset, the text detection result is as follows: |SAST|ResNet50_vd|91.39%|83.77%|87.42%|[trained model](https://paddleocr.bj.bcebos.com/dygraph_v2.0/en/det_r50_vd_sast_icdar15_v2.0_train.tar)| |PSE|ResNet50_vd|85.81%|79.53%|82.55%|[trianed model](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/det_r50_vd_pse_v2.0_train.tar)| |PSE|MobileNetV3|82.20%|70.48%|75.89%|[trianed model](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/det_mv3_pse_v2.0_train.tar)| +|DB++|ResNet50|90.89%|82.66%|86.58%|[合成数据预训练模型](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/ResNet50_dcn_asf_synthtext_pretrained.pdparams)/[训练模型](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/det_r50_db%2B%2B_icdar15_train.tar)| On Total-Text dataset, the text detection result is as follows: From 78ec2de9b988c132052082196392c25772ecb1c8 Mon Sep 17 00:00:00 2001 From: wangjingyeye <1025993141@qq.com> Date: Wed, 24 Aug 2022 12:20:40 +0000 Subject: [PATCH 2/7] db++ doc --- doc/doc_ch/algorithm_overview.md | 1 - 1 file changed, 1 deletion(-) diff --git a/doc/doc_ch/algorithm_overview.md b/doc/doc_ch/algorithm_overview.md index cda8b7a927..a12e597063 100755 --- a/doc/doc_ch/algorithm_overview.md +++ b/doc/doc_ch/algorithm_overview.md @@ -34,7 +34,6 @@ |SAST|ResNet50_vd|91.39%|83.77%|87.42%|[训练模型](https://paddleocr.bj.bcebos.com/dygraph_v2.0/en/det_r50_vd_sast_icdar15_v2.0_train.tar)| |PSE|ResNet50_vd|85.81%|79.53%|82.55%|[训练模型](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/det_r50_vd_pse_v2.0_train.tar)| |PSE|MobileNetV3|82.20%|70.48%|75.89%|[训练模型](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/det_mv3_pse_v2.0_train.tar)| -|DB|ResNet50|86.41%|78.72%|82.38%|[训练模型](https://paddleocr.bj.bcebos.com/dygraph_v2.0/en/det_r50_vd_db_v2.0_train.tar)| |DB++|ResNet50|90.89%|82.66%|86.58%|[合成数据预训练模型](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/ResNet50_dcn_asf_synthtext_pretrained.pdparams)/[训练模型](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/det_r50_db%2B%2B_icdar15_train.tar)| 在Total-text文本检测公开数据集上,算法效果如下: From 16c08fae51e9747c108dc07c1175f249ad4fafc6 Mon Sep 17 00:00:00 2001 From: wangjingyeye <1025993141@qq.com> Date: Wed, 24 Aug 2022 12:22:18 +0000 Subject: [PATCH 3/7] db++ doc --- doc/doc_en/algorithm_det_db_en.md | 2 +- doc/doc_en/algorithm_overview_en.md | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/doc/doc_en/algorithm_det_db_en.md b/doc/doc_en/algorithm_det_db_en.md index 0bd0152ce3..b0c332c2ed 100644 --- a/doc/doc_en/algorithm_det_db_en.md +++ b/doc/doc_en/algorithm_det_db_en.md @@ -37,7 +37,7 @@ On the TD_TR dataset, the text detection result is as follows: |Model|Backbone|Configuration|Precision|Recall|Hmean|Download| | --- | --- | --- | --- | --- | --- | --- | -|DB++|ResNet50|[configs/det/det_r50_db++_td_tr.yml](../../configs/det/det_r50_db++_td_tr.yml)|92.92%|86.48%|89.58%|[合成数据预训练模型](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/ResNet50_dcn_asf_synthtext_pretrained.pdparams)/[训练模型](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/det_r50_db%2B%2B_td_tr_train.tar)| +|DB++|ResNet50|[configs/det/det_r50_db++_td_tr.yml](../../configs/det/det_r50_db++_td_tr.yml)|92.92%|86.48%|89.58%|[pretrained model](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/ResNet50_dcn_asf_synthtext_pretrained.pdparams)/[trained model](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/det_r50_db%2B%2B_td_tr_train.tar)| ## 2. Environment diff --git a/doc/doc_en/algorithm_overview_en.md b/doc/doc_en/algorithm_overview_en.md index d8f9428aae..e7e6758541 100755 --- a/doc/doc_en/algorithm_overview_en.md +++ b/doc/doc_en/algorithm_overview_en.md @@ -34,7 +34,7 @@ On the ICDAR2015 dataset, the text detection result is as follows: |SAST|ResNet50_vd|91.39%|83.77%|87.42%|[trained model](https://paddleocr.bj.bcebos.com/dygraph_v2.0/en/det_r50_vd_sast_icdar15_v2.0_train.tar)| |PSE|ResNet50_vd|85.81%|79.53%|82.55%|[trianed model](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/det_r50_vd_pse_v2.0_train.tar)| |PSE|MobileNetV3|82.20%|70.48%|75.89%|[trianed model](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/det_mv3_pse_v2.0_train.tar)| -|DB++|ResNet50|90.89%|82.66%|86.58%|[合成数据预训练模型](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/ResNet50_dcn_asf_synthtext_pretrained.pdparams)/[训练模型](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/det_r50_db%2B%2B_icdar15_train.tar)| +|DB++|ResNet50|90.89%|82.66%|86.58%|[pretrained model](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/ResNet50_dcn_asf_synthtext_pretrained.pdparams)/[trained model](https://paddleocr.bj.bcebos.com/dygraph_v2.1/en_det/det_r50_db%2B%2B_icdar15_train.tar)| On Total-Text dataset, the text detection result is as follows: From 04aaaa748f9f340bca87b75b8446d21abcd96f19 Mon Sep 17 00:00:00 2001 From: wangjingyeye <1025993141@qq.com> Date: Wed, 24 Aug 2022 12:29:30 +0000 Subject: [PATCH 4/7] db++ doc --- doc/doc_en/algorithm_det_db_en.md | 2 +- doc/doc_en/algorithm_overview_en.md | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/doc/doc_en/algorithm_det_db_en.md b/doc/doc_en/algorithm_det_db_en.md index b0c332c2ed..fde344c357 100644 --- a/doc/doc_en/algorithm_det_db_en.md +++ b/doc/doc_en/algorithm_det_db_en.md @@ -1,4 +1,4 @@ -# DB and DB++ +# DB && DB++ - [1. Introduction](#1) - [2. Environment](#2) diff --git a/doc/doc_en/algorithm_overview_en.md b/doc/doc_en/algorithm_overview_en.md index e7e6758541..21c7426d6e 100755 --- a/doc/doc_en/algorithm_overview_en.md +++ b/doc/doc_en/algorithm_overview_en.md @@ -17,7 +17,7 @@ This tutorial lists the OCR algorithms supported by PaddleOCR, as well as the mo ### 1.1 Text Detection Algorithms Supported text detection algorithms (Click the link to get the tutorial): -- [x] [DB and DB++](./algorithm_det_db_en.md) +- [x] [DB && DB++](./algorithm_det_db_en.md) - [x] [EAST](./algorithm_det_east_en.md) - [x] [SAST](./algorithm_det_sast_en.md) - [x] [PSENet](./algorithm_det_psenet_en.md) From 929b4f4557ab3aa1ed9e20a33ada9319ca52542a Mon Sep 17 00:00:00 2001 From: wangjingyeye <1025993141@qq.com> Date: Tue, 30 Aug 2022 05:58:39 +0000 Subject: [PATCH 5/7] update pgnet --- configs/e2e/e2e_r50_vd_pg.yml | 11 +- ppocr/data/imaug/pg_process.py | 182 +++++++++++++++--- ppocr/losses/e2e_pg_loss.py | 9 +- ppocr/modeling/heads/e2e_pg_head.py | 4 +- ppocr/postprocess/pg_postprocess.py | 12 +- .../utils/e2e_utils/extract_textpoint_fast.py | 40 +++- ppocr/utils/e2e_utils/pgnet_pp_utils.py | 13 +- tools/infer_e2e.py | 54 +++++- 8 files changed, 278 insertions(+), 47 deletions(-) diff --git a/configs/e2e/e2e_r50_vd_pg.yml b/configs/e2e/e2e_r50_vd_pg.yml index c4c5226e79..5f1fde6bbc 100644 --- a/configs/e2e/e2e_r50_vd_pg.yml +++ b/configs/e2e/e2e_r50_vd_pg.yml @@ -13,6 +13,7 @@ Global: save_inference_dir: use_visualdl: False infer_img: + infer_visual_type: EN # two mode: EN is for english datasets, CN is for chinese datasets valid_set: totaltext # two mode: totaltext valid curved words, partvgg valid non-curved words save_res_path: ./output/pgnet_r50_vd_totaltext/predicts_pgnet.txt character_dict_path: ppocr/utils/ic15_dict.txt @@ -32,6 +33,7 @@ Architecture: name: PGFPN Head: name: PGHead + tcc_channels: 37 # the length of character dict Loss: name: PGLoss @@ -45,16 +47,18 @@ Optimizer: beta1: 0.9 beta2: 0.999 lr: + name: Cosine learning_rate: 0.001 + warmup_epoch: 50 regularizer: name: 'L2' - factor: 0 - + factor: 0.0001 PostProcess: name: PGPostProcess score_thresh: 0.5 mode: fast # fast or slow two ways + tcc_type: v3 # same as PGProcessTrain: tcc_type Metric: name: E2EMetric @@ -76,9 +80,12 @@ Train: - E2ELabelEncodeTrain: - PGProcessTrain: batch_size: 14 # same as loader: batch_size_per_card + use_resize: True + use_random_crop: False min_crop_size: 24 min_text_size: 4 max_text_size: 512 + tcc_type: v3 # two ways, v2 is original code, v3 is updated code - KeepKeys: keep_keys: [ 'images', 'tcl_maps', 'tcl_label_maps', 'border_maps','direction_maps', 'training_masks', 'label_list', 'pos_list', 'pos_mask' ] # dataloader will return list in this order loader: diff --git a/ppocr/data/imaug/pg_process.py b/ppocr/data/imaug/pg_process.py index 53031064c0..2c8f88215c 100644 --- a/ppocr/data/imaug/pg_process.py +++ b/ppocr/data/imaug/pg_process.py @@ -15,6 +15,8 @@ import math import cv2 import numpy as np +from skimage.morphology._skeletonize import thin +from ppocr.utils.e2e_utils.extract_textpoint_fast import sort_and_expand_with_direction_v2 __all__ = ['PGProcessTrain'] @@ -26,17 +28,24 @@ class PGProcessTrain(object): max_text_nums, tcl_len, batch_size=14, + use_resize=True, + use_random_crop=False, min_crop_size=24, min_text_size=4, max_text_size=512, + tcc_type='v3', **kwargs): self.tcl_len = tcl_len self.max_text_length = max_text_length self.max_text_nums = max_text_nums self.batch_size = batch_size - self.min_crop_size = min_crop_size + if use_random_crop is True: + self.min_crop_size = min_crop_size + self.use_random_crop = use_random_crop self.min_text_size = min_text_size self.max_text_size = max_text_size + self.use_resize = use_resize + self.tcc_type = tcc_type self.Lexicon_Table = self.get_dict(character_dict_path) self.pad_num = len(self.Lexicon_Table) self.img_id = 0 @@ -282,6 +291,95 @@ class PGProcessTrain(object): pos_m[:keep] = 1.0 return pos_l, pos_m + def fit_and_gather_tcl_points_v3(self, + min_area_quad, + poly, + max_h, + max_w, + fixed_point_num=64, + img_id=0, + reference_height=3): + """ + Find the center point of poly as key_points, then fit and gather. + """ + det_mask = np.zeros((int(max_h / self.ds_ratio), + int(max_w / self.ds_ratio))).astype(np.float32) + + # score_big_map + cv2.fillPoly(det_mask, + np.round(poly / self.ds_ratio).astype(np.int32), 1.0) + det_mask = cv2.resize( + det_mask, dsize=None, fx=self.ds_ratio, fy=self.ds_ratio) + det_mask = np.array(det_mask > 1e-3, dtype='float32') + + f_direction = self.f_direction + skeleton_map = thin(det_mask.astype(np.uint8)) + instance_count, instance_label_map = cv2.connectedComponents( + skeleton_map.astype(np.uint8), connectivity=8) + + ys, xs = np.where(instance_label_map == 1) + pos_list = list(zip(ys, xs)) + if len(pos_list) < 3: + return None + pos_list_sorted = sort_and_expand_with_direction_v2( + pos_list, f_direction, det_mask) + + pos_list_sorted = np.array(pos_list_sorted) + length = len(pos_list_sorted) - 1 + insert_num = 0 + for index in range(length): + stride_y = np.abs(pos_list_sorted[index + insert_num][0] - + pos_list_sorted[index + 1 + insert_num][0]) + stride_x = np.abs(pos_list_sorted[index + insert_num][1] - + pos_list_sorted[index + 1 + insert_num][1]) + max_points = int(max(stride_x, stride_y)) + + stride = (pos_list_sorted[index + insert_num] - + pos_list_sorted[index + 1 + insert_num]) / (max_points) + insert_num_temp = max_points - 1 + + for i in range(int(insert_num_temp)): + insert_value = pos_list_sorted[index + insert_num] - (i + 1 + ) * stride + insert_index = index + i + 1 + insert_num + pos_list_sorted = np.insert( + pos_list_sorted, insert_index, insert_value, axis=0) + insert_num += insert_num_temp + + pos_info = np.array(pos_list_sorted).reshape(-1, 2).astype( + np.float32) # xy-> yx + + point_num = len(pos_info) + if point_num > fixed_point_num: + keep_ids = [ + int((point_num * 1.0 / fixed_point_num) * x) + for x in range(fixed_point_num) + ] + pos_info = pos_info[keep_ids, :] + + keep = int(min(len(pos_info), fixed_point_num)) + reference_width = (np.abs(poly[0, 0, 0] - poly[-1, 1, 0]) + + np.abs(poly[0, 3, 0] - poly[-1, 2, 0])) // 2 + if np.random.rand() < 1: + dh = (np.random.rand(keep) - 0.5) * reference_height + offset = np.random.rand() - 0.5 + dw = np.array([[0, offset * reference_width * 0.2]]) + random_float_h = np.array([1, 0]).reshape([1, 2]) * dh.reshape( + [keep, 1]) + random_float_w = dw.repeat(keep, axis=0) + pos_info += random_float_h + pos_info += random_float_w + pos_info[:, 0] = np.clip(pos_info[:, 0], 0, max_h - 1) + pos_info[:, 1] = np.clip(pos_info[:, 1], 0, max_w - 1) + + # padding to fixed length + pos_l = np.zeros((self.tcl_len, 3), dtype=np.int32) + pos_l[:, 0] = np.ones((self.tcl_len, )) * img_id + pos_m = np.zeros((self.tcl_len, 1), dtype=np.float32) + pos_l[:keep, 1:] = np.round(pos_info).astype(np.int32) + pos_m[:keep] = 1.0 + return pos_l, pos_m + def generate_direction_map(self, poly_quads, n_char, direction_map): """ """ @@ -334,6 +432,7 @@ class PGProcessTrain(object): """ Generate polygon. """ + self.ds_ratio = ds_ratio score_map_big = np.zeros( ( h, @@ -384,7 +483,6 @@ class PGProcessTrain(object): text_label = text_strs[poly_idx] text_label = self.prepare_text_label(text_label, self.Lexicon_Table) - text_label_index_list = [[self.Lexicon_Table.index(c_)] for c_ in text_label if c_ in self.Lexicon_Table] @@ -432,14 +530,30 @@ class PGProcessTrain(object): # pos info average_shrink_height = self.calculate_average_height( stcl_quads) - pos_l, pos_m = self.fit_and_gather_tcl_points_v2( - min_area_quad, - poly, - max_h=h, - max_w=w, - fixed_point_num=64, - img_id=self.img_id, - reference_height=average_shrink_height) + + if self.tcc_type == 'v3': + self.f_direction = direction_map[:, :, :-1].copy() + pos_res = self.fit_and_gather_tcl_points_v3( + min_area_quad, + stcl_quads, + max_h=h, + max_w=w, + fixed_point_num=64, + img_id=self.img_id, + reference_height=average_shrink_height) + if pos_res is None: + continue + pos_l, pos_m = pos_res[0], pos_res[1] + + elif self.tcc_type == 'v2': + pos_l, pos_m = self.fit_and_gather_tcl_points_v2( + min_area_quad, + poly, + max_h=h, + max_w=w, + fixed_point_num=64, + img_id=self.img_id, + reference_height=average_shrink_height) label_l = text_label_index_list if len(text_label_index_list) < 2: @@ -770,27 +884,41 @@ class PGProcessTrain(object): text_polys[:, :, 0] *= asp_wx text_polys[:, :, 1] *= asp_hy - h, w, _ = im.shape - if max(h, w) > 2048: - rd_scale = 2048.0 / max(h, w) - im = cv2.resize(im, dsize=None, fx=rd_scale, fy=rd_scale) - text_polys *= rd_scale - h, w, _ = im.shape - if min(h, w) < 16: - return None + if self.use_resize is True: + ori_h, ori_w, _ = im.shape + if max(ori_h, ori_w) < 200: + ratio = 200 / max(ori_h, ori_w) + im = cv2.resize(im, (int(ori_w * ratio), int(ori_h * ratio))) + text_polys[:, :, 0] *= ratio + text_polys[:, :, 1] *= ratio - # no background - im, text_polys, text_tags, hv_tags, text_strs = self.crop_area( - im, - text_polys, - text_tags, - hv_tags, - text_strs, - crop_background=False) + if max(ori_h, ori_w) > 512: + ratio = 512 / max(ori_h, ori_w) + im = cv2.resize(im, (int(ori_w * ratio), int(ori_h * ratio))) + text_polys[:, :, 0] *= ratio + text_polys[:, :, 1] *= ratio + elif self.use_random_crop is True: + h, w, _ = im.shape + if max(h, w) > 2048: + rd_scale = 2048.0 / max(h, w) + im = cv2.resize(im, dsize=None, fx=rd_scale, fy=rd_scale) + text_polys *= rd_scale + h, w, _ = im.shape + if min(h, w) < 16: + return None + + # no background + im, text_polys, text_tags, hv_tags, text_strs = self.crop_area( + im, + text_polys, + text_tags, + hv_tags, + text_strs, + crop_background=False) if text_polys.shape[0] == 0: return None - # # continue for all ignore case + # continue for all ignore case if np.sum((text_tags * 1.0)) >= text_tags.size: return None new_h, new_w, _ = im.shape diff --git a/ppocr/losses/e2e_pg_loss.py b/ppocr/losses/e2e_pg_loss.py index 10a8ed0aa9..aff67b7ce3 100644 --- a/ppocr/losses/e2e_pg_loss.py +++ b/ppocr/losses/e2e_pg_loss.py @@ -89,12 +89,13 @@ class PGLoss(nn.Layer): tcl_pos = paddle.reshape(tcl_pos, [-1, 3]) tcl_pos = paddle.cast(tcl_pos, dtype=int) f_tcl_char = paddle.gather_nd(f_char, tcl_pos) - f_tcl_char = paddle.reshape(f_tcl_char, - [-1, 64, 37]) # len(Lexicon_Table)+1 - f_tcl_char_fg, f_tcl_char_bg = paddle.split(f_tcl_char, [36, 1], axis=2) + f_tcl_char = paddle.reshape( + f_tcl_char, [-1, 64, self.pad_num + 1]) # len(Lexicon_Table)+1 + f_tcl_char_fg, f_tcl_char_bg = paddle.split( + f_tcl_char, [self.pad_num, 1], axis=2) f_tcl_char_bg = f_tcl_char_bg * tcl_mask + (1.0 - tcl_mask) * 20.0 b, c, l = tcl_mask.shape - tcl_mask_fg = paddle.expand(x=tcl_mask, shape=[b, c, 36 * l]) + tcl_mask_fg = paddle.expand(x=tcl_mask, shape=[b, c, self.pad_num * l]) tcl_mask_fg.stop_gradient = True f_tcl_char_fg = f_tcl_char_fg * tcl_mask_fg + (1.0 - tcl_mask_fg) * ( -20.0) diff --git a/ppocr/modeling/heads/e2e_pg_head.py b/ppocr/modeling/heads/e2e_pg_head.py index 274e1cdac5..4bdabeb4d8 100644 --- a/ppocr/modeling/heads/e2e_pg_head.py +++ b/ppocr/modeling/heads/e2e_pg_head.py @@ -66,7 +66,7 @@ class PGHead(nn.Layer): """ """ - def __init__(self, in_channels, **kwargs): + def __init__(self, in_channels, tcc_channels=37, **kwargs): super(PGHead, self).__init__() self.conv_f_score1 = ConvBNLayer( in_channels=in_channels, @@ -178,7 +178,7 @@ class PGHead(nn.Layer): name="conv_f_char{}".format(5)) self.conv3 = nn.Conv2D( in_channels=256, - out_channels=37, + out_channels=tcc_channels, kernel_size=3, stride=1, padding=1, diff --git a/ppocr/postprocess/pg_postprocess.py b/ppocr/postprocess/pg_postprocess.py index 0b1455181f..7f17579b74 100644 --- a/ppocr/postprocess/pg_postprocess.py +++ b/ppocr/postprocess/pg_postprocess.py @@ -31,11 +31,12 @@ class PGPostProcess(object): """ def __init__(self, character_dict_path, valid_set, score_thresh, mode, - **kwargs): + tcc_type, **kwargs): self.character_dict_path = character_dict_path self.valid_set = valid_set self.score_thresh = score_thresh self.mode = mode + self.tcc_type = tcc_type # c++ la-nms is faster, but only support python 3.5 self.is_python35 = False @@ -43,8 +44,13 @@ class PGPostProcess(object): self.is_python35 = True def __call__(self, outs_dict, shape_list): - post = PGNet_PostProcess(self.character_dict_path, self.valid_set, - self.score_thresh, outs_dict, shape_list) + post = PGNet_PostProcess( + self.character_dict_path, + self.valid_set, + self.score_thresh, + outs_dict, + shape_list, + tcc_type=self.tcc_type) if self.mode == 'fast': data = post.pg_postprocess_fast() else: diff --git a/ppocr/utils/e2e_utils/extract_textpoint_fast.py b/ppocr/utils/e2e_utils/extract_textpoint_fast.py index 787cd3017f..fee4145f2d 100644 --- a/ppocr/utils/e2e_utils/extract_textpoint_fast.py +++ b/ppocr/utils/e2e_utils/extract_textpoint_fast.py @@ -88,8 +88,33 @@ def ctc_greedy_decoder(probs_seq, blank=95, keep_blank_in_idxs=True): return dst_str, keep_idx_list -def instance_ctc_greedy_decoder(gather_info, logits_map, pts_num=4): +def instance_ctc_greedy_decoder(gather_info, + logits_map, + pts_num=4, + tcc_type='v3'): _, _, C = logits_map.shape + if tcc_type == 'v3': + insert_num = 0 + gather_info = np.array(gather_info) + length = len(gather_info) - 1 + for index in range(length): + stride_y = np.abs(gather_info[index + insert_num][0] - gather_info[ + index + 1 + insert_num][0]) + stride_x = np.abs(gather_info[index + insert_num][1] - gather_info[ + index + 1 + insert_num][1]) + max_points = int(max(stride_x, stride_y)) + stride = (gather_info[index + insert_num] - + gather_info[index + 1 + insert_num]) / (max_points) + insert_num_temp = max_points - 1 + + for i in range(int(insert_num_temp)): + insert_value = gather_info[index + insert_num] - (i + 1 + ) * stride + insert_index = index + i + 1 + insert_num + gather_info = np.insert( + gather_info, insert_index, insert_value, axis=0) + insert_num += insert_num_temp + gather_info = gather_info.tolist() ys, xs = zip(*gather_info) logits_seq = logits_map[list(ys), list(xs)] probs_seq = logits_seq @@ -104,7 +129,8 @@ def instance_ctc_greedy_decoder(gather_info, logits_map, pts_num=4): def ctc_decoder_for_image(gather_info_list, logits_map, Lexicon_Table, - pts_num=6): + pts_num=6, + tcc_type='v3'): """ CTC decoder using multiple processes. """ @@ -114,7 +140,7 @@ def ctc_decoder_for_image(gather_info_list, if len(gather_info) < pts_num: continue dst_str, xys_list = instance_ctc_greedy_decoder( - gather_info, logits_map, pts_num=pts_num) + gather_info, logits_map, pts_num=pts_num, tcc_type='v3') dst_str_readable = ''.join([Lexicon_Table[idx] for idx in dst_str]) if len(dst_str_readable) < 2: continue @@ -356,7 +382,8 @@ def generate_pivot_list_fast(p_score, p_char_maps, f_direction, Lexicon_Table, - score_thresh=0.5): + score_thresh=0.5, + tcc_type='v3'): """ return center point and end point of TCL instance; filter with the char maps; """ @@ -384,7 +411,10 @@ def generate_pivot_list_fast(p_score, p_char_maps = p_char_maps.transpose([1, 2, 0]) decoded_str, keep_yxs_list = ctc_decoder_for_image( - all_pos_yxs, logits_map=p_char_maps, Lexicon_Table=Lexicon_Table) + all_pos_yxs, + logits_map=p_char_maps, + Lexicon_Table=Lexicon_Table, + tcc_type='v3') return keep_yxs_list, decoded_str diff --git a/ppocr/utils/e2e_utils/pgnet_pp_utils.py b/ppocr/utils/e2e_utils/pgnet_pp_utils.py index a15503c0a8..605ab0e1d2 100644 --- a/ppocr/utils/e2e_utils/pgnet_pp_utils.py +++ b/ppocr/utils/e2e_utils/pgnet_pp_utils.py @@ -28,13 +28,19 @@ from extract_textpoint_fast import generate_pivot_list_fast, restore_poly class PGNet_PostProcess(object): # two different post-process - def __init__(self, character_dict_path, valid_set, score_thresh, outs_dict, - shape_list): + def __init__(self, + character_dict_path, + valid_set, + score_thresh, + outs_dict, + shape_list, + tcc_type='v3'): self.Lexicon_Table = get_dict(character_dict_path) self.valid_set = valid_set self.score_thresh = score_thresh self.outs_dict = outs_dict self.shape_list = shape_list + self.tcc_type = tcc_type def pg_postprocess_fast(self): p_score = self.outs_dict['f_score'] @@ -58,7 +64,8 @@ class PGNet_PostProcess(object): p_char, p_direction, self.Lexicon_Table, - score_thresh=self.score_thresh) + score_thresh=self.score_thresh, + tcc_type=self.tcc_type) poly_list, keep_str_list = restore_poly(instance_yxs_list, seq_strs, p_border, ratio_w, ratio_h, src_w, src_h, self.valid_set) diff --git a/tools/infer_e2e.py b/tools/infer_e2e.py index d3e6b28fca..37fdcbaadc 100755 --- a/tools/infer_e2e.py +++ b/tools/infer_e2e.py @@ -37,6 +37,46 @@ from ppocr.postprocess import build_post_process from ppocr.utils.save_load import load_model from ppocr.utils.utility import get_image_file_list import tools.program as program +from PIL import Image, ImageDraw, ImageFont +import math + + +def draw_e2e_res_for_chinese(image, + boxes, + txts, + config, + img_name, + font_path="./doc/simfang.ttf"): + h, w = image.height, image.width + img_left = image.copy() + img_right = Image.new('RGB', (w, h), (255, 255, 255)) + + import random + + random.seed(0) + draw_left = ImageDraw.Draw(img_left) + draw_right = ImageDraw.Draw(img_right) + for idx, (box, txt) in enumerate(zip(boxes, txts)): + box = np.array(box) + box = [tuple(x) for x in box] + color = (random.randint(0, 255), random.randint(0, 255), + random.randint(0, 255)) + draw_left.polygon(box, fill=color) + draw_right.polygon(box, outline=color) + font = ImageFont.truetype(font_path, 15, encoding="utf-8") + draw_right.text([box[0][0], box[0][1]], txt, fill=(0, 0, 0), font=font) + img_left = Image.blend(image, img_left, 0.5) + img_show = Image.new('RGB', (w * 2, h), (255, 255, 255)) + img_show.paste(img_left, (0, 0, w, h)) + img_show.paste(img_right, (w, 0, w * 2, h)) + + save_e2e_path = os.path.dirname(config['Global'][ + 'save_res_path']) + "/e2e_results/" + if not os.path.exists(save_e2e_path): + os.makedirs(save_e2e_path) + save_path = os.path.join(save_e2e_path, os.path.basename(img_name)) + cv2.imwrite(save_path, np.array(img_show)[:, :, ::-1]) + logger.info("The e2e Image saved in {}".format(save_path)) def draw_e2e_res(dt_boxes, strs, config, img, img_name): @@ -113,7 +153,19 @@ def main(): otstr = file + "\t" + json.dumps(dt_boxes_json) + "\n" fout.write(otstr.encode()) src_img = cv2.imread(file) - draw_e2e_res(points, strs, config, src_img, file) + if global_config['infer_visual_type'] == 'EN': + draw_e2e_res(points, strs, config, src_img, file) + elif global_config['infer_visual_type'] == 'CN': + src_img = Image.fromarray( + cv2.cvtColor(src_img, cv2.COLOR_BGR2RGB)) + draw_e2e_res_for_chinese( + src_img, + points, + strs, + config, + file, + font_path="./doc/fonts/simfang.ttf") + logger.info("success!") From 4c0b08733d41946d4c4817878511f23d2f68feb0 Mon Sep 17 00:00:00 2001 From: wangjingyeye <1025993141@qq.com> Date: Mon, 5 Sep 2022 07:03:16 +0000 Subject: [PATCH 6/7] update pgnet --- configs/e2e/e2e_r50_vd_pg.yml | 6 +++--- ppocr/data/imaug/pg_process.py | 8 ++++---- ppocr/modeling/heads/e2e_pg_head.py | 13 +++++++++++-- ppocr/postprocess/pg_postprocess.py | 6 +++--- ppocr/utils/e2e_utils/extract_textpoint_fast.py | 12 ++++++------ ppocr/utils/e2e_utils/pgnet_pp_utils.py | 6 +++--- 6 files changed, 30 insertions(+), 21 deletions(-) diff --git a/configs/e2e/e2e_r50_vd_pg.yml b/configs/e2e/e2e_r50_vd_pg.yml index 5f1fde6bbc..4adbd2d430 100644 --- a/configs/e2e/e2e_r50_vd_pg.yml +++ b/configs/e2e/e2e_r50_vd_pg.yml @@ -33,7 +33,7 @@ Architecture: name: PGFPN Head: name: PGHead - tcc_channels: 37 # the length of character dict + character_dict_path: ppocr/utils/ic15_dict.txt # the same as Global:character_dict_path Loss: name: PGLoss @@ -58,7 +58,7 @@ PostProcess: name: PGPostProcess score_thresh: 0.5 mode: fast # fast or slow two ways - tcc_type: v3 # same as PGProcessTrain: tcc_type + point_gather_mode: v3 # same as PGProcessTrain: point_gather_mode Metric: name: E2EMetric @@ -85,7 +85,7 @@ Train: min_crop_size: 24 min_text_size: 4 max_text_size: 512 - tcc_type: v3 # two ways, v2 is original code, v3 is updated code + point_gather_mode: v3 # two ways, v2 is original code, v3 is updated code - KeepKeys: keep_keys: [ 'images', 'tcl_maps', 'tcl_label_maps', 'border_maps','direction_maps', 'training_masks', 'label_list', 'pos_list', 'pos_mask' ] # dataloader will return list in this order loader: diff --git a/ppocr/data/imaug/pg_process.py b/ppocr/data/imaug/pg_process.py index 2c8f88215c..622a5a68f1 100644 --- a/ppocr/data/imaug/pg_process.py +++ b/ppocr/data/imaug/pg_process.py @@ -33,7 +33,7 @@ class PGProcessTrain(object): min_crop_size=24, min_text_size=4, max_text_size=512, - tcc_type='v3', + point_gather_mode='v3', **kwargs): self.tcl_len = tcl_len self.max_text_length = max_text_length @@ -45,7 +45,7 @@ class PGProcessTrain(object): self.min_text_size = min_text_size self.max_text_size = max_text_size self.use_resize = use_resize - self.tcc_type = tcc_type + self.point_gather_mode = point_gather_mode self.Lexicon_Table = self.get_dict(character_dict_path) self.pad_num = len(self.Lexicon_Table) self.img_id = 0 @@ -531,7 +531,7 @@ class PGProcessTrain(object): average_shrink_height = self.calculate_average_height( stcl_quads) - if self.tcc_type == 'v3': + if self.point_gather_mode == 'v3': self.f_direction = direction_map[:, :, :-1].copy() pos_res = self.fit_and_gather_tcl_points_v3( min_area_quad, @@ -545,7 +545,7 @@ class PGProcessTrain(object): continue pos_l, pos_m = pos_res[0], pos_res[1] - elif self.tcc_type == 'v2': + elif self.point_gather_mode == 'v2': pos_l, pos_m = self.fit_and_gather_tcl_points_v2( min_area_quad, poly, diff --git a/ppocr/modeling/heads/e2e_pg_head.py b/ppocr/modeling/heads/e2e_pg_head.py index 4bdabeb4d8..514962ef97 100644 --- a/ppocr/modeling/heads/e2e_pg_head.py +++ b/ppocr/modeling/heads/e2e_pg_head.py @@ -66,8 +66,17 @@ class PGHead(nn.Layer): """ """ - def __init__(self, in_channels, tcc_channels=37, **kwargs): + def __init__(self, + in_channels, + character_dict_path='ppocr/utils/ic15_dict.txt', + **kwargs): super(PGHead, self).__init__() + + # get character_length + with open(character_dict_path, "rb") as fin: + lines = fin.readlines() + character_length = len(lines) + 1 + self.conv_f_score1 = ConvBNLayer( in_channels=in_channels, out_channels=64, @@ -178,7 +187,7 @@ class PGHead(nn.Layer): name="conv_f_char{}".format(5)) self.conv3 = nn.Conv2D( in_channels=256, - out_channels=tcc_channels, + out_channels=character_length, kernel_size=3, stride=1, padding=1, diff --git a/ppocr/postprocess/pg_postprocess.py b/ppocr/postprocess/pg_postprocess.py index 7f17579b74..1a52979c14 100644 --- a/ppocr/postprocess/pg_postprocess.py +++ b/ppocr/postprocess/pg_postprocess.py @@ -31,12 +31,12 @@ class PGPostProcess(object): """ def __init__(self, character_dict_path, valid_set, score_thresh, mode, - tcc_type, **kwargs): + point_gather_mode, **kwargs): self.character_dict_path = character_dict_path self.valid_set = valid_set self.score_thresh = score_thresh self.mode = mode - self.tcc_type = tcc_type + self.point_gather_mode = point_gather_mode # c++ la-nms is faster, but only support python 3.5 self.is_python35 = False @@ -50,7 +50,7 @@ class PGPostProcess(object): self.score_thresh, outs_dict, shape_list, - tcc_type=self.tcc_type) + point_gather_mode=self.point_gather_mode) if self.mode == 'fast': data = post.pg_postprocess_fast() else: diff --git a/ppocr/utils/e2e_utils/extract_textpoint_fast.py b/ppocr/utils/e2e_utils/extract_textpoint_fast.py index fee4145f2d..6cf3eb8453 100644 --- a/ppocr/utils/e2e_utils/extract_textpoint_fast.py +++ b/ppocr/utils/e2e_utils/extract_textpoint_fast.py @@ -91,9 +91,9 @@ def ctc_greedy_decoder(probs_seq, blank=95, keep_blank_in_idxs=True): def instance_ctc_greedy_decoder(gather_info, logits_map, pts_num=4, - tcc_type='v3'): + point_gather_mode='v3'): _, _, C = logits_map.shape - if tcc_type == 'v3': + if point_gather_mode == 'v3': insert_num = 0 gather_info = np.array(gather_info) length = len(gather_info) - 1 @@ -130,7 +130,7 @@ def ctc_decoder_for_image(gather_info_list, logits_map, Lexicon_Table, pts_num=6, - tcc_type='v3'): + point_gather_mode='v3'): """ CTC decoder using multiple processes. """ @@ -140,7 +140,7 @@ def ctc_decoder_for_image(gather_info_list, if len(gather_info) < pts_num: continue dst_str, xys_list = instance_ctc_greedy_decoder( - gather_info, logits_map, pts_num=pts_num, tcc_type='v3') + gather_info, logits_map, pts_num=pts_num, point_gather_mode='v3') dst_str_readable = ''.join([Lexicon_Table[idx] for idx in dst_str]) if len(dst_str_readable) < 2: continue @@ -383,7 +383,7 @@ def generate_pivot_list_fast(p_score, f_direction, Lexicon_Table, score_thresh=0.5, - tcc_type='v3'): + point_gather_mode='v3'): """ return center point and end point of TCL instance; filter with the char maps; """ @@ -414,7 +414,7 @@ def generate_pivot_list_fast(p_score, all_pos_yxs, logits_map=p_char_maps, Lexicon_Table=Lexicon_Table, - tcc_type='v3') + point_gather_mode='v3') return keep_yxs_list, decoded_str diff --git a/ppocr/utils/e2e_utils/pgnet_pp_utils.py b/ppocr/utils/e2e_utils/pgnet_pp_utils.py index 605ab0e1d2..12f9dac5f3 100644 --- a/ppocr/utils/e2e_utils/pgnet_pp_utils.py +++ b/ppocr/utils/e2e_utils/pgnet_pp_utils.py @@ -34,13 +34,13 @@ class PGNet_PostProcess(object): score_thresh, outs_dict, shape_list, - tcc_type='v3'): + point_gather_mode='v3'): self.Lexicon_Table = get_dict(character_dict_path) self.valid_set = valid_set self.score_thresh = score_thresh self.outs_dict = outs_dict self.shape_list = shape_list - self.tcc_type = tcc_type + self.point_gather_mode = point_gather_mode def pg_postprocess_fast(self): p_score = self.outs_dict['f_score'] @@ -65,7 +65,7 @@ class PGNet_PostProcess(object): p_direction, self.Lexicon_Table, score_thresh=self.score_thresh, - tcc_type=self.tcc_type) + point_gather_mode=self.point_gather_mode) poly_list, keep_str_list = restore_poly(instance_yxs_list, seq_strs, p_border, ratio_w, ratio_h, src_w, src_h, self.valid_set) From 0fd122b674bc619a4c9b50e6c01fe1e76771b987 Mon Sep 17 00:00:00 2001 From: wangjingyeye <1025993141@qq.com> Date: Mon, 5 Sep 2022 09:06:17 +0000 Subject: [PATCH 7/7] update pgnet --- configs/e2e/e2e_r50_vd_pg.yml | 4 ++-- ppocr/data/imaug/pg_process.py | 6 +++--- ppocr/postprocess/pg_postprocess.py | 9 +++++++-- ppocr/utils/e2e_utils/extract_textpoint_fast.py | 17 +++++++++++------ ppocr/utils/e2e_utils/pgnet_pp_utils.py | 2 +- 5 files changed, 24 insertions(+), 14 deletions(-) diff --git a/configs/e2e/e2e_r50_vd_pg.yml b/configs/e2e/e2e_r50_vd_pg.yml index 4adbd2d430..4642f54486 100644 --- a/configs/e2e/e2e_r50_vd_pg.yml +++ b/configs/e2e/e2e_r50_vd_pg.yml @@ -58,7 +58,7 @@ PostProcess: name: PGPostProcess score_thresh: 0.5 mode: fast # fast or slow two ways - point_gather_mode: v3 # same as PGProcessTrain: point_gather_mode + point_gather_mode: align # same as PGProcessTrain: point_gather_mode Metric: name: E2EMetric @@ -85,7 +85,7 @@ Train: min_crop_size: 24 min_text_size: 4 max_text_size: 512 - point_gather_mode: v3 # two ways, v2 is original code, v3 is updated code + point_gather_mode: align # two mode: align and none, align mode is better than none mode - KeepKeys: keep_keys: [ 'images', 'tcl_maps', 'tcl_label_maps', 'border_maps','direction_maps', 'training_masks', 'label_list', 'pos_list', 'pos_mask' ] # dataloader will return list in this order loader: diff --git a/ppocr/data/imaug/pg_process.py b/ppocr/data/imaug/pg_process.py index 622a5a68f1..f1e5f912b7 100644 --- a/ppocr/data/imaug/pg_process.py +++ b/ppocr/data/imaug/pg_process.py @@ -33,7 +33,7 @@ class PGProcessTrain(object): min_crop_size=24, min_text_size=4, max_text_size=512, - point_gather_mode='v3', + point_gather_mode=None, **kwargs): self.tcl_len = tcl_len self.max_text_length = max_text_length @@ -531,7 +531,7 @@ class PGProcessTrain(object): average_shrink_height = self.calculate_average_height( stcl_quads) - if self.point_gather_mode == 'v3': + if self.point_gather_mode == 'align': self.f_direction = direction_map[:, :, :-1].copy() pos_res = self.fit_and_gather_tcl_points_v3( min_area_quad, @@ -545,7 +545,7 @@ class PGProcessTrain(object): continue pos_l, pos_m = pos_res[0], pos_res[1] - elif self.point_gather_mode == 'v2': + else: pos_l, pos_m = self.fit_and_gather_tcl_points_v2( min_area_quad, poly, diff --git a/ppocr/postprocess/pg_postprocess.py b/ppocr/postprocess/pg_postprocess.py index 1a52979c14..058cf8b907 100644 --- a/ppocr/postprocess/pg_postprocess.py +++ b/ppocr/postprocess/pg_postprocess.py @@ -30,8 +30,13 @@ class PGPostProcess(object): The post process for PGNet. """ - def __init__(self, character_dict_path, valid_set, score_thresh, mode, - point_gather_mode, **kwargs): + def __init__(self, + character_dict_path, + valid_set, + score_thresh, + mode, + point_gather_mode=None, + **kwargs): self.character_dict_path = character_dict_path self.valid_set = valid_set self.score_thresh = score_thresh diff --git a/ppocr/utils/e2e_utils/extract_textpoint_fast.py b/ppocr/utils/e2e_utils/extract_textpoint_fast.py index 6cf3eb8453..a85b8e78ea 100644 --- a/ppocr/utils/e2e_utils/extract_textpoint_fast.py +++ b/ppocr/utils/e2e_utils/extract_textpoint_fast.py @@ -91,9 +91,9 @@ def ctc_greedy_decoder(probs_seq, blank=95, keep_blank_in_idxs=True): def instance_ctc_greedy_decoder(gather_info, logits_map, pts_num=4, - point_gather_mode='v3'): + point_gather_mode=None): _, _, C = logits_map.shape - if point_gather_mode == 'v3': + if point_gather_mode == 'align': insert_num = 0 gather_info = np.array(gather_info) length = len(gather_info) - 1 @@ -115,6 +115,8 @@ def instance_ctc_greedy_decoder(gather_info, gather_info, insert_index, insert_value, axis=0) insert_num += insert_num_temp gather_info = gather_info.tolist() + else: + pass ys, xs = zip(*gather_info) logits_seq = logits_map[list(ys), list(xs)] probs_seq = logits_seq @@ -130,7 +132,7 @@ def ctc_decoder_for_image(gather_info_list, logits_map, Lexicon_Table, pts_num=6, - point_gather_mode='v3'): + point_gather_mode=None): """ CTC decoder using multiple processes. """ @@ -140,7 +142,10 @@ def ctc_decoder_for_image(gather_info_list, if len(gather_info) < pts_num: continue dst_str, xys_list = instance_ctc_greedy_decoder( - gather_info, logits_map, pts_num=pts_num, point_gather_mode='v3') + gather_info, + logits_map, + pts_num=pts_num, + point_gather_mode=point_gather_mode) dst_str_readable = ''.join([Lexicon_Table[idx] for idx in dst_str]) if len(dst_str_readable) < 2: continue @@ -383,7 +388,7 @@ def generate_pivot_list_fast(p_score, f_direction, Lexicon_Table, score_thresh=0.5, - point_gather_mode='v3'): + point_gather_mode=None): """ return center point and end point of TCL instance; filter with the char maps; """ @@ -414,7 +419,7 @@ def generate_pivot_list_fast(p_score, all_pos_yxs, logits_map=p_char_maps, Lexicon_Table=Lexicon_Table, - point_gather_mode='v3') + point_gather_mode=point_gather_mode) return keep_yxs_list, decoded_str diff --git a/ppocr/utils/e2e_utils/pgnet_pp_utils.py b/ppocr/utils/e2e_utils/pgnet_pp_utils.py index 12f9dac5f3..06a766b0e7 100644 --- a/ppocr/utils/e2e_utils/pgnet_pp_utils.py +++ b/ppocr/utils/e2e_utils/pgnet_pp_utils.py @@ -34,7 +34,7 @@ class PGNet_PostProcess(object): score_thresh, outs_dict, shape_list, - point_gather_mode='v3'): + point_gather_mode=None): self.Lexicon_Table = get_dict(character_dict_path) self.valid_set = valid_set self.score_thresh = score_thresh