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