veriscript 1.5.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.
- deva_crnn/__init__.py +1 -0
- deva_crnn/augment.py +242 -0
- deva_crnn/charset.py +31 -0
- deva_crnn/data.py +48 -0
- deva_crnn/gate.py +97 -0
- deva_crnn/model.py +44 -0
- deva_crnn/predict.py +127 -0
- deva_crnn/train.py +254 -0
- veriscript/__init__.py +9 -0
- veriscript/__main__.py +4 -0
- veriscript/branding.py +26 -0
- veriscript/calibration.py +74 -0
- veriscript/cli.py +473 -0
- veriscript/core/__init__.py +1 -0
- veriscript/core/data.py +1586 -0
- veriscript/core/degradation.py +230 -0
- veriscript/core/metrics.py +591 -0
- veriscript/deva/__init__.py +1 -0
- veriscript/deva/reader.py +218 -0
- veriscript/document/__init__.py +1 -0
- veriscript/document/confusions.py +142 -0
- veriscript/document/corrections.py +378 -0
- veriscript/document/export.py +388 -0
- veriscript/document/layout.py +528 -0
- veriscript/document/memory.py +561 -0
- veriscript/document/ocr.py +1295 -0
- veriscript/document/orientation.py +247 -0
- veriscript/document/pipeline.py +367 -0
- veriscript/document/reconcile.py +225 -0
- veriscript/document/restore.py +222 -0
- veriscript/document/router.py +76 -0
- veriscript/document/verifier.py +132 -0
- veriscript/lexicon.py +105 -0
- veriscript/logging_setup.py +44 -0
- veriscript/paths.py +13 -0
- veriscript/photo/__init__.py +1 -0
- veriscript/photo/hybrid.py +736 -0
- veriscript/photo/rrdbnet.py +70 -0
- veriscript/photo/sr_engine.py +390 -0
- veriscript/photo/srvggnet.py +69 -0
- veriscript/photo/tv_refinement.py +142 -0
- veriscript/photo/upscaler.py +360 -0
- veriscript-1.5.0.data/data/calibration/rapidocr_devanagari_v1.json +60 -0
- veriscript-1.5.0.data/data/fonts/Mukta-Regular.ttf +0 -0
- veriscript-1.5.0.data/data/fonts/OFL.txt +93 -0
- veriscript-1.5.0.dist-info/METADATA +342 -0
- veriscript-1.5.0.dist-info/RECORD +51 -0
- veriscript-1.5.0.dist-info/WHEEL +5 -0
- veriscript-1.5.0.dist-info/entry_points.txt +3 -0
- veriscript-1.5.0.dist-info/licenses/LICENSE +21 -0
- veriscript-1.5.0.dist-info/top_level.txt +2 -0
deva_crnn/__init__.py
ADDED
|
@@ -0,0 +1 @@
|
|
|
1
|
+
"""deva_crnn — small Devanagari line recognizer (CRNN+CTC) for the W1 fine-tune."""
|
deva_crnn/augment.py
ADDED
|
@@ -0,0 +1,242 @@
|
|
|
1
|
+
"""
|
|
2
|
+
augment.py — print-degradation augmentation for line crops (local use).
|
|
3
|
+
|
|
4
|
+
Kept out of `data.py` so the training path never imports OpenCV: the Vertex
|
|
5
|
+
PyTorch container has torch/PIL/numpy but not cv2.
|
|
6
|
+
|
|
7
|
+
`level="light"` is the attempt-1 augmentation (blur/noise/JPEG/brightness).
|
|
8
|
+
`level="heavy"` adds the artifacts that actually separate a rendered line from
|
|
9
|
+
a scanned one: ink spread/erosion, resolution loss (downscale-upscale), uneven
|
|
10
|
+
illumination, stronger blur/noise/JPEG and a slight rotation.
|
|
11
|
+
`level="xheavy"` is the stress set for few-shot real crops (N5 digit lines):
|
|
12
|
+
affine + perspective + elastic geometry, local shadow fields, gamma and faded
|
|
13
|
+
print, motion blur, salt-and-pepper and speckle noise, tighter crop jitter.
|
|
14
|
+
Existing levels are frozen; new work extends by adding a level.
|
|
15
|
+
"""
|
|
16
|
+
from __future__ import annotations
|
|
17
|
+
|
|
18
|
+
import cv2
|
|
19
|
+
import numpy as np
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def augment_line(img_bgr: np.ndarray, rng: np.random.Generator,
|
|
23
|
+
level: str = "light") -> np.ndarray:
|
|
24
|
+
"""Deterministic print-degradation augmentation (given `rng`)."""
|
|
25
|
+
if level == "xheavy":
|
|
26
|
+
return _xheavy(img_bgr, rng)
|
|
27
|
+
if level == "heavy":
|
|
28
|
+
return _heavy(img_bgr, rng)
|
|
29
|
+
return _light(img_bgr, rng)
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def _light(img_bgr: np.ndarray, rng: np.random.Generator) -> np.ndarray:
|
|
33
|
+
"""Cheap print-degradation augmentation (blur/noise/JPEG/brightness)."""
|
|
34
|
+
out = img_bgr.astype(np.float32)
|
|
35
|
+
if rng.random() < 0.5:
|
|
36
|
+
sigma = float(rng.uniform(0.4, 1.3))
|
|
37
|
+
out = cv2.GaussianBlur(out, (0, 0), sigmaX=sigma)
|
|
38
|
+
if rng.random() < 0.4:
|
|
39
|
+
out = out + rng.normal(0.0, float(rng.uniform(3, 12)), out.shape)
|
|
40
|
+
out = np.clip(out * float(rng.uniform(0.85, 1.15)) +
|
|
41
|
+
float(rng.uniform(-12, 12)), 0, 255).astype(np.uint8)
|
|
42
|
+
if rng.random() < 0.5:
|
|
43
|
+
q = int(rng.integers(40, 85))
|
|
44
|
+
ok, enc = cv2.imencode(".jpg", out, [cv2.IMWRITE_JPEG_QUALITY, q])
|
|
45
|
+
if ok:
|
|
46
|
+
out = cv2.imdecode(enc, cv2.IMREAD_COLOR)
|
|
47
|
+
return out
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def _heavy(img_bgr: np.ndarray, rng: np.random.Generator) -> np.ndarray:
|
|
51
|
+
out = img_bgr.astype(np.float32)
|
|
52
|
+
h, w = out.shape[:2]
|
|
53
|
+
|
|
54
|
+
# Ink spread (heavy inking / letterpress bleed) or erosion (thin print).
|
|
55
|
+
if rng.random() < 0.55:
|
|
56
|
+
k = np.ones((2, 2), np.uint8)
|
|
57
|
+
if rng.random() < 0.5:
|
|
58
|
+
out = cv2.dilate(out, k, iterations=1).astype(np.float32)
|
|
59
|
+
else:
|
|
60
|
+
out = cv2.erode(out, k, iterations=1).astype(np.float32)
|
|
61
|
+
|
|
62
|
+
# Resolution loss: a 300-dpi scan of small print is not crisp.
|
|
63
|
+
if rng.random() < 0.5:
|
|
64
|
+
f = float(rng.uniform(0.45, 0.8))
|
|
65
|
+
small = cv2.resize(out, (max(4, int(w * f)), max(4, int(h * f))),
|
|
66
|
+
interpolation=cv2.INTER_AREA)
|
|
67
|
+
out = cv2.resize(small, (w, h), interpolation=cv2.INTER_LINEAR)
|
|
68
|
+
|
|
69
|
+
# Anisotropic squeeze/stretch: detector boxes are tight around cells and
|
|
70
|
+
# long lines alike, so the recognizer must not assume a fixed aspect.
|
|
71
|
+
if rng.random() < 0.4:
|
|
72
|
+
fx = float(rng.uniform(0.65, 1.45))
|
|
73
|
+
out = cv2.resize(out, (max(4, int(w * fx)), h),
|
|
74
|
+
interpolation=cv2.INTER_LINEAR)
|
|
75
|
+
h, w = out.shape[:2]
|
|
76
|
+
|
|
77
|
+
# Uneven illumination (page curvature / phone shadow).
|
|
78
|
+
if rng.random() < 0.4:
|
|
79
|
+
ramp = np.linspace(float(rng.uniform(0.72, 0.88)),
|
|
80
|
+
float(rng.uniform(1.0, 1.12)), w,
|
|
81
|
+
dtype=np.float32)
|
|
82
|
+
out = out * ramp[None, :, None]
|
|
83
|
+
|
|
84
|
+
if rng.random() < 0.8:
|
|
85
|
+
sigma = float(rng.uniform(0.5, 2.2))
|
|
86
|
+
out = cv2.GaussianBlur(out, (0, 0), sigmaX=sigma)
|
|
87
|
+
if rng.random() < 0.7:
|
|
88
|
+
out = out + rng.normal(0.0, float(rng.uniform(4, 20)), out.shape)
|
|
89
|
+
out = np.clip(out * float(rng.uniform(0.7, 1.2)) +
|
|
90
|
+
float(rng.uniform(-20, 20)), 0, 255)
|
|
91
|
+
|
|
92
|
+
if rng.random() < 0.7:
|
|
93
|
+
q = int(rng.integers(28, 72))
|
|
94
|
+
ok, enc = cv2.imencode(".jpg", out.astype(np.uint8),
|
|
95
|
+
[cv2.IMWRITE_JPEG_QUALITY, q])
|
|
96
|
+
if ok:
|
|
97
|
+
out = cv2.imdecode(enc, cv2.IMREAD_COLOR).astype(np.float32)
|
|
98
|
+
|
|
99
|
+
if rng.random() < 0.5:
|
|
100
|
+
ang = float(rng.uniform(-1.5, 1.5))
|
|
101
|
+
m = cv2.getRotationMatrix2D((w / 2.0, h / 2.0), ang, 1.0)
|
|
102
|
+
out = cv2.warpAffine(out, m, (w, h), flags=cv2.INTER_LINEAR,
|
|
103
|
+
borderMode=cv2.BORDER_REPLICATE)
|
|
104
|
+
return np.clip(out, 0, 255).astype(np.uint8)
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
def _motion_blur(img: np.ndarray, rng: np.random.Generator) -> np.ndarray:
|
|
108
|
+
"""Random-angle linear motion blur (hand-held phone captures)."""
|
|
109
|
+
k = int(rng.integers(3, 8)) | 1
|
|
110
|
+
kernel = np.zeros((k, k), np.float32)
|
|
111
|
+
kernel[k // 2, :] = 1.0
|
|
112
|
+
angle = float(rng.uniform(0, 180))
|
|
113
|
+
m = cv2.getRotationMatrix2D((k / 2.0 - 0.5, k / 2.0 - 0.5), angle, 1.0)
|
|
114
|
+
kernel = cv2.warpAffine(kernel, m, (k, k))
|
|
115
|
+
total = float(kernel.sum())
|
|
116
|
+
if total <= 0:
|
|
117
|
+
return img
|
|
118
|
+
return cv2.filter2D(img, -1, kernel / total)
|
|
119
|
+
|
|
120
|
+
|
|
121
|
+
def _elastic(img: np.ndarray, rng: np.random.Generator, alpha: float,
|
|
122
|
+
sigma: float) -> np.ndarray:
|
|
123
|
+
"""Mild smooth displacement field (paper curl, lens distortion)."""
|
|
124
|
+
h, w = img.shape[:2]
|
|
125
|
+
dx = cv2.GaussianBlur(
|
|
126
|
+
rng.uniform(-1, 1, (h, w)).astype(np.float32), (0, 0),
|
|
127
|
+
sigmaX=sigma, sigmaY=sigma) * alpha
|
|
128
|
+
dy = cv2.GaussianBlur(
|
|
129
|
+
rng.uniform(-1, 1, (h, w)).astype(np.float32), (0, 0),
|
|
130
|
+
sigmaX=sigma, sigmaY=sigma) * alpha
|
|
131
|
+
xs = np.tile(np.arange(w, dtype=np.float32), (h, 1)) + dx
|
|
132
|
+
ys = np.tile(np.arange(h, dtype=np.float32)[:, None], (1, w)) + dy
|
|
133
|
+
return cv2.remap(img, xs, ys, cv2.INTER_LINEAR,
|
|
134
|
+
borderMode=cv2.BORDER_REPLICATE)
|
|
135
|
+
|
|
136
|
+
|
|
137
|
+
def _shadow_field(img: np.ndarray, rng: np.random.Generator) -> np.ndarray:
|
|
138
|
+
"""Smooth local darkening (a hand/phone shadow across the page)."""
|
|
139
|
+
h, w = img.shape[:2]
|
|
140
|
+
low = rng.uniform(0.55, 1.0, (4, 6)).astype(np.float32)
|
|
141
|
+
field = cv2.resize(low, (w, h), interpolation=cv2.INTER_CUBIC)
|
|
142
|
+
field = cv2.GaussianBlur(field, (0, 0), sigmaX=max(2.0, w / 40.0))
|
|
143
|
+
return img * field[:, :, None]
|
|
144
|
+
|
|
145
|
+
|
|
146
|
+
def _xheavy(img_bgr: np.ndarray, rng: np.random.Generator) -> np.ndarray:
|
|
147
|
+
"""Stress augmentation for few-shot real crops (N5 digit lines).
|
|
148
|
+
|
|
149
|
+
Chain: geometry (affine/perspective/elastic) -> ink -> resolution ->
|
|
150
|
+
illumination (shadow fields, gamma, fade) -> optics (motion/gaussian
|
|
151
|
+
blur) -> sensor (gaussian/salt-pepper/speckle) -> JPEG last. Every step
|
|
152
|
+
is applied with its own probability and bounded so the label is
|
|
153
|
+
preserved: no flips, rotation <= 2.5 deg, displacement <= ~3 px.
|
|
154
|
+
"""
|
|
155
|
+
out = img_bgr.astype(np.float32)
|
|
156
|
+
h, w = out.shape[:2]
|
|
157
|
+
|
|
158
|
+
# --- geometry ---
|
|
159
|
+
if rng.random() < 0.7:
|
|
160
|
+
ang = float(rng.uniform(-2.5, 2.5))
|
|
161
|
+
scale = float(rng.uniform(0.92, 1.08))
|
|
162
|
+
shear = float(rng.uniform(-0.035, 0.035))
|
|
163
|
+
m = cv2.getRotationMatrix2D((w / 2.0, h / 2.0), ang, scale)
|
|
164
|
+
m[0, 1] += shear
|
|
165
|
+
m[0, 2] += float(rng.uniform(-0.02, 0.02)) * w
|
|
166
|
+
m[1, 2] += float(rng.uniform(-0.02, 0.02)) * h
|
|
167
|
+
out = cv2.warpAffine(out, m, (w, h), flags=cv2.INTER_LINEAR,
|
|
168
|
+
borderMode=cv2.BORDER_REPLICATE)
|
|
169
|
+
if rng.random() < 0.3:
|
|
170
|
+
j = 0.015
|
|
171
|
+
src = np.float32([[0, 0], [w, 0], [w, h], [0, h]])
|
|
172
|
+
dst = src + (rng.uniform(-j, j, src.shape).astype(np.float32)
|
|
173
|
+
* np.float32([w, h]))
|
|
174
|
+
m = cv2.getPerspectiveTransform(src, dst)
|
|
175
|
+
out = cv2.warpPerspective(out, m, (w, h), flags=cv2.INTER_LINEAR,
|
|
176
|
+
borderMode=cv2.BORDER_REPLICATE)
|
|
177
|
+
if rng.random() < 0.3 and h >= 24 and w >= 40:
|
|
178
|
+
out = _elastic(out, rng, alpha=float(rng.uniform(1.5, 3.0)),
|
|
179
|
+
sigma=float(rng.uniform(6.0, 9.0)))
|
|
180
|
+
|
|
181
|
+
# --- ink ---
|
|
182
|
+
if rng.random() < 0.6:
|
|
183
|
+
k = (np.ones((3, 3), np.uint8) if rng.random() < 0.4
|
|
184
|
+
else np.ones((2, 2), np.uint8))
|
|
185
|
+
if rng.random() < 0.5:
|
|
186
|
+
out = cv2.dilate(out, k, iterations=1).astype(np.float32)
|
|
187
|
+
else:
|
|
188
|
+
out = cv2.erode(out, k, iterations=1).astype(np.float32)
|
|
189
|
+
|
|
190
|
+
# --- resolution ---
|
|
191
|
+
if rng.random() < 0.55:
|
|
192
|
+
f = float(rng.uniform(0.4, 0.75))
|
|
193
|
+
small = cv2.resize(out, (max(4, int(w * f)), max(4, int(h * f))),
|
|
194
|
+
interpolation=cv2.INTER_AREA)
|
|
195
|
+
out = cv2.resize(small, (w, h), interpolation=cv2.INTER_LINEAR)
|
|
196
|
+
if rng.random() < 0.45:
|
|
197
|
+
fx = float(rng.uniform(0.6, 1.5))
|
|
198
|
+
out = cv2.resize(out, (max(4, int(w * fx)), h),
|
|
199
|
+
interpolation=cv2.INTER_LINEAR)
|
|
200
|
+
h, w = out.shape[:2]
|
|
201
|
+
|
|
202
|
+
# --- illumination ---
|
|
203
|
+
if rng.random() < 0.5:
|
|
204
|
+
out = _shadow_field(out, rng)
|
|
205
|
+
if rng.random() < 0.4:
|
|
206
|
+
ramp = np.linspace(float(rng.uniform(0.7, 0.9)),
|
|
207
|
+
float(rng.uniform(1.0, 1.12)), w,
|
|
208
|
+
dtype=np.float32)
|
|
209
|
+
out = out * ramp[None, :, None]
|
|
210
|
+
if rng.random() < 0.4:
|
|
211
|
+
gamma = float(rng.uniform(0.6, 1.6))
|
|
212
|
+
out = 255.0 * np.power(np.clip(out, 0, 255) / 255.0, gamma)
|
|
213
|
+
if rng.random() < 0.25:
|
|
214
|
+
out = out * 0.55 + 245.0 * 0.45 # faded print toward paper tone
|
|
215
|
+
out = np.clip(out * float(rng.uniform(0.7, 1.2)) +
|
|
216
|
+
float(rng.uniform(-20, 20)), 0, 255)
|
|
217
|
+
|
|
218
|
+
# --- optics / sensor ---
|
|
219
|
+
if rng.random() < 0.3:
|
|
220
|
+
out = _motion_blur(out, rng)
|
|
221
|
+
if rng.random() < 0.8:
|
|
222
|
+
out = cv2.GaussianBlur(out, (0, 0),
|
|
223
|
+
sigmaX=float(rng.uniform(0.4, 2.6)))
|
|
224
|
+
if rng.random() < 0.75:
|
|
225
|
+
out = out + rng.normal(0.0, float(rng.uniform(3, 22)), out.shape)
|
|
226
|
+
if rng.random() < 0.25:
|
|
227
|
+
p = float(rng.uniform(0.002, 0.012))
|
|
228
|
+
mask = rng.random((out.shape[0], out.shape[1]))
|
|
229
|
+
dark = (mask < p * 0.5)[:, :, None]
|
|
230
|
+
bright = ((mask >= p * 0.5) & (mask < p))[:, :, None]
|
|
231
|
+
out = np.where(dark, 0.0, np.where(bright, 255.0, out))
|
|
232
|
+
if rng.random() < 0.2:
|
|
233
|
+
out = out * (1.0 + rng.normal(0.0, 0.06, out.shape))
|
|
234
|
+
out = np.clip(out, 0, 255).astype(np.uint8)
|
|
235
|
+
|
|
236
|
+
# --- compression last ---
|
|
237
|
+
if rng.random() < 0.75:
|
|
238
|
+
q = int(rng.integers(20, 66))
|
|
239
|
+
ok, enc = cv2.imencode(".jpg", out, [cv2.IMWRITE_JPEG_QUALITY, q])
|
|
240
|
+
if ok:
|
|
241
|
+
out = cv2.imdecode(enc, cv2.IMREAD_COLOR)
|
|
242
|
+
return out
|
deva_crnn/charset.py
ADDED
|
@@ -0,0 +1,31 @@
|
|
|
1
|
+
"""
|
|
2
|
+
charset.py — CTC charset for Devanagari line recognition.
|
|
3
|
+
|
|
4
|
+
Characters come from the training labels; the blank is index 0 (CTC). Encoding
|
|
5
|
+
is per-codepoint (no shaping - the image carries the shaping).
|
|
6
|
+
"""
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
from typing import Dict, Iterable, List, Tuple
|
|
10
|
+
|
|
11
|
+
BLANK = 0
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def build_charset(texts: Iterable[str]) -> List[str]:
|
|
15
|
+
chars = sorted({c for t in texts for c in t})
|
|
16
|
+
return ["\u0000"] + chars
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def encode(text: str, charset: List[str]) -> List[int]:
|
|
20
|
+
idx = {c: i for i, c in enumerate(charset)}
|
|
21
|
+
return [idx[c] for c in text if c in idx]
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def decode(ids: Iterable[int], charset: List[str]) -> str:
|
|
25
|
+
out: List[str] = []
|
|
26
|
+
prev = None
|
|
27
|
+
for i in ids:
|
|
28
|
+
if i != prev and i != BLANK:
|
|
29
|
+
out.append(charset[i])
|
|
30
|
+
prev = i
|
|
31
|
+
return "".join(out)
|
deva_crnn/data.py
ADDED
|
@@ -0,0 +1,48 @@
|
|
|
1
|
+
"""
|
|
2
|
+
data.py — line-image dataset for the CRNN: npz export + torch Dataset.
|
|
3
|
+
|
|
4
|
+
The npz format (`images` uint8 NxHxW, `texts` unicode array) is compact enough
|
|
5
|
+
to ship inside the Vertex AI python package and fast to load.
|
|
6
|
+
"""
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
import os
|
|
10
|
+
from typing import Dict, List, Tuple
|
|
11
|
+
|
|
12
|
+
import numpy as np
|
|
13
|
+
from PIL import Image
|
|
14
|
+
|
|
15
|
+
IN_H = 32
|
|
16
|
+
IN_W = 256
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def normalize_line(img_bgr: np.ndarray, h: int = IN_H, w: int = IN_W) -> np.ndarray:
|
|
20
|
+
"""Grayscale, height-normalized, right-padded line image (uint8).
|
|
21
|
+
|
|
22
|
+
PIL-only on purpose: the Vertex training container has no OpenCV.
|
|
23
|
+
"""
|
|
24
|
+
if img_bgr.ndim == 3:
|
|
25
|
+
gray = np.asarray(Image.fromarray(img_bgr[:, :, ::-1]).convert("L"))
|
|
26
|
+
else:
|
|
27
|
+
gray = img_bgr
|
|
28
|
+
scale = h / gray.shape[0]
|
|
29
|
+
new_w = min(w, max(4, int(round(gray.shape[1] * scale))))
|
|
30
|
+
resized = np.asarray(Image.fromarray(gray).resize((new_w, h),
|
|
31
|
+
Image.BILINEAR))
|
|
32
|
+
out = np.full((h, w), 255, dtype=np.uint8)
|
|
33
|
+
out[:, :new_w] = resized
|
|
34
|
+
return out
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
def export_npz(images: List[np.ndarray], texts: List[str], path: str,
|
|
38
|
+
h: int = IN_H, w: int = IN_W) -> str:
|
|
39
|
+
os.makedirs(os.path.dirname(os.path.abspath(path)), exist_ok=True)
|
|
40
|
+
arr = np.stack([normalize_line(im, h=h, w=w) for im in images])
|
|
41
|
+
np.savez_compressed(path, images=arr,
|
|
42
|
+
texts=np.array(texts, dtype=object).astype(str))
|
|
43
|
+
return path
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
def load_npz(path: str) -> Tuple[np.ndarray, List[str]]:
|
|
47
|
+
d = np.load(path, allow_pickle=False)
|
|
48
|
+
return d["images"], [str(t) for t in d["texts"]]
|
deva_crnn/gate.py
ADDED
|
@@ -0,0 +1,97 @@
|
|
|
1
|
+
"""
|
|
2
|
+
gate.py — W1 adopt-if gate: line-level digit-exact + bagCER on frozen sets.
|
|
3
|
+
|
|
4
|
+
Gate (pre-registered): digit-exact >= 0.75 on digit-bearing lines, bagCER not
|
|
5
|
+
worse than the RapidOCR line baseline, <= +1 s/page. Sets:
|
|
6
|
+
* heiDATA frozen pages -> ALTO line crops (human GT)
|
|
7
|
+
* nepali_pdf_v2 anchor pages -> RapidOCR line boxes, Gemini anchor GT
|
|
8
|
+
"""
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
import json
|
|
12
|
+
import os
|
|
13
|
+
from typing import Dict, List, Optional, Tuple
|
|
14
|
+
|
|
15
|
+
import numpy as np
|
|
16
|
+
|
|
17
|
+
from .data import normalize_line # noqa: E402
|
|
18
|
+
from veriscript.core.metrics import cer, cer_bag, digit_string, digit_tokens # noqa: E402
|
|
19
|
+
|
|
20
|
+
DIGIT_EXACT_BAR = 0.75
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def line_stats(gts: List[str], hyps: List[str]) -> Dict:
|
|
24
|
+
"""Digit-exact rate over digit-bearing lines + mean CER/bagCER."""
|
|
25
|
+
pairs = [(g, h) for g, h in zip(gts, hyps) if g.strip()]
|
|
26
|
+
digit_pairs = [(g, h) for g, h in pairs if digit_tokens(g)]
|
|
27
|
+
exact = [1.0 if digit_string(g) == digit_string(h) else 0.0
|
|
28
|
+
for g, h in digit_pairs]
|
|
29
|
+
return {
|
|
30
|
+
"lines": len(pairs),
|
|
31
|
+
"digit_lines": len(digit_pairs),
|
|
32
|
+
"digit_exact": float(np.mean(exact)) if exact else None,
|
|
33
|
+
"cer": float(np.mean([cer(g, h) for g, h in pairs])) if pairs else None,
|
|
34
|
+
"bagcer": (float(np.mean([cer_bag(g, h) for g, h in pairs]))
|
|
35
|
+
if pairs else None),
|
|
36
|
+
}
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def line_values(gts: List[str], hyps: List[str]) -> Dict[str, List[float]]:
|
|
40
|
+
"""Per-line values for bootstrap CIs (digit-exact 0/1 + CER)."""
|
|
41
|
+
pairs = [(g, h) for g, h in zip(gts, hyps) if g.strip()]
|
|
42
|
+
digit_pairs = [(g, h) for g, h in pairs if digit_tokens(g)]
|
|
43
|
+
return {
|
|
44
|
+
"digit_exact": [1.0 if digit_string(g) == digit_string(h) else 0.0
|
|
45
|
+
for g, h in digit_pairs],
|
|
46
|
+
"cer": [cer(g, h) for g, h in pairs],
|
|
47
|
+
"bagcer": [cer_bag(g, h) for g, h in pairs],
|
|
48
|
+
}
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def heidata_line_crops(data_dir: str, limit: int = 0) -> Tuple[List[str], List[np.ndarray]]:
|
|
52
|
+
"""(texts, crops) from frozen heiDATA pages + ALTO line boxes."""
|
|
53
|
+
from veriscript.core import data as doc_data
|
|
54
|
+
manifest = doc_data.load_dataset(data_dir)
|
|
55
|
+
texts: List[str] = []
|
|
56
|
+
crops: List[np.ndarray] = []
|
|
57
|
+
for e in manifest["entries"]:
|
|
58
|
+
img = doc_data.imread_safe(e["_degraded_path"])
|
|
59
|
+
if img is None:
|
|
60
|
+
continue
|
|
61
|
+
boxes = json.load(open(e["_boxes_path"], encoding="utf-8"))
|
|
62
|
+
for b in boxes:
|
|
63
|
+
x0, y0, x1, y1 = b["bbox"]
|
|
64
|
+
pad = 4
|
|
65
|
+
crop = img[max(0, y0 - pad):y1 + pad, max(0, x0 - pad):x1 + pad]
|
|
66
|
+
if crop.size == 0:
|
|
67
|
+
continue
|
|
68
|
+
texts.append(b["text"])
|
|
69
|
+
crops.append(crop)
|
|
70
|
+
if limit and len(crops) >= limit:
|
|
71
|
+
break
|
|
72
|
+
return texts, crops
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def v2_anchor_pages(data_dir: str, readings_path: str,
|
|
76
|
+
limit: int = 0) -> List[Dict]:
|
|
77
|
+
"""Per anchor page: RapidOCR line boxes (image + crops) + Gemini GT."""
|
|
78
|
+
from veriscript.core import data as doc_data
|
|
79
|
+
from veriscript.document.ocr import ocr_page
|
|
80
|
+
manifest = doc_data.load_dataset(data_dir)
|
|
81
|
+
readings = {r["page"]: r["readings"]["gemini"]
|
|
82
|
+
for r in json.load(open(readings_path, encoding="utf-8"))["rows"]}
|
|
83
|
+
pages: List[Dict] = []
|
|
84
|
+
for e in manifest["entries"]:
|
|
85
|
+
img = doc_data.imread_safe(e["_degraded_path"])
|
|
86
|
+
if img is None or e["id"] not in readings:
|
|
87
|
+
continue
|
|
88
|
+
rec = ocr_page(img, backend="rapidocr", lang="ne")
|
|
89
|
+
crops = []
|
|
90
|
+
for t in rec.tokens:
|
|
91
|
+
x0, y0, x1, y1 = t.bbox
|
|
92
|
+
crops.append(img[max(0, y0 - 3):y1 + 3, max(0, x0 - 3):x1 + 3])
|
|
93
|
+
pages.append({"page": e["id"], "gt": readings[e["id"]],
|
|
94
|
+
"boxes": [t.bbox for t in rec.tokens], "crops": crops})
|
|
95
|
+
if limit and len(pages) >= limit:
|
|
96
|
+
break
|
|
97
|
+
return pages
|
deva_crnn/model.py
ADDED
|
@@ -0,0 +1,44 @@
|
|
|
1
|
+
"""
|
|
2
|
+
model.py — CRNN (CNN + BiLSTM + CTC) for Devanagari line crops.
|
|
3
|
+
|
|
4
|
+
Input: grayscale 32 x W (padded to 256), output: log-probabilities over the
|
|
5
|
+
charset per timestep. Small on purpose: trains on a single T4 in minutes and
|
|
6
|
+
runs on CPU for inference.
|
|
7
|
+
"""
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import torch
|
|
11
|
+
import torch.nn as nn
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class CRNN(nn.Module):
|
|
15
|
+
def __init__(self, n_classes: int, hidden: int = 256, in_h: int = 32):
|
|
16
|
+
super().__init__()
|
|
17
|
+
# GroupNorm, not BatchNorm: small-batch CTC training diverged to NaN
|
|
18
|
+
# with BN (measured), GroupNorm is stable at batch 8 and 64 alike.
|
|
19
|
+
def block(cin, cout):
|
|
20
|
+
return [nn.Conv2d(cin, cout, 3, padding=1),
|
|
21
|
+
nn.GroupNorm(8, cout), nn.ReLU()]
|
|
22
|
+
|
|
23
|
+
# The stack must collapse the input height to exactly 1: 32 and 48 need
|
|
24
|
+
# a final (2,1) pool, 64 needs (4,1) (64 -> 32 -> 16 -> 8 -> 4 -> 1).
|
|
25
|
+
last_pool = (4, 1) if in_h >= 64 else (2, 1)
|
|
26
|
+
self.cnn = nn.Sequential(
|
|
27
|
+
*block(1, 32), nn.MaxPool2d(2), # 16 x W/2
|
|
28
|
+
*block(32, 64), nn.MaxPool2d(2), # 8 x W/4
|
|
29
|
+
*block(64, 128), nn.MaxPool2d((2, 1)), # 4 x W/4
|
|
30
|
+
*block(128, 128), nn.MaxPool2d((2, 1)), # 2 x W/4
|
|
31
|
+
*block(128, 256), nn.MaxPool2d(last_pool), # 1 x W/4
|
|
32
|
+
)
|
|
33
|
+
self.rnn = nn.LSTM(256, hidden, num_layers=2, bidirectional=True,
|
|
34
|
+
batch_first=True)
|
|
35
|
+
self.head = nn.Linear(2 * hidden, n_classes)
|
|
36
|
+
|
|
37
|
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
38
|
+
"""x: (B, 1, H, W) -> log-probs (T, B, C) for CTCLoss."""
|
|
39
|
+
f = self.cnn(x) # B, 256, 1, W/4
|
|
40
|
+
b, c, h, w = f.shape
|
|
41
|
+
f = f.squeeze(2).permute(0, 2, 1) # B, W/4, 256
|
|
42
|
+
out, _ = self.rnn(f)
|
|
43
|
+
logits = self.head(out) # B, T, C
|
|
44
|
+
return logits.log_softmax(-1).permute(1, 0, 2)
|
deva_crnn/predict.py
ADDED
|
@@ -0,0 +1,127 @@
|
|
|
1
|
+
"""
|
|
2
|
+
predict.py — run a trained deva_crnn checkpoint on line crops.
|
|
3
|
+
|
|
4
|
+
Usage:
|
|
5
|
+
python -m deva_crnn.predict --ckpt out/deva_crnn/ckpt.pt \
|
|
6
|
+
--images-dir data/doc_eval/heidata_lines/pages --out texts.tsv
|
|
7
|
+
"""
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import argparse
|
|
11
|
+
import os
|
|
12
|
+
import sys
|
|
13
|
+
from typing import List, Tuple
|
|
14
|
+
|
|
15
|
+
import cv2
|
|
16
|
+
import numpy as np
|
|
17
|
+
import torch
|
|
18
|
+
|
|
19
|
+
BASE_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
|
20
|
+
if BASE_DIR not in sys.path:
|
|
21
|
+
sys.path.insert(0, BASE_DIR)
|
|
22
|
+
|
|
23
|
+
from deva_crnn.charset import decode
|
|
24
|
+
from deva_crnn.data import normalize_line
|
|
25
|
+
from deva_crnn.model import CRNN
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
def load_model(ckpt_path: str, device: str = "cpu"):
|
|
29
|
+
ckpt = torch.load(ckpt_path, map_location=device, weights_only=False)
|
|
30
|
+
in_h = int(ckpt.get("in_h", 32))
|
|
31
|
+
in_w = int(ckpt.get("in_w", 256))
|
|
32
|
+
model = CRNN(n_classes=len(ckpt["charset"]), in_h=in_h)
|
|
33
|
+
model.load_state_dict(ckpt["model"])
|
|
34
|
+
model.eval()
|
|
35
|
+
model.to(device)
|
|
36
|
+
model.in_h = in_h # recognized by recognize_lines for normalization
|
|
37
|
+
model.in_w = in_w
|
|
38
|
+
return model, ckpt["charset"]
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def beam_search_decode(log_probs: np.ndarray, charset: List[str],
|
|
42
|
+
beam_width: int = 8) -> str:
|
|
43
|
+
"""CTC prefix beam search over (T, C) log-probabilities.
|
|
44
|
+
|
|
45
|
+
Greedy argmax commits to one alignment per frame; on noisy letterpress
|
|
46
|
+
digits the second-best prefix is often the right one, and beam search
|
|
47
|
+
keeps those alternatives. `charset[0]` is the CTC blank.
|
|
48
|
+
"""
|
|
49
|
+
beams = {(): (0.0, -np.inf)} # prefix -> (p_blank, p_nonblank) log-probs
|
|
50
|
+
for t in range(log_probs.shape[0]):
|
|
51
|
+
lp = log_probs[t]
|
|
52
|
+
nxt: dict = {}
|
|
53
|
+
for prefix, (pb, pnb) in beams.items():
|
|
54
|
+
p_total = np.logaddexp(pb, pnb)
|
|
55
|
+
for c in range(lp.shape[0]):
|
|
56
|
+
p = lp[c]
|
|
57
|
+
if c == 0: # blank: prefix stays
|
|
58
|
+
cur = nxt.get(prefix, (-np.inf, -np.inf))
|
|
59
|
+
nxt[prefix] = (np.logaddexp(cur[0], p_total + p), cur[1])
|
|
60
|
+
continue
|
|
61
|
+
last = prefix[-1] if prefix else None
|
|
62
|
+
if c == last: # repeat: only extend via a blank in between
|
|
63
|
+
cur = nxt.get(prefix, (-np.inf, -np.inf))
|
|
64
|
+
nxt[prefix] = (cur[0], np.logaddexp(cur[1], pnb + p))
|
|
65
|
+
new_prefix = prefix + (c,)
|
|
66
|
+
cur = nxt.get(new_prefix, (-np.inf, -np.inf))
|
|
67
|
+
nxt[new_prefix] = (cur[0],
|
|
68
|
+
np.logaddexp(cur[1], pb + p))
|
|
69
|
+
else:
|
|
70
|
+
new_prefix = prefix + (c,)
|
|
71
|
+
cur = nxt.get(new_prefix, (-np.inf, -np.inf))
|
|
72
|
+
nxt[new_prefix] = (cur[0],
|
|
73
|
+
np.logaddexp(cur[1], p_total + p))
|
|
74
|
+
beams = dict(sorted(nxt.items(),
|
|
75
|
+
key=lambda kv: -np.logaddexp(kv[1][0], kv[1][1])
|
|
76
|
+
)[:beam_width])
|
|
77
|
+
best = max(beams.items(), key=lambda kv: np.logaddexp(kv[1][0], kv[1][1]))
|
|
78
|
+
return "".join(charset[c] for c in best[0])
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
def recognize_lines(model, charset: List[str], images_bgr: List[np.ndarray],
|
|
82
|
+
device: str = "cpu", decode_mode: str = "greedy",
|
|
83
|
+
beam_width: int = 8) -> List[str]:
|
|
84
|
+
if not images_bgr:
|
|
85
|
+
return []
|
|
86
|
+
in_h = int(getattr(model, "in_h", 32))
|
|
87
|
+
in_w = int(getattr(model, "in_w", 256))
|
|
88
|
+
batch = np.stack([normalize_line(im, h=in_h, w=in_w)
|
|
89
|
+
for im in images_bgr]).astype(np.float32)
|
|
90
|
+
x = torch.from_numpy(batch / 255.0).unsqueeze(1)
|
|
91
|
+
x = ((x - 0.5) / 0.5).to(device)
|
|
92
|
+
with torch.no_grad():
|
|
93
|
+
logits = model(x)
|
|
94
|
+
if decode_mode == "beam":
|
|
95
|
+
log_probs = logits.log_softmax(-1).cpu().numpy() # (T, B, C)
|
|
96
|
+
return [beam_search_decode(log_probs[:, b, :], charset,
|
|
97
|
+
beam_width=beam_width)
|
|
98
|
+
for b in range(log_probs.shape[1])]
|
|
99
|
+
preds = logits.argmax(-1).permute(1, 0)
|
|
100
|
+
return [decode(p.tolist(), charset) for p in preds]
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
def main():
|
|
104
|
+
ap = argparse.ArgumentParser(description="deva_crnn line inference")
|
|
105
|
+
ap.add_argument("--ckpt", required=True)
|
|
106
|
+
ap.add_argument("--images-dir", required=True)
|
|
107
|
+
ap.add_argument("--out", required=True)
|
|
108
|
+
ap.add_argument("--device", default="cpu")
|
|
109
|
+
ap.add_argument("--decode", choices=("greedy", "beam"), default="greedy")
|
|
110
|
+
ap.add_argument("--beam-width", type=int, default=8)
|
|
111
|
+
args = ap.parse_args()
|
|
112
|
+
|
|
113
|
+
model, charset = load_model(args.ckpt, args.device)
|
|
114
|
+
names = sorted(f for f in os.listdir(args.images_dir)
|
|
115
|
+
if f.lower().endswith((".png", ".jpg", ".jpeg")))
|
|
116
|
+
imgs = [cv2.imread(os.path.join(args.images_dir, n)) for n in names]
|
|
117
|
+
keep = [(n, im) for n, im in zip(names, imgs) if im is not None]
|
|
118
|
+
texts = recognize_lines(model, charset, [im for _, im in keep], args.device,
|
|
119
|
+
decode_mode=args.decode, beam_width=args.beam_width)
|
|
120
|
+
with open(args.out, "w", encoding="utf-8") as f:
|
|
121
|
+
for (n, _), t in zip(keep, texts):
|
|
122
|
+
f.write(f"{n}\t{t}\n")
|
|
123
|
+
print(f"wrote {args.out} ({len(keep)} lines)")
|
|
124
|
+
|
|
125
|
+
|
|
126
|
+
if __name__ == "__main__":
|
|
127
|
+
main()
|