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.
Files changed (51) hide show
  1. deva_crnn/__init__.py +1 -0
  2. deva_crnn/augment.py +242 -0
  3. deva_crnn/charset.py +31 -0
  4. deva_crnn/data.py +48 -0
  5. deva_crnn/gate.py +97 -0
  6. deva_crnn/model.py +44 -0
  7. deva_crnn/predict.py +127 -0
  8. deva_crnn/train.py +254 -0
  9. veriscript/__init__.py +9 -0
  10. veriscript/__main__.py +4 -0
  11. veriscript/branding.py +26 -0
  12. veriscript/calibration.py +74 -0
  13. veriscript/cli.py +473 -0
  14. veriscript/core/__init__.py +1 -0
  15. veriscript/core/data.py +1586 -0
  16. veriscript/core/degradation.py +230 -0
  17. veriscript/core/metrics.py +591 -0
  18. veriscript/deva/__init__.py +1 -0
  19. veriscript/deva/reader.py +218 -0
  20. veriscript/document/__init__.py +1 -0
  21. veriscript/document/confusions.py +142 -0
  22. veriscript/document/corrections.py +378 -0
  23. veriscript/document/export.py +388 -0
  24. veriscript/document/layout.py +528 -0
  25. veriscript/document/memory.py +561 -0
  26. veriscript/document/ocr.py +1295 -0
  27. veriscript/document/orientation.py +247 -0
  28. veriscript/document/pipeline.py +367 -0
  29. veriscript/document/reconcile.py +225 -0
  30. veriscript/document/restore.py +222 -0
  31. veriscript/document/router.py +76 -0
  32. veriscript/document/verifier.py +132 -0
  33. veriscript/lexicon.py +105 -0
  34. veriscript/logging_setup.py +44 -0
  35. veriscript/paths.py +13 -0
  36. veriscript/photo/__init__.py +1 -0
  37. veriscript/photo/hybrid.py +736 -0
  38. veriscript/photo/rrdbnet.py +70 -0
  39. veriscript/photo/sr_engine.py +390 -0
  40. veriscript/photo/srvggnet.py +69 -0
  41. veriscript/photo/tv_refinement.py +142 -0
  42. veriscript/photo/upscaler.py +360 -0
  43. veriscript-1.5.0.data/data/calibration/rapidocr_devanagari_v1.json +60 -0
  44. veriscript-1.5.0.data/data/fonts/Mukta-Regular.ttf +0 -0
  45. veriscript-1.5.0.data/data/fonts/OFL.txt +93 -0
  46. veriscript-1.5.0.dist-info/METADATA +342 -0
  47. veriscript-1.5.0.dist-info/RECORD +51 -0
  48. veriscript-1.5.0.dist-info/WHEEL +5 -0
  49. veriscript-1.5.0.dist-info/entry_points.txt +3 -0
  50. veriscript-1.5.0.dist-info/licenses/LICENSE +21 -0
  51. 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()