mirror of
https://github.com/PaddlePaddle/PaddleOCR.git
synced 2026-09-24 23:33:08 +08:00
Merge branch 'dygraph' of https://github.com/PaddlePaddle/PaddleOCR into drrg_branch
This commit is contained in:
@@ -77,7 +77,7 @@ def export_single_model(model,
|
||||
elif arch_config["algorithm"] == "PREN":
|
||||
other_shape = [
|
||||
paddle.static.InputSpec(
|
||||
shape=[None, 3, 64, 512], dtype="float32"),
|
||||
shape=[None, 3, 64, 256], dtype="float32"),
|
||||
]
|
||||
model = to_static(model, input_spec=other_shape)
|
||||
elif arch_config["model_type"] == "sr":
|
||||
@@ -99,7 +99,7 @@ def export_single_model(model,
|
||||
]
|
||||
# print([None, 3, 32, 128])
|
||||
model = to_static(model, input_spec=other_shape)
|
||||
elif arch_config["algorithm"] in ["NRTR", "SPIN"]:
|
||||
elif arch_config["algorithm"] in ["NRTR", "SPIN", 'RFL']:
|
||||
other_shape = [
|
||||
paddle.static.InputSpec(
|
||||
shape=[None, 1, 32, 100], dtype="float32"),
|
||||
|
||||
@@ -100,6 +100,14 @@ class TextRecognizer(object):
|
||||
"use_space_char": args.use_space_char,
|
||||
"rm_symbol": True
|
||||
}
|
||||
elif self.rec_algorithm == 'RFL':
|
||||
postprocess_params = {
|
||||
'name': 'RFLLabelDecode',
|
||||
"character_dict_path": None,
|
||||
"use_space_char": args.use_space_char
|
||||
}
|
||||
elif self.rec_algorithm == "PREN":
|
||||
postprocess_params = {'name': 'PRENLabelDecode'}
|
||||
self.postprocess_op = build_post_process(postprocess_params)
|
||||
self.predictor, self.input_tensor, self.output_tensors, self.config = \
|
||||
utility.create_predictor(args, 'rec', logger)
|
||||
@@ -143,6 +151,16 @@ class TextRecognizer(object):
|
||||
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]
|
||||
imgW = int((imgH * max_wh_ratio))
|
||||
@@ -384,7 +402,7 @@ class TextRecognizer(object):
|
||||
self.rec_image_shape)
|
||||
norm_img = norm_img[np.newaxis, :]
|
||||
norm_img_batch.append(norm_img)
|
||||
elif self.rec_algorithm == "VisionLAN":
|
||||
elif self.rec_algorithm in ["VisionLAN", "PREN"]:
|
||||
norm_img = self.resize_norm_img_vl(img_list[indices[ino]],
|
||||
self.rec_image_shape)
|
||||
norm_img = norm_img[np.newaxis, :]
|
||||
|
||||
+10
-4
@@ -97,7 +97,8 @@ def main():
|
||||
elif config['Architecture']['algorithm'] == "SAR":
|
||||
op[op_name]['keep_keys'] = ['image', 'valid_ratio']
|
||||
elif config['Architecture']['algorithm'] == "RobustScanner":
|
||||
op[op_name]['keep_keys'] = ['image', 'valid_ratio', 'word_positons']
|
||||
op[op_name][
|
||||
'keep_keys'] = ['image', 'valid_ratio', 'word_positons']
|
||||
else:
|
||||
op[op_name]['keep_keys'] = ['image']
|
||||
transforms.append(op)
|
||||
@@ -136,9 +137,10 @@ def main():
|
||||
if config['Architecture']['algorithm'] == "RobustScanner":
|
||||
valid_ratio = np.expand_dims(batch[1], axis=0)
|
||||
word_positons = np.expand_dims(batch[2], axis=0)
|
||||
img_metas = [paddle.to_tensor(valid_ratio),
|
||||
paddle.to_tensor(word_positons),
|
||||
]
|
||||
img_metas = [
|
||||
paddle.to_tensor(valid_ratio),
|
||||
paddle.to_tensor(word_positons),
|
||||
]
|
||||
images = np.expand_dims(batch[0], axis=0)
|
||||
images = paddle.to_tensor(images)
|
||||
if config['Architecture']['algorithm'] == "SRN":
|
||||
@@ -160,6 +162,10 @@ def main():
|
||||
"score": float(post_result[key][0][1]),
|
||||
}
|
||||
info = json.dumps(rec_info, ensure_ascii=False)
|
||||
elif isinstance(post_result, list) and isinstance(post_result[0],
|
||||
int):
|
||||
# for RFLearning CNT branch
|
||||
info = str(post_result[0])
|
||||
else:
|
||||
if len(post_result[0]) >= 2:
|
||||
info = post_result[0][0] + "\t" + str(post_result[0][1])
|
||||
|
||||
+10
-4
@@ -114,7 +114,7 @@ def merge_config(config, opts):
|
||||
return config
|
||||
|
||||
|
||||
def check_device(use_gpu, use_xpu=False, use_npu=False):
|
||||
def check_device(use_gpu, use_xpu=False, use_npu=False, use_mlu=False):
|
||||
"""
|
||||
Log error and exit when set use_gpu=true in paddlepaddle
|
||||
cpu version.
|
||||
@@ -137,6 +137,9 @@ def check_device(use_gpu, use_xpu=False, use_npu=False):
|
||||
if use_npu and not paddle.device.is_compiled_with_npu():
|
||||
print(err.format("use_npu", "npu", "npu", "use_npu"))
|
||||
sys.exit(1)
|
||||
if use_mlu and not paddle.device.is_compiled_with_mlu():
|
||||
print(err.format("use_mlu", "mlu", "mlu", "use_mlu"))
|
||||
sys.exit(1)
|
||||
except Exception as e:
|
||||
pass
|
||||
|
||||
@@ -217,7 +220,7 @@ def train(config,
|
||||
use_srn = config['Architecture']['algorithm'] == "SRN"
|
||||
extra_input_models = [
|
||||
"SRN", "NRTR", "SAR", "SEED", "SVTR", "SPIN", "VisionLAN",
|
||||
"RobustScanner", 'DRRG'
|
||||
"RobustScanner", "RFL", 'DRRG'
|
||||
]
|
||||
extra_input = False
|
||||
if config['Architecture']['algorithm'] == 'Distillation':
|
||||
@@ -618,6 +621,7 @@ def preprocess(is_train=False):
|
||||
use_gpu = config['Global'].get('use_gpu', False)
|
||||
use_xpu = config['Global'].get('use_xpu', False)
|
||||
use_npu = config['Global'].get('use_npu', False)
|
||||
use_mlu = config['Global'].get('use_mlu', False)
|
||||
|
||||
alg = config['Architecture']['algorithm']
|
||||
assert alg in [
|
||||
@@ -625,17 +629,19 @@ def preprocess(is_train=False):
|
||||
'CLS', 'PGNet', 'Distillation', 'NRTR', 'TableAttn', 'SAR', 'PSE',
|
||||
'SEED', 'SDMGR', 'LayoutXLM', 'LayoutLM', 'LayoutLMv2', 'PREN', 'FCE',
|
||||
'SVTR', 'ViTSTR', 'ABINet', 'DB++', 'TableMaster', 'SPIN', 'VisionLAN',
|
||||
'Gestalt', 'SLANet', 'RobustScanner', 'CT', 'DRRG'
|
||||
'Gestalt', 'SLANet', 'RobustScanner', 'CT', 'RFL', 'DRRG'
|
||||
]
|
||||
|
||||
if use_xpu:
|
||||
device = 'xpu:{0}'.format(os.getenv('FLAGS_selected_xpus', 0))
|
||||
elif use_npu:
|
||||
device = 'npu:{0}'.format(os.getenv('FLAGS_selected_npus', 0))
|
||||
elif use_mlu:
|
||||
device = 'mlu:{0}'.format(os.getenv('FLAGS_selected_mlus', 0))
|
||||
else:
|
||||
device = 'gpu:{}'.format(dist.ParallelEnv()
|
||||
.dev_id) if use_gpu else 'cpu'
|
||||
check_device(use_gpu, use_xpu, use_npu)
|
||||
check_device(use_gpu, use_xpu, use_npu, use_mlu)
|
||||
|
||||
device = paddle.set_device(device)
|
||||
|
||||
|
||||
+5
-4
@@ -149,10 +149,11 @@ def main(config, device, logger, vdl_writer):
|
||||
amp_level = config["Global"].get("amp_level", 'O2')
|
||||
amp_custom_black_list = config['Global'].get('amp_custom_black_list', [])
|
||||
if use_amp:
|
||||
AMP_RELATED_FLAGS_SETTING = {
|
||||
'FLAGS_cudnn_batchnorm_spatial_persistent': 1,
|
||||
'FLAGS_max_inplace_grad_add': 8,
|
||||
}
|
||||
AMP_RELATED_FLAGS_SETTING = {'FLAGS_max_inplace_grad_add': 8, }
|
||||
if paddle.is_compiled_with_cuda():
|
||||
AMP_RELATED_FLAGS_SETTING.update({
|
||||
'FLAGS_cudnn_batchnorm_spatial_persistent': 1
|
||||
})
|
||||
paddle.fluid.set_flags(AMP_RELATED_FLAGS_SETTING)
|
||||
scale_loss = config["Global"].get("scale_loss", 1.0)
|
||||
use_dynamic_loss_scaling = config["Global"].get(
|
||||
|
||||
Reference in New Issue
Block a user