bobine 0.2.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.
- bobine/__init__.py +84 -0
- bobine/_vendor/rapid_latex_ocr/__init__.py +9 -0
- bobine/_vendor/rapid_latex_ocr/main.py +214 -0
- bobine/_vendor/rapid_latex_ocr/models.py +154 -0
- bobine/_vendor/rapid_latex_ocr/utils.py +185 -0
- bobine/_vendor/rapid_latex_ocr/utils_load.py +172 -0
- bobine/assets.py +90 -0
- bobine/config.py +94 -0
- bobine/converter.py +829 -0
- bobine/documents.py +180 -0
- bobine/engine.py +239 -0
- bobine/markdown.py +103 -0
- bobine/pipeline.py +278 -0
- bobine/tables.py +87 -0
- bobine/versions.py +148 -0
- bobine-0.2.0.dist-info/METADATA +227 -0
- bobine-0.2.0.dist-info/RECORD +22 -0
- bobine-0.2.0.dist-info/WHEEL +5 -0
- bobine-0.2.0.dist-info/licenses/LICENSE +6 -0
- bobine-0.2.0.dist-info/licenses/LICENSES/Apache-2.0.txt +105 -0
- bobine-0.2.0.dist-info/licenses/LICENSES/MIT.txt +21 -0
- bobine-0.2.0.dist-info/top_level.txt +1 -0
bobine/__init__.py
ADDED
|
@@ -0,0 +1,84 @@
|
|
|
1
|
+
"""bobine — PDF / Office / text → Markdown ingestion engine.
|
|
2
|
+
|
|
3
|
+
Standalone module extracted from the OKFgraph project. Runs on a single
|
|
4
|
+
``onnxruntime`` wheel with no CUDA-version coupling (RapidAI family + pdf_oxide),
|
|
5
|
+
and needs nothing but the stdlib for plain text / markdown documents.
|
|
6
|
+
|
|
7
|
+
Public API
|
|
8
|
+
----------
|
|
9
|
+
Conversion (PDF/Office)
|
|
10
|
+
- ``ConverterConfig`` / ``RoutingMode`` — configuration and routing modes.
|
|
11
|
+
- ``HybridConverter`` — core conversion pipeline (Qt-independent).
|
|
12
|
+
- ``OnnxRapidEngine`` — lazy ONNX model manager.
|
|
13
|
+
- ``html_tables_to_gfm`` — HTML table → GFM pipe-table converter.
|
|
14
|
+
|
|
15
|
+
Images
|
|
16
|
+
- ``stage_images_as_okf_assets`` — okf-asset:// staging for extracted images.
|
|
17
|
+
- ``asset_id`` / ``ASSET_STORE_DIRNAME`` — asset naming convention.
|
|
18
|
+
|
|
19
|
+
Text-type documents
|
|
20
|
+
- ``Document`` — normalized document model (id, title, body, tags, …).
|
|
21
|
+
- ``load_markdown_document`` — frontmatter-aware .md loading.
|
|
22
|
+
- ``wrap_thoughts`` — raw reasoning text → OKF-compliant markdown.
|
|
23
|
+
- ``lint_markdown`` / ``lint_markdown_file`` — mordant linting (guarded).
|
|
24
|
+
|
|
25
|
+
Orchestration
|
|
26
|
+
- ``convert_to_markdown`` — dispatch PDF / Office / text → markdown string.
|
|
27
|
+
- ``stage_images`` — collect extracted images into an asset store.
|
|
28
|
+
- ``ingest_document`` — full pipeline: convert → write .md → stage assets
|
|
29
|
+
→ lint, returns a ``ConvertedDocument``.
|
|
30
|
+
|
|
31
|
+
Versioning
|
|
32
|
+
- ``check_rapid_versions`` — runtime version check for RapidAI packages.
|
|
33
|
+
"""
|
|
34
|
+
|
|
35
|
+
from bobine.assets import (
|
|
36
|
+
ASSET_STORE_DIRNAME,
|
|
37
|
+
asset_id,
|
|
38
|
+
stage_images_as_okf_assets,
|
|
39
|
+
)
|
|
40
|
+
from bobine.config import ConverterConfig, RoutingMode
|
|
41
|
+
from bobine.converter import HybridConverter
|
|
42
|
+
from bobine.documents import Document, load_markdown_document, wrap_thoughts
|
|
43
|
+
from bobine.engine import OnnxRapidEngine
|
|
44
|
+
from bobine.markdown import lint_markdown, lint_markdown_file
|
|
45
|
+
from bobine.pipeline import (
|
|
46
|
+
OFFICE_EXTS,
|
|
47
|
+
SUPPORTED_EXTENSIONS,
|
|
48
|
+
TEXT_EXTS,
|
|
49
|
+
ConvertedDocument,
|
|
50
|
+
convert_directory,
|
|
51
|
+
convert_to_markdown,
|
|
52
|
+
ingest_document,
|
|
53
|
+
stage_images,
|
|
54
|
+
)
|
|
55
|
+
from bobine.tables import html_tables_to_gfm
|
|
56
|
+
from bobine.versions import check_rapid_versions
|
|
57
|
+
|
|
58
|
+
__version__ = "0.2.0"
|
|
59
|
+
|
|
60
|
+
__all__ = [
|
|
61
|
+
"ASSET_STORE_DIRNAME",
|
|
62
|
+
"OFFICE_EXTS",
|
|
63
|
+
"SUPPORTED_EXTENSIONS",
|
|
64
|
+
"TEXT_EXTS",
|
|
65
|
+
"ConvertedDocument",
|
|
66
|
+
"ConverterConfig",
|
|
67
|
+
"Document",
|
|
68
|
+
"HybridConverter",
|
|
69
|
+
"OnnxRapidEngine",
|
|
70
|
+
"RoutingMode",
|
|
71
|
+
"__version__",
|
|
72
|
+
"asset_id",
|
|
73
|
+
"check_rapid_versions",
|
|
74
|
+
"convert_directory",
|
|
75
|
+
"convert_to_markdown",
|
|
76
|
+
"html_tables_to_gfm",
|
|
77
|
+
"ingest_document",
|
|
78
|
+
"lint_markdown",
|
|
79
|
+
"lint_markdown_file",
|
|
80
|
+
"load_markdown_document",
|
|
81
|
+
"stage_images",
|
|
82
|
+
"stage_images_as_okf_assets",
|
|
83
|
+
"wrap_thoughts",
|
|
84
|
+
]
|
|
@@ -0,0 +1,9 @@
|
|
|
1
|
+
# -*- encoding: utf-8 -*-
|
|
2
|
+
# @Author: SWHL (vendored into bobine, MIT licensed)
|
|
3
|
+
from .main import LaTeXOCR
|
|
4
|
+
|
|
5
|
+
# Legacy alias: older builds shipped the class as `LatexOCR`; keep both so
|
|
6
|
+
# bobine's engine import works regardless of which build was vendored.
|
|
7
|
+
LatexOCR = LaTeXOCR
|
|
8
|
+
|
|
9
|
+
__all__ = ["LaTeXOCR", "LatexOCR"]
|
|
@@ -0,0 +1,214 @@
|
|
|
1
|
+
# -*- encoding: utf-8 -*-
|
|
2
|
+
# @Author: SWHL
|
|
3
|
+
# @Contact: liekkaskono@163.com
|
|
4
|
+
import argparse
|
|
5
|
+
import re
|
|
6
|
+
import time
|
|
7
|
+
import traceback
|
|
8
|
+
from dataclasses import dataclass
|
|
9
|
+
from pathlib import Path
|
|
10
|
+
from typing import Tuple, Union
|
|
11
|
+
|
|
12
|
+
import numpy as np
|
|
13
|
+
import yaml
|
|
14
|
+
from PIL import Image
|
|
15
|
+
|
|
16
|
+
from .models import EncoderDecoder
|
|
17
|
+
from .utils import DownloadModel, PreProcess, TokenizerCls, get_file_encode
|
|
18
|
+
from .utils_load import InputType, LoadImage, LoadImageError, OrtInferSession
|
|
19
|
+
|
|
20
|
+
cur_dir = Path(__file__).resolve().parent
|
|
21
|
+
DEFAULT_CONFIG = cur_dir / "config.yaml"
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
@dataclass
|
|
25
|
+
class LaTeXOCRInput:
|
|
26
|
+
max_width: int = 672
|
|
27
|
+
max_height: int = 192
|
|
28
|
+
min_height: int = 32
|
|
29
|
+
min_width: int = 32
|
|
30
|
+
bos_token: int = 1
|
|
31
|
+
max_seq_len: int = 512
|
|
32
|
+
eos_token: int = 2
|
|
33
|
+
temperature: float = 0.00001
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
class LaTeXOCR:
|
|
37
|
+
def __init__(
|
|
38
|
+
self,
|
|
39
|
+
config_path: Union[str, Path] = DEFAULT_CONFIG,
|
|
40
|
+
image_resizer_path: Union[str, Path] = None,
|
|
41
|
+
encoder_path: Union[str, Path] = None,
|
|
42
|
+
decoder_path: Union[str, Path] = None,
|
|
43
|
+
tokenizer_json: Union[str, Path] = None,
|
|
44
|
+
):
|
|
45
|
+
self.image_resizer_path = image_resizer_path
|
|
46
|
+
self.encoder_path = encoder_path
|
|
47
|
+
self.decoder_path = decoder_path
|
|
48
|
+
self.tokenizer_json = tokenizer_json
|
|
49
|
+
|
|
50
|
+
self.get_model_path()
|
|
51
|
+
|
|
52
|
+
file_encode = get_file_encode(config_path)
|
|
53
|
+
with open(config_path, "r", encoding=file_encode) as f:
|
|
54
|
+
args = yaml.load(f, Loader=yaml.FullLoader)
|
|
55
|
+
input_params = LaTeXOCRInput(**args)
|
|
56
|
+
|
|
57
|
+
self.max_dims = [input_params.max_width, input_params.max_height]
|
|
58
|
+
self.min_dims = [input_params.min_width, input_params.min_height]
|
|
59
|
+
self.temperature = input_params.temperature
|
|
60
|
+
|
|
61
|
+
self.load_img = LoadImage()
|
|
62
|
+
|
|
63
|
+
self.pre_pro = PreProcess(max_dims=self.max_dims, min_dims=self.min_dims)
|
|
64
|
+
|
|
65
|
+
self.image_resizer = OrtInferSession(self.image_resizer_path)
|
|
66
|
+
|
|
67
|
+
self.encoder_decoder = EncoderDecoder(
|
|
68
|
+
encoder_path=self.encoder_path,
|
|
69
|
+
decoder_path=self.decoder_path,
|
|
70
|
+
bos_token=input_params.bos_token,
|
|
71
|
+
eos_token=input_params.eos_token,
|
|
72
|
+
max_seq_len=input_params.max_seq_len,
|
|
73
|
+
)
|
|
74
|
+
self.tokenizer = TokenizerCls(self.tokenizer_json)
|
|
75
|
+
|
|
76
|
+
def get_model_path(
|
|
77
|
+
self,
|
|
78
|
+
) -> Tuple[str]:
|
|
79
|
+
def try_download(file_name):
|
|
80
|
+
save_path = default_model_dir / file_name
|
|
81
|
+
if save_path.exists() or downloader(file_name):
|
|
82
|
+
return save_path
|
|
83
|
+
raise FileNotFoundError(f"{file_name} must not be None.")
|
|
84
|
+
|
|
85
|
+
downloader = DownloadModel()
|
|
86
|
+
decoder_name = "decoder.onnx"
|
|
87
|
+
encoder_name = "encoder.onnx"
|
|
88
|
+
resizer_name = "image_resizer.onnx"
|
|
89
|
+
tokenizer_name = "tokenizer.json"
|
|
90
|
+
|
|
91
|
+
default_model_dir = cur_dir / "models"
|
|
92
|
+
|
|
93
|
+
if self.image_resizer_path is None:
|
|
94
|
+
self.image_resizer_path = try_download(resizer_name)
|
|
95
|
+
|
|
96
|
+
if self.encoder_path is None:
|
|
97
|
+
self.encoder_path = try_download(encoder_name)
|
|
98
|
+
|
|
99
|
+
if self.decoder_path is None:
|
|
100
|
+
self.decoder_path = try_download(decoder_name)
|
|
101
|
+
|
|
102
|
+
if self.tokenizer_json is None:
|
|
103
|
+
self.tokenizer_json = try_download(tokenizer_name)
|
|
104
|
+
|
|
105
|
+
def __call__(self, img: InputType) -> Tuple[str, float]:
|
|
106
|
+
s = time.perf_counter()
|
|
107
|
+
|
|
108
|
+
try:
|
|
109
|
+
img = self.load_img(img)
|
|
110
|
+
except LoadImageError as exc:
|
|
111
|
+
error_info = traceback.format_exc()
|
|
112
|
+
raise LoadImageError(f"Load the img meets error. Error info is {error_info}") from exc
|
|
113
|
+
|
|
114
|
+
try:
|
|
115
|
+
resizered_img = self.loop_image_resizer(img)
|
|
116
|
+
except Exception as e:
|
|
117
|
+
error_info = traceback.format_exc()
|
|
118
|
+
raise ValueError(f"image resizer meets error. Error info is {error_info}") from e
|
|
119
|
+
|
|
120
|
+
try:
|
|
121
|
+
dec = self.encoder_decoder(resizered_img, temperature=self.temperature)
|
|
122
|
+
except Exception as e:
|
|
123
|
+
error_info = traceback.format_exc()
|
|
124
|
+
raise ValueError(f"EncoderDecoder meets error. Error info is {error_info}") from e
|
|
125
|
+
|
|
126
|
+
decode = self.tokenizer.token2str(dec)
|
|
127
|
+
pred = self.post_process(decode[0])
|
|
128
|
+
|
|
129
|
+
elapse = time.perf_counter() - s
|
|
130
|
+
return pred, elapse
|
|
131
|
+
|
|
132
|
+
def loop_image_resizer(self, img: np.ndarray) -> np.ndarray:
|
|
133
|
+
pillow_img = Image.fromarray(img)
|
|
134
|
+
pad_img = self.pre_pro.pad(pillow_img)
|
|
135
|
+
input_image = self.pre_pro.minmax_size(pad_img).convert("RGB")
|
|
136
|
+
r, w, h = 1, input_image.size[0], input_image.size[1]
|
|
137
|
+
for _ in range(10):
|
|
138
|
+
h = int(h * r)
|
|
139
|
+
final_img, pad_img = self.pre_process(input_image, r, w, h)
|
|
140
|
+
|
|
141
|
+
resizer_res = self.image_resizer([final_img.astype(np.float32)])[0]
|
|
142
|
+
|
|
143
|
+
argmax_idx = int(np.argmax(resizer_res, axis=-1).item())
|
|
144
|
+
w = (argmax_idx + 1) * 32
|
|
145
|
+
if w == pad_img.size[0]:
|
|
146
|
+
break
|
|
147
|
+
|
|
148
|
+
r = w / pad_img.size[0]
|
|
149
|
+
return final_img
|
|
150
|
+
|
|
151
|
+
def pre_process(self, input_image: Image.Image, r, w, h) -> Tuple[np.ndarray, Image.Image]:
|
|
152
|
+
if r > 1:
|
|
153
|
+
resize_func = Image.Resampling.BILINEAR
|
|
154
|
+
else:
|
|
155
|
+
resize_func = Image.Resampling.LANCZOS
|
|
156
|
+
|
|
157
|
+
resize_img = input_image.resize((w, h), resize_func)
|
|
158
|
+
pad_img = self.pre_pro.pad(self.pre_pro.minmax_size(resize_img))
|
|
159
|
+
cvt_img = np.array(pad_img.convert("RGB"))
|
|
160
|
+
|
|
161
|
+
gray_img = self.pre_pro.to_gray(cvt_img)
|
|
162
|
+
normal_img = self.pre_pro.normalize(gray_img)
|
|
163
|
+
final_img = self.pre_pro.transpose_and_four_dim(normal_img)
|
|
164
|
+
return final_img, pad_img
|
|
165
|
+
|
|
166
|
+
@staticmethod
|
|
167
|
+
def post_process(s: str) -> str:
|
|
168
|
+
"""Remove unnecessary whitespace from LaTeX code.
|
|
169
|
+
|
|
170
|
+
Args:
|
|
171
|
+
s (str): Input string
|
|
172
|
+
|
|
173
|
+
Returns:
|
|
174
|
+
str: Processed image
|
|
175
|
+
"""
|
|
176
|
+
text_reg = r"(\\(operatorname|mathrm|text|mathbf)\s?\*? {.*?})"
|
|
177
|
+
letter = "[a-zA-Z]"
|
|
178
|
+
noletter = r"[\W_^\d]"
|
|
179
|
+
names = [x[0].replace(" ", "") for x in re.findall(text_reg, s)]
|
|
180
|
+
s = re.sub(text_reg, lambda match: str(names.pop(0)), s)
|
|
181
|
+
news = s
|
|
182
|
+
while True:
|
|
183
|
+
s = news
|
|
184
|
+
news = re.sub(r"(?!\\ )(%s)\s+?(%s)" % (noletter, noletter), r"\1\2", s)
|
|
185
|
+
news = re.sub(r"(?!\\ )(%s)\s+?(%s)" % (noletter, letter), r"\1\2", news)
|
|
186
|
+
news = re.sub(r"(%s)\s+?(%s)" % (letter, noletter), r"\1\2", news)
|
|
187
|
+
if news == s:
|
|
188
|
+
break
|
|
189
|
+
return s
|
|
190
|
+
|
|
191
|
+
|
|
192
|
+
def main():
|
|
193
|
+
parser = argparse.ArgumentParser()
|
|
194
|
+
parser.add_argument("-img_resizer", "--image_resizer_path", type=str, default=None)
|
|
195
|
+
parser.add_argument("-encdoer", "--encoder_path", type=str, default=None)
|
|
196
|
+
parser.add_argument("-decoder", "--decoder_path", type=str, default=None)
|
|
197
|
+
parser.add_argument("-tokenizer", "--tokenizer_json", type=str, default=None)
|
|
198
|
+
parser.add_argument("img_path", type=str, help="Only img path of the formula.")
|
|
199
|
+
args = parser.parse_args()
|
|
200
|
+
|
|
201
|
+
engine = LaTeXOCR(
|
|
202
|
+
image_resizer_path=args.image_resizer_path,
|
|
203
|
+
encoder_path=args.encoder_path,
|
|
204
|
+
decoder_path=args.decoder_path,
|
|
205
|
+
tokenizer_json=args.tokenizer_json,
|
|
206
|
+
)
|
|
207
|
+
|
|
208
|
+
result, elapse = engine(args.img_path)
|
|
209
|
+
print(result)
|
|
210
|
+
print(f"cost: {elapse:.5f}")
|
|
211
|
+
|
|
212
|
+
|
|
213
|
+
if __name__ == "__main__":
|
|
214
|
+
main()
|
|
@@ -0,0 +1,154 @@
|
|
|
1
|
+
# -*- encoding: utf-8 -*-
|
|
2
|
+
# @Author: SWHL
|
|
3
|
+
# @Contact: liekkaskono@163.com
|
|
4
|
+
from pathlib import Path
|
|
5
|
+
from typing import Tuple, Union
|
|
6
|
+
|
|
7
|
+
import numpy as np
|
|
8
|
+
|
|
9
|
+
from .utils_load import OrtInferSession
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class EncoderDecoder:
|
|
13
|
+
def __init__(
|
|
14
|
+
self,
|
|
15
|
+
encoder_path: Union[Path, str],
|
|
16
|
+
decoder_path: Union[Path, str],
|
|
17
|
+
bos_token: int,
|
|
18
|
+
eos_token: int,
|
|
19
|
+
max_seq_len: int,
|
|
20
|
+
):
|
|
21
|
+
self.bos_token = bos_token
|
|
22
|
+
self.eos_token = eos_token
|
|
23
|
+
self.max_seq_len = max_seq_len
|
|
24
|
+
|
|
25
|
+
self.encoder = OrtInferSession(encoder_path)
|
|
26
|
+
self.decoder = Decoder(decoder_path)
|
|
27
|
+
|
|
28
|
+
def __call__(self, x: np.ndarray, temperature: float = 0.25):
|
|
29
|
+
ort_input_data = np.array([self.bos_token] * len(x))[:, None]
|
|
30
|
+
context = self.encoder([x])[0]
|
|
31
|
+
output = self.decoder(
|
|
32
|
+
ort_input_data,
|
|
33
|
+
self.max_seq_len,
|
|
34
|
+
eos_token=self.eos_token,
|
|
35
|
+
context=context,
|
|
36
|
+
temperature=temperature,
|
|
37
|
+
)
|
|
38
|
+
return output
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
class Decoder:
|
|
42
|
+
def __init__(self, decoder_path: Union[Path, str]):
|
|
43
|
+
self.max_seq_len = 512
|
|
44
|
+
self.session = OrtInferSession(decoder_path)
|
|
45
|
+
|
|
46
|
+
def __call__(
|
|
47
|
+
self,
|
|
48
|
+
start_tokens,
|
|
49
|
+
seq_len=256,
|
|
50
|
+
eos_token=None,
|
|
51
|
+
temperature=1.0,
|
|
52
|
+
filter_thres=0.9,
|
|
53
|
+
context=None,
|
|
54
|
+
):
|
|
55
|
+
num_dims = len(start_tokens.shape)
|
|
56
|
+
|
|
57
|
+
b, t = start_tokens.shape
|
|
58
|
+
|
|
59
|
+
out = start_tokens
|
|
60
|
+
mask = np.full_like(start_tokens, True, dtype=bool)
|
|
61
|
+
|
|
62
|
+
for _ in range(seq_len):
|
|
63
|
+
x = out[:, -self.max_seq_len :]
|
|
64
|
+
mask = mask[:, -self.max_seq_len :]
|
|
65
|
+
|
|
66
|
+
ort_outs = self.session([x.astype(np.int64), mask, context])[0]
|
|
67
|
+
np_preds = ort_outs
|
|
68
|
+
np_logits = np_preds[:, -1, :]
|
|
69
|
+
|
|
70
|
+
np_filtered_logits = self.npp_top_k(np_logits, thres=filter_thres)
|
|
71
|
+
np_probs = self.softmax(np_filtered_logits / temperature, axis=-1)
|
|
72
|
+
|
|
73
|
+
sample = self.multinomial(np_probs.squeeze(), 1)[None, ...]
|
|
74
|
+
|
|
75
|
+
out = np.concatenate([out, sample], axis=-1)
|
|
76
|
+
mask = np.pad(mask, [(0, 0), (0, 1)], "constant", constant_values=True)
|
|
77
|
+
|
|
78
|
+
if eos_token is not None and (np.cumsum(out == eos_token, axis=1)[:, -1] >= 1).all():
|
|
79
|
+
break
|
|
80
|
+
|
|
81
|
+
out = out[:, t:]
|
|
82
|
+
if num_dims == 1:
|
|
83
|
+
out = out.squeeze(0)
|
|
84
|
+
return out
|
|
85
|
+
|
|
86
|
+
@staticmethod
|
|
87
|
+
def softmax(x, axis=None) -> float:
|
|
88
|
+
def logsumexp(a, axis=None, b=None, keepdims=False):
|
|
89
|
+
a_max = np.amax(a, axis=axis, keepdims=True)
|
|
90
|
+
|
|
91
|
+
if a_max.ndim > 0:
|
|
92
|
+
a_max[~np.isfinite(a_max)] = 0
|
|
93
|
+
elif not np.isfinite(a_max):
|
|
94
|
+
a_max = 0
|
|
95
|
+
|
|
96
|
+
tmp = np.exp(a - a_max)
|
|
97
|
+
|
|
98
|
+
# suppress warnings about log of zero
|
|
99
|
+
with np.errstate(divide="ignore"):
|
|
100
|
+
s = np.sum(tmp, axis=axis, keepdims=keepdims)
|
|
101
|
+
out = np.log(s)
|
|
102
|
+
|
|
103
|
+
if not keepdims:
|
|
104
|
+
a_max = np.squeeze(a_max, axis=axis)
|
|
105
|
+
out += a_max
|
|
106
|
+
return out
|
|
107
|
+
|
|
108
|
+
return np.exp(x - logsumexp(x, axis=axis, keepdims=True))
|
|
109
|
+
|
|
110
|
+
def npp_top_k(self, logits, thres=0.9):
|
|
111
|
+
k = int((1 - thres) * logits.shape[-1])
|
|
112
|
+
val, ind = self.np_top_k(logits, k)
|
|
113
|
+
probs = np.full_like(logits, float("-inf"))
|
|
114
|
+
np.put_along_axis(probs, ind, val, axis=1)
|
|
115
|
+
return probs
|
|
116
|
+
|
|
117
|
+
@staticmethod
|
|
118
|
+
def np_top_k(
|
|
119
|
+
a: np.ndarray, k: int, axis=-1, largest=True, sorted=True
|
|
120
|
+
) -> Tuple[np.ndarray, np.ndarray]:
|
|
121
|
+
if axis is None:
|
|
122
|
+
axis_size = a.size
|
|
123
|
+
else:
|
|
124
|
+
axis_size = a.shape[axis]
|
|
125
|
+
|
|
126
|
+
assert 1 <= k <= axis_size
|
|
127
|
+
|
|
128
|
+
a = np.asanyarray(a)
|
|
129
|
+
if largest:
|
|
130
|
+
index_array = np.argpartition(a, axis_size - k, axis=axis)
|
|
131
|
+
topk_indices = np.take(index_array, -np.arange(k) - 1, axis=axis)
|
|
132
|
+
else:
|
|
133
|
+
index_array = np.argpartition(a, k - 1, axis=axis)
|
|
134
|
+
topk_indices = np.take(index_array, np.arange(k), axis=axis)
|
|
135
|
+
|
|
136
|
+
topk_values = np.take_along_axis(a, topk_indices, axis=axis)
|
|
137
|
+
if sorted:
|
|
138
|
+
sorted_indices_in_topk = np.argsort(topk_values, axis=axis)
|
|
139
|
+
if largest:
|
|
140
|
+
sorted_indices_in_topk = np.flip(sorted_indices_in_topk, axis=axis)
|
|
141
|
+
sorted_topk_values = np.take_along_axis(topk_values, sorted_indices_in_topk, axis=axis)
|
|
142
|
+
sorted_topk_indices = np.take_along_axis(
|
|
143
|
+
topk_indices, sorted_indices_in_topk, axis=axis
|
|
144
|
+
)
|
|
145
|
+
return sorted_topk_values, sorted_topk_indices
|
|
146
|
+
return topk_values, topk_indices
|
|
147
|
+
|
|
148
|
+
@staticmethod
|
|
149
|
+
def multinomial(weights, num_samples, replacement=True):
|
|
150
|
+
weights = np.asarray(weights)
|
|
151
|
+
weights /= np.sum(weights) # 确保权重之和为1
|
|
152
|
+
indices = np.arange(len(weights))
|
|
153
|
+
samples = np.random.choice(indices, size=num_samples, replace=replacement, p=weights)
|
|
154
|
+
return samples
|
|
@@ -0,0 +1,185 @@
|
|
|
1
|
+
# -*- encoding: utf-8 -*-
|
|
2
|
+
# @Author: SWHL
|
|
3
|
+
# @Contact: liekkaskono@163.com
|
|
4
|
+
import io
|
|
5
|
+
from pathlib import Path
|
|
6
|
+
from typing import List, Optional, Union
|
|
7
|
+
|
|
8
|
+
import chardet
|
|
9
|
+
import cv2
|
|
10
|
+
import numpy as np
|
|
11
|
+
import requests
|
|
12
|
+
import tqdm
|
|
13
|
+
from PIL import Image
|
|
14
|
+
from tokenizers import Tokenizer
|
|
15
|
+
from tokenizers.models import BPE
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class PreProcess:
|
|
19
|
+
def __init__(self, max_dims: List[int], min_dims: List[int]):
|
|
20
|
+
self.max_dims, self.min_dims = max_dims, min_dims
|
|
21
|
+
self.mean = np.array([0.7931, 0.7931, 0.7931]).astype(np.float32)
|
|
22
|
+
self.std = np.array([0.1738, 0.1738, 0.1738]).astype(np.float32)
|
|
23
|
+
|
|
24
|
+
@staticmethod
|
|
25
|
+
def pad(img: Image.Image, divable: int = 32) -> Image.Image:
|
|
26
|
+
"""Pad an Image to the next full divisible value of `divable`. Also normalizes the image and invert if needed.
|
|
27
|
+
|
|
28
|
+
Args:
|
|
29
|
+
img (PIL.Image): input image
|
|
30
|
+
divable (int, optional): . Defaults to 32.
|
|
31
|
+
|
|
32
|
+
Returns:
|
|
33
|
+
PIL.Image
|
|
34
|
+
"""
|
|
35
|
+
threshold = 128
|
|
36
|
+
data = np.array(img.convert("LA"))
|
|
37
|
+
if data[..., -1].var() == 0:
|
|
38
|
+
data = (data[..., 0]).astype(np.uint8)
|
|
39
|
+
else:
|
|
40
|
+
data = (255 - data[..., -1]).astype(np.uint8)
|
|
41
|
+
|
|
42
|
+
data = (data - data.min()) / (data.max() - data.min()) * 255
|
|
43
|
+
if data.mean() > threshold:
|
|
44
|
+
# To invert the text to white
|
|
45
|
+
gray = 255 * (data < threshold).astype(np.uint8)
|
|
46
|
+
else:
|
|
47
|
+
gray = 255 * (data > threshold).astype(np.uint8)
|
|
48
|
+
data = 255 - data
|
|
49
|
+
|
|
50
|
+
coords = cv2.findNonZero(gray) # Find all non-zero points (text)
|
|
51
|
+
a, b, w, h = cv2.boundingRect(coords) # Find minimum spanning bounding box
|
|
52
|
+
rect = data[b : b + h, a : a + w]
|
|
53
|
+
im = Image.fromarray(rect).convert("L")
|
|
54
|
+
dims: List[Union[int, int]] = []
|
|
55
|
+
for x in [w, h]:
|
|
56
|
+
div, mod = divmod(x, divable)
|
|
57
|
+
dims.append(divable * (div + (1 if mod > 0 else 0)))
|
|
58
|
+
|
|
59
|
+
padded = Image.new("L", tuple(dims), 255)
|
|
60
|
+
padded.paste(im, (0, 0, im.size[0], im.size[1]))
|
|
61
|
+
return padded
|
|
62
|
+
|
|
63
|
+
def minmax_size(
|
|
64
|
+
self,
|
|
65
|
+
img: Image.Image,
|
|
66
|
+
) -> Image.Image:
|
|
67
|
+
"""Resize or pad an image to fit into given dimensions
|
|
68
|
+
|
|
69
|
+
Args:
|
|
70
|
+
img (Image): Image to scale up/down.
|
|
71
|
+
|
|
72
|
+
Returns:
|
|
73
|
+
Image: Image with correct dimensionality
|
|
74
|
+
"""
|
|
75
|
+
if self.max_dims is not None:
|
|
76
|
+
ratios = [a / b for a, b in zip(img.size, self.max_dims)]
|
|
77
|
+
if any([r > 1 for r in ratios]):
|
|
78
|
+
size = np.array(img.size) // max(ratios)
|
|
79
|
+
size = np.maximum(size, 1)
|
|
80
|
+
img = img.resize(tuple(size.astype(int)), Image.BILINEAR)
|
|
81
|
+
|
|
82
|
+
if self.min_dims is not None:
|
|
83
|
+
padded_size: List[Union[int, int]] = [
|
|
84
|
+
max(img_dim, min_dim) for img_dim, min_dim in zip(img.size, self.min_dims)
|
|
85
|
+
]
|
|
86
|
+
|
|
87
|
+
new_pad_size = tuple(padded_size)
|
|
88
|
+
if new_pad_size != img.size: # assert hypothesis
|
|
89
|
+
padded_im = Image.new("L", new_pad_size, 255)
|
|
90
|
+
padded_im.paste(img, img.getbbox())
|
|
91
|
+
img = padded_im
|
|
92
|
+
return img
|
|
93
|
+
|
|
94
|
+
def normalize(self, img: np.ndarray, max_pixel_value=255.0) -> np.ndarray:
|
|
95
|
+
mean = self.mean * max_pixel_value
|
|
96
|
+
std = self.std * max_pixel_value
|
|
97
|
+
denominator = np.reciprocal(std, dtype=np.float32)
|
|
98
|
+
img = img.astype(np.float32)
|
|
99
|
+
img -= mean
|
|
100
|
+
img *= denominator
|
|
101
|
+
return img
|
|
102
|
+
|
|
103
|
+
@staticmethod
|
|
104
|
+
def to_gray(img) -> np.ndarray:
|
|
105
|
+
gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)
|
|
106
|
+
return cv2.cvtColor(gray, cv2.COLOR_GRAY2RGB)
|
|
107
|
+
|
|
108
|
+
@staticmethod
|
|
109
|
+
def transpose_and_four_dim(img: np.ndarray) -> np.ndarray:
|
|
110
|
+
return img.transpose(2, 0, 1)[:1][None, ...]
|
|
111
|
+
|
|
112
|
+
|
|
113
|
+
class TokenizerCls:
|
|
114
|
+
def __init__(self, json_file: Union[Path, str]):
|
|
115
|
+
self.tokenizer = Tokenizer(BPE()).from_file(str(json_file))
|
|
116
|
+
|
|
117
|
+
def token2str(self, tokens) -> List[str]:
|
|
118
|
+
if len(tokens.shape) == 1:
|
|
119
|
+
tokens = tokens[None, :]
|
|
120
|
+
|
|
121
|
+
dec = [self.tokenizer.decode(tok.tolist()) for tok in tokens]
|
|
122
|
+
return [
|
|
123
|
+
"".join(detok.split(" "))
|
|
124
|
+
.replace("Ġ", " ")
|
|
125
|
+
.replace("[EOS]", "")
|
|
126
|
+
.replace("[BOS]", "")
|
|
127
|
+
.replace("[PAD]", "")
|
|
128
|
+
.strip()
|
|
129
|
+
for detok in dec
|
|
130
|
+
]
|
|
131
|
+
|
|
132
|
+
|
|
133
|
+
class DownloadModel:
|
|
134
|
+
"""Modified from https://github.com/lukas-blecher/LaTeX-OCR/blob/1781514fb8c92ea9f94057295fdae0e683f4648e/pix2tex/model/checkpoints/get_latest_checkpoint.py"""
|
|
135
|
+
|
|
136
|
+
def __init__(self) -> None:
|
|
137
|
+
self.url = "https://github.com/RapidAI/RapidLaTeXOCR/releases/download/v0.0.0"
|
|
138
|
+
self.cur_dir = Path(__file__).resolve().parent
|
|
139
|
+
|
|
140
|
+
def __call__(self, file_name: str) -> bool:
|
|
141
|
+
save_dir = self.cur_dir / "models"
|
|
142
|
+
save_dir.mkdir(parents=True, exist_ok=True)
|
|
143
|
+
|
|
144
|
+
full_path = f"{self.url}/{file_name}"
|
|
145
|
+
print(f"Download {full_path} to {self.cur_dir}/models")
|
|
146
|
+
|
|
147
|
+
try:
|
|
148
|
+
file = self.download_as_bytes_with_progress(full_path, file_name)
|
|
149
|
+
save_file_path = save_dir / file_name
|
|
150
|
+
self.save_file(save_file_path, file)
|
|
151
|
+
except Exception:
|
|
152
|
+
return False
|
|
153
|
+
return True
|
|
154
|
+
|
|
155
|
+
@staticmethod
|
|
156
|
+
def download_as_bytes_with_progress(url: str, name: Optional[str] = None) -> bytes:
|
|
157
|
+
resp = requests.get(url, stream=True, allow_redirects=True)
|
|
158
|
+
total = int(resp.headers.get("content-length", 0))
|
|
159
|
+
bio = io.BytesIO()
|
|
160
|
+
with tqdm.tqdm(desc=name, total=total, unit="b", unit_scale=True, unit_divisor=1024) as bar:
|
|
161
|
+
for chunk in resp.iter_content(chunk_size=65536):
|
|
162
|
+
bar.update(len(chunk))
|
|
163
|
+
bio.write(chunk)
|
|
164
|
+
return bio.getvalue()
|
|
165
|
+
|
|
166
|
+
@staticmethod
|
|
167
|
+
def save_file(save_path: Union[str, Path], file: bytes):
|
|
168
|
+
with open(save_path, "wb") as f:
|
|
169
|
+
f.write(file)
|
|
170
|
+
|
|
171
|
+
|
|
172
|
+
def get_file_encode(file_path: Union[str, Path]) -> str:
|
|
173
|
+
try:
|
|
174
|
+
with open(file_path, "rb") as f:
|
|
175
|
+
raw_data = f.read(100)
|
|
176
|
+
result = chardet.detect(raw_data)
|
|
177
|
+
encoding = result["encoding"]
|
|
178
|
+
return encoding
|
|
179
|
+
except Exception:
|
|
180
|
+
return "utf-8"
|
|
181
|
+
|
|
182
|
+
|
|
183
|
+
if __name__ == "__main__":
|
|
184
|
+
downloader = DownloadModel()
|
|
185
|
+
downloader("decoder.onnx")
|