This commit is contained in:
121
ppocr/postprocess/__init__.py
Normal file
121
ppocr/postprocess/__init__.py
Normal file
@@ -0,0 +1,121 @@
|
||||
# copyright (c) 2020 PaddlePaddle Authors. All Rights Reserve.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from __future__ import absolute_import
|
||||
from __future__ import division
|
||||
from __future__ import print_function
|
||||
from __future__ import unicode_literals
|
||||
|
||||
import os
|
||||
import copy
|
||||
|
||||
__all__ = ["build_post_process"]
|
||||
|
||||
from .db_postprocess import DBPostProcess, DistillationDBPostProcess
|
||||
from .east_postprocess import EASTPostProcess
|
||||
from .sast_postprocess import SASTPostProcess
|
||||
from .fce_postprocess import FCEPostProcess
|
||||
from .rec_postprocess import (
|
||||
CTCLabelDecode,
|
||||
AttnLabelDecode,
|
||||
SRNLabelDecode,
|
||||
DistillationCTCLabelDecode,
|
||||
NRTRLabelDecode,
|
||||
SARLabelDecode,
|
||||
SEEDLabelDecode,
|
||||
PRENLabelDecode,
|
||||
ViTSTRLabelDecode,
|
||||
ABINetLabelDecode,
|
||||
SPINLabelDecode,
|
||||
VLLabelDecode,
|
||||
RFLLabelDecode,
|
||||
SATRNLabelDecode,
|
||||
ParseQLabelDecode,
|
||||
CPPDLabelDecode,
|
||||
LaTeXOCRDecode,
|
||||
UniMERNetDecode,
|
||||
)
|
||||
from .cls_postprocess import ClsPostProcess
|
||||
from .pg_postprocess import PGPostProcess
|
||||
from .vqa_token_ser_layoutlm_postprocess import (
|
||||
VQASerTokenLayoutLMPostProcess,
|
||||
DistillationSerPostProcess,
|
||||
)
|
||||
from .vqa_token_re_layoutlm_postprocess import (
|
||||
VQAReTokenLayoutLMPostProcess,
|
||||
DistillationRePostProcess,
|
||||
)
|
||||
from .table_postprocess import TableMasterLabelDecode, TableLabelDecode
|
||||
from .picodet_postprocess import PicoDetPostProcess
|
||||
from .ct_postprocess import CTPostProcess
|
||||
from .drrg_postprocess import DRRGPostprocess
|
||||
from .rec_postprocess import CANLabelDecode
|
||||
|
||||
|
||||
def build_post_process(config, global_config=None):
|
||||
support_dict = [
|
||||
"DBPostProcess",
|
||||
"EASTPostProcess",
|
||||
"SASTPostProcess",
|
||||
"FCEPostProcess",
|
||||
"CTCLabelDecode",
|
||||
"AttnLabelDecode",
|
||||
"ClsPostProcess",
|
||||
"SRNLabelDecode",
|
||||
"PGPostProcess",
|
||||
"DistillationCTCLabelDecode",
|
||||
"TableLabelDecode",
|
||||
"DistillationDBPostProcess",
|
||||
"NRTRLabelDecode",
|
||||
"SARLabelDecode",
|
||||
"SEEDLabelDecode",
|
||||
"VQASerTokenLayoutLMPostProcess",
|
||||
"VQAReTokenLayoutLMPostProcess",
|
||||
"PRENLabelDecode",
|
||||
"DistillationSARLabelDecode",
|
||||
"ViTSTRLabelDecode",
|
||||
"ABINetLabelDecode",
|
||||
"TableMasterLabelDecode",
|
||||
"SPINLabelDecode",
|
||||
"DistillationSerPostProcess",
|
||||
"DistillationRePostProcess",
|
||||
"VLLabelDecode",
|
||||
"PicoDetPostProcess",
|
||||
"CTPostProcess",
|
||||
"RFLLabelDecode",
|
||||
"DRRGPostprocess",
|
||||
"CANLabelDecode",
|
||||
"SATRNLabelDecode",
|
||||
"ParseQLabelDecode",
|
||||
"CPPDLabelDecode",
|
||||
"LaTeXOCRDecode",
|
||||
"UniMERNetDecode",
|
||||
]
|
||||
|
||||
if config["name"] == "PSEPostProcess":
|
||||
from .pse_postprocess import PSEPostProcess
|
||||
|
||||
support_dict.append("PSEPostProcess")
|
||||
|
||||
config = copy.deepcopy(config)
|
||||
module_name = config.pop("name")
|
||||
if module_name == "None":
|
||||
return
|
||||
if global_config is not None:
|
||||
config.update(global_config)
|
||||
assert module_name in support_dict, Exception(
|
||||
"post process only support {}".format(support_dict)
|
||||
)
|
||||
module_class = eval(module_name)(**config)
|
||||
return module_class
|
||||
43
ppocr/postprocess/cls_postprocess.py
Normal file
43
ppocr/postprocess/cls_postprocess.py
Normal file
@@ -0,0 +1,43 @@
|
||||
# copyright (c) 2020 PaddlePaddle Authors. All Rights Reserve.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
import paddle
|
||||
|
||||
|
||||
class ClsPostProcess(object):
|
||||
"""Convert between text-label and text-index"""
|
||||
|
||||
def __init__(self, label_list=None, key=None, **kwargs):
|
||||
super(ClsPostProcess, self).__init__()
|
||||
self.label_list = label_list
|
||||
self.key = key
|
||||
|
||||
def __call__(self, preds, label=None, *args, **kwargs):
|
||||
if self.key is not None:
|
||||
preds = preds[self.key]
|
||||
|
||||
label_list = self.label_list
|
||||
if label_list is None:
|
||||
label_list = {idx: idx for idx in range(preds.shape[-1])}
|
||||
|
||||
if isinstance(preds, paddle.Tensor):
|
||||
preds = preds.numpy()
|
||||
|
||||
pred_idxs = preds.argmax(axis=1)
|
||||
decode_out = [
|
||||
(label_list[idx], preds[i, idx]) for i, idx in enumerate(pred_idxs)
|
||||
]
|
||||
if label is None:
|
||||
return decode_out
|
||||
label = [(label_list[idx], 1.0) for idx in label]
|
||||
return decode_out, label
|
||||
158
ppocr/postprocess/ct_postprocess.py
Executable file
158
ppocr/postprocess/ct_postprocess.py
Executable file
@@ -0,0 +1,158 @@
|
||||
# Copyright (c) 2020 PaddlePaddle Authors. All Rights Reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""
|
||||
This code is referred from:
|
||||
https://github.com/shengtao96/CentripetalText/blob/main/test.py
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import
|
||||
from __future__ import division
|
||||
from __future__ import print_function
|
||||
|
||||
import os
|
||||
import os.path as osp
|
||||
import numpy as np
|
||||
import cv2
|
||||
import paddle
|
||||
import pyclipper
|
||||
|
||||
|
||||
class CTPostProcess(object):
|
||||
"""
|
||||
The post process for Centripetal Text (CT).
|
||||
"""
|
||||
|
||||
def __init__(self, min_score=0.88, min_area=16, box_type="poly", **kwargs):
|
||||
self.min_score = min_score
|
||||
self.min_area = min_area
|
||||
self.box_type = box_type
|
||||
|
||||
self.coord = np.zeros((2, 300, 300), dtype=np.int32)
|
||||
for i in range(300):
|
||||
for j in range(300):
|
||||
self.coord[0, i, j] = j
|
||||
self.coord[1, i, j] = i
|
||||
|
||||
def __call__(self, preds, batch):
|
||||
outs = preds["maps"]
|
||||
out_scores = preds["score"]
|
||||
|
||||
if isinstance(outs, paddle.Tensor):
|
||||
outs = outs.numpy()
|
||||
if isinstance(out_scores, paddle.Tensor):
|
||||
out_scores = out_scores.numpy()
|
||||
|
||||
batch_size = outs.shape[0]
|
||||
boxes_batch = []
|
||||
for idx in range(batch_size):
|
||||
bboxes = []
|
||||
scores = []
|
||||
|
||||
img_shape = batch[idx]
|
||||
|
||||
org_img_size = img_shape[:3]
|
||||
img_shape = img_shape[3:]
|
||||
img_size = img_shape[:2]
|
||||
|
||||
out = np.expand_dims(outs[idx], axis=0)
|
||||
outputs = dict()
|
||||
|
||||
score = np.expand_dims(out_scores[idx], axis=0)
|
||||
|
||||
kernel = out[:, 0, :, :] > 0.2
|
||||
loc = out[:, 1:, :, :].astype("float32")
|
||||
|
||||
score = score[0].astype(np.float32)
|
||||
kernel = kernel[0].astype(np.uint8)
|
||||
loc = loc[0].astype(np.float32)
|
||||
|
||||
label_num, label_kernel = cv2.connectedComponents(kernel, connectivity=4)
|
||||
|
||||
for i in range(1, label_num):
|
||||
ind = label_kernel == i
|
||||
if ind.sum() < 10: # pixel number less than 10, treated as background
|
||||
label_kernel[ind] = 0
|
||||
|
||||
label = np.zeros_like(label_kernel)
|
||||
h, w = label_kernel.shape
|
||||
pixels = self.coord[:, :h, :w].reshape(2, -1)
|
||||
points = pixels.transpose([1, 0]).astype(np.float32)
|
||||
|
||||
off_points = (points + 10.0 / 4.0 * loc[:, pixels[1], pixels[0]].T).astype(
|
||||
np.int32
|
||||
)
|
||||
off_points[:, 0] = np.clip(off_points[:, 0], 0, label.shape[1] - 1)
|
||||
off_points[:, 1] = np.clip(off_points[:, 1], 0, label.shape[0] - 1)
|
||||
|
||||
label[pixels[1], pixels[0]] = label_kernel[
|
||||
off_points[:, 1], off_points[:, 0]
|
||||
]
|
||||
label[label_kernel > 0] = label_kernel[label_kernel > 0]
|
||||
|
||||
score_pocket = [0.0]
|
||||
for i in range(1, label_num):
|
||||
ind = label_kernel == i
|
||||
if ind.sum() == 0:
|
||||
score_pocket.append(0.0)
|
||||
continue
|
||||
score_i = np.mean(score[ind])
|
||||
score_pocket.append(score_i)
|
||||
|
||||
label_num = np.max(label) + 1
|
||||
label = cv2.resize(
|
||||
label, (img_size[1], img_size[0]), interpolation=cv2.INTER_NEAREST
|
||||
)
|
||||
|
||||
scale = (
|
||||
float(org_img_size[1]) / float(img_size[1]),
|
||||
float(org_img_size[0]) / float(img_size[0]),
|
||||
)
|
||||
|
||||
for i in range(1, label_num):
|
||||
ind = label == i
|
||||
points = np.array(np.where(ind)).transpose((1, 0))
|
||||
|
||||
if points.shape[0] < self.min_area:
|
||||
continue
|
||||
|
||||
score_i = score_pocket[i]
|
||||
if score_i < self.min_score:
|
||||
continue
|
||||
|
||||
if self.box_type == "rect":
|
||||
rect = cv2.minAreaRect(points[:, ::-1])
|
||||
bbox = cv2.boxPoints(rect) * scale
|
||||
z = bbox.mean(0)
|
||||
bbox = z + (bbox - z) * 0.85
|
||||
elif self.box_type == "poly":
|
||||
binary = np.zeros(label.shape, dtype="uint8")
|
||||
binary[ind] = 1
|
||||
try:
|
||||
_, contours, _ = cv2.findContours(
|
||||
binary, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE
|
||||
)
|
||||
except BaseException:
|
||||
contours, _ = cv2.findContours(
|
||||
binary, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE
|
||||
)
|
||||
|
||||
bbox = contours[0] * scale
|
||||
|
||||
bbox = bbox.astype("int32")
|
||||
bboxes.append(bbox.reshape(-1, 2))
|
||||
scores.append(score_i)
|
||||
|
||||
boxes_batch.append({"points": bboxes})
|
||||
|
||||
return boxes_batch
|
||||
289
ppocr/postprocess/db_postprocess.py
Executable file
289
ppocr/postprocess/db_postprocess.py
Executable file
@@ -0,0 +1,289 @@
|
||||
# Copyright (c) 2020 PaddlePaddle Authors. All Rights Reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""
|
||||
This code is referred from:
|
||||
https://github.com/WenmuZhou/DBNet.pytorch/blob/master/post_processing/seg_detector_representer.py
|
||||
"""
|
||||
from __future__ import absolute_import
|
||||
from __future__ import division
|
||||
from __future__ import print_function
|
||||
|
||||
import numpy as np
|
||||
import cv2
|
||||
import paddle
|
||||
from shapely.geometry import Polygon
|
||||
import pyclipper
|
||||
|
||||
|
||||
class DBPostProcess(object):
|
||||
"""
|
||||
The post process for Differentiable Binarization (DB).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
thresh=0.3,
|
||||
box_thresh=0.7,
|
||||
max_candidates=1000,
|
||||
unclip_ratio=2.0,
|
||||
use_dilation=False,
|
||||
score_mode="fast",
|
||||
box_type="quad",
|
||||
**kwargs,
|
||||
):
|
||||
self.thresh = thresh
|
||||
self.box_thresh = box_thresh
|
||||
self.max_candidates = max_candidates
|
||||
self.unclip_ratio = unclip_ratio
|
||||
self.min_size = 3
|
||||
self.score_mode = score_mode
|
||||
self.box_type = box_type
|
||||
assert score_mode in [
|
||||
"slow",
|
||||
"fast",
|
||||
], "Score mode must be in [slow, fast] but got: {}".format(score_mode)
|
||||
|
||||
self.dilation_kernel = None if not use_dilation else np.array([[1, 1], [1, 1]])
|
||||
|
||||
def polygons_from_bitmap(self, pred, _bitmap, dest_width, dest_height):
|
||||
"""
|
||||
_bitmap: single map with shape (1, H, W),
|
||||
whose values are binarized as {0, 1}
|
||||
"""
|
||||
|
||||
bitmap = _bitmap
|
||||
height, width = bitmap.shape
|
||||
|
||||
boxes = []
|
||||
scores = []
|
||||
|
||||
contours, _ = cv2.findContours(
|
||||
(bitmap * 255).astype(np.uint8), cv2.RETR_LIST, cv2.CHAIN_APPROX_SIMPLE
|
||||
)
|
||||
|
||||
for contour in contours[: self.max_candidates]:
|
||||
epsilon = 0.002 * cv2.arcLength(contour, True)
|
||||
approx = cv2.approxPolyDP(contour, epsilon, True)
|
||||
points = approx.reshape((-1, 2))
|
||||
if points.shape[0] < 4:
|
||||
continue
|
||||
|
||||
score = self.box_score_fast(pred, points.reshape(-1, 2))
|
||||
if self.box_thresh > score:
|
||||
continue
|
||||
|
||||
if points.shape[0] > 2:
|
||||
box = self.unclip(points, self.unclip_ratio)
|
||||
if len(box) > 1:
|
||||
continue
|
||||
else:
|
||||
continue
|
||||
box = np.array(box).reshape(-1, 2)
|
||||
if len(box) == 0:
|
||||
continue
|
||||
|
||||
_, sside = self.get_mini_boxes(box.reshape((-1, 1, 2)))
|
||||
if sside < self.min_size + 2:
|
||||
continue
|
||||
|
||||
box = np.array(box)
|
||||
box[:, 0] = np.clip(np.round(box[:, 0] / width * dest_width), 0, dest_width)
|
||||
box[:, 1] = np.clip(
|
||||
np.round(box[:, 1] / height * dest_height), 0, dest_height
|
||||
)
|
||||
boxes.append(box.tolist())
|
||||
scores.append(score)
|
||||
return boxes, scores
|
||||
|
||||
def boxes_from_bitmap(self, pred, _bitmap, dest_width, dest_height):
|
||||
"""
|
||||
_bitmap: single map with shape (1, H, W),
|
||||
whose values are binarized as {0, 1}
|
||||
"""
|
||||
|
||||
bitmap = _bitmap
|
||||
height, width = bitmap.shape
|
||||
|
||||
outs = cv2.findContours(
|
||||
(bitmap * 255).astype(np.uint8), cv2.RETR_LIST, cv2.CHAIN_APPROX_SIMPLE
|
||||
)
|
||||
if len(outs) == 3:
|
||||
img, contours, _ = outs[0], outs[1], outs[2]
|
||||
elif len(outs) == 2:
|
||||
contours, _ = outs[0], outs[1]
|
||||
|
||||
num_contours = min(len(contours), self.max_candidates)
|
||||
|
||||
boxes = []
|
||||
scores = []
|
||||
for index in range(num_contours):
|
||||
contour = contours[index]
|
||||
points, sside = self.get_mini_boxes(contour)
|
||||
if sside < self.min_size:
|
||||
continue
|
||||
points = np.array(points)
|
||||
if self.score_mode == "fast":
|
||||
score = self.box_score_fast(pred, points.reshape(-1, 2))
|
||||
else:
|
||||
score = self.box_score_slow(pred, contour)
|
||||
if self.box_thresh > score:
|
||||
continue
|
||||
|
||||
box = self.unclip(points, self.unclip_ratio)
|
||||
if len(box) > 1:
|
||||
continue
|
||||
box = np.array(box).reshape(-1, 1, 2)
|
||||
box, sside = self.get_mini_boxes(box)
|
||||
if sside < self.min_size + 2:
|
||||
continue
|
||||
box = np.array(box)
|
||||
|
||||
box[:, 0] = np.clip(np.round(box[:, 0] / width * dest_width), 0, dest_width)
|
||||
box[:, 1] = np.clip(
|
||||
np.round(box[:, 1] / height * dest_height), 0, dest_height
|
||||
)
|
||||
boxes.append(box.astype("int32"))
|
||||
scores.append(score)
|
||||
return np.array(boxes, dtype="int32"), scores
|
||||
|
||||
def unclip(self, box, unclip_ratio):
|
||||
poly = Polygon(box)
|
||||
distance = poly.area * unclip_ratio / poly.length
|
||||
offset = pyclipper.PyclipperOffset()
|
||||
offset.AddPath(box, pyclipper.JT_ROUND, pyclipper.ET_CLOSEDPOLYGON)
|
||||
expanded = offset.Execute(distance)
|
||||
return expanded
|
||||
|
||||
def get_mini_boxes(self, contour):
|
||||
bounding_box = cv2.minAreaRect(contour)
|
||||
points = sorted(list(cv2.boxPoints(bounding_box)), key=lambda x: x[0])
|
||||
|
||||
index_1, index_2, index_3, index_4 = 0, 1, 2, 3
|
||||
if points[1][1] > points[0][1]:
|
||||
index_1 = 0
|
||||
index_4 = 1
|
||||
else:
|
||||
index_1 = 1
|
||||
index_4 = 0
|
||||
if points[3][1] > points[2][1]:
|
||||
index_2 = 2
|
||||
index_3 = 3
|
||||
else:
|
||||
index_2 = 3
|
||||
index_3 = 2
|
||||
|
||||
box = [points[index_1], points[index_2], points[index_3], points[index_4]]
|
||||
return box, min(bounding_box[1])
|
||||
|
||||
def box_score_fast(self, bitmap, _box):
|
||||
"""
|
||||
box_score_fast: use bbox mean score as the mean score
|
||||
"""
|
||||
h, w = bitmap.shape[:2]
|
||||
box = _box.copy()
|
||||
xmin = np.clip(np.floor(box[:, 0].min()).astype("int32"), 0, w - 1)
|
||||
xmax = np.clip(np.ceil(box[:, 0].max()).astype("int32"), 0, w - 1)
|
||||
ymin = np.clip(np.floor(box[:, 1].min()).astype("int32"), 0, h - 1)
|
||||
ymax = np.clip(np.ceil(box[:, 1].max()).astype("int32"), 0, h - 1)
|
||||
|
||||
mask = np.zeros((ymax - ymin + 1, xmax - xmin + 1), dtype=np.uint8)
|
||||
box[:, 0] = box[:, 0] - xmin
|
||||
box[:, 1] = box[:, 1] - ymin
|
||||
cv2.fillPoly(mask, box.reshape(1, -1, 2).astype("int32"), 1)
|
||||
return cv2.mean(bitmap[ymin : ymax + 1, xmin : xmax + 1], mask)[0]
|
||||
|
||||
def box_score_slow(self, bitmap, contour):
|
||||
"""
|
||||
box_score_slow: use polyon mean score as the mean score
|
||||
"""
|
||||
h, w = bitmap.shape[:2]
|
||||
contour = contour.copy()
|
||||
contour = np.reshape(contour, (-1, 2))
|
||||
|
||||
xmin = np.clip(np.min(contour[:, 0]), 0, w - 1)
|
||||
xmax = np.clip(np.max(contour[:, 0]), 0, w - 1)
|
||||
ymin = np.clip(np.min(contour[:, 1]), 0, h - 1)
|
||||
ymax = np.clip(np.max(contour[:, 1]), 0, h - 1)
|
||||
|
||||
mask = np.zeros((ymax - ymin + 1, xmax - xmin + 1), dtype=np.uint8)
|
||||
|
||||
contour[:, 0] = contour[:, 0] - xmin
|
||||
contour[:, 1] = contour[:, 1] - ymin
|
||||
|
||||
cv2.fillPoly(mask, contour.reshape(1, -1, 2).astype("int32"), 1)
|
||||
return cv2.mean(bitmap[ymin : ymax + 1, xmin : xmax + 1], mask)[0]
|
||||
|
||||
def __call__(self, outs_dict, shape_list):
|
||||
pred = outs_dict["maps"]
|
||||
if isinstance(pred, paddle.Tensor):
|
||||
pred = pred.numpy()
|
||||
pred = pred[:, 0, :, :]
|
||||
segmentation = pred > self.thresh
|
||||
|
||||
boxes_batch = []
|
||||
for batch_index in range(pred.shape[0]):
|
||||
src_h, src_w, ratio_h, ratio_w = shape_list[batch_index]
|
||||
if self.dilation_kernel is not None:
|
||||
mask = cv2.dilate(
|
||||
np.array(segmentation[batch_index]).astype(np.uint8),
|
||||
self.dilation_kernel,
|
||||
)
|
||||
else:
|
||||
mask = segmentation[batch_index]
|
||||
if self.box_type == "poly":
|
||||
boxes, scores = self.polygons_from_bitmap(
|
||||
pred[batch_index], mask, src_w, src_h
|
||||
)
|
||||
elif self.box_type == "quad":
|
||||
boxes, scores = self.boxes_from_bitmap(
|
||||
pred[batch_index], mask, src_w, src_h
|
||||
)
|
||||
else:
|
||||
raise ValueError("box_type can only be one of ['quad', 'poly']")
|
||||
|
||||
boxes_batch.append({"points": boxes})
|
||||
return boxes_batch
|
||||
|
||||
|
||||
class DistillationDBPostProcess(object):
|
||||
def __init__(
|
||||
self,
|
||||
model_name=["student"],
|
||||
key=None,
|
||||
thresh=0.3,
|
||||
box_thresh=0.6,
|
||||
max_candidates=1000,
|
||||
unclip_ratio=1.5,
|
||||
use_dilation=False,
|
||||
score_mode="fast",
|
||||
box_type="quad",
|
||||
**kwargs,
|
||||
):
|
||||
self.model_name = model_name
|
||||
self.key = key
|
||||
self.post_process = DBPostProcess(
|
||||
thresh=thresh,
|
||||
box_thresh=box_thresh,
|
||||
max_candidates=max_candidates,
|
||||
unclip_ratio=unclip_ratio,
|
||||
use_dilation=use_dilation,
|
||||
score_mode=score_mode,
|
||||
box_type=box_type,
|
||||
)
|
||||
|
||||
def __call__(self, predicts, shape_list):
|
||||
results = {}
|
||||
for k in self.model_name:
|
||||
results[k] = self.post_process(predicts[k], shape_list=shape_list)
|
||||
return results
|
||||
338
ppocr/postprocess/drrg_postprocess.py
Normal file
338
ppocr/postprocess/drrg_postprocess.py
Normal file
@@ -0,0 +1,338 @@
|
||||
# copyright (c) 2022 PaddlePaddle Authors. All Rights Reserve.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""
|
||||
This code is refer from:
|
||||
https://github.com/open-mmlab/mmocr/blob/main/mmocr/models/textdet/postprocess/drrg_postprocessor.py
|
||||
"""
|
||||
|
||||
import functools
|
||||
import operator
|
||||
|
||||
import numpy as np
|
||||
import paddle
|
||||
from numpy.linalg import norm
|
||||
import cv2
|
||||
|
||||
|
||||
class Node:
|
||||
def __init__(self, ind):
|
||||
self.__ind = ind
|
||||
self.__links = set()
|
||||
|
||||
@property
|
||||
def ind(self):
|
||||
return self.__ind
|
||||
|
||||
@property
|
||||
def links(self):
|
||||
return set(self.__links)
|
||||
|
||||
def add_link(self, link_node):
|
||||
self.__links.add(link_node)
|
||||
link_node.__links.add(self)
|
||||
|
||||
|
||||
def graph_propagation(edges, scores, text_comps, edge_len_thr=50.0):
|
||||
assert edges.ndim == 2
|
||||
assert edges.shape[1] == 2
|
||||
assert edges.shape[0] == scores.shape[0]
|
||||
assert text_comps.ndim == 2
|
||||
assert isinstance(edge_len_thr, float)
|
||||
|
||||
edges = np.sort(edges, axis=1)
|
||||
score_dict = {}
|
||||
for i, edge in enumerate(edges):
|
||||
if text_comps is not None:
|
||||
box1 = text_comps[edge[0], :8].reshape(4, 2)
|
||||
box2 = text_comps[edge[1], :8].reshape(4, 2)
|
||||
center1 = np.mean(box1, axis=0)
|
||||
center2 = np.mean(box2, axis=0)
|
||||
distance = norm(center1 - center2)
|
||||
if distance > edge_len_thr:
|
||||
scores[i] = 0
|
||||
if (edge[0], edge[1]) in score_dict:
|
||||
score_dict[edge[0], edge[1]] = 0.5 * (
|
||||
score_dict[edge[0], edge[1]] + scores[i]
|
||||
)
|
||||
else:
|
||||
score_dict[edge[0], edge[1]] = scores[i]
|
||||
|
||||
nodes = np.sort(np.unique(edges.flatten()))
|
||||
mapping = -1 * np.ones((np.max(nodes) + 1), dtype=np.int32)
|
||||
mapping[nodes] = np.arange(nodes.shape[0])
|
||||
order_inds = mapping[edges]
|
||||
vertices = [Node(node) for node in nodes]
|
||||
for ind in order_inds:
|
||||
vertices[ind[0]].add_link(vertices[ind[1]])
|
||||
|
||||
return vertices, score_dict
|
||||
|
||||
|
||||
def connected_components(nodes, score_dict, link_thr):
|
||||
assert isinstance(nodes, list)
|
||||
assert all([isinstance(node, Node) for node in nodes])
|
||||
assert isinstance(score_dict, dict)
|
||||
assert isinstance(link_thr, float)
|
||||
|
||||
clusters = []
|
||||
nodes = set(nodes)
|
||||
while nodes:
|
||||
node = nodes.pop()
|
||||
cluster = {node}
|
||||
node_queue = [node]
|
||||
while node_queue:
|
||||
node = node_queue.pop(0)
|
||||
neighbors = set(
|
||||
[
|
||||
neighbor
|
||||
for neighbor in node.links
|
||||
if score_dict[tuple(sorted([node.ind, neighbor.ind]))] >= link_thr
|
||||
]
|
||||
)
|
||||
neighbors.difference_update(cluster)
|
||||
nodes.difference_update(neighbors)
|
||||
cluster.update(neighbors)
|
||||
node_queue.extend(neighbors)
|
||||
clusters.append(list(cluster))
|
||||
return clusters
|
||||
|
||||
|
||||
def clusters2labels(clusters, num_nodes):
|
||||
assert isinstance(clusters, list)
|
||||
assert all([isinstance(cluster, list) for cluster in clusters])
|
||||
assert all([isinstance(node, Node) for cluster in clusters for node in cluster])
|
||||
assert isinstance(num_nodes, int)
|
||||
|
||||
node_labels = np.zeros(num_nodes)
|
||||
for cluster_ind, cluster in enumerate(clusters):
|
||||
for node in cluster:
|
||||
node_labels[node.ind] = cluster_ind
|
||||
return node_labels
|
||||
|
||||
|
||||
def remove_single(text_comps, comp_pred_labels):
|
||||
assert text_comps.ndim == 2
|
||||
assert text_comps.shape[0] == comp_pred_labels.shape[0]
|
||||
|
||||
single_flags = np.zeros_like(comp_pred_labels)
|
||||
pred_labels = np.unique(comp_pred_labels)
|
||||
for label in pred_labels:
|
||||
current_label_flag = comp_pred_labels == label
|
||||
if np.sum(current_label_flag) == 1:
|
||||
single_flags[np.where(current_label_flag)[0][0]] = 1
|
||||
keep_ind = [i for i in range(len(comp_pred_labels)) if not single_flags[i]]
|
||||
filtered_text_comps = text_comps[keep_ind, :]
|
||||
filtered_labels = comp_pred_labels[keep_ind]
|
||||
|
||||
return filtered_text_comps, filtered_labels
|
||||
|
||||
|
||||
def norm2(point1, point2):
|
||||
return ((point1[0] - point2[0]) ** 2 + (point1[1] - point2[1]) ** 2) ** 0.5
|
||||
|
||||
|
||||
def min_connect_path(points):
|
||||
assert isinstance(points, list)
|
||||
assert all([isinstance(point, list) for point in points])
|
||||
assert all([isinstance(coord, int) for point in points for coord in point])
|
||||
|
||||
points_queue = points.copy()
|
||||
shortest_path = []
|
||||
current_edge = [[], []]
|
||||
|
||||
edge_dict0 = {}
|
||||
edge_dict1 = {}
|
||||
current_edge[0] = points_queue[0]
|
||||
current_edge[1] = points_queue[0]
|
||||
points_queue.remove(points_queue[0])
|
||||
while points_queue:
|
||||
for point in points_queue:
|
||||
length0 = norm2(point, current_edge[0])
|
||||
edge_dict0[length0] = [point, current_edge[0]]
|
||||
length1 = norm2(current_edge[1], point)
|
||||
edge_dict1[length1] = [current_edge[1], point]
|
||||
key0 = min(edge_dict0.keys())
|
||||
key1 = min(edge_dict1.keys())
|
||||
|
||||
if key0 <= key1:
|
||||
start = edge_dict0[key0][0]
|
||||
end = edge_dict0[key0][1]
|
||||
shortest_path.insert(0, [points.index(start), points.index(end)])
|
||||
points_queue.remove(start)
|
||||
current_edge[0] = start
|
||||
else:
|
||||
start = edge_dict1[key1][0]
|
||||
end = edge_dict1[key1][1]
|
||||
shortest_path.append([points.index(start), points.index(end)])
|
||||
points_queue.remove(end)
|
||||
current_edge[1] = end
|
||||
|
||||
edge_dict0 = {}
|
||||
edge_dict1 = {}
|
||||
|
||||
shortest_path = functools.reduce(operator.concat, shortest_path)
|
||||
shortest_path = sorted(set(shortest_path), key=shortest_path.index)
|
||||
|
||||
return shortest_path
|
||||
|
||||
|
||||
def in_contour(cont, point):
|
||||
x, y = point
|
||||
is_inner = cv2.pointPolygonTest(cont, (int(x), int(y)), False) > 0.5
|
||||
return is_inner
|
||||
|
||||
|
||||
def fix_corner(top_line, bot_line, start_box, end_box):
|
||||
assert isinstance(top_line, list)
|
||||
assert all(isinstance(point, list) for point in top_line)
|
||||
assert isinstance(bot_line, list)
|
||||
assert all(isinstance(point, list) for point in bot_line)
|
||||
assert start_box.shape == end_box.shape == (4, 2)
|
||||
|
||||
contour = np.array(top_line + bot_line[::-1])
|
||||
start_left_mid = (start_box[0] + start_box[3]) / 2
|
||||
start_right_mid = (start_box[1] + start_box[2]) / 2
|
||||
end_left_mid = (end_box[0] + end_box[3]) / 2
|
||||
end_right_mid = (end_box[1] + end_box[2]) / 2
|
||||
if not in_contour(contour, start_left_mid):
|
||||
top_line.insert(0, start_box[0].tolist())
|
||||
bot_line.insert(0, start_box[3].tolist())
|
||||
elif not in_contour(contour, start_right_mid):
|
||||
top_line.insert(0, start_box[1].tolist())
|
||||
bot_line.insert(0, start_box[2].tolist())
|
||||
if not in_contour(contour, end_left_mid):
|
||||
top_line.append(end_box[0].tolist())
|
||||
bot_line.append(end_box[3].tolist())
|
||||
elif not in_contour(contour, end_right_mid):
|
||||
top_line.append(end_box[1].tolist())
|
||||
bot_line.append(end_box[2].tolist())
|
||||
return top_line, bot_line
|
||||
|
||||
|
||||
def comps2boundaries(text_comps, comp_pred_labels):
|
||||
assert text_comps.ndim == 2
|
||||
assert len(text_comps) == len(comp_pred_labels)
|
||||
boundaries = []
|
||||
if len(text_comps) < 1:
|
||||
return boundaries
|
||||
for cluster_ind in range(0, int(np.max(comp_pred_labels)) + 1):
|
||||
cluster_comp_inds = np.where(comp_pred_labels == cluster_ind)
|
||||
text_comp_boxes = (
|
||||
text_comps[cluster_comp_inds, :8].reshape((-1, 4, 2)).astype(np.int32)
|
||||
)
|
||||
score = np.mean(text_comps[cluster_comp_inds, -1])
|
||||
|
||||
if text_comp_boxes.shape[0] < 1:
|
||||
continue
|
||||
|
||||
elif text_comp_boxes.shape[0] > 1:
|
||||
centers = np.mean(text_comp_boxes, axis=1).astype(np.int32).tolist()
|
||||
shortest_path = min_connect_path(centers)
|
||||
text_comp_boxes = text_comp_boxes[shortest_path]
|
||||
top_line = (
|
||||
np.mean(text_comp_boxes[:, 0:2, :], axis=1).astype(np.int32).tolist()
|
||||
)
|
||||
bot_line = (
|
||||
np.mean(text_comp_boxes[:, 2:4, :], axis=1).astype(np.int32).tolist()
|
||||
)
|
||||
top_line, bot_line = fix_corner(
|
||||
top_line, bot_line, text_comp_boxes[0], text_comp_boxes[-1]
|
||||
)
|
||||
boundary_points = top_line + bot_line[::-1]
|
||||
|
||||
else:
|
||||
top_line = text_comp_boxes[0, 0:2, :].astype(np.int32).tolist()
|
||||
bot_line = text_comp_boxes[0, 2:4:-1, :].astype(np.int32).tolist()
|
||||
boundary_points = top_line + bot_line
|
||||
|
||||
boundary = [p for coord in boundary_points for p in coord] + [score]
|
||||
boundaries.append(boundary)
|
||||
|
||||
return boundaries
|
||||
|
||||
|
||||
class DRRGPostprocess(object):
|
||||
"""Merge text components and construct boundaries of text instances.
|
||||
|
||||
Args:
|
||||
link_thr (float): The edge score threshold.
|
||||
"""
|
||||
|
||||
def __init__(self, link_thr, **kwargs):
|
||||
assert isinstance(link_thr, float)
|
||||
self.link_thr = link_thr
|
||||
|
||||
def __call__(self, preds, shape_list):
|
||||
"""
|
||||
Args:
|
||||
edges (ndarray): The edge array of shape N * 2, each row is a node
|
||||
index pair that makes up an edge in graph.
|
||||
scores (ndarray): The edge score array of shape (N,).
|
||||
text_comps (ndarray): The text components.
|
||||
|
||||
Returns:
|
||||
List[list[float]]: The predicted boundaries of text instances.
|
||||
"""
|
||||
edges, scores, text_comps = preds
|
||||
if edges is not None:
|
||||
if isinstance(edges, paddle.Tensor):
|
||||
edges = edges.numpy()
|
||||
if isinstance(scores, paddle.Tensor):
|
||||
scores = scores.numpy()
|
||||
if isinstance(text_comps, paddle.Tensor):
|
||||
text_comps = text_comps.numpy()
|
||||
assert len(edges) == len(scores)
|
||||
assert text_comps.ndim == 2
|
||||
assert text_comps.shape[1] == 9
|
||||
|
||||
vertices, score_dict = graph_propagation(edges, scores, text_comps)
|
||||
clusters = connected_components(vertices, score_dict, self.link_thr)
|
||||
pred_labels = clusters2labels(clusters, text_comps.shape[0])
|
||||
text_comps, pred_labels = remove_single(text_comps, pred_labels)
|
||||
boundaries = comps2boundaries(text_comps, pred_labels)
|
||||
else:
|
||||
boundaries = []
|
||||
|
||||
boundaries, scores = self.resize_boundary(
|
||||
boundaries, (1 / shape_list[0, 2:]).tolist()[::-1]
|
||||
)
|
||||
boxes_batch = [dict(points=boundaries, scores=scores)]
|
||||
return boxes_batch
|
||||
|
||||
def resize_boundary(self, boundaries, scale_factor):
|
||||
"""Rescale boundaries via scale_factor.
|
||||
|
||||
Args:
|
||||
boundaries (list[list[float]]): The boundary list. Each boundary
|
||||
with size 2k+1 with k>=4.
|
||||
scale_factor(ndarray): The scale factor of size (4,).
|
||||
|
||||
Returns:
|
||||
boundaries (list[list[float]]): The scaled boundaries.
|
||||
"""
|
||||
boxes = []
|
||||
scores = []
|
||||
for b in boundaries:
|
||||
sz = len(b)
|
||||
scores.append(b[-1])
|
||||
b = (
|
||||
(
|
||||
np.array(b[: sz - 1])
|
||||
* (np.tile(scale_factor[:2], int((sz - 1) / 2)).reshape(1, sz - 1))
|
||||
)
|
||||
.flatten()
|
||||
.tolist()
|
||||
)
|
||||
boxes.append(np.array(b).reshape([-1, 2]))
|
||||
return boxes, scores
|
||||
141
ppocr/postprocess/east_postprocess.py
Executable file
141
ppocr/postprocess/east_postprocess.py
Executable file
@@ -0,0 +1,141 @@
|
||||
# Copyright (c) 2020 PaddlePaddle Authors. All Rights Reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from __future__ import absolute_import
|
||||
from __future__ import division
|
||||
from __future__ import print_function
|
||||
|
||||
import numpy as np
|
||||
from .locality_aware_nms import nms_locality
|
||||
import cv2
|
||||
import paddle
|
||||
|
||||
import os
|
||||
from ppocr.utils.utility import check_install
|
||||
import sys
|
||||
|
||||
|
||||
class EASTPostProcess(object):
|
||||
"""
|
||||
The post process for EAST.
|
||||
"""
|
||||
|
||||
def __init__(self, score_thresh=0.8, cover_thresh=0.1, nms_thresh=0.2, **kwargs):
|
||||
self.score_thresh = score_thresh
|
||||
self.cover_thresh = cover_thresh
|
||||
self.nms_thresh = nms_thresh
|
||||
|
||||
def restore_rectangle_quad(self, origin, geometry):
|
||||
"""
|
||||
Restore rectangle from quadrangle.
|
||||
"""
|
||||
# quad
|
||||
origin_concat = np.concatenate(
|
||||
(origin, origin, origin, origin), axis=1
|
||||
) # (n, 8)
|
||||
pred_quads = origin_concat - geometry
|
||||
pred_quads = pred_quads.reshape((-1, 4, 2)) # (n, 4, 2)
|
||||
return pred_quads
|
||||
|
||||
def detect(
|
||||
self, score_map, geo_map, score_thresh=0.8, cover_thresh=0.1, nms_thresh=0.2
|
||||
):
|
||||
"""
|
||||
restore text boxes from score map and geo map
|
||||
"""
|
||||
|
||||
score_map = score_map[0]
|
||||
geo_map = np.swapaxes(geo_map, 1, 0)
|
||||
geo_map = np.swapaxes(geo_map, 1, 2)
|
||||
# filter the score map
|
||||
xy_text = np.argwhere(score_map > score_thresh)
|
||||
if len(xy_text) == 0:
|
||||
return []
|
||||
# sort the text boxes via the y axis
|
||||
xy_text = xy_text[np.argsort(xy_text[:, 0])]
|
||||
# restore quad proposals
|
||||
text_box_restored = self.restore_rectangle_quad(
|
||||
xy_text[:, ::-1] * 4, geo_map[xy_text[:, 0], xy_text[:, 1], :]
|
||||
)
|
||||
boxes = np.zeros((text_box_restored.shape[0], 9), dtype=np.float32)
|
||||
boxes[:, :8] = text_box_restored.reshape((-1, 8))
|
||||
boxes[:, 8] = score_map[xy_text[:, 0], xy_text[:, 1]]
|
||||
|
||||
try:
|
||||
check_install("lanms", "lanms-nova")
|
||||
import lanms
|
||||
|
||||
boxes = lanms.merge_quadrangle_n9(boxes, nms_thresh)
|
||||
except:
|
||||
print(
|
||||
"You should install lanms by pip3 install lanms-nova to speed up nms_locality"
|
||||
)
|
||||
boxes = nms_locality(boxes.astype(np.float64), nms_thresh)
|
||||
if boxes.shape[0] == 0:
|
||||
return []
|
||||
# Here we filter some low score boxes by the average score map,
|
||||
# this is different from the original paper.
|
||||
for i, box in enumerate(boxes):
|
||||
mask = np.zeros_like(score_map, dtype=np.uint8)
|
||||
cv2.fillPoly(mask, box[:8].reshape((-1, 4, 2)).astype(np.int32) // 4, 1)
|
||||
boxes[i, 8] = cv2.mean(score_map, mask)[0]
|
||||
boxes = boxes[boxes[:, 8] > cover_thresh]
|
||||
return boxes
|
||||
|
||||
def sort_poly(self, p):
|
||||
"""
|
||||
Sort polygons.
|
||||
"""
|
||||
min_axis = np.argmin(np.sum(p, axis=1))
|
||||
p = p[[min_axis, (min_axis + 1) % 4, (min_axis + 2) % 4, (min_axis + 3) % 4]]
|
||||
if abs(p[0, 0] - p[1, 0]) > abs(p[0, 1] - p[1, 1]):
|
||||
return p
|
||||
else:
|
||||
return p[[0, 3, 2, 1]]
|
||||
|
||||
def __call__(self, outs_dict, shape_list):
|
||||
score_list = outs_dict["f_score"]
|
||||
geo_list = outs_dict["f_geo"]
|
||||
if isinstance(score_list, paddle.Tensor):
|
||||
score_list = score_list.numpy()
|
||||
geo_list = geo_list.numpy()
|
||||
img_num = len(shape_list)
|
||||
dt_boxes_list = []
|
||||
for ino in range(img_num):
|
||||
score = score_list[ino]
|
||||
geo = geo_list[ino]
|
||||
boxes = self.detect(
|
||||
score_map=score,
|
||||
geo_map=geo,
|
||||
score_thresh=self.score_thresh,
|
||||
cover_thresh=self.cover_thresh,
|
||||
nms_thresh=self.nms_thresh,
|
||||
)
|
||||
boxes_norm = []
|
||||
if len(boxes) > 0:
|
||||
h, w = score.shape[1:]
|
||||
src_h, src_w, ratio_h, ratio_w = shape_list[ino]
|
||||
boxes = boxes[:, :8].reshape((-1, 4, 2))
|
||||
boxes[:, :, 0] /= ratio_w
|
||||
boxes[:, :, 1] /= ratio_h
|
||||
for i_box, box in enumerate(boxes):
|
||||
box = self.sort_poly(box.astype(np.int32))
|
||||
if (
|
||||
np.linalg.norm(box[0] - box[1]) < 5
|
||||
or np.linalg.norm(box[3] - box[0]) < 5
|
||||
):
|
||||
continue
|
||||
boxes_norm.append(box)
|
||||
dt_boxes_list.append({"points": np.array(boxes_norm)})
|
||||
return dt_boxes_list
|
||||
250
ppocr/postprocess/fce_postprocess.py
Executable file
250
ppocr/postprocess/fce_postprocess.py
Executable file
@@ -0,0 +1,250 @@
|
||||
# copyright (c) 2022 PaddlePaddle Authors. All Rights Reserve.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""
|
||||
This code is refer from:
|
||||
https://github.com/open-mmlab/mmocr/blob/v0.3.0/mmocr/models/textdet/postprocess/wrapper.py
|
||||
"""
|
||||
|
||||
import cv2
|
||||
import paddle
|
||||
import numpy as np
|
||||
from numpy.fft import ifft
|
||||
from ppocr.utils.poly_nms import poly_nms, valid_boundary
|
||||
|
||||
|
||||
def fill_hole(input_mask):
|
||||
h, w = input_mask.shape
|
||||
canvas = np.zeros((h + 2, w + 2), np.uint8)
|
||||
canvas[1 : h + 1, 1 : w + 1] = input_mask.copy()
|
||||
|
||||
mask = np.zeros((h + 4, w + 4), np.uint8)
|
||||
|
||||
cv2.floodFill(canvas, mask, (0, 0), 1)
|
||||
canvas = canvas[1 : h + 1, 1 : w + 1].astype(np.bool_)
|
||||
|
||||
return ~canvas | input_mask
|
||||
|
||||
|
||||
def fourier2poly(fourier_coeff, num_reconstr_points=50):
|
||||
"""Inverse Fourier transform
|
||||
Args:
|
||||
fourier_coeff (ndarray): Fourier coefficients shaped (n, 2k+1),
|
||||
with n and k being candidates number and Fourier degree
|
||||
respectively.
|
||||
num_reconstr_points (int): Number of reconstructed polygon points.
|
||||
Returns:
|
||||
Polygons (ndarray): The reconstructed polygons shaped (n, n')
|
||||
"""
|
||||
|
||||
a = np.zeros((len(fourier_coeff), num_reconstr_points), dtype="complex")
|
||||
k = (len(fourier_coeff[0]) - 1) // 2
|
||||
|
||||
a[:, 0 : k + 1] = fourier_coeff[:, k:]
|
||||
a[:, -k:] = fourier_coeff[:, :k]
|
||||
|
||||
poly_complex = ifft(a) * num_reconstr_points
|
||||
polygon = np.zeros((len(fourier_coeff), num_reconstr_points, 2))
|
||||
polygon[:, :, 0] = poly_complex.real
|
||||
polygon[:, :, 1] = poly_complex.imag
|
||||
return polygon.astype("int32").reshape((len(fourier_coeff), -1))
|
||||
|
||||
|
||||
class FCEPostProcess(object):
|
||||
"""
|
||||
The post process for FCENet.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
scales,
|
||||
fourier_degree=5,
|
||||
num_reconstr_points=50,
|
||||
decoding_type="fcenet",
|
||||
score_thr=0.3,
|
||||
nms_thr=0.1,
|
||||
alpha=1.0,
|
||||
beta=1.0,
|
||||
box_type="poly",
|
||||
**kwargs,
|
||||
):
|
||||
self.scales = scales
|
||||
self.fourier_degree = fourier_degree
|
||||
self.num_reconstr_points = num_reconstr_points
|
||||
self.decoding_type = decoding_type
|
||||
self.score_thr = score_thr
|
||||
self.nms_thr = nms_thr
|
||||
self.alpha = alpha
|
||||
self.beta = beta
|
||||
self.box_type = box_type
|
||||
|
||||
def __call__(self, preds, shape_list):
|
||||
score_maps = []
|
||||
for key, value in preds.items():
|
||||
if isinstance(value, paddle.Tensor):
|
||||
value = value.numpy()
|
||||
cls_res = value[:, :4, :, :]
|
||||
reg_res = value[:, 4:, :, :]
|
||||
score_maps.append([cls_res, reg_res])
|
||||
|
||||
return self.get_boundary(score_maps, shape_list)
|
||||
|
||||
def resize_boundary(self, boundaries, scale_factor):
|
||||
"""Rescale boundaries via scale_factor.
|
||||
|
||||
Args:
|
||||
boundaries (list[list[float]]): The boundary list. Each boundary
|
||||
with size 2k+1 with k>=4.
|
||||
scale_factor(ndarray): The scale factor of size (4,).
|
||||
|
||||
Returns:
|
||||
boundaries (list[list[float]]): The scaled boundaries.
|
||||
"""
|
||||
boxes = []
|
||||
scores = []
|
||||
for b in boundaries:
|
||||
sz = len(b)
|
||||
valid_boundary(b, True)
|
||||
scores.append(b[-1])
|
||||
b = (
|
||||
(
|
||||
np.array(b[: sz - 1])
|
||||
* (np.tile(scale_factor[:2], int((sz - 1) / 2)).reshape(1, sz - 1))
|
||||
)
|
||||
.flatten()
|
||||
.tolist()
|
||||
)
|
||||
boxes.append(np.array(b).reshape([-1, 2]))
|
||||
|
||||
return np.array(boxes, dtype=np.float32), scores
|
||||
|
||||
def get_boundary(self, score_maps, shape_list):
|
||||
assert len(score_maps) == len(self.scales)
|
||||
boundaries = []
|
||||
for idx, score_map in enumerate(score_maps):
|
||||
scale = self.scales[idx]
|
||||
boundaries = boundaries + self._get_boundary_single(score_map, scale)
|
||||
|
||||
# nms
|
||||
boundaries = poly_nms(boundaries, self.nms_thr)
|
||||
boundaries, scores = self.resize_boundary(
|
||||
boundaries, (1 / shape_list[0, 2:]).tolist()[::-1]
|
||||
)
|
||||
|
||||
boxes_batch = [dict(points=boundaries, scores=scores)]
|
||||
return boxes_batch
|
||||
|
||||
def _get_boundary_single(self, score_map, scale):
|
||||
assert len(score_map) == 2
|
||||
assert score_map[1].shape[1] == 4 * self.fourier_degree + 2
|
||||
|
||||
return self.fcenet_decode(
|
||||
preds=score_map,
|
||||
fourier_degree=self.fourier_degree,
|
||||
num_reconstr_points=self.num_reconstr_points,
|
||||
scale=scale,
|
||||
alpha=self.alpha,
|
||||
beta=self.beta,
|
||||
box_type=self.box_type,
|
||||
score_thr=self.score_thr,
|
||||
nms_thr=self.nms_thr,
|
||||
)
|
||||
|
||||
def fcenet_decode(
|
||||
self,
|
||||
preds,
|
||||
fourier_degree,
|
||||
num_reconstr_points,
|
||||
scale,
|
||||
alpha=1.0,
|
||||
beta=2.0,
|
||||
box_type="poly",
|
||||
score_thr=0.3,
|
||||
nms_thr=0.1,
|
||||
):
|
||||
"""Decoding predictions of FCENet to instances.
|
||||
|
||||
Args:
|
||||
preds (list(Tensor)): The head output tensors.
|
||||
fourier_degree (int): The maximum Fourier transform degree k.
|
||||
num_reconstr_points (int): The points number of the polygon
|
||||
reconstructed from predicted Fourier coefficients.
|
||||
scale (int): The down-sample scale of the prediction.
|
||||
alpha (float) : The parameter to calculate final scores. Score_{final}
|
||||
= (Score_{text region} ^ alpha)
|
||||
* (Score_{text center region}^ beta)
|
||||
beta (float) : The parameter to calculate final score.
|
||||
box_type (str): Boundary encoding type 'poly' or 'quad'.
|
||||
score_thr (float) : The threshold used to filter out the final
|
||||
candidates.
|
||||
nms_thr (float) : The threshold of nms.
|
||||
|
||||
Returns:
|
||||
boundaries (list[list[float]]): The instance boundary and confidence
|
||||
list.
|
||||
"""
|
||||
assert isinstance(preds, list)
|
||||
assert len(preds) == 2
|
||||
assert box_type in ["poly", "quad"]
|
||||
|
||||
cls_pred = preds[0][0]
|
||||
tr_pred = cls_pred[0:2]
|
||||
tcl_pred = cls_pred[2:]
|
||||
|
||||
reg_pred = preds[1][0].transpose([1, 2, 0])
|
||||
x_pred = reg_pred[:, :, : 2 * fourier_degree + 1]
|
||||
y_pred = reg_pred[:, :, 2 * fourier_degree + 1 :]
|
||||
|
||||
score_pred = (tr_pred[1] ** alpha) * (tcl_pred[1] ** beta)
|
||||
tr_pred_mask = (score_pred) > score_thr
|
||||
tr_mask = fill_hole(tr_pred_mask)
|
||||
|
||||
tr_contours, _ = cv2.findContours(
|
||||
tr_mask.astype(np.uint8), cv2.RETR_TREE, cv2.CHAIN_APPROX_SIMPLE
|
||||
) # opencv4
|
||||
|
||||
mask = np.zeros_like(tr_mask)
|
||||
boundaries = []
|
||||
for cont in tr_contours:
|
||||
deal_map = mask.copy().astype(np.int8)
|
||||
cv2.drawContours(deal_map, [cont], -1, 1, -1)
|
||||
|
||||
score_map = score_pred * deal_map
|
||||
score_mask = score_map > 0
|
||||
xy_text = np.argwhere(score_mask)
|
||||
dxy = xy_text[:, 1] + xy_text[:, 0] * 1j
|
||||
|
||||
x, y = x_pred[score_mask], y_pred[score_mask]
|
||||
c = x + y * 1j
|
||||
c[:, fourier_degree] = c[:, fourier_degree] + dxy
|
||||
c *= scale
|
||||
|
||||
polygons = fourier2poly(c, num_reconstr_points)
|
||||
score = score_map[score_mask].reshape(-1, 1)
|
||||
polygons = poly_nms(np.hstack((polygons, score)).tolist(), nms_thr)
|
||||
|
||||
boundaries = boundaries + polygons
|
||||
|
||||
boundaries = poly_nms(boundaries, nms_thr)
|
||||
|
||||
if box_type == "quad":
|
||||
new_boundaries = []
|
||||
for boundary in boundaries:
|
||||
poly = np.array(boundary[:-1]).reshape(-1, 2).astype(np.float32)
|
||||
score = boundary[-1]
|
||||
points = cv2.boxPoints(cv2.minAreaRect(poly))
|
||||
points = np.int64(points)
|
||||
new_boundaries.append(points.reshape(-1).tolist() + [score])
|
||||
boundaries = new_boundaries
|
||||
|
||||
return boundaries
|
||||
198
ppocr/postprocess/locality_aware_nms.py
Normal file
198
ppocr/postprocess/locality_aware_nms.py
Normal file
@@ -0,0 +1,198 @@
|
||||
"""
|
||||
Locality aware nms.
|
||||
This code is referred from: https://github.com/songdejia/EAST/blob/master/locality_aware_nms.py
|
||||
"""
|
||||
|
||||
import numpy as np
|
||||
from shapely.geometry import Polygon
|
||||
|
||||
|
||||
def intersection(g, p):
|
||||
"""
|
||||
Intersection.
|
||||
"""
|
||||
g = Polygon(g[:8].reshape((4, 2)))
|
||||
p = Polygon(p[:8].reshape((4, 2)))
|
||||
g = g.buffer(0)
|
||||
p = p.buffer(0)
|
||||
if not g.is_valid or not p.is_valid:
|
||||
return 0
|
||||
inter = Polygon(g).intersection(Polygon(p)).area
|
||||
union = g.area + p.area - inter
|
||||
if union == 0:
|
||||
return 0
|
||||
else:
|
||||
return inter / union
|
||||
|
||||
|
||||
def intersection_iog(g, p):
|
||||
"""
|
||||
Intersection_iog.
|
||||
"""
|
||||
g = Polygon(g[:8].reshape((4, 2)))
|
||||
p = Polygon(p[:8].reshape((4, 2)))
|
||||
if not g.is_valid or not p.is_valid:
|
||||
return 0
|
||||
inter = Polygon(g).intersection(Polygon(p)).area
|
||||
# union = g.area + p.area - inter
|
||||
union = p.area
|
||||
if union == 0:
|
||||
print("p_area is very small")
|
||||
return 0
|
||||
else:
|
||||
return inter / union
|
||||
|
||||
|
||||
def weighted_merge(g, p):
|
||||
"""
|
||||
Weighted merge.
|
||||
"""
|
||||
g[:8] = (g[8] * g[:8] + p[8] * p[:8]) / (g[8] + p[8])
|
||||
g[8] = g[8] + p[8]
|
||||
return g
|
||||
|
||||
|
||||
def standard_nms(S, thres):
|
||||
"""
|
||||
Standard nms.
|
||||
"""
|
||||
order = np.argsort(S[:, 8])[::-1]
|
||||
keep = []
|
||||
while order.size > 0:
|
||||
i = order[0]
|
||||
keep.append(i)
|
||||
ovr = np.array([intersection(S[i], S[t]) for t in order[1:]])
|
||||
|
||||
inds = np.where(ovr <= thres)[0]
|
||||
order = order[inds + 1]
|
||||
|
||||
return S[keep]
|
||||
|
||||
|
||||
def standard_nms_inds(S, thres):
|
||||
"""
|
||||
Standard nms, return inds.
|
||||
"""
|
||||
order = np.argsort(S[:, 8])[::-1]
|
||||
keep = []
|
||||
while order.size > 0:
|
||||
i = order[0]
|
||||
keep.append(i)
|
||||
ovr = np.array([intersection(S[i], S[t]) for t in order[1:]])
|
||||
|
||||
inds = np.where(ovr <= thres)[0]
|
||||
order = order[inds + 1]
|
||||
|
||||
return keep
|
||||
|
||||
|
||||
def nms(S, thres):
|
||||
"""
|
||||
nms.
|
||||
"""
|
||||
order = np.argsort(S[:, 8])[::-1]
|
||||
keep = []
|
||||
while order.size > 0:
|
||||
i = order[0]
|
||||
keep.append(i)
|
||||
ovr = np.array([intersection(S[i], S[t]) for t in order[1:]])
|
||||
|
||||
inds = np.where(ovr <= thres)[0]
|
||||
order = order[inds + 1]
|
||||
|
||||
return keep
|
||||
|
||||
|
||||
def soft_nms(boxes_in, Nt_thres=0.3, threshold=0.8, sigma=0.5, method=2):
|
||||
"""
|
||||
soft_nms
|
||||
:para boxes_in, N x 9 (coords + score)
|
||||
:para threshould, eliminate cases min score(0.001)
|
||||
:para Nt_thres, iou_threshi
|
||||
:para sigma, gaussian weght
|
||||
:method, linear or gaussian
|
||||
"""
|
||||
boxes = boxes_in.copy()
|
||||
N = boxes.shape[0]
|
||||
if N is None or N < 1:
|
||||
return np.array([])
|
||||
pos, maxpos = 0, 0
|
||||
weight = 0.0
|
||||
inds = np.arange(N)
|
||||
tbox, sbox = boxes[0].copy(), boxes[0].copy()
|
||||
for i in range(N):
|
||||
maxscore = boxes[i, 8]
|
||||
maxpos = i
|
||||
tbox = boxes[i].copy()
|
||||
ti = inds[i]
|
||||
pos = i + 1
|
||||
# get max box
|
||||
while pos < N:
|
||||
if maxscore < boxes[pos, 8]:
|
||||
maxscore = boxes[pos, 8]
|
||||
maxpos = pos
|
||||
pos = pos + 1
|
||||
# add max box as a detection
|
||||
boxes[i, :] = boxes[maxpos, :]
|
||||
inds[i] = inds[maxpos]
|
||||
# swap
|
||||
boxes[maxpos, :] = tbox
|
||||
inds[maxpos] = ti
|
||||
tbox = boxes[i].copy()
|
||||
pos = i + 1
|
||||
# NMS iteration
|
||||
while pos < N:
|
||||
sbox = boxes[pos].copy()
|
||||
ts_iou_val = intersection(tbox, sbox)
|
||||
if ts_iou_val > 0:
|
||||
if method == 1:
|
||||
if ts_iou_val > Nt_thres:
|
||||
weight = 1 - ts_iou_val
|
||||
else:
|
||||
weight = 1
|
||||
elif method == 2:
|
||||
weight = np.exp(-1.0 * ts_iou_val**2 / sigma)
|
||||
else:
|
||||
if ts_iou_val > Nt_thres:
|
||||
weight = 0
|
||||
else:
|
||||
weight = 1
|
||||
boxes[pos, 8] = weight * boxes[pos, 8]
|
||||
# if box score falls below threshold, discard the box by
|
||||
# swapping last box update N
|
||||
if boxes[pos, 8] < threshold:
|
||||
boxes[pos, :] = boxes[N - 1, :]
|
||||
inds[pos] = inds[N - 1]
|
||||
N = N - 1
|
||||
pos = pos - 1
|
||||
pos = pos + 1
|
||||
|
||||
return boxes[:N]
|
||||
|
||||
|
||||
def nms_locality(polys, thres=0.3):
|
||||
"""
|
||||
locality aware nms of EAST
|
||||
:param polys: a N*9 numpy array. first 8 coordinates, then prob
|
||||
:return: boxes after nms
|
||||
"""
|
||||
S = []
|
||||
p = None
|
||||
for g in polys:
|
||||
if p is not None and intersection(g, p) > thres:
|
||||
p = weighted_merge(g, p)
|
||||
else:
|
||||
if p is not None:
|
||||
S.append(p)
|
||||
p = g
|
||||
if p is not None:
|
||||
S.append(p)
|
||||
|
||||
if len(S) == 0:
|
||||
return np.array([])
|
||||
return standard_nms(np.array(S), thres)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# 343,350,448,135,474,143,369,359
|
||||
print(Polygon(np.array([[343, 350], [448, 135], [474, 143], [369, 359]])).area)
|
||||
66
ppocr/postprocess/pg_postprocess.py
Normal file
66
ppocr/postprocess/pg_postprocess.py
Normal file
@@ -0,0 +1,66 @@
|
||||
# Copyright (c) 2021 PaddlePaddle Authors. All Rights Reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from __future__ import absolute_import
|
||||
from __future__ import division
|
||||
from __future__ import print_function
|
||||
|
||||
import os
|
||||
import sys
|
||||
|
||||
__dir__ = os.path.dirname(__file__)
|
||||
sys.path.append(__dir__)
|
||||
sys.path.append(os.path.join(__dir__, ".."))
|
||||
from ppocr.utils.e2e_utils.pgnet_pp_utils import PGNet_PostProcess
|
||||
|
||||
|
||||
class PGPostProcess(object):
|
||||
"""
|
||||
The post process for PGNet.
|
||||
"""
|
||||
|
||||
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
|
||||
self.mode = mode
|
||||
self.point_gather_mode = point_gather_mode
|
||||
|
||||
# c++ la-nms is faster, but only support python 3.5
|
||||
self.is_python35 = False
|
||||
if sys.version_info.major == 3 and sys.version_info.minor == 5:
|
||||
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,
|
||||
point_gather_mode=self.point_gather_mode,
|
||||
)
|
||||
if self.mode == "fast":
|
||||
data = post.pg_postprocess_fast()
|
||||
else:
|
||||
data = post.pg_postprocess_slow()
|
||||
return data
|
||||
297
ppocr/postprocess/picodet_postprocess.py
Normal file
297
ppocr/postprocess/picodet_postprocess.py
Normal file
@@ -0,0 +1,297 @@
|
||||
# Copyright (c) 2021 PaddlePaddle Authors. All Rights Reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import numpy as np
|
||||
from scipy.special import softmax
|
||||
|
||||
|
||||
def hard_nms(box_scores, iou_threshold, top_k=-1, candidate_size=200):
|
||||
"""
|
||||
Args:
|
||||
box_scores (N, 5): boxes in corner-form and probabilities.
|
||||
iou_threshold: intersection over union threshold.
|
||||
top_k: keep top_k results. If k <= 0, keep all the results.
|
||||
candidate_size: only consider the candidates with the highest scores.
|
||||
Returns:
|
||||
picked: a list of indexes of the kept boxes
|
||||
"""
|
||||
scores = box_scores[:, -1]
|
||||
boxes = box_scores[:, :-1]
|
||||
picked = []
|
||||
indexes = np.argsort(scores)
|
||||
indexes = indexes[-candidate_size:]
|
||||
while len(indexes) > 0:
|
||||
current = indexes[-1]
|
||||
picked.append(current)
|
||||
if 0 < top_k == len(picked) or len(indexes) == 1:
|
||||
break
|
||||
current_box = boxes[current, :]
|
||||
indexes = indexes[:-1]
|
||||
rest_boxes = boxes[indexes, :]
|
||||
iou = iou_of(
|
||||
rest_boxes,
|
||||
np.expand_dims(current_box, axis=0),
|
||||
)
|
||||
indexes = indexes[iou <= iou_threshold]
|
||||
|
||||
return box_scores[picked, :]
|
||||
|
||||
|
||||
def iou_of(boxes0, boxes1, eps=1e-5):
|
||||
"""Return intersection-over-union (Jaccard index) of boxes.
|
||||
Args:
|
||||
boxes0 (N, 4): ground truth boxes.
|
||||
boxes1 (N or 1, 4): predicted boxes.
|
||||
eps: a small number to avoid 0 as denominator.
|
||||
Returns:
|
||||
iou (N): IoU values.
|
||||
"""
|
||||
overlap_left_top = np.maximum(boxes0[..., :2], boxes1[..., :2])
|
||||
overlap_right_bottom = np.minimum(boxes0[..., 2:], boxes1[..., 2:])
|
||||
|
||||
overlap_area = area_of(overlap_left_top, overlap_right_bottom)
|
||||
area0 = area_of(boxes0[..., :2], boxes0[..., 2:])
|
||||
area1 = area_of(boxes1[..., :2], boxes1[..., 2:])
|
||||
return overlap_area / (area0 + area1 - overlap_area + eps)
|
||||
|
||||
|
||||
def area_of(left_top, right_bottom):
|
||||
"""Compute the areas of rectangles given two corners.
|
||||
Args:
|
||||
left_top (N, 2): left top corner.
|
||||
right_bottom (N, 2): right bottom corner.
|
||||
Returns:
|
||||
area (N): return the area.
|
||||
"""
|
||||
hw = np.clip(right_bottom - left_top, 0.0, None)
|
||||
return hw[..., 0] * hw[..., 1]
|
||||
|
||||
|
||||
def calculate_containment(boxes0, boxes1):
|
||||
"""
|
||||
Calculate the containment of the boxes.
|
||||
Args:
|
||||
boxes0 (N, 4): ground truth boxes.
|
||||
boxes1 (N or 1, 4): predicted boxes.
|
||||
Returns:
|
||||
containment (N): containment values.
|
||||
"""
|
||||
overlap_left_top = np.maximum(boxes0[..., :2], boxes1[..., :2])
|
||||
overlap_right_bottom = np.minimum(boxes0[..., 2:], boxes1[..., 2:])
|
||||
|
||||
overlap_area = area_of(overlap_left_top, overlap_right_bottom)
|
||||
area0 = area_of(boxes0[..., :2], boxes0[..., 2:])
|
||||
area1 = area_of(boxes1[..., :2], boxes1[..., 2:])
|
||||
return overlap_area / np.minimum(area0, np.expand_dims(area1, axis=0))
|
||||
|
||||
|
||||
class PicoDetPostProcess(object):
|
||||
"""
|
||||
Args:
|
||||
input_shape (int): network input image size
|
||||
ori_shape (int): ori image shape of before padding
|
||||
scale_factor (float): scale factor of ori image
|
||||
enable_mkldnn (bool): whether to open MKLDNN
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
layout_dict_path,
|
||||
strides=[8, 16, 32, 64],
|
||||
score_threshold=0.4,
|
||||
nms_threshold=0.5,
|
||||
nms_top_k=1000,
|
||||
keep_top_k=100,
|
||||
):
|
||||
self.labels = self.load_layout_dict(layout_dict_path)
|
||||
self.strides = strides
|
||||
self.score_threshold = score_threshold
|
||||
self.nms_threshold = nms_threshold
|
||||
self.nms_top_k = nms_top_k
|
||||
self.keep_top_k = keep_top_k
|
||||
|
||||
def load_layout_dict(self, layout_dict_path):
|
||||
with open(layout_dict_path, "r", encoding="utf-8") as fp:
|
||||
labels = fp.readlines()
|
||||
return [label.strip("\n") for label in labels]
|
||||
|
||||
def warp_boxes(self, boxes, ori_shape):
|
||||
"""Apply transform to boxes"""
|
||||
width, height = ori_shape[1], ori_shape[0]
|
||||
n = len(boxes)
|
||||
if n:
|
||||
# warp points
|
||||
xy = np.ones((n * 4, 3))
|
||||
xy[:, :2] = boxes[:, [0, 1, 2, 3, 0, 3, 2, 1]].reshape(
|
||||
n * 4, 2
|
||||
) # x1y1, x2y2, x1y2, x2y1
|
||||
# xy = xy @ M.T # transform
|
||||
xy = (xy[:, :2] / xy[:, 2:3]).reshape(n, 8) # rescale
|
||||
# create new boxes
|
||||
x = xy[:, [0, 2, 4, 6]]
|
||||
y = xy[:, [1, 3, 5, 7]]
|
||||
xy = (
|
||||
np.concatenate((x.min(1), y.min(1), x.max(1), y.max(1))).reshape(4, n).T
|
||||
)
|
||||
# clip boxes
|
||||
xy[:, [0, 2]] = xy[:, [0, 2]].clip(0, width)
|
||||
xy[:, [1, 3]] = xy[:, [1, 3]].clip(0, height)
|
||||
return xy.astype(np.float32)
|
||||
else:
|
||||
return boxes
|
||||
|
||||
def img_info(self, ori_img, img):
|
||||
origin_shape = ori_img.shape
|
||||
resize_shape = img.shape
|
||||
im_scale_y = resize_shape[2] / float(origin_shape[0])
|
||||
im_scale_x = resize_shape[3] / float(origin_shape[1])
|
||||
scale_factor = np.array([im_scale_y, im_scale_x], dtype=np.float32)
|
||||
img_shape = np.array(img.shape[2:], dtype=np.float32)
|
||||
|
||||
input_shape = np.array(img).astype("float32").shape[2:]
|
||||
ori_shape = np.array((img_shape,)).astype("float32")
|
||||
scale_factor = np.array((scale_factor,)).astype("float32")
|
||||
return ori_shape, input_shape, scale_factor
|
||||
|
||||
def __call__(self, ori_img, img, preds):
|
||||
scores, raw_boxes = preds["boxes"], preds["boxes_num"]
|
||||
batch_size = raw_boxes[0].shape[0]
|
||||
reg_max = int(raw_boxes[0].shape[-1] / 4 - 1)
|
||||
out_boxes_num = []
|
||||
out_boxes_list = []
|
||||
results = []
|
||||
ori_shape, input_shape, scale_factor = self.img_info(ori_img, img)
|
||||
|
||||
for batch_id in range(batch_size):
|
||||
# generate centers
|
||||
decode_boxes = []
|
||||
select_scores = []
|
||||
for stride, box_distribute, score in zip(self.strides, raw_boxes, scores):
|
||||
box_distribute = box_distribute[batch_id]
|
||||
score = score[batch_id]
|
||||
# centers
|
||||
fm_h = input_shape[0] / stride
|
||||
fm_w = input_shape[1] / stride
|
||||
h_range = np.arange(fm_h)
|
||||
w_range = np.arange(fm_w)
|
||||
ww, hh = np.meshgrid(w_range, h_range)
|
||||
ct_row = (hh.flatten() + 0.5) * stride
|
||||
ct_col = (ww.flatten() + 0.5) * stride
|
||||
center = np.stack((ct_col, ct_row, ct_col, ct_row), axis=1)
|
||||
|
||||
# box distribution to distance
|
||||
reg_range = np.arange(reg_max + 1)
|
||||
box_distance = box_distribute.reshape((-1, reg_max + 1))
|
||||
box_distance = softmax(box_distance, axis=1)
|
||||
box_distance = box_distance * np.expand_dims(reg_range, axis=0)
|
||||
box_distance = np.sum(box_distance, axis=1).reshape((-1, 4))
|
||||
box_distance = box_distance * stride
|
||||
|
||||
# top K candidate
|
||||
topk_idx = np.argsort(score.max(axis=1))[::-1]
|
||||
topk_idx = topk_idx[: self.nms_top_k]
|
||||
center = center[topk_idx]
|
||||
score = score[topk_idx]
|
||||
box_distance = box_distance[topk_idx]
|
||||
|
||||
# decode box
|
||||
decode_box = center + [-1, -1, 1, 1] * box_distance
|
||||
|
||||
select_scores.append(score)
|
||||
decode_boxes.append(decode_box)
|
||||
|
||||
# nms
|
||||
bboxes = np.concatenate(decode_boxes, axis=0)
|
||||
confidences = np.concatenate(select_scores, axis=0)
|
||||
picked_box_probs = []
|
||||
picked_labels = []
|
||||
for class_index in range(0, confidences.shape[1]):
|
||||
probs = confidences[:, class_index]
|
||||
mask = probs > self.score_threshold
|
||||
probs = probs[mask]
|
||||
if probs.shape[0] == 0:
|
||||
continue
|
||||
subset_boxes = bboxes[mask, :]
|
||||
box_probs = np.concatenate([subset_boxes, probs.reshape(-1, 1)], axis=1)
|
||||
box_probs = hard_nms(
|
||||
box_probs,
|
||||
iou_threshold=self.nms_threshold,
|
||||
top_k=self.keep_top_k,
|
||||
)
|
||||
picked_box_probs.append(box_probs)
|
||||
picked_labels.extend([class_index] * box_probs.shape[0])
|
||||
|
||||
if len(picked_box_probs) == 0:
|
||||
out_boxes_list.append(np.empty((0, 4)))
|
||||
out_boxes_num.append(0)
|
||||
|
||||
else:
|
||||
picked_box_probs = np.concatenate(picked_box_probs)
|
||||
|
||||
# resize output boxes
|
||||
picked_box_probs[:, :4] = self.warp_boxes(
|
||||
picked_box_probs[:, :4], ori_shape[batch_id]
|
||||
)
|
||||
im_scale = np.concatenate(
|
||||
[scale_factor[batch_id][::-1], scale_factor[batch_id][::-1]]
|
||||
)
|
||||
picked_box_probs[:, :4] /= im_scale
|
||||
# clas score box
|
||||
out_boxes_list.append(
|
||||
np.concatenate(
|
||||
[
|
||||
np.expand_dims(np.array(picked_labels), axis=-1),
|
||||
np.expand_dims(picked_box_probs[:, 4], axis=-1),
|
||||
picked_box_probs[:, :4],
|
||||
],
|
||||
axis=1,
|
||||
)
|
||||
)
|
||||
out_boxes_num.append(len(picked_labels))
|
||||
|
||||
out_boxes_list = np.concatenate(out_boxes_list, axis=0)
|
||||
out_boxes_num = np.asarray(out_boxes_num).astype(np.int32)
|
||||
|
||||
for dt in out_boxes_list:
|
||||
clsid, bbox, score = int(dt[0]), dt[2:], dt[1]
|
||||
label = self.labels[clsid]
|
||||
result = {"bbox": bbox, "label": label, "score": score}
|
||||
results.append(result)
|
||||
|
||||
# Handle conflict where a box is simultaneously recognized as multiple labels.
|
||||
# Use IoU to find similar boxes. Prioritize labels as table, text, and others when deduplicate similar boxes.
|
||||
bboxes = np.array([x["bbox"] for x in results])
|
||||
duplicate_idx = list()
|
||||
for i in range(len(results)):
|
||||
if i in duplicate_idx:
|
||||
continue
|
||||
containments = calculate_containment(bboxes, bboxes[i, ...])
|
||||
overlaps = np.where(containments > 0.5)[0]
|
||||
if len(overlaps) > 1:
|
||||
table_box = [x for x in overlaps if results[x]["label"] == "table"]
|
||||
if len(table_box) > 0:
|
||||
keep = sorted(
|
||||
[(x, results[x]) for x in table_box],
|
||||
key=lambda x: x[1]["score"],
|
||||
reverse=True,
|
||||
)[0][0]
|
||||
else:
|
||||
keep = sorted(
|
||||
[(x, results[x]) for x in overlaps],
|
||||
key=lambda x: x[1]["score"],
|
||||
reverse=True,
|
||||
)[0][0]
|
||||
duplicate_idx.extend([x for x in overlaps if x != keep])
|
||||
results = [x for i, x in enumerate(results) if i not in duplicate_idx]
|
||||
return results
|
||||
15
ppocr/postprocess/pse_postprocess/__init__.py
Normal file
15
ppocr/postprocess/pse_postprocess/__init__.py
Normal file
@@ -0,0 +1,15 @@
|
||||
# Copyright (c) 2021 PaddlePaddle Authors. All Rights Reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from .pse_postprocess import PSEPostProcess
|
||||
6
ppocr/postprocess/pse_postprocess/pse/README.md
Normal file
6
ppocr/postprocess/pse_postprocess/pse/README.md
Normal file
@@ -0,0 +1,6 @@
|
||||
## 编译
|
||||
This code is refer from:
|
||||
https://github.com/whai362/PSENet/blob/python3/models/post_processing/pse
|
||||
```python
|
||||
python3 setup.py build_ext --inplace
|
||||
```
|
||||
33
ppocr/postprocess/pse_postprocess/pse/__init__.py
Normal file
33
ppocr/postprocess/pse_postprocess/pse/__init__.py
Normal file
@@ -0,0 +1,33 @@
|
||||
# copyright (c) 2020 PaddlePaddle Authors. All Rights Reserve.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
import sys
|
||||
import os
|
||||
import subprocess
|
||||
|
||||
python_path = sys.executable
|
||||
|
||||
ori_path = os.getcwd()
|
||||
os.chdir("ppocr/postprocess/pse_postprocess/pse")
|
||||
if (
|
||||
subprocess.call("{} setup.py build_ext --inplace".format(python_path), shell=True)
|
||||
!= 0
|
||||
):
|
||||
raise RuntimeError(
|
||||
"Cannot compile pse: {}, if your system is windows, you need to install all the default components of `desktop development using C++` in visual studio 2019+".format(
|
||||
os.path.dirname(os.path.realpath(__file__))
|
||||
)
|
||||
)
|
||||
os.chdir(ori_path)
|
||||
|
||||
from .pse import pse
|
||||
70
ppocr/postprocess/pse_postprocess/pse/pse.pyx
Normal file
70
ppocr/postprocess/pse_postprocess/pse/pse.pyx
Normal file
@@ -0,0 +1,70 @@
|
||||
|
||||
import numpy as np
|
||||
import cv2
|
||||
cimport numpy as np
|
||||
cimport cython
|
||||
cimport libcpp
|
||||
cimport libcpp.pair
|
||||
cimport libcpp.queue
|
||||
from libcpp.pair cimport *
|
||||
from libcpp.queue cimport *
|
||||
|
||||
@cython.boundscheck(False)
|
||||
@cython.wraparound(False)
|
||||
cdef np.ndarray[np.int32_t, ndim=2] _pse(np.ndarray[np.uint8_t, ndim=3] kernels,
|
||||
np.ndarray[np.int32_t, ndim=2] label,
|
||||
int kernel_num,
|
||||
int label_num,
|
||||
float min_area=0):
|
||||
cdef np.ndarray[np.int32_t, ndim=2] pred
|
||||
pred = np.zeros((label.shape[0], label.shape[1]), dtype=np.int32)
|
||||
|
||||
for label_idx in range(1, label_num):
|
||||
if np.sum(label == label_idx) < min_area:
|
||||
label[label == label_idx] = 0
|
||||
|
||||
cdef libcpp.queue.queue[libcpp.pair.pair[np.int16_t,np.int16_t]] que = \
|
||||
queue[libcpp.pair.pair[np.int16_t,np.int16_t]]()
|
||||
cdef libcpp.queue.queue[libcpp.pair.pair[np.int16_t,np.int16_t]] nxt_que = \
|
||||
queue[libcpp.pair.pair[np.int16_t,np.int16_t]]()
|
||||
cdef np.int16_t* dx = [-1, 1, 0, 0]
|
||||
cdef np.int16_t* dy = [0, 0, -1, 1]
|
||||
cdef np.int16_t tmpx, tmpy
|
||||
|
||||
points = np.array(np.where(label > 0)).transpose((1, 0))
|
||||
for point_idx in range(points.shape[0]):
|
||||
tmpx, tmpy = points[point_idx, 0], points[point_idx, 1]
|
||||
que.push(pair[np.int16_t,np.int16_t](tmpx, tmpy))
|
||||
pred[tmpx, tmpy] = label[tmpx, tmpy]
|
||||
|
||||
cdef libcpp.pair.pair[np.int16_t,np.int16_t] cur
|
||||
cdef int cur_label
|
||||
for kernel_idx in range(kernel_num - 1, -1, -1):
|
||||
while not que.empty():
|
||||
cur = que.front()
|
||||
que.pop()
|
||||
cur_label = pred[cur.first, cur.second]
|
||||
|
||||
is_edge = True
|
||||
for j in range(4):
|
||||
tmpx = cur.first + dx[j]
|
||||
tmpy = cur.second + dy[j]
|
||||
if tmpx < 0 or tmpx >= label.shape[0] or tmpy < 0 or tmpy >= label.shape[1]:
|
||||
continue
|
||||
if kernels[kernel_idx, tmpx, tmpy] == 0 or pred[tmpx, tmpy] > 0:
|
||||
continue
|
||||
|
||||
que.push(pair[np.int16_t,np.int16_t](tmpx, tmpy))
|
||||
pred[tmpx, tmpy] = cur_label
|
||||
is_edge = False
|
||||
if is_edge:
|
||||
nxt_que.push(cur)
|
||||
|
||||
que, nxt_que = nxt_que, que
|
||||
|
||||
return pred
|
||||
|
||||
def pse(kernels, min_area):
|
||||
kernel_num = kernels.shape[0]
|
||||
label_num, label = cv2.connectedComponents(kernels[-1], connectivity=4)
|
||||
return _pse(kernels[:-1], label, kernel_num, label_num, min_area)
|
||||
18
ppocr/postprocess/pse_postprocess/pse/setup.py
Normal file
18
ppocr/postprocess/pse_postprocess/pse/setup.py
Normal file
@@ -0,0 +1,18 @@
|
||||
from distutils.core import setup, Extension
|
||||
from Cython.Build import cythonize
|
||||
import numpy
|
||||
|
||||
setup(
|
||||
ext_modules=cythonize(
|
||||
Extension(
|
||||
"pse",
|
||||
sources=["pse.pyx"],
|
||||
language="c++",
|
||||
include_dirs=[numpy.get_include()],
|
||||
library_dirs=[],
|
||||
libraries=[],
|
||||
extra_compile_args=["-O3"],
|
||||
extra_link_args=[],
|
||||
)
|
||||
)
|
||||
)
|
||||
122
ppocr/postprocess/pse_postprocess/pse_postprocess.py
Executable file
122
ppocr/postprocess/pse_postprocess/pse_postprocess.py
Executable file
@@ -0,0 +1,122 @@
|
||||
# copyright (c) 2021 PaddlePaddle Authors. All Rights Reserve.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""
|
||||
This code is refer from:
|
||||
https://github.com/whai362/PSENet/blob/python3/models/head/psenet_head.py
|
||||
"""
|
||||
|
||||
from __future__ import absolute_import
|
||||
from __future__ import division
|
||||
from __future__ import print_function
|
||||
|
||||
import numpy as np
|
||||
import cv2
|
||||
import paddle
|
||||
from paddle.nn import functional as F
|
||||
|
||||
from ppocr.postprocess.pse_postprocess.pse import pse
|
||||
|
||||
|
||||
class PSEPostProcess(object):
|
||||
"""
|
||||
The post process for PSE.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
thresh=0.5,
|
||||
box_thresh=0.85,
|
||||
min_area=16,
|
||||
box_type="quad",
|
||||
scale=4,
|
||||
**kwargs,
|
||||
):
|
||||
assert box_type in ["quad", "poly"], "Only quad and poly is supported"
|
||||
self.thresh = thresh
|
||||
self.box_thresh = box_thresh
|
||||
self.min_area = min_area
|
||||
self.box_type = box_type
|
||||
self.scale = scale
|
||||
|
||||
def __call__(self, outs_dict, shape_list):
|
||||
pred = outs_dict["maps"]
|
||||
if not isinstance(pred, paddle.Tensor):
|
||||
pred = paddle.to_tensor(pred)
|
||||
pred = F.interpolate(pred, scale_factor=4 // self.scale, mode="bilinear")
|
||||
|
||||
score = F.sigmoid(pred[:, 0, :, :])
|
||||
|
||||
kernels = (pred > self.thresh).astype("float32")
|
||||
text_mask = kernels[:, 0, :, :]
|
||||
text_mask = paddle.unsqueeze(text_mask, axis=1)
|
||||
|
||||
kernels[:, 0:, :, :] = kernels[:, 0:, :, :] * text_mask
|
||||
|
||||
score = score.numpy()
|
||||
kernels = kernels.numpy().astype(np.uint8)
|
||||
|
||||
boxes_batch = []
|
||||
for batch_index in range(pred.shape[0]):
|
||||
boxes, scores = self.boxes_from_bitmap(
|
||||
score[batch_index], kernels[batch_index], shape_list[batch_index]
|
||||
)
|
||||
|
||||
boxes_batch.append({"points": boxes, "scores": scores})
|
||||
return boxes_batch
|
||||
|
||||
def boxes_from_bitmap(self, score, kernels, shape):
|
||||
label = pse(kernels, self.min_area)
|
||||
return self.generate_box(score, label, shape)
|
||||
|
||||
def generate_box(self, score, label, shape):
|
||||
src_h, src_w, ratio_h, ratio_w = shape
|
||||
label_num = np.max(label) + 1
|
||||
|
||||
boxes = []
|
||||
scores = []
|
||||
for i in range(1, label_num):
|
||||
ind = label == i
|
||||
points = np.array(np.where(ind)).transpose((1, 0))[:, ::-1]
|
||||
|
||||
if points.shape[0] < self.min_area:
|
||||
label[ind] = 0
|
||||
continue
|
||||
|
||||
score_i = np.mean(score[ind])
|
||||
if score_i < self.box_thresh:
|
||||
label[ind] = 0
|
||||
continue
|
||||
|
||||
if self.box_type == "quad":
|
||||
rect = cv2.minAreaRect(points)
|
||||
bbox = cv2.boxPoints(rect)
|
||||
elif self.box_type == "poly":
|
||||
box_height = np.max(points[:, 1]) + 10
|
||||
box_width = np.max(points[:, 0]) + 10
|
||||
|
||||
mask = np.zeros((box_height, box_width), np.uint8)
|
||||
mask[points[:, 1], points[:, 0]] = 255
|
||||
|
||||
contours, _ = cv2.findContours(
|
||||
mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE
|
||||
)
|
||||
bbox = np.squeeze(contours[0], 1)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
bbox[:, 0] = np.clip(np.round(bbox[:, 0] / ratio_w), 0, src_w)
|
||||
bbox[:, 1] = np.clip(np.round(bbox[:, 1] / ratio_h), 0, src_h)
|
||||
boxes.append(bbox)
|
||||
scores.append(score_i)
|
||||
return boxes, scores
|
||||
1551
ppocr/postprocess/rec_postprocess.py
Normal file
1551
ppocr/postprocess/rec_postprocess.py
Normal file
File diff suppressed because it is too large
Load Diff
371
ppocr/postprocess/sast_postprocess.py
Executable file
371
ppocr/postprocess/sast_postprocess.py
Executable file
@@ -0,0 +1,371 @@
|
||||
# Copyright (c) 2020 PaddlePaddle Authors. All Rights Reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from __future__ import absolute_import
|
||||
from __future__ import division
|
||||
from __future__ import print_function
|
||||
|
||||
import os
|
||||
import sys
|
||||
|
||||
__dir__ = os.path.dirname(__file__)
|
||||
sys.path.append(__dir__)
|
||||
sys.path.append(os.path.join(__dir__, ".."))
|
||||
|
||||
import numpy as np
|
||||
from .locality_aware_nms import nms_locality
|
||||
import paddle
|
||||
import cv2
|
||||
import time
|
||||
|
||||
|
||||
class SASTPostProcess(object):
|
||||
"""
|
||||
The post process for SAST.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
score_thresh=0.5,
|
||||
nms_thresh=0.2,
|
||||
sample_pts_num=2,
|
||||
shrink_ratio_of_width=0.3,
|
||||
expand_scale=1.0,
|
||||
tcl_map_thresh=0.5,
|
||||
**kwargs,
|
||||
):
|
||||
self.score_thresh = score_thresh
|
||||
self.nms_thresh = nms_thresh
|
||||
self.sample_pts_num = sample_pts_num
|
||||
self.shrink_ratio_of_width = shrink_ratio_of_width
|
||||
self.expand_scale = expand_scale
|
||||
self.tcl_map_thresh = tcl_map_thresh
|
||||
|
||||
# c++ la-nms is faster, but only support python 3.5
|
||||
self.is_python35 = False
|
||||
if sys.version_info.major == 3 and sys.version_info.minor == 5:
|
||||
self.is_python35 = True
|
||||
|
||||
def point_pair2poly(self, point_pair_list):
|
||||
"""
|
||||
Transfer vertical point_pairs into poly point in clockwise.
|
||||
"""
|
||||
# construct poly
|
||||
point_num = len(point_pair_list) * 2
|
||||
point_list = [0] * point_num
|
||||
for idx, point_pair in enumerate(point_pair_list):
|
||||
point_list[idx] = point_pair[0]
|
||||
point_list[point_num - 1 - idx] = point_pair[1]
|
||||
return np.array(point_list).reshape(-1, 2)
|
||||
|
||||
def shrink_quad_along_width(self, quad, begin_width_ratio=0.0, end_width_ratio=1.0):
|
||||
"""
|
||||
Generate shrink_quad_along_width.
|
||||
"""
|
||||
ratio_pair = np.array(
|
||||
[[begin_width_ratio], [end_width_ratio]], dtype=np.float32
|
||||
)
|
||||
p0_1 = quad[0] + (quad[1] - quad[0]) * ratio_pair
|
||||
p3_2 = quad[3] + (quad[2] - quad[3]) * ratio_pair
|
||||
return np.array([p0_1[0], p0_1[1], p3_2[1], p3_2[0]])
|
||||
|
||||
def expand_poly_along_width(self, poly, shrink_ratio_of_width=0.3):
|
||||
"""
|
||||
expand poly along width.
|
||||
"""
|
||||
point_num = poly.shape[0]
|
||||
left_quad = np.array([poly[0], poly[1], poly[-2], poly[-1]], dtype=np.float32)
|
||||
left_ratio = (
|
||||
-shrink_ratio_of_width
|
||||
* np.linalg.norm(left_quad[0] - left_quad[3])
|
||||
/ (np.linalg.norm(left_quad[0] - left_quad[1]) + 1e-6)
|
||||
)
|
||||
left_quad_expand = self.shrink_quad_along_width(left_quad, left_ratio, 1.0)
|
||||
right_quad = np.array(
|
||||
[
|
||||
poly[point_num // 2 - 2],
|
||||
poly[point_num // 2 - 1],
|
||||
poly[point_num // 2],
|
||||
poly[point_num // 2 + 1],
|
||||
],
|
||||
dtype=np.float32,
|
||||
)
|
||||
right_ratio = 1.0 + shrink_ratio_of_width * np.linalg.norm(
|
||||
right_quad[0] - right_quad[3]
|
||||
) / (np.linalg.norm(right_quad[0] - right_quad[1]) + 1e-6)
|
||||
right_quad_expand = self.shrink_quad_along_width(right_quad, 0.0, right_ratio)
|
||||
poly[0] = left_quad_expand[0]
|
||||
poly[-1] = left_quad_expand[-1]
|
||||
poly[point_num // 2 - 1] = right_quad_expand[1]
|
||||
poly[point_num // 2] = right_quad_expand[2]
|
||||
return poly
|
||||
|
||||
def restore_quad(self, tcl_map, tcl_map_thresh, tvo_map):
|
||||
"""Restore quad."""
|
||||
xy_text = np.argwhere(tcl_map[:, :, 0] > tcl_map_thresh)
|
||||
xy_text = xy_text[:, ::-1] # (n, 2)
|
||||
|
||||
# Sort the text boxes via the y axis
|
||||
xy_text = xy_text[np.argsort(xy_text[:, 1])]
|
||||
|
||||
scores = tcl_map[xy_text[:, 1], xy_text[:, 0], 0]
|
||||
scores = scores[:, np.newaxis]
|
||||
|
||||
# Restore
|
||||
point_num = int(tvo_map.shape[-1] / 2)
|
||||
assert point_num == 4
|
||||
tvo_map = tvo_map[xy_text[:, 1], xy_text[:, 0], :]
|
||||
xy_text_tile = np.tile(xy_text, (1, point_num)) # (n, point_num * 2)
|
||||
quads = xy_text_tile - tvo_map
|
||||
|
||||
return scores, quads, xy_text
|
||||
|
||||
def quad_area(self, quad):
|
||||
"""
|
||||
compute area of a quad.
|
||||
"""
|
||||
edge = [
|
||||
(quad[1][0] - quad[0][0]) * (quad[1][1] + quad[0][1]),
|
||||
(quad[2][0] - quad[1][0]) * (quad[2][1] + quad[1][1]),
|
||||
(quad[3][0] - quad[2][0]) * (quad[3][1] + quad[2][1]),
|
||||
(quad[0][0] - quad[3][0]) * (quad[0][1] + quad[3][1]),
|
||||
]
|
||||
return np.sum(edge) / 2.0
|
||||
|
||||
def nms(self, dets):
|
||||
if self.is_python35:
|
||||
from ppocr.utils.utility import check_install
|
||||
|
||||
check_install("lanms", "lanms-nova")
|
||||
import lanms
|
||||
|
||||
dets = lanms.merge_quadrangle_n9(dets, self.nms_thresh)
|
||||
else:
|
||||
dets = nms_locality(dets, self.nms_thresh)
|
||||
return dets
|
||||
|
||||
def cluster_by_quads_tco(self, tcl_map, tcl_map_thresh, quads, tco_map):
|
||||
"""
|
||||
Cluster pixels in tcl_map based on quads.
|
||||
"""
|
||||
instance_count = quads.shape[0] + 1 # contain background
|
||||
instance_label_map = np.zeros(tcl_map.shape[:2], dtype=np.int32)
|
||||
if instance_count == 1:
|
||||
return instance_count, instance_label_map
|
||||
|
||||
# predict text center
|
||||
xy_text = np.argwhere(tcl_map[:, :, 0] > tcl_map_thresh)
|
||||
n = xy_text.shape[0]
|
||||
xy_text = xy_text[:, ::-1] # (n, 2)
|
||||
tco = tco_map[xy_text[:, 1], xy_text[:, 0], :] # (n, 2)
|
||||
pred_tc = xy_text - tco
|
||||
|
||||
# get gt text center
|
||||
m = quads.shape[0]
|
||||
gt_tc = np.mean(quads, axis=1) # (m, 2)
|
||||
|
||||
pred_tc_tile = np.tile(pred_tc[:, np.newaxis, :], (1, m, 1)) # (n, m, 2)
|
||||
gt_tc_tile = np.tile(gt_tc[np.newaxis, :, :], (n, 1, 1)) # (n, m, 2)
|
||||
dist_mat = np.linalg.norm(pred_tc_tile - gt_tc_tile, axis=2) # (n, m)
|
||||
xy_text_assign = np.argmin(dist_mat, axis=1) + 1 # (n,)
|
||||
|
||||
instance_label_map[xy_text[:, 1], xy_text[:, 0]] = xy_text_assign
|
||||
return instance_count, instance_label_map
|
||||
|
||||
def estimate_sample_pts_num(self, quad, xy_text):
|
||||
"""
|
||||
Estimate sample points number.
|
||||
"""
|
||||
eh = (
|
||||
np.linalg.norm(quad[0] - quad[3]) + np.linalg.norm(quad[1] - quad[2])
|
||||
) / 2.0
|
||||
ew = (
|
||||
np.linalg.norm(quad[0] - quad[1]) + np.linalg.norm(quad[2] - quad[3])
|
||||
) / 2.0
|
||||
|
||||
dense_sample_pts_num = max(2, int(ew))
|
||||
dense_xy_center_line = xy_text[
|
||||
np.linspace(
|
||||
0,
|
||||
xy_text.shape[0] - 1,
|
||||
dense_sample_pts_num,
|
||||
endpoint=True,
|
||||
dtype=np.float32,
|
||||
).astype(np.int32)
|
||||
]
|
||||
|
||||
dense_xy_center_line_diff = dense_xy_center_line[1:] - dense_xy_center_line[:-1]
|
||||
estimate_arc_len = np.sum(np.linalg.norm(dense_xy_center_line_diff, axis=1))
|
||||
|
||||
sample_pts_num = max(2, int(estimate_arc_len / eh))
|
||||
return sample_pts_num
|
||||
|
||||
def detect_sast(
|
||||
self,
|
||||
tcl_map,
|
||||
tvo_map,
|
||||
tbo_map,
|
||||
tco_map,
|
||||
ratio_w,
|
||||
ratio_h,
|
||||
src_w,
|
||||
src_h,
|
||||
shrink_ratio_of_width=0.3,
|
||||
tcl_map_thresh=0.5,
|
||||
offset_expand=1.0,
|
||||
out_strid=4.0,
|
||||
):
|
||||
"""
|
||||
first resize the tcl_map, tvo_map and tbo_map to the input_size, then restore the polys
|
||||
"""
|
||||
# restore quad
|
||||
scores, quads, xy_text = self.restore_quad(tcl_map, tcl_map_thresh, tvo_map)
|
||||
dets = np.hstack((quads, scores)).astype(np.float32, copy=False)
|
||||
dets = self.nms(dets)
|
||||
if dets.shape[0] == 0:
|
||||
return []
|
||||
quads = dets[:, :-1].reshape(-1, 4, 2)
|
||||
|
||||
# Compute quad area
|
||||
quad_areas = []
|
||||
for quad in quads:
|
||||
quad_areas.append(-self.quad_area(quad))
|
||||
|
||||
# instance segmentation
|
||||
# instance_count, instance_label_map = cv2.connectedComponents(tcl_map.astype(np.uint8), connectivity=8)
|
||||
instance_count, instance_label_map = self.cluster_by_quads_tco(
|
||||
tcl_map, tcl_map_thresh, quads, tco_map
|
||||
)
|
||||
|
||||
# restore single poly with tcl instance.
|
||||
poly_list = []
|
||||
for instance_idx in range(1, instance_count):
|
||||
xy_text = np.argwhere(instance_label_map == instance_idx)[:, ::-1]
|
||||
quad = quads[instance_idx - 1]
|
||||
q_area = quad_areas[instance_idx - 1]
|
||||
if q_area < 5:
|
||||
continue
|
||||
|
||||
#
|
||||
len1 = float(np.linalg.norm(quad[0] - quad[1]))
|
||||
len2 = float(np.linalg.norm(quad[1] - quad[2]))
|
||||
min_len = min(len1, len2)
|
||||
if min_len < 3:
|
||||
continue
|
||||
|
||||
# filter small CC
|
||||
if xy_text.shape[0] <= 0:
|
||||
continue
|
||||
|
||||
# filter low confidence instance
|
||||
xy_text_scores = tcl_map[xy_text[:, 1], xy_text[:, 0], 0]
|
||||
if np.sum(xy_text_scores) / quad_areas[instance_idx - 1] < 0.1:
|
||||
# if np.sum(xy_text_scores) / quad_areas[instance_idx - 1] < 0.05:
|
||||
continue
|
||||
|
||||
# sort xy_text
|
||||
left_center_pt = np.array(
|
||||
[[(quad[0, 0] + quad[-1, 0]) / 2.0, (quad[0, 1] + quad[-1, 1]) / 2.0]]
|
||||
) # (1, 2)
|
||||
right_center_pt = np.array(
|
||||
[[(quad[1, 0] + quad[2, 0]) / 2.0, (quad[1, 1] + quad[2, 1]) / 2.0]]
|
||||
) # (1, 2)
|
||||
proj_unit_vec = (right_center_pt - left_center_pt) / (
|
||||
np.linalg.norm(right_center_pt - left_center_pt) + 1e-6
|
||||
)
|
||||
proj_value = np.sum(xy_text * proj_unit_vec, axis=1)
|
||||
xy_text = xy_text[np.argsort(proj_value)]
|
||||
|
||||
# Sample pts in tcl map
|
||||
if self.sample_pts_num == 0:
|
||||
sample_pts_num = self.estimate_sample_pts_num(quad, xy_text)
|
||||
else:
|
||||
sample_pts_num = self.sample_pts_num
|
||||
xy_center_line = xy_text[
|
||||
np.linspace(
|
||||
0,
|
||||
xy_text.shape[0] - 1,
|
||||
sample_pts_num,
|
||||
endpoint=True,
|
||||
dtype=np.float32,
|
||||
).astype(np.int32)
|
||||
]
|
||||
|
||||
point_pair_list = []
|
||||
for x, y in xy_center_line:
|
||||
# get corresponding offset
|
||||
offset = tbo_map[y, x, :].reshape(2, 2)
|
||||
if offset_expand != 1.0:
|
||||
offset_length = np.linalg.norm(offset, axis=1, keepdims=True)
|
||||
expand_length = np.clip(
|
||||
offset_length * (offset_expand - 1), a_min=0.5, a_max=3.0
|
||||
)
|
||||
offset_detal = offset / offset_length * expand_length
|
||||
offset = offset + offset_detal
|
||||
# original point
|
||||
ori_yx = np.array([y, x], dtype=np.float32)
|
||||
point_pair = (
|
||||
(ori_yx + offset)[:, ::-1]
|
||||
* out_strid
|
||||
/ np.array([ratio_w, ratio_h]).reshape(-1, 2)
|
||||
)
|
||||
point_pair_list.append(point_pair)
|
||||
|
||||
# ndarry: (x, 2), expand poly along width
|
||||
detected_poly = self.point_pair2poly(point_pair_list)
|
||||
detected_poly = self.expand_poly_along_width(
|
||||
detected_poly, shrink_ratio_of_width
|
||||
)
|
||||
detected_poly[:, 0] = np.clip(detected_poly[:, 0], a_min=0, a_max=src_w)
|
||||
detected_poly[:, 1] = np.clip(detected_poly[:, 1], a_min=0, a_max=src_h)
|
||||
poly_list.append(detected_poly)
|
||||
|
||||
return poly_list
|
||||
|
||||
def __call__(self, outs_dict, shape_list):
|
||||
score_list = outs_dict["f_score"]
|
||||
border_list = outs_dict["f_border"]
|
||||
tvo_list = outs_dict["f_tvo"]
|
||||
tco_list = outs_dict["f_tco"]
|
||||
if isinstance(score_list, paddle.Tensor):
|
||||
score_list = score_list.numpy()
|
||||
border_list = border_list.numpy()
|
||||
tvo_list = tvo_list.numpy()
|
||||
tco_list = tco_list.numpy()
|
||||
|
||||
img_num = len(shape_list)
|
||||
poly_lists = []
|
||||
for ino in range(img_num):
|
||||
p_score = score_list[ino].transpose((1, 2, 0))
|
||||
p_border = border_list[ino].transpose((1, 2, 0))
|
||||
p_tvo = tvo_list[ino].transpose((1, 2, 0))
|
||||
p_tco = tco_list[ino].transpose((1, 2, 0))
|
||||
src_h, src_w, ratio_h, ratio_w = shape_list[ino]
|
||||
|
||||
poly_list = self.detect_sast(
|
||||
p_score,
|
||||
p_tvo,
|
||||
p_border,
|
||||
p_tco,
|
||||
ratio_w,
|
||||
ratio_h,
|
||||
src_w,
|
||||
src_h,
|
||||
shrink_ratio_of_width=self.shrink_ratio_of_width,
|
||||
tcl_map_thresh=self.tcl_map_thresh,
|
||||
offset_expand=self.expand_scale,
|
||||
)
|
||||
poly_lists.append({"points": np.array(poly_list)})
|
||||
|
||||
return poly_lists
|
||||
191
ppocr/postprocess/table_postprocess.py
Normal file
191
ppocr/postprocess/table_postprocess.py
Normal file
@@ -0,0 +1,191 @@
|
||||
# Copyright (c) 2021 PaddlePaddle Authors. All Rights Reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import numpy as np
|
||||
import paddle
|
||||
|
||||
from .rec_postprocess import AttnLabelDecode
|
||||
|
||||
|
||||
class TableLabelDecode(AttnLabelDecode):
|
||||
""" """
|
||||
|
||||
def __init__(self, character_dict_path, merge_no_span_structure=False, **kwargs):
|
||||
dict_character = []
|
||||
with open(character_dict_path, "rb") as fin:
|
||||
lines = fin.readlines()
|
||||
for line in lines:
|
||||
line = line.decode("utf-8").strip("\n").strip("\r\n")
|
||||
dict_character.append(line)
|
||||
|
||||
if merge_no_span_structure:
|
||||
if "<td></td>" not in dict_character:
|
||||
dict_character.append("<td></td>")
|
||||
if "<td>" in dict_character:
|
||||
dict_character.remove("<td>")
|
||||
|
||||
dict_character = self.add_special_char(dict_character)
|
||||
self.dict = {}
|
||||
for i, char in enumerate(dict_character):
|
||||
self.dict[char] = i
|
||||
self.character = dict_character
|
||||
self.td_token = ["<td>", "<td", "<td></td>"]
|
||||
|
||||
def __call__(self, preds, batch=None):
|
||||
structure_probs = preds["structure_probs"]
|
||||
bbox_preds = preds["loc_preds"]
|
||||
if isinstance(structure_probs, paddle.Tensor):
|
||||
structure_probs = structure_probs.numpy()
|
||||
if isinstance(bbox_preds, paddle.Tensor):
|
||||
bbox_preds = bbox_preds.numpy()
|
||||
shape_list = batch[-1]
|
||||
result = self.decode(structure_probs, bbox_preds, shape_list)
|
||||
if len(batch) == 1: # only contains shape
|
||||
return result
|
||||
|
||||
label_decode_result = self.decode_label(batch)
|
||||
return result, label_decode_result
|
||||
|
||||
def decode(self, structure_probs, bbox_preds, shape_list):
|
||||
"""convert text-label into text-index."""
|
||||
ignored_tokens = self.get_ignored_tokens()
|
||||
end_idx = self.dict[self.end_str]
|
||||
|
||||
structure_idx = structure_probs.argmax(axis=2)
|
||||
structure_probs = structure_probs.max(axis=2)
|
||||
|
||||
structure_batch_list = []
|
||||
bbox_batch_list = []
|
||||
batch_size = len(structure_idx)
|
||||
for batch_idx in range(batch_size):
|
||||
structure_list = []
|
||||
bbox_list = []
|
||||
score_list = []
|
||||
for idx in range(len(structure_idx[batch_idx])):
|
||||
char_idx = int(structure_idx[batch_idx][idx])
|
||||
if idx > 0 and char_idx == end_idx:
|
||||
break
|
||||
if char_idx in ignored_tokens:
|
||||
continue
|
||||
text = self.character[char_idx]
|
||||
if text in self.td_token:
|
||||
bbox = bbox_preds[batch_idx, idx]
|
||||
bbox = self._bbox_decode(bbox, shape_list[batch_idx])
|
||||
bbox_list.append(bbox)
|
||||
structure_list.append(text)
|
||||
score_list.append(structure_probs[batch_idx, idx])
|
||||
structure_batch_list.append([structure_list, np.mean(score_list)])
|
||||
bbox_batch_list.append(np.array(bbox_list))
|
||||
result = {
|
||||
"bbox_batch_list": bbox_batch_list,
|
||||
"structure_batch_list": structure_batch_list,
|
||||
}
|
||||
return result
|
||||
|
||||
def decode_label(self, batch):
|
||||
"""convert text-label into text-index."""
|
||||
structure_idx = batch[1]
|
||||
gt_bbox_list = batch[2]
|
||||
shape_list = batch[-1]
|
||||
ignored_tokens = self.get_ignored_tokens()
|
||||
end_idx = self.dict[self.end_str]
|
||||
|
||||
structure_batch_list = []
|
||||
bbox_batch_list = []
|
||||
batch_size = len(structure_idx)
|
||||
for batch_idx in range(batch_size):
|
||||
structure_list = []
|
||||
bbox_list = []
|
||||
for idx in range(len(structure_idx[batch_idx])):
|
||||
char_idx = int(structure_idx[batch_idx][idx])
|
||||
if idx > 0 and char_idx == end_idx:
|
||||
break
|
||||
if char_idx in ignored_tokens:
|
||||
continue
|
||||
structure_list.append(self.character[char_idx])
|
||||
|
||||
bbox = gt_bbox_list[batch_idx][idx]
|
||||
if bbox.sum() != 0:
|
||||
bbox = self._bbox_decode(bbox, shape_list[batch_idx])
|
||||
bbox_list.append(bbox)
|
||||
structure_batch_list.append(structure_list)
|
||||
bbox_batch_list.append(bbox_list)
|
||||
result = {
|
||||
"bbox_batch_list": bbox_batch_list,
|
||||
"structure_batch_list": structure_batch_list,
|
||||
}
|
||||
return result
|
||||
|
||||
def _bbox_decode(self, bbox, shape):
|
||||
h, w, ratio_h, ratio_w, pad_h, pad_w = shape
|
||||
h, w = pad_h, pad_w
|
||||
bbox[0::2] *= w
|
||||
bbox[1::2] *= h
|
||||
bbox[0::2] /= ratio_w
|
||||
bbox[1::2] /= ratio_h
|
||||
return bbox
|
||||
|
||||
|
||||
class TableMasterLabelDecode(TableLabelDecode):
|
||||
""" """
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
character_dict_path,
|
||||
box_shape="ori",
|
||||
merge_no_span_structure=True,
|
||||
**kwargs,
|
||||
):
|
||||
super(TableMasterLabelDecode, self).__init__(
|
||||
character_dict_path, merge_no_span_structure
|
||||
)
|
||||
self.box_shape = box_shape
|
||||
assert box_shape in [
|
||||
"ori",
|
||||
"pad",
|
||||
], "The shape used for box normalization must be ori or pad"
|
||||
|
||||
def add_special_char(self, dict_character):
|
||||
self.beg_str = "<SOS>"
|
||||
self.end_str = "<EOS>"
|
||||
self.unknown_str = "<UKN>"
|
||||
self.pad_str = "<PAD>"
|
||||
dict_character = dict_character
|
||||
dict_character = dict_character + [
|
||||
self.unknown_str,
|
||||
self.beg_str,
|
||||
self.end_str,
|
||||
self.pad_str,
|
||||
]
|
||||
return dict_character
|
||||
|
||||
def get_ignored_tokens(self):
|
||||
pad_idx = self.dict[self.pad_str]
|
||||
start_idx = self.dict[self.beg_str]
|
||||
end_idx = self.dict[self.end_str]
|
||||
unknown_idx = self.dict[self.unknown_str]
|
||||
return [start_idx, end_idx, pad_idx, unknown_idx]
|
||||
|
||||
def _bbox_decode(self, bbox, shape):
|
||||
h, w, ratio_h, ratio_w, pad_h, pad_w = shape
|
||||
if self.box_shape == "pad":
|
||||
h, w = pad_h, pad_w
|
||||
bbox[0::2] *= w
|
||||
bbox[1::2] *= h
|
||||
bbox[0::2] /= ratio_w
|
||||
bbox[1::2] /= ratio_h
|
||||
x, y, w, h = bbox
|
||||
x1, y1, x2, y2 = x - w // 2, y - h // 2, x + w // 2, y + h // 2
|
||||
bbox = np.array([x1, y1, x2, y2])
|
||||
return bbox
|
||||
96
ppocr/postprocess/vqa_token_re_layoutlm_postprocess.py
Normal file
96
ppocr/postprocess/vqa_token_re_layoutlm_postprocess.py
Normal file
@@ -0,0 +1,96 @@
|
||||
# Copyright (c) 2021 PaddlePaddle Authors. All Rights Reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
import paddle
|
||||
|
||||
|
||||
class VQAReTokenLayoutLMPostProcess(object):
|
||||
"""Convert between text-label and text-index"""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super(VQAReTokenLayoutLMPostProcess, self).__init__()
|
||||
|
||||
def __call__(self, preds, label=None, *args, **kwargs):
|
||||
pred_relations = preds["pred_relations"]
|
||||
if isinstance(preds["pred_relations"], paddle.Tensor):
|
||||
pred_relations = pred_relations.numpy()
|
||||
pred_relations = self.decode_pred(pred_relations)
|
||||
|
||||
if label is not None:
|
||||
return self._metric(pred_relations, label)
|
||||
else:
|
||||
return self._infer(pred_relations, *args, **kwargs)
|
||||
|
||||
def _metric(self, pred_relations, label):
|
||||
return pred_relations, label[-1], label[-2]
|
||||
|
||||
def _infer(self, pred_relations, *args, **kwargs):
|
||||
ser_results = kwargs["ser_results"]
|
||||
entity_idx_dict_batch = kwargs["entity_idx_dict_batch"]
|
||||
|
||||
# merge relations and ocr info
|
||||
results = []
|
||||
for pred_relation, ser_result, entity_idx_dict in zip(
|
||||
pred_relations, ser_results, entity_idx_dict_batch
|
||||
):
|
||||
result = []
|
||||
used_tail_id = []
|
||||
for relation in pred_relation:
|
||||
if relation["tail_id"] in used_tail_id:
|
||||
continue
|
||||
used_tail_id.append(relation["tail_id"])
|
||||
ocr_info_head = ser_result[entity_idx_dict[relation["head_id"]]]
|
||||
ocr_info_tail = ser_result[entity_idx_dict[relation["tail_id"]]]
|
||||
result.append((ocr_info_head, ocr_info_tail))
|
||||
results.append(result)
|
||||
return results
|
||||
|
||||
def decode_pred(self, pred_relations):
|
||||
pred_relations_new = []
|
||||
for pred_relation in pred_relations:
|
||||
pred_relation_new = []
|
||||
pred_relation = pred_relation[1 : pred_relation[0, 0, 0] + 1]
|
||||
for relation in pred_relation:
|
||||
relation_new = dict()
|
||||
relation_new["head_id"] = relation[0, 0]
|
||||
relation_new["head"] = tuple(relation[1])
|
||||
relation_new["head_type"] = relation[2, 0]
|
||||
relation_new["tail_id"] = relation[3, 0]
|
||||
relation_new["tail"] = tuple(relation[4])
|
||||
relation_new["tail_type"] = relation[5, 0]
|
||||
relation_new["type"] = relation[6, 0]
|
||||
pred_relation_new.append(relation_new)
|
||||
pred_relations_new.append(pred_relation_new)
|
||||
return pred_relations_new
|
||||
|
||||
|
||||
class DistillationRePostProcess(VQAReTokenLayoutLMPostProcess):
|
||||
"""
|
||||
DistillationRePostProcess
|
||||
"""
|
||||
|
||||
def __init__(self, model_name=["Student"], key=None, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
if not isinstance(model_name, list):
|
||||
model_name = [model_name]
|
||||
self.model_name = model_name
|
||||
self.key = key
|
||||
|
||||
def __call__(self, preds, *args, **kwargs):
|
||||
output = dict()
|
||||
for name in self.model_name:
|
||||
pred = preds[name]
|
||||
if self.key is not None:
|
||||
pred = pred[self.key]
|
||||
output[name] = super().__call__(pred, *args, **kwargs)
|
||||
return output
|
||||
116
ppocr/postprocess/vqa_token_ser_layoutlm_postprocess.py
Normal file
116
ppocr/postprocess/vqa_token_ser_layoutlm_postprocess.py
Normal file
@@ -0,0 +1,116 @@
|
||||
# Copyright (c) 2021 PaddlePaddle Authors. All Rights Reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
import numpy as np
|
||||
import paddle
|
||||
from ppocr.utils.utility import load_vqa_bio_label_maps
|
||||
|
||||
|
||||
class VQASerTokenLayoutLMPostProcess(object):
|
||||
"""Convert between text-label and text-index"""
|
||||
|
||||
def __init__(self, class_path, **kwargs):
|
||||
super(VQASerTokenLayoutLMPostProcess, self).__init__()
|
||||
label2id_map, self.id2label_map = load_vqa_bio_label_maps(class_path)
|
||||
|
||||
self.label2id_map_for_draw = dict()
|
||||
for key in label2id_map:
|
||||
if key.startswith("I-"):
|
||||
self.label2id_map_for_draw[key] = label2id_map["B" + key[1:]]
|
||||
else:
|
||||
self.label2id_map_for_draw[key] = label2id_map[key]
|
||||
|
||||
self.id2label_map_for_show = dict()
|
||||
for key in self.label2id_map_for_draw:
|
||||
val = self.label2id_map_for_draw[key]
|
||||
if key == "O":
|
||||
self.id2label_map_for_show[val] = key
|
||||
if key.startswith("B-") or key.startswith("I-"):
|
||||
self.id2label_map_for_show[val] = key[2:]
|
||||
else:
|
||||
self.id2label_map_for_show[val] = key
|
||||
|
||||
def __call__(self, preds, batch=None, *args, **kwargs):
|
||||
if isinstance(preds, tuple):
|
||||
preds = preds[0]
|
||||
if isinstance(preds, paddle.Tensor):
|
||||
preds = preds.numpy()
|
||||
|
||||
if batch is not None:
|
||||
return self._metric(preds, batch[5])
|
||||
else:
|
||||
return self._infer(preds, **kwargs)
|
||||
|
||||
def _metric(self, preds, label):
|
||||
pred_idxs = preds.argmax(axis=2)
|
||||
decode_out_list = [[] for _ in range(pred_idxs.shape[0])]
|
||||
label_decode_out_list = [[] for _ in range(pred_idxs.shape[0])]
|
||||
|
||||
for i in range(pred_idxs.shape[0]):
|
||||
for j in range(pred_idxs.shape[1]):
|
||||
if label[i, j] != -100:
|
||||
label_decode_out_list[i].append(self.id2label_map[label[i, j]])
|
||||
decode_out_list[i].append(self.id2label_map[pred_idxs[i, j]])
|
||||
return decode_out_list, label_decode_out_list
|
||||
|
||||
def _infer(self, preds, segment_offset_ids, ocr_infos):
|
||||
results = []
|
||||
|
||||
for pred, segment_offset_id, ocr_info in zip(
|
||||
preds, segment_offset_ids, ocr_infos
|
||||
):
|
||||
pred = np.argmax(pred, axis=1)
|
||||
pred = [self.id2label_map[idx] for idx in pred]
|
||||
|
||||
for idx in range(len(segment_offset_id)):
|
||||
if idx == 0:
|
||||
start_id = 0
|
||||
else:
|
||||
start_id = segment_offset_id[idx - 1]
|
||||
|
||||
end_id = segment_offset_id[idx]
|
||||
|
||||
curr_pred = pred[start_id:end_id]
|
||||
curr_pred = [self.label2id_map_for_draw[p] for p in curr_pred]
|
||||
|
||||
if len(curr_pred) <= 0:
|
||||
pred_id = 0
|
||||
else:
|
||||
counts = np.bincount(curr_pred)
|
||||
pred_id = np.argmax(counts)
|
||||
ocr_info[idx]["pred_id"] = int(pred_id)
|
||||
ocr_info[idx]["pred"] = self.id2label_map_for_show[int(pred_id)]
|
||||
results.append(ocr_info)
|
||||
return results
|
||||
|
||||
|
||||
class DistillationSerPostProcess(VQASerTokenLayoutLMPostProcess):
|
||||
"""
|
||||
DistillationSerPostProcess
|
||||
"""
|
||||
|
||||
def __init__(self, class_path, model_name=["Student"], key=None, **kwargs):
|
||||
super().__init__(class_path, **kwargs)
|
||||
if not isinstance(model_name, list):
|
||||
model_name = [model_name]
|
||||
self.model_name = model_name
|
||||
self.key = key
|
||||
|
||||
def __call__(self, preds, batch=None, *args, **kwargs):
|
||||
output = dict()
|
||||
for name in self.model_name:
|
||||
pred = preds[name]
|
||||
if self.key is not None:
|
||||
pred = pred[self.key]
|
||||
output[name] = super().__call__(pred, batch=batch, *args, **kwargs)
|
||||
return output
|
||||
Reference in New Issue
Block a user