first commit
Some checks are pending
Build/Publish Develop Docs / deploy (push) Waiting to run

This commit is contained in:
2025-07-02 08:57:16 +03:00
commit 56532cc9a9
1901 changed files with 457695 additions and 0 deletions

14
tools/__init__.py Normal file
View File

@@ -0,0 +1,14 @@
# Copyright (c) 2021 PaddlePaddle Authors. All Rights Reserved.
# Copyright 2018 The Google AI Language Team Authors and The HuggingFace Inc. team.
#
# 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.

View File

@@ -0,0 +1,103 @@
# Copyright (c) 2022 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 json
import os
def poly_to_string(poly):
if len(poly.shape) > 1:
poly = np.array(poly).flatten()
string = "\t".join(str(i) for i in poly)
return string
def convert_label(label_dir, mode="gt", save_dir="./save_results/"):
if not os.path.exists(label_dir):
raise ValueError(f"The file {label_dir} does not exist!")
assert label_dir != save_dir, "hahahhaha"
label_file = open(label_dir, "r")
data = label_file.readlines()
gt_dict = {}
for line in data:
try:
tmp = line.split("\t")
assert len(tmp) == 2, ""
except:
tmp = line.strip().split(" ")
gt_lists = []
if tmp[0].split("/")[0] is not None:
img_path = tmp[0]
anno = json.loads(tmp[1])
gt_collect = []
for dic in anno:
# txt = dic['transcription'].replace(' ', '') # ignore blank
txt = dic["transcription"]
if "score" in dic and float(dic["score"]) < 0.5:
continue
if "\u3000" in txt:
txt = txt.replace("\u3000", " ")
# while ' ' in txt:
# txt = txt.replace(' ', '')
poly = np.array(dic["points"]).flatten()
if txt == "###":
txt_tag = 1 ## ignore 1
else:
txt_tag = 0
if mode == "gt":
gt_label = (
poly_to_string(poly) + "\t" + str(txt_tag) + "\t" + txt + "\n"
)
else:
gt_label = poly_to_string(poly) + "\t" + txt + "\n"
gt_lists.append(gt_label)
gt_dict[img_path] = gt_lists
else:
continue
if not os.path.exists(save_dir):
os.makedirs(save_dir)
for img_name in gt_dict.keys():
save_name = img_name.split("/")[-1]
save_file = os.path.join(save_dir, save_name + ".txt")
with open(save_file, "w") as f:
f.writelines(gt_dict[img_name])
print("The convert label saved in {}".format(save_dir))
def parse_args():
import argparse
parser = argparse.ArgumentParser(description="args")
parser.add_argument("--label_path", type=str, required=True)
parser.add_argument("--save_folder", type=str, required=True)
parser.add_argument("--mode", type=str, default=False)
args = parser.parse_args()
return args
if __name__ == "__main__":
args = parse_args()
convert_label(args.label_path, args.mode, args.save_folder)

View File

@@ -0,0 +1,72 @@
# Copyright (c) 2022 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 os
import argparse
def str2bool(v):
return v.lower() in ("true", "t", "1")
def init_args():
parser = argparse.ArgumentParser()
parser.add_argument("--image_dir", type=str, default="")
parser.add_argument("--save_html_path", type=str, default="./default.html")
parser.add_argument("--width", type=int, default=640)
return parser
def parse_args():
parser = init_args()
return parser.parse_args()
def draw_debug_img(args):
html_path = args.save_html_path
err_cnt = 0
with open(html_path, "w") as html:
html.write("<html>\n<body>\n")
html.write('<table border="1">\n')
html.write(
'<meta http-equiv="Content-Type" content="text/html; charset=utf-8" />'
)
image_list = []
path = args.image_dir
for i, filename in enumerate(sorted(os.listdir(path))):
if filename.endswith("txt"):
continue
# The image path
base = "{}/{}".format(path, filename)
html.write("<tr>\n")
html.write(f"<td> {filename}\n GT")
html.write(f'<td>GT\n<img src="{base}" width={args.width}></td>')
html.write("</tr>\n")
html.write("<style>\n")
html.write("span {\n")
html.write(" color: red;\n")
html.write("}\n")
html.write("</style>\n")
html.write("</table>\n")
html.write("</html>\n</body>\n")
print(f"The html file saved in {html_path}")
return
if __name__ == "__main__":
args = parse_args()
draw_debug_img(args)

View File

@@ -0,0 +1,191 @@
# Copyright (c) 2022 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 os
import re
import sys
import shapely
from shapely.geometry import Polygon
import numpy as np
from collections import defaultdict
import operator
import editdistance
def strQ2B(ustring):
rstring = ""
for uchar in ustring:
inside_code = ord(uchar)
if inside_code == 12288:
inside_code = 32
elif inside_code >= 65281 and inside_code <= 65374:
inside_code -= 65248
rstring += chr(inside_code)
return rstring
def polygon_from_str(polygon_points):
"""
Create a shapely polygon object from gt or dt line.
"""
polygon_points = np.array(polygon_points).reshape(4, 2)
polygon = Polygon(polygon_points).convex_hull
return polygon
def polygon_iou(poly1, poly2):
"""
Intersection over union between two shapely polygons.
"""
if not poly1.intersects(poly2): # this test is fast and can accelerate calculation
iou = 0
else:
try:
inter_area = poly1.intersection(poly2).area
union_area = poly1.area + poly2.area - inter_area
iou = float(inter_area) / union_area
except shapely.geos.TopologicalError:
# except Exception as e:
# print(e)
print("shapely.geos.TopologicalError occurred, iou set to 0")
iou = 0
return iou
def ed(str1, str2):
return editdistance.eval(str1, str2)
def e2e_eval(gt_dir, res_dir, ignore_blank=False):
print("start testing...")
iou_thresh = 0.5
val_names = os.listdir(gt_dir)
num_gt_chars = 0
gt_count = 0
dt_count = 0
hit = 0
ed_sum = 0
for i, val_name in enumerate(val_names):
with open(os.path.join(gt_dir, val_name), encoding="utf-8") as f:
gt_lines = [o.strip() for o in f.readlines()]
gts = []
ignore_masks = []
for line in gt_lines:
parts = line.strip().split("\t")
# ignore illegal data
if len(parts) < 9:
continue
assert len(parts) < 11
if len(parts) == 9:
gts.append(parts[:8] + [""])
else:
gts.append(parts[:8] + [parts[-1]])
ignore_masks.append(parts[8])
val_path = os.path.join(res_dir, val_name)
if not os.path.exists(val_path):
dt_lines = []
else:
with open(val_path, encoding="utf-8") as f:
dt_lines = [o.strip() for o in f.readlines()]
dts = []
for line in dt_lines:
# print(line)
parts = line.strip().split("\t")
assert len(parts) < 10, "line error: {}".format(line)
if len(parts) == 8:
dts.append(parts + [""])
else:
dts.append(parts)
dt_match = [False] * len(dts)
gt_match = [False] * len(gts)
all_ious = defaultdict(tuple)
for index_gt, gt in enumerate(gts):
gt_coors = [float(gt_coor) for gt_coor in gt[0:8]]
gt_poly = polygon_from_str(gt_coors)
for index_dt, dt in enumerate(dts):
dt_coors = [float(dt_coor) for dt_coor in dt[0:8]]
dt_poly = polygon_from_str(dt_coors)
iou = polygon_iou(dt_poly, gt_poly)
if iou >= iou_thresh:
all_ious[(index_gt, index_dt)] = iou
sorted_ious = sorted(all_ious.items(), key=operator.itemgetter(1), reverse=True)
sorted_gt_dt_pairs = [item[0] for item in sorted_ious]
# matched gt and dt
for gt_dt_pair in sorted_gt_dt_pairs:
index_gt, index_dt = gt_dt_pair
if gt_match[index_gt] == False and dt_match[index_dt] == False:
gt_match[index_gt] = True
dt_match[index_dt] = True
if ignore_blank:
gt_str = strQ2B(gts[index_gt][8]).replace(" ", "")
dt_str = strQ2B(dts[index_dt][8]).replace(" ", "")
else:
gt_str = strQ2B(gts[index_gt][8])
dt_str = strQ2B(dts[index_dt][8])
if ignore_masks[index_gt] == "0":
ed_sum += ed(gt_str, dt_str)
num_gt_chars += len(gt_str)
if gt_str == dt_str:
hit += 1
gt_count += 1
dt_count += 1
# unmatched dt
for tindex, dt_match_flag in enumerate(dt_match):
if dt_match_flag == False:
dt_str = dts[tindex][8]
gt_str = ""
ed_sum += ed(dt_str, gt_str)
dt_count += 1
# unmatched gt
for tindex, gt_match_flag in enumerate(gt_match):
if gt_match_flag == False and ignore_masks[tindex] == "0":
dt_str = ""
gt_str = gts[tindex][8]
ed_sum += ed(gt_str, dt_str)
num_gt_chars += len(gt_str)
gt_count += 1
eps = 1e-9
print("hit, dt_count, gt_count", hit, dt_count, gt_count)
precision = hit / (dt_count + eps)
recall = hit / (gt_count + eps)
fmeasure = 2.0 * precision * recall / (precision + recall + eps)
avg_edit_dist_img = ed_sum / len(val_names)
avg_edit_dist_field = ed_sum / (gt_count + eps)
character_acc = 1 - ed_sum / (num_gt_chars + eps)
print("character_acc: %.2f" % (character_acc * 100) + "%")
print("avg_edit_dist_field: %.2f" % (avg_edit_dist_field))
print("avg_edit_dist_img: %.2f" % (avg_edit_dist_img))
print("precision: %.2f" % (precision * 100) + "%")
print("recall: %.2f" % (recall * 100) + "%")
print("fmeasure: %.2f" % (fmeasure * 100) + "%")
if __name__ == "__main__":
# if len(sys.argv) != 3:
# print("python3 ocr_e2e_eval.py gt_dir res_dir")
# exit(-1)
# gt_folder = sys.argv[1]
# pred_folder = sys.argv[2]
gt_folder = sys.argv[1]
pred_folder = sys.argv[2]
e2e_eval(gt_folder, pred_folder)

63
tools/end2end/readme.md Normal file
View File

@@ -0,0 +1,63 @@
# 简介
`tools/end2end`目录下存放了文本检测+文本识别pipeline串联预测的指标评测代码以及可视化工具。本节介绍文本检测+文本识别的端对端指标评估方式。
## 端对端评测步骤
**步骤一:**
运行`tools/infer/predict_system.py`,得到保存的结果:
```
python3 tools/infer/predict_system.py --det_model_dir=./ch_PP-OCRv2_det_infer/ --rec_model_dir=./ch_PP-OCRv2_rec_infer/ --image_dir=./datasets/img_dir/ --draw_img_save_dir=./ch_PP-OCRv2_results/ --is_visualize=True
```
文本检测识别可视化图默认保存在`./ch_PP-OCRv2_results/`目录下,预测结果默认保存在`./ch_PP-OCRv2_results/system_results.txt`中,格式如下:
```
all-sum-510/00224225.jpg [{"transcription": "超赞", "points": [[8.0, 48.0], [157.0, 44.0], [159.0, 115.0], [10.0, 119.0]], "score": "0.99396634"}, {"transcription": "中", "points": [[202.0, 152.0], [230.0, 152.0], [230.0, 163.0], [202.0, 163.0]], "score": "0.09310734"}, {"transcription": "58.0m", "points": [[196.0, 192.0], [444.0, 192.0], [444.0, 240.0], [196.0, 240.0]], "score": "0.44041982"}, {"transcription": "汽配", "points": [[55.0, 263.0], [95.0, 263.0], [95.0, 281.0], [55.0, 281.0]], "score": "0.9986651"}, {"transcription": "成总店", "points": [[120.0, 262.0], [176.0, 262.0], [176.0, 283.0], [120.0, 283.0]], "score": "0.9929402"}, {"transcription": "K", "points": [[237.0, 286.0], [311.0, 286.0], [311.0, 345.0], [237.0, 345.0]], "score": "0.6074794"}, {"transcription": "88-8", "points": [[203.0, 405.0], [477.0, 414.0], [475.0, 459.0], [201.0, 450.0]], "score": "0.7106863"}]
```
**步骤二:**
将步骤一保存的数据转换为端对端评测需要的数据格式:
修改 `tools/end2end/convert_ppocr_label.py`中的代码convert_label函数中设置输入标签路径Mode保存标签路径等对预测数据的GTlabel和预测结果的label格式进行转换。
```
python3 tools/end2end/convert_ppocr_label.py --mode=gt --label_path=path/to/label_txt --save_folder=save_gt_label
python3 tools/end2end/convert_ppocr_label.py --mode=pred --label_path=path/to/pred_txt --save_folder=save_PPOCRV2_infer
```
得到如下结果:
```
├── ./save_gt_label/
├── ./save_PPOCRV2_infer/
```
**步骤三:**
执行端对端评测,运行`tools/eval_end2end.py`计算端对端指标,运行方式如下:
```
python3 tools/eval_end2end.py "gt_label_dir" "predict_label_dir"
```
比如:
```
python3 tools/eval_end2end.py ./save_gt_label/ ./save_PPOCRV2_infer/
```
将得到如下结果fmeasure为主要关注的指标
```
hit, dt_count, gt_count 1557 2693 3283
character_acc: 61.77%
avg_edit_dist_field: 3.08
avg_edit_dist_img: 51.82
precision: 57.82%
recall: 47.43%
fmeasure: 52.11%
```

181
tools/eval.py Executable file
View File

@@ -0,0 +1,181 @@
# 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(os.path.abspath(__file__))
sys.path.insert(0, __dir__)
sys.path.insert(0, os.path.abspath(os.path.join(__dir__, "..")))
import paddle
from ppocr.data import build_dataloader, set_signal_handlers
from ppocr.modeling.architectures import build_model
from ppocr.postprocess import build_post_process
from ppocr.metrics import build_metric
from ppocr.utils.save_load import load_model
import tools.program as program
def main():
global_config = config["Global"]
# build dataloader
set_signal_handlers()
valid_dataloader = build_dataloader(config, "Eval", device, logger)
# build post process
post_process_class = build_post_process(config["PostProcess"], global_config)
# build model
# for rec algorithm
if hasattr(post_process_class, "character"):
char_num = len(getattr(post_process_class, "character"))
if config["Architecture"]["algorithm"] in [
"Distillation",
]: # distillation model
for key in config["Architecture"]["Models"]:
if (
config["Architecture"]["Models"][key]["Head"]["name"] == "MultiHead"
): # for multi head
out_channels_list = {}
if config["PostProcess"]["name"] == "DistillationSARLabelDecode":
char_num = char_num - 2
if config["PostProcess"]["name"] == "DistillationNRTRLabelDecode":
char_num = char_num - 3
out_channels_list["CTCLabelDecode"] = char_num
out_channels_list["SARLabelDecode"] = char_num + 2
out_channels_list["NRTRLabelDecode"] = char_num + 3
config["Architecture"]["Models"][key]["Head"][
"out_channels_list"
] = out_channels_list
else:
config["Architecture"]["Models"][key]["Head"][
"out_channels"
] = char_num
elif config["Architecture"]["Head"]["name"] == "MultiHead": # for multi head
out_channels_list = {}
if config["PostProcess"]["name"] == "SARLabelDecode":
char_num = char_num - 2
if config["PostProcess"]["name"] == "NRTRLabelDecode":
char_num = char_num - 3
out_channels_list["CTCLabelDecode"] = char_num
out_channels_list["SARLabelDecode"] = char_num + 2
out_channels_list["NRTRLabelDecode"] = char_num + 3
config["Architecture"]["Head"]["out_channels_list"] = out_channels_list
else: # base rec model
config["Architecture"]["Head"]["out_channels"] = char_num
model = build_model(config["Architecture"])
extra_input_models = [
"SRN",
"NRTR",
"SAR",
"SEED",
"SVTR",
"SVTR_LCNet",
"VisionLAN",
"RobustScanner",
"SVTR_HGNet",
]
extra_input = False
if config["Architecture"]["algorithm"] == "Distillation":
for key in config["Architecture"]["Models"]:
extra_input = (
extra_input
or config["Architecture"]["Models"][key]["algorithm"]
in extra_input_models
)
else:
extra_input = config["Architecture"]["algorithm"] in extra_input_models
if "model_type" in config["Architecture"].keys():
if config["Architecture"]["algorithm"] == "CAN":
model_type = "can"
elif config["Architecture"]["algorithm"] == "LaTeXOCR":
model_type = "latexocr"
config["Metric"]["cal_bleu_score"] = True
elif config["Architecture"]["algorithm"] == "UniMERNet":
model_type = "unimernet"
config["Metric"]["cal_bleu_score"] = True
elif config["Architecture"]["algorithm"] in [
"PP-FormulaNet-S",
"PP-FormulaNet-L",
"PP-FormulaNet_plus-S",
"PP-FormulaNet_plus-M",
"PP-FormulaNet_plus-L",
]:
model_type = "pp_formulanet"
config["Metric"]["cal_bleu_score"] = True
else:
model_type = config["Architecture"]["model_type"]
else:
model_type = None
# build metric
eval_class = build_metric(config["Metric"])
# amp
use_amp = config["Global"].get("use_amp", False)
amp_level = config["Global"].get("amp_level", "O2")
amp_custom_black_list = config["Global"].get("amp_custom_black_list", [])
if use_amp:
AMP_RELATED_FLAGS_SETTING = {
"FLAGS_cudnn_batchnorm_spatial_persistent": 1,
}
paddle.set_flags(AMP_RELATED_FLAGS_SETTING)
scale_loss = config["Global"].get("scale_loss", 1.0)
use_dynamic_loss_scaling = config["Global"].get(
"use_dynamic_loss_scaling", False
)
scaler = paddle.amp.GradScaler(
init_loss_scaling=scale_loss,
use_dynamic_loss_scaling=use_dynamic_loss_scaling,
)
if amp_level == "O2":
model = paddle.amp.decorate(
models=model, level=amp_level, master_weight=True
)
else:
scaler = None
best_model_dict = load_model(
config, model, model_type=config["Architecture"]["model_type"]
)
if len(best_model_dict):
logger.info("metric in ckpt ***************")
for k, v in best_model_dict.items():
logger.info("{}:{}".format(k, v))
# start eval
metric = program.eval(
model,
valid_dataloader,
post_process_class,
eval_class,
model_type,
extra_input,
scaler,
amp_level,
amp_custom_black_list,
)
logger.info("metric eval ***************")
for k, v in metric.items():
logger.info("{}:{}".format(k, v))
if __name__ == "__main__":
config, device, logger, vdl_writer = program.preprocess()
main()

77
tools/export_center.py Normal file
View File

@@ -0,0 +1,77 @@
# 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
import pickle
__dir__ = os.path.dirname(os.path.abspath(__file__))
sys.path.append(__dir__)
sys.path.append(os.path.abspath(os.path.join(__dir__, "..")))
from ppocr.data import build_dataloader, set_signal_handlers
from ppocr.modeling.architectures import build_model
from ppocr.postprocess import build_post_process
from ppocr.utils.save_load import load_model
from ppocr.utils.utility import print_dict
import tools.program as program
def main():
global_config = config["Global"]
# build dataloader
config["Eval"]["dataset"]["name"] = config["Train"]["dataset"]["name"]
config["Eval"]["dataset"]["data_dir"] = config["Train"]["dataset"]["data_dir"]
config["Eval"]["dataset"]["label_file_list"] = config["Train"]["dataset"][
"label_file_list"
]
set_signal_handlers()
eval_dataloader = build_dataloader(config, "Eval", device, logger)
# build post process
post_process_class = build_post_process(config["PostProcess"], global_config)
# build model
# for rec algorithm
if hasattr(post_process_class, "character"):
char_num = len(getattr(post_process_class, "character"))
config["Architecture"]["Head"]["out_channels"] = char_num
# set return_features = True
config["Architecture"]["Head"]["return_feats"] = True
model = build_model(config["Architecture"])
best_model_dict = load_model(config, model)
if len(best_model_dict):
logger.info("metric in ckpt ***************")
for k, v in best_model_dict.items():
logger.info("{}:{}".format(k, v))
# get features from train data
char_center = program.get_center(model, eval_dataloader, post_process_class)
# serialize to disk
with open("train_center.pkl", "wb") as f:
pickle.dump(char_center, f)
return
if __name__ == "__main__":
config, device, logger, vdl_writer = program.preprocess()
main()

37
tools/export_model.py Executable file
View File

@@ -0,0 +1,37 @@
# 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.
import os
import sys
__dir__ = os.path.dirname(os.path.abspath(__file__))
sys.path.append(__dir__)
sys.path.insert(0, os.path.abspath(os.path.join(__dir__, "..")))
import argparse
from tools.program import load_config, merge_config, ArgsParser
from ppocr.utils.export_model import export
def main():
FLAGS = ArgsParser().parse_args()
config = load_config(FLAGS.config)
config = merge_config(config, FLAGS.opt)
# export model
export(config)
if __name__ == "__main__":
main()

164
tools/infer/predict_cls.py Executable file
View File

