litealpr 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
litealpr/__init__.py ADDED
@@ -0,0 +1,5 @@
1
+ __version__ = '0.1.0'
2
+
3
+ from .pipeline import LiteALPR
4
+
5
+ __all__ = ['LiteALPR']
@@ -0,0 +1 @@
1
+ # det package
@@ -0,0 +1,37 @@
1
+ 0
2
+ 1
3
+ 2
4
+ 3
5
+ 4
6
+ 5
7
+ 6
8
+ 7
9
+ 8
10
+ 9
11
+ A
12
+ B
13
+ C
14
+ D
15
+ Đ
16
+ E
17
+ F
18
+ G
19
+ H
20
+ I
21
+ J
22
+ K
23
+ L
24
+ M
25
+ N
26
+ O
27
+ P
28
+ Q
29
+ R
30
+ S
31
+ T
32
+ U
33
+ V
34
+ W
35
+ X
36
+ Y
37
+ Z
litealpr/pipeline.py ADDED
@@ -0,0 +1,187 @@
1
+ import os
2
+ import sys
3
+ import warnings
4
+
5
+ import cv2
6
+ import torch
7
+
8
+ warnings.filterwarnings("ignore")
9
+
10
+ # This ensures that we can import det and rec if running from this file
11
+ __dir__ = os.path.dirname(os.path.abspath(__file__))
12
+ root_dir = os.path.abspath(os.path.join(__dir__, '..'))
13
+ if root_dir not in sys.path:
14
+ sys.path.insert(0, root_dir)
15
+
16
+ from ultralytics import YOLO
17
+
18
+ from litealpr.rec.modeling import build_model
19
+ from litealpr.rec.postprocess import build_post_process
20
+
21
+
22
+ def download_from_hf(filename):
23
+ try:
24
+ from huggingface_hub import hf_hub_download
25
+ print(f"[LiteALPR] Downloading/Verifying {filename} from HuggingFace...")
26
+ return hf_hub_download(repo_id="anhone3/LiteALPR", filename=filename)
27
+ except ImportError:
28
+ raise ImportError("Please install huggingface_hub to auto-download models: pip install huggingface_hub")
29
+
30
+
31
+ class LiteALPR:
32
+ def __init__(self,
33
+ use_det=True,
34
+ use_rec=True,
35
+ det_model_path=None,
36
+ rec_model_path=None,
37
+ device=None):
38
+
39
+ self.device = torch.device(device if device else ('cuda:0' if torch.cuda.is_available() else 'cpu'))
40
+ print(f"[LiteALPR] Initializing on {self.device}...")
41
+
42
+ self.det_model = None
43
+ self.rec_model = None
44
+ self.post_process_class = None
45
+
46
+ # 1. Initialize DET (YOLO)
47
+ if use_det:
48
+ if det_model_path is None:
49
+ det_model_path = download_from_hf("yolov8n_efficient/best.pt")
50
+
51
+ print(f"[LiteALPR] Loading Detection Model: {det_model_path}")
52
+ self.det_model = YOLO(det_model_path)
53
+ self.det_model.to(self.device)
54
+ self.det_model.fuse()
55
+
56
+ # 2. Initialize REC (SVTR26)
57
+ if use_rec:
58
+ if rec_model_path is None:
59
+ rec_model_path = download_from_hf("svtr26_tiny/best.pth")
60
+
61
+ print(f"[LiteALPR] Loading Recognition Model: {rec_model_path}")
62
+ checkpoint = torch.load(rec_model_path, map_location='cpu')
63
+ cfg = checkpoint['config']
64
+
65
+ # Monkey patch the dictionary path
66
+ dict_path = os.path.join(os.path.dirname(__file__), 'license_plate_dict.txt')
67
+ cfg['Global']['character_dict_path'] = dict_path
68
+
69
+ self.post_process_class = build_post_process(cfg['PostProcess'], cfg['Global'])
70
+ cfg['Architecture']['Decoder']['out_channels'] = self.post_process_class.get_character_num()
71
+
72
+ self.rec_model = build_model(cfg['Architecture'])
73
+ self.rec_model.load_state_dict(checkpoint['state_dict'], strict=True)
74
+ self.rec_model.to(self.device)
75
+ self.rec_model.eval()
76
+
77
+ if self.det_model or self.rec_model:
78
+ print("[LiteALPR] Models loaded successfully!")
79
+ else:
80
+ print("[LiteALPR] WARNING: No models loaded. Please provide det_model_path or rec_model_path.")
81
+
82
+ def _preprocess_crop(self, img_crop, max_ratio=12, base_shape=None, base_h=32):
83
+ """Preprocesses cropped image for SVTR with RatioRecTVResize logic using PIL (to match training)."""
84
+ from PIL import Image
85
+ from torchvision import transforms as T
86
+ from torchvision.transforms import functional as F
87
+
88
+ # Convert OpenCV BGR to PIL RGB
89
+ if base_shape is None:
90
+ base_shape = [[64, 64], [96, 48], [112, 40], [128, 32]]
91
+ img_rgb = cv2.cvtColor(img_crop, cv2.COLOR_BGR2RGB)
92
+ img = Image.fromarray(img_rgb)
93
+
94
+ w, h = img.size
95
+ gen_ratio = int(float(w) / float(h)) + 1
96
+ ratio_resize = min(gen_ratio, max_ratio)
97
+
98
+ if ratio_resize <= 4:
99
+ imgW, imgH = base_shape[ratio_resize - 1]
100
+ else:
101
+ imgW, imgH = [base_h * ratio_resize, base_h]
102
+
103
+ # SVTR is sensitive to interpolation algorithms; we MUST use PIL BICUBIC like during training
104
+ resized_image = F.resize(img, (imgH, imgW), interpolation=T.InterpolationMode.BICUBIC)
105
+
106
+ transforms = T.Compose([
107
+ T.ToTensor(),
108
+ T.Normalize(0.5, 0.5)
109
+ ])
110
+
111
+ tensor = transforms(resized_image)
112
+ tensor = tensor.unsqueeze(0)
113
+ return tensor.to(self.device)
114
+
115
+ def detect(self, img, conf_thresh=0.25):
116
+ """
117
+ Detect license plates in an image.
118
+ Returns: list of bounding boxes [[x1, y1, x2, y2], ...]
119
+ """
120
+ if self.det_model is None:
121
+ raise ValueError("Detection model not loaded! Initialize with det_model_path.")
122
+
123
+ if isinstance(img, str):
124
+ img = cv2.imread(img)
125
+
126
+ det_results = self.det_model(img, verbose=False, conf=conf_thresh)[0]
127
+ boxes = det_results.boxes.data.cpu().numpy() # [x1, y1, x2, y2, conf, cls]
128
+
129
+ results = []
130
+ for box in boxes:
131
+ x1, y1, x2, y2, _conf, _cls = box
132
+ results.append([int(x1), int(y1), int(x2), int(y2)])
133
+ return results
134
+
135
+ def recognize(self, crop_img):
136
+ """
137
+ Recognize text in a cropped license plate image.
138
+ Returns: tuple (text, confidence_score)
139
+ """
140
+ if self.rec_model is None:
141
+ raise ValueError("Recognition model not loaded! Initialize with rec_model_path.")
142
+
143
+ if crop_img.size == 0:
144
+ return "", 0.0
145
+
146
+ tensor = self._preprocess_crop(crop_img)
147
+ with torch.no_grad():
148
+ preds = self.rec_model(tensor)
149
+
150
+ post_result = self.post_process_class(preds)
151
+ text, score = post_result[0]
152
+ return text, float(score)
153
+
154
+ def read(self, image_path, conf_thresh=0.25):
155
+ """
156
+ End-to-End inference.
157
+ Returns: list of dicts [{'box': [x1,y1,x2,y2], 'text': '51F1234', 'score': 0.99}]
158
+ """
159
+ if self.det_model is None or self.rec_model is None:
160
+ raise ValueError("End-to-End read() requires BOTH det_model_path and rec_model_path to be loaded.")
161
+
162
+ img = cv2.imread(image_path)
163
+ if img is None:
164
+ raise ValueError(f"Could not read image: {image_path}")
165
+
166
+ boxes = self.detect(img, conf_thresh)
167
+ final_results = []
168
+
169
+ h, w, _ = img.shape
170
+ for box in boxes:
171
+ x1, y1, x2, y2 = box
172
+
173
+ # Ensure bounds
174
+ x1_c, y1_c = max(0, x1), max(0, y1)
175
+ x2_c, y2_c = min(w, x2), min(h, y2)
176
+
177
+ crop_img = img[y1_c:y2_c, x1_c:x2_c]
178
+
179
+ text, score = self.recognize(crop_img)
180
+
181
+ final_results.append({
182
+ 'box': box,
183
+ 'text': text,
184
+ 'score': score
185
+ })
186
+
187
+ return final_results
@@ -0,0 +1,71 @@
1
+ import copy
2
+ from importlib import import_module
3
+
4
+ from torch import nn
5
+
6
+ name_to_module = {
7
+ 'ABINetLoss': '.abinet_loss',
8
+ 'ARLoss': '.ar_loss',
9
+ 'CDistNetLoss': '.cdistnet_loss',
10
+ 'CELoss': '.ce_loss',
11
+ 'CPPDLoss': '.cppd_loss',
12
+ 'CTCLoss': '.ctc_loss',
13
+ 'IGTRLoss': '.igtr_loss',
14
+ 'LISTERLoss': '.lister_loss',
15
+ 'LPVLoss': '.lpv_loss',
16
+ 'MGPLoss': '.mgp_loss',
17
+ 'PARSeqLoss': '.parseq_loss',
18
+ 'RobustScannerLoss': '.robustscanner_loss',
19
+ 'SEEDLoss': '.seed_loss',
20
+ 'SMTRLoss': '.smtr_loss',
21
+ 'SRNLoss': '.srn_loss',
22
+ 'VisionLANLoss': '.visionlan_loss',
23
+ 'CAMLoss': '.cam_loss',
24
+ 'MDiffLoss': '.mdiff_loss',
25
+ 'UniRecLoss': '.unirec_loss',
26
+ 'CMERLoss': '.cmer_loss',
27
+ }
28
+
29
+
30
+ def build_loss(config):
31
+ config = copy.deepcopy(config)
32
+ module_name = config.pop('name')
33
+
34
+ if module_name in globals():
35
+ module_class = globals()[module_name]
36
+ else:
37
+ assert module_name in name_to_module, Exception(
38
+ f'{module_name} is not supported. The losses in {list(name_to_module.keys())} are supportes')
39
+ module_path = name_to_module[module_name]
40
+ module = import_module(module_path, package=__package__)
41
+ module_class = getattr(module, module_name)
42
+
43
+ return module_class(**config)
44
+
45
+
46
+ class GTCLoss(nn.Module):
47
+
48
+ def __init__(self,
49
+ gtc_loss,
50
+ gtc_weight=1.0,
51
+ ctc_weight=1.0,
52
+ zero_infinity=True,
53
+ **kwargs):
54
+ super().__init__()
55
+ # Dynamically build CTCLoss
56
+ ctc_config = {'name': 'CTCLoss', 'zero_infinity': zero_infinity}
57
+ self.ctc_loss = build_loss(ctc_config)
58
+ # Build GTC loss
59
+ self.gtc_loss = build_loss(gtc_loss)
60
+ self.gtc_weight = gtc_weight
61
+ self.ctc_weight = ctc_weight
62
+
63
+ def forward(self, predicts, batch):
64
+ ctc_loss = self.ctc_loss(predicts['ctc_pred'],
65
+ [None] + batch[-2:])['loss']
66
+ gtc_loss = self.gtc_loss(predicts['gtc_pred'], batch[:-2])['loss']
67
+ return {
68
+ 'loss': self.ctc_weight * ctc_loss + self.gtc_weight * gtc_loss,
69
+ 'ctc_loss': ctc_loss,
70
+ 'gtc_loss': gtc_loss
71
+ }
@@ -0,0 +1,33 @@
1
+ import torch
2
+ from torch import nn
3
+
4
+
5
+ class CTCLoss(nn.Module):
6
+
7
+ def __init__(self, use_focal_loss=False, zero_infinity=False, **kwargs):
8
+ super().__init__()
9
+ self.loss_func = nn.CTCLoss(blank=0,
10
+ reduction='none',
11
+ zero_infinity=zero_infinity)
12
+ self.use_focal_loss = use_focal_loss
13
+
14
+ def forward(self, predicts, batch):
15
+ # predicts = predicts['res']
16
+
17
+ batch_size = predicts.size(0)
18
+ label, label_length = batch[1], batch[2]
19
+ predicts = predicts.log_softmax(2)
20
+ predicts = predicts.permute(1, 0, 2)
21
+ preds_lengths = torch.tensor([predicts.size(0)] * batch_size,
22
+ dtype=torch.long)
23
+ loss = self.loss_func(predicts, label, preds_lengths, label_length)
24
+
25
+ if self.use_focal_loss:
26
+ # Use torch.clamp to limit the range of loss, avoiding overflow in exponential calculation
27
+ clamped_loss = torch.clamp(loss, min=-20, max=20)
28
+ weight = 1 - torch.exp(-clamped_loss)
29
+ weight = torch.square(weight)
30
+ # Use torch.where to avoid multiplying by zero weight
31
+ loss = torch.where(weight > 0, loss * weight, loss)
32
+ loss = loss.mean()
33
+ return {'loss': loss}
@@ -0,0 +1,33 @@
1
+ import copy
2
+
3
+ __all__ = ['build_metric']
4
+
5
+ support_dict = [
6
+ 'RecMetric', 'RecMetricLong', 'RecGTCMetric', 'RecMPGMetric', 'CMERMetric'
7
+ ]
8
+
9
+
10
+ def build_metric(config):
11
+ config = copy.deepcopy(config)
12
+ module_name = config.pop('name')
13
+ assert module_name in support_dict, Exception(
14
+ f'metric only support {support_dict}')
15
+
16
+ # Lazy import
17
+ if module_name == 'RecMetric':
18
+ from .rec_metric import RecMetric
19
+ module_class = RecMetric(**config)
20
+ elif module_name == 'RecGTCMetric':
21
+ from .rec_metric_gtc import RecGTCMetric
22
+ module_class = RecGTCMetric(**config)
23
+ elif module_name == 'RecMetricLong':
24
+ from .rec_metric_long import RecMetricLong
25
+ module_class = RecMetricLong(**config)
26
+ elif module_name == 'RecMPGMetric':
27
+ from .rec_metric_mgp import RecMPGMetric
28
+ module_class = RecMPGMetric(**config)
29
+ elif module_name == 'CMERMetric':
30
+ from .rec_metric_cmer import CMERMetric
31
+ module_class = CMERMetric(**config)
32
+
33
+ return module_class
@@ -0,0 +1,287 @@
1
+ import string
2
+
3
+ import numpy as np
4
+ from rapidfuzz.distance import Levenshtein
5
+
6
+
7
+ def match_ss(ss1, ss2):
8
+ s1_len = len(ss1)
9
+ for c_i in range(s1_len):
10
+ if ss1[c_i:] == ss2[:s1_len - c_i]:
11
+ return ss2[s1_len - c_i:]
12
+ return ss2
13
+
14
+
15
+ def stream_match(text):
16
+ bs = len(text)
17
+ s_list = []
18
+ conf_list = []
19
+ for s_conf in text:
20
+ s_list.append(s_conf[0])
21
+ conf_list.append(s_conf[1])
22
+ s_n = bs
23
+ s_start = s_list[0][:-1]
24
+ s_new = s_start
25
+ for s_i in range(1, s_n):
26
+ s_start = match_ss(
27
+ s_start, s_list[s_i][1:-1] if s_i < s_n - 1 else s_list[s_i][1:])
28
+ s_new += s_start
29
+ return s_new, sum(conf_list) / bs
30
+
31
+
32
+ class RecMetric:
33
+
34
+ def __init__(self,
35
+ main_indicator='acc',
36
+ is_filter=False,
37
+ is_lower=True,
38
+ ignore_space=True,
39
+ stream=False,
40
+ with_ratio=False,
41
+ max_len=25,
42
+ max_ratio=4,
43
+ **kwargs):
44
+ self.main_indicator = main_indicator
45
+ self.is_filter = is_filter
46
+ self.is_lower = is_lower
47
+ self.ignore_space = ignore_space
48
+ self.stream = stream
49
+ self.eps = 1e-5
50
+ self.with_ratio = with_ratio
51
+ self.max_len = max_len
52
+ self.max_ratio = max_ratio
53
+ self.reset()
54
+
55
+ def _normalize_text(self, text):
56
+ text = ''.join(
57
+ filter(lambda x: x in (string.digits + string.ascii_letters),
58
+ text))
59
+ return text
60
+
61
+ def __call__(self,
62
+ pred_label,
63
+ batch=None,
64
+ training=False,
65
+ *args,
66
+ **kwargs):
67
+ if self.with_ratio and not training:
68
+ return self.eval_all_metric(pred_label, batch)
69
+ else:
70
+ return self.eval_metric(pred_label)
71
+
72
+ def eval_metric(self, pred_label, *args, **kwargs):
73
+ preds, labels = pred_label
74
+ correct_num = 0
75
+ all_num = 0
76
+ norm_edit_dis = 0.0
77
+ total_edit_dis = 0.0
78
+ total_char_len = 0
79
+ for (pred, pred_conf), (target, _) in zip(preds, labels):
80
+ if self.stream:
81
+ assert len(labels) == 1
82
+ pred, _ = stream_match(preds)
83
+ if self.ignore_space:
84
+ pred = pred.replace(' ', '')
85
+ target = target.replace(' ', '')
86
+ if self.is_filter:
87
+ pred = self._normalize_text(pred)
88
+ target = self._normalize_text(target)
89
+ if self.is_lower:
90
+ pred = pred.lower()
91
+ target = target.lower()
92
+
93
+ # For CER: absolute Levenshtein distance
94
+ ed = Levenshtein.distance(pred, target)
95
+ total_edit_dis += ed
96
+ total_char_len += len(target)
97
+
98
+ norm_edit_dis += Levenshtein.normalized_distance(pred, target)
99
+ if pred == target:
100
+ correct_num += 1
101
+ all_num += 1
102
+ self.correct_num += correct_num
103
+ self.all_num += all_num
104
+ self.norm_edit_dis += norm_edit_dis
105
+ self.total_edit_dis += total_edit_dis
106
+ self.total_char_len += total_char_len
107
+
108
+ return {
109
+ 'acc': correct_num / (all_num + self.eps),
110
+ 'norm_edit_dis': 1 - norm_edit_dis / (all_num + self.eps),
111
+ 'cer': total_edit_dis / (total_char_len + self.eps),
112
+ }
113
+
114
+ def eval_all_metric(self, pred_label, batch=None, *args, **kwargs):
115
+ if self.with_ratio:
116
+ ratio = batch[-1]
117
+ preds, labels = pred_label
118
+ correct_num = 0
119
+ correct_num_real = 0
120
+ correct_num_lower = 0
121
+ correct_num_ignore_space = 0
122
+ correct_num_ignore_space_lower = 0
123
+ correct_num_ignore_space_symbol = 0
124
+ all_num = 0
125
+ norm_edit_dis = 0.0
126
+ total_edit_dis = 0.0
127
+ total_char_len = 0
128
+ each_len_num = [0 for _ in range(self.max_len)]
129
+ each_len_correct_num = [0 for _ in range(self.max_len)]
130
+ each_len_norm_edit_dis = [0 for _ in range(self.max_len)]
131
+ each_ratio_num = [0 for _ in range(self.max_ratio)]
132
+ each_ratio_correct_num = [0 for _ in range(self.max_ratio)]
133
+ each_ratio_norm_edit_dis = [0 for _ in range(self.max_ratio)]
134
+ for (pred, pred_conf), (target, _) in zip(preds, labels):
135
+ if self.stream:
136
+ assert len(labels) == 1
137
+ pred, _ = stream_match(preds)
138
+ if pred == target:
139
+ correct_num_real += 1
140
+
141
+ if pred.lower() == target.lower():
142
+ correct_num_lower += 1
143
+
144
+ if self.ignore_space:
145
+ pred = pred.replace(' ', '')
146
+ target = target.replace(' ', '')
147
+ if pred == target:
148
+ correct_num_ignore_space += 1
149
+
150
+ if pred.lower() == target.lower():
151
+ correct_num_ignore_space_lower += 1
152
+
153
+ if self.is_filter:
154
+ pred = self._normalize_text(pred)
155
+ target = self._normalize_text(target)
156
+ if pred == target:
157
+ correct_num_ignore_space_symbol += 1
158
+
159
+ if self.is_lower:
160
+ pred = pred.lower()
161
+ target = target.lower()
162
+
163
+ ed = Levenshtein.distance(pred, target)
164
+ total_edit_dis += ed
165
+ total_char_len += len(target)
166
+
167
+ dis = Levenshtein.normalized_distance(pred, target)
168
+ norm_edit_dis += dis
169
+ ratio_i = ratio[all_num] - 1 if ratio[
170
+ all_num] < self.max_ratio else self.max_ratio - 1
171
+ len_i = max(0, min(self.max_len, len(target)) - 1)
172
+ if pred == target:
173
+ correct_num += 1
174
+ each_len_correct_num[len_i] += 1
175
+ each_ratio_correct_num[ratio_i] += 1
176
+ each_len_num[len_i] += 1
177
+ each_len_norm_edit_dis[len_i] += dis
178
+
179
+ each_ratio_num[ratio_i] += 1
180
+ each_ratio_norm_edit_dis[ratio_i] += dis
181
+ all_num += 1
182
+ self.correct_num += correct_num
183
+ self.correct_num_real += correct_num_real
184
+ self.correct_num_lower += correct_num_lower
185
+ self.correct_num_ignore_space += correct_num_ignore_space
186
+ self.correct_num_ignore_space_lower += correct_num_ignore_space_lower
187
+ self.correct_num_ignore_space_symbol += correct_num_ignore_space_symbol
188
+ self.all_num += all_num
189
+ self.norm_edit_dis += norm_edit_dis
190
+ self.total_edit_dis += total_edit_dis
191
+ self.total_char_len += total_char_len
192
+ self.each_len_num = self.each_len_num + np.array(each_len_num)
193
+ self.each_len_correct_num = self.each_len_correct_num + np.array(
194
+ each_len_correct_num)
195
+ self.each_len_norm_edit_dis = self.each_len_norm_edit_dis + np.array(
196
+ each_len_norm_edit_dis)
197
+ self.each_ratio_num = self.each_ratio_num + np.array(each_ratio_num)
198
+ self.each_ratio_correct_num = self.each_ratio_correct_num + np.array(
199
+ each_ratio_correct_num)
200
+ self.each_ratio_norm_edit_dis = self.each_ratio_norm_edit_dis + np.array(
201
+ each_ratio_norm_edit_dis)
202
+ return {
203
+ 'acc': correct_num / (all_num + self.eps),
204
+ 'norm_edit_dis': 1 - norm_edit_dis / (all_num + self.eps),
205
+ 'cer': total_edit_dis / (total_char_len + self.eps),
206
+ }
207
+
208
+ def get_metric(self, training=False):
209
+ if self.with_ratio and not training:
210
+ return self.get_all_metric()
211
+ acc = 1.0 * self.correct_num / (self.all_num + self.eps)
212
+ norm_edit_dis = 1 - self.norm_edit_dis / (self.all_num + self.eps)
213
+ cer = self.total_edit_dis / (self.total_char_len + self.eps)
214
+ num_samples = self.all_num
215
+ self.reset()
216
+ return {
217
+ 'acc': acc,
218
+ 'norm_edit_dis': norm_edit_dis,
219
+ 'cer': cer,
220
+ 'num_samples': num_samples
221
+ }
222
+
223
+ def get_all_metric(self):
224
+ acc = 1.0 * self.correct_num / (self.all_num + self.eps)
225
+ acc_real = 1.0 * self.correct_num_real / (self.all_num + self.eps)
226
+ acc_lower = 1.0 * self.correct_num_lower / (self.all_num + self.eps)
227
+ acc_ignore_space = 1.0 * self.correct_num_ignore_space / (
228
+ self.all_num + self.eps)
229
+ acc_ignore_space_lower = 1.0 * self.correct_num_ignore_space_lower / (
230
+ self.all_num + self.eps)
231
+ acc_ignore_space_symbol = 1.0 * self.correct_num_ignore_space_symbol / (
232
+ self.all_num + self.eps)
233
+
234
+ norm_edit_dis = 1 - self.norm_edit_dis / (self.all_num + self.eps)
235
+ cer = self.total_edit_dis / (self.total_char_len + self.eps)
236
+ num_samples = self.all_num
237
+ each_len_acc = (self.each_len_correct_num /
238
+ (self.each_len_num + self.eps)).tolist()
239
+ each_len_norm_edit_dis = (1 -
240
+ ((self.each_len_norm_edit_dis) /
241
+ ((self.each_len_num) + self.eps))).tolist()
242
+ each_len_num = self.each_len_num.tolist()
243
+ each_ratio_acc = (self.each_ratio_correct_num /
244
+ (self.each_ratio_num + self.eps)).tolist()
245
+ each_ratio_norm_edit_dis = (1 - ((self.each_ratio_norm_edit_dis) / (
246
+ (self.each_ratio_num) + self.eps))).tolist()
247
+ each_ratio_num = self.each_ratio_num.tolist()
248
+ self.reset()
249
+ return {
250
+ 'acc': acc,
251
+ 'acc_real': acc_real,
252
+ 'acc_lower': acc_lower,
253
+ 'acc_ignore_space': acc_ignore_space,
254
+ 'acc_ignore_space_lower': acc_ignore_space_lower,
255
+ 'acc_ignore_space_symbol': acc_ignore_space_symbol,
256
+ 'acc_ignore_space_lower_symbol': acc,
257
+ 'each_len_num': each_len_num,
258
+ 'each_len_acc': each_len_acc,
259
+ 'each_len_norm_edit_dis': each_len_norm_edit_dis,
260
+ 'each_ratio_num': each_ratio_num,
261
+ 'each_ratio_acc': each_ratio_acc,
262
+ 'each_ratio_norm_edit_dis': each_ratio_norm_edit_dis,
263
+ 'norm_edit_dis': norm_edit_dis,
264
+ 'cer': cer,
265
+ 'num_samples': num_samples
266
+ }
267
+
268
+ def reset(self):
269
+ self.correct_num = 0
270
+ self.all_num = 0
271
+ self.norm_edit_dis = 0
272
+ self.total_edit_dis = 0.0
273
+ self.total_char_len = 0
274
+ self.correct_num_real = 0
275
+ self.correct_num_lower = 0
276
+ self.correct_num_ignore_space = 0
277
+ self.correct_num_ignore_space_lower = 0
278
+ self.correct_num_ignore_space_symbol = 0
279
+ self.each_len_num = np.array([0 for _ in range(self.max_len)])
280
+ self.each_len_correct_num = np.array([0 for _ in range(self.max_len)])
281
+ self.each_len_norm_edit_dis = np.array(
282
+ [0. for _ in range(self.max_len)])
283
+ self.each_ratio_num = np.array([0 for _ in range(self.max_ratio)])
284
+ self.each_ratio_correct_num = np.array(
285
+ [0 for _ in range(self.max_ratio)])
286
+ self.each_ratio_norm_edit_dis = np.array(
287
+ [0. for _ in range(self.max_ratio)])
@@ -0,0 +1,11 @@
1
+ import copy
2
+
3
+ from .base_recognizer import BaseRecognizer
4
+
5
+ __all__ = ['build_model']
6
+
7
+
8
+ def build_model(config):
9
+ config = copy.deepcopy(config)
10
+ rec_model = BaseRecognizer(config)
11
+ return rec_model