| import multiprocessing |
| import glob |
| import time |
| import json |
| from tqdm import tqdm |
| import argparse |
| import os |
| from os.path import join as pjoin, exists |
| import cv2 |
|
|
| import detect_compo.ip_region_proposal as ip |
|
|
| def resize_height_by_longest_edge(img_path, resize_length=800): |
| org = cv2.imread(img_path) |
| height, width = org.shape[:2] |
| if height > width: |
| return resize_length |
| else: |
| return int(resize_length * (height / width)) |
|
|
|
|
| def main(args): |
| |
| |
| |
| |
| input_img_root = args.input_img_root |
| output_root = args.output_root |
| with open(args.ann, 'r') as f: |
| data = json.load(f) |
|
|
| os.makedirs(output_root, exist_ok=True) |
| |
| |
| input_imgs = [pjoin(input_img_root, img['file_name']) for img in data['images']] |
|
|
| |
| |
| key_params = {'min-grad': 10, 'ffl-block': 5, 'min-ele-area': 50, 'merge-contained-ele': True, |
| 'max-word-inline-gap': 10, 'max-line-ingraph-gap': 4, 'remove-bar': False} |
| |
|
|
| is_ip = True |
| is_clf = False |
| is_ocr = True |
| is_merge = True |
|
|
| |
| compo_classifier = None |
| if is_ip and is_clf: |
| compo_classifier = {} |
| from cnn.CNN import CNN |
| |
| compo_classifier['Elements'] = CNN('Elements') |
| |
| ocr_model = None |
| if is_ocr: |
| import detect_text.text_detection as text |
|
|
| |
| num = 0 |
| |
| |
| start_time = time.time() |
| for input_img in tqdm(input_imgs, total=len(input_imgs)): |
| name = input_img.split('/')[-1].split('.')[-2] |
| resized_height = resize_height_by_longest_edge(input_img) |
| |
| |
| |
| |
| |
|
|
| if is_ocr: |
| text.text_detection(input_img, output_root, show=False, method='paddle') |
|
|
| if is_ip: |
| ip.compo_detection(input_img, output_root, key_params, classifier=compo_classifier, resize_by_height=resized_height, show=False) |
|
|
| if is_merge: |
| import detect_merge.merge as merge |
| |
| |
| compo_path = pjoin(output_root, 'ip', name + '.json') |
| ocr_path = pjoin(output_root, 'ocr', name + '.json') |
| merge.merge(input_img, compo_path, ocr_path, output_root, is_remove_bar=key_params['remove-bar'], show=False) |
|
|
| num += 1 |
| |
| end_time = time.time() |
| print(f'Total inference time: {end_time - start_time} s') |
| os.makedirs(pjoin(output_root, 'time'), exist_ok=True) |
| with open(pjoin(output_root, 'time', 'inference time.txt'), 'w') as f: |
| f.write(f'First stage total inference time: {end_time - start_time} s') |
|
|
|
|
| if __name__ == '__main__': |
| parser = argparse.ArgumentParser() |
| parser.add_argument('--input_img_root', type=str, required=True) |
| parser.add_argument('--output_root', type=str, required=True) |
| parser.add_argument('--ann', type=str, required=True) |
| args = parser.parse_args() |
| main(args) |
|
|