mirror of
https://github.com/opendatalab/MinerU.git
synced 2026-08-29 03:44:03 +08:00
refactor: reorganize project structure and update import paths
This commit is contained in:
@@ -0,0 +1 @@
|
||||
# Copyright (c) Opendatalab. All rights reserved.
|
||||
@@ -0,0 +1,214 @@
|
||||
import cv2
|
||||
from loguru import logger
|
||||
from tqdm import tqdm
|
||||
|
||||
from .model_init import AtomModelSingleton
|
||||
from ...utils.model_utils import crop_img, get_res_list_from_layout_res, get_coords_and_area
|
||||
from ...utils.ocr_utils import get_adjusted_mfdetrec_res, get_ocr_result_list
|
||||
|
||||
YOLO_LAYOUT_BASE_BATCH_SIZE = 1
|
||||
MFD_BASE_BATCH_SIZE = 1
|
||||
MFR_BASE_BATCH_SIZE = 16
|
||||
|
||||
|
||||
class BatchAnalyze:
|
||||
def __init__(self, model_manager, batch_ratio: int, formula_enable, table_enable):
|
||||
self.batch_ratio = batch_ratio
|
||||
self.formula_enable = formula_enable
|
||||
self.table_enable = table_enable
|
||||
self.model_manager = model_manager
|
||||
|
||||
def __call__(self, images_with_extra_info: list) -> list:
|
||||
if len(images_with_extra_info) == 0:
|
||||
return []
|
||||
|
||||
images_layout_res = []
|
||||
|
||||
self.model = self.model_manager.get_model(
|
||||
lang=None,
|
||||
formula_enable=self.formula_enable,
|
||||
table_enable=self.table_enable,
|
||||
)
|
||||
atom_model_manager = AtomModelSingleton()
|
||||
|
||||
images = [image for image, _, _ in images_with_extra_info]
|
||||
|
||||
# doclayout_yolo
|
||||
layout_images = []
|
||||
for image_index, image in enumerate(images):
|
||||
layout_images.append(image)
|
||||
|
||||
|
||||
images_layout_res += self.model.layout_model.batch_predict(
|
||||
layout_images, YOLO_LAYOUT_BASE_BATCH_SIZE
|
||||
)
|
||||
|
||||
if self.formula_enable:
|
||||
# 公式检测
|
||||
images_mfd_res = self.model.mfd_model.batch_predict(
|
||||
images, MFD_BASE_BATCH_SIZE
|
||||
)
|
||||
|
||||
# 公式识别
|
||||
images_formula_list = self.model.mfr_model.batch_predict(
|
||||
images_mfd_res,
|
||||
images,
|
||||
batch_size=self.batch_ratio * MFR_BASE_BATCH_SIZE,
|
||||
)
|
||||
mfr_count = 0
|
||||
for image_index in range(len(images)):
|
||||
images_layout_res[image_index] += images_formula_list[image_index]
|
||||
mfr_count += len(images_formula_list[image_index])
|
||||
|
||||
# 清理显存
|
||||
# clean_vram(self.model.device, vram_threshold=8)
|
||||
|
||||
ocr_res_list_all_page = []
|
||||
table_res_list_all_page = []
|
||||
for index in range(len(images)):
|
||||
_, ocr_enable, _lang = images_with_extra_info[index]
|
||||
layout_res = images_layout_res[index]
|
||||
np_array_img = images[index]
|
||||
|
||||
ocr_res_list, table_res_list, single_page_mfdetrec_res = (
|
||||
get_res_list_from_layout_res(layout_res)
|
||||
)
|
||||
|
||||
ocr_res_list_all_page.append({'ocr_res_list':ocr_res_list,
|
||||
'lang':_lang,
|
||||
'ocr_enable':ocr_enable,
|
||||
'np_array_img':np_array_img,
|
||||
'single_page_mfdetrec_res':single_page_mfdetrec_res,
|
||||
'layout_res':layout_res,
|
||||
})
|
||||
|
||||
for table_res in table_res_list:
|
||||
table_img, _ = crop_img(table_res, np_array_img)
|
||||
table_res_list_all_page.append({'table_res':table_res,
|
||||
'lang':_lang,
|
||||
'table_img':table_img,
|
||||
})
|
||||
|
||||
# 文本框检测
|
||||
|
||||
for ocr_res_list_dict in tqdm(ocr_res_list_all_page, desc="OCR-det Predict"):
|
||||
# Process each area that requires OCR processing
|
||||
_lang = ocr_res_list_dict['lang']
|
||||
# Get OCR results for this language's images
|
||||
ocr_model = atom_model_manager.get_atom_model(
|
||||
atom_model_name='ocr',
|
||||
det_db_box_thresh=0.3,
|
||||
lang=_lang
|
||||
)
|
||||
for res in ocr_res_list_dict['ocr_res_list']:
|
||||
new_image, useful_list = crop_img(
|
||||
res, ocr_res_list_dict['np_array_img'], crop_paste_x=50, crop_paste_y=50
|
||||
)
|
||||
adjusted_mfdetrec_res = get_adjusted_mfdetrec_res(
|
||||
ocr_res_list_dict['single_page_mfdetrec_res'], useful_list
|
||||
)
|
||||
|
||||
# OCR-det
|
||||
new_image = cv2.cvtColor(new_image, cv2.COLOR_RGB2BGR)
|
||||
ocr_res = ocr_model.ocr(
|
||||
new_image, mfd_res=adjusted_mfdetrec_res, rec=False
|
||||
)[0]
|
||||
|
||||
# Integration results
|
||||
if ocr_res:
|
||||
ocr_result_list = get_ocr_result_list(ocr_res, useful_list, ocr_res_list_dict['ocr_enable'], new_image, _lang)
|
||||
|
||||
if res["category_id"] == 3:
|
||||
# ocr_result_list中所有bbox的面积之和
|
||||
ocr_res_area = sum(get_coords_and_area(ocr_res_item)[4] for ocr_res_item in ocr_result_list if 'poly' in ocr_res_item)
|
||||
# 求ocr_res_area和res的面积的比值
|
||||
res_area = get_coords_and_area(res)[4]
|
||||
if res_area > 0:
|
||||
ratio = ocr_res_area / res_area
|
||||
if ratio > 0.25:
|
||||
res["category_id"] = 1
|
||||
else:
|
||||
continue
|
||||
|
||||
ocr_res_list_dict['layout_res'].extend(ocr_result_list)
|
||||
|
||||
# 表格识别 table recognition
|
||||
if self.table_enable:
|
||||
for table_res_dict in tqdm(table_res_list_all_page, desc="Table Predict"):
|
||||
_lang = table_res_dict['lang']
|
||||
table_model = atom_model_manager.get_atom_model(
|
||||
atom_model_name='table',
|
||||
device='cpu',
|
||||
lang=_lang,
|
||||
table_sub_model_name='slanet_plus'
|
||||
)
|
||||
html_code, table_cell_bboxes, logic_points, elapse = table_model.predict(table_res_dict['table_img'])
|
||||
# 判断是否返回正常
|
||||
if html_code:
|
||||
expected_ending = html_code.strip().endswith(
|
||||
'</html>'
|
||||
) or html_code.strip().endswith('</table>')
|
||||
if expected_ending:
|
||||
table_res_dict['table_res']['html'] = html_code
|
||||
else:
|
||||
logger.warning(
|
||||
'table recognition processing fails, not found expected HTML table end'
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
'table recognition processing fails, not get html return'
|
||||
)
|
||||
|
||||
# Create dictionaries to store items by language
|
||||
need_ocr_lists_by_lang = {} # Dict of lists for each language
|
||||
img_crop_lists_by_lang = {} # Dict of lists for each language
|
||||
|
||||
for layout_res in images_layout_res:
|
||||
for layout_res_item in layout_res:
|
||||
if layout_res_item['category_id'] in [15]:
|
||||
if 'np_img' in layout_res_item and 'lang' in layout_res_item:
|
||||
lang = layout_res_item['lang']
|
||||
|
||||
# Initialize lists for this language if not exist
|
||||
if lang not in need_ocr_lists_by_lang:
|
||||
need_ocr_lists_by_lang[lang] = []
|
||||
img_crop_lists_by_lang[lang] = []
|
||||
|
||||
# Add to the appropriate language-specific lists
|
||||
need_ocr_lists_by_lang[lang].append(layout_res_item)
|
||||
img_crop_lists_by_lang[lang].append(layout_res_item['np_img'])
|
||||
|
||||
# Remove the fields after adding to lists
|
||||
layout_res_item.pop('np_img')
|
||||
layout_res_item.pop('lang')
|
||||
|
||||
if len(img_crop_lists_by_lang) > 0:
|
||||
|
||||
# Process OCR by language
|
||||
total_processed = 0
|
||||
|
||||
# Process each language separately
|
||||
for lang, img_crop_list in img_crop_lists_by_lang.items():
|
||||
if len(img_crop_list) > 0:
|
||||
# Get OCR results for this language's images
|
||||
|
||||
ocr_model = atom_model_manager.get_atom_model(
|
||||
atom_model_name='ocr',
|
||||
det_db_box_thresh=0.3,
|
||||
lang=lang
|
||||
)
|
||||
ocr_res_list = ocr_model.ocr(img_crop_list, det=False, tqdm_enable=True)[0]
|
||||
|
||||
# Verify we have matching counts
|
||||
assert len(ocr_res_list) == len(
|
||||
need_ocr_lists_by_lang[lang]), f'ocr_res_list: {len(ocr_res_list)}, need_ocr_list: {len(need_ocr_lists_by_lang[lang])} for lang: {lang}'
|
||||
|
||||
# Process OCR results for this language
|
||||
for index, layout_res_item in enumerate(need_ocr_lists_by_lang[lang]):
|
||||
ocr_text, ocr_score = ocr_res_list[index]
|
||||
layout_res_item['text'] = ocr_text
|
||||
layout_res_item['score'] = float(f"{ocr_score:.3f}")
|
||||
|
||||
total_processed += len(img_crop_list)
|
||||
|
||||
return images_layout_res
|
||||
@@ -0,0 +1,235 @@
|
||||
import os
|
||||
import time
|
||||
import numpy as np
|
||||
import torch
|
||||
from mineru.backend.pipeline.model_init import MineruPipelineModel
|
||||
|
||||
os.environ['FLAGS_npu_jit_compile'] = '0' # 关闭paddle的jit编译
|
||||
os.environ['FLAGS_use_stride_kernel'] = '0'
|
||||
os.environ['PYTORCH_ENABLE_MPS_FALLBACK'] = '1' # 让mps可以fallback
|
||||
os.environ['NO_ALBUMENTATIONS_UPDATE'] = '1' # 禁止albumentations检查更新
|
||||
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from ...utils.model_utils import get_vram, clean_memory
|
||||
from magic_pdf.libs.config_reader import (get_device, get_formula_config,
|
||||
get_layout_config,
|
||||
get_local_models_dir,
|
||||
get_table_recog_config)
|
||||
|
||||
class ModelSingleton:
|
||||
_instance = None
|
||||
_models = {}
|
||||
|
||||
def __new__(cls, *args, **kwargs):
|
||||
if cls._instance is None:
|
||||
cls._instance = super().__new__(cls)
|
||||
return cls._instance
|
||||
|
||||
def get_model(
|
||||
self,
|
||||
lang=None,
|
||||
formula_enable=None,
|
||||
table_enable=None,
|
||||
):
|
||||
key = (lang, formula_enable, table_enable)
|
||||
if key not in self._models:
|
||||
self._models[key] = custom_model_init(
|
||||
lang=lang,
|
||||
formula_enable=formula_enable,
|
||||
table_enable=table_enable,
|
||||
)
|
||||
return self._models[key]
|
||||
|
||||
|
||||
def custom_model_init(
|
||||
lang=None,
|
||||
formula_enable=None,
|
||||
table_enable=None,
|
||||
):
|
||||
model_init_start = time.time()
|
||||
# 从配置文件读取model-dir和device
|
||||
local_models_dir = get_local_models_dir()
|
||||
device = get_device()
|
||||
|
||||
formula_config = get_formula_config()
|
||||
if formula_enable is not None:
|
||||
formula_config['enable'] = formula_enable
|
||||
|
||||
table_config = get_table_recog_config()
|
||||
if table_enable is not None:
|
||||
table_config['enable'] = table_enable
|
||||
|
||||
model_input = {
|
||||
'models_dir': local_models_dir,
|
||||
'device': device,
|
||||
'table_config': table_config,
|
||||
'formula_config': formula_config,
|
||||
'lang': lang,
|
||||
}
|
||||
|
||||
custom_model = MineruPipelineModel(**model_input)
|
||||
|
||||
model_init_cost = time.time() - model_init_start
|
||||
logger.info(f'model init cost: {model_init_cost}')
|
||||
|
||||
return custom_model
|
||||
|
||||
def doc_analyze(
|
||||
dataset: Dataset,
|
||||
ocr: bool = False,
|
||||
start_page_id=0,
|
||||
end_page_id=None,
|
||||
lang=None,
|
||||
formula_enable=None,
|
||||
table_enable=None,
|
||||
):
|
||||
end_page_id = (
|
||||
end_page_id
|
||||
if end_page_id is not None and end_page_id >= 0
|
||||
else len(dataset) - 1
|
||||
)
|
||||
|
||||
MIN_BATCH_INFERENCE_SIZE = int(os.environ.get('MINERU_MIN_BATCH_INFERENCE_SIZE', 100))
|
||||
images = []
|
||||
page_wh_list = []
|
||||
for index in range(len(dataset)):
|
||||
if start_page_id <= index <= end_page_id:
|
||||
page_data = dataset.get_page(index)
|
||||
img_dict = page_data.get_image()
|
||||
images.append(img_dict['img'])
|
||||
page_wh_list.append((img_dict['width'], img_dict['height']))
|
||||
|
||||
images_with_extra_info = [(images[index], ocr, dataset._lang) for index in range(len(images))]
|
||||
|
||||
if len(images) >= MIN_BATCH_INFERENCE_SIZE:
|
||||
batch_size = MIN_BATCH_INFERENCE_SIZE
|
||||
batch_images = [images_with_extra_info[i:i+batch_size] for i in range(0, len(images_with_extra_info), batch_size)]
|
||||
else:
|
||||
batch_images = [images_with_extra_info]
|
||||
|
||||
results = []
|
||||
processed_images_count = 0
|
||||
for index, batch_image in enumerate(batch_images):
|
||||
processed_images_count += len(batch_image)
|
||||
logger.info(f'Batch {index + 1}/{len(batch_images)}: {processed_images_count} pages/{len(images_with_extra_info)} pages')
|
||||
result = may_batch_image_analyze(batch_image, formula_enable, table_enable)
|
||||
results.extend(result)
|
||||
|
||||
model_json = []
|
||||
for index in range(len(dataset)):
|
||||
if start_page_id <= index <= end_page_id:
|
||||
result = results.pop(0)
|
||||
page_width, page_height = page_wh_list.pop(0)
|
||||
else:
|
||||
result = []
|
||||
page_height = 0
|
||||
page_width = 0
|
||||
|
||||
page_info = {'page_no': index, 'width': page_width, 'height': page_height}
|
||||
page_dict = {'layout_dets': result, 'page_info': page_info}
|
||||
model_json.append(page_dict)
|
||||
|
||||
return model_json
|
||||
|
||||
def batch_doc_analyze(
|
||||
datasets: list[Dataset],
|
||||
parse_method: str = 'auto',
|
||||
lang=None,
|
||||
formula_enable=None,
|
||||
table_enable=None,
|
||||
):
|
||||
MIN_BATCH_INFERENCE_SIZE = int(os.environ.get('MINERU_MIN_BATCH_INFERENCE_SIZE', 100))
|
||||
batch_size = MIN_BATCH_INFERENCE_SIZE
|
||||
page_wh_list = []
|
||||
|
||||
images_with_extra_info = []
|
||||
for dataset in datasets:
|
||||
|
||||
ocr = False
|
||||
if parse_method == 'auto':
|
||||
if dataset.classify() == 'txt':
|
||||
ocr = False
|
||||
elif dataset.classify() == 'ocr':
|
||||
ocr = True
|
||||
elif parse_method == 'ocr':
|
||||
ocr = True
|
||||
elif parse_method == 'txt':
|
||||
ocr = False
|
||||
|
||||
_lang = dataset._lang
|
||||
|
||||
for index in range(len(dataset)):
|
||||
page_data = dataset.get_page(index)
|
||||
img_dict = page_data.get_image()
|
||||
page_wh_list.append((img_dict['width'], img_dict['height']))
|
||||
images_with_extra_info.append((img_dict['img'], ocr, _lang))
|
||||
|
||||
batch_images = [images_with_extra_info[i:i+batch_size] for i in range(0, len(images_with_extra_info), batch_size)]
|
||||
results = []
|
||||
processed_images_count = 0
|
||||
for index, batch_image in enumerate(batch_images):
|
||||
processed_images_count += len(batch_image)
|
||||
logger.info(f'Batch {index + 1}/{len(batch_images)}: {processed_images_count} pages/{len(images_with_extra_info)} pages')
|
||||
result = may_batch_image_analyze(batch_image, formula_enable, table_enable)
|
||||
results.extend(result)
|
||||
|
||||
infer_results = []
|
||||
for index in range(len(datasets)):
|
||||
dataset = datasets[index]
|
||||
model_json = []
|
||||
for i in range(len(dataset)):
|
||||
result = results.pop(0)
|
||||
page_width, page_height = page_wh_list.pop(0)
|
||||
page_info = {'page_no': i, 'width': page_width, 'height': page_height}
|
||||
page_dict = {'layout_dets': result, 'page_info': page_info}
|
||||
model_json.append(page_dict)
|
||||
infer_results.append(model_json)
|
||||
return infer_results
|
||||
|
||||
|
||||
def may_batch_image_analyze(
|
||||
images_with_extra_info: list[(np.ndarray, bool, str)],
|
||||
formula_enable=None,
|
||||
table_enable=None):
|
||||
# os.environ['CUDA_VISIBLE_DEVICES'] = str(idx)
|
||||
|
||||
from .batch_analyze import BatchAnalyze
|
||||
|
||||
model_manager = ModelSingleton()
|
||||
|
||||
batch_ratio = 1
|
||||
device = get_device()
|
||||
|
||||
if str(device).startswith('npu'):
|
||||
import torch_npu
|
||||
if torch_npu.npu.is_available():
|
||||
torch.npu.set_compile_mode(jit_compile=False)
|
||||
|
||||
if str(device).startswith('npu') or str(device).startswith('cuda'):
|
||||
vram = get_vram(device)
|
||||
if vram is not None:
|
||||
gpu_memory = int(os.getenv('VIRTUAL_VRAM_SIZE', round(vram)))
|
||||
if gpu_memory >= 16:
|
||||
batch_ratio = 16
|
||||
elif gpu_memory >= 12:
|
||||
batch_ratio = 8
|
||||
elif gpu_memory >= 8:
|
||||
batch_ratio = 4
|
||||
elif gpu_memory >= 6:
|
||||
batch_ratio = 2
|
||||
else:
|
||||
batch_ratio = 1
|
||||
logger.info(f'gpu_memory: {gpu_memory} GB, batch_ratio: {batch_ratio}')
|
||||
else:
|
||||
# Default batch_ratio when VRAM can't be determined
|
||||
batch_ratio = 1
|
||||
logger.info(f'Could not determine GPU memory, using default batch_ratio: {batch_ratio}')
|
||||
|
||||
batch_model = BatchAnalyze(model_manager, batch_ratio, formula_enable, table_enable)
|
||||
results = batch_model(images_with_extra_info)
|
||||
|
||||
clean_memory(get_device())
|
||||
|
||||
return results
|
||||
@@ -0,0 +1,771 @@
|
||||
import enum
|
||||
|
||||
from magic_pdf.config.model_block_type import ModelBlockTypeEnum
|
||||
from magic_pdf.config.ocr_content_type import CategoryId, ContentType
|
||||
from magic_pdf.data.dataset import Dataset
|
||||
from magic_pdf.libs.boxbase import (_is_in, bbox_distance, bbox_relative_pos,
|
||||
calculate_iou)
|
||||
from magic_pdf.libs.coordinate_transform import get_scale_ratio
|
||||
from magic_pdf.pre_proc.remove_bbox_overlap import _remove_overlap_between_bbox
|
||||
|
||||
CAPATION_OVERLAP_AREA_RATIO = 0.6
|
||||
MERGE_BOX_OVERLAP_AREA_RATIO = 1.1
|
||||
|
||||
|
||||
class PosRelationEnum(enum.Enum):
|
||||
LEFT = 'left'
|
||||
RIGHT = 'right'
|
||||
UP = 'up'
|
||||
BOTTOM = 'bottom'
|
||||
ALL = 'all'
|
||||
|
||||
|
||||
class MagicModel:
|
||||
"""每个函数没有得到元素的时候返回空list."""
|
||||
|
||||
def __fix_axis(self):
|
||||
for model_page_info in self.__model_list:
|
||||
need_remove_list = []
|
||||
page_no = model_page_info['page_info']['page_no']
|
||||
horizontal_scale_ratio, vertical_scale_ratio = get_scale_ratio(
|
||||
model_page_info, self.__docs.get_page(page_no)
|
||||
)
|
||||
layout_dets = model_page_info['layout_dets']
|
||||
for layout_det in layout_dets:
|
||||
|
||||
if layout_det.get('bbox') is not None:
|
||||
# 兼容直接输出bbox的模型数据,如paddle
|
||||
x0, y0, x1, y1 = layout_det['bbox']
|
||||
else:
|
||||
# 兼容直接输出poly的模型数据,如xxx
|
||||
x0, y0, _, _, x1, y1, _, _ = layout_det['poly']
|
||||
|
||||
bbox = [
|
||||
int(x0 / horizontal_scale_ratio),
|
||||
int(y0 / vertical_scale_ratio),
|
||||
int(x1 / horizontal_scale_ratio),
|
||||
int(y1 / vertical_scale_ratio),
|
||||
]
|
||||
layout_det['bbox'] = bbox
|
||||
# 删除高度或者宽度小于等于0的spans
|
||||
if bbox[2] - bbox[0] <= 0 or bbox[3] - bbox[1] <= 0:
|
||||
need_remove_list.append(layout_det)
|
||||
for need_remove in need_remove_list:
|
||||
layout_dets.remove(need_remove)
|
||||
|
||||
def __fix_by_remove_low_confidence(self):
|
||||
for model_page_info in self.__model_list:
|
||||
need_remove_list = []
|
||||
layout_dets = model_page_info['layout_dets']
|
||||
for layout_det in layout_dets:
|
||||
if layout_det['score'] <= 0.05:
|
||||
need_remove_list.append(layout_det)
|
||||
else:
|
||||
continue
|
||||
for need_remove in need_remove_list:
|
||||
layout_dets.remove(need_remove)
|
||||
|
||||
def __fix_by_remove_high_iou_and_low_confidence(self):
|
||||
for model_page_info in self.__model_list:
|
||||
need_remove_list = []
|
||||
layout_dets = model_page_info['layout_dets']
|
||||
for layout_det1 in layout_dets:
|
||||
for layout_det2 in layout_dets:
|
||||
if layout_det1 == layout_det2:
|
||||
continue
|
||||
if layout_det1['category_id'] in [
|
||||
0,
|
||||
1,
|
||||
2,
|
||||
3,
|
||||
4,
|
||||
5,
|
||||
6,
|
||||
7,
|
||||
8,
|
||||
9,
|
||||
] and layout_det2['category_id'] in [0, 1, 2, 3, 4, 5, 6, 7, 8, 9]:
|
||||
if (
|
||||
calculate_iou(layout_det1['bbox'], layout_det2['bbox'])
|
||||
> 0.9
|
||||
):
|
||||
if layout_det1['score'] < layout_det2['score']:
|
||||
layout_det_need_remove = layout_det1
|
||||
else:
|
||||
layout_det_need_remove = layout_det2
|
||||
|
||||
if layout_det_need_remove not in need_remove_list:
|
||||
need_remove_list.append(layout_det_need_remove)
|
||||
else:
|
||||
continue
|
||||
else:
|
||||
continue
|
||||
for need_remove in need_remove_list:
|
||||
layout_dets.remove(need_remove)
|
||||
|
||||
def __init__(self, model_list: list, docs: Dataset):
|
||||
self.__model_list = model_list
|
||||
self.__docs = docs
|
||||
"""为所有模型数据添加bbox信息(缩放,poly->bbox)"""
|
||||
self.__fix_axis()
|
||||
"""删除置信度特别低的模型数据(<0.05),提高质量"""
|
||||
self.__fix_by_remove_low_confidence()
|
||||
"""删除高iou(>0.9)数据中置信度较低的那个"""
|
||||
self.__fix_by_remove_high_iou_and_low_confidence()
|
||||
self.__fix_footnote()
|
||||
|
||||
def _bbox_distance(self, bbox1, bbox2):
|
||||
left, right, bottom, top = bbox_relative_pos(bbox1, bbox2)
|
||||
flags = [left, right, bottom, top]
|
||||
count = sum([1 if v else 0 for v in flags])
|
||||
if count > 1:
|
||||
return float('inf')
|
||||
if left or right:
|
||||
l1 = bbox1[3] - bbox1[1]
|
||||
l2 = bbox2[3] - bbox2[1]
|
||||
else:
|
||||
l1 = bbox1[2] - bbox1[0]
|
||||
l2 = bbox2[2] - bbox2[0]
|
||||
|
||||
if l2 > l1 and (l2 - l1) / l1 > 0.3:
|
||||
return float('inf')
|
||||
|
||||
return bbox_distance(bbox1, bbox2)
|
||||
|
||||
def __fix_footnote(self):
|
||||
# 3: figure, 5: table, 7: footnote
|
||||
for model_page_info in self.__model_list:
|
||||
footnotes = []
|
||||
figures = []
|
||||
tables = []
|
||||
|
||||
for obj in model_page_info['layout_dets']:
|
||||
if obj['category_id'] == 7:
|
||||
footnotes.append(obj)
|
||||
elif obj['category_id'] == 3:
|
||||
figures.append(obj)
|
||||
elif obj['category_id'] == 5:
|
||||
tables.append(obj)
|
||||
if len(footnotes) * len(figures) == 0:
|
||||
continue
|
||||
dis_figure_footnote = {}
|
||||
dis_table_footnote = {}
|
||||
|
||||
for i in range(len(footnotes)):
|
||||
for j in range(len(figures)):
|
||||
pos_flag_count = sum(
|
||||
list(
|
||||
map(
|
||||
lambda x: 1 if x else 0,
|
||||
bbox_relative_pos(
|
||||
footnotes[i]['bbox'], figures[j]['bbox']
|
||||
),
|
||||
)
|
||||
)
|
||||
)
|
||||
if pos_flag_count > 1:
|
||||
continue
|
||||
dis_figure_footnote[i] = min(
|
||||
self._bbox_distance(figures[j]['bbox'], footnotes[i]['bbox']),
|
||||
dis_figure_footnote.get(i, float('inf')),
|
||||
)
|
||||
for i in range(len(footnotes)):
|
||||
for j in range(len(tables)):
|
||||
pos_flag_count = sum(
|
||||
list(
|
||||
map(
|
||||
lambda x: 1 if x else 0,
|
||||
bbox_relative_pos(
|
||||
footnotes[i]['bbox'], tables[j]['bbox']
|
||||
),
|
||||
)
|
||||
)
|
||||
)
|
||||
if pos_flag_count > 1:
|
||||
continue
|
||||
|
||||
dis_table_footnote[i] = min(
|
||||
self._bbox_distance(tables[j]['bbox'], footnotes[i]['bbox']),
|
||||
dis_table_footnote.get(i, float('inf')),
|
||||
)
|
||||
for i in range(len(footnotes)):
|
||||
if i not in dis_figure_footnote:
|
||||
continue
|
||||
if dis_table_footnote.get(i, float('inf')) > dis_figure_footnote[i]:
|
||||
footnotes[i]['category_id'] = CategoryId.ImageFootnote
|
||||
|
||||
def __reduct_overlap(self, bboxes):
|
||||
N = len(bboxes)
|
||||
keep = [True] * N
|
||||
for i in range(N):
|
||||
for j in range(N):
|
||||
if i == j:
|
||||
continue
|
||||
if _is_in(bboxes[i]['bbox'], bboxes[j]['bbox']):
|
||||
keep[i] = False
|
||||
return [bboxes[i] for i in range(N) if keep[i]]
|
||||
|
||||
def __tie_up_category_by_distance_v2(
|
||||
self,
|
||||
page_no: int,
|
||||
subject_category_id: int,
|
||||
object_category_id: int,
|
||||
priority_pos: PosRelationEnum,
|
||||
):
|
||||
"""_summary_
|
||||
|
||||
Args:
|
||||
page_no (int): _description_
|
||||
subject_category_id (int): _description_
|
||||
object_category_id (int): _description_
|
||||
priority_pos (PosRelationEnum): _description_
|
||||
|
||||
Returns:
|
||||
_type_: _description_
|
||||
"""
|
||||
AXIS_MULPLICITY = 0.5
|
||||
subjects = self.__reduct_overlap(
|
||||
list(
|
||||
map(
|
||||
lambda x: {'bbox': x['bbox'], 'score': x['score']},
|
||||
filter(
|
||||
lambda x: x['category_id'] == subject_category_id,
|
||||
self.__model_list[page_no]['layout_dets'],
|
||||
),
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
objects = self.__reduct_overlap(
|
||||
list(
|
||||
map(
|
||||
lambda x: {'bbox': x['bbox'], 'score': x['score']},
|
||||
filter(
|
||||
lambda x: x['category_id'] == object_category_id,
|
||||
self.__model_list[page_no]['layout_dets'],
|
||||
),
|
||||
)
|
||||
)
|
||||
)
|
||||
M = len(objects)
|
||||
|
||||
subjects.sort(key=lambda x: x['bbox'][0] ** 2 + x['bbox'][1] ** 2)
|
||||
objects.sort(key=lambda x: x['bbox'][0] ** 2 + x['bbox'][1] ** 2)
|
||||
|
||||
sub_obj_map_h = {i: [] for i in range(len(subjects))}
|
||||
|
||||
dis_by_directions = {
|
||||
'top': [[-1, float('inf')]] * M,
|
||||
'bottom': [[-1, float('inf')]] * M,
|
||||
'left': [[-1, float('inf')]] * M,
|
||||
'right': [[-1, float('inf')]] * M,
|
||||
}
|
||||
|
||||
for i, obj in enumerate(objects):
|
||||
l_x_axis, l_y_axis = (
|
||||
obj['bbox'][2] - obj['bbox'][0],
|
||||
obj['bbox'][3] - obj['bbox'][1],
|
||||
)
|
||||
axis_unit = min(l_x_axis, l_y_axis)
|
||||
for j, sub in enumerate(subjects):
|
||||
|
||||
bbox1, bbox2, _ = _remove_overlap_between_bbox(
|
||||
objects[i]['bbox'], subjects[j]['bbox']
|
||||
)
|
||||
left, right, bottom, top = bbox_relative_pos(bbox1, bbox2)
|
||||
flags = [left, right, bottom, top]
|
||||
if sum([1 if v else 0 for v in flags]) > 1:
|
||||
continue
|
||||
|
||||
if left:
|
||||
if dis_by_directions['left'][i][1] > bbox_distance(
|
||||
obj['bbox'], sub['bbox']
|
||||
):
|
||||
dis_by_directions['left'][i] = [
|
||||
j,
|
||||
bbox_distance(obj['bbox'], sub['bbox']),
|
||||
]
|
||||
if right:
|
||||
if dis_by_directions['right'][i][1] > bbox_distance(
|
||||
obj['bbox'], sub['bbox']
|
||||
):
|
||||
dis_by_directions['right'][i] = [
|
||||
j,
|
||||
bbox_distance(obj['bbox'], sub['bbox']),
|
||||
]
|
||||
if bottom:
|
||||
if dis_by_directions['bottom'][i][1] > bbox_distance(
|
||||
obj['bbox'], sub['bbox']
|
||||
):
|
||||
dis_by_directions['bottom'][i] = [
|
||||
j,
|
||||
bbox_distance(obj['bbox'], sub['bbox']),
|
||||
]
|
||||
if top:
|
||||
if dis_by_directions['top'][i][1] > bbox_distance(
|
||||
obj['bbox'], sub['bbox']
|
||||
):
|
||||
dis_by_directions['top'][i] = [
|
||||
j,
|
||||
bbox_distance(obj['bbox'], sub['bbox']),
|
||||
]
|
||||
|
||||
if (
|
||||
dis_by_directions['top'][i][1] != float('inf')
|
||||
and dis_by_directions['bottom'][i][1] != float('inf')
|
||||
and priority_pos in (PosRelationEnum.BOTTOM, PosRelationEnum.UP)
|
||||
):
|
||||
RATIO = 3
|
||||
if (
|
||||
abs(
|
||||
dis_by_directions['top'][i][1]
|
||||
- dis_by_directions['bottom'][i][1]
|
||||
)
|
||||
< RATIO * axis_unit
|
||||
):
|
||||
|
||||
if priority_pos == PosRelationEnum.BOTTOM:
|
||||
sub_obj_map_h[dis_by_directions['bottom'][i][0]].append(i)
|
||||
else:
|
||||
sub_obj_map_h[dis_by_directions['top'][i][0]].append(i)
|
||||
continue
|
||||
|
||||
if dis_by_directions['left'][i][1] != float('inf') or dis_by_directions[
|
||||
'right'
|
||||
][i][1] != float('inf'):
|
||||
if dis_by_directions['left'][i][1] != float(
|
||||
'inf'
|
||||
) and dis_by_directions['right'][i][1] != float('inf'):
|
||||
if AXIS_MULPLICITY * axis_unit >= abs(
|
||||
dis_by_directions['left'][i][1]
|
||||
- dis_by_directions['right'][i][1]
|
||||
):
|
||||
left_sub_bbox = subjects[dis_by_directions['left'][i][0]][
|
||||
'bbox'
|
||||
]
|
||||
right_sub_bbox = subjects[dis_by_directions['right'][i][0]][
|
||||
'bbox'
|
||||
]
|
||||
|
||||
left_sub_bbox_y_axis = left_sub_bbox[3] - left_sub_bbox[1]
|
||||
right_sub_bbox_y_axis = right_sub_bbox[3] - right_sub_bbox[1]
|
||||
|
||||
if (
|
||||
abs(left_sub_bbox_y_axis - l_y_axis)
|
||||
+ dis_by_directions['left'][i][0]
|
||||
> abs(right_sub_bbox_y_axis - l_y_axis)
|
||||
+ dis_by_directions['right'][i][0]
|
||||
):
|
||||
left_or_right = dis_by_directions['right'][i]
|
||||
else:
|
||||
left_or_right = dis_by_directions['left'][i]
|
||||
else:
|
||||
left_or_right = dis_by_directions['left'][i]
|
||||
if left_or_right[1] > dis_by_directions['right'][i][1]:
|
||||
left_or_right = dis_by_directions['right'][i]
|
||||
else:
|
||||
left_or_right = dis_by_directions['left'][i]
|
||||
if left_or_right[1] == float('inf'):
|
||||
left_or_right = dis_by_directions['right'][i]
|
||||
else:
|
||||
left_or_right = [-1, float('inf')]
|
||||
|
||||
if dis_by_directions['top'][i][1] != float('inf') or dis_by_directions[
|
||||
'bottom'
|
||||
][i][1] != float('inf'):
|
||||
if dis_by_directions['top'][i][1] != float('inf') and dis_by_directions[
|
||||
'bottom'
|
||||
][i][1] != float('inf'):
|
||||
if AXIS_MULPLICITY * axis_unit >= abs(
|
||||
dis_by_directions['top'][i][1]
|
||||
- dis_by_directions['bottom'][i][1]
|
||||
):
|
||||
top_bottom = subjects[dis_by_directions['bottom'][i][0]]['bbox']
|
||||
bottom_top = subjects[dis_by_directions['top'][i][0]]['bbox']
|
||||
|
||||
top_bottom_x_axis = top_bottom[2] - top_bottom[0]
|
||||
bottom_top_x_axis = bottom_top[2] - bottom_top[0]
|
||||
if (
|
||||
abs(top_bottom_x_axis - l_x_axis)
|
||||
+ dis_by_directions['bottom'][i][1]
|
||||
> abs(bottom_top_x_axis - l_x_axis)
|
||||
+ dis_by_directions['top'][i][1]
|
||||
):
|
||||
top_or_bottom = dis_by_directions['top'][i]
|
||||
else:
|
||||
top_or_bottom = dis_by_directions['bottom'][i]
|
||||
else:
|
||||
top_or_bottom = dis_by_directions['top'][i]
|
||||
if top_or_bottom[1] > dis_by_directions['bottom'][i][1]:
|
||||
top_or_bottom = dis_by_directions['bottom'][i]
|
||||
else:
|
||||
top_or_bottom = dis_by_directions['top'][i]
|
||||
if top_or_bottom[1] == float('inf'):
|
||||
top_or_bottom = dis_by_directions['bottom'][i]
|
||||
else:
|
||||
top_or_bottom = [-1, float('inf')]
|
||||
|
||||
if left_or_right[1] != float('inf') or top_or_bottom[1] != float('inf'):
|
||||
if left_or_right[1] != float('inf') and top_or_bottom[1] != float(
|
||||
'inf'
|
||||
):
|
||||
if AXIS_MULPLICITY * axis_unit >= abs(
|
||||
left_or_right[1] - top_or_bottom[1]
|
||||
):
|
||||
y_axis_bbox = subjects[left_or_right[0]]['bbox']
|
||||
x_axis_bbox = subjects[top_or_bottom[0]]['bbox']
|
||||
|
||||
if (
|
||||
abs((x_axis_bbox[2] - x_axis_bbox[0]) - l_x_axis) / l_x_axis
|
||||
> abs((y_axis_bbox[3] - y_axis_bbox[1]) - l_y_axis)
|
||||
/ l_y_axis
|
||||
):
|
||||
sub_obj_map_h[left_or_right[0]].append(i)
|
||||
else:
|
||||
sub_obj_map_h[top_or_bottom[0]].append(i)
|
||||
else:
|
||||
if left_or_right[1] > top_or_bottom[1]:
|
||||
sub_obj_map_h[top_or_bottom[0]].append(i)
|
||||
else:
|
||||
sub_obj_map_h[left_or_right[0]].append(i)
|
||||
else:
|
||||
if left_or_right[1] != float('inf'):
|
||||
sub_obj_map_h[left_or_right[0]].append(i)
|
||||
else:
|
||||
sub_obj_map_h[top_or_bottom[0]].append(i)
|
||||
ret = []
|
||||
for i in sub_obj_map_h.keys():
|
||||
ret.append(
|
||||
{
|
||||
'sub_bbox': {
|
||||
'bbox': subjects[i]['bbox'],
|
||||
'score': subjects[i]['score'],
|
||||
},
|
||||
'obj_bboxes': [
|
||||
{'score': objects[j]['score'], 'bbox': objects[j]['bbox']}
|
||||
for j in sub_obj_map_h[i]
|
||||
],
|
||||
'sub_idx': i,
|
||||
}
|
||||
)
|
||||
return ret
|
||||
|
||||
|
||||
def __tie_up_category_by_distance_v3(
|
||||
self,
|
||||
page_no: int,
|
||||
subject_category_id: int,
|
||||
object_category_id: int,
|
||||
priority_pos: PosRelationEnum,
|
||||
):
|
||||
subjects = self.__reduct_overlap(
|
||||
list(
|
||||
map(
|
||||
lambda x: {'bbox': x['bbox'], 'score': x['score']},
|
||||
filter(
|
||||
lambda x: x['category_id'] == subject_category_id,
|
||||
self.__model_list[page_no]['layout_dets'],
|
||||
),
|
||||
)
|
||||
)
|
||||
)
|
||||
objects = self.__reduct_overlap(
|
||||
list(
|
||||
map(
|
||||
lambda x: {'bbox': x['bbox'], 'score': x['score']},
|
||||
filter(
|
||||
lambda x: x['category_id'] == object_category_id,
|
||||
self.__model_list[page_no]['layout_dets'],
|
||||
),
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
ret = []
|
||||
N, M = len(subjects), len(objects)
|
||||
subjects.sort(key=lambda x: x['bbox'][0] ** 2 + x['bbox'][1] ** 2)
|
||||
objects.sort(key=lambda x: x['bbox'][0] ** 2 + x['bbox'][1] ** 2)
|
||||
|
||||
OBJ_IDX_OFFSET = 10000
|
||||
SUB_BIT_KIND, OBJ_BIT_KIND = 0, 1
|
||||
|
||||
all_boxes_with_idx = [(i, SUB_BIT_KIND, sub['bbox'][0], sub['bbox'][1]) for i, sub in enumerate(subjects)] + [(i + OBJ_IDX_OFFSET , OBJ_BIT_KIND, obj['bbox'][0], obj['bbox'][1]) for i, obj in enumerate(objects)]
|
||||
seen_idx = set()
|
||||
seen_sub_idx = set()
|
||||
|
||||
while N > len(seen_sub_idx):
|
||||
candidates = []
|
||||
for idx, kind, x0, y0 in all_boxes_with_idx:
|
||||
if idx in seen_idx:
|
||||
continue
|
||||
candidates.append((idx, kind, x0, y0))
|
||||
|
||||
if len(candidates) == 0:
|
||||
break
|
||||
left_x = min([v[2] for v in candidates])
|
||||
top_y = min([v[3] for v in candidates])
|
||||
|
||||
candidates.sort(key=lambda x: (x[2]-left_x) ** 2 + (x[3] - top_y) ** 2)
|
||||
|
||||
|
||||
fst_idx, fst_kind, left_x, top_y = candidates[0]
|
||||
candidates.sort(key=lambda x: (x[2] - left_x) ** 2 + (x[3] - top_y)**2)
|
||||
nxt = None
|
||||
|
||||
for i in range(1, len(candidates)):
|
||||
if candidates[i][1] ^ fst_kind == 1:
|
||||
nxt = candidates[i]
|
||||
break
|
||||
if nxt is None:
|
||||
break
|
||||
|
||||
if fst_kind == SUB_BIT_KIND:
|
||||
sub_idx, obj_idx = fst_idx, nxt[0] - OBJ_IDX_OFFSET
|
||||
|
||||
else:
|
||||
sub_idx, obj_idx = nxt[0], fst_idx - OBJ_IDX_OFFSET
|
||||
|
||||
pair_dis = bbox_distance(subjects[sub_idx]['bbox'], objects[obj_idx]['bbox'])
|
||||
nearest_dis = float('inf')
|
||||
for i in range(N):
|
||||
if i in seen_idx or i == sub_idx:continue
|
||||
nearest_dis = min(nearest_dis, bbox_distance(subjects[i]['bbox'], objects[obj_idx]['bbox']))
|
||||
|
||||
if pair_dis >= 3*nearest_dis:
|
||||
seen_idx.add(sub_idx)
|
||||
continue
|
||||
|
||||
seen_idx.add(sub_idx)
|
||||
seen_idx.add(obj_idx + OBJ_IDX_OFFSET)
|
||||
seen_sub_idx.add(sub_idx)
|
||||
|
||||
ret.append(
|
||||
{
|
||||
'sub_bbox': {
|
||||
'bbox': subjects[sub_idx]['bbox'],
|
||||
'score': subjects[sub_idx]['score'],
|
||||
},
|
||||
'obj_bboxes': [
|
||||
{'score': objects[obj_idx]['score'], 'bbox': objects[obj_idx]['bbox']}
|
||||
],
|
||||
'sub_idx': sub_idx,
|
||||
}
|
||||
)
|
||||
|
||||
for i in range(len(objects)):
|
||||
j = i + OBJ_IDX_OFFSET
|
||||
if j in seen_idx:
|
||||
continue
|
||||
seen_idx.add(j)
|
||||
nearest_dis, nearest_sub_idx = float('inf'), -1
|
||||
for k in range(len(subjects)):
|
||||
dis = bbox_distance(objects[i]['bbox'], subjects[k]['bbox'])
|
||||
if dis < nearest_dis:
|
||||
nearest_dis = dis
|
||||
nearest_sub_idx = k
|
||||
|
||||
for k in range(len(subjects)):
|
||||
if k != nearest_sub_idx: continue
|
||||
if k in seen_sub_idx:
|
||||
for kk in range(len(ret)):
|
||||
if ret[kk]['sub_idx'] == k:
|
||||
ret[kk]['obj_bboxes'].append({'score': objects[i]['score'], 'bbox': objects[i]['bbox']})
|
||||
break
|
||||
else:
|
||||
ret.append(
|
||||
{
|
||||
'sub_bbox': {
|
||||
'bbox': subjects[k]['bbox'],
|
||||
'score': subjects[k]['score'],
|
||||
},
|
||||
'obj_bboxes': [
|
||||
{'score': objects[i]['score'], 'bbox': objects[i]['bbox']}
|
||||
],
|
||||
'sub_idx': k,
|
||||
}
|
||||
)
|
||||
seen_sub_idx.add(k)
|
||||
seen_idx.add(k)
|
||||
|
||||
|
||||
for i in range(len(subjects)):
|
||||
if i in seen_sub_idx:
|
||||
continue
|
||||
ret.append(
|
||||
{
|
||||
'sub_bbox': {
|
||||
'bbox': subjects[i]['bbox'],
|
||||
'score': subjects[i]['score'],
|
||||
},
|
||||
'obj_bboxes': [],
|
||||
'sub_idx': i,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
return ret
|
||||
|
||||
|
||||
def get_imgs_v2(self, page_no: int):
|
||||
with_captions = self.__tie_up_category_by_distance_v3(
|
||||
page_no, 3, 4, PosRelationEnum.BOTTOM
|
||||
)
|
||||
with_footnotes = self.__tie_up_category_by_distance_v3(
|
||||
page_no, 3, CategoryId.ImageFootnote, PosRelationEnum.ALL
|
||||
)
|
||||
ret = []
|
||||
for v in with_captions:
|
||||
record = {
|
||||
'image_body': v['sub_bbox'],
|
||||
'image_caption_list': v['obj_bboxes'],
|
||||
}
|
||||
filter_idx = v['sub_idx']
|
||||
d = next(filter(lambda x: x['sub_idx'] == filter_idx, with_footnotes))
|
||||
record['image_footnote_list'] = d['obj_bboxes']
|
||||
ret.append(record)
|
||||
return ret
|
||||
|
||||
def get_tables_v2(self, page_no: int) -> list:
|
||||
with_captions = self.__tie_up_category_by_distance_v3(
|
||||
page_no, 5, 6, PosRelationEnum.UP
|
||||
)
|
||||
with_footnotes = self.__tie_up_category_by_distance_v3(
|
||||
page_no, 5, 7, PosRelationEnum.ALL
|
||||
)
|
||||
ret = []
|
||||
for v in with_captions:
|
||||
record = {
|
||||
'table_body': v['sub_bbox'],
|
||||
'table_caption_list': v['obj_bboxes'],
|
||||
}
|
||||
filter_idx = v['sub_idx']
|
||||
d = next(filter(lambda x: x['sub_idx'] == filter_idx, with_footnotes))
|
||||
record['table_footnote_list'] = d['obj_bboxes']
|
||||
ret.append(record)
|
||||
return ret
|
||||
|
||||
def get_imgs(self, page_no: int):
|
||||
return self.get_imgs_v2(page_no)
|
||||
|
||||
def get_tables(
|
||||
self, page_no: int
|
||||
) -> list: # 3个坐标, caption, table主体,table-note
|
||||
return self.get_tables_v2(page_no)
|
||||
|
||||
def get_equations(self, page_no: int) -> list: # 有坐标,也有字
|
||||
inline_equations = self.__get_blocks_by_type(
|
||||
ModelBlockTypeEnum.EMBEDDING.value, page_no, ['latex']
|
||||
)
|
||||
interline_equations = self.__get_blocks_by_type(
|
||||
ModelBlockTypeEnum.ISOLATED.value, page_no, ['latex']
|
||||
)
|
||||
interline_equations_blocks = self.__get_blocks_by_type(
|
||||
ModelBlockTypeEnum.ISOLATE_FORMULA.value, page_no
|
||||
)
|
||||
return inline_equations, interline_equations, interline_equations_blocks
|
||||
|
||||
def get_discarded(self, page_no: int) -> list: # 自研模型,只有坐标
|
||||
blocks = self.__get_blocks_by_type(ModelBlockTypeEnum.ABANDON.value, page_no)
|
||||
return blocks
|
||||
|
||||
def get_text_blocks(self, page_no: int) -> list: # 自研模型搞的,只有坐标,没有字
|
||||
blocks = self.__get_blocks_by_type(ModelBlockTypeEnum.PLAIN_TEXT.value, page_no)
|
||||
return blocks
|
||||
|
||||
def get_title_blocks(self, page_no: int) -> list: # 自研模型,只有坐标,没字
|
||||
blocks = self.__get_blocks_by_type(ModelBlockTypeEnum.TITLE.value, page_no)
|
||||
return blocks
|
||||
|
||||
def get_ocr_text(self, page_no: int) -> list: # paddle 搞的,有字也有坐标
|
||||
text_spans = []
|
||||
model_page_info = self.__model_list[page_no]
|
||||
layout_dets = model_page_info['layout_dets']
|
||||
for layout_det in layout_dets:
|
||||
if layout_det['category_id'] == '15':
|
||||
span = {
|
||||
'bbox': layout_det['bbox'],
|
||||
'content': layout_det['text'],
|
||||
}
|
||||
text_spans.append(span)
|
||||
return text_spans
|
||||
|
||||
def get_all_spans(self, page_no: int) -> list:
|
||||
|
||||
def remove_duplicate_spans(spans):
|
||||
new_spans = []
|
||||
for span in spans:
|
||||
if not any(span == existing_span for existing_span in new_spans):
|
||||
new_spans.append(span)
|
||||
return new_spans
|
||||
|
||||
all_spans = []
|
||||
model_page_info = self.__model_list[page_no]
|
||||
layout_dets = model_page_info['layout_dets']
|
||||
allow_category_id_list = [3, 5, 13, 14, 15]
|
||||
"""当成span拼接的"""
|
||||
# 3: 'image', # 图片
|
||||
# 5: 'table', # 表格
|
||||
# 13: 'inline_equation', # 行内公式
|
||||
# 14: 'interline_equation', # 行间公式
|
||||
# 15: 'text', # ocr识别文本
|
||||
for layout_det in layout_dets:
|
||||
category_id = layout_det['category_id']
|
||||
if category_id in allow_category_id_list:
|
||||
span = {'bbox': layout_det['bbox'], 'score': layout_det['score']}
|
||||
if category_id == 3:
|
||||
span['type'] = ContentType.Image
|
||||
elif category_id == 5:
|
||||
# 获取table模型结果
|
||||
latex = layout_det.get('latex', None)
|
||||
html = layout_det.get('html', None)
|
||||
if latex:
|
||||
span['latex'] = latex
|
||||
elif html:
|
||||
span['html'] = html
|
||||
span['type'] = ContentType.Table
|
||||
elif category_id == 13:
|
||||
span['content'] = layout_det['latex']
|
||||
span['type'] = ContentType.InlineEquation
|
||||
elif category_id == 14:
|
||||
span['content'] = layout_det['latex']
|
||||
span['type'] = ContentType.InterlineEquation
|
||||
elif category_id == 15:
|
||||
span['content'] = layout_det['text']
|
||||
span['type'] = ContentType.Text
|
||||
all_spans.append(span)
|
||||
return remove_duplicate_spans(all_spans)
|
||||
|
||||
def get_page_size(self, page_no: int): # 获取页面宽高
|
||||
# 获取当前页的page对象
|
||||
page = self.__docs.get_page(page_no).get_page_info()
|
||||
# 获取当前页的宽高
|
||||
page_w = page.w
|
||||
page_h = page.h
|
||||
return page_w, page_h
|
||||
|
||||
def __get_blocks_by_type(
|
||||
self, type: int, page_no: int, extra_col: list[str] = []
|
||||
) -> list:
|
||||
blocks = []
|
||||
for page_dict in self.__model_list:
|
||||
layout_dets = page_dict.get('layout_dets', [])
|
||||
page_info = page_dict.get('page_info', {})
|
||||
page_number = page_info.get('page_no', -1)
|
||||
if page_no != page_number:
|
||||
continue
|
||||
for item in layout_dets:
|
||||
category_id = item.get('category_id', -1)
|
||||
bbox = item.get('bbox', None)
|
||||
|
||||
if category_id == type:
|
||||
block = {
|
||||
'bbox': bbox,
|
||||
'score': item.get('score'),
|
||||
}
|
||||
for col in extra_col:
|
||||
block[col] = item.get(col, None)
|
||||
blocks.append(block)
|
||||
return blocks
|
||||
|
||||
def get_model_list(self, page_no):
|
||||
return self.__model_list[page_no]
|
||||
@@ -0,0 +1,190 @@
|
||||
import os
|
||||
|
||||
import torch
|
||||
from loguru import logger
|
||||
|
||||
from .model_list import AtomicModel
|
||||
from ...model.layout.doclayout_yolo import DocLayoutYOLOModel
|
||||
from ...model.mfd.yolo_v8 import YOLOv8MFDModel
|
||||
from ...model.mfr.unimernet.Unimernet import UnimernetModel
|
||||
from ...model.ocr.paddleocr2pytorch.pytorch_paddle import PytorchPaddleOCR
|
||||
from ...model.table.rapid_table import RapidTableModel
|
||||
|
||||
doclayout_yolo = "Layout/YOLO/doclayout_yolo_docstructbench_imgsz1280_2501.pt"
|
||||
yolo_v8_mfd = "MFD/YOLO/yolo_v8_ft.pt"
|
||||
unimernet_small = "MFR/unimernet_hf_small_2503"
|
||||
|
||||
|
||||
def table_model_init(lang=None):
|
||||
atom_model_manager = AtomModelSingleton()
|
||||
ocr_engine = atom_model_manager.get_atom_model(
|
||||
atom_model_name='ocr',
|
||||
det_db_box_thresh=0.5,
|
||||
det_db_unclip_ratio=1.6,
|
||||
lang=lang
|
||||
)
|
||||
table_model = RapidTableModel(ocr_engine)
|
||||
return table_model
|
||||
|
||||
|
||||
def mfd_model_init(weight, device='cpu'):
|
||||
if str(device).startswith('npu'):
|
||||
device = torch.device(device)
|
||||
mfd_model = YOLOv8MFDModel(weight, device)
|
||||
return mfd_model
|
||||
|
||||
|
||||
def mfr_model_init(weight_dir, device='cpu'):
|
||||
mfr_model = UnimernetModel(weight_dir, device)
|
||||
return mfr_model
|
||||
|
||||
|
||||
def doclayout_yolo_model_init(weight, device='cpu'):
|
||||
if str(device).startswith('npu'):
|
||||
device = torch.device(device)
|
||||
model = DocLayoutYOLOModel(weight, device)
|
||||
return model
|
||||
|
||||
def ocr_model_init(det_db_box_thresh=0.3,
|
||||
lang=None,
|
||||
use_dilation=True,
|
||||
det_db_unclip_ratio=1.8,
|
||||
):
|
||||
if lang is not None and lang != '':
|
||||
model = PytorchPaddleOCR(
|
||||
det_db_box_thresh=det_db_box_thresh,
|
||||
lang=lang,
|
||||
use_dilation=use_dilation,
|
||||
det_db_unclip_ratio=det_db_unclip_ratio,
|
||||
)
|
||||
else:
|
||||
model = PytorchPaddleOCR(
|
||||
det_db_box_thresh=det_db_box_thresh,
|
||||
use_dilation=use_dilation,
|
||||
det_db_unclip_ratio=det_db_unclip_ratio,
|
||||
)
|
||||
return model
|
||||
|
||||
|
||||
class AtomModelSingleton:
|
||||
_instance = None
|
||||
_models = {}
|
||||
|
||||
def __new__(cls, *args, **kwargs):
|
||||
if cls._instance is None:
|
||||
cls._instance = super().__new__(cls)
|
||||
return cls._instance
|
||||
|
||||
def get_atom_model(self, atom_model_name: str, **kwargs):
|
||||
|
||||
lang = kwargs.get('lang', None)
|
||||
table_model_name = kwargs.get('table_model_name', None)
|
||||
|
||||
if atom_model_name in [AtomicModel.OCR]:
|
||||
key = (atom_model_name, lang)
|
||||
elif atom_model_name in [AtomicModel.Table]:
|
||||
key = (atom_model_name, table_model_name, lang)
|
||||
else:
|
||||
key = atom_model_name
|
||||
|
||||
if key not in self._models:
|
||||
self._models[key] = atom_model_init(model_name=atom_model_name, **kwargs)
|
||||
return self._models[key]
|
||||
|
||||
def atom_model_init(model_name: str, **kwargs):
|
||||
atom_model = None
|
||||
if model_name == AtomicModel.Layout:
|
||||
atom_model = doclayout_yolo_model_init(
|
||||
kwargs.get('doclayout_yolo_weights'),
|
||||
kwargs.get('device')
|
||||
)
|
||||
elif model_name == AtomicModel.MFD:
|
||||
atom_model = mfd_model_init(
|
||||
kwargs.get('mfd_weights'),
|
||||
kwargs.get('device')
|
||||
)
|
||||
elif model_name == AtomicModel.MFR:
|
||||
atom_model = mfr_model_init(
|
||||
kwargs.get('mfr_weight_dir'),
|
||||
kwargs.get('device')
|
||||
)
|
||||
elif model_name == AtomicModel.OCR:
|
||||
atom_model = ocr_model_init(
|
||||
kwargs.get('det_db_box_thresh'),
|
||||
kwargs.get('lang'),
|
||||
)
|
||||
elif model_name == AtomicModel.Table:
|
||||
atom_model = table_model_init(
|
||||
kwargs.get('lang'),
|
||||
)
|
||||
else:
|
||||
logger.error('model name not allow')
|
||||
exit(1)
|
||||
|
||||
if atom_model is None:
|
||||
logger.error('model init failed')
|
||||
exit(1)
|
||||
else:
|
||||
return atom_model
|
||||
|
||||
|
||||
class MineruPipelineModel:
|
||||
def __init__(self, **kwargs):
|
||||
self.formula_config = kwargs.get('formula_config')
|
||||
self.apply_formula = self.formula_config.get('enable', True)
|
||||
self.table_config = kwargs.get('table_config')
|
||||
self.apply_table = self.table_config.get('enable', True)
|
||||
self.lang = kwargs.get('lang', None)
|
||||
self.device = kwargs.get('device', 'cpu')
|
||||
logger.info(
|
||||
'DocAnalysis init, this may take some times......'
|
||||
)
|
||||
atom_model_manager = AtomModelSingleton()
|
||||
models_dir = kwargs.get('models_dir', "")
|
||||
if not models_dir:
|
||||
logger.error("can't found models_dir, please set models_dir")
|
||||
exit(1)
|
||||
|
||||
if self.apply_formula:
|
||||
# 初始化公式检测模型
|
||||
self.mfd_model = atom_model_manager.get_atom_model(
|
||||
atom_model_name=AtomicModel.MFD,
|
||||
mfd_weights=str(
|
||||
os.path.join(models_dir, yolo_v8_mfd)
|
||||
),
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
# 初始化公式解析模型
|
||||
mfr_weight_dir = str(
|
||||
os.path.join(models_dir, unimernet_small)
|
||||
)
|
||||
|
||||
self.mfr_model = atom_model_manager.get_atom_model(
|
||||
atom_model_name=AtomicModel.MFR,
|
||||
mfr_weight_dir=mfr_weight_dir,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
# 初始化layout模型
|
||||
self.layout_model = atom_model_manager.get_atom_model(
|
||||
atom_model_name=AtomicModel.Layout,
|
||||
doclayout_yolo_weights=str(
|
||||
os.path.join(models_dir, doclayout_yolo)
|
||||
),
|
||||
device=self.device,
|
||||
)
|
||||
# 初始化ocr
|
||||
self.ocr_model = atom_model_manager.get_atom_model(
|
||||
atom_model_name=AtomicModel.OCR,
|
||||
det_db_box_thresh=0.3,
|
||||
lang=self.lang
|
||||
)
|
||||
# init table model
|
||||
if self.apply_table:
|
||||
self.table_model = atom_model_manager.get_atom_model(
|
||||
atom_model_name=AtomicModel.Table,
|
||||
lang=self.lang,
|
||||
)
|
||||
|
||||
logger.info('DocAnalysis init done!')
|
||||
@@ -0,0 +1,6 @@
|
||||
class AtomicModel:
|
||||
Layout = "layout"
|
||||
MFD = "mfd"
|
||||
MFR = "mfr"
|
||||
OCR = "ocr"
|
||||
Table = "table"
|
||||
@@ -1,10 +1,10 @@
|
||||
import re
|
||||
|
||||
from ...libs.cut_image import cut_image_and_table
|
||||
from ...libs.enum_class import BlockType, ContentType
|
||||
from ...libs.hash_utils import str_md5
|
||||
from ...libs.magic_model import fix_two_layer_blocks
|
||||
from ...libs.version import __version__
|
||||
from mineru.utils.cut_image import cut_image_and_table
|
||||
from mineru.utils.enum_class import BlockType, ContentType
|
||||
from mineru.utils.hash_utils import str_md5
|
||||
from mineru.utils.magic_model import fix_two_layer_blocks
|
||||
from mineru.version import __version__
|
||||
|
||||
|
||||
def token_to_page_info(token, image_dict, page, image_writer, page_index) -> dict:
|
||||
|
||||
@@ -4,7 +4,7 @@ import time
|
||||
from loguru import logger
|
||||
|
||||
from ...data.data_reader_writer import DataWriter
|
||||
from ...libs.pdf_image_tools import load_images_from_pdf
|
||||
from mineru.utils.pdf_image_tools import load_images_from_pdf
|
||||
from .base_predictor import BasePredictor
|
||||
from .predictor import get_predictor
|
||||
from .token_to_middle_json import result_to_middle_json
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
# Copyright (c) Opendatalab. All rights reserved.
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
# Copyright (c) Opendatalab. All rights reserved.
|
||||
@@ -0,0 +1,64 @@
|
||||
from doclayout_yolo import YOLOv10
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
class DocLayoutYOLOModel(object):
|
||||
def __init__(self, weight, device):
|
||||
self.model = YOLOv10(weight)
|
||||
self.device = device
|
||||
|
||||
def predict(self, image):
|
||||
layout_res = []
|
||||
doclayout_yolo_res = self.model.predict(
|
||||
image,
|
||||
imgsz=1280,
|
||||
conf=0.10,
|
||||
iou=0.45,
|
||||
verbose=False, device=self.device
|
||||
)[0]
|
||||
for xyxy, conf, cla in zip(
|
||||
doclayout_yolo_res.boxes.xyxy.cpu(),
|
||||
doclayout_yolo_res.boxes.conf.cpu(),
|
||||
doclayout_yolo_res.boxes.cls.cpu(),
|
||||
):
|
||||
xmin, ymin, xmax, ymax = [int(p.item()) for p in xyxy]
|
||||
new_item = {
|
||||
"category_id": int(cla.item()),
|
||||
"poly": [xmin, ymin, xmax, ymin, xmax, ymax, xmin, ymax],
|
||||
"score": round(float(conf.item()), 3),
|
||||
}
|
||||
layout_res.append(new_item)
|
||||
return layout_res
|
||||
|
||||
def batch_predict(self, images: list, batch_size: int) -> list:
|
||||
images_layout_res = []
|
||||
# for index in range(0, len(images), batch_size):
|
||||
for index in tqdm(range(0, len(images), batch_size), desc="Layout Predict"):
|
||||
doclayout_yolo_res = [
|
||||
image_res.cpu()
|
||||
for image_res in self.model.predict(
|
||||
images[index : index + batch_size],
|
||||
imgsz=1280,
|
||||
conf=0.10,
|
||||
iou=0.45,
|
||||
verbose=False,
|
||||
device=self.device,
|
||||
)
|
||||
]
|
||||
for image_res in doclayout_yolo_res:
|
||||
layout_res = []
|
||||
for xyxy, conf, cla in zip(
|
||||
image_res.boxes.xyxy,
|
||||
image_res.boxes.conf,
|
||||
image_res.boxes.cls,
|
||||
):
|
||||
xmin, ymin, xmax, ymax = [int(p.item()) for p in xyxy]
|
||||
new_item = {
|
||||
"category_id": int(cla.item()),
|
||||
"poly": [xmin, ymin, xmax, ymin, xmax, ymax, xmin, ymax],
|
||||
"score": round(float(conf.item()), 3),
|
||||
}
|
||||
layout_res.append(new_item)
|
||||
images_layout_res.append(layout_res)
|
||||
|
||||
return images_layout_res
|
||||
@@ -0,0 +1 @@
|
||||
# Copyright (c) Opendatalab. All rights reserved.
|
||||
@@ -0,0 +1,33 @@
|
||||
from tqdm import tqdm
|
||||
from ultralytics import YOLO
|
||||
|
||||
|
||||
class YOLOv8MFDModel(object):
|
||||
def __init__(self, weight, device="cpu"):
|
||||
self.mfd_model = YOLO(weight)
|
||||
self.device = device
|
||||
|
||||
def predict(self, image):
|
||||
mfd_res = self.mfd_model.predict(
|
||||
image, imgsz=1888, conf=0.25, iou=0.45, verbose=False, device=self.device
|
||||
)[0]
|
||||
return mfd_res
|
||||
|
||||
def batch_predict(self, images: list, batch_size: int) -> list:
|
||||
images_mfd_res = []
|
||||
# for index in range(0, len(images), batch_size):
|
||||
for index in tqdm(range(0, len(images), batch_size), desc="MFD Predict"):
|
||||
mfd_res = [
|
||||
image_res.cpu()
|
||||
for image_res in self.mfd_model.predict(
|
||||
images[index : index + batch_size],
|
||||
imgsz=1888,
|
||||
conf=0.25,
|
||||
iou=0.45,
|
||||
verbose=False,
|
||||
device=self.device,
|
||||
)
|
||||
]
|
||||
for image_res in mfd_res:
|
||||
images_mfd_res.append(image_res)
|
||||
return images_mfd_res
|
||||
@@ -0,0 +1 @@
|
||||
# Copyright (c) Opendatalab. All rights reserved.
|
||||
@@ -0,0 +1,135 @@
|
||||
import torch
|
||||
from torch.utils.data import DataLoader, Dataset
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
class MathDataset(Dataset):
|
||||
def __init__(self, image_paths, transform=None):
|
||||
self.image_paths = image_paths
|
||||
self.transform = transform
|
||||
|
||||
def __len__(self):
|
||||
return len(self.image_paths)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
raw_image = self.image_paths[idx]
|
||||
if self.transform:
|
||||
image = self.transform(raw_image)
|
||||
return image
|
||||
|
||||
|
||||
class UnimernetModel(object):
|
||||
def __init__(self, weight_dir, cfg_path, _device_="cpu"):
|
||||
from .unimernet_hf import UnimernetModel
|
||||
if _device_.startswith("mps"):
|
||||
self.model = UnimernetModel.from_pretrained(weight_dir, attn_implementation="eager")
|
||||
else:
|
||||
self.model = UnimernetModel.from_pretrained(weight_dir)
|
||||
self.device = _device_
|
||||
self.model.to(_device_)
|
||||
if not _device_.startswith("cpu"):
|
||||
self.model = self.model.to(dtype=torch.float16)
|
||||
self.model.eval()
|
||||
|
||||
def predict(self, mfd_res, image):
|
||||
formula_list = []
|
||||
mf_image_list = []
|
||||
for xyxy, conf, cla in zip(
|
||||
mfd_res.boxes.xyxy.cpu(), mfd_res.boxes.conf.cpu(), mfd_res.boxes.cls.cpu()
|
||||
):
|
||||
xmin, ymin, xmax, ymax = [int(p.item()) for p in xyxy]
|
||||
new_item = {
|
||||
"category_id": 13 + int(cla.item()),
|
||||
"poly": [xmin, ymin, xmax, ymin, xmax, ymax, xmin, ymax],
|
||||
"score": round(float(conf.item()), 2),
|
||||
"latex": "",
|
||||
}
|
||||
formula_list.append(new_item)
|
||||
bbox_img = image[ymin:ymax, xmin:xmax]
|
||||
mf_image_list.append(bbox_img)
|
||||
|
||||
dataset = MathDataset(mf_image_list, transform=self.model.transform)
|
||||
dataloader = DataLoader(dataset, batch_size=32, num_workers=0)
|
||||
mfr_res = []
|
||||
for mf_img in dataloader:
|
||||
mf_img = mf_img.to(dtype=self.model.dtype)
|
||||
mf_img = mf_img.to(self.device)
|
||||
with torch.no_grad():
|
||||
output = self.model.generate({"image": mf_img})
|
||||
mfr_res.extend(output["fixed_str"])
|
||||
for res, latex in zip(formula_list, mfr_res):
|
||||
res["latex"] = latex
|
||||
return formula_list
|
||||
|
||||
def batch_predict(self, images_mfd_res: list, images: list, batch_size: int = 64) -> list:
|
||||
images_formula_list = []
|
||||
mf_image_list = []
|
||||
backfill_list = []
|
||||
image_info = [] # Store (area, original_index, image) tuples
|
||||
|
||||
# Collect images with their original indices
|
||||
for image_index in range(len(images_mfd_res)):
|
||||
mfd_res = images_mfd_res[image_index]
|
||||
np_array_image = images[image_index]
|
||||
formula_list = []
|
||||
|
||||
for idx, (xyxy, conf, cla) in enumerate(zip(
|
||||
mfd_res.boxes.xyxy, mfd_res.boxes.conf, mfd_res.boxes.cls
|
||||
)):
|
||||
xmin, ymin, xmax, ymax = [int(p.item()) for p in xyxy]
|
||||
new_item = {
|
||||
"category_id": 13 + int(cla.item()),
|
||||
"poly": [xmin, ymin, xmax, ymin, xmax, ymax, xmin, ymax],
|
||||
"score": round(float(conf.item()), 2),
|
||||
"latex": "",
|
||||
}
|
||||
formula_list.append(new_item)
|
||||
bbox_img = np_array_image[ymin:ymax, xmin:xmax]
|
||||
area = (xmax - xmin) * (ymax - ymin)
|
||||
|
||||
curr_idx = len(mf_image_list)
|
||||
image_info.append((area, curr_idx, bbox_img))
|
||||
mf_image_list.append(bbox_img)
|
||||
|
||||
images_formula_list.append(formula_list)
|
||||
backfill_list += formula_list
|
||||
|
||||
# Stable sort by area
|
||||
image_info.sort(key=lambda x: x[0]) # sort by area
|
||||
sorted_indices = [x[1] for x in image_info]
|
||||
sorted_images = [x[2] for x in image_info]
|
||||
|
||||
# Create mapping for results
|
||||
index_mapping = {new_idx: old_idx for new_idx, old_idx in enumerate(sorted_indices)}
|
||||
|
||||
# Create dataset with sorted images
|
||||
dataset = MathDataset(sorted_images, transform=self.model.transform)
|
||||
dataloader = DataLoader(dataset, batch_size=batch_size, num_workers=0)
|
||||
|
||||
# Process batches and store results
|
||||
mfr_res = []
|
||||
# for mf_img in dataloader:
|
||||
|
||||
with tqdm(total=len(sorted_images), desc="MFR Predict") as pbar:
|
||||
for index, mf_img in enumerate(dataloader):
|
||||
mf_img = mf_img.to(dtype=self.model.dtype)
|
||||
mf_img = mf_img.to(self.device)
|
||||
with torch.no_grad():
|
||||
output = self.model.generate({"image": mf_img})
|
||||
mfr_res.extend(output["fixed_str"])
|
||||
|
||||
# 更新进度条,每次增加batch_size,但要注意最后一个batch可能不足batch_size
|
||||
current_batch_size = min(batch_size, len(sorted_images) - index * batch_size)
|
||||
pbar.update(current_batch_size)
|
||||
|
||||
# Restore original order
|
||||
unsorted_results = [""] * len(mfr_res)
|
||||
for new_idx, latex in enumerate(mfr_res):
|
||||
original_idx = index_mapping[new_idx]
|
||||
unsorted_results[original_idx] = latex
|
||||
|
||||
# Fill results back
|
||||
for res, latex in zip(backfill_list, unsorted_results):
|
||||
res["latex"] = latex
|
||||
|
||||
return images_formula_list
|
||||
@@ -0,0 +1,13 @@
|
||||
from .unimer_swin import UnimerSwinConfig, UnimerSwinModel, UnimerSwinImageProcessor
|
||||
from .unimer_mbart import UnimerMBartConfig, UnimerMBartModel, UnimerMBartForCausalLM
|
||||
from .modeling_unimernet import UnimernetModel
|
||||
|
||||
__all__ = [
|
||||
"UnimerSwinConfig",
|
||||
"UnimerSwinModel",
|
||||
"UnimerSwinImageProcessor",
|
||||
"UnimerMBartConfig",
|
||||
"UnimerMBartModel",
|
||||
"UnimerMBartForCausalLM",
|
||||
"UnimernetModel",
|
||||
]
|
||||
@@ -0,0 +1,490 @@
|
||||
import os
|
||||
import re
|
||||
import warnings
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
from ftfy import fix_text
|
||||
from loguru import logger
|
||||
|
||||
from transformers import AutoConfig, AutoModel, AutoModelForCausalLM, AutoTokenizer, PretrainedConfig, PreTrainedModel
|
||||
from transformers import VisionEncoderDecoderConfig, VisionEncoderDecoderModel
|
||||
from transformers.models.vision_encoder_decoder.modeling_vision_encoder_decoder import logger as base_model_logger
|
||||
|
||||
from .unimer_swin import UnimerSwinConfig, UnimerSwinModel, UnimerSwinImageProcessor
|
||||
from .unimer_mbart import UnimerMBartConfig, UnimerMBartForCausalLM
|
||||
|
||||
AutoConfig.register(UnimerSwinConfig.model_type, UnimerSwinConfig)
|
||||
AutoConfig.register(UnimerMBartConfig.model_type, UnimerMBartConfig)
|
||||
AutoModel.register(UnimerSwinConfig, UnimerSwinModel)
|
||||
AutoModelForCausalLM.register(UnimerMBartConfig, UnimerMBartForCausalLM)
|
||||
|
||||
|
||||
# TODO: rewrite tokenizer
|
||||
class TokenizerWrapper:
|
||||
def __init__(self, tokenizer):
|
||||
self.tokenizer = tokenizer
|
||||
self.pad_token_id = self.tokenizer.pad_token_id
|
||||
self.bos_token_id = self.tokenizer.bos_token_id
|
||||
self.eos_token_id = self.tokenizer.eos_token_id
|
||||
|
||||
def __len__(self):
|
||||
return len(self.tokenizer)
|
||||
|
||||
def tokenize(self, text, **kwargs):
|
||||
return self.tokenizer(
|
||||
text,
|
||||
return_token_type_ids=False,
|
||||
return_tensors="pt",
|
||||
padding="longest",
|
||||
truncation=True,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def token2str(self, tokens) -> list:
|
||||
generated_text = self.tokenizer.batch_decode(tokens, skip_special_tokens=True)
|
||||
generated_text = [fix_text(text) for text in generated_text]
|
||||
return generated_text
|
||||
|
||||
def detokenize(self, tokens):
|
||||
toks = [self.tokenizer.convert_ids_to_tokens(tok) for tok in tokens]
|
||||
for b in range(len(toks)):
|
||||
for i in reversed(range(len(toks[b]))):
|
||||
if toks[b][i] is None:
|
||||
toks[b][i] = ''
|
||||
toks[b][i] = toks[b][i].replace('Ġ', ' ').strip()
|
||||
if toks[b][i] in ([self.tokenizer.bos_token, self.tokenizer.eos_token, self.tokenizer.pad_token]):
|
||||
del toks[b][i]
|
||||
return toks
|
||||
|
||||
|
||||
LEFT_PATTERN = re.compile(r'(\\left)(\S*)')
|
||||
RIGHT_PATTERN = re.compile(r'(\\right)(\S*)')
|
||||
LEFT_COUNT_PATTERN = re.compile(r'\\left(?![a-zA-Z])')
|
||||
RIGHT_COUNT_PATTERN = re.compile(r'\\right(?![a-zA-Z])')
|
||||
LEFT_RIGHT_REMOVE_PATTERN = re.compile(r'\\left\.?|\\right\.?')
|
||||
|
||||
def fix_latex_left_right(s):
|
||||
"""
|
||||
修复LaTeX中的\\left和\\right命令
|
||||
1. 确保它们后面跟有效分隔符
|
||||
2. 平衡\\left和\\right的数量
|
||||
"""
|
||||
# 白名单分隔符
|
||||
valid_delims_list = [r'(', r')', r'[', r']', r'{', r'}', r'/', r'|',
|
||||
r'\{', r'\}', r'\lceil', r'\rceil', r'\lfloor',
|
||||
r'\rfloor', r'\backslash', r'\uparrow', r'\downarrow',
|
||||
r'\Uparrow', r'\Downarrow', r'\|', r'\.']
|
||||
|
||||
# 为\left后缺失有效分隔符的情况添加点
|
||||
def fix_delim(match, is_left=True):
|
||||
cmd = match.group(1) # \left 或 \right
|
||||
rest = match.group(2) if len(match.groups()) > 1 else ""
|
||||
if not rest or rest not in valid_delims_list:
|
||||
return cmd + "."
|
||||
return match.group(0)
|
||||
|
||||
# 使用更精确的模式匹配\left和\right命令
|
||||
# 确保它们是独立的命令,不是其他命令的一部分
|
||||
# 使用预编译正则和统一回调函数
|
||||
s = LEFT_PATTERN.sub(lambda m: fix_delim(m, True), s)
|
||||
s = RIGHT_PATTERN.sub(lambda m: fix_delim(m, False), s)
|
||||
|
||||
# 更精确地计算\left和\right的数量
|
||||
left_count = len(LEFT_COUNT_PATTERN.findall(s)) # 不匹配\lefteqn等
|
||||
right_count = len(RIGHT_COUNT_PATTERN.findall(s)) # 不匹配\rightarrow等
|
||||
|
||||
if left_count == right_count:
|
||||
# 如果数量相等,检查是否在同一组
|
||||
return fix_left_right_pairs(s)
|
||||
else:
|
||||
# 如果数量不等,移除所有\left和\right
|
||||
# logger.debug(f"latex:{s}")
|
||||
# logger.warning(f"left_count: {left_count}, right_count: {right_count}")
|
||||
return LEFT_RIGHT_REMOVE_PATTERN.sub('', s)
|
||||
|
||||
|
||||
def fix_left_right_pairs(latex_formula):
|
||||
"""
|
||||
检测并修复LaTeX公式中\\left和\\right不在同一组的情况
|
||||
|
||||
Args:
|
||||
latex_formula (str): 输入的LaTeX公式
|
||||
|
||||
Returns:
|
||||
str: 修复后的LaTeX公式
|
||||
"""
|
||||
# 用于跟踪花括号嵌套层级
|
||||
brace_stack = []
|
||||
# 用于存储\left信息: (位置, 深度, 分隔符)
|
||||
left_stack = []
|
||||
# 存储需要调整的\right信息: (开始位置, 结束位置, 目标位置)
|
||||
adjustments = []
|
||||
|
||||
i = 0
|
||||
while i < len(latex_formula):
|
||||
# 检查是否是转义字符
|
||||
if i > 0 and latex_formula[i - 1] == '\\':
|
||||
backslash_count = 0
|
||||
j = i - 1
|
||||
while j >= 0 and latex_formula[j] == '\\':
|
||||
backslash_count += 1
|
||||
j -= 1
|
||||
|
||||
if backslash_count % 2 == 1:
|
||||
i += 1
|
||||
continue
|
||||
|
||||
# 检测\left命令
|
||||
if i + 5 < len(latex_formula) and latex_formula[i:i + 5] == "\\left" and i + 5 < len(latex_formula):
|
||||
delimiter = latex_formula[i + 5]
|
||||
left_stack.append((i, len(brace_stack), delimiter))
|
||||
i += 6 # 跳过\left和分隔符
|
||||
continue
|
||||
|
||||
# 检测\right命令
|
||||
elif i + 6 < len(latex_formula) and latex_formula[i:i + 6] == "\\right" and i + 6 < len(latex_formula):
|
||||
delimiter = latex_formula[i + 6]
|
||||
|
||||
if left_stack:
|
||||
left_pos, left_depth, left_delim = left_stack.pop()
|
||||
|
||||
# 如果\left和\right不在同一花括号深度
|
||||
if left_depth != len(brace_stack):
|
||||
# 找到\left所在花括号组的结束位置
|
||||
target_pos = find_group_end(latex_formula, left_pos, left_depth)
|
||||
if target_pos != -1:
|
||||
# 记录需要移动的\right
|
||||
adjustments.append((i, i + 7, target_pos))
|
||||
|
||||
i += 7 # 跳过\right和分隔符
|
||||
continue
|
||||
|
||||
# 处理花括号
|
||||
if latex_formula[i] == '{':
|
||||
brace_stack.append(i)
|
||||
elif latex_formula[i] == '}':
|
||||
if brace_stack:
|
||||
brace_stack.pop()
|
||||
|
||||
i += 1
|
||||
|
||||
# 应用调整,从后向前处理以避免索引变化
|
||||
if not adjustments:
|
||||
return latex_formula
|
||||
|
||||
result = list(latex_formula)
|
||||
adjustments.sort(reverse=True, key=lambda x: x[0])
|
||||
|
||||
for start, end, target in adjustments:
|
||||
# 提取\right部分
|
||||
right_part = result[start:end]
|
||||
# 从原位置删除
|
||||
del result[start:end]
|
||||
# 在目标位置插入
|
||||
result.insert(target, ''.join(right_part))
|
||||
|
||||
return ''.join(result)
|
||||
|
||||
|
||||
def find_group_end(text, pos, depth):
|
||||
"""查找特定深度的花括号组的结束位置"""
|
||||
current_depth = depth
|
||||
i = pos
|
||||
|
||||
while i < len(text):
|
||||
if text[i] == '{' and (i == 0 or not is_escaped(text, i)):
|
||||
current_depth += 1
|
||||
elif text[i] == '}' and (i == 0 or not is_escaped(text, i)):
|
||||
current_depth -= 1
|
||||
if current_depth < depth:
|
||||
return i
|
||||
i += 1
|
||||
|
||||
return -1 # 未找到对应结束位置
|
||||
|
||||
|
||||
def is_escaped(text, pos):
|
||||
"""检查字符是否被转义"""
|
||||
backslash_count = 0
|
||||
j = pos - 1
|
||||
while j >= 0 and text[j] == '\\':
|
||||
backslash_count += 1
|
||||
j -= 1
|
||||
|
||||
return backslash_count % 2 == 1
|
||||
|
||||
|
||||
def fix_unbalanced_braces(latex_formula):
|
||||
"""
|
||||
检测LaTeX公式中的花括号是否闭合,并删除无法配对的花括号
|
||||
|
||||
Args:
|
||||
latex_formula (str): 输入的LaTeX公式
|
||||
|
||||
Returns:
|
||||
str: 删除无法配对的花括号后的LaTeX公式
|
||||
"""
|
||||
stack = [] # 存储左括号的索引
|
||||
unmatched = set() # 存储不匹配括号的索引
|
||||
i = 0
|
||||
|
||||
while i < len(latex_formula):
|
||||
# 检查是否是转义的花括号
|
||||
if latex_formula[i] in ['{', '}']:
|
||||
# 计算前面连续的反斜杠数量
|
||||
backslash_count = 0
|
||||
j = i - 1
|
||||
while j >= 0 and latex_formula[j] == '\\':
|
||||
backslash_count += 1
|
||||
j -= 1
|
||||
|
||||
# 如果前面有奇数个反斜杠,则该花括号是转义的,不参与匹配
|
||||
if backslash_count % 2 == 1:
|
||||
i += 1
|
||||
continue
|
||||
|
||||
# 否则,该花括号参与匹配
|
||||
if latex_formula[i] == '{':
|
||||
stack.append(i)
|
||||
else: # latex_formula[i] == '}'
|
||||
if stack: # 有对应的左括号
|
||||
stack.pop()
|
||||
else: # 没有对应的左括号
|
||||
unmatched.add(i)
|
||||
|
||||
i += 1
|
||||
|
||||
# 所有未匹配的左括号
|
||||
unmatched.update(stack)
|
||||
|
||||
# 构建新字符串,删除不匹配的括号
|
||||
return ''.join(char for i, char in enumerate(latex_formula) if i not in unmatched)
|
||||
|
||||
|
||||
def process_latex(input_string):
|
||||
"""
|
||||
处理LaTeX公式中的反斜杠:
|
||||
1. 如果\后跟特殊字符(#$%&~_^\\{})或空格,保持不变
|
||||
2. 如果\后跟两个小写字母,保持不变
|
||||
3. 其他情况,在\后添加空格
|
||||
|
||||
Args:
|
||||
input_string (str): 输入的LaTeX公式
|
||||
|
||||
Returns:
|
||||
str: 处理后的LaTeX公式
|
||||
"""
|
||||
|
||||
def replace_func(match):
|
||||
# 获取\后面的字符
|
||||
next_char = match.group(1)
|
||||
|
||||
# 如果是特殊字符或空格,保持不变
|
||||
if next_char in "#$%&~_^|\\{} \t\n\r\v\f":
|
||||
return match.group(0)
|
||||
|
||||
# 如果是字母,检查下一个字符
|
||||
if 'a' <= next_char <= 'z' or 'A' <= next_char <= 'Z':
|
||||
pos = match.start() + 2 # \x后的位置
|
||||
if pos < len(input_string) and ('a' <= input_string[pos] <= 'z' or 'A' <= input_string[pos] <= 'Z'):
|
||||
# 下一个字符也是字母,保持不变
|
||||
return match.group(0)
|
||||
|
||||
# 其他情况,在\后添加空格
|
||||
return '\\' + ' ' + next_char
|
||||
|
||||
# 匹配\后面跟一个字符的情况
|
||||
pattern = r'\\(.)'
|
||||
|
||||
return re.sub(pattern, replace_func, input_string)
|
||||
|
||||
# 常见的在KaTeX/MathJax中可用的数学环境
|
||||
ENV_TYPES = ['array', 'matrix', 'pmatrix', 'bmatrix', 'vmatrix',
|
||||
'Bmatrix', 'Vmatrix', 'cases', 'aligned', 'gathered']
|
||||
ENV_BEGIN_PATTERNS = {env: re.compile(r'\\begin\{' + env + r'\}') for env in ENV_TYPES}
|
||||
ENV_END_PATTERNS = {env: re.compile(r'\\end\{' + env + r'\}') for env in ENV_TYPES}
|
||||
ENV_FORMAT_PATTERNS = {env: re.compile(r'\\begin\{' + env + r'\}\{([^}]*)\}') for env in ENV_TYPES}
|
||||
|
||||
def fix_latex_environments(s):
|
||||
"""
|
||||
检测LaTeX中环境(如array)的\\begin和\\end是否匹配
|
||||
1. 如果缺少\\begin标签则在开头添加
|
||||
2. 如果缺少\\end标签则在末尾添加
|
||||
"""
|
||||
for env in ENV_TYPES:
|
||||
begin_count = len(ENV_BEGIN_PATTERNS[env].findall(s))
|
||||
end_count = len(ENV_END_PATTERNS[env].findall(s))
|
||||
|
||||
if begin_count != end_count:
|
||||
if end_count > begin_count:
|
||||
format_match = ENV_FORMAT_PATTERNS[env].search(s)
|
||||
default_format = '{c}' if env == 'array' else ''
|
||||
format_str = '{' + format_match.group(1) + '}' if format_match else default_format
|
||||
|
||||
missing_count = end_count - begin_count
|
||||
begin_command = '\\begin{' + env + '}' + format_str + ' '
|
||||
s = begin_command * missing_count + s
|
||||
else:
|
||||
missing_count = begin_count - end_count
|
||||
s = s + (' \\end{' + env + '}') * missing_count
|
||||
|
||||
return s
|
||||
|
||||
|
||||
UP_PATTERN = re.compile(r'\\up([a-zA-Z]+)')
|
||||
COMMANDS_TO_REMOVE_PATTERN = re.compile(
|
||||
r'\\(?:lefteqn|boldmath|ensuremath|centering|textsubscript|sides|textsl|textcent|emph|protect|null)')
|
||||
REPLACEMENTS_PATTERNS = {
|
||||
re.compile(r'\\underbar'): r'\\underline',
|
||||
re.compile(r'\\Bar'): r'\\hat',
|
||||
re.compile(r'\\Hat'): r'\\hat',
|
||||
re.compile(r'\\Tilde'): r'\\tilde',
|
||||
re.compile(r'\\slash'): r'/',
|
||||
re.compile(r'\\textperthousand'): r'‰',
|
||||
re.compile(r'\\sun'): r'☉',
|
||||
re.compile(r'\\textunderscore'): r'\\_',
|
||||
re.compile(r'\\fint'): r'⨏',
|
||||
re.compile(r'\\up '): r'\\ ',
|
||||
re.compile(r'\\vline = '): r'\\models ',
|
||||
re.compile(r'\\vDash '): r'\\models ',
|
||||
re.compile(r'\\sq \\sqcup '): r'\\square ',
|
||||
}
|
||||
QQUAD_PATTERN = re.compile(r'\\qquad(?!\s)')
|
||||
|
||||
def latex_rm_whitespace(s: str):
|
||||
"""Remove unnecessary whitespace from LaTeX code."""
|
||||
s = fix_unbalanced_braces(s)
|
||||
s = fix_latex_left_right(s)
|
||||
s = fix_latex_environments(s)
|
||||
|
||||
# 使用预编译的正则表达式
|
||||
s = UP_PATTERN.sub(
|
||||
lambda m: m.group(0) if m.group(1) in ["arrow", "downarrow", "lus", "silon"] else f"\\{m.group(1)}", s
|
||||
)
|
||||
s = COMMANDS_TO_REMOVE_PATTERN.sub('', s)
|
||||
|
||||
# 应用所有替换
|
||||
for pattern, replacement in REPLACEMENTS_PATTERNS.items():
|
||||
s = pattern.sub(replacement, s)
|
||||
|
||||
# 处理LaTeX中的反斜杠和空格
|
||||
s = process_latex(s)
|
||||
|
||||
# \qquad后补空格
|
||||
s = QQUAD_PATTERN.sub(r'\\qquad ', s)
|
||||
|
||||
return s
|
||||
|
||||
|
||||
class UnimernetModel(VisionEncoderDecoderModel):
|
||||
def __init__(
|
||||
self,
|
||||
config: Optional[PretrainedConfig] = None,
|
||||
encoder: Optional[PreTrainedModel] = None,
|
||||
decoder: Optional[PreTrainedModel] = None,
|
||||
):
|
||||
# VisionEncoderDecoderModel's checking log has bug, disable for temp.
|
||||
base_model_logger.disabled = True
|
||||
try:
|
||||
super().__init__(config, encoder, decoder)
|
||||
finally:
|
||||
base_model_logger.disabled = False
|
||||
|
||||
if not config or not hasattr(config, "_name_or_path"):
|
||||
raise RuntimeError("config._name_or_path is required by UnimernetModel.")
|
||||
|
||||
model_path = config._name_or_path
|
||||
self.transform = UnimerSwinImageProcessor()
|
||||
self.tokenizer = TokenizerWrapper(AutoTokenizer.from_pretrained(model_path))
|
||||
self._post_check()
|
||||
|
||||
def _post_check(self):
|
||||
tokenizer = self.tokenizer
|
||||
|
||||
if tokenizer.tokenizer.model_max_length != self.config.decoder.max_position_embeddings:
|
||||
warnings.warn(
|
||||
f"decoder.max_position_embeddings={self.config.decoder.max_position_embeddings}," +
|
||||
f" but tokenizer.model_max_length={tokenizer.tokenizer.model_max_length}, will set" +
|
||||
f" tokenizer.model_max_length to {self.config.decoder.max_position_embeddings}.")
|
||||
tokenizer.tokenizer.model_max_length = self.config.decoder.max_position_embeddings
|
||||
|
||||
assert self.config.decoder.vocab_size == len(tokenizer)
|
||||
assert self.config.decoder_start_token_id == tokenizer.bos_token_id
|
||||
assert self.config.pad_token_id == tokenizer.pad_token_id
|
||||
|
||||
@classmethod
|
||||
def from_checkpoint(cls, model_path: str, model_filename: str = "pytorch_model.pth", state_dict_strip_prefix="model.model."):
|
||||
config = VisionEncoderDecoderConfig.from_pretrained(model_path)
|
||||
config._name_or_path = model_path
|
||||
config.encoder = UnimerSwinConfig(**vars(config.encoder))
|
||||
config.decoder = UnimerMBartConfig(**vars(config.decoder))
|
||||
|
||||
encoder = UnimerSwinModel(config.encoder)
|
||||
decoder = UnimerMBartForCausalLM(config.decoder)
|
||||
model = cls(config, encoder, decoder)
|
||||
|
||||
# load model weights
|
||||
model_file_path = os.path.join(model_path, model_filename)
|
||||
checkpoint = torch.load(model_file_path, map_location="cpu", weights_only=True)
|
||||
state_dict = checkpoint["model"] if "model" in checkpoint else checkpoint
|
||||
if not state_dict:
|
||||
raise RuntimeError("state_dict is empty.")
|
||||
if state_dict_strip_prefix:
|
||||
state_dict = {
|
||||
k[len(state_dict_strip_prefix):] if k.startswith(state_dict_strip_prefix) else k: v
|
||||
for k, v in state_dict.items()
|
||||
}
|
||||
missing_keys, unexpected_keys = model.load_state_dict(state_dict, strict=False)
|
||||
if len(unexpected_keys) > 0:
|
||||
warnings.warn("Unexpected key(s) in state_dict: {}.".format(", ".join(f'"{k}"' for k in unexpected_keys)))
|
||||
if len(missing_keys) > 0:
|
||||
raise RuntimeError("Missing key(s) in state_dict: {}.".format(", ".join(f'"{k}"' for k in missing_keys)))
|
||||
return model
|
||||
|
||||
def forward_bak(self, samples):
|
||||
pixel_values, text = samples["image"], samples["text_input"]
|
||||
|
||||
text_inputs = self.tokenizer.tokenize(text).to(pixel_values.device)
|
||||
decoder_input_ids, decoder_attention_mask = text_inputs["input_ids"], text_inputs["attention_mask"]
|
||||
|
||||
num_channels = pixel_values.shape[1]
|
||||
if num_channels == 1:
|
||||
pixel_values = pixel_values.repeat(1, 3, 1, 1)
|
||||
|
||||
labels = decoder_input_ids * 1
|
||||
labels = labels.masked_fill(labels == self.tokenizer.pad_token_id, -100)
|
||||
|
||||
loss = self.model(
|
||||
pixel_values=pixel_values,
|
||||
decoder_input_ids=decoder_input_ids[:, :-1],
|
||||
decoder_attention_mask=decoder_attention_mask[:, :-1],
|
||||
labels=labels[:, 1:],
|
||||
).loss
|
||||
return {"loss": loss}
|
||||
|
||||
def generate(self, samples, do_sample: bool = False, temperature: float = 0.2, top_p: float = 0.95):
|
||||
pixel_values = samples["image"]
|
||||
num_channels = pixel_values.shape[1]
|
||||
if num_channels == 1:
|
||||
pixel_values = pixel_values.repeat(1, 3, 1, 1)
|
||||
|
||||
kwargs = {}
|
||||
if do_sample:
|
||||
kwargs["temperature"] = temperature
|
||||
kwargs["top_p"] = top_p
|
||||
|
||||
outputs = super().generate(
|
||||
pixel_values=pixel_values,
|
||||
max_new_tokens=self.tokenizer.tokenizer.model_max_length, # required
|
||||
decoder_start_token_id=self.tokenizer.tokenizer.bos_token_id,
|
||||
do_sample=do_sample,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
outputs = outputs[:, 1:].cpu().numpy()
|
||||
pred_tokens = self.tokenizer.detokenize(outputs)
|
||||
pred_str = self.tokenizer.token2str(outputs)
|
||||
fixed_str = [latex_rm_whitespace(s) for s in pred_str]
|
||||
return {"pred_ids": outputs, "pred_tokens": pred_tokens, "pred_str": pred_str, "fixed_str": fixed_str}
|
||||
|
||||
@@ -0,0 +1,8 @@
|
||||
from .configuration_unimer_mbart import UnimerMBartConfig
|
||||
from .modeling_unimer_mbart import UnimerMBartModel, UnimerMBartForCausalLM
|
||||
|
||||
__all__ = [
|
||||
"UnimerMBartConfig",
|
||||
"UnimerMBartModel",
|
||||
"UnimerMBartForCausalLM",
|
||||
]
|
||||
@@ -0,0 +1,163 @@
|
||||
# coding=utf-8
|
||||
# Copyright 2021, The Facebook AI Research Team and The HuggingFace Inc. team. 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.
|
||||
"""UnimerMBART model configuration"""
|
||||
|
||||
from transformers.configuration_utils import PretrainedConfig
|
||||
from transformers.utils import logging
|
||||
|
||||
|
||||
logger = logging.get_logger(__name__)
|
||||
|
||||
|
||||
class UnimerMBartConfig(PretrainedConfig):
|
||||
r"""
|
||||
This is the configuration class to store the configuration of a [`MBartModel`]. It is used to instantiate an MBART
|
||||
model according to the specified arguments, defining the model architecture. Instantiating a configuration with the
|
||||
defaults will yield a similar configuration to that of the MBART
|
||||
[facebook/mbart-large-cc25](https://huggingface.co/facebook/mbart-large-cc25) architecture.
|
||||
|
||||
Configuration objects inherit from [`PretrainedConfig`] and can be used to control the model outputs. Read the
|
||||
documentation from [`PretrainedConfig`] for more information.
|
||||
|
||||
|
||||
Args:
|
||||
vocab_size (`int`, *optional*, defaults to 50265):
|
||||
Vocabulary size of the MBART model. Defines the number of different tokens that can be represented by the
|
||||
`inputs_ids` passed when calling [`MBartModel`] or [`TFMBartModel`].
|
||||
d_model (`int`, *optional*, defaults to 1024):
|
||||
Dimensionality of the layers and the pooler layer.
|
||||
qk_squeeze (`int`, *optional*, defaults to 2):
|
||||
Squeeze ratio for query/key's output dimension. See the [UniMERNet paper](https://arxiv.org/abs/2404.15254).
|
||||
Squeeze Attention maps the query and key to a lower-dimensional space without excessive loss of information,
|
||||
thereby accelerating the computation of attention.
|
||||
encoder_layers (`int`, *optional*, defaults to 12):
|
||||
Number of encoder layers.
|
||||
decoder_layers (`int`, *optional*, defaults to 12):
|
||||
Number of decoder layers.
|
||||
encoder_attention_heads (`int`, *optional*, defaults to 16):
|
||||
Number of attention heads for each attention layer in the Transformer encoder.
|
||||
decoder_attention_heads (`int`, *optional*, defaults to 16):
|
||||
Number of attention heads for each attention layer in the Transformer decoder.
|
||||
decoder_ffn_dim (`int`, *optional*, defaults to 4096):
|
||||
Dimensionality of the "intermediate" (often named feed-forward) layer in decoder.
|
||||
encoder_ffn_dim (`int`, *optional*, defaults to 4096):
|
||||
Dimensionality of the "intermediate" (often named feed-forward) layer in decoder.
|
||||
activation_function (`str` or `function`, *optional*, defaults to `"gelu"`):
|
||||
The non-linear activation function (function or string) in the encoder and pooler. If string, `"gelu"`,
|
||||
`"relu"`, `"silu"` and `"gelu_new"` are supported.
|
||||
dropout (`float`, *optional*, defaults to 0.1):
|
||||
The dropout probability for all fully connected layers in the embeddings, encoder, and pooler.
|
||||
attention_dropout (`float`, *optional*, defaults to 0.0):
|
||||
The dropout ratio for the attention probabilities.
|
||||
activation_dropout (`float`, *optional*, defaults to 0.0):
|
||||
The dropout ratio for activations inside the fully connected layer.
|
||||
classifier_dropout (`float`, *optional*, defaults to 0.0):
|
||||
The dropout ratio for classifier.
|
||||
max_position_embeddings (`int`, *optional*, defaults to 1024):
|
||||
The maximum sequence length that this model might ever be used with. Typically set this to something large
|
||||
just in case (e.g., 512 or 1024 or 2048).
|
||||
init_std (`float`, *optional*, defaults to 0.02):
|
||||
The standard deviation of the truncated_normal_initializer for initializing all weight matrices.
|
||||
encoder_layerdrop (`float`, *optional*, defaults to 0.0):
|
||||
The LayerDrop probability for the encoder. See the [LayerDrop paper](see https://arxiv.org/abs/1909.11556)
|
||||
for more details.
|
||||
decoder_layerdrop (`float`, *optional*, defaults to 0.0):
|
||||
The LayerDrop probability for the decoder. See the [LayerDrop paper](see https://arxiv.org/abs/1909.11556)
|
||||
for more details.
|
||||
scale_embedding (`bool`, *optional*, defaults to `False`):
|
||||
Scale embeddings by diving by sqrt(d_model).
|
||||
use_cache (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not the model should return the last key/values attentions (not used by all models)
|
||||
forced_eos_token_id (`int`, *optional*, defaults to 2):
|
||||
The id of the token to force as the last generated token when `max_length` is reached. Usually set to
|
||||
`eos_token_id`.
|
||||
|
||||
Example:
|
||||
|
||||
```python
|
||||
>>> from transformers import MBartConfig, MBartModel
|
||||
|
||||
>>> # Initializing a MBART facebook/mbart-large-cc25 style configuration
|
||||
>>> configuration = MBartConfig()
|
||||
|
||||
>>> # Initializing a model (with random weights) from the facebook/mbart-large-cc25 style configuration
|
||||
>>> model = MBartModel(configuration)
|
||||
|
||||
>>> # Accessing the model configuration
|
||||
>>> configuration = model.config
|
||||
```"""
|
||||
|
||||
model_type = "unimer-mbart"
|
||||
keys_to_ignore_at_inference = ["past_key_values"]
|
||||
attribute_map = {"num_attention_heads": "encoder_attention_heads", "hidden_size": "d_model"}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vocab_size=50265,
|
||||
max_position_embeddings=1024,
|
||||
encoder_layers=12,
|
||||
encoder_ffn_dim=4096,
|
||||
encoder_attention_heads=16,
|
||||
decoder_layers=12,
|
||||
decoder_ffn_dim=4096,
|
||||
decoder_attention_heads=16,
|
||||
encoder_layerdrop=0.0,
|
||||
decoder_layerdrop=0.0,
|
||||
use_cache=True,
|
||||
is_encoder_decoder=True,
|
||||
activation_function="gelu",
|
||||
d_model=1024,
|
||||
qk_squeeze=2,
|
||||
dropout=0.1,
|
||||
attention_dropout=0.0,
|
||||
activation_dropout=0.0,
|
||||
init_std=0.02,
|
||||
classifier_dropout=0.0,
|
||||
scale_embedding=False,
|
||||
pad_token_id=1,
|
||||
bos_token_id=0,
|
||||
eos_token_id=2,
|
||||
forced_eos_token_id=2,
|
||||
**kwargs,
|
||||
):
|
||||
self.vocab_size = vocab_size
|
||||
self.max_position_embeddings = max_position_embeddings
|
||||
self.d_model = d_model
|
||||
self.qk_squeeze = qk_squeeze
|
||||
self.encoder_ffn_dim = encoder_ffn_dim
|
||||
self.encoder_layers = encoder_layers
|
||||
self.encoder_attention_heads = encoder_attention_heads
|
||||
self.decoder_ffn_dim = decoder_ffn_dim
|
||||
self.decoder_layers = decoder_layers
|
||||
self.decoder_attention_heads = decoder_attention_heads
|
||||
self.dropout = dropout
|
||||
self.attention_dropout = attention_dropout
|
||||
self.activation_dropout = activation_dropout
|
||||
self.activation_function = activation_function
|
||||
self.init_std = init_std
|
||||
self.encoder_layerdrop = encoder_layerdrop
|
||||
self.decoder_layerdrop = decoder_layerdrop
|
||||
self.classifier_dropout = classifier_dropout
|
||||
self.use_cache = use_cache
|
||||
self.num_hidden_layers = encoder_layers
|
||||
self.scale_embedding = scale_embedding # scale factor will be sqrt(d_model) if True
|
||||
super().__init__(
|
||||
pad_token_id=pad_token_id,
|
||||
bos_token_id=bos_token_id,
|
||||
eos_token_id=eos_token_id,
|
||||
is_encoder_decoder=is_encoder_decoder,
|
||||
forced_eos_token_id=forced_eos_token_id,
|
||||
**kwargs,
|
||||
)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,9 @@
|
||||
from .configuration_unimer_swin import UnimerSwinConfig
|
||||
from .modeling_unimer_swin import UnimerSwinModel
|
||||
from .image_processing_unimer_swin import UnimerSwinImageProcessor
|
||||
|
||||
__all__ = [
|
||||
"UnimerSwinConfig",
|
||||
"UnimerSwinModel",
|
||||
"UnimerSwinImageProcessor",
|
||||
]
|
||||
@@ -0,0 +1,132 @@
|
||||
# coding=utf-8
|
||||
# Copyright 2022 The HuggingFace Inc. team. 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.
|
||||
"""Donut Swin Transformer model configuration"""
|
||||
|
||||
from transformers.configuration_utils import PretrainedConfig
|
||||
from transformers.utils import logging
|
||||
|
||||
|
||||
logger = logging.get_logger(__name__)
|
||||
|
||||
|
||||
class UnimerSwinConfig(PretrainedConfig):
|
||||
r"""
|
||||
This is the configuration class to store the configuration of a [`UnimerSwinModel`]. It is used to instantiate a
|
||||
Donut model according to the specified arguments, defining the model architecture. Instantiating a configuration
|
||||
with the defaults will yield a similar configuration to that of the Donut
|
||||
[naver-clova-ix/donut-base](https://huggingface.co/naver-clova-ix/donut-base) architecture.
|
||||
|
||||
Configuration objects inherit from [`PretrainedConfig`] and can be used to control the model outputs. Read the
|
||||
documentation from [`PretrainedConfig`] for more information.
|
||||
|
||||
Args:
|
||||
image_size (`int`, *optional*, defaults to 224):
|
||||
The size (resolution) of each image.
|
||||
patch_size (`int`, *optional*, defaults to 4):
|
||||
The size (resolution) of each patch.
|
||||
num_channels (`int`, *optional*, defaults to 3):
|
||||
The number of input channels.
|
||||
embed_dim (`int`, *optional*, defaults to 96):
|
||||
Dimensionality of patch embedding.
|
||||
depths (`list(int)`, *optional*, defaults to `[2, 2, 6, 2]`):
|
||||
Depth of each layer in the Transformer encoder.
|
||||
num_heads (`list(int)`, *optional*, defaults to `[3, 6, 12, 24]`):
|
||||
Number of attention heads in each layer of the Transformer encoder.
|
||||
window_size (`int`, *optional*, defaults to 7):
|
||||
Size of windows.
|
||||
mlp_ratio (`float`, *optional*, defaults to 4.0):
|
||||
Ratio of MLP hidden dimensionality to embedding dimensionality.
|
||||
qkv_bias (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not a learnable bias should be added to the queries, keys and values.
|
||||
hidden_dropout_prob (`float`, *optional*, defaults to 0.0):
|
||||
The dropout probability for all fully connected layers in the embeddings and encoder.
|
||||
attention_probs_dropout_prob (`float`, *optional*, defaults to 0.0):
|
||||
The dropout ratio for the attention probabilities.
|
||||
drop_path_rate (`float`, *optional*, defaults to 0.1):
|
||||
Stochastic depth rate.
|
||||
hidden_act (`str` or `function`, *optional*, defaults to `"gelu"`):
|
||||
The non-linear activation function (function or string) in the encoder. If string, `"gelu"`, `"relu"`,
|
||||
`"selu"` and `"gelu_new"` are supported.
|
||||
use_absolute_embeddings (`bool`, *optional*, defaults to `False`):
|
||||
Whether or not to add absolute position embeddings to the patch embeddings.
|
||||
initializer_range (`float`, *optional*, defaults to 0.02):
|
||||
The standard deviation of the truncated_normal_initializer for initializing all weight matrices.
|
||||
layer_norm_eps (`float`, *optional*, defaults to 1e-05):
|
||||
The epsilon used by the layer normalization layers.
|
||||
|
||||
Example:
|
||||
|
||||
```python
|
||||
>>> from transformers import UnimerSwinConfig, UnimerSwinModel
|
||||
|
||||
>>> # Initializing a Donut naver-clova-ix/donut-base style configuration
|
||||
>>> configuration = UnimerSwinConfig()
|
||||
|
||||
>>> # Randomly initializing a model from the naver-clova-ix/donut-base style configuration
|
||||
>>> model = UnimerSwinModel(configuration)
|
||||
|
||||
>>> # Accessing the model configuration
|
||||
>>> configuration = model.config
|
||||
```"""
|
||||
|
||||
model_type = "unimer-swin"
|
||||
|
||||
attribute_map = {
|
||||
"num_attention_heads": "num_heads",
|
||||
"num_hidden_layers": "num_layers",
|
||||
}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
image_size=224,
|
||||
patch_size=4,
|
||||
num_channels=3,
|
||||
embed_dim=96,
|
||||
depths=[2, 2, 6, 2],
|
||||
num_heads=[3, 6, 12, 24],
|
||||
window_size=7,
|
||||
mlp_ratio=4.0,
|
||||
qkv_bias=True,
|
||||
hidden_dropout_prob=0.0,
|
||||
attention_probs_dropout_prob=0.0,
|
||||
drop_path_rate=0.1,
|
||||
hidden_act="gelu",
|
||||
use_absolute_embeddings=False,
|
||||
initializer_range=0.02,
|
||||
layer_norm_eps=1e-5,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
|
||||
self.image_size = image_size
|
||||
self.patch_size = patch_size
|
||||
self.num_channels = num_channels
|
||||
self.embed_dim = embed_dim
|
||||
self.depths = depths
|
||||
self.num_layers = len(depths)
|
||||
self.num_heads = num_heads
|
||||
self.window_size = window_size
|
||||
self.mlp_ratio = mlp_ratio
|
||||
self.qkv_bias = qkv_bias
|
||||
self.hidden_dropout_prob = hidden_dropout_prob
|
||||
self.attention_probs_dropout_prob = attention_probs_dropout_prob
|
||||
self.drop_path_rate = drop_path_rate
|
||||
self.hidden_act = hidden_act
|
||||
self.use_absolute_embeddings = use_absolute_embeddings
|
||||
self.layer_norm_eps = layer_norm_eps
|
||||
self.initializer_range = initializer_range
|
||||
# we set the hidden_size attribute in order to make Swin work with VisionEncoderDecoderModel
|
||||
# this indicates the channel dimension after the last stage of the model
|
||||
self.hidden_size = int(embed_dim * 2 ** (len(depths) - 1))
|
||||
@@ -0,0 +1,132 @@
|
||||
from transformers.image_processing_utils import BaseImageProcessor
|
||||
import numpy as np
|
||||
import cv2
|
||||
import albumentations as alb
|
||||
from albumentations.pytorch import ToTensorV2
|
||||
|
||||
|
||||
# TODO: dereference cv2 if possible
|
||||
class UnimerSwinImageProcessor(BaseImageProcessor):
|
||||
def __init__(
|
||||
self,
|
||||
image_size = (192, 672),
|
||||
):
|
||||
self.input_size = [int(_) for _ in image_size]
|
||||
assert len(self.input_size) == 2
|
||||
|
||||
self.transform = alb.Compose(
|
||||
[
|
||||
alb.ToGray(),
|
||||
alb.Normalize((0.7931, 0.7931, 0.7931), (0.1738, 0.1738, 0.1738)),
|
||||
# alb.Sharpen()
|
||||
ToTensorV2(),
|
||||
]
|
||||
)
|
||||
|
||||
def __call__(self, item):
|
||||
image = self.prepare_input(item)
|
||||
return self.transform(image=image)['image'][:1]
|
||||
|
||||
@staticmethod
|
||||
def crop_margin_numpy(img: np.ndarray) -> np.ndarray:
|
||||
"""Crop margins of image using NumPy operations"""
|
||||
# Convert to grayscale if it's a color image
|
||||
if len(img.shape) == 3 and img.shape[2] == 3:
|
||||
gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)
|
||||
else:
|
||||
gray = img.copy()
|
||||
|
||||
# Normalize and threshold
|
||||
if gray.max() == gray.min():
|
||||
return img
|
||||
|
||||
normalized = (((gray - gray.min()) / (gray.max() - gray.min())) * 255).astype(np.uint8)
|
||||
binary = 255 * (normalized < 200).astype(np.uint8)
|
||||
|
||||
# Find bounding box
|
||||
coords = cv2.findNonZero(binary) # Find all non-zero points (text)
|
||||
x, y, w, h = cv2.boundingRect(coords) # Find minimum spanning bounding box
|
||||
|
||||
# Return cropped image
|
||||
return img[y:y + h, x:x + w]
|
||||
|
||||
def prepare_input(self, img, random_padding: bool = False):
|
||||
"""
|
||||
Convert PIL Image or numpy array to properly sized and padded image after:
|
||||
- crop margins
|
||||
- resize while maintaining aspect ratio
|
||||
- pad to target size
|
||||
"""
|
||||
if img is None:
|
||||
return None
|
||||
|
||||
# try:
|
||||
# img = self.crop_margin_numpy(img)
|
||||
# except Exception:
|
||||
# # might throw an error for broken files
|
||||
# return None
|
||||
|
||||
if img.shape[0] == 0 or img.shape[1] == 0:
|
||||
return None
|
||||
|
||||
# Get current dimensions
|
||||
h, w = img.shape[:2]
|
||||
target_h, target_w = self.input_size
|
||||
|
||||
# Calculate scale to preserve aspect ratio (equivalent to resize + thumbnail)
|
||||
scale = min(target_h / h, target_w / w)
|
||||
|
||||
# Calculate new dimensions
|
||||
new_h, new_w = int(h * scale), int(w * scale)
|
||||
|
||||
# Resize the image while preserving aspect ratio
|
||||
resized_img = cv2.resize(img, (new_w, new_h))
|
||||
|
||||
# Calculate padding values using the existing method
|
||||
delta_width = target_w - new_w
|
||||
delta_height = target_h - new_h
|
||||
|
||||
pad_width, pad_height = self._get_padding_values(new_w, new_h, random_padding)
|
||||
|
||||
# Apply padding (convert PIL padding format to OpenCV format)
|
||||
padding_color = [0, 0, 0] if len(img.shape) == 3 else [0]
|
||||
|
||||
padded_img = cv2.copyMakeBorder(
|
||||
resized_img,
|
||||
pad_height, # top
|
||||
delta_height - pad_height, # bottom
|
||||
pad_width, # left
|
||||
delta_width - pad_width, # right
|
||||
cv2.BORDER_CONSTANT,
|
||||
value=padding_color
|
||||
)
|
||||
|
||||
return padded_img
|
||||
|
||||
def _calculate_padding(self, new_w, new_h, random_padding):
|
||||
"""Calculate padding values for PIL images"""
|
||||
delta_width = self.input_size[1] - new_w
|
||||
delta_height = self.input_size[0] - new_h
|
||||
|
||||
pad_width, pad_height = self._get_padding_values(new_w, new_h, random_padding)
|
||||
|
||||
return (
|
||||
pad_width,
|
||||
pad_height,
|
||||
delta_width - pad_width,
|
||||
delta_height - pad_height,
|
||||
)
|
||||
|
||||
def _get_padding_values(self, new_w, new_h, random_padding):
|
||||
"""Get padding values based on image dimensions and padding strategy"""
|
||||
delta_width = self.input_size[1] - new_w
|
||||
delta_height = self.input_size[0] - new_h
|
||||
|
||||
if random_padding:
|
||||
pad_width = np.random.randint(low=0, high=delta_width + 1)
|
||||
pad_height = np.random.randint(low=0, high=delta_height + 1)
|
||||
else:
|
||||
pad_width = delta_width // 2
|
||||
pad_height = delta_height // 2
|
||||
|
||||
return pad_width, pad_height
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1 @@
|
||||
# Copyright (c) Opendatalab. All rights reserved.
|
||||
@@ -0,0 +1 @@
|
||||
# Copyright (c) Opendatalab. All rights reserved.
|
||||
@@ -0,0 +1,199 @@
|
||||
# Copyright (c) Opendatalab. All rights reserved.
|
||||
import copy
|
||||
import os.path
|
||||
import warnings
|
||||
from pathlib import Path
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import yaml
|
||||
from loguru import logger
|
||||
|
||||
from magic_pdf.libs.config_reader import get_device, get_local_models_dir
|
||||
from ....utils.ocr_utils import check_img, preprocess_image, sorted_boxes, merge_det_boxes, update_det_boxes, get_rotate_crop_image
|
||||
from .tools.infer.predict_system import TextSystem
|
||||
from .tools.infer import pytorchocr_utility as utility
|
||||
import argparse
|
||||
|
||||
|
||||
latin_lang = [
|
||||
'af', 'az', 'bs', 'cs', 'cy', 'da', 'de', 'es', 'et', 'fr', 'ga', 'hr', # noqa: E126
|
||||
'hu', 'id', 'is', 'it', 'ku', 'la', 'lt', 'lv', 'mi', 'ms', 'mt', 'nl',
|
||||
'no', 'oc', 'pi', 'pl', 'pt', 'ro', 'rs_latin', 'sk', 'sl', 'sq', 'sv',
|
||||
'sw', 'tl', 'tr', 'uz', 'vi', 'french', 'german'
|
||||
]
|
||||
arabic_lang = ['ar', 'fa', 'ug', 'ur']
|
||||
cyrillic_lang = [
|
||||
'ru', 'rs_cyrillic', 'be', 'bg', 'uk', 'mn', 'abq', 'ady', 'kbd', 'ava', # noqa: E126
|
||||
'dar', 'inh', 'che', 'lbe', 'lez', 'tab'
|
||||
]
|
||||
devanagari_lang = [
|
||||
'hi', 'mr', 'ne', 'bh', 'mai', 'ang', 'bho', 'mah', 'sck', 'new', 'gom', # noqa: E126
|
||||
'sa', 'bgc'
|
||||
]
|
||||
|
||||
|
||||
def get_model_params(lang, config):
|
||||
if lang in config['lang']:
|
||||
params = config['lang'][lang]
|
||||
det = params.get('det')
|
||||
rec = params.get('rec')
|
||||
dict_file = params.get('dict')
|
||||
return det, rec, dict_file
|
||||
else:
|
||||
raise Exception (f'Language {lang} not supported')
|
||||
|
||||
|
||||
root_dir = Path(__file__).resolve().parent
|
||||
|
||||
|
||||
class PytorchPaddleOCR(TextSystem):
|
||||
def __init__(self, *args, **kwargs):
|
||||
parser = utility.init_args()
|
||||
args = parser.parse_args(args)
|
||||
|
||||
self.lang = kwargs.get('lang', 'ch')
|
||||
|
||||
device = get_device()
|
||||
if device == 'cpu' and self.lang in ['ch', 'ch_server']:
|
||||
logger.warning("The current device in use is CPU. To ensure the speed of parsing, the language is automatically switched to ch_lite.")
|
||||
self.lang = 'ch_lite'
|
||||
|
||||
if self.lang in latin_lang:
|
||||
self.lang = 'latin'
|
||||
elif self.lang in arabic_lang:
|
||||
self.lang = 'arabic'
|
||||
elif self.lang in cyrillic_lang:
|
||||
self.lang = 'cyrillic'
|
||||
elif self.lang in devanagari_lang:
|
||||
self.lang = 'devanagari'
|
||||
else:
|
||||
pass
|
||||
|
||||
models_config_path = os.path.join(root_dir, 'pytorchocr', 'utils', 'resources', 'models_config.yml')
|
||||
with open(models_config_path) as file:
|
||||
config = yaml.safe_load(file)
|
||||
det, rec, dict_file = get_model_params(self.lang, config)
|
||||
ocr_models_dir = os.path.join(get_local_models_dir(), 'OCR', 'paddleocr_torch')
|
||||
kwargs['det_model_path'] = os.path.join(ocr_models_dir, det)
|
||||
kwargs['rec_model_path'] = os.path.join(ocr_models_dir, rec)
|
||||
kwargs['rec_char_dict_path'] = os.path.join(root_dir, 'pytorchocr', 'utils', 'resources', 'dict', dict_file)
|
||||
# kwargs['rec_batch_num'] = 8
|
||||
|
||||
kwargs['device'] = device
|
||||
|
||||
default_args = vars(args)
|
||||
default_args.update(kwargs)
|
||||
args = argparse.Namespace(**default_args)
|
||||
|
||||
super().__init__(args)
|
||||
|
||||
def ocr(self,
|
||||
img,
|
||||
det=True,
|
||||
rec=True,
|
||||
mfd_res=None,
|
||||
tqdm_enable=False,
|
||||
):
|
||||
assert isinstance(img, (np.ndarray, list, str, bytes))
|
||||
if isinstance(img, list) and det == True:
|
||||
logger.error('When input a list of images, det must be false')
|
||||
exit(0)
|
||||
img = check_img(img)
|
||||
imgs = [img]
|
||||
with warnings.catch_warnings():
|
||||
warnings.simplefilter("ignore", category=RuntimeWarning)
|
||||
if det and rec:
|
||||
ocr_res = []
|
||||
for img in imgs:
|
||||
img = preprocess_image(img)
|
||||
dt_boxes, rec_res = self.__call__(img, mfd_res=mfd_res)
|
||||
if not dt_boxes and not rec_res:
|
||||
ocr_res.append(None)
|
||||
continue
|
||||
tmp_res = [[box.tolist(), res] for box, res in zip(dt_boxes, rec_res)]
|
||||
ocr_res.append(tmp_res)
|
||||
return ocr_res
|
||||
elif det and not rec:
|
||||
ocr_res = []
|
||||
for img in imgs:
|
||||
img = preprocess_image(img)
|
||||
dt_boxes, elapse = self.text_detector(img)
|
||||
# logger.debug("dt_boxes num : {}, elapsed : {}".format(len(dt_boxes), elapse))
|
||||
if dt_boxes is None:
|
||||
ocr_res.append(None)
|
||||
continue
|
||||
dt_boxes = sorted_boxes(dt_boxes)
|
||||
# merge_det_boxes 和 update_det_boxes 都会把poly转成bbox再转回poly,因此需要过滤所有倾斜程度较大的文本框
|
||||
dt_boxes = merge_det_boxes(dt_boxes)
|
||||
if mfd_res:
|
||||
dt_boxes = update_det_boxes(dt_boxes, mfd_res)
|
||||
tmp_res = [box.tolist() for box in dt_boxes]
|
||||
ocr_res.append(tmp_res)
|
||||
return ocr_res
|
||||
elif not det and rec:
|
||||
ocr_res = []
|
||||
for img in imgs:
|
||||
if not isinstance(img, list):
|
||||
img = preprocess_image(img)
|
||||
img = [img]
|
||||
rec_res, elapse = self.text_recognizer(img, tqdm_enable=tqdm_enable)
|
||||
# logger.debug("rec_res num : {}, elapsed : {}".format(len(rec_res), elapse))
|
||||
ocr_res.append(rec_res)
|
||||
return ocr_res
|
||||
|
||||
def __call__(self, img, mfd_res=None):
|
||||
|
||||
if img is None:
|
||||
logger.debug("no valid image provided")
|
||||
return None, None
|
||||
|
||||
ori_im = img.copy()
|
||||
dt_boxes, elapse = self.text_detector(img)
|
||||
|
||||
if dt_boxes is None:
|
||||
logger.debug("no dt_boxes found, elapsed : {}".format(elapse))
|
||||
return None, None
|
||||
else:
|
||||
pass
|
||||
# logger.debug("dt_boxes num : {}, elapsed : {}".format(len(dt_boxes), elapse))
|
||||
img_crop_list = []
|
||||
|
||||
dt_boxes = sorted_boxes(dt_boxes)
|
||||
|
||||
# merge_det_boxes 和 update_det_boxes 都会把poly转成bbox再转回poly,因此需要过滤所有倾斜程度较大的文本框
|
||||
dt_boxes = merge_det_boxes(dt_boxes)
|
||||
|
||||
if mfd_res:
|
||||
dt_boxes = update_det_boxes(dt_boxes, mfd_res)
|
||||
|
||||
for bno in range(len(dt_boxes)):
|
||||
tmp_box = copy.deepcopy(dt_boxes[bno])
|
||||
img_crop = get_rotate_crop_image(ori_im, tmp_box)
|
||||
img_crop_list.append(img_crop)
|
||||
|
||||
rec_res, elapse = self.text_recognizer(img_crop_list)
|
||||
# logger.debug("rec_res num : {}, elapsed : {}".format(len(rec_res), elapse))
|
||||
|
||||
filter_boxes, filter_rec_res = [], []
|
||||
for box, rec_result in zip(dt_boxes, rec_res):
|
||||
text, score = rec_result
|
||||
if score >= self.drop_score:
|
||||
filter_boxes.append(box)
|
||||
filter_rec_res.append(rec_result)
|
||||
|
||||
return filter_boxes, filter_rec_res
|
||||
|
||||
if __name__ == '__main__':
|
||||
pytorch_paddle_ocr = PytorchPaddleOCR()
|
||||
img = cv2.imread("/Users/myhloli/Downloads/screenshot-20250326-194348.png")
|
||||
dt_boxes, rec_res = pytorch_paddle_ocr(img)
|
||||
ocr_res = []
|
||||
if not dt_boxes and not rec_res:
|
||||
ocr_res.append(None)
|
||||
else:
|
||||
tmp_res = [[box.tolist(), res] for box, res in zip(dt_boxes, rec_res)]
|
||||
ocr_res.append(tmp_res)
|
||||
print(ocr_res)
|
||||
|
||||
|
||||
@@ -0,0 +1,39 @@
|
||||
import os
|
||||
import torch
|
||||
from .modeling.architectures.base_model import BaseModel
|
||||
|
||||
class BaseOCRV20:
|
||||
def __init__(self, config, **kwargs):
|
||||
self.config = config
|
||||
self.build_net(**kwargs)
|
||||
self.net.eval()
|
||||
|
||||
|
||||
def build_net(self, **kwargs):
|
||||
self.net = BaseModel(self.config, **kwargs)
|
||||
|
||||
def read_pytorch_weights(self, weights_path):
|
||||
if not os.path.exists(weights_path):
|
||||
raise FileNotFoundError('{} is not existed.'.format(weights_path))
|
||||
weights = torch.load(weights_path)
|
||||
return weights
|
||||
|
||||
def get_out_channels(self, weights):
|
||||
if list(weights.keys())[-1].endswith('.weight') and len(list(weights.values())[-1].shape) == 2:
|
||||
out_channels = list(weights.values())[-1].numpy().shape[1]
|
||||
else:
|
||||
out_channels = list(weights.values())[-1].numpy().shape[0]
|
||||
return out_channels
|
||||
|
||||
def load_state_dict(self, weights):
|
||||
self.net.load_state_dict(weights)
|
||||
# print('weights is loaded.')
|
||||
|
||||
def load_pytorch_weights(self, weights_path):
|
||||
self.net.load_state_dict(torch.load(weights_path, weights_only=True))
|
||||
# print('model is loaded: {}'.format(weights_path))
|
||||
|
||||
def inference(self, inputs):
|
||||
with torch.no_grad():
|
||||
infer = self.net(inputs)
|
||||
return infer
|
||||
@@ -0,0 +1,8 @@
|
||||
from __future__ import absolute_import
|
||||
from __future__ import division
|
||||
from __future__ import print_function
|
||||
from __future__ import unicode_literals
|
||||
|
||||
from .imaug import transform, create_operators
|
||||
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
from __future__ import absolute_import
|
||||
from __future__ import division
|
||||
from __future__ import print_function
|
||||
from __future__ import unicode_literals
|
||||
|
||||
# from .iaa_augment import IaaAugment
|
||||
# from .make_border_map import MakeBorderMap
|
||||
# from .make_shrink_map import MakeShrinkMap
|
||||
# from .random_crop_data import EastRandomCropData, PSERandomCrop
|
||||
|
||||
# from .rec_img_aug import RecAug, RecResizeImg, ClsResizeImg
|
||||
# from .randaugment import RandAugment
|
||||
from .operators import *
|
||||
# from .label_ops import *
|
||||
|
||||
# from .east_process import *
|
||||
# from .sast_process import *
|
||||
# from .gen_table_mask import *
|
||||
|
||||
def transform(data, ops=None):
|
||||
""" transform """
|
||||
if ops is None:
|
||||
ops = []
|
||||
for op in ops:
|
||||
data = op(data)
|
||||
if data is None:
|
||||
return None
|
||||
return data
|
||||
|
||||
|
||||
def create_operators(op_param_list, global_config=None):
|
||||
"""
|
||||
create operators based on the config
|
||||
Args:
|
||||
params(list): a dict list, used to create some operators
|
||||
"""
|
||||
assert isinstance(op_param_list, list), ('operator config should be a list')
|
||||
ops = []
|
||||
for operator in op_param_list:
|
||||
assert isinstance(operator,
|
||||
dict) and len(operator) == 1, "yaml format error"
|
||||
op_name = list(operator)[0]
|
||||
param = {} if operator[op_name] is None else operator[op_name]
|
||||
if global_config is not None:
|
||||
param.update(global_config)
|
||||
op = eval(op_name)(**param)
|
||||
ops.append(op)
|
||||
return ops
|
||||
@@ -0,0 +1,418 @@
|
||||
"""
|
||||
# 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
|
||||
from __future__ import unicode_literals
|
||||
|
||||
import sys
|
||||
import six
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
|
||||
class DecodeImage(object):
|
||||
""" decode image """
|
||||
|
||||
def __init__(self, img_mode='RGB', channel_first=False, **kwargs):
|
||||
self.img_mode = img_mode
|
||||
self.channel_first = channel_first
|
||||
|
||||
def __call__(self, data):
|
||||
img = data['image']
|
||||
if six.PY2:
|
||||
assert type(img) is str and len(
|
||||
img) > 0, "invalid input 'img' in DecodeImage"
|
||||
else:
|
||||
assert type(img) is bytes and len(
|
||||
img) > 0, "invalid input 'img' in DecodeImage"
|
||||
img = np.frombuffer(img, dtype='uint8')
|
||||
img = cv2.imdecode(img, 1)
|
||||
if img is None:
|
||||
return None
|
||||
if self.img_mode == 'GRAY':
|
||||
img = cv2.cvtColor(img, cv2.COLOR_GRAY2BGR)
|
||||
elif self.img_mode == 'RGB':
|
||||
assert img.shape[2] == 3, 'invalid shape of image[%s]' % (img.shape)
|
||||
img = img[:, :, ::-1]
|
||||
|
||||
if self.channel_first:
|
||||
img = img.transpose((2, 0, 1))
|
||||
|
||||
data['image'] = img
|
||||
return data
|
||||
|
||||
|
||||
class NRTRDecodeImage(object):
|
||||
""" decode image """
|
||||
|
||||
def __init__(self, img_mode='RGB', channel_first=False, **kwargs):
|
||||
self.img_mode = img_mode
|
||||
self.channel_first = channel_first
|
||||
|
||||
def __call__(self, data):
|
||||
img = data['image']
|
||||
if six.PY2:
|
||||
assert type(img) is str and len(
|
||||
img) > 0, "invalid input 'img' in DecodeImage"
|
||||
else:
|
||||
assert type(img) is bytes and len(
|
||||
img) > 0, "invalid input 'img' in DecodeImage"
|
||||
img = np.frombuffer(img, dtype='uint8')
|
||||
|
||||
img = cv2.imdecode(img, 1)
|
||||
|
||||
if img is None:
|
||||
return None
|
||||
if self.img_mode == 'GRAY':
|
||||
img = cv2.cvtColor(img, cv2.COLOR_GRAY2BGR)
|
||||
elif self.img_mode == 'RGB':
|
||||
assert img.shape[2] == 3, 'invalid shape of image[%s]' % (img.shape)
|
||||
img = img[:, :, ::-1]
|
||||
img = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)
|
||||
if self.channel_first:
|
||||
img = img.transpose((2, 0, 1))
|
||||
data['image'] = img
|
||||
return data
|
||||
|
||||
|
||||
class NormalizeImage(object):
|
||||
""" normalize image such as substract mean, divide std
|
||||
"""
|
||||
|
||||
def __init__(self, scale=None, mean=None, std=None, order='chw', **kwargs):
|
||||
if isinstance(scale, str):
|
||||
scale = eval(scale)
|
||||
self.scale = np.float32(scale if scale is not None else 1.0 / 255.0)
|
||||
mean = mean if mean is not None else [0.485, 0.456, 0.406]
|
||||
std = std if std is not None else [0.229, 0.224, 0.225]
|
||||
|
||||
shape = (3, 1, 1) if order == 'chw' else (1, 1, 3)
|
||||
self.mean = np.array(mean).reshape(shape).astype('float32')
|
||||
self.std = np.array(std).reshape(shape).astype('float32')
|
||||
|
||||
def __call__(self, data):
|
||||
img = data['image']
|
||||
from PIL import Image
|
||||
if isinstance(img, Image.Image):
|
||||
img = np.array(img)
|
||||
assert isinstance(img,
|
||||
np.ndarray), "invalid input 'img' in NormalizeImage"
|
||||
data['image'] = (
|
||||
img.astype('float32') * self.scale - self.mean) / self.std
|
||||
return data
|
||||
|
||||
|
||||
class ToCHWImage(object):
|
||||
""" convert hwc image to chw image
|
||||
"""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
pass
|
||||
|
||||
def __call__(self, data):
|
||||
img = data['image']
|
||||
from PIL import Image
|
||||
if isinstance(img, Image.Image):
|
||||
img = np.array(img)
|
||||
data['image'] = img.transpose((2, 0, 1))
|
||||
return data
|
||||
|
||||
|
||||
class Fasttext(object):
|
||||
def __init__(self, path="None", **kwargs):
|
||||
import fasttext
|
||||
self.fast_model = fasttext.load_model(path)
|
||||
|
||||
def __call__(self, data):
|
||||
label = data['label']
|
||||
fast_label = self.fast_model[label]
|
||||
data['fast_label'] = fast_label
|
||||
return data
|
||||
|
||||
|
||||
class KeepKeys(object):
|
||||
def __init__(self, keep_keys, **kwargs):
|
||||
self.keep_keys = keep_keys
|
||||
|
||||
def __call__(self, data):
|
||||
data_list = []
|
||||
for key in self.keep_keys:
|
||||
data_list.append(data[key])
|
||||
return data_list
|
||||
|
||||
|
||||
class Resize(object):
|
||||
def __init__(self, size=(640, 640), **kwargs):
|
||||
self.size = size
|
||||
|
||||
def resize_image(self, img):
|
||||
resize_h, resize_w = self.size
|
||||
ori_h, ori_w = img.shape[:2] # (h, w, c)
|
||||
ratio_h = float(resize_h) / ori_h
|
||||
ratio_w = float(resize_w) / ori_w
|
||||
img = cv2.resize(img, (int(resize_w), int(resize_h)))
|
||||
return img, [ratio_h, ratio_w]
|
||||
|
||||
def __call__(self, data):
|
||||
img = data['image']
|
||||
text_polys = data['polys']
|
||||
|
||||
img_resize, [ratio_h, ratio_w] = self.resize_image(img)
|
||||
new_boxes = []
|
||||
for box in text_polys:
|
||||
new_box = []
|
||||
for cord in box:
|
||||
new_box.append([cord[0] * ratio_w, cord[1] * ratio_h])
|
||||
new_boxes.append(new_box)
|
||||
data['image'] = img_resize
|
||||
data['polys'] = np.array(new_boxes, dtype=np.float32)
|
||||
return data
|
||||
|
||||
|
||||
class DetResizeForTest(object):
|
||||
def __init__(self, **kwargs):
|
||||
super(DetResizeForTest, self).__init__()
|
||||
self.resize_type = 0
|
||||
if 'image_shape' in kwargs:
|
||||
self.image_shape = kwargs['image_shape']
|
||||
self.resize_type = 1
|
||||
elif 'limit_side_len' in kwargs:
|
||||
self.limit_side_len = kwargs['limit_side_len']
|
||||
self.limit_type = kwargs.get('limit_type', 'min')
|
||||
elif 'resize_long' in kwargs:
|
||||
self.resize_type = 2
|
||||
self.resize_long = kwargs.get('resize_long', 960)
|
||||
else:
|
||||
self.limit_side_len = 736
|
||||
self.limit_type = 'min'
|
||||
|
||||
def __call__(self, data):
|
||||
img = data['image']
|
||||
src_h, src_w, _ = img.shape
|
||||
|
||||
if self.resize_type == 0:
|
||||
# img, shape = self.resize_image_type0(img)
|
||||
img, [ratio_h, ratio_w] = self.resize_image_type0(img)
|
||||
elif self.resize_type == 2:
|
||||
img, [ratio_h, ratio_w] = self.resize_image_type2(img)
|
||||
else:
|
||||
# img, shape = self.resize_image_type1(img)
|
||||
img, [ratio_h, ratio_w] = self.resize_image_type1(img)
|
||||
data['image'] = img
|
||||
data['shape'] = np.array([src_h, src_w, ratio_h, ratio_w])
|
||||
return data
|
||||
|
||||
def resize_image_type1(self, img):
|
||||
resize_h, resize_w = self.image_shape
|
||||
ori_h, ori_w = img.shape[:2] # (h, w, c)
|
||||
ratio_h = float(resize_h) / ori_h
|
||||
ratio_w = float(resize_w) / ori_w
|
||||
img = cv2.resize(img, (int(resize_w), int(resize_h)))
|
||||
# return img, np.array([ori_h, ori_w])
|
||||
return img, [ratio_h, ratio_w]
|
||||
|
||||
def resize_image_type0(self, img):
|
||||
"""
|
||||
resize image to a size multiple of 32 which is required by the network
|
||||
args:
|
||||
img(array): array with shape [h, w, c]
|
||||
return(tuple):
|
||||
img, (ratio_h, ratio_w)
|
||||
"""
|
||||
limit_side_len = self.limit_side_len
|
||||
h, w, c = img.shape
|
||||
|
||||
# limit the max side
|
||||
if self.limit_type == 'max':
|
||||
if max(h, w) > limit_side_len:
|
||||
if h > w:
|
||||
ratio = float(limit_side_len) / h
|
||||
else:
|
||||
ratio = float(limit_side_len) / w
|
||||
else:
|
||||
ratio = 1.
|
||||
elif self.limit_type == 'min':
|
||||
if min(h, w) < limit_side_len:
|
||||
if h < w:
|
||||
ratio = float(limit_side_len) / h
|
||||
else:
|
||||
ratio = float(limit_side_len) / w
|
||||
else:
|
||||
ratio = 1.
|
||||
elif self.limit_type == 'resize_long':
|
||||
ratio = float(limit_side_len) / max(h, w)
|
||||
else:
|
||||
raise Exception('not support limit type, image ')
|
||||
resize_h = int(h * ratio)
|
||||
resize_w = int(w * ratio)
|
||||
|
||||
resize_h = max(int(round(resize_h / 32) * 32), 32)
|
||||
resize_w = max(int(round(resize_w / 32) * 32), 32)
|
||||
|
||||
try:
|
||||
if int(resize_w) <= 0 or int(resize_h) <= 0:
|
||||
return None, (None, None)
|
||||
img = cv2.resize(img, (int(resize_w), int(resize_h)))
|
||||
except:
|
||||
print(img.shape, resize_w, resize_h)
|
||||
sys.exit(0)
|
||||
ratio_h = resize_h / float(h)
|
||||
ratio_w = resize_w / float(w)
|
||||
return img, [ratio_h, ratio_w]
|
||||
|
||||
def resize_image_type2(self, img):
|
||||
h, w, _ = img.shape
|
||||
|
||||
resize_w = w
|
||||
resize_h = h
|
||||
|
||||
if resize_h > resize_w:
|
||||
ratio = float(self.resize_long) / resize_h
|
||||
else:
|
||||
ratio = float(self.resize_long) / resize_w
|
||||
|
||||
resize_h = int(resize_h * ratio)
|
||||
resize_w = int(resize_w * ratio)
|
||||
|
||||
max_stride = 128
|
||||
resize_h = (resize_h + max_stride - 1) // max_stride * max_stride
|
||||
resize_w = (resize_w + max_stride - 1) // max_stride * max_stride
|
||||
img = cv2.resize(img, (int(resize_w), int(resize_h)))
|
||||
ratio_h = resize_h / float(h)
|
||||
ratio_w = resize_w / float(w)
|
||||
|
||||
return img, [ratio_h, ratio_w]
|
||||
|
||||
|
||||
class E2EResizeForTest(object):
|
||||
def __init__(self, **kwargs):
|
||||
super(E2EResizeForTest, self).__init__()
|
||||
self.max_side_len = kwargs['max_side_len']
|
||||
self.valid_set = kwargs['valid_set']
|
||||
|
||||
def __call__(self, data):
|
||||
img = data['image']
|
||||
src_h, src_w, _ = img.shape
|
||||
if self.valid_set == 'totaltext':
|
||||
im_resized, [ratio_h, ratio_w] = self.resize_image_for_totaltext(
|
||||
img, max_side_len=self.max_side_len)
|
||||
else:
|
||||
im_resized, (ratio_h, ratio_w) = self.resize_image(
|
||||
img, max_side_len=self.max_side_len)
|
||||
data['image'] = im_resized
|
||||
data['shape'] = np.array([src_h, src_w, ratio_h, ratio_w])
|
||||
return data
|
||||
|
||||
def resize_image_for_totaltext(self, im, max_side_len=512):
|
||||
|
||||
h, w, _ = im.shape
|
||||
resize_w = w
|
||||
resize_h = h
|
||||
ratio = 1.25
|
||||
if h * ratio > max_side_len:
|
||||
ratio = float(max_side_len) / resize_h
|
||||
resize_h = int(resize_h * ratio)
|
||||
resize_w = int(resize_w * ratio)
|
||||
|
||||
max_stride = 128
|
||||
resize_h = (resize_h + max_stride - 1) // max_stride * max_stride
|
||||
resize_w = (resize_w + max_stride - 1) // max_stride * max_stride
|
||||
im = cv2.resize(im, (int(resize_w), int(resize_h)))
|
||||
ratio_h = resize_h / float(h)
|
||||
ratio_w = resize_w / float(w)
|
||||
return im, (ratio_h, ratio_w)
|
||||
|
||||
def resize_image(self, im, max_side_len=512):
|
||||
"""
|
||||
resize image to a size multiple of max_stride which is required by the network
|
||||
:param im: the resized image
|
||||
:param max_side_len: limit of max image size to avoid out of memory in gpu
|
||||
:return: the resized image and the resize ratio
|
||||
"""
|
||||
h, w, _ = im.shape
|
||||
|
||||
resize_w = w
|
||||
resize_h = h
|
||||
|
||||
# Fix the longer side
|
||||
if resize_h > resize_w:
|
||||
ratio = float(max_side_len) / resize_h
|
||||
else:
|
||||
ratio = float(max_side_len) / resize_w
|
||||
|
||||
resize_h = int(resize_h * ratio)
|
||||
resize_w = int(resize_w * ratio)
|
||||
|
||||
max_stride = 128
|
||||
resize_h = (resize_h + max_stride - 1) // max_stride * max_stride
|
||||
resize_w = (resize_w + max_stride - 1) // max_stride * max_stride
|
||||
im = cv2.resize(im, (int(resize_w), int(resize_h)))
|
||||
ratio_h = resize_h / float(h)
|
||||
ratio_w = resize_w / float(w)
|
||||
|
||||
return im, (ratio_h, ratio_w)
|
||||
|
||||
|
||||
class KieResize(object):
|
||||
def __init__(self, **kwargs):
|
||||
super(KieResize, self).__init__()
|
||||
self.max_side, self.min_side = kwargs['img_scale'][0], kwargs[
|
||||
'img_scale'][1]
|
||||
|
||||
def __call__(self, data):
|
||||
img = data['image']
|
||||
points = data['points']
|
||||
src_h, src_w, _ = img.shape
|
||||
im_resized, scale_factor, [ratio_h, ratio_w
|
||||
], [new_h, new_w] = self.resize_image(img)
|
||||
resize_points = self.resize_boxes(img, points, scale_factor)
|
||||
data['ori_image'] = img
|
||||
data['ori_boxes'] = points
|
||||
data['points'] = resize_points
|
||||
data['image'] = im_resized
|
||||
data['shape'] = np.array([new_h, new_w])
|
||||
return data
|
||||
|
||||
def resize_image(self, img):
|
||||
norm_img = np.zeros([1024, 1024, 3], dtype='float32')
|
||||
scale = [512, 1024]
|
||||
h, w = img.shape[:2]
|
||||
max_long_edge = max(scale)
|
||||
max_short_edge = min(scale)
|
||||
scale_factor = min(max_long_edge / max(h, w),
|
||||
max_short_edge / min(h, w))
|
||||
resize_w, resize_h = int(w * float(scale_factor) + 0.5), int(h * float(
|
||||
scale_factor) + 0.5)
|
||||
max_stride = 32
|
||||
resize_h = (resize_h + max_stride - 1) // max_stride * max_stride
|
||||
resize_w = (resize_w + max_stride - 1) // max_stride * max_stride
|
||||
im = cv2.resize(img, (resize_w, resize_h))
|
||||
new_h, new_w = im.shape[:2]
|
||||
w_scale = new_w / w
|
||||
h_scale = new_h / h
|
||||
scale_factor = np.array(
|
||||
[w_scale, h_scale, w_scale, h_scale], dtype=np.float32)
|
||||
norm_img[:new_h, :new_w, :] = im
|
||||
return norm_img, scale_factor, [h_scale, w_scale], [new_h, new_w]
|
||||
|
||||
def resize_boxes(self, im, points, scale_factor):
|
||||
points = points * scale_factor
|
||||
img_shape = im.shape[:2]
|
||||
points[:, 0::2] = np.clip(points[:, 0::2], 0, img_shape[1])
|
||||
points[:, 1::2] = np.clip(points[:, 1::2], 0, img_shape[0])
|
||||
return points
|
||||
@@ -0,0 +1,25 @@
|
||||
# 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 copy
|
||||
|
||||
__all__ = ["build_model"]
|
||||
|
||||
|
||||
def build_model(config, **kwargs):
|
||||
from .base_model import BaseModel
|
||||
|
||||
config = copy.deepcopy(config)
|
||||
module_class = BaseModel(config, **kwargs)
|
||||
return module_class
|
||||
@@ -0,0 +1,105 @@
|
||||
from torch import nn
|
||||
|
||||
from ..backbones import build_backbone
|
||||
from ..heads import build_head
|
||||
from ..necks import build_neck
|
||||
|
||||
|
||||
class BaseModel(nn.Module):
|
||||
def __init__(self, config, **kwargs):
|
||||
"""
|
||||
the module for OCR.
|
||||
args:
|
||||
config (dict): the super parameters for module.
|
||||
"""
|
||||
super(BaseModel, self).__init__()
|
||||
|
||||
in_channels = config.get("in_channels", 3)
|
||||
model_type = config["model_type"]
|
||||
# build backbone, backbone is need for del, rec and cls
|
||||
if "Backbone" not in config or config["Backbone"] is None:
|
||||
self.use_backbone = False
|
||||
else:
|
||||
self.use_backbone = True
|
||||
config["Backbone"]["in_channels"] = in_channels
|
||||
self.backbone = build_backbone(config["Backbone"], model_type)
|
||||
in_channels = self.backbone.out_channels
|
||||
|
||||
# build neck
|
||||
# for rec, neck can be cnn,rnn or reshape(None)
|
||||
# for det, neck can be FPN, BIFPN and so on.
|
||||
# for cls, neck should be none
|
||||
if "Neck" not in config or config["Neck"] is None:
|
||||
self.use_neck = False
|
||||
else:
|
||||
self.use_neck = True
|
||||
config["Neck"]["in_channels"] = in_channels
|
||||
self.neck = build_neck(config["Neck"])
|
||||
in_channels = self.neck.out_channels
|
||||
|
||||
# # build head, head is need for det, rec and cls
|
||||
if "Head" not in config or config["Head"] is None:
|
||||
self.use_head = False
|
||||
else:
|
||||
self.use_head = True
|
||||
config["Head"]["in_channels"] = in_channels
|
||||
self.head = build_head(config["Head"], **kwargs)
|
||||
|
||||
self.return_all_feats = config.get("return_all_feats", False)
|
||||
|
||||
self._initialize_weights()
|
||||
|
||||
def _initialize_weights(self):
|
||||
# weight initialization
|
||||
for m in self.modules():
|
||||
if isinstance(m, nn.Conv2d):
|
||||
nn.init.kaiming_normal_(m.weight, mode="fan_out")
|
||||
if m.bias is not None:
|
||||
nn.init.zeros_(m.bias)
|
||||
elif isinstance(m, nn.BatchNorm2d):
|
||||
nn.init.ones_(m.weight)
|
||||
nn.init.zeros_(m.bias)
|
||||
elif isinstance(m, nn.Linear):
|
||||
nn.init.normal_(m.weight, 0, 0.01)
|
||||
if m.bias is not None:
|
||||
nn.init.zeros_(m.bias)
|
||||
elif isinstance(m, nn.ConvTranspose2d):
|
||||
nn.init.kaiming_normal_(m.weight, mode="fan_out")
|
||||
if m.bias is not None:
|
||||
nn.init.zeros_(m.bias)
|
||||
|
||||
def forward(self, x):
|
||||
y = dict()
|
||||
if self.use_backbone:
|
||||
x = self.backbone(x)
|
||||
if isinstance(x, dict):
|
||||
y.update(x)
|
||||
else:
|
||||
y["backbone_out"] = x
|
||||
final_name = "backbone_out"
|
||||
if self.use_neck:
|
||||
x = self.neck(x)
|
||||
if isinstance(x, dict):
|
||||
y.update(x)
|
||||
else:
|
||||
y["neck_out"] = x
|
||||
final_name = "neck_out"
|
||||
if self.use_head:
|
||||
x = self.head(x)
|
||||
# for multi head, save ctc neck out for udml
|
||||
if isinstance(x, dict) and "ctc_nect" in x.keys():
|
||||
y["neck_out"] = x["ctc_neck"]
|
||||
y["head_out"] = x
|
||||
elif isinstance(x, dict):
|
||||
y.update(x)
|
||||
else:
|
||||
y["head_out"] = x
|
||||
if self.return_all_feats:
|
||||
if self.training:
|
||||
return y
|
||||
elif isinstance(x, dict):
|
||||
return x
|
||||
else:
|
||||
return {final_name: x}
|
||||
else:
|
||||
return x
|
||||
@@ -0,0 +1,63 @@
|
||||
# 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.
|
||||
|
||||
__all__ = ["build_backbone"]
|
||||
|
||||
|
||||
def build_backbone(config, model_type):
|
||||
if model_type == "det":
|
||||
from .det_mobilenet_v3 import MobileNetV3
|
||||
from .rec_hgnet import PPHGNet_small
|
||||
from .rec_lcnetv3 import PPLCNetV3
|
||||
|
||||
support_dict = [
|
||||
"MobileNetV3",
|
||||
"ResNet",
|
||||
"ResNet_vd",
|
||||
"ResNet_SAST",
|
||||
"PPLCNetV3",
|
||||
"PPHGNet_small",
|
||||
]
|
||||
elif model_type == "rec" or model_type == "cls":
|
||||
from .rec_hgnet import PPHGNet_small
|
||||
from .rec_lcnetv3 import PPLCNetV3
|
||||
from .rec_mobilenet_v3 import MobileNetV3
|
||||
from .rec_svtrnet import SVTRNet
|
||||
from .rec_mv1_enhance import MobileNetV1Enhance
|
||||
from .rec_pphgnetv2 import PPHGNetV2_B4
|
||||
support_dict = [
|
||||
"MobileNetV1Enhance",
|
||||
"MobileNetV3",
|
||||
"ResNet",
|
||||
"ResNetFPN",
|
||||
"MTB",
|
||||
"ResNet31",
|
||||
"SVTRNet",
|
||||
"ViTSTR",
|
||||
"DenseNet",
|
||||
"PPLCNetV3",
|
||||
"PPHGNet_small",
|
||||
"PPHGNetV2_B4",
|
||||
]
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
module_name = config.pop("name")
|
||||
assert module_name in support_dict, Exception(
|
||||
"when model typs is {}, backbone only support {}".format(
|
||||
model_type, support_dict
|
||||
)
|
||||
)
|
||||
module_class = eval(module_name)(**config)
|
||||
return module_class
|
||||
@@ -0,0 +1,269 @@
|
||||
from torch import nn
|
||||
|
||||
from ..common import Activation
|
||||
|
||||
|
||||
def make_divisible(v, divisor=8, min_value=None):
|
||||
if min_value is None:
|
||||
min_value = divisor
|
||||
new_v = max(min_value, int(v + divisor / 2) // divisor * divisor)
|
||||
if new_v < 0.9 * v:
|
||||
new_v += divisor
|
||||
return new_v
|
||||
|
||||
|
||||
class ConvBNLayer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
stride,
|
||||
padding,
|
||||
groups=1,
|
||||
if_act=True,
|
||||
act=None,
|
||||
name=None,
|
||||
):
|
||||
super(ConvBNLayer, self).__init__()
|
||||
self.if_act = if_act
|
||||
self.conv = nn.Conv2d(
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=kernel_size,
|
||||
stride=stride,
|
||||
padding=padding,
|
||||
groups=groups,
|
||||
bias=False,
|
||||
)
|
||||
|
||||
self.bn = nn.BatchNorm2d(
|
||||
out_channels,
|
||||
)
|
||||
if self.if_act:
|
||||
self.act = Activation(act_type=act, inplace=True)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv(x)
|
||||
x = self.bn(x)
|
||||
if self.if_act:
|
||||
x = self.act(x)
|
||||
return x
|
||||
|
||||
|
||||
class SEModule(nn.Module):
|
||||
def __init__(self, in_channels, reduction=4, name=""):
|
||||
super(SEModule, self).__init__()
|
||||
self.avg_pool = nn.AdaptiveAvgPool2d(1)
|
||||
self.conv1 = nn.Conv2d(
|
||||
in_channels=in_channels,
|
||||
out_channels=in_channels // reduction,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0,
|
||||
bias=True,
|
||||
)
|
||||
self.relu1 = Activation(act_type="relu", inplace=True)
|
||||
self.conv2 = nn.Conv2d(
|
||||
in_channels=in_channels // reduction,
|
||||
out_channels=in_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0,
|
||||
bias=True,
|
||||
)
|
||||
self.hard_sigmoid = Activation(act_type="hard_sigmoid", inplace=True)
|
||||
|
||||
def forward(self, inputs):
|
||||
outputs = self.avg_pool(inputs)
|
||||
outputs = self.conv1(outputs)
|
||||
outputs = self.relu1(outputs)
|
||||
outputs = self.conv2(outputs)
|
||||
outputs = self.hard_sigmoid(outputs)
|
||||
outputs = inputs * outputs
|
||||
return outputs
|
||||
|
||||
|
||||
class ResidualUnit(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
mid_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
stride,
|
||||
use_se,
|
||||
act=None,
|
||||
name="",
|
||||
):
|
||||
super(ResidualUnit, self).__init__()
|
||||
self.if_shortcut = stride == 1 and in_channels == out_channels
|
||||
self.if_se = use_se
|
||||
|
||||
self.expand_conv = ConvBNLayer(
|
||||
in_channels=in_channels,
|
||||
out_channels=mid_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0,
|
||||
if_act=True,
|
||||
act=act,
|
||||
name=name + "_expand",
|
||||
)
|
||||
self.bottleneck_conv = ConvBNLayer(
|
||||
in_channels=mid_channels,
|
||||
out_channels=mid_channels,
|
||||
kernel_size=kernel_size,
|
||||
stride=stride,
|
||||
padding=int((kernel_size - 1) // 2),
|
||||
groups=mid_channels,
|
||||
if_act=True,
|
||||
act=act,
|
||||
name=name + "_depthwise",
|
||||
)
|
||||
if self.if_se:
|
||||
self.mid_se = SEModule(mid_channels, name=name + "_se")
|
||||
self.linear_conv = ConvBNLayer(
|
||||
in_channels=mid_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0,
|
||||
if_act=False,
|
||||
act=None,
|
||||
name=name + "_linear",
|
||||
)
|
||||
|
||||
def forward(self, inputs):
|
||||
x = self.expand_conv(inputs)
|
||||
x = self.bottleneck_conv(x)
|
||||
if self.if_se:
|
||||
x = self.mid_se(x)
|
||||
x = self.linear_conv(x)
|
||||
if self.if_shortcut:
|
||||
x = inputs + x
|
||||
return x
|
||||
|
||||
|
||||
class MobileNetV3(nn.Module):
|
||||
def __init__(
|
||||
self, in_channels=3, model_name="large", scale=0.5, disable_se=False, **kwargs
|
||||
):
|
||||
"""
|
||||
the MobilenetV3 backbone network for detection module.
|
||||
Args:
|
||||
params(dict): the super parameters for build network
|
||||
"""
|
||||
super(MobileNetV3, self).__init__()
|
||||
|
||||
self.disable_se = disable_se
|
||||
|
||||
if model_name == "large":
|
||||
cfg = [
|
||||
# k, exp, c, se, nl, s,
|
||||
[3, 16, 16, False, "relu", 1],
|
||||
[3, 64, 24, False, "relu", 2],
|
||||
[3, 72, 24, False, "relu", 1],
|
||||
[5, 72, 40, True, "relu", 2],
|
||||
[5, 120, 40, True, "relu", 1],
|
||||
[5, 120, 40, True, "relu", 1],
|
||||
[3, 240, 80, False, "hard_swish", 2],
|
||||
[3, 200, 80, False, "hard_swish", 1],
|
||||
[3, 184, 80, False, "hard_swish", 1],
|
||||
[3, 184, 80, False, "hard_swish", 1],
|
||||
[3, 480, 112, True, "hard_swish", 1],
|
||||
[3, 672, 112, True, "hard_swish", 1],
|
||||
[5, 672, 160, True, "hard_swish", 2],
|
||||
[5, 960, 160, True, "hard_swish", 1],
|
||||
[5, 960, 160, True, "hard_swish", 1],
|
||||
]
|
||||
cls_ch_squeeze = 960
|
||||
elif model_name == "small":
|
||||
cfg = [
|
||||
# k, exp, c, se, nl, s,
|
||||
[3, 16, 16, True, "relu", 2],
|
||||
[3, 72, 24, False, "relu", 2],
|
||||
[3, 88, 24, False, "relu", 1],
|
||||
[5, 96, 40, True, "hard_swish", 2],
|
||||
[5, 240, 40, True, "hard_swish", 1],
|
||||
[5, 240, 40, True, "hard_swish", 1],
|
||||
[5, 120, 48, True, "hard_swish", 1],
|
||||
[5, 144, 48, True, "hard_swish", 1],
|
||||
[5, 288, 96, True, "hard_swish", 2],
|
||||
[5, 576, 96, True, "hard_swish", 1],
|
||||
[5, 576, 96, True, "hard_swish", 1],
|
||||
]
|
||||
cls_ch_squeeze = 576
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
"mode[" + model_name + "_model] is not implemented!"
|
||||
)
|
||||
|
||||
supported_scale = [0.35, 0.5, 0.75, 1.0, 1.25]
|
||||
assert (
|
||||
scale in supported_scale
|
||||
), "supported scale are {} but input scale is {}".format(supported_scale, scale)
|
||||
inplanes = 16
|
||||
# conv1
|
||||
self.conv = ConvBNLayer(
|
||||
in_channels=in_channels,
|
||||
out_channels=make_divisible(inplanes * scale),
|
||||
kernel_size=3,
|
||||
stride=2,
|
||||
padding=1,
|
||||
groups=1,
|
||||
if_act=True,
|
||||
act="hard_swish",
|
||||
name="conv1",
|
||||
)
|
||||
|
||||
self.stages = nn.ModuleList()
|
||||
self.out_channels = []
|
||||
block_list = []
|
||||
i = 0
|
||||
inplanes = make_divisible(inplanes * scale)
|
||||
for k, exp, c, se, nl, s in cfg:
|
||||
se = se and not self.disable_se
|
||||
if s == 2 and i > 2:
|
||||
self.out_channels.append(inplanes)
|
||||
self.stages.append(nn.Sequential(*block_list))
|
||||
block_list = []
|
||||
block_list.append(
|
||||
ResidualUnit(
|
||||
in_channels=inplanes,
|
||||
mid_channels=make_divisible(scale * exp),
|
||||
out_channels=make_divisible(scale * c),
|
||||
kernel_size=k,
|
||||
stride=s,
|
||||
use_se=se,
|
||||
act=nl,
|
||||
name="conv" + str(i + 2),
|
||||
)
|
||||
)
|
||||
inplanes = make_divisible(scale * c)
|
||||
i += 1
|
||||
block_list.append(
|
||||
ConvBNLayer(
|
||||
in_channels=inplanes,
|
||||
out_channels=make_divisible(scale * cls_ch_squeeze),
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0,
|
||||
groups=1,
|
||||
if_act=True,
|
||||
act="hard_swish",
|
||||
name="conv_last",
|
||||
)
|
||||
)
|
||||
self.stages.append(nn.Sequential(*block_list))
|
||||
self.out_channels.append(make_divisible(scale * cls_ch_squeeze))
|
||||
# for i, stage in enumerate(self.stages):
|
||||
# self.add_sublayer(sublayer=stage, name="stage{}".format(i))
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv(x)
|
||||
out_list = []
|
||||
for stage in self.stages:
|
||||
x = stage(x)
|
||||
out_list.append(x)
|
||||
return out_list
|
||||
@@ -0,0 +1,290 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
|
||||
class ConvBNAct(nn.Module):
|
||||
def __init__(
|
||||
self, in_channels, out_channels, kernel_size, stride, groups=1, use_act=True
|
||||
):
|
||||
super().__init__()
|
||||
self.use_act = use_act
|
||||
self.conv = nn.Conv2d(
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
stride,
|
||||
padding=(kernel_size - 1) // 2,
|
||||
groups=groups,
|
||||
bias=False,
|
||||
)
|
||||
self.bn = nn.BatchNorm2d(out_channels)
|
||||
if self.use_act:
|
||||
self.act = nn.ReLU()
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv(x)
|
||||
x = self.bn(x)
|
||||
if self.use_act:
|
||||
x = self.act(x)
|
||||
return x
|
||||
|
||||
|
||||
class ESEModule(nn.Module):
|
||||
def __init__(self, channels):
|
||||
super().__init__()
|
||||
self.avg_pool = nn.AdaptiveAvgPool2d(1)
|
||||
self.conv = nn.Conv2d(
|
||||
in_channels=channels,
|
||||
out_channels=channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0,
|
||||
)
|
||||
self.sigmoid = nn.Sigmoid()
|
||||
|
||||
def forward(self, x):
|
||||
identity = x
|
||||
x = self.avg_pool(x)
|
||||
x = self.conv(x)
|
||||
x = self.sigmoid(x)
|
||||
return x * identity
|
||||
|
||||
|
||||
class HG_Block(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
mid_channels,
|
||||
out_channels,
|
||||
layer_num,
|
||||
identity=False,
|
||||
):
|
||||
super().__init__()
|
||||
self.identity = identity
|
||||
|
||||
self.layers = nn.ModuleList()
|
||||
self.layers.append(
|
||||
ConvBNAct(
|
||||
in_channels=in_channels,
|
||||
out_channels=mid_channels,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
)
|
||||
)
|
||||
for _ in range(layer_num - 1):
|
||||
self.layers.append(
|
||||
ConvBNAct(
|
||||
in_channels=mid_channels,
|
||||
out_channels=mid_channels,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
)
|
||||
)
|
||||
|
||||
# feature aggregation
|
||||
total_channels = in_channels + layer_num * mid_channels
|
||||
self.aggregation_conv = ConvBNAct(
|
||||
in_channels=total_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
)
|
||||
self.att = ESEModule(out_channels)
|
||||
|
||||
def forward(self, x):
|
||||
identity = x
|
||||
output = []
|
||||
output.append(x)
|
||||
for layer in self.layers:
|
||||
x = layer(x)
|
||||
output.append(x)
|
||||
x = torch.cat(output, dim=1)
|
||||
x = self.aggregation_conv(x)
|
||||
x = self.att(x)
|
||||
if self.identity:
|
||||
x += identity
|
||||
return x
|
||||
|
||||
|
||||
class HG_Stage(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
mid_channels,
|
||||
out_channels,
|
||||
block_num,
|
||||
layer_num,
|
||||
downsample=True,
|
||||
stride=[2, 1],
|
||||
):
|
||||
super().__init__()
|
||||
self.downsample = downsample
|
||||
if downsample:
|
||||
self.downsample = ConvBNAct(
|
||||
in_channels=in_channels,
|
||||
out_channels=in_channels,
|
||||
kernel_size=3,
|
||||
stride=stride,
|
||||
groups=in_channels,
|
||||
use_act=False,
|
||||
)
|
||||
|
||||
blocks_list = []
|
||||
blocks_list.append(
|
||||
HG_Block(in_channels, mid_channels, out_channels, layer_num, identity=False)
|
||||
)
|
||||
for _ in range(block_num - 1):
|
||||
blocks_list.append(
|
||||
HG_Block(
|
||||
out_channels, mid_channels, out_channels, layer_num, identity=True
|
||||
)
|
||||
)
|
||||
self.blocks = nn.Sequential(*blocks_list)
|
||||
|
||||
def forward(self, x):
|
||||
if self.downsample:
|
||||
x = self.downsample(x)
|
||||
x = self.blocks(x)
|
||||
return x
|
||||
|
||||
|
||||
class PPHGNet(nn.Module):
|
||||
"""
|
||||
PPHGNet
|
||||
Args:
|
||||
stem_channels: list. Stem channel list of PPHGNet.
|
||||
stage_config: dict. The configuration of each stage of PPHGNet. such as the number of channels, stride, etc.
|
||||
layer_num: int. Number of layers of HG_Block.
|
||||
use_last_conv: boolean. Whether to use a 1x1 convolutional layer before the classification layer.
|
||||
class_expand: int=2048. Number of channels for the last 1x1 convolutional layer.
|
||||
dropout_prob: float. Parameters of dropout, 0.0 means dropout is not used.
|
||||
class_num: int=1000. The number of classes.
|
||||
Returns:
|
||||
model: nn.Layer. Specific PPHGNet model depends on args.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
stem_channels,
|
||||
stage_config,
|
||||
layer_num,
|
||||
in_channels=3,
|
||||
det=False,
|
||||
out_indices=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.det = det
|
||||
self.out_indices = out_indices if out_indices is not None else [0, 1, 2, 3]
|
||||
|
||||
# stem
|
||||
stem_channels.insert(0, in_channels)
|
||||
self.stem = nn.Sequential(
|
||||
*[
|
||||
ConvBNAct(
|
||||
in_channels=stem_channels[i],
|
||||
out_channels=stem_channels[i + 1],
|
||||
kernel_size=3,
|
||||
stride=2 if i == 0 else 1,
|
||||
)
|
||||
for i in range(len(stem_channels) - 1)
|
||||
]
|
||||
)
|
||||
|
||||
if self.det:
|
||||
self.pool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
|
||||
# stages
|
||||
self.stages = nn.ModuleList()
|
||||
self.out_channels = []
|
||||
for block_id, k in enumerate(stage_config):
|
||||
(
|
||||
in_channels,
|
||||
mid_channels,
|
||||
out_channels,
|
||||
block_num,
|
||||
downsample,
|
||||
stride,
|
||||
) = stage_config[k]
|
||||
self.stages.append(
|
||||
HG_Stage(
|
||||
in_channels,
|
||||
mid_channels,
|
||||
out_channels,
|
||||
block_num,
|
||||
layer_num,
|
||||
downsample,
|
||||
stride,
|
||||
)
|
||||
)
|
||||
if block_id in self.out_indices:
|
||||
self.out_channels.append(out_channels)
|
||||
|
||||
if not self.det:
|
||||
self.out_channels = stage_config["stage4"][2]
|
||||
|
||||
self._init_weights()
|
||||
|
||||
def _init_weights(self):
|
||||
for m in self.modules():
|
||||
if isinstance(m, nn.Conv2d):
|
||||
nn.init.kaiming_normal_(m.weight)
|
||||
elif isinstance(m, nn.BatchNorm2d):
|
||||
nn.init.ones_(m.weight)
|
||||
nn.init.zeros_(m.bias)
|
||||
elif isinstance(m, nn.Linear):
|
||||
nn.init.zeros_(m.bias)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.stem(x)
|
||||
if self.det:
|
||||
x = self.pool(x)
|
||||
|
||||
out = []
|
||||
for i, stage in enumerate(self.stages):
|
||||
x = stage(x)
|
||||
if self.det and i in self.out_indices:
|
||||
out.append(x)
|
||||
if self.det:
|
||||
return out
|
||||
|
||||
if self.training:
|
||||
x = F.adaptive_avg_pool2d(x, [1, 40])
|
||||
else:
|
||||
x = F.avg_pool2d(x, [3, 2])
|
||||
return x
|
||||
|
||||
|
||||
def PPHGNet_small(pretrained=False, use_ssld=False, det=False, **kwargs):
|
||||
"""
|
||||
PPHGNet_small
|
||||
Args:
|
||||
pretrained: bool=False or str. If `True` load pretrained parameters, `False` otherwise.
|
||||
If str, means the path of the pretrained model.
|
||||
use_ssld: bool=False. Whether using distillation pretrained model when pretrained=True.
|
||||
Returns:
|
||||
model: nn.Layer. Specific `PPHGNet_small` model depends on args.
|
||||
"""
|
||||
stage_config_det = {
|
||||
# in_channels, mid_channels, out_channels, blocks, downsample
|
||||
"stage1": [128, 128, 256, 1, False, 2],
|
||||
"stage2": [256, 160, 512, 1, True, 2],
|
||||
"stage3": [512, 192, 768, 2, True, 2],
|
||||
"stage4": [768, 224, 1024, 1, True, 2],
|
||||
}
|
||||
|
||||
stage_config_rec = {
|
||||
# in_channels, mid_channels, out_channels, blocks, downsample
|
||||
"stage1": [128, 128, 256, 1, True, [2, 1]],
|
||||
"stage2": [256, 160, 512, 1, True, [1, 2]],
|
||||
"stage3": [512, 192, 768, 2, True, [2, 1]],
|
||||
"stage4": [768, 224, 1024, 1, True, [2, 1]],
|
||||
}
|
||||
|
||||
model = PPHGNet(
|
||||
stem_channels=[64, 64, 128],
|
||||
stage_config=stage_config_det if det else stage_config_rec,
|
||||
layer_num=6,
|
||||
det=det,
|
||||
**kwargs
|
||||
)
|
||||
return model
|
||||
@@ -0,0 +1,516 @@
|
||||
# copyright (c) 2021 PaddlePaddle Authors. All Rights Reserve.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from __future__ import absolute_import, division, print_function
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
from ..common import Activation
|
||||
|
||||
NET_CONFIG_det = {
|
||||
"blocks2":
|
||||
# k, in_c, out_c, s, use_se
|
||||
[[3, 16, 32, 1, False]],
|
||||
"blocks3": [[3, 32, 64, 2, False], [3, 64, 64, 1, False]],
|
||||
"blocks4": [[3, 64, 128, 2, False], [3, 128, 128, 1, False]],
|
||||
"blocks5": [
|
||||
[3, 128, 256, 2, False],
|
||||
[5, 256, 256, 1, False],
|
||||
[5, 256, 256, 1, False],
|
||||
[5, 256, 256, 1, False],
|
||||
[5, 256, 256, 1, False],
|
||||
],
|
||||
"blocks6": [
|
||||
[5, 256, 512, 2, True],
|
||||
[5, 512, 512, 1, True],
|
||||
[5, 512, 512, 1, False],
|
||||
[5, 512, 512, 1, False],
|
||||
],
|
||||
}
|
||||
|
||||
NET_CONFIG_rec = {
|
||||
"blocks2":
|
||||
# k, in_c, out_c, s, use_se
|
||||
[[3, 16, 32, 1, False]],
|
||||
"blocks3": [[3, 32, 64, 1, False], [3, 64, 64, 1, False]],
|
||||
"blocks4": [[3, 64, 128, (2, 1), False], [3, 128, 128, 1, False]],
|
||||
"blocks5": [
|
||||
[3, 128, 256, (1, 2), False],
|
||||
[5, 256, 256, 1, False],
|
||||
[5, 256, 256, 1, False],
|
||||
[5, 256, 256, 1, False],
|
||||
[5, 256, 256, 1, False],
|
||||
],
|
||||
"blocks6": [
|
||||
[5, 256, 512, (2, 1), True],
|
||||
[5, 512, 512, 1, True],
|
||||
[5, 512, 512, (2, 1), False],
|
||||
[5, 512, 512, 1, False],
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
def make_divisible(v, divisor=16, min_value=None):
|
||||
if min_value is None:
|
||||
min_value = divisor
|
||||
new_v = max(min_value, int(v + divisor / 2) // divisor * divisor)
|
||||
if new_v < 0.9 * v:
|
||||
new_v += divisor
|
||||
return new_v
|
||||
|
||||
|
||||
class LearnableAffineBlock(nn.Module):
|
||||
def __init__(self, scale_value=1.0, bias_value=0.0, lr_mult=1.0, lab_lr=0.1):
|
||||
super().__init__()
|
||||
self.scale = nn.Parameter(torch.Tensor([scale_value]))
|
||||
self.bias = nn.Parameter(torch.Tensor([bias_value]))
|
||||
|
||||
def forward(self, x):
|
||||
return self.scale * x + self.bias
|
||||
|
||||
|
||||
class ConvBNLayer(nn.Module):
|
||||
def __init__(
|
||||
self, in_channels, out_channels, kernel_size, stride, groups=1, lr_mult=1.0
|
||||
):
|
||||
super().__init__()
|
||||
self.conv = nn.Conv2d(
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=kernel_size,
|
||||
stride=stride,
|
||||
padding=(kernel_size - 1) // 2,
|
||||
groups=groups,
|
||||
bias=False,
|
||||
)
|
||||
|
||||
self.bn = nn.BatchNorm2d(
|
||||
out_channels,
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv(x)
|
||||
x = self.bn(x)
|
||||
return x
|
||||
|
||||
|
||||
class Act(nn.Module):
|
||||
def __init__(self, act="hswish", lr_mult=1.0, lab_lr=0.1):
|
||||
super().__init__()
|
||||
if act == "hswish":
|
||||
self.act = nn.Hardswish(inplace=True)
|
||||
else:
|
||||
assert act == "relu"
|
||||
self.act = Activation(act)
|
||||
self.lab = LearnableAffineBlock(lr_mult=lr_mult, lab_lr=lab_lr)
|
||||
|
||||
def forward(self, x):
|
||||
return self.lab(self.act(x))
|
||||
|
||||
|
||||
class LearnableRepLayer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
stride=1,
|
||||
groups=1,
|
||||
num_conv_branches=1,
|
||||
lr_mult=1.0,
|
||||
lab_lr=0.1,
|
||||
):
|
||||
super().__init__()
|
||||
self.is_repped = False
|
||||
self.groups = groups
|
||||
self.stride = stride
|
||||
self.kernel_size = kernel_size
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = out_channels
|
||||
self.num_conv_branches = num_conv_branches
|
||||
self.padding = (kernel_size - 1) // 2
|
||||
|
||||
self.identity = (
|
||||
nn.BatchNorm2d(
|
||||
num_features=in_channels,
|
||||
)
|
||||
if out_channels == in_channels and stride == 1
|
||||
else None
|
||||
)
|
||||
|
||||
self.conv_kxk = nn.ModuleList(
|
||||
[
|
||||
ConvBNLayer(
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
stride,
|
||||
groups=groups,
|
||||
lr_mult=lr_mult,
|
||||
)
|
||||
for _ in range(self.num_conv_branches)
|
||||
]
|
||||
)
|
||||
|
||||
self.conv_1x1 = (
|
||||
ConvBNLayer(
|
||||
in_channels, out_channels, 1, stride, groups=groups, lr_mult=lr_mult
|
||||
)
|
||||
if kernel_size > 1
|
||||
else None
|
||||
)
|
||||
|
||||
self.lab = LearnableAffineBlock(lr_mult=lr_mult, lab_lr=lab_lr)
|
||||
self.act = Act(lr_mult=lr_mult, lab_lr=lab_lr)
|
||||
|
||||
def forward(self, x):
|
||||
# for export
|
||||
if self.is_repped:
|
||||
out = self.lab(self.reparam_conv(x))
|
||||
if self.stride != 2:
|
||||
out = self.act(out)
|
||||
return out
|
||||
|
||||
out = 0
|
||||
if self.identity is not None:
|
||||
out += self.identity(x)
|
||||
|
||||
if self.conv_1x1 is not None:
|
||||
out += self.conv_1x1(x)
|
||||
|
||||
for conv in self.conv_kxk:
|
||||
out += conv(x)
|
||||
|
||||
out = self.lab(out)
|
||||
if self.stride != 2:
|
||||
out = self.act(out)
|
||||
return out
|
||||
|
||||
def rep(self):
|
||||
if self.is_repped:
|
||||
return
|
||||
kernel, bias = self._get_kernel_bias()
|
||||
self.reparam_conv = nn.Conv2d(
|
||||
in_channels=self.in_channels,
|
||||
out_channels=self.out_channels,
|
||||
kernel_size=self.kernel_size,
|
||||
stride=self.stride,
|
||||
padding=self.padding,
|
||||
groups=self.groups,
|
||||
)
|
||||
self.reparam_conv.weight.data = kernel
|
||||
self.reparam_conv.bias.data = bias
|
||||
self.is_repped = True
|
||||
|
||||
def _pad_kernel_1x1_to_kxk(self, kernel1x1, pad):
|
||||
if not isinstance(kernel1x1, torch.Tensor):
|
||||
return 0
|
||||
else:
|
||||
return nn.functional.pad(kernel1x1, [pad, pad, pad, pad])
|
||||
|
||||
def _get_kernel_bias(self):
|
||||
kernel_conv_1x1, bias_conv_1x1 = self._fuse_bn_tensor(self.conv_1x1)
|
||||
kernel_conv_1x1 = self._pad_kernel_1x1_to_kxk(
|
||||
kernel_conv_1x1, self.kernel_size // 2
|
||||
)
|
||||
|
||||
kernel_identity, bias_identity = self._fuse_bn_tensor(self.identity)
|
||||
|
||||
kernel_conv_kxk = 0
|
||||
bias_conv_kxk = 0
|
||||
for conv in self.conv_kxk:
|
||||
kernel, bias = self._fuse_bn_tensor(conv)
|
||||
kernel_conv_kxk += kernel
|
||||
bias_conv_kxk += bias
|
||||
|
||||
kernel_reparam = kernel_conv_kxk + kernel_conv_1x1 + kernel_identity
|
||||
bias_reparam = bias_conv_kxk + bias_conv_1x1 + bias_identity
|
||||
return kernel_reparam, bias_reparam
|
||||
|
||||
def _fuse_bn_tensor(self, branch):
|
||||
if not branch:
|
||||
return 0, 0
|
||||
elif isinstance(branch, ConvBNLayer):
|
||||
kernel = branch.conv.weight
|
||||
running_mean = branch.bn._mean
|
||||
running_var = branch.bn._variance
|
||||
gamma = branch.bn.weight
|
||||
beta = branch.bn.bias
|
||||
eps = branch.bn._epsilon
|
||||
else:
|
||||
assert isinstance(branch, nn.BatchNorm2d)
|
||||
if not hasattr(self, "id_tensor"):
|
||||
input_dim = self.in_channels // self.groups
|
||||
kernel_value = torch.zeros(
|
||||
(self.in_channels, input_dim, self.kernel_size, self.kernel_size),
|
||||
dtype=branch.weight.dtype,
|
||||
)
|
||||
for i in range(self.in_channels):
|
||||
kernel_value[
|
||||
i, i % input_dim, self.kernel_size // 2, self.kernel_size // 2
|
||||
] = 1
|
||||
self.id_tensor = kernel_value
|
||||
kernel = self.id_tensor
|
||||
running_mean = branch._mean
|
||||
running_var = branch._variance
|
||||
gamma = branch.weight
|
||||
beta = branch.bias
|
||||
eps = branch._epsilon
|
||||
std = (running_var + eps).sqrt()
|
||||
t = (gamma / std).reshape((-1, 1, 1, 1))
|
||||
return kernel * t, beta - running_mean * gamma / std
|
||||
|
||||
|
||||
class SELayer(nn.Module):
|
||||
def __init__(self, channel, reduction=4, lr_mult=1.0):
|
||||
super().__init__()
|
||||
self.avg_pool = nn.AdaptiveAvgPool2d(1)
|
||||
self.conv1 = nn.Conv2d(
|
||||
in_channels=channel,
|
||||
out_channels=channel // reduction,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0,
|
||||
)
|
||||
self.relu = nn.ReLU()
|
||||
self.conv2 = nn.Conv2d(
|
||||
in_channels=channel // reduction,
|
||||
out_channels=channel,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0,
|
||||
)
|
||||
self.hardsigmoid = nn.Hardsigmoid(inplace=True)
|
||||
|
||||
def forward(self, x):
|
||||
identity = x
|
||||
x = self.avg_pool(x)
|
||||
x = self.conv1(x)
|
||||
x = self.relu(x)
|
||||
x = self.conv2(x)
|
||||
x = self.hardsigmoid(x)
|
||||
x = identity * x
|
||||
return x
|
||||
|
||||
|
||||
class LCNetV3Block(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
stride,
|
||||
dw_size,
|
||||
use_se=False,
|
||||
conv_kxk_num=4,
|
||||
lr_mult=1.0,
|
||||
lab_lr=0.1,
|
||||
):
|
||||
super().__init__()
|
||||
self.use_se = use_se
|
||||
self.dw_conv = LearnableRepLayer(
|
||||
in_channels=in_channels,
|
||||
out_channels=in_channels,
|
||||
kernel_size=dw_size,
|
||||
stride=stride,
|
||||
groups=in_channels,
|
||||
num_conv_branches=conv_kxk_num,
|
||||
lr_mult=lr_mult,
|
||||
lab_lr=lab_lr,
|
||||
)
|
||||
if use_se:
|
||||
self.se = SELayer(in_channels, lr_mult=lr_mult)
|
||||
self.pw_conv = LearnableRepLayer(
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
num_conv_branches=conv_kxk_num,
|
||||
lr_mult=lr_mult,
|
||||
lab_lr=lab_lr,
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.dw_conv(x)
|
||||
if self.use_se:
|
||||
x = self.se(x)
|
||||
x = self.pw_conv(x)
|
||||
return x
|
||||
|
||||
|
||||
class PPLCNetV3(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
scale=1.0,
|
||||
conv_kxk_num=4,
|
||||
lr_mult_list=[1.0, 1.0, 1.0, 1.0, 1.0, 1.0],
|
||||
lab_lr=0.1,
|
||||
det=False,
|
||||
**kwargs
|
||||
):
|
||||
super().__init__()
|
||||
self.scale = scale
|
||||
self.lr_mult_list = lr_mult_list
|
||||
self.det = det
|
||||
|
||||
self.net_config = NET_CONFIG_det if self.det else NET_CONFIG_rec
|
||||
|
||||
assert isinstance(
|
||||
self.lr_mult_list, (list, tuple)
|
||||
), "lr_mult_list should be in (list, tuple) but got {}".format(
|
||||
type(self.lr_mult_list)
|
||||
)
|
||||
assert (
|
||||
len(self.lr_mult_list) == 6
|
||||
), "lr_mult_list length should be 6 but got {}".format(len(self.lr_mult_list))
|
||||
|
||||
self.conv1 = ConvBNLayer(
|
||||
in_channels=3,
|
||||
out_channels=make_divisible(16 * scale),
|
||||
kernel_size=3,
|
||||
stride=2,
|
||||
lr_mult=self.lr_mult_list[0],
|
||||
)
|
||||
|
||||
self.blocks2 = nn.Sequential(
|
||||
*[
|
||||
LCNetV3Block(
|
||||
in_channels=make_divisible(in_c * scale),
|
||||
out_channels=make_divisible(out_c * scale),
|
||||
dw_size=k,
|
||||
stride=s,
|
||||
use_se=se,
|
||||
conv_kxk_num=conv_kxk_num,
|
||||
lr_mult=self.lr_mult_list[1],
|
||||
lab_lr=lab_lr,
|
||||
)
|
||||
for i, (k, in_c, out_c, s, se) in enumerate(self.net_config["blocks2"])
|
||||
]
|
||||
)
|
||||
|
||||
self.blocks3 = nn.Sequential(
|
||||
*[
|
||||
LCNetV3Block(
|
||||
in_channels=make_divisible(in_c * scale),
|
||||
out_channels=make_divisible(out_c * scale),
|
||||
dw_size=k,
|
||||
stride=s,
|
||||
use_se=se,
|
||||
conv_kxk_num=conv_kxk_num,
|
||||
lr_mult=self.lr_mult_list[2],
|
||||
lab_lr=lab_lr,
|
||||
)
|
||||
for i, (k, in_c, out_c, s, se) in enumerate(self.net_config["blocks3"])
|
||||
]
|
||||
)
|
||||
|
||||
self.blocks4 = nn.Sequential(
|
||||
*[
|
||||
LCNetV3Block(
|
||||
in_channels=make_divisible(in_c * scale),
|
||||
out_channels=make_divisible(out_c * scale),
|
||||
dw_size=k,
|
||||
stride=s,
|
||||
use_se=se,
|
||||
conv_kxk_num=conv_kxk_num,
|
||||
lr_mult=self.lr_mult_list[3],
|
||||
lab_lr=lab_lr,
|
||||
)
|
||||
for i, (k, in_c, out_c, s, se) in enumerate(self.net_config["blocks4"])
|
||||
]
|
||||
)
|
||||
|
||||
self.blocks5 = nn.Sequential(
|
||||
*[
|
||||
LCNetV3Block(
|
||||
in_channels=make_divisible(in_c * scale),
|
||||
out_channels=make_divisible(out_c * scale),
|
||||
dw_size=k,
|
||||
stride=s,
|
||||
use_se=se,
|
||||
conv_kxk_num=conv_kxk_num,
|
||||
lr_mult=self.lr_mult_list[4],
|
||||
lab_lr=lab_lr,
|
||||
)
|
||||
for i, (k, in_c, out_c, s, se) in enumerate(self.net_config["blocks5"])
|
||||
]
|
||||
)
|
||||
|
||||
self.blocks6 = nn.Sequential(
|
||||
*[
|
||||
LCNetV3Block(
|
||||
in_channels=make_divisible(in_c * scale),
|
||||
out_channels=make_divisible(out_c * scale),
|
||||
dw_size=k,
|
||||
stride=s,
|
||||
use_se=se,
|
||||
conv_kxk_num=conv_kxk_num,
|
||||
lr_mult=self.lr_mult_list[5],
|
||||
lab_lr=lab_lr,
|
||||
)
|
||||
for i, (k, in_c, out_c, s, se) in enumerate(self.net_config["blocks6"])
|
||||
]
|
||||
)
|
||||
self.out_channels = make_divisible(512 * scale)
|
||||
|
||||
if self.det:
|
||||
mv_c = [16, 24, 56, 480]
|
||||
self.out_channels = [
|
||||
make_divisible(self.net_config["blocks3"][-1][2] * scale),
|
||||
make_divisible(self.net_config["blocks4"][-1][2] * scale),
|
||||
make_divisible(self.net_config["blocks5"][-1][2] * scale),
|
||||
make_divisible(self.net_config["blocks6"][-1][2] * scale),
|
||||
]
|
||||
|
||||
self.layer_list = nn.ModuleList(
|
||||
[
|
||||
nn.Conv2d(self.out_channels[0], int(mv_c[0] * scale), 1, 1, 0),
|
||||
nn.Conv2d(self.out_channels[1], int(mv_c[1] * scale), 1, 1, 0),
|
||||
nn.Conv2d(self.out_channels[2], int(mv_c[2] * scale), 1, 1, 0),
|
||||
nn.Conv2d(self.out_channels[3], int(mv_c[3] * scale), 1, 1, 0),
|
||||
]
|
||||
)
|
||||
self.out_channels = [
|
||||
int(mv_c[0] * scale),
|
||||
int(mv_c[1] * scale),
|
||||
int(mv_c[2] * scale),
|
||||
int(mv_c[3] * scale),
|
||||
]
|
||||
|
||||
def forward(self, x):
|
||||
out_list = []
|
||||
x = self.conv1(x)
|
||||
x = self.blocks2(x)
|
||||
x = self.blocks3(x)
|
||||
out_list.append(x)
|
||||
x = self.blocks4(x)
|
||||
out_list.append(x)
|
||||
x = self.blocks5(x)
|
||||
out_list.append(x)
|
||||
x = self.blocks6(x)
|
||||
out_list.append(x)
|
||||
|
||||
if self.det:
|
||||
out_list[0] = self.layer_list[0](out_list[0])
|
||||
out_list[1] = self.layer_list[1](out_list[1])
|
||||
out_list[2] = self.layer_list[2](out_list[2])
|
||||
out_list[3] = self.layer_list[3](out_list[3])
|
||||
return out_list
|
||||
|
||||
if self.training:
|
||||
x = F.adaptive_avg_pool2d(x, [1, 40])
|
||||
else:
|
||||
x = F.avg_pool2d(x, [3, 2])
|
||||
return x
|
||||
@@ -0,0 +1,136 @@
|
||||
from torch import nn
|
||||
|
||||
from .det_mobilenet_v3 import ConvBNLayer, ResidualUnit, make_divisible
|
||||
|
||||
|
||||
class MobileNetV3(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels=3,
|
||||
model_name="small",
|
||||
scale=0.5,
|
||||
large_stride=None,
|
||||
small_stride=None,
|
||||
**kwargs
|
||||
):
|
||||
super(MobileNetV3, self).__init__()
|
||||
if small_stride is None:
|
||||
small_stride = [2, 2, 2, 2]
|
||||
if large_stride is None:
|
||||
large_stride = [1, 2, 2, 2]
|
||||
|
||||
assert isinstance(
|
||||
large_stride, list
|
||||
), "large_stride type must " "be list but got {}".format(type(large_stride))
|
||||
assert isinstance(
|
||||
small_stride, list
|
||||
), "small_stride type must " "be list but got {}".format(type(small_stride))
|
||||
assert (
|
||||
len(large_stride) == 4
|
||||
), "large_stride length must be " "4 but got {}".format(len(large_stride))
|
||||
assert (
|
||||
len(small_stride) == 4
|
||||
), "small_stride length must be " "4 but got {}".format(len(small_stride))
|
||||
|
||||
if model_name == "large":
|
||||
cfg = [
|
||||
# k, exp, c, se, nl, s,
|
||||
[3, 16, 16, False, "relu", large_stride[0]],
|
||||
[3, 64, 24, False, "relu", (large_stride[1], 1)],
|
||||
[3, 72, 24, False, "relu", 1],
|
||||
[5, 72, 40, True, "relu", (large_stride[2], 1)],
|
||||
[5, 120, 40, True, "relu", 1],
|
||||
[5, 120, 40, True, "relu", 1],
|
||||
[3, 240, 80, False, "hard_swish", 1],
|
||||
[3, 200, 80, False, "hard_swish", 1],
|
||||
[3, 184, 80, False, "hard_swish", 1],
|
||||
[3, 184, 80, False, "hard_swish", 1],
|
||||
[3, 480, 112, True, "hard_swish", 1],
|
||||
[3, 672, 112, True, "hard_swish", 1],
|
||||
[5, 672, 160, True, "hard_swish", (large_stride[3], 1)],
|
||||
[5, 960, 160, True, "hard_swish", 1],
|
||||
[5, 960, 160, True, "hard_swish", 1],
|
||||
]
|
||||
cls_ch_squeeze = 960
|
||||
elif model_name == "small":
|
||||
cfg = [
|
||||
# k, exp, c, se, nl, s,
|
||||
[3, 16, 16, True, "relu", (small_stride[0], 1)],
|
||||
[3, 72, 24, False, "relu", (small_stride[1], 1)],
|
||||
[3, 88, 24, False, "relu", 1],
|
||||
[5, 96, 40, True, "hard_swish", (small_stride[2], 1)],
|
||||
[5, 240, 40, True, "hard_swish", 1],
|
||||
[5, 240, 40, True, "hard_swish", 1],
|
||||
[5, 120, 48, True, "hard_swish", 1],
|
||||
[5, 144, 48, True, "hard_swish", 1],
|
||||
[5, 288, 96, True, "hard_swish", (small_stride[3], 1)],
|
||||
[5, 576, 96, True, "hard_swish", 1],
|
||||
[5, 576, 96, True, "hard_swish", 1],
|
||||
]
|
||||
cls_ch_squeeze = 576
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
"mode[" + model_name + "_model] is not implemented!"
|
||||
)
|
||||
|
||||
supported_scale = [0.35, 0.5, 0.75, 1.0, 1.25]
|
||||
assert (
|
||||
scale in supported_scale
|
||||
), "supported scales are {} but input scale is {}".format(
|
||||
supported_scale, scale
|
||||
)
|
||||
|
||||
inplanes = 16
|
||||
# conv1
|
||||
self.conv1 = ConvBNLayer(
|
||||
in_channels=in_channels,
|
||||
out_channels=make_divisible(inplanes * scale),
|
||||
kernel_size=3,
|
||||
stride=2,
|
||||
padding=1,
|
||||
groups=1,
|
||||
if_act=True,
|
||||
act="hard_swish",
|
||||
name="conv1",
|
||||
)
|
||||
i = 0
|
||||
block_list = []
|
||||
inplanes = make_divisible(inplanes * scale)
|
||||
for k, exp, c, se, nl, s in cfg:
|
||||
block_list.append(
|
||||
ResidualUnit(
|
||||
in_channels=inplanes,
|
||||
mid_channels=make_divisible(scale * exp),
|
||||
out_channels=make_divisible(scale * c),
|
||||
kernel_size=k,
|
||||
stride=s,
|
||||
use_se=se,
|
||||
act=nl,
|
||||
name="conv" + str(i + 2),
|
||||
)
|
||||
)
|
||||
inplanes = make_divisible(scale * c)
|
||||
i += 1
|
||||
self.blocks = nn.Sequential(*block_list)
|
||||
|
||||
self.conv2 = ConvBNLayer(
|
||||
in_channels=inplanes,
|
||||
out_channels=make_divisible(scale * cls_ch_squeeze),
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0,
|
||||
groups=1,
|
||||
if_act=True,
|
||||
act="hard_swish",
|
||||
name="conv_last",
|
||||
)
|
||||
|
||||
self.pool = nn.MaxPool2d(kernel_size=2, stride=2, padding=0)
|
||||
self.out_channels = make_divisible(scale * cls_ch_squeeze)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv1(x)
|
||||
x = self.blocks(x)
|
||||
x = self.conv2(x)
|
||||
x = self.pool(x)
|
||||
return x
|
||||
@@ -0,0 +1,234 @@
|
||||
import os, sys
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from ..common import Activation
|
||||
|
||||
|
||||
class ConvBNLayer(nn.Module):
|
||||
def __init__(self,
|
||||
num_channels,
|
||||
filter_size,
|
||||
num_filters,
|
||||
stride,
|
||||
padding,
|
||||
channels=None,
|
||||
num_groups=1,
|
||||
act='hard_swish'):
|
||||
super(ConvBNLayer, self).__init__()
|
||||
self.act = act
|
||||
self._conv = nn.Conv2d(
|
||||
in_channels=num_channels,
|
||||
out_channels=num_filters,
|
||||
kernel_size=filter_size,
|
||||
stride=stride,
|
||||
padding=padding,
|
||||
groups=num_groups,
|
||||
bias=False)
|
||||
|
||||
self._batch_norm = nn.BatchNorm2d(
|
||||
num_filters,
|
||||
)
|
||||
if self.act is not None:
|
||||
self._act = Activation(act_type=act, inplace=True)
|
||||
|
||||
def forward(self, inputs):
|
||||
y = self._conv(inputs)
|
||||
y = self._batch_norm(y)
|
||||
if self.act is not None:
|
||||
y = self._act(y)
|
||||
return y
|
||||
|
||||
|
||||
class DepthwiseSeparable(nn.Module):
|
||||
def __init__(self,
|
||||
num_channels,
|
||||
num_filters1,
|
||||
num_filters2,
|
||||
num_groups,
|
||||
stride,
|
||||
scale,
|
||||
dw_size=3,
|
||||
padding=1,
|
||||
use_se=False):
|
||||
super(DepthwiseSeparable, self).__init__()
|
||||
self.use_se = use_se
|
||||
self._depthwise_conv = ConvBNLayer(
|
||||
num_channels=num_channels,
|
||||
num_filters=int(num_filters1 * scale),
|
||||
filter_size=dw_size,
|
||||
stride=stride,
|
||||
padding=padding,
|
||||
num_groups=int(num_groups * scale))
|
||||
if use_se:
|
||||
self._se = SEModule(int(num_filters1 * scale))
|
||||
self._pointwise_conv = ConvBNLayer(
|
||||
num_channels=int(num_filters1 * scale),
|
||||
filter_size=1,
|
||||
num_filters=int(num_filters2 * scale),
|
||||
stride=1,
|
||||
padding=0)
|
||||
|
||||
def forward(self, inputs):
|
||||
y = self._depthwise_conv(inputs)
|
||||
if self.use_se:
|
||||
y = self._se(y)
|
||||
y = self._pointwise_conv(y)
|
||||
return y
|
||||
|
||||
|
||||
class MobileNetV1Enhance(nn.Module):
|
||||
def __init__(self,
|
||||
in_channels=3,
|
||||
scale=0.5,
|
||||
last_conv_stride=1,
|
||||
last_pool_type='max',
|
||||
**kwargs):
|
||||
super().__init__()
|
||||
self.scale = scale
|
||||
self.block_list = []
|
||||
|
||||
self.conv1 = ConvBNLayer(
|
||||
num_channels=in_channels,
|
||||
filter_size=3,
|
||||
channels=3,
|
||||
num_filters=int(32 * scale),
|
||||
stride=2,
|
||||
padding=1)
|
||||
|
||||
conv2_1 = DepthwiseSeparable(
|
||||
num_channels=int(32 * scale),
|
||||
num_filters1=32,
|
||||
num_filters2=64,
|
||||
num_groups=32,
|
||||
stride=1,
|
||||
scale=scale)
|
||||
self.block_list.append(conv2_1)
|
||||
|
||||
conv2_2 = DepthwiseSeparable(
|
||||
num_channels=int(64 * scale),
|
||||
num_filters1=64,
|
||||
num_filters2=128,
|
||||
num_groups=64,
|
||||
stride=1,
|
||||
scale=scale)
|
||||
self.block_list.append(conv2_2)
|
||||
|
||||
conv3_1 = DepthwiseSeparable(
|
||||
num_channels=int(128 * scale),
|
||||
num_filters1=128,
|
||||
num_filters2=128,
|
||||
num_groups=128,
|
||||
stride=1,
|
||||
scale=scale)
|
||||
self.block_list.append(conv3_1)
|
||||
|
||||
conv3_2 = DepthwiseSeparable(
|
||||
num_channels=int(128 * scale),
|
||||
num_filters1=128,
|
||||
num_filters2=256,
|
||||
num_groups=128,
|
||||
stride=(2, 1),
|
||||
scale=scale)
|
||||
self.block_list.append(conv3_2)
|
||||
|
||||
conv4_1 = DepthwiseSeparable(
|
||||
num_channels=int(256 * scale),
|
||||
num_filters1=256,
|
||||
num_filters2=256,
|
||||
num_groups=256,
|
||||
stride=1,
|
||||
scale=scale)
|
||||
self.block_list.append(conv4_1)
|
||||
|
||||
conv4_2 = DepthwiseSeparable(
|
||||
num_channels=int(256 * scale),
|
||||
num_filters1=256,
|
||||
num_filters2=512,
|
||||
num_groups=256,
|
||||
stride=(2, 1),
|
||||
scale=scale)
|
||||
self.block_list.append(conv4_2)
|
||||
|
||||
for _ in range(5):
|
||||
conv5 = DepthwiseSeparable(
|
||||
num_channels=int(512 * scale),
|
||||
num_filters1=512,
|
||||
num_filters2=512,
|
||||
num_groups=512,
|
||||
stride=1,
|
||||
dw_size=5,
|
||||
padding=2,
|
||||
scale=scale,
|
||||
use_se=False)
|
||||
self.block_list.append(conv5)
|
||||
|
||||
conv5_6 = DepthwiseSeparable(
|
||||
num_channels=int(512 * scale),
|
||||
num_filters1=512,
|
||||
num_filters2=1024,
|
||||
num_groups=512,
|
||||
stride=(2, 1),
|
||||
dw_size=5,
|
||||
padding=2,
|
||||
scale=scale,
|
||||
use_se=True)
|
||||
self.block_list.append(conv5_6)
|
||||
|
||||
conv6 = DepthwiseSeparable(
|
||||
num_channels=int(1024 * scale),
|
||||
num_filters1=1024,
|
||||
num_filters2=1024,
|
||||
num_groups=1024,
|
||||
stride=last_conv_stride,
|
||||
dw_size=5,
|
||||
padding=2,
|
||||
use_se=True,
|
||||
scale=scale)
|
||||
self.block_list.append(conv6)
|
||||
|
||||
self.block_list = nn.Sequential(*self.block_list)
|
||||
if last_pool_type == 'avg':
|
||||
self.pool = nn.AvgPool2d(kernel_size=2, stride=2, padding=0)
|
||||
else:
|
||||
self.pool = nn.MaxPool2d(kernel_size=2, stride=2, padding=0)
|
||||
self.out_channels = int(1024 * scale)
|
||||
|
||||
def forward(self, inputs):
|
||||
y = self.conv1(inputs)
|
||||
y = self.block_list(y)
|
||||
y = self.pool(y)
|
||||
return y
|
||||
|
||||
def hardsigmoid(x):
|
||||
return F.relu6(x + 3., inplace=True) / 6.
|
||||
|
||||
class SEModule(nn.Module):
|
||||
def __init__(self, channel, reduction=4):
|
||||
super(SEModule, self).__init__()
|
||||
self.avg_pool = nn.AdaptiveAvgPool2d(1)
|
||||
self.conv1 = nn.Conv2d(
|
||||
in_channels=channel,
|
||||
out_channels=channel // reduction,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0,
|
||||
bias=True)
|
||||
self.conv2 = nn.Conv2d(
|
||||
in_channels=channel // reduction,
|
||||
out_channels=channel,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0,
|
||||
bias=True)
|
||||
|
||||
def forward(self, inputs):
|
||||
outputs = self.avg_pool(inputs)
|
||||
outputs = self.conv1(outputs)
|
||||
outputs = F.relu(outputs)
|
||||
outputs = self.conv2(outputs)
|
||||
outputs = hardsigmoid(outputs)
|
||||
x = torch.mul(inputs, outputs)
|
||||
|
||||
return x
|
||||
@@ -0,0 +1,810 @@
|
||||
import math
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class AdaptiveAvgPool2D(nn.AdaptiveAvgPool2d):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
if isinstance(self.output_size, int) and self.output_size == 1:
|
||||
self._gap = True
|
||||
elif (
|
||||
isinstance(self.output_size, tuple)
|
||||
and self.output_size[0] == 1
|
||||
and self.output_size[1] == 1
|
||||
):
|
||||
self._gap = True
|
||||
else:
|
||||
self._gap = False
|
||||
|
||||
def forward(self, x):
|
||||
if self._gap:
|
||||
# Global Average Pooling
|
||||
N, C, _, _ = x.shape
|
||||
x_mean = torch.mean(x, dim=[2, 3])
|
||||
x_mean = torch.reshape(x_mean, [N, C, 1, 1])
|
||||
return x_mean
|
||||
else:
|
||||
return F.adaptive_avg_pool2d(
|
||||
x,
|
||||
output_size=self.output_size
|
||||
)
|
||||
|
||||
class LearnableAffineBlock(nn.Module):
|
||||
"""
|
||||
Create a learnable affine block module. This module can significantly improve accuracy on smaller models.
|
||||
|
||||
Args:
|
||||
scale_value (float): The initial value of the scale parameter, default is 1.0.
|
||||
bias_value (float): The initial value of the bias parameter, default is 0.0.
|
||||
lr_mult (float): The learning rate multiplier, default is 1.0.
|
||||
lab_lr (float): The learning rate, default is 0.01.
|
||||
"""
|
||||
|
||||
def __init__(self, scale_value=1.0, bias_value=0.0, lr_mult=1.0, lab_lr=0.01):
|
||||
super().__init__()
|
||||
self.scale = nn.Parameter(torch.Tensor([scale_value]))
|
||||
self.bias = nn.Parameter(torch.Tensor([bias_value]))
|
||||
|
||||
def forward(self, x):
|
||||
return self.scale * x + self.bias
|
||||
|
||||
|
||||
class ConvBNAct(nn.Module):
|
||||
"""
|
||||
ConvBNAct is a combination of convolution and batchnorm layers.
|
||||
|
||||
Args:
|
||||
in_channels (int): Number of input channels.
|
||||
out_channels (int): Number of output channels.
|
||||
kernel_size (int): Size of the convolution kernel. Defaults to 3.
|
||||
stride (int): Stride of the convolution. Defaults to 1.
|
||||
padding (int/str): Padding or padding type for the convolution. Defaults to 1.
|
||||
groups (int): Number of groups for the convolution. Defaults to 1.
|
||||
use_act: (bool): Whether to use activation function. Defaults to True.
|
||||
use_lab (bool): Whether to use the LAB operation. Defaults to False.
|
||||
lr_mult (float): Learning rate multiplier for the layer. Defaults to 1.0.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1,
|
||||
groups=1,
|
||||
use_act=True,
|
||||
use_lab=False,
|
||||
lr_mult=1.0,
|
||||
):
|
||||
super().__init__()
|
||||
self.use_act = use_act
|
||||
self.use_lab = use_lab
|
||||
|
||||
self.conv = nn.Conv2d(
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
stride,
|
||||
padding=padding if isinstance(padding, str) else (kernel_size - 1) // 2,
|
||||
# padding=(kernel_size - 1) // 2,
|
||||
groups=groups,
|
||||
bias=False,
|
||||
)
|
||||
self.bn = nn.BatchNorm2d(
|
||||
out_channels,
|
||||
)
|
||||
if self.use_act:
|
||||
self.act = nn.ReLU()
|
||||
if self.use_lab:
|
||||
self.lab = LearnableAffineBlock(lr_mult=lr_mult)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv(x)
|
||||
x = self.bn(x)
|
||||
if self.use_act:
|
||||
x = self.act(x)
|
||||
if self.use_lab:
|
||||
x = self.lab(x)
|
||||
return x
|
||||
|
||||
|
||||
class LightConvBNAct(nn.Module):
|
||||
"""
|
||||
LightConvBNAct is a combination of pw and dw layers.
|
||||
|
||||
Args:
|
||||
in_channels (int): Number of input channels.
|
||||
out_channels (int): Number of output channels.
|
||||
kernel_size (int): Size of the depth-wise convolution kernel.
|
||||
use_lab (bool): Whether to use the LAB operation. Defaults to False.
|
||||
lr_mult (float): Learning rate multiplier for the layer. Defaults to 1.0.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
use_lab=False,
|
||||
lr_mult=1.0,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.conv1 = ConvBNAct(
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=1,
|
||||
use_act=False,
|
||||
use_lab=use_lab,
|
||||
lr_mult=lr_mult,
|
||||
)
|
||||
self.conv2 = ConvBNAct(
|
||||
in_channels=out_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=kernel_size,
|
||||
groups=out_channels,
|
||||
use_act=True,
|
||||
use_lab=use_lab,
|
||||
lr_mult=lr_mult,
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv1(x)
|
||||
x = self.conv2(x)
|
||||
return x
|
||||
|
||||
|
||||
class CustomMaxPool2d(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
kernel_size,
|
||||
stride=None,
|
||||
padding=0,
|
||||
dilation=1,
|
||||
return_indices=False,
|
||||
ceil_mode=False,
|
||||
data_format="NCHW",
|
||||
):
|
||||
super(CustomMaxPool2d, self).__init__()
|
||||
self.kernel_size = kernel_size if isinstance(kernel_size, (tuple, list)) else (kernel_size, kernel_size)
|
||||
self.stride = stride if stride is not None else self.kernel_size
|
||||
self.stride = self.stride if isinstance(self.stride, (tuple, list)) else (self.stride, self.stride)
|
||||
self.dilation = dilation if isinstance(dilation, (tuple, list)) else (dilation, dilation)
|
||||
self.return_indices = return_indices
|
||||
self.ceil_mode = ceil_mode
|
||||
self.padding_mode = padding
|
||||
|
||||
# 当padding不是"same"时使用标准MaxPool2d
|
||||
if padding != "same":
|
||||
self.padding = padding if isinstance(padding, (tuple, list)) else (padding, padding)
|
||||
self.pool = nn.MaxPool2d(
|
||||
kernel_size=self.kernel_size,
|
||||
stride=self.stride,
|
||||
padding=self.padding,
|
||||
dilation=self.dilation,
|
||||
return_indices=self.return_indices,
|
||||
ceil_mode=self.ceil_mode
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
# 处理same padding
|
||||
if self.padding_mode == "same":
|
||||
input_height, input_width = x.size(2), x.size(3)
|
||||
|
||||
# 计算期望的输出尺寸
|
||||
out_height = math.ceil(input_height / self.stride[0])
|
||||
out_width = math.ceil(input_width / self.stride[1])
|
||||
|
||||
# 计算需要的padding
|
||||
pad_height = max((out_height - 1) * self.stride[0] + self.kernel_size[0] - input_height, 0)
|
||||
pad_width = max((out_width - 1) * self.stride[1] + self.kernel_size[1] - input_width, 0)
|
||||
|
||||
# 将padding分配到两边
|
||||
pad_top = pad_height // 2
|
||||
pad_bottom = pad_height - pad_top
|
||||
pad_left = pad_width // 2
|
||||
pad_right = pad_width - pad_left
|
||||
|
||||
# 应用padding
|
||||
x = F.pad(x, (pad_left, pad_right, pad_top, pad_bottom))
|
||||
|
||||
# 使用标准max_pool2d函数
|
||||
if self.return_indices:
|
||||
return F.max_pool2d_with_indices(
|
||||
x,
|
||||
kernel_size=self.kernel_size,
|
||||
stride=self.stride,
|
||||
padding=0, # 已经手动pad过了
|
||||
dilation=self.dilation,
|
||||
ceil_mode=self.ceil_mode
|
||||
)
|
||||
else:
|
||||
return F.max_pool2d(
|
||||
x,
|
||||
kernel_size=self.kernel_size,
|
||||
stride=self.stride,
|
||||
padding=0, # 已经手动pad过了
|
||||
dilation=self.dilation,
|
||||
ceil_mode=self.ceil_mode
|
||||
)
|
||||
else:
|
||||
# 使用预定义的MaxPool2d
|
||||
return self.pool(x)
|
||||
|
||||
class StemBlock(nn.Module):
|
||||
"""
|
||||
StemBlock for PP-HGNetV2.
|
||||
|
||||
Args:
|
||||
in_channels (int): Number of input channels.
|
||||
mid_channels (int): Number of middle channels.
|
||||
out_channels (int): Number of output channels.
|
||||
use_lab (bool): Whether to use the LAB operation. Defaults to False.
|
||||
lr_mult (float): Learning rate multiplier for the layer. Defaults to 1.0.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
mid_channels,
|
||||
out_channels,
|
||||
use_lab=False,
|
||||
lr_mult=1.0,
|
||||
text_rec=False,
|
||||
):
|
||||
super().__init__()
|
||||
self.stem1 = ConvBNAct(
|
||||
in_channels=in_channels,
|
||||
out_channels=mid_channels,
|
||||
kernel_size=3,
|
||||
stride=2,
|
||||
use_lab=use_lab,
|
||||
lr_mult=lr_mult,
|
||||
)
|
||||
self.stem2a = ConvBNAct(
|
||||
in_channels=mid_channels,
|
||||
out_channels=mid_channels // 2,
|
||||
kernel_size=2,
|
||||
stride=1,
|
||||
padding="same",
|
||||
use_lab=use_lab,
|
||||
lr_mult=lr_mult,
|
||||
)
|
||||
self.stem2b = ConvBNAct(
|
||||
in_channels=mid_channels // 2,
|
||||
out_channels=mid_channels,
|
||||
kernel_size=2,
|
||||
stride=1,
|
||||
padding="same",
|
||||
use_lab=use_lab,
|
||||
lr_mult=lr_mult,
|
||||
)
|
||||
self.stem3 = ConvBNAct(
|
||||
in_channels=mid_channels * 2,
|
||||
out_channels=mid_channels,
|
||||
kernel_size=3,
|
||||
stride=1 if text_rec else 2,
|
||||
use_lab=use_lab,
|
||||
lr_mult=lr_mult,
|
||||
)
|
||||
self.stem4 = ConvBNAct(
|
||||
in_channels=mid_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
use_lab=use_lab,
|
||||
lr_mult=lr_mult,
|
||||
)
|
||||
self.pool = CustomMaxPool2d(
|
||||
kernel_size=2, stride=1, ceil_mode=True, padding="same"
|
||||
)
|
||||
# self.pool = nn.MaxPool2d(
|
||||
# kernel_size=2, stride=1, ceil_mode=True, padding=1
|
||||
# )
|
||||
|
||||
def forward(self, x):
|
||||
x = self.stem1(x)
|
||||
x2 = self.stem2a(x)
|
||||
x2 = self.stem2b(x2)
|
||||
x1 = self.pool(x)
|
||||
|
||||
# if x1.shape[2:] != x2.shape[2:]:
|
||||
# x1 = F.interpolate(x1, size=x2.shape[2:], mode='bilinear', align_corners=False)
|
||||
|
||||
x = torch.cat([x1, x2], 1)
|
||||
x = self.stem3(x)
|
||||
x = self.stem4(x)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class HGV2_Block(nn.Module):
|
||||
"""
|
||||
HGV2_Block, the basic unit that constitutes the HGV2_Stage.
|
||||
|
||||
Args:
|
||||
in_channels (int): Number of input channels.
|
||||
mid_channels (int): Number of middle channels.
|
||||
out_channels (int): Number of output channels.
|
||||
kernel_size (int): Size of the convolution kernel. Defaults to 3.
|
||||
layer_num (int): Number of layers in the HGV2 block. Defaults to 6.
|
||||
stride (int): Stride of the convolution. Defaults to 1.
|
||||
padding (int/str): Padding or padding type for the convolution. Defaults to 1.
|
||||
groups (int): Number of groups for the convolution. Defaults to 1.
|
||||
use_act (bool): Whether to use activation function. Defaults to True.
|
||||
use_lab (bool): Whether to use the LAB operation. Defaults to False.
|
||||
lr_mult (float): Learning rate multiplier for the layer. Defaults to 1.0.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
mid_channels,
|
||||
out_channels,
|
||||
kernel_size=3,
|
||||
layer_num=6,
|
||||
identity=False,
|
||||
light_block=True,
|
||||
use_lab=False,
|
||||
lr_mult=1.0,
|
||||
):
|
||||
super().__init__()
|
||||
self.identity = identity
|
||||
|
||||
self.layers = nn.ModuleList()
|
||||
block_type = "LightConvBNAct" if light_block else "ConvBNAct"
|
||||
for i in range(layer_num):
|
||||
self.layers.append(
|
||||
eval(block_type)(
|
||||
in_channels=in_channels if i == 0 else mid_channels,
|
||||
out_channels=mid_channels,
|
||||
stride=1,
|
||||
kernel_size=kernel_size,
|
||||
use_lab=use_lab,
|
||||
lr_mult=lr_mult,
|
||||
)
|
||||
)
|
||||
# feature aggregation
|
||||
total_channels = in_channels + layer_num * mid_channels
|
||||
self.aggregation_squeeze_conv = ConvBNAct(
|
||||
in_channels=total_channels,
|
||||
out_channels=out_channels // 2,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
use_lab=use_lab,
|
||||
lr_mult=lr_mult,
|
||||
)
|
||||
self.aggregation_excitation_conv = ConvBNAct(
|
||||
in_channels=out_channels // 2,
|
||||
out_channels=out_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
use_lab=use_lab,
|
||||
lr_mult=lr_mult,
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
identity = x
|
||||
output = []
|
||||
output.append(x)
|
||||
for layer in self.layers:
|
||||
x = layer(x)
|
||||
output.append(x)
|
||||
x = torch.cat(output, dim=1)
|
||||
x = self.aggregation_squeeze_conv(x)
|
||||
x = self.aggregation_excitation_conv(x)
|
||||
if self.identity:
|
||||
x += identity
|
||||
return x
|
||||
|
||||
|
||||
class HGV2_Stage(nn.Module):
|
||||
"""
|
||||
HGV2_Stage, the basic unit that constitutes the PPHGNetV2.
|
||||
|
||||
Args:
|
||||
in_channels (int): Number of input channels.
|
||||
mid_channels (int): Number of middle channels.
|
||||
out_channels (int): Number of output channels.
|
||||
block_num (int): Number of blocks in the HGV2 stage.
|
||||
layer_num (int): Number of layers in the HGV2 block. Defaults to 6.
|
||||
is_downsample (bool): Whether to use downsampling operation. Defaults to False.
|
||||
light_block (bool): Whether to use light block. Defaults to True.
|
||||
kernel_size (int): Size of the convolution kernel. Defaults to 3.
|
||||
use_lab (bool, optional): Whether to use the LAB operation. Defaults to False.
|
||||
lr_mult (float, optional): Learning rate multiplier for the layer. Defaults to 1.0.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
mid_channels,
|
||||
out_channels,
|
||||
block_num,
|
||||
layer_num=6,
|
||||
is_downsample=True,
|
||||
light_block=True,
|
||||
kernel_size=3,
|
||||
use_lab=False,
|
||||
stride=2,
|
||||
lr_mult=1.0,
|
||||
):
|
||||
|
||||
super().__init__()
|
||||
self.is_downsample = is_downsample
|
||||
if self.is_downsample:
|
||||
self.downsample = ConvBNAct(
|
||||
in_channels=in_channels,
|
||||
out_channels=in_channels,
|
||||
kernel_size=3,
|
||||
stride=stride,
|
||||
groups=in_channels,
|
||||
use_act=False,
|
||||
use_lab=use_lab,
|
||||
lr_mult=lr_mult,
|
||||
)
|
||||
|
||||
blocks_list = []
|
||||
for i in range(block_num):
|
||||
blocks_list.append(
|
||||
HGV2_Block(
|
||||
in_channels=in_channels if i == 0 else out_channels,
|
||||
mid_channels=mid_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=kernel_size,
|
||||
layer_num=layer_num,
|
||||
identity=False if i == 0 else True,
|
||||
light_block=light_block,
|
||||
use_lab=use_lab,
|
||||
lr_mult=lr_mult,
|
||||
)
|
||||
)
|
||||
self.blocks = nn.Sequential(*blocks_list)
|
||||
|
||||
def forward(self, x):
|
||||
if self.is_downsample:
|
||||
x = self.downsample(x)
|
||||
x = self.blocks(x)
|
||||
return x
|
||||
|
||||
|
||||
class DropoutInferDownscale(nn.Module):
|
||||
"""
|
||||
实现与Paddle的mode="downscale_in_infer"等效的Dropout
|
||||
训练模式:out = input * mask(直接应用掩码,不进行放大)
|
||||
推理模式:out = input * (1.0 - p)(在推理时按概率缩小)
|
||||
"""
|
||||
|
||||
def __init__(self, p=0.5):
|
||||
super().__init__()
|
||||
self.p = p
|
||||
|
||||
def forward(self, x):
|
||||
if self.training:
|
||||
# 训练时:应用随机mask但不放大
|
||||
return F.dropout(x, self.p, training=True) * (1.0 - self.p)
|
||||
else:
|
||||
# 推理时:按照dropout概率缩小输出
|
||||
return x * (1.0 - self.p)
|
||||
|
||||
class PPHGNetV2(nn.Module):
|
||||
"""
|
||||
PPHGNetV2
|
||||
|
||||
Args:
|
||||
stage_config (dict): Config for PPHGNetV2 stages. such as the number of channels, stride, etc.
|
||||
stem_channels: (list): Number of channels of the stem of the PPHGNetV2.
|
||||
use_lab (bool): Whether to use the LAB operation. Defaults to False.
|
||||
use_last_conv (bool): Whether to use the last conv layer as the output channel. Defaults to True.
|
||||
class_expand (int): Number of channels for the last 1x1 convolutional layer.
|
||||
drop_prob (float): Dropout probability for the last 1x1 convolutional layer. Defaults to 0.0.
|
||||
class_num (int): The number of classes for the classification layer. Defaults to 1000.
|
||||
lr_mult_list (list): Learning rate multiplier for the stages. Defaults to [1.0, 1.0, 1.0, 1.0, 1.0].
|
||||
Returns:
|
||||
model: nn.Layer. Specific PPHGNetV2 model depends on args.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
stage_config,
|
||||
stem_channels=[3, 32, 64],
|
||||
use_lab=False,
|
||||
use_last_conv=True,
|
||||
class_expand=2048,
|
||||
dropout_prob=0.0,
|
||||
class_num=1000,
|
||||
lr_mult_list=[1.0, 1.0, 1.0, 1.0, 1.0],
|
||||
det=False,
|
||||
text_rec=False,
|
||||
out_indices=None,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
self.det = det
|
||||
self.text_rec = text_rec
|
||||
self.use_lab = use_lab
|
||||
self.use_last_conv = use_last_conv
|
||||
self.class_expand = class_expand
|
||||
self.class_num = class_num
|
||||
self.out_indices = out_indices if out_indices is not None else [0, 1, 2, 3]
|
||||
self.out_channels = []
|
||||
|
||||
# stem
|
||||
self.stem = StemBlock(
|
||||
in_channels=stem_channels[0],
|
||||
mid_channels=stem_channels[1],
|
||||
out_channels=stem_channels[2],
|
||||
use_lab=use_lab,
|
||||
lr_mult=lr_mult_list[0],
|
||||
text_rec=text_rec,
|
||||
)
|
||||
|
||||
# stages
|
||||
self.stages = nn.ModuleList()
|
||||
for i, k in enumerate(stage_config):
|
||||
(
|
||||
in_channels,
|
||||
mid_channels,
|
||||
out_channels,
|
||||
block_num,
|
||||
is_downsample,
|
||||
light_block,
|
||||
kernel_size,
|
||||
layer_num,
|
||||
stride,
|
||||
) = stage_config[k]
|
||||
self.stages.append(
|
||||
HGV2_Stage(
|
||||
in_channels,
|
||||
mid_channels,
|
||||
out_channels,
|
||||
block_num,
|
||||
layer_num,
|
||||
is_downsample,
|
||||
light_block,
|
||||
kernel_size,
|
||||
use_lab,
|
||||
stride,
|
||||
lr_mult=lr_mult_list[i + 1],
|
||||
)
|
||||
)
|
||||
if i in self.out_indices:
|
||||
self.out_channels.append(out_channels)
|
||||
if not self.det:
|
||||
self.out_channels = stage_config["stage4"][2]
|
||||
|
||||
self.avg_pool = AdaptiveAvgPool2D(1)
|
||||
|
||||
if self.use_last_conv:
|
||||
self.last_conv = nn.Conv2d(
|
||||
in_channels=out_channels,
|
||||
out_channels=self.class_expand,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0,
|
||||
bias=False,
|
||||
)
|
||||
self.act = nn.ReLU()
|
||||
if self.use_lab:
|
||||
self.lab = LearnableAffineBlock()
|
||||
self.dropout = DropoutInferDownscale(p=dropout_prob)
|
||||
|
||||
self.flatten = nn.Flatten(start_dim=1, end_dim=-1)
|
||||
if not self.det:
|
||||
self.fc = nn.Linear(
|
||||
self.class_expand if self.use_last_conv else out_channels,
|
||||
self.class_num,
|
||||
)
|
||||
|
||||
self._init_weights()
|
||||
|
||||
def _init_weights(self):
|
||||
for m in self.modules():
|
||||
if isinstance(m, nn.Conv2d):
|
||||
nn.init.kaiming_normal_(m.weight)
|
||||
elif isinstance(m, nn.BatchNorm2d):
|
||||
nn.init.ones_(m.weight)
|
||||
nn.init.zeros_(m.bias)
|
||||
elif isinstance(m, nn.Linear):
|
||||
nn.init.zeros_(m.bias)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.stem(x)
|
||||
out = []
|
||||
for i, stage in enumerate(self.stages):
|
||||
x = stage(x)
|
||||
if self.det and i in self.out_indices:
|
||||
out.append(x)
|
||||
if self.det:
|
||||
return out
|
||||
|
||||
if self.text_rec:
|
||||
if self.training:
|
||||
x = F.adaptive_avg_pool2d(x, [1, 40])
|
||||
else:
|
||||
x = F.avg_pool2d(x, [3, 2])
|
||||
return x
|
||||
|
||||
|
||||
def PPHGNetV2_B0(pretrained=False, use_ssld=False, **kwargs):
|
||||
"""
|
||||
PPHGNetV2_B0
|
||||
Args:
|
||||
pretrained (bool/str): If `True` load pretrained parameters, `False` otherwise.
|
||||
If str, means the path of the pretrained model.
|
||||
use_ssld (bool) Whether using ssld pretrained model when pretrained is True.
|
||||
Returns:
|
||||
model: nn.Layer. Specific `PPHGNetV2_B0` model depends on args.
|
||||
"""
|
||||
stage_config = {
|
||||
# in_channels, mid_channels, out_channels, num_blocks, is_downsample, light_block, kernel_size, layer_num
|
||||
"stage1": [16, 16, 64, 1, False, False, 3, 3],
|
||||
"stage2": [64, 32, 256, 1, True, False, 3, 3],
|
||||
"stage3": [256, 64, 512, 2, True, True, 5, 3],
|
||||
"stage4": [512, 128, 1024, 1, True, True, 5, 3],
|
||||
}
|
||||
|
||||
model = PPHGNetV2(
|
||||
stem_channels=[3, 16, 16], stage_config=stage_config, use_lab=True, **kwargs
|
||||
)
|
||||
return model
|
||||
|
||||
|
||||
def PPHGNetV2_B1(pretrained=False, use_ssld=False, **kwargs):
|
||||
"""
|
||||
PPHGNetV2_B1
|
||||
Args:
|
||||
pretrained (bool/str): If `True` load pretrained parameters, `False` otherwise.
|
||||
If str, means the path of the pretrained model.
|
||||
use_ssld (bool) Whether using ssld pretrained model when pretrained is True.
|
||||
Returns:
|
||||
model: nn.Layer. Specific `PPHGNetV2_B1` model depends on args.
|
||||
"""
|
||||
stage_config = {
|
||||
# in_channels, mid_channels, out_channels, num_blocks, is_downsample, light_block, kernel_size, layer_num
|
||||
"stage1": [32, 32, 64, 1, False, False, 3, 3],
|
||||
"stage2": [64, 48, 256, 1, True, False, 3, 3],
|
||||
"stage3": [256, 96, 512, 2, True, True, 5, 3],
|
||||
"stage4": [512, 192, 1024, 1, True, True, 5, 3],
|
||||
}
|
||||
|
||||
model = PPHGNetV2(
|
||||
stem_channels=[3, 24, 32], stage_config=stage_config, use_lab=True, **kwargs
|
||||
)
|
||||
return model
|
||||
|
||||
|
||||
def PPHGNetV2_B2(pretrained=False, use_ssld=False, **kwargs):
|
||||
"""
|
||||
PPHGNetV2_B2
|
||||
Args:
|
||||
pretrained (bool/str): If `True` load pretrained parameters, `False` otherwise.
|
||||
If str, means the path of the pretrained model.
|
||||
use_ssld (bool) Whether using ssld pretrained model when pretrained is True.
|
||||
Returns:
|
||||
model: nn.Layer. Specific `PPHGNetV2_B2` model depends on args.
|
||||
"""
|
||||
stage_config = {
|
||||
# in_channels, mid_channels, out_channels, num_blocks, is_downsample, light_block, kernel_size, layer_num
|
||||
"stage1": [32, 32, 96, 1, False, False, 3, 4],
|
||||
"stage2": [96, 64, 384, 1, True, False, 3, 4],
|
||||
"stage3": [384, 128, 768, 3, True, True, 5, 4],
|
||||
"stage4": [768, 256, 1536, 1, True, True, 5, 4],
|
||||
}
|
||||
|
||||
model = PPHGNetV2(
|
||||
stem_channels=[3, 24, 32], stage_config=stage_config, use_lab=True, **kwargs
|
||||
)
|
||||
return model
|
||||
|
||||
|
||||
def PPHGNetV2_B3(pretrained=False, use_ssld=False, **kwargs):
|
||||
"""
|
||||
PPHGNetV2_B3
|
||||
Args:
|
||||
pretrained (bool/str): If `True` load pretrained parameters, `False` otherwise.
|
||||
If str, means the path of the pretrained model.
|
||||
use_ssld (bool) Whether using ssld pretrained model when pretrained is True.
|
||||
Returns:
|
||||
model: nn.Layer. Specific `PPHGNetV2_B3` model depends on args.
|
||||
"""
|
||||
stage_config = {
|
||||
# in_channels, mid_channels, out_channels, num_blocks, is_downsample, light_block, kernel_size, layer_num
|
||||
"stage1": [32, 32, 128, 1, False, False, 3, 5],
|
||||
"stage2": [128, 64, 512, 1, True, False, 3, 5],
|
||||
"stage3": [512, 128, 1024, 3, True, True, 5, 5],
|
||||
"stage4": [1024, 256, 2048, 1, True, True, 5, 5],
|
||||
}
|
||||
|
||||
model = PPHGNetV2(
|
||||
stem_channels=[3, 24, 32], stage_config=stage_config, use_lab=True, **kwargs
|
||||
)
|
||||
return model
|
||||
|
||||
|
||||
def PPHGNetV2_B4(pretrained=False, use_ssld=False, det=False, text_rec=False, **kwargs):
|
||||
"""
|
||||
PPHGNetV2_B4
|
||||
Args:
|
||||
pretrained (bool/str): If `True` load pretrained parameters, `False` otherwise.
|
||||
If str, means the path of the pretrained model.
|
||||
use_ssld (bool) Whether using ssld pretrained model when pretrained is True.
|
||||
Returns:
|
||||
model: nn.Layer. Specific `PPHGNetV2_B4` model depends on args.
|
||||
"""
|
||||
stage_config_rec = {
|
||||
# in_channels, mid_channels, out_channels, num_blocks, is_downsample, light_block, kernel_size, layer_num, stride
|
||||
"stage1": [48, 48, 128, 1, True, False, 3, 6, [2, 1]],
|
||||
"stage2": [128, 96, 512, 1, True, False, 3, 6, [1, 2]],
|
||||
"stage3": [512, 192, 1024, 3, True, True, 5, 6, [2, 1]],
|
||||
"stage4": [1024, 384, 2048, 1, True, True, 5, 6, [2, 1]],
|
||||
}
|
||||
|
||||
stage_config_det = {
|
||||
# in_channels, mid_channels, out_channels, num_blocks, is_downsample, light_block, kernel_size, layer_num
|
||||
"stage1": [48, 48, 128, 1, False, False, 3, 6, 2],
|
||||
"stage2": [128, 96, 512, 1, True, False, 3, 6, 2],
|
||||
"stage3": [512, 192, 1024, 3, True, True, 5, 6, 2],
|
||||
"stage4": [1024, 384, 2048, 1, True, True, 5, 6, 2],
|
||||
}
|
||||
model = PPHGNetV2(
|
||||
stem_channels=[3, 32, 48],
|
||||
stage_config=stage_config_det if det else stage_config_rec,
|
||||
use_lab=False,
|
||||
det=det,
|
||||
text_rec=text_rec,
|
||||
**kwargs,
|
||||
)
|
||||
return model
|
||||
|
||||
|
||||
def PPHGNetV2_B5(pretrained=False, use_ssld=False, **kwargs):
|
||||
"""
|
||||
PPHGNetV2_B5
|
||||
Args:
|
||||
pretrained (bool/str): If `True` load pretrained parameters, `False` otherwise.
|
||||
If str, means the path of the pretrained model.
|
||||
use_ssld (bool) Whether using ssld pretrained model when pretrained is True.
|
||||
Returns:
|
||||
model: nn.Layer. Specific `PPHGNetV2_B5` model depends on args.
|
||||
"""
|
||||
stage_config = {
|
||||
# in_channels, mid_channels, out_channels, num_blocks, is_downsample, light_block, kernel_size, layer_num
|
||||
"stage1": [64, 64, 128, 1, False, False, 3, 6],
|
||||
"stage2": [128, 128, 512, 2, True, False, 3, 6],
|
||||
"stage3": [512, 256, 1024, 5, True, True, 5, 6],
|
||||
"stage4": [1024, 512, 2048, 2, True, True, 5, 6],
|
||||
}
|
||||
|
||||
model = PPHGNetV2(
|
||||
stem_channels=[3, 32, 64], stage_config=stage_config, use_lab=False, **kwargs
|
||||
)
|
||||
return model
|
||||
|
||||
|
||||
def PPHGNetV2_B6(pretrained=False, use_ssld=False, **kwargs):
|
||||
"""
|
||||
PPHGNetV2_B6
|
||||
Args:
|
||||
pretrained (bool/str): If `True` load pretrained parameters, `False` otherwise.
|
||||
If str, means the path of the pretrained model.
|
||||
use_ssld (bool) Whether using ssld pretrained model when pretrained is True.
|
||||
Returns:
|
||||
model: nn.Layer. Specific `PPHGNetV2_B6` model depends on args.
|
||||
"""
|
||||
stage_config = {
|
||||
# in_channels, mid_channels, out_channels, num_blocks, is_downsample, light_block, kernel_size, layer_num
|
||||
"stage1": [96, 96, 192, 2, False, False, 3, 6],
|
||||
"stage2": [192, 192, 512, 3, True, False, 3, 6],
|
||||
"stage3": [512, 384, 1024, 6, True, True, 5, 6],
|
||||
"stage4": [1024, 768, 2048, 3, True, True, 5, 6],
|
||||
}
|
||||
|
||||
model = PPHGNetV2(
|
||||
stem_channels=[3, 48, 96], stage_config=stage_config, use_lab=False, **kwargs
|
||||
)
|
||||
return model
|
||||
@@ -0,0 +1,638 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from ..common import Activation
|
||||
|
||||
|
||||
def drop_path(x, drop_prob=0.0, training=False):
|
||||
"""Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).
|
||||
the original name is misleading as 'Drop Connect' is a different form of dropout in a separate paper...
|
||||
See discussion: https://github.com/tensorflow/tpu/issues/494#issuecomment-532968956 ...
|
||||
"""
|
||||
if drop_prob == 0.0 or not training:
|
||||
return x
|
||||
keep_prob = torch.as_tensor(1 - drop_prob)
|
||||
shape = (x.shape[0],) + (1,) * (x.ndim - 1)
|
||||
random_tensor = keep_prob + torch.rand(shape, dtype=x.dtype)
|
||||
random_tensor = torch.floor(random_tensor) # binarize
|
||||
output = x.divide(keep_prob) * random_tensor
|
||||
return output
|
||||
|
||||
|
||||
class ConvBNLayer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=0,
|
||||
bias_attr=False,
|
||||
groups=1,
|
||||
act="gelu",
|
||||
):
|
||||
super().__init__()
|
||||
self.conv = nn.Conv2d(
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=kernel_size,
|
||||
stride=stride,
|
||||
padding=padding,
|
||||
groups=groups,
|
||||
bias=bias_attr,
|
||||
)
|
||||
self.norm = nn.BatchNorm2d(out_channels)
|
||||
self.act = Activation(act_type=act, inplace=True)
|
||||
|
||||
def forward(self, inputs):
|
||||
out = self.conv(inputs)
|
||||
out = self.norm(out)
|
||||
out = self.act(out)
|
||||
return out
|
||||
|
||||
|
||||
class DropPath(nn.Module):
|
||||
"""Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks)."""
|
||||
|
||||
def __init__(self, drop_prob=None):
|
||||
super(DropPath, self).__init__()
|
||||
self.drop_prob = drop_prob
|
||||
|
||||
def forward(self, x):
|
||||
return drop_path(x, self.drop_prob, self.training)
|
||||
|
||||
|
||||
class Identity(nn.Module):
|
||||
def __init__(self):
|
||||
super(Identity, self).__init__()
|
||||
|
||||
def forward(self, input):
|
||||
return input
|
||||
|
||||
|
||||
class Mlp(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_features,
|
||||
hidden_features=None,
|
||||
out_features=None,
|
||||
act_layer="gelu",
|
||||
drop=0.0,
|
||||
):
|
||||
super().__init__()
|
||||
out_features = out_features or in_features
|
||||
hidden_features = hidden_features or in_features
|
||||
self.fc1 = nn.Linear(in_features, hidden_features)
|
||||
self.act = Activation(act_type=act_layer, inplace=True)
|
||||
self.fc2 = nn.Linear(hidden_features, out_features)
|
||||
self.drop = nn.Dropout(drop)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.fc1(x)
|
||||
x = self.act(x)
|
||||
x = self.drop(x)
|
||||
x = self.fc2(x)
|
||||
x = self.drop(x)
|
||||
return x
|
||||
|
||||
|
||||
class ConvMixer(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
num_heads=8,
|
||||
HW=[8, 25],
|
||||
local_k=[3, 3],
|
||||
):
|
||||
super().__init__()
|
||||
self.HW = HW
|
||||
self.dim = dim
|
||||
self.local_mixer = nn.Conv2d(
|
||||
dim,
|
||||
dim,
|
||||
local_k,
|
||||
1,
|
||||
[local_k[0] // 2, local_k[1] // 2],
|
||||
groups=num_heads,
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
h = self.HW[0]
|
||||
w = self.HW[1]
|
||||
x = x.transpose([0, 2, 1]).reshape([0, self.dim, h, w])
|
||||
x = self.local_mixer(x)
|
||||
x = x.flatten(2).permute(0, 2, 1)
|
||||
return x
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
num_heads=8,
|
||||
mixer="Global",
|
||||
HW=[8, 25],
|
||||
local_k=[7, 11],
|
||||
qkv_bias=False,
|
||||
qk_scale=None,
|
||||
attn_drop=0.0,
|
||||
proj_drop=0.0,
|
||||
):
|
||||
super().__init__()
|
||||
self.num_heads = num_heads
|
||||
head_dim = dim // num_heads
|
||||
self.scale = qk_scale or head_dim**-0.5
|
||||
self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
|
||||
self.attn_drop = nn.Dropout(attn_drop)
|
||||
self.proj = nn.Linear(dim, dim)
|
||||
self.proj_drop = nn.Dropout(proj_drop)
|
||||
self.HW = HW
|
||||
if HW is not None:
|
||||
H = HW[0]
|
||||
W = HW[1]
|
||||
self.N = H * W
|
||||
self.C = dim
|
||||
if mixer == "Local" and HW is not None:
|
||||
hk = local_k[0]
|
||||
wk = local_k[1]
|
||||
mask = torch.ones(H * W, H + hk - 1, W + wk - 1, dtype=torch.float32)
|
||||
for h in range(0, H):
|
||||
for w in range(0, W):
|
||||
mask[h * W + w, h : h + hk, w : w + wk] = 0.0
|
||||
mask_paddle = mask[:, hk // 2 : H + hk // 2, wk // 2 : W + wk // 2].flatten(
|
||||
1
|
||||
)
|
||||
mask_inf = torch.full(
|
||||
[H * W, H * W], fill_value=float("-Inf"), dtype=torch.float32
|
||||
)
|
||||
mask = torch.where(mask_paddle < 1, mask_paddle, mask_inf)
|
||||
self.mask = mask.unsqueeze(0).unsqueeze(1)
|
||||
# self.mask = mask[None, None, :]
|
||||
self.mixer = mixer
|
||||
|
||||
def forward(self, x):
|
||||
if self.HW is not None:
|
||||
N = self.N
|
||||
C = self.C
|
||||
else:
|
||||
_, N, C = x.shape
|
||||
qkv = self.qkv(x)
|
||||
qkv = qkv.reshape((-1, N, 3, self.num_heads, C // self.num_heads)).permute(
|
||||
2, 0, 3, 1, 4
|
||||
)
|
||||
q, k, v = qkv[0] * self.scale, qkv[1], qkv[2]
|
||||
|
||||
attn = q.matmul(k.permute(0, 1, 3, 2))
|
||||
if self.mixer == "Local":
|
||||
attn += self.mask
|
||||
attn = nn.functional.softmax(attn, dim=-1)
|
||||
attn = self.attn_drop(attn)
|
||||
|
||||
x = (attn.matmul(v)).permute(0, 2, 1, 3).reshape((-1, N, C))
|
||||
x = self.proj(x)
|
||||
x = self.proj_drop(x)
|
||||
return x
|
||||
|
||||
|
||||
class Block(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
num_heads,
|
||||
mixer="Global",
|
||||
local_mixer=[7, 11],
|
||||
HW=None,
|
||||
mlp_ratio=4.0,
|
||||
qkv_bias=False,
|
||||
qk_scale=None,
|
||||
drop=0.0,
|
||||
attn_drop=0.0,
|
||||
drop_path=0.0,
|
||||
act_layer="gelu",
|
||||
norm_layer="nn.LayerNorm",
|
||||
epsilon=1e-6,
|
||||
prenorm=True,
|
||||
):
|
||||
super().__init__()
|
||||
if isinstance(norm_layer, str):
|
||||
self.norm1 = eval(norm_layer)(dim, eps=epsilon)
|
||||
else:
|
||||
self.norm1 = norm_layer(dim)
|
||||
if mixer == "Global" or mixer == "Local":
|
||||
self.mixer = Attention(
|
||||
dim,
|
||||
num_heads=num_heads,
|
||||
mixer=mixer,
|
||||
HW=HW,
|
||||
local_k=local_mixer,
|
||||
qkv_bias=qkv_bias,
|
||||
qk_scale=qk_scale,
|
||||
attn_drop=attn_drop,
|
||||
proj_drop=drop,
|
||||
)
|
||||
elif mixer == "Conv":
|
||||
self.mixer = ConvMixer(dim, num_heads=num_heads, HW=HW, local_k=local_mixer)
|
||||
else:
|
||||
raise TypeError("The mixer must be one of [Global, Local, Conv]")
|
||||
|
||||
self.drop_path = DropPath(drop_path) if drop_path > 0.0 else Identity()
|
||||
if isinstance(norm_layer, str):
|
||||
self.norm2 = eval(norm_layer)(dim, eps=epsilon)
|
||||
else:
|
||||
self.norm2 = norm_layer(dim)
|
||||
mlp_hidden_dim = int(dim * mlp_ratio)
|
||||
self.mlp_ratio = mlp_ratio
|
||||
self.mlp = Mlp(
|
||||
in_features=dim,
|
||||
hidden_features=mlp_hidden_dim,
|
||||
act_layer=act_layer,
|
||||
drop=drop,
|
||||
)
|
||||
self.prenorm = prenorm
|
||||
|
||||
def forward(self, x):
|
||||
if self.prenorm:
|
||||
x = self.norm1(x + self.drop_path(self.mixer(x)))
|
||||
x = self.norm2(x + self.drop_path(self.mlp(x)))
|
||||
else:
|
||||
x = x + self.drop_path(self.mixer(self.norm1(x)))
|
||||
x = x + self.drop_path(self.mlp(self.norm2(x)))
|
||||
return x
|
||||
|
||||
|
||||
class PatchEmbed(nn.Module):
|
||||
"""Image to Patch Embedding"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
img_size=[32, 100],
|
||||
in_channels=3,
|
||||
embed_dim=768,
|
||||
sub_num=2,
|
||||
patch_size=[4, 4],
|
||||
mode="pope",
|
||||
):
|
||||
super().__init__()
|
||||
num_patches = (img_size[1] // (2**sub_num)) * (img_size[0] // (2**sub_num))
|
||||
self.img_size = img_size
|
||||
self.num_patches = num_patches
|
||||
self.embed_dim = embed_dim
|
||||
self.norm = None
|
||||
if mode == "pope":
|
||||
if sub_num == 2:
|
||||
self.proj = nn.Sequential(
|
||||
ConvBNLayer(
|
||||
in_channels=in_channels,
|
||||
out_channels=embed_dim // 2,
|
||||
kernel_size=3,
|
||||
stride=2,
|
||||
padding=1,
|
||||
act="gelu",
|
||||
bias_attr=True,
|
||||
),
|
||||
ConvBNLayer(
|
||||
in_channels=embed_dim // 2,
|
||||
out_channels=embed_dim,
|
||||
kernel_size=3,
|
||||
stride=2,
|
||||
padding=1,
|
||||
act="gelu",
|
||||
bias_attr=True,
|
||||
),
|
||||
)
|
||||
if sub_num == 3:
|
||||
self.proj = nn.Sequential(
|
||||
ConvBNLayer(
|
||||
in_channels=in_channels,
|
||||
out_channels=embed_dim // 4,
|
||||
kernel_size=3,
|
||||
stride=2,
|
||||
padding=1,
|
||||
act="gelu",
|
||||
bias_attr=True,
|
||||
),
|
||||
ConvBNLayer(
|
||||
in_channels=embed_dim // 4,
|
||||
out_channels=embed_dim // 2,
|
||||
kernel_size=3,
|
||||
stride=2,
|
||||
padding=1,
|
||||
act="gelu",
|
||||
bias_attr=True,
|
||||
),
|
||||
ConvBNLayer(
|
||||
in_channels=embed_dim // 2,
|
||||
out_channels=embed_dim,
|
||||
kernel_size=3,
|
||||
stride=2,
|
||||
padding=1,
|
||||
act="gelu",
|
||||
bias_attr=True,
|
||||
),
|
||||
)
|
||||
elif mode == "linear":
|
||||
self.proj = nn.Conv2d(
|
||||
1, embed_dim, kernel_size=patch_size, stride=patch_size
|
||||
)
|
||||
self.num_patches = (
|
||||
img_size[0] // patch_size[0] * img_size[1] // patch_size[1]
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
B, C, H, W = x.shape
|
||||
assert (
|
||||
H == self.img_size[0] and W == self.img_size[1]
|
||||
), "Input image size ({}*{}) doesn't match model ({}*{}).".format(
|
||||
H, W, self.img_size[0], self.img_size[1]
|
||||
)
|
||||
x = self.proj(x).flatten(2).permute(0, 2, 1)
|
||||
return x
|
||||
|
||||
|
||||
class SubSample(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
types="Pool",
|
||||
stride=[2, 1],
|
||||
sub_norm="nn.LayerNorm",
|
||||
act=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.types = types
|
||||
if types == "Pool":
|
||||
self.avgpool = nn.AvgPool2d(
|
||||
kernel_size=[3, 5], stride=stride, padding=[1, 2]
|
||||
)
|
||||
self.maxpool = nn.MaxPool2d(
|
||||
kernel_size=[3, 5], stride=stride, padding=[1, 2]
|
||||
)
|
||||
self.proj = nn.Linear(in_channels, out_channels)
|
||||
else:
|
||||
self.conv = nn.Conv2d(
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size=3,
|
||||
stride=stride,
|
||||
padding=1,
|
||||
)
|
||||
self.norm = eval(sub_norm)(out_channels)
|
||||
if act is not None:
|
||||
self.act = act()
|
||||
else:
|
||||
self.act = None
|
||||
|
||||
def forward(self, x):
|
||||
if self.types == "Pool":
|
||||
x1 = self.avgpool(x)
|
||||
x2 = self.maxpool(x)
|
||||
x = (x1 + x2) * 0.5
|
||||
out = self.proj(x.flatten(2).permute(0, 2, 1))
|
||||
else:
|
||||
x = self.conv(x)
|
||||
out = x.flatten(2).permute(0, 2, 1)
|
||||
out = self.norm(out)
|
||||
if self.act is not None:
|
||||
out = self.act(out)
|
||||
|
||||
return out
|
||||
|
||||
|
||||
class SVTRNet(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
img_size=[32, 100],
|
||||
in_channels=3,
|
||||
embed_dim=[64, 128, 256],
|
||||
depth=[3, 6, 3],
|
||||
num_heads=[2, 4, 8],
|
||||
mixer=["Local"] * 6 + ["Global"] * 6, # Local atten, Global atten, Conv
|
||||
local_mixer=[[7, 11], [7, 11], [7, 11]],
|
||||
patch_merging="Conv", # Conv, Pool, None
|
||||
mlp_ratio=4,
|
||||
qkv_bias=True,
|
||||
qk_scale=None,
|
||||
drop_rate=0.0,
|
||||
last_drop=0.0,
|
||||
attn_drop_rate=0.0,
|
||||
drop_path_rate=0.1,
|
||||
norm_layer="nn.LayerNorm",
|
||||
sub_norm="nn.LayerNorm",
|
||||
epsilon=1e-6,
|
||||
out_channels=192,
|
||||
out_char_num=25,
|
||||
block_unit="Block",
|
||||
act="gelu",
|
||||
last_stage=True,
|
||||
sub_num=2,
|
||||
prenorm=True,
|
||||
use_lenhead=False,
|
||||
**kwargs
|
||||
):
|
||||
super().__init__()
|
||||
self.img_size = img_size
|
||||
self.embed_dim = embed_dim
|
||||
self.out_channels = out_channels
|
||||
self.prenorm = prenorm
|
||||
patch_merging = (
|
||||
None
|
||||
if patch_merging != "Conv" and patch_merging != "Pool"
|
||||
else patch_merging
|
||||
)
|
||||
self.patch_embed = PatchEmbed(
|
||||
img_size=img_size,
|
||||
in_channels=in_channels,
|
||||
embed_dim=embed_dim[0],
|
||||
sub_num=sub_num,
|
||||
)
|
||||
num_patches = self.patch_embed.num_patches
|
||||
self.HW = [img_size[0] // (2**sub_num), img_size[1] // (2**sub_num)]
|
||||
self.pos_embed = nn.Parameter(torch.zeros(1, num_patches, embed_dim[0]))
|
||||
self.pos_drop = nn.Dropout(p=drop_rate)
|
||||
Block_unit = eval(block_unit)
|
||||
|
||||
dpr = np.linspace(0, drop_path_rate, sum(depth))
|
||||
self.blocks1 = nn.ModuleList(
|
||||
[
|
||||
Block_unit(
|
||||
dim=embed_dim[0],
|
||||
num_heads=num_heads[0],
|
||||
mixer=mixer[0 : depth[0]][i],
|
||||
HW=self.HW,
|
||||
local_mixer=local_mixer[0],
|
||||
mlp_ratio=mlp_ratio,
|
||||
qkv_bias=qkv_bias,
|
||||
qk_scale=qk_scale,
|
||||
drop=drop_rate,
|
||||
act_layer=act,
|
||||
attn_drop=attn_drop_rate,
|
||||
drop_path=dpr[0 : depth[0]][i],
|
||||
norm_layer=norm_layer,
|
||||
epsilon=epsilon,
|
||||
prenorm=prenorm,
|
||||
)
|
||||
for i in range(depth[0])
|
||||
]
|
||||
)
|
||||
if patch_merging is not None:
|
||||
self.sub_sample1 = SubSample(
|
||||
embed_dim[0],
|
||||
embed_dim[1],
|
||||
sub_norm=sub_norm,
|
||||
stride=[2, 1],
|
||||
types=patch_merging,
|
||||
)
|
||||
HW = [self.HW[0] // 2, self.HW[1]]
|
||||
else:
|
||||
HW = self.HW
|
||||
self.patch_merging = patch_merging
|
||||
self.blocks2 = nn.ModuleList(
|
||||
[
|
||||
Block_unit(
|
||||
dim=embed_dim[1],
|
||||
num_heads=num_heads[1],
|
||||
mixer=mixer[depth[0] : depth[0] + depth[1]][i],
|
||||
HW=HW,
|
||||
local_mixer=local_mixer[1],
|
||||
mlp_ratio=mlp_ratio,
|
||||
qkv_bias=qkv_bias,
|
||||
qk_scale=qk_scale,
|
||||
drop=drop_rate,
|
||||
act_layer=act,
|
||||
attn_drop=attn_drop_rate,
|
||||
drop_path=dpr[depth[0] : depth[0] + depth[1]][i],
|
||||
norm_layer=norm_layer,
|
||||
epsilon=epsilon,
|
||||
prenorm=prenorm,
|
||||
)
|
||||
for i in range(depth[1])
|
||||
]
|
||||
)
|
||||
if patch_merging is not None:
|
||||
self.sub_sample2 = SubSample(
|
||||
embed_dim[1],
|
||||
embed_dim[2],
|
||||
sub_norm=sub_norm,
|
||||
stride=[2, 1],
|
||||
types=patch_merging,
|
||||
)
|
||||
HW = [self.HW[0] // 4, self.HW[1]]
|
||||
else:
|
||||
HW = self.HW
|
||||
self.blocks3 = nn.ModuleList(
|
||||
[
|
||||
Block_unit(
|
||||
dim=embed_dim[2],
|
||||
num_heads=num_heads[2],
|
||||
mixer=mixer[depth[0] + depth[1] :][i],
|
||||
HW=HW,
|
||||
local_mixer=local_mixer[2],
|
||||
mlp_ratio=mlp_ratio,
|
||||
qkv_bias=qkv_bias,
|
||||
qk_scale=qk_scale,
|
||||
drop=drop_rate,
|
||||
act_layer=act,
|
||||
attn_drop=attn_drop_rate,
|
||||
drop_path=dpr[depth[0] + depth[1] :][i],
|
||||
norm_layer=norm_layer,
|
||||
epsilon=epsilon,
|
||||
prenorm=prenorm,
|
||||
)
|
||||
for i in range(depth[2])
|
||||
]
|
||||
)
|
||||
self.last_stage = last_stage
|
||||
if last_stage:
|
||||
self.avg_pool = nn.AdaptiveAvgPool2d([1, out_char_num])
|
||||
self.last_conv = nn.Conv2d(
|
||||
in_channels=embed_dim[2],
|
||||
out_channels=self.out_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0,
|
||||
bias=False,
|
||||
)
|
||||
self.hardswish = Activation("hard_swish", inplace=True) # nn.Hardswish()
|
||||
# self.dropout = nn.Dropout(p=last_drop, mode="downscale_in_infer")
|
||||
self.dropout = nn.Dropout(p=last_drop)
|
||||
if not prenorm:
|
||||
self.norm = eval(norm_layer)(embed_dim[-1], eps=epsilon)
|
||||
self.use_lenhead = use_lenhead
|
||||
if use_lenhead:
|
||||
self.len_conv = nn.Linear(embed_dim[2], self.out_channels)
|
||||
self.hardswish_len = Activation(
|
||||
"hard_swish", inplace=True
|
||||
) # nn.Hardswish()
|
||||
self.dropout_len = nn.Dropout(p=last_drop)
|
||||
|
||||
torch.nn.init.xavier_normal_(self.pos_embed)
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
# weight initialization
|
||||
if isinstance(m, nn.Conv2d):
|
||||
nn.init.kaiming_normal_(m.weight, mode="fan_out")
|
||||
if m.bias is not None:
|
||||
nn.init.zeros_(m.bias)
|
||||
elif isinstance(m, nn.BatchNorm2d):
|
||||
nn.init.ones_(m.weight)
|
||||
nn.init.zeros_(m.bias)
|
||||
elif isinstance(m, nn.Linear):
|
||||
nn.init.normal_(m.weight, 0, 0.01)
|
||||
if m.bias is not None:
|
||||
nn.init.zeros_(m.bias)
|
||||
elif isinstance(m, nn.ConvTranspose2d):
|
||||
nn.init.kaiming_normal_(m.weight, mode="fan_out")
|
||||
if m.bias is not None:
|
||||
nn.init.zeros_(m.bias)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.ones_(m.weight)
|
||||
nn.init.zeros_(m.bias)
|
||||
|
||||
def forward_features(self, x):
|
||||
x = self.patch_embed(x)
|
||||
x = x + self.pos_embed
|
||||
x = self.pos_drop(x)
|
||||
for blk in self.blocks1:
|
||||
x = blk(x)
|
||||
if self.patch_merging is not None:
|
||||
x = self.sub_sample1(
|
||||
x.permute(0, 2, 1).reshape(
|
||||
[-1, self.embed_dim[0], self.HW[0], self.HW[1]]
|
||||
)
|
||||
)
|
||||
for blk in self.blocks2:
|
||||
x = blk(x)
|
||||
if self.patch_merging is not None:
|
||||
x = self.sub_sample2(
|
||||
x.permute(0, 2, 1).reshape(
|
||||
[-1, self.embed_dim[1], self.HW[0] // 2, self.HW[1]]
|
||||
)
|
||||
)
|
||||
for blk in self.blocks3:
|
||||
x = blk(x)
|
||||
if not self.prenorm:
|
||||
x = self.norm(x)
|
||||
return x
|
||||
|
||||
def forward(self, x):
|
||||
x = self.forward_features(x)
|
||||
if self.use_lenhead:
|
||||
len_x = self.len_conv(x.mean(1))
|
||||
len_x = self.dropout_len(self.hardswish_len(len_x))
|
||||
if self.last_stage:
|
||||
if self.patch_merging is not None:
|
||||
h = self.HW[0] // 4
|
||||
else:
|
||||
h = self.HW[0]
|
||||
x = self.avg_pool(
|
||||
x.permute(0, 2, 1).reshape([-1, self.embed_dim[2], h, self.HW[1]])
|
||||
)
|
||||
x = self.last_conv(x)
|
||||
x = self.hardswish(x)
|
||||
x = self.dropout(x)
|
||||
if self.use_lenhead:
|
||||
return x, len_x
|
||||
return x
|
||||
@@ -0,0 +1,76 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
|
||||
class Hswish(nn.Module):
|
||||
def __init__(self, inplace=True):
|
||||
super(Hswish, self).__init__()
|
||||
self.inplace = inplace
|
||||
|
||||
def forward(self, x):
|
||||
return x * F.relu6(x + 3.0, inplace=self.inplace) / 6.0
|
||||
|
||||
|
||||
# out = max(0, min(1, slop*x+offset))
|
||||
# paddle.fluid.layers.hard_sigmoid(x, slope=0.2, offset=0.5, name=None)
|
||||
class Hsigmoid(nn.Module):
|
||||
def __init__(self, inplace=True):
|
||||
super(Hsigmoid, self).__init__()
|
||||
self.inplace = inplace
|
||||
|
||||
def forward(self, x):
|
||||
# torch: F.relu6(x + 3., inplace=self.inplace) / 6.
|
||||
# paddle: F.relu6(1.2 * x + 3., inplace=self.inplace) / 6.
|
||||
return F.relu6(1.2 * x + 3.0, inplace=self.inplace) / 6.0
|
||||
|
||||
|
||||
class GELU(nn.Module):
|
||||
def __init__(self, inplace=True):
|
||||
super(GELU, self).__init__()
|
||||
self.inplace = inplace
|
||||
|
||||
def forward(self, x):
|
||||
return torch.nn.functional.gelu(x)
|
||||
|
||||
|
||||
class Swish(nn.Module):
|
||||
def __init__(self, inplace=True):
|
||||
super(Swish, self).__init__()
|
||||
self.inplace = inplace
|
||||
|
||||
def forward(self, x):
|
||||
if self.inplace:
|
||||
x.mul_(torch.sigmoid(x))
|
||||
return x
|
||||
else:
|
||||
return x * torch.sigmoid(x)
|
||||
|
||||
|
||||
class Activation(nn.Module):
|
||||
def __init__(self, act_type, inplace=True):
|
||||
super(Activation, self).__init__()
|
||||
act_type = act_type.lower()
|
||||
if act_type == "relu":
|
||||
self.act = nn.ReLU(inplace=inplace)
|
||||
elif act_type == "relu6":
|
||||
self.act = nn.ReLU6(inplace=inplace)
|
||||
elif act_type == "sigmoid":
|
||||
raise NotImplementedError
|
||||
elif act_type == "hard_sigmoid":
|
||||
self.act = Hsigmoid(
|
||||
inplace
|
||||
) # nn.Hardsigmoid(inplace=inplace)#Hsigmoid(inplace)#
|
||||
elif act_type == "hard_swish" or act_type == "hswish":
|
||||
self.act = Hswish(inplace=inplace)
|
||||
elif act_type == "leakyrelu":
|
||||
self.act = nn.LeakyReLU(inplace=inplace)
|
||||
elif act_type == "gelu":
|
||||
self.act = GELU(inplace=inplace)
|
||||
elif act_type == "swish":
|
||||
self.act = Swish(inplace=inplace)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
def forward(self, inputs):
|
||||
return self.act(inputs)
|
||||
@@ -0,0 +1,43 @@
|
||||
# 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.
|
||||
|
||||
__all__ = ["build_head"]
|
||||
|
||||
|
||||
def build_head(config, **kwargs):
|
||||
# det head
|
||||
from .det_db_head import DBHead, PFHeadLocal
|
||||
|
||||
# rec head
|
||||
from .rec_ctc_head import CTCHead
|
||||
from .rec_multi_head import MultiHead
|
||||
|
||||
# cls head
|
||||
from .cls_head import ClsHead
|
||||
|
||||
support_dict = [
|
||||
"DBHead",
|
||||
"CTCHead",
|
||||
"ClsHead",
|
||||
"MultiHead",
|
||||
"PFHeadLocal",
|
||||
]
|
||||
|
||||
module_name = config.pop("name")
|
||||
char_num = config.pop("char_num", 6625)
|
||||
assert module_name in support_dict, Exception(
|
||||
"head only support {}".format(support_dict)
|
||||
)
|
||||
module_class = eval(module_name)(**config, **kwargs)
|
||||
return module_class
|
||||
@@ -0,0 +1,23 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
|
||||
class ClsHead(nn.Module):
|
||||
"""
|
||||
Class orientation
|
||||
Args:
|
||||
params(dict): super parameters for build Class network
|
||||
"""
|
||||
|
||||
def __init__(self, in_channels, class_dim, **kwargs):
|
||||
super(ClsHead, self).__init__()
|
||||
self.pool = nn.AdaptiveAvgPool2d(1)
|
||||
self.fc = nn.Linear(in_channels, class_dim, bias=True)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.pool(x)
|
||||
x = torch.reshape(x, shape=[x.shape[0], x.shape[1]])
|
||||
x = self.fc(x)
|
||||
x = F.softmax(x, dim=1)
|
||||
return x
|
||||
@@ -0,0 +1,109 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from ..common import Activation
|
||||
from ..backbones.det_mobilenet_v3 import ConvBNLayer
|
||||
|
||||
class Head(nn.Module):
|
||||
def __init__(self, in_channels, **kwargs):
|
||||
super(Head, self).__init__()
|
||||
self.conv1 = nn.Conv2d(
|
||||
in_channels=in_channels,
|
||||
out_channels=in_channels // 4,
|
||||
kernel_size=3,
|
||||
padding=1,
|
||||
bias=False)
|
||||
self.conv_bn1 = nn.BatchNorm2d(
|
||||
in_channels // 4)
|
||||
self.relu1 = Activation(act_type='relu')
|
||||
|
||||
self.conv2 = nn.ConvTranspose2d(
|
||||
in_channels=in_channels // 4,
|
||||
out_channels=in_channels // 4,
|
||||
kernel_size=2,
|
||||
stride=2)
|
||||
self.conv_bn2 = nn.BatchNorm2d(
|
||||
in_channels // 4)
|
||||
self.relu2 = Activation(act_type='relu')
|
||||
|
||||
self.conv3 = nn.ConvTranspose2d(
|
||||
in_channels=in_channels // 4,
|
||||
out_channels=1,
|
||||
kernel_size=2,
|
||||
stride=2)
|
||||
|
||||
def forward(self, x, return_f=False):
|
||||
x = self.conv1(x)
|
||||
x = self.conv_bn1(x)
|
||||
x = self.relu1(x)
|
||||
x = self.conv2(x)
|
||||
x = self.conv_bn2(x)
|
||||
x = self.relu2(x)
|
||||
if return_f is True:
|
||||
f = x
|
||||
x = self.conv3(x)
|
||||
x = torch.sigmoid(x)
|
||||
if return_f is True:
|
||||
return x, f
|
||||
return x
|
||||
|
||||
|
||||
class DBHead(nn.Module):
|
||||
"""
|
||||
Differentiable Binarization (DB) for text detection:
|
||||
see https://arxiv.org/abs/1911.08947
|
||||
args:
|
||||
params(dict): super parameters for build DB network
|
||||
"""
|
||||
|
||||
def __init__(self, in_channels, k=50, **kwargs):
|
||||
super(DBHead, self).__init__()
|
||||
self.k = k
|
||||
binarize_name_list = [
|
||||
'conv2d_56', 'batch_norm_47', 'conv2d_transpose_0', 'batch_norm_48',
|
||||
'conv2d_transpose_1', 'binarize'
|
||||
]
|
||||
thresh_name_list = [
|
||||
'conv2d_57', 'batch_norm_49', 'conv2d_transpose_2', 'batch_norm_50',
|
||||
'conv2d_transpose_3', 'thresh'
|
||||
]
|
||||
self.binarize = Head(in_channels, **kwargs)# binarize_name_list)
|
||||
self.thresh = Head(in_channels, **kwargs)#thresh_name_list)
|
||||
|
||||
def step_function(self, x, y):
|
||||
return torch.reciprocal(1 + torch.exp(-self.k * (x - y)))
|
||||
|
||||
def forward(self, x):
|
||||
shrink_maps = self.binarize(x)
|
||||
return {'maps': shrink_maps}
|
||||
|
||||
|
||||
class LocalModule(nn.Module):
|
||||
def __init__(self, in_c, mid_c, use_distance=True):
|
||||
super(self.__class__, self).__init__()
|
||||
self.last_3 = ConvBNLayer(in_c + 1, mid_c, 3, 1, 1, act='relu')
|
||||
self.last_1 = nn.Conv2d(mid_c, 1, 1, 1, 0)
|
||||
|
||||
def forward(self, x, init_map, distance_map):
|
||||
outf = torch.cat([init_map, x], dim=1)
|
||||
# last Conv
|
||||
out = self.last_1(self.last_3(outf))
|
||||
return out
|
||||
|
||||
class PFHeadLocal(DBHead):
|
||||
def __init__(self, in_channels, k=50, mode='small', **kwargs):
|
||||
super(PFHeadLocal, self).__init__(in_channels, k, **kwargs)
|
||||
self.mode = mode
|
||||
|
||||
self.up_conv = nn.Upsample(scale_factor=2, mode="nearest")
|
||||
if self.mode == 'large':
|
||||
self.cbn_layer = LocalModule(in_channels // 4, in_channels // 4)
|
||||
elif self.mode == 'small':
|
||||
self.cbn_layer = LocalModule(in_channels // 4, in_channels // 8)
|
||||
|
||||
def forward(self, x, targets=None):
|
||||
shrink_maps, f = self.binarize(x, return_f=True)
|
||||
base_maps = shrink_maps
|
||||
cbn_maps = self.cbn_layer(self.up_conv(f), shrink_maps, None)
|
||||
cbn_maps = F.sigmoid(cbn_maps)
|
||||
return {'maps': 0.5 * (base_maps + cbn_maps), 'cbn_maps': cbn_maps}
|
||||
@@ -0,0 +1,54 @@
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
|
||||
class CTCHead(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
out_channels=6625,
|
||||
fc_decay=0.0004,
|
||||
mid_channels=None,
|
||||
return_feats=False,
|
||||
**kwargs
|
||||
):
|
||||
super(CTCHead, self).__init__()
|
||||
if mid_channels is None:
|
||||
self.fc = nn.Linear(
|
||||
in_channels,
|
||||
out_channels,
|
||||
bias=True,
|
||||
)
|
||||
else:
|
||||
self.fc1 = nn.Linear(
|
||||
in_channels,
|
||||
mid_channels,
|
||||
bias=True,
|
||||
)
|
||||
self.fc2 = nn.Linear(
|
||||
mid_channels,
|
||||
out_channels,
|
||||
bias=True,
|
||||
)
|
||||
|
||||
self.out_channels = out_channels
|
||||
self.mid_channels = mid_channels
|
||||
self.return_feats = return_feats
|
||||
|
||||
def forward(self, x, labels=None):
|
||||
if self.mid_channels is None:
|
||||
predicts = self.fc(x)
|
||||
else:
|
||||
x = self.fc1(x)
|
||||
predicts = self.fc2(x)
|
||||
|
||||
if self.return_feats:
|
||||
result = (x, predicts)
|
||||
else:
|
||||
result = predicts
|
||||
|
||||
if not self.training:
|
||||
predicts = F.softmax(predicts, dim=2)
|
||||
result = predicts
|
||||
|
||||
return result
|
||||
@@ -0,0 +1,58 @@
|
||||
from torch import nn
|
||||
|
||||
from ..necks.rnn import Im2Seq, SequenceEncoder
|
||||
from .rec_ctc_head import CTCHead
|
||||
|
||||
|
||||
class FCTranspose(nn.Module):
|
||||
def __init__(self, in_channels, out_channels, only_transpose=False):
|
||||
super().__init__()
|
||||
self.only_transpose = only_transpose
|
||||
if not self.only_transpose:
|
||||
self.fc = nn.Linear(in_channels, out_channels, bias=False)
|
||||
|
||||
def forward(self, x):
|
||||
if self.only_transpose:
|
||||
return x.permute([0, 2, 1])
|
||||
else:
|
||||
return self.fc(x.permute([0, 2, 1]))
|
||||
|
||||
|
||||
class MultiHead(nn.Module):
|
||||
def __init__(self, in_channels, out_channels_list, **kwargs):
|
||||
super().__init__()
|
||||
self.head_list = kwargs.pop("head_list")
|
||||
|
||||
self.gtc_head = "sar"
|
||||
assert len(self.head_list) >= 2
|
||||
for idx, head_name in enumerate(self.head_list):
|
||||
name = list(head_name)[0]
|
||||
if name == "SARHead":
|
||||
pass
|
||||
|
||||
elif name == "NRTRHead":
|
||||
pass
|
||||
elif name == "CTCHead":
|
||||
# ctc neck
|
||||
self.encoder_reshape = Im2Seq(in_channels)
|
||||
neck_args = self.head_list[idx][name]["Neck"]
|
||||
encoder_type = neck_args.pop("name")
|
||||
self.ctc_encoder = SequenceEncoder(
|
||||
in_channels=in_channels, encoder_type=encoder_type, **neck_args
|
||||
)
|
||||
# ctc head
|
||||
head_args = self.head_list[idx][name].get("Head", {})
|
||||
if head_args is None:
|
||||
head_args = {}
|
||||
|
||||
self.ctc_head = CTCHead(
|
||||
in_channels=self.ctc_encoder.out_channels,
|
||||
out_channels=out_channels_list["CTCLabelDecode"],
|
||||
**head_args,
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(f"{name} is not supported in MultiHead yet")
|
||||
|
||||
def forward(self, x, data=None):
|
||||
ctc_encoder = self.ctc_encoder(x)
|
||||
return self.ctc_head(ctc_encoder)
|
||||
@@ -0,0 +1,29 @@
|
||||
# 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.
|
||||
|
||||
__all__ = ["build_neck"]
|
||||
|
||||
|
||||
def build_neck(config):
|
||||
from .db_fpn import DBFPN, LKPAN, RSEFPN
|
||||
from .rnn import SequenceEncoder
|
||||
|
||||
support_dict = ["DBFPN", "SequenceEncoder", "RSEFPN", "LKPAN"]
|
||||
|
||||
module_name = config.pop("name")
|
||||
assert module_name in support_dict, Exception(
|
||||
"neck only support {}".format(support_dict)
|
||||
)
|
||||
module_class = eval(module_name)(**config)
|
||||
return module_class
|
||||
@@ -0,0 +1,456 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
from ..backbones.det_mobilenet_v3 import SEModule
|
||||
from ..necks.intracl import IntraCLBlock
|
||||
|
||||
|
||||
def hard_swish(x, inplace=True):
|
||||
return x * F.relu6(x + 3.0, inplace=inplace) / 6.0
|
||||
|
||||
|
||||
class DSConv(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
padding,
|
||||
stride=1,
|
||||
groups=None,
|
||||
if_act=True,
|
||||
act="relu",
|
||||
**kwargs
|
||||
):
|
||||
super(DSConv, self).__init__()
|
||||
if groups == None:
|
||||
groups = in_channels
|
||||
self.if_act = if_act
|
||||
self.act = act
|
||||
self.conv1 = nn.Conv2d(
|
||||
in_channels=in_channels,
|
||||
out_channels=in_channels,
|
||||
kernel_size=kernel_size,
|
||||
stride=stride,
|
||||
padding=padding,
|
||||
groups=groups,
|
||||
bias=False,
|
||||
)
|
||||
|
||||
self.bn1 = nn.BatchNorm2d(in_channels)
|
||||
|
||||
self.conv2 = nn.Conv2d(
|
||||
in_channels=in_channels,
|
||||
out_channels=int(in_channels * 4),
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
bias=False,
|
||||
)
|
||||
|
||||
self.bn2 = nn.BatchNorm2d(int(in_channels * 4))
|
||||
|
||||
self.conv3 = nn.Conv2d(
|
||||
in_channels=int(in_channels * 4),
|
||||
out_channels=out_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
bias=False,
|
||||
)
|
||||
self._c = [in_channels, out_channels]
|
||||
if in_channels != out_channels:
|
||||
self.conv_end = nn.Conv2d(
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
bias=False,
|
||||
)
|
||||
|
||||
def forward(self, inputs):
|
||||
x = self.conv1(inputs)
|
||||
x = self.bn1(x)
|
||||
|
||||
x = self.conv2(x)
|
||||
x = self.bn2(x)
|
||||
if self.if_act:
|
||||
if self.act == "relu":
|
||||
x = F.relu(x)
|
||||
elif self.act == "hardswish":
|
||||
x = hard_swish(x)
|
||||
else:
|
||||
print(
|
||||
"The activation function({}) is selected incorrectly.".format(
|
||||
self.act
|
||||
)
|
||||
)
|
||||
exit()
|
||||
|
||||
x = self.conv3(x)
|
||||
if self._c[0] != self._c[1]:
|
||||
x = x + self.conv_end(inputs)
|
||||
return x
|
||||
|
||||
|
||||
class DBFPN(nn.Module):
|
||||
def __init__(self, in_channels, out_channels, use_asf=False, **kwargs):
|
||||
super(DBFPN, self).__init__()
|
||||
self.out_channels = out_channels
|
||||
self.use_asf = use_asf
|
||||
|
||||
self.in2_conv = nn.Conv2d(
|
||||
in_channels=in_channels[0],
|
||||
out_channels=self.out_channels,
|
||||
kernel_size=1,
|
||||
bias=False,
|
||||
)
|
||||
self.in3_conv = nn.Conv2d(
|
||||
in_channels=in_channels[1],
|
||||
out_channels=self.out_channels,
|
||||
kernel_size=1,
|
||||
bias=False,
|
||||
)
|
||||
self.in4_conv = nn.Conv2d(
|
||||
in_channels=in_channels[2],
|
||||
out_channels=self.out_channels,
|
||||
kernel_size=1,
|
||||
bias=False,
|
||||
)
|
||||
self.in5_conv = nn.Conv2d(
|
||||
in_channels=in_channels[3],
|
||||
out_channels=self.out_channels,
|
||||
kernel_size=1,
|
||||
bias=False,
|
||||
)
|
||||
self.p5_conv = nn.Conv2d(
|
||||
in_channels=self.out_channels,
|
||||
out_channels=self.out_channels // 4,
|
||||
kernel_size=3,
|
||||
padding=1,
|
||||
bias=False,
|
||||
)
|
||||
self.p4_conv = nn.Conv2d(
|
||||
in_channels=self.out_channels,
|
||||
out_channels=self.out_channels // 4,
|
||||
kernel_size=3,
|
||||
padding=1,
|
||||
bias=False,
|
||||
)
|
||||
self.p3_conv = nn.Conv2d(
|
||||
in_channels=self.out_channels,
|
||||
out_channels=self.out_channels // 4,
|
||||
kernel_size=3,
|
||||
padding=1,
|
||||
bias=False,
|
||||
)
|
||||
self.p2_conv = nn.Conv2d(
|
||||
in_channels=self.out_channels,
|
||||
out_channels=self.out_channels // 4,
|
||||
kernel_size=3,
|
||||
padding=1,
|
||||
bias=False,
|
||||
)
|
||||
|
||||
if self.use_asf is True:
|
||||
self.asf = ASFBlock(self.out_channels, self.out_channels // 4)
|
||||
|
||||
def forward(self, x):
|
||||
c2, c3, c4, c5 = x
|
||||
|
||||
in5 = self.in5_conv(c5)
|
||||
in4 = self.in4_conv(c4)
|
||||
in3 = self.in3_conv(c3)
|
||||
in2 = self.in2_conv(c2)
|
||||
|
||||
out4 = in4 + F.interpolate(
|
||||
in5,
|
||||
scale_factor=2,
|
||||
mode="nearest",
|
||||
) # align_mode=1) # 1/16
|
||||
out3 = in3 + F.interpolate(
|
||||
out4,
|
||||
scale_factor=2,
|
||||
mode="nearest",
|
||||
) # align_mode=1) # 1/8
|
||||
out2 = in2 + F.interpolate(
|
||||
out3,
|
||||
scale_factor=2,
|
||||
mode="nearest",
|
||||
) # align_mode=1) # 1/4
|
||||
|
||||
p5 = self.p5_conv(in5)
|
||||
p4 = self.p4_conv(out4)
|
||||
p3 = self.p3_conv(out3)
|
||||
p2 = self.p2_conv(out2)
|
||||
p5 = F.interpolate(
|
||||
p5,
|
||||
scale_factor=8,
|
||||
mode="nearest",
|
||||
) # align_mode=1)
|
||||
p4 = F.interpolate(
|
||||
p4,
|
||||
scale_factor=4,
|
||||
mode="nearest",
|
||||
) # align_mode=1)
|
||||
p3 = F.interpolate(
|
||||
p3,
|
||||
scale_factor=2,
|
||||
mode="nearest",
|
||||
) # align_mode=1)
|
||||
|
||||
fuse = torch.cat([p5, p4, p3, p2], dim=1)
|
||||
|
||||
if self.use_asf is True:
|
||||
fuse = self.asf(fuse, [p5, p4, p3, p2])
|
||||
|
||||
return fuse
|
||||
|
||||
|
||||
class RSELayer(nn.Module):
|
||||
def __init__(self, in_channels, out_channels, kernel_size, shortcut=True):
|
||||
super(RSELayer, self).__init__()
|
||||
self.out_channels = out_channels
|
||||
self.in_conv = nn.Conv2d(
|
||||
in_channels=in_channels,
|
||||
out_channels=self.out_channels,
|
||||
kernel_size=kernel_size,
|
||||
padding=int(kernel_size // 2),
|
||||
bias=False,
|
||||
)
|
||||
self.se_block = SEModule(self.out_channels)
|
||||
self.shortcut = shortcut
|
||||
|
||||
def forward(self, ins):
|
||||
x = self.in_conv(ins)
|
||||
if self.shortcut:
|
||||
out = x + self.se_block(x)
|
||||
else:
|
||||
out = self.se_block(x)
|
||||
return out
|
||||
|
||||
|
||||
class RSEFPN(nn.Module):
|
||||
def __init__(self, in_channels, out_channels, shortcut=True, **kwargs):
|
||||
super(RSEFPN, self).__init__()
|
||||
self.out_channels = out_channels
|
||||
self.ins_conv = nn.ModuleList()
|
||||
self.inp_conv = nn.ModuleList()
|
||||
self.intracl = False
|
||||
if "intracl" in kwargs.keys() and kwargs["intracl"] is True:
|
||||
self.intracl = kwargs["intracl"]
|
||||
self.incl1 = IntraCLBlock(self.out_channels // 4, reduce_factor=2)
|
||||
self.incl2 = IntraCLBlock(self.out_channels // 4, reduce_factor=2)
|
||||
self.incl3 = IntraCLBlock(self.out_channels // 4, reduce_factor=2)
|
||||
self.incl4 = IntraCLBlock(self.out_channels // 4, reduce_factor=2)
|
||||
|
||||
for i in range(len(in_channels)):
|
||||
self.ins_conv.append(
|
||||
RSELayer(in_channels[i], out_channels, kernel_size=1, shortcut=shortcut)
|
||||
)
|
||||
self.inp_conv.append(
|
||||
RSELayer(
|
||||
out_channels, out_channels // 4, kernel_size=3, shortcut=shortcut
|
||||
)
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
c2, c3, c4, c5 = x
|
||||
|
||||
in5 = self.ins_conv[3](c5)
|
||||
in4 = self.ins_conv[2](c4)
|
||||
in3 = self.ins_conv[1](c3)
|
||||
in2 = self.ins_conv[0](c2)
|
||||
|
||||
out4 = in4 + F.interpolate(in5, scale_factor=2, mode="nearest") # 1/16
|
||||
out3 = in3 + F.interpolate(out4, scale_factor=2, mode="nearest") # 1/8
|
||||
out2 = in2 + F.interpolate(out3, scale_factor=2, mode="nearest") # 1/4
|
||||
|
||||
p5 = self.inp_conv[3](in5)
|
||||
p4 = self.inp_conv[2](out4)
|
||||
p3 = self.inp_conv[1](out3)
|
||||
p2 = self.inp_conv[0](out2)
|
||||
|
||||
if self.intracl is True:
|
||||
p5 = self.incl4(p5)
|
||||
p4 = self.incl3(p4)
|
||||
p3 = self.incl2(p3)
|
||||
p2 = self.incl1(p2)
|
||||
|
||||
p5 = F.interpolate(p5, scale_factor=8, mode="nearest")
|
||||
p4 = F.interpolate(p4, scale_factor=4, mode="nearest")
|
||||
p3 = F.interpolate(p3, scale_factor=2, mode="nearest")
|
||||
|
||||
fuse = torch.cat([p5, p4, p3, p2], dim=1)
|
||||
return fuse
|
||||
|
||||
|
||||
class LKPAN(nn.Module):
|
||||
def __init__(self, in_channels, out_channels, mode="large", **kwargs):
|
||||
super(LKPAN, self).__init__()
|
||||
self.out_channels = out_channels
|
||||
|
||||
self.ins_conv = nn.ModuleList()
|
||||
self.inp_conv = nn.ModuleList()
|
||||
# pan head
|
||||
self.pan_head_conv = nn.ModuleList()
|
||||
self.pan_lat_conv = nn.ModuleList()
|
||||
|
||||
if mode.lower() == "lite":
|
||||
p_layer = DSConv
|
||||
elif mode.lower() == "large":
|
||||
p_layer = nn.Conv2d
|
||||
else:
|
||||
raise ValueError(
|
||||
"mode can only be one of ['lite', 'large'], but received {}".format(
|
||||
mode
|
||||
)
|
||||
)
|
||||
|
||||
for i in range(len(in_channels)):
|
||||
self.ins_conv.append(
|
||||
nn.Conv2d(
|
||||
in_channels=in_channels[i],
|
||||
out_channels=self.out_channels,
|
||||
kernel_size=1,
|
||||
bias=False,
|
||||
)
|
||||
)
|
||||
|
||||
self.inp_conv.append(
|
||||
p_layer(
|
||||
in_channels=self.out_channels,
|
||||
out_channels=self.out_channels // 4,
|
||||
kernel_size=9,
|
||||
padding=4,
|
||||
bias=False,
|
||||
)
|
||||
)
|
||||
|
||||
if i > 0:
|
||||
self.pan_head_conv.append(
|
||||
nn.Conv2d(
|
||||
in_channels=self.out_channels // 4,
|
||||
out_channels=self.out_channels // 4,
|
||||
kernel_size=3,
|
||||
padding=1,
|
||||
stride=2,
|
||||
bias=False,
|
||||
)
|
||||
)
|
||||
self.pan_lat_conv.append(
|
||||
p_layer(
|
||||
in_channels=self.out_channels // 4,
|
||||
out_channels=self.out_channels // 4,
|
||||
kernel_size=9,
|
||||
padding=4,
|
||||
bias=False,
|
||||
)
|
||||
)
|
||||
self.intracl = False
|
||||
if "intracl" in kwargs.keys() and kwargs["intracl"] is True:
|
||||
self.intracl = kwargs["intracl"]
|
||||
self.incl1 = IntraCLBlock(self.out_channels // 4, reduce_factor=2)
|
||||
self.incl2 = IntraCLBlock(self.out_channels // 4, reduce_factor=2)
|
||||
self.incl3 = IntraCLBlock(self.out_channels // 4, reduce_factor=2)
|
||||
self.incl4 = IntraCLBlock(self.out_channels // 4, reduce_factor=2)
|
||||
|
||||
def forward(self, x):
|
||||
c2, c3, c4, c5 = x
|
||||
|
||||
in5 = self.ins_conv[3](c5)
|
||||
in4 = self.ins_conv[2](c4)
|
||||
in3 = self.ins_conv[1](c3)
|
||||
in2 = self.ins_conv[0](c2)
|
||||
|
||||
out4 = in4 + F.interpolate(in5, scale_factor=2, mode="nearest") # 1/16
|
||||
out3 = in3 + F.interpolate(out4, scale_factor=2, mode="nearest") # 1/8
|
||||
out2 = in2 + F.interpolate(out3, scale_factor=2, mode="nearest") # 1/4
|
||||
|
||||
f5 = self.inp_conv[3](in5)
|
||||
f4 = self.inp_conv[2](out4)
|
||||
f3 = self.inp_conv[1](out3)
|
||||
f2 = self.inp_conv[0](out2)
|
||||
|
||||
pan3 = f3 + self.pan_head_conv[0](f2)
|
||||
pan4 = f4 + self.pan_head_conv[1](pan3)
|
||||
pan5 = f5 + self.pan_head_conv[2](pan4)
|
||||
|
||||
p2 = self.pan_lat_conv[0](f2)
|
||||
p3 = self.pan_lat_conv[1](pan3)
|
||||
p4 = self.pan_lat_conv[2](pan4)
|
||||
p5 = self.pan_lat_conv[3](pan5)
|
||||
|
||||
if self.intracl is True:
|
||||
p5 = self.incl4(p5)
|
||||
p4 = self.incl3(p4)
|
||||
p3 = self.incl2(p3)
|
||||
p2 = self.incl1(p2)
|
||||
|
||||
p5 = F.interpolate(p5, scale_factor=8, mode="nearest")
|
||||
p4 = F.interpolate(p4, scale_factor=4, mode="nearest")
|
||||
p3 = F.interpolate(p3, scale_factor=2, mode="nearest")
|
||||
|
||||
fuse = torch.cat([p5, p4, p3, p2], dim=1)
|
||||
return fuse
|
||||
|
||||
|
||||
class ASFBlock(nn.Module):
|
||||
"""
|
||||
This code is refered from:
|
||||
https://github.com/MhLiao/DB/blob/master/decoders/feature_attention.py
|
||||
"""
|
||||
|
||||
def __init__(self, in_channels, inter_channels, out_features_num=4):
|
||||
"""
|
||||
Adaptive Scale Fusion (ASF) block of DBNet++
|
||||
Args:
|
||||
in_channels: the number of channels in the input data
|
||||
inter_channels: the number of middle channels
|
||||
out_features_num: the number of fused stages
|
||||
"""
|
||||
super(ASFBlock, self).__init__()
|
||||
self.in_channels = in_channels
|
||||
self.inter_channels = inter_channels
|
||||
self.out_features_num = out_features_num
|
||||
self.conv = nn.Conv2d(in_channels, inter_channels, 3, padding=1)
|
||||
|
||||
self.spatial_scale = nn.Sequential(
|
||||
# Nx1xHxW
|
||||
nn.Conv2d(
|
||||
in_channels=1,
|
||||
out_channels=1,
|
||||
kernel_size=3,
|
||||
bias=False,
|
||||
padding=1,
|
||||
),
|
||||
nn.ReLU(),
|
||||
nn.Conv2d(
|
||||
in_channels=1,
|
||||
out_channels=1,
|
||||
kernel_size=1,
|
||||
bias=False,
|
||||
),
|
||||
nn.Sigmoid(),
|
||||
)
|
||||
|
||||
self.channel_scale = nn.Sequential(
|
||||
nn.Conv2d(
|
||||
in_channels=inter_channels,
|
||||
out_channels=out_features_num,
|
||||
kernel_size=1,
|
||||
bias=False,
|
||||
),
|
||||
nn.Sigmoid(),
|
||||
)
|
||||
|
||||
def forward(self, fuse_features, features_list):
|
||||
fuse_features = self.conv(fuse_features)
|
||||
spatial_x = torch.mean(fuse_features, dim=1, keepdim=True)
|
||||
attention_scores = self.spatial_scale(spatial_x) + fuse_features
|
||||
attention_scores = self.channel_scale(attention_scores)
|
||||
assert len(features_list) == self.out_features_num
|
||||
|
||||
out_list = []
|
||||
for i in range(self.out_features_num):
|
||||
out_list.append(attention_scores[:, i : i + 1] * features_list[i])
|
||||
return torch.cat(out_list, dim=1)
|
||||
@@ -0,0 +1,117 @@
|
||||
from torch import nn
|
||||
|
||||
|
||||
class IntraCLBlock(nn.Module):
|
||||
def __init__(self, in_channels=96, reduce_factor=4):
|
||||
super(IntraCLBlock, self).__init__()
|
||||
self.channels = in_channels
|
||||
self.rf = reduce_factor
|
||||
self.conv1x1_reduce_channel = nn.Conv2d(
|
||||
self.channels, self.channels // self.rf, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
self.conv1x1_return_channel = nn.Conv2d(
|
||||
self.channels // self.rf, self.channels, kernel_size=1, stride=1, padding=0
|
||||
)
|
||||
|
||||
self.v_layer_7x1 = nn.Conv2d(
|
||||
self.channels // self.rf,
|
||||
self.channels // self.rf,
|
||||
kernel_size=(7, 1),
|
||||
stride=(1, 1),
|
||||
padding=(3, 0),
|
||||
)
|
||||
self.v_layer_5x1 = nn.Conv2d(
|
||||
self.channels // self.rf,
|
||||
self.channels // self.rf,
|
||||
kernel_size=(5, 1),
|
||||
stride=(1, 1),
|
||||
padding=(2, 0),
|
||||
)
|
||||
self.v_layer_3x1 = nn.Conv2d(
|
||||
self.channels // self.rf,
|
||||
self.channels // self.rf,
|
||||
kernel_size=(3, 1),
|
||||
stride=(1, 1),
|
||||
padding=(1, 0),
|
||||
)
|
||||
|
||||
self.q_layer_1x7 = nn.Conv2d(
|
||||
self.channels // self.rf,
|
||||
self.channels // self.rf,
|
||||
kernel_size=(1, 7),
|
||||
stride=(1, 1),
|
||||
padding=(0, 3),
|
||||
)
|
||||
self.q_layer_1x5 = nn.Conv2d(
|
||||
self.channels // self.rf,
|
||||
self.channels // self.rf,
|
||||
kernel_size=(1, 5),
|
||||
stride=(1, 1),
|
||||
padding=(0, 2),
|
||||
)
|
||||
self.q_layer_1x3 = nn.Conv2d(
|
||||
self.channels // self.rf,
|
||||
self.channels // self.rf,
|
||||
kernel_size=(1, 3),
|
||||
stride=(1, 1),
|
||||
padding=(0, 1),
|
||||
)
|
||||
|
||||
# base
|
||||
self.c_layer_7x7 = nn.Conv2d(
|
||||
self.channels // self.rf,
|
||||
self.channels // self.rf,
|
||||
kernel_size=(7, 7),
|
||||
stride=(1, 1),
|
||||
padding=(3, 3),
|
||||
)
|
||||
self.c_layer_5x5 = nn.Conv2d(
|
||||
self.channels // self.rf,
|
||||
self.channels // self.rf,
|
||||
kernel_size=(5, 5),
|
||||
stride=(1, 1),
|
||||
padding=(2, 2),
|
||||
)
|
||||
self.c_layer_3x3 = nn.Conv2d(
|
||||
self.channels // self.rf,
|
||||
self.channels // self.rf,
|
||||
kernel_size=(3, 3),
|
||||
stride=(1, 1),
|
||||
padding=(1, 1),
|
||||
)
|
||||
|
||||
self.bn = nn.BatchNorm2d(self.channels)
|
||||
self.relu = nn.ReLU()
|
||||
|
||||
def forward(self, x):
|
||||
x_new = self.conv1x1_reduce_channel(x)
|
||||
|
||||
x_7_c = self.c_layer_7x7(x_new)
|
||||
x_7_v = self.v_layer_7x1(x_new)
|
||||
x_7_q = self.q_layer_1x7(x_new)
|
||||
x_7 = x_7_c + x_7_v + x_7_q
|
||||
|
||||
x_5_c = self.c_layer_5x5(x_7)
|
||||
x_5_v = self.v_layer_5x1(x_7)
|
||||
x_5_q = self.q_layer_1x5(x_7)
|
||||
x_5 = x_5_c + x_5_v + x_5_q
|
||||
|
||||
x_3_c = self.c_layer_3x3(x_5)
|
||||
x_3_v = self.v_layer_3x1(x_5)
|
||||
x_3_q = self.q_layer_1x3(x_5)
|
||||
x_3 = x_3_c + x_3_v + x_3_q
|
||||
|
||||
x_relation = self.conv1x1_return_channel(x_3)
|
||||
|
||||
x_relation = self.bn(x_relation)
|
||||
x_relation = self.relu(x_relation)
|
||||
|
||||
return x + x_relation
|
||||
|
||||
|
||||
def build_intraclblock_list(num_block):
|
||||
IntraCLBlock_list = nn.ModuleList()
|
||||
for i in range(num_block):
|
||||
IntraCLBlock_list.append(IntraCLBlock())
|
||||
|
||||
return IntraCLBlock_list
|
||||
@@ -0,0 +1,241 @@
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from ..backbones.rec_svtrnet import Block, ConvBNLayer
|
||||
|
||||
|
||||
class Im2Seq(nn.Module):
|
||||
def __init__(self, in_channels, **kwargs):
|
||||
super().__init__()
|
||||
self.out_channels = in_channels
|
||||
|
||||
# def forward(self, x):
|
||||
# B, C, H, W = x.shape
|
||||
# # assert H == 1
|
||||
# x = x.squeeze(dim=2)
|
||||
# # x = x.transpose([0, 2, 1]) # paddle (NTC)(batch, width, channels)
|
||||
# x = x.permute(0, 2, 1)
|
||||
# return x
|
||||
|
||||
def forward(self, x):
|
||||
B, C, H, W = x.shape
|
||||
# 处理四维张量,将空间维度展平为序列
|
||||
if H == 1:
|
||||
# 原来的处理逻辑,适用于H=1的情况
|
||||
x = x.squeeze(dim=2)
|
||||
x = x.permute(0, 2, 1) # (B, W, C)
|
||||
else:
|
||||
# 处理H不为1的情况
|
||||
x = x.permute(0, 2, 3, 1) # (B, H, W, C)
|
||||
x = x.reshape(B, H * W, C) # (B, H*W, C)
|
||||
|
||||
return x
|
||||
|
||||
class EncoderWithRNN_(nn.Module):
|
||||
def __init__(self, in_channels, hidden_size):
|
||||
super(EncoderWithRNN_, self).__init__()
|
||||
self.out_channels = hidden_size * 2
|
||||
self.rnn1 = nn.LSTM(
|
||||
in_channels,
|
||||
hidden_size,
|
||||
bidirectional=False,
|
||||
batch_first=True,
|
||||
num_layers=2,
|
||||
)
|
||||
self.rnn2 = nn.LSTM(
|
||||
in_channels,
|
||||
hidden_size,
|
||||
bidirectional=False,
|
||||
batch_first=True,
|
||||
num_layers=2,
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
self.rnn1.flatten_parameters()
|
||||
self.rnn2.flatten_parameters()
|
||||
out1, h1 = self.rnn1(x)
|
||||
out2, h2 = self.rnn2(torch.flip(x, [1]))
|
||||
return torch.cat([out1, torch.flip(out2, [1])], 2)
|
||||
|
||||
|
||||
class EncoderWithRNN(nn.Module):
|
||||
def __init__(self, in_channels, hidden_size):
|
||||
super(EncoderWithRNN, self).__init__()
|
||||
self.out_channels = hidden_size * 2
|
||||
self.lstm = nn.LSTM(
|
||||
in_channels, hidden_size, num_layers=2, batch_first=True, bidirectional=True
|
||||
) # batch_first:=True
|
||||
|
||||
def forward(self, x):
|
||||
x, _ = self.lstm(x)
|
||||
return x
|
||||
|
||||
|
||||
class EncoderWithFC(nn.Module):
|
||||
def __init__(self, in_channels, hidden_size):
|
||||
super(EncoderWithFC, self).__init__()
|
||||
self.out_channels = hidden_size
|
||||
self.fc = nn.Linear(
|
||||
in_channels,
|
||||
hidden_size,
|
||||
bias=True,
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.fc(x)
|
||||
return x
|
||||
|
||||
|
||||
class EncoderWithSVTR(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
dims=64, # XS
|
||||
depth=2,
|
||||
hidden_dims=120,
|
||||
use_guide=False,
|
||||
num_heads=8,
|
||||
qkv_bias=True,
|
||||
mlp_ratio=2.0,
|
||||
drop_rate=0.1,
|
||||
kernel_size=[3, 3],
|
||||
attn_drop_rate=0.1,
|
||||
drop_path=0.0,
|
||||
qk_scale=None,
|
||||
):
|
||||
super(EncoderWithSVTR, self).__init__()
|
||||
self.depth = depth
|
||||
self.use_guide = use_guide
|
||||
self.conv1 = ConvBNLayer(
|
||||
in_channels,
|
||||
in_channels // 8,
|
||||
kernel_size=kernel_size,
|
||||
padding=[kernel_size[0] // 2, kernel_size[1] // 2],
|
||||
act="swish",
|
||||
)
|
||||
self.conv2 = ConvBNLayer(
|
||||
in_channels // 8, hidden_dims, kernel_size=1, act="swish"
|
||||
)
|
||||
|
||||
self.svtr_block = nn.ModuleList(
|
||||
[
|
||||
Block(
|
||||
dim=hidden_dims,
|
||||
num_heads=num_heads,
|
||||
mixer="Global",
|
||||
HW=None,
|
||||
mlp_ratio=mlp_ratio,
|
||||
qkv_bias=qkv_bias,
|
||||
qk_scale=qk_scale,
|
||||
drop=drop_rate,
|
||||
act_layer="swish",
|
||||
attn_drop=attn_drop_rate,
|
||||
drop_path=drop_path,
|
||||
norm_layer="nn.LayerNorm",
|
||||
epsilon=1e-05,
|
||||
prenorm=False,
|
||||
)
|
||||
for i in range(depth)
|
||||
]
|
||||
)
|
||||
self.norm = nn.LayerNorm(hidden_dims, eps=1e-6)
|
||||
self.conv3 = ConvBNLayer(hidden_dims, in_channels, kernel_size=1, act="swish")
|
||||
# last conv-nxn, the input is concat of input tensor and conv3 output tensor
|
||||
self.conv4 = ConvBNLayer(
|
||||
2 * in_channels, in_channels // 8, padding=1, act="swish"
|
||||
)
|
||||
|
||||
self.conv1x1 = ConvBNLayer(in_channels // 8, dims, kernel_size=1, act="swish")
|
||||
self.out_channels = dims
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
# weight initialization
|
||||
if isinstance(m, nn.Conv2d):
|
||||
nn.init.kaiming_normal_(m.weight, mode="fan_out")
|
||||
if m.bias is not None:
|
||||
nn.init.zeros_(m.bias)
|
||||
elif isinstance(m, nn.BatchNorm2d):
|
||||
nn.init.ones_(m.weight)
|
||||
nn.init.zeros_(m.bias)
|
||||
elif isinstance(m, nn.Linear):
|
||||
nn.init.normal_(m.weight, 0, 0.01)
|
||||
if m.bias is not None:
|
||||
nn.init.zeros_(m.bias)
|
||||
elif isinstance(m, nn.ConvTranspose2d):
|
||||
nn.init.kaiming_normal_(m.weight, mode="fan_out")
|
||||
if m.bias is not None:
|
||||
nn.init.zeros_(m.bias)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.ones_(m.weight)
|
||||
nn.init.zeros_(m.bias)
|
||||
|
||||
def forward(self, x):
|
||||
# for use guide
|
||||
if self.use_guide:
|
||||
z = x.clone()
|
||||
z.stop_gradient = True
|
||||
else:
|
||||
z = x
|
||||
# for short cut
|
||||
h = z
|
||||
# reduce dim
|
||||
z = self.conv1(z)
|
||||
z = self.conv2(z)
|
||||
# SVTR global block
|
||||
B, C, H, W = z.shape
|
||||
z = z.flatten(2).permute(0, 2, 1)
|
||||
|
||||
for blk in self.svtr_block:
|
||||
z = blk(z)
|
||||
|
||||
z = self.norm(z)
|
||||
# last stage
|
||||
z = z.reshape([-1, H, W, C]).permute(0, 3, 1, 2)
|
||||
z = self.conv3(z)
|
||||
z = torch.cat((h, z), dim=1)
|
||||
z = self.conv1x1(self.conv4(z))
|
||||
|
||||
return z
|
||||
|
||||
|
||||
class SequenceEncoder(nn.Module):
|
||||
def __init__(self, in_channels, encoder_type, hidden_size=48, **kwargs):
|
||||
super(SequenceEncoder, self).__init__()
|
||||
self.encoder_reshape = Im2Seq(in_channels)
|
||||
self.out_channels = self.encoder_reshape.out_channels
|
||||
self.encoder_type = encoder_type
|
||||
if encoder_type == "reshape":
|
||||
self.only_reshape = True
|
||||
else:
|
||||
support_encoder_dict = {
|
||||
"reshape": Im2Seq,
|
||||
"fc": EncoderWithFC,
|
||||
"rnn": EncoderWithRNN,
|
||||
"svtr": EncoderWithSVTR,
|
||||
}
|
||||
assert encoder_type in support_encoder_dict, "{} must in {}".format(
|
||||
encoder_type, support_encoder_dict.keys()
|
||||
)
|
||||
|
||||
if encoder_type == "svtr":
|
||||
self.encoder = support_encoder_dict[encoder_type](
|
||||
self.encoder_reshape.out_channels, **kwargs
|
||||
)
|
||||
else:
|
||||
self.encoder = support_encoder_dict[encoder_type](
|
||||
self.encoder_reshape.out_channels, hidden_size
|
||||
)
|
||||
self.out_channels = self.encoder.out_channels
|
||||
self.only_reshape = False
|
||||
|
||||
def forward(self, x):
|
||||
if self.encoder_type != "svtr":
|
||||
x = self.encoder_reshape(x)
|
||||
if not self.only_reshape:
|
||||
x = self.encoder(x)
|
||||
return x
|
||||
else:
|
||||
x = self.encoder(x)
|
||||
x = self.encoder_reshape(x)
|
||||
return x
|
||||
@@ -0,0 +1,33 @@
|
||||
|
||||
from __future__ import absolute_import
|
||||
from __future__ import division
|
||||
from __future__ import print_function
|
||||
from __future__ import unicode_literals
|
||||
|
||||
import copy
|
||||
|
||||
__all__ = ['build_post_process']
|
||||
|
||||
|
||||
def build_post_process(config, global_config=None):
|
||||
from .db_postprocess import DBPostProcess
|
||||
from .rec_postprocess import CTCLabelDecode, AttnLabelDecode, SRNLabelDecode, TableLabelDecode, \
|
||||
NRTRLabelDecode, SARLabelDecode, ViTSTRLabelDecode, RFLLabelDecode
|
||||
from .cls_postprocess import ClsPostProcess
|
||||
from .rec_postprocess import CANLabelDecode
|
||||
|
||||
support_dict = [
|
||||
'DBPostProcess', 'CTCLabelDecode',
|
||||
'AttnLabelDecode', 'ClsPostProcess', 'SRNLabelDecode',
|
||||
'TableLabelDecode', 'NRTRLabelDecode', 'SARLabelDecode',
|
||||
'ViTSTRLabelDecode','CANLabelDecode', 'RFLLabelDecode'
|
||||
]
|
||||
|
||||
config = copy.deepcopy(config)
|
||||
module_name = config.pop('name')
|
||||
if global_config is not None:
|
||||
config.update(global_config)
|
||||
assert module_name in support_dict, Exception(
|
||||
'post process only support {}, but got {}'.format(support_dict, module_name))
|
||||
module_class = eval(module_name)(**config)
|
||||
return module_class
|
||||
+20
@@ -0,0 +1,20 @@
|
||||
import torch
|
||||
|
||||
|
||||
class ClsPostProcess(object):
|
||||
""" Convert between text-label and text-index """
|
||||
|
||||
def __init__(self, label_list, **kwargs):
|
||||
super(ClsPostProcess, self).__init__()
|
||||
self.label_list = label_list
|
||||
|
||||
def __call__(self, preds, label=None, *args, **kwargs):
|
||||
if isinstance(preds, torch.Tensor):
|
||||
preds = preds.cpu().numpy()
|
||||
pred_idxs = preds.argmax(axis=1)
|
||||
decode_out = [(self.label_list[idx], preds[i, idx])
|
||||
for i, idx in enumerate(pred_idxs)]
|
||||
if label is None:
|
||||
return decode_out
|
||||
label = [(self.label_list[idx], 1.0) for idx in label]
|
||||
return decode_out, label
|
||||
+179
@@ -0,0 +1,179 @@
|
||||
"""
|
||||
This code is refered from:
|
||||
https://github.com/WenmuZhou/DBNet.pytorch/blob/master/post_processing/seg_detector_representer.py
|
||||
"""
|
||||
from __future__ import absolute_import
|
||||
from __future__ import division
|
||||
from __future__ import print_function
|
||||
|
||||
import numpy as np
|
||||
import cv2
|
||||
import torch
|
||||
from shapely.geometry import Polygon
|
||||
import pyclipper
|
||||
|
||||
|
||||
class DBPostProcess(object):
|
||||
"""
|
||||
The post process for Differentiable Binarization (DB).
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
thresh=0.3,
|
||||
box_thresh=0.7,
|
||||
max_candidates=1000,
|
||||
unclip_ratio=2.0,
|
||||
use_dilation=False,
|
||||
score_mode="fast",
|
||||
**kwargs):
|
||||
self.thresh = thresh
|
||||
self.box_thresh = box_thresh
|
||||
self.max_candidates = max_candidates
|
||||
self.unclip_ratio = unclip_ratio
|
||||
self.min_size = 3
|
||||
self.score_mode = score_mode
|
||||
assert score_mode in [
|
||||
"slow", "fast"
|
||||
], "Score mode must be in [slow, fast] but got: {}".format(score_mode)
|
||||
|
||||
self.dilation_kernel = None if not use_dilation else np.array(
|
||||
[[1, 1], [1, 1]])
|
||||
|
||||
def boxes_from_bitmap(self, pred, _bitmap, dest_width, dest_height):
|
||||
'''
|
||||
_bitmap: single map with shape (1, H, W),
|
||||
whose values are binarized as {0, 1}
|
||||
'''
|
||||
|
||||
bitmap = _bitmap
|
||||
height, width = bitmap.shape
|
||||
|
||||
outs = cv2.findContours((bitmap * 255).astype(np.uint8), cv2.RETR_LIST,
|
||||
cv2.CHAIN_APPROX_SIMPLE)
|
||||
if len(outs) == 3:
|
||||
img, contours, _ = outs[0], outs[1], outs[2]
|
||||
elif len(outs) == 2:
|
||||
contours, _ = outs[0], outs[1]
|
||||
|
||||
num_contours = min(len(contours), self.max_candidates)
|
||||
|
||||
boxes = []
|
||||
scores = []
|
||||
for index in range(num_contours):
|
||||
contour = contours[index]
|
||||
points, sside = self.get_mini_boxes(contour)
|
||||
if sside < self.min_size:
|
||||
continue
|
||||
points = np.array(points)
|
||||
if self.score_mode == "fast":
|
||||
score = self.box_score_fast(pred, points.reshape(-1, 2))
|
||||
else:
|
||||
score = self.box_score_slow(pred, contour)
|
||||
if self.box_thresh > score:
|
||||
continue
|
||||
|
||||
box = self.unclip(points).reshape(-1, 1, 2)
|
||||
box, sside = self.get_mini_boxes(box)
|
||||
if sside < self.min_size + 2:
|
||||
continue
|
||||
box = np.array(box)
|
||||
|
||||
box[:, 0] = np.clip(
|
||||
np.round(box[:, 0] / width * dest_width), 0, dest_width)
|
||||
box[:, 1] = np.clip(
|
||||
np.round(box[:, 1] / height * dest_height), 0, dest_height)
|
||||
boxes.append(box.astype(np.int16))
|
||||
scores.append(score)
|
||||
return np.array(boxes, dtype=np.int16), scores
|
||||
|
||||
def unclip(self, box):
|
||||
unclip_ratio = self.unclip_ratio
|
||||
poly = Polygon(box)
|
||||
distance = poly.area * unclip_ratio / poly.length
|
||||
offset = pyclipper.PyclipperOffset()
|
||||
offset.AddPath(box, pyclipper.JT_ROUND, pyclipper.ET_CLOSEDPOLYGON)
|
||||
expanded = np.array(offset.Execute(distance))
|
||||
return expanded
|
||||
|
||||
def get_mini_boxes(self, contour):
|
||||
bounding_box = cv2.minAreaRect(contour)
|
||||
points = sorted(list(cv2.boxPoints(bounding_box)), key=lambda x: x[0])
|
||||
|
||||
index_1, index_2, index_3, index_4 = 0, 1, 2, 3
|
||||
if points[1][1] > points[0][1]:
|
||||
index_1 = 0
|
||||
index_4 = 1
|
||||
else:
|
||||
index_1 = 1
|
||||
index_4 = 0
|
||||
if points[3][1] > points[2][1]:
|
||||
index_2 = 2
|
||||
index_3 = 3
|
||||
else:
|
||||
index_2 = 3
|
||||
index_3 = 2
|
||||
|
||||
box = [
|
||||
points[index_1], points[index_2], points[index_3], points[index_4]
|
||||
]
|
||||
return box, min(bounding_box[1])
|
||||
|
||||
def box_score_fast(self, bitmap, _box):
|
||||
'''
|
||||
box_score_fast: use bbox mean score as the mean score
|
||||
'''
|
||||
h, w = bitmap.shape[:2]
|
||||
box = _box.copy()
|
||||
xmin = np.clip(np.floor(box[:, 0].min()).astype(np.int64), 0, w - 1)
|
||||
xmax = np.clip(np.ceil(box[:, 0].max()).astype(np.int64), 0, w - 1)
|
||||
ymin = np.clip(np.floor(box[:, 1].min()).astype(np.int64), 0, h - 1)
|
||||
ymax = np.clip(np.ceil(box[:, 1].max()).astype(np.int64), 0, h - 1)
|
||||
|
||||
mask = np.zeros((ymax - ymin + 1, xmax - xmin + 1), dtype=np.uint8)
|
||||
box[:, 0] = box[:, 0] - xmin
|
||||
box[:, 1] = box[:, 1] - ymin
|
||||
cv2.fillPoly(mask, box.reshape(1, -1, 2).astype(np.int32), 1)
|
||||
return cv2.mean(bitmap[ymin:ymax + 1, xmin:xmax + 1], mask)[0]
|
||||
|
||||
def box_score_slow(self, bitmap, contour):
|
||||
'''
|
||||
box_score_slow: use polyon mean score as the mean score
|
||||
'''
|
||||
h, w = bitmap.shape[:2]
|
||||
contour = contour.copy()
|
||||
contour = np.reshape(contour, (-1, 2))
|
||||
|
||||
xmin = np.clip(np.min(contour[:, 0]), 0, w - 1)
|
||||
xmax = np.clip(np.max(contour[:, 0]), 0, w - 1)
|
||||
ymin = np.clip(np.min(contour[:, 1]), 0, h - 1)
|
||||
ymax = np.clip(np.max(contour[:, 1]), 0, h - 1)
|
||||
|
||||
mask = np.zeros((ymax - ymin + 1, xmax - xmin + 1), dtype=np.uint8)
|
||||
|
||||
contour[:, 0] = contour[:, 0] - xmin
|
||||
contour[:, 1] = contour[:, 1] - ymin
|
||||
|
||||
cv2.fillPoly(mask, contour.reshape(1, -1, 2).astype(np.int32), 1)
|
||||
return cv2.mean(bitmap[ymin:ymax + 1, xmin:xmax + 1], mask)[0]
|
||||
|
||||
def __call__(self, outs_dict, shape_list):
|
||||
pred = outs_dict['maps']
|
||||
if isinstance(pred, torch.Tensor):
|
||||
pred = pred.cpu().numpy()
|
||||
pred = pred[:, 0, :, :]
|
||||
segmentation = pred > self.thresh
|
||||
|
||||
boxes_batch = []
|
||||
for batch_index in range(pred.shape[0]):
|
||||
src_h, src_w, ratio_h, ratio_w = shape_list[batch_index]
|
||||
if self.dilation_kernel is not None:
|
||||
mask = cv2.dilate(
|
||||
np.array(segmentation[batch_index]).astype(np.uint8),
|
||||
self.dilation_kernel)
|
||||
else:
|
||||
mask = segmentation[batch_index]
|
||||
boxes, scores = self.boxes_from_bitmap(pred[batch_index], mask,
|
||||
src_w, src_h)
|
||||
|
||||
boxes_batch.append({'points': boxes})
|
||||
return boxes_batch
|
||||
+690
@@ -0,0 +1,690 @@
|
||||
# copyright (c) 2020 PaddlePaddle Authors. All Rights Reserve.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
|
||||
class BaseRecLabelDecode(object):
|
||||
""" Convert between text-label and text-index """
|
||||
|
||||
def __init__(self,
|
||||
character_dict_path=None,
|
||||
use_space_char=False):
|
||||
|
||||
self.beg_str = "sos"
|
||||
self.end_str = "eos"
|
||||
|
||||
self.character_str = []
|
||||
if character_dict_path is None:
|
||||
self.character_str = "0123456789abcdefghijklmnopqrstuvwxyz"
|
||||
dict_character = list(self.character_str)
|
||||
else:
|
||||
with open(character_dict_path, "rb") as fin:
|
||||
lines = fin.readlines()
|
||||
for line in lines:
|
||||
line = line.decode('utf-8').strip("\n").strip("\r\n")
|
||||
self.character_str.append(line)
|
||||
if use_space_char:
|
||||
self.character_str.append(" ")
|
||||
dict_character = list(self.character_str)
|
||||
|
||||
dict_character = self.add_special_char(dict_character)
|
||||
self.dict = {}
|
||||
for i, char in enumerate(dict_character):
|
||||
self.dict[char] = i
|
||||
self.character = dict_character
|
||||
|
||||
def add_special_char(self, dict_character):
|
||||
return dict_character
|
||||
|
||||
def decode(self, text_index, text_prob=None, is_remove_duplicate=False):
|
||||
""" convert text-index into text-label. """
|
||||
result_list = []
|
||||
ignored_tokens = self.get_ignored_tokens()
|
||||
batch_size = len(text_index)
|
||||
for batch_idx in range(batch_size):
|
||||
char_list = []
|
||||
conf_list = []
|
||||
for idx in range(len(text_index[batch_idx])):
|
||||
if text_index[batch_idx][idx] in ignored_tokens:
|
||||
continue
|
||||
if is_remove_duplicate:
|
||||
# only for predict
|
||||
if idx > 0 and text_index[batch_idx][idx - 1] == text_index[
|
||||
batch_idx][idx]:
|
||||
continue
|
||||
char_list.append(self.character[int(text_index[batch_idx][
|
||||
idx])])
|
||||
if text_prob is not None:
|
||||
conf_list.append(text_prob[batch_idx][idx])
|
||||
else:
|
||||
conf_list.append(1)
|
||||
text = ''.join(char_list)
|
||||
result_list.append((text, np.mean(conf_list)))
|
||||
return result_list
|
||||
|
||||
def get_ignored_tokens(self):
|
||||
return [0] # for ctc blank
|
||||
|
||||
|
||||
class CTCLabelDecode(BaseRecLabelDecode):
|
||||
""" Convert between text-label and text-index """
|
||||
|
||||
def __init__(self,
|
||||
character_dict_path=None,
|
||||
use_space_char=False,
|
||||
**kwargs):
|
||||
super(CTCLabelDecode, self).__init__(character_dict_path,
|
||||
use_space_char)
|
||||
|
||||
def __call__(self, preds, label=None, *args, **kwargs):
|
||||
if isinstance(preds, torch.Tensor):
|
||||
preds = preds.numpy()
|
||||
preds_idx = preds.argmax(axis=2)
|
||||
preds_prob = preds.max(axis=2)
|
||||
text = self.decode(preds_idx, preds_prob, is_remove_duplicate=True)
|
||||
|
||||
if label is None:
|
||||
return text
|
||||
label = self.decode(label)
|
||||
return text, label
|
||||
|
||||
def add_special_char(self, dict_character):
|
||||
dict_character = ['blank'] + dict_character
|
||||
return dict_character
|
||||
|
||||
|
||||
class NRTRLabelDecode(BaseRecLabelDecode):
|
||||
""" Convert between text-label and text-index """
|
||||
|
||||
def __init__(self, character_dict_path=None, use_space_char=True, **kwargs):
|
||||
super(NRTRLabelDecode, self).__init__(character_dict_path,
|
||||
use_space_char)
|
||||
|
||||
def __call__(self, preds, label=None, *args, **kwargs):
|
||||
|
||||
if len(preds) == 2:
|
||||
preds_id = preds[0]
|
||||
preds_prob = preds[1]
|
||||
if isinstance(preds_id, torch.Tensor):
|
||||
preds_id = preds_id.numpy()
|
||||
if isinstance(preds_prob, torch.Tensor):
|
||||
preds_prob = preds_prob.numpy()
|
||||
if preds_id[0][0] == 2:
|
||||
preds_idx = preds_id[:, 1:]
|
||||
preds_prob = preds_prob[:, 1:]
|
||||
else:
|
||||
preds_idx = preds_id
|
||||
text = self.decode(preds_idx, preds_prob, is_remove_duplicate=False)
|
||||
if label is None:
|
||||
return text
|
||||
label = self.decode(label[:, 1:])
|
||||
else:
|
||||
if isinstance(preds, torch.Tensor):
|
||||
preds = preds.numpy()
|
||||
preds_idx = preds.argmax(axis=2)
|
||||
preds_prob = preds.max(axis=2)
|
||||
text = self.decode(preds_idx, preds_prob, is_remove_duplicate=False)
|
||||
if label is None:
|
||||
return text
|
||||
label = self.decode(label[:, 1:])
|
||||
return text, label
|
||||
|
||||
def add_special_char(self, dict_character):
|
||||
dict_character = ['blank', '<unk>', '<s>', '</s>'] + dict_character
|
||||
return dict_character
|
||||
|
||||
def decode(self, text_index, text_prob=None, is_remove_duplicate=False):
|
||||
""" convert text-index into text-label. """
|
||||
result_list = []
|
||||
batch_size = len(text_index)
|
||||
for batch_idx in range(batch_size):
|
||||
char_list = []
|
||||
conf_list = []
|
||||
for idx in range(len(text_index[batch_idx])):
|
||||
try:
|
||||
char_idx = self.character[int(text_index[batch_idx][idx])]
|
||||
except:
|
||||
continue
|
||||
if char_idx == '</s>': # end
|
||||
break
|
||||
char_list.append(char_idx)
|
||||
if text_prob is not None:
|
||||
conf_list.append(text_prob[batch_idx][idx])
|
||||
else:
|
||||
conf_list.append(1)
|
||||
text = ''.join(char_list)
|
||||
result_list.append((text.lower(), np.mean(conf_list).tolist()))
|
||||
return result_list
|
||||
|
||||
class ViTSTRLabelDecode(NRTRLabelDecode):
|
||||
""" Convert between text-label and text-index """
|
||||
|
||||
def __init__(self, character_dict_path=None, use_space_char=False,
|
||||
**kwargs):
|
||||
super(ViTSTRLabelDecode, self).__init__(character_dict_path,
|
||||
use_space_char)
|
||||
|
||||
def __call__(self, preds, label=None, *args, **kwargs):
|
||||
if isinstance(preds, torch.Tensor):
|
||||
preds = preds[:, 1:].numpy()
|
||||
else:
|
||||
preds = preds[:, 1:]
|
||||
preds_idx = preds.argmax(axis=2)
|
||||
preds_prob = preds.max(axis=2)
|
||||
text = self.decode(preds_idx, preds_prob, is_remove_duplicate=False)
|
||||
if label is None:
|
||||
return text
|
||||
label = self.decode(label[:, 1:])
|
||||
return text, label
|
||||
|
||||
def add_special_char(self, dict_character):
|
||||
dict_character = ['<s>', '</s>'] + dict_character
|
||||
return dict_character
|
||||
|
||||
|
||||
class AttnLabelDecode(BaseRecLabelDecode):
|
||||
""" Convert between text-label and text-index """
|
||||
|
||||
def __init__(self,
|
||||
character_dict_path=None,
|
||||
use_space_char=False,
|
||||
**kwargs):
|
||||
super(AttnLabelDecode, self).__init__(character_dict_path,
|
||||
use_space_char)
|
||||
|
||||
def add_special_char(self, dict_character):
|
||||
self.beg_str = "sos"
|
||||
self.end_str = "eos"
|
||||
dict_character = dict_character
|
||||
dict_character = [self.beg_str] + dict_character + [self.end_str]
|
||||
return dict_character
|
||||
|
||||
def decode(self, text_index, text_prob=None, is_remove_duplicate=False):
|
||||
""" convert text-index into text-label. """
|
||||
result_list = []
|
||||
ignored_tokens = self.get_ignored_tokens()
|
||||
[beg_idx, end_idx] = self.get_ignored_tokens()
|
||||
batch_size = len(text_index)
|
||||
for batch_idx in range(batch_size):
|
||||
char_list = []
|
||||
conf_list = []
|
||||
for idx in range(len(text_index[batch_idx])):
|
||||
if text_index[batch_idx][idx] in ignored_tokens:
|
||||
continue
|
||||
if int(text_index[batch_idx][idx]) == int(end_idx):
|
||||
break
|
||||
if is_remove_duplicate:
|
||||
# only for predict
|
||||
if idx > 0 and text_index[batch_idx][idx - 1] == text_index[
|
||||
batch_idx][idx]:
|
||||
continue
|
||||
char_list.append(self.character[int(text_index[batch_idx][
|
||||
idx])])
|
||||
if text_prob is not None:
|
||||
conf_list.append(text_prob[batch_idx][idx])
|
||||
else:
|
||||
conf_list.append(1)
|
||||
text = ''.join(char_list)
|
||||
result_list.append((text, np.mean(conf_list)))
|
||||
return result_list
|
||||
|
||||
def __call__(self, preds, label=None, *args, **kwargs):
|
||||
"""
|
||||
text = self.decode(text)
|
||||
if label is None:
|
||||
return text
|
||||
else:
|
||||
label = self.decode(label, is_remove_duplicate=False)
|
||||
return text, label
|
||||
"""
|
||||
if isinstance(preds, torch.Tensor):
|
||||
preds = preds.cpu().numpy()
|
||||
|
||||
preds_idx = preds.argmax(axis=2)
|
||||
preds_prob = preds.max(axis=2)
|
||||
text = self.decode(preds_idx, preds_prob, is_remove_duplicate=False)
|
||||
if label is None:
|
||||
return text
|
||||
label = self.decode(label, is_remove_duplicate=False)
|
||||
return text, label
|
||||
|
||||
def get_ignored_tokens(self):
|
||||
beg_idx = self.get_beg_end_flag_idx("beg")
|
||||
end_idx = self.get_beg_end_flag_idx("end")
|
||||
return [beg_idx, end_idx]
|
||||
|
||||
def get_beg_end_flag_idx(self, beg_or_end):
|
||||
if beg_or_end == "beg":
|
||||
idx = np.array(self.dict[self.beg_str])
|
||||
elif beg_or_end == "end":
|
||||
idx = np.array(self.dict[self.end_str])
|
||||
else:
|
||||
assert False, "unsupport type %s in get_beg_end_flag_idx" \
|
||||
% beg_or_end
|
||||
return idx
|
||||
|
||||
|
||||
class RFLLabelDecode(BaseRecLabelDecode):
|
||||
""" Convert between text-label and text-index """
|
||||
|
||||
def __init__(self, character_dict_path=None, use_space_char=False,
|
||||
**kwargs):
|
||||
super(RFLLabelDecode, self).__init__(character_dict_path,
|
||||
use_space_char)
|
||||
|
||||
def add_special_char(self, dict_character):
|
||||
self.beg_str = "sos"
|
||||
self.end_str = "eos"
|
||||
dict_character = dict_character
|
||||
dict_character = [self.beg_str] + dict_character + [self.end_str]
|
||||
return dict_character
|
||||
|
||||
def decode(self, text_index, text_prob=None, is_remove_duplicate=False):
|
||||
""" convert text-index into text-label. """
|
||||
result_list = []
|
||||
ignored_tokens = self.get_ignored_tokens()
|
||||
[beg_idx, end_idx] = self.get_ignored_tokens()
|
||||
batch_size = len(text_index)
|
||||
for batch_idx in range(batch_size):
|
||||
char_list = []
|
||||
conf_list = []
|
||||
for idx in range(len(text_index[batch_idx])):
|
||||
if text_index[batch_idx][idx] in ignored_tokens:
|
||||
continue
|
||||
if int(text_index[batch_idx][idx]) == int(end_idx):
|
||||
break
|
||||
if is_remove_duplicate:
|
||||
# only for predict
|
||||
if idx > 0 and text_index[batch_idx][idx - 1] == text_index[
|
||||
batch_idx][idx]:
|
||||
continue
|
||||
char_list.append(self.character[int(text_index[batch_idx][
|
||||
idx])])
|
||||
if text_prob is not None:
|
||||
conf_list.append(text_prob[batch_idx][idx])
|
||||
else:
|
||||
conf_list.append(1)
|
||||
text = ''.join(char_list)
|
||||
result_list.append((text, np.mean(conf_list).tolist()))
|
||||
return result_list
|
||||
|
||||
def __call__(self, preds, label=None, *args, **kwargs):
|
||||
# if seq_outputs is not None:
|
||||
if isinstance(preds, tuple) or isinstance(preds, list):
|
||||
cnt_outputs, seq_outputs = preds
|
||||
if isinstance(seq_outputs, torch.Tensor):
|
||||
seq_outputs = seq_outputs.numpy()
|
||||
preds_idx = seq_outputs.argmax(axis=2)
|
||||
preds_prob = seq_outputs.max(axis=2)
|
||||
text = self.decode(preds_idx, preds_prob, is_remove_duplicate=False)
|
||||
|
||||
if label is None:
|
||||
return text
|
||||
label = self.decode(label, is_remove_duplicate=False)
|
||||
return text, label
|
||||
|
||||
else:
|
||||
cnt_outputs = preds
|
||||
if isinstance(cnt_outputs, torch.Tensor):
|
||||
cnt_outputs = cnt_outputs.numpy()
|
||||
cnt_length = []
|
||||
for lens in cnt_outputs:
|
||||
length = round(np.sum(lens))
|
||||
cnt_length.append(length)
|
||||
if label is None:
|
||||
return cnt_length
|
||||
label = self.decode(label, is_remove_duplicate=False)
|
||||
length = [len(res[0]) for res in label]
|
||||
return cnt_length, length
|
||||
|
||||
def get_ignored_tokens(self):
|
||||
beg_idx = self.get_beg_end_flag_idx("beg")
|
||||
end_idx = self.get_beg_end_flag_idx("end")
|
||||
return [beg_idx, end_idx]
|
||||
|
||||
def get_beg_end_flag_idx(self, beg_or_end):
|
||||
if beg_or_end == "beg":
|
||||
idx = np.array(self.dict[self.beg_str])
|
||||
elif beg_or_end == "end":
|
||||
idx = np.array(self.dict[self.end_str])
|
||||
else:
|
||||
assert False, "unsupport type %s in get_beg_end_flag_idx" \
|
||||
% beg_or_end
|
||||
return idx
|
||||
|
||||
|
||||
class SRNLabelDecode(BaseRecLabelDecode):
|
||||
""" Convert between text-label and text-index """
|
||||
|
||||
def __init__(self,
|
||||
character_dict_path=None,
|
||||
use_space_char=False,
|
||||
**kwargs):
|
||||
self.max_text_length = kwargs.get('max_text_length', 25)
|
||||
super(SRNLabelDecode, self).__init__(character_dict_path,
|
||||
use_space_char)
|
||||
|
||||
def __call__(self, preds, label=None, *args, **kwargs):
|
||||
pred = preds['predict']
|
||||
char_num = len(self.character_str) + 2
|
||||
if isinstance(pred, torch.Tensor):
|
||||
pred = pred.numpy()
|
||||
pred = np.reshape(pred, [-1, char_num])
|
||||
|
||||
preds_idx = np.argmax(pred, axis=1)
|
||||
preds_prob = np.max(pred, axis=1)
|
||||
|
||||
preds_idx = np.reshape(preds_idx, [-1, self.max_text_length])
|
||||
|
||||
preds_prob = np.reshape(preds_prob, [-1, self.max_text_length])
|
||||
|
||||
text = self.decode(preds_idx, preds_prob)
|
||||
|
||||
if label is None:
|
||||
text = self.decode(preds_idx, preds_prob, is_remove_duplicate=False)
|
||||
return text
|
||||
label = self.decode(label)
|
||||
return text, label
|
||||
|
||||
def decode(self, text_index, text_prob=None, is_remove_duplicate=False):
|
||||
""" convert text-index into text-label. """
|
||||
result_list = []
|
||||
ignored_tokens = self.get_ignored_tokens()
|
||||
batch_size = len(text_index)
|
||||
|
||||
for batch_idx in range(batch_size):
|
||||
char_list = []
|
||||
conf_list = []
|
||||
for idx in range(len(text_index[batch_idx])):
|
||||
if text_index[batch_idx][idx] in ignored_tokens:
|
||||
continue
|
||||
if is_remove_duplicate:
|
||||
# only for predict
|
||||
if idx > 0 and text_index[batch_idx][idx - 1] == text_index[
|
||||
batch_idx][idx]:
|
||||
continue
|
||||
char_list.append(self.character[int(text_index[batch_idx][
|
||||
idx])])
|
||||
if text_prob is not None:
|
||||
conf_list.append(text_prob[batch_idx][idx])
|
||||
else:
|
||||
conf_list.append(1)
|
||||
|
||||
text = ''.join(char_list)
|
||||
result_list.append((text, np.mean(conf_list)))
|
||||
return result_list
|
||||
|
||||
def add_special_char(self, dict_character):
|
||||
dict_character = dict_character + [self.beg_str, self.end_str]
|
||||
return dict_character
|
||||
|
||||
def get_ignored_tokens(self):
|
||||
beg_idx = self.get_beg_end_flag_idx("beg")
|
||||
end_idx = self.get_beg_end_flag_idx("end")
|
||||
return [beg_idx, end_idx]
|
||||
|
||||
def get_beg_end_flag_idx(self, beg_or_end):
|
||||
if beg_or_end == "beg":
|
||||
idx = np.array(self.dict[self.beg_str])
|
||||
elif beg_or_end == "end":
|
||||
idx = np.array(self.dict[self.end_str])
|
||||
else:
|
||||
assert False, "unsupport type %s in get_beg_end_flag_idx" \
|
||||
% beg_or_end
|
||||
return idx
|
||||
|
||||
|
||||
class TableLabelDecode(object):
|
||||
""" """
|
||||
|
||||
def __init__(self,
|
||||
character_dict_path,
|
||||
**kwargs):
|
||||
list_character, list_elem = self.load_char_elem_dict(character_dict_path)
|
||||
list_character = self.add_special_char(list_character)
|
||||
list_elem = self.add_special_char(list_elem)
|
||||
self.dict_character = {}
|
||||
self.dict_idx_character = {}
|
||||
for i, char in enumerate(list_character):
|
||||
self.dict_idx_character[i] = char
|
||||
self.dict_character[char] = i
|
||||
self.dict_elem = {}
|
||||
self.dict_idx_elem = {}
|
||||
for i, elem in enumerate(list_elem):
|
||||
self.dict_idx_elem[i] = elem
|
||||
self.dict_elem[elem] = i
|
||||
|
||||
def load_char_elem_dict(self, character_dict_path):
|
||||
list_character = []
|
||||
list_elem = []
|
||||
with open(character_dict_path, "rb") as fin:
|
||||
lines = fin.readlines()
|
||||
substr = lines[0].decode('utf-8').strip("\n").strip("\r\n").split("\t")
|
||||
character_num = int(substr[0])
|
||||
elem_num = int(substr[1])
|
||||
for cno in range(1, 1 + character_num):
|
||||
character = lines[cno].decode('utf-8').strip("\n").strip("\r\n")
|
||||
list_character.append(character)
|
||||
for eno in range(1 + character_num, 1 + character_num + elem_num):
|
||||
elem = lines[eno].decode('utf-8').strip("\n").strip("\r\n")
|
||||
list_elem.append(elem)
|
||||
return list_character, list_elem
|
||||
|
||||
def add_special_char(self, list_character):
|
||||
self.beg_str = "sos"
|
||||
self.end_str = "eos"
|
||||
list_character = [self.beg_str] + list_character + [self.end_str]
|
||||
return list_character
|
||||
|
||||
def __call__(self, preds):
|
||||
structure_probs = preds['structure_probs']
|
||||
loc_preds = preds['loc_preds']
|
||||
if isinstance(structure_probs,torch.Tensor):
|
||||
structure_probs = structure_probs.numpy()
|
||||
if isinstance(loc_preds,torch.Tensor):
|
||||
loc_preds = loc_preds.numpy()
|
||||
structure_idx = structure_probs.argmax(axis=2)
|
||||
structure_probs = structure_probs.max(axis=2)
|
||||
structure_str, structure_pos, result_score_list, result_elem_idx_list = self.decode(structure_idx,
|
||||
structure_probs, 'elem')
|
||||
res_html_code_list = []
|
||||
res_loc_list = []
|
||||
batch_num = len(structure_str)
|
||||
for bno in range(batch_num):
|
||||
res_loc = []
|
||||
for sno in range(len(structure_str[bno])):
|
||||
text = structure_str[bno][sno]
|
||||
if text in ['<td>', '<td']:
|
||||
pos = structure_pos[bno][sno]
|
||||
res_loc.append(loc_preds[bno, pos])
|
||||
res_html_code = ''.join(structure_str[bno])
|
||||
res_loc = np.array(res_loc)
|
||||
res_html_code_list.append(res_html_code)
|
||||
res_loc_list.append(res_loc)
|
||||
return {'res_html_code': res_html_code_list, 'res_loc': res_loc_list, 'res_score_list': result_score_list,
|
||||
'res_elem_idx_list': result_elem_idx_list,'structure_str_list':structure_str}
|
||||
|
||||
def decode(self, text_index, structure_probs, char_or_elem):
|
||||
"""convert text-label into text-index.
|
||||
"""
|
||||
if char_or_elem == "char":
|
||||
current_dict = self.dict_idx_character
|
||||
else:
|
||||
current_dict = self.dict_idx_elem
|
||||
ignored_tokens = self.get_ignored_tokens('elem')
|
||||
beg_idx, end_idx = ignored_tokens
|
||||
|
||||
result_list = []
|
||||
result_pos_list = []
|
||||
result_score_list = []
|
||||
result_elem_idx_list = []
|
||||
batch_size = len(text_index)
|
||||
for batch_idx in range(batch_size):
|
||||
char_list = []
|
||||
elem_pos_list = []
|
||||
elem_idx_list = []
|
||||
score_list = []
|
||||
for idx in range(len(text_index[batch_idx])):
|
||||
tmp_elem_idx = int(text_index[batch_idx][idx])
|
||||
if idx > 0 and tmp_elem_idx == end_idx:
|
||||
break
|
||||
if tmp_elem_idx in ignored_tokens:
|
||||
continue
|
||||
|
||||
char_list.append(current_dict[tmp_elem_idx])
|
||||
elem_pos_list.append(idx)
|
||||
score_list.append(structure_probs[batch_idx, idx])
|
||||
elem_idx_list.append(tmp_elem_idx)
|
||||
result_list.append(char_list)
|
||||
result_pos_list.append(elem_pos_list)
|
||||
result_score_list.append(score_list)
|
||||
result_elem_idx_list.append(elem_idx_list)
|
||||
return result_list, result_pos_list, result_score_list, result_elem_idx_list
|
||||
|
||||
def get_ignored_tokens(self, char_or_elem):
|
||||
beg_idx = self.get_beg_end_flag_idx("beg", char_or_elem)
|
||||
end_idx = self.get_beg_end_flag_idx("end", char_or_elem)
|
||||
return [beg_idx, end_idx]
|
||||
|
||||
def get_beg_end_flag_idx(self, beg_or_end, char_or_elem):
|
||||
if char_or_elem == "char":
|
||||
if beg_or_end == "beg":
|
||||
idx = self.dict_character[self.beg_str]
|
||||
elif beg_or_end == "end":
|
||||
idx = self.dict_character[self.end_str]
|
||||
else:
|
||||
assert False, "Unsupport type %s in get_beg_end_flag_idx of char" \
|
||||
% beg_or_end
|
||||
elif char_or_elem == "elem":
|
||||
if beg_or_end == "beg":
|
||||
idx = self.dict_elem[self.beg_str]
|
||||
elif beg_or_end == "end":
|
||||
idx = self.dict_elem[self.end_str]
|
||||
else:
|
||||
assert False, "Unsupport type %s in get_beg_end_flag_idx of elem" \
|
||||
% beg_or_end
|
||||
else:
|
||||
assert False, "Unsupport type %s in char_or_elem" \
|
||||
% char_or_elem
|
||||
return idx
|
||||
|
||||
|
||||
class SARLabelDecode(BaseRecLabelDecode):
|
||||
""" Convert between text-label and text-index """
|
||||
|
||||
def __init__(self, character_dict_path=None, use_space_char=False,
|
||||
**kwargs):
|
||||
super(SARLabelDecode, self).__init__(character_dict_path,
|
||||
use_space_char)
|
||||
|
||||
self.rm_symbol = kwargs.get('rm_symbol', False)
|
||||
|
||||
def add_special_char(self, dict_character):
|
||||
beg_end_str = "<BOS/EOS>"
|
||||
unknown_str = "<UKN>"
|
||||
padding_str = "<PAD>"
|
||||
dict_character = dict_character + [unknown_str]
|
||||
self.unknown_idx = len(dict_character) - 1
|
||||
dict_character = dict_character + [beg_end_str]
|
||||
self.start_idx = len(dict_character) - 1
|
||||
self.end_idx = len(dict_character) - 1
|
||||
dict_character = dict_character + [padding_str]
|
||||
self.padding_idx = len(dict_character) - 1
|
||||
return dict_character
|
||||
|
||||
def decode(self, text_index, text_prob=None, is_remove_duplicate=False):
|
||||
""" convert text-index into text-label. """
|
||||
result_list = []
|
||||
ignored_tokens = self.get_ignored_tokens()
|
||||
|
||||
batch_size = len(text_index)
|
||||
for batch_idx in range(batch_size):
|
||||
char_list = []
|
||||
conf_list = []
|
||||
for idx in range(len(text_index[batch_idx])):
|
||||
if text_index[batch_idx][idx] in ignored_tokens:
|
||||
continue
|
||||
if int(text_index[batch_idx][idx]) == int(self.end_idx):
|
||||
if text_prob is None and idx == 0:
|
||||
continue
|
||||
else:
|
||||
break
|
||||
if is_remove_duplicate:
|
||||
# only for predict
|
||||
if idx > 0 and text_index[batch_idx][idx - 1] == text_index[
|
||||
batch_idx][idx]:
|
||||
continue
|
||||
char_list.append(self.character[int(text_index[batch_idx][
|
||||
idx])])
|
||||
if text_prob is not None:
|
||||
conf_list.append(text_prob[batch_idx][idx])
|
||||
else:
|
||||
conf_list.append(1)
|
||||
text = ''.join(char_list)
|
||||
if self.rm_symbol:
|
||||
comp = re.compile('[^A-Z^a-z^0-9^\u4e00-\u9fa5]')
|
||||
text = text.lower()
|
||||
text = comp.sub('', text)
|
||||
result_list.append((text, np.mean(conf_list).tolist()))
|
||||
return result_list
|
||||
|
||||
def __call__(self, preds, label=None, *args, **kwargs):
|
||||
if isinstance(preds, torch.Tensor):
|
||||
preds = preds.cpu().numpy()
|
||||
preds_idx = preds.argmax(axis=2)
|
||||
preds_prob = preds.max(axis=2)
|
||||
|
||||
text = self.decode(preds_idx, preds_prob, is_remove_duplicate=False)
|
||||
|
||||
if label is None:
|
||||
return text
|
||||
label = self.decode(label, is_remove_duplicate=False)
|
||||
return text, label
|
||||
|
||||
def get_ignored_tokens(self):
|
||||
return [self.padding_idx]
|
||||
|
||||
|
||||
class CANLabelDecode(BaseRecLabelDecode):
|
||||
""" Convert between latex-symbol and symbol-index """
|
||||
|
||||
def __init__(self, character_dict_path=None, use_space_char=False,
|
||||
**kwargs):
|
||||
super(CANLabelDecode, self).__init__(character_dict_path,
|
||||
use_space_char)
|
||||
|
||||
def decode(self, text_index, preds_prob=None):
|
||||
result_list = []
|
||||
batch_size = len(text_index)
|
||||
for batch_idx in range(batch_size):
|
||||
seq_end = text_index[batch_idx].argmin(0)
|
||||
idx_list = text_index[batch_idx][:seq_end].tolist()
|
||||
symbol_list = [self.character[idx] for idx in idx_list]
|
||||
probs = []
|
||||
if preds_prob is not None:
|
||||
probs = preds_prob[batch_idx][:len(symbol_list)].tolist()
|
||||
|
||||
result_list.append([' '.join(symbol_list), probs])
|
||||
return result_list
|
||||
|
||||
def __call__(self, preds, label=None, *args, **kwargs):
|
||||
pred_prob, _, _, _ = preds
|
||||
preds_idx = pred_prob.argmax(axis=2)
|
||||
|
||||
text = self.decode(preds_idx)
|
||||
if label is None:
|
||||
return text
|
||||
label = self.decode(label)
|
||||
return text, label
|
||||
@@ -0,0 +1,476 @@
|
||||
ch_ptocr_mobile_v2.0_cls_infer:
|
||||
model_type: cls
|
||||
algorithm: CLS
|
||||
Transform:
|
||||
Backbone:
|
||||
name: MobileNetV3
|
||||
scale: 0.35
|
||||
model_name: small
|
||||
Neck:
|
||||
Head:
|
||||
name: ClsHead
|
||||
class_dim: 2
|
||||
|
||||
Multilingual_PP-OCRv3_det_infer:
|
||||
model_type: det
|
||||
algorithm: DB
|
||||
Transform:
|
||||
Backbone:
|
||||
name: MobileNetV3
|
||||
scale: 0.5
|
||||
model_name: large
|
||||
disable_se: True
|
||||
Neck:
|
||||
name: RSEFPN
|
||||
out_channels: 96
|
||||
shortcut: True
|
||||
Head:
|
||||
name: DBHead
|
||||
k: 50
|
||||
|
||||
en_PP-OCRv3_det_infer:
|
||||
model_type: det
|
||||
algorithm: DB
|
||||
Transform:
|
||||
Backbone:
|
||||
name: MobileNetV3
|
||||
scale: 0.5
|
||||
model_name: large
|
||||
disable_se: True
|
||||
Neck:
|
||||
name: RSEFPN
|
||||
out_channels: 96
|
||||
shortcut: True
|
||||
Head:
|
||||
name: DBHead
|
||||
k: 50
|
||||
|
||||
ch_PP-OCRv3_det_infer:
|
||||
model_type: det
|
||||
algorithm: DB
|
||||
Transform:
|
||||
Backbone:
|
||||
name: MobileNetV3
|
||||
scale: 0.5
|
||||
model_name: large
|
||||
disable_se: True
|
||||
Neck:
|
||||
name: RSEFPN
|
||||
out_channels: 96
|
||||
shortcut: True
|
||||
Head:
|
||||
name: DBHead
|
||||
k: 50
|
||||
|
||||
en_PP-OCRv4_rec_infer:
|
||||
model_type: rec
|
||||
algorithm: SVTR_LCNet
|
||||
Transform:
|
||||
Backbone:
|
||||
name: PPLCNetV3
|
||||
scale: 0.95
|
||||
Head:
|
||||
name: MultiHead
|
||||
out_channels_list:
|
||||
CTCLabelDecode: 97 #'blank' + ...(62) + ' '
|
||||
head_list:
|
||||
- CTCHead:
|
||||
Neck:
|
||||
name: svtr
|
||||
dims: 120
|
||||
depth: 2
|
||||
hidden_dims: 120
|
||||
kernel_size: [ 1, 3 ]
|
||||
use_guide: True
|
||||
Head:
|
||||
fc_decay: 0.00001
|
||||
- NRTRHead:
|
||||
nrtr_dim: 384
|
||||
max_text_length: 25
|
||||
|
||||
ch_PP-OCRv4_det_infer:
|
||||
model_type: det
|
||||
algorithm: DB
|
||||
Transform: null
|
||||
Backbone:
|
||||
name: PPLCNetV3
|
||||
scale: 0.75
|
||||
det: True
|
||||
Neck:
|
||||
name: RSEFPN
|
||||
out_channels: 96
|
||||
shortcut: True
|
||||
Head:
|
||||
name: DBHead
|
||||
k: 50
|
||||
|
||||
ch_PP-OCRv5_det_infer:
|
||||
model_type: det
|
||||
algorithm: DB
|
||||
Transform: null
|
||||
Backbone:
|
||||
name: PPLCNetV3
|
||||
scale: 0.75
|
||||
det: True
|
||||
Neck:
|
||||
name: RSEFPN
|
||||
out_channels: 96
|
||||
shortcut: True
|
||||
Head:
|
||||
name: DBHead
|
||||
k: 50
|
||||
|
||||
ch_PP-OCRv4_det_server_infer:
|
||||
model_type: det
|
||||
algorithm: DB
|
||||
Transform: null
|
||||
Backbone:
|
||||
name: PPHGNet_small
|
||||
det: True
|
||||
Neck:
|
||||
name: LKPAN
|
||||
out_channels: 256
|
||||
intracl: true
|
||||
Head:
|
||||
name: PFHeadLocal
|
||||
k: 50
|
||||
mode: "large"
|
||||
|
||||
ch_PP-OCRv4_rec_infer:
|
||||
model_type: rec
|
||||
algorithm: SVTR_LCNet
|
||||
Transform:
|
||||
Backbone:
|
||||
name: PPLCNetV3
|
||||
scale: 0.95
|
||||
Head:
|
||||
name: MultiHead
|
||||
out_channels_list:
|
||||
CTCLabelDecode: 6625 #'blank' + ...(6623) + ' '
|
||||
head_list:
|
||||
- CTCHead:
|
||||
Neck:
|
||||
name: svtr
|
||||
dims: 120
|
||||
depth: 2
|
||||
hidden_dims: 120
|
||||
kernel_size: [ 1, 3 ]
|
||||
use_guide: True
|
||||
Head:
|
||||
fc_decay: 0.00001
|
||||
- NRTRHead:
|
||||
nrtr_dim: 384
|
||||
max_text_length: 25
|
||||
|
||||
ch_PP-OCRv4_rec_server_infer:
|
||||
model_type: rec
|
||||
algorithm: SVTR_HGNet
|
||||
Transform:
|
||||
Backbone:
|
||||
name: PPHGNet_small
|
||||
Head:
|
||||
name: MultiHead
|
||||
out_channels_list:
|
||||
CTCLabelDecode: 6625 #'blank' + ...(6623) + ' '
|
||||
head_list:
|
||||
- CTCHead:
|
||||
Neck:
|
||||
name: svtr
|
||||
dims: 120
|
||||
depth: 2
|
||||
hidden_dims: 120
|
||||
kernel_size: [ 1, 3 ]
|
||||
use_guide: True
|
||||
Head:
|
||||
fc_decay: 0.00001
|
||||
- NRTRHead:
|
||||
nrtr_dim: 384
|
||||
max_text_length: 25
|
||||
|
||||
ch_PP-OCRv4_rec_server_doc_infer:
|
||||
model_type: rec
|
||||
algorithm: SVTR_HGNet
|
||||
Transform:
|
||||
Backbone:
|
||||
name: PPHGNet_small
|
||||
Head:
|
||||
name: MultiHead
|
||||
out_channels_list:
|
||||
CTCLabelDecode: 15631
|
||||
head_list:
|
||||
- CTCHead:
|
||||
Neck:
|
||||
name: svtr
|
||||
dims: 120
|
||||
depth: 2
|
||||
hidden_dims: 120
|
||||
kernel_size: [ 1, 3 ]
|
||||
use_guide: True
|
||||
Head:
|
||||
fc_decay: 0.00001
|
||||
- NRTRHead:
|
||||
nrtr_dim: 384
|
||||
max_text_length: 25
|
||||
|
||||
ch_PP-OCRv5_rec_server_infer:
|
||||
model_type: rec
|
||||
algorithm: SVTR_HGNet
|
||||
Transform:
|
||||
Backbone:
|
||||
name: PPHGNetV2_B4
|
||||
text_rec: True
|
||||
Head:
|
||||
name: MultiHead
|
||||
out_channels_list:
|
||||
CTCLabelDecode: 18385
|
||||
head_list:
|
||||
- CTCHead:
|
||||
Neck:
|
||||
name: svtr
|
||||
dims: 120
|
||||
depth: 2
|
||||
hidden_dims: 120
|
||||
kernel_size: [ 1, 3 ]
|
||||
use_guide: True
|
||||
Head:
|
||||
fc_decay: 0.00001
|
||||
- NRTRHead:
|
||||
nrtr_dim: 384
|
||||
max_text_length: 25
|
||||
|
||||
ch_PP-OCRv5_rec_infer:
|
||||
model_type: rec
|
||||
algorithm: SVTR_HGNet
|
||||
Transform:
|
||||
Backbone:
|
||||
name: PPLCNetV3
|
||||
scale: 0.95
|
||||
Head:
|
||||
name: MultiHead
|
||||
out_channels_list:
|
||||
CTCLabelDecode: 18385
|
||||
head_list:
|
||||
- CTCHead:
|
||||
Neck:
|
||||
name: svtr
|
||||
dims: 120
|
||||
depth: 2
|
||||
hidden_dims: 120
|
||||
kernel_size: [ 1, 3 ]
|
||||
use_guide: True
|
||||
Head:
|
||||
fc_decay: 0.00001
|
||||
- NRTRHead:
|
||||
nrtr_dim: 384
|
||||
max_text_length: 25
|
||||
|
||||
chinese_cht_PP-OCRv3_rec_infer:
|
||||
model_type: rec
|
||||
algorithm: SVTR
|
||||
Transform:
|
||||
Backbone:
|
||||
name: MobileNetV1Enhance
|
||||
scale: 0.5
|
||||
last_conv_stride: [1, 2]
|
||||
last_pool_type: avg
|
||||
Neck:
|
||||
name: SequenceEncoder
|
||||
encoder_type: svtr
|
||||
dims: 64
|
||||
depth: 2
|
||||
hidden_dims: 120
|
||||
use_guide: True
|
||||
Head:
|
||||
name: CTCHead
|
||||
# out_channels: 8423
|
||||
fc_decay: 0.00001
|
||||
|
||||
latin_PP-OCRv3_rec_infer:
|
||||
model_type: rec
|
||||
algorithm: SVTR
|
||||
Transform:
|
||||
Backbone:
|
||||
name: MobileNetV1Enhance
|
||||
scale: 0.5
|
||||
last_conv_stride: [ 1, 2 ]
|
||||
last_pool_type: avg
|
||||
Neck:
|
||||
name: SequenceEncoder
|
||||
encoder_type: svtr
|
||||
dims: 64
|
||||
depth: 2
|
||||
hidden_dims: 120
|
||||
use_guide: True
|
||||
Head:
|
||||
name: CTCHead
|
||||
# out_channels: 187
|
||||
fc_decay: 0.00001
|
||||
|
||||
cyrillic_PP-OCRv3_rec_infer:
|
||||
model_type: rec
|
||||
algorithm: SVTR
|
||||
Transform:
|
||||
Backbone:
|
||||
name: MobileNetV1Enhance
|
||||
scale: 0.5
|
||||
last_conv_stride: [ 1, 2 ]
|
||||
last_pool_type: avg
|
||||
Neck:
|
||||
name: SequenceEncoder
|
||||
encoder_type: svtr
|
||||
dims: 64
|
||||
depth: 2
|
||||
hidden_dims: 120
|
||||
use_guide: True
|
||||
Head:
|
||||
name: CTCHead
|
||||
# out_channels: 165
|
||||
fc_decay: 0.00001
|
||||
|
||||
arabic_PP-OCRv3_rec_infer:
|
||||
model_type: rec
|
||||
algorithm: SVTR
|
||||
Transform:
|
||||
Backbone:
|
||||
name: MobileNetV1Enhance
|
||||
scale: 0.5
|
||||
last_conv_stride: [ 1, 2 ]
|
||||
last_pool_type: avg
|
||||
Neck:
|
||||
name: SequenceEncoder
|
||||
encoder_type: svtr
|
||||
dims: 64
|
||||
depth: 2
|
||||
hidden_dims: 120
|
||||
use_guide: True
|
||||
Head:
|
||||
name: CTCHead
|
||||
# out_channels: 164
|
||||
fc_decay: 0.00001
|
||||
|
||||
korean_PP-OCRv3_rec_infer:
|
||||
model_type: rec
|
||||
algorithm: SVTR
|
||||
Transform:
|
||||
Backbone:
|
||||
name: MobileNetV1Enhance
|
||||
scale: 0.5
|
||||
last_conv_stride: [ 1, 2 ]
|
||||
last_pool_type: avg
|
||||
Neck:
|
||||
name: SequenceEncoder
|
||||
encoder_type: svtr
|
||||
dims: 64
|
||||
depth: 2
|
||||
hidden_dims: 120
|
||||
use_guide: True
|
||||
Head:
|
||||
name: CTCHead
|
||||
# out_channels: 3690
|
||||
fc_decay: 0.00001
|
||||
|
||||
japan_PP-OCRv3_rec_infer:
|
||||
model_type: rec
|
||||
algorithm: SVTR
|
||||
Transform:
|
||||
Backbone:
|
||||
name: MobileNetV1Enhance
|
||||
scale: 0.5
|
||||
last_conv_stride: [ 1, 2 ]
|
||||
last_pool_type: avg
|
||||
Neck:
|
||||
name: SequenceEncoder
|
||||
encoder_type: svtr
|
||||
dims: 64
|
||||
depth: 2
|
||||
hidden_dims: 120
|
||||
use_guide: True
|
||||
Head:
|
||||
name: CTCHead
|
||||
# out_channels: 4401
|
||||
fc_decay: 0.00001
|
||||
|
||||
ta_PP-OCRv3_rec_infer:
|
||||
model_type: rec
|
||||
algorithm: SVTR
|
||||
Transform:
|
||||
Backbone:
|
||||
name: MobileNetV1Enhance
|
||||
scale: 0.5
|
||||
last_conv_stride: [ 1, 2 ]
|
||||
last_pool_type: avg
|
||||
Neck:
|
||||
name: SequenceEncoder
|
||||
encoder_type: svtr
|
||||
dims: 64
|
||||
depth: 2
|
||||
hidden_dims: 120
|
||||
use_guide: True
|
||||
Head:
|
||||
name: CTCHead
|
||||
# out_channels: 130
|
||||
fc_decay: 0.00001
|
||||
|
||||
te_PP-OCRv3_rec_infer:
|
||||
model_type: rec
|
||||
algorithm: SVTR
|
||||
Transform:
|
||||
Backbone:
|
||||
name: MobileNetV1Enhance
|
||||
scale: 0.5
|
||||
last_conv_stride: [ 1, 2 ]
|
||||
last_pool_type: avg
|
||||
Neck:
|
||||
name: SequenceEncoder
|
||||
encoder_type: svtr
|
||||
dims: 64
|
||||
depth: 2
|
||||
hidden_dims: 120
|
||||
use_guide: True
|
||||
Head:
|
||||
name: CTCHead
|
||||
# out_channels: 153
|
||||
fc_decay: 0.00001
|
||||
|
||||
ka_PP-OCRv3_rec_infer:
|
||||
model_type: rec
|
||||
algorithm: SVTR
|
||||
Transform:
|
||||
Backbone:
|
||||
name: MobileNetV1Enhance
|
||||
scale: 0.5
|
||||
last_conv_stride: [ 1, 2 ]
|
||||
last_pool_type: avg
|
||||
Neck:
|
||||
name: SequenceEncoder
|
||||
encoder_type: svtr
|
||||
dims: 64
|
||||
depth: 2
|
||||
hidden_dims: 120
|
||||
use_guide: True
|
||||
Head:
|
||||
name: CTCHead
|
||||
# out_channels: 155
|
||||
fc_decay: 0.00001
|
||||
|
||||
devanagari_PP-OCRv3_rec_infer:
|
||||
model_type: rec
|
||||
algorithm: SVTR
|
||||
Transform:
|
||||
Backbone:
|
||||
name: MobileNetV1Enhance
|
||||
scale: 0.5
|
||||
last_conv_stride: [ 1, 2 ]
|
||||
last_pool_type: avg
|
||||
Neck:
|
||||
name: SequenceEncoder
|
||||
encoder_type: svtr
|
||||
dims: 64
|
||||
depth: 2
|
||||
hidden_dims: 120
|
||||
use_guide: True
|
||||
Head:
|
||||
name: CTCHead
|
||||
# out_channels: 169
|
||||
fc_decay: 0.00001
|
||||
|
||||
@@ -0,0 +1,162 @@
|
||||
|
||||
!
|
||||
#
|
||||
$
|
||||
%
|
||||
&
|
||||
'
|
||||
(
|
||||
+
|
||||
,
|
||||
-
|
||||
.
|
||||
/
|
||||
0
|
||||
1
|
||||
2
|
||||
3
|
||||
4
|
||||
5
|
||||
6
|
||||
7
|
||||
8
|
||||
9
|
||||
:
|
||||
?
|
||||
@
|
||||
A
|
||||
B
|
||||
C
|
||||
D
|
||||
E
|
||||
F
|
||||
G
|
||||
H
|
||||
I
|
||||
J
|
||||
K
|
||||
L
|
||||
M
|
||||
N
|
||||
O
|
||||
P
|
||||
Q
|
||||
R
|
||||
S
|
||||
T
|
||||
U
|
||||
V
|
||||
W
|
||||
X
|
||||
Y
|
||||
Z
|
||||
_
|
||||
a
|
||||
b
|
||||
c
|
||||
d
|
||||
e
|
||||
f
|
||||
g
|
||||
h
|
||||
i
|
||||
j
|
||||
k
|
||||
l
|
||||
m
|
||||
n
|
||||
o
|
||||
p
|
||||
q
|
||||
r
|
||||
s
|
||||
t
|
||||
u
|
||||
v
|
||||
w
|
||||
x
|
||||
y
|
||||
z
|
||||
É
|
||||
é
|
||||
ء
|
||||
آ
|
||||
أ
|
||||
ؤ
|
||||
إ
|
||||
ئ
|
||||
ا
|
||||
ب
|
||||
ة
|
||||
ت
|
||||
ث
|
||||
ج
|
||||
ح
|
||||
خ
|
||||
د
|
||||
ذ
|
||||
ر
|
||||
ز
|
||||
س
|
||||
ش
|
||||
ص
|
||||
ض
|
||||
ط
|
||||
ظ
|
||||
ع
|
||||
غ
|
||||
ف
|
||||
ق
|
||||
ك
|
||||
ل
|
||||
م
|
||||
ن
|
||||
ه
|
||||
و
|
||||
ى
|
||||
ي
|
||||
ً
|
||||
ٌ
|
||||
ٍ
|
||||
َ
|
||||
ُ
|
||||
ِ
|
||||
ّ
|
||||
ْ
|
||||
ٓ
|
||||
ٔ
|
||||
ٰ
|
||||
ٱ
|
||||
ٹ
|
||||
پ
|
||||
چ
|
||||
ڈ
|
||||
ڑ
|
||||
ژ
|
||||
ک
|
||||
ڭ
|
||||
گ
|
||||
ں
|
||||
ھ
|
||||
ۀ
|
||||
ہ
|
||||
ۂ
|
||||
ۃ
|
||||
ۆ
|
||||
ۇ
|
||||
ۈ
|
||||
ۋ
|
||||
ی
|
||||
ې
|
||||
ے
|
||||
ۓ
|
||||
ە
|
||||
١
|
||||
٢
|
||||
٣
|
||||
٤
|
||||
٥
|
||||
٦
|
||||
٧
|
||||
٨
|
||||
٩
|
||||
+8421
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,163 @@
|
||||
|
||||
!
|
||||
#
|
||||
$
|
||||
%
|
||||
&
|
||||
'
|
||||
(
|
||||
+
|
||||
,
|
||||
-
|
||||
.
|
||||
/
|
||||
0
|
||||
1
|
||||
2
|
||||
3
|
||||
4
|
||||
5
|
||||
6
|
||||
7
|
||||
8
|
||||
9
|
||||
:
|
||||
?
|
||||
@
|
||||
A
|
||||
B
|
||||
C
|
||||
D
|
||||
E
|
||||
F
|
||||
G
|
||||
H
|
||||
I
|
||||
J
|
||||
K
|
||||
L
|
||||
M
|
||||
N
|
||||
O
|
||||
P
|
||||
Q
|
||||
R
|
||||
S
|
||||
T
|
||||
U
|
||||
V
|
||||
W
|
||||
X
|
||||
Y
|
||||
Z
|
||||
_
|
||||
a
|
||||
b
|
||||
c
|
||||
d
|
||||
e
|
||||
f
|
||||
g
|
||||
h
|
||||
i
|
||||
j
|
||||
k
|
||||
l
|
||||
m
|
||||
n
|
||||
o
|
||||
p
|
||||
q
|
||||
r
|
||||
s
|
||||
t
|
||||
u
|
||||
v
|
||||
w
|
||||
x
|
||||
y
|
||||
z
|
||||
É
|
||||
é
|
||||
Ё
|
||||
Є
|
||||
І
|
||||
Ј
|
||||
Љ
|
||||
Ў
|
||||
А
|
||||
Б
|
||||
В
|
||||
Г
|
||||
Д
|
||||
Е
|
||||
Ж
|
||||
З
|
||||
И
|
||||
Й
|
||||
К
|
||||
Л
|
||||
М
|
||||
Н
|
||||
О
|
||||
П
|
||||
Р
|
||||
С
|
||||
Т
|
||||
У
|
||||
Ф
|
||||
Х
|
||||
Ц
|
||||
Ч
|
||||
Ш
|
||||
Щ
|
||||
Ъ
|
||||
Ы
|
||||
Ь
|
||||
Э
|
||||
Ю
|
||||
Я
|
||||
а
|
||||
б
|
||||
в
|
||||
г
|
||||
д
|
||||
е
|
||||
ж
|
||||
з
|
||||
и
|
||||
й
|
||||
к
|
||||
л
|
||||
м
|
||||
н
|
||||
о
|
||||
п
|
||||
р
|
||||
с
|
||||
т
|
||||
у
|
||||
ф
|
||||
х
|
||||
ц
|
||||
ч
|
||||
ш
|
||||
щ
|
||||
ъ
|
||||
ы
|
||||
ь
|
||||
э
|
||||
ю
|
||||
я
|
||||
ё
|
||||
ђ
|
||||
є
|
||||
і
|
||||
ј
|
||||
љ
|
||||
њ
|
||||
ћ
|
||||
ў
|
||||
џ
|
||||
Ґ
|
||||
ґ
|
||||
+167
@@ -0,0 +1,167 @@
|
||||
|
||||
!
|
||||
#
|
||||
$
|
||||
%
|
||||
&
|
||||
'
|
||||
(
|
||||
+
|
||||
,
|
||||
-
|
||||
.
|
||||
/
|
||||
0
|
||||
1
|
||||
2
|
||||
3
|
||||
4
|
||||
5
|
||||
6
|
||||
7
|
||||
8
|
||||
9
|
||||
:
|
||||
?
|
||||
@
|
||||
A
|
||||
B
|
||||
C
|
||||
D
|
||||
E
|
||||
F
|
||||
G
|
||||
H
|
||||
I
|
||||
J
|
||||
K
|
||||
L
|
||||
M
|
||||
N
|
||||
O
|
||||
P
|
||||
Q
|
||||
R
|
||||
S
|
||||
T
|
||||
U
|
||||
V
|
||||
W
|
||||
X
|
||||
Y
|
||||
Z
|
||||
_
|
||||
a
|
||||
b
|
||||
c
|
||||
d
|
||||
e
|
||||
f
|
||||
g
|
||||
h
|
||||
i
|
||||
j
|
||||
k
|
||||
l
|
||||
m
|
||||
n
|
||||
o
|
||||
p
|
||||
q
|
||||
r
|
||||
s
|
||||
t
|
||||
u
|
||||
v
|
||||
w
|
||||
x
|
||||
y
|
||||
z
|
||||
É
|
||||
é
|
||||
ँ
|
||||
ं
|
||||
ः
|
||||
अ
|
||||
आ
|
||||
इ
|
||||
ई
|
||||
उ
|
||||
ऊ
|
||||
ऋ
|
||||
ए
|
||||
ऐ
|
||||
ऑ
|
||||
ओ
|
||||
औ
|
||||
क
|
||||
ख
|
||||
ग
|
||||
घ
|
||||
ङ
|
||||
च
|
||||
छ
|
||||
ज
|
||||
झ
|
||||
ञ
|
||||
ट
|
||||
ठ
|
||||
ड
|
||||
ढ
|
||||
ण
|
||||
त
|
||||
थ
|
||||
द
|
||||
ध
|
||||
न
|
||||
ऩ
|
||||
प
|
||||
फ
|
||||
ब
|
||||
भ
|
||||
म
|
||||
य
|
||||
र
|
||||
ऱ
|
||||
ल
|
||||
ळ
|
||||
व
|
||||
श
|
||||
ष
|
||||
स
|
||||
ह
|
||||
़
|
||||
ा
|
||||
ि
|
||||
ी
|
||||
ु
|
||||
ू
|
||||
ृ
|
||||
ॅ
|
||||
े
|
||||
ै
|
||||
ॉ
|
||||
ो
|
||||
ौ
|
||||
्
|
||||
॒
|
||||
क़
|
||||
ख़
|
||||
ग़
|
||||
ज़
|
||||
ड़
|
||||
ढ़
|
||||
फ़
|
||||
ॠ
|
||||
।
|
||||
०
|
||||
१
|
||||
२
|
||||
३
|
||||
४
|
||||
५
|
||||
६
|
||||
७
|
||||
८
|
||||
९
|
||||
॰
|
||||
@@ -0,0 +1,95 @@
|
||||
0
|
||||
1
|
||||
2
|
||||
3
|
||||
4
|
||||
5
|
||||
6
|
||||
7
|
||||
8
|
||||
9
|
||||
:
|
||||
;
|
||||
<
|
||||
=
|
||||
>
|
||||
?
|
||||
@
|
||||
A
|
||||
B
|
||||
C
|
||||
D
|
||||
E
|
||||
F
|
||||
G
|
||||
H
|
||||
I
|
||||
J
|
||||
K
|
||||
L
|
||||
M
|
||||
N
|
||||
O
|
||||
P
|
||||
Q
|
||||
R
|
||||
S
|
||||
T
|
||||
U
|
||||
V
|
||||
W
|
||||
X
|
||||
Y
|
||||
Z
|
||||
[
|
||||
\
|
||||
]
|
||||
^
|
||||
_
|
||||
`
|
||||
a
|
||||
b
|
||||
c
|
||||
d
|
||||
e
|
||||
f
|
||||
g
|
||||
h
|
||||
i
|
||||
j
|
||||
k
|
||||
l
|
||||
m
|
||||
n
|
||||
o
|
||||
p
|
||||
q
|
||||
r
|
||||
s
|
||||
t
|
||||
u
|
||||
v
|
||||
w
|
||||
x
|
||||
y
|
||||
z
|
||||
{
|
||||
|
|
||||
}
|
||||
~
|
||||
!
|
||||
"
|
||||
#
|
||||
$
|
||||
%
|
||||
&
|
||||
'
|
||||
(
|
||||
)
|
||||
*
|
||||
+
|
||||
,
|
||||
-
|
||||
.
|
||||
/
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,153 @@
|
||||
k
|
||||
a
|
||||
_
|
||||
i
|
||||
m
|
||||
g
|
||||
/
|
||||
1
|
||||
2
|
||||
I
|
||||
L
|
||||
S
|
||||
V
|
||||
R
|
||||
C
|
||||
0
|
||||
v
|
||||
l
|
||||
6
|
||||
4
|
||||
8
|
||||
.
|
||||
j
|
||||
p
|
||||
ಗ
|
||||
ು
|
||||
ಣ
|
||||
ಪ
|
||||
ಡ
|
||||
ಿ
|
||||
ಸ
|
||||
ಲ
|
||||
ಾ
|
||||
ದ
|
||||
್
|
||||
7
|
||||
5
|
||||
3
|
||||
ವ
|
||||
ಷ
|
||||
ಬ
|
||||
ಹ
|
||||
ೆ
|
||||
9
|
||||
ಅ
|
||||
ಳ
|
||||
ನ
|
||||
ರ
|
||||
ಉ
|
||||
ಕ
|
||||
ಎ
|
||||
ೇ
|
||||
ಂ
|
||||
ೈ
|
||||
ೊ
|
||||
ೀ
|
||||
ಯ
|
||||
ೋ
|
||||
ತ
|
||||
ಶ
|
||||
ಭ
|
||||
ಧ
|
||||
ಚ
|
||||
ಜ
|
||||
ೂ
|
||||
ಮ
|
||||
ಒ
|
||||
ೃ
|
||||
ಥ
|
||||
ಇ
|
||||
ಟ
|
||||
ಖ
|
||||
ಆ
|
||||
ಞ
|
||||
ಫ
|
||||
-
|
||||
ಢ
|
||||
ಊ
|
||||
ಓ
|
||||
ಐ
|
||||
ಃ
|
||||
ಘ
|
||||
ಝ
|
||||
ೌ
|
||||
ಠ
|
||||
ಛ
|
||||
ಔ
|
||||
ಏ
|
||||
ಈ
|
||||
ಋ
|
||||
೨
|
||||
೦
|
||||
೧
|
||||
೮
|
||||
೯
|
||||
೪
|
||||
,
|
||||
೫
|
||||
೭
|
||||
೩
|
||||
೬
|
||||
ಙ
|
||||
s
|
||||
c
|
||||
e
|
||||
n
|
||||
w
|
||||
o
|
||||
u
|
||||
t
|
||||
d
|
||||
E
|
||||
A
|
||||
T
|
||||
B
|
||||
Z
|
||||
N
|
||||
G
|
||||
O
|
||||
q
|
||||
z
|
||||
r
|
||||
x
|
||||
P
|
||||
K
|
||||
M
|
||||
J
|
||||
U
|
||||
D
|
||||
f
|
||||
F
|
||||
h
|
||||
b
|
||||
W
|
||||
Y
|
||||
y
|
||||
H
|
||||
X
|
||||
Q
|
||||
'
|
||||
#
|
||||
&
|
||||
!
|
||||
@
|
||||
$
|
||||
:
|
||||
%
|
||||
é
|
||||
É
|
||||
(
|
||||
?
|
||||
+
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,185 @@
|
||||
|
||||
!
|
||||
"
|
||||
#
|
||||
$
|
||||
%
|
||||
&
|
||||
'
|
||||
(
|
||||
)
|
||||
*
|
||||
+
|
||||
,
|
||||
-
|
||||
.
|
||||
/
|
||||
0
|
||||
1
|
||||
2
|
||||
3
|
||||
4
|
||||
5
|
||||
6
|
||||
7
|
||||
8
|
||||
9
|
||||
:
|
||||
;
|
||||
<
|
||||
=
|
||||
>
|
||||
?
|
||||
@
|
||||
A
|
||||
B
|
||||
C
|
||||
D
|
||||
E
|
||||
F
|
||||
G
|
||||
H
|
||||
I
|
||||
J
|
||||
K
|
||||
L
|
||||
M
|
||||
N
|
||||
O
|
||||
P
|
||||
Q
|
||||
R
|
||||
S
|
||||
T
|
||||
U
|
||||
V
|
||||
W
|
||||
X
|
||||
Y
|
||||
Z
|
||||
[
|
||||
]
|
||||
_
|
||||
`
|
||||
a
|
||||
b
|
||||
c
|
||||
d
|
||||
e
|
||||
f
|
||||
g
|
||||
h
|
||||
i
|
||||
j
|
||||
k
|
||||
l
|
||||
m
|
||||
n
|
||||
o
|
||||
p
|
||||
q
|
||||
r
|
||||
s
|
||||
t
|
||||
u
|
||||
v
|
||||
w
|
||||
x
|
||||
y
|
||||
z
|
||||
{
|
||||
}
|
||||
¡
|
||||
£
|
||||
§
|
||||
ª
|
||||
«
|
||||
|
||||
°
|
||||
²
|
||||
³
|
||||
´
|
||||
µ
|
||||
·
|
||||
º
|
||||
»
|
||||
¿
|
||||
À
|
||||
Á
|
||||
Â
|
||||
Ä
|
||||
Å
|
||||
Ç
|
||||
È
|
||||
É
|
||||
Ê
|
||||
Ë
|
||||
Ì
|
||||
Í
|
||||
Î
|
||||
Ï
|
||||
Ò
|
||||
Ó
|
||||
Ô
|
||||
Õ
|
||||
Ö
|
||||
Ú
|
||||
Ü
|
||||
Ý
|
||||
ß
|
||||
à
|
||||
á
|
||||
â
|
||||
ã
|
||||
ä
|
||||
å
|
||||
æ
|
||||
ç
|
||||
è
|
||||
é
|
||||
ê
|
||||
ë
|
||||
ì
|
||||
í
|
||||
î
|
||||
ï
|
||||
ñ
|
||||
ò
|
||||
ó
|
||||
ô
|
||||
õ
|
||||
ö
|
||||
ø
|
||||
ù
|
||||
ú
|
||||
û
|
||||
ü
|
||||
ý
|
||||
ą
|
||||
Ć
|
||||
ć
|
||||
Č
|
||||
č
|
||||
Đ
|
||||
đ
|
||||
ę
|
||||
ı
|
||||
Ł
|
||||
ł
|
||||
ō
|
||||
Œ
|
||||
œ
|
||||
Š
|
||||
š
|
||||
Ÿ
|
||||
Ž
|
||||
ž
|
||||
ʒ
|
||||
β
|
||||
δ
|
||||
ε
|
||||
з
|
||||
Ṡ
|
||||
‘
|
||||
€
|
||||
™
|
||||
+6623
File diff suppressed because it is too large
Load Diff
+15629
File diff suppressed because it is too large
Load Diff
+18383
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,128 @@
|
||||
t
|
||||
a
|
||||
_
|
||||
i
|
||||
m
|
||||
g
|
||||
/
|
||||
3
|
||||
I
|
||||
L
|
||||
S
|
||||
V
|
||||
R
|
||||
C
|
||||
2
|
||||
0
|
||||
1
|
||||
v
|
||||
l
|
||||
9
|
||||
7
|
||||
8
|
||||
.
|
||||
j
|
||||
p
|
||||
ப
|
||||
ூ
|
||||
த
|
||||
ம
|
||||
ி
|
||||
வ
|
||||
ர
|
||||
்
|
||||
ந
|
||||
ோ
|
||||
ன
|
||||
6
|
||||
ஆ
|
||||
ற
|
||||
ல
|
||||
5
|
||||
ள
|
||||
ா
|
||||
ொ
|
||||
ழ
|
||||
ு
|
||||
4
|
||||
ெ
|
||||
ண
|
||||
க
|
||||
ட
|
||||
ை
|
||||
ே
|
||||
ச
|
||||
ய
|
||||
ஒ
|
||||
இ
|
||||
அ
|
||||
ங
|
||||
உ
|
||||
ீ
|
||||
ஞ
|
||||
எ
|
||||
ஓ
|
||||
ஃ
|
||||
ஜ
|
||||
ஷ
|
||||
ஸ
|
||||
ஏ
|
||||
ஊ
|
||||
ஹ
|
||||
ஈ
|
||||
ஐ
|
||||
ௌ
|
||||
ஔ
|
||||
s
|
||||
c
|
||||
e
|
||||
n
|
||||
w
|
||||
F
|
||||
T
|
||||
O
|
||||
P
|
||||
K
|
||||
A
|
||||
N
|
||||
G
|
||||
Y
|
||||
E
|
||||
M
|
||||
H
|
||||
U
|
||||
B
|
||||
o
|
||||
b
|
||||
D
|
||||
d
|
||||
r
|
||||
W
|
||||
u
|
||||
y
|
||||
f
|
||||
X
|
||||
k
|
||||
q
|
||||
h
|
||||
J
|
||||
z
|
||||
Z
|
||||
Q
|
||||
x
|
||||
-
|
||||
'
|
||||
$
|
||||
,
|
||||
%
|
||||
@
|
||||
é
|
||||
!
|
||||
#
|
||||
+
|
||||
É
|
||||
&
|
||||
:
|
||||
(
|
||||
?
|
||||
|
||||
@@ -0,0 +1,151 @@
|
||||
t
|
||||
e
|
||||
_
|
||||
i
|
||||
m
|
||||
g
|
||||
/
|
||||
5
|
||||
I
|
||||
L
|
||||
S
|
||||
V
|
||||
R
|
||||
C
|
||||
2
|
||||
0
|
||||
1
|
||||
v
|
||||
a
|
||||
l
|
||||
3
|
||||
4
|
||||
8
|
||||
9
|
||||
.
|
||||
j
|
||||
p
|
||||
త
|
||||
ె
|
||||
ర
|
||||
క
|
||||
్
|
||||
ి
|
||||
ం
|
||||
చ
|
||||
ే
|
||||
ద
|
||||
ు
|
||||
7
|
||||
6
|
||||
ఉ
|
||||
ా
|
||||
మ
|
||||
ట
|
||||
ో
|
||||
వ
|
||||
ప
|
||||
ల
|
||||
శ
|
||||
ఆ
|
||||
య
|
||||
ై
|
||||
భ
|
||||
'
|
||||
ీ
|
||||
గ
|
||||
ూ
|
||||
డ
|
||||
ధ
|
||||
హ
|
||||
న
|
||||
జ
|
||||
స
|
||||
[
|
||||
|
||||
ష
|
||||
అ
|
||||
ణ
|
||||
ఫ
|
||||
బ
|
||||
ఎ
|
||||
;
|
||||
ళ
|
||||
థ
|
||||
ొ
|
||||
ఠ
|
||||
ృ
|
||||
ఒ
|
||||
ఇ
|
||||
ః
|
||||
ఊ
|
||||
ఖ
|
||||
-
|
||||
ఐ
|
||||
ఘ
|
||||
ౌ
|
||||
ఏ
|
||||
ఈ
|
||||
ఛ
|
||||
,
|
||||
ఓ
|
||||
ఞ
|
||||
|
|
||||
?
|
||||
:
|
||||
ఢ
|
||||
"
|
||||
(
|
||||
”
|
||||
!
|
||||
+
|
||||
)
|
||||
*
|
||||
=
|
||||
&
|
||||
“
|
||||
€
|
||||
]
|
||||
£
|
||||
$
|
||||
s
|
||||
c
|
||||
n
|
||||
w
|
||||
k
|
||||
J
|
||||
G
|
||||
u
|
||||
d
|
||||
r
|
||||
E
|
||||
o
|
||||
h
|
||||
y
|
||||
b
|
||||
f
|
||||
B
|
||||
M
|
||||
O
|
||||
T
|
||||
N
|
||||
D
|
||||
P
|
||||
A
|
||||
F
|
||||
x
|
||||
W
|
||||
Y
|
||||
U
|
||||
H
|
||||
K
|
||||
X
|
||||
z
|
||||
Z
|
||||
Q
|
||||
q
|
||||
É
|
||||
%
|
||||
#
|
||||
@
|
||||
é
|
||||
@@ -0,0 +1,65 @@
|
||||
lang:
|
||||
ch_lite:
|
||||
det: ch_PP-OCRv3_det_infer.pth
|
||||
rec: ch_PP-OCRv5_rec_infer.pth
|
||||
dict: ppocrv5_dict.txt
|
||||
ch_lite_v4:
|
||||
det: ch_PP-OCRv3_det_infer.pth
|
||||
rec: ch_PP-OCRv4_rec_infer.pth
|
||||
dict: ppocr_keys_v1.txt
|
||||
ch_server:
|
||||
det: ch_PP-OCRv3_det_infer.pth
|
||||
rec: ch_PP-OCRv5_rec_server_infer.pth
|
||||
dict: ppocrv5_dict.txt
|
||||
ch_server_v4:
|
||||
det: ch_PP-OCRv3_det_infer.pth
|
||||
rec: ch_PP-OCRv4_rec_server_infer.pth
|
||||
dict: ppocr_keys_v1.txt
|
||||
ch:
|
||||
det: ch_PP-OCRv3_det_infer.pth
|
||||
rec: ch_PP-OCRv4_rec_server_doc_infer.pth
|
||||
dict: ppocrv4_doc_dict.txt
|
||||
en:
|
||||
det: en_PP-OCRv3_det_infer.pth
|
||||
rec: en_PP-OCRv4_rec_infer.pth
|
||||
dict: en_dict.txt
|
||||
korean:
|
||||
det: Multilingual_PP-OCRv3_det_infer.pth
|
||||
rec: korean_PP-OCRv3_rec_infer.pth
|
||||
dict: korean_dict.txt
|
||||
japan:
|
||||
det: Multilingual_PP-OCRv3_det_infer.pth
|
||||
rec: japan_PP-OCRv3_rec_infer.pth
|
||||
dict: japan_dict.txt
|
||||
chinese_cht:
|
||||
det: Multilingual_PP-OCRv3_det_infer.pth
|
||||
rec: chinese_cht_PP-OCRv3_rec_infer.pth
|
||||
dict: chinese_cht_dict.txt
|
||||
ta:
|
||||
det: Multilingual_PP-OCRv3_det_infer.pth
|
||||
rec: ta_PP-OCRv3_rec_infer.pth
|
||||
dict: ta_dict.txt
|
||||
te:
|
||||
det: Multilingual_PP-OCRv3_det_infer.pth
|
||||
rec: te_PP-OCRv3_rec_infer.pth
|
||||
dict: te_dict.txt
|
||||
ka:
|
||||
det: Multilingual_PP-OCRv3_det_infer.pth
|
||||
rec: ka_PP-OCRv3_rec_infer.pth
|
||||
dict: ka_dict.txt
|
||||
latin:
|
||||
det: en_PP-OCRv3_det_infer.pth
|
||||
rec: latin_PP-OCRv3_rec_infer.pth
|
||||
dict: latin_dict.txt
|
||||
arabic:
|
||||
det: Multilingual_PP-OCRv3_det_infer.pth
|
||||
rec: arabic_PP-OCRv3_rec_infer.pth
|
||||
dict: arabic_dict.txt
|
||||
cyrillic:
|
||||
det: Multilingual_PP-OCRv3_det_infer.pth
|
||||
rec: cyrillic_PP-OCRv3_rec_infer.pth
|
||||
dict: cyrillic_dict.txt
|
||||
devanagari:
|
||||
det: Multilingual_PP-OCRv3_det_infer.pth
|
||||
rec: devanagari_PP-OCRv3_rec_infer.pth
|
||||
dict: devanagari_dict.txt
|
||||
@@ -0,0 +1 @@
|
||||
# Copyright (c) Opendatalab. All rights reserved.
|
||||
@@ -0,0 +1 @@
|
||||
# Copyright (c) Opendatalab. All rights reserved.
|
||||
@@ -0,0 +1,106 @@
|
||||
import cv2
|
||||
import copy
|
||||
import numpy as np
|
||||
import math
|
||||
import time
|
||||
import torch
|
||||
from ...pytorchocr.base_ocr_v20 import BaseOCRV20
|
||||
from . import pytorchocr_utility as utility
|
||||
from ...pytorchocr.postprocess import build_post_process
|
||||
|
||||
|
||||
class TextClassifier(BaseOCRV20):
|
||||
def __init__(self, args, **kwargs):
|
||||
self.device = args.device
|
||||
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.weights_path = args.cls_model_path
|
||||
self.yaml_path = args.cls_yaml_path
|
||||
network_config = utility.get_arch_config(self.weights_path)
|
||||
super(TextClassifier, self).__init__(network_config, **kwargs)
|
||||
|
||||
self.cls_image_shape = [int(v) for v in args.cls_image_shape.split(",")]
|
||||
|
||||
self.limited_max_width = args.limited_max_width
|
||||
self.limited_min_width = args.limited_min_width
|
||||
|
||||
self.load_pytorch_weights(self.weights_path)
|
||||
self.net.eval()
|
||||
self.net.to(self.device)
|
||||
|
||||
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)
|
||||
imgW = max(min(imgW, self.limited_max_width), self.limited_min_width)
|
||||
ratio_imgH = math.ceil(imgH * ratio)
|
||||
ratio_imgH = max(ratio_imgH, self.limited_min_width)
|
||||
if ratio_imgH > 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
|
||||
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()
|
||||
starttime = time.time()
|
||||
|
||||
with torch.no_grad():
|
||||
inp = torch.from_numpy(norm_img_batch)
|
||||
inp = inp.to(self.device)
|
||||
prob_out = self.net(inp)
|
||||
prob_out = prob_out.cpu().numpy()
|
||||
|
||||
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
|
||||
@@ -0,0 +1,217 @@
|
||||
import sys
|
||||
|
||||
import numpy as np
|
||||
import time
|
||||
import torch
|
||||
from ...pytorchocr.base_ocr_v20 import BaseOCRV20
|
||||
from . import pytorchocr_utility as utility
|
||||
from ...pytorchocr.data import create_operators, transform
|
||||
from ...pytorchocr.postprocess import build_post_process
|
||||
|
||||
|
||||
class TextDetector(BaseOCRV20):
|
||||
def __init__(self, args, **kwargs):
|
||||
self.args = args
|
||||
self.det_algorithm = args.det_algorithm
|
||||
self.device = args.device
|
||||
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
|
||||
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
|
||||
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
|
||||
self.det_sast_polygon = args.det_sast_polygon
|
||||
if self.det_sast_polygon:
|
||||
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_pse_box_type
|
||||
postprocess_params["scale"] = args.det_pse_scale
|
||||
self.det_pse_box_type = args.det_pse_box_type
|
||||
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_fce_box_type
|
||||
else:
|
||||
print("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.weights_path = args.det_model_path
|
||||
self.yaml_path = args.det_yaml_path
|
||||
network_config = utility.get_arch_config(self.weights_path)
|
||||
super(TextDetector, self).__init__(network_config, **kwargs)
|
||||
self.load_pytorch_weights(self.weights_path)
|
||||
self.net.eval()
|
||||
self.net.to(self.device)
|
||||
|
||||
def order_points_clockwise(self, pts):
|
||||
"""
|
||||
reference from: https://github.com/jrosebr1/imutils/blob/master/imutils/perspective.py
|
||||
# sort the points based on their x-coordinates
|
||||
"""
|
||||
xSorted = pts[np.argsort(pts[:, 0]), :]
|
||||
|
||||
# grab the left-most and right-most points from the sorted
|
||||
# x-roodinate points
|
||||
leftMost = xSorted[:2, :]
|
||||
rightMost = xSorted[2:, :]
|
||||
|
||||
# now, sort the left-most coordinates according to their
|
||||
# y-coordinates so we can grab the top-left and bottom-left
|
||||
# points, respectively
|
||||
leftMost = leftMost[np.argsort(leftMost[:, 1]), :]
|
||||
(tl, bl) = leftMost
|
||||
|
||||
rightMost = rightMost[np.argsort(rightMost[:, 1]), :]
|
||||
(tr, br) = rightMost
|
||||
|
||||
rect = np.array([tl, tr, br, bl], dtype="float32")
|
||||
return rect
|
||||
|
||||
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:
|
||||
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:
|
||||
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()
|
||||
|
||||
with torch.no_grad():
|
||||
inp = torch.from_numpy(img)
|
||||
inp = inp.to(self.device)
|
||||
outputs = self.net(inp)
|
||||
|
||||
preds = {}
|
||||
if self.det_algorithm == "EAST":
|
||||
preds['f_geo'] = outputs['f_geo'].cpu().numpy()
|
||||
preds['f_score'] = outputs['f_score'].cpu().numpy()
|
||||
elif self.det_algorithm == 'SAST':
|
||||
preds['f_border'] = outputs['f_border'].cpu().numpy()
|
||||
preds['f_score'] = outputs['f_score'].cpu().numpy()
|
||||
preds['f_tco'] = outputs['f_tco'].cpu().numpy()
|
||||
preds['f_tvo'] = outputs['f_tvo'].cpu().numpy()
|
||||
elif self.det_algorithm in ['DB', 'PSE', 'DB++']:
|
||||
preds['maps'] = outputs['maps'].cpu().numpy()
|
||||
elif self.det_algorithm == 'FCE':
|
||||
for i, (k, output) in enumerate(outputs.items()):
|
||||
preds['level_{}'.format(i)] = output
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
post_result = self.postprocess_op(preds, shape_list)
|
||||
dt_boxes = post_result[0]['points']
|
||||
if (self.det_algorithm == "SAST" and
|
||||
self.det_sast_polygon) or (self.det_algorithm in ["PSE", "FCE"] and
|
||||
self.postprocess_op.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)
|
||||
|
||||
elapse = time.time() - starttime
|
||||
return dt_boxes, elapse
|
||||
@@ -0,0 +1,446 @@
|
||||
from PIL import Image
|
||||
import cv2
|
||||
import numpy as np
|
||||
import math
|
||||
import time
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from ...pytorchocr.base_ocr_v20 import BaseOCRV20
|
||||
from . import pytorchocr_utility as utility
|
||||
from ...pytorchocr.postprocess import build_post_process
|
||||
|
||||
|
||||
class TextRecognizer(BaseOCRV20):
|
||||
def __init__(self, args, **kwargs):
|
||||
self.device = args.device
|
||||
self.rec_image_shape = [int(v) for v in args.rec_image_shape.split(",")]
|
||||
self.character_type = args.rec_char_type
|
||||
self.rec_batch_num = args.rec_batch_num
|
||||
self.rec_algorithm = args.rec_algorithm
|
||||
self.max_text_length = args.max_text_length
|
||||
postprocess_params = {
|
||||
'name': 'CTCLabelDecode',
|
||||
"character_type": args.rec_char_type,
|
||||
"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_type": args.rec_char_type,
|
||||
"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_type": args.rec_char_type,
|
||||
"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 == 'ViTSTR':
|
||||
postprocess_params = {
|
||||
'name': 'ViTSTRLabelDecode',
|
||||
"character_dict_path": args.rec_char_dict_path,
|
||||
"use_space_char": args.use_space_char
|
||||
}
|
||||
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 == 'RFL':
|
||||
postprocess_params = {
|
||||
'name': 'RFLLabelDecode',
|
||||
"character_dict_path": None,
|
||||
"use_space_char": args.use_space_char
|
||||
}
|
||||
self.postprocess_op = build_post_process(postprocess_params)
|
||||
|
||||
self.limited_max_width = args.limited_max_width
|
||||
self.limited_min_width = args.limited_min_width
|
||||
|
||||
self.weights_path = args.rec_model_path
|
||||
self.yaml_path = args.rec_yaml_path
|
||||
|
||||
network_config = utility.get_arch_config(self.weights_path)
|
||||
weights = self.read_pytorch_weights(self.weights_path)
|
||||
|
||||
self.out_channels = self.get_out_channels(weights)
|
||||
if self.rec_algorithm == 'NRTR':
|
||||
self.out_channels = list(weights.values())[-1].numpy().shape[0]
|
||||
elif self.rec_algorithm == 'SAR':
|
||||
self.out_channels = list(weights.values())[-3].numpy().shape[0]
|
||||
|
||||
kwargs['out_channels'] = self.out_channels
|
||||
super(TextRecognizer, self).__init__(network_config, **kwargs)
|
||||
|
||||
self.load_state_dict(weights)
|
||||
self.net.eval()
|
||||
self.net.to(self.device)
|
||||
|
||||
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.ANTIALIAS)
|
||||
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.
|
||||
else:
|
||||
norm_img = norm_img.astype(np.float32) / 128. - 1.
|
||||
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]
|
||||
max_wh_ratio = max(max_wh_ratio, imgW / imgH)
|
||||
imgW = int((imgH * max_wh_ratio))
|
||||
imgW = max(min(imgW, self.limited_max_width), self.limited_min_width)
|
||||
h, w = img.shape[:2]
|
||||
ratio = w / float(h)
|
||||
ratio_imgH = math.ceil(imgH * ratio)
|
||||
ratio_imgH = max(ratio_imgH, self.limited_min_width)
|
||||
if ratio_imgH > imgW:
|
||||
resized_w = imgW
|
||||
else:
|
||||
resized_w = int(ratio_imgH)
|
||||
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_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_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 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 __call__(self, img_list, tqdm_enable=False):
|
||||
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 = []
|
||||
rec_res = [['', 0.0]] * img_num
|
||||
batch_num = self.rec_batch_num
|
||||
elapse = 0
|
||||
# for beg_img_no in range(0, img_num, batch_num):
|
||||
with tqdm(total=img_num, desc='OCR-rec Predict', disable=not tqdm_enable) as pbar:
|
||||
index = 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
|
||||
for ino in range(beg_img_no, end_img_no):
|
||||
# h, w = img_list[ino].shape[0:2]
|
||||
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):
|
||||
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 = []
|
||||
valid_ratios.append(valid_ratio)
|
||||
norm_img_batch.append(norm_img)
|
||||
|
||||
elif self.rec_algorithm == "SVTR":
|
||||
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 == "SRN":
|
||||
norm_img = self.process_image_srn(img_list[indices[ino]],
|
||||
self.rec_image_shape, 8,
|
||||
self.max_text_length)
|
||||
encoder_word_pos_list = []
|
||||
gsrm_word_pos_list = []
|
||||
gsrm_slf_attn_bias1_list = []
|
||||
gsrm_slf_attn_bias2_list = []
|
||||
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 == "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)
|
||||
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.rec_algorithm == "SRN":
|
||||
starttime = time.time()
|
||||
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)
|
||||
|
||||
with torch.no_grad():
|
||||
inp = torch.from_numpy(norm_img_batch)
|
||||
encoder_word_pos_inp = torch.from_numpy(encoder_word_pos_list)
|
||||
gsrm_word_pos_inp = torch.from_numpy(gsrm_word_pos_list)
|
||||
gsrm_slf_attn_bias1_inp = torch.from_numpy(gsrm_slf_attn_bias1_list)
|
||||
gsrm_slf_attn_bias2_inp = torch.from_numpy(gsrm_slf_attn_bias2_list)
|
||||
|
||||
inp = inp.to(self.device)
|
||||
encoder_word_pos_inp = encoder_word_pos_inp.to(self.device)
|
||||
gsrm_word_pos_inp = gsrm_word_pos_inp.to(self.device)
|
||||
gsrm_slf_attn_bias1_inp = gsrm_slf_attn_bias1_inp.to(self.device)
|
||||
gsrm_slf_attn_bias2_inp = gsrm_slf_attn_bias2_inp.to(self.device)
|
||||
|
||||
backbone_out = self.net.backbone(inp) # backbone_feat
|
||||
prob_out = self.net.head(backbone_out, [encoder_word_pos_inp, gsrm_word_pos_inp, gsrm_slf_attn_bias1_inp, gsrm_slf_attn_bias2_inp])
|
||||
# preds = {"predict": prob_out[2]}
|
||||
preds = {"predict": prob_out["predict"]}
|
||||
|
||||
elif self.rec_algorithm == "SAR":
|
||||
starttime = time.time()
|
||||
# valid_ratios = np.concatenate(valid_ratios)
|
||||
# inputs = [
|
||||
# norm_img_batch,
|
||||
# valid_ratios,
|
||||
# ]
|
||||
|
||||
with torch.no_grad():
|
||||
inp = torch.from_numpy(norm_img_batch)
|
||||
inp = inp.to(self.device)
|
||||
preds = self.net(inp)
|
||||
|
||||
elif self.rec_algorithm == "CAN":
|
||||
starttime = time.time()
|
||||
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]
|
||||
|
||||
inp = [torch.from_numpy(e_i) for e_i in inputs]
|
||||
inp = [e_i.to(self.device) for e_i in inp]
|
||||
with torch.no_grad():
|
||||
outputs = self.net(inp)
|
||||
outputs = [v.cpu().numpy() for k, v in enumerate(outputs)]
|
||||
|
||||
preds = outputs
|
||||
|
||||
else:
|
||||
starttime = time.time()
|
||||
|
||||
with torch.no_grad():
|
||||
inp = torch.from_numpy(norm_img_batch)
|
||||
inp = inp.to(self.device)
|
||||
prob_out = self.net(inp)
|
||||
|
||||
if isinstance(prob_out, list):
|
||||
preds = [v.cpu().numpy() for v in prob_out]
|
||||
else:
|
||||
preds = prob_out.cpu().numpy()
|
||||
|
||||
rec_result = self.postprocess_op(preds)
|
||||
for rno in range(len(rec_result)):
|
||||
rec_res[indices[beg_img_no + rno]] = rec_result[rno]
|
||||
elapse += time.time() - starttime
|
||||
|
||||
# 更新进度条,每次增加batch_size,但要注意最后一个batch可能不足batch_size
|
||||
current_batch_size = min(batch_num, img_num - index * batch_num)
|
||||
index += 1
|
||||
pbar.update(current_batch_size)
|
||||
|
||||
# Fix NaN values in recognition results
|
||||
for i in range(len(rec_res)):
|
||||
text, score = rec_res[i]
|
||||
if isinstance(score, float) and math.isnan(score):
|
||||
rec_res[i] = (text, 0.0)
|
||||
|
||||
return rec_res, elapse
|
||||
@@ -0,0 +1,104 @@
|
||||
import cv2
|
||||
import copy
|
||||
import numpy as np
|
||||
|
||||
from . import predict_rec
|
||||
from . import predict_det
|
||||
from . import predict_cls
|
||||
|
||||
|
||||
class TextSystem(object):
|
||||
def __init__(self, args, **kwargs):
|
||||
self.text_detector = predict_det.TextDetector(args, **kwargs)
|
||||
self.text_recognizer = predict_rec.TextRecognizer(args, **kwargs)
|
||||
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, **kwargs)
|
||||
|
||||
def get_rotate_crop_image(self, img, points):
|
||||
'''
|
||||
img_height, img_width = img.shape[0:2]
|
||||
left = int(np.min(points[:, 0]))
|
||||
right = int(np.max(points[:, 0]))
|
||||
top = int(np.min(points[:, 1]))
|
||||
bottom = int(np.max(points[:, 1]))
|
||||
img_crop = img[top:bottom, left:right, :].copy()
|
||||
points[:, 0] = points[:, 0] - left
|
||||
points[:, 1] = points[:, 1] - top
|
||||
'''
|
||||
img_crop_width = int(
|
||||
max(
|
||||
np.linalg.norm(points[0] - points[1]),
|
||||
np.linalg.norm(points[2] - points[3])))
|
||||
img_crop_height = int(
|
||||
max(
|
||||
np.linalg.norm(points[0] - points[3]),
|
||||
np.linalg.norm(points[1] - points[2])))
|
||||
pts_std = np.float32([[0, 0], [img_crop_width, 0],
|
||||
[img_crop_width, img_crop_height],
|
||||
[0, img_crop_height]])
|
||||
M = cv2.getPerspectiveTransform(points, pts_std)
|
||||
dst_img = cv2.warpPerspective(
|
||||
img,
|
||||
M, (img_crop_width, img_crop_height),
|
||||
borderMode=cv2.BORDER_REPLICATE,
|
||||
flags=cv2.INTER_CUBIC)
|
||||
dst_img_height, dst_img_width = dst_img.shape[0:2]
|
||||
if dst_img_height * 1.0 / dst_img_width >= 1.5:
|
||||
dst_img = np.rot90(dst_img)
|
||||
return dst_img
|
||||
|
||||
def __call__(self, img):
|
||||
ori_im = img.copy()
|
||||
dt_boxes, elapse = self.text_detector(img)
|
||||
print("dt_boxes num : {}, elapse : {}".format(
|
||||
len(dt_boxes), elapse))
|
||||
if dt_boxes is None:
|
||||
return None, None
|
||||
img_crop_list = []
|
||||
|
||||
dt_boxes = sorted_boxes(dt_boxes)
|
||||
|
||||
for bno in range(len(dt_boxes)):
|
||||
tmp_box = copy.deepcopy(dt_boxes[bno])
|
||||
img_crop = self.get_rotate_crop_image(ori_im, tmp_box)
|
||||
img_crop_list.append(img_crop)
|
||||
if self.use_angle_cls:
|
||||
img_crop_list, angle_list, elapse = self.text_classifier(
|
||||
img_crop_list)
|
||||
print("cls num : {}, elapse : {}".format(
|
||||
len(img_crop_list), elapse))
|
||||
|
||||
rec_res, elapse = self.text_recognizer(img_crop_list)
|
||||
print("rec_res num : {}, elapse : {}".format(
|
||||
len(rec_res), elapse))
|
||||
# self.print_draw_crop_rec_res(img_crop_list, rec_res)
|
||||
filter_boxes, filter_rec_res = [], []
|
||||
for box, rec_reuslt in zip(dt_boxes, rec_res):
|
||||
text, score = rec_reuslt
|
||||
if score >= self.drop_score:
|
||||
filter_boxes.append(box)
|
||||
filter_rec_res.append(rec_reuslt)
|
||||
return filter_boxes, filter_rec_res
|
||||
|
||||
|
||||
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):
|
||||
if abs(_boxes[i + 1][0][1] - _boxes[i][0][1]) < 10 and \
|
||||
(_boxes[i + 1][0][0] < _boxes[i][0][0]):
|
||||
tmp = _boxes[i]
|
||||
_boxes[i] = _boxes[i + 1]
|
||||
_boxes[i + 1] = tmp
|
||||
return _boxes
|
||||
@@ -0,0 +1,227 @@
|
||||
import os
|
||||
import math
|
||||
from pathlib import Path
|
||||
import numpy as np
|
||||
import cv2
|
||||
import argparse
|
||||
|
||||
|
||||
root_dir = Path(__file__).resolve().parent.parent.parent
|
||||
DEFAULT_CFG_PATH = root_dir / "pytorchocr" / "utils" / "resources" / "arch_config.yaml"
|
||||
|
||||
|
||||
def init_args():
|
||||
def str2bool(v):
|
||||
return v.lower() in ("true", "t", "1")
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
# params for prediction engine
|
||||
parser.add_argument("--use_gpu", type=str2bool, default=False)
|
||||
parser.add_argument("--det", type=str2bool, default=True)
|
||||
parser.add_argument("--rec", type=str2bool, default=True)
|
||||
parser.add_argument("--device", type=str, default='cpu')
|
||||
# parser.add_argument("--ir_optim", type=str2bool, default=True)
|
||||
# parser.add_argument("--use_tensorrt", type=str2bool, default=False)
|
||||
# parser.add_argument("--use_fp16", type=str2bool, default=False)
|
||||
parser.add_argument("--gpu_mem", type=int, default=500)
|
||||
parser.add_argument("--warmup", type=str2bool, default=False)
|
||||
|
||||
# params for text detector
|
||||
parser.add_argument("--image_dir", type=str)
|
||||
parser.add_argument("--det_algorithm", type=str, default='DB')
|
||||
parser.add_argument("--det_model_path", type=str)
|
||||
parser.add_argument("--det_limit_side_len", type=float, default=960)
|
||||
parser.add_argument("--det_limit_type", type=str, default='max')
|
||||
|
||||
# DB parmas
|
||||
parser.add_argument("--det_db_thresh", type=float, default=0.3)
|
||||
parser.add_argument("--det_db_box_thresh", type=float, default=0.6)
|
||||
parser.add_argument("--det_db_unclip_ratio", type=float, default=1.5)
|
||||
parser.add_argument("--max_batch_size", type=int, default=10)
|
||||
parser.add_argument("--use_dilation", type=str2bool, default=False)
|
||||
parser.add_argument("--det_db_score_mode", type=str, default="fast")
|
||||
|
||||
# EAST parmas
|
||||
parser.add_argument("--det_east_score_thresh", type=float, default=0.8)
|
||||
parser.add_argument("--det_east_cover_thresh", type=float, default=0.1)
|
||||
parser.add_argument("--det_east_nms_thresh", type=float, default=0.2)
|
||||
|
||||
# SAST parmas
|
||||
parser.add_argument("--det_sast_score_thresh", type=float, default=0.5)
|
||||
parser.add_argument("--det_sast_nms_thresh", type=float, default=0.2)
|
||||
parser.add_argument("--det_sast_polygon", type=str2bool, default=False)
|
||||
|
||||
# PSE parmas
|
||||
parser.add_argument("--det_pse_thresh", type=float, default=0)
|
||||
parser.add_argument("--det_pse_box_thresh", type=float, default=0.85)
|
||||
parser.add_argument("--det_pse_min_area", type=float, default=16)
|
||||
parser.add_argument("--det_pse_box_type", type=str, default='box')
|
||||
parser.add_argument("--det_pse_scale", type=int, default=1)
|
||||
|
||||
# FCE parmas
|
||||
parser.add_argument("--scales", type=list, default=[8, 16, 32])
|
||||
parser.add_argument("--alpha", type=float, default=1.0)
|
||||
parser.add_argument("--beta", type=float, default=1.0)
|
||||
parser.add_argument("--fourier_degree", type=int, default=5)
|
||||
parser.add_argument("--det_fce_box_type", type=str, default='poly')
|
||||
|
||||
# params for text recognizer
|
||||
parser.add_argument("--rec_algorithm", type=str, default='CRNN')
|
||||
parser.add_argument("--rec_model_path", type=str)
|
||||
parser.add_argument("--rec_image_inverse", type=str2bool, default=True)
|
||||
parser.add_argument("--rec_image_shape", type=str, default="3, 48, 320")
|
||||
parser.add_argument("--rec_char_type", type=str, default='ch')
|
||||
parser.add_argument("--rec_batch_num", type=int, default=6)
|
||||
parser.add_argument("--max_text_length", type=int, default=25)
|
||||
|
||||
parser.add_argument("--use_space_char", type=str2bool, default=True)
|
||||
parser.add_argument("--drop_score", type=float, default=0.5)
|
||||
parser.add_argument("--limited_max_width", type=int, default=1280)
|
||||
parser.add_argument("--limited_min_width", type=int, default=16)
|
||||
|
||||
parser.add_argument(
|
||||
"--vis_font_path", type=str,
|
||||
default=os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))), 'doc/fonts/simfang.ttf'))
|
||||
parser.add_argument(
|
||||
"--rec_char_dict_path",
|
||||
type=str,
|
||||
default=os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))),
|
||||
'pytorchocr/utils/ppocr_keys_v1.txt'))
|
||||
|
||||
# params for text classifier
|
||||
parser.add_argument("--use_angle_cls", type=str2bool, default=False)
|
||||
parser.add_argument("--cls_model_path", type=str)
|
||||
parser.add_argument("--cls_image_shape", type=str, default="3, 48, 192")
|
||||
parser.add_argument("--label_list", type=list, default=['0', '180'])
|
||||
parser.add_argument("--cls_batch_num", type=int, default=6)
|
||||
parser.add_argument("--cls_thresh", type=float, default=0.9)
|
||||
|
||||
parser.add_argument("--enable_mkldnn", type=str2bool, default=False)
|
||||
parser.add_argument("--use_pdserving", type=str2bool, default=False)
|
||||
|
||||
# params for e2e
|
||||
parser.add_argument("--e2e_algorithm", type=str, default='PGNet')
|
||||
parser.add_argument("--e2e_model_path", type=str)
|
||||
parser.add_argument("--e2e_limit_side_len", type=float, default=768)
|
||||
parser.add_argument("--e2e_limit_type", type=str, default='max')
|
||||
|
||||
# PGNet parmas
|
||||
parser.add_argument("--e2e_pgnet_score_thresh", type=float, default=0.5)
|
||||
parser.add_argument(
|
||||
"--e2e_char_dict_path", type=str,
|
||||
default=os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))),
|
||||
'pytorchocr/utils/ic15_dict.txt'))
|
||||
parser.add_argument("--e2e_pgnet_valid_set", type=str, default='totaltext')
|
||||
parser.add_argument("--e2e_pgnet_polygon", type=bool, default=True)
|
||||
parser.add_argument("--e2e_pgnet_mode", type=str, default='fast')
|
||||
|
||||
# SR parmas
|
||||
parser.add_argument("--sr_model_path", type=str)
|
||||
parser.add_argument("--sr_image_shape", type=str, default="3, 32, 128")
|
||||
parser.add_argument("--sr_batch_num", type=int, default=1)
|
||||
|
||||
# params .yaml
|
||||
parser.add_argument("--det_yaml_path", type=str, default=None)
|
||||
parser.add_argument("--rec_yaml_path", type=str, default=None)
|
||||
parser.add_argument("--cls_yaml_path", type=str, default=None)
|
||||
parser.add_argument("--e2e_yaml_path", type=str, default=None)
|
||||
parser.add_argument("--sr_yaml_path", type=str, default=None)
|
||||
|
||||
# multi-process
|
||||
parser.add_argument("--use_mp", type=str2bool, default=False)
|
||||
parser.add_argument("--total_process_num", type=int, default=1)
|
||||
parser.add_argument("--process_id", type=int, default=0)
|
||||
|
||||
parser.add_argument("--benchmark", type=str2bool, default=False)
|
||||
parser.add_argument("--save_log_path", type=str, default="./log_output/")
|
||||
|
||||
parser.add_argument("--show_log", type=str2bool, default=True)
|
||||
|
||||
return parser
|
||||
|
||||
def parse_args():
|
||||
parser = init_args()
|
||||
return parser.parse_args()
|
||||
|
||||
def get_default_config(args):
|
||||
return vars(args)
|
||||
|
||||
|
||||
def read_network_config_from_yaml(yaml_path, char_num=None):
|
||||
if not os.path.exists(yaml_path):
|
||||
raise FileNotFoundError('{} is not existed.'.format(yaml_path))
|
||||
import yaml
|
||||
with open(yaml_path, encoding='utf-8') as f:
|
||||
res = yaml.safe_load(f)
|
||||
if res.get('Architecture') is None:
|
||||
raise ValueError('{} has no Architecture'.format(yaml_path))
|
||||
if res['Architecture']['Head']['name'] == 'MultiHead' and char_num is not None:
|
||||
res['Architecture']['Head']['out_channels_list'] = {
|
||||
'CTCLabelDecode': char_num,
|
||||
'SARLabelDecode': char_num + 2,
|
||||
'NRTRLabelDecode': char_num + 3
|
||||
}
|
||||
return res['Architecture']
|
||||
|
||||
def AnalysisConfig(weights_path, yaml_path=None, char_num=None):
|
||||
if not os.path.exists(os.path.abspath(weights_path)):
|
||||
raise FileNotFoundError('{} is not found.'.format(weights_path))
|
||||
|
||||
if yaml_path is not None:
|
||||
return read_network_config_from_yaml(yaml_path, char_num=char_num)
|
||||
|
||||
|
||||
def resize_img(img, input_size=600):
|
||||
"""
|
||||
resize img and limit the longest side of the image to input_size
|
||||
"""
|
||||
img = np.array(img)
|
||||
im_shape = img.shape
|
||||
im_size_max = np.max(im_shape[0:2])
|
||||
im_scale = float(input_size) / float(im_size_max)
|
||||
img = cv2.resize(img, None, None, fx=im_scale, fy=im_scale)
|
||||
return img
|
||||
|
||||
|
||||
def str_count(s):
|
||||
"""
|
||||
Count the number of Chinese characters,
|
||||
a single English character and a single number
|
||||
equal to half the length of Chinese characters.
|
||||
args:
|
||||
s(string): the input of string
|
||||
return(int):
|
||||
the number of Chinese characters
|
||||
"""
|
||||
import string
|
||||
count_zh = count_pu = 0
|
||||
s_len = len(s)
|
||||
en_dg_count = 0
|
||||
for c in s:
|
||||
if c in string.ascii_letters or c.isdigit() or c.isspace():
|
||||
en_dg_count += 1
|
||||
elif c.isalpha():
|
||||
count_zh += 1
|
||||
else:
|
||||
count_pu += 1
|
||||
return s_len - math.ceil(en_dg_count / 2)
|
||||
|
||||
|
||||
def base64_to_cv2(b64str):
|
||||
import base64
|
||||
data = base64.b64decode(b64str.encode('utf8'))
|
||||
data = np.fromstring(data, np.uint8)
|
||||
data = cv2.imdecode(data, cv2.IMREAD_COLOR)
|
||||
return data
|
||||
|
||||
|
||||
def get_arch_config(model_path):
|
||||
from omegaconf import OmegaConf
|
||||
all_arch_config = OmegaConf.load(DEFAULT_CFG_PATH)
|
||||
path = Path(model_path)
|
||||
file_name = path.stem
|
||||
if file_name not in all_arch_config:
|
||||
raise ValueError(f"architecture {file_name} is not in arch_config.yaml")
|
||||
|
||||
arch_config = all_arch_config[file_name]
|
||||
return arch_config
|
||||
@@ -0,0 +1 @@
|
||||
# Copyright (c) Opendatalab. All rights reserved.
|
||||
@@ -0,0 +1,125 @@
|
||||
from collections import defaultdict
|
||||
from typing import List, Dict
|
||||
|
||||
import torch
|
||||
from transformers import LayoutLMv3ForTokenClassification
|
||||
|
||||
MAX_LEN = 510
|
||||
CLS_TOKEN_ID = 0
|
||||
UNK_TOKEN_ID = 3
|
||||
EOS_TOKEN_ID = 2
|
||||
|
||||
|
||||
class DataCollator:
|
||||
def __call__(self, features: List[dict]) -> Dict[str, torch.Tensor]:
|
||||
bbox = []
|
||||
labels = []
|
||||
input_ids = []
|
||||
attention_mask = []
|
||||
|
||||
# clip bbox and labels to max length, build input_ids and attention_mask
|
||||
for feature in features:
|
||||
_bbox = feature["source_boxes"]
|
||||
if len(_bbox) > MAX_LEN:
|
||||
_bbox = _bbox[:MAX_LEN]
|
||||
_labels = feature["target_index"]
|
||||
if len(_labels) > MAX_LEN:
|
||||
_labels = _labels[:MAX_LEN]
|
||||
_input_ids = [UNK_TOKEN_ID] * len(_bbox)
|
||||
_attention_mask = [1] * len(_bbox)
|
||||
assert len(_bbox) == len(_labels) == len(_input_ids) == len(_attention_mask)
|
||||
bbox.append(_bbox)
|
||||
labels.append(_labels)
|
||||
input_ids.append(_input_ids)
|
||||
attention_mask.append(_attention_mask)
|
||||
|
||||
# add CLS and EOS tokens
|
||||
for i in range(len(bbox)):
|
||||
bbox[i] = [[0, 0, 0, 0]] + bbox[i] + [[0, 0, 0, 0]]
|
||||
labels[i] = [-100] + labels[i] + [-100]
|
||||
input_ids[i] = [CLS_TOKEN_ID] + input_ids[i] + [EOS_TOKEN_ID]
|
||||
attention_mask[i] = [1] + attention_mask[i] + [1]
|
||||
|
||||
# padding to max length
|
||||
max_len = max(len(x) for x in bbox)
|
||||
for i in range(len(bbox)):
|
||||
bbox[i] = bbox[i] + [[0, 0, 0, 0]] * (max_len - len(bbox[i]))
|
||||
labels[i] = labels[i] + [-100] * (max_len - len(labels[i]))
|
||||
input_ids[i] = input_ids[i] + [EOS_TOKEN_ID] * (max_len - len(input_ids[i]))
|
||||
attention_mask[i] = attention_mask[i] + [0] * (
|
||||
max_len - len(attention_mask[i])
|
||||
)
|
||||
|
||||
ret = {
|
||||
"bbox": torch.tensor(bbox),
|
||||
"attention_mask": torch.tensor(attention_mask),
|
||||
"labels": torch.tensor(labels),
|
||||
"input_ids": torch.tensor(input_ids),
|
||||
}
|
||||
# set label > MAX_LEN to -100, because original labels may be > MAX_LEN
|
||||
ret["labels"][ret["labels"] > MAX_LEN] = -100
|
||||
# set label > 0 to label-1, because original labels are 1-indexed
|
||||
ret["labels"][ret["labels"] > 0] -= 1
|
||||
return ret
|
||||
|
||||
|
||||
def boxes2inputs(boxes: List[List[int]]) -> Dict[str, torch.Tensor]:
|
||||
bbox = [[0, 0, 0, 0]] + boxes + [[0, 0, 0, 0]]
|
||||
input_ids = [CLS_TOKEN_ID] + [UNK_TOKEN_ID] * len(boxes) + [EOS_TOKEN_ID]
|
||||
attention_mask = [1] + [1] * len(boxes) + [1]
|
||||
return {
|
||||
"bbox": torch.tensor([bbox]),
|
||||
"attention_mask": torch.tensor([attention_mask]),
|
||||
"input_ids": torch.tensor([input_ids]),
|
||||
}
|
||||
|
||||
|
||||
def prepare_inputs(
|
||||
inputs: Dict[str, torch.Tensor], model: LayoutLMv3ForTokenClassification
|
||||
) -> Dict[str, torch.Tensor]:
|
||||
ret = {}
|
||||
for k, v in inputs.items():
|
||||
v = v.to(model.device)
|
||||
if torch.is_floating_point(v):
|
||||
v = v.to(model.dtype)
|
||||
ret[k] = v
|
||||
return ret
|
||||
|
||||
|
||||
def parse_logits(logits: torch.Tensor, length: int) -> List[int]:
|
||||
"""
|
||||
parse logits to orders
|
||||
|
||||
:param logits: logits from model
|
||||
:param length: input length
|
||||
:return: orders
|
||||
"""
|
||||
logits = logits[1 : length + 1, :length]
|
||||
orders = logits.argsort(descending=False).tolist()
|
||||
ret = [o.pop() for o in orders]
|
||||
while True:
|
||||
order_to_idxes = defaultdict(list)
|
||||
for idx, order in enumerate(ret):
|
||||
order_to_idxes[order].append(idx)
|
||||
# filter idxes len > 1
|
||||
order_to_idxes = {k: v for k, v in order_to_idxes.items() if len(v) > 1}
|
||||
if not order_to_idxes:
|
||||
break
|
||||
# filter
|
||||
for order, idxes in order_to_idxes.items():
|
||||
# find original logits of idxes
|
||||
idxes_to_logit = {}
|
||||
for idx in idxes:
|
||||
idxes_to_logit[idx] = logits[idx, order]
|
||||
idxes_to_logit = sorted(
|
||||
idxes_to_logit.items(), key=lambda x: x[1], reverse=True
|
||||
)
|
||||
# keep the highest logit as order, set others to next candidate
|
||||
for idx, _ in idxes_to_logit[1:]:
|
||||
ret[idx] = orders[idx].pop()
|
||||
|
||||
return ret
|
||||
|
||||
|
||||
def check_duplicate(a: List[int]) -> bool:
|
||||
return len(a) != len(set(a))
|
||||
@@ -0,0 +1,242 @@
|
||||
from typing import List
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
|
||||
def projection_by_bboxes(boxes: np.array, axis: int) -> np.ndarray:
|
||||
"""
|
||||
通过一组 bbox 获得投影直方图,最后以 per-pixel 形式输出
|
||||
|
||||
Args:
|
||||
boxes: [N, 4]
|
||||
axis: 0-x坐标向水平方向投影, 1-y坐标向垂直方向投影
|
||||
|
||||
Returns:
|
||||
1D 投影直方图,长度为投影方向坐标的最大值(我们不需要图片的实际边长,因为只是要找文本框的间隔)
|
||||
|
||||
"""
|
||||
assert axis in [0, 1]
|
||||
length = np.max(boxes[:, axis::2])
|
||||
res = np.zeros(length, dtype=int)
|
||||
# TODO: how to remove for loop?
|
||||
for start, end in boxes[:, axis::2]:
|
||||
res[start:end] += 1
|
||||
return res
|
||||
|
||||
|
||||
# from: https://dothinking.github.io/2021-06-19-%E9%80%92%E5%BD%92%E6%8A%95%E5%BD%B1%E5%88%86%E5%89%B2%E7%AE%97%E6%B3%95/#:~:text=%E9%80%92%E5%BD%92%E6%8A%95%E5%BD%B1%E5%88%86%E5%89%B2%EF%BC%88Recursive%20XY,%EF%BC%8C%E5%8F%AF%E4%BB%A5%E5%88%92%E5%88%86%E6%AE%B5%E8%90%BD%E3%80%81%E8%A1%8C%E3%80%82
|
||||
def split_projection_profile(arr_values: np.array, min_value: float, min_gap: float):
|
||||
"""Split projection profile:
|
||||
|
||||
```
|
||||
┌──┐
|
||||
arr_values │ │ ┌─┐───
|
||||
┌──┐ │ │ │ │ |
|
||||
│ │ │ │ ┌───┐ │ │min_value
|
||||
│ │<- min_gap ->│ │ │ │ │ │ |
|
||||
────┴──┴─────────────┴──┴─┴───┴─┴─┴─┴───
|
||||
0 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16
|
||||
```
|
||||
|
||||
Args:
|
||||
arr_values (np.array): 1-d array representing the projection profile.
|
||||
min_value (float): Ignore the profile if `arr_value` is less than `min_value`.
|
||||
min_gap (float): Ignore the gap if less than this value.
|
||||
|
||||
Returns:
|
||||
tuple: Start indexes and end indexes of split groups.
|
||||
"""
|
||||
# all indexes with projection height exceeding the threshold
|
||||
arr_index = np.where(arr_values > min_value)[0]
|
||||
if not len(arr_index):
|
||||
return
|
||||
|
||||
# find zero intervals between adjacent projections
|
||||
# | | ||
|
||||
# ||||<- zero-interval -> |||||
|
||||
arr_diff = arr_index[1:] - arr_index[0:-1]
|
||||
arr_diff_index = np.where(arr_diff > min_gap)[0]
|
||||
arr_zero_intvl_start = arr_index[arr_diff_index]
|
||||
arr_zero_intvl_end = arr_index[arr_diff_index + 1]
|
||||
|
||||
# convert to index of projection range:
|
||||
# the start index of zero interval is the end index of projection
|
||||
arr_start = np.insert(arr_zero_intvl_end, 0, arr_index[0])
|
||||
arr_end = np.append(arr_zero_intvl_start, arr_index[-1])
|
||||
arr_end += 1 # end index will be excluded as index slice
|
||||
|
||||
return arr_start, arr_end
|
||||
|
||||
|
||||
def recursive_xy_cut(boxes: np.ndarray, indices: List[int], res: List[int]):
|
||||
"""
|
||||
|
||||
Args:
|
||||
boxes: (N, 4)
|
||||
indices: 递归过程中始终表示 box 在原始数据中的索引
|
||||
res: 保存输出结果
|
||||
|
||||
"""
|
||||
# 向 y 轴投影
|
||||
assert len(boxes) == len(indices)
|
||||
|
||||
_indices = boxes[:, 1].argsort()
|
||||
y_sorted_boxes = boxes[_indices]
|
||||
y_sorted_indices = indices[_indices]
|
||||
|
||||
# debug_vis(y_sorted_boxes, y_sorted_indices)
|
||||
|
||||
y_projection = projection_by_bboxes(boxes=y_sorted_boxes, axis=1)
|
||||
pos_y = split_projection_profile(y_projection, 0, 1)
|
||||
if not pos_y:
|
||||
return
|
||||
|
||||
arr_y0, arr_y1 = pos_y
|
||||
for r0, r1 in zip(arr_y0, arr_y1):
|
||||
# [r0, r1] 表示按照水平切分,有 bbox 的区域,对这些区域会再进行垂直切分
|
||||
_indices = (r0 <= y_sorted_boxes[:, 1]) & (y_sorted_boxes[:, 1] < r1)
|
||||
|
||||
y_sorted_boxes_chunk = y_sorted_boxes[_indices]
|
||||
y_sorted_indices_chunk = y_sorted_indices[_indices]
|
||||
|
||||
_indices = y_sorted_boxes_chunk[:, 0].argsort()
|
||||
x_sorted_boxes_chunk = y_sorted_boxes_chunk[_indices]
|
||||
x_sorted_indices_chunk = y_sorted_indices_chunk[_indices]
|
||||
|
||||
# 往 x 方向投影
|
||||
x_projection = projection_by_bboxes(boxes=x_sorted_boxes_chunk, axis=0)
|
||||
pos_x = split_projection_profile(x_projection, 0, 1)
|
||||
if not pos_x:
|
||||
continue
|
||||
|
||||
arr_x0, arr_x1 = pos_x
|
||||
if len(arr_x0) == 1:
|
||||
# x 方向无法切分
|
||||
res.extend(x_sorted_indices_chunk)
|
||||
continue
|
||||
|
||||
# x 方向上能分开,继续递归调用
|
||||
for c0, c1 in zip(arr_x0, arr_x1):
|
||||
_indices = (c0 <= x_sorted_boxes_chunk[:, 0]) & (
|
||||
x_sorted_boxes_chunk[:, 0] < c1
|
||||
)
|
||||
recursive_xy_cut(
|
||||
x_sorted_boxes_chunk[_indices], x_sorted_indices_chunk[_indices], res
|
||||
)
|
||||
|
||||
|
||||
def points_to_bbox(points):
|
||||
assert len(points) == 8
|
||||
|
||||
# [x1,y1,x2,y2,x3,y3,x4,y4]
|
||||
left = min(points[::2])
|
||||
right = max(points[::2])
|
||||
top = min(points[1::2])
|
||||
bottom = max(points[1::2])
|
||||
|
||||
left = max(left, 0)
|
||||
top = max(top, 0)
|
||||
right = max(right, 0)
|
||||
bottom = max(bottom, 0)
|
||||
return [left, top, right, bottom]
|
||||
|
||||
|
||||
def bbox2points(bbox):
|
||||
left, top, right, bottom = bbox
|
||||
return [left, top, right, top, right, bottom, left, bottom]
|
||||
|
||||
|
||||
def vis_polygon(img, points, thickness=2, color=None):
|
||||
br2bl_color = color
|
||||
tl2tr_color = color
|
||||
tr2br_color = color
|
||||
bl2tl_color = color
|
||||
cv2.line(
|
||||
img,
|
||||
(points[0][0], points[0][1]),
|
||||
(points[1][0], points[1][1]),
|
||||
color=tl2tr_color,
|
||||
thickness=thickness,
|
||||
)
|
||||
|
||||
cv2.line(
|
||||
img,
|
||||
(points[1][0], points[1][1]),
|
||||
(points[2][0], points[2][1]),
|
||||
color=tr2br_color,
|
||||
thickness=thickness,
|
||||
)
|
||||
|
||||
cv2.line(
|
||||
img,
|
||||
(points[2][0], points[2][1]),
|
||||
(points[3][0], points[3][1]),
|
||||
color=br2bl_color,
|
||||
thickness=thickness,
|
||||
)
|
||||
|
||||
cv2.line(
|
||||
img,
|
||||
(points[3][0], points[3][1]),
|
||||
(points[0][0], points[0][1]),
|
||||
color=bl2tl_color,
|
||||
thickness=thickness,
|
||||
)
|
||||
return img
|
||||
|
||||
|
||||
def vis_points(
|
||||
img: np.ndarray, points, texts: List[str] = None, color=(0, 200, 0)
|
||||
) -> np.ndarray:
|
||||
"""
|
||||
|
||||
Args:
|
||||
img:
|
||||
points: [N, 8] 8: x1,y1,x2,y2,x3,y3,x3,y4
|
||||
texts:
|
||||
color:
|
||||
|
||||
Returns:
|
||||
|
||||
"""
|
||||
points = np.array(points)
|
||||
if texts is not None:
|
||||
assert len(texts) == points.shape[0]
|
||||
|
||||
for i, _points in enumerate(points):
|
||||
vis_polygon(img, _points.reshape(-1, 2), thickness=2, color=color)
|
||||
bbox = points_to_bbox(_points)
|
||||
left, top, right, bottom = bbox
|
||||
cx = (left + right) // 2
|
||||
cy = (top + bottom) // 2
|
||||
|
||||
txt = texts[i]
|
||||
font = cv2.FONT_HERSHEY_SIMPLEX
|
||||
cat_size = cv2.getTextSize(txt, font, 0.5, 2)[0]
|
||||
|
||||
img = cv2.rectangle(
|
||||
img,
|
||||
(cx - 5 * len(txt), cy - cat_size[1] - 5),
|
||||
(cx - 5 * len(txt) + cat_size[0], cy - 5),
|
||||
color,
|
||||
-1,
|
||||
)
|
||||
|
||||
img = cv2.putText(
|
||||
img,
|
||||
txt,
|
||||
(cx - 5 * len(txt), cy - 5),
|
||||
font,
|
||||
0.5,
|
||||
(255, 255, 255),
|
||||
thickness=1,
|
||||
lineType=cv2.LINE_AA,
|
||||
)
|
||||
|
||||
return img
|
||||
|
||||
|
||||
def vis_polygons_with_index(image, points):
|
||||
texts = [str(i) for i in range(len(points))]
|
||||
res_img = vis_points(image.copy(), points, texts)
|
||||
return res_img
|
||||
@@ -0,0 +1 @@
|
||||
# Copyright (c) Opendatalab. All rights reserved.
|
||||
@@ -0,0 +1,79 @@
|
||||
import os
|
||||
from pathlib import Path
|
||||
import cv2
|
||||
import numpy as np
|
||||
from loguru import logger
|
||||
from rapid_table import RapidTable, RapidTableInput
|
||||
|
||||
|
||||
class RapidTableModel(object):
|
||||
def __init__(self, ocr_engine):
|
||||
root_dir = Path(__file__).absolute().parent.parent.parent.parent.parent
|
||||
slanet_plus_model_path = os.path.join(root_dir, 'resources', 'slanet_plus', 'slanet-plus.onnx')
|
||||
input_args = RapidTableInput(model_type='slanet_plus', model_path=slanet_plus_model_path)
|
||||
self.table_model = RapidTable(input_args)
|
||||
self.ocr_engine = ocr_engine
|
||||
|
||||
|
||||
def predict(self, image):
|
||||
bgr_image = cv2.cvtColor(np.asarray(image), cv2.COLOR_RGB2BGR)
|
||||
|
||||
# First check the overall image aspect ratio (height/width)
|
||||
img_height, img_width = bgr_image.shape[:2]
|
||||
img_aspect_ratio = img_height / img_width if img_width > 0 else 1.0
|
||||
img_is_portrait = img_aspect_ratio > 1.2
|
||||
|
||||
if img_is_portrait:
|
||||
|
||||
det_res = self.ocr_engine.ocr(bgr_image, rec=False)[0]
|
||||
# Check if table is rotated by analyzing text box aspect ratios
|
||||
is_rotated = False
|
||||
if det_res:
|
||||
vertical_count = 0
|
||||
|
||||
for box_ocr_res in det_res:
|
||||
p1, p2, p3, p4 = box_ocr_res
|
||||
|
||||
# Calculate width and height
|
||||
width = p3[0] - p1[0]
|
||||
height = p3[1] - p1[1]
|
||||
|
||||
aspect_ratio = width / height if height > 0 else 1.0
|
||||
|
||||
# Count vertical vs horizontal text boxes
|
||||
if aspect_ratio < 0.8: # Taller than wide - vertical text
|
||||
vertical_count += 1
|
||||
# elif aspect_ratio > 1.2: # Wider than tall - horizontal text
|
||||
# horizontal_count += 1
|
||||
|
||||
# If we have more vertical text boxes than horizontal ones,
|
||||
# and vertical ones are significant, table might be rotated
|
||||
if vertical_count >= len(det_res) * 0.3:
|
||||
is_rotated = True
|
||||
|
||||
# logger.debug(f"Text orientation analysis: vertical={vertical_count}, det_res={len(det_res)}, rotated={is_rotated}")
|
||||
|
||||
# Rotate image if necessary
|
||||
if is_rotated:
|
||||
# logger.debug("Table appears to be in portrait orientation, rotating 90 degrees clockwise")
|
||||
image = cv2.rotate(np.asarray(image), cv2.ROTATE_90_CLOCKWISE)
|
||||
bgr_image = cv2.cvtColor(image, cv2.COLOR_RGB2BGR)
|
||||
|
||||
# Continue with OCR on potentially rotated image
|
||||
ocr_result = self.ocr_engine.ocr(bgr_image)[0]
|
||||
if ocr_result:
|
||||
ocr_result = [[item[0], item[1][0], item[1][1]] for item in ocr_result if
|
||||
len(item) == 2 and isinstance(item[1], tuple)]
|
||||
else:
|
||||
ocr_result = None
|
||||
|
||||
|
||||
if ocr_result:
|
||||
table_results = self.table_model(np.asarray(image), ocr_result)
|
||||
html_code = table_results.pred_html
|
||||
table_cell_bboxes = table_results.cell_bboxes
|
||||
logic_points = table_results.logic_points
|
||||
elapse = table_results.elapse
|
||||
return html_code, table_cell_bboxes, logic_points, elapse
|
||||
else:
|
||||
return None, None, None, None
|
||||
@@ -0,0 +1,323 @@
|
||||
import time
|
||||
import torch
|
||||
import gc
|
||||
from loguru import logger
|
||||
import numpy as np
|
||||
|
||||
from magic_pdf.libs.boxbase import get_minbox_if_overlap_by_ratio
|
||||
|
||||
|
||||
def crop_img(input_res, input_np_img, crop_paste_x=0, crop_paste_y=0):
|
||||
|
||||
crop_xmin, crop_ymin = int(input_res['poly'][0]), int(input_res['poly'][1])
|
||||
crop_xmax, crop_ymax = int(input_res['poly'][4]), int(input_res['poly'][5])
|
||||
|
||||
# Calculate new dimensions
|
||||
crop_new_width = crop_xmax - crop_xmin + crop_paste_x * 2
|
||||
crop_new_height = crop_ymax - crop_ymin + crop_paste_y * 2
|
||||
|
||||
# Create a white background array
|
||||
return_image = np.ones((crop_new_height, crop_new_width, 3), dtype=np.uint8) * 255
|
||||
|
||||
# Crop the original image using numpy slicing
|
||||
cropped_img = input_np_img[crop_ymin:crop_ymax, crop_xmin:crop_xmax]
|
||||
|
||||
# Paste the cropped image onto the white background
|
||||
return_image[crop_paste_y:crop_paste_y + (crop_ymax - crop_ymin),
|
||||
crop_paste_x:crop_paste_x + (crop_xmax - crop_xmin)] = cropped_img
|
||||
|
||||
return_list = [crop_paste_x, crop_paste_y, crop_xmin, crop_ymin, crop_xmax, crop_ymax, crop_new_width,
|
||||
crop_new_height]
|
||||
return return_image, return_list
|
||||
|
||||
|
||||
def get_coords_and_area(block_with_poly):
|
||||
"""Extract coordinates and area from a table."""
|
||||
xmin, ymin = int(block_with_poly['poly'][0]), int(block_with_poly['poly'][1])
|
||||
xmax, ymax = int(block_with_poly['poly'][4]), int(block_with_poly['poly'][5])
|
||||
area = (xmax - xmin) * (ymax - ymin)
|
||||
return xmin, ymin, xmax, ymax, area
|
||||
|
||||
|
||||
def calculate_intersection(box1, box2):
|
||||
"""Calculate intersection coordinates between two boxes."""
|
||||
intersection_xmin = max(box1[0], box2[0])
|
||||
intersection_ymin = max(box1[1], box2[1])
|
||||
intersection_xmax = min(box1[2], box2[2])
|
||||
intersection_ymax = min(box1[3], box2[3])
|
||||
|
||||
# Check if intersection is valid
|
||||
if intersection_xmax <= intersection_xmin or intersection_ymax <= intersection_ymin:
|
||||
return None
|
||||
|
||||
return intersection_xmin, intersection_ymin, intersection_xmax, intersection_ymax
|
||||
|
||||
|
||||
def calculate_iou(box1, box2):
|
||||
"""Calculate IoU between two boxes."""
|
||||
intersection = calculate_intersection(box1[:4], box2[:4])
|
||||
|
||||
if not intersection:
|
||||
return 0
|
||||
|
||||
intersection_xmin, intersection_ymin, intersection_xmax, intersection_ymax = intersection
|
||||
intersection_area = (intersection_xmax - intersection_xmin) * (intersection_ymax - intersection_ymin)
|
||||
|
||||
area1, area2 = box1[4], box2[4]
|
||||
union_area = area1 + area2 - intersection_area
|
||||
|
||||
return intersection_area / union_area if union_area > 0 else 0
|
||||
|
||||
|
||||
def is_inside(small_box, big_box, overlap_threshold=0.8):
|
||||
"""Check if small_box is inside big_box by at least overlap_threshold."""
|
||||
intersection = calculate_intersection(small_box[:4], big_box[:4])
|
||||
|
||||
if not intersection:
|
||||
return False
|
||||
|
||||
intersection_xmin, intersection_ymin, intersection_xmax, intersection_ymax = intersection
|
||||
intersection_area = (intersection_xmax - intersection_xmin) * (intersection_ymax - intersection_ymin)
|
||||
|
||||
# Check if overlap exceeds threshold
|
||||
return intersection_area >= overlap_threshold * small_box[4]
|
||||
|
||||
|
||||
def do_overlap(box1, box2):
|
||||
"""Check if two boxes overlap."""
|
||||
return calculate_intersection(box1[:4], box2[:4]) is not None
|
||||
|
||||
|
||||
def merge_high_iou_tables(table_res_list, layout_res, table_indices, iou_threshold=0.7):
|
||||
"""Merge tables with IoU > threshold."""
|
||||
if len(table_res_list) < 2:
|
||||
return table_res_list, table_indices
|
||||
|
||||
table_info = [get_coords_and_area(table) for table in table_res_list]
|
||||
merged = True
|
||||
|
||||
while merged:
|
||||
merged = False
|
||||
i = 0
|
||||
while i < len(table_res_list) - 1:
|
||||
j = i + 1
|
||||
while j < len(table_res_list):
|
||||
iou = calculate_iou(table_info[i], table_info[j])
|
||||
|
||||
if iou > iou_threshold:
|
||||
# Merge tables by taking their union
|
||||
x1_min, y1_min, x1_max, y1_max, _ = table_info[i]
|
||||
x2_min, y2_min, x2_max, y2_max, _ = table_info[j]
|
||||
|
||||
union_xmin = min(x1_min, x2_min)
|
||||
union_ymin = min(y1_min, y2_min)
|
||||
union_xmax = max(x1_max, x2_max)
|
||||
union_ymax = max(y1_max, y2_max)
|
||||
|
||||
# Create merged table
|
||||
merged_table = table_res_list[i].copy()
|
||||
merged_table['poly'][0] = union_xmin
|
||||
merged_table['poly'][1] = union_ymin
|
||||
merged_table['poly'][2] = union_xmax
|
||||
merged_table['poly'][3] = union_ymin
|
||||
merged_table['poly'][4] = union_xmax
|
||||
merged_table['poly'][5] = union_ymax
|
||||
merged_table['poly'][6] = union_xmin
|
||||
merged_table['poly'][7] = union_ymax
|
||||
|
||||
# Update layout_res
|
||||
to_remove = [table_indices[j], table_indices[i]]
|
||||
for idx in sorted(to_remove, reverse=True):
|
||||
del layout_res[idx]
|
||||
layout_res.append(merged_table)
|
||||
|
||||
# Update tracking lists
|
||||
table_indices = [k if k < min(to_remove) else
|
||||
k - 1 if k < max(to_remove) else
|
||||
k - 2 if k > max(to_remove) else
|
||||
len(layout_res) - 1
|
||||
for k in table_indices
|
||||
if k not in to_remove]
|
||||
table_indices.append(len(layout_res) - 1)
|
||||
|
||||
# Update table lists
|
||||
table_res_list.pop(j)
|
||||
table_res_list.pop(i)
|
||||
table_res_list.append(merged_table)
|
||||
|
||||
# Update table_info
|
||||
table_info = [get_coords_and_area(table) for table in table_res_list]
|
||||
|
||||
merged = True
|
||||
break
|
||||
j += 1
|
||||
|
||||
if merged:
|
||||
break
|
||||
i += 1
|
||||
|
||||
return table_res_list, table_indices
|
||||
|
||||
|
||||
def filter_nested_tables(table_res_list, overlap_threshold=0.8, area_threshold=0.8):
|
||||
"""Remove big tables containing multiple smaller tables within them."""
|
||||
if len(table_res_list) < 3:
|
||||
return table_res_list
|
||||
|
||||
table_info = [get_coords_and_area(table) for table in table_res_list]
|
||||
big_tables_idx = []
|
||||
|
||||
for i in range(len(table_res_list)):
|
||||
# Find tables inside this one
|
||||
tables_inside = [j for j in range(len(table_res_list))
|
||||
if i != j and is_inside(table_info[j], table_info[i], overlap_threshold)]
|
||||
|
||||
# Continue if there are at least 3 tables inside
|
||||
if len(tables_inside) >= 3:
|
||||
# Check if inside tables overlap with each other
|
||||
tables_overlap = any(do_overlap(table_info[tables_inside[idx1]], table_info[tables_inside[idx2]])
|
||||
for idx1 in range(len(tables_inside))
|
||||
for idx2 in range(idx1 + 1, len(tables_inside)))
|
||||
|
||||
# If no overlaps, check area condition
|
||||
if not tables_overlap:
|
||||
total_inside_area = sum(table_info[j][4] for j in tables_inside)
|
||||
big_table_area = table_info[i][4]
|
||||
|
||||
if total_inside_area > area_threshold * big_table_area:
|
||||
big_tables_idx.append(i)
|
||||
|
||||
return [table for i, table in enumerate(table_res_list) if i not in big_tables_idx]
|
||||
|
||||
|
||||
def remove_overlaps_min_blocks(res_list):
|
||||
# 重叠block,小的不能直接删除,需要和大的那个合并成一个更大的。
|
||||
# 删除重叠blocks中较小的那些
|
||||
need_remove = []
|
||||
for res1 in res_list:
|
||||
for res2 in res_list:
|
||||
if res1 != res2:
|
||||
overlap_box = get_minbox_if_overlap_by_ratio(
|
||||
res1['bbox'], res2['bbox'], 0.8
|
||||
)
|
||||
if overlap_box is not None:
|
||||
res_to_remove = next(
|
||||
(res for res in res_list if res['bbox'] == overlap_box),
|
||||
None,
|
||||
)
|
||||
if (
|
||||
res_to_remove is not None
|
||||
and res_to_remove not in need_remove
|
||||
):
|
||||
large_res = res1 if res1 != res_to_remove else res2
|
||||
x1, y1, x2, y2 = large_res['bbox']
|
||||
sx1, sy1, sx2, sy2 = res_to_remove['bbox']
|
||||
x1 = min(x1, sx1)
|
||||
y1 = min(y1, sy1)
|
||||
x2 = max(x2, sx2)
|
||||
y2 = max(y2, sy2)
|
||||
large_res['bbox'] = [x1, y1, x2, y2]
|
||||
need_remove.append(res_to_remove)
|
||||
|
||||
if len(need_remove) > 0:
|
||||
for res in need_remove:
|
||||
res_list.remove(res)
|
||||
|
||||
return res_list, need_remove
|
||||
|
||||
|
||||
def get_res_list_from_layout_res(layout_res, iou_threshold=0.7, overlap_threshold=0.8, area_threshold=0.8):
|
||||
"""Extract OCR, table and other regions from layout results."""
|
||||
ocr_res_list = []
|
||||
text_res_list = []
|
||||
table_res_list = []
|
||||
table_indices = []
|
||||
single_page_mfdetrec_res = []
|
||||
|
||||
# Categorize regions
|
||||
for i, res in enumerate(layout_res):
|
||||
category_id = int(res['category_id'])
|
||||
|
||||
if category_id in [13, 14]: # Formula regions
|
||||
single_page_mfdetrec_res.append({
|
||||
"bbox": [int(res['poly'][0]), int(res['poly'][1]),
|
||||
int(res['poly'][4]), int(res['poly'][5])],
|
||||
})
|
||||
elif category_id in [0, 2, 4, 6, 7, 3]: # OCR regions
|
||||
ocr_res_list.append(res)
|
||||
elif category_id == 5: # Table regions
|
||||
table_res_list.append(res)
|
||||
table_indices.append(i)
|
||||
elif category_id in [1]: # Text regions
|
||||
res['bbox'] = [int(res['poly'][0]), int(res['poly'][1]), int(res['poly'][4]), int(res['poly'][5])]
|
||||
text_res_list.append(res)
|
||||
|
||||
# Process tables: merge high IoU tables first, then filter nested tables
|
||||
table_res_list, table_indices = merge_high_iou_tables(
|
||||
table_res_list, layout_res, table_indices, iou_threshold)
|
||||
|
||||
filtered_table_res_list = filter_nested_tables(
|
||||
table_res_list, overlap_threshold, area_threshold)
|
||||
|
||||
# Remove filtered out tables from layout_res
|
||||
if len(filtered_table_res_list) < len(table_res_list):
|
||||
kept_tables = set(id(table) for table in filtered_table_res_list)
|
||||
to_remove = [table_indices[i] for i, table in enumerate(table_res_list)
|
||||
if id(table) not in kept_tables]
|
||||
|
||||
for idx in sorted(to_remove, reverse=True):
|
||||
del layout_res[idx]
|
||||
|
||||
# Remove overlaps in OCR and text regions
|
||||
text_res_list, need_remove = remove_overlaps_min_blocks(text_res_list)
|
||||
for res in text_res_list:
|
||||
# 将res的poly使用bbox重构
|
||||
res['poly'] = [res['bbox'][0], res['bbox'][1], res['bbox'][2], res['bbox'][1],
|
||||
res['bbox'][2], res['bbox'][3], res['bbox'][0], res['bbox'][3]]
|
||||
# 删除res的bbox
|
||||
del res['bbox']
|
||||
|
||||
ocr_res_list.extend(text_res_list)
|
||||
|
||||
if len(need_remove) > 0:
|
||||
for res in need_remove:
|
||||
del res['bbox']
|
||||
layout_res.remove(res)
|
||||
|
||||
return ocr_res_list, filtered_table_res_list, single_page_mfdetrec_res
|
||||
|
||||
|
||||
def clean_memory(device='cuda'):
|
||||
if device == 'cuda':
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
elif str(device).startswith("npu"):
|
||||
import torch_npu
|
||||
if torch_npu.npu.is_available():
|
||||
torch_npu.npu.empty_cache()
|
||||
elif str(device).startswith("mps"):
|
||||
torch.mps.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
def clean_vram(device, vram_threshold=8):
|
||||
total_memory = get_vram(device)
|
||||
if total_memory and total_memory <= vram_threshold:
|
||||
gc_start = time.time()
|
||||
clean_memory(device)
|
||||
gc_time = round(time.time() - gc_start, 2)
|
||||
logger.info(f"gc time: {gc_time}")
|
||||
|
||||
|
||||
def get_vram(device):
|
||||
if torch.cuda.is_available() and str(device).startswith("cuda"):
|
||||
total_memory = torch.cuda.get_device_properties(device).total_memory / (1024 ** 3) # 将字节转换为 GB
|
||||
return total_memory
|
||||
elif str(device).startswith("npu"):
|
||||
import torch_npu
|
||||
if torch_npu.npu.is_available():
|
||||
total_memory = torch_npu.npu.get_device_properties(device).total_memory / (1024 ** 3) # 转为 GB
|
||||
return total_memory
|
||||
else:
|
||||
return None
|
||||
@@ -0,0 +1,401 @@
|
||||
# Copyright (c) Opendatalab. All rights reserved.
|
||||
import copy
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
|
||||
def merge_spans_to_line(spans, threshold=0.6):
|
||||
if len(spans) == 0:
|
||||
return []
|
||||
else:
|
||||
# 按照y0坐标排序
|
||||
spans.sort(key=lambda span: span['bbox'][1])
|
||||
|
||||
lines = []
|
||||
current_line = [spans[0]]
|
||||
for span in spans[1:]:
|
||||
# 如果当前的span与当前行的最后一个span在y轴上重叠,则添加到当前行
|
||||
if __is_overlaps_y_exceeds_threshold(span['bbox'], current_line[-1]['bbox'], threshold):
|
||||
current_line.append(span)
|
||||
else:
|
||||
# 否则,开始新行
|
||||
lines.append(current_line)
|
||||
current_line = [span]
|
||||
|
||||
# 添加最后一行
|
||||
if current_line:
|
||||
lines.append(current_line)
|
||||
|
||||
return lines
|
||||
|
||||
def __is_overlaps_y_exceeds_threshold(bbox1,
|
||||
bbox2,
|
||||
overlap_ratio_threshold=0.8):
|
||||
"""检查两个bbox在y轴上是否有重叠,并且该重叠区域的高度占两个bbox高度更低的那个超过80%"""
|
||||
_, y0_1, _, y1_1 = bbox1
|
||||
_, y0_2, _, y1_2 = bbox2
|
||||
|
||||
overlap = max(0, min(y1_1, y1_2) - max(y0_1, y0_2))
|
||||
height1, height2 = y1_1 - y0_1, y1_2 - y0_2
|
||||
# max_height = max(height1, height2)
|
||||
min_height = min(height1, height2)
|
||||
|
||||
return (overlap / min_height) > overlap_ratio_threshold
|
||||
|
||||
|
||||
def img_decode(content: bytes):
|
||||
np_arr = np.frombuffer(content, dtype=np.uint8)
|
||||
return cv2.imdecode(np_arr, cv2.IMREAD_UNCHANGED)
|
||||
|
||||
def check_img(img):
|
||||
if isinstance(img, bytes):
|
||||
img = img_decode(img)
|
||||
if isinstance(img, np.ndarray) and len(img.shape) == 2:
|
||||
img = cv2.cvtColor(img, cv2.COLOR_GRAY2BGR)
|
||||
return img
|
||||
|
||||
|
||||
def alpha_to_color(img, alpha_color=(255, 255, 255)):
|
||||
if len(img.shape) == 3 and img.shape[2] == 4:
|
||||
B, G, R, A = cv2.split(img)
|
||||
alpha = A / 255
|
||||
|
||||
R = (alpha_color[0] * (1 - alpha) + R * alpha).astype(np.uint8)
|
||||
G = (alpha_color[1] * (1 - alpha) + G * alpha).astype(np.uint8)
|
||||
B = (alpha_color[2] * (1 - alpha) + B * alpha).astype(np.uint8)
|
||||
|
||||
img = cv2.merge((B, G, R))
|
||||
return img
|
||||
|
||||
|
||||
def preprocess_image(_image):
|
||||
alpha_color = (255, 255, 255)
|
||||
_image = alpha_to_color(_image, alpha_color)
|
||||
return _image
|
||||
|
||||
|
||||
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 bbox_to_points(bbox):
|
||||
""" 将bbox格式转换为四个顶点的数组 """
|
||||
x0, y0, x1, y1 = bbox
|
||||
return np.array([[x0, y0], [x1, y0], [x1, y1], [x0, y1]]).astype('float32')
|
||||
|
||||
|
||||
def points_to_bbox(points):
|
||||
""" 将四个顶点的数组转换为bbox格式 """
|
||||
x0, y0 = points[0]
|
||||
x1, _ = points[1]
|
||||
_, y1 = points[2]
|
||||
return [x0, y0, x1, y1]
|
||||
|
||||
|
||||
def merge_intervals(intervals):
|
||||
# Sort the intervals based on the start value
|
||||
intervals.sort(key=lambda x: x[0])
|
||||
|
||||
merged = []
|
||||
for interval in intervals:
|
||||
# If the list of merged intervals is empty or if the current
|
||||
# interval does not overlap with the previous, simply append it.
|
||||
if not merged or merged[-1][1] < interval[0]:
|
||||
merged.append(interval)
|
||||
else:
|
||||
# Otherwise, there is overlap, so we merge the current and previous intervals.
|
||||
merged[-1][1] = max(merged[-1][1], interval[1])
|
||||
|
||||
return merged
|
||||
|
||||
|
||||
def remove_intervals(original, masks):
|
||||
# Merge all mask intervals
|
||||
merged_masks = merge_intervals(masks)
|
||||
|
||||
result = []
|
||||
original_start, original_end = original
|
||||
|
||||
for mask in merged_masks:
|
||||
mask_start, mask_end = mask
|
||||
|
||||
# If the mask starts after the original range, ignore it
|
||||
if mask_start > original_end:
|
||||
continue
|
||||
|
||||
# If the mask ends before the original range starts, ignore it
|
||||
if mask_end < original_start:
|
||||
continue
|
||||
|
||||
# Remove the masked part from the original range
|
||||
if original_start < mask_start:
|
||||
result.append([original_start, mask_start - 1])
|
||||
|
||||
original_start = max(mask_end + 1, original_start)
|
||||
|
||||
# Add the remaining part of the original range, if any
|
||||
if original_start <= original_end:
|
||||
result.append([original_start, original_end])
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def update_det_boxes(dt_boxes, mfd_res):
|
||||
new_dt_boxes = []
|
||||
angle_boxes_list = []
|
||||
for text_box in dt_boxes:
|
||||
|
||||
if calculate_is_angle(text_box):
|
||||
angle_boxes_list.append(text_box)
|
||||
continue
|
||||
|
||||
text_bbox = points_to_bbox(text_box)
|
||||
masks_list = []
|
||||
for mf_box in mfd_res:
|
||||
mf_bbox = mf_box['bbox']
|
||||
if __is_overlaps_y_exceeds_threshold(text_bbox, mf_bbox):
|
||||
masks_list.append([mf_bbox[0], mf_bbox[2]])
|
||||
text_x_range = [text_bbox[0], text_bbox[2]]
|
||||
text_remove_mask_range = remove_intervals(text_x_range, masks_list)
|
||||
temp_dt_box = []
|
||||
for text_remove_mask in text_remove_mask_range:
|
||||
temp_dt_box.append(bbox_to_points([text_remove_mask[0], text_bbox[1], text_remove_mask[1], text_bbox[3]]))
|
||||
if len(temp_dt_box) > 0:
|
||||
new_dt_boxes.extend(temp_dt_box)
|
||||
|
||||
new_dt_boxes.extend(angle_boxes_list)
|
||||
|
||||
return new_dt_boxes
|
||||
|
||||
|
||||
def merge_overlapping_spans(spans):
|
||||
"""
|
||||
Merges overlapping spans on the same line.
|
||||
|
||||
:param spans: A list of span coordinates [(x1, y1, x2, y2), ...]
|
||||
:return: A list of merged spans
|
||||
"""
|
||||
# Return an empty list if the input spans list is empty
|
||||
if not spans:
|
||||
return []
|
||||
|
||||
# Sort spans by their starting x-coordinate
|
||||
spans.sort(key=lambda x: x[0])
|
||||
|
||||
# Initialize the list of merged spans
|
||||
merged = []
|
||||
for span in spans:
|
||||
# Unpack span coordinates
|
||||
x1, y1, x2, y2 = span
|
||||
# If the merged list is empty or there's no horizontal overlap, add the span directly
|
||||
if not merged or merged[-1][2] < x1:
|
||||
merged.append(span)
|
||||
else:
|
||||
# If there is horizontal overlap, merge the current span with the previous one
|
||||
last_span = merged.pop()
|
||||
# Update the merged span's top-left corner to the smaller (x1, y1) and bottom-right to the larger (x2, y2)
|
||||
x1 = min(last_span[0], x1)
|
||||
y1 = min(last_span[1], y1)
|
||||
x2 = max(last_span[2], x2)
|
||||
y2 = max(last_span[3], y2)
|
||||
# Add the merged span back to the list
|
||||
merged.append((x1, y1, x2, y2))
|
||||
|
||||
# Return the list of merged spans
|
||||
return merged
|
||||
|
||||
|
||||
def merge_det_boxes(dt_boxes):
|
||||
"""
|
||||
Merge detection boxes.
|
||||
|
||||
This function takes a list of detected bounding boxes, each represented by four corner points.
|
||||
The goal is to merge these bounding boxes into larger text regions.
|
||||
|
||||
Parameters:
|
||||
dt_boxes (list): A list containing multiple text detection boxes, where each box is defined by four corner points.
|
||||
|
||||
Returns:
|
||||
list: A list containing the merged text regions, where each region is represented by four corner points.
|
||||
"""
|
||||
# Convert the detection boxes into a dictionary format with bounding boxes and type
|
||||
dt_boxes_dict_list = []
|
||||
angle_boxes_list = []
|
||||
for text_box in dt_boxes:
|
||||
text_bbox = points_to_bbox(text_box)
|
||||
|
||||
if calculate_is_angle(text_box):
|
||||
angle_boxes_list.append(text_box)
|
||||
continue
|
||||
|
||||
text_box_dict = {'bbox': text_bbox}
|
||||
dt_boxes_dict_list.append(text_box_dict)
|
||||
|
||||
# Merge adjacent text regions into lines
|
||||
lines = merge_spans_to_line(dt_boxes_dict_list)
|
||||
|
||||
# Initialize a new list for storing the merged text regions
|
||||
new_dt_boxes = []
|
||||
for line in lines:
|
||||
line_bbox_list = []
|
||||
for span in line:
|
||||
line_bbox_list.append(span['bbox'])
|
||||
|
||||
# Merge overlapping text regions within the same line
|
||||
merged_spans = merge_overlapping_spans(line_bbox_list)
|
||||
|
||||
# Convert the merged text regions back to point format and add them to the new detection box list
|
||||
for span in merged_spans:
|
||||
new_dt_boxes.append(bbox_to_points(span))
|
||||
|
||||
new_dt_boxes.extend(angle_boxes_list)
|
||||
|
||||
return new_dt_boxes
|
||||
|
||||
|
||||
def get_adjusted_mfdetrec_res(single_page_mfdetrec_res, useful_list):
|
||||
paste_x, paste_y, xmin, ymin, xmax, ymax, new_width, new_height = useful_list
|
||||
# Adjust the coordinates of the formula area
|
||||
adjusted_mfdetrec_res = []
|
||||
for mf_res in single_page_mfdetrec_res:
|
||||
mf_xmin, mf_ymin, mf_xmax, mf_ymax = mf_res["bbox"]
|
||||
# Adjust the coordinates of the formula area to the coordinates relative to the cropping area
|
||||
x0 = mf_xmin - xmin + paste_x
|
||||
y0 = mf_ymin - ymin + paste_y
|
||||
x1 = mf_xmax - xmin + paste_x
|
||||
y1 = mf_ymax - ymin + paste_y
|
||||
# Filter formula blocks outside the graph
|
||||
if any([x1 < 0, y1 < 0]) or any([x0 > new_width, y0 > new_height]):
|
||||
continue
|
||||
else:
|
||||
adjusted_mfdetrec_res.append({
|
||||
"bbox": [x0, y0, x1, y1],
|
||||
})
|
||||
return adjusted_mfdetrec_res
|
||||
|
||||
|
||||
def get_ocr_result_list(ocr_res, useful_list, ocr_enable, new_image, lang):
|
||||
paste_x, paste_y, xmin, ymin, xmax, ymax, new_width, new_height = useful_list
|
||||
ocr_result_list = []
|
||||
ori_im = new_image.copy()
|
||||
for box_ocr_res in ocr_res:
|
||||
|
||||
if len(box_ocr_res) == 2:
|
||||
p1, p2, p3, p4 = box_ocr_res[0]
|
||||
text, score = box_ocr_res[1]
|
||||
# logger.info(f"text: {text}, score: {score}")
|
||||
if score < 0.6: # 过滤低置信度的结果
|
||||
continue
|
||||
else:
|
||||
p1, p2, p3, p4 = box_ocr_res
|
||||
text, score = "", 1
|
||||
|
||||
if ocr_enable:
|
||||
tmp_box = copy.deepcopy(np.array([p1, p2, p3, p4]).astype('float32'))
|
||||
img_crop = get_rotate_crop_image(ori_im, tmp_box)
|
||||
|
||||
# average_angle_degrees = calculate_angle_degrees(box_ocr_res[0])
|
||||
# if average_angle_degrees > 0.5:
|
||||
poly = [p1, p2, p3, p4]
|
||||
if calculate_is_angle(poly):
|
||||
# logger.info(f"average_angle_degrees: {average_angle_degrees}, text: {text}")
|
||||
# 与x轴的夹角超过0.5度,对边界做一下矫正
|
||||
# 计算几何中心
|
||||
x_center = sum(point[0] for point in poly) / 4
|
||||
y_center = sum(point[1] for point in poly) / 4
|
||||
new_height = ((p4[1] - p1[1]) + (p3[1] - p2[1])) / 2
|
||||
new_width = p3[0] - p1[0]
|
||||
p1 = [x_center - new_width / 2, y_center - new_height / 2]
|
||||
p2 = [x_center + new_width / 2, y_center - new_height / 2]
|
||||
p3 = [x_center + new_width / 2, y_center + new_height / 2]
|
||||
p4 = [x_center - new_width / 2, y_center + new_height / 2]
|
||||
|
||||
# Convert the coordinates back to the original coordinate system
|
||||
p1 = [p1[0] - paste_x + xmin, p1[1] - paste_y + ymin]
|
||||
p2 = [p2[0] - paste_x + xmin, p2[1] - paste_y + ymin]
|
||||
p3 = [p3[0] - paste_x + xmin, p3[1] - paste_y + ymin]
|
||||
p4 = [p4[0] - paste_x + xmin, p4[1] - paste_y + ymin]
|
||||
|
||||
if ocr_enable:
|
||||
ocr_result_list.append({
|
||||
'category_id': 15,
|
||||
'poly': p1 + p2 + p3 + p4,
|
||||
'score': 1,
|
||||
'text': text,
|
||||
'np_img': img_crop,
|
||||
'lang': lang,
|
||||
})
|
||||
else:
|
||||
ocr_result_list.append({
|
||||
'category_id': 15,
|
||||
'poly': p1 + p2 + p3 + p4,
|
||||
'score': float(round(score, 2)),
|
||||
'text': text,
|
||||
})
|
||||
|
||||
return ocr_result_list
|
||||
|
||||
|
||||
def calculate_is_angle(poly):
|
||||
p1, p2, p3, p4 = poly
|
||||
height = ((p4[1] - p1[1]) + (p3[1] - p2[1])) / 2
|
||||
if 0.8 * height <= (p3[1] - p1[1]) <= 1.2 * height:
|
||||
return False
|
||||
else:
|
||||
# logger.info((p3[1] - p1[1])/height)
|
||||
return True
|
||||
|
||||
|
||||
def get_rotate_crop_image(img, points):
|
||||
'''
|
||||
img_height, img_width = img.shape[0:2]
|
||||
left = int(np.min(points[:, 0]))
|
||||
right = int(np.max(points[:, 0]))
|
||||
top = int(np.min(points[:, 1]))
|
||||
bottom = int(np.max(points[:, 1]))
|
||||
img_crop = img[top:bottom, left:right, :].copy()
|
||||
points[:, 0] = points[:, 0] - left
|
||||
points[:, 1] = points[:, 1] - top
|
||||
'''
|
||||
assert len(points) == 4, "shape of points must be 4*2"
|
||||
img_crop_width = int(
|
||||
max(
|
||||
np.linalg.norm(points[0] - points[1]),
|
||||
np.linalg.norm(points[2] - points[3])))
|
||||
img_crop_height = int(
|
||||
max(
|
||||
np.linalg.norm(points[0] - points[3]),
|
||||
np.linalg.norm(points[1] - points[2])))
|
||||
pts_std = np.float32([[0, 0], [img_crop_width, 0],
|
||||
[img_crop_width, img_crop_height],
|
||||
[0, img_crop_height]])
|
||||
M = cv2.getPerspectiveTransform(points, pts_std)
|
||||
dst_img = cv2.warpPerspective(
|
||||
img,
|
||||
M, (img_crop_width, img_crop_height),
|
||||
borderMode=cv2.BORDER_REPLICATE,
|
||||
flags=cv2.INTER_CUBIC)
|
||||
dst_img_height, dst_img_width = dst_img.shape[0:2]
|
||||
if dst_img_height * 1.0 / dst_img_width >= 1.5:
|
||||
dst_img = np.rot90(dst_img)
|
||||
return dst_img
|
||||
@@ -5,8 +5,8 @@ import pypdfium2 as pdfium
|
||||
from loguru import logger
|
||||
from PIL import Image
|
||||
|
||||
from ..data.data_reader_writer import FileBasedDataWriter
|
||||
from ..utils.pdf_reader import image_to_b64str, image_to_bytes, page_to_image
|
||||
from mineru.data.data_reader_writer import FileBasedDataWriter
|
||||
from mineru.utils.pdf_reader import image_to_b64str, image_to_bytes, page_to_image
|
||||
from .hash_utils import str_sha256
|
||||
|
||||
|
||||
Reference in New Issue
Block a user