@@ -0,0 +1,164 @@
# 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.
import os
import sys
__dir__ = os.path.dirname(os.path.abspath(__file__))
sys.path.append(__dir__)
sys.path.insert(0, os.path.abspath(os.path.join(__dir__, "../..")))
os.environ["FLAGS_allocator_strategy"] = "auto_growth"
import cv2
import copy
import numpy as np
import math
import time
import traceback
import tools.infer.utility as utility
from ppocr.postprocess import build_post_process
from ppocr.utils.logging import get_logger
from ppocr.utils.utility import get_image_file_list, check_and_read
logger = get_logger()
class TextClassifier(object):
def __init__(self, args):
if os.path.exists(f"{args.cls_model_dir}/inference.yml"):
model_config = utility.load_config(f"{args.cls_model_dir}/inference.yml")
model_name = model_config.get("Global", {}).get("model_name", "")
if model_name:
raise ValueError(
f"{model_name} is not supported. Please check if the model is supported by the PaddleOCR wheel."
)
self.cls_image_shape = [int(v) for v in args.cls_image_shape.split(",")]
self.cls_batch_num = args.cls_batch_num
self.cls_thresh = args.cls_thresh
postprocess_params = {
"name": "ClsPostProcess",
"label_list": args.label_list,
}
self.postprocess_op = build_post_process(postprocess_params)
(
self.predictor,
self.input_tensor,
self.output_tensors,
_,
) = utility.create_predictor(args, "cls", logger)
self.use_onnx = args.use_onnx
def resize_norm_img(self, img):
imgC, imgH, imgW = self.cls_image_shape
h = img.shape[0]
w = img.shape[1]
ratio = w / float(h)
if math.ceil(imgH * ratio) > imgW:
resized_w = imgW
else:
resized_w = int(math.ceil(imgH * ratio))
resized_image = cv2.resize(img, (resized_w, imgH))
resized_image = resized_image.astype("float32")
if self.cls_image_shape[0] == 1:
resized_image = resized_image / 255
resized_image = resized_image[np.newaxis, :]
else:
resized_image = resized_image.transpose((2, 0, 1)) / 255
resized_image -= 0.5
resized_image /= 0.5
padding_im = np.zeros((imgC, imgH, imgW), dtype=np.float32)
padding_im[:, :, 0:resized_w] = resized_image
return padding_im
def __call__(self, img_list):
img_list = copy.deepcopy(img_list)
img_num = len(img_list)
# Calculate the aspect ratio of all text bars
width_list = []
for img in img_list:
width_list.append(img.shape[1] / float(img.shape[0]))
# Sorting can speed up the cls process
indices = np.argsort(np.array(width_list))
cls_res = [["", 0.0]] * img_num
batch_num = self.cls_batch_num
elapse = 0
for beg_img_no in range(0, img_num, batch_num):
end_img_no = min(img_num, beg_img_no + batch_num)
norm_img_batch = []
max_wh_ratio = 0
starttime = time.time()
for ino in range(beg_img_no, end_img_no):
h, w = img_list[indices[ino]].shape[0:2]
wh_ratio = w * 1.0 / h
max_wh_ratio = max(max_wh_ratio, wh_ratio)
for ino in range(beg_img_no, end_img_no):
norm_img = self.resize_norm_img(img_list[indices[ino]])
norm_img = norm_img[np.newaxis, :]
norm_img_batch.append(norm_img)
norm_img_batch = np.concatenate(norm_img_batch)
norm_img_batch = norm_img_batch.copy()
if self.use_onnx:
input_dict = {}
input_dict[self.input_tensor.name] = norm_img_batch
outputs = self.predictor.run(self.output_tensors, input_dict)
prob_out = outputs[0]
else:
self.input_tensor.copy_from_cpu(norm_img_batch)
self.predictor.run()
prob_out = self.output_tensors[0].copy_to_cpu()
self.predictor.try_shrink_memory()
cls_result = self.postprocess_op(prob_out)
elapse += time.time() - starttime
for rno in range(len(cls_result)):
label, score = cls_result[rno]
cls_res[indices[beg_img_no + rno]] = [label, score]
if "180" in label and score > self.cls_thresh:
img_list[indices[beg_img_no + rno]] = cv2.rotate(
img_list[indices[beg_img_no + rno]], 1
)
return img_list, cls_res, elapse
def main(args):
image_file_list = get_image_file_list(args.image_dir)
text_classifier = TextClassifier(args)
valid_image_file_list = []
img_list = []
for image_file in image_file_list:
img, flag, _ = check_and_read(image_file)
if not flag:
img = cv2.imread(image_file)
if img is None:
logger.info("error in loading image:{}".format(image_file))
continue
valid_image_file_list.append(image_file)
img_list.append(img)
try:
img_list, cls_res, predict_time = text_classifier(img_list)
except Exception as E:
logger.info(traceback.format_exc())
logger.info(E)
exit()
for ino in range(len(img_list)):
logger.info(
"Predicts of {}:{}".format(valid_image_file_list[ino], cls_res[ino])
)
if __name__ == "__main__":
main(utility.parse_args())

501
tools/infer/predict_det.py Executable file
View File

@@ -0,0 +1,501 @@
# 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.
import os
import sys
__dir__ = os.path.dirname(os.path.abspath(__file__))
sys.path.append(__dir__)
sys.path.insert(0, os.path.abspath(os.path.join(__dir__, "../..")))
os.environ["FLAGS_allocator_strategy"] = "auto_growth"
import cv2
import numpy as np
import time
import sys
import tools.infer.utility as utility
from ppocr.utils.logging import get_logger
from ppocr.utils.utility import get_image_file_list, check_and_read
from ppocr.data import create_operators, transform
from ppocr.postprocess import build_post_process
import json
class TextDetector(object):
def __init__(self, args, logger=None):
if os.path.exists(f"{args.det_model_dir}/inference.yml"):
model_config = utility.load_config(f"{args.det_model_dir}/inference.yml")
model_name = model_config.get("Global", {}).get("model_name", "")
if model_name and model_name not in [
"PP-OCRv5_mobile_det",
"PP-OCRv5_server_det",
]:
raise ValueError(
f"{model_name} is not supported. Please check if the model is supported by the PaddleOCR wheel."
)
if logger is None:
logger = get_logger()
self.args = args
self.det_algorithm = args.det_algorithm
self.use_onnx = args.use_onnx
pre_process_list = [
{
"DetResizeForTest": {
"limit_side_len": args.det_limit_side_len,
"limit_type": args.det_limit_type,
}
},
{
"NormalizeImage": {
"std": [0.229, 0.224, 0.225],
"mean": [0.485, 0.456, 0.406],
"scale": "1./255.",
"order": "hwc",
}
},
{"ToCHWImage": None},
{"KeepKeys": {"keep_keys": ["image", "shape"]}},
]
postprocess_params = {}
if self.det_algorithm == "DB":
postprocess_params["name"] = "DBPostProcess"
postprocess_params["thresh"] = args.det_db_thresh
postprocess_params["box_thresh"] = args.det_db_box_thresh
postprocess_params["max_candidates"] = 1000
postprocess_params["unclip_ratio"] = args.det_db_unclip_ratio
postprocess_params["use_dilation"] = args.use_dilation
postprocess_params["score_mode"] = args.det_db_score_mode
postprocess_params["box_type"] = args.det_box_type
elif self.det_algorithm == "DB++":
postprocess_params["name"] = "DBPostProcess"
postprocess_params["thresh"] = args.det_db_thresh
postprocess_params["box_thresh"] = args.det_db_box_thresh
postprocess_params["max_candidates"] = 1000
postprocess_params["unclip_ratio"] = args.det_db_unclip_ratio
postprocess_params["use_dilation"] = args.use_dilation
postprocess_params["score_mode"] = args.det_db_score_mode
postprocess_params["box_type"] = args.det_box_type
pre_process_list[1] = {
"NormalizeImage": {
"std": [1.0, 1.0, 1.0],
"mean": [0.48109378172549, 0.45752457890196, 0.40787054090196],
"scale": "1./255.",
"order": "hwc",
}
}
elif self.det_algorithm == "EAST":
postprocess_params["name"] = "EASTPostProcess"
postprocess_params["score_thresh"] = args.det_east_score_thresh
postprocess_params["cover_thresh"] = args.det_east_cover_thresh
postprocess_params["nms_thresh"] = args.det_east_nms_thresh
elif self.det_algorithm == "SAST":
pre_process_list[0] = {
"DetResizeForTest": {"resize_long": args.det_limit_side_len}
}
postprocess_params["name"] = "SASTPostProcess"
postprocess_params["score_thresh"] = args.det_sast_score_thresh
postprocess_params["nms_thresh"] = args.det_sast_nms_thresh
if args.det_box_type == "poly":
postprocess_params["sample_pts_num"] = 6
postprocess_params["expand_scale"] = 1.2
postprocess_params["shrink_ratio_of_width"] = 0.2
else:
postprocess_params["sample_pts_num"] = 2
postprocess_params["expand_scale"] = 1.0
postprocess_params["shrink_ratio_of_width"] = 0.3
elif self.det_algorithm == "PSE":
postprocess_params["name"] = "PSEPostProcess"
postprocess_params["thresh"] = args.det_pse_thresh
postprocess_params["box_thresh"] = args.det_pse_box_thresh
postprocess_params["min_area"] = args.det_pse_min_area
postprocess_params["box_type"] = args.det_box_type
postprocess_params["scale"] = args.det_pse_scale
elif self.det_algorithm == "FCE":
pre_process_list[0] = {"DetResizeForTest": {"rescale_img": [1080, 736]}}
postprocess_params["name"] = "FCEPostProcess"
postprocess_params["scales"] = args.scales
postprocess_params["alpha"] = args.alpha
postprocess_params["beta"] = args.beta
postprocess_params["fourier_degree"] = args.fourier_degree
postprocess_params["box_type"] = args.det_box_type
elif self.det_algorithm == "CT":
pre_process_list[0] = {"ScaleAlignedShort": {"short_size": 640}}
postprocess_params["name"] = "CTPostProcess"
else:
logger.info("unknown det_algorithm:{}".format(self.det_algorithm))
sys.exit(0)
self.preprocess_op = create_operators(pre_process_list)
self.postprocess_op = build_post_process(postprocess_params)
(
self.predictor,
self.input_tensor,
self.output_tensors,
self.config,
) = utility.create_predictor(args, "det", logger)
if self.use_onnx:
img_h, img_w = self.input_tensor.shape[2:]
if isinstance(img_h, str) or isinstance(img_w, str):
pass
elif img_h is not None and img_w is not None and img_h > 0 and img_w > 0:
pre_process_list[0] = {
"DetResizeForTest": {"image_shape": [img_h, img_w]}
}
self.preprocess_op = create_operators(pre_process_list)
if args.benchmark:
import auto_log
pid = os.getpid()
gpu_id = utility.get_infer_gpuid()
self.autolog = auto_log.AutoLogger(
model_name="det",
model_precision=args.precision,
batch_size=1,
data_shape="dynamic",
save_path=None, # not used if logger is not None
inference_config=self.config,
pids=pid,
process_name=None,
gpu_ids=gpu_id if args.use_gpu else None,
time_keys=["preprocess_time", "inference_time", "postprocess_time"],
warmup=2,
logger=logger,
)
def order_points_clockwise(self, pts):
rect = np.zeros((4, 2), dtype="float32")
s = pts.sum(axis=1)
rect[0] = pts[np.argmin(s)]
rect[2] = pts[np.argmax(s)]
tmp = np.delete(pts, (np.argmin(s), np.argmax(s)), axis=0)
diff = np.diff(np.array(tmp), axis=1)
rect[1] = tmp[np.argmin(diff)]
rect[3] = tmp[np.argmax(diff)]
return rect
def pad_polygons(self, polygon, max_points):
padding_size = max_points - len(polygon)
if padding_size == 0:
return polygon
last_point = polygon[-1]
padding = np.repeat([last_point], padding_size, axis=0)
return np.vstack([polygon, padding])
def clip_det_res(self, points, img_height, img_width):
for pno in range(points.shape[0]):
points[pno, 0] = int(min(max(points[pno, 0], 0), img_width - 1))
points[pno, 1] = int(min(max(points[pno, 1], 0), img_height - 1))
return points
def filter_tag_det_res(self, dt_boxes, image_shape):
img_height, img_width = image_shape[0:2]
dt_boxes_new = []
for box in dt_boxes:
if type(box) is list:
box = np.array(box)
box = self.order_points_clockwise(box)
box = self.clip_det_res(box, img_height, img_width)
rect_width = int(np.linalg.norm(box[0] - box[1]))
rect_height = int(np.linalg.norm(box[0] - box[3]))
if rect_width <= 3 or rect_height <= 3:
continue
dt_boxes_new.append(box)
dt_boxes = np.array(dt_boxes_new)
return dt_boxes
def filter_tag_det_res_only_clip(self, dt_boxes, image_shape):
img_height, img_width = image_shape[0:2]
dt_boxes_new = []
for box in dt_boxes:
if type(box) is list:
box = np.array(box)
box = self.clip_det_res(box, img_height, img_width)
dt_boxes_new.append(box)
if len(dt_boxes_new) > 0:
max_points = max(len(polygon) for polygon in dt_boxes_new)
dt_boxes_new = [
self.pad_polygons(polygon, max_points) for polygon in dt_boxes_new
]
dt_boxes = np.array(dt_boxes_new)
return dt_boxes
def predict(self, img):
ori_im = img.copy()
data = {"image": img}
st = time.time()
if self.args.benchmark:
self.autolog.times.start()
data = transform(data, self.preprocess_op)
img, shape_list = data
if img is None:
return None, 0
img = np.expand_dims(img, axis=0)
shape_list = np.expand_dims(shape_list, axis=0)
img = img.copy()
if self.args.benchmark:
self.autolog.times.stamp()
if self.use_onnx:
input_dict = {}
input_dict[self.input_tensor.name] = img
outputs = self.predictor.run(self.output_tensors, input_dict)
else:
self.input_tensor.copy_from_cpu(img)
self.predictor.run()
outputs = []
for output_tensor in self.output_tensors:
output = output_tensor.copy_to_cpu()
outputs.append(output)
if self.args.benchmark:
self.autolog.times.stamp()
preds = {}
if self.det_algorithm == "EAST":
preds["f_geo"] = outputs[0]
preds["f_score"] = outputs[1]
elif self.det_algorithm == "SAST":
preds["f_border"] = outputs[0]
preds["f_score"] = outputs[1]
preds["f_tco"] = outputs[2]
preds["f_tvo"] = outputs[3]
elif self.det_algorithm in ["DB", "PSE", "DB++"]:
preds["maps"] = outputs[0]
elif self.det_algorithm == "FCE":
for i, output in enumerate(outputs):
preds["level_{}".format(i)] = output
elif self.det_algorithm == "CT":
preds["maps"] = outputs[0]
preds["score"] = outputs[1]
else:
raise NotImplementedError
post_result = self.postprocess_op(preds, shape_list)
dt_boxes = post_result[0]["points"]
if self.args.det_box_type == "poly":
dt_boxes = self.filter_tag_det_res_only_clip(dt_boxes, ori_im.shape)
else:
dt_boxes = self.filter_tag_det_res(dt_boxes, ori_im.shape)
if self.args.benchmark:
self.autolog.times.end(stamp=True)
et = time.time()
return dt_boxes, et - st
def __call__(self, img, use_slice=False):
# For image like poster with one side much greater than the other side,
# splitting recursively and processing with overlap to enhance performance.
MIN_BOUND_DISTANCE = 50
dt_boxes = np.zeros((0, 4, 2), dtype=np.float32)
elapse = 0
if (
img.shape[0] / img.shape[1] > 2
and img.shape[0] > self.args.det_limit_side_len
and use_slice
):
start_h = 0
end_h = 0
while end_h <= img.shape[0]:
end_h = start_h + img.shape[1] * 3 // 4
subimg = img[start_h:end_h, :]
if len(subimg) == 0:
break
sub_dt_boxes, sub_elapse = self.predict(subimg)
offset = start_h
# To prevent text blocks from being cut off, roll back a certain buffer area.
if (
len(sub_dt_boxes) == 0
or img.shape[1] - max([x[-1][1] for x in sub_dt_boxes])
> MIN_BOUND_DISTANCE
):
start_h = end_h
else:
sorted_indices = np.argsort(sub_dt_boxes[:, 2, 1])
sub_dt_boxes = sub_dt_boxes[sorted_indices]
bottom_line = (
0
if len(sub_dt_boxes) <= 1
else int(np.max(sub_dt_boxes[:-1, 2, 1]))
)
if bottom_line > 0:
start_h += bottom_line
sub_dt_boxes = sub_dt_boxes[
sub_dt_boxes[:, 2, 1] <= bottom_line
]
else:
start_h = end_h
if len(sub_dt_boxes) > 0:
if dt_boxes.shape[0] == 0:
dt_boxes = sub_dt_boxes + np.array(
[0, offset], dtype=np.float32
)
else:
dt_boxes = np.append(
dt_boxes,
sub_dt_boxes + np.array([0, offset], dtype=np.float32),
axis=0,
)
elapse += sub_elapse
elif (
img.shape[1] / img.shape[0] > 3
and img.shape[1] > self.args.det_limit_side_len * 3
and use_slice
):
start_w = 0
end_w = 0
while end_w <= img.shape[1]:
end_w = start_w + img.shape[0] * 3 // 4
subimg = img[:, start_w:end_w]
if len(subimg) == 0:
break
sub_dt_boxes, sub_elapse = self.predict(subimg)
offset = start_w
if (
len(sub_dt_boxes) == 0
or img.shape[0] - max([x[-1][0] for x in sub_dt_boxes])
> MIN_BOUND_DISTANCE
):
start_w = end_w
else:
sorted_indices = np.argsort(sub_dt_boxes[:, 2, 0])
sub_dt_boxes = sub_dt_boxes[sorted_indices]
right_line = (
0
if len(sub_dt_boxes) <= 1
else int(np.max(sub_dt_boxes[:-1, 1, 0]))
)
if right_line > 0:
start_w += right_line
sub_dt_boxes = sub_dt_boxes[sub_dt_boxes[:, 1, 0] <= right_line]
else:
start_w = end_w
if len(sub_dt_boxes) > 0:
if dt_boxes.shape[0] == 0:
dt_boxes = sub_dt_boxes + np.array(
[offset, 0], dtype=np.float32
)
else:
dt_boxes = np.append(
dt_boxes,
sub_dt_boxes + np.array([offset, 0], dtype=np.float32),
axis=0,
)
elapse += sub_elapse
else:
dt_boxes, elapse = self.predict(img)
return dt_boxes, elapse
if __name__ == "__main__":
args = utility.parse_args()
image_file_list = get_image_file_list(args.image_dir)
total_time = 0
draw_img_save_dir = args.draw_img_save_dir
os.makedirs(draw_img_save_dir, exist_ok=True)
# logger
log_file = args.save_log_path
if os.path.isdir(args.save_log_path) or (
not os.path.exists(args.save_log_path) and args.save_log_path.endswith("/")
):
log_file = os.path.join(log_file, "benchmark_detection.log")
logger = get_logger(log_file=log_file)
# create text detector
text_detector = TextDetector(args, logger)
if args.warmup:
img = np.random.uniform(0, 255, [640, 640, 3]).astype(np.uint8)
for i in range(2):
res = text_detector(img)
save_results = []
for idx, image_file in enumerate(image_file_list):
img, flag_gif, flag_pdf = check_and_read(image_file)
if not flag_gif and not flag_pdf:
img = cv2.imread(image_file)
if not flag_pdf:
if img is None:
logger.debug("error in loading image:{}".format(image_file))
continue
imgs = [img]
else:
page_num = args.page_num
if page_num > len(img) or page_num == 0:
page_num = len(img)
imgs = img[:page_num]
for index, img in enumerate(imgs):
st = time.time()
dt_boxes, _ = text_detector(img)
elapse = time.time() - st
total_time += elapse
if len(imgs) > 1:
save_pred = (
os.path.basename(image_file)
+ "_"
+ str(index)
+ "\t"
+ str(json.dumps([x.tolist() for x in dt_boxes]))
+ "\n"
)
else:
save_pred = (
os.path.basename(image_file)
+ "\t"
+ str(json.dumps([x.tolist() for x in dt_boxes]))
+ "\n"
)
save_results.append(save_pred)
logger.info(save_pred)
if len(imgs) > 1:
logger.info(
"{}_{} The predict time of {}: {}".format(
idx, index, image_file, elapse
)
)
else:
logger.info(
"{} The predict time of {}: {}".format(idx, image_file, elapse)
)
src_im = utility.draw_text_det_res(dt_boxes, img)
if flag_gif:
save_file = image_file[:-3] + "png"
elif flag_pdf:
save_file = image_file.replace(".pdf", "_" + str(index) + ".png")
else:
save_file = image_file
img_path = os.path.join(
draw_img_save_dir, "det_res_{}".format(os.path.basename(save_file))
)
cv2.imwrite(img_path, src_im)
logger.info("The visualized image saved in {}".format(img_path))
with open(os.path.join(draw_img_save_dir, "det_results.txt"), "w") as f:
f.writelines(save_results)
f.close()
if args.benchmark:
text_detector.autolog.report()

178
tools/infer/predict_e2e.py Executable file
View File

@@ -0,0 +1,178 @@
# 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.
import os
import sys
__dir__ = os.path.dirname(os.path.abspath(__file__))
sys.path.append(__dir__)
sys.path.insert(0, os.path.abspath(os.path.join(__dir__, "../..")))
os.environ["FLAGS_allocator_strategy"] = "auto_growth"
import cv2
import numpy as np
import time
import sys
import tools.infer.utility as utility
from ppocr.utils.logging import get_logger
from ppocr.utils.utility import get_image_file_list, check_and_read
from ppocr.data import create_operators, transform
from ppocr.postprocess import build_post_process
logger = get_logger()
class TextE2E(object):
def __init__(self, args):
if os.path.exists(f"{args.e2e_model_dir}/inference.yml"):
model_config = utility.load_config(f"{args.e2e_model_dir}/inference.yml")
model_name = model_config.get("Global", {}).get("model_name", "")
if model_name:
raise ValueError(
f"{model_name} is not supported. Please check if the model is supported by the PaddleOCR wheel."
)
self.args = args
self.e2e_algorithm = args.e2e_algorithm
self.use_onnx = args.use_onnx
pre_process_list = [
{"E2EResizeForTest": {}},
{
"NormalizeImage": {
"std": [0.229, 0.224, 0.225],
"mean": [0.485, 0.456, 0.406],
"scale": "1./255.",
"order": "hwc",
}
},
{"ToCHWImage": None},
{"KeepKeys": {"keep_keys": ["image", "shape"]}},
]
postprocess_params = {}
if self.e2e_algorithm == "PGNet":
pre_process_list[0] = {
"E2EResizeForTest": {
"max_side_len": args.e2e_limit_side_len,
"valid_set": "totaltext",
}
}
postprocess_params["name"] = "PGPostProcess"
postprocess_params["score_thresh"] = args.e2e_pgnet_score_thresh
postprocess_params["character_dict_path"] = args.e2e_char_dict_path
postprocess_params["valid_set"] = args.e2e_pgnet_valid_set
postprocess_params["mode"] = args.e2e_pgnet_mode
else:
logger.info("unknown e2e_algorithm:{}".format(self.e2e_algorithm))
sys.exit(0)
self.preprocess_op = create_operators(pre_process_list)
self.postprocess_op = build_post_process(postprocess_params)
(
self.predictor,
self.input_tensor,
self.output_tensors,
_,
) = utility.create_predictor(
args, "e2e", logger
) # paddle.jit.load(args.det_model_dir)
# self.predictor.eval()
def clip_det_res(self, points, img_height, img_width):
for pno in range(points.shape[0]):
points[pno, 0] = int(min(max(points[pno, 0], 0), img_width - 1))
points[pno, 1] = int(min(max(points[pno, 1], 0), img_height - 1))
return points
def filter_tag_det_res_only_clip(self, dt_boxes, image_shape):
img_height, img_width = image_shape[0:2]
dt_boxes_new = []
for box in dt_boxes:
box = self.clip_det_res(box, img_height, img_width)
dt_boxes_new.append(box)
dt_boxes = np.array(dt_boxes_new)
return dt_boxes
def __call__(self, img):
ori_im = img.copy()
data = {"image": img}
data = transform(data, self.preprocess_op)
img, shape_list = data
if img is None:
return None, 0
img = np.expand_dims(img, axis=0)
shape_list = np.expand_dims(shape_list, axis=0)
img = img.copy()
starttime = time.time()
if self.use_onnx:
input_dict = {}
input_dict[self.input_tensor.name] = img
outputs = self.predictor.run(self.output_tensors, input_dict)
preds = {}
preds["f_border"] = outputs[0]
preds["f_char"] = outputs[1]
preds["f_direction"] = outputs[2]
preds["f_score"] = outputs[3]
else:
self.input_tensor.copy_from_cpu(img)
self.predictor.run()
outputs = []
for output_tensor in self.output_tensors:
output = output_tensor.copy_to_cpu()
outputs.append(output)
preds = {}
if self.e2e_algorithm == "PGNet":
preds["f_border"] = outputs[0]
preds["f_char"] = outputs[1]
preds["f_direction"] = outputs[2]
preds["f_score"] = outputs[3]
else:
raise NotImplementedError
post_result = self.postprocess_op(preds, shape_list)
points, strs = post_result["points"], post_result["texts"]
dt_boxes = self.filter_tag_det_res_only_clip(points, ori_im.shape)
elapse = time.time() - starttime
return dt_boxes, strs, elapse
if __name__ == "__main__":
args = utility.parse_args()
image_file_list = get_image_file_list(args.image_dir)
text_detector = TextE2E(args)
count = 0
total_time = 0
draw_img_save = "./inference_results"
if not os.path.exists(draw_img_save):
os.makedirs(draw_img_save)
for image_file in image_file_list:
img, flag, _ = check_and_read(image_file)
if not flag:
img = cv2.imread(image_file)
if img is None:
logger.info("error in loading image:{}".format(image_file))
continue
points, strs, elapse = text_detector(img)
if count > 0:
total_time += elapse
count += 1
logger.info("Predict time of {}: {}".format(image_file, elapse))
src_im = utility.draw_e2e_res(points, strs, image_file)
img_name_pure = os.path.split(image_file)[-1]
img_path = os.path.join(draw_img_save, "e2e_res_{}".format(img_name_pure))
cv2.imwrite(img_path, src_im)
logger.info("The visualized image saved in {}".format(img_path))
if count > 1:
logger.info("Avg Time: {}".format(total_time / (count - 1)))

896
tools/infer/predict_rec.py Executable file
View File

@@ -0,0 +1,896 @@
# 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.
import os
import sys
from PIL import Image
__dir__ = os.path.dirname(os.path.abspath(__file__))
sys.path.append(__dir__)
sys.path.insert(0, os.path.abspath(os.path.join(__dir__, "../..")))
os.environ["FLAGS_allocator_strategy"] = "auto_growth"
import cv2
import numpy as np
import math
import time
import traceback
import paddle
import tools.infer.utility as utility
from ppocr.postprocess import build_post_process
from ppocr.utils.logging import get_logger
from ppocr.utils.utility import get_image_file_list, check_and_read
logger = get_logger()
class TextRecognizer(object):
def __init__(self, args, logger=None):
if os.path.exists(f"{args.rec_model_dir}/inference.yml"):
model_config = utility.load_config(f"{args.rec_model_dir}/inference.yml")
model_name = model_config.get("Global", {}).get("model_name", "")
if model_name and model_name not in [
"PP-OCRv5_mobile_rec",
"PP-OCRv5_server_rec",
]:
raise ValueError(
f"{model_name} is not supported. Please check if the model is supported by the PaddleOCR wheel."
)
if args.rec_char_dict_path == "./ppocr/utils/ppocr_keys_v1.txt":
rec_char_list = model_config.get("PostProcess", {}).get(
"character_dict", []
)
if rec_char_list:
new_rec_char_dict_path = f"{args.rec_model_dir}/ppocr_keys.txt"
with open(new_rec_char_dict_path, "w", encoding="utf-8") as f:
f.writelines([char + "\n" for char in rec_char_list])
args.rec_char_dict_path = new_rec_char_dict_path
if logger is None:
logger = get_logger()
self.rec_image_shape = [int(v) for v in args.rec_image_shape.split(",")]
self.rec_batch_num = args.rec_batch_num
self.rec_algorithm = args.rec_algorithm
postprocess_params = {
"name": "CTCLabelDecode",
"character_dict_path": args.rec_char_dict_path,
"use_space_char": args.use_space_char,
}
if self.rec_algorithm == "SRN":
postprocess_params = {
"name": "SRNLabelDecode",
"character_dict_path": args.rec_char_dict_path,
"use_space_char": args.use_space_char,
}
elif self.rec_algorithm == "RARE":
postprocess_params = {
"name": "AttnLabelDecode",
"character_dict_path": args.rec_char_dict_path,
"use_space_char": args.use_space_char,
}
elif self.rec_algorithm == "NRTR":
postprocess_params = {
"name": "NRTRLabelDecode",
"character_dict_path": args.rec_char_dict_path,
"use_space_char": args.use_space_char,
}
elif self.rec_algorithm == "SAR":
postprocess_params = {
"name": "SARLabelDecode",
"character_dict_path": args.rec_char_dict_path,
"use_space_char": args.use_space_char,
}
elif self.rec_algorithm == "VisionLAN":
postprocess_params = {
"name": "VLLabelDecode",
"character_dict_path": args.rec_char_dict_path,
"use_space_char": args.use_space_char,
"max_text_length": args.max_text_length,
}
elif self.rec_algorithm == "ViTSTR":
postprocess_params = {
"name": "ViTSTRLabelDecode",
"character_dict_path": args.rec_char_dict_path,
"use_space_char": args.use_space_char,
}
elif self.rec_algorithm == "ABINet":
postprocess_params = {
"name": "ABINetLabelDecode",
"character_dict_path": args.rec_char_dict_path,
"use_space_char": args.use_space_char,
}
elif self.rec_algorithm == "SPIN":
postprocess_params = {
"name": "SPINLabelDecode",
"character_dict_path": args.rec_char_dict_path,
"use_space_char": args.use_space_char,
}
elif self.rec_algorithm == "RobustScanner":
postprocess_params = {
"name": "SARLabelDecode",
"character_dict_path": args.rec_char_dict_path,
"use_space_char": args.use_space_char,
"rm_symbol": True,
}
elif self.rec_algorithm == "RFL":
postprocess_params = {
"name": "RFLLabelDecode",
"character_dict_path": None,
"use_space_char": args.use_space_char,
}
elif self.rec_algorithm == "SATRN":
postprocess_params = {
"name": "SATRNLabelDecode",
"character_dict_path": args.rec_char_dict_path,
"use_space_char": args.use_space_char,
"rm_symbol": True,
}
elif self.rec_algorithm in ["CPPD", "CPPDPadding"]:
postprocess_params = {
"name": "CPPDLabelDecode",
"character_dict_path": args.rec_char_dict_path,
"use_space_char": args.use_space_char,
"rm_symbol": True,
}
elif self.rec_algorithm == "PREN":
postprocess_params = {"name": "PRENLabelDecode"}
elif self.rec_algorithm == "CAN":
self.inverse = args.rec_image_inverse
postprocess_params = {
"name": "CANLabelDecode",
"character_dict_path": args.rec_char_dict_path,
"use_space_char": args.use_space_char,
}
elif self.rec_algorithm == "LaTeXOCR":
postprocess_params = {
"name": "LaTeXOCRDecode",
"rec_char_dict_path": args.rec_char_dict_path,
}
elif self.rec_algorithm == "ParseQ":
postprocess_params = {
"name": "ParseQLabelDecode",
"character_dict_path": args.rec_char_dict_path,
"use_space_char": args.use_space_char,
}
self.postprocess_op = build_post_process(postprocess_params)
self.postprocess_params = postprocess_params
(
self.predictor,
self.input_tensor,
self.output_tensors,
self.config,
) = utility.create_predictor(args, "rec", logger)
self.benchmark = args.benchmark
self.use_onnx = args.use_onnx
if args.benchmark:
import auto_log
pid = os.getpid()
gpu_id = utility.get_infer_gpuid()
self.autolog = auto_log.AutoLogger(
model_name="rec",
model_precision=args.precision,
batch_size=args.rec_batch_num,
data_shape="dynamic",
save_path=None, # not used if logger is not None
inference_config=self.config,
pids=pid,
process_name=None,
gpu_ids=gpu_id if args.use_gpu else None,
time_keys=["preprocess_time", "inference_time", "postprocess_time"],
warmup=0,
logger=logger,
)
self.return_word_box = args.return_word_box
def resize_norm_img(self, img, max_wh_ratio):
imgC, imgH, imgW = self.rec_image_shape
if self.rec_algorithm == "NRTR" or self.rec_algorithm == "ViTSTR":
img = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)
# return padding_im
image_pil = Image.fromarray(np.uint8(img))
if self.rec_algorithm == "ViTSTR":
img = image_pil.resize([imgW, imgH], Image.BICUBIC)
else:
img = image_pil.resize([imgW, imgH], Image.Resampling.LANCZOS)
img = np.array(img)
norm_img = np.expand_dims(img, -1)
norm_img = norm_img.transpose((2, 0, 1))
if self.rec_algorithm == "ViTSTR":
norm_img = norm_img.astype(np.float32) / 255.0
else:
norm_img = norm_img.astype(np.float32) / 128.0 - 1.0
return norm_img
elif self.rec_algorithm == "RFL":
img = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)
resized_image = cv2.resize(img, (imgW, imgH), interpolation=cv2.INTER_CUBIC)
resized_image = resized_image.astype("float32")
resized_image = resized_image / 255
resized_image = resized_image[np.newaxis, :]
resized_image -= 0.5
resized_image /= 0.5
return resized_image
assert imgC == img.shape[2]
imgW = int((imgH * max_wh_ratio))
if self.use_onnx:
w = self.input_tensor.shape[3:][0]
if isinstance(w, str):
pass
elif w is not None and w > 0:
imgW = w
h, w = img.shape[:2]
ratio = w / float(h)
if math.ceil(imgH * ratio) > imgW:
resized_w = imgW
else:
resized_w = int(math.ceil(imgH * ratio))
if self.rec_algorithm == "RARE":
if resized_w > self.rec_image_shape[2]:
resized_w = self.rec_image_shape[2]
imgW = self.rec_image_shape[2]
resized_image = cv2.resize(img, (resized_w, imgH))
resized_image = resized_image.astype("float32")
resized_image = resized_image.transpose((2, 0, 1)) / 255
resized_image -= 0.5
resized_image /= 0.5
padding_im = np.zeros((imgC, imgH, imgW), dtype=np.float32)
padding_im[:, :, 0:resized_w] = resized_image
return padding_im
def resize_norm_img_vl(self, img, image_shape):
imgC, imgH, imgW = image_shape
img = img[:, :, ::-1] # bgr2rgb
resized_image = cv2.resize(img, (imgW, imgH), interpolation=cv2.INTER_LINEAR)
resized_image = resized_image.astype("float32")
resized_image = resized_image.transpose((2, 0, 1)) / 255
return resized_image
def resize_norm_img_srn(self, img, image_shape):
imgC, imgH, imgW = image_shape
img_black = np.zeros((imgH, imgW))
im_hei = img.shape[0]
im_wid = img.shape[1]
if im_wid <= im_hei * 1:
img_new = cv2.resize(img, (imgH * 1, imgH))
elif im_wid <= im_hei * 2:
img_new = cv2.resize(img, (imgH * 2, imgH))
elif im_wid <= im_hei * 3:
img_new = cv2.resize(img, (imgH * 3, imgH))
else:
img_new = cv2.resize(img, (imgW, imgH))
img_np = np.asarray(img_new)
img_np = cv2.cvtColor(img_np, cv2.COLOR_BGR2GRAY)
img_black[:, 0 : img_np.shape[1]] = img_np
img_black = img_black[:, :, np.newaxis]
row, col, c = img_black.shape
c = 1
return np.reshape(img_black, (c, row, col)).astype(np.float32)
def srn_other_inputs(self, image_shape, num_heads, max_text_length):
imgC, imgH, imgW = image_shape
feature_dim = int((imgH / 8) * (imgW / 8))
encoder_word_pos = (
np.array(range(0, feature_dim)).reshape((feature_dim, 1)).astype("int64")
)
gsrm_word_pos = (
np.array(range(0, max_text_length))
.reshape((max_text_length, 1))
.astype("int64")
)
gsrm_attn_bias_data = np.ones((1, max_text_length, max_text_length))
gsrm_slf_attn_bias1 = np.triu(gsrm_attn_bias_data, 1).reshape(
[-1, 1, max_text_length, max_text_length]
)
gsrm_slf_attn_bias1 = np.tile(gsrm_slf_attn_bias1, [1, num_heads, 1, 1]).astype(
"float32"
) * [-1e9]
gsrm_slf_attn_bias2 = np.tril(gsrm_attn_bias_data, -1).reshape(
[-1, 1, max_text_length, max_text_length]
)
gsrm_slf_attn_bias2 = np.tile(gsrm_slf_attn_bias2, [1, num_heads, 1, 1]).astype(
"float32"
) * [-1e9]
encoder_word_pos = encoder_word_pos[np.newaxis, :]
gsrm_word_pos = gsrm_word_pos[np.newaxis, :]
return [
encoder_word_pos,
gsrm_word_pos,
gsrm_slf_attn_bias1,
gsrm_slf_attn_bias2,
]
def process_image_srn(self, img, image_shape, num_heads, max_text_length):
norm_img = self.resize_norm_img_srn(img, image_shape)
norm_img = norm_img[np.newaxis, :]
[
encoder_word_pos,
gsrm_word_pos,
gsrm_slf_attn_bias1,
gsrm_slf_attn_bias2,
] = self.srn_other_inputs(image_shape, num_heads, max_text_length)
gsrm_slf_attn_bias1 = gsrm_slf_attn_bias1.astype(np.float32)
gsrm_slf_attn_bias2 = gsrm_slf_attn_bias2.astype(np.float32)
encoder_word_pos = encoder_word_pos.astype(np.int64)
gsrm_word_pos = gsrm_word_pos.astype(np.int64)
return (
norm_img,
encoder_word_pos,
gsrm_word_pos,
gsrm_slf_attn_bias1,
gsrm_slf_attn_bias2,
)
def resize_norm_img_sar(self, img, image_shape, width_downsample_ratio=0.25):
imgC, imgH, imgW_min, imgW_max = image_shape
h = img.shape[0]
w = img.shape[1]
valid_ratio = 1.0
# make sure new_width is an integral multiple of width_divisor.
width_divisor = int(1 / width_downsample_ratio)
# resize
ratio = w / float(h)
resize_w = math.ceil(imgH * ratio)
if resize_w % width_divisor != 0:
resize_w = round(resize_w / width_divisor) * width_divisor
if imgW_min is not None:
resize_w = max(imgW_min, resize_w)
if imgW_max is not None:
valid_ratio = min(1.0, 1.0 * resize_w / imgW_max)
resize_w = min(imgW_max, resize_w)
resized_image = cv2.resize(img, (resize_w, imgH))
resized_image = resized_image.astype("float32")
# norm
if image_shape[0] == 1:
resized_image = resized_image / 255
resized_image = resized_image[np.newaxis, :]
else:
resized_image = resized_image.transpose((2, 0, 1)) / 255
resized_image -= 0.5
resized_image /= 0.5
resize_shape = resized_image.shape
padding_im = -1.0 * np.ones((imgC, imgH, imgW_max), dtype=np.float32)
padding_im[:, :, 0:resize_w] = resized_image
pad_shape = padding_im.shape
return padding_im, resize_shape, pad_shape, valid_ratio
def resize_norm_img_spin(self, img):
img = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)
# return padding_im
img = cv2.resize(img, tuple([100, 32]), cv2.INTER_CUBIC)
img = np.array(img, np.float32)
img = np.expand_dims(img, -1)
img = img.transpose((2, 0, 1))
mean = [127.5]
std = [127.5]
mean = np.array(mean, dtype=np.float32)
std = np.array(std, dtype=np.float32)
mean = np.float32(mean.reshape(1, -1))
stdinv = 1 / np.float32(std.reshape(1, -1))
img -= mean
img *= stdinv
return img
def resize_norm_img_svtr(self, img, image_shape):
imgC, imgH, imgW = image_shape
resized_image = cv2.resize(img, (imgW, imgH), interpolation=cv2.INTER_LINEAR)
resized_image = resized_image.astype("float32")
resized_image = resized_image.transpose((2, 0, 1)) / 255
resized_image -= 0.5
resized_image /= 0.5
return resized_image
def resize_norm_img_cppd_padding(
self, img, image_shape, padding=True, interpolation=cv2.INTER_LINEAR
):
imgC, imgH, imgW = image_shape
h = img.shape[0]
w = img.shape[1]
if not padding:
resized_image = cv2.resize(img, (imgW, imgH), interpolation=interpolation)
resized_w = imgW
else:
ratio = w / float(h)
if math.ceil(imgH * ratio) > imgW:
resized_w = imgW
else:
resized_w = int(math.ceil(imgH * ratio))
resized_image = cv2.resize(img, (resized_w, imgH))
resized_image = resized_image.astype("float32")
if image_shape[0] == 1:
resized_image = resized_image / 255
resized_image = resized_image[np.newaxis, :]
else:
resized_image = resized_image.transpose((2, 0, 1)) / 255
resized_image -= 0.5
resized_image /= 0.5
padding_im = np.zeros((imgC, imgH, imgW), dtype=np.float32)
padding_im[:, :, 0:resized_w] = resized_image
return padding_im
def resize_norm_img_abinet(self, img, image_shape):
imgC, imgH, imgW = image_shape
resized_image = cv2.resize(img, (imgW, imgH), interpolation=cv2.INTER_LINEAR)
resized_image = resized_image.astype("float32")
resized_image = resized_image / 255.0
mean = np.array([0.485, 0.456, 0.406])
std = np.array([0.229, 0.224, 0.225])
resized_image = (resized_image - mean[None, None, ...]) / std[None, None, ...]
resized_image = resized_image.transpose((2, 0, 1))
resized_image = resized_image.astype("float32")
return resized_image
def norm_img_can(self, img, image_shape):
img = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY) # CAN only predict gray scale image
if self.inverse:
img = 255 - img
if self.rec_image_shape[0] == 1:
h, w = img.shape
_, imgH, imgW = self.rec_image_shape
if h < imgH or w < imgW:
padding_h = max(imgH - h, 0)
padding_w = max(imgW - w, 0)
img_padded = np.pad(
img,
((0, padding_h), (0, padding_w)),
"constant",
constant_values=(255),
)
img = img_padded
img = np.expand_dims(img, 0) / 255.0 # h,w,c -> c,h,w
img = img.astype("float32")
return img
def pad_(self, img, divable=32):
threshold = 128
data = np.array(img.convert("LA"))
if data[..., -1].var() == 0:
data = (data[..., 0]).astype(np.uint8)
else:
data = (255 - data[..., -1]).astype(np.uint8)
data = (data - data.min()) / (data.max() - data.min()) * 255
if data.mean() > threshold:
# To invert the text to white
gray = 255 * (data < threshold).astype(np.uint8)
else:
gray = 255 * (data > threshold).astype(np.uint8)
data = 255 - data
coords = cv2.findNonZero(gray) # Find all non-zero points (text)
a, b, w, h = cv2.boundingRect(coords) # Find minimum spanning bounding box
rect = data[b : b + h, a : a + w]
im = Image.fromarray(rect).convert("L")
dims = []
for x in [w, h]:
div, mod = divmod(x, divable)
dims.append(divable * (div + (1 if mod > 0 else 0)))
padded = Image.new("L", dims, 255)
padded.paste(im, (0, 0, im.size[0], im.size[1]))
return padded
def minmax_size_(
self,
img,
max_dimensions,
min_dimensions,
):
if max_dimensions is not None:
ratios = [a / b for a, b in zip(img.size, max_dimensions)]
if any([r > 1 for r in ratios]):
size = np.array(img.size) // max(ratios)
img = img.resize(tuple(size.astype(int)), Image.BILINEAR)
if min_dimensions is not None:
# hypothesis: there is a dim in img smaller than min_dimensions, and return a proper dim >= min_dimensions
padded_size = [
max(img_dim, min_dim)
for img_dim, min_dim in zip(img.size, min_dimensions)
]
if padded_size != list(img.size): # assert hypothesis
padded_im = Image.new("L", padded_size, 255)
padded_im.paste(img, img.getbbox())
img = padded_im
return img
def norm_img_latexocr(self, img):
# CAN only predict gray scale image
shape = (1, 1, 3)
mean = [0.7931, 0.7931, 0.7931]
std = [0.1738, 0.1738, 0.1738]
scale = np.float32(1.0 / 255.0)
min_dimensions = [32, 32]
max_dimensions = [672, 192]
mean = np.array(mean).reshape(shape).astype("float32")
std = np.array(std).reshape(shape).astype("float32")
im_h, im_w = img.shape[:2]
if (
min_dimensions[0] <= im_w <= max_dimensions[0]
and min_dimensions[1] <= im_h <= max_dimensions[1]
):
pass
else:
img = Image.fromarray(np.uint8(img))
img = self.minmax_size_(self.pad_(img), max_dimensions, min_dimensions)
img = np.array(img)
im_h, im_w = img.shape[:2]
img = np.dstack([img, img, img])
img = (img.astype("float32") * scale - mean) / std
img = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)
divide_h = math.ceil(im_h / 16) * 16
divide_w = math.ceil(im_w / 16) * 16
img = np.pad(
img, ((0, divide_h - im_h), (0, divide_w - im_w)), constant_values=(1, 1)
)
img = img[:, :, np.newaxis].transpose(2, 0, 1)
img = img.astype("float32")
return img
def __call__(self, img_list):
img_num = len(img_list)
# Calculate the aspect ratio of all text bars
width_list = []
for img in img_list:
width_list.append(img.shape[1] / float(img.shape[0]))
# Sorting can speed up the recognition process
indices = np.argsort(np.array(width_list))
rec_res = [["", 0.0]] * img_num
batch_num = self.rec_batch_num
st = time.time()
if self.benchmark:
self.autolog.times.start()
for beg_img_no in range(0, img_num, batch_num):
end_img_no = min(img_num, beg_img_no + batch_num)
norm_img_batch = []
if self.rec_algorithm == "SRN":
encoder_word_pos_list = []
gsrm_word_pos_list = []
gsrm_slf_attn_bias1_list = []
gsrm_slf_attn_bias2_list = []
if self.rec_algorithm == "SAR":
valid_ratios = []
imgC, imgH, imgW = self.rec_image_shape[:3]
max_wh_ratio = imgW / imgH
wh_ratio_list = []
for ino in range(beg_img_no, end_img_no):
h, w = img_list[indices[ino]].shape[0:2]
wh_ratio = w * 1.0 / h
max_wh_ratio = max(max_wh_ratio, wh_ratio)
wh_ratio_list.append(wh_ratio)
for ino in range(beg_img_no, end_img_no):
if self.rec_algorithm == "SAR":
norm_img, _, _, valid_ratio = self.resize_norm_img_sar(
img_list[indices[ino]], self.rec_image_shape
)
norm_img = norm_img[np.newaxis, :]
valid_ratio = np.expand_dims(valid_ratio, axis=0)
valid_ratios.append(valid_ratio)
norm_img_batch.append(norm_img)
elif self.rec_algorithm == "SRN":
norm_img = self.process_image_srn(
img_list[indices[ino]], self.rec_image_shape, 8, 25
)
encoder_word_pos_list.append(norm_img[1])
gsrm_word_pos_list.append(norm_img[2])
gsrm_slf_attn_bias1_list.append(norm_img[3])
gsrm_slf_attn_bias2_list.append(norm_img[4])
norm_img_batch.append(norm_img[0])
elif self.rec_algorithm in ["SVTR", "SATRN", "ParseQ", "CPPD"]:
norm_img = self.resize_norm_img_svtr(
img_list[indices[ino]], self.rec_image_shape
)
norm_img = norm_img[np.newaxis, :]
norm_img_batch.append(norm_img)
elif self.rec_algorithm in ["CPPDPadding"]:
norm_img = self.resize_norm_img_cppd_padding(
img_list[indices[ino]], self.rec_image_shape
)
norm_img = norm_img[np.newaxis, :]
norm_img_batch.append(norm_img)
elif self.rec_algorithm in ["VisionLAN", "PREN"]:
norm_img = self.resize_norm_img_vl(
img_list[indices[ino]], self.rec_image_shape
)
norm_img = norm_img[np.newaxis, :]
norm_img_batch.append(norm_img)
elif self.rec_algorithm == "SPIN":
norm_img = self.resize_norm_img_spin(img_list[indices[ino]])
norm_img = norm_img[np.newaxis, :]
norm_img_batch.append(norm_img)
elif self.rec_algorithm == "ABINet":
norm_img = self.resize_norm_img_abinet(
img_list[indices[ino]], self.rec_image_shape
)
norm_img = norm_img[np.newaxis, :]
norm_img_batch.append(norm_img)
elif self.rec_algorithm == "RobustScanner":
norm_img, _, _, valid_ratio = self.resize_norm_img_sar(
img_list[indices[ino]],
self.rec_image_shape,
width_downsample_ratio=0.25,
)
norm_img = norm_img[np.newaxis, :]
valid_ratio = np.expand_dims(valid_ratio, axis=0)
valid_ratios = []
valid_ratios.append(valid_ratio)
norm_img_batch.append(norm_img)
word_positions_list = []
word_positions = np.array(range(0, 40)).astype("int64")
word_positions = np.expand_dims(word_positions, axis=0)
word_positions_list.append(word_positions)
elif self.rec_algorithm == "CAN":
norm_img = self.norm_img_can(img_list[indices[ino]], max_wh_ratio)
norm_img = norm_img[np.newaxis, :]
norm_img_batch.append(norm_img)
norm_image_mask = np.ones(norm_img.shape, dtype="float32")
word_label = np.ones([1, 36], dtype="int64")
norm_img_mask_batch = []
word_label_list = []
norm_img_mask_batch.append(norm_image_mask)
word_label_list.append(word_label)
elif self.rec_algorithm == "LaTeXOCR":
norm_img = self.norm_img_latexocr(img_list[indices[ino]])
norm_img = norm_img[np.newaxis, :]
norm_img_batch.append(norm_img)
else:
norm_img = self.resize_norm_img(
img_list[indices[ino]], max_wh_ratio
)
norm_img = norm_img[np.newaxis, :]
norm_img_batch.append(norm_img)
norm_img_batch = np.concatenate(norm_img_batch)
norm_img_batch = norm_img_batch.copy()
if self.benchmark:
self.autolog.times.stamp()
if self.rec_algorithm == "SRN":
encoder_word_pos_list = np.concatenate(encoder_word_pos_list)
gsrm_word_pos_list = np.concatenate(gsrm_word_pos_list)
gsrm_slf_attn_bias1_list = np.concatenate(gsrm_slf_attn_bias1_list)
gsrm_slf_attn_bias2_list = np.concatenate(gsrm_slf_attn_bias2_list)
inputs = [
norm_img_batch,
encoder_word_pos_list,
gsrm_word_pos_list,
gsrm_slf_attn_bias1_list,
gsrm_slf_attn_bias2_list,
]
if self.use_onnx:
input_dict = {}
input_dict[self.input_tensor.name] = norm_img_batch
outputs = self.predictor.run(self.output_tensors, input_dict)
preds = {"predict": outputs[2]}
else:
input_names = self.predictor.get_input_names()
for i in range(len(input_names)):
input_tensor = self.predictor.get_input_handle(input_names[i])
input_tensor.copy_from_cpu(inputs[i])
self.predictor.run()
outputs = []
for output_tensor in self.output_tensors:
output = output_tensor.copy_to_cpu()
outputs.append(output)
if self.benchmark:
self.autolog.times.stamp()
preds = {"predict": outputs[2]}
elif self.rec_algorithm == "SAR":
valid_ratios = np.concatenate(valid_ratios)
inputs = [
norm_img_batch,
np.array([valid_ratios], dtype=np.float32).T,
]
if self.use_onnx:
input_dict = {}
input_dict[self.input_tensor.name] = norm_img_batch
outputs = self.predictor.run(self.output_tensors, input_dict)
preds = outputs[0]
else:
input_names = self.predictor.get_input_names()
for i in range(len(input_names)):
input_tensor = self.predictor.get_input_handle(input_names[i])
input_tensor.copy_from_cpu(inputs[i])
self.predictor.run()
outputs = []
for output_tensor in self.output_tensors:
output = output_tensor.copy_to_cpu()
outputs.append(output)
if self.benchmark:
self.autolog.times.stamp()
preds = outputs[0]
elif self.rec_algorithm == "RobustScanner":
valid_ratios = np.concatenate(valid_ratios)
word_positions_list = np.concatenate(word_positions_list)
inputs = [norm_img_batch, valid_ratios, word_positions_list]
if self.use_onnx:
input_dict = {}
input_dict[self.input_tensor.name] = norm_img_batch
outputs = self.predictor.run(self.output_tensors, input_dict)
preds = outputs[0]
else:
input_names = self.predictor.get_input_names()
for i in range(len(input_names)):
input_tensor = self.predictor.get_input_handle(input_names[i])
input_tensor.copy_from_cpu(inputs[i])
self.predictor.run()
outputs = []
for output_tensor in self.output_tensors:
output = output_tensor.copy_to_cpu()
outputs.append(output)
if self.benchmark:
self.autolog.times.stamp()
preds = outputs[0]
elif self.rec_algorithm == "CAN":
norm_img_mask_batch = np.concatenate(norm_img_mask_batch)
word_label_list = np.concatenate(word_label_list)
inputs = [norm_img_batch, norm_img_mask_batch, word_label_list]
if self.use_onnx:
input_dict = {}
input_dict[self.input_tensor.name] = norm_img_batch
outputs = self.predictor.run(self.output_tensors, input_dict)
preds = outputs
else:
input_names = self.predictor.get_input_names()
input_tensor = []
for i in range(len(input_names)):
input_tensor_i = self.predictor.get_input_handle(input_names[i])
input_tensor_i.copy_from_cpu(inputs[i])
input_tensor.append(input_tensor_i)
self.input_tensor = input_tensor
self.predictor.run()
outputs = []
for output_tensor in self.output_tensors:
output = output_tensor.copy_to_cpu()
outputs.append(output)
if self.benchmark:
self.autolog.times.stamp()
preds = outputs
elif self.rec_algorithm == "LaTeXOCR":
inputs = [norm_img_batch]
if self.use_onnx:
input_dict = {}
input_dict[self.input_tensor.name] = norm_img_batch
outputs = self.predictor.run(self.output_tensors, input_dict)
preds = outputs
else:
input_names = self.predictor.get_input_names()
input_tensor = []
for i in range(len(input_names)):
input_tensor_i = self.predictor.get_input_handle(input_names[i])
input_tensor_i.copy_from_cpu(inputs[i])
input_tensor.append(input_tensor_i)
self.input_tensor = input_tensor
self.predictor.run()
outputs = []
for output_tensor in self.output_tensors:
output = output_tensor.copy_to_cpu()
outputs.append(output)
if self.benchmark:
self.autolog.times.stamp()
preds = outputs
else:
if self.use_onnx:
input_dict = {}
input_dict[self.input_tensor.name] = norm_img_batch
outputs = self.predictor.run(self.output_tensors, input_dict)
preds = outputs[0]
else:
self.input_tensor.copy_from_cpu(norm_img_batch)
self.predictor.run()
outputs = []
for output_tensor in self.output_tensors:
output = output_tensor.copy_to_cpu()
outputs.append(output)
if self.benchmark:
self.autolog.times.stamp()
if len(outputs) != 1:
preds = outputs
else:
preds = outputs[0]
if self.postprocess_params["name"] == "CTCLabelDecode":
rec_result = self.postprocess_op(
preds,
return_word_box=self.return_word_box,
wh_ratio_list=wh_ratio_list,
max_wh_ratio=max_wh_ratio,
)
elif self.postprocess_params["name"] == "LaTeXOCRDecode":
preds = [p.reshape([-1]) for p in preds]
rec_result = self.postprocess_op(preds)
else:
rec_result = self.postprocess_op(preds)
for rno in range(len(rec_result)):
rec_res[indices[beg_img_no + rno]] = rec_result[rno]
if self.benchmark:
self.autolog.times.end(stamp=True)
return rec_res, time.time() - st
def main(args):
image_file_list = get_image_file_list(args.image_dir)
valid_image_file_list = []
img_list = []
# logger
log_file = args.save_log_path
if os.path.isdir(args.save_log_path) or (
not os.path.exists(args.save_log_path) and args.save_log_path.endswith("/")
):
log_file = os.path.join(log_file, "benchmark_recognition.log")
logger = get_logger(log_file=log_file)
# create text recognizer
text_recognizer = TextRecognizer(args)
logger.info(
"In PP-OCRv3, rec_image_shape parameter defaults to '3, 48, 320', "
"if you are using recognition model with PP-OCRv2 or an older version, please set --rec_image_shape='3,32,320"
)
# warmup 2 times
if args.warmup:
img = np.random.uniform(0, 255, [48, 320, 3]).astype(np.uint8)
for i in range(2):
res = text_recognizer([img] * int(args.rec_batch_num))
for image_file in image_file_list:
img, flag, _ = check_and_read(image_file)
if not flag:
img = cv2.imread(image_file)
if img is None:
logger.info("error in loading image:{}".format(image_file))
continue
valid_image_file_list.append(image_file)
img_list.append(img)
try:
rec_res, _ = text_recognizer(img_list)
except Exception as E:
logger.info(traceback.format_exc())
logger.info(E)
exit()
for ino in range(len(img_list)):
logger.info(
"Predicts of {}:{}".format(valid_image_file_list[ino], rec_res[ino])
)
if args.benchmark:
text_recognizer.autolog.report()
if __name__ == "__main__":
main(utility.parse_args())

