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 +5 -0
- litealpr/det/__init__.py +1 -0
- litealpr/license_plate_dict.txt +37 -0
- litealpr/pipeline.py +187 -0
- litealpr/rec/losses/__init__.py +71 -0
- litealpr/rec/losses/ctc_loss.py +33 -0
- litealpr/rec/metrics/__init__.py +33 -0
- litealpr/rec/metrics/rec_metric.py +287 -0
- litealpr/rec/modeling/__init__.py +11 -0
- litealpr/rec/modeling/base_recognizer.py +57 -0
- litealpr/rec/modeling/common.py +242 -0
- litealpr/rec/modeling/decoders/__init__.py +22 -0
- litealpr/rec/modeling/decoders/efficient_rctc_decoder.py +56 -0
- litealpr/rec/modeling/encoders/__init__.py +22 -0
- litealpr/rec/modeling/encoders/svtrv2.py +495 -0
- litealpr/rec/modeling/encoders/svtrv2_lnconv.py +524 -0
- litealpr/rec/modeling/encoders/svtrv2_lnconv_two33.py +542 -0
- litealpr/rec/optimizer/__init__.py +75 -0
- litealpr/rec/optimizer/lr.py +276 -0
- litealpr/rec/postprocess/__init__.py +26 -0
- litealpr/rec/postprocess/ctc_postprocess.py +118 -0
- litealpr/rec/preprocess/__init__.py +78 -0
- litealpr/rec/preprocess/auto_augment.py +493 -0
- litealpr/rec/preprocess/ctc_label_encode.py +123 -0
- litealpr/rec/preprocess/parseq_aug.py +156 -0
- litealpr/rec/preprocess/rec_aug.py +12 -0
- litealpr-0.1.0.dist-info/METADATA +261 -0
- litealpr-0.1.0.dist-info/RECORD +31 -0
- litealpr-0.1.0.dist-info/WHEEL +5 -0
- litealpr-0.1.0.dist-info/licenses/LICENSE +201 -0
- litealpr-0.1.0.dist-info/top_level.txt +1 -0
litealpr/__init__.py
ADDED
litealpr/det/__init__.py
ADDED
|
@@ -0,0 +1 @@
|
|
|
1
|
+
# det package
|
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)])
|