V4Rec code pr (#9725)

* v4rec code

* v4rec add nrtrloss

* Add V4rec backbone file

* Add V4Rec config file.

* Fix V4rec reparameters when export_model

* convert lvnetv3

* fix codestyle

* fix infer_rec v4rec
This commit is contained in:
topduke
2023-04-19 16:04:23 +08:00
committed by GitHub
parent 385a1f99d4
commit 43abe2fa4b
15 changed files with 787 additions and 71 deletions
+15 -6
View File
@@ -1241,27 +1241,36 @@ class MultiLabelEncode(BaseRecLabelEncode):
max_text_length,
character_dict_path=None,
use_space_char=False,
gtc_encode=None,
**kwargs):
super(MultiLabelEncode, self).__init__(
max_text_length, character_dict_path, use_space_char)
self.ctc_encode = CTCLabelEncode(max_text_length, character_dict_path,
use_space_char, **kwargs)
self.sar_encode = SARLabelEncode(max_text_length, character_dict_path,
use_space_char, **kwargs)
self.gtc_encode_type = gtc_encode
if gtc_encode is None:
self.gtc_encode = SARLabelEncode(
max_text_length, character_dict_path, use_space_char, **kwargs)
else:
self.gtc_encode = eval(gtc_encode)(
max_text_length, character_dict_path, use_space_char, **kwargs)
def __call__(self, data):
data_ctc = copy.deepcopy(data)
data_sar = copy.deepcopy(data)
data_gtc = copy.deepcopy(data)
data_out = dict()
data_out['img_path'] = data.get('img_path', None)
data_out['image'] = data['image']
ctc = self.ctc_encode.__call__(data_ctc)
sar = self.sar_encode.__call__(data_sar)
if ctc is None or sar is None:
gtc = self.gtc_encode.__call__(data_gtc)
if ctc is None or gtc is None:
return None
data_out['label_ctc'] = ctc['label']
data_out['label_sar'] = sar['label']
if self.gtc_encode_type is not None:
data_out['label_gtc'] = gtc['label']
else:
data_out['label_sar'] = gtc['label']
data_out['length'] = ctc['length']
return data_out