173
tools/infer/predict_sr.py Executable file
View File

@@ -0,0 +1,173 @@
# 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.
import os
import sys
from PIL import Image
__dir__ = os.path.dirname(os.path.abspath(__file__))
sys.path.insert(0, __dir__)
sys.path.insert(0, os.path.abspath(os.path.join(__dir__, "../..")))
os.environ["FLAGS_allocator_strategy"] = "auto_growth"
import cv2
import numpy as np
import math
import time
import traceback
import paddle
import tools.infer.utility as utility
from ppocr.postprocess import build_post_process
from ppocr.utils.logging import get_logger
from ppocr.utils.utility import get_image_file_list, check_and_read
logger = get_logger()
class TextSR(object):
def __init__(self, args):
if os.path.exists(f"{args.sr_model_dir}/inference.yml"):
model_config = utility.load_config(f"{args.sr_model_dir}/inference.yml")
model_name = model_config.get("Global", {}).get("model_name", "")
if model_name:
raise ValueError(
f"{model_name} is not supported. Please check if the model is supported by the PaddleOCR wheel."
)
self.sr_image_shape = [int(v) for v in args.sr_image_shape.split(",")]
self.sr_batch_num = args.sr_batch_num
(
self.predictor,
self.input_tensor,
self.output_tensors,
self.config,
) = utility.create_predictor(args, "sr", logger)
self.benchmark = args.benchmark
if args.benchmark:
import auto_log
pid = os.getpid()
gpu_id = utility.get_infer_gpuid()
self.autolog = auto_log.AutoLogger(
model_name="sr",
model_precision=args.precision,
batch_size=args.sr_batch_num,
data_shape="dynamic",
save_path=None, # args.save_log_path,
inference_config=self.config,
pids=pid,
process_name=None,
gpu_ids=gpu_id if args.use_gpu else None,
time_keys=["preprocess_time", "inference_time", "postprocess_time"],
warmup=0,
logger=logger,
)
def resize_norm_img(self, img):
imgC, imgH, imgW = self.sr_image_shape
img = img.resize((imgW // 2, imgH // 2), Image.BICUBIC)
img_numpy = np.array(img).astype("float32")
img_numpy = img_numpy.transpose((2, 0, 1)) / 255
return img_numpy
def __call__(self, img_list):
img_num = len(img_list)
batch_num = self.sr_batch_num
st = time.time()
st = time.time()
all_result = [] * img_num
if self.benchmark:
self.autolog.times.start()
for beg_img_no in range(0, img_num, batch_num):
end_img_no = min(img_num, beg_img_no + batch_num)
norm_img_batch = []
imgC, imgH, imgW = self.sr_image_shape
for ino in range(beg_img_no, end_img_no):
norm_img = self.resize_norm_img(img_list[ino])
norm_img = norm_img[np.newaxis, :]
norm_img_batch.append(norm_img)
norm_img_batch = np.concatenate(norm_img_batch)
norm_img_batch = norm_img_batch.copy()
if self.benchmark:
self.autolog.times.stamp()
self.input_tensor.copy_from_cpu(norm_img_batch)
self.predictor.run()
outputs = []
for output_tensor in self.output_tensors:
output = output_tensor.copy_to_cpu()
outputs.append(output)
if len(outputs) != 1:
preds = outputs
else:
preds = outputs[0]
all_result.append(outputs)
if self.benchmark:
self.autolog.times.end(stamp=True)
return all_result, time.time() - st
def main(args):
image_file_list = get_image_file_list(args.image_dir)
text_recognizer = TextSR(args)
valid_image_file_list = []
img_list = []
# warmup 2 times
if args.warmup:
img = np.random.uniform(0, 255, [16, 64, 3]).astype(np.uint8)
for i in range(2):
res = text_recognizer([img] * int(args.sr_batch_num))
for image_file in image_file_list:
img, flag, _ = check_and_read(image_file)
if not flag:
img = Image.open(image_file).convert("RGB")
if img is None:
logger.info("error in loading image:{}".format(image_file))
continue
valid_image_file_list.append(image_file)
img_list.append(img)
try:
preds, _ = text_recognizer(img_list)
for beg_no in range(len(preds)):
sr_img = preds[beg_no][1]
lr_img = preds[beg_no][0]
for i in range(sr_img.shape[0]):
fm_sr = (sr_img[i] * 255).transpose(1, 2, 0).astype(np.uint8)
fm_lr = (lr_img[i] * 255).transpose(1, 2, 0).astype(np.uint8)
img_name_pure = os.path.split(
valid_image_file_list[beg_no * args.sr_batch_num + i]
)[-1]
cv2.imwrite(
"infer_result/sr_{}".format(img_name_pure), fm_sr[:, :, ::-1]
)
logger.info(
"The visualized image saved in infer_result/sr_{}".format(
img_name_pure
)
)
except Exception as E:
logger.info(traceback.format_exc())
logger.info(E)
exit()
if args.benchmark:
text_recognizer.autolog.report()
if __name__ == "__main__":
main(utility.parse_args())

326
tools/infer/predict_system.py Executable file
View File

@@ -0,0 +1,326 @@
# 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.
import os
import sys
import subprocess
__dir__ = os.path.dirname(os.path.abspath(__file__))
sys.path.append(__dir__)
sys.path.insert(0, os.path.abspath(os.path.join(__dir__, "../..")))
os.environ["FLAGS_allocator_strategy"] = "auto_growth"
import cv2
import copy
import numpy as np
import json
import time
import logging
from PIL import Image
import tools.infer.utility as utility
import tools.infer.predict_rec as predict_rec
import tools.infer.predict_det as predict_det
import tools.infer.predict_cls as predict_cls
from ppocr.utils.utility import get_image_file_list, check_and_read
from ppocr.utils.logging import get_logger
from tools.infer.utility import (
draw_ocr_box_txt,
get_rotate_crop_image,
get_minarea_rect_crop,
slice_generator,
merge_fragmented,
)
logger = get_logger()
class TextSystem(object):
def __init__(self, args):
if not args.show_log:
logger.setLevel(logging.INFO)
self.text_detector = predict_det.TextDetector(args)
self.text_recognizer = predict_rec.TextRecognizer(args)
self.use_angle_cls = args.use_angle_cls
self.drop_score = args.drop_score
if self.use_angle_cls:
self.text_classifier = predict_cls.TextClassifier(args)
self.args = args
self.crop_image_res_index = 0
def draw_crop_rec_res(self, output_dir, img_crop_list, rec_res):
os.makedirs(output_dir, exist_ok=True)
bbox_num = len(img_crop_list)
for bno in range(bbox_num):
cv2.imwrite(
os.path.join(
output_dir, f"mg_crop_{bno+self.crop_image_res_index}.jpg"
),
img_crop_list[bno],
)
logger.debug(f"{bno}, {rec_res[bno]}")
self.crop_image_res_index += bbox_num
def __call__(self, img, cls=True, slice={}):
time_dict = {"det": 0, "rec": 0, "cls": 0, "all": 0}
if img is None:
logger.debug("no valid image provided")
return None, None, time_dict
start = time.time()
ori_im = img.copy()
if slice:
slice_gen = slice_generator(
img,
horizontal_stride=slice["horizontal_stride"],
vertical_stride=slice["vertical_stride"],
)
elapsed = []
dt_slice_boxes = []
for slice_crop, v_start, h_start in slice_gen:
dt_boxes, elapse = self.text_detector(slice_crop, use_slice=True)
if dt_boxes.size:
dt_boxes[:, :, 0] += h_start
dt_boxes[:, :, 1] += v_start
dt_slice_boxes.append(dt_boxes)
elapsed.append(elapse)
dt_boxes = np.concatenate(dt_slice_boxes)
dt_boxes = merge_fragmented(
boxes=dt_boxes,
x_threshold=slice["merge_x_thres"],
y_threshold=slice["merge_y_thres"],
)
elapse = sum(elapsed)
else:
dt_boxes, elapse = self.text_detector(img)
time_dict["det"] = elapse
if dt_boxes is None:
logger.debug("no dt_boxes found, elapsed : {}".format(elapse))
end = time.time()
time_dict["all"] = end - start
return None, None, time_dict
else:
logger.debug(
"dt_boxes num : {}, elapsed : {}".format(len(dt_boxes), elapse)
)
img_crop_list = []
dt_boxes = sorted_boxes(dt_boxes)
for bno in range(len(dt_boxes)):
tmp_box = copy.deepcopy(dt_boxes[bno])
if self.args.det_box_type == "quad":
img_crop = get_rotate_crop_image(ori_im, tmp_box)
else:
img_crop = get_minarea_rect_crop(ori_im, tmp_box)
img_crop_list.append(img_crop)
if self.use_angle_cls and cls:
img_crop_list, angle_list, elapse = self.text_classifier(img_crop_list)
time_dict["cls"] = elapse
logger.debug(
"cls num : {}, elapsed : {}".format(len(img_crop_list), elapse)
)
if len(img_crop_list) > 1000:
logger.debug(
f"rec crops num: {len(img_crop_list)}, time and memory cost may be large."
)
rec_res, elapse = self.text_recognizer(img_crop_list)
time_dict["rec"] = elapse
logger.debug("rec_res num : {}, elapsed : {}".format(len(rec_res), elapse))
if self.args.save_crop_res:
self.draw_crop_rec_res(self.args.crop_res_save_dir, img_crop_list, rec_res)
filter_boxes, filter_rec_res = [], []
for box, rec_result in zip(dt_boxes, rec_res):
text, score = rec_result[0], rec_result[1]
if score >= self.drop_score:
filter_boxes.append(box)
filter_rec_res.append(rec_result)
end = time.time()
time_dict["all"] = end - start
return filter_boxes, filter_rec_res, time_dict
def sorted_boxes(dt_boxes):
"""
Sort text boxes in order from top to bottom, left to right
args:
dt_boxes(array):detected text boxes with shape [4, 2]
return:
sorted boxes(array) with shape [4, 2]
"""
num_boxes = dt_boxes.shape[0]
sorted_boxes = sorted(dt_boxes, key=lambda x: (x[0][1], x[0][0]))
_boxes = list(sorted_boxes)
for i in range(num_boxes - 1):
for j in range(i, -1, -1):
if abs(_boxes[j + 1][0][1] - _boxes[j][0][1]) < 10 and (
_boxes[j + 1][0][0] < _boxes[j][0][0]
):
tmp = _boxes[j]
_boxes[j] = _boxes[j + 1]
_boxes[j + 1] = tmp
else:
break
return _boxes
def main(args):
image_file_list = get_image_file_list(args.image_dir)
image_file_list = image_file_list[args.process_id :: args.total_process_num]
text_sys = TextSystem(args)
is_visualize = True
font_path = args.vis_font_path
drop_score = args.drop_score
draw_img_save_dir = args.draw_img_save_dir
os.makedirs(draw_img_save_dir, exist_ok=True)
save_results = []
logger.info(
"In PP-OCRv3, rec_image_shape parameter defaults to '3, 48, 320', "
"if you are using recognition model with PP-OCRv2 or an older version, please set --rec_image_shape='3,32,320"
)
# warm up 10 times
if args.warmup:
img = np.random.uniform(0, 255, [640, 640, 3]).astype(np.uint8)
for i in range(10):
res = text_sys(img)
total_time = 0
cpu_mem, gpu_mem, gpu_util = 0, 0, 0
_st = time.time()
count = 0
for idx, image_file in enumerate(image_file_list):
img, flag_gif, flag_pdf = check_and_read(image_file)
if not flag_gif and not flag_pdf:
img = cv2.imread(image_file)
if not flag_pdf:
if img is None:
logger.debug("error in loading image:{}".format(image_file))
continue
imgs = [img]
else:
page_num = args.page_num
if page_num > len(img) or page_num == 0:
page_num = len(img)
imgs = img[:page_num]
for index, img in enumerate(imgs):
starttime = time.time()
dt_boxes, rec_res, time_dict = text_sys(img)
elapse = time.time() - starttime
total_time += elapse
if len(imgs) > 1:
logger.debug(
str(idx)
+ "_"
+ str(index)
+ " Predict time of %s: %.3fs" % (image_file, elapse)
)
else:
logger.debug(
str(idx) + " Predict time of %s: %.3fs" % (image_file, elapse)
)
for text, score in rec_res:
logger.debug("{}, {:.3f}".format(text, score))
res = [
{
"transcription": rec_res[i][0],
"points": np.array(dt_boxes[i]).astype(np.int32).tolist(),
}
for i in range(len(dt_boxes))
]
if len(imgs) > 1:
save_pred = (
os.path.basename(image_file)
+ "_"
+ str(index)
+ "\t"
+ json.dumps(res, ensure_ascii=False)
+ "\n"
)
else:
save_pred = (
os.path.basename(image_file)
+ "\t"
+ json.dumps(res, ensure_ascii=False)
+ "\n"
)
save_results.append(save_pred)
if is_visualize:
image = Image.fromarray(cv2.cvtColor(img, cv2.COLOR_BGR2RGB))
boxes = dt_boxes
txts = [rec_res[i][0] for i in range(len(rec_res))]
scores = [rec_res[i][1] for i in range(len(rec_res))]
draw_img = draw_ocr_box_txt(
image,
boxes,
txts,
scores,
drop_score=drop_score,
font_path=font_path,
)
if flag_gif:
save_file = image_file[:-3] + "png"
elif flag_pdf:
save_file = image_file.replace(".pdf", "_" + str(index) + ".png")
else:
save_file = image_file
cv2.imwrite(
os.path.join(draw_img_save_dir, os.path.basename(save_file)),
draw_img[:, :, ::-1],
)
logger.debug(
"The visualized image saved in {}".format(
os.path.join(draw_img_save_dir, os.path.basename(save_file))
)
)
logger.info("The predict total time is {}".format(time.time() - _st))
if args.benchmark:
text_sys.text_detector.autolog.report()
text_sys.text_recognizer.autolog.report()
with open(
os.path.join(draw_img_save_dir, "system_results.txt"), "w", encoding="utf-8"
) as f:
f.writelines(save_results)
if __name__ == "__main__":
args = utility.parse_args()
if args.use_mp:
p_list = []
total_process_num = args.total_process_num
for process_id in range(total_process_num):
cmd = (
[sys.executable, "-u"]
+ sys.argv
+ ["--process_id={}".format(process_id), "--use_mp={}".format(False)]
)
p = subprocess.Popen(cmd, stdout=sys.stdout, stderr=sys.stdout)
p_list.append(p)
for p in p_list:
p.wait()
else:
main(args)

1030
tools/infer/utility.py Normal file

File diff suppressed because it is too large Load Diff

84
tools/infer_cls.py Executable file
View File

@@ -0,0 +1,84 @@
# 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
import os
import sys
__dir__ = os.path.dirname(os.path.abspath(__file__))
sys.path.append(__dir__)
sys.path.insert(0, os.path.abspath(os.path.join(__dir__, "..")))
os.environ["FLAGS_allocator_strategy"] = "auto_growth"
import paddle
from ppocr.data import create_operators, transform
from ppocr.modeling.architectures import build_model
from ppocr.postprocess import build_post_process
from ppocr.utils.save_load import load_model
from ppocr.utils.utility import get_image_file_list
import tools.program as program
def main():
global_config = config["Global"]
# build post process
post_process_class = build_post_process(config["PostProcess"], global_config)
# build model
model = build_model(config["Architecture"])
load_model(config, model)
# create data ops
transforms = []
for op in config["Eval"]["dataset"]["transforms"]:
op_name = list(op)[0]
if "Label" in op_name:
continue
elif op_name == "KeepKeys":
op[op_name]["keep_keys"] = ["image"]
elif op_name == "SSLRotateResize":
op[op_name]["mode"] = "test"
transforms.append(op)
global_config["infer_mode"] = True
ops = create_operators(transforms, global_config)
model.eval()
for file in get_image_file_list(config["Global"]["infer_img"]):
logger.info("infer_img: {}".format(file))
with open(file, "rb") as f:
img = f.read()
data = {"image": img}
batch = transform(data, ops)
images = np.expand_dims(batch[0], axis=0)
images = paddle.to_tensor(images)
preds = model(images)
post_result = post_process_class(preds)
for rec_result in post_result:
logger.info("\t result: {}".format(rec_result))
logger.info("success!")
if __name__ == "__main__":
config, device, logger, vdl_writer = program.preprocess()
main()

136
tools/infer_det.py Executable file
View File

@@ -0,0 +1,136 @@
# 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
import os
import sys
__dir__ = os.path.dirname(os.path.abspath(__file__))
sys.path.append(__dir__)
sys.path.insert(0, os.path.abspath(os.path.join(__dir__, "..")))
os.environ["FLAGS_allocator_strategy"] = "auto_growth"
import cv2
import json
import paddle
from ppocr.data import create_operators, transform
from ppocr.modeling.architectures import build_model
from ppocr.postprocess import build_post_process
from ppocr.utils.save_load import load_model
from ppocr.utils.utility import get_image_file_list
import tools.program as program
def draw_det_res(dt_boxes, config, img, img_name, save_path):
import cv2
src_im = img
for box in dt_boxes:
box = np.array(box).astype(np.int32).reshape((-1, 1, 2))
cv2.polylines(src_im, [box], True, color=(255, 255, 0), thickness=2)
if not os.path.exists(save_path):
os.makedirs(save_path)
save_path = os.path.join(save_path, os.path.basename(img_name))
cv2.imwrite(save_path, src_im)
logger.info("The detected Image saved in {}".format(save_path))
@paddle.no_grad()
def main():
global_config = config["Global"]
# build model
model = build_model(config["Architecture"])
load_model(config, model)
# build post process
post_process_class = build_post_process(config["PostProcess"])
# create data ops
transforms = []
for op in config["Eval"]["dataset"]["transforms"]:
op_name = list(op)[0]
if "Label" in op_name:
continue
elif op_name == "KeepKeys":
op[op_name]["keep_keys"] = ["image", "shape"]
transforms.append(op)
ops = create_operators(transforms, global_config)
save_res_path = config["Global"]["save_res_path"]
if not os.path.exists(os.path.dirname(save_res_path)):
os.makedirs(os.path.dirname(save_res_path))
model.eval()
with open(save_res_path, "wb") as fout:
for file in get_image_file_list(config["Global"]["infer_img"]):
logger.info("infer_img: {}".format(file))
with open(file, "rb") as f:
img = f.read()
data = {"image": img}
batch = transform(data, ops)
images = np.expand_dims(batch[0], axis=0)
shape_list = np.expand_dims(batch[1], axis=0)
images = paddle.to_tensor(images)
preds = model(images)
post_result = post_process_class(preds, shape_list)
src_img = cv2.imread(file)
dt_boxes_json = []
# parser boxes if post_result is dict
if isinstance(post_result, dict):
det_box_json = {}
for k in post_result.keys():
boxes = post_result[k][0]["points"]
dt_boxes_list = []
for box in boxes:
tmp_json = {"transcription": ""}
tmp_json["points"] = np.array(box).tolist()
dt_boxes_list.append(tmp_json)
det_box_json[k] = dt_boxes_list
save_det_path = os.path.dirname(
config["Global"]["save_res_path"]
) + "/det_results_{}/".format(k)
draw_det_res(boxes, config, src_img, file, save_det_path)
else:
boxes = post_result[0]["points"]
dt_boxes_json = []
# write result
for box in boxes:
tmp_json = {"transcription": ""}
tmp_json["points"] = np.array(box).tolist()
dt_boxes_json.append(tmp_json)
save_det_path = (
os.path.dirname(config["Global"]["save_res_path"]) + "/det_results/"
)
draw_det_res(boxes, config, src_img, file, save_det_path)
otstr = file + "\t" + json.dumps(dt_boxes_json) + "\n"
fout.write(otstr.encode())
logger.info("success!")
if __name__ == "__main__":
config, device, logger, vdl_writer = program.preprocess()
main()

170
tools/infer_e2e.py Executable file
View File

@@ -0,0 +1,170 @@
# 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
import os
import sys
__dir__ = os.path.dirname(os.path.abspath(__file__))
sys.path.append(__dir__)
sys.path.insert(0, os.path.abspath(os.path.join(__dir__, "..")))
os.environ["FLAGS_allocator_strategy"] = "auto_growth"
import cv2
import json
import paddle
from ppocr.data import create_operators, transform
from ppocr.modeling.architectures import build_model
from ppocr.postprocess import build_post_process
from ppocr.utils.save_load import load_model
from ppocr.utils.utility import get_image_file_list
import tools.program as program
from PIL import Image, ImageDraw, ImageFont
import math
def draw_e2e_res_for_chinese(
image, boxes, txts, config, img_name, font_path="./doc/simfang.ttf"
):
h, w = image.height, image.width
img_left = image.copy()
img_right = Image.new("RGB", (w, h), (255, 255, 255))
import random
random.seed(0)
draw_left = ImageDraw.Draw(img_left)
draw_right = ImageDraw.Draw(img_right)
for idx, (box, txt) in enumerate(zip(boxes, txts)):
box = np.array(box)
box = [tuple(x) for x in box]
color = (random.randint(0, 255), random.randint(0, 255), random.randint(0, 255))
draw_left.polygon(box, fill=color)
draw_right.polygon(box, outline=color)
font = ImageFont.truetype(font_path, 15, encoding="utf-8")
draw_right.text([box[0][0], box[0][1]], txt, fill=(0, 0, 0), font=font)
img_left = Image.blend(image, img_left, 0.5)
img_show = Image.new("RGB", (w * 2, h), (255, 255, 255))
img_show.paste(img_left, (0, 0, w, h))
img_show.paste(img_right, (w, 0, w * 2, h))
save_e2e_path = os.path.dirname(config["Global"]["save_res_path"]) + "/e2e_results/"
if not os.path.exists(save_e2e_path):
os.makedirs(save_e2e_path)
save_path = os.path.join(save_e2e_path, os.path.basename(img_name))
cv2.imwrite(save_path, np.array(img_show)[:, :, ::-1])
logger.info("The e2e Image saved in {}".format(save_path))
def draw_e2e_res(dt_boxes, strs, config, img, img_name):
if len(dt_boxes) > 0:
src_im = img
for box, str in zip(dt_boxes, strs):
box = box.astype(np.int32).reshape((-1, 1, 2))
cv2.polylines(src_im, [box], True, color=(255, 255, 0), thickness=2)
cv2.putText(
src_im,
str,
org=(int(box[0, 0, 0]), int(box[0, 0, 1])),
fontFace=cv2.FONT_HERSHEY_COMPLEX,
fontScale=0.7,
color=(0, 255, 0),
thickness=1,
)
save_det_path = (
os.path.dirname(config["Global"]["save_res_path"]) + "/e2e_results/"
)
if not os.path.exists(save_det_path):
os.makedirs(save_det_path)
save_path = os.path.join(save_det_path, os.path.basename(img_name))
cv2.imwrite(save_path, src_im)
logger.info("The e2e Image saved in {}".format(save_path))
def main():
global_config = config["Global"]
# build model
model = build_model(config["Architecture"])
load_model(config, model)
# build post process
post_process_class = build_post_process(config["PostProcess"], global_config)
# create data ops
transforms = []
for op in config["Eval"]["dataset"]["transforms"]:
op_name = list(op)[0]
if "Label" in op_name:
continue
elif op_name == "KeepKeys":
op[op_name]["keep_keys"] = ["image", "shape"]
transforms.append(op)
ops = create_operators(transforms, global_config)
save_res_path = config["Global"]["save_res_path"]
if not os.path.exists(os.path.dirname(save_res_path)):
os.makedirs(os.path.dirname(save_res_path))
model.eval()
with open(save_res_path, "wb") as fout:
for file in get_image_file_list(config["Global"]["infer_img"]):
logger.info("infer_img: {}".format(file))
with open(file, "rb") as f:
img = f.read()
data = {"image": img}
batch = transform(data, ops)
images = np.expand_dims(batch[0], axis=0)
shape_list = np.expand_dims(batch[1], axis=0)
images = paddle.to_tensor(images)
preds = model(images)
post_result = post_process_class(preds, shape_list)
points, strs = post_result["points"], post_result["texts"]
# write result
dt_boxes_json = []
for poly, str in zip(points, strs):
tmp_json = {"transcription": str}
tmp_json["points"] = poly.tolist()
dt_boxes_json.append(tmp_json)
otstr = file + "\t" + json.dumps(dt_boxes_json) + "\n"
fout.write(otstr.encode())
src_img = cv2.imread(file)
if global_config["infer_visual_type"] == "EN":
draw_e2e_res(points, strs, config, src_img, file)
elif global_config["infer_visual_type"] == "CN":
src_img = Image.fromarray(cv2.cvtColor(src_img, cv2.COLOR_BGR2RGB))
draw_e2e_res_for_chinese(
src_img,
points,
strs,
config,
file,
font_path="./doc/fonts/simfang.ttf",
)
logger.info("success!")
if __name__ == "__main__":
config, device, logger, vdl_writer = program.preprocess()
main()

186
tools/infer_kie.py Executable file
View File

@@ -0,0 +1,186 @@
# 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
import paddle.nn.functional as F
import os
import sys
__dir__ = os.path.dirname(os.path.abspath(__file__))
sys.path.append(__dir__)
sys.path.insert(0, os.path.abspath(os.path.join(__dir__, "..")))
os.environ["FLAGS_allocator_strategy"] = "auto_growth"
import cv2
import paddle
from ppocr.data import create_operators, transform
from ppocr.modeling.architectures import build_model
from ppocr.utils.save_load import load_model
import tools.program as program
import time
def read_class_list(filepath):
ret = {}
with open(filepath, "r") as f:
lines = f.readlines()
for idx, line in enumerate(lines):
ret[idx] = line.strip("\n")
return ret
def draw_kie_result(batch, node, idx_to_cls, count):
img = batch[6].copy()
boxes = batch[7]
h, w = img.shape[:2]
pred_img = np.ones((h, w * 2, 3), dtype=np.uint8) * 255
max_value, max_idx = paddle.max(node, -1), paddle.argmax(node, -1)
node_pred_label = max_idx.numpy().tolist()
node_pred_score = max_value.numpy().tolist()
for i, box in enumerate(boxes):
if i >= len(node_pred_label):
break
new_box = [
[box[0], box[1]],
[box[2], box[1]],
[box[2], box[3]],
[box[0], box[3]],
]
Pts = np.array([new_box], np.int32)
cv2.polylines(
img, [Pts.reshape((-1, 1, 2))], True, color=(255, 255, 0), thickness=1
)
x_min = int(min([point[0] for point in new_box]))
y_min = int(min([point[1] for point in new_box]))
pred_label = node_pred_label[i]
if pred_label in idx_to_cls:
pred_label = idx_to_cls[pred_label]
pred_score = "{:.2f}".format(node_pred_score[i])
text = pred_label + "(" + pred_score + ")"
cv2.putText(
pred_img,
text,
(x_min * 2, y_min),
cv2.FONT_HERSHEY_SIMPLEX,
0.5,
(255, 0, 0),
1,
)
vis_img = np.ones((h, w * 3, 3), dtype=np.uint8) * 255
vis_img[:, :w] = img
vis_img[:, w:] = pred_img
save_kie_path = os.path.dirname(config["Global"]["save_res_path"]) + "/kie_results/"
if not os.path.exists(save_kie_path):
os.makedirs(save_kie_path)
save_path = os.path.join(save_kie_path, str(count) + ".png")
cv2.imwrite(save_path, vis_img)
logger.info("The Kie Image saved in {}".format(save_path))
def write_kie_result(fout, node, data):
"""
Write infer result to output file, sorted by the predict label of each line.
The format keeps the same as the input with additional score attribute.
"""
import json
label = data["label"]
annotations = json.loads(label)
max_value, max_idx = paddle.max(node, -1), paddle.argmax(node, -1)
node_pred_label = max_idx.numpy().tolist()
node_pred_score = max_value.numpy().tolist()
res = []
for i, label in enumerate(node_pred_label):
pred_score = "{:.2f}".format(node_pred_score[i])
pred_res = {
"label": label,
"transcription": annotations[i]["transcription"],
"score": pred_score,
"points": annotations[i]["points"],
}
res.append(pred_res)
res.sort(key=lambda x: x["label"])
fout.writelines([json.dumps(res, ensure_ascii=False) + "\n"])
def main():
global_config = config["Global"]
# build model
model = build_model(config["Architecture"])
load_model(config, model)
# create data ops
transforms = []
for op in config["Eval"]["dataset"]["transforms"]:
transforms.append(op)
data_dir = config["Eval"]["dataset"]["data_dir"]
ops = create_operators(transforms, global_config)
save_res_path = config["Global"]["save_res_path"]
class_path = config["Global"]["class_path"]
idx_to_cls = read_class_list(class_path)
os.makedirs(os.path.dirname(save_res_path), exist_ok=True)
model.eval()
warmup_times = 0
count_t = []
with open(save_res_path, "w") as fout:
with open(config["Global"]["infer_img"], "rb") as f:
lines = f.readlines()
for index, data_line in enumerate(lines):
if index == 10:
warmup_t = time.time()
data_line = data_line.decode("utf-8")
substr = data_line.strip("\n").split("\t")
img_path, label = data_dir + "/" + substr[0], substr[1]
data = {"img_path": img_path, "label": label}
with open(data["img_path"], "rb") as f:
img = f.read()
data["image"] = img
st = time.time()
batch = transform(data, ops)
batch_pred = [0] * len(batch)
for i in range(len(batch)):
batch_pred[i] = paddle.to_tensor(np.expand_dims(batch[i], axis=0))
st = time.time()
node, edge = model(batch_pred)
node = F.softmax(node, -1)
count_t.append(time.time() - st)
draw_kie_result(batch, node, idx_to_cls, index)
write_kie_result(fout, node, data)
fout.close()
logger.info("success!")
logger.info(
"It took {} s for predict {} images.".format(np.sum(count_t), len(count_t))
)
ips = len(count_t[warmup_times:]) / np.sum(count_t[warmup_times:])
logger.info("The ips is {} images/s".format(ips))
if __name__ == "__main__":
config, device, logger, vdl_writer = program.preprocess()
main()

178
tools/infer_kie_token_ser.py Executable file
View File

@@ -0,0 +1,178 @@
# 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
import os
import sys
__dir__ = os.path.dirname(os.path.abspath(__file__))
sys.path.append(__dir__)
sys.path.insert(0, os.path.abspath(os.path.join(__dir__, "..")))
os.environ["FLAGS_allocator_strategy"] = "auto_growth"
import cv2
import json
import paddle
from ppocr.data import create_operators, transform
from ppocr.modeling.architectures import build_model
from ppocr.postprocess import build_post_process
from ppocr.utils.save_load import load_model
from ppocr.utils.visual import draw_ser_results
from ppocr.utils.utility import get_image_file_list, load_vqa_bio_label_maps
import tools.program as program
def to_tensor(data):
import numbers
from collections import defaultdict
data_dict = defaultdict(list)
to_tensor_idxs = []
for idx, v in enumerate(data):
if isinstance(v, (np.ndarray, paddle.Tensor, numbers.Number)):
if idx not in to_tensor_idxs:
to_tensor_idxs.append(idx)
data_dict[idx].append(v)
for idx in to_tensor_idxs:
data_dict[idx] = paddle.to_tensor(data_dict[idx])
return list(data_dict.values())
class SerPredictor(object):
def __init__(self, config):
global_config = config["Global"]
self.algorithm = config["Architecture"]["algorithm"]
# build post process
self.post_process_class = build_post_process(
config["PostProcess"], global_config
)
# build model
self.model = build_model(config["Architecture"])
load_model(config, self.model, model_type=config["Architecture"]["model_type"])
from paddleocr import PaddleOCR
self.ocr_engine = PaddleOCR(
use_angle_cls=False,
show_log=False,
rec_model_dir=global_config.get("kie_rec_model_dir", None),
det_model_dir=global_config.get("kie_det_model_dir", None),
use_gpu=global_config["use_gpu"],
)
# create data ops
transforms = []
for op in config["Eval"]["dataset"]["transforms"]:
op_name = list(op)[0]
if "Label" in op_name:
op[op_name]["ocr_engine"] = self.ocr_engine
elif op_name == "KeepKeys":
op[op_name]["keep_keys"] = [
"input_ids",
"bbox",
"attention_mask",
"token_type_ids",
"image",
"labels",
"segment_offset_id",
"ocr_info",
"entities",
]
transforms.append(op)
if config["Global"].get("infer_mode", None) is None:
global_config["infer_mode"] = True
self.ops = create_operators(
config["Eval"]["dataset"]["transforms"], global_config
)
self.model.eval()
def __call__(self, data):
with open(data["img_path"], "rb") as f:
img = f.read()
data["image"] = img
batch = transform(data, self.ops)
batch = to_tensor(batch)
preds = self.model(batch)
post_result = self.post_process_class(
preds, segment_offset_ids=batch[6], ocr_infos=batch[7]
)
return post_result, batch
if __name__ == "__main__":
config, device, logger, vdl_writer = program.preprocess()
os.makedirs(config["Global"]["save_res_path"], exist_ok=True)
ser_engine = SerPredictor(config)
if config["Global"].get("infer_mode", None) is False:
data_dir = config["Eval"]["dataset"]["data_dir"]
with open(config["Global"]["infer_img"], "rb") as f:
infer_imgs = f.readlines()
else:
infer_imgs = get_image_file_list(config["Global"]["infer_img"])
with open(
os.path.join(config["Global"]["save_res_path"], "infer_results.txt"),
"w",
encoding="utf-8",
) as fout:
for idx, info in enumerate(infer_imgs):
if config["Global"].get("infer_mode", None) is False:
data_line = info.decode("utf-8")
substr = data_line.strip("\n").split("\t")
img_path = os.path.join(data_dir, substr[0])
data = {"img_path": img_path, "label": substr[1]}
else:
img_path = info
data = {"img_path": img_path}
save_img_path = os.path.join(
config["Global"]["save_res_path"],
os.path.splitext(os.path.basename(img_path))[0] + "_ser.jpg",
)
result, _ = ser_engine(data)
result = result[0]
fout.write(
img_path
+ "\t"
+ json.dumps(
{
"ocr_info": result,
},
ensure_ascii=False,
)
+ "\n"
)
img_res = draw_ser_results(img_path, result)
cv2.imwrite(save_img_path, img_res)
logger.info(
"process: [{}/{}], save result to {}".format(
idx, len(infer_imgs), save_img_path
)
)

226
tools/infer_kie_token_ser_re.py Executable file
View File

@@ -0,0 +1,226 @@
# 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
import os
import sys
__dir__ = os.path.dirname(os.path.abspath(__file__))
sys.path.append(__dir__)
sys.path.insert(0, os.path.abspath(os.path.join(__dir__, "..")))
os.environ["FLAGS_allocator_strategy"] = "auto_growth"
import cv2
import json
import paddle
import paddle.distributed as dist
from ppocr.data import create_operators, transform
from ppocr.modeling.architectures import build_model
from ppocr.postprocess import build_post_process
from ppocr.utils.save_load import load_model
from ppocr.utils.visual import draw_re_results
from ppocr.utils.logging import get_logger
from ppocr.utils.utility import get_image_file_list, load_vqa_bio_label_maps, print_dict
from tools.program import ArgsParser, load_config, merge_config
from tools.infer_kie_token_ser import SerPredictor
class ReArgsParser(ArgsParser):
def __init__(self):
super(ReArgsParser, self).__init__()
self.add_argument(
"-c_ser", "--config_ser", help="ser configuration file to use"
)
self.add_argument(
"-o_ser", "--opt_ser", nargs="+", help="set ser configuration options "
)
def parse_args(self, argv=None):
args = super(ReArgsParser, self).parse_args(argv)
assert (
args.config_ser is not None
), "Please specify --config_ser=ser_configure_file_path."
args.opt_ser = self._parse_opt(args.opt_ser)
return args
def make_input(ser_inputs, ser_results):
entities_labels = {"HEADER": 0, "QUESTION": 1, "ANSWER": 2}
batch_size, max_seq_len = ser_inputs[0].shape[:2]
entities = ser_inputs[8][0]
ser_results = ser_results[0]
assert len(entities) == len(ser_results)
# entities
start = []
end = []
label = []
entity_idx_dict = {}
for i, (res, entity) in enumerate(zip(ser_results, entities)):
if res["pred"] == "O":
continue
entity_idx_dict[len(start)] = i
start.append(entity["start"])
end.append(entity["end"])
label.append(entities_labels[res["pred"]])
entities = np.full([max_seq_len + 1, 3], fill_value=-1, dtype=np.int64)
entities[0, 0] = len(start)
entities[1 : len(start) + 1, 0] = start
entities[0, 1] = len(end)
entities[1 : len(end) + 1, 1] = end
entities[0, 2] = len(label)
entities[1 : len(label) + 1, 2] = label
# relations
head = []
tail = []
for i in range(len(label)):
for j in range(len(label)):
if label[i] == 1 and label[j] == 2:
head.append(i)
tail.append(j)
relations = np.full([len(head) + 1, 2], fill_value=-1, dtype=np.int64)
relations[0, 0] = len(head)
relations[1 : len(head) + 1, 0] = head
relations[0, 1] = len(tail)
relations[1 : len(tail) + 1, 1] = tail
entities = np.expand_dims(entities, axis=0)
entities = np.repeat(entities, batch_size, axis=0)
relations = np.expand_dims(relations, axis=0)
relations = np.repeat(relations, batch_size, axis=0)
# remove ocr_info segment_offset_id and label in ser input
if isinstance(ser_inputs[0], paddle.Tensor):
entities = paddle.to_tensor(entities)
relations = paddle.to_tensor(relations)
ser_inputs = ser_inputs[:5] + [entities, relations]
entity_idx_dict_batch = []
for b in range(batch_size):
entity_idx_dict_batch.append(entity_idx_dict)
return ser_inputs, entity_idx_dict_batch
class SerRePredictor(object):
def __init__(self, config, ser_config):
global_config = config["Global"]
if "infer_mode" in global_config:
ser_config["Global"]["infer_mode"] = global_config["infer_mode"]
self.ser_engine = SerPredictor(ser_config)
# init re model
# build post process
self.post_process_class = build_post_process(
config["PostProcess"], global_config
)
# build model
self.model = build_model(config["Architecture"])
load_model(config, self.model, model_type=config["Architecture"]["model_type"])
self.model.eval()
def __call__(self, data):
ser_results, ser_inputs = self.ser_engine(data)
re_input, entity_idx_dict_batch = make_input(ser_inputs, ser_results)
if self.model.backbone.use_visual_backbone is False:
re_input.pop(4)
preds = self.model(re_input)
post_result = self.post_process_class(
preds, ser_results=ser_results, entity_idx_dict_batch=entity_idx_dict_batch
)
return post_result
def preprocess():
FLAGS = ReArgsParser().parse_args()
config = load_config(FLAGS.config)
config = merge_config(config, FLAGS.opt)
ser_config = load_config(FLAGS.config_ser)
ser_config = merge_config(ser_config, FLAGS.opt_ser)
logger = get_logger()
# check if set use_gpu=True in paddlepaddle cpu version
use_gpu = config["Global"]["use_gpu"]
device = "gpu:{}".format(dist.ParallelEnv().dev_id) if use_gpu else "cpu"
device = paddle.set_device(device)
logger.info("{} re config {}".format("*" * 10, "*" * 10))
print_dict(config, logger)
logger.info("\n")
logger.info("{} ser config {}".format("*" * 10, "*" * 10))
print_dict(ser_config, logger)
logger.info("train with paddle {} and device {}".format(paddle.__version__, device))
return config, ser_config, device, logger
if __name__ == "__main__":
config, ser_config, device, logger = preprocess()
os.makedirs(config["Global"]["save_res_path"], exist_ok=True)
ser_re_engine = SerRePredictor(config, ser_config)
if config["Global"].get("infer_mode", None) is False:
data_dir = config["Eval"]["dataset"]["data_dir"]
with open(config["Global"]["infer_img"], "rb") as f:
infer_imgs = f.readlines()
else:
infer_imgs = get_image_file_list(config["Global"]["infer_img"])
with open(
os.path.join(config["Global"]["save_res_path"], "infer_results.txt"),
"w",
encoding="utf-8",
) as fout:
for idx, info in enumerate(infer_imgs):
if config["Global"].get("infer_mode", None) is False:
data_line = info.decode("utf-8")
substr = data_line.strip("\n").split("\t")
img_path = os.path.join(data_dir, substr[0])
data = {"img_path": img_path, "label": substr[1]}
else:
img_path = info
data = {"img_path": img_path}
save_img_path = os.path.join(
config["Global"]["save_res_path"],
os.path.splitext(os.path.basename(img_path))[0] + "_ser_re.jpg",
)
result = ser_re_engine(data)
result = result[0]
fout.write(img_path + "\t" + json.dumps(result, ensure_ascii=False) + "\n")
img_res = draw_re_results(img_path, result)
cv2.imwrite(save_img_path, img_res)
logger.info(
"process: [{}/{}], save result to {}".format(
idx, len(infer_imgs), save_img_path
)
)

232
tools/infer_rec.py Executable file
View File

@@ -0,0 +1,232 @@
# 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
import os
import sys
import json
__dir__ = os.path.dirname(os.path.abspath(__file__))
sys.path.append(__dir__)
sys.path.insert(0, os.path.abspath(os.path.join(__dir__, "..")))
os.environ["FLAGS_allocator_strategy"] = "auto_growth"
import paddle
from ppocr.data import create_operators, transform
from ppocr.modeling.architectures import build_model
from ppocr.postprocess import build_post_process
from ppocr.utils.save_load import load_model
from ppocr.utils.utility import get_image_file_list
import tools.program as program
def main():
global_config = config["Global"]
if config["Architecture"].get("algorithm") in [
"UniMERNet",
"PP-FormulaNet-S",
"PP-FormulaNet-L",
"PP-FormulaNet_plus-S",
"PP-FormulaNet_plus-M",
"PP-FormulaNet_plus-L",
]:
config["PostProcess"]["is_infer"] = True
# build post process
post_process_class = build_post_process(config["PostProcess"], global_config)
# build model
if hasattr(post_process_class, "character"):
char_num = len(getattr(post_process_class, "character"))
if config["Architecture"]["algorithm"] in [
"Distillation",
]: # distillation model
for key in config["Architecture"]["Models"]:
if (
config["Architecture"]["Models"][key]["Head"]["name"] == "MultiHead"
): # multi head
out_channels_list = {}
if config["PostProcess"]["name"] == "DistillationSARLabelDecode":
char_num = char_num - 2
if config["PostProcess"]["name"] == "DistillationNRTRLabelDecode":
char_num = char_num - 3
out_channels_list["CTCLabelDecode"] = char_num
out_channels_list["SARLabelDecode"] = char_num + 2
out_channels_list["NRTRLabelDecode"] = char_num + 3
config["Architecture"]["Models"][key]["Head"][
"out_channels_list"
] = out_channels_list
else:
config["Architecture"]["Models"][key]["Head"][
"out_channels"
] = char_num
elif config["Architecture"]["Head"]["name"] == "MultiHead": # multi head
out_channels_list = {}
char_num = len(getattr(post_process_class, "character"))
if config["PostProcess"]["name"] == "SARLabelDecode":
char_num = char_num - 2
if config["PostProcess"]["name"] == "NRTRLabelDecode":
char_num = char_num - 3
out_channels_list["CTCLabelDecode"] = char_num
out_channels_list["SARLabelDecode"] = char_num + 2
out_channels_list["NRTRLabelDecode"] = char_num + 3
config["Architecture"]["Head"]["out_channels_list"] = out_channels_list
else: # base rec model
config["Architecture"]["Head"]["out_channels"] = char_num
if config["Architecture"].get("algorithm") in ["LaTeXOCR"]:
config["Architecture"]["Backbone"]["is_predict"] = True
config["Architecture"]["Backbone"]["is_export"] = True
config["Architecture"]["Head"]["is_export"] = True
model = build_model(config["Architecture"])
load_model(config, model)
# create data ops
transforms = []
for op in config["Eval"]["dataset"]["transforms"]:
op_name = list(op)[0]
if "Label" in op_name:
continue
elif op_name in ["RecResizeImg"]:
op[op_name]["infer_mode"] = True
elif op_name == "KeepKeys":
if config["Architecture"]["algorithm"] == "SRN":
op[op_name]["keep_keys"] = [
"image",
"encoder_word_pos",
"gsrm_word_pos",
"gsrm_slf_attn_bias1",
"gsrm_slf_attn_bias2",
]
elif config["Architecture"]["algorithm"] == "SAR":
op[op_name]["keep_keys"] = ["image", "valid_ratio"]
elif config["Architecture"]["algorithm"] == "RobustScanner":
op[op_name]["keep_keys"] = ["image", "valid_ratio", "word_positons"]
else:
op[op_name]["keep_keys"] = ["image"]
transforms.append(op)
global_config["infer_mode"] = True
ops = create_operators(transforms, global_config)
save_res_path = config["Global"].get(
"save_res_path", "./output/rec/predicts_rec.txt"
)
if not os.path.exists(os.path.dirname(save_res_path)):
os.makedirs(os.path.dirname(save_res_path))
model.eval()
infer_imgs = config["Global"]["infer_img"]
infer_list = config["Global"].get("infer_list", None)
with open(save_res_path, "w") as fout:
for file in get_image_file_list(infer_imgs, infer_list=infer_list):
logger.info("infer_img: {}".format(file))
with open(file, "rb") as f:
img = f.read()
if config["Architecture"]["algorithm"] in [
"UniMERNet",
"PP-FormulaNet-S",
"PP-FormulaNet-L",
"PP-FormulaNet_plus-S",
"PP-FormulaNet_plus-M",
"PP-FormulaNet_plus-L",
]:
data = {"image": img, "filename": file}
else:
data = {"image": img}
batch = transform(data, ops)
if config["Architecture"]["algorithm"] == "SRN":
encoder_word_pos_list = np.expand_dims(batch[1], axis=0)
gsrm_word_pos_list = np.expand_dims(batch[2], axis=0)
gsrm_slf_attn_bias1_list = np.expand_dims(batch[3], axis=0)
gsrm_slf_attn_bias2_list = np.expand_dims(batch[4], axis=0)
others = [
paddle.to_tensor(encoder_word_pos_list),
paddle.to_tensor(gsrm_word_pos_list),
paddle.to_tensor(gsrm_slf_attn_bias1_list),
paddle.to_tensor(gsrm_slf_attn_bias2_list),
]
if config["Architecture"]["algorithm"] == "SAR":
valid_ratio = np.expand_dims(batch[-1], axis=0)
img_metas = [paddle.to_tensor(valid_ratio)]
if config["Architecture"]["algorithm"] == "RobustScanner":
valid_ratio = np.expand_dims(batch[1], axis=0)
word_positons = np.expand_dims(batch[2], axis=0)
img_metas = [
paddle.to_tensor(valid_ratio),
paddle.to_tensor(word_positons),
]
if config["Architecture"]["algorithm"] == "CAN":
image_mask = paddle.ones(
(np.expand_dims(batch[0], axis=0).shape), dtype="float32"
)
label = paddle.ones((1, 36), dtype="int64")
images = np.expand_dims(batch[0], axis=0)
images = paddle.to_tensor(images)
if config["Architecture"]["algorithm"] == "SRN":
preds = model(images, others)
elif config["Architecture"]["algorithm"] == "SAR":
preds = model(images, img_metas)
elif config["Architecture"]["algorithm"] == "RobustScanner":
preds = model(images, img_metas)
elif config["Architecture"]["algorithm"] == "CAN":
preds = model([images, image_mask, label])
else:
preds = model(images)
post_result = post_process_class(preds)
info = None
if isinstance(post_result, dict):
rec_info = dict()
for key in post_result:
if len(post_result[key][0]) >= 2:
rec_info[key] = {
"label": post_result[key][0][0],
"score": float(post_result[key][0][1]),
}
info = json.dumps(rec_info, ensure_ascii=False)
elif isinstance(post_result, list) and isinstance(post_result[0], int):
# for RFLearning CNT branch
info = str(post_result[0])
elif config["Architecture"]["algorithm"] in [
"LaTeXOCR",
"UniMERNet",
"PP-FormulaNet-S",
"PP-FormulaNet-L",
"PP-FormulaNet_plus-S",
"PP-FormulaNet_plus-M",
"PP-FormulaNet_plus-L",
]:
info = str(post_result[0])
else:
if len(post_result[0]) >= 2:
info = post_result[0][0] + "\t" + str(post_result[0][1])
if info is not None:
logger.info("\t result: {}".format(info))
fout.write(file + "\t" + info + "\n")
logger.info("success!")
if __name__ == "__main__":
config, device, logger, vdl_writer = program.preprocess()
main()

101
tools/infer_sr.py Executable file
View File

@@ -0,0 +1,101 @@
# 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
import os
import sys
import json
from PIL import Image
import cv2
__dir__ = os.path.dirname(os.path.abspath(__file__))
sys.path.insert(0, __dir__)
sys.path.insert(0, os.path.abspath(os.path.join(__dir__, "..")))
os.environ["FLAGS_allocator_strategy"] = "auto_growth"
import paddle
from ppocr.data import create_operators, transform
from ppocr.modeling.architectures import build_model
from ppocr.postprocess import build_post_process
from ppocr.utils.save_load import load_model
from ppocr.utils.utility import get_image_file_list
import tools.program as program
def main():
global_config = config["Global"]
# build post process
post_process_class = build_post_process(config["PostProcess"], global_config)
# sr transform
config["Architecture"]["Transform"]["infer_mode"] = True
model = build_model(config["Architecture"])
load_model(config, model)
# create data ops
transforms = []
for op in config["Eval"]["dataset"]["transforms"]:
op_name = list(op)[0]
if "Label" in op_name:
continue
elif op_name in ["SRResize"]:
op[op_name]["infer_mode"] = True
elif op_name == "KeepKeys":
op[op_name]["keep_keys"] = ["img_lr"]
transforms.append(op)
global_config["infer_mode"] = True
ops = create_operators(transforms, global_config)
save_visual_path = config["Global"].get("save_visual", "infer_result/")
if not os.path.exists(os.path.dirname(save_visual_path)):
os.makedirs(os.path.dirname(save_visual_path))
model.eval()
for file in get_image_file_list(config["Global"]["infer_img"]):
logger.info("infer_img: {}".format(file))
img = Image.open(file).convert("RGB")
data = {"image_lr": img}
batch = transform(data, ops)
images = np.expand_dims(batch[0], axis=0)
images = paddle.to_tensor(images)
preds = model(images)
sr_img = preds["sr_img"][0]
lr_img = preds["lr_img"][0]
fm_sr = (sr_img.numpy() * 255).transpose(1, 2, 0).astype(np.uint8)
fm_lr = (lr_img.numpy() * 255).transpose(1, 2, 0).astype(np.uint8)
img_name_pure = os.path.split(file)[-1]
cv2.imwrite(
"{}/sr_{}".format(save_visual_path, img_name_pure), fm_sr[:, :, ::-1]
)
logger.info(
"The visualized image saved in infer_result/sr_{}".format(img_name_pure)
)
logger.info("success!")
if __name__ == "__main__":
config, device, logger, vdl_writer = program.preprocess()
main()

120
tools/infer_table.py Normal file
View File

@@ -0,0 +1,120 @@
# 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
import os
import sys
import json
__dir__ = os.path.dirname(os.path.abspath(__file__))
sys.path.append(__dir__)
sys.path.insert(0, os.path.abspath(os.path.join(__dir__, "..")))
os.environ["FLAGS_allocator_strategy"] = "auto_growth"
import paddle
from paddle.jit import to_static
from ppocr.data import create_operators, transform
from ppocr.modeling.architectures import build_model
from ppocr.postprocess import build_post_process
from ppocr.utils.save_load import load_model
from ppocr.utils.utility import get_image_file_list
from ppocr.utils.visual import draw_rectangle
from tools.infer.utility import draw_boxes
import tools.program as program
import cv2
@paddle.no_grad()
def main(config, device, logger, vdl_writer):
global_config = config["Global"]
# build post process
post_process_class = build_post_process(config["PostProcess"], global_config)
# build model
if hasattr(post_process_class, "character"):
config["Architecture"]["Head"]["out_channels"] = len(
getattr(post_process_class, "character")
)
model = build_model(config["Architecture"])
algorithm = config["Architecture"]["algorithm"]
load_model(config, model)
# create data ops
transforms = []
for op in config["Eval"]["dataset"]["transforms"]:
op_name = list(op)[0]
if "Encode" in op_name:
continue
if op_name == "KeepKeys":
op[op_name]["keep_keys"] = ["image", "shape"]
transforms.append(op)
global_config["infer_mode"] = True
ops = create_operators(transforms, global_config)
save_res_path = config["Global"]["save_res_path"]
os.makedirs(save_res_path, exist_ok=True)
model.eval()
with open(
os.path.join(save_res_path, "infer.txt"), mode="w", encoding="utf-8"
) as f_w:
for file in get_image_file_list(config["Global"]["infer_img"]):
logger.info("infer_img: {}".format(file))
with open(file, "rb") as f:
img = f.read()
data = {"image": img}
batch = transform(data, ops)
images = np.expand_dims(batch[0], axis=0)
shape_list = np.expand_dims(batch[1], axis=0)
images = paddle.to_tensor(images)
preds = model(images)
post_result = post_process_class(preds, [shape_list])
structure_str_list = post_result["structure_batch_list"][0]
bbox_list = post_result["bbox_batch_list"][0]
structure_str_list = structure_str_list[0]
structure_str_list = (
["<html>", "<body>", "<table>"]
+ structure_str_list
+ ["</table>", "</body>", "</html>"]
)
bbox_list_str = json.dumps(bbox_list.tolist())
logger.info("result: {}, {}".format(structure_str_list, bbox_list_str))
f_w.write("result: {}, {}\n".format(structure_str_list, bbox_list_str))
if len(bbox_list) > 0 and len(bbox_list[0]) == 4:
img = draw_rectangle(file, bbox_list)
else:
img = draw_boxes(cv2.imread(file), bbox_list)
cv2.imwrite(os.path.join(save_res_path, os.path.basename(file)), img)
logger.info("save result to {}".format(save_res_path))
logger.info("success!")
if __name__ == "__main__":
config, device, logger, vdl_writer = program.preprocess()
main(config, device, logger, vdl_writer)

121
tools/naive_sync_bn.py Normal file
View File

@@ -0,0 +1,121 @@
# Copyright (c) 2024 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.distributed as dist
import math
import paddle
import paddle.nn as nn
class _AllReduce(paddle.autograd.PyLayer):
@staticmethod
def forward(ctx, input):
input_list = [paddle.zeros_like(input) for k in range(dist.get_world_size())]
# Use allgather instead of allreduce since I don't trust in-place operations ..
dist.all_gather(input_list, input, sync_op=True)
inputs = paddle.stack(input_list, axis=0)
return paddle.sum(inputs, axis=0)
@staticmethod
def backward(ctx, grad_output):
dist.all_reduce(grad_output, sync_op=True)
return grad_output
def differentiable_all_reduce(input):
"""
Differentiable counterpart of `dist.all_reduce`.
"""
if (
not dist.is_available()
or not dist.is_initialized()
or dist.get_world_size() == 1
):
return input
return _AllReduce.apply(input)
class NaiveSyncBatchNorm(nn.BatchNorm2D):
def __init__(self, *args, stats_mode="", **kwargs):
super().__init__(*args, **kwargs)
assert stats_mode in ["", "N"]
self._stats_mode = stats_mode
def forward(self, input):
if dist.get_world_size() == 1 or not self.training:
return super().forward(input)
B, C = input.shape[0], input.shape[1]
mean = paddle.mean(input, axis=[0, 2, 3])
meansqr = paddle.mean(input * input, axis=[0, 2, 3])
if self._stats_mode == "":
assert (
B > 0
), 'SyncBatchNorm(stats_mode="") does not support zero batch size.'
vec = paddle.concat([mean, meansqr], axis=0)
vec = differentiable_all_reduce(vec) * (1.0 / dist.get_world_size())
mean, meansqr = paddle.split(vec, [C, C])
momentum = (
1 - self._momentum
) # NOTE: paddle has reverse momentum definition
else:
if B == 0:
vec = paddle.zeros([2 * C + 1], dtype=mean.dtype)
vec = vec + input.sum() # make sure there is gradient w.r.t input
else:
vec = paddle.concat(
[
mean,
meansqr,
paddle.ones([1], dtype=mean.dtype),
],
axis=0,
)
vec = differentiable_all_reduce(vec * B)
total_batch = vec[-1].detach()
momentum = total_batch.clip(max=1) * (
1 - self._momentum
) # no update if total_batch is 0
mean, meansqr, _ = paddle.split(
vec / total_batch.clip(min=1), [C, C, int(vec.shape[0] - 2 * C)]
) # avoid div-by-zero
var = meansqr - mean * mean
invstd = paddle.rsqrt(var + self._epsilon)
scale = self.weight * invstd
bias = self.bias - mean * scale
scale = scale.reshape([1, -1, 1, 1])
bias = bias.reshape([1, -1, 1, 1])
tmp_mean = self._mean + momentum * (mean.detach() - self._mean)
self._mean.set_value(tmp_mean)
tmp_variance = self._variance + (momentum * (var.detach() - self._variance))
self._variance.set_value(tmp_variance)
ret = input * scale + bias
return ret
def convert_syncbn(model):
for n, m in model.named_children():
if isinstance(m, nn.layer.norm._BatchNormBase):
syncbn = NaiveSyncBatchNorm(
m._num_features, m._momentum, m._epsilon, m._weight_attr, m._bias_attr
)
setattr(model, n, syncbn)
else:
convert_syncbn(m)

937
tools/program.py Executable file
View File

@@ -0,0 +1,937 @@
# 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 gc
import sys
import platform
import yaml
import time
import datetime
import paddle
import paddle.distributed as dist
from tqdm import tqdm
import cv2
import numpy as np
import copy
from argparse import ArgumentParser, RawDescriptionHelpFormatter
from ppocr.utils.stats import TrainingStats
from ppocr.utils.save_load import save_model
from ppocr.utils.utility import print_dict, AverageMeter
from ppocr.utils.logging import get_logger
from ppocr.utils.loggers import WandbLogger, Loggers
from ppocr.utils import profiler
from ppocr.data import build_dataloader
from ppocr.utils.export_model import export
class ArgsParser(ArgumentParser):
def __init__(self):
super(ArgsParser, self).__init__(formatter_class=RawDescriptionHelpFormatter)
self.add_argument("-c", "--config", help="configuration file to use")
self.add_argument("-o", "--opt", nargs="+", help="set configuration options")
self.add_argument(
"-p",
"--profiler_options",
type=str,
default=None,
help="The option of profiler, which should be in format "
'"key1=value1;key2=value2;key3=value3".',
)
def parse_args(self, argv=None):
args = super(ArgsParser, self).parse_args(argv)
assert args.config is not None, "Please specify --config=configure_file_path."
args.opt = self._parse_opt(args.opt)
return args
def _parse_opt(self, opts):
config = {}
if not opts:
return config
for s in opts:
s = s.strip()
k, v = s.split("=")
config[k] = yaml.load(v, Loader=yaml.Loader)
return config
def load_config(file_path):
"""
Load config from yml/yaml file.
Args:
file_path (str): Path of the config file to be loaded.
Returns: global config
"""
_, ext = os.path.splitext(file_path)
assert ext in [".yml", ".yaml"], "only support yaml files for now"
config = yaml.load(open(file_path, "rb"), Loader=yaml.Loader)
return config
def merge_config(config, opts):
"""
Merge config into global config.
Args:
config (dict): Config to be merged.
Returns: global config
"""
for key, value in opts.items():
if "." not in key:
if isinstance(value, dict) and key in config:
config[key].update(value)
else:
config[key] = value
else:
sub_keys = key.split(".")
assert sub_keys[0] in config, (
"the sub_keys can only be one of global_config: {}, but get: "
"{}, please check your running command".format(
config.keys(), sub_keys[0]
)
)
cur = config[sub_keys[0]]
for idx, sub_key in enumerate(sub_keys[1:]):
if idx == len(sub_keys) - 2:
cur[sub_key] = value
else:
cur = cur[sub_key]
return config
def check_device(use_gpu, use_xpu=False, use_npu=False, use_mlu=False, use_gcu=False):
"""
Log error and exit when set use_gpu=true in paddlepaddle
cpu version.
"""
err = (
"Config {} cannot be set as true while your paddle "
"is not compiled with {} ! \nPlease try: \n"
"\t1. Install paddlepaddle to run model on {} \n"
"\t2. Set {} as false in config file to run "
"model on CPU"
)
try:
if use_gpu and use_xpu:
print("use_xpu and use_gpu can not both be true.")
if use_gpu and not paddle.is_compiled_with_cuda():
print(err.format("use_gpu", "cuda", "gpu", "use_gpu"))
sys.exit(1)
if use_xpu and not paddle.device.is_compiled_with_xpu():
print(err.format("use_xpu", "xpu", "xpu", "use_xpu"))
sys.exit(1)
if use_npu:
if (
int(paddle.version.major) != 0
and int(paddle.version.major) <= 2
and int(paddle.version.minor) <= 4
):
if not paddle.device.is_compiled_with_npu():
print(err.format("use_npu", "npu", "npu", "use_npu"))
sys.exit(1)
# is_compiled_with_npu() has been updated after paddle-2.4
else:
if not paddle.device.is_compiled_with_custom_device("npu"):
print(err.format("use_npu", "npu", "npu", "use_npu"))
sys.exit(1)
if use_mlu and not paddle.device.is_compiled_with_mlu():
print(err.format("use_mlu", "mlu", "mlu", "use_mlu"))
sys.exit(1)
if use_gcu and not paddle.device.is_compiled_with_custom_device("gcu"):
print(err.format("use_gcu", "gcu", "gcu", "use_gcu"))
sys.exit(1)
except Exception as e:
pass
def to_float32(preds):
if isinstance(preds, dict):
for k in preds:
if isinstance(preds[k], dict) or isinstance(preds[k], list):
preds[k] = to_float32(preds[k])
elif isinstance(preds[k], paddle.Tensor):
preds[k] = preds[k].astype(paddle.float32)
elif isinstance(preds, list):
for k in range(len(preds)):
if isinstance(preds[k], dict):
preds[k] = to_float32(preds[k])
elif isinstance(preds[k], list):
preds[k] = to_float32(preds[k])
elif isinstance(preds[k], paddle.Tensor):
preds[k] = preds[k].astype(paddle.float32)
elif isinstance(preds, paddle.Tensor):
preds = preds.astype(paddle.float32)
return preds
def train(
config,
train_dataloader,
valid_dataloader,
device,
model,
loss_class,
optimizer,
lr_scheduler,
post_process_class,
eval_class,
pre_best_model_dict,
logger,
step_pre_epoch,
log_writer=None,
scaler=None,
amp_level="O2",
amp_custom_black_list=[],
amp_custom_white_list=[],
amp_dtype="float16",
):
cal_metric_during_train = config["Global"].get("cal_metric_during_train", False)
calc_epoch_interval = config["Global"].get("calc_epoch_interval", 1)
log_smooth_window = config["Global"]["log_smooth_window"]
epoch_num = config["Global"]["epoch_num"]
print_batch_step = config["Global"]["print_batch_step"]
eval_batch_step = config["Global"]["eval_batch_step"]
eval_batch_epoch = config["Global"].get("eval_batch_epoch", None)
profiler_options = config["profiler_options"]
print_mem_info = config["Global"].get("print_mem_info", True)
uniform_output_enabled = config["Global"].get("uniform_output_enabled", False)
global_step = 0
if "global_step" in pre_best_model_dict:
global_step = pre_best_model_dict["global_step"]
start_eval_step = 0
if isinstance(eval_batch_step, list) and len(eval_batch_step) >= 2:
start_eval_step = eval_batch_step[0] if not eval_batch_epoch else 0
eval_batch_step = (
eval_batch_step[1]
if not eval_batch_epoch
else step_pre_epoch * eval_batch_epoch
)
if len(valid_dataloader) == 0:
logger.info(
"No Images in eval dataset, evaluation during training "
"will be disabled"
)
start_eval_step = 1e111
logger.info(
"During the training process, after the {}th iteration, "
"an evaluation is run every {} iterations".format(
start_eval_step, eval_batch_step
)
)
save_epoch_step = config["Global"]["save_epoch_step"]
save_model_dir = config["Global"]["save_model_dir"]
if not os.path.exists(save_model_dir):
os.makedirs(save_model_dir)
main_indicator = eval_class.main_indicator
best_model_dict = {main_indicator: 0}
best_model_dict.update(pre_best_model_dict)
train_stats = TrainingStats(log_smooth_window, ["lr"])
model_average = False
model.train()
use_srn = config["Architecture"]["algorithm"] == "SRN"
extra_input_models = [
"SRN",
"NRTR",
"SAR",
"SEED",
"SVTR",
"SVTR_LCNet",
"SPIN",
"VisionLAN",
"RobustScanner",
"RFL",
"DRRG",
"SATRN",
"SVTR_HGNet",
"ParseQ",
"CPPD",
]
extra_input = False
if config["Architecture"]["algorithm"] == "Distillation":
for key in config["Architecture"]["Models"]:
extra_input = (
extra_input
or config["Architecture"]["Models"][key]["algorithm"]
in extra_input_models
)
else:
extra_input = config["Architecture"]["algorithm"] in extra_input_models
try:
model_type = config["Architecture"]["model_type"]
except:
model_type = None
algorithm = config["Architecture"]["algorithm"]
start_epoch = (
best_model_dict["start_epoch"] if "start_epoch" in best_model_dict else 1
)
total_samples = 0
train_reader_cost = 0.0
train_batch_cost = 0.0
reader_start = time.time()
eta_meter = AverageMeter()
max_iter = (
len(train_dataloader) - 1
if platform.system() == "Windows"
else len(train_dataloader)
)
for epoch in range(start_epoch, epoch_num + 1):
if train_dataloader.dataset.need_reset:
train_dataloader = build_dataloader(
config, "Train", device, logger, seed=epoch
)
max_iter = (
len(train_dataloader) - 1
if platform.system() == "Windows"
else len(train_dataloader)
)
for idx, batch in enumerate(train_dataloader):
model.train()
profiler.add_profiler_step(profiler_options)
train_reader_cost += time.time() - reader_start
if idx >= max_iter:
break
lr = optimizer.get_lr()
images = batch[0]
if use_srn:
model_average = True
# use amp
if scaler:
with paddle.amp.auto_cast(
level=amp_level,
custom_black_list=amp_custom_black_list,
custom_white_list=amp_custom_white_list,
dtype=amp_dtype,
):
if model_type == "table" or extra_input:
preds = model(images, data=batch[1:])
elif model_type in ["kie"]:
preds = model(batch)
elif algorithm in ["CAN"]:
preds = model(batch[:3])
elif algorithm in [
"LaTeXOCR",
"UniMERNet",
"PP-FormulaNet-S",
"PP-FormulaNet-L",
"PP-FormulaNet_plus-S",
"PP-FormulaNet_plus-M",
"PP-FormulaNet_plus-L",
]:
preds = model(batch)
else:
preds = model(images)
preds = to_float32(preds)
loss = loss_class(preds, batch)
avg_loss = loss["loss"]
scaled_avg_loss = scaler.scale(avg_loss)
scaled_avg_loss.backward()
scaler.minimize(optimizer, scaled_avg_loss)
else:
if model_type == "table" or extra_input:
preds = model(images, data=batch[1:])
elif model_type in ["kie", "sr"]:
preds = model(batch)
elif algorithm in ["CAN"]:
preds = model(batch[:3])
elif algorithm in [
"LaTeXOCR",
"UniMERNet",
"PP-FormulaNet-S",
"PP-FormulaNet-L",
"PP-FormulaNet_plus-S",
"PP-FormulaNet_plus-M",
"PP-FormulaNet_plus-L",
]:
preds = model(batch)
else:
preds = model(images)
loss = loss_class(preds, batch)
avg_loss = loss["loss"]
avg_loss.backward()
optimizer.step()
optimizer.clear_grad()
if (
cal_metric_during_train and epoch % calc_epoch_interval == 0
): # only rec and cls need
batch = [item.numpy() for item in batch]
if model_type in ["kie", "sr"]:
eval_class(preds, batch)
elif model_type in ["table"]:
post_result = post_process_class(preds, batch)
eval_class(post_result, batch)
elif algorithm in ["CAN"]:
model_type = "can"
eval_class(preds[0], batch[2:], epoch_reset=(idx == 0))
elif algorithm in ["LaTeXOCR"]:
model_type = "latexocr"
post_result = post_process_class(preds, batch[1], mode="train")
eval_class(post_result[0], post_result[1], epoch_reset=(idx == 0))
elif algorithm in ["UniMERNet"]:
model_type = "unimernet"
post_result = post_process_class(preds[0], batch[1], mode="train")
eval_class(post_result[0], post_result[1], epoch_reset=(idx == 0))
elif algorithm in [
"PP-FormulaNet-S",
"PP-FormulaNet-L",
"PP-FormulaNet_plus-S",
"PP-FormulaNet_plus-M",
"PP-FormulaNet_plus-L",
]:
model_type = "pp_formulanet"
post_result = post_process_class(preds[0], batch[1], mode="train")
eval_class(post_result[0], post_result[1], epoch_reset=(idx == 0))
else:
if config["Loss"]["name"] in [
"MultiLoss",
"MultiLoss_v2",
]: # for multi head loss
post_result = post_process_class(
preds["ctc"], batch[1]
) # for CTC head out
elif config["Loss"]["name"] in ["VLLoss"]:
post_result = post_process_class(preds, batch[1], batch[-1])
else:
post_result = post_process_class(preds, batch[1])
eval_class(post_result, batch)
metric = eval_class.get_metric()
train_stats.update(metric)
train_batch_time = time.time() - reader_start
train_batch_cost += train_batch_time
eta_meter.update(train_batch_time)
global_step += 1
total_samples += len(images)
if not isinstance(lr_scheduler, float):
lr_scheduler.step()
# logger and visualdl
stats = {
k: float(v) if v.shape == [] else v.numpy().mean()
for k, v in loss.items()
}
stats["lr"] = lr
train_stats.update(stats)
if log_writer is not None and dist.get_rank() == 0:
log_writer.log_metrics(
metrics=train_stats.get(), prefix="TRAIN", step=global_step
)
if (global_step > 0 and global_step % print_batch_step == 0) or (
idx >= len(train_dataloader) - 1
):
logs = train_stats.log()
eta_sec = (
(epoch_num + 1 - epoch) * len(train_dataloader) - idx - 1
) * eta_meter.avg
eta_sec_format = str(datetime.timedelta(seconds=int(eta_sec)))
max_mem_reserved_str = ""
max_mem_allocated_str = ""
if paddle.device.is_compiled_with_cuda() and print_mem_info:
max_mem_reserved_str = f", max_mem_reserved: {paddle.device.cuda.max_memory_reserved() // (1024 ** 2)} MB,"
max_mem_allocated_str = f" max_mem_allocated: {paddle.device.cuda.max_memory_allocated() // (1024 ** 2)} MB"
strs = (
"epoch: [{}/{}], global_step: {}, {}, avg_reader_cost: "
"{:.5f} s, avg_batch_cost: {:.5f} s, avg_samples: {}, "
"ips: {:.5f} samples/s, eta: {}{}{}".format(
epoch,
epoch_num,
global_step,
logs,
train_reader_cost / print_batch_step,
train_batch_cost / print_batch_step,
total_samples / print_batch_step,
total_samples / train_batch_cost,
eta_sec_format,
max_mem_reserved_str,
max_mem_allocated_str,
)
)
logger.info(strs)
total_samples = 0
train_reader_cost = 0.0
train_batch_cost = 0.0
# eval
if (
global_step > start_eval_step
and (global_step - start_eval_step) % eval_batch_step == 0
and dist.get_rank() == 0
):
if model_average:
Model_Average = paddle.incubate.ModelAverage(
0.15,
parameters=model.parameters(),
min_average_window=10000,
max_average_window=15625,
)
Model_Average.apply()
cur_metric = eval(
model,
valid_dataloader,
post_process_class,
eval_class,
model_type,
extra_input=extra_input,
scaler=scaler,
amp_level=amp_level,
amp_custom_black_list=amp_custom_black_list,
amp_custom_white_list=amp_custom_white_list,
amp_dtype=amp_dtype,
)
cur_metric_str = "cur metric, {}".format(
", ".join(["{}: {}".format(k, v) for k, v in cur_metric.items()])
)
logger.info(cur_metric_str)
# logger metric
if log_writer is not None:
log_writer.log_metrics(
metrics=cur_metric, prefix="EVAL", step=global_step
)
if cur_metric[main_indicator] >= best_model_dict[main_indicator]:
best_model_dict.update(cur_metric)
best_model_dict["best_epoch"] = epoch
prefix = "best_accuracy"
if uniform_output_enabled:
export(
config,
model,
os.path.join(save_model_dir, prefix, "inference"),
)
gc.collect()
model_info = {"epoch": epoch, "metric": best_model_dict}
else:
model_info = None
save_model(
model,
optimizer,
(
os.path.join(save_model_dir, prefix)
if uniform_output_enabled
else save_model_dir
),
logger,
config,
is_best=True,
prefix=prefix,
save_model_info=model_info,
best_model_dict=best_model_dict,
epoch=epoch,
global_step=global_step,
)
best_str = "best metric, {}".format(
", ".join(
["{}: {}".format(k, v) for k, v in best_model_dict.items()]
)
)
logger.info(best_str)
# logger best metric
if log_writer is not None:
log_writer.log_metrics(
metrics={
"best_{}".format(main_indicator): best_model_dict[
main_indicator
]
},
prefix="EVAL",
step=global_step,
)
log_writer.log_model(
is_best=True, prefix="best_accuracy", metadata=best_model_dict
)
reader_start = time.time()
if dist.get_rank() == 0:
prefix = "latest"
if uniform_output_enabled:
export(config, model, os.path.join(save_model_dir, prefix, "inference"))
gc.collect()
model_info = {"epoch": epoch, "metric": best_model_dict}
else:
model_info = None
save_model(
model,
optimizer,
(
os.path.join(save_model_dir, prefix)
if uniform_output_enabled
else save_model_dir
),
logger,
config,
is_best=False,
prefix=prefix,
save_model_info=model_info,
best_model_dict=best_model_dict,
epoch=epoch,
global_step=global_step,
)
if log_writer is not None:
log_writer.log_model(is_best=False, prefix="latest")
if dist.get_rank() == 0 and epoch > 0 and epoch % save_epoch_step == 0:
prefix = "iter_epoch_{}".format(epoch)
if uniform_output_enabled:
export(config, model, os.path.join(save_model_dir, prefix, "inference"))
gc.collect()
model_info = {"epoch": epoch, "metric": best_model_dict}
else:
model_info = None
save_model(
model,
optimizer,
(
os.path.join(save_model_dir, prefix)
if uniform_output_enabled
else save_model_dir
),
logger,
config,
is_best=False,
prefix=prefix,
save_model_info=model_info,
best_model_dict=best_model_dict,
epoch=epoch,
global_step=global_step,
done_flag=epoch == config["Global"]["epoch_num"],
)
if log_writer is not None:
log_writer.log_model(
is_best=False, prefix="iter_epoch_{}".format(epoch)
)
best_str = "best metric, {}".format(
", ".join(["{}: {}".format(k, v) for k, v in best_model_dict.items()])
)
logger.info(best_str)
if dist.get_rank() == 0 and log_writer is not None:
log_writer.close()
return
def eval(
model,
valid_dataloader,
post_process_class,
eval_class,
model_type=None,
extra_input=False,
scaler=None,
amp_level="O2",
amp_custom_black_list=[],
amp_custom_white_list=[],
amp_dtype="float16",
):
model.eval()
with paddle.no_grad():
total_frame = 0.0
total_time = 0.0
pbar = tqdm(
total=len(valid_dataloader), desc="eval model:", position=0, leave=True
)
max_iter = (
len(valid_dataloader) - 1
if platform.system() == "Windows"
else len(valid_dataloader)
)
sum_images = 0
for idx, batch in enumerate(valid_dataloader):
if idx >= max_iter:
break
images = batch[0]
start = time.time()
# use amp
if scaler:
with paddle.amp.auto_cast(
level=amp_level,
custom_black_list=amp_custom_black_list,
dtype=amp_dtype,
):
if model_type == "table" or extra_input:
preds = model(images, data=batch[1:])
elif model_type in ["kie"]:
preds = model(batch)
elif model_type in ["can"]:
preds = model(batch[:3])
elif model_type in ["latexocr"]:
preds = model(batch)
elif model_type in ["sr"]:
preds = model(batch)
sr_img = preds["sr_img"]
lr_img = preds["lr_img"]
else:
preds = model(images)
preds = to_float32(preds)
else:
if model_type == "table" or extra_input:
preds = model(images, data=batch[1:])
elif model_type in ["kie"]:
preds = model(batch)
elif model_type in ["can"]:
preds = model(batch[:3])
elif model_type in ["latexocr", "unimernet", "pp_formulanet"]:
preds = model(batch)
elif model_type in ["sr"]:
preds = model(batch)
sr_img = preds["sr_img"]
lr_img = preds["lr_img"]
else:
preds = model(images)
batch_numpy = []
for item in batch:
if isinstance(item, paddle.Tensor):
batch_numpy.append(item.numpy())
else:
batch_numpy.append(item)
# Obtain usable results from post-processing methods
total_time += time.time() - start
# Evaluate the results of the current batch
if model_type in ["table", "kie"]:
if post_process_class is None:
eval_class(preds, batch_numpy)
else:
post_result = post_process_class(preds, batch_numpy)
eval_class(post_result, batch_numpy)
elif model_type in ["sr"]:
eval_class(preds, batch_numpy)
elif model_type in ["can"]:
eval_class(preds[0], batch_numpy[2:], epoch_reset=(idx == 0))
elif model_type in ["latexocr", "unimernet", "pp_formulanet"]:
post_result = post_process_class(preds, batch[1], "eval")
eval_class(post_result[0], post_result[1], epoch_reset=(idx == 0))
else:
post_result = post_process_class(preds, batch_numpy[1])
eval_class(post_result, batch_numpy)
pbar.update(1)
total_frame += len(images)
sum_images += 1
# Get final metriceg. acc or hmean
metric = eval_class.get_metric()
pbar.close()
model.train()
# Avoid ZeroDivisionError
if total_time > 0:
metric["fps"] = total_frame / total_time
else:
metric["fps"] = 0 # or set to a fallback value
return metric
def update_center(char_center, post_result, preds):
result, label = post_result
feats, logits = preds
logits = paddle.argmax(logits, axis=-1)
feats = feats.numpy()
logits = logits.numpy()
for idx_sample in range(len(label)):
if result[idx_sample][0] == label[idx_sample][0]:
feat = feats[idx_sample]
logit = logits[idx_sample]
for idx_time in range(len(logit)):
index = logit[idx_time]
if index in char_center.keys():
char_center[index][0] = (
char_center[index][0] * char_center[index][1] + feat[idx_time]
) / (char_center[index][1] + 1)
char_center[index][1] += 1
else:
char_center[index] = [feat[idx_time], 1]
return char_center
def get_center(model, eval_dataloader, post_process_class):
pbar = tqdm(total=len(eval_dataloader), desc="get center:")
max_iter = (
len(eval_dataloader) - 1
if platform.system() == "Windows"
else len(eval_dataloader)
)
char_center = dict()
for idx, batch in enumerate(eval_dataloader):
if idx >= max_iter:
break
images = batch[0]
start = time.time()
preds = model(images)
batch = [item.numpy() for item in batch]
# Obtain usable results from post-processing methods
post_result = post_process_class(preds, batch[1])
# update char_center
char_center = update_center(char_center, post_result, preds)
pbar.update(1)
pbar.close()
for key in char_center.keys():
char_center[key] = char_center[key][0]
return char_center
def preprocess(is_train=False):
FLAGS = ArgsParser().parse_args()
profiler_options = FLAGS.profiler_options
config = load_config(FLAGS.config)
config = merge_config(config, FLAGS.opt)
profile_dic = {"profiler_options": FLAGS.profiler_options}
config = merge_config(config, profile_dic)
if is_train:
# save_config
save_model_dir = config["Global"]["save_model_dir"]
os.makedirs(save_model_dir, exist_ok=True)
with open(os.path.join(save_model_dir, "config.yml"), "w") as f:
yaml.dump(dict(config), f, default_flow_style=False, sort_keys=False)
log_file = "{}/train.log".format(save_model_dir)
else:
log_file = None
log_ranks = config["Global"].get("log_ranks", "0")
logger = get_logger(log_file=log_file, log_ranks=log_ranks)
# check if set use_gpu=True in paddlepaddle cpu version
use_gpu = config["Global"].get("use_gpu", False)
use_xpu = config["Global"].get("use_xpu", False)
use_npu = config["Global"].get("use_npu", False)
use_mlu = config["Global"].get("use_mlu", False)
use_gcu = config["Global"].get("use_gcu", False)
alg = config["Architecture"]["algorithm"]
assert alg in [
"EAST",
"DB",
"SAST",
"Rosetta",
"CRNN",
"STARNet",
"RARE",
"SRN",
"CLS",
"PGNet",
"Distillation",
"NRTR",
"TableAttn",
"SAR",
"PSE",
"SEED",
"SDMGR",
"LayoutXLM",
"LayoutLM",
"LayoutLMv2",
"PREN",
"FCE",
"SVTR",
"SVTR_LCNet",
"ViTSTR",
"ABINet",
"DB++",
"TableMaster",
"SPIN",
"VisionLAN",
"Gestalt",
"SLANet",
"RobustScanner",
"CT",
"RFL",
"DRRG",
"CAN",
"Telescope",
"SATRN",
"SVTR_HGNet",
"ParseQ",
"CPPD",
"LaTeXOCR",
"UniMERNet",
"SLANeXt",
"PP-FormulaNet-S",
"PP-FormulaNet-L",
"PP-FormulaNet_plus-S",
"PP-FormulaNet_plus-M",
"PP-FormulaNet_plus-L",
]
if use_xpu:
device = "xpu:{0}".format(os.getenv("FLAGS_selected_xpus", 0))
elif use_npu:
device = "npu:{0}".format(os.getenv("FLAGS_selected_npus", 0))
elif use_mlu:
device = "mlu:{0}".format(os.getenv("FLAGS_selected_mlus", 0))
elif use_gcu: # Use Enflame GCU(General Compute Unit)
device = "gcu:{0}".format(os.getenv("FLAGS_selected_gcus", 0))
else:
device = "gpu:{}".format(dist.ParallelEnv().dev_id) if use_gpu else "cpu"
check_device(use_gpu, use_xpu, use_npu, use_mlu, use_gcu)
device = paddle.set_device(device)
config["Global"]["distributed"] = dist.get_world_size() != 1
loggers = []
if "use_visualdl" in config["Global"] and config["Global"]["use_visualdl"]:
logger.warning(
"You are using VisualDL, the VisualDL is deprecated and "
"removed in ppocr!"
)
log_writer = None
if (
"use_wandb" in config["Global"] and config["Global"]["use_wandb"]
) or "wandb" in config:
save_dir = config["Global"]["save_model_dir"]
wandb_writer_path = "{}/wandb".format(save_dir)
if "wandb" in config:
wandb_params = config["wandb"]
else:
wandb_params = dict()
wandb_params.update({"save_dir": save_dir})
log_writer = WandbLogger(**wandb_params, config=config)
loggers.append(log_writer)
else:
log_writer = None
print_dict(config, logger)
if loggers:
log_writer = Loggers(loggers)
else:
log_writer = None
logger.info("train with paddle {} and device {}".format(paddle.__version__, device))
return config, device, logger, log_writer

162
tools/test_hubserving.py Executable file
View File

@@ -0,0 +1,162 @@
# 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.
import os
import sys
__dir__ = os.path.dirname(os.path.abspath(__file__))
sys.path.append(__dir__)
sys.path.append(os.path.abspath(os.path.join(__dir__, "..")))
from ppocr.utils.logging import get_logger
logger = get_logger()
import cv2
import numpy as np
import time
from PIL import Image
from ppocr.utils.utility import get_image_file_list
from tools.infer.utility import draw_ocr, draw_boxes, str2bool
from ppstructure.utility import draw_structure_result
from ppstructure.predict_system import to_excel
import requests
import json
import base64
def cv2_to_base64(image):
return base64.b64encode(image).decode("utf8")
def draw_server_result(image_file, res):
img = cv2.imread(image_file)
image = Image.fromarray(cv2.cvtColor(img, cv2.COLOR_BGR2RGB))
if len(res) == 0:
return np.array(image)
keys = res[0].keys()
if "text_region" not in keys: # for ocr_rec, draw function is invalid
logger.info("draw function is invalid for ocr_rec!")
return None
elif "text" not in keys: # for ocr_det
logger.info("draw text boxes only!")
boxes = []
for dno in range(len(res)):
boxes.append(res[dno]["text_region"])
boxes = np.array(boxes)
draw_img = draw_boxes(image, boxes)
return draw_img
else: # for ocr_system
logger.info("draw boxes and texts!")
boxes = []
texts = []
scores = []
for dno in range(len(res)):
boxes.append(res[dno]["text_region"])
texts.append(res[dno]["text"])
scores.append(res[dno]["confidence"])
boxes = np.array(boxes)
scores = np.array(scores)
draw_img = draw_ocr(image, boxes, texts, scores, draw_txt=True, drop_score=0.5)
return draw_img
def save_structure_res(res, save_folder, image_file):
img = cv2.imread(image_file)
excel_save_folder = os.path.join(save_folder, os.path.basename(image_file))
os.makedirs(excel_save_folder, exist_ok=True)
# save res
with open(os.path.join(excel_save_folder, "res.txt"), "w", encoding="utf8") as f:
for region in res:
if region["type"] == "Table":
excel_path = os.path.join(
excel_save_folder, "{}.xlsx".format(region["bbox"])
)
to_excel(region["res"], excel_path)
elif region["type"] == "Figure":
x1, y1, x2, y2 = region["bbox"]
print(region["bbox"])
roi_img = img[y1:y2, x1:x2, :]
img_path = os.path.join(
excel_save_folder, "{}.jpg".format(region["bbox"])
)
cv2.imwrite(img_path, roi_img)
else:
for text_result in region["res"]:
f.write("{}\n".format(json.dumps(text_result)))
def main(args):
image_file_list = get_image_file_list(args.image_dir)
is_visualize = False
headers = {"Content-type": "application/json"}
cnt = 0
total_time = 0
for image_file in image_file_list:
img = open(image_file, "rb").read()
if img is None:
logger.info("error in loading image:{}".format(image_file))
continue
img_name = os.path.basename(image_file)
# seed http request
starttime = time.time()
data = {"images": [cv2_to_base64(img)]}
r = requests.post(url=args.server_url, headers=headers, data=json.dumps(data))
elapse = time.time() - starttime
total_time += elapse
logger.info("Predict time of %s: %.3fs" % (image_file, elapse))
res = r.json()["results"][0]
logger.info(res)
if args.visualize:
draw_img = None
if "structure_table" in args.server_url:
to_excel(res["html"], "./{}.xlsx".format(img_name))
elif "structure_system" in args.server_url:
save_structure_res(res["regions"], args.output, image_file)
else:
draw_img = draw_server_result(image_file, res)
if draw_img is not None:
if not os.path.exists(args.output):
os.makedirs(args.output)
cv2.imwrite(
os.path.join(args.output, os.path.basename(image_file)),
draw_img[:, :, ::-1],
)
logger.info(
"The visualized image saved in {}".format(
os.path.join(args.output, os.path.basename(image_file))
)
)
cnt += 1
if cnt % 100 == 0:
logger.info("{} processed".format(cnt))
logger.info("avg time cost: {}".format(float(total_time) / cnt))
def parse_args():
import argparse
parser = argparse.ArgumentParser(description="args for hub serving")
parser.add_argument("--server_url", type=str, required=True)
parser.add_argument("--image_dir", type=str, required=True)
parser.add_argument("--visualize", type=str2bool, default=False)
parser.add_argument("--output", type=str, default="./hubserving_result")
args = parser.parse_args()
return args
if __name__ == "__main__":
args = parse_args()
main(args)

273
tools/train.py Executable file
View File

@@ -0,0 +1,273 @@
# 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(os.path.abspath(__file__))
sys.path.append(__dir__)
sys.path.insert(0, os.path.abspath(os.path.join(__dir__, "..")))
import yaml
import paddle
import paddle.distributed as dist
from ppocr.data import build_dataloader, set_signal_handlers
from ppocr.modeling.architectures import build_model
from ppocr.losses import build_loss
from ppocr.optimizer import build_optimizer
from ppocr.postprocess import build_post_process
from ppocr.metrics import build_metric
from ppocr.utils.save_load import load_model
from ppocr.utils.utility import set_seed
from ppocr.modeling.architectures import apply_to_static
import tools.program as program
import tools.naive_sync_bn as naive_sync_bn
dist.get_world_size()
def main(config, device, logger, vdl_writer, seed):
# init dist environment
if config["Global"]["distributed"]:
dist.init_parallel_env()
global_config = config["Global"]
# build dataloader
set_signal_handlers()
train_dataloader = build_dataloader(config, "Train", device, logger, seed)
if len(train_dataloader) == 0:
logger.error(
"No Images in train dataset, please ensure\n"
+ "\t1. The images num in the train label_file_list should be larger than or equal with batch size.\n"
+ "\t2. The annotation file and path in the configuration file are provided normally."
)
return
if config["Eval"]:
valid_dataloader = build_dataloader(config, "Eval", device, logger, seed)
else:
valid_dataloader = None
step_pre_epoch = len(train_dataloader)
# build post process
post_process_class = build_post_process(config["PostProcess"], global_config)
# build model
# for rec algorithm
if hasattr(post_process_class, "character"):
char_num = len(getattr(post_process_class, "character"))
if config["Architecture"]["algorithm"] in [
"Distillation",
]: # distillation model
for key in config["Architecture"]["Models"]:
if (
config["Architecture"]["Models"][key]["Head"]["name"] == "MultiHead"
): # for multi head
if config["PostProcess"]["name"] == "DistillationSARLabelDecode":
char_num = char_num - 2
if config["PostProcess"]["name"] == "DistillationNRTRLabelDecode":
char_num = char_num - 3
out_channels_list = {}
out_channels_list["CTCLabelDecode"] = char_num
# update SARLoss params
if (
list(config["Loss"]["loss_config_list"][-1].keys())[0]
== "DistillationSARLoss"
):
config["Loss"]["loss_config_list"][-1]["DistillationSARLoss"][
"ignore_index"
] = (char_num + 1)
out_channels_list["SARLabelDecode"] = char_num + 2
elif any(
"DistillationNRTRLoss" in d
for d in config["Loss"]["loss_config_list"]
):
out_channels_list["NRTRLabelDecode"] = char_num + 3
config["Architecture"]["Models"][key]["Head"][
"out_channels_list"
] = out_channels_list
else:
config["Architecture"]["Models"][key]["Head"][
"out_channels"
] = char_num
elif config["Architecture"]["Head"]["name"] == "MultiHead": # for multi head
if config["PostProcess"]["name"] == "SARLabelDecode":
char_num = char_num - 2
if config["PostProcess"]["name"] == "NRTRLabelDecode":
char_num = char_num - 3
out_channels_list = {}
out_channels_list["CTCLabelDecode"] = char_num
# update SARLoss params
if list(config["Loss"]["loss_config_list"][1].keys())[0] == "SARLoss":
if config["Loss"]["loss_config_list"][1]["SARLoss"] is None:
config["Loss"]["loss_config_list"][1]["SARLoss"] = {
"ignore_index": char_num + 1
}
else:
config["Loss"]["loss_config_list"][1]["SARLoss"]["ignore_index"] = (
char_num + 1
)
out_channels_list["SARLabelDecode"] = char_num + 2
elif list(config["Loss"]["loss_config_list"][1].keys())[0] == "NRTRLoss":
out_channels_list["NRTRLabelDecode"] = char_num + 3
config["Architecture"]["Head"]["out_channels_list"] = out_channels_list
else: # base rec model
config["Architecture"]["Head"]["out_channels"] = char_num
if config["PostProcess"]["name"] == "SARLabelDecode": # for SAR model
config["Loss"]["ignore_index"] = char_num - 1
model = build_model(config["Architecture"])
use_sync_bn = config["Global"].get("use_sync_bn", False)
if use_sync_bn:
if config["Global"].get("use_npu", False) or config["Global"].get(
"use_xpu", False
):
naive_sync_bn.convert_syncbn(model)
else:
model = paddle.nn.SyncBatchNorm.convert_sync_batchnorm(model)
logger.info("convert_sync_batchnorm")
model = apply_to_static(model, config, logger)
# build loss
loss_class = build_loss(config["Loss"])
# build optim
optimizer, lr_scheduler = build_optimizer(
config["Optimizer"],
epochs=config["Global"]["epoch_num"],
step_each_epoch=len(train_dataloader),
model=model,
)
# build metric
eval_class = build_metric(config["Metric"])
logger.info("train dataloader has {} iters".format(len(train_dataloader)))
if valid_dataloader is not None:
logger.info("valid dataloader has {} iters".format(len(valid_dataloader)))
use_amp = config["Global"].get("use_amp", False)
amp_level = config["Global"].get("amp_level", "O2")
amp_dtype = config["Global"].get("amp_dtype", "float16")
amp_custom_black_list = config["Global"].get("amp_custom_black_list", [])
amp_custom_white_list = config["Global"].get("amp_custom_white_list", [])
if os.path.exists(
os.path.join(config["Global"]["save_model_dir"], "train_result.json")
):
try:
os.remove(
os.path.join(config["Global"]["save_model_dir"], "train_result.json")
)
except:
pass
if use_amp:
AMP_RELATED_FLAGS_SETTING = {}
if paddle.is_compiled_with_cuda():
AMP_RELATED_FLAGS_SETTING.update(
{
"FLAGS_cudnn_batchnorm_spatial_persistent": 1,
"FLAGS_gemm_use_half_precision_compute_type": 0,
}
)
paddle.set_flags(AMP_RELATED_FLAGS_SETTING)
scale_loss = config["Global"].get("scale_loss", 1.0)
use_dynamic_loss_scaling = config["Global"].get(
"use_dynamic_loss_scaling", False
)
scaler = paddle.amp.GradScaler(
init_loss_scaling=scale_loss,
use_dynamic_loss_scaling=use_dynamic_loss_scaling,
)
if amp_level == "O2":
model, optimizer = paddle.amp.decorate(
models=model,
optimizers=optimizer,
level=amp_level,
master_weight=True,
dtype=amp_dtype,
)
else:
scaler = None
# load pretrain model
pre_best_model_dict = load_model(
config, model, optimizer, config["Architecture"]["model_type"]
)
if config["Global"]["distributed"]:
find_unused_parameters = config["Global"].get("find_unused_parameters", False)
model = paddle.DataParallel(
model, find_unused_parameters=find_unused_parameters
)
# start train
program.train(
config,
train_dataloader,
valid_dataloader,
device,
model,
loss_class,
optimizer,
lr_scheduler,
post_process_class,
eval_class,
pre_best_model_dict,
logger,
step_pre_epoch,
vdl_writer,
scaler,
amp_level,
amp_custom_black_list,
amp_custom_white_list,
amp_dtype,
)
def test_reader(config, device, logger):
loader = build_dataloader(config, "Train", device, logger)
import time
starttime = time.time()
count = 0
try:
for data in loader():
count += 1
if count % 1 == 0:
batch_time = time.time() - starttime
starttime = time.time()
logger.info(
"reader: {}, {}, {}".format(count, len(data[0]), batch_time)
)
except Exception as e:
logger.info(e)
logger.info("finish reader: {}, Success!".format(count))
if __name__ == "__main__":
config, device, logger, vdl_writer = program.preprocess(is_train=True)
seed = config["Global"]["seed"] if "seed" in config["Global"] else 1024
set_seed(seed)
main(config, device, logger, vdl_writer, seed)
# test_reader(config, device, logger)