mirror of
https://github.com/PaddlePaddle/PaddleOCR.git
synced 2026-09-24 23:33:08 +08:00
Add rec algo VisionLAN (#6943)
* add vl * add vl * add vl * add ref * fix head out * add visionlan doc * fix vl infer * update dict
This commit is contained in:
+8
-4
@@ -227,7 +227,9 @@ def train(config,
|
||||
model.train()
|
||||
|
||||
use_srn = config['Architecture']['algorithm'] == "SRN"
|
||||
extra_input_models = ["SRN", "NRTR", "SAR", "SEED", "SVTR", "SPIN"]
|
||||
extra_input_models = [
|
||||
"SRN", "NRTR", "SAR", "SEED", "SVTR", "SPIN", "VisionLAN"
|
||||
]
|
||||
extra_input = False
|
||||
if config['Architecture']['algorithm'] == 'Distillation':
|
||||
for key in config['Architecture']["Models"]:
|
||||
@@ -269,7 +271,6 @@ def train(config,
|
||||
images = batch[0]
|
||||
if use_srn:
|
||||
model_average = True
|
||||
|
||||
# use amp
|
||||
if scaler:
|
||||
with paddle.amp.auto_cast(level='O2'):
|
||||
@@ -310,6 +311,9 @@ def train(config,
|
||||
]: # for multi head loss
|
||||
post_result = post_process_class(
|
||||
preds['ctc'], batch[1]) # for CTC head out
|
||||
elif config['Loss']['name'] in ['VLLoss']:
|
||||
post_result = post_process_class(preds, batch[1],
|
||||
batch[-1])
|
||||
else:
|
||||
post_result = post_process_class(preds, batch[1])
|
||||
eval_class(post_result, batch)
|
||||
@@ -612,7 +616,7 @@ def preprocess(is_train=False):
|
||||
'EAST', 'DB', 'SAST', 'Rosetta', 'CRNN', 'STARNet', 'RARE', 'SRN',
|
||||
'CLS', 'PGNet', 'Distillation', 'NRTR', 'TableAttn', 'SAR', 'PSE',
|
||||
'SEED', 'SDMGR', 'LayoutXLM', 'LayoutLM', 'LayoutLMv2', 'PREN', 'FCE',
|
||||
'SVTR', 'ViTSTR', 'ABINet', 'DB++', 'TableMaster', 'SPIN'
|
||||
'SVTR', 'ViTSTR', 'ABINet', 'DB++', 'TableMaster', 'SPIN', 'VisionLAN'
|
||||
]
|
||||
|
||||
if use_xpu:
|
||||
@@ -631,7 +635,7 @@ def preprocess(is_train=False):
|
||||
if 'use_visualdl' in config['Global'] and config['Global']['use_visualdl']:
|
||||
save_model_dir = config['Global']['save_model_dir']
|
||||
vdl_writer_path = '{}/vdl/'.format(save_model_dir)
|
||||
log_writer = VDLLogger(save_model_dir)
|
||||
log_writer = VDLLogger(vdl_writer_path)
|
||||
loggers.append(log_writer)
|
||||
if ('use_wandb' in config['Global'] and
|
||||
config['Global']['use_wandb']) or 'wandb' in config:
|
||||
|
||||
Reference in New Issue
Block